From 0bf7e094d431c4da0625f069c53300d074d5fded Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 25 Sep 2026 17:01:58 -0700 Subject: [PATCH 01/10] Add worldmodel transformer plan and action heads --- torchtitan/experiments/worldmodel/dataset.py | 10 ++ torchtitan/experiments/worldmodel/loss.py | 19 ++++ torchtitan/experiments/worldmodel/model.py | 105 ++++++++++++++++-- .../worldmodel/model_for_inference.py | 47 ++++++-- torchtitan/experiments/worldmodel/trainer.py | 19 +++- 5 files changed, 176 insertions(+), 24 deletions(-) diff --git a/torchtitan/experiments/worldmodel/dataset.py b/torchtitan/experiments/worldmodel/dataset.py index ab24339394e..f1d04c2f481 100644 --- a/torchtitan/experiments/worldmodel/dataset.py +++ b/torchtitan/experiments/worldmodel/dataset.py @@ -34,6 +34,7 @@ class _DiffusionConfig: train_skip: int val_skip: int nan_engaged_plans: bool + action_targets: bool = False def skip(self, val: bool) -> int: return self.val_skip if val else self.train_skip @@ -75,6 +76,13 @@ def _mock_batch(self) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]: } ) targets = {"plan": np.random.randn(batch, PLAN_SIZE).astype(np.float32)} + if cfg.action_targets: + from xx.training.path.dataloader import compute_action_target_from_plan + + inputs["action_t"] = np.random.uniform(0, 1, (batch, 2)).astype(np.float32) + targets["action"] = compute_action_target_from_plan( + targets["plan"][:, :495].reshape(-1, 33, 15), inputs["action_t"] + ) return inputs, targets @@ -106,6 +114,7 @@ class Config(BaseDataLoader.Config): limit: int | None mock_data: bool mock_segment_batch_size: int + action_targets: bool = False def __post_init__(self) -> None: total_frames = self.context_size_frames + self.future_size_frames @@ -226,6 +235,7 @@ def _build_dataset(config: Config, *, val: bool, global_rank: int, global_world_ train_skip=config.train_skip, val_skip=config.val_skip, nan_engaged_plans=config.nan_engaged_plans, + action_targets=config.action_targets, ), val=val, local_rank=int(os.environ.get("LOCAL_RANK", "0")), diff --git a/torchtitan/experiments/worldmodel/loss.py b/torchtitan/experiments/worldmodel/loss.py index f433229a5ea..007190820c2 100644 --- a/torchtitan/experiments/worldmodel/loss.py +++ b/torchtitan/experiments/worldmodel/loss.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + from __future__ import annotations import math @@ -41,6 +47,7 @@ def compute_worldmodel_losses( targets: dict[str, torch.Tensor], *, plan_loss_weight: float, + action_loss_weight: float = 1.0, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: loss: torch.Tensor | None = None terms: dict[str, torch.Tensor] = {} @@ -66,6 +73,16 @@ def compute_worldmodel_losses( terms["plan_loss"] = plan_loss.detach() terms["plan_mse"] = plan_mse.detach() + if "action" in outputs: + pred = outputs["action"].float() + target = targets["action"].to(device=pred.device, dtype=torch.float32) + values, err, mask = laplacian_density_loss(target, pred) + action_loss = values.mean(dim=-1) + weighted = action_loss_weight * action_loss + loss = weighted if loss is None else loss + weighted + terms["action_loss"] = action_loss.detach() + terms["action_mse"] = (err.square() * mask).sum(-1).div(mask.sum(-1).clamp_min(1)).detach() + if loss is None: raise RuntimeError("worldmodel produced no trainable outputs") terms["loss"] = loss.detach() @@ -78,6 +95,7 @@ class WorldModelLoss(BaseLoss): @dataclass(kw_only=True, slots=True) class Config(BaseLoss.Config): plan_loss_weight: float + action_loss_weight: float = 1.0 def __init__(self, config: Config, *, compile_config: CompileConfig | None = None): plan_loss_weight = config.plan_loss_weight @@ -90,6 +108,7 @@ def loss_fn( outputs, targets, plan_loss_weight=plan_loss_weight, + action_loss_weight=config.action_loss_weight, ) self.fn = loss_fn diff --git a/torchtitan/experiments/worldmodel/model.py b/torchtitan/experiments/worldmodel/model.py index 9aa0f04a523..df55c0732c8 100644 --- a/torchtitan/experiments/worldmodel/model.py +++ b/torchtitan/experiments/worldmodel/model.py @@ -282,7 +282,7 @@ def _init_plan_bias_(bias: torch.Tensor) -> None: local = _local_tensor(bias) with torch.no_grad(): local.zero_() - split = PLAN_SIZE // 2 + split = bias.shape[0] // 2 if not isinstance(bias, DTensor): local[split:].fill_(math.log(PLAN_HEAD_INIT_LOG_SIGMA_SCALE)) return @@ -611,12 +611,22 @@ def __init__(self, config: "WorldModel.Config", linears: PlanHeadLinearsConfig): ) self.head = linears.head.build() self.scale_layer = ScaleLayer(PLAN_SIZE) + self.action_t_encoder = nn.Linear(2, config.plan_head.n_embd) if config.plan_head_action else None + if self.action_t_encoder is not None: + self.action_head = nn.Linear(config.plan_head.n_embd, 4) + self.action_scale = ScaleLayer(4) self.init_weights() - def forward(self, x: torch.Tensor) -> torch.Tensor: + def forward(self, x: torch.Tensor, action_t: torch.Tensor | None = None) -> torch.Tensor | dict[str, torch.Tensor]: + if self.action_t_encoder is not None: + assert action_t is not None + x = x + self.action_t_encoder(action_t.to(x.dtype)) for mlp in self.mlps: x = mlp(x) - return self.scale_layer(self.head(x)) + plan = self.scale_layer(self.head(x)) + if self.action_t_encoder is not None: + return {"plan": plan, "action": self.action_scale(self.action_head(x))} + return plan def init_weights(self) -> None: for module in self.mlps.modules(): @@ -626,6 +636,62 @@ def init_weights(self) -> None: if self.head.bias is not None: _init_plan_bias_(self.head.bias) self.scale_layer.init_weights() + if self.action_t_encoder is not None: + init_transformer_linear_weights(self.action_t_encoder) + _init_normal_(self.action_head.weight, std=PLAN_HEAD_INIT_STD) + _init_plan_bias_(self.action_head.bias) + self.action_scale.init_weights() + + +class PlanTransformerBlock(nn.Module): + def __init__(self, config: TransformerConfig, linears: FFNLinearsConfig): + super().__init__() + self.attn = SelfAttention(config, self_attention_linears_config(config)) + self.mlp = residual_ffn(config, linears) + self.init_weights() + + def forward(self, x, input_mask=None, cache_pos=None, cache_seq_length=None): + if cache_pos is None: + attention = self.attn(x, input_mask=input_mask) + else: + attention = self.attn(x, input_mask=input_mask, cache_pos=cache_pos, cache_seq_length=cache_seq_length) + return self.mlp(x + attention) + + def init_weights(self): + for module in self.modules(): + if isinstance(module, nn.Linear): + init_transformer_linear_weights(module) + + +class TransformerPlanHead(nn.Module): + def __init__(self, config: "WorldModel.Config", linears: PlanHeadLinearsConfig): + super().__init__() + width = config.plan_head.n_embd + self.mlps = nn.ModuleList() + self.blocks = nn.ModuleList(PlanTransformerBlock(config.plan_head, block) for block in linears.blocks) + self.action_t_encoder = nn.Linear(2, width) + self.head = linears.head.build() + self.scale_layer = ScaleLayer(PLAN_SIZE) + self.action_head = nn.Linear(width, 4) + self.action_scale = ScaleLayer(4) + self.init_weights() + + def forward(self, x, action_t, input_mask=None, cache_pos=None, cache_seq_length=None): + x = x + self.action_t_encoder(action_t.to(x.dtype))[:, None] + for block in self.blocks: + x = block(x, input_mask, cache_pos, cache_seq_length) + x = x[:, -1].to(self.head.weight.dtype) + return {"plan": self.scale_layer(self.head(x)), "action": self.action_scale(self.action_head(x))} + + def init_weights(self): + for block in self.blocks: + block.init_weights() + init_transformer_linear_weights(self.action_t_encoder) + for head in (self.head, self.action_head): + _init_normal_(head.weight, std=PLAN_HEAD_INIT_STD) + _init_plan_bias_(head.bias) + self.scale_layer.init_weights() + self.action_scale.init_weights() class DiTBlock(nn.Module): @@ -725,6 +791,8 @@ class Config(BaseModel.Config): transformer: TransformerConfig plan_head: TransformerConfig experimental_pose_only_xy: bool + plan_head_transformer: bool = False + plan_head_action: bool = False x_embedder: PatchEmbedderLinearsConfig = field(init=False) augments_pos_ref_augment_embedder: ConditioningEmbedderLinearsConfig = field(init=False) ref_augment_from_augments_euler_embedder: ConditioningEmbedderLinearsConfig = field(init=False) @@ -754,6 +822,9 @@ def _sync_derived_fields(self) -> None: self.transformer.block_size = self.num_patches self.transformer.attention_mask_mini_block_size = self.num_spatial_patches self.plan_head.n_embd = self.transformer.n_embd + if self.plan_head_transformer: + self.plan_head.block_size = self.num_patches + self.plan_head.attention_mask_mini_block_size = self.num_spatial_patches hidden = self.transformer.n_embd pose_half = self.pose_size // 2 current_blocks = getattr(self, "blocks", []) @@ -850,7 +921,8 @@ def __init__(self, config: Config): self.fidx_embedder = DiscreteEmbedder(50, config.transformer.n_embd, config.fidx_embedder) self.blocks = nn.ModuleList(DiTBlock(config, config.blocks[i]) for i in range(config.transformer.n_layer)) self.final_layer = FinalLayer(config, config.final_layer) if config.final_layer is not None else None - self.plan_head = PlanHead(config, config.plan_head_linears) if config.plan_head_linears is not None else None + head_cls = TransformerPlanHead if config.plan_head_transformer else PlanHead + self.plan_head = head_cls(config, config.plan_head_linears) if config.plan_head_linears is not None else None self.register_buffer("pos_embed", torch.empty(1, config.num_patches, config.transformer.n_embd)) self.mask: TensorOrMask | None = None self.init_states(buffer_device=self.pos_embed.device) @@ -938,6 +1010,7 @@ def forward( cache_pos: torch.Tensor | None = None, cache_seq_length: int | None = None, input_mask: TensorOrMask | None = None, + action_t: torch.Tensor | None = None, ) -> dict[str, torch.Tensor]: if input_pos is None: input_mask = self.mask @@ -968,7 +1041,13 @@ def forward( x = block(x, t6, input_pos_t, input_mask, cache_pos, cache_seq_length) outputs = {} if return_plan and self.plan_head is not None: - outputs["plan"] = self.plan_head(x[:, -1, :]) + if isinstance(self.plan_head, TransformerPlanHead): + assert action_t is not None + outputs.update(self.plan_head(x, action_t, input_mask, cache_pos, cache_seq_length)) + elif self.config.plan_head_action: + outputs.update(self.plan_head(x[:, -1, :], action_t)) + else: + outputs["plan"] = self.plan_head(x[:, -1, :]) if self.final_layer is not None: outputs["sample"] = self.unpatchify(self.final_layer(x, t2, input_pos_t)) return outputs @@ -1037,11 +1116,11 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module: wrap(block, f"blocks.{layer_id}"), ) if model.plan_head is not None: - for layer_id, block in model.plan_head.mlps.named_children(): - model.plan_head.mlps.register_module( - layer_id, - wrap(block, f"plan_head.mlps.{layer_id}"), - ) + for name in ("mlps", "blocks"): + layers = getattr(model.plan_head, name, None) + if layers is not None: + for layer_id, block in layers.named_children(): + layers.register_module(layer_id, wrap(block, f"plan_head.{name}.{layer_id}")) logger.info(f"Applied {mode} activation checkpointing to the worldmodel") @@ -1066,6 +1145,9 @@ def _apply_compile(model: WorldModel, compile_config: CompileConfig) -> None: if model.plan_head is not None: for block in model.plan_head.mlps: block.compile(backend=compile_config.backend, fullgraph=True) + if isinstance(model.plan_head, TransformerPlanHead): + for block in model.plan_head.blocks: + block.compile(backend=compile_config.backend, fullgraph=True) logger.info("Compiling worldmodel components with torch.compile") @@ -1113,6 +1195,9 @@ def _apply_fsdp( if model.plan_head is not None: for block in model.plan_head.mlps: fully_shard(block, **fsdp_config, reshard_after_forward=reshard_after_forward) + if isinstance(model.plan_head, TransformerPlanHead): + for block in model.plan_head.blocks: + fully_shard(block, **fsdp_config, reshard_after_forward=reshard_after_forward) fully_shard(model.plan_head, **fsdp_config, reshard_after_forward=reshard_after_forward) if model.final_layer is not None: fully_shard( diff --git a/torchtitan/experiments/worldmodel/model_for_inference.py b/torchtitan/experiments/worldmodel/model_for_inference.py index dae30f99797..61c43c6d026 100644 --- a/torchtitan/experiments/worldmodel/model_for_inference.py +++ b/torchtitan/experiments/worldmodel/model_for_inference.py @@ -228,7 +228,7 @@ def __init__( default_kv_cache_dtype: KVCacheDType = FP8_KV_CACHE_DTYPE, ): super().__init__(config) - for block in self.blocks: + for block in self._attention_blocks(): block.attn.__class__ = InferenceSelfAttention self.inference_masks: dict[tuple[str, int, int, bool], tuple[TensorOrMask | None, TensorOrMask | None]] = {} self.max_batch_size = -1 @@ -236,6 +236,9 @@ def __init__( self.cache_dtype: torch.dtype | None = None self.default_kv_cache_dtype = default_kv_cache_dtype + def _attention_blocks(self): + return (*self.blocks, *getattr(self.plan_head, "blocks", ())) + @staticmethod def _cast_plan_head_input_to_float32( _module: nn.Module, @@ -247,21 +250,33 @@ def _ensure_plan_head_float32(self) -> None: if self.plan_head is None or getattr(self, "_plan_head_fp32_ready", False): return - self.plan_head.float() - self.plan_head.register_forward_pre_hook(self._cast_plan_head_input_to_float32) + if self.config.plan_head_transformer: + for module in ( + self.plan_head.head, + self.plan_head.scale_layer, + self.plan_head.action_head, + self.plan_head.action_scale, + ): + module.float() + else: + self.plan_head.float() + self.plan_head.register_forward_pre_hook(self._cast_plan_head_input_to_float32) self._plan_head_fp32_ready = True @staticmethod def input_shapes(config: WorldModel.Config, batch_size: int = 1) -> dict[str, tuple[int, ...]]: frames, height, width = config.input_size pose_size = config.pose_size // 2 - return { + shapes: dict[str, tuple[int, ...]] = { "latents": (batch_size, frames, config.in_channels, height, width), "augments_pos_ref_augment": (batch_size, frames, pose_size), "ref_augment_from_augments_euler": (batch_size, frames, pose_size), "pose_mask": (batch_size, frames), "fidxs": (batch_size, frames), } + if config.plan_head_transformer or config.plan_head_action: + shapes["action_t"] = (batch_size, 2) + return shapes @staticmethod def input_dtypes(dtype: torch.dtype = torch.bfloat16) -> dict[str, torch.dtype]: @@ -271,6 +286,7 @@ def input_dtypes(dtype: torch.dtype = torch.bfloat16) -> dict[str, torch.dtype]: "ref_augment_from_augments_euler": dtype, "pose_mask": torch.int64, "fidxs": torch.int64, + "action_t": dtype, } @classmethod @@ -318,7 +334,7 @@ def get_model_io( def compile_for_inference(self) -> None: self._ensure_plan_head_float32() - for block in self.blocks: + for block in self._attention_blocks(): block.compile(mode="max-autotune-no-cudagraphs") def quantize_for_inference(self, weight_format: WeightFormat = "fp8_nvfp4") -> None: @@ -433,7 +449,7 @@ def _has_compatible_caches( return False return all( (cache := block.attn.kv_cache) is not None and cache.dtype == dtype and cache.k_cache.device == device - for block in self.blocks + for block in self._attention_blocks() ) def setup_caches( @@ -450,7 +466,7 @@ def setup_caches( self.max_seq_length = max_seq_length self.max_batch_size = max_batch_size self.cache_dtype = dtype - for block in self.blocks: + for block in self._attention_blocks(): block.attn.kv_cache = KVCache( max_batch_size, max_seq_length, @@ -461,7 +477,7 @@ def setup_caches( ) def cleanup_caches(self) -> None: - for block in self.blocks: + for block in self._attention_blocks(): block.attn.kv_cache = None self.max_batch_size = -1 self.max_seq_length = -1 @@ -484,6 +500,7 @@ def forward_n_steps( scheduler: RFScheduler, steps: int, return_trajectory: bool = False, + action_t: torch.Tensor | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor], torch.Tensor | None]: device = x.device batch, frames = x.shape[:2] @@ -508,6 +525,7 @@ def forward_n_steps( cache_pos=cache_pos, cache_seq_length=cache_seq_length, input_mask=input_mask, + action_t=action_t, ) velocity = model_output["sample"] if cfg > 0.0: @@ -536,8 +554,11 @@ def forward_n_steps( cache_pos=cache_pos, cache_seq_length=cache_seq_length, input_mask=input_mask, + action_t=action_t, ) model_output["plan"] = clean_output["plan"] + if "action" in clean_output: + model_output["action"] = clean_output["action"] return x, model_output, torch.stack(trajectory, dim=1) if trajectory is not None else None @@ -558,6 +579,7 @@ def generate( cfg: float = 0.0, return_trajectory: bool = False, kv_cache_dtype: KVCacheDType | None = None, + action_t: torch.Tensor | None = None, **scheduler_kwargs: Any, ) -> dict[str, torch.Tensor]: self._ensure_plan_head_float32() @@ -594,12 +616,15 @@ def generate( ref_augment_from_augments_euler, pose_mask, fidxs, + action_t=action_t, ) start = max(0, num_prefill_frames - 1) output_latents = self.unscale_latents(latents[:, start:]) outputs = {"latents": output_latents} if "plan" in model_output: outputs["plan"] = model_output["plan"] + if "action" in model_output: + outputs["action"] = model_output["action"] if return_trajectory: outputs["trajectory"] = output_latents.unsqueeze(1) return outputs @@ -643,6 +668,7 @@ def generate( num_prefill_frames=num_prefill_frames, cache_seq_length=cache_seq_length, prefill_mask=prefill_mask, + action_t=action_t, ) decode_frames = latents[:, num_prefill_frames:] @@ -664,6 +690,7 @@ def generate( scheduler=scheduler, steps=steps, return_trajectory=return_trajectory, + action_t=action_t, ) if not is_meta and not all(torch.isfinite(value).all() for value in model_output.values()): @@ -673,6 +700,8 @@ def generate( outputs = {"latents": self.unscale_latents(decode_frames)} if "plan" in model_output: outputs["plan"] = model_output["plan"] + if "action" in model_output: + outputs["action"] = model_output["action"] if trajectory is not None: outputs["trajectory"] = self.unscale_latents(trajectory) return outputs @@ -689,6 +718,7 @@ def _prefill( num_prefill_frames: int, cache_seq_length: int, prefill_mask: TensorOrMask | None, + action_t: torch.Tensor | None = None, ) -> dict[str, torch.Tensor]: if num_prefill_frames <= 0: return {} @@ -710,6 +740,7 @@ def _prefill( cache_pos=prefix_pos, cache_seq_length=cache_seq_length, input_mask=prefill_mask, + action_t=action_t, ) diff --git a/torchtitan/experiments/worldmodel/trainer.py b/torchtitan/experiments/worldmodel/trainer.py index be18733c194..c639a1b9c9b 100644 --- a/torchtitan/experiments/worldmodel/trainer.py +++ b/torchtitan/experiments/worldmodel/trainer.py @@ -140,14 +140,17 @@ def _prepare_worldmodel_batch( noisy_latents = scheduler.add_noise(latents, noise, fake_timesteps) targets = {**targets, "v": latents - noise, "mask": mask} - return { + model_inputs = { "x": noisy_latents, "t": timesteps, "augments_pos_ref_augment": augments, "ref_augment_from_augments_euler": eulers, "pose_mask": pose_mask.to(dtype=torch.int64), "fidx": fidxs, - }, targets + } + if "action_t" in input_dict: + model_inputs["action_t"] = input_dict["action_t"].to(device=device, dtype=dtype) + return model_inputs, targets class WorldModelValidator(BaseValidator): @@ -300,6 +303,7 @@ class Config(Trainer.Config): no_noise_prefill_frames_prob: float fake_timesteps_prob: float enable_rollout_report: bool = True + reports: list[Report] = field(default_factory=list) def __post_init__(self) -> None: Trainer.Config.__post_init__(self) @@ -332,8 +336,9 @@ def __init__(self, config: Config): config.training.steps, } ) - self.report_runner = ReportRunner( - [ + reports = list(config.reports) + if config.enable_rollout_report: + reports.append( Report( test_cls=AnalyseWorldmodel, test_config=AnalyseWorldmodelConfig(format=ReportFormat.HTML, save_tmp=False), @@ -343,12 +348,14 @@ def __init__(self, config: Config): steps=report_steps, wait_for_ckpt_keys=["model.fp8.torchpackage", "model.fp8_nvfp4.torchpackage"], ) - ], + ) + self.report_runner = ReportRunner( + reports, metrics_processor=self.metrics_processor, miniray={"codedir": config.codedir}, training_id=training_id, enabled=( - config.enable_rollout_report + (config.enable_rollout_report or bool(config.reports)) and config.metrics.enable_reporterv2 and config.checkpoint.enable and not config.checkpoint.load_only From 5abefee2a450002e574918f85ee7b3f566e287c0 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 25 Sep 2026 17:07:35 -0700 Subject: [PATCH 02/10] Limit worldmodel changes to transformer action-head training --- torchtitan/experiments/worldmodel/model.py | 22 ++----------------- .../worldmodel/model_for_inference.py | 2 +- torchtitan/experiments/worldmodel/trainer.py | 12 ++++------ 3 files changed, 7 insertions(+), 29 deletions(-) diff --git a/torchtitan/experiments/worldmodel/model.py b/torchtitan/experiments/worldmodel/model.py index df55c0732c8..82a2d68bce5 100644 --- a/torchtitan/experiments/worldmodel/model.py +++ b/torchtitan/experiments/worldmodel/model.py @@ -611,22 +611,12 @@ def __init__(self, config: "WorldModel.Config", linears: PlanHeadLinearsConfig): ) self.head = linears.head.build() self.scale_layer = ScaleLayer(PLAN_SIZE) - self.action_t_encoder = nn.Linear(2, config.plan_head.n_embd) if config.plan_head_action else None - if self.action_t_encoder is not None: - self.action_head = nn.Linear(config.plan_head.n_embd, 4) - self.action_scale = ScaleLayer(4) self.init_weights() - def forward(self, x: torch.Tensor, action_t: torch.Tensor | None = None) -> torch.Tensor | dict[str, torch.Tensor]: - if self.action_t_encoder is not None: - assert action_t is not None - x = x + self.action_t_encoder(action_t.to(x.dtype)) + def forward(self, x: torch.Tensor) -> torch.Tensor: for mlp in self.mlps: x = mlp(x) - plan = self.scale_layer(self.head(x)) - if self.action_t_encoder is not None: - return {"plan": plan, "action": self.action_scale(self.action_head(x))} - return plan + return self.scale_layer(self.head(x)) def init_weights(self) -> None: for module in self.mlps.modules(): @@ -636,11 +626,6 @@ def init_weights(self) -> None: if self.head.bias is not None: _init_plan_bias_(self.head.bias) self.scale_layer.init_weights() - if self.action_t_encoder is not None: - init_transformer_linear_weights(self.action_t_encoder) - _init_normal_(self.action_head.weight, std=PLAN_HEAD_INIT_STD) - _init_plan_bias_(self.action_head.bias) - self.action_scale.init_weights() class PlanTransformerBlock(nn.Module): @@ -792,7 +777,6 @@ class Config(BaseModel.Config): plan_head: TransformerConfig experimental_pose_only_xy: bool plan_head_transformer: bool = False - plan_head_action: bool = False x_embedder: PatchEmbedderLinearsConfig = field(init=False) augments_pos_ref_augment_embedder: ConditioningEmbedderLinearsConfig = field(init=False) ref_augment_from_augments_euler_embedder: ConditioningEmbedderLinearsConfig = field(init=False) @@ -1044,8 +1028,6 @@ def forward( if isinstance(self.plan_head, TransformerPlanHead): assert action_t is not None outputs.update(self.plan_head(x, action_t, input_mask, cache_pos, cache_seq_length)) - elif self.config.plan_head_action: - outputs.update(self.plan_head(x[:, -1, :], action_t)) else: outputs["plan"] = self.plan_head(x[:, -1, :]) if self.final_layer is not None: diff --git a/torchtitan/experiments/worldmodel/model_for_inference.py b/torchtitan/experiments/worldmodel/model_for_inference.py index 61c43c6d026..c12489f27ad 100644 --- a/torchtitan/experiments/worldmodel/model_for_inference.py +++ b/torchtitan/experiments/worldmodel/model_for_inference.py @@ -274,7 +274,7 @@ def input_shapes(config: WorldModel.Config, batch_size: int = 1) -> dict[str, tu "pose_mask": (batch_size, frames), "fidxs": (batch_size, frames), } - if config.plan_head_transformer or config.plan_head_action: + if config.plan_head_transformer: shapes["action_t"] = (batch_size, 2) return shapes diff --git a/torchtitan/experiments/worldmodel/trainer.py b/torchtitan/experiments/worldmodel/trainer.py index c639a1b9c9b..312eedd1836 100644 --- a/torchtitan/experiments/worldmodel/trainer.py +++ b/torchtitan/experiments/worldmodel/trainer.py @@ -303,7 +303,6 @@ class Config(Trainer.Config): no_noise_prefill_frames_prob: float fake_timesteps_prob: float enable_rollout_report: bool = True - reports: list[Report] = field(default_factory=list) def __post_init__(self) -> None: Trainer.Config.__post_init__(self) @@ -336,9 +335,8 @@ def __init__(self, config: Config): config.training.steps, } ) - reports = list(config.reports) - if config.enable_rollout_report: - reports.append( + self.report_runner = ReportRunner( + [ Report( test_cls=AnalyseWorldmodel, test_config=AnalyseWorldmodelConfig(format=ReportFormat.HTML, save_tmp=False), @@ -348,14 +346,12 @@ def __init__(self, config: Config): steps=report_steps, wait_for_ckpt_keys=["model.fp8.torchpackage", "model.fp8_nvfp4.torchpackage"], ) - ) - self.report_runner = ReportRunner( - reports, + ], metrics_processor=self.metrics_processor, miniray={"codedir": config.codedir}, training_id=training_id, enabled=( - (config.enable_rollout_report or bool(config.reports)) + config.enable_rollout_report and config.metrics.enable_reporterv2 and config.checkpoint.enable and not config.checkpoint.load_only From 7d63985a28a2e98f91d373c3561318e5d57a67d6 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 25 Sep 2026 17:15:18 -0700 Subject: [PATCH 03/10] Reuse worldmodel plan-head setup and supervised loss calculation --- torchtitan/experiments/worldmodel/loss.py | 32 ++++----- torchtitan/experiments/worldmodel/model.py | 81 +++++++++------------- 2 files changed, 45 insertions(+), 68 deletions(-) diff --git a/torchtitan/experiments/worldmodel/loss.py b/torchtitan/experiments/worldmodel/loss.py index 007190820c2..4420bb56b2e 100644 --- a/torchtitan/experiments/worldmodel/loss.py +++ b/torchtitan/experiments/worldmodel/loss.py @@ -61,27 +61,19 @@ def compute_worldmodel_losses( loss = diffusion_loss terms["diffusion_loss"] = diffusion_loss.detach() - if "plan" in outputs and "plan" in targets: - pred = outputs["plan"] - target = targets["plan"].to(device=pred.device, dtype=pred.dtype) - plan_loss_values, plan_err, plan_mask = laplacian_density_loss(target.float(), pred.float()) - plan_loss = plan_loss_values.flatten(1).mean(dim=1) - flat_mask = plan_mask.flatten(1).float() - plan_mse = (plan_err.square().flatten(1) * flat_mask).sum(dim=1) / flat_mask.sum(dim=1).clamp_min(1.0) - weighted_plan_loss = (plan_loss_weight if "sample" in outputs else 1.0) * plan_loss - loss = weighted_plan_loss if loss is None else loss + weighted_plan_loss - terms["plan_loss"] = plan_loss.detach() - terms["plan_mse"] = plan_mse.detach() - - if "action" in outputs: - pred = outputs["action"].float() - target = targets["action"].to(device=pred.device, dtype=torch.float32) - values, err, mask = laplacian_density_loss(target, pred) - action_loss = values.mean(dim=-1) - weighted = action_loss_weight * action_loss + for name, weight in (("plan", plan_loss_weight if "sample" in outputs else 1.0), ("action", action_loss_weight)): + if name not in outputs or (name == "plan" and name not in targets): + continue + pred = outputs[name] + target = targets[name].to(device=pred.device, dtype=pred.dtype if name == "plan" else torch.float32) + values, err, mask = laplacian_density_loss(target.float(), pred.float()) + head_loss = values.flatten(1).mean(dim=1) + flat_mask = mask.flatten(1).float() + head_mse = (err.square().flatten(1) * flat_mask).sum(dim=1) / flat_mask.sum(dim=1).clamp_min(1.0) + weighted = weight * head_loss loss = weighted if loss is None else loss + weighted - terms["action_loss"] = action_loss.detach() - terms["action_mse"] = (err.square() * mask).sum(-1).div(mask.sum(-1).clamp_min(1)).detach() + terms[f"{name}_loss"] = head_loss.detach() + terms[f"{name}_mse"] = head_mse.detach() if loss is None: raise RuntimeError("worldmodel produced no trainable outputs") diff --git a/torchtitan/experiments/worldmodel/model.py b/torchtitan/experiments/worldmodel/model.py index 82a2d68bce5..5c87c9211a1 100644 --- a/torchtitan/experiments/worldmodel/model.py +++ b/torchtitan/experiments/worldmodel/model.py @@ -607,13 +607,29 @@ class PlanHead(nn.Module): def __init__(self, config: "WorldModel.Config", linears: PlanHeadLinearsConfig): super().__init__() self.mlps = nn.ModuleList( - residual_ffn(config.plan_head, linears.blocks[i]) for i in range(config.plan_head.n_layer) + [] if config.plan_head_transformer else [residual_ffn(config.plan_head, block) for block in linears.blocks] ) + self.blocks = nn.ModuleList( + [PlanTransformerBlock(config.plan_head, block) for block in linears.blocks] + if config.plan_head_transformer + else [] + ) + if self.blocks: + self.action_t_encoder = nn.Linear(2, config.plan_head.n_embd) self.head = linears.head.build() self.scale_layer = ScaleLayer(PLAN_SIZE) + if self.blocks: + self.action_head = nn.Linear(config.plan_head.n_embd, 4) + self.action_scale = ScaleLayer(4) self.init_weights() - def forward(self, x: torch.Tensor) -> torch.Tensor: + def forward(self, x, action_t=None, input_mask=None, cache_pos=None, cache_seq_length=None): + if self.blocks: + x = x + self.action_t_encoder(action_t.to(x.dtype))[:, None] + for block in self.blocks: + x = block(x, input_mask, cache_pos, cache_seq_length) + x = x[:, -1].to(self.head.weight.dtype) + return {"plan": self.scale_layer(self.head(x)), "action": self.action_scale(self.action_head(x))} for mlp in self.mlps: x = mlp(x) return self.scale_layer(self.head(x)) @@ -622,10 +638,18 @@ def init_weights(self) -> None: for module in self.mlps.modules(): if isinstance(module, nn.Linear): init_transformer_linear_weights(module) + for block in self.blocks: + block.init_weights() + if self.blocks: + init_transformer_linear_weights(self.action_t_encoder) _init_normal_(self.head.weight, std=PLAN_HEAD_INIT_STD) if self.head.bias is not None: _init_plan_bias_(self.head.bias) self.scale_layer.init_weights() + if self.blocks: + _init_normal_(self.action_head.weight, std=PLAN_HEAD_INIT_STD) + _init_plan_bias_(self.action_head.bias) + self.action_scale.init_weights() class PlanTransformerBlock(nn.Module): @@ -648,37 +672,6 @@ def init_weights(self): init_transformer_linear_weights(module) -class TransformerPlanHead(nn.Module): - def __init__(self, config: "WorldModel.Config", linears: PlanHeadLinearsConfig): - super().__init__() - width = config.plan_head.n_embd - self.mlps = nn.ModuleList() - self.blocks = nn.ModuleList(PlanTransformerBlock(config.plan_head, block) for block in linears.blocks) - self.action_t_encoder = nn.Linear(2, width) - self.head = linears.head.build() - self.scale_layer = ScaleLayer(PLAN_SIZE) - self.action_head = nn.Linear(width, 4) - self.action_scale = ScaleLayer(4) - self.init_weights() - - def forward(self, x, action_t, input_mask=None, cache_pos=None, cache_seq_length=None): - x = x + self.action_t_encoder(action_t.to(x.dtype))[:, None] - for block in self.blocks: - x = block(x, input_mask, cache_pos, cache_seq_length) - x = x[:, -1].to(self.head.weight.dtype) - return {"plan": self.scale_layer(self.head(x)), "action": self.action_scale(self.action_head(x))} - - def init_weights(self): - for block in self.blocks: - block.init_weights() - init_transformer_linear_weights(self.action_t_encoder) - for head in (self.head, self.action_head): - _init_normal_(head.weight, std=PLAN_HEAD_INIT_STD) - _init_plan_bias_(head.bias) - self.scale_layer.init_weights() - self.action_scale.init_weights() - - class DiTBlock(nn.Module): def __init__(self, config: "WorldModel.Config", linears: DiTBlockLinearsConfig): super().__init__() @@ -905,8 +898,7 @@ def __init__(self, config: Config): self.fidx_embedder = DiscreteEmbedder(50, config.transformer.n_embd, config.fidx_embedder) self.blocks = nn.ModuleList(DiTBlock(config, config.blocks[i]) for i in range(config.transformer.n_layer)) self.final_layer = FinalLayer(config, config.final_layer) if config.final_layer is not None else None - head_cls = TransformerPlanHead if config.plan_head_transformer else PlanHead - self.plan_head = head_cls(config, config.plan_head_linears) if config.plan_head_linears is not None else None + self.plan_head = PlanHead(config, config.plan_head_linears) if config.plan_head_linears is not None else None self.register_buffer("pos_embed", torch.empty(1, config.num_patches, config.transformer.n_embd)) self.mask: TensorOrMask | None = None self.init_states(buffer_device=self.pos_embed.device) @@ -1025,7 +1017,7 @@ def forward( x = block(x, t6, input_pos_t, input_mask, cache_pos, cache_seq_length) outputs = {} if return_plan and self.plan_head is not None: - if isinstance(self.plan_head, TransformerPlanHead): + if self.config.plan_head_transformer: assert action_t is not None outputs.update(self.plan_head(x, action_t, input_mask, cache_pos, cache_seq_length)) else: @@ -1099,10 +1091,9 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module: ) if model.plan_head is not None: for name in ("mlps", "blocks"): - layers = getattr(model.plan_head, name, None) - if layers is not None: - for layer_id, block in layers.named_children(): - layers.register_module(layer_id, wrap(block, f"plan_head.{name}.{layer_id}")) + layers = getattr(model.plan_head, name) + for layer_id, block in layers.named_children(): + layers.register_module(layer_id, wrap(block, f"plan_head.{name}.{layer_id}")) logger.info(f"Applied {mode} activation checkpointing to the worldmodel") @@ -1125,11 +1116,8 @@ def _apply_compile(model: WorldModel, compile_config: CompileConfig) -> None: if model.final_layer is not None: model.final_layer.compile(backend=compile_config.backend, fullgraph=True) if model.plan_head is not None: - for block in model.plan_head.mlps: + for block in (*model.plan_head.mlps, *model.plan_head.blocks): block.compile(backend=compile_config.backend, fullgraph=True) - if isinstance(model.plan_head, TransformerPlanHead): - for block in model.plan_head.blocks: - block.compile(backend=compile_config.backend, fullgraph=True) logger.info("Compiling worldmodel components with torch.compile") @@ -1175,11 +1163,8 @@ def _apply_fsdp( for block in model.blocks: fully_shard(block, **fsdp_config, reshard_after_forward=reshard_after_forward) if model.plan_head is not None: - for block in model.plan_head.mlps: + for block in (*model.plan_head.mlps, *model.plan_head.blocks): fully_shard(block, **fsdp_config, reshard_after_forward=reshard_after_forward) - if isinstance(model.plan_head, TransformerPlanHead): - for block in model.plan_head.blocks: - fully_shard(block, **fsdp_config, reshard_after_forward=reshard_after_forward) fully_shard(model.plan_head, **fsdp_config, reshard_after_forward=reshard_after_forward) if model.final_layer is not None: fully_shard( From 54269df5030e9638cc2d42664b783dd91dc16464 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 25 Sep 2026 17:22:07 -0700 Subject: [PATCH 04/10] Simplify plan-head construction and inference outputs --- torchtitan/experiments/worldmodel/model.py | 12 +++----- .../worldmodel/model_for_inference.py | 28 ++++++------------- 2 files changed, 12 insertions(+), 28 deletions(-) diff --git a/torchtitan/experiments/worldmodel/model.py b/torchtitan/experiments/worldmodel/model.py index 5c87c9211a1..a0b75ca3251 100644 --- a/torchtitan/experiments/worldmodel/model.py +++ b/torchtitan/experiments/worldmodel/model.py @@ -606,14 +606,10 @@ def residual_ffn(config: TransformerConfig, linears: FFNLinearsConfig) -> Residu class PlanHead(nn.Module): def __init__(self, config: "WorldModel.Config", linears: PlanHeadLinearsConfig): super().__init__() - self.mlps = nn.ModuleList( - [] if config.plan_head_transformer else [residual_ffn(config.plan_head, block) for block in linears.blocks] - ) - self.blocks = nn.ModuleList( - [PlanTransformerBlock(config.plan_head, block) for block in linears.blocks] - if config.plan_head_transformer - else [] - ) + self.mlps, self.blocks = nn.ModuleList(), nn.ModuleList() + build_block = PlanTransformerBlock if config.plan_head_transformer else residual_ffn + layers = self.blocks if config.plan_head_transformer else self.mlps + layers.extend(build_block(config.plan_head, block) for block in linears.blocks) if self.blocks: self.action_t_encoder = nn.Linear(2, config.plan_head.n_embd) self.head = linears.head.build() diff --git a/torchtitan/experiments/worldmodel/model_for_inference.py b/torchtitan/experiments/worldmodel/model_for_inference.py index c12489f27ad..bdef38c341b 100644 --- a/torchtitan/experiments/worldmodel/model_for_inference.py +++ b/torchtitan/experiments/worldmodel/model_for_inference.py @@ -247,20 +247,16 @@ def _cast_plan_head_input_to_float32( return (args[0].float(), *args[1:]) def _ensure_plan_head_float32(self) -> None: - if self.plan_head is None or getattr(self, "_plan_head_fp32_ready", False): + head = self.plan_head + if head is None or getattr(self, "_plan_head_fp32_ready", False): return if self.config.plan_head_transformer: - for module in ( - self.plan_head.head, - self.plan_head.scale_layer, - self.plan_head.action_head, - self.plan_head.action_scale, - ): + for module in (head.head, head.scale_layer, head.action_head, head.action_scale): module.float() else: - self.plan_head.float() - self.plan_head.register_forward_pre_hook(self._cast_plan_head_input_to_float32) + head.float() + head.register_forward_pre_hook(self._cast_plan_head_input_to_float32) self._plan_head_fp32_ready = True @staticmethod @@ -556,9 +552,7 @@ def forward_n_steps( input_mask=input_mask, action_t=action_t, ) - model_output["plan"] = clean_output["plan"] - if "action" in clean_output: - model_output["action"] = clean_output["action"] + model_output.update({key: value for key, value in clean_output.items() if key != "sample"}) return x, model_output, torch.stack(trajectory, dim=1) if trajectory is not None else None @@ -621,10 +615,7 @@ def generate( start = max(0, num_prefill_frames - 1) output_latents = self.unscale_latents(latents[:, start:]) outputs = {"latents": output_latents} - if "plan" in model_output: - outputs["plan"] = model_output["plan"] - if "action" in model_output: - outputs["action"] = model_output["action"] + outputs.update({key: value for key, value in model_output.items() if key != "sample"}) if return_trajectory: outputs["trajectory"] = output_latents.unsqueeze(1) return outputs @@ -698,10 +689,7 @@ def generate( raise ValueError("model outputs contain inf/nan") outputs = {"latents": self.unscale_latents(decode_frames)} - if "plan" in model_output: - outputs["plan"] = model_output["plan"] - if "action" in model_output: - outputs["action"] = model_output["action"] + outputs.update({key: value for key, value in model_output.items() if key != "sample"}) if trajectory is not None: outputs["trajectory"] = self.unscale_latents(trajectory) return outputs From f7b6b792247e590bbf9765f3e9056c118a2f0c1a Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 25 Sep 2026 23:29:00 -0700 Subject: [PATCH 05/10] Diffuse frozen plan latents alongside image latents --- torchtitan/experiments/worldmodel/loss.py | 18 +++++--- torchtitan/experiments/worldmodel/model.py | 43 ++++++++++++++++--- .../worldmodel/model_for_inference.py | 40 ++++++++++++++--- .../experiments/worldmodel/tokenizer.py | 15 +++++++ torchtitan/experiments/worldmodel/trainer.py | 12 ++++++ 5 files changed, 111 insertions(+), 17 deletions(-) diff --git a/torchtitan/experiments/worldmodel/loss.py b/torchtitan/experiments/worldmodel/loss.py index 4420bb56b2e..ac6027e8150 100644 --- a/torchtitan/experiments/worldmodel/loss.py +++ b/torchtitan/experiments/worldmodel/loss.py @@ -52,14 +52,20 @@ def compute_worldmodel_losses( loss: torch.Tensor | None = None terms: dict[str, torch.Tensor] = {} - if "sample" in outputs: - pred = outputs["sample"] - target = targets["v"].to(device=pred.device, dtype=pred.dtype) - mask = targets["mask"].to(device=pred.device).flatten(1).float() + for output, target_key, mask_key, weight, term in ( + ("sample", "v", "mask", 1.0, "diffusion_loss"), + ("plan_v", "plan_v", "plan_mask", plan_loss_weight, "plan_diffusion_loss"), + ): + if output not in outputs: + continue + pred = outputs[output] if output == "sample" else outputs[output][:, -1] + target = targets[target_key].to(device=pred.device, dtype=pred.dtype) + mask = targets[mask_key].to(device=pred.device).flatten(1).float() mse = F.mse_loss(pred.float(), target.float(), reduction="none").flatten(1) diffusion_loss = (mse * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1.0) - loss = diffusion_loss - terms["diffusion_loss"] = diffusion_loss.detach() + weighted = weight * diffusion_loss + loss = weighted if loss is None else loss + weighted + terms[term] = diffusion_loss.detach() for name, weight in (("plan", plan_loss_weight if "sample" in outputs else 1.0), ("action", action_loss_weight)): if name not in outputs or (name == "plan" and name not in targets): diff --git a/torchtitan/experiments/worldmodel/model.py b/torchtitan/experiments/worldmodel/model.py index a0b75ca3251..3bac69bec36 100644 --- a/torchtitan/experiments/worldmodel/model.py +++ b/torchtitan/experiments/worldmodel/model.py @@ -766,6 +766,7 @@ class Config(BaseModel.Config): plan_head: TransformerConfig experimental_pose_only_xy: bool plan_head_transformer: bool = False + plan_latent_shape: tuple[int, int] = (0, 0) x_embedder: PatchEmbedderLinearsConfig = field(init=False) augments_pos_ref_augment_embedder: ConditioningEmbedderLinearsConfig = field(init=False) ref_augment_from_augments_euler_embedder: ConditioningEmbedderLinearsConfig = field(init=False) @@ -784,20 +785,24 @@ def num_spatial_patches(self) -> int: def num_temporal_patches(self) -> int: return self.input_size[0] // self.patch_size[0] + @property + def num_frame_tokens(self) -> int: + return self.num_spatial_patches + self.plan_latent_shape[0] + @property def num_patches(self) -> int: - return self.num_spatial_patches * self.num_temporal_patches + return self.num_frame_tokens * self.num_temporal_patches def __post_init__(self) -> None: self._sync_derived_fields() def _sync_derived_fields(self) -> None: self.transformer.block_size = self.num_patches - self.transformer.attention_mask_mini_block_size = self.num_spatial_patches + self.transformer.attention_mask_mini_block_size = self.num_frame_tokens self.plan_head.n_embd = self.transformer.n_embd if self.plan_head_transformer: self.plan_head.block_size = self.num_patches - self.plan_head.attention_mask_mini_block_size = self.num_spatial_patches + self.plan_head.attention_mask_mini_block_size = self.num_frame_tokens hidden = self.transformer.n_embd pose_half = self.pose_size // 2 current_blocks = getattr(self, "blocks", []) @@ -895,6 +900,14 @@ def __init__(self, config: Config): self.blocks = nn.ModuleList(DiTBlock(config, config.blocks[i]) for i in range(config.transformer.n_layer)) self.final_layer = FinalLayer(config, config.final_layer) if config.final_layer is not None else None self.plan_head = PlanHead(config, config.plan_head_linears) if config.plan_head_linears is not None else None + if config.plan_latent_shape[0]: + self.plan_embedder = nn.Linear(config.plan_latent_shape[1], config.transformer.n_embd) + self.plan_final_layer = FinalLayer( + config, + FinalLayerLinearsConfig( + linear=linear_config(config.transformer.n_embd, config.plan_latent_shape[1], bias=True) + ), + ) self.register_buffer("pos_embed", torch.empty(1, config.num_patches, config.transformer.n_embd)) self.mask: TensorOrMask | None = None self.init_states(buffer_device=self.pos_embed.device) @@ -917,11 +930,14 @@ def reset_parameters(self) -> None: self.config.input_size[2] // self.config.patch_size[2], ) spatial = torch.from_numpy(get_2d_sincos_pos_embed(self.pos_embed.shape[-1], spatial_grid)) + if self.config.plan_latent_shape[0]: + plan_pos = get_1d_sincos_pos_embed(self.pos_embed.shape[-1], self.config.plan_latent_shape[0]) + spatial = torch.cat((spatial, torch.from_numpy(plan_pos)), dim=0) spatial = spatial.to(dtype=self.pos_embed.dtype, device=self.pos_embed.device).unsqueeze(0) spatial = einops.repeat(spatial, "() n d -> () (t n) d", t=self.config.num_temporal_patches) temporal = torch.from_numpy(get_1d_sincos_pos_embed(self.pos_embed.shape[-1], self.config.num_temporal_patches)) temporal = temporal.to(dtype=self.pos_embed.dtype, device=self.pos_embed.device).unsqueeze(0) - temporal = einops.repeat(temporal, "() t d -> () (t n) d", n=self.config.num_spatial_patches) + temporal = einops.repeat(temporal, "() t d -> () (t n) d", n=self.config.num_frame_tokens) self.pos_embed[:] = spatial + temporal def init_states(self, *, buffer_device: torch.device | None = None) -> None: @@ -983,6 +999,7 @@ def forward( cache_seq_length: int | None = None, input_mask: TensorOrMask | None = None, action_t: torch.Tensor | None = None, + plan_latents: torch.Tensor | None = None, ) -> dict[str, torch.Tensor]: if input_pos is None: input_mask = self.mask @@ -994,12 +1011,20 @@ def forward( ) pos_embed = self.pos_embed[:, input_pos] if input_pos is not None else self.pos_embed input_pos_t = ( - input_pos[:: self.config.num_spatial_patches] // self.config.num_spatial_patches + input_pos[:: self.config.num_frame_tokens] // self.config.num_frame_tokens if input_pos is not None else None ) - x = self.x_embedder(x) + pos_embed + batch, frames = x.shape[:2] + x = self.x_embedder(x) + if self.config.plan_latent_shape[0]: + if plan_latents is None: + plan_latents = x.new_zeros(batch, frames, *self.config.plan_latent_shape) + x = torch.cat( + (x.unflatten(1, (frames, self.config.num_spatial_patches)), self.plan_embedder(plan_latents)), dim=2 + ).flatten(1, 2) + x = x + pos_embed augments_pos_ref_augment = self.position_scale(augments_pos_ref_augment) ref_augment_from_augments_euler = self.euler_scale(ref_augment_from_augments_euler) t6, t2 = self.t_embedder(t) @@ -1018,6 +1043,12 @@ def forward( outputs.update(self.plan_head(x, action_t, input_mask, cache_pos, cache_seq_length)) else: outputs["plan"] = self.plan_head(x[:, -1, :]) + if self.config.plan_latent_shape[0]: + x = x.unflatten(1, (frames, self.config.num_frame_tokens)) + outputs["plan_v"] = self.plan_final_layer( + x[:, :, self.config.num_spatial_patches :].flatten(1, 2), t2, input_pos_t + ).unflatten(1, (frames, self.config.plan_latent_shape[0])) + x = x[:, :, : self.config.num_spatial_patches].flatten(1, 2) if self.final_layer is not None: outputs["sample"] = self.unpatchify(self.final_layer(x, t2, input_pos_t)) return outputs diff --git a/torchtitan/experiments/worldmodel/model_for_inference.py b/torchtitan/experiments/worldmodel/model_for_inference.py index bdef38c341b..9abcd0fec29 100644 --- a/torchtitan/experiments/worldmodel/model_for_inference.py +++ b/torchtitan/experiments/worldmodel/model_for_inference.py @@ -64,6 +64,13 @@ def pack_cfg_inputs( ) +def pack_plan_latents(plan_latents: torch.Tensor | None, frames: int, cfg: float) -> torch.Tensor | None: + if plan_latents is None: + return None + plan_latents = F.pad(plan_latents[:, None], (0, 0, 0, 0, frames - 1, 0)) + return torch.cat((plan_latents, plan_latents), dim=1) if cfg > 0.0 else plan_latents + + def prefill_mask_predicate( mask_fn: Callable | None, prefix_tokens: int, @@ -272,6 +279,8 @@ def input_shapes(config: WorldModel.Config, batch_size: int = 1) -> dict[str, tu } if config.plan_head_transformer: shapes["action_t"] = (batch_size, 2) + if config.plan_latent_shape[0]: + shapes["plan_latents"] = (batch_size, *config.plan_latent_shape) return shapes @staticmethod @@ -283,6 +292,7 @@ def input_dtypes(dtype: torch.dtype = torch.bfloat16) -> dict[str, torch.dtype]: "pose_mask": torch.int64, "fidxs": torch.int64, "action_t": dtype, + "plan_latents": dtype, } @classmethod @@ -497,6 +507,7 @@ def forward_n_steps( steps: int, return_trajectory: bool = False, action_t: torch.Tensor | None = None, + plan_latents: torch.Tensor | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor], torch.Tensor | None]: device = x.device batch, frames = x.shape[:2] @@ -522,6 +533,7 @@ def forward_n_steps( cache_seq_length=cache_seq_length, input_mask=input_mask, action_t=action_t, + plan_latents=pack_plan_latents(plan_latents, frames, cfg), ) velocity = model_output["sample"] if cfg > 0.0: @@ -529,6 +541,12 @@ def forward_n_steps( velocity = unconditional + cfg * (conditional - unconditional) model_output["sample"] = velocity x = scheduler.step(velocity, step_idx, x).to(x.dtype) + if plan_latents is not None: + plan_velocity = model_output.pop("plan_v") + if cfg > 0.0: + unconditional, conditional = plan_velocity.chunk(2, dim=1) + plan_velocity = unconditional + cfg * (conditional - unconditional) + plan_latents = scheduler.step(plan_velocity[:, -1], step_idx, plan_latents).to(plan_latents.dtype) if trajectory is not None: trajectory.append(x.clone()) @@ -551,9 +569,12 @@ def forward_n_steps( cache_seq_length=cache_seq_length, input_mask=input_mask, action_t=action_t, + plan_latents=pack_plan_latents(plan_latents, frames, cfg), ) - model_output.update({key: value for key, value in clean_output.items() if key != "sample"}) + model_output.update({key: value for key, value in clean_output.items() if key not in ("sample", "plan_v")}) + if plan_latents is not None: + model_output["plan_latents"] = plan_latents return x, model_output, torch.stack(trajectory, dim=1) if trajectory is not None else None @torch.inference_mode() @@ -574,6 +595,7 @@ def generate( return_trajectory: bool = False, kv_cache_dtype: KVCacheDType | None = None, action_t: torch.Tensor | None = None, + plan_latents: torch.Tensor | None = None, **scheduler_kwargs: Any, ) -> dict[str, torch.Tensor]: self._ensure_plan_head_float32() @@ -592,6 +614,10 @@ def generate( latents = latents.to(dtype=dtype) device = latents.device is_meta = latents.is_meta + if self.config.plan_latent_shape[0]: + if plan_latents is None: + raise ValueError("plan diffusion requires plan_latents initialized with Gaussian noise") + plan_latents = plan_latents.to(device=device, dtype=dtype) if steps <= 0: self.cleanup_caches() @@ -611,11 +637,14 @@ def generate( pose_mask, fidxs, action_t=action_t, + plan_latents=pack_plan_latents(plan_latents, frames, 0.0), ) start = max(0, num_prefill_frames - 1) output_latents = self.unscale_latents(latents[:, start:]) outputs = {"latents": output_latents} - outputs.update({key: value for key, value in model_output.items() if key != "sample"}) + outputs.update({key: value for key, value in model_output.items() if key not in ("sample", "plan_v")}) + if plan_latents is not None: + outputs["plan_latents"] = plan_latents if return_trajectory: outputs["trajectory"] = output_latents.unsqueeze(1) return outputs @@ -628,8 +657,8 @@ def generate( scheduler = RFScheduler(steps=steps, inference_schedule=inference_schedule, **scheduler_kwargs).to( device=device ) - prefix_tokens = num_prefill_frames * self.config.num_spatial_patches - decode_tokens = (frames - num_prefill_frames) * self.config.num_spatial_patches + prefix_tokens = num_prefill_frames * self.config.num_frame_tokens + decode_tokens = (frames - num_prefill_frames) * self.config.num_frame_tokens packed_decode_tokens = decode_tokens * (2 if cfg > 0.0 else 1) cache_seq_length = prefix_tokens + packed_decode_tokens requested_kv_cache_dtype = ( @@ -682,6 +711,7 @@ def generate( steps=steps, return_trajectory=return_trajectory, action_t=action_t, + plan_latents=plan_latents, ) if not is_meta and not all(torch.isfinite(value).all() for value in model_output.values()): @@ -713,7 +743,7 @@ def _prefill( batch = latents.shape[0] device = latents.device - prefix_pos = torch.arange(0, num_prefill_frames * self.config.num_spatial_patches, device=device) + prefix_pos = torch.arange(0, num_prefill_frames * self.config.num_frame_tokens, device=device) timesteps = ( torch.ones((batch, num_prefill_frames), device=device, dtype=torch.float32) * scheduler.no_noise_timestep ) diff --git a/torchtitan/experiments/worldmodel/tokenizer.py b/torchtitan/experiments/worldmodel/tokenizer.py index f4693c68bd7..2bc78d91669 100644 --- a/torchtitan/experiments/worldmodel/tokenizer.py +++ b/torchtitan/experiments/worldmodel/tokenizer.py @@ -1,3 +1,9 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + from __future__ import annotations import io @@ -16,6 +22,7 @@ class WorldModelTokenizer(BaseTokenizer): class Config(BaseTokenizer.Config): compressor_model: str = "" compressor_in_channels: Literal[3, 6, "auto"] = "auto" + plan_encoder: str = "" def __init__( self, @@ -28,6 +35,14 @@ def __init__( self.config = config self._encoder: torch.nn.Module | None = None self._encoder_key: tuple[torch.device, torch.dtype] | None = None + self._plan_encoder: torch.nn.Module | None = None + + @torch.no_grad() + def encode_plan(self, plan: torch.Tensor) -> torch.Tensor: + if self._plan_encoder is None: + self._plan_encoder = torch.export.load(self.config.plan_encoder).module().requires_grad_(False) + self._plan_encoder.to(device=plan.device) + return self._plan_encoder(plan.float()) def encode( self, diff --git a/torchtitan/experiments/worldmodel/trainer.py b/torchtitan/experiments/worldmodel/trainer.py index 312eedd1836..7d5a198fb14 100644 --- a/torchtitan/experiments/worldmodel/trainer.py +++ b/torchtitan/experiments/worldmodel/trainer.py @@ -150,6 +150,18 @@ def _prepare_worldmodel_batch( } if "action_t" in input_dict: model_inputs["action_t"] = input_dict["action_t"].to(device=device, dtype=dtype) + if model.config.plan_latent_shape[0]: + plan = targets["plan"][:, :495] + valid = torch.isfinite(plan).all(dim=-1) + plan_latents = tokenizer.encode_plan(plan.masked_fill(~valid[:, None], 0)) + plan_latents = plan_latents.reshape(batch_size, *model.config.plan_latent_shape) + plan_noise = torch.randn_like(plan_latents) + noisy_plan = scheduler.add_noise(plan_latents, plan_noise, fake_timesteps[:, -1]) + # Historical plans would reveal the future, so only the target frame gets a plan latent. + model_inputs["plan_latents"] = latents.new_zeros(batch_size, num_frames, *model.config.plan_latent_shape) + model_inputs["plan_latents"][:, -1] = noisy_plan.to(dtype) + targets["plan_v"] = plan_latents - plan_noise + targets["plan_mask"] = valid[:, None, None].expand_as(plan_latents) return model_inputs, targets From 34858f3aa695e163ae5a63835ec39814460b68d1 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 25 Sep 2026 23:54:42 -0700 Subject: [PATCH 06/10] Run configured policy reports for worldmodel checkpoints --- torchtitan/experiments/worldmodel/trainer.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/torchtitan/experiments/worldmodel/trainer.py b/torchtitan/experiments/worldmodel/trainer.py index 7d5a198fb14..23107cda6c2 100644 --- a/torchtitan/experiments/worldmodel/trainer.py +++ b/torchtitan/experiments/worldmodel/trainer.py @@ -315,6 +315,7 @@ class Config(Trainer.Config): no_noise_prefill_frames_prob: float fake_timesteps_prob: float enable_rollout_report: bool = True + reports: list[Report] = field(default_factory=list) def __post_init__(self) -> None: Trainer.Config.__post_init__(self) @@ -347,8 +348,9 @@ def __init__(self, config: Config): config.training.steps, } ) - self.report_runner = ReportRunner( - [ + reports = list(config.reports) + if config.enable_rollout_report: + reports.append( Report( test_cls=AnalyseWorldmodel, test_config=AnalyseWorldmodelConfig(format=ReportFormat.HTML, save_tmp=False), @@ -358,12 +360,14 @@ def __init__(self, config: Config): steps=report_steps, wait_for_ckpt_keys=["model.fp8.torchpackage", "model.fp8_nvfp4.torchpackage"], ) - ], + ) + self.report_runner = ReportRunner( + reports, metrics_processor=self.metrics_processor, miniray={"codedir": config.codedir}, training_id=training_id, enabled=( - config.enable_rollout_report + (config.enable_rollout_report or bool(config.reports)) and config.metrics.enable_reporterv2 and config.checkpoint.enable and not config.checkpoint.load_only From 62d52d603ab9838ecbbc049fc40e88591016a8e4 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Sat, 26 Sep 2026 20:49:33 -0700 Subject: [PATCH 07/10] Support image codec resolution and bounded encoder batches --- torchtitan/experiments/worldmodel/dataset.py | 2 ++ torchtitan/experiments/worldmodel/tokenizer.py | 17 +++++++++++------ 2 files changed, 13 insertions(+), 6 deletions(-) diff --git a/torchtitan/experiments/worldmodel/dataset.py b/torchtitan/experiments/worldmodel/dataset.py index f1d04c2f481..2fd1616898f 100644 --- a/torchtitan/experiments/worldmodel/dataset.py +++ b/torchtitan/experiments/worldmodel/dataset.py @@ -34,6 +34,7 @@ class _DiffusionConfig: train_skip: int val_skip: int nan_engaged_plans: bool + image_size: tuple[int, int] = (128, 256) action_targets: bool = False def skip(self, val: bool) -> int: @@ -235,6 +236,7 @@ def _build_dataset(config: Config, *, val: bool, global_rank: int, global_world_ train_skip=config.train_skip, val_skip=config.val_skip, nan_engaged_plans=config.nan_engaged_plans, + image_size=config.image_size, action_targets=config.action_targets, ), val=val, diff --git a/torchtitan/experiments/worldmodel/tokenizer.py b/torchtitan/experiments/worldmodel/tokenizer.py index 2bc78d91669..f050067d71c 100644 --- a/torchtitan/experiments/worldmodel/tokenizer.py +++ b/torchtitan/experiments/worldmodel/tokenizer.py @@ -22,6 +22,8 @@ class WorldModelTokenizer(BaseTokenizer): class Config(BaseTokenizer.Config): compressor_model: str = "" compressor_in_channels: Literal[3, 6, "auto"] = "auto" + encode_batch_size: int | None = None + encode_dtype: Literal["float32", "bfloat16"] | None = None plan_encoder: str = "" def __init__( @@ -54,7 +56,8 @@ def encode( if "latents" in inputs: return inputs["latents"].to(device=device, dtype=dtype) - encoder = self._encoder_on(device=device, dtype=dtype) + encode_dtype = getattr(torch, self.config.encode_dtype) if self.config.encode_dtype else dtype + encoder = self._encoder_on(device=device, dtype=encode_dtype) imgs = inputs["imgs"] big_imgs = inputs["big_imgs"] batch, timesteps = imgs.shape[:2] @@ -75,18 +78,20 @@ def encode( nc=2, b=batch, t=timesteps, - ).to(device=device, dtype=dtype) + ).to(device=device, dtype=encode_dtype) x = x.div(255.0).mul(2).sub(1).clamp(-1, 1) - latents = encoder(x) - if isinstance(latents, tuple): - latents = latents[0] + encoded = [] + for chunk in x.split(self.config.encode_batch_size or x.shape[0]): + latents = encoder(chunk) + encoded.append(latents[0] if isinstance(latents, tuple) else latents) + latents = torch.cat(encoded) if len(encoded) > 1 else encoded[0] return einops.rearrange( latents, inverse_spec, nc=2, b=batch, t=timesteps, - ) + ).to(dtype=dtype) def decode(self, *args: Any, **kwargs: Any) -> str: return "" From b312aef8be16960849931457d36e4cc18025ead7 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Sat, 26 Sep 2026 20:53:19 -0700 Subject: [PATCH 08/10] Normalize pixels before applying codec precision override --- torchtitan/experiments/worldmodel/tokenizer.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/torchtitan/experiments/worldmodel/tokenizer.py b/torchtitan/experiments/worldmodel/tokenizer.py index f050067d71c..41539e47d0c 100644 --- a/torchtitan/experiments/worldmodel/tokenizer.py +++ b/torchtitan/experiments/worldmodel/tokenizer.py @@ -78,8 +78,8 @@ def encode( nc=2, b=batch, t=timesteps, - ).to(device=device, dtype=encode_dtype) - x = x.div(255.0).mul(2).sub(1).clamp(-1, 1) + ).to(device=device, dtype=torch.float32 if self.config.encode_dtype else dtype) + x = x.div(255.0).mul(2).sub(1).clamp(-1, 1).to(dtype=encode_dtype) encoded = [] for chunk in x.split(self.config.encode_batch_size or x.shape[0]): latents = encoder(chunk) From f913d0b36312f580a6eece59d67e284c454531f7 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Sun, 27 Sep 2026 21:22:26 -0700 Subject: [PATCH 09/10] Add a transformer branch for denoising plan latents --- torchtitan/experiments/worldmodel/model.py | 33 ++++++++++++++----- .../worldmodel/model_for_inference.py | 15 +++++---- 2 files changed, 33 insertions(+), 15 deletions(-) diff --git a/torchtitan/experiments/worldmodel/model.py b/torchtitan/experiments/worldmodel/model.py index 3bac69bec36..27aedcea357 100644 --- a/torchtitan/experiments/worldmodel/model.py +++ b/torchtitan/experiments/worldmodel/model.py @@ -767,6 +767,7 @@ class Config(BaseModel.Config): experimental_pose_only_xy: bool plan_head_transformer: bool = False plan_latent_shape: tuple[int, int] = (0, 0) + plan_transformer_layers: int = 0 x_embedder: PatchEmbedderLinearsConfig = field(init=False) augments_pos_ref_augment_embedder: ConditioningEmbedderLinearsConfig = field(init=False) ref_augment_from_augments_euler_embedder: ConditioningEmbedderLinearsConfig = field(init=False) @@ -774,6 +775,7 @@ class Config(BaseModel.Config): t_embedder: ConditioningEmbedderLinearsConfig = field(init=False) fidx_embedder: ConditioningEmbedderLinearsConfig = field(init=False) blocks: list[DiTBlockLinearsConfig] = field(init=False) + plan_blocks: list[DiTBlockLinearsConfig] = field(init=False) final_layer: FinalLayerLinearsConfig | None = field(init=False) plan_head_linears: PlanHeadLinearsConfig | None = field(init=False) @@ -806,6 +808,7 @@ def _sync_derived_fields(self) -> None: hidden = self.transformer.n_embd pose_half = self.pose_size // 2 current_blocks = getattr(self, "blocks", []) + current_plan_blocks = getattr(self, "plan_blocks", []) current_final = getattr(self, "final_layer", None) current_plan = getattr(self, "plan_head_linears", None) self.x_embedder = PatchEmbedderLinearsConfig( @@ -837,6 +840,13 @@ def _sync_derived_fields(self) -> None: ) for i in range(self.transformer.n_layer) ] + self.plan_blocks = [ + dit_block_linears_config( + self.transformer, + current_plan_blocks[i] if i < len(current_plan_blocks) else None, + ) + for i in range(self.plan_transformer_layers) + ] self.final_layer = ( FinalLayerLinearsConfig( linear=linear_config( @@ -880,7 +890,8 @@ def update_from_config(self, *, config: Any, **kwargs: Any) -> None: def get_nparams_and_flops(self, model: nn.Module, seq_len: int) -> tuple[int, int]: del seq_len nparams = sum(p.numel() for p in model.parameters()) - return nparams, 6 * nparams + attn_flops(self.transformer) // max(1, self.num_patches) + plan_attn_flops = 12 * self.plan_transformer_layers * self.transformer.n_embd * self.num_patches**2 + return nparams, 6 * nparams + (attn_flops(self.transformer) + plan_attn_flops) // max(1, self.num_patches) def __init__(self, config: Config): super().__init__() @@ -898,6 +909,7 @@ def __init__(self, config: Config): self.t_embedder = TimestepEmbedder(config.t_embedder, time_factor=config.time_factor) self.fidx_embedder = DiscreteEmbedder(50, config.transformer.n_embd, config.fidx_embedder) self.blocks = nn.ModuleList(DiTBlock(config, config.blocks[i]) for i in range(config.transformer.n_layer)) + self.plan_blocks = nn.ModuleList(DiTBlock(config, block) for block in config.plan_blocks) self.final_layer = FinalLayer(config, config.final_layer) if config.final_layer is not None else None self.plan_head = PlanHead(config, config.plan_head_linears) if config.plan_head_linears is not None else None if config.plan_latent_shape[0]: @@ -1044,9 +1056,13 @@ def forward( else: outputs["plan"] = self.plan_head(x[:, -1, :]) if self.config.plan_latent_shape[0]: + plan_x = x + for block in self.plan_blocks: + plan_x = block(plan_x, t6, input_pos_t, input_mask, cache_pos, cache_seq_length) + plan_x = plan_x.unflatten(1, (frames, self.config.num_frame_tokens)) x = x.unflatten(1, (frames, self.config.num_frame_tokens)) outputs["plan_v"] = self.plan_final_layer( - x[:, :, self.config.num_spatial_patches :].flatten(1, 2), t2, input_pos_t + plan_x[:, :, self.config.num_spatial_patches :].flatten(1, 2), t2, input_pos_t ).unflatten(1, (frames, self.config.plan_latent_shape[0])) x = x[:, :, : self.config.num_spatial_patches].flatten(1, 2) if self.final_layer is not None: @@ -1111,11 +1127,10 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module: mode = "full" if isinstance(ac_policy, FullAC) else "selective" - for layer_id, block in model.blocks.named_children(): - model.blocks.register_module( - layer_id, - wrap(block, f"blocks.{layer_id}"), - ) + for name in ("blocks", "plan_blocks"): + layers = getattr(model, name) + for layer_id, block in layers.named_children(): + layers.register_module(layer_id, wrap(block, f"{name}.{layer_id}")) if model.plan_head is not None: for name in ("mlps", "blocks"): layers = getattr(model.plan_head, name) @@ -1138,7 +1153,7 @@ def _apply_compile(model: WorldModel, compile_config: CompileConfig) -> None: model.fidx_embedder, ): module.compile(backend=compile_config.backend, fullgraph=True) - for block in model.blocks: + for block in (*model.blocks, *model.plan_blocks): block.compile(backend=compile_config.backend, fullgraph=True) if model.final_layer is not None: model.final_layer.compile(backend=compile_config.backend, fullgraph=True) @@ -1187,7 +1202,7 @@ def _apply_fsdp( ): fully_shard(module, **fsdp_config, reshard_after_forward=reshard_after_forward) - for block in model.blocks: + for block in (*model.blocks, *model.plan_blocks): fully_shard(block, **fsdp_config, reshard_after_forward=reshard_after_forward) if model.plan_head is not None: for block in (*model.plan_head.mlps, *model.plan_head.blocks): diff --git a/torchtitan/experiments/worldmodel/model_for_inference.py b/torchtitan/experiments/worldmodel/model_for_inference.py index 9abcd0fec29..29242666886 100644 --- a/torchtitan/experiments/worldmodel/model_for_inference.py +++ b/torchtitan/experiments/worldmodel/model_for_inference.py @@ -244,7 +244,7 @@ def __init__( self.default_kv_cache_dtype = default_kv_cache_dtype def _attention_blocks(self): - return (*self.blocks, *getattr(self.plan_head, "blocks", ())) + return (*self.blocks, *self.plan_blocks, *getattr(self.plan_head, "blocks", ())) @staticmethod def _cast_plan_head_input_to_float32( @@ -357,6 +357,7 @@ def quantize_for_inference(self, weight_format: WeightFormat = "fp8_nvfp4") -> N if weight_format == "fp8": quantize_(self.blocks, fp8_config) + quantize_(self.plan_blocks, fp8_config) return from torchao.prototype.mx_formats import NVFP4DynamicActivationNVFP4WeightConfig @@ -372,15 +373,17 @@ def is_mlp_linear(module: nn.Module, fqn: str) -> bool: fp8_config, filter_fn=is_attention_linear, ) + quantize_(self.plan_blocks, fp8_config, filter_fn=is_attention_linear) nvfp4_config = NVFP4DynamicActivationNVFP4WeightConfig( use_dynamic_per_tensor_scale=True, use_triton_kernel=False, ) - for fqn, module in self.blocks.named_modules(): - if is_mlp_linear(module, fqn): - module.to(device="cuda") - quantize_(module, nvfp4_config) - module.to(device="cpu") + for blocks in (self.blocks, self.plan_blocks): + for fqn, module in blocks.named_modules(): + if is_mlp_linear(module, fqn): + module.to(device="cuda") + quantize_(module, nvfp4_config) + module.to(device="cpu") @torch.no_grad() def get_inference_masks( From 5ef544daeb2b7ae6403865e8b0c9419e7063d597 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Mon, 28 Sep 2026 07:50:35 -0700 Subject: [PATCH 10/10] Allow the plan transformer branch to predict trajectories directly --- torchtitan/experiments/worldmodel/model.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/torchtitan/experiments/worldmodel/model.py b/torchtitan/experiments/worldmodel/model.py index 27aedcea357..e75d3953dee 100644 --- a/torchtitan/experiments/worldmodel/model.py +++ b/torchtitan/experiments/worldmodel/model.py @@ -1048,17 +1048,17 @@ def forward( t2 = t2 + pos2 + euler2 + pose_mask2 + fidx2 for block in self.blocks: x = block(x, t6, input_pos_t, input_mask, cache_pos, cache_seq_length) + plan_x = x + for block in self.plan_blocks: + plan_x = block(plan_x, t6, input_pos_t, input_mask, cache_pos, cache_seq_length) outputs = {} if return_plan and self.plan_head is not None: if self.config.plan_head_transformer: assert action_t is not None - outputs.update(self.plan_head(x, action_t, input_mask, cache_pos, cache_seq_length)) + outputs.update(self.plan_head(plan_x, action_t, input_mask, cache_pos, cache_seq_length)) else: - outputs["plan"] = self.plan_head(x[:, -1, :]) + outputs["plan"] = self.plan_head(plan_x[:, -1, :]) if self.config.plan_latent_shape[0]: - plan_x = x - for block in self.plan_blocks: - plan_x = block(plan_x, t6, input_pos_t, input_mask, cache_pos, cache_seq_length) plan_x = plan_x.unflatten(1, (frames, self.config.num_frame_tokens)) x = x.unflatten(1, (frames, self.config.num_frame_tokens)) outputs["plan_v"] = self.plan_final_layer(