Rate this Page
★ ★ ★ ★ ★

torchrl.trainers.algorithms.configs.modules.TdMpc2PolicyPriorConfig#

class torchrl.trainers.algorithms.configs.modules.TdMpc2PolicyPriorConfig(_partial_: bool = False, in_keys: Any = None, out_keys: Any = None, shared: bool = False, latent_dim: int = '???', action_dim: int = '???', mlp_dim: int = 512, log_std_min: float = -10.0, log_std_max: float = 2.0, device: Any = None, _target_: str = 'torchrl.trainers.algorithms.configs.modules._make_tdmpc2_policy_prior')[source]#

Configuration for a TD-MPC2 policy prior.

The policy prior maps a latent state to a sampled action and the distribution statistics used by the TD-MPC2 objective. By default, the module reads "latent" and writes "action", "mean"`, ``"log_std", "entropy", and "scaled_entropy".

Example

>>> import torch
>>> from hydra.utils import instantiate
>>> from tensordict import TensorDict
>>> from torchrl.trainers.algorithms.configs import TdMpc2PolicyPriorConfig
>>> cfg = TdMpc2PolicyPriorConfig(latent_dim=8, action_dim=2, mlp_dim=16)
>>> policy = instantiate(cfg)
>>> td = TensorDict({"latent": torch.randn(3, 8)}, batch_size=[3])
>>> assert policy(td)["action"].shape == (3, 2)