from gguf build chat_template and tokenizer
This commit is contained in:
Generated
+1
@@ -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",
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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}")),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
pub mod common;
|
||||||
|
pub mod qwen3_5;
|
||||||
@@ -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,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;
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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(())
|
||||||
|
}
|
||||||
@@ -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
@@ -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)?;
|
||||||
|
|||||||
Reference in New Issue
Block a user