diff --git a/tests/e2e/pull_request/four_card/context_parallel/test_accuracy.py b/tests/e2e/pull_request/four_card/context_parallel/test_accuracy.py index 296411858a05..38faab54b71c 100644 --- a/tests/e2e/pull_request/four_card/context_parallel/test_accuracy.py +++ b/tests/e2e/pull_request/four_card/context_parallel/test_accuracy.py @@ -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: @@ -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", + }, }, ), ] @@ -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" diff --git a/vllm_ascend/attention/context_parallel/dsa_cp.py b/vllm_ascend/attention/context_parallel/dsa_cp.py index 89675fef4457..9f4f47150427 100644 --- a/vllm_ascend/attention/context_parallel/dsa_cp.py +++ b/vllm_ascend/attention/context_parallel/dsa_cp.py @@ -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, @@ -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) # +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 @@ -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, ( @@ -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 @@ -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, @@ -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] @@ -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, @@ -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] @@ -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 @@ -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, @@ -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, @@ -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, @@ -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