From f338c9b034f29774c833ad8a7107e5fdfa9f7ec7 Mon Sep 17 00:00:00 2001 From: BANMA-00403 Date: Tue, 1 Sep 2026 09:53:35 +0800 Subject: [PATCH 1/3] fix: support QSA with sequence parallelism --- src/mcore_bridge/model/gpts/qwen4_exp.py | 2 +- src/mcore_bridge/model/modules/qsa_indexer.py | 19 +++++++++++++++---- 2 files changed, 16 insertions(+), 5 deletions(-) diff --git a/src/mcore_bridge/model/gpts/qwen4_exp.py b/src/mcore_bridge/model/gpts/qwen4_exp.py index 7ce61f8..077ccef 100644 --- a/src/mcore_bridge/model/gpts/qwen4_exp.py +++ b/src/mcore_bridge/model/gpts/qwen4_exp.py @@ -58,7 +58,7 @@ def __init__(self, config, submodules, layer_number: int = 1, **kwargs): config, config.ple_layer_ids.index(self.layer_number), pg_collection=self.pg_collection) is_linear_attention = config.linear_attention_freq[self.layer_number - 1] if not is_linear_attention and getattr(config, 'indexer_n_heads', None) is not None: - self.self_attention.indexer = QSAIndexer(config) + self.self_attention.indexer = QSAIndexer(config, tp_group=self.tp_group) self.attn_hyper_connection = Qwen4ExpTextGatedResidual(config) self.mlp_hyper_connection = Qwen4ExpTextGatedResidual(config) diff --git a/src/mcore_bridge/model/modules/qsa_indexer.py b/src/mcore_bridge/model/modules/qsa_indexer.py index 193bcbb..5a2c706 100644 --- a/src/mcore_bridge/model/modules/qsa_indexer.py +++ b/src/mcore_bridge/model/modules/qsa_indexer.py @@ -1,6 +1,7 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import math import torch +from megatron.core.tensor_parallel import gather_from_sequence_parallel_region from torch import nn @@ -28,9 +29,10 @@ def _rotate_half(x: torch.Tensor) -> torch.Tensor: class QSAIndexer(nn.Module): # refer: transformers Qwen4ExpTextQSAIndexer - def __init__(self, config): + def __init__(self, config, tp_group=None): super().__init__() self.config = config + self.tp_group = tp_group self.index_n_heads = config.indexer_n_heads self.index_kv_heads = config.indexer_kv_heads self.index_head_dim = config.indexer_head_dim @@ -63,7 +65,8 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch Args: hidden_states: ``[s, b, h]`` (mcore layout), pre-attention input -- - the same tensor the reference indexer consumes. + the same tensor the reference indexer consumes. ``s`` is the + local sequence length when sequence parallelism is enabled. freqs: mcore rotary frequencies ``[s, 1, 1, rot_dim]``. mcore stores angles rather than cos/sin, so they are materialized here the way ``_patch_apply_rotary_pos_emb`` does (``cos(freqs) * mscale``), @@ -78,7 +81,9 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch Only the causal, unpacked layout is handled; callers must not invoke this for packed/THD or context-parallel inputs (see the guard in the layer). """ - s, b, _ = hidden_states.shape + local_s, b, _ = hidden_states.shape + tp_size = self.config.tensor_model_parallel_size if self.config.sequence_parallel else 1 + s = local_s * tp_size R = self.compress_ratio max_blocks = s // R # Selection is a no-op while the causal prefix never exceeds the budget: @@ -90,8 +95,14 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch device = hidden_states.device # ---- project to indexer q/k ---- - # [s, b, h] -> [s, b, (nh + nkv) * d] + # Project the local SP shard first, then gather only the compact indexer + # activations. For Qwen3.8-Flash-Next this communicates 640 values/token + # instead of gathering all 2560 hidden values/token. + # [local_s, b, h] -> [local_s, b, (nh + nkv) * d] qk = self.index_qk_proj(hidden_states) + if tp_size > 1: + qk = gather_from_sequence_parallel_region(qk, tensor_parallel_output_grad=False, group=self.tp_group) + assert qk.shape[0] == s, f'QSA indexer expected sequence length {s}, got {qk.shape[0]}' q, token_k = torch.split( qk, [self.index_n_heads * self.index_head_dim, self.index_kv_heads * self.index_head_dim], dim=-1) # -> [b, s, nh, d] / [b, s, d] From 069e57d050c7e2974ddca5d4dffd4f7976dc769d Mon Sep 17 00:00:00 2001 From: BANMA-00403 Date: Tue, 1 Sep 2026 10:16:02 +0800 Subject: [PATCH 2/3] fix: derive QSA sequence length from TP group --- src/mcore_bridge/model/modules/qsa_indexer.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/src/mcore_bridge/model/modules/qsa_indexer.py b/src/mcore_bridge/model/modules/qsa_indexer.py index 5a2c706..51c51fb 100644 --- a/src/mcore_bridge/model/modules/qsa_indexer.py +++ b/src/mcore_bridge/model/modules/qsa_indexer.py @@ -67,10 +67,11 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch hidden_states: ``[s, b, h]`` (mcore layout), pre-attention input -- the same tensor the reference indexer consumes. ``s`` is the local sequence length when sequence parallelism is enabled. - freqs: mcore rotary frequencies ``[s, 1, 1, rot_dim]``. mcore stores - angles rather than cos/sin, so they are materialized here the way - ``_patch_apply_rotary_pos_emb`` does (``cos(freqs) * mscale``), - keeping the indexer's RoPE identical to the attention's. + freqs: Full-sequence mcore rotary frequencies + ``[S, 1, 1, rot_dim]``. mcore stores angles rather than cos/sin, + so they are materialized here the way ``_patch_apply_rotary_pos_emb`` + does (``cos(freqs) * mscale``), keeping the indexer's RoPE identical + to the attention's. Returns: ``[b, 1, s, s]`` bool mask where True marks a *masked-out* key (TE's @@ -82,7 +83,12 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch for packed/THD or context-parallel inputs (see the guard in the layer). """ local_s, b, _ = hidden_states.shape - tp_size = self.config.tensor_model_parallel_size if self.config.sequence_parallel else 1 + if self.config.sequence_parallel: + if self.tp_group is None: + raise RuntimeError('QSA sequence parallelism requires a tensor-parallel process group.') + tp_size = self.tp_group.size() + else: + tp_size = 1 s = local_s * tp_size R = self.compress_ratio max_blocks = s // R @@ -102,7 +108,8 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch qk = self.index_qk_proj(hidden_states) if tp_size > 1: qk = gather_from_sequence_parallel_region(qk, tensor_parallel_output_grad=False, group=self.tp_group) - assert qk.shape[0] == s, f'QSA indexer expected sequence length {s}, got {qk.shape[0]}' + if qk.shape[0] != s: + raise RuntimeError(f'QSA indexer expected sequence length {s}, got {qk.shape[0]}.') q, token_k = torch.split( qk, [self.index_n_heads * self.index_head_dim, self.index_kv_heads * self.index_head_dim], dim=-1) # -> [b, s, nh, d] / [b, s, d] From 192116792ca92adfe0f5ca3b67ae5c47a1e36150 Mon Sep 17 00:00:00 2001 From: BANMA-00403 Date: Tue, 1 Sep 2026 10:19:45 +0800 Subject: [PATCH 3/3] Revert "fix: derive QSA sequence length from TP group" This reverts commit 069e57d050c7e2974ddca5d4dffd4f7976dc769d. --- src/mcore_bridge/model/modules/qsa_indexer.py | 19 ++++++------------- 1 file changed, 6 insertions(+), 13 deletions(-) diff --git a/src/mcore_bridge/model/modules/qsa_indexer.py b/src/mcore_bridge/model/modules/qsa_indexer.py index 51c51fb..5a2c706 100644 --- a/src/mcore_bridge/model/modules/qsa_indexer.py +++ b/src/mcore_bridge/model/modules/qsa_indexer.py @@ -67,11 +67,10 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch hidden_states: ``[s, b, h]`` (mcore layout), pre-attention input -- the same tensor the reference indexer consumes. ``s`` is the local sequence length when sequence parallelism is enabled. - freqs: Full-sequence mcore rotary frequencies - ``[S, 1, 1, rot_dim]``. mcore stores angles rather than cos/sin, - so they are materialized here the way ``_patch_apply_rotary_pos_emb`` - does (``cos(freqs) * mscale``), keeping the indexer's RoPE identical - to the attention's. + freqs: mcore rotary frequencies ``[s, 1, 1, rot_dim]``. mcore stores + angles rather than cos/sin, so they are materialized here the way + ``_patch_apply_rotary_pos_emb`` does (``cos(freqs) * mscale``), + keeping the indexer's RoPE identical to the attention's. Returns: ``[b, 1, s, s]`` bool mask where True marks a *masked-out* key (TE's @@ -83,12 +82,7 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch for packed/THD or context-parallel inputs (see the guard in the layer). """ local_s, b, _ = hidden_states.shape - if self.config.sequence_parallel: - if self.tp_group is None: - raise RuntimeError('QSA sequence parallelism requires a tensor-parallel process group.') - tp_size = self.tp_group.size() - else: - tp_size = 1 + tp_size = self.config.tensor_model_parallel_size if self.config.sequence_parallel else 1 s = local_s * tp_size R = self.compress_ratio max_blocks = s // R @@ -108,8 +102,7 @@ def select_mask(self, hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch qk = self.index_qk_proj(hidden_states) if tp_size > 1: qk = gather_from_sequence_parallel_region(qk, tensor_parallel_output_grad=False, group=self.tp_group) - if qk.shape[0] != s: - raise RuntimeError(f'QSA indexer expected sequence length {s}, got {qk.shape[0]}.') + assert qk.shape[0] == s, f'QSA indexer expected sequence length {s}, got {qk.shape[0]}' q, token_k = torch.split( qk, [self.index_n_heads * self.index_head_dim, self.index_kv_heads * self.index_head_dim], dim=-1) # -> [b, s, nh, d] / [b, s, d]