diff --git a/src/generators.cpp b/src/generators.cpp index 3657810442..388e63501e 100644 --- a/src/generators.cpp +++ b/src/generators.cpp @@ -420,6 +420,8 @@ void Generator::InitializeSamplingMethod(const GeneratorParams& params) { throw std::runtime_error("top_p must be between 0.0 and 1.0"); if (search.top_k < 0) throw std::runtime_error("top_k must be 0 or greater"); + if (search.top_k > params.config.model.vocab_size) + throw std::runtime_error("top_k (" + std::to_string(search.top_k) + ") must be less than or equal to vocab_size (" + std::to_string(params.config.model.vocab_size) + ")"); if (search.top_p > 0.0f && search.top_p < 1.0f && search.top_k > 1) { sampling_method_ = SamplingMethod::kTopKTopP; } else if (search.top_k > 1) { diff --git a/src/search.cpp b/src/search.cpp index bbc27644c2..b8ca196e60 100644 --- a/src/search.cpp +++ b/src/search.cpp @@ -170,6 +170,7 @@ void GreedySearch_Cpu::SelectTop() { void GreedySearch_Cpu::SampleTopK(int k, float temperature) { const int vocab_size = params_->config.model.vocab_size; + k = std::min(k, vocab_size); std::vector indices(vocab_size); std::vector top_k_scores(k); @@ -329,6 +330,9 @@ void GreedySearch_Cpu::SampleTopP(float p, float temperature) { void GreedySearch_Cpu::SampleTopKTopP(int k, float p, float temperature) { assert(temperature > 0.0f); + // Clamp k to vocab_size to prevent out-of-bounds access in partial_sort + k = std::min(k, params_->config.model.vocab_size); + // --- Buffers allocated once to avoid re-allocations in the batch loop --- std::vector indices(params_->config.model.vocab_size); diff --git a/test/sampling_tests.cpp b/test/sampling_tests.cpp index 60bac47f40..afca86345a 100644 --- a/test/sampling_tests.cpp +++ b/test/sampling_tests.cpp @@ -460,6 +460,19 @@ TEST(SamplingTests, RepetitionPenaltyCorrectnessCpu) { << " Unpenalized per-token avg=" << unpenalized_avg; } +TEST(SamplingTests, TopKExceedingVocabSizeIsRejected) { + auto config = OgaConfig::Create(MODEL_PATH "hf-internal-testing/tiny-random-gpt2-fp32"); + config->ClearProviders(); + auto model = OgaModel::Create(*config); + + auto params = OgaGeneratorParams::Create(*model); + params->SetSearchOption("max_length", 25); + params->SetSearchOptionBool("do_sample", true); + params->SetSearchOption("top_k", 5000); // vocab_size is 1000 + + EXPECT_THROW(OgaGenerator::Create(*model, *params), std::runtime_error); +} + #if USE_CUDA TEST(SamplingTests, BatchedSamplingTopPCuda) { std::vector input_ids{0, 1, 2, 3}; @@ -480,6 +493,7 @@ TEST(SamplingTests, BatchedSamplingTopPCuda) { auto params = OgaGeneratorParams::Create(*model); params->SetSearchOption("max_length", 10); params->SetSearchOptionBool("do_sample", true); + params->SetSearchOption("top_k", 1); params->SetSearchOption("top_p", 0.25f); params->SetSearchOption("batch_size", batch_size);