diff --git a/torchtitan/experiments/worldmodel/dataset.py b/torchtitan/experiments/worldmodel/dataset.py index ab24339394e..2fd1616898f 100644 --- a/torchtitan/experiments/worldmodel/dataset.py +++ b/torchtitan/experiments/worldmodel/dataset.py @@ -34,6 +34,8 @@ 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: return self.val_skip if val else self.train_skip @@ -75,6 +77,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 +115,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 +236,8 @@ 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, 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..ac6027e8150 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,30 +47,39 @@ 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] = {} - 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() - - 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() + 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): + 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[f"{name}_loss"] = head_loss.detach() + terms[f"{name}_mse"] = head_mse.detach() if loss is None: raise RuntimeError("worldmodel produced no trainable outputs") @@ -78,6 +93,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 +106,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..e75d3953dee 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 @@ -606,14 +606,26 @@ 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( - residual_ffn(config.plan_head, linears.blocks[i]) for i in range(config.plan_head.n_layer) - ) + 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() 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 +634,38 @@ 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): + 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 DiTBlock(nn.Module): @@ -725,6 +765,9 @@ class Config(BaseModel.Config): transformer: TransformerConfig plan_head: TransformerConfig 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) @@ -732,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) @@ -743,20 +787,28 @@ 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_frame_tokens 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( @@ -788,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( @@ -831,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__() @@ -849,8 +909,17 @@ 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]: + 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) @@ -873,11 +942,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: @@ -938,6 +1010,8 @@ 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, + plan_latents: torch.Tensor | None = None, ) -> dict[str, torch.Tensor]: if input_pos is None: input_mask = self.mask @@ -949,12 +1023,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) @@ -966,9 +1048,23 @@ 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: - outputs["plan"] = self.plan_head(x[:, -1, :]) + if self.config.plan_head_transformer: + assert action_t is not None + outputs.update(self.plan_head(plan_x, action_t, input_mask, cache_pos, cache_seq_length)) + else: + outputs["plan"] = self.plan_head(plan_x[:, -1, :]) + if self.config.plan_latent_shape[0]: + 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( + 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: outputs["sample"] = self.unpatchify(self.final_layer(x, t2, input_pos_t)) return outputs @@ -1031,17 +1127,15 @@ 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 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) + 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") @@ -1059,12 +1153,12 @@ 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) 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) logger.info("Compiling worldmodel components with torch.compile") @@ -1108,10 +1202,10 @@ 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: + for block in (*model.plan_head.mlps, *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: diff --git a/torchtitan/experiments/worldmodel/model_for_inference.py b/torchtitan/experiments/worldmodel/model_for_inference.py index dae30f99797..29242666886 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, @@ -228,7 +235,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 +243,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, *self.plan_blocks, *getattr(self.plan_head, "blocks", ())) + @staticmethod def _cast_plan_head_input_to_float32( _module: nn.Module, @@ -244,24 +254,34 @@ 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 - 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 (head.head, head.scale_layer, head.action_head, head.action_scale): + module.float() + else: + head.float() + 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: + shapes["action_t"] = (batch_size, 2) + if config.plan_latent_shape[0]: + shapes["plan_latents"] = (batch_size, *config.plan_latent_shape) + return shapes @staticmethod def input_dtypes(dtype: torch.dtype = torch.bfloat16) -> dict[str, torch.dtype]: @@ -271,6 +291,8 @@ 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, + "plan_latents": dtype, } @classmethod @@ -318,7 +340,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: @@ -335,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 @@ -350,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( @@ -433,7 +458,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 +475,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 +486,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 +509,8 @@ def forward_n_steps( scheduler: RFScheduler, 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] @@ -508,6 +535,8 @@ def forward_n_steps( cache_pos=cache_pos, 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: @@ -515,6 +544,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()) @@ -536,9 +571,13 @@ def forward_n_steps( cache_pos=cache_pos, 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["plan"] = clean_output["plan"] + 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() @@ -558,6 +597,8 @@ def generate( cfg: float = 0.0, 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() @@ -576,6 +617,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() @@ -594,12 +639,15 @@ def generate( ref_augment_from_augments_euler, 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} - if "plan" in model_output: - outputs["plan"] = model_output["plan"] + 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 @@ -612,8 +660,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 = ( @@ -643,6 +691,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 +713,8 @@ def generate( scheduler=scheduler, 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()): @@ -671,8 +722,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"] + 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 @@ -689,13 +739,14 @@ 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 {} 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 ) @@ -710,6 +761,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/tokenizer.py b/torchtitan/experiments/worldmodel/tokenizer.py index f4693c68bd7..41539e47d0c 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,9 @@ 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__( self, @@ -28,6 +37,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, @@ -39,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] @@ -60,18 +78,20 @@ def encode( nc=2, b=batch, t=timesteps, - ).to(device=device, dtype=dtype) - x = x.div(255.0).mul(2).sub(1).clamp(-1, 1) - latents = encoder(x) - if isinstance(latents, tuple): - latents = latents[0] + ).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) + 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 "" diff --git a/torchtitan/experiments/worldmodel/trainer.py b/torchtitan/experiments/worldmodel/trainer.py index be18733c194..23107cda6c2 100644 --- a/torchtitan/experiments/worldmodel/trainer.py +++ b/torchtitan/experiments/worldmodel/trainer.py @@ -140,14 +140,29 @@ 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) + 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 class WorldModelValidator(BaseValidator): @@ -300,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) @@ -332,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), @@ -343,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