Rate this Page
★ ★ ★ ★ ★

FQLTrainer#

class torchrl.trainers.algorithms.FQLTrainer(*args, **kwargs)[source]#

Train FQL with shared offline-to-online replay and Trainer lifecycle hooks.

See FQLTrainerConfig for the Hydra configuration counterpart.

The standard optimizer hook updates all networks before the target-network hook runs. Unlike reference FQL, the target EMA uses the updated critic. Allocate replay capacity for the offline dataset and online transitions. Online collection performs one update per batch, so use one frame per batch for one update per interaction. update() performs a single replay update.

collector may be omitted for offline-only training or provided as a factory initialized after pretraining. compile_loss compiles the loss without changing its checkpoint keys. See OfflineToOnlineTrainer for logging, checkpointing and hook options.

Parameters:
  • loss_module (FQLLoss) – flow, actor and critic objective.

  • optimizer (torch.optim.Optimizer) – optimizer for the loss parameters.

  • replay_buffer (ReplayBuffer) – offline dataset and online replay storage.

  • target_net_updater (TargetNetUpdater) – critic target update rule.

  • offline_steps (int) – gradient updates before online collection.

Keyword Arguments:
  • logger (Logger, optional) – scalar logger. Defaults to None.

  • checkpoint (Checkpoint, optional) – checkpoint component registry.

  • checkpoint_rotation (CheckpointRotation, optional) – checkpoint retention policy. Use with checkpoint; alternatively set save_trainer_file.

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 accept strict to control missing or incompatible components. Arguments are ignored when CKPT_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=torch checkpoints, weights_only=True is the default for safer deserialization. Pass weights_only=False explicitly only if the state dict contains custom objects. On torch < 2.4 the default is weights_only=False because the weights-only unpickler of those versions cannot deserialize the torch.device instances contained in TensorDict state-dicts.

Note

Explicit torch payloads and CKPT_BACKEND=torch checkpoints use mmap=True by default. Pass mmap=False for legacy pre-zipfile torch.save files or file-like objects. On Windows the default is mmap=False because 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_location to 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

file may also be a CheckpointRotation directory, 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.

shutdown() → None#

Release the online collector, if it was initialized.

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 calls request_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 raises KeyboardInterrupt. Previous handlers are restored on exit.

Parameters:

signals (Collection[int], optional) – signal numbers to handle. Defaults to SIGINT and SIGTERM.

Examples

>>> with trainer.stop_on_signal():  
...     trainer.train()
update() → TensorDictBase#

Run one replay update through the standard optimization hooks.