diff --git a/Cargo.toml b/Cargo.toml index 45f738dd..181cdd33 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "candle-vllm" -version = "0.7.1" +version = "0.7.2" edition = "2021" default-run = "candle-vllm" @@ -46,7 +46,7 @@ dirs = "5.0.1" minijinja = { version = "2.10.2", features = ["builtins", "json"] } minijinja-contrib = { version = "2.10.2", features = ["pycompat"] } thiserror = "1.0.58" -attention-rs = { git = "https://github.com/guoqingbao/attention.rs.git", version="0.5.2", rev = "36ba0bb" } +attention-rs = { git = "https://github.com/guoqingbao/attention.rs.git", version="0.5.2", rev = "2ea587f" } metal = { version = "0.27.0", features = ["mps"], optional = true } lazy_static = {version = "1.4.0"} interprocess = "2.2.2" @@ -85,8 +85,7 @@ metal = ["candle-core/metal", "candle-nn/metal", "dep:metal", "attention-rs/meta cudnn = ["candle-core/cudnn"] flashattn = ["attention-rs/flashattn", "attention-rs/no-fp8-kvcache"] flashinfer = ["attention-rs/flashinfer"] -trtllm = ["flashinfer", "attention-rs/trtllm"] mkl = ["dep:intel-mkl-src", "candle-core/mkl", "candle-nn/mkl"] nccl = ["candle-core/nccl"] mpi = ["candle-core/nccl", "dep:mpi"] -graph = ["attention-rs/graph", "candle-core/graph"] +graph = ["cuda", "attention-rs/graph", "candle-core/graph"] diff --git a/README-CN.md b/README-CN.md index 941ecd84..9d56742c 100644 --- a/README-CN.md +++ b/README-CN.md @@ -41,7 +41,7 @@ | #4 | **QWen2/Qwen3 Dense** |96 tks/s (8B)|135 tks/s **(8B, Q4k)**| | #5 | **QWen3 MoE** |92 tks/s **(30B)**|114 tks/s **(30B, Q4K)** | | #6 | **QWen3-Next MoE** |71 tks/s **(80B, BF16, tp=2)**|TBD| - | #7 | **QWen3.5 Dense** |30 tks/s **(27B, BF16)**|~42 tks/s **(27B, Q4K / FP8)** | + | #7 | **QWen3.5/3.6 Dense** |30 tks/s **(27B, BF16)**|~42 tks/s **(27B, Q4K / FP8)** | | #8 | **QWen3.5/3.6 MoE** |82 tks/s **(35B)**|93 tks/s **(35B, Q4K)** | | #9 | **Yi** |148 tks/s (6B)| 180 tks/s (6B, Q4k)| | #10 | **StableLM** |223 tks/s (3B)|-| @@ -138,16 +138,16 @@ cargo install --features metal --path . **示例:** ```shell - [RUST_LOG=warn] cargo run [--release --features cuda,nccl,flashinfer,cutlass,graph] -- [--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.7 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0] [--m Qwen/Qwen3.5-27B-FP8] [--fp8-kvcache] [--ui-server] + [RUST_LOG=warn] cargo run [--release --features cuda,nccl,flashinfer,cutlass,graph] -- [--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.5 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0] [--m Qwen/Qwen3.6-27B-FP8] [--fp8-kvcache] [--ui-server] ``` `ENV_PARAM`: RUST_LOG=warn `BUILD_PARAM`: --release --features cuda,nccl,flashinfer,cutlass,graph - `PROGRAM_PARAM`:--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.7 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0 + `PROGRAM_PARAM`:--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.5 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0 - `MODEL_ID/MODEL_WEIGHT_PATH`: --m Qwen/Qwen3.5-27B-FP8(或使用 `--w` 指定本地模型路径) + `MODEL_ID/MODEL_WEIGHT_PATH`: --m Qwen/Qwen3.6-27B-FP8(或使用 `--w` 指定本地模型路径) `CACHE CONFIG`: --fp8-kvcache @@ -184,23 +184,23 @@ docker run --rm -it --gpus all --network host -v /home:/home -v /data:/data cand **本地路径 (ISQ量化, +UI Server)** ```shell - candle-vllm --p 8000 --d 0,1 --w /home/Qwen3.5-27B/ --isq q4k --ui-server --prefix-cache + candle-vllm --p 8000 --d 0,1 --w /home/Qwen3.6-27B/ --isq q4k --ui-server --prefix-cache ``` **模型ID(从Huggingface下载)** ```shell - candle-vllm --m Qwen/Qwen3.5-35B-A3B --ui-server --prefix-cache + candle-vllm --m Qwen/Qwen3.6-35B-A3B --ui-server --prefix-cache ``` **手动设置 YaRN 缩放** ```shell - candle-vllm --m Qwen/Qwen3.5-35B-A3B --yarn-scaling-factor 4.0 --ui-server --prefix-cache + candle-vllm --m Qwen/Qwen3.6-35B-A3B --yarn-scaling-factor 4.0 --ui-server --prefix-cache ``` **FP8 模型** (block-wise量化, 通过增加`cutlass`特性构建) ```shell - candle-vllm --m Qwen/Qwen3.5-27B-FP8 --ui-server --prefix-cache + candle-vllm --m Qwen/Qwen3.6-27B-FP8 --ui-server --prefix-cache ``` **FP4 模型** (MXFP4/NVFP4, 暂不支持MLX量化格式) ```shell @@ -255,7 +255,7 @@ docker run --rm -it --gpus all --network host -v /home:/home -v /data:/data cand **只需在运行未量化模型时添加`isq`参数** ```shell - candle-vllm --m Qwen/Qwen3.5-27B --isq q4k + candle-vllm --m Qwen/Qwen3.6-27B --isq q4k ``` 注:原位量化加载可能需要更长的加载时间,原位`isq`参数选项:["q4_0", "q4_1", "q5_0", "q5_1", "q8_0", "q2k", "q3k","q4k","q5k","q6k"] @@ -653,7 +653,7 @@ docker run --rm -it --gpus all --network host -v /home:/home -v /data:/data cand 显示详情 `--mem` (`kvcache-mem-gpu`) 用于以 MB 为单位设置固定 KV Cache 预算,默认值为 `4096` MB。 - `--gpu-memory-fraction` 提供一个更轻量的自动模式。不显式指定时,默认值为 `0.7`。模型加载完成后,candle-vllm 会探测每张已加载的 CUDA 或 Metal 设备,并按以下公式计算 KV Cache 预算: + `--gpu-memory-fraction` 提供一个更轻量的自动模式。不显式指定时,默认值为 `0.5`。模型加载完成后,candle-vllm 会探测每张已加载的 CUDA 或 Metal 设备,并按以下公式计算 KV Cache 预算: ``` gpu_memory_fraction * total_gpu_memory - current_memory_usage @@ -662,7 +662,7 @@ docker run --rm -it --gpus all --network host -v /home:/home -v /data:/data cand 多卡场景下,会取所有 rank 中最小的结果作为每个 rank 的 KV Cache 预算。例如: ``` - candle-vllm --w /home/Qwen3-Coder-30B-A3B-Instruct-FP8 --d 0,1 --gpu-memory-fraction 0.7 + candle-vllm --w /home/Qwen3-Coder-30B-A3B-Instruct-FP8 --d 0,1 --gpu-memory-fraction 0.5 ``` 当你需要显式固定缓存预算时,用 `--mem`。当你希望服务根据模型加载后的可用显存自动调整时,用 `--gpu-memory-fraction`。 diff --git a/README.md b/README.md index 34c4da97..edf73105 100644 --- a/README.md +++ b/README.md @@ -42,7 +42,7 @@ Efficient, easy-to-use platform for inference and serving local LLMs including a | #4 | **QWen2/Qwen3 Dense** |96 tks/s (8B)|135 tks/s **(8B, Q4k)**| | #5 | **QWen3 MoE** |92 tks/s **(30B)**|114 tks/s **(30B, Q4K)** | | #6 | **QWen3-Next MoE** |71 tks/s **(80B, BF16, tp=2)**|TBD| - | #7 | **QWen3.5 Dense** |30 tks/s **(27B, BF16)**|~42 tks/s **(27B, Q4K / FP8)** | + | #7 | **QWen3.5/3.6 Dense** |30 tks/s **(27B, BF16)**|~42 tks/s **(27B, Q4K / FP8)** | | #8 | **QWen3.5/3.6 MoE** |82 tks/s **(35B)**|93 tks/s **(35B, Q4K)** | | #9 | **Yi** |148 tks/s (6B)| 180 tks/s (6B, Q4k)| | #10 | **StableLM** |223 tks/s (3B)|-| @@ -141,16 +141,16 @@ cargo install --features metal --path . **Example:** ```shell - [RUST_LOG=warn] cargo run [--release --features cuda,nccl,flashinfer,cutlass,graph] -- [--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.7 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0] [--m Qwen/Qwen3.5-27B-FP8] [--fp8-kvcache] [--ui-server] + [RUST_LOG=warn] cargo run [--release --features cuda,nccl,flashinfer,cutlass,graph] -- [--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.5 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0] [--m Qwen/Qwen3.6-27B-FP8] [--fp8-kvcache] [--ui-server] ``` `ENV_PARAM`: RUST_LOG=warn `BUILD_PARAM`: --release --features cuda,nccl,flashinfer,cutlass,graph - `PROGRAM_PARAM`:--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.7 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0 + `PROGRAM_PARAM`:--log --dtype bf16 --p 2000 --d 0,1 --gpu-memory-fraction 0.5 --isq q4k --prefill-chunk-size 8192 --frequency-penalty 1.1 --presence-penalty 1.1 --enforce-parser qwen_coder --yarn-scaling-factor 4.0 - `MODEL_ID/MODEL_WEIGHT_PATH`: --m Qwen/Qwen3.5-27B-FP8 (or `--w` specify local model path) + `MODEL_ID/MODEL_WEIGHT_PATH`: --m Qwen/Qwen3.6-27B-FP8 (or `--w` specify local model path) `CACHE CONFIG`: --fp8-kvcache @@ -188,23 +188,23 @@ docker run --rm -it --gpus all --network host -v /home:/home -v /data:/data cand **Local Path (ISQ, +UI Server)** ```shell - candle-vllm --p 8000 --d 0,1 --w /home/Qwen3.5-27B/ --isq q4k --ui-server --prefix-cache + candle-vllm --p 8000 --d 0,1 --w /home/Qwen3.6-27B/ --isq q4k --ui-server --prefix-cache ``` **Model-ID (download from Huggingface)** ```shell - candle-vllm --m Qwen/Qwen3.5-35B-A3B --ui-server --prefix-cache + candle-vllm --m Qwen/Qwen3.6-35B-A3B --ui-server --prefix-cache ``` **Manual YaRN scaling** ```shell - candle-vllm --m Qwen/Qwen3.5-35B-A3B --yarn-scaling-factor 4.0 --ui-server --prefix-cache + candle-vllm --m Qwen/Qwen3.6-35B-A3B --yarn-scaling-factor 4.0 --ui-server --prefix-cache ``` **FP8 Model** (block-wise quant, build with `cutlass` feature) ```shell - candle-vllm --m Qwen/Qwen3.5-27B-FP8 --ui-server --prefix-cache + candle-vllm --m Qwen/Qwen3.6-27B-FP8 --ui-server --prefix-cache ``` ```shell @@ -265,7 +265,7 @@ docker run --rm -it --gpus all --network host -v /home:/home -v /data:/data cand **Simply add `isq` parameter when running unquantized models** ```shell - candle-vllm --p 2000 --m Qwen/Qwen3.5-27B --isq q4k + candle-vllm --p 2000 --m Qwen/Qwen3.6-27B --isq q4k ``` Options for in-site `isq` parameters: ["q4_0", "q4_1", "q5_0", "q5_1", "q8_0", "q2k", "q3k","q4k","q5k","q6k"] @@ -659,7 +659,7 @@ Chat frontend (any frontend compatible with openai API, simple options available Show details The `--mem` (`kvcache-mem-gpu`) parameter sets a fixed KV cache budget in MB. By default this is `4096` MB. - The `--gpu-memory-fraction` parameter is a lighter-weight auto mode. When omitted, it defaults to `0.7`. After the model finishes loading, candle-vllm probes each loaded CUDA or Metal device and computes the KV cache budget as: + The `--gpu-memory-fraction` parameter is a lighter-weight auto mode. When omitted, it defaults to `0.5`. After the model finishes loading, candle-vllm probes each loaded CUDA or Metal device and computes the KV cache budget as: ``` gpu_memory_fraction * remaining_gpu_memory_after_model_load @@ -668,7 +668,7 @@ Chat frontend (any frontend compatible with openai API, simple options available This means the fraction directly controls how much of the free GPU memory left after model load can be used for the combined GPU cache budget. The minimum detected budget across ranks is used as the KV cache budget per rank. For example: ``` - candle-vllm --w /home/Qwen3-Coder-30B-A3B-Instruct-FP8 --d 0,1 --gpu-memory-fraction 0.7 + candle-vllm --w /home/Qwen3-Coder-30B-A3B-Instruct-FP8 --d 0,1 --gpu-memory-fraction 0.5 ``` Use `--mem` when you want an explicit fixed budget. Use `--gpu-memory-fraction` when you want the server to adapt to the currently available GPU memory after model load. diff --git a/docs/kilocode.md b/docs/kilocode.md index f2c0c3ed..83a7c95e 100644 --- a/docs/kilocode.md +++ b/docs/kilocode.md @@ -10,11 +10,11 @@ Kilo Code -> Candle-vLLM (OpenAI-compatible) ```bash cargo run --release --features cuda,nccl,graph,flashinfer,cutlass -- \ - --m Qwen/Qwen3.5-27B-FP8 \ + --m Qwen/Qwen3.6-27B-FP8 \ --d 0 \ --prefix-cache \ --p 8000 \ - --gpu-memory-fraction 0.7 \ + --gpu-memory-fraction 0.5 \ --enforce-parser qwen_coder ``` @@ -42,7 +42,7 @@ Create `~/.config/kilo/config.json`: }, "models": { "qwen3-coder": { - "name": "Qwen/Qwen3.5-27B-FP8" + "name": "Qwen/Qwen3.6-27B-FP8" } } } diff --git a/docs/opencode.md b/docs/opencode.md index 742f2979..56632d98 100644 --- a/docs/opencode.md +++ b/docs/opencode.md @@ -10,11 +10,11 @@ OpenCode -> Candle-vLLM (OpenAI-compatible) ```bash cargo run --release --features cuda,nccl,graph,flashinfer,cutlass -- \ - --m Qwen/Qwen3.5-27B-FP8 \ + --m Qwen/Qwen3.6-27B-FP8 \ --d 0 \ --prefix-cache \ --p 8000 \ - --gpu-memory-fraction 0.7 \ + --gpu-memory-fraction 0.5 \ --enforce-parser qwen_coder ``` @@ -56,7 +56,7 @@ Create `~/.config/opencode/config.json`: }, "models": { "qwen3-coder": { - "name": "Qwen/Qwen3.5-27B-FP8" + "name": "Qwen/Qwen3.6-27B-FP8" } } } diff --git a/src/api.rs b/src/api.rs index 27a5d8b4..590e8740 100644 --- a/src/api.rs +++ b/src/api.rs @@ -9,7 +9,9 @@ use crate::openai::PipelineConfig; use crate::scheduler::cache_engine::{CacheConfig, CacheEngine}; use crate::scheduler::prefix_cache::PrefixCacheConfig; use crate::scheduler::SchedulerConfig; -use crate::tools::ToolFormat; +use crate::tools::helpers::{ + build_invalid_tool_call_feedback, build_tool_schema_map, filter_tool_calls, +}; use candle_core::{DType, Result}; use parking_lot::RwLock; use std::collections::HashMap; @@ -66,7 +68,7 @@ impl EngineBuilder { max_num_seqs: 16, block_size: if cfg!(feature = "cuda") { 64 } else { 32 }, kvcache_mem_gpu: 4096, - gpu_memory_fraction: Some(0.7), + gpu_memory_fraction: Some(0.5), kvcache_mem_cpu: 128, temperature: None, top_p: None, @@ -446,12 +448,13 @@ impl Engine { .map_err(candle_core::Error::wrap)?; } - let (prompt, tokenizer, image_data) = { + let (prompt, tokenizer, image_data, resolved_tools) = { let e = self.engine.read(); let (pipeline, _) = e.get_pipeline(0).unwrap(); let tool_config = resolve_tools_for_request(&request.tools, &request.tool_choice, None) .map_err(candle_core::Error::wrap)?; + let resolved_tools = tool_config.tools.clone(); // tokenizer is inside DefaultPipeline let mut conversation = pipeline.conversation.clone(); @@ -499,35 +502,15 @@ impl Engine { } }; - if !tool_config.tools.is_empty() { - let mut tools_prompt = ToolFormat::get_tool_prompt( - &pipeline.tool_config, - &pipeline.tool_model_type, - &pipeline.tool_parser_model_id, - pipeline.enforce_parser.as_deref(), - ); - - // Enforce tool_choice=function - if let crate::openai::ToolChoiceKind::Function(name) = &tool_config.choice { - tools_prompt = format!( - "IMPORTANT: You MUST call the tool \"{}\". Do not respond with plain text.\n\n{}", - name, tools_prompt - ); - } - - let current_system = conversation.get_system_message().unwrap_or_default(); - let new_system = if current_system.is_empty() { - tools_prompt - } else { - format!("{}\n\n{}", current_system, tools_prompt) - }; - conversation.set_system_message(Some(new_system)); - } - let enable_thinking = request.thinking.unwrap_or(true); let prompt = conversation.get_prompt(enable_thinking, &tool_config.tools); - (prompt, pipeline.tokenizer.clone(), image_data) + ( + prompt, + pipeline.tokenizer.clone(), + image_data, + resolved_tools, + ) }; let request_id = format!("cmpl-{}", uuid::Uuid::new_v4()); @@ -553,7 +536,7 @@ impl Engine { let prefilled_reasoning_end = crate::tools::stream_parser::detect_prefilled_reasoning_end_marker(&prompt); - let has_tools = request.tools.as_ref().is_some_and(|t| !t.is_empty()); + let has_tools = !resolved_tools.is_empty(); { let mut e = self.engine.write(); let mut sampling_params = SamplingParams::new( @@ -589,7 +572,7 @@ impl Engine { false, // is_embedding crate::openai::requests::EncodingFormat::default(), crate::openai::requests::EmbeddingType::default(), - request.tools.clone().unwrap_or_default(), + resolved_tools.clone(), image_data, None, // streamer Some(req_notify.clone()), @@ -635,6 +618,45 @@ impl Engine { } } } + if has_tools { + let parser = crate::tools::parser::ToolParser::new(); + let tool_schemas = build_tool_schema_map(&resolved_tools); + for choice in &mut choices { + let parsed_calls = if let Some(calls) = choice.message.tool_calls.take() { + calls + } else if let Some(content) = &choice.message.content { + parser.parse(content) + } else { + Vec::new() + }; + + if parsed_calls.is_empty() { + continue; + } + + let (valid_calls, invalid_calls) = + filter_tool_calls(&parsed_calls, &tool_schemas); + if !invalid_calls.is_empty() { + tracing::warn!( + "Dropped {} invalid tool call(s) before response", + invalid_calls.len() + ); + } + if valid_calls.is_empty() { + if let Some(feedback) = + build_invalid_tool_call_feedback(&invalid_calls, &tool_schemas, None) + { + choice.message.content = Some(feedback); + } + choice.finish_reason = Some("stop".to_string()); + continue; + } + + choice.message.tool_calls = Some(valid_calls); + choice.message.content = None; + choice.finish_reason = Some("tool_calls".to_string()); + } + } Ok(ChatCompletionResponse { id: request_id, choices, diff --git a/src/backend/graph.rs b/src/backend/graph.rs index 323a2737..5d462001 100644 --- a/src/backend/graph.rs +++ b/src/backend/graph.rs @@ -368,6 +368,47 @@ pub fn planned_graph_capture_batches(max_num_seqs: usize) -> Vec { graph_bs } +#[cfg(feature = "flashinfer")] +fn graph_decode_plan( + device: &Device, + params: &FlashInferKvParams, + indptr_host: &[u32], + last_len_host: &[u32], + kv_len_arr_host: &[u32], + batch_size: usize, + is_mla: bool, + enable_cuda_graph: bool, +) -> Result<(Option>, Option>)> { + if is_mla { + let plan = attention_rs::mla::mla_decode_plan( + device, + params.kv_dtype, + indptr_host, + batch_size, + params.num_qo_heads, + params.page_size, + enable_cuda_graph, + )?; + Ok((None, Some(plan))) + } else { + let plan = attention_rs::flashinfer::decode_plan( + device, + params.kv_dtype, + params.out_dtype, + indptr_host, + Some(last_len_host), + Some(kv_len_arr_host), + batch_size, + params.num_qo_heads, + params.num_kv_heads, + params.head_dim, + params.page_size, + enable_cuda_graph, + )?; + Ok((Some(plan), None)) + } +} + impl GraphCapturer { pub fn new( model: M, @@ -439,6 +480,89 @@ impl GraphCapturer { last_len, ) }; + + // Some CUDA kernels lazily initialize shape-specific resources on first use. + // Run the planned decode shapes once before stream capture so those one-time + // allocations do not invalidate CUDA graph capture. + for &bs in self.graph_bs.iter().rev() { + #[cfg(feature = "flashinfer")] + let flashinfer_metadata = { + let mut indptr_host = Vec::with_capacity(bs + 1); + indptr_host.push(0u32); + for i in 0..bs { + indptr_host.push(((i + 1) * max_num_blocks) as u32); + } + + let (decode_plan_info, mla_decode_plan_info, kv_len_arr_host) = + if let Some(params) = self.flashinfer_kv_params { + let mut kv_len_arr_host_bs = Vec::with_capacity(bs); + for i in 0..bs { + let num_pages = indptr_host[i + 1] - indptr_host[i]; + if num_pages == 0 { + kv_len_arr_host_bs.push(0); + } else { + let full = (num_pages - 1) * params.page_size as u32; + kv_len_arr_host_bs.push(full + last_len_host[i]); + } + } + let (dp, mdp) = graph_decode_plan( + device, + ¶ms, + &indptr_host, + &last_len_host[..bs], + &kv_len_arr_host_bs, + bs, + self.is_mla, + false, + )?; + (dp, mdp, Some(kv_len_arr_host_bs)) + } else { + (None, None, None) + }; + + Some(attention_rs::FlashInferMetadata { + indptr: flashinfer_indptr.narrow(0, 0, bs + 1)?, + indptr_host, + indices: flashinfer_indices.narrow(0, 0, bs * max_num_blocks)?, + last_len: flashinfer_last_len.narrow(0, 0, bs)?, + last_len_host: Some(last_len_host[..bs].to_vec()), + kv_len_arr_host, + total_num_rows: None, + batch_indices: None, + positions: None, + use_cuda_graph: false, + decode_plan_info, + prefill_plan_info: None, + mla_decode_plan_info, + mla_prefill_plan_info: None, + }) + }; + #[cfg(not(feature = "flashinfer"))] + let flashinfer_metadata = None; + let input_metadata = InputMetadata { + is_prefill: false, + is_mla: self.is_mla, + sequence_ids: None, + mamba_slot_mapping: Some(mamba_slot_mapping.narrow(0, 0, bs)?), + slot_mapping: slot_mapping.narrow(0, 0, bs)?, + block_tables: Some(block_tables.narrow(0, 0, bs)?), + context_lens: Some(context_lens.narrow(0, 0, bs)?), + cu_seqlens_q: None, + cu_seqlens_k: None, + max_seqlen_q: 0, + max_seqlen_k: 0, + max_context_len: self.max_model_len, + disable_flash_attn: None, + seqlens: None, + flashinfer_metadata, + }; + let input_ids_bs = input_ids.narrow(0, 0, bs)?; + let positions_bs = positions.narrow(0, 0, bs)?; + let _ = self + .model + .forward(&input_ids_bs, &positions_bs, kv_caches, &input_metadata)?; + } + let mut outputs = BTreeMap::::new(); for i in tqdm(0..self.graph_bs.len()).desc(Some("Graph capturing")) { let bs = self.graph_bs[self.graph_bs.len() - i - 1]; @@ -464,35 +588,17 @@ impl GraphCapturer { kv_len_arr_host_bs.push(full + last_len_host[i]); } } - let kv_len_arr_host = kv_len_arr_host_bs.clone(); - if self.is_mla { - let plan = attention_rs::mla::mla_decode_plan( - device, - params.kv_dtype, - &indptr_host, - bs, - params.num_qo_heads, - params.page_size, - true, - )?; - (None, Some(plan), Some(kv_len_arr_host)) - } else { - let plan = attention_rs::flashinfer::decode_plan( - device, - params.kv_dtype, - params.out_dtype, - &indptr_host, - Some(&last_len_host[..bs]), - Some(kv_len_arr_host_bs.as_slice()), - bs, - params.num_qo_heads, - params.num_kv_heads, - params.head_dim, - params.page_size, - true, - )?; - (Some(plan), None, Some(kv_len_arr_host)) - } + let (dp, mdp) = graph_decode_plan( + device, + ¶ms, + &indptr_host, + &last_len_host[..bs], + &kv_len_arr_host_bs, + bs, + self.is_mla, + true, + )?; + (dp, mdp, Some(kv_len_arr_host_bs)) } else { (None, None, None) }; @@ -596,8 +702,7 @@ impl GraphCapturer { } let max_num_blocks = (self.max_model_len + self.block_size - 1) / self.block_size; let input_batch = input_ids.dim(0)?; - let require_exact_batch = input_metadata.mamba_slot_mapping.is_some() - || input_metadata.flashinfer_metadata.is_some(); + let require_exact_batch = input_metadata.mamba_slot_mapping.is_some(); if let Some(graph_vars) = &self.graph_vars { let selected_batch = if require_exact_batch { graph_vars @@ -671,32 +776,22 @@ impl GraphCapturer { .device .as_ref() .ok_or_else(|| candle_core::Error::msg("graph device is missing"))?; - if self.is_mla { - let _ = attention_rs::mla::mla_decode_plan( - dev, - params.kv_dtype, - &indptr_host, - batch, - params.num_qo_heads, - params.page_size, - fm.use_cuda_graph, - )?; - } else { - let _ = attention_rs::flashinfer::decode_plan( - dev, - params.kv_dtype, - params.out_dtype, - &indptr_host, - fm.last_len_host.as_deref(), - fm.kv_len_arr_host.as_deref(), - batch, - params.num_qo_heads, - params.num_kv_heads, - params.head_dim, - params.page_size, - fm.use_cuda_graph, - )?; - } + let last_len_host = fm.last_len_host.as_deref().ok_or_else(|| { + candle_core::Error::msg("graph replay requires last_len_host") + })?; + let kv_len_arr_host = fm.kv_len_arr_host.as_deref().ok_or_else(|| { + candle_core::Error::msg("graph replay requires kv_len_arr_host") + })?; + let _ = graph_decode_plan( + dev, + ¶ms, + &indptr_host, + last_len_host, + kv_len_arr_host, + batch, + self.is_mla, + fm.use_cuda_graph, + )?; } } diff --git a/src/main.rs b/src/main.rs index 82ebe7d5..9a88035c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -84,7 +84,7 @@ struct Args { kvcache_mem_gpu: usize, /// Auto-size GPU KV cache after model load using `fraction * remaining_gpu_mem`. - /// Defaults to 0.7 and takes priority over `--mem` on CUDA/Metal. + /// Defaults to 0.5 and takes priority over `--mem` on CUDA/Metal. #[arg(long)] gpu_memory_fraction: Option, @@ -405,11 +405,9 @@ async fn main() -> Result<()> { .expect("at least one pipeline must be loaded"); let first_config = first_pipeline.get_model_config(); let first_model_dtype = first_pipeline.dtype; - let requested_gpu_memory_fraction = args - .gpu_memory_fraction - .unwrap_or(if cfg!(feature = "cuda") { 0.7 } else { 0.5 }); + let requested_gpu_memory_fraction = args.gpu_memory_fraction.unwrap_or(0.5); let explicit_gpu_memory_fraction = args.gpu_memory_fraction.is_some(); - + let (kvcache_mem_gpu, mamba_cache_budget_bytes, kvcache_budget_desc) = match candle_vllm::detect_kvcache_mem_gpu_mb_for_devices( &devices, diff --git a/src/openai/conversation/default_conversation.rs b/src/openai/conversation/default_conversation.rs index 689903db..977cd1f5 100644 --- a/src/openai/conversation/default_conversation.rs +++ b/src/openai/conversation/default_conversation.rs @@ -2,8 +2,10 @@ use crate::tools::{Tool, ToolCall}; use super::{ApplyChatTemplateError, Message}; use minijinja::{context, value::Kwargs, Environment, Error, ErrorKind, Value as JinjaValue}; +use regex::Regex; use serde::Serialize; use serde_json::Value as JsonValue; +use std::sync::OnceLock; use tokenizers::Tokenizer; pub const ROLES: (&str, &str) = ("USER", "ASSISTANT"); @@ -203,6 +205,24 @@ fn should_escape_nested_xml_tool_markers(tool_markers: &[&str]) -> bool { .any(|marker| marker.starts_with('<') && marker.contains("tool_call")) } +fn normalize_template_source(source: &str) -> String { + static TOJSON_ENSURE_ASCII_RE: OnceLock = OnceLock::new(); + let regex = TOJSON_ENSURE_ASCII_RE.get_or_init(|| { + Regex::new( + r#"(?x) + \| + (?P\s*) + tojson + \( + \s*ensure_ascii\s*=\s*(?:false|true|False|True)\s* + \) + "#, + ) + .expect("valid tojson ensure_ascii regex") + }); + regex.replace_all(source, "|${ws}tojson").into_owned() +} + impl DefaultConversation { pub fn collect_escape_tokens(tokenizer: &Tokenizer, tool_markers: &[&str]) -> Vec { let mut tokens = tokenizer @@ -255,17 +275,43 @@ impl DefaultConversation { escape_special_tokens_in_text(content, &self.escape_tokens, &self.preserve_tokens) } - fn escaped_messages_for_render(&self) -> Vec { - if self.escape_tokens.is_empty() { - return self.messages.clone(); - } - + fn escaped_messages_for_render(&self, enable_thinking: bool) -> Vec { + let need_escape = !self.escape_tokens.is_empty(); + let is_thinking_model = enable_thinking + && self + .escape_tokens + .iter() + .any(|token| token.to_lowercase().contains("think")); self.messages .iter() .map(|message| { let mut escaped = message.clone(); - if !matches!(escaped.role.as_str(), "system" | "developer") { - escaped.content = self.escape_text(&escaped.content); + match escaped.role.as_str() { + "system" | "developer" => {} + "assistant" => { + if let Some((reasoning, remaining)) = + crate::tools::stream_parser::extract_reasoning_content(&escaped.content) + { + if escaped.reasoning_content.is_none() { + escaped.reasoning_content = Some(reasoning); + } + escaped.content = remaining; + } + if is_thinking_model + && escaped.reasoning_content.is_none() + && escaped.tool_calls.is_some() + { + escaped.reasoning_content = Some("...".to_string()); + } + if need_escape { + escaped.content = self.escape_text(&escaped.content); + } + } + _ => { + if need_escape { + escaped.content = self.escape_text(&escaped.content); + } + } } escaped }) @@ -375,7 +421,7 @@ impl DefaultConversation { env.set_lstrip_blocks(true); env.set_trim_blocks(true); env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback); - let template = self.chat_template.as_ref().unwrap(); + let template = normalize_template_source(self.chat_template.as_ref().unwrap()); let mut template = template.replace("[::-1]", "|reverse"); if template.contains("{{ meta }}") { template = template.replace("{%- set meta = message.get(\"metadata\", \"\") %}", ""); @@ -394,7 +440,7 @@ impl DefaultConversation { let template = env .get_template(&self.name) .map_err(ApplyChatTemplateError::GetTemplateError)?; - let render_messages = self.escaped_messages_for_render(); + let render_messages = self.escaped_messages_for_render(enable_thinking); template .render(context! { messages => render_messages, @@ -435,14 +481,23 @@ impl DefaultConversation { Ok(prompt) => prompt, Err(e) => { if self.chat_template.is_some() { - tracing::warn!("apply chat template failed {:?}", e); + if !tools.is_empty() { + tracing::error!( + "Chat template rendering FAILED (tools will be dropped!): {:?}. \ + Tools provided: {:?}. Using built-in fallback which does NOT support tools.", + e, + tools.iter().map(|t| &t.function.name).collect::>() + ); + } else { + tracing::warn!("apply chat template failed {:?}", e); + } } //no chat template exists? using the built-in template let system_prompt = self .system_message .as_ref() .map_or("".to_string(), |msg| format!("<|system|>\n {msg}")); - let render_messages = self.escaped_messages_for_render(); + let render_messages = self.escaped_messages_for_render(thinking); match self.sep_style { SeparatorStyle::AddColonSingle @@ -785,4 +840,25 @@ mod tests { assert!(escaped.contains(&"".to_string())); assert!(!escaped.contains(&"plain_text".to_string())); } + + #[test] + fn normalize_template_source_strips_unsupported_ensure_ascii_kwarg() { + let source = "{{ value | tojson(ensure_ascii=False) }}"; + + let mut raw_env = Environment::new(); + raw_env.add_template("raw", source).unwrap(); + let raw_template = raw_env.get_template("raw").unwrap(); + let raw_err = raw_template + .render(context! { value => "hello" }) + .unwrap_err(); + assert!( + raw_err + .to_string() + .contains("unknown keyword argument 'ensure_ascii'"), + "unexpected raw error: {raw_err}" + ); + + let normalized = normalize_template_source(source); + assert_eq!(normalized, "{{ value | tojson }}"); + } } diff --git a/src/openai/models/deepseek.rs b/src/openai/models/deepseek.rs index 16b280d9..2bf30914 100644 --- a/src/openai/models/deepseek.rs +++ b/src/openai/models/deepseek.rs @@ -218,7 +218,7 @@ impl DeepSeekDecoderLayer { config.hidden_size, config.rms_norm_eps, vb.pp("input_layernorm"), - dtype, + DType::F32, false, )?; @@ -226,7 +226,7 @@ impl DeepSeekDecoderLayer { config.hidden_size, config.rms_norm_eps, vb.pp("post_attention_layernorm"), - dtype, + DType::F32, false, )?; @@ -440,12 +440,8 @@ impl DeepSeek { let mut mla_rope_cfg = cfg.clone(); mla_rope_cfg.head_dim = Some(mla_cfg.qk_rope_head_dim); mla_rope_cfg.partial_rotary_factor = None; - let rotary_dtype = if let Some(qcfg) = &cfg.quantization_config { - if matches!(qcfg.quant_method.as_str(), "nvfp4" | "mxfp4" | "fp8") { - dtype - } else { - DType::F32 - } + let rotary_dtype = if cfg.isq_quant.is_some() || cfg.higher_precision_required() { + DType::F32 } else { dtype }; @@ -476,7 +472,7 @@ impl DeepSeek { cfg.hidden_size, cfg.rms_norm_eps, vb_m.pp("norm"), - dtype, + DType::F32, false, )?; diff --git a/src/openai/models/gemma4.rs b/src/openai/models/gemma4.rs index b03cdb21..4cb1b254 100644 --- a/src/openai/models/gemma4.rs +++ b/src/openai/models/gemma4.rs @@ -7,6 +7,7 @@ use crate::openai::distributed::{embedding, Comm, ReplicatedLinear, VarBuilder}; use crate::openai::models::layers::moe::{ FusedMoe, FusedMoeFp8, FusedMoeISQ, FusedMoeMxfp4, FusedMoeNvfp4, }; +use crate::openai::models::layers::others::{rms_norm, NormX}; use crate::openai::models::mask::get_attention_causal_mask; use crate::openai::models::rotary_emb::DefaultRotaryEmbedding; use crate::openai::models::ScalingValue; @@ -14,7 +15,7 @@ use crate::openai::models::TokenID; use crate::InputMetadata; use candle::{DType, Device, Module, Result, Tensor}; use candle_core as candle; -use candle_nn::{Activation, Linear, RmsNorm}; +use candle_nn::{Activation, Linear}; use either::Either; use parking_lot::RwLock; use std::collections::HashMap; @@ -118,41 +119,6 @@ pub struct Gemma4Config { pub text_config: Gemma4TextConfig, } -struct Gemma4RmsNorm { - inner: RmsNorm, - weight_dtype: DType, -} - -impl Gemma4RmsNorm { - fn new(weight: Tensor, eps: f64) -> Self { - let weight_dtype = weight.dtype(); - Self { - inner: RmsNorm::new(weight, eps), - weight_dtype, - } - } - - fn forward(&self, xs: &Tensor) -> Result { - let in_dtype = xs.dtype(); - let xs = if in_dtype != self.weight_dtype { - xs.to_dtype(self.weight_dtype)? - } else { - xs.clone() - }; - let out = self.inner.forward(&xs)?; - if out.dtype() != in_dtype { - out.to_dtype(in_dtype) - } else { - Ok(out) - } - } -} - -fn rms_norm(dim: usize, eps: f64, vb: VarBuilder) -> Result { - let weight = vb.get(dim, "weight")?; - Ok(Gemma4RmsNorm::new(weight, eps)) -} - struct Gemma4Router { scale: Tensor, proj: Linear, @@ -272,14 +238,14 @@ struct Gemma4DecoderLayer { mlp: Mlp, moe: Option, gemma4_router: Option, - input_layernorm: Gemma4RmsNorm, - post_attention_layernorm: Gemma4RmsNorm, - pre_feedforward_layernorm: Gemma4RmsNorm, - post_feedforward_layernorm: Gemma4RmsNorm, - post_feedforward_layernorm_1: Option, - post_feedforward_layernorm_2: Option, - pre_feedforward_layernorm_2: Option, - post_per_layer_input_norm: Option, + input_layernorm: NormX, + post_attention_layernorm: NormX, + pre_feedforward_layernorm: NormX, + post_feedforward_layernorm: NormX, + post_feedforward_layernorm_1: Option, + post_feedforward_layernorm_2: Option, + pre_feedforward_layernorm_2: Option, + post_per_layer_input_norm: Option, per_layer_input_gate: Option, per_layer_projection: Option, layer_scalar: Tensor, @@ -328,6 +294,7 @@ impl Gemma4DecoderLayer { comm.clone(), sliding_window, k_eq_v, + false, Some(1.0), )?; @@ -402,22 +369,33 @@ impl Gemma4DecoderLayer { (None, None) }; - let input_layernorm = - rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; + let input_layernorm = rms_norm( + cfg.hidden_size, + cfg.rms_norm_eps, + vb.pp("input_layernorm"), + DType::F32, + false, + )?; let post_attention_layernorm = rms_norm( cfg.hidden_size, cfg.rms_norm_eps, vb.pp("post_attention_layernorm"), + DType::F32, + false, )?; let pre_feedforward_layernorm = rms_norm( cfg.hidden_size, cfg.rms_norm_eps, vb.pp("pre_feedforward_layernorm"), + DType::F32, + false, )?; let post_feedforward_layernorm = rms_norm( cfg.hidden_size, cfg.rms_norm_eps, vb.pp("post_feedforward_layernorm"), + DType::F32, + false, )?; let ( @@ -430,16 +408,22 @@ impl Gemma4DecoderLayer { cfg.hidden_size, cfg.rms_norm_eps, vb.pp("post_feedforward_layernorm_1"), + DType::F32, + false, )?), Some(rms_norm( cfg.hidden_size, cfg.rms_norm_eps, vb.pp("post_feedforward_layernorm_2"), + DType::F32, + false, )?), Some(rms_norm( cfg.hidden_size, cfg.rms_norm_eps, vb.pp("pre_feedforward_layernorm_2"), + DType::F32, + false, )?), ) } else { @@ -452,6 +436,8 @@ impl Gemma4DecoderLayer { cfg.hidden_size, cfg.rms_norm_eps, vb.pp("post_per_layer_input_norm"), + DType::F32, + false, )?; let gate = ReplicatedLinear::load_no_bias( cfg.hidden_size, @@ -608,9 +594,9 @@ pub struct Gemma4 { embed_tokens: candle_nn::Embedding, embed_tokens_per_layer: Option, per_layer_model_projection: Option, - per_layer_projection_norm: Option, + per_layer_projection_norm: Option, layers: Vec, - norm: Gemma4RmsNorm, + norm: NormX, lm_head: ReplicatedLinear, device: Device, dtype: DType, @@ -874,6 +860,8 @@ impl Gemma4 { pli_dim, cfg.rms_norm_eps, vb_m.pp("per_layer_projection_norm"), + DType::F32, + false, )?; (Some(emb), Some(proj), Some(norm)) @@ -884,39 +872,74 @@ impl Gemma4 { let (global_rope_theta, partial_rotary_factor) = { let mut theta = cfg.rope_theta; let mut prf = cfg.partial_rotary_factor.unwrap_or(0.25) as f64; - if let Some(fa) = text_cfg - .get("rope_parameters") - .and_then(|rp| rp.get("full_attention")) - { - if let Some(t) = fa.get("rope_theta").and_then(|v| v.as_f64()) { - theta = t; - } - if let Some(p) = fa.get("partial_rotary_factor").and_then(|v| v.as_f64()) { - prf = p; + if let Some(extra) = &cfg.extra_config_json { + let v: serde_json::Value = + serde_json::from_str(extra).unwrap_or(serde_json::Value::Null); + let fa = v + .get("text_config") + .and_then(|tc| tc.get("rope_parameters")) + .and_then(|rp| rp.get("full_attention")); + if let Some(fa) = fa { + if let Some(t) = fa.get("rope_theta").and_then(|v| v.as_f64()) { + theta = t; + } + if let Some(p) = fa.get("partial_rotary_factor").and_then(|v| v.as_f64()) { + prf = p; + } } } (theta, prf) }; + let rope_angles = (partial_rotary_factor * global_head_dim as f64 / 2.0) as usize; + let half_dim = global_head_dim / 2; - let rotary_emb = Arc::new(Self::create_partial_rotary_emb( - DType::F32, - global_head_dim, - partial_rotary_factor, - global_rope_theta, - cfg.max_seq_len, - device, - )?); - - let rotary_emb_local = Arc::new(ScalingRotaryEmbedding::new_sliding( - DType::F32, - cfg.sliding_window, - &Config { - head_dim: Some(swa_head_dim), - rope_theta: rope_local_base_freq, - ..cfg.clone() + let mut inv_freq_vec: Vec = Vec::with_capacity(half_dim); + for i in 0..rope_angles { + inv_freq_vec.push( + 1.0f32 / (global_rope_theta as f32).powf((2 * i) as f32 / global_head_dim as f32), + ); + } + for _ in rope_angles..half_dim { + inv_freq_vec.push(0.0f32); + } + + let inv_freq = Tensor::from_vec(inv_freq_vec, (1, half_dim), &vb.device())?; + let t = Tensor::arange( + 0u32, + cfg.max_position_embeddings.unwrap() as u32, + &vb.device(), + )? + .to_dtype(DType::F32)? + .reshape((cfg.max_position_embeddings.unwrap(), 1))?; + let freqs = t.matmul(&inv_freq)?; + let rotary_emb = Arc::new(ScalingRotaryEmbedding { + 0: DefaultRotaryEmbedding { + cos: freqs.cos()?.to_dtype(DType::F32)?, + sin: freqs.sin()?.to_dtype(DType::F32)?, + is_gpt_neox: true, + rotary_dim: None, }, - device, - )?); + }); + + let swa_head_dim_for_rope = if let Some(extra) = &cfg.extra_config_json { + let v: serde_json::Value = + serde_json::from_str(extra).unwrap_or(serde_json::Value::Null); + v.get("swa_head_dim") + .or_else(|| v.get("text_config").and_then(|tc| tc.get("head_dim"))) + .and_then(|v| v.as_u64()) + .unwrap_or(256) as usize + } else { + 256 + }; + + let mut local_config = cfg.clone(); + local_config.head_dim = Some(swa_head_dim_for_rope); + local_config.partial_rotary_factor = None; + local_config.rope_theta = rope_local_base_freq; + + let rotary_emb_local = Arc::new(ScalingRotaryEmbedding { + 0: DefaultRotaryEmbedding::new(DType::F32, &local_config, &vb.device(), true)?, + }); let mut layers = Vec::with_capacity(cfg.num_hidden_layers); let vb_l = vb_m.pp("layers"); @@ -960,18 +983,24 @@ impl Gemma4 { reporter.write().set_progress(layer_idx + 1); } - let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb_m.pp("norm"))?; - let lm_head = if cfg.tie_word_embeddings { - ReplicatedLinear::from_weight_bias(embed_tokens.embeddings().clone(), None)? - } else { - ReplicatedLinear::load_no_bias( - cfg.hidden_size, - cfg.vocab_size, - vb.pp("lm_head"), - &cfg.isq_quant, - &cfg.quantization_config, - )? - }; + let norm = rms_norm( + cfg.hidden_size, + cfg.rms_norm_eps, + vb_m.pp("norm"), + DType::F32, + false, + )?; + let lm_head = ReplicatedLinear::load_no_bias( + cfg.hidden_size, + cfg.vocab_size, + if cfg.tie_word_embeddings { + vb_m.pp("embed_tokens") + } else { + vb.pp("lm_head") + }, + &None, + &None, + )?; Ok(Self { embed_tokens, @@ -991,41 +1020,6 @@ impl Gemma4 { }) } - fn create_partial_rotary_emb( - dtype: DType, - head_dim: usize, - partial_rotary_factor: f64, - rope_theta: f64, - max_seq_len: usize, - device: &Device, - ) -> Result { - let rope_angles = (partial_rotary_factor * head_dim as f64 / 2.0) as usize; - let half_dim = head_dim / 2; - - let mut inv_freq_vec: Vec = Vec::with_capacity(half_dim); - for i in 0..rope_angles { - inv_freq_vec.push(1.0f32 / (rope_theta as f32).powf((2 * i) as f32 / head_dim as f32)); - } - for _ in rope_angles..half_dim { - inv_freq_vec.push(0.0f32); - } - - let inv_freq = Tensor::from_vec(inv_freq_vec, (1, half_dim), device)?; - let t = Tensor::arange(0u32, max_seq_len as u32, device)? - .to_dtype(DType::F32)? - .reshape((max_seq_len, 1))?; - let freqs = t.matmul(&inv_freq)?; - let sin = freqs.sin()?.to_dtype(dtype)?; - let cos = freqs.cos()?.to_dtype(dtype)?; - - Ok(ScalingRotaryEmbedding(DefaultRotaryEmbedding { - cos, - sin, - is_gpt_neox: true, - rotary_dim: None, - })) - } - fn create_attention_masks( &self, seqlens: &[u32], @@ -1211,7 +1205,10 @@ impl Gemma4 { return xs.to_dtype(DType::F32); } - let logits = self.lm_head.forward(&xs)?.to_dtype(DType::F32)?; + let logits = self + .lm_head + .forward(&xs.to_dtype(self.dtype)?)? + .to_dtype(DType::F32)?; let logits = match self.cfg.final_logit_softcapping { None => logits, diff --git a/src/openai/models/glm4_moe_lite.rs b/src/openai/models/glm4_moe_lite.rs index e71f895b..4d84a82d 100644 --- a/src/openai/models/glm4_moe_lite.rs +++ b/src/openai/models/glm4_moe_lite.rs @@ -212,7 +212,7 @@ impl GLM4MoeLiteDecoderLayer { config.hidden_size, config.rms_norm_eps, vb.pp("input_layernorm"), - dtype, + DType::F32, false, )?; @@ -220,7 +220,7 @@ impl GLM4MoeLiteDecoderLayer { config.hidden_size, config.rms_norm_eps, vb.pp("post_attention_layernorm"), - dtype, + DType::F32, false, )?; @@ -349,12 +349,9 @@ impl GLM4MoeLiteForCausalLM { let mut mla_rope_cfg = config.clone(); mla_rope_cfg.head_dim = Some(mla_cfg.qk_rope_head_dim); mla_rope_cfg.partial_rotary_factor = None; - let rotary_dtype = if let Some(qcfg) = &config.quantization_config { - if matches!(qcfg.quant_method.as_str(), "nvfp4" | "mxfp4" | "fp8") { - dtype - } else { - DType::F32 - } + let is_qvar_builder = config.isq_quant.is_some(); + let rotary_dtype = if is_qvar_builder || config.higher_precision_required() { + DType::F32 } else { dtype }; @@ -385,7 +382,7 @@ impl GLM4MoeLiteForCausalLM { config.hidden_size, config.rms_norm_eps, vb_m.pp("norm"), - dtype, + DType::F32, false, )?; diff --git a/src/openai/models/layers/attention.rs b/src/openai/models/layers/attention.rs index 257b8c59..1a192a9f 100644 --- a/src/openai/models/layers/attention.rs +++ b/src/openai/models/layers/attention.rs @@ -364,6 +364,7 @@ impl Attention { comm, sliding_window, false, + false, None, ) } @@ -375,6 +376,7 @@ impl Attention { comm: Rc, sliding_window: Option, k_eq_v: bool, + qk_l2_norm: bool, attention_scale: Option, ) -> Result { let hidden_sz = cfg.hidden_size; @@ -585,15 +587,11 @@ impl Attention { attn_output_gate, no_per_head_norm: no_per_head_norm_models.contains(&arch), full_dim_qk_norm, - qk_l2_norm: false, + qk_l2_norm, v_norm_eps, }) } - pub fn set_qk_l2_norm(&mut self, enable: bool) { - self.qk_l2_norm = enable; - } - pub fn forward_ext( &self, xs: &Tensor, @@ -653,8 +651,6 @@ impl Attention { let k = key_states.reshape((seq_len, self.num_kv_heads, self.head_dim))?; let v = value_states.reshape((seq_len, self.num_kv_heads, self.head_dim))?; - // Q/K norm weights are loaded in F32 for Qwen3.5/Next; cast activations - // to keep CUDA RMSNorm dtype-consistent. let (q, k) = if q.dtype() != DType::F32 { (q.to_dtype(DType::F32)?, k.to_dtype(DType::F32)?) } else { diff --git a/src/openai/models/layers/deltanet.rs b/src/openai/models/layers/deltanet.rs index 0040c61c..ca2705d3 100644 --- a/src/openai/models/layers/deltanet.rs +++ b/src/openai/models/layers/deltanet.rs @@ -557,7 +557,7 @@ impl GatedDeltaNet { // Get mutable reference to global state for in-place update (optimized prefill) let global_state = mamba_cache.recurrent_state_mut(self.gdn_layer_idx); - xs.device().synchronize()?; + // xs.device().synchronize()?; gdn::gated_delta_rule_recurrence_varlen( &q_scaled, diff --git a/src/openai/models/layers/mla_attention.rs b/src/openai/models/layers/mla_attention.rs index f68b53b5..0884976f 100644 --- a/src/openai/models/layers/mla_attention.rs +++ b/src/openai/models/layers/mla_attention.rs @@ -79,6 +79,7 @@ pub struct MlaAttention { sm_scale: f32, rope_scale: f32, rope_theta: f32, + promote_qk_to_f32: bool, dtype: DType, } @@ -97,6 +98,12 @@ impl MlaAttention { let qk_rope_head_dim = mla_cfg.qk_rope_head_dim; let v_head_dim = mla_cfg.v_head_dim; let q_head_dim = qk_nope_head_dim + qk_rope_head_dim; + let is_qvar_builder = config.isq_quant.is_some(); + let norm_dtype = if is_qvar_builder || config.higher_precision_required() { + DType::F32 + } else { + dtype + }; let (q_a_proj, q_a_layernorm, q_b_proj, q_proj) = if let Some(q_lora_rank) = mla_cfg.q_lora_rank { @@ -112,7 +119,7 @@ impl MlaAttention { q_lora_rank, mla_cfg.rms_norm_eps, vb.pp("q_a_layernorm"), - dtype, + norm_dtype, false, )?; let q_b = ReplicatedLinear::load_b( @@ -149,7 +156,7 @@ impl MlaAttention { kv_lora_rank, mla_cfg.rms_norm_eps, vb.pp("kv_a_layernorm"), - dtype, + norm_dtype, false, )?; @@ -238,6 +245,7 @@ impl MlaAttention { sm_scale, rope_scale, rope_theta: config.rope_theta as f32, + promote_qk_to_f32: is_qvar_builder || config.higher_precision_required(), dtype, }) } @@ -292,6 +300,15 @@ impl MlaAttention { let k_pe = k_pe_raw.reshape((seq_len, 1, self.qk_rope_head_dim))?; let q_pe_for_rope = q_pe.contiguous()?; + + let (q_pe_for_rope, k_pe) = if self.promote_qk_to_f32 { + ( + q_pe_for_rope.to_dtype(DType::F32)?, + k_pe.to_dtype(DType::F32)?, + ) + } else { + (q_pe_for_rope, k_pe) + }; let (q_pe, k_pe) = if let Some(rotary_emb) = &rotary_emb { let (q_new, k_new) = rotary_emb.apply_rotary_emb(&q_pe_for_rope, &k_pe, positions)?; (q_new, k_new) diff --git a/src/openai/models/layers/moe.rs b/src/openai/models/layers/moe.rs index c6205199..7967207a 100644 --- a/src/openai/models/layers/moe.rs +++ b/src/openai/models/layers/moe.rs @@ -6,6 +6,7 @@ use crate::openai::models::{Config, MoEConfig, QuantConfig, QwenMoEConfig}; use attention_rs::moe; use attention_rs::moe::moe_gemm_fp8; use attention_rs::silu_and_mul::silu_and_mul; +use attention_rs::sort::ArgSortOp; use candle::{DType, Module, Result, Tensor, D}; use candle_core as candle; use candle_core::quantized::GgmlDType; @@ -31,6 +32,26 @@ fn gated_activation(gate_up: &Tensor, half_dim: usize, act: &Activation) -> Resu } } +fn sort_expert_assignments(topk_ids: &Tensor, is_prefill: bool) -> Result<(Tensor, Tensor)> { + let flat = topk_ids.flatten_all()?; + if is_prefill { + flat.sort(true) + } else { + flat.sort_last_dim(true) + } +} + +fn presorted_expert_assignments( + topk_ids: &Tensor, + is_prefill: bool, +) -> Result> { + if !is_prefill { + return Ok(None); + } + let (expert_ids, sorted_token_ids) = sort_expert_assignments(topk_ids, true)?; + Ok(Some((sorted_token_ids, expert_ids))) +} + #[derive(Clone, Copy, Debug)] enum PackedGateUpLayout { // [experts, hidden, 2*intermediate] @@ -414,17 +435,7 @@ impl FusedMoe { ) -> Result { let (num_tokens, hidden_dim) = xs.dims2()?; - let (expert_ids, sorted_token_ids) = if is_prefill { - #[cfg(feature = "cuda")] - { - use attention_rs::sort::ArgSortOp; - topk_ids.flatten_all()?.sort(true)? - } - #[cfg(not(feature = "cuda"))] - topk_ids.flatten_all()?.sort_last_dim(true)? - } else { - topk_ids.flatten_all()?.sort_last_dim(true)? - }; + let (expert_ids, sorted_token_ids) = sort_expert_assignments(&topk_ids, is_prefill)?; let gate_up = moe::moe_gemm( &xs, @@ -724,17 +735,7 @@ impl FusedMoeISQ { xs.to_owned() }; - let (expert_ids, sorted_token_ids) = if is_prefill { - #[cfg(feature = "cuda")] - { - use attention_rs::sort::ArgSortOp; - topk_ids.flatten_all()?.sort(true)? - } - #[cfg(not(feature = "cuda"))] - topk_ids.flatten_all()?.sort_last_dim(true)? - } else { - topk_ids.flatten_all()?.sort_last_dim(true)? - }; + let (expert_ids, sorted_token_ids) = sort_expert_assignments(&topk_ids, is_prefill)?; let ys = { let gate = moe::moe_gemm_gguf( @@ -1267,17 +1268,7 @@ impl FusedMoeFp8 { ) -> Result { let (num_tokens, hidden_dim) = xs.dims2()?; - let (expert_ids, sorted_token_ids) = if is_prefill { - #[cfg(feature = "cuda")] - { - use attention_rs::sort::ArgSortOp; - topk_ids.flatten_all()?.sort(true)? - } - #[cfg(not(feature = "cuda"))] - topk_ids.flatten_all()?.sort_last_dim(true)? - } else { - topk_ids.flatten_all()?.sort_last_dim(true)? - }; + let (expert_ids, sorted_token_ids) = sort_expert_assignments(&topk_ids, is_prefill)?; let xs = if xs.dtype() == DType::F32 { xs.to_dtype(DType::BF16)? @@ -2177,14 +2168,7 @@ impl FusedMoeNvfp4 { xs }; - let pre_sorted = if is_prefill { - use attention_rs::sort::ArgSortOp; - let flat = topk_ids.flatten_all()?.contiguous()?; - let (eids, tids) = flat.sort(true)?; - Some((tids, eids)) - } else { - None - }; + let pre_sorted = presorted_expert_assignments(&topk_ids, is_prefill)?; let pre_sorted_refs = pre_sorted.as_ref().map(|(a, b)| (a, b)); let gate_up = moe::moe_gemm_nvfp4( diff --git a/src/openai/models/llama4/mod.rs b/src/openai/models/llama4/mod.rs index c31868ff..71b1f91f 100644 --- a/src/openai/models/llama4/mod.rs +++ b/src/openai/models/llama4/mod.rs @@ -187,18 +187,18 @@ impl LLama4DecoderLayer { None }; - let mut self_attn = Attention::new( + let qk_l2_norm = text_cfg.use_qk_norm && use_rope; + let self_attn = Attention::new_with_option( rotary_emb.clone(), config, vb.pp("self_attn"), comm.clone(), sliding_window, + false, + qk_l2_norm, + None, )?; - if text_cfg.use_qk_norm && use_rope { - self_attn.set_qk_l2_norm(true); - } - let moe_layers = text_cfg.moe_layers(); let is_moe_layer = moe_layers.contains(&layer_idx); @@ -220,14 +220,14 @@ impl LLama4DecoderLayer { config.hidden_size, config.rms_norm_eps, vb.pp("input_layernorm"), - dtype, + DType::F32, false, )?; let post_attention_layernorm = rms_norm( config.hidden_size, config.rms_norm_eps, vb.pp("post_attention_layernorm"), - dtype, + DType::F32, false, )?; diff --git a/src/openai/models/mod.rs b/src/openai/models/mod.rs index 4a1c5daa..39a26928 100644 --- a/src/openai/models/mod.rs +++ b/src/openai/models/mod.rs @@ -537,6 +537,14 @@ impl Config { None } + pub fn higher_precision_required(&self) -> bool { + self.isq_quant.is_some() + || self + .quantization_config + .as_ref() + .is_some_and(|cfg| matches!(cfg.quant_method.as_str(), "mxfp4" | "nvfp4")) + } + pub fn effective_max_seq_len(&self) -> usize { let base_max_position_embeddings = self.base_max_position_embeddings(); let Some(rope_scaling) = &self.rope_scaling else { diff --git a/src/openai/models/quantized_qwen3_5.rs b/src/openai/models/quantized_qwen3_5.rs index 7b274b82..46c1b024 100644 --- a/src/openai/models/quantized_qwen3_5.rs +++ b/src/openai/models/quantized_qwen3_5.rs @@ -404,7 +404,7 @@ impl QuantizedGatedDeltaNet { .as_ref() .expect("cu_seqlens_q must be present in prefill!"); let global_state = mamba_cache.recurrent_state_mut(self.gdn_layer_idx); - xs.device().synchronize()?; + // xs.device().synchronize()?; gdn::gated_delta_rule_recurrence_varlen( &q_scaled, &k, diff --git a/src/openai/openai_server.rs b/src/openai/openai_server.rs index 95d47075..66c54d17 100644 --- a/src/openai/openai_server.rs +++ b/src/openai/openai_server.rs @@ -10,10 +10,12 @@ use super::streaming::{ChatResponse, Streamer, StreamingStatus}; use super::OpenAIServerData; use crate::openai::multimodal::{build_messages_and_images, ImageData}; use crate::openai::{resolve_tools_for_request, ResolvedToolConfig}; +use crate::tools::helpers::{ + build_invalid_tool_call_feedback, build_tool_schema_map, filter_tool_calls, +}; use crate::tools::stream_parser::{ detect_prefilled_reasoning_end_marker, extract_reasoning_content, }; -use crate::tools::ToolFormat; use axum::response::sse::KeepAlive; use axum::{ extract::{Json, State}, @@ -103,31 +105,6 @@ async fn get_gen_prompt( } } - if !tool_config.tools.is_empty() { - let mut tools_prompt = ToolFormat::get_tool_prompt( - &pipeline.0.tool_config, - &pipeline.0.tool_model_type, - &pipeline.0.tool_parser_model_id, - pipeline.0.enforce_parser.as_deref(), - ); - - // Enforce tool_choice=function by prepending a mandatory instruction - if let crate::openai::ToolChoiceKind::Function(name) = &tool_config.choice { - tools_prompt = format!( - "IMPORTANT: You MUST call the tool \"{}\". Do not respond with plain text.\n\n{}", - name, tools_prompt - ); - } - - let current_system = conversation.get_system_message().unwrap_or_default(); - let new_system = if current_system.is_empty() { - tools_prompt - } else { - format!("{}\n\n{}", current_system, tools_prompt) - }; - conversation.set_system_message(Some(new_system)); - } - let enable_thinking = request.thinking.unwrap_or(true); let prompt = conversation.get_prompt(enable_thinking, &tool_config.tools); @@ -335,6 +312,8 @@ pub async fn chat_completions( } else { Some(Arc::clone(&sync_notify)) }; + let request_tools_for_engine = tool_config.tools.clone(); + let response_tools = tool_config.tools.clone(); let _ = tokio::task::spawn_blocking(move || { tokio::runtime::Handle::current().block_on(async move { @@ -349,7 +328,7 @@ pub async fn chat_completions( false, EncodingFormat::default(), EmbeddingType::default(), - tool_config.tools.clone(), + request_tools_for_engine.clone(), image_data, if stream_request { Some(Arc::new(response_tx)) @@ -430,18 +409,40 @@ pub async fn chat_completions( if has_tools { let parser = crate::tools::parser::ToolParser::new(); + let tool_schemas = build_tool_schema_map(&response_tools); for choice in &mut final_choices { - if choice.message.tool_calls.is_some() { + let parsed_calls = if let Some(calls) = choice.message.tool_calls.take() { + calls + } else if let Some(content) = &choice.message.content { + parser.parse(content) + } else { + Vec::new() + }; + + if parsed_calls.is_empty() { continue; } - if let Some(content) = &choice.message.content { - let calls = parser.parse(content); - if !calls.is_empty() { - choice.message.tool_calls = Some(calls); - choice.message.content = None; - choice.finish_reason = Some("tool_calls".to_string()); + + let (valid_calls, invalid_calls) = filter_tool_calls(&parsed_calls, &tool_schemas); + if !invalid_calls.is_empty() { + tracing::warn!( + "Dropped {} invalid tool call(s) before response", + invalid_calls.len() + ); + } + if valid_calls.is_empty() { + if let Some(feedback) = + build_invalid_tool_call_feedback(&invalid_calls, &tool_schemas, None) + { + choice.message.content = Some(feedback); } + choice.finish_reason = Some("stop".to_string()); + continue; } + + choice.message.tool_calls = Some(valid_calls); + choice.message.content = None; + choice.finish_reason = Some("tool_calls".to_string()); } } diff --git a/src/openai/pipelines/inputs.rs b/src/openai/pipelines/inputs.rs index 8ee9b7d4..3a59a779 100644 --- a/src/openai/pipelines/inputs.rs +++ b/src/openai/pipelines/inputs.rs @@ -9,6 +9,14 @@ use super::{LLMEngine, PreparedInputs, Sequence, SequenceGroup, PREFILL_CHUNK_SI use crate::InputMetadata; impl LLMEngine { + fn used_blocks_for_len(seq_len: usize, block_size: usize, table_len: usize) -> usize { + if seq_len == 0 { + 0 + } else { + seq_len.div_ceil(block_size).min(table_len) + } + } + pub fn prepare_block_tables( &self, groups: &VecDeque>, @@ -250,7 +258,7 @@ impl LLMEngine { let mamba_slot_mapping = self.prepare_mamba_slot_mapping(&sequence_ids, true, rank, device)?; #[cfg(feature = "flashinfer")] - let flashinfer_metadata = if self.flashinfer_kv_params_for_rank(rank)?.is_some() { + let flashinfer_metadata = if let Some(params) = self.flashinfer_kv_params_for_rank(rank)? { let mut indptr = vec![0u32]; let mut indices = Vec::new(); let mut last_len = Vec::new(); @@ -313,47 +321,37 @@ impl LLMEngine { kv_len_arr_host.push(full + last_len_host[i]); } } - let num_qo_heads = self.config.num_attention_heads / self.num_shards; - let num_kv_heads = self - .config - .num_key_value_heads - .unwrap_or(self.config.num_attention_heads) - / self.num_shards; - let head_dim = self.config.head_dim.unwrap_or(128); - let page_size = self.cache_config.block_size; let total_num_rows = *cu_seqlens_q_vec.last().unwrap(); let mut prefill_plan_info = None; let mut mla_prefill_plan_info = None; if self.config.is_mla() { - mla_prefill_plan_info = attention_rs::mla::mla_prefill_plan( + mla_prefill_plan_info = Some(attention_rs::mla::mla_prefill_plan( device, &cu_seqlens_q_vec, &indptr_host, &kv_len_arr_host, last_len_host.len(), - num_qo_heads, - head_dim, + params.num_qo_heads, + params.head_dim, true, - ) - .ok(); + )?); } else { - prefill_plan_info = attention_rs::flashinfer::prefill_plan( + prefill_plan_info = Some(attention_rs::flashinfer::prefill_plan( device, &cu_seqlens_q_vec, &indptr_host, &kv_len_arr_host, total_num_rows, last_len_host.len(), - num_qo_heads, - num_kv_heads, - head_dim, - page_size, - self.cache_config.dtype, + params.num_qo_heads, + params.num_kv_heads, + params.head_dim, + params.page_size, + params.out_dtype, None, - ) - .ok(); + )?); } Some(FlashInferMetadata { @@ -456,6 +454,13 @@ impl LLMEngine { let slot: i64 = slot.try_into().unwrap(); slot_mapping.push(slot); + let used_blocks = Self::used_blocks_for_len( + seq.deref().get_len(), + self.cache_config.block_size, + table.len(), + ); + let table = table.get(..used_blocks).unwrap_or(&[]).to_vec(); + if let Some(sliding_window) = self.config.sliding_window { let sliding_window_blocks = sliding_window / self.cache_config.block_size; let slide_idx = if table.len() > sliding_window_blocks { @@ -498,7 +503,16 @@ impl LLMEngine { let (pipeline, _) = self.get_pipeline(rank).ok_or_else(|| { candle_core::Error::msg(format!("missing pipeline for rank {rank}")) })?; - pipeline.capturer.is_exact_captured(length) + // Match vllm.rs: only mamba models need exact batch size graphs. + // Non-mamba models (including MLA like GLM4) can use any captured + // graph >= the current batch size, keeping the use_cuda_graph flag + // consistent with the graph replay decision in forward(). + let require_exact_graph = mamba_slot_mapping.is_some(); + if require_exact_graph { + pipeline.capturer.is_exact_captured(length) + } else { + pipeline.capturer.is_captured(length) + } }; #[cfg(not(all(feature = "cuda", feature = "graph")))] let use_cuda_graph = false; @@ -517,9 +531,11 @@ impl LLMEngine { .iter() .map(|block| block.deref_mut().block_id as u32) .collect::>(); - indices.extend(table); - indptr.push(indices.len() as u32); let len = seq.deref().get_len(); + let used_blocks = + Self::used_blocks_for_len(len, self.cache_config.block_size, table.len()); + indices.extend(table.iter().take(used_blocks).copied()); + indptr.push(indices.len() as u32); let last = if len == 0 { 0 } else { diff --git a/src/openai/pipelines/pipeline.rs b/src/openai/pipelines/pipeline.rs index ea8f5284..2fad8cfc 100644 --- a/src/openai/pipelines/pipeline.rs +++ b/src/openai/pipelines/pipeline.rs @@ -107,7 +107,8 @@ pub enum LLMModel { fn tool_model_type_for(model: &LLMModel) -> ToolModelType { match model { - LLMModel::Llama(_) | LLMModel::LLaMa4(_) | LLMModel::LlamaGGUF(_) => ToolModelType::LLaMa, + LLMModel::Llama(_) | LLMModel::LlamaGGUF(_) => ToolModelType::LLaMa, + LLMModel::LLaMa4(_) => ToolModelType::LLaMa4, LLMModel::Qwen(_) | LLMModel::Qwen3_5(_) | LLMModel::Qwen3VL(_) @@ -118,7 +119,8 @@ fn tool_model_type_for(model: &LLMModel) -> ToolModelType { | LLMModel::QWenGGUFMoE(_) | LLMModel::QWen3_5GGUFMoE(_) => ToolModelType::Qwen3MoE, LLMModel::Gemma(_) => ToolModelType::Gemma, - LLMModel::Gemma3(_) | LLMModel::Gemma3VL(_) | LLMModel::Gemma4(_) => ToolModelType::Gemma3, + LLMModel::Gemma3(_) | LLMModel::Gemma3VL(_) => ToolModelType::Gemma3, + LLMModel::Gemma4(_) => ToolModelType::Gemma4, LLMModel::MiniMax(_) => ToolModelType::MiniMax, LLMModel::Mistral(_) | LLMModel::Mistral3VL(_) => ToolModelType::Mistral, LLMModel::Yi(_) => ToolModelType::Yi, @@ -1775,8 +1777,7 @@ impl DefaultPipeline { #[cfg(all(feature = "cuda", feature = "graph"))] if !input_metadata.is_prefill { let input_batch = input_tokens.dim(0)?; - let require_exact_graph = input_metadata.mamba_slot_mapping.is_some() - || input_metadata.flashinfer_metadata.is_some(); + let require_exact_graph = input_metadata.mamba_slot_mapping.is_some(); let can_replay = if require_exact_graph { self.capturer.is_exact_captured(input_batch) } else { @@ -2180,7 +2181,9 @@ impl DefaultPipeline { seq.deref_mut().deref_mut().pending_finish_logprobs = Some(finish_logprobs); return Right("tool_calls".to_string()); } - if self.json_end_token_id == Some(next_token) { + if self.tool_call_end_token_ids.is_empty() + && self.json_end_token_id == Some(next_token) + { let mut output_tokens: Vec = seq .deref() .get_output_tokens() diff --git a/src/openai/sampling_params.rs b/src/openai/sampling_params.rs index 26f45eea..741d4565 100644 --- a/src/openai/sampling_params.rs +++ b/src/openai/sampling_params.rs @@ -256,12 +256,6 @@ impl SamplingParams { )); } - if self.top_k.is_some_and(|k| k != -1) { - return Err(APIError::new_str( - "top_k must be -1 when using greedy sampling (no temperature specified).", - )); - } - Ok(()) } } diff --git a/src/tools/helpers.rs b/src/tools/helpers.rs index 06614b22..22bcc399 100644 --- a/src/tools/helpers.rs +++ b/src/tools/helpers.rs @@ -8,16 +8,104 @@ use super::{FunctionCall, Tool, ToolCall}; use regex::Regex; use serde_json::Value; use std::collections::HashMap; +use std::env; use std::sync::OnceLock; +static STRICT_TOOL_CALL_VALIDATION: OnceLock = OnceLock::new(); + +/// Returns whether strict server-side tool schema validation is enabled. +/// When disabled, parsed tool calls are still checked for known tool names before +/// being returned to clients, but argument-schema failures are allowed through. +pub fn strict_tool_call_validation_enabled() -> bool { + *STRICT_TOOL_CALL_VALIDATION.get_or_init(|| { + env::var("VLLM_RS_STRICT_TOOL_CALL") + .ok() + .map(|raw| { + let normalized = raw.trim().to_ascii_lowercase(); + matches!(normalized.as_str(), "1" | "true" | "yes" | "on") + }) + .unwrap_or(false) + }) +} + /// Build a map of tool names to their parameter schemas pub fn build_tool_schema_map(tools: &[Tool]) -> HashMap { tools .iter() - .map(|tool| (tool.function.name.clone(), tool.function.parameters.clone())) + .map(|tool| { + let mut schema = tool.function.parameters.clone(); + if tool.function.strict == Some(false) { + if let Some(obj) = schema.as_object_mut() { + obj.insert("x-vllm-rs-lenient".to_string(), Value::Bool(true)); + } + } + (tool.function.name.clone(), schema) + }) .collect() } +/// Enforce `tool_choice=function` by retaining only calls that match `forced_tool_name`. +/// Returns the number of dropped calls. +pub fn retain_tool_calls_forced_name( + tool_calls: &mut Vec, + forced_tool_name: Option<&str>, +) -> usize { + let Some(forced_name) = forced_tool_name else { + return 0; + }; + + let before = tool_calls.len(); + tool_calls.retain(|call| call.function.name == forced_name); + before - tool_calls.len() +} + +/// Build a model-facing fallback message when tool calls were parsed but rejected. +pub fn build_invalid_tool_call_feedback( + invalid_calls: &[ToolCall], + schemas: &HashMap, + forced_tool_name: Option<&str>, +) -> Option { + if invalid_calls.is_empty() { + return None; + } + + let mut rejected_tools: Vec = invalid_calls + .iter() + .map(|call| call.function.name.trim()) + .filter(|name| !name.is_empty()) + .map(ToOwned::to_owned) + .collect(); + rejected_tools.sort(); + rejected_tools.dedup(); + + let mut allowed_tools: Vec = schemas.keys().cloned().collect(); + allowed_tools.sort(); + + let rejected_summary = if rejected_tools.is_empty() { + "Rejected tool call(s).".to_string() + } else { + format!("Rejected tool call(s): {}.", rejected_tools.join(", ")) + }; + + let mut parts = vec![rejected_summary]; + if let Some(name) = forced_tool_name { + if !name.trim().is_empty() { + parts.push(format!("Required tool_choice is '{}'.", name)); + } + } + if allowed_tools.is_empty() { + parts.push("No callable tools are available for this turn.".to_string()); + } else { + parts.push(format!("Allowed tools: {}.", allowed_tools.join(", "))); + } + parts.push( + "Retry with one valid tool call using a JSON object that matches the tool schema." + .to_string(), + ); + + Some(parts.join(" ")) +} + /// Filter tool calls into valid and invalid based on schema validation. /// /// Valid calls have their arguments parsed, repaired, normalized, and coerced. diff --git a/src/tools/mod.rs b/src/tools/mod.rs index 28e34c68..766607a9 100644 --- a/src/tools/mod.rs +++ b/src/tools/mod.rs @@ -279,113 +279,9 @@ impl ToolResult { } } -/// Format tool definitions for injection into the prompt -#[derive(Debug, Clone)] -pub struct ToolFormat {} - -impl ToolFormat { - fn parser_name_for_prompt( - tool_config: &crate::tools::stream_parser::ToolConfig, - model_type: &crate::tools::stream_parser::ToolModelType, - model_id: &str, - enforce_parser: Option<&str>, - ) -> &'static str { - if let Some(name) = enforce_parser.map(str::trim).filter(|s| !s.is_empty()) { - return match name { - "qwen_coder" => "qwen_coder", - "json" => "json", - "qwen" => "qwen", - "mistral" => "mistral", - "llama" => "llama", - "glm47_moe" => "glm47_moe", - "deepseek" => "deepseek", - _ => { - if tool_config.start_token_str.contains("tool_call") { - "qwen" - } else { - "json" - } - } - }; - } - - let model_lower = model_id.to_ascii_lowercase(); - match model_type { - crate::tools::stream_parser::ToolModelType::LLaMa => "llama", - crate::tools::stream_parser::ToolModelType::Mistral => "mistral", - crate::tools::stream_parser::ToolModelType::Qwen - | crate::tools::stream_parser::ToolModelType::Qwen3MoE => { - if model_lower.contains("coder") || model_lower.contains("qwen3.5") { - "qwen_coder" - } else { - "qwen" - } - } - crate::tools::stream_parser::ToolModelType::Gemma - | crate::tools::stream_parser::ToolModelType::Gemma3 - | crate::tools::stream_parser::ToolModelType::GLM4 => "json", - crate::tools::stream_parser::ToolModelType::Phi - | crate::tools::stream_parser::ToolModelType::Phi4 - | crate::tools::stream_parser::ToolModelType::Yi - | crate::tools::stream_parser::ToolModelType::StableLM => "qwen", - crate::tools::stream_parser::ToolModelType::DeepSeek => "deepseek", - crate::tools::stream_parser::ToolModelType::MiniMax => "minimax_m2", - } - } - - /// Get tool prompt for a specific tool config (model-aware tags). - /// Tool definitions are injected by the chat template — this only provides usage instructions. - pub fn get_tool_prompt( - tool_config: &crate::tools::stream_parser::ToolConfig, - model_type: &crate::tools::stream_parser::ToolModelType, - model_id: &str, - enforce_parser: Option<&str>, - ) -> String { - let start_tag = &tool_config.start_token_str; - let end_tag = &tool_config.end_token_str; - match Self::parser_name_for_prompt(tool_config, model_type, model_id, enforce_parser) { - "qwen_coder" => format!( - "MOST IMPORTANT INSTRUCTION, **MUST** FOLLOW:\n\ - For each function call, you MUST wrap the function block in {start_tag}{end_tag} tags.\n\n\ - Do NOT USE ANY code blocks. Required format:\n\ - {start_tag}\n\ - >\n\ - >value\n\ - >value\n\ - \n\ - {end_tag}\n\n\ - Rules:\n\ - - Use XML-style and tags inside {start_tag}{end_tag}\n\ - - Do NOT emit JSON for tool calls\n\ - - Required parameters MUST be provided using their own blocks\n\ - - Tool-use must be placed **at the end** of your response (AFTER reasoning), **top-level**, and not nested within other tags.\n\ - - Do NOT USE ANY code blocks\n\ - - MUST FOLLOW the above instruction when using tool call!" - ), - _ => format!( - "MOST IMPORTANT INSTRUCTION, **MUST** FOLLOW:\n\ - For each function call, you MUST wrap function name and arguments in {start_tag}{end_tag} tags.\n\n\ - Do NOT USE ANY code blocks. Required format:\n\ - {start_tag}\n\ - {{\"name\": \"\", \"arguments\": }}\n\ - {end_tag}\n\n\ - Rules:\n\ - - Wrap function name and arguments with {start_tag} and {end_tag} tags\n\ - - Always use the exact {start_tag}{end_tag} format shown above\n\ - - Do NOT USE ANY code blocks\n\ - - Tool-use must be placed **at the end** of your response (AFTER reasoning), **top-level**, and not nested within other tags.\n\ - - Always adhere to this format for the tool use to ensure proper parsing and execution.\n\ - - The \"name\" and \"arguments\" are necessary fields\n\ - - MUST FOLLOW the above instruction when using tool call!" - ), - } - } -} - #[cfg(test)] mod tests { use super::*; - use crate::tools::stream_parser::{ToolConfig, ToolModelType}; #[test] fn tool_choice_deserializes_string_modes() { @@ -423,31 +319,4 @@ mod tests { assert!(id.starts_with("call_")); assert_eq!(id.len(), 5 + 16); // "call_" + 16 hex chars } - - #[test] - fn tool_prompt_uses_xml_for_qwen_coder_models() { - let prompt = ToolFormat::get_tool_prompt( - &ToolConfig::for_model_type(&ToolModelType::Qwen), - &ToolModelType::Qwen, - "qwen3-coder", - None, - ); - assert!(prompt.contains(">")); - assert!(prompt.contains(">value")); - assert!(prompt.contains("Do NOT emit JSON for tool calls")); - } - - #[test] - fn tool_prompt_uses_json_for_regular_qwen_models() { - let prompt = ToolFormat::get_tool_prompt( - &ToolConfig::for_model_type(&ToolModelType::Qwen), - &ToolModelType::Qwen, - "qwen3-instruct", - None, - ); - assert!( - prompt.contains("{\"name\": \"\", \"arguments\": }") - ); - assert!(!prompt.contains("Do NOT emit JSON for tool calls")); - } } diff --git a/src/tools/stream_parser.rs b/src/tools/stream_parser.rs index 8011eec6..f2a3f1f7 100644 --- a/src/tools/stream_parser.rs +++ b/src/tools/stream_parser.rs @@ -11,6 +11,119 @@ use tool_parser::{ ParserFactory, ToolParser as ExternalToolParser, }; +/// Look up the JSON schema types for a parameter. +/// Supports direct `type`, compound schemas, and enum values. +fn extract_schema_types(schema: &Value) -> Vec { + let Some(obj) = schema.as_object() else { + return vec!["string".to_string()]; + }; + + let mut types = Vec::new(); + if let Some(t) = obj.get("type") { + match t { + Value::String(s) => types.push(s.clone()), + Value::Array(arr) => { + types.extend( + arr.iter() + .filter_map(|item| item.as_str().map(str::to_string)), + ); + } + _ => {} + } + } + + for key in ["anyOf", "oneOf", "allOf"] { + if let Some(Value::Array(choices)) = obj.get(key) { + for choice in choices { + types.extend(extract_schema_types(choice)); + } + } + } + + if let Some(Value::Array(enum_vals)) = obj.get("enum") { + for val in enum_vals { + let ty = match val { + Value::Null => "null", + Value::Bool(_) => "boolean", + Value::Number(n) if n.is_i64() || n.is_u64() => "integer", + Value::Number(_) => "number", + Value::String(_) => "string", + Value::Array(_) => "array", + Value::Object(_) => "object", + }; + types.push(ty.to_string()); + } + } + + if types.is_empty() { + types.push("string".to_string()); + } + types.sort(); + types.dedup(); + types +} + +fn coerce_param_value(raw: &str, schema_types: &[String]) -> Value { + let raw = raw.trim(); + let lower = raw.to_ascii_lowercase(); + if matches!(lower.as_str(), "null" | "none" | "nil") { + return Value::Null; + } + + let has_explicit_non_string = schema_types + .iter() + .any(|t| !matches!(t.as_str(), "string" | "str" | "text")); + if has_explicit_non_string { + for ptype in ["integer", "number", "boolean", "object", "array", "string"] { + if !schema_types.iter().any(|t| t == ptype) { + continue; + } + match ptype { + "integer" => { + if let Ok(n) = raw.parse::() { + return Value::Number(n.into()); + } + } + "number" => { + if let Ok(f) = raw.parse::() { + if f == (f as i64) as f64 { + return Value::Number((f as i64).into()); + } + if let Some(n) = serde_json::Number::from_f64(f) { + return Value::Number(n); + } + } + } + "boolean" => match lower.as_str() { + "true" | "1" | "yes" | "on" => return Value::Bool(true), + "false" | "0" | "no" | "off" => return Value::Bool(false), + _ => {} + }, + "object" | "array" => { + if let Ok(v) = serde_json::from_str::(raw) { + return v; + } + } + "string" => return Value::String(raw.to_string()), + _ => {} + } + } + } + + serde_json::from_str::(raw).unwrap_or_else(|_| Value::String(raw.to_string())) +} + +fn resolve_param_properties<'a>( + function_name: &str, + tools: &'a [openai_protocol::common::Tool], +) -> Option<&'a serde_json::Map> { + tools + .iter() + .find(|tool| tool.function.name == function_name) + .and_then(|tool| tool.function.parameters.get("properties")) + .and_then(Value::as_object) +} + /// Convert our local Tool to openai_protocol::Tool for the tool-parser crate. fn to_openai_tools(tools: &[crate::tools::Tool]) -> Vec { tools @@ -29,7 +142,10 @@ fn to_openai_tools(tools: &[crate::tools::Tool]) -> Vec...` -fn parse_minimax_xml_tool_calls(text: &str) -> Vec { +fn parse_minimax_xml_tool_calls( + text: &str, + tools: &[openai_protocol::common::Tool], +) -> Vec { let mut calls = Vec::new(); let mut search_from = 0; @@ -64,6 +180,7 @@ fn parse_minimax_xml_tool_calls(text: &str) -> Vec { }; let invoke_block = &text[abs_invoke_start..invoke_end]; + let param_props = resolve_param_properties(function_name, tools); // Extract parameters from ... let mut args = Map::new(); @@ -107,12 +224,11 @@ fn parse_minimax_xml_tool_calls(text: &str) -> Vec { .unwrap_or(value_section.len()); let param_value = value_section[..value_end].trim(); - // Try to parse value as JSON, otherwise use as string - let json_value = if let Ok(parsed) = serde_json::from_str::(param_value) { - parsed - } else { - Value::String(param_value.to_string()) - }; + let schema_types = param_props + .and_then(|props| props.get(param_name)) + .map(extract_schema_types) + .unwrap_or_else(|| vec!["string".to_string()]); + let json_value = coerce_param_value(param_value, &schema_types); args.insert(param_name.to_string(), json_value); param_search = abs_param_start + value_start_rel + value_end; @@ -138,11 +254,13 @@ fn parse_minimax_xml_tool_calls(text: &str) -> Vec { #[derive(Clone, Debug, PartialEq)] pub enum ToolModelType { LLaMa, + LLaMa4, Qwen, Qwen3MoE, Mistral, Gemma, Gemma3, + Gemma4, Phi, Phi4, GLM4, @@ -196,6 +314,17 @@ impl ToolConfig { end_token_str: "<|eom_id|>".into(), } } + ToolModelType::LLaMa4 => { + start_ids.insert(200016); + end_ids.insert(200007); + end_ids.insert(200008); + Self { + start_token_ids: start_ids, + end_token_ids: end_ids, + start_token_str: "<|python_start|>".into(), + end_token_str: "<|eom|>".into(), + } + } ToolModelType::Qwen | ToolModelType::Qwen3MoE => { start_ids.insert(151657); end_ids.insert(151658); @@ -221,6 +350,16 @@ impl ToolConfig { start_token_str: "".into(), end_token_str: "".into(), }, + ToolModelType::Gemma4 => { + start_ids.insert(48); + end_ids.insert(49); + Self { + start_token_ids: start_ids, + end_token_ids: end_ids, + start_token_str: "<|tool_call>".into(), + end_token_str: "".into(), + } + } ToolModelType::MiniMax => { // MiniMax tokenizer ships dedicated tool envelope tokens: // 200052 => @@ -254,41 +393,72 @@ impl ToolConfig { } pub fn validate_with_tokenizer(&mut self, tokenizer: &Tokenizer, model_type: &ToolModelType) { - if self.has_start_tokens() - && !Self::matches_single_token(tokenizer, &self.start_token_str, &self.start_token_ids) - { - if Self::try_rebind_single_token_id( + if self.has_start_tokens() { + if !Self::matches_single_token(tokenizer, &self.start_token_str, &self.start_token_ids) + { + if Self::try_rebind_single_token_id( + tokenizer, + &self.start_token_str, + &mut self.start_token_ids, + ) { + tracing::warn!( + "Tool start token IDs corrected for model {:?}: {:?}", + model_type, + self.start_token_ids + ); + } else { + tracing::warn!( + "Tool start token IDs not supported for model {:?}, falling back to text matching", + model_type + ); + self.start_token_ids.clear(); + } + } + } else if !self.start_token_str.is_empty() + && Self::try_rebind_single_token_id( tokenizer, &self.start_token_str, &mut self.start_token_ids, - ) { - tracing::warn!( - "Tool start token IDs corrected for model {:?}: {:?}", - model_type, - self.start_token_ids - ); - } else { - tracing::warn!("Tool start token IDs not supported for model {:?}, falling back to text matching", model_type); - self.start_token_ids.clear(); - } - } - if self.has_end_tokens() - && !Self::matches_single_token(tokenizer, &self.end_token_str, &self.end_token_ids) + ) { - if Self::try_rebind_single_token_id( + tracing::info!( + "Tool start token IDs auto-populated from tokenizer for model {:?}: {:?}", + model_type, + self.start_token_ids + ); + } + if self.has_end_tokens() { + if !Self::matches_single_token(tokenizer, &self.end_token_str, &self.end_token_ids) { + if Self::try_rebind_single_token_id( + tokenizer, + &self.end_token_str, + &mut self.end_token_ids, + ) { + tracing::warn!( + "Tool end token IDs corrected for model {:?}: {:?}", + model_type, + self.end_token_ids + ); + } else { + tracing::warn!( + "Tool end token IDs not supported for model {:?}, falling back to text matching", + model_type + ); + self.end_token_ids.clear(); + } + } + } else if !self.end_token_str.is_empty() + && Self::try_rebind_single_token_id( tokenizer, &self.end_token_str, &mut self.end_token_ids, - ) { - tracing::warn!( - "Tool end token IDs corrected for model {:?}: {:?}", - model_type, - self.end_token_ids - ); - } else { - tracing::warn!("Tool end token IDs not supported for model {:?}, falling back to text matching", model_type); - self.end_token_ids.clear(); - } + ) + { + tracing::info!( + "Tool end token IDs auto-populated from tokenizer for model {:?}: {:?}", + model_type, + self.end_token_ids + ); } } @@ -400,6 +570,7 @@ const REASONING_MARKERS: &[(&str, &str)] = &[ ("<|think|>", "<|/think|>"), ("[THINK]", "[/THINK]"), ("", ""), + ("<|channel>", ""), ]; pub fn reasoning_markers() -> &'static [(&'static str, &'static str)] { @@ -499,6 +670,7 @@ pub struct StreamToolParser { config: ToolConfig, state: ParserState, buffer: String, + model_id: String, parse_strategy: String, parser: Box, tools: Vec, @@ -527,12 +699,14 @@ impl StreamToolParser { tools: Vec, enforce_parser: Option, ) -> Self { - let openai_tools = to_openai_tools(&tools); let parse_strategy = match model_type { ToolModelType::Mistral => "mistral_list", + ToolModelType::Gemma4 => "gemma4", + ToolModelType::LLaMa4 => "pythonic", _ => "json", } .to_string(); + let openai_tools = to_openai_tools(&tools); let factory = ParserFactory::new(); let parser_name = if let Some(name) = enforce_parser.as_ref().and_then(|s| { @@ -575,6 +749,7 @@ impl StreamToolParser { config, state: ParserState::Normal, buffer: String::new(), + model_id, parse_strategy, parser, tools: openai_tools, @@ -939,6 +1114,23 @@ impl StreamToolParser { } pub async fn parse_complete_with_fallback(&self, text: &str) -> Vec { + if self.parse_strategy == "gemma4" { + match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + Self::parse_gemma4_tool_calls(text) + })) { + Ok(Some(calls)) => return calls, + Ok(None) => {} + Err(e) => { + let msg = e + .downcast_ref::() + .map(|s| s.as_str()) + .or_else(|| e.downcast_ref::<&str>().copied()) + .unwrap_or("unknown"); + tracing::warn!("Gemma4 tool call parse panicked: {}", msg); + } + } + } + let mut parsed_calls = match self.parser.parse_complete(text).await { Ok((_normal_text, calls)) => calls, Err(err) => { @@ -947,6 +1139,15 @@ impl StreamToolParser { } }; + if parsed_calls.is_empty() && self.parse_strategy == "pythonic" { + let factory = ParserFactory::new(); + if let Some(pythonic_parser) = factory.registry().create_parser("pythonic") { + if let Ok((_normal_text, calls)) = pythonic_parser.parse_complete(text).await { + parsed_calls = calls; + } + } + } + if parsed_calls.is_empty() && text.contains("` style special tokens (e.g. Gemma4's <|tool_call>/) + // are not XML envelopes; skip structural validation for them. + if self.config.start_token_str.starts_with("<|") + || self.config.end_token_str.starts_with("<|") + { + return true; + } + let Some(start_idx) = self.buffer.find(&self.config.start_token_str) else { return true; }; @@ -1214,9 +1424,9 @@ impl StreamToolParser { let Some(invoke_end_rel) = invoke_section.rfind("") else { return false; }; - let invoke_block = - &block[invoke_start..invoke_start + invoke_end_rel + "".len()]; - return Self::has_balanced_minimax_parameter_tags(invoke_block); + let invoke_end = invoke_start + invoke_end_rel + "".len(); + let invoke_block = &block[invoke_start..invoke_end]; + return Self::has_balanced_parameter_tags(invoke_block, "") else { return false; }; - let func_block = &block[fs..fs + fe_rel + "".len()]; - return Self::has_balanced_parameter_tags(func_block); + let func_end = fs + fe_rel + "".len(); + let func_block = &block[fs..func_end]; + return Self::has_balanced_parameter_tags(func_block, "") || block.contains("") { + return block.contains("") + && Self::has_balanced_xml_tags(block, "", "") + && Self::has_balanced_xml_tags(block, "", ""); } if inner.is_empty() { return false; @@ -1236,55 +1452,18 @@ impl StreamToolParser { serde_json::from_str::(inner).is_ok() } - fn has_balanced_minimax_parameter_tags(invoke_block: &str) -> bool { + fn has_balanced_parameter_tags(function_block: &str, open_tag: &str) -> bool { let mut idx = 0usize; let mut open_count = 0usize; - const OPEN: &str = ""; - while idx < invoke_block.len() { - let open_pos = invoke_block[idx..].find(OPEN).map(|p| idx + p); - let close_pos = invoke_block[idx..].find(CLOSE).map(|p| idx + p); - match (open_pos, close_pos) { - (None, None) => break, - (Some(op), None) => { - open_count += 1; - idx = op + OPEN.len(); - } - (None, Some(cp)) => { - if open_count > 0 { - open_count -= 1; - } - idx = cp + CLOSE.len(); - } - (Some(op), Some(cp)) => { - if op < cp { - open_count += 1; - idx = op + OPEN.len(); - } else { - if open_count > 0 { - open_count -= 1; - } - idx = cp + CLOSE.len(); - } - } - } - } - open_count == 0 - } - - fn has_balanced_parameter_tags(function_block: &str) -> bool { - let mut idx = 0usize; - let mut open_count = 0usize; - const OPEN: &str = " break, (Some(op), None) => { open_count += 1; - idx = op + OPEN.len(); + idx = op + open_tag.len(); } (None, Some(cp)) => { if open_count > 0 { @@ -1295,7 +1474,7 @@ impl StreamToolParser { (Some(op), Some(cp)) => { if op < cp { open_count += 1; - idx = op + OPEN.len(); + idx = op + open_tag.len(); } else { if open_count > 0 { open_count -= 1; @@ -1308,6 +1487,12 @@ impl StreamToolParser { open_count == 0 } + fn has_balanced_xml_tags(block: &str, open: &str, close: &str) -> bool { + let open_count = block.matches(open).count(); + let close_count = block.matches(close).count(); + open_count > 0 && open_count == close_count + } + fn finalize_streamed_arguments(&self, raw: &str) -> String { let trimmed = raw.trim(); if trimmed.is_empty() { @@ -1351,23 +1536,368 @@ impl StreamToolParser { let model_lower = model_id.to_ascii_lowercase(); match model_type { ToolModelType::LLaMa => "llama", + ToolModelType::LLaMa4 => "pythonic", ToolModelType::Mistral => "mistral", ToolModelType::Qwen | ToolModelType::Qwen3MoE => { - if model_lower.contains("coder") || model_lower.contains("qwen3.5") { + if model_lower.contains("coder") + || model_lower.contains("qwen3.5") + || model_lower.contains("qwen3.6") + { "qwen_coder" } else { "qwen" } } - ToolModelType::Gemma | ToolModelType::Gemma3 => "json", + ToolModelType::Gemma | ToolModelType::Gemma3 | ToolModelType::Gemma4 => "json", ToolModelType::Phi | ToolModelType::Phi4 => "qwen", - ToolModelType::GLM4 => "json", + ToolModelType::GLM4 => "glm47_moe", ToolModelType::Yi | ToolModelType::StableLM => "qwen", ToolModelType::DeepSeek => "deepseek", ToolModelType::MiniMax => "minimax_m2", } } + /// Parse Gemma4 tool calls: `<|tool_call>call:NAME{key:<|"|>value<|"|>,...}`. + fn parse_gemma4_tool_calls(text: &str) -> Option> { + const PREFIX: &str = "<|tool_call>call:"; + const PREFIX_STRIPPED: &str = "call:"; + const SUFFIX: &str = ""; + + let text = text.trim_end(); + let text = text + .strip_suffix("<|tool_response>") + .or_else(|| text.strip_suffix("")) + .unwrap_or(text); + + let has_full_prefix = text.contains(PREFIX); + let has_stripped_prefix = !has_full_prefix && text.contains(PREFIX_STRIPPED); + if !has_full_prefix && !has_stripped_prefix { + return None; + } + let active_prefix = if has_full_prefix { + PREFIX + } else { + PREFIX_STRIPPED + }; + + let mut calls = Vec::new(); + let mut search_start = 0usize; + while let Some(rel_pos) = text[search_start..].find(active_prefix) { + let abs_start = search_start + rel_pos + active_prefix.len(); + let Some(brace_rel) = text[abs_start..].find('{') else { + break; + }; + let name = text[abs_start..abs_start + brace_rel].trim(); + let brace_abs = abs_start + brace_rel; + let Some((inner, after_brace)) = Self::gemma4_extract_braces(text, brace_abs) else { + break; + }; + let arguments = Self::gemma4_parse_args(inner); + calls.push(crate::tools::new_tool_call( + crate::tools::generate_tool_call_id(), + name.to_string(), + serde_json::to_string(&arguments).unwrap_or_else(|_| "{}".to_string()), + )); + + let remaining = &text[after_brace..]; + search_start = if let Some(suf_pos) = remaining.find(SUFFIX) { + after_brace + suf_pos + SUFFIX.len() + } else { + after_brace + }; + } + + (!calls.is_empty()).then_some(calls) + } + + fn gemma4_extract_braces(s: &str, start: usize) -> Option<(&str, usize)> { + const DELIM: &str = "<|\"|>"; + if !s.is_char_boundary(start) || s.as_bytes().get(start) != Some(&b'{') { + return None; + } + + let mut depth = 0usize; + let mut in_delim_string = false; + let mut in_regular_string = false; + let tail = &s[start..]; + let mut iter = tail.char_indices(); + while let Some((offset, ch)) = iter.next() { + let abs = start + offset; + if in_delim_string { + if tail[offset..].starts_with(DELIM) { + in_delim_string = false; + for _ in 0..DELIM.len().saturating_sub(ch.len_utf8()) { + iter.next(); + } + } + continue; + } + if in_regular_string { + if ch == '"' && (offset == 0 || tail.as_bytes()[offset - 1] != b'\\') { + in_regular_string = false; + } + continue; + } + if tail[offset..].starts_with(DELIM) { + in_delim_string = true; + for _ in 0..DELIM.len().saturating_sub(ch.len_utf8()) { + iter.next(); + } + continue; + } + match ch { + '"' => in_regular_string = true, + '{' => depth += 1, + '}' => { + depth = depth.saturating_sub(1); + if depth == 0 { + let inner_start = start + '{'.len_utf8(); + return Some((&s[inner_start..abs], abs + '}'.len_utf8())); + } + } + _ => {} + } + } + None + } + + fn gemma4_parse_args(args_str: &str) -> Value { + if args_str.trim().is_empty() { + return Value::Object(Map::new()); + } + + let cleaned = args_str.replace("<|\"|>", "\""); + if let Ok(v) = serde_json::from_str::(&format!("{{{cleaned}}}")) { + return v; + } + + let mut map = Map::new(); + let chars: Vec<(usize, char)> = args_str.char_indices().collect(); + let n = chars.len(); + let mut ci = 0usize; + while ci < n { + while ci < n && matches!(chars[ci].1, ' ' | ',' | '\n' | '\t') { + ci += 1; + } + if ci >= n { + break; + } + + let key_start = chars[ci].0; + while ci < n && chars[ci].1 != ':' { + ci += 1; + } + if ci >= n { + break; + } + let key = args_str[key_start..chars[ci].0].trim().trim_matches('"'); + ci += 1; + while ci < n && matches!(chars[ci].1, ' ' | '\n' | '\t') { + ci += 1; + } + if ci >= n { + map.insert(key.to_string(), Value::String(String::new())); + break; + } + + const DELIM: &str = "<|\"|>"; + let byte_pos = chars[ci].0; + if args_str[byte_pos..].starts_with(DELIM) { + let delim_char_len = DELIM.chars().count(); + ci += delim_char_len; + let val_start = if ci < n { chars[ci].0 } else { args_str.len() }; + match args_str[val_start..].find(DELIM) { + Some(rel) => { + let val = &args_str[val_start..val_start + rel]; + map.insert(key.to_string(), Value::String(val.to_string())); + let after = val_start + rel + DELIM.len(); + ci = chars.iter().position(|&(b, _)| b >= after).unwrap_or(n); + } + None => { + map.insert( + key.to_string(), + Value::String(args_str[val_start..].to_string()), + ); + break; + } + } + } else if chars[ci].1 == '"' { + ci += 1; + let val_start = if ci < n { chars[ci].0 } else { args_str.len() }; + let mut end_ci = ci; + while end_ci < n { + if chars[end_ci].1 == '"' && (end_ci == 0 || chars[end_ci - 1].1 != '\\') { + break; + } + end_ci += 1; + } + let val_end = if end_ci < n { + chars[end_ci].0 + } else { + args_str.len() + }; + map.insert( + key.to_string(), + Value::String(args_str[val_start..val_end].to_string()), + ); + ci = if end_ci < n { end_ci + 1 } else { n }; + } else if chars[ci].1 == '{' { + let (inner, after_ci) = Self::gemma4_scan_nested(&chars, ci, '{', '}', n, args_str); + map.insert(key.to_string(), Self::gemma4_parse_args(inner)); + ci = after_ci; + } else if chars[ci].1 == '[' { + let (inner, after_ci) = Self::gemma4_scan_nested(&chars, ci, '[', ']', n, args_str); + map.insert(key.to_string(), Self::gemma4_parse_array(inner)); + ci = after_ci; + } else { + let val_start = chars[ci].0; + while ci < n && !matches!(chars[ci].1, ',' | '}' | ']') { + ci += 1; + } + let val_end = if ci < n { chars[ci].0 } else { args_str.len() }; + map.insert( + key.to_string(), + Self::gemma4_parse_bare_value(args_str[val_start..val_end].trim()), + ); + } + } + + Value::Object(map) + } + + fn gemma4_scan_nested<'a>( + chars: &[(usize, char)], + start_ci: usize, + open: char, + close: char, + n: usize, + source: &'a str, + ) -> (&'a str, usize) { + const DELIM: &str = "<|\"|>"; + let delim_char_len = DELIM.chars().count(); + let mut depth = 1usize; + let mut ci = start_ci + 1; + let inner_start = if ci < n { chars[ci].0 } else { source.len() }; + while ci < n && depth > 0 { + let byte_pos = chars[ci].0; + if source[byte_pos..].starts_with(DELIM) { + ci += delim_char_len; + while ci < n { + let bp = chars[ci].0; + if source[bp..].starts_with(DELIM) { + ci += delim_char_len; + break; + } + ci += 1; + } + continue; + } + if chars[ci].1 == open { + depth += 1; + } else if chars[ci].1 == close { + depth = depth.saturating_sub(1); + } + ci += 1; + } + let inner_end = if depth == 0 && ci > 0 { + chars[ci - 1].0 + } else if ci < n { + chars[ci].0 + } else { + source.len() + }; + (&source[inner_start..inner_end], ci) + } + + fn gemma4_parse_array(arr_str: &str) -> Value { + const DELIM: &str = "<|\"|>"; + let mut items = Vec::new(); + let chars: Vec<(usize, char)> = arr_str.char_indices().collect(); + let n = chars.len(); + let mut ci = 0usize; + while ci < n { + while ci < n && matches!(chars[ci].1, ' ' | ',' | '\n' | '\t') { + ci += 1; + } + if ci >= n { + break; + } + let byte_pos = chars[ci].0; + if arr_str[byte_pos..].starts_with(DELIM) { + let delim_char_len = DELIM.chars().count(); + ci += delim_char_len; + let val_start = if ci < n { chars[ci].0 } else { arr_str.len() }; + match arr_str[val_start..].find(DELIM) { + Some(rel) => { + items.push(Value::String( + arr_str[val_start..val_start + rel].to_string(), + )); + let after = val_start + rel + DELIM.len(); + ci = chars.iter().position(|&(b, _)| b >= after).unwrap_or(n); + } + None => { + items.push(Value::String(arr_str[val_start..].to_string())); + break; + } + } + } else if chars[ci].1 == '"' { + ci += 1; + let val_start = if ci < n { chars[ci].0 } else { arr_str.len() }; + let mut end_ci = ci; + while end_ci < n + && !(chars[end_ci].1 == '"' && (end_ci == 0 || chars[end_ci - 1].1 != '\\')) + { + end_ci += 1; + } + let val_end = if end_ci < n { + chars[end_ci].0 + } else { + arr_str.len() + }; + items.push(Value::String(arr_str[val_start..val_end].to_string())); + ci = if end_ci < n { end_ci + 1 } else { n }; + } else if chars[ci].1 == '{' { + let (inner, after_ci) = Self::gemma4_scan_nested(&chars, ci, '{', '}', n, arr_str); + items.push(Self::gemma4_parse_args(inner)); + ci = after_ci; + } else if chars[ci].1 == '[' { + let (inner, after_ci) = Self::gemma4_scan_nested(&chars, ci, '[', ']', n, arr_str); + items.push(Self::gemma4_parse_array(inner)); + ci = after_ci; + } else { + let val_start = chars[ci].0; + while ci < n && !matches!(chars[ci].1, ',' | ']') { + ci += 1; + } + let val_end = if ci < n { chars[ci].0 } else { arr_str.len() }; + let val = arr_str[val_start..val_end].trim(); + if !val.is_empty() { + items.push(Self::gemma4_parse_bare_value(val)); + } + } + } + Value::Array(items) + } + + fn gemma4_parse_bare_value(val: &str) -> Value { + match val.to_ascii_lowercase().as_str() { + "true" => Value::Bool(true), + "false" => Value::Bool(false), + "null" | "none" | "nil" => Value::Null, + _ => { + if let Ok(n) = val.parse::() { + Value::Number(n.into()) + } else if let Ok(f) = val.parse::() { + serde_json::Number::from_f64(f) + .map(Value::Number) + .unwrap_or_else(|| Value::String(val.to_string())) + } else { + Value::String(val.to_string()) + } + } + } + } + fn strip_tool_tags(&self, text: &str) -> String { let mut output = text.to_string(); if !self.config.start_token_str.is_empty() { @@ -1487,11 +2017,19 @@ impl StreamToolParser { } else if self.config.start_token_str.contains("tool_call") && self.config.end_token_str.contains("tool_call") { - markers.extend( - ["", ""] - .into_iter() - .map(|s| s.to_string()), - ); + if self.uses_glm_xml() { + markers.extend( + ["", "", "", ""] + .into_iter() + .map(|s| s.to_string()), + ); + } else { + markers.extend( + ["", ""] + .into_iter() + .map(|s| s.to_string()), + ); + } } markers } @@ -1524,8 +2062,8 @@ impl StreamToolParser { }; let mut merged_any = false; for (key, value) in recovered { - if !args_obj.contains_key(&key) && !value.is_empty() { - args_obj.insert(key, Value::String(value)); + if !args_obj.contains_key(&key) { + args_obj.insert(key, value); merged_any = true; } } @@ -1541,6 +2079,14 @@ impl StreamToolParser { && self.config.end_token_str == "" } + fn uses_glm_xml(&self) -> bool { + let id = self.model_id.to_ascii_lowercase(); + id.contains("glm") + && !self.uses_minimax_xml() + && self.config.start_token_str == "" + && self.config.end_token_str == "" + } + fn xml_parameter_open_prefix(uses_minimax_xml: bool) -> &'static str { if uses_minimax_xml { ""#); @@ -1601,19 +2147,51 @@ impl StreamToolParser { let value_end = value_start + value_end_rel; let value = section[value_start..value_end] .trim_matches(|c| c == '\n' || c == '\r') + .trim() .to_string(); - recovered.insert(parameter_name, value); + recovered.insert(parameter_name, Self::parse_recovered_xml_value(&value)); cursor = value_end + PARAM_END.len(); } else { let value = section[value_start..] .trim_matches(|c| c == '\n' || c == '\r') + .trim() .to_string(); - recovered.insert(parameter_name, value); + recovered.insert(parameter_name, Self::parse_recovered_xml_value(&value)); break; } } recovered } + + fn parse_recovered_xml_value(raw: &str) -> Value { + let decoded = raw + .replace("<", "<") + .replace(">", ">") + .replace("&", "&") + .replace(""", "\"") + .replace("'", "'"); + + match decoded.as_str() { + "true" | "True" => return Value::Bool(true), + "false" | "False" => return Value::Bool(false), + "null" | "None" => return Value::Null, + _ => {} + } + if (decoded.starts_with('{') || decoded.starts_with('[')) + && serde_json::from_str::(&decoded).is_ok() + { + return serde_json::from_str::(&decoded).unwrap(); + } + if let Ok(num) = decoded.parse::() { + return Value::Number(num.into()); + } + if let Ok(num) = decoded.parse::() { + if let Some(num) = serde_json::Number::from_f64(num) { + return Value::Number(num); + } + } + Value::String(decoded) + } } fn repair_streamed_json_arguments(raw: &str) -> String { @@ -2280,7 +2858,7 @@ Some markdown text here. "#; - let calls = parse_minimax_xml_tool_calls(text); + let calls = parse_minimax_xml_tool_calls(text, &[]); assert_eq!(calls.len(), 1); assert_eq!(calls[0].function.name, "write"); assert!(calls[0].function.arguments.contains("content")); @@ -2294,7 +2872,7 @@ Some markdown text here. /root/AGENTS.md "#; - let calls = parse_minimax_xml_tool_calls(text); + let calls = parse_minimax_xml_tool_calls(text, &[]); assert_eq!(calls.len(), 1); assert_eq!(calls[0].function.name, "read"); assert!(calls[0].function.arguments.contains("filePath")); @@ -2306,10 +2884,213 @@ Some markdown text here. ["rust", "programming"] "#; - let calls = parse_minimax_xml_tool_calls(text); + let calls = parse_minimax_xml_tool_calls(text, &[]); assert_eq!(calls.len(), 1); assert_eq!(calls[0].function.name, "search"); let args: Value = serde_json::from_str(&calls[0].function.arguments).unwrap(); assert!(args["tags"].is_array()); } + + #[test] + fn test_parse_minimax_xml_type_coercion_with_schema() { + let tools = to_openai_tools(&[crate::tools::function_tool("get_weather", "desc") + .parameters_schema(serde_json::json!({ + "type": "object", + "properties": { + "days": {"type": "integer"}, + "include_hourly": {"type": "boolean"}, + "units": {"type": "string"} + } + })) + .build()]); + let text = r#" +3 +true +metric +"#; + + let calls = parse_minimax_xml_tool_calls(text, &tools); + let args: Value = serde_json::from_str(&calls[0].function.arguments).unwrap(); + assert_eq!(args["days"], 3); + assert_eq!(args["include_hourly"], true); + assert_eq!(args["units"], "metric"); + } + + #[test] + fn test_tool_config_llama4() { + let config = ToolConfig::for_model_type(&ToolModelType::LLaMa4); + assert!(config.start_token_ids.contains(&200016)); + assert!(config.end_token_ids.contains(&200007)); + assert!(config.end_token_ids.contains(&200008)); + assert_eq!(config.start_token_str, "<|python_start|>"); + } + + #[test] + fn test_llama4_uses_pythonic_parser() { + assert_eq!( + StreamToolParser::parser_name_for_model( + &ToolModelType::LLaMa4, + "meta-llama/Llama-4-Scout" + ), + "pythonic" + ); + } + + #[test] + fn test_llama4_parse_pythonic_tool_call() { + let tools = vec![crate::tools::function_tool("get_weather", "desc").build()]; + let parser = StreamToolParser::new_with_config( + &ToolModelType::LLaMa4, + "llama4".to_string(), + ToolConfig::for_model_type(&ToolModelType::LLaMa4), + tools, + None, + ); + + let calls = futures::executor::block_on(parser.parse_complete_with_fallback( + r#"[get_weather(location="Vancouver", units="celsius")]"#, + )); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].function.name, "get_weather"); + let args: Value = serde_json::from_str(&calls[0].function.arguments).unwrap(); + assert_eq!(args["location"], "Vancouver"); + assert_eq!(args["units"], "celsius"); + } + + #[test] + fn test_envelope_glm47_xml_format() { + let tools = vec![crate::tools::function_tool("read", "Read a file").build()]; + let mut parser = StreamToolParser::new_with_config( + &ToolModelType::GLM4, + "glm-4.7-flash".to_string(), + ToolConfig::for_model_type(&ToolModelType::GLM4), + tools, + None, + ); + + parser.buffer = + "readfilePath/tmp/test.rs" + .to_string(); + assert!(parser.has_complete_tool_envelope()); + + parser.buffer = + "readfilePath/tmp/test.rs" + .to_string(); + assert!(!parser.has_complete_tool_envelope()); + } + + #[test] + fn test_glm47_display_escape_markers() { + let tools = vec![crate::tools::function_tool("read", "Read a file").build()]; + let parser = StreamToolParser::new_with_config( + &ToolModelType::GLM4, + "glm-4.7-flash".to_string(), + ToolConfig::for_model_type(&ToolModelType::GLM4), + tools, + None, + ); + + let markers = parser.display_escape_markers(); + assert!(markers.iter().any(|m| m == "")); + assert!(markers.iter().any(|m| m == "")); + assert!(markers.iter().any(|m| m == "")); + assert!(markers.iter().any(|m| m == "")); + } + + #[test] + fn test_gemma4_parse_bare_value_case_insensitive() { + assert_eq!( + StreamToolParser::gemma4_parse_bare_value("TRUE"), + Value::Bool(true) + ); + assert_eq!( + StreamToolParser::gemma4_parse_bare_value("False"), + Value::Bool(false) + ); + assert_eq!( + StreamToolParser::gemma4_parse_bare_value("None"), + Value::Null + ); + assert_eq!( + StreamToolParser::gemma4_parse_bare_value("42"), + Value::Number(42.into()) + ); + } + + #[test] + fn test_gemma4_tool_call_parse() { + let tools = vec![crate::tools::function_tool("search", "desc").build()]; + let parser = StreamToolParser::new_with_config( + &ToolModelType::Gemma4, + "gemma4".to_string(), + ToolConfig::for_model_type(&ToolModelType::Gemma4), + tools, + None, + ); + + let text = + r#"<|tool_call>call:search{query:<|"|>rust programming<|"|>,count:5}"#; + let calls = futures::executor::block_on(parser.parse_complete_with_fallback(text)); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].function.name, "search"); + let args: Value = serde_json::from_str(&calls[0].function.arguments).unwrap(); + assert_eq!(args["query"], "rust programming"); + assert_eq!(args["count"], 5); + } + + #[test] + fn test_gemma4_streaming_tool_call_with_reasoning() { + // Simulates Gemma4 output: reasoning in <|channel>..., then tool call + let tools = vec![crate::tools::function_tool("search", "desc").build()]; + let mut parser = StreamToolParser::new_with_config( + &ToolModelType::Gemma4, + "gemma4".to_string(), + ToolConfig::for_model_type(&ToolModelType::Gemma4), + tools, + None, + ); + parser.set_detect_tools_in_reasoning(true); + + // Reasoning block tokens (non-special IDs for text content) + let tokens: Vec<(u32, &str)> = vec![ + (100, "<|channel>"), + (101, "thought"), + (102, "\n"), + (103, "I should search for this."), + (104, "\n"), + (105, ""), + // Tool call start (special token ID 48) + (48, "<|tool_call>"), + (106, "call:search{query:"), + (107, "<|\"|>"), + (108, "rust programming"), + (109, "<|\"|>"), + (110, ",count:5}"), + // Tool call end (special token ID 49) + (49, ""), + ]; + + let mut got_tool_calls = false; + for (id, text) in &tokens { + match parser.process_token(*id, text) { + StreamResult::ToolCalls(calls) => { + assert_eq!(calls.len(), 1, "Expected exactly one tool call"); + assert_eq!(calls[0].function.name, "search"); + let args: Value = serde_json::from_str(&calls[0].function.arguments).unwrap(); + assert_eq!(args["query"], "rust programming"); + assert_eq!(args["count"], 5); + got_tool_calls = true; + } + StreamResult::Buffering => {} + StreamResult::Content(_) => {} + StreamResult::FlushBuffer(_) => {} + } + } + assert!(got_tool_calls, "Expected tool calls to be emitted"); + // After tool call, parser should be back in Normal state + assert!( + matches!(parser.state(), ParserState::Normal), + "Parser should return to Normal state after tool call" + ); + } }