from gguf build chat_template and tokenizer

This commit is contained in:
jhqxxx
2026-03-11 17:30:28 +08:00
parent 86141322bb
commit 4d544ab68f
14 changed files with 374 additions and 118 deletions
Generated
+1
View File
@@ -33,6 +33,7 @@ name = "aha"
version = "0.2.2" version = "0.2.2"
dependencies = [ dependencies = [
"aha_openai_dive", "aha_openai_dive",
"ahash",
"anyhow", "anyhow",
"base64 0.22.1", "base64 0.22.1",
"byteorder", "byteorder",
+1
View File
@@ -42,6 +42,7 @@ zip = "7.2.0"
half = "2.7.1" half = "2.7.1"
byteorder = "1.5.0" byteorder = "1.5.0"
sentencepiece = "0.13.1" sentencepiece = "0.13.1"
ahash = "0.8.12"
[features] [features]
flash-attn = ["candle-flash-attn"] flash-attn = ["candle-flash-attn"]
+55 -72
View File
@@ -2,7 +2,35 @@ use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use anyhow::{Result, anyhow}; use anyhow::{Result, anyhow};
use minijinja::{Environment, Value as MiniJinjaValue, context}; 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('<tool_response>')",
"content is startingwith('<tool_response>')", // 使用minijinja中的 is startingwith 替换
)
.replace(
"content.endswith('</tool_response>')",
"content is endingwith('</tool_response>')", // 使用minijinja中的 is endingwith 替换
)
.replace(
"content.split('</think>')[0].rstrip('\\n').split('<think>')[-1].lstrip('\\n')",
"((content | split('</think>'))[0] | rstrip('\\n') | split('<think>'))[-1] | lstrip('\\n')", // 使用自定义的split, rstrip, lstrip过滤器替换
)
.replace(
"content.split('</think>')[-1].lstrip('\\n')",
"(content | split('</think>'))[-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<String> { pub fn get_template(path: String) -> Result<String> {
let tokenizer_config_file = path.clone() + "/tokenizer_config.json"; let tokenizer_config_file = path.clone() + "/tokenizer_config.json";
@@ -47,31 +75,7 @@ pub fn get_template(path: String) -> Result<String> {
}; };
let chat_template = chat_template.ok_or(anyhow!(format!("chat_template is none")))?; let chat_template = chat_template.ok_or(anyhow!(format!("chat_template is none")))?;
// 修复模板中的问题行 // 修复模板中的问题行
let fixed_template = chat_template let fixed_template = fix_template(&chat_template);
.replace(
"content.startswith('<tool_response>')",
"content is startingwith('<tool_response>')", // 使用minijinja中的 is startingwith 替换
)
.replace(
"content.endswith('</tool_response>')",
"content is endingwith('</tool_response>')", // 使用minijinja中的 is endingwith 替换
)
.replace(
"content.split('</think>')[0].rstrip('\\n').split('<think>')[-1].lstrip('\\n')",
"((content | split('</think>'))[0] | rstrip('\\n') | split('<think>'))[-1] | lstrip('\\n')", // 使用自定义的split, rstrip, lstrip过滤器替换
)
.replace(
"content.split('</think>')[-1].lstrip('\\n')",
"(content | split('</think>'))[-1] | lstrip('\\n')", // 使用自定义的过滤器替换
)
.replace(
"reasoning_content.strip('\\n')",
"reasoning_content | strip('\\n')", // 使用自定义的过滤器替换
)
.replace(
"content.lstrip('\\n')",
"content | lstrip('\\n')", // 使用自定义的过滤器替换
);
Ok(fixed_template) Ok(fixed_template)
} }
@@ -80,29 +84,7 @@ pub struct ChatTemplate<'a> {
} }
impl<'a> ChatTemplate<'a> { impl<'a> ChatTemplate<'a> {
pub fn init(path: &str) -> Result<Self> { fn setup_environment(env: &mut Environment<'a>) {
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();
// 添加自定义过滤器
env.add_filter("tojson", |v: MiniJinjaValue| { env.add_filter("tojson", |v: MiniJinjaValue| {
serde_json::to_string(&v).unwrap() serde_json::to_string(&v).unwrap()
}); });
@@ -113,44 +95,44 @@ impl<'a> ChatTemplate<'a> {
.collect::<Vec<String>>() .collect::<Vec<String>>()
}); });
// 添加 lstrip 过滤器
env.add_filter("lstrip", |s: String, chars: Option<String>| match chars { env.add_filter("lstrip", |s: String, chars: Option<String>| match chars {
Some(chars_str) => s.trim_start_matches(chars_str.as_str()).to_string(), Some(chars_str) => s.trim_start_matches(chars_str.as_str()).to_string(),
None => s.trim_start().to_string(), None => s.trim_start().to_string(),
}); });
// 添加 rstrip 过滤器
env.add_filter("rstrip", |s: String, chars: Option<String>| match chars { env.add_filter("rstrip", |s: String, chars: Option<String>| match chars {
Some(chars_str) => s.trim_end_matches(chars_str.as_str()).to_string(), Some(chars_str) => s.trim_end_matches(chars_str.as_str()).to_string(),
None => s.trim_end().to_string(), None => s.trim_end().to_string(),
}); });
// let template = get_template(path.to_string())?; }
pub fn init(path: &str) -> Result<Self> {
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<Self> {
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); let _ = env.add_template("chat", template);
Ok(Self { env }) Ok(Self { env })
} }
pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result<String> { pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result<String> {
let context = context! { let enable_thinking = extract_metadata_value::<bool>(&messages.metadata, "enable_thinking");
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<bool>,
) -> Result<String> {
let context = context! { let context = context! {
messages => &messages.messages, messages => &messages.messages,
tools => &messages.tools.as_ref(), tools => &messages.tools.as_ref(),
@@ -166,4 +148,5 @@ impl<'a> ChatTemplate<'a> {
.map_err(|e| anyhow!(format!("render template error {}", e)))?; .map_err(|e| anyhow!(format!("render template error {}", e)))?;
Ok(message_str) Ok(message_str)
} }
} }
+136
View File
@@ -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<R: Read + Seek> {
ct: gguf_file::Content,
reader: R,
device: Device,
}
impl<R: Read + Seek> Gguf<R> {
pub fn new(ct: gguf_file::Content, reader: R, device: Device) -> Self {
Self { ct, reader, device }
}
pub fn get_matedata(&self, name: &str) -> Result<Value> {
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<QMatMul> {
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<RmsNorm> {
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<String, gguf_file::Value> {
&self.ct.metadata
}
pub fn tensor(&mut self, name: &str) -> Result<QTensor> {
Ok(self.ct.tensor(&mut self.reader, name, &self.device)?)
}
pub fn build_tokenizer(
&self,
add_prefix_space: Option<bool>,
trim_offsets: Option<bool>,
use_regex: Option<bool>,
) -> Result<TokenizerModel> {
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<String> = vocab
.into_iter()
.map(|tokens| tokens.to_string().map(|x| x.clone()))
.collect::<Result<Vec<String>, 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<String> = merges
.into_iter()
.map(|tokens| tokens.to_string().map(|x| x.clone()))
.collect::<Result<Vec<String>, 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::<Result<Vec<i32>, 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}")),
}
}
}
+2
View File
@@ -0,0 +1,2 @@
pub mod common;
pub mod qwen3_5;
+87
View File
@@ -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<Self> {
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<R: Read + Seek>(
content: gguf_file::Content,
reader: &mut R,
device: &Device,
) -> Result<Self> {
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(())
}
}
+1
View File
@@ -1,5 +1,6 @@
pub mod chat_template; pub mod chat_template;
pub mod exec; pub mod exec;
pub mod gguf_models;
pub mod models; pub mod models;
pub mod position_embed; pub mod position_embed;
pub mod process; pub mod process;
+11 -9
View File
@@ -66,11 +66,12 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64; let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
// let mes_render = self.chat_template.apply_chat_template(&mes)?; let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mes_render = self // let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
.chat_template // let mes_render = self
.apply_chat_temp_think(&mes, enable_thinking)?; // .chat_template
// .apply_chat_temp_think(&mes, enable_thinking)?;
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32; 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 seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking"); let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mes_render = self // let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
.chat_template // let mes_render = self
.apply_chat_temp_think(&mes, enable_thinking)?; // .chat_template
// .apply_chat_temp_think(&mes, enable_thinking)?;
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let mut seqlen_offset = 0; let mut seqlen_offset = 0;
+13 -9
View File
@@ -60,16 +60,19 @@ impl<'a> Qwen3_5GenerateModel<'a> {
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> { impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let seed = mes.seed.unwrap_or(34562) as u64; let seed = mes.seed.unwrap_or(34562) as u64;
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let mes_render = self let mes_render = self.chat_template.apply_chat_template(&mes)?;
.chat_template println!("mes_render: {}", mes_render);
.apply_chat_temp_think(&mes, enable_thinking)?; // let enable_thinking = extract_metadata_value::<bool>(&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 input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self let mut input_ids = self
.tokenizer .tokenizer
.text_encode(input.replace_text.clone(), &self.device)?; .text_encode(input.replace_text.clone(), &self.device)?;
println!("input_ids: {}", input_ids);
let mut seq_len = input_ids.dim(1)?; let mut seq_len = input_ids.dim(1)?;
let prompt_tokens = seq_len as u32; let prompt_tokens = seq_len as u32;
let mut seqlen_offset = 0; 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 seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking"); let mes_render = self.chat_template.apply_chat_template(&mes)?;
let mes_render = self // let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
.chat_template // let mes_render = self
.apply_chat_temp_think(&mes, enable_thinking)?; // .chat_template
// .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?; let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self let mut input_ids = self
.tokenizer .tokenizer
+10 -10
View File
@@ -73,11 +73,11 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64; let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking"); let mes_render = self.chat_template.apply_chat_template(&mes)?;
// let mes_render = self.chat_template.apply_chat_template(&mes)?; // let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
let mes_render = self // let mes_render = self
.chat_template // .chat_template
.apply_chat_temp_think(&mes, enable_thinking)?; // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?; let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self let mut input_ids = self
.tokenizer .tokenizer
@@ -142,11 +142,11 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
let seed = mes.seed.unwrap_or(34562) as u64; let seed = mes.seed.unwrap_or(34562) as u64;
let mut logit_processor = let mut logit_processor =
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed); get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking"); let mes_render = self.chat_template.apply_chat_template(&mes)?;
// let mes_render = self.chat_template.apply_chat_template(&mes)?; // let enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
let mes_render = self // let mes_render = self
.chat_template // .chat_template
.apply_chat_temp_think(&mes, enable_thinking)?; // .apply_chat_temp_think(&mes, enable_thinking)?;
let input = self.pre_processor.process_info(&mes, &mes_render)?; let input = self.pre_processor.process_info(&mes, &mes_render)?;
let mut input_ids = self let mut input_ids = self
.tokenizer .tokenizer
+6 -6
View File
@@ -12,6 +12,10 @@ pub struct TokenizerModel {
} }
impl TokenizerModel { impl TokenizerModel {
pub fn new(tokenizer: Tokenizer) -> Self {
Self { tokenizer }
}
pub fn init(path: &str) -> Result<Self> { pub fn init(path: &str) -> Result<Self> {
let path = path.to_string(); let path = path.to_string();
assert!( assert!(
@@ -81,6 +85,8 @@ impl TokenizerModel {
} }
tokenizer tokenizer
}; };
let len = tokenizer.get_vocab_size(true);
println!("len: {}", len);
Ok(Self { tokenizer }) Ok(Self { tokenizer })
} }
@@ -94,12 +100,6 @@ impl TokenizerModel {
Ok(token_id) Ok(token_id)
} }
pub fn text_encode(&self, text: String, device: &Device) -> Result<Tensor> { pub fn text_encode(&self, text: String, device: &Device) -> Result<Tensor> {
// 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_id = self.text_encode_vec(text, true)?;
let token_tensor = Tensor::from_slice(&token_id, (1, token_id.len()), device)?; let token_tensor = Tensor::from_slice(&token_id, (1, token_id.len()), device)?;
Ok(token_tensor) Ok(token_tensor)
+47
View File
@@ -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
// <think>
// </think>
// 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(())
}
+1 -1
View File
@@ -22,7 +22,7 @@ fn glm_asr_nano_generate() -> Result<()> {
"type": "audio", "type": "audio",
"audio_url": "audio_url":
{ {
"url": "https://package-release.coderbox.cn/aiway/test/other/%E5%93%AA%E5%90%92.wav" "url": "file://./assets/audio/zh.mp3"
} }
}, },
{ {
+3 -11
View File
@@ -19,22 +19,14 @@ fn qwen3_5_generate() -> Result<()> {
"messages": [ "messages": [
{ {
"role": "user", "role": "user",
"content": [ "content": [
{
"type": "image",
"image_url":
{
"url": "https://www.lifeberrys.com/img/article/tourist-attraction-3-1644590220-lb.jpg"
}
},
{ {
"type": "text", "type": "text",
"text": "描述这张图片." "text": "你好啊"
} }
] ]
} }
], ]
"metadata": {"enable_thinking": "true"}
} }
"#; "#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?; let mes: ChatCompletionParameters = serde_json::from_str(message)?;