[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
|
--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.
|
#### 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)
|
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.attention.nsa.utils import is_nsa_enable_prefill_cp
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
from sglang.srt.layers.radix_attention import AttentionType
|
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.speculative.spec_info import SpecInput
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -207,6 +208,49 @@ class AscendAttnMaskBuilder:
|
|||||||
return attn_mask
|
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):
|
class AscendAttnBackend(AttentionBackend):
|
||||||
|
|
||||||
def __init__(self, model_runner: ModelRunner, speculative_step_id: int = 0):
|
def __init__(self, model_runner: ModelRunner, speculative_step_id: int = 0):
|
||||||
@@ -286,6 +330,8 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
self.is_dllm_model = True
|
self.is_dllm_model = True
|
||||||
self.dllm_block_size = self.dllm_config.block_size
|
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):
|
def get_verify_buffers_to_fill_after_draft(self):
|
||||||
"""
|
"""
|
||||||
Return buffers for verify attention kernels that needs to be filled after draft.
|
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)
|
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(
|
def forward_sparse(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -906,8 +1029,21 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if not self.use_mla:
|
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
|
# 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:
|
if save_kv_cache and k is not None and v is not None:
|
||||||
|
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
|
# support cross attention
|
||||||
cache_loc = (
|
cache_loc = (
|
||||||
forward_batch.out_cache_loc
|
forward_batch.out_cache_loc
|
||||||
@@ -940,6 +1076,18 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
return attn_out
|
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:
|
if self.use_fia:
|
||||||
"""FIA will support multi-bs in the later version of CANN"""
|
"""FIA will support multi-bs in the later version of CANN"""
|
||||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
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