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