TdMpc2QEnsemble#
- class torchrl.modules.TdMpc2QEnsemble(*args, **kwargs)[source]#
Vectorized TD-MPC2 ensemble of distributional Q-functions.
Each Q-function receives the concatenation of a latent state and an action and returns logits for a scalar categorical representation. The ensemble dimension is exposed immediately before the category dimension.
forward()returns the logits of every Q-function.reduce()follows TD-MPC2’s value-estimation rule by selecting two Q-functions at random, decoding their logits, and taking either their minimum or average. Online, detached, and target parameter sources are available for both operations.- Parameters:
q_networks – Sequence of identically shaped single-network Q-functions.
num_bins – Number of categorical bins. Must be greater than 1.
vmin – Minimum value of the symlog-space categorical support.
vmax – Maximum value of the symlog-space categorical support.
in_keys – Two TensorDict keys for the latent state and action. Defaults to
["latent", "action"].out_keys – One TensorDict key for the ensemble logits. Defaults to
["q_logits"].q_value_key – TensorDict key written by
reduce(). Defaults to"q_value".
Examples
>>> import torch >>> from tensordict import TensorDict >>> from torchrl.modules import MLP >>> from torchrl.modules.models.tdmpc2 import TdMpc2QEnsemble >>> q_networks = [ ... MLP(in_features=6, out_features=5, depth=2, num_cells=8) ... for _ in range(5) ... ] >>> q_ensemble = TdMpc2QEnsemble( ... q_networks, num_bins=5, vmin=-10.0, vmax=10.0 ... ) >>> data = TensorDict( ... {"latent": torch.randn(4, 4), "action": torch.randn(4, 2)}, ... batch_size=[4], ... ) >>> data = q_ensemble(data) >>> data["q_logits"].shape torch.Size([4, 5, 5]) >>> q_ensemble.reduce(data, reduction="min")["q_value"].shape torch.Size([4, 1])
- forward(tensordict: TensorDictBase, *, source: Literal['online', 'detached', 'target'] = 'online') TensorDictBase[source]#
Write the logits of all Q-functions to the input TensorDict.
- Parameters:
tensordict – TensorDict containing the latent state and action.
source – Parameter source used for the Q-functions. Defaults to
"online".
- reduce(tensordict: TensorDictBase, *, reduction: Literal['min', 'avg'], source: Literal['online', 'detached', 'target'] = 'online') TensorDictBase[source]#
Write a TD-MPC2 two-Q value reduction to the input TensorDict.
- Parameters:
tensordict – TensorDict containing the latent state and action.
reduction – Either
"min"or"avg"for the two sampled Q-functions.source – Parameter source used for the Q-functions. Defaults to
"online".