From f2919c05bf8c49345424c5058e5d1c98161265b5 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Mon, 22 Jun 2026 17:41:46 -0700 Subject: [PATCH 01/28] plan_vit: add the muP / scaling-study ViT as a torchtitan experiment Self-contained plan ViT for the prune-10m muP and scaling study, mirroring path's structure: model + config_registry (standard and muP flavors, n_embd 128..2048 at head_dim 64) + a thin trainer. Two cameras are channel-stacked into in_channels=24, matching the production worldmodel I/O. Registered as the "plan_vit" experiment so it launches like path: run.sh torchtitan/run_train.sh -e MODULE=plan_vit -e CONFIG=plan_vit_mup_w512 --- torchtitan/experiments/__init__.py | 1 + torchtitan/experiments/plan_vit/__init__.py | 4 + .../experiments/plan_vit/config_registry.py | 302 ++++++++++++++++++ torchtitan/experiments/plan_vit/model.py | 246 ++++++++++++++ torchtitan/experiments/plan_vit/trainer.py | 109 +++++++ 5 files changed, 662 insertions(+) create mode 100644 torchtitan/experiments/plan_vit/__init__.py create mode 100644 torchtitan/experiments/plan_vit/config_registry.py create mode 100644 torchtitan/experiments/plan_vit/model.py create mode 100644 torchtitan/experiments/plan_vit/trainer.py diff --git a/torchtitan/experiments/__init__.py b/torchtitan/experiments/__init__.py index 4b0e9ce7a5..8c5261ac73 100644 --- a/torchtitan/experiments/__init__.py +++ b/torchtitan/experiments/__init__.py @@ -13,6 +13,7 @@ "autoparallel.llama3", "autoparallel.local_map_deepseek_v3", "path", + "plan_vit", "torchft.llama3", "rl", # RL examples own a per-example config_registry under rl/examples/; diff --git a/torchtitan/experiments/plan_vit/__init__.py b/torchtitan/experiments/plan_vit/__init__.py new file mode 100644 index 0000000000..ba08e2ee97 --- /dev/null +++ b/torchtitan/experiments/plan_vit/__init__.py @@ -0,0 +1,4 @@ +from .config_registry import model_registry +from .model import parallelize_plan_vit, PlanViT + +__all__ = ["PlanViT", "model_registry", "parallelize_plan_vit"] diff --git a/torchtitan/experiments/plan_vit/config_registry.py b/torchtitan/experiments/plan_vit/config_registry.py new file mode 100644 index 0000000000..e5665ca34f --- /dev/null +++ b/torchtitan/experiments/plan_vit/config_registry.py @@ -0,0 +1,302 @@ +"""Config assembly + flavors for plan_vit, mirroring path/config_registry.py. + +Width flavors scale n_head at fixed head_dim=64 (the clean muP axis); base = w256. Two cameras are +channel-stacked into in_channels=24 (no VAE). The trainer-side config functions live below the model side. +""" +from __future__ import annotations + +import math +import os +from functools import partial + +import torch.nn as nn + +from torchtitan.components.checkpoint import CheckpointManager +from torchtitan.components.lr_scheduler import LRSchedulersContainer +from torchtitan.components.metrics import MetricsProcessor +from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig +from torchtitan.components.tokenizer import NoOpTokenizer +from torchtitan.config import DebugConfig, ParallelismConfig, TrainingConfig +from torchtitan.experiments.path.dataset import PathDataLoader +from torchtitan.experiments.path.loss import PathLoss +from torchtitan.models.common import Embedding, LayerNorm, Linear +from torchtitan.models.common.attention import ScaledDotProductAttention +from torchtitan.protocols.model_spec import ModelSpec +from xx.ml_tools.constants.model import SUPERCOMBO_FPS + +from .model import ( + parallelize_plan_vit, + PatchEmbed, + PlanHead, + PlanViT, + PlanViTAttention, + PlanViTBlock, + PlanViTMLP, +) + +from .trainer import PlanViTTrainer + +_LINEAR_INIT = { + "weight": partial(nn.init.normal_, mean=0.0, std=0.02), + "bias": nn.init.zeros_, +} +_NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} + +HEAD_DIM = 64 +N_LAYER = 8 +INPUT_SIZE = ( + 1, + 128, + 256, +) # current frame; spatial ViT (temporal history is a later variant) +PATCH_SIZE = (1, 16, 8) +IN_CHANNELS = 24 # two cameras (IMG + BIG_IMG), 12 YUV channels each, channel-stacked +PLAN_SIZE = 15 * 33 * 2 # 990, laplacian mu+log-sigma +BASE_WIDTH = 256 +PLAN_VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} + + +def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Config: + return Linear.Config( + in_features=in_f, + out_features=out_f, + bias=bias, + param_init={ + "weight": partial(nn.init.normal_, mean=0.0, std=std), + "bias": nn.init.zeros_, + }, + ) + + +def _hidden_std(fan_in: int, *, mup: bool) -> float: + # muP shrinks hidden/output init to 1/sqrt(fan_in) so pre-activations stay O(1) as width grows; + # standard param holds the base-width variance 1/sqrt(BASE_WIDTH), so it fans out with width. + return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 + + +def _ln(dim: int) -> LayerNorm.Config: + return LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT) + + +def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: + return multiple_of * math.ceil(int(dim * mult) / multiple_of) + + +def _attention(dim: int, n_head: int, *, mup: bool, qk_norm: bool = True) -> PlanViTAttention.Config: + head_dim = dim // n_head + return PlanViTAttention.Config( + norm=_ln(dim), + q_norm=_ln(head_dim) if qk_norm else None, + k_norm=_ln(head_dim) if qk_norm else None, + c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), + c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * N_LAYER)), + inner_attention=ScaledDotProductAttention.Config(), + n_head=n_head, + head_dim=head_dim, + dropout=0.0, + ) + + +def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PlanViTMLP.Config: + hidden = _hidden(dim, mult) + return PlanViTMLP.Config( + norm=_ln(dim), + c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), + c_proj=_lin(hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER)), + act="gelu_tanh", + dropout=0.0, + ) + + +def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: + n_embd = PLAN_VIT_WIDTHS[flavor] + n_head = n_embd // HEAD_DIM + pt, ph, pw = PATCH_SIZE + patch_dim = pt * IN_CHANNELS * ph * pw + t, h, w = INPUT_SIZE + num_patches = (t // pt) * (h // ph) * (w // pw) + return PlanViT.Config( + input_size=INPUT_SIZE, + patch_size=PATCH_SIZE, + in_channels=IN_CHANNELS, + n_embd=n_embd, + output_mult=(BASE_WIDTH / n_embd) if mup else 1.0, # muP readout multiplier 1/m + patch_embed=PatchEmbed.Config( + proj=_lin(patch_dim, n_embd, std=patch_dim**-0.5), # input embed: width-independent + patch_size=PATCH_SIZE, + ), + pos_embedding=Embedding.Config( + num_embeddings=num_patches, embedding_dim=n_embd, param_init=_LINEAR_INIT + ), + blocks=[ + PlanViTBlock.Config( + attention=_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), mlp=_mlp(n_embd, mup=mup) + ) + for _ in range(N_LAYER) + ], + norm=_ln(n_embd), + plan_head=PlanHead.Config( + norm=_ln(n_embd), + head=_lin(n_embd, PLAN_SIZE, std=_hidden_std(n_embd, mup=mup)), + ), + ) + + +def model_registry(flavor: str, *, mup: bool) -> ModelSpec: + return ModelSpec( + name="plan_vit", + flavor=flavor, + model=_model_config(flavor, mup=mup), + parallelize_fn=parallelize_plan_vit, + pipelining_fn=None, + post_optimizer_build_fn=None, + state_dict_adapter=None, + ) + + +STEPS = 512 # per-run step budget; override with training.steps=N on the CLI +# learning rate is the muTransfer sweep axis: one run per (flavor, lr); set with `-e PLAN_VIT_LR=...` +SWEEP_LR = float(os.getenv("PLAN_VIT_LR", "3e-4")) +# hidden + output matrix weights get muP lr eta/m; input embed, norms, biases get eta +MUP_PATTERN = ( + r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" + r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj|plan_head\.head)\.weight$" +) + + +def _si_int(value: str | int) -> int: + suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} + value = str(value).strip().lower() + return ( + int(float(value[:-1]) * suffixes[value[-1]]) + if value[-1] in suffixes + else int(value) + ) + + +def _dataloader_config(*, split: str) -> PathDataLoader.Config: + from xx.common.basedir import XX_BASEDIR + from xx.datasets.constants import BASE_DIR_GT_10M + from xx.training.path.config import DatasetConfig as XXPathDatasetConfig + + base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) + return PathDataLoader.Config( + # prune-10M study data: a seeded random 10k sample of the 10M store (training_2026_02) + dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), + split=split, + shuffle_size=_si_int(base.shuffle_size), + min_mixing=base.min_mixing, + num_writers=base.num_writers, + num_readers=base.num_readers, + fps=base.fps, + pipeline_dir=BASE_DIR_GT_10M, # the 10M store, not the 2.5M big-train list + plan_only=base.plan_only, + limit=base.limit, + n_frames=base.n_frames, + rgb=base.rgb, + unvision=base.unvision, + ) + + +def _optimizer_config( + flavor: str, *, mup: bool, lr: float, wd: float +) -> OptimizersContainer.Config: + m = PLAN_VIT_WIDTHS[flavor] / BASE_WIDTH + common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} + if mup: + groups = [ + ParamGroupConfig( + pattern=MUP_PATTERN, + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr / m}, + ), + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr}, + ), + ] + else: + groups = [ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr}, + ) + ] + return OptimizersContainer.Config( + implementation="fused_opt_states_bf16", param_groups=groups + ) + + +def _plan_vit( + flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 +) -> PlanViTTrainer.Config: + return PlanViTTrainer.Config( + loss=PathLoss.Config(), + model_spec=model_registry(flavor, mup=mup), + tokenizer=NoOpTokenizer.Config(), + dataloader=_dataloader_config(split="train"), + optimizer=_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), + lr_scheduler=LRSchedulersContainer.Config( + warmup_steps=round(STEPS * 0.1), + total_steps=STEPS, + decay_ratio=0.8, + decay_type="cosine", + min_lr_factor=0.0, + ), + training=TrainingConfig( + local_batch_size=16, + global_batch_size=-1, + seq_len=1, + steps=STEPS, + max_norm=1.0, + dtype="float32", + mixed_precision_param="bfloat16", + mixed_precision_reduce="float32", + ), + parallelism=ParallelismConfig( + data_parallel_replicate_degree=1, data_parallel_shard_degree=8 + ), + checkpoint=CheckpointManager.Config(enable=False), + metrics=MetricsProcessor.Config( + log_freq=10, enable_reporterv2=True, save_freq=STEPS + ), + debug=DebugConfig(seed=0), + ) + + +def plan_vit_standard_w256() -> PlanViTTrainer.Config: + return _plan_vit("w256", mup=False) + + +def plan_vit_standard_w512() -> PlanViTTrainer.Config: + return _plan_vit("w512", mup=False) + + +def plan_vit_standard_w1024() -> PlanViTTrainer.Config: + return _plan_vit("w1024", mup=False) + + +def plan_vit_standard_w2048() -> PlanViTTrainer.Config: + return _plan_vit("w2048", mup=False) + + +def plan_vit_mup_w256() -> PlanViTTrainer.Config: + return _plan_vit("w256", mup=True) + + +def plan_vit_mup_w512() -> PlanViTTrainer.Config: + return _plan_vit("w512", mup=True) + + +def plan_vit_mup_w1024() -> PlanViTTrainer.Config: + return _plan_vit("w1024", mup=True) + + +def plan_vit_mup_w2048() -> PlanViTTrainer.Config: + return _plan_vit("w2048", mup=True) + + +def plan_vit() -> PlanViTTrainer.Config: + return plan_vit_mup_w256() diff --git a/torchtitan/experiments/plan_vit/model.py b/torchtitan/experiments/plan_vit/model.py new file mode 100644 index 0000000000..a301f3a756 --- /dev/null +++ b/torchtitan/experiments/plan_vit/model.py @@ -0,0 +1,246 @@ +"""Plan ViT: raw camera frames -> patches -> transformer -> plan. NO VAE. + +A self-contained planning model for the muP + scaling study, built from torchtitan.models.common +blocks the same way path/model.py is. Scales cleanly by width (n_embd / n_head) for muTransfer. +""" +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn as nn +from einops import rearrange +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy + +from torchtitan.config import ( + CompileConfig, + ParallelismConfig, + TORCH_DTYPE_MAP, + TrainingConfig, +) +from torchtitan.distributed import ParallelDims +from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig +from torchtitan.models.common import Embedding, LayerNorm, Linear, RMSNorm +from torchtitan.models.common.attention import ScaledDotProductAttention +from torchtitan.protocols.model import BaseModel +from torchtitan.protocols.module import Module, ModuleList +from torchtitan.tools.logging import logger +from xx.ml_tools.constants.model import ModelInputs + + +class PlanViTMLP(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + norm: LayerNorm.Config | RMSNorm.Config + c_fc: Linear.Config + c_proj: Linear.Config + act: str + dropout: float + + def __init__(self, config: Config): + super().__init__() + self.norm = config.norm.build() + self.c_fc = config.c_fc.build() + self.act = ( + nn.GELU(approximate="tanh") if config.act == "gelu_tanh" else nn.GELU() + ) + self.c_proj = config.c_proj.build() + self.dropout = nn.Dropout(config.dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.dropout(self.c_proj(self.act(self.c_fc(self.norm(x))))) + + +class PlanViTAttention(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + norm: LayerNorm.Config | RMSNorm.Config + q_norm: LayerNorm.Config | RMSNorm.Config | None + k_norm: LayerNorm.Config | RMSNorm.Config | None + c_attn: Linear.Config + c_proj: Linear.Config + inner_attention: ScaledDotProductAttention.Config + n_head: int + head_dim: int + dropout: float + + def __init__(self, config: Config): + super().__init__() + self.n_head = config.n_head + self.head_dim = config.head_dim + self.norm = config.norm.build() + self.q_norm = ( + config.q_norm.build() if config.q_norm is not None else nn.Identity() + ) + self.k_norm = ( + config.k_norm.build() if config.k_norm is not None else nn.Identity() + ) + self.c_attn = config.c_attn.build() + self.c_proj = config.c_proj.build() + self.inner_attention = config.inner_attention.build() + self.dropout = nn.Dropout(config.dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + b, t, _ = x.shape + qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) + q, k, v = qkv.unbind(2) + q, k = self.q_norm(q), self.k_norm(k) + x = self.inner_attention( + q, k, v, is_causal=False + ) # ViT: bidirectional over patches + return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) + + +class PlanViTBlock(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + attention: PlanViTAttention.Config + mlp: PlanViTMLP.Config + + def __init__(self, config: Config): + super().__init__() + self.attention = config.attention.build() + self.mlp = config.mlp.build() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x + self.attention(x) + return x + self.mlp(x) + + +class PatchEmbed(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + proj: Linear.Config + patch_size: tuple[int, int, int] # (pt, ph, pw) + + def __init__(self, config: Config): + super().__init__() + self.patch_size = config.patch_size + self.proj = config.proj.build() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + # x: (B, T, C, H, W) raw frames -> (B, num_patches, patch_dim) -> (B, num_patches, n_embd) + pt, ph, pw = self.patch_size + x = rearrange( + x, "b (t pt) c (h ph) (w pw) -> b (t h w) (pt c ph pw)", pt=pt, ph=ph, pw=pw + ) + return self.proj(x.to(self.proj.weight.dtype)) # match the bf16 (mp) weights, like path's vision + + +class PlanHead(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + norm: LayerNorm.Config | RMSNorm.Config + head: Linear.Config + + def __init__(self, config: Config): + super().__init__() + self.norm = config.norm.build() + self.head = config.head.build() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.head(self.norm(x)) + + +class PlanViT(BaseModel): + @dataclass(kw_only=True, slots=True) + class Config(BaseModel.Config): + input_size: tuple[int, int, int] # (n_frames, H, W) + patch_size: tuple[int, int, int] + in_channels: int + n_embd: int + output_mult: float # muP readout multiplier 1/m (m = n_embd / base); 1.0 for standard param + patch_embed: PatchEmbed.Config + pos_embedding: Embedding.Config + blocks: list[PlanViTBlock.Config] + norm: LayerNorm.Config | RMSNorm.Config + plan_head: PlanHead.Config + + @property + def num_patches(self) -> int: + t, h, w = self.input_size + pt, ph, pw = self.patch_size + return (t // pt) * (h // ph) * (w // pw) + + def update_from_config(self, *, config, **kwargs) -> None: + parallelism = config.parallelism + for name, degree in { + "tensor parallel": parallelism.tensor_parallel_degree, + "context parallel": parallelism.context_parallel_degree, + "pipeline parallel": parallelism.pipeline_parallel_degree, + "expert parallel": parallelism.expert_parallel_degree, + }.items(): + if degree > 1: + raise ValueError(f"plan_vit does not support {name}") + + def get_nparams_and_flops(self, model: Module, seq_len: int) -> tuple[int, int]: + nparams = sum(p.numel() for p in model.parameters()) + return nparams, 6 * nparams + + def __init__(self, config: Config): + super().__init__() + self.config = config + self.patch_embed = config.patch_embed.build() + self.pos_embedding = config.pos_embedding.build() + self.blocks = ModuleList([block.build() for block in config.blocks]) + self.norm = config.norm.build() + self.plan_head = config.plan_head.build() + + def verify_module_protocol(self) -> None: + pass # nn.Dropout/GELU/Identity are plain nn.Module, like path + + def _frames(self, inputs: dict[str, torch.Tensor] | torch.Tensor) -> torch.Tensor: + # production input: two cameras IMG, BIG_IMG, each (B, T, 12, H, W) YUV. Take the current frame of each, + # channel-stack -> (B, 1, 24, H, W). NO VAE. A plain tensor (testing) is passed through unchanged. + if isinstance(inputs, torch.Tensor): + return inputs + img, big = inputs[ModelInputs.IMG], inputs[ModelInputs.BIG_IMG] + frame = torch.cat([img[:, -1], big[:, -1]], dim=1).unsqueeze(1) + return (frame.float() - 127.5) / 63.75 # uint8 YUV -> normalized float (mean 255/2, std 255/4 like path) + + def forward( + self, inputs: dict[str, torch.Tensor] | torch.Tensor + ) -> dict[str, torch.Tensor]: + x = self.patch_embed(self._frames(inputs)) + pos = self.pos_embedding(torch.arange(x.shape[1], device=x.device)) + x = x + rearrange(pos, "t c -> () t c") + for block in self.blocks: + x = block(x) + x = self.norm(x) + # global-pool the patches -> plan; the muP readout multiplier keeps the output width-stable + return {"plan": self.plan_head(x.mean(dim=1)) * self.config.output_mult} + + +def parallelize_plan_vit( + model: PlanViT, + *, + parallel_dims: ParallelDims, + training: TrainingConfig, + parallelism: ParallelismConfig, + compile_config: CompileConfig, + ac_config: ActivationCheckpointingConfig, + dump_folder: str, +) -> PlanViT: + if ( + parallel_dims.tp_enabled + or parallel_dims.cp_enabled + or parallel_dims.pp_enabled + or parallel_dims.ep_enabled + ): + raise ValueError("plan_vit supports data parallelism only") + names = ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] + dp_mesh: DeviceMesh = parallel_dims.get_mesh(names) + mp_policy = MixedPrecisionPolicy( + param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param], + reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], + cast_forward_inputs=True, + ) + fsdp_config = {"mesh": dp_mesh, "mp_policy": mp_policy} + for idx, block in enumerate(model.blocks): + fully_shard( + block, **fsdp_config, reshard_after_forward=(idx < len(model.blocks) - 1) + ) + fully_shard(model, **fsdp_config) + logger.info("Applied FSDP to plan_vit") + return model diff --git a/torchtitan/experiments/plan_vit/trainer.py b/torchtitan/experiments/plan_vit/trainer.py new file mode 100644 index 0000000000..b6df10a2b8 --- /dev/null +++ b/torchtitan/experiments/plan_vit/trainer.py @@ -0,0 +1,109 @@ +"""Thin trainer for plan_vit: model(inputs_dict) -> pred dict, PathLoss(pred, targets), backward. + +Mirrors PathTrainer's generic core without the path-specific ONNX export / driving validator / reports a +scaling study doesn't need. loss + dataloader come from the base Trainer.Config (set in config_registry). +""" +from __future__ import annotations + +import time +from collections.abc import Iterable, Iterator +from dataclasses import dataclass + +import torch + +from torchtitan.components.dataloader import DataloaderExhaustedError +from torchtitan.distributed import utils as dist_utils +from torchtitan.observability import structured_logger as sl +from torchtitan.trainer import Trainer + + +class PlanViTTrainer(Trainer): + @dataclass(kw_only=True, slots=True) + class Config(Trainer.Config): + pass + + def __init__(self, config: "PlanViTTrainer.Config"): + self.ntokens_seen = 0 + self._metrics: dict[str, torch.Tensor] = {} + super().__init__(config) + self.loss_fn.to(self.device) + + def batch_generator( + self, + data_iterable: Iterable[ + tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] + ], + ) -> Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]]: + data_iterator = iter(data_iterable) + while True: + t0 = time.perf_counter() + try: + input_dict, targets = next(data_iterator) + except StopIteration as ex: + raise DataloaderExhaustedError() from ex + self.metrics_processor.ntokens_since_last_log += next( + iter(input_dict.values()) + ).shape[0] + self.metrics_processor.data_loading_times.append(time.perf_counter() - t0) + yield input_dict, targets + + @sl.log_trace_span("fwd_bwd") + def forward_backward_step( + self, *, input_dict: dict[str, torch.Tensor], labels: dict[str, torch.Tensor] + ) -> torch.Tensor: + assert len(self.model_parts) == 1 + with self.train_context(): + pred = self.model_parts[0](input_dict) + # plan_vit is single-frame: it predicts the current (last) frame's plan, so supervise + # against the last temporal position of the dense target (path trains all positions). + labels = {**labels, "plan": labels["plan"][:, -1]} + loss_vec, metrics = self.loss_fn(pred, labels) + loss = loss_vec.mean() + self._metrics = metrics + loss.backward() + return loss + + def train_step( + self, + data_iterator: Iterator[ + tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] + ], + ) -> None: + self.optimizers.zero_grad() + lr = self.lr_schedulers.schedulers[0].get_last_lr()[0] + + input_dict, targets = next(data_iterator) + input_dict = {k: v.to(self.device) for k, v in input_dict.items()} + targets = {k: v.to(self.device) for k, v in targets.items()} + self.ntokens_seen += next(iter(input_dict.values())).shape[0] + loss = self.forward_backward_step(input_dict=input_dict, labels=targets) + + grad_norm = dist_utils.clip_grad_norm_( + [p for m in self.model_parts for p in m.parameters()], + self.config.training.max_norm, + foreach=True, + pp_mesh=self.parallel_dims.get_optional_mesh("pp"), + ep_enabled=self.parallel_dims.ep_enabled, + ) + self.checkpointer.maybe_wait_for_staging() + self.optimizers.step() + self.lr_schedulers.step() + + if not self.metrics_processor.should_log(self.step): + return + + local_loss = loss.detach() + if self.parallel_dims.dp_cp_enabled: + loss_mesh = self.parallel_dims.get_optional_mesh("loss") + global_avg_loss = dist_utils.dist_mean(local_loss, loss_mesh) + global_max_loss = dist_utils.dist_max(local_loss, loss_mesh) + else: + global_avg_loss = global_max_loss = float(local_loss.item()) + + self.metrics_processor.log( + self.step, + global_avg_loss, + global_max_loss, + float(grad_norm.item()), + extra_metrics={"metrics/lr/": lr}, + ) From 2e5d27e04a2e7cf33aad2b6256a0b3feede820d6 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Mon, 22 Jun 2026 19:42:51 -0700 Subject: [PATCH 02/28] plan_vit: add license headers and ufmt formatting --- torchtitan/experiments/plan_vit/__init__.py | 6 +++++ .../experiments/plan_vit/config_registry.py | 26 ++++++++++++++----- torchtitan/experiments/plan_vit/model.py | 17 +++++++++--- torchtitan/experiments/plan_vit/trainer.py | 7 +++++ 4 files changed, 46 insertions(+), 10 deletions(-) diff --git a/torchtitan/experiments/plan_vit/__init__.py b/torchtitan/experiments/plan_vit/__init__.py index ba08e2ee97..be18ef68a2 100644 --- a/torchtitan/experiments/plan_vit/__init__.py +++ b/torchtitan/experiments/plan_vit/__init__.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 .config_registry import model_registry from .model import parallelize_plan_vit, PlanViT diff --git a/torchtitan/experiments/plan_vit/config_registry.py b/torchtitan/experiments/plan_vit/config_registry.py index e5665ca34f..498a280868 100644 --- a/torchtitan/experiments/plan_vit/config_registry.py +++ b/torchtitan/experiments/plan_vit/config_registry.py @@ -1,13 +1,21 @@ +# 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. + """Config assembly + flavors for plan_vit, mirroring path/config_registry.py. Width flavors scale n_head at fixed head_dim=64 (the clean muP axis); base = w256. Two cameras are channel-stacked into in_channels=24 (no VAE). The trainer-side config functions live below the model side. """ + from __future__ import annotations import math import os from functools import partial +from xx.ml_tools.constants.model import SUPERCOMBO_FPS import torch.nn as nn @@ -22,8 +30,6 @@ from torchtitan.models.common import Embedding, LayerNorm, Linear from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model_spec import ModelSpec -from xx.ml_tools.constants.model import SUPERCOMBO_FPS - from .model import ( parallelize_plan_vit, PatchEmbed, @@ -33,7 +39,6 @@ PlanViTBlock, PlanViTMLP, ) - from .trainer import PlanViTTrainer _LINEAR_INIT = { @@ -82,7 +87,9 @@ def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: return multiple_of * math.ceil(int(dim * mult) / multiple_of) -def _attention(dim: int, n_head: int, *, mup: bool, qk_norm: bool = True) -> PlanViTAttention.Config: +def _attention( + dim: int, n_head: int, *, mup: bool, qk_norm: bool = True +) -> PlanViTAttention.Config: head_dim = dim // n_head return PlanViTAttention.Config( norm=_ln(dim), @@ -102,7 +109,9 @@ def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PlanViTMLP.Config: return PlanViTMLP.Config( norm=_ln(dim), c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), - c_proj=_lin(hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER)), + c_proj=_lin( + hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER) + ), act="gelu_tanh", dropout=0.0, ) @@ -122,7 +131,9 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co n_embd=n_embd, output_mult=(BASE_WIDTH / n_embd) if mup else 1.0, # muP readout multiplier 1/m patch_embed=PatchEmbed.Config( - proj=_lin(patch_dim, n_embd, std=patch_dim**-0.5), # input embed: width-independent + proj=_lin( + patch_dim, n_embd, std=patch_dim**-0.5 + ), # input embed: width-independent patch_size=PATCH_SIZE, ), pos_embedding=Embedding.Config( @@ -130,7 +141,8 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co ), blocks=[ PlanViTBlock.Config( - attention=_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), mlp=_mlp(n_embd, mup=mup) + attention=_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), + mlp=_mlp(n_embd, mup=mup), ) for _ in range(N_LAYER) ], diff --git a/torchtitan/experiments/plan_vit/model.py b/torchtitan/experiments/plan_vit/model.py index a301f3a756..a5cd17facf 100644 --- a/torchtitan/experiments/plan_vit/model.py +++ b/torchtitan/experiments/plan_vit/model.py @@ -1,11 +1,19 @@ +# 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. + """Plan ViT: raw camera frames -> patches -> transformer -> plan. NO VAE. A self-contained planning model for the muP + scaling study, built from torchtitan.models.common blocks the same way path/model.py is. Scales cleanly by width (n_embd / n_head) for muTransfer. """ + from __future__ import annotations from dataclasses import dataclass +from xx.ml_tools.constants.model import ModelInputs import torch import torch.nn as nn @@ -26,7 +34,6 @@ from torchtitan.protocols.model import BaseModel from torchtitan.protocols.module import Module, ModuleList from torchtitan.tools.logging import logger -from xx.ml_tools.constants.model import ModelInputs class PlanViTMLP(Module): @@ -125,7 +132,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: x = rearrange( x, "b (t pt) c (h ph) (w pw) -> b (t h w) (pt c ph pw)", pt=pt, ph=ph, pw=pw ) - return self.proj(x.to(self.proj.weight.dtype)) # match the bf16 (mp) weights, like path's vision + return self.proj( + x.to(self.proj.weight.dtype) + ) # match the bf16 (mp) weights, like path's vision class PlanHead(Module): @@ -197,7 +206,9 @@ def _frames(self, inputs: dict[str, torch.Tensor] | torch.Tensor) -> torch.Tenso return inputs img, big = inputs[ModelInputs.IMG], inputs[ModelInputs.BIG_IMG] frame = torch.cat([img[:, -1], big[:, -1]], dim=1).unsqueeze(1) - return (frame.float() - 127.5) / 63.75 # uint8 YUV -> normalized float (mean 255/2, std 255/4 like path) + return ( + frame.float() - 127.5 + ) / 63.75 # uint8 YUV -> normalized float (mean 255/2, std 255/4 like path) def forward( self, inputs: dict[str, torch.Tensor] | torch.Tensor diff --git a/torchtitan/experiments/plan_vit/trainer.py b/torchtitan/experiments/plan_vit/trainer.py index b6df10a2b8..1f536f033f 100644 --- a/torchtitan/experiments/plan_vit/trainer.py +++ b/torchtitan/experiments/plan_vit/trainer.py @@ -1,8 +1,15 @@ +# 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. + """Thin trainer for plan_vit: model(inputs_dict) -> pred dict, PathLoss(pred, targets), backward. Mirrors PathTrainer's generic core without the path-specific ONNX export / driving validator / reports a scaling study doesn't need. loss + dataloader come from the base Trainer.Config (set in config_registry). """ + from __future__ import annotations import time From 8fd8c52a64634a3bc8d9741ff72709dbd045b79a Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Tue, 23 Jun 2026 19:39:37 -0700 Subject: [PATCH 03/28] plan_vit: drop the 1/m readout multiplier, output now width-stable --- torchtitan/experiments/plan_vit/config_registry.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/torchtitan/experiments/plan_vit/config_registry.py b/torchtitan/experiments/plan_vit/config_registry.py index 498a280868..c758ffd61b 100644 --- a/torchtitan/experiments/plan_vit/config_registry.py +++ b/torchtitan/experiments/plan_vit/config_registry.py @@ -129,7 +129,7 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co patch_size=PATCH_SIZE, in_channels=IN_CHANNELS, n_embd=n_embd, - output_mult=(BASE_WIDTH / n_embd) if mup else 1.0, # muP readout multiplier 1/m + output_mult=1.0, # no readout multiplier: the 1/m mult broke output width-stability (coord check) patch_embed=PatchEmbed.Config( proj=_lin( patch_dim, n_embd, std=patch_dim**-0.5 From b446573fa0c38d6fc974df97af173404a32d96f7 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Tue, 23 Jun 2026 21:16:54 -0700 Subject: [PATCH 04/28] plan_vit: derive data parallelism from the launch so N>1 nodes validate the config hardcoded dp_shard=8 (only valid at world_size=8, i.e. N=1). launching N=2 (world_size=16) tripped the parallel-dims assertion at startup. derive replicate=num_nodes, shard=local_world_size from env like path does. --- torchtitan/experiments/plan_vit/config_registry.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/torchtitan/experiments/plan_vit/config_registry.py b/torchtitan/experiments/plan_vit/config_registry.py index c758ffd61b..3cf4ac34f5 100644 --- a/torchtitan/experiments/plan_vit/config_registry.py +++ b/torchtitan/experiments/plan_vit/config_registry.py @@ -244,6 +244,10 @@ def _optimizer_config( def _plan_vit( flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 ) -> PlanViTTrainer.Config: + # derive data parallelism from the launch (like path), so any N nodes x GPUs validate + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) + world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) + num_nodes = int(os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size))) return PlanViTTrainer.Config( loss=PathLoss.Config(), model_spec=model_registry(flavor, mup=mup), @@ -268,7 +272,8 @@ def _plan_vit( mixed_precision_reduce="float32", ), parallelism=ParallelismConfig( - data_parallel_replicate_degree=1, data_parallel_shard_degree=8 + data_parallel_replicate_degree=num_nodes, + data_parallel_shard_degree=local_world_size, ), checkpoint=CheckpointManager.Config(enable=False), metrics=MetricsProcessor.Config( From e160f7eec467ed900a7be7a6d216cbf7ba2ff2c1 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Wed, 24 Jun 2026 14:57:44 -0700 Subject: [PATCH 05/28] plan_vit: unset lr_scheduler total_steps so it tracks training.steps A pinned total_steps wrapped the cosine schedule, making the LR oscillate when training.steps exceeded it. None falls back to the real training_steps. --- torchtitan/experiments/plan_vit/config_registry.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/torchtitan/experiments/plan_vit/config_registry.py b/torchtitan/experiments/plan_vit/config_registry.py index 3cf4ac34f5..135b15f09a 100644 --- a/torchtitan/experiments/plan_vit/config_registry.py +++ b/torchtitan/experiments/plan_vit/config_registry.py @@ -247,7 +247,9 @@ def _plan_vit( # derive data parallelism from the launch (like path), so any N nodes x GPUs validate local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int(os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size))) + num_nodes = int( + os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) + ) return PlanViTTrainer.Config( loss=PathLoss.Config(), model_spec=model_registry(flavor, mup=mup), @@ -256,7 +258,7 @@ def _plan_vit( optimizer=_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), lr_scheduler=LRSchedulersContainer.Config( warmup_steps=round(STEPS * 0.1), - total_steps=STEPS, + total_steps=None, # use the real training.steps; a fixed value wraps the cosine on longer runs decay_ratio=0.8, decay_type="cosine", min_lr_factor=0.0, From a430cef94919f5314cb21f61dab4cd63ac958e08 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Wed, 24 Jun 2026 20:32:58 -0700 Subject: [PATCH 06/28] plan_vit: restore canonical muP readout (1/m mult, base-width init, base lr) output_mult=1 made the coord check flat for the wrong reason (compensating errors that cancel only at low step count). Canonical muP readout: forward multiplier 1/m, base-width init, base lr (vector-like under Adam, ninf==1). Verified by coord check: init output slopes ~1/sqrt(m), trained output flat. --- torchtitan/experiments/plan_vit/config_registry.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/torchtitan/experiments/plan_vit/config_registry.py b/torchtitan/experiments/plan_vit/config_registry.py index 135b15f09a..6eb2bb5fc1 100644 --- a/torchtitan/experiments/plan_vit/config_registry.py +++ b/torchtitan/experiments/plan_vit/config_registry.py @@ -129,7 +129,9 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co patch_size=PATCH_SIZE, in_channels=IN_CHANNELS, n_embd=n_embd, - output_mult=1.0, # no readout multiplier: the 1/m mult broke output width-stability (coord check) + output_mult=(BASE_WIDTH / n_embd) + if mup + else 1.0, # muP readout fwd mult 1/m (init output slopes ~1/sqrt(m)) patch_embed=PatchEmbed.Config( proj=_lin( patch_dim, n_embd, std=patch_dim**-0.5 @@ -149,7 +151,9 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co norm=_ln(n_embd), plan_head=PlanHead.Config( norm=_ln(n_embd), - head=_lin(n_embd, PLAN_SIZE, std=_hidden_std(n_embd, mup=mup)), + head=_lin( + n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5 + ), # muP readout: base-width init ), ) @@ -169,10 +173,11 @@ def model_registry(flavor: str, *, mup: bool) -> ModelSpec: STEPS = 512 # per-run step budget; override with training.steps=N on the CLI # learning rate is the muTransfer sweep axis: one run per (flavor, lr); set with `-e PLAN_VIT_LR=...` SWEEP_LR = float(os.getenv("PLAN_VIT_LR", "3e-4")) -# hidden + output matrix weights get muP lr eta/m; input embed, norms, biases get eta +# hidden matrix weights get muP lr eta/m; input embed, readout, norms, biases get base eta +# (readout is fan_in-infinite only, so Adam treats it vector-like -> base lr, not eta/m) MUP_PATTERN = ( r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" - r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj|plan_head\.head)\.weight$" + r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" ) From c6dbad228062d79f02508d8a628e3b8f24e5a511 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Thu, 25 Jun 2026 15:32:07 -0700 Subject: [PATCH 07/28] path: muP plan_vit as path/vit.py on PathTrainer (validator/onnx off) adds default-off plan_target_last_frame flag so the single-frame ViT supervises the last plan frame; convnext unchanged when off --- .../experiments/path/config_registry.py | 166 ++++++--- torchtitan/experiments/path/trainer.py | 97 ++++- torchtitan/experiments/path/vit.py | 259 +++++++++++++ .../experiments/path/vit_config_registry.py | 341 ++++++++++++++++++ 4 files changed, 801 insertions(+), 62 deletions(-) create mode 100644 torchtitan/experiments/path/vit.py create mode 100644 torchtitan/experiments/path/vit_config_registry.py diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index bff5fa9786..dc716b545d 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -1,8 +1,31 @@ +# 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 import os from functools import partial +from xx.datasets.helpers import DEFAULT_BIG_TRAIN_LIST +from xx.ml_tools.constants.model import ( + frame_constants_from_fps, + FRAME_TYPE, + INPUT_FRAMES_NAMES, + ModelInputs, + N_FRAMES, + SUPERCOMBO_FPS, + TEMPORAL_INPUTS, +) +from xx.training.path.config import DatasetConfig as XXPathDatasetConfig +from xx.training.path.hydra_configs import ( + DRIVING_HEADS, + META_HEADS, + POSE_HEADS, + TEMPORAL_META_HEADS, +) import torch.nn as nn @@ -20,23 +43,13 @@ from torchtitan.models.common import Embedding, LayerNorm, Linear from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model_spec import ModelSpec -from xx.datasets.helpers import DEFAULT_BIG_TRAIN_LIST -from xx.ml_tools.constants.model import ( - SUPERCOMBO_FPS, - FRAME_TYPE, - INPUT_FRAMES_NAMES, - N_FRAMES, - TEMPORAL_INPUTS, - ModelInputs, - frame_constants_from_fps, -) -from xx.training.path.config import DatasetConfig as XXPathDatasetConfig -from xx.training.path.hydra_configs import DRIVING_HEADS, META_HEADS, POSE_HEADS, TEMPORAL_META_HEADS from .dataset import PathDataLoader +from .loss import PathLoss from .model import ( Hydra, LinearEncoder, + parallelize_path, PathHead, PathMLP, PathModel, @@ -49,15 +62,29 @@ TemporalPolicy, TemporalSummarizer, Vision, - parallelize_path, ) -from .loss import PathLoss from .onnx_checkpoint import PathOnnxCheckpointManager from .trainer import PathTrainer from .validate import PathValidator +# Plan ViT flavors ride PathTrainer too; re-exported so `--module path --config vit_*` resolves here, +# the same way convnext_* do (the config manager looks the name up on this module). +from .vit_config_registry import ( # noqa: F401 + vit_mup_w1024, + vit_mup_w2048, + vit_mup_w256, + vit_mup_w512, + vit_standard_w1024, + vit_standard_w2048, + vit_standard_w256, + vit_standard_w512, +) + -_LINEAR_INIT = {"weight": partial(nn.init.normal_, mean=0.0, std=0.02), "bias": nn.init.zeros_} +_LINEAR_INIT = { + "weight": partial(nn.init.normal_, mean=0.0, std=0.02), + "bias": nn.init.zeros_, +} _NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} @@ -90,10 +117,10 @@ def convnext_xxlarge() -> PathTrainer.Config: def _path(flavor: str) -> PathTrainer.Config: - steps = 1024*100 + steps = 1024 * 100 validation_freq = 1024 reports = { - name: [validation_freq, steps //2 , steps] + name: [validation_freq, steps // 2, steps] for name in ( "analyse_driving", "analyse_lat_no_noise", @@ -107,10 +134,14 @@ def _path(flavor: str) -> PathTrainer.Config: mixed_precision_param = "bfloat16" local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int(os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size))) + num_nodes = int( + os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) + ) reporterv2_host = os.getenv("REPORTERV2_HOST") reporterv2_training_id = os.getenv("REPORTERV2_TRAINING_ID") - checkpoint_base_folder = f"{reporterv2_host.rstrip('/')}/checkpoint" if reporterv2_host else "" + checkpoint_base_folder = ( + f"{reporterv2_host.rstrip('/')}/checkpoint" if reporterv2_host else "" + ) fps = SUPERCOMBO_FPS plan_only = False return PathTrainer.Config( @@ -153,7 +184,9 @@ def _path(flavor: str) -> PathTrainer.Config: fps=fps, activation_checkpoint=FullAC.Config(), compile=CompileConfig(enable=True, components=["model"]), - metrics=MetricsProcessor.Config(log_freq=16, enable_reporterv2=True, save_freq=validation_freq), + metrics=MetricsProcessor.Config( + log_freq=16, enable_reporterv2=True, save_freq=validation_freq + ), validator=PathValidator.Config( enable=True, freq=validation_freq, @@ -171,8 +204,12 @@ def _model_config(flavor: str) -> PathModel.Config: n_frames_input = N_FRAMES input_frame_names = INPUT_FRAMES_NAMES input_frame_type = FRAME_TYPE - frame_constants = frame_constants_from_fps(n_frames=n_frames_input, frame_type=input_frame_type) - in_channels = sum(frame_constants["frame_shapes"][name][0] for name in input_frame_names) + frame_constants = frame_constants_from_fps( + n_frames=n_frames_input, frame_type=input_frame_type + ) + in_channels = sum( + frame_constants["frame_shapes"][name][0] for name in input_frame_names + ) block_size = len(frame_constants["history_idxs"]) temporal_len = frame_constants["temporal_len"] dim = vision_features @@ -202,9 +239,13 @@ def _model_config(flavor: str) -> PathModel.Config: temporal_summarizer=TemporalSummarizer.Config( mlp1=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), mlp2=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - desire_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * temporal_len, dim), + desire_encoder=_encoder( + TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * temporal_len, dim + ), traffic_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.TRAFFIC][0], dim), - action_t_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], dim), + action_t_encoder=_encoder( + TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], dim + ), transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( @@ -214,17 +255,25 @@ def _model_config(flavor: str) -> PathModel.Config: for _ in range(4) ] ), - pos_embedding=Embedding.Config(num_embeddings=block_size, embedding_dim=dim, param_init=_LINEAR_INIT), + pos_embedding=Embedding.Config( + num_embeddings=block_size, + embedding_dim=dim, + param_init=_LINEAR_INIT, + ), block_size=block_size, dense_training_outputs=True, ), - temporal_hydra=_hydra(_heads(DRIVING_HEADS + TEMPORAL_META_HEADS), in_features=dim, mlp_mult=2), + temporal_hydra=_hydra( + _heads(DRIVING_HEADS + TEMPORAL_META_HEADS), in_features=dim, mlp_mult=2 + ), history_idxs=tuple(int(x) for x in frame_constants["history_idxs"]), ), ) -def _dataloader_config(*, split: str, fps: int, plan_only: bool) -> PathDataLoader.Config: +def _dataloader_config( + *, split: str, fps: int, plan_only: bool +) -> PathDataLoader.Config: base = XXPathDatasetConfig(fps=fps, plan_only=plan_only) return PathDataLoader.Config( dataset=DEFAULT_BIG_TRAIN_LIST, @@ -243,7 +292,9 @@ def _dataloader_config(*, split: str, fps: int, plan_only: bool) -> PathDataLoad ) -def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnxCheckpointManager.Config: +def _checkpoint_config( + folder: str, base_folder: str, interval: int +) -> PathOnnxCheckpointManager.Config: frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) temporal_len = frame_constants["temporal_len"] vision_input_names = [ModelInputs.IMG, ModelInputs.BIG_IMG] @@ -266,10 +317,10 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.ACTION_T][0]], ] return PathOnnxCheckpointManager.Config( - keep_latest_k=0, # keep all checkpoints + keep_latest_k=0, # keep all checkpoints enable=True, checkpoint_base_folder=base_folder, - save_model_state_dict=True, # another copy of full state dict + save_model_state_dict=True, # another copy of full state dict export_onnx=True, enable_first_step_checkpoint=True, folder=folder, @@ -277,7 +328,7 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx input_names=input_names, input_shapes=input_shapes, input_dtypes=["float16"] * len(input_names), - onnx_model_dtype="float16", # WIP: test if fp16 doesn't degrade performance + onnx_model_dtype="float16", # WIP: test if fp16 doesn't degrade performance vision_input_names=vision_input_names, temporal_policy_input_names=temporal_policy_input_names, ) @@ -286,7 +337,11 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx def _si_int(value: str | int) -> int: suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} value = str(value).strip().lower() - return int(float(value[:-1]) * suffixes[value[-1]]) if value[-1] in suffixes else int(value) + return ( + int(float(value[:-1]) * suffixes[value[-1]]) + if value[-1] in suffixes + else int(value) + ) def _optimizer_config() -> OptimizersContainer.Config: @@ -310,7 +365,9 @@ def _optimizer_config() -> OptimizersContainer.Config: def _heads(heads) -> tuple[PathHead, ...]: - return tuple(PathHead(head.name, head.output_size, head.mlp, head.scale) for head in heads) + return tuple( + PathHead(head.name, head.output_size, head.mlp, head.scale) for head in heads + ) def _hidden_dim(dim: int, mlp_mult: float, multiple_of: int = 256) -> int: @@ -322,8 +379,12 @@ def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Co hidden = _hidden_dim(dim, mlp_mult) return PathMLP.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - c_fc=Linear.Config(in_features=dim, out_features=hidden, bias=bias, param_init=_LINEAR_INIT), - c_proj=Linear.Config(in_features=hidden, out_features=dim, bias=bias, param_init=_LINEAR_INIT), + c_fc=Linear.Config( + in_features=dim, out_features=hidden, bias=bias, param_init=_LINEAR_INIT + ), + c_proj=Linear.Config( + in_features=hidden, out_features=dim, bias=bias, param_init=_LINEAR_INIT + ), act="gelu_tanh", dropout=dropout, ) @@ -331,8 +392,15 @@ def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Co def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: return LinearEncoder.Config( - in_layer=Linear.Config(in_features=in_features, out_features=dim, bias=True, param_init=_LINEAR_INIT), - out_layer=Linear.Config(in_features=dim, out_features=dim, bias=False, param_init=_LINEAR_INIT), + in_layer=Linear.Config( + in_features=in_features, + out_features=dim, + bias=True, + param_init=_LINEAR_INIT, + ), + out_layer=Linear.Config( + in_features=dim, out_features=dim, bias=False, param_init=_LINEAR_INIT + ), ) @@ -342,8 +410,12 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), - c_attn=Linear.Config(in_features=dim, out_features=3 * dim, bias=True, param_init=_LINEAR_INIT), - c_proj=Linear.Config(in_features=dim, out_features=dim, bias=True, param_init=_LINEAR_INIT), + c_attn=Linear.Config( + in_features=dim, out_features=3 * dim, bias=True, param_init=_LINEAR_INIT + ), + c_proj=Linear.Config( + in_features=dim, out_features=dim, bias=True, param_init=_LINEAR_INIT + ), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, head_dim=head_dim, @@ -351,10 +423,16 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co ) -def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> Hydra.Config: +def _hydra( + heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float +) -> Hydra.Config: return Hydra.Config( heads=heads, - head_mlps={head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) for head in heads if head.mlp}, + head_mlps={ + head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) + for head in heads + if head.mlp + }, final_layers={ head.name: Linear.Config( in_features=in_features, @@ -364,5 +442,9 @@ def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> ) for head in heads }, - scale_layers={head.name: ScaleLayer.Config(n_features=head.output_size) for head in heads if head.scale}, + scale_layers={ + head.name: ScaleLayer.Config(n_features=head.output_size) + for head in heads + if head.scale + }, ) diff --git a/torchtitan/experiments/path/trainer.py b/torchtitan/experiments/path/trainer.py index d47da3d537..d7c06bbb78 100644 --- a/torchtitan/experiments/path/trainer.py +++ b/torchtitan/experiments/path/trainer.py @@ -27,22 +27,32 @@ class Config(Trainer.Config): checkpoint: PathOnnxCheckpointManager.Config miniray: dict[str, Any] = field(default_factory=dict) fps: int + # single-frame plan supervision: slice the dense plan target to the last frame before the + # loss. off by default (dense; convnext/worldmodel unchanged); on for the single-frame plan_vit. + plan_target_last_frame: bool = False def __post_init__(self) -> None: Trainer.Config.__post_init__(self) if self.codedir: self.miniray = {**self.miniray, "codedir": self.codedir} - self.validator.miniray = {**self.validator.miniray, "codedir": self.codedir} + self.validator.miniray = { + **self.validator.miniray, + "codedir": self.codedir, + } def __init__(self, config: Config): super().__init__(config) training_id = os.getenv("REPORTERV2_TRAINING_ID") or "local" - self.unique_segment_counter = StringUniqueCounter(f"unique_ids:{training_id}:path:train") + self.unique_segment_counter = StringUniqueCounter( + f"unique_ids:{training_id}:path:train" + ) self.loss_fn.to(self.device) def batch_generator( self, - data_iterable: Iterable[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]], + data_iterable: Iterable[ + tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] + ], ) -> Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]]: data_iterator = iter(data_iterable) while True: @@ -51,8 +61,12 @@ def batch_generator( input_dict, targets = next(data_iterator) except StopIteration as ex: raise DataloaderExhaustedError() from ex - self.metrics_processor.ntokens_since_last_log += next(iter(input_dict.values())).shape[0] - self.metrics_processor.data_loading_times.append(time.perf_counter() - data_load_start) + self.metrics_processor.ntokens_since_last_log += next( + iter(input_dict.values()) + ).shape[0] + self.metrics_processor.data_loading_times.append( + time.perf_counter() - data_load_start + ) yield input_dict, targets @sl.log_trace_span("post_dataloading_process") @@ -76,6 +90,9 @@ def forward_backward_step( assert len(self.model_parts) == 1 with self.train_context(): pred = self.model_parts[0](inputs) + if self.config.plan_target_last_frame: + # single-frame plan models (plan_vit) predict the last frame's plan; supervise that frame + labels = {**labels, "plan": labels["plan"][:, -1]} loss_vec, metrics = self.loss_fn(pred, labels) loss = loss_vec.sum() / local_samples del pred @@ -84,7 +101,9 @@ def forward_backward_step( def train_step( self, - data_iterator: Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]], + data_iterator: Iterator[ + tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] + ], ) -> None: self.optimizers.zero_grad() lr_metrics = self.lr_schedulers.get_metrics() @@ -99,7 +118,9 @@ def train_step( input_dict, targets = next(data_iterator) local_samples += next(iter(input_dict.values())).shape[0] if "info" in input_dict: - step_segment_names.update(segment_names_from_info(input_dict["info"])) + step_segment_names.update( + segment_names_from_info(input_dict["info"]) + ) microbatches.append((input_dict, targets)) sl.log_trace_scalar({"local_samples": int(local_samples)}) @@ -108,7 +129,9 @@ def train_step( global_samples = dist_utils.dist_sum(local_samples, batch_mesh) else: global_samples = local_samples.float() - global_samples = torch.as_tensor(global_samples, dtype=torch.float32, device=self.device) + global_samples = torch.as_tensor( + global_samples, dtype=torch.float32, device=self.device + ) global_samples_value = float(global_samples.item()) accumulated_losses = [] @@ -125,7 +148,10 @@ def train_step( for name, value in metrics.items(): if name == "loss": continue - metric_sums[name] = metric_sums.get(name, torch.zeros((), device=self.device)) + value.float().sum() + metric_sums[name] = ( + metric_sums.get(name, torch.zeros((), device=self.device)) + + value.float().sum() + ) with sl.log_trace_span("optim"): grad_norm = dist_utils.clip_grad_norm_( @@ -149,17 +175,30 @@ def train_step( loss_mesh = parallel_dims.get_optional_mesh("loss") local_loss_sum = loss * local_samples global_avg_loss, global_max_loss, global_samples_seen = ( - dist_utils.dist_sum(local_loss_sum.detach(), loss_mesh) / global_samples_value, + dist_utils.dist_sum(local_loss_sum.detach(), loss_mesh) + / global_samples_value, dist_utils.dist_max(loss.detach(), loss_mesh), - dist_utils.dist_sum(torch.tensor(self.ntokens_seen, dtype=torch.int64, device=self.device), loss_mesh), + dist_utils.dist_sum( + torch.tensor( + self.ntokens_seen, dtype=torch.int64, device=self.device + ), + loss_mesh, + ), ) - metric_sums = {k: dist_utils.dist_sum(v, loss_mesh) for k, v in metric_sums.items()} + metric_sums = { + k: dist_utils.dist_sum(v, loss_mesh) for k, v in metric_sums.items() + } else: global_avg_loss = global_max_loss = float(loss.detach().item()) global_samples_seen = self.ntokens_seen path_metrics = { - f"path/{k}": float((torch.as_tensor(v, dtype=torch.float32, device=self.device) / global_samples).item()) + f"path/{k}": float( + ( + torch.as_tensor(v, dtype=torch.float32, device=self.device) + / global_samples + ).item() + ) for k, v in metric_sums.items() } unique_segments_seen = ( @@ -170,7 +209,12 @@ def train_step( dataset_metrics = { "dataset/unique_segments_seen": unique_segments_seen, } - extra_metrics = {"n_samples_seen": global_samples_seen, **lr_metrics, **path_metrics, **dataset_metrics} + extra_metrics = { + "n_samples_seen": global_samples_seen, + **lr_metrics, + **path_metrics, + **dataset_metrics, + } self.metrics_processor.log( self.step, global_avg_loss, @@ -188,15 +232,28 @@ def close(self) -> None: def state_dict(self) -> dict[str, Any]: state = super().state_dict() state["unique_segment_counter"] = self.unique_segment_counter.state_dict() - validator_unique_segment_counter = getattr(getattr(self, "validator", None), "unique_segment_counter", None) + validator_unique_segment_counter = getattr( + getattr(self, "validator", None), "unique_segment_counter", None + ) if validator_unique_segment_counter is not None: - state["validation_unique_segment_counter"] = validator_unique_segment_counter.state_dict() + state[ + "validation_unique_segment_counter" + ] = validator_unique_segment_counter.state_dict() return state def load_state_dict(self, state_dict: dict[str, Any]) -> None: super().load_state_dict(state_dict) if "unique_segment_counter" in state_dict: - self.unique_segment_counter.load_state_dict(state_dict["unique_segment_counter"]) - validator_unique_segment_counter = getattr(getattr(self, "validator", None), "unique_segment_counter", None) - if validator_unique_segment_counter is not None and "validation_unique_segment_counter" in state_dict: - validator_unique_segment_counter.load_state_dict(state_dict["validation_unique_segment_counter"]) + self.unique_segment_counter.load_state_dict( + state_dict["unique_segment_counter"] + ) + validator_unique_segment_counter = getattr( + getattr(self, "validator", None), "unique_segment_counter", None + ) + if ( + validator_unique_segment_counter is not None + and "validation_unique_segment_counter" in state_dict + ): + validator_unique_segment_counter.load_state_dict( + state_dict["validation_unique_segment_counter"] + ) diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py new file mode 100644 index 0000000000..1d3bc3fed7 --- /dev/null +++ b/torchtitan/experiments/path/vit.py @@ -0,0 +1,259 @@ +# 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. + +"""Plan ViT for the path experiment: raw camera frames -> patches -> transformer -> plan. NO VAE. + +Ported verbatim (behavior-identical) from experiments/plan_vit/model.py so the ViT can ride +PathTrainer via config instead of the standalone PlanViTTrainer. A self-contained planning model +for the muP + scaling study, built from torchtitan.models.common blocks the same way path/model.py +is. Scales cleanly by width (n_embd / n_head) for muTransfer. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from xx.ml_tools.constants.model import ModelInputs + +import torch +import torch.nn as nn +from einops import rearrange +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy + +from torchtitan.config import ( + CompileConfig, + ParallelismConfig, + TORCH_DTYPE_MAP, + TrainingConfig, +) +from torchtitan.distributed import ParallelDims +from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig +from torchtitan.models.common import Embedding, LayerNorm, Linear, RMSNorm +from torchtitan.models.common.attention import ScaledDotProductAttention +from torchtitan.protocols.model import BaseModel +from torchtitan.protocols.module import Module, ModuleList +from torchtitan.tools.logging import logger + + +class PlanViTMLP(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + norm: LayerNorm.Config | RMSNorm.Config + c_fc: Linear.Config + c_proj: Linear.Config + act: str + dropout: float + + def __init__(self, config: Config): + super().__init__() + self.norm = config.norm.build() + self.c_fc = config.c_fc.build() + self.act = ( + nn.GELU(approximate="tanh") if config.act == "gelu_tanh" else nn.GELU() + ) + self.c_proj = config.c_proj.build() + self.dropout = nn.Dropout(config.dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.dropout(self.c_proj(self.act(self.c_fc(self.norm(x))))) + + +class PlanViTAttention(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + norm: LayerNorm.Config | RMSNorm.Config + q_norm: LayerNorm.Config | RMSNorm.Config | None + k_norm: LayerNorm.Config | RMSNorm.Config | None + c_attn: Linear.Config + c_proj: Linear.Config + inner_attention: ScaledDotProductAttention.Config + n_head: int + head_dim: int + dropout: float + + def __init__(self, config: Config): + super().__init__() + self.n_head = config.n_head + self.head_dim = config.head_dim + self.norm = config.norm.build() + self.q_norm = ( + config.q_norm.build() if config.q_norm is not None else nn.Identity() + ) + self.k_norm = ( + config.k_norm.build() if config.k_norm is not None else nn.Identity() + ) + self.c_attn = config.c_attn.build() + self.c_proj = config.c_proj.build() + self.inner_attention = config.inner_attention.build() + self.dropout = nn.Dropout(config.dropout) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + b, t, _ = x.shape + qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) + q, k, v = qkv.unbind(2) + q, k = self.q_norm(q), self.k_norm(k) + x = self.inner_attention( + q, k, v, is_causal=False + ) # ViT: bidirectional over patches + return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) + + +class PlanViTBlock(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + attention: PlanViTAttention.Config + mlp: PlanViTMLP.Config + + def __init__(self, config: Config): + super().__init__() + self.attention = config.attention.build() + self.mlp = config.mlp.build() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x + self.attention(x) + return x + self.mlp(x) + + +class PatchEmbed(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + proj: Linear.Config + patch_size: tuple[int, int, int] # (pt, ph, pw) + + def __init__(self, config: Config): + super().__init__() + self.patch_size = config.patch_size + self.proj = config.proj.build() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + # x: (B, T, C, H, W) raw frames -> (B, num_patches, patch_dim) -> (B, num_patches, n_embd) + pt, ph, pw = self.patch_size + x = rearrange( + x, "b (t pt) c (h ph) (w pw) -> b (t h w) (pt c ph pw)", pt=pt, ph=ph, pw=pw + ) + return self.proj( + x.to(self.proj.weight.dtype) + ) # match the bf16 (mp) weights, like path's vision + + +class PlanHead(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + norm: LayerNorm.Config | RMSNorm.Config + head: Linear.Config + + def __init__(self, config: Config): + super().__init__() + self.norm = config.norm.build() + self.head = config.head.build() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.head(self.norm(x)) + + +class PlanViT(BaseModel): + @dataclass(kw_only=True, slots=True) + class Config(BaseModel.Config): + input_size: tuple[int, int, int] # (n_frames, H, W) + patch_size: tuple[int, int, int] + in_channels: int + n_embd: int + output_mult: float # muP readout multiplier 1/m (m = n_embd / base); 1.0 for standard param + patch_embed: PatchEmbed.Config + pos_embedding: Embedding.Config + blocks: list[PlanViTBlock.Config] + norm: LayerNorm.Config | RMSNorm.Config + plan_head: PlanHead.Config + + @property + def num_patches(self) -> int: + t, h, w = self.input_size + pt, ph, pw = self.patch_size + return (t // pt) * (h // ph) * (w // pw) + + def update_from_config(self, *, config, **kwargs) -> None: + parallelism = config.parallelism + for name, degree in { + "tensor parallel": parallelism.tensor_parallel_degree, + "context parallel": parallelism.context_parallel_degree, + "pipeline parallel": parallelism.pipeline_parallel_degree, + "expert parallel": parallelism.expert_parallel_degree, + }.items(): + if degree > 1: + raise ValueError(f"plan_vit does not support {name}") + + def get_nparams_and_flops(self, model: Module, seq_len: int) -> tuple[int, int]: + nparams = sum(p.numel() for p in model.parameters()) + return nparams, 6 * nparams + + def __init__(self, config: Config): + super().__init__() + self.config = config + self.patch_embed = config.patch_embed.build() + self.pos_embedding = config.pos_embedding.build() + self.blocks = ModuleList([block.build() for block in config.blocks]) + self.norm = config.norm.build() + self.plan_head = config.plan_head.build() + + def verify_module_protocol(self) -> None: + pass # nn.Dropout/GELU/Identity are plain nn.Module, like path + + def _frames(self, inputs: dict[str, torch.Tensor] | torch.Tensor) -> torch.Tensor: + # production input: two cameras IMG, BIG_IMG, each (B, T, 12, H, W) YUV. Take the current frame of each, + # channel-stack -> (B, 1, 24, H, W). NO VAE. A plain tensor (testing) is passed through unchanged. + if isinstance(inputs, torch.Tensor): + return inputs + img, big = inputs[ModelInputs.IMG], inputs[ModelInputs.BIG_IMG] + frame = torch.cat([img[:, -1], big[:, -1]], dim=1).unsqueeze(1) + return ( + frame.float() - 127.5 + ) / 63.75 # uint8 YUV -> normalized float (mean 255/2, std 255/4 like path) + + def forward( + self, inputs: dict[str, torch.Tensor] | torch.Tensor + ) -> dict[str, torch.Tensor]: + x = self.patch_embed(self._frames(inputs)) + pos = self.pos_embedding(torch.arange(x.shape[1], device=x.device)) + x = x + rearrange(pos, "t c -> () t c") + for block in self.blocks: + x = block(x) + x = self.norm(x) + # global-pool the patches -> plan; the muP readout multiplier keeps the output width-stable + return {"plan": self.plan_head(x.mean(dim=1)) * self.config.output_mult} + + +def parallelize_vit( + model: PlanViT, + *, + parallel_dims: ParallelDims, + training: TrainingConfig, + parallelism: ParallelismConfig, + compile_config: CompileConfig, + ac_config: ActivationCheckpointingConfig, + dump_folder: str, +) -> PlanViT: + if ( + parallel_dims.tp_enabled + or parallel_dims.cp_enabled + or parallel_dims.pp_enabled + or parallel_dims.ep_enabled + ): + raise ValueError("plan_vit supports data parallelism only") + names = ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] + dp_mesh: DeviceMesh = parallel_dims.get_mesh(names) + mp_policy = MixedPrecisionPolicy( + param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param], + reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], + cast_forward_inputs=True, + ) + fsdp_config = {"mesh": dp_mesh, "mp_policy": mp_policy} + for idx, block in enumerate(model.blocks): + fully_shard( + block, **fsdp_config, reshard_after_forward=(idx < len(model.blocks) - 1) + ) + fully_shard(model, **fsdp_config) + logger.info("Applied FSDP to plan_vit") + return model diff --git a/torchtitan/experiments/path/vit_config_registry.py b/torchtitan/experiments/path/vit_config_registry.py new file mode 100644 index 0000000000..2c9e7f8fb3 --- /dev/null +++ b/torchtitan/experiments/path/vit_config_registry.py @@ -0,0 +1,341 @@ +# 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. + +"""Config assembly + flavors for the path ViT, riding PathTrainer. + +The muP recipe (readout init/mult, eta/m optimizer groups, qk-norm, scheduler, widths, training) is +carried over verbatim from experiments/plan_vit/config_registry.py so it stays identical to the proven +config. The only change is the trainer wiring: these flavors return PathTrainer.Config (not the standalone +PlanViTTrainer.Config) with the path-specifics turned off in config -- the driving validator is disabled +and the checkpoint manager is the plain CheckpointManager with onnx export off. + +Width flavors scale n_head at fixed head_dim=64 (the clean muP axis); base = w256. Two cameras are +channel-stacked into in_channels=24 (no VAE). +""" + +from __future__ import annotations + +import math +import os +from functools import partial +from xx.ml_tools.constants.model import SUPERCOMBO_FPS + +import torch.nn as nn + +from torchtitan.components.checkpoint import CheckpointManager +from torchtitan.components.lr_scheduler import LRSchedulersContainer +from torchtitan.components.metrics import MetricsProcessor +from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig +from torchtitan.components.tokenizer import NoOpTokenizer +from torchtitan.config import DebugConfig, ParallelismConfig, TrainingConfig +from torchtitan.models.common import Embedding, LayerNorm, Linear +from torchtitan.models.common.attention import ScaledDotProductAttention +from torchtitan.protocols.model_spec import ModelSpec + +from .dataset import PathDataLoader +from .loss import PathLoss +from .trainer import PathTrainer +from .validate import PathValidator +from .vit import ( + parallelize_vit, + PatchEmbed, + PlanHead, + PlanViT, + PlanViTAttention, + PlanViTBlock, + PlanViTMLP, +) + +_LINEAR_INIT = { + "weight": partial(nn.init.normal_, mean=0.0, std=0.02), + "bias": nn.init.zeros_, +} +_NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} + +HEAD_DIM = 64 +N_LAYER = 8 +INPUT_SIZE = ( + 1, + 128, + 256, +) # current frame; spatial ViT (temporal history is a later variant) +PATCH_SIZE = (1, 16, 8) +IN_CHANNELS = 24 # two cameras (IMG + BIG_IMG), 12 YUV channels each, channel-stacked +PLAN_SIZE = 15 * 33 * 2 # 990, laplacian mu+log-sigma +BASE_WIDTH = 256 +PLAN_VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} + + +def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Config: + return Linear.Config( + in_features=in_f, + out_features=out_f, + bias=bias, + param_init={ + "weight": partial(nn.init.normal_, mean=0.0, std=std), + "bias": nn.init.zeros_, + }, + ) + + +def _hidden_std(fan_in: int, *, mup: bool) -> float: + # muP shrinks hidden/output init to 1/sqrt(fan_in) so pre-activations stay O(1) as width grows; + # standard param holds the base-width variance 1/sqrt(BASE_WIDTH), so it fans out with width. + return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 + + +def _ln(dim: int) -> LayerNorm.Config: + return LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT) + + +def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: + return multiple_of * math.ceil(int(dim * mult) / multiple_of) + + +def _attention( + dim: int, n_head: int, *, mup: bool, qk_norm: bool = True +) -> PlanViTAttention.Config: + head_dim = dim // n_head + return PlanViTAttention.Config( + norm=_ln(dim), + q_norm=_ln(head_dim) if qk_norm else None, + k_norm=_ln(head_dim) if qk_norm else None, + c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), + c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * N_LAYER)), + inner_attention=ScaledDotProductAttention.Config(), + n_head=n_head, + head_dim=head_dim, + dropout=0.0, + ) + + +def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PlanViTMLP.Config: + hidden = _hidden(dim, mult) + return PlanViTMLP.Config( + norm=_ln(dim), + c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), + c_proj=_lin( + hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER) + ), + act="gelu_tanh", + dropout=0.0, + ) + + +def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: + n_embd = PLAN_VIT_WIDTHS[flavor] + n_head = n_embd // HEAD_DIM + pt, ph, pw = PATCH_SIZE + patch_dim = pt * IN_CHANNELS * ph * pw + t, h, w = INPUT_SIZE + num_patches = (t // pt) * (h // ph) * (w // pw) + return PlanViT.Config( + input_size=INPUT_SIZE, + patch_size=PATCH_SIZE, + in_channels=IN_CHANNELS, + n_embd=n_embd, + output_mult=(BASE_WIDTH / n_embd) + if mup + else 1.0, # muP readout fwd mult 1/m (init output slopes ~1/sqrt(m)) + patch_embed=PatchEmbed.Config( + proj=_lin( + patch_dim, n_embd, std=patch_dim**-0.5 + ), # input embed: width-independent + patch_size=PATCH_SIZE, + ), + pos_embedding=Embedding.Config( + num_embeddings=num_patches, embedding_dim=n_embd, param_init=_LINEAR_INIT + ), + blocks=[ + PlanViTBlock.Config( + attention=_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), + mlp=_mlp(n_embd, mup=mup), + ) + for _ in range(N_LAYER) + ], + norm=_ln(n_embd), + plan_head=PlanHead.Config( + norm=_ln(n_embd), + head=_lin( + n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5 + ), # muP readout: base-width init + ), + ) + + +def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: + return ModelSpec( + name="path", + flavor=flavor, + model=_model_config(flavor, mup=mup), + parallelize_fn=parallelize_vit, + pipelining_fn=None, + post_optimizer_build_fn=None, + state_dict_adapter=None, + ) + + +STEPS = 512 # per-run step budget; override with training.steps=N on the CLI +# learning rate is the muTransfer sweep axis: one run per (flavor, lr); set with `-e PLAN_VIT_LR=...` +SWEEP_LR = float(os.getenv("PLAN_VIT_LR", "3e-4")) +# hidden matrix weights get muP lr eta/m; input embed, readout, norms, biases get base eta +# (readout is fan_in-infinite only, so Adam treats it vector-like -> base lr, not eta/m) +MUP_PATTERN = ( + r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" + r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" +) + + +def _si_int(value: str | int) -> int: + suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} + value = str(value).strip().lower() + return ( + int(float(value[:-1]) * suffixes[value[-1]]) + if value[-1] in suffixes + else int(value) + ) + + +def _dataloader_config(*, split: str) -> PathDataLoader.Config: + from xx.common.basedir import XX_BASEDIR + from xx.datasets.constants import BASE_DIR_GT_10M + from xx.training.path.config import DatasetConfig as XXPathDatasetConfig + + base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) + return PathDataLoader.Config( + # prune-10M study data: a seeded random 10k sample of the 10M store (training_2026_02) + dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), + split=split, + shuffle_size=_si_int(base.shuffle_size), + min_mixing=base.min_mixing, + num_writers=base.num_writers, + num_readers=base.num_readers, + fps=base.fps, + pipeline_dir=BASE_DIR_GT_10M, # the 10M store, not the 2.5M big-train list + plan_only=base.plan_only, + limit=base.limit, + n_frames=base.n_frames, + rgb=base.rgb, + unvision=base.unvision, + ) + + +def _optimizer_config( + flavor: str, *, mup: bool, lr: float, wd: float +) -> OptimizersContainer.Config: + m = PLAN_VIT_WIDTHS[flavor] / BASE_WIDTH + common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} + if mup: + groups = [ + ParamGroupConfig( + pattern=MUP_PATTERN, + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr / m}, + ), + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr}, + ), + ] + else: + groups = [ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr}, + ) + ] + return OptimizersContainer.Config( + implementation="fused_opt_states_bf16", param_groups=groups + ) + + +def _vit( + flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 +) -> PathTrainer.Config: + # derive data parallelism from the launch (like path), so any N nodes x GPUs validate + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) + world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) + num_nodes = int( + os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) + ) + return PathTrainer.Config( + loss=PathLoss.Config(), + model_spec=vit_model_registry(flavor, mup=mup), + tokenizer=NoOpTokenizer.Config(), + dataloader=_dataloader_config(split="train"), + optimizer=_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), + lr_scheduler=LRSchedulersContainer.Config( + warmup_steps=round(STEPS * 0.1), + total_steps=None, # use the real training.steps; a fixed value wraps the cosine on longer runs + decay_ratio=0.8, + decay_type="cosine", + min_lr_factor=0.0, + ), + training=TrainingConfig( + local_batch_size=16, + global_batch_size=-1, + seq_len=1, + steps=STEPS, + max_norm=1.0, + dtype="float32", + mixed_precision_param="bfloat16", + mixed_precision_reduce="float32", + ), + parallelism=ParallelismConfig( + data_parallel_replicate_degree=num_nodes, + data_parallel_shard_degree=local_world_size, + ), + # plain CheckpointManager with onnx export off (path-specific PathOnnxCheckpointManager disabled) + checkpoint=CheckpointManager.Config(enable=False), + metrics=MetricsProcessor.Config( + log_freq=10, enable_reporterv2=True, save_freq=STEPS + ), + # path-specific driving validator disabled; dataloader is required by the dataclass but never + # built while enable=False (Trainer builds the validator only when validator.enable is True) + validator=PathValidator.Config( + enable=False, + steps=-1, + dataloader=_dataloader_config(split="val"), + mixed_precision_param="bfloat16", + ), + fps=SUPERCOMBO_FPS, + plan_target_last_frame=True, # ViT predicts a single-frame plan; supervise the last frame + debug=DebugConfig(seed=0), + ) + + +def vit_standard_w256() -> PathTrainer.Config: + return _vit("w256", mup=False) + + +def vit_standard_w512() -> PathTrainer.Config: + return _vit("w512", mup=False) + + +def vit_standard_w1024() -> PathTrainer.Config: + return _vit("w1024", mup=False) + + +def vit_standard_w2048() -> PathTrainer.Config: + return _vit("w2048", mup=False) + + +def vit_mup_w256() -> PathTrainer.Config: + return _vit("w256", mup=True) + + +def vit_mup_w512() -> PathTrainer.Config: + return _vit("w512", mup=True) + + +def vit_mup_w1024() -> PathTrainer.Config: + return _vit("w1024", mup=True) + + +def vit_mup_w2048() -> PathTrainer.Config: + return _vit("w2048", mup=True) From 9a235456871415af9eec262d2a1b09cb4ffefada Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Thu, 25 Jun 2026 15:36:21 -0700 Subject: [PATCH 08/28] path: remove standalone plan_vit experiment (moved into path/vit.py) --- torchtitan/experiments/plan_vit/__init__.py | 10 - .../experiments/plan_vit/config_registry.py | 326 ------------------ torchtitan/experiments/plan_vit/model.py | 257 -------------- torchtitan/experiments/plan_vit/trainer.py | 116 ------- 4 files changed, 709 deletions(-) delete mode 100644 torchtitan/experiments/plan_vit/__init__.py delete mode 100644 torchtitan/experiments/plan_vit/config_registry.py delete mode 100644 torchtitan/experiments/plan_vit/model.py delete mode 100644 torchtitan/experiments/plan_vit/trainer.py diff --git a/torchtitan/experiments/plan_vit/__init__.py b/torchtitan/experiments/plan_vit/__init__.py deleted file mode 100644 index be18ef68a2..0000000000 --- a/torchtitan/experiments/plan_vit/__init__.py +++ /dev/null @@ -1,10 +0,0 @@ -# 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 .config_registry import model_registry -from .model import parallelize_plan_vit, PlanViT - -__all__ = ["PlanViT", "model_registry", "parallelize_plan_vit"] diff --git a/torchtitan/experiments/plan_vit/config_registry.py b/torchtitan/experiments/plan_vit/config_registry.py deleted file mode 100644 index 6eb2bb5fc1..0000000000 --- a/torchtitan/experiments/plan_vit/config_registry.py +++ /dev/null @@ -1,326 +0,0 @@ -# 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. - -"""Config assembly + flavors for plan_vit, mirroring path/config_registry.py. - -Width flavors scale n_head at fixed head_dim=64 (the clean muP axis); base = w256. Two cameras are -channel-stacked into in_channels=24 (no VAE). The trainer-side config functions live below the model side. -""" - -from __future__ import annotations - -import math -import os -from functools import partial -from xx.ml_tools.constants.model import SUPERCOMBO_FPS - -import torch.nn as nn - -from torchtitan.components.checkpoint import CheckpointManager -from torchtitan.components.lr_scheduler import LRSchedulersContainer -from torchtitan.components.metrics import MetricsProcessor -from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig -from torchtitan.components.tokenizer import NoOpTokenizer -from torchtitan.config import DebugConfig, ParallelismConfig, TrainingConfig -from torchtitan.experiments.path.dataset import PathDataLoader -from torchtitan.experiments.path.loss import PathLoss -from torchtitan.models.common import Embedding, LayerNorm, Linear -from torchtitan.models.common.attention import ScaledDotProductAttention -from torchtitan.protocols.model_spec import ModelSpec -from .model import ( - parallelize_plan_vit, - PatchEmbed, - PlanHead, - PlanViT, - PlanViTAttention, - PlanViTBlock, - PlanViTMLP, -) -from .trainer import PlanViTTrainer - -_LINEAR_INIT = { - "weight": partial(nn.init.normal_, mean=0.0, std=0.02), - "bias": nn.init.zeros_, -} -_NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} - -HEAD_DIM = 64 -N_LAYER = 8 -INPUT_SIZE = ( - 1, - 128, - 256, -) # current frame; spatial ViT (temporal history is a later variant) -PATCH_SIZE = (1, 16, 8) -IN_CHANNELS = 24 # two cameras (IMG + BIG_IMG), 12 YUV channels each, channel-stacked -PLAN_SIZE = 15 * 33 * 2 # 990, laplacian mu+log-sigma -BASE_WIDTH = 256 -PLAN_VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} - - -def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Config: - return Linear.Config( - in_features=in_f, - out_features=out_f, - bias=bias, - param_init={ - "weight": partial(nn.init.normal_, mean=0.0, std=std), - "bias": nn.init.zeros_, - }, - ) - - -def _hidden_std(fan_in: int, *, mup: bool) -> float: - # muP shrinks hidden/output init to 1/sqrt(fan_in) so pre-activations stay O(1) as width grows; - # standard param holds the base-width variance 1/sqrt(BASE_WIDTH), so it fans out with width. - return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 - - -def _ln(dim: int) -> LayerNorm.Config: - return LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT) - - -def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: - return multiple_of * math.ceil(int(dim * mult) / multiple_of) - - -def _attention( - dim: int, n_head: int, *, mup: bool, qk_norm: bool = True -) -> PlanViTAttention.Config: - head_dim = dim // n_head - return PlanViTAttention.Config( - norm=_ln(dim), - q_norm=_ln(head_dim) if qk_norm else None, - k_norm=_ln(head_dim) if qk_norm else None, - c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), - c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * N_LAYER)), - inner_attention=ScaledDotProductAttention.Config(), - n_head=n_head, - head_dim=head_dim, - dropout=0.0, - ) - - -def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PlanViTMLP.Config: - hidden = _hidden(dim, mult) - return PlanViTMLP.Config( - norm=_ln(dim), - c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), - c_proj=_lin( - hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER) - ), - act="gelu_tanh", - dropout=0.0, - ) - - -def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: - n_embd = PLAN_VIT_WIDTHS[flavor] - n_head = n_embd // HEAD_DIM - pt, ph, pw = PATCH_SIZE - patch_dim = pt * IN_CHANNELS * ph * pw - t, h, w = INPUT_SIZE - num_patches = (t // pt) * (h // ph) * (w // pw) - return PlanViT.Config( - input_size=INPUT_SIZE, - patch_size=PATCH_SIZE, - in_channels=IN_CHANNELS, - n_embd=n_embd, - output_mult=(BASE_WIDTH / n_embd) - if mup - else 1.0, # muP readout fwd mult 1/m (init output slopes ~1/sqrt(m)) - patch_embed=PatchEmbed.Config( - proj=_lin( - patch_dim, n_embd, std=patch_dim**-0.5 - ), # input embed: width-independent - patch_size=PATCH_SIZE, - ), - pos_embedding=Embedding.Config( - num_embeddings=num_patches, embedding_dim=n_embd, param_init=_LINEAR_INIT - ), - blocks=[ - PlanViTBlock.Config( - attention=_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), - mlp=_mlp(n_embd, mup=mup), - ) - for _ in range(N_LAYER) - ], - norm=_ln(n_embd), - plan_head=PlanHead.Config( - norm=_ln(n_embd), - head=_lin( - n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5 - ), # muP readout: base-width init - ), - ) - - -def model_registry(flavor: str, *, mup: bool) -> ModelSpec: - return ModelSpec( - name="plan_vit", - flavor=flavor, - model=_model_config(flavor, mup=mup), - parallelize_fn=parallelize_plan_vit, - pipelining_fn=None, - post_optimizer_build_fn=None, - state_dict_adapter=None, - ) - - -STEPS = 512 # per-run step budget; override with training.steps=N on the CLI -# learning rate is the muTransfer sweep axis: one run per (flavor, lr); set with `-e PLAN_VIT_LR=...` -SWEEP_LR = float(os.getenv("PLAN_VIT_LR", "3e-4")) -# hidden matrix weights get muP lr eta/m; input embed, readout, norms, biases get base eta -# (readout is fan_in-infinite only, so Adam treats it vector-like -> base lr, not eta/m) -MUP_PATTERN = ( - r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" - r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" -) - - -def _si_int(value: str | int) -> int: - suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} - value = str(value).strip().lower() - return ( - int(float(value[:-1]) * suffixes[value[-1]]) - if value[-1] in suffixes - else int(value) - ) - - -def _dataloader_config(*, split: str) -> PathDataLoader.Config: - from xx.common.basedir import XX_BASEDIR - from xx.datasets.constants import BASE_DIR_GT_10M - from xx.training.path.config import DatasetConfig as XXPathDatasetConfig - - base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) - return PathDataLoader.Config( - # prune-10M study data: a seeded random 10k sample of the 10M store (training_2026_02) - dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), - split=split, - shuffle_size=_si_int(base.shuffle_size), - min_mixing=base.min_mixing, - num_writers=base.num_writers, - num_readers=base.num_readers, - fps=base.fps, - pipeline_dir=BASE_DIR_GT_10M, # the 10M store, not the 2.5M big-train list - plan_only=base.plan_only, - limit=base.limit, - n_frames=base.n_frames, - rgb=base.rgb, - unvision=base.unvision, - ) - - -def _optimizer_config( - flavor: str, *, mup: bool, lr: float, wd: float -) -> OptimizersContainer.Config: - m = PLAN_VIT_WIDTHS[flavor] / BASE_WIDTH - common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} - if mup: - groups = [ - ParamGroupConfig( - pattern=MUP_PATTERN, - optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr / m}, - ), - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr}, - ), - ] - else: - groups = [ - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr}, - ) - ] - return OptimizersContainer.Config( - implementation="fused_opt_states_bf16", param_groups=groups - ) - - -def _plan_vit( - flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 -) -> PlanViTTrainer.Config: - # derive data parallelism from the launch (like path), so any N nodes x GPUs validate - local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) - world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int( - os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) - ) - return PlanViTTrainer.Config( - loss=PathLoss.Config(), - model_spec=model_registry(flavor, mup=mup), - tokenizer=NoOpTokenizer.Config(), - dataloader=_dataloader_config(split="train"), - optimizer=_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), - lr_scheduler=LRSchedulersContainer.Config( - warmup_steps=round(STEPS * 0.1), - total_steps=None, # use the real training.steps; a fixed value wraps the cosine on longer runs - decay_ratio=0.8, - decay_type="cosine", - min_lr_factor=0.0, - ), - training=TrainingConfig( - local_batch_size=16, - global_batch_size=-1, - seq_len=1, - steps=STEPS, - max_norm=1.0, - dtype="float32", - mixed_precision_param="bfloat16", - mixed_precision_reduce="float32", - ), - parallelism=ParallelismConfig( - data_parallel_replicate_degree=num_nodes, - data_parallel_shard_degree=local_world_size, - ), - checkpoint=CheckpointManager.Config(enable=False), - metrics=MetricsProcessor.Config( - log_freq=10, enable_reporterv2=True, save_freq=STEPS - ), - debug=DebugConfig(seed=0), - ) - - -def plan_vit_standard_w256() -> PlanViTTrainer.Config: - return _plan_vit("w256", mup=False) - - -def plan_vit_standard_w512() -> PlanViTTrainer.Config: - return _plan_vit("w512", mup=False) - - -def plan_vit_standard_w1024() -> PlanViTTrainer.Config: - return _plan_vit("w1024", mup=False) - - -def plan_vit_standard_w2048() -> PlanViTTrainer.Config: - return _plan_vit("w2048", mup=False) - - -def plan_vit_mup_w256() -> PlanViTTrainer.Config: - return _plan_vit("w256", mup=True) - - -def plan_vit_mup_w512() -> PlanViTTrainer.Config: - return _plan_vit("w512", mup=True) - - -def plan_vit_mup_w1024() -> PlanViTTrainer.Config: - return _plan_vit("w1024", mup=True) - - -def plan_vit_mup_w2048() -> PlanViTTrainer.Config: - return _plan_vit("w2048", mup=True) - - -def plan_vit() -> PlanViTTrainer.Config: - return plan_vit_mup_w256() diff --git a/torchtitan/experiments/plan_vit/model.py b/torchtitan/experiments/plan_vit/model.py deleted file mode 100644 index a5cd17facf..0000000000 --- a/torchtitan/experiments/plan_vit/model.py +++ /dev/null @@ -1,257 +0,0 @@ -# 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. - -"""Plan ViT: raw camera frames -> patches -> transformer -> plan. NO VAE. - -A self-contained planning model for the muP + scaling study, built from torchtitan.models.common -blocks the same way path/model.py is. Scales cleanly by width (n_embd / n_head) for muTransfer. -""" - -from __future__ import annotations - -from dataclasses import dataclass -from xx.ml_tools.constants.model import ModelInputs - -import torch -import torch.nn as nn -from einops import rearrange -from torch.distributed.device_mesh import DeviceMesh -from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy - -from torchtitan.config import ( - CompileConfig, - ParallelismConfig, - TORCH_DTYPE_MAP, - TrainingConfig, -) -from torchtitan.distributed import ParallelDims -from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig -from torchtitan.models.common import Embedding, LayerNorm, Linear, RMSNorm -from torchtitan.models.common.attention import ScaledDotProductAttention -from torchtitan.protocols.model import BaseModel -from torchtitan.protocols.module import Module, ModuleList -from torchtitan.tools.logging import logger - - -class PlanViTMLP(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - norm: LayerNorm.Config | RMSNorm.Config - c_fc: Linear.Config - c_proj: Linear.Config - act: str - dropout: float - - def __init__(self, config: Config): - super().__init__() - self.norm = config.norm.build() - self.c_fc = config.c_fc.build() - self.act = ( - nn.GELU(approximate="tanh") if config.act == "gelu_tanh" else nn.GELU() - ) - self.c_proj = config.c_proj.build() - self.dropout = nn.Dropout(config.dropout) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.dropout(self.c_proj(self.act(self.c_fc(self.norm(x))))) - - -class PlanViTAttention(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - norm: LayerNorm.Config | RMSNorm.Config - q_norm: LayerNorm.Config | RMSNorm.Config | None - k_norm: LayerNorm.Config | RMSNorm.Config | None - c_attn: Linear.Config - c_proj: Linear.Config - inner_attention: ScaledDotProductAttention.Config - n_head: int - head_dim: int - dropout: float - - def __init__(self, config: Config): - super().__init__() - self.n_head = config.n_head - self.head_dim = config.head_dim - self.norm = config.norm.build() - self.q_norm = ( - config.q_norm.build() if config.q_norm is not None else nn.Identity() - ) - self.k_norm = ( - config.k_norm.build() if config.k_norm is not None else nn.Identity() - ) - self.c_attn = config.c_attn.build() - self.c_proj = config.c_proj.build() - self.inner_attention = config.inner_attention.build() - self.dropout = nn.Dropout(config.dropout) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - b, t, _ = x.shape - qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) - q, k, v = qkv.unbind(2) - q, k = self.q_norm(q), self.k_norm(k) - x = self.inner_attention( - q, k, v, is_causal=False - ) # ViT: bidirectional over patches - return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) - - -class PlanViTBlock(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - attention: PlanViTAttention.Config - mlp: PlanViTMLP.Config - - def __init__(self, config: Config): - super().__init__() - self.attention = config.attention.build() - self.mlp = config.mlp.build() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = x + self.attention(x) - return x + self.mlp(x) - - -class PatchEmbed(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - proj: Linear.Config - patch_size: tuple[int, int, int] # (pt, ph, pw) - - def __init__(self, config: Config): - super().__init__() - self.patch_size = config.patch_size - self.proj = config.proj.build() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - # x: (B, T, C, H, W) raw frames -> (B, num_patches, patch_dim) -> (B, num_patches, n_embd) - pt, ph, pw = self.patch_size - x = rearrange( - x, "b (t pt) c (h ph) (w pw) -> b (t h w) (pt c ph pw)", pt=pt, ph=ph, pw=pw - ) - return self.proj( - x.to(self.proj.weight.dtype) - ) # match the bf16 (mp) weights, like path's vision - - -class PlanHead(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - norm: LayerNorm.Config | RMSNorm.Config - head: Linear.Config - - def __init__(self, config: Config): - super().__init__() - self.norm = config.norm.build() - self.head = config.head.build() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.head(self.norm(x)) - - -class PlanViT(BaseModel): - @dataclass(kw_only=True, slots=True) - class Config(BaseModel.Config): - input_size: tuple[int, int, int] # (n_frames, H, W) - patch_size: tuple[int, int, int] - in_channels: int - n_embd: int - output_mult: float # muP readout multiplier 1/m (m = n_embd / base); 1.0 for standard param - patch_embed: PatchEmbed.Config - pos_embedding: Embedding.Config - blocks: list[PlanViTBlock.Config] - norm: LayerNorm.Config | RMSNorm.Config - plan_head: PlanHead.Config - - @property - def num_patches(self) -> int: - t, h, w = self.input_size - pt, ph, pw = self.patch_size - return (t // pt) * (h // ph) * (w // pw) - - def update_from_config(self, *, config, **kwargs) -> None: - parallelism = config.parallelism - for name, degree in { - "tensor parallel": parallelism.tensor_parallel_degree, - "context parallel": parallelism.context_parallel_degree, - "pipeline parallel": parallelism.pipeline_parallel_degree, - "expert parallel": parallelism.expert_parallel_degree, - }.items(): - if degree > 1: - raise ValueError(f"plan_vit does not support {name}") - - def get_nparams_and_flops(self, model: Module, seq_len: int) -> tuple[int, int]: - nparams = sum(p.numel() for p in model.parameters()) - return nparams, 6 * nparams - - def __init__(self, config: Config): - super().__init__() - self.config = config - self.patch_embed = config.patch_embed.build() - self.pos_embedding = config.pos_embedding.build() - self.blocks = ModuleList([block.build() for block in config.blocks]) - self.norm = config.norm.build() - self.plan_head = config.plan_head.build() - - def verify_module_protocol(self) -> None: - pass # nn.Dropout/GELU/Identity are plain nn.Module, like path - - def _frames(self, inputs: dict[str, torch.Tensor] | torch.Tensor) -> torch.Tensor: - # production input: two cameras IMG, BIG_IMG, each (B, T, 12, H, W) YUV. Take the current frame of each, - # channel-stack -> (B, 1, 24, H, W). NO VAE. A plain tensor (testing) is passed through unchanged. - if isinstance(inputs, torch.Tensor): - return inputs - img, big = inputs[ModelInputs.IMG], inputs[ModelInputs.BIG_IMG] - frame = torch.cat([img[:, -1], big[:, -1]], dim=1).unsqueeze(1) - return ( - frame.float() - 127.5 - ) / 63.75 # uint8 YUV -> normalized float (mean 255/2, std 255/4 like path) - - def forward( - self, inputs: dict[str, torch.Tensor] | torch.Tensor - ) -> dict[str, torch.Tensor]: - x = self.patch_embed(self._frames(inputs)) - pos = self.pos_embedding(torch.arange(x.shape[1], device=x.device)) - x = x + rearrange(pos, "t c -> () t c") - for block in self.blocks: - x = block(x) - x = self.norm(x) - # global-pool the patches -> plan; the muP readout multiplier keeps the output width-stable - return {"plan": self.plan_head(x.mean(dim=1)) * self.config.output_mult} - - -def parallelize_plan_vit( - model: PlanViT, - *, - parallel_dims: ParallelDims, - training: TrainingConfig, - parallelism: ParallelismConfig, - compile_config: CompileConfig, - ac_config: ActivationCheckpointingConfig, - dump_folder: str, -) -> PlanViT: - if ( - parallel_dims.tp_enabled - or parallel_dims.cp_enabled - or parallel_dims.pp_enabled - or parallel_dims.ep_enabled - ): - raise ValueError("plan_vit supports data parallelism only") - names = ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] - dp_mesh: DeviceMesh = parallel_dims.get_mesh(names) - mp_policy = MixedPrecisionPolicy( - param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param], - reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], - cast_forward_inputs=True, - ) - fsdp_config = {"mesh": dp_mesh, "mp_policy": mp_policy} - for idx, block in enumerate(model.blocks): - fully_shard( - block, **fsdp_config, reshard_after_forward=(idx < len(model.blocks) - 1) - ) - fully_shard(model, **fsdp_config) - logger.info("Applied FSDP to plan_vit") - return model diff --git a/torchtitan/experiments/plan_vit/trainer.py b/torchtitan/experiments/plan_vit/trainer.py deleted file mode 100644 index 1f536f033f..0000000000 --- a/torchtitan/experiments/plan_vit/trainer.py +++ /dev/null @@ -1,116 +0,0 @@ -# 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. - -"""Thin trainer for plan_vit: model(inputs_dict) -> pred dict, PathLoss(pred, targets), backward. - -Mirrors PathTrainer's generic core without the path-specific ONNX export / driving validator / reports a -scaling study doesn't need. loss + dataloader come from the base Trainer.Config (set in config_registry). -""" - -from __future__ import annotations - -import time -from collections.abc import Iterable, Iterator -from dataclasses import dataclass - -import torch - -from torchtitan.components.dataloader import DataloaderExhaustedError -from torchtitan.distributed import utils as dist_utils -from torchtitan.observability import structured_logger as sl -from torchtitan.trainer import Trainer - - -class PlanViTTrainer(Trainer): - @dataclass(kw_only=True, slots=True) - class Config(Trainer.Config): - pass - - def __init__(self, config: "PlanViTTrainer.Config"): - self.ntokens_seen = 0 - self._metrics: dict[str, torch.Tensor] = {} - super().__init__(config) - self.loss_fn.to(self.device) - - def batch_generator( - self, - data_iterable: Iterable[ - tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] - ], - ) -> Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]]: - data_iterator = iter(data_iterable) - while True: - t0 = time.perf_counter() - try: - input_dict, targets = next(data_iterator) - except StopIteration as ex: - raise DataloaderExhaustedError() from ex - self.metrics_processor.ntokens_since_last_log += next( - iter(input_dict.values()) - ).shape[0] - self.metrics_processor.data_loading_times.append(time.perf_counter() - t0) - yield input_dict, targets - - @sl.log_trace_span("fwd_bwd") - def forward_backward_step( - self, *, input_dict: dict[str, torch.Tensor], labels: dict[str, torch.Tensor] - ) -> torch.Tensor: - assert len(self.model_parts) == 1 - with self.train_context(): - pred = self.model_parts[0](input_dict) - # plan_vit is single-frame: it predicts the current (last) frame's plan, so supervise - # against the last temporal position of the dense target (path trains all positions). - labels = {**labels, "plan": labels["plan"][:, -1]} - loss_vec, metrics = self.loss_fn(pred, labels) - loss = loss_vec.mean() - self._metrics = metrics - loss.backward() - return loss - - def train_step( - self, - data_iterator: Iterator[ - tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] - ], - ) -> None: - self.optimizers.zero_grad() - lr = self.lr_schedulers.schedulers[0].get_last_lr()[0] - - input_dict, targets = next(data_iterator) - input_dict = {k: v.to(self.device) for k, v in input_dict.items()} - targets = {k: v.to(self.device) for k, v in targets.items()} - self.ntokens_seen += next(iter(input_dict.values())).shape[0] - loss = self.forward_backward_step(input_dict=input_dict, labels=targets) - - grad_norm = dist_utils.clip_grad_norm_( - [p for m in self.model_parts for p in m.parameters()], - self.config.training.max_norm, - foreach=True, - pp_mesh=self.parallel_dims.get_optional_mesh("pp"), - ep_enabled=self.parallel_dims.ep_enabled, - ) - self.checkpointer.maybe_wait_for_staging() - self.optimizers.step() - self.lr_schedulers.step() - - if not self.metrics_processor.should_log(self.step): - return - - local_loss = loss.detach() - if self.parallel_dims.dp_cp_enabled: - loss_mesh = self.parallel_dims.get_optional_mesh("loss") - global_avg_loss = dist_utils.dist_mean(local_loss, loss_mesh) - global_max_loss = dist_utils.dist_max(local_loss, loss_mesh) - else: - global_avg_loss = global_max_loss = float(local_loss.item()) - - self.metrics_processor.log( - self.step, - global_avg_loss, - global_max_loss, - float(grad_norm.item()), - extra_metrics={"metrics/lr/": lr}, - ) From 1baf6c616c6656783c595572b3140702194ad07a Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Thu, 25 Jun 2026 18:52:10 -0700 Subject: [PATCH 09/28] path: trim the vit experiment to a minimal surface drop the Meta copyright headers, the inherited-Meta formatter churn, and the dangling plan_vit registry entry; the vit resolves via --module path --config vit_*. --- torchtitan/experiments/__init__.py | 1 - .../experiments/path/config_registry.py | 165 +++++------------- torchtitan/experiments/path/trainer.py | 95 +++------- torchtitan/experiments/path/vit.py | 19 +- .../experiments/path/vit_config_registry.py | 23 +-- 5 files changed, 84 insertions(+), 219 deletions(-) diff --git a/torchtitan/experiments/__init__.py b/torchtitan/experiments/__init__.py index da99bc24be..530a90ba40 100644 --- a/torchtitan/experiments/__init__.py +++ b/torchtitan/experiments/__init__.py @@ -13,7 +13,6 @@ "autoparallel.llama3", "autoparallel.local_map_deepseek_v3", "path", - "plan_vit", "worldmodel", "torchft.llama3", "rl", diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index dc716b545d..6bf627063f 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -1,31 +1,8 @@ -# 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 import os from functools import partial -from xx.datasets.helpers import DEFAULT_BIG_TRAIN_LIST -from xx.ml_tools.constants.model import ( - frame_constants_from_fps, - FRAME_TYPE, - INPUT_FRAMES_NAMES, - ModelInputs, - N_FRAMES, - SUPERCOMBO_FPS, - TEMPORAL_INPUTS, -) -from xx.training.path.config import DatasetConfig as XXPathDatasetConfig -from xx.training.path.hydra_configs import ( - DRIVING_HEADS, - META_HEADS, - POSE_HEADS, - TEMPORAL_META_HEADS, -) import torch.nn as nn @@ -43,13 +20,23 @@ from torchtitan.models.common import Embedding, LayerNorm, Linear from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model_spec import ModelSpec +from xx.datasets.helpers import DEFAULT_BIG_TRAIN_LIST +from xx.ml_tools.constants.model import ( + SUPERCOMBO_FPS, + FRAME_TYPE, + INPUT_FRAMES_NAMES, + N_FRAMES, + TEMPORAL_INPUTS, + ModelInputs, + frame_constants_from_fps, +) +from xx.training.path.config import DatasetConfig as XXPathDatasetConfig +from xx.training.path.hydra_configs import DRIVING_HEADS, META_HEADS, POSE_HEADS, TEMPORAL_META_HEADS from .dataset import PathDataLoader -from .loss import PathLoss from .model import ( Hydra, LinearEncoder, - parallelize_path, PathHead, PathMLP, PathModel, @@ -62,29 +49,28 @@ TemporalPolicy, TemporalSummarizer, Vision, + parallelize_path, ) +from .loss import PathLoss from .onnx_checkpoint import PathOnnxCheckpointManager from .trainer import PathTrainer from .validate import PathValidator -# Plan ViT flavors ride PathTrainer too; re-exported so `--module path --config vit_*` resolves here, -# the same way convnext_* do (the config manager looks the name up on this module). +# Path ViT flavors ride PathTrainer too; re-exported so `--module path --config vit_*` resolves +# here, the same way convnext_* do (the config manager looks the name up on this module). from .vit_config_registry import ( # noqa: F401 - vit_mup_w1024, - vit_mup_w2048, vit_mup_w256, vit_mup_w512, - vit_standard_w1024, - vit_standard_w2048, + vit_mup_w1024, + vit_mup_w2048, vit_standard_w256, vit_standard_w512, + vit_standard_w1024, + vit_standard_w2048, ) -_LINEAR_INIT = { - "weight": partial(nn.init.normal_, mean=0.0, std=0.02), - "bias": nn.init.zeros_, -} +_LINEAR_INIT = {"weight": partial(nn.init.normal_, mean=0.0, std=0.02), "bias": nn.init.zeros_} _NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} @@ -117,10 +103,10 @@ def convnext_xxlarge() -> PathTrainer.Config: def _path(flavor: str) -> PathTrainer.Config: - steps = 1024 * 100 + steps = 1024*100 validation_freq = 1024 reports = { - name: [validation_freq, steps // 2, steps] + name: [validation_freq, steps //2 , steps] for name in ( "analyse_driving", "analyse_lat_no_noise", @@ -134,14 +120,10 @@ def _path(flavor: str) -> PathTrainer.Config: mixed_precision_param = "bfloat16" local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int( - os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) - ) + num_nodes = int(os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size))) reporterv2_host = os.getenv("REPORTERV2_HOST") reporterv2_training_id = os.getenv("REPORTERV2_TRAINING_ID") - checkpoint_base_folder = ( - f"{reporterv2_host.rstrip('/')}/checkpoint" if reporterv2_host else "" - ) + checkpoint_base_folder = f"{reporterv2_host.rstrip('/')}/checkpoint" if reporterv2_host else "" fps = SUPERCOMBO_FPS plan_only = False return PathTrainer.Config( @@ -184,9 +166,7 @@ def _path(flavor: str) -> PathTrainer.Config: fps=fps, activation_checkpoint=FullAC.Config(), compile=CompileConfig(enable=True, components=["model"]), - metrics=MetricsProcessor.Config( - log_freq=16, enable_reporterv2=True, save_freq=validation_freq - ), + metrics=MetricsProcessor.Config(log_freq=16, enable_reporterv2=True, save_freq=validation_freq), validator=PathValidator.Config( enable=True, freq=validation_freq, @@ -204,12 +184,8 @@ def _model_config(flavor: str) -> PathModel.Config: n_frames_input = N_FRAMES input_frame_names = INPUT_FRAMES_NAMES input_frame_type = FRAME_TYPE - frame_constants = frame_constants_from_fps( - n_frames=n_frames_input, frame_type=input_frame_type - ) - in_channels = sum( - frame_constants["frame_shapes"][name][0] for name in input_frame_names - ) + frame_constants = frame_constants_from_fps(n_frames=n_frames_input, frame_type=input_frame_type) + in_channels = sum(frame_constants["frame_shapes"][name][0] for name in input_frame_names) block_size = len(frame_constants["history_idxs"]) temporal_len = frame_constants["temporal_len"] dim = vision_features @@ -239,13 +215,9 @@ def _model_config(flavor: str) -> PathModel.Config: temporal_summarizer=TemporalSummarizer.Config( mlp1=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), mlp2=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - desire_encoder=_encoder( - TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * temporal_len, dim - ), + desire_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * temporal_len, dim), traffic_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.TRAFFIC][0], dim), - action_t_encoder=_encoder( - TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], dim - ), + action_t_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], dim), transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( @@ -255,25 +227,17 @@ def _model_config(flavor: str) -> PathModel.Config: for _ in range(4) ] ), - pos_embedding=Embedding.Config( - num_embeddings=block_size, - embedding_dim=dim, - param_init=_LINEAR_INIT, - ), + pos_embedding=Embedding.Config(num_embeddings=block_size, embedding_dim=dim, param_init=_LINEAR_INIT), block_size=block_size, dense_training_outputs=True, ), - temporal_hydra=_hydra( - _heads(DRIVING_HEADS + TEMPORAL_META_HEADS), in_features=dim, mlp_mult=2 - ), + temporal_hydra=_hydra(_heads(DRIVING_HEADS + TEMPORAL_META_HEADS), in_features=dim, mlp_mult=2), history_idxs=tuple(int(x) for x in frame_constants["history_idxs"]), ), ) -def _dataloader_config( - *, split: str, fps: int, plan_only: bool -) -> PathDataLoader.Config: +def _dataloader_config(*, split: str, fps: int, plan_only: bool) -> PathDataLoader.Config: base = XXPathDatasetConfig(fps=fps, plan_only=plan_only) return PathDataLoader.Config( dataset=DEFAULT_BIG_TRAIN_LIST, @@ -292,9 +256,7 @@ def _dataloader_config( ) -def _checkpoint_config( - folder: str, base_folder: str, interval: int -) -> PathOnnxCheckpointManager.Config: +def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnxCheckpointManager.Config: frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) temporal_len = frame_constants["temporal_len"] vision_input_names = [ModelInputs.IMG, ModelInputs.BIG_IMG] @@ -317,10 +279,10 @@ def _checkpoint_config( [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.ACTION_T][0]], ] return PathOnnxCheckpointManager.Config( - keep_latest_k=0, # keep all checkpoints + keep_latest_k=0, # keep all checkpoints enable=True, checkpoint_base_folder=base_folder, - save_model_state_dict=True, # another copy of full state dict + save_model_state_dict=True, # another copy of full state dict export_onnx=True, enable_first_step_checkpoint=True, folder=folder, @@ -328,7 +290,7 @@ def _checkpoint_config( input_names=input_names, input_shapes=input_shapes, input_dtypes=["float16"] * len(input_names), - onnx_model_dtype="float16", # WIP: test if fp16 doesn't degrade performance + onnx_model_dtype="float16", # WIP: test if fp16 doesn't degrade performance vision_input_names=vision_input_names, temporal_policy_input_names=temporal_policy_input_names, ) @@ -337,11 +299,7 @@ def _checkpoint_config( def _si_int(value: str | int) -> int: suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} value = str(value).strip().lower() - return ( - int(float(value[:-1]) * suffixes[value[-1]]) - if value[-1] in suffixes - else int(value) - ) + return int(float(value[:-1]) * suffixes[value[-1]]) if value[-1] in suffixes else int(value) def _optimizer_config() -> OptimizersContainer.Config: @@ -365,9 +323,7 @@ def _optimizer_config() -> OptimizersContainer.Config: def _heads(heads) -> tuple[PathHead, ...]: - return tuple( - PathHead(head.name, head.output_size, head.mlp, head.scale) for head in heads - ) + return tuple(PathHead(head.name, head.output_size, head.mlp, head.scale) for head in heads) def _hidden_dim(dim: int, mlp_mult: float, multiple_of: int = 256) -> int: @@ -379,12 +335,8 @@ def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Co hidden = _hidden_dim(dim, mlp_mult) return PathMLP.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - c_fc=Linear.Config( - in_features=dim, out_features=hidden, bias=bias, param_init=_LINEAR_INIT - ), - c_proj=Linear.Config( - in_features=hidden, out_features=dim, bias=bias, param_init=_LINEAR_INIT - ), + c_fc=Linear.Config(in_features=dim, out_features=hidden, bias=bias, param_init=_LINEAR_INIT), + c_proj=Linear.Config(in_features=hidden, out_features=dim, bias=bias, param_init=_LINEAR_INIT), act="gelu_tanh", dropout=dropout, ) @@ -392,15 +344,8 @@ def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Co def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: return LinearEncoder.Config( - in_layer=Linear.Config( - in_features=in_features, - out_features=dim, - bias=True, - param_init=_LINEAR_INIT, - ), - out_layer=Linear.Config( - in_features=dim, out_features=dim, bias=False, param_init=_LINEAR_INIT - ), + in_layer=Linear.Config(in_features=in_features, out_features=dim, bias=True, param_init=_LINEAR_INIT), + out_layer=Linear.Config(in_features=dim, out_features=dim, bias=False, param_init=_LINEAR_INIT), ) @@ -410,12 +355,8 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), - c_attn=Linear.Config( - in_features=dim, out_features=3 * dim, bias=True, param_init=_LINEAR_INIT - ), - c_proj=Linear.Config( - in_features=dim, out_features=dim, bias=True, param_init=_LINEAR_INIT - ), + c_attn=Linear.Config(in_features=dim, out_features=3 * dim, bias=True, param_init=_LINEAR_INIT), + c_proj=Linear.Config(in_features=dim, out_features=dim, bias=True, param_init=_LINEAR_INIT), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, head_dim=head_dim, @@ -423,16 +364,10 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co ) -def _hydra( - heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float -) -> Hydra.Config: +def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> Hydra.Config: return Hydra.Config( heads=heads, - head_mlps={ - head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) - for head in heads - if head.mlp - }, + head_mlps={head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) for head in heads if head.mlp}, final_layers={ head.name: Linear.Config( in_features=in_features, @@ -442,9 +377,5 @@ def _hydra( ) for head in heads }, - scale_layers={ - head.name: ScaleLayer.Config(n_features=head.output_size) - for head in heads - if head.scale - }, + scale_layers={head.name: ScaleLayer.Config(n_features=head.output_size) for head in heads if head.scale}, ) diff --git a/torchtitan/experiments/path/trainer.py b/torchtitan/experiments/path/trainer.py index d7c06bbb78..8ea74ba9ad 100644 --- a/torchtitan/experiments/path/trainer.py +++ b/torchtitan/experiments/path/trainer.py @@ -28,31 +28,24 @@ class Config(Trainer.Config): miniray: dict[str, Any] = field(default_factory=dict) fps: int # single-frame plan supervision: slice the dense plan target to the last frame before the - # loss. off by default (dense; convnext/worldmodel unchanged); on for the single-frame plan_vit. + # loss. off by default (dense; convnext/worldmodel unchanged); on for the single-frame path vit. plan_target_last_frame: bool = False def __post_init__(self) -> None: Trainer.Config.__post_init__(self) if self.codedir: self.miniray = {**self.miniray, "codedir": self.codedir} - self.validator.miniray = { - **self.validator.miniray, - "codedir": self.codedir, - } + self.validator.miniray = {**self.validator.miniray, "codedir": self.codedir} def __init__(self, config: Config): super().__init__(config) training_id = os.getenv("REPORTERV2_TRAINING_ID") or "local" - self.unique_segment_counter = StringUniqueCounter( - f"unique_ids:{training_id}:path:train" - ) + self.unique_segment_counter = StringUniqueCounter(f"unique_ids:{training_id}:path:train") self.loss_fn.to(self.device) def batch_generator( self, - data_iterable: Iterable[ - tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] - ], + data_iterable: Iterable[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]], ) -> Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]]: data_iterator = iter(data_iterable) while True: @@ -61,12 +54,8 @@ def batch_generator( input_dict, targets = next(data_iterator) except StopIteration as ex: raise DataloaderExhaustedError() from ex - self.metrics_processor.ntokens_since_last_log += next( - iter(input_dict.values()) - ).shape[0] - self.metrics_processor.data_loading_times.append( - time.perf_counter() - data_load_start - ) + self.metrics_processor.ntokens_since_last_log += next(iter(input_dict.values())).shape[0] + self.metrics_processor.data_loading_times.append(time.perf_counter() - data_load_start) yield input_dict, targets @sl.log_trace_span("post_dataloading_process") @@ -91,7 +80,7 @@ def forward_backward_step( with self.train_context(): pred = self.model_parts[0](inputs) if self.config.plan_target_last_frame: - # single-frame plan models (plan_vit) predict the last frame's plan; supervise that frame + # single-frame plan models predict the last frame's plan; supervise that frame labels = {**labels, "plan": labels["plan"][:, -1]} loss_vec, metrics = self.loss_fn(pred, labels) loss = loss_vec.sum() / local_samples @@ -101,9 +90,7 @@ def forward_backward_step( def train_step( self, - data_iterator: Iterator[ - tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]] - ], + data_iterator: Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]], ) -> None: self.optimizers.zero_grad() lr_metrics = self.lr_schedulers.get_metrics() @@ -118,9 +105,7 @@ def train_step( input_dict, targets = next(data_iterator) local_samples += next(iter(input_dict.values())).shape[0] if "info" in input_dict: - step_segment_names.update( - segment_names_from_info(input_dict["info"]) - ) + step_segment_names.update(segment_names_from_info(input_dict["info"])) microbatches.append((input_dict, targets)) sl.log_trace_scalar({"local_samples": int(local_samples)}) @@ -129,9 +114,7 @@ def train_step( global_samples = dist_utils.dist_sum(local_samples, batch_mesh) else: global_samples = local_samples.float() - global_samples = torch.as_tensor( - global_samples, dtype=torch.float32, device=self.device - ) + global_samples = torch.as_tensor(global_samples, dtype=torch.float32, device=self.device) global_samples_value = float(global_samples.item()) accumulated_losses = [] @@ -148,10 +131,7 @@ def train_step( for name, value in metrics.items(): if name == "loss": continue - metric_sums[name] = ( - metric_sums.get(name, torch.zeros((), device=self.device)) - + value.float().sum() - ) + metric_sums[name] = metric_sums.get(name, torch.zeros((), device=self.device)) + value.float().sum() with sl.log_trace_span("optim"): grad_norm = dist_utils.clip_grad_norm_( @@ -175,30 +155,17 @@ def train_step( loss_mesh = parallel_dims.get_optional_mesh("loss") local_loss_sum = loss * local_samples global_avg_loss, global_max_loss, global_samples_seen = ( - dist_utils.dist_sum(local_loss_sum.detach(), loss_mesh) - / global_samples_value, + dist_utils.dist_sum(local_loss_sum.detach(), loss_mesh) / global_samples_value, dist_utils.dist_max(loss.detach(), loss_mesh), - dist_utils.dist_sum( - torch.tensor( - self.ntokens_seen, dtype=torch.int64, device=self.device - ), - loss_mesh, - ), + dist_utils.dist_sum(torch.tensor(self.ntokens_seen, dtype=torch.int64, device=self.device), loss_mesh), ) - metric_sums = { - k: dist_utils.dist_sum(v, loss_mesh) for k, v in metric_sums.items() - } + metric_sums = {k: dist_utils.dist_sum(v, loss_mesh) for k, v in metric_sums.items()} else: global_avg_loss = global_max_loss = float(loss.detach().item()) global_samples_seen = self.ntokens_seen path_metrics = { - f"path/{k}": float( - ( - torch.as_tensor(v, dtype=torch.float32, device=self.device) - / global_samples - ).item() - ) + f"path/{k}": float((torch.as_tensor(v, dtype=torch.float32, device=self.device) / global_samples).item()) for k, v in metric_sums.items() } unique_segments_seen = ( @@ -209,12 +176,7 @@ def train_step( dataset_metrics = { "dataset/unique_segments_seen": unique_segments_seen, } - extra_metrics = { - "n_samples_seen": global_samples_seen, - **lr_metrics, - **path_metrics, - **dataset_metrics, - } + extra_metrics = {"n_samples_seen": global_samples_seen, **lr_metrics, **path_metrics, **dataset_metrics} self.metrics_processor.log( self.step, global_avg_loss, @@ -232,28 +194,15 @@ def close(self) -> None: def state_dict(self) -> dict[str, Any]: state = super().state_dict() state["unique_segment_counter"] = self.unique_segment_counter.state_dict() - validator_unique_segment_counter = getattr( - getattr(self, "validator", None), "unique_segment_counter", None - ) + validator_unique_segment_counter = getattr(getattr(self, "validator", None), "unique_segment_counter", None) if validator_unique_segment_counter is not None: - state[ - "validation_unique_segment_counter" - ] = validator_unique_segment_counter.state_dict() + state["validation_unique_segment_counter"] = validator_unique_segment_counter.state_dict() return state def load_state_dict(self, state_dict: dict[str, Any]) -> None: super().load_state_dict(state_dict) if "unique_segment_counter" in state_dict: - self.unique_segment_counter.load_state_dict( - state_dict["unique_segment_counter"] - ) - validator_unique_segment_counter = getattr( - getattr(self, "validator", None), "unique_segment_counter", None - ) - if ( - validator_unique_segment_counter is not None - and "validation_unique_segment_counter" in state_dict - ): - validator_unique_segment_counter.load_state_dict( - state_dict["validation_unique_segment_counter"] - ) + self.unique_segment_counter.load_state_dict(state_dict["unique_segment_counter"]) + validator_unique_segment_counter = getattr(getattr(self, "validator", None), "unique_segment_counter", None) + if validator_unique_segment_counter is not None and "validation_unique_segment_counter" in state_dict: + validator_unique_segment_counter.load_state_dict(state_dict["validation_unique_segment_counter"]) diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py index 1d3bc3fed7..3a1b774063 100644 --- a/torchtitan/experiments/path/vit.py +++ b/torchtitan/experiments/path/vit.py @@ -1,15 +1,8 @@ -# 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. - """Plan ViT for the path experiment: raw camera frames -> patches -> transformer -> plan. NO VAE. -Ported verbatim (behavior-identical) from experiments/plan_vit/model.py so the ViT can ride -PathTrainer via config instead of the standalone PlanViTTrainer. A self-contained planning model -for the muP + scaling study, built from torchtitan.models.common blocks the same way path/model.py -is. Scales cleanly by width (n_embd / n_head) for muTransfer. +Rides PathTrainer via config. A self-contained planning model for the muP + scaling study, built +from torchtitan.models.common blocks the same way path/model.py is. Scales cleanly by width +(n_embd / n_head) for muTransfer. """ from __future__ import annotations @@ -183,7 +176,7 @@ def update_from_config(self, *, config, **kwargs) -> None: "expert parallel": parallelism.expert_parallel_degree, }.items(): if degree > 1: - raise ValueError(f"plan_vit does not support {name}") + raise ValueError(f"PlanViT does not support {name}") def get_nparams_and_flops(self, model: Module, seq_len: int) -> tuple[int, int]: nparams = sum(p.numel() for p in model.parameters()) @@ -241,7 +234,7 @@ def parallelize_vit( or parallel_dims.pp_enabled or parallel_dims.ep_enabled ): - raise ValueError("plan_vit supports data parallelism only") + raise ValueError("PlanViT supports data parallelism only") names = ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] dp_mesh: DeviceMesh = parallel_dims.get_mesh(names) mp_policy = MixedPrecisionPolicy( @@ -255,5 +248,5 @@ def parallelize_vit( block, **fsdp_config, reshard_after_forward=(idx < len(model.blocks) - 1) ) fully_shard(model, **fsdp_config) - logger.info("Applied FSDP to plan_vit") + logger.info("Applied FSDP to PlanViT") return model diff --git a/torchtitan/experiments/path/vit_config_registry.py b/torchtitan/experiments/path/vit_config_registry.py index 2c9e7f8fb3..d52c227829 100644 --- a/torchtitan/experiments/path/vit_config_registry.py +++ b/torchtitan/experiments/path/vit_config_registry.py @@ -1,16 +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. - """Config assembly + flavors for the path ViT, riding PathTrainer. The muP recipe (readout init/mult, eta/m optimizer groups, qk-norm, scheduler, widths, training) is -carried over verbatim from experiments/plan_vit/config_registry.py so it stays identical to the proven -config. The only change is the trainer wiring: these flavors return PathTrainer.Config (not the standalone -PlanViTTrainer.Config) with the path-specifics turned off in config -- the driving validator is disabled -and the checkpoint manager is the plain CheckpointManager with onnx export off. +assembled here. These flavors return PathTrainer.Config with the path-specifics turned off in config +-- the driving validator is disabled and the checkpoint manager is the plain CheckpointManager with +onnx export off. Width flavors scale n_head at fixed head_dim=64 (the clean muP axis); base = w256. Two cameras are channel-stacked into in_channels=24 (no VAE). @@ -66,7 +59,7 @@ IN_CHANNELS = 24 # two cameras (IMG + BIG_IMG), 12 YUV channels each, channel-stacked PLAN_SIZE = 15 * 33 * 2 # 990, laplacian mu+log-sigma BASE_WIDTH = 256 -PLAN_VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} +VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Config: @@ -126,7 +119,7 @@ def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PlanViTMLP.Config: def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: - n_embd = PLAN_VIT_WIDTHS[flavor] + n_embd = VIT_WIDTHS[flavor] n_head = n_embd // HEAD_DIM pt, ph, pw = PATCH_SIZE patch_dim = pt * IN_CHANNELS * ph * pw @@ -179,8 +172,8 @@ def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: STEPS = 512 # per-run step budget; override with training.steps=N on the CLI -# learning rate is the muTransfer sweep axis: one run per (flavor, lr); set with `-e PLAN_VIT_LR=...` -SWEEP_LR = float(os.getenv("PLAN_VIT_LR", "3e-4")) +# learning rate is the muTransfer sweep axis: one run per (flavor, lr); set with `-e VIT_LR=...` +SWEEP_LR = float(os.getenv("VIT_LR", "3e-4")) # hidden matrix weights get muP lr eta/m; input embed, readout, norms, biases get base eta # (readout is fan_in-infinite only, so Adam treats it vector-like -> base lr, not eta/m) MUP_PATTERN = ( @@ -226,7 +219,7 @@ def _dataloader_config(*, split: str) -> PathDataLoader.Config: def _optimizer_config( flavor: str, *, mup: bool, lr: float, wd: float ) -> OptimizersContainer.Config: - m = PLAN_VIT_WIDTHS[flavor] / BASE_WIDTH + m = VIT_WIDTHS[flavor] / BASE_WIDTH common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} if mup: groups = [ From 754be013e823ae44361093721ee7a96143201675 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Thu, 25 Jun 2026 22:17:19 -0700 Subject: [PATCH 10/28] path: drive the muP base lr via --mup_base_lr so run.sh's native sweep works --- torchtitan/experiments/path/trainer.py | 19 +++++++++++++++++++ .../experiments/path/vit_config_registry.py | 4 ++++ 2 files changed, 23 insertions(+) diff --git a/torchtitan/experiments/path/trainer.py b/torchtitan/experiments/path/trainer.py index 8ea74ba9ad..12020f21e9 100644 --- a/torchtitan/experiments/path/trainer.py +++ b/torchtitan/experiments/path/trainer.py @@ -30,12 +30,31 @@ class Config(Trainer.Config): # single-frame plan supervision: slice the dense plan target to the last frame before the # loss. off by default (dense; convnext/worldmodel unchanged); on for the single-frame path vit. plan_target_last_frame: bool = False + # muP base lr, the muTransfer sweep axis. tyro-overridable scalar (--mup_base_lr=X) so run.sh's + # native sweep can drive it; the eta/m split lives in optimizer.param_groups and is re-derived + # post-tyro by rescaling every group's lr by mup_base_lr/built_base_lr (see __post_init__). + mup_base_lr: float | None = None + # the base lr the optimizer param groups were built with, recorded by _optimizer_config so the + # post-tyro rescale knows the ratio. left None for non-vit flavors (no rescale). + built_base_lr: float | None = None def __post_init__(self) -> None: Trainer.Config.__post_init__(self) if self.codedir: self.miniray = {**self.miniray, "codedir": self.codedir} self.validator.miniray = {**self.validator.miniray, "codedir": self.codedir} + # re-derive the muP eta/m split for a swept base lr after tyro overlays --mup_base_lr. + # one shared ratio scales every group, so hidden stays base/m and others stay base; a + # no-op when the swept lr equals the build-time lr (default, non-swept runs). + if ( + self.mup_base_lr is not None + and self.built_base_lr is not None + and self.mup_base_lr != self.built_base_lr + ): + ratio = self.mup_base_lr / self.built_base_lr + for group in self.optimizer.param_groups: + group.optimizer_kwargs["lr"] *= ratio + self.built_base_lr = self.mup_base_lr # idempotent: a re-run sees no ratio def __init__(self, config: Config): super().__init__(config) diff --git a/torchtitan/experiments/path/vit_config_registry.py b/torchtitan/experiments/path/vit_config_registry.py index d52c227829..c15e1bdaf5 100644 --- a/torchtitan/experiments/path/vit_config_registry.py +++ b/torchtitan/experiments/path/vit_config_registry.py @@ -299,6 +299,10 @@ def _vit( fps=SUPERCOMBO_FPS, plan_target_last_frame=True, # ViT predicts a single-frame plan; supervise the last frame debug=DebugConfig(seed=0), + # record the build-time base lr so a swept --mup_base_lr can re-derive the eta/m split + # post-tyro (PathTrainer.Config.__post_init__). seeded equal -> default runs are a no-op. + built_base_lr=lr, + mup_base_lr=lr, ) From 16a3e4d22db82872aa440e90e03c4362f199670a Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Thu, 25 Jun 2026 22:26:15 -0700 Subject: [PATCH 11/28] cleanup --- .../experiments/path/config_registry.py | 2 - torchtitan/experiments/path/trainer.py | 13 +---- torchtitan/experiments/path/vit.py | 31 +++--------- .../experiments/path/vit_config_registry.py | 49 ++++--------------- 4 files changed, 18 insertions(+), 77 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 6bf627063f..9422c794e8 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -56,8 +56,6 @@ from .trainer import PathTrainer from .validate import PathValidator -# Path ViT flavors ride PathTrainer too; re-exported so `--module path --config vit_*` resolves -# here, the same way convnext_* do (the config manager looks the name up on this module). from .vit_config_registry import ( # noqa: F401 vit_mup_w256, vit_mup_w512, diff --git a/torchtitan/experiments/path/trainer.py b/torchtitan/experiments/path/trainer.py index 12020f21e9..2cb6655724 100644 --- a/torchtitan/experiments/path/trainer.py +++ b/torchtitan/experiments/path/trainer.py @@ -27,15 +27,8 @@ class Config(Trainer.Config): checkpoint: PathOnnxCheckpointManager.Config miniray: dict[str, Any] = field(default_factory=dict) fps: int - # single-frame plan supervision: slice the dense plan target to the last frame before the - # loss. off by default (dense; convnext/worldmodel unchanged); on for the single-frame path vit. plan_target_last_frame: bool = False - # muP base lr, the muTransfer sweep axis. tyro-overridable scalar (--mup_base_lr=X) so run.sh's - # native sweep can drive it; the eta/m split lives in optimizer.param_groups and is re-derived - # post-tyro by rescaling every group's lr by mup_base_lr/built_base_lr (see __post_init__). mup_base_lr: float | None = None - # the base lr the optimizer param groups were built with, recorded by _optimizer_config so the - # post-tyro rescale knows the ratio. left None for non-vit flavors (no rescale). built_base_lr: float | None = None def __post_init__(self) -> None: @@ -43,9 +36,6 @@ def __post_init__(self) -> None: if self.codedir: self.miniray = {**self.miniray, "codedir": self.codedir} self.validator.miniray = {**self.validator.miniray, "codedir": self.codedir} - # re-derive the muP eta/m split for a swept base lr after tyro overlays --mup_base_lr. - # one shared ratio scales every group, so hidden stays base/m and others stay base; a - # no-op when the swept lr equals the build-time lr (default, non-swept runs). if ( self.mup_base_lr is not None and self.built_base_lr is not None @@ -54,7 +44,7 @@ def __post_init__(self) -> None: ratio = self.mup_base_lr / self.built_base_lr for group in self.optimizer.param_groups: group.optimizer_kwargs["lr"] *= ratio - self.built_base_lr = self.mup_base_lr # idempotent: a re-run sees no ratio + self.built_base_lr = self.mup_base_lr def __init__(self, config: Config): super().__init__(config) @@ -99,7 +89,6 @@ def forward_backward_step( with self.train_context(): pred = self.model_parts[0](inputs) if self.config.plan_target_last_frame: - # single-frame plan models predict the last frame's plan; supervise that frame labels = {**labels, "plan": labels["plan"][:, -1]} loss_vec, metrics = self.loss_fn(pred, labels) loss = loss_vec.sum() / local_samples diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py index 3a1b774063..edfc0346d1 100644 --- a/torchtitan/experiments/path/vit.py +++ b/torchtitan/experiments/path/vit.py @@ -1,10 +1,3 @@ -"""Plan ViT for the path experiment: raw camera frames -> patches -> transformer -> plan. NO VAE. - -Rides PathTrainer via config. A self-contained planning model for the muP + scaling study, built -from torchtitan.models.common blocks the same way path/model.py is. Scales cleanly by width -(n_embd / n_head) for muTransfer. -""" - from __future__ import annotations from dataclasses import dataclass @@ -88,9 +81,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) q, k, v = qkv.unbind(2) q, k = self.q_norm(q), self.k_norm(k) - x = self.inner_attention( - q, k, v, is_causal=False - ) # ViT: bidirectional over patches + x = self.inner_attention(q, k, v, is_causal=False) return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) @@ -114,7 +105,7 @@ class PatchEmbed(Module): @dataclass(kw_only=True, slots=True) class Config(Module.Config): proj: Linear.Config - patch_size: tuple[int, int, int] # (pt, ph, pw) + patch_size: tuple[int, int, int] def __init__(self, config: Config): super().__init__() @@ -122,14 +113,11 @@ def __init__(self, config: Config): self.proj = config.proj.build() def forward(self, x: torch.Tensor) -> torch.Tensor: - # x: (B, T, C, H, W) raw frames -> (B, num_patches, patch_dim) -> (B, num_patches, n_embd) pt, ph, pw = self.patch_size x = rearrange( x, "b (t pt) c (h ph) (w pw) -> b (t h w) (pt c ph pw)", pt=pt, ph=ph, pw=pw ) - return self.proj( - x.to(self.proj.weight.dtype) - ) # match the bf16 (mp) weights, like path's vision + return self.proj(x.to(self.proj.weight.dtype)) class PlanHead(Module): @@ -150,11 +138,11 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class PlanViT(BaseModel): @dataclass(kw_only=True, slots=True) class Config(BaseModel.Config): - input_size: tuple[int, int, int] # (n_frames, H, W) + input_size: tuple[int, int, int] patch_size: tuple[int, int, int] in_channels: int n_embd: int - output_mult: float # muP readout multiplier 1/m (m = n_embd / base); 1.0 for standard param + output_mult: float patch_embed: PatchEmbed.Config pos_embedding: Embedding.Config blocks: list[PlanViTBlock.Config] @@ -192,18 +180,14 @@ def __init__(self, config: Config): self.plan_head = config.plan_head.build() def verify_module_protocol(self) -> None: - pass # nn.Dropout/GELU/Identity are plain nn.Module, like path + pass def _frames(self, inputs: dict[str, torch.Tensor] | torch.Tensor) -> torch.Tensor: - # production input: two cameras IMG, BIG_IMG, each (B, T, 12, H, W) YUV. Take the current frame of each, - # channel-stack -> (B, 1, 24, H, W). NO VAE. A plain tensor (testing) is passed through unchanged. if isinstance(inputs, torch.Tensor): return inputs img, big = inputs[ModelInputs.IMG], inputs[ModelInputs.BIG_IMG] frame = torch.cat([img[:, -1], big[:, -1]], dim=1).unsqueeze(1) - return ( - frame.float() - 127.5 - ) / 63.75 # uint8 YUV -> normalized float (mean 255/2, std 255/4 like path) + return (frame.float() - 127.5) / 63.75 def forward( self, inputs: dict[str, torch.Tensor] | torch.Tensor @@ -214,7 +198,6 @@ def forward( for block in self.blocks: x = block(x) x = self.norm(x) - # global-pool the patches -> plan; the muP readout multiplier keeps the output width-stable return {"plan": self.plan_head(x.mean(dim=1)) * self.config.output_mult} diff --git a/torchtitan/experiments/path/vit_config_registry.py b/torchtitan/experiments/path/vit_config_registry.py index c15e1bdaf5..63c69b1817 100644 --- a/torchtitan/experiments/path/vit_config_registry.py +++ b/torchtitan/experiments/path/vit_config_registry.py @@ -1,14 +1,3 @@ -"""Config assembly + flavors for the path ViT, riding PathTrainer. - -The muP recipe (readout init/mult, eta/m optimizer groups, qk-norm, scheduler, widths, training) is -assembled here. These flavors return PathTrainer.Config with the path-specifics turned off in config --- the driving validator is disabled and the checkpoint manager is the plain CheckpointManager with -onnx export off. - -Width flavors scale n_head at fixed head_dim=64 (the clean muP axis); base = w256. Two cameras are -channel-stacked into in_channels=24 (no VAE). -""" - from __future__ import annotations import math @@ -54,10 +43,10 @@ 1, 128, 256, -) # current frame; spatial ViT (temporal history is a later variant) +) PATCH_SIZE = (1, 16, 8) -IN_CHANNELS = 24 # two cameras (IMG + BIG_IMG), 12 YUV channels each, channel-stacked -PLAN_SIZE = 15 * 33 * 2 # 990, laplacian mu+log-sigma +IN_CHANNELS = 24 +PLAN_SIZE = 15 * 33 * 2 BASE_WIDTH = 256 VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} @@ -75,8 +64,6 @@ def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Conf def _hidden_std(fan_in: int, *, mup: bool) -> float: - # muP shrinks hidden/output init to 1/sqrt(fan_in) so pre-activations stay O(1) as width grows; - # standard param holds the base-width variance 1/sqrt(BASE_WIDTH), so it fans out with width. return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 @@ -130,13 +117,9 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co patch_size=PATCH_SIZE, in_channels=IN_CHANNELS, n_embd=n_embd, - output_mult=(BASE_WIDTH / n_embd) - if mup - else 1.0, # muP readout fwd mult 1/m (init output slopes ~1/sqrt(m)) + output_mult=(BASE_WIDTH / n_embd) if mup else 1.0, patch_embed=PatchEmbed.Config( - proj=_lin( - patch_dim, n_embd, std=patch_dim**-0.5 - ), # input embed: width-independent + proj=_lin(patch_dim, n_embd, std=patch_dim**-0.5), patch_size=PATCH_SIZE, ), pos_embedding=Embedding.Config( @@ -152,9 +135,7 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co norm=_ln(n_embd), plan_head=PlanHead.Config( norm=_ln(n_embd), - head=_lin( - n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5 - ), # muP readout: base-width init + head=_lin(n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5), ), ) @@ -171,11 +152,8 @@ def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: ) -STEPS = 512 # per-run step budget; override with training.steps=N on the CLI -# learning rate is the muTransfer sweep axis: one run per (flavor, lr); set with `-e VIT_LR=...` +STEPS = 512 SWEEP_LR = float(os.getenv("VIT_LR", "3e-4")) -# hidden matrix weights get muP lr eta/m; input embed, readout, norms, biases get base eta -# (readout is fan_in-infinite only, so Adam treats it vector-like -> base lr, not eta/m) MUP_PATTERN = ( r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" @@ -199,7 +177,6 @@ def _dataloader_config(*, split: str) -> PathDataLoader.Config: base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) return PathDataLoader.Config( - # prune-10M study data: a seeded random 10k sample of the 10M store (training_2026_02) dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), split=split, shuffle_size=_si_int(base.shuffle_size), @@ -207,7 +184,7 @@ def _dataloader_config(*, split: str) -> PathDataLoader.Config: num_writers=base.num_writers, num_readers=base.num_readers, fps=base.fps, - pipeline_dir=BASE_DIR_GT_10M, # the 10M store, not the 2.5M big-train list + pipeline_dir=BASE_DIR_GT_10M, plan_only=base.plan_only, limit=base.limit, n_frames=base.n_frames, @@ -250,7 +227,6 @@ def _optimizer_config( def _vit( flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 ) -> PathTrainer.Config: - # derive data parallelism from the launch (like path), so any N nodes x GPUs validate local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) num_nodes = int( @@ -264,7 +240,7 @@ def _vit( optimizer=_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), lr_scheduler=LRSchedulersContainer.Config( warmup_steps=round(STEPS * 0.1), - total_steps=None, # use the real training.steps; a fixed value wraps the cosine on longer runs + total_steps=None, decay_ratio=0.8, decay_type="cosine", min_lr_factor=0.0, @@ -283,13 +259,10 @@ def _vit( data_parallel_replicate_degree=num_nodes, data_parallel_shard_degree=local_world_size, ), - # plain CheckpointManager with onnx export off (path-specific PathOnnxCheckpointManager disabled) checkpoint=CheckpointManager.Config(enable=False), metrics=MetricsProcessor.Config( log_freq=10, enable_reporterv2=True, save_freq=STEPS ), - # path-specific driving validator disabled; dataloader is required by the dataclass but never - # built while enable=False (Trainer builds the validator only when validator.enable is True) validator=PathValidator.Config( enable=False, steps=-1, @@ -297,10 +270,8 @@ def _vit( mixed_precision_param="bfloat16", ), fps=SUPERCOMBO_FPS, - plan_target_last_frame=True, # ViT predicts a single-frame plan; supervise the last frame + plan_target_last_frame=True, debug=DebugConfig(seed=0), - # record the build-time base lr so a swept --mup_base_lr can re-derive the eta/m split - # post-tyro (PathTrainer.Config.__post_init__). seeded equal -> default runs are a no-op. built_base_lr=lr, mup_base_lr=lr, ) From 8e0b8c8ce20cab1f5b54672ed3129d189df9ed8c Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 10:07:26 -0700 Subject: [PATCH 12/28] path: build the vit from path's transformer blocks --- torchtitan/experiments/path/model.py | 4 +- torchtitan/experiments/path/vit.py | 80 +------------------ .../experiments/path/vit_config_registry.py | 15 ++-- 3 files changed, 12 insertions(+), 87 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index c3b02b1c38..4fbbb2aec0 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -90,11 +90,13 @@ class Config(Module.Config): n_head: int head_dim: int dropout: float + is_causal: bool = True def __init__(self, config: Config): super().__init__() self.n_head = config.n_head self.head_dim = config.head_dim + self.is_causal = config.is_causal self.norm = config.norm.build() self.q_norm = config.q_norm.build() if config.q_norm is not None else nn.Identity() self.k_norm = config.k_norm.build() if config.k_norm is not None else nn.Identity() @@ -108,7 +110,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) q, k, v = qkv.unbind(2) q, k = self.q_norm(q), self.k_norm(k) - x = self.inner_attention(q, k, v, is_causal=True) + x = self.inner_attention(q, k, v, is_causal=self.is_causal) return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py index edfc0346d1..652136fcf3 100644 --- a/torchtitan/experiments/path/vit.py +++ b/torchtitan/experiments/path/vit.py @@ -18,87 +18,11 @@ from torchtitan.distributed import ParallelDims from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig from torchtitan.models.common import Embedding, LayerNorm, Linear, RMSNorm -from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model import BaseModel from torchtitan.protocols.module import Module, ModuleList from torchtitan.tools.logging import logger - -class PlanViTMLP(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - norm: LayerNorm.Config | RMSNorm.Config - c_fc: Linear.Config - c_proj: Linear.Config - act: str - dropout: float - - def __init__(self, config: Config): - super().__init__() - self.norm = config.norm.build() - self.c_fc = config.c_fc.build() - self.act = ( - nn.GELU(approximate="tanh") if config.act == "gelu_tanh" else nn.GELU() - ) - self.c_proj = config.c_proj.build() - self.dropout = nn.Dropout(config.dropout) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.dropout(self.c_proj(self.act(self.c_fc(self.norm(x))))) - - -class PlanViTAttention(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - norm: LayerNorm.Config | RMSNorm.Config - q_norm: LayerNorm.Config | RMSNorm.Config | None - k_norm: LayerNorm.Config | RMSNorm.Config | None - c_attn: Linear.Config - c_proj: Linear.Config - inner_attention: ScaledDotProductAttention.Config - n_head: int - head_dim: int - dropout: float - - def __init__(self, config: Config): - super().__init__() - self.n_head = config.n_head - self.head_dim = config.head_dim - self.norm = config.norm.build() - self.q_norm = ( - config.q_norm.build() if config.q_norm is not None else nn.Identity() - ) - self.k_norm = ( - config.k_norm.build() if config.k_norm is not None else nn.Identity() - ) - self.c_attn = config.c_attn.build() - self.c_proj = config.c_proj.build() - self.inner_attention = config.inner_attention.build() - self.dropout = nn.Dropout(config.dropout) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - b, t, _ = x.shape - qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) - q, k, v = qkv.unbind(2) - q, k = self.q_norm(q), self.k_norm(k) - x = self.inner_attention(q, k, v, is_causal=False) - return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) - - -class PlanViTBlock(Module): - @dataclass(kw_only=True, slots=True) - class Config(Module.Config): - attention: PlanViTAttention.Config - mlp: PlanViTMLP.Config - - def __init__(self, config: Config): - super().__init__() - self.attention = config.attention.build() - self.mlp = config.mlp.build() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x = x + self.attention(x) - return x + self.mlp(x) +from .model import PathMLP, PathSelfAttention, PathTransformerBlock class PatchEmbed(Module): @@ -145,7 +69,7 @@ class Config(BaseModel.Config): output_mult: float patch_embed: PatchEmbed.Config pos_embedding: Embedding.Config - blocks: list[PlanViTBlock.Config] + blocks: list[PathTransformerBlock.Config] norm: LayerNorm.Config | RMSNorm.Config plan_head: PlanHead.Config diff --git a/torchtitan/experiments/path/vit_config_registry.py b/torchtitan/experiments/path/vit_config_registry.py index 63c69b1817..1171545a5a 100644 --- a/torchtitan/experiments/path/vit_config_registry.py +++ b/torchtitan/experiments/path/vit_config_registry.py @@ -21,14 +21,12 @@ from .loss import PathLoss from .trainer import PathTrainer from .validate import PathValidator +from .model import PathMLP, PathSelfAttention, PathTransformerBlock from .vit import ( parallelize_vit, PatchEmbed, PlanHead, PlanViT, - PlanViTAttention, - PlanViTBlock, - PlanViTMLP, ) _LINEAR_INIT = { @@ -77,9 +75,9 @@ def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: def _attention( dim: int, n_head: int, *, mup: bool, qk_norm: bool = True -) -> PlanViTAttention.Config: +) -> PathSelfAttention.Config: head_dim = dim // n_head - return PlanViTAttention.Config( + return PathSelfAttention.Config( norm=_ln(dim), q_norm=_ln(head_dim) if qk_norm else None, k_norm=_ln(head_dim) if qk_norm else None, @@ -89,12 +87,13 @@ def _attention( n_head=n_head, head_dim=head_dim, dropout=0.0, + is_causal=False, ) -def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PlanViTMLP.Config: +def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: hidden = _hidden(dim, mult) - return PlanViTMLP.Config( + return PathMLP.Config( norm=_ln(dim), c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), c_proj=_lin( @@ -126,7 +125,7 @@ def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Co num_embeddings=num_patches, embedding_dim=n_embd, param_init=_LINEAR_INIT ), blocks=[ - PlanViTBlock.Config( + PathTransformerBlock.Config( attention=_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), mlp=_mlp(n_embd, mup=mup), ) From 3aad5e8ea941c0a4092515b11a90c38d452c454b Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 10:33:21 -0700 Subject: [PATCH 13/28] path: fold the vit configs into config_registry --- .../experiments/path/vit_config_registry.py | 308 ------------------ 1 file changed, 308 deletions(-) delete mode 100644 torchtitan/experiments/path/vit_config_registry.py diff --git a/torchtitan/experiments/path/vit_config_registry.py b/torchtitan/experiments/path/vit_config_registry.py deleted file mode 100644 index 1171545a5a..0000000000 --- a/torchtitan/experiments/path/vit_config_registry.py +++ /dev/null @@ -1,308 +0,0 @@ -from __future__ import annotations - -import math -import os -from functools import partial -from xx.ml_tools.constants.model import SUPERCOMBO_FPS - -import torch.nn as nn - -from torchtitan.components.checkpoint import CheckpointManager -from torchtitan.components.lr_scheduler import LRSchedulersContainer -from torchtitan.components.metrics import MetricsProcessor -from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig -from torchtitan.components.tokenizer import NoOpTokenizer -from torchtitan.config import DebugConfig, ParallelismConfig, TrainingConfig -from torchtitan.models.common import Embedding, LayerNorm, Linear -from torchtitan.models.common.attention import ScaledDotProductAttention -from torchtitan.protocols.model_spec import ModelSpec - -from .dataset import PathDataLoader -from .loss import PathLoss -from .trainer import PathTrainer -from .validate import PathValidator -from .model import PathMLP, PathSelfAttention, PathTransformerBlock -from .vit import ( - parallelize_vit, - PatchEmbed, - PlanHead, - PlanViT, -) - -_LINEAR_INIT = { - "weight": partial(nn.init.normal_, mean=0.0, std=0.02), - "bias": nn.init.zeros_, -} -_NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} - -HEAD_DIM = 64 -N_LAYER = 8 -INPUT_SIZE = ( - 1, - 128, - 256, -) -PATCH_SIZE = (1, 16, 8) -IN_CHANNELS = 24 -PLAN_SIZE = 15 * 33 * 2 -BASE_WIDTH = 256 -VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} - - -def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Config: - return Linear.Config( - in_features=in_f, - out_features=out_f, - bias=bias, - param_init={ - "weight": partial(nn.init.normal_, mean=0.0, std=std), - "bias": nn.init.zeros_, - }, - ) - - -def _hidden_std(fan_in: int, *, mup: bool) -> float: - return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 - - -def _ln(dim: int) -> LayerNorm.Config: - return LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT) - - -def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: - return multiple_of * math.ceil(int(dim * mult) / multiple_of) - - -def _attention( - dim: int, n_head: int, *, mup: bool, qk_norm: bool = True -) -> PathSelfAttention.Config: - head_dim = dim // n_head - return PathSelfAttention.Config( - norm=_ln(dim), - q_norm=_ln(head_dim) if qk_norm else None, - k_norm=_ln(head_dim) if qk_norm else None, - c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), - c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * N_LAYER)), - inner_attention=ScaledDotProductAttention.Config(), - n_head=n_head, - head_dim=head_dim, - dropout=0.0, - is_causal=False, - ) - - -def _mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: - hidden = _hidden(dim, mult) - return PathMLP.Config( - norm=_ln(dim), - c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), - c_proj=_lin( - hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER) - ), - act="gelu_tanh", - dropout=0.0, - ) - - -def _model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: - n_embd = VIT_WIDTHS[flavor] - n_head = n_embd // HEAD_DIM - pt, ph, pw = PATCH_SIZE - patch_dim = pt * IN_CHANNELS * ph * pw - t, h, w = INPUT_SIZE - num_patches = (t // pt) * (h // ph) * (w // pw) - return PlanViT.Config( - input_size=INPUT_SIZE, - patch_size=PATCH_SIZE, - in_channels=IN_CHANNELS, - n_embd=n_embd, - output_mult=(BASE_WIDTH / n_embd) if mup else 1.0, - patch_embed=PatchEmbed.Config( - proj=_lin(patch_dim, n_embd, std=patch_dim**-0.5), - patch_size=PATCH_SIZE, - ), - pos_embedding=Embedding.Config( - num_embeddings=num_patches, embedding_dim=n_embd, param_init=_LINEAR_INIT - ), - blocks=[ - PathTransformerBlock.Config( - attention=_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), - mlp=_mlp(n_embd, mup=mup), - ) - for _ in range(N_LAYER) - ], - norm=_ln(n_embd), - plan_head=PlanHead.Config( - norm=_ln(n_embd), - head=_lin(n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5), - ), - ) - - -def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: - return ModelSpec( - name="path", - flavor=flavor, - model=_model_config(flavor, mup=mup), - parallelize_fn=parallelize_vit, - pipelining_fn=None, - post_optimizer_build_fn=None, - state_dict_adapter=None, - ) - - -STEPS = 512 -SWEEP_LR = float(os.getenv("VIT_LR", "3e-4")) -MUP_PATTERN = ( - r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" - r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" -) - - -def _si_int(value: str | int) -> int: - suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} - value = str(value).strip().lower() - return ( - int(float(value[:-1]) * suffixes[value[-1]]) - if value[-1] in suffixes - else int(value) - ) - - -def _dataloader_config(*, split: str) -> PathDataLoader.Config: - from xx.common.basedir import XX_BASEDIR - from xx.datasets.constants import BASE_DIR_GT_10M - from xx.training.path.config import DatasetConfig as XXPathDatasetConfig - - base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) - return PathDataLoader.Config( - dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), - split=split, - shuffle_size=_si_int(base.shuffle_size), - min_mixing=base.min_mixing, - num_writers=base.num_writers, - num_readers=base.num_readers, - fps=base.fps, - pipeline_dir=BASE_DIR_GT_10M, - plan_only=base.plan_only, - limit=base.limit, - n_frames=base.n_frames, - rgb=base.rgb, - unvision=base.unvision, - ) - - -def _optimizer_config( - flavor: str, *, mup: bool, lr: float, wd: float -) -> OptimizersContainer.Config: - m = VIT_WIDTHS[flavor] / BASE_WIDTH - common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} - if mup: - groups = [ - ParamGroupConfig( - pattern=MUP_PATTERN, - optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr / m}, - ), - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr}, - ), - ] - else: - groups = [ - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr}, - ) - ] - return OptimizersContainer.Config( - implementation="fused_opt_states_bf16", param_groups=groups - ) - - -def _vit( - flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 -) -> PathTrainer.Config: - local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) - world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int( - os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) - ) - return PathTrainer.Config( - loss=PathLoss.Config(), - model_spec=vit_model_registry(flavor, mup=mup), - tokenizer=NoOpTokenizer.Config(), - dataloader=_dataloader_config(split="train"), - optimizer=_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), - lr_scheduler=LRSchedulersContainer.Config( - warmup_steps=round(STEPS * 0.1), - total_steps=None, - decay_ratio=0.8, - decay_type="cosine", - min_lr_factor=0.0, - ), - training=TrainingConfig( - local_batch_size=16, - global_batch_size=-1, - seq_len=1, - steps=STEPS, - max_norm=1.0, - dtype="float32", - mixed_precision_param="bfloat16", - mixed_precision_reduce="float32", - ), - parallelism=ParallelismConfig( - data_parallel_replicate_degree=num_nodes, - data_parallel_shard_degree=local_world_size, - ), - checkpoint=CheckpointManager.Config(enable=False), - metrics=MetricsProcessor.Config( - log_freq=10, enable_reporterv2=True, save_freq=STEPS - ), - validator=PathValidator.Config( - enable=False, - steps=-1, - dataloader=_dataloader_config(split="val"), - mixed_precision_param="bfloat16", - ), - fps=SUPERCOMBO_FPS, - plan_target_last_frame=True, - debug=DebugConfig(seed=0), - built_base_lr=lr, - mup_base_lr=lr, - ) - - -def vit_standard_w256() -> PathTrainer.Config: - return _vit("w256", mup=False) - - -def vit_standard_w512() -> PathTrainer.Config: - return _vit("w512", mup=False) - - -def vit_standard_w1024() -> PathTrainer.Config: - return _vit("w1024", mup=False) - - -def vit_standard_w2048() -> PathTrainer.Config: - return _vit("w2048", mup=False) - - -def vit_mup_w256() -> PathTrainer.Config: - return _vit("w256", mup=True) - - -def vit_mup_w512() -> PathTrainer.Config: - return _vit("w512", mup=True) - - -def vit_mup_w1024() -> PathTrainer.Config: - return _vit("w1024", mup=True) - - -def vit_mup_w2048() -> PathTrainer.Config: - return _vit("w2048", mup=True) From 77355af1aa0377d39f8273a2465c5c774ea642f7 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 10:38:12 -0700 Subject: [PATCH 14/28] path: inline the vit configs --- .../experiments/path/config_registry.py | 277 +++++++++++++++++- 1 file changed, 267 insertions(+), 10 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 9422c794e8..298a7c3a9c 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -6,6 +6,7 @@ import torch.nn as nn +from torchtitan.components.checkpoint import CheckpointManager from torchtitan.components.lr_scheduler import LRSchedulersContainer from torchtitan.components.metrics import MetricsProcessor from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig @@ -55,16 +56,11 @@ from .onnx_checkpoint import PathOnnxCheckpointManager from .trainer import PathTrainer from .validate import PathValidator - -from .vit_config_registry import ( # noqa: F401 - vit_mup_w256, - vit_mup_w512, - vit_mup_w1024, - vit_mup_w2048, - vit_standard_w256, - vit_standard_w512, - vit_standard_w1024, - vit_standard_w2048, +from .vit import ( + parallelize_vit, + PatchEmbed, + PlanHead, + PlanViT, ) @@ -377,3 +373,264 @@ def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> }, scale_layers={head.name: ScaleLayer.Config(n_features=head.output_size) for head in heads if head.scale}, ) + + +HEAD_DIM = 64 +N_LAYER = 8 +INPUT_SIZE = ( + 1, + 128, + 256, +) +PATCH_SIZE = (1, 16, 8) +IN_CHANNELS = 24 +PLAN_SIZE = 15 * 33 * 2 +BASE_WIDTH = 256 +VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} +STEPS = 512 +SWEEP_LR = float(os.getenv("VIT_LR", "3e-4")) +MUP_PATTERN = ( + r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" + r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" +) + + +def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Config: + return Linear.Config( + in_features=in_f, + out_features=out_f, + bias=bias, + param_init={ + "weight": partial(nn.init.normal_, mean=0.0, std=std), + "bias": nn.init.zeros_, + }, + ) + + +def _hidden_std(fan_in: int, *, mup: bool) -> float: + return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 + + +def _ln(dim: int) -> LayerNorm.Config: + return LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT) + + +def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: + return multiple_of * math.ceil(int(dim * mult) / multiple_of) + + +def _vit_attention( + dim: int, n_head: int, *, mup: bool, qk_norm: bool = True +) -> PathSelfAttention.Config: + head_dim = dim // n_head + return PathSelfAttention.Config( + norm=_ln(dim), + q_norm=_ln(head_dim) if qk_norm else None, + k_norm=_ln(head_dim) if qk_norm else None, + c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), + c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * N_LAYER)), + inner_attention=ScaledDotProductAttention.Config(), + n_head=n_head, + head_dim=head_dim, + dropout=0.0, + is_causal=False, + ) + + +def _vit_mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: + hidden = _hidden(dim, mult) + return PathMLP.Config( + norm=_ln(dim), + c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), + c_proj=_lin( + hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER) + ), + act="gelu_tanh", + dropout=0.0, + ) + + +def _vit_model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: + n_embd = VIT_WIDTHS[flavor] + n_head = n_embd // HEAD_DIM + pt, ph, pw = PATCH_SIZE + patch_dim = pt * IN_CHANNELS * ph * pw + t, h, w = INPUT_SIZE + num_patches = (t // pt) * (h // ph) * (w // pw) + return PlanViT.Config( + input_size=INPUT_SIZE, + patch_size=PATCH_SIZE, + in_channels=IN_CHANNELS, + n_embd=n_embd, + output_mult=(BASE_WIDTH / n_embd) if mup else 1.0, + patch_embed=PatchEmbed.Config( + proj=_lin(patch_dim, n_embd, std=patch_dim**-0.5), + patch_size=PATCH_SIZE, + ), + pos_embedding=Embedding.Config( + num_embeddings=num_patches, embedding_dim=n_embd, param_init=_LINEAR_INIT + ), + blocks=[ + PathTransformerBlock.Config( + attention=_vit_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), + mlp=_vit_mlp(n_embd, mup=mup), + ) + for _ in range(N_LAYER) + ], + norm=_ln(n_embd), + plan_head=PlanHead.Config( + norm=_ln(n_embd), + head=_lin(n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5), + ), + ) + + +def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: + return ModelSpec( + name="path", + flavor=flavor, + model=_vit_model_config(flavor, mup=mup), + parallelize_fn=parallelize_vit, + pipelining_fn=None, + post_optimizer_build_fn=None, + state_dict_adapter=None, + ) + + +def _vit_dataloader_config(*, split: str) -> PathDataLoader.Config: + from xx.common.basedir import XX_BASEDIR + from xx.datasets.constants import BASE_DIR_GT_10M + from xx.training.path.config import DatasetConfig as XXPathDatasetConfig + + base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) + return PathDataLoader.Config( + dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), + split=split, + shuffle_size=_si_int(base.shuffle_size), + min_mixing=base.min_mixing, + num_writers=base.num_writers, + num_readers=base.num_readers, + fps=base.fps, + pipeline_dir=BASE_DIR_GT_10M, + plan_only=base.plan_only, + limit=base.limit, + n_frames=base.n_frames, + rgb=base.rgb, + unvision=base.unvision, + ) + + +def _vit_optimizer_config( + flavor: str, *, mup: bool, lr: float, wd: float +) -> OptimizersContainer.Config: + m = VIT_WIDTHS[flavor] / BASE_WIDTH + common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} + if mup: + groups = [ + ParamGroupConfig( + pattern=MUP_PATTERN, + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr / m}, + ), + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr}, + ), + ] + else: + groups = [ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common, "lr": lr}, + ) + ] + return OptimizersContainer.Config( + implementation="fused_opt_states_bf16", param_groups=groups + ) + + +def _vit( + flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 +) -> PathTrainer.Config: + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) + world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) + num_nodes = int( + os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) + ) + return PathTrainer.Config( + loss=PathLoss.Config(), + model_spec=vit_model_registry(flavor, mup=mup), + tokenizer=NoOpTokenizer.Config(), + dataloader=_vit_dataloader_config(split="train"), + optimizer=_vit_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), + lr_scheduler=LRSchedulersContainer.Config( + warmup_steps=round(STEPS * 0.1), + total_steps=None, + decay_ratio=0.8, + decay_type="cosine", + min_lr_factor=0.0, + ), + training=TrainingConfig( + local_batch_size=16, + global_batch_size=-1, + seq_len=1, + steps=STEPS, + max_norm=1.0, + dtype="float32", + mixed_precision_param="bfloat16", + mixed_precision_reduce="float32", + ), + parallelism=ParallelismConfig( + data_parallel_replicate_degree=num_nodes, + data_parallel_shard_degree=local_world_size, + ), + checkpoint=CheckpointManager.Config(enable=False), + metrics=MetricsProcessor.Config( + log_freq=10, enable_reporterv2=True, save_freq=STEPS + ), + validator=PathValidator.Config( + enable=False, + steps=-1, + dataloader=_vit_dataloader_config(split="val"), + mixed_precision_param="bfloat16", + ), + fps=SUPERCOMBO_FPS, + plan_target_last_frame=True, + debug=DebugConfig(seed=0), + built_base_lr=lr, + mup_base_lr=lr, + ) + + +def vit_standard_w256() -> PathTrainer.Config: + return _vit("w256", mup=False) + + +def vit_standard_w512() -> PathTrainer.Config: + return _vit("w512", mup=False) + + +def vit_standard_w1024() -> PathTrainer.Config: + return _vit("w1024", mup=False) + + +def vit_standard_w2048() -> PathTrainer.Config: + return _vit("w2048", mup=False) + + +def vit_mup_w256() -> PathTrainer.Config: + return _vit("w256", mup=True) + + +def vit_mup_w512() -> PathTrainer.Config: + return _vit("w512", mup=True) + + +def vit_mup_w1024() -> PathTrainer.Config: + return _vit("w1024", mup=True) + + +def vit_mup_w2048() -> PathTrainer.Config: + return _vit("w2048", mup=True) From 3cfbb59db6966db24e59e591d59e82cfcb5af53e Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 11:26:55 -0700 Subject: [PATCH 15/28] path: follow house patterns; muP lr via --optimizer.lr --- torchtitan/components/optimizer.py | 35 +++++- .../experiments/path/config_registry.py | 113 ++++++++---------- torchtitan/experiments/path/trainer.py | 14 --- torchtitan/experiments/path/vit.py | 55 +++++---- 4 files changed, 114 insertions(+), 103 deletions(-) diff --git a/torchtitan/components/optimizer.py b/torchtitan/components/optimizer.py index bd399fe50b..f17869dc2f 100644 --- a/torchtitan/components/optimizer.py +++ b/torchtitan/components/optimizer.py @@ -65,7 +65,14 @@ class ParamGroupConfig: optimizer_kwargs: dict[str, Any] = field(default_factory=dict) """Keyword arguments passed to the optimizer constructor. - Must include all required kwargs (e.g. ``lr``). No implicit defaults.""" + Must include all required kwargs (e.g. ``lr``) unless the container sets a + base ``lr`` (see ``OptimizersContainer.Config.lr``). No implicit defaults.""" + + lr_mult: float = 1.0 + """Multiplier on the container's base ``lr`` for this group. Only applies when + ``OptimizersContainer.Config.lr`` is set; the group's learning rate is then + ``lr * lr_mult`` and ``lr`` must not also appear in ``optimizer_kwargs``. + Defaults to 1.0 (group uses the base lr unscaled).""" T = TypeVar("T", bound=Optimizer) @@ -107,6 +114,14 @@ class Config(Configurable.Config): regex pattern and a self-contained optimizer setup. Patterns are checked in order; first match wins.""" + lr: float | None = None + """Optional base learning rate shared by all param groups. When set, each + group's learning rate is ``lr * ParamGroupConfig.lr_mult`` and groups must + not set their own ``lr`` in ``optimizer_kwargs``. When ``None`` (default), + each group provides its own ``lr`` in ``optimizer_kwargs``. A single base + lr lets the native ``--optimizer.lr`` override (and lr sweeps) scale every + group at once while preserving per-group ratios (e.g. muP lr scaling).""" + implementation: Literal[ "for-loop", "foreach", "fused", "fused_opt_states_bf16" ] = "fused" @@ -158,10 +173,13 @@ def _build_param_groups( model: nn.Module, param_group_configs: list[ParamGroupConfig], impl_kwargs: dict[str, Any], + base_lr: float | None = None, ) -> tuple[dict[str, list[dict[str, Any]]], dict[str, list[str]]]: """Build PyTorch param groups from model parameters, partitioned by optimizer. Each parameter is assigned to the first matching ParamGroupConfig pattern. + When ``base_lr`` is set, each group's learning rate is ``base_lr * + ParamGroupConfig.lr_mult`` and groups must not set their own ``lr``. Returns two dicts keyed by optimizer name and aligned by index: the param group dicts to pass to the optimizer constructor, and the regex pattern of @@ -192,12 +210,21 @@ def _build_param_groups( f"matched no parameters" ) + group_kwargs = {**impl_kwargs, **pg.optimizer_kwargs} + if base_lr is not None: + if "lr" in pg.optimizer_kwargs: + raise ValueError( + f"Optimizer param_groups pattern '{pg.pattern}' sets 'lr' in " + f"optimizer_kwargs while the optimizer Config also sets a base " + f"lr; use lr_mult to scale from the base lr instead" + ) + group_kwargs["lr"] = base_lr * pg.lr_mult + groups[pg.optimizer_name].append( { "params": params, "param_names": param_names, - **impl_kwargs, - **pg.optimizer_kwargs, + **group_kwargs, } ) patterns[pg.optimizer_name].append(pg.pattern) @@ -213,7 +240,7 @@ def __init__(self, config: Config, *, model_parts: list[nn.Module]) -> None: for part_idx, model in enumerate(self.model_parts): groups_by_opt_name, patterns_by_opt_name = self._build_param_groups( - model, param_group_configs, impl_kwargs + model, param_group_configs, impl_kwargs, base_lr=config.lr ) for opt_name, opt_param_groups in groups_by_opt_name.items(): optimizer = self._resolve_optimizer_cls(opt_name)(opt_param_groups) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 298a7c3a9c..1c8b304205 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -21,6 +21,8 @@ from torchtitan.models.common import Embedding, LayerNorm, Linear from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model_spec import ModelSpec +from xx.common.basedir import XX_BASEDIR +from xx.datasets.constants import BASE_DIR_GT_10M from xx.datasets.helpers import DEFAULT_BIG_TRAIN_LIST from xx.ml_tools.constants.model import ( SUPERCOMBO_FPS, @@ -61,12 +63,28 @@ PatchEmbed, PlanHead, PlanViT, + PlanViTLoss, ) _LINEAR_INIT = {"weight": partial(nn.init.normal_, mean=0.0, std=0.02), "bias": nn.init.zeros_} _NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} +# PlanViT (single-frame plan ViT) architecture and muP constants +HEAD_DIM = 64 +NUM_LAYERS = 8 +INPUT_SIZE = (1, 128, 256) +PATCH_SIZE = (1, 16, 8) +IN_CHANNELS = 24 +PLAN_SIZE = 15 * 33 * 2 +BASE_WIDTH = 256 +VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} +STEPS = 512 +MUP_PATTERN = ( + r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" + r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" +) + def model_registry(flavor: str) -> ModelSpec: return ModelSpec( @@ -375,26 +393,6 @@ def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> ) -HEAD_DIM = 64 -N_LAYER = 8 -INPUT_SIZE = ( - 1, - 128, - 256, -) -PATCH_SIZE = (1, 16, 8) -IN_CHANNELS = 24 -PLAN_SIZE = 15 * 33 * 2 -BASE_WIDTH = 256 -VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} -STEPS = 512 -SWEEP_LR = float(os.getenv("VIT_LR", "3e-4")) -MUP_PATTERN = ( - r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" - r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" -) - - def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Config: return Linear.Config( in_features=in_f, @@ -411,24 +409,16 @@ def _hidden_std(fan_in: int, *, mup: bool) -> float: return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 -def _ln(dim: int) -> LayerNorm.Config: - return LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT) - - -def _hidden(dim: int, mult: float, multiple_of: int = 256) -> int: - return multiple_of * math.ceil(int(dim * mult) / multiple_of) - - def _vit_attention( - dim: int, n_head: int, *, mup: bool, qk_norm: bool = True + dim: int, *, n_head: int, mup: bool, qk_norm: bool = True ) -> PathSelfAttention.Config: head_dim = dim // n_head return PathSelfAttention.Config( - norm=_ln(dim), - q_norm=_ln(head_dim) if qk_norm else None, - k_norm=_ln(head_dim) if qk_norm else None, + norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), + q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) if qk_norm else None, + k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) if qk_norm else None, c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), - c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * N_LAYER)), + c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * NUM_LAYERS)), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, head_dim=head_dim, @@ -438,12 +428,12 @@ def _vit_attention( def _vit_mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: - hidden = _hidden(dim, mult) + hidden = _hidden_dim(dim, mult) return PathMLP.Config( - norm=_ln(dim), + norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), c_proj=_lin( - hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * N_LAYER) + hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * NUM_LAYERS) ), act="gelu_tanh", dropout=0.0, @@ -451,36 +441,35 @@ def _vit_mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: def _vit_model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: - n_embd = VIT_WIDTHS[flavor] - n_head = n_embd // HEAD_DIM + dim = VIT_WIDTHS[flavor] + n_head = dim // HEAD_DIM pt, ph, pw = PATCH_SIZE patch_dim = pt * IN_CHANNELS * ph * pw t, h, w = INPUT_SIZE num_patches = (t // pt) * (h // ph) * (w // pw) return PlanViT.Config( - input_size=INPUT_SIZE, - patch_size=PATCH_SIZE, - in_channels=IN_CHANNELS, - n_embd=n_embd, - output_mult=(BASE_WIDTH / n_embd) if mup else 1.0, + dim=dim, + output_mult=(BASE_WIDTH / dim) if mup else 1.0, + mean=255 / 2, + std=255 / 4, patch_embed=PatchEmbed.Config( - proj=_lin(patch_dim, n_embd, std=patch_dim**-0.5), + proj=_lin(patch_dim, dim, std=patch_dim**-0.5), patch_size=PATCH_SIZE, ), pos_embedding=Embedding.Config( - num_embeddings=num_patches, embedding_dim=n_embd, param_init=_LINEAR_INIT + num_embeddings=num_patches, embedding_dim=dim, param_init=_LINEAR_INIT ), blocks=[ PathTransformerBlock.Config( - attention=_vit_attention(n_embd, n_head, mup=mup, qk_norm=qk_norm), - mlp=_vit_mlp(n_embd, mup=mup), + attention=_vit_attention(dim, n_head=n_head, mup=mup, qk_norm=qk_norm), + mlp=_vit_mlp(dim, mup=mup), ) - for _ in range(N_LAYER) + for _ in range(NUM_LAYERS) ], - norm=_ln(n_embd), + norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), plan_head=PlanHead.Config( - norm=_ln(n_embd), - head=_lin(n_embd, PLAN_SIZE, std=BASE_WIDTH**-0.5), + norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), + head=_lin(dim, PLAN_SIZE, std=BASE_WIDTH**-0.5), ), ) @@ -498,10 +487,6 @@ def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: def _vit_dataloader_config(*, split: str) -> PathDataLoader.Config: - from xx.common.basedir import XX_BASEDIR - from xx.datasets.constants import BASE_DIR_GT_10M - from xx.training.path.config import DatasetConfig as XXPathDatasetConfig - base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) return PathDataLoader.Config( dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), @@ -523,6 +508,8 @@ def _vit_dataloader_config(*, split: str) -> PathDataLoader.Config: def _vit_optimizer_config( flavor: str, *, mup: bool, lr: float, wd: float ) -> OptimizersContainer.Config: + # base lr is carried on the container so --optimizer.lr can sweep every group + # at once; muP scales the hidden matmuls down by lr_mult = 1/m (m = width ratio). m = VIT_WIDTHS[flavor] / BASE_WIDTH common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} if mup: @@ -530,12 +517,13 @@ def _vit_optimizer_config( ParamGroupConfig( pattern=MUP_PATTERN, optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr / m}, + lr_mult=1.0 / m, + optimizer_kwargs={**common}, ), ParamGroupConfig( pattern=r".*", optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr}, + optimizer_kwargs={**common}, ), ] else: @@ -543,16 +531,16 @@ def _vit_optimizer_config( ParamGroupConfig( pattern=r".*", optimizer_name="AdamW", - optimizer_kwargs={**common, "lr": lr}, + optimizer_kwargs={**common}, ) ] return OptimizersContainer.Config( - implementation="fused_opt_states_bf16", param_groups=groups + implementation="fused_opt_states_bf16", lr=lr, param_groups=groups ) def _vit( - flavor: str, *, mup: bool, lr: float = SWEEP_LR, wd: float = 3e-2 + flavor: str, *, mup: bool, lr: float = 3e-4, wd: float = 3e-2 ) -> PathTrainer.Config: local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) @@ -560,7 +548,7 @@ def _vit( os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) ) return PathTrainer.Config( - loss=PathLoss.Config(), + loss=PlanViTLoss.Config(), model_spec=vit_model_registry(flavor, mup=mup), tokenizer=NoOpTokenizer.Config(), dataloader=_vit_dataloader_config(split="train"), @@ -597,10 +585,7 @@ def _vit( mixed_precision_param="bfloat16", ), fps=SUPERCOMBO_FPS, - plan_target_last_frame=True, debug=DebugConfig(seed=0), - built_base_lr=lr, - mup_base_lr=lr, ) diff --git a/torchtitan/experiments/path/trainer.py b/torchtitan/experiments/path/trainer.py index 2cb6655724..d47da3d537 100644 --- a/torchtitan/experiments/path/trainer.py +++ b/torchtitan/experiments/path/trainer.py @@ -27,24 +27,12 @@ class Config(Trainer.Config): checkpoint: PathOnnxCheckpointManager.Config miniray: dict[str, Any] = field(default_factory=dict) fps: int - plan_target_last_frame: bool = False - mup_base_lr: float | None = None - built_base_lr: float | None = None def __post_init__(self) -> None: Trainer.Config.__post_init__(self) if self.codedir: self.miniray = {**self.miniray, "codedir": self.codedir} self.validator.miniray = {**self.validator.miniray, "codedir": self.codedir} - if ( - self.mup_base_lr is not None - and self.built_base_lr is not None - and self.mup_base_lr != self.built_base_lr - ): - ratio = self.mup_base_lr / self.built_base_lr - for group in self.optimizer.param_groups: - group.optimizer_kwargs["lr"] *= ratio - self.built_base_lr = self.mup_base_lr def __init__(self, config: Config): super().__init__(config) @@ -88,8 +76,6 @@ def forward_backward_step( assert len(self.model_parts) == 1 with self.train_context(): pred = self.model_parts[0](inputs) - if self.config.plan_target_last_frame: - labels = {**labels, "plan": labels["plan"][:, -1]} loss_vec, metrics = self.loss_fn(pred, labels) loss = loss_vec.sum() / local_samples del pred diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py index 652136fcf3..ec4c4cd329 100644 --- a/torchtitan/experiments/path/vit.py +++ b/torchtitan/experiments/path/vit.py @@ -1,10 +1,8 @@ from __future__ import annotations from dataclasses import dataclass -from xx.ml_tools.constants.model import ModelInputs import torch -import torch.nn as nn from einops import rearrange from torch.distributed.device_mesh import DeviceMesh from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy @@ -21,8 +19,10 @@ from torchtitan.protocols.model import BaseModel from torchtitan.protocols.module import Module, ModuleList from torchtitan.tools.logging import logger +from xx.ml_tools.constants.model import ModelInputs -from .model import PathMLP, PathSelfAttention, PathTransformerBlock +from .loss import PathLoss +from .model import PathTransformerBlock class PatchEmbed(Module): @@ -62,23 +62,16 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class PlanViT(BaseModel): @dataclass(kw_only=True, slots=True) class Config(BaseModel.Config): - input_size: tuple[int, int, int] - patch_size: tuple[int, int, int] - in_channels: int - n_embd: int + dim: int output_mult: float + mean: float + std: float patch_embed: PatchEmbed.Config pos_embedding: Embedding.Config blocks: list[PathTransformerBlock.Config] norm: LayerNorm.Config | RMSNorm.Config plan_head: PlanHead.Config - @property - def num_patches(self) -> int: - t, h, w = self.input_size - pt, ph, pw = self.patch_size - return (t // pt) * (h // ph) * (w // pw) - def update_from_config(self, *, config, **kwargs) -> None: parallelism = config.parallelism for name, degree in { @@ -106,17 +99,16 @@ def __init__(self, config: Config): def verify_module_protocol(self) -> None: pass - def _frames(self, inputs: dict[str, torch.Tensor] | torch.Tensor) -> torch.Tensor: - if isinstance(inputs, torch.Tensor): - return inputs - img, big = inputs[ModelInputs.IMG], inputs[ModelInputs.BIG_IMG] - frame = torch.cat([img[:, -1], big[:, -1]], dim=1).unsqueeze(1) - return (frame.float() - 127.5) / 63.75 - def forward( self, inputs: dict[str, torch.Tensor] | torch.Tensor ) -> dict[str, torch.Tensor]: - x = self.patch_embed(self._frames(inputs)) + if isinstance(inputs, torch.Tensor): + frame = inputs + else: + img, big = inputs[ModelInputs.IMG], inputs[ModelInputs.BIG_IMG] + frame = torch.cat([img[:, -1], big[:, -1]], dim=1).unsqueeze(1) + frame = (frame.float() - self.config.mean) / self.config.std + x = self.patch_embed(frame) pos = self.pos_embedding(torch.arange(x.shape[1], device=x.device)) x = x + rearrange(pos, "t c -> () t c") for block in self.blocks: @@ -125,6 +117,27 @@ def forward( return {"plan": self.plan_head(x.mean(dim=1)) * self.config.output_mult} +class PlanViTLoss(PathLoss): + """PathLoss for the single-frame plan ViT. + + PlanViT predicts one frame, so the temporal plan target is reduced to its + last frame before scoring. Everything else matches PathLoss. + """ + + @dataclass(kw_only=True, slots=True) + class Config(PathLoss.Config): + pass + + def __call__( + self, + pred: dict[str, torch.Tensor], + targets: dict[str, torch.Tensor], + global_valid_tokens: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: + targets = {**targets, "plan": targets["plan"][:, -1]} + return super().__call__(pred, targets, global_valid_tokens) + + def parallelize_vit( model: PlanViT, *, From a83821bf59887f9749242454e5ea8769fa3830cc Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 16:42:44 -0700 Subject: [PATCH 16/28] path: trim vit muP config surface --- .../experiments/path/config_registry.py | 29 ++++++++----------- torchtitan/experiments/path/vit.py | 12 -------- 2 files changed, 12 insertions(+), 29 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 1c8b304205..9c904cbbc5 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -78,7 +78,7 @@ IN_CHANNELS = 24 PLAN_SIZE = 15 * 33 * 2 BASE_WIDTH = 256 -VIT_WIDTHS = {"w128": 128, "w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} +VIT_WIDTHS = {"w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} STEPS = 512 MUP_PATTERN = ( r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" @@ -448,7 +448,6 @@ def _vit_model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanVi t, h, w = INPUT_SIZE num_patches = (t // pt) * (h // ph) * (w // pw) return PlanViT.Config( - dim=dim, output_mult=(BASE_WIDTH / dim) if mup else 1.0, mean=255 / 2, std=255 / 4, @@ -512,28 +511,24 @@ def _vit_optimizer_config( # at once; muP scales the hidden matmuls down by lr_mult = 1/m (m = width ratio). m = VIT_WIDTHS[flavor] / BASE_WIDTH common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} + groups = [ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common}, + ) + ] if mup: - groups = [ + # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all + groups.insert( + 0, ParamGroupConfig( pattern=MUP_PATTERN, optimizer_name="AdamW", lr_mult=1.0 / m, optimizer_kwargs={**common}, ), - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common}, - ), - ] - else: - groups = [ - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common}, - ) - ] + ) return OptimizersContainer.Config( implementation="fused_opt_states_bf16", lr=lr, param_groups=groups ) diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py index ec4c4cd329..6a222f5fbd 100644 --- a/torchtitan/experiments/path/vit.py +++ b/torchtitan/experiments/path/vit.py @@ -62,7 +62,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class PlanViT(BaseModel): @dataclass(kw_only=True, slots=True) class Config(BaseModel.Config): - dim: int output_mult: float mean: float std: float @@ -72,17 +71,6 @@ class Config(BaseModel.Config): norm: LayerNorm.Config | RMSNorm.Config plan_head: PlanHead.Config - def update_from_config(self, *, config, **kwargs) -> None: - parallelism = config.parallelism - for name, degree in { - "tensor parallel": parallelism.tensor_parallel_degree, - "context parallel": parallelism.context_parallel_degree, - "pipeline parallel": parallelism.pipeline_parallel_degree, - "expert parallel": parallelism.expert_parallel_degree, - }.items(): - if degree > 1: - raise ValueError(f"PlanViT does not support {name}") - def get_nparams_and_flops(self, model: Module, seq_len: int) -> tuple[int, int]: nparams = sum(p.numel() for p in model.parameters()) return nparams, 6 * nparams From 63809b3225aab2b63d3935f572e2ba0806a026fb Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 17:42:10 -0700 Subject: [PATCH 17/28] path: license headers and ufmt formatting --- .../experiments/path/config_registry.py | 181 ++++++++++++------ torchtitan/experiments/path/model.py | 122 +++++++++--- torchtitan/experiments/path/vit.py | 8 +- 3 files changed, 229 insertions(+), 82 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 9c904cbbc5..78a9f3a783 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -1,8 +1,33 @@ +# 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 import os from functools import partial +from xx.common.basedir import XX_BASEDIR +from xx.datasets.constants import BASE_DIR_GT_10M +from xx.datasets.helpers import DEFAULT_BIG_TRAIN_LIST +from xx.ml_tools.constants.model import ( + frame_constants_from_fps, + FRAME_TYPE, + INPUT_FRAMES_NAMES, + ModelInputs, + N_FRAMES, + SUPERCOMBO_FPS, + TEMPORAL_INPUTS, +) +from xx.training.path.config import DatasetConfig as XXPathDatasetConfig +from xx.training.path.hydra_configs import ( + DRIVING_HEADS, + META_HEADS, + POSE_HEADS, + TEMPORAL_META_HEADS, +) import torch.nn as nn @@ -21,25 +46,13 @@ from torchtitan.models.common import Embedding, LayerNorm, Linear from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model_spec import ModelSpec -from xx.common.basedir import XX_BASEDIR -from xx.datasets.constants import BASE_DIR_GT_10M -from xx.datasets.helpers import DEFAULT_BIG_TRAIN_LIST -from xx.ml_tools.constants.model import ( - SUPERCOMBO_FPS, - FRAME_TYPE, - INPUT_FRAMES_NAMES, - N_FRAMES, - TEMPORAL_INPUTS, - ModelInputs, - frame_constants_from_fps, -) -from xx.training.path.config import DatasetConfig as XXPathDatasetConfig -from xx.training.path.hydra_configs import DRIVING_HEADS, META_HEADS, POSE_HEADS, TEMPORAL_META_HEADS from .dataset import PathDataLoader +from .loss import PathLoss from .model import ( Hydra, LinearEncoder, + parallelize_path, PathHead, PathMLP, PathModel, @@ -52,22 +65,17 @@ TemporalPolicy, TemporalSummarizer, Vision, - parallelize_path, ) -from .loss import PathLoss from .onnx_checkpoint import PathOnnxCheckpointManager from .trainer import PathTrainer from .validate import PathValidator -from .vit import ( - parallelize_vit, - PatchEmbed, - PlanHead, - PlanViT, - PlanViTLoss, -) +from .vit import parallelize_vit, PatchEmbed, PlanHead, PlanViT, PlanViTLoss -_LINEAR_INIT = {"weight": partial(nn.init.normal_, mean=0.0, std=0.02), "bias": nn.init.zeros_} +_LINEAR_INIT = { + "weight": partial(nn.init.normal_, mean=0.0, std=0.02), + "bias": nn.init.zeros_, +} _NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} # PlanViT (single-frame plan ViT) architecture and muP constants @@ -115,10 +123,10 @@ def convnext_xxlarge() -> PathTrainer.Config: def _path(flavor: str) -> PathTrainer.Config: - steps = 1024*100 + steps = 1024 * 100 validation_freq = 1024 reports = { - name: [validation_freq, steps //2 , steps] + name: [validation_freq, steps // 2, steps] for name in ( "analyse_driving", "analyse_lat_no_noise", @@ -132,10 +140,14 @@ def _path(flavor: str) -> PathTrainer.Config: mixed_precision_param = "bfloat16" local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int(os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size))) + num_nodes = int( + os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) + ) reporterv2_host = os.getenv("REPORTERV2_HOST") reporterv2_training_id = os.getenv("REPORTERV2_TRAINING_ID") - checkpoint_base_folder = f"{reporterv2_host.rstrip('/')}/checkpoint" if reporterv2_host else "" + checkpoint_base_folder = ( + f"{reporterv2_host.rstrip('/')}/checkpoint" if reporterv2_host else "" + ) fps = SUPERCOMBO_FPS plan_only = False return PathTrainer.Config( @@ -178,7 +190,9 @@ def _path(flavor: str) -> PathTrainer.Config: fps=fps, activation_checkpoint=FullAC.Config(), compile=CompileConfig(enable=True, components=["model"]), - metrics=MetricsProcessor.Config(log_freq=16, enable_reporterv2=True, save_freq=validation_freq), + metrics=MetricsProcessor.Config( + log_freq=16, enable_reporterv2=True, save_freq=validation_freq + ), validator=PathValidator.Config( enable=True, freq=validation_freq, @@ -196,8 +210,12 @@ def _model_config(flavor: str) -> PathModel.Config: n_frames_input = N_FRAMES input_frame_names = INPUT_FRAMES_NAMES input_frame_type = FRAME_TYPE - frame_constants = frame_constants_from_fps(n_frames=n_frames_input, frame_type=input_frame_type) - in_channels = sum(frame_constants["frame_shapes"][name][0] for name in input_frame_names) + frame_constants = frame_constants_from_fps( + n_frames=n_frames_input, frame_type=input_frame_type + ) + in_channels = sum( + frame_constants["frame_shapes"][name][0] for name in input_frame_names + ) block_size = len(frame_constants["history_idxs"]) temporal_len = frame_constants["temporal_len"] dim = vision_features @@ -227,9 +245,13 @@ def _model_config(flavor: str) -> PathModel.Config: temporal_summarizer=TemporalSummarizer.Config( mlp1=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), mlp2=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - desire_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * temporal_len, dim), + desire_encoder=_encoder( + TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * temporal_len, dim + ), traffic_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.TRAFFIC][0], dim), - action_t_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], dim), + action_t_encoder=_encoder( + TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], dim + ), transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( @@ -239,17 +261,25 @@ def _model_config(flavor: str) -> PathModel.Config: for _ in range(4) ] ), - pos_embedding=Embedding.Config(num_embeddings=block_size, embedding_dim=dim, param_init=_LINEAR_INIT), + pos_embedding=Embedding.Config( + num_embeddings=block_size, + embedding_dim=dim, + param_init=_LINEAR_INIT, + ), block_size=block_size, dense_training_outputs=True, ), - temporal_hydra=_hydra(_heads(DRIVING_HEADS + TEMPORAL_META_HEADS), in_features=dim, mlp_mult=2), + temporal_hydra=_hydra( + _heads(DRIVING_HEADS + TEMPORAL_META_HEADS), in_features=dim, mlp_mult=2 + ), history_idxs=tuple(int(x) for x in frame_constants["history_idxs"]), ), ) -def _dataloader_config(*, split: str, fps: int, plan_only: bool) -> PathDataLoader.Config: +def _dataloader_config( + *, split: str, fps: int, plan_only: bool +) -> PathDataLoader.Config: base = XXPathDatasetConfig(fps=fps, plan_only=plan_only) return PathDataLoader.Config( dataset=DEFAULT_BIG_TRAIN_LIST, @@ -268,7 +298,9 @@ def _dataloader_config(*, split: str, fps: int, plan_only: bool) -> PathDataLoad ) -def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnxCheckpointManager.Config: +def _checkpoint_config( + folder: str, base_folder: str, interval: int +) -> PathOnnxCheckpointManager.Config: frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) temporal_len = frame_constants["temporal_len"] vision_input_names = [ModelInputs.IMG, ModelInputs.BIG_IMG] @@ -291,10 +323,10 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.ACTION_T][0]], ] return PathOnnxCheckpointManager.Config( - keep_latest_k=0, # keep all checkpoints + keep_latest_k=0, # keep all checkpoints enable=True, checkpoint_base_folder=base_folder, - save_model_state_dict=True, # another copy of full state dict + save_model_state_dict=True, # another copy of full state dict export_onnx=True, enable_first_step_checkpoint=True, folder=folder, @@ -302,7 +334,7 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx input_names=input_names, input_shapes=input_shapes, input_dtypes=["float16"] * len(input_names), - onnx_model_dtype="float16", # WIP: test if fp16 doesn't degrade performance + onnx_model_dtype="float16", # WIP: test if fp16 doesn't degrade performance vision_input_names=vision_input_names, temporal_policy_input_names=temporal_policy_input_names, ) @@ -311,7 +343,11 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx def _si_int(value: str | int) -> int: suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} value = str(value).strip().lower() - return int(float(value[:-1]) * suffixes[value[-1]]) if value[-1] in suffixes else int(value) + return ( + int(float(value[:-1]) * suffixes[value[-1]]) + if value[-1] in suffixes + else int(value) + ) def _optimizer_config() -> OptimizersContainer.Config: @@ -335,7 +371,9 @@ def _optimizer_config() -> OptimizersContainer.Config: def _heads(heads) -> tuple[PathHead, ...]: - return tuple(PathHead(head.name, head.output_size, head.mlp, head.scale) for head in heads) + return tuple( + PathHead(head.name, head.output_size, head.mlp, head.scale) for head in heads + ) def _hidden_dim(dim: int, mlp_mult: float, multiple_of: int = 256) -> int: @@ -347,8 +385,12 @@ def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Co hidden = _hidden_dim(dim, mlp_mult) return PathMLP.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - c_fc=Linear.Config(in_features=dim, out_features=hidden, bias=bias, param_init=_LINEAR_INIT), - c_proj=Linear.Config(in_features=hidden, out_features=dim, bias=bias, param_init=_LINEAR_INIT), + c_fc=Linear.Config( + in_features=dim, out_features=hidden, bias=bias, param_init=_LINEAR_INIT + ), + c_proj=Linear.Config( + in_features=hidden, out_features=dim, bias=bias, param_init=_LINEAR_INIT + ), act="gelu_tanh", dropout=dropout, ) @@ -356,8 +398,15 @@ def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Co def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: return LinearEncoder.Config( - in_layer=Linear.Config(in_features=in_features, out_features=dim, bias=True, param_init=_LINEAR_INIT), - out_layer=Linear.Config(in_features=dim, out_features=dim, bias=False, param_init=_LINEAR_INIT), + in_layer=Linear.Config( + in_features=in_features, + out_features=dim, + bias=True, + param_init=_LINEAR_INIT, + ), + out_layer=Linear.Config( + in_features=dim, out_features=dim, bias=False, param_init=_LINEAR_INIT + ), ) @@ -367,8 +416,12 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), - c_attn=Linear.Config(in_features=dim, out_features=3 * dim, bias=True, param_init=_LINEAR_INIT), - c_proj=Linear.Config(in_features=dim, out_features=dim, bias=True, param_init=_LINEAR_INIT), + c_attn=Linear.Config( + in_features=dim, out_features=3 * dim, bias=True, param_init=_LINEAR_INIT + ), + c_proj=Linear.Config( + in_features=dim, out_features=dim, bias=True, param_init=_LINEAR_INIT + ), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, head_dim=head_dim, @@ -376,10 +429,16 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co ) -def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> Hydra.Config: +def _hydra( + heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float +) -> Hydra.Config: return Hydra.Config( heads=heads, - head_mlps={head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) for head in heads if head.mlp}, + head_mlps={ + head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) + for head in heads + if head.mlp + }, final_layers={ head.name: Linear.Config( in_features=in_features, @@ -389,7 +448,11 @@ def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> ) for head in heads }, - scale_layers={head.name: ScaleLayer.Config(n_features=head.output_size) for head in heads if head.scale}, + scale_layers={ + head.name: ScaleLayer.Config(n_features=head.output_size) + for head in heads + if head.scale + }, ) @@ -415,10 +478,16 @@ def _vit_attention( head_dim = dim // n_head return PathSelfAttention.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) if qk_norm else None, - k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) if qk_norm else None, + q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) + if qk_norm + else None, + k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) + if qk_norm + else None, c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), - c_proj=_lin(dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * NUM_LAYERS)), + c_proj=_lin( + dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * NUM_LAYERS) + ), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, head_dim=head_dim, @@ -440,7 +509,9 @@ def _vit_mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: ) -def _vit_model_config(flavor: str, *, mup: bool, qk_norm: bool = True) -> PlanViT.Config: +def _vit_model_config( + flavor: str, *, mup: bool, qk_norm: bool = True +) -> PlanViT.Config: dim = VIT_WIDTHS[flavor] n_head = dim // HEAD_DIM pt, ph, pw = PATCH_SIZE diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 4fbbb2aec0..17fd24262c 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -1,13 +1,20 @@ +# 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 from dataclasses import dataclass +from xx.ml_tools.constants.model import ModelInputs import torch import torch.nn as nn from einops import rearrange from torch.distributed.device_mesh import DeviceMesh from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy -from torch.distributed.tensor import DTensor, distribute_tensor +from torch.distributed.tensor import distribute_tensor, DTensor from torchtitan.config import ( CompileConfig, @@ -22,13 +29,15 @@ FullAC, MemoryBudgetAC, ) -from torchtitan.distributed.fsdp import enable_fsdp_symm_mem, get_fsdp_reshard_after_forward_policy +from torchtitan.distributed.fsdp import ( + enable_fsdp_symm_mem, + get_fsdp_reshard_after_forward_policy, +) from torchtitan.models.common import Embedding, LayerNorm, Linear, RMSNorm, SiLU from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model import BaseModel from torchtitan.protocols.module import Module, ModuleDict, ModuleList, Sequential from torchtitan.tools.logging import logger -from xx.ml_tools.constants.model import ModelInputs from . import convnext @@ -70,7 +79,9 @@ def __init__(self, config: Config): super().__init__() self.norm = config.norm.build() self.c_fc = config.c_fc.build() - self.act = nn.GELU(approximate="tanh") if config.act == "gelu_tanh" else nn.GELU() + self.act = ( + nn.GELU(approximate="tanh") if config.act == "gelu_tanh" else nn.GELU() + ) self.c_proj = config.c_proj.build() self.dropout = nn.Dropout(config.dropout) @@ -98,8 +109,12 @@ def __init__(self, config: Config): self.head_dim = config.head_dim self.is_causal = config.is_causal self.norm = config.norm.build() - self.q_norm = config.q_norm.build() if config.q_norm is not None else nn.Identity() - self.k_norm = config.k_norm.build() if config.k_norm is not None else nn.Identity() + self.q_norm = ( + config.q_norm.build() if config.q_norm is not None else nn.Identity() + ) + self.k_norm = ( + config.k_norm.build() if config.k_norm is not None else nn.Identity() + ) self.c_attn = config.c_attn.build() self.c_proj = config.c_proj.build() self.inner_attention = config.inner_attention.build() @@ -177,7 +192,9 @@ class Config(Module.Config): def __init__(self, config: Config): super().__init__() - self.net = Sequential(config.in_layer.build(), SiLU.Config().build(), config.out_layer.build()) + self.net = Sequential( + config.in_layer.build(), SiLU.Config().build(), config.out_layer.build() + ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.net(x) @@ -218,10 +235,18 @@ def forward( feats = self.mlp1(feats) + feats feats = self.mlp2(feats) + feats desire = rearrange(self.desire_encoder(desire), "b c -> b () c") - traffic_convention = rearrange(self.traffic_encoder(traffic_convention), "b c -> b () c") + traffic_convention = rearrange( + self.traffic_encoder(traffic_convention), "b c -> b () c" + ) action_t = rearrange(self.action_t_encoder(action_t), "b c -> b () c") pos = self.pos_embedding(torch.arange(self.block_size, device=feats.device)) - x = feats + rearrange(pos, "t c -> () t c") + desire + traffic_convention + action_t + x = ( + feats + + rearrange(pos, "t c -> () t c") + + desire + + traffic_convention + + action_t + ) x = self.transformer(x) return x if self.dense_training_outputs else x[:, self.block_size - 1] @@ -237,14 +262,24 @@ class Config(Module.Config): def __init__(self, config: Config): super().__init__() self.heads = config.heads - self.head_mlp = ModuleDict({name: cfg.build() for name, cfg in config.head_mlps.items()}) - self.final_layer = ModuleDict({name: cfg.build() for name, cfg in config.final_layers.items()}) - self.scale_layer = ModuleDict({name: cfg.build() for name, cfg in config.scale_layers.items()}) + self.head_mlp = ModuleDict( + {name: cfg.build() for name, cfg in config.head_mlps.items()} + ) + self.final_layer = ModuleDict( + {name: cfg.build() for name, cfg in config.final_layers.items()} + ) + self.scale_layer = ModuleDict( + {name: cfg.build() for name, cfg in config.scale_layers.items()} + ) def forward(self, in_feats: torch.Tensor) -> dict[str, torch.Tensor]: ret = {} for name, layer in self.final_layer.items(): - feats = self.head_mlp[name](in_feats) + in_feats if name in self.head_mlp else in_feats + feats = ( + self.head_mlp[name](in_feats) + in_feats + if name in self.head_mlp + else in_feats + ) ret[name] = layer(feats) for name, layer in self.scale_layer.items(): ret[name] = layer(ret[name]) @@ -278,11 +313,19 @@ def __init__(self, config: Config): self.config = config self.temporal_summarizer = config.temporal_summarizer.build() self.temporal_hydra = config.temporal_hydra.build() - self.register_buffer("history_idxs", torch.tensor(config.history_idxs, dtype=torch.long), persistent=False) + self.register_buffer( + "history_idxs", + torch.tensor(config.history_idxs, dtype=torch.long), + persistent=False, + ) def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None: - device = buffer_device if buffer_device is not None else self.history_idxs.device - self.history_idxs = torch.tensor(self.config.history_idxs, dtype=torch.long, device=device) + device = ( + buffer_device if buffer_device is not None else self.history_idxs.device + ) + self.history_idxs = torch.tensor( + self.config.history_idxs, dtype=torch.long, device=device + ) def forward( self, @@ -324,13 +367,21 @@ def __init__(self, config: Config): num_classes=config.vision_features, drop_path_rate=config.drop_path_rate, ) - self.register_buffer("_mean", torch.empty(1, config.in_channels, 1, 1), persistent=True) - self.register_buffer("_std", torch.empty(1, config.in_channels, 1, 1), persistent=True) + self.register_buffer( + "_mean", torch.empty(1, config.in_channels, 1, 1), persistent=True + ) + self.register_buffer( + "_std", torch.empty(1, config.in_channels, 1, 1), persistent=True + ) def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None: device = buffer_device if buffer_device is not None else self._mean.device - self._mean = torch.full((1, self.config.in_channels, 1, 1), self.config.mean, device=device) - self._std = torch.full((1, self.config.in_channels, 1, 1), self.config.std, device=device) + self._mean = torch.full( + (1, self.config.in_channels, 1, 1), self.config.mean, device=device + ) + self._std = torch.full( + (1, self.config.in_channels, 1, 1), self.config.std, device=device + ) def load_pretrained(self) -> None: if not self.config.pretrained: @@ -366,13 +417,21 @@ def _pretrained_state_dict(self) -> dict[str, torch.Tensor]: ) state_dict = convnext.checkpoint_filter_fn(state_dict, self.encoder) if self.config.in_channels != 3: - state_dict["stem.0.weight"] = adapt_input_conv(self.config.in_channels, state_dict["stem.0.weight"]) + state_dict["stem.0.weight"] = adapt_input_conv( + self.config.in_channels, state_dict["stem.0.weight"] + ) return state_dict @staticmethod - def _move_pretrained_value(value: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + def _move_pretrained_value( + value: torch.Tensor, target: torch.Tensor + ) -> torch.Tensor: if isinstance(target, DTensor): - return distribute_tensor(value.to(dtype=target.dtype), target.device_mesh, list(target.placements)) + return distribute_tensor( + value.to(dtype=target.dtype), + target.device_mesh, + list(target.placements), + ) return value.to(device=target.device, dtype=target.dtype) def forward(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: @@ -478,10 +537,17 @@ def parallelize_path( ) -> PathModel: if parallelism.spmd_backend == "full_dtensor": raise ValueError("path v1 does not support full DTensor") - if parallel_dims.tp_enabled or parallel_dims.cp_enabled or parallel_dims.pp_enabled or parallel_dims.ep_enabled: + if ( + parallel_dims.tp_enabled + or parallel_dims.cp_enabled + or parallel_dims.pp_enabled + or parallel_dims.ep_enabled + ): raise ValueError("path v1 supports data parallelism only") - model_compile_enabled = compile_config.enable and "model" in compile_config.components + model_compile_enabled = ( + compile_config.enable and "model" in compile_config.components + ) if ac_config is not None: _apply_activation_checkpointing(model, ac_config, dump_folder=dump_folder) @@ -500,7 +566,11 @@ def parallelize_path( enable_symm_mem=parallelism.enable_fsdp_symm_mem, ) - logger.info("Applied HSDP to the path model" if parallel_dims.dp_replicate_enabled else "Applied FSDP to the path model") + logger.info( + "Applied HSDP to the path model" + if parallel_dims.dp_replicate_enabled + else "Applied FSDP to the path model" + ) if training.enable_cpu_offload: logger.info("Applied CPU Offloading to the path model") return model diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py index 6a222f5fbd..62a5f1b3ad 100644 --- a/torchtitan/experiments/path/vit.py +++ b/torchtitan/experiments/path/vit.py @@ -1,6 +1,13 @@ +# 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 from dataclasses import dataclass +from xx.ml_tools.constants.model import ModelInputs import torch from einops import rearrange @@ -19,7 +26,6 @@ from torchtitan.protocols.model import BaseModel from torchtitan.protocols.module import Module, ModuleList from torchtitan.tools.logging import logger -from xx.ml_tools.constants.model import ModelInputs from .loss import PathLoss from .model import PathTransformerBlock From 08450a53c3f219e83587acc7c3f50595fe9608cf Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 18:15:24 -0700 Subject: [PATCH 18/28] optimizer: test base lr with per-group lr_mult --- .../unit_tests/test_optimizer_param_groups.py | 61 +++++++++++++++++++ 1 file changed, 61 insertions(+) diff --git a/tests/unit_tests/test_optimizer_param_groups.py b/tests/unit_tests/test_optimizer_param_groups.py index 58d3e7d0df..a4ade9965d 100644 --- a/tests/unit_tests/test_optimizer_param_groups.py +++ b/tests/unit_tests/test_optimizer_param_groups.py @@ -263,6 +263,67 @@ def test_lr_override(self): # Default group self.assertEqual(groups[1]["lr"], 1e-3) + def test_base_lr_with_lr_mult(self): + """A base lr scales each group by its lr_mult.""" + model = SimpleModel() + config = OptimizersContainer.Config( + implementation="for-loop", + lr=1e-3, + param_groups=[ + ParamGroupConfig( + pattern=r"embed_tokens\.", + optimizer_name="AdamW", + optimizer_kwargs={"weight_decay": 0.1}, + lr_mult=0.25, + ), + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={"weight_decay": 0.1}, + lr_mult=1.0, + ), + ], + ) + opt = config.build(model_parts=[model]).optimizers[0] + + self.assertAlmostEqual(opt.param_groups[0]["lr"], 1e-3 * 0.25) + self.assertAlmostEqual(opt.param_groups[1]["lr"], 1e-3 * 1.0) + + def test_base_lr_default_lr_mult_unscaled(self): + """lr_mult defaults to 1.0 so a single base lr applies unscaled.""" + model = SimpleModel() + config = OptimizersContainer.Config( + implementation="for-loop", + lr=5e-4, + param_groups=[ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={"weight_decay": 0.1}, + ), + ], + ) + opt = config.build(model_parts=[model]).optimizers[0] + + self.assertAlmostEqual(opt.param_groups[0]["lr"], 5e-4) + + def test_base_lr_rejects_group_lr(self): + """Setting lr in optimizer_kwargs while a base lr is set raises ValueError.""" + model = SimpleModel() + config = OptimizersContainer.Config( + implementation="for-loop", + lr=1e-3, + param_groups=[ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={"lr": 1e-4, "weight_decay": 0.1}, + ), + ], + ) + with self.assertRaises(ValueError): + config.build(model_parts=[model]) + def test_first_match_wins(self): """When patterns overlap, the first match wins.""" model = SimpleModel() From 6fe28a91239f2bc29761f2a064cf822eccf46481 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 18:15:24 -0700 Subject: [PATCH 19/28] path: extract muP into a reusable module --- .../experiments/path/config_registry.py | 48 +++++---------- torchtitan/experiments/path/mup.py | 59 +++++++++++++++++++ 2 files changed, 75 insertions(+), 32 deletions(-) create mode 100644 torchtitan/experiments/path/mup.py diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 78a9f3a783..ea02d23354 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -66,6 +66,7 @@ TemporalSummarizer, Vision, ) +from .mup import hidden_std, MuPSpec, output_mult, param_groups, residual_std from .onnx_checkpoint import PathOnnxCheckpointManager from .trainer import PathTrainer from .validate import PathValidator @@ -92,6 +93,9 @@ r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" ) +VIT_MUP = MuPSpec( + base_width=BASE_WIDTH, num_layers=NUM_LAYERS, hidden_pattern=MUP_PATTERN +) def model_registry(flavor: str) -> ModelSpec: @@ -468,10 +472,6 @@ def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Conf ) -def _hidden_std(fan_in: int, *, mup: bool) -> float: - return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 - - def _vit_attention( dim: int, *, n_head: int, mup: bool, qk_norm: bool = True ) -> PathSelfAttention.Config: @@ -484,10 +484,8 @@ def _vit_attention( k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) if qk_norm else None, - c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), - c_proj=_lin( - dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * NUM_LAYERS) - ), + c_attn=_lin(dim, 3 * dim, std=hidden_std(dim, VIT_MUP, mup=mup)), + c_proj=_lin(dim, dim, std=residual_std(dim, VIT_MUP, mup=mup)), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, head_dim=head_dim, @@ -500,10 +498,8 @@ def _vit_mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: hidden = _hidden_dim(dim, mult) return PathMLP.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), - c_proj=_lin( - hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * NUM_LAYERS) - ), + c_fc=_lin(dim, hidden, std=hidden_std(dim, VIT_MUP, mup=mup)), + c_proj=_lin(hidden, dim, std=residual_std(hidden, VIT_MUP, mup=mup)), act="gelu_tanh", dropout=0.0, ) @@ -519,7 +515,7 @@ def _vit_model_config( t, h, w = INPUT_SIZE num_patches = (t // pt) * (h // ph) * (w // pw) return PlanViT.Config( - output_mult=(BASE_WIDTH / dim) if mup else 1.0, + output_mult=output_mult(dim, VIT_MUP, mup=mup), mean=255 / 2, std=255 / 4, patch_embed=PatchEmbed.Config( @@ -580,26 +576,14 @@ def _vit_optimizer_config( ) -> OptimizersContainer.Config: # base lr is carried on the container so --optimizer.lr can sweep every group # at once; muP scales the hidden matmuls down by lr_mult = 1/m (m = width ratio). - m = VIT_WIDTHS[flavor] / BASE_WIDTH common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} - groups = [ - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common}, - ) - ] - if mup: - # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all - groups.insert( - 0, - ParamGroupConfig( - pattern=MUP_PATTERN, - optimizer_name="AdamW", - lr_mult=1.0 / m, - optimizer_kwargs={**common}, - ), - ) + groups = param_groups( + VIT_WIDTHS[flavor], + VIT_MUP, + mup=mup, + optimizer_name="AdamW", + optimizer_kwargs=common, + ) return OptimizersContainer.Config( implementation="fused_opt_states_bf16", lr=lr, param_groups=groups ) diff --git a/torchtitan/experiments/path/mup.py b/torchtitan/experiments/path/mup.py new file mode 100644 index 0000000000..7ef0e1e414 --- /dev/null +++ b/torchtitan/experiments/path/mup.py @@ -0,0 +1,59 @@ +# 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. + +import math +from dataclasses import dataclass + +from torchtitan.components.optimizer import ParamGroupConfig + + +@dataclass(frozen=True) +class MuPSpec: + base_width: int + num_layers: int + hidden_pattern: str + + +def hidden_std(fan_in: int, spec: MuPSpec, *, mup: bool) -> float: + return fan_in**-0.5 if mup else spec.base_width**-0.5 + + +def residual_std(fan_in: int, spec: MuPSpec, *, mup: bool) -> float: + return hidden_std(fan_in, spec, mup=mup) / math.sqrt(2 * spec.num_layers) + + +def output_mult(width: int, spec: MuPSpec, *, mup: bool) -> float: + return (spec.base_width / width) if mup else 1.0 + + +def param_groups( + width: int, + spec: MuPSpec, + *, + mup: bool, + optimizer_name: str, + optimizer_kwargs: dict, +) -> list[ParamGroupConfig]: + groups = [ + ParamGroupConfig( + pattern=r".*", + optimizer_name=optimizer_name, + optimizer_kwargs=dict(optimizer_kwargs), + ) + ] + if mup: + # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all + m = width / spec.base_width + groups.insert( + 0, + ParamGroupConfig( + pattern=spec.hidden_pattern, + optimizer_name=optimizer_name, + lr_mult=1.0 / m, + optimizer_kwargs=dict(optimizer_kwargs), + ), + ) + return groups From 4a3b5a6cc60ee9466a9d1b11d63e3545b4e26131 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Fri, 26 Jun 2026 19:32:22 -0700 Subject: [PATCH 20/28] mup: move to components as a reusable technique --- torchtitan/{experiments/path => components}/mup.py | 0 torchtitan/experiments/path/config_registry.py | 8 +++++++- 2 files changed, 7 insertions(+), 1 deletion(-) rename torchtitan/{experiments/path => components}/mup.py (100%) diff --git a/torchtitan/experiments/path/mup.py b/torchtitan/components/mup.py similarity index 100% rename from torchtitan/experiments/path/mup.py rename to torchtitan/components/mup.py diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index ea02d23354..76012f0b2c 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -34,6 +34,13 @@ from torchtitan.components.checkpoint import CheckpointManager from torchtitan.components.lr_scheduler import LRSchedulersContainer from torchtitan.components.metrics import MetricsProcessor +from torchtitan.components.mup import ( + hidden_std, + MuPSpec, + output_mult, + param_groups, + residual_std, +) from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig from torchtitan.components.tokenizer import NoOpTokenizer from torchtitan.config import ( @@ -66,7 +73,6 @@ TemporalSummarizer, Vision, ) -from .mup import hidden_std, MuPSpec, output_mult, param_groups, residual_std from .onnx_checkpoint import PathOnnxCheckpointManager from .trainer import PathTrainer from .validate import PathValidator From a174e6f2228d2e43bb660f9fc433b550eb65db77 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Sat, 27 Jun 2026 17:07:02 -0700 Subject: [PATCH 21/28] mup: move parametrization to experiments/mup, trim comments --- .../{components/mup.py => experiments/mup/parametrization.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename torchtitan/{components/mup.py => experiments/mup/parametrization.py} (100%) diff --git a/torchtitan/components/mup.py b/torchtitan/experiments/mup/parametrization.py similarity index 100% rename from torchtitan/components/mup.py rename to torchtitan/experiments/mup/parametrization.py From c81fe060a95c56433f08a7eb149365bfb60a909e Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Mon, 29 Jun 2026 10:21:49 -0700 Subject: [PATCH 22/28] path: re-inline vit muP, drop the muP module --- torchtitan/experiments/mup/parametrization.py | 59 ------------------- .../experiments/path/config_registry.py | 54 ++++++++++------- 2 files changed, 32 insertions(+), 81 deletions(-) delete mode 100644 torchtitan/experiments/mup/parametrization.py diff --git a/torchtitan/experiments/mup/parametrization.py b/torchtitan/experiments/mup/parametrization.py deleted file mode 100644 index 7ef0e1e414..0000000000 --- a/torchtitan/experiments/mup/parametrization.py +++ /dev/null @@ -1,59 +0,0 @@ -# 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. - -import math -from dataclasses import dataclass - -from torchtitan.components.optimizer import ParamGroupConfig - - -@dataclass(frozen=True) -class MuPSpec: - base_width: int - num_layers: int - hidden_pattern: str - - -def hidden_std(fan_in: int, spec: MuPSpec, *, mup: bool) -> float: - return fan_in**-0.5 if mup else spec.base_width**-0.5 - - -def residual_std(fan_in: int, spec: MuPSpec, *, mup: bool) -> float: - return hidden_std(fan_in, spec, mup=mup) / math.sqrt(2 * spec.num_layers) - - -def output_mult(width: int, spec: MuPSpec, *, mup: bool) -> float: - return (spec.base_width / width) if mup else 1.0 - - -def param_groups( - width: int, - spec: MuPSpec, - *, - mup: bool, - optimizer_name: str, - optimizer_kwargs: dict, -) -> list[ParamGroupConfig]: - groups = [ - ParamGroupConfig( - pattern=r".*", - optimizer_name=optimizer_name, - optimizer_kwargs=dict(optimizer_kwargs), - ) - ] - if mup: - # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all - m = width / spec.base_width - groups.insert( - 0, - ParamGroupConfig( - pattern=spec.hidden_pattern, - optimizer_name=optimizer_name, - lr_mult=1.0 / m, - optimizer_kwargs=dict(optimizer_kwargs), - ), - ) - return groups diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 76012f0b2c..78a9f3a783 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -34,13 +34,6 @@ from torchtitan.components.checkpoint import CheckpointManager from torchtitan.components.lr_scheduler import LRSchedulersContainer from torchtitan.components.metrics import MetricsProcessor -from torchtitan.components.mup import ( - hidden_std, - MuPSpec, - output_mult, - param_groups, - residual_std, -) from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig from torchtitan.components.tokenizer import NoOpTokenizer from torchtitan.config import ( @@ -99,9 +92,6 @@ r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" ) -VIT_MUP = MuPSpec( - base_width=BASE_WIDTH, num_layers=NUM_LAYERS, hidden_pattern=MUP_PATTERN -) def model_registry(flavor: str) -> ModelSpec: @@ -478,6 +468,10 @@ def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Conf ) +def _hidden_std(fan_in: int, *, mup: bool) -> float: + return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 + + def _vit_attention( dim: int, *, n_head: int, mup: bool, qk_norm: bool = True ) -> PathSelfAttention.Config: @@ -490,8 +484,10 @@ def _vit_attention( k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) if qk_norm else None, - c_attn=_lin(dim, 3 * dim, std=hidden_std(dim, VIT_MUP, mup=mup)), - c_proj=_lin(dim, dim, std=residual_std(dim, VIT_MUP, mup=mup)), + c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), + c_proj=_lin( + dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * NUM_LAYERS) + ), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, head_dim=head_dim, @@ -504,8 +500,10 @@ def _vit_mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: hidden = _hidden_dim(dim, mult) return PathMLP.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - c_fc=_lin(dim, hidden, std=hidden_std(dim, VIT_MUP, mup=mup)), - c_proj=_lin(hidden, dim, std=residual_std(hidden, VIT_MUP, mup=mup)), + c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), + c_proj=_lin( + hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * NUM_LAYERS) + ), act="gelu_tanh", dropout=0.0, ) @@ -521,7 +519,7 @@ def _vit_model_config( t, h, w = INPUT_SIZE num_patches = (t // pt) * (h // ph) * (w // pw) return PlanViT.Config( - output_mult=output_mult(dim, VIT_MUP, mup=mup), + output_mult=(BASE_WIDTH / dim) if mup else 1.0, mean=255 / 2, std=255 / 4, patch_embed=PatchEmbed.Config( @@ -582,14 +580,26 @@ def _vit_optimizer_config( ) -> OptimizersContainer.Config: # base lr is carried on the container so --optimizer.lr can sweep every group # at once; muP scales the hidden matmuls down by lr_mult = 1/m (m = width ratio). + m = VIT_WIDTHS[flavor] / BASE_WIDTH common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} - groups = param_groups( - VIT_WIDTHS[flavor], - VIT_MUP, - mup=mup, - optimizer_name="AdamW", - optimizer_kwargs=common, - ) + groups = [ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={**common}, + ) + ] + if mup: + # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all + groups.insert( + 0, + ParamGroupConfig( + pattern=MUP_PATTERN, + optimizer_name="AdamW", + lr_mult=1.0 / m, + optimizer_kwargs={**common}, + ), + ) return OptimizersContainer.Config( implementation="fused_opt_states_bf16", lr=lr, param_groups=groups ) From 2c8fd746336eba286682d89a02a7b16552c3d60a Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Tue, 30 Jun 2026 11:21:13 -0700 Subject: [PATCH 23/28] path: vit muP proxy on 100k sample slice --- torchtitan/experiments/path/config_registry.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 78a9f3a783..29630b0aab 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -559,7 +559,7 @@ def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: def _vit_dataloader_config(*, split: str) -> PathDataLoader.Config: base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) return PathDataLoader.Config( - dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random10k_seed0.txt"), + dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random100k_seed0.txt"), split=split, shuffle_size=_si_int(base.shuffle_size), min_mixing=base.min_mixing, From 9cd5773a9fbd33e72272b08f582e1a69fbaeee02 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Tue, 30 Jun 2026 15:26:53 -0700 Subject: [PATCH 24/28] path: VIT_WD env override for weight-decay ablation --- torchtitan/experiments/path/config_registry.py | 1 + 1 file changed, 1 insertion(+) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 29630b0aab..310abb1c0e 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -608,6 +608,7 @@ def _vit_optimizer_config( def _vit( flavor: str, *, mup: bool, lr: float = 3e-4, wd: float = 3e-2 ) -> PathTrainer.Config: + wd = float(os.environ.get("VIT_WD", wd)) local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) num_nodes = int( From 8523e1aedd28ab3c20a9dc07908b9fc048a06e15 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Tue, 30 Jun 2026 23:45:26 -0700 Subject: [PATCH 25/28] prune_10m: width-scaled muP weight decay, drop VIT_WD --- torchtitan/experiments/path/config_registry.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 310abb1c0e..a5cb75bb39 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -597,7 +597,7 @@ def _vit_optimizer_config( pattern=MUP_PATTERN, optimizer_name="AdamW", lr_mult=1.0 / m, - optimizer_kwargs={**common}, + optimizer_kwargs={**common, "weight_decay": wd * m}, ), ) return OptimizersContainer.Config( @@ -606,9 +606,8 @@ def _vit_optimizer_config( def _vit( - flavor: str, *, mup: bool, lr: float = 3e-4, wd: float = 3e-2 + flavor: str, *, mup: bool, lr: float = 3e-4, wd: float = 0.0125 ) -> PathTrainer.Config: - wd = float(os.environ.get("VIT_WD", wd)) local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) num_nodes = int( From 23027927a1e11887e7d5e1e0da5641cdd73c7d51 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Wed, 1 Jul 2026 12:38:59 -0700 Subject: [PATCH 26/28] prune_10m: ufmt wrap dataset line --- torchtitan/experiments/path/config_registry.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index a5cb75bb39..5b7eb6c611 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -559,7 +559,9 @@ def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: def _vit_dataloader_config(*, split: str) -> PathDataLoader.Config: base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) return PathDataLoader.Config( - dataset=os.path.join(XX_BASEDIR, "datasets/lists/prune10m_random100k_seed0.txt"), + dataset=os.path.join( + XX_BASEDIR, "datasets/lists/prune10m_random100k_seed0.txt" + ), split=split, shuffle_size=_si_int(base.shuffle_size), min_mixing=base.min_mixing, From 7631b2b0cdf901bd6d16c65607b006d7e19cbe84 Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Wed, 1 Jul 2026 13:20:04 -0700 Subject: [PATCH 27/28] prune_10m: dedup dp degrees, ordered mup groups, real muP guard test --- .../unit_tests/test_optimizer_param_groups.py | 55 +++++++++++++++++++ .../experiments/path/config_registry.py | 52 ++++++++---------- 2 files changed, 79 insertions(+), 28 deletions(-) diff --git a/tests/unit_tests/test_optimizer_param_groups.py b/tests/unit_tests/test_optimizer_param_groups.py index a4ade9965d..6b2b4b0f10 100644 --- a/tests/unit_tests/test_optimizer_param_groups.py +++ b/tests/unit_tests/test_optimizer_param_groups.py @@ -762,5 +762,60 @@ def test_mixed_optimizer_same_schedule_different_base_lr(self): self.assertAlmostEqual(base_lr, 1e-3, places=6) +class TestViTMuPParamGroups(unittest.TestCase): + def test_mup_group_scales_lr_and_wd_across_widths(self): + """muP scales the hidden-matmul group by 1/m and its wd by m; catch-all unscaled.""" + from torchtitan.experiments.path.config_registry import ( + _vit_optimizer_config, + BASE_WIDTH, + MUP_PATTERN, + ) + + base_lr, base_wd = 3e-4, 0.0125 + for flavor, width in (("w256", 256), ("w512", 512)): + config = _vit_optimizer_config(flavor, mup=True, lr=base_lr, wd=base_wd) + mup_group, catch_all = config.param_groups[0], config.param_groups[-1] + + self.assertEqual(mup_group.pattern, MUP_PATTERN) + self.assertAlmostEqual(mup_group.lr_mult, BASE_WIDTH / width) + self.assertAlmostEqual( + mup_group.optimizer_kwargs["weight_decay"], base_wd * width / BASE_WIDTH + ) + + self.assertEqual(catch_all.pattern, r".*") + self.assertAlmostEqual(catch_all.lr_mult, 1.0) + self.assertAlmostEqual(catch_all.optimizer_kwargs["weight_decay"], base_wd) + + def test_mup_pattern_matches_hidden_matmuls_exactly(self): + """MUP_PATTERN selects exactly the per-block attention/mlp matmul weights.""" + import re + + from torchtitan.experiments.path.config_registry import ( + _vit_model_config, + MUP_PATTERN, + ) + + model = _vit_model_config("w256", mup=True).build() + param_names = {name for name, _ in model.named_parameters()} + expected = { + f"blocks.{i}.{submodule}.{leaf}.weight" + for i in range(len(model.blocks)) + for submodule, leaf in ( + ("attention", "c_attn"), + ("attention", "c_proj"), + ("mlp", "c_fc"), + ("mlp", "c_proj"), + ) + } + self.assertEqual(len(expected), 4 * len(model.blocks)) + self.assertTrue( + expected <= param_names, + f"hidden matmul weights missing from model: {expected - param_names}", + ) + + matched = {name for name in param_names if re.search(MUP_PATTERN, name)} + self.assertEqual(matched, expected) + + if __name__ == "__main__": unittest.main() diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 5b7eb6c611..d9a987b01c 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -122,6 +122,15 @@ def convnext_xxlarge() -> PathTrainer.Config: return _path("convnext_xxlarge") +def _dp_degrees() -> tuple[int, int]: + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) + world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) + num_nodes = int( + os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) + ) + return num_nodes, local_world_size + + def _path(flavor: str) -> PathTrainer.Config: steps = 1024 * 100 validation_freq = 1024 @@ -138,11 +147,7 @@ def _path(flavor: str) -> PathTrainer.Config: } reports["analyse_dataset"] = [validation_freq] mixed_precision_param = "bfloat16" - local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) - world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int( - os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) - ) + num_nodes, local_world_size = _dp_degrees() reporterv2_host = os.getenv("REPORTERV2_HOST") reporterv2_training_id = os.getenv("REPORTERV2_TRAINING_ID") checkpoint_base_folder = ( @@ -584,24 +589,19 @@ def _vit_optimizer_config( # at once; muP scales the hidden matmuls down by lr_mult = 1/m (m = width ratio). m = VIT_WIDTHS[flavor] / BASE_WIDTH common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} - groups = [ - ParamGroupConfig( - pattern=r".*", - optimizer_name="AdamW", - optimizer_kwargs={**common}, - ) - ] - if mup: - # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all - groups.insert( - 0, - ParamGroupConfig( - pattern=MUP_PATTERN, - optimizer_name="AdamW", - lr_mult=1.0 / m, - optimizer_kwargs={**common, "weight_decay": wd * m}, - ), - ) + catch_all = ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs=common, + ) + # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all + mup_group = ParamGroupConfig( + pattern=MUP_PATTERN, + optimizer_name="AdamW", + lr_mult=1.0 / m, + optimizer_kwargs={**common, "weight_decay": wd * m}, + ) + groups = [mup_group, catch_all] if mup else [catch_all] return OptimizersContainer.Config( implementation="fused_opt_states_bf16", lr=lr, param_groups=groups ) @@ -610,11 +610,7 @@ def _vit_optimizer_config( def _vit( flavor: str, *, mup: bool, lr: float = 3e-4, wd: float = 0.0125 ) -> PathTrainer.Config: - local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) - world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) - num_nodes = int( - os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size)) - ) + num_nodes, local_world_size = _dp_degrees() return PathTrainer.Config( loss=PlanViTLoss.Config(), model_spec=vit_model_registry(flavor, mup=mup), From ccb637660dcb63a7f2a218f044ac2836ce5dd17b Mon Sep 17 00:00:00 2001 From: Utkarsh Gill Date: Wed, 1 Jul 2026 16:59:51 -0700 Subject: [PATCH 28/28] prune_10m: move vit mup tests under experiments, raise on lr_mult without base lr, VIT_ prefix constants, dedup vit dataloader config --- .../unit_tests/test_optimizer_param_groups.py | 72 ++++------------ torchtitan/components/optimizer.py | 6 ++ .../experiments/path/config_registry.py | 86 +++++++------------ torchtitan/experiments/path/tests/__init__.py | 5 ++ .../experiments/path/tests/test_vit_mup.py | 62 +++++++++++++ torchtitan/experiments/path/vit.py | 6 -- 6 files changed, 123 insertions(+), 114 deletions(-) create mode 100644 torchtitan/experiments/path/tests/__init__.py create mode 100644 torchtitan/experiments/path/tests/test_vit_mup.py diff --git a/tests/unit_tests/test_optimizer_param_groups.py b/tests/unit_tests/test_optimizer_param_groups.py index 6b2b4b0f10..8d08895cf7 100644 --- a/tests/unit_tests/test_optimizer_param_groups.py +++ b/tests/unit_tests/test_optimizer_param_groups.py @@ -324,6 +324,23 @@ def test_base_lr_rejects_group_lr(self): with self.assertRaises(ValueError): config.build(model_parts=[model]) + def test_lr_mult_without_base_lr_raises(self): + """Setting lr_mult without a base lr raises ValueError.""" + model = SimpleModel() + config = OptimizersContainer.Config( + implementation="for-loop", + param_groups=[ + ParamGroupConfig( + pattern=r".*", + optimizer_name="AdamW", + optimizer_kwargs={"lr": 1e-3, "weight_decay": 0.1}, + lr_mult=0.5, + ), + ], + ) + with self.assertRaises(ValueError): + config.build(model_parts=[model]) + def test_first_match_wins(self): """When patterns overlap, the first match wins.""" model = SimpleModel() @@ -762,60 +779,5 @@ def test_mixed_optimizer_same_schedule_different_base_lr(self): self.assertAlmostEqual(base_lr, 1e-3, places=6) -class TestViTMuPParamGroups(unittest.TestCase): - def test_mup_group_scales_lr_and_wd_across_widths(self): - """muP scales the hidden-matmul group by 1/m and its wd by m; catch-all unscaled.""" - from torchtitan.experiments.path.config_registry import ( - _vit_optimizer_config, - BASE_WIDTH, - MUP_PATTERN, - ) - - base_lr, base_wd = 3e-4, 0.0125 - for flavor, width in (("w256", 256), ("w512", 512)): - config = _vit_optimizer_config(flavor, mup=True, lr=base_lr, wd=base_wd) - mup_group, catch_all = config.param_groups[0], config.param_groups[-1] - - self.assertEqual(mup_group.pattern, MUP_PATTERN) - self.assertAlmostEqual(mup_group.lr_mult, BASE_WIDTH / width) - self.assertAlmostEqual( - mup_group.optimizer_kwargs["weight_decay"], base_wd * width / BASE_WIDTH - ) - - self.assertEqual(catch_all.pattern, r".*") - self.assertAlmostEqual(catch_all.lr_mult, 1.0) - self.assertAlmostEqual(catch_all.optimizer_kwargs["weight_decay"], base_wd) - - def test_mup_pattern_matches_hidden_matmuls_exactly(self): - """MUP_PATTERN selects exactly the per-block attention/mlp matmul weights.""" - import re - - from torchtitan.experiments.path.config_registry import ( - _vit_model_config, - MUP_PATTERN, - ) - - model = _vit_model_config("w256", mup=True).build() - param_names = {name for name, _ in model.named_parameters()} - expected = { - f"blocks.{i}.{submodule}.{leaf}.weight" - for i in range(len(model.blocks)) - for submodule, leaf in ( - ("attention", "c_attn"), - ("attention", "c_proj"), - ("mlp", "c_fc"), - ("mlp", "c_proj"), - ) - } - self.assertEqual(len(expected), 4 * len(model.blocks)) - self.assertTrue( - expected <= param_names, - f"hidden matmul weights missing from model: {expected - param_names}", - ) - - matched = {name for name in param_names if re.search(MUP_PATTERN, name)} - self.assertEqual(matched, expected) - - if __name__ == "__main__": unittest.main() diff --git a/torchtitan/components/optimizer.py b/torchtitan/components/optimizer.py index f17869dc2f..0fe641ee95 100644 --- a/torchtitan/components/optimizer.py +++ b/torchtitan/components/optimizer.py @@ -219,6 +219,12 @@ def _build_param_groups( f"lr; use lr_mult to scale from the base lr instead" ) group_kwargs["lr"] = base_lr * pg.lr_mult + elif pg.lr_mult != 1.0: + raise ValueError( + f"Optimizer param_groups pattern '{pg.pattern}' sets lr_mult " + f"but the optimizer Config sets no base lr; lr_mult only " + f"applies when the Config sets a base lr" + ) groups[pg.optimizer_name].append( { diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index d9a987b01c..512da23362 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -6,6 +6,7 @@ from __future__ import annotations +import dataclasses import math import os from functools import partial @@ -25,6 +26,7 @@ from xx.training.path.hydra_configs import ( DRIVING_HEADS, META_HEADS, + PLAN_HEAD_SIZE, POSE_HEADS, TEMPORAL_META_HEADS, ) @@ -78,16 +80,14 @@ } _NORM_INIT = {"weight": nn.init.ones_, "bias": nn.init.zeros_} -# PlanViT (single-frame plan ViT) architecture and muP constants -HEAD_DIM = 64 -NUM_LAYERS = 8 -INPUT_SIZE = (1, 128, 256) -PATCH_SIZE = (1, 16, 8) -IN_CHANNELS = 24 -PLAN_SIZE = 15 * 33 * 2 -BASE_WIDTH = 256 +VIT_HEAD_DIM = 64 +VIT_NUM_LAYERS = 8 +VIT_INPUT_SIZE = (1, 128, 256) +VIT_PATCH_SIZE = (1, 16, 8) +VIT_IN_CHANNELS = 24 +VIT_BASE_WIDTH = 256 VIT_WIDTHS = {"w256": 256, "w512": 512, "w1024": 1024, "w2048": 2048} -STEPS = 512 +VIT_STEPS = 512 MUP_PATTERN = ( r"^(blocks\.\d+\.attention\.c_attn|blocks\.\d+\.attention\.c_proj" r"|blocks\.\d+\.mlp\.c_fc|blocks\.\d+\.mlp\.c_proj)\.weight$" @@ -474,24 +474,18 @@ def _lin(in_f: int, out_f: int, *, std: float, bias: bool = True) -> Linear.Conf def _hidden_std(fan_in: int, *, mup: bool) -> float: - return fan_in**-0.5 if mup else BASE_WIDTH**-0.5 + return fan_in**-0.5 if mup else VIT_BASE_WIDTH**-0.5 -def _vit_attention( - dim: int, *, n_head: int, mup: bool, qk_norm: bool = True -) -> PathSelfAttention.Config: +def _vit_attention(dim: int, *, n_head: int, mup: bool) -> PathSelfAttention.Config: head_dim = dim // n_head return PathSelfAttention.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) - if qk_norm - else None, - k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT) - if qk_norm - else None, + q_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), + k_norm=LayerNorm.Config(normalized_shape=head_dim, param_init=_NORM_INIT), c_attn=_lin(dim, 3 * dim, std=_hidden_std(dim, mup=mup)), c_proj=_lin( - dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * NUM_LAYERS) + dim, dim, std=_hidden_std(dim, mup=mup) / math.sqrt(2 * VIT_NUM_LAYERS) ), inner_attention=ScaledDotProductAttention.Config(), n_head=n_head, @@ -507,44 +501,44 @@ def _vit_mlp(dim: int, *, mup: bool, mult: float = 4.0) -> PathMLP.Config: norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), c_fc=_lin(dim, hidden, std=_hidden_std(dim, mup=mup)), c_proj=_lin( - hidden, dim, std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * NUM_LAYERS) + hidden, + dim, + std=_hidden_std(hidden, mup=mup) / math.sqrt(2 * VIT_NUM_LAYERS), ), act="gelu_tanh", dropout=0.0, ) -def _vit_model_config( - flavor: str, *, mup: bool, qk_norm: bool = True -) -> PlanViT.Config: +def _vit_model_config(flavor: str, *, mup: bool) -> PlanViT.Config: dim = VIT_WIDTHS[flavor] - n_head = dim // HEAD_DIM - pt, ph, pw = PATCH_SIZE - patch_dim = pt * IN_CHANNELS * ph * pw - t, h, w = INPUT_SIZE + n_head = dim // VIT_HEAD_DIM + pt, ph, pw = VIT_PATCH_SIZE + patch_dim = pt * VIT_IN_CHANNELS * ph * pw + t, h, w = VIT_INPUT_SIZE num_patches = (t // pt) * (h // ph) * (w // pw) return PlanViT.Config( - output_mult=(BASE_WIDTH / dim) if mup else 1.0, + output_mult=(VIT_BASE_WIDTH / dim) if mup else 1.0, mean=255 / 2, std=255 / 4, patch_embed=PatchEmbed.Config( proj=_lin(patch_dim, dim, std=patch_dim**-0.5), - patch_size=PATCH_SIZE, + patch_size=VIT_PATCH_SIZE, ), pos_embedding=Embedding.Config( num_embeddings=num_patches, embedding_dim=dim, param_init=_LINEAR_INIT ), blocks=[ PathTransformerBlock.Config( - attention=_vit_attention(dim, n_head=n_head, mup=mup, qk_norm=qk_norm), + attention=_vit_attention(dim, n_head=n_head, mup=mup), mlp=_vit_mlp(dim, mup=mup), ) - for _ in range(NUM_LAYERS) + for _ in range(VIT_NUM_LAYERS) ], norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), plan_head=PlanHead.Config( norm=LayerNorm.Config(normalized_shape=dim, param_init=_NORM_INIT), - head=_lin(dim, PLAN_SIZE, std=BASE_WIDTH**-0.5), + head=_lin(dim, PLAN_HEAD_SIZE, std=VIT_BASE_WIDTH**-0.5), ), ) @@ -562,39 +556,25 @@ def vit_model_registry(flavor: str, *, mup: bool) -> ModelSpec: def _vit_dataloader_config(*, split: str) -> PathDataLoader.Config: - base = XXPathDatasetConfig(fps=SUPERCOMBO_FPS, plan_only=True) - return PathDataLoader.Config( + return dataclasses.replace( + _dataloader_config(split=split, fps=SUPERCOMBO_FPS, plan_only=True), dataset=os.path.join( XX_BASEDIR, "datasets/lists/prune10m_random100k_seed0.txt" ), - split=split, - shuffle_size=_si_int(base.shuffle_size), - min_mixing=base.min_mixing, - num_writers=base.num_writers, - num_readers=base.num_readers, - fps=base.fps, pipeline_dir=BASE_DIR_GT_10M, - plan_only=base.plan_only, - limit=base.limit, - n_frames=base.n_frames, - rgb=base.rgb, - unvision=base.unvision, ) def _vit_optimizer_config( flavor: str, *, mup: bool, lr: float, wd: float ) -> OptimizersContainer.Config: - # base lr is carried on the container so --optimizer.lr can sweep every group - # at once; muP scales the hidden matmuls down by lr_mult = 1/m (m = width ratio). - m = VIT_WIDTHS[flavor] / BASE_WIDTH + m = VIT_WIDTHS[flavor] / VIT_BASE_WIDTH common = {"betas": (0.9, 0.95), "eps": 1e-8, "weight_decay": wd} catch_all = ParamGroupConfig( pattern=r".*", optimizer_name="AdamW", optimizer_kwargs=common, ) - # first-match-wins: scale hidden matmuls by lr_mult = 1/m before the catch-all mup_group = ParamGroupConfig( pattern=MUP_PATTERN, optimizer_name="AdamW", @@ -618,7 +598,7 @@ def _vit( dataloader=_vit_dataloader_config(split="train"), optimizer=_vit_optimizer_config(flavor, mup=mup, lr=lr, wd=wd), lr_scheduler=LRSchedulersContainer.Config( - warmup_steps=round(STEPS * 0.1), + warmup_steps=round(VIT_STEPS * 0.1), total_steps=None, decay_ratio=0.8, decay_type="cosine", @@ -628,7 +608,7 @@ def _vit( local_batch_size=16, global_batch_size=-1, seq_len=1, - steps=STEPS, + steps=VIT_STEPS, max_norm=1.0, dtype="float32", mixed_precision_param="bfloat16", @@ -640,7 +620,7 @@ def _vit( ), checkpoint=CheckpointManager.Config(enable=False), metrics=MetricsProcessor.Config( - log_freq=10, enable_reporterv2=True, save_freq=STEPS + log_freq=10, enable_reporterv2=True, save_freq=VIT_STEPS ), validator=PathValidator.Config( enable=False, diff --git a/torchtitan/experiments/path/tests/__init__.py b/torchtitan/experiments/path/tests/__init__.py new file mode 100644 index 0000000000..2e41cd717f --- /dev/null +++ b/torchtitan/experiments/path/tests/__init__.py @@ -0,0 +1,5 @@ +# 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. diff --git a/torchtitan/experiments/path/tests/test_vit_mup.py b/torchtitan/experiments/path/tests/test_vit_mup.py new file mode 100644 index 0000000000..80453df22a --- /dev/null +++ b/torchtitan/experiments/path/tests/test_vit_mup.py @@ -0,0 +1,62 @@ +# 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. + +import re +import unittest + +from torchtitan.experiments.path.config_registry import ( + _vit_model_config, + _vit_optimizer_config, + MUP_PATTERN, + VIT_BASE_WIDTH, +) + + +class TestViTMuPParamGroups(unittest.TestCase): + def test_mup_group_scales_lr_and_wd_across_widths(self): + """muP scales the hidden-matmul group by 1/m and its wd by m; catch-all unscaled.""" + base_lr, base_wd = 3e-4, 0.0125 + for flavor, width in (("w256", 256), ("w512", 512)): + config = _vit_optimizer_config(flavor, mup=True, lr=base_lr, wd=base_wd) + mup_group, catch_all = config.param_groups[0], config.param_groups[-1] + + self.assertEqual(mup_group.pattern, MUP_PATTERN) + self.assertAlmostEqual(mup_group.lr_mult, VIT_BASE_WIDTH / width) + self.assertAlmostEqual( + mup_group.optimizer_kwargs["weight_decay"], + base_wd * width / VIT_BASE_WIDTH, + ) + + self.assertEqual(catch_all.pattern, r".*") + self.assertAlmostEqual(catch_all.lr_mult, 1.0) + self.assertAlmostEqual(catch_all.optimizer_kwargs["weight_decay"], base_wd) + + def test_mup_pattern_matches_hidden_matmuls_exactly(self): + """MUP_PATTERN selects exactly the per-block attention/mlp matmul weights.""" + model = _vit_model_config("w256", mup=True).build() + param_names = {name for name, _ in model.named_parameters()} + expected = { + f"blocks.{i}.{submodule}.{leaf}.weight" + for i in range(len(model.blocks)) + for submodule, leaf in ( + ("attention", "c_attn"), + ("attention", "c_proj"), + ("mlp", "c_fc"), + ("mlp", "c_proj"), + ) + } + self.assertEqual(len(expected), 4 * len(model.blocks)) + self.assertTrue( + expected <= param_names, + f"hidden matmul weights missing from model: {expected - param_names}", + ) + + matched = {name for name in param_names if re.search(MUP_PATTERN, name)} + self.assertEqual(matched, expected) + + +if __name__ == "__main__": + unittest.main() diff --git a/torchtitan/experiments/path/vit.py b/torchtitan/experiments/path/vit.py index 62a5f1b3ad..bf17840561 100644 --- a/torchtitan/experiments/path/vit.py +++ b/torchtitan/experiments/path/vit.py @@ -112,12 +112,6 @@ def forward( class PlanViTLoss(PathLoss): - """PathLoss for the single-frame plan ViT. - - PlanViT predicts one frame, so the temporal plan target is reduced to its - last frame before scoring. Everything else matches PathLoss. - """ - @dataclass(kw_only=True, slots=True) class Config(PathLoss.Config): pass