diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs index f435113..43dfad7 100644 --- a/src/models/deepseek_ocr/generate.rs +++ b/src/models/deepseek_ocr/generate.rs @@ -42,9 +42,10 @@ impl DeepseekOCRGenerateModel { let device = &get_device(device); let dtype = get_dtype(dtype, &cfg_dtype); let model_name = std::path::Path::new(path) - .file_stem() // 获取文件名主干(不含扩展名) + .file_name() .and_then(|s| s.to_str()) - .unwrap_or("deepseek-ocr"); + .unwrap_or("deepseek-ocr") + .to_string(); let version = if model_name.contains("2") || cfg.vision_config.width.qwen2_0_5b.is_some() { 2usize } else { diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs index a192826..521be97 100644 --- a/src/models/fun_asr_nano/generate.rs +++ b/src/models/fun_asr_nano/generate.rs @@ -78,6 +78,11 @@ impl FunAsrNanoGenerateModel { } let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device); let fun_asr_nano = FunAsrNanoModel::new(vb, &cfg, &llm_cfg)?; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("fun-asr-nano") + .to_string(); Ok(Self { tokenizer, processor, @@ -87,7 +92,7 @@ impl FunAsrNanoGenerateModel { eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, - model_name: "fun-asr-nano".to_string(), + model_name, }) } } diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs index ca723fb..0d783ad 100644 --- a/src/models/glm_asr_nano/generate.rs +++ b/src/models/glm_asr_nano/generate.rs @@ -48,6 +48,11 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; let glm_asr_nano = GlmAsrNanoModel::new(vb, cfg)?; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("glm-asr-nano") + .to_string(); Ok(Self { chat_template, tokenizer, @@ -58,7 +63,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { eos_token_id1: 59246, eos_token_id2: 59253, eos_token_id3: 59255, - model_name: "glm-asr-nano".to_string(), + model_name, }) } } diff --git a/src/models/glm_ocr/generate.rs b/src/models/glm_ocr/generate.rs index 2f072a7..f602d4c 100644 --- a/src/models/glm_ocr/generate.rs +++ b/src/models/glm_ocr/generate.rs @@ -57,7 +57,11 @@ impl GlmOcrGenerateModel { let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: GlmOcrGenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; - + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("glm-ocr") + .to_string(); Ok(Self { // chat_template, tokenizer, @@ -65,7 +69,7 @@ impl GlmOcrGenerateModel { model, device, eos_token_ids: generation_config.eos_token_id.clone(), - model_name: "glm-ocr".to_string(), + model_name, image_token_id: cfg.image_token_id, image_start_token_id: cfg.image_start_token_id, image_end_token_id: cfg.image_end_token_id, diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs index 0bc7539..dbf9e96 100644 --- a/src/models/hunyuan_ocr/generate.rs +++ b/src/models/hunyuan_ocr/generate.rs @@ -52,6 +52,11 @@ impl<'a> HunyuanOCRGenerateModel<'a> { let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: HunyuanOCRGenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("hunyuan_ocr") + .to_string(); Ok(Self { chat_template, tokenizer, @@ -61,7 +66,7 @@ impl<'a> HunyuanOCRGenerateModel<'a> { eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, - model_name: "hunyuan_ocr".to_string(), + model_name, }) } } diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 46cf783..c4b2dd0 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -40,7 +40,11 @@ impl<'a> MiniCPMGenerateModel<'a> { let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; let minicpm = MiniCPMModel::new(vb, cfg)?; - + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("minicpm4") + .to_string(); Ok(MiniCPMGenerateModel { chat_template, tokenizer, @@ -48,7 +52,7 @@ impl<'a> MiniCPMGenerateModel<'a> { device: device.clone(), endoftext_id, im_end_id, - model_name: "minicpm4".to_string(), + model_name, }) } } diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs index f809358..ba13a04 100644 --- a/src/models/paddleocr_vl/generate.rs +++ b/src/models/paddleocr_vl/generate.rs @@ -45,7 +45,11 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; let paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb)?; - + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("paddleocr_vl") + .to_string(); Ok(PaddleOCRVLGenerateModel { chat_template, tokenizer, @@ -54,7 +58,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { cfg, device: device.clone(), end_token_id, - model_name: "paddleocr_vl".to_string(), + model_name, }) } } diff --git a/src/models/qwen2_5vl/config.rs b/src/models/qwen2_5vl/config.rs index 011e80b..3a07d10 100644 --- a/src/models/qwen2_5vl/config.rs +++ b/src/models/qwen2_5vl/config.rs @@ -1,5 +1,7 @@ use candle_nn::Activation; +use crate::models::qwen2::Qwen2Config; + #[derive(Debug, Clone, PartialEq, serde::Deserialize)] pub struct VisionConfig { pub depth: usize, @@ -53,6 +55,27 @@ pub struct Qwen2_5VLConfig { pub vocab_size: usize, } +impl Qwen2_5VLConfig { + pub fn to_qwen2cfg(&self) -> Qwen2Config { + Qwen2Config { + vocab_size: self.vocab_size, + hidden_size: self.hidden_size, + intermediate_size: self.intermediate_size, + num_hidden_layers: self.num_hidden_layers, + num_attention_heads: self.num_attention_heads, + num_key_value_heads: self.num_key_value_heads, + max_position_embeddings: self.max_position_embeddings, + sliding_window: self.sliding_window, + max_window_layers: self.max_window_layers, + tie_word_embeddings: self.tie_word_embeddings, + rope_theta: self.rope_theta, + rms_norm_eps: self.rms_norm_eps, + use_sliding_window: self.use_sliding_window, + hidden_act: self.hidden_act, + } + } +} + pub struct VisionSetting { pub image_factor: u32, pub min_pixels: u32, diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 3bb7224..43b60a3 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -48,7 +48,11 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? }; let qwen2_5_vl = Qwen2_5VLModel::new(cfg, vb)?; - + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("qwen2.5vl") + .to_string(); Ok(Qwen2_5VLGenerateModel { chat_template, tokenizer, @@ -57,7 +61,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { device: device.clone(), endoftext_id, im_end_id, - model_name: "qwen2.5vl".to_string(), + model_name, }) } } diff --git a/src/models/qwen2_5vl/model.rs b/src/models/qwen2_5vl/model.rs index 2cd2ea6..c7185c7 100644 --- a/src/models/qwen2_5vl/model.rs +++ b/src/models/qwen2_5vl/model.rs @@ -4,8 +4,7 @@ use candle_nn::{Init, Linear, Module, RmsNorm, VarBuilder, linear, linear_no_bia use crate::{ models::{ - common::{GateUpDownMLP, eager_attention_forward}, - qwen2_5vl::config::{Qwen2_5VLConfig, RopeScaling}, + common::{GateUpDownMLP, eager_attention_forward}, qwen2::Qwen2DecoderLayer, qwen2_5vl::config::{Qwen2_5VLConfig, RopeScaling} }, position_embed::rope::{ Qwen2_5VLTextRotaryEmbedding, Qwen2_5VisionRotaryEmbedding, apply_rotary_pos_emb, @@ -512,158 +511,11 @@ impl Qwen2_5VLVisionModel { } } -#[derive(Debug, Clone)] -struct Qwen2_5VLTextAttention { - q_proj: Linear, - k_proj: Linear, - v_proj: Linear, - o_proj: Linear, - num_heads: usize, - num_kv_heads: usize, - num_kv_groups: usize, - head_dim: usize, - hidden_size: usize, - kv_cache: Option<(Tensor, Tensor)>, -} - -impl Qwen2_5VLTextAttention { - fn new(cfg: &Qwen2_5VLConfig, vb: VarBuilder) -> Result { - let hidden_size = cfg.hidden_size; - let num_heads = cfg.num_attention_heads; - let num_kv_heads = cfg.num_key_value_heads; - let num_kv_groups = num_heads / num_kv_heads; - let head_dim = hidden_size / num_heads; - let q_proj = linear(hidden_size, num_heads * head_dim, vb.pp("q_proj"))?; - let k_proj = linear(hidden_size, num_kv_heads * head_dim, vb.pp("k_proj"))?; - let v_proj = linear(hidden_size, num_kv_heads * head_dim, vb.pp("v_proj"))?; - let o_proj = linear_no_bias(hidden_size, hidden_size, vb.pp("o_proj"))?; - Ok(Self { - q_proj, - k_proj, - v_proj, - o_proj, - num_heads, - num_kv_heads, - num_kv_groups, - head_dim, - hidden_size, - kv_cache: None, - }) - } - - fn forward( - &mut self, - xs: &Tensor, - cos: &Tensor, - sin: &Tensor, - attention_mask: Option<&Tensor>, - ) -> Result { - let (b_sz, q_len, _) = xs.dims3()?; - let query_states = self.q_proj.forward(xs)?; - let key_states = self.k_proj.forward(xs)?; - let value_states = self.v_proj.forward(xs)?; - let query_states = query_states - .reshape((b_sz, q_len, self.num_heads, self.head_dim))? - .transpose(1, 2)?; - let key_states = key_states - .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? - .transpose(1, 2)?; - let value_states = value_states - .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))? - .transpose(1, 2)?; - let (query_states, key_states) = - apply_rotary_pos_emb(&query_states, &key_states, cos, sin, false)?; - let (key_states, value_states) = match &self.kv_cache { - None => (key_states, value_states), - Some((prev_k, prev_v)) => { - let key_states = Tensor::cat(&[prev_k, &key_states], 2)?; - let value_states = Tensor::cat(&[prev_v, &value_states], 2)?; - (key_states, value_states) - } - }; - - self.kv_cache = Some((key_states.clone(), value_states.clone())); - let scale = 1f64 / f64::sqrt(self.head_dim as f64); - let attn_output = eager_attention_forward( - &query_states, - &key_states, - &value_states, - Some(self.num_kv_groups), - attention_mask, - scale, - )?; - let attn_output = attn_output.reshape((b_sz, q_len, self.hidden_size))?; - let attn_output = attn_output.apply(&self.o_proj)?; - Ok(attn_output) - } - - fn clear_kv_cache(&mut self) { - self.kv_cache = None - } -} - -#[derive(Debug, Clone)] -struct Qwen2_5VLTextDecoderLayer { - self_attn: Qwen2_5VLTextAttention, - mlp: GateUpDownMLP, - input_layernorm: RmsNorm, - post_attention_layernorm: RmsNorm, -} - -impl Qwen2_5VLTextDecoderLayer { - fn new(cfg: &Qwen2_5VLConfig, vb: VarBuilder) -> Result { - let self_attn = Qwen2_5VLTextAttention::new(cfg, vb.pp("self_attn"))?; - let mlp = GateUpDownMLP::new( - vb.pp("mlp"), - cfg.hidden_size, - cfg.intermediate_size, - cfg.hidden_act, - false, - None, - None, - None, - )?; - let input_layernorm = - rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?; - let post_attention_layernorm = rms_norm( - cfg.hidden_size, - cfg.rms_norm_eps, - vb.pp("post_attention_layernorm"), - )?; - Ok(Self { - self_attn, - mlp, - input_layernorm, - post_attention_layernorm, - }) - } - - fn forward( - &mut self, - xs: &Tensor, - cos: &Tensor, - sin: &Tensor, - attention_mask: Option<&Tensor>, - ) -> Result { - let residual = xs; - let xs = self.input_layernorm.forward(xs)?; - let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?; - let xs = (xs + residual)?; - let residual = &xs; - let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?; - let xs = (residual + xs)?; - Ok(xs) - } - - fn clear_kv_cache(&mut self) { - self.self_attn.clear_kv_cache() - } -} #[derive(Debug, Clone)] pub struct Qwen2_5VLTextModel { pub embed_tokens: candle_nn::Embedding, - layers: Vec, + layers: Vec, norm: RmsNorm, rotary_emb: Qwen2_5VLTextRotaryEmbedding, dtype: DType, @@ -677,9 +529,10 @@ impl Qwen2_5VLTextModel { let head_dim = cfg.hidden_size / cfg.num_attention_heads; let rotary_emb = Qwen2_5VLTextRotaryEmbedding::new(head_dim, cfg.rope_theta); let mut layers = Vec::with_capacity(cfg.num_hidden_layers); + let qwen2cfg = cfg.to_qwen2cfg(); let vb_l = vb.pp("layers"); for layer_idx in 0..cfg.num_hidden_layers { - let layer = Qwen2_5VLTextDecoderLayer::new(cfg, vb_l.pp(layer_idx))?; + let layer = Qwen2DecoderLayer::new(&qwen2cfg, vb_l.pp(layer_idx))?; layers.push(layer) } let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?; diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs index f06d85c..94c3ac9 100644 --- a/src/models/qwen3/generate.rs +++ b/src/models/qwen3/generate.rs @@ -42,7 +42,11 @@ impl<'a> Qwen3GenerateModel<'a> { let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; - + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("qwen3") + .to_string(); Ok(Qwen3GenerateModel { chat_template, tokenizer, @@ -51,7 +55,7 @@ impl<'a> Qwen3GenerateModel<'a> { eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, - model_name: "qwen3".to_string(), + model_name, }) } } diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index d366411..742e4d3 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -57,7 +57,11 @@ impl<'a> Qwen3AsrGenerateModel<'a> { let model_list = find_type_files(path, "safetensors")?; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? }; let qwen3_asr = Qwen3ASRModel::new(vb, &cfg)?; - + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("qwen3-asr") + .to_string(); Ok(Self { chat_template, tokenizer, @@ -68,7 +72,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> { eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, - model_name: "qwen3-asr".to_string(), + model_name, }) } } diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index b29f660..59dd5fa 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -49,6 +49,11 @@ impl<'a> Qwen3VLGenerateModel<'a> { let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("qwen3vl") + .to_string(); Ok(Self { chat_template, tokenizer, @@ -58,7 +63,7 @@ impl<'a> Qwen3VLGenerateModel<'a> { eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32, generation_config, - model_name: "qwen3vl".to_string(), + model_name, }) } } diff --git a/src/models/rmbg2_0/generate.rs b/src/models/rmbg2_0/generate.rs index 1822a3d..4c44626 100644 --- a/src/models/rmbg2_0/generate.rs +++ b/src/models/rmbg2_0/generate.rs @@ -41,6 +41,11 @@ impl RMBG2_0Model { Tensor::from_slice(&[0.485, 0.456, 0.406], (3, 1, 1), &device)?.to_dtype(dtype)?; let img_std = Tensor::from_slice(&[0.229, 0.224, 0.225], (3, 1, 1), &device)?.to_dtype(dtype)?; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("rmbg2.0") + .to_string(); Ok(Self { model, h: 1024, @@ -49,7 +54,7 @@ impl RMBG2_0Model { img_std, device, dtype, - model_name: "rmbg2.0".to_string(), + model_name, }) } diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs index 302dae2..1cf1f21 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -62,11 +62,16 @@ impl VoxCPMGenerate { sample_rate: 16000, }, }; - let model_name = if audio_config.sample_rate == 16000 { - "VoxCPM".to_string() - } else { - "VoxCPM1.5".to_string() - }; + let model_name = std::path::Path::new(path) + .file_name() + .and_then(|s| s.to_str()) + .unwrap_or("VoxCPM") + .to_string(); + // let model_name = if audio_config.sample_rate == 16000 { + // "VoxCPM".to_string() + // } else { + // "VoxCPM1.5".to_string() + // }; let audio_vae = AudioVAE::new( vb_vae, audio_config.encoder_dim,