add qwen3.5 gguf

This commit is contained in:
jhqxxx
2026-03-14 19:27:01 +08:00
parent 4d544ab68f
commit aad0e62151
15 changed files with 656 additions and 241 deletions
-1
View File
@@ -148,5 +148,4 @@ 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)
} }
} }
-2
View File
@@ -1,2 +0,0 @@
pub mod common;
pub mod qwen3_5;
-87
View File
@@ -1,87 +0,0 @@
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,6 +1,5 @@
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;
@@ -3,13 +3,13 @@ use std::io::{Read, Seek};
use ahash::AHashMap; use ahash::AHashMap;
use anyhow::{Result, anyhow}; use anyhow::{Result, anyhow};
use candle_core::{ use candle_core::{
Device, Device, Tensor,
quantized::{ quantized::{
QMatMul, QTensor, QMatMul, QTensor,
gguf_file::{self, Value}, gguf_file::{self, Value},
}, },
}; };
use candle_nn::RmsNorm; use candle_nn::{Conv1d, Conv1dConfig, Linear, Module, RmsNorm, VarBuilder, linear_b};
use tokenizers::{self, AddedToken, Tokenizer, models::bpe::BPE}; use tokenizers::{self, AddedToken, Tokenizer, models::bpe::BPE};
use crate::tokenizer::TokenizerModel; use crate::tokenizer::TokenizerModel;
@@ -34,7 +34,7 @@ impl<R: Read + Seek> Gguf<R> {
pub fn qmatmul(&mut self, name: &str) -> Result<QMatMul> { pub fn qmatmul(&mut self, name: &str) -> Result<QMatMul> {
let ws = self.ct.tensor(&mut self.reader, name, &self.device)?; let ws = self.ct.tensor(&mut self.reader, name, &self.device)?;
Ok(QMatMul::from_arc(ws.into())?) Ok(QMatMul::from_qtensor(ws)?)
} }
pub fn rms_norm(&mut self, name: &str, eps: f64) -> Result<RmsNorm> { pub fn rms_norm(&mut self, name: &str, eps: f64) -> Result<RmsNorm> {
@@ -51,6 +51,35 @@ impl<R: Read + Seek> Gguf<R> {
Ok(self.ct.tensor(&mut self.reader, name, &self.device)?) 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( pub fn build_tokenizer(
&self, &self,
add_prefix_space: Option<bool>, add_prefix_space: Option<bool>,
@@ -69,7 +98,7 @@ impl<R: Read + Seek> Gguf<R> {
.clone(); .clone();
let vocab: Vec<String> = vocab let vocab: Vec<String> = vocab
.into_iter() .into_iter()
.map(|tokens| tokens.to_string().map(|x| x.clone())) .map(|tokens| tokens.to_string().cloned())
.collect::<Result<Vec<String>, candle_core::Error>>()?; .collect::<Result<Vec<String>, candle_core::Error>>()?;
let mut vocab_map = AHashMap::new(); let mut vocab_map = AHashMap::new();
for (id, token) in vocab.iter().enumerate() { for (id, token) in vocab.iter().enumerate() {
@@ -82,7 +111,7 @@ impl<R: Read + Seek> Gguf<R> {
.clone(); .clone();
let merges: Vec<String> = merges let merges: Vec<String> = merges
.into_iter() .into_iter()
.map(|tokens| tokens.to_string().map(|x| x.clone())) .map(|tokens| tokens.to_string().cloned())
.collect::<Result<Vec<String>, candle_core::Error>>()?; .collect::<Result<Vec<String>, candle_core::Error>>()?;
let merges: Vec<(String, String)> = merges let merges: Vec<(String, String)> = merges
.into_iter() .into_iter()
@@ -119,11 +148,16 @@ impl<R: Read + Seek> Gguf<R> {
let mut add_tokens = vec![]; let mut add_tokens = vec![];
for (id, type_) in token_types.into_iter().enumerate() { for (id, type_) in token_types.into_iter().enumerate() {
if type_ == 3 || type_ == 4 { if type_ == 3
if let Some(token_str) = vocab.get(id) { && let Some(token_str) = vocab.get(id)
let add_token = AddedToken::from(token_str.clone(), true); {
add_tokens.push(add_token); 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.add_special_tokens(&add_tokens);
@@ -134,3 +168,67 @@ impl<R: Read + Seek> Gguf<R> {
} }
} }
} }
#[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)?)
}
}
+2
View File
@@ -7,6 +7,8 @@ use candle_nn::{
embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm, embedding, layer_norm, linear_b, linear_no_bias, ops::sigmoid, rms_norm,
}; };
pub mod gguf;
use crate::{ use crate::{
position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer}, position_embed::rope::{RoPE, apply_rotary_pos_emb, apply_rotary_pos_emb_roformer},
utils::tensor_utils::{prepare_causal_attention_mask, repeat_kv}, utils::tensor_utils::{prepare_causal_attention_mask, repeat_kv},
+3 -3
View File
@@ -11,8 +11,8 @@ use crate::models::qwen3::config::{Qwen3Config, Qwen3GenerationConfig};
use crate::models::qwen3::model::Qwen3Model; use crate::models::qwen3::model::Qwen3Model;
// use crate::models::GenerateStream; // use crate::models::GenerateStream;
use crate::utils::{ use crate::utils::{
build_completion_chunk_response, build_completion_response, extract_metadata_value, build_completion_chunk_response, build_completion_response, find_type_files, get_device,
find_type_files, get_device, get_dtype, get_logit_processor, get_dtype, get_logit_processor,
}; };
use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel}; use crate::{chat_template::ChatTemplate, models::GenerateModel, tokenizer::TokenizerModel};
@@ -66,7 +66,7 @@ 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 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 enable_thinking = extract_metadata_value::<bool>(&mes.metadata, "enable_thinking");
// let mes_render = self // let mes_render = self
+98 -17
View File
@@ -2,7 +2,7 @@ use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse, ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
}; };
use anyhow::{Result, anyhow}; use anyhow::{Result, anyhow};
use candle_core::{DType, Device, Tensor}; use candle_core::{DType, Device, Tensor, quantized::gguf_file};
use candle_nn::VarBuilder; use candle_nn::VarBuilder;
use rocket::async_stream::stream; use rocket::async_stream::stream;
use rocket::futures::Stream; use rocket::futures::Stream;
@@ -11,13 +11,14 @@ use crate::{
chat_template::ChatTemplate, chat_template::ChatTemplate,
models::{ models::{
GenerateModel, GenerateModel,
common::gguf::Gguf,
qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model}, qwen3_5::{config::Qwen3_5Config, model::Qwen3_5Model},
qwen3vl::processor::Qwen3VLProcessor, qwen3vl::processor::Qwen3VLProcessor,
}, },
tokenizer::TokenizerModel, tokenizer::TokenizerModel,
utils::{ utils::{
build_completion_chunk_response, build_completion_response, extract_metadata_value, build_completion_chunk_response, build_completion_response, find_type_files, get_device,
find_type_files, get_device, get_dtype, get_logit_processor, get_dtype, get_logit_processor,
}, },
}; };
@@ -29,10 +30,17 @@ pub struct Qwen3_5GenerateModel<'a> {
device: Device, device: Device,
eos_token_id: u32, eos_token_id: u32,
model_name: String, model_name: String,
repeat_penalty: f32,
repeat_last_n: usize,
} }
impl<'a> Qwen3_5GenerateModel<'a> { impl<'a> Qwen3_5GenerateModel<'a> {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> { 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 chat_template = ChatTemplate::init(path)?;
let tokenizer = TokenizerModel::init(path)?; let tokenizer = TokenizerModel::init(path)?;
let config_path = path.to_string() + "/config.json"; 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 model_list = find_type_files(path, "safetensors")?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
let eos_token_id = cfg.text_config.eos_token_id; 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 { Ok(Self {
chat_template, chat_template,
@@ -53,26 +61,81 @@ impl<'a> Qwen3_5GenerateModel<'a> {
qwen3_5, qwen3_5,
device, device,
eos_token_id, 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> { 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(32768) as u64;
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed); 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 mes_render = self.chat_template.apply_chat_template(&mes)?;
println!("mes_render: {}", mes_render);
// 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;
@@ -92,6 +155,16 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
seqlen_offset, seqlen_offset,
)?; )?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; 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)?; let next_token = logit_processor.sample(&logits)?;
generate.push(next_token); generate.push(next_token);
if next_token == self.eos_token_id { if next_token == self.eos_token_id {
@@ -129,10 +202,6 @@ 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 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
// .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
@@ -148,6 +217,7 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
let video_grid_thw = input.video_grid_thw.as_ref(); let video_grid_thw = input.video_grid_thw.as_ref();
let mut tool_call_id = None; let mut tool_call_id = None;
let mut tool_call_content = String::new(); let mut tool_call_content = String::new();
let mut generate = Vec::new();
for _ in 0..sample_len { for _ in 0..sample_len {
let logits = self.qwen3_5.forward( let logits = self.qwen3_5.forward(
&input_ids, &input_ids,
@@ -158,7 +228,18 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
seqlen_offset, seqlen_offset,
)?; )?;
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; 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)?; let next_token = logit_processor.sample(&logits)?;
generate.push(next_token);
let mut decode_ids = Vec::new(); let mut decode_ids = Vec::new();
if !error_tokens.is_empty() { if !error_tokens.is_empty() {
decode_ids.extend_from_slice(&error_tokens); decode_ids.extend_from_slice(&error_tokens);
+347 -85
View File
@@ -1,5 +1,7 @@
use std::io::{Read, Seek};
use anyhow::{Result, anyhow}; use anyhow::{Result, anyhow};
use candle_core::{D, IndexOp, Tensor}; use candle_core::{D, DType, Device, IndexOp, Tensor, quantized::QMatMul};
use candle_nn::{ use candle_nn::{
Conv1d, Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias, Conv1d, Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias,
ops::sigmoid, rms_norm, ops::sigmoid, rms_norm,
@@ -7,7 +9,11 @@ use candle_nn::{
use crate::{ use crate::{
models::{ 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}, qwen3_5::config::{Qwen3_5Config, Qwen3_5TextConfig},
qwen3vl::model::Qwen3VLVisionModel, qwen3vl::model::Qwen3VLVisionModel,
}, },
@@ -30,6 +36,10 @@ impl Qwen3_5RMSNorm {
Ok(Self { eps, weight }) 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> { pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
let x = xs.to_dtype(candle_core::DType::F32)?; let x = xs.to_dtype(candle_core::DType::F32)?;
let norm_ = x let norm_ = x
@@ -53,6 +63,11 @@ impl Qwen3_5RMSNormGated {
Ok(Self { norm }) 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> { pub fn forward(&self, xs: &Tensor, gate: Option<&Tensor>) -> Result<Tensor> {
let mut xs = self.norm.forward(xs)?; let mut xs = self.norm.forward(xs)?;
if let Some(gate) = gate { if let Some(gate) = gate {
@@ -62,6 +77,7 @@ impl Qwen3_5RMSNormGated {
} }
} }
#[macro_export]
macro_rules! transmute_tensors { macro_rules! transmute_tensors {
($($tensor:expr),*) => { ($($tensor:expr),*) => {
($( ($(
@@ -69,7 +85,7 @@ macro_rules! transmute_tensors {
)*) )*)
}; };
} }
#[macro_export]
macro_rules! right_pad_zero_tensor { macro_rules! right_pad_zero_tensor {
($dim:expr, $pad_size:expr, $($tensor:expr),+) => { ($dim:expr, $pad_size:expr, $($tensor:expr),+) => {
($( ($(
@@ -78,6 +94,7 @@ macro_rules! right_pad_zero_tensor {
}; };
} }
#[macro_export]
macro_rules! reshape_chunk_tensor { macro_rules! reshape_chunk_tensor {
($chunk_size:expr, $($tensor:expr),*) => { ($chunk_size:expr, $($tensor:expr),*) => {
($( ($(
@@ -107,30 +124,30 @@ pub struct Qwen3_5GatedDeltaNet {
dt_bias: Tensor, dt_bias: Tensor,
a_log: Tensor, a_log: Tensor,
norm: Qwen3_5RMSNormGated, norm: Qwen3_5RMSNormGated,
out_proj: Linear, out_proj: ProjKind,
// Z, B, A 投影 // Z, B, A 投影
in_proj_qkv: candle_nn::Linear, in_proj_qkv: ProjKind,
in_proj_z: candle_nn::Linear, in_proj_z: ProjKind,
in_proj_b: candle_nn::Linear, in_proj_b: ProjKind,
in_proj_a: candle_nn::Linear, in_proj_a: ProjKind,
conv_state_cache: Option<Tensor>, conv_state_cache: Option<Tensor>,
recurrent_state_cache: Option<Tensor>, recurrent_state_cache: Option<Tensor>,
} }
impl Qwen3_5GatedDeltaNet { impl Qwen3_5GatedDeltaNet {
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 hidden_size = config.hidden_size; // 1024
let num_v_heads = config.linear_num_value_heads; let num_v_heads = config.linear_num_value_heads; // 16
let num_k_heads = config.linear_num_key_heads; let num_k_heads = config.linear_num_key_heads; // 16
let head_k_dim = config.linear_key_head_dim; let head_k_dim = config.linear_key_head_dim; // 128
let head_v_dim = config.linear_value_head_dim; let head_v_dim = config.linear_value_head_dim; // 128
let key_dim = head_k_dim * num_k_heads; let key_dim = head_k_dim * num_k_heads; // 2048
let value_dim = head_v_dim * num_v_heads; let value_dim = head_v_dim * num_v_heads; // 2048
let conv_kernel_size = config.linear_conv_kernel_dim; let conv_kernel_size = config.linear_conv_kernel_dim; // 4
// let activation = config.hidden_act; // let activation = config.hidden_act;
let layer_norm_epsilon = config.rms_norm_eps; 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( let conv1d = get_conv1d(
vb.pp("conv1d"), vb.pp("conv1d"),
conv_dim, conv_dim,
@@ -146,12 +163,17 @@ impl Qwen3_5GatedDeltaNet {
let a_log = vb.get(num_v_heads, "A_log")?; let a_log = vb.get(num_v_heads, "A_log")?;
let norm = Qwen3_5RMSNormGated::new(vb.pp("norm"), head_v_dim, layer_norm_epsilon)?; 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"))?; 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"))?; 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"))?; 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"))?; 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"))?; let in_proj_a = linear_no_bias(hidden_size, num_v_heads, vb.pp("in_proj_a"))?;
Ok(Self { Ok(Self {
// hidden_size, // hidden_size,
num_v_heads, num_v_heads,
@@ -168,11 +190,64 @@ impl Qwen3_5GatedDeltaNet {
dt_bias, dt_bias,
a_log, a_log,
norm, norm,
out_proj, out_proj: ProjKind::LinearProj(out_proj),
in_proj_qkv, in_proj_qkv: ProjKind::LinearProj(in_proj_qkv),
in_proj_z, in_proj_z: ProjKind::LinearProj(in_proj_z),
in_proj_b, in_proj_b: ProjKind::LinearProj(in_proj_b),
in_proj_a, 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, conv_state_cache: None,
recurrent_state_cache: None, recurrent_state_cache: None,
}) })
@@ -466,7 +541,6 @@ impl Qwen3_5GatedDeltaNet {
&[self.key_dim, self.key_dim, self.value_dim], &[self.key_dim, self.key_dim, self.value_dim],
D::Minus1, D::Minus1,
)?; )?;
let mut query = qkv_split[0].reshape((bs, seq_len, (), self.head_k_dim))?; 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 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))?; let value = qkv_split[2].reshape((bs, seq_len, (), self.head_v_dim))?;
@@ -502,10 +576,10 @@ impl Qwen3_5GatedDeltaNet {
} }
pub struct Qwen3_5Attention { pub struct Qwen3_5Attention {
q_proj: Linear, q_proj: ProjKind,
k_proj: Linear, k_proj: ProjKind,
v_proj: Linear, v_proj: ProjKind,
o_proj: Linear, o_proj: ProjKind,
q_norm: Qwen3_5RMSNorm, q_norm: Qwen3_5RMSNorm,
k_norm: Qwen3_5RMSNorm, k_norm: Qwen3_5RMSNorm,
num_attention_heads: usize, num_attention_heads: usize,
@@ -517,7 +591,7 @@ pub struct Qwen3_5Attention {
} }
impl 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 hidden_size = config.hidden_size;
let num_attention_heads = config.num_attention_heads; let num_attention_heads = config.num_attention_heads;
let head_dim = config.head_dim; 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 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)?; let k_norm = Qwen3_5RMSNorm::new(vb.pp("k_norm"), head_dim, config.rms_norm_eps)?;
Ok(Self { Ok(Self {
q_proj, q_proj: ProjKind::LinearProj(q_proj),
k_proj, k_proj: ProjKind::LinearProj(k_proj),
v_proj, v_proj: ProjKind::LinearProj(v_proj),
o_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, q_norm,
k_norm, k_norm,
num_attention_heads, 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 { pub struct Qwen3_5DecoderLayer {
// hidden_size: usize, // hidden_size: usize,
layer_type: String, layer_type: String,
linear_attn: Option<Qwen3_5GatedDeltaNet>, attn: AttnKind,
self_attn: Option<Qwen3_5Attention>, mlp: GateUpDownMLPGguf,
mlp: GateUpDownMLP,
input_layernorm: Qwen3_5RMSNorm, input_layernorm: Qwen3_5RMSNorm,
post_attention_layernorm: Qwen3_5RMSNorm, post_attention_layernorm: Qwen3_5RMSNorm,
} }
impl Qwen3_5DecoderLayer { 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 hidden_size = config.hidden_size;
let layer_type = config.layer_types[layer_idx].clone(); let layer_type = config.layer_types[layer_idx].clone();
let (linear_attn, self_attn) = if layer_type.eq("linear_attention") { let attn = if layer_type.eq("linear_attention") {
let linear_attn = Qwen3_5GatedDeltaNet::new(vb.pp("linear_attn"), config)?; let attn = Qwen3_5GatedDeltaNet::new_from_vb(vb.pp("linear_attn"), config)?;
(Some(linear_attn), None) AttnKind::LinearAttn(attn)
} else { } else {
let self_attn = Qwen3_5Attention::new(vb.pp("self_attn"), config)?; let attn = Qwen3_5Attention::new_from_vb(vb.pp("self_attn"), config)?;
(None, Some(self_attn)) AttnKind::SelfAttn(attn)
}; };
let mlp = GateUpDownMLP::new( let mlp = GateUpDownMLPGguf::new_from_vb(
vb.pp("mlp"), vb.pp("mlp"),
hidden_size, hidden_size,
config.intermediate_size, config.intermediate_size,
config.hidden_act,
false, false,
None, None,
None, None,
@@ -668,8 +811,36 @@ impl Qwen3_5DecoderLayer {
Ok(Self { Ok(Self {
// hidden_size, // hidden_size,
layer_type, layer_type,
linear_attn, attn,
self_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, mlp,
input_layernorm, input_layernorm,
post_attention_layernorm, post_attention_layernorm,
@@ -685,16 +856,17 @@ impl Qwen3_5DecoderLayer {
) -> Result<Tensor> { ) -> Result<Tensor> {
let residual = xs.clone(); let residual = xs.clone();
let mut xs = self.input_layernorm.forward(xs)?; let mut xs = self.input_layernorm.forward(xs)?;
if self.layer_type.eq("linear_attention") xs = self.attn.forward(&xs, cos, sin, attention_mask)?;
&& let Some(linear_attn) = self.linear_attn.as_mut() // 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() // xs = linear_attn.forward(&xs, attention_mask)?;
&& let Some(cos) = cos // } else if let Some(self_attn) = self.self_attn.as_mut()
&& let Some(sin) = sin // && 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)?;
// }
let residual = xs.add(&residual)?; let residual = xs.add(&residual)?;
xs = self.post_attention_layernorm.forward(&residual)?; xs = self.post_attention_layernorm.forward(&residual)?;
xs = self.mlp.forward(&xs)?; xs = self.mlp.forward(&xs)?;
@@ -703,11 +875,13 @@ impl Qwen3_5DecoderLayer {
} }
pub fn clear_cache(&mut self) { pub fn clear_cache(&mut self) {
if let Some(linear_attn) = self.linear_attn.as_mut() { match &mut self.attn {
linear_attn.clear_cache(); AttnKind::LinearAttn(attn) => {
} attn.clear_cache();
if let Some(self_attn) = self.self_attn.as_mut() { }
self_attn.clear_kv_cache(); AttnKind::SelfAttn(attn) => {
attn.clear_kv_cache();
}
} }
} }
} }
@@ -718,15 +892,17 @@ pub struct Qwen3_5TextModel {
norm: Qwen3_5RMSNorm, norm: Qwen3_5RMSNorm,
rotary_emb: Qwen3VLTextRotaryEmbedding, rotary_emb: Qwen3VLTextRotaryEmbedding,
mrope_section: Vec<usize>, mrope_section: Vec<usize>,
dtype: DType,
} }
impl Qwen3_5TextModel { 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 embed_tokens = embedding(config.vocab_size, config.hidden_size, vb.pp("embed_tokens"))?;
let mut layers = vec![]; let mut layers = vec![];
let vb_layers = vb.pp("layers"); let vb_layers = vb.pp("layers");
for i in 0..config.num_hidden_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); layers.push(layer);
} }
let norm = Qwen3_5RMSNorm::new(vb.pp("norm"), config.hidden_size, config.rms_norm_eps)?; let norm = Qwen3_5RMSNorm::new(vb.pp("norm"), config.hidden_size, config.rms_norm_eps)?;
@@ -740,17 +916,70 @@ impl Qwen3_5TextModel {
norm, norm,
rotary_emb, rotary_emb,
mrope_section: config.rope_parameters.mrope_section.clone(), 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> { pub fn forward(&mut self, inputs_embeds: &Tensor, position_ids: &Tensor) -> Result<Tensor> {
let (b_size, seq_len, _) = inputs_embeds.dims3()?; let (b_size, seq_len, _) = inputs_embeds.dims3()?;
let (cos, sin) = self.rotary_emb.forward( let (cos, sin) =
position_ids, self.rotary_emb
inputs_embeds.dtype(), .forward(position_ids, self.dtype, self.mrope_section.clone())?;
self.mrope_section.clone(),
)?;
let mut xs = inputs_embeds.clone(); let mut xs = inputs_embeds.clone();
let attention_mask: Option<Tensor> = { let attention_mask: Option<Tensor> = {
if seq_len <= 1 { if seq_len <= 1 {
@@ -785,18 +1014,23 @@ impl Qwen3_5TextModel {
} }
pub struct Qwen3_5Model { pub struct Qwen3_5Model {
config: Qwen3_5Config, // config: Qwen3_5Config,
visual: Qwen3VLVisionModel, spatial_merge_size: usize,
image_token_id: u32,
video_token_id: u32,
vision_start_token_id: u32,
visual: Option<Qwen3VLVisionModel>,
language_model: Qwen3_5TextModel, language_model: Qwen3_5TextModel,
lm_head: Linear, lm_head: ProjKind,
rope_deltas: Option<Tensor>, rope_deltas: Option<Tensor>,
} }
impl Qwen3_5Model { 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 vb_m = vb.pp("model");
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?; 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 { let lm_head = if config.tie_word_embeddings {
Linear::new(language_model.embed_tokens.embeddings().clone(), None) Linear::new(language_model.embed_tokens.embeddings().clone(), None)
} else { } else {
@@ -807,10 +1041,36 @@ impl Qwen3_5Model {
)? )?
}; };
Ok(Self { Ok(Self {
config, spatial_merge_size: config.vision_config.spatial_merge_size,
visual, 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, 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, rope_deltas: None,
}) })
} }
@@ -842,10 +1102,10 @@ impl Qwen3_5Model {
None => None, None => None,
}; };
let spatial_merge_size = self.config.vision_config.spatial_merge_size; let spatial_merge_size = self.spatial_merge_size;
let image_token_id = self.config.image_token_id; let image_token_id = self.image_token_id;
let video_token_id = self.config.video_token_id; let video_token_id = self.video_token_id;
let vision_start_token_id = self.config.vision_start_token_id; let vision_start_token_id = self.vision_start_token_id;
let mut mrope_position_deltas = vec![]; let mut mrope_position_deltas = vec![];
if image_grid_thw.is_some() || video_grid_thw.is_some() { if image_grid_thw.is_some() || video_grid_thw.is_some() {
let total_input_ids = input_ids.clone(); 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)?; let mut inputs_embeds = self.language_model.embed_tokens.forward(input_ids)?;
if let Some(pixel_values) = pixel_values if let Some(pixel_values) = pixel_values
&& let Some(image_grid_thw) = image_grid_thw && 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 (image_embeds, _) = visual.forward(pixel_values, image_grid_thw)?;
let vision_mask = get_equal_mask(input_ids, self.config.image_token_id)?; let vision_mask = get_equal_mask(input_ids, self.image_token_id)?;
let n_image_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?; let n_image_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
if n_image_tokens as usize != image_embeds.dim(0)? { if n_image_tokens as usize != image_embeds.dim(0)? {
return Err(anyhow!(format!( return Err(anyhow!(format!(
@@ -1108,9 +1369,10 @@ impl Qwen3_5Model {
} }
if let Some(pixel_values_video) = pixel_values_video if let Some(pixel_values_video) = pixel_values_video
&& let Some(video_grid_thw) = video_grid_thw && 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 (video_embeds, _) = visual.forward(pixel_values_video, video_grid_thw)?;
let vision_mask = get_equal_mask(input_ids, self.config.video_token_id)?; let vision_mask = get_equal_mask(input_ids, self.video_token_id)?;
let n_video_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?; let n_video_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
if n_video_tokens as usize != video_embeds.dim(0)? { if n_video_tokens as usize != video_embeds.dim(0)? {
return Err(anyhow!(format!( return Err(anyhow!(format!(
+30
View File
@@ -18,6 +18,36 @@ pub struct PreprocessorConfig {
pub image_std: Vec<f32>, 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)] #[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct RopeScaling { pub struct RopeScaling {
pub rope_type: String, pub rope_type: String,
+2 -6
View File
@@ -16,8 +16,8 @@ use crate::{
}, },
tokenizer::TokenizerModel, tokenizer::TokenizerModel,
utils::{ utils::{
build_completion_chunk_response, build_completion_response, extract_metadata_value, build_completion_chunk_response, build_completion_response, find_type_files, get_device,
find_type_files, get_device, get_dtype, get_logit_processor, get_dtype, get_logit_processor,
}, },
}; };
@@ -143,10 +143,6 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
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 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
// .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
+22
View File
@@ -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( pub fn extract_vision_info(
&self, &self,
mes: &ChatCompletionParameters, mes: &ChatCompletionParameters,
-2
View File
@@ -85,8 +85,6 @@ impl TokenizerModel {
} }
tokenizer tokenizer
}; };
let len = tokenizer.get_vocab_size(true);
println!("len: {}", len);
Ok(Self { tokenizer }) Ok(Self { tokenizer })
} }
+41 -24
View File
@@ -1,18 +1,32 @@
use aha::{chat::ChatCompletionParameters, gguf_models::qwen3_5::GgufQwen3_5}; use std::time::Instant;
use aha::{
chat::ChatCompletionParameters,
models::{GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel},
};
use anyhow::Result; use anyhow::Result;
use candle_core::{Device, quantized::gguf_file}; // use candle_core::{Device, quantized::gguf_file};
#[test] #[test]
fn gguf_test() -> Result<()> { fn gguf_test() -> Result<()> {
// cargo test -r -F cuda --test test_gguf_qwen3_5 gguf_test -- --nocapture // 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 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 mut file = std::fs::File::open(path)?;
let model = gguf_file::Content::read(&mut file)?; // let model = gguf_file::Content::read(&mut file)?;
let device = Device::new_cuda(0)?; // 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!("model: {:?}", model.magic);
// println!("generat.type: {:#?}", model.metadata.keys()); // println!("generat.type: {:#?}", model.metadata.keys());
// println!("tokenizer.ggml.model: {:#?}", model.metadata.get("tokenizer.ggml.model")); // gpt2 // println!("tokenizer.ggml.eos_token_id: {:#?}", model.metadata.get("tokenizer.ggml.eos_token_id"));
// // println!("model: {:?}", model.tensor_infos); // println!("model: {:#?}", model.tensor_infos.keys());
let message = r#" let message = r#"
{ {
"model": "qwen3.5", "model": "qwen3.5",
@@ -22,26 +36,29 @@ fn gguf_test() -> Result<()> {
"content": [ "content": [
{ {
"type": "text", "type": "text",
"text": "你好啊" "text": "你如何看待AI"
} }
] ]
} }
] ]
} }
"#; "#;
// 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 mes: ChatCompletionParameters = serde_json::from_str(message)?;
let mut gguf_qwen3_5 = GgufQwen3_5::new(&path, None)?; let i_start = Instant::now();
let _ = gguf_qwen3_5.generate(mes)?; 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(()) Ok(())
} }
+3 -3
View File
@@ -22,7 +22,7 @@ fn qwen3_5_generate() -> Result<()> {
"content": [ "content": [
{ {
"type": "text", "type": "text",
"text": "你好啊" "text": "你好啊,你是谁"
} }
] ]
} }
@@ -31,12 +31,12 @@ fn qwen3_5_generate() -> Result<()> {
"#; "#;
let mes: ChatCompletionParameters = serde_json::from_str(message)?; let mes: ChatCompletionParameters = serde_json::from_str(message)?;
let i_start = Instant::now(); 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(); let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration); println!("Time elapsed in load model is: {:?}", i_duration);
let i_start = Instant::now(); let i_start = Instant::now();
let res = qwen3vl.generate(mes)?; let res = qwen3_5.generate(mes)?;
let i_duration = i_start.elapsed(); let i_duration = i_start.elapsed();
println!("generate: \n {:?}", res); println!("generate: \n {:?}", res);
if res.usage.is_some() { if res.usage.is_some() {