@@ -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+ )
0 commit comments