Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
75 changes: 75 additions & 0 deletions tests/e2e/pull_request/four_card/context_parallel/test_accuracy.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,21 @@
DEEPSEEK_V4_GOLDEN = ["Hello, my name is {name} and I", 'What is the meaning of life?",\n "What is']
DEEPSEEK_V4_MODEL = "gdydems/DeepSeek-V4-Flash-w4a8-mtp"

# max_num_seqs deliberately smaller than the cudagraph capture bucket so that
# FULL-decode graph padding pushes the draft-path num_reqs beyond max_num_seqs.
CONCURRENT_MAX_NUM_SEQS = 4
CONCURRENT_CAPTURE_SIZES = [16]
CONCURRENT_PROMPTS = [
"The capital of France is",
"Hello, my name is",
"What is the meaning of life?",
"The president of United States is",
"Write a short story about a robot.",
"Explain quantum computing simply.",
"List three primary colors.",
"Describe a rainy day in Paris.",
]


@dataclass(frozen=True)
class AccuracyCase:
Expand Down Expand Up @@ -174,6 +189,10 @@ def match_outputs_with_goldens(outputs: list[tuple[list[int], str]], goldens: Se
"additional_config": {
"enable_dsa_cp": True,
},
"speculative_config": {
"num_speculative_tokens": 3,
"method": "mtp",
},
},
),
]
Expand Down Expand Up @@ -242,3 +261,59 @@ def test_deepseek_v4_dsa_cp_prefill_decode_accuracy() -> None:
@pytest.mark.parametrize("case", FULL_FEATURE_MODEL_CASES, ids=lambda case: case.name)
def test_models_dcp_full_feature_accuracy(case: AccuracyCase) -> None:
_run_accuracy_case(case)


@patch.dict(
os.environ,
{
"HCCL_BUFFSIZE": "768",
"PYTORCH_NPU_ALLOC_CONF": "expandable_segments:True",
},
)
@wait_until_npu_memory_free(target_free_percentage=0.8)
def test_models_dcp_full_graph_concurrent_requests() -> None:
"""Accuracy guard for concurrent requests under FULL-decode aclgraph.

Regression test for the DSA-CP per-request metadata buffers: in
FULL_DECODE_ONLY graph mode the draft path receives ``num_reqs`` padded to
the cudagraph capture bucket (plus the FIA dummy request), which exceeds
``max_num_seqs`` here (CONCURRENT_CAPTURE_SIZES=[16] vs
CONCURRENT_MAX_NUM_SEQS=4). Buffers sized by ``max_num_seqs`` get their
``[:num_reqs]`` views truncated — copy_() then fails with a shape mismatch
(EZ1007) or triton kernels write past the buffer end and corrupt adjacent
device memory (EZ9999 MTE illegal GM address).

With MTP(num_speculative_tokens=3) each decode step of 4 concurrent
requests yields 16 tokens, so the batch is padded to the 16-token capture
bucket and every draft step exercises the padded path.
"""
runner_kwargs: dict[str, Any] = {
"max_model_len": 8192,
"max_num_seqs": CONCURRENT_MAX_NUM_SEQS,
"max_num_batched_tokens": 4096,
"dtype": "auto",
"tensor_parallel_size": 4,
"decode_context_parallel_size": 1,
"enable_expert_parallel": True,
"gpu_memory_utilization": 0.9,
"quantization": "ascend",
"tokenizer_mode": "deepseek_v4",
"block_size": 128,
"compilation_config": {
"cudagraph_mode": "FULL_DECODE_ONLY",
"cudagraph_capture_sizes": CONCURRENT_CAPTURE_SIZES,
},
"additional_config": {
"enable_dsa_cp": True,
},
"speculative_config": {
"num_speculative_tokens": 3,
"method": "mtp",
},
}
with VllmRunner("gdydems/DeepSeek-V4-Flash-w4a8-mtp", **runner_kwargs) as runner:
outputs = runner.generate_greedy(CONCURRENT_PROMPTS, 5)

assert len(outputs) == len(CONCURRENT_PROMPTS)
for _, output_str in outputs:
assert output_str, "Concurrent request produced an empty output"
117 changes: 82 additions & 35 deletions vllm_ascend/attention/context_parallel/dsa_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,9 @@
# Legacy DSA-CP implementation (TP/SP group)
# =============================================================================

# Fixed-size contract required by the underlying _C_ascend ops (must be exactly 1024).
SAS_METADATA_SIZE = 1024


def hadamard_transform_ref(
x: torch.Tensor,
Expand Down Expand Up @@ -252,29 +255,31 @@ def __init__(
torch.bfloat16
),
)
self.start_pos_prefill = torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
self.req_sas_metadata = torch.zeros(1024, dtype=torch.int32, device=self.device)
self.req_qli_metadata = torch.zeros(1024, dtype=torch.int32, device=self.device)
# Full-decode graphs pad the request count beyond max_num_seqs
# (cudagraph capture sizes plus the FIA dummy request), so size the
# per-request QLI buffers for the graph-mode maximum.
max_qli_reqs = scheduler_config.max_num_seqs
# (cudagraph capture sizes plus the FIA dummy request), so size all
# per-request buffers for the graph-mode maximum. Otherwise the
# [:num_reqs] views taken below get truncated, copy_() into them
# fails with shape mismatches, or triton kernels write past the
# buffer end and corrupt adjacent device memory.
compilation_config = self.vllm_config.compilation_config
max_padded_reqs = scheduler_config.max_num_seqs
if compilation_config.cudagraph_mode != CUDAGraphMode.NONE and compilation_config.cudagraph_capture_sizes:
max_qli_reqs = max(max_qli_reqs, compilation_config.max_cudagraph_capture_size)
max_padded_reqs = max(max_padded_reqs, compilation_config.max_cudagraph_capture_size)
Comment thread
yiz-liu marked this conversation as resolved.
# +1 holds the FIA dummy request inserted by mixed-batch padding.
self.qli_seqused_k = torch.zeros(max_qli_reqs + 1, dtype=torch.int32, device=self.device)
self.qli_cmp_residual_k = torch.zeros(max_qli_reqs + 1, dtype=torch.int32, device=self.device)
max_padded_reqs += 1
self.start_pos_prefill = torch.zeros(max_padded_reqs, dtype=torch.int32, device=self.device)
self.req_sas_metadata = torch.zeros(SAS_METADATA_SIZE, dtype=torch.int32, device=self.device)
self.req_qli_metadata = torch.zeros(SAS_METADATA_SIZE, dtype=torch.int32, device=self.device)
self.qli_seqused_k = torch.zeros(max_padded_reqs, dtype=torch.int32, device=self.device)
self.qli_cmp_residual_k = torch.zeros(max_padded_reqs, dtype=torch.int32, device=self.device)
self._device_metadata_enabled = False
self._device_metadata_tasks: tuple[DeviceMetadataTask, ...] = ()
self.cu_seqlens_ori_kv = torch.tensor([], device=self.device)
self.cu_seqlens_cmp_kv = torch.tensor([], device=self.device)
self.seqused_q = torch.tensor([], device=self.device)
self._zero_i32 = torch.tensor([0], device=self.device, dtype=torch.int32)
self.local_query_start_loc = torch.zeros(
scheduler_config.max_num_seqs + 1, dtype=torch.int32, device=self.device
)
self.local_seq_lens = torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
self.local_query_start_loc = torch.zeros(max_padded_reqs + 1, dtype=torch.int32, device=self.device)
self.local_seq_lens = torch.zeros(max_padded_reqs, dtype=torch.int32, device=self.device)

self.speculative_config = vllm_config.speculative_config
self.decode_threshold = 1
Expand All @@ -292,12 +297,16 @@ def __init__(
for _ in range(spec_token_num)
]
self.spec_local_query_start_loc = [
torch.zeros(scheduler_config.max_num_seqs + 1, dtype=torch.int32, device=self.device)
for _ in range(spec_token_num)
torch.zeros(max_padded_reqs + 1, dtype=torch.int32, device=self.device) for _ in range(spec_token_num)
]
self.spec_local_seq_lens = [
torch.zeros(scheduler_config.max_num_seqs, dtype=torch.int32, device=self.device)
for _ in range(spec_token_num)
torch.zeros(max_padded_reqs, dtype=torch.int32, device=self.device) for _ in range(spec_token_num)
]
self.spec_sas_metadata = [
torch.zeros(SAS_METADATA_SIZE, dtype=torch.int32, device=self.device) for _ in range(spec_token_num)
]
self.spec_start_pos = [
torch.zeros(max_padded_reqs, dtype=torch.int32, device=self.device) for _ in range(spec_token_num)
]
self.decode_threshold += spec_token_num
assert self.decode_threshold <= 16, (
Expand Down Expand Up @@ -429,6 +438,7 @@ def build_for_drafting(
**kwargs,
) -> AscendDSAMetadata:
assert self.compressor_ratio <= 1, "vLLM-Ascend only support SWA-layer for Deepseek-V4 now."
# NOTE(Csrayz): num_reqs here is the padded token count from graph dispatch
num_reqs = common_attn_metadata.num_reqs
num_input_tokens = common_attn_metadata.num_input_tokens
# Cross-kv-cache-group metadata cache. The spec-decode proposer passes
Expand Down Expand Up @@ -462,9 +472,10 @@ def build_for_drafting(
treat_short_extends_as_decodes=False,
)
input_positions = common_attn_metadata.positions[:num_input_tokens].long()
# Draft steps update positions independently. Reusing the global RoPE
# cache can let later draft steps overwrite step-0 metadata.
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=False)
# Use per-draft-index RoPE buffer so tensor addresses stay stable
# across graph capture/replay; the per-step cache below then lets
# sibling kv-cache groups reuse the same stable tensors.
cos, sin = get_cos_and_sin_dsa(input_positions, use_cache=True, draft_index=draft_index)
if metadata_cache is not None:
metadata_cache.update(
num_decodes=num_decodes,
Expand Down Expand Up @@ -492,6 +503,13 @@ def build_for_drafting(
self.seq_lens_cpu = self.seq_lens.cpu()
if metadata_cache is not None:
metadata_cache["seq_lens_cpu"] = self.seq_lens_cpu
# In aclgraph mode num_reqs is the padded request count; keep the
# tensor's batch dim stable at num_reqs across capture and replay.
num_real_reqs = self.seq_lens_cpu.shape[0]
if num_real_reqs < num_reqs:
self.seq_lens_cpu = F.pad(self.seq_lens_cpu, (0, num_reqs - num_real_reqs))
else:
self.seq_lens_cpu = self.seq_lens_cpu[:num_reqs]

slot_mapping = common_attn_metadata.slot_mapping[:num_input_tokens]

Expand Down Expand Up @@ -549,19 +567,25 @@ def build_req_metadata_for_drafting(

if metadata_cache is not None and "local_query_start_loc" in metadata_cache:
# Hit: local token metadata computed by a sibling kv-cache group
# of the same draft step. Values are identical across groups, so
# the cached (cloned) tensors are reused directly.
# of the same draft step. Values are identical across groups; copy
# them into this builder's stable per-draft buffers so the returned
# views keep fixed addresses across aclgraph capture/replay.
local_start = metadata_cache["local_start"]
local_end_with_pad = metadata_cache["local_end_with_pad"]
tokens_per_rank = metadata_cache["tokens_per_rank"]
num_tokens_pad = metadata_cache["num_tokens_pad"]
local_query_start_loc = metadata_cache["local_query_start_loc"]
local_seq_lens = metadata_cache["local_seq_lens"]
max_local_query_len = metadata_cache["max_local_query_len"]
max_local_seq_lens = metadata_cache["max_local_seq_lens"]
local_cos = metadata_cache["local_cos"]
local_sin = metadata_cache["local_sin"]
start_pos = metadata_cache["start_pos"]
self.spec_local_query_start_loc[draft_index - 1][: num_reqs + 1].copy_(
metadata_cache["local_query_start_loc"]
)
self.spec_local_seq_lens[draft_index - 1][:num_reqs].copy_(metadata_cache["local_seq_lens"])
self.spec_start_pos[draft_index - 1][:num_reqs].copy_(metadata_cache["start_pos"])
local_query_start_loc = self.spec_local_query_start_loc[draft_index - 1][: num_reqs + 1]
local_seq_lens = self.spec_local_seq_lens[draft_index - 1][:num_reqs]
start_pos = self.spec_start_pos[draft_index - 1][:num_reqs]
else:
(
local_start,
Expand All @@ -579,8 +603,11 @@ def build_req_metadata_for_drafting(
local_seq_lens=self.spec_local_seq_lens[draft_index - 1],
is_noncausal=is_noncausal,
)
local_query_start_loc = local_query_start_loc.clone()
local_seq_lens = local_seq_lens.clone()
# NOTE: no .clone() here. ACL graph capture bakes the addresses of
# these tensors into the draft graph, so every replay must read the
# freshly built values from the same (stable) buffer addresses.
# Returning a fresh clone would leave the graph reading stale
# capture-time metadata and cause illegal device memory accesses.
local_cos = cos.pad_to(num_tokens_pad)[local_start:local_end_with_pad]
local_sin = sin.pad_to(num_tokens_pad)[local_start:local_end_with_pad]

Expand All @@ -595,21 +622,24 @@ def build_req_metadata_for_drafting(
max_local_query_len = max(1, int(local_seq_lens_q_cpu.max().item()))
max_local_seq_lens = max(1, int(local_seq_lens_cpu.max().item()))

start_pos = self.seq_lens[:num_reqs] - seq_lens_q
start_pos = self.spec_start_pos[draft_index - 1][:num_reqs]
start_pos.copy_(self.seq_lens[:num_reqs] - seq_lens_q)

if metadata_cache is not None:
# Store value snapshots: the persistent buffers are refilled on
# every step, so sibling groups must copy the values now.
metadata_cache.update(
local_start=local_start,
local_end_with_pad=local_end_with_pad,
tokens_per_rank=tokens_per_rank,
num_tokens_pad=num_tokens_pad,
local_query_start_loc=local_query_start_loc,
local_seq_lens=local_seq_lens,
local_query_start_loc=local_query_start_loc.clone(),
local_seq_lens=local_seq_lens.clone(),
max_local_query_len=max_local_query_len,
max_local_seq_lens=max_local_seq_lens,
local_cos=local_cos,
local_sin=local_sin,
start_pos=start_pos,
start_pos=start_pos.clone(),
)

dspark_swa_indices = None
Expand Down Expand Up @@ -700,6 +730,11 @@ def build_req_metadata_for_drafting(
sas_head_dim=head_dim,
sas_metadata=sas_metadata,
)
# Cache sas_metadata in the per-draft-index buffer so the tensor
# address stays stable across aclgraph capture/replay, regardless of
# whether it came from the device metadata kernel or a sibling group.
self.spec_sas_metadata[draft_index - 1][:SAS_METADATA_SIZE].copy_(sas_metadata[:SAS_METADATA_SIZE])
sas_metadata = self.spec_sas_metadata[draft_index - 1]

cp_metadata = DSACPMetadata(
local_query_start_loc=local_query_start_loc,
Expand Down Expand Up @@ -1258,8 +1293,8 @@ def _build_sas_metadata(

metadata = metadata_op(**kw)
self.common_ratio_to_sas_metadata[cache_key] = metadata
self.req_sas_metadata[:1024] = metadata
return self.req_sas_metadata[:1024]
self.req_sas_metadata[:SAS_METADATA_SIZE] = metadata
return self.req_sas_metadata[:SAS_METADATA_SIZE]

def _build_qli_metadata(
self,
Expand Down Expand Up @@ -1308,8 +1343,8 @@ def _build_qli_metadata(
device=str(self.seqused_q.device),
)
self.common_ratio_to_sas_metadata[cache_key] = metadata
self.req_qli_metadata[:1024] = metadata
return self.req_qli_metadata[:1024]
self.req_qli_metadata[:SAS_METADATA_SIZE] = metadata
return self.req_qli_metadata[:SAS_METADATA_SIZE]

def build_for_graph_capture(
self,
Expand Down Expand Up @@ -1493,6 +1528,18 @@ def process_weights_after_loading(self, act_dtype: torch.dtype):
if self.enable_dsa_cp_full_o_proj:
self._enable_o_proj_full_weight_switch()

@staticmethod
def update_graph_params(
update_stream,
forward_context,
num_tokens,
vllm_config=None,
speculative_config=None,
draft_attn_metadatas=None,
):
# DSA-CP does not need to update graph params.
pass

@staticmethod
def _get_weight_switch_method(layer: torch.nn.Module) -> WeightSwitchMixin:
quant_method = layer.quant_method
Expand Down
Loading