From 3e1394f82bbfe59decc4fb83df0c075360f85d8f Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Thu, 11 Jun 2026 16:04:05 -0700 Subject: [PATCH 1/3] Add upper bounds --- src/generators.cpp | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/generators.cpp b/src/generators.cpp index 3657810442..8ecc6fb8fe 100644 --- a/src/generators.cpp +++ b/src/generators.cpp @@ -380,8 +380,10 @@ Generator::Generator(const Model& model, const GeneratorParams& params) : model_ throw std::runtime_error("search max_length is 0"); if (params.search.max_length > model.config_->model.context_length) throw std::runtime_error("max_length (" + std::to_string(params.search.max_length) + ") cannot be greater than model context_length (" + std::to_string(model.config_->model.context_length) + ")"); - if (params.search.batch_size < 1) - throw std::runtime_error("batch_size must be 1 or greater, is " + std::to_string(params.search.batch_size)); + if (params.search.batch_size < 1 || params.search.batch_size > 256) + throw std::runtime_error("batch_size (" + std::to_string(params.search.batch_size) + ") must be in [1, 256]"); + if (params.search.num_beams < 1 || params.search.num_beams > 256) + throw std::runtime_error("num_beams (" + std::to_string(params.search.num_beams) + ") must be in [1, 256]"); if (params.config.model.vocab_size < 1) throw std::runtime_error("vocab_size must be 1 or greater, is " + std::to_string(params.config.model.vocab_size)); From 2828894a8f06a7d224896282bf38f68ae72ada61 Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Thu, 11 Jun 2026 16:20:14 -0700 Subject: [PATCH 2/3] Add unit tests --- src/generators.cpp | 18 ++++++++++++++---- test/model_tests.cpp | 41 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 55 insertions(+), 4 deletions(-) diff --git a/src/generators.cpp b/src/generators.cpp index 8ecc6fb8fe..38416fbad3 100644 --- a/src/generators.cpp +++ b/src/generators.cpp @@ -380,10 +380,20 @@ Generator::Generator(const Model& model, const GeneratorParams& params) : model_ throw std::runtime_error("search max_length is 0"); if (params.search.max_length > model.config_->model.context_length) throw std::runtime_error("max_length (" + std::to_string(params.search.max_length) + ") cannot be greater than model context_length (" + std::to_string(model.config_->model.context_length) + ")"); - if (params.search.batch_size < 1 || params.search.batch_size > 256) - throw std::runtime_error("batch_size (" + std::to_string(params.search.batch_size) + ") must be in [1, 256]"); - if (params.search.num_beams < 1 || params.search.num_beams > 256) - throw std::runtime_error("num_beams (" + std::to_string(params.search.num_beams) + ") must be in [1, 256]"); + + constexpr int kMaxBatchSize = 256; + constexpr int kMaxNumBeams = 256; + constexpr int kMaxNumBeamsCuda = 32; + + if (params.search.batch_size < 1 || params.search.batch_size > kMaxBatchSize) + throw std::runtime_error("batch_size (" + std::to_string(params.search.batch_size) + ") must be in [1, " + std::to_string(kMaxBatchSize) + "]"); + + const int max_num_beams = (params.search.num_beams > 1 && + (params.p_device->GetType() == DeviceType::CUDA || params.p_device->GetType() == DeviceType::NvTensorRtRtx)) + ? kMaxNumBeamsCuda + : kMaxNumBeams; + if (params.search.num_beams < 1 || params.search.num_beams > max_num_beams) + throw std::runtime_error("num_beams (" + std::to_string(params.search.num_beams) + ") must be in [1, " + std::to_string(max_num_beams) + "]"); if (params.config.model.vocab_size < 1) throw std::runtime_error("vocab_size must be 1 or greater, is " + std::to_string(params.config.model.vocab_size)); diff --git a/test/model_tests.cpp b/test/model_tests.cpp index 7910f41122..c161b5566b 100644 --- a/test/model_tests.cpp +++ b/test/model_tests.cpp @@ -477,4 +477,45 @@ Print all primes between 1 and n std::cout << tokenizer->Decode(result) << "\r\n"; } +#endif + +// Validation tests for search parameter bounds +#if !USE_DML +TEST(ModelTests, NumBeamsUpperBoundThrows) { + auto model = OgaModel::Create(MODEL_PATH "hf-internal-testing/tiny-random-gpt2-fp32"); + auto params = OgaGeneratorParams::Create(*model); + params->SetSearchOption("max_length", 20); + params->SetSearchOption("batch_size", 1); + params->SetSearchOption("num_beams", 512); // exceeds upper bound of 256 + + EXPECT_THROW(OgaGenerator::Create(*model, *params), std::runtime_error); +} + +TEST(ModelTests, BatchSizeUpperBoundThrows) { + auto model = OgaModel::Create(MODEL_PATH "hf-internal-testing/tiny-random-gpt2-fp32"); + auto params = OgaGeneratorParams::Create(*model); + params->SetSearchOption("max_length", 20); + params->SetSearchOption("batch_size", 512); // exceeds upper bound of 256 + + EXPECT_THROW(OgaGenerator::Create(*model, *params), std::runtime_error); +} + +TEST(ModelTests, NumBeamsZeroThrows) { + auto model = OgaModel::Create(MODEL_PATH "hf-internal-testing/tiny-random-gpt2-fp32"); + auto params = OgaGeneratorParams::Create(*model); + params->SetSearchOption("max_length", 20); + params->SetSearchOption("batch_size", 1); + params->SetSearchOption("num_beams", 0); // below lower bound of 1 + + EXPECT_THROW(OgaGenerator::Create(*model, *params), std::runtime_error); +} + +TEST(ModelTests, BatchSizeZeroThrows) { + auto model = OgaModel::Create(MODEL_PATH "hf-internal-testing/tiny-random-gpt2-fp32"); + auto params = OgaGeneratorParams::Create(*model); + params->SetSearchOption("max_length", 20); + params->SetSearchOption("batch_size", 0); // below lower bound of 1 + + EXPECT_THROW(OgaGenerator::Create(*model, *params), std::runtime_error); +} #endif \ No newline at end of file From bcfb861fccd12b5177942e2f840f7d78c20d32e7 Mon Sep 17 00:00:00 2001 From: Akshay Sonawane Date: Fri, 12 Jun 2026 16:43:14 -0700 Subject: [PATCH 3/3] Update batch_size to 32 --- src/generators.cpp | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/generators.cpp b/src/generators.cpp index 38416fbad3..22fa98ca80 100644 --- a/src/generators.cpp +++ b/src/generators.cpp @@ -381,8 +381,8 @@ Generator::Generator(const Model& model, const GeneratorParams& params) : model_ if (params.search.max_length > model.config_->model.context_length) throw std::runtime_error("max_length (" + std::to_string(params.search.max_length) + ") cannot be greater than model context_length (" + std::to_string(model.config_->model.context_length) + ")"); - constexpr int kMaxBatchSize = 256; - constexpr int kMaxNumBeams = 256; + constexpr int kMaxBatchSize = 32; + constexpr int kMaxNumBeams = 32; constexpr int kMaxNumBeamsCuda = 32; if (params.search.batch_size < 1 || params.search.batch_size > kMaxBatchSize)