Merge branch 'add_model'
This commit is contained in:
Generated
+1
@@ -33,6 +33,7 @@ name = "aha"
|
||||
version = "0.2.2"
|
||||
dependencies = [
|
||||
"aha_openai_dive",
|
||||
"ahash",
|
||||
"anyhow",
|
||||
"base64 0.22.1",
|
||||
"byteorder",
|
||||
|
||||
@@ -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"]
|
||||
|
||||
+54
-72
@@ -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('<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> {
|
||||
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 fixed_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')", // 使用自定义的过滤器替换
|
||||
);
|
||||
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<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 = 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::<Vec<String>>()
|
||||
});
|
||||
|
||||
// 添加 lstrip 过滤器
|
||||
env.add_filter("lstrip", |s: String, chars: Option<String>| 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<String>| 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<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);
|
||||
|
||||
Ok(Self { env })
|
||||
}
|
||||
|
||||
pub fn apply_chat_template(&self, messages: &ChatCompletionParameters) -> Result<String> {
|
||||
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<bool>,
|
||||
) -> Result<String> {
|
||||
let enable_thinking = extract_metadata_value::<bool>(&messages.metadata, "enable_thinking");
|
||||
let context = context! {
|
||||
messages => &messages.messages,
|
||||
tools => &messages.tools.as_ref(),
|
||||
|
||||
@@ -0,0 +1,234 @@
|
||||
use std::io::{Read, Seek};
|
||||
|
||||
use ahash::AHashMap;
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{
|
||||
Device, Tensor,
|
||||
quantized::{
|
||||
QMatMul, QTensor,
|
||||
gguf_file::{self, Value},
|
||||
},
|
||||
};
|
||||
use candle_nn::{Conv1d, Conv1dConfig, Linear, Module, RmsNorm, VarBuilder, linear_b};
|
||||
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_qtensor(ws)?)
|
||||
}
|
||||
|
||||
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 get_dequantized(&mut self, name: &str) -> Result<Tensor> {
|
||||
Ok(self.tensor(name)?.dequantize(&self.device)?)
|
||||
}
|
||||
|
||||
pub fn conv1d(
|
||||
&mut self,
|
||||
prefix: &str,
|
||||
padding: usize,
|
||||
stride: usize,
|
||||
dilation: usize,
|
||||
groups: usize,
|
||||
bias: bool,
|
||||
) -> Result<Conv1d> {
|
||||
let weight = self.get_dequantized(&format!("{prefix}.weight"))?;
|
||||
let bias = if bias {
|
||||
self.get_dequantized(&format!("{prefix}.bias")).ok()
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let cfg = Conv1dConfig {
|
||||
padding,
|
||||
stride,
|
||||
dilation,
|
||||
groups,
|
||||
cudnn_fwd_algo: None,
|
||||
};
|
||||
Ok(Conv1d::new(weight, bias, cfg))
|
||||
}
|
||||
|
||||
pub fn build_tokenizer(
|
||||
&self,
|
||||
add_prefix_space: Option<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().cloned())
|
||||
.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().cloned())
|
||||
.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
|
||||
&& let Some(token_str) = vocab.get(id)
|
||||
{
|
||||
let add_token = AddedToken::from(token_str.clone(), true);
|
||||
add_tokens.push(add_token);
|
||||
} else if type_ == 4
|
||||
&& let Some(token_str) = vocab.get(id)
|
||||
{
|
||||
let add_token = AddedToken::from(token_str.clone(), false);
|
||||
add_tokens.push(add_token);
|
||||
}
|
||||
}
|
||||
let _ = tokenizer.add_special_tokens(&add_tokens);
|
||||
let tokenizer_model = TokenizerModel::new(tokenizer);
|
||||
Ok(tokenizer_model)
|
||||
}
|
||||
_ => Err(anyhow!("Unsupported tokenizer model type: {model_type}")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub enum ProjKind {
|
||||
QuantizedProj(QMatMul),
|
||||
LinearProj(Linear),
|
||||
}
|
||||
|
||||
impl Module for ProjKind {
|
||||
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
|
||||
match self {
|
||||
ProjKind::QuantizedProj(q) => q.forward(xs),
|
||||
ProjKind::LinearProj(l) => l.forward(xs),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct GateUpDownMLPGguf {
|
||||
gate_proj: ProjKind, // ffn_gate.weight
|
||||
up_proj: ProjKind, // ffn_up.weight
|
||||
down_proj: ProjKind, // ffn_down.weight
|
||||
}
|
||||
|
||||
impl GateUpDownMLPGguf {
|
||||
pub fn new_from_gguf<R: Read + Seek>(gguf: &mut Gguf<R>, prefix: &str) -> Result<Self> {
|
||||
let gate_proj = gguf.qmatmul(&format!("{prefix}.ffn_gate.weight"))?;
|
||||
let up_proj = gguf.qmatmul(&format!("{prefix}.ffn_up.weight"))?;
|
||||
let down_proj = gguf.qmatmul(&format!("{prefix}.ffn_down.weight"))?;
|
||||
Ok(Self {
|
||||
gate_proj: ProjKind::QuantizedProj(gate_proj),
|
||||
up_proj: ProjKind::QuantizedProj(up_proj),
|
||||
down_proj: ProjKind::QuantizedProj(down_proj),
|
||||
})
|
||||
}
|
||||
pub fn new_from_vb(
|
||||
vb: VarBuilder,
|
||||
hidden_size: usize,
|
||||
intermediate_size: usize,
|
||||
bias: bool,
|
||||
gate_pp_name: Option<&str>,
|
||||
up_pp_name: Option<&str>,
|
||||
down_pp_name: Option<&str>,
|
||||
) -> Result<Self> {
|
||||
let gate_pp_name = gate_pp_name.unwrap_or("gate_proj");
|
||||
let up_pp_name = up_pp_name.unwrap_or("up_proj");
|
||||
let down_pp_name = down_pp_name.unwrap_or("down_proj");
|
||||
let gate_proj = linear_b(hidden_size, intermediate_size, bias, vb.pp(gate_pp_name))?;
|
||||
let up_proj = linear_b(hidden_size, intermediate_size, bias, vb.pp(up_pp_name))?;
|
||||
let down_proj = linear_b(intermediate_size, hidden_size, bias, vb.pp(down_pp_name))?;
|
||||
Ok(Self {
|
||||
gate_proj: ProjKind::LinearProj(gate_proj),
|
||||
up_proj: ProjKind::LinearProj(up_proj),
|
||||
down_proj: ProjKind::LinearProj(down_proj),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Module for GateUpDownMLPGguf {
|
||||
fn forward(&self, xs: &Tensor) -> candle_core::Result<Tensor> {
|
||||
let w1 = self.gate_proj.forward(xs)?;
|
||||
let w3 = self.up_proj.forward(xs)?;
|
||||
self.down_proj.forward(&(candle_nn::ops::silu(&w1)? * w3)?)
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,8 @@ use candle_nn::{
|
||||
embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm,
|
||||
};
|
||||
|
||||
pub mod gguf;
|
||||
|
||||
use crate::{
|
||||
position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer},
|
||||
utils::tensor_utils::{prepare_causal_attention_mask, repeat_kv},
|
||||
|
||||
@@ -11,8 +11,8 @@ use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
|
||||
use crate::models::qwen3::model::Qwen3Model;
|
||||
// use crate::models::GenerateStream;
|
||||
use crate::utils::{
|
||||
build_completion_chunk_response, build_completion_response, extract_metadata_value,
|
||||
find_type_files, get_device, get_dtype, get_logit_processor,
|
||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||
get_dtype, get_logit_processor,
|
||||
};
|
||||
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
|
||||
|
||||
@@ -66,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::<bool>(&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::<bool>(&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::<bool>(&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::<bool>(&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;
|
||||
|
||||
+100
-15
@@ -2,7 +2,7 @@ use aha_openai_dive::v1::resources::chat::{
|
||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||
};
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{DType, Device, Tensor};
|
||||
use candle_core::{DType, Device, Tensor, quantized::gguf_file};
|
||||
use candle_nn::VarBuilder;
|
||||
use rocket::async_stream::stream;
|
||||
use rocket::futures::Stream;
|
||||
@@ -11,13 +11,14 @@ use crate::{
|
||||
chat_template::ChatTemplate,
|
||||
models::{
|
||||
GenerateModel,
|
||||
common::gguf::Gguf,
|
||||
qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model},
|
||||
qwen3vl::processor::Qwen3VLProcessor,
|
||||
},
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{
|
||||
build_completion_chunk_response, build_completion_response, extract_metadata_value,
|
||||
find_type_files, get_device, get_dtype, get_logit_processor,
|
||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||
get_dtype, get_logit_processor,
|
||||
},
|
||||
};
|
||||
|
||||
@@ -29,10 +30,17 @@ pub struct Qwen3_5GenerateModel<'a> {
|
||||
device: Device,
|
||||
eos_token_id: u32,
|
||||
model_name: String,
|
||||
repeat_penalty: f32,
|
||||
repeat_last_n: usize,
|
||||
}
|
||||
|
||||
impl<'a> Qwen3_5GenerateModel<'a> {
|
||||
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||
let model_name = path
|
||||
.split("/")
|
||||
.collect::<Vec<&str>>()
|
||||
.pop()
|
||||
.unwrap_or("qwen3.5");
|
||||
let chat_template = ChatTemplate::init(path)?;
|
||||
let tokenizer = TokenizerModel::init(path)?;
|
||||
let config_path = path.to_string() + "/config.json";
|
||||
@@ -44,7 +52,7 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
||||
let model_list = find_type_files(path, "safetensors")?;
|
||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||
let eos_token_id = cfg.text_config.eos_token_id;
|
||||
let qwen3_5 = Qwen3_5Model::new(vb, cfg)?;
|
||||
let qwen3_5 = Qwen3_5Model::new_from_vb(vb, cfg)?;
|
||||
|
||||
Ok(Self {
|
||||
chat_template,
|
||||
@@ -53,19 +61,77 @@ impl<'a> Qwen3_5GenerateModel<'a> {
|
||||
qwen3_5,
|
||||
device,
|
||||
eos_token_id,
|
||||
model_name: "qwen3.5".to_string(),
|
||||
model_name: model_name.to_string(),
|
||||
repeat_penalty: 1.01,
|
||||
repeat_last_n: 64,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn init_from_gguf(
|
||||
model_file: &str,
|
||||
mmproj_file: Option<&str>,
|
||||
device: Option<&Device>,
|
||||
) -> Result<Self> {
|
||||
if !model_file.contains("Qwen3.5") || !model_file.ends_with("gguf") {
|
||||
return Err(anyhow!("Qwen3.5 gguf model file name illigal {model_file}"));
|
||||
}
|
||||
if let Some(mmproj) = mmproj_file
|
||||
&& (!mmproj.contains("mmproj") || !mmproj.ends_with("gguf"))
|
||||
{
|
||||
return Err(anyhow!("Qwen3.5 mmproj_file name illigal {model_file}"));
|
||||
}
|
||||
|
||||
let mut reader = std::fs::File::open(model_file)?;
|
||||
let content = gguf_file::Content::read(&mut reader)?;
|
||||
let device = get_device(device);
|
||||
let mut gguf = Gguf::new(content, reader, device.clone());
|
||||
|
||||
let chat_template_str = gguf
|
||||
.get_matedata("tokenizer.chat_template")?
|
||||
.to_string()?
|
||||
.clone();
|
||||
let chat_template = ChatTemplate::str_init(&chat_template_str)?;
|
||||
let tokenizer = gguf.build_tokenizer(Some(false), Some(false), Some(false))?;
|
||||
let dtype = match gguf.get_matedata("general.type") {
|
||||
Ok(v) => match v.to_u32() {
|
||||
Ok(0) => DType::F32,
|
||||
Ok(1) => DType::F16,
|
||||
_ => DType::F16,
|
||||
},
|
||||
Err(_) => DType::F16,
|
||||
};
|
||||
let pre_processor = Qwen3VLProcessor::new_qwen3_5_default(&device, dtype)?;
|
||||
|
||||
// let eos_token_id = gguf.get_matedata("tokenizer.ggml.eos_token_id")?.to_u32()?;
|
||||
let qwen3_5 = Qwen3_5Model::new_from_gguf(&mut gguf, &device)?;
|
||||
let stem = std::path::Path::new(model_file)
|
||||
.file_stem() // 获取文件名主干(不含扩展名)
|
||||
.and_then(|s| s.to_str())
|
||||
.unwrap_or("qwen3.5");
|
||||
Ok(Self {
|
||||
chat_template,
|
||||
tokenizer,
|
||||
pre_processor,
|
||||
qwen3_5,
|
||||
device,
|
||||
eos_token_id: 248044,
|
||||
model_name: stem.to_string(),
|
||||
repeat_penalty: 1.1,
|
||||
repeat_last_n: 64,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||
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 mes_render = self
|
||||
.chat_template
|
||||
.apply_chat_temp_think(&mes, enable_thinking)?;
|
||||
let seed = mes.seed.unwrap_or(32768) as u64;
|
||||
let temperature = mes.temperature.unwrap_or(0.6);
|
||||
let top_p = mes.top_p.unwrap_or(0.95);
|
||||
let mut logit_processor =
|
||||
get_logit_processor(temperature.into(), top_p.into(), Some(20), seed);
|
||||
// let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
|
||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
let mut input_ids = self
|
||||
.tokenizer
|
||||
@@ -89,6 +155,16 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
||||
seqlen_offset,
|
||||
)?;
|
||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||
let logits = if self.repeat_penalty == 1. {
|
||||
logits
|
||||
} else {
|
||||
let start_at = generate.len().saturating_sub(self.repeat_last_n);
|
||||
candle_transformers::utils::apply_repeat_penalty(
|
||||
&logits,
|
||||
self.repeat_penalty,
|
||||
&generate[start_at..],
|
||||
)?
|
||||
};
|
||||
let next_token = logit_processor.sample(&logits)?;
|
||||
generate.push(next_token);
|
||||
if next_token == self.eos_token_id {
|
||||
@@ -125,10 +201,7 @@ 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::<bool>(&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 input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
let mut input_ids = self
|
||||
.tokenizer
|
||||
@@ -144,6 +217,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
||||
let video_grid_thw = input.video_grid_thw.as_ref();
|
||||
let mut tool_call_id = None;
|
||||
let mut tool_call_content = String::new();
|
||||
let mut generate = Vec::new();
|
||||
for _ in 0..sample_len {
|
||||
let logits = self.qwen3_5.forward(
|
||||
&input_ids,
|
||||
@@ -154,7 +228,18 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
|
||||
seqlen_offset,
|
||||
)?;
|
||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||
let logits = if self.repeat_penalty == 1. {
|
||||
logits
|
||||
} else {
|
||||
let start_at = generate.len().saturating_sub(self.repeat_last_n);
|
||||
candle_transformers::utils::apply_repeat_penalty(
|
||||
&logits,
|
||||
self.repeat_penalty,
|
||||
&generate[start_at..],
|
||||
)?
|
||||
};
|
||||
let next_token = logit_processor.sample(&logits)?;
|
||||
generate.push(next_token);
|
||||
let mut decode_ids = Vec::new();
|
||||
if !error_tokens.is_empty() {
|
||||
decode_ids.extend_from_slice(&error_tokens);
|
||||
|
||||
+347
-85
@@ -1,5 +1,7 @@
|
||||
use std::io::{Read, Seek};
|
||||
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{D, IndexOp, Tensor};
|
||||
use candle_core::{D, DType, Device, IndexOp, Tensor, quantized::QMatMul};
|
||||
use candle_nn::{
|
||||
Conv1d, Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias,
|
||||
ops::sigmoid, rms_norm,
|
||||
@@ -7,7 +9,11 @@ use candle_nn::{
|
||||
|
||||
use crate::{
|
||||
models::{
|
||||
common::{GateUpDownMLP, conv1d_depthwise, eager_attention_forward, get_conv1d, softplus},
|
||||
common::{
|
||||
conv1d_depthwise, eager_attention_forward, get_conv1d,
|
||||
gguf::{GateUpDownMLPGguf, Gguf, ProjKind},
|
||||
softplus,
|
||||
},
|
||||
qwen3_5::config::{Qwen3_5Config, Qwen3_5TextConfig},
|
||||
qwen3vl::model::Qwen3VLVisionModel,
|
||||
},
|
||||
@@ -30,6 +36,10 @@ impl Qwen3_5RMSNorm {
|
||||
Ok(Self { eps, weight })
|
||||
}
|
||||
|
||||
pub fn from_weight(weight: Tensor, eps: f64) -> Result<Self> {
|
||||
Ok(Self { eps, weight })
|
||||
}
|
||||
|
||||
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
|
||||
let x = xs.to_dtype(candle_core::DType::F32)?;
|
||||
let norm_ = x
|
||||
@@ -53,6 +63,11 @@ impl Qwen3_5RMSNormGated {
|
||||
Ok(Self { norm })
|
||||
}
|
||||
|
||||
pub fn from_weight(weight: Tensor, eps: f64) -> Result<Self> {
|
||||
let norm = RmsNorm::new(weight, eps);
|
||||
Ok(Self { norm })
|
||||
}
|
||||
|
||||
pub fn forward(&self, xs: &Tensor, gate: Option<&Tensor>) -> Result<Tensor> {
|
||||
let mut xs = self.norm.forward(xs)?;
|
||||
if let Some(gate) = gate {
|
||||
@@ -62,6 +77,7 @@ impl Qwen3_5RMSNormGated {
|
||||
}
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! transmute_tensors {
|
||||
($($tensor:expr),*) => {
|
||||
($(
|
||||
@@ -69,7 +85,7 @@ macro_rules! transmute_tensors {
|
||||
)*)
|
||||
};
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! right_pad_zero_tensor {
|
||||
($dim:expr, $pad_size:expr, $($tensor:expr),+) => {
|
||||
($(
|
||||
@@ -78,6 +94,7 @@ macro_rules! right_pad_zero_tensor {
|
||||
};
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! reshape_chunk_tensor {
|
||||
($chunk_size:expr, $($tensor:expr),*) => {
|
||||
($(
|
||||
@@ -107,30 +124,30 @@ pub struct Qwen3_5GatedDeltaNet {
|
||||
dt_bias: Tensor,
|
||||
a_log: Tensor,
|
||||
norm: Qwen3_5RMSNormGated,
|
||||
out_proj: Linear,
|
||||
out_proj: ProjKind,
|
||||
|
||||
// Z, B, A 投影
|
||||
in_proj_qkv: candle_nn::Linear,
|
||||
in_proj_z: candle_nn::Linear,
|
||||
in_proj_b: candle_nn::Linear,
|
||||
in_proj_a: candle_nn::Linear,
|
||||
in_proj_qkv: ProjKind,
|
||||
in_proj_z: ProjKind,
|
||||
in_proj_b: ProjKind,
|
||||
in_proj_a: ProjKind,
|
||||
conv_state_cache: Option<Tensor>,
|
||||
recurrent_state_cache: Option<Tensor>,
|
||||
}
|
||||
|
||||
impl Qwen3_5GatedDeltaNet {
|
||||
pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result<Self> {
|
||||
let hidden_size = config.hidden_size;
|
||||
let num_v_heads = config.linear_num_value_heads;
|
||||
let num_k_heads = config.linear_num_key_heads;
|
||||
let head_k_dim = config.linear_key_head_dim;
|
||||
let head_v_dim = config.linear_value_head_dim;
|
||||
let key_dim = head_k_dim * num_k_heads;
|
||||
let value_dim = head_v_dim * num_v_heads;
|
||||
let conv_kernel_size = config.linear_conv_kernel_dim;
|
||||
pub fn new_from_vb(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result<Self> {
|
||||
let hidden_size = config.hidden_size; // 1024
|
||||
let num_v_heads = config.linear_num_value_heads; // 16
|
||||
let num_k_heads = config.linear_num_key_heads; // 16
|
||||
let head_k_dim = config.linear_key_head_dim; // 128
|
||||
let head_v_dim = config.linear_value_head_dim; // 128
|
||||
let key_dim = head_k_dim * num_k_heads; // 2048
|
||||
let value_dim = head_v_dim * num_v_heads; // 2048
|
||||
let conv_kernel_size = config.linear_conv_kernel_dim; // 4
|
||||
// let activation = config.hidden_act;
|
||||
let layer_norm_epsilon = config.rms_norm_eps;
|
||||
let conv_dim = key_dim * 2 + value_dim;
|
||||
let conv_dim = key_dim * 2 + value_dim; // 6144
|
||||
let conv1d = get_conv1d(
|
||||
vb.pp("conv1d"),
|
||||
conv_dim,
|
||||
@@ -146,12 +163,17 @@ impl Qwen3_5GatedDeltaNet {
|
||||
let a_log = vb.get(num_v_heads, "A_log")?;
|
||||
let norm = Qwen3_5RMSNormGated::new(vb.pp("norm"), head_v_dim, layer_norm_epsilon)?;
|
||||
|
||||
// 2048, 1024
|
||||
let out_proj = linear_no_bias(value_dim, hidden_size, vb.pp("out_proj"))?;
|
||||
|
||||
// 1024, 6144
|
||||
let in_proj_qkv = linear_no_bias(hidden_size, conv_dim, vb.pp("in_proj_qkv"))?;
|
||||
// 1024, 2048
|
||||
let in_proj_z = linear_no_bias(hidden_size, value_dim, vb.pp("in_proj_z"))?;
|
||||
// 1024, 16
|
||||
let in_proj_b = linear_no_bias(hidden_size, num_v_heads, vb.pp("in_proj_b"))?;
|
||||
// 1024, 16
|
||||
let in_proj_a = linear_no_bias(hidden_size, num_v_heads, vb.pp("in_proj_a"))?;
|
||||
|
||||
Ok(Self {
|
||||
// hidden_size,
|
||||
num_v_heads,
|
||||
@@ -168,11 +190,64 @@ impl Qwen3_5GatedDeltaNet {
|
||||
dt_bias,
|
||||
a_log,
|
||||
norm,
|
||||
out_proj,
|
||||
in_proj_qkv,
|
||||
in_proj_z,
|
||||
in_proj_b,
|
||||
in_proj_a,
|
||||
out_proj: ProjKind::LinearProj(out_proj),
|
||||
in_proj_qkv: ProjKind::LinearProj(in_proj_qkv),
|
||||
in_proj_z: ProjKind::LinearProj(in_proj_z),
|
||||
in_proj_b: ProjKind::LinearProj(in_proj_b),
|
||||
in_proj_a: ProjKind::LinearProj(in_proj_a),
|
||||
conv_state_cache: None,
|
||||
recurrent_state_cache: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn new_from_gguf<R: Read + Seek>(
|
||||
gguf: &mut Gguf<R>,
|
||||
prefix: &str,
|
||||
rms_norm_eps: f64,
|
||||
) -> Result<Self> {
|
||||
let num_k_heads = gguf.get_matedata("qwen35.ssm.group_count")?.to_u32()? as usize;
|
||||
let num_v_heads = gguf.get_matedata("qwen35.ssm.time_step_rank")?.to_u32()? as usize;
|
||||
let conv_kernel_size = gguf.get_matedata("qwen35.ssm.conv_kernel")?.to_u32()? as usize;
|
||||
let head_k_dim = gguf.get_matedata("qwen35.ssm.state_size")?.to_u32()? as usize;
|
||||
let head_v_dim = head_k_dim;
|
||||
let key_dim = head_k_dim * num_k_heads;
|
||||
let value_dim = head_v_dim * num_v_heads;
|
||||
let conv_dim = key_dim * 2 + value_dim;
|
||||
let conv1d = gguf.conv1d(
|
||||
&format!("{prefix}.ssm_conv1d"),
|
||||
conv_kernel_size - 1,
|
||||
1,
|
||||
1,
|
||||
conv_dim,
|
||||
false,
|
||||
)?;
|
||||
let dt_bias = gguf.get_dequantized(&format!("{prefix}.ssm_dt.bias"))?;
|
||||
let a_log = gguf.get_dequantized(&format!("{prefix}.ssm_a"))?;
|
||||
let norm_weight = gguf.get_dequantized(&format!("{prefix}.ssm_norm.weight"))?;
|
||||
let norm = Qwen3_5RMSNormGated::from_weight(norm_weight, rms_norm_eps)?;
|
||||
let out_proj = gguf.qmatmul(&format!("{prefix}.ssm_out.weight"))?;
|
||||
let in_proj_qkv = gguf.qmatmul(&format!("{prefix}.attn_qkv.weight"))?;
|
||||
let in_proj_z = gguf.qmatmul(&format!("{prefix}.attn_gate.weight"))?;
|
||||
let in_proj_b = gguf.qmatmul(&format!("{prefix}.ssm_beta.weight"))?;
|
||||
let in_proj_a = gguf.qmatmul(&format!("{prefix}.ssm_alpha.weight"))?;
|
||||
|
||||
Ok(Self {
|
||||
num_v_heads,
|
||||
num_k_heads,
|
||||
head_k_dim,
|
||||
head_v_dim,
|
||||
key_dim,
|
||||
value_dim,
|
||||
conv_kernel_size,
|
||||
conv1d,
|
||||
dt_bias,
|
||||
a_log,
|
||||
norm,
|
||||
out_proj: ProjKind::QuantizedProj(out_proj),
|
||||
in_proj_qkv: ProjKind::QuantizedProj(in_proj_qkv),
|
||||
in_proj_z: ProjKind::QuantizedProj(in_proj_z),
|
||||
in_proj_b: ProjKind::QuantizedProj(in_proj_b),
|
||||
in_proj_a: ProjKind::QuantizedProj(in_proj_a),
|
||||
conv_state_cache: None,
|
||||
recurrent_state_cache: None,
|
||||
})
|
||||
@@ -466,7 +541,6 @@ impl Qwen3_5GatedDeltaNet {
|
||||
&[self.key_dim, self.key_dim, self.value_dim],
|
||||
D::Minus1,
|
||||
)?;
|
||||
|
||||
let mut query = qkv_split[0].reshape((bs, seq_len, (), self.head_k_dim))?;
|
||||
let mut key = qkv_split[1].reshape((bs, seq_len, (), self.head_k_dim))?;
|
||||
let value = qkv_split[2].reshape((bs, seq_len, (), self.head_v_dim))?;
|
||||
@@ -502,10 +576,10 @@ impl Qwen3_5GatedDeltaNet {
|
||||
}
|
||||
|
||||
pub struct Qwen3_5Attention {
|
||||
q_proj: Linear,
|
||||
k_proj: Linear,
|
||||
v_proj: Linear,
|
||||
o_proj: Linear,
|
||||
q_proj: ProjKind,
|
||||
k_proj: ProjKind,
|
||||
v_proj: ProjKind,
|
||||
o_proj: ProjKind,
|
||||
q_norm: Qwen3_5RMSNorm,
|
||||
k_norm: Qwen3_5RMSNorm,
|
||||
num_attention_heads: usize,
|
||||
@@ -517,7 +591,7 @@ pub struct Qwen3_5Attention {
|
||||
}
|
||||
|
||||
impl Qwen3_5Attention {
|
||||
pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result<Self> {
|
||||
pub fn new_from_vb(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result<Self> {
|
||||
let hidden_size = config.hidden_size;
|
||||
let num_attention_heads = config.num_attention_heads;
|
||||
let head_dim = config.head_dim;
|
||||
@@ -550,11 +624,50 @@ impl Qwen3_5Attention {
|
||||
)?;
|
||||
let q_norm = Qwen3_5RMSNorm::new(vb.pp("q_norm"), head_dim, config.rms_norm_eps)?;
|
||||
let k_norm = Qwen3_5RMSNorm::new(vb.pp("k_norm"), head_dim, config.rms_norm_eps)?;
|
||||
|
||||
Ok(Self {
|
||||
q_proj,
|
||||
k_proj,
|
||||
v_proj,
|
||||
o_proj,
|
||||
q_proj: ProjKind::LinearProj(q_proj),
|
||||
k_proj: ProjKind::LinearProj(k_proj),
|
||||
v_proj: ProjKind::LinearProj(v_proj),
|
||||
o_proj: ProjKind::LinearProj(o_proj),
|
||||
q_norm,
|
||||
k_norm,
|
||||
num_attention_heads,
|
||||
num_key_value_heads,
|
||||
num_kv_groups,
|
||||
head_dim,
|
||||
scaling,
|
||||
kv_cache: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn new_from_gguf<R: Read + Seek>(
|
||||
gguf: &mut Gguf<R>,
|
||||
prefix: &str,
|
||||
rms_norm_eps: f64,
|
||||
) -> Result<Self> {
|
||||
let num_attention_heads =
|
||||
gguf.get_matedata("qwen35.attention.head_count")?.to_u32()? as usize;
|
||||
let num_key_value_heads = gguf
|
||||
.get_matedata("qwen35.attention.head_count_kv")?
|
||||
.to_u32()? as usize;
|
||||
let num_kv_groups = num_attention_heads / num_key_value_heads;
|
||||
let head_dim = gguf.get_matedata("qwen35.attention.key_length")?.to_u32()? as usize;
|
||||
let scaling = 1f64 / f64::sqrt(head_dim as f64);
|
||||
let q_proj = gguf.qmatmul(&format!("{prefix}.attn_q.weight"))?;
|
||||
let k_proj = gguf.qmatmul(&format!("{prefix}.attn_k.weight"))?;
|
||||
let v_proj = gguf.qmatmul(&format!("{prefix}.attn_v.weight"))?;
|
||||
let o_proj = gguf.qmatmul(&format!("{prefix}.attn_output.weight"))?;
|
||||
let q_norm_weight = gguf.get_dequantized(&format!("{prefix}.attn_q_norm.weight"))?;
|
||||
let q_norm = Qwen3_5RMSNorm::from_weight(q_norm_weight, rms_norm_eps)?;
|
||||
let k_norm_weight = gguf.get_dequantized(&format!("{prefix}.attn_k_norm.weight"))?;
|
||||
let k_norm = Qwen3_5RMSNorm::from_weight(k_norm_weight, rms_norm_eps)?;
|
||||
|
||||
Ok(Self {
|
||||
q_proj: ProjKind::QuantizedProj(q_proj),
|
||||
k_proj: ProjKind::QuantizedProj(k_proj),
|
||||
v_proj: ProjKind::QuantizedProj(v_proj),
|
||||
o_proj: ProjKind::QuantizedProj(o_proj),
|
||||
q_norm,
|
||||
k_norm,
|
||||
num_attention_heads,
|
||||
@@ -627,32 +740,62 @@ impl Qwen3_5Attention {
|
||||
}
|
||||
}
|
||||
|
||||
enum AttnKind {
|
||||
LinearAttn(Qwen3_5GatedDeltaNet),
|
||||
SelfAttn(Qwen3_5Attention),
|
||||
}
|
||||
|
||||
impl AttnKind {
|
||||
fn forward(
|
||||
&mut self,
|
||||
xs: &Tensor,
|
||||
cos: Option<&Tensor>,
|
||||
sin: Option<&Tensor>,
|
||||
attention_mask: Option<&Tensor>,
|
||||
) -> Result<Tensor> {
|
||||
match self {
|
||||
AttnKind::LinearAttn(attn) => attn.forward(xs, attention_mask),
|
||||
AttnKind::SelfAttn(attn) => {
|
||||
if let Some(cos) = cos
|
||||
&& let Some(sin) = sin
|
||||
{
|
||||
attn.forward(xs, cos, sin, attention_mask)
|
||||
} else {
|
||||
Err(anyhow!("Qwen3_5 self attn cos and sin is all need"))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct Qwen3_5DecoderLayer {
|
||||
// hidden_size: usize,
|
||||
layer_type: String,
|
||||
linear_attn: Option<Qwen3_5GatedDeltaNet>,
|
||||
self_attn: Option<Qwen3_5Attention>,
|
||||
mlp: GateUpDownMLP,
|
||||
attn: AttnKind,
|
||||
mlp: GateUpDownMLPGguf,
|
||||
input_layernorm: Qwen3_5RMSNorm,
|
||||
post_attention_layernorm: Qwen3_5RMSNorm,
|
||||
}
|
||||
|
||||
impl Qwen3_5DecoderLayer {
|
||||
pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig, layer_idx: usize) -> Result<Self> {
|
||||
pub fn new_from_vb(
|
||||
vb: VarBuilder,
|
||||
config: &Qwen3_5TextConfig,
|
||||
layer_idx: usize,
|
||||
) -> Result<Self> {
|
||||
let hidden_size = config.hidden_size;
|
||||
let layer_type = config.layer_types[layer_idx].clone();
|
||||
let (linear_attn, self_attn) = if layer_type.eq("linear_attention") {
|
||||
let linear_attn = Qwen3_5GatedDeltaNet::new(vb.pp("linear_attn"), config)?;
|
||||
(Some(linear_attn), None)
|
||||
let attn = if layer_type.eq("linear_attention") {
|
||||
let attn = Qwen3_5GatedDeltaNet::new_from_vb(vb.pp("linear_attn"), config)?;
|
||||
AttnKind::LinearAttn(attn)
|
||||
} else {
|
||||
let self_attn = Qwen3_5Attention::new(vb.pp("self_attn"), config)?;
|
||||
(None, Some(self_attn))
|
||||
let attn = Qwen3_5Attention::new_from_vb(vb.pp("self_attn"), config)?;
|
||||
AttnKind::SelfAttn(attn)
|
||||
};
|
||||
let mlp = GateUpDownMLP::new(
|
||||
let mlp = GateUpDownMLPGguf::new_from_vb(
|
||||
vb.pp("mlp"),
|
||||
hidden_size,
|
||||
config.intermediate_size,
|
||||
config.hidden_act,
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
@@ -668,8 +811,36 @@ impl Qwen3_5DecoderLayer {
|
||||
Ok(Self {
|
||||
// hidden_size,
|
||||
layer_type,
|
||||
linear_attn,
|
||||
self_attn,
|
||||
attn,
|
||||
mlp,
|
||||
input_layernorm,
|
||||
post_attention_layernorm,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn new_from_gguf<R: Read + Seek>(
|
||||
gguf: &mut Gguf<R>,
|
||||
prefix: &str,
|
||||
layer_type: &str,
|
||||
rms_norm_eps: f64,
|
||||
) -> Result<Self> {
|
||||
let attn = if layer_type.eq("linear_attention") {
|
||||
let attn = Qwen3_5GatedDeltaNet::new_from_gguf(gguf, prefix, rms_norm_eps)?;
|
||||
AttnKind::LinearAttn(attn)
|
||||
} else {
|
||||
let attn = Qwen3_5Attention::new_from_gguf(gguf, prefix, rms_norm_eps)?;
|
||||
AttnKind::SelfAttn(attn)
|
||||
};
|
||||
let mlp = GateUpDownMLPGguf::new_from_gguf(gguf, prefix)?;
|
||||
let input_norm_weight = gguf.get_dequantized(&format!("{prefix}.attn_norm.weight"))?;
|
||||
let input_layernorm = Qwen3_5RMSNorm::from_weight(input_norm_weight, rms_norm_eps)?;
|
||||
let post_norm_weight =
|
||||
gguf.get_dequantized(&format!("{prefix}.post_attention_norm.weight"))?;
|
||||
let post_attention_layernorm = Qwen3_5RMSNorm::from_weight(post_norm_weight, rms_norm_eps)?;
|
||||
Ok(Self {
|
||||
// hidden_size,
|
||||
layer_type: layer_type.to_string(),
|
||||
attn,
|
||||
mlp,
|
||||
input_layernorm,
|
||||
post_attention_layernorm,
|
||||
@@ -685,16 +856,17 @@ impl Qwen3_5DecoderLayer {
|
||||
) -> Result<Tensor> {
|
||||
let residual = xs.clone();
|
||||
let mut xs = self.input_layernorm.forward(xs)?;
|
||||
if self.layer_type.eq("linear_attention")
|
||||
&& let Some(linear_attn) = self.linear_attn.as_mut()
|
||||
{
|
||||
xs = linear_attn.forward(&xs, attention_mask)?;
|
||||
} else if let Some(self_attn) = self.self_attn.as_mut()
|
||||
&& let Some(cos) = cos
|
||||
&& let Some(sin) = sin
|
||||
{
|
||||
xs = self_attn.forward(&xs, cos, sin, attention_mask)?;
|
||||
}
|
||||
xs = self.attn.forward(&xs, cos, sin, attention_mask)?;
|
||||
// if self.layer_type.eq("linear_attention")
|
||||
// && let Some(linear_attn) = self.linear_attn.as_mut()
|
||||
// {
|
||||
// xs = linear_attn.forward(&xs, attention_mask)?;
|
||||
// } else if let Some(self_attn) = self.self_attn.as_mut()
|
||||
// && let Some(cos) = cos
|
||||
// && let Some(sin) = sin
|
||||
// {
|
||||
// xs = self_attn.forward(&xs, cos, sin, attention_mask)?;
|
||||
// }
|
||||
let residual = xs.add(&residual)?;
|
||||
xs = self.post_attention_layernorm.forward(&residual)?;
|
||||
xs = self.mlp.forward(&xs)?;
|
||||
@@ -703,11 +875,13 @@ impl Qwen3_5DecoderLayer {
|
||||
}
|
||||
|
||||
pub fn clear_cache(&mut self) {
|
||||
if let Some(linear_attn) = self.linear_attn.as_mut() {
|
||||
linear_attn.clear_cache();
|
||||
}
|
||||
if let Some(self_attn) = self.self_attn.as_mut() {
|
||||
self_attn.clear_kv_cache();
|
||||
match &mut self.attn {
|
||||
AttnKind::LinearAttn(attn) => {
|
||||
attn.clear_cache();
|
||||
}
|
||||
AttnKind::SelfAttn(attn) => {
|
||||
attn.clear_kv_cache();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -718,15 +892,17 @@ pub struct Qwen3_5TextModel {
|
||||
norm: Qwen3_5RMSNorm,
|
||||
rotary_emb: Qwen3VLTextRotaryEmbedding,
|
||||
mrope_section: Vec<usize>,
|
||||
dtype: DType,
|
||||
}
|
||||
|
||||
impl Qwen3_5TextModel {
|
||||
pub fn new(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result<Self> {
|
||||
pub fn new_from_vb(vb: VarBuilder, config: &Qwen3_5TextConfig) -> Result<Self> {
|
||||
let embed_tokens = embedding(config.vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
|
||||
let mut layers = vec![];
|
||||
let vb_layers = vb.pp("layers");
|
||||
for i in 0..config.num_hidden_layers {
|
||||
let layer = Qwen3_5DecoderLayer::new(vb_layers.pp(i), config, i)?;
|
||||
// for i in 0..4 {
|
||||
let layer = Qwen3_5DecoderLayer::new_from_vb(vb_layers.pp(i), config, i)?;
|
||||
layers.push(layer);
|
||||
}
|
||||
let norm = Qwen3_5RMSNorm::new(vb.pp("norm"), config.hidden_size, config.rms_norm_eps)?;
|
||||
@@ -740,17 +916,70 @@ impl Qwen3_5TextModel {
|
||||
norm,
|
||||
rotary_emb,
|
||||
mrope_section: config.rope_parameters.mrope_section.clone(),
|
||||
dtype: vb.dtype(),
|
||||
})
|
||||
}
|
||||
pub fn new_from_gguf<R: Read + Seek>(gguf: &mut Gguf<R>, device: &Device) -> Result<Self> {
|
||||
let num_layers = gguf.get_matedata("qwen35.block_count")?.to_u32()? as usize;
|
||||
let full_attention_interval = gguf
|
||||
.get_matedata("qwen35.full_attention_interval")?
|
||||
.to_u32()? as usize;
|
||||
let rope_freq_base = gguf.get_matedata("qwen35.rope.freq_base")?.to_f32()?;
|
||||
let rope_dimension_count =
|
||||
gguf.get_matedata("qwen35.rope.dimension_count")?.to_u32()? as usize;
|
||||
let mut mrope_section = gguf
|
||||
.get_matedata("qwen35.rope.dimension_sections")?
|
||||
.to_vec()?
|
||||
.iter()
|
||||
.map(|v| v.to_i32().map(|x| x as usize))
|
||||
.collect::<Result<Vec<usize>, candle_core::Error>>()?;
|
||||
let _ = mrope_section.pop();
|
||||
let rms_norm_eps = gguf
|
||||
.get_matedata("qwen35.attention.layer_norm_rms_epsilon")?
|
||||
.to_f32()? as f64;
|
||||
let hidden_size = gguf.get_matedata("qwen35.embedding_length")?.to_u32()? as usize; // 1024
|
||||
let embed_tensor = gguf.tensor("token_embd.weight")?;
|
||||
let embed_tokens = Embedding::new(embed_tensor.dequantize(device)?, hidden_size);
|
||||
let mut layers = vec![];
|
||||
for i in 0..num_layers {
|
||||
// for i in 0..4 {
|
||||
let prefix = format!("blk.{i}");
|
||||
let layer_type = if (i + 1) % full_attention_interval == 0 {
|
||||
"full_attention".to_string()
|
||||
} else {
|
||||
"linear_attention".to_string()
|
||||
};
|
||||
let layer =
|
||||
Qwen3_5DecoderLayer::new_from_gguf(gguf, &prefix, &layer_type, rms_norm_eps)?;
|
||||
layers.push(layer);
|
||||
}
|
||||
let norm_weight = gguf.get_dequantized("output_norm.weight")?;
|
||||
let norm = Qwen3_5RMSNorm::from_weight(norm_weight, rms_norm_eps)?;
|
||||
let rotary_emb = Qwen3VLTextRotaryEmbedding::new(rope_dimension_count, rope_freq_base);
|
||||
let dtype = match gguf.get_matedata("general.type") {
|
||||
Ok(v) => match v.to_u32() {
|
||||
Ok(0) => DType::F32,
|
||||
Ok(1) => DType::F16,
|
||||
_ => DType::F16,
|
||||
},
|
||||
Err(_) => DType::F16,
|
||||
};
|
||||
Ok(Self {
|
||||
embed_tokens,
|
||||
layers,
|
||||
norm,
|
||||
rotary_emb,
|
||||
mrope_section,
|
||||
dtype,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn forward(&mut self, inputs_embeds: &Tensor, position_ids: &Tensor) -> Result<Tensor> {
|
||||
let (b_size, seq_len, _) = inputs_embeds.dims3()?;
|
||||
|
||||
let (cos, sin) = self.rotary_emb.forward(
|
||||
position_ids,
|
||||
inputs_embeds.dtype(),
|
||||
self.mrope_section.clone(),
|
||||
)?;
|
||||
let (cos, sin) =
|
||||
self.rotary_emb
|
||||
.forward(position_ids, self.dtype, self.mrope_section.clone())?;
|
||||
let mut xs = inputs_embeds.clone();
|
||||
let attention_mask: Option<Tensor> = {
|
||||
if seq_len <= 1 {
|
||||
@@ -785,18 +1014,23 @@ impl Qwen3_5TextModel {
|
||||
}
|
||||
|
||||
pub struct Qwen3_5Model {
|
||||
config: Qwen3_5Config,
|
||||
visual: Qwen3VLVisionModel,
|
||||
// config: Qwen3_5Config,
|
||||
spatial_merge_size: usize,
|
||||
image_token_id: u32,
|
||||
video_token_id: u32,
|
||||
vision_start_token_id: u32,
|
||||
visual: Option<Qwen3VLVisionModel>,
|
||||
language_model: Qwen3_5TextModel,
|
||||
lm_head: Linear,
|
||||
lm_head: ProjKind,
|
||||
rope_deltas: Option<Tensor>,
|
||||
}
|
||||
|
||||
impl Qwen3_5Model {
|
||||
pub fn new(vb: VarBuilder, config: Qwen3_5Config) -> Result<Self> {
|
||||
pub fn new_from_vb(vb: VarBuilder, config: Qwen3_5Config) -> Result<Self> {
|
||||
let vb_m = vb.pp("model");
|
||||
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
|
||||
let language_model = Qwen3_5TextModel::new(vb_m.pp("language_model"), &config.text_config)?;
|
||||
let language_model =
|
||||
Qwen3_5TextModel::new_from_vb(vb_m.pp("language_model"), &config.text_config)?;
|
||||
let lm_head = if config.tie_word_embeddings {
|
||||
Linear::new(language_model.embed_tokens.embeddings().clone(), None)
|
||||
} else {
|
||||
@@ -807,10 +1041,36 @@ impl Qwen3_5Model {
|
||||
)?
|
||||
};
|
||||
Ok(Self {
|
||||
config,
|
||||
visual,
|
||||
spatial_merge_size: config.vision_config.spatial_merge_size,
|
||||
image_token_id: config.image_token_id,
|
||||
video_token_id: config.video_token_id,
|
||||
vision_start_token_id: config.vision_start_token_id,
|
||||
visual: Some(visual),
|
||||
language_model,
|
||||
lm_head,
|
||||
lm_head: ProjKind::LinearProj(lm_head),
|
||||
rope_deltas: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn new_from_gguf<R: Read + Seek>(gguf: &mut Gguf<R>, device: &Device) -> Result<Self> {
|
||||
let spatial_merge_size = 2usize;
|
||||
let image_token_id = 248056u32;
|
||||
let video_token_id = 248057u32;
|
||||
let vision_start_token_id = 248053u32;
|
||||
let language_model = Qwen3_5TextModel::new_from_gguf(gguf, device)?;
|
||||
let lm_head_tensor = match gguf.tensor("output.weight") {
|
||||
Ok(tensor) => tensor,
|
||||
Err(_) => gguf.tensor("token_embd.weight")?,
|
||||
};
|
||||
let lm_head = QMatMul::from_qtensor(lm_head_tensor)?;
|
||||
Ok(Self {
|
||||
spatial_merge_size,
|
||||
image_token_id,
|
||||
video_token_id,
|
||||
vision_start_token_id,
|
||||
visual: None,
|
||||
language_model,
|
||||
lm_head: ProjKind::QuantizedProj(lm_head),
|
||||
rope_deltas: None,
|
||||
})
|
||||
}
|
||||
@@ -842,10 +1102,10 @@ impl Qwen3_5Model {
|
||||
None => None,
|
||||
};
|
||||
|
||||
let spatial_merge_size = self.config.vision_config.spatial_merge_size;
|
||||
let image_token_id = self.config.image_token_id;
|
||||
let video_token_id = self.config.video_token_id;
|
||||
let vision_start_token_id = self.config.vision_start_token_id;
|
||||
let spatial_merge_size = self.spatial_merge_size;
|
||||
let image_token_id = self.image_token_id;
|
||||
let video_token_id = self.video_token_id;
|
||||
let vision_start_token_id = self.vision_start_token_id;
|
||||
let mut mrope_position_deltas = vec![];
|
||||
if image_grid_thw.is_some() || video_grid_thw.is_some() {
|
||||
let total_input_ids = input_ids.clone();
|
||||
@@ -1093,9 +1353,10 @@ impl Qwen3_5Model {
|
||||
let mut inputs_embeds = self.language_model.embed_tokens.forward(input_ids)?;
|
||||
if let Some(pixel_values) = pixel_values
|
||||
&& let Some(image_grid_thw) = image_grid_thw
|
||||
&& let Some(visual) = self.visual.as_ref()
|
||||
{
|
||||
let (image_embeds, _) = self.visual.forward(pixel_values, image_grid_thw)?;
|
||||
let vision_mask = get_equal_mask(input_ids, self.config.image_token_id)?;
|
||||
let (image_embeds, _) = visual.forward(pixel_values, image_grid_thw)?;
|
||||
let vision_mask = get_equal_mask(input_ids, self.image_token_id)?;
|
||||
let n_image_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
|
||||
if n_image_tokens as usize != image_embeds.dim(0)? {
|
||||
return Err(anyhow!(format!(
|
||||
@@ -1108,9 +1369,10 @@ impl Qwen3_5Model {
|
||||
}
|
||||
if let Some(pixel_values_video) = pixel_values_video
|
||||
&& let Some(video_grid_thw) = video_grid_thw
|
||||
&& let Some(visual) = self.visual.as_ref()
|
||||
{
|
||||
let (video_embeds, _) = self.visual.forward(pixel_values_video, video_grid_thw)?;
|
||||
let vision_mask = get_equal_mask(input_ids, self.config.video_token_id)?;
|
||||
let (video_embeds, _) = visual.forward(pixel_values_video, video_grid_thw)?;
|
||||
let vision_mask = get_equal_mask(input_ids, self.video_token_id)?;
|
||||
let n_video_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
|
||||
if n_video_tokens as usize != video_embeds.dim(0)? {
|
||||
return Err(anyhow!(format!(
|
||||
|
||||
@@ -18,6 +18,36 @@ pub struct PreprocessorConfig {
|
||||
pub image_std: Vec<f32>,
|
||||
}
|
||||
|
||||
impl PreprocessorConfig {
|
||||
pub fn qwen3_5_img_default() -> Self {
|
||||
Self {
|
||||
size: Size {
|
||||
longest_edge: 16777216,
|
||||
shortest_edge: 65536,
|
||||
},
|
||||
patch_size: 16,
|
||||
temporal_patch_size: 2,
|
||||
merge_size: 2,
|
||||
image_mean: vec![0.5, 0.5, 0.5],
|
||||
image_std: vec![0.5, 0.5, 0.5],
|
||||
}
|
||||
}
|
||||
|
||||
pub fn qwen3_5_video_default() -> Self {
|
||||
Self {
|
||||
size: Size {
|
||||
longest_edge: 25165824,
|
||||
shortest_edge: 4096,
|
||||
},
|
||||
patch_size: 16,
|
||||
temporal_patch_size: 2,
|
||||
merge_size: 2,
|
||||
image_mean: vec![0.5, 0.5, 0.5],
|
||||
image_std: vec![0.5, 0.5, 0.5],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct RopeScaling {
|
||||
pub rope_type: String,
|
||||
|
||||
@@ -16,8 +16,8 @@ use crate::{
|
||||
},
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{
|
||||
build_completion_chunk_response, build_completion_response, extract_metadata_value,
|
||||
find_type_files, get_device, get_dtype, get_logit_processor,
|
||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||
get_dtype, get_logit_processor,
|
||||
},
|
||||
};
|
||||
|
||||
@@ -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::<bool>(&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::<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 mut input_ids = self
|
||||
.tokenizer
|
||||
@@ -142,11 +142,7 @@ 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::<bool>(&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 input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
let mut input_ids = self
|
||||
.tokenizer
|
||||
|
||||
@@ -99,6 +99,28 @@ impl Qwen3VLProcessor {
|
||||
})
|
||||
}
|
||||
|
||||
pub fn new_qwen3_5_default(device: &Device, dtype: DType) -> Result<Self> {
|
||||
let img_process_cfg = PreprocessorConfig::qwen3_5_img_default();
|
||||
let video_process_cfg = PreprocessorConfig::qwen3_5_video_default();
|
||||
let image_token = "<|image_pad|>".to_string();
|
||||
let video_token = "<|video_pad|>".to_string();
|
||||
let vision_start_token = "<|vision_start|>".to_string();
|
||||
let vision_end_token = "<|vision_end|>".to_string();
|
||||
Ok(Self {
|
||||
img_process_cfg,
|
||||
video_process_cfg,
|
||||
device: device.clone(),
|
||||
dtype,
|
||||
image_token,
|
||||
video_token,
|
||||
vision_start_token,
|
||||
vision_end_token,
|
||||
fps: 2,
|
||||
min_frames: 4,
|
||||
max_frames: 768,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn extract_vision_info(
|
||||
&self,
|
||||
mes: &ChatCompletionParameters,
|
||||
|
||||
@@ -12,6 +12,10 @@ pub struct TokenizerModel {
|
||||
}
|
||||
|
||||
impl TokenizerModel {
|
||||
pub fn new(tokenizer: Tokenizer) -> Self {
|
||||
Self { tokenizer }
|
||||
}
|
||||
|
||||
pub fn init(path: &str) -> Result<Self> {
|
||||
let path = path.to_string();
|
||||
assert!(
|
||||
@@ -94,12 +98,6 @@ impl TokenizerModel {
|
||||
Ok(token_id)
|
||||
}
|
||||
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_tensor = Tensor::from_slice(&token_id, (1, token_id.len()), device)?;
|
||||
Ok(token_tensor)
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
use std::time::Instant;
|
||||
|
||||
use aha::{
|
||||
chat::ChatCompletionParameters,
|
||||
models::{GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel},
|
||||
};
|
||||
use anyhow::Result;
|
||||
// use candle_core::{Device, quantized::gguf_file};
|
||||
#[test]
|
||||
fn gguf_test() -> Result<()> {
|
||||
// cargo test -r -F cuda --test test_gguf_qwen3_5 gguf_test -- --nocapture
|
||||
// let path = "/home/jhq/.aha/Qwen/Qwen3.5-4B-GGUF/Qwen3.5-4B-Q5_K_M.gguf"; // 有问题
|
||||
// let path = "/home/jhq/.aha/Qwen/Qwen3.5-2B-GGUF/Qwen3.5-2B-Q6_K.gguf";
|
||||
let path = "/home/jhq/.aha/Qwen/Qwen3.5-0.8B-GGUF/Qwen3.5-0.8B-Q4_K_M.gguf";
|
||||
// let mut file = std::fs::File::open(path)?;
|
||||
// let model = gguf_file::Content::read(&mut file)?;
|
||||
// println!("group_count: {:?}", model.metadata.get("qwen35.ssm.group_count"));
|
||||
// println!("time_step_rank: {:?}", model.metadata.get("qwen35.ssm.time_step_rank"));
|
||||
// println!("state_size: {:?}", model.metadata.get("qwen35.ssm.state_size"));
|
||||
// for (key, value) in model.metadata {
|
||||
// if key.contains("tokenizer") {
|
||||
// continue;
|
||||
// }
|
||||
// println!("{key}: {:#?}", value);
|
||||
// }
|
||||
// println!("model: {:?}", model.magic);
|
||||
// println!("generat.type: {:#?}", model.metadata.keys());
|
||||
// println!("tokenizer.ggml.eos_token_id: {:#?}", model.metadata.get("tokenizer.ggml.eos_token_id"));
|
||||
// println!("model: {:#?}", model.tensor_infos.keys());
|
||||
let message = r#"
|
||||
{
|
||||
"model": "qwen3.5",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "你如何看待AI"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
"#;
|
||||
let mes: ChatCompletionParameters = serde_json::from_str(message)?;
|
||||
let i_start = Instant::now();
|
||||
let mut gguf_qwen3_5 = Qwen3_5GenerateModel::init_from_gguf(path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let res = gguf_qwen3_5.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if res.usage.is_some() {
|
||||
let num_token = res.usage.as_ref().unwrap().total_tokens;
|
||||
let duration_secs = i_duration.as_secs_f64();
|
||||
let tps = num_token as f64 / duration_secs;
|
||||
println!("Tokens per second (TPS): {:.2}", tps);
|
||||
}
|
||||
println!("Time elapsed in generate is: {:?}", i_duration);
|
||||
Ok(())
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
},
|
||||
{
|
||||
|
||||
+5
-13
@@ -19,32 +19,24 @@ 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)?;
|
||||
let i_start = Instant::now();
|
||||
let mut qwen3vl = Qwen3_5GenerateModel::init(&model_path, None, None)?;
|
||||
let mut qwen3_5 = Qwen3_5GenerateModel::init(&model_path, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let i_start = Instant::now();
|
||||
let res = qwen3vl.generate(mes)?;
|
||||
let res = qwen3_5.generate(mes)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("generate: \n {:?}", res);
|
||||
if res.usage.is_some() {
|
||||
|
||||
Reference in New Issue
Block a user