data
PlaNetOutput
extends
ModelOutputPlaNetOutput(observation: Tensor, reward: Tensor, posterior_stoch: Tensor, posterior_mean: Tensor, posterior_std: Tensor, prior_mean: Tensor, prior_std: Tensor, deter: Tensor, loss: Tensor | None = None, recon_loss: Tensor | None = None, reward_loss: Tensor | None = None, kl_loss: Tensor | None = None, overshoot_kl_loss: Tensor | None = None, overshoot_reward_loss: Tensor | None = None)Forward output of the latent dynamics model.
Attributes
observationTensorReconstruction shaped
(B, T, C, 64, 64), decoded
from the posterior.rewardTensorPredicted reward shaped
(B, T).posterior_stochTensorThe sampled latent ,
(B, T, stoch_size).posterior_mean, posterior_stdTensorParameters of .
prior_mean, prior_stdTensorParameters of — what the dynamics predicted
before seeing the frame.
deterTensorThe deterministic path ,
(B, T, deter_size). Carried
once, not twice: prior and posterior share it by construction,
because the observation refines the belief about
without changing the path that led there.loss(Tensor or None, optional)recon_loss + reward_loss + kl_weight * kl_loss. None on
PlaNetModel, which builds no objective.recon_loss, reward_loss, kl_loss(Tensor or None, optional)The one-step terms separately.
kl_loss is already clamped
at free_nats.overshoot_kl_loss, overshoot_reward_loss(Tensor or None, optional)The latent-overshooting terms, averaged over distances.
None
when overshooting is off or the sequence is too short. Reported
separately because they enter loss with their own weights
and are otherwise invisible.Notes
Returned by both PlaNetModel.forward (losses None) and
PlaNetForWorldModeling.forward (losses populated).
Examples
>>> import lucid
>>> from lucid.models.generative.planet import PlaNetConfig, PlaNetModel
>>> cfg = PlaNetConfig(action_dim=2, stoch_size=4, deter_size=8,
... hidden_size=8, cnn_depth=4, reward_hidden=8)
>>> model = PlaNetModel(cfg).eval()
>>> out = model(lucid.randn((1, 3, 3, 64, 64)), lucid.zeros(1, 3, 2))
>>> out.observation.shape, out.reward.shape
((1, 3, 3, 64, 64), (1, 3))Used by 1
Constructors
1dunder
__init__
→None__init__(observation: Tensor, reward: Tensor, posterior_stoch: Tensor, posterior_mean: Tensor, posterior_std: Tensor, prior_mean: Tensor, prior_std: Tensor, deter: Tensor, loss: Tensor | None = None, recon_loss: Tensor | None = None, reward_loss: Tensor | None = None, kl_loss: Tensor | None = None, overshoot_kl_loss: Tensor | None = None, overshoot_reward_loss: Tensor | None = None)