From f585f0339600e941981d7510441b5a23e351db3d Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Wed, 22 Jul 2026 20:06:04 +0800 Subject: [PATCH 01/27] core: add DSA accuracy compatibility option --- megatron/core/transformer/transformer_config.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index bbcf413baee..a6beb137947 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -318,6 +318,11 @@ class TransformerConfig(ModelParallelConfig): ``none`` disables fused DSA kernels. Explicit ``tilelang`` or ``cudnn`` enables only that backend. Unsupported DSA layouts continue to use the PyTorch fallback.""" + dsa_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--dsa-accuracy-compatible"]}} + ) + """Use the full-score DSA fallback with explicit softmax backward for alignment.""" + dsa_indexer_rope_interleaved: bool = False """Whether DSA indexer RoPE should use MLA-style interleaving.""" From 6fa5e02ac4cf1c774399812eb44cb5571dc34235 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Wed, 22 Jul 2026 20:06:33 +0800 Subject: [PATCH 02/27] core: add accuracy-compatible DSA fallback --- .../experimental_attention_variant/dsa.py | 60 +++++++++++++++++-- .../test_attention_variant_dsa.py | 55 +++++++++++++++++ 2 files changed, 110 insertions(+), 5 deletions(-) diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index dde238635c2..c7f3a212558 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -64,6 +64,7 @@ def _unfused_absorbed_dsa_fn( varlen_starts: Optional[torch.Tensor] = None, varlen_ends: Optional[torch.Tensor] = None, key_positions: Optional[torch.Tensor] = None, + accuracy_compatible: bool = False, ) -> torch.Tensor: """Unfused absorbed-MLA attention: output stays [sq, b, np, v_channels].""" sq, b, np, hn = query.size() @@ -99,10 +100,15 @@ def _unfused_absorbed_dsa_fn( ) attention_scores = attention_scores + index_mask.unsqueeze(1) - valid_index_mask = torch.isfinite(index_mask) - attention_scores = dsa_masking.masked_softmax( - attention_scores.float(), valid_index_mask.unsqueeze(1).expand(b, np, sq, skv), dim=-1 - ) + valid_index_mask = torch.isfinite(index_mask).unsqueeze(1).expand(b, np, sq, skv) + if accuracy_compatible: + attention_scores = _AccuracyCompatibleSoftmax.apply( + attention_scores.float(), valid_index_mask + ) + else: + attention_scores = dsa_masking.masked_softmax( + attention_scores.float(), valid_index_mask, dim=-1 + ) # Latent value is the first v_channels slice of absorbed key cache. value = key[..., :v_channels].permute(1, 2, 0, 3) # [b,1,skv,v] @@ -110,6 +116,25 @@ def _unfused_absorbed_dsa_fn( return output.permute(2, 0, 1, 3).contiguous() +class _AccuracyCompatibleSoftmax(torch.autograd.Function): + """Masked softmax with an explicit backward formula for DSA alignment.""" + + @staticmethod + def forward(ctx, logits: torch.Tensor, valid_mask: torch.Tensor) -> torch.Tensor: + probabilities = torch.softmax(logits.masked_fill(~valid_mask, float("-inf")), dim=-1) + probabilities = probabilities.masked_fill(~valid_mask, 0.0) + ctx.save_for_backward(probabilities, valid_mask) + return probabilities + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + probabilities, valid_mask = ctx.saved_tensors + grad_logits = probabilities * ( + grad_output - (grad_output * probabilities).sum(dim=-1, keepdim=True) + ) + return grad_logits.masked_fill(~valid_mask, 0.0), None + + def _run_sparse_attention( *, absorbed_mla: bool, @@ -127,6 +152,7 @@ def _run_sparse_attention( topk_length: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Run sparse attention for absorbed and non-absorbed MLA paths.""" + accuracy_compatible = bool(getattr(config, "dsa_accuracy_compatible", False)) if absorbed_mla: latent_v_channels = int(getattr(config, "kv_lora_rank", 0) or 0) if latent_v_channels <= 0: @@ -143,7 +169,7 @@ def _run_sparse_attention( "Received absorbed layout with explicit value tensor." ) output = None - if dsa_kernels.use_fused_dsa_kernels(config): + if not accuracy_compatible and dsa_kernels.use_fused_dsa_kernels(config): output = dsa_kernels.run_fused_absorbed_sparse_attention( config, query, @@ -166,6 +192,7 @@ def _run_sparse_attention( varlen_starts=varlen_starts, varlen_ends=varlen_ends, key_positions=key_positions, + accuracy_compatible=accuracy_compatible, ) assert output is not None output = torch.einsum("sbhc,hdc->sbhd", output, up_v_weight).contiguous() @@ -182,6 +209,7 @@ def _run_sparse_attention( varlen_starts=varlen_starts, varlen_ends=varlen_ends, key_positions=key_positions, + accuracy_compatible=accuracy_compatible, ) @@ -1411,6 +1439,7 @@ def unfused_dsa_fn( varlen_starts: Optional[torch.Tensor] = None, varlen_ends: Optional[torch.Tensor] = None, key_positions: Optional[torch.Tensor] = None, + accuracy_compatible: bool = False, ): """ Unfused sparse attention implementation. @@ -1457,6 +1486,27 @@ def unfused_dsa_fn( device=query.device, ) + if accuracy_compatible: + index_mask = torch.full((b, sq, skv), float("-inf"), device=query.device) + dsa_masking.scatter_topk_into_index_mask(index_mask, topk_indices) + index_mask = dsa_masking.apply_sparse_validity_to_index_mask( + index_mask, + row_mask=row_mask, + varlen_starts=varlen_starts, + varlen_ends=varlen_ends, + key_positions=key_positions, + ) + valid_index_mask = torch.isfinite(index_mask).unsqueeze(1).expand(b, np, sq, skv) + attention_scores = ( + torch.matmul(query_b.float(), key_b.float().transpose(-1, -2)) * softmax_scale + ) + attention_probs = _AccuracyCompatibleSoftmax.apply( + attention_scores + index_mask.unsqueeze(1), valid_index_mask + ) + output = torch.matmul(attention_probs.to(value_b.dtype), value_b) + output = output.permute(2, 0, 1, 3).contiguous().view(sq, b, np * hnv) + return output.squeeze(1) if query_was_thd else output + seq_chunk_size = 512 head_chunk_size = 16 topk_chunk_size = 1024 diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py index 642aeeb126f..67ae7530216 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_attention_variant_dsa.py @@ -27,6 +27,7 @@ DSAttention, DSAttentionSubmodules, FusedDSAIndexerLoss, + _AccuracyCompatibleSoftmax, _run_sparse_attention, _validate_nonpacked_cp_uniform_length, compute_dsa_indexer_loss, @@ -68,6 +69,60 @@ def mock_hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor return x * scale +class TestAccuracyCompatibleDSA: + """Test the opt-in full-score DSA alignment path.""" + + def test_explicit_softmax_backward_matches_formula(self): + logits = torch.randn(2, 3, 5, device="cuda", requires_grad=True) + valid_mask = torch.ones_like(logits, dtype=torch.bool) + valid_mask[..., -1] = False + grad_output = torch.randn_like(logits) + + probabilities = _AccuracyCompatibleSoftmax.apply(logits, valid_mask) + probabilities.backward(grad_output) + expected = probabilities.detach() * ( + grad_output - (grad_output * probabilities.detach()).sum(dim=-1, keepdim=True) + ) + expected = expected.masked_fill(~valid_mask, 0.0) + + assert torch.equal(logits.grad, expected) + assert torch.equal(probabilities[..., -1], torch.zeros_like(probabilities[..., -1])) + + def test_accuracy_compatible_switch_defaults_off(self, monkeypatch): + query = torch.randn(8, 1, 2, 8, device="cuda", dtype=torch.bfloat16) + key = torch.randn_like(query) + value = torch.randn(8, 1, 2, 4, device="cuda", dtype=torch.bfloat16) + indices = torch.arange(8, device="cuda").view(1, 8, 1) + original = unfused_dsa_fn + calls = [] + + def capture(*args, **kwargs): + calls.append(kwargs.get("accuracy_compatible")) + return original(*args, **kwargs) + + monkeypatch.setattr( + "megatron.core.transformer.experimental_attention_variant.dsa.unfused_dsa_fn", + capture, + ) + common = dict( + absorbed_mla=False, + query=query, + key=key, + value=value, + up_v_weight=None, + topk_indices=indices, + softmax_scale=query.size(-1) ** -0.5, + mask=None, + varlen_starts=None, + varlen_ends=None, + key_positions=None, + ) + _run_sparse_attention(config=SimpleNamespace(), **common) + _run_sparse_attention(config=SimpleNamespace(dsa_accuracy_compatible=True), **common) + + assert calls == [False, True] + + class TestDSAIndexShareHelpers: """Test cross-layer top-k sharing helpers.""" From 21a9d382f4d8e6e8c66425876cebf4cada964744 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Thu, 23 Jul 2026 06:30:05 +0800 Subject: [PATCH 03/27] core:add-RMSNorm-compatibility-option --- megatron/core/transformer/transformer_config.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index a6beb137947..777a10ecab3 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -186,6 +186,11 @@ class TransformerConfig(ModelParallelConfig): ) """Epsilon value for any LayerNorm/RMSNorm operations.""" + norm_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--norm-accuracy-compatible"]}} + ) + """Use explicit fp32 normalization formulas instead of native norm kernels for alignment.""" + layernorm_zero_centered_gamma: bool = field( default=False, metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}} ) From cee7604604268be11fe60eacf52684e55915556a Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Thu, 23 Jul 2026 06:30:05 +0800 Subject: [PATCH 04/27] core:add-accuracy-compatible-RMSNorm --- ...rimental_attention_variant_module_specs.py | 9 +++- megatron/core/transformer/torch_norm.py | 30 +++++++++++++ ...rimental_attention_variant_module_specs.py | 18 ++++++++ .../unit_tests/transformer/test_torch_norm.py | 44 +++++++++++++++++++ 4 files changed, 100 insertions(+), 1 deletion(-) create mode 100644 tests/unit_tests/transformer/test_torch_norm.py diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index a76fe6e3a23..a4848f6cb2d 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -20,6 +20,7 @@ ) from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.spec_utils import ModuleSpec +from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm from megatron.core.transformer.transformer_block import ( TransformerBlockSubmodules, get_num_layers_to_build, @@ -107,7 +108,13 @@ def get_dsa_module_spec_for_backend( # DSA indexer requires normalized q as input, so here we cannot fuse qk layernorm # with linear projection and have to use unfused qk layernorm. qk_norm = ( - backend.layer_norm(rms_norm=rms_norm, for_qk=True) if config.qk_layernorm else IdentityOp + ( + AccuracyCompatibleRMSNorm + if config.norm_accuracy_compatible + else backend.layer_norm(rms_norm=rms_norm, for_qk=True) + ) + if config.qk_layernorm + else IdentityOp ) attention = ModuleSpec( diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py index 5948ae600f9..3c735565d2e 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -24,6 +24,34 @@ def __call__( ) -> LayerNormInterface: ... +class AccuracyCompatibleRMSNorm(torch.nn.Module, LayerNormInterface): + """RMSNorm with explicit fp32 reduction and one output cast.""" + + def __init__( + self, + normalized_shape: int | None = None, + eps: float = 1e-5, + *, + hidden_size: int | None = None, + config: TransformerConfig | None = None, + **kwargs, + ): + super().__init__() + normalized_shape = hidden_size if normalized_shape is None else normalized_shape + if normalized_shape is None: + raise ValueError("normalized_shape or hidden_size is required") + self.normalized_shape = (normalized_shape,) + self.eps = eps + dtype = config.params_dtype if config is not None else None + self.weight = torch.nn.Parameter(torch.ones(normalized_shape, dtype=dtype)) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x_float = x.float() + variance = x_float.pow(2).mean(dim=-1, keepdim=True) + output = x_float * torch.rsqrt(variance + self.eps) + return (output * self.weight.float()).to(x.dtype) + + class WrappedTorchNorm: """ A conditional wrapper to initialize an instance of PyTorch's @@ -56,6 +84,8 @@ def __new__( if config.normalization == "LayerNorm": norm_cls = torch.nn.LayerNorm elif config.normalization == "RMSNorm": + if config.norm_accuracy_compatible: + return AccuracyCompatibleRMSNorm(normalized_shape=hidden_size, eps=eps) assert is_torch_min_version( "2.4.0a0" ), 'Torch RMSNorm requires PyTorch version >= 2.4.0' diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py index 0a454b5d7ff..5fca6010a30 100644 --- a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -65,6 +65,7 @@ def _make_config(**overrides): defaults = dict( num_layers=4, normalization="RMSNorm", + norm_accuracy_compatible=False, qk_layernorm=False, multi_latent_attention=False, qk_l2_norm=False, @@ -369,6 +370,23 @@ def test_qk_layernorm_enabled(self, normalization): assert spec.submodules.q_layernorm is spec.submodules.kv_layernorm backend.layer_norm.assert_any_call(rms_norm=expected_rms, for_qk=True) + def test_accuracy_compatible_qk_rmsnorm(self): + """Verify DSA q/kv norms can use the explicit fp32 RMSNorm path.""" + from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm + + backend = _make_backend() + cfg = _make_config( + multi_latent_attention=True, + qk_l2_norm=False, + qk_layernorm=True, + normalization="RMSNorm", + norm_accuracy_compatible=True, + ) + spec = self._call(cfg=cfg, backend=backend) + + assert spec.submodules.q_layernorm is AccuracyCompatibleRMSNorm + assert spec.submodules.kv_layernorm is AccuracyCompatibleRMSNorm + def test_qk_layernorm_disabled(self): """Verify q/kv layernorm becomes IdentityOp, skipping backend.layer_norm for qk.""" backend = _make_backend() diff --git a/tests/unit_tests/transformer/test_torch_norm.py b/tests/unit_tests/transformer/test_torch_norm.py new file mode 100644 index 00000000000..14b077abe98 --- /dev/null +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -0,0 +1,44 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +import torch + +from megatron.core.transformer.torch_norm import ( + AccuracyCompatibleRMSNorm, + WrappedTorchNorm, +) +from megatron.core.transformer.transformer_config import TransformerConfig + + +def _config(**overrides): + values = { + "num_layers": 1, + "hidden_size": 64, + "num_attention_heads": 4, + "normalization": "RMSNorm", + } + values.update(overrides) + return TransformerConfig(**values) + + +def test_accuracy_compatible_rmsnorm_matches_explicit_formula(): + config = _config(norm_accuracy_compatible=True) + norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5).cuda().bfloat16() + x = torch.randn(2, 3, 64, device="cuda", dtype=torch.bfloat16) + + output = norm(x) + x_float = x.float() + expected = ( + x_float + * torch.rsqrt(x_float.pow(2).mean(dim=-1, keepdim=True) + 1e-5) + * norm.weight.float() + ).to(torch.bfloat16) + + assert isinstance(norm, AccuracyCompatibleRMSNorm) + assert torch.equal(output, expected) + + +def test_default_rmsnorm_stays_native(): + config = _config(norm_accuracy_compatible=False) + norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) + + assert isinstance(norm, torch.nn.RMSNorm) From 4ec435a92cf85e0773da7d4391243d71cc4aced6 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Thu, 23 Jul 2026 06:47:30 +0800 Subject: [PATCH 05/27] core:add-router-compatibility-option --- megatron/core/transformer/transformer_config.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 777a10ecab3..b2fb6da4b05 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -191,6 +191,11 @@ class TransformerConfig(ModelParallelConfig): ) """Use explicit fp32 normalization formulas instead of native norm kernels for alignment.""" + router_accuracy_compatible: bool = field( + default=False, metadata={"argparse_meta": {"arg_names": ["--router-accuracy-compatible"]}} + ) + """Use an explicit fp32 router GEMM instead of the fused Transformer Engine path.""" + layernorm_zero_centered_gamma: bool = field( default=False, metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}} ) From 2f213f97f7fead3b58b85277d0599a6127f5b861 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Thu, 23 Jul 2026 06:47:30 +0800 Subject: [PATCH 06/27] core:add-accuracy-compatible-router-gating --- megatron/core/transformer/moe/router.py | 8 +++++ .../transformer/moe/test_routers.py | 36 +++++++++++++++++++ 2 files changed, 44 insertions(+) diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 03317b65f1c..e25339b5ff1 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -104,6 +104,14 @@ def gating(self, input: torch.Tensor): router_dtype = torch.float32 elif self.config.moe_router_dtype == 'fp64': router_dtype = torch.float64 + if self.config.router_accuracy_compatible: + inp_shape = input.shape + logits = torch.mm( + input.reshape(-1, inp_shape[-1]).float(), self.weight.float().t() + ) + if self.bias is not None: + logits = logits + self.bias.float() + return logits.view(*inp_shape[:-1], -1) logits = router_gating_linear(input, self.weight, self.bias, router_dtype) return logits diff --git a/tests/unit_tests/transformer/moe/test_routers.py b/tests/unit_tests/transformer/moe/test_routers.py index 9f33dd01920..da05d34b937 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -62,6 +62,42 @@ def test_constructor(self): num_weights = sum([p.numel() for p in self.router.parameters()]) assert num_weights == 12 * 4, num_weights + @pytest.mark.internal + def test_router_accuracy_compatible_gating(self): + hidden_states = torch.randn( + (3, 1, self.router.config.hidden_size), device="cuda", dtype=torch.bfloat16 + ) + self.router.config.router_accuracy_compatible = True + + logits = self.router.gating(hidden_states) + expected = torch.mm( + hidden_states.reshape(-1, hidden_states.shape[-1]).float(), + self.router.weight.float().t(), + ).view(3, 1, -1) + + assert logits.dtype == torch.float32 + assert torch.equal(logits, expected) + + @pytest.mark.internal + def test_default_router_gating_stays_native(self, monkeypatch): + expected = torch.randn((3, 1, self.router.config.num_moe_experts)) + called = False + + def fake_router_gating_linear(inp, weight, bias, router_dtype): + nonlocal called + called = True + return expected + + monkeypatch.setattr( + "megatron.core.transformer.moe.router.router_gating_linear", + fake_router_gating_linear, + ) + hidden_states = torch.randn((3, 1, self.router.config.hidden_size), dtype=torch.bfloat16) + + assert self.router.config.router_accuracy_compatible is False + assert self.router.gating(hidden_states) is expected + assert called + @pytest.mark.internal @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.parametrize("moe_router_pre_softmax", [(True), (False)]) From 09d0c169c341d4d7280f7ab889a0e4b6beaf87f8 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Thu, 23 Jul 2026 15:55:39 +0800 Subject: [PATCH 07/27] core-align-MTP-RMSNorm-backward --- megatron/core/models/gpt/gpt_layer_specs.py | 2 +- .../transformer/multi_token_prediction.py | 17 ++++++++--- megatron/core/transformer/torch_norm.py | 24 +++++++++++++-- .../test_multi_token_prediction.py | 29 +++++++++++++++++++ .../unit_tests/transformer/test_torch_norm.py | 15 ++++++++++ 5 files changed, 80 insertions(+), 7 deletions(-) diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py index 984840b3a87..63a1aa51b02 100755 --- a/megatron/core/models/gpt/gpt_layer_specs.py +++ b/megatron/core/models/gpt/gpt_layer_specs.py @@ -773,7 +773,7 @@ def get_gpt_mtp_block_spec_for_backend( raise ValueError(f"Invalid spec: {spec}") mtp_layer_spec = get_mtp_layer_spec_for_backend( - mtp_model_layer_spec=transformer_layer_spec, backend=backend + mtp_model_layer_spec=transformer_layer_spec, backend=backend, config=config ) mtp_num_layers = config.mtp_num_layers if config.mtp_num_layers else 0 if config.mtp_use_repeated_layer: diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index b20514ce6a4..536c3dfd15e 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -30,7 +30,7 @@ from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module -from megatron.core.transformer.torch_norm import LayerNormBuilder +from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm, LayerNormBuilder from megatron.core.transformer.transformer_block import TransformerBlockSubmodules from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.typed_torch import apply_module @@ -577,7 +577,9 @@ class MultiTokenPredictionLayerSubmodules: def get_mtp_layer_spec( - mtp_model_layer_spec: ModuleSpec, use_transformer_engine: bool + mtp_model_layer_spec: ModuleSpec, + use_transformer_engine: bool, + config: Optional[TransformerConfig] = None, ) -> ModuleSpec: """Get the MTP layer spec. @@ -587,11 +589,14 @@ def get_mtp_layer_spec( return get_mtp_layer_spec_for_backend( mtp_model_layer_spec, backend=TESpecProvider() if use_transformer_engine else LocalSpecProvider(), + config=config, ) def get_mtp_layer_spec_for_backend( - mtp_model_layer_spec: ModuleSpec, backend: BackendSpecProvider + mtp_model_layer_spec: ModuleSpec, + backend: BackendSpecProvider, + config: Optional[TransformerConfig] = None, ) -> ModuleSpec: """Get the MTP layer spec. @@ -599,7 +604,11 @@ def get_mtp_layer_spec_for_backend( ModuleSpec: Module specification with modules from the backend. """ column_parallel_linear_impl: type = backend.column_parallel_linear() - layer_norm_impl = backend.layer_norm() + layer_norm_impl = ( + AccuracyCompatibleRMSNorm + if config is not None and config.norm_accuracy_compatible + else backend.layer_norm() + ) mtp_layer_spec = ModuleSpec( module=MultiTokenPredictionLayer, submodules=MultiTokenPredictionLayerSubmodules( diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py index 3c735565d2e..c0ddf0763fc 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -24,6 +24,27 @@ def __call__( ) -> LayerNormInterface: ... +class _AccuracyCompatibleRMSNormFunction(torch.autograd.Function): + """RMSNorm core with a stable fp32 backward and canonical zero gradients.""" + + @staticmethod + def forward(ctx, x: torch.Tensor, eps: float) -> torch.Tensor: + variance = x.pow(2).mean(dim=-1, keepdim=True) + inv_rms = torch.rsqrt(variance + eps) + ctx.save_for_backward(x, inv_rms) + return x * inv_rms + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + x, inv_rms = ctx.saved_tensors + dot = (grad_output * x).sum(dim=-1, keepdim=True) + correction_scale = dot * (-0.5) * inv_rms.pow(3) / x.shape[-1] + correction = (correction_scale * x) * 2.0 + grad_input = grad_output * inv_rms + correction + grad_input = torch.where(grad_input == 0, torch.zeros_like(grad_input), grad_input) + return grad_input, None + + class AccuracyCompatibleRMSNorm(torch.nn.Module, LayerNormInterface): """RMSNorm with explicit fp32 reduction and one output cast.""" @@ -47,8 +68,7 @@ def __init__( def forward(self, x: torch.Tensor) -> torch.Tensor: x_float = x.float() - variance = x_float.pow(2).mean(dim=-1, keepdim=True) - output = x_float * torch.rsqrt(variance + self.eps) + output = _AccuracyCompatibleRMSNormFunction.apply(x_float, self.eps) return (output * self.weight.float()).to(x.dtype) diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index c3c3944e007..9014d0896fb 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -29,6 +29,7 @@ process_mtp_loss, roll_tensor, ) +from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import get_batch_on_this_cp_rank, is_te_min_version, unwrap_model from megatron.training.argument_utils import gpt_config_from_args, hybrid_config_from_args @@ -82,6 +83,34 @@ def _create_config_and_mtp_block_spec(self, tp, cp, use_te=False): ) return config, mtp_block_spec + def test_accuracy_compatible_norms_override_te_mtp_norms(self): + """Accuracy mode must route all MTP-owned norms through the explicit RMSNorm.""" + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + config = TransformerConfig( + mtp_num_layers=1, + num_layers=1, + hidden_size=64, + num_attention_heads=8, + normalization="RMSNorm", + norm_accuracy_compatible=True, + use_cpu_initialization=True, + ) + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec() + mtp_block_spec = get_gpt_mtp_block_spec( + config=config, + spec=transformer_layer_spec, + use_transformer_engine=True, + ) + mtp_layer_spec = mtp_block_spec.layer_specs[0] + + assert mtp_layer_spec.submodules.enorm is AccuracyCompatibleRMSNorm + assert mtp_layer_spec.submodules.hnorm is AccuracyCompatibleRMSNorm + assert mtp_layer_spec.submodules.layer_norm is AccuracyCompatibleRMSNorm + final_norm = mtp_layer_spec.submodules.layer_norm( + config=config, hidden_size=config.hidden_size, eps=config.layernorm_epsilon + ) + assert isinstance(final_norm, AccuracyCompatibleRMSNorm) + def test_mtp_detach_heads_config(self): """Test that mtp_detach_heads config defaults to False.""" config = TransformerConfig( diff --git a/tests/unit_tests/transformer/test_torch_norm.py b/tests/unit_tests/transformer/test_torch_norm.py index 14b077abe98..d3e69a59c6e 100644 --- a/tests/unit_tests/transformer/test_torch_norm.py +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -37,6 +37,21 @@ def test_accuracy_compatible_rmsnorm_matches_explicit_formula(): assert torch.equal(output, expected) +def test_accuracy_compatible_rmsnorm_canonicalizes_zero_input_gradients(): + config = _config(norm_accuracy_compatible=True) + norm = WrappedTorchNorm(config=config, hidden_size=4, eps=1e-5).cuda().bfloat16() + with torch.no_grad(): + norm.weight.copy_(torch.tensor([-1.0, 1.0, -2.0, 2.0], device="cuda")) + x = torch.tensor( + [[[1.0, -1.0, 2.0, -2.0]]], device="cuda", dtype=torch.bfloat16, requires_grad=True + ) + + norm(x).backward(torch.zeros_like(x)) + + assert torch.equal(x.grad, torch.zeros_like(x.grad)) + assert torch.equal(x.grad.view(torch.uint16), torch.zeros_like(x.grad.view(torch.uint16))) + + def test_default_rmsnorm_stays_native(): config = _config(norm_accuracy_compatible=False) norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) From 1a266f0b7a37ea55182e394836ca5e4a3ea9b310 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Tue, 11 Aug 2026 17:11:58 +0800 Subject: [PATCH 08/27] feat(glm52): align accuracy-compatible training paths --- ...rimental_attention_variant_module_specs.py | 79 ++- megatron/core/optimizer/__init__.py | 3 + megatron/core/optimizer/optimizer_config.py | 3 + .../transformer/multi_token_prediction.py | 347 +++++++--- megatron/core/transformer/torch_norm.py | 91 +-- .../core/transformer/transformer_config.py | 611 +++++++++++------- ...rimental_attention_variant_module_specs.py | 115 +++- .../test_multi_token_prediction.py | 467 ++++++++----- .../unit_tests/transformer/test_torch_norm.py | 42 +- 9 files changed, 1114 insertions(+), 644 deletions(-) diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index a4848f6cb2d..a8cc4886717 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -20,7 +20,7 @@ ) from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.spec_utils import ModuleSpec -from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm +from megatron.core.transformer.torch_norm import WrappedTorchNorm from megatron.core.transformer.transformer_block import ( TransformerBlockSubmodules, get_num_layers_to_build, @@ -58,6 +58,15 @@ ########## +def _get_standalone_norm( + config: TransformerConfig, backend: BackendSpecProvider, *, for_qk=False +): + rms_norm = config.normalization == "RMSNorm" + if rms_norm and config.norm_accuracy_compatible: + return WrappedTorchNorm + return backend.layer_norm(rms_norm=rms_norm, for_qk=for_qk) + + def get_gated_delta_net_module_spec( config: TransformerConfig, backend: BackendSpecProvider = None ) -> ModuleSpec: @@ -66,12 +75,11 @@ def get_gated_delta_net_module_spec( if backend is None: backend = _get_backend_spec_provider(config=config) - rms_norm = config.normalization == "RMSNorm" attention = ModuleSpec( module=GatedDeltaNet, submodules=GatedDeltaNetSubmodules( in_proj=backend.column_parallel_layer_norm_linear(), - out_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False), + out_norm=_get_standalone_norm(config, backend), out_proj=backend.row_parallel_linear(), ), metainfo={"fuse_input_layernorm": True}, @@ -83,7 +91,9 @@ def get_dsa_module_spec_for_backend( config: TransformerConfig, backend: BackendSpecProvider = None ) -> ModuleSpec: """Helper function to get module spec for Sparse Attention.""" - assert config.multi_latent_attention, "Currently only MLA supports sparse attention." + assert config.multi_latent_attention, ( + "Currently only MLA supports sparse attention." + ) assert config.qk_l2_norm is False, "qk_l2_norm is not supported with MLA." # Because TransformerEngine does not support sparse attention yet, we use local @@ -103,16 +113,10 @@ def get_dsa_module_spec_for_backend( ), ) - # Adjust for RMS norm. - rms_norm = config.normalization == "RMSNorm" # DSA indexer requires normalized q as input, so here we cannot fuse qk layernorm # with linear projection and have to use unfused qk layernorm. qk_norm = ( - ( - AccuracyCompatibleRMSNorm - if config.norm_accuracy_compatible - else backend.layer_norm(rms_norm=rms_norm, for_qk=True) - ) + _get_standalone_norm(config, backend, for_qk=True) if config.qk_layernorm else IdentityOp ) @@ -210,7 +214,9 @@ def get_transformer_layer_with_experimental_attention_variant_spec( experimental_attention_spec = None if 0 in experimental_attention_pattern: - standard_attention_spec = _get_self_attention_module_spec(config=config, backend=backend) + standard_attention_spec = _get_self_attention_module_spec( + config=config, backend=backend + ) else: standard_attention_spec = None @@ -235,7 +241,6 @@ def get_transformer_layer_with_experimental_attention_variant_spec( dense_mlp_layer_spec, fuse_layernorm_pre_dense = None, False # Get GPT decoder block layer specs - rms_norm = config.normalization == "RMSNorm" layer_specs = [] for layer_number in range(config.num_layers): attention = ( @@ -243,7 +248,11 @@ def get_transformer_layer_with_experimental_attention_variant_spec( if experimental_attention_pattern[layer_number] == 1 else standard_attention_spec ) - mlp = moe_layer_spec if moe_layer_pattern[layer_number] == 1 else dense_mlp_layer_spec + mlp = ( + moe_layer_spec + if moe_layer_pattern[layer_number] == 1 + else dense_mlp_layer_spec + ) fuse_pre_mlp_layernorm = ( fuse_layernorm_pre_moe if moe_layer_pattern[layer_number] == 1 @@ -252,12 +261,12 @@ def get_transformer_layer_with_experimental_attention_variant_spec( input_layernorm = ( IdentityOp if attention.metainfo["fuse_input_layernorm"] - else backend.layer_norm(rms_norm=rms_norm, for_qk=False) + else _get_standalone_norm(config, backend) ) pre_mlp_layernorm = ( IdentityOp if fuse_pre_mlp_layernorm - else backend.layer_norm(rms_norm=rms_norm, for_qk=False) + else _get_standalone_norm(config, backend) ) layer_specs.append( @@ -278,7 +287,9 @@ def get_transformer_layer_with_experimental_attention_variant_spec( def get_transformer_block_with_experimental_attention_variant_spec( - config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None + config: TransformerConfig, + vp_stage: Optional[int] = None, + pp_rank: Optional[int] = None, ) -> TransformerBlockSubmodules: """Build transformer block spec with experimental attention variants (e.g., linear attention). @@ -316,17 +327,20 @@ def get_transformer_block_with_experimental_attention_variant_spec( layer_type=LayerType.decoder, vp_stage=vp_stage, pp_rank=pp_rank ) else: - offset = get_transformer_layer_offset(config, vp_stage=vp_stage, pp_rank=pp_rank) - num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage, pp_rank=pp_rank) + offset = get_transformer_layer_offset( + config, vp_stage=vp_stage, pp_rank=pp_rank + ) + num_layers_to_build = get_num_layers_to_build( + config, vp_stage=vp_stage, pp_rank=pp_rank + ) local_layer_ids = range(offset, offset + num_layers_to_build) _validate_dsa_index_share_pipeline_split(config, local_layer_ids) layer_specs = [layer_specs[layer_id] for layer_id in local_layer_ids] # Get GPT decoder block spec - rms_norm = config.normalization == "RMSNorm" gpt_decoder_block_spec = TransformerBlockSubmodules( - layer_specs=layer_specs, layer_norm=backend.layer_norm(rms_norm=rms_norm, for_qk=False) + layer_specs=layer_specs, layer_norm=_get_standalone_norm(config, backend) ) return gpt_decoder_block_spec @@ -343,7 +357,9 @@ def is_linear_attention_variant(experimental_attention_variant: Optional[str]) - return experimental_attention_variant in linear_attention_variants -def _validate_dsa_index_share_pipeline_split(config: TransformerConfig, local_layer_ids) -> None: +def _validate_dsa_index_share_pipeline_split( + config: TransformerConfig, local_layer_ids +) -> None: """Ensure DSA top-k sharing does not require top-k indices from another PP stage.""" if ( config.experimental_attention_variant != "dsa" @@ -358,12 +374,16 @@ def _validate_dsa_index_share_pipeline_split(config: TransformerConfig, local_la for position, layer_id in enumerate(local_layer_ids): layer_number = layer_id + 1 if not is_dsa_skip_topk_layer( - layer_number, config.dsa_indexer_skip_topk_offset, config.dsa_indexer_topk_freq + layer_number, + config.dsa_indexer_skip_topk_offset, + config.dsa_indexer_topk_freq, ): continue source_layer_number = source_dsa_compute_layer( - layer_number, config.dsa_indexer_skip_topk_offset, config.dsa_indexer_topk_freq + layer_number, + config.dsa_indexer_skip_topk_offset, + config.dsa_indexer_topk_freq, ) source_layer_id = source_layer_number - 1 if ( @@ -390,7 +410,8 @@ def get_moe_layer_pattern(config: TransformerConfig) -> List[int]: if isinstance(config.moe_layer_freq, int): # [1,0,0,...,0,1,0,0,...,0,...] moe_layer_pattern = [ - 1 if (i % config.moe_layer_freq == 0) else 0 for i in range(config.num_layers) + 1 if (i % config.moe_layer_freq == 0) else 0 + for i in range(config.num_layers) ] elif isinstance(config.moe_layer_freq, list): moe_layer_pattern = config.moe_layer_freq @@ -478,7 +499,9 @@ def _get_self_attention_module_spec( if backend is None: backend = _get_backend_spec_provider(config=config) - from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_spec, + ) layer_spec = get_gpt_layer_with_transformer_engine_spec( num_experts=config.num_moe_experts, @@ -539,7 +562,9 @@ def _get_moe_module_spec( if backend is None: backend = _get_backend_spec_provider(config=config) - from megatron.core.models.gpt.moe_module_specs import get_moe_module_spec_for_backend + from megatron.core.models.gpt.moe_module_specs import ( + get_moe_module_spec_for_backend, + ) return ( get_moe_module_spec_for_backend( diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 32a61cf7efc..c7406191dc2 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -556,6 +556,9 @@ def _get_megatron_optimizer_based_on_param_groups( # on source of optimizer (Torch or TE/Apex) if USING_PYTORCH_OPTIMIZER: adam_cls = torch.optim.AdamW if config.decoupled_weight_decay else torch.optim.Adam + elif config.native_unfused_adamw: + adam_cls = torch.optim.AdamW if config.decoupled_weight_decay else torch.optim.Adam + kwargs.update({"foreach": False, "fused": False}) else: kwargs["adam_w_mode"] = config.decoupled_weight_decay adam_cls = Adam diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 24f9a032c47..a11f5c8cc3b 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -247,6 +247,9 @@ class OptimizerConfig: adam_eps: float = 1e-08 """Term added to the denominator to improve numerical stability in Adam optimizer.""" + native_unfused_adamw: bool = False + """Use torch.optim.AdamW with foreach=False and fused=False instead of TE/Apex Adam.""" + decoupled_weight_decay: bool = True """If true, decouples weight decay from the gradient update, equivalent to AdamW. If false, original Adam update rule will be used. Defaults to True. diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 536c3dfd15e..ee92afe4954 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -11,7 +11,10 @@ from megatron.core import InferenceParams, parallel_state, tensor_parallel from megatron.core.dist_checkpointing.mapping import ShardedStateDict -from megatron.core.dist_checkpointing.utils import apply_prefix_mapping, replace_prefix_for_sharding +from megatron.core.dist_checkpointing.utils import ( + apply_prefix_mapping, + replace_prefix_for_sharding, +) from megatron.core.enums import Fp8Recipe from megatron.core.extensions.transformer_engine import HAVE_TE from megatron.core.fp8_utils import get_fp8_context @@ -30,7 +33,7 @@ from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module -from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm, LayerNormBuilder +from megatron.core.transformer.torch_norm import LayerNormBuilder, WrappedTorchNorm from megatron.core.transformer.transformer_block import TransformerBlockSubmodules from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.typed_torch import apply_module @@ -61,7 +64,9 @@ else: TESpecProvider = None -from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout +from megatron.core.transformer.pipeline_parallel_layer_layout import ( + PipelineParallelLayerLayout, +) def tie_word_embeddings_state_dict( @@ -165,7 +170,9 @@ def roll_tensor(tensor, shifts=-1, dims=-1, cp_group=None, packed_seq_params=Non # Handle packed sequences cases if packed_seq_params is not None: - return _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group) + return _roll_tensor_packed_seq( + tensor, shifts, dims, packed_seq_params, cp_group + ) # Standard rolling behavior when CP is not enabled (cp_group is None or size=1) if cp_group is None or cp_group.size() == 1: @@ -202,17 +209,25 @@ def roll_tensor(tensor, shifts=-1, dims=-1, cp_group=None, packed_seq_params=Non # Start send and recv ops ops = [] if local_rank != 0: - req_send_first_part = torch.distributed.isend(tensor=tensor_send_list[0], dst=prev_rank) + req_send_first_part = torch.distributed.isend( + tensor=tensor_send_list[0], dst=prev_rank + ) ops.append(req_send_first_part) - req_recv_second_part = torch.distributed.irecv(tensor=tensor_recv_list[1], src=prev_rank) + req_recv_second_part = torch.distributed.irecv( + tensor=tensor_recv_list[1], src=prev_rank + ) ops.append(req_recv_second_part) else: # Inserted elements are set to be 0.0. tensor_recv_list[1] = 0 if local_rank != len(global_ranks) - 1: - req_recv_first_part = torch.distributed.irecv(tensor=tensor_recv_list[0], src=next_rank) + req_recv_first_part = torch.distributed.irecv( + tensor=tensor_recv_list[0], src=next_rank + ) ops.append(req_recv_first_part) - req_send_second_part = torch.distributed.isend(tensor=tensor_send_list[1], dst=next_rank) + req_send_second_part = torch.distributed.isend( + tensor=tensor_send_list[1], dst=next_rank + ) ops.append(req_send_second_part) else: # For the last CP rank, the removed elements of second part go into the first part @@ -242,12 +257,14 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No # Notice: This is a naive implementation to test the correctness, # a better solution will only sync the boundary tokens once. - assert ( - dims == -1 or dims == tensor.dim() - 1 - ), "Packed sequence roll only supports the last dimension." + assert dims == -1 or dims == tensor.dim() - 1, ( + "Packed sequence roll only supports the last dimension." + ) assert shifts == -1, "Packed sequence roll only supports a single-token left shift." cu_seqlens = packed_seq_params.cu_seqlens_q - assert cu_seqlens is not None, "Packed sequence parameters must provide cu_seqlens_q." + assert cu_seqlens is not None, ( + "Packed sequence parameters must provide cu_seqlens_q." + ) rolled_tensor = tensor.clone() @@ -289,7 +306,9 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No # The following code is very similar as the code in roll_tensor function local_chunks = tensor_slice.chunk(2, dim=dims) - rolled_chunks = [torch.roll(chunk, shifts=shifts, dims=dims) for chunk in local_chunks] + rolled_chunks = [ + torch.roll(chunk, shifts=shifts, dims=dims) for chunk in local_chunks + ] tensor_send_list = [] tensor_recv_list = [] @@ -297,10 +316,14 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No # Skip empty chunks that can occur when the sequence slice is very small if chunk.size(dims) == 0: tensor_send_list.append( - torch.empty(chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device) + torch.empty( + chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device + ) ) tensor_recv_list.append( - torch.empty(chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device) + torch.empty( + chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device + ) ) continue boundary = chunk.select(dims, shifts).contiguous().clone() @@ -309,14 +332,22 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No ops = [] if local_rank != 0: - ops.append(torch.distributed.isend(tensor=tensor_send_list[0], dst=prev_rank)) - ops.append(torch.distributed.irecv(tensor=tensor_recv_list[1], src=prev_rank)) + ops.append( + torch.distributed.isend(tensor=tensor_send_list[0], dst=prev_rank) + ) + ops.append( + torch.distributed.irecv(tensor=tensor_recv_list[1], src=prev_rank) + ) else: tensor_recv_list[1].zero_() if local_rank != cp_size - 1: - ops.append(torch.distributed.irecv(tensor=tensor_recv_list[0], src=next_rank)) - ops.append(torch.distributed.isend(tensor=tensor_send_list[1], dst=next_rank)) + ops.append( + torch.distributed.irecv(tensor=tensor_recv_list[0], src=next_rank) + ) + ops.append( + torch.distributed.isend(tensor=tensor_send_list[1], dst=next_rank) + ) else: tensor_recv_list[0].copy_(tensor_send_list[1]) @@ -371,11 +402,17 @@ def save_metrics_to_tracker( tracker = MTPLossLoggingHelper.tracker if "loss_values" not in tracker: - tracker["loss_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) + tracker["loss_values"] = torch.zeros( + num_layers, device=torch.cuda.current_device() + ) if "correct_values" not in tracker: - tracker["correct_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) + tracker["correct_values"] = torch.zeros( + num_layers, device=torch.cuda.current_device() + ) if "total_values" not in tracker: - tracker["total_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) + tracker["total_values"] = torch.zeros( + num_layers, device=torch.cuda.current_device() + ) tracker["loss_values"][layer_number] += loss.detach() tracker["correct_values"][layer_number] += correct.detach() @@ -404,26 +441,32 @@ def reduce_metrics_in_tracker(): return loss_values = tracker["loss_values"] - if tracker.get('reduce_group') is not None: - torch.distributed.all_reduce(loss_values, group=tracker.get('reduce_group')) - if tracker.get('avg_group') is not None: + if tracker.get("reduce_group") is not None: + torch.distributed.all_reduce(loss_values, group=tracker.get("reduce_group")) + if tracker.get("avg_group") is not None: torch.distributed.all_reduce( - loss_values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.AVG + loss_values, + group=tracker["avg_group"], + op=torch.distributed.ReduceOp.AVG, ) for key in ["correct_values", "total_values"]: if key not in tracker: continue values = tracker[key] - if tracker.get('reduce_group') is not None: - torch.distributed.all_reduce(values, group=tracker.get('reduce_group')) - if tracker.get('avg_group') is not None: + if tracker.get("reduce_group") is not None: + torch.distributed.all_reduce(values, group=tracker.get("reduce_group")) + if tracker.get("avg_group") is not None: torch.distributed.all_reduce( - values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.SUM + values, + group=tracker["avg_group"], + op=torch.distributed.ReduceOp.SUM, ) @staticmethod - def track_mtp_metrics(loss_scale, iteration, writer, wandb_writer=None, total_loss_dict=None): + def track_mtp_metrics( + loss_scale, iteration, writer, wandb_writer=None, total_loss_dict=None + ): """Track the Multi-Token Prediction (MTP) metrics for logging.""" MTPLossLoggingHelper.reduce_metrics_in_tracker() tracker = MTPLossLoggingHelper.tracker @@ -453,15 +496,16 @@ def track_mtp_metrics(loss_scale, iteration, writer, wandb_writer=None, total_lo mtp_num_layers = mtp_losses.shape[0] for i in range(mtp_num_layers): - loss_name = f"mtp_{i+1} loss" - step_acc_name = f"mtp_{i+1}_acceptance_rate" - cum_acc_name = f"mtp_{i+1}_cumulative_acceptance_rate" + loss_name = f"mtp_{i + 1} loss" + step_acc_name = f"mtp_{i + 1}_acceptance_rate" + cum_acc_name = f"mtp_{i + 1}_cumulative_acceptance_rate" loss = mtp_losses[i] # Empty masks can leave no valid MTP positions, so clamp denominators to avoid NaNs. step_rate = (mtp_corrects[i] / torch.clamp(mtp_totals[i], min=1)) * 100.0 cum_rate = ( - mtp_cumulative_corrects[i] / torch.clamp(mtp_cumulative_totals[i], min=1) + mtp_cumulative_corrects[i] + / torch.clamp(mtp_cumulative_totals[i], min=1) ) * 100.0 if total_loss_dict is not None: @@ -491,7 +535,9 @@ def _mtp_logits_are_vocab_sharded( def _vocab_parallel_argmax( - vocab_parallel_logits: Tensor, tp_group: torch.distributed.ProcessGroup, tp_size: int + vocab_parallel_logits: Tensor, + tp_group: torch.distributed.ProcessGroup, + tp_size: int, ) -> Tensor: """Return global argmax ids from logits sharded across the vocab dimension.""" vocab_shard_size = vocab_parallel_logits.size(-1) @@ -505,9 +551,9 @@ def _vocab_parallel_argmax( stacked_max_vals = torch.stack(gathered_max_vals, dim=0) stacked_argmax = torch.stack(gathered_argmax, dim=0) winning_rank = stacked_max_vals.argmax(dim=0) # [s, b] - winning_local_argmax = torch.gather(stacked_argmax, 0, winning_rank.unsqueeze(0)).squeeze( - 0 - ) # [s, b] + winning_local_argmax = torch.gather( + stacked_argmax, 0, winning_rank.unsqueeze(0) + ).squeeze(0) # [s, b] return winning_rank * vocab_shard_size + winning_local_argmax # [s, b] @@ -534,7 +580,11 @@ def _compute_mtp_acceptance_counts( "tp_group must be provided when computing MTP acceptance counts " "from vocab-sharded logits under tensor model parallelism." ) - tp_size = torch.distributed.get_world_size(group=tp_group) if tp_group is not None else 1 + tp_size = ( + torch.distributed.get_world_size(group=tp_group) + if tp_group is not None + else 1 + ) # Apply TP rank offsets only when logits are vocab-sharded; gathered logits already # contain global vocab ids in their last dimension. @@ -605,7 +655,7 @@ def get_mtp_layer_spec_for_backend( """ column_parallel_linear_impl: type = backend.column_parallel_linear() layer_norm_impl = ( - AccuracyCompatibleRMSNorm + WrappedTorchNorm if config is not None and config.norm_accuracy_compatible else backend.layer_norm() ) @@ -646,14 +696,19 @@ def mtp_on_this_rank( # with custom PP layout, we support put MTP layers on any pipeline stage if ( not ignore_virtual - and parallel_state.get_virtual_pipeline_model_parallel_world_size() is not None + and parallel_state.get_virtual_pipeline_model_parallel_world_size() + is not None ): - assert vp_stage is not None, "vp_stage must be passed if virtual pipeline is enabled" + assert vp_stage is not None, ( + "vp_stage must be passed if virtual pipeline is enabled" + ) num_layers_to_build = layout.layout[pp_rank][vp_stage].count(LayerType.mtp) mtp_on_this_rank = num_layers_to_build > 0 else: for vpp_rank in range(len(layout.layout[pp_rank])): - num_layers_to_build = layout.layout[pp_rank][vpp_rank].count(LayerType.mtp) + num_layers_to_build = layout.layout[pp_rank][vpp_rank].count( + LayerType.mtp + ) if num_layers_to_build > 0: mtp_on_this_rank = True break @@ -684,7 +739,9 @@ def get_mtp_ranks(pp_ranks: List[int], config: TransformerConfig) -> List[int]: return list(mtp_ranks) -def get_mtp_layer_offset(config: TransformerConfig, vp_stage: Optional[int] = None) -> int: +def get_mtp_layer_offset( + config: TransformerConfig, vp_stage: Optional[int] = None +) -> int: """Get the offset of the MTP layer.""" if config.pipeline_model_parallel_size > 1: if config.pipeline_model_parallel_layout: @@ -699,21 +756,29 @@ def get_mtp_layer_offset(config: TransformerConfig, vp_stage: Optional[int] = No def get_mtp_num_layers_to_build( - config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None + config: TransformerConfig, + vp_stage: Optional[int] = None, + pp_rank: Optional[int] = None, ) -> int: """Get the number of MTP layers to build.""" if config.pipeline_model_parallel_layout is not None: # If we have a custom PP layout, get the number of mtp layers in the layout array. - num_layers_to_build = config.pipeline_model_parallel_layout.get_num_layers_to_build( - layer_type=LayerType.mtp, vp_stage=vp_stage + num_layers_to_build = ( + config.pipeline_model_parallel_layout.get_num_layers_to_build( + layer_type=LayerType.mtp, vp_stage=vp_stage + ) ) - assert num_layers_to_build == config.mtp_num_layers or num_layers_to_build == 0, ( + assert ( + num_layers_to_build == config.mtp_num_layers or num_layers_to_build == 0 + ), ( f"Currently, we only support put all of MTP layers on the last pipeline stage, " f"so the number of MTP layers to build ({num_layers_to_build}) must match " f"mtp_num_layers ({config.mtp_num_layers}) or be 0." ) else: - if parallel_state.is_pipeline_last_stage(ignore_virtual=False, vp_stage=vp_stage): + if parallel_state.is_pipeline_last_stage( + ignore_virtual=False, vp_stage=vp_stage + ): num_layers_to_build = config.mtp_num_layers if config.mtp_num_layers else 0 else: num_layers_to_build = 0 @@ -820,7 +885,11 @@ def process_mtp_loss( if input_ids is None: return hidden_states labels, _ = roll_tensor( - input_ids, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params + input_ids, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, ) derived_labels_from_input_ids = True @@ -838,7 +907,11 @@ def process_mtp_loss( # label is fabricated (zeroed). Roll loss_mask in lockstep with the # input_ids -> labels shift so that boundary position is masked. loss_mask, _ = roll_tensor( - loss_mask, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params + loss_mask, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, ) # Store the original number of tokens before rolling for proper normalization @@ -855,10 +928,18 @@ def process_mtp_loss( if scale_logits_fn is not None: mtp_logits = scale_logits_fn(mtp_logits) mtp_labels, _ = roll_tensor( - mtp_labels, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params + mtp_labels, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, ) loss_mask, num_tokens = roll_tensor( - loss_mask, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params + loss_mask, + shifts=-1, + dims=-1, + cp_group=cp_group, + packed_seq_params=packed_seq_params, ) mtp_loss = compute_language_model_loss(mtp_labels, mtp_logits) @@ -870,7 +951,12 @@ def process_mtp_loss( torch.sum(mtp_loss) * (num_tokens > 0).to(mtp_loss.dtype) ) / num_tokens.clamp(min=1) correct, total = _compute_mtp_acceptance_counts( - mtp_logits, mtp_labels, loss_mask, output_layer, runtime_gather_output, tp_group + mtp_logits, + mtp_labels, + loss_mask, + output_layer, + runtime_gather_output, + tp_group, ) MTPLossLoggingHelper.save_metrics_to_tracker( @@ -879,7 +965,9 @@ def process_mtp_loss( total, mtp_layer_number, config.mtp_num_layers, - avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), + avg_group=parallel_state.get_data_parallel_group( + with_context_parallel=True + ), ) mtp_loss_scale = config.mtp_loss_scaling_factor / config.mtp_num_layers if config.calculate_per_token_loss: @@ -963,18 +1051,28 @@ def __init__( # Validate attention mask type if using transformer-based inner layers if self.submodules.mtp_model_layer is not None and hasattr( - self.submodules.mtp_model_layer, 'submodules' + self.submodules.mtp_model_layer, "submodules" ): from megatron.core.models.hybrid.hybrid_block import HybridStackSubmodules - from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules + from megatron.core.transformer.transformer_layer import ( + TransformerLayerSubmodules, + ) layer_submodules = None - if isinstance(self.submodules.mtp_model_layer.submodules, HybridStackSubmodules): - attention_layer_spec = self.submodules.mtp_model_layer.submodules.attention_layer - if hasattr(attention_layer_spec, 'submodules'): - assert isinstance(attention_layer_spec.submodules, TransformerLayerSubmodules) + if isinstance( + self.submodules.mtp_model_layer.submodules, HybridStackSubmodules + ): + attention_layer_spec = ( + self.submodules.mtp_model_layer.submodules.attention_layer + ) + if hasattr(attention_layer_spec, "submodules"): + assert isinstance( + attention_layer_spec.submodules, TransformerLayerSubmodules + ) layer_submodules = attention_layer_spec.submodules - elif isinstance(self.submodules.mtp_model_layer.submodules, TransformerLayerSubmodules): + elif isinstance( + self.submodules.mtp_model_layer.submodules, TransformerLayerSubmodules + ): layer_submodules = self.submodules.mtp_model_layer.submodules else: raise ValueError( @@ -982,7 +1080,7 @@ def __init__( ) if layer_submodules: self_attention_spec = layer_submodules.self_attention - attn_mask_type = self_attention_spec.params.get('attn_mask_type', '') + attn_mask_type = self_attention_spec.params.get("attn_mask_type", "") assert attn_mask_type in SUPPORTED_ATTN_MASK, ( f"Multi-Token Prediction (MTP) is not yet supported with " f"{attn_mask_type} attention mask type. " @@ -1026,7 +1124,9 @@ def __init__( # 2. GPT path: single TransformerLayer if mtp_layer_pattern is not None and hybrid_submodules is not None: from megatron.core.models.hybrid.hybrid_block import HybridStack - from megatron.core.models.hybrid.hybrid_layer_allocation import validate_segment_layers + from megatron.core.models.hybrid.hybrid_layer_allocation import ( + validate_segment_layers, + ) self.mtp_model_layer = HybridStack( config=self.config, @@ -1115,7 +1215,9 @@ def _get_embeddings( if self.config.mtp_detach_heads: decoder_input = decoder_input.detach() - hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # make_viewless_tensor no-ops when hidden_states is not a view (_base is None), # which happens after detach() with mtp_detach_heads. Activation # checkpointing (CheckpointFunction.apply) requires at least one input tensor @@ -1126,14 +1228,20 @@ def _get_embeddings( return input_ids, position_ids, padding_mask, decoder_input, hidden_states - def _concat_embeddings(self, hidden_states: torch.Tensor, decoder_input: torch.Tensor): + def _concat_embeddings( + self, hidden_states: torch.Tensor, decoder_input: torch.Tensor + ): """ Concatenate the tokens before sending to transformer layer. """ decoder_input = apply_module(self.enorm)(decoder_input) - decoder_input = make_viewless_tensor(inp=decoder_input, requires_grad=True, keep_graph=True) + decoder_input = make_viewless_tensor( + inp=decoder_input, requires_grad=True, keep_graph=True + ) hidden_states = apply_module(self.hnorm)(hidden_states) - hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # At the (k - 1)-th MTP module, concatenates the i-th token's hidden_states # and the (i + K)-th token's embedding, and combine them with linear projection. hidden_states = torch.cat((decoder_input, hidden_states), -1) @@ -1150,7 +1258,9 @@ def _concat_embeddings(self, hidden_states: torch.Tensor, decoder_input: torch.T ) # For sequence parallel, scatter after linear_fc and before transformer layer. if self.sequence_parallel: - hidden_states = scatter_to_sequence_parallel_region(hidden_states, group=self.tp_group) + hidden_states = scatter_to_sequence_parallel_region( + hidden_states, group=self.tp_group + ) return hidden_states def _proj_and_transformer_layer( @@ -1235,7 +1345,9 @@ def _postprocess(self, hidden_states: torch.Tensor): # TENorm produces a "viewed" tensor. This will result in schedule.py's # deallocate_output_tensor() throwing an error, so a viewless tensor is # created to prevent this. - hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) return hidden_states @@ -1413,19 +1525,20 @@ def checkpoint_handler(): sequence_len_offset, ) - if self.config.recompute_method == 'uniform': + if self.config.recompute_method == "uniform": # Uniformly divide the total number of Transformer layers and checkpoint # the input activation of each divided chunk. # A method to further reduce memory usage reducing checkpoints. - assert ( - self.config.recompute_num_layers == 1 - ), "recompute_num_layers must be 1 for MTP recompute" + assert self.config.recompute_num_layers == 1, ( + "recompute_num_layers must be 1 for MTP recompute" + ) with outer_quantization_context: outputs = checkpoint_handler() - elif self.config.recompute_method == 'block': + elif self.config.recompute_method == "block": # TODO: implement block-based recompute for MTP warnings.warn( - "recompute_method == 'block' is not supported for MTP yet." " Skipping recompute." + "recompute_method == 'block' is not supported for MTP yet." + " Skipping recompute." ) outputs = self._proj_and_transformer_layer( hidden_states=hidden_states, @@ -1487,17 +1600,21 @@ def forward( Union[Tensor, Tuple[Tensor, Tensor]]: The output hidden states tensor of shape [s, b, h], and optionally the updated context tensor if cross-attention is used. """ - assert context is None, "multi token prediction + cross attention is not yet supported." - input_ids, position_ids, padding_mask, decoder_input, hidden_states = self._get_embeddings( - input_ids=input_ids, - position_ids=position_ids, - padding_mask=padding_mask, - embedding=embedding, - hidden_states=hidden_states, - packed_seq_params=packed_seq_params, + assert context is None, ( + "multi token prediction + cross attention is not yet supported." + ) + input_ids, position_ids, padding_mask, decoder_input, hidden_states = ( + self._get_embeddings( + input_ids=input_ids, + position_ids=position_ids, + padding_mask=padding_mask, + embedding=embedding, + hidden_states=hidden_states, + packed_seq_params=packed_seq_params, + ) ) - if self.config.recompute_granularity == 'full' and self.training: + if self.config.recompute_granularity == "full" and self.training: hidden_states = self._checkpointed_forward( hidden_states=hidden_states, decoder_input=decoder_input, @@ -1533,7 +1650,10 @@ def forward( return hidden_states, input_ids, position_ids, padding_mask def sharded_state_dict( - self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None + self, + prefix: str = "", + sharded_offsets: tuple = (), + metadata: Optional[dict] = None, ) -> ShardedStateDict: """ Generate a sharded state dictionary for the multi token prediction layer. @@ -1547,7 +1667,9 @@ def sharded_state_dict( ShardedStateDict: A dictionary containing the sharded state of the multi token prediction layer. """ - sharded_state_dict = super().sharded_state_dict(prefix, sharded_offsets, metadata) + sharded_state_dict = super().sharded_state_dict( + prefix, sharded_offsets, metadata + ) # Backward compatibility: GPT MTP checkpoints were saved with the submodule # named 'transformer_layer'. Remap checkpoint keys so old checkpoints load @@ -1555,7 +1677,8 @@ def sharded_state_dict( # since no older checkpoints exist for them. if self.mtp_layer_pattern is None: apply_prefix_mapping( - sharded_state_dict, {f'{prefix}mtp_model_layer.': f'{prefix}transformer_layer.'} + sharded_state_dict, + {f"{prefix}mtp_model_layer.": f"{prefix}transformer_layer."}, ) return sharded_state_dict @@ -1580,7 +1703,8 @@ class MultiTokenPredictionBlockSubmodules: def _get_mtp_block_submodules( - config: TransformerConfig, spec: Union[MultiTokenPredictionBlockSubmodules, ModuleSpec] + config: TransformerConfig, + spec: Union[MultiTokenPredictionBlockSubmodules, ModuleSpec], ) -> MultiTokenPredictionBlockSubmodules: """ Retrieve or construct MultiTokenPredictionBlockSubmodules based on the provided specification. @@ -1681,21 +1805,25 @@ def __init__( # to the roll_tensor function for proper boundary communication if pg_collection is None: # Use default MPU process groups if not provided - pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp', 'tp']) + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["cp", "tp"] + ) else: # Ensure the provided process groups include CP - assert hasattr( - pg_collection, 'cp' - ), "MultiTokenPredictionBlock pg_collection must have cp process group" + assert hasattr(pg_collection, "cp"), ( + "MultiTokenPredictionBlock pg_collection must have cp process group" + ) self._build_layers(pg_collection) - assert len(self.layers) > 0, "MultiTokenPredictionBlock must have at least one layer." + assert len(self.layers) > 0, ( + "MultiTokenPredictionBlock must have at least one layer." + ) self.cp_group = pg_collection.cp if self.config.mtp_detach_heads: # Tag MTP params so the optimizer can clip their gradients separately. for param in self.parameters(): - param.grad_norm_group = 'mtp' + param.grad_norm_group = "mtp" def _build_layers(self, pg_collection): # Determine number of depths to build @@ -1715,7 +1843,9 @@ def build_layer_legacy(layer_spec, layer_number): vp_stage=self.vp_stage, pg_collection=pg_collection, mtp_layer_pattern=self.mtp_layer_pattern, - name=(self.name + f".layers.{layer_number}") if self.name is not None else None, + name=(self.name + f".layers.{layer_number}") + if self.name is not None + else None, ) return module @@ -1733,7 +1863,9 @@ def build_layer_with_pattern( pg_collection=pg_collection, mtp_layer_pattern=mtp_layer_pattern, hybrid_submodules=hybrid_submodules, - name=(self.name + f".layers.{layer_number}") if self.name is not None else None, + name=(self.name + f".layers.{layer_number}") + if self.name is not None + else None, ) return module @@ -1825,7 +1957,9 @@ def forward( for iteration in range(self.config.mtp_num_layers): layer_idx = 0 if self.mtp_use_repeated_layer else iteration - (hidden_states, input_ids, position_ids, padding_mask) = self.layers[layer_idx]( + (hidden_states, input_ids, position_ids, padding_mask) = self.layers[ + layer_idx + ]( input_ids=input_ids, position_ids=position_ids, hidden_states=hidden_states, @@ -1850,7 +1984,10 @@ def forward( return hidden_states def sharded_state_dict( - self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None + self, + prefix: str = "", + sharded_offsets: tuple = (), + metadata: Optional[dict] = None, ) -> ShardedStateDict: """ Generate a sharded state dictionary for the multi token prediction module. @@ -1865,16 +2002,18 @@ def sharded_state_dict( token prediction module. """ sharded_state_dict = {} - layer_prefix = f'{prefix}layers.' + layer_prefix = f"{prefix}layers." for layer in self.layers: offset = get_mtp_layer_offset(self.config, self.vp_stage) - sharded_prefix = f'{layer_prefix}{layer.layer_number - 1}.' + sharded_prefix = f"{layer_prefix}{layer.layer_number - 1}." - state_dict_prefix = f'{layer_prefix}{layer.layer_number - 1 - offset}.' + state_dict_prefix = f"{layer_prefix}{layer.layer_number - 1 - offset}." sharded_pp_offset = [] layer_sharded_state_dict = layer.sharded_state_dict( state_dict_prefix, sharded_pp_offset, metadata ) - replace_prefix_for_sharding(layer_sharded_state_dict, state_dict_prefix, sharded_prefix) + replace_prefix_for_sharding( + layer_sharded_state_dict, state_dict_prefix, sharded_prefix + ) sharded_state_dict.update(layer_sharded_state_dict) return sharded_state_dict diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py index c0ddf0763fc..c75525dcd59 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -24,54 +24,6 @@ def __call__( ) -> LayerNormInterface: ... -class _AccuracyCompatibleRMSNormFunction(torch.autograd.Function): - """RMSNorm core with a stable fp32 backward and canonical zero gradients.""" - - @staticmethod - def forward(ctx, x: torch.Tensor, eps: float) -> torch.Tensor: - variance = x.pow(2).mean(dim=-1, keepdim=True) - inv_rms = torch.rsqrt(variance + eps) - ctx.save_for_backward(x, inv_rms) - return x * inv_rms - - @staticmethod - def backward(ctx, grad_output: torch.Tensor): - x, inv_rms = ctx.saved_tensors - dot = (grad_output * x).sum(dim=-1, keepdim=True) - correction_scale = dot * (-0.5) * inv_rms.pow(3) / x.shape[-1] - correction = (correction_scale * x) * 2.0 - grad_input = grad_output * inv_rms + correction - grad_input = torch.where(grad_input == 0, torch.zeros_like(grad_input), grad_input) - return grad_input, None - - -class AccuracyCompatibleRMSNorm(torch.nn.Module, LayerNormInterface): - """RMSNorm with explicit fp32 reduction and one output cast.""" - - def __init__( - self, - normalized_shape: int | None = None, - eps: float = 1e-5, - *, - hidden_size: int | None = None, - config: TransformerConfig | None = None, - **kwargs, - ): - super().__init__() - normalized_shape = hidden_size if normalized_shape is None else normalized_shape - if normalized_shape is None: - raise ValueError("normalized_shape or hidden_size is required") - self.normalized_shape = (normalized_shape,) - self.eps = eps - dtype = config.params_dtype if config is not None else None - self.weight = torch.nn.Parameter(torch.ones(normalized_shape, dtype=dtype)) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - x_float = x.float() - output = _AccuracyCompatibleRMSNormFunction.apply(x_float, self.eps) - return (output * self.weight.float()).to(x.dtype) - - class WrappedTorchNorm: """ A conditional wrapper to initialize an instance of PyTorch's @@ -89,34 +41,41 @@ def __new__( zero_centered_gamma: bool = False, normalization: str = "LayerNorm", ) -> LayerNormInterface: - assert ( - not config.layernorm_zero_centered_gamma - ), f"zero_centered_gamma not supported by torch LayerNorm" - - assert not config.persist_layer_norm, f"persist_layer_norm not supported by torch LayerNorm" + assert not config.layernorm_zero_centered_gamma, ( + f"zero_centered_gamma not supported by torch LayerNorm" + ) - assert not config.sequence_parallel, f"sequence parallel not supported by torch LayerNorm" + assert not config.persist_layer_norm, ( + f"persist_layer_norm not supported by torch LayerNorm" + ) - assert ( - not config.memory_efficient_layer_norm - ), f"memory_efficient_layer_norm not supported by torch LayerNorm" + assert not config.memory_efficient_layer_norm, ( + f"memory_efficient_layer_norm not supported by torch LayerNorm" + ) if config.normalization == "LayerNorm": norm_cls = torch.nn.LayerNorm elif config.normalization == "RMSNorm": - if config.norm_accuracy_compatible: - return AccuracyCompatibleRMSNorm(normalized_shape=hidden_size, eps=eps) - assert is_torch_min_version( - "2.4.0a0" - ), 'Torch RMSNorm requires PyTorch version >= 2.4.0' + assert is_torch_min_version("2.4.0a0"), ( + "Torch RMSNorm requires PyTorch version >= 2.4.0" + ) norm_cls = torch.nn.RMSNorm elif config.normalization == "L2Norm": norm_cls = torch.nn.L2Norm else: - raise Exception("Only LayerNorm, RMSNorm and L2Norm are currently supported") + raise Exception( + "Only LayerNorm, RMSNorm and L2Norm are currently supported" + ) - return norm_cls(normalized_shape=hidden_size, eps=eps) + factory_kwargs = {} + if config.normalization == "RMSNorm" and config.norm_accuracy_compatible: + factory_kwargs["dtype"] = config.params_dtype + norm = norm_cls(normalized_shape=hidden_size, eps=eps, **factory_kwargs) + if config.sequence_parallel: + for parameter in norm.parameters(): + parameter.sequence_parallel = True + return norm class L2Norm(torch.nn.Module, LayerNormInterface): @@ -149,7 +108,9 @@ def _norm(self, x: torch.Tensor) -> torch.Tensor: torch.Tensor: The L2-normalized tensor. """ x_float = x.float() - return (x_float * torch.rsqrt(x_float.pow(2).mean(-1, keepdim=True) + self.eps)).type_as(x) + return ( + x_float * torch.rsqrt(x_float.pow(2).mean(-1, keepdim=True) + self.eps) + ).type_as(x) def forward(self, x: torch.Tensor) -> torch.Tensor: """ diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index b2fb6da4b05..a8e083341c2 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -25,7 +25,9 @@ CudaGraphScope, InferenceCudaGraphScope, ) -from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout +from megatron.core.transformer.pipeline_parallel_layer_layout import ( + PipelineParallelLayerLayout, +) from .._rank_utils import log_single_rank from ..fusions.fused_bias_geglu import quick_gelu @@ -101,7 +103,9 @@ class TransformerConfig(ModelParallelConfig): """Number of transformer layers on last pipeline stage. None implies equal layer division across PP ranks.""" - pipeline_model_parallel_layout: Optional[Union[str, list, PipelineParallelLayerLayout]] = None + pipeline_model_parallel_layout: Optional[ + Union[str, list, PipelineParallelLayerLayout] + ] = None """Custom definition of the pipeline parallel partitioning. Support type: - str: e.g., 'Et*3|(tt|)*29,m|L'. Stages are split by '|', replicated stages or layers @@ -138,7 +142,9 @@ class TransformerConfig(ModelParallelConfig): hidden_size: int = field(default=0, metadata={"argparse_meta": {"default": None}}) """Transformer hidden size.""" - num_attention_heads: int = field(default=0, metadata={"argparse_meta": {"default": None}}) + num_attention_heads: int = field( + default=0, metadata={"argparse_meta": {"default": None}} + ) """Number of transformer attention heads.""" attention_backend: AttnBackend = AttnBackend.auto @@ -150,7 +156,7 @@ class TransformerConfig(ModelParallelConfig): softmax_scale: Optional[float] = None """Softmax scale for attention scaling.""" - softmax_type: Literal['vanilla', 'off-by-one', 'learnable'] = 'vanilla' + softmax_type: Literal["vanilla", "off-by-one", "learnable"] = "vanilla" """Applies modified softmax from https://www.evanmiller.org/attention-is-off-by-one.html. Supports both TE FusedAttention and local unfused attention. Supports both a fixed offset and and learnable offset.""" @@ -187,23 +193,27 @@ class TransformerConfig(ModelParallelConfig): """Epsilon value for any LayerNorm/RMSNorm operations.""" norm_accuracy_compatible: bool = field( - default=False, metadata={"argparse_meta": {"arg_names": ["--norm-accuracy-compatible"]}} + default=False, + metadata={"argparse_meta": {"arg_names": ["--norm-accuracy-compatible"]}}, ) - """Use explicit fp32 normalization formulas instead of native norm kernels for alignment.""" + """Use native Torch RMSNorm modules instead of Transformer Engine norm modules for alignment.""" router_accuracy_compatible: bool = field( - default=False, metadata={"argparse_meta": {"arg_names": ["--router-accuracy-compatible"]}} + default=False, + metadata={"argparse_meta": {"arg_names": ["--router-accuracy-compatible"]}}, ) """Use an explicit fp32 router GEMM instead of the fused Transformer Engine path.""" layernorm_zero_centered_gamma: bool = field( - default=False, metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}} + default=False, + metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}}, ) """If set to True, the LayerNorm is adjusted to center the gamma values around 0. This improves numerical stability.""" add_bias_linear: bool = field( - default=True, metadata={"argparse_meta": {"arg_names": ["--disable-bias-linear"]}} + default=True, + metadata={"argparse_meta": {"arg_names": ["--disable-bias-linear"]}}, ) """Include/exclude a bias term in all linear layers (QKV projections, after core attention, and two in MLP layer).""" @@ -246,7 +256,7 @@ class TransformerConfig(ModelParallelConfig): - An integer N: Represents a (N-1):1 ratio, one full attention layer after (N-1) SWA layers. - A list that defines a custom pattern, e.g.: [1,1,1,1,0,0,0,0], where 1 represents SWA. """ - normalization: Literal['LayerNorm', 'RMSNorm'] = "LayerNorm" + normalization: Literal["LayerNorm", "RMSNorm"] = "LayerNorm" """Which norm to use for normalization layers, valid options are `LayerNorm` and `RMSNorm`.""" qk_layernorm: bool = False @@ -291,10 +301,12 @@ class TransformerConfig(ModelParallelConfig): #################### # attention variant #################### - experimental_attention_variant: Optional[Literal['gated_delta_net', 'dsa']] = None + experimental_attention_variant: Optional[Literal["gated_delta_net", "dsa"]] = None """Type of attention variant to use. Currently support gated_delta_net and dsa.""" - experimental_attention_variant_loss_scale_func: Optional[Callable[[torch.Tensor], None]] = None + experimental_attention_variant_loss_scale_func: Optional[ + Callable[[torch.Tensor], None] + ] = None """Optional hook for experimental attention variants to receive the main loss scale.""" #################### @@ -329,7 +341,8 @@ class TransformerConfig(ModelParallelConfig): backend. Unsupported DSA layouts continue to use the PyTorch fallback.""" dsa_accuracy_compatible: bool = field( - default=False, metadata={"argparse_meta": {"arg_names": ["--dsa-accuracy-compatible"]}} + default=False, + metadata={"argparse_meta": {"arg_names": ["--dsa-accuracy-compatible"]}}, ) """Use the full-score DSA fallback with explicit softmax backward for alignment.""" @@ -520,7 +533,7 @@ class TransformerConfig(ModelParallelConfig): #################### # activation recomputation #################### - recompute_granularity: Optional[Literal['full', 'selective']] = None + recompute_granularity: Optional[Literal["full", "selective"]] = None """Determines which type of activation recompute to use. Megatron-core supports 'selective' activation checkpointing where the submodules set in --recompute-modules is checkpointed. The default is "core_attn" which is the memory intensive part of attention. @@ -531,7 +544,7 @@ class TransformerConfig(ModelParallelConfig): If set, must be 'selective' or 'full'. 'selective' always uses all layers. """ - recompute_method: Optional[Literal['uniform', 'block']] = None + recompute_method: Optional[Literal["uniform", "block"]] = None """Determines which transformer layers will be recomputed. uniform will uniformly divide the total number of transformer layers in a transformer block and recompute the input activation of each divided chunk at the specified granularity. block will recompute the input activations for @@ -568,16 +581,16 @@ class TransformerConfig(ModelParallelConfig): #################### # fp8 related #################### - fp8: Optional[Literal['e4m3', 'hybrid']] = field( + fp8: Optional[Literal["e4m3", "hybrid"]] = field( default=None, metadata={"argparse_meta": {"arg_names": ["--fp8-format"]}} ) """If set, enables the use of FP8 precision through Transformer Engine. There are 2 predefined choices (1) 'e4m3' uniformly uses e4m3 for all FP8 tensors, (2) 'hybrid' uses e4m3 for all FP8 activation and weight tensors and e5m2 for all FP8 output activation gradient tensors.""" - fp8_recipe: Optional[Literal['tensorwise', 'delayed', 'mxfp8', 'blockwise', 'custom']] = ( - "delayed" - ) + fp8_recipe: Optional[ + Literal["tensorwise", "delayed", "mxfp8", "blockwise", "custom"] + ] = "delayed" """If set, enables the use of FP8 precision through Transformer Engine. There are 5 predefined choices (1) 'tensorwise' uses per tensor current scaling recipe, (2) 'delayed' uses delayed scaling recipe, 3) 'mxfp8' for Blackwell architecture only, @@ -605,7 +618,7 @@ class TransformerConfig(ModelParallelConfig): fp8_amax_history_len: int = 1 """The length of the amax history window used for scaling factor computation.""" - fp8_amax_compute_algo: Literal['most_recent', 'max'] = "most_recent" + fp8_amax_compute_algo: Literal["most_recent", "max"] = "most_recent" """Algorithm used for choosing the `amax` value for the scaling factor computation. There are 2 predefined choices: `max` chooses the largest `amax` in the history window, while `most_recent` always chooses the most recently seen value. @@ -654,13 +667,13 @@ class TransformerConfig(ModelParallelConfig): #################### # fp4 related #################### - fp4: Optional[Literal['e2m1']] = field( + fp4: Optional[Literal["e2m1"]] = field( default=None, metadata={"argparse_meta": {"arg_names": ["--fp4-format"]}} ) """If set, enables the use of FP4 precision through Transformer Engine. Currently only supports 'nvfp4' which uses NVFP4BlockScaling recipe (requires TE >= 2.7.0.dev0).""" - fp4_recipe: Optional[Literal['nvfp4', 'custom']] = "nvfp4" + fp4_recipe: Optional[Literal["nvfp4", "custom"]] = "nvfp4" """If set, enables the use of FP4 precision through Transformer Engine. Currently only 'nvfp4' is supported which uses NVFP4BlockScaling recipe for Blackwell+ architecture.""" @@ -774,10 +787,10 @@ class TransformerConfig(ModelParallelConfig): """Scaling factor for routing score in top-k selection, only works when moe_router_pre_softmax enabled. Defaults to None, which means no scaling.""" - moe_router_score_function: Literal['softmax', 'sigmoid', 'sqrtsoftplus'] = "softmax" + moe_router_score_function: Literal["softmax", "sigmoid", "sqrtsoftplus"] = "softmax" """Score function for MoE routing. Can be "softmax", "sigmoid" or "sqrtsoftplus".""" - moe_router_dtype: Optional[Literal['fp32', 'fp64']] = None + moe_router_dtype: Optional[Literal["fp32", "fp64"]] = None """Data type for routing and expert output weighted averaging. Using fp32 or fp64 can improve stability especially when the number of experts is large (e.g. finegrained-moe). None means no changes for dtype.""" @@ -831,7 +844,9 @@ class TransformerConfig(ModelParallelConfig): If a list of load balancing types is provided for `moe_router_load_balancing_type`, a corresponding list of coefficients should be provided here.""" - moe_z_loss_coeff: Optional[float] = None # 1e-3 would be a good start value for z-loss + moe_z_loss_coeff: Optional[float] = ( + None # 1e-3 would be a good start value for z-loss + ) """Scaling coefficient for the z-loss. A starting value of 1e-3 is recommended.""" moe_input_jitter_eps: Optional[float] = None @@ -842,14 +857,14 @@ class TransformerConfig(ModelParallelConfig): specified capacity, similar to GShard, Switch-Transformer, and DeepSpeed-MoE. Note that this is currently unsupported so should remain False.""" - moe_token_dispatcher_type: Literal['allgather', 'alltoall', 'flex'] = "allgather" + moe_token_dispatcher_type: Literal["allgather", "alltoall", "flex"] = "allgather" """The type of token dispatcher to use. The default is 'allgather'. Options are 'allgather','alltoall' and 'flex'.""" moe_enable_deepep: bool = False """[Experimental] Enable DeepEP for efficient token dispatching and combine in MoE models.""" - moe_flex_dispatcher_backend: Literal['deepep', 'hybridep'] = "deepep" + moe_flex_dispatcher_backend: Literal["deepep", "hybridep"] = "deepep" """[Experimental] The backend to use for flex token dispatcher. The default is "deepep". Options are "deepep" and "hybridep". Currently only "hybridep" backend supports the MNNVL case.""" @@ -875,7 +890,7 @@ class TransformerConfig(ModelParallelConfig): max that an expert could see during inference so no tokens are actually dropped. The default setting is False.""" - moe_token_drop_policy: Literal['probs', 'position'] = "probs" + moe_token_drop_policy: Literal["probs", "position"] = "probs" """The policy to drop tokens. Can be either "probs" or "position". If "probs", the tokens with the lowest probabilities will be dropped. If "position", tokens at the end of each batch will be dropped. @@ -977,7 +992,9 @@ class TransformerConfig(ModelParallelConfig): """DEPRECATED and replaced by cuda_graph_impl. When set to true, TransformerLayer layers are swapped with user provided CUDA graphs.""" - cuda_graph_impl: Literal['none', 'local', 'transformer_engine', 'full_iteration'] = "none" + cuda_graph_impl: Literal[ + "none", "local", "transformer_engine", "full_iteration" + ] = "none" """Determines the CUDA graph capture implementation. "none": no CUDA graph. "local": MCore CUDA graph implementation. During training, graphable modules own per-layer @@ -992,7 +1009,9 @@ class TransformerConfig(ModelParallelConfig): cuda_graph_modules has no effect when cuda_graph_impl="none" and must be empty when cuda_graph_impl="full_iteration".""" - cuda_graph_modules: Union[str, CudaGraphModule, List[str], List[CudaGraphModule]] = "full" + cuda_graph_modules: Union[ + str, CudaGraphModule, List[str], List[CudaGraphModule] + ] = "full" """Selects training capture coverage within per-layer CUDA graphs (local and transformer_engine implementations). Valid values are "attn", "mlp", "moe", "moe_router", "moe_preprocess", and "mamba": @@ -1033,7 +1052,10 @@ class TransformerConfig(ModelParallelConfig): cuda_graph_scope: Optional[ Union[ - str, CudaGraphModule, CudaGraphScope, List[Union[str, CudaGraphModule, CudaGraphScope]] + str, + CudaGraphModule, + CudaGraphScope, + List[Union[str, CudaGraphModule, CudaGraphScope]], ] ] = None """Deprecated: renamed to cuda_graph_modules. Accepted for backward compatibility and @@ -1076,7 +1098,9 @@ class TransformerConfig(ModelParallelConfig): inference_sampling_seed: int = 42 """ Random seed to use for sampling during inference. """ - symmetric_ar_type: Optional[Literal['two_shot', "one_shot", "multimem_all_reduce"]] = None + symmetric_ar_type: Optional[ + Literal["two_shot", "one_shot", "multimem_all_reduce"] + ] = None """What type of symmetric all reduce to use. The default is None which is no use of symmetric memory. """ @@ -1093,7 +1117,7 @@ class TransformerConfig(ModelParallelConfig): inference_disable_triton_nvls_kernels: bool = False """ If true, disables the use of Triton NVLS kernels during inference. """ - inference_grouped_gemm_backend: Literal['flashinfer', 'torch', 'vllm'] = "vllm" + inference_grouped_gemm_backend: Literal["flashinfer", "torch", "vllm"] = "vllm" """Specifies the backend to use for grouped GEMM operations during inference. Options: - 'flashinfer': Uses FlashInfer cutlass_fused_moe. Not compatible with MXFP8. @@ -1109,7 +1133,7 @@ class TransformerConfig(ModelParallelConfig): fp8_recipe='mxfp8'. Set to True to disable fusion and use separate kernel launches (useful for debugging).""" - inference_moe_token_dispatcher_type: Literal['nccl', 'nvls'] = 'nvls' + inference_moe_token_dispatcher_type: Literal["nccl", "nvls"] = "nvls" """Token dispatcher to use for MoE expert parallelism during inference. - 'nccl': AllGather/ReduceScatter via NCCL. Fixed token counts per rank; requires decode-only CUDA graphs (forced automatically). @@ -1142,7 +1166,8 @@ class TransformerConfig(ModelParallelConfig): None causes the states to follow the activation dtype.""" use_mamba_mem_eff_path: bool = field( - default=True, metadata={"argparse_meta": {"arg_names": ["--disable-mamba-mem-eff-path"]}} + default=True, + metadata={"argparse_meta": {"arg_names": ["--disable-mamba-mem-eff-path"]}}, ) """Controls usage of the memory efficient path for Mamba layers.""" @@ -1165,7 +1190,7 @@ class TransformerConfig(ModelParallelConfig): quant_recipe: Optional[RecipeConfig] = None """Configuration of any per-module quantization settings to be applied to the model""" - transformer_impl: Literal['local', 'transformer_engine', 'inference_optimized'] = ( + transformer_impl: Literal["local", "transformer_engine", "inference_optimized"] = ( "transformer_engine" ) """Transformer implementation to use. @@ -1293,26 +1318,26 @@ def __post_init__(self): ) if self.experimental_attention_variant == "gated_delta_net": - assert ( - self.linear_attention_freq is not None - ), f"linear_attention_freq must be set for linear gated_delta_net." + assert self.linear_attention_freq is not None, ( + f"linear_attention_freq must be set for linear gated_delta_net." + ) # Check required parameters - assert ( - self.linear_conv_kernel_dim is not None - ), "linear_conv_kernel_dim must be set for gated delta net." - assert ( - self.linear_key_head_dim is not None - ), "linear_key_head_dim must be set for gated delta net." - assert ( - self.linear_value_head_dim is not None - ), "linear_value_head_dim must be set for gated delta net." - assert ( - self.linear_num_key_heads is not None - ), "linear_num_key_heads must be set for gated delta net." - assert ( - self.linear_num_value_heads is not None - ), "linear_num_value_heads must be set for gated delta net." + assert self.linear_conv_kernel_dim is not None, ( + "linear_conv_kernel_dim must be set for gated delta net." + ) + assert self.linear_key_head_dim is not None, ( + "linear_key_head_dim must be set for gated delta net." + ) + assert self.linear_value_head_dim is not None, ( + "linear_value_head_dim must be set for gated delta net." + ) + assert self.linear_num_key_heads is not None, ( + "linear_num_key_heads must be set for gated delta net." + ) + assert self.linear_num_value_heads is not None, ( + "linear_num_value_heads must be set for gated delta net." + ) assert self.linear_num_value_heads % self.linear_num_key_heads == 0, ( f"linear_num_value_heads ({self.linear_num_value_heads}) must be a multiple of " f"linear_num_key_heads ({self.linear_num_key_heads})." @@ -1348,7 +1373,9 @@ def __post_init__(self): if self.fp8: # cannot support first last layer bf16 with delayed scaling if self.first_last_layers_bf16 and self.fp8_recipe == Fp8Recipe.delayed: - raise ValueError("Delayed scaling does not support first / last layer in BF16.") + raise ValueError( + "Delayed scaling does not support first / last layer in BF16." + ) # max bf16 layers per pipeline stage max_bf16_layers_per_pipeline_stage = ( @@ -1359,7 +1386,8 @@ def __post_init__(self): if self.first_last_layers_bf16: if ( self.num_layers_at_start_in_bf16 < 0 - or self.num_layers_at_start_in_bf16 > max_bf16_layers_per_pipeline_stage + or self.num_layers_at_start_in_bf16 + > max_bf16_layers_per_pipeline_stage ): raise ValueError( f"num_layers_at_start_in_bf16 ({self.num_layers_at_start_in_bf16}) must be " @@ -1368,7 +1396,8 @@ def __post_init__(self): ) if ( self.num_layers_at_end_in_bf16 < 0 - or self.num_layers_at_end_in_bf16 > max_bf16_layers_per_pipeline_stage + or self.num_layers_at_end_in_bf16 + > max_bf16_layers_per_pipeline_stage ): raise ValueError( f"num_layers_at_end_in_bf16 ({self.num_layers_at_end_in_bf16}) must be " @@ -1392,7 +1421,8 @@ def __post_init__(self): raise ValueError("fp8_output_proj must be used together with fp8 mode.") if self.fp8_recipe != Fp8Recipe.mxfp8: raise ValueError( - f"fp8_output_proj requires fp8_recipe='mxfp8', got " f"'{self.fp8_recipe}'." + f"fp8_output_proj requires fp8_recipe='mxfp8', got " + f"'{self.fp8_recipe}'." ) # FP4 validation @@ -1400,7 +1430,9 @@ def __post_init__(self): raise ValueError("fp4_param must be used together with fp4 mode.") if self.fp4 and self.fp8: - raise ValueError("fp4 and fp8 cannot be used simultaneously. Please choose one.") + raise ValueError( + "fp4 and fp8 cannot be used simultaneously. Please choose one." + ) if self.fp4 and self.fp4_recipe == Fp4Recipe.custom: if not self.fp4_quantizer_factory: @@ -1416,13 +1448,18 @@ def __post_init__(self): if self.expert_model_parallel_size > 1 and self.num_moe_experts is None: raise ValueError("num_moe_experts must be non None to use expert-parallel.") - if self.transformer_impl == "inference_optimized" and self.num_moe_experts is not None: + if ( + self.transformer_impl == "inference_optimized" + and self.num_moe_experts is not None + ): if self.expert_tensor_parallel_size > 1: raise ValueError( "Inference-optimized MoE layers does not support expert tensor parallelism." ) if self.moe_expert_capacity_factor is not None: - raise ValueError("Inference-optimized MoE layers only support dropless MoE ") + raise ValueError( + "Inference-optimized MoE layers only support dropless MoE " + ) if self.moe_router_padding_for_quantization: raise ValueError( "Inference-optimized MoE layers do not support padded " @@ -1459,7 +1496,8 @@ def __post_init__(self): ) if ( - self.inference_grouped_gemm_backend == InferenceGroupedGemmBackend.FLASHINFER + self.inference_grouped_gemm_backend + == InferenceGroupedGemmBackend.FLASHINFER and self.fp8 == "mxfp8" ): raise ValueError( @@ -1481,7 +1519,9 @@ def __post_init__(self): if self.num_moe_experts is not None and self.moe_ffn_hidden_size is None: self.moe_ffn_hidden_size = self.ffn_hidden_size - warnings.warn("moe_ffn_hidden_size is not set, using ffn_hidden_size instead.") + warnings.warn( + "moe_ffn_hidden_size is not set, using ffn_hidden_size instead." + ) if self.num_moe_experts is None and self.moe_ffn_hidden_size is not None: is_mixed_model = ( @@ -1529,9 +1569,13 @@ def __post_init__(self): if self.moe_enable_deepep: if self.moe_token_dispatcher_type != "flex": - raise ValueError("DeepEP backend is only supported with flex token dispatcher.") + raise ValueError( + "DeepEP backend is only supported with flex token dispatcher." + ) if self.moe_flex_dispatcher_backend == "hybridep": - raise ValueError("Only one backend is supported for flex token dispatcher.") + raise ValueError( + "Only one backend is supported for flex token dispatcher." + ) self.moe_flex_dispatcher_backend = "deepep" warnings.warn( "moe_enable_deepep is deprecated." @@ -1554,10 +1598,14 @@ def __post_init__(self): f"num_shared_experts * ffn_size_of_each_shared_expert, " f"but got {self.moe_shared_expert_intermediate_size}" ) - if self.moe_shared_expert_overlap and self.moe_token_dispatcher_type not in [ - "alltoall", - "flex", - ]: + if ( + self.moe_shared_expert_overlap + and self.moe_token_dispatcher_type + not in [ + "alltoall", + "flex", + ] + ): raise ValueError( f"moe_shared_expert_overlap only works with alltoall or flex token dispatcher." ) @@ -1615,7 +1663,8 @@ def __post_init__(self): ) if self.cpu_offloading and ( - self.cpu_offloading_num_layers < 0 or self.cpu_offloading_num_layers >= self.num_layers + self.cpu_offloading_num_layers < 0 + or self.cpu_offloading_num_layers >= self.num_layers ): raise ValueError( f"CPU offloading can be done only for layers less than {self.num_layers}" @@ -1649,7 +1698,10 @@ def __post_init__(self): 'recompute_method must be "block" or "uniform"' ) - if self.recompute_granularity != "selective" and self.recompute_num_layers is None: + if ( + self.recompute_granularity != "selective" + and self.recompute_num_layers is None + ): raise ValueError( f"When using recompute_granularity: {self.recompute_granularity} " "recompute_num_layers must be between " @@ -1657,7 +1709,8 @@ def __post_init__(self): f"{self.num_layers // self.pipeline_model_parallel_size}" ) elif ( - self.recompute_granularity == "selective" and self.recompute_num_layers is not None + self.recompute_granularity == "selective" + and self.recompute_num_layers is not None ): raise ValueError( f"When using recompute_granularity: {self.recompute_granularity} " @@ -1696,7 +1749,10 @@ def __post_init__(self): "moe_act in recompute_modules is only supported with moe_grouped_gemm." ) - if "mla_up_proj" in self.recompute_modules and not self.multi_latent_attention: + if ( + "mla_up_proj" in self.recompute_modules + and not self.multi_latent_attention + ): raise ValueError( "mla_up_proj in recompute_modules is only supported with " "multi_latent_attention." @@ -1729,8 +1785,11 @@ def __post_init__(self): ) if self.fp8: - if "moe_act" in self.recompute_modules or "layernorm" in self.recompute_modules: - if self.fp8_recipe == 'delayed': + if ( + "moe_act" in self.recompute_modules + or "layernorm" in self.recompute_modules + ): + if self.fp8_recipe == "delayed": raise ValueError( "Delayed scaling does not support moe_act and layernorm recompute " "for fp8." @@ -1756,9 +1815,9 @@ def __post_init__(self): self.recompute_modules.append("moe") if self.fine_grained_activation_offloading: - assert ( - not self.cpu_offloading - ), "fine_grained_activation_offloading cannot be enabled with cpu_offloading." + assert not self.cpu_offloading, ( + "fine_grained_activation_offloading cannot be enabled with cpu_offloading." + ) assert self.offload_modules is not None and len(self.offload_modules) > 0 allowed_modules = { "core_attn", @@ -1772,16 +1831,22 @@ def __post_init__(self): } invalid_modules = set(self.offload_modules) - allowed_modules assert not invalid_modules, ( - f'Invalid choices for offload_modules: {invalid_modules}. ' - f'Allowed modules are: {allowed_modules}' + f"Invalid choices for offload_modules: {invalid_modules}. " + f"Allowed modules are: {allowed_modules}" ) - if "attn_proj" in self.offload_modules and "core_attn" not in self.offload_modules: + if ( + "attn_proj" in self.offload_modules + and "core_attn" not in self.offload_modules + ): raise ValueError( "attn_proj cannot be set to offload_modules alone without core_attn " "because the input of attn_proj is the output of core_attn, " "which is needed in core_attn.backward()." ) - if self.recompute_granularity == "selective" and "moe" in self.recompute_modules: + if ( + self.recompute_granularity == "selective" + and "moe" in self.recompute_modules + ): offload_inside_moe = {"moe_act", "expert_fc1", "fused_group_mlp"} & set( self.offload_modules ) @@ -1792,20 +1857,25 @@ def __post_init__(self): f"Either remove 'moe' from --recompute-modules or remove " f"{offload_inside_moe} from --offload-modules." ) + assert self.min_offloaded_tensor_size >= 0, ( + "min_offloaded_tensor_size must be non-negative." + ) assert ( - self.min_offloaded_tensor_size >= 0 - ), "min_offloaded_tensor_size must be non-negative." - assert ( - self.activation_offload_fraction >= 0 and self.activation_offload_fraction <= 1 + self.activation_offload_fraction >= 0 + and self.activation_offload_fraction <= 1 ), "activation_offload_fraction must be in range [0, 1]." - assert ( - self.delta_offload_bytes_across_pp_ranks >= 0 - ), "delta_offload_bytes_across_pp_ranks must be non-negative." + assert self.delta_offload_bytes_across_pp_ranks >= 0, ( + "delta_offload_bytes_across_pp_ranks must be non-negative." + ) if "fused_group_mlp" in self.offload_modules: if not self.use_transformer_engine_op_fuser: - raise ValueError("fused_group_mlp requires use_transformer_engine_op_fuser.") - moe_partial_offload = {"expert_fc1", "moe_act"} & set(self.offload_modules) + raise ValueError( + "fused_group_mlp requires use_transformer_engine_op_fuser." + ) + moe_partial_offload = {"expert_fc1", "moe_act"} & set( + self.offload_modules + ) if moe_partial_offload: raise ValueError( "fused_group_mlp offloads the whole fused grouped MLP and cannot be " @@ -1813,7 +1883,9 @@ def __post_init__(self): ) if self.moe_paged_stash: if self.cpu_offloading: - raise ValueError("moe_paged_stash cannot be enabled with cpu_offloading.") + raise ValueError( + "moe_paged_stash cannot be enabled with cpu_offloading." + ) if self.moe_expert_rank_capacity_factor is None: raise ValueError( "moe_paged_stash requires moe_expert_rank_capacity_factor to be set; " @@ -1834,7 +1906,8 @@ def __post_init__(self): self.num_layers_in_first_pipeline_stage is not None or self.num_layers_in_last_pipeline_stage is not None ) and ( - self.account_for_embedding_in_pipeline_split or self.account_for_loss_in_pipeline_split + self.account_for_embedding_in_pipeline_split + or self.account_for_loss_in_pipeline_split ): raise ValueError( "num_layers_in_first_pipeline_stage and num_layers_in_last_pipeline_stage cannot be" @@ -1865,9 +1938,11 @@ def __post_init__(self): # Transfer pipeline_model_parallel_layout from str or list to # PipelineParallelLayerLayout if isinstance(self.pipeline_model_parallel_layout, str): - self.pipeline_model_parallel_layout = PipelineParallelLayerLayout.from_str( - layout=self.pipeline_model_parallel_layout, - pipeline_model_parallel_size=self.pipeline_model_parallel_size, + self.pipeline_model_parallel_layout = ( + PipelineParallelLayerLayout.from_str( + layout=self.pipeline_model_parallel_layout, + pipeline_model_parallel_size=self.pipeline_model_parallel_size, + ) ) elif isinstance(self.pipeline_model_parallel_layout, list): # Since list is not hashable, the initialization will not be cached. @@ -1891,8 +1966,10 @@ def __post_init__(self): self.virtual_pipeline_model_parallel_size = detected_vpp_size # Check whether the layout is valid. - self.mtp_standalone = self.pipeline_model_parallel_layout.validate_layer_layout( - num_layers=self.num_layers, mtp_num_layers=self.mtp_num_layers + self.mtp_standalone = ( + self.pipeline_model_parallel_layout.validate_layer_layout( + num_layers=self.num_layers, mtp_num_layers=self.mtp_num_layers + ) ) # Uneven PP @@ -1905,7 +1982,9 @@ def __post_init__(self): if self.num_layers_in_first_pipeline_stage is not None: if self.num_layers_in_first_pipeline_stage <= 0: - raise ValueError("num_layers_in_first_pipeline_stage must be larger than 0") + raise ValueError( + "num_layers_in_first_pipeline_stage must be larger than 0" + ) if self.virtual_pipeline_model_parallel_size is not None: if ( @@ -1924,7 +2003,9 @@ def __post_init__(self): if self.num_layers_in_last_pipeline_stage is not None: if self.num_layers_in_last_pipeline_stage <= 0: - raise ValueError("num_layers_in_last_pipeline_stage must be larger than 0") + raise ValueError( + "num_layers_in_last_pipeline_stage must be larger than 0" + ) if self.virtual_pipeline_model_parallel_size is not None: if ( @@ -1958,8 +2039,13 @@ def __post_init__(self): # If there are middle PP stages, check number of layers # on each middle PP rank is divisible by VPP size. - if pipeline_parallel_size and self.virtual_pipeline_model_parallel_size is not None: - num_layers_per_middle_pipeline_rank = num_layers // pipeline_parallel_size + if ( + pipeline_parallel_size + and self.virtual_pipeline_model_parallel_size is not None + ): + num_layers_per_middle_pipeline_rank = ( + num_layers // pipeline_parallel_size + ) if ( not num_layers_per_middle_pipeline_rank % self.virtual_pipeline_model_parallel_size @@ -1972,7 +2058,8 @@ def __post_init__(self): ) elif ( - self.account_for_embedding_in_pipeline_split or self.account_for_loss_in_pipeline_split + self.account_for_embedding_in_pipeline_split + or self.account_for_loss_in_pipeline_split ): if self.virtual_pipeline_model_parallel_size is None: num_layers = self.num_layers @@ -2005,9 +2092,12 @@ def __post_init__(self): f"{self.pipeline_model_parallel_size}" ) - num_layers_per_pipeline_rank = num_layers // self.pipeline_model_parallel_size + num_layers_per_pipeline_rank = ( + num_layers // self.pipeline_model_parallel_size + ) if ( - not num_layers_per_pipeline_rank % self.virtual_pipeline_model_parallel_size + not num_layers_per_pipeline_rank + % self.virtual_pipeline_model_parallel_size == 0 ): raise ValueError( @@ -2070,7 +2160,9 @@ def __post_init__(self): if self.activation_func_fp8_input_store: if self.activation_func != F.silu or not self.gated_linear_unit: - raise ValueError("Storing activation input in FP8 is supported only for SwiGLU.") + raise ValueError( + "Storing activation input in FP8 is supported only for SwiGLU." + ) if self.apply_rope_fusion: if self.multi_latent_attention: @@ -2091,33 +2183,46 @@ def __post_init__(self): fused_apply_rotary_pos_emb_thd, ) - if fused_apply_rotary_pos_emb is None and fused_apply_rotary_pos_emb_thd is None: + if ( + fused_apply_rotary_pos_emb is None + and fused_apply_rotary_pos_emb_thd is None + ): raise ValueError( "apply_rope_fusion is not available. Please install TE >= 1.4." ) if self.fused_single_qkv_rope: if self.attention_output_gate: - raise ValueError("fused_single_qkv_rope does not support gated attention for now.") + raise ValueError( + "fused_single_qkv_rope does not support gated attention for now." + ) if self.multi_latent_attention and self.rotary_interleaved: - raise ValueError("rotary_interleaved does not work with multi_latent_attention.") + raise ValueError( + "rotary_interleaved does not work with multi_latent_attention." + ) # MuP (Maximal Update Parameterization) configuration if self.use_mup: # Default base_hidden_size to hidden_size (base model case, width_mult=1.0) if self.mup_base_hidden_size is None: self.mup_base_hidden_size = self.hidden_size - assert self.mup_base_hidden_size > 0, "--mup-base-hidden-size must be positive." + assert self.mup_base_hidden_size > 0, ( + "--mup-base-hidden-size must be positive." + ) # Compute width multiplier self.mup_width_mult = self.hidden_size / self.mup_base_hidden_size # MuP attention scaling: 1/d_head instead of 1/sqrt(d_head). if self.softmax_scale is None: base_head_scale = ( - 1.0 if self.mup_base_head_dim is None else self.mup_base_head_dim**0.5 + 1.0 + if self.mup_base_head_dim is None + else self.mup_base_head_dim**0.5 + ) + self.softmax_scale = base_head_scale / ( + self.kv_channels**self.mup_attn_scale_power ) - self.softmax_scale = base_head_scale / (self.kv_channels**self.mup_attn_scale_power) # MuP output scaling: scale logits by 1/width_mult to keep outputs O(1). # Only auto-set if user hasn't explicitly configured it. @@ -2150,10 +2255,14 @@ def __post_init__(self): self.embedding_init_method_std = self.init_method_std if self.embedding_init_method is None: - if self.init_method is None or (self.embedding_init_method_std != self.init_method_std): + if self.init_method is None or ( + self.embedding_init_method_std != self.init_method_std + ): # In this case, we set both the init method and the embedding init method to # whatever std value requested (or defaulted) for the embedding_init_layer - self.embedding_init_method = init_method_normal(self.embedding_init_method_std) + self.embedding_init_method = init_method_normal( + self.embedding_init_method_std + ) else: # Replicate the current behavior where if you are not changing the std of the # embedding init differently and the init method is set, we fallback to the @@ -2187,13 +2296,17 @@ def __post_init__(self): ) if self.num_moe_experts is not None and self.add_bias_linear: - assert ( - self.expert_tensor_parallel_size == 1 - ), "Bias in Moe is only supported when ETP==1" + assert self.expert_tensor_parallel_size == 1, ( + "Bias in Moe is only supported when ETP==1" + ) - if self.moe_router_enable_expert_bias and self.moe_router_score_function not in ( - "sigmoid", - "sqrtsoftplus", + if ( + self.moe_router_enable_expert_bias + and self.moe_router_score_function + not in ( + "sigmoid", + "sqrtsoftplus", + ) ): raise ValueError( "Expert bias for aux-loss-free routing only supports 'sigmoid' and 'sqrtsoftplus' " @@ -2273,20 +2386,22 @@ def __post_init__(self): self.moe_router_num_groups = self.expert_model_parallel_size if self.enable_cuda_graph or self.external_cuda_graph: - assert ( - self.cuda_graph_impl == "none" - ), "Do not use enable_cuda_graph or external_cuda_graph with cuda_graph_impl." - assert ( - not self.enable_cuda_graph or not self.external_cuda_graph - ), "enable_cuda_graph and external_cuda_graph cannot be enabled at the same time." + assert self.cuda_graph_impl == "none", ( + "Do not use enable_cuda_graph or external_cuda_graph with cuda_graph_impl." + ) + assert not self.enable_cuda_graph or not self.external_cuda_graph, ( + "enable_cuda_graph and external_cuda_graph cannot be enabled at the same time." + ) if self.enable_cuda_graph: - warnings.warn('enable_cuda_graph is deprecated, use cuda_graph_impl=local instead.') + warnings.warn( + "enable_cuda_graph is deprecated, use cuda_graph_impl=local instead." + ) self.cuda_graph_impl = "local" if self.external_cuda_graph: warnings.warn( - 'external_cuda_graph is deprecated, ' - 'use cuda_graph_impl=transformer_engine instead.' + "external_cuda_graph is deprecated, " + "use cuda_graph_impl=transformer_engine instead." ) self.cuda_graph_impl = "transformer_engine" @@ -2315,8 +2430,8 @@ def _scope_to_str(s): self.cuda_graph_modules = _scope_to_str(scope) self.cuda_graph_scope = None - normalized_scopes, deprecated_scopes, used_full_scope = normalize_cuda_graph_modules( - self.cuda_graph_modules + normalized_scopes, deprecated_scopes, used_full_scope = ( + normalize_cuda_graph_modules(self.cuda_graph_modules) ) validate_deprecated_cuda_graph_modules_migration_inputs( deprecated_scopes, self.cuda_graph_impl, self.inference_cuda_graph_scope @@ -2350,7 +2465,9 @@ def _scope_to_str(s): self.cuda_graph_modules = normalized_scopes assert all( isinstance(scope, CudaGraphModule) for scope in self.cuda_graph_modules - ), f"cuda_graph_modules must be a list of CudaGraphModule, got {self.cuda_graph_modules}." + ), ( + f"cuda_graph_modules must be a list of CudaGraphModule, got {self.cuda_graph_modules}." + ) assert self.cuda_graph_impl in [ "none", @@ -2363,7 +2480,10 @@ def _scope_to_str(s): self.inference_cuda_graph_scope, self.cuda_graph_impl ) - assert self.inference_cuda_graph_scope in ALLOWED_INFERENCE_SCOPES[self.cuda_graph_impl], ( + assert ( + self.inference_cuda_graph_scope + in ALLOWED_INFERENCE_SCOPES[self.cuda_graph_impl] + ), ( "Invalid inference CUDA graph scope " f"{self.inference_cuda_graph_scope.name!r} for cuda_graph_impl=" f"{self.cuda_graph_impl!r}." @@ -2373,7 +2493,6 @@ def _scope_to_str(s): ), 'cuda_graph_modules must be empty when cuda_graph_impl="full_iteration".' if self.cuda_graph_impl != "none": - if self.cpu_offloading and self.cuda_graph_impl != "full_iteration": raise ValueError("CUDA graphs not supported with CPU offloading.") @@ -2388,51 +2507,60 @@ def _scope_to_str(s): ): if CudaGraphModule.moe_router not in self.cuda_graph_modules: self.cuda_graph_modules.append(CudaGraphModule.moe_router) - if CudaGraphModule.moe_preprocess not in self.cuda_graph_modules: - self.cuda_graph_modules.append(CudaGraphModule.moe_preprocess) + if ( + CudaGraphModule.moe_preprocess + not in self.cuda_graph_modules + ): + self.cuda_graph_modules.append( + CudaGraphModule.moe_preprocess + ) assert ( CudaGraphModule.moe not in self.cuda_graph_modules or CudaGraphModule.moe_router not in self.cuda_graph_modules - ), 'cuda_graph_modules must not contain both moe and moe_router.' + ), "cuda_graph_modules must not contain both moe and moe_router." if CudaGraphModule.moe_preprocess in self.cuda_graph_modules: - assert ( - CudaGraphModule.moe_router in self.cuda_graph_modules - ), 'moe_preprocess cuda graph is only supported with moe_router cuda graph.' + assert CudaGraphModule.moe_router in self.cuda_graph_modules, ( + "moe_preprocess cuda graph is only supported with moe_router cuda graph." + ) if self.num_moe_experts is None or self.num_moe_experts <= 1: assert ( CudaGraphModule.moe not in self.cuda_graph_modules and CudaGraphModule.moe_router not in self.cuda_graph_modules - ), 'moe cuda graph is only supported for MoE.' + ), "moe cuda graph is only supported for MoE." else: if self.moe_layer_freq == 1 or ( - isinstance(self.moe_layer_freq, list) and 0 not in self.moe_layer_freq + isinstance(self.moe_layer_freq, list) + and 0 not in self.moe_layer_freq ): assert CudaGraphModule.mlp not in self.cuda_graph_modules, ( - 'mlp cuda graph is only supported for dense layers, ' - 'but not found in the model.' + "mlp cuda graph is only supported for dense layers, " + "but not found in the model." ) if ( self.moe_expert_capacity_factor is None or not self.moe_pad_expert_input_to_capacity ): - assert ( - CudaGraphModule.moe not in self.cuda_graph_modules - ), 'moe cuda graph is only supported with drop-padding MoE.' - if self.moe_token_dispatcher_type == 'alltoall' and ( + assert CudaGraphModule.moe not in self.cuda_graph_modules, ( + "moe cuda graph is only supported with drop-padding MoE." + ) + if self.moe_token_dispatcher_type == "alltoall" and ( self.moe_expert_capacity_factor is not None or self.moe_router_padding_for_fp8 ): - assert CudaGraphModule.moe_preprocess not in self.cuda_graph_modules, ( - 'moe_preprocess cuda graph is not supported when there are ' - 'DtoH copies and synchronizations in the preprocess step.' + assert ( + CudaGraphModule.moe_preprocess + not in self.cuda_graph_modules + ), ( + "moe_preprocess cuda graph is not supported when there are " + "DtoH copies and synchronizations in the preprocess step." ) if self.recompute_granularity: if self.recompute_granularity != "selective": - assert ( - self.cuda_graph_impl == "full_iteration" - ), "full recompute is only supported with full iteration CUDA graph." + assert self.cuda_graph_impl == "full_iteration", ( + "full recompute is only supported with full iteration CUDA graph." + ) else: # The recompute module should be inside or outside of the graph scope. # Recompute module coverring graph scope is not allowed. @@ -2442,13 +2570,18 @@ def _scope_to_str(s): ): assert ( CudaGraphModule.moe_router not in self.cuda_graph_modules - ), "moe recompute is not supported with moe_router CUDA graph with: " + ), ( + "moe recompute is not supported with moe_router CUDA graph with: " + ) "--cuda-graph-impl transformer_engine." # Graphed recompute module doesn't accept random number. # full_cudagraph means either full_iteration impl or an empty per-layer scope # (which captures the whole layer). - if self.cuda_graph_impl == "full_iteration" or not self.cuda_graph_modules: + if ( + self.cuda_graph_impl == "full_iteration" + or not self.cuda_graph_modules + ): full_cudagraph = True else: full_cudagraph = False @@ -2473,7 +2606,9 @@ def _scope_to_str(s): and CudaGraphModule.moe not in self.cuda_graph_modules ) or "moe" not in self.recompute_modules - ), "hidden dropout is not supported with graphed MLP/MoE recomputation." + ), ( + "hidden dropout is not supported with graphed MLP/MoE recomputation." + ) if self.moe_input_jitter_eps is not None: assert ( not full_cudagraph @@ -2485,8 +2620,14 @@ def _scope_to_str(s): if self.fine_grained_activation_offloading: offload_modules = set(self.offload_modules or []) if self.cuda_graph_impl == "local": - local_supported_offload_modules = {"expert_fc1", "moe_act", "fused_group_mlp"} - unsupported_offload_modules = offload_modules - local_supported_offload_modules + local_supported_offload_modules = { + "expert_fc1", + "moe_act", + "fused_group_mlp", + } + unsupported_offload_modules = ( + offload_modules - local_supported_offload_modules + ) assert not unsupported_offload_modules, ( "fine-grained activation offloading with cuda_graph_impl='local' " "only supports offload_modules 'expert_fc1', 'moe_act', and " @@ -2513,9 +2654,9 @@ def _scope_to_str(s): "are supported only for expert_fc1, moe_act, or fused_group_mlp " "offload when the full MoE module is not captured." ) - assert ( - CudaGraphModule.moe not in self.cuda_graph_modules - ), "Token-drop MoE is temporarily not supported with activation offloading." + assert CudaGraphModule.moe not in self.cuda_graph_modules, ( + "Token-drop MoE is temporarily not supported with activation offloading." + ) assert self.cuda_graph_warmup_steps > 0, ( "cuda_graph_warmup_steps must be greater than 0 when enabling " "fine-grained activation offloading." @@ -2553,51 +2694,55 @@ def _scope_to_str(s): or fused_sort_chunks_by_index_with_probs is None or fused_unpermute is None ): - raise ValueError("fused permutation is not available. Please install TE >= 2.1.0.") + raise ValueError( + "fused permutation is not available. Please install TE >= 2.1.0." + ) if self.overlap_moe_expert_parallel_comm: # TODO: remove this after we fix the hang issue with torch version < 2.6.0 - assert is_torch_min_version( - "2.6.0" - ), "A2A Overlap encounters hang issue with torch version < 2.6.0" + assert is_torch_min_version("2.6.0"), ( + "A2A Overlap encounters hang issue with torch version < 2.6.0" + ) if self.pipeline_model_parallel_size > 1: assert self.virtual_pipeline_model_parallel_size is not None, ( "If enabling EP A2A overlap, virtual_pipeline_model_parallel_size " "must be specified when pipeline_model_parallel_size > 1" ) # Expert model parallelism requirements - assert ( - self.expert_model_parallel_size > 1 - ), 'overlap_moe_expert_parallel_comm is only supported with expert model parallelism' + assert self.expert_model_parallel_size > 1, ( + "overlap_moe_expert_parallel_comm is only supported with expert model parallelism" + ) assert self.moe_token_dispatcher_type in [ - 'alltoall', - 'flex', - ], 'overlap_moe_expert_parallel_comm is supported with alltoall/flex token dispatcher' + "alltoall", + "flex", + ], ( + "overlap_moe_expert_parallel_comm is supported with alltoall/flex token dispatcher" + ) - assert ( - self.recompute_granularity != 'full' - ), 'disable full recomputation when enabling overlap_moe_expert_parallel_comm' - assert ( - self.recompute_method is None - ), 'disable recomputation method when enabling overlap_moe_expert_parallel_comm' - assert ( - self.recompute_num_layers is None - ), 'recompute_num_layers must be None when enabling overlap_moe_expert_parallel_comm' - assert ( - "moe" not in self.recompute_modules - ), 'disable moe in recompute_modules when enabling overlap_moe_expert_parallel_comm' + assert self.recompute_granularity != "full", ( + "disable full recomputation when enabling overlap_moe_expert_parallel_comm" + ) + assert self.recompute_method is None, ( + "disable recomputation method when enabling overlap_moe_expert_parallel_comm" + ) + assert self.recompute_num_layers is None, ( + "recompute_num_layers must be None when enabling overlap_moe_expert_parallel_comm" + ) + assert "moe" not in self.recompute_modules, ( + "disable moe in recompute_modules when enabling overlap_moe_expert_parallel_comm" + ) # Check if bf16 or fp16 is used - assert ( - self.bf16 or self.fp16 - ), 'overlap_moe_expert_parallel_comm is only supported with bf16 or fp16 model' + assert self.bf16 or self.fp16, ( + "overlap_moe_expert_parallel_comm is only supported with bf16 or fp16 model" + ) - assert ( - not self.moe_shared_expert_overlap - ), 'disable moe_shared_expert_overlap when enabling overlap_moe_expert_parallel_comm' - assert ( - self.mtp_num_layers is None or self.mtp_num_layers == 1 - ), 'MTP layernum only supports 1 when enabling overlap_moe_expert_parallel_comm.' + assert not self.moe_shared_expert_overlap, ( + "disable moe_shared_expert_overlap when enabling overlap_moe_expert_parallel_comm" + ) + assert self.mtp_num_layers is None or self.mtp_num_layers == 1, ( + "MTP layernum only supports 1 when enabling overlap_moe_expert_parallel_comm." + ) if self.cuda_graph_impl != "none": if self.cuda_graph_impl == "transformer_engine": @@ -2605,38 +2750,38 @@ def _scope_to_str(s): CudaGraphModule.moe not in self.cuda_graph_modules and CudaGraphModule.mlp not in self.cuda_graph_modules ), ( - 'CUDA graph scope on moe and mlp is not ' - 'supported with overlap_moe_expert_parallel_comm' + "CUDA graph scope on moe and mlp is not " + "supported with overlap_moe_expert_parallel_comm" ) # Check delay_wgrad_compute compatibility if self.delay_wgrad_compute: - assert ( - self.overlap_moe_expert_parallel_comm - ), 'overlap_moe_expert_parallel_comm must be enabled when enabling delay_wgrad_compute' + assert self.overlap_moe_expert_parallel_comm, ( + "overlap_moe_expert_parallel_comm must be enabled when enabling delay_wgrad_compute" + ) if self.cuda_graph_impl == "transformer_engine": assert is_te_min_version("2.10.0"), ( - 'TE version >= 2.10.0 is required for delay_wgrad_compute with ' - 'partial cuda graph' + "TE version >= 2.10.0 is required for delay_wgrad_compute with " + "partial cuda graph" ) if self.overlap_dispatch_backward_with_experts_wgrad: assert not self.overlap_moe_expert_parallel_comm, ( - 'overlap_moe_expert_parallel_comm must be disabled when enabling ' - 'overlap_dispatch_backward_with_experts_wgrad.' + "overlap_moe_expert_parallel_comm must be disabled when enabling " + "overlap_dispatch_backward_with_experts_wgrad." + ) + assert is_te_min_version("2.3.0"), ( + "TE version >= 2.3.0 is required for overlap_dispatch_backward_with_experts_wgrad" ) - assert is_te_min_version( - "2.3.0" - ), 'TE version >= 2.3.0 is required for overlap_dispatch_backward_with_experts_wgrad' assert not self.delay_wgrad_compute, ( - 'delay_wgrad_compute and overlap_dispatch_backward_with_experts_wgrad ' - 'are mutually exclusive; use only one' + "delay_wgrad_compute and overlap_dispatch_backward_with_experts_wgrad " + "are mutually exclusive; use only one" ) if self.ep_overlap_early_attn_memory_release: assert self.overlap_moe_expert_parallel_comm, ( - 'overlap_moe_expert_parallel_comm must be enabled when enabling ' - 'ep_overlap_early_attn_memory_release' + "overlap_moe_expert_parallel_comm must be enabled when enabling " + "ep_overlap_early_attn_memory_release" ) if self.context_parallel_size > 1 and self.cp_comm_type is not None: @@ -2646,14 +2791,14 @@ def _scope_to_str(s): f"the total number of transformer layers ({self.num_layers})!" ) else: - assert isinstance( - self.cp_comm_type, str - ), "Unsupported communication type for context parallelism!" + assert isinstance(self.cp_comm_type, str), ( + "Unsupported communication type for context parallelism!" + ) - assert ( - self.pipeline_model_parallel_size > 0 - ), f"Pipeline model parallel size must be larger than 0 \ + assert self.pipeline_model_parallel_size > 0, ( + f"Pipeline model parallel size must be larger than 0 \ when enable --standalone-embedding-stage and --standalone-loss-stage" + ) if ( self.num_moe_experts is not None @@ -2669,10 +2814,14 @@ def _scope_to_str(s): raise ImportError( "packaging is not installed. Please install it with `pip install packaging`." ) - assert is_torch_min_version("2.7.0a0"), "Must have at least torch version 2.7 or higher" + assert is_torch_min_version("2.7.0a0"), ( + "Must have at least torch version 2.7 or higher" + ) assert is_te_min_version("2.3.0") or get_te_version() == PkgVersion( "2.3.0.dev0+39c0e70" - ), "Must have at least TE version 2.3 or higher to use symmetric memory all reduce" + ), ( + "Must have at least TE version 2.3 or higher to use symmetric memory all reduce" + ) if self.no_rope_freq: assert not self.flash_decode, "flash_decode cannot be used with no_rope." @@ -2699,7 +2848,9 @@ def _scope_to_str(s): assert not self.use_kitchen if self.experimental_attention_variant == "dsa": - assert not self.apply_rope_fusion, "RoPE fusion is not supported for DSAttention" + assert not self.apply_rope_fusion, ( + "RoPE fusion is not supported for DSAttention" + ) if self.context_parallel_size > 1: cp_comm_types = ( self.cp_comm_type @@ -2720,9 +2871,9 @@ def _scope_to_str(s): "inference_fuse_tp_communication is only supported " "for inference_optimized transformer implementation." ) - assert ( - self.num_moe_experts is None - ), "--inference-fuse-tp-communication is not supported for MoE models." + assert self.num_moe_experts is None, ( + "--inference-fuse-tp-communication is not supported for MoE models." + ) if self.inference_disable_triton_nvls_kernels: assert self.transformer_impl == "inference_optimized", ( @@ -2731,9 +2882,9 @@ def _scope_to_str(s): ) if self.batch_invariant_mode: - assert ( - self.attention_backend == AttnBackend.flash - ), "Batch invariant mode only supports FlashAttention" + assert self.attention_backend == AttnBackend.flash, ( + "Batch invariant mode only supports FlashAttention" + ) @dataclass @@ -2804,13 +2955,17 @@ class MLATransformerConfig(TransformerConfig): def __post_init__(self): super().__post_init__() - if self.multi_latent_attention and self.apply_rope_fusion and self.rope_type != "yarn": + if ( + self.multi_latent_attention + and self.apply_rope_fusion + and self.rope_type != "yarn" + ): raise ValueError("apply_rope_fusion for MLA only works with YARN RoPE.") if self.attention_output_gate: raise NotImplementedError("Output gate is not supported for MLA yet.") if self.cache_mla_latents: - assert ( - self.apply_rope_fusion is False - ), "Rope Fusion is not compatible with caching latents" + assert self.apply_rope_fusion is False, ( + "Rope Fusion is not compatible with caching latents" + ) diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py index 5fca6010a30..d05ab24df38 100644 --- a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -49,7 +49,9 @@ def _make_backend(fuse_layernorm=True): backend.linear.return_value = _FakeLinear backend.column_parallel_linear.return_value = _FakeColumnParallelLinear backend.row_parallel_linear.return_value = _FakeRowParallelLinear - backend.column_parallel_layer_norm_linear.return_value = _FakeLayerNormColumnParallelLinear + backend.column_parallel_layer_norm_linear.return_value = ( + _FakeLayerNormColumnParallelLinear + ) backend.fuse_layernorm_and_linear.return_value = fuse_layernorm backend.core_attention.return_value = _FakeCoreAttention @@ -107,7 +109,12 @@ def _fn(variant): @pytest.mark.parametrize( "variant, expected", - [("gated_delta_net", True), ("dsa", False), (None, False), ("some_unknown_variant", False)], + [ + ("gated_delta_net", True), + ("dsa", False), + (None, False), + ("some_unknown_variant", False), + ], ) def test_variants(self, variant, expected): """Validate linear-attention variant classification across supported and unsupported names.""" @@ -199,7 +206,9 @@ def test_list_freq_wrong_length_raises(self): def test_none_for_non_linear_variant(self): """Verify non-linear variants default to all-standard attention when freq is None.""" cfg = _make_config( - num_layers=4, linear_attention_freq=None, experimental_attention_variant="dsa" + num_layers=4, + linear_attention_freq=None, + experimental_attention_variant="dsa", ) assert self._fn(cfg) == [0, 0, 0, 0] @@ -295,7 +304,9 @@ def _call(self, cfg=None, backend=None): ) if cfg is None: - cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True) + cfg = _make_config( + multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True + ) if backend is None: backend = _make_backend() return get_dsa_module_spec_for_backend(cfg, backend=backend) @@ -333,7 +344,9 @@ def test_returns_absorbed_mla_self_attention_spec(self): def test_core_attention_is_dsa(self): """Verify MLA core_attention is wrapped with DSAttention.""" - from megatron.core.transformer.experimental_attention_variant.dsa import DSAttention + from megatron.core.transformer.experimental_attention_variant.dsa import ( + DSAttention, + ) spec = self._call() core = spec.submodules.core_attention @@ -341,7 +354,9 @@ def test_core_attention_is_dsa(self): def test_dsa_indexer_structure(self): """Verify DSA indexer wiring uses expected backend linear/norm modules.""" - from megatron.core.transformer.experimental_attention_variant.dsa import DSAIndexer + from megatron.core.transformer.experimental_attention_variant.dsa import ( + DSAIndexer, + ) spec = self._call() indexer = spec.submodules.core_attention.submodules.indexer @@ -371,8 +386,8 @@ def test_qk_layernorm_enabled(self, normalization): backend.layer_norm.assert_any_call(rms_norm=expected_rms, for_qk=True) def test_accuracy_compatible_qk_rmsnorm(self): - """Verify DSA q/kv norms can use the explicit fp32 RMSNorm path.""" - from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm + """Verify DSA q/kv norms can use the native Torch RMSNorm builder.""" + from megatron.core.transformer.torch_norm import WrappedTorchNorm backend = _make_backend() cfg = _make_config( @@ -384,28 +399,34 @@ def test_accuracy_compatible_qk_rmsnorm(self): ) spec = self._call(cfg=cfg, backend=backend) - assert spec.submodules.q_layernorm is AccuracyCompatibleRMSNorm - assert spec.submodules.kv_layernorm is AccuracyCompatibleRMSNorm + assert spec.submodules.q_layernorm is WrappedTorchNorm + assert spec.submodules.kv_layernorm is WrappedTorchNorm def test_qk_layernorm_disabled(self): """Verify q/kv layernorm becomes IdentityOp, skipping backend.layer_norm for qk.""" backend = _make_backend() - cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=False) + cfg = _make_config( + multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=False + ) spec = self._call(cfg=cfg, backend=backend) assert spec.submodules.q_layernorm is IdentityOp assert spec.submodules.kv_layernorm is IdentityOp # backend.layer_norm is still called for the indexer k_norm (for_qk=True at line 94), # but NOT for the outer qk_norm (line 105-107 takes the else branch). # Exactly one for_qk=True call should exist (from the indexer, not from qk_norm). - qk_calls = [c for c in backend.layer_norm.call_args_list if c.kwargs.get("for_qk")] - assert ( - len(qk_calls) == 1 - ), f"Expected 1 for_qk=True call (indexer only), got {len(qk_calls)}" + qk_calls = [ + c for c in backend.layer_norm.call_args_list if c.kwargs.get("for_qk") + ] + assert len(qk_calls) == 1, ( + f"Expected 1 for_qk=True call (indexer only), got {len(qk_calls)}" + ) def test_linear_projections(self): """Verify Q/KV projection slots and backend.column_parallel_linear call count.""" backend = _make_backend() - cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True) + cfg = _make_config( + multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True + ) spec = self._call(cfg=cfg, backend=backend) subs = spec.submodules assert subs.linear_q_proj == _FakeColumnParallelLinear @@ -437,14 +458,18 @@ class TestGetExperimentalAttentionVariantModuleSpec: def test_dispatches_to_variant_handler(self, variant, target_fn): """Verify dispatcher routes each variant name to its corresponding builder function.""" backend = _make_backend() - cfg = _make_config(experimental_attention_variant=variant, normalization="RMSNorm") + cfg = _make_config( + experimental_attention_variant=variant, normalization="RMSNorm" + ) with patch(f"{self.MODULE}.{target_fn}") as mock_fn: mock_fn.return_value = ModuleSpec(module=MagicMock) from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( get_experimental_attention_variant_module_spec, ) - result = get_experimental_attention_variant_module_spec(cfg, backend=backend) + result = get_experimental_attention_variant_module_spec( + cfg, backend=backend + ) mock_fn.assert_called_once_with(config=cfg, backend=backend) assert result is mock_fn.return_value @@ -469,12 +494,15 @@ class TestGetTransformerLayerWithExperimentalAttentionVariantSpec: def _make_attention_spec(self, fuse_input_layernorm=True): """Construct a mock attention spec with configurable fuse metadata.""" - return ModuleSpec(module=MagicMock, metainfo={"fuse_input_layernorm": fuse_input_layernorm}) + return ModuleSpec( + module=MagicMock, metainfo={"fuse_input_layernorm": fuse_input_layernorm} + ) def _make_mlp_spec(self, fuse_pre_mlp_layernorm=True): """Construct a mock MLP spec with configurable fuse metadata.""" return ModuleSpec( - module=MagicMock, metainfo={"fuse_pre_mlp_layernorm": fuse_pre_mlp_layernorm} + module=MagicMock, + metainfo={"fuse_pre_mlp_layernorm": fuse_pre_mlp_layernorm}, ) def test_all_experimental_no_moe(self): @@ -498,7 +526,10 @@ def test_all_experimental_no_moe(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=attn_spec, ), - patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(mlp_spec, True)), + patch( + f"{self.MODULE}._get_dense_mlp_module_spec", + return_value=(mlp_spec, True), + ), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -534,8 +565,14 @@ def test_hybrid_attention_pattern(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=exp_attn_spec, ), - patch(f"{self.MODULE}._get_self_attention_module_spec", return_value=std_attn_spec), - patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(mlp_spec, True)), + patch( + f"{self.MODULE}._get_self_attention_module_spec", + return_value=std_attn_spec, + ), + patch( + f"{self.MODULE}._get_dense_mlp_module_spec", + return_value=(mlp_spec, True), + ), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -571,8 +608,13 @@ def test_hybrid_moe_pattern(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=attn_spec, ), - patch(f"{self.MODULE}._get_moe_module_spec", return_value=(moe_spec, False)), - patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(dense_spec, True)), + patch( + f"{self.MODULE}._get_moe_module_spec", return_value=(moe_spec, False) + ), + patch( + f"{self.MODULE}._get_dense_mlp_module_spec", + return_value=(dense_spec, True), + ), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -636,7 +678,8 @@ def test_get_transformer_block_with_experimental_attention_variant_spec( ) backend = _make_backend() fake_layer_specs = [ - ModuleSpec(module=TransformerLayer, submodules=MagicMock()) for _ in range(num_layers) + ModuleSpec(module=TransformerLayer, submodules=MagicMock()) + for _ in range(num_layers) ] with ( @@ -657,17 +700,25 @@ def test_get_transformer_block_with_experimental_attention_variant_spec( # Without explicit layout, slicing comes from offset + num_layers_to_build. with ( patch( - f"{self.MODULE}.get_transformer_layer_offset", return_value=offset + f"{self.MODULE}.get_transformer_layer_offset", + return_value=offset, ) as mock_offset, patch( - f"{self.MODULE}.get_num_layers_to_build", return_value=num_layers_to_build + f"{self.MODULE}.get_num_layers_to_build", + return_value=num_layers_to_build, ) as mock_num_layers, ): - result = get_transformer_block_with_experimental_attention_variant_spec( - cfg, vp_stage=vp_stage, pp_rank=pp_rank + result = ( + get_transformer_block_with_experimental_attention_variant_spec( + cfg, vp_stage=vp_stage, pp_rank=pp_rank + ) ) - mock_offset.assert_called_once_with(cfg, vp_stage=vp_stage, pp_rank=pp_rank) - mock_num_layers.assert_called_once_with(cfg, vp_stage=vp_stage, pp_rank=pp_rank) + mock_offset.assert_called_once_with( + cfg, vp_stage=vp_stage, pp_rank=pp_rank + ) + mock_num_layers.assert_called_once_with( + cfg, vp_stage=vp_stage, pp_rank=pp_rank + ) assert isinstance(result, TransformerBlockSubmodules) assert result.layer_specs == [fake_layer_specs[i] for i in expected_ids] diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index 9014d0896fb..e4c422a5108 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -17,7 +17,9 @@ from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec from megatron.core.models.hybrid.hybrid_model import HybridModel -from megatron.core.num_microbatches_calculator import destroy_num_microbatches_calculator +from megatron.core.num_microbatches_calculator import ( + destroy_num_microbatches_calculator, +) from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.parallel_state import get_context_parallel_group from megatron.core.process_groups_config import ProcessGroupCollection @@ -29,11 +31,22 @@ process_mtp_loss, roll_tensor, ) -from megatron.core.transformer.torch_norm import AccuracyCompatibleRMSNorm +from megatron.core.transformer.torch_norm import WrappedTorchNorm from megatron.core.transformer.transformer_config import TransformerConfig -from megatron.core.utils import get_batch_on_this_cp_rank, is_te_min_version, unwrap_model -from megatron.training.argument_utils import gpt_config_from_args, hybrid_config_from_args -from megatron.training.arguments import core_transformer_config_from_args, parse_args, validate_args +from megatron.core.utils import ( + get_batch_on_this_cp_rank, + is_te_min_version, + unwrap_model, +) +from megatron.training.argument_utils import ( + gpt_config_from_args, + hybrid_config_from_args, +) +from megatron.training.arguments import ( + core_transformer_config_from_args, + parse_args, + validate_args, +) from megatron.training.checkpointing import load_checkpoint, save_checkpoint from megatron.training.global_vars import ( destroy_global_vars, @@ -46,7 +59,9 @@ from tests.unit_tests.test_utilities import Utils if HAVE_TE: - from megatron.core.extensions.transformer_engine import TEColumnParallelGroupedLinear + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelGroupedLinear, + ) else: TEColumnParallelGroupedLinear = None @@ -55,7 +70,7 @@ class TestMultiTokenPredictionLayer: def setup_method(self, method): - os.environ['CUDA_DEVICE_MAX_CONNECTIONS'] = '1' + os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "1" def teardown_method(self, method): Utils.destroy_model_parallel() @@ -63,7 +78,9 @@ def teardown_method(self, method): destroy_num_microbatches_calculator() def _create_config_and_mtp_block_spec(self, tp, cp, use_te=False): - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) config = TransformerConfig( mtp_num_layers=2, num_layers=4, @@ -84,8 +101,10 @@ def _create_config_and_mtp_block_spec(self, tp, cp, use_te=False): return config, mtp_block_spec def test_accuracy_compatible_norms_override_te_mtp_norms(self): - """Accuracy mode must route all MTP-owned norms through the explicit RMSNorm.""" - Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + """Accuracy mode routes all MTP-owned norms through native Torch RMSNorm.""" + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, context_parallel_size=1 + ) config = TransformerConfig( mtp_num_layers=1, num_layers=1, @@ -103,25 +122,30 @@ def test_accuracy_compatible_norms_override_te_mtp_norms(self): ) mtp_layer_spec = mtp_block_spec.layer_specs[0] - assert mtp_layer_spec.submodules.enorm is AccuracyCompatibleRMSNorm - assert mtp_layer_spec.submodules.hnorm is AccuracyCompatibleRMSNorm - assert mtp_layer_spec.submodules.layer_norm is AccuracyCompatibleRMSNorm + assert mtp_layer_spec.submodules.enorm is WrappedTorchNorm + assert mtp_layer_spec.submodules.hnorm is WrappedTorchNorm + assert mtp_layer_spec.submodules.layer_norm is WrappedTorchNorm final_norm = mtp_layer_spec.submodules.layer_norm( config=config, hidden_size=config.hidden_size, eps=config.layernorm_epsilon ) - assert isinstance(final_norm, AccuracyCompatibleRMSNorm) + assert isinstance(final_norm, torch.nn.RMSNorm) def test_mtp_detach_heads_config(self): """Test that mtp_detach_heads config defaults to False.""" config = TransformerConfig( - num_layers=4, hidden_size=64, num_attention_heads=8, use_cpu_initialization=True + num_layers=4, + hidden_size=64, + num_attention_heads=8, + use_cpu_initialization=True, ) assert config.mtp_detach_heads is False def test_constructor_with_detach_heads(self): """Test construction of MTP module with mtp_detach_heads=True.""" torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, context_parallel_size=1 + ) config = TransformerConfig( mtp_num_layers=2, num_layers=4, @@ -141,11 +165,11 @@ def test_constructor_with_detach_heads(self): # Verify all parameters are tagged for separate MTP grad-norm handling. for name, param in mtp.named_parameters(): - assert ( - getattr(param, 'grad_norm_group', None) == 'mtp' - ), f"Parameter {name} missing grad_norm_group attribute" + assert getattr(param, "grad_norm_group", None) == "mtp", ( + f"Parameter {name} missing grad_norm_group attribute" + ) - @pytest.mark.parametrize(('tp'), [(1), (2), (4)]) + @pytest.mark.parametrize(("tp"), [(1), (2), (4)]) def test_constructor_local(self, tp): """Test basic construction of MTP module.""" @@ -171,12 +195,16 @@ def test_constructor_local(self, tp): assert num_weights == 15216 * config.mtp_num_layers @pytest.mark.skipif(not HAVE_TE, reason="transformer_engine not available") - @pytest.mark.parametrize(('tp', 'cp'), [(1, 1), (1, 2), (2, 1), (2, 2)]) + @pytest.mark.parametrize(("tp", "cp"), [(1, 1), (1, 2), (2, 1), (2, 2)]) def test_constructor_ues_te(self, tp, cp): """Test basic construction of MTP module.""" torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) - config, mtp_block_spec = self._create_config_and_mtp_block_spec(tp, cp, use_te=True) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) + config, mtp_block_spec = self._create_config_and_mtp_block_spec( + tp, cp, use_te=True + ) mtp = MultiTokenPredictionBlock(config=config, spec=mtp_block_spec) assert isinstance(mtp, MultiTokenPredictionBlock) @@ -205,15 +233,22 @@ def test_get_embeddings_rolls_padding_mask(self): seq_len = 6 batch_size = 2 - input_ids = torch.tensor([[1, 2, 3, 4, 0, 0], [5, 6, 7, 0, 0, 0]], dtype=torch.int64) + input_ids = torch.tensor( + [[1, 2, 3, 4, 0, 0], [5, 6, 7, 0, 0, 0]], dtype=torch.int64 + ) position_ids = torch.arange(seq_len, dtype=torch.int64).repeat(batch_size, 1) padding_mask = torch.tensor( - [[True, True, True, True, False, False], [True, True, True, False, False, False]] + [ + [True, True, True, True, False, False], + [True, True, True, False, False, False], + ] ) hidden_states = torch.randn(seq_len, batch_size, config.hidden_size) def fake_embedding(input_ids, position_ids): - return torch.zeros(seq_len, batch_size, config.hidden_size, dtype=hidden_states.dtype) + return torch.zeros( + seq_len, batch_size, config.hidden_size, dtype=hidden_states.dtype + ) rolled_input_ids, rolled_position_ids, rolled_padding_mask, _, _ = ( mtp_layer._get_embeddings( @@ -245,13 +280,17 @@ def test_forward_propagates_rolled_padding_mask(self, monkeypatch): batch_size = 2 input_ids = torch.tensor([[1, 2, 3, 0], [4, 5, 0, 0]], dtype=torch.int64) position_ids = torch.arange(seq_len, dtype=torch.int64).repeat(batch_size, 1) - padding_mask = torch.tensor([[True, True, True, False], [True, True, False, False]]) + padding_mask = torch.tensor( + [[True, True, True, False], [True, True, False, False]] + ) hidden_states = torch.randn(seq_len, batch_size, config.hidden_size) attention_mask = torch.ones((batch_size, 1, seq_len, seq_len), dtype=torch.bool) seen = {} def fake_embedding(input_ids, position_ids): - return torch.zeros(seq_len, batch_size, config.hidden_size, dtype=hidden_states.dtype) + return torch.zeros( + seq_len, batch_size, config.hidden_size, dtype=hidden_states.dtype + ) def fake_proj_and_transformer_layer( self, @@ -307,7 +346,9 @@ def test_get_embeddings_detaches_decoder_input(self): position_ids = torch.arange(seq_len, dtype=torch.int64).repeat(batch_size, 1) # hidden_states arrives without requires_grad (it is detached upstream by the block). hidden_states = torch.randn(seq_len, batch_size, config.hidden_size) - emb_weight = torch.nn.Parameter(torch.randn(seq_len, batch_size, config.hidden_size)) + emb_weight = torch.nn.Parameter( + torch.randn(seq_len, batch_size, config.hidden_size) + ) def fake_embedding(input_ids, position_ids): return emb_weight.clone() @@ -352,8 +393,12 @@ def forward(self, hidden_states, **kwargs): seq_len = 4 batch_size = 2 input_ids = torch.tensor([[1, 2, 3, 0], [4, 5, 0, 0]], dtype=torch.int64).cuda() - position_ids = torch.arange(seq_len, dtype=torch.int64).repeat(batch_size, 1).cuda() - attention_mask = torch.ones((batch_size, 1, seq_len, seq_len), dtype=torch.bool).cuda() + position_ids = ( + torch.arange(seq_len, dtype=torch.int64).repeat(batch_size, 1).cuda() + ) + attention_mask = torch.ones( + (batch_size, 1, seq_len, seq_len), dtype=torch.bool + ).cuda() hidden_states = torch.randn( seq_len, batch_size, config.hidden_size, device="cuda", requires_grad=True ) @@ -388,7 +433,9 @@ def fake_embedding(input_ids, position_ids): # The returned block output still includes the original hidden-state # chunk, so autograd may allocate a zero grad for it through cat(). if hidden_states.grad is not None: - torch.testing.assert_close(hidden_states.grad, torch.zeros_like(hidden_states)) + torch.testing.assert_close( + hidden_states.grad, torch.zeros_like(hidden_states) + ) assert emb_weight.grad is None else: assert hidden_states.grad is not None @@ -399,7 +446,9 @@ def test_process_mtp_loss_detaches_output_weight(self, detach_heads): """process_mtp_loss must detach the output-head weight when mtp_detach_heads=True so the MTP loss does not update the (shared) output projection weight.""" torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, context_parallel_size=1 + ) config = TransformerConfig( mtp_num_layers=2, num_layers=4, @@ -455,7 +504,7 @@ class TestMultiTokenPrediction: def setup_method(self, method): self.seq_length = 32 self.micro_batch_size = 2 - os.environ['CUDA_DEVICE_MAX_CONNECTIONS'] = '1' + os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "1" def teardown_method(self, method): Utils.destroy_model_parallel() @@ -502,7 +551,7 @@ def create_test_args( destroy_global_vars() destroy_num_microbatches_calculator() - sys.argv = ['test_multi_token_predictioin.py'] + sys.argv = ["test_multi_token_predictioin.py"] args = parse_args() args.num_layers = 2 args.mtp_num_layers = 2 @@ -517,10 +566,10 @@ def create_test_args( args.tensor_model_parallel_size = tp args.sequence_parallel = True if tp > 1 else False args.context_parallel_size = cp - args.position_embedding_type = 'rope' + args.position_embedding_type = "rope" args.num_experts = 8 args.train_iters = 1 - args.ckpt_format = 'torch_dist' + args.ckpt_format = "torch_dist" args.moe_router_topk = 2 args.moe_router_pre_softmax = False args.lr = 3e-5 @@ -536,10 +585,10 @@ def create_test_args( args.moe_grouped_gemm = False args.bf16 = True if fp8 is not None: - args.fp8 = 'e4m3' + args.fp8 = "e4m3" if full_recompute: - args.recompute_granularity = 'full' - args.recompute_method = 'uniform' + args.recompute_granularity = "full" + args.recompute_method = "uniform" args.recompute_num_layers = 1 else: args.recompute_granularity = None @@ -552,19 +601,26 @@ def create_test_args( def get_batch(self, seq_length, micro_batch_size): data = list(range(seq_length)) - input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - labels = 1 + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - position_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + input_ids = ( + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + ) + labels = ( + 1 + + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + ) + position_ids = ( + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + ) attention_mask = torch.ones( (micro_batch_size, 1, seq_length, seq_length), dtype=bool ).cuda() loss_mask = torch.ones(seq_length).repeat((micro_batch_size, 1)).cuda() batch = { - 'tokens': input_ids, - 'labels': labels, - 'loss_mask': loss_mask, - 'attention_mask': attention_mask, - 'position_ids': position_ids, + "tokens": input_ids, + "labels": labels, + "loss_mask": loss_mask, + "attention_mask": attention_mask, + "position_ids": position_ids, } return batch @@ -595,7 +651,9 @@ def get_packed_batch(self, seq_lengths, micro_batch_size): # Convert to tensors with shape [batch, total_seq_length] input_ids = torch.tensor(input_ids_list, dtype=torch.int64).unsqueeze(0).cuda() labels = torch.tensor(labels_list, dtype=torch.int64).unsqueeze(0).cuda() - position_ids = torch.tensor(position_ids_list, dtype=torch.int64).unsqueeze(0).cuda() + position_ids = ( + torch.tensor(position_ids_list, dtype=torch.int64).unsqueeze(0).cuda() + ) # Create attention mask for packed sequences (all ones for simplicity) attention_mask = torch.ones( @@ -607,7 +665,8 @@ def get_packed_batch(self, seq_lengths, micro_batch_size): # Create cumulative sequence lengths for PackedSeqParams cu_seqlens = torch.tensor( - [0] + [sum(seq_lengths[: i + 1]) for i in range(len(seq_lengths))], dtype=torch.int32 + [0] + [sum(seq_lengths[: i + 1]) for i in range(len(seq_lengths))], + dtype=torch.int32, ).cuda() packed_seq_params = PackedSeqParams( @@ -615,16 +674,16 @@ def get_packed_batch(self, seq_lengths, micro_batch_size): cu_seqlens_kv=cu_seqlens, max_seqlen_q=max(seq_lengths), max_seqlen_kv=max(seq_lengths), - qkv_format='thd', + qkv_format="thd", ) batch = { - 'tokens': input_ids, - 'labels': labels, - 'loss_mask': loss_mask, - 'attention_mask': attention_mask, - 'position_ids': position_ids, - 'packed_seq_params': packed_seq_params, + "tokens": input_ids, + "labels": labels, + "loss_mask": loss_mask, + "attention_mask": attention_mask, + "position_ids": position_ids, + "packed_seq_params": packed_seq_params, } return batch @@ -638,7 +697,9 @@ def test_sharded_state_dict(self, tp, cp): args = self.create_test_args(tp, cp, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) model_parallel_cuda_manual_seed(_SEED) pg_collection = ProcessGroupCollection.use_mpu_process_groups() @@ -667,7 +728,9 @@ def test_forward_backward(self, tmp_path_dist_ckpt, tp, cp, full_recompute): """Test MTP forward and backward with gptmodel.""" tp_ref = 1 cp_ref = 1 - args = self.create_test_args(tp_ref, cp_ref, self.seq_length, self.micro_batch_size) + args = self.create_test_args( + tp_ref, cp_ref, self.seq_length, self.micro_batch_size + ) set_args(args) torch.manual_seed(_SEED) Utils.initialize_model_parallel( @@ -688,7 +751,7 @@ def test_forward_backward(self, tmp_path_dist_ckpt, tp, cp, full_recompute): tracker = MTPLossLoggingHelper.tracker mtp_loss_ref = None assert "loss_values" in tracker - mtp_loss_ref = tracker['loss_values'].clone() + mtp_loss_ref = tracker["loss_values"].clone() MTPLossLoggingHelper.clean_metrics_in_tracker() iteration = 123 @@ -699,7 +762,7 @@ def set_ckpt_path(ckpt_path): args.load = ckpt_path with TempNamedDir( - tmp_path_dist_ckpt / 'test_mtp_model_reconfiguration_model_A' + tmp_path_dist_ckpt / "test_mtp_model_reconfiguration_model_A" ) as ckpt_dir_A: set_ckpt_path(ckpt_dir_A) save_checkpoint( @@ -716,12 +779,18 @@ def set_ckpt_path(ckpt_path): # Test with different TP/CP configuration Utils.destroy_model_parallel() args = self.create_test_args( - tp, cp, self.seq_length, self.micro_batch_size, full_recompute=full_recompute + tp, + cp, + self.seq_length, + self.micro_batch_size, + full_recompute=full_recompute, ) set_args(args) set_ckpt_path(ckpt_dir_A) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) gpt_model, optimizer, opt_param_scheduler = setup_model_and_optimizer( ModelType.encoder_or_decoder, self.model_provider ) @@ -731,7 +800,9 @@ def set_ckpt_path(ckpt_path): batch = get_batch_on_this_cp_rank( batch, is_hybrid_cp=False, cp_group=get_context_parallel_group() ) - tokens, labels, loss_mask, attention_mask, position_ids, output_ref = batch.values() + tokens, labels, loss_mask, attention_mask, position_ids, output_ref = ( + batch.values() + ) output = gpt_model[0].forward( input_ids=tokens, position_ids=position_ids, @@ -741,9 +812,11 @@ def set_ckpt_path(ckpt_path): ) tracker = MTPLossLoggingHelper.tracker assert "loss_values" in tracker - mtp_loss = tracker['loss_values'].clone() + mtp_loss = tracker["loss_values"].clone() # Average MTP loss across CP ranks for comparison with reference - pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp']) + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["cp"] + ) torch.distributed.all_reduce( mtp_loss, group=pg_collection.cp, op=torch.distributed.ReduceOp.AVG ) @@ -772,14 +845,21 @@ def test_fp8_support(self, full_recompute): """Test MTP with FP8 training enabled.""" tp = 1 cp = 1 - fp8 = 'e4m3' + fp8 = "e4m3" args = self.create_test_args( - tp, cp, self.seq_length, self.micro_batch_size, fp8, full_recompute=full_recompute + tp, + cp, + self.seq_length, + self.micro_batch_size, + fp8, + full_recompute=full_recompute, ) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) batch = self.get_batch(self.seq_length, self.micro_batch_size) tokens, labels, loss_mask, attention_mask, position_ids = batch.values() gpt_model, optimizer, opt_param_scheduler = setup_model_and_optimizer( @@ -794,7 +874,9 @@ def test_fp8_support(self, full_recompute): loss_mask=loss_mask, ) - assert output.dtype == torch.float32 # Output should be converted back to float32 + assert ( + output.dtype == torch.float32 + ) # Output should be converted back to float32 loss = output.mean() loss.backward() @@ -814,16 +896,18 @@ def test_packed_sequences(self, tp, cp): set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) # Get packed batch batch = self.get_packed_batch(seq_lengths, micro_batch_size=1) - tokens = batch['tokens'] - labels = batch['labels'] - loss_mask = batch['loss_mask'] - attention_mask = batch['attention_mask'] - position_ids = batch['position_ids'] - packed_seq_params = batch['packed_seq_params'] + tokens = batch["tokens"] + labels = batch["labels"] + loss_mask = batch["loss_mask"] + attention_mask = batch["attention_mask"] + position_ids = batch["position_ids"] + packed_seq_params = batch["packed_seq_params"] # Create model model_parallel_cuda_manual_seed(_SEED) @@ -853,7 +937,7 @@ def test_packed_sequences(self, tp, cp): # Verify MTP loss was computed tracker = MTPLossLoggingHelper.tracker assert "loss_values" in tracker - mtp_loss = tracker['loss_values'].clone() + mtp_loss = tracker["loss_values"].clone() assert mtp_loss.shape[0] == args.mtp_num_layers MTPLossLoggingHelper.clean_metrics_in_tracker() @@ -885,12 +969,18 @@ def test_packed_sequences_with_full_recompute(self): total_seq_length = sum(seq_lengths) args = self.create_test_args( - tp=1, cp=1, sequence_length=total_seq_length, micro_batch_size=1, full_recompute=True + tp=1, + cp=1, + sequence_length=total_seq_length, + micro_batch_size=1, + full_recompute=True, ) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, context_parallel_size=1 + ) batch = self.get_packed_batch(seq_lengths, micro_batch_size=1) @@ -905,12 +995,12 @@ def test_packed_sequences_with_full_recompute(self): ) output = gpt_model[0].forward( - input_ids=batch['tokens'], - position_ids=batch['position_ids'], - attention_mask=batch['attention_mask'], - labels=batch['labels'], - loss_mask=batch['loss_mask'], - packed_seq_params=batch['packed_seq_params'], + input_ids=batch["tokens"], + position_ids=batch["position_ids"], + attention_mask=batch["attention_mask"], + labels=batch["labels"], + loss_mask=batch["loss_mask"], + packed_seq_params=batch["packed_seq_params"], ) # Backward must run end-to-end through the recomputed MTP layer. @@ -922,7 +1012,9 @@ def test_packed_sequences_with_full_recompute(self): def test_roll_tensor_none_input(self): """Test that roll_tensor returns (None, None) when given None input.""" - Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, context_parallel_size=1 + ) result, sum_val = roll_tensor(None, shifts=-1, dims=-1) assert result is None assert sum_val is None @@ -935,7 +1027,9 @@ def test_roll_tensor_shifts_left_and_zeroes_last(self): are not provided (RL training): label[i] = input_id[i+1], last position zeroed. The end-to-end derivation is covered by process_mtp_loss (see input_ids path). """ - Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, context_parallel_size=1 + ) # Simulate input_ids [batch=2, seq=5] input_ids = torch.tensor( [[10, 20, 30, 40, 50], [60, 70, 80, 90, 100]], dtype=torch.int64 @@ -955,13 +1049,13 @@ def test_process_mtp_loss_skips_when_no_labels_and_no_input_ids(self): hidden_size=8, num_layers=2, num_attention_heads=2, mtp_num_layers=1 ) hidden_states = torch.ones(2, 1, 4) - called = {'value': False} + called = {"value": False} def output_layer(hidden, weight=None, runtime_gather_output=None): return hidden.clone(), None def compute_language_model_loss(mtp_labels, mtp_logits): - called['value'] = True + called["value"] = True return torch.ones_like(mtp_labels, dtype=mtp_logits.dtype) out = process_mtp_loss( @@ -980,7 +1074,7 @@ def compute_language_model_loss(mtp_labels, mtp_logits): ) # First chunk is returned unchanged and the loss is never computed. - assert not called['value'] + assert not called["value"] assert torch.equal(out, torch.chunk(hidden_states, 2, dim=0)[0]) def test_process_mtp_loss_derives_labels_from_input_ids(self): @@ -996,13 +1090,13 @@ def test_process_mtp_loss_derives_labels_from_input_ids(self): # hidden_states is chunked into (1 + mtp_num_layers) along dim 0. hidden_states = torch.ones(2, 1, 5) input_ids = torch.tensor([[10, 20, 30, 40, 50]], dtype=torch.long) - seen = {'labels': None, 'masked_loss': None} + seen = {"labels": None, "masked_loss": None} def output_layer(hidden, weight=None, runtime_gather_output=None): return hidden.clone(), None def compute_language_model_loss(mtp_labels, mtp_logits): - seen['labels'] = mtp_labels.clone() + seen["labels"] = mtp_labels.clone() # Per-position loss of 1.0 so loss_mask * loss exposes the active mask. return torch.ones_like(mtp_labels, dtype=torch.float32) @@ -1023,8 +1117,10 @@ def compute_language_model_loss(mtp_labels, mtp_logits): # input_ids rolled twice (once to SFT format, once in the MTP layer loop): # [10,20,30,40,50] -> [20,30,40,50,0] -> [30,40,50,0,0]. - assert seen['labels'] is not None, "loss should be computed in RL mode" - assert torch.equal(seen['labels'], torch.tensor([[30, 40, 50, 0, 0]], dtype=torch.long)) + assert seen["labels"] is not None, "loss should be computed in RL mode" + assert torch.equal( + seen["labels"], torch.tensor([[30, 40, 50, 0, 0]], dtype=torch.long) + ) @pytest.mark.parametrize("cp", [1, 2]) def test_roll_tensor_with_packed_sequences(self, cp): @@ -1033,9 +1129,13 @@ def test_roll_tensor_with_packed_sequences(self, cp): For CP=1: Tests standard packed sequence rolling with verified expected values For CP=2: Tests CP-enabled rolling executes without errors """ - Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, context_parallel_size=cp + ) cp_group = get_context_parallel_group() if cp > 1 else None - cp_rank = torch.distributed.get_rank(group=cp_group) if cp_group is not None else 0 + cp_rank = ( + torch.distributed.get_rank(group=cp_group) if cp_group is not None else 0 + ) if cp == 1: # Test case: Simple packed sequences (CP disabled) @@ -1047,12 +1147,16 @@ def test_roll_tensor_with_packed_sequences(self, cp): cu_seqlens_kv=cu_seqlens, max_seqlen_q=3, max_seqlen_kv=3, - qkv_format='thd', + qkv_format="thd", ) # Roll by -1 (shift left) rolled, sum_val = roll_tensor( - tensor, shifts=-1, dims=0, cp_group=cp_group, packed_seq_params=packed_seq_params + tensor, + shifts=-1, + dims=0, + cp_group=cp_group, + packed_seq_params=packed_seq_params, ) # Expected: [2, 3, 0, 5, 0] - boundaries at indices 2 and 4 are zeroed @@ -1088,21 +1192,25 @@ def test_roll_tensor_with_packed_sequences(self, cp): cu_seqlens_kv=cu_seqlens, max_seqlen_q=6, # max(4, 6) - max local seq length per sequence max_seqlen_kv=6, - qkv_format='thd', + qkv_format="thd", ) # Roll by -1 (shift left) with CP communication rolled, sum_val = roll_tensor( - tensor, shifts=-1, dims=0, cp_group=cp_group, packed_seq_params=packed_seq_params + tensor, + shifts=-1, + dims=0, + cp_group=cp_group, + packed_seq_params=packed_seq_params, ) # Verify the rolled tensor matches expected values - assert ( - rolled.shape == expected.shape - ), f"Shape mismatch: expected {expected.shape}, got {rolled.shape}" - assert torch.equal( - rolled, expected - ), f"CP Rank {cp_rank}: Expected\n{expected}\nbut got\n{rolled}\nDiff:\n{rolled - expected}" + assert rolled.shape == expected.shape, ( + f"Shape mismatch: expected {expected.shape}, got {rolled.shape}" + ) + assert torch.equal(rolled, expected), ( + f"CP Rank {cp_rank}: Expected\n{expected}\nbut got\n{rolled}\nDiff:\n{rolled - expected}" + ) # Verify sum is correct assert sum_val.numel() == 1, "Sum should be a scalar" @@ -1152,10 +1260,22 @@ class DummyOutputLayer: def __init__(self, gather_output): self.gather_output = gather_output - assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), None) is False - assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=False), None) is True - assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), True) is False - assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), False) is True + assert ( + _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), None) + is False + ) + assert ( + _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=False), None) + is True + ) + assert ( + _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), True) + is False + ) + assert ( + _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), False) + is True + ) def test_track_mtp_metrics(self): """Test tracking MTP metrics including acceptance rate.""" @@ -1166,7 +1286,11 @@ def test_track_mtp_metrics(self): for i in range(num_layers): MTPLossLoggingHelper.save_metrics_to_tracker( - loss=loss, correct=correct, total=total, layer_number=i, num_layers=num_layers + loss=loss, + correct=correct, + total=total, + layer_number=i, + num_layers=num_layers, ) class DummyWriter: @@ -1197,20 +1321,25 @@ def log(self, metrics, iteration): # Verify loss uses the legacy normalized MTP loss scaled by loss_scale. expected_loss = loss * loss_scale for i in range(num_layers): - assert f"mtp_{i+1} loss" in writer.scalars - assert torch.isclose(torch.as_tensor(writer.scalars[f"mtp_{i+1} loss"]), expected_loss) - assert torch.isclose(total_loss_dict[f"mtp_{i+1} loss"], expected_loss) + assert f"mtp_{i + 1} loss" in writer.scalars + assert torch.isclose( + torch.as_tensor(writer.scalars[f"mtp_{i + 1} loss"]), expected_loss + ) + assert torch.isclose(total_loss_dict[f"mtp_{i + 1} loss"], expected_loss) # Verify acceptance rate is computed as (correct / total) * 100 expected_rate = (correct / total) * 100.0 for i in range(num_layers): - assert f"mtp_{i+1}_acceptance_rate" in writer.scalars + assert f"mtp_{i + 1}_acceptance_rate" in writer.scalars assert torch.isclose( - torch.as_tensor(writer.scalars[f"mtp_{i+1}_acceptance_rate"]), expected_rate + torch.as_tensor(writer.scalars[f"mtp_{i + 1}_acceptance_rate"]), + expected_rate, ) - assert f"mtp_{i+1}_cumulative_acceptance_rate" in writer.scalars + assert f"mtp_{i + 1}_cumulative_acceptance_rate" in writer.scalars assert torch.isclose( - torch.as_tensor(writer.scalars[f"mtp_{i+1}_cumulative_acceptance_rate"]), + torch.as_tensor( + writer.scalars[f"mtp_{i + 1}_cumulative_acceptance_rate"] + ), expected_rate, ) @@ -1237,16 +1366,23 @@ def log(self, metrics, iteration): ) expected_second_rate = (second_correct / second_total) * 100.0 - expected_cumulative_rate = ((correct + second_correct) / (total + second_total)) * 100.0 + expected_cumulative_rate = ( + (correct + second_correct) / (total + second_total) + ) * 100.0 for i in range(num_layers): assert torch.isclose( - torch.as_tensor(writer.scalars[f"mtp_{i+1}_acceptance_rate"]), expected_second_rate + torch.as_tensor(writer.scalars[f"mtp_{i + 1}_acceptance_rate"]), + expected_second_rate, ) assert torch.isclose( - torch.as_tensor(writer.scalars[f"mtp_{i+1}_cumulative_acceptance_rate"]), + torch.as_tensor( + writer.scalars[f"mtp_{i + 1}_cumulative_acceptance_rate"] + ), expected_cumulative_rate, ) - assert torch.isclose(total_loss_dict[f"mtp_{i+1} loss"], expected_loss * 2) + assert torch.isclose( + total_loss_dict[f"mtp_{i + 1} loss"], expected_loss * 2 + ) # Verify tracker is cleaned assert torch.all(MTPLossLoggingHelper.tracker["loss_values"] == 0) @@ -1263,10 +1399,18 @@ def test_track_mtp_loss_preserves_legacy_normalized_loss_semantics(self): layer_number = 0 MTPLossLoggingHelper.save_metrics_to_tracker( - loss=first_loss, correct=correct, total=total, layer_number=layer_number, num_layers=1 + loss=first_loss, + correct=correct, + total=total, + layer_number=layer_number, + num_layers=1, ) MTPLossLoggingHelper.save_metrics_to_tracker( - loss=second_loss, correct=correct, total=total, layer_number=layer_number, num_layers=1 + loss=second_loss, + correct=correct, + total=total, + layer_number=layer_number, + num_layers=1, ) class DummyWriter: @@ -1294,7 +1438,7 @@ class TestMultiTokenPredictionHybrid: def setup_method(self, method): self.seq_length = 32 self.micro_batch_size = 2 - os.environ['CUDA_DEVICE_MAX_CONNECTIONS'] = '1' + os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "1" def teardown_method(self, method): Utils.destroy_model_parallel() @@ -1337,7 +1481,7 @@ def create_test_args( destroy_global_vars() destroy_num_microbatches_calculator() - sys.argv = ['test_multi_token_prediction_hybrid.py'] + sys.argv = ["test_multi_token_prediction_hybrid.py"] args = parse_args() args.mtp_num_layers = 2 args.mtp_loss_scaling_factor = 0.1 @@ -1353,9 +1497,9 @@ def create_test_args( args.tensor_model_parallel_size = tp args.sequence_parallel = True if tp > 1 else False args.context_parallel_size = cp - args.position_embedding_type = 'rope' + args.position_embedding_type = "rope" args.train_iters = 1 - args.ckpt_format = 'torch_dist' + args.ckpt_format = "torch_dist" args.lr = 3e-5 args.attention_dropout = 0.0 args.hidden_dropout = 0.0 @@ -1367,10 +1511,10 @@ def create_test_args( args.hybrid_layer_pattern = "M*M*/M*/M*" if fp8 is not None: - args.fp8 = 'e4m3' + args.fp8 = "e4m3" if full_recompute: - args.recompute_granularity = 'full' - args.recompute_method = 'uniform' + args.recompute_granularity = "full" + args.recompute_method = "uniform" args.recompute_num_layers = 1 else: args.recompute_granularity = None @@ -1383,19 +1527,26 @@ def create_test_args( def get_batch(self, seq_length, micro_batch_size): data = list(range(seq_length)) - input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - labels = 1 + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - position_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + input_ids = ( + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + ) + labels = ( + 1 + + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + ) + position_ids = ( + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + ) attention_mask = torch.ones( (micro_batch_size, 1, seq_length, seq_length), dtype=bool ).cuda() loss_mask = torch.ones(seq_length).repeat((micro_batch_size, 1)).cuda() batch = { - 'tokens': input_ids, - 'labels': labels, - 'loss_mask': loss_mask, - 'attention_mask': attention_mask, - 'position_ids': position_ids, + "tokens": input_ids, + "labels": labels, + "loss_mask": loss_mask, + "attention_mask": attention_mask, + "position_ids": position_ids, } return batch @@ -1406,7 +1557,9 @@ def test_sharded_state_dict_mamba(self, tp, cp): args = self.create_test_args(tp, cp, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) model_parallel_cuda_manual_seed(_SEED) pg_collection = ProcessGroupCollection.use_mpu_process_groups() @@ -1430,7 +1583,9 @@ def test_forward_backward_mamba(self, tmp_path_dist_ckpt, tp, cp): """Test MTP forward and backward with Mamba hybrid model.""" tp_ref = 1 cp_ref = 1 - args = self.create_test_args(tp_ref, cp_ref, self.seq_length, self.micro_batch_size) + args = self.create_test_args( + tp_ref, cp_ref, self.seq_length, self.micro_batch_size + ) set_args(args) torch.manual_seed(_SEED) Utils.initialize_model_parallel( @@ -1459,7 +1614,7 @@ def test_forward_backward_mamba(self, tmp_path_dist_ckpt, tp, cp): tracker = MTPLossLoggingHelper.tracker mtp_loss_ref = None assert "loss_values" in tracker - mtp_loss_ref = tracker['loss_values'].clone() + mtp_loss_ref = tracker["loss_values"].clone() MTPLossLoggingHelper.clean_metrics_in_tracker() iteration = 123 @@ -1469,7 +1624,9 @@ def set_ckpt_path(ckpt_path): args.save = ckpt_path args.load = ckpt_path - with TempNamedDir(tmp_path_dist_ckpt / 'test_mtp_mamba_model_reconfiguration') as ckpt_dir: + with TempNamedDir( + tmp_path_dist_ckpt / "test_mtp_mamba_model_reconfiguration" + ) as ckpt_dir: set_ckpt_path(ckpt_dir) save_checkpoint( iteration, @@ -1487,7 +1644,9 @@ def set_ckpt_path(ckpt_path): set_args(args) set_ckpt_path(ckpt_dir) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) model_parallel_cuda_manual_seed(_SEED) cfg_container = Utils.pretrain_config_from_global_args(args, "hybrid") @@ -1504,7 +1663,9 @@ def set_ckpt_path(ckpt_path): batch = get_batch_on_this_cp_rank( batch, is_hybrid_cp=False, cp_group=get_context_parallel_group() ) - tokens, labels, loss_mask, attention_mask, position_ids, output_ref = batch.values() + tokens, labels, loss_mask, attention_mask, position_ids, output_ref = ( + batch.values() + ) output = mamba_model[0].forward( input_ids=tokens, position_ids=position_ids, @@ -1514,8 +1675,10 @@ def set_ckpt_path(ckpt_path): ) tracker = MTPLossLoggingHelper.tracker assert "loss_values" in tracker - mtp_loss = tracker['loss_values'].clone() - pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp']) + mtp_loss = tracker["loss_values"].clone() + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["cp"] + ) torch.distributed.all_reduce( mtp_loss, group=pg_collection.cp, op=torch.distributed.ReduceOp.AVG ) @@ -1539,7 +1702,9 @@ def test_attention_mask_validation_mamba(self): args = self.create_test_args(tp, cp, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, context_parallel_size=cp + ) pg_collection = ProcessGroupCollection.use_mpu_process_groups() model_cfg = hybrid_config_from_args(args) builder_cls = model_cfg.get_builder_cls() @@ -1554,6 +1719,8 @@ def test_attention_mask_validation_mamba(self): assert mamba_model[0].mtp is not None except AssertionError as e: if "Multi-Token Prediction (MTP) is not yet supported" in str(e): - pytest.fail(f"Attention mask validation failed for Mamba hybrid model: {e}") + pytest.fail( + f"Attention mask validation failed for Mamba hybrid model: {e}" + ) else: raise diff --git a/tests/unit_tests/transformer/test_torch_norm.py b/tests/unit_tests/transformer/test_torch_norm.py index d3e69a59c6e..8951eb2f365 100644 --- a/tests/unit_tests/transformer/test_torch_norm.py +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -2,10 +2,7 @@ import torch -from megatron.core.transformer.torch_norm import ( - AccuracyCompatibleRMSNorm, - WrappedTorchNorm, -) +from megatron.core.transformer.torch_norm import WrappedTorchNorm from megatron.core.transformer.transformer_config import TransformerConfig @@ -20,40 +17,9 @@ def _config(**overrides): return TransformerConfig(**values) -def test_accuracy_compatible_rmsnorm_matches_explicit_formula(): - config = _config(norm_accuracy_compatible=True) - norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5).cuda().bfloat16() - x = torch.randn(2, 3, 64, device="cuda", dtype=torch.bfloat16) - - output = norm(x) - x_float = x.float() - expected = ( - x_float - * torch.rsqrt(x_float.pow(2).mean(dim=-1, keepdim=True) + 1e-5) - * norm.weight.float() - ).to(torch.bfloat16) - - assert isinstance(norm, AccuracyCompatibleRMSNorm) - assert torch.equal(output, expected) - - -def test_accuracy_compatible_rmsnorm_canonicalizes_zero_input_gradients(): - config = _config(norm_accuracy_compatible=True) - norm = WrappedTorchNorm(config=config, hidden_size=4, eps=1e-5).cuda().bfloat16() - with torch.no_grad(): - norm.weight.copy_(torch.tensor([-1.0, 1.0, -2.0, 2.0], device="cuda")) - x = torch.tensor( - [[[1.0, -1.0, 2.0, -2.0]]], device="cuda", dtype=torch.bfloat16, requires_grad=True - ) - - norm(x).backward(torch.zeros_like(x)) - - assert torch.equal(x.grad, torch.zeros_like(x.grad)) - assert torch.equal(x.grad.view(torch.uint16), torch.zeros_like(x.grad.view(torch.uint16))) - - -def test_default_rmsnorm_stays_native(): - config = _config(norm_accuracy_compatible=False) +def test_rmsnorm_uses_native_torch_implementation(): + config = _config(norm_accuracy_compatible=True, params_dtype=torch.bfloat16) norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) assert isinstance(norm, torch.nn.RMSNorm) + assert norm.weight.dtype == torch.bfloat16 From c91e01173d9dcab31bcc324c5c78d5e20745eb48 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Thu, 27 Aug 2026 17:01:42 +0800 Subject: [PATCH 09/27] Expose LocalSpecProvider.linear for DSA/MLA down-projections backend.linear() is required when DSA is built with the local (non-TE) spec. Return TELinear in duplicated mode, distinct from column-parallel. --- megatron/core/models/backends.py | 8 ++++++++ .../models/test_local_spec_provider_linear.py | 13 +++++++++++++ 2 files changed, 21 insertions(+) create mode 100644 tests/unit_tests/models/test_local_spec_provider_linear.py diff --git a/megatron/core/models/backends.py b/megatron/core/models/backends.py index a270161ddd6..918252a1d15 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -99,6 +99,14 @@ def activation_func(self) -> TEActivationFunctionBuilder | None: class LocalSpecProvider(BackendSpecProvider): """A protocol for providing Local submodules used in Spec building.""" + def linear(self) -> type: + """TP-replicated Linear (mcore TELinear parallel_mode=duplicated). + + DSA indexer / MLA down-projections call backend.linear(). Without this + method, mcore_bridge can only build those layers via TESpecProvider. + """ + return TELinear + def column_parallel_linear(self) -> type: """Which column parallel linear module the backend uses""" return ColumnParallelLinear diff --git a/tests/unit_tests/models/test_local_spec_provider_linear.py b/tests/unit_tests/models/test_local_spec_provider_linear.py new file mode 100644 index 00000000000..63d9b3ab727 --- /dev/null +++ b/tests/unit_tests/models/test_local_spec_provider_linear.py @@ -0,0 +1,13 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +"""LocalSpecProvider must expose backend.linear() for DSA/MLA down-projections.""" + +from megatron.core.extensions.transformer_engine import TELinear +from megatron.core.models.backends import LocalSpecProvider +from megatron.core.tensor_parallel.layers import ColumnParallelLinear + + +def test_local_spec_provider_linear_is_replicated_te_linear(): + backend = LocalSpecProvider() + assert backend.linear() is TELinear + assert backend.column_parallel_linear() is ColumnParallelLinear + assert backend.linear() is not backend.column_parallel_linear() From dc12dbcfa872e13e6be2343af14cd27d6d419166 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Thu, 27 Aug 2026 18:32:37 +0800 Subject: [PATCH 10/27] Return modelopt Linear from LocalSpecProvider.linear, not TELinear TESpecProvider.linear stays TELinear, so the TE-on default path is unchanged. LocalSpecProvider.linear is the TE-off counterpart used when accuracy-compatible forces LocalSpecProvider for DSA/MLA down-projections. --- megatron/core/models/backends.py | 10 ++++++---- .../models/test_local_spec_provider_linear.py | 8 +++++--- 2 files changed, 11 insertions(+), 7 deletions(-) diff --git a/megatron/core/models/backends.py b/megatron/core/models/backends.py index 918252a1d15..c543a49e266 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -10,6 +10,7 @@ TEColumnParallelGroupedLinear, TERowParallelGroupedLinear, ) +from megatron.core.post_training.modelopt.layers import Linear from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.transformer.dot_product_attention import DotProductAttention from megatron.core.transformer.mlp import MLPSubmodules, TEActivationFunctionBuilder @@ -100,12 +101,13 @@ class LocalSpecProvider(BackendSpecProvider): """A protocol for providing Local submodules used in Spec building.""" def linear(self) -> type: - """TP-replicated Linear (mcore TELinear parallel_mode=duplicated). + """TP-replicated local Linear (modelopt Linear, not TELinear). - DSA indexer / MLA down-projections call backend.linear(). Without this - method, mcore_bridge can only build those layers via TESpecProvider. + DSA indexer / MLA down-projections call backend.linear(). TESpecProvider + still returns TELinear; this method is the TE-off counterpart so a + LocalSpecProvider DSA spec does not re-enter Transformer Engine. """ - return TELinear + return Linear def column_parallel_linear(self) -> type: """Which column parallel linear module the backend uses""" diff --git a/tests/unit_tests/models/test_local_spec_provider_linear.py b/tests/unit_tests/models/test_local_spec_provider_linear.py index 63d9b3ab727..4564ce0adf8 100644 --- a/tests/unit_tests/models/test_local_spec_provider_linear.py +++ b/tests/unit_tests/models/test_local_spec_provider_linear.py @@ -1,13 +1,15 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. -"""LocalSpecProvider must expose backend.linear() for DSA/MLA down-projections.""" +"""LocalSpecProvider must expose a non-TE backend.linear() for DSA/MLA.""" from megatron.core.extensions.transformer_engine import TELinear from megatron.core.models.backends import LocalSpecProvider +from megatron.core.post_training.modelopt.layers import Linear from megatron.core.tensor_parallel.layers import ColumnParallelLinear -def test_local_spec_provider_linear_is_replicated_te_linear(): +def test_local_spec_provider_linear_is_replicated_local_linear(): backend = LocalSpecProvider() - assert backend.linear() is TELinear + assert backend.linear() is Linear + assert backend.linear() is not TELinear assert backend.column_parallel_linear() is ColumnParallelLinear assert backend.linear() is not backend.column_parallel_linear() From 4e81432a0ac965a9fb08f7a58ec6b3a2409152af Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Sat, 5 Sep 2026 07:39:18 +0800 Subject: [PATCH 11/27] Keep 1/num_tokens gradient scaling under accuracy-compatible calculate_per_token_loss=true backprops the raw sum(loss*mask). Skipping mcore's 1/num_tokens scale under UAC left every gradient ~44x too large at step 1 (grad_norm 2511 vs 57) and drifted IEEE at step 2. Restore the scale; keep the fp32 gate-wgrad DP all-reduce. Signed-off-by: Zhan Rongrui --- .../core/distributed/finalize_model_grads.py | 33 +++++++------------ .../distributed/test_finalize_model_grads.py | 10 ++++++ 2 files changed, 22 insertions(+), 21 deletions(-) diff --git a/megatron/core/distributed/finalize_model_grads.py b/megatron/core/distributed/finalize_model_grads.py index 60397e9af9f..14cd83186e4 100644 --- a/megatron/core/distributed/finalize_model_grads.py +++ b/megatron/core/distributed/finalize_model_grads.py @@ -3,7 +3,6 @@ from functools import partial from typing import Callable, Dict, List, Optional, Union -import os import torch from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors @@ -460,17 +459,15 @@ def finalize_model_grads( """ config = get_model_config(model[0]) - - # [对齐修复] use_accuracy_compatible=1: PaddleFleet 的 fixed-loss 路径已在 autograd 图内除过 - # 本地有效 token 数, MCore 这里再用全局 num_tokens 缩放会引入 ~global_token / local_token - # 倍的额外因子 (实测 ~74.64x)。在对齐模式下跳过 num_tokens 全局缩放, 改为 grad sync 后做 - # 1/dp_size 平均, 与 Paddle DP 平均语义对齐; 同时对 RouterGatingLinearFunction 记录的 - # fp32 gate wgrad 做一次 DP all-reduce, 与参考实现一致。 + + # E-172: keep mcore 1/num_tokens scaling under use_accuracy_compatible. + # calculate_per_token_loss=true backprops raw sum(loss*mask); skipping this + # left every gradient num_tokens times too large (44.0x at step 1). Formal + # N=5 with the skip matched step-1 IEEE but drifted at steps 2-5. from ..transformer.module import _use_accuracy_compatible - loss_normalized_in_graph = _use_accuracy_compatible() and (num_tokens is not None) - if loss_normalized_in_graph: - num_tokens = None - + + loss_normalized_in_graph = False + tp_dp_cp_group = None if pg_collection is not None: assert hasattr(pg_collection, 'tp') @@ -563,16 +560,10 @@ def finalize_model_grads( reset_model_temporary_tensors(config, model) - # [对齐修复] 对应上方 loss_normalized_in_graph 早跳过分支: 跳过全局 num_tokens 缩放后, - # 这里改为 grad sync 后做 1/dp_size 平均, 与 PaddleFleet DP 平均语义对齐; - # 并对 RouterGatingLinearFunction 记录到 param._run_torch_gate_fp32_wgrad 的 fp32 gate - # wgrad 做一次 DP all-reduce (若该机制未启用则 getattr 为 None, 代码为 no-op)。 - if loss_normalized_in_graph: - dp_size = parallel_state.get_data_parallel_world_size(with_context_parallel=True) - if dp_size > 1: - for model_chunk in model: - model_chunk.scale_gradients(1.0 / dp_size) - + # All-reduce the fp32 gate wgrad that RouterGatingLinearFunction records on + # ``param._run_torch_gate_fp32_wgrad`` across DP, matching the reference + # implementation. A no-op when that mechanism is not enabled. + if _use_accuracy_compatible(): for model_chunk in model: for param in model_chunk.parameters(): gate_wgrad = getattr(param, "_run_torch_gate_fp32_wgrad", None) diff --git a/tests/unit_tests/distributed/test_finalize_model_grads.py b/tests/unit_tests/distributed/test_finalize_model_grads.py index ee535c29baf..431ef59aa77 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -21,6 +21,16 @@ from tests.unit_tests.test_utilities import Utils +def test_uac_keeps_mcore_num_tokens_scaling(): + """Shipped finalize_model_grads must not skip 1/num_tokens under UAC (E-172).""" + src = inspect.getsource(finalize_model_grads) + assert "loss_normalized_in_graph = False" in src + assert "num_tokens = None" not in src.split("loss_normalized_in_graph = False", 1)[1][ + :800 + ] + assert "scale_gradients(1.0 / dp_size)" not in src + + class _RouterExpertBiasModel(torch.nn.Module): def __init__(self, config, local_tokens_per_expert): super().__init__() From 9815d1583b05b97dd0bb4cabaa93843aaa2d4451 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Sat, 5 Sep 2026 09:24:46 +0800 Subject: [PATCH 12/27] Route ETP --- .../core/extensions/transformer_engine.py | 18 +++++-- megatron/core/tensor_parallel/layers.py | 49 +++++++++++++++++-- .../unit_tests/tensor_parallel/test_layers.py | 48 +++++++++++++++++- 3 files changed, 107 insertions(+), 8 deletions(-) diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index b7de1013695..103c2a9a7d6 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -916,8 +916,12 @@ def __init__( for param in self.parameters(): if is_expert: - # Reduce the gradient on the expert_data_parallel group for expert linear layers - setattr(param, "allreduce", not self.expert_parallel) + # Reduce the gradient on the expert_data_parallel group for expert linear layers. + # See _expert_grads_need_own_dp_domain in tensor_parallel/layers.py: ETP < TP + # also puts expert grads in their own (larger) data-parallel domain. + from ..tensor_parallel.layers import _expert_grads_need_own_dp_domain + + setattr(param, "allreduce", not _expert_grads_need_own_dp_domain(self.config)) else: # Reduce the gradient on DP group setattr(param, "allreduce", True) @@ -2035,7 +2039,15 @@ def __init__( ) for param in self.parameters(): - setattr(param, "allreduce", not (is_expert and self.expert_parallel)) + # See _expert_grads_need_own_dp_domain in tensor_parallel/layers.py: + # ETP < TP also puts expert grads in their own data-parallel domain. + from ..tensor_parallel.layers import _expert_grads_need_own_dp_domain + + setattr( + param, + "allreduce", + not (is_expert and _expert_grads_need_own_dp_domain(self.config)), + ) # Explicitly stamp partition_dim and partition_stride on expert weight # tensors when explicit_expert_comm cleared parallel_mode. TE ≤2.12 diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 2f927d5218d..2e3f024c2b0 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -777,6 +777,31 @@ def linear_with_grad_accumulation_and_async_allreduce( linear_with_grad_accumulation_and_async_allreduce.warned = False +def _expert_grads_need_own_dp_domain(config) -> bool: + """Whether expert parameters must be reduced over the expert-data-parallel group. + + mcore normally decides this from ``expert_model_parallel_size > 1`` alone, which + is correct only when the expert tensor-parallel size equals the dense one: the + expert parameter is then sharded exactly like a dense parameter and its + data-parallel domain coincides with ``dp_cp``. + + With ``expert_tensor_parallel_size < tensor_model_parallel_size`` (the accuracy + -compatible topology uses ETP=1 with TP=2) every rank in the tensor-parallel + group holds a FULL copy of the expert weight while consuming only its own + sequence-parallel shard of the tokens. Those partial weight gradients live in a + larger data-parallel domain (``expt_dp`` = TP/ETP times ``dp_cp``) and must be + summed. Leaving ``allreduce=True`` puts them in the dense bucket, which is + reduced over ``dp_cp`` — size 1 in this topology — so the reduction silently + never happens and every expert gradient stays a per-rank partial sum. + """ + if config.expert_model_parallel_size > 1: + return True + etp = getattr(config, 'expert_tensor_parallel_size', None) + if etp is None: + return False + return etp != config.tensor_model_parallel_size + + class ColumnParallelLinear(torch.nn.Module): """Linear layer with column parallelism. @@ -921,7 +946,11 @@ def __init__( tensor=self.weight, is_parallel=True, dim=0, stride=stride ) - setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.weight, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) else: self.weight = None @@ -943,7 +972,11 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.bias, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) else: self.register_parameter("bias", None) @@ -1271,7 +1304,11 @@ def __init__( set_tensor_model_parallel_attributes( tensor=self.weight, is_parallel=True, dim=1, stride=stride ) - setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.weight, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) if bias: if config.use_cpu_initialization: @@ -1289,7 +1326,11 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr( + self.bias, + "allreduce", + not (self.is_expert and _expert_grads_need_own_dp_domain(config)), + ) setattr(self.bias, "sequence_parallel", self.sequence_parallel) else: self.register_parameter("bias", None) diff --git a/tests/unit_tests/tensor_parallel/test_layers.py b/tests/unit_tests/tensor_parallel/test_layers.py index dbc27f502c6..5c4398aefcf 100644 --- a/tests/unit_tests/tensor_parallel/test_layers.py +++ b/tests/unit_tests/tensor_parallel/test_layers.py @@ -1,12 +1,58 @@ # Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. +from pathlib import Path +from types import SimpleNamespace + import pytest import torch -from megatron.core.tensor_parallel.layers import linear_with_frozen_weight +from megatron.core.tensor_parallel.layers import ( + _expert_grads_need_own_dp_domain, + linear_with_frozen_weight, +) from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region from tests.unit_tests.test_utilities import Utils +def test_expert_grads_need_own_dp_domain_etp_lt_tp(): + """EP=1 / ETP=1 / TP=2 expert wgrad must leave the dense dp_cp bucket.""" + frozen = SimpleNamespace( + expert_model_parallel_size=1, + tensor_model_parallel_size=2, + expert_tensor_parallel_size=1, + ) + assert _expert_grads_need_own_dp_domain(frozen) is True + eq = SimpleNamespace( + expert_model_parallel_size=1, + tensor_model_parallel_size=2, + expert_tensor_parallel_size=2, + ) + assert _expert_grads_need_own_dp_domain(eq) is False + ep2 = SimpleNamespace( + expert_model_parallel_size=2, + tensor_model_parallel_size=2, + expert_tensor_parallel_size=1, + ) + assert _expert_grads_need_own_dp_domain(ep2) is True + missing = SimpleNamespace( + expert_model_parallel_size=1, + tensor_model_parallel_size=2, + expert_tensor_parallel_size=None, + ) + assert _expert_grads_need_own_dp_domain(missing) is False + + +def test_expert_dp_domain_is_wired_into_linear_allreduce(): + """Column/Row/TE expert allreduce must use the ETP Date: Sat, 5 Sep 2026 16:12:55 +0800 Subject: [PATCH 13/27] Select stack-paired PaddleFleet pin for alignment CI Keep pull_request on the historical CodeSync/develop tarball and develop/latest wheels. workflow_dispatch can opt into fail-closed SHA checkout plus artifact digest checks; pairing remains unproven unless this job built the wheels from the pin. MinimaxV2.5_EP2 and GLM45Air_EP2 are unchanged. --- .../workflows/alignment_model_accuracy.yml | 105 ++- scripts/select_paddlefleet_alignment_pin.sh | 638 ++++++++++++++++++ 2 files changed, 725 insertions(+), 18 deletions(-) create mode 100755 scripts/select_paddlefleet_alignment_pin.sh diff --git a/.github/workflows/alignment_model_accuracy.yml b/.github/workflows/alignment_model_accuracy.yml index 185330c16a9..02d2b3e7264 100644 --- a/.github/workflows/alignment_model_accuracy.yml +++ b/.github/workflows/alignment_model_accuracy.yml @@ -3,6 +3,36 @@ name: Alignment Model Accuracy on: pull_request: workflow_dispatch: + inputs: + paddlefleet_mode: + description: "develop keeps the historical CodeSync tarball. stack-paired fail-closed SHA check." + type: choice + default: develop + options: [develop, stack-paired] + paddlefleet_pin_sha: + description: "40-hex PaddleFleet commit (required for stack-paired)" + required: false + type: string + paddlefleet_git_url: + description: "Git URL that contains paddlefleet_pin_sha" + required: false + type: string + paddlefleet_wheel_url: + description: "Optional wheel URL from Build Fleet whl / Actions artifact metadata" + required: false + type: string + paddlefleet_wheel_sha256: + description: "sha256 for paddlefleet_wheel_url" + required: false + type: string + paddlefleet_ops_wheel_url: + description: "Optional paddlefleet_ops wheel URL" + required: false + type: string + paddlefleet_ops_wheel_sha256: + description: "sha256 for paddlefleet_ops_wheel_url" + required: false + type: string concurrency: group: Alignment-${{ github.workflow }}-${{ github.event.pull_request.number }} @@ -17,6 +47,14 @@ env: TASK: Megatron-LM-${{ github.sha }}-alignment CE_name: alignment-Megatron-LM no_proxy: "localhost,bj.bcebos.com,su.bcebos.com,bcebos.com,apiin.im.baidu.com,gitee.com,aliyun.com,.baidu.com,.tuna.tsinghua.edu.cn" + # pull_request stays on historical develop. stack-paired is workflow_dispatch only. + ALIGNMENT_PADDLEFLEET_MODE: ${{ github.event_name == 'workflow_dispatch' && github.event.inputs.paddlefleet_mode || 'develop' }} + PADDLEFLEET_PIN_SHA: ${{ github.event.inputs.paddlefleet_pin_sha }} + PADDLEFLEET_GIT_URL: ${{ github.event.inputs.paddlefleet_git_url }} + PADDLEFLEET_WHEEL_URL: ${{ github.event.inputs.paddlefleet_wheel_url }} + PADDLEFLEET_WHEEL_SHA256: ${{ github.event.inputs.paddlefleet_wheel_sha256 }} + PADDLEFLEET_OPS_WHEEL_URL: ${{ github.event.inputs.paddlefleet_ops_wheel_url }} + PADDLEFLEET_OPS_WHEEL_SHA256: ${{ github.event.inputs.paddlefleet_ops_wheel_sha256 }} defaults: run: @@ -55,17 +93,28 @@ jobs: -e no_proxy \ -e CE_name \ -e python_version \ + -e ALIGNMENT_PADDLEFLEET_MODE \ + -e PADDLEFLEET_PIN_SHA \ + -e PADDLEFLEET_GIT_URL \ + -e PADDLEFLEET_WHEEL_URL \ + -e PADDLEFLEET_WHEEL_SHA256 \ + -e PADDLEFLEET_OPS_WHEEL_URL \ + -e PADDLEFLEET_OPS_WHEEL_SHA256 \ -w /workspace $IMAGE_NAME - name: Checkout Code run: | docker exec -t $container_name /bin/bash -c ' rm -rf * .[^.]* source $work_dir/../../../proxy - echo "Download PaddleFleet form https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar" - wget -q --no-proxy https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar --no-check-certificate - rm -rf PaddleFleet && tar xf PaddleFleet.tar && rm -rf PaddleFleet.tar - cd PaddleFleet && git pull && cd - - + if [ "${ALIGNMENT_PADDLEFLEET_MODE:-develop}" = "stack-paired" ]; then + echo "stack-paired: skip CodeSync/develop PaddleFleet.tar; selector runs after python" + else + echo "Download PaddleFleet form https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar" + wget -q --no-proxy https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar --no-check-certificate + rm -rf PaddleFleet && tar xf PaddleFleet.tar && rm -rf PaddleFleet.tar + cd PaddleFleet && git pull && cd - + fi + echo "Download Megatron-LM form https://paddle-github-action.bj.bcebos.com/whl/Megatron-LM.tar.gz" wget -q --no-proxy https://paddle-github-action.bj.bcebos.com/whl/Megatron-LM.tar.gz --no-check-certificate rm -rf Megatron-LM && tar zxf Megatron-LM.tar.gz && rm -rf Megatron-LM.tar.gz @@ -110,18 +159,34 @@ jobs: ldconfig BOS=https://paddle-github-action.bj.bcebos.com - echo "::group::Download paddlefleet / paddlefleet_ops / ms_swift / mcore-bridge wheels from BOS" cd /workspace - for url in \ - $BOS/PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl \ - $BOS/PaddleFleet/develop/latest/cu130/paddle-release/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl \ - $BOS/whl/ms_swift-0.0.0-py3-none-any.whl \ - $BOS/whl/mcore_bridge-0.0.0-py3-none-any.whl ; do - echo "Downloading $url" - wget -q --no-proxy --no-check-certificate --tries=3 --timeout=60 "$url" - done + if [ "${ALIGNMENT_PADDLEFLEET_MODE:-develop}" = "stack-paired" ]; then + echo "::group::Select stack-paired PaddleFleet (fail-closed SHA)" + python -m pip install uv + bash /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh --dest /workspace + cat /workspace/paddlefleet_alignment_pin_receipt.json + echo "::endgroup::" + echo "::group::Download remaining wheels from BOS" + for url in \ + $BOS/whl/ms_swift-0.0.0-py3-none-any.whl \ + $BOS/whl/mcore_bridge-0.0.0-py3-none-any.whl ; do + echo "Downloading $url" + wget -q --no-proxy --no-check-certificate --tries=3 --timeout=60 "$url" + done + echo "::endgroup::" + else + echo "::group::Download paddlefleet / paddlefleet_ops / ms_swift / mcore-bridge wheels from BOS" + for url in \ + $BOS/PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl \ + $BOS/PaddleFleet/develop/latest/cu130/paddle-release/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl \ + $BOS/whl/ms_swift-0.0.0-py3-none-any.whl \ + $BOS/whl/mcore_bridge-0.0.0-py3-none-any.whl ; do + echo "Downloading $url" + wget -q --no-proxy --no-check-certificate --tries=3 --timeout=60 "$url" + done + echo "::endgroup::" + fi ls -l /workspace/ - echo "::endgroup::" echo "::group::Build megatron-core wheel" cd /workspace/Megatron-LM @@ -139,8 +204,12 @@ jobs: ldconfig source $work_dir/../../../proxy export PROXY_URL="${http_proxy}" - export PADDLEFLEET_WHEEL_PATH="/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" - export PADDLEFLEET_OPS_WHEEL_PATH="/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" + if [ -f /workspace/paddlefleet_alignment_pin.env ]; then + . /workspace/paddlefleet_alignment_pin.env + echo "loaded pin receipt ${PADDLEFLEET_PIN_RECEIPT:-/workspace/paddlefleet_alignment_pin_receipt.json}" + fi + export PADDLEFLEET_WHEEL_PATH="${PADDLEFLEET_WHEEL_PATH:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" + export PADDLEFLEET_OPS_WHEEL_PATH="${PADDLEFLEET_OPS_WHEEL_PATH:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" export MEGATRON_CORE_WHEEL_PATH=/workspace/upload/megatron_core-0.0.0-cp312-cp312-linux_x86_64.whl export MS_SWIFT_WHEEL_PATH=/workspace/ms_swift-0.0.0-py3-none-any.whl export MCORE_BRIDGE_WHEEL_PATH=/workspace/mcore_bridge-0.0.0-py3-none-any.whl @@ -150,7 +219,7 @@ jobs: for whl in "$PADDLEFLEET_WHEEL_PATH" "$PADDLEFLEET_OPS_WHEEL_PATH" \ "$MS_SWIFT_WHEEL_PATH" "$MEGATRON_CORE_WHEEL_PATH" \ "$MCORE_BRIDGE_WHEEL_PATH"; do - [ -f "$whl" ] || { echo "::error:: missing wheel: $whl"; exit 1; } + [ -f "$whl" ] || [ -d "$whl" ] || { echo "::error:: missing wheel: $whl"; exit 1; } echo "using $whl" done python -m pip install uv diff --git a/scripts/select_paddlefleet_alignment_pin.sh b/scripts/select_paddlefleet_alignment_pin.sh new file mode 100755 index 00000000000..71c2131318e --- /dev/null +++ b/scripts/select_paddlefleet_alignment_pin.sh @@ -0,0 +1,638 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Select PaddleFleet source + wheels for alignment_model_accuracy. +# +# Default (ALIGNMENT_PADDLEFLEET_MODE=develop or unset): +# historical CodeSync/develop tarball + BOS develop/latest wheels. +# Cases are not filtered. +# +# Explicit (ALIGNMENT_PADDLEFLEET_MODE=stack-paired): fail-closed. +# Checkout PADDLEFLEET_PIN_SHA; git rev-parse HEAD must equal the pin. +# Artifacts: caller URL+sha256 (from Build Fleet whl / Actions metadata) +# or build from the checked-out tree. Independent digest matches do not +# prove the wheels were produced from that commit — receipt records +# source_commit vs artifact sha256 separately and pairing as unproven +# unless this invocation built the files from the pin. +# git/download/checkout failures still write an error receipt. + +set -euo pipefail + +usage() { + cat <<'EOF' +Usage: select_paddlefleet_alignment_pin.sh [--dest DIR] [--self-test] + +--self-test ignores --dest and uses an offline fixture (no network). + +Env: + ALIGNMENT_PADDLEFLEET_MODE develop (default) | stack-paired + PADDLEFLEET_PIN_SHA required 40-hex commit in stack-paired + PADDLEFLEET_GIT_URL git remote or local repo (stack-paired) + PADDLEFLEET_WHEEL_URL optional explicit wheel (https or local path) + PADDLEFLEET_WHEEL_SHA256 required with WHEEL_URL + PADDLEFLEET_OPS_WHEEL_URL optional explicit ops wheel + PADDLEFLEET_OPS_WHEEL_SHA256 required with OPS URL + PADDLEFLEET_BUILD_CMD optional; default uv build paddlefleet + PADDLEFLEET_BUILD_OPS_CMD optional; default uv build paddlefleet-ops + ALIGNMENT_PADDLEFLEET_DEST output directory (default /workspace) +EOF +} + +MODE="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" +DEST="${ALIGNMENT_PADDLEFLEET_DEST:-/workspace}" +RUN_SELF_TEST=0 +while [[ $# -gt 0 ]]; do + case "$1" in + --dest) + DEST="${2:?--dest requires a path}" + shift 2 + ;; + --self-test) + RUN_SELF_TEST=1 + shift + ;; + -h|--help) + usage + exit 0 + ;; + *) + echo "unknown arg: $1" >&2 + usage + exit 2 + ;; + esac +done + +BOS="${PADDLEFLEET_BOS:-https://paddle-github-action.bj.bcebos.com}" +DEFAULT_TAR_URL="https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar" +DEFAULT_WHL_URL="${BOS}/PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" +DEFAULT_OPS_URL="${BOS}/PaddleFleet/develop/latest/cu130/paddle-release/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" +GIT_URL="${PADDLEFLEET_GIT_URL:-https://github.com/PaddlePaddle/PaddleFleet.git}" +PIN_SHA="${PADDLEFLEET_PIN_SHA:-}" +WHEEL_URL="${PADDLEFLEET_WHEEL_URL:-}" +WHEEL_SHA="${PADDLEFLEET_WHEEL_SHA256:-}" +OPS_URL="${PADDLEFLEET_OPS_WHEEL_URL:-}" +OPS_SHA="${PADDLEFLEET_OPS_WHEEL_SHA256:-}" +BUILD_CMD="${PADDLEFLEET_BUILD_CMD:-}" +BUILD_OPS_CMD="${PADDLEFLEET_BUILD_OPS_CMD:-}" + +ACTUAL_SHA="" +SOURCE_VERIFIED=false +PADDLEFLEET_WHEEL_PATH="" +PADDLEFLEET_OPS_WHEEL_PATH="" +ACTUAL_WHEEL_SHA="" +ACTUAL_OPS_SHA="" +WHEEL_DIGEST_VERIFIED=false +OPS_DIGEST_VERIFIED=false +WHEEL_ORIGIN="" +OPS_ORIGIN="" +WHEEL_BUILT_FROM_COMMIT="" +OPS_BUILT_FROM_COMMIT="" +LOADED_FROM="" +RECEIPT_WRITTEN=0 + +log() { echo "[paddlefleet-pin] $*" >&2; } + +sha256_file() { sha256sum -- "$1" | awk '{print $1}'; } + +unpaired_url() { + case "$1" in + *"/develop/latest/"*|*"CodeSync/develop/"*) return 0 ;; + *) return 1 ;; + esac +} + +pairing_fields() { + local proven=false + local status="unproven" + local reason="wheel/ops digest match does not prove production from source_commit" + if [[ "${MODE}" == "develop" ]]; then + status="unpaired_default" + reason="develop tarball and develop/latest wheels; not a stack pin" + elif [[ "${WHEEL_ORIGIN}" == "build" && "${OPS_ORIGIN}" == "build" \ + && "${WHEEL_BUILT_FROM_COMMIT}" == "${ACTUAL_SHA}" \ + && "${OPS_BUILT_FROM_COMMIT}" == "${ACTUAL_SHA}" \ + && "${SOURCE_VERIFIED}" == "true" ]]; then + status="built_from_checked_out_pin" + reason="this invocation built both artifacts from checked-out source_commit; not a remote-stack proof" + fi + printf '%s\t%s\t%s\n' "${proven}" "${status}" "${reason}" +} + +write_receipt() { + local status="$1" detail="${2:-}" + mkdir -p "${DEST}" + local receipt="${DEST}/paddlefleet_alignment_pin_receipt.json" + local pair + pair="$(pairing_fields)" + local stack_proven pairing_status pairing_reason + stack_proven="${pair%%$'\t'*}" + pair="${pair#*$'\t'}" + pairing_status="${pair%%$'\t'*}" + pairing_reason="${pair#*$'\t'}" + if ! command -v python3 >/dev/null 2>&1; then + printf '{"schema":"paddlefleet-alignment-pin/v1","status":"%s","detail":"%s"}\n' \ + "${status}" "${detail}" >"${receipt}" + RECEIPT_WRITTEN=1 + return 0 + fi + python3 - "${receipt}" "${status}" "${detail}" "${stack_proven}" \ + "${pairing_status}" "${pairing_reason}" <<'PY' +import json, os, sys +from datetime import datetime, timezone +path, status, detail, stack_proven, pairing_status, pairing_reason = sys.argv[1:7] + +def art(name, pth, url, exp, act, digest_ok, origin, built_from): + if not pth: + return None + return { + "name": name, + "path": pth, + "url": url or None, + "expected_sha256": exp or None, + "actual_sha256": act or None, + "digest_verified": digest_ok == "true", + "origin": origin or None, + "built_from_commit": built_from or None, + } + +arts = [a for a in ( + art("paddlefleet", os.environ.get("PADDLEFLEET_WHEEL_PATH", ""), + os.environ.get("WHEEL_URL", ""), os.environ.get("WHEEL_SHA", ""), + os.environ.get("ACTUAL_WHEEL_SHA", ""), os.environ.get("WHEEL_DIGEST_VERIFIED", "false"), + os.environ.get("WHEEL_ORIGIN", ""), os.environ.get("WHEEL_BUILT_FROM_COMMIT", "")), + art("paddlefleet_ops", os.environ.get("PADDLEFLEET_OPS_WHEEL_PATH", ""), + os.environ.get("OPS_URL", ""), os.environ.get("OPS_SHA", ""), + os.environ.get("ACTUAL_OPS_SHA", ""), os.environ.get("OPS_DIGEST_VERIFIED", "false"), + os.environ.get("OPS_ORIGIN", ""), os.environ.get("OPS_BUILT_FROM_COMMIT", "")), +) if a] + +doc = { + "schema": "paddlefleet-alignment-pin/v1", + "status": status, + "detail": detail, + "captured_at": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"), + "mode": os.environ.get("MODE"), + "dest": os.environ.get("DEST"), + "source": { + "git_url": os.environ.get("GIT_URL") or None, + "expected_commit": os.environ.get("PIN_SHA") or None, + "actual_commit": os.environ.get("ACTUAL_SHA") or None, + "commit_verified": os.environ.get("SOURCE_VERIFIED") == "true", + }, + "loaded_from": os.environ.get("LOADED_FROM") or None, + "artifacts": arts, + "pairing": { + "stack_paired_proven": stack_proven == "true", + "status": pairing_status, + "reason": pairing_reason, + }, + "default_urls": { + "source_tar": os.environ.get("DEFAULT_TAR_URL"), + "paddlefleet_wheel": os.environ.get("DEFAULT_WHL_URL"), + "paddlefleet_ops_wheel": os.environ.get("DEFAULT_OPS_URL"), + }, + "cases_preserved": ["MinimaxV2.5_EP2", "GLM45Air_EP2"], + "unpaired_develop_rejected_in_stack_paired": True, +} +open(path, "w", encoding="utf-8").write(json.dumps(doc, indent=2) + "\n") +print("[paddlefleet-pin] receipt", path, file=sys.stderr) +PY + RECEIPT_WRITTEN=1 +} + +export_receipt_env() { + export MODE DEST GIT_URL PIN_SHA ACTUAL_SHA SOURCE_VERIFIED LOADED_FROM + export PADDLEFLEET_WHEEL_PATH PADDLEFLEET_OPS_WHEEL_PATH + export WHEEL_URL WHEEL_SHA ACTUAL_WHEEL_SHA WHEEL_DIGEST_VERIFIED WHEEL_ORIGIN WHEEL_BUILT_FROM_COMMIT + export OPS_URL OPS_SHA ACTUAL_OPS_SHA OPS_DIGEST_VERIFIED OPS_ORIGIN OPS_BUILT_FROM_COMMIT + export DEFAULT_TAR_URL DEFAULT_WHL_URL DEFAULT_OPS_URL +} + +fail() { + trap - ERR + local msg="$1" + log "FAIL: ${msg}" + export_receipt_env + write_receipt "error" "${msg}" + echo "::error:: ${msg}" >&2 + exit 1 +} + +on_err() { + local rc=$? + if [[ "${RECEIPT_WRITTEN}" == 1 || "${RUN_SELF_TEST}" == 1 ]]; then + return "${rc}" + fi + fail "command failed rc=${rc}" +} +trap 'on_err' ERR + +write_envfile() { + cat >"${DEST}/paddlefleet_alignment_pin.env" < ${out}" + mkdir -p "$(dirname "${out}")" + if [[ "${url}" == file://* ]]; then + local src="${url#file://}" + [[ -f "${src}" ]] || fail "download failed, local file missing: ${src}" + cp -- "${src}" "${out}" || fail "download copy failed: ${src}" + return 0 + fi + if [[ "${url}" == /* ]]; then + [[ -f "${url}" ]] || fail "download failed, local file missing: ${url}" + cp -- "${url}" "${out}" || fail "download copy failed: ${url}" + return 0 + fi + if wget -q --no-proxy --no-check-certificate --tries=2 --timeout=15 -O "${out}" "${url}"; then + return 0 + fi + fail "download failed: ${url}" +} + +# Must not run inside $(); fail() has to exit this shell. +require_digest() { + local path="$1" expected="$2" label="$3" actual="$4" + [[ -f "${path}" ]] || fail "missing ${label}: ${path}" + [[ -n "${expected}" ]] || fail "stack-paired missing ${label} sha256" + if [[ "${actual}" != "${expected}" ]]; then + fail "stack-paired ${label} sha256 mismatch expected=${expected} actual=${actual}" + fi +} + +checkout_pin() { + [[ "${PIN_SHA}" =~ ^[0-9a-fA-F]{40}$ ]] || fail "stack-paired requires PADDLEFLEET_PIN_SHA (40 hex), got '${PIN_SHA}'" + PIN_SHA="$(printf '%s' "${PIN_SHA}" | tr 'A-F' 'a-f')" + rm -rf "${DEST}/PaddleFleet" + log "clone ${GIT_URL}" + if ! git clone --quiet "${GIT_URL}" "${DEST}/PaddleFleet" >/dev/null 2>"${DEST}/.git-clone.err"; then + fail "git clone failed: $(tr '\n' ' ' <"${DEST}/.git-clone.err")" + fi + git -C "${DEST}/PaddleFleet" config advice.detachedHead false || true + log "checkout ${PIN_SHA}" + if ! git -C "${DEST}/PaddleFleet" checkout --quiet --force "${PIN_SHA}" >/dev/null 2>"${DEST}/.git-co.err"; then + fail "git checkout failed for ${PIN_SHA}: $(tr '\n' ' ' <"${DEST}/.git-co.err")" + fi + ACTUAL_SHA="$(git -C "${DEST}/PaddleFleet" rev-parse HEAD)" + if [[ "${ACTUAL_SHA}" != "${PIN_SHA}" ]]; then + SOURCE_VERIFIED=false + fail "stack-paired source SHA mismatch expected=${PIN_SHA} actual=${ACTUAL_SHA}" + fi + SOURCE_VERIFIED=true + log "source commit verified ${ACTUAL_SHA}" +} + +# Sets DEST_PATH and DEST_SHA in the caller. Must run in this shell so +# fail() writes the receipt (never wrap this in $()). +acquire_explicit() { + local url="$1" expected="$2" dest_name="$3" label="$4" + unpaired_url "${url}" && fail "stack-paired rejects unpaired ${label} URL: ${url}" + [[ -n "${expected}" ]] || fail "stack-paired ${label} URL requires matching sha256" + download "${url}" "${DEST}/${dest_name}" + DEST_PATH="${DEST}/${dest_name}" + DEST_SHA="$(sha256_file "${DEST_PATH}")" + # Record path/digest before require_digest so a mismatch receipt still has them. + if [[ "${label}" == paddlefleet\ wheel ]]; then + PADDLEFLEET_WHEEL_PATH="${DEST_PATH}" + ACTUAL_WHEEL_SHA="${DEST_SHA}" + WHEEL_ORIGIN="ci_metadata" + else + PADDLEFLEET_OPS_WHEEL_PATH="${DEST_PATH}" + ACTUAL_OPS_SHA="${DEST_SHA}" + OPS_ORIGIN="ci_metadata" + fi + require_digest "${DEST_PATH}" "${expected}" "${label}" "${DEST_SHA}" + LOADED_FROM="${LOADED_FROM:+${LOADED_FROM};}${DEST_PATH} from ${url}" +} + +run_build() { + local cmd="$1" glob="$2" label="$3" + mkdir -p "${DEST}/dist" + log "build ${label}: ${cmd}" + if ! (cd "${DEST}/PaddleFleet" && bash -lc "${cmd}"); then + fail "stack-paired build failed for ${label}" + fi + local built + built="$(ls -1 ${glob} 2>/dev/null | head -n 1 || true)" + [[ -n "${built}" && -f "${built}" ]] || fail "stack-paired build produced no ${label} (glob ${glob})" + local dest_name + dest_name="$(basename "${built}")" + cp -f -- "${built}" "${DEST}/${dest_name}" + DEST_PATH="${DEST}/${dest_name}" + DEST_SHA="$(sha256_file "${DEST_PATH}")" + LOADED_FROM="${LOADED_FROM:+${LOADED_FROM};}built ${DEST_PATH} from ${ACTUAL_SHA}" +} + +acquire_wheel() { + if [[ -n "${WHEEL_URL}" ]]; then + acquire_explicit "${WHEEL_URL}" "${WHEEL_SHA}" "paddlefleet.whl" "paddlefleet wheel" + PADDLEFLEET_WHEEL_PATH="${DEST_PATH}" + ACTUAL_WHEEL_SHA="${DEST_SHA}" + WHEEL_DIGEST_VERIFIED=true + WHEEL_ORIGIN="ci_metadata" + WHEEL_BUILT_FROM_COMMIT="" + else + local cmd="${BUILD_CMD:-uv build --wheel --package paddlefleet --out-dir '${DEST}/dist' --clear}" + run_build "${cmd}" "${DEST}/dist/paddlefleet-*.whl" "paddlefleet wheel" + PADDLEFLEET_WHEEL_PATH="${DEST_PATH}" + ACTUAL_WHEEL_SHA="${DEST_SHA}" + WHEEL_DIGEST_VERIFIED=true + WHEEL_ORIGIN="build" + WHEEL_BUILT_FROM_COMMIT="${ACTUAL_SHA}" + fi +} + +acquire_ops() { + if [[ -n "${OPS_URL}" ]]; then + acquire_explicit "${OPS_URL}" "${OPS_SHA}" "paddlefleet_ops.whl" "paddlefleet_ops wheel" + PADDLEFLEET_OPS_WHEEL_PATH="${DEST_PATH}" + ACTUAL_OPS_SHA="${DEST_SHA}" + OPS_DIGEST_VERIFIED=true + OPS_ORIGIN="ci_metadata" + OPS_BUILT_FROM_COMMIT="" + else + local cmd="${BUILD_OPS_CMD:-uv build --wheel --package paddlefleet-ops --out-dir '${DEST}/dist' --no-build-isolation}" + run_build "${cmd}" "${DEST}/dist/paddlefleet_ops-*.whl" "paddlefleet_ops wheel" + PADDLEFLEET_OPS_WHEEL_PATH="${DEST_PATH}" + ACTUAL_OPS_SHA="${DEST_SHA}" + OPS_DIGEST_VERIFIED=true + OPS_ORIGIN="build" + OPS_BUILT_FROM_COMMIT="${ACTUAL_SHA}" + fi +} + +fetch_default() { + log "mode=develop (historical unpaired CodeSync tarball)" + download "${DEFAULT_TAR_URL}" "${DEST}/PaddleFleet.tar" + rm -rf "${DEST}/PaddleFleet" + tar xf "${DEST}/PaddleFleet.tar" -C "${DEST}" + rm -f "${DEST}/PaddleFleet.tar" + if [[ -d "${DEST}/PaddleFleet/.git" ]]; then + git -C "${DEST}/PaddleFleet" pull || log "git pull skipped" + ACTUAL_SHA="$(git -C "${DEST}/PaddleFleet" rev-parse HEAD 2>/dev/null || true)" + fi + download "${DEFAULT_WHL_URL}" "${DEST}/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" + download "${DEFAULT_OPS_URL}" "${DEST}/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" + PADDLEFLEET_WHEEL_PATH="${DEST}/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" + PADDLEFLEET_OPS_WHEEL_PATH="${DEST}/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" + WHEEL_URL="${DEFAULT_WHL_URL}" + OPS_URL="${DEFAULT_OPS_URL}" + ACTUAL_WHEEL_SHA="$(sha256_file "${PADDLEFLEET_WHEEL_PATH}")" + ACTUAL_OPS_SHA="$(sha256_file "${PADDLEFLEET_OPS_WHEEL_PATH}")" + WHEEL_ORIGIN="develop_latest" + OPS_ORIGIN="develop_latest" + SOURCE_VERIFIED=false + WHEEL_DIGEST_VERIFIED=false + OPS_DIGEST_VERIFIED=false + LOADED_FROM="develop_latest ${DEFAULT_WHL_URL} ${DEFAULT_OPS_URL}" + export_receipt_env + write_envfile + write_receipt "ok" "develop tarball and develop/latest wheels; unpaired with a stack pin" +} + +fetch_stack_paired() { + log "mode=stack-paired" + checkout_pin + acquire_wheel + acquire_ops + export_receipt_env + write_envfile + local pair pairing_status + pair="$(pairing_fields)" + pair="${pair#*$'\t'}" + pairing_status="${pair%%$'\t'*}" + write_receipt "ok" "source_commit checked out; pairing.status=${pairing_status}; stack_paired_proven=false" +} + +install_offline_stubs() { + local bin="$1" + mkdir -p "${bin}" + cat >"${bin}/wget" <<'WGET' +#!/usr/bin/env bash +out="" +url="" +while [[ $# -gt 0 ]]; do + case "$1" in + -O) out="$2"; shift 2 ;; + --*) shift ;; + *) url="$1"; shift ;; + esac +done +if [[ -z "${url}" || -z "${out}" ]]; then + echo "wget-stub: missing url/out" >&2 + exit 1 +fi +if [[ "${url}" == http://* || "${url}" == https://* ]]; then + echo "wget-stub: blocked network ${url}" >&2 + exit 1 +fi +src="${url#file://}" +if [[ -f "${src}" ]]; then + cp -- "${src}" "${out}" + exit 0 +fi +echo "wget-stub: not a local file ${url}" >&2 +exit 1 +WGET + chmod +x "${bin}/wget" +} + +run_self_test() { + trap - ERR + local root script + root="$(mktemp -d)" + script="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)/$(basename -- "${BASH_SOURCE[0]}")" + trap 'rm -rf "${root}"' RETURN + install_offline_stubs "${root}/bin" + export PATH="${root}/bin:${PATH}" + + git init -q "${root}/upstream" + git -C "${root}/upstream" config user.email test@example.com + git -C "${root}/upstream" config user.name test + echo source-a >"${root}/upstream/README" + git -C "${root}/upstream" add README + git -C "${root}/upstream" commit -q -m a + local sha_a sha_b + sha_a="$(git -C "${root}/upstream" rev-parse HEAD)" + echo source-b >"${root}/upstream/README" + git -C "${root}/upstream" add README + git -C "${root}/upstream" commit -q -m b + sha_b="$(git -C "${root}/upstream" rev-parse HEAD)" + + mkdir -p "${root}/art" + echo py-body >"${root}/art/py.whl" + echo ops-body >"${root}/art/ops.whl" + local py_sha ops_sha + py_sha="$(sha256_file "${root}/art/py.whl")" + ops_sha="$(sha256_file "${root}/art/ops.whl")" + + expect_fail() { + local dest="$1" + local needle="$2" + shift 2 + mkdir -p "${dest}" + if "$@"; then + echo "self-test FAIL: expected failure (${needle})" >&2 + exit 1 + fi + local rec="${dest}/paddlefleet_alignment_pin_receipt.json" + [[ -f "${rec}" ]] || { echo "self-test FAIL: missing error receipt ${rec}" >&2; exit 1; } + grep -q '"status": "error"' "${rec}" + grep -q "${needle}" "${rec}" + echo "[self-test] fail-closed ${dest}: ${needle}" + } + + local run + run() { env PATH="${root}/bin:${PATH}" "$@"; } + + expect_fail "${root}/m1" "PADDLEFLEET_PIN_SHA" \ + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA= \ + bash "${script}" --dest "${root}/m1" + + expect_fail "${root}/m2" "rejects unpaired" \ + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + PADDLEFLEET_WHEEL_URL="${DEFAULT_WHL_URL}" \ + PADDLEFLEET_WHEEL_SHA256="${py_sha}" \ + PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ + PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ + bash "${script}" --dest "${root}/m2" + + expect_fail "${root}/m3" "git checkout failed" \ + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="0000000000000000000000000000000000000000" \ + PADDLEFLEET_GIT_URL="${root}/upstream" \ + bash "${script}" --dest "${root}/m3" + + # Real checksum mismatch after a successful local copy (not a wget miss). + expect_fail "${root}/m4" "sha256 mismatch" \ + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + PADDLEFLEET_WHEEL_URL="${root}/art/py.whl" \ + PADDLEFLEET_WHEEL_SHA256="deadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeef" \ + PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ + PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ + bash "${script}" --dest "${root}/m4" + python3 - "${root}/m4/paddlefleet_alignment_pin_receipt.json" "${py_sha}" <<'PY' +import json, sys +doc = json.load(open(sys.argv[1])) +assert doc["status"] == "error" +wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") +assert wheel["actual_sha256"] == sys.argv[2] +assert wheel["digest_verified"] is False +assert wheel["actual_sha256"] != (wheel.get("expected_sha256") or "") +print("m4 checksum-mismatch receipt has actual digest, not a download miss") +PY + + expect_fail "${root}/m5" "produced no paddlefleet" \ + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + PADDLEFLEET_BUILD_CMD="mkdir -p '${root}/m5/dist'" \ + PADDLEFLEET_BUILD_OPS_CMD="true" \ + bash "${script}" --dest "${root}/m5" + + expect_fail "${root}/m6" "download failed" \ + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + PADDLEFLEET_WHEEL_URL="https://example.invalid/paddlefleet.whl" \ + PADDLEFLEET_WHEEL_SHA256="${py_sha}" \ + PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ + PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ + bash "${script}" --dest "${root}/m6" + + expect_fail "${root}/m7" "git clone failed" \ + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" \ + PADDLEFLEET_GIT_URL="${root}/no-such-remote" \ + bash "${script}" --dest "${root}/m7" + + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + PADDLEFLEET_WHEEL_URL="${root}/art/py.whl" \ + PADDLEFLEET_WHEEL_SHA256="${py_sha}" \ + PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ + PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ + bash "${script}" --dest "${root}/ok-url" + python3 - "${root}/ok-url/paddlefleet_alignment_pin_receipt.json" "${sha_b}" "${py_sha}" <<'PY' +import json, sys +doc = json.load(open(sys.argv[1])) +sha_b, py_sha = sys.argv[2], sys.argv[3] +assert doc["status"] == "ok" +assert doc["source"]["actual_commit"] == sha_b +assert doc["source"]["commit_verified"] is True +wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") +assert wheel["actual_sha256"] == py_sha +assert wheel["actual_sha256"] != sha_b +assert wheel["digest_verified"] is True +assert wheel.get("built_from_commit") in (None, "") +assert "verified" not in wheel +assert doc["pairing"]["stack_paired_proven"] is False +assert doc["pairing"]["status"] == "unproven" +assert "MinimaxV2.5_EP2" in doc["cases_preserved"] +assert "GLM45Air_EP2" in doc["cases_preserved"] +print("ok-url receipt fields checked") +PY + [[ "$(git -C "${root}/ok-url/PaddleFleet" rev-parse HEAD)" == "${sha_b}" ]] + + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + PADDLEFLEET_BUILD_CMD="mkdir -p '${root}/ok-build/dist' && cp '${root}/art/py.whl' '${root}/ok-build/dist/paddlefleet-0.0.0-py3-none-any.whl'" \ + PADDLEFLEET_BUILD_OPS_CMD="mkdir -p '${root}/ok-build/dist' && cp '${root}/art/ops.whl' '${root}/ok-build/dist/paddlefleet_ops-0.0.0-py3-none-any.whl'" \ + bash "${script}" --dest "${root}/ok-build" + python3 - "${root}/ok-build/paddlefleet_alignment_pin_receipt.json" "${sha_b}" "${py_sha}" <<'PY' +import json, sys +doc = json.load(open(sys.argv[1])) +sha_b, py_sha = sys.argv[2], sys.argv[3] +assert doc["status"] == "ok" +wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") +assert wheel["actual_sha256"] == py_sha +assert wheel["actual_sha256"] != sha_b +assert wheel["origin"] == "build" +assert wheel["built_from_commit"] == sha_b +assert doc["source"]["actual_commit"] == sha_b +assert doc["pairing"]["stack_paired_proven"] is False +assert doc["pairing"]["status"] == "built_from_checked_out_pin" +print("ok-build receipt fields checked") +PY + grep -q "PADDLEFLEET_SOURCE_COMMIT=${sha_b}" "${root}/ok-build/paddlefleet_alignment_pin.env" + + grep -q 'CodeSync/develop/PaddleFleet.tar' "${script}" + grep -q 'PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl' "${script}" + + echo "select_paddlefleet_alignment_pin self-test OK" +} + +if [[ "${RUN_SELF_TEST}" == 1 ]]; then + run_self_test + exit 0 +fi + +mkdir -p "${DEST}" +case "${MODE}" in + stack-paired) fetch_stack_paired ;; + develop) fetch_default ;; + *) fail "unknown ALIGNMENT_PADDLEFLEET_MODE=${MODE} (develop|stack-paired)" ;; +esac From 7694fca818c2ceb88bed84c3b3e0e809e2ad7dfb Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Sat, 5 Sep 2026 21:03:30 +0800 Subject: [PATCH 14/27] Pass stack-paired PaddleFleet source paths across docker exec steps workflow_dispatch left PR_ID empty so checkout stayed on the BOS main tarball and the next exec defaulted to paddlefleet-0.0.0.whl before setup_venvs. Checkout COMMIT_ID, export source trees from the pin selector, and consume that env in the alignment step without the 0.0.0 fallback. --- .../workflows/alignment_model_accuracy.yml | 35 +++-- scripts/consume_paddlefleet_alignment_pin.sh | 143 ++++++++++++++++++ scripts/select_paddlefleet_alignment_pin.sh | 65 +++++++- scripts/test_paddlefleet_pin_handoff.sh | 78 ++++++++++ 4 files changed, 306 insertions(+), 15 deletions(-) create mode 100755 scripts/consume_paddlefleet_alignment_pin.sh create mode 100755 scripts/test_paddlefleet_pin_handoff.sh diff --git a/.github/workflows/alignment_model_accuracy.yml b/.github/workflows/alignment_model_accuracy.yml index 02d2b3e7264..82d7b8ea184 100644 --- a/.github/workflows/alignment_model_accuracy.yml +++ b/.github/workflows/alignment_model_accuracy.yml @@ -122,17 +122,25 @@ jobs: git config --global --add safe.directory /workspace/Megatron-LM git pull git submodule update --init --recursive --force + git remote add upstream https://github.com/PFCCLab/Megatron-LM.git || true if [ -n "$PR_ID" ] && [ "$PR_ID" != "0" ]; then git fetch origin pull/${PR_ID}/head - git checkout -b PR_${PR_ID} FETCH_HEAD - git remote add upstream https://github.com/PFCCLab/Megatron-LM.git + git checkout -B PR_${PR_ID} FETCH_HEAD echo "Checking out ${BRANCH}..." - git fetch upstream ${BRANCH}:${BRANCH} - git merge ${BRANCH} --no-edit + git fetch upstream ${BRANCH}:${BRANCH} || true + git merge ${BRANCH} --no-edit || true git diff --numstat ${BRANCH} -- | awk "{print \$NF}" + elif [ -n "$COMMIT_ID" ]; then + echo "workflow_dispatch: checkout COMMIT_ID=${COMMIT_ID} (not BOS main)" + git fetch --all --tags || true + git fetch origin "$COMMIT_ID" || git fetch upstream "$COMMIT_ID" || true + git checkout --force "$COMMIT_ID" + git submodule update --init --recursive --force else - echo "Not in a pull_request event. Skipping PR-specific operations." + echo "Not in a pull_request event and COMMIT_ID empty. Leaving tarball HEAD." fi + echo "checked_out_head=$(git rev-parse HEAD)" + test -z "$COMMIT_ID" || test "$(git rev-parse HEAD)" = "$COMMIT_ID" git log --pretty=oneline -10 ' - name: Change python version @@ -163,8 +171,11 @@ jobs: if [ "${ALIGNMENT_PADDLEFLEET_MODE:-develop}" = "stack-paired" ]; then echo "::group::Select stack-paired PaddleFleet (fail-closed SHA)" python -m pip install uv + test -x /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh \ + || { echo "::error:: selector missing; checkout did not land COMMIT_ID"; exit 1; } bash /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh --dest /workspace cat /workspace/paddlefleet_alignment_pin_receipt.json + test -f /workspace/paddlefleet_alignment_pin.env echo "::endgroup::" echo "::group::Download remaining wheels from BOS" for url in \ @@ -204,12 +215,12 @@ jobs: ldconfig source $work_dir/../../../proxy export PROXY_URL="${http_proxy}" - if [ -f /workspace/paddlefleet_alignment_pin.env ]; then - . /workspace/paddlefleet_alignment_pin.env - echo "loaded pin receipt ${PADDLEFLEET_PIN_RECEIPT:-/workspace/paddlefleet_alignment_pin_receipt.json}" - fi - export PADDLEFLEET_WHEEL_PATH="${PADDLEFLEET_WHEEL_PATH:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" - export PADDLEFLEET_OPS_WHEEL_PATH="${PADDLEFLEET_OPS_WHEEL_PATH:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" + bash /workspace/Megatron-LM/scripts/consume_paddlefleet_alignment_pin.sh \ + --env /workspace/paddlefleet_alignment_pin.env \ + --out /workspace/paddlefleet_alignment_pin.consumed.env + set -a + . /workspace/paddlefleet_alignment_pin.consumed.env + set +a export MEGATRON_CORE_WHEEL_PATH=/workspace/upload/megatron_core-0.0.0-cp312-cp312-linux_x86_64.whl export MS_SWIFT_WHEEL_PATH=/workspace/ms_swift-0.0.0-py3-none-any.whl export MCORE_BRIDGE_WHEEL_PATH=/workspace/mcore_bridge-0.0.0-py3-none-any.whl @@ -219,7 +230,7 @@ jobs: for whl in "$PADDLEFLEET_WHEEL_PATH" "$PADDLEFLEET_OPS_WHEEL_PATH" \ "$MS_SWIFT_WHEEL_PATH" "$MEGATRON_CORE_WHEEL_PATH" \ "$MCORE_BRIDGE_WHEEL_PATH"; do - [ -f "$whl" ] || [ -d "$whl" ] || { echo "::error:: missing wheel: $whl"; exit 1; } + [ -e "$whl" ] || { echo "::error:: missing wheel: $whl"; exit 1; } echo "using $whl" done python -m pip install uv diff --git a/scripts/consume_paddlefleet_alignment_pin.sh b/scripts/consume_paddlefleet_alignment_pin.sh new file mode 100755 index 00000000000..d7be1857c40 --- /dev/null +++ b/scripts/consume_paddlefleet_alignment_pin.sh @@ -0,0 +1,143 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# Consume selector output in a later docker exec. Source-mode paths must +# survive the step boundary; stack-paired must not fall back to the +# hardcoded 0.0.0 wheel filenames used by unpaired develop. + +set -euo pipefail + +usage() { + cat <<'EOF' +Usage: consume_paddlefleet_alignment_pin.sh [--env FILE] [--out FILE] [--self-test] + +Reads paddlefleet_alignment_pin.env written by select_paddlefleet_alignment_pin.sh +and writes a consumed env file for setup_venvs.sh. Stack-paired refuses the +0.0.0 develop filename and requires existing file-or-directory paths. +EOF +} + +ENVFILE="${PADDLEFLEET_PIN_ENV:-/workspace/paddlefleet_alignment_pin.env}" +OUTFILE="${PADDLEFLEET_CONSUMED_ENV:-/workspace/paddlefleet_alignment_pin.consumed.env}" +RUN_SELF_TEST=0 +while [[ $# -gt 0 ]]; do + case "$1" in + --env) ENVFILE="${2:?}"; shift 2 ;; + --out) OUTFILE="${2:?}"; shift 2 ;; + --self-test) RUN_SELF_TEST=1; shift ;; + -h|--help) usage; exit 0 ;; + *) echo "unknown arg: $1" >&2; usage; exit 2 ;; + esac +done + +hardcoded_develop_wheel() { + case "$1" in + */paddlefleet-0.0.0-py3-none-linux_x86_64.whl) return 0 ;; + */paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl) return 0 ;; + *) return 1 ;; + esac +} + +write_consumed() { + mkdir -p "$(dirname "${OUTFILE}")" + cat >"${OUTFILE}" <&2 + echo "[paddlefleet-pin-consume] wheel=${PADDLEFLEET_WHEEL_PATH} ops=${PADDLEFLEET_OPS_WHEEL_PATH} mode=${ALIGNMENT_PADDLEFLEET_MODE} origin=${PADDLEFLEET_WHEEL_ORIGIN:-}" >&2 +} + +consume() { + local mode="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" + if [[ -f "${ENVFILE}" ]]; then + # shellcheck disable=SC1090 + set -a + # shellcheck disable=SC1090 + . "${ENVFILE}" + set +a + echo "[paddlefleet-pin-consume] loaded ${ENVFILE} receipt=${PADDLEFLEET_PIN_RECEIPT:-}" >&2 + mode="${ALIGNMENT_PADDLEFLEET_MODE:-${mode}}" + fi + ALIGNMENT_PADDLEFLEET_MODE="${mode}" + + if [[ "${mode}" == "stack-paired" ]]; then + [[ -f "${ENVFILE}" ]] || { + echo "::error:: stack-paired missing ${ENVFILE}; selector export did not cross docker exec" >&2 + exit 1 + } + [[ -n "${PADDLEFLEET_WHEEL_PATH:-}" && -n "${PADDLEFLEET_OPS_WHEEL_PATH:-}" ]] || { + echo "::error:: stack-paired env missing PADDLEFLEET_WHEEL_PATH or OPS path" >&2 + exit 1 + } + if hardcoded_develop_wheel "${PADDLEFLEET_WHEEL_PATH}" || hardcoded_develop_wheel "${PADDLEFLEET_OPS_WHEEL_PATH}"; then + echo "::error:: stack-paired refused hardcoded 0.0.0 wheel fallback: ${PADDLEFLEET_WHEEL_PATH} ${PADDLEFLEET_OPS_WHEEL_PATH}" >&2 + exit 1 + fi + else + PADDLEFLEET_WHEEL_PATH="${PADDLEFLEET_WHEEL_PATH:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" + PADDLEFLEET_OPS_WHEEL_PATH="${PADDLEFLEET_OPS_WHEEL_PATH:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" + fi + + if [[ ! -e "${PADDLEFLEET_WHEEL_PATH}" ]]; then + echo "::error:: missing paddlefleet path: ${PADDLEFLEET_WHEEL_PATH}" >&2 + exit 1 + fi + if [[ ! -e "${PADDLEFLEET_OPS_WHEEL_PATH}" ]]; then + echo "::error:: missing paddlefleet_ops path: ${PADDLEFLEET_OPS_WHEEL_PATH}" >&2 + exit 1 + fi + write_consumed +} + +run_self_test() { + local root self + self="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)/$(basename -- "${BASH_SOURCE[0]}")" + root="$(mktemp -d)" + trap 'rm -rf "${root}"' RETURN + mkdir -p "${root}/PaddleFleet/packages/paddlefleet_ops" + echo tree >"${root}/PaddleFleet/pyproject.toml" + echo ops >"${root}/PaddleFleet/packages/paddlefleet_ops/pyproject.toml" + + cat >"${root}/pin.env" <"${root}/missing.env" <&2 + exit 1 + fi + echo "consume_paddlefleet_alignment_pin self-test OK" +} + +if [[ "${RUN_SELF_TEST}" == 1 ]]; then + run_self_test + exit 0 +fi +consume diff --git a/scripts/select_paddlefleet_alignment_pin.sh b/scripts/select_paddlefleet_alignment_pin.sh index 71c2131318e..2fdc7b2d069 100755 --- a/scripts/select_paddlefleet_alignment_pin.sh +++ b/scripts/select_paddlefleet_alignment_pin.sh @@ -121,6 +121,10 @@ pairing_fields() { if [[ "${MODE}" == "develop" ]]; then status="unpaired_default" reason="develop tarball and develop/latest wheels; not a stack pin" + elif [[ "${WHEEL_ORIGIN}" == "source_tree" && "${OPS_ORIGIN}" == "source_tree" \ + && "${SOURCE_VERIFIED}" == "true" ]]; then + status="source_tree_from_checked_out_pin" + reason="this invocation exported checked-out source trees; not a wheel digest proof" elif [[ "${WHEEL_ORIGIN}" == "build" && "${OPS_ORIGIN}" == "build" \ && "${WHEEL_BUILT_FROM_COMMIT}" == "${ACTUAL_SHA}" \ && "${OPS_BUILT_FROM_COMMIT}" == "${ACTUAL_SHA}" \ @@ -247,6 +251,9 @@ PADDLEFLEET_OPS_WHEEL_PATH=${PADDLEFLEET_OPS_WHEEL_PATH} ALIGNMENT_PADDLEFLEET_MODE=${MODE} PADDLEFLEET_PIN_RECEIPT=${DEST}/paddlefleet_alignment_pin_receipt.json PADDLEFLEET_SOURCE_COMMIT=${ACTUAL_SHA} +PADDLEFLEET_PIN_SHA=${PIN_SHA} +PADDLEFLEET_WHEEL_ORIGIN=${WHEEL_ORIGIN} +PADDLEFLEET_OPS_ORIGIN=${OPS_ORIGIN} EOF } @@ -413,11 +420,35 @@ fetch_default() { write_receipt "ok" "develop tarball and develop/latest wheels; unpaired with a stack pin" } +acquire_source_tree() { + local src="${DEST}/PaddleFleet" + local ops="${src}/packages/paddlefleet_ops" + [[ -d "${src}" ]] || fail "stack-paired source tree missing: ${src}" + [[ -f "${src}/pyproject.toml" ]] || fail "stack-paired source tree missing pyproject.toml: ${src}" + [[ -d "${ops}" ]] || fail "stack-paired ops source tree missing: ${ops}" + PADDLEFLEET_WHEEL_PATH="${src}" + PADDLEFLEET_OPS_WHEEL_PATH="${ops}" + WHEEL_ORIGIN="source_tree" + OPS_ORIGIN="source_tree" + WHEEL_BUILT_FROM_COMMIT="${ACTUAL_SHA}" + OPS_BUILT_FROM_COMMIT="${ACTUAL_SHA}" + WHEEL_DIGEST_VERIFIED=false + OPS_DIGEST_VERIFIED=false + LOADED_FROM="source_tree ${src} ${ops} from ${ACTUAL_SHA}" + log "source-tree paths ${src} ${ops}" +} + fetch_stack_paired() { log "mode=stack-paired" checkout_pin - acquire_wheel - acquire_ops + if [[ -n "${WHEEL_URL}" || -n "${OPS_URL}" || -n "${BUILD_CMD}" || -n "${BUILD_OPS_CMD}" ]]; then + acquire_wheel + acquire_ops + else + # No CI wheel URL and no explicit build: export the checked-out trees. + # A later docker exec must source paddlefleet_alignment_pin.env. + acquire_source_tree + fi export_receipt_env write_envfile local pair pairing_status @@ -473,7 +504,10 @@ run_self_test() { git -C "${root}/upstream" config user.email test@example.com git -C "${root}/upstream" config user.name test echo source-a >"${root}/upstream/README" - git -C "${root}/upstream" add README + mkdir -p "${root}/upstream/packages/paddlefleet_ops" + printf '%s\n' '[project]' 'name = "paddlefleet"' >"${root}/upstream/pyproject.toml" + printf '%s\n' '[project]' 'name = "paddlefleet-ops"' >"${root}/upstream/packages/paddlefleet_ops/pyproject.toml" + git -C "${root}/upstream" add README pyproject.toml packages git -C "${root}/upstream" commit -q -m a local sha_a sha_b sha_a="$(git -C "${root}/upstream" rev-parse HEAD)" @@ -619,6 +653,31 @@ print("ok-build receipt fields checked") PY grep -q "PADDLEFLEET_SOURCE_COMMIT=${sha_b}" "${root}/ok-build/paddlefleet_alignment_pin.env" + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + bash "${script}" --dest "${root}/ok-source" + python3 - "${root}/ok-source/paddlefleet_alignment_pin_receipt.json" "${sha_b}" "${root}/ok-source" <<'PY' +import json, sys +doc = json.load(open(sys.argv[1])) +sha_b, dest = sys.argv[2], sys.argv[3] +assert doc["status"] == "ok" +assert doc["source"]["actual_commit"] == sha_b +assert doc["source"]["commit_verified"] is True +wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") +ops = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet_ops") +assert wheel["path"] == f"{dest}/PaddleFleet" +assert ops["path"] == f"{dest}/PaddleFleet/packages/paddlefleet_ops" +assert wheel["origin"] == "source_tree" +assert ops["origin"] == "source_tree" +assert doc["pairing"]["status"] == "source_tree_from_checked_out_pin" +assert doc["pairing"]["stack_paired_proven"] is False +print("ok-source receipt fields checked") +PY + grep -q "PADDLEFLEET_WHEEL_PATH=${root}/ok-source/PaddleFleet$" "${root}/ok-source/paddlefleet_alignment_pin.env" + grep -q "PADDLEFLEET_OPS_WHEEL_PATH=${root}/ok-source/PaddleFleet/packages/paddlefleet_ops$" "${root}/ok-source/paddlefleet_alignment_pin.env" + grep -q "PADDLEFLEET_WHEEL_ORIGIN=source_tree" "${root}/ok-source/paddlefleet_alignment_pin.env" + grep -q "PADDLEFLEET_SOURCE_COMMIT=${sha_b}" "${root}/ok-source/paddlefleet_alignment_pin.env" + grep -q 'CodeSync/develop/PaddleFleet.tar' "${script}" grep -q 'PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl' "${script}" diff --git a/scripts/test_paddlefleet_pin_handoff.sh b/scripts/test_paddlefleet_pin_handoff.sh new file mode 100755 index 00000000000..948dcbe3188 --- /dev/null +++ b/scripts/test_paddlefleet_pin_handoff.sh @@ -0,0 +1,78 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# End-to-end: selector source-mode -> new-shell consume -> setup_venvs consumer. +# This is the docker-exec step boundary. Isolated selector --self-test is not enough. + +set -euo pipefail + +ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" +SELECTOR="${ROOT}/select_paddlefleet_alignment_pin.sh" +CONSUME="${ROOT}/consume_paddlefleet_alignment_pin.sh" + +tmp="$(mktemp -d)" +trap 'rm -rf "${tmp}"' EXIT + +git init -q "${tmp}/upstream" +git -C "${tmp}/upstream" config user.email test@example.com +git -C "${tmp}/upstream" config user.name test +mkdir -p "${tmp}/upstream/packages/paddlefleet_ops" +printf '%s\n' '[project]' 'name = "paddlefleet"' >"${tmp}/upstream/pyproject.toml" +printf '%s\n' '[project]' 'name = "paddlefleet-ops"' >"${tmp}/upstream/packages/paddlefleet_ops/pyproject.toml" +echo src >"${tmp}/upstream/README" +git -C "${tmp}/upstream" add README pyproject.toml packages +git -C "${tmp}/upstream" commit -q -m pin +PIN="$(git -C "${tmp}/upstream" rev-parse HEAD)" + +# Step A: Get Whl equivalent (selector). +ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${PIN}" \ + PADDLEFLEET_GIT_URL="${tmp}/upstream" \ + bash "${SELECTOR}" --dest "${tmp}/ws" + +test -f "${tmp}/ws/paddlefleet_alignment_pin.env" +test -d "${tmp}/ws/PaddleFleet" + +# Step B: new docker exec — drop selector shell state, keep only files. +unset PADDLEFLEET_WHEEL_PATH PADDLEFLEET_OPS_WHEEL_PATH PADDLEFLEET_SOURCE_COMMIT || true +export ALIGNMENT_PADDLEFLEET_MODE=stack-paired +bash "${CONSUME}" --env "${tmp}/ws/paddlefleet_alignment_pin.env" --out "${tmp}/ws/consumed.env" + +# Step C: setup_venvs consumer. Fail if it would install the 0.0.0 develop wheel. +stub_setup="${tmp}/setup_venvs.sh" +cat >"${stub_setup}" <<'STUB' +#!/usr/bin/env bash +set -euo pipefail +# Mirrors setup_venvs.sh: it only sees exported PADDLEFLEET_WHEEL_PATH. +PADDLEFLEET_WHEEL="${PADDLEFLEET_WHEEL_PATH:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" +PADDLEFLEET_OPS_WHEEL="${PADDLEFLEET_OPS_WHEEL_PATH:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" +case "${PADDLEFLEET_WHEEL}" in + */paddlefleet-0.0.0-py3-none-linux_x86_64.whl) + echo "STUB_SETUP used hardcoded 0.0.0 wheel" >&2 + exit 1 + ;; +esac +[[ -d "${PADDLEFLEET_WHEEL}" || -f "${PADDLEFLEET_WHEEL}" ]] || { echo "missing ${PADDLEFLEET_WHEEL}" >&2; exit 1; } +[[ -d "${PADDLEFLEET_OPS_WHEEL}" || -f "${PADDLEFLEET_OPS_WHEEL}" ]] || { echo "missing ${PADDLEFLEET_OPS_WHEEL}" >&2; exit 1; } +echo "STUB_SETUP paddlefleet=${PADDLEFLEET_WHEEL}" +echo "STUB_SETUP ops=${PADDLEFLEET_OPS_WHEEL}" +STUB +chmod +x "${stub_setup}" + +set -a +# shellcheck disable=SC1090 +. "${tmp}/ws/consumed.env" +set +a +bash "${stub_setup}" | tee "${tmp}/setup.out" +grep -q "STUB_SETUP paddlefleet=${tmp}/ws/PaddleFleet" "${tmp}/setup.out" +grep -q "packages/paddlefleet_ops" "${tmp}/setup.out" + +# Negative: alignment step that ignores consumed.env and defaults 0.0.0 must fail the stub. +if PADDLEFLEET_WHEEL_PATH=/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl \ + PADDLEFLEET_OPS_WHEEL_PATH=/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl \ + bash "${stub_setup}"; then + echo "handoff FAIL: stub accepted 0.0.0 default" >&2 + exit 1 +fi + +echo "paddlefleet pin handoff OK pin=${PIN}" From 1164e14111f2730e9a2bb1d692cf998e77d15e84 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Sat, 5 Sep 2026 21:08:56 +0800 Subject: [PATCH 15/27] Keep requested stack-paired mode/pin when consuming selector env Do not let a leftover develop env overwrite the caller. Refuse unproven develop/latest fallback, not a verified 0.0.0 artifact name. Handoff fixture remains a path-pass, not a uv install. --- .../workflows/alignment_model_accuracy.yml | 2 + scripts/consume_paddlefleet_alignment_pin.sh | 294 +++++++++++++++--- scripts/select_paddlefleet_alignment_pin.sh | 3 + scripts/test_paddlefleet_pin_handoff.sh | 62 ++-- 4 files changed, 283 insertions(+), 78 deletions(-) diff --git a/.github/workflows/alignment_model_accuracy.yml b/.github/workflows/alignment_model_accuracy.yml index 82d7b8ea184..b03f4f7b4ff 100644 --- a/.github/workflows/alignment_model_accuracy.yml +++ b/.github/workflows/alignment_model_accuracy.yml @@ -215,6 +215,8 @@ jobs: ldconfig source $work_dir/../../../proxy export PROXY_URL="${http_proxy}" + ALIGNMENT_PADDLEFLEET_MODE="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" \ + PADDLEFLEET_PIN_SHA="${PADDLEFLEET_PIN_SHA:-}" \ bash /workspace/Megatron-LM/scripts/consume_paddlefleet_alignment_pin.sh \ --env /workspace/paddlefleet_alignment_pin.env \ --out /workspace/paddlefleet_alignment_pin.consumed.env diff --git a/scripts/consume_paddlefleet_alignment_pin.sh b/scripts/consume_paddlefleet_alignment_pin.sh index d7be1857c40..f9181d03ba4 100755 --- a/scripts/consume_paddlefleet_alignment_pin.sh +++ b/scripts/consume_paddlefleet_alignment_pin.sh @@ -1,9 +1,11 @@ #!/usr/bin/env bash # Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. # -# Consume selector output in a later docker exec. Source-mode paths must -# survive the step boundary; stack-paired must not fall back to the -# hardcoded 0.0.0 wheel filenames used by unpaired develop. +# Consume selector output in a later docker exec. The caller mode/pin stay +# authoritative: a leftover develop env must not silently downgrade +# stack-paired. A 0.0.0 filename is allowed when the receipt proves an +# explicit URL+sha256 (or a source tree / in-invocation build). Unproven +# develop/latest fallback is refused. set -euo pipefail @@ -11,9 +13,11 @@ usage() { cat <<'EOF' Usage: consume_paddlefleet_alignment_pin.sh [--env FILE] [--out FILE] [--self-test] -Reads paddlefleet_alignment_pin.env written by select_paddlefleet_alignment_pin.sh -and writes a consumed env file for setup_venvs.sh. Stack-paired refuses the -0.0.0 develop filename and requires existing file-or-directory paths. +Reads paddlefleet_alignment_pin.env + receipt from select_paddlefleet_alignment_pin.sh +and writes a consumed env file for setup_venvs.sh. + +Caller ALIGNMENT_PADDLEFLEET_MODE / PADDLEFLEET_PIN_SHA are the request. +The env file must match that request; it does not override them. EOF } @@ -30,12 +34,101 @@ while [[ $# -gt 0 ]]; do esac done -hardcoded_develop_wheel() { - case "$1" in - */paddlefleet-0.0.0-py3-none-linux_x86_64.whl) return 0 ;; - */paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl) return 0 ;; - *) return 1 ;; +fail() { + echo "::error:: $*" >&2 + exit 1 +} + +parse_envfile() { + FILE_MODE="" + FILE_PIN="" + FILE_SOURCE="" + FILE_RECEIPT="" + FILE_WHEEL="" + FILE_OPS="" + FILE_WHEEL_ORIGIN="" + FILE_OPS_ORIGIN="" + FILE_WHEEL_DIGEST="" + FILE_OPS_DIGEST="" + [[ -f "${ENVFILE}" ]] || return 0 + local line k v + while IFS= read -r line || [[ -n "${line}" ]]; do + [[ -z "${line}" || "${line}" == \#* ]] && continue + k="${line%%=*}" + v="${line#*=}" + case "${k}" in + ALIGNMENT_PADDLEFLEET_MODE) FILE_MODE="${v}" ;; + PADDLEFLEET_PIN_SHA) FILE_PIN="${v}" ;; + PADDLEFLEET_SOURCE_COMMIT) FILE_SOURCE="${v}" ;; + PADDLEFLEET_PIN_RECEIPT) FILE_RECEIPT="${v}" ;; + PADDLEFLEET_WHEEL_PATH) FILE_WHEEL="${v}" ;; + PADDLEFLEET_OPS_WHEEL_PATH) FILE_OPS="${v}" ;; + PADDLEFLEET_WHEEL_ORIGIN) FILE_WHEEL_ORIGIN="${v}" ;; + PADDLEFLEET_OPS_ORIGIN) FILE_OPS_ORIGIN="${v}" ;; + PADDLEFLEET_WHEEL_DIGEST_VERIFIED) FILE_WHEEL_DIGEST="${v}" ;; + PADDLEFLEET_OPS_DIGEST_VERIFIED) FILE_OPS_DIGEST="${v}" ;; + esac + done <"${ENVFILE}" +} + +# Unproven develop fallback: develop_latest origin, or a default 0.0.0 +# filename with no digest proof. A verified ci_metadata/build artifact may +# legally keep the 0.0.0 filename. +unproven_develop_fallback() { + local path="$1" origin="$2" digest="$3" + case "${origin}" in + develop_latest) return 0 ;; + source_tree|build|ci_metadata) + [[ "${origin}" == develop_latest ]] && return 0 + return 1 + ;; + esac + case "${path}" in + */paddlefleet-0.0.0-py3-none-linux_x86_64.whl|*/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl) + [[ "${digest}" == "true" ]] && return 1 + return 0 + ;; esac + return 1 +} + +check_receipt() { + local receipt="$1" requested_mode="$2" requested_pin="$3" + [[ -f "${receipt}" ]] || fail "stack-paired missing receipt ${receipt}" + python3 - "${receipt}" "${requested_mode}" "${requested_pin}" <<'PY' +import json, sys +path, requested_mode, requested_pin = sys.argv[1:4] +doc = json.load(open(path, encoding="utf-8")) +if doc.get("status") != "ok": + raise SystemExit(f"receipt status={doc.get('status')!r} is not ok") +if requested_mode == "stack-paired": + if doc.get("mode") != "stack-paired": + raise SystemExit( + f"requested stack-paired but receipt mode={doc.get('mode')!r}" + ) + src = doc.get("source") or {} + if requested_pin: + exp = src.get("expected_commit") or "" + act = src.get("actual_commit") or "" + if requested_pin not in (exp, act): + raise SystemExit( + f"receipt source pin mismatch requested={requested_pin} " + f"expected={exp} actual={act}" + ) + if src.get("commit_verified") is not True: + raise SystemExit("receipt source commit_verified is not true") + pairing = (doc.get("pairing") or {}).get("status") + if pairing == "unpaired_default": + raise SystemExit("receipt pairing.status=unpaired_default") + for art in doc.get("artifacts") or []: + origin = art.get("origin") or "" + url = art.get("url") or "" + if origin == "develop_latest": + raise SystemExit(f"artifact {art.get('name')} origin=develop_latest") + if "/develop/latest/" in url or "CodeSync/develop/" in url: + raise SystemExit(f"artifact {art.get('name')} unpaired develop URL") +print("receipt matches request") +PY } write_consumed() { @@ -46,55 +139,97 @@ PADDLEFLEET_OPS_WHEEL_PATH=${PADDLEFLEET_OPS_WHEEL_PATH} ALIGNMENT_PADDLEFLEET_MODE=${ALIGNMENT_PADDLEFLEET_MODE} PADDLEFLEET_PIN_RECEIPT=${PADDLEFLEET_PIN_RECEIPT:-} PADDLEFLEET_SOURCE_COMMIT=${PADDLEFLEET_SOURCE_COMMIT:-} +PADDLEFLEET_PIN_SHA=${PADDLEFLEET_PIN_SHA:-} PADDLEFLEET_WHEEL_ORIGIN=${PADDLEFLEET_WHEEL_ORIGIN:-} PADDLEFLEET_OPS_ORIGIN=${PADDLEFLEET_OPS_ORIGIN:-} EOF echo "[paddlefleet-pin-consume] wrote ${OUTFILE}" >&2 - echo "[paddlefleet-pin-consume] wheel=${PADDLEFLEET_WHEEL_PATH} ops=${PADDLEFLEET_OPS_WHEEL_PATH} mode=${ALIGNMENT_PADDLEFLEET_MODE} origin=${PADDLEFLEET_WHEEL_ORIGIN:-}" >&2 + echo "[paddlefleet-pin-consume] requested_mode=${ALIGNMENT_PADDLEFLEET_MODE} pin=${PADDLEFLEET_PIN_SHA:-} wheel=${PADDLEFLEET_WHEEL_PATH} origin=${PADDLEFLEET_WHEEL_ORIGIN:-}" >&2 } consume() { - local mode="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" - if [[ -f "${ENVFILE}" ]]; then - # shellcheck disable=SC1090 - set -a - # shellcheck disable=SC1090 - . "${ENVFILE}" - set +a - echo "[paddlefleet-pin-consume] loaded ${ENVFILE} receipt=${PADDLEFLEET_PIN_RECEIPT:-}" >&2 - mode="${ALIGNMENT_PADDLEFLEET_MODE:-${mode}}" - fi - ALIGNMENT_PADDLEFLEET_MODE="${mode}" - - if [[ "${mode}" == "stack-paired" ]]; then - [[ -f "${ENVFILE}" ]] || { - echo "::error:: stack-paired missing ${ENVFILE}; selector export did not cross docker exec" >&2 - exit 1 - } - [[ -n "${PADDLEFLEET_WHEEL_PATH:-}" && -n "${PADDLEFLEET_OPS_WHEEL_PATH:-}" ]] || { - echo "::error:: stack-paired env missing PADDLEFLEET_WHEEL_PATH or OPS path" >&2 - exit 1 - } - if hardcoded_develop_wheel "${PADDLEFLEET_WHEEL_PATH}" || hardcoded_develop_wheel "${PADDLEFLEET_OPS_WHEEL_PATH}"; then - echo "::error:: stack-paired refused hardcoded 0.0.0 wheel fallback: ${PADDLEFLEET_WHEEL_PATH} ${PADDLEFLEET_OPS_WHEEL_PATH}" >&2 - exit 1 + local requested_mode="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" + local requested_pin="${PADDLEFLEET_PIN_SHA:-}" + + parse_envfile + + if [[ "${requested_mode}" == "stack-paired" ]]; then + [[ -f "${ENVFILE}" ]] || fail "stack-paired missing ${ENVFILE}; selector export did not cross docker exec" + [[ "${FILE_MODE}" == "stack-paired" ]] || fail "requested stack-paired but env mode=${FILE_MODE:-empty} (will not consume a develop leftover)" + if [[ -n "${requested_pin}" ]]; then + local file_id="${FILE_PIN:-${FILE_SOURCE}}" + [[ "${file_id}" == "${requested_pin}" ]] || fail "requested pin ${requested_pin} != env pin/source ${file_id:-empty}" + fi + local receipt="${FILE_RECEIPT:-}" + if [[ -z "${receipt}" && -f "${ENVFILE%/*}/paddlefleet_alignment_pin_receipt.json" ]]; then + receipt="${ENVFILE%/*}/paddlefleet_alignment_pin_receipt.json" + fi + check_receipt "${receipt}" "${requested_mode}" "${requested_pin}" + PADDLEFLEET_WHEEL_PATH="${FILE_WHEEL}" + PADDLEFLEET_OPS_WHEEL_PATH="${FILE_OPS}" + PADDLEFLEET_WHEEL_ORIGIN="${FILE_WHEEL_ORIGIN}" + PADDLEFLEET_OPS_ORIGIN="${FILE_OPS_ORIGIN}" + PADDLEFLEET_PIN_RECEIPT="${receipt}" + PADDLEFLEET_SOURCE_COMMIT="${FILE_SOURCE}" + PADDLEFLEET_PIN_SHA="${requested_pin:-${FILE_PIN}}" + ALIGNMENT_PADDLEFLEET_MODE="stack-paired" + [[ -n "${PADDLEFLEET_WHEEL_PATH}" && -n "${PADDLEFLEET_OPS_WHEEL_PATH}" ]] \ + || fail "stack-paired env missing PADDLEFLEET_WHEEL_PATH or OPS path" + if unproven_develop_fallback "${PADDLEFLEET_WHEEL_PATH}" "${FILE_WHEEL_ORIGIN}" "${FILE_WHEEL_DIGEST}"; then + fail "stack-paired refused unproven develop fallback for paddlefleet: path=${PADDLEFLEET_WHEEL_PATH} origin=${FILE_WHEEL_ORIGIN:-empty} digest_verified=${FILE_WHEEL_DIGEST:-false}" + fi + if unproven_develop_fallback "${PADDLEFLEET_OPS_WHEEL_PATH}" "${FILE_OPS_ORIGIN}" "${FILE_OPS_DIGEST}"; then + fail "stack-paired refused unproven develop fallback for paddlefleet_ops: path=${PADDLEFLEET_OPS_WHEEL_PATH} origin=${FILE_OPS_ORIGIN:-empty} digest_verified=${FILE_OPS_DIGEST:-false}" fi else - PADDLEFLEET_WHEEL_PATH="${PADDLEFLEET_WHEEL_PATH:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" - PADDLEFLEET_OPS_WHEEL_PATH="${PADDLEFLEET_OPS_WHEEL_PATH:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" + ALIGNMENT_PADDLEFLEET_MODE="develop" + PADDLEFLEET_WHEEL_PATH="${FILE_WHEEL:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" + PADDLEFLEET_OPS_WHEEL_PATH="${FILE_OPS:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" + PADDLEFLEET_WHEEL_ORIGIN="${FILE_WHEEL_ORIGIN:-develop_latest}" + PADDLEFLEET_OPS_ORIGIN="${FILE_OPS_ORIGIN:-develop_latest}" + PADDLEFLEET_PIN_RECEIPT="${FILE_RECEIPT:-}" + PADDLEFLEET_SOURCE_COMMIT="${FILE_SOURCE:-}" + PADDLEFLEET_PIN_SHA="${requested_pin}" fi if [[ ! -e "${PADDLEFLEET_WHEEL_PATH}" ]]; then - echo "::error:: missing paddlefleet path: ${PADDLEFLEET_WHEEL_PATH}" >&2 - exit 1 + fail "missing paddlefleet path: ${PADDLEFLEET_WHEEL_PATH}" fi if [[ ! -e "${PADDLEFLEET_OPS_WHEEL_PATH}" ]]; then - echo "::error:: missing paddlefleet_ops path: ${PADDLEFLEET_OPS_WHEEL_PATH}" >&2 - exit 1 + fail "missing paddlefleet_ops path: ${PADDLEFLEET_OPS_WHEEL_PATH}" fi write_consumed } +write_min_receipt() { + local path="$1" mode="$2" pin="$3" wheel="$4" ops="$5" origin="$6" digest="$7" pairing="$8" + python3 - "${path}" "${mode}" "${pin}" "${wheel}" "${ops}" "${origin}" "${digest}" "${pairing}" <<'PY' +import json, sys +path, mode, pin, wheel, ops, origin, digest, pairing = sys.argv[1:9] +digest_ok = digest == "true" +doc = { + "schema": "paddlefleet-alignment-pin/v1", + "status": "ok", + "mode": mode, + "source": { + "expected_commit": pin or None, + "actual_commit": pin or None, + "commit_verified": bool(pin) and mode == "stack-paired", + }, + "artifacts": [ + {"name": "paddlefleet", "path": wheel, "url": None, "origin": origin, + "digest_verified": digest_ok, "expected_sha256": "abc" if digest_ok else None, + "actual_sha256": "abc" if digest_ok else None}, + {"name": "paddlefleet_ops", "path": ops, "url": None, "origin": origin, + "digest_verified": digest_ok, "expected_sha256": "def" if digest_ok else None, + "actual_sha256": "def" if digest_ok else None}, + ], + "pairing": {"stack_paired_proven": False, "status": pairing, "reason": "fixture"}, +} +open(path, "w", encoding="utf-8").write(json.dumps(doc, indent=2) + "\n") +PY +} + run_self_test() { local root self self="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)/$(basename -- "${BASH_SOURCE[0]}")" @@ -103,36 +238,93 @@ run_self_test() { mkdir -p "${root}/PaddleFleet/packages/paddlefleet_ops" echo tree >"${root}/PaddleFleet/pyproject.toml" echo ops >"${root}/PaddleFleet/packages/paddlefleet_ops/pyproject.toml" + local pin="aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" + write_min_receipt "${root}/receipt.json" stack-paired "${pin}" \ + "${root}/PaddleFleet" "${root}/PaddleFleet/packages/paddlefleet_ops" \ + source_tree false source_tree_from_checked_out_pin cat >"${root}/pin.env" <"${root}/missing.env" <"${root}/devwhl/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" + echo dummy >"${root}/devwhl/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" + write_min_receipt "${root}/develop-receipt.json" develop "" \ + "${root}/devwhl/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" \ + "${root}/devwhl/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" \ + develop_latest false unpaired_default + cat >"${root}/develop.env" <&2 + exit 1 + fi + + # 2) verified URL+hash artifact may keep the 0.0.0 filename. + write_min_receipt "${root}/named-receipt.json" stack-paired "${pin}" \ + "${root}/devwhl/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" \ + "${root}/devwhl/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" \ + ci_metadata true unproven + cat >"${root}/named.env" <&2 + ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${pin}" \ + bash "${self}" --env "${root}/named.env" --out "${root}/named.consumed.env" + grep -q "paddlefleet-0.0.0-py3-none-linux_x86_64.whl" "${root}/named.consumed.env" + + # Unproven 0.0.0 fallback (no origin, no digest) still fails. + cat >"${root}/bare.env" <&2 exit 1 fi + echo "consume_paddlefleet_alignment_pin self-test OK" } diff --git a/scripts/select_paddlefleet_alignment_pin.sh b/scripts/select_paddlefleet_alignment_pin.sh index 2fdc7b2d069..55e5f7f3625 100755 --- a/scripts/select_paddlefleet_alignment_pin.sh +++ b/scripts/select_paddlefleet_alignment_pin.sh @@ -254,6 +254,8 @@ PADDLEFLEET_SOURCE_COMMIT=${ACTUAL_SHA} PADDLEFLEET_PIN_SHA=${PIN_SHA} PADDLEFLEET_WHEEL_ORIGIN=${WHEEL_ORIGIN} PADDLEFLEET_OPS_ORIGIN=${OPS_ORIGIN} +PADDLEFLEET_WHEEL_DIGEST_VERIFIED=${WHEEL_DIGEST_VERIFIED} +PADDLEFLEET_OPS_DIGEST_VERIFIED=${OPS_DIGEST_VERIFIED} EOF } @@ -676,6 +678,7 @@ PY grep -q "PADDLEFLEET_WHEEL_PATH=${root}/ok-source/PaddleFleet$" "${root}/ok-source/paddlefleet_alignment_pin.env" grep -q "PADDLEFLEET_OPS_WHEEL_PATH=${root}/ok-source/PaddleFleet/packages/paddlefleet_ops$" "${root}/ok-source/paddlefleet_alignment_pin.env" grep -q "PADDLEFLEET_WHEEL_ORIGIN=source_tree" "${root}/ok-source/paddlefleet_alignment_pin.env" + grep -q "PADDLEFLEET_WHEEL_DIGEST_VERIFIED=false" "${root}/ok-source/paddlefleet_alignment_pin.env" grep -q "PADDLEFLEET_SOURCE_COMMIT=${sha_b}" "${root}/ok-source/paddlefleet_alignment_pin.env" grep -q 'CodeSync/develop/PaddleFleet.tar' "${script}" diff --git a/scripts/test_paddlefleet_pin_handoff.sh b/scripts/test_paddlefleet_pin_handoff.sh index 948dcbe3188..990d9924f8c 100755 --- a/scripts/test_paddlefleet_pin_handoff.sh +++ b/scripts/test_paddlefleet_pin_handoff.sh @@ -1,8 +1,10 @@ #!/usr/bin/env bash # Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. # -# End-to-end: selector source-mode -> new-shell consume -> setup_venvs consumer. -# This is the docker-exec step boundary. Isolated selector --self-test is not enough. +# End-to-end path check: selector source-mode -> new-shell consume -> +# setup_venvs *path* consumer. This proves docker-exec handoff of source +# paths. It is NOT a uv install / real setup_venvs / numerical CI run. +# Isolated selector --self-test is not enough for the step boundary. set -euo pipefail @@ -34,28 +36,41 @@ test -f "${tmp}/ws/paddlefleet_alignment_pin.env" test -d "${tmp}/ws/PaddleFleet" # Step B: new docker exec — drop selector shell state, keep only files. +# Requested mode/pin stay on the caller; leftover develop env must not win. unset PADDLEFLEET_WHEEL_PATH PADDLEFLEET_OPS_WHEEL_PATH PADDLEFLEET_SOURCE_COMMIT || true -export ALIGNMENT_PADDLEFLEET_MODE=stack-paired -bash "${CONSUME}" --env "${tmp}/ws/paddlefleet_alignment_pin.env" --out "${tmp}/ws/consumed.env" +ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${PIN}" \ + bash "${CONSUME}" --env "${tmp}/ws/paddlefleet_alignment_pin.env" --out "${tmp}/ws/consumed.env" -# Step C: setup_venvs consumer. Fail if it would install the 0.0.0 develop wheel. -stub_setup="${tmp}/setup_venvs.sh" +# Negative: leftover develop env + requested stack-paired must fail closed. +mkdir -p "${tmp}/dev" +echo dummy >"${tmp}/dev/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" +echo dummy >"${tmp}/dev/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" +cat >"${tmp}/dev.env" <&2 + exit 1 +fi + +# Step C: path consumer only. Does not run uv or setup_venvs.sh. +stub_setup="${tmp}/setup_path_consumer.sh" cat >"${stub_setup}" <<'STUB' #!/usr/bin/env bash set -euo pipefail -# Mirrors setup_venvs.sh: it only sees exported PADDLEFLEET_WHEEL_PATH. -PADDLEFLEET_WHEEL="${PADDLEFLEET_WHEEL_PATH:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" -PADDLEFLEET_OPS_WHEEL="${PADDLEFLEET_OPS_WHEEL_PATH:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" -case "${PADDLEFLEET_WHEEL}" in - */paddlefleet-0.0.0-py3-none-linux_x86_64.whl) - echo "STUB_SETUP used hardcoded 0.0.0 wheel" >&2 - exit 1 - ;; -esac +# Mirrors setup_venvs.sh reading PADDLEFLEET_WHEEL_PATH. Path presence only. +PADDLEFLEET_WHEEL="${PADDLEFLEET_WHEEL_PATH:?missing PADDLEFLEET_WHEEL_PATH}" +PADDLEFLEET_OPS_WHEEL="${PADDLEFLEET_OPS_WHEEL_PATH:?missing PADDLEFLEET_OPS_WHEEL_PATH}" [[ -d "${PADDLEFLEET_WHEEL}" || -f "${PADDLEFLEET_WHEEL}" ]] || { echo "missing ${PADDLEFLEET_WHEEL}" >&2; exit 1; } [[ -d "${PADDLEFLEET_OPS_WHEEL}" || -f "${PADDLEFLEET_OPS_WHEEL}" ]] || { echo "missing ${PADDLEFLEET_OPS_WHEEL}" >&2; exit 1; } -echo "STUB_SETUP paddlefleet=${PADDLEFLEET_WHEEL}" -echo "STUB_SETUP ops=${PADDLEFLEET_OPS_WHEEL}" +echo "PATH_CONSUMER paddlefleet=${PADDLEFLEET_WHEEL}" +echo "PATH_CONSUMER ops=${PADDLEFLEET_OPS_WHEEL}" +echo "PATH_CONSUMER not_uv_install=true" STUB chmod +x "${stub_setup}" @@ -64,15 +79,8 @@ set -a . "${tmp}/ws/consumed.env" set +a bash "${stub_setup}" | tee "${tmp}/setup.out" -grep -q "STUB_SETUP paddlefleet=${tmp}/ws/PaddleFleet" "${tmp}/setup.out" +grep -q "PATH_CONSUMER paddlefleet=${tmp}/ws/PaddleFleet" "${tmp}/setup.out" grep -q "packages/paddlefleet_ops" "${tmp}/setup.out" +grep -q "not_uv_install=true" "${tmp}/setup.out" -# Negative: alignment step that ignores consumed.env and defaults 0.0.0 must fail the stub. -if PADDLEFLEET_WHEEL_PATH=/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl \ - PADDLEFLEET_OPS_WHEEL_PATH=/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl \ - bash "${stub_setup}"; then - echo "handoff FAIL: stub accepted 0.0.0 default" >&2 - exit 1 -fi - -echo "paddlefleet pin handoff OK pin=${PIN}" +echo "paddlefleet pin handoff PATH_PASS pin=${PIN} (not uv install, not CI)" From ed9c57cd0296b18eeb2057e4a89f5dd2dd62a6f1 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Sat, 5 Sep 2026 23:02:04 +0800 Subject: [PATCH 16/27] Stop Get Whl after selector failure and drop nested quotes Selector clone failure must fail the docker exec (set -e) instead of continuing wget/build. Receipt checks live in a helper so the single-quoted -c script has no nested quotes. Missing pin.env after an error receipt is selector failure, not env-handoff. --- .../workflows/alignment_model_accuracy.yml | 9 +- scripts/consume_paddlefleet_alignment_pin.sh | 37 ++++- scripts/require_paddlefleet_selector_ok.sh | 40 +++++ scripts/test_alignment_workflow_shell.sh | 153 ++++++++++++++++++ scripts/test_paddlefleet_pin_handoff.sh | 25 +++ 5 files changed, 260 insertions(+), 4 deletions(-) create mode 100755 scripts/require_paddlefleet_selector_ok.sh create mode 100755 scripts/test_alignment_workflow_shell.sh diff --git a/.github/workflows/alignment_model_accuracy.yml b/.github/workflows/alignment_model_accuracy.yml index b03f4f7b4ff..c3d73b3e82d 100644 --- a/.github/workflows/alignment_model_accuracy.yml +++ b/.github/workflows/alignment_model_accuracy.yml @@ -161,6 +161,7 @@ jobs: - name: Get Whl run: | docker exec -t $container_name /bin/bash -c ' + set -eo pipefail . /opt/conda/etc/profile.d/conda.sh conda activate py_$python_version python --version @@ -173,9 +174,10 @@ jobs: python -m pip install uv test -x /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh \ || { echo "::error:: selector missing; checkout did not land COMMIT_ID"; exit 1; } - bash /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh --dest /workspace - cat /workspace/paddlefleet_alignment_pin_receipt.json - test -f /workspace/paddlefleet_alignment_pin.env + bash /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh --dest /workspace \ + || { echo "::error:: selector failed; stop Get Whl (do not download remaining wheels or build megatron-core)"; exit 1; } + bash /workspace/Megatron-LM/scripts/require_paddlefleet_selector_ok.sh /workspace \ + || { echo "::error:: selector receipt status is not ok; stop Get Whl"; exit 1; } echo "::endgroup::" echo "::group::Download remaining wheels from BOS" for url in \ @@ -209,6 +211,7 @@ jobs: - name: alignment_model_accuracy run: | docker exec -t $container_name /bin/bash -c ' + set -eo pipefail . /opt/conda/etc/profile.d/conda.sh conda activate py_$python_version python --version diff --git a/scripts/consume_paddlefleet_alignment_pin.sh b/scripts/consume_paddlefleet_alignment_pin.sh index f9181d03ba4..1ba0b8e56b5 100755 --- a/scripts/consume_paddlefleet_alignment_pin.sh +++ b/scripts/consume_paddlefleet_alignment_pin.sh @@ -154,7 +154,24 @@ consume() { parse_envfile if [[ "${requested_mode}" == "stack-paired" ]]; then - [[ -f "${ENVFILE}" ]] || fail "stack-paired missing ${ENVFILE}; selector export did not cross docker exec" + if [[ ! -f "${ENVFILE}" ]]; then + local rec="${ENVFILE%/*}/paddlefleet_alignment_pin_receipt.json" + local rec_status="" + if [[ -f "${rec}" ]]; then + rec_status="$(python3 - "${rec}" <<'PY' +import json, sys +print(json.load(open(sys.argv[1])).get("status") or "") +PY +)" + fi + if [[ "${rec_status}" == "error" ]]; then + fail "stack-paired missing ${ENVFILE} because selector failed (error receipt ${rec}); Get Whl must exit on selector failure. This is not proof that a generated env failed to cross docker exec" + fi + if [[ "${rec_status}" == "ok" ]]; then + fail "stack-paired missing ${ENVFILE}; selector wrote ok receipt but env did not cross docker exec" + fi + fail "stack-paired missing ${ENVFILE} and no selector receipt" + fi [[ "${FILE_MODE}" == "stack-paired" ]] || fail "requested stack-paired but env mode=${FILE_MODE:-empty} (will not consume a develop leftover)" if [[ -n "${requested_pin}" ]]; then local file_id="${FILE_PIN:-${FILE_SOURCE}}" @@ -325,6 +342,24 @@ EOF exit 1 fi + # Selector clone/fail: error receipt, no env. Missing env is the + # consequence, not proof a generated env failed to cross docker exec. + mkdir -p "${root}/sel-fail" + cat >"${root}/sel-fail/paddlefleet_alignment_pin_receipt.json" <<'EOF' +{"schema":"paddlefleet-alignment-pin/v1","status":"error","detail":"git clone failed: github.com:443","mode":"stack-paired"} +EOF + if ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${pin}" \ + bash "${self}" --env "${root}/sel-fail/paddlefleet_alignment_pin.env" \ + --out "${root}/sel-fail.consumed.env" 2>"${root}/sel-fail.err"; then + echo "self-test FAIL: selector error receipt was consumed" >&2 + exit 1 + fi + grep -q "because selector failed" "${root}/sel-fail.err" + if grep -q "selector wrote ok receipt but env did not cross docker exec" "${root}/sel-fail.err"; then + echo "self-test FAIL: selector error misclassified as env-handoff" >&2 + exit 1 + fi + echo "consume_paddlefleet_alignment_pin self-test OK" } diff --git a/scripts/require_paddlefleet_selector_ok.sh b/scripts/require_paddlefleet_selector_ok.sh new file mode 100755 index 00000000000..f34221204ae --- /dev/null +++ b/scripts/require_paddlefleet_selector_ok.sh @@ -0,0 +1,40 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# Called from Get Whl after select_paddlefleet_alignment_pin.sh. +# Keep this file free of nested quotes so the workflow docker exec +# single-quoted -c script can invoke it without breaking the host shell. +# A selector error receipt is not an env-handoff failure. + +set -euo pipefail + +DEST="${1:-/workspace}" +REC="${DEST}/paddlefleet_alignment_pin_receipt.json" +ENVF="${DEST}/paddlefleet_alignment_pin.env" + +if [[ ! -f "${REC}" ]]; then + echo "::error:: missing selector receipt ${REC}" >&2 + exit 1 +fi + +python3 - "${REC}" <<'PY' +import json +import sys + +path = sys.argv[1] +doc = json.load(open(path, encoding="utf-8")) +status = doc.get("status") +if status != "ok": + raise SystemExit( + f"selector receipt status={status!r} is not ok; " + "Get Whl must stop (do not download remaining wheels or build)" + ) +print("selector receipt status=ok") +PY + +if [[ ! -f "${ENVF}" ]]; then + echo "::error:: selector did not write ${ENVF}" >&2 + exit 1 +fi + +cat "${REC}" diff --git a/scripts/test_alignment_workflow_shell.sh b/scripts/test_alignment_workflow_shell.sh new file mode 100755 index 00000000000..4b666863c5f --- /dev/null +++ b/scripts/test_alignment_workflow_shell.sh @@ -0,0 +1,153 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# +# Syntax-check every workflow `run:` block, then execute the extracted +# Get Whl docker-exec body against a failing selector. Independent +# helper --self-test is not this check. + +set -euo pipefail + +ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" +YAML="${ROOT}/../.github/workflows/alignment_model_accuracy.yml" +SELECTOR_PATH="/workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh" +REQUIRE_PATH="/workspace/Megatron-LM/scripts/require_paddlefleet_selector_ok.sh" +BUILD_MARKER="build megatron-core" + +python3 - "${YAML}" "${ROOT}" "${SELECTOR_PATH}" "${REQUIRE_PATH}" "${BUILD_MARKER}" <<'PY' +import os, re, subprocess, sys, tempfile, textwrap, pathlib, json, stat + +yaml_path, scripts_root, selector_path, require_path, build_marker = sys.argv[1:6] +text = pathlib.Path(yaml_path).read_text() +lines = text.splitlines(True) + +blocks = [] +i = 0 +while i < len(lines): + m = re.match(r"^(\s*)run:\s*\|\s*$", lines[i]) + if not m: + i += 1 + continue + indent = len(m.group(1)) + name = "unnamed" + for j in range(i, -1, -1): + nm = re.match(r"^\s+- name:\s*(.*)$", lines[j]) + if nm: + name = nm.group(1).strip() + break + i += 1 + body = [] + while i < len(lines): + line = lines[i] + if line.strip() == "": + body.append(line) + i += 1 + continue + lead = len(line) - len(line.lstrip(" ")) + if lead <= indent and line.strip(): + break + body.append(line[indent + 2 :] if lead >= indent + 2 else line.lstrip()) + i += 1 + blocks.append((name, "".join(body))) + +if not blocks: + raise SystemExit(f"no run: | blocks in {yaml_path}") + +tmp = pathlib.Path(tempfile.mkdtemp(prefix="yaml-run-")) +print(f"extracted {len(blocks)} run blocks from {yaml_path}") +for n, body in blocks: + p = tmp / (re.sub(r"[^A-Za-z0-9._-]+", "_", n) + ".sh") + p.write_text("#!/usr/bin/env bash\n" + body) + r = subprocess.run(["bash", "-n", str(p)], capture_output=True, text=True) + if r.returncode != 0: + raise SystemExit(f"bash -n FAIL {n}: {r.stderr}") + print(f"bash -n OK run:{n}") + + for m in re.finditer(r"""/bin/bash -c\s+'""", body): + start = m.end() + end = body.find("'\n", start) + if end < 0: + end = body.rfind("'") + inner = body[start:end] + if "'" in inner: + raise SystemExit( + f"nested single quote inside docker exec -c in {n}: " + f"{inner[inner.find(chr(39))-40:inner.find(chr(39))+40]!r}" + ) + inner_p = tmp / (p.stem + ".docker-inner.sh") + inner_p.write_text("#!/usr/bin/env bash\n" + inner + "\n") + r = subprocess.run(["bash", "-n", str(inner_p)], capture_output=True, text=True) + if r.returncode != 0: + raise SystemExit(f"bash -n FAIL docker-inner {n}: {r.stderr}") + print(f"bash -n OK docker-inner:{n} (no nested single quotes)") + +getwhl = next((b for n, b in blocks if n == "Get Whl"), None) +if getwhl is None: + raise SystemExit("Get Whl run block missing") +m = re.search(r"""/bin/bash -c\s+'""", getwhl) +if not m: + raise SystemExit("Get Whl docker exec -c missing") +inner = getwhl[m.end():] +end = inner.rfind("'") +inner = inner[:end] + +ws = tmp / "ws" +(ws / "Megatron-LM/scripts").mkdir(parents=True) +(ws / "upload").mkdir(parents=True) +selector = ws / "Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh" +require_src = pathlib.Path(scripts_root) / "require_paddlefleet_selector_ok.sh" +require_dst = ws / "Megatron-LM/scripts/require_paddlefleet_selector_ok.sh" +require_dst.write_text(require_src.read_text()) +require_dst.chmod(require_dst.stat().st_mode | stat.S_IXUSR) +selector.write_text(textwrap.dedent("""\ + #!/usr/bin/env bash + set -euo pipefail + dest="${2:-/workspace}" + mkdir -p "${dest}" + cat >"${dest}/paddlefleet_alignment_pin_receipt.json" <<'EOF' + {"schema":"paddlefleet-alignment-pin/v1","status":"error","detail":"git clone failed: github.com:443","mode":"stack-paired"} + EOF + echo "[paddlefleet-pin] FAIL: git clone failed: github.com:443" >&2 + echo "::error:: git clone failed: github.com:443" >&2 + exit 1 + """)) +selector.chmod(selector.stat().st_mode | stat.S_IXUSR) +(ws / "Megatron-LM/scripts/dependence").mkdir(parents=True, exist_ok=True) +(ws / "Megatron-LM/scripts/dependence/build.sh").write_text("#!/usr/bin/env bash\necho BUILD_RAN > /workspace/upload/BUILD_RAN\n") +(ws / "Megatron-LM/scripts/dependence/build.sh").chmod(0o755) + +rewritten = inner.replace("/workspace", str(ws)) +rewritten = rewritten.replace("conda activate py_$python_version", "true") +rewritten = rewritten.replace(". /opt/conda/etc/profile.d/conda.sh", "true") +rewritten = "#!/usr/bin/env bash\nexport ALIGNMENT_PADDLEFLEET_MODE=stack-paired\nexport python_version=3.12\n" + rewritten +bin = tmp / "bin" +bin.mkdir() +(bin / "wget").write_text("#!/usr/bin/env bash\necho WGET_RAN \"$@\" >> '%s/WGET_RAN'\nexit 0\n" % ws) +(bin / "python").write_text("#!/usr/bin/env bash\necho PY_RAN \"$@\" >> '%s/PY_RAN'\nexit 0\n" % ws) +(bin / "python3").write_text("#!/usr/bin/env bash\nexec /usr/bin/python3 \"$@\"\n") +(bin / "pip").write_text("#!/usr/bin/env bash\necho PIP_RAN \"$@\" >> '%s/PIP_RAN'\nexit 0\n" % ws) +(bin / "ldconfig").write_text("#!/usr/bin/env bash\nexit 0\n") +for f in bin.iterdir(): + f.chmod(0o755) + +script = tmp / "getwhl.extracted.sh" +script.write_text(rewritten) +script.chmod(0o755) +env = os.environ.copy() +env["PATH"] = str(bin) + ":" + env.get("PATH", "") +env["ALIGNMENT_PADDLEFLEET_MODE"] = "stack-paired" +r = subprocess.run(["bash", str(script)], capture_output=True, text=True, env=env) +log = (r.stdout or "") + (r.stderr or "") +print("extracted Get Whl rc=", r.returncode) +print(log[-2000:]) +if r.returncode == 0: + raise SystemExit("FAIL: extracted Get Whl continued after selector failure") +if (ws / "WGET_RAN").exists(): + raise SystemExit("FAIL: wget ran after selector failure") +if (ws / "upload/BUILD_RAN").exists(): + raise SystemExit("FAIL: build.sh ran after selector failure") +if "selector failed; stop Get Whl" not in log and "git clone failed" not in log: + raise SystemExit("FAIL: extracted Get Whl did not surface selector failure") +print("extracted Get Whl fixture: selector fail stopped remaining wheels and build") +print("workflow shell checks OK") +PY +echo "alignment workflow shell PATH_PASS (extracted YAML, not helper-only)" diff --git a/scripts/test_paddlefleet_pin_handoff.sh b/scripts/test_paddlefleet_pin_handoff.sh index 990d9924f8c..9b01769b882 100755 --- a/scripts/test_paddlefleet_pin_handoff.sh +++ b/scripts/test_paddlefleet_pin_handoff.sh @@ -58,6 +58,31 @@ if ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${PIN}" \ exit 1 fi +# Negative: selector clone/fail writes error receipt and no env. Consume +# must name selector failure. Missing env here is the consequence, not +# proof that a generated env failed to cross docker exec. +REQUIRE="${ROOT}/require_paddlefleet_selector_ok.sh" +mkdir -p "${tmp}/sel-fail" +cat >"${tmp}/sel-fail/paddlefleet_alignment_pin_receipt.json" <<'EOF' +{"schema":"paddlefleet-alignment-pin/v1","status":"error","detail":"git clone failed: github.com:443","mode":"stack-paired"} +EOF +if bash "${REQUIRE}" "${tmp}/sel-fail" >"${tmp}/sel-fail.require.out" 2>"${tmp}/sel-fail.require.err"; then + echo "handoff FAIL: require_ok accepted error receipt" >&2 + exit 1 +fi +grep -q "selector receipt status=" "${tmp}/sel-fail.require.err" +if ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${PIN}" \ + bash "${CONSUME}" --env "${tmp}/sel-fail/paddlefleet_alignment_pin.env" \ + --out "${tmp}/sel-fail.consumed.env" 2>"${tmp}/sel-fail.err"; then + echo "handoff FAIL: selector error receipt was consumed" >&2 + exit 1 +fi +grep -q "because selector failed" "${tmp}/sel-fail.err" +if grep -q "selector wrote ok receipt but env did not cross docker exec" "${tmp}/sel-fail.err"; then + echo "handoff FAIL: selector error misclassified as env-handoff" >&2 + exit 1 +fi + # Step C: path consumer only. Does not run uv or setup_venvs.sh. stub_setup="${tmp}/setup_path_consumer.sh" cat >"${stub_setup}" <<'STUB' From 9b48c090ab88a1c9603f6ebaeb86cc155fafce2e Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Sun, 6 Sep 2026 22:57:32 +0800 Subject: [PATCH 17/27] fix: fetch PaddleFleet pin at depth 1 with bounded retries stack-paired checkout_pin used a full default-branch git clone with no retry. Swift 34038640242 failed Get Whl on curl 56 / early EOF before ops, so Fleet 09bb4bd4 was not evaluated. Fetch origin $PIN_SHA at --depth=1 --no-tags, detach-checkout, and keep exact HEAD. Transient RPC/curl-56 retries up to 3 clean dests; HTTP 401/403, missing SHA, and missing repo fail closed on attempt 1 with commit_verified=false. --- scripts/select_paddlefleet_alignment_pin.sh | 231 ++++++++++++++++++-- 1 file changed, 210 insertions(+), 21 deletions(-) diff --git a/scripts/select_paddlefleet_alignment_pin.sh b/scripts/select_paddlefleet_alignment_pin.sh index 55e5f7f3625..cfa778fa9d3 100755 --- a/scripts/select_paddlefleet_alignment_pin.sh +++ b/scripts/select_paddlefleet_alignment_pin.sh @@ -20,7 +20,9 @@ # Cases are not filtered. # # Explicit (ALIGNMENT_PADDLEFLEET_MODE=stack-paired): fail-closed. -# Checkout PADDLEFLEET_PIN_SHA; git rev-parse HEAD must equal the pin. +# Fetch PADDLEFLEET_PIN_SHA at depth 1 (not a full default-branch clone). +# Transient git RPC/curl failures retry up to 3 clean dests; other +# fetch errors fail closed. git rev-parse HEAD must equal the pin. # Artifacts: caller URL+sha256 (from Build Fleet whl / Actions metadata) # or build from the checked-out tree. Independent digest matches do not # prove the wheels were produced from that commit — receipt records @@ -292,26 +294,65 @@ require_digest() { fi } +# Auth / missing-object / missing-repo are permanent. Do not treat a +# generic "RPC failed" as transient: HTTP 401/403 also say RPC failed. +_is_permanent_git_err() { + case "$1" in + *"HTTP 401"*|*"HTTP 403"*|*"Authentication failed"*|*"access denied"*|*"Access denied"*|*"Permission denied"*|*"not our ref"*|*"does not appear to be a git repository"*|*"Could not read from remote repository"*|*"remote: Write access"*) + return 0 + ;; + esac + return 1 +} + +_is_transient_git_err() { + _is_permanent_git_err "$1" && return 1 + case "$1" in + *"curl 56"*|*"Connection timed out"*|*"Couldn't connect to server"*|*"Failed to connect to github.com port 443"*|*"early EOF"*|*"fetch-pack: unexpected disconnect"*|*"bytes of body are still expected"*|*"invalid index-pack output"*|*"Connection reset by peer"*) + return 0 + ;; + esac + return 1 +} + checkout_pin() { [[ "${PIN_SHA}" =~ ^[0-9a-fA-F]{40}$ ]] || fail "stack-paired requires PADDLEFLEET_PIN_SHA (40 hex), got '${PIN_SHA}'" PIN_SHA="$(printf '%s' "${PIN_SHA}" | tr 'A-F' 'a-f')" - rm -rf "${DEST}/PaddleFleet" - log "clone ${GIT_URL}" - if ! git clone --quiet "${GIT_URL}" "${DEST}/PaddleFleet" >/dev/null 2>"${DEST}/.git-clone.err"; then - fail "git clone failed: $(tr '\n' ' ' <"${DEST}/.git-clone.err")" - fi - git -C "${DEST}/PaddleFleet" config advice.detachedHead false || true - log "checkout ${PIN_SHA}" - if ! git -C "${DEST}/PaddleFleet" checkout --quiet --force "${PIN_SHA}" >/dev/null 2>"${DEST}/.git-co.err"; then - fail "git checkout failed for ${PIN_SHA}: $(tr '\n' ' ' <"${DEST}/.git-co.err")" - fi - ACTUAL_SHA="$(git -C "${DEST}/PaddleFleet" rev-parse HEAD)" - if [[ "${ACTUAL_SHA}" != "${PIN_SHA}" ]]; then - SOURCE_VERIFIED=false - fail "stack-paired source SHA mismatch expected=${PIN_SHA} actual=${ACTUAL_SHA}" - fi - SOURCE_VERIFIED=true - log "source commit verified ${ACTUAL_SHA}" + local dest="${DEST}/PaddleFleet" + local attempt max_attempts=3 + local last_err="" + local fetch_err="${DEST}/.git-fetch.err" + for attempt in $(seq 1 "${max_attempts}"); do + rm -rf "${dest}" + log "fetch ${GIT_URL} ${PIN_SHA} (attempt ${attempt}/${max_attempts})" + if ! git init --quiet "${dest}" >/dev/null 2>"${DEST}/.git-init.err"; then + fail "git init failed: $(tr '\n' ' ' <"${DEST}/.git-init.err")" + fi + if ! git -C "${dest}" remote add origin "${GIT_URL}" >/dev/null 2>"${DEST}/.git-remote.err"; then + fail "git remote add failed: $(tr '\n' ' ' <"${DEST}/.git-remote.err")" + fi + git -C "${dest}" config advice.detachedHead false || true + if GIT_TERMINAL_PROMPT=0 git -C "${dest}" fetch --depth=1 --no-tags origin "${PIN_SHA}" \ + >/dev/null 2>"${fetch_err}"; then + if ! git -C "${dest}" checkout --quiet --force --detach "${PIN_SHA}" >/dev/null 2>"${DEST}/.git-co.err"; then + fail "git checkout failed for ${PIN_SHA}: $(tr '\n' ' ' <"${DEST}/.git-co.err")" + fi + ACTUAL_SHA="$(git -C "${dest}" rev-parse HEAD)" + if [[ "${ACTUAL_SHA}" != "${PIN_SHA}" ]]; then + SOURCE_VERIFIED=false + fail "stack-paired source SHA mismatch expected=${PIN_SHA} actual=${ACTUAL_SHA}" + fi + SOURCE_VERIFIED=true + log "source commit verified ${ACTUAL_SHA}" + return 0 + fi + last_err="$(tr '\n' ' ' <"${fetch_err}")" + if ! _is_transient_git_err "${last_err}"; then + fail "git fetch failed: ${last_err}" + fi + log "transient fetch failure attempt ${attempt}/${max_attempts}: ${last_err}" + done + fail "git fetch failed: ${last_err}" } # Sets DEST_PATH and DEST_SHA in the caller. Must run in this shell so @@ -557,7 +598,7 @@ run_self_test() { PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ bash "${script}" --dest "${root}/m2" - expect_fail "${root}/m3" "git checkout failed" \ + expect_fail "${root}/m3" "git fetch failed" \ run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ PADDLEFLEET_PIN_SHA="0000000000000000000000000000000000000000" \ PADDLEFLEET_GIT_URL="${root}/upstream" \ @@ -599,12 +640,159 @@ PY PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ bash "${script}" --dest "${root}/m6" - expect_fail "${root}/m7" "git clone failed" \ + expect_fail "${root}/m7" "git fetch failed" \ run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ PADDLEFLEET_PIN_SHA="${sha_b}" \ PADDLEFLEET_GIT_URL="${root}/no-such-remote" \ bash "${script}" --dest "${root}/m7" + assert_error_unverified() { + python3 - "$1" <<'PY' +import json, sys +doc = json.load(open(sys.argv[1])) +assert doc["status"] == "error", doc +assert doc["source"]["commit_verified"] is False, doc["source"] +print("error receipt commit_verified=false") +PY + } + + assert_pin_checkout() { + local repo="$1" sha="$2" + [[ "$(git -C "${repo}" rev-parse HEAD)" == "${sha}" ]] + [[ -f "${repo}/.git/shallow" ]] || { echo "self-test FAIL: missing ${repo}/.git/shallow" >&2; exit 1; } + [[ "$(git -C "${repo}" rev-list --count HEAD)" == 1 ]] || { + echo "self-test FAIL: ${repo} rev-list count != 1 (not depth=1)" >&2 + exit 1 + } + if git -C "${repo}" symbolic-ref -q HEAD >/dev/null; then + echo "self-test FAIL: ${repo} HEAD is a branch, not detached pin" >&2 + git -C "${repo}" symbolic-ref HEAD >&2 + exit 1 + fi + } + + write_git_wrapper() { + local bindir="$1" mode="$2" + local real_git + real_git="$(command -v git)" + mkdir -p "${bindir}" + cat >"${bindir}/git" <>"\$STALE" + exit 99 + fi + if [[ "\$has_depth" -ne 1 || "\$has_no_tags" -ne 1 ]]; then + echo "fetch missing --depth=1/--no-tags: \${args[*]}" >>"\$STALE" + exit 99 + fi + count=\$((count + 1)) + printf '%s\n' "\$count" >"\$STATE" + case "${mode}" in + 401) + echo "RPC failed; HTTP 401 curl 22 The requested URL returned error: 401" >&2 + echo "fatal: Authentication failed" >&2 + printf 'sentinel\n' >"\$workdir/.retry-sentinel" + exit 128 + ;; + exhaust) + echo "error: RPC failed; curl 56 Recv failure: Connection timed out" >&2 + echo "error: 9515 bytes of body are still expected" >&2 + echo "fatal: early EOF" >&2 + printf 'sentinel\n' >"\$workdir/.retry-sentinel" + exit 128 + ;; + retry) + if [[ "\$count" -lt 3 ]]; then + echo "error: RPC failed; curl 56 Recv failure: Connection timed out" >&2 + echo "error: 9515 bytes of body are still expected" >&2 + echo "fetch-pack: unexpected disconnect while reading sideband packet" >&2 + echo "fatal: early EOF" >&2 + echo "fatal: fetch-pack: invalid index-pack output" >&2 + printf 'sentinel\n' >"\$workdir/.retry-sentinel" + exit 128 + fi + ;; + esac +fi +exec "\$real" "\$@" +GITWRAP + chmod +x "${bindir}/git" + } + + assert_error_unverified "${root}/m3/paddlefleet_alignment_pin_receipt.json" + assert_error_unverified "${root}/m7/paddlefleet_alignment_pin_receipt.json" + + # Permanent HTTP 401 (also says RPC failed) must not retry. + write_git_wrapper "${root}/bin-401" 401 + expect_fail "${root}/m8" "git fetch failed" \ + env PATH="${root}/bin-401:${root}/bin:${PATH}" \ + ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + bash "${script}" --dest "${root}/m8" + [[ "$(cat "${root}/bin-401/count")" == 1 ]] + [[ ! -f "${root}/bin-401/stale" ]] + assert_error_unverified "${root}/m8/paddlefleet_alignment_pin_receipt.json" + + # First two fetches emit Swift 34038640242 curl-56 / early-EOF and leave + # a dest sentinel; the next fetch must see a clean dest. Third fetch is + # real git. Depth=1 and detached pin, not a default-branch clone. + write_git_wrapper "${root}/bin-retry" retry + run PATH="${root}/bin-retry:${root}/bin:${PATH}" \ + ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + bash "${script}" --dest "${root}/ok-retry" + [[ "$(cat "${root}/bin-retry/count")" == 3 ]] + [[ ! -f "${root}/bin-retry/stale" ]] + [[ ! -e "${root}/ok-retry/PaddleFleet/.retry-sentinel" ]] + assert_pin_checkout "${root}/ok-retry/PaddleFleet" "${sha_b}" + python3 - "${root}/ok-retry/paddlefleet_alignment_pin_receipt.json" "${sha_b}" <<'PY' +import json, sys +doc = json.load(open(sys.argv[1])) +assert doc["status"] == "ok" +assert doc["source"]["actual_commit"] == sys.argv[2] +assert doc["source"]["commit_verified"] is True +print("ok-retry receipt exact HEAD verified") +PY + + # Exhausted transient retries: wrapper count is 3, dest cleaned between + # attempts, error receipt stays unverified. + write_git_wrapper "${root}/bin-fail" exhaust + expect_fail "${root}/m9" "git fetch failed" \ + env PATH="${root}/bin-fail:${root}/bin:${PATH}" \ + ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ + PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ + bash "${script}" --dest "${root}/m9" + [[ "$(cat "${root}/bin-fail/count")" == 3 ]] + [[ ! -f "${root}/bin-fail/stale" ]] + assert_error_unverified "${root}/m9/paddlefleet_alignment_pin_receipt.json" + run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ PADDLEFLEET_WHEEL_URL="${root}/art/py.whl" \ @@ -631,7 +819,7 @@ assert "MinimaxV2.5_EP2" in doc["cases_preserved"] assert "GLM45Air_EP2" in doc["cases_preserved"] print("ok-url receipt fields checked") PY - [[ "$(git -C "${root}/ok-url/PaddleFleet" rev-parse HEAD)" == "${sha_b}" ]] + assert_pin_checkout "${root}/ok-url/PaddleFleet" "${sha_b}" run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ @@ -680,6 +868,7 @@ PY grep -q "PADDLEFLEET_WHEEL_ORIGIN=source_tree" "${root}/ok-source/paddlefleet_alignment_pin.env" grep -q "PADDLEFLEET_WHEEL_DIGEST_VERIFIED=false" "${root}/ok-source/paddlefleet_alignment_pin.env" grep -q "PADDLEFLEET_SOURCE_COMMIT=${sha_b}" "${root}/ok-source/paddlefleet_alignment_pin.env" + assert_pin_checkout "${root}/ok-source/PaddleFleet" "${sha_b}" grep -q 'CodeSync/develop/PaddleFleet.tar' "${script}" grep -q 'PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl' "${script}" From 6eac17b1c0509e913d9cb1b945d653f986ca9fbc Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Tue, 8 Sep 2026 08:20:21 -0700 Subject: [PATCH 18/27] Bound ops submodule preparation before alignment setup Bound source submodule setup before environment preparation. Keep wheel installs on their existing path. Signed-off-by: Zhan Rongrui --- .../workflows/alignment_model_accuracy.yml | 2 + scripts/prepare_paddlefleet_ops_submodules.sh | 29 +++++++ scripts/test_ops_submodule_preflight.sh | 78 +++++++++++++++++++ 3 files changed, 109 insertions(+) create mode 100644 scripts/prepare_paddlefleet_ops_submodules.sh create mode 100644 scripts/test_ops_submodule_preflight.sh diff --git a/.github/workflows/alignment_model_accuracy.yml b/.github/workflows/alignment_model_accuracy.yml index c3d73b3e82d..8bf2696fd85 100644 --- a/.github/workflows/alignment_model_accuracy.yml +++ b/.github/workflows/alignment_model_accuracy.yml @@ -239,6 +239,8 @@ jobs: echo "using $whl" done python -m pip install uv + bash Megatron-LM/scripts/prepare_paddlefleet_ops_submodules.sh \ + "$PADDLEFLEET_OPS_WHEEL_PATH" PaddleFleet/scripts/alignment_model_accuracy/setup_venvs.sh bash PaddleFleet/scripts/alignment_model_accuracy/setup_venvs.sh bash -x PaddleFleet/scripts/alignment_model_accuracy/run_alignment_test.sh diff --git a/scripts/prepare_paddlefleet_ops_submodules.sh b/scripts/prepare_paddlefleet_ops_submodules.sh new file mode 100644 index 00000000000..7da5c9c32bd --- /dev/null +++ b/scripts/prepare_paddlefleet_ops_submodules.sh @@ -0,0 +1,29 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +# Retry source-tree network preparation before setup/build, retaining the +# pinned Fleet helper's exact and nested gitlink verification. +set -euo pipefail +ops_path="${1:?ops-path}" +setup_script="${2:?pinned-setup-script}" +if [[ ! -d "$ops_path" ]]; then + echo "[ops-submodules] wheel input; no source preparation" + exit 0 +fi +command -v timeout >/dev/null +[[ -f "$setup_script" ]] || { echo "[ops-submodules] pinned setup script missing" >&2; exit 1; } +for attempt in 1 2 3; do + echo "[ops-submodules] prepare and verify recorded gitlinks, attempt $attempt/3" + if timeout --kill-after=30s 15m bash "$setup_script" --prepare-ops-submodules "$ops_path"; then + echo "[ops-submodules] recorded gitlinks prepared and verified" + exit 0 + else + rc=$? + fi + if [[ "$attempt" == 3 ]]; then + echo "[ops-submodules] preparation failed after 3 attempts (exit $rc); stop before setup/build/training" >&2 + exit "$rc" + fi + delay=$((attempt * 15)) + echo "[ops-submodules] preparation exited $rc; retry in ${delay}s" + sleep "$delay" +done diff --git a/scripts/test_ops_submodule_preflight.sh b/scripts/test_ops_submodule_preflight.sh new file mode 100644 index 00000000000..1f08b7253dc --- /dev/null +++ b/scripts/test_ops_submodule_preflight.sh @@ -0,0 +1,78 @@ +#!/usr/bin/env bash +# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. +set -euo pipefail +ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)" +"${PYTHON_BIN:-python3}" - "$ROOT" <<'PY' +import os +import pathlib +import shutil +import subprocess +import sys +import tempfile + +root = pathlib.Path(sys.argv[1]) +workflow = (root / '.github/workflows/alignment_model_accuracy.yml').read_text() +start = workflow.index(' bash Megatron-LM/scripts/prepare_paddlefleet_ops_submodules.sh') +end = workflow.index('\n model_acc_align_exit_code=', start) +commands = workflow[start:end] +assert commands.count('setup_venvs.sh') == 2 +assert commands.count('run_alignment_test.sh') == 1 +assert ' set -eo pipefail' in workflow[:start] +real_timeout = shutil.which('timeout') +assert real_timeout +with tempfile.TemporaryDirectory() as directory: + work = pathlib.Path(directory) + (work / 'bin').mkdir() + (work / 'ops').mkdir() + (work / 'Megatron-LM').symlink_to(root) + fleet = work / 'PaddleFleet/scripts/alignment_model_accuracy' + fleet.mkdir(parents=True) + setup = fleet / 'setup_venvs.sh' + setup.write_text('''#!/usr/bin/env bash +set -eu +if [[ ${1:-} != --prepare-ops-submodules ]]; then echo setup >> "$TRACE"; exit 0; fi +[[ $2 == "$OPS" ]] +echo prepare >> "$TRACE" +n=$(grep -c prepare "$TRACE") +case "$CASE" in + success) exit 0;; + transient) [[ $n -ge 2 ]];; + persistent) exit 23;; + timeout) /bin/sleep 5;; +esac +''') + (fleet / 'run_alignment_test.sh').write_text('echo train >> "$TRACE"\n') + sleeper = work / 'bin/sleep' + sleeper.write_text('#!/usr/bin/env bash\necho "sleep:$1" >> "$TRACE"\n') + timer = work / 'bin/timeout' + timer.write_text('''#!/usr/bin/env bash +set -eu +[[ $1 == --kill-after=30s && $2 == 15m ]] +echo bound >> "$TRACE" +shift 2 +if [[ $CASE == timeout ]]; then exec "$REAL_TIMEOUT" --kill-after=0.1s 0.1s "$@"; fi +exec "$@" +''') + sleeper.chmod(0o755) + timer.chmod(0o755) + env = dict(os.environ, PATH=str(work / 'bin') + ':' + os.environ['PATH'], + TRACE=str(work / 'trace'), OPS=str(work / 'ops'), REAL_TIMEOUT=real_timeout) + for case, count, code in [('success', 1, 0), ('transient', 2, 0), + ('persistent', 3, 23), ('timeout', 3, 124), + ('wheel', 0, 0)]: + trace = work / 'trace' + trace.write_text('') + env.update(CASE=case, PADDLEFLEET_OPS_WHEEL_PATH=env['OPS'] if case != 'wheel' else str(work / 'ops.whl')) + result = subprocess.run(['bash', '-eo', 'pipefail', '-c', commands], cwd=work, + env=env, capture_output=True, text=True, timeout=10) + lines = trace.read_text().splitlines() + assert result.returncode == code, (case, result.returncode, result.stderr) + assert lines.count('prepare') == count, (case, lines) + assert lines.count('bound') == count, (case, lines) + assert [x for x in lines if x.startswith('sleep:')] == ['sleep:15', 'sleep:30'][:max(0, count - 1)], (case, lines) + assert ('setup' in lines) == (code == 0), (case, lines) + assert ('train' in lines) == (code == 0), (case, lines) + assert lines.count('setup') <= 1 and lines.count('train') <= 1 + print(f'PASS: extracted workflow {case}, attempts={count}, exit={code}') +print('All source-submodule preflight fixtures passed') +PY From 18efc53324b9bffae51f52ab0ac79853ba5fb05c Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Tue, 8 Sep 2026 08:21:24 -0700 Subject: [PATCH 19/27] Preserve accuracy routing and expert accumulation semantics Keep FP32 accuracy routing, shared expert ordering and native expert accumulation behavior consistent with the validated candidates. Signed-off-by: Zhan Rongrui --- megatron/core/transformer/moe/moe_layer.py | 20 +- megatron/core/transformer/moe/moe_utils.py | 6 +- .../core/transformer/moe/token_dispatcher.py | 7 + .../moe/test_accuracy_migration.py | 212 ++++++++++++++++++ 4 files changed, 240 insertions(+), 5 deletions(-) create mode 100644 tests/unit_tests/transformer/moe/test_accuracy_migration.py diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index 5eb8de35f09..027948f2470 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -576,7 +576,11 @@ def postprocess(self, output: torch.Tensor, shared_expert_output: Optional[torch output, _ = self.fc2_latent_proj(output) if shared_expert_output is not None: - output = output + shared_expert_output + if _use_accuracy_compatible(): + orig_dtype = output.dtype + output = (output.float() + shared_expert_output.float()).to(orig_dtype) + else: + output = output + shared_expert_output elif ( isinstance(self.token_dispatcher, NVLSAllGatherVDispatcher) and self._latent_shared_expert_output is not None @@ -660,7 +664,13 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): hidden_states_router = hidden_states hidden_states_dispatch = hidden_states - shared_expert_output = self.shared_experts_compute(hidden_states_shared) + if _use_accuracy_compatible() and not self.shared_expert_overlap: + self._accuracy_shared_input = hidden_states_shared + shared_expert_output = None + else: + shared_expert_output = self.shared_experts_compute( + hidden_states_shared + ) probs, routing_map = self.route(hidden_states_router, padding_mask) hidden_states, probs = self.preprocess( hidden_states_dispatch, probs, routing_map @@ -695,6 +705,12 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): if intermediate_tensors is not None: output, shared_expert_output = intermediate_tensors + if _use_accuracy_compatible(): + shared_input = getattr(self, "_accuracy_shared_input", None) + if shared_input is not None: + shared_expert_output = self.shared_experts_compute(shared_input) + self._accuracy_shared_input = None + output = self.postprocess(output, shared_expert_output) if intermediate_tensors is not None: diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index 8baa816e697..2deec6bdd48 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -880,13 +880,13 @@ def compute_topk(scores, topk, num_groups=None, group_topk=None): elif score_function in ("sigmoid", "sqrtsoftplus"): if _use_accuracy_compatible(): if score_function == "sigmoid": - scores = torch.sigmoid(logits.float()).type_as(logits) + scores = torch.sigmoid(logits.float()) else: - scores = torch.nn.functional.softplus(logits.float()).sqrt().type_as(logits) + scores = torch.nn.functional.softplus(logits.float()).sqrt() if expert_bias is not None: scores_for_routing = scores + expert_bias _, top_indices = compute_topk(scores_for_routing, topk, num_groups, group_topk) - scores = torch.gather(scores, dim=1, index=top_indices).type_as(logits) + scores = torch.gather(scores, dim=1, index=top_indices) else: scores, top_indices = compute_topk(scores, topk, num_groups, group_topk) _scores_f64 = scores.double() diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 490e4d8cc6d..e4c72b30ff3 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -17,6 +17,7 @@ reduce_scatter_to_sequence_parallel_region, ) from megatron.core.transformer.enums import CudaGraphModule +from megatron.core.transformer.module import _use_accuracy_compatible from megatron.core.transformer.moe.fused_a2a import ( fused_combine, fused_dispatch, @@ -1526,6 +1527,9 @@ def token_dispatch( """ if self.shared_experts is not None: self.shared_experts.wait_current_stream() + if _use_accuracy_compatible(): + async_finish = False + allocate_on_comm_stream = False dispatched_hidden_states = self._comm_manager.dispatch( hidden_states, async_finish, allocate_on_comm_stream ) @@ -1585,6 +1589,9 @@ def token_combine( # when CUDA_DEVICE_MAX_CONNECTIONS>1. if self.shared_experts is not None: self.shared_experts.wait_current_stream() + if _use_accuracy_compatible(): + async_finish = False + allocate_on_comm_stream = False return self._comm_manager.combine(hidden_states, async_finish, allocate_on_comm_stream) def combine_postprocess(self, hidden_states: torch.Tensor): diff --git a/tests/unit_tests/transformer/moe/test_accuracy_migration.py b/tests/unit_tests/transformer/moe/test_accuracy_migration.py new file mode 100644 index 00000000000..1a2e6af8478 --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_accuracy_migration.py @@ -0,0 +1,212 @@ +"""CUDA unit tests for migrated MoE accuracy-compatible production paths.""" +from __future__ import annotations + +import ast +import unittest +from pathlib import Path +from types import SimpleNamespace +from typing import Optional + +import torch + +ROOT = Path(__file__).resolve().parents[4] +_UAC = {"on": False} + + +def _use_accuracy_compatible(): + return _UAC["on"] + + +class MoECudaGraphPartialCaptureSignal(Exception): + pass + + +class NVLSAllGatherVDispatcher: + pass + + +def _load_fn(rel: str, name: str): + src = (ROOT / rel).read_text() + tree = ast.parse(src) + class_name = ("MoEFlexTokenDispatcher" if rel.endswith("token_dispatcher.py") + else "MoELayer" if rel.endswith("moe_layer.py") else None) + body = tree.body + if class_name: + body = next(node.body for node in tree.body + if isinstance(node, ast.ClassDef) and node.name == class_name) + target = next(node for node in body if isinstance(node, ast.FunctionDef) and node.name == name) + target.decorator_list = [] + mod = ast.Module(body=[target], type_ignores=[]) + ast.fix_missing_locations(mod) + ns = { + "torch": torch, + "Optional": Optional, + "_use_accuracy_compatible": _use_accuracy_compatible, + "MoECudaGraphPartialCaptureSignal": MoECudaGraphPartialCaptureSignal, + "NVLSAllGatherVDispatcher": NVLSAllGatherVDispatcher, + "Tuple": tuple, + } + exec(compile(mod, rel, "exec"), ns) + return ns[name] + + +token_dispatch = _load_fn("megatron/core/transformer/moe/token_dispatcher.py", "token_dispatch") +token_combine = _load_fn("megatron/core/transformer/moe/token_dispatcher.py", "token_combine") +moe_forward = _load_fn("megatron/core/transformer/moe/moe_layer.py", "forward") +moe_postprocess = _load_fn("megatron/core/transformer/moe/moe_layer.py", "postprocess") +topk_routing_with_score_function = _load_fn( + "megatron/core/transformer/moe/moe_utils.py", "topk_routing_with_score_function" +) + + +class _FakeComm: + def __init__(self): + self.calls = [] + self.dispatched_probs = None + + def dispatch(self, hidden, async_finish, allocate_on_comm_stream): + self.calls.append(("dispatch", async_finish, allocate_on_comm_stream)) + self.dispatched_probs = hidden + return hidden + + def combine(self, hidden, async_finish, allocate_on_comm_stream): + self.calls.append(("combine", async_finish, allocate_on_comm_stream)) + return hidden + + def combine_postprocess(self, output): + return output + + +class _Flex: + def __init__(self): + self.shared_experts = None + self._comm_manager = _FakeComm() + + token_dispatch = token_dispatch + token_combine = token_combine + + +class _MoE: + def __init__(self): + self.training = False + self.attn_tp_group = SimpleNamespace(size=lambda: 1) + self.config = SimpleNamespace( + sequence_parallel=True, moe_shared_expert_overlap=False, moe_latent_size=0, fp8=False, fp4=False + ) + self.token_dispatcher = SimpleNamespace(combine_postprocess=lambda x: x) + self.shared_expert_overlap = False + self.fwd_execution_map = {"route", "expert_compute", "postprocess"} + self.moe_layer_recompute = False + self._accuracy_shared_input = None + self.order = [] + + def shared_experts_compute(self, x): + self.order.append("shared") + return x * 2 + + def route(self, x, padding_mask): + self.order.append("route") + return x, x + + def preprocess(self, x, probs, routing_map): + self.order.append("preprocess") + return x, probs + + def dispatch(self, x, probs): + self.order.append("dispatch") + return x, probs + + def routed_experts_compute(self, x, probs): + self.order.append("routed") + return x * 3, None + + def combine(self, x): + self.order.append("combine") + return x + + postprocess = moe_postprocess + forward = moe_forward + + +class TestAccuracyMigration(unittest.TestCase): + @classmethod + def setUpClass(cls): + if not torch.cuda.is_available(): + raise RuntimeError("CUDA unavailable") + torch.cuda.set_device(0) + + def setUp(self): + torch.cuda.set_device(0) + if not torch.cuda.is_available(): + raise RuntimeError("CUDA unavailable") + + def test_token_dispatch_combine_flags(self): + flex = _Flex() + x = torch.ones(2, 4, device="cuda") + _UAC["on"] = True + flex.token_dispatch(x, async_finish=True, allocate_on_comm_stream=True) + flex.token_combine(x, async_finish=True, allocate_on_comm_stream=True) + self.assertEqual(flex._comm_manager.calls[0][1:], (False, False)) + self.assertEqual(flex._comm_manager.calls[1][1:], (False, False)) + flex._comm_manager.calls.clear() + _UAC["on"] = False + flex.token_dispatch(x, async_finish=True, allocate_on_comm_stream=False) + flex.token_combine(x, async_finish=False, allocate_on_comm_stream=True) + self.assertEqual(flex._comm_manager.calls[0][1:], (True, False)) + self.assertEqual(flex._comm_manager.calls[1][1:], (False, True)) + + def test_forward_shared_order_and_value(self): + layer = _MoE() + x = torch.ones(2, 3, 4, device="cuda", requires_grad=True) + _UAC["on"] = True + out, _ = layer.forward(x) + self.assertEqual(layer.order, ["route", "preprocess", "dispatch", "routed", "combine", "shared"]) + self.assertTrue(torch.equal(out, 5 * x.detach())) + out.sum().backward() + self.assertTrue(torch.equal(x.grad, torch.full_like(x, 5))) + self.assertIsNone(getattr(layer, "_accuracy_shared_input", None)) + layer.order.clear() + x2 = torch.ones(2, 3, 4, device="cuda", requires_grad=True) + out2, _ = layer.forward(x2) + self.assertTrue(torch.equal(out2, 5 * x2.detach())) + self.assertIsNone(getattr(layer, "_accuracy_shared_input", None)) + layer = _MoE() + _UAC["on"] = False + x3 = torch.ones(2, 3, 4, device="cuda", requires_grad=True) + out3, _ = layer.forward(x3) + self.assertEqual(layer.order, ["shared", "route", "preprocess", "dispatch", "routed", "combine"]) + self.assertTrue(torch.equal(out3, 5 * x3.detach())) + + def test_postprocess_mixed_dtype(self): + layer = _MoE() + routed = torch.ones(2, 4, device="cuda", dtype=torch.bfloat16, requires_grad=True) + shared = torch.ones(2, 4, device="cuda", dtype=torch.float32, requires_grad=True) + _UAC["on"] = True + out = layer.postprocess(routed, shared) + self.assertEqual(out.dtype, torch.bfloat16) + out.float().sum().backward() + self.assertIsNotNone(routed.grad) + self.assertIsNotNone(shared.grad) + routed2 = torch.ones(2, 4, device="cuda", dtype=torch.bfloat16, requires_grad=True) + shared2 = torch.ones(2, 4, device="cuda", dtype=torch.float32, requires_grad=True) + _UAC["on"] = False + out2 = layer.postprocess(routed2, shared2) + self.assertEqual(out2.dtype, torch.float32) + + def test_topk_routing_output_dtype_and_gradients(self): + _UAC["on"] = True + logits = torch.tensor([[1.0, 2.0, 0.5, 3.0], [0.2, 4.0, 1.5, 0.8]], device="cuda", dtype=torch.bfloat16) + logits = logits.clone().requires_grad_(True) + for score in ("sigmoid", "sqrtsoftplus"): + probs, _idx = topk_routing_with_score_function( + logits, topk=2, score_function=score, dense_output=True, fused=False, router_replay=None + ) + self.assertEqual(probs.dtype, logits.dtype) + self.assertTrue(torch.allclose(probs.float().sum(dim=-1), torch.ones(probs.size(0), device="cuda"), atol=0.01)) + probs.sum().backward() + self.assertTrue(torch.isfinite(logits.grad).all()) + logits.grad = None + + +if __name__ == "__main__": + unittest.main() From 9716115ebc5ec0950a5b291b282e8a7009db216e Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Tue, 8 Sep 2026 08:22:12 -0700 Subject: [PATCH 20/27] Preserve TP1 embedding and parallel autograd behavior Retain embedding gradient accumulation and TP1 autograd structure in accuracy mode, with focused native behavior tests. Signed-off-by: Zhan Rongrui --- megatron/core/pipeline_parallel/schedules.py | 16 +- megatron/core/tensor_parallel/layers.py | 83 +++- megatron/core/tensor_parallel/mappings.py | 4 + .../test_accuracy_tp1_migration.py | 374 ++++++++++++++++++ 4 files changed, 473 insertions(+), 4 deletions(-) create mode 100644 tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py diff --git a/megatron/core/pipeline_parallel/schedules.py b/megatron/core/pipeline_parallel/schedules.py index e67c498e2cc..b75fe63ab8e 100644 --- a/megatron/core/pipeline_parallel/schedules.py +++ b/megatron/core/pipeline_parallel/schedules.py @@ -24,6 +24,7 @@ ProcessGroupCollection, ) from megatron.core.transformer.cuda_graphs import create_cudagraphs, set_current_microbatch +from megatron.core.transformer.module import _use_accuracy_compatible from megatron.core.transformer.moe.paged_stash import paged_stash_reset from megatron.core.transformer.moe.router import MoEAuxLossAutoScaler from megatron.core.utils import ( @@ -177,6 +178,11 @@ def deallocate_output_tensor(out, deallocate_pipeline_outputs=False): ''' if (out is None) or (not deallocate_pipeline_outputs): return + if _use_accuracy_compatible(): + # Compatibility fallback: callers supply only a tensor and deallocation flag. + _tp_size = int(parallel_state.get_tensor_model_parallel_world_size() or 1) + if _tp_size <= 1: + return # Handle dict format (multi-module pipelines) if isinstance(out, dict): @@ -568,7 +574,10 @@ def backward_step(input_tensor, output_tensor, output_tensor_grad, config): # This results in a tensor that does not require gradients. # In such cases, we intentionally skip the backward pass while preserving zero gradients. if output_tensor[0].requires_grad: - if config.deallocate_pipeline_outputs: + _tp_size = int(getattr(config, "tensor_model_parallel_size", 1) or 1) + if config.deallocate_pipeline_outputs and ( + not _use_accuracy_compatible() or _tp_size > 1 + ): custom_backward(output_tensor[0], output_tensor_grad[0]) else: torch.autograd.backward(output_tensor[0], grad_tensors=output_tensor_grad[0]) @@ -641,7 +650,10 @@ def _unwrap_single_tensor_list(tensor): # In multi-modal models like VLM, some batches may not have images. # In such cases, skip backward while preserving zero gradients. if output_tensor_module is not None and output_tensor_module.requires_grad: - if config.deallocate_pipeline_outputs: + _tp_size = int(getattr(config, "tensor_model_parallel_size", 1) or 1) + if config.deallocate_pipeline_outputs and ( + not _use_accuracy_compatible() or _tp_size > 1 + ): custom_backward(output_tensor_module, output_tensor_grad_module) else: torch.autograd.backward( diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 2e3f024c2b0..9bd553ef9dd 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -59,6 +59,53 @@ from megatron.core.transformer.module import _use_accuracy_compatible + +class _EmbedFp32MainGrad(torch.autograd.Function): + """UAC embedding lookup whose wgrad lands in fp32 main_grad. + + Forward is ``weight[ids]`` (bf16 activation unchanged). Backward uses a + unique-row clone plus ``autograd.grad``, then ``index_add_`` into an fp32 + ``main_grad``. Returns None for weight.grad so MixPrecision cannot add_(bf16). + """ + + @staticmethod + def forward(ctx, weight, ids): + ctx.save_for_backward(ids) + ctx.weight_ref = weight + return weight[ids] + + @staticmethod + def backward(ctx, grad_output): + (ids,) = ctx.saved_tensors + weight = ctx.weight_ref + prev = torch.is_grad_enabled() + torch.set_grad_enabled(True) + try: + ids_flat = ids.reshape(-1) + unique_ids, inv = torch.unique(ids_flat, return_inverse=True) + uniq_w = weight.detach()[unique_ids].clone().requires_grad_(True) + looked = uniq_w[inv.reshape(ids.shape)] + (gw,) = torch.autograd.grad( + looked, uniq_w, grad_outputs=grad_output, allow_unused=True + ) + finally: + torch.set_grad_enabled(prev) + if gw is None: + return None, None + fp = gw.float() + if hasattr(weight, "main_grad") and weight.main_grad is not None: + weight.main_grad.index_add_(0, unique_ids, fp) + else: + acc = torch.zeros( + weight.shape, dtype=torch.float32, device=weight.device + ) + acc.index_add_(0, unique_ids, fp) + weight.main_grad = acc + if hasattr(weight, "grad_added_to_main_grad"): + weight.grad_added_to_main_grad = True + return None, None + + _MODEL_PARALLEL_ATTRIBUTE_DEFAULTS = { "expert_tp": False, "is_qkv": False, @@ -299,7 +346,15 @@ def forward(self, input_): masked_input = input_ # Get the embeddings. if self.deterministic_mode: - output_parallel = self.weight[masked_input] + _tp_size = 1 if self.tp_group is None else self.tp_group.size() + if ( + _use_accuracy_compatible() + and _tp_size <= 1 + and os.environ.get("MODEL_REPRO_TWO_FP32_ACCUM", "") == "1" + ): + output_parallel = _EmbedFp32MainGrad.apply(self.weight, masked_input) + else: + output_parallel = self.weight[masked_input] else: # F.embedding currently has a non-deterministic backward function output_parallel = F.embedding(masked_input, self.weight) @@ -740,6 +795,17 @@ def linear_with_grad_accumulation_and_async_allreduce( """ tp_group = get_tensor_model_parallel_group_if_none(tp_group) + _tp_size = 1 if tp_group is None else tp_group.size() + if ( + _use_accuracy_compatible() + and _tp_size <= 1 + and not sequence_parallel + and not allreduce_dgrad + ): + output = torch.matmul(input, weight.t()) + if bias is not None: + output = output + bias + return output args = [ input, @@ -1072,6 +1138,11 @@ def forward( or self.disable_grad_reduce ): input_parallel = input_ + elif ( + _use_accuracy_compatible() + and (self.tp_group is None or self.tp_group.size() <= 1) + ): + input_parallel = input_ else: input_parallel = copy_to_tensor_model_parallel_region(input_, group=self.tp_group) @@ -1117,7 +1188,10 @@ def forward( if runtime_gather_output is not None: gather_output = runtime_gather_output - if gather_output: + if gather_output and ( + not _use_accuracy_compatible() + or (self.tp_group is not None and self.tp_group.size() > 1) + ): # All-gather across the partitions. if self.use_inference_optimized_all_gather and not self.training: # Deferred to avoid circular import: inference_layers → TE → layers. @@ -1396,6 +1470,11 @@ def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: output_ = reduce_scatter_to_sequence_parallel_region( output_parallel, group=self.tp_group ) + elif ( + _use_accuracy_compatible() + and (self.tp_group is None or self.tp_group.size() <= 1) + ): + output_ = output_parallel else: output_ = reduce_from_tensor_model_parallel_region(output_parallel, group=self.tp_group) if not self.skip_bias_add: diff --git a/megatron/core/tensor_parallel/mappings.py b/megatron/core/tensor_parallel/mappings.py index 6a1605d08a7..267ec7d61a8 100644 --- a/megatron/core/tensor_parallel/mappings.py +++ b/megatron/core/tensor_parallel/mappings.py @@ -510,6 +510,10 @@ def scatter_to_tensor_model_parallel_region(input_, group=None): def gather_from_tensor_model_parallel_region(input_, group=None): """Wrapper for autograd function: forward: AG, backward: split """ group = get_tensor_model_parallel_group_if_none(group) + from megatron.core.transformer.module import _use_accuracy_compatible + + if _use_accuracy_compatible() and (group is None or group.size() <= 1): + return input_ return _GatherFromModelParallelRegion.apply(input_, group) diff --git a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py new file mode 100644 index 00000000000..3399ed7cf60 --- /dev/null +++ b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py @@ -0,0 +1,374 @@ +"""CUDA unit tests for native C2 TP1 accuracy-compatible migration.""" +from __future__ import annotations + +import ast +import sys +import unittest +from pathlib import Path +from types import ModuleType, SimpleNamespace +from typing import List, Optional +from unittest.mock import patch + +import torch +import torch.nn.functional as F + +ROOT = Path(__file__).resolve().parents[3] +_UAC = {"on": False} +_TP = {"size": 1} +_CUSTOM_BWD = {"calls": []} + + +def _use_accuracy_compatible(): + return _UAC["on"] + + +def _custom_backward(output, grad): + _CUSTOM_BWD["calls"].append((output, grad)) + output.backward(grad) + + +class _SentinelApply(torch.autograd.Function): + last = None + + @staticmethod + def forward(ctx, *args): + _SentinelApply.last = args + inp = args[0] + return inp.new_zeros(inp.shape[:-1] + (args[1].shape[0],)) + + @staticmethod + def backward(ctx, grad_output): + return (grad_output,) + (None,) * 8 + + +class _SentinelGather(torch.autograd.Function): + last = None + + @staticmethod + def forward(ctx, input_, group): + _SentinelGather.last = (input_, group) + return input_ * 2 + + @staticmethod + def backward(ctx, grad_output): + return grad_output, None + + +class _FakeGroup: + def __init__(self, size): + self._size = size + + def size(self): + return self._size + + def rank(self): + return 0 + + +def _load_named(rel: str, name: str, extra_ns=None, class_name=None): + src = (ROOT / rel).read_text() + tree = ast.parse(src) + body = tree.body + if class_name: + body = next( + node.body + for node in tree.body + if isinstance(node, ast.ClassDef) and node.name == class_name + ) + target = next( + node + for node in body + if isinstance(node, ast.ClassDef) and node.name == name + or isinstance(node, ast.FunctionDef) and node.name == name + ) + if isinstance(target, ast.FunctionDef): + target.decorator_list = [] + mod = ast.Module(body=[target], type_ignores=[]) + ast.fix_missing_locations(mod) + ns = { + "torch": torch, + "F": F, + "Optional": Optional, + "List": List, + "os": __import__("os"), + "warnings": __import__("warnings"), + "_use_accuracy_compatible": _use_accuracy_compatible, + "get_tensor_model_parallel_group_if_none": lambda g: g, + "LinearWithGradAccumulationAndAsyncCommunication": _SentinelApply, + "_GatherFromModelParallelRegion": _SentinelGather, + "parallel_state": SimpleNamespace( + get_tensor_model_parallel_world_size=lambda: _TP["size"] + ), + "custom_backward": _custom_backward, + "Variable": torch.autograd.Variable, + } + if extra_ns: + ns.update(extra_ns) + exec(compile(mod, rel, "exec"), ns) + return ns[name] + + +_EmbedFp32MainGrad = _load_named( + "megatron/core/tensor_parallel/layers.py", + "_EmbedFp32MainGrad", +) +linear_with_grad_accumulation_and_async_allreduce = _load_named( + "megatron/core/tensor_parallel/layers.py", + "linear_with_grad_accumulation_and_async_allreduce", +) +linear_with_grad_accumulation_and_async_allreduce.warned = True +gather_from_tensor_model_parallel_region = _load_named( + "megatron/core/tensor_parallel/mappings.py", + "gather_from_tensor_model_parallel_region", +) +deallocate_output_tensor = _load_named( + "megatron/core/pipeline_parallel/schedules.py", + "deallocate_output_tensor", +) +backward_step = _load_named( + "megatron/core/pipeline_parallel/schedules.py", + "backward_step", +) + + +def _ref_embed_fp32_wgrad(weight_bf16, ids, grad_out): + table = weight_bf16.detach().clone().requires_grad_(True) + looked = F.embedding(ids, table) + (gw,) = torch.autograd.grad(looked, table, grad_outputs=grad_out) + return gw.float() + + +def _cuda_bf16(values, shape, device): + return torch.tensor(values, device=device, dtype=torch.bfloat16).reshape(shape) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestEmbedFp32MainGradCuda(unittest.TestCase): + def test_repeated_indices_matches_independent_full_table_autograd(self): + device = torch.device("cuda") + vocab, dim = 8, 4 + weight = _cuda_bf16( + [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, + 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32], + (vocab, dim), + device, + ) + ids = torch.tensor([[1, 3, 1, 5], [3, 3, 0, 1]], device=device) + grad_out = _cuda_bf16( + list(range(1, 33)), (2, 4, dim), device + ) + w = weight.clone().requires_grad_(True) + w.main_grad = torch.zeros(vocab, dim, device=device, dtype=torch.float32) + w.grad_added_to_main_grad = False + out = _EmbedFp32MainGrad.apply(w, ids) + self.assertEqual(out.dtype, torch.bfloat16) + torch.testing.assert_close(out, w[ids], atol=0, rtol=0) + out.backward(grad_out) + ref = _ref_embed_fp32_wgrad(weight, ids, grad_out) + torch.testing.assert_close(w.main_grad, ref, atol=0, rtol=0) + self.assertIsNone(w.grad) + self.assertTrue(w.grad_added_to_main_grad) + unused = [i for i in range(vocab) if i not in set(ids.reshape(-1).tolist())] + self.assertTrue((w.main_grad[unused] == 0).all()) + + def test_main_grad_accumulates_two_backwards_unused_rows_stay_zero(self): + device = torch.device("cuda") + vocab, dim = 6, 3 + weight = _cuda_bf16( + list(range(1, 19)), (vocab, dim), device + ) + ids_a = torch.tensor([2, 2, 4], device=device) + ids_b = torch.tensor([4, 1, 2], device=device) + go_a = _cuda_bf16([1, 2, 3, 4, 5, 6, 7, 8, 9], (3, dim), device) + go_b = _cuda_bf16([2, 1, 0, 1, 2, 3, 4, 5, 6], (3, dim), device) + w = weight.clone().requires_grad_(True) + w.main_grad = torch.zeros(vocab, dim, device=device, dtype=torch.float32) + w.grad_added_to_main_grad = False + _EmbedFp32MainGrad.apply(w, ids_a).backward(go_a) + first = w.main_grad.clone() + self.assertTrue(w.grad_added_to_main_grad) + _EmbedFp32MainGrad.apply(w, ids_b).backward(go_b) + ref = _ref_embed_fp32_wgrad(weight, ids_a, go_a) + _ref_embed_fp32_wgrad( + weight, ids_b, go_b + ) + torch.testing.assert_close(w.main_grad, ref, atol=0, rtol=0) + self.assertFalse(torch.equal(first, w.main_grad)) + self.assertTrue((w.main_grad[0] == 0).all()) + self.assertTrue((w.main_grad[3] == 0).all()) + self.assertTrue((w.main_grad[5] == 0).all()) + self.assertIsNone(w.grad) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestLinearTp1Native(unittest.TestCase): + def setUp(self): + _UAC["on"] = True + _SentinelApply.last = None + + def tearDown(self): + _UAC["on"] = False + + def test_tp1_forward_dgrad_wgrad_bias_matches_f_linear(self): + device = torch.device("cuda") + x = _cuda_bf16( + [1, 2, 0, -1, 1, 0, 2, 1, -2, 0, 1, 1, 2, 0, 1], + (5, 3), + device, + ).requires_grad_(True) + w = _cuda_bf16( + [1, 0, -1, 0, 1, 1, 1, -1, 0, 0, 1, -1], + (4, 3), + device, + ).requires_grad_(True) + b = _cuda_bf16([1, -1, 0, 2], (4,), device).requires_grad_(True) + out = linear_with_grad_accumulation_and_async_allreduce( + x, w, b, False, False, False, None, 0, _FakeGroup(1) + ) + xref = x.detach().clone().requires_grad_(True) + wref = w.detach().clone().requires_grad_(True) + bref = b.detach().clone().requires_grad_(True) + ref = F.linear(xref, wref, bref) + torch.testing.assert_close(out, ref, atol=0, rtol=0) + self.assertIsNone(_SentinelApply.last) + go = _cuda_bf16([1, 1, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 0, 1, 1], (5, 4), device) + out.backward(go) + F.linear(xref, wref, bref).backward(go) + torch.testing.assert_close(x.grad, xref.grad, atol=0, rtol=0) + torch.testing.assert_close(w.grad, wref.grad, atol=0, rtol=0) + torch.testing.assert_close(b.grad, bref.grad, atol=0, rtol=0) + + def test_off_delegates_to_native_custom_function(self): + _UAC["on"] = False + device = torch.device("cuda") + x = torch.ones(2, 3, device=device, dtype=torch.bfloat16, requires_grad=True) + w = torch.ones(4, 3, device=device, dtype=torch.bfloat16) + linear_with_grad_accumulation_and_async_allreduce( + x, w, None, False, False, False, None, 0, _FakeGroup(1) + ) + self.assertIsNotNone(_SentinelApply.last) + self.assertTrue(torch.equal(_SentinelApply.last[0], x)) + + def test_tp2_or_allreduce_skips_tp1_matmul_path(self): + device = torch.device("cuda") + x = torch.ones(2, 3, device=device, dtype=torch.bfloat16, requires_grad=True) + w = torch.ones(4, 3, device=device, dtype=torch.bfloat16) + linear_with_grad_accumulation_and_async_allreduce( + x, w, None, False, False, False, None, 0, _FakeGroup(2) + ) + self.assertIsNotNone(_SentinelApply.last) + _SentinelApply.last = None + linear_with_grad_accumulation_and_async_allreduce( + x, w, None, False, True, False, None, 0, _FakeGroup(1) + ) + self.assertIsNotNone(_SentinelApply.last) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestGatherTp1Identity(unittest.TestCase): + def setUp(self): + # The production wrapper imports the gate lazily; isolate only that + # dependency, retaining the actual gather wrapper and routing branch. + module = ModuleType("megatron.core.transformer.module") + module._use_accuracy_compatible = _use_accuracy_compatible + replacement = patch.dict(sys.modules, {module.__name__: module}) + replacement.start() + self.addCleanup(replacement.stop) + + def tearDown(self): + _UAC["on"] = False + _SentinelGather.last = None + + def test_uac_tp1_returns_same_tensor(self): + _UAC["on"] = True + t = torch.arange(6.0, device="cuda").reshape(2, 3) + out = gather_from_tensor_model_parallel_region(t, _FakeGroup(1)) + self.assertIs(out, t) + self.assertIsNone(_SentinelGather.last) + out_none = gather_from_tensor_model_parallel_region(t, None) + self.assertIs(out_none, t) + + def test_tp2_or_off_delegates_to_gather_function(self): + _UAC["on"] = True + t = torch.arange(6.0, device="cuda").reshape(2, 3) + out = gather_from_tensor_model_parallel_region(t, _FakeGroup(2)) + self.assertIsNotNone(_SentinelGather.last) + self.assertTrue(torch.equal(out, t * 2)) + _SentinelGather.last = None + _UAC["on"] = False + gather_from_tensor_model_parallel_region(t, _FakeGroup(1)) + self.assertIsNotNone(_SentinelGather.last) + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestPipelineHelpersTp1(unittest.TestCase): + def tearDown(self): + _UAC["on"] = False + _TP["size"] = 1 + _CUSTOM_BWD["calls"] = [] + + def test_deallocate_uac_tp1_preserves_storage(self): + _UAC["on"] = True + _TP["size"] = 1 + t = torch.arange(4.0, device="cuda", requires_grad=True) + data_before = t.data.clone() + deallocate_output_tensor(t, True) + torch.testing.assert_close(t.data, data_before, atol=0, rtol=0) + self.assertEqual(tuple(t.shape), (4,)) + + def test_deallocate_off_or_tp2_still_frees(self): + _UAC["on"] = False + _TP["size"] = 1 + t = torch.arange(4.0, device="cuda") + deallocate_output_tensor(t, True) + self.assertEqual(tuple(t.shape), (1,)) + _UAC["on"] = True + _TP["size"] = 2 + t2 = torch.arange(4.0, device="cuda") + deallocate_output_tensor(t2, True) + self.assertEqual(tuple(t2.shape), (1,)) + + def test_backward_step_uac_tp1_uses_autograd_not_custom(self): + _UAC["on"] = True + _CUSTOM_BWD["calls"] = [] + x = torch.tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], device="cuda", requires_grad=True) + y = x * 3 + go = torch.ones_like(y) + cfg = SimpleNamespace( + timers=None, + grad_scale_func=None, + deallocate_pipeline_outputs=True, + tensor_model_parallel_size=1, + ) + gin = backward_step(x, y, go, cfg) + torch.testing.assert_close(gin, go * 3, atol=0, rtol=0) + self.assertEqual(_CUSTOM_BWD["calls"], []) + + def test_backward_step_off_or_tp2_uses_custom_backward(self): + x = torch.tensor([[1.0, 2.0], [3.0, 4.0]], device="cuda", requires_grad=True) + y = x * 2 + go = torch.ones_like(y) + cfg = SimpleNamespace( + timers=None, + grad_scale_func=None, + deallocate_pipeline_outputs=True, + tensor_model_parallel_size=1, + ) + _UAC["on"] = False + _CUSTOM_BWD["calls"] = [] + gin = backward_step(x, y, go, cfg) + self.assertEqual(len(_CUSTOM_BWD["calls"]), 1) + torch.testing.assert_close(gin, go * 2, atol=0, rtol=0) + + x2 = torch.tensor([[1.0, 2.0], [3.0, 4.0]], device="cuda", requires_grad=True) + y2 = x2 * 2 + go2 = torch.ones_like(y2) + cfg.tensor_model_parallel_size = 2 + _UAC["on"] = True + _CUSTOM_BWD["calls"] = [] + gin2 = backward_step(x2, y2, go2, cfg) + self.assertEqual(len(_CUSTOM_BWD["calls"]), 1) + torch.testing.assert_close(gin2, go2 * 2, atol=0, rtol=0) + + +if __name__ == "__main__": + unittest.main() From b0c2cd099283b365f4dd115d44ffc89d56f06fd7 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Tue, 8 Sep 2026 08:23:15 -0700 Subject: [PATCH 21/27] Preserve native MTP and transformer accuracy graph Preserve MTP loss attachment, rotary indexing, transformer exit and attention graph behavior in native accuracy mode. Signed-off-by: Zhan Rongrui --- .../experimental_attention_variant/dsa.py | 9 ++-- .../transformer/multi_token_prediction.py | 44 +++++++++---------- .../core/transformer/transformer_block.py | 20 ++++++--- .../core/transformer/transformer_layer.py | 11 +++-- 4 files changed, 48 insertions(+), 36 deletions(-) diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index c7f3a212558..6c42eb358c1 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -22,7 +22,7 @@ dsa_layout, dsa_masking, ) -from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_config import TransformerConfig @@ -1827,8 +1827,11 @@ def forward( skv = key.size(0) # Detach x and qr to prevent gradients of indexer from flowing back to the main model. - x = x.detach() - qr = qr.detach() + _tp_group = getattr(self.pg_collection, "tp", None) + _tp_size = 1 if _tp_group is None else _tp_group.size() + if not (_use_accuracy_compatible() and _tp_size <= 1): + x = x.detach() + qr = qr.detach() indexer_loss_coeff = self.config.dsa_indexer_loss_coeff or 0.0 computes_topk = not self.skip_topk diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index ee92afe4954..c35747cda95 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -11,10 +11,7 @@ from megatron.core import InferenceParams, parallel_state, tensor_parallel from megatron.core.dist_checkpointing.mapping import ShardedStateDict -from megatron.core.dist_checkpointing.utils import ( - apply_prefix_mapping, - replace_prefix_for_sharding, -) +from megatron.core.dist_checkpointing.utils import apply_prefix_mapping, replace_prefix_for_sharding from megatron.core.enums import Fp8Recipe from megatron.core.extensions.transformer_engine import HAVE_TE from megatron.core.fp8_utils import get_fp8_context @@ -31,7 +28,7 @@ inference_all_gather_from_tensor_model_parallel_region, ) from megatron.core.transformer.enums import AttnMaskType, LayerType -from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder, WrappedTorchNorm from megatron.core.transformer.transformer_block import TransformerBlockSubmodules @@ -64,9 +61,7 @@ else: TESpecProvider = None -from megatron.core.transformer.pipeline_parallel_layer_layout import ( - PipelineParallelLayerLayout, -) +from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout def tie_word_embeddings_state_dict( @@ -1054,9 +1049,7 @@ def __init__( self.submodules.mtp_model_layer, "submodules" ): from megatron.core.models.hybrid.hybrid_block import HybridStackSubmodules - from megatron.core.transformer.transformer_layer import ( - TransformerLayerSubmodules, - ) + from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules layer_submodules = None if isinstance( @@ -1124,9 +1117,7 @@ def __init__( # 2. GPT path: single TransformerLayer if mtp_layer_pattern is not None and hybrid_submodules is not None: from megatron.core.models.hybrid.hybrid_block import HybridStack - from megatron.core.models.hybrid.hybrid_layer_allocation import ( - validate_segment_layers, - ) + from megatron.core.models.hybrid.hybrid_layer_allocation import validate_segment_layers self.mtp_model_layer = HybridStack( config=self.config, @@ -1215,9 +1206,11 @@ def _get_embeddings( if self.config.mtp_detach_heads: decoder_input = decoder_input.detach() - hidden_states = make_viewless_tensor( - inp=hidden_states, requires_grad=True, keep_graph=True - ) + _tp_size = 1 if self.tp_group is None else self.tp_group.size() + if not (_use_accuracy_compatible() and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # make_viewless_tensor no-ops when hidden_states is not a view (_base is None), # which happens after detach() with mtp_detach_heads. Activation # checkpointing (CheckpointFunction.apply) requires at least one input tensor @@ -1234,14 +1227,17 @@ def _concat_embeddings( """ Concatenate the tokens before sending to transformer layer. """ + _tp_size = 1 if self.tp_group is None else self.tp_group.size() decoder_input = apply_module(self.enorm)(decoder_input) - decoder_input = make_viewless_tensor( - inp=decoder_input, requires_grad=True, keep_graph=True - ) + if not (_use_accuracy_compatible() and _tp_size <= 1): + decoder_input = make_viewless_tensor( + inp=decoder_input, requires_grad=True, keep_graph=True + ) hidden_states = apply_module(self.hnorm)(hidden_states) - hidden_states = make_viewless_tensor( - inp=hidden_states, requires_grad=True, keep_graph=True - ) + if not (_use_accuracy_compatible() and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # At the (k - 1)-th MTP module, concatenates the i-th token's hidden_states # and the (i + K)-th token's embedding, and combine them with linear projection. hidden_states = torch.cat((decoder_input, hidden_states), -1) @@ -1252,7 +1248,7 @@ def _concat_embeddings( hidden_states = inference_all_gather_from_tensor_model_parallel_region( hidden_states, self.tp_group, self.config ) - else: + elif not (_use_accuracy_compatible() and _tp_size <= 1): hidden_states = gather_from_tensor_model_parallel_region( hidden_states, group=self.tp_group ) diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index 0415035ffbe..f01cf55c1fc 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -23,7 +23,11 @@ from megatron.core.recompute import checkpointed_forward from megatron.core.transformer.cuda_graphs import annotate_first_last_layer from megatron.core.transformer.enums import InferenceCudaGraphScope, LayerType -from megatron.core.transformer.module import GraphableMegatronModule, MegatronModule +from megatron.core.transformer.module import ( + GraphableMegatronModule, + MegatronModule, + _use_accuracy_compatible, +) from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder from megatron.core.transformer.transformer_config import TransformerConfig @@ -590,7 +594,11 @@ def forward( # likely redundant, since p2p_communication.py (likely originator) # already creates viewless tensors. That said, make_viewless_tensor() # is called here to be future-proof and corner-case-proof. - hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) + _tp_size = int(getattr(self.config, "tensor_model_parallel_size", 1) or 1) + if not (_use_accuracy_compatible() and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) if self.config.sequence_parallel: rng_context = tensor_parallel.get_cuda_rng_tracker().fork() @@ -694,9 +702,11 @@ def forward( # TENorm produces a "viewed" tensor. This will result in schedule.py's # deallocate_output_tensor() throwing an error, so a viewless tensor is # created to prevent this. - hidden_states = make_viewless_tensor( - inp=hidden_states, requires_grad=True, keep_graph=True - ) + _tp_size = int(getattr(self.config, "tensor_model_parallel_size", 1) or 1) + if not (_use_accuracy_compatible() and _tp_size <= 1): + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) # If this TransformerBlock is empty, input and output hidden states will be the same node # on the computational graph and will lead to unexpected errors in pipeline schedules. diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index 904912c18d8..82e2051ff5b 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -22,7 +22,7 @@ from megatron.core.transformer.enums import CudaGraphModule, InferenceCudaGraphScope, LayerType from megatron.core.transformer.identity_op import IdentityFuncOp, IdentityOp from megatron.core.transformer.mlp import MLP -from megatron.core.transformer.module import GraphableMegatronModule +from megatron.core.transformer.module import GraphableMegatronModule, _use_accuracy_compatible from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder from megatron.core.transformer.transformer_config import TransformerConfig @@ -941,9 +941,12 @@ def _forward_post_mlp( # won't result in memory savings (like the data loader, or # p2p_communication), it serves to document the origin of this # 'view' tensor. - output = make_viewless_tensor( - inp=hidden_states, requires_grad=hidden_states.requires_grad, keep_graph=True - ) + if _use_accuracy_compatible() and self.config.tensor_model_parallel_size <= 1: + output = hidden_states + else: + output = make_viewless_tensor( + inp=hidden_states, requires_grad=hidden_states.requires_grad, keep_graph=True + ) return output From d4bd982fa07ab3d512650d748e2416f5c26c863f Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Tue, 8 Sep 2026 22:04:40 -0700 Subject: [PATCH 22/27] Remove unrelated CI changes from GLM alignment PR Signed-off-by: Zhan Rongrui --- .../workflows/alignment_model_accuracy.yml | 133 +-- scripts/consume_paddlefleet_alignment_pin.sh | 370 -------- scripts/prepare_paddlefleet_ops_submodules.sh | 29 - scripts/require_paddlefleet_selector_ok.sh | 40 - scripts/select_paddlefleet_alignment_pin.sh | 889 ------------------ scripts/test_alignment_workflow_shell.sh | 153 --- scripts/test_ops_submodule_preflight.sh | 78 -- scripts/test_paddlefleet_pin_handoff.sh | 111 --- 8 files changed, 23 insertions(+), 1780 deletions(-) delete mode 100755 scripts/consume_paddlefleet_alignment_pin.sh delete mode 100644 scripts/prepare_paddlefleet_ops_submodules.sh delete mode 100755 scripts/require_paddlefleet_selector_ok.sh delete mode 100755 scripts/select_paddlefleet_alignment_pin.sh delete mode 100755 scripts/test_alignment_workflow_shell.sh delete mode 100644 scripts/test_ops_submodule_preflight.sh delete mode 100755 scripts/test_paddlefleet_pin_handoff.sh diff --git a/.github/workflows/alignment_model_accuracy.yml b/.github/workflows/alignment_model_accuracy.yml index 8bf2696fd85..185330c16a9 100644 --- a/.github/workflows/alignment_model_accuracy.yml +++ b/.github/workflows/alignment_model_accuracy.yml @@ -3,36 +3,6 @@ name: Alignment Model Accuracy on: pull_request: workflow_dispatch: - inputs: - paddlefleet_mode: - description: "develop keeps the historical CodeSync tarball. stack-paired fail-closed SHA check." - type: choice - default: develop - options: [develop, stack-paired] - paddlefleet_pin_sha: - description: "40-hex PaddleFleet commit (required for stack-paired)" - required: false - type: string - paddlefleet_git_url: - description: "Git URL that contains paddlefleet_pin_sha" - required: false - type: string - paddlefleet_wheel_url: - description: "Optional wheel URL from Build Fleet whl / Actions artifact metadata" - required: false - type: string - paddlefleet_wheel_sha256: - description: "sha256 for paddlefleet_wheel_url" - required: false - type: string - paddlefleet_ops_wheel_url: - description: "Optional paddlefleet_ops wheel URL" - required: false - type: string - paddlefleet_ops_wheel_sha256: - description: "sha256 for paddlefleet_ops_wheel_url" - required: false - type: string concurrency: group: Alignment-${{ github.workflow }}-${{ github.event.pull_request.number }} @@ -47,14 +17,6 @@ env: TASK: Megatron-LM-${{ github.sha }}-alignment CE_name: alignment-Megatron-LM no_proxy: "localhost,bj.bcebos.com,su.bcebos.com,bcebos.com,apiin.im.baidu.com,gitee.com,aliyun.com,.baidu.com,.tuna.tsinghua.edu.cn" - # pull_request stays on historical develop. stack-paired is workflow_dispatch only. - ALIGNMENT_PADDLEFLEET_MODE: ${{ github.event_name == 'workflow_dispatch' && github.event.inputs.paddlefleet_mode || 'develop' }} - PADDLEFLEET_PIN_SHA: ${{ github.event.inputs.paddlefleet_pin_sha }} - PADDLEFLEET_GIT_URL: ${{ github.event.inputs.paddlefleet_git_url }} - PADDLEFLEET_WHEEL_URL: ${{ github.event.inputs.paddlefleet_wheel_url }} - PADDLEFLEET_WHEEL_SHA256: ${{ github.event.inputs.paddlefleet_wheel_sha256 }} - PADDLEFLEET_OPS_WHEEL_URL: ${{ github.event.inputs.paddlefleet_ops_wheel_url }} - PADDLEFLEET_OPS_WHEEL_SHA256: ${{ github.event.inputs.paddlefleet_ops_wheel_sha256 }} defaults: run: @@ -93,28 +55,17 @@ jobs: -e no_proxy \ -e CE_name \ -e python_version \ - -e ALIGNMENT_PADDLEFLEET_MODE \ - -e PADDLEFLEET_PIN_SHA \ - -e PADDLEFLEET_GIT_URL \ - -e PADDLEFLEET_WHEEL_URL \ - -e PADDLEFLEET_WHEEL_SHA256 \ - -e PADDLEFLEET_OPS_WHEEL_URL \ - -e PADDLEFLEET_OPS_WHEEL_SHA256 \ -w /workspace $IMAGE_NAME - name: Checkout Code run: | docker exec -t $container_name /bin/bash -c ' rm -rf * .[^.]* source $work_dir/../../../proxy - if [ "${ALIGNMENT_PADDLEFLEET_MODE:-develop}" = "stack-paired" ]; then - echo "stack-paired: skip CodeSync/develop PaddleFleet.tar; selector runs after python" - else - echo "Download PaddleFleet form https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar" - wget -q --no-proxy https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar --no-check-certificate - rm -rf PaddleFleet && tar xf PaddleFleet.tar && rm -rf PaddleFleet.tar - cd PaddleFleet && git pull && cd - - fi - + echo "Download PaddleFleet form https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar" + wget -q --no-proxy https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar --no-check-certificate + rm -rf PaddleFleet && tar xf PaddleFleet.tar && rm -rf PaddleFleet.tar + cd PaddleFleet && git pull && cd - + echo "Download Megatron-LM form https://paddle-github-action.bj.bcebos.com/whl/Megatron-LM.tar.gz" wget -q --no-proxy https://paddle-github-action.bj.bcebos.com/whl/Megatron-LM.tar.gz --no-check-certificate rm -rf Megatron-LM && tar zxf Megatron-LM.tar.gz && rm -rf Megatron-LM.tar.gz @@ -122,25 +73,17 @@ jobs: git config --global --add safe.directory /workspace/Megatron-LM git pull git submodule update --init --recursive --force - git remote add upstream https://github.com/PFCCLab/Megatron-LM.git || true if [ -n "$PR_ID" ] && [ "$PR_ID" != "0" ]; then git fetch origin pull/${PR_ID}/head - git checkout -B PR_${PR_ID} FETCH_HEAD + git checkout -b PR_${PR_ID} FETCH_HEAD + git remote add upstream https://github.com/PFCCLab/Megatron-LM.git echo "Checking out ${BRANCH}..." - git fetch upstream ${BRANCH}:${BRANCH} || true - git merge ${BRANCH} --no-edit || true + git fetch upstream ${BRANCH}:${BRANCH} + git merge ${BRANCH} --no-edit git diff --numstat ${BRANCH} -- | awk "{print \$NF}" - elif [ -n "$COMMIT_ID" ]; then - echo "workflow_dispatch: checkout COMMIT_ID=${COMMIT_ID} (not BOS main)" - git fetch --all --tags || true - git fetch origin "$COMMIT_ID" || git fetch upstream "$COMMIT_ID" || true - git checkout --force "$COMMIT_ID" - git submodule update --init --recursive --force else - echo "Not in a pull_request event and COMMIT_ID empty. Leaving tarball HEAD." + echo "Not in a pull_request event. Skipping PR-specific operations." fi - echo "checked_out_head=$(git rev-parse HEAD)" - test -z "$COMMIT_ID" || test "$(git rev-parse HEAD)" = "$COMMIT_ID" git log --pretty=oneline -10 ' - name: Change python version @@ -161,45 +104,24 @@ jobs: - name: Get Whl run: | docker exec -t $container_name /bin/bash -c ' - set -eo pipefail . /opt/conda/etc/profile.d/conda.sh conda activate py_$python_version python --version ldconfig BOS=https://paddle-github-action.bj.bcebos.com + echo "::group::Download paddlefleet / paddlefleet_ops / ms_swift / mcore-bridge wheels from BOS" cd /workspace - if [ "${ALIGNMENT_PADDLEFLEET_MODE:-develop}" = "stack-paired" ]; then - echo "::group::Select stack-paired PaddleFleet (fail-closed SHA)" - python -m pip install uv - test -x /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh \ - || { echo "::error:: selector missing; checkout did not land COMMIT_ID"; exit 1; } - bash /workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh --dest /workspace \ - || { echo "::error:: selector failed; stop Get Whl (do not download remaining wheels or build megatron-core)"; exit 1; } - bash /workspace/Megatron-LM/scripts/require_paddlefleet_selector_ok.sh /workspace \ - || { echo "::error:: selector receipt status is not ok; stop Get Whl"; exit 1; } - echo "::endgroup::" - echo "::group::Download remaining wheels from BOS" - for url in \ - $BOS/whl/ms_swift-0.0.0-py3-none-any.whl \ - $BOS/whl/mcore_bridge-0.0.0-py3-none-any.whl ; do - echo "Downloading $url" - wget -q --no-proxy --no-check-certificate --tries=3 --timeout=60 "$url" - done - echo "::endgroup::" - else - echo "::group::Download paddlefleet / paddlefleet_ops / ms_swift / mcore-bridge wheels from BOS" - for url in \ - $BOS/PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl \ - $BOS/PaddleFleet/develop/latest/cu130/paddle-release/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl \ - $BOS/whl/ms_swift-0.0.0-py3-none-any.whl \ - $BOS/whl/mcore_bridge-0.0.0-py3-none-any.whl ; do - echo "Downloading $url" - wget -q --no-proxy --no-check-certificate --tries=3 --timeout=60 "$url" - done - echo "::endgroup::" - fi + for url in \ + $BOS/PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl \ + $BOS/PaddleFleet/develop/latest/cu130/paddle-release/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl \ + $BOS/whl/ms_swift-0.0.0-py3-none-any.whl \ + $BOS/whl/mcore_bridge-0.0.0-py3-none-any.whl ; do + echo "Downloading $url" + wget -q --no-proxy --no-check-certificate --tries=3 --timeout=60 "$url" + done ls -l /workspace/ + echo "::endgroup::" echo "::group::Build megatron-core wheel" cd /workspace/Megatron-LM @@ -211,21 +133,14 @@ jobs: - name: alignment_model_accuracy run: | docker exec -t $container_name /bin/bash -c ' - set -eo pipefail . /opt/conda/etc/profile.d/conda.sh conda activate py_$python_version python --version ldconfig source $work_dir/../../../proxy export PROXY_URL="${http_proxy}" - ALIGNMENT_PADDLEFLEET_MODE="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" \ - PADDLEFLEET_PIN_SHA="${PADDLEFLEET_PIN_SHA:-}" \ - bash /workspace/Megatron-LM/scripts/consume_paddlefleet_alignment_pin.sh \ - --env /workspace/paddlefleet_alignment_pin.env \ - --out /workspace/paddlefleet_alignment_pin.consumed.env - set -a - . /workspace/paddlefleet_alignment_pin.consumed.env - set +a + export PADDLEFLEET_WHEEL_PATH="/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" + export PADDLEFLEET_OPS_WHEEL_PATH="/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" export MEGATRON_CORE_WHEEL_PATH=/workspace/upload/megatron_core-0.0.0-cp312-cp312-linux_x86_64.whl export MS_SWIFT_WHEEL_PATH=/workspace/ms_swift-0.0.0-py3-none-any.whl export MCORE_BRIDGE_WHEEL_PATH=/workspace/mcore_bridge-0.0.0-py3-none-any.whl @@ -235,12 +150,10 @@ jobs: for whl in "$PADDLEFLEET_WHEEL_PATH" "$PADDLEFLEET_OPS_WHEEL_PATH" \ "$MS_SWIFT_WHEEL_PATH" "$MEGATRON_CORE_WHEEL_PATH" \ "$MCORE_BRIDGE_WHEEL_PATH"; do - [ -e "$whl" ] || { echo "::error:: missing wheel: $whl"; exit 1; } + [ -f "$whl" ] || { echo "::error:: missing wheel: $whl"; exit 1; } echo "using $whl" done python -m pip install uv - bash Megatron-LM/scripts/prepare_paddlefleet_ops_submodules.sh \ - "$PADDLEFLEET_OPS_WHEEL_PATH" PaddleFleet/scripts/alignment_model_accuracy/setup_venvs.sh bash PaddleFleet/scripts/alignment_model_accuracy/setup_venvs.sh bash -x PaddleFleet/scripts/alignment_model_accuracy/run_alignment_test.sh diff --git a/scripts/consume_paddlefleet_alignment_pin.sh b/scripts/consume_paddlefleet_alignment_pin.sh deleted file mode 100755 index 1ba0b8e56b5..00000000000 --- a/scripts/consume_paddlefleet_alignment_pin.sh +++ /dev/null @@ -1,370 +0,0 @@ -#!/usr/bin/env bash -# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. -# -# Consume selector output in a later docker exec. The caller mode/pin stay -# authoritative: a leftover develop env must not silently downgrade -# stack-paired. A 0.0.0 filename is allowed when the receipt proves an -# explicit URL+sha256 (or a source tree / in-invocation build). Unproven -# develop/latest fallback is refused. - -set -euo pipefail - -usage() { - cat <<'EOF' -Usage: consume_paddlefleet_alignment_pin.sh [--env FILE] [--out FILE] [--self-test] - -Reads paddlefleet_alignment_pin.env + receipt from select_paddlefleet_alignment_pin.sh -and writes a consumed env file for setup_venvs.sh. - -Caller ALIGNMENT_PADDLEFLEET_MODE / PADDLEFLEET_PIN_SHA are the request. -The env file must match that request; it does not override them. -EOF -} - -ENVFILE="${PADDLEFLEET_PIN_ENV:-/workspace/paddlefleet_alignment_pin.env}" -OUTFILE="${PADDLEFLEET_CONSUMED_ENV:-/workspace/paddlefleet_alignment_pin.consumed.env}" -RUN_SELF_TEST=0 -while [[ $# -gt 0 ]]; do - case "$1" in - --env) ENVFILE="${2:?}"; shift 2 ;; - --out) OUTFILE="${2:?}"; shift 2 ;; - --self-test) RUN_SELF_TEST=1; shift ;; - -h|--help) usage; exit 0 ;; - *) echo "unknown arg: $1" >&2; usage; exit 2 ;; - esac -done - -fail() { - echo "::error:: $*" >&2 - exit 1 -} - -parse_envfile() { - FILE_MODE="" - FILE_PIN="" - FILE_SOURCE="" - FILE_RECEIPT="" - FILE_WHEEL="" - FILE_OPS="" - FILE_WHEEL_ORIGIN="" - FILE_OPS_ORIGIN="" - FILE_WHEEL_DIGEST="" - FILE_OPS_DIGEST="" - [[ -f "${ENVFILE}" ]] || return 0 - local line k v - while IFS= read -r line || [[ -n "${line}" ]]; do - [[ -z "${line}" || "${line}" == \#* ]] && continue - k="${line%%=*}" - v="${line#*=}" - case "${k}" in - ALIGNMENT_PADDLEFLEET_MODE) FILE_MODE="${v}" ;; - PADDLEFLEET_PIN_SHA) FILE_PIN="${v}" ;; - PADDLEFLEET_SOURCE_COMMIT) FILE_SOURCE="${v}" ;; - PADDLEFLEET_PIN_RECEIPT) FILE_RECEIPT="${v}" ;; - PADDLEFLEET_WHEEL_PATH) FILE_WHEEL="${v}" ;; - PADDLEFLEET_OPS_WHEEL_PATH) FILE_OPS="${v}" ;; - PADDLEFLEET_WHEEL_ORIGIN) FILE_WHEEL_ORIGIN="${v}" ;; - PADDLEFLEET_OPS_ORIGIN) FILE_OPS_ORIGIN="${v}" ;; - PADDLEFLEET_WHEEL_DIGEST_VERIFIED) FILE_WHEEL_DIGEST="${v}" ;; - PADDLEFLEET_OPS_DIGEST_VERIFIED) FILE_OPS_DIGEST="${v}" ;; - esac - done <"${ENVFILE}" -} - -# Unproven develop fallback: develop_latest origin, or a default 0.0.0 -# filename with no digest proof. A verified ci_metadata/build artifact may -# legally keep the 0.0.0 filename. -unproven_develop_fallback() { - local path="$1" origin="$2" digest="$3" - case "${origin}" in - develop_latest) return 0 ;; - source_tree|build|ci_metadata) - [[ "${origin}" == develop_latest ]] && return 0 - return 1 - ;; - esac - case "${path}" in - */paddlefleet-0.0.0-py3-none-linux_x86_64.whl|*/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl) - [[ "${digest}" == "true" ]] && return 1 - return 0 - ;; - esac - return 1 -} - -check_receipt() { - local receipt="$1" requested_mode="$2" requested_pin="$3" - [[ -f "${receipt}" ]] || fail "stack-paired missing receipt ${receipt}" - python3 - "${receipt}" "${requested_mode}" "${requested_pin}" <<'PY' -import json, sys -path, requested_mode, requested_pin = sys.argv[1:4] -doc = json.load(open(path, encoding="utf-8")) -if doc.get("status") != "ok": - raise SystemExit(f"receipt status={doc.get('status')!r} is not ok") -if requested_mode == "stack-paired": - if doc.get("mode") != "stack-paired": - raise SystemExit( - f"requested stack-paired but receipt mode={doc.get('mode')!r}" - ) - src = doc.get("source") or {} - if requested_pin: - exp = src.get("expected_commit") or "" - act = src.get("actual_commit") or "" - if requested_pin not in (exp, act): - raise SystemExit( - f"receipt source pin mismatch requested={requested_pin} " - f"expected={exp} actual={act}" - ) - if src.get("commit_verified") is not True: - raise SystemExit("receipt source commit_verified is not true") - pairing = (doc.get("pairing") or {}).get("status") - if pairing == "unpaired_default": - raise SystemExit("receipt pairing.status=unpaired_default") - for art in doc.get("artifacts") or []: - origin = art.get("origin") or "" - url = art.get("url") or "" - if origin == "develop_latest": - raise SystemExit(f"artifact {art.get('name')} origin=develop_latest") - if "/develop/latest/" in url or "CodeSync/develop/" in url: - raise SystemExit(f"artifact {art.get('name')} unpaired develop URL") -print("receipt matches request") -PY -} - -write_consumed() { - mkdir -p "$(dirname "${OUTFILE}")" - cat >"${OUTFILE}" <&2 - echo "[paddlefleet-pin-consume] requested_mode=${ALIGNMENT_PADDLEFLEET_MODE} pin=${PADDLEFLEET_PIN_SHA:-} wheel=${PADDLEFLEET_WHEEL_PATH} origin=${PADDLEFLEET_WHEEL_ORIGIN:-}" >&2 -} - -consume() { - local requested_mode="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" - local requested_pin="${PADDLEFLEET_PIN_SHA:-}" - - parse_envfile - - if [[ "${requested_mode}" == "stack-paired" ]]; then - if [[ ! -f "${ENVFILE}" ]]; then - local rec="${ENVFILE%/*}/paddlefleet_alignment_pin_receipt.json" - local rec_status="" - if [[ -f "${rec}" ]]; then - rec_status="$(python3 - "${rec}" <<'PY' -import json, sys -print(json.load(open(sys.argv[1])).get("status") or "") -PY -)" - fi - if [[ "${rec_status}" == "error" ]]; then - fail "stack-paired missing ${ENVFILE} because selector failed (error receipt ${rec}); Get Whl must exit on selector failure. This is not proof that a generated env failed to cross docker exec" - fi - if [[ "${rec_status}" == "ok" ]]; then - fail "stack-paired missing ${ENVFILE}; selector wrote ok receipt but env did not cross docker exec" - fi - fail "stack-paired missing ${ENVFILE} and no selector receipt" - fi - [[ "${FILE_MODE}" == "stack-paired" ]] || fail "requested stack-paired but env mode=${FILE_MODE:-empty} (will not consume a develop leftover)" - if [[ -n "${requested_pin}" ]]; then - local file_id="${FILE_PIN:-${FILE_SOURCE}}" - [[ "${file_id}" == "${requested_pin}" ]] || fail "requested pin ${requested_pin} != env pin/source ${file_id:-empty}" - fi - local receipt="${FILE_RECEIPT:-}" - if [[ -z "${receipt}" && -f "${ENVFILE%/*}/paddlefleet_alignment_pin_receipt.json" ]]; then - receipt="${ENVFILE%/*}/paddlefleet_alignment_pin_receipt.json" - fi - check_receipt "${receipt}" "${requested_mode}" "${requested_pin}" - PADDLEFLEET_WHEEL_PATH="${FILE_WHEEL}" - PADDLEFLEET_OPS_WHEEL_PATH="${FILE_OPS}" - PADDLEFLEET_WHEEL_ORIGIN="${FILE_WHEEL_ORIGIN}" - PADDLEFLEET_OPS_ORIGIN="${FILE_OPS_ORIGIN}" - PADDLEFLEET_PIN_RECEIPT="${receipt}" - PADDLEFLEET_SOURCE_COMMIT="${FILE_SOURCE}" - PADDLEFLEET_PIN_SHA="${requested_pin:-${FILE_PIN}}" - ALIGNMENT_PADDLEFLEET_MODE="stack-paired" - [[ -n "${PADDLEFLEET_WHEEL_PATH}" && -n "${PADDLEFLEET_OPS_WHEEL_PATH}" ]] \ - || fail "stack-paired env missing PADDLEFLEET_WHEEL_PATH or OPS path" - if unproven_develop_fallback "${PADDLEFLEET_WHEEL_PATH}" "${FILE_WHEEL_ORIGIN}" "${FILE_WHEEL_DIGEST}"; then - fail "stack-paired refused unproven develop fallback for paddlefleet: path=${PADDLEFLEET_WHEEL_PATH} origin=${FILE_WHEEL_ORIGIN:-empty} digest_verified=${FILE_WHEEL_DIGEST:-false}" - fi - if unproven_develop_fallback "${PADDLEFLEET_OPS_WHEEL_PATH}" "${FILE_OPS_ORIGIN}" "${FILE_OPS_DIGEST}"; then - fail "stack-paired refused unproven develop fallback for paddlefleet_ops: path=${PADDLEFLEET_OPS_WHEEL_PATH} origin=${FILE_OPS_ORIGIN:-empty} digest_verified=${FILE_OPS_DIGEST:-false}" - fi - else - ALIGNMENT_PADDLEFLEET_MODE="develop" - PADDLEFLEET_WHEEL_PATH="${FILE_WHEEL:-/workspace/paddlefleet-0.0.0-py3-none-linux_x86_64.whl}" - PADDLEFLEET_OPS_WHEEL_PATH="${FILE_OPS:-/workspace/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl}" - PADDLEFLEET_WHEEL_ORIGIN="${FILE_WHEEL_ORIGIN:-develop_latest}" - PADDLEFLEET_OPS_ORIGIN="${FILE_OPS_ORIGIN:-develop_latest}" - PADDLEFLEET_PIN_RECEIPT="${FILE_RECEIPT:-}" - PADDLEFLEET_SOURCE_COMMIT="${FILE_SOURCE:-}" - PADDLEFLEET_PIN_SHA="${requested_pin}" - fi - - if [[ ! -e "${PADDLEFLEET_WHEEL_PATH}" ]]; then - fail "missing paddlefleet path: ${PADDLEFLEET_WHEEL_PATH}" - fi - if [[ ! -e "${PADDLEFLEET_OPS_WHEEL_PATH}" ]]; then - fail "missing paddlefleet_ops path: ${PADDLEFLEET_OPS_WHEEL_PATH}" - fi - write_consumed -} - -write_min_receipt() { - local path="$1" mode="$2" pin="$3" wheel="$4" ops="$5" origin="$6" digest="$7" pairing="$8" - python3 - "${path}" "${mode}" "${pin}" "${wheel}" "${ops}" "${origin}" "${digest}" "${pairing}" <<'PY' -import json, sys -path, mode, pin, wheel, ops, origin, digest, pairing = sys.argv[1:9] -digest_ok = digest == "true" -doc = { - "schema": "paddlefleet-alignment-pin/v1", - "status": "ok", - "mode": mode, - "source": { - "expected_commit": pin or None, - "actual_commit": pin or None, - "commit_verified": bool(pin) and mode == "stack-paired", - }, - "artifacts": [ - {"name": "paddlefleet", "path": wheel, "url": None, "origin": origin, - "digest_verified": digest_ok, "expected_sha256": "abc" if digest_ok else None, - "actual_sha256": "abc" if digest_ok else None}, - {"name": "paddlefleet_ops", "path": ops, "url": None, "origin": origin, - "digest_verified": digest_ok, "expected_sha256": "def" if digest_ok else None, - "actual_sha256": "def" if digest_ok else None}, - ], - "pairing": {"stack_paired_proven": False, "status": pairing, "reason": "fixture"}, -} -open(path, "w", encoding="utf-8").write(json.dumps(doc, indent=2) + "\n") -PY -} - -run_self_test() { - local root self - self="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)/$(basename -- "${BASH_SOURCE[0]}")" - root="$(mktemp -d)" - trap 'rm -rf "${root}"' RETURN - mkdir -p "${root}/PaddleFleet/packages/paddlefleet_ops" - echo tree >"${root}/PaddleFleet/pyproject.toml" - echo ops >"${root}/PaddleFleet/packages/paddlefleet_ops/pyproject.toml" - local pin="aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa" - - write_min_receipt "${root}/receipt.json" stack-paired "${pin}" \ - "${root}/PaddleFleet" "${root}/PaddleFleet/packages/paddlefleet_ops" \ - source_tree false source_tree_from_checked_out_pin - cat >"${root}/pin.env" <"${root}/devwhl/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" - echo dummy >"${root}/devwhl/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" - write_min_receipt "${root}/develop-receipt.json" develop "" \ - "${root}/devwhl/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" \ - "${root}/devwhl/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" \ - develop_latest false unpaired_default - cat >"${root}/develop.env" <&2 - exit 1 - fi - - # 2) verified URL+hash artifact may keep the 0.0.0 filename. - write_min_receipt "${root}/named-receipt.json" stack-paired "${pin}" \ - "${root}/devwhl/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" \ - "${root}/devwhl/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" \ - ci_metadata true unproven - cat >"${root}/named.env" <"${root}/bare.env" <&2 - exit 1 - fi - - # Selector clone/fail: error receipt, no env. Missing env is the - # consequence, not proof a generated env failed to cross docker exec. - mkdir -p "${root}/sel-fail" - cat >"${root}/sel-fail/paddlefleet_alignment_pin_receipt.json" <<'EOF' -{"schema":"paddlefleet-alignment-pin/v1","status":"error","detail":"git clone failed: github.com:443","mode":"stack-paired"} -EOF - if ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${pin}" \ - bash "${self}" --env "${root}/sel-fail/paddlefleet_alignment_pin.env" \ - --out "${root}/sel-fail.consumed.env" 2>"${root}/sel-fail.err"; then - echo "self-test FAIL: selector error receipt was consumed" >&2 - exit 1 - fi - grep -q "because selector failed" "${root}/sel-fail.err" - if grep -q "selector wrote ok receipt but env did not cross docker exec" "${root}/sel-fail.err"; then - echo "self-test FAIL: selector error misclassified as env-handoff" >&2 - exit 1 - fi - - echo "consume_paddlefleet_alignment_pin self-test OK" -} - -if [[ "${RUN_SELF_TEST}" == 1 ]]; then - run_self_test - exit 0 -fi -consume diff --git a/scripts/prepare_paddlefleet_ops_submodules.sh b/scripts/prepare_paddlefleet_ops_submodules.sh deleted file mode 100644 index 7da5c9c32bd..00000000000 --- a/scripts/prepare_paddlefleet_ops_submodules.sh +++ /dev/null @@ -1,29 +0,0 @@ -#!/usr/bin/env bash -# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. -# Retry source-tree network preparation before setup/build, retaining the -# pinned Fleet helper's exact and nested gitlink verification. -set -euo pipefail -ops_path="${1:?ops-path}" -setup_script="${2:?pinned-setup-script}" -if [[ ! -d "$ops_path" ]]; then - echo "[ops-submodules] wheel input; no source preparation" - exit 0 -fi -command -v timeout >/dev/null -[[ -f "$setup_script" ]] || { echo "[ops-submodules] pinned setup script missing" >&2; exit 1; } -for attempt in 1 2 3; do - echo "[ops-submodules] prepare and verify recorded gitlinks, attempt $attempt/3" - if timeout --kill-after=30s 15m bash "$setup_script" --prepare-ops-submodules "$ops_path"; then - echo "[ops-submodules] recorded gitlinks prepared and verified" - exit 0 - else - rc=$? - fi - if [[ "$attempt" == 3 ]]; then - echo "[ops-submodules] preparation failed after 3 attempts (exit $rc); stop before setup/build/training" >&2 - exit "$rc" - fi - delay=$((attempt * 15)) - echo "[ops-submodules] preparation exited $rc; retry in ${delay}s" - sleep "$delay" -done diff --git a/scripts/require_paddlefleet_selector_ok.sh b/scripts/require_paddlefleet_selector_ok.sh deleted file mode 100755 index f34221204ae..00000000000 --- a/scripts/require_paddlefleet_selector_ok.sh +++ /dev/null @@ -1,40 +0,0 @@ -#!/usr/bin/env bash -# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. -# -# Called from Get Whl after select_paddlefleet_alignment_pin.sh. -# Keep this file free of nested quotes so the workflow docker exec -# single-quoted -c script can invoke it without breaking the host shell. -# A selector error receipt is not an env-handoff failure. - -set -euo pipefail - -DEST="${1:-/workspace}" -REC="${DEST}/paddlefleet_alignment_pin_receipt.json" -ENVF="${DEST}/paddlefleet_alignment_pin.env" - -if [[ ! -f "${REC}" ]]; then - echo "::error:: missing selector receipt ${REC}" >&2 - exit 1 -fi - -python3 - "${REC}" <<'PY' -import json -import sys - -path = sys.argv[1] -doc = json.load(open(path, encoding="utf-8")) -status = doc.get("status") -if status != "ok": - raise SystemExit( - f"selector receipt status={status!r} is not ok; " - "Get Whl must stop (do not download remaining wheels or build)" - ) -print("selector receipt status=ok") -PY - -if [[ ! -f "${ENVF}" ]]; then - echo "::error:: selector did not write ${ENVF}" >&2 - exit 1 -fi - -cat "${REC}" diff --git a/scripts/select_paddlefleet_alignment_pin.sh b/scripts/select_paddlefleet_alignment_pin.sh deleted file mode 100755 index cfa778fa9d3..00000000000 --- a/scripts/select_paddlefleet_alignment_pin.sh +++ /dev/null @@ -1,889 +0,0 @@ -#!/usr/bin/env bash -# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -# Select PaddleFleet source + wheels for alignment_model_accuracy. -# -# Default (ALIGNMENT_PADDLEFLEET_MODE=develop or unset): -# historical CodeSync/develop tarball + BOS develop/latest wheels. -# Cases are not filtered. -# -# Explicit (ALIGNMENT_PADDLEFLEET_MODE=stack-paired): fail-closed. -# Fetch PADDLEFLEET_PIN_SHA at depth 1 (not a full default-branch clone). -# Transient git RPC/curl failures retry up to 3 clean dests; other -# fetch errors fail closed. git rev-parse HEAD must equal the pin. -# Artifacts: caller URL+sha256 (from Build Fleet whl / Actions metadata) -# or build from the checked-out tree. Independent digest matches do not -# prove the wheels were produced from that commit — receipt records -# source_commit vs artifact sha256 separately and pairing as unproven -# unless this invocation built the files from the pin. -# git/download/checkout failures still write an error receipt. - -set -euo pipefail - -usage() { - cat <<'EOF' -Usage: select_paddlefleet_alignment_pin.sh [--dest DIR] [--self-test] - ---self-test ignores --dest and uses an offline fixture (no network). - -Env: - ALIGNMENT_PADDLEFLEET_MODE develop (default) | stack-paired - PADDLEFLEET_PIN_SHA required 40-hex commit in stack-paired - PADDLEFLEET_GIT_URL git remote or local repo (stack-paired) - PADDLEFLEET_WHEEL_URL optional explicit wheel (https or local path) - PADDLEFLEET_WHEEL_SHA256 required with WHEEL_URL - PADDLEFLEET_OPS_WHEEL_URL optional explicit ops wheel - PADDLEFLEET_OPS_WHEEL_SHA256 required with OPS URL - PADDLEFLEET_BUILD_CMD optional; default uv build paddlefleet - PADDLEFLEET_BUILD_OPS_CMD optional; default uv build paddlefleet-ops - ALIGNMENT_PADDLEFLEET_DEST output directory (default /workspace) -EOF -} - -MODE="${ALIGNMENT_PADDLEFLEET_MODE:-develop}" -DEST="${ALIGNMENT_PADDLEFLEET_DEST:-/workspace}" -RUN_SELF_TEST=0 -while [[ $# -gt 0 ]]; do - case "$1" in - --dest) - DEST="${2:?--dest requires a path}" - shift 2 - ;; - --self-test) - RUN_SELF_TEST=1 - shift - ;; - -h|--help) - usage - exit 0 - ;; - *) - echo "unknown arg: $1" >&2 - usage - exit 2 - ;; - esac -done - -BOS="${PADDLEFLEET_BOS:-https://paddle-github-action.bj.bcebos.com}" -DEFAULT_TAR_URL="https://paddle-qa.bj.bcebos.com/CodeSync/develop/PaddleFleet.tar" -DEFAULT_WHL_URL="${BOS}/PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" -DEFAULT_OPS_URL="${BOS}/PaddleFleet/develop/latest/cu130/paddle-release/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" -GIT_URL="${PADDLEFLEET_GIT_URL:-https://github.com/PaddlePaddle/PaddleFleet.git}" -PIN_SHA="${PADDLEFLEET_PIN_SHA:-}" -WHEEL_URL="${PADDLEFLEET_WHEEL_URL:-}" -WHEEL_SHA="${PADDLEFLEET_WHEEL_SHA256:-}" -OPS_URL="${PADDLEFLEET_OPS_WHEEL_URL:-}" -OPS_SHA="${PADDLEFLEET_OPS_WHEEL_SHA256:-}" -BUILD_CMD="${PADDLEFLEET_BUILD_CMD:-}" -BUILD_OPS_CMD="${PADDLEFLEET_BUILD_OPS_CMD:-}" - -ACTUAL_SHA="" -SOURCE_VERIFIED=false -PADDLEFLEET_WHEEL_PATH="" -PADDLEFLEET_OPS_WHEEL_PATH="" -ACTUAL_WHEEL_SHA="" -ACTUAL_OPS_SHA="" -WHEEL_DIGEST_VERIFIED=false -OPS_DIGEST_VERIFIED=false -WHEEL_ORIGIN="" -OPS_ORIGIN="" -WHEEL_BUILT_FROM_COMMIT="" -OPS_BUILT_FROM_COMMIT="" -LOADED_FROM="" -RECEIPT_WRITTEN=0 - -log() { echo "[paddlefleet-pin] $*" >&2; } - -sha256_file() { sha256sum -- "$1" | awk '{print $1}'; } - -unpaired_url() { - case "$1" in - *"/develop/latest/"*|*"CodeSync/develop/"*) return 0 ;; - *) return 1 ;; - esac -} - -pairing_fields() { - local proven=false - local status="unproven" - local reason="wheel/ops digest match does not prove production from source_commit" - if [[ "${MODE}" == "develop" ]]; then - status="unpaired_default" - reason="develop tarball and develop/latest wheels; not a stack pin" - elif [[ "${WHEEL_ORIGIN}" == "source_tree" && "${OPS_ORIGIN}" == "source_tree" \ - && "${SOURCE_VERIFIED}" == "true" ]]; then - status="source_tree_from_checked_out_pin" - reason="this invocation exported checked-out source trees; not a wheel digest proof" - elif [[ "${WHEEL_ORIGIN}" == "build" && "${OPS_ORIGIN}" == "build" \ - && "${WHEEL_BUILT_FROM_COMMIT}" == "${ACTUAL_SHA}" \ - && "${OPS_BUILT_FROM_COMMIT}" == "${ACTUAL_SHA}" \ - && "${SOURCE_VERIFIED}" == "true" ]]; then - status="built_from_checked_out_pin" - reason="this invocation built both artifacts from checked-out source_commit; not a remote-stack proof" - fi - printf '%s\t%s\t%s\n' "${proven}" "${status}" "${reason}" -} - -write_receipt() { - local status="$1" detail="${2:-}" - mkdir -p "${DEST}" - local receipt="${DEST}/paddlefleet_alignment_pin_receipt.json" - local pair - pair="$(pairing_fields)" - local stack_proven pairing_status pairing_reason - stack_proven="${pair%%$'\t'*}" - pair="${pair#*$'\t'}" - pairing_status="${pair%%$'\t'*}" - pairing_reason="${pair#*$'\t'}" - if ! command -v python3 >/dev/null 2>&1; then - printf '{"schema":"paddlefleet-alignment-pin/v1","status":"%s","detail":"%s"}\n' \ - "${status}" "${detail}" >"${receipt}" - RECEIPT_WRITTEN=1 - return 0 - fi - python3 - "${receipt}" "${status}" "${detail}" "${stack_proven}" \ - "${pairing_status}" "${pairing_reason}" <<'PY' -import json, os, sys -from datetime import datetime, timezone -path, status, detail, stack_proven, pairing_status, pairing_reason = sys.argv[1:7] - -def art(name, pth, url, exp, act, digest_ok, origin, built_from): - if not pth: - return None - return { - "name": name, - "path": pth, - "url": url or None, - "expected_sha256": exp or None, - "actual_sha256": act or None, - "digest_verified": digest_ok == "true", - "origin": origin or None, - "built_from_commit": built_from or None, - } - -arts = [a for a in ( - art("paddlefleet", os.environ.get("PADDLEFLEET_WHEEL_PATH", ""), - os.environ.get("WHEEL_URL", ""), os.environ.get("WHEEL_SHA", ""), - os.environ.get("ACTUAL_WHEEL_SHA", ""), os.environ.get("WHEEL_DIGEST_VERIFIED", "false"), - os.environ.get("WHEEL_ORIGIN", ""), os.environ.get("WHEEL_BUILT_FROM_COMMIT", "")), - art("paddlefleet_ops", os.environ.get("PADDLEFLEET_OPS_WHEEL_PATH", ""), - os.environ.get("OPS_URL", ""), os.environ.get("OPS_SHA", ""), - os.environ.get("ACTUAL_OPS_SHA", ""), os.environ.get("OPS_DIGEST_VERIFIED", "false"), - os.environ.get("OPS_ORIGIN", ""), os.environ.get("OPS_BUILT_FROM_COMMIT", "")), -) if a] - -doc = { - "schema": "paddlefleet-alignment-pin/v1", - "status": status, - "detail": detail, - "captured_at": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"), - "mode": os.environ.get("MODE"), - "dest": os.environ.get("DEST"), - "source": { - "git_url": os.environ.get("GIT_URL") or None, - "expected_commit": os.environ.get("PIN_SHA") or None, - "actual_commit": os.environ.get("ACTUAL_SHA") or None, - "commit_verified": os.environ.get("SOURCE_VERIFIED") == "true", - }, - "loaded_from": os.environ.get("LOADED_FROM") or None, - "artifacts": arts, - "pairing": { - "stack_paired_proven": stack_proven == "true", - "status": pairing_status, - "reason": pairing_reason, - }, - "default_urls": { - "source_tar": os.environ.get("DEFAULT_TAR_URL"), - "paddlefleet_wheel": os.environ.get("DEFAULT_WHL_URL"), - "paddlefleet_ops_wheel": os.environ.get("DEFAULT_OPS_URL"), - }, - "cases_preserved": ["MinimaxV2.5_EP2", "GLM45Air_EP2"], - "unpaired_develop_rejected_in_stack_paired": True, -} -open(path, "w", encoding="utf-8").write(json.dumps(doc, indent=2) + "\n") -print("[paddlefleet-pin] receipt", path, file=sys.stderr) -PY - RECEIPT_WRITTEN=1 -} - -export_receipt_env() { - export MODE DEST GIT_URL PIN_SHA ACTUAL_SHA SOURCE_VERIFIED LOADED_FROM - export PADDLEFLEET_WHEEL_PATH PADDLEFLEET_OPS_WHEEL_PATH - export WHEEL_URL WHEEL_SHA ACTUAL_WHEEL_SHA WHEEL_DIGEST_VERIFIED WHEEL_ORIGIN WHEEL_BUILT_FROM_COMMIT - export OPS_URL OPS_SHA ACTUAL_OPS_SHA OPS_DIGEST_VERIFIED OPS_ORIGIN OPS_BUILT_FROM_COMMIT - export DEFAULT_TAR_URL DEFAULT_WHL_URL DEFAULT_OPS_URL -} - -fail() { - trap - ERR - local msg="$1" - log "FAIL: ${msg}" - export_receipt_env - write_receipt "error" "${msg}" - echo "::error:: ${msg}" >&2 - exit 1 -} - -on_err() { - local rc=$? - if [[ "${RECEIPT_WRITTEN}" == 1 || "${RUN_SELF_TEST}" == 1 ]]; then - return "${rc}" - fi - fail "command failed rc=${rc}" -} -trap 'on_err' ERR - -write_envfile() { - cat >"${DEST}/paddlefleet_alignment_pin.env" < ${out}" - mkdir -p "$(dirname "${out}")" - if [[ "${url}" == file://* ]]; then - local src="${url#file://}" - [[ -f "${src}" ]] || fail "download failed, local file missing: ${src}" - cp -- "${src}" "${out}" || fail "download copy failed: ${src}" - return 0 - fi - if [[ "${url}" == /* ]]; then - [[ -f "${url}" ]] || fail "download failed, local file missing: ${url}" - cp -- "${url}" "${out}" || fail "download copy failed: ${url}" - return 0 - fi - if wget -q --no-proxy --no-check-certificate --tries=2 --timeout=15 -O "${out}" "${url}"; then - return 0 - fi - fail "download failed: ${url}" -} - -# Must not run inside $(); fail() has to exit this shell. -require_digest() { - local path="$1" expected="$2" label="$3" actual="$4" - [[ -f "${path}" ]] || fail "missing ${label}: ${path}" - [[ -n "${expected}" ]] || fail "stack-paired missing ${label} sha256" - if [[ "${actual}" != "${expected}" ]]; then - fail "stack-paired ${label} sha256 mismatch expected=${expected} actual=${actual}" - fi -} - -# Auth / missing-object / missing-repo are permanent. Do not treat a -# generic "RPC failed" as transient: HTTP 401/403 also say RPC failed. -_is_permanent_git_err() { - case "$1" in - *"HTTP 401"*|*"HTTP 403"*|*"Authentication failed"*|*"access denied"*|*"Access denied"*|*"Permission denied"*|*"not our ref"*|*"does not appear to be a git repository"*|*"Could not read from remote repository"*|*"remote: Write access"*) - return 0 - ;; - esac - return 1 -} - -_is_transient_git_err() { - _is_permanent_git_err "$1" && return 1 - case "$1" in - *"curl 56"*|*"Connection timed out"*|*"Couldn't connect to server"*|*"Failed to connect to github.com port 443"*|*"early EOF"*|*"fetch-pack: unexpected disconnect"*|*"bytes of body are still expected"*|*"invalid index-pack output"*|*"Connection reset by peer"*) - return 0 - ;; - esac - return 1 -} - -checkout_pin() { - [[ "${PIN_SHA}" =~ ^[0-9a-fA-F]{40}$ ]] || fail "stack-paired requires PADDLEFLEET_PIN_SHA (40 hex), got '${PIN_SHA}'" - PIN_SHA="$(printf '%s' "${PIN_SHA}" | tr 'A-F' 'a-f')" - local dest="${DEST}/PaddleFleet" - local attempt max_attempts=3 - local last_err="" - local fetch_err="${DEST}/.git-fetch.err" - for attempt in $(seq 1 "${max_attempts}"); do - rm -rf "${dest}" - log "fetch ${GIT_URL} ${PIN_SHA} (attempt ${attempt}/${max_attempts})" - if ! git init --quiet "${dest}" >/dev/null 2>"${DEST}/.git-init.err"; then - fail "git init failed: $(tr '\n' ' ' <"${DEST}/.git-init.err")" - fi - if ! git -C "${dest}" remote add origin "${GIT_URL}" >/dev/null 2>"${DEST}/.git-remote.err"; then - fail "git remote add failed: $(tr '\n' ' ' <"${DEST}/.git-remote.err")" - fi - git -C "${dest}" config advice.detachedHead false || true - if GIT_TERMINAL_PROMPT=0 git -C "${dest}" fetch --depth=1 --no-tags origin "${PIN_SHA}" \ - >/dev/null 2>"${fetch_err}"; then - if ! git -C "${dest}" checkout --quiet --force --detach "${PIN_SHA}" >/dev/null 2>"${DEST}/.git-co.err"; then - fail "git checkout failed for ${PIN_SHA}: $(tr '\n' ' ' <"${DEST}/.git-co.err")" - fi - ACTUAL_SHA="$(git -C "${dest}" rev-parse HEAD)" - if [[ "${ACTUAL_SHA}" != "${PIN_SHA}" ]]; then - SOURCE_VERIFIED=false - fail "stack-paired source SHA mismatch expected=${PIN_SHA} actual=${ACTUAL_SHA}" - fi - SOURCE_VERIFIED=true - log "source commit verified ${ACTUAL_SHA}" - return 0 - fi - last_err="$(tr '\n' ' ' <"${fetch_err}")" - if ! _is_transient_git_err "${last_err}"; then - fail "git fetch failed: ${last_err}" - fi - log "transient fetch failure attempt ${attempt}/${max_attempts}: ${last_err}" - done - fail "git fetch failed: ${last_err}" -} - -# Sets DEST_PATH and DEST_SHA in the caller. Must run in this shell so -# fail() writes the receipt (never wrap this in $()). -acquire_explicit() { - local url="$1" expected="$2" dest_name="$3" label="$4" - unpaired_url "${url}" && fail "stack-paired rejects unpaired ${label} URL: ${url}" - [[ -n "${expected}" ]] || fail "stack-paired ${label} URL requires matching sha256" - download "${url}" "${DEST}/${dest_name}" - DEST_PATH="${DEST}/${dest_name}" - DEST_SHA="$(sha256_file "${DEST_PATH}")" - # Record path/digest before require_digest so a mismatch receipt still has them. - if [[ "${label}" == paddlefleet\ wheel ]]; then - PADDLEFLEET_WHEEL_PATH="${DEST_PATH}" - ACTUAL_WHEEL_SHA="${DEST_SHA}" - WHEEL_ORIGIN="ci_metadata" - else - PADDLEFLEET_OPS_WHEEL_PATH="${DEST_PATH}" - ACTUAL_OPS_SHA="${DEST_SHA}" - OPS_ORIGIN="ci_metadata" - fi - require_digest "${DEST_PATH}" "${expected}" "${label}" "${DEST_SHA}" - LOADED_FROM="${LOADED_FROM:+${LOADED_FROM};}${DEST_PATH} from ${url}" -} - -run_build() { - local cmd="$1" glob="$2" label="$3" - mkdir -p "${DEST}/dist" - log "build ${label}: ${cmd}" - if ! (cd "${DEST}/PaddleFleet" && bash -lc "${cmd}"); then - fail "stack-paired build failed for ${label}" - fi - local built - built="$(ls -1 ${glob} 2>/dev/null | head -n 1 || true)" - [[ -n "${built}" && -f "${built}" ]] || fail "stack-paired build produced no ${label} (glob ${glob})" - local dest_name - dest_name="$(basename "${built}")" - cp -f -- "${built}" "${DEST}/${dest_name}" - DEST_PATH="${DEST}/${dest_name}" - DEST_SHA="$(sha256_file "${DEST_PATH}")" - LOADED_FROM="${LOADED_FROM:+${LOADED_FROM};}built ${DEST_PATH} from ${ACTUAL_SHA}" -} - -acquire_wheel() { - if [[ -n "${WHEEL_URL}" ]]; then - acquire_explicit "${WHEEL_URL}" "${WHEEL_SHA}" "paddlefleet.whl" "paddlefleet wheel" - PADDLEFLEET_WHEEL_PATH="${DEST_PATH}" - ACTUAL_WHEEL_SHA="${DEST_SHA}" - WHEEL_DIGEST_VERIFIED=true - WHEEL_ORIGIN="ci_metadata" - WHEEL_BUILT_FROM_COMMIT="" - else - local cmd="${BUILD_CMD:-uv build --wheel --package paddlefleet --out-dir '${DEST}/dist' --clear}" - run_build "${cmd}" "${DEST}/dist/paddlefleet-*.whl" "paddlefleet wheel" - PADDLEFLEET_WHEEL_PATH="${DEST_PATH}" - ACTUAL_WHEEL_SHA="${DEST_SHA}" - WHEEL_DIGEST_VERIFIED=true - WHEEL_ORIGIN="build" - WHEEL_BUILT_FROM_COMMIT="${ACTUAL_SHA}" - fi -} - -acquire_ops() { - if [[ -n "${OPS_URL}" ]]; then - acquire_explicit "${OPS_URL}" "${OPS_SHA}" "paddlefleet_ops.whl" "paddlefleet_ops wheel" - PADDLEFLEET_OPS_WHEEL_PATH="${DEST_PATH}" - ACTUAL_OPS_SHA="${DEST_SHA}" - OPS_DIGEST_VERIFIED=true - OPS_ORIGIN="ci_metadata" - OPS_BUILT_FROM_COMMIT="" - else - local cmd="${BUILD_OPS_CMD:-uv build --wheel --package paddlefleet-ops --out-dir '${DEST}/dist' --no-build-isolation}" - run_build "${cmd}" "${DEST}/dist/paddlefleet_ops-*.whl" "paddlefleet_ops wheel" - PADDLEFLEET_OPS_WHEEL_PATH="${DEST_PATH}" - ACTUAL_OPS_SHA="${DEST_SHA}" - OPS_DIGEST_VERIFIED=true - OPS_ORIGIN="build" - OPS_BUILT_FROM_COMMIT="${ACTUAL_SHA}" - fi -} - -fetch_default() { - log "mode=develop (historical unpaired CodeSync tarball)" - download "${DEFAULT_TAR_URL}" "${DEST}/PaddleFleet.tar" - rm -rf "${DEST}/PaddleFleet" - tar xf "${DEST}/PaddleFleet.tar" -C "${DEST}" - rm -f "${DEST}/PaddleFleet.tar" - if [[ -d "${DEST}/PaddleFleet/.git" ]]; then - git -C "${DEST}/PaddleFleet" pull || log "git pull skipped" - ACTUAL_SHA="$(git -C "${DEST}/PaddleFleet" rev-parse HEAD 2>/dev/null || true)" - fi - download "${DEFAULT_WHL_URL}" "${DEST}/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" - download "${DEFAULT_OPS_URL}" "${DEST}/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" - PADDLEFLEET_WHEEL_PATH="${DEST}/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" - PADDLEFLEET_OPS_WHEEL_PATH="${DEST}/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" - WHEEL_URL="${DEFAULT_WHL_URL}" - OPS_URL="${DEFAULT_OPS_URL}" - ACTUAL_WHEEL_SHA="$(sha256_file "${PADDLEFLEET_WHEEL_PATH}")" - ACTUAL_OPS_SHA="$(sha256_file "${PADDLEFLEET_OPS_WHEEL_PATH}")" - WHEEL_ORIGIN="develop_latest" - OPS_ORIGIN="develop_latest" - SOURCE_VERIFIED=false - WHEEL_DIGEST_VERIFIED=false - OPS_DIGEST_VERIFIED=false - LOADED_FROM="develop_latest ${DEFAULT_WHL_URL} ${DEFAULT_OPS_URL}" - export_receipt_env - write_envfile - write_receipt "ok" "develop tarball and develop/latest wheels; unpaired with a stack pin" -} - -acquire_source_tree() { - local src="${DEST}/PaddleFleet" - local ops="${src}/packages/paddlefleet_ops" - [[ -d "${src}" ]] || fail "stack-paired source tree missing: ${src}" - [[ -f "${src}/pyproject.toml" ]] || fail "stack-paired source tree missing pyproject.toml: ${src}" - [[ -d "${ops}" ]] || fail "stack-paired ops source tree missing: ${ops}" - PADDLEFLEET_WHEEL_PATH="${src}" - PADDLEFLEET_OPS_WHEEL_PATH="${ops}" - WHEEL_ORIGIN="source_tree" - OPS_ORIGIN="source_tree" - WHEEL_BUILT_FROM_COMMIT="${ACTUAL_SHA}" - OPS_BUILT_FROM_COMMIT="${ACTUAL_SHA}" - WHEEL_DIGEST_VERIFIED=false - OPS_DIGEST_VERIFIED=false - LOADED_FROM="source_tree ${src} ${ops} from ${ACTUAL_SHA}" - log "source-tree paths ${src} ${ops}" -} - -fetch_stack_paired() { - log "mode=stack-paired" - checkout_pin - if [[ -n "${WHEEL_URL}" || -n "${OPS_URL}" || -n "${BUILD_CMD}" || -n "${BUILD_OPS_CMD}" ]]; then - acquire_wheel - acquire_ops - else - # No CI wheel URL and no explicit build: export the checked-out trees. - # A later docker exec must source paddlefleet_alignment_pin.env. - acquire_source_tree - fi - export_receipt_env - write_envfile - local pair pairing_status - pair="$(pairing_fields)" - pair="${pair#*$'\t'}" - pairing_status="${pair%%$'\t'*}" - write_receipt "ok" "source_commit checked out; pairing.status=${pairing_status}; stack_paired_proven=false" -} - -install_offline_stubs() { - local bin="$1" - mkdir -p "${bin}" - cat >"${bin}/wget" <<'WGET' -#!/usr/bin/env bash -out="" -url="" -while [[ $# -gt 0 ]]; do - case "$1" in - -O) out="$2"; shift 2 ;; - --*) shift ;; - *) url="$1"; shift ;; - esac -done -if [[ -z "${url}" || -z "${out}" ]]; then - echo "wget-stub: missing url/out" >&2 - exit 1 -fi -if [[ "${url}" == http://* || "${url}" == https://* ]]; then - echo "wget-stub: blocked network ${url}" >&2 - exit 1 -fi -src="${url#file://}" -if [[ -f "${src}" ]]; then - cp -- "${src}" "${out}" - exit 0 -fi -echo "wget-stub: not a local file ${url}" >&2 -exit 1 -WGET - chmod +x "${bin}/wget" -} - -run_self_test() { - trap - ERR - local root script - root="$(mktemp -d)" - script="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)/$(basename -- "${BASH_SOURCE[0]}")" - trap 'rm -rf "${root}"' RETURN - install_offline_stubs "${root}/bin" - export PATH="${root}/bin:${PATH}" - - git init -q "${root}/upstream" - git -C "${root}/upstream" config user.email test@example.com - git -C "${root}/upstream" config user.name test - echo source-a >"${root}/upstream/README" - mkdir -p "${root}/upstream/packages/paddlefleet_ops" - printf '%s\n' '[project]' 'name = "paddlefleet"' >"${root}/upstream/pyproject.toml" - printf '%s\n' '[project]' 'name = "paddlefleet-ops"' >"${root}/upstream/packages/paddlefleet_ops/pyproject.toml" - git -C "${root}/upstream" add README pyproject.toml packages - git -C "${root}/upstream" commit -q -m a - local sha_a sha_b - sha_a="$(git -C "${root}/upstream" rev-parse HEAD)" - echo source-b >"${root}/upstream/README" - git -C "${root}/upstream" add README - git -C "${root}/upstream" commit -q -m b - sha_b="$(git -C "${root}/upstream" rev-parse HEAD)" - - mkdir -p "${root}/art" - echo py-body >"${root}/art/py.whl" - echo ops-body >"${root}/art/ops.whl" - local py_sha ops_sha - py_sha="$(sha256_file "${root}/art/py.whl")" - ops_sha="$(sha256_file "${root}/art/ops.whl")" - - expect_fail() { - local dest="$1" - local needle="$2" - shift 2 - mkdir -p "${dest}" - if "$@"; then - echo "self-test FAIL: expected failure (${needle})" >&2 - exit 1 - fi - local rec="${dest}/paddlefleet_alignment_pin_receipt.json" - [[ -f "${rec}" ]] || { echo "self-test FAIL: missing error receipt ${rec}" >&2; exit 1; } - grep -q '"status": "error"' "${rec}" - grep -q "${needle}" "${rec}" - echo "[self-test] fail-closed ${dest}: ${needle}" - } - - local run - run() { env PATH="${root}/bin:${PATH}" "$@"; } - - expect_fail "${root}/m1" "PADDLEFLEET_PIN_SHA" \ - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA= \ - bash "${script}" --dest "${root}/m1" - - expect_fail "${root}/m2" "rejects unpaired" \ - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - PADDLEFLEET_WHEEL_URL="${DEFAULT_WHL_URL}" \ - PADDLEFLEET_WHEEL_SHA256="${py_sha}" \ - PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ - PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ - bash "${script}" --dest "${root}/m2" - - expect_fail "${root}/m3" "git fetch failed" \ - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="0000000000000000000000000000000000000000" \ - PADDLEFLEET_GIT_URL="${root}/upstream" \ - bash "${script}" --dest "${root}/m3" - - # Real checksum mismatch after a successful local copy (not a wget miss). - expect_fail "${root}/m4" "sha256 mismatch" \ - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - PADDLEFLEET_WHEEL_URL="${root}/art/py.whl" \ - PADDLEFLEET_WHEEL_SHA256="deadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeefdeadbeef" \ - PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ - PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ - bash "${script}" --dest "${root}/m4" - python3 - "${root}/m4/paddlefleet_alignment_pin_receipt.json" "${py_sha}" <<'PY' -import json, sys -doc = json.load(open(sys.argv[1])) -assert doc["status"] == "error" -wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") -assert wheel["actual_sha256"] == sys.argv[2] -assert wheel["digest_verified"] is False -assert wheel["actual_sha256"] != (wheel.get("expected_sha256") or "") -print("m4 checksum-mismatch receipt has actual digest, not a download miss") -PY - - expect_fail "${root}/m5" "produced no paddlefleet" \ - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - PADDLEFLEET_BUILD_CMD="mkdir -p '${root}/m5/dist'" \ - PADDLEFLEET_BUILD_OPS_CMD="true" \ - bash "${script}" --dest "${root}/m5" - - expect_fail "${root}/m6" "download failed" \ - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - PADDLEFLEET_WHEEL_URL="https://example.invalid/paddlefleet.whl" \ - PADDLEFLEET_WHEEL_SHA256="${py_sha}" \ - PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ - PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ - bash "${script}" --dest "${root}/m6" - - expect_fail "${root}/m7" "git fetch failed" \ - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" \ - PADDLEFLEET_GIT_URL="${root}/no-such-remote" \ - bash "${script}" --dest "${root}/m7" - - assert_error_unverified() { - python3 - "$1" <<'PY' -import json, sys -doc = json.load(open(sys.argv[1])) -assert doc["status"] == "error", doc -assert doc["source"]["commit_verified"] is False, doc["source"] -print("error receipt commit_verified=false") -PY - } - - assert_pin_checkout() { - local repo="$1" sha="$2" - [[ "$(git -C "${repo}" rev-parse HEAD)" == "${sha}" ]] - [[ -f "${repo}/.git/shallow" ]] || { echo "self-test FAIL: missing ${repo}/.git/shallow" >&2; exit 1; } - [[ "$(git -C "${repo}" rev-list --count HEAD)" == 1 ]] || { - echo "self-test FAIL: ${repo} rev-list count != 1 (not depth=1)" >&2 - exit 1 - } - if git -C "${repo}" symbolic-ref -q HEAD >/dev/null; then - echo "self-test FAIL: ${repo} HEAD is a branch, not detached pin" >&2 - git -C "${repo}" symbolic-ref HEAD >&2 - exit 1 - fi - } - - write_git_wrapper() { - local bindir="$1" mode="$2" - local real_git - real_git="$(command -v git)" - mkdir -p "${bindir}" - cat >"${bindir}/git" <>"\$STALE" - exit 99 - fi - if [[ "\$has_depth" -ne 1 || "\$has_no_tags" -ne 1 ]]; then - echo "fetch missing --depth=1/--no-tags: \${args[*]}" >>"\$STALE" - exit 99 - fi - count=\$((count + 1)) - printf '%s\n' "\$count" >"\$STATE" - case "${mode}" in - 401) - echo "RPC failed; HTTP 401 curl 22 The requested URL returned error: 401" >&2 - echo "fatal: Authentication failed" >&2 - printf 'sentinel\n' >"\$workdir/.retry-sentinel" - exit 128 - ;; - exhaust) - echo "error: RPC failed; curl 56 Recv failure: Connection timed out" >&2 - echo "error: 9515 bytes of body are still expected" >&2 - echo "fatal: early EOF" >&2 - printf 'sentinel\n' >"\$workdir/.retry-sentinel" - exit 128 - ;; - retry) - if [[ "\$count" -lt 3 ]]; then - echo "error: RPC failed; curl 56 Recv failure: Connection timed out" >&2 - echo "error: 9515 bytes of body are still expected" >&2 - echo "fetch-pack: unexpected disconnect while reading sideband packet" >&2 - echo "fatal: early EOF" >&2 - echo "fatal: fetch-pack: invalid index-pack output" >&2 - printf 'sentinel\n' >"\$workdir/.retry-sentinel" - exit 128 - fi - ;; - esac -fi -exec "\$real" "\$@" -GITWRAP - chmod +x "${bindir}/git" - } - - assert_error_unverified "${root}/m3/paddlefleet_alignment_pin_receipt.json" - assert_error_unverified "${root}/m7/paddlefleet_alignment_pin_receipt.json" - - # Permanent HTTP 401 (also says RPC failed) must not retry. - write_git_wrapper "${root}/bin-401" 401 - expect_fail "${root}/m8" "git fetch failed" \ - env PATH="${root}/bin-401:${root}/bin:${PATH}" \ - ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - bash "${script}" --dest "${root}/m8" - [[ "$(cat "${root}/bin-401/count")" == 1 ]] - [[ ! -f "${root}/bin-401/stale" ]] - assert_error_unverified "${root}/m8/paddlefleet_alignment_pin_receipt.json" - - # First two fetches emit Swift 34038640242 curl-56 / early-EOF and leave - # a dest sentinel; the next fetch must see a clean dest. Third fetch is - # real git. Depth=1 and detached pin, not a default-branch clone. - write_git_wrapper "${root}/bin-retry" retry - run PATH="${root}/bin-retry:${root}/bin:${PATH}" \ - ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - bash "${script}" --dest "${root}/ok-retry" - [[ "$(cat "${root}/bin-retry/count")" == 3 ]] - [[ ! -f "${root}/bin-retry/stale" ]] - [[ ! -e "${root}/ok-retry/PaddleFleet/.retry-sentinel" ]] - assert_pin_checkout "${root}/ok-retry/PaddleFleet" "${sha_b}" - python3 - "${root}/ok-retry/paddlefleet_alignment_pin_receipt.json" "${sha_b}" <<'PY' -import json, sys -doc = json.load(open(sys.argv[1])) -assert doc["status"] == "ok" -assert doc["source"]["actual_commit"] == sys.argv[2] -assert doc["source"]["commit_verified"] is True -print("ok-retry receipt exact HEAD verified") -PY - - # Exhausted transient retries: wrapper count is 3, dest cleaned between - # attempts, error receipt stays unverified. - write_git_wrapper "${root}/bin-fail" exhaust - expect_fail "${root}/m9" "git fetch failed" \ - env PATH="${root}/bin-fail:${root}/bin:${PATH}" \ - ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - bash "${script}" --dest "${root}/m9" - [[ "$(cat "${root}/bin-fail/count")" == 3 ]] - [[ ! -f "${root}/bin-fail/stale" ]] - assert_error_unverified "${root}/m9/paddlefleet_alignment_pin_receipt.json" - - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - PADDLEFLEET_WHEEL_URL="${root}/art/py.whl" \ - PADDLEFLEET_WHEEL_SHA256="${py_sha}" \ - PADDLEFLEET_OPS_WHEEL_URL="${root}/art/ops.whl" \ - PADDLEFLEET_OPS_WHEEL_SHA256="${ops_sha}" \ - bash "${script}" --dest "${root}/ok-url" - python3 - "${root}/ok-url/paddlefleet_alignment_pin_receipt.json" "${sha_b}" "${py_sha}" <<'PY' -import json, sys -doc = json.load(open(sys.argv[1])) -sha_b, py_sha = sys.argv[2], sys.argv[3] -assert doc["status"] == "ok" -assert doc["source"]["actual_commit"] == sha_b -assert doc["source"]["commit_verified"] is True -wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") -assert wheel["actual_sha256"] == py_sha -assert wheel["actual_sha256"] != sha_b -assert wheel["digest_verified"] is True -assert wheel.get("built_from_commit") in (None, "") -assert "verified" not in wheel -assert doc["pairing"]["stack_paired_proven"] is False -assert doc["pairing"]["status"] == "unproven" -assert "MinimaxV2.5_EP2" in doc["cases_preserved"] -assert "GLM45Air_EP2" in doc["cases_preserved"] -print("ok-url receipt fields checked") -PY - assert_pin_checkout "${root}/ok-url/PaddleFleet" "${sha_b}" - - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - PADDLEFLEET_BUILD_CMD="mkdir -p '${root}/ok-build/dist' && cp '${root}/art/py.whl' '${root}/ok-build/dist/paddlefleet-0.0.0-py3-none-any.whl'" \ - PADDLEFLEET_BUILD_OPS_CMD="mkdir -p '${root}/ok-build/dist' && cp '${root}/art/ops.whl' '${root}/ok-build/dist/paddlefleet_ops-0.0.0-py3-none-any.whl'" \ - bash "${script}" --dest "${root}/ok-build" - python3 - "${root}/ok-build/paddlefleet_alignment_pin_receipt.json" "${sha_b}" "${py_sha}" <<'PY' -import json, sys -doc = json.load(open(sys.argv[1])) -sha_b, py_sha = sys.argv[2], sys.argv[3] -assert doc["status"] == "ok" -wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") -assert wheel["actual_sha256"] == py_sha -assert wheel["actual_sha256"] != sha_b -assert wheel["origin"] == "build" -assert wheel["built_from_commit"] == sha_b -assert doc["source"]["actual_commit"] == sha_b -assert doc["pairing"]["stack_paired_proven"] is False -assert doc["pairing"]["status"] == "built_from_checked_out_pin" -print("ok-build receipt fields checked") -PY - grep -q "PADDLEFLEET_SOURCE_COMMIT=${sha_b}" "${root}/ok-build/paddlefleet_alignment_pin.env" - - run ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${sha_b}" PADDLEFLEET_GIT_URL="${root}/upstream" \ - bash "${script}" --dest "${root}/ok-source" - python3 - "${root}/ok-source/paddlefleet_alignment_pin_receipt.json" "${sha_b}" "${root}/ok-source" <<'PY' -import json, sys -doc = json.load(open(sys.argv[1])) -sha_b, dest = sys.argv[2], sys.argv[3] -assert doc["status"] == "ok" -assert doc["source"]["actual_commit"] == sha_b -assert doc["source"]["commit_verified"] is True -wheel = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet") -ops = next(a for a in doc["artifacts"] if a["name"] == "paddlefleet_ops") -assert wheel["path"] == f"{dest}/PaddleFleet" -assert ops["path"] == f"{dest}/PaddleFleet/packages/paddlefleet_ops" -assert wheel["origin"] == "source_tree" -assert ops["origin"] == "source_tree" -assert doc["pairing"]["status"] == "source_tree_from_checked_out_pin" -assert doc["pairing"]["stack_paired_proven"] is False -print("ok-source receipt fields checked") -PY - grep -q "PADDLEFLEET_WHEEL_PATH=${root}/ok-source/PaddleFleet$" "${root}/ok-source/paddlefleet_alignment_pin.env" - grep -q "PADDLEFLEET_OPS_WHEEL_PATH=${root}/ok-source/PaddleFleet/packages/paddlefleet_ops$" "${root}/ok-source/paddlefleet_alignment_pin.env" - grep -q "PADDLEFLEET_WHEEL_ORIGIN=source_tree" "${root}/ok-source/paddlefleet_alignment_pin.env" - grep -q "PADDLEFLEET_WHEEL_DIGEST_VERIFIED=false" "${root}/ok-source/paddlefleet_alignment_pin.env" - grep -q "PADDLEFLEET_SOURCE_COMMIT=${sha_b}" "${root}/ok-source/paddlefleet_alignment_pin.env" - assert_pin_checkout "${root}/ok-source/PaddleFleet" "${sha_b}" - - grep -q 'CodeSync/develop/PaddleFleet.tar' "${script}" - grep -q 'PaddleFleet/develop/latest/paddlefleet-0.0.0-py3-none-linux_x86_64.whl' "${script}" - - echo "select_paddlefleet_alignment_pin self-test OK" -} - -if [[ "${RUN_SELF_TEST}" == 1 ]]; then - run_self_test - exit 0 -fi - -mkdir -p "${DEST}" -case "${MODE}" in - stack-paired) fetch_stack_paired ;; - develop) fetch_default ;; - *) fail "unknown ALIGNMENT_PADDLEFLEET_MODE=${MODE} (develop|stack-paired)" ;; -esac diff --git a/scripts/test_alignment_workflow_shell.sh b/scripts/test_alignment_workflow_shell.sh deleted file mode 100755 index 4b666863c5f..00000000000 --- a/scripts/test_alignment_workflow_shell.sh +++ /dev/null @@ -1,153 +0,0 @@ -#!/usr/bin/env bash -# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. -# -# Syntax-check every workflow `run:` block, then execute the extracted -# Get Whl docker-exec body against a failing selector. Independent -# helper --self-test is not this check. - -set -euo pipefail - -ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" -YAML="${ROOT}/../.github/workflows/alignment_model_accuracy.yml" -SELECTOR_PATH="/workspace/Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh" -REQUIRE_PATH="/workspace/Megatron-LM/scripts/require_paddlefleet_selector_ok.sh" -BUILD_MARKER="build megatron-core" - -python3 - "${YAML}" "${ROOT}" "${SELECTOR_PATH}" "${REQUIRE_PATH}" "${BUILD_MARKER}" <<'PY' -import os, re, subprocess, sys, tempfile, textwrap, pathlib, json, stat - -yaml_path, scripts_root, selector_path, require_path, build_marker = sys.argv[1:6] -text = pathlib.Path(yaml_path).read_text() -lines = text.splitlines(True) - -blocks = [] -i = 0 -while i < len(lines): - m = re.match(r"^(\s*)run:\s*\|\s*$", lines[i]) - if not m: - i += 1 - continue - indent = len(m.group(1)) - name = "unnamed" - for j in range(i, -1, -1): - nm = re.match(r"^\s+- name:\s*(.*)$", lines[j]) - if nm: - name = nm.group(1).strip() - break - i += 1 - body = [] - while i < len(lines): - line = lines[i] - if line.strip() == "": - body.append(line) - i += 1 - continue - lead = len(line) - len(line.lstrip(" ")) - if lead <= indent and line.strip(): - break - body.append(line[indent + 2 :] if lead >= indent + 2 else line.lstrip()) - i += 1 - blocks.append((name, "".join(body))) - -if not blocks: - raise SystemExit(f"no run: | blocks in {yaml_path}") - -tmp = pathlib.Path(tempfile.mkdtemp(prefix="yaml-run-")) -print(f"extracted {len(blocks)} run blocks from {yaml_path}") -for n, body in blocks: - p = tmp / (re.sub(r"[^A-Za-z0-9._-]+", "_", n) + ".sh") - p.write_text("#!/usr/bin/env bash\n" + body) - r = subprocess.run(["bash", "-n", str(p)], capture_output=True, text=True) - if r.returncode != 0: - raise SystemExit(f"bash -n FAIL {n}: {r.stderr}") - print(f"bash -n OK run:{n}") - - for m in re.finditer(r"""/bin/bash -c\s+'""", body): - start = m.end() - end = body.find("'\n", start) - if end < 0: - end = body.rfind("'") - inner = body[start:end] - if "'" in inner: - raise SystemExit( - f"nested single quote inside docker exec -c in {n}: " - f"{inner[inner.find(chr(39))-40:inner.find(chr(39))+40]!r}" - ) - inner_p = tmp / (p.stem + ".docker-inner.sh") - inner_p.write_text("#!/usr/bin/env bash\n" + inner + "\n") - r = subprocess.run(["bash", "-n", str(inner_p)], capture_output=True, text=True) - if r.returncode != 0: - raise SystemExit(f"bash -n FAIL docker-inner {n}: {r.stderr}") - print(f"bash -n OK docker-inner:{n} (no nested single quotes)") - -getwhl = next((b for n, b in blocks if n == "Get Whl"), None) -if getwhl is None: - raise SystemExit("Get Whl run block missing") -m = re.search(r"""/bin/bash -c\s+'""", getwhl) -if not m: - raise SystemExit("Get Whl docker exec -c missing") -inner = getwhl[m.end():] -end = inner.rfind("'") -inner = inner[:end] - -ws = tmp / "ws" -(ws / "Megatron-LM/scripts").mkdir(parents=True) -(ws / "upload").mkdir(parents=True) -selector = ws / "Megatron-LM/scripts/select_paddlefleet_alignment_pin.sh" -require_src = pathlib.Path(scripts_root) / "require_paddlefleet_selector_ok.sh" -require_dst = ws / "Megatron-LM/scripts/require_paddlefleet_selector_ok.sh" -require_dst.write_text(require_src.read_text()) -require_dst.chmod(require_dst.stat().st_mode | stat.S_IXUSR) -selector.write_text(textwrap.dedent("""\ - #!/usr/bin/env bash - set -euo pipefail - dest="${2:-/workspace}" - mkdir -p "${dest}" - cat >"${dest}/paddlefleet_alignment_pin_receipt.json" <<'EOF' - {"schema":"paddlefleet-alignment-pin/v1","status":"error","detail":"git clone failed: github.com:443","mode":"stack-paired"} - EOF - echo "[paddlefleet-pin] FAIL: git clone failed: github.com:443" >&2 - echo "::error:: git clone failed: github.com:443" >&2 - exit 1 - """)) -selector.chmod(selector.stat().st_mode | stat.S_IXUSR) -(ws / "Megatron-LM/scripts/dependence").mkdir(parents=True, exist_ok=True) -(ws / "Megatron-LM/scripts/dependence/build.sh").write_text("#!/usr/bin/env bash\necho BUILD_RAN > /workspace/upload/BUILD_RAN\n") -(ws / "Megatron-LM/scripts/dependence/build.sh").chmod(0o755) - -rewritten = inner.replace("/workspace", str(ws)) -rewritten = rewritten.replace("conda activate py_$python_version", "true") -rewritten = rewritten.replace(". /opt/conda/etc/profile.d/conda.sh", "true") -rewritten = "#!/usr/bin/env bash\nexport ALIGNMENT_PADDLEFLEET_MODE=stack-paired\nexport python_version=3.12\n" + rewritten -bin = tmp / "bin" -bin.mkdir() -(bin / "wget").write_text("#!/usr/bin/env bash\necho WGET_RAN \"$@\" >> '%s/WGET_RAN'\nexit 0\n" % ws) -(bin / "python").write_text("#!/usr/bin/env bash\necho PY_RAN \"$@\" >> '%s/PY_RAN'\nexit 0\n" % ws) -(bin / "python3").write_text("#!/usr/bin/env bash\nexec /usr/bin/python3 \"$@\"\n") -(bin / "pip").write_text("#!/usr/bin/env bash\necho PIP_RAN \"$@\" >> '%s/PIP_RAN'\nexit 0\n" % ws) -(bin / "ldconfig").write_text("#!/usr/bin/env bash\nexit 0\n") -for f in bin.iterdir(): - f.chmod(0o755) - -script = tmp / "getwhl.extracted.sh" -script.write_text(rewritten) -script.chmod(0o755) -env = os.environ.copy() -env["PATH"] = str(bin) + ":" + env.get("PATH", "") -env["ALIGNMENT_PADDLEFLEET_MODE"] = "stack-paired" -r = subprocess.run(["bash", str(script)], capture_output=True, text=True, env=env) -log = (r.stdout or "") + (r.stderr or "") -print("extracted Get Whl rc=", r.returncode) -print(log[-2000:]) -if r.returncode == 0: - raise SystemExit("FAIL: extracted Get Whl continued after selector failure") -if (ws / "WGET_RAN").exists(): - raise SystemExit("FAIL: wget ran after selector failure") -if (ws / "upload/BUILD_RAN").exists(): - raise SystemExit("FAIL: build.sh ran after selector failure") -if "selector failed; stop Get Whl" not in log and "git clone failed" not in log: - raise SystemExit("FAIL: extracted Get Whl did not surface selector failure") -print("extracted Get Whl fixture: selector fail stopped remaining wheels and build") -print("workflow shell checks OK") -PY -echo "alignment workflow shell PATH_PASS (extracted YAML, not helper-only)" diff --git a/scripts/test_ops_submodule_preflight.sh b/scripts/test_ops_submodule_preflight.sh deleted file mode 100644 index 1f08b7253dc..00000000000 --- a/scripts/test_ops_submodule_preflight.sh +++ /dev/null @@ -1,78 +0,0 @@ -#!/usr/bin/env bash -# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. -set -euo pipefail -ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")/.." && pwd)" -"${PYTHON_BIN:-python3}" - "$ROOT" <<'PY' -import os -import pathlib -import shutil -import subprocess -import sys -import tempfile - -root = pathlib.Path(sys.argv[1]) -workflow = (root / '.github/workflows/alignment_model_accuracy.yml').read_text() -start = workflow.index(' bash Megatron-LM/scripts/prepare_paddlefleet_ops_submodules.sh') -end = workflow.index('\n model_acc_align_exit_code=', start) -commands = workflow[start:end] -assert commands.count('setup_venvs.sh') == 2 -assert commands.count('run_alignment_test.sh') == 1 -assert ' set -eo pipefail' in workflow[:start] -real_timeout = shutil.which('timeout') -assert real_timeout -with tempfile.TemporaryDirectory() as directory: - work = pathlib.Path(directory) - (work / 'bin').mkdir() - (work / 'ops').mkdir() - (work / 'Megatron-LM').symlink_to(root) - fleet = work / 'PaddleFleet/scripts/alignment_model_accuracy' - fleet.mkdir(parents=True) - setup = fleet / 'setup_venvs.sh' - setup.write_text('''#!/usr/bin/env bash -set -eu -if [[ ${1:-} != --prepare-ops-submodules ]]; then echo setup >> "$TRACE"; exit 0; fi -[[ $2 == "$OPS" ]] -echo prepare >> "$TRACE" -n=$(grep -c prepare "$TRACE") -case "$CASE" in - success) exit 0;; - transient) [[ $n -ge 2 ]];; - persistent) exit 23;; - timeout) /bin/sleep 5;; -esac -''') - (fleet / 'run_alignment_test.sh').write_text('echo train >> "$TRACE"\n') - sleeper = work / 'bin/sleep' - sleeper.write_text('#!/usr/bin/env bash\necho "sleep:$1" >> "$TRACE"\n') - timer = work / 'bin/timeout' - timer.write_text('''#!/usr/bin/env bash -set -eu -[[ $1 == --kill-after=30s && $2 == 15m ]] -echo bound >> "$TRACE" -shift 2 -if [[ $CASE == timeout ]]; then exec "$REAL_TIMEOUT" --kill-after=0.1s 0.1s "$@"; fi -exec "$@" -''') - sleeper.chmod(0o755) - timer.chmod(0o755) - env = dict(os.environ, PATH=str(work / 'bin') + ':' + os.environ['PATH'], - TRACE=str(work / 'trace'), OPS=str(work / 'ops'), REAL_TIMEOUT=real_timeout) - for case, count, code in [('success', 1, 0), ('transient', 2, 0), - ('persistent', 3, 23), ('timeout', 3, 124), - ('wheel', 0, 0)]: - trace = work / 'trace' - trace.write_text('') - env.update(CASE=case, PADDLEFLEET_OPS_WHEEL_PATH=env['OPS'] if case != 'wheel' else str(work / 'ops.whl')) - result = subprocess.run(['bash', '-eo', 'pipefail', '-c', commands], cwd=work, - env=env, capture_output=True, text=True, timeout=10) - lines = trace.read_text().splitlines() - assert result.returncode == code, (case, result.returncode, result.stderr) - assert lines.count('prepare') == count, (case, lines) - assert lines.count('bound') == count, (case, lines) - assert [x for x in lines if x.startswith('sleep:')] == ['sleep:15', 'sleep:30'][:max(0, count - 1)], (case, lines) - assert ('setup' in lines) == (code == 0), (case, lines) - assert ('train' in lines) == (code == 0), (case, lines) - assert lines.count('setup') <= 1 and lines.count('train') <= 1 - print(f'PASS: extracted workflow {case}, attempts={count}, exit={code}') -print('All source-submodule preflight fixtures passed') -PY diff --git a/scripts/test_paddlefleet_pin_handoff.sh b/scripts/test_paddlefleet_pin_handoff.sh deleted file mode 100755 index 9b01769b882..00000000000 --- a/scripts/test_paddlefleet_pin_handoff.sh +++ /dev/null @@ -1,111 +0,0 @@ -#!/usr/bin/env bash -# Copyright (c) 2026 PaddlePaddle Authors. All Rights Reserved. -# -# End-to-end path check: selector source-mode -> new-shell consume -> -# setup_venvs *path* consumer. This proves docker-exec handoff of source -# paths. It is NOT a uv install / real setup_venvs / numerical CI run. -# Isolated selector --self-test is not enough for the step boundary. - -set -euo pipefail - -ROOT="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" -SELECTOR="${ROOT}/select_paddlefleet_alignment_pin.sh" -CONSUME="${ROOT}/consume_paddlefleet_alignment_pin.sh" - -tmp="$(mktemp -d)" -trap 'rm -rf "${tmp}"' EXIT - -git init -q "${tmp}/upstream" -git -C "${tmp}/upstream" config user.email test@example.com -git -C "${tmp}/upstream" config user.name test -mkdir -p "${tmp}/upstream/packages/paddlefleet_ops" -printf '%s\n' '[project]' 'name = "paddlefleet"' >"${tmp}/upstream/pyproject.toml" -printf '%s\n' '[project]' 'name = "paddlefleet-ops"' >"${tmp}/upstream/packages/paddlefleet_ops/pyproject.toml" -echo src >"${tmp}/upstream/README" -git -C "${tmp}/upstream" add README pyproject.toml packages -git -C "${tmp}/upstream" commit -q -m pin -PIN="$(git -C "${tmp}/upstream" rev-parse HEAD)" - -# Step A: Get Whl equivalent (selector). -ALIGNMENT_PADDLEFLEET_MODE=stack-paired \ - PADDLEFLEET_PIN_SHA="${PIN}" \ - PADDLEFLEET_GIT_URL="${tmp}/upstream" \ - bash "${SELECTOR}" --dest "${tmp}/ws" - -test -f "${tmp}/ws/paddlefleet_alignment_pin.env" -test -d "${tmp}/ws/PaddleFleet" - -# Step B: new docker exec — drop selector shell state, keep only files. -# Requested mode/pin stay on the caller; leftover develop env must not win. -unset PADDLEFLEET_WHEEL_PATH PADDLEFLEET_OPS_WHEEL_PATH PADDLEFLEET_SOURCE_COMMIT || true -ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${PIN}" \ - bash "${CONSUME}" --env "${tmp}/ws/paddlefleet_alignment_pin.env" --out "${tmp}/ws/consumed.env" - -# Negative: leftover develop env + requested stack-paired must fail closed. -mkdir -p "${tmp}/dev" -echo dummy >"${tmp}/dev/paddlefleet-0.0.0-py3-none-linux_x86_64.whl" -echo dummy >"${tmp}/dev/paddlefleet_ops-0.0.0-cp312-cp312-linux_x86_64.whl" -cat >"${tmp}/dev.env" <&2 - exit 1 -fi - -# Negative: selector clone/fail writes error receipt and no env. Consume -# must name selector failure. Missing env here is the consequence, not -# proof that a generated env failed to cross docker exec. -REQUIRE="${ROOT}/require_paddlefleet_selector_ok.sh" -mkdir -p "${tmp}/sel-fail" -cat >"${tmp}/sel-fail/paddlefleet_alignment_pin_receipt.json" <<'EOF' -{"schema":"paddlefleet-alignment-pin/v1","status":"error","detail":"git clone failed: github.com:443","mode":"stack-paired"} -EOF -if bash "${REQUIRE}" "${tmp}/sel-fail" >"${tmp}/sel-fail.require.out" 2>"${tmp}/sel-fail.require.err"; then - echo "handoff FAIL: require_ok accepted error receipt" >&2 - exit 1 -fi -grep -q "selector receipt status=" "${tmp}/sel-fail.require.err" -if ALIGNMENT_PADDLEFLEET_MODE=stack-paired PADDLEFLEET_PIN_SHA="${PIN}" \ - bash "${CONSUME}" --env "${tmp}/sel-fail/paddlefleet_alignment_pin.env" \ - --out "${tmp}/sel-fail.consumed.env" 2>"${tmp}/sel-fail.err"; then - echo "handoff FAIL: selector error receipt was consumed" >&2 - exit 1 -fi -grep -q "because selector failed" "${tmp}/sel-fail.err" -if grep -q "selector wrote ok receipt but env did not cross docker exec" "${tmp}/sel-fail.err"; then - echo "handoff FAIL: selector error misclassified as env-handoff" >&2 - exit 1 -fi - -# Step C: path consumer only. Does not run uv or setup_venvs.sh. -stub_setup="${tmp}/setup_path_consumer.sh" -cat >"${stub_setup}" <<'STUB' -#!/usr/bin/env bash -set -euo pipefail -# Mirrors setup_venvs.sh reading PADDLEFLEET_WHEEL_PATH. Path presence only. -PADDLEFLEET_WHEEL="${PADDLEFLEET_WHEEL_PATH:?missing PADDLEFLEET_WHEEL_PATH}" -PADDLEFLEET_OPS_WHEEL="${PADDLEFLEET_OPS_WHEEL_PATH:?missing PADDLEFLEET_OPS_WHEEL_PATH}" -[[ -d "${PADDLEFLEET_WHEEL}" || -f "${PADDLEFLEET_WHEEL}" ]] || { echo "missing ${PADDLEFLEET_WHEEL}" >&2; exit 1; } -[[ -d "${PADDLEFLEET_OPS_WHEEL}" || -f "${PADDLEFLEET_OPS_WHEEL}" ]] || { echo "missing ${PADDLEFLEET_OPS_WHEEL}" >&2; exit 1; } -echo "PATH_CONSUMER paddlefleet=${PADDLEFLEET_WHEEL}" -echo "PATH_CONSUMER ops=${PADDLEFLEET_OPS_WHEEL}" -echo "PATH_CONSUMER not_uv_install=true" -STUB -chmod +x "${stub_setup}" - -set -a -# shellcheck disable=SC1090 -. "${tmp}/ws/consumed.env" -set +a -bash "${stub_setup}" | tee "${tmp}/setup.out" -grep -q "PATH_CONSUMER paddlefleet=${tmp}/ws/PaddleFleet" "${tmp}/setup.out" -grep -q "packages/paddlefleet_ops" "${tmp}/setup.out" -grep -q "not_uv_install=true" "${tmp}/setup.out" - -echo "paddlefleet pin handoff PATH_PASS pin=${PIN} (not uv install, not CI)" From d8b6841697350023f7b71a5792dd636d9cf63023 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Tue, 8 Sep 2026 23:04:39 -0700 Subject: [PATCH 23/27] refactor: select embedding accumulation from model configuration Signed-off-by: Zhan Rongrui --- megatron/core/model_parallel_config.py | 3 + megatron/core/tensor_parallel/layers.py | 28 ++-- .../test_accuracy_tp1_migration.py | 131 +++++++++++++----- 3 files changed, 106 insertions(+), 56 deletions(-) diff --git a/megatron/core/model_parallel_config.py b/megatron/core/model_parallel_config.py index 88bb070e105..1e1c367d89c 100644 --- a/megatron/core/model_parallel_config.py +++ b/megatron/core/model_parallel_config.py @@ -150,6 +150,9 @@ class ModelParallelConfig: be synchronized. """ + use_accuracy_compatible: bool = False + """Use explicit accuracy-compatible arithmetic in model layers.""" + deterministic_mode: bool = False """If true, code that has deterministic execution will be chosen. This usually means slower execution, but is good for debugging and testing. Defaults to False.""" diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 9bd553ef9dd..6b6d8aaf22d 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -70,12 +70,14 @@ class _EmbedFp32MainGrad(torch.autograd.Function): @staticmethod def forward(ctx, weight, ids): + """Look up embeddings while retaining the accumulator owner.""" ctx.save_for_backward(ids) ctx.weight_ref = weight return weight[ids] @staticmethod def backward(ctx, grad_output): + """Accumulate repeated-index gradients into the FP32 master buffer.""" (ids,) = ctx.saved_tensors weight = ctx.weight_ref prev = torch.is_grad_enabled() @@ -85,9 +87,7 @@ def backward(ctx, grad_output): unique_ids, inv = torch.unique(ids_flat, return_inverse=True) uniq_w = weight.detach()[unique_ids].clone().requires_grad_(True) looked = uniq_w[inv.reshape(ids.shape)] - (gw,) = torch.autograd.grad( - looked, uniq_w, grad_outputs=grad_output, allow_unused=True - ) + (gw,) = torch.autograd.grad(looked, uniq_w, grad_outputs=grad_output, allow_unused=True) finally: torch.set_grad_enabled(prev) if gw is None: @@ -96,9 +96,7 @@ def backward(ctx, grad_output): if hasattr(weight, "main_grad") and weight.main_grad is not None: weight.main_grad.index_add_(0, unique_ids, fp) else: - acc = torch.zeros( - weight.shape, dtype=torch.float32, device=weight.device - ) + acc = torch.zeros(weight.shape, dtype=torch.float32, device=weight.device) acc.index_add_(0, unique_ids, fp) weight.main_grad = acc if hasattr(weight, "grad_added_to_main_grad"): @@ -284,7 +282,7 @@ def __init__( ) ) self.num_embeddings_per_partition = self.vocab_end_index - self.vocab_start_index - self.deterministic_mode = config.deterministic_mode or _use_accuracy_compatible() + self.deterministic_mode = config.deterministic_mode or config.use_accuracy_compatible self.config = config self.use_inference_optimized_reduce_scatter = ( @@ -347,11 +345,7 @@ def forward(self, input_): # Get the embeddings. if self.deterministic_mode: _tp_size = 1 if self.tp_group is None else self.tp_group.size() - if ( - _use_accuracy_compatible() - and _tp_size <= 1 - and os.environ.get("MODEL_REPRO_TWO_FP32_ACCUM", "") == "1" - ): + if self.config.use_accuracy_compatible and _tp_size <= 1: output_parallel = _EmbedFp32MainGrad.apply(self.weight, masked_input) else: output_parallel = self.weight[masked_input] @@ -1138,10 +1132,7 @@ def forward( or self.disable_grad_reduce ): input_parallel = input_ - elif ( - _use_accuracy_compatible() - and (self.tp_group is None or self.tp_group.size() <= 1) - ): + elif _use_accuracy_compatible() and (self.tp_group is None or self.tp_group.size() <= 1): input_parallel = input_ else: input_parallel = copy_to_tensor_model_parallel_region(input_, group=self.tp_group) @@ -1470,10 +1461,7 @@ def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: output_ = reduce_scatter_to_sequence_parallel_region( output_parallel, group=self.tp_group ) - elif ( - _use_accuracy_compatible() - and (self.tp_group is None or self.tp_group.size() <= 1) - ): + elif _use_accuracy_compatible() and (self.tp_group is None or self.tp_group.size() <= 1): output_ = output_parallel else: output_ = reduce_from_tensor_model_parallel_region(output_parallel, group=self.tp_group) diff --git a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py index 3399ed7cf60..d66edc26411 100644 --- a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py +++ b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py @@ -1,7 +1,9 @@ """CUDA unit tests for native C2 TP1 accuracy-compatible migration.""" + from __future__ import annotations import ast +import os import sys import unittest from pathlib import Path @@ -78,8 +80,10 @@ def _load_named(rel: str, name: str, extra_ns=None, class_name=None): target = next( node for node in body - if isinstance(node, ast.ClassDef) and node.name == name - or isinstance(node, ast.FunctionDef) and node.name == name + if isinstance(node, ast.ClassDef) + and node.name == name + or isinstance(node, ast.FunctionDef) + and node.name == name ) if isinstance(target, ast.FunctionDef): target.decorator_list = [] @@ -96,9 +100,7 @@ def _load_named(rel: str, name: str, extra_ns=None, class_name=None): "get_tensor_model_parallel_group_if_none": lambda g: g, "LinearWithGradAccumulationAndAsyncCommunication": _SentinelApply, "_GatherFromModelParallelRegion": _SentinelGather, - "parallel_state": SimpleNamespace( - get_tensor_model_parallel_world_size=lambda: _TP["size"] - ), + "parallel_state": SimpleNamespace(get_tensor_model_parallel_world_size=lambda: _TP["size"]), "custom_backward": _custom_backward, "Variable": torch.autograd.Variable, } @@ -108,27 +110,18 @@ def _load_named(rel: str, name: str, extra_ns=None, class_name=None): return ns[name] -_EmbedFp32MainGrad = _load_named( - "megatron/core/tensor_parallel/layers.py", - "_EmbedFp32MainGrad", -) +_EmbedFp32MainGrad = _load_named("megatron/core/tensor_parallel/layers.py", "_EmbedFp32MainGrad") linear_with_grad_accumulation_and_async_allreduce = _load_named( - "megatron/core/tensor_parallel/layers.py", - "linear_with_grad_accumulation_and_async_allreduce", + "megatron/core/tensor_parallel/layers.py", "linear_with_grad_accumulation_and_async_allreduce" ) linear_with_grad_accumulation_and_async_allreduce.warned = True gather_from_tensor_model_parallel_region = _load_named( - "megatron/core/tensor_parallel/mappings.py", - "gather_from_tensor_model_parallel_region", + "megatron/core/tensor_parallel/mappings.py", "gather_from_tensor_model_parallel_region" ) deallocate_output_tensor = _load_named( - "megatron/core/pipeline_parallel/schedules.py", - "deallocate_output_tensor", -) -backward_step = _load_named( - "megatron/core/pipeline_parallel/schedules.py", - "backward_step", + "megatron/core/pipeline_parallel/schedules.py", "deallocate_output_tensor" ) +backward_step = _load_named("megatron/core/pipeline_parallel/schedules.py", "backward_step") def _ref_embed_fp32_wgrad(weight_bf16, ids, grad_out): @@ -144,19 +137,91 @@ def _cuda_bf16(values, shape, device): @unittest.skipUnless(torch.cuda.is_available(), "CUDA required") class TestEmbedFp32MainGradCuda(unittest.TestCase): + def test_embedding_configuration_controls_gradient_destination(self): + forward = _load_named( + "megatron/core/tensor_parallel/layers.py", + "forward", + {"_EmbedFp32MainGrad": _EmbedFp32MainGrad}, + class_name="VocabParallelEmbedding", + ) + ids = torch.tensor([1, 1, 3], device="cuda") + instances = [] + for enabled in (False, True): + weight = torch.ones(4, 8, device="cuda", dtype=torch.bfloat16, requires_grad=True) + weight.main_grad = torch.zeros_like(weight, dtype=torch.float32) + instances.append( + SimpleNamespace( + config=SimpleNamespace(use_accuracy_compatible=enabled), + deterministic_mode=True, + tp_group=_FakeGroup(1), + weight=weight, + reduce_scatter_embeddings=False, + ) + ) + for instance in instances: + enabled = instance.config.use_accuracy_compatible + with patch.dict( + os.environ, + { + "MODEL_REPRO_TWO_FP32_ACCUM": str(int(not enabled)), + "USE_ACCURACY_COMPATIBLE": str(int(not enabled)), + }, + ): + output = forward(instance, ids) + output.sum().backward() + expected = torch.zeros(4, 8, device="cuda", dtype=torch.float32) + expected[1] = 2 + expected[3] = 1 + if enabled: + self.assertIsNone(instance.weight.grad) + torch.testing.assert_close(instance.weight.main_grad, expected, atol=0, rtol=0) + else: + torch.testing.assert_close(instance.weight.grad.float(), expected, atol=0, rtol=0) + self.assertEqual(torch.count_nonzero(instance.weight.main_grad).item(), 0) + def test_repeated_indices_matches_independent_full_table_autograd(self): device = torch.device("cuda") vocab, dim = 8, 4 weight = _cuda_bf16( - [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, - 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32], + [ + 1, + 2, + 3, + 4, + 5, + 6, + 7, + 8, + 9, + 10, + 11, + 12, + 13, + 14, + 15, + 16, + 17, + 18, + 19, + 20, + 21, + 22, + 23, + 24, + 25, + 26, + 27, + 28, + 29, + 30, + 31, + 32, + ], (vocab, dim), device, ) ids = torch.tensor([[1, 3, 1, 5], [3, 3, 0, 1]], device=device) - grad_out = _cuda_bf16( - list(range(1, 33)), (2, 4, dim), device - ) + grad_out = _cuda_bf16(list(range(1, 33)), (2, 4, dim), device) w = weight.clone().requires_grad_(True) w.main_grad = torch.zeros(vocab, dim, device=device, dtype=torch.float32) w.grad_added_to_main_grad = False @@ -174,9 +239,7 @@ def test_repeated_indices_matches_independent_full_table_autograd(self): def test_main_grad_accumulates_two_backwards_unused_rows_stay_zero(self): device = torch.device("cuda") vocab, dim = 6, 3 - weight = _cuda_bf16( - list(range(1, 19)), (vocab, dim), device - ) + weight = _cuda_bf16(list(range(1, 19)), (vocab, dim), device) ids_a = torch.tensor([2, 2, 4], device=device) ids_b = torch.tensor([4, 1, 2], device=device) go_a = _cuda_bf16([1, 2, 3, 4, 5, 6, 7, 8, 9], (3, dim), device) @@ -211,15 +274,9 @@ def tearDown(self): def test_tp1_forward_dgrad_wgrad_bias_matches_f_linear(self): device = torch.device("cuda") x = _cuda_bf16( - [1, 2, 0, -1, 1, 0, 2, 1, -2, 0, 1, 1, 2, 0, 1], - (5, 3), - device, - ).requires_grad_(True) - w = _cuda_bf16( - [1, 0, -1, 0, 1, 1, 1, -1, 0, 0, 1, -1], - (4, 3), - device, + [1, 2, 0, -1, 1, 0, 2, 1, -2, 0, 1, 1, 2, 0, 1], (5, 3), device ).requires_grad_(True) + w = _cuda_bf16([1, 0, -1, 0, 1, 1, 1, -1, 0, 0, 1, -1], (4, 3), device).requires_grad_(True) b = _cuda_bf16([1, -1, 0, 2], (4,), device).requires_grad_(True) out = linear_with_grad_accumulation_and_async_allreduce( x, w, b, False, False, False, None, 0, _FakeGroup(1) @@ -230,7 +287,9 @@ def test_tp1_forward_dgrad_wgrad_bias_matches_f_linear(self): ref = F.linear(xref, wref, bref) torch.testing.assert_close(out, ref, atol=0, rtol=0) self.assertIsNone(_SentinelApply.last) - go = _cuda_bf16([1, 1, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 0, 1, 1], (5, 4), device) + go = _cuda_bf16( + [1, 1, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 1, 1, 0, 1, 0, 0, 1, 1], (5, 4), device + ) out.backward(go) F.linear(xref, wref, bref).backward(go) torch.testing.assert_close(x.grad, xref.grad, atol=0, rtol=0) From d0e8ae477b3c6e0024541932cb74b13ab24578f7 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Thu, 10 Sep 2026 01:46:54 -0700 Subject: [PATCH 24/27] fix: preserve accuracy-compatible expert padding and loss defaults Signed-off-by: Zhan Rongrui --- megatron/core/transformer/moe/experts.py | 26 +++++-- .../core/transformer/transformer_config.py | 7 ++ .../moe/test_sequential_expert_padding.py | 76 +++++++++++++++++++ 3 files changed, 104 insertions(+), 5 deletions(-) create mode 100644 tests/unit_tests/transformer/moe/test_sequential_expert_padding.py diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 85b0e657b24..64e7115747f 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -1180,16 +1180,17 @@ class _SeqMLPProxy: Required by upper layers (e.g. ms-swift's ``GPTBridge._set_mlp_state``) that probe ``mg_mlp.linear_fc1`` / ``linear_fc2`` like GroupedMLP. """ + def __init__(self, experts, attr): self._experts = experts self._attr = attr def __getattr__(self, name): if name.startswith('weight'): - idx = int(name[len('weight'):]) + idx = int(name[len('weight') :]) return getattr(self._experts[idx], self._attr).weight if name.startswith('bias'): - idx = int(name[len('bias'):]) + idx = int(name[len('bias') :]) return getattr(self._experts[idx], self._attr).bias raise AttributeError(name) @@ -1275,8 +1276,7 @@ def _grad_hook(grad, _w=weight, _i=saved_inp): return grad with torch.no_grad(): wg = torch.matmul( - grad.detach().to(torch.float32).transpose(0, 1), - _i.to(torch.float32), + grad.detach().to(torch.float32).transpose(0, 1), _i.to(torch.float32) ) prev = getattr(_w, '_run_torch_expert_fp32_wgrad', None) if prev is None: @@ -1293,7 +1293,6 @@ def _grad_hook(grad, _w=weight, _i=saved_inp): for lin in (expert.linear_fc1, expert.linear_fc2): lin.register_forward_hook(_make_forward_hook(lin)) - def _pad_tensor_for_quantization(self, hidden, probs): """Padding tensor shape to multiples of 16/32.""" actual_num_tokens = hidden.shape[0] @@ -1349,12 +1348,29 @@ def forward( output_local_list = [] for expert, tokens, probs in zip(self.local_experts, tokens_list, probs_list): + # The unfused Paddle expert pads tiny GEMMs to 32 rows. The + # grouped-storage fallback uses real token counts instead. + num_real_tokens = tokens.shape[0] + pad_small_expert = ( + self.config.use_accuracy_compatible + and not self.config.moe_grouped_gemm + and not (self.config.fp8 or self.config.fp4) + and 0 < num_real_tokens < 17 + ) + if pad_small_expert: + num_pad_tokens = 32 - num_real_tokens + tokens = torch.cat( + (tokens, tokens.new_zeros(num_pad_tokens, tokens.shape[1])), dim=0 + ) + probs = torch.cat((probs, probs.new_zeros(num_pad_tokens)), dim=0) if self.config.fp8 or self.config.fp4: hidden, probs = self._pad_tensor_for_quantization(tokens, probs) output, output_bias = expert(hidden, probs) output = output[: tokens.shape[0]] else: output, output_bias = expert(tokens, probs) + if pad_small_expert: + output = output[:num_real_tokens] output_local_list.append(output) output_local = torch.cat(output_local_list, dim=0) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index a8e083341c2..74904e5ef57 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -288,6 +288,13 @@ class TransformerConfig(ModelParallelConfig): """Whether cross entropy loss is calculated over the actual number of non-padded tokens in the global batch, versus the default behavior of assuming all tokens are non-padded.""" + accuracy_compatible_loss_sum_dtype: Literal["float32", "float64"] = "float64" + """Token-loss accumulation dtype in accuracy-compatible training. + + Preserve FP64 accumulation by default. Model providers can select FP32 + when that is the reference loss-reduction contract. + """ + multi_latent_attention: bool = False """Whether to use multi-latent attention.""" diff --git a/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py b/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py new file mode 100644 index 00000000000..65bb907d052 --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_sequential_expert_padding.py @@ -0,0 +1,76 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import pytest +import torch + +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.mlp import MLPSubmodules +from megatron.core.transformer.moe.experts import SequentialMLP +from megatron.core.transformer.transformer_config import TransformerConfig +from tests.unit_tests.test_utilities import Utils + + +class TestSequentialExpertPadding: + def setup_method(self): + Utils.initialize_model_parallel(tensor_model_parallel_size=1, expert_model_parallel_size=1) + model_parallel_cuda_manual_seed(123) + + def teardown_method(self): + Utils.destroy_model_parallel() + + @pytest.mark.parametrize("enabled,grouped", [(False, False), (True, False), (True, True)]) + @pytest.mark.parametrize("first_count", [0, 3, 16, 17]) + def test_expert_gemm_rows_preserve_storage_contract(self, enabled, grouped, first_count): + config = TransformerConfig( + num_layers=1, + hidden_size=32, + num_attention_heads=4, + ffn_hidden_size=64, + moe_ffn_hidden_size=64, + num_moe_experts=2, + moe_router_topk=1, + moe_router_pre_softmax=True, + add_bias_linear=False, + gated_linear_unit=True, + activation_func=torch.nn.functional.silu, + bias_activation_fusion=False, + params_dtype=torch.bfloat16, + use_accuracy_compatible=enabled, + moe_grouped_gemm=grouped, + ) + experts = SequentialMLP( + 2, + config, + MLPSubmodules(linear_fc1=ColumnParallelLinear, linear_fc2=RowParallelLinear), + pg_collection=ProcessGroupCollection.use_mpu_process_groups(), + ).cuda() + observed = [] + + def record_rows(module, inputs): + tokens, probs = inputs + observed.append((tokens.shape[0], probs.shape[0])) + + handles = [ + expert.register_forward_pre_hook(record_rows) for expert in experts.local_experts + ] + counts = torch.tensor([first_count, 19], dtype=torch.int64) + tokens = torch.randn( + first_count + 19, 32, device="cuda", dtype=torch.bfloat16, requires_grad=True + ) + probs = torch.ones(first_count + 19, device="cuda", dtype=torch.float32, requires_grad=True) + try: + output, bias = experts(tokens, counts, probs) + expected_rows = 32 if enabled and not grouped and 0 < first_count < 17 else first_count + assert observed == [(expected_rows, expected_rows), (19, 19)] + assert output.shape == tokens.shape + assert bias is None + output.float().sum().backward() + assert tokens.grad.shape == tokens.shape + assert probs.grad.shape == probs.shape + assert torch.isfinite(tokens.grad).all() + assert torch.isfinite(probs.grad).all() + finally: + for handle in handles: + handle.remove() From 69cb9568e0d0428eac768becd3d517541a7a5b47 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui <46243324+zrr1999@users.noreply.github.com> Date: Thu, 10 Sep 2026 03:53:56 -0700 Subject: [PATCH 25/27] fix: make accuracy-compatible clipping independent of gradient partitions Accumulate FP32 squares in exact integer bins on device, reduce over existing owner groups, and round once before FP32 clipping. Signed-off-by: Zhan Rongrui --- megatron/core/optimizer/clip_grads.py | 45 +++++- megatron/core/optimizer/optimizer.py | 38 +++++- megatron/core/optimizer/optimizer_config.py | 3 + megatron/core/optimizer/reproducible_norm.py | 128 ++++++++++++++++++ .../optimizer/test_reproducible_norm.py | 109 +++++++++++++++ 5 files changed, 313 insertions(+), 10 deletions(-) create mode 100644 megatron/core/optimizer/reproducible_norm.py create mode 100644 tests/unit_tests/optimizer/test_reproducible_norm.py diff --git a/megatron/core/optimizer/clip_grads.py b/megatron/core/optimizer/clip_grads.py index 3c5491d39a1..6345991e057 100644 --- a/megatron/core/optimizer/clip_grads.py +++ b/megatron/core/optimizer/clip_grads.py @@ -50,12 +50,37 @@ from ..tensor_parallel import param_is_not_tensor_parallel_duplicate from ..transformer.module import param_is_not_shared from ..utils import get_data_parallel_group_if_dtensor, to_local_if_dtensor +from .reproducible_norm import ReproducibleL2Norm + + +@torch.no_grad() +def get_reproducible_grad_norm_bins( + grads_for_norm: List[torch.Tensor], + grad_stats_parallel_group: torch.distributed.ProcessGroup | None, +) -> torch.Tensor: + """Reduce exact FP32-square bins using the existing gradient ownership groups.""" + grads_for_norm = list(grads_for_norm) + data_parallel_group = None + for grad in grads_for_norm: + data_parallel_group = get_data_parallel_group_if_dtensor(grad, data_parallel_group) + grads_for_norm = [to_local_if_dtensor(grad) for grad in grads_for_norm] + accumulator = ReproducibleL2Norm(grads_for_norm[0].device if grads_for_norm else None) + bins = accumulator.zeros() + for grad in grads_for_norm: + if grad.layout != torch.strided: + raise TypeError("Reproducible clipping requires dense FP32 gradients") + bins = accumulator.accumulate(bins, grad) + if data_parallel_group: + torch.distributed.all_reduce(bins, group=data_parallel_group) + torch.distributed.all_reduce(bins, group=grad_stats_parallel_group) + return bins def get_grad_norm_fp32( grads_for_norm: Union[List[torch.Tensor], torch.Tensor], norm_type: Union[int, float] = 2, grad_stats_parallel_group: Optional[torch.distributed.ProcessGroup] = None, + use_accuracy_compatible: bool = False, ) -> float: """Calculate the p-norm of gradients in FP32 precision. @@ -80,6 +105,12 @@ def get_grad_norm_fp32( if isinstance(grads_for_norm, torch.Tensor): grads_for_norm = [grads_for_norm] + if use_accuracy_compatible: + if float(norm_type) != 2.0: + raise ValueError("Reproducible clipping supports only the L2 norm") + bins = get_reproducible_grad_norm_bins(grads_for_norm, grad_stats_parallel_group) + return ReproducibleL2Norm(bins.device).finish(bins)[0] + data_parallel_group = None for grad in grads_for_norm: data_parallel_group = get_data_parallel_group_if_dtensor(grad, data_parallel_group) @@ -184,12 +215,14 @@ def clip_grad_by_total_norm_fp32( dummy_overflow_buf = torch.zeros(1, dtype=torch.int, device='cuda') if isinstance(clip_coeff, torch.Tensor): clip_coeff.clamp_max_(1.0) - assert ( - multi_tensor_scale_tensor_impl is not None - ), "clip_coeff is tensor type. But multi_tensor_scale_tensor not available." - multi_tensor_applier( - multi_tensor_scale_tensor_impl, dummy_overflow_buf, [grads, grads], clip_coeff - ) + if multi_tensor_scale_tensor_impl is not None: + multi_tensor_applier( + multi_tensor_scale_tensor_impl, dummy_overflow_buf, [grads, grads], clip_coeff + ) + else: + multi_tensor_applier( + multi_tensor_scale_impl, dummy_overflow_buf, [grads, grads], clip_coeff.item() + ) elif clip_coeff < 1.0: multi_tensor_applier( multi_tensor_scale_impl, dummy_overflow_buf, [grads, grads], clip_coeff diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index 4a74328d0d9..e00c17b200f 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -48,9 +48,15 @@ from ..dist_checkpointing.utils import add_prefix_for_sharding from ..transformer.module import param_is_not_shared from ..utils import log_single_rank -from .clip_grads import clip_grad_by_total_norm_fp32, count_zeros_fp32, get_grad_norm_fp32 +from .clip_grads import ( + clip_grad_by_total_norm_fp32, + count_zeros_fp32, + get_grad_norm_fp32, + get_reproducible_grad_norm_bins, +) from .grad_scaler import MegatronGradScaler from .optimizer_config import OptimizerConfig +from .reproducible_norm import ReproducibleL2Norm logger = getLogger(__name__) @@ -296,7 +302,10 @@ def get_grad_norm(self): """Compute and return grad norm.""" grads_for_norm = self.get_grads_for_grad_norm() total_norm = get_grad_norm_fp32( - grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group() + grads_for_norm, + grad_stats_parallel_group=self.get_grad_stats_parallel_group(), + use_accuracy_compatible=self.config.use_accuracy_compatible + and self.config.clip_grad > 0, ) return total_norm @@ -308,7 +317,10 @@ def _compute_grad_norms_by_group(self) -> Dict[str, float]: if self.has_grad_norm_group(grad_norm_group): grouped_grads = self.get_grads_for_grad_norm(grad_norm_group) group_grad_norm = get_grad_norm_fp32( - grouped_grads, grad_stats_parallel_group=self.get_grad_stats_parallel_group() + grouped_grads, + grad_stats_parallel_group=self.get_grad_stats_parallel_group(), + use_accuracy_compatible=self.config.use_accuracy_compatible + and self.config.clip_grad > 0, ) self.grad_norms_by_group[grad_norm_group] = group_grad_norm return self.grad_norms_by_group @@ -326,7 +338,10 @@ def clip_grad_norm(self, clip_grad: float) -> float: else: grads_for_norm = [] grad_norm = get_grad_norm_fp32( - grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group() + grads_for_norm, + grad_stats_parallel_group=self.get_grad_stats_parallel_group(), + use_accuracy_compatible=self.config.use_accuracy_compatible + and self.config.clip_grad > 0, ) if clip_grad > 0.0 and params: @@ -1566,8 +1581,21 @@ def get_grad_stats_parallel_group(self) -> torch.distributed.ProcessGroup: ) return self.chained_optimizers[0].get_grad_stats_parallel_group() + @torch.no_grad() + def _get_reproducible_grad_norm(self, grad_norm_group=None): + bins = None + for optimizer in self.chained_optimizers: + part = get_reproducible_grad_norm_bins( + optimizer.get_grads_for_grad_norm(grad_norm_group), + optimizer.get_grad_stats_parallel_group(), + ) + bins = part if bins is None else bins + part + return ReproducibleL2Norm(bins.device).finish(bins)[0] + @torch.no_grad() def get_grad_norm(self): + if self.config.use_accuracy_compatible and self.config.clip_grad > 0: + return self._get_reproducible_grad_norm() if len(self.chained_optimizers) == 1: return self.chained_optimizers[0].get_grad_norm() if self.grads_states_parallel_group_is_shared(): @@ -1631,6 +1659,8 @@ def has_grad_norm_group(self, grad_norm_group: str) -> bool: def _get_grad_norm_for_group(self, grad_norm_group: str): """Compute gradient norm for a named parameter group.""" _validate_grad_norm_group(grad_norm_group) + if self.config.use_accuracy_compatible and self.config.clip_grad > 0: + return self._get_reproducible_grad_norm(grad_norm_group) if self.grads_states_parallel_group_is_shared(): grouped_grads = [] for optimizer in self.chained_optimizers: diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index a11f5c8cc3b..115916f3325 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -371,6 +371,9 @@ class OptimizerConfig: ################ # Miscellaneous ################ + use_accuracy_compatible: bool = False + """Use a partition-independent FP32 gradient norm when clipping is enabled.""" + clip_grad: float = 1.0 """Gradient clipping based on global L2 norm.""" diff --git a/megatron/core/optimizer/reproducible_norm.py b/megatron/core/optimizer/reproducible_norm.py new file mode 100644 index 00000000000..1bd91386ac9 --- /dev/null +++ b/megatron/core/optimizer/reproducible_norm.py @@ -0,0 +1,128 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Layout- and partition-independent L2 norm of FP32 gradients. + +Squares are computed in FP32 on the gradient device. Their mantissas are +accumulated into base-65536 integer bins, which may be SUM-reduced across +owners before rounding the total once to FP32 and applying the native sqrt. +No gradient values or floating-point norm computation leave the device. + +The 20 limbs cover FP32 squares and the supported 2**40 global element count. +Each uncarried limb is bounded by 2**40 * 65535 < 2**56, so arbitrary parameter +and rank reduction orders cannot overflow int64. The last three bins count +infinities, NaNs, and elements. This deliberately costs more than a fused norm +and is intended only for explicitly enabled accuracy compatibility. +""" + +import torch + + +class ReproducibleL2Norm: + """Accumulate on a single device; reduce bins before calling ``finish``.""" + + def __init__(self, device: torch.device | None = None) -> None: + self.device = ( + device if device is not None else torch.device("cuda", torch.cuda.current_device()) + ) + + def cast(self, value: torch.Tensor, dtype: str) -> torch.Tensor: + return value.to(getattr(torch, dtype)) + + def zeros(self, count: int = 23) -> torch.Tensor: + return torch.zeros(count, dtype=torch.int64, device=self.device) + + def tensor( + self, value: list[int | float] | int | float, dtype: str = "float32" + ) -> torch.Tensor: + return torch.tensor(value, dtype=getattr(torch, dtype), device=self.device) + + def view(self, value: torch.Tensor, dtype: str) -> torch.Tensor: + return value.view(getattr(torch, dtype)) + + def add(self, bins: torch.Tensor, indices: torch.Tensor, values: torch.Tensor) -> torch.Tensor: + return bins.scatter_add_(0, indices, values) + + def accumulate( + self, bins: torch.Tensor, gradient: torch.Tensor, chunk_size: int = 1048576 + ) -> torch.Tensor: + if gradient.layout != torch.strided or gradient.dtype != torch.float32: + raise TypeError("Reproducible clipping requires FP32 gradients") + if chunk_size <= 0: + raise ValueError("chunk_size must be positive") + flat = gradient.reshape([-1]) + if flat.shape[0] > 2**40: + raise OverflowError("Reproducible norm supports at most 2**40 global elements") + bins = self.add(bins, self.tensor([22], "int64"), self.tensor([flat.shape[0]], "int64")) + for start in range(0, flat.shape[0], chunk_size): + value = flat[start : start + chunk_size] + bits = self.cast(self.view(value * value, "int32"), "int64") + exponent = (bits >> 23) & self.tensor(255, "int64") + fraction = bits & self.tensor(0x7FFFFF, "int64") + finite = exponent != 255 + mantissa = fraction | torch.where( + exponent > 0, torch.full_like(exponent, 0x800000), torch.zeros_like(exponent) + ) + shift = torch.maximum(exponent - 1, torch.zeros_like(exponent)) + limb = shift // 16 + shifted = (mantissa << (shift % 16)) * self.cast(finite, "int64") + for offset in range(3): + bins = self.add( + bins, limb + offset, (shifted >> (16 * offset)) & self.tensor(65535, "int64") + ) + flags = torch.stack( + [ + self.cast((exponent == 255) & (fraction == 0), "int64").sum(), + self.cast((exponent == 255) & (fraction != 0), "int64").sum(), + ] + ) + bins = self.add(bins, self.tensor([20, 21], "int64"), flags) + return bins + + def finish(self, bins: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + if int(bins[22].item()) > 2**40: + raise OverflowError("Reproducible norm supports at most 2**40 global elements") + carry = self.zeros(1)[0] + digits = [] + for i in range(20): + value = bins[i] + carry + digits.append(value & self.tensor(65535, "int64")) + carry = value >> 16 + digits = torch.stack(digits) + index = self.tensor(list(range(20)), "int64") + top = torch.where(digits != 0, index, torch.full_like(index, -1)).max() + + def get(i): + return torch.where(index == i, digits, torch.zeros_like(digits)).sum() + + word = get(top) + leading = self.zeros(1)[0] + for width in [8, 4, 2, 1]: + take = word >= (1 << width) + word = torch.where(take, word >> width, word) + leading = leading + self.cast(take, "int64") * width + highest = top * 16 + leading + cut = torch.maximum(highest - 23, self.zeros(1)[0]) + limb, shift = cut // 16, cut % 16 + significand = ( + (get(limb) >> shift) | (get(limb + 1) << (16 - shift)) | (get(limb + 2) << (32 - shift)) + ) & self.tensor(0xFFFFFF, "int64") + round_position = torch.maximum(cut - 1, self.zeros(1)[0]) + round_limb, round_shift = round_position // 16, round_position % 16 + round_word = get(round_limb) + round_bit = ((round_word >> round_shift) & self.tensor(1, "int64")) * self.cast( + cut > 0, "int64" + ) + sticky = torch.where(index < round_limb, digits, torch.zeros_like(digits)).sum() != 0 + sticky = sticky | ((round_word & ((self.tensor(1, "int64") << round_shift) - 1)) != 0) + significand = significand + round_bit * self.cast( + sticky | ((significand & self.tensor(1, "int64")) != 0), "int64" + ) + exponent = highest - 22 + self.cast(significand == 0x1000000, "int64") + raw = (exponent << 23) | (significand & self.tensor(0x7FFFFF, "int64")) + raw = torch.where(exponent >= 255, torch.full_like(raw, 0x7F800000), raw) + raw = torch.where(highest < 23, get(0) | (get(1) << 16), raw) + raw = torch.where(top < 0, torch.zeros_like(raw), raw) + raw = torch.where(bins[20] != 0, torch.full_like(raw, 0x7F800000), raw) + raw = torch.where(bins[21] != 0, torch.full_like(raw, 0x7FC00000), raw) + square_sum = self.view(self.cast(raw.reshape([1]), "int32"), "float32") + return torch.sqrt(square_sum), square_sum diff --git a/tests/unit_tests/optimizer/test_reproducible_norm.py b/tests/unit_tests/optimizer/test_reproducible_norm.py new file mode 100644 index 00000000000..0b672f3355b --- /dev/null +++ b/tests/unit_tests/optimizer/test_reproducible_norm.py @@ -0,0 +1,109 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +from types import SimpleNamespace + +import pytest +import torch + +from megatron.core.optimizer.clip_grads import clip_grad_by_total_norm_fp32 +from megatron.core.optimizer.optimizer import ChainedOptimizer +from megatron.core.optimizer.optimizer_config import OptimizerConfig +from megatron.core.optimizer.reproducible_norm import ReproducibleL2Norm + + +@pytest.mark.parametrize( + 'values,expected', + [ + ([0.0], 0.0), + ([3.0, 4.0], 25.0), + ([4096.0, 1.0], 16777216.0), + ([4096.0, 1.0, 1.0, 1.0], 16777220.0), + ([2.0**-70], 2.0**-140), + ([float('inf')], float('inf')), + ([3e38], float('inf')), + ], +) +def test_sum_squares_rounding(values, expected): + norm = ReproducibleL2Norm() + _, squared = norm.finish(norm.accumulate(norm.zeros(), norm.tensor(values))) + assert squared.item() == expected + + +def test_nan_and_invalid_input(): + norm = ReproducibleL2Norm() + actual, _ = norm.finish( + norm.accumulate(norm.zeros(), norm.tensor([float('inf'), float('nan')])) + ) + assert torch.isnan(actual).item() + with pytest.raises(TypeError, match='FP32'): + norm.accumulate(norm.zeros(), norm.tensor([1.0]).bfloat16()) + bins = norm.zeros() + bins[22] = 2**40 + 1 + with pytest.raises(OverflowError): + norm.finish(bins) + + +def test_layout_and_chunk_invariance(): + norm = ReproducibleL2Norm() + gradient = torch.arange(1, 12289, device='cuda', dtype=torch.float32).reshape(96, 128) / 16384 + whole = norm.accumulate(norm.zeros(), gradient) + transposed = norm.accumulate(norm.zeros(), gradient.T.contiguous(), chunk_size=123) + split = norm.accumulate(norm.zeros(), gradient.flatten()[:1234]) + split += norm.accumulate(norm.zeros(), gradient.flatten()[1234:]) + assert torch.equal(whole, transposed) + assert torch.equal(whole, split) + + +def test_chained_owner_groups_and_clipping(monkeypatch): + rank = torch.distributed.get_rank() + world = torch.distributed.get_world_size() + singleton = None + for member in range(world): + group = torch.distributed.new_group([member]) + if member == rank: + singleton = group + config = OptimizerConfig(use_accuracy_compatible=True, clip_grad=1.0) + # Dense 3 and 4 have distinct owners; the expert 12 is replicated between + # singleton stats groups. Finishing each child separately loses this contract. + dense = [torch.tensor([3.0 if rank == 0 else 4.0], device='cuda')] if rank < 2 else [] + if world == 1: + dense = [torch.tensor([3.0, 4.0], device='cuda')] + expert = [torch.tensor([12.0], device='cuda')] + children = [ + SimpleNamespace( + config=config, + get_grads_for_grad_norm=lambda _group=None: dense, + get_grad_stats_parallel_group=lambda: torch.distributed.group.WORLD, + ), + SimpleNamespace( + config=config, + get_grads_for_grad_norm=lambda _group=None: expert, + get_grad_stats_parallel_group=lambda: singleton, + ), + ] + actual = ChainedOptimizer(children).get_grad_norm() + assert actual.item() == 13.0 + parameter = torch.nn.Parameter(torch.zeros(2, device='cuda')) + parameter.grad = torch.tensor([3.0, 4.0], device='cuda') + monkeypatch.setattr('megatron.core.optimizer.clip_grads.multi_tensor_scale_tensor_impl', None) + clip_grad_by_total_norm_fp32([parameter], 1.0, actual) + expected = torch.tensor([3.0, 4.0], device='cuda') * (1.0 / (actual + 1e-6)) + assert torch.equal(parameter.grad, expected) + + +@pytest.fixture(scope='module', autouse=True) +def distributed_norm_device(): + import os + + torch.cuda.set_device(int(os.environ.get('LOCAL_RANK', '0'))) + owns_group = not torch.distributed.is_initialized() + if owns_group: + torch.distributed.init_process_group(backend='nccl') + yield + if owns_group: + torch.distributed.destroy_process_group() + + +@pytest.fixture(scope='session') +def ensure_test_data(): + """The norm tests are self-contained and do not consume external datasets.""" From 11e8be093f79080292c6e0e0073b93782bc3fcbc Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Mon, 14 Sep 2026 16:34:33 +0800 Subject: [PATCH 26/27] refactor(glm52): scope reference numerics and trim unrelated changes --- .../core/distributed/finalize_model_grads.py | 41 +- megatron/core/models/backends.py | 3 +- .../common/language_module/language_module.py | 17 +- ...rimental_attention_variant_module_specs.py | 61 +- megatron/core/optimizer/clip_grads.py | 45 +- megatron/core/optimizer/distrib_optimizer.py | 10 +- megatron/core/optimizer/optimizer.py | 38 +- megatron/core/optimizer/optimizer_config.py | 3 - megatron/core/optimizer/reproducible_norm.py | 128 ---- megatron/core/pipeline_parallel/schedules.py | 36 +- megatron/core/tensor_parallel/layers.py | 64 +- megatron/core/tensor_parallel/mappings.py | 4 - .../experimental_attention_variant/dsa.py | 4 +- megatron/core/transformer/moe/experts.py | 16 +- megatron/core/transformer/moe/moe_layer.py | 14 +- megatron/core/transformer/moe/moe_utils.py | 126 +++- megatron/core/transformer/moe/router.py | 15 +- .../core/transformer/moe/token_dispatcher.py | 8 +- .../transformer/multi_token_prediction.py | 324 +++------ megatron/core/transformer/torch_norm.py | 34 +- .../core/transformer/transformer_block.py | 10 +- .../core/transformer/transformer_config.py | 614 +++++++----------- .../core/transformer/transformer_layer.py | 4 +- .../distributed/test_finalize_model_grads.py | 55 +- ...rimental_attention_variant_module_specs.py | 107 +-- .../optimizer/test_reproducible_norm.py | 109 ---- .../test_accuracy_tp1_migration.py | 77 +-- .../unit_tests/tensor_parallel/test_layers.py | 28 +- .../test_attention_variant_dsa.py | 3 +- .../moe/test_accuracy_migration.py | 59 +- .../transformer/moe/test_routers.py | 3 +- .../moe/test_sequential_expert_padding.py | 1 + .../test_multi_token_prediction.py | 459 +++++-------- .../unit_tests/transformer/test_torch_norm.py | 15 + 34 files changed, 917 insertions(+), 1618 deletions(-) delete mode 100644 megatron/core/optimizer/reproducible_norm.py delete mode 100644 tests/unit_tests/optimizer/test_reproducible_norm.py diff --git a/megatron/core/distributed/finalize_model_grads.py b/megatron/core/distributed/finalize_model_grads.py index 9f7134ca3e5..d5b1e77066f 100644 --- a/megatron/core/distributed/finalize_model_grads.py +++ b/megatron/core/distributed/finalize_model_grads.py @@ -460,14 +460,18 @@ def finalize_model_grads( config = get_model_config(model[0]) - # E-172: keep mcore 1/num_tokens scaling under use_accuracy_compatible. - # calculate_per_token_loss=true backprops raw sum(loss*mask); skipping this - # left every gradient num_tokens times too large (44.0x at step 1). Formal - # N=5 with the skip matched step-1 IEEE but drifted at steps 2-5. Do not - # take main's GLM-4.5 skip (loss_normalized_in_graph=True / num_tokens=None). + # [对齐修复] use_accuracy_compatible=1: PaddleFleet 的 fixed-loss 路径已在 autograd 图内除过 + # 本地有效 token 数, MCore 这里再用全局 num_tokens 缩放会引入 ~global_token / local_token + # 倍的额外因子 (实测 ~74.64x)。在对齐模式下跳过 num_tokens 全局缩放, 改为 grad sync 后做 + # 1/dp_size 平均, 与 Paddle DP 平均语义对齐; 同时对 RouterGatingLinearFunction 记录的 + # fp32 gate wgrad 做一次 DP all-reduce, 与参考实现一致。 from ..transformer.module import _use_accuracy_compatible - loss_normalized_in_graph = False + loss_normalized_in_graph = ( + _use_accuracy_compatible() and not config.dsa_accuracy_compatible and num_tokens is not None + ) + if loss_normalized_in_graph: + num_tokens = None tp_dp_cp_group = None if pg_collection is not None: @@ -553,18 +557,25 @@ def finalize_model_grads( if not _use_accuracy_compatible(): _update_router_expert_bias(model, config) - if pg_collection is None: - tp_dp_cp_group = parallel_state.get_tensor_and_data_parallel_group( - with_context_parallel=True - ) - _update_router_expert_bias(model, config, tp_dp_cp_group=tp_dp_cp_group) + if pg_collection is None: + tp_dp_cp_group = parallel_state.get_tensor_and_data_parallel_group( + with_context_parallel=True + ) + _update_router_expert_bias(model, config, tp_dp_cp_group=tp_dp_cp_group) reset_model_temporary_tensors(config, model) - # All-reduce the fp32 gate wgrad that RouterGatingLinearFunction records on - # ``param._run_torch_gate_fp32_wgrad`` across DP, matching the reference - # implementation. A no-op when that mechanism is not enabled. - if _use_accuracy_compatible(): + # [对齐修复] 对应上方 loss_normalized_in_graph 早跳过分支: 跳过全局 num_tokens 缩放后, + # 这里改为 grad sync 后做 1/dp_size 平均, 与 PaddleFleet DP 平均语义对齐; + # 并对 RouterGatingLinearFunction 记录到 param._run_torch_gate_fp32_wgrad 的 fp32 gate + # wgrad 做一次 DP all-reduce (若该机制未启用则 getattr 为 None, 代码为 no-op)。 + if loss_normalized_in_graph: + dp_size = parallel_state.get_data_parallel_world_size(with_context_parallel=True) + if dp_size > 1: + for model_chunk in model: + model_chunk.scale_gradients(1.0 / dp_size) + + if loss_normalized_in_graph or config.dsa_accuracy_compatible: for model_chunk in model: for param in model_chunk.parameters(): gate_wgrad = getattr(param, "_run_torch_gate_fp32_wgrad", None) diff --git a/megatron/core/models/backends.py b/megatron/core/models/backends.py index c543a49e266..e71db578f14 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -10,7 +10,6 @@ TEColumnParallelGroupedLinear, TERowParallelGroupedLinear, ) -from megatron.core.post_training.modelopt.layers import Linear from megatron.core.tensor_parallel.layers import ColumnParallelLinear, RowParallelLinear from megatron.core.transformer.dot_product_attention import DotProductAttention from megatron.core.transformer.mlp import MLPSubmodules, TEActivationFunctionBuilder @@ -107,6 +106,8 @@ def linear(self) -> type: still returns TELinear; this method is the TE-off counterpart so a LocalSpecProvider DSA spec does not re-enter Transformer Engine. """ + from megatron.core.post_training.modelopt.layers import Linear + return Linear def column_parallel_linear(self) -> type: diff --git a/megatron/core/models/common/language_module/language_module.py b/megatron/core/models/common/language_module/language_module.py index b77f3219f55..615e2c45196 100644 --- a/megatron/core/models/common/language_module/language_module.py +++ b/megatron/core/models/common/language_module/language_module.py @@ -206,9 +206,19 @@ def compute_language_model_loss(self, labels: Tensor, logits: Tensor) -> Tensor: elif self.config.cross_entropy_fusion_impl == 'native': loss = fused_vocab_parallel_cross_entropy(logits, labels, self.pg_collection.tp) else: - loss = tensor_parallel.vocab_parallel_cross_entropy( - logits, labels, tp_group=self.tp_group - ) + if _use_accuracy_compatible() and not self.config.dsa_accuracy_compatible: + s, b = labels.shape + loss = torch.nn.functional.cross_entropy( + logits.float().reshape(s * b, -1), # [s*b, vocab] + labels.reshape(s * b), # [s*b] + reduction='none', + ).reshape( + s, b + ) # [s, b] + else: + loss = tensor_parallel.vocab_parallel_cross_entropy( + logits, labels, tp_group=self.tp_group + ) # [s b] => [b, s] loss = loss.transpose(0, 1).contiguous() @@ -227,6 +237,7 @@ def compute_language_model_loss(self, labels: Tensor, logits: Tensor) -> Tensor: f"md5={_hashlib.md5(_l.cpu().numpy().tobytes()).hexdigest()}", flush=True, ) + return loss def setup_embeddings_and_output_layer(self) -> None: diff --git a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py index a8cc4886717..4b01ba561b6 100644 --- a/megatron/core/models/gpt/experimental_attention_variant_module_specs.py +++ b/megatron/core/models/gpt/experimental_attention_variant_module_specs.py @@ -58,9 +58,7 @@ ########## -def _get_standalone_norm( - config: TransformerConfig, backend: BackendSpecProvider, *, for_qk=False -): +def _get_standalone_norm(config: TransformerConfig, backend: BackendSpecProvider, *, for_qk=False): rms_norm = config.normalization == "RMSNorm" if rms_norm and config.norm_accuracy_compatible: return WrappedTorchNorm @@ -91,9 +89,7 @@ def get_dsa_module_spec_for_backend( config: TransformerConfig, backend: BackendSpecProvider = None ) -> ModuleSpec: """Helper function to get module spec for Sparse Attention.""" - assert config.multi_latent_attention, ( - "Currently only MLA supports sparse attention." - ) + assert config.multi_latent_attention, "Currently only MLA supports sparse attention." assert config.qk_l2_norm is False, "qk_l2_norm is not supported with MLA." # Because TransformerEngine does not support sparse attention yet, we use local @@ -116,9 +112,7 @@ def get_dsa_module_spec_for_backend( # DSA indexer requires normalized q as input, so here we cannot fuse qk layernorm # with linear projection and have to use unfused qk layernorm. qk_norm = ( - _get_standalone_norm(config, backend, for_qk=True) - if config.qk_layernorm - else IdentityOp + _get_standalone_norm(config, backend, for_qk=True) if config.qk_layernorm else IdentityOp ) attention = ModuleSpec( @@ -214,9 +208,7 @@ def get_transformer_layer_with_experimental_attention_variant_spec( experimental_attention_spec = None if 0 in experimental_attention_pattern: - standard_attention_spec = _get_self_attention_module_spec( - config=config, backend=backend - ) + standard_attention_spec = _get_self_attention_module_spec(config=config, backend=backend) else: standard_attention_spec = None @@ -248,11 +240,7 @@ def get_transformer_layer_with_experimental_attention_variant_spec( if experimental_attention_pattern[layer_number] == 1 else standard_attention_spec ) - mlp = ( - moe_layer_spec - if moe_layer_pattern[layer_number] == 1 - else dense_mlp_layer_spec - ) + mlp = moe_layer_spec if moe_layer_pattern[layer_number] == 1 else dense_mlp_layer_spec fuse_pre_mlp_layernorm = ( fuse_layernorm_pre_moe if moe_layer_pattern[layer_number] == 1 @@ -264,9 +252,7 @@ def get_transformer_layer_with_experimental_attention_variant_spec( else _get_standalone_norm(config, backend) ) pre_mlp_layernorm = ( - IdentityOp - if fuse_pre_mlp_layernorm - else _get_standalone_norm(config, backend) + IdentityOp if fuse_pre_mlp_layernorm else _get_standalone_norm(config, backend) ) layer_specs.append( @@ -287,9 +273,7 @@ def get_transformer_layer_with_experimental_attention_variant_spec( def get_transformer_block_with_experimental_attention_variant_spec( - config: TransformerConfig, - vp_stage: Optional[int] = None, - pp_rank: Optional[int] = None, + config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None ) -> TransformerBlockSubmodules: """Build transformer block spec with experimental attention variants (e.g., linear attention). @@ -327,12 +311,8 @@ def get_transformer_block_with_experimental_attention_variant_spec( layer_type=LayerType.decoder, vp_stage=vp_stage, pp_rank=pp_rank ) else: - offset = get_transformer_layer_offset( - config, vp_stage=vp_stage, pp_rank=pp_rank - ) - num_layers_to_build = get_num_layers_to_build( - config, vp_stage=vp_stage, pp_rank=pp_rank - ) + offset = get_transformer_layer_offset(config, vp_stage=vp_stage, pp_rank=pp_rank) + num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage, pp_rank=pp_rank) local_layer_ids = range(offset, offset + num_layers_to_build) _validate_dsa_index_share_pipeline_split(config, local_layer_ids) @@ -357,9 +337,7 @@ def is_linear_attention_variant(experimental_attention_variant: Optional[str]) - return experimental_attention_variant in linear_attention_variants -def _validate_dsa_index_share_pipeline_split( - config: TransformerConfig, local_layer_ids -) -> None: +def _validate_dsa_index_share_pipeline_split(config: TransformerConfig, local_layer_ids) -> None: """Ensure DSA top-k sharing does not require top-k indices from another PP stage.""" if ( config.experimental_attention_variant != "dsa" @@ -374,16 +352,12 @@ def _validate_dsa_index_share_pipeline_split( for position, layer_id in enumerate(local_layer_ids): layer_number = layer_id + 1 if not is_dsa_skip_topk_layer( - layer_number, - config.dsa_indexer_skip_topk_offset, - config.dsa_indexer_topk_freq, + layer_number, config.dsa_indexer_skip_topk_offset, config.dsa_indexer_topk_freq ): continue source_layer_number = source_dsa_compute_layer( - layer_number, - config.dsa_indexer_skip_topk_offset, - config.dsa_indexer_topk_freq, + layer_number, config.dsa_indexer_skip_topk_offset, config.dsa_indexer_topk_freq ) source_layer_id = source_layer_number - 1 if ( @@ -410,8 +384,7 @@ def get_moe_layer_pattern(config: TransformerConfig) -> List[int]: if isinstance(config.moe_layer_freq, int): # [1,0,0,...,0,1,0,0,...,0,...] moe_layer_pattern = [ - 1 if (i % config.moe_layer_freq == 0) else 0 - for i in range(config.num_layers) + 1 if (i % config.moe_layer_freq == 0) else 0 for i in range(config.num_layers) ] elif isinstance(config.moe_layer_freq, list): moe_layer_pattern = config.moe_layer_freq @@ -499,9 +472,7 @@ def _get_self_attention_module_spec( if backend is None: backend = _get_backend_spec_provider(config=config) - from megatron.core.models.gpt.gpt_layer_specs import ( - get_gpt_layer_with_transformer_engine_spec, - ) + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec layer_spec = get_gpt_layer_with_transformer_engine_spec( num_experts=config.num_moe_experts, @@ -562,9 +533,7 @@ def _get_moe_module_spec( if backend is None: backend = _get_backend_spec_provider(config=config) - from megatron.core.models.gpt.moe_module_specs import ( - get_moe_module_spec_for_backend, - ) + from megatron.core.models.gpt.moe_module_specs import get_moe_module_spec_for_backend return ( get_moe_module_spec_for_backend( diff --git a/megatron/core/optimizer/clip_grads.py b/megatron/core/optimizer/clip_grads.py index 6345991e057..3c5491d39a1 100644 --- a/megatron/core/optimizer/clip_grads.py +++ b/megatron/core/optimizer/clip_grads.py @@ -50,37 +50,12 @@ from ..tensor_parallel import param_is_not_tensor_parallel_duplicate from ..transformer.module import param_is_not_shared from ..utils import get_data_parallel_group_if_dtensor, to_local_if_dtensor -from .reproducible_norm import ReproducibleL2Norm - - -@torch.no_grad() -def get_reproducible_grad_norm_bins( - grads_for_norm: List[torch.Tensor], - grad_stats_parallel_group: torch.distributed.ProcessGroup | None, -) -> torch.Tensor: - """Reduce exact FP32-square bins using the existing gradient ownership groups.""" - grads_for_norm = list(grads_for_norm) - data_parallel_group = None - for grad in grads_for_norm: - data_parallel_group = get_data_parallel_group_if_dtensor(grad, data_parallel_group) - grads_for_norm = [to_local_if_dtensor(grad) for grad in grads_for_norm] - accumulator = ReproducibleL2Norm(grads_for_norm[0].device if grads_for_norm else None) - bins = accumulator.zeros() - for grad in grads_for_norm: - if grad.layout != torch.strided: - raise TypeError("Reproducible clipping requires dense FP32 gradients") - bins = accumulator.accumulate(bins, grad) - if data_parallel_group: - torch.distributed.all_reduce(bins, group=data_parallel_group) - torch.distributed.all_reduce(bins, group=grad_stats_parallel_group) - return bins def get_grad_norm_fp32( grads_for_norm: Union[List[torch.Tensor], torch.Tensor], norm_type: Union[int, float] = 2, grad_stats_parallel_group: Optional[torch.distributed.ProcessGroup] = None, - use_accuracy_compatible: bool = False, ) -> float: """Calculate the p-norm of gradients in FP32 precision. @@ -105,12 +80,6 @@ def get_grad_norm_fp32( if isinstance(grads_for_norm, torch.Tensor): grads_for_norm = [grads_for_norm] - if use_accuracy_compatible: - if float(norm_type) != 2.0: - raise ValueError("Reproducible clipping supports only the L2 norm") - bins = get_reproducible_grad_norm_bins(grads_for_norm, grad_stats_parallel_group) - return ReproducibleL2Norm(bins.device).finish(bins)[0] - data_parallel_group = None for grad in grads_for_norm: data_parallel_group = get_data_parallel_group_if_dtensor(grad, data_parallel_group) @@ -215,14 +184,12 @@ def clip_grad_by_total_norm_fp32( dummy_overflow_buf = torch.zeros(1, dtype=torch.int, device='cuda') if isinstance(clip_coeff, torch.Tensor): clip_coeff.clamp_max_(1.0) - if multi_tensor_scale_tensor_impl is not None: - multi_tensor_applier( - multi_tensor_scale_tensor_impl, dummy_overflow_buf, [grads, grads], clip_coeff - ) - else: - multi_tensor_applier( - multi_tensor_scale_impl, dummy_overflow_buf, [grads, grads], clip_coeff.item() - ) + assert ( + multi_tensor_scale_tensor_impl is not None + ), "clip_coeff is tensor type. But multi_tensor_scale_tensor not available." + multi_tensor_applier( + multi_tensor_scale_tensor_impl, dummy_overflow_buf, [grads, grads], clip_coeff + ) elif clip_coeff < 1.0: multi_tensor_applier( multi_tensor_scale_impl, dummy_overflow_buf, [grads, grads], clip_coeff diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 9e030a6b17f..258b1b8b3bb 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -663,10 +663,12 @@ def __init__( assert self.ddp_config == model_chunk.ddp_config self.distributed_optimizer_instance_id = distributed_optimizer_instance_id - assert ( - isinstance(optimizer, (Adam, torch.optim.AdamW, HybridDeviceOptimizer)) - or optimizer is None - ), ( + from megatron.core.transformer.module import _use_accuracy_compatible + + _allowed_optim_types = (Adam, HybridDeviceOptimizer) + if _use_accuracy_compatible(): + _allowed_optim_types = (Adam, torch.optim.AdamW, HybridDeviceOptimizer) + assert isinstance(optimizer, _allowed_optim_types) or optimizer is None, ( "Only Adam and HybridDeviceOptimizer currently supported, " "due to checkpointing requirements." ) diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index e00c17b200f..4a74328d0d9 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -48,15 +48,9 @@ from ..dist_checkpointing.utils import add_prefix_for_sharding from ..transformer.module import param_is_not_shared from ..utils import log_single_rank -from .clip_grads import ( - clip_grad_by_total_norm_fp32, - count_zeros_fp32, - get_grad_norm_fp32, - get_reproducible_grad_norm_bins, -) +from .clip_grads import clip_grad_by_total_norm_fp32, count_zeros_fp32, get_grad_norm_fp32 from .grad_scaler import MegatronGradScaler from .optimizer_config import OptimizerConfig -from .reproducible_norm import ReproducibleL2Norm logger = getLogger(__name__) @@ -302,10 +296,7 @@ def get_grad_norm(self): """Compute and return grad norm.""" grads_for_norm = self.get_grads_for_grad_norm() total_norm = get_grad_norm_fp32( - grads_for_norm, - grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - use_accuracy_compatible=self.config.use_accuracy_compatible - and self.config.clip_grad > 0, + grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group() ) return total_norm @@ -317,10 +308,7 @@ def _compute_grad_norms_by_group(self) -> Dict[str, float]: if self.has_grad_norm_group(grad_norm_group): grouped_grads = self.get_grads_for_grad_norm(grad_norm_group) group_grad_norm = get_grad_norm_fp32( - grouped_grads, - grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - use_accuracy_compatible=self.config.use_accuracy_compatible - and self.config.clip_grad > 0, + grouped_grads, grad_stats_parallel_group=self.get_grad_stats_parallel_group() ) self.grad_norms_by_group[grad_norm_group] = group_grad_norm return self.grad_norms_by_group @@ -338,10 +326,7 @@ def clip_grad_norm(self, clip_grad: float) -> float: else: grads_for_norm = [] grad_norm = get_grad_norm_fp32( - grads_for_norm, - grad_stats_parallel_group=self.get_grad_stats_parallel_group(), - use_accuracy_compatible=self.config.use_accuracy_compatible - and self.config.clip_grad > 0, + grads_for_norm, grad_stats_parallel_group=self.get_grad_stats_parallel_group() ) if clip_grad > 0.0 and params: @@ -1581,21 +1566,8 @@ def get_grad_stats_parallel_group(self) -> torch.distributed.ProcessGroup: ) return self.chained_optimizers[0].get_grad_stats_parallel_group() - @torch.no_grad() - def _get_reproducible_grad_norm(self, grad_norm_group=None): - bins = None - for optimizer in self.chained_optimizers: - part = get_reproducible_grad_norm_bins( - optimizer.get_grads_for_grad_norm(grad_norm_group), - optimizer.get_grad_stats_parallel_group(), - ) - bins = part if bins is None else bins + part - return ReproducibleL2Norm(bins.device).finish(bins)[0] - @torch.no_grad() def get_grad_norm(self): - if self.config.use_accuracy_compatible and self.config.clip_grad > 0: - return self._get_reproducible_grad_norm() if len(self.chained_optimizers) == 1: return self.chained_optimizers[0].get_grad_norm() if self.grads_states_parallel_group_is_shared(): @@ -1659,8 +1631,6 @@ def has_grad_norm_group(self, grad_norm_group: str) -> bool: def _get_grad_norm_for_group(self, grad_norm_group: str): """Compute gradient norm for a named parameter group.""" _validate_grad_norm_group(grad_norm_group) - if self.config.use_accuracy_compatible and self.config.clip_grad > 0: - return self._get_reproducible_grad_norm(grad_norm_group) if self.grads_states_parallel_group_is_shared(): grouped_grads = [] for optimizer in self.chained_optimizers: diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 115916f3325..a11f5c8cc3b 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -371,9 +371,6 @@ class OptimizerConfig: ################ # Miscellaneous ################ - use_accuracy_compatible: bool = False - """Use a partition-independent FP32 gradient norm when clipping is enabled.""" - clip_grad: float = 1.0 """Gradient clipping based on global L2 norm.""" diff --git a/megatron/core/optimizer/reproducible_norm.py b/megatron/core/optimizer/reproducible_norm.py deleted file mode 100644 index 1bd91386ac9..00000000000 --- a/megatron/core/optimizer/reproducible_norm.py +++ /dev/null @@ -1,128 +0,0 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - -"""Layout- and partition-independent L2 norm of FP32 gradients. - -Squares are computed in FP32 on the gradient device. Their mantissas are -accumulated into base-65536 integer bins, which may be SUM-reduced across -owners before rounding the total once to FP32 and applying the native sqrt. -No gradient values or floating-point norm computation leave the device. - -The 20 limbs cover FP32 squares and the supported 2**40 global element count. -Each uncarried limb is bounded by 2**40 * 65535 < 2**56, so arbitrary parameter -and rank reduction orders cannot overflow int64. The last three bins count -infinities, NaNs, and elements. This deliberately costs more than a fused norm -and is intended only for explicitly enabled accuracy compatibility. -""" - -import torch - - -class ReproducibleL2Norm: - """Accumulate on a single device; reduce bins before calling ``finish``.""" - - def __init__(self, device: torch.device | None = None) -> None: - self.device = ( - device if device is not None else torch.device("cuda", torch.cuda.current_device()) - ) - - def cast(self, value: torch.Tensor, dtype: str) -> torch.Tensor: - return value.to(getattr(torch, dtype)) - - def zeros(self, count: int = 23) -> torch.Tensor: - return torch.zeros(count, dtype=torch.int64, device=self.device) - - def tensor( - self, value: list[int | float] | int | float, dtype: str = "float32" - ) -> torch.Tensor: - return torch.tensor(value, dtype=getattr(torch, dtype), device=self.device) - - def view(self, value: torch.Tensor, dtype: str) -> torch.Tensor: - return value.view(getattr(torch, dtype)) - - def add(self, bins: torch.Tensor, indices: torch.Tensor, values: torch.Tensor) -> torch.Tensor: - return bins.scatter_add_(0, indices, values) - - def accumulate( - self, bins: torch.Tensor, gradient: torch.Tensor, chunk_size: int = 1048576 - ) -> torch.Tensor: - if gradient.layout != torch.strided or gradient.dtype != torch.float32: - raise TypeError("Reproducible clipping requires FP32 gradients") - if chunk_size <= 0: - raise ValueError("chunk_size must be positive") - flat = gradient.reshape([-1]) - if flat.shape[0] > 2**40: - raise OverflowError("Reproducible norm supports at most 2**40 global elements") - bins = self.add(bins, self.tensor([22], "int64"), self.tensor([flat.shape[0]], "int64")) - for start in range(0, flat.shape[0], chunk_size): - value = flat[start : start + chunk_size] - bits = self.cast(self.view(value * value, "int32"), "int64") - exponent = (bits >> 23) & self.tensor(255, "int64") - fraction = bits & self.tensor(0x7FFFFF, "int64") - finite = exponent != 255 - mantissa = fraction | torch.where( - exponent > 0, torch.full_like(exponent, 0x800000), torch.zeros_like(exponent) - ) - shift = torch.maximum(exponent - 1, torch.zeros_like(exponent)) - limb = shift // 16 - shifted = (mantissa << (shift % 16)) * self.cast(finite, "int64") - for offset in range(3): - bins = self.add( - bins, limb + offset, (shifted >> (16 * offset)) & self.tensor(65535, "int64") - ) - flags = torch.stack( - [ - self.cast((exponent == 255) & (fraction == 0), "int64").sum(), - self.cast((exponent == 255) & (fraction != 0), "int64").sum(), - ] - ) - bins = self.add(bins, self.tensor([20, 21], "int64"), flags) - return bins - - def finish(self, bins: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - if int(bins[22].item()) > 2**40: - raise OverflowError("Reproducible norm supports at most 2**40 global elements") - carry = self.zeros(1)[0] - digits = [] - for i in range(20): - value = bins[i] + carry - digits.append(value & self.tensor(65535, "int64")) - carry = value >> 16 - digits = torch.stack(digits) - index = self.tensor(list(range(20)), "int64") - top = torch.where(digits != 0, index, torch.full_like(index, -1)).max() - - def get(i): - return torch.where(index == i, digits, torch.zeros_like(digits)).sum() - - word = get(top) - leading = self.zeros(1)[0] - for width in [8, 4, 2, 1]: - take = word >= (1 << width) - word = torch.where(take, word >> width, word) - leading = leading + self.cast(take, "int64") * width - highest = top * 16 + leading - cut = torch.maximum(highest - 23, self.zeros(1)[0]) - limb, shift = cut // 16, cut % 16 - significand = ( - (get(limb) >> shift) | (get(limb + 1) << (16 - shift)) | (get(limb + 2) << (32 - shift)) - ) & self.tensor(0xFFFFFF, "int64") - round_position = torch.maximum(cut - 1, self.zeros(1)[0]) - round_limb, round_shift = round_position // 16, round_position % 16 - round_word = get(round_limb) - round_bit = ((round_word >> round_shift) & self.tensor(1, "int64")) * self.cast( - cut > 0, "int64" - ) - sticky = torch.where(index < round_limb, digits, torch.zeros_like(digits)).sum() != 0 - sticky = sticky | ((round_word & ((self.tensor(1, "int64") << round_shift) - 1)) != 0) - significand = significand + round_bit * self.cast( - sticky | ((significand & self.tensor(1, "int64")) != 0), "int64" - ) - exponent = highest - 22 + self.cast(significand == 0x1000000, "int64") - raw = (exponent << 23) | (significand & self.tensor(0x7FFFFF, "int64")) - raw = torch.where(exponent >= 255, torch.full_like(raw, 0x7F800000), raw) - raw = torch.where(highest < 23, get(0) | (get(1) << 16), raw) - raw = torch.where(top < 0, torch.zeros_like(raw), raw) - raw = torch.where(bins[20] != 0, torch.full_like(raw, 0x7F800000), raw) - raw = torch.where(bins[21] != 0, torch.full_like(raw, 0x7FC00000), raw) - square_sum = self.view(self.cast(raw.reshape([1]), "int32"), "float32") - return torch.sqrt(square_sum), square_sum diff --git a/megatron/core/pipeline_parallel/schedules.py b/megatron/core/pipeline_parallel/schedules.py index b75fe63ab8e..f576bc05a58 100644 --- a/megatron/core/pipeline_parallel/schedules.py +++ b/megatron/core/pipeline_parallel/schedules.py @@ -24,7 +24,6 @@ ProcessGroupCollection, ) from megatron.core.transformer.cuda_graphs import create_cudagraphs, set_current_microbatch -from megatron.core.transformer.module import _use_accuracy_compatible from megatron.core.transformer.moe.paged_stash import paged_stash_reset from megatron.core.transformer.moe.router import MoEAuxLossAutoScaler from megatron.core.utils import ( @@ -164,7 +163,7 @@ def forward_step(data_iterator, model): return forward_backward_func -def deallocate_output_tensor(out, deallocate_pipeline_outputs=False): +def deallocate_output_tensor(out, deallocate_pipeline_outputs=False, config=None): '''Pseudo-deallocate (i.e., set to scalar) the output tensor's '.data' field. This method should be called right after the output tensor has been @@ -178,22 +177,23 @@ def deallocate_output_tensor(out, deallocate_pipeline_outputs=False): ''' if (out is None) or (not deallocate_pipeline_outputs): return - if _use_accuracy_compatible(): - # Compatibility fallback: callers supply only a tensor and deallocation flag. - _tp_size = int(parallel_state.get_tensor_model_parallel_world_size() or 1) - if _tp_size <= 1: - return + if ( + config is not None + and config.dsa_accuracy_compatible + and config.tensor_model_parallel_size <= 1 + ): + return # Handle dict format (multi-module pipelines) if isinstance(out, dict): for value in out.values(): - deallocate_output_tensor(value, deallocate_pipeline_outputs) + deallocate_output_tensor(value, deallocate_pipeline_outputs, config) return # Handle list format if isinstance(out, list): for item in out: - deallocate_output_tensor(item, deallocate_pipeline_outputs) + deallocate_output_tensor(item, deallocate_pipeline_outputs, config) return # Base case: deallocate tensor @@ -576,7 +576,7 @@ def backward_step(input_tensor, output_tensor, output_tensor_grad, config): if output_tensor[0].requires_grad: _tp_size = int(getattr(config, "tensor_model_parallel_size", 1) or 1) if config.deallocate_pipeline_outputs and ( - not _use_accuracy_compatible() or _tp_size > 1 + not config.dsa_accuracy_compatible or _tp_size > 1 ): custom_backward(output_tensor[0], output_tensor_grad[0]) else: @@ -652,7 +652,7 @@ def _unwrap_single_tensor_list(tensor): if output_tensor_module is not None and output_tensor_module.requires_grad: _tp_size = int(getattr(config, "tensor_model_parallel_size", 1) or 1) if config.deallocate_pipeline_outputs and ( - not _use_accuracy_compatible() or _tp_size > 1 + not config.dsa_accuracy_compatible or _tp_size > 1 ): custom_backward(output_tensor_module, output_tensor_grad_module) else: @@ -1649,7 +1649,7 @@ def forward_backward_helper_wrapper( ) if recv_prev: input_tensors[next_forward_model_chunk_id].append(input_tensor) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) else: if not is_pp_first_stage(p2p_communicator.pp_group): # Send only since recv prefetched. @@ -1679,7 +1679,7 @@ def forward_backward_helper_wrapper( send_next_wait_handle.wait() send_next_wait_handle = None - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) if recv_prev: input_tensors[next_forward_model_chunk_id].append( fwd_recv_buffer[k % fwd_recv_buffer_size] @@ -1757,7 +1757,7 @@ def pp_pre_forward(vp_stage=None): recv_prev_wait_handle = recv_prev_wait_handles.pop(0) recv_prev_wait_handle.wait() - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Async forward send / receive def pp_post_forward(output_tensor, vp_stage=None): @@ -1932,7 +1932,7 @@ def pp_post_backward(input_tensor_grad, vp_stage=None): tensor_shape=tensor_shape, ) ) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Put input_tensor and output_tensor_grad in data structures in the # right location. if recv_prev: @@ -1940,7 +1940,7 @@ def pp_post_backward(input_tensor_grad, vp_stage=None): if recv_next: output_tensor_grads[next_backward_model_chunk_id].append(output_tensor_grad) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) nvtx_range_pop(suffix="steady") # Run cooldown backward passes (flush out pipeline) for the last model chunk. @@ -2365,7 +2365,7 @@ def enable_grad_sync(): if not forward_only: input_tensors.append(input_tensor) output_tensors.append(output_tensor) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Before running 1F1B, need to receive first forward tensor. # If all microbatches are run in warmup / cooldown phase, then no need to @@ -2420,7 +2420,7 @@ def enable_grad_sync(): # Add input_tensor and output_tensor to end of list. input_tensors.append(input_tensor) output_tensors.append(output_tensor) - deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs) + deallocate_output_tensor(output_tensor, config.deallocate_pipeline_outputs, config) # Pop input_tensor and output_tensor from the start of the list for # the backward pass. diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 6b6d8aaf22d..b613c581ed2 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -282,7 +282,11 @@ def __init__( ) ) self.num_embeddings_per_partition = self.vocab_end_index - self.vocab_start_index - self.deterministic_mode = config.deterministic_mode or config.use_accuracy_compatible + self.deterministic_mode = ( + config.deterministic_mode + or _use_accuracy_compatible() + or config.use_accuracy_compatible + ) self.config = config self.use_inference_optimized_reduce_scatter = ( @@ -345,7 +349,7 @@ def forward(self, input_): # Get the embeddings. if self.deterministic_mode: _tp_size = 1 if self.tp_group is None else self.tp_group.size() - if self.config.use_accuracy_compatible and _tp_size <= 1: + if getattr(self.config, "dsa_accuracy_compatible", False) and _tp_size <= 1: output_parallel = _EmbedFp32MainGrad.apply(self.weight, masked_input) else: output_parallel = self.weight[masked_input] @@ -724,6 +728,7 @@ def linear_with_grad_accumulation_and_async_allreduce( grad_output_buffer: Optional[List[torch.Tensor]] = None, wgrad_deferral_limit: Optional[int] = 0, tp_group: Optional[torch.distributed.ProcessGroup] = None, + dsa_accuracy_compatible: bool = False, ) -> torch.Tensor: """Linear layer execution with asynchronous communication and gradient accumulation fusion in backprop. @@ -790,12 +795,7 @@ def linear_with_grad_accumulation_and_async_allreduce( tp_group = get_tensor_model_parallel_group_if_none(tp_group) _tp_size = 1 if tp_group is None else tp_group.size() - if ( - _use_accuracy_compatible() - and _tp_size <= 1 - and not sequence_parallel - and not allreduce_dgrad - ): + if dsa_accuracy_compatible and _tp_size <= 1 and not sequence_parallel and not allreduce_dgrad: output = torch.matmul(input, weight.t()) if bias is not None: output = output + bias @@ -838,24 +838,16 @@ def linear_with_grad_accumulation_and_async_allreduce( def _expert_grads_need_own_dp_domain(config) -> bool: - """Whether expert parameters must be reduced over the expert-data-parallel group. - - mcore normally decides this from ``expert_model_parallel_size > 1`` alone, which - is correct only when the expert tensor-parallel size equals the dense one: the - expert parameter is then sharded exactly like a dense parameter and its - data-parallel domain coincides with ``dp_cp``. - - With ``expert_tensor_parallel_size < tensor_model_parallel_size`` (the accuracy - -compatible topology uses ETP=1 with TP=2) every rank in the tensor-parallel - group holds a FULL copy of the expert weight while consuming only its own - sequence-parallel shard of the tokens. Those partial weight gradients live in a - larger data-parallel domain (``expt_dp`` = TP/ETP times ``dp_cp``) and must be - summed. Leaving ``allreduce=True`` puts them in the dense bucket, which is - reduced over ``dp_cp`` — size 1 in this topology — so the reduction silently - never happens and every expert gradient stays a per-rank partial sum. + """Use expert DP for split expert tensor groups in DSA alignment. + + The default retains EP-only grouping. With DSA TP2/ETP1, expert replicas + consume different sequence shards: reduce their gradients over expt_dp + instead of the dense dp_cp group. """ if config.expert_model_parallel_size > 1: return True + if not getattr(config, 'dsa_accuracy_compatible', False): + return False etp = getattr(config, 'expert_tensor_parallel_size', None) if etp is None: return False @@ -1084,7 +1076,13 @@ def _forward_impl(self, input, weight, *args, **kwargs): if not weight.requires_grad: return linear_with_frozen_weight(input, weight, *args, **kwargs) else: - return linear_with_grad_accumulation_and_async_allreduce(input, weight, *args, **kwargs) + return linear_with_grad_accumulation_and_async_allreduce( + input, + weight, + *args, + dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), + **kwargs, + ) def forward( self, @@ -1132,7 +1130,9 @@ def forward( or self.disable_grad_reduce ): input_parallel = input_ - elif _use_accuracy_compatible() and (self.tp_group is None or self.tp_group.size() <= 1): + elif getattr(self.config, "dsa_accuracy_compatible", False) and ( + self.tp_group is None or self.tp_group.size() <= 1 + ): input_parallel = input_ else: input_parallel = copy_to_tensor_model_parallel_region(input_, group=self.tp_group) @@ -1180,7 +1180,7 @@ def forward( gather_output = runtime_gather_output if gather_output and ( - not _use_accuracy_compatible() + not getattr(self.config, "dsa_accuracy_compatible", False) or (self.tp_group is not None and self.tp_group.size() > 1) ): # All-gather across the partitions. @@ -1411,7 +1411,13 @@ def _forward_impl(self, input, weight, *args, **kwargs): if not weight.requires_grad: return linear_with_frozen_weight(input, weight, *args, **kwargs) else: - return linear_with_grad_accumulation_and_async_allreduce(input, weight, *args, **kwargs) + return linear_with_grad_accumulation_and_async_allreduce( + input, + weight, + *args, + dsa_accuracy_compatible=getattr(self.config, "dsa_accuracy_compatible", False), + **kwargs, + ) def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Forward of RowParallelLinear @@ -1461,7 +1467,9 @@ def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: output_ = reduce_scatter_to_sequence_parallel_region( output_parallel, group=self.tp_group ) - elif _use_accuracy_compatible() and (self.tp_group is None or self.tp_group.size() <= 1): + elif getattr(self.config, "dsa_accuracy_compatible", False) and ( + self.tp_group is None or self.tp_group.size() <= 1 + ): output_ = output_parallel else: output_ = reduce_from_tensor_model_parallel_region(output_parallel, group=self.tp_group) diff --git a/megatron/core/tensor_parallel/mappings.py b/megatron/core/tensor_parallel/mappings.py index 267ec7d61a8..6a1605d08a7 100644 --- a/megatron/core/tensor_parallel/mappings.py +++ b/megatron/core/tensor_parallel/mappings.py @@ -510,10 +510,6 @@ def scatter_to_tensor_model_parallel_region(input_, group=None): def gather_from_tensor_model_parallel_region(input_, group=None): """Wrapper for autograd function: forward: AG, backward: split """ group = get_tensor_model_parallel_group_if_none(group) - from megatron.core.transformer.module import _use_accuracy_compatible - - if _use_accuracy_compatible() and (group is None or group.size() <= 1): - return input_ return _GatherFromModelParallelRegion.apply(input_, group) diff --git a/megatron/core/transformer/experimental_attention_variant/dsa.py b/megatron/core/transformer/experimental_attention_variant/dsa.py index 6c42eb358c1..8e0c135d725 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa.py @@ -22,7 +22,7 @@ dsa_layout, dsa_masking, ) -from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible +from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_config import TransformerConfig @@ -1829,7 +1829,7 @@ def forward( # Detach x and qr to prevent gradients of indexer from flowing back to the main model. _tp_group = getattr(self.pg_collection, "tp", None) _tp_size = 1 if _tp_group is None else _tp_group.size() - if not (_use_accuracy_compatible() and _tp_size <= 1): + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): x = x.detach() qr = qr.detach() diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index 64e7115747f..67fc3106187 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -34,7 +34,7 @@ TEActivationFunctionBuilder, apply_swiglu_sharded_factory, ) -from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible from megatron.core.transformer.moe.moe_utils import ( ProcessGroupCollection, get_align_size_for_quantization, @@ -1351,12 +1351,14 @@ def forward( # The unfused Paddle expert pads tiny GEMMs to 32 rows. The # grouped-storage fallback uses real token counts instead. num_real_tokens = tokens.shape[0] - pad_small_expert = ( - self.config.use_accuracy_compatible - and not self.config.moe_grouped_gemm - and not (self.config.fp8 or self.config.fp4) - and 0 < num_real_tokens < 17 - ) + pad_small_expert = _use_accuracy_compatible() and 0 < num_real_tokens < 17 + if self.config.dsa_accuracy_compatible: + pad_small_expert = ( + self.config.use_accuracy_compatible + and not self.config.moe_grouped_gemm + and not (self.config.fp8 or self.config.fp4) + and 0 < num_real_tokens < 17 + ) if pad_small_expert: num_pad_tokens = 32 - num_real_tokens tokens = torch.cat( diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index 027948f2470..ba9b2cdfd6c 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -12,7 +12,7 @@ from megatron.core.extensions.transformer_engine import HAVE_TE from megatron.core.inference.utils import InferenceMode from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible +from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.moe.moe_utils import ( MoECudaGraphPartialCaptureSignal, MoECudaGraphTensorStore, @@ -576,7 +576,7 @@ def postprocess(self, output: torch.Tensor, shared_expert_output: Optional[torch output, _ = self.fc2_latent_proj(output) if shared_expert_output is not None: - if _use_accuracy_compatible(): + if self.config.dsa_accuracy_compatible: orig_dtype = output.dtype output = (output.float() + shared_expert_output.float()).to(orig_dtype) else: @@ -651,7 +651,7 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): # logging probe: removing the nodes changes bf16 gradient sum # order at the shared input, and makes the router input grad a # 3-way accumulated value that PF's ThreePathCloneAlignMG splits. - if _use_accuracy_compatible() and hidden_states.requires_grad: + if self.config.dsa_accuracy_compatible and hidden_states.requires_grad: _hs_router_path_mg = hidden_states.clone() _hs_dispatcher_path_mg = hidden_states.clone() _hs_shared_path_mg = hidden_states.clone() @@ -664,13 +664,11 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): hidden_states_router = hidden_states hidden_states_dispatch = hidden_states - if _use_accuracy_compatible() and not self.shared_expert_overlap: + if self.config.dsa_accuracy_compatible and not self.shared_expert_overlap: self._accuracy_shared_input = hidden_states_shared shared_expert_output = None else: - shared_expert_output = self.shared_experts_compute( - hidden_states_shared - ) + shared_expert_output = self.shared_experts_compute(hidden_states_shared) probs, routing_map = self.route(hidden_states_router, padding_mask) hidden_states, probs = self.preprocess( hidden_states_dispatch, probs, routing_map @@ -705,7 +703,7 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): if intermediate_tensors is not None: output, shared_expert_output = intermediate_tensors - if _use_accuracy_compatible(): + if self.config.dsa_accuracy_compatible: shared_input = getattr(self, "_accuracy_shared_input", None) if shared_input is not None: shared_expert_output = self.shared_experts_compute(shared_input) diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index 2deec6bdd48..6da367bb659 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -21,6 +21,7 @@ from megatron.core.tensor_parallel.mappings import reduce_from_tensor_model_parallel_region from megatron.core.transformer.cuda_graphs import is_graph_capturing from megatron.core.transformer.enums import CudaGraphModule +from megatron.core.transformer.module import _use_accuracy_compatible from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker from megatron.core.transformer.moe.router_replay import RouterReplay from megatron.core.transformer.transformer_config import TransformerConfig @@ -94,6 +95,7 @@ class _Fp32BackwardIndexSelect(torch.autograd.Function): @staticmethod def forward(ctx, tokens, sorted_indices): + """Forward: index_select tokens by sorted_indices.""" ctx.save_for_backward(sorted_indices) ctx.num_tokens = tokens.shape[0] ctx.hidden = tokens.shape[1] @@ -102,6 +104,7 @@ def forward(ctx, tokens, sorted_indices): @staticmethod def backward(ctx, grad_output): + """Backward: fp32 scatter_add for deterministic accumulation.""" (sorted_indices,) = ctx.saved_tensors grad_tokens = torch.zeros( (ctx.num_tokens, ctx.hidden), dtype=torch.float32, device=grad_output.device @@ -363,6 +366,35 @@ def set_loss_scale(scale: torch.Tensor) -> None: MoEAuxLossAutoScaler.main_loss_backward_scale.copy_(scale) +class _PermuteAlignedAutogradFn(torch.autograd.Function): + """MG-aligned deterministic permute (matches PF _PermuteAlignedPyLayer). + + Forward: tokens.index_select(0, sorted_indices) + Backward: gather(reverse_indices) -> reshape [N, topk, H] -> sum(dim=1) + with fp32 internal accumulation. + """ + + @staticmethod + def forward(ctx, tokens, sorted_indices, reverse_indices_flat, num_tokens, topk, hidden): + """Forward: permute tokens by sorted_indices.""" + ctx.input_dtype = tokens.dtype + ctx.num_tokens = num_tokens + ctx.topk = topk + ctx.hidden = hidden + ctx.save_for_backward(reverse_indices_flat) + permuted_input = tokens.index_select(0, sorted_indices) + return permuted_input + + @staticmethod + def backward(ctx, grad_permuted): + """Backward: fp32 gather-reshape-sum for deterministic unpermute.""" + (reverse_indices_flat,) = ctx.saved_tensors + gathered = grad_permuted.float().index_select(0, reverse_indices_flat) + gathered = gathered.reshape(ctx.num_tokens, ctx.topk, ctx.hidden) + grad_tokens = gathered.sum(dim=1) + return grad_tokens.to(ctx.input_dtype), None, None, None, None, None + + def permute( tokens: torch.Tensor, routing_map: torch.Tensor, @@ -372,6 +404,7 @@ def permute( drop_and_pad: bool = False, tokens_per_expert: Optional[torch.Tensor] = None, align_size: int = 0, + dsa_accuracy_compatible: bool = False, ) -> Tuple[ torch.Tensor, Optional[torch.Tensor], @@ -477,8 +510,9 @@ def permute( num_out_tokens is not None ), "num_out_tokens is required for the argsort-based permute" + rm_orig_bool = routing_map.bool() # [num_tokens, num_experts] # mask [num_tokens, num_experts] -> [num_experts, num_tokens] - routing_map = routing_map.bool().T.contiguous() + routing_map = rm_orig_bool.T.contiguous() # Use argsort to get indices of non-zero entries in row-major order. # This is equivalent to masked_select but produces fixed-shape output, @@ -490,8 +524,32 @@ def permute( if probs is not None: permuted_probs = probs.T.contiguous().reshape(-1)[flat_sorted] - # use the mapping to permute the tokens - if _use_accuracy_compatible() and not drop_and_pad: + # === BIT-EXACT permute backward (gated by MOE_DETERMINISTIC_UNPERMUTE) === + if ( + _use_accuracy_compatible() + and not dsa_accuracy_compatible + and not (drop_and_pad and num_out_tokens is not None) + ): + rm_T_int = routing_map.long() # [num_experts, num_tokens] + tokens_per_expert_local = rm_T_int.sum(dim=-1) # [num_experts] + expert_offsets = torch.zeros(num_experts + 1, dtype=torch.long, device=tokens.device) + expert_offsets[1:] = torch.cumsum(tokens_per_expert_local, dim=0) + position_in_expert_T = rm_T_int.cumsum(dim=-1) - 1 # [num_experts, num_tokens] + global_position = ( + position_in_expert_T + expert_offsets[:-1].unsqueeze(1) + ).T # [num_tokens, num_experts] + + topk_val = int(rm_orig_bool.long().sum(dim=-1)[0].item()) + valid_positions = global_position * rm_orig_bool.long() + reverse_indices_flat = torch.masked_select(valid_positions, rm_orig_bool).reshape( + num_tokens * topk_val + ) + reverse_indices_flat.requires_grad_(False) + + permuted_input = _PermuteAlignedAutogradFn.apply( + tokens, sorted_indices, reverse_indices_flat, num_tokens, topk_val, hidden + ) + elif _use_accuracy_compatible() and not drop_and_pad: # fp32 确定性反向累积,复刻 PF 侧 permute backward permuted_input = _fp32_backward_index_select(tokens, sorted_indices) else: @@ -590,24 +648,50 @@ def unpermute( # allocation. permuted_tokens = permuted_tokens * permuted_probs.unsqueeze(-1) - # Create an output tensor filled with zeros - output_tokens = torch.zeros( - restore_shape, dtype=permuted_tokens.dtype, device=permuted_tokens.device - ) - if torch.are_deterministic_algorithms_enabled(): - # Use index_add which is deterministic when deterministic algorithms are enabled - # and is CUDA graph compatible - output_tokens = torch.zeros( - restore_shape, dtype=permuted_tokens.dtype, device=permuted_tokens.device + _use_deterministic = _use_accuracy_compatible() + + if _use_deterministic and routing_map is not None: + # 确定性 gather+sum 实现(用于逐位对齐) + num_tokens = restore_shape[0] + num_experts = routing_map.shape[1] + routing_map_bool = routing_map.bool() + routing_map_T = routing_map_bool.T.contiguous() # [num_experts, num_tokens] + tokens_per_expert_local = routing_map_T.long().sum(dim=-1) # [num_experts] + expert_offsets = torch.zeros( + num_experts + 1, dtype=torch.long, device=permuted_tokens.device ) - # index_add is deterministic when torch.use_deterministic_algorithms(True) is set - # and is CUDA graph compatible unlike scatter_add - output_tokens.index_add_(0, sorted_indices, permuted_tokens) + expert_offsets[1:] = torch.cumsum(tokens_per_expert_local, dim=0) + position_in_expert_T = routing_map_T.long().cumsum(dim=-1) - 1 # [num_experts, num_tokens] + global_position = position_in_expert_T + expert_offsets[:-1].unsqueeze( + 1 + ) # [num_experts, num_tokens] + global_position_per_token = global_position.T # [num_tokens, num_experts] + topk = int(routing_map_bool.long().sum(dim=-1)[0].item()) + valid_positions = global_position_per_token * routing_map_bool.long() + reverse_indices = valid_positions[routing_map_bool].reshape(num_tokens, topk) + # 用 embedding lookup 替代 index_select(反向是确定性的 scatter,无累加) + gathered = torch.nn.functional.embedding(reverse_indices.reshape(-1), permuted_tokens) + gathered = gathered.reshape(num_tokens, topk, hidden) else: - # Scatter add the permuted_input back to the original positions - output_tokens.scatter_add_( - 0, sorted_indices.unsqueeze(1).expand(-1, hidden), permuted_tokens + # Create an output tensor filled with zeros + output_tokens = torch.zeros( + restore_shape, dtype=permuted_tokens.dtype, device=permuted_tokens.device ) + if torch.are_deterministic_algorithms_enabled(): + # Use index_add which is deterministic when deterministic algorithms are enabled + # and is CUDA graph compatible + output_tokens = torch.zeros( + restore_shape, dtype=permuted_tokens.dtype, device=permuted_tokens.device + ) + # index_add is deterministic when torch.use_deterministic_algorithms(True) is set + # and is CUDA graph compatible unlike scatter_add + output_tokens.index_add_(0, sorted_indices, permuted_tokens) + else: + # Scatter add the permuted_input back to the original positions + output_tokens.scatter_add_( + 0, sorted_indices.unsqueeze(1).expand(-1, hidden), permuted_tokens + ) + return output_tokens.to(dtype=input_dtype) @@ -880,13 +964,13 @@ def compute_topk(scores, topk, num_groups=None, group_topk=None): elif score_function in ("sigmoid", "sqrtsoftplus"): if _use_accuracy_compatible(): if score_function == "sigmoid": - scores = torch.sigmoid(logits.float()) + scores = torch.sigmoid(logits.float()).type_as(logits) else: - scores = torch.nn.functional.softplus(logits.float()).sqrt() + scores = torch.nn.functional.softplus(logits.float()).sqrt().type_as(logits) if expert_bias is not None: scores_for_routing = scores + expert_bias _, top_indices = compute_topk(scores_for_routing, topk, num_groups, group_topk) - scores = torch.gather(scores, dim=1, index=top_indices) + scores = torch.gather(scores, dim=1, index=top_indices).type_as(logits) else: scores, top_indices = compute_topk(scores, topk, num_groups, group_topk) _scores_f64 = scores.double() diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index e25339b5ff1..a1ca6cdbe42 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -7,7 +7,7 @@ from megatron.core.inference.utils import InferenceMode from megatron.core.jit import jit_fuser -from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible from megatron.core.transformer.moe.moe_logging import get_moe_metrics_tracker from megatron.core.transformer.moe.moe_utils import ( MoEAuxLossAutoScaler, @@ -106,9 +106,7 @@ def gating(self, input: torch.Tensor): router_dtype = torch.float64 if self.config.router_accuracy_compatible: inp_shape = input.shape - logits = torch.mm( - input.reshape(-1, inp_shape[-1]).float(), self.weight.float().t() - ) + logits = torch.mm(input.reshape(-1, inp_shape[-1]).float(), self.weight.float().t()) if self.bias is not None: logits = logits + self.bias.float() return logits.view(*inp_shape[:-1], -1) @@ -615,10 +613,11 @@ def _apply_expert_bias( Prevent extra local tokens accumulation on evaluation or activation recomputation """ if self.enable_expert_bias and torch.is_grad_enabled(): - with torch.no_grad(): - if padding_mask is not None: - routing_map = routing_map & (~padding_mask) - self.local_tokens_per_expert += routing_map.sum(dim=0) + if not _use_accuracy_compatible(): + with torch.no_grad(): + if padding_mask is not None: + routing_map = routing_map & (~padding_mask) + self.local_tokens_per_expert += routing_map.sum(dim=0) def routing(self, logits: torch.Tensor, padding_mask: Optional[torch.Tensor] = None): """Top-k routing function diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index e4c72b30ff3..1b26b87e4e4 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -17,7 +17,6 @@ reduce_scatter_to_sequence_parallel_region, ) from megatron.core.transformer.enums import CudaGraphModule -from megatron.core.transformer.module import _use_accuracy_compatible from megatron.core.transformer.moe.fused_a2a import ( fused_combine, fused_dispatch, @@ -305,6 +304,7 @@ def dispatch_postprocess(self, hidden_states, probs): self.local_map, num_out_tokens=tokens_per_expert.sum().item(), fused=self.config.moe_permute_fusion, + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) ) @@ -652,6 +652,7 @@ def dispatch_preprocess( num_out_tokens=self.num_out_tokens, fused=self.config.moe_permute_fusion, drop_and_pad=self.drop_and_pad, + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) return permutated_local_input_tokens, permuted_probs @@ -1378,6 +1379,7 @@ def get_permuted_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> fused=self.permute_fusion, tokens_per_expert=self.tokens_per_expert, align_size=get_align_size_for_quantization(self.config), + dsa_accuracy_compatible=self.config.dsa_accuracy_compatible, ) if self.router_dtype == "fp64": permuted_probs = permuted_probs.to(torch.float64) @@ -1527,7 +1529,7 @@ def token_dispatch( """ if self.shared_experts is not None: self.shared_experts.wait_current_stream() - if _use_accuracy_compatible(): + if self.config.dsa_accuracy_compatible: async_finish = False allocate_on_comm_stream = False dispatched_hidden_states = self._comm_manager.dispatch( @@ -1589,7 +1591,7 @@ def token_combine( # when CUDA_DEVICE_MAX_CONNECTIONS>1. if self.shared_experts is not None: self.shared_experts.wait_current_stream() - if _use_accuracy_compatible(): + if self.config.dsa_accuracy_compatible: async_finish = False allocate_on_comm_stream = False return self._comm_manager.combine(hidden_states, async_finish, allocate_on_comm_stream) diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index c35747cda95..5a91a7da037 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -28,7 +28,7 @@ inference_all_gather_from_tensor_model_parallel_region, ) from megatron.core.transformer.enums import AttnMaskType, LayerType -from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible +from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder, WrappedTorchNorm from megatron.core.transformer.transformer_block import TransformerBlockSubmodules @@ -165,9 +165,7 @@ def roll_tensor(tensor, shifts=-1, dims=-1, cp_group=None, packed_seq_params=Non # Handle packed sequences cases if packed_seq_params is not None: - return _roll_tensor_packed_seq( - tensor, shifts, dims, packed_seq_params, cp_group - ) + return _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group) # Standard rolling behavior when CP is not enabled (cp_group is None or size=1) if cp_group is None or cp_group.size() == 1: @@ -204,25 +202,17 @@ def roll_tensor(tensor, shifts=-1, dims=-1, cp_group=None, packed_seq_params=Non # Start send and recv ops ops = [] if local_rank != 0: - req_send_first_part = torch.distributed.isend( - tensor=tensor_send_list[0], dst=prev_rank - ) + req_send_first_part = torch.distributed.isend(tensor=tensor_send_list[0], dst=prev_rank) ops.append(req_send_first_part) - req_recv_second_part = torch.distributed.irecv( - tensor=tensor_recv_list[1], src=prev_rank - ) + req_recv_second_part = torch.distributed.irecv(tensor=tensor_recv_list[1], src=prev_rank) ops.append(req_recv_second_part) else: # Inserted elements are set to be 0.0. tensor_recv_list[1] = 0 if local_rank != len(global_ranks) - 1: - req_recv_first_part = torch.distributed.irecv( - tensor=tensor_recv_list[0], src=next_rank - ) + req_recv_first_part = torch.distributed.irecv(tensor=tensor_recv_list[0], src=next_rank) ops.append(req_recv_first_part) - req_send_second_part = torch.distributed.isend( - tensor=tensor_send_list[1], dst=next_rank - ) + req_send_second_part = torch.distributed.isend(tensor=tensor_send_list[1], dst=next_rank) ops.append(req_send_second_part) else: # For the last CP rank, the removed elements of second part go into the first part @@ -252,14 +242,12 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No # Notice: This is a naive implementation to test the correctness, # a better solution will only sync the boundary tokens once. - assert dims == -1 or dims == tensor.dim() - 1, ( - "Packed sequence roll only supports the last dimension." - ) + assert ( + dims == -1 or dims == tensor.dim() - 1 + ), "Packed sequence roll only supports the last dimension." assert shifts == -1, "Packed sequence roll only supports a single-token left shift." cu_seqlens = packed_seq_params.cu_seqlens_q - assert cu_seqlens is not None, ( - "Packed sequence parameters must provide cu_seqlens_q." - ) + assert cu_seqlens is not None, "Packed sequence parameters must provide cu_seqlens_q." rolled_tensor = tensor.clone() @@ -301,9 +289,7 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No # The following code is very similar as the code in roll_tensor function local_chunks = tensor_slice.chunk(2, dim=dims) - rolled_chunks = [ - torch.roll(chunk, shifts=shifts, dims=dims) for chunk in local_chunks - ] + rolled_chunks = [torch.roll(chunk, shifts=shifts, dims=dims) for chunk in local_chunks] tensor_send_list = [] tensor_recv_list = [] @@ -311,14 +297,10 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No # Skip empty chunks that can occur when the sequence slice is very small if chunk.size(dims) == 0: tensor_send_list.append( - torch.empty( - chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device - ) + torch.empty(chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device) ) tensor_recv_list.append( - torch.empty( - chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device - ) + torch.empty(chunk.shape[:-1], dtype=chunk.dtype, device=chunk.device) ) continue boundary = chunk.select(dims, shifts).contiguous().clone() @@ -327,22 +309,14 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No ops = [] if local_rank != 0: - ops.append( - torch.distributed.isend(tensor=tensor_send_list[0], dst=prev_rank) - ) - ops.append( - torch.distributed.irecv(tensor=tensor_recv_list[1], src=prev_rank) - ) + ops.append(torch.distributed.isend(tensor=tensor_send_list[0], dst=prev_rank)) + ops.append(torch.distributed.irecv(tensor=tensor_recv_list[1], src=prev_rank)) else: tensor_recv_list[1].zero_() if local_rank != cp_size - 1: - ops.append( - torch.distributed.irecv(tensor=tensor_recv_list[0], src=next_rank) - ) - ops.append( - torch.distributed.isend(tensor=tensor_send_list[1], dst=next_rank) - ) + ops.append(torch.distributed.irecv(tensor=tensor_recv_list[0], src=next_rank)) + ops.append(torch.distributed.isend(tensor=tensor_send_list[1], dst=next_rank)) else: tensor_recv_list[0].copy_(tensor_send_list[1]) @@ -397,17 +371,11 @@ def save_metrics_to_tracker( tracker = MTPLossLoggingHelper.tracker if "loss_values" not in tracker: - tracker["loss_values"] = torch.zeros( - num_layers, device=torch.cuda.current_device() - ) + tracker["loss_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) if "correct_values" not in tracker: - tracker["correct_values"] = torch.zeros( - num_layers, device=torch.cuda.current_device() - ) + tracker["correct_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) if "total_values" not in tracker: - tracker["total_values"] = torch.zeros( - num_layers, device=torch.cuda.current_device() - ) + tracker["total_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) tracker["loss_values"][layer_number] += loss.detach() tracker["correct_values"][layer_number] += correct.detach() @@ -436,32 +404,26 @@ def reduce_metrics_in_tracker(): return loss_values = tracker["loss_values"] - if tracker.get("reduce_group") is not None: - torch.distributed.all_reduce(loss_values, group=tracker.get("reduce_group")) - if tracker.get("avg_group") is not None: + if tracker.get('reduce_group') is not None: + torch.distributed.all_reduce(loss_values, group=tracker.get('reduce_group')) + if tracker.get('avg_group') is not None: torch.distributed.all_reduce( - loss_values, - group=tracker["avg_group"], - op=torch.distributed.ReduceOp.AVG, + loss_values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.AVG ) for key in ["correct_values", "total_values"]: if key not in tracker: continue values = tracker[key] - if tracker.get("reduce_group") is not None: - torch.distributed.all_reduce(values, group=tracker.get("reduce_group")) - if tracker.get("avg_group") is not None: + if tracker.get('reduce_group') is not None: + torch.distributed.all_reduce(values, group=tracker.get('reduce_group')) + if tracker.get('avg_group') is not None: torch.distributed.all_reduce( - values, - group=tracker["avg_group"], - op=torch.distributed.ReduceOp.SUM, + values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.SUM ) @staticmethod - def track_mtp_metrics( - loss_scale, iteration, writer, wandb_writer=None, total_loss_dict=None - ): + def track_mtp_metrics(loss_scale, iteration, writer, wandb_writer=None, total_loss_dict=None): """Track the Multi-Token Prediction (MTP) metrics for logging.""" MTPLossLoggingHelper.reduce_metrics_in_tracker() tracker = MTPLossLoggingHelper.tracker @@ -491,16 +453,15 @@ def track_mtp_metrics( mtp_num_layers = mtp_losses.shape[0] for i in range(mtp_num_layers): - loss_name = f"mtp_{i + 1} loss" - step_acc_name = f"mtp_{i + 1}_acceptance_rate" - cum_acc_name = f"mtp_{i + 1}_cumulative_acceptance_rate" + loss_name = f"mtp_{i+1} loss" + step_acc_name = f"mtp_{i+1}_acceptance_rate" + cum_acc_name = f"mtp_{i+1}_cumulative_acceptance_rate" loss = mtp_losses[i] # Empty masks can leave no valid MTP positions, so clamp denominators to avoid NaNs. step_rate = (mtp_corrects[i] / torch.clamp(mtp_totals[i], min=1)) * 100.0 cum_rate = ( - mtp_cumulative_corrects[i] - / torch.clamp(mtp_cumulative_totals[i], min=1) + mtp_cumulative_corrects[i] / torch.clamp(mtp_cumulative_totals[i], min=1) ) * 100.0 if total_loss_dict is not None: @@ -530,9 +491,7 @@ def _mtp_logits_are_vocab_sharded( def _vocab_parallel_argmax( - vocab_parallel_logits: Tensor, - tp_group: torch.distributed.ProcessGroup, - tp_size: int, + vocab_parallel_logits: Tensor, tp_group: torch.distributed.ProcessGroup, tp_size: int ) -> Tensor: """Return global argmax ids from logits sharded across the vocab dimension.""" vocab_shard_size = vocab_parallel_logits.size(-1) @@ -546,9 +505,9 @@ def _vocab_parallel_argmax( stacked_max_vals = torch.stack(gathered_max_vals, dim=0) stacked_argmax = torch.stack(gathered_argmax, dim=0) winning_rank = stacked_max_vals.argmax(dim=0) # [s, b] - winning_local_argmax = torch.gather( - stacked_argmax, 0, winning_rank.unsqueeze(0) - ).squeeze(0) # [s, b] + winning_local_argmax = torch.gather(stacked_argmax, 0, winning_rank.unsqueeze(0)).squeeze( + 0 + ) # [s, b] return winning_rank * vocab_shard_size + winning_local_argmax # [s, b] @@ -575,11 +534,7 @@ def _compute_mtp_acceptance_counts( "tp_group must be provided when computing MTP acceptance counts " "from vocab-sharded logits under tensor model parallelism." ) - tp_size = ( - torch.distributed.get_world_size(group=tp_group) - if tp_group is not None - else 1 - ) + tp_size = torch.distributed.get_world_size(group=tp_group) if tp_group is not None else 1 # Apply TP rank offsets only when logits are vocab-sharded; gathered logits already # contain global vocab ids in their last dimension. @@ -691,19 +646,14 @@ def mtp_on_this_rank( # with custom PP layout, we support put MTP layers on any pipeline stage if ( not ignore_virtual - and parallel_state.get_virtual_pipeline_model_parallel_world_size() - is not None + and parallel_state.get_virtual_pipeline_model_parallel_world_size() is not None ): - assert vp_stage is not None, ( - "vp_stage must be passed if virtual pipeline is enabled" - ) + assert vp_stage is not None, "vp_stage must be passed if virtual pipeline is enabled" num_layers_to_build = layout.layout[pp_rank][vp_stage].count(LayerType.mtp) mtp_on_this_rank = num_layers_to_build > 0 else: for vpp_rank in range(len(layout.layout[pp_rank])): - num_layers_to_build = layout.layout[pp_rank][vpp_rank].count( - LayerType.mtp - ) + num_layers_to_build = layout.layout[pp_rank][vpp_rank].count(LayerType.mtp) if num_layers_to_build > 0: mtp_on_this_rank = True break @@ -734,9 +684,7 @@ def get_mtp_ranks(pp_ranks: List[int], config: TransformerConfig) -> List[int]: return list(mtp_ranks) -def get_mtp_layer_offset( - config: TransformerConfig, vp_stage: Optional[int] = None -) -> int: +def get_mtp_layer_offset(config: TransformerConfig, vp_stage: Optional[int] = None) -> int: """Get the offset of the MTP layer.""" if config.pipeline_model_parallel_size > 1: if config.pipeline_model_parallel_layout: @@ -751,29 +699,21 @@ def get_mtp_layer_offset( def get_mtp_num_layers_to_build( - config: TransformerConfig, - vp_stage: Optional[int] = None, - pp_rank: Optional[int] = None, + config: TransformerConfig, vp_stage: Optional[int] = None, pp_rank: Optional[int] = None ) -> int: """Get the number of MTP layers to build.""" if config.pipeline_model_parallel_layout is not None: # If we have a custom PP layout, get the number of mtp layers in the layout array. - num_layers_to_build = ( - config.pipeline_model_parallel_layout.get_num_layers_to_build( - layer_type=LayerType.mtp, vp_stage=vp_stage - ) + num_layers_to_build = config.pipeline_model_parallel_layout.get_num_layers_to_build( + layer_type=LayerType.mtp, vp_stage=vp_stage ) - assert ( - num_layers_to_build == config.mtp_num_layers or num_layers_to_build == 0 - ), ( + assert num_layers_to_build == config.mtp_num_layers or num_layers_to_build == 0, ( f"Currently, we only support put all of MTP layers on the last pipeline stage, " f"so the number of MTP layers to build ({num_layers_to_build}) must match " f"mtp_num_layers ({config.mtp_num_layers}) or be 0." ) else: - if parallel_state.is_pipeline_last_stage( - ignore_virtual=False, vp_stage=vp_stage - ): + if parallel_state.is_pipeline_last_stage(ignore_virtual=False, vp_stage=vp_stage): num_layers_to_build = config.mtp_num_layers if config.mtp_num_layers else 0 else: num_layers_to_build = 0 @@ -880,11 +820,7 @@ def process_mtp_loss( if input_ids is None: return hidden_states labels, _ = roll_tensor( - input_ids, - shifts=-1, - dims=-1, - cp_group=cp_group, - packed_seq_params=packed_seq_params, + input_ids, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params ) derived_labels_from_input_ids = True @@ -902,11 +838,7 @@ def process_mtp_loss( # label is fabricated (zeroed). Roll loss_mask in lockstep with the # input_ids -> labels shift so that boundary position is masked. loss_mask, _ = roll_tensor( - loss_mask, - shifts=-1, - dims=-1, - cp_group=cp_group, - packed_seq_params=packed_seq_params, + loss_mask, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params ) # Store the original number of tokens before rolling for proper normalization @@ -923,18 +855,10 @@ def process_mtp_loss( if scale_logits_fn is not None: mtp_logits = scale_logits_fn(mtp_logits) mtp_labels, _ = roll_tensor( - mtp_labels, - shifts=-1, - dims=-1, - cp_group=cp_group, - packed_seq_params=packed_seq_params, + mtp_labels, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params ) loss_mask, num_tokens = roll_tensor( - loss_mask, - shifts=-1, - dims=-1, - cp_group=cp_group, - packed_seq_params=packed_seq_params, + loss_mask, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params ) mtp_loss = compute_language_model_loss(mtp_labels, mtp_logits) @@ -946,12 +870,7 @@ def process_mtp_loss( torch.sum(mtp_loss) * (num_tokens > 0).to(mtp_loss.dtype) ) / num_tokens.clamp(min=1) correct, total = _compute_mtp_acceptance_counts( - mtp_logits, - mtp_labels, - loss_mask, - output_layer, - runtime_gather_output, - tp_group, + mtp_logits, mtp_labels, loss_mask, output_layer, runtime_gather_output, tp_group ) MTPLossLoggingHelper.save_metrics_to_tracker( @@ -960,9 +879,7 @@ def process_mtp_loss( total, mtp_layer_number, config.mtp_num_layers, - avg_group=parallel_state.get_data_parallel_group( - with_context_parallel=True - ), + avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), ) mtp_loss_scale = config.mtp_loss_scaling_factor / config.mtp_num_layers if config.calculate_per_token_loss: @@ -1046,26 +963,18 @@ def __init__( # Validate attention mask type if using transformer-based inner layers if self.submodules.mtp_model_layer is not None and hasattr( - self.submodules.mtp_model_layer, "submodules" + self.submodules.mtp_model_layer, 'submodules' ): from megatron.core.models.hybrid.hybrid_block import HybridStackSubmodules from megatron.core.transformer.transformer_layer import TransformerLayerSubmodules layer_submodules = None - if isinstance( - self.submodules.mtp_model_layer.submodules, HybridStackSubmodules - ): - attention_layer_spec = ( - self.submodules.mtp_model_layer.submodules.attention_layer - ) - if hasattr(attention_layer_spec, "submodules"): - assert isinstance( - attention_layer_spec.submodules, TransformerLayerSubmodules - ) + if isinstance(self.submodules.mtp_model_layer.submodules, HybridStackSubmodules): + attention_layer_spec = self.submodules.mtp_model_layer.submodules.attention_layer + if hasattr(attention_layer_spec, 'submodules'): + assert isinstance(attention_layer_spec.submodules, TransformerLayerSubmodules) layer_submodules = attention_layer_spec.submodules - elif isinstance( - self.submodules.mtp_model_layer.submodules, TransformerLayerSubmodules - ): + elif isinstance(self.submodules.mtp_model_layer.submodules, TransformerLayerSubmodules): layer_submodules = self.submodules.mtp_model_layer.submodules else: raise ValueError( @@ -1073,7 +982,7 @@ def __init__( ) if layer_submodules: self_attention_spec = layer_submodules.self_attention - attn_mask_type = self_attention_spec.params.get("attn_mask_type", "") + attn_mask_type = self_attention_spec.params.get('attn_mask_type', '') assert attn_mask_type in SUPPORTED_ATTN_MASK, ( f"Multi-Token Prediction (MTP) is not yet supported with " f"{attn_mask_type} attention mask type. " @@ -1207,7 +1116,7 @@ def _get_embeddings( decoder_input = decoder_input.detach() _tp_size = 1 if self.tp_group is None else self.tp_group.size() - if not (_use_accuracy_compatible() and _tp_size <= 1): + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): hidden_states = make_viewless_tensor( inp=hidden_states, requires_grad=True, keep_graph=True ) @@ -1221,20 +1130,18 @@ def _get_embeddings( return input_ids, position_ids, padding_mask, decoder_input, hidden_states - def _concat_embeddings( - self, hidden_states: torch.Tensor, decoder_input: torch.Tensor - ): + def _concat_embeddings(self, hidden_states: torch.Tensor, decoder_input: torch.Tensor): """ Concatenate the tokens before sending to transformer layer. """ _tp_size = 1 if self.tp_group is None else self.tp_group.size() decoder_input = apply_module(self.enorm)(decoder_input) - if not (_use_accuracy_compatible() and _tp_size <= 1): + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): decoder_input = make_viewless_tensor( inp=decoder_input, requires_grad=True, keep_graph=True ) hidden_states = apply_module(self.hnorm)(hidden_states) - if not (_use_accuracy_compatible() and _tp_size <= 1): + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): hidden_states = make_viewless_tensor( inp=hidden_states, requires_grad=True, keep_graph=True ) @@ -1248,15 +1155,13 @@ def _concat_embeddings( hidden_states = inference_all_gather_from_tensor_model_parallel_region( hidden_states, self.tp_group, self.config ) - elif not (_use_accuracy_compatible() and _tp_size <= 1): + elif not (self.config.dsa_accuracy_compatible and _tp_size <= 1): hidden_states = gather_from_tensor_model_parallel_region( hidden_states, group=self.tp_group ) # For sequence parallel, scatter after linear_fc and before transformer layer. if self.sequence_parallel: - hidden_states = scatter_to_sequence_parallel_region( - hidden_states, group=self.tp_group - ) + hidden_states = scatter_to_sequence_parallel_region(hidden_states, group=self.tp_group) return hidden_states def _proj_and_transformer_layer( @@ -1341,9 +1246,7 @@ def _postprocess(self, hidden_states: torch.Tensor): # TENorm produces a "viewed" tensor. This will result in schedule.py's # deallocate_output_tensor() throwing an error, so a viewless tensor is # created to prevent this. - hidden_states = make_viewless_tensor( - inp=hidden_states, requires_grad=True, keep_graph=True - ) + hidden_states = make_viewless_tensor(inp=hidden_states, requires_grad=True, keep_graph=True) return hidden_states @@ -1521,20 +1424,19 @@ def checkpoint_handler(): sequence_len_offset, ) - if self.config.recompute_method == "uniform": + if self.config.recompute_method == 'uniform': # Uniformly divide the total number of Transformer layers and checkpoint # the input activation of each divided chunk. # A method to further reduce memory usage reducing checkpoints. - assert self.config.recompute_num_layers == 1, ( - "recompute_num_layers must be 1 for MTP recompute" - ) + assert ( + self.config.recompute_num_layers == 1 + ), "recompute_num_layers must be 1 for MTP recompute" with outer_quantization_context: outputs = checkpoint_handler() - elif self.config.recompute_method == "block": + elif self.config.recompute_method == 'block': # TODO: implement block-based recompute for MTP warnings.warn( - "recompute_method == 'block' is not supported for MTP yet." - " Skipping recompute." + "recompute_method == 'block' is not supported for MTP yet." " Skipping recompute." ) outputs = self._proj_and_transformer_layer( hidden_states=hidden_states, @@ -1596,21 +1498,17 @@ def forward( Union[Tensor, Tuple[Tensor, Tensor]]: The output hidden states tensor of shape [s, b, h], and optionally the updated context tensor if cross-attention is used. """ - assert context is None, ( - "multi token prediction + cross attention is not yet supported." - ) - input_ids, position_ids, padding_mask, decoder_input, hidden_states = ( - self._get_embeddings( - input_ids=input_ids, - position_ids=position_ids, - padding_mask=padding_mask, - embedding=embedding, - hidden_states=hidden_states, - packed_seq_params=packed_seq_params, - ) + assert context is None, "multi token prediction + cross attention is not yet supported." + input_ids, position_ids, padding_mask, decoder_input, hidden_states = self._get_embeddings( + input_ids=input_ids, + position_ids=position_ids, + padding_mask=padding_mask, + embedding=embedding, + hidden_states=hidden_states, + packed_seq_params=packed_seq_params, ) - if self.config.recompute_granularity == "full" and self.training: + if self.config.recompute_granularity == 'full' and self.training: hidden_states = self._checkpointed_forward( hidden_states=hidden_states, decoder_input=decoder_input, @@ -1646,10 +1544,7 @@ def forward( return hidden_states, input_ids, position_ids, padding_mask def sharded_state_dict( - self, - prefix: str = "", - sharded_offsets: tuple = (), - metadata: Optional[dict] = None, + self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None ) -> ShardedStateDict: """ Generate a sharded state dictionary for the multi token prediction layer. @@ -1663,9 +1558,7 @@ def sharded_state_dict( ShardedStateDict: A dictionary containing the sharded state of the multi token prediction layer. """ - sharded_state_dict = super().sharded_state_dict( - prefix, sharded_offsets, metadata - ) + sharded_state_dict = super().sharded_state_dict(prefix, sharded_offsets, metadata) # Backward compatibility: GPT MTP checkpoints were saved with the submodule # named 'transformer_layer'. Remap checkpoint keys so old checkpoints load @@ -1673,8 +1566,7 @@ def sharded_state_dict( # since no older checkpoints exist for them. if self.mtp_layer_pattern is None: apply_prefix_mapping( - sharded_state_dict, - {f"{prefix}mtp_model_layer.": f"{prefix}transformer_layer."}, + sharded_state_dict, {f'{prefix}mtp_model_layer.': f'{prefix}transformer_layer.'} ) return sharded_state_dict @@ -1699,8 +1591,7 @@ class MultiTokenPredictionBlockSubmodules: def _get_mtp_block_submodules( - config: TransformerConfig, - spec: Union[MultiTokenPredictionBlockSubmodules, ModuleSpec], + config: TransformerConfig, spec: Union[MultiTokenPredictionBlockSubmodules, ModuleSpec] ) -> MultiTokenPredictionBlockSubmodules: """ Retrieve or construct MultiTokenPredictionBlockSubmodules based on the provided specification. @@ -1801,25 +1692,21 @@ def __init__( # to the roll_tensor function for proper boundary communication if pg_collection is None: # Use default MPU process groups if not provided - pg_collection = ProcessGroupCollection.use_mpu_process_groups( - required_pgs=["cp", "tp"] - ) + pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp', 'tp']) else: # Ensure the provided process groups include CP - assert hasattr(pg_collection, "cp"), ( - "MultiTokenPredictionBlock pg_collection must have cp process group" - ) + assert hasattr( + pg_collection, 'cp' + ), "MultiTokenPredictionBlock pg_collection must have cp process group" self._build_layers(pg_collection) - assert len(self.layers) > 0, ( - "MultiTokenPredictionBlock must have at least one layer." - ) + assert len(self.layers) > 0, "MultiTokenPredictionBlock must have at least one layer." self.cp_group = pg_collection.cp if self.config.mtp_detach_heads: # Tag MTP params so the optimizer can clip their gradients separately. for param in self.parameters(): - param.grad_norm_group = "mtp" + param.grad_norm_group = 'mtp' def _build_layers(self, pg_collection): # Determine number of depths to build @@ -1839,9 +1726,7 @@ def build_layer_legacy(layer_spec, layer_number): vp_stage=self.vp_stage, pg_collection=pg_collection, mtp_layer_pattern=self.mtp_layer_pattern, - name=(self.name + f".layers.{layer_number}") - if self.name is not None - else None, + name=(self.name + f".layers.{layer_number}") if self.name is not None else None, ) return module @@ -1859,9 +1744,7 @@ def build_layer_with_pattern( pg_collection=pg_collection, mtp_layer_pattern=mtp_layer_pattern, hybrid_submodules=hybrid_submodules, - name=(self.name + f".layers.{layer_number}") - if self.name is not None - else None, + name=(self.name + f".layers.{layer_number}") if self.name is not None else None, ) return module @@ -1953,9 +1836,7 @@ def forward( for iteration in range(self.config.mtp_num_layers): layer_idx = 0 if self.mtp_use_repeated_layer else iteration - (hidden_states, input_ids, position_ids, padding_mask) = self.layers[ - layer_idx - ]( + (hidden_states, input_ids, position_ids, padding_mask) = self.layers[layer_idx]( input_ids=input_ids, position_ids=position_ids, hidden_states=hidden_states, @@ -1980,10 +1861,7 @@ def forward( return hidden_states def sharded_state_dict( - self, - prefix: str = "", - sharded_offsets: tuple = (), - metadata: Optional[dict] = None, + self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None ) -> ShardedStateDict: """ Generate a sharded state dictionary for the multi token prediction module. @@ -1998,18 +1876,16 @@ def sharded_state_dict( token prediction module. """ sharded_state_dict = {} - layer_prefix = f"{prefix}layers." + layer_prefix = f'{prefix}layers.' for layer in self.layers: offset = get_mtp_layer_offset(self.config, self.vp_stage) - sharded_prefix = f"{layer_prefix}{layer.layer_number - 1}." + sharded_prefix = f'{layer_prefix}{layer.layer_number - 1}.' - state_dict_prefix = f"{layer_prefix}{layer.layer_number - 1 - offset}." + state_dict_prefix = f'{layer_prefix}{layer.layer_number - 1 - offset}.' sharded_pp_offset = [] layer_sharded_state_dict = layer.sharded_state_dict( state_dict_prefix, sharded_pp_offset, metadata ) - replace_prefix_for_sharding( - layer_sharded_state_dict, state_dict_prefix, sharded_prefix - ) + replace_prefix_for_sharding(layer_sharded_state_dict, state_dict_prefix, sharded_prefix) sharded_state_dict.update(layer_sharded_state_dict) return sharded_state_dict diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py index c75525dcd59..a7f9db5e60a 100644 --- a/megatron/core/transformer/torch_norm.py +++ b/megatron/core/transformer/torch_norm.py @@ -41,32 +41,32 @@ def __new__( zero_centered_gamma: bool = False, normalization: str = "LayerNorm", ) -> LayerNormInterface: - assert not config.layernorm_zero_centered_gamma, ( - f"zero_centered_gamma not supported by torch LayerNorm" - ) + assert ( + not config.layernorm_zero_centered_gamma + ), f"zero_centered_gamma not supported by torch LayerNorm" - assert not config.persist_layer_norm, ( - f"persist_layer_norm not supported by torch LayerNorm" - ) + assert not config.persist_layer_norm, f"persist_layer_norm not supported by torch LayerNorm" - assert not config.memory_efficient_layer_norm, ( - f"memory_efficient_layer_norm not supported by torch LayerNorm" - ) + assert ( + config.norm_accuracy_compatible or not config.sequence_parallel + ), "sequence parallel not supported by torch LayerNorm" + + assert ( + not config.memory_efficient_layer_norm + ), f"memory_efficient_layer_norm not supported by torch LayerNorm" if config.normalization == "LayerNorm": norm_cls = torch.nn.LayerNorm elif config.normalization == "RMSNorm": - assert is_torch_min_version("2.4.0a0"), ( - "Torch RMSNorm requires PyTorch version >= 2.4.0" - ) + assert is_torch_min_version( + "2.4.0a0" + ), 'Torch RMSNorm requires PyTorch version >= 2.4.0' norm_cls = torch.nn.RMSNorm elif config.normalization == "L2Norm": norm_cls = torch.nn.L2Norm else: - raise Exception( - "Only LayerNorm, RMSNorm and L2Norm are currently supported" - ) + raise Exception("Only LayerNorm, RMSNorm and L2Norm are currently supported") factory_kwargs = {} if config.normalization == "RMSNorm" and config.norm_accuracy_compatible: @@ -108,9 +108,7 @@ def _norm(self, x: torch.Tensor) -> torch.Tensor: torch.Tensor: The L2-normalized tensor. """ x_float = x.float() - return ( - x_float * torch.rsqrt(x_float.pow(2).mean(-1, keepdim=True) + self.eps) - ).type_as(x) + return (x_float * torch.rsqrt(x_float.pow(2).mean(-1, keepdim=True) + self.eps)).type_as(x) def forward(self, x: torch.Tensor) -> torch.Tensor: """ diff --git a/megatron/core/transformer/transformer_block.py b/megatron/core/transformer/transformer_block.py index f01cf55c1fc..b3b977f8053 100755 --- a/megatron/core/transformer/transformer_block.py +++ b/megatron/core/transformer/transformer_block.py @@ -23,11 +23,7 @@ from megatron.core.recompute import checkpointed_forward from megatron.core.transformer.cuda_graphs import annotate_first_last_layer from megatron.core.transformer.enums import InferenceCudaGraphScope, LayerType -from megatron.core.transformer.module import ( - GraphableMegatronModule, - MegatronModule, - _use_accuracy_compatible, -) +from megatron.core.transformer.module import GraphableMegatronModule, MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder from megatron.core.transformer.transformer_config import TransformerConfig @@ -595,7 +591,7 @@ def forward( # already creates viewless tensors. That said, make_viewless_tensor() # is called here to be future-proof and corner-case-proof. _tp_size = int(getattr(self.config, "tensor_model_parallel_size", 1) or 1) - if not (_use_accuracy_compatible() and _tp_size <= 1): + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): hidden_states = make_viewless_tensor( inp=hidden_states, requires_grad=True, keep_graph=True ) @@ -703,7 +699,7 @@ def forward( # deallocate_output_tensor() throwing an error, so a viewless tensor is # created to prevent this. _tp_size = int(getattr(self.config, "tensor_model_parallel_size", 1) or 1) - if not (_use_accuracy_compatible() and _tp_size <= 1): + if not (self.config.dsa_accuracy_compatible and _tp_size <= 1): hidden_states = make_viewless_tensor( inp=hidden_states, requires_grad=True, keep_graph=True ) diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index 74904e5ef57..34c02443365 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -25,9 +25,7 @@ CudaGraphScope, InferenceCudaGraphScope, ) -from megatron.core.transformer.pipeline_parallel_layer_layout import ( - PipelineParallelLayerLayout, -) +from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout from .._rank_utils import log_single_rank from ..fusions.fused_bias_geglu import quick_gelu @@ -103,9 +101,7 @@ class TransformerConfig(ModelParallelConfig): """Number of transformer layers on last pipeline stage. None implies equal layer division across PP ranks.""" - pipeline_model_parallel_layout: Optional[ - Union[str, list, PipelineParallelLayerLayout] - ] = None + pipeline_model_parallel_layout: Optional[Union[str, list, PipelineParallelLayerLayout]] = None """Custom definition of the pipeline parallel partitioning. Support type: - str: e.g., 'Et*3|(tt|)*29,m|L'. Stages are split by '|', replicated stages or layers @@ -142,9 +138,7 @@ class TransformerConfig(ModelParallelConfig): hidden_size: int = field(default=0, metadata={"argparse_meta": {"default": None}}) """Transformer hidden size.""" - num_attention_heads: int = field( - default=0, metadata={"argparse_meta": {"default": None}} - ) + num_attention_heads: int = field(default=0, metadata={"argparse_meta": {"default": None}}) """Number of transformer attention heads.""" attention_backend: AttnBackend = AttnBackend.auto @@ -156,7 +150,7 @@ class TransformerConfig(ModelParallelConfig): softmax_scale: Optional[float] = None """Softmax scale for attention scaling.""" - softmax_type: Literal["vanilla", "off-by-one", "learnable"] = "vanilla" + softmax_type: Literal['vanilla', 'off-by-one', 'learnable'] = 'vanilla' """Applies modified softmax from https://www.evanmiller.org/attention-is-off-by-one.html. Supports both TE FusedAttention and local unfused attention. Supports both a fixed offset and and learnable offset.""" @@ -193,27 +187,23 @@ class TransformerConfig(ModelParallelConfig): """Epsilon value for any LayerNorm/RMSNorm operations.""" norm_accuracy_compatible: bool = field( - default=False, - metadata={"argparse_meta": {"arg_names": ["--norm-accuracy-compatible"]}}, + default=False, metadata={"argparse_meta": {"arg_names": ["--norm-accuracy-compatible"]}} ) """Use native Torch RMSNorm modules instead of Transformer Engine norm modules for alignment.""" router_accuracy_compatible: bool = field( - default=False, - metadata={"argparse_meta": {"arg_names": ["--router-accuracy-compatible"]}}, + default=False, metadata={"argparse_meta": {"arg_names": ["--router-accuracy-compatible"]}} ) """Use an explicit fp32 router GEMM instead of the fused Transformer Engine path.""" layernorm_zero_centered_gamma: bool = field( - default=False, - metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}}, + default=False, metadata={"argparse_meta": {"arg_names": ["--apply-layernorm-1p"]}} ) """If set to True, the LayerNorm is adjusted to center the gamma values around 0. This improves numerical stability.""" add_bias_linear: bool = field( - default=True, - metadata={"argparse_meta": {"arg_names": ["--disable-bias-linear"]}}, + default=True, metadata={"argparse_meta": {"arg_names": ["--disable-bias-linear"]}} ) """Include/exclude a bias term in all linear layers (QKV projections, after core attention, and two in MLP layer).""" @@ -256,7 +246,7 @@ class TransformerConfig(ModelParallelConfig): - An integer N: Represents a (N-1):1 ratio, one full attention layer after (N-1) SWA layers. - A list that defines a custom pattern, e.g.: [1,1,1,1,0,0,0,0], where 1 represents SWA. """ - normalization: Literal["LayerNorm", "RMSNorm"] = "LayerNorm" + normalization: Literal['LayerNorm', 'RMSNorm'] = "LayerNorm" """Which norm to use for normalization layers, valid options are `LayerNorm` and `RMSNorm`.""" qk_layernorm: bool = False @@ -308,12 +298,10 @@ class TransformerConfig(ModelParallelConfig): #################### # attention variant #################### - experimental_attention_variant: Optional[Literal["gated_delta_net", "dsa"]] = None + experimental_attention_variant: Optional[Literal['gated_delta_net', 'dsa']] = None """Type of attention variant to use. Currently support gated_delta_net and dsa.""" - experimental_attention_variant_loss_scale_func: Optional[ - Callable[[torch.Tensor], None] - ] = None + experimental_attention_variant_loss_scale_func: Optional[Callable[[torch.Tensor], None]] = None """Optional hook for experimental attention variants to receive the main loss scale.""" #################### @@ -348,10 +336,12 @@ class TransformerConfig(ModelParallelConfig): backend. Unsupported DSA layouts continue to use the PyTorch fallback.""" dsa_accuracy_compatible: bool = field( - default=False, - metadata={"argparse_meta": {"arg_names": ["--dsa-accuracy-compatible"]}}, + default=False, metadata={"argparse_meta": {"arg_names": ["--dsa-accuracy-compatible"]}} ) - """Use the full-score DSA fallback with explicit softmax backward for alignment.""" + """Use DSA reference numerics: explicit softmax backward, deferred token-loss + normalization, FP32 MoE accumulation and TP1 autograd paths. Disabled by + default to preserve existing models' accuracy-compatible behavior. + """ dsa_indexer_rope_interleaved: bool = False """Whether DSA indexer RoPE should use MLA-style interleaving.""" @@ -540,7 +530,7 @@ class TransformerConfig(ModelParallelConfig): #################### # activation recomputation #################### - recompute_granularity: Optional[Literal["full", "selective"]] = None + recompute_granularity: Optional[Literal['full', 'selective']] = None """Determines which type of activation recompute to use. Megatron-core supports 'selective' activation checkpointing where the submodules set in --recompute-modules is checkpointed. The default is "core_attn" which is the memory intensive part of attention. @@ -551,7 +541,7 @@ class TransformerConfig(ModelParallelConfig): If set, must be 'selective' or 'full'. 'selective' always uses all layers. """ - recompute_method: Optional[Literal["uniform", "block"]] = None + recompute_method: Optional[Literal['uniform', 'block']] = None """Determines which transformer layers will be recomputed. uniform will uniformly divide the total number of transformer layers in a transformer block and recompute the input activation of each divided chunk at the specified granularity. block will recompute the input activations for @@ -588,16 +578,16 @@ class TransformerConfig(ModelParallelConfig): #################### # fp8 related #################### - fp8: Optional[Literal["e4m3", "hybrid"]] = field( + fp8: Optional[Literal['e4m3', 'hybrid']] = field( default=None, metadata={"argparse_meta": {"arg_names": ["--fp8-format"]}} ) """If set, enables the use of FP8 precision through Transformer Engine. There are 2 predefined choices (1) 'e4m3' uniformly uses e4m3 for all FP8 tensors, (2) 'hybrid' uses e4m3 for all FP8 activation and weight tensors and e5m2 for all FP8 output activation gradient tensors.""" - fp8_recipe: Optional[ - Literal["tensorwise", "delayed", "mxfp8", "blockwise", "custom"] - ] = "delayed" + fp8_recipe: Optional[Literal['tensorwise', 'delayed', 'mxfp8', 'blockwise', 'custom']] = ( + "delayed" + ) """If set, enables the use of FP8 precision through Transformer Engine. There are 5 predefined choices (1) 'tensorwise' uses per tensor current scaling recipe, (2) 'delayed' uses delayed scaling recipe, 3) 'mxfp8' for Blackwell architecture only, @@ -625,7 +615,7 @@ class TransformerConfig(ModelParallelConfig): fp8_amax_history_len: int = 1 """The length of the amax history window used for scaling factor computation.""" - fp8_amax_compute_algo: Literal["most_recent", "max"] = "most_recent" + fp8_amax_compute_algo: Literal['most_recent', 'max'] = "most_recent" """Algorithm used for choosing the `amax` value for the scaling factor computation. There are 2 predefined choices: `max` chooses the largest `amax` in the history window, while `most_recent` always chooses the most recently seen value. @@ -674,13 +664,13 @@ class TransformerConfig(ModelParallelConfig): #################### # fp4 related #################### - fp4: Optional[Literal["e2m1"]] = field( + fp4: Optional[Literal['e2m1']] = field( default=None, metadata={"argparse_meta": {"arg_names": ["--fp4-format"]}} ) """If set, enables the use of FP4 precision through Transformer Engine. Currently only supports 'nvfp4' which uses NVFP4BlockScaling recipe (requires TE >= 2.7.0.dev0).""" - fp4_recipe: Optional[Literal["nvfp4", "custom"]] = "nvfp4" + fp4_recipe: Optional[Literal['nvfp4', 'custom']] = "nvfp4" """If set, enables the use of FP4 precision through Transformer Engine. Currently only 'nvfp4' is supported which uses NVFP4BlockScaling recipe for Blackwell+ architecture.""" @@ -794,10 +784,10 @@ class TransformerConfig(ModelParallelConfig): """Scaling factor for routing score in top-k selection, only works when moe_router_pre_softmax enabled. Defaults to None, which means no scaling.""" - moe_router_score_function: Literal["softmax", "sigmoid", "sqrtsoftplus"] = "softmax" + moe_router_score_function: Literal['softmax', 'sigmoid', 'sqrtsoftplus'] = "softmax" """Score function for MoE routing. Can be "softmax", "sigmoid" or "sqrtsoftplus".""" - moe_router_dtype: Optional[Literal["fp32", "fp64"]] = None + moe_router_dtype: Optional[Literal['fp32', 'fp64']] = None """Data type for routing and expert output weighted averaging. Using fp32 or fp64 can improve stability especially when the number of experts is large (e.g. finegrained-moe). None means no changes for dtype.""" @@ -851,9 +841,7 @@ class TransformerConfig(ModelParallelConfig): If a list of load balancing types is provided for `moe_router_load_balancing_type`, a corresponding list of coefficients should be provided here.""" - moe_z_loss_coeff: Optional[float] = ( - None # 1e-3 would be a good start value for z-loss - ) + moe_z_loss_coeff: Optional[float] = None # 1e-3 would be a good start value for z-loss """Scaling coefficient for the z-loss. A starting value of 1e-3 is recommended.""" moe_input_jitter_eps: Optional[float] = None @@ -864,14 +852,14 @@ class TransformerConfig(ModelParallelConfig): specified capacity, similar to GShard, Switch-Transformer, and DeepSpeed-MoE. Note that this is currently unsupported so should remain False.""" - moe_token_dispatcher_type: Literal["allgather", "alltoall", "flex"] = "allgather" + moe_token_dispatcher_type: Literal['allgather', 'alltoall', 'flex'] = "allgather" """The type of token dispatcher to use. The default is 'allgather'. Options are 'allgather','alltoall' and 'flex'.""" moe_enable_deepep: bool = False """[Experimental] Enable DeepEP for efficient token dispatching and combine in MoE models.""" - moe_flex_dispatcher_backend: Literal["deepep", "hybridep"] = "deepep" + moe_flex_dispatcher_backend: Literal['deepep', 'hybridep'] = "deepep" """[Experimental] The backend to use for flex token dispatcher. The default is "deepep". Options are "deepep" and "hybridep". Currently only "hybridep" backend supports the MNNVL case.""" @@ -897,7 +885,7 @@ class TransformerConfig(ModelParallelConfig): max that an expert could see during inference so no tokens are actually dropped. The default setting is False.""" - moe_token_drop_policy: Literal["probs", "position"] = "probs" + moe_token_drop_policy: Literal['probs', 'position'] = "probs" """The policy to drop tokens. Can be either "probs" or "position". If "probs", the tokens with the lowest probabilities will be dropped. If "position", tokens at the end of each batch will be dropped. @@ -999,9 +987,7 @@ class TransformerConfig(ModelParallelConfig): """DEPRECATED and replaced by cuda_graph_impl. When set to true, TransformerLayer layers are swapped with user provided CUDA graphs.""" - cuda_graph_impl: Literal[ - "none", "local", "transformer_engine", "full_iteration" - ] = "none" + cuda_graph_impl: Literal['none', 'local', 'transformer_engine', 'full_iteration'] = "none" """Determines the CUDA graph capture implementation. "none": no CUDA graph. "local": MCore CUDA graph implementation. During training, graphable modules own per-layer @@ -1016,9 +1002,7 @@ class TransformerConfig(ModelParallelConfig): cuda_graph_modules has no effect when cuda_graph_impl="none" and must be empty when cuda_graph_impl="full_iteration".""" - cuda_graph_modules: Union[ - str, CudaGraphModule, List[str], List[CudaGraphModule] - ] = "full" + cuda_graph_modules: Union[str, CudaGraphModule, List[str], List[CudaGraphModule]] = "full" """Selects training capture coverage within per-layer CUDA graphs (local and transformer_engine implementations). Valid values are "attn", "mlp", "moe", "moe_router", "moe_preprocess", and "mamba": @@ -1059,10 +1043,7 @@ class TransformerConfig(ModelParallelConfig): cuda_graph_scope: Optional[ Union[ - str, - CudaGraphModule, - CudaGraphScope, - List[Union[str, CudaGraphModule, CudaGraphScope]], + str, CudaGraphModule, CudaGraphScope, List[Union[str, CudaGraphModule, CudaGraphScope]] ] ] = None """Deprecated: renamed to cuda_graph_modules. Accepted for backward compatibility and @@ -1105,9 +1086,7 @@ class TransformerConfig(ModelParallelConfig): inference_sampling_seed: int = 42 """ Random seed to use for sampling during inference. """ - symmetric_ar_type: Optional[ - Literal["two_shot", "one_shot", "multimem_all_reduce"] - ] = None + symmetric_ar_type: Optional[Literal['two_shot', "one_shot", "multimem_all_reduce"]] = None """What type of symmetric all reduce to use. The default is None which is no use of symmetric memory. """ @@ -1124,7 +1103,7 @@ class TransformerConfig(ModelParallelConfig): inference_disable_triton_nvls_kernels: bool = False """ If true, disables the use of Triton NVLS kernels during inference. """ - inference_grouped_gemm_backend: Literal["flashinfer", "torch", "vllm"] = "vllm" + inference_grouped_gemm_backend: Literal['flashinfer', 'torch', 'vllm'] = "vllm" """Specifies the backend to use for grouped GEMM operations during inference. Options: - 'flashinfer': Uses FlashInfer cutlass_fused_moe. Not compatible with MXFP8. @@ -1140,7 +1119,7 @@ class TransformerConfig(ModelParallelConfig): fp8_recipe='mxfp8'. Set to True to disable fusion and use separate kernel launches (useful for debugging).""" - inference_moe_token_dispatcher_type: Literal["nccl", "nvls"] = "nvls" + inference_moe_token_dispatcher_type: Literal['nccl', 'nvls'] = 'nvls' """Token dispatcher to use for MoE expert parallelism during inference. - 'nccl': AllGather/ReduceScatter via NCCL. Fixed token counts per rank; requires decode-only CUDA graphs (forced automatically). @@ -1173,8 +1152,7 @@ class TransformerConfig(ModelParallelConfig): None causes the states to follow the activation dtype.""" use_mamba_mem_eff_path: bool = field( - default=True, - metadata={"argparse_meta": {"arg_names": ["--disable-mamba-mem-eff-path"]}}, + default=True, metadata={"argparse_meta": {"arg_names": ["--disable-mamba-mem-eff-path"]}} ) """Controls usage of the memory efficient path for Mamba layers.""" @@ -1197,7 +1175,7 @@ class TransformerConfig(ModelParallelConfig): quant_recipe: Optional[RecipeConfig] = None """Configuration of any per-module quantization settings to be applied to the model""" - transformer_impl: Literal["local", "transformer_engine", "inference_optimized"] = ( + transformer_impl: Literal['local', 'transformer_engine', 'inference_optimized'] = ( "transformer_engine" ) """Transformer implementation to use. @@ -1325,26 +1303,26 @@ def __post_init__(self): ) if self.experimental_attention_variant == "gated_delta_net": - assert self.linear_attention_freq is not None, ( - f"linear_attention_freq must be set for linear gated_delta_net." - ) + assert ( + self.linear_attention_freq is not None + ), f"linear_attention_freq must be set for linear gated_delta_net." # Check required parameters - assert self.linear_conv_kernel_dim is not None, ( - "linear_conv_kernel_dim must be set for gated delta net." - ) - assert self.linear_key_head_dim is not None, ( - "linear_key_head_dim must be set for gated delta net." - ) - assert self.linear_value_head_dim is not None, ( - "linear_value_head_dim must be set for gated delta net." - ) - assert self.linear_num_key_heads is not None, ( - "linear_num_key_heads must be set for gated delta net." - ) - assert self.linear_num_value_heads is not None, ( - "linear_num_value_heads must be set for gated delta net." - ) + assert ( + self.linear_conv_kernel_dim is not None + ), "linear_conv_kernel_dim must be set for gated delta net." + assert ( + self.linear_key_head_dim is not None + ), "linear_key_head_dim must be set for gated delta net." + assert ( + self.linear_value_head_dim is not None + ), "linear_value_head_dim must be set for gated delta net." + assert ( + self.linear_num_key_heads is not None + ), "linear_num_key_heads must be set for gated delta net." + assert ( + self.linear_num_value_heads is not None + ), "linear_num_value_heads must be set for gated delta net." assert self.linear_num_value_heads % self.linear_num_key_heads == 0, ( f"linear_num_value_heads ({self.linear_num_value_heads}) must be a multiple of " f"linear_num_key_heads ({self.linear_num_key_heads})." @@ -1380,9 +1358,7 @@ def __post_init__(self): if self.fp8: # cannot support first last layer bf16 with delayed scaling if self.first_last_layers_bf16 and self.fp8_recipe == Fp8Recipe.delayed: - raise ValueError( - "Delayed scaling does not support first / last layer in BF16." - ) + raise ValueError("Delayed scaling does not support first / last layer in BF16.") # max bf16 layers per pipeline stage max_bf16_layers_per_pipeline_stage = ( @@ -1393,8 +1369,7 @@ def __post_init__(self): if self.first_last_layers_bf16: if ( self.num_layers_at_start_in_bf16 < 0 - or self.num_layers_at_start_in_bf16 - > max_bf16_layers_per_pipeline_stage + or self.num_layers_at_start_in_bf16 > max_bf16_layers_per_pipeline_stage ): raise ValueError( f"num_layers_at_start_in_bf16 ({self.num_layers_at_start_in_bf16}) must be " @@ -1403,8 +1378,7 @@ def __post_init__(self): ) if ( self.num_layers_at_end_in_bf16 < 0 - or self.num_layers_at_end_in_bf16 - > max_bf16_layers_per_pipeline_stage + or self.num_layers_at_end_in_bf16 > max_bf16_layers_per_pipeline_stage ): raise ValueError( f"num_layers_at_end_in_bf16 ({self.num_layers_at_end_in_bf16}) must be " @@ -1428,8 +1402,7 @@ def __post_init__(self): raise ValueError("fp8_output_proj must be used together with fp8 mode.") if self.fp8_recipe != Fp8Recipe.mxfp8: raise ValueError( - f"fp8_output_proj requires fp8_recipe='mxfp8', got " - f"'{self.fp8_recipe}'." + f"fp8_output_proj requires fp8_recipe='mxfp8', got " f"'{self.fp8_recipe}'." ) # FP4 validation @@ -1437,9 +1410,7 @@ def __post_init__(self): raise ValueError("fp4_param must be used together with fp4 mode.") if self.fp4 and self.fp8: - raise ValueError( - "fp4 and fp8 cannot be used simultaneously. Please choose one." - ) + raise ValueError("fp4 and fp8 cannot be used simultaneously. Please choose one.") if self.fp4 and self.fp4_recipe == Fp4Recipe.custom: if not self.fp4_quantizer_factory: @@ -1455,18 +1426,13 @@ def __post_init__(self): if self.expert_model_parallel_size > 1 and self.num_moe_experts is None: raise ValueError("num_moe_experts must be non None to use expert-parallel.") - if ( - self.transformer_impl == "inference_optimized" - and self.num_moe_experts is not None - ): + if self.transformer_impl == "inference_optimized" and self.num_moe_experts is not None: if self.expert_tensor_parallel_size > 1: raise ValueError( "Inference-optimized MoE layers does not support expert tensor parallelism." ) if self.moe_expert_capacity_factor is not None: - raise ValueError( - "Inference-optimized MoE layers only support dropless MoE " - ) + raise ValueError("Inference-optimized MoE layers only support dropless MoE ") if self.moe_router_padding_for_quantization: raise ValueError( "Inference-optimized MoE layers do not support padded " @@ -1503,8 +1469,7 @@ def __post_init__(self): ) if ( - self.inference_grouped_gemm_backend - == InferenceGroupedGemmBackend.FLASHINFER + self.inference_grouped_gemm_backend == InferenceGroupedGemmBackend.FLASHINFER and self.fp8 == "mxfp8" ): raise ValueError( @@ -1526,9 +1491,7 @@ def __post_init__(self): if self.num_moe_experts is not None and self.moe_ffn_hidden_size is None: self.moe_ffn_hidden_size = self.ffn_hidden_size - warnings.warn( - "moe_ffn_hidden_size is not set, using ffn_hidden_size instead." - ) + warnings.warn("moe_ffn_hidden_size is not set, using ffn_hidden_size instead.") if self.num_moe_experts is None and self.moe_ffn_hidden_size is not None: is_mixed_model = ( @@ -1576,13 +1539,9 @@ def __post_init__(self): if self.moe_enable_deepep: if self.moe_token_dispatcher_type != "flex": - raise ValueError( - "DeepEP backend is only supported with flex token dispatcher." - ) + raise ValueError("DeepEP backend is only supported with flex token dispatcher.") if self.moe_flex_dispatcher_backend == "hybridep": - raise ValueError( - "Only one backend is supported for flex token dispatcher." - ) + raise ValueError("Only one backend is supported for flex token dispatcher.") self.moe_flex_dispatcher_backend = "deepep" warnings.warn( "moe_enable_deepep is deprecated." @@ -1605,14 +1564,10 @@ def __post_init__(self): f"num_shared_experts * ffn_size_of_each_shared_expert, " f"but got {self.moe_shared_expert_intermediate_size}" ) - if ( - self.moe_shared_expert_overlap - and self.moe_token_dispatcher_type - not in [ - "alltoall", - "flex", - ] - ): + if self.moe_shared_expert_overlap and self.moe_token_dispatcher_type not in [ + "alltoall", + "flex", + ]: raise ValueError( f"moe_shared_expert_overlap only works with alltoall or flex token dispatcher." ) @@ -1670,8 +1625,7 @@ def __post_init__(self): ) if self.cpu_offloading and ( - self.cpu_offloading_num_layers < 0 - or self.cpu_offloading_num_layers >= self.num_layers + self.cpu_offloading_num_layers < 0 or self.cpu_offloading_num_layers >= self.num_layers ): raise ValueError( f"CPU offloading can be done only for layers less than {self.num_layers}" @@ -1705,10 +1659,7 @@ def __post_init__(self): 'recompute_method must be "block" or "uniform"' ) - if ( - self.recompute_granularity != "selective" - and self.recompute_num_layers is None - ): + if self.recompute_granularity != "selective" and self.recompute_num_layers is None: raise ValueError( f"When using recompute_granularity: {self.recompute_granularity} " "recompute_num_layers must be between " @@ -1716,8 +1667,7 @@ def __post_init__(self): f"{self.num_layers // self.pipeline_model_parallel_size}" ) elif ( - self.recompute_granularity == "selective" - and self.recompute_num_layers is not None + self.recompute_granularity == "selective" and self.recompute_num_layers is not None ): raise ValueError( f"When using recompute_granularity: {self.recompute_granularity} " @@ -1756,10 +1706,7 @@ def __post_init__(self): "moe_act in recompute_modules is only supported with moe_grouped_gemm." ) - if ( - "mla_up_proj" in self.recompute_modules - and not self.multi_latent_attention - ): + if "mla_up_proj" in self.recompute_modules and not self.multi_latent_attention: raise ValueError( "mla_up_proj in recompute_modules is only supported with " "multi_latent_attention." @@ -1792,11 +1739,8 @@ def __post_init__(self): ) if self.fp8: - if ( - "moe_act" in self.recompute_modules - or "layernorm" in self.recompute_modules - ): - if self.fp8_recipe == "delayed": + if "moe_act" in self.recompute_modules or "layernorm" in self.recompute_modules: + if self.fp8_recipe == 'delayed': raise ValueError( "Delayed scaling does not support moe_act and layernorm recompute " "for fp8." @@ -1822,9 +1766,9 @@ def __post_init__(self): self.recompute_modules.append("moe") if self.fine_grained_activation_offloading: - assert not self.cpu_offloading, ( - "fine_grained_activation_offloading cannot be enabled with cpu_offloading." - ) + assert ( + not self.cpu_offloading + ), "fine_grained_activation_offloading cannot be enabled with cpu_offloading." assert self.offload_modules is not None and len(self.offload_modules) > 0 allowed_modules = { "core_attn", @@ -1838,22 +1782,16 @@ def __post_init__(self): } invalid_modules = set(self.offload_modules) - allowed_modules assert not invalid_modules, ( - f"Invalid choices for offload_modules: {invalid_modules}. " - f"Allowed modules are: {allowed_modules}" + f'Invalid choices for offload_modules: {invalid_modules}. ' + f'Allowed modules are: {allowed_modules}' ) - if ( - "attn_proj" in self.offload_modules - and "core_attn" not in self.offload_modules - ): + if "attn_proj" in self.offload_modules and "core_attn" not in self.offload_modules: raise ValueError( "attn_proj cannot be set to offload_modules alone without core_attn " "because the input of attn_proj is the output of core_attn, " "which is needed in core_attn.backward()." ) - if ( - self.recompute_granularity == "selective" - and "moe" in self.recompute_modules - ): + if self.recompute_granularity == "selective" and "moe" in self.recompute_modules: offload_inside_moe = {"moe_act", "expert_fc1", "fused_group_mlp"} & set( self.offload_modules ) @@ -1864,25 +1802,20 @@ def __post_init__(self): f"Either remove 'moe' from --recompute-modules or remove " f"{offload_inside_moe} from --offload-modules." ) - assert self.min_offloaded_tensor_size >= 0, ( - "min_offloaded_tensor_size must be non-negative." - ) assert ( - self.activation_offload_fraction >= 0 - and self.activation_offload_fraction <= 1 + self.min_offloaded_tensor_size >= 0 + ), "min_offloaded_tensor_size must be non-negative." + assert ( + self.activation_offload_fraction >= 0 and self.activation_offload_fraction <= 1 ), "activation_offload_fraction must be in range [0, 1]." - assert self.delta_offload_bytes_across_pp_ranks >= 0, ( - "delta_offload_bytes_across_pp_ranks must be non-negative." - ) + assert ( + self.delta_offload_bytes_across_pp_ranks >= 0 + ), "delta_offload_bytes_across_pp_ranks must be non-negative." if "fused_group_mlp" in self.offload_modules: if not self.use_transformer_engine_op_fuser: - raise ValueError( - "fused_group_mlp requires use_transformer_engine_op_fuser." - ) - moe_partial_offload = {"expert_fc1", "moe_act"} & set( - self.offload_modules - ) + raise ValueError("fused_group_mlp requires use_transformer_engine_op_fuser.") + moe_partial_offload = {"expert_fc1", "moe_act"} & set(self.offload_modules) if moe_partial_offload: raise ValueError( "fused_group_mlp offloads the whole fused grouped MLP and cannot be " @@ -1890,9 +1823,7 @@ def __post_init__(self): ) if self.moe_paged_stash: if self.cpu_offloading: - raise ValueError( - "moe_paged_stash cannot be enabled with cpu_offloading." - ) + raise ValueError("moe_paged_stash cannot be enabled with cpu_offloading.") if self.moe_expert_rank_capacity_factor is None: raise ValueError( "moe_paged_stash requires moe_expert_rank_capacity_factor to be set; " @@ -1913,8 +1844,7 @@ def __post_init__(self): self.num_layers_in_first_pipeline_stage is not None or self.num_layers_in_last_pipeline_stage is not None ) and ( - self.account_for_embedding_in_pipeline_split - or self.account_for_loss_in_pipeline_split + self.account_for_embedding_in_pipeline_split or self.account_for_loss_in_pipeline_split ): raise ValueError( "num_layers_in_first_pipeline_stage and num_layers_in_last_pipeline_stage cannot be" @@ -1945,11 +1875,9 @@ def __post_init__(self): # Transfer pipeline_model_parallel_layout from str or list to # PipelineParallelLayerLayout if isinstance(self.pipeline_model_parallel_layout, str): - self.pipeline_model_parallel_layout = ( - PipelineParallelLayerLayout.from_str( - layout=self.pipeline_model_parallel_layout, - pipeline_model_parallel_size=self.pipeline_model_parallel_size, - ) + self.pipeline_model_parallel_layout = PipelineParallelLayerLayout.from_str( + layout=self.pipeline_model_parallel_layout, + pipeline_model_parallel_size=self.pipeline_model_parallel_size, ) elif isinstance(self.pipeline_model_parallel_layout, list): # Since list is not hashable, the initialization will not be cached. @@ -1973,10 +1901,8 @@ def __post_init__(self): self.virtual_pipeline_model_parallel_size = detected_vpp_size # Check whether the layout is valid. - self.mtp_standalone = ( - self.pipeline_model_parallel_layout.validate_layer_layout( - num_layers=self.num_layers, mtp_num_layers=self.mtp_num_layers - ) + self.mtp_standalone = self.pipeline_model_parallel_layout.validate_layer_layout( + num_layers=self.num_layers, mtp_num_layers=self.mtp_num_layers ) # Uneven PP @@ -1989,9 +1915,7 @@ def __post_init__(self): if self.num_layers_in_first_pipeline_stage is not None: if self.num_layers_in_first_pipeline_stage <= 0: - raise ValueError( - "num_layers_in_first_pipeline_stage must be larger than 0" - ) + raise ValueError("num_layers_in_first_pipeline_stage must be larger than 0") if self.virtual_pipeline_model_parallel_size is not None: if ( @@ -2010,9 +1934,7 @@ def __post_init__(self): if self.num_layers_in_last_pipeline_stage is not None: if self.num_layers_in_last_pipeline_stage <= 0: - raise ValueError( - "num_layers_in_last_pipeline_stage must be larger than 0" - ) + raise ValueError("num_layers_in_last_pipeline_stage must be larger than 0") if self.virtual_pipeline_model_parallel_size is not None: if ( @@ -2046,13 +1968,8 @@ def __post_init__(self): # If there are middle PP stages, check number of layers # on each middle PP rank is divisible by VPP size. - if ( - pipeline_parallel_size - and self.virtual_pipeline_model_parallel_size is not None - ): - num_layers_per_middle_pipeline_rank = ( - num_layers // pipeline_parallel_size - ) + if pipeline_parallel_size and self.virtual_pipeline_model_parallel_size is not None: + num_layers_per_middle_pipeline_rank = num_layers // pipeline_parallel_size if ( not num_layers_per_middle_pipeline_rank % self.virtual_pipeline_model_parallel_size @@ -2065,8 +1982,7 @@ def __post_init__(self): ) elif ( - self.account_for_embedding_in_pipeline_split - or self.account_for_loss_in_pipeline_split + self.account_for_embedding_in_pipeline_split or self.account_for_loss_in_pipeline_split ): if self.virtual_pipeline_model_parallel_size is None: num_layers = self.num_layers @@ -2099,12 +2015,9 @@ def __post_init__(self): f"{self.pipeline_model_parallel_size}" ) - num_layers_per_pipeline_rank = ( - num_layers // self.pipeline_model_parallel_size - ) + num_layers_per_pipeline_rank = num_layers // self.pipeline_model_parallel_size if ( - not num_layers_per_pipeline_rank - % self.virtual_pipeline_model_parallel_size + not num_layers_per_pipeline_rank % self.virtual_pipeline_model_parallel_size == 0 ): raise ValueError( @@ -2167,9 +2080,7 @@ def __post_init__(self): if self.activation_func_fp8_input_store: if self.activation_func != F.silu or not self.gated_linear_unit: - raise ValueError( - "Storing activation input in FP8 is supported only for SwiGLU." - ) + raise ValueError("Storing activation input in FP8 is supported only for SwiGLU.") if self.apply_rope_fusion: if self.multi_latent_attention: @@ -2190,46 +2101,33 @@ def __post_init__(self): fused_apply_rotary_pos_emb_thd, ) - if ( - fused_apply_rotary_pos_emb is None - and fused_apply_rotary_pos_emb_thd is None - ): + if fused_apply_rotary_pos_emb is None and fused_apply_rotary_pos_emb_thd is None: raise ValueError( "apply_rope_fusion is not available. Please install TE >= 1.4." ) if self.fused_single_qkv_rope: if self.attention_output_gate: - raise ValueError( - "fused_single_qkv_rope does not support gated attention for now." - ) + raise ValueError("fused_single_qkv_rope does not support gated attention for now.") if self.multi_latent_attention and self.rotary_interleaved: - raise ValueError( - "rotary_interleaved does not work with multi_latent_attention." - ) + raise ValueError("rotary_interleaved does not work with multi_latent_attention.") # MuP (Maximal Update Parameterization) configuration if self.use_mup: # Default base_hidden_size to hidden_size (base model case, width_mult=1.0) if self.mup_base_hidden_size is None: self.mup_base_hidden_size = self.hidden_size - assert self.mup_base_hidden_size > 0, ( - "--mup-base-hidden-size must be positive." - ) + assert self.mup_base_hidden_size > 0, "--mup-base-hidden-size must be positive." # Compute width multiplier self.mup_width_mult = self.hidden_size / self.mup_base_hidden_size # MuP attention scaling: 1/d_head instead of 1/sqrt(d_head). if self.softmax_scale is None: base_head_scale = ( - 1.0 - if self.mup_base_head_dim is None - else self.mup_base_head_dim**0.5 - ) - self.softmax_scale = base_head_scale / ( - self.kv_channels**self.mup_attn_scale_power + 1.0 if self.mup_base_head_dim is None else self.mup_base_head_dim**0.5 ) + self.softmax_scale = base_head_scale / (self.kv_channels**self.mup_attn_scale_power) # MuP output scaling: scale logits by 1/width_mult to keep outputs O(1). # Only auto-set if user hasn't explicitly configured it. @@ -2262,14 +2160,10 @@ def __post_init__(self): self.embedding_init_method_std = self.init_method_std if self.embedding_init_method is None: - if self.init_method is None or ( - self.embedding_init_method_std != self.init_method_std - ): + if self.init_method is None or (self.embedding_init_method_std != self.init_method_std): # In this case, we set both the init method and the embedding init method to # whatever std value requested (or defaulted) for the embedding_init_layer - self.embedding_init_method = init_method_normal( - self.embedding_init_method_std - ) + self.embedding_init_method = init_method_normal(self.embedding_init_method_std) else: # Replicate the current behavior where if you are not changing the std of the # embedding init differently and the init method is set, we fallback to the @@ -2303,17 +2197,13 @@ def __post_init__(self): ) if self.num_moe_experts is not None and self.add_bias_linear: - assert self.expert_tensor_parallel_size == 1, ( - "Bias in Moe is only supported when ETP==1" - ) + assert ( + self.expert_tensor_parallel_size == 1 + ), "Bias in Moe is only supported when ETP==1" - if ( - self.moe_router_enable_expert_bias - and self.moe_router_score_function - not in ( - "sigmoid", - "sqrtsoftplus", - ) + if self.moe_router_enable_expert_bias and self.moe_router_score_function not in ( + "sigmoid", + "sqrtsoftplus", ): raise ValueError( "Expert bias for aux-loss-free routing only supports 'sigmoid' and 'sqrtsoftplus' " @@ -2393,22 +2283,20 @@ def __post_init__(self): self.moe_router_num_groups = self.expert_model_parallel_size if self.enable_cuda_graph or self.external_cuda_graph: - assert self.cuda_graph_impl == "none", ( - "Do not use enable_cuda_graph or external_cuda_graph with cuda_graph_impl." - ) - assert not self.enable_cuda_graph or not self.external_cuda_graph, ( - "enable_cuda_graph and external_cuda_graph cannot be enabled at the same time." - ) + assert ( + self.cuda_graph_impl == "none" + ), "Do not use enable_cuda_graph or external_cuda_graph with cuda_graph_impl." + assert ( + not self.enable_cuda_graph or not self.external_cuda_graph + ), "enable_cuda_graph and external_cuda_graph cannot be enabled at the same time." if self.enable_cuda_graph: - warnings.warn( - "enable_cuda_graph is deprecated, use cuda_graph_impl=local instead." - ) + warnings.warn('enable_cuda_graph is deprecated, use cuda_graph_impl=local instead.') self.cuda_graph_impl = "local" if self.external_cuda_graph: warnings.warn( - "external_cuda_graph is deprecated, " - "use cuda_graph_impl=transformer_engine instead." + 'external_cuda_graph is deprecated, ' + 'use cuda_graph_impl=transformer_engine instead.' ) self.cuda_graph_impl = "transformer_engine" @@ -2437,8 +2325,8 @@ def _scope_to_str(s): self.cuda_graph_modules = _scope_to_str(scope) self.cuda_graph_scope = None - normalized_scopes, deprecated_scopes, used_full_scope = ( - normalize_cuda_graph_modules(self.cuda_graph_modules) + normalized_scopes, deprecated_scopes, used_full_scope = normalize_cuda_graph_modules( + self.cuda_graph_modules ) validate_deprecated_cuda_graph_modules_migration_inputs( deprecated_scopes, self.cuda_graph_impl, self.inference_cuda_graph_scope @@ -2472,9 +2360,7 @@ def _scope_to_str(s): self.cuda_graph_modules = normalized_scopes assert all( isinstance(scope, CudaGraphModule) for scope in self.cuda_graph_modules - ), ( - f"cuda_graph_modules must be a list of CudaGraphModule, got {self.cuda_graph_modules}." - ) + ), f"cuda_graph_modules must be a list of CudaGraphModule, got {self.cuda_graph_modules}." assert self.cuda_graph_impl in [ "none", @@ -2487,10 +2373,7 @@ def _scope_to_str(s): self.inference_cuda_graph_scope, self.cuda_graph_impl ) - assert ( - self.inference_cuda_graph_scope - in ALLOWED_INFERENCE_SCOPES[self.cuda_graph_impl] - ), ( + assert self.inference_cuda_graph_scope in ALLOWED_INFERENCE_SCOPES[self.cuda_graph_impl], ( "Invalid inference CUDA graph scope " f"{self.inference_cuda_graph_scope.name!r} for cuda_graph_impl=" f"{self.cuda_graph_impl!r}." @@ -2500,6 +2383,7 @@ def _scope_to_str(s): ), 'cuda_graph_modules must be empty when cuda_graph_impl="full_iteration".' if self.cuda_graph_impl != "none": + if self.cpu_offloading and self.cuda_graph_impl != "full_iteration": raise ValueError("CUDA graphs not supported with CPU offloading.") @@ -2514,60 +2398,51 @@ def _scope_to_str(s): ): if CudaGraphModule.moe_router not in self.cuda_graph_modules: self.cuda_graph_modules.append(CudaGraphModule.moe_router) - if ( - CudaGraphModule.moe_preprocess - not in self.cuda_graph_modules - ): - self.cuda_graph_modules.append( - CudaGraphModule.moe_preprocess - ) + if CudaGraphModule.moe_preprocess not in self.cuda_graph_modules: + self.cuda_graph_modules.append(CudaGraphModule.moe_preprocess) assert ( CudaGraphModule.moe not in self.cuda_graph_modules or CudaGraphModule.moe_router not in self.cuda_graph_modules - ), "cuda_graph_modules must not contain both moe and moe_router." + ), 'cuda_graph_modules must not contain both moe and moe_router.' if CudaGraphModule.moe_preprocess in self.cuda_graph_modules: - assert CudaGraphModule.moe_router in self.cuda_graph_modules, ( - "moe_preprocess cuda graph is only supported with moe_router cuda graph." - ) + assert ( + CudaGraphModule.moe_router in self.cuda_graph_modules + ), 'moe_preprocess cuda graph is only supported with moe_router cuda graph.' if self.num_moe_experts is None or self.num_moe_experts <= 1: assert ( CudaGraphModule.moe not in self.cuda_graph_modules and CudaGraphModule.moe_router not in self.cuda_graph_modules - ), "moe cuda graph is only supported for MoE." + ), 'moe cuda graph is only supported for MoE.' else: if self.moe_layer_freq == 1 or ( - isinstance(self.moe_layer_freq, list) - and 0 not in self.moe_layer_freq + isinstance(self.moe_layer_freq, list) and 0 not in self.moe_layer_freq ): assert CudaGraphModule.mlp not in self.cuda_graph_modules, ( - "mlp cuda graph is only supported for dense layers, " - "but not found in the model." + 'mlp cuda graph is only supported for dense layers, ' + 'but not found in the model.' ) if ( self.moe_expert_capacity_factor is None or not self.moe_pad_expert_input_to_capacity ): - assert CudaGraphModule.moe not in self.cuda_graph_modules, ( - "moe cuda graph is only supported with drop-padding MoE." - ) - if self.moe_token_dispatcher_type == "alltoall" and ( + assert ( + CudaGraphModule.moe not in self.cuda_graph_modules + ), 'moe cuda graph is only supported with drop-padding MoE.' + if self.moe_token_dispatcher_type == 'alltoall' and ( self.moe_expert_capacity_factor is not None or self.moe_router_padding_for_fp8 ): - assert ( - CudaGraphModule.moe_preprocess - not in self.cuda_graph_modules - ), ( - "moe_preprocess cuda graph is not supported when there are " - "DtoH copies and synchronizations in the preprocess step." + assert CudaGraphModule.moe_preprocess not in self.cuda_graph_modules, ( + 'moe_preprocess cuda graph is not supported when there are ' + 'DtoH copies and synchronizations in the preprocess step.' ) if self.recompute_granularity: if self.recompute_granularity != "selective": - assert self.cuda_graph_impl == "full_iteration", ( - "full recompute is only supported with full iteration CUDA graph." - ) + assert ( + self.cuda_graph_impl == "full_iteration" + ), "full recompute is only supported with full iteration CUDA graph." else: # The recompute module should be inside or outside of the graph scope. # Recompute module coverring graph scope is not allowed. @@ -2577,18 +2452,13 @@ def _scope_to_str(s): ): assert ( CudaGraphModule.moe_router not in self.cuda_graph_modules - ), ( - "moe recompute is not supported with moe_router CUDA graph with: " - ) + ), "moe recompute is not supported with moe_router CUDA graph with: " "--cuda-graph-impl transformer_engine." # Graphed recompute module doesn't accept random number. # full_cudagraph means either full_iteration impl or an empty per-layer scope # (which captures the whole layer). - if ( - self.cuda_graph_impl == "full_iteration" - or not self.cuda_graph_modules - ): + if self.cuda_graph_impl == "full_iteration" or not self.cuda_graph_modules: full_cudagraph = True else: full_cudagraph = False @@ -2613,9 +2483,7 @@ def _scope_to_str(s): and CudaGraphModule.moe not in self.cuda_graph_modules ) or "moe" not in self.recompute_modules - ), ( - "hidden dropout is not supported with graphed MLP/MoE recomputation." - ) + ), "hidden dropout is not supported with graphed MLP/MoE recomputation." if self.moe_input_jitter_eps is not None: assert ( not full_cudagraph @@ -2627,14 +2495,8 @@ def _scope_to_str(s): if self.fine_grained_activation_offloading: offload_modules = set(self.offload_modules or []) if self.cuda_graph_impl == "local": - local_supported_offload_modules = { - "expert_fc1", - "moe_act", - "fused_group_mlp", - } - unsupported_offload_modules = ( - offload_modules - local_supported_offload_modules - ) + local_supported_offload_modules = {"expert_fc1", "moe_act", "fused_group_mlp"} + unsupported_offload_modules = offload_modules - local_supported_offload_modules assert not unsupported_offload_modules, ( "fine-grained activation offloading with cuda_graph_impl='local' " "only supports offload_modules 'expert_fc1', 'moe_act', and " @@ -2661,9 +2523,9 @@ def _scope_to_str(s): "are supported only for expert_fc1, moe_act, or fused_group_mlp " "offload when the full MoE module is not captured." ) - assert CudaGraphModule.moe not in self.cuda_graph_modules, ( - "Token-drop MoE is temporarily not supported with activation offloading." - ) + assert ( + CudaGraphModule.moe not in self.cuda_graph_modules + ), "Token-drop MoE is temporarily not supported with activation offloading." assert self.cuda_graph_warmup_steps > 0, ( "cuda_graph_warmup_steps must be greater than 0 when enabling " "fine-grained activation offloading." @@ -2701,55 +2563,51 @@ def _scope_to_str(s): or fused_sort_chunks_by_index_with_probs is None or fused_unpermute is None ): - raise ValueError( - "fused permutation is not available. Please install TE >= 2.1.0." - ) + raise ValueError("fused permutation is not available. Please install TE >= 2.1.0.") if self.overlap_moe_expert_parallel_comm: # TODO: remove this after we fix the hang issue with torch version < 2.6.0 - assert is_torch_min_version("2.6.0"), ( - "A2A Overlap encounters hang issue with torch version < 2.6.0" - ) + assert is_torch_min_version( + "2.6.0" + ), "A2A Overlap encounters hang issue with torch version < 2.6.0" if self.pipeline_model_parallel_size > 1: assert self.virtual_pipeline_model_parallel_size is not None, ( "If enabling EP A2A overlap, virtual_pipeline_model_parallel_size " "must be specified when pipeline_model_parallel_size > 1" ) # Expert model parallelism requirements - assert self.expert_model_parallel_size > 1, ( - "overlap_moe_expert_parallel_comm is only supported with expert model parallelism" - ) + assert ( + self.expert_model_parallel_size > 1 + ), 'overlap_moe_expert_parallel_comm is only supported with expert model parallelism' assert self.moe_token_dispatcher_type in [ - "alltoall", - "flex", - ], ( - "overlap_moe_expert_parallel_comm is supported with alltoall/flex token dispatcher" - ) + 'alltoall', + 'flex', + ], 'overlap_moe_expert_parallel_comm is supported with alltoall/flex token dispatcher' - assert self.recompute_granularity != "full", ( - "disable full recomputation when enabling overlap_moe_expert_parallel_comm" - ) - assert self.recompute_method is None, ( - "disable recomputation method when enabling overlap_moe_expert_parallel_comm" - ) - assert self.recompute_num_layers is None, ( - "recompute_num_layers must be None when enabling overlap_moe_expert_parallel_comm" - ) - assert "moe" not in self.recompute_modules, ( - "disable moe in recompute_modules when enabling overlap_moe_expert_parallel_comm" - ) + assert ( + self.recompute_granularity != 'full' + ), 'disable full recomputation when enabling overlap_moe_expert_parallel_comm' + assert ( + self.recompute_method is None + ), 'disable recomputation method when enabling overlap_moe_expert_parallel_comm' + assert ( + self.recompute_num_layers is None + ), 'recompute_num_layers must be None when enabling overlap_moe_expert_parallel_comm' + assert ( + "moe" not in self.recompute_modules + ), 'disable moe in recompute_modules when enabling overlap_moe_expert_parallel_comm' # Check if bf16 or fp16 is used - assert self.bf16 or self.fp16, ( - "overlap_moe_expert_parallel_comm is only supported with bf16 or fp16 model" - ) + assert ( + self.bf16 or self.fp16 + ), 'overlap_moe_expert_parallel_comm is only supported with bf16 or fp16 model' - assert not self.moe_shared_expert_overlap, ( - "disable moe_shared_expert_overlap when enabling overlap_moe_expert_parallel_comm" - ) - assert self.mtp_num_layers is None or self.mtp_num_layers == 1, ( - "MTP layernum only supports 1 when enabling overlap_moe_expert_parallel_comm." - ) + assert ( + not self.moe_shared_expert_overlap + ), 'disable moe_shared_expert_overlap when enabling overlap_moe_expert_parallel_comm' + assert ( + self.mtp_num_layers is None or self.mtp_num_layers == 1 + ), 'MTP layernum only supports 1 when enabling overlap_moe_expert_parallel_comm.' if self.cuda_graph_impl != "none": if self.cuda_graph_impl == "transformer_engine": @@ -2757,38 +2615,38 @@ def _scope_to_str(s): CudaGraphModule.moe not in self.cuda_graph_modules and CudaGraphModule.mlp not in self.cuda_graph_modules ), ( - "CUDA graph scope on moe and mlp is not " - "supported with overlap_moe_expert_parallel_comm" + 'CUDA graph scope on moe and mlp is not ' + 'supported with overlap_moe_expert_parallel_comm' ) # Check delay_wgrad_compute compatibility if self.delay_wgrad_compute: - assert self.overlap_moe_expert_parallel_comm, ( - "overlap_moe_expert_parallel_comm must be enabled when enabling delay_wgrad_compute" - ) + assert ( + self.overlap_moe_expert_parallel_comm + ), 'overlap_moe_expert_parallel_comm must be enabled when enabling delay_wgrad_compute' if self.cuda_graph_impl == "transformer_engine": assert is_te_min_version("2.10.0"), ( - "TE version >= 2.10.0 is required for delay_wgrad_compute with " - "partial cuda graph" + 'TE version >= 2.10.0 is required for delay_wgrad_compute with ' + 'partial cuda graph' ) if self.overlap_dispatch_backward_with_experts_wgrad: assert not self.overlap_moe_expert_parallel_comm, ( - "overlap_moe_expert_parallel_comm must be disabled when enabling " - "overlap_dispatch_backward_with_experts_wgrad." - ) - assert is_te_min_version("2.3.0"), ( - "TE version >= 2.3.0 is required for overlap_dispatch_backward_with_experts_wgrad" + 'overlap_moe_expert_parallel_comm must be disabled when enabling ' + 'overlap_dispatch_backward_with_experts_wgrad.' ) + assert is_te_min_version( + "2.3.0" + ), 'TE version >= 2.3.0 is required for overlap_dispatch_backward_with_experts_wgrad' assert not self.delay_wgrad_compute, ( - "delay_wgrad_compute and overlap_dispatch_backward_with_experts_wgrad " - "are mutually exclusive; use only one" + 'delay_wgrad_compute and overlap_dispatch_backward_with_experts_wgrad ' + 'are mutually exclusive; use only one' ) if self.ep_overlap_early_attn_memory_release: assert self.overlap_moe_expert_parallel_comm, ( - "overlap_moe_expert_parallel_comm must be enabled when enabling " - "ep_overlap_early_attn_memory_release" + 'overlap_moe_expert_parallel_comm must be enabled when enabling ' + 'ep_overlap_early_attn_memory_release' ) if self.context_parallel_size > 1 and self.cp_comm_type is not None: @@ -2798,14 +2656,14 @@ def _scope_to_str(s): f"the total number of transformer layers ({self.num_layers})!" ) else: - assert isinstance(self.cp_comm_type, str), ( - "Unsupported communication type for context parallelism!" - ) + assert isinstance( + self.cp_comm_type, str + ), "Unsupported communication type for context parallelism!" - assert self.pipeline_model_parallel_size > 0, ( - f"Pipeline model parallel size must be larger than 0 \ + assert ( + self.pipeline_model_parallel_size > 0 + ), f"Pipeline model parallel size must be larger than 0 \ when enable --standalone-embedding-stage and --standalone-loss-stage" - ) if ( self.num_moe_experts is not None @@ -2821,14 +2679,10 @@ def _scope_to_str(s): raise ImportError( "packaging is not installed. Please install it with `pip install packaging`." ) - assert is_torch_min_version("2.7.0a0"), ( - "Must have at least torch version 2.7 or higher" - ) + assert is_torch_min_version("2.7.0a0"), "Must have at least torch version 2.7 or higher" assert is_te_min_version("2.3.0") or get_te_version() == PkgVersion( "2.3.0.dev0+39c0e70" - ), ( - "Must have at least TE version 2.3 or higher to use symmetric memory all reduce" - ) + ), "Must have at least TE version 2.3 or higher to use symmetric memory all reduce" if self.no_rope_freq: assert not self.flash_decode, "flash_decode cannot be used with no_rope." @@ -2855,9 +2709,7 @@ def _scope_to_str(s): assert not self.use_kitchen if self.experimental_attention_variant == "dsa": - assert not self.apply_rope_fusion, ( - "RoPE fusion is not supported for DSAttention" - ) + assert not self.apply_rope_fusion, "RoPE fusion is not supported for DSAttention" if self.context_parallel_size > 1: cp_comm_types = ( self.cp_comm_type @@ -2878,9 +2730,9 @@ def _scope_to_str(s): "inference_fuse_tp_communication is only supported " "for inference_optimized transformer implementation." ) - assert self.num_moe_experts is None, ( - "--inference-fuse-tp-communication is not supported for MoE models." - ) + assert ( + self.num_moe_experts is None + ), "--inference-fuse-tp-communication is not supported for MoE models." if self.inference_disable_triton_nvls_kernels: assert self.transformer_impl == "inference_optimized", ( @@ -2889,9 +2741,9 @@ def _scope_to_str(s): ) if self.batch_invariant_mode: - assert self.attention_backend == AttnBackend.flash, ( - "Batch invariant mode only supports FlashAttention" - ) + assert ( + self.attention_backend == AttnBackend.flash + ), "Batch invariant mode only supports FlashAttention" @dataclass @@ -2962,17 +2814,13 @@ class MLATransformerConfig(TransformerConfig): def __post_init__(self): super().__post_init__() - if ( - self.multi_latent_attention - and self.apply_rope_fusion - and self.rope_type != "yarn" - ): + if self.multi_latent_attention and self.apply_rope_fusion and self.rope_type != "yarn": raise ValueError("apply_rope_fusion for MLA only works with YARN RoPE.") if self.attention_output_gate: raise NotImplementedError("Output gate is not supported for MLA yet.") if self.cache_mla_latents: - assert self.apply_rope_fusion is False, ( - "Rope Fusion is not compatible with caching latents" - ) + assert ( + self.apply_rope_fusion is False + ), "Rope Fusion is not compatible with caching latents" diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index 82e2051ff5b..5a584e59f42 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -22,7 +22,7 @@ from megatron.core.transformer.enums import CudaGraphModule, InferenceCudaGraphScope, LayerType from megatron.core.transformer.identity_op import IdentityFuncOp, IdentityOp from megatron.core.transformer.mlp import MLP -from megatron.core.transformer.module import GraphableMegatronModule, _use_accuracy_compatible +from megatron.core.transformer.module import GraphableMegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder from megatron.core.transformer.transformer_config import TransformerConfig @@ -941,7 +941,7 @@ def _forward_post_mlp( # won't result in memory savings (like the data loader, or # p2p_communication), it serves to document the origin of this # 'view' tensor. - if _use_accuracy_compatible() and self.config.tensor_model_parallel_size <= 1: + if self.config.dsa_accuracy_compatible and self.config.tensor_model_parallel_size <= 1: output = hidden_states else: output = make_viewless_tensor( diff --git a/tests/unit_tests/distributed/test_finalize_model_grads.py b/tests/unit_tests/distributed/test_finalize_model_grads.py index 431ef59aa77..762e1fdd81d 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -1,5 +1,4 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. -import inspect import os import pytest @@ -21,14 +20,52 @@ from tests.unit_tests.test_utilities import Utils -def test_uac_keeps_mcore_num_tokens_scaling(): - """Shipped finalize_model_grads must not skip 1/num_tokens under UAC (E-172).""" - src = inspect.getsource(finalize_model_grads) - assert "loss_normalized_in_graph = False" in src - assert "num_tokens = None" not in src.split("loss_normalized_in_graph = False", 1)[1][ - :800 - ] - assert "scale_gradients(1.0 / dp_size)" not in src +@pytest.mark.parametrize( + "accuracy,dsa,expected", [(False, False, 1.0), (True, False, 4.0), (True, True, 1.0)] +) +def test_token_normalization_preserves_legacy_and_dsa_contracts( + monkeypatch, accuracy, dsa, expected +): + """Legacy in-graph normalization averages DP; DSA divides by global tokens.""" + import importlib + from types import SimpleNamespace + + implementation = importlib.import_module("megatron.core.distributed.finalize_model_grads") + module = importlib.import_module("megatron.core.transformer.module") + monkeypatch.setattr(module, "_use_accuracy_compatible", lambda: accuracy) + config = TransformerConfig( + num_layers=1, hidden_size=8, num_attention_heads=1, dsa_accuracy_compatible=dsa + ) + gradient = torch.tensor(8.0, device="cuda") + events = [] + + def scale(value): + events.append("scale") + gradient.mul_(value) + + model = SimpleNamespace( + config=config, + parameters=lambda: (), + scale_gradients=scale, + finish_grad_sync=lambda **kwargs: events.append("sync"), + ) + monkeypatch.setattr(implementation, "get_model_config", lambda model: model.config) + for name in ( + "_allreduce_conditional_embedding_grads", + "_allreduce_non_tensor_model_parallel_grads", + "_allreduce_word_embedding_grads", + "_allreduce_position_embedding_grads", + "reset_model_temporary_tensors", + ): + monkeypatch.setattr(implementation, name, lambda *args: None) + monkeypatch.setattr(parallel_state, "get_data_parallel_world_size", lambda **kwargs: 2) + monkeypatch.setattr(implementation, "get_pp_last_rank", lambda group: 0) + monkeypatch.setattr(dist, "broadcast", lambda tensor, **kwargs: None) + monkeypatch.setattr(dist, "all_reduce", lambda tensor, **kwargs: tensor.mul_(2)) + groups = SimpleNamespace(tp=None, pp=None, embd=None, pos_embd=None, dp_cp=None) + finalize_model_grads([model], num_tokens=torch.tensor(4.0, device="cuda"), pg_collection=groups) + assert gradient.item() == expected + assert events == ["sync", "scale"] class _RouterExpertBiasModel(torch.nn.Module): diff --git a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py index d05ab24df38..13cc7419c45 100644 --- a/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py +++ b/tests/unit_tests/models/test_experimental_attention_variant_module_specs.py @@ -49,9 +49,7 @@ def _make_backend(fuse_layernorm=True): backend.linear.return_value = _FakeLinear backend.column_parallel_linear.return_value = _FakeColumnParallelLinear backend.row_parallel_linear.return_value = _FakeRowParallelLinear - backend.column_parallel_layer_norm_linear.return_value = ( - _FakeLayerNormColumnParallelLinear - ) + backend.column_parallel_layer_norm_linear.return_value = _FakeLayerNormColumnParallelLinear backend.fuse_layernorm_and_linear.return_value = fuse_layernorm backend.core_attention.return_value = _FakeCoreAttention @@ -109,12 +107,7 @@ def _fn(variant): @pytest.mark.parametrize( "variant, expected", - [ - ("gated_delta_net", True), - ("dsa", False), - (None, False), - ("some_unknown_variant", False), - ], + [("gated_delta_net", True), ("dsa", False), (None, False), ("some_unknown_variant", False)], ) def test_variants(self, variant, expected): """Validate linear-attention variant classification across supported and unsupported names.""" @@ -206,9 +199,7 @@ def test_list_freq_wrong_length_raises(self): def test_none_for_non_linear_variant(self): """Verify non-linear variants default to all-standard attention when freq is None.""" cfg = _make_config( - num_layers=4, - linear_attention_freq=None, - experimental_attention_variant="dsa", + num_layers=4, linear_attention_freq=None, experimental_attention_variant="dsa" ) assert self._fn(cfg) == [0, 0, 0, 0] @@ -304,9 +295,7 @@ def _call(self, cfg=None, backend=None): ) if cfg is None: - cfg = _make_config( - multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True - ) + cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True) if backend is None: backend = _make_backend() return get_dsa_module_spec_for_backend(cfg, backend=backend) @@ -344,9 +333,7 @@ def test_returns_absorbed_mla_self_attention_spec(self): def test_core_attention_is_dsa(self): """Verify MLA core_attention is wrapped with DSAttention.""" - from megatron.core.transformer.experimental_attention_variant.dsa import ( - DSAttention, - ) + from megatron.core.transformer.experimental_attention_variant.dsa import DSAttention spec = self._call() core = spec.submodules.core_attention @@ -354,9 +341,7 @@ def test_core_attention_is_dsa(self): def test_dsa_indexer_structure(self): """Verify DSA indexer wiring uses expected backend linear/norm modules.""" - from megatron.core.transformer.experimental_attention_variant.dsa import ( - DSAIndexer, - ) + from megatron.core.transformer.experimental_attention_variant.dsa import DSAIndexer spec = self._call() indexer = spec.submodules.core_attention.submodules.indexer @@ -405,28 +390,22 @@ def test_accuracy_compatible_qk_rmsnorm(self): def test_qk_layernorm_disabled(self): """Verify q/kv layernorm becomes IdentityOp, skipping backend.layer_norm for qk.""" backend = _make_backend() - cfg = _make_config( - multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=False - ) + cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=False) spec = self._call(cfg=cfg, backend=backend) assert spec.submodules.q_layernorm is IdentityOp assert spec.submodules.kv_layernorm is IdentityOp # backend.layer_norm is still called for the indexer k_norm (for_qk=True at line 94), # but NOT for the outer qk_norm (line 105-107 takes the else branch). # Exactly one for_qk=True call should exist (from the indexer, not from qk_norm). - qk_calls = [ - c for c in backend.layer_norm.call_args_list if c.kwargs.get("for_qk") - ] - assert len(qk_calls) == 1, ( - f"Expected 1 for_qk=True call (indexer only), got {len(qk_calls)}" - ) + qk_calls = [c for c in backend.layer_norm.call_args_list if c.kwargs.get("for_qk")] + assert ( + len(qk_calls) == 1 + ), f"Expected 1 for_qk=True call (indexer only), got {len(qk_calls)}" def test_linear_projections(self): """Verify Q/KV projection slots and backend.column_parallel_linear call count.""" backend = _make_backend() - cfg = _make_config( - multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True - ) + cfg = _make_config(multi_latent_attention=True, qk_l2_norm=False, qk_layernorm=True) spec = self._call(cfg=cfg, backend=backend) subs = spec.submodules assert subs.linear_q_proj == _FakeColumnParallelLinear @@ -458,18 +437,14 @@ class TestGetExperimentalAttentionVariantModuleSpec: def test_dispatches_to_variant_handler(self, variant, target_fn): """Verify dispatcher routes each variant name to its corresponding builder function.""" backend = _make_backend() - cfg = _make_config( - experimental_attention_variant=variant, normalization="RMSNorm" - ) + cfg = _make_config(experimental_attention_variant=variant, normalization="RMSNorm") with patch(f"{self.MODULE}.{target_fn}") as mock_fn: mock_fn.return_value = ModuleSpec(module=MagicMock) from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( get_experimental_attention_variant_module_spec, ) - result = get_experimental_attention_variant_module_spec( - cfg, backend=backend - ) + result = get_experimental_attention_variant_module_spec(cfg, backend=backend) mock_fn.assert_called_once_with(config=cfg, backend=backend) assert result is mock_fn.return_value @@ -494,15 +469,12 @@ class TestGetTransformerLayerWithExperimentalAttentionVariantSpec: def _make_attention_spec(self, fuse_input_layernorm=True): """Construct a mock attention spec with configurable fuse metadata.""" - return ModuleSpec( - module=MagicMock, metainfo={"fuse_input_layernorm": fuse_input_layernorm} - ) + return ModuleSpec(module=MagicMock, metainfo={"fuse_input_layernorm": fuse_input_layernorm}) def _make_mlp_spec(self, fuse_pre_mlp_layernorm=True): """Construct a mock MLP spec with configurable fuse metadata.""" return ModuleSpec( - module=MagicMock, - metainfo={"fuse_pre_mlp_layernorm": fuse_pre_mlp_layernorm}, + module=MagicMock, metainfo={"fuse_pre_mlp_layernorm": fuse_pre_mlp_layernorm} ) def test_all_experimental_no_moe(self): @@ -526,10 +498,7 @@ def test_all_experimental_no_moe(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=attn_spec, ), - patch( - f"{self.MODULE}._get_dense_mlp_module_spec", - return_value=(mlp_spec, True), - ), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(mlp_spec, True)), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -565,14 +534,8 @@ def test_hybrid_attention_pattern(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=exp_attn_spec, ), - patch( - f"{self.MODULE}._get_self_attention_module_spec", - return_value=std_attn_spec, - ), - patch( - f"{self.MODULE}._get_dense_mlp_module_spec", - return_value=(mlp_spec, True), - ), + patch(f"{self.MODULE}._get_self_attention_module_spec", return_value=std_attn_spec), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(mlp_spec, True)), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -608,13 +571,8 @@ def test_hybrid_moe_pattern(self): f"{self.MODULE}.get_experimental_attention_variant_module_spec", return_value=attn_spec, ), - patch( - f"{self.MODULE}._get_moe_module_spec", return_value=(moe_spec, False) - ), - patch( - f"{self.MODULE}._get_dense_mlp_module_spec", - return_value=(dense_spec, True), - ), + patch(f"{self.MODULE}._get_moe_module_spec", return_value=(moe_spec, False)), + patch(f"{self.MODULE}._get_dense_mlp_module_spec", return_value=(dense_spec, True)), ): specs = get_transformer_layer_with_experimental_attention_variant_spec( cfg, backend=backend @@ -678,8 +636,7 @@ def test_get_transformer_block_with_experimental_attention_variant_spec( ) backend = _make_backend() fake_layer_specs = [ - ModuleSpec(module=TransformerLayer, submodules=MagicMock()) - for _ in range(num_layers) + ModuleSpec(module=TransformerLayer, submodules=MagicMock()) for _ in range(num_layers) ] with ( @@ -700,25 +657,17 @@ def test_get_transformer_block_with_experimental_attention_variant_spec( # Without explicit layout, slicing comes from offset + num_layers_to_build. with ( patch( - f"{self.MODULE}.get_transformer_layer_offset", - return_value=offset, + f"{self.MODULE}.get_transformer_layer_offset", return_value=offset ) as mock_offset, patch( - f"{self.MODULE}.get_num_layers_to_build", - return_value=num_layers_to_build, + f"{self.MODULE}.get_num_layers_to_build", return_value=num_layers_to_build ) as mock_num_layers, ): - result = ( - get_transformer_block_with_experimental_attention_variant_spec( - cfg, vp_stage=vp_stage, pp_rank=pp_rank - ) + result = get_transformer_block_with_experimental_attention_variant_spec( + cfg, vp_stage=vp_stage, pp_rank=pp_rank ) - mock_offset.assert_called_once_with( - cfg, vp_stage=vp_stage, pp_rank=pp_rank - ) - mock_num_layers.assert_called_once_with( - cfg, vp_stage=vp_stage, pp_rank=pp_rank - ) + mock_offset.assert_called_once_with(cfg, vp_stage=vp_stage, pp_rank=pp_rank) + mock_num_layers.assert_called_once_with(cfg, vp_stage=vp_stage, pp_rank=pp_rank) assert isinstance(result, TransformerBlockSubmodules) assert result.layer_specs == [fake_layer_specs[i] for i in expected_ids] diff --git a/tests/unit_tests/optimizer/test_reproducible_norm.py b/tests/unit_tests/optimizer/test_reproducible_norm.py deleted file mode 100644 index 0b672f3355b..00000000000 --- a/tests/unit_tests/optimizer/test_reproducible_norm.py +++ /dev/null @@ -1,109 +0,0 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. - -from types import SimpleNamespace - -import pytest -import torch - -from megatron.core.optimizer.clip_grads import clip_grad_by_total_norm_fp32 -from megatron.core.optimizer.optimizer import ChainedOptimizer -from megatron.core.optimizer.optimizer_config import OptimizerConfig -from megatron.core.optimizer.reproducible_norm import ReproducibleL2Norm - - -@pytest.mark.parametrize( - 'values,expected', - [ - ([0.0], 0.0), - ([3.0, 4.0], 25.0), - ([4096.0, 1.0], 16777216.0), - ([4096.0, 1.0, 1.0, 1.0], 16777220.0), - ([2.0**-70], 2.0**-140), - ([float('inf')], float('inf')), - ([3e38], float('inf')), - ], -) -def test_sum_squares_rounding(values, expected): - norm = ReproducibleL2Norm() - _, squared = norm.finish(norm.accumulate(norm.zeros(), norm.tensor(values))) - assert squared.item() == expected - - -def test_nan_and_invalid_input(): - norm = ReproducibleL2Norm() - actual, _ = norm.finish( - norm.accumulate(norm.zeros(), norm.tensor([float('inf'), float('nan')])) - ) - assert torch.isnan(actual).item() - with pytest.raises(TypeError, match='FP32'): - norm.accumulate(norm.zeros(), norm.tensor([1.0]).bfloat16()) - bins = norm.zeros() - bins[22] = 2**40 + 1 - with pytest.raises(OverflowError): - norm.finish(bins) - - -def test_layout_and_chunk_invariance(): - norm = ReproducibleL2Norm() - gradient = torch.arange(1, 12289, device='cuda', dtype=torch.float32).reshape(96, 128) / 16384 - whole = norm.accumulate(norm.zeros(), gradient) - transposed = norm.accumulate(norm.zeros(), gradient.T.contiguous(), chunk_size=123) - split = norm.accumulate(norm.zeros(), gradient.flatten()[:1234]) - split += norm.accumulate(norm.zeros(), gradient.flatten()[1234:]) - assert torch.equal(whole, transposed) - assert torch.equal(whole, split) - - -def test_chained_owner_groups_and_clipping(monkeypatch): - rank = torch.distributed.get_rank() - world = torch.distributed.get_world_size() - singleton = None - for member in range(world): - group = torch.distributed.new_group([member]) - if member == rank: - singleton = group - config = OptimizerConfig(use_accuracy_compatible=True, clip_grad=1.0) - # Dense 3 and 4 have distinct owners; the expert 12 is replicated between - # singleton stats groups. Finishing each child separately loses this contract. - dense = [torch.tensor([3.0 if rank == 0 else 4.0], device='cuda')] if rank < 2 else [] - if world == 1: - dense = [torch.tensor([3.0, 4.0], device='cuda')] - expert = [torch.tensor([12.0], device='cuda')] - children = [ - SimpleNamespace( - config=config, - get_grads_for_grad_norm=lambda _group=None: dense, - get_grad_stats_parallel_group=lambda: torch.distributed.group.WORLD, - ), - SimpleNamespace( - config=config, - get_grads_for_grad_norm=lambda _group=None: expert, - get_grad_stats_parallel_group=lambda: singleton, - ), - ] - actual = ChainedOptimizer(children).get_grad_norm() - assert actual.item() == 13.0 - parameter = torch.nn.Parameter(torch.zeros(2, device='cuda')) - parameter.grad = torch.tensor([3.0, 4.0], device='cuda') - monkeypatch.setattr('megatron.core.optimizer.clip_grads.multi_tensor_scale_tensor_impl', None) - clip_grad_by_total_norm_fp32([parameter], 1.0, actual) - expected = torch.tensor([3.0, 4.0], device='cuda') * (1.0 / (actual + 1e-6)) - assert torch.equal(parameter.grad, expected) - - -@pytest.fixture(scope='module', autouse=True) -def distributed_norm_device(): - import os - - torch.cuda.set_device(int(os.environ.get('LOCAL_RANK', '0'))) - owns_group = not torch.distributed.is_initialized() - if owns_group: - torch.distributed.init_process_group(backend='nccl') - yield - if owns_group: - torch.distributed.destroy_process_group() - - -@pytest.fixture(scope='session') -def ensure_test_data(): - """The norm tests are self-contained and do not consume external datasets.""" diff --git a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py index d66edc26411..979fff5282f 100644 --- a/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py +++ b/tests/unit_tests/tensor_parallel/test_accuracy_tp1_migration.py @@ -4,10 +4,9 @@ import ast import os -import sys import unittest from pathlib import Path -from types import ModuleType, SimpleNamespace +from types import SimpleNamespace from typing import List, Optional from unittest.mock import patch @@ -43,19 +42,6 @@ def backward(ctx, grad_output): return (grad_output,) + (None,) * 8 -class _SentinelGather(torch.autograd.Function): - last = None - - @staticmethod - def forward(ctx, input_, group): - _SentinelGather.last = (input_, group) - return input_ * 2 - - @staticmethod - def backward(ctx, grad_output): - return grad_output, None - - class _FakeGroup: def __init__(self, size): self._size = size @@ -99,7 +85,6 @@ def _load_named(rel: str, name: str, extra_ns=None, class_name=None): "_use_accuracy_compatible": _use_accuracy_compatible, "get_tensor_model_parallel_group_if_none": lambda g: g, "LinearWithGradAccumulationAndAsyncCommunication": _SentinelApply, - "_GatherFromModelParallelRegion": _SentinelGather, "parallel_state": SimpleNamespace(get_tensor_model_parallel_world_size=lambda: _TP["size"]), "custom_backward": _custom_backward, "Variable": torch.autograd.Variable, @@ -115,9 +100,6 @@ def _load_named(rel: str, name: str, extra_ns=None, class_name=None): "megatron/core/tensor_parallel/layers.py", "linear_with_grad_accumulation_and_async_allreduce" ) linear_with_grad_accumulation_and_async_allreduce.warned = True -gather_from_tensor_model_parallel_region = _load_named( - "megatron/core/tensor_parallel/mappings.py", "gather_from_tensor_model_parallel_region" -) deallocate_output_tensor = _load_named( "megatron/core/pipeline_parallel/schedules.py", "deallocate_output_tensor" ) @@ -151,7 +133,7 @@ def test_embedding_configuration_controls_gradient_destination(self): weight.main_grad = torch.zeros_like(weight, dtype=torch.float32) instances.append( SimpleNamespace( - config=SimpleNamespace(use_accuracy_compatible=enabled), + config=SimpleNamespace(dsa_accuracy_compatible=enabled), deterministic_mode=True, tp_group=_FakeGroup(1), weight=weight, @@ -159,7 +141,7 @@ def test_embedding_configuration_controls_gradient_destination(self): ) ) for instance in instances: - enabled = instance.config.use_accuracy_compatible + enabled = instance.config.dsa_accuracy_compatible with patch.dict( os.environ, { @@ -279,7 +261,7 @@ def test_tp1_forward_dgrad_wgrad_bias_matches_f_linear(self): w = _cuda_bf16([1, 0, -1, 0, 1, 1, 1, -1, 0, 0, 1, -1], (4, 3), device).requires_grad_(True) b = _cuda_bf16([1, -1, 0, 2], (4,), device).requires_grad_(True) out = linear_with_grad_accumulation_and_async_allreduce( - x, w, b, False, False, False, None, 0, _FakeGroup(1) + x, w, b, False, False, False, None, 0, _FakeGroup(1), dsa_accuracy_compatible=True ) xref = x.detach().clone().requires_grad_(True) wref = w.detach().clone().requires_grad_(True) @@ -296,8 +278,8 @@ def test_tp1_forward_dgrad_wgrad_bias_matches_f_linear(self): torch.testing.assert_close(w.grad, wref.grad, atol=0, rtol=0) torch.testing.assert_close(b.grad, bref.grad, atol=0, rtol=0) - def test_off_delegates_to_native_custom_function(self): - _UAC["on"] = False + def test_global_accuracy_without_dsa_keeps_native_custom_function(self): + _UAC["on"] = True device = torch.device("cuda") x = torch.ones(2, 3, device=device, dtype=torch.bfloat16, requires_grad=True) w = torch.ones(4, 3, device=device, dtype=torch.bfloat16) @@ -312,52 +294,16 @@ def test_tp2_or_allreduce_skips_tp1_matmul_path(self): x = torch.ones(2, 3, device=device, dtype=torch.bfloat16, requires_grad=True) w = torch.ones(4, 3, device=device, dtype=torch.bfloat16) linear_with_grad_accumulation_and_async_allreduce( - x, w, None, False, False, False, None, 0, _FakeGroup(2) + x, w, None, False, False, False, None, 0, _FakeGroup(2), dsa_accuracy_compatible=True ) self.assertIsNotNone(_SentinelApply.last) _SentinelApply.last = None linear_with_grad_accumulation_and_async_allreduce( - x, w, None, False, True, False, None, 0, _FakeGroup(1) + x, w, None, False, True, False, None, 0, _FakeGroup(1), dsa_accuracy_compatible=True ) self.assertIsNotNone(_SentinelApply.last) -@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") -class TestGatherTp1Identity(unittest.TestCase): - def setUp(self): - # The production wrapper imports the gate lazily; isolate only that - # dependency, retaining the actual gather wrapper and routing branch. - module = ModuleType("megatron.core.transformer.module") - module._use_accuracy_compatible = _use_accuracy_compatible - replacement = patch.dict(sys.modules, {module.__name__: module}) - replacement.start() - self.addCleanup(replacement.stop) - - def tearDown(self): - _UAC["on"] = False - _SentinelGather.last = None - - def test_uac_tp1_returns_same_tensor(self): - _UAC["on"] = True - t = torch.arange(6.0, device="cuda").reshape(2, 3) - out = gather_from_tensor_model_parallel_region(t, _FakeGroup(1)) - self.assertIs(out, t) - self.assertIsNone(_SentinelGather.last) - out_none = gather_from_tensor_model_parallel_region(t, None) - self.assertIs(out_none, t) - - def test_tp2_or_off_delegates_to_gather_function(self): - _UAC["on"] = True - t = torch.arange(6.0, device="cuda").reshape(2, 3) - out = gather_from_tensor_model_parallel_region(t, _FakeGroup(2)) - self.assertIsNotNone(_SentinelGather.last) - self.assertTrue(torch.equal(out, t * 2)) - _SentinelGather.last = None - _UAC["on"] = False - gather_from_tensor_model_parallel_region(t, _FakeGroup(1)) - self.assertIsNotNone(_SentinelGather.last) - - @unittest.skipUnless(torch.cuda.is_available(), "CUDA required") class TestPipelineHelpersTp1(unittest.TestCase): def tearDown(self): @@ -370,7 +316,9 @@ def test_deallocate_uac_tp1_preserves_storage(self): _TP["size"] = 1 t = torch.arange(4.0, device="cuda", requires_grad=True) data_before = t.data.clone() - deallocate_output_tensor(t, True) + deallocate_output_tensor( + t, True, SimpleNamespace(dsa_accuracy_compatible=True, tensor_model_parallel_size=1) + ) torch.testing.assert_close(t.data, data_before, atol=0, rtol=0) self.assertEqual(tuple(t.shape), (4,)) @@ -397,6 +345,7 @@ def test_backward_step_uac_tp1_uses_autograd_not_custom(self): grad_scale_func=None, deallocate_pipeline_outputs=True, tensor_model_parallel_size=1, + dsa_accuracy_compatible=True, ) gin = backward_step(x, y, go, cfg) torch.testing.assert_close(gin, go * 3, atol=0, rtol=0) @@ -411,6 +360,7 @@ def test_backward_step_off_or_tp2_uses_custom_backward(self): grad_scale_func=None, deallocate_pipeline_outputs=True, tensor_model_parallel_size=1, + dsa_accuracy_compatible=False, ) _UAC["on"] = False _CUSTOM_BWD["calls"] = [] @@ -421,6 +371,7 @@ def test_backward_step_off_or_tp2_uses_custom_backward(self): x2 = torch.tensor([[1.0, 2.0], [3.0, 4.0]], device="cuda", requires_grad=True) y2 = x2 * 2 go2 = torch.ones_like(y2) + cfg.dsa_accuracy_compatible = True cfg.tensor_model_parallel_size = 2 _UAC["on"] = True _CUSTOM_BWD["calls"] = [] diff --git a/tests/unit_tests/tensor_parallel/test_layers.py b/tests/unit_tests/tensor_parallel/test_layers.py index 5c4398aefcf..be5795b2642 100644 --- a/tests/unit_tests/tensor_parallel/test_layers.py +++ b/tests/unit_tests/tensor_parallel/test_layers.py @@ -1,5 +1,4 @@ # Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. -from pathlib import Path from types import SimpleNamespace import pytest @@ -19,40 +18,25 @@ def test_expert_grads_need_own_dp_domain_etp_lt_tp(): expert_model_parallel_size=1, tensor_model_parallel_size=2, expert_tensor_parallel_size=1, + dsa_accuracy_compatible=True, ) assert _expert_grads_need_own_dp_domain(frozen) is True + frozen.dsa_accuracy_compatible = False + assert _expert_grads_need_own_dp_domain(frozen) is False eq = SimpleNamespace( - expert_model_parallel_size=1, - tensor_model_parallel_size=2, - expert_tensor_parallel_size=2, + expert_model_parallel_size=1, tensor_model_parallel_size=2, expert_tensor_parallel_size=2 ) assert _expert_grads_need_own_dp_domain(eq) is False ep2 = SimpleNamespace( - expert_model_parallel_size=2, - tensor_model_parallel_size=2, - expert_tensor_parallel_size=1, + expert_model_parallel_size=2, tensor_model_parallel_size=2, expert_tensor_parallel_size=1 ) assert _expert_grads_need_own_dp_domain(ep2) is True missing = SimpleNamespace( - expert_model_parallel_size=1, - tensor_model_parallel_size=2, - expert_tensor_parallel_size=None, + expert_model_parallel_size=1, tensor_model_parallel_size=2, expert_tensor_parallel_size=None ) assert _expert_grads_need_own_dp_domain(missing) is False -def test_expert_dp_domain_is_wired_into_linear_allreduce(): - """Column/Row/TE expert allreduce must use the ETP 1 else False args.context_parallel_size = cp - args.position_embedding_type = "rope" + args.position_embedding_type = 'rope' args.num_experts = 8 args.train_iters = 1 - args.ckpt_format = "torch_dist" + args.ckpt_format = 'torch_dist' args.moe_router_topk = 2 args.moe_router_pre_softmax = False args.lr = 3e-5 @@ -585,10 +534,10 @@ def create_test_args( args.moe_grouped_gemm = False args.bf16 = True if fp8 is not None: - args.fp8 = "e4m3" + args.fp8 = 'e4m3' if full_recompute: - args.recompute_granularity = "full" - args.recompute_method = "uniform" + args.recompute_granularity = 'full' + args.recompute_method = 'uniform' args.recompute_num_layers = 1 else: args.recompute_granularity = None @@ -601,26 +550,19 @@ def create_test_args( def get_batch(self, seq_length, micro_batch_size): data = list(range(seq_length)) - input_ids = ( - torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - ) - labels = ( - 1 - + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - ) - position_ids = ( - torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - ) + input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + labels = 1 + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + position_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() attention_mask = torch.ones( (micro_batch_size, 1, seq_length, seq_length), dtype=bool ).cuda() loss_mask = torch.ones(seq_length).repeat((micro_batch_size, 1)).cuda() batch = { - "tokens": input_ids, - "labels": labels, - "loss_mask": loss_mask, - "attention_mask": attention_mask, - "position_ids": position_ids, + 'tokens': input_ids, + 'labels': labels, + 'loss_mask': loss_mask, + 'attention_mask': attention_mask, + 'position_ids': position_ids, } return batch @@ -651,9 +593,7 @@ def get_packed_batch(self, seq_lengths, micro_batch_size): # Convert to tensors with shape [batch, total_seq_length] input_ids = torch.tensor(input_ids_list, dtype=torch.int64).unsqueeze(0).cuda() labels = torch.tensor(labels_list, dtype=torch.int64).unsqueeze(0).cuda() - position_ids = ( - torch.tensor(position_ids_list, dtype=torch.int64).unsqueeze(0).cuda() - ) + position_ids = torch.tensor(position_ids_list, dtype=torch.int64).unsqueeze(0).cuda() # Create attention mask for packed sequences (all ones for simplicity) attention_mask = torch.ones( @@ -665,8 +605,7 @@ def get_packed_batch(self, seq_lengths, micro_batch_size): # Create cumulative sequence lengths for PackedSeqParams cu_seqlens = torch.tensor( - [0] + [sum(seq_lengths[: i + 1]) for i in range(len(seq_lengths))], - dtype=torch.int32, + [0] + [sum(seq_lengths[: i + 1]) for i in range(len(seq_lengths))], dtype=torch.int32 ).cuda() packed_seq_params = PackedSeqParams( @@ -674,16 +613,16 @@ def get_packed_batch(self, seq_lengths, micro_batch_size): cu_seqlens_kv=cu_seqlens, max_seqlen_q=max(seq_lengths), max_seqlen_kv=max(seq_lengths), - qkv_format="thd", + qkv_format='thd', ) batch = { - "tokens": input_ids, - "labels": labels, - "loss_mask": loss_mask, - "attention_mask": attention_mask, - "position_ids": position_ids, - "packed_seq_params": packed_seq_params, + 'tokens': input_ids, + 'labels': labels, + 'loss_mask': loss_mask, + 'attention_mask': attention_mask, + 'position_ids': position_ids, + 'packed_seq_params': packed_seq_params, } return batch @@ -697,9 +636,7 @@ def test_sharded_state_dict(self, tp, cp): args = self.create_test_args(tp, cp, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) model_parallel_cuda_manual_seed(_SEED) pg_collection = ProcessGroupCollection.use_mpu_process_groups() @@ -728,9 +665,7 @@ def test_forward_backward(self, tmp_path_dist_ckpt, tp, cp, full_recompute): """Test MTP forward and backward with gptmodel.""" tp_ref = 1 cp_ref = 1 - args = self.create_test_args( - tp_ref, cp_ref, self.seq_length, self.micro_batch_size - ) + args = self.create_test_args(tp_ref, cp_ref, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) Utils.initialize_model_parallel( @@ -751,7 +686,7 @@ def test_forward_backward(self, tmp_path_dist_ckpt, tp, cp, full_recompute): tracker = MTPLossLoggingHelper.tracker mtp_loss_ref = None assert "loss_values" in tracker - mtp_loss_ref = tracker["loss_values"].clone() + mtp_loss_ref = tracker['loss_values'].clone() MTPLossLoggingHelper.clean_metrics_in_tracker() iteration = 123 @@ -762,7 +697,7 @@ def set_ckpt_path(ckpt_path): args.load = ckpt_path with TempNamedDir( - tmp_path_dist_ckpt / "test_mtp_model_reconfiguration_model_A" + tmp_path_dist_ckpt / 'test_mtp_model_reconfiguration_model_A' ) as ckpt_dir_A: set_ckpt_path(ckpt_dir_A) save_checkpoint( @@ -779,18 +714,12 @@ def set_ckpt_path(ckpt_path): # Test with different TP/CP configuration Utils.destroy_model_parallel() args = self.create_test_args( - tp, - cp, - self.seq_length, - self.micro_batch_size, - full_recompute=full_recompute, + tp, cp, self.seq_length, self.micro_batch_size, full_recompute=full_recompute ) set_args(args) set_ckpt_path(ckpt_dir_A) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) gpt_model, optimizer, opt_param_scheduler = setup_model_and_optimizer( ModelType.encoder_or_decoder, self.model_provider ) @@ -800,9 +729,7 @@ def set_ckpt_path(ckpt_path): batch = get_batch_on_this_cp_rank( batch, is_hybrid_cp=False, cp_group=get_context_parallel_group() ) - tokens, labels, loss_mask, attention_mask, position_ids, output_ref = ( - batch.values() - ) + tokens, labels, loss_mask, attention_mask, position_ids, output_ref = batch.values() output = gpt_model[0].forward( input_ids=tokens, position_ids=position_ids, @@ -812,11 +739,9 @@ def set_ckpt_path(ckpt_path): ) tracker = MTPLossLoggingHelper.tracker assert "loss_values" in tracker - mtp_loss = tracker["loss_values"].clone() + mtp_loss = tracker['loss_values'].clone() # Average MTP loss across CP ranks for comparison with reference - pg_collection = ProcessGroupCollection.use_mpu_process_groups( - required_pgs=["cp"] - ) + pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp']) torch.distributed.all_reduce( mtp_loss, group=pg_collection.cp, op=torch.distributed.ReduceOp.AVG ) @@ -845,21 +770,14 @@ def test_fp8_support(self, full_recompute): """Test MTP with FP8 training enabled.""" tp = 1 cp = 1 - fp8 = "e4m3" + fp8 = 'e4m3' args = self.create_test_args( - tp, - cp, - self.seq_length, - self.micro_batch_size, - fp8, - full_recompute=full_recompute, + tp, cp, self.seq_length, self.micro_batch_size, fp8, full_recompute=full_recompute ) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) batch = self.get_batch(self.seq_length, self.micro_batch_size) tokens, labels, loss_mask, attention_mask, position_ids = batch.values() gpt_model, optimizer, opt_param_scheduler = setup_model_and_optimizer( @@ -874,9 +792,7 @@ def test_fp8_support(self, full_recompute): loss_mask=loss_mask, ) - assert ( - output.dtype == torch.float32 - ) # Output should be converted back to float32 + assert output.dtype == torch.float32 # Output should be converted back to float32 loss = output.mean() loss.backward() @@ -896,18 +812,16 @@ def test_packed_sequences(self, tp, cp): set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) # Get packed batch batch = self.get_packed_batch(seq_lengths, micro_batch_size=1) - tokens = batch["tokens"] - labels = batch["labels"] - loss_mask = batch["loss_mask"] - attention_mask = batch["attention_mask"] - position_ids = batch["position_ids"] - packed_seq_params = batch["packed_seq_params"] + tokens = batch['tokens'] + labels = batch['labels'] + loss_mask = batch['loss_mask'] + attention_mask = batch['attention_mask'] + position_ids = batch['position_ids'] + packed_seq_params = batch['packed_seq_params'] # Create model model_parallel_cuda_manual_seed(_SEED) @@ -937,7 +851,7 @@ def test_packed_sequences(self, tp, cp): # Verify MTP loss was computed tracker = MTPLossLoggingHelper.tracker assert "loss_values" in tracker - mtp_loss = tracker["loss_values"].clone() + mtp_loss = tracker['loss_values'].clone() assert mtp_loss.shape[0] == args.mtp_num_layers MTPLossLoggingHelper.clean_metrics_in_tracker() @@ -969,18 +883,12 @@ def test_packed_sequences_with_full_recompute(self): total_seq_length = sum(seq_lengths) args = self.create_test_args( - tp=1, - cp=1, - sequence_length=total_seq_length, - micro_batch_size=1, - full_recompute=True, + tp=1, cp=1, sequence_length=total_seq_length, micro_batch_size=1, full_recompute=True ) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=1, context_parallel_size=1 - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) batch = self.get_packed_batch(seq_lengths, micro_batch_size=1) @@ -995,12 +903,12 @@ def test_packed_sequences_with_full_recompute(self): ) output = gpt_model[0].forward( - input_ids=batch["tokens"], - position_ids=batch["position_ids"], - attention_mask=batch["attention_mask"], - labels=batch["labels"], - loss_mask=batch["loss_mask"], - packed_seq_params=batch["packed_seq_params"], + input_ids=batch['tokens'], + position_ids=batch['position_ids'], + attention_mask=batch['attention_mask'], + labels=batch['labels'], + loss_mask=batch['loss_mask'], + packed_seq_params=batch['packed_seq_params'], ) # Backward must run end-to-end through the recomputed MTP layer. @@ -1012,9 +920,7 @@ def test_packed_sequences_with_full_recompute(self): def test_roll_tensor_none_input(self): """Test that roll_tensor returns (None, None) when given None input.""" - Utils.initialize_model_parallel( - tensor_model_parallel_size=1, context_parallel_size=1 - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) result, sum_val = roll_tensor(None, shifts=-1, dims=-1) assert result is None assert sum_val is None @@ -1027,9 +933,7 @@ def test_roll_tensor_shifts_left_and_zeroes_last(self): are not provided (RL training): label[i] = input_id[i+1], last position zeroed. The end-to-end derivation is covered by process_mtp_loss (see input_ids path). """ - Utils.initialize_model_parallel( - tensor_model_parallel_size=1, context_parallel_size=1 - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) # Simulate input_ids [batch=2, seq=5] input_ids = torch.tensor( [[10, 20, 30, 40, 50], [60, 70, 80, 90, 100]], dtype=torch.int64 @@ -1049,13 +953,13 @@ def test_process_mtp_loss_skips_when_no_labels_and_no_input_ids(self): hidden_size=8, num_layers=2, num_attention_heads=2, mtp_num_layers=1 ) hidden_states = torch.ones(2, 1, 4) - called = {"value": False} + called = {'value': False} def output_layer(hidden, weight=None, runtime_gather_output=None): return hidden.clone(), None def compute_language_model_loss(mtp_labels, mtp_logits): - called["value"] = True + called['value'] = True return torch.ones_like(mtp_labels, dtype=mtp_logits.dtype) out = process_mtp_loss( @@ -1074,7 +978,7 @@ def compute_language_model_loss(mtp_labels, mtp_logits): ) # First chunk is returned unchanged and the loss is never computed. - assert not called["value"] + assert not called['value'] assert torch.equal(out, torch.chunk(hidden_states, 2, dim=0)[0]) def test_process_mtp_loss_derives_labels_from_input_ids(self): @@ -1090,13 +994,13 @@ def test_process_mtp_loss_derives_labels_from_input_ids(self): # hidden_states is chunked into (1 + mtp_num_layers) along dim 0. hidden_states = torch.ones(2, 1, 5) input_ids = torch.tensor([[10, 20, 30, 40, 50]], dtype=torch.long) - seen = {"labels": None, "masked_loss": None} + seen = {'labels': None, 'masked_loss': None} def output_layer(hidden, weight=None, runtime_gather_output=None): return hidden.clone(), None def compute_language_model_loss(mtp_labels, mtp_logits): - seen["labels"] = mtp_labels.clone() + seen['labels'] = mtp_labels.clone() # Per-position loss of 1.0 so loss_mask * loss exposes the active mask. return torch.ones_like(mtp_labels, dtype=torch.float32) @@ -1117,10 +1021,8 @@ def compute_language_model_loss(mtp_labels, mtp_logits): # input_ids rolled twice (once to SFT format, once in the MTP layer loop): # [10,20,30,40,50] -> [20,30,40,50,0] -> [30,40,50,0,0]. - assert seen["labels"] is not None, "loss should be computed in RL mode" - assert torch.equal( - seen["labels"], torch.tensor([[30, 40, 50, 0, 0]], dtype=torch.long) - ) + assert seen['labels'] is not None, "loss should be computed in RL mode" + assert torch.equal(seen['labels'], torch.tensor([[30, 40, 50, 0, 0]], dtype=torch.long)) @pytest.mark.parametrize("cp", [1, 2]) def test_roll_tensor_with_packed_sequences(self, cp): @@ -1129,13 +1031,9 @@ def test_roll_tensor_with_packed_sequences(self, cp): For CP=1: Tests standard packed sequence rolling with verified expected values For CP=2: Tests CP-enabled rolling executes without errors """ - Utils.initialize_model_parallel( - tensor_model_parallel_size=1, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=cp) cp_group = get_context_parallel_group() if cp > 1 else None - cp_rank = ( - torch.distributed.get_rank(group=cp_group) if cp_group is not None else 0 - ) + cp_rank = torch.distributed.get_rank(group=cp_group) if cp_group is not None else 0 if cp == 1: # Test case: Simple packed sequences (CP disabled) @@ -1147,16 +1045,12 @@ def test_roll_tensor_with_packed_sequences(self, cp): cu_seqlens_kv=cu_seqlens, max_seqlen_q=3, max_seqlen_kv=3, - qkv_format="thd", + qkv_format='thd', ) # Roll by -1 (shift left) rolled, sum_val = roll_tensor( - tensor, - shifts=-1, - dims=0, - cp_group=cp_group, - packed_seq_params=packed_seq_params, + tensor, shifts=-1, dims=0, cp_group=cp_group, packed_seq_params=packed_seq_params ) # Expected: [2, 3, 0, 5, 0] - boundaries at indices 2 and 4 are zeroed @@ -1192,25 +1086,21 @@ def test_roll_tensor_with_packed_sequences(self, cp): cu_seqlens_kv=cu_seqlens, max_seqlen_q=6, # max(4, 6) - max local seq length per sequence max_seqlen_kv=6, - qkv_format="thd", + qkv_format='thd', ) # Roll by -1 (shift left) with CP communication rolled, sum_val = roll_tensor( - tensor, - shifts=-1, - dims=0, - cp_group=cp_group, - packed_seq_params=packed_seq_params, + tensor, shifts=-1, dims=0, cp_group=cp_group, packed_seq_params=packed_seq_params ) # Verify the rolled tensor matches expected values - assert rolled.shape == expected.shape, ( - f"Shape mismatch: expected {expected.shape}, got {rolled.shape}" - ) - assert torch.equal(rolled, expected), ( - f"CP Rank {cp_rank}: Expected\n{expected}\nbut got\n{rolled}\nDiff:\n{rolled - expected}" - ) + assert ( + rolled.shape == expected.shape + ), f"Shape mismatch: expected {expected.shape}, got {rolled.shape}" + assert torch.equal( + rolled, expected + ), f"CP Rank {cp_rank}: Expected\n{expected}\nbut got\n{rolled}\nDiff:\n{rolled - expected}" # Verify sum is correct assert sum_val.numel() == 1, "Sum should be a scalar" @@ -1260,22 +1150,10 @@ class DummyOutputLayer: def __init__(self, gather_output): self.gather_output = gather_output - assert ( - _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), None) - is False - ) - assert ( - _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=False), None) - is True - ) - assert ( - _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), True) - is False - ) - assert ( - _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), False) - is True - ) + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), None) is False + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=False), None) is True + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), True) is False + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), False) is True def test_track_mtp_metrics(self): """Test tracking MTP metrics including acceptance rate.""" @@ -1286,11 +1164,7 @@ def test_track_mtp_metrics(self): for i in range(num_layers): MTPLossLoggingHelper.save_metrics_to_tracker( - loss=loss, - correct=correct, - total=total, - layer_number=i, - num_layers=num_layers, + loss=loss, correct=correct, total=total, layer_number=i, num_layers=num_layers ) class DummyWriter: @@ -1321,25 +1195,20 @@ def log(self, metrics, iteration): # Verify loss uses the legacy normalized MTP loss scaled by loss_scale. expected_loss = loss * loss_scale for i in range(num_layers): - assert f"mtp_{i + 1} loss" in writer.scalars - assert torch.isclose( - torch.as_tensor(writer.scalars[f"mtp_{i + 1} loss"]), expected_loss - ) - assert torch.isclose(total_loss_dict[f"mtp_{i + 1} loss"], expected_loss) + assert f"mtp_{i+1} loss" in writer.scalars + assert torch.isclose(torch.as_tensor(writer.scalars[f"mtp_{i+1} loss"]), expected_loss) + assert torch.isclose(total_loss_dict[f"mtp_{i+1} loss"], expected_loss) # Verify acceptance rate is computed as (correct / total) * 100 expected_rate = (correct / total) * 100.0 for i in range(num_layers): - assert f"mtp_{i + 1}_acceptance_rate" in writer.scalars + assert f"mtp_{i+1}_acceptance_rate" in writer.scalars assert torch.isclose( - torch.as_tensor(writer.scalars[f"mtp_{i + 1}_acceptance_rate"]), - expected_rate, + torch.as_tensor(writer.scalars[f"mtp_{i+1}_acceptance_rate"]), expected_rate ) - assert f"mtp_{i + 1}_cumulative_acceptance_rate" in writer.scalars + assert f"mtp_{i+1}_cumulative_acceptance_rate" in writer.scalars assert torch.isclose( - torch.as_tensor( - writer.scalars[f"mtp_{i + 1}_cumulative_acceptance_rate"] - ), + torch.as_tensor(writer.scalars[f"mtp_{i+1}_cumulative_acceptance_rate"]), expected_rate, ) @@ -1366,23 +1235,16 @@ def log(self, metrics, iteration): ) expected_second_rate = (second_correct / second_total) * 100.0 - expected_cumulative_rate = ( - (correct + second_correct) / (total + second_total) - ) * 100.0 + expected_cumulative_rate = ((correct + second_correct) / (total + second_total)) * 100.0 for i in range(num_layers): assert torch.isclose( - torch.as_tensor(writer.scalars[f"mtp_{i + 1}_acceptance_rate"]), - expected_second_rate, + torch.as_tensor(writer.scalars[f"mtp_{i+1}_acceptance_rate"]), expected_second_rate ) assert torch.isclose( - torch.as_tensor( - writer.scalars[f"mtp_{i + 1}_cumulative_acceptance_rate"] - ), + torch.as_tensor(writer.scalars[f"mtp_{i+1}_cumulative_acceptance_rate"]), expected_cumulative_rate, ) - assert torch.isclose( - total_loss_dict[f"mtp_{i + 1} loss"], expected_loss * 2 - ) + assert torch.isclose(total_loss_dict[f"mtp_{i+1} loss"], expected_loss * 2) # Verify tracker is cleaned assert torch.all(MTPLossLoggingHelper.tracker["loss_values"] == 0) @@ -1399,18 +1261,10 @@ def test_track_mtp_loss_preserves_legacy_normalized_loss_semantics(self): layer_number = 0 MTPLossLoggingHelper.save_metrics_to_tracker( - loss=first_loss, - correct=correct, - total=total, - layer_number=layer_number, - num_layers=1, + loss=first_loss, correct=correct, total=total, layer_number=layer_number, num_layers=1 ) MTPLossLoggingHelper.save_metrics_to_tracker( - loss=second_loss, - correct=correct, - total=total, - layer_number=layer_number, - num_layers=1, + loss=second_loss, correct=correct, total=total, layer_number=layer_number, num_layers=1 ) class DummyWriter: @@ -1438,7 +1292,7 @@ class TestMultiTokenPredictionHybrid: def setup_method(self, method): self.seq_length = 32 self.micro_batch_size = 2 - os.environ["CUDA_DEVICE_MAX_CONNECTIONS"] = "1" + os.environ['CUDA_DEVICE_MAX_CONNECTIONS'] = '1' def teardown_method(self, method): Utils.destroy_model_parallel() @@ -1481,7 +1335,7 @@ def create_test_args( destroy_global_vars() destroy_num_microbatches_calculator() - sys.argv = ["test_multi_token_prediction_hybrid.py"] + sys.argv = ['test_multi_token_prediction_hybrid.py'] args = parse_args() args.mtp_num_layers = 2 args.mtp_loss_scaling_factor = 0.1 @@ -1497,9 +1351,9 @@ def create_test_args( args.tensor_model_parallel_size = tp args.sequence_parallel = True if tp > 1 else False args.context_parallel_size = cp - args.position_embedding_type = "rope" + args.position_embedding_type = 'rope' args.train_iters = 1 - args.ckpt_format = "torch_dist" + args.ckpt_format = 'torch_dist' args.lr = 3e-5 args.attention_dropout = 0.0 args.hidden_dropout = 0.0 @@ -1511,10 +1365,10 @@ def create_test_args( args.hybrid_layer_pattern = "M*M*/M*/M*" if fp8 is not None: - args.fp8 = "e4m3" + args.fp8 = 'e4m3' if full_recompute: - args.recompute_granularity = "full" - args.recompute_method = "uniform" + args.recompute_granularity = 'full' + args.recompute_method = 'uniform' args.recompute_num_layers = 1 else: args.recompute_granularity = None @@ -1527,26 +1381,19 @@ def create_test_args( def get_batch(self, seq_length, micro_batch_size): data = list(range(seq_length)) - input_ids = ( - torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - ) - labels = ( - 1 - + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - ) - position_ids = ( - torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() - ) + input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + labels = 1 + torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + position_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() attention_mask = torch.ones( (micro_batch_size, 1, seq_length, seq_length), dtype=bool ).cuda() loss_mask = torch.ones(seq_length).repeat((micro_batch_size, 1)).cuda() batch = { - "tokens": input_ids, - "labels": labels, - "loss_mask": loss_mask, - "attention_mask": attention_mask, - "position_ids": position_ids, + 'tokens': input_ids, + 'labels': labels, + 'loss_mask': loss_mask, + 'attention_mask': attention_mask, + 'position_ids': position_ids, } return batch @@ -1557,9 +1404,7 @@ def test_sharded_state_dict_mamba(self, tp, cp): args = self.create_test_args(tp, cp, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) model_parallel_cuda_manual_seed(_SEED) pg_collection = ProcessGroupCollection.use_mpu_process_groups() @@ -1583,9 +1428,7 @@ def test_forward_backward_mamba(self, tmp_path_dist_ckpt, tp, cp): """Test MTP forward and backward with Mamba hybrid model.""" tp_ref = 1 cp_ref = 1 - args = self.create_test_args( - tp_ref, cp_ref, self.seq_length, self.micro_batch_size - ) + args = self.create_test_args(tp_ref, cp_ref, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) Utils.initialize_model_parallel( @@ -1614,7 +1457,7 @@ def test_forward_backward_mamba(self, tmp_path_dist_ckpt, tp, cp): tracker = MTPLossLoggingHelper.tracker mtp_loss_ref = None assert "loss_values" in tracker - mtp_loss_ref = tracker["loss_values"].clone() + mtp_loss_ref = tracker['loss_values'].clone() MTPLossLoggingHelper.clean_metrics_in_tracker() iteration = 123 @@ -1624,9 +1467,7 @@ def set_ckpt_path(ckpt_path): args.save = ckpt_path args.load = ckpt_path - with TempNamedDir( - tmp_path_dist_ckpt / "test_mtp_mamba_model_reconfiguration" - ) as ckpt_dir: + with TempNamedDir(tmp_path_dist_ckpt / 'test_mtp_mamba_model_reconfiguration') as ckpt_dir: set_ckpt_path(ckpt_dir) save_checkpoint( iteration, @@ -1644,9 +1485,7 @@ def set_ckpt_path(ckpt_path): set_args(args) set_ckpt_path(ckpt_dir) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) model_parallel_cuda_manual_seed(_SEED) cfg_container = Utils.pretrain_config_from_global_args(args, "hybrid") @@ -1663,9 +1502,7 @@ def set_ckpt_path(ckpt_path): batch = get_batch_on_this_cp_rank( batch, is_hybrid_cp=False, cp_group=get_context_parallel_group() ) - tokens, labels, loss_mask, attention_mask, position_ids, output_ref = ( - batch.values() - ) + tokens, labels, loss_mask, attention_mask, position_ids, output_ref = batch.values() output = mamba_model[0].forward( input_ids=tokens, position_ids=position_ids, @@ -1675,10 +1512,8 @@ def set_ckpt_path(ckpt_path): ) tracker = MTPLossLoggingHelper.tracker assert "loss_values" in tracker - mtp_loss = tracker["loss_values"].clone() - pg_collection = ProcessGroupCollection.use_mpu_process_groups( - required_pgs=["cp"] - ) + mtp_loss = tracker['loss_values'].clone() + pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp']) torch.distributed.all_reduce( mtp_loss, group=pg_collection.cp, op=torch.distributed.ReduceOp.AVG ) @@ -1702,9 +1537,7 @@ def test_attention_mask_validation_mamba(self): args = self.create_test_args(tp, cp, self.seq_length, self.micro_batch_size) set_args(args) torch.manual_seed(_SEED) - Utils.initialize_model_parallel( - tensor_model_parallel_size=tp, context_parallel_size=cp - ) + Utils.initialize_model_parallel(tensor_model_parallel_size=tp, context_parallel_size=cp) pg_collection = ProcessGroupCollection.use_mpu_process_groups() model_cfg = hybrid_config_from_args(args) builder_cls = model_cfg.get_builder_cls() @@ -1719,8 +1552,6 @@ def test_attention_mask_validation_mamba(self): assert mamba_model[0].mtp is not None except AssertionError as e: if "Multi-Token Prediction (MTP) is not yet supported" in str(e): - pytest.fail( - f"Attention mask validation failed for Mamba hybrid model: {e}" - ) + pytest.fail(f"Attention mask validation failed for Mamba hybrid model: {e}") else: raise diff --git a/tests/unit_tests/transformer/test_torch_norm.py b/tests/unit_tests/transformer/test_torch_norm.py index 8951eb2f365..34d37ae4662 100644 --- a/tests/unit_tests/transformer/test_torch_norm.py +++ b/tests/unit_tests/transformer/test_torch_norm.py @@ -1,5 +1,6 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +import pytest import torch from megatron.core.transformer.torch_norm import WrappedTorchNorm @@ -23,3 +24,17 @@ def test_rmsnorm_uses_native_torch_implementation(): assert isinstance(norm, torch.nn.RMSNorm) assert norm.weight.dtype == torch.bfloat16 + + +def test_sequence_parallel_is_opt_in_for_native_norm(): + with pytest.raises(AssertionError, match="sequence parallel"): + WrappedTorchNorm( + config=_config(sequence_parallel=True, tensor_model_parallel_size=2), + hidden_size=64, + eps=1e-5, + ) + config = _config( + sequence_parallel=True, tensor_model_parallel_size=2, norm_accuracy_compatible=True + ) + norm = WrappedTorchNorm(config=config, hidden_size=64, eps=1e-5) + assert all(parameter.sequence_parallel for parameter in norm.parameters()) From 76dd11a1f4887ef5c3735153003a9890d08d3904 Mon Sep 17 00:00:00 2001 From: Zhan Rongrui Date: Mon, 14 Sep 2026 16:46:14 +0800 Subject: [PATCH 27/27] fix: preserve upstream accuracy MoE gradient branches --- megatron/core/transformer/moe/moe_layer.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index ba9b2cdfd6c..798eb4cabd6 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -12,7 +12,7 @@ from megatron.core.extensions.transformer_engine import HAVE_TE from megatron.core.inference.utils import InferenceMode from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.transformer.module import MegatronModule +from megatron.core.transformer.module import MegatronModule, _use_accuracy_compatible from megatron.core.transformer.moe.moe_utils import ( MoECudaGraphPartialCaptureSignal, MoECudaGraphTensorStore, @@ -651,7 +651,7 @@ def custom_forward(hidden_states, intermediate_tensors=None, padding_mask=None): # logging probe: removing the nodes changes bf16 gradient sum # order at the shared input, and makes the router input grad a # 3-way accumulated value that PF's ThreePathCloneAlignMG splits. - if self.config.dsa_accuracy_compatible and hidden_states.requires_grad: + if _use_accuracy_compatible() and hidden_states.requires_grad: _hs_router_path_mg = hidden_states.clone() _hs_dispatcher_path_mg = hidden_states.clone() _hs_shared_path_mg = hidden_states.clone()