Skip to content

Commit 261a3e9

Browse files
czc-unacChen zhichaozhangxinyuehfad
authored
[Misc] M2m 0724 continuation (vllm-project#12648)
### What this PR does / why we need it? Adapt vllm-ascend to vLLM main commits up to July 24. ### Changes | Files | Upstream vLLM Change | vllm-ascend Adaptation | |---|---|---| | `vllm_ascend/worker/v2/attn_utils.py`<br>`vllm_ascend/worker/v2/spec_decode/dflash/speculator.py`<br>`vllm_ascend/worker/v2/spec_decode/dspark/speculator.py` | [vllm#44492](vllm-project/vllm#44492) ([4ec199b6](vllm-project/vllm@4ec199b6)) — Populate draft `seq_lens_cpu_upper_bound` for spec-decode attention metadata. Upstream threads `seq_lens_cpu_upper_bound` through `_multi_step_decode()` (new 5th positional arg) and `_build_draft_attn_metadata()` (new `seq_lens_cpu_upper_bound` + `step` kwargs). | `attn_utils.py`: Added `build_attn_metadata` with `seq_lens_cpu_upper_bound` param (defaults to `seq_lens_cpu`) and `build_attn_metadata_wrapper()` context manager for Ascend NPU attention metadata building.<br>`dflash/speculator.py`: `propose()` caches `self.input_batch` before `super().propose()`; `build_draft_attn_metadatas()` passes actual `num_reqs`, padded counts, `seq_lens_cpu_upper_bound`, and `step` to `_build_draft_attn_metadata`.<br>`dspark/speculator.py`: `build_draft_attn_metadatas()` passes `seq_lens_cpu_upper_bound` + `step`; `propose()` caches `self.input_batch` (version-gated for main only, since v0.25.1 never reaches this code path via `run_fullgraph`). | | `vllm_ascend/worker/v2/spec_decode/dflash/aclgraph.py` | [vllm#47914](vllm-project/vllm#47914) ([0d12618e](vllm-project/vllm@0d12618e)) — `causal` moved from `DFlashCudaGraphManager.__init__` to `capture()`. | `capture()` is version-gated: passes `causal` kwarg only on main branch; v0.25.1 receives it via `__init__`. `run_fullgraph()` accesses `self.speculator.input_batch.seq_lens_cpu_upper_bound` for building draft attention metadata during full-graph replay, then updates Ascend NPU full-graph params. | | `vllm_ascend/core/recompute_scheduler.py` | [vllm#47312](vllm-project/vllm#47312) ([12213c67](vllm-project/vllm@12213c67)) — Handle grammar compilation failures. `finish_requests()` return type changed from `list[tuple(str, int)]` to `list[Request]`. | Assertion in `_finish_recomputed_request` is version-gated: on main branch extracts `(request_id, client_index)` from `Request` objects; v0.25.1 retains the direct tuple comparison. | | `vllm_ascend/_310p/model_runner_310p.py` | General vLLM V1 API evolution (model loading, KV cache allocation, `determine_available_memory` flow). | New file. 310P hardware platform model runner extending `NPUModelRunner`. Adapts 310P-specific ACL graph capture/replay, splitfuse attention, NPU input batch, and FRACTAL_NZ weight layout to the upstream V1 API. | | `.github/vllm-main-verified.commit` | N/A (CI tracking) | Updated to `d02df748bf9efd99022f1a062597dc3cb3808485` to record the verified upstream vLLM commit. | ### Does this PR introduce _any_ user-facing change? ### How was this patch tested? - vLLM version: v0.25.1 - vLLM main: vllm-project/vllm@fe784ff --------- Signed-off-by: Chen zhichao <chenzhichao33@h-partners.com> Signed-off-by: hfadzxy <starmoon_zhang@163.com> Co-authored-by: Chen zhichao <chenzhichao33@h-partners.com> Co-authored-by: hfadzxy <starmoon_zhang@163.com>
1 parent 686732c commit 261a3e9

10 files changed

Lines changed: 84 additions & 19 deletions

File tree

.github/vllm-main-verified.commit

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
fe784ff22e630a31fd798f392b01e0a75c18f047
1+
d02df748bf9efd99022f1a062597dc3cb3808485

vllm_ascend/_310p/model_runner_310p.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,11 @@
5454
from vllm_ascend.spec_decode.utils import (
5555
update_num_computed_tokens_for_batch_change,
5656
)
57-
from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ, is_rc_device, lmhead_tp_enable
57+
from vllm_ascend.utils import (
58+
ACL_FORMAT_FRACTAL_NZ,
59+
is_rc_device,
60+
lmhead_tp_enable,
61+
)
5862
from vllm_ascend.worker.model_runner_v1 import NPUModelRunner
5963

6064
_NGRAM_GRAPH_UNIFORM_DECODE_QUERY_LEN = 1

vllm_ascend/core/recompute_scheduler.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -143,7 +143,12 @@ def _finish_recomputed_request(
143143
request.request_id,
144144
RequestStatus.FINISHED_ABORTED,
145145
)
146-
assert finished_reqs == [(request.request_id, request.client_index)]
146+
if vllm_version_is("0.25.1"):
147+
assert finished_reqs == [(request.request_id, request.client_index)]
148+
else:
149+
assert [(r.request_id, r.client_index) for r in finished_reqs] == [
150+
(request.request_id, request.client_index)
151+
]
147152

148153
def schedule(self, throttle_prefills: bool = False) -> RecomputeSchedulerOutput:
149154
self.current_step += 1

vllm_ascend/worker/v2/attn_utils.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,7 @@
4545
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata
4646
from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec
4747
from vllm_ascend.quantization.utils import enable_fa_quant
48-
from vllm_ascend.utils import calc_split_factor
48+
from vllm_ascend.utils import calc_split_factor, vllm_version_is
4949

5050
_ATTENTION_MASK_BUILDER = None
5151

@@ -109,6 +109,7 @@ def build_attn_metadata(
109109
dcp_local_seq_lens: torch.Tensor | None = None,
110110
# extra attributes for ascend npus.
111111
seq_lens_np: np.ndarray | None = None,
112+
seq_lens_cpu_upper_bound: torch.Tensor | None = None,
112113
num_computed_tokens_cpu: torch.Tensor | None = None,
113114
positions: torch.Tensor | None = None,
114115
attn_state: Any | None = None,
@@ -127,6 +128,8 @@ def build_attn_metadata(
127128
if seq_lens_np is None:
128129
seq_lens_np = np.full(num_reqs, max_seq_len, dtype=np.int32)
129130
seq_lens_cpu = torch.from_numpy(seq_lens_np)[:num_reqs]
131+
if not vllm_version_is("0.25.1") and seq_lens_cpu_upper_bound is None:
132+
seq_lens_cpu_upper_bound = seq_lens_cpu
130133

131134
attn_metadata: dict[str, Any] = {}
132135
kv_cache_groups = kv_cache_config.kv_cache_groups
@@ -145,7 +148,7 @@ def build_attn_metadata(
145148
query_start_loc=query_start_loc_gpu,
146149
query_start_loc_cpu=query_start_loc_cpu,
147150
seq_lens_cpu=seq_lens_cpu,
148-
seq_lens_cpu_upper_bound=seq_lens_cpu,
151+
seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound,
149152
seq_lens=seq_lens[:num_reqs],
150153
num_reqs=num_reqs,
151154
num_actual_tokens=num_tokens,

vllm_ascend/worker/v2/spec_decode/autoregressive/speculator.py

Lines changed: 16 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@
3535
from vllm.v1.worker.utils import AttentionGroup
3636

3737
from vllm_ascend.attention.attention_v1 import AscendAttentionState
38+
from vllm_ascend.utils import vllm_version_is
3839
from vllm_ascend.worker.v2.attn_utils import build_attn_metadata_wrapper
3940
from vllm_ascend.worker.v2.input_batch import AscendInputBuffers
4041

@@ -275,6 +276,7 @@ def _multi_step_decode(
275276
skip_attn: bool,
276277
batch_desc: BatchExecutionDescriptor,
277278
num_tokens_across_dp: torch.Tensor | None,
279+
seq_lens_cpu_upper_bound: torch.Tensor | None = None,
278280
) -> None:
279281
"""Minimal override to handle the merged multi-step graph in FULL mode.
280282
@@ -288,23 +290,29 @@ def _multi_step_decode(
288290
assert self.decode_cudagraph_manager is not None
289291
self.decode_cudagraph_manager.run_fullgraph(batch_desc)
290292
return
291-
super()._multi_step_decode(num_reqs, skip_attn, batch_desc, num_tokens_across_dp)
293+
if vllm_version_is("0.25.1"):
294+
super()._multi_step_decode(num_reqs, skip_attn, batch_desc, num_tokens_across_dp)
295+
else:
296+
super()._multi_step_decode(num_reqs, skip_attn, batch_desc, num_tokens_across_dp, seq_lens_cpu_upper_bound)
292297

293298
def _build_draft_attn_metadata(
294299
self,
295300
num_reqs: int,
296301
num_reqs_padded: int,
297302
num_tokens_padded: int,
303+
seq_lens_cpu_upper_bound: torch.Tensor | None = None,
304+
step: int = 1,
298305
num_query_per_req: int = 1,
299306
causal: bool = True,
300307
) -> dict[str, Any] | None:
301-
attn_metadata = super()._build_draft_attn_metadata(
302-
num_reqs,
303-
num_reqs_padded,
304-
num_tokens_padded,
305-
num_query_per_req,
306-
causal,
307-
)
308+
if vllm_version_is("0.25.1"):
309+
attn_metadata = super()._build_draft_attn_metadata(
310+
num_reqs, num_reqs_padded, num_tokens_padded, num_query_per_req, causal
311+
)
312+
else:
313+
attn_metadata = super()._build_draft_attn_metadata(
314+
num_reqs, num_reqs_padded, num_tokens_padded, seq_lens_cpu_upper_bound, step, num_query_per_req, causal
315+
)
308316
if attn_metadata is not None:
309317
# Ascend-specific: force DecodeOnly attention state for the draft model.
310318
for metadata in attn_metadata.values():

vllm_ascend/worker/v2/spec_decode/dflash/aclgraph.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -105,7 +105,10 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[
105105
"""Override run_fullgraph to update full graph params in run_fullgraph."""
106106
num_tokens = desc.num_tokens
107107

108-
draft_attn_metadatas = self.speculator.build_draft_attn_metadatas(desc.num_reqs)
108+
draft_attn_metadatas = self.speculator.build_draft_attn_metadatas(
109+
desc.num_reqs,
110+
self.speculator.input_batch.seq_lens_cpu_upper_bound,
111+
)
109112

110113
ret = super().run_fullgraph(desc)
111114

vllm_ascend/worker/v2/spec_decode/dflash/speculator.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
DFlashSpeculator,
1717
)
1818

19+
from vllm_ascend.utils import vllm_version_is
1920
from vllm_ascend.worker.v2.attn_utils import build_attn_metadata_wrapper
2021

2122
logger = logging.getLogger(__name__)
@@ -77,13 +78,15 @@ def set_attn(
7778

7879
# NOTE: upstream vLLM named this to _build_draft_attn_metadatas;
7980
# keep the current name for now as upstream may change it again.
80-
def build_draft_attn_metadatas(self, num_reqs_padded):
81+
def build_draft_attn_metadatas(self, num_reqs_padded, seq_lens_cpu_upper_bound):
8182
num_tokens_padded = num_reqs_padded * self.num_query_per_req
8283
with build_attn_metadata_wrapper():
8384
attn_metadata = self._build_draft_attn_metadata(
8485
num_reqs=num_reqs_padded,
8586
num_reqs_padded=num_reqs_padded,
8687
num_tokens_padded=num_tokens_padded,
88+
seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound,
89+
step=self.num_query_per_req,
8790
causal=self._group_causal,
8891
)
8992
return [attn_metadata]
@@ -107,6 +110,11 @@ def propose(
107110
mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None,
108111
is_profile: bool = False,
109112
) -> torch.Tensor:
113+
# TODO: Remove the if not vllm_version_is guard when dropping v0.25.1.
114+
# input_batch is cached here so that DFlashAclGraphManager can access
115+
# seq_lens_cpu_upper_bound during full-graph replay.
116+
if not vllm_version_is("0.25.1"):
117+
self.input_batch = input_batch
110118
with build_attn_metadata_wrapper():
111119
return super().propose(
112120
input_batch,

vllm_ascend/worker/v2/spec_decode/dspark/speculator.py

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
DSparkSpeculator,
2929
)
3030

31+
from vllm_ascend.utils import vllm_version_is
3132
from vllm_ascend.worker.v2.attn_utils import build_attn_metadata_wrapper
3233

3334

@@ -36,6 +37,11 @@ class AscendDSparkSpeculator(DSparkSpeculator):
3637

3738
def __init__(self, vllm_config: VllmConfig, device: torch.device):
3839
super().__init__(vllm_config, device)
40+
# TODO: Remove the if not vllm_version_is guard when dropping v0.25.1.
41+
# input_batch is cached in propose() so that DFlashAclGraphManager
42+
# can access seq_lens_cpu_upper_bound during full-graph replay.
43+
if not vllm_version_is("0.25.1"):
44+
self.input_batch: InputBatch | None = None
3945

4046
# we need to update full graph params in run_fullgraph,
4147
# so create a stream to update full graph params.
@@ -82,13 +88,15 @@ def set_attn(
8288

8389
self.attn_backends = attn_backends
8490

85-
def build_draft_attn_metadatas(self, num_reqs_padded):
91+
def build_draft_attn_metadatas(self, num_reqs_padded, seq_lens_cpu_upper_bound):
8692
num_tokens_padded = num_reqs_padded * self.num_query_per_req
8793
with build_attn_metadata_wrapper():
8894
attn_metadata = self._build_draft_attn_metadata(
8995
num_reqs=num_reqs_padded,
9096
num_reqs_padded=num_reqs_padded,
9197
num_tokens_padded=num_tokens_padded,
98+
seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound,
99+
step=self.num_query_per_req,
92100
causal=self._group_causal,
93101
)
94102
return [attn_metadata]
@@ -112,6 +120,11 @@ def propose(
112120
mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None,
113121
is_profile: bool = False,
114122
) -> torch.Tensor:
123+
# TODO: Remove the if not vllm_version_is guard when dropping v0.25.1.
124+
# input_batch is cached here so that DFlashAclGraphManager can access
125+
# seq_lens_cpu_upper_bound during full-graph replay.
126+
if not vllm_version_is("0.25.1"):
127+
self.input_batch = input_batch
115128
with build_attn_metadata_wrapper():
116129
return super().propose(
117130
input_batch,

vllm_ascend/worker/v2/spec_decode/eagle/aclgraph.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
set_draft_graph_prefill_params,
2727
update_full_graph_params,
2828
)
29+
from vllm_ascend.utils import vllm_version_is
2930
from vllm_ascend.worker.v2.aclgraph_utils import collect_sorted_captured_token_sizes, model_capture_wrapper
3031
from vllm_ascend.worker.v2.utils import communicator_switch
3132

@@ -111,11 +112,16 @@ def create_forward_fn(desc: BatchExecutionDescriptor, warmup: bool):
111112
kv_cache_config,
112113
skip_attn=(desc.cg_mode == CUDAGraphMode.PIECEWISE),
113114
)
115+
if vllm_version_is("0.25.1"):
116+
seq_lens_cpu_upper_bound = None
117+
else:
118+
seq_lens_cpu_upper_bound = input_buffers.seq_lens_cpu[:num_reqs]
114119
return lambda cg_mode: forward_fn(
115120
num_reqs,
116121
cg_mode == CUDAGraphMode.PIECEWISE,
117122
BatchExecutionDescriptor(cg_mode=cg_mode, num_tokens=num_tokens, num_reqs=num_reqs),
118123
num_tokens_across_dp,
124+
seq_lens_cpu_upper_bound,
119125
)
120126

121127
CudaGraphManager.capture(self, create_forward_fn, progress_bar_desc=progress_bar_desc)

vllm_ascend/worker/v2/spec_decode/mtp/speculator.py

Lines changed: 18 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -23,9 +23,11 @@
2323
from contextlib import contextmanager
2424
from typing import Any
2525

26+
import torch
2627
import vllm.v1.worker.gpu.spec_decode.speculator as _upstream_speculator
2728
from vllm.v1.worker.gpu.spec_decode.mtp.speculator import MTPSpeculator
2829

30+
from vllm_ascend.utils import vllm_version_is
2931
from vllm_ascend.worker.v2.spec_decode.autoregressive.speculator import AscendAutoRegressiveSpeculator
3032

3133

@@ -76,6 +78,8 @@ def _build_draft_attn_metadata(
7678
num_reqs: int,
7779
num_reqs_padded: int,
7880
num_tokens_padded: int,
81+
seq_lens_cpu_upper_bound: torch.Tensor | None = None,
82+
step: int = 1,
7983
num_query_per_req: int = 1,
8084
causal: bool = True,
8185
) -> dict[str, Any] | None:
@@ -84,6 +88,17 @@ def _build_draft_attn_metadata(
8488
# super() path does not forward them. Wrap build_attn_metadata to inject
8589
# positions[:num_tokens_padded] and reuse super() (no arg duplication).
8690
with build_position_wrapper(self.input_buffers.positions, num_tokens_padded):
87-
return super()._build_draft_attn_metadata(
88-
num_reqs, num_reqs_padded, num_tokens_padded, num_query_per_req, causal
89-
)
91+
if vllm_version_is("0.25.1"):
92+
return super()._build_draft_attn_metadata(
93+
num_reqs, num_reqs_padded, num_tokens_padded, num_query_per_req, causal
94+
)
95+
else:
96+
return super()._build_draft_attn_metadata(
97+
num_reqs,
98+
num_reqs_padded,
99+
num_tokens_padded,
100+
seq_lens_cpu_upper_bound,
101+
step,
102+
num_query_per_req,
103+
causal,
104+
)

0 commit comments

Comments
 (0)