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));