Rate this Page
★ ★ ★ ★ ★

TdMpc2Loss#

class torchrl.objectives.TdMpc2Loss(*args, **kwargs)[source]#

Compute the TD-MPC2 model-learning objective.

Parameters:
  • world_model – TensorDict-native encoder, dynamics, and reward model.

  • policy_prior – TensorDict module producing sampled actions and policy statistics from a latent state.

  • q_ensemble – Distributional Q-function ensemble.

  • horizon – Number of transitions in each sampled sequence.

  • discount – Scalar discount factor used for TD targets.

  • rho – Temporal weighting factor for model and actor losses.

  • consistency_coef – Weight of latent consistency loss.

  • reward_coef – Weight of distributional reward loss.

  • value_coef – Weight of distributional value loss.

  • entropy_coef – Entropy coefficient in the actor objective.

  • scale_tau – Exponential averaging factor for the running value scale.

  • observation_key – Current observation key.

  • action_key – Action key.

  • reward_key – Reward key under "next".

  • terminated_key – Termination key under "next".

Examples

>>> import torch
>>> from tensordict import TensorDict
>>> from torchrl.objectives import TdMpc2Loss
>>> # The three TensorDict-native components are built from the
>>> # corresponding TD-MPC2 configuration classes.
>>> loss = TdMpc2Loss(world_model, policy_prior, q_ensemble)  
actor_loss_from_latents(latent_sequence: Tensor) → tuple[Tensor, TensorDictBase][source]#

Compute the actor objective from detached imagined latents.

default_keys#

alias of _AcceptedKeys

forward(tensordict: TensorDictBase = None) → TensorDictBase[source]#

Return the independent weighted model-loss components.

property in_keys: list[NestedKey]#

Return the canonical current and next transition keys.

model_loss(sample: TensorDictBase) → tuple[Tensor, TensorDictBase][source]#

Compute the model, reward, value, and consistency objectives.

Parameters:

sample – Canonical batch-major transition sequence with final batch dimension equal to horizon.

Returns:

The weighted model loss and detached latent/model metadata for the subsequent actor phase.

property out_keys: list[NestedKey]#

Return the scalar loss and diagnostic keys written by forward.