From f3f370ea8d91ed7f6995218fd8151b11b5240900 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Sun, 19 Jul 2026 11:58:13 +0100 Subject: [PATCH 1/3] add configurations Signed-off-by: kerthcet --- README.md | 68 +++++++++++++++---------- docs/configuration.md | 36 ++++++++++++++ src/api/chat.rs | 4 +- src/api/completions.rs | 2 +- src/api/tests.rs | 3 +- src/backend/llm_engine.rs | 87 ++++++++++++++++++++++++++++---- src/backend/mod.rs | 2 +- src/cli/commands.rs | 101 ++++++++++++++++++++++++++++++++++++-- src/cli/serve.rs | 4 +- 9 files changed, 262 insertions(+), 45 deletions(-) create mode 100644 docs/configuration.md diff --git a/README.md b/README.md index b2558c2..430dbd3 100644 --- a/README.md +++ b/README.md @@ -51,26 +51,29 @@ make build ```bash # Download a model -puma pull inftyai/tiny-random-gpt2 +puma pull qwen/qwen2.5-0.5b + +# Run a model in an interactive chat +puma run qwen/qwen2.5-0.5b # List all models puma ls # Inspect model details -puma inspect inftyai/tiny-random-gpt2 +puma inspect qwen/qwen2.5-0.5b # Check system info puma info # Remove a model -puma rm inftyai/tiny-random-gpt2 +puma rm qwen/qwen2.5-0.5b ``` ### API Server ```bash # Start the inference server with a model -puma serve inftyai/tiny-random-gpt2 +puma serve qwen/qwen2.5-0.5b # Server will start on http://0.0.0.0:8000 # API endpoints: @@ -91,7 +94,7 @@ curl http://localhost:8000/health curl http://localhost:8000/v1/chat/completions \ -H "Content-Type: application/json" \ -d '{ - "model": "inftyai/tiny-random-gpt2", + "model": "qwen/qwen2.5-0.5b", "messages": [{"role": "user", "content": "Hello!"}] }' @@ -109,9 +112,9 @@ curl http://localhost:8000/v1/chat/completions \ | `rm ` | ✅ | Remove model and cache | | `info` | ✅ | Display system information | | `version` | ✅ | Show PUMA version | +| `run ` | ✅ | Run a model in an interactive chat | | `serve ` | ✅ | Start OpenAI-compatible API server with a model | | `ps` | 🚧 | List running models | -| `run` | 🚧 | Start model inference | | `stop` | 🚧 | Stop running model | ## Advanced Usage @@ -144,6 +147,17 @@ puma ls llama -l author=meta **Available filters:** `author`, `task`, `license`, `provider`, `model_series` +### Engine Tuning + +Both `run` and `serve` accept flags to tune the inference engine (KV-cache pool, +block size, batch size, default token budget): + +```bash +puma run qwen/qwen2.5-0.5b --max-batch-size 64 --default-max-tokens 256 +``` + +See [docs/configuration.md](docs/configuration.md) for the full list of flags and defaults. + ## API Server PUMA provides an OpenAI-compatible API server for model inference. @@ -152,15 +166,17 @@ PUMA provides an OpenAI-compatible API server for model inference. ```bash # Start server with a model (default: 0.0.0.0:8000) -puma serve inftyai/tiny-random-gpt2 +puma serve qwen/qwen2.5-0.5b # Custom host and port -puma serve inftyai/tiny-random-gpt2 --host 127.0.0.1 --port 3000 +puma serve qwen/qwen2.5-0.5b --host 127.0.0.1 --port 3000 # Model must be pulled first -puma pull inftyai/tiny-random-gpt2 +puma pull qwen/qwen2.5-0.5b ``` +Engine parameters (KV-cache, batch size, etc.) can be tuned with additional flags — see [docs/configuration.md](docs/configuration.md). + ### API Endpoints #### Chat Completions (Recommended) @@ -168,7 +184,7 @@ puma pull inftyai/tiny-random-gpt2 curl http://localhost:8000/v1/chat/completions \ -H "Content-Type: application/json" \ -d '{ - "model": "inftyai/tiny-random-gpt2", + "model": "qwen/qwen2.5-0.5b", "messages": [ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Hello!"} @@ -183,7 +199,7 @@ curl http://localhost:8000/v1/chat/completions \ curl http://localhost:8000/v1/chat/completions \ -H "Content-Type: application/json" \ -d '{ - "model": "inftyai/tiny-random-gpt2", + "model": "qwen/qwen2.5-0.5b", "messages": [{"role": "user", "content": "Tell me a story"}], "stream": true }' @@ -214,7 +230,7 @@ client = OpenAI( ) response = client.chat.completions.create( - model="inftyai/tiny-random-gpt2", + model="qwen/qwen2.5-0.5b", messages=[ {"role": "user", "content": "Hello!"} ] @@ -226,28 +242,28 @@ print(response.choices[0].message.content) ### Inspect Output ```bash -$ puma inspect inftyai/tiny-random-gpt2 +$ puma inspect qwen/qwen2.5-0.5b -name: inftyai/tiny-random-gpt2 +name: qwen/qwen2.5-0.5b kind: Model spec: - author: inftyai - model_series: gpt2 + author: qwen + model_series: qwen2 task: text-generation - license: MIT - context_window: 2.05K + license: APACHE-2.0 + context_window: 32.77K safetensors: - total: 7.00B + total: 494.03M parameters: - f32: 7.00B - provider: huggingface + bf16: 494.03M + provider: huggingface cache: - revision: abc123de - size: 1.24 GB - path: ~/.puma/cache/... + revision: 060db6499f32faf8b98477b0a26969ef7d8b9987 + size: 988.10 MB + path: ~/.puma/cache/huggingface/models--qwen--qwen2.5-0.5b status: - created: 2 hours ago - updated: 2 hours ago + created: 2 months ago + updated: 2 months ago ``` ## Model Management diff --git a/docs/configuration.md b/docs/configuration.md new file mode 100644 index 0000000..0f0fba8 --- /dev/null +++ b/docs/configuration.md @@ -0,0 +1,36 @@ +# Configuration + +PUMA's inference engine exposes a set of tunable parameters. Both `puma run` +and `puma serve` accept the same engine flags — all are optional and fall back +to the defaults below. + +## Engine flags + +| Flag | Default | Description | +|------|---------|-------------| +| `--memory-pool-bytes` | `104857600` (100 MB) | Total KV-cache memory pool, in bytes | +| `--block-size-bytes` | `512` | Size of a single KV block, in bytes | +| `--tokens-per-block` | `16` | Number of tokens whose KV state fits in one block | +| `--max-batch-size` | `32` | Maximum number of sequences batched together per step | +| `--default-max-tokens` | `100` | Completion-token budget when a request omits `max_tokens` | + +## Examples + +```bash +# Larger KV-cache pool and batch size, bigger default generation budget +puma serve qwen/qwen2.5-0.5b \ + --memory-pool-bytes 209715200 \ + --max-batch-size 64 \ + --default-max-tokens 256 + +# The same flags work for interactive run +puma run qwen/qwen2.5-0.5b --tokens-per-block 32 +``` + +## Notes + +- **`--tokens-per-block`** must be greater than `0`. +- **`--default-max-tokens`** only applies when a request omits `max_tokens`; an + explicit per-request `max_tokens` always takes precedence. +- The flag defaults are sourced from `EngineConfig` in + `src/backend/llm_engine.rs`, so the CLI and the library stay in sync. diff --git a/src/api/chat.rs b/src/api/chat.rs index c87f1f3..38b4a95 100644 --- a/src/api/chat.rs +++ b/src/api/chat.rs @@ -97,7 +97,7 @@ async fn chat_completions_non_stream( .generate( &req.model, &prompt, - req.max_tokens.unwrap_or(100), + req.max_tokens.unwrap_or_else(|| engine.default_max_tokens()), req.temperature.unwrap_or(0.7), ) .await?; @@ -169,7 +169,7 @@ async fn chat_completions_stream( .generate_stream( &model, &prompt, - req.max_tokens.unwrap_or(100), + req.max_tokens.unwrap_or_else(|| engine.default_max_tokens()), req.temperature.unwrap_or(0.7), ) .await diff --git a/src/api/completions.rs b/src/api/completions.rs index c9bedf8..2a103e0 100644 --- a/src/api/completions.rs +++ b/src/api/completions.rs @@ -74,7 +74,7 @@ pub async fn completions( .generate( &req.model, &prompt, - req.max_tokens.unwrap_or(100), + req.max_tokens.unwrap_or_else(|| engine.default_max_tokens()), req.temperature.unwrap_or(0.7), ) .await diff --git a/src/api/tests.rs b/src/api/tests.rs index fd8939e..41a18a1 100644 --- a/src/api/tests.rs +++ b/src/api/tests.rs @@ -13,8 +13,8 @@ use tempfile::TempDir; use tower::util::ServiceExt; // for `oneshot` and `ready` use super::routes::create_router; -use crate::backend::engine; use crate::backend::mock::MockEngine; +use crate::backend::{engine, EngineConfig}; use crate::registry::model_registry::{CacheInfo, ModelInfo, ModelMetadata, ModelRegistry}; /// Helper to create test app with a pre-registered test model @@ -25,6 +25,7 @@ fn create_test_app() -> (axum::Router, TempDir) { MockEngine::new(), create_test_tokenizer(), "test-model".to_string(), + EngineConfig::default(), ); tokio::spawn(runner.serve()); diff --git a/src/backend/llm_engine.rs b/src/backend/llm_engine.rs index 21e71bc..b5eae91 100644 --- a/src/backend/llm_engine.rs +++ b/src/backend/llm_engine.rs @@ -38,6 +38,8 @@ pub struct EngineHandle { event_tx: mpsc::UnboundedSender, seq_id_counter: Arc, model: String, + /// Completion-token budget to use when a request omits `max_tokens`. + default_max_tokens: usize, } impl EngineHandle { @@ -129,6 +131,11 @@ impl EngineHandle { pub fn model(&self) -> &str { &self.model } + + /// Completion-token budget to apply when a request omits `max_tokens`. + pub fn default_max_tokens(&self) -> usize { + self.default_max_tokens + } } /// The engine itself: owns the scheduler and backend, drives the event loop. @@ -336,26 +343,80 @@ fn stream_flush(decoded: &str, sent_len: usize) -> Option<&str> { } } +/// Tunable engine parameters. +/// +/// These were previously hardcoded inside [`engine`]. Construct via +/// [`EngineConfig::default`] and override fields as needed: +/// +/// ``` +/// # use puma::backend::llm_engine::EngineConfig; +/// let cfg = EngineConfig { max_batch_size: 64, ..Default::default() }; +/// ``` +#[derive(Debug, Clone)] +pub struct EngineConfig { + /// Total KV-cache memory pool, in bytes. + pub memory_pool_bytes: usize, + /// Size of a single KV block, in bytes. + pub block_size_bytes: usize, + /// Number of tokens whose KV state fits in one block. Must be > 0. + pub tokens_per_block: usize, + /// Maximum number of sequences batched together per step. + pub max_batch_size: usize, + /// Completion-token budget applied when a request does not specify one. + pub default_max_tokens: usize, +} + +impl EngineConfig { + /// Default KV-cache memory pool: 100 MB. + pub const DEFAULT_MEMORY_POOL_BYTES: usize = 1024 * 1024 * 100; + /// Default KV block size, in bytes. + pub const DEFAULT_BLOCK_SIZE_BYTES: usize = 512; + /// Default tokens per KV block. + pub const DEFAULT_TOKENS_PER_BLOCK: usize = 16; + /// Default maximum batch size. + pub const DEFAULT_MAX_BATCH_SIZE: usize = 32; + /// Default completion-token budget when a request omits `max_tokens`. + pub const DEFAULT_MAX_TOKENS: usize = 100; +} + +impl Default for EngineConfig { + fn default() -> Self { + Self { + memory_pool_bytes: Self::DEFAULT_MEMORY_POOL_BYTES, + block_size_bytes: Self::DEFAULT_BLOCK_SIZE_BYTES, + tokens_per_block: Self::DEFAULT_TOKENS_PER_BLOCK, + max_batch_size: Self::DEFAULT_MAX_BATCH_SIZE, + default_max_tokens: Self::DEFAULT_MAX_TOKENS, + } + } +} + /// Construct a paired [`EngineHandle`] and [`EngineRunner`]. /// -/// Spawn `runner.serve()` on a task and share the returned handle with the API / -/// CLI. All request submission goes through events, so the handle never -/// touches the scheduler directly. +/// Pass [`EngineConfig::default`] for the standard settings, or override fields +/// to tune memory/batching. Spawn `runner.serve()` on a task and share the +/// returned handle with the API / CLI. All request submission goes through +/// events, so the handle never touches the scheduler directly. pub fn engine( backend: B, tokenizer: Tokenizer, model: String, + config: EngineConfig, ) -> (EngineHandle, EngineRunner) { - // Create block manager (100MB memory pool, 512 bytes per block) - let allocator = Box::new(CpuAllocator::new(1024 * 1024 * 100)); - let block_manager = BlockManager::new(allocator, 512); + // KV-cache memory pool + block manager. + let allocator = Box::new(CpuAllocator::new(config.memory_pool_bytes)); + let block_manager = BlockManager::new(allocator, config.block_size_bytes); // Event channel: the handle produces events, the scheduler consumes them. let (event_tx, event_rx) = mpsc::unbounded_channel(); - // Create scheduler (max 32 batch size, 16 tokens per block); it owns the - // event receiver. - let scheduler = Scheduler::new(block_manager, event_rx, 32, 16); + // The scheduler owns the event receiver. + let scheduler = Scheduler::new( + block_manager, + event_rx, + config.max_batch_size, + config.tokens_per_block, + ); let tokenizer = Arc::new(tokenizer); @@ -364,6 +425,7 @@ pub fn engine( event_tx, seq_id_counter: Arc::new(AtomicU64::new(1)), model, + default_max_tokens: config.default_max_tokens, }; let runner = EngineRunner { @@ -392,7 +454,12 @@ mod tests { async fn test_llm_engine() { let backend = MockEngine::new(); let tokenizer = create_test_tokenizer(); - let (handle, runner) = engine(backend, tokenizer, "test-model".to_string()); + let (handle, runner) = engine( + backend, + tokenizer, + "test-model".to_string(), + EngineConfig::default(), + ); tokio::spawn(runner.serve()); diff --git a/src/backend/mod.rs b/src/backend/mod.rs index 68e2dab..c2c7de0 100644 --- a/src/backend/mod.rs +++ b/src/backend/mod.rs @@ -2,4 +2,4 @@ pub mod engine; pub mod llm_engine; pub mod mock; -pub use llm_engine::{engine, EngineHandle}; +pub use llm_engine::{engine, EngineConfig, EngineHandle}; diff --git a/src/cli/commands.rs b/src/cli/commands.rs index c1ab516..8a07fff 100644 --- a/src/cli/commands.rs +++ b/src/cli/commands.rs @@ -4,8 +4,8 @@ use prettytable::{format, row, Table}; use tokenizers::Tokenizer; -use crate::backend::engine; use crate::backend::mock::MockEngine; +use crate::backend::{engine, EngineConfig}; use crate::cli::{chat, inspect, ls, rm}; use crate::downloader::{self, Provider}; use crate::registry::model_registry::ModelRegistry; @@ -69,6 +69,9 @@ struct RunArgs { default_value = "huggingface" )] provider: Provider, + + #[command(flatten)] + engine: EngineArgs, } #[derive(Parser)] @@ -83,6 +86,48 @@ struct ServeArgs { /// Port to listen on #[arg(short, long, default_value = "8000")] port: u16, + + #[command(flatten)] + engine: EngineArgs, +} + +/// Engine tuning flags shared by `run` and `serve`. +/// +/// Defaults mirror [`EngineConfig::default`] via its `DEFAULT_*` constants, so +/// the two never drift. +#[derive(Parser, Clone)] +struct EngineArgs { + /// Total KV-cache memory pool, in bytes + #[arg(long, default_value_t = EngineConfig::DEFAULT_MEMORY_POOL_BYTES)] + memory_pool_bytes: usize, + + /// Size of a single KV block, in bytes + #[arg(long, default_value_t = EngineConfig::DEFAULT_BLOCK_SIZE_BYTES)] + block_size_bytes: usize, + + /// Number of tokens whose KV state fits in one block + #[arg(long, default_value_t = EngineConfig::DEFAULT_TOKENS_PER_BLOCK)] + tokens_per_block: usize, + + /// Maximum number of sequences batched together per step + #[arg(long, default_value_t = EngineConfig::DEFAULT_MAX_BATCH_SIZE)] + max_batch_size: usize, + + /// Completion-token budget when a request omits max_tokens + #[arg(long, default_value_t = EngineConfig::DEFAULT_MAX_TOKENS)] + default_max_tokens: usize, +} + +impl EngineArgs { + fn to_config(&self) -> EngineConfig { + EngineConfig { + memory_pool_bytes: self.memory_pool_bytes, + block_size_bytes: self.block_size_bytes, + tokens_per_block: self.tokens_per_block, + max_batch_size: self.max_batch_size, + default_max_tokens: self.default_max_tokens, + } + } } #[derive(Parser)] @@ -246,7 +291,8 @@ pub async fn run(cli: Cli) { let backend = MockEngine::new(); // Create engine: cheap send-side handle + runner that owns the scheduler - let (handle, runner) = engine(backend, tokenizer, args.model.clone()); + let (handle, runner) = + engine(backend, tokenizer, args.model.clone(), args.engine.to_config()); // Spawn the runner's event loop; the handle submits work via events tokio::spawn(runner.serve()); @@ -310,7 +356,10 @@ pub async fn run(cli: Cli) { } } - if let Err(e) = crate::cli::serve::execute(&args.host, args.port, &args.model).await { + if let Err(e) = + crate::cli::serve::execute(&args.host, args.port, &args.model, args.engine.to_config()) + .await + { eprintln!("Error starting server: {}", e); std::process::exit(1); } @@ -577,4 +626,50 @@ mod tests { let result = app.try_get_matches_from(vec!["puma", "run", "test/model", "-p", "ms"]); assert!(result.is_ok()); } + + #[test] + fn test_engine_args_default_to_config() { + use clap::Parser; + let cli = Cli::parse_from(vec!["puma", "serve", "test/model"]); + let Commands::SERVE(args) = cli.command else { + panic!("expected SERVE"); + }; + let cfg = args.engine.to_config(); + // With no flags, the config matches EngineConfig::default(). + let default = EngineConfig::default(); + assert_eq!(cfg.memory_pool_bytes, default.memory_pool_bytes); + assert_eq!(cfg.block_size_bytes, default.block_size_bytes); + assert_eq!(cfg.tokens_per_block, default.tokens_per_block); + assert_eq!(cfg.max_batch_size, default.max_batch_size); + assert_eq!(cfg.default_max_tokens, default.default_max_tokens); + } + + #[test] + fn test_engine_args_overrides_flow_to_config() { + use clap::Parser; + let cli = Cli::parse_from(vec![ + "puma", + "run", + "test/model", + "--block-size-bytes", + "1024", + "--tokens-per-block", + "32", + "--max-batch-size", + "8", + "--default-max-tokens", + "256", + "--memory-pool-bytes", + "2048", + ]); + let Commands::RUN(args) = cli.command else { + panic!("expected RUN"); + }; + let cfg = args.engine.to_config(); + assert_eq!(cfg.memory_pool_bytes, 2048); + assert_eq!(cfg.block_size_bytes, 1024); + assert_eq!(cfg.tokens_per_block, 32); + assert_eq!(cfg.max_batch_size, 8); + assert_eq!(cfg.default_max_tokens, 256); + } } diff --git a/src/cli/serve.rs b/src/cli/serve.rs index 55aab77..a2538a5 100644 --- a/src/cli/serve.rs +++ b/src/cli/serve.rs @@ -7,6 +7,7 @@ use tracing::{debug, info}; use crate::api::routes::create_router; use crate::backend::engine; use crate::backend::mock::MockEngine; +use crate::backend::EngineConfig; use crate::registry::model_registry::ModelRegistry; /// Execute the serve command @@ -14,6 +15,7 @@ pub async fn execute( host: &str, port: u16, model_name: &str, + config: EngineConfig, ) -> Result<(), Box> { println!( "{}", @@ -40,7 +42,7 @@ pub async fn execute( let tokenizer = Tokenizer::new(BPE::default()); // Create engine: cheap send-side handle + runner that owns the scheduler - let (handle, runner) = engine(backend, tokenizer, model_name.to_string()); + let (handle, runner) = engine(backend, tokenizer, model_name.to_string(), config); // Spawn the runner's event loop; the handle submits work via events tokio::spawn(runner.serve()); From 2b90c4a603f9236e412b33fbc9f7ff459621ba81 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Sun, 19 Jul 2026 12:02:20 +0100 Subject: [PATCH 2/3] fix lint Signed-off-by: kerthcet --- src/api/chat.rs | 6 ++++-- src/api/completions.rs | 3 ++- src/cli/commands.rs | 18 +++++++++++++----- 3 files changed, 19 insertions(+), 8 deletions(-) diff --git a/src/api/chat.rs b/src/api/chat.rs index 38b4a95..e155733 100644 --- a/src/api/chat.rs +++ b/src/api/chat.rs @@ -97,7 +97,8 @@ async fn chat_completions_non_stream( .generate( &req.model, &prompt, - req.max_tokens.unwrap_or_else(|| engine.default_max_tokens()), + req.max_tokens + .unwrap_or_else(|| engine.default_max_tokens()), req.temperature.unwrap_or(0.7), ) .await?; @@ -169,7 +170,8 @@ async fn chat_completions_stream( .generate_stream( &model, &prompt, - req.max_tokens.unwrap_or_else(|| engine.default_max_tokens()), + req.max_tokens + .unwrap_or_else(|| engine.default_max_tokens()), req.temperature.unwrap_or(0.7), ) .await diff --git a/src/api/completions.rs b/src/api/completions.rs index 2a103e0..2e4ca1e 100644 --- a/src/api/completions.rs +++ b/src/api/completions.rs @@ -74,7 +74,8 @@ pub async fn completions( .generate( &req.model, &prompt, - req.max_tokens.unwrap_or_else(|| engine.default_max_tokens()), + req.max_tokens + .unwrap_or_else(|| engine.default_max_tokens()), req.temperature.unwrap_or(0.7), ) .await diff --git a/src/cli/commands.rs b/src/cli/commands.rs index 8a07fff..a7b8ae6 100644 --- a/src/cli/commands.rs +++ b/src/cli/commands.rs @@ -291,8 +291,12 @@ pub async fn run(cli: Cli) { let backend = MockEngine::new(); // Create engine: cheap send-side handle + runner that owns the scheduler - let (handle, runner) = - engine(backend, tokenizer, args.model.clone(), args.engine.to_config()); + let (handle, runner) = engine( + backend, + tokenizer, + args.model.clone(), + args.engine.to_config(), + ); // Spawn the runner's event loop; the handle submits work via events tokio::spawn(runner.serve()); @@ -356,9 +360,13 @@ pub async fn run(cli: Cli) { } } - if let Err(e) = - crate::cli::serve::execute(&args.host, args.port, &args.model, args.engine.to_config()) - .await + if let Err(e) = crate::cli::serve::execute( + &args.host, + args.port, + &args.model, + args.engine.to_config(), + ) + .await { eprintln!("Error starting server: {}", e); std::process::exit(1); From 3afa889c71fc898b654524744bdd28f723916566 Mon Sep 17 00:00:00 2001 From: kerthcet Date: Sun, 19 Jul 2026 12:08:24 +0100 Subject: [PATCH 3/3] fix test Signed-off-by: kerthcet --- src/cli/commands.rs | 78 ++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 70 insertions(+), 8 deletions(-) diff --git a/src/cli/commands.rs b/src/cli/commands.rs index a7b8ae6..e55457c 100644 --- a/src/cli/commands.rs +++ b/src/cli/commands.rs @@ -97,20 +97,36 @@ struct ServeArgs { /// the two never drift. #[derive(Parser, Clone)] struct EngineArgs { - /// Total KV-cache memory pool, in bytes - #[arg(long, default_value_t = EngineConfig::DEFAULT_MEMORY_POOL_BYTES)] + /// Total KV-cache memory pool, in bytes (must be > 0) + #[arg( + long, + default_value_t = EngineConfig::DEFAULT_MEMORY_POOL_BYTES, + value_parser = parse_positive_usize + )] memory_pool_bytes: usize, - /// Size of a single KV block, in bytes - #[arg(long, default_value_t = EngineConfig::DEFAULT_BLOCK_SIZE_BYTES)] + /// Size of a single KV block, in bytes (must be > 0) + #[arg( + long, + default_value_t = EngineConfig::DEFAULT_BLOCK_SIZE_BYTES, + value_parser = parse_positive_usize + )] block_size_bytes: usize, - /// Number of tokens whose KV state fits in one block - #[arg(long, default_value_t = EngineConfig::DEFAULT_TOKENS_PER_BLOCK)] + /// Number of tokens whose KV state fits in one block (must be > 0) + #[arg( + long, + default_value_t = EngineConfig::DEFAULT_TOKENS_PER_BLOCK, + value_parser = parse_positive_usize + )] tokens_per_block: usize, - /// Maximum number of sequences batched together per step - #[arg(long, default_value_t = EngineConfig::DEFAULT_MAX_BATCH_SIZE)] + /// Maximum number of sequences batched together per step (must be > 0) + #[arg( + long, + default_value_t = EngineConfig::DEFAULT_MAX_BATCH_SIZE, + value_parser = parse_positive_usize + )] max_batch_size: usize, /// Completion-token budget when a request omits max_tokens @@ -118,6 +134,19 @@ struct EngineArgs { default_max_tokens: usize, } +/// Parse a `usize` argument and reject `0`, since these engine dimensions are +/// used as divisors / capacities that must be positive. +fn parse_positive_usize(s: &str) -> Result { + let value: usize = s + .parse() + .map_err(|_| format!("`{s}` is not a valid number"))?; + if value == 0 { + Err("value must be greater than 0".to_string()) + } else { + Ok(value) + } +} + impl EngineArgs { fn to_config(&self) -> EngineConfig { EngineConfig { @@ -680,4 +709,37 @@ mod tests { assert_eq!(cfg.max_batch_size, 8); assert_eq!(cfg.default_max_tokens, 256); } + + #[test] + fn test_engine_args_reject_zero() { + use clap::CommandFactory; + let app = Cli::command(); + // A zero for a must-be-positive dimension is rejected at parse time. + let result = app.clone().try_get_matches_from(vec![ + "puma", + "serve", + "test/model", + "--tokens-per-block", + "0", + ]); + assert!(result.is_err()); + + // A non-numeric value is also rejected. + let result = app.try_get_matches_from(vec![ + "puma", + "serve", + "test/model", + "--max-batch-size", + "abc", + ]); + assert!(result.is_err()); + } + + #[test] + fn test_parse_positive_usize() { + assert_eq!(parse_positive_usize("16"), Ok(16)); + assert!(parse_positive_usize("0").is_err()); + assert!(parse_positive_usize("-1").is_err()); + assert!(parse_positive_usize("abc").is_err()); + } }