@@ -708,13 +708,21 @@ def _train_iter_off_policy(self):
708708 # 1. replay buffer is not ready (initial collect steps not reached)
709709 # 2. utd ratio is too high (training is too fast; wait for more data)
710710 next_replay_buffer_log_time = time .time ()
711+ utd_throttle_start_time = None
712+ utd_throttle_total_time = 0.
711713 while True :
712714 replay_buffer_size = self ._replay_buffer .total_size
713715 replay_buffer_not_ready = (replay_buffer_size
714716 < self ._config .initial_collect_steps )
715717 utd = self .utd ()
716718 utd_exceeded = utd > self ._max_utd_ratio
717719 now = time .time ()
720+ if utd_exceeded and utd_throttle_start_time is None :
721+ utd_throttle_start_time = now
722+ elif not utd_exceeded and utd_throttle_start_time is not None :
723+ wait_time = now - utd_throttle_start_time
724+ utd_throttle_total_time += wait_time
725+ utd_throttle_start_time = None
718726 if now >= next_replay_buffer_log_time :
719727 logging .info (
720728 f"Rank { self ._ddp_rank } replay buffer steps="
@@ -727,6 +735,10 @@ def _train_iter_off_policy(self):
727735 break
728736 time .sleep (0.01 )
729737
738+ if utd_throttle_total_time > 0 :
739+ alf .summary .scalar ("time/trainer_wait_for_utd" ,
740+ utd_throttle_total_time )
741+
730742 steps = super ()._train_iter_off_policy ()
731743 self ._total_updates += self ._config .num_updates_per_train_iter
732744
0 commit comments