@@ -69,15 +69,22 @@ SampleOutput RejectionSampler::forward(const torch::Tensor& draft_token_ids,
6969 bool mask_out_rejected_tokens) const {
7070 CHECK_EQ (draft_token_ids.size (0 ), do_sample_.size (0 ))
7171 << " batch size mismatch" ;
72- DCHECK_EQ (draft_token_ids.size (1 ), draft_probs.size (1 ));
72+ if (!all_greedy_sample_) {
73+ CHECK (draft_probs.defined ())
74+ << " draft_probs must be defined for random sampling" ;
75+ CHECK_EQ (draft_token_ids.size (1 ), draft_probs.size (1 ));
76+ }
7377 // DCHECK_EQ(draft_probs.sizes(), target_probs.sizes());
7478
75- // [batch_size, n_speculative_tokens + 1, vocab_size] FloatTensor
76- auto target_probs =
77- torch::softmax (target_logits, /* dim=*/ -1 , /* dtype=*/ torch::kFloat32 );
78- // filter out probs for bonus tokens
79- target_probs = target_probs.slice (
80- /* dim=*/ 1 , /* start=*/ 0 , /* end=*/ target_probs.size (1 ) - 1 );
79+ torch::Tensor target_probs;
80+ if (!all_greedy_sample_) {
81+ // The bonus row is sampled by the target worker and does not participate
82+ // in random rejection sampling.
83+ torch::Tensor target_draft_logits = target_logits.slice (
84+ /* dim=*/ 1 , /* start=*/ 0 , /* end=*/ target_logits.size (1 ) - 1 );
85+ target_probs = torch::softmax (
86+ target_draft_logits, /* dim=*/ -1 , /* dtype=*/ torch::kFloat32 );
87+ }
8188
8289 // Determine whether we need to restore rejected tokens.
8390 // IMPORTANT: The fused kernel implementation only supports masking out
@@ -97,9 +104,11 @@ SampleOutput RejectionSampler::forward(const torch::Tensor& draft_token_ids,
97104 torch::Tensor accepted_token_ids;
98105 torch::Tensor masked_accepted_token_ids;
99106 if (all_greedy_sample_) {
107+ torch::Tensor target_draft_logits = target_logits.slice (
108+ /* dim=*/ 1 , /* start=*/ 0 , /* end=*/ target_logits.size (1 ) - 1 );
100109 std::tie (accepted_token_ids, masked_accepted_token_ids) =
101110 greedy_sample (draft_token_ids,
102- target_probs ,
111+ target_draft_logits ,
103112 bonus_token_ids,
104113 mask_out_rejected_tokens);
105114 } else if (all_random_sample_) {
@@ -331,14 +340,27 @@ std::tuple<torch::Tensor, torch::Tensor> RejectionSampler::random_sample_fused(
331340
332341std::tuple<torch::Tensor, torch::Tensor> RejectionSampler::greedy_sample (
333342 const torch::Tensor& draft_token_ids,
334- const torch::Tensor& target_probs ,
343+ const torch::Tensor& target_scores ,
335344 const torch::Tensor& bonus_token_ids,
336345 bool mask_out_rejected_tokens) {
337- auto target_token_ids = Sampler::greedy_sample (target_probs);
346+ torch::Tensor target_token_ids = Sampler::greedy_sample (target_scores);
347+ return greedy_sample_from_token_ids (draft_token_ids,
348+ target_token_ids,
349+ bonus_token_ids,
350+ mask_out_rejected_tokens);
351+ }
338352
353+ std::tuple<torch::Tensor, torch::Tensor>
354+ RejectionSampler::greedy_sample_from_token_ids (
355+ const torch::Tensor& draft_token_ids,
356+ const torch::Tensor& target_token_ids,
357+ const torch::Tensor& bonus_token_ids,
358+ bool mask_out_rejected_tokens) {
359+ CHECK_EQ (target_token_ids.sizes (), draft_token_ids.sizes ())
360+ << " target and draft token shapes must match" ;
339361 // mask out the rejected tokens with -1
340362 // [batch_size, n_speculative_tokens + 1]
341- auto accepted_token_ids =
363+ torch::Tensor accepted_token_ids =
342364 torch::cat ({target_token_ids, bonus_token_ids}, /* dim=*/ -1 );
343365 torch::Tensor masked_accepted_token_ids;
344366 if (mask_out_rejected_tokens) {
0 commit comments