Skip to content
Open
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
151 changes: 134 additions & 17 deletions megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
import warnings
from collections import defaultdict
from dataclasses import astuple
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union

import torch
from torch.optim import SGD as CPUSGD
Expand Down Expand Up @@ -722,6 +722,136 @@ def check_config_overrides_consistency(
return True


def _tag_muon_split_qkv_parameters(
config: OptimizerConfig, model_chunks: list[MegatronModule]
) -> None:
"""Tag standard QKV and MLA up-projection weights with Muon split metadata.

Note that MLA down-projection weights aren't split with the current implementation.

The applied tags are:

- `is_qkv` (bool): whether the layer is a fused/packed collection of standard Multi-Head
Attention Q, K, and V projections. The layer contains a single weight, which includes all of
the Q, K, and V projections.
- `qkv_split_shapes` (tuple[int, ...] | None): per-group row counts for the projection slices
packed in this weight, if applicable. The gradients for the tagged parameter will be split
along the row dimension. The splits will be orthogonalized independently.

For example, without an Attention output gate, we have the following for the fused/packed
standard Multi-Head Attention QKV projection when not using tensor model parallelism:

```
# Remember that PyTorch uses transposed weights.
linear_qkv.weight.shape == (
(
num_attention_heads * kv_channels # <- Qs
+ kv_channels * num_query_groups # <- Ks
+ kv_channels * num_query_groups # <- Vs
),
hidden_size,
)

linear_qkv.qkv_split_shapes = (
num_attention_heads // num_query_groups * kv_channels, # <- Q
kv_channels, # <- K
kv_channels, # <- V
)
```

In the optimizer step, we transform the gradient like the following:
```
# The initial shape is like `linear_qkv.weight.shape` before.
linear_qkv.grad.shape == linear_qkv.weight.shape

qs_grad, ks_grad, vs_grad = muon_split_grads(linear_qkv.grad, linear_qkv.qkv_split_shapes)

qs_grad.shape == (
num_query_groups,
num_attention_heads // num_query_groups * kv_channels,
hidden_size,
)
ks_grad.shape == vs_grad.shape == (
num_query_groups,
kv_channels,
hidden_size,
)

# Contract first two dimensions.
qs_grad = qs_grad.reshape(num_attention_heads * kv_channels, hidden_size)
qs_grad = orthogonalize(qs_grad)

ks_grad = ks_grad.reshape(num_query_groups * kv_channels, hidden_size)
ks_grad = orthogonalize(ks_grad)

vs_grad = vs_grad.reshape(num_query_groups * kv_channels, hidden_size)
vs_grad = orthogonalize(vs_grad)

result_grad = muon_concat_grads(qs_grad, ks_grad, vs_grad)
result_grad.shape == linear_qkv.grad.shape
```
"""
split_qkv = getattr(config, 'muon_split_qkv', True)

def _is_muon_split_metadata_managed_param(
name: str, attn_variant: Literal["mha", "mla"]
) -> bool:
"""Return whether the parameter with the given name should be managed by Muon split tags.

E.g., standard QKV or MLA projection weights.
"""
if 'linear_qkv.weight' in name:
return True
if attn_variant != "mla":
return False
return any(
f'{projection}.weight' in name
for projection in (
'linear_q_proj',
'linear_q_up_proj',
'linear_kv_up_proj',
'linear_q_down_proj',
'linear_kv_down_proj',
'linear_qkv_down_proj',
)
)

for model_chunk in model_chunks:
model_cfg = get_model_config(model_chunk)
attn_variant = "mla" if getattr(model_cfg, 'multi_latent_attention', False) else "mha"
standard_qkv_split_shapes = None
mla_q_split_shapes = None
mla_kv_split_shapes = None
if attn_variant == "mla":
mla_q_split_shapes = (model_cfg.qk_head_dim, model_cfg.qk_pos_emb_head_dim)
mla_kv_split_shapes = (model_cfg.qk_head_dim, model_cfg.v_head_dim)

for name, param in model_chunk.named_parameters():
if not param.requires_grad:
continue

is_managed_param = _is_muon_split_metadata_managed_param(name, attn_variant)
if is_managed_param:
# Remove existing Muon split tags from a parameter.
for attr in ("is_qkv", "qkv_split_shapes"):
if hasattr(param, attr):
delattr(param, attr)
if not is_managed_param or not split_qkv or len(param.shape) != 2:
continue

if 'linear_qkv.weight' in name:
if standard_qkv_split_shapes is None:
standard_qkv_split_shapes = tuple(_get_qkv_split_shapes(model_cfg))
param.qkv_split_shapes = standard_qkv_split_shapes
param.is_qkv = True
elif mla_q_split_shapes is not None and (
'linear_q_up_proj.weight' in name or 'linear_q_proj.weight' in name
):
param.qkv_split_shapes = mla_q_split_shapes
elif mla_kv_split_shapes is not None and 'linear_kv_up_proj.weight' in name:
param.qkv_split_shapes = mla_kv_split_shapes


def _get_megatron_emerging_optimizer(
config: OptimizerConfig,
model_chunks: List[MegatronModule],
Expand Down Expand Up @@ -774,28 +904,15 @@ def _get_megatron_emerging_optimizer(

log_single_rank(logger, logging.INFO, f'Setting up emerging optimizer with config {config}')

# Tag parameters with optimizer-specific attributes (expert_tp, is_qkv).
# Tag parameters with optimizer-specific attributes.
for model_chunk in model_chunks:
qkv_split_shapes = None
for name, param in model_chunk.named_parameters():
if not param.requires_grad:
continue
if 'experts' in name and 'shared' not in name:
param.expert_tp = True
# TODO(deyuf): support MLA
if 'linear_qkv.weight' in name and len(param.shape) == 2:
if qkv_split_shapes is None:
qkv_split_shapes = _get_qkv_split_shapes(model_chunk.config)
if param.shape[0] % sum(qkv_split_shapes) == 0:
param.is_qkv = True
param.qkv_split_shapes = qkv_split_shapes
else:
log_single_rank(
logger,
logging.DEBUG,
f"Emerging optimizer QKV split skipped for {name}: "
f"shape={tuple(param.shape)}, split_shapes={qkv_split_shapes}",
)
if eopt_name in ('muon', 'adaptive_muon'):
_tag_muon_split_qkv_parameters(config, model_chunks)

# Apply optimizer-specific default param overrides (e.g. muon: non-linear -> adam).
# For Muon-family optimizers, the scalar optimizer that handles non-linear/embedding
Expand Down
91 changes: 60 additions & 31 deletions megatron/core/optimizer/emerging_optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -266,6 +266,58 @@ def scaled_orthogonalize_fn_with_gtp_remat(self, p, grad, tp_group, partition_di
shard_size = gathered_grad.shape[0] // gtp_remat_size
return gathered_grad[gtp_rank * shard_size : (gtp_rank + 1) * shard_size].contiguous()

def _get_muon_split_shapes(self, p: torch.Tensor) -> tuple[int, ...] | None:
"""Return per-parameter Muon split shapes, if this parameter should be split."""
split_shapes: tuple[int, ...] | None = getattr(p, "qkv_split_shapes", None)
if split_shapes is None:
if self.is_qkv_fn is None or not self.is_qkv_fn(p):
return None
split_shapes = self.qkv_split_shapes
if split_shapes is None:
raise ValueError(
"Muon QKV split was requested for a parameter, but `qkv_split_shapes` is not "
"set."
)

if not split_shapes or any(shape <= 0 for shape in split_shapes):
raise ValueError(f"Muon split shapes must be positive integers, got {split_shapes}.")
return split_shapes

def _orthogonalize_split_grad(
self,
param: torch.Tensor,
grad: torch.Tensor,
split_shapes: tuple[int, ...],
tp_group: torch.distributed.ProcessGroup | None,
partition_dim: int | None,
) -> torch.Tensor:
"""Orthogonalize a fused projection gradient by splitting its row layout first."""
grad_shape = grad.shape
split_size = sum(split_shapes)
if grad_shape[0] % split_size != 0:
raise ValueError(
f"Muon split parameter has incompatible grad shape `{tuple(grad_shape)}` "
f"for split shapes `{split_shapes}`: `grad.shape[0]` must be divisible by "
f"`sum(split_shapes)={split_size}`."
)

log_single_rank(
logger,
logging.DEBUG,
f'muon split grad shape `{grad_shape}`, split shapes `{split_shapes}`',
)
num_groups = grad_shape[0] // split_size
split_grads = torch.split(grad.view(num_groups, split_size, -1), split_shapes, dim=1)
split_grads = [g.reshape(-1, grad_shape[-1]) for g in split_grads]

split_grads = [
self.scaled_orthogonalize_fn_with_gtp_remat(param, g, tp_group, partition_dim).view(
num_groups, -1, grad_shape[-1]
)
for g in split_grads
]
return torch.cat(split_grads, dim=1).view(grad_shape)

def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> torch.Tensor:
"""Orthogonalize the momentum.

Expand All @@ -291,38 +343,15 @@ def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> t
if partition_dim == -1:
partition_dim = None

if self.split_qkv and self.is_qkv_fn(p): # type: ignore[misc]
grad_shape = grad.shape
qkv_split_shapes = getattr(p, "qkv_split_shapes", None)
if qkv_split_shapes is None:
qkv_split_shapes = self.qkv_split_shapes
if qkv_split_shapes is None:
raise RuntimeError("Muon QKV split requested but qkv_split_shapes is not set")
qkv_split_dim = sum(qkv_split_shapes)
if grad_shape[0] % qkv_split_dim != 0:
raise RuntimeError(
f"Muon QKV split shape mismatch: grad_shape={tuple(grad_shape)}, "
f"split_shapes={qkv_split_shapes}"
grad_was_split = False
if self.split_qkv:
split_shapes = self._get_muon_split_shapes(p)
if split_shapes is not None:
grad = self._orthogonalize_split_grad(
p, grad, split_shapes, tp_group, partition_dim
)
log_single_rank(
logger,
logging.DEBUG,
f'qkv split grad shape {grad_shape}, split shapes {qkv_split_shapes}',
)
num_query_groups = grad_shape[0] // qkv_split_dim
qkv_grads = torch.split(
grad.view(num_query_groups, qkv_split_dim, -1), qkv_split_shapes, dim=1
)
qkv_grads = [g.reshape(-1, grad_shape[-1]) for g in qkv_grads]

qkv_grads = [
self.scaled_orthogonalize_fn_with_gtp_remat(p, g, tp_group, partition_dim).view(
num_query_groups, -1, grad_shape[-1]
)
for g in qkv_grads
]
grad = torch.cat(qkv_grads, dim=1).view(grad_shape)
else:
grad_was_split = True
if not grad_was_split:
grad = self.scaled_orthogonalize_fn_with_gtp_remat(p, grad, tp_group, partition_dim)
return grad

Expand Down
95 changes: 95 additions & 0 deletions tests/unit_tests/test_emerging_optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -720,6 +720,101 @@ def test_muon_optimizer_qkv_split():
), "Weights should be different between split_qkv=True and split_qkv=False"


def test_muon_optimizer_qkv_split_uses_legacy_is_qkv_metadata():
"""Test that legacy is_qkv + qkv_split_shapes still uses the split path."""
split_shapes = (3, 2, 1)
num_groups = 4
hidden_size = 5
param = torch.nn.Parameter(
torch.zeros(num_groups * sum(split_shapes), hidden_size, device='cuda')
)
param.is_qkv = True
optimizer = TensorParallelMuon(
params=[param],
lr=0.01,
split_qkv=True,
is_qkv_fn=lambda p: getattr(p, 'is_qkv', False),
qkv_split_shapes=split_shapes,
num_ns_steps=1,
pg_collection=None,
tp_mode="duplicated",
)

calls = []

def fake_orthogonalize(grad, tp_group, partition_dim):
calls.append(tuple(grad.shape))
return torch.full_like(grad, len(calls))

optimizer.scaled_orthogonalize_fn = fake_orthogonalize
result = optimizer.orthogonalize(param, torch.zeros_like(param))

assert calls == [
(num_groups * split_shapes[0], hidden_size),
(num_groups * split_shapes[1], hidden_size),
(num_groups * split_shapes[2], hidden_size),
]
result = result.view(num_groups, sum(split_shapes), hidden_size)
assert torch.all(result[:, : split_shapes[0]] == 1).item()
assert torch.all(result[:, split_shapes[0] : split_shapes[0] + split_shapes[1]] == 2).item()
assert torch.all(result[:, -split_shapes[2] :] == 3).item()


def test_muon_optimizer_mla_kv_split_uses_parameter_metadata():
"""Test that MLA-style 2-way split metadata uses the split path."""
split_shapes = (3, 5)
num_heads = 4
hidden_size = 7
param = torch.nn.Parameter(
torch.zeros(num_heads * sum(split_shapes), hidden_size, device='cuda')
)
param.qkv_split_shapes = split_shapes
optimizer = TensorParallelMuon(
params=[param],
lr=0.01,
split_qkv=True,
num_ns_steps=1,
pg_collection=None,
tp_mode="duplicated",
)

calls = []

def fake_orthogonalize(grad, tp_group, partition_dim):
calls.append(tuple(grad.shape))
return torch.full_like(grad, len(calls))

optimizer.scaled_orthogonalize_fn = fake_orthogonalize
result = optimizer.orthogonalize(param, torch.zeros_like(param))

assert calls == [
(num_heads * split_shapes[0], hidden_size),
(num_heads * split_shapes[1], hidden_size),
]
result = result.view(num_heads, sum(split_shapes), hidden_size)
assert torch.all(result[:, : split_shapes[0]] == 1).item()
assert torch.all(result[:, split_shapes[0] :] == 2).item()


def test_muon_optimizer_mla_split_metadata_validates_grad_shape():
"""Test that invalid MLA split metadata reports an incompatible layout."""
split_shapes = (3, 5)
hidden_size = 7
param = torch.nn.Parameter(torch.zeros(sum(split_shapes) + 1, hidden_size, device='cuda'))
param.qkv_split_shapes = split_shapes
optimizer = TensorParallelMuon(
params=[param],
lr=0.01,
split_qkv=True,
num_ns_steps=1,
pg_collection=None,
tp_mode="duplicated",
)

with pytest.raises(ValueError, match=r"grad\.shape\[0\].*sum\(split_shapes\)=8"):
optimizer.orthogonalize(param, torch.zeros_like(param))


def test_muon_optimizer_extra_scale_factor():
"""Test TensorParallelMuon optimizer with different extra_scale_factor values."""
model = torch.nn.Linear(80, 40, bias=False, dtype=torch.float32, device='cuda')
Expand Down
Loading
Loading