TdMpc2OptimizationStepper#
- class torchrl.trainers.algorithms.TdMpc2OptimizationStepper(loss_module: TdMpc2Loss, optimizer_model: Optimizer, optimizer_actor: Optimizer, *, target_tau: float = 0.01, zero_grad_set_to_none: bool = True)[source]#
Execute the two-phase TD-MPC2 learner update.
The model optimizer updates the world model and online Q-functions first; the actor optimizer then updates the policy objective on the detached imagined latent sequence captured by that model update. The policy update uses the model/Q parameters after the first optimizer step. The target Q-functions are soft-updated last.
- Parameters:
loss_module – TD-MPC2 loss providing model and actor objectives.
optimizer_model – Optimizer for the world model and online Q-functions.
optimizer_actor – Optimizer for the policy prior.
target_tau – Target-Q Polyak averaging factor. Defaults to
0.01.zero_grad_set_to_none – Whether optimizer
zero_gradcalls set gradients toNone. Defaults toTrue.
See also
TdMpc2OptimizationStepperConfig, TD-MPC2: Scalable, Robust World Models for Continuous Control.- register(trainer: Trainer, name: str = 'optimization_stepper') None[source]#
Register the stepper and validate exclusive trainer ownership.
- step(trainer: Trainer, sub_batch: TensorDictBase) TensorDictBase[source]#
Run model, actor, and target-Q updates and return detached metrics.
- property update_count: int#
Number of completed TD-MPC2 updates.