@@ -1503,6 +1503,10 @@ def train_from_replay_buffer(self, update_global_counter=False):
15031503 ``config.update_counter_every_mini_batch=True``.
15041504 """
15051505 config : TrainerConfig = self ._config
1506+ num_updates_per_train_iter = int (config .num_updates_per_train_iter ())
1507+
1508+ print (num_updates_per_train_iter )
1509+ print ("-----" )
15061510
15071511 # returns 0 if haven't started training yet, when ``_replay_buffer`` is
15081512 # not None and the number of samples in the buffer is less than
@@ -1535,7 +1539,7 @@ def _replay(batch_multiplier):
15351539 if config .whole_replay_buffer_training :
15361540 experience , batch_info = self ._replay_buffer .gather_all (
15371541 ignore_earliest_frames = True )
1538- num_updates = config . num_updates_per_train_iter
1542+ num_updates = num_updates_per_train_iter
15391543 else :
15401544 assert config .mini_batch_length is not None , (
15411545 "No mini_batch_length is specified for off-policy training"
@@ -1550,7 +1554,7 @@ def _replay(batch_multiplier):
15501554 if (config .sample_mini_batch_per_update
15511555 and not config .whole_replay_buffer_training ):
15521556 train_steps = 0
1553- for _ in range (config . num_updates_per_train_iter ):
1557+ for _ in range (num_updates_per_train_iter ):
15541558 experience , batch_info , num_updates , mini_batch_size = _replay (
15551559 1 )
15561560 with record_time ("time/train" ):
@@ -1567,7 +1571,7 @@ def _replay(batch_multiplier):
15671571 return train_steps
15681572 else :
15691573 experience , batch_info , num_updates , mini_batch_size = _replay (
1570- config . num_updates_per_train_iter )
1574+ num_updates_per_train_iter )
15711575 with record_time ("time/train" ):
15721576 return self ._train_experience (
15731577 experience ,
@@ -1597,7 +1601,7 @@ def _replay(batch_multiplier):
15971601 if (config .sample_mini_batch_per_update
15981602 and not config .whole_replay_buffer_training ):
15991603 train_steps = 0
1600- for _ in range (config . num_updates_per_train_iter ):
1604+ for _ in range (num_updates_per_train_iter ):
16011605 if self ._RL_train :
16021606 experience , batch_info , num_updates , mini_batch_size = _replay (
16031607 1 )
@@ -1623,7 +1627,7 @@ def _replay(batch_multiplier):
16231627 else :
16241628 if self ._RL_train :
16251629 experience , batch_info , num_updates , mini_batch_size = _replay (
1626- config . num_updates_per_train_iter )
1630+ num_updates_per_train_iter )
16271631 else :
16281632 experience = None
16291633 batch_info = None
@@ -1633,7 +1637,7 @@ def _replay(batch_multiplier):
16331637 with record_time ("time/offline_replay" ):
16341638 offline_experience , offline_batch_info = self ._offline_replay_buffer .get_batch (
16351639 batch_size = (mini_batch_size *
1636- config . num_updates_per_train_iter ),
1640+ num_updates_per_train_iter ),
16371641 batch_length = config .mini_batch_length )
16381642 # train hybrid
16391643 with record_time ("time/offline_train" ):
0 commit comments