mirror of
https://github.com/wassname/DeepRL.git
synced 2026-10-10 11:50:27 +08:00
Specify a gpu for a network
This commit is contained in:
1 parent
fef04b6bf6
commit
e91ec3be45
12 files changed
+93
-103
No files matched your search
+2
-2
@@ -78,8 +78,8 @@ class DeterministicPolicyGradient:
|
||||
experiences = self.replay.sample()
|
||||
states, actions, rewards, next_states, terminals = experiences
|
||||
q_next = target_critic.predict(next_states, target_actor.predict(next_states))
|
||||
terminals = critic.to_torch_variable(terminals).unsqueeze(1)
|
||||
rewards = critic.to_torch_variable(rewards).unsqueeze(1)
|
||||
terminals = critic.variable(terminals).unsqueeze(1)
|
||||
rewards = critic.variable(rewards).unsqueeze(1)
|
||||
q_next = config.discount * q_next * (1 - terminals)
|
||||
q_next.add_(rewards)
|
||||
q_next = q_next.detach()
|
||||
|
||||
+5
-5
@@ -92,10 +92,10 @@ class ProximalPolicyOptimization:
|
||||
R = critic_net.predict(np.stack([state])).data
|
||||
|
||||
|
||||
values.append(actor_net.to_torch_variable(R))
|
||||
A = actor_net.to_torch_variable(torch.zeros((1, 1)))
|
||||
values.append(actor_net.variable(R))
|
||||
A = actor_net.variable(torch.zeros((1, 1)))
|
||||
for i in reversed(range(len(rewards))):
|
||||
R = actor_net.to_torch_variable([[rewards[i]]])
|
||||
R = actor_net.variable([[rewards[i]]])
|
||||
ret = R + self.config.discount * values[i + 1]
|
||||
A = ret - values[i] + self.config.discount * self.config.gae_tau * A
|
||||
advantages.append(A.detach())
|
||||
@@ -123,8 +123,8 @@ class ProximalPolicyOptimization:
|
||||
self.worker_network.load_state_dict(self.shared_network.state_dict())
|
||||
|
||||
states, actions, returns, advantages = replay.sample()
|
||||
states = actor_net.to_torch_variable(np.stack(states))
|
||||
actions = actor_net.to_torch_variable(np.stack(actions))
|
||||
states = actor_net.variable(np.stack(states))
|
||||
actions = actor_net.variable(np.stack(actions))
|
||||
returns = torch.cat(returns, 0)
|
||||
advantages = torch.cat(advantages, 0).squeeze(1)
|
||||
advantages = (advantages - advantages.mean()) / advantages.std()
|
||||
|
||||
Reference in new issue
Block a user