From 5b5b32adbbc69606b1073816f24c5fa2d0f59908 Mon Sep 17 00:00:00 2001 From: lizhihan1 Date: Tue, 12 May 2026 00:12:16 +0800 Subject: [PATCH 1/2] Fix Muon QKV split for gated attention --- megatron/core/optimizer/__init__.py | 16 +++++++- .../core/optimizer/emerging_optimizers.py | 38 +++++++++++++------ megatron/core/tensor_parallel/layers.py | 2 + .../test_tp_attrs_without_init.py | 15 ++++++++ tests/unit_tests/test_emerging_optimizers.py | 17 +++++++++ 5 files changed, 75 insertions(+), 13 deletions(-) diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index c6d3e41aed5..3a47d8ef713 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -65,6 +65,7 @@ _EMERGING_OPTIMIZERS, HAVE_EMERGING_OPTIMIZERS, _create_emerging_optimizer, + _get_qkv_split_shapes, ) from .grad_scaler import ConstantGradScaler, DynamicGradScaler from .layer_wise_optimizer import LayerWiseDistributedOptimizer @@ -769,6 +770,7 @@ def _get_megatron_emerging_optimizer( # Tag parameters with optimizer-specific attributes (expert_tp, is_qkv). for model_chunk in model_chunks: + qkv_split_shapes = None for name, param in model_chunk.named_parameters(): if not param.requires_grad: continue @@ -776,7 +778,19 @@ def _get_megatron_emerging_optimizer( param.expert_tp = True # TODO(deyuf): support MLA if 'linear_qkv.weight' in name and len(param.shape) == 2: - param.is_qkv = True + 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 + param.muon_param_name = name + 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}", + ) # Apply optimizer-specific default param overrides (e.g. muon: non-linear -> adam). config_overrides.update(_EMERGING_OPTIMIZERS[eopt_name].default_param_overrides) diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py index cc218d6ba40..878cadf74ec 100644 --- a/megatron/core/optimizer/emerging_optimizers.py +++ b/megatron/core/optimizer/emerging_optimizers.py @@ -132,11 +132,17 @@ def _is_nonlinear_or_embedding(param): def _get_qkv_split_shapes(model_cfg) -> List[int]: """Compute QKV split shapes from model config.""" - return [ - model_cfg.num_attention_heads // model_cfg.num_query_groups * model_cfg.kv_channels, - model_cfg.kv_channels, - model_cfg.kv_channels, - ] + query_projection_size = ( + model_cfg.num_attention_heads // model_cfg.num_query_groups * model_cfg.kv_channels + ) + if getattr(model_cfg, 'attention_output_gate', False): + return [ + query_projection_size, + query_projection_size, + model_cfg.kv_channels, + model_cfg.kv_channels, + ] + return [query_projection_size, model_cfg.kv_channels, model_cfg.kv_channels] # =========================================================================== @@ -164,7 +170,7 @@ def __init__( use_decoupled_weight_decay: bool = True, split_qkv: bool = False, is_qkv_fn: Callable[[torch.Tensor], bool] | None = None, - qkv_split_shapes: tuple[int, int, int] | None = None, + qkv_split_shapes: List[int] | None = None, fp32_matmul_prec: str = "medium", coefficient_type: str = "quintic", num_ns_steps: int = 5, @@ -251,16 +257,24 @@ def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> t 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", 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: + param_name = getattr(p, "muon_param_name", "") + raise RuntimeError( + f"Muon QKV split shape mismatch for {param_name}: " + f"grad_shape={tuple(grad_shape)}, split_shapes={qkv_split_shapes}" + ) log_single_rank( logger, logging.DEBUG, - f'qkv split grad shape {grad_shape}, ' f'split shapes {self.qkv_split_shapes}', + f'qkv split grad shape {grad_shape}, split shapes {qkv_split_shapes}', ) - num_query_groups = grad_shape[0] // sum(self.qkv_split_shapes) + num_query_groups = grad_shape[0] // qkv_split_dim qkv_grads = torch.split( - grad.view(num_query_groups, sum(self.qkv_split_shapes), -1), - self.qkv_split_shapes, - dim=1, + 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] @@ -317,7 +331,7 @@ def __init__( use_decoupled_weight_decay: bool = True, split_qkv: bool = False, is_qkv_fn: Callable[[torch.Tensor], bool] | None = None, - qkv_split_shapes: tuple[int, int, int] | None = None, + qkv_split_shapes: List[int] | None = None, fp32_matmul_prec: str = "medium", coefficient_type: str = "quintic", num_ns_steps: int = 5, diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 662373064cf..75960df5d89 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -60,6 +60,8 @@ _MODEL_PARALLEL_ATTRIBUTE_DEFAULTS = { "expert_tp": False, "is_qkv": False, + "qkv_split_shapes": None, + "muon_param_name": None, "tensor_model_parallel": False, "partition_dim": -1, "partition_stride": 1, diff --git a/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py b/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py index f8c345ed78b..64968cb4da7 100644 --- a/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py +++ b/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py @@ -7,6 +7,7 @@ ColumnParallelLinear, RowParallelLinear, VocabParallelEmbedding, + copy_tensor_model_parallel_attributes, ) from megatron.core.transformer.transformer_config import TransformerConfig from tests.unit_tests.test_utilities import Utils @@ -87,3 +88,17 @@ def test_row_parallel_linear_tp_attrs_no_init(self, use_cpu_init): assert hasattr(w, "tensor_model_parallel") and w.tensor_model_parallel is True assert hasattr(w, "partition_dim") and w.partition_dim == 1 assert hasattr(w, "partition_stride") and w.partition_stride == 1 + + +def test_copy_tensor_model_parallel_attributes_preserves_muon_metadata(): + source = torch.empty(4, 4) + destination = torch.empty_like(source) + source.is_qkv = True + source.qkv_split_shapes = [256, 64, 64] + source.muon_param_name = "decoder.layers.0.self_attention.linear_qkv.weight" + + copy_tensor_model_parallel_attributes(destination, source) + + assert destination.is_qkv is True + assert destination.qkv_split_shapes == source.qkv_split_shapes + assert destination.muon_param_name == source.muon_param_name diff --git a/tests/unit_tests/test_emerging_optimizers.py b/tests/unit_tests/test_emerging_optimizers.py index cbd2bed5ec6..df6180d7e34 100644 --- a/tests/unit_tests/test_emerging_optimizers.py +++ b/tests/unit_tests/test_emerging_optimizers.py @@ -15,6 +15,7 @@ HAVE_EMERGING_OPTIMIZERS, TensorParallelAdaptiveMuon, TensorParallelMuon, + _get_qkv_split_shapes, get_supported_coefficient_types, validate_coefficient_type, ) @@ -65,6 +66,22 @@ def forward(self, x): # =========================================================================== +def test_muon_qkv_split_shapes(): + config = TransformerConfig( + num_layers=1, hidden_size=1024, num_attention_heads=16, num_query_groups=8 + ) + gated_config = TransformerConfig( + num_layers=1, + hidden_size=1024, + num_attention_heads=16, + num_query_groups=8, + attention_output_gate=True, + ) + + assert _get_qkv_split_shapes(config) == [128, 64, 64] + assert _get_qkv_split_shapes(gated_config) == [128, 128, 64, 64] + + def test_muon_optimizer_smoke(): """Smoke test for TensorParallelMuon optimizer.""" # Create a simple linear model for testing From 1d8ed4f7fadd3baf5cd5808096ee5f455e151e9d Mon Sep 17 00:00:00 2001 From: Moozy Date: Thu, 28 May 2026 23:31:39 +0800 Subject: [PATCH 2/2] Update QKV split shape lookup and remove unnecessary param metadata --- megatron/core/optimizer/__init__.py | 1 - megatron/core/optimizer/emerging_optimizers.py | 17 +++++++++-------- megatron/core/tensor_parallel/layers.py | 1 - .../test_tp_attrs_without_init.py | 4 +--- 4 files changed, 10 insertions(+), 13 deletions(-) diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 3a47d8ef713..ea8fe4a1e55 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -783,7 +783,6 @@ def _get_megatron_emerging_optimizer( if param.shape[0] % sum(qkv_split_shapes) == 0: param.is_qkv = True param.qkv_split_shapes = qkv_split_shapes - param.muon_param_name = name else: log_single_rank( logger, diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py index 878cadf74ec..4dfed6199b3 100644 --- a/megatron/core/optimizer/emerging_optimizers.py +++ b/megatron/core/optimizer/emerging_optimizers.py @@ -11,7 +11,7 @@ import inspect import logging from dataclasses import dataclass, field -from typing import Any, Callable, Dict, List, Literal, Optional, get_args +from typing import Any, Callable, Dict, Literal, Optional, get_args import torch from torch.optim.optimizer import ParamsT @@ -130,7 +130,7 @@ def _is_nonlinear_or_embedding(param): return getattr(param, 'is_embedding_or_output_parameter', False) or len(param.shape) != 2 -def _get_qkv_split_shapes(model_cfg) -> List[int]: +def _get_qkv_split_shapes(model_cfg) -> list[int]: """Compute QKV split shapes from model config.""" query_projection_size = ( model_cfg.num_attention_heads // model_cfg.num_query_groups * model_cfg.kv_channels @@ -170,7 +170,7 @@ def __init__( use_decoupled_weight_decay: bool = True, split_qkv: bool = False, is_qkv_fn: Callable[[torch.Tensor], bool] | None = None, - qkv_split_shapes: List[int] | None = None, + qkv_split_shapes: list[int] | None = None, fp32_matmul_prec: str = "medium", coefficient_type: str = "quintic", num_ns_steps: int = 5, @@ -257,15 +257,16 @@ def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> t 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", self.qkv_split_shapes) + 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: - param_name = getattr(p, "muon_param_name", "") raise RuntimeError( - f"Muon QKV split shape mismatch for {param_name}: " - f"grad_shape={tuple(grad_shape)}, split_shapes={qkv_split_shapes}" + f"Muon QKV split shape mismatch: grad_shape={tuple(grad_shape)}, " + f"split_shapes={qkv_split_shapes}" ) log_single_rank( logger, @@ -331,7 +332,7 @@ def __init__( use_decoupled_weight_decay: bool = True, split_qkv: bool = False, is_qkv_fn: Callable[[torch.Tensor], bool] | None = None, - qkv_split_shapes: List[int] | None = None, + qkv_split_shapes: list[int] | None = None, fp32_matmul_prec: str = "medium", coefficient_type: str = "quintic", num_ns_steps: int = 5, diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 75960df5d89..52484afff8e 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -61,7 +61,6 @@ "expert_tp": False, "is_qkv": False, "qkv_split_shapes": None, - "muon_param_name": None, "tensor_model_parallel": False, "partition_dim": -1, "partition_stride": 1, diff --git a/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py b/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py index 64968cb4da7..44d6fa21178 100644 --- a/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py +++ b/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py @@ -90,15 +90,13 @@ def test_row_parallel_linear_tp_attrs_no_init(self, use_cpu_init): assert hasattr(w, "partition_stride") and w.partition_stride == 1 -def test_copy_tensor_model_parallel_attributes_preserves_muon_metadata(): +def test_copy_tensor_model_parallel_attributes_preserves_qkv_split_shapes(): source = torch.empty(4, 4) destination = torch.empty_like(source) source.is_qkv = True source.qkv_split_shapes = [256, 64, 64] - source.muon_param_name = "decoder.layers.0.self_attention.linear_qkv.weight" copy_tensor_model_parallel_attributes(destination, source) assert destination.is_qkv is True assert destination.qkv_split_shapes == source.qkv_split_shapes - assert destination.muon_param_name == source.muon_param_name