From d4f63dfdb6b40e331264cfcf0255982316dd3143 Mon Sep 17 00:00:00 2001 From: xiaotaoliu Date: Sat, 21 Feb 2026 12:04:41 +0800 Subject: [PATCH 1/5] fixed the bug where gate was not sliced when kv_head < tp_size. --- megatron/core/transformer/attention.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index 28e3dde01c4..b5d3add9a54 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -1498,6 +1498,14 @@ def get_query_key_value_tensors( if output_gate: # Gate [sq, b, ng, np/ng * hn] -> [sq, b, np, hn] gate = gate.reshape(*gate.shape[:2], -1, self.hidden_size_per_attention_head) + if self.config.num_query_groups < self.world_size: + idx = get_tensor_model_parallel_rank() % ( + self.world_size // self.config.num_query_groups + ) + size = self.num_attention_heads_per_partition // ( + self.world_size // self.config.num_query_groups + ) + gate = gate[:, :, idx * size : (idx + 1) * size, :] return query, key, value, gate return query, key, value From 984355cd7823b0c1d07d1e33aeccb960fffc024e Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Mon, 23 Feb 2026 23:38:28 -0800 Subject: [PATCH 2/5] add UTs --- .../unit_tests/transformer/test_attention.py | 260 ++++++++++++++++++ 1 file changed, 260 insertions(+) diff --git a/tests/unit_tests/transformer/test_attention.py b/tests/unit_tests/transformer/test_attention.py index 704cc78d004..2e06ceb300f 100644 --- a/tests/unit_tests/transformer/test_attention.py +++ b/tests/unit_tests/transformer/test_attention.py @@ -449,3 +449,263 @@ def test_self_attention_independent_pg_smoke(self): pg_collection = ProcessGroupCollection(tp=tp_group, cp=cp_group) self.run_self_attention(pg_collection) + + +def _test_parallel_attention_correctness( + transformer_config, + transformer_layer_spec, + tmp_path_dist_ckpt, + atol, + rtol, + tp=1, + sp=False, + cp=1, + seed=123, + sequence_length=256, + micro_batch_size=4, +): + # Model initialization function + def initialize_gpt_model( + config, pre_process=True, post_process=True, vp_stage=None, pg_collection=None + ): + gpt_model = GPTModel( + config=config, + transformer_layer_spec=transformer_layer_spec, + vocab_size=128, + max_sequence_length=sequence_length, + pre_process=pre_process, + post_process=post_process, + vp_stage=vp_stage, + pg_collection=pg_collection, + ) + return gpt_model + + # Initialize baseline parallel state + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, context_parallel_size=1 + ) + + # Initialize input hidden states + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + input_hidden_states = ( + torch.rand((sequence_length, micro_batch_size, transformer_config.hidden_size)) + .cuda() + .bfloat16() + .requires_grad_(True) + ) + + with TempNamedDir(tmp_path_dist_ckpt / 'test_parallel_attn', sync=True) as ckpt_dir: + # Set argument + mock_args = parse_args(ignore_unknown_args=True) + set_args(mock_args) + + # Initialize baseline model + init_basic_mock_args(mock_args, 1, 1, bf16=True) + mock_args.context_parallel_size = 1 + mock_args.sequence_parallel = 1 + gpt_model = unwrap_model(get_model(initialize_gpt_model, config=transformer_config)) + + # Initialize args and save checkpoint + init_checkpointing_mock_args(mock_args, ckpt_dir, False) + mock_args.no_save_optim = True + mock_args.no_save_rng = True + mock_args.no_load_optim = True + mock_args.no_load_rng = True + save_checkpoint(10, gpt_model, None, None, 0) + + # Calculate baseline output + attention = gpt_model[0].decoder.layers[0].self_attention + output_hidden_states_baseline, bias_hidden_states_baseline = attention( + input_hidden_states, attention_mask=None + ) + output_hidden_states_baseline.sum().backward() + + # Save baseline output + input_grad_baseline = input_hidden_states.grad.detach() + output_hidden_states_baseline = output_hidden_states_baseline.detach() + bias_hidden_states_baseline = bias_hidden_states_baseline + if bias_hidden_states_baseline is not None: + bias_hidden_states_baseline = bias_hidden_states_baseline.detach() + has_bias = True + else: + has_bias = False + + # Initialize parallel model + Utils.destroy_model_parallel() + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp, pipeline_model_parallel_size=1, context_parallel_size=cp + ) + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + transformer_config.context_parallel_size = cp + transformer_config.tensor_model_parallel_size = tp + transformer_config.sequence_parallel = sp + init_basic_mock_args(mock_args, tp, 1, bf16=True) + mock_args.context_parallel_size = cp + mock_args.sequence_parallel = sp + gpt_model = unwrap_model(get_model(initialize_gpt_model, config=transformer_config)) + with mock.patch('megatron.training.checkpointing.check_checkpoint_args'): + with mock.patch('megatron.training.checkpointing.update_num_microbatches'): + load_checkpoint(gpt_model, None, None) + + # Function to get tensor on this tp and cp rank + cp_group = parallel_state.get_context_parallel_group() + tp_rank = parallel_state.get_tensor_model_parallel_rank() + + def get_tensor_on_this_rank(tensor): + if cp > 1: + tensor = get_tensor_on_this_cp_rank(tensor, 0, cp_group) + if tp > 1 and sp: + sp_seg = sequence_length // tp // cp + tensor = tensor[tp_rank * sp_seg : (tp_rank + 1) * sp_seg] + return tensor + + # Calculate parallel model output + input_hidden_states = get_tensor_on_this_rank(input_hidden_states) + input_hidden_states = input_hidden_states.detach().requires_grad_(True) + parallel_attention = gpt_model[0].decoder.layers[0].self_attention + output_hidden_states_parallel, bias_hidden_states_parallel = parallel_attention( + input_hidden_states, attention_mask=None + ) + output_hidden_states_parallel.sum().backward() + input_grad_parallel = input_hidden_states.grad.detach() + + # Check if the output is close + output_hidden_states_baseline = get_tensor_on_this_rank(output_hidden_states_baseline) + input_grad_baseline = get_tensor_on_this_rank(input_grad_baseline) + + assert torch.all( + ~torch.isnan(output_hidden_states_baseline) + ), "output_hidden_states_baseline contains nan" + assert torch.all( + ~torch.isinf(output_hidden_states_baseline) + ), "output_hidden_states_baseline contains inf" + assert torch.all(~torch.isnan(input_grad_baseline)), "input_grad_baseline contains nan" + assert torch.all(~torch.isinf(input_grad_baseline)), "input_grad_baseline contains inf" + assert torch.all( + ~torch.isnan(output_hidden_states_parallel) + ), "output_hidden_states_parallel contains nan" + assert torch.all( + ~torch.isinf(output_hidden_states_parallel) + ), "output_hidden_states_parallel contains inf" + assert torch.all(~torch.isnan(input_grad_parallel)), "input_grad_parallel contains nan" + assert torch.all(~torch.isinf(input_grad_parallel)), "input_grad_parallel contains inf" + if has_bias: + assert torch.all( + ~torch.isnan(bias_hidden_states_baseline) + ), "bias_hidden_states_baseline contains nan" + assert torch.all( + ~torch.isinf(bias_hidden_states_baseline) + ), "bias_hidden_states_baseline contains inf" + assert torch.all( + ~torch.isnan(bias_hidden_states_parallel) + ), "bias_hidden_states_parallel contains nan" + assert torch.all( + ~torch.isinf(bias_hidden_states_parallel) + ), "bias_hidden_states_parallel contains inf" + + torch.testing.assert_close( + output_hidden_states_baseline, + output_hidden_states_parallel, + atol=atol, + rtol=rtol, + msg=lambda msg: f"Mismatch in output_hidden_states: {msg}", + ) + torch.testing.assert_close( + input_grad_baseline, + input_grad_parallel, + atol=atol, + rtol=rtol, + msg=lambda msg: f"Mismatch in input_grad: {msg}", + ) + if has_bias: + torch.testing.assert_close( + bias_hidden_states_baseline, + bias_hidden_states_parallel, + atol=atol, + rtol=rtol, + msg=lambda msg: f"Mismatch in bias_hidden_states: {msg}", + ) + + Utils.destroy_model_parallel() + + +@pytest.mark.parametrize("apply_rope_fusion", [False, True]) +@pytest.mark.parametrize( + ("tp", "sp", "cp"), + [ + (4, False, 1), # TP w/o SP + (4, True, 1), # TP w/ SP + (1, False, 4), # CP + (2, False, 2), # CP + TP w/o SP + (2, True, 2), # CP + TP w/ SP + ], +) +@pytest.mark.parametrize("qk_layernorm", [False, True]) +@pytest.mark.parametrize("output_gate", [False, True]) +def test_parallel_attention_correctness( + tmp_path_dist_ckpt, apply_rope_fusion, tp, sp, cp, qk_layernorm, output_gate +): + transformer_config = TransformerConfig( + num_layers=1, + hidden_size=128, + num_attention_heads=4, + normalization="RMSNorm", + bf16=True, + qk_layernorm=qk_layernorm, + apply_rope_fusion=apply_rope_fusion, + attention_output_gate=output_gate, + hidden_dropout=0.0, + attention_dropout=0.0, + ) + + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec(qk_layernorm=qk_layernorm) + atol, rtol = 1e-2, 1e-2 + + _test_parallel_attention_correctness( + transformer_config, + transformer_layer_spec, + tmp_path_dist_ckpt, + atol=atol, + rtol=rtol, + tp=tp, + sp=sp, + cp=cp, + seed=123, + sequence_length=256, + ) + + +@pytest.mark.parametrize("sp", [True, False]) +@pytest.mark.parametrize("output_gate", [False, True]) +def test_parallel_attention_correctness_num_query_groups_less_than_tp_size( + tmp_path_dist_ckpt, sp, output_gate +): + transformer_config = TransformerConfig( + num_layers=1, + hidden_size=128, + num_attention_heads=8, + num_query_groups=2, + normalization="RMSNorm", + bf16=True, + attention_output_gate=output_gate, + hidden_dropout=0.0, + attention_dropout=0.0, + ) + + transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec() + atol, rtol = 1e-2, 1e-2 + + _test_parallel_attention_correctness( + transformer_config, + transformer_layer_spec, + tmp_path_dist_ckpt, + atol=atol, + rtol=rtol, + tp=4, + sp=sp, + seed=123, + sequence_length=256, + ) + From c217222dcff6f376eff3ff89111090959a745a20 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Fri, 27 Feb 2026 00:39:32 -0800 Subject: [PATCH 3/5] fix lint --- tests/unit_tests/transformer/test_attention.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/unit_tests/transformer/test_attention.py b/tests/unit_tests/transformer/test_attention.py index 2e06ceb300f..e0abb0e6125 100644 --- a/tests/unit_tests/transformer/test_attention.py +++ b/tests/unit_tests/transformer/test_attention.py @@ -708,4 +708,3 @@ def test_parallel_attention_correctness_num_query_groups_less_than_tp_size( seed=123, sequence_length=256, ) - From 719f6158a96bfcf26559e48f591b35ca139aad51 Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Fri, 27 Feb 2026 01:50:35 -0800 Subject: [PATCH 4/5] fix ut --- tests/unit_tests/transformer/test_attention.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/unit_tests/transformer/test_attention.py b/tests/unit_tests/transformer/test_attention.py index e0abb0e6125..c1fe11621ff 100644 --- a/tests/unit_tests/transformer/test_attention.py +++ b/tests/unit_tests/transformer/test_attention.py @@ -9,6 +9,7 @@ import megatron.core.parallel_state as parallel_state from megatron.core.hyper_comm_grid import HyperCommGrid from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_with_transformer_engine_spec, get_gpt_layer_with_transformer_engine_submodules, ) from megatron.core.process_groups_config import ProcessGroupCollection From 5f0403795efff5b9be901f1bc9a226fe96dfaf9b Mon Sep 17 00:00:00 2001 From: Yuzhong Wang Date: Sun, 1 Mar 2026 21:44:24 -0800 Subject: [PATCH 5/5] fix import --- tests/unit_tests/transformer/test_attention.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/tests/unit_tests/transformer/test_attention.py b/tests/unit_tests/transformer/test_attention.py index c1fe11621ff..15d95128bac 100644 --- a/tests/unit_tests/transformer/test_attention.py +++ b/tests/unit_tests/transformer/test_attention.py @@ -1,6 +1,7 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. import copy +from unittest import mock import pytest import torch @@ -8,16 +9,30 @@ import megatron.core.parallel_state as parallel_state from megatron.core.hyper_comm_grid import HyperCommGrid +from megatron.core.models.common.embeddings.rope_utils import ( + get_pos_emb_on_this_cp_rank as get_tensor_on_this_cp_rank, +) from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_with_transformer_engine_spec, get_gpt_layer_with_transformer_engine_submodules, ) +from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig from megatron.core.transformer.attention import SelfAttention from megatron.core.transformer.enums import AttnMaskType from megatron.core.utils import is_te_min_version +from megatron.training.arguments import parse_args +from megatron.training.checkpointing import load_checkpoint, save_checkpoint +from megatron.training.global_vars import set_args +from megatron.training.training import get_model +from megatron.training.utils import unwrap_model +from tests.unit_tests.dist_checkpointing import ( + TempNamedDir, + init_basic_mock_args, + init_checkpointing_mock_args, +) from tests.unit_tests.test_utilities import Utils try: