diff --git a/tests/unit_tests/test_optimizer_param_groups.py b/tests/unit_tests/test_optimizer_param_groups.py index 58d3e7d0df..8d08895cf7 100644 --- a/tests/unit_tests/test_optimizer_param_groups.py +++ b/tests/unit_tests/test_optimizer_param_groups.py @@ -263,6 +263,84 @@ 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_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() diff --git a/torchtitan/components/optimizer.py b/torchtitan/components/optimizer.py index bd399fe50b..0fe641ee95 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,27 @@ 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 + 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( { "params": params, "param_names": param_names, - **impl_kwargs, - **pg.optimizer_kwargs, + **group_kwargs, } ) patterns[pg.optimizer_name].append(pg.pattern) @@ -213,7 +246,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 bff5fa9786..512da23362 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -1,11 +1,39 @@ +# 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 dataclasses 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, + PLAN_HEAD_SIZE, + POSE_HEADS, + TEMPORAL_META_HEADS, +) 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 @@ -20,23 +48,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,17 +67,32 @@ 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 -_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_} +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} +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$" +) + def model_registry(flavor: str) -> ModelSpec: return ModelSpec( @@ -89,11 +122,20 @@ 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 + 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", @@ -105,12 +147,12 @@ 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 = 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 +195,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 +215,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 +250,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 +266,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 +303,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 +328,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 +339,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 +348,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 +376,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 +390,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 +403,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 +421,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 +434,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 +453,213 @@ 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 + }, ) + + +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 VIT_BASE_WIDTH**-0.5 + + +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), + 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 * VIT_NUM_LAYERS) + ), + 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(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 * VIT_NUM_LAYERS), + ), + act="gelu_tanh", + dropout=0.0, + ) + + +def _vit_model_config(flavor: str, *, mup: bool) -> PlanViT.Config: + dim = VIT_WIDTHS[flavor] + 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=(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=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), + mlp=_vit_mlp(dim, mup=mup), + ) + 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_HEAD_SIZE, std=VIT_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: + 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" + ), + pipeline_dir=BASE_DIR_GT_10M, + ) + + +def _vit_optimizer_config( + flavor: str, *, mup: bool, lr: float, wd: float +) -> OptimizersContainer.Config: + 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, + ) + 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 + ) + + +def _vit( + flavor: str, *, mup: bool, lr: float = 3e-4, wd: float = 0.0125 +) -> PathTrainer.Config: + num_nodes, local_world_size = _dp_degrees() + return PathTrainer.Config( + loss=PlanViTLoss.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(VIT_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=VIT_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=VIT_STEPS + ), + validator=PathValidator.Config( + enable=False, + steps=-1, + dataloader=_vit_dataloader_config(split="val"), + mixed_precision_param="bfloat16", + ), + fps=SUPERCOMBO_FPS, + 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) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index c3b02b1c38..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) @@ -90,14 +101,20 @@ 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() + 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() @@ -108,7 +125,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))) @@ -175,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) @@ -216,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] @@ -235,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]) @@ -276,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, @@ -322,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: @@ -364,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: @@ -476,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) @@ -498,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/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 new file mode 100644 index 0000000000..bf17840561 --- /dev/null +++ b/torchtitan/experiments/path/vit.py @@ -0,0 +1,160 @@ +# 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 +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.protocols.model import BaseModel +from torchtitan.protocols.module import Module, ModuleList +from torchtitan.tools.logging import logger + +from .loss import PathLoss +from .model import PathTransformerBlock + + +class PatchEmbed(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + proj: Linear.Config + patch_size: tuple[int, int, int] + + 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: + 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)) + + +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): + 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 + + 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 + + def forward( + self, inputs: dict[str, torch.Tensor] | torch.Tensor + ) -> dict[str, torch.Tensor]: + 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: + x = block(x) + x = self.norm(x) + return {"plan": self.plan_head(x.mean(dim=1)) * self.config.output_mult} + + +class PlanViTLoss(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, + *, + 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("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( + 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 PlanViT") + return model