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() {