mirror of
https://github.com/wassname/DeepRL.git
synced 2026-09-09 11:13:47 +08:00
tensorboard logging
This commit is contained in:
@@ -65,6 +65,7 @@ class AsyncAgent:
|
||||
self.config.steps_lock = mp.Lock()
|
||||
self.config.network_lock = mp.Lock()
|
||||
self.config.total_steps = mp.Value('i', 0)
|
||||
self.config.total_epochs = mp.Value('i', 0)
|
||||
self.config.stop_signal = mp.Value('i', False)
|
||||
|
||||
def run(self):
|
||||
|
||||
+12
-2
@@ -138,8 +138,9 @@ class ProximalPolicyOptimization:
|
||||
obj = ratio * advantages
|
||||
obj_clipped = ratio.clamp(1.0 - self.config.ppo_ratio_clip, 1.0 + self.config.ppo_ratio_clip) * advantages
|
||||
policy_loss = -torch.min(obj, obj_clipped).mean(0)
|
||||
entropy = actor_net.entropy(std)
|
||||
if config.entropy_weight:
|
||||
policy_loss += -config.entropy_weight * actor_net.entropy(std)
|
||||
policy_loss += -config.entropy_weight * entropy
|
||||
|
||||
v = critic_net.predict(states)
|
||||
value_loss = 0.5 * (returns - v).pow(2).mean()
|
||||
@@ -148,7 +149,15 @@ class ProximalPolicyOptimization:
|
||||
self.worker_network.zero_grad()
|
||||
policy_loss.backward()
|
||||
value_loss.backward()
|
||||
nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
|
||||
grad_norm = nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip)
|
||||
|
||||
config.logger.scalar_summary('loss_policy', policy_loss, config.total_epochs.value)
|
||||
config.logger.scalar_summary('loss_value', value_loss, config.total_epochs.value)
|
||||
config.logger.scalar_summary('entropy', entropy, config.total_epochs.value)
|
||||
config.logger.histo_summary('ppo_ratio', ratio, config.total_epochs.value)
|
||||
config.logger.histo_summary('advantages', advantages, config.total_epochs.value)
|
||||
config.logger.histo_summary('grad_norm', grad_norm, config.total_epochs.value)
|
||||
|
||||
with config.network_lock:
|
||||
self.shared_network.zero_grad()
|
||||
self.actor_opt.zero_grad()
|
||||
@@ -156,5 +165,6 @@ class ProximalPolicyOptimization:
|
||||
sync_grad(self.shared_network, self.worker_network)
|
||||
self.actor_opt.step()
|
||||
self.critic_opt.step()
|
||||
config.total_epochs.value += 1
|
||||
|
||||
return batched_steps, batched_rewards
|
||||
|
||||
Reference in New Issue
Block a user