Skip to content

Commit 1c4e530

Browse files
committed
Make num_updates_per_train_iter scheduleable
1 parent 11463e7 commit 1c4e530

6 files changed

Lines changed: 23 additions & 12 deletions

File tree

alf/algorithms/algorithm.py

Lines changed: 10 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -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"):

alf/algorithms/config.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -285,8 +285,8 @@ def __init__(self,
285285
initial_collect_steps (int): if positive, number of steps each single
286286
environment steps before perform first update. Only used
287287
by ``OffPolicyAlgorithm``.
288-
num_updates_per_train_iter (int): number of optimization steps for
289-
one iteration. Only used by ``OffPolicyAlgorithm``.
288+
num_updates_per_train_iter (int|Scheduler): number of optimization
289+
steps for one iteration. Only used by ``OffPolicyAlgorithm``.
290290
sample_mini_batch_per_update (bool): If True and
291291
``whole_replay_buffer_training`` is False, sample one minibatch
292292
per update instead of sampling ``num_updates_per_train_iter``
@@ -440,7 +440,8 @@ def __init__(self,
440440
self.summarize_action_distributions = summarize_action_distributions
441441
self.summarize_output = summarize_output
442442
self.initial_collect_steps = initial_collect_steps
443-
self.num_updates_per_train_iter = num_updates_per_train_iter
443+
self.num_updates_per_train_iter = as_scheduler(
444+
num_updates_per_train_iter)
444445
self.sample_mini_batch_per_update = sample_mini_batch_per_update
445446
self.mini_batch_length = mini_batch_length
446447
self.mini_batch_size = mini_batch_size

alf/algorithms/distributed_off_policy_algorithm.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -560,7 +560,7 @@ def _train_iter_off_policy(self):
560560
time.sleep(0.01)
561561

562562
steps = super()._train_iter_off_policy()
563-
self._total_updates += self._config.num_updates_per_train_iter
563+
self._total_updates += int(self._config.num_updates_per_train_iter())
564564

565565
with record_time("time/trainer_send_params_to_unroller"):
566566
if (self._total_updates %

alf/algorithms/muzero_representation_learner.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1027,7 +1027,8 @@ def __init__(self,
10271027
updated.clear_replay_buffer = False
10281028
updated.mini_batch_length = training_options.mini_batch_length
10291029
updated.mini_batch_size = training_options.mini_batch_size
1030-
updated.num_updates_per_train_iter = training_options.num_updates_per_train_iter
1030+
updated.num_updates_per_train_iter = as_scheduler(
1031+
training_options.num_updates_per_train_iter)
10311032
updated.replay_buffer_length = training_options.replay_buffer_length
10321033
updated.initial_collect_steps = training_options.initial_collect_steps
10331034
updated.priority_replay = training_options.priority_replay

alf/algorithms/ppg/ppg_aux_algorithm.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
from alf.data_structures import TimeStep, AlgStep, LossInfo
2525
from alf.experience_replayers.replay_buffer import ReplayBuffer
2626
from alf.utils import dist_utils
27+
from alf.utils.schedulers import as_scheduler
2728
from alf.tensor_specs import TensorSpec
2829

2930
# Data structure to store the options for PPG's auxiliary phase
@@ -109,7 +110,8 @@ def __init__(self,
109110
updated_config.mini_batch_length = (aux_options.mini_batch_length
110111
or config.unroll_length)
111112
updated_config.mini_batch_size = aux_options.mini_batch_size
112-
updated_config.num_updates_per_train_iter = aux_options.num_updates_per_train_iter
113+
updated_config.num_updates_per_train_iter = as_scheduler(
114+
aux_options.num_updates_per_train_iter)
113115

114116
# Since we are going to store already-transformed experience in the
115117
# replay buffer, the aux algorithm shall not inherit the data

alf/algorithms/rlpd_algorithm.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -186,6 +186,9 @@ def __init__(self,
186186
else:
187187
total_utd = alf.config_util.get_config_value(
188188
"num_updates_per_train_iter")
189+
if callable(total_utd):
190+
total_utd = total_utd()
191+
total_utd = int(total_utd)
189192
if critic_utd is not None:
190193
assert critic_utd < total_utd, (
191194
"critic_utd should be less than num_updates_per_train_iter"

0 commit comments

Comments
 (0)