GRPOTrainer#
- class torchrl.trainers.algorithms.GRPOTrainer(*args, **kwargs)[source]#
A trainer for LLM alignment using GRPO (or compatible) objectives.
See also
GRPOTrainerConfigfor the Hydra configuration counterpart.Warning
This is an experimental/prototype feature. The API may change in future versions. Please report any issues or feedback to help improve this implementation.
This trainer integrates the full GRPO training loop — mixed-precision, gradient accumulation, inference-weight synchronization, and LLM-specific logging — into the standard
Trainerhook system. Scalar diagnostics emitted by the loss (e.g.ESS,clip_fraction,kl_approxforGRPOLoss) are logged automatically after each optimization loop.It is designed to work with:
GRPOLoss(or anyLossModulewhose outputs start with"loss_")RayReplayBuffer(or anyReplayBuffer)
The weight-sync sender is intentionally decoupled from the trainer so that neither
vllmnorsglangneed to be imported by the core library.- Parameters:
collector (BaseCollector) – The data collector (typically a
RayLLMCollector).total_frames (int) – Total number of frames / dialog turns.
frame_skip (int) – Frame skip value (set to 1 for LLM tasks).
optim_steps_per_batch (int, optional) – Number of micro-batches drawn from the replay buffer per collected batch and epoch.
None(default) iterates over the whole replay buffer once per epoch.loss_module (LossModule) – The GRPO loss module.
optimizer (optim.Optimizer, optional) – Optimizer. Required when
optimization_stepperis not provided.optimization_stepper (OptimizationStepper, optional) – Custom stepper. If omitted, a
MixedPrecisionOptimizationStepperis constructed automatically fromoptimizerand the mixed-precision arguments below.weight_sync_sender (optional) – Object with an
update_weights()method used to push training weights to the inference engine. PassNoneto disable weight synchronization (useful for offline testing).weight_update_frequency (int, optional) – Optimizer steps between weight pushes to the inference engine when
async_collection=True(registered at thepost_optimstage throughUpdateWeights). In sync mode weights are pushed once per collected batch and this value is unused. Default:1.empty_replay_buffer_on_weight_update (bool, optional) – If
True, the replay buffer is emptied after each weight push (sync GRPO). Default:False.replay_buffer (ReplayBuffer, optional) – The replay buffer used for sampling.
batch_size (int, optional) – Override the replay buffer’s batch size.
device (torch.device, optional) – Device on which sampled batches are placed before the loss forward pass (typically the training device).
Noneleaves samples on their storage device.mixed_precision (bool, optional) – Enable autocast + GradScaler. Default:
False.autocast_dtype (torch.dtype, optional) – dtype for
autocast. Default:torch.bfloat16.gradient_accumulation_steps (int, optional) – Gradient accumulation. Default:
1.logger (Logger, optional) – Logger (e.g.
WandbLogger).clip_norm (float, optional) – Gradient clip norm, applied by the stepper. Default:
1.0.progress_bar (bool, optional) – Show a
tqdmprogress bar.seed (int, optional) – Random seed.
save_trainer_interval (int, optional) – Frame interval between saves.
log_interval (int, optional) – Frame interval between logs.
save_trainer_file (str | Path, optional) – Path for legacy saves.
checkpoint (Checkpoint, optional) – Unified checkpoint object.
checkpoint_rotation (CheckpointRotation, optional) – Rotation policy.
checkpoint_metadata (Callable, optional) – Extra metadata callback.
num_epochs (int, optional) – Epochs per collected batch. Default:
1.async_collection (bool, optional) – Whether data is collected asynchronously (
grpo-asyncmode). Default:False.log_timings (bool, optional) – Log timing of each hook. Default:
False.auto_log_optim_steps (bool, optional) – Log
optim_stepsafter each optimization loop. Default:True.log_rewards (bool, optional) – Log reward / return statistics. Default:
True.log_kl (bool, optional) – Log KL-divergence keys from the loss output. Default:
True.
Examples
>>> from torchrl.trainers.algorithms.grpo import GRPOTrainer >>> # Assuming you have a collector, loss_fn, optimizer, replay_buffer, >>> # and weight_sync_sender already constructed (see SOTA scripts): >>> trainer = GRPOTrainer( ... collector=collector, ... total_frames=cfg.train.total_dialog_turns, ... frame_skip=1, ... optim_steps_per_batch=cfg.train.epochs, ... loss_module=loss_fn, ... optimizer=optimizer, ... weight_sync_sender=sender, ... weight_update_frequency=1, ... empty_replay_buffer_on_weight_update=cfg.train.empty_replay_buffer, ... replay_buffer=replay_buffer, ... mixed_precision=cfg.train.mixed_precision, ... gradient_accumulation_steps=cfg.train.gradient_accumulation_steps, ... clip_norm=cfg.optimizer.clip_grad_norm, ... logger=wandb_logger, ... ) >>> trainer.train()
- 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()