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] 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)?;