From d12f292405270cf9c282911a41c40159379877c7 Mon Sep 17 00:00:00 2001 From: goodruyas <211121886+goodruyas@users.noreply.github.com> Date: Wed, 9 Sep 2026 16:27:36 +0800 Subject: [PATCH] common : reject non-positive batch size (#28525) --- common/arg.cpp | 3 +++ src/llama-context.cpp | 4 ++-- tests/test-arg-parser.cpp | 7 +++++++ 3 files changed, 12 insertions(+), 2 deletions(-) diff --git a/common/arg.cpp b/common/arg.cpp index 74241f931285..d555663c47ef 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -1677,6 +1677,9 @@ common_params_context common_params_parser_init(common_params & params, llama_ex {"-b", "--batch-size"}, "N", string_format("logical maximum batch size (default: %d)", params.n_batch), [](common_params & params, int value) { + if (value <= 0) { + throw std::invalid_argument("value must be positive"); + } params.n_batch = value; } ).set_env("LLAMA_ARG_BATCH")); diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 21501574a911..ee3e4bf0e135 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3665,8 +3665,8 @@ llama_context * llama_init_from_model( return nullptr; } - if (params.n_batch == 0 && params.n_ubatch == 0) { - LLAMA_LOG_ERROR("%s: n_batch and n_ubatch cannot both be zero\n", __func__); + if (params.n_batch == 0) { + LLAMA_LOG_ERROR("%s: n_batch cannot be zero\n", __func__); return nullptr; } diff --git a/tests/test-arg-parser.cpp b/tests/test-arg-parser.cpp index e0907631abd8..10f63cd7ebbd 100644 --- a/tests/test-arg-parser.cpp +++ b/tests/test-arg-parser.cpp @@ -182,6 +182,13 @@ static void test(void) { argv = {"binary_name", "-ngl", "hello"}; assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON)); + // batch-size must be positive + argv = {"binary_name", "--batch-size", "0"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON)); + + argv = {"binary_name", "--batch-size", "-1"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON)); + // wrong value (enum) argv = {"binary_name", "-sm", "hello"}; assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));