Skip to content

Commit c6343b1

Browse files
SunnyLee151064Your Name
andauthored
[Feature] qwen3.5 support pcp+mtp (vllm-project#10487)
### What this PR does / why we need it? This PR introduces support for hybrid attention PCP (Parallel Context Processing) by adding linear-splitting of prefill inputs (`_split_pcp_input_hybrid`) and updating the model runner to handle hybrid attention logits indices. It also corrects several index offsets in `pcp_utils.py` (switching between decode requests and decode tokens) and updates Triton operations for Mamba and FLA chunk-gated delta rules. However, several critical issues were identified during the review: - In `causal_conv1d.py`, `num_prefills` is initialized to `0` but never updated from `attn_metadata`, which breaks the PCP + MTP logic. - In `model_runner_v1.py`, calculating `num_scheduled_tokens` via slice differences misses the first request's token count, and an out-of-bounds/incorrect index access occurs when `num_decode_reqs` is `0`. - In `chunk.py`, calling `.item()` on an NPU tensor triggers an expensive device-to-host synchronization in the hot path. ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? No testing details were provided in the patch. - vLLM version: v0.22.1 - vLLM main: vllm-project/vllm@967c5c3 --------- Signed-off-by: Your Name <you@example.com> Co-authored-by: Your Name <you@example.com>
1 parent 6e9fabb commit c6343b1

10 files changed

Lines changed: 435 additions & 33 deletions

File tree

tests/ut/spec_decode/a2/test_eagle_proposer.py

Lines changed: 177 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3825,6 +3825,7 @@ def test_set_inputs_first_pass_pcp_dcp_mixed(self):
38253825
self.runner.input_batch.req_ids = req_ids
38263826
# maybe not reasonable just to run test
38273827
self.runner.logits_indices = torch.arange(12, dtype=torch.int32, device=self.device)
3828+
self.runner.pcp_manager.pcp_use_hybrid_attn = False
38283829

38293830
proposer, vllm_config = self._create_proposer(
38303831
method="eagle",
@@ -4205,3 +4206,179 @@ def test_set_inputs_first_pass_draft_model(self):
42054206

42064207
for attrition in attrs_from_cad:
42074208
assert_attr_equal(attrition, expected_cad, out_cad)
4209+
4210+
4211+
def _build_split_pcp_input_hybrid_inputs(req_scheduled_tokens: dict[str, int], hidden_size: int):
4212+
"""Build input_ids and target_hidden_states where row i is filled with i.
4213+
4214+
Using the row index as the value makes it easy to assert which original
4215+
tokens survive the per-rank slice and which slots are padding.
4216+
"""
4217+
total_tokens = sum(req_scheduled_tokens.values())
4218+
input_ids = torch.arange(total_tokens, dtype=torch.int32)
4219+
target_hidden_states = torch.arange(total_tokens, dtype=torch.float32).unsqueeze(-1).repeat(1, hidden_size)
4220+
return input_ids, target_hidden_states
4221+
4222+
4223+
# yapf: disable
4224+
@pytest.mark.parametrize(
4225+
"pcp_size, pcp_rank, req_scheduled_tokens, hidden_size,"
4226+
" expected_num_tokens, expected_input_ids, expected_hidden_first_col,"
4227+
" expected_seq_lens, expected_cu_num_tokens, expected_max_query_len",
4228+
[
4229+
# Case 1: single req, perfectly aligned to 2*pcp_size.
4230+
# ori=8, pcp=2 -> padded=8, pcp_tokens=4
4231+
# rank 0: [0,1,2,3], rank 1: [4,5,6,7]
4232+
(
4233+
2, 0, {"0": 8}, 4,
4234+
4, [0, 1, 2, 3], [0.0, 1.0, 2.0, 3.0],
4235+
[4], [0, 4], 4,
4236+
),
4237+
(
4238+
2, 1, {"0": 8}, 4,
4239+
4, [4, 5, 6, 7], [4.0, 5.0, 6.0, 7.0],
4240+
[4], [0, 4], 4,
4241+
),
4242+
# Case 2: single req, needs padding.
4243+
# ori=7, pcp=2 -> padded=8, pcp_tokens=4, num_pads=1
4244+
# rank 0: [0,1,2,3] (all valid)
4245+
# rank 1: [4,5,6,PAD] -> input_ids=[4,5,6,0], hidden_first_col=[4,5,6,0]
4246+
(
4247+
2, 0, {"0": 7}, 4,
4248+
4, [0, 1, 2, 3], [0.0, 1.0, 2.0, 3.0],
4249+
[4], [0, 4], 4,
4250+
),
4251+
(
4252+
2, 1, {"0": 7}, 4,
4253+
4, [4, 5, 6, 0], [4.0, 5.0, 6.0, 0.0],
4254+
[4], [0, 4], 4,
4255+
),
4256+
# Case 3: multiple reqs with different padding needs.
4257+
# req 0: ori=4 -> padded=4, pcp_tokens=2
4258+
# req 1: ori=6 -> padded=8, pcp_tokens=4, num_pads=2
4259+
# rank 0: [0,1] + [4,5,6,7]
4260+
# rank 1: [2,3] + [8,9,PAD,PAD]
4261+
(
4262+
2, 0, {"0": 4, "1": 6}, 2,
4263+
6, [0, 1, 4, 5, 6, 7], [0.0, 1.0, 4.0, 5.0, 6.0, 7.0],
4264+
[2, 4], [0, 2, 6], 4,
4265+
),
4266+
(
4267+
2, 1, {"0": 4, "1": 6}, 2,
4268+
6, [2, 3, 8, 9, 0, 0], [2.0, 3.0, 8.0, 9.0, 0.0, 0.0],
4269+
[2, 4], [0, 2, 6], 4,
4270+
),
4271+
# Case 4: pcp_size=4
4272+
# ori=9, pcp=4 -> padded=16, pcp_tokens=4, num_pads=7
4273+
# rank 0: tokens [0,1,2,3]
4274+
# rank 1: tokens [4,5,6,7]
4275+
# rank 2: tokens [8,PAD,PAD,PAD]
4276+
# rank 3: tokens [PAD,PAD,PAD,PAD]
4277+
(
4278+
4, 0, {"0": 9}, 2,
4279+
4, [0, 1, 2, 3], [0.0, 1.0, 2.0, 3.0],
4280+
[4], [0, 4], 4,
4281+
),
4282+
(
4283+
4, 1, {"0": 9}, 2,
4284+
4, [4, 5, 6, 7], [4.0, 5.0, 6.0, 7.0],
4285+
[4], [0, 4], 4,
4286+
),
4287+
(
4288+
4, 2, {"0": 9}, 2,
4289+
4, [8, 0, 0, 0], [8.0, 0.0, 0.0, 0.0],
4290+
[4], [0, 4], 4,
4291+
),
4292+
(
4293+
4, 3, {"0": 9}, 2,
4294+
4, [0, 0, 0, 0], [0.0, 0.0, 0.0, 0.0],
4295+
[4], [0, 4], 4,
4296+
),
4297+
# Case 5: minimal request - single token.
4298+
# ori=1, pcp=2 -> padded=4, pcp_tokens=2, num_pads=3
4299+
# rank 0: [0, PAD]
4300+
# rank 1: [PAD, PAD]
4301+
(
4302+
2, 0, {"0": 1}, 2,
4303+
2, [0, 0], [0.0, 0.0],
4304+
[2], [0, 2], 2,
4305+
),
4306+
(
4307+
2, 1, {"0": 1}, 2,
4308+
2, [0, 0], [0.0, 0.0],
4309+
[2], [0, 2], 2,
4310+
),
4311+
],
4312+
)
4313+
# yapf: enable
4314+
def test_split_pcp_input_hybrid(
4315+
pcp_size,
4316+
pcp_rank,
4317+
req_scheduled_tokens,
4318+
hidden_size,
4319+
expected_num_tokens,
4320+
expected_input_ids,
4321+
expected_hidden_first_col,
4322+
expected_seq_lens,
4323+
expected_cu_num_tokens,
4324+
expected_max_query_len,
4325+
):
4326+
input_ids, target_hidden_states = _build_split_pcp_input_hybrid_inputs(
4327+
req_scheduled_tokens, hidden_size
4328+
)
4329+
4330+
# _split_pcp_input_hybrid only reads self.pcp_size and self.pcp_rank,
4331+
# so a MagicMock with just those attributes is enough to drive the
4332+
# unbound method without instantiating the full proposer.
4333+
mock_self = MagicMock()
4334+
mock_self.pcp_size = pcp_size
4335+
mock_self.pcp_rank = pcp_rank
4336+
4337+
(
4338+
num_tokens,
4339+
out_input_ids,
4340+
out_hidden_states,
4341+
max_query_len,
4342+
seq_lens,
4343+
cu_num_tokens,
4344+
) = llm_base_proposer.AscendSpecDecodeBaseProposer._split_pcp_input_hybrid(
4345+
mock_self, req_scheduled_tokens, input_ids, target_hidden_states
4346+
)
4347+
4348+
assert num_tokens == expected_num_tokens
4349+
assert max_query_len == expected_max_query_len
4350+
assert torch.equal(
4351+
out_input_ids, torch.tensor(expected_input_ids, dtype=torch.int32)
4352+
)
4353+
assert torch.equal(seq_lens, torch.tensor(expected_seq_lens, dtype=torch.int32))
4354+
assert torch.equal(
4355+
cu_num_tokens, torch.tensor(expected_cu_num_tokens, dtype=torch.int64)
4356+
)
4357+
4358+
# hidden_states shape: [num_tokens, hidden_size]
4359+
assert out_hidden_states.shape == (expected_num_tokens, hidden_size)
4360+
assert torch.equal(
4361+
out_hidden_states[:, 0],
4362+
torch.tensor(expected_hidden_first_col, dtype=torch.float32),
4363+
)
4364+
4365+
4366+
def test_split_pcp_input_hybrid_preserves_hidden_size():
4367+
"""Hidden states must come back with the same hidden dim as the input."""
4368+
hidden_size = 7
4369+
req_scheduled_tokens = {"0": 6}
4370+
input_ids, target_hidden_states = _build_split_pcp_input_hybrid_inputs(
4371+
req_scheduled_tokens, hidden_size
4372+
)
4373+
4374+
mock_self = MagicMock()
4375+
mock_self.pcp_size = 2
4376+
mock_self.pcp_rank = 0
4377+
4378+
_, _, out_hidden_states, _, _, _ = (
4379+
llm_base_proposer.AscendSpecDecodeBaseProposer._split_pcp_input_hybrid(
4380+
mock_self, req_scheduled_tokens, input_ids, target_hidden_states
4381+
)
4382+
)
4383+
4384+
assert out_hidden_states.shape[1] == hidden_size

tests/ut/worker/test_pcp_manager.py

Lines changed: 114 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
# limitations under the License.
1212
# This file is a part of the vllm-ascend project.
1313

14+
import math
1415
from unittest.mock import MagicMock
1516

1617
import numpy as np
@@ -424,3 +425,116 @@ def test_generate_pcp_mtp_input(
424425
target_input_ids_pcp_full)
425426
assert torch.equal(pcp_manager.query_start_loc_pcp_full.cpu[:num_reqs + 1],
426427
target_query_start_loc_pcp_full)
428+
429+
430+
def _realistic_num_pcp_pads(
431+
num_scheduled_tokens: list[int],
432+
pcp_world_size: int,
433+
num_decode_reqs: int,
434+
) -> list[int]:
435+
"""Compute num_pcp_pads exactly as PCPManager.update_tokens_for_pcp would.
436+
437+
Decode reqs duplicate tokens across pcp_world_size ranks, so their pad
438+
count is num_tokens * (pcp_world_size - 1). Prefill reqs are padded up
439+
to a multiple of 2 * pcp_world_size, so their pad count is the difference
440+
between the padded and original length.
441+
"""
442+
pads: list[int] = []
443+
pad_multiple = 2 * pcp_world_size
444+
for i, n in enumerate(num_scheduled_tokens):
445+
if i < num_decode_reqs:
446+
pads.append(n * (pcp_world_size - 1))
447+
else:
448+
padded = math.ceil(n / pad_multiple) * pad_multiple
449+
pads.append(padded - n)
450+
return pads
451+
452+
453+
# yapf: disable
454+
@pytest.mark.parametrize(
455+
"pcp_size, num_scheduled_tokens, num_decode_reqs,"
456+
" expected_cu_num_scheduled_tokens",
457+
[
458+
# Case 1: no prefill reqs -> returned unchanged.
459+
(2, [1, 1], 2, [1, 2]),
460+
# Case 2: prefill only (num_decode_reqs == 0). Exercises the case where
461+
# the diff-based per-req length derivation would otherwise drop the
462+
# first req's length.
463+
# num_scheduled=[3, 5], realistic pads=[1, 3] -> cumsum=[1, 4]
464+
# cu=[3, 8]; padded prefill_lens=[4, 8]; base=0
465+
# prefill_cu=[4, 12]; final = [4*2-1, 12*2-4] = [7, 20]
466+
(2, [3, 5], 0, [7, 20]),
467+
# Case 3: mix decode + prefill, pcp_size=2.
468+
# num_scheduled=[1, 1, 3, 5], realistic pads=[1, 1, 1, 3]
469+
# prefill pads cumsum=[1, 4]
470+
# cu=[1, 2, 5, 10]; padded prefill_lens=[4, 8]; base=cu[1]=2
471+
# prefill_cu=[6, 14]; final[2:] = [12, 28] - [1, 4] = [11, 24]
472+
(2, [1, 1, 3, 5], 2, [1, 2, 11, 24]),
473+
# Case 4: pcp_size=4, mix decode + prefill with uneven prefill tokens.
474+
# num_scheduled=[1, 1, 5, 9], realistic pads=[3, 3, 3, 7]
475+
# prefill pads cumsum=[3, 10]
476+
# cu=[1, 2, 7, 16]; padded prefill_lens=[8, 16]; base=cu[1]=2
477+
# prefill_cu=[10, 26]; final[2:] = [40, 104] - [3, 10] = [37, 94]
478+
(4, [1, 1, 5, 9], 2, [1, 2, 37, 94]),
479+
# Case 5: single prefill req, pcp_size=2.
480+
# num_scheduled=[7], realistic pads=[1]
481+
# cu=[7]; padded prefill_lens=[8]; base=0
482+
# prefill_cu=[8]; final = [8*2-1] = [15]
483+
(2, [7], 0, [15]),
484+
# Case 6: prefill req that's already aligned to 2*pcp_size.
485+
# num_scheduled=[8], realistic pads=[0]
486+
# cu=[8]; padded prefill_lens=[8]; base=0
487+
# prefill_cu=[8]; final = [8*2-0] = [16]
488+
(2, [8], 0, [16]),
489+
],
490+
)
491+
# yapf: enable
492+
def test_adjust_cu_num_scheduled_tokens_for_pcp(
493+
pcp_size,
494+
num_scheduled_tokens,
495+
num_decode_reqs,
496+
expected_cu_num_scheduled_tokens,
497+
):
498+
vllm_config = MagicMock()
499+
vllm_config.model_config = MagicMock()
500+
vllm_config.speculative_config.num_speculative_tokens = 0
501+
502+
pcp_manager = PCPManager(
503+
pcp_world_size=pcp_size,
504+
pcp_rank=0,
505+
dcp_world_size=1,
506+
dcp_rank=0,
507+
max_buffer_num_tokens=10000,
508+
max_num_reqs=1000,
509+
device="cpu",
510+
vllm_config=vllm_config,
511+
use_async_scheduling=False,
512+
pin_memory=False,
513+
)
514+
515+
num_reqs = len(num_scheduled_tokens)
516+
cu_num_scheduled_tokens = np.cumsum(
517+
np.array(num_scheduled_tokens, dtype=np.int32)
518+
)
519+
# Use realistic pads that match what update_tokens_for_pcp would produce
520+
# for the given num_scheduled_tokens / pcp_world_size combination.
521+
# Tests previously fed zeros for prefill pads, which made the math look
522+
# correct but never exercised the real runtime path.
523+
num_pcp_pads = np.array(
524+
_realistic_num_pcp_pads(num_scheduled_tokens, pcp_size, num_decode_reqs),
525+
dtype=np.int32,
526+
)
527+
528+
# Seed the manager state normally populated by init_batch_info.
529+
pcp_manager.num_decode_reqs = num_decode_reqs
530+
pcp_manager.num_prefill_reqs = num_reqs - num_decode_reqs
531+
532+
result = pcp_manager.adjust_cu_num_scheduled_tokens_for_pcp(
533+
cu_num_scheduled_tokens, num_pcp_pads
534+
)
535+
536+
assert np.array_equal(
537+
result, np.array(expected_cu_num_scheduled_tokens, dtype=np.int32)
538+
), (
539+
f"Expected {expected_cu_num_scheduled_tokens}, got {result.tolist()}"
540+
)

vllm_ascend/attention/context_parallel/attention_cp.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -252,9 +252,9 @@ def build(
252252
num_decodes:
253253
]
254254
else:
255-
actual_seq_lengths_q = [self.decode_threshold * (i + 1) for i in range(num_decodes)] + query_start_loc_cpu[
256-
num_decodes + 1 :
257-
].tolist()
255+
actual_seq_lengths_q = (
256+
query_start_loc_cpu[1 : num_decodes + 1].tolist() + query_start_loc_cpu[num_decodes + 1 :].tolist()
257+
)
258258

259259
attn_metadata = AscendMetadata(
260260
num_actual_tokens=num_actual_tokens,
@@ -898,7 +898,9 @@ def _gather_and_restore_pcp_qkv(
898898
decode_offset = attn_metadata.num_decode_tokens * self.pcp_size
899899
qkv_fa_padding_workspace[:decode_offset] = actual_qkv[:decode_offset]
900900

901-
pcp_unpad_mask = attn_metadata.prefill.pcp_metadata.pcp_unpad_mask[attn_metadata.num_decodes * self.pcp_size :]
901+
pcp_unpad_mask = attn_metadata.prefill.pcp_metadata.pcp_unpad_mask[
902+
attn_metadata.num_decode_tokens * self.pcp_size :
903+
]
902904
qkv_fa_padding_workspace[decode_offset:][pcp_unpad_mask] = actual_qkv[decode_offset:]
903905
qkv_fa_padding_workspace[decode_offset:][~pcp_unpad_mask] = 0
904906

vllm_ascend/ops/gdn_attn_builder.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -753,17 +753,18 @@ def __init__(
753753
self.vllm_config.scheduler_config.max_num_seqs,
754754
self.decode_cudagraph_max_bs,
755755
)
756+
757+
self.spec_sequence_masks: torch.Tensor = torch.empty(
758+
(sequence_index_capacity,), dtype=torch.bool, device=device
759+
)
760+
756761
self.spec_sequence_masks_cpu: torch.Tensor = torch.empty(
757762
(sequence_index_capacity,),
758763
dtype=torch.bool,
759764
device="cpu",
760765
pin_memory=device.type != "cpu",
761766
)
762-
self.spec_sequence_masks: torch.Tensor = torch.empty(
763-
(sequence_index_capacity,),
764-
dtype=torch.bool,
765-
device=device,
766-
)
767+
767768
self.spec_sequence_indices_cpu: torch.Tensor = torch.empty(
768769
(sequence_index_capacity,),
769770
dtype=torch.int64,
@@ -796,7 +797,7 @@ def _init_reorder_batch_threshold(
796797
super()._init_reorder_batch_threshold(
797798
reorder_batch_threshold,
798799
supports_spec_as_decode,
799-
supports_dcp_with_varlen,
800+
True,
800801
)
801802
if self.reorder_batch_threshold != 1: # type: ignore
802803
speculative_config = self.vllm_config.speculative_config
@@ -1156,6 +1157,7 @@ def build( # type: ignore[override]
11561157
and num_spec_decode_tokens <= self.decode_cudagraph_max_bs
11571158
):
11581159
assert spec_sequence_masks is not None
1160+
self.spec_state_indices_tensor[batch_size:].fill_(NULL_BLOCK_ID)
11591161
self.spec_state_indices_tensor[:num_spec_decodes].copy_(
11601162
spec_state_indices_tensor,
11611163
non_blocking=True,

vllm_ascend/ops/triton/fla/chunk.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -151,7 +151,12 @@ def chunk_gated_delta_rule_fwd(
151151

152152
if get_pcp_group().rank_in_group > 0:
153153
rerun_initial_state = initial_state.clone()
154-
prefill_slice = slice(num_decodes, final_state.shape[0])
154+
if cu_seqlens is not None:
155+
_ns_lens = cu_seqlens[1:] - cu_seqlens[:-1]
156+
prefill_seq_offset = int(((_ns_lens > 0) & (_ns_lens <= 1)).sum().item())
157+
else:
158+
prefill_seq_offset = num_decodes
159+
prefill_slice = slice(prefill_seq_offset, final_state.shape[0])
155160
rerun_initial_state[prefill_slice] = updated_h_state[prefill_slice]
156161
h, v_new, _ = chunk_gated_delta_rule_fwd_h(
157162
k=k,

0 commit comments

Comments
 (0)