mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-11 11:53:01 +08:00
Optimize DQN
This commit is contained in:
+19
-6
@@ -80,11 +80,19 @@ class DQNAgent:
|
||||
states, actions, rewards, next_states, terminals = experiences
|
||||
states = self.task.normalize_state(states)
|
||||
next_states = self.task.normalize_state(next_states)
|
||||
q_next = self.target_network.predict(next_states)
|
||||
q_next = np.max(q_next, axis=1)
|
||||
q_next = np.where(terminals, 0, q_next)
|
||||
q_next = rewards + self.discount * q_next
|
||||
self.learning_network.learn(states, actions, q_next)
|
||||
q_next = self.target_network.predict(next_states, False).detach()
|
||||
q_next, _ = q_next.max(1)
|
||||
terminals = self.learning_network.to_torch_variable(terminals).unsqueeze(1)
|
||||
rewards = self.learning_network.to_torch_variable(rewards).unsqueeze(1)
|
||||
q_next = q_next * (1 - terminals)
|
||||
q_next.add_(rewards)
|
||||
actions = self.learning_network.to_torch_variable(actions, 'int64').unsqueeze(1)
|
||||
q = self.learning_network.predict(states, False)
|
||||
q = q.gather(1, actions)
|
||||
loss = self.learning_network.criterion(q, q_next)
|
||||
self.learning_network.zero_grad()
|
||||
loss.backward()
|
||||
self.learning_network.optimizer.step()
|
||||
if not deterministic and self.total_steps % self.target_network_update_freq == 0:
|
||||
self.target_network.load_state_dict(self.learning_network.state_dict())
|
||||
if not deterministic and self.total_steps > self.explore_steps:
|
||||
@@ -103,6 +111,7 @@ class DQNAgent:
|
||||
window_size = 100
|
||||
ep = 0
|
||||
rewards = []
|
||||
avg_test_rewards = []
|
||||
while True:
|
||||
ep += 1
|
||||
reward = self.episode()
|
||||
@@ -113,12 +122,16 @@ class DQNAgent:
|
||||
|
||||
if ep % self.test_interval == 0:
|
||||
self.logger.info('Testing...')
|
||||
self.save('data/dqn-model.bin')
|
||||
self.save('data/dqn-model-%s.bin' % (self.task.name))
|
||||
test_rewards = []
|
||||
for _ in range(self.test_repetitions):
|
||||
test_rewards.append(self.episode(True))
|
||||
avg_reward = np.mean(test_rewards)
|
||||
avg_test_rewards.append(avg_reward)
|
||||
self.logger.info('Avg reward %f(%f)' % (
|
||||
avg_reward, np.std(test_rewards) / np.sqrt(self.test_repetitions)))
|
||||
with open('data/dqn-statistics-%s.bin' % (self.task.name), 'wb') as f:
|
||||
pickle.dump({'rewards': rewards,
|
||||
'test_rewards': avg_test_rewards}, f)
|
||||
if avg_reward > self.task.success_threshold:
|
||||
break
|
||||
@@ -43,7 +43,7 @@ def async_lunar_lander():
|
||||
def dqn_cart_pole():
|
||||
config = dict()
|
||||
config['task_fn'] = lambda: CartPole()
|
||||
config['optimizer_fn'] = lambda params: torch.optim.SGD(params, 0.001)
|
||||
config['optimizer_fn'] = lambda params: torch.optim.RMSprop(params, 0.001)
|
||||
config['network_fn'] = lambda optimizer_fn: FullyConnectedNet([8, 50, 200, 2], optimizer_fn)
|
||||
config['policy_fn'] = lambda: GreedyPolicy(epsilon=1.0, final_step=10000, min_epsilon=0.1)
|
||||
config['replay_fn'] = lambda: Replay(memory_size=10000, batch_size=10)
|
||||
|
||||
+25
-70
@@ -10,14 +10,8 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
|
||||
class FullyConnectedNet(nn.Module):
|
||||
def __init__(self, dims, optimizer_fn=None, gpu=True):
|
||||
super(FullyConnectedNet, self).__init__()
|
||||
self.fc1 = nn.Linear(dims[0], dims[1])
|
||||
self.fc2 = nn.Linear(dims[1], dims[2])
|
||||
self.fc3 = nn.Linear(dims[2], dims[3])
|
||||
self.criterion = nn.MSELoss()
|
||||
class BasicNet:
|
||||
def __init__(self, optimizer_fn, gpu):
|
||||
if optimizer_fn is not None:
|
||||
self.optimizer = optimizer_fn(self.parameters())
|
||||
self.gpu = gpu and torch.cuda.is_available()
|
||||
@@ -26,51 +20,40 @@ class FullyConnectedNet(nn.Module):
|
||||
self.cuda()
|
||||
print 'Network transferred.'
|
||||
|
||||
def forward(self, x):
|
||||
x = x.reshape((x.shape[0], -1))
|
||||
x = self.to_torch_variable(x)
|
||||
|
||||
y = F.relu(self.fc1(x))
|
||||
y = F.relu(self.fc2(y))
|
||||
y = self.fc3(y)
|
||||
return y
|
||||
|
||||
def predict(self, x):
|
||||
return self.forward(x).cpu().data.numpy()
|
||||
|
||||
def to_torch_variable(self, x, dtype='float32'):
|
||||
x = torch.from_numpy(np.asarray(x, dtype=dtype))
|
||||
if self.gpu:
|
||||
x = x.cuda()
|
||||
return Variable(x)
|
||||
|
||||
def learn(self, x, actions, targets):
|
||||
self.zero_grad()
|
||||
self.gradient(x, actions, targets)
|
||||
self.optimizer.step()
|
||||
|
||||
# def clippedLearn(self, x, actions, targets):
|
||||
# y = self.forward(x)
|
||||
# actions = self.to_torch_variable(actions, 'int64').unsqueeze(1)
|
||||
# targets = self.to_torch_variable(targets).unsqueeze(1)
|
||||
# y = y.gather(1, actions)
|
||||
# bellman_error = targets - y
|
||||
# bellman_error = bellman_error.clamp(-1, 1) * -1
|
||||
# self.zero_grad()
|
||||
# y.backward(bellman_error.data)
|
||||
# self.optimizer.step()
|
||||
def predict(self, x, to_numpy=True):
|
||||
y = self.forward(x)
|
||||
if to_numpy:
|
||||
y = y.cpu().data.numpy()
|
||||
return y
|
||||
|
||||
def gradient(self, x, actions, targets):
|
||||
y = self.forward(x)
|
||||
actions = self.to_torch_variable(actions, 'int64').unsqueeze(1)
|
||||
targets = self.to_torch_variable(targets).unsqueeze(1)
|
||||
y = y.gather(1, actions)
|
||||
loss = self.criterion(y, targets)
|
||||
loss.backward()
|
||||
|
||||
def output_transfer(self, y):
|
||||
return y
|
||||
class FullyConnectedNet(nn.Module, BasicNet):
|
||||
def __init__(self, dims, optimizer_fn=None, gpu=True):
|
||||
super(FullyConnectedNet, self).__init__()
|
||||
self.fc1 = nn.Linear(dims[0], dims[1])
|
||||
self.fc2 = nn.Linear(dims[1], dims[2])
|
||||
self.fc3 = nn.Linear(dims[2], dims[3])
|
||||
self.criterion = nn.MSELoss()
|
||||
BasicNet.__init__(self, optimizer_fn, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
x = x.view(x.size(0), -1)
|
||||
y = F.relu(self.fc1(x))
|
||||
y = F.relu(self.fc2(y))
|
||||
y = self.fc3(y)
|
||||
return y
|
||||
|
||||
class ActorCriticNet(nn.Module):
|
||||
def __init__(self, dims, gpu=True):
|
||||
@@ -117,8 +100,7 @@ class ActorCriticNet(nn.Module):
|
||||
phi = self.forward(x)
|
||||
return self.fc_critic(phi).cpu().data.numpy()
|
||||
|
||||
|
||||
class ConvNet(nn.Module):
|
||||
class ConvNet(nn.Module, BasicNet):
|
||||
def __init__(self, in_channels, n_actions, optimizer_fn=None, gpu=True):
|
||||
super(ConvNet, self).__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels, 32, kernel_size=8, stride=4)
|
||||
@@ -126,24 +108,11 @@ class ConvNet(nn.Module):
|
||||
self.conv3 = nn.Conv2d(64, 64, kernel_size=3, stride=1)
|
||||
self.fc4 = nn.Linear(7 * 7 * 64, 512)
|
||||
self.fc5 = nn.Linear(512, n_actions)
|
||||
|
||||
self.criterion = nn.MSELoss()
|
||||
if optimizer_fn is not None:
|
||||
self.optimizer = optimizer_fn(self.parameters())
|
||||
|
||||
self.gpu = gpu and torch.cuda.is_available()
|
||||
if self.gpu:
|
||||
print 'Transferring network to GPU...'
|
||||
self.cuda()
|
||||
print 'Network transferred.'
|
||||
|
||||
def to_torch_variable(self, x, dtype='float32'):
|
||||
x = torch.from_numpy(np.asarray(x, dtype=dtype))
|
||||
if self.gpu:
|
||||
x = x.cuda()
|
||||
return Variable(x)
|
||||
BasicNet.__init__(self, optimizer_fn, gpu)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.to_torch_variable(x)
|
||||
y = F.relu(self.conv1(x))
|
||||
y = F.relu(self.conv2(y))
|
||||
y = F.relu(self.conv3(y))
|
||||
@@ -151,18 +120,4 @@ class ConvNet(nn.Module):
|
||||
y = F.relu(self.fc4(y))
|
||||
return self.fc5(y)
|
||||
|
||||
def predict(self, x):
|
||||
return self.forward(self.to_torch_variable(x)).cpu().data.numpy()
|
||||
|
||||
def learn(self, x, actions, targets):
|
||||
self.zero_grad()
|
||||
self.gradient(x, actions, targets)
|
||||
self.optimizer.step()
|
||||
|
||||
def gradient(self, x, actions, targets):
|
||||
y = self.forward(self.to_torch_variable(x))
|
||||
actions = self.to_torch_variable(actions, 'int64').unsqueeze(1)
|
||||
targets = self.to_torch_variable(targets).unsqueeze(1)
|
||||
y = y.gather(1, actions)
|
||||
loss = self.criterion(y, targets)
|
||||
loss.backward()
|
||||
|
||||
Reference in New Issue
Block a user