Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
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
39 changes: 36 additions & 3 deletions torchtitan/components/optimizer/optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import logging
import re
from collections import defaultdict
from collections.abc import Callable, Iterator
from collections.abc import Callable, Iterator, Mapping
from dataclasses import dataclass, field
from typing import Any, cast, Generic, Literal, overload, Protocol, TypeVar

Expand All @@ -19,7 +19,11 @@
from torchtitan.components.checkpointer.utils import canonical_fqn
from torchtitan.config import Configurable
from torchtitan.distributed import ParallelDims
from torchtitan.distributed.flex_shard import build_dist_muon
from torchtitan.distributed.flex_shard import (
build_dist_muon,
validate_dist_muon_assignments,
)
from torchtitan.distributed.flex_shard.dist_muon import DistMuon

from .utils import (
get_flat_optim_state_dict,
Expand All @@ -38,6 +42,10 @@
"register_moe_quantile_balancing_hook",
]

_ASSIGNMENT_VALIDATORS: dict[
str, Callable[[Mapping[str, str | None], Mapping[str, Any]], None]
] = {DistMuon.__name__: validate_dist_muon_assignments}


@dataclass(kw_only=True, slots=True)
class ParamGroupConfig:
Expand Down Expand Up @@ -150,7 +158,7 @@ def _resolve_optimizer_factory(name: str) -> Callable[..., Optimizer]:
optimizer_factories: dict[str, Callable[..., Optimizer]] = {
"Adam": torch.optim.Adam,
"AdamW": torch.optim.AdamW,
"DistMuon": build_dist_muon,
DistMuon.__name__: build_dist_muon,
}
if name not in optimizer_factories:
raise NotImplementedError(f"Optimizer {name} not added.")
Expand Down Expand Up @@ -233,6 +241,9 @@ def __init__(self, config: Config, *, model_parts: list[nn.Module]) -> None:
groups_by_opt_name, patterns_by_opt_name = self._build_param_groups(
model, param_group_configs, impl_kwargs
)
self._validate_optimizer_assignments(
model, groups_by_opt_name, config.optimizer_factory_kwargs_by_name
)
for opt_name, opt_param_groups in groups_by_opt_name.items():
optimizer = self._resolve_optimizer_factory(opt_name)(
opt_param_groups,
Expand Down Expand Up @@ -272,6 +283,28 @@ def _log_optimizer(
f"{num_params} params [{pattern}] {kwargs}"
)

@staticmethod
def _validate_optimizer_assignments(
model: nn.Module,
groups_by_opt_name: dict[str, list[dict[str, Any]]],
factory_kwargs_by_name: dict[str, dict[str, Any]],
) -> None:
"""Run registered checks on all local trainable parameter assignments."""
optimizer_by_fqn: dict[str, str | None] = {
canonical_fqn(name): None
for name, param in model.named_parameters()
if param.requires_grad
}
for opt_name, groups in groups_by_opt_name.items():
for group in groups:
optimizer_by_fqn.update((fqn, opt_name) for fqn in group["param_names"])
# Include metadata-only entries so misassigned parameters cannot skip
# validation merely because their intended optimizer has no groups.
for opt_name in dict.fromkeys((*groups_by_opt_name, *factory_kwargs_by_name)):
validator = _ASSIGNMENT_VALIDATORS.get(opt_name)
if validator is not None:
validator(optimizer_by_fqn, factory_kwargs_by_name.get(opt_name, {}))

def _validate_params(self, all_params: list[nn.Parameter]) -> None:
"""Verify every trainable param is assigned to exactly one optimizer."""
expected = {
Expand Down
3 changes: 2 additions & 1 deletion torchtitan/distributed/flex_shard/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,12 @@

"""Flexible storage-to-compute redistribution APIs."""

from .dist_muon import build_dist_muon
from .dist_muon import build_dist_muon, validate_dist_muon_assignments
from .optimizer_reshard import BlockShard, BucketConfig, ComputeLayout, Owned

__all__ = [
"build_dist_muon",
"validate_dist_muon_assignments",
"BlockShard",
"BucketConfig",
"ComputeLayout",
Expand Down
24 changes: 24 additions & 0 deletions torchtitan/distributed/flex_shard/dist_muon.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@

__all__ = [
"build_dist_muon",
"validate_dist_muon_assignments",
]


Expand Down Expand Up @@ -110,6 +111,29 @@ def _normalize_param_groups(
return normalized_param_groups


def validate_dist_muon_assignments(
optimizer_by_fqn: Mapping[str, str | None],
factory_kwargs: Mapping[str, Any],
) -> None:
"""Require local trainable parameters with Muon layouts to use DistMuon.

The caller supplies canonical FQNs for local trainable parameters, including
unassigned parameters with a value of None. Frozen parameters and entries
for other pipeline stages are excluded from this assignment map.
"""
compute_sharding_fqns = factory_kwargs.get("compute_sharding_by_fqn", {})
for fqn, opt_name in optimizer_by_fqn.items():
if fqn not in compute_sharding_fqns:
continue
if opt_name != DistMuon.__name__:
assignment = (
f"assigned to {opt_name}"
if opt_name is not None
else "not assigned to an optimizer"
)
raise ValueError(f"{fqn} has a DistMuon compute layout but is {assignment}")


def _validate_compute_sharding_configuration(
compute_sharding_by_fqn: Mapping[str, ComputeLayout],
) -> None:
Expand Down
Loading