Skip to content

Commit 367b8e6

Browse files
authored
[BugFix][v0.20.2rc]Reduce sampling is reconstructed to eliminate all patch behaviors and support DFlash and MTP (#9946)
<!-- Thanks for sending a pull request! BEFORE SUBMITTING, PLEASE READ https://docs.vllm.ai/en/latest/contributing/overview.html --> ### What this PR does / why we need it? <!-- - Please clarify what changes you are proposing. The purpose of this section is to outline the changes and how this PR fixes the issue. If possible, please consider writing useful notes for better and faster reviews in your PR. - Please clarify why the changes are needed. For instance, the use case and bug description. - Fixes # --> The Reduce_sampling optimization is reconstructed to eliminate all patch behaviors and support the DFlash and MTP. When sampling optimization is enabled, if speculative decoding is Eagle3 or DFlash, sampling can be optimized for both the main model and the MTP layer. If the speculative decoding method is MTP, some models can optimize sampling for both the main model and the MTP layer, while others can only optimize sampling for the main model. ### Does this PR introduce _any_ user-facing change? <!-- Note that it means *any* user-facing change including all aspects such as API, interface or other behavior changes. Documentation-only updates are not considered user-facing changes. --> ### How was this patch tested? <!-- CI passed with new added/existing test. If it was tested in a way different from regular unit tests, please clarify how you tested step by step, ideally copy and paste-able, so that other reviewers can test and check, and descendants can verify in the future. If tests were not added, please describe why they were not added and/or why it was difficult to add. --> --------- Signed-off-by: hzx55906 <513464215@qq.com>
1 parent c942e37 commit 367b8e6

8 files changed

Lines changed: 265 additions & 156 deletions

File tree

‎tests/ut/spec_decode/test_eagle_proposer.py‎

Lines changed: 14 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -2161,9 +2161,8 @@ def test_intermediate_variables_precision(self):
21612161
class MockDraftModel:
21622162
"""Draft model that records prepared forward inputs."""
21632163

2164-
def __init__(self, returns_tuple=True, enable_reduce_sample=False, vocab_size=200000):
2164+
def __init__(self, returns_tuple=True, vocab_size=200000):
21652165
self.returns_tuple = returns_tuple
2166-
self.enable_reduce_sample = enable_reduce_sample
21672166
self.vocab_size = vocab_size
21682167
self.calls = []
21692168
self.logit_inputs = []
@@ -2187,13 +2186,11 @@ def __call__(self, **kwargs):
21872186
return last_hidden_states, hidden_states
21882187
return last_hidden_states
21892188

2190-
def compute_logits(self, sample_hidden_states, enable_reduce_sample=False):
2189+
def compute_logits(self, sample_hidden_states):
21912190
self.logit_inputs.append(sample_hidden_states.clone())
21922191
token_ids = sample_hidden_states[:, 0].to(torch.long)
21932192
logits = torch.full((sample_hidden_states.shape[0], self.vocab_size), -1000.0)
21942193
logits[torch.arange(sample_hidden_states.shape[0]), token_ids] = 1000.0
2195-
if self.enable_reduce_sample:
2196-
logits = logits.argmax(dim=-1)
21972194
return logits
21982195

21992196
def embed_input_ids(self, input_ids):
@@ -2417,7 +2414,7 @@ def check_mock(self):
24172414
assert sig_name == ["self", "input_ids", "positions", "hidden_states", "inputs_embeds"]
24182415
sig = inspect.signature(RunnerCls.compute_logits)
24192416
sig_name = self.get_param_names(sig)
2420-
assert sig_name == ["self", "hidden_states", "enable_reduce_sample"]
2417+
assert sig_name == ["self", "hidden_states"]
24212418

24222419
import vllm_ascend.ascend_forward_context
24232420

@@ -2513,7 +2510,17 @@ def get_param_names(self, sig):
25132510
return [p.name for p in sig.parameters.values()]
25142511

25152512
def test_run_merged_draft_eagle3_decode_prepares_each_forward_input(self):
2516-
self.proposer.model = MockDraftModel(returns_tuple=True, enable_reduce_sample= True)
2513+
self.proposer.model = MockDraftModel(returns_tuple=True)
2514+
2515+
def compute_draft_token_ids(sample_hidden_states):
2516+
self.proposer.model.logit_inputs.append(sample_hidden_states.clone())
2517+
token_ids = sample_hidden_states[:, 0].to(torch.long)
2518+
logits = torch.full((sample_hidden_states.shape[0], self.proposer.model.vocab_size), -1000.0)
2519+
logits[torch.arange(sample_hidden_states.shape[0]), token_ids] = 1000.0
2520+
logits = logits.argmax(dim=-1)
2521+
return logits
2522+
2523+
self.proposer.compute_draft_token_ids = compute_draft_token_ids
25172524
self.proposer.supports_mm_inputs = True
25182525
initial_input_ids = torch.tensor(
25192526
[279, 1196, 374, 8014, 151667, 198, 32313, 11, 151667, 198, 32313, 11],

‎vllm_ascend/ops/triton/bincount.py‎

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
from vllm.distributed.parallel_state import get_tp_group
2525
from vllm.triton_utils import tl, triton
2626

27+
from vllm_ascend.ascend_config import get_ascend_config
2728
from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num
2829

2930

@@ -80,10 +81,10 @@ def token_bin_counts_and_mask_kernel(
8081
)
8182

8283
local_token = token - vocab_start_idx
83-
8484
token_in_range = pos_mask & (token >= vocab_start_idx) & (local_token < vocab_size)
8585

86-
count_ptr = batch_counts_start + local_token * counts_vocab_stride
86+
safe_local_token = tl.where(token_in_range, local_token, 0)
87+
count_ptr = batch_counts_start + safe_local_token * counts_vocab_stride
8788
tl.atomic_add(count_ptr, 1, mask=token_in_range)
8889

8990

@@ -129,8 +130,11 @@ def get_token_bin_counts_and_mask_triton(
129130
progs = min(core_num, total_blocks)
130131
grid = (progs, triton.cdiv(total_blocks, progs))
131132

132-
tp_group = get_tp_group()
133-
tp_rank = tp_group.rank_in_group
133+
if get_ascend_config().enable_reduce_sample:
134+
tp_group = get_tp_group()
135+
tp_rank = tp_group.rank_in_group
136+
else:
137+
tp_rank = 0
134138
token_bin_counts_and_mask_kernel[grid](
135139
tokens,
136140
tokens.stride(0),

‎vllm_ascend/ops/triton/reject_sample.py‎

Lines changed: 161 additions & 64 deletions
Original file line numberDiff line numberDiff line change
@@ -179,10 +179,10 @@ def rejection_random_sample_kernel(
179179

180180
for pos in range(num_draft_tokens):
181181
if not rejected:
182-
token_idx = start_idx + pos
183-
draft_token_id = tl.load(draft_token_ids_ptr + token_idx)
184-
185182
if ENABLE_REDUCE_SAMPLING:
183+
token_idx = start_idx + pos
184+
draft_token_id = tl.load(draft_token_ids_ptr + token_idx)
185+
186186
target_prob = 0.0
187187
found = False
188188

@@ -209,30 +209,47 @@ def rejection_random_sample_kernel(
209209
if current_match_prob > 0.0:
210210
target_prob = current_match_prob
211211
found = True
212-
else:
213-
target_prob = tl.load(target_probs_ptr + token_idx * vocab_size + draft_token_id)
214212

215-
if NO_DRAFT_PROBS:
216-
draft_prob = 1.0
217-
else:
218-
vocab_for_draft = global_vocab_size if ENABLE_REDUCE_SAMPLING else vocab_size
219-
draft_prob = tl.load(draft_probs_ptr + token_idx * vocab_for_draft + draft_token_id)
213+
if NO_DRAFT_PROBS:
214+
draft_prob = 1
215+
else:
216+
draft_prob = tl.load(draft_probs_ptr + token_idx * global_vocab_size + draft_token_id)
220217

221-
uniform_prob = tl.load(uniform_probs_ptr + token_idx)
218+
uniform_prob = tl.load(uniform_probs_ptr + token_idx)
222219

223-
# Acceptance condition
224-
if draft_prob > 0 and target_prob / draft_prob >= uniform_prob:
225-
# Accept
226-
token_id = draft_token_id
227-
else:
228-
# Reject - use recovered token
229-
rejected = True
230-
token_id = tl.load(recovered_token_ids_ptr + token_idx)
220+
# Acceptance condition
221+
if draft_prob > 0 and target_prob / draft_prob >= uniform_prob:
222+
# Accept
223+
token_id = draft_token_id
224+
else:
225+
# Reject - use recovered token
226+
rejected = True
227+
token_id = tl.load(recovered_token_ids_ptr + token_idx)
231228

232-
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + pos, token_id)
229+
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + pos, token_id)
230+
else:
231+
draft_token_id = tl.load(draft_token_ids_ptr + start_idx + pos)
232+
target_prob = tl.load(target_probs_ptr + (start_idx + pos) * global_vocab_size + draft_token_id)
233+
if NO_DRAFT_PROBS:
234+
draft_prob = 1
235+
else:
236+
draft_prob = tl.load(
237+
draft_probs_ptr + (start_idx + pos) * global_vocab_size + draft_token_id
238+
)
239+
uniform_prob = tl.load(uniform_probs_ptr + start_idx + pos)
240+
# NOTE(woosuk): While the draft probability should never be 0,
241+
# we check it to avoid NaNs. If it happens to be 0, we reject.
242+
if draft_prob > 0 and target_prob / draft_prob >= uniform_prob:
243+
# Accept.
244+
token_id = draft_token_id
245+
else:
246+
# Reject. Use recovered token.
247+
rejected = True
248+
token_id = tl.load(recovered_token_ids_ptr + start_idx + pos)
249+
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + pos, token_id)
233250

234251
if not rejected:
235-
# All tokens accepted - append bonus token
252+
# If all tokens are accepted, append the bonus token.
236253
bonus_token_id = tl.load(bonus_token_ids_ptr + req_idx)
237254
tl.store(
238255
output_token_ids_ptr + req_idx * (max_spec_len + 1) + num_draft_tokens,
@@ -470,38 +487,40 @@ def rejection_random_sample_block_verify_kernel(
470487
NO_DRAFT_PROBS: tl.constexpr,
471488
ENABLE_REDUCE_SAMPLING: tl.constexpr, # Whether using reduce_sampling
472489
BLOCK_SIZE: tl.constexpr,
473-
VOCAB_BLOCK_SIZE: tl.constexpr = 512,
490+
SUB_BLOCK: tl.constexpr = 512,
474491
):
475492
block_idx = tl.program_id(0)
476493
offsets = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
477494
mask = offsets < vec_len
478495
is_greedy = tl.load(is_greedy_ptr + offsets, mask, other=1)
479496
not_greedy_mask = is_greedy == 0
480-
start_idxs = tl.where(offsets == 0, 0, tl.load(cu_num_draft_tokens_ptr + offsets - 1, not_greedy_mask))
497+
prev_mask = not_greedy_mask & (offsets > 0)
498+
prev_end_idxs = tl.load(cu_num_draft_tokens_ptr + offsets - 1, prev_mask, other=0)
499+
start_idxs = tl.where(offsets == 0, 0, prev_end_idxs)
481500
end_idxs = tl.load(cu_num_draft_tokens_ptr + offsets, not_greedy_mask)
482501
n_num_draft_tokens = end_idxs - start_idxs
483502

484-
for req_i in range(BLOCK_SIZE):
485-
not_greedy = get_element(not_greedy_mask, (req_i,))
486-
if not_greedy:
487-
pi = 1.0
488-
uniform_prob = 1.0
489-
last_accepted_token_pos = -1
490-
start_idx = get_element(start_idxs, (req_i,))
491-
req_idx = block_idx * BLOCK_SIZE + req_i
492-
num_draft_tokens = get_element(n_num_draft_tokens, (req_i,))
493-
494-
for pos in range(num_draft_tokens):
495-
token_idx = start_idx + pos
496-
draft_token_id = tl.load(draft_token_ids_ptr + token_idx)
503+
if ENABLE_REDUCE_SAMPLING:
504+
for req_i in range(BLOCK_SIZE):
505+
not_greedy = get_element(not_greedy_mask, (req_i,))
506+
if not_greedy:
507+
pi = 1.0
508+
uniform_prob = 1.0
509+
last_accepted_token_pos = -1
510+
start_idx = get_element(start_idxs, (req_i,))
511+
req_idx = block_idx * BLOCK_SIZE + req_i
512+
num_draft_tokens = get_element(n_num_draft_tokens, (req_i,))
513+
514+
for pos in range(num_draft_tokens):
515+
token_idx = start_idx + pos
516+
draft_token_id = tl.load(draft_token_ids_ptr + token_idx)
497517

498-
if ENABLE_REDUCE_SAMPLING:
499518
target_prob = 0.0
500519
found = False
501520

502-
for v_offset in range(0, vocab_size, VOCAB_BLOCK_SIZE):
521+
for v_offset in range(0, vocab_size, SUB_BLOCK):
503522
if not found:
504-
vocab_offsets = v_offset + tl.arange(0, VOCAB_BLOCK_SIZE)
523+
vocab_offsets = v_offset + tl.arange(0, SUB_BLOCK)
505524
vocab_mask = vocab_offsets < vocab_size
506525

507526
candidate_indices = tl.load(
@@ -519,37 +538,115 @@ def rejection_random_sample_block_verify_kernel(
519538
if current_match_prob > 0.0:
520539
target_prob = current_match_prob
521540
found = True
541+
542+
tmp_uniform_prob = tl.load(uniform_probs_ptr + token_idx)
543+
uniform_prob = uniform_prob * tmp_uniform_prob
544+
545+
if NO_DRAFT_PROBS:
546+
draft_prob = 1.0
547+
else:
548+
draft_prob = tl.load(draft_probs_ptr + token_idx * global_vocab_size + draft_token_id)
549+
550+
pi = min(pi * target_prob / draft_prob, 1.0)
551+
if draft_prob > 0 and pi >= uniform_prob:
552+
last_accepted_token_pos = pos
553+
554+
# Store accepted tokens
555+
if last_accepted_token_pos > -1:
556+
for pos in range(last_accepted_token_pos + 1):
557+
token_id = tl.load(draft_token_ids_ptr + start_idx + pos)
558+
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + pos, token_id)
559+
560+
# Store recovered or bonus token
561+
if last_accepted_token_pos + 1 < num_draft_tokens:
562+
# Rejected - store recovered token
563+
recovered_token_id = tl.load(recovered_token_ids_ptr + start_idx + last_accepted_token_pos + 1)
564+
tl.store(
565+
output_token_ids_ptr + req_idx * (max_spec_len + 1) + last_accepted_token_pos + 1,
566+
recovered_token_id,
567+
)
522568
else:
569+
# All accepted - store bonus token
570+
bonus_token_id = tl.load(bonus_token_ids_ptr + req_idx)
571+
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + num_draft_tokens, bonus_token_id)
572+
else:
573+
vocab_size = global_vocab_size
574+
loop = (vocab_size + SUB_BLOCK - 1) // SUB_BLOCK
575+
for req_i in range(BLOCK_SIZE):
576+
not_greedy = get_element(not_greedy_mask, (req_i,))
577+
if not_greedy:
578+
start_idx = get_element(start_idxs, (req_i,))
579+
req_idx = block_idx * BLOCK_SIZE + req_i
580+
num_draft_tokens = get_element(n_num_draft_tokens, (req_i,))
581+
if num_draft_tokens == 0:
582+
bonus_token_id = tl.load(bonus_token_ids_ptr + req_idx)
583+
tl.store(
584+
output_token_ids_ptr + req_idx * (max_spec_len + 1),
585+
bonus_token_id,
586+
)
587+
continue
588+
589+
accepted_len = 0
590+
prefix_prob = 1.0
591+
for pos in range(num_draft_tokens):
592+
token_idx = start_idx + pos
593+
draft_token_id = tl.load(draft_token_ids_ptr + token_idx)
523594
target_prob = tl.load(target_probs_ptr + token_idx * vocab_size + draft_token_id)
524595

525-
tmp_uniform_prob = tl.load(uniform_probs_ptr + token_idx)
526-
uniform_prob = uniform_prob * tmp_uniform_prob
596+
if NO_DRAFT_PROBS:
597+
draft_prob = 1.0
598+
else:
599+
draft_prob = tl.load(draft_probs_ptr + token_idx * vocab_size + draft_token_id)
527600

528-
if NO_DRAFT_PROBS:
529-
draft_prob = 1.0
530-
else:
531-
vocab_for_draft = global_vocab_size if ENABLE_REDUCE_SAMPLING else vocab_size
532-
draft_prob = tl.load(draft_probs_ptr + token_idx * vocab_for_draft + draft_token_id)
601+
if draft_prob > 0:
602+
prefix_prob = min(prefix_prob * target_prob / draft_prob, 1.0)
603+
else:
604+
prefix_prob = 0.0
533605

534-
pi = min(pi * target_prob / draft_prob, 1.0)
535-
if draft_prob > 0 and pi >= uniform_prob:
536-
last_accepted_token_pos = pos
606+
if pos == num_draft_tokens - 1:
607+
h_block = prefix_prob
608+
else:
609+
next_token_idx = token_idx + 1
610+
if NO_DRAFT_PROBS:
611+
next_draft_token_id = tl.load(draft_token_ids_ptr + next_token_idx)
612+
next_target_prob = tl.load(
613+
target_probs_ptr + next_token_idx * vocab_size + next_draft_token_id
614+
)
615+
residual_mass = prefix_prob * (1.0 - next_target_prob)
616+
else:
617+
residual_mass = 0.0
618+
for loop_i in range(loop):
619+
vocab_start = loop_i * SUB_BLOCK
620+
vocab_offset = vocab_start + tl.arange(0, SUB_BLOCK)
621+
next_draft_prob = tl.load(
622+
draft_probs_ptr + next_token_idx * vocab_size + vocab_offset,
623+
mask=vocab_offset < vocab_size,
624+
other=0,
625+
)
626+
next_target_prob = tl.load(
627+
target_probs_ptr + next_token_idx * vocab_size + vocab_offset,
628+
mask=vocab_offset < vocab_size,
629+
other=0,
630+
)
631+
residual_prob = tl.maximum(prefix_prob * next_target_prob - next_draft_prob, 0.0)
632+
residual_mass += tl.sum(residual_prob, axis=0)
633+
denom = residual_mass + 1.0 - prefix_prob
634+
h_block = residual_mass / denom if denom > 0 else 0.0
537635

538-
# Store accepted tokens
539-
if last_accepted_token_pos > -1:
540-
for pos in range(last_accepted_token_pos + 1):
636+
uniform_prob = tl.load(uniform_probs_ptr + token_idx)
637+
if uniform_prob <= h_block:
638+
accepted_len = pos + 1
639+
640+
for pos in range(accepted_len):
541641
token_id = tl.load(draft_token_ids_ptr + start_idx + pos)
542642
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + pos, token_id)
543643

544-
# Store recovered or bonus token
545-
if last_accepted_token_pos + 1 < num_draft_tokens:
546-
# Rejected - store recovered token
547-
recovered_token_id = tl.load(recovered_token_ids_ptr + start_idx + last_accepted_token_pos + 1)
548-
tl.store(
549-
output_token_ids_ptr + req_idx * (max_spec_len + 1) + last_accepted_token_pos + 1,
550-
recovered_token_id,
551-
)
552-
else:
553-
# All accepted - store bonus token
554-
bonus_token_id = tl.load(bonus_token_ids_ptr + req_idx)
555-
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + num_draft_tokens, bonus_token_id)
644+
if accepted_len == num_draft_tokens:
645+
bonus_token_id = tl.load(bonus_token_ids_ptr + req_idx)
646+
tl.store(output_token_ids_ptr + req_idx * (max_spec_len + 1) + num_draft_tokens, bonus_token_id)
647+
else:
648+
recovered_token_id = tl.load(recovered_token_ids_ptr + start_idx + accepted_len)
649+
tl.store(
650+
output_token_ids_ptr + req_idx * (max_spec_len + 1) + accepted_len,
651+
recovered_token_id,
652+
)

‎vllm_ascend/patch/worker/__init__.py‎

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -45,9 +45,6 @@
4545
else:
4646
import vllm_ascend.patch.worker.patch_idex_310 # noqa
4747
import vllm_ascend.patch.worker.patch_rejection_sampler # noqa
48-
49-
# TODO(MengqingCao): remove after the upstream community is modified
50-
import vllm_ascend.patch.worker.patch_llama_eagle3 # noqa
5148
import vllm_ascend.patch.worker.patch_npugraph_ex_triton # noqa
5249
import vllm_ascend.patch.worker.patch_kimi_k25 # noqa
5350
import vllm_ascend.patch.worker.patch_draft_quarot # noqa

0 commit comments

Comments
 (0)