Skip to content

Commit 9076e99

Browse files
authored
Revert "[BugFix][SpecDecode] Fix low MTP acceptance rate for SFA+DSA_CP (MTP>1) (vllm-project#10825)" (vllm-project#10868)
### What this PR does / why we need it? This reverts commit 21e8dcf, which breaks tests/e2e/pull_request/two_card/spec_decode/test_spec_decode.py::test_qwen3_eagle3_pcp2_tp1 - vLLM version: v0.22.1 - vLLM main: vllm-project/vllm@967c5c3
1 parent eb81619 commit 9076e99

2 files changed

Lines changed: 7 additions & 48 deletions

File tree

vllm_ascend/attention/sfa_v1.py

Lines changed: 7 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -200,11 +200,6 @@ def __init__(
200200

201201
self.speculative_config = vllm_config.speculative_config
202202
self.decode_threshold = 1
203-
max_num_reqs = vllm_config.scheduler_config.max_num_seqs
204-
self.actual_seq_lengths_query = torch.zeros(max_num_reqs + 1, dtype=torch.int32, device=device)
205-
self.actual_seq_lengths_key = torch.empty_like(self.actual_seq_lengths_query)
206-
self.spec_actual_seq_lengths_query: list[torch.Tensor] | None = None
207-
self.spec_actual_seq_lengths_key: list[torch.Tensor] | None = None
208203
if self.speculative_config:
209204
spec_token_num = self.speculative_config.num_speculative_tokens
210205
self.decode_threshold += spec_token_num
@@ -213,20 +208,15 @@ def __init__(
213208
npu_fused_infer_attention_score TND layout's limit of 16, \
214209
got {self.decode_threshold}"
215210
)
216-
self.spec_actual_seq_lengths_query = [
217-
torch.zeros(max_num_reqs * (spec_token_num + 1) + 1, dtype=torch.int32, device=device)
218-
for _ in range(spec_token_num)
219-
]
220-
self.spec_actual_seq_lengths_key = [
221-
torch.zeros(max_num_reqs * (spec_token_num + 1) + 1, dtype=torch.int32, device=device)
222-
for _ in range(spec_token_num)
223-
]
224-
225211
self.reorder_batch_threshold = self.decode_threshold
226212
self.attn_mask_builder = AttentionMaskBuilder(self.device)
227213
self.rope_dim = self.model_config.hf_text_config.qk_rope_head_dim
228214
self.enable_dsa_cp = enable_dsa_cp()
229215

216+
max_num_reqs = vllm_config.scheduler_config.max_num_seqs
217+
self.actual_seq_lengths_query = torch.zeros(max_num_reqs + 1, dtype=torch.int32, device=device)
218+
self.actual_seq_lengths_key = torch.empty_like(self.actual_seq_lengths_query)
219+
230220
@staticmethod
231221
def determine_chunked_prefill_workspace_size(vllm_config: VllmConfig) -> int:
232222
return ascend_chunked_prefill_workspace_size(vllm_config)
@@ -250,22 +240,6 @@ def build(
250240
common_prefix_len: int,
251241
common_attn_metadata: AscendCommonAttentionMetadata,
252242
fast_build: bool = False,
253-
) -> AscendSFAMetadata:
254-
# common_prefix_len / fast_build are unused; kept for API compatibility.
255-
return self._build(common_attn_metadata, draft_step=None)
256-
257-
def build_for_drafting(
258-
self,
259-
draft_step: int,
260-
common_attn_metadata: AscendCommonAttentionMetadata,
261-
**kwargs,
262-
) -> AscendSFAMetadata:
263-
return self._build(common_attn_metadata, draft_step=draft_step)
264-
265-
def _build(
266-
self,
267-
common_attn_metadata: AscendCommonAttentionMetadata,
268-
draft_step: int | None = None,
269243
) -> AscendSFAMetadata:
270244
num_reqs = common_attn_metadata.num_reqs
271245
num_actual_tokens = common_attn_metadata.num_actual_tokens
@@ -291,7 +265,7 @@ def _build(
291265
else:
292266
seq_lens_cpu = common_attn_metadata.seq_lens[:num_reqs].to("cpu")
293267

294-
cos, sin = get_cos_and_sin_mla(input_positions, use_cache=(draft_step is None))
268+
cos, sin = get_cos_and_sin_mla(input_positions, True)
295269

296270
dsa_cp_context = None
297271
if self.enable_dsa_cp:
@@ -333,16 +307,8 @@ def _build(
333307
got {slot_mapping.shape[0]} and {num_tokens_pad}"
334308
)
335309

336-
if draft_step is not None:
337-
assert self.spec_actual_seq_lengths_query is not None
338-
assert self.spec_actual_seq_lengths_key is not None
339-
# Per-draft-step buffers: independent, graph-stable storage so
340-
# later draft steps don't clobber earlier ones' metadata.
341-
actual_seq_lengths_query = self.spec_actual_seq_lengths_query[draft_step - 1]
342-
actual_seq_lengths_key = self.spec_actual_seq_lengths_key[draft_step - 1]
343-
else:
344-
actual_seq_lengths_query = self.actual_seq_lengths_query
345-
actual_seq_lengths_key = self.actual_seq_lengths_key
310+
actual_seq_lengths_query = self.actual_seq_lengths_query
311+
actual_seq_lengths_key = self.actual_seq_lengths_key
346312

347313
num_segs = cum_query_lens.shape[0]
348314
last_token = 0

vllm_ascend/spec_decode/llm_base_proposer.py

Lines changed: 0 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1652,13 +1652,6 @@ def attn_update_stack_num_spec_norm(
16521652
common_attn_metadata,
16531653
**extra_attn_metadata_args,
16541654
)
1655-
elif hasattr(attn_metadata_builder, "build_for_drafting"):
1656-
# e.g. SFA (dsa_cp): route draft steps through a draft-aware build so
1657-
# per-draft-step buffers are used and cross-step aliasing is avoided.
1658-
attn_metadata = attn_metadata_builder.build_for_drafting(
1659-
draft_step,
1660-
common_attn_metadata,
1661-
)
16621655
else:
16631656
attn_metadata = attn_metadata_builder.build(
16641657
0,

0 commit comments

Comments
 (0)