From 4d544ab68ffb305429cdd38baab667b843f6ebcc Mon Sep 17 00:00:00 2001
From: jhqxxx <18280426169@163.com>
Date: Wed, 11 Mar 2026 17:30:28 +0800
Subject: [PATCH 1/2] from gguf build chat_template and tokenizer
---
Cargo.lock | 1 +
Cargo.toml | 1 +
src/chat_template/mod.rs | 127 +++++++++++++-----------------
src/gguf_models/common/mod.rs | 136 +++++++++++++++++++++++++++++++++
src/gguf_models/mod.rs | 2 +
src/gguf_models/qwen3_5/mod.rs | 87 +++++++++++++++++++++
src/lib.rs | 1 +
src/models/qwen3/generate.rs | 20 ++---
src/models/qwen3_5/generate.rs | 22 +++---
src/models/qwen3vl/generate.rs | 20 ++---
src/tokenizer/mod.rs | 12 +--
tests/test_gguf_qwen3_5.rs | 47 ++++++++++++
tests/test_glm_asr_nano.rs | 2 +-
tests/test_qwen3_5.rs | 14 +---
14 files changed, 374 insertions(+), 118 deletions(-)
create mode 100644 src/gguf_models/common/mod.rs
create mode 100644 src/gguf_models/mod.rs
create mode 100644 src/gguf_models/qwen3_5/mod.rs
create mode 100644 tests/test_gguf_qwen3_5.rs
diff --git a/Cargo.lock b/Cargo.lock
index 1af6034..ea7b2ec 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -33,6 +33,7 @@ name = "aha"
version = "0.2.2"
dependencies = [
"aha_openai_dive",
+ "ahash",
"anyhow",
"base64 0.22.1",
"byteorder",
diff --git a/Cargo.toml b/Cargo.toml
index 134e639..fd78630 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -42,6 +42,7 @@ zip = "7.2.0"
half = "2.7.1"
byteorder = "1.5.0"
sentencepiece = "0.13.1"
+ahash = "0.8.12"
[features]
flash-attn = ["candle-flash-attn"]
diff --git a/src/chat_template/mod.rs b/src/chat_template/mod.rs
index cbb3f57..33243d2 100644
--- a/src/chat_template/mod.rs
+++ b/src/chat_template/mod.rs
@@ -2,7 +2,35 @@ use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::{Result, anyhow};
use minijinja::{Environment, Value as MiniJinjaValue, context};
-use crate::utils::string_to_static_str;
+use crate::utils::{extract_metadata_value, string_to_static_str};
+
+pub fn fix_template(chat_template: &str) -> String {
+ chat_template
+ .replace(
+ "content.startswith('')",
+ "content is startingwith('')", // 使用minijinja中的 is startingwith 替换
+ )
+ .replace(
+ "content.endswith('')",
+ "content is endingwith('')", // 使用minijinja中的 is endingwith 替换
+ )
+ .replace(
+ "content.split('')[0].rstrip('\\n').split('')[-1].lstrip('\\n')",
+ "((content | split(''))[0] | rstrip('\\n') | split(''))[-1] | lstrip('\\n')", // 使用自定义的split, rstrip, lstrip过滤器替换
+ )
+ .replace(
+ "content.split('')[-1].lstrip('\\n')",
+ "(content | split(''))[-1] | lstrip('\\n')", // 使用自定义的过滤器替换
+ )
+ .replace(
+ "reasoning_content.strip('\\n')",
+ "reasoning_content | strip('\\n')", // 使用自定义的过滤器替换
+ )
+ .replace(
+ "content.lstrip('\\n')",
+ "content | lstrip('\\n')", // 使用自定义的过滤器替换
+ )
+}
pub fn get_template(path: String) -> Result {
let tokenizer_config_file = path.clone() + "/tokenizer_config.json";
@@ -47,31 +75,7 @@ pub fn get_template(path: String) -> Result {
};
let chat_template = chat_template.ok_or(anyhow!(format!("chat_template is none")))?;
// 修复模板中的问题行
- let fixed_template = chat_template
- .replace(
- "content.startswith('')",
- "content is startingwith('')", // 使用minijinja中的 is startingwith 替换
- )
- .replace(
- "content.endswith('')",
- "content is endingwith('')", // 使用minijinja中的 is endingwith 替换
- )
- .replace(
- "content.split('')[0].rstrip('\\n').split('')[-1].lstrip('\\n')",
- "((content | split(''))[0] | rstrip('\\n') | split(''))[-1] | lstrip('\\n')", // 使用自定义的split, rstrip, lstrip过滤器替换
- )
- .replace(
- "content.split('')[-1].lstrip('\\n')",
- "(content | split(''))[-1] | lstrip('\\n')", // 使用自定义的过滤器替换
- )
- .replace(
- "reasoning_content.strip('\\n')",
- "reasoning_content | strip('\\n')", // 使用自定义的过滤器替换
- )
- .replace(
- "content.lstrip('\\n')",
- "content | lstrip('\\n')", // 使用自定义的过滤器替换
- );
+ let fixed_template = fix_template(&chat_template);
Ok(fixed_template)
}
@@ -80,29 +84,7 @@ pub struct ChatTemplate<'a> {
}
impl<'a> ChatTemplate<'a> {
- pub fn init(path: &str) -> Result {
- let path: String = path.to_string();
- if !std::path::Path::new(&path).exists() {
- return Err(anyhow!("model path not found"));
- }
- let template = get_template(path.clone())?;
- // let template = match get_template(path.clone()) {
- // Ok(template) => template,
- // Err(e) => {
- // let jinja_path = path + "/chat_template.jinja";
- // if !std::path::Path::new(&jinja_path).exists() {
- // return Err(anyhow!(
- // "get_template err {e} and chat_template.jinja not found"
- // ));
- // }
- // std::fs::read_to_string(&jinja_path)
- // .map_err(|e| anyhow!("Failed to read chat_template.jinja: {}", e))?
- // }
- // };
- let template = string_to_static_str(template);
- // 加载jinjaenv处理chat_template
- let mut env = Environment::new();
- // 添加自定义过滤器
+ fn setup_environment(env: &mut Environment<'a>) {
env.add_filter("tojson", |v: MiniJinjaValue| {
serde_json::to_string(&v).unwrap()
});
@@ -113,44 +95,44 @@ impl<'a> ChatTemplate<'a> {
.collect::>()
});
- // 添加 lstrip 过滤器
env.add_filter("lstrip", |s: String, chars: Option| match chars {
Some(chars_str) => s.trim_start_matches(chars_str.as_str()).to_string(),
None => s.trim_start().to_string(),
});
- // 添加 rstrip 过滤器
env.add_filter("rstrip", |s: String, chars: Option| match chars {
Some(chars_str) => s.trim_end_matches(chars_str.as_str()).to_string(),
None => s.trim_end().to_string(),
});
- // let template = get_template(path.to_string())?;
+ }
+ pub fn init(path: &str) -> Result {
+ let path: String = path.to_string();
+ if !std::path::Path::new(&path).exists() {
+ return Err(anyhow!("model path not found"));
+ }
+ let template = get_template(path.clone())?;
+ let template = string_to_static_str(template);
+ // 加载jinjaenv处理chat_template
+ let mut env = Environment::new();
+ Self::setup_environment(&mut env);
+ let _ = env.add_template("chat", template);
+
+ Ok(Self { env })
+ }
+
+ pub fn str_init(chat_template: &str) -> Result {
+ let fixed_template = fix_template(chat_template);
+ let template = string_to_static_str(fixed_template);
+ // 加载jinjaenv处理chat_template
+ let mut env = Environment::new();
+ Self::setup_environment(&mut env);
let _ = env.add_template("chat", template);
Ok(Self { env })
}
pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result {
- let context = context! {
- messages => &messages.messages,
- tools => &messages.tools.as_ref(),
- add_generation_prompt => true,
- };
- let template = self
- .env
- .get_template("chat")
- .map_err(|e| anyhow!(format!("render template error {}", e)))?;
- let message_str = template
- .render(context)
- .map_err(|e| anyhow!(format!("render template error {}", e)))?;
- Ok(message_str)
- }
-
- pub fn apply_chat_temp_think(
- &self,
- messages: &ChatCompletionParameters,
- enable_thinking: Option,
- ) -> Result {
+ let enable_thinking = extract_metadata_value::(&messages.metadata, "enable_thinking");
let context = context! {
messages => &messages.messages,
tools => &messages.tools.as_ref(),
@@ -166,4 +148,5 @@ impl<'a> ChatTemplate<'a> {
.map_err(|e| anyhow!(format!("render template error {}", e)))?;
Ok(message_str)
}
+
}
diff --git a/src/gguf_models/common/mod.rs b/src/gguf_models/common/mod.rs
new file mode 100644
index 0000000..01ed82f
--- /dev/null
+++ b/src/gguf_models/common/mod.rs
@@ -0,0 +1,136 @@
+use std::io::{Read, Seek};
+
+use ahash::AHashMap;
+use anyhow::{Result, anyhow};
+use candle_core::{
+ Device,
+ quantized::{
+ QMatMul, QTensor,
+ gguf_file::{self, Value},
+ },
+};
+use candle_nn::RmsNorm;
+use tokenizers::{self, AddedToken, Tokenizer, models::bpe::BPE};
+
+use crate::tokenizer::TokenizerModel;
+
+pub struct Gguf {
+ ct: gguf_file::Content,
+ reader: R,
+ device: Device,
+}
+
+impl Gguf {
+ pub fn new(ct: gguf_file::Content, reader: R, device: Device) -> Self {
+ Self { ct, reader, device }
+ }
+
+ pub fn get_matedata(&self, name: &str) -> Result {
+ match self.ct.metadata.get(name) {
+ None => Err(anyhow!("cannot find {name} in metadata")),
+ Some(v) => Ok(v.clone()),
+ }
+ }
+
+ pub fn qmatmul(&mut self, name: &str) -> Result {
+ let ws = self.ct.tensor(&mut self.reader, name, &self.device)?;
+ Ok(QMatMul::from_arc(ws.into())?)
+ }
+
+ pub fn rms_norm(&mut self, name: &str, eps: f64) -> Result {
+ let ws = self.ct.tensor(&mut self.reader, name, &self.device)?;
+ let weight = ws.dequantize(&self.device)?;
+ Ok(RmsNorm::new(weight, eps))
+ }
+
+ pub fn metadata(&self) -> &std::collections::HashMap {
+ &self.ct.metadata
+ }
+
+ pub fn tensor(&mut self, name: &str) -> Result {
+ Ok(self.ct.tensor(&mut self.reader, name, &self.device)?)
+ }
+
+ pub fn build_tokenizer(
+ &self,
+ add_prefix_space: Option,
+ trim_offsets: Option,
+ use_regex: Option,
+ ) -> Result {
+ let model_type = self
+ .get_matedata("tokenizer.ggml.model")?
+ .to_string()?
+ .clone();
+ match model_type.as_str() {
+ "gpt2" | "llama" => {
+ let vocab = self
+ .get_matedata("tokenizer.ggml.tokens")?
+ .to_vec()?
+ .clone();
+ let vocab: Vec = vocab
+ .into_iter()
+ .map(|tokens| tokens.to_string().map(|x| x.clone()))
+ .collect::, candle_core::Error>>()?;
+ let mut vocab_map = AHashMap::new();
+ for (id, token) in vocab.iter().enumerate() {
+ vocab_map.insert(token.clone(), id as u32);
+ }
+
+ let merges = self
+ .get_matedata("tokenizer.ggml.merges")?
+ .to_vec()?
+ .clone();
+ let merges: Vec = merges
+ .into_iter()
+ .map(|tokens| tokens.to_string().map(|x| x.clone()))
+ .collect::, candle_core::Error>>()?;
+ let merges: Vec<(String, String)> = merges
+ .into_iter()
+ .map(|token_merge| {
+ let merge: Vec<&str> = token_merge.split(" ").collect();
+ if merge.len() != 2 {
+ // 处理格式不正确的merge规则
+ return ("".to_string(), "".to_string());
+ }
+ (merge[0].to_string(), merge[1].to_string())
+ })
+ .filter(|(a, b)| !a.is_empty() && !b.is_empty())
+ .collect();
+ let bpe_model = BPE::new(vocab_map, merges);
+ let mut tokenizer = Tokenizer::new(bpe_model);
+ let add_prefix_space = add_prefix_space.unwrap_or(false);
+ let trim_offsets = trim_offsets.unwrap_or(false);
+ let use_regex = use_regex.unwrap_or(false);
+ let pre_byte_level = tokenizers::pre_tokenizers::byte_level::ByteLevel::default()
+ .add_prefix_space(add_prefix_space) // 是否在文本开头添加空格,gpt-2默认是true
+ .trim_offsets(trim_offsets) // 是否删除首尾空白字符
+ .use_regex(use_regex); // 是否使用正则表达式来分割特殊字符
+ tokenizer.with_pre_tokenizer(Some(pre_byte_level));
+ let dec_byte_level = tokenizers::decoders::byte_level::ByteLevel::default();
+ tokenizer.with_decoder(Some(dec_byte_level));
+ let token_types = self
+ .get_matedata("tokenizer.ggml.token_type")?
+ .to_vec()?
+ .clone();
+ let token_types = token_types
+ .into_iter()
+ .map(|types| types.to_i32())
+ .collect::, candle_core::Error>>()?;
+
+ let mut add_tokens = vec![];
+ for (id, type_) in token_types.into_iter().enumerate() {
+ if type_ == 3 || type_ == 4 {
+ if let Some(token_str) = vocab.get(id) {
+ let add_token = AddedToken::from(token_str.clone(), true);
+ add_tokens.push(add_token);
+ }
+ }
+ }
+ let _ = tokenizer.add_special_tokens(&add_tokens);
+ let tokenizer_model = TokenizerModel::new(tokenizer);
+ Ok(tokenizer_model)
+ }
+ _ => Err(anyhow!("Unsupported tokenizer model type: {model_type}")),
+ }
+ }
+}
diff --git a/src/gguf_models/mod.rs b/src/gguf_models/mod.rs
new file mode 100644
index 0000000..19761fe
--- /dev/null
+++ b/src/gguf_models/mod.rs
@@ -0,0 +1,2 @@
+pub mod common;
+pub mod qwen3_5;
\ No newline at end of file
diff --git a/src/gguf_models/qwen3_5/mod.rs b/src/gguf_models/qwen3_5/mod.rs
new file mode 100644
index 0000000..8425af5
--- /dev/null
+++ b/src/gguf_models/qwen3_5/mod.rs
@@ -0,0 +1,87 @@
+use std::io::{Read, Seek};
+
+use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
+use anyhow::{Result, anyhow};
+use candle_core::{DType, Device, quantized::gguf_file};
+use candle_nn::Embedding;
+
+use crate::{
+ chat_template::ChatTemplate, gguf_models::common::Gguf, tokenizer::TokenizerModel,
+ utils::get_device,
+};
+
+pub struct GgufQwen3_5<'a> {
+ chat_template: ChatTemplate<'a>,
+ tokenizer: TokenizerModel,
+ embed_tokens: Embedding,
+ device: Device,
+ dtype: DType,
+}
+
+impl<'a> GgufQwen3_5<'a> {
+ pub fn new(file_path: &str, device: Option<&Device>) -> Result {
+ if !file_path.ends_with("gguf") {
+ return Err(anyhow!("model file suffix must be gguf: {file_path}"));
+ }
+ let mut reader = std::fs::File::open(file_path)?;
+ let content = gguf_file::Content::read(&mut reader)?;
+ let device = get_device(device);
+ Self::from_gguf(content, &mut reader, &device)
+ }
+ pub fn from_gguf(
+ content: gguf_file::Content,
+ reader: &mut R,
+ device: &Device,
+ ) -> Result {
+ let mut gguf = Gguf::new(content, reader, device.clone());
+
+ let chat_template_str = gguf
+ .get_matedata("tokenizer.chat_template")?
+ .to_string()?
+ .clone();
+ let chat_template = ChatTemplate::str_init(&chat_template_str)?;
+ let tokenizer = gguf.build_tokenizer(Some(false), Some(false), Some(false))?;
+
+ let num_attention_heads =
+ gguf.get_matedata("qwen35.attention.head_count")?.to_u32()? as usize;
+ let num_kv_heads = gguf
+ .get_matedata("qwen35.attention.head_count_kv")?
+ .to_u32()? as usize;
+ let head_dim = gguf.get_matedata("qwen35.attention.key_length")?.to_u32()? as usize;
+ let num_layers = gguf.get_matedata("qwen35.block_count")?.to_u32()? as usize;
+ let hidden_size = gguf.get_matedata("qwen35.embedding_length")?.to_u32()? as usize;
+ let max_position_embeddings =
+ gguf.get_matedata("qwen35.context_length")?.to_u32()? as usize;
+ let rms_norm_eps = gguf
+ .get_matedata("qwen35.attention.layer_norm_rms_epsilon")?
+ .to_f32()? as f64;
+ let rope_freq_base = gguf.get_matedata("qwen35.rope.freq_base")?.to_f32()? as f64;
+
+ let dtype = match gguf.get_matedata("general.type") {
+ Ok(v) => match v.to_u32() {
+ Ok(0) => DType::F32,
+ Ok(1) => DType::F16,
+ _ => DType::F16,
+ },
+ Err(_) => DType::F16,
+ };
+
+ let embed_tensor = gguf.tensor("token_embd.weight")?;
+ let embed_tokens = Embedding::new(embed_tensor.dequantize(device)?, hidden_size);
+ Ok(Self {
+ chat_template,
+ tokenizer,
+ embed_tokens,
+ device: device.clone(),
+ dtype,
+ })
+ }
+
+ pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<()> {
+ let render = self.chat_template.apply_chat_template(&mes)?;
+ println!("render: {}", render);
+ let input_ids = self.tokenizer.text_encode(render, &self.device)?;
+ println!("input_ids: {}", input_ids);
+ Ok(())
+ }
+}
diff --git a/src/lib.rs b/src/lib.rs
index e454ceb..1f25a30 100644
--- a/src/lib.rs
+++ b/src/lib.rs
@@ -1,5 +1,6 @@
pub mod chat_template;
pub mod exec;
+pub mod gguf_models;
pub mod models;
pub mod position_embed;
pub mod process;
diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs
index 2a56a39..f792670 100644
--- a/src/models/qwen3/generate.rs
+++ b/src/models/qwen3/generate.rs
@@ -66,11 +66,12 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
- let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // let mes_render = self.chat_template.apply_chat_template(&mes)?;
- let mes_render = self
- .chat_template
- .apply_chat_temp_think(&mes, enable_thinking)?;
+
+ let mes_render = self.chat_template.apply_chat_template(&mes)?;
+ // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
+ // let mes_render = self
+ // .chat_template
+ // .apply_chat_temp_think(&mes, enable_thinking)?;
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
@@ -115,10 +116,11 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
- let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- let mes_render = self
- .chat_template
- .apply_chat_temp_think(&mes, enable_thinking)?;
+ let mes_render = self.chat_template.apply_chat_template(&mes)?;
+ // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
+ // let mes_render = self
+ // .chat_template
+ // .apply_chat_temp_think(&mes, enable_thinking)?;
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0;
diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs
index 82b23ed..d7c747f 100644
--- a/src/models/qwen3_5/generate.rs
+++ b/src/models/qwen3_5/generate.rs
@@ -60,16 +60,19 @@ impl<'a> Qwen3_5GenerateModel<'a> {
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
- let seed = mes.seed.unwrap_or(34562) as u64;
- let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
+ let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
- let mes_render = self
- .chat_template
- .apply_chat_temp_think(&mes, enable_thinking)?;
+ let mes_render = self.chat_template.apply_chat_template(&mes)?;
+ println!("mes_render: {}", mes_render);
+ // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
+ // let mes_render = self
+ // .chat_template
+ // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self
.tokenizer
.text_encode(input.replace_text.clone(), &self.device)?;
+ println!("input_ids: {}", input_ids);
let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0;
@@ -125,10 +128,11 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
> {
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
- let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- let mes_render = self
- .chat_template
- .apply_chat_temp_think(&mes, enable_thinking)?;
+ let mes_render = self.chat_template.apply_chat_template(&mes)?;
+ // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
+ // let mes_render = self
+ // .chat_template
+ // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self
.tokenizer
diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs
index 3918d63..032ce84 100644
--- a/src/models/qwen3vl/generate.rs
+++ b/src/models/qwen3vl/generate.rs
@@ -73,11 +73,11 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
- let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // let mes_render = self.chat_template.apply_chat_template(&mes)?;
- let mes_render = self
- .chat_template
- .apply_chat_temp_think(&mes, enable_thinking)?;
+ let mes_render = self.chat_template.apply_chat_template(&mes)?;
+ // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
+ // let mes_render = self
+ // .chat_template
+ // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self
.tokenizer
@@ -142,11 +142,11 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
- let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // let mes_render = self.chat_template.apply_chat_template(&mes)?;
- let mes_render = self
- .chat_template
- .apply_chat_temp_think(&mes, enable_thinking)?;
+ let mes_render = self.chat_template.apply_chat_template(&mes)?;
+ // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
+ // let mes_render = self
+ // .chat_template
+ // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self
.tokenizer
diff --git a/src/tokenizer/mod.rs b/src/tokenizer/mod.rs
index 1f1a5e8..fd8fdbc 100644
--- a/src/tokenizer/mod.rs
+++ b/src/tokenizer/mod.rs
@@ -12,6 +12,10 @@ pub struct TokenizerModel {
}
impl TokenizerModel {
+ pub fn new(tokenizer: Tokenizer) -> Self {
+ Self { tokenizer }
+ }
+
pub fn init(path: &str) -> Result {
let path = path.to_string();
assert!(
@@ -81,6 +85,8 @@ impl TokenizerModel {
}
tokenizer
};
+ let len = tokenizer.get_vocab_size(true);
+ println!("len: {}", len);
Ok(Self { tokenizer })
}
@@ -94,12 +100,6 @@ impl TokenizerModel {
Ok(token_id)
}
pub fn text_encode(&self, text: String, device: &Device) -> Result {
- // let token_id = self
- // .tokenizer
- // .encode(text, true)
- // .map_err(|e| anyhow!(format!("tokenizer encode error: {}", e)))?
- // .get_ids()
- // .to_vec();
let token_id = self.text_encode_vec(text, true)?;
let token_tensor = Tensor::from_slice(&token_id, (1, token_id.len()), device)?;
Ok(token_tensor)
diff --git a/tests/test_gguf_qwen3_5.rs b/tests/test_gguf_qwen3_5.rs
new file mode 100644
index 0000000..236b904
--- /dev/null
+++ b/tests/test_gguf_qwen3_5.rs
@@ -0,0 +1,47 @@
+use aha::{chat::ChatCompletionParameters, gguf_models::qwen3_5::GgufQwen3_5};
+use anyhow::Result;
+use candle_core::{Device, quantized::gguf_file};
+#[test]
+fn gguf_test() -> Result<()> {
+ // cargo test -r -F cuda --test test_gguf_qwen3_5 gguf_test -- --nocapture
+ let path = "/home/jhq/.aha/Qwen/Qwen3.5-0.8B-GGUF/Qwen3.5-0.8B-Q4_K_M.gguf";
+ let mut file = std::fs::File::open(path)?;
+ let model = gguf_file::Content::read(&mut file)?;
+ let device = Device::new_cuda(0)?;
+ // println!("model: {:?}", model.magic);
+ // println!("generat.type: {:#?}", model.metadata.keys());
+ // println!("tokenizer.ggml.model: {:#?}", model.metadata.get("tokenizer.ggml.model")); // gpt2
+ // // println!("model: {:?}", model.tensor_infos);
+
+ let message = r#"
+ {
+ "model": "qwen3.5",
+ "messages": [
+ {
+ "role": "user",
+ "content": [
+ {
+ "type": "text",
+ "text": "你好啊"
+ }
+ ]
+ }
+ ]
+ }
+ "#;
+// render: <|im_start|>user
+// 你好啊<|im_end|>
+// <|im_start|>assistant
+//
+
+//
+
+
+// input_ids: [[248045, 846, 198, 109266, 98710, 248046, 198, 248045, 74455, 198,
+// 248068, 271, 248069, 271]]
+// Tensor[[1, 14], u32, cuda:0]
+ let mes: ChatCompletionParameters = serde_json::from_str(message)?;
+ let mut gguf_qwen3_5 = GgufQwen3_5::new(&path, None)?;
+ let _ = gguf_qwen3_5.generate(mes)?;
+ Ok(())
+}
\ No newline at end of file
diff --git a/tests/test_glm_asr_nano.rs b/tests/test_glm_asr_nano.rs
index 8e75549..49e9ebe 100644
--- a/tests/test_glm_asr_nano.rs
+++ b/tests/test_glm_asr_nano.rs
@@ -22,7 +22,7 @@ fn glm_asr_nano_generate() -> Result<()> {
"type": "audio",
"audio_url":
{
- "url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav"
+ "url": "file://./assets/audio/zh.mp3"
}
},
{
diff --git a/tests/test_qwen3_5.rs b/tests/test_qwen3_5.rs
index 7e3b7f7..c8b108d 100644
--- a/tests/test_qwen3_5.rs
+++ b/tests/test_qwen3_5.rs
@@ -19,22 +19,14 @@ fn qwen3_5_generate() -> Result<()> {
"messages": [
{
"role": "user",
- "content": [
- {
- "type": "image",
- "image_url":
- {
- "url": "https://www.lifeberrys.com/img/article/tourist-attraction-3-1644590220-lb.jpg"
- }
- },
+ "content": [
{
"type": "text",
- "text": "描述这张图片."
+ "text": "你好啊"
}
]
}
- ],
- "metadata": {"enable_thinking": "true"}
+ ]
}
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
From aad0e62151a5104c0962ff087aad79f96ff82d06 Mon Sep 17 00:00:00 2001
From: jhqxxx <18280426169@163.com>
Date: Sat, 14 Mar 2026 19:27:01 +0800
Subject: [PATCH 2/2] add qwen3.5 gguf
---
src/chat_template/mod.rs | 1 -
src/gguf_models/mod.rs | 2 -
src/gguf_models/qwen3_5/mod.rs | 87 ----
src/lib.rs | 1 -
.../common/mod.rs => models/common/gguf.rs} | 118 ++++-
src/models/common/mod.rs | 2 +
src/models/qwen3/generate.rs | 6 +-
src/models/qwen3_5/generate.rs | 115 ++++-
src/models/qwen3_5/model.rs | 432 ++++++++++++++----
src/models/qwen3vl/config.rs | 30 ++
src/models/qwen3vl/generate.rs | 8 +-
src/models/qwen3vl/processor.rs | 22 +
src/tokenizer/mod.rs | 2 -
tests/test_gguf_qwen3_5.rs | 65 ++-
tests/test_qwen3_5.rs | 6 +-
15 files changed, 656 insertions(+), 241 deletions(-)
delete mode 100644 src/gguf_models/mod.rs
delete mode 100644 src/gguf_models/qwen3_5/mod.rs
rename src/{gguf_models/common/mod.rs => models/common/gguf.rs} (56%)
diff --git a/src/chat_template/mod.rs b/src/chat_template/mod.rs
index 33243d2..4b3b6f0 100644
--- a/src/chat_template/mod.rs
+++ b/src/chat_template/mod.rs
@@ -148,5 +148,4 @@ impl<'a> ChatTemplate<'a> {
.map_err(|e| anyhow!(format!("render template error {}", e)))?;
Ok(message_str)
}
-
}
diff --git a/src/gguf_models/mod.rs b/src/gguf_models/mod.rs
deleted file mode 100644
index 19761fe..0000000
--- a/src/gguf_models/mod.rs
+++ /dev/null
@@ -1,2 +0,0 @@
-pub mod common;
-pub mod qwen3_5;
\ No newline at end of file
diff --git a/src/gguf_models/qwen3_5/mod.rs b/src/gguf_models/qwen3_5/mod.rs
deleted file mode 100644
index 8425af5..0000000
--- a/src/gguf_models/qwen3_5/mod.rs
+++ /dev/null
@@ -1,87 +0,0 @@
-use std::io::{Read, Seek};
-
-use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
-use anyhow::{Result, anyhow};
-use candle_core::{DType, Device, quantized::gguf_file};
-use candle_nn::Embedding;
-
-use crate::{
- chat_template::ChatTemplate, gguf_models::common::Gguf, tokenizer::TokenizerModel,
- utils::get_device,
-};
-
-pub struct GgufQwen3_5<'a> {
- chat_template: ChatTemplate<'a>,
- tokenizer: TokenizerModel,
- embed_tokens: Embedding,
- device: Device,
- dtype: DType,
-}
-
-impl<'a> GgufQwen3_5<'a> {
- pub fn new(file_path: &str, device: Option<&Device>) -> Result {
- if !file_path.ends_with("gguf") {
- return Err(anyhow!("model file suffix must be gguf: {file_path}"));
- }
- let mut reader = std::fs::File::open(file_path)?;
- let content = gguf_file::Content::read(&mut reader)?;
- let device = get_device(device);
- Self::from_gguf(content, &mut reader, &device)
- }
- pub fn from_gguf(
- content: gguf_file::Content,
- reader: &mut R,
- device: &Device,
- ) -> Result {
- let mut gguf = Gguf::new(content, reader, device.clone());
-
- let chat_template_str = gguf
- .get_matedata("tokenizer.chat_template")?
- .to_string()?
- .clone();
- let chat_template = ChatTemplate::str_init(&chat_template_str)?;
- let tokenizer = gguf.build_tokenizer(Some(false), Some(false), Some(false))?;
-
- let num_attention_heads =
- gguf.get_matedata("qwen35.attention.head_count")?.to_u32()? as usize;
- let num_kv_heads = gguf
- .get_matedata("qwen35.attention.head_count_kv")?
- .to_u32()? as usize;
- let head_dim = gguf.get_matedata("qwen35.attention.key_length")?.to_u32()? as usize;
- let num_layers = gguf.get_matedata("qwen35.block_count")?.to_u32()? as usize;
- let hidden_size = gguf.get_matedata("qwen35.embedding_length")?.to_u32()? as usize;
- let max_position_embeddings =
- gguf.get_matedata("qwen35.context_length")?.to_u32()? as usize;
- let rms_norm_eps = gguf
- .get_matedata("qwen35.attention.layer_norm_rms_epsilon")?
- .to_f32()? as f64;
- let rope_freq_base = gguf.get_matedata("qwen35.rope.freq_base")?.to_f32()? as f64;
-
- let dtype = match gguf.get_matedata("general.type") {
- Ok(v) => match v.to_u32() {
- Ok(0) => DType::F32,
- Ok(1) => DType::F16,
- _ => DType::F16,
- },
- Err(_) => DType::F16,
- };
-
- let embed_tensor = gguf.tensor("token_embd.weight")?;
- let embed_tokens = Embedding::new(embed_tensor.dequantize(device)?, hidden_size);
- Ok(Self {
- chat_template,
- tokenizer,
- embed_tokens,
- device: device.clone(),
- dtype,
- })
- }
-
- pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<()> {
- let render = self.chat_template.apply_chat_template(&mes)?;
- println!("render: {}", render);
- let input_ids = self.tokenizer.text_encode(render, &self.device)?;
- println!("input_ids: {}", input_ids);
- Ok(())
- }
-}
diff --git a/src/lib.rs b/src/lib.rs
index 1f25a30..e454ceb 100644
--- a/src/lib.rs
+++ b/src/lib.rs
@@ -1,6 +1,5 @@
pub mod chat_template;
pub mod exec;
-pub mod gguf_models;
pub mod models;
pub mod position_embed;
pub mod process;
diff --git a/src/gguf_models/common/mod.rs b/src/models/common/gguf.rs
similarity index 56%
rename from src/gguf_models/common/mod.rs
rename to src/models/common/gguf.rs
index 01ed82f..cb54c27 100644
--- a/src/gguf_models/common/mod.rs
+++ b/src/models/common/gguf.rs
@@ -3,13 +3,13 @@ use std::io::{Read, Seek};
use ahash::AHashMap;
use anyhow::{Result, anyhow};
use candle_core::{
- Device,
+ Device, Tensor,
quantized::{
QMatMul, QTensor,
gguf_file::{self, Value},
},
};
-use candle_nn::RmsNorm;
+use candle_nn::{Conv1d, Conv1dConfig, Linear, Module, RmsNorm, VarBuilder, linear_b};
use tokenizers::{self, AddedToken, Tokenizer, models::bpe::BPE};
use crate::tokenizer::TokenizerModel;
@@ -34,7 +34,7 @@ impl Gguf {
pub fn qmatmul(&mut self, name: &str) -> Result {
let ws = self.ct.tensor(&mut self.reader, name, &self.device)?;
- Ok(QMatMul::from_arc(ws.into())?)
+ Ok(QMatMul::from_qtensor(ws)?)
}
pub fn rms_norm(&mut self, name: &str, eps: f64) -> Result {
@@ -51,6 +51,35 @@ impl Gguf {
Ok(self.ct.tensor(&mut self.reader, name, &self.device)?)
}
+ pub fn get_dequantized(&mut self, name: &str) -> Result {
+ Ok(self.tensor(name)?.dequantize(&self.device)?)
+ }
+
+ pub fn conv1d(
+ &mut self,
+ prefix: &str,
+ padding: usize,
+ stride: usize,
+ dilation: usize,
+ groups: usize,
+ bias: bool,
+ ) -> Result {
+ let weight = self.get_dequantized(&format!("{prefix}.weight"))?;
+ let bias = if bias {
+ self.get_dequantized(&format!("{prefix}.bias")).ok()
+ } else {
+ None
+ };
+ let cfg = Conv1dConfig {
+ padding,
+ stride,
+ dilation,
+ groups,
+ cudnn_fwd_algo: None,
+ };
+ Ok(Conv1d::new(weight, bias, cfg))
+ }
+
pub fn build_tokenizer(
&self,
add_prefix_space: Option,
@@ -69,7 +98,7 @@ impl Gguf {
.clone();
let vocab: Vec = vocab
.into_iter()
- .map(|tokens| tokens.to_string().map(|x| x.clone()))
+ .map(|tokens| tokens.to_string().cloned())
.collect::, candle_core::Error>>()?;
let mut vocab_map = AHashMap::new();
for (id, token) in vocab.iter().enumerate() {
@@ -82,7 +111,7 @@ impl Gguf {
.clone();
let merges: Vec = merges
.into_iter()
- .map(|tokens| tokens.to_string().map(|x| x.clone()))
+ .map(|tokens| tokens.to_string().cloned())
.collect::, candle_core::Error>>()?;
let merges: Vec<(String, String)> = merges
.into_iter()
@@ -119,11 +148,16 @@ impl Gguf {
let mut add_tokens = vec![];
for (id, type_) in token_types.into_iter().enumerate() {
- if type_ == 3 || type_ == 4 {
- if let Some(token_str) = vocab.get(id) {
- let add_token = AddedToken::from(token_str.clone(), true);
- add_tokens.push(add_token);
- }
+ if type_ == 3
+ && let Some(token_str) = vocab.get(id)
+ {
+ let add_token = AddedToken::from(token_str.clone(), true);
+ add_tokens.push(add_token);
+ } else if type_ == 4
+ && let Some(token_str) = vocab.get(id)
+ {
+ let add_token = AddedToken::from(token_str.clone(), false);
+ add_tokens.push(add_token);
}
}
let _ = tokenizer.add_special_tokens(&add_tokens);
@@ -134,3 +168,67 @@ impl Gguf {
}
}
}
+
+#[derive(Debug, Clone)]
+pub enum ProjKind {
+ QuantizedProj(QMatMul),
+ LinearProj(Linear),
+}
+
+impl Module for ProjKind {
+ fn forward(&self, xs: &Tensor) -> candle_core::Result {
+ match self {
+ ProjKind::QuantizedProj(q) => q.forward(xs),
+ ProjKind::LinearProj(l) => l.forward(xs),
+ }
+ }
+}
+
+#[derive(Debug, Clone)]
+pub struct GateUpDownMLPGguf {
+ gate_proj: ProjKind, // ffn_gate.weight
+ up_proj: ProjKind, // ffn_up.weight
+ down_proj: ProjKind, // ffn_down.weight
+}
+
+impl GateUpDownMLPGguf {
+ pub fn new_from_gguf(gguf: &mut Gguf, prefix: &str) -> Result {
+ let gate_proj = gguf.qmatmul(&format!("{prefix}.ffn_gate.weight"))?;
+ let up_proj = gguf.qmatmul(&format!("{prefix}.ffn_up.weight"))?;
+ let down_proj = gguf.qmatmul(&format!("{prefix}.ffn_down.weight"))?;
+ Ok(Self {
+ gate_proj: ProjKind::QuantizedProj(gate_proj),
+ up_proj: ProjKind::QuantizedProj(up_proj),
+ down_proj: ProjKind::QuantizedProj(down_proj),
+ })
+ }
+ pub fn new_from_vb(
+ vb: VarBuilder,
+ hidden_size: usize,
+ intermediate_size: usize,
+ bias: bool,
+ gate_pp_name: Option<&str>,
+ up_pp_name: Option<&str>,
+ down_pp_name: Option<&str>,
+ ) -> Result {
+ let gate_pp_name = gate_pp_name.unwrap_or("gate_proj");
+ let up_pp_name = up_pp_name.unwrap_or("up_proj");
+ let down_pp_name = down_pp_name.unwrap_or("down_proj");
+ let gate_proj = linear_b(hidden_size, intermediate_size, bias, vb.pp(gate_pp_name))?;
+ let up_proj = linear_b(hidden_size, intermediate_size, bias, vb.pp(up_pp_name))?;
+ let down_proj = linear_b(intermediate_size, hidden_size, bias, vb.pp(down_pp_name))?;
+ Ok(Self {
+ gate_proj: ProjKind::LinearProj(gate_proj),
+ up_proj: ProjKind::LinearProj(up_proj),
+ down_proj: ProjKind::LinearProj(down_proj),
+ })
+ }
+}
+
+impl Module for GateUpDownMLPGguf {
+ fn forward(&self, xs: &Tensor) -> candle_core::Result {
+ let w1 = self.gate_proj.forward(xs)?;
+ let w3 = self.up_proj.forward(xs)?;
+ self.down_proj.forward(&(candle_nn::ops::silu(&w1)? * w3)?)
+ }
+}
diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs
index 05abc32..0a26948 100644
--- a/src/models/common/mod.rs
+++ b/src/models/common/mod.rs
@@ -7,6 +7,8 @@ use candle_nn::{
embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm,
};
+pub mod gguf;
+
use crate::{
position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer},
utils::tensor_utils::{prepare_causal_attention_mask, repeat_kv},
diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs
index f792670..f06d85c 100644
--- a/src/models/qwen3/generate.rs
+++ b/src/models/qwen3/generate.rs
@@ -11,8 +11,8 @@ use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
use crate::models::qwen3::model::Qwen3Model;
// use crate::models::GenerateStream;
use crate::utils::{
- build_completion_chunk_response, build_completion_response, extract_metadata_value,
- find_type_files, get_device, get_dtype, get_logit_processor,
+ build_completion_chunk_response, build_completion_response, find_type_files, get_device,
+ get_dtype, get_logit_processor,
};
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
@@ -66,7 +66,7 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
-
+
let mes_render = self.chat_template.apply_chat_template(&mes)?;
// let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
// let mes_render = self
diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs
index d7c747f..15483df 100644
--- a/src/models/qwen3_5/generate.rs
+++ b/src/models/qwen3_5/generate.rs
@@ -2,7 +2,7 @@ use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
use anyhow::{Result, anyhow};
-use candle_core::{DType, Device, Tensor};
+use candle_core::{DType, Device, Tensor, quantized::gguf_file};
use candle_nn::VarBuilder;
use rocket::async_stream::stream;
use rocket::futures::Stream;
@@ -11,13 +11,14 @@ use crate::{
chat_template::ChatTemplate,
models::{
GenerateModel,
+ common::gguf::Gguf,
qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model},
qwen3vl::processor::Qwen3VLProcessor,
},
tokenizer::TokenizerModel,
utils::{
- build_completion_chunk_response, build_completion_response, extract_metadata_value,
- find_type_files, get_device, get_dtype, get_logit_processor,
+ build_completion_chunk_response, build_completion_response, find_type_files, get_device,
+ get_dtype, get_logit_processor,
},
};
@@ -29,10 +30,17 @@ pub struct Qwen3_5GenerateModel<'a> {
device: Device,
eos_token_id: u32,
model_name: String,
+ repeat_penalty: f32,
+ repeat_last_n: usize,
}
impl<'a> Qwen3_5GenerateModel<'a> {
pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result {
+ let model_name = path
+ .split("/")
+ .collect::>()
+ .pop()
+ .unwrap_or("qwen3.5");
let chat_template = ChatTemplate::init(path)?;
let tokenizer = TokenizerModel::init(path)?;
let config_path = path.to_string() + "/config.json";
@@ -44,7 +52,7 @@ impl<'a> Qwen3_5GenerateModel<'a> {
let model_list = find_type_files(path, "safetensors")?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
let eos_token_id = cfg.text_config.eos_token_id;
- let qwen3_5 = Qwen3_5Model::new(vb, cfg)?;
+ let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg)?;
Ok(Self {
chat_template,
@@ -53,26 +61,81 @@ impl<'a> Qwen3_5GenerateModel<'a> {
qwen3_5,
device,
eos_token_id,
- model_name: "qwen3.5".to_string(),
+ model_name: model_name.to_string(),
+ repeat_penalty: 1.01,
+ repeat_last_n: 64,
+ })
+ }
+
+ pub fn init_from_gguf(
+ model_file: &str,
+ mmproj_file: Option<&str>,
+ device: Option<&Device>,
+ ) -> Result {
+ if !model_file.contains("Qwen3.5") || !model_file.ends_with("gguf") {
+ return Err(anyhow!("Qwen3.5 gguf model file name illigal {model_file}"));
+ }
+ if let Some(mmproj) = mmproj_file
+ && (!mmproj.contains("mmproj") || !mmproj.ends_with("gguf"))
+ {
+ return Err(anyhow!("Qwen3.5 mmproj_file name illigal {model_file}"));
+ }
+
+ let mut reader = std::fs::File::open(model_file)?;
+ let content = gguf_file::Content::read(&mut reader)?;
+ let device = get_device(device);
+ let mut gguf = Gguf::new(content, reader, device.clone());
+
+ let chat_template_str = gguf
+ .get_matedata("tokenizer.chat_template")?
+ .to_string()?
+ .clone();
+ let chat_template = ChatTemplate::str_init(&chat_template_str)?;
+ let tokenizer = gguf.build_tokenizer(Some(false), Some(false), Some(false))?;
+ let dtype = match gguf.get_matedata("general.type") {
+ Ok(v) => match v.to_u32() {
+ Ok(0) => DType::F32,
+ Ok(1) => DType::F16,
+ _ => DType::F16,
+ },
+ Err(_) => DType::F16,
+ };
+ let pre_processor = Qwen3VLProcessor::new_qwen3_5_default(&device, dtype)?;
+
+ // let eos_token_id = gguf.get_matedata("tokenizer.ggml.eos_token_id")?.to_u32()?;
+ let qwen3_5 = Qwen3_5Model::new_from_gguf(&mut gguf, &device)?;
+ let stem = std::path::Path::new(model_file)
+ .file_stem() // 获取文件名主干(不含扩展名)
+ .and_then(|s| s.to_str())
+ .unwrap_or("qwen3.5");
+ Ok(Self {
+ chat_template,
+ tokenizer,
+ pre_processor,
+ qwen3_5,
+ device,
+ eos_token_id: 248044,
+ model_name: stem.to_string(),
+ repeat_penalty: 1.1,
+ repeat_last_n: 64,
})
}
}
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result {
- let seed = mes.seed.unwrap_or(34562) as u64;
- let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
+ let seed = mes.seed.unwrap_or(32768) as u64;
+ let temperature = mes.temperature.unwrap_or(0.6);
+ let top_p = mes.top_p.unwrap_or(0.95);
+ let mut logit_processor =
+ get_logit_processor(temperature.into(), top_p.into(), Some(20), seed);
+ // let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let mes_render = self.chat_template.apply_chat_template(&mes)?;
- println!("mes_render: {}", mes_render);
- // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // let mes_render = self
- // .chat_template
- // .apply_chat_temp_think(&mes, enable_thinking)?;
+
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self
.tokenizer
.text_encode(input.replace_text.clone(), &self.device)?;
- println!("input_ids: {}", input_ids);
let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0;
@@ -92,6 +155,16 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
seqlen_offset,
)?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
+ let logits = if self.repeat_penalty == 1. {
+ logits
+ } else {
+ let start_at = generate.len().saturating_sub(self.repeat_last_n);
+ candle_transformers::utils::apply_repeat_penalty(
+ &logits,
+ self.repeat_penalty,
+ &generate[start_at..],
+ )?
+ };
let next_token = logit_processor.sample(&logits)?;
generate.push(next_token);
if next_token == self.eos_token_id {
@@ -129,10 +202,6 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let mes_render = self.chat_template.apply_chat_template(&mes)?;
- // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // let mes_render = self
- // .chat_template
- // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self
.tokenizer
@@ -148,6 +217,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
let video_grid_thw = input.video_grid_thw.as_ref();
let mut tool_call_id = None;
let mut tool_call_content = String::new();
+ let mut generate = Vec::new();
for _ in 0..sample_len {
let logits = self.qwen3_5.forward(
&input_ids,
@@ -158,7 +228,18 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
seqlen_offset,
)?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
+ let logits = if self.repeat_penalty == 1. {
+ logits
+ } else {
+ let start_at = generate.len().saturating_sub(self.repeat_last_n);
+ candle_transformers::utils::apply_repeat_penalty(
+ &logits,
+ self.repeat_penalty,
+ &generate[start_at..],
+ )?
+ };
let next_token = logit_processor.sample(&logits)?;
+ generate.push(next_token);
let mut decode_ids = Vec::new();
if !error_tokens.is_empty() {
decode_ids.extend_from_slice(&error_tokens);
diff --git a/src/models/qwen3_5/model.rs b/src/models/qwen3_5/model.rs
index ce86819..1e281d8 100644
--- a/src/models/qwen3_5/model.rs
+++ b/src/models/qwen3_5/model.rs
@@ -1,5 +1,7 @@
+use std::io::{Read, Seek};
+
use anyhow::{Result, anyhow};
-use candle_core::{D, IndexOp, Tensor};
+use candle_core::{D, DType, Device, IndexOp, Tensor, quantized::QMatMul};
use candle_nn::{
Conv1d, Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias,
ops::sigmoid, rms_norm,
@@ -7,7 +9,11 @@ use candle_nn::{
use crate::{
models::{
- common::{GateUpDownMLP, conv1d_depthwise, eager_attention_forward, get_conv1d, softplus},
+ common::{
+ conv1d_depthwise, eager_attention_forward, get_conv1d,
+ gguf::{GateUpDownMLPGguf, Gguf, ProjKind},
+ softplus,
+ },
qwen3_5::config::{Qwen3_5Config, Qwen3_5TextConfig},
qwen3vl::model::Qwen3VLVisionModel,
},
@@ -30,6 +36,10 @@ impl Qwen3_5RMSNorm {
Ok(Self { eps, weight })
}
+ pub fn from_weight(weight: Tensor, eps: f64) -> Result {
+ Ok(Self { eps, weight })
+ }
+
pub fn forward(&self, xs: &Tensor) -> Result {
let x = xs.to_dtype(candle_core::DType::F32)?;
let norm_ = x
@@ -53,6 +63,11 @@ impl Qwen3_5RMSNormGated {
Ok(Self { norm })
}
+ pub fn from_weight(weight: Tensor, eps: f64) -> Result {
+ let norm = RmsNorm::new(weight, eps);
+ Ok(Self { norm })
+ }
+
pub fn forward(&self, xs: &Tensor, gate: Option<&Tensor>) -> Result {
let mut xs = self.norm.forward(xs)?;
if let Some(gate) = gate {
@@ -62,6 +77,7 @@ impl Qwen3_5RMSNormGated {
}
}
+#[macro_export]
macro_rules! transmute_tensors {
($($tensor:expr),*) => {
($(
@@ -69,7 +85,7 @@ macro_rules! transmute_tensors {
)*)
};
}
-
+#[macro_export]
macro_rules! right_pad_zero_tensor {
($dim:expr, $pad_size:expr, $($tensor:expr),+) => {
($(
@@ -78,6 +94,7 @@ macro_rules! right_pad_zero_tensor {
};
}
+#[macro_export]
macro_rules! reshape_chunk_tensor {
($chunk_size:expr, $($tensor:expr),*) => {
($(
@@ -107,30 +124,30 @@ pub struct Qwen3_5GatedDeltaNet {
dt_bias: Tensor,
a_log: Tensor,
norm: Qwen3_5RMSNormGated,
- out_proj: Linear,
+ out_proj: ProjKind,
// Z, B, A 投影
- in_proj_qkv: candle_nn::Linear,
- in_proj_z: candle_nn::Linear,
- in_proj_b: candle_nn::Linear,
- in_proj_a: candle_nn::Linear,
+ in_proj_qkv: ProjKind,
+ in_proj_z: ProjKind,
+ in_proj_b: ProjKind,
+ in_proj_a: ProjKind,
conv_state_cache: Option,
recurrent_state_cache: Option,
}
impl Qwen3_5GatedDeltaNet {
- pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result {
- let hidden_size = config.hidden_size;
- let num_v_heads = config.linear_num_value_heads;
- let num_k_heads = config.linear_num_key_heads;
- let head_k_dim = config.linear_key_head_dim;
- let head_v_dim = config.linear_value_head_dim;
- let key_dim = head_k_dim * num_k_heads;
- let value_dim = head_v_dim * num_v_heads;
- let conv_kernel_size = config.linear_conv_kernel_dim;
+ pub fn new_from_vb(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result {
+ let hidden_size = config.hidden_size; // 1024
+ let num_v_heads = config.linear_num_value_heads; // 16
+ let num_k_heads = config.linear_num_key_heads; // 16
+ let head_k_dim = config.linear_key_head_dim; // 128
+ let head_v_dim = config.linear_value_head_dim; // 128
+ let key_dim = head_k_dim * num_k_heads; // 2048
+ let value_dim = head_v_dim * num_v_heads; // 2048
+ let conv_kernel_size = config.linear_conv_kernel_dim; // 4
// let activation = config.hidden_act;
let layer_norm_epsilon = config.rms_norm_eps;
- let conv_dim = key_dim * 2 + value_dim;
+ let conv_dim = key_dim * 2 + value_dim; // 6144
let conv1d = get_conv1d(
vb.pp("conv1d"),
conv_dim,
@@ -146,12 +163,17 @@ impl Qwen3_5GatedDeltaNet {
let a_log = vb.get(num_v_heads, "A_log")?;
let norm = Qwen3_5RMSNormGated::new(vb.pp("norm"), head_v_dim, layer_norm_epsilon)?;
+ // 2048, 1024
let out_proj = linear_no_bias(value_dim, hidden_size, vb.pp("out_proj"))?;
-
+ // 1024, 6144
let in_proj_qkv = linear_no_bias(hidden_size, conv_dim, vb.pp("in_proj_qkv"))?;
+ // 1024, 2048
let in_proj_z = linear_no_bias(hidden_size, value_dim, vb.pp("in_proj_z"))?;
+ // 1024, 16
let in_proj_b = linear_no_bias(hidden_size, num_v_heads, vb.pp("in_proj_b"))?;
+ // 1024, 16
let in_proj_a = linear_no_bias(hidden_size, num_v_heads, vb.pp("in_proj_a"))?;
+
Ok(Self {
// hidden_size,
num_v_heads,
@@ -168,11 +190,64 @@ impl Qwen3_5GatedDeltaNet {
dt_bias,
a_log,
norm,
- out_proj,
- in_proj_qkv,
- in_proj_z,
- in_proj_b,
- in_proj_a,
+ out_proj: ProjKind::LinearProj(out_proj),
+ in_proj_qkv: ProjKind::LinearProj(in_proj_qkv),
+ in_proj_z: ProjKind::LinearProj(in_proj_z),
+ in_proj_b: ProjKind::LinearProj(in_proj_b),
+ in_proj_a: ProjKind::LinearProj(in_proj_a),
+ conv_state_cache: None,
+ recurrent_state_cache: None,
+ })
+ }
+
+ pub fn new_from_gguf(
+ gguf: &mut Gguf,
+ prefix: &str,
+ rms_norm_eps: f64,
+ ) -> Result {
+ let num_k_heads = gguf.get_matedata("qwen35.ssm.group_count")?.to_u32()? as usize;
+ let num_v_heads = gguf.get_matedata("qwen35.ssm.time_step_rank")?.to_u32()? as usize;
+ let conv_kernel_size = gguf.get_matedata("qwen35.ssm.conv_kernel")?.to_u32()? as usize;
+ let head_k_dim = gguf.get_matedata("qwen35.ssm.state_size")?.to_u32()? as usize;
+ let head_v_dim = head_k_dim;
+ let key_dim = head_k_dim * num_k_heads;
+ let value_dim = head_v_dim * num_v_heads;
+ let conv_dim = key_dim * 2 + value_dim;
+ let conv1d = gguf.conv1d(
+ &format!("{prefix}.ssm_conv1d"),
+ conv_kernel_size - 1,
+ 1,
+ 1,
+ conv_dim,
+ false,
+ )?;
+ let dt_bias = gguf.get_dequantized(&format!("{prefix}.ssm_dt.bias"))?;
+ let a_log = gguf.get_dequantized(&format!("{prefix}.ssm_a"))?;
+ let norm_weight = gguf.get_dequantized(&format!("{prefix}.ssm_norm.weight"))?;
+ let norm = Qwen3_5RMSNormGated::from_weight(norm_weight, rms_norm_eps)?;
+ let out_proj = gguf.qmatmul(&format!("{prefix}.ssm_out.weight"))?;
+ let in_proj_qkv = gguf.qmatmul(&format!("{prefix}.attn_qkv.weight"))?;
+ let in_proj_z = gguf.qmatmul(&format!("{prefix}.attn_gate.weight"))?;
+ let in_proj_b = gguf.qmatmul(&format!("{prefix}.ssm_beta.weight"))?;
+ let in_proj_a = gguf.qmatmul(&format!("{prefix}.ssm_alpha.weight"))?;
+
+ Ok(Self {
+ num_v_heads,
+ num_k_heads,
+ head_k_dim,
+ head_v_dim,
+ key_dim,
+ value_dim,
+ conv_kernel_size,
+ conv1d,
+ dt_bias,
+ a_log,
+ norm,
+ out_proj: ProjKind::QuantizedProj(out_proj),
+ in_proj_qkv: ProjKind::QuantizedProj(in_proj_qkv),
+ in_proj_z: ProjKind::QuantizedProj(in_proj_z),
+ in_proj_b: ProjKind::QuantizedProj(in_proj_b),
+ in_proj_a: ProjKind::QuantizedProj(in_proj_a),
conv_state_cache: None,
recurrent_state_cache: None,
})
@@ -466,7 +541,6 @@ impl Qwen3_5GatedDeltaNet {
&[self.key_dim, self.key_dim, self.value_dim],
D::Minus1,
)?;
-
let mut query = qkv_split[0].reshape((bs, seq_len, (), self.head_k_dim))?;
let mut key = qkv_split[1].reshape((bs, seq_len, (), self.head_k_dim))?;
let value = qkv_split[2].reshape((bs, seq_len, (), self.head_v_dim))?;
@@ -502,10 +576,10 @@ impl Qwen3_5GatedDeltaNet {
}
pub struct Qwen3_5Attention {
- q_proj: Linear,
- k_proj: Linear,
- v_proj: Linear,
- o_proj: Linear,
+ q_proj: ProjKind,
+ k_proj: ProjKind,
+ v_proj: ProjKind,
+ o_proj: ProjKind,
q_norm: Qwen3_5RMSNorm,
k_norm: Qwen3_5RMSNorm,
num_attention_heads: usize,
@@ -517,7 +591,7 @@ pub struct Qwen3_5Attention {
}
impl Qwen3_5Attention {
- pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result {
+ pub fn new_from_vb(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result {
let hidden_size = config.hidden_size;
let num_attention_heads = config.num_attention_heads;
let head_dim = config.head_dim;
@@ -550,11 +624,50 @@ impl Qwen3_5Attention {
)?;
let q_norm = Qwen3_5RMSNorm::new(vb.pp("q_norm"), head_dim, config.rms_norm_eps)?;
let k_norm = Qwen3_5RMSNorm::new(vb.pp("k_norm"), head_dim, config.rms_norm_eps)?;
+
Ok(Self {
- q_proj,
- k_proj,
- v_proj,
- o_proj,
+ q_proj: ProjKind::LinearProj(q_proj),
+ k_proj: ProjKind::LinearProj(k_proj),
+ v_proj: ProjKind::LinearProj(v_proj),
+ o_proj: ProjKind::LinearProj(o_proj),
+ q_norm,
+ k_norm,
+ num_attention_heads,
+ num_key_value_heads,
+ num_kv_groups,
+ head_dim,
+ scaling,
+ kv_cache: None,
+ })
+ }
+
+ pub fn new_from_gguf(
+ gguf: &mut Gguf,
+ prefix: &str,
+ rms_norm_eps: f64,
+ ) -> Result {
+ let num_attention_heads =
+ gguf.get_matedata("qwen35.attention.head_count")?.to_u32()? as usize;
+ let num_key_value_heads = gguf
+ .get_matedata("qwen35.attention.head_count_kv")?
+ .to_u32()? as usize;
+ let num_kv_groups = num_attention_heads / num_key_value_heads;
+ let head_dim = gguf.get_matedata("qwen35.attention.key_length")?.to_u32()? as usize;
+ let scaling = 1f64 / f64::sqrt(head_dim as f64);
+ let q_proj = gguf.qmatmul(&format!("{prefix}.attn_q.weight"))?;
+ let k_proj = gguf.qmatmul(&format!("{prefix}.attn_k.weight"))?;
+ let v_proj = gguf.qmatmul(&format!("{prefix}.attn_v.weight"))?;
+ let o_proj = gguf.qmatmul(&format!("{prefix}.attn_output.weight"))?;
+ let q_norm_weight = gguf.get_dequantized(&format!("{prefix}.attn_q_norm.weight"))?;
+ let q_norm = Qwen3_5RMSNorm::from_weight(q_norm_weight, rms_norm_eps)?;
+ let k_norm_weight = gguf.get_dequantized(&format!("{prefix}.attn_k_norm.weight"))?;
+ let k_norm = Qwen3_5RMSNorm::from_weight(k_norm_weight, rms_norm_eps)?;
+
+ Ok(Self {
+ q_proj: ProjKind::QuantizedProj(q_proj),
+ k_proj: ProjKind::QuantizedProj(k_proj),
+ v_proj: ProjKind::QuantizedProj(v_proj),
+ o_proj: ProjKind::QuantizedProj(o_proj),
q_norm,
k_norm,
num_attention_heads,
@@ -627,32 +740,62 @@ impl Qwen3_5Attention {
}
}
+enum AttnKind {
+ LinearAttn(Qwen3_5GatedDeltaNet),
+ SelfAttn(Qwen3_5Attention),
+}
+
+impl AttnKind {
+ fn forward(
+ &mut self,
+ xs: &Tensor,
+ cos: Option<&Tensor>,
+ sin: Option<&Tensor>,
+ attention_mask: Option<&Tensor>,
+ ) -> Result {
+ match self {
+ AttnKind::LinearAttn(attn) => attn.forward(xs, attention_mask),
+ AttnKind::SelfAttn(attn) => {
+ if let Some(cos) = cos
+ && let Some(sin) = sin
+ {
+ attn.forward(xs, cos, sin, attention_mask)
+ } else {
+ Err(anyhow!("Qwen3_5 self attn cos and sin is all need"))
+ }
+ }
+ }
+ }
+}
+
pub struct Qwen3_5DecoderLayer {
// hidden_size: usize,
layer_type: String,
- linear_attn: Option,
- self_attn: Option,
- mlp: GateUpDownMLP,
+ attn: AttnKind,
+ mlp: GateUpDownMLPGguf,
input_layernorm: Qwen3_5RMSNorm,
post_attention_layernorm: Qwen3_5RMSNorm,
}
impl Qwen3_5DecoderLayer {
- pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig, layer_idx: usize) -> Result {
+ pub fn new_from_vb(
+ vb: VarBuilder,
+ config: &Qwen3_5TextConfig,
+ layer_idx: usize,
+ ) -> Result {
let hidden_size = config.hidden_size;
let layer_type = config.layer_types[layer_idx].clone();
- let (linear_attn, self_attn) = if layer_type.eq("linear_attention") {
- let linear_attn = Qwen3_5GatedDeltaNet::new(vb.pp("linear_attn"), config)?;
- (Some(linear_attn), None)
+ let attn = if layer_type.eq("linear_attention") {
+ let attn = Qwen3_5GatedDeltaNet::new_from_vb(vb.pp("linear_attn"), config)?;
+ AttnKind::LinearAttn(attn)
} else {
- let self_attn = Qwen3_5Attention::new(vb.pp("self_attn"), config)?;
- (None, Some(self_attn))
+ let attn = Qwen3_5Attention::new_from_vb(vb.pp("self_attn"), config)?;
+ AttnKind::SelfAttn(attn)
};
- let mlp = GateUpDownMLP::new(
+ let mlp = GateUpDownMLPGguf::new_from_vb(
vb.pp("mlp"),
hidden_size,
config.intermediate_size,
- config.hidden_act,
false,
None,
None,
@@ -668,8 +811,36 @@ impl Qwen3_5DecoderLayer {
Ok(Self {
// hidden_size,
layer_type,
- linear_attn,
- self_attn,
+ attn,
+ mlp,
+ input_layernorm,
+ post_attention_layernorm,
+ })
+ }
+
+ pub fn new_from_gguf(
+ gguf: &mut Gguf,
+ prefix: &str,
+ layer_type: &str,
+ rms_norm_eps: f64,
+ ) -> Result {
+ let attn = if layer_type.eq("linear_attention") {
+ let attn = Qwen3_5GatedDeltaNet::new_from_gguf(gguf, prefix, rms_norm_eps)?;
+ AttnKind::LinearAttn(attn)
+ } else {
+ let attn = Qwen3_5Attention::new_from_gguf(gguf, prefix, rms_norm_eps)?;
+ AttnKind::SelfAttn(attn)
+ };
+ let mlp = GateUpDownMLPGguf::new_from_gguf(gguf, prefix)?;
+ let input_norm_weight = gguf.get_dequantized(&format!("{prefix}.attn_norm.weight"))?;
+ let input_layernorm = Qwen3_5RMSNorm::from_weight(input_norm_weight, rms_norm_eps)?;
+ let post_norm_weight =
+ gguf.get_dequantized(&format!("{prefix}.post_attention_norm.weight"))?;
+ let post_attention_layernorm = Qwen3_5RMSNorm::from_weight(post_norm_weight, rms_norm_eps)?;
+ Ok(Self {
+ // hidden_size,
+ layer_type: layer_type.to_string(),
+ attn,
mlp,
input_layernorm,
post_attention_layernorm,
@@ -685,16 +856,17 @@ impl Qwen3_5DecoderLayer {
) -> Result {
let residual = xs.clone();
let mut xs = self.input_layernorm.forward(xs)?;
- if self.layer_type.eq("linear_attention")
- && let Some(linear_attn) = self.linear_attn.as_mut()
- {
- xs = linear_attn.forward(&xs, attention_mask)?;
- } else if let Some(self_attn) = self.self_attn.as_mut()
- && let Some(cos) = cos
- && let Some(sin) = sin
- {
- xs = self_attn.forward(&xs, cos, sin, attention_mask)?;
- }
+ xs = self.attn.forward(&xs, cos, sin, attention_mask)?;
+ // if self.layer_type.eq("linear_attention")
+ // && let Some(linear_attn) = self.linear_attn.as_mut()
+ // {
+ // xs = linear_attn.forward(&xs, attention_mask)?;
+ // } else if let Some(self_attn) = self.self_attn.as_mut()
+ // && let Some(cos) = cos
+ // && let Some(sin) = sin
+ // {
+ // xs = self_attn.forward(&xs, cos, sin, attention_mask)?;
+ // }
let residual = xs.add(&residual)?;
xs = self.post_attention_layernorm.forward(&residual)?;
xs = self.mlp.forward(&xs)?;
@@ -703,11 +875,13 @@ impl Qwen3_5DecoderLayer {
}
pub fn clear_cache(&mut self) {
- if let Some(linear_attn) = self.linear_attn.as_mut() {
- linear_attn.clear_cache();
- }
- if let Some(self_attn) = self.self_attn.as_mut() {
- self_attn.clear_kv_cache();
+ match &mut self.attn {
+ AttnKind::LinearAttn(attn) => {
+ attn.clear_cache();
+ }
+ AttnKind::SelfAttn(attn) => {
+ attn.clear_kv_cache();
+ }
}
}
}
@@ -718,15 +892,17 @@ pub struct Qwen3_5TextModel {
norm: Qwen3_5RMSNorm,
rotary_emb: Qwen3VLTextRotaryEmbedding,
mrope_section: Vec,
+ dtype: DType,
}
impl Qwen3_5TextModel {
- pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result {
+ pub fn new_from_vb(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result {
let embed_tokens = embedding(config.vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
let mut layers = vec![];
let vb_layers = vb.pp("layers");
for i in 0..config.num_hidden_layers {
- let layer = Qwen3_5DecoderLayer::new(vb_layers.pp(i), config, i)?;
+ // for i in 0..4 {
+ let layer = Qwen3_5DecoderLayer::new_from_vb(vb_layers.pp(i), config, i)?;
layers.push(layer);
}
let norm = Qwen3_5RMSNorm::new(vb.pp("norm"), config.hidden_size, config.rms_norm_eps)?;
@@ -740,17 +916,70 @@ impl Qwen3_5TextModel {
norm,
rotary_emb,
mrope_section: config.rope_parameters.mrope_section.clone(),
+ dtype: vb.dtype(),
+ })
+ }
+ pub fn new_from_gguf(gguf: &mut Gguf, device: &Device) -> Result {
+ let num_layers = gguf.get_matedata("qwen35.block_count")?.to_u32()? as usize;
+ let full_attention_interval = gguf
+ .get_matedata("qwen35.full_attention_interval")?
+ .to_u32()? as usize;
+ let rope_freq_base = gguf.get_matedata("qwen35.rope.freq_base")?.to_f32()?;
+ let rope_dimension_count =
+ gguf.get_matedata("qwen35.rope.dimension_count")?.to_u32()? as usize;
+ let mut mrope_section = gguf
+ .get_matedata("qwen35.rope.dimension_sections")?
+ .to_vec()?
+ .iter()
+ .map(|v| v.to_i32().map(|x| x as usize))
+ .collect::, candle_core::Error>>()?;
+ let _ = mrope_section.pop();
+ let rms_norm_eps = gguf
+ .get_matedata("qwen35.attention.layer_norm_rms_epsilon")?
+ .to_f32()? as f64;
+ let hidden_size = gguf.get_matedata("qwen35.embedding_length")?.to_u32()? as usize; // 1024
+ let embed_tensor = gguf.tensor("token_embd.weight")?;
+ let embed_tokens = Embedding::new(embed_tensor.dequantize(device)?, hidden_size);
+ let mut layers = vec![];
+ for i in 0..num_layers {
+ // for i in 0..4 {
+ let prefix = format!("blk.{i}");
+ let layer_type = if (i + 1) % full_attention_interval == 0 {
+ "full_attention".to_string()
+ } else {
+ "linear_attention".to_string()
+ };
+ let layer =
+ Qwen3_5DecoderLayer::new_from_gguf(gguf, &prefix, &layer_type, rms_norm_eps)?;
+ layers.push(layer);
+ }
+ let norm_weight = gguf.get_dequantized("output_norm.weight")?;
+ let norm = Qwen3_5RMSNorm::from_weight(norm_weight, rms_norm_eps)?;
+ let rotary_emb = Qwen3VLTextRotaryEmbedding::new(rope_dimension_count, rope_freq_base);
+ let dtype = match gguf.get_matedata("general.type") {
+ Ok(v) => match v.to_u32() {
+ Ok(0) => DType::F32,
+ Ok(1) => DType::F16,
+ _ => DType::F16,
+ },
+ Err(_) => DType::F16,
+ };
+ Ok(Self {
+ embed_tokens,
+ layers,
+ norm,
+ rotary_emb,
+ mrope_section,
+ dtype,
})
}
pub fn forward(&mut self, inputs_embeds: &Tensor, position_ids: &Tensor) -> Result {
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
- let (cos, sin) = self.rotary_emb.forward(
- position_ids,
- inputs_embeds.dtype(),
- self.mrope_section.clone(),
- )?;
+ let (cos, sin) =
+ self.rotary_emb
+ .forward(position_ids, self.dtype, self.mrope_section.clone())?;
let mut xs = inputs_embeds.clone();
let attention_mask: Option = {
if seq_len <= 1 {
@@ -785,18 +1014,23 @@ impl Qwen3_5TextModel {
}
pub struct Qwen3_5Model {
- config: Qwen3_5Config,
- visual: Qwen3VLVisionModel,
+ // config: Qwen3_5Config,
+ spatial_merge_size: usize,
+ image_token_id: u32,
+ video_token_id: u32,
+ vision_start_token_id: u32,
+ visual: Option,
language_model: Qwen3_5TextModel,
- lm_head: Linear,
+ lm_head: ProjKind,
rope_deltas: Option,
}
impl Qwen3_5Model {
- pub fn new(vb: VarBuilder, config: Qwen3_5Config) -> Result {
+ pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config) -> Result {
let vb_m = vb.pp("model");
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
- let language_model = Qwen3_5TextModel::new(vb_m.pp("language_model"), &config.text_config)?;
+ let language_model =
+ Qwen3_5TextModel::new_from_vb(vb_m.pp("language_model"), &config.text_config)?;
let lm_head = if config.tie_word_embeddings {
Linear::new(language_model.embed_tokens.embeddings().clone(), None)
} else {
@@ -807,10 +1041,36 @@ impl Qwen3_5Model {
)?
};
Ok(Self {
- config,
- visual,
+ spatial_merge_size: config.vision_config.spatial_merge_size,
+ image_token_id: config.image_token_id,
+ video_token_id: config.video_token_id,
+ vision_start_token_id: config.vision_start_token_id,
+ visual: Some(visual),
language_model,
- lm_head,
+ lm_head: ProjKind::LinearProj(lm_head),
+ rope_deltas: None,
+ })
+ }
+
+ pub fn new_from_gguf(gguf: &mut Gguf, device: &Device) -> Result {
+ let spatial_merge_size = 2usize;
+ let image_token_id = 248056u32;
+ let video_token_id = 248057u32;
+ let vision_start_token_id = 248053u32;
+ let language_model = Qwen3_5TextModel::new_from_gguf(gguf, device)?;
+ let lm_head_tensor = match gguf.tensor("output.weight") {
+ Ok(tensor) => tensor,
+ Err(_) => gguf.tensor("token_embd.weight")?,
+ };
+ let lm_head = QMatMul::from_qtensor(lm_head_tensor)?;
+ Ok(Self {
+ spatial_merge_size,
+ image_token_id,
+ video_token_id,
+ vision_start_token_id,
+ visual: None,
+ language_model,
+ lm_head: ProjKind::QuantizedProj(lm_head),
rope_deltas: None,
})
}
@@ -842,10 +1102,10 @@ impl Qwen3_5Model {
None => None,
};
- let spatial_merge_size = self.config.vision_config.spatial_merge_size;
- let image_token_id = self.config.image_token_id;
- let video_token_id = self.config.video_token_id;
- let vision_start_token_id = self.config.vision_start_token_id;
+ let spatial_merge_size = self.spatial_merge_size;
+ let image_token_id = self.image_token_id;
+ let video_token_id = self.video_token_id;
+ let vision_start_token_id = self.vision_start_token_id;
let mut mrope_position_deltas = vec![];
if image_grid_thw.is_some() || video_grid_thw.is_some() {
let total_input_ids = input_ids.clone();
@@ -1093,9 +1353,10 @@ impl Qwen3_5Model {
let mut inputs_embeds = self.language_model.embed_tokens.forward(input_ids)?;
if let Some(pixel_values) = pixel_values
&& let Some(image_grid_thw) = image_grid_thw
+ && let Some(visual) = self.visual.as_ref()
{
- let (image_embeds, _) = self.visual.forward(pixel_values, image_grid_thw)?;
- let vision_mask = get_equal_mask(input_ids, self.config.image_token_id)?;
+ let (image_embeds, _) = visual.forward(pixel_values, image_grid_thw)?;
+ let vision_mask = get_equal_mask(input_ids, self.image_token_id)?;
let n_image_tokens = vision_mask.sum_all()?.to_scalar::()?;
if n_image_tokens as usize != image_embeds.dim(0)? {
return Err(anyhow!(format!(
@@ -1108,9 +1369,10 @@ impl Qwen3_5Model {
}
if let Some(pixel_values_video) = pixel_values_video
&& let Some(video_grid_thw) = video_grid_thw
+ && let Some(visual) = self.visual.as_ref()
{
- let (video_embeds, _) = self.visual.forward(pixel_values_video, video_grid_thw)?;
- let vision_mask = get_equal_mask(input_ids, self.config.video_token_id)?;
+ let (video_embeds, _) = visual.forward(pixel_values_video, video_grid_thw)?;
+ let vision_mask = get_equal_mask(input_ids, self.video_token_id)?;
let n_video_tokens = vision_mask.sum_all()?.to_scalar::()?;
if n_video_tokens as usize != video_embeds.dim(0)? {
return Err(anyhow!(format!(
diff --git a/src/models/qwen3vl/config.rs b/src/models/qwen3vl/config.rs
index 91ee0d8..4e2f104 100644
--- a/src/models/qwen3vl/config.rs
+++ b/src/models/qwen3vl/config.rs
@@ -18,6 +18,36 @@ pub struct PreprocessorConfig {
pub image_std: Vec,
}
+impl PreprocessorConfig {
+ pub fn qwen3_5_img_default() -> Self {
+ Self {
+ size: Size {
+ longest_edge: 16777216,
+ shortest_edge: 65536,
+ },
+ patch_size: 16,
+ temporal_patch_size: 2,
+ merge_size: 2,
+ image_mean: vec![0.5, 0.5, 0.5],
+ image_std: vec![0.5, 0.5, 0.5],
+ }
+ }
+
+ pub fn qwen3_5_video_default() -> Self {
+ Self {
+ size: Size {
+ longest_edge: 25165824,
+ shortest_edge: 4096,
+ },
+ patch_size: 16,
+ temporal_patch_size: 2,
+ merge_size: 2,
+ image_mean: vec![0.5, 0.5, 0.5],
+ image_std: vec![0.5, 0.5, 0.5],
+ }
+ }
+}
+
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct RopeScaling {
pub rope_type: String,
diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs
index 032ce84..b29f660 100644
--- a/src/models/qwen3vl/generate.rs
+++ b/src/models/qwen3vl/generate.rs
@@ -16,8 +16,8 @@ use crate::{
},
tokenizer::TokenizerModel,
utils::{
- build_completion_chunk_response, build_completion_response, extract_metadata_value,
- find_type_files, get_device, get_dtype, get_logit_processor,
+ build_completion_chunk_response, build_completion_response, find_type_files, get_device,
+ get_dtype, get_logit_processor,
},
};
@@ -143,10 +143,6 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let mes_render = self.chat_template.apply_chat_template(&mes)?;
- // let enable_thinking = extract_metadata_value::(&mes.metadata, "enable_thinking");
- // let mes_render = self
- // .chat_template
- // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self
.tokenizer
diff --git a/src/models/qwen3vl/processor.rs b/src/models/qwen3vl/processor.rs
index 613f409..509bb70 100644
--- a/src/models/qwen3vl/processor.rs
+++ b/src/models/qwen3vl/processor.rs
@@ -99,6 +99,28 @@ impl Qwen3VLProcessor {
})
}
+ pub fn new_qwen3_5_default(device: &Device, dtype: DType) -> Result {
+ let img_process_cfg = PreprocessorConfig::qwen3_5_img_default();
+ let video_process_cfg = PreprocessorConfig::qwen3_5_video_default();
+ let image_token = "<|image_pad|>".to_string();
+ let video_token = "<|video_pad|>".to_string();
+ let vision_start_token = "<|vision_start|>".to_string();
+ let vision_end_token = "<|vision_end|>".to_string();
+ Ok(Self {
+ img_process_cfg,
+ video_process_cfg,
+ device: device.clone(),
+ dtype,
+ image_token,
+ video_token,
+ vision_start_token,
+ vision_end_token,
+ fps: 2,
+ min_frames: 4,
+ max_frames: 768,
+ })
+ }
+
pub fn extract_vision_info(
&self,
mes: &ChatCompletionParameters,
diff --git a/src/tokenizer/mod.rs b/src/tokenizer/mod.rs
index fd8fdbc..1e1d8e3 100644
--- a/src/tokenizer/mod.rs
+++ b/src/tokenizer/mod.rs
@@ -85,8 +85,6 @@ impl TokenizerModel {
}
tokenizer
};
- let len = tokenizer.get_vocab_size(true);
- println!("len: {}", len);
Ok(Self { tokenizer })
}
diff --git a/tests/test_gguf_qwen3_5.rs b/tests/test_gguf_qwen3_5.rs
index 236b904..03d4f96 100644
--- a/tests/test_gguf_qwen3_5.rs
+++ b/tests/test_gguf_qwen3_5.rs
@@ -1,18 +1,32 @@
-use aha::{chat::ChatCompletionParameters, gguf_models::qwen3_5::GgufQwen3_5};
+use std::time::Instant;
+
+use aha::{
+ chat::ChatCompletionParameters,
+ models::{GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel},
+};
use anyhow::Result;
-use candle_core::{Device, quantized::gguf_file};
+// use candle_core::{Device, quantized::gguf_file};
#[test]
fn gguf_test() -> Result<()> {
// cargo test -r -F cuda --test test_gguf_qwen3_5 gguf_test -- --nocapture
+ // let path = "/home/jhq/.aha/Qwen/Qwen3.5-4B-GGUF/Qwen3.5-4B-Q5_K_M.gguf"; // 有问题
+ // let path = "/home/jhq/.aha/Qwen/Qwen3.5-2B-GGUF/Qwen3.5-2B-Q6_K.gguf";
let path = "/home/jhq/.aha/Qwen/Qwen3.5-0.8B-GGUF/Qwen3.5-0.8B-Q4_K_M.gguf";
- let mut file = std::fs::File::open(path)?;
- let model = gguf_file::Content::read(&mut file)?;
- let device = Device::new_cuda(0)?;
+ // let mut file = std::fs::File::open(path)?;
+ // let model = gguf_file::Content::read(&mut file)?;
+ // println!("group_count: {:?}", model.metadata.get("qwen35.ssm.group_count"));
+ // println!("time_step_rank: {:?}", model.metadata.get("qwen35.ssm.time_step_rank"));
+ // println!("state_size: {:?}", model.metadata.get("qwen35.ssm.state_size"));
+ // for (key, value) in model.metadata {
+ // if key.contains("tokenizer") {
+ // continue;
+ // }
+ // println!("{key}: {:#?}", value);
+ // }
// println!("model: {:?}", model.magic);
- // println!("generat.type: {:#?}", model.metadata.keys());
- // println!("tokenizer.ggml.model: {:#?}", model.metadata.get("tokenizer.ggml.model")); // gpt2
- // // println!("model: {:?}", model.tensor_infos);
-
+ // println!("generat.type: {:#?}", model.metadata.keys());
+ // println!("tokenizer.ggml.eos_token_id: {:#?}", model.metadata.get("tokenizer.ggml.eos_token_id"));
+ // println!("model: {:#?}", model.tensor_infos.keys());
let message = r#"
{
"model": "qwen3.5",
@@ -22,26 +36,29 @@ fn gguf_test() -> Result<()> {
"content": [
{
"type": "text",
- "text": "你好啊"
+ "text": "你如何看待AI"
}
]
}
]
}
"#;
-// render: <|im_start|>user
-// 你好啊<|im_end|>
-// <|im_start|>assistant
-//
-
-//
-
-
-// input_ids: [[248045, 846, 198, 109266, 98710, 248046, 198, 248045, 74455, 198,
-// 248068, 271, 248069, 271]]
-// Tensor[[1, 14], u32, cuda:0]
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
- let mut gguf_qwen3_5 = GgufQwen3_5::new(&path, None)?;
- let _ = gguf_qwen3_5.generate(mes)?;
+ let i_start = Instant::now();
+ let mut gguf_qwen3_5 = Qwen3_5GenerateModel::init_from_gguf(path, None, None)?;
+ let i_duration = i_start.elapsed();
+ println!("Time elapsed in load model is: {:?}", i_duration);
+
+ let i_start = Instant::now();
+ let res = gguf_qwen3_5.generate(mes)?;
+ let i_duration = i_start.elapsed();
+ println!("generate: \n {:?}", res);
+ if res.usage.is_some() {
+ let num_token = res.usage.as_ref().unwrap().total_tokens;
+ let duration_secs = i_duration.as_secs_f64();
+ let tps = num_token as f64 / duration_secs;
+ println!("Tokens per second (TPS): {:.2}", tps);
+ }
+ println!("Time elapsed in generate is: {:?}", i_duration);
Ok(())
-}
\ No newline at end of file
+}
diff --git a/tests/test_qwen3_5.rs b/tests/test_qwen3_5.rs
index c8b108d..fb98621 100644
--- a/tests/test_qwen3_5.rs
+++ b/tests/test_qwen3_5.rs
@@ -22,7 +22,7 @@ fn qwen3_5_generate() -> Result<()> {
"content": [
{
"type": "text",
- "text": "你好啊"
+ "text": "你好啊,你是谁"
}
]
}
@@ -31,12 +31,12 @@ fn qwen3_5_generate() -> Result<()> {
"#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
let i_start = Instant::now();
- let mut qwen3vl = Qwen3_5GenerateModel::init(&model_path, None, None)?;
+ let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?;
let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration);
let i_start = Instant::now();
- let res = qwen3vl.generate(mes)?;
+ let res = qwen3_5.generate(mes)?;
let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res);
if res.usage.is_some() {