diff --git a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py index f658538248b..26528942bea 100644 --- a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py +++ b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py @@ -669,8 +669,13 @@ def qkv_up_proj_and_rope_apply(q_compressed, kv_compressed, k_pos_emb, rotary_po return q_absorbed, kv_compressed if self.recompute_up_proj: + # Quantized replay is safe here for the same reason as in MLASelfAttention: + # CheckpointWithoutOutput records the forward recipe/amax state and replays + # under the recorded fp8_autocast, and the only quantized op inside + # qkv_up_proj_and_rope_apply is the Q up projection. The absorption einsum + # reads the K up-projection weight directly, which is a persistent parameter + # and therefore identical between the forward pass and the replay. quantization = self.config.fp8 or self.config.fp4 - assert not quantization, "FP8/FP4 is not supported for AbsorbedMLA" self.qkv_up_checkpoint = tensor_parallel.CheckpointWithoutOutput(fp8=quantization) q_absorbed, kv_compressed = self.qkv_up_checkpoint.checkpoint( qkv_up_proj_and_rope_apply, q_compressed, kv_compressed, k_pos_emb, rotary_pos_emb diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py b/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py index fc1778f649f..d6d98910e41 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py @@ -4,12 +4,16 @@ from types import SimpleNamespace from typing import List, Optional, Tuple +import numpy as np import pytest import torch import torch.distributed as dist +from transformer_engine.pytorch.fp8 import check_fp8_support, check_nvfp4_support from megatron.core import parallel_state from megatron.core.extensions.transformer_engine_spec_provider import TESpecProvider +from megatron.core.fp4_utils import get_fp4_context +from megatron.core.fp8_utils import get_fp8_context from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnMaskType @@ -29,6 +33,37 @@ from megatron.core.utils import init_method_normal, scaled_init_method_normal from tests.unit_tests.test_utilities import Utils +fp8_available, reason_for_no_fp8 = check_fp8_support() +nvfp4_available, reason_for_no_nvfp4 = check_nvfp4_support() + + +# Inlined from tests.unit_tests.determinism.utils rather than imported: that module pulls +# in torch.testing._internal.common_utils, whose import calls +# torch.backends.disable_global_flags() and breaks later tests in this pytest session that +# assign torch.backends.* flags (e.g. test_te_layers_batch_invariant.py). +def capture_rng_state() -> dict: + """Snapshot every RNG that the framework consumes during a fwd+bwd pass.""" + from megatron.core.tensor_parallel.random import get_cuda_rng_tracker + + return { + "random": random.getstate(), + "numpy": np.random.get_state(), + "torch_cpu": torch.get_rng_state(), + "torch_cuda": torch.cuda.get_rng_state(), + "mpu_tracker": get_cuda_rng_tracker().get_states(), + } + + +def restore_rng_state(state: dict) -> None: + """Inverse of ``capture_rng_state``.""" + from megatron.core.tensor_parallel.random import get_cuda_rng_tracker + + random.setstate(state["random"]) + np.random.set_state(state["numpy"]) + torch.set_rng_state(state["torch_cpu"]) + torch.cuda.set_rng_state(state["torch_cuda"]) + get_cuda_rng_tracker().set_states(state["mpu_tracker"]) + class MockCoreAttention(torch.nn.Module): """Mock core attention for testing MLA computation flow.""" @@ -125,7 +160,14 @@ def _forward_thd(self, q, k, v, packed_seq_params): def get_mock_mla_config( - tensor_model_parallel_size: int, context_parallel_size: int, qk_layernorm: bool + tensor_model_parallel_size: int, + context_parallel_size: int, + qk_layernorm: bool, + recompute_mla_up_proj: bool = False, + fp8: Optional[str] = None, + fp8_recipe: str = "delayed", + fp4: Optional[str] = None, + fp4_recipe: str = "nvfp4", ) -> MLATransformerConfig: """Create test config with all attributes used in MLA.""" return MLATransformerConfig( @@ -158,11 +200,14 @@ def get_mock_mla_config( beta_fast=32, beta_slow=1, rotary_interleaved=False, - recompute_granularity=None, + recompute_granularity="selective" if recompute_mla_up_proj else None, + recompute_modules=["mla_up_proj"] if recompute_mla_up_proj else [], fine_grained_activation_offloading=False, gradient_accumulation_fusion=False, - fp8=False, - fp4=False, + fp8=fp8 if fp8 else False, + fp8_recipe=fp8_recipe, + fp4=fp4, + fp4_recipe=fp4_recipe, init_method=init_method_normal(0.02), output_layer_init_method=scaled_init_method_normal(0.02, 61, multiplier=2.0), kv_channels=56, @@ -528,3 +573,174 @@ def _calculate_tensor_similarity(x, y): assert _calculate_tensor_similarity(absorbed_grad, standard_grad) > 0.9999 Utils.destroy_model_parallel() + + +# Hopper = SM 9.0, Blackwell = SM 10.0+. mxfp8 needs Blackwell. +_IS_BLACKWELL = torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 10 + +_QUANT_RECIPES = [ + pytest.param( + {"fp8": "hybrid", "fp8_recipe": "tensorwise"}, + id="fp8-tensorwise", + marks=pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8), + ), + pytest.param( + {"fp8": "hybrid", "fp8_recipe": "mxfp8"}, + id="fp8-mxfp8", + marks=pytest.mark.skipif(not fp8_available or not _IS_BLACKWELL, reason="needs Blackwell"), + ), + pytest.param( + {"fp8": "hybrid", "fp8_recipe": "blockwise"}, + id="fp8-blockwise", + marks=pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8), + ), + pytest.param( + {"fp4": "e2m1", "fp4_recipe": "nvfp4"}, + id="fp4-nvfp4", + marks=pytest.mark.skipif(not nvfp4_available, reason=reason_for_no_nvfp4), + ), +] + + +def _build_absorbed_mla(config, state_dict=None): + """Build an AbsorbedMLASelfAttention, optionally loading a shared state dict.""" + model = AbsorbedMLASelfAttention( + config=config, + submodules=get_absorbed_mla_submodules( + down_proj_use_column_parallel=False, qk_layernorm=True, rms_norm=True + ), + layer_number=0, + attn_mask_type=AttnMaskType.causal, + cp_comm_type=None, + pg_collection=None, + ).cuda() + if state_dict is not None: + model.load_state_dict(state_dict) + return model + + +@pytest.mark.parametrize("quant_overrides", _QUANT_RECIPES) +@pytest.mark.parametrize("qkv_format", ['sbhd', 'thd']) +def test_quantized_up_proj_recompute_parity(quant_overrides: dict, qkv_format: str): + """`mla_up_proj` recompute must be an exact replay under FP8 and FP4. + + The absorbed up-projection recompute used to be blocked for FP8/FP4. It is + allowed now, so pin the contract: running the same inputs through two + identically-initialized modules — one recomputing the up projection, one not — + must produce bitwise-identical outputs and parameter gradients. Any drift means + the replay used a different quantization scale or a stale weight. + + A bf16 reference run guards against the parity check passing vacuously, that is, + both modules silently falling back to bf16. + """ + Utils.initialize_model_parallel(tensor_model_parallel_size=1, context_parallel_size=1) + model_parallel_cuda_manual_seed(123) + + def make_config(recompute: bool) -> MLATransformerConfig: + return get_mock_mla_config( + tensor_model_parallel_size=1, + context_parallel_size=1, + qk_layernorm=True, + recompute_mla_up_proj=recompute, + **quant_overrides, + ) + + baseline = _build_absorbed_mla(make_config(False)) + recomputed = _build_absorbed_mla(make_config(True), state_dict=baseline.state_dict()) + assert recomputed.recompute_up_proj and not baseline.recompute_up_proj + + # Quantized GEMMs need the token dimension aligned; keep every sequence a multiple + # of 128, which satisfies both the FP8 (16) and FP4 (32) alignment requirements. + if qkv_format == 'thd': + random.seed(42) + seqlens = [random.randint(1, 8) * 128 for _ in range(3)] + cu_seqlens = [0] + for length in seqlens: + cu_seqlens.append(cu_seqlens[-1] + length) + total_tokens = cu_seqlens[-1] + packed_seq_params = PackedSeqParams( + cu_seqlens_q=torch.IntTensor(cu_seqlens).cuda(), + cu_seqlens_q_padded=torch.IntTensor(cu_seqlens).cuda(), + cu_seqlens_kv=torch.IntTensor(cu_seqlens).cuda(), + cu_seqlens_kv_padded=torch.IntTensor(cu_seqlens).cuda(), + max_seqlen_q=max(seqlens), + max_seqlen_kv=max(seqlens), + qkv_format='thd', + ) + hidden_states = torch.randn( + (total_tokens, 1, baseline.config.hidden_size), dtype=torch.bfloat16, device='cuda' + ) + else: + packed_seq_params = None + hidden_states = torch.randn( + (1024, 2, baseline.config.hidden_size), dtype=torch.bfloat16, device='cuda' + ) + grads = torch.randn_like(hidden_states) + + def quantization_context(config): + # Mirrors the dispatch in TransformerBlock. + return get_fp8_context(config) if config.fp8 else get_fp4_context(config) + + def run(model): + with quantization_context(model.config): + output, _ = model( + hidden_states, attention_mask=None, packed_seq_params=packed_seq_params + ) + output.backward(grads) + return output, {name: param.grad for name, param in model.named_parameters()} + + # NVFP4 draws randomness in the backward (stochastic rounding), so the two runs have + # to start from the same RNG state to be comparable at all. Same ritual as + # BitExactRunner._two_runs. + rng_state = capture_rng_state() + baseline_output, baseline_grads = run(baseline) + restore_rng_state(rng_state) + recomputed_output, recomputed_grads = run(recomputed) + + torch.testing.assert_close(recomputed_output, baseline_output, atol=0, rtol=0) + + assert recomputed_grads.keys() == baseline_grads.keys() + for name, baseline_grad in baseline_grads.items(): + assert baseline_grad is not None, f"{name} has no gradient" + recomputed_grad = recomputed_grads[name] + assert recomputed_grad is not None, f"{name} has no gradient with up-proj recompute" + torch.testing.assert_close( + recomputed_grad, baseline_grad, atol=0, rtol=0, msg=lambda m, n=name: f"{n}: {m}" + ) + + # The assertions above hold trivially if the recipe never engaged, since two bf16 + # replays also match bitwise. Run the same weights through a bf16 module: the + # quantized output must track it closely (the math is right) yet differ from it + # (the values really went through quantization). + bf16_model = _build_absorbed_mla( + get_mock_mla_config( + tensor_model_parallel_size=1, context_parallel_size=1, qk_layernorm=True + ) + ) + # Copy parameters directly rather than via state_dict, which also carries the + # quantization `_extra_state` blobs that a bf16 module has no use for. + baseline_params = dict(baseline.named_parameters()) + with torch.no_grad(): + for name, param in bf16_model.named_parameters(): + param.copy_(baseline_params[name]) + bf16_output, _ = bf16_model( + hidden_states, attention_mask=None, packed_seq_params=packed_seq_params + ) + + assert not torch.equal( + recomputed_output, bf16_output + ), "quantized output is bitwise equal to bf16, so the recipe never engaged" + # The bound stays loose on purpose: it only has to catch a quantized path that + # produces garbage, and it must hold across recipes and architectures whose + # quantization error differs by an order of magnitude — hence the separate, coarser + # bound for 4-bit elements. Accuracy itself is pinned by the exact comparisons above + # and by test_functionality. + threshold = 0.9 if quant_overrides.get("fp4") else 0.99 + cosine_sim = torch.nn.functional.cosine_similarity( + recomputed_output.flatten().float().unsqueeze(0), bf16_output.flatten().float().unsqueeze(0) + ).item() + assert ( + cosine_sim > threshold + ), f"{'FP8' if quant_overrides.get('fp8') else 'FP4'} quantized output diverges from bf16: cosine similarity = {cosine_sim}" + + Utils.destroy_model_parallel()