Rate this Page
★ ★ ★ ★ ★

Trainer Basics#

Core trainer classes and builder utilities.

Trainer and hooks#

Trainer(*args, **kwargs)

A generic Trainer class.

TrainerHookBase()

An abstract hooking class for torchrl Trainer class.

MixedPrecisionOptimizationStepper(optimizer, *)

Optimization step with mixed precision and gradient accumulation.

Algorithm-specific trainers#

On-policy telemetry#

On-policy trainers expose telemetry="standard" by default. Standard mode adds diagnostics under the training/ logger namespace for collected and batch frames, completed episodes, terminal rates, reward and complete-episode summaries, optimizer learning rate and gradient norm, collection and optimizer throughput, and cheap collector or replay-buffer statistics when those values are available. Optional metrics are omitted when the collected batch does not contain enough information to compute them; for example, complete-episode returns require trajectory identifiers and reset markers.

Standard mode uses training/rewards/{min,mean,std,max} for transition reward summaries, without emitting legacy reward or terminal aliases. Synchronous training summarizes the collected batch. Fully asynchronous training summarizes valid transitions in replay samples instead, since no collected batch reaches the learner. Episode and terminal metrics are omitted in that mode: replay slices may be incomplete or carry artificial boundaries used for advantage estimation.

Set telemetry="minimal" to retain the legacy metric set without querying collector or replay statistics or computing the additional reductions. Legacy metric names such as r_training and done_percentage are emitted only in minimal mode.

OnPolicyTrainer(*args, **kwargs)

Shared implementation for on-policy trainers (PPO, A2C, REINFORCE).

A2CTrainer(*args, **kwargs)

A2C (Advantage Actor-Critic) trainer implementation.

PPOTrainer(*args, **kwargs)

PPO (Proximal Policy Optimization) trainer implementation.

ReinforceTrainer(*args, **kwargs)

REINFORCE (policy gradient with baseline) trainer implementation.

SACTrainer(*args, **kwargs)

A trainer class for Soft Actor-Critic (SAC) algorithm.

OfflineToOnlineTrainer(*args, **kwargs)

Train from offline replay before optional online fine-tuning.

DQNTrainer(*args, **kwargs)

A trainer class for Deep Q-Network (DQN) algorithm.

DDPGTrainer(*args, **kwargs)

A trainer class for Deep Deterministic Policy Gradient (DDPG) algorithm.

IQLTrainer(*args, **kwargs)

A trainer class for Implicit Q-Learning (IQL) algorithm.

FQLTrainer(*args, **kwargs)

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

CQLTrainer(*args, **kwargs)

A trainer class for Conservative Q-Learning (CQL) algorithm.

TD3Trainer(*args, **kwargs)

A trainer class for Twin Delayed DDPG (TD3) algorithm.

GRPOTrainer(*args, **kwargs)

A trainer for LLM alignment using GRPO (or compatible) objectives.

TdMpc2OptimizationStepper(loss_module, ...)

Execute the two-phase TD-MPC2 learner update.

Offline and online FQL#

FQLTrainer uses OfflineToOnlineTrainer replay hooks and the standard Trainer lifecycle. Supply a preloaded replay buffer and offline_steps for offline training. An optional collector factory and total_frames add online fine-tuning; the factory is called after pretraining. An empty replay buffer is supported when offline_steps=0.

Use logger, checkpoint and checkpoint_rotation as with other trainers. When pretraining is configured, checkpoint intervals and rotation filenames use optimization steps across both phases. load_from_file restores the completed budgets, optimizer, replay and target updater before training continues. compile_loss=True preserves the loss module’s checkpoint keys.

Target networks use the standard post-optimizer update. This differs from the reference FQL implementation’s pre-optimizer EMA; learning equivalence requires a matched training comparison.

PPO from an environment#

PPOTrainer.from_env() builds the standard collector, clipped PPO loss, Adam optimizer, GAE and minibatches. Supply the networks and training budget:

trainer = PPOTrainer.from_env(
    env, actor=actor, critic=critic,
    total_frames=1_000_000,
    frames_per_batch=4096,
    minibatch_size=256,
)
trainer.train()

For a recurrent policy, keep consecutive time windows and, when appropriate, normalize advantages separately within each task:

trainer = PPOTrainer.from_env(
    env, actor=actor, critic=critic,
    total_frames=1_000_000,
    frames_per_batch=4096,
    minibatch_size=256,
    sub_traj_len=64,
    gae_kwargs={"group_key": "task_id", "average_gae": True},
)

The collector installs missing policy primers and initialization tracking. The environment’s unique action and reward keys are inferred, including nested keys; value_key selects the critic output. Episode boundaries default to siblings of the reward, falling back to root-level keys. Pass explicit trainer key arguments when a task needs different boundaries. Multi-agent rewards and boundaries must have compatible shapes, as required by GAE.

sub_traj_len counts consecutive steps per environment, while frames_per_batch and minibatch_size count transitions across all environments. Feedforward PPO shuffles individual transitions; recurrent PPO samples contiguous windows. Training closes the collector’s environment: create a fresh environment for evaluation.

Existing logging and checkpoint options are forwarded to the trainer. Call trainer.load_from_file(path) before train() to resume a saved run. Use the ordinary constructor when supplying custom training components.

Hydra can target the same factory without duplicating its defaults:

_target_: torchrl.trainers.algorithms.PPOTrainer.from_env
total_frames: 1000000
frames_per_batch: 4096
minibatch_size: 256
sub_traj_len: 64
gae_kwargs:
  group_key: task_id
  average_gae: true

Instantiate it with instantiate(cfg, env=env, actor=actor, critic=critic). The existing PPOTrainerConfig continues to support explicit component configuration.

classmethod PPOTrainer.from_env(env: EnvBase, *, actor: TensorDictModuleBase, critic: TensorDictModuleBase, total_frames: int, frames_per_batch: int = 1024, minibatch_size: int = 256, sub_traj_len: int | None = None, learning_rate: float = 0.0003, value_key: NestedKey = 'state_value', loss_kwargs: Mapping[str, Any] | None = None, gae_kwargs: Mapping[str, Any] | None = None, collector_kwargs: Mapping[str, Any] | None = None, **trainer_kwargs: Any) → PPOTrainer[source]#

Build a PPO trainer from an environment, actor and critic.

Constructs a Collector, ClipPPOLoss, Adam optimizer, GAE and minibatch sampling. Use the ordinary constructor to supply custom collectors, losses, optimizers or replay buffers.

Parameters:
  • env (EnvBase) – Environment owned by the resulting collector. Training closes it. The collector installs missing policy primers and initialization tracking by default.

  • actor (TensorDictModuleBase) – Probabilistic actor returning actions and their log probabilities.

  • critic (TensorDictModuleBase) – Value network writing value_key.

  • total_frames (int) – Total environment transitions to collect. For closed-loop action deployment these count high-level decisions.

  • frames_per_batch (int, optional) – Transitions collected per update, across all environments. Defaults to 1024.

  • minibatch_size (int, optional) – Transitions per optimization step. Clamped to the collected batch size. Defaults to 256.

  • sub_traj_len (int, optional) – Consecutive time steps per recurrent training window. When set, uses BatchSubSampler and recurrent-mode GAE instead of flattening time into replay. The minibatch size must be a multiple of this length. Defaults to None (feedforward PPO).

  • learning_rate (float, optional) – Adam learning rate. Defaults to 3e-4.

  • value_key (NestedKey, optional) – Critic output key, also configured on the loss and GAE. Defaults to "state_value".

  • loss_kwargs (Mapping, optional) – Extra ClipPPOLoss arguments. Advantage normalization defaults to True, or False when GAE already normalizes advantages (e.g. within each task).

  • gae_kwargs (Mapping, optional) – Extra GAE arguments. For per-task normalization, pass {"group_key": "task_id", "average_gae": True}. Recurrent windows default to shifted=False, deactivate_vmap=True because their value networks may depend on recurrent state that cannot be reconstructed by shifting observations alone.

  • collector_kwargs (Mapping, optional) – Extra Collector arguments, such as policy_device and storing_device.

  • **trainer_kwargs – Additional OnPolicyTrainer options, such as num_epochs, gamma, lmbda, logging and checkpointing. Action/reward keys default to the environment’s unique keys. Done/terminated keys use the reward’s namespace when available, otherwise the root namespace; explicit key overrides take precedence. frame_skip defaults to 1 and clip_norm to 1.0.

Returns:

Configured trainer; call train() to start learning.

Return type:

PPOTrainer

Examples

>>> import torch
>>> from tensordict.nn import TensorDictModule, NormalParamExtractor
>>> from torchrl.modules import ProbabilisticActor, TanhNormal
>>> from torchrl.testing.mocking_classes import ContinuousActionVecMockEnv
>>> env = ContinuousActionVecMockEnv()
>>> obs_dim = env.observation_spec["observation"].shape[-1]
>>> action_dim = env.action_spec.shape[-1]
>>> actor = ProbabilisticActor(
...     TensorDictModule(
...         torch.nn.Sequential(torch.nn.Linear(obs_dim, 2 * action_dim), NormalParamExtractor()),
...         in_keys=["observation"], out_keys=["loc", "scale"],
...     ),
...     in_keys=["loc", "scale"], distribution_class=TanhNormal,
...     return_log_prob=True,
... )
>>> critic = TensorDictModule(
...     torch.nn.Linear(obs_dim, 1), in_keys=["observation"], out_keys=["state_value"],
... )
>>> trainer = PPOTrainer.from_env(
...     env, actor=actor, critic=critic, total_frames=32,
...     frames_per_batch=16, minibatch_size=8, progress_bar=False,
... )
>>> trainer.train()

Builders#

make_collector_offpolicy(make_env, ...[, ...])

Returns a data collector for off-policy sota-implementations.

make_collector_onpolicy(make_env, ...[, ...])

Makes a collector in on-policy settings.

make_dqn_loss(model, cfg)

Builds the DQN loss module.

make_replay_buffer(device, cfg)

Builds a replay buffer using the config built from ReplayArgsConfig.

make_target_updater(cfg, loss_module)

Builds a target network weight update object.

make_trainer(collector, loss_module[, ...])

Creates a Trainer instance given its constituents.

parallel_env_constructor(cfg, **kwargs)

Returns a parallel environment from an argparse.Namespace built with the appropriate parser constructor.

sync_async_collector(env_fns, env_kwargs[, ...])

Runs asynchronous collectors, each running synchronous environments.

sync_sync_collector(env_fns, env_kwargs[, ...])

Runs synchronous collectors, each running synchronous environments.

transformed_env_constructor(cfg[, ...])

Returns an environment creator from an argparse.Namespace built with the appropriate parser constructor.

Utils#

correct_for_frame_skip(cfg)

Correct the arguments for the input frame_skip, by dividing all the arguments that reflect a count of frames by the frame_skip.

get_stats_random_rollout(cfg[, ...])

Gathers stas (loc and scale) from an environment using random rollouts.