Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
29 commits
Select commit Hold shift + click to select a range
f2919c0
plan_vit: add the muP / scaling-study ViT as a torchtitan experiment
utkarshgill Jun 23, 2026
2e5d27e
plan_vit: add license headers and ufmt formatting
utkarshgill Jun 23, 2026
8fd8c52
plan_vit: drop the 1/m readout multiplier, output now width-stable
utkarshgill Jun 24, 2026
b446573
plan_vit: derive data parallelism from the launch so N>1 nodes validate
utkarshgill Jun 24, 2026
e160f7e
plan_vit: unset lr_scheduler total_steps so it tracks training.steps
utkarshgill Jun 24, 2026
a430cef
plan_vit: restore canonical muP readout (1/m mult, base-width init, b…
utkarshgill Jun 25, 2026
c6dbad2
path: muP plan_vit as path/vit.py on PathTrainer (validator/onnx off)
utkarshgill Jun 25, 2026
9a23545
path: remove standalone plan_vit experiment (moved into path/vit.py)
utkarshgill Jun 25, 2026
9a4aa78
Merge remote-tracking branch 'origin/main' into plan-vit-experiment
utkarshgill Jun 26, 2026
1baf6c6
path: trim the vit experiment to a minimal surface
utkarshgill Jun 26, 2026
754be01
path: drive the muP base lr via --mup_base_lr so run.sh's native swee…
utkarshgill Jun 26, 2026
16a3e4d
cleanup
utkarshgill Jun 26, 2026
8e0b8c8
path: build the vit from path's transformer blocks
utkarshgill Jun 26, 2026
3aad5e8
path: fold the vit configs into config_registry
utkarshgill Jun 26, 2026
77355af
path: inline the vit configs
utkarshgill Jun 26, 2026
3cfbb59
path: follow house patterns; muP lr via --optimizer.lr
utkarshgill Jun 26, 2026
a83821b
path: trim vit muP config surface
utkarshgill Jun 26, 2026
63809b3
path: license headers and ufmt formatting
utkarshgill Jun 27, 2026
08450a5
optimizer: test base lr with per-group lr_mult
utkarshgill Jun 27, 2026
6fe28a9
path: extract muP into a reusable module
utkarshgill Jun 27, 2026
4a3b5a6
mup: move to components as a reusable technique
utkarshgill Jun 27, 2026
a174e6f
mup: move parametrization to experiments/mup, trim comments
utkarshgill Jun 28, 2026
c81fe06
path: re-inline vit muP, drop the muP module
utkarshgill Jun 29, 2026
2c8fd74
path: vit muP proxy on 100k sample slice
utkarshgill Jun 30, 2026
9cd5773
path: VIT_WD env override for weight-decay ablation
utkarshgill Jun 30, 2026
8523e1a
prune_10m: width-scaled muP weight decay, drop VIT_WD
utkarshgill Jul 1, 2026
2302792
prune_10m: ufmt wrap dataset line
utkarshgill Jul 1, 2026
7631b2b
prune_10m: dedup dp degrees, ordered mup groups, real muP guard test
utkarshgill Jul 1, 2026
ccb6376
prune_10m: move vit mup tests under experiments, raise on lr_mult wit…
utkarshgill Jul 1, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 78 additions & 0 deletions tests/unit_tests/test_optimizer_param_groups.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
41 changes: 37 additions & 4 deletions torchtitan/components/optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down
Loading
Loading