[NPU] ascend backend support qwen3 moe attention cp (#21685)
This commit is contained in:
@@ -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 <PREFILL_HOST_IP>:**
|
||||
|
||||
```shell
|
||||
export SGLANG_SET_CPU_AFFINITY=1
|
||||
export ASCEND_MF_STORE_URL="tcp://<PREFILL_HOST_IP>: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 <PREFILL_HOST_IP> \
|
||||
--port 8000 \
|
||||
--nnodes 1 \
|
||||
--node-rank 0 \
|
||||
--dist-init-addr <PREFILL_HOST_IP>: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 (<DECODE_HOST_IP>):**
|
||||
|
||||
```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 <DECODE_HOST_IP> \
|
||||
--port 8001 \
|
||||
--nnodes 1 \
|
||||
--node-rank 0 \
|
||||
--dist-init-addr <DECODE_HOST_IP>: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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user