@@ -452,8 +452,11 @@ def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = Fals
452452 start = time .time ()
453453 # Rollout
454454
455- mean_style_reward_log = 0
456- mean_task_reward_log = 0
455+ # Accumulate on-device (zero-dim CUDA tensors) so that we do NOT
456+ # force a GPU->CPU synchronization on every environment step. A single
457+ # .item() sync is performed once per iteration, after the rollout.
458+ mean_style_reward_log = torch .zeros ((), device = self .device )
459+ mean_task_reward_log = torch .zeros ((), device = self .device )
457460
458461 with torch .inference_mode ():
459462 for _ in range (self .num_steps_per_env ):
@@ -471,8 +474,8 @@ def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = Fals
471474 amp_obs , next_amp_obs
472475 )
473476
474- mean_task_reward_log += rewards .mean (). item ()
475- mean_style_reward_log += style_rewards .mean (). item ()
477+ mean_task_reward_log += rewards .mean ()
478+ mean_style_reward_log += style_rewards .mean ()
476479
477480 rewards = (
478481 1 - self .style_weight
@@ -507,8 +510,9 @@ def learn(self, num_learning_iterations: int, init_at_random_ep_len: bool = Fals
507510 start = stop
508511 self .alg .compute_returns (obs )
509512
510- mean_style_reward_log /= self .num_steps_per_env
511- mean_task_reward_log /= self .num_steps_per_env
513+ # Single synchronization point for the whole rollout.
514+ mean_style_reward_log = mean_style_reward_log .item () / self .num_steps_per_env
515+ mean_task_reward_log = mean_task_reward_log .item () / self .num_steps_per_env
512516
513517 (
514518 mean_value_loss ,
0 commit comments