From 04f50389097a3ae5825f49d81050ebc0c3f96604 Mon Sep 17 00:00:00 2001 From: brian Date: Mon, 20 Jul 2026 15:54:19 +0800 Subject: [PATCH] fix(docker): slice output gate for GQA TP --- docker/patch/latest/megatron.patch | 17 +++++++++++++++++ tests/test_qwen3.5_0.8B_gsm8k_short.py | 3 ++- 2 files changed, 19 insertions(+), 1 deletion(-) diff --git a/docker/patch/latest/megatron.patch b/docker/patch/latest/megatron.patch index 4b4c74d28e..e278962eca 100644 --- a/docker/patch/latest/megatron.patch +++ b/docker/patch/latest/megatron.patch @@ -731,6 +731,23 @@ index ac839c21f..f18309217 100644 ) ops.append(recv_next_op) if len(ops) > 0: +diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py +index bc5e4e2..739450c 100644 +--- a/megatron/core/transformer/attention.py ++++ b/megatron/core/transformer/attention.py +@@ -1491,4 +1491,12 @@ class SelfAttention(Attention): + 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 diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index 75825cd37..445b3fb84 100644 --- a/megatron/core/transformer/moe/moe_utils.py diff --git a/tests/test_qwen3.5_0.8B_gsm8k_short.py b/tests/test_qwen3.5_0.8B_gsm8k_short.py index 93306eadd3..05aff114cc 100644 --- a/tests/test_qwen3.5_0.8B_gsm8k_short.py +++ b/tests/test_qwen3.5_0.8B_gsm8k_short.py @@ -51,8 +51,9 @@ def execute(): "--eval-top-k 1 " ) + # Qwen3.5-0.8B has 2 query groups; TP=4 covers gated attention with query groups < TP. perf_args = ( - "--tensor-model-parallel-size 1 " + "--tensor-model-parallel-size 4 " "--sequence-parallel " "--pipeline-model-parallel-size 1 " "--context-parallel-size 1 "