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