OfflineToOnlineTrainer#
- class torchrl.trainers.algorithms.OfflineToOnlineTrainer(*args, **kwargs)[source]#
Train from offline replay before optional online fine-tuning.
See also
OfflineToOnlineTrainerConfigfor the Hydra configuration counterpart.Builds on
SACTrainertarget updates, collector weight synchronization and logging. Offline updates precede online collection within the standard Trainer lifecycle. With a mixed buffer, online collection anneals the offline sampling fraction overanneal_frames.- Parameters:
collector (BaseCollector or callable, optional) – online collector or a factory called after offline training, unless loss initialization requires environment metadata first. Defaults to
None.total_frames (int) – online frames to collect. Defaults to zero.
frame_skip (int) – frames skipped between policy updates. Defaults to one.
optim_steps_per_batch (int) – updates per collected batch. Defaults to one.
loss_module (LossModule) – actor-critic objective with an
actor_network.replay_buffer (ReplayBuffer or OfflineToOnlineReplayBuffer) – regular replay storage or independently sampled offline and online buffers.
- Keyword Arguments:
anneal_frames (int, optional) – frames over which
offline_fractiondecays to 0. Defaults tototal_frames; pass<= 0to keep the fraction fixed.batch_size (int, optional) – replay-buffer sampling batch size.
offline_steps (int) – gradient updates before online collection (default zero).
device (device, optional) – device for sampled training batches.
compile_loss (bool) – compile the loss module, retaining its checkpoint keys.
A regular replay buffer retains offline and online transitions together. A mixed offline-to-online buffer instead controls their sampling fractions. The collector may be omitted for offline-only training or supplied as a factory, created when collection or environment metadata requires it. Losses are logged by update count during pretraining and by collected frames during online training.
See
SACTrainerfor the remaining keyword arguments.Note
Experimental/prototype feature; the API may change.
- property checkpoint_step: int#
Use update counts across both phases when pretraining is configured.
- property completed_steps: int#
Number of completed offline and online updates.
- compute_loss(sub_batch: TensorDictBase, method: str | None = None) TensorDictBase | tuple[Any, ...]#
Evaluate the configured loss through the active execution boundary.
- load_from_file(file: str | Path, **kwargs) Trainer#
Loads a file and its state-dict in the trainer.
Keyword arguments are passed to the
load()function for legacy torch checkpoints and unified components explicitly saved with the torch state-dict payload format. Unified checkpoints additionally acceptstrictto control missing or incompatible components. Arguments are ignored whenCKPT_BACKEND=memmap.Note
Unified state-dict components use TensorDict storage by default and do not invoke the pickle loader. For explicit torch payloads and
CKPT_BACKEND=torchcheckpoints,weights_only=Trueis the default for safer deserialization. Passweights_only=Falseexplicitly only if the state dict contains custom objects. On torch < 2.4 the default isweights_only=Falsebecause the weights-only unpickler of those versions cannot deserialize thetorch.deviceinstances contained in TensorDict state-dicts.Note
Explicit torch payloads and
CKPT_BACKEND=torchcheckpoints usemmap=Trueby default. Passmmap=Falsefor legacy pre-zipfiletorch.savefiles or file-like objects. On Windows the default ismmap=Falsebecause a mapped checkpoint keeps the file locked, preventing deletion or re-save.Note
Unified checkpoint tensors are mapped to CPU by default. Pass an explicit
map_locationto select another device mapping.Note
After restoring an independently registered policy component, the trainer synchronizes the collector once so local policy copies and remote workers observe the restored learner weights.
Note
filemay also be aCheckpointRotationdirectory, in which case its newest checkpoint is restored.
- optim_steps(batch: ~tensordict.base.TensorDictBase, *, optim_steps_per_batch: int | None | object = <object object>, num_epochs: int | object = <object object>) None#
Run the configured optimization loop for one collected batch.
Keyword overrides are applied only to this call and do not change the trainer configuration. They are useful for algorithms that need a one-time optimization schedule while retaining the standard Trainer hooks and logging behavior.
- request_stop(reason: str | None = None) None#
Signal that training should stop at the next loop boundary.
- stop_on_signal(signals: Collection[int] = (Signals.SIGINT, Signals.SIGTERM))#
Stop training cleanly when the process receives a termination signal.
Wrap
train()in this context. The first signal callsrequest_stop(), so the loop finishes the current batch, writes a final checkpoint when a save destination is configured, shuts the collector down and returns. A second signal raisesKeyboardInterrupt. Previous handlers are restored on exit.- Parameters:
signals (Collection[int], optional) – signal numbers to handle. Defaults to
SIGINTandSIGTERM.
Examples
>>> with trainer.stop_on_signal(): ... trainer.train()
- update() TensorDictBase[source]#
Run one replay update through the standard optimization hooks.