Skip to content

Commit 65e6750

Browse files
committed
perf: remove unused greedy mtp probability work.
- reuse target token IDs for greedy MTP validation - skip draft probabilities only for all-greedy sampling - preserve random, mixed, and logprob sampling paths
1 parent d638401 commit 65e6750

7 files changed

Lines changed: 156 additions & 26 deletions

File tree

tests/core/framework/sampling/rejection_sampler_test.cpp

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -202,6 +202,53 @@ TEST(RejectionSamplerTest, Greedy) {
202202
bonus_token_ids));
203203
}
204204

205+
TEST(RejectionSamplerTest, GreedyFromTokenIds) {
206+
const auto device = get_test_device();
207+
const auto draft_token_ids = torch::tensor({{1, 2, 3}, {3, 1, 0}}, device);
208+
const auto target_token_ids = torch::tensor({{1, 4, 3}, {3, 2, 0}}, device);
209+
const auto bonus_token_ids = torch::tensor({{5}, {6}}, device);
210+
211+
auto [output, masked_output] = RejectionSampler::greedy_sample_from_token_ids(
212+
draft_token_ids,
213+
target_token_ids,
214+
bonus_token_ids,
215+
/*mask_out_rejected_tokens=*/true);
216+
217+
EXPECT_TRUE(torch::equal(
218+
output, torch::tensor({{1, 4, 3, 5}, {3, 2, 0, 6}}, device)));
219+
EXPECT_TRUE(torch::equal(
220+
masked_output, torch::tensor({{1, 4, -1, -1}, {3, 2, -1, -1}}, device)));
221+
}
222+
223+
TEST(RejectionSamplerTest, GreedyForwardAllowsUndefinedDraftProbs) {
224+
const auto options = get_test_options(torch::kFloat32);
225+
const auto device = get_test_device();
226+
const auto do_sample = torch::tensor({false, false}, device);
227+
RejectionSampler sampler(do_sample,
228+
/*all_random_sample=*/false,
229+
/*all_greedy_sample=*/true,
230+
/*logprobs=*/false,
231+
/*max_top_logprobs=*/0);
232+
const auto draft_token_ids = torch::tensor({{1, 2}, {3, 4}}, device);
233+
const auto target_logits = torch::tensor({{{0.0f, 4.0f, 1.0f, 2.0f, 3.0f},
234+
{0.0f, 1.0f, 2.0f, 4.0f, 3.0f},
235+
{4.0f, 0.0f, 1.0f, 2.0f, 3.0f}},
236+
{{0.0f, 1.0f, 2.0f, 4.0f, 3.0f},
237+
{0.0f, 1.0f, 2.0f, 3.0f, 4.0f},
238+
{0.0f, 4.0f, 1.0f, 2.0f, 3.0f}}},
239+
options);
240+
const auto bonus_token_ids = torch::tensor({{0}, {1}}, device);
241+
242+
SampleOutput output = sampler.forward(draft_token_ids,
243+
torch::Tensor(),
244+
target_logits,
245+
bonus_token_ids,
246+
/*mask_out_rejected_tokens=*/true);
247+
248+
EXPECT_TRUE(torch::equal(output.next_tokens,
249+
torch::tensor({{1, 3, -1}, {3, 4, 1}}, device)));
250+
}
251+
205252
TEST(RejectionSamplerTest, LogProbs) {
206253
const auto options = get_test_options(torch::kFloat32);
207254
const auto device = get_test_device();

tests/core/runtime/spec_input_builder_test.cpp

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -475,6 +475,25 @@ TEST(DraftProbsBuilderTest, BuildValidateTensorsSelectedOnly) {
475475
torch::tensor({{0.3f, 0.5f}, {0.4f, 0.6f}}, torch::kFloat32)));
476476
}
477477

478+
TEST(DraftProbsBuilderTest, BuildValidateTensorsSkipsGreedyProbs) {
479+
std::vector<torch::Tensor> token_steps = {
480+
torch::tensor({3, 4}, torch::kInt64),
481+
torch::tensor({5, 6}, torch::kInt64)};
482+
std::vector<torch::Tensor> probs_steps(token_steps.size());
483+
484+
auto [draft_token_ids, draft_probs] =
485+
draftProbs::build_validate_tensors(token_steps,
486+
probs_steps,
487+
/*batch_size=*/2,
488+
/*vocab_size=*/8,
489+
/*enable_opt_validate_probs=*/true,
490+
/*draft_probs_required=*/false);
491+
492+
EXPECT_TRUE(torch::equal(draft_token_ids,
493+
torch::tensor({{3, 5}, {4, 6}}, torch::kInt64)));
494+
EXPECT_FALSE(draft_probs.defined());
495+
}
496+
478497
TEST(DraftProbsBuilderTest, BuildValidateTensorsRecoveredDense) {
479498
std::vector<torch::Tensor> token_steps = {
480499
torch::tensor({1, 2}, torch::kInt64),

xllm/core/framework/sampling/rejection_sampler.cpp

Lines changed: 33 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -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

332341
std::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) {

xllm/core/framework/sampling/rejection_sampler.h

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ class RejectionSampler final {
4242
// draft_probs:
4343
// 1) dense format: [batch_size, n_speculative_tokens, vocab_size]
4444
// 2) selected-only format: [batch_size, n_speculative_tokens]
45+
// 3) undefined for all-greedy sampling
4546
// target_logits: [batch_size, n_speculative_tokens + 1, vocab_size]
4647
// bonus_token_ids: [batch_size, 1]
4748
SampleOutput forward(const torch::Tensor& draft_token_ids,
@@ -73,7 +74,13 @@ class RejectionSampler final {
7374

7475
static std::tuple<torch::Tensor, torch::Tensor> greedy_sample(
7576
const torch::Tensor& draft_token_ids,
76-
const torch::Tensor& target_probs,
77+
const torch::Tensor& target_scores,
78+
const torch::Tensor& bonus_token_ids,
79+
bool mask_out_rejected_tokens);
80+
81+
static std::tuple<torch::Tensor, torch::Tensor> greedy_sample_from_token_ids(
82+
const torch::Tensor& draft_token_ids,
83+
const torch::Tensor& target_token_ids,
7784
const torch::Tensor& bonus_token_ids,
7885
bool mask_out_rejected_tokens);
7986

xllm/core/runtime/mtp_worker_impl.cpp

Lines changed: 35 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1549,7 +1549,8 @@ void MTPWorkerImpl::prepare_draft_extend_inputs(
15491549
ForwardInput& extend_input) {
15501550
c10::StreamGuard stream_guard = prepare_stream_->set_stream_guard();
15511551
extend_input = base_input;
1552-
extend_input.sampling_params.return_probs = true;
1552+
extend_input.sampling_params.return_probs =
1553+
!extend_input.sampling_params.all_greedy_sample;
15531554
clear_ready_events(extend_input);
15541555
extend_input.device_tensors_ready = false;
15551556
auto& input_params = extend_input.input_params;
@@ -1764,7 +1765,8 @@ void MTPWorkerImpl::prepare_draft_inputs(const ForwardInput& input,
17641765
int32_t position_offset) {
17651766
c10::StreamGuard stream_guard = prepare_stream_->set_stream_guard();
17661767
draft_input = input;
1767-
draft_input.sampling_params.return_probs = true;
1768+
draft_input.sampling_params.return_probs =
1769+
!draft_input.sampling_params.all_greedy_sample;
17681770
clear_ready_events(draft_input);
17691771
draft_input.device_tensors_ready = false;
17701772

@@ -1844,7 +1846,8 @@ SampleOutput MTPWorkerImpl::validate(
18441846
draft_probs_steps,
18451847
batch_size,
18461848
vocab_size,
1847-
enable_opt_validate_probs_);
1849+
enable_opt_validate_probs_,
1850+
/*draft_probs_required=*/!sampling_params.all_greedy_sample);
18481851
return validate(sampling_params, draft_token_ids, draft_probs, target_output);
18491852
}
18501853

@@ -1866,6 +1869,28 @@ SampleOutput MTPWorkerImpl::validate(const SamplingParameters& sampling_params,
18661869
.index({"...", ISlice(num_val_tokens - 1, None, num_val_tokens)})
18671870
.view({-1, 1});
18681871

1872+
if (sampling_params.all_greedy_sample && !target_output.logprobs) {
1873+
torch::Tensor target_token_ids =
1874+
target_output.sample_output.next_tokens.view(
1875+
{batch_size, num_val_tokens});
1876+
torch::Tensor target_draft_token_ids = target_token_ids.slice(
1877+
/*dim=*/1, /*start=*/0, /*end=*/num_val_tokens - 1);
1878+
auto [accepted_token_ids, masked_accepted_token_ids] =
1879+
RejectionSampler::greedy_sample_from_token_ids(
1880+
draft_token_ids.to(target_draft_token_ids),
1881+
target_draft_token_ids,
1882+
bonus_token_ids,
1883+
/*mask_out_rejected_tokens=*/true);
1884+
(void)accepted_token_ids;
1885+
1886+
SampleOutput sample_output;
1887+
sample_output.next_tokens = masked_accepted_token_ids;
1888+
torch::Tensor embeddings = target_output.sample_output.embeddings;
1889+
sample_output.embeddings =
1890+
embeddings.view({batch_size, num_val_tokens, embeddings.size(-1)});
1891+
return sample_output;
1892+
}
1893+
18691894
auto target_logits =
18701895
target_output.logits.view({batch_size, num_val_tokens, vocab_size});
18711896

@@ -1879,12 +1904,13 @@ SampleOutput MTPWorkerImpl::validate(const SamplingParameters& sampling_params,
18791904
enable_fused_kernel_);
18801905

18811906
// get the accepted tokens
1882-
SampleOutput sample_output =
1883-
rejection_sampler->forward(draft_token_ids.to(bonus_token_ids),
1884-
draft_probs.to(target_logits.device()),
1885-
target_logits,
1886-
bonus_token_ids,
1887-
/*mask_out_rejected_tokens=*/true);
1907+
SampleOutput sample_output = rejection_sampler->forward(
1908+
draft_token_ids.to(bonus_token_ids),
1909+
draft_probs.defined() ? draft_probs.to(target_logits.device())
1910+
: torch::Tensor(),
1911+
target_logits,
1912+
bonus_token_ids,
1913+
/*mask_out_rejected_tokens=*/true);
18881914

18891915
// process embedding
18901916
auto embeddings = target_output.sample_output.embeddings;

xllm/core/runtime/spec_input_builder.cpp

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -468,7 +468,8 @@ std::pair<torch::Tensor, torch::Tensor> build_validate_tensors(
468468
const std::vector<torch::Tensor>& draft_probs_steps,
469469
int32_t batch_size,
470470
int32_t vocab_size,
471-
bool enable_opt_validate_probs) {
471+
bool enable_opt_validate_probs,
472+
bool draft_probs_required) {
472473
CHECK_GT(batch_size, 0) << "batch_size must be > 0";
473474
CHECK_GT(vocab_size, 0) << "vocab_size must be > 0";
474475
CHECK_EQ(draft_token_ids_steps.size(), draft_probs_steps.size())
@@ -483,11 +484,14 @@ std::pair<torch::Tensor, torch::Tensor> build_validate_tensors(
483484
for (size_t i = 0; i < draft_token_ids_steps.size(); ++i) {
484485
auto draft_token_ids =
485486
draft_token_ids_steps[i].view({batch_size, 1}).to(torch::kLong);
486-
auto selected_probs =
487+
token_ids_vec.emplace_back(draft_token_ids);
488+
if (!draft_probs_required) {
489+
continue;
490+
}
491+
492+
torch::Tensor selected_probs =
487493
extract_selected_probs(draft_probs_steps[i], draft_token_ids)
488494
.view({batch_size, 1});
489-
490-
token_ids_vec.emplace_back(draft_token_ids);
491495
if (enable_opt_validate_probs) {
492496
probs_vec.emplace_back(selected_probs);
493497
} else {
@@ -502,6 +506,9 @@ std::pair<torch::Tensor, torch::Tensor> build_validate_tensors(
502506
}
503507

504508
auto draft_token_ids = torch::cat(token_ids_vec, /*dim=*/1);
509+
if (!draft_probs_required) {
510+
return {draft_token_ids, torch::Tensor()};
511+
}
505512
auto draft_probs = torch::cat(probs_vec, /*dim=*/1);
506513
return {draft_token_ids, draft_probs};
507514
}

xllm/core/runtime/spec_input_builder.h

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -170,12 +170,14 @@ torch::Tensor compress_for_cache(const torch::Tensor& draft_probs,
170170
// enable_opt_validate_probs=true
171171
// * recovered-dense [batch_size, n_speculative_tokens, vocab_size], if
172172
// enable_opt_validate_probs=false
173+
// * undefined, if draft_probs_required=false
173174
std::pair<torch::Tensor, torch::Tensor> build_validate_tensors(
174175
const std::vector<torch::Tensor>& draft_token_ids_steps,
175176
const std::vector<torch::Tensor>& draft_probs_steps,
176177
int32_t batch_size,
177178
int32_t vocab_size,
178-
bool enable_opt_validate_probs);
179+
bool enable_opt_validate_probs,
180+
bool draft_probs_required = true);
179181

180182
} // namespace draftProbs
181183

0 commit comments

Comments
 (0)