TD3Trainer#
- class torchrl.trainers.algorithms.TD3Trainer(*args, **kwargs)[source]#
A trainer class for Twin Delayed DDPG (TD3) algorithm.
See also
TD3TrainerConfigfor the Hydra configuration counterpart.This trainer implements the TD3 algorithm, an off-policy actor-critic method that builds on DDPG with improvements for stability including: - Clipped double Q-learning - Delayed policy updates - Target policy smoothing
The trainer handles: - Replay buffer management for off-policy learning - Target network updates (typically SoftUpdate) for stable training - Policy weight updates to the data collector - Comprehensive logging of training metrics
- Parameters:
collector (BaseCollector) – The data collector used to gather environment interactions.
total_frames (int) – Total number of frames to collect during training.
frame_skip (int) – Number of frames to skip between policy updates.
optim_steps_per_batch (int) – Number of optimization steps per collected batch.
loss_module (LossModule | Callable) – The TD3 loss module or a callable that computes losses.
optimizer (optim.Optimizer, optional) – Fallback optimizer for training. Defaults to None.
optimization_stepper (TD3OptimizationStepper, optional) – Custom optimization stepper controlling delayed actor/critic updates. Defaults to None.
logger (Logger, optional) – Logger for recording training metrics. Defaults to None.
clip_grad_norm (bool, optional) – Whether to clip gradient norms. Defaults to True.
clip_norm (float, optional) – Maximum gradient norm for clipping. Defaults to None.
progress_bar (bool, optional) – Whether to show a progress bar during training. Defaults to True.
seed (int, optional) – Random seed for reproducibility. Defaults to None.
save_trainer_interval (int, optional) – Interval for saving trainer state. Defaults to 10000.
log_interval (int, optional) – Interval for logging metrics. Defaults to 10000.
save_trainer_file (str | pathlib.Path, optional) – File path for saving trainer state. Defaults to None.
num_epochs (int, optional) – Number of epochs per batch. Defaults to 1 (typical for off-policy).
replay_buffer (ReplayBuffer, optional) – Replay buffer for storing and sampling experiences. Defaults to None.
batch_size (int, optional) – Global learner batch size. Defaults to the replay buffer batch size.
learner_backend (str) – Optimization placement,
"local"or"ray".learner_backend_options (dict, optional) – Ray world size and resources.
learner_poll_interval (float) – Remote replay polling interval.
enable_logging (bool, optional) – Whether to enable metric logging. Defaults to True.
log_rewards (bool, optional) – Whether to log reward statistics. Defaults to True.
log_actions (bool, optional) – Whether to log action statistics. Defaults to True.
log_observations (bool, optional) – Whether to log observation statistics. Defaults to False.
async_collection (bool, optional) – Whether to use async collection. Defaults to False.
log_timings (bool, optional) – Whether to log timing information. Defaults to False.
target_net_updater (TargetNetUpdater) – Target network updater (typically SoftUpdate).
exploration_module (torch.nn.Module, optional) – Optional exploration module appended to actor weights when syncing policy parameters to the collector. Defaults to None.
Note
This is an experimental/prototype feature. The API may change in future versions. TD3 is particularly effective for continuous control tasks.
- 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.
- request_stop(reason: str | None = None) None#
Signal that training should stop at the next loop boundary.