From 4c1eefca4fa978fd64cad626eefab4bcbdd6130d Mon Sep 17 00:00:00 2001 From: AndyLi429 <68410213+AndyLi429@users.noreply.github.com> Date: Wed, 29 Apr 2026 19:25:17 +0800 Subject: [PATCH] [NPU] ascend backend support qwen3 moe attention cp (#21685) --- .../ascend/ascend_npu_qwen3_examples.md | 80 +++++++++ .../npu/attention/ascend_backend.py | 164 +++++++++++++++++- .../llm_models/test_npu_qwen3_30b_attn_cp.py | 95 ++++++++++ 3 files changed, 331 insertions(+), 8 deletions(-) create mode 100644 test/registered/ascend/llm_models/test_npu_qwen3_30b_attn_cp.py diff --git a/docs/platforms/ascend/ascend_npu_qwen3_examples.md b/docs/platforms/ascend/ascend_npu_qwen3_examples.md index 7ceedd351..f17ed6b71 100644 --- a/docs/platforms/ascend/ascend_npu_qwen3_examples.md +++ b/docs/platforms/ascend/ascend_npu_qwen3_examples.md @@ -185,6 +185,86 @@ python3 -m sglang_router.launch_router \ --prometheus-port 29010 ``` +#### Running Qwen3-235B-A22B-Instruct-2507-W8A8 with Prefill Context Parallel (CP) on 2 x Atlas 800I A3 + +This example enables **Prefill Context Parallel** (`--enable-prefill-context-parallel`) to split the context across CP ranks during prefill, reducing per-device memory pressure and improving TTFT for long sequences. PD disaggregation is required. + +> **Constraints** +> - Prefill side must set `--max-running-requests 1` (PCP only supports batch_size=1) +> - `--attn-cp-size` must evenly divide `--tp-size`; each CP rank occupies `tp_size / cp_size` NPUs + +**Prefill node :** + +```shell +export SGLANG_SET_CPU_AFFINITY=1 +export ASCEND_MF_STORE_URL="tcp://:23456" +export ASCEND_USE_FIA=True + +python3 -m sglang.launch_server \ + --model-path /mnt/share/weights/Qwen3-235B-A22B-Instruct-2507-W8A8 \ + --trust-remote-code \ + --disaggregation-mode prefill \ + --disaggregation-transfer-backend ascend \ + --disaggregation-bootstrap-port 8995 \ + --quantization modelslim \ + --attention-backend ascend \ + --skip-server-warmup \ + --mem-fraction-static 0.7 \ + --chunked-prefill-size 32768 \ + --device npu \ + --base-gpu-id 0 \ + --tp-size 16 \ + --enable-prefill-context-parallel \ + --attn-cp-size 2 \ + --moe-dp-size 2 \ + --max-running-requests 1 \ + --host \ + --port 8000 \ + --nnodes 1 \ + --node-rank 0 \ + --dist-init-addr :6688 +``` + +Key parameters for PCP: + +| Parameter | Value | Description | +|-----------|-------|-------------| +| `--enable-prefill-context-parallel` | flag | Enable PCP feature | +| `--attn-cp-size` | 2 | Split context across 2 CP ranks (each rank handles half the sequence) | +| `--moe-dp-size` | 2 | MoE DP size, should match `--attn-cp-size` | +| `--max-running-requests` | 1 | Required by PCP (batch_size=1 constraint) | + +**Decode node ():** + +```shell +export ASCEND_MF_STORE_URL="tcp://141.61.39.231:23456" +export ASCEND_USE_FIA=True + +python3 -m sglang.launch_server \ + --model-path /mnt/share/weights/Qwen3-235B-A22B-Instruct-2507-W8A8 \ + --trust-remote-code \ + --disaggregation-mode decode \ + --disaggregation-transfer-backend ascend \ + --quantization modelslim \ + --attention-backend ascend \ + --disable-radix-cache \ + --disable-cuda-graph \ + --mem-fraction-static 0.7 \ + --chunked-prefill-size 32768 \ + --skip-server-warmup \ + --device npu \ + --base-gpu-id 0 \ + --tp-size 8 \ + --max-running-requests 32 \ + --host \ + --port 8001 \ + --nnodes 1 \ + --node-rank 0 \ + --dist-init-addr :6688 +``` + +> **Note:** `ASCEND_MF_STORE_URL` on both nodes must point to the same KV store (typically the Prefill node IP). `ASCEND_USE_FIA=True` enables fast interconnect aggregation for KV transfer. PCP is a Prefill-only feature; the Decode side needs no CP-related flags. + #### Running Qwen3-VL-8B-Instruct on 1 x Atlas 800I A3. Model weights could be found [here](https://huggingface.co/Qwen/Qwen3-VL-8B-Instruct) diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index 7bc7cb2b1..7ed67370c 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -23,9 +23,10 @@ from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.layers.radix_attention import AttentionType +from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.speculative.spec_info import SpecInput -from sglang.srt.utils import get_bool_env_var +from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -207,6 +208,49 @@ class AscendAttnMaskBuilder: return attn_mask +def _cp_allgather_and_save_kv_npu(forward_batch, layer, k, v, cp_size): + """NPU-compatible CP KV all-gather with merged K/V communication. + + Merges K and V along the feature dimension so only one all-gather is + needed instead of two, halving communication latency. + + k shape: [S_local, tp_k_head_num, qk_head_dim] + v shape: [S_local, tp_v_head_num, v_head_dim] + + Equivalent to cp_allgather_and_save_kv_cache() in cp_utils.py, but uses + a single all-gather for both K and V. + """ + cache_loc = ( + forward_batch.out_cache_loc + if not layer.is_cross_attention + else forward_batch.encoder_out_cache_loc + ) + # Save original trailing shapes for reshape after gather. + k_tail = k.shape[1:] # (tp_k_head_num, qk_head_dim) + v_tail = v.shape[1:] # (tp_v_head_num, v_head_dim) + + # Flatten trailing dims then concat → one all-gather instead of two. + # Works for GQA where tp_k_head_num != tp_v_head_num. + k_flat = k.contiguous().reshape(k.shape[0], -1) # [S_local, k_feat] + v_flat = v.contiguous().reshape(v.shape[0], -1) # [S_local, v_feat] + k_feat_size = k_flat.shape[-1] + kv_flat = torch.cat([k_flat, v_flat], dim=-1) # [S_local, k_feat + v_feat] + + kv_full = cp_all_gather_rerange_kv_cache( + kv_flat, cp_size, forward_batch, get_current_device_stream_fast() + ) # [S_full, k_feat + v_feat] + + key_cache_full = kv_full[..., :k_feat_size].reshape(-1, *k_tail) + value_cache_full = kv_full[..., k_feat_size:].reshape(-1, *v_tail) + + forward_batch.token_to_kv_pool.set_kv_buffer( + layer, + cache_loc, + key_cache_full, + value_cache_full, + ) + + class AscendAttnBackend(AttentionBackend): def __init__(self, model_runner: ModelRunner, speculative_step_id: int = 0): @@ -286,6 +330,8 @@ class AscendAttnBackend(AttentionBackend): self.is_dllm_model = True self.dllm_block_size = self.dllm_config.block_size + self.attn_cp_size = model_runner.attn_cp_size + def get_verify_buffers_to_fill_after_draft(self): """ Return buffers for verify attention kernels that needs to be filled after draft. @@ -736,6 +782,83 @@ class AscendAttnBackend(AttentionBackend): ) return torch.cat([attn_out_prev, attn_out_next], dim=0) + def do_cp_attn_fia( + self, + q: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + layer: "RadixAttention", + forward_batch: ForwardBatch, + ) -> torch.Tensor: + """CP-aware attention for standard (non-MLA) models using FIA on Ascend NPU. + + Uses npu_fused_infer_attention_score with paged KV cache (block_table). + The KV cache must already contain the full gathered sequence + (written by _cp_allgather_and_save_kv_npu before this call). + + Args: + q: Query tensor, shape [total_q_tokens, tp_q_head_num * qk_head_dim] + k_cache: Full key cache from token_to_kv_pool + v_cache: Full value cache from token_to_kv_pool + layer: RadixAttention layer + forward_batch: ForwardBatch with attn_cp_metadata populated + + Returns: + attn_output [total_q_tokens, tp_q_head_num * v_head_dim] + """ + cp_meta = forward_batch.attn_cp_metadata + + # Split Q into prev/next halves per zigzag pattern. + # torch.chunk(q, 2) gives ceil(n/2) and floor(n/2), matching + # actual_seq_q_prev and actual_seq_q_next. + q_prev, q_next = torch.chunk(q, 2, dim=0) + q_prev = q_prev.contiguous().reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) + q_next = q_next.contiguous().reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) + + k_cache_paged = k_cache.view( + -1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim + ) + v_cache_paged = v_cache.view( + -1, self.page_size, layer.tp_v_head_num * layer.v_head_dim + ) + + attn_out_prev, _ = torch.ops.npu.npu_fused_infer_attention_score( + q_prev, + k_cache_paged, + v_cache_paged, + block_table=self.forward_metadata.block_tables, + block_size=self.page_size, + num_heads=layer.tp_q_head_num, + num_key_value_heads=layer.tp_k_head_num, + input_layout="TND", + atten_mask=self.fia_mask, + sparse_mode=3, + next_tokens=0, + scale=layer.scaling, + actual_seq_lengths=[cp_meta.actual_seq_q_prev], + actual_seq_lengths_kv=[cp_meta.kv_len_prev], + ) + + attn_out_next, _ = torch.ops.npu.npu_fused_infer_attention_score( + q_next, + k_cache_paged, + v_cache_paged, + block_table=self.forward_metadata.block_tables, + block_size=self.page_size, + num_heads=layer.tp_q_head_num, + num_key_value_heads=layer.tp_k_head_num, + input_layout="TND", + atten_mask=self.fia_mask, + sparse_mode=3, + next_tokens=0, + scale=layer.scaling, + actual_seq_lengths=[cp_meta.actual_seq_q_next], + actual_seq_lengths_kv=[cp_meta.kv_len_next], + ) + + attn_out = torch.cat([attn_out_prev, attn_out_next], dim=0) + return attn_out.view(-1, layer.tp_q_head_num * layer.v_head_dim) + def forward_sparse( self, q: torch.Tensor, @@ -906,15 +1029,28 @@ class AscendAttnBackend(AttentionBackend): ) if not self.use_mla: + # Detect CP mode for prefill (context parallel) + is_cp_mode = ( + forward_batch.forward_mode.is_context_parallel_extend() + and forward_batch.attn_cp_metadata is not None + and self.attn_cp_size > 1 + ) + # In cross attention layer, when there is no vision input,the values of k and v is None if save_kv_cache and k is not None and v is not None: - # support cross attention - cache_loc = ( - forward_batch.out_cache_loc - if not layer.is_cross_attention - else forward_batch.encoder_out_cache_loc - ) - forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) + if is_cp_mode: + # All-gather K/V from all CP ranks and write full sequence to KV pool + _cp_allgather_and_save_kv_npu( + forward_batch, layer, k, v, self.attn_cp_size + ) + else: + # support cross attention + cache_loc = ( + forward_batch.out_cache_loc + if not layer.is_cross_attention + else forward_batch.encoder_out_cache_loc + ) + forward_batch.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v) k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) @@ -940,6 +1076,18 @@ class AscendAttnBackend(AttentionBackend): ) return attn_out + if is_cp_mode: + if self.use_fia: + attn_output = self.do_cp_attn_fia( + q, k_cache, v_cache, layer, forward_batch + ) + else: + raise NotImplementedError( + "CP attention for non-FIA path on Ascend is not yet implemented. " + "Set ASCEND_USE_FIA=1 to use FIA-based CP attention." + ) + return attn_output + if self.use_fia: """FIA will support multi-bs in the later version of CANN""" q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim) diff --git a/test/registered/ascend/llm_models/test_npu_qwen3_30b_attn_cp.py b/test/registered/ascend/llm_models/test_npu_qwen3_30b_attn_cp.py new file mode 100644 index 000000000..32604a3d8 --- /dev/null +++ b/test/registered/ascend/llm_models/test_npu_qwen3_30b_attn_cp.py @@ -0,0 +1,95 @@ +import os +import unittest +from types import SimpleNamespace + +from python.sglang.test.ascend.test_ascend_utils import QWEN3_30B_A3B_WEIGHTS_PATH +from sglang.test.ci.ci_register import register_npu_ci +from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + kill_process_tree, + popen_launch_server, +) + +register_npu_ci(est_time=500, suite="nightly-4-npu-a3", nightly=True) + +QWEN3_30B_MODEL = QWEN3_30B_A3B_WEIGHTS_PATH +GSM8K_MIN_ACCURACY = 0.92 +GSM8K_NUM_QUESTIONS = 100 + +_NPU_ENV_VARS = { + "ASCEND_USE_FIA": "1", +} + + +class TestQwen330BAttnCP(CustomTestCase): + """GSM8K accuracy test for Qwen3-30B-A3B mixed deployment on 4 NPUs. + + The test uses: + - TP = 4 + - MOE_DP = 2 + - ATTN_CP = 2 + - prefill context parallel enabled + + This is the mixed/co-located deployment variant and reuses the Ascend + environment variables from the PD GSM8K test. + """ + + @classmethod + def setUpClass(cls): + cls.model = QWEN3_30B_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.npu_env = {**os.environ, **_NPU_ENV_VARS} + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--trust-remote-code", + "--mem-fraction-static", + "0.7", + "--max-running-requests", + "32", + "--attention-backend", + "ascend", + "--tp-size", + "4", + "--moe-dp-size", + "2", + "--attn-cp-size", + "2", + "--cuda-graph-max-bs", + "32", + "--enable-prefill-context-parallel", + ], + env=cls.npu_env, + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process is not None: + kill_process_tree(cls.process.pid) + + def test_gsm8k_accuracy(self): + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=GSM8K_NUM_QUESTIONS, + max_new_tokens=512, + parallel=32, + host="http://127.0.0.1", + port=int(self.base_url.split(":")[-1]), + ) + metrics = run_eval_few_shot_gsm8k(args) + print( + "GSM8K accuracy " + f"(mixed TP=4 MOE_DP=2 ATTN_CP=2, {GSM8K_NUM_QUESTIONS} samples): " + f"{metrics['accuracy']:.3f}" + ) + self.assertGreaterEqual(metrics["accuracy"], GSM8K_MIN_ACCURACY) + + +if __name__ == "__main__": + unittest.main()