From 363833d8fb1e677dc73864e02b58fe661e34df0a Mon Sep 17 00:00:00 2001 From: jhqxxx <18280426169@163.com> Date: Mon, 30 Mar 2026 20:36:57 +0800 Subject: [PATCH] delete some code --- src/models/deepseek_ocr/generate.rs | 2 +- src/models/glm_asr_nano/generate.rs | 2 +- src/models/glm_ocr/generate.rs | 2 +- src/models/hunyuan_ocr/generate.rs | 2 +- src/models/minicpm4/generate.rs | 2 +- src/models/paddleocr_vl/generate.rs | 2 +- src/models/qwen2_5vl/generate.rs | 2 +- src/models/qwen2_5vl/model.rs | 8 +- src/models/qwen3/generate.rs | 2 +- src/models/qwen3/model.rs | 129 +--------------------------- src/models/qwen3_asr/generate.rs | 2 +- src/models/qwen3vl/generate.rs | 2 +- src/models/rmbg2_0/generate.rs | 2 +- src/models/voxcpm/generate.rs | 2 +- 14 files changed, 19 insertions(+), 142 deletions(-) diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs index 43dfad7..9876c15 100644 --- a/src/models/deepseek_ocr/generate.rs +++ b/src/models/deepseek_ocr/generate.rs @@ -42,7 +42,7 @@ impl DeepseekOCRGenerateModel { let device = &get_device(device); let dtype = get_dtype(dtype, &cfg_dtype); let model_name = std::path::Path::new(path) - .file_name() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("deepseek-ocr") .to_string(); diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs index 0d783ad..0ac2f05 100644 --- a/src/models/glm_asr_nano/generate.rs +++ b/src/models/glm_asr_nano/generate.rs @@ -49,7 +49,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> { 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() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("glm-asr-nano") .to_string(); diff --git a/src/models/glm_ocr/generate.rs b/src/models/glm_ocr/generate.rs index f602d4c..ae2f594 100644 --- a/src/models/glm_ocr/generate.rs +++ b/src/models/glm_ocr/generate.rs @@ -58,7 +58,7 @@ impl GlmOcrGenerateModel { let generation_config: GlmOcrGenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; let model_name = std::path::Path::new(path) - .file_name() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("glm-ocr") .to_string(); diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs index dbf9e96..fd7fa69 100644 --- a/src/models/hunyuan_ocr/generate.rs +++ b/src/models/hunyuan_ocr/generate.rs @@ -53,7 +53,7 @@ impl<'a> HunyuanOCRGenerateModel<'a> { let generation_config: HunyuanOCRGenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; let model_name = std::path::Path::new(path) - .file_name() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("hunyuan_ocr") .to_string(); diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index c4b2dd0..3fa1af9 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -41,7 +41,7 @@ impl<'a> MiniCPMGenerateModel<'a> { 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() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("minicpm4") .to_string(); diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs index ba13a04..699a73f 100644 --- a/src/models/paddleocr_vl/generate.rs +++ b/src/models/paddleocr_vl/generate.rs @@ -46,7 +46,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> { 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() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("paddleocr_vl") .to_string(); diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index 43b60a3..b521833 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -49,7 +49,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { 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() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("qwen2.5vl") .to_string(); diff --git a/src/models/qwen2_5vl/model.rs b/src/models/qwen2_5vl/model.rs index c7185c7..0a20ed1 100644 --- a/src/models/qwen2_5vl/model.rs +++ b/src/models/qwen2_5vl/model.rs @@ -4,11 +4,12 @@ use candle_nn::{Init, Linear, Module, RmsNorm, VarBuilder, linear, linear_no_bia use crate::{ models::{ - common::{GateUpDownMLP, eager_attention_forward}, qwen2::Qwen2DecoderLayer, 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, - apply_rotary_pos_emb_vision, + Qwen2_5VLTextRotaryEmbedding, Qwen2_5VisionRotaryEmbedding, apply_rotary_pos_emb_vision, }, utils::tensor_utils::{ get_equal_mask, get_vision_next_indices, masked_scatter_dim0, nonzero_index, @@ -511,7 +512,6 @@ impl Qwen2_5VLVisionModel { } } - #[derive(Debug, Clone)] pub struct Qwen2_5VLTextModel { pub embed_tokens: candle_nn::Embedding, diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs index 94c3ac9..46da4aa 100644 --- a/src/models/qwen3/generate.rs +++ b/src/models/qwen3/generate.rs @@ -43,7 +43,7 @@ impl<'a> Qwen3GenerateModel<'a> { let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; let model_name = std::path::Path::new(path) - .file_name() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("qwen3") .to_string(); diff --git a/src/models/qwen3/model.rs b/src/models/qwen3/model.rs index bd67077..9541763 100644 --- a/src/models/qwen3/model.rs +++ b/src/models/qwen3/model.rs @@ -1,141 +1,19 @@ use anyhow::Result; use candle_core::Tensor; use candle_nn::{ - Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_b, linear_no_bias, rms_norm, + Embedding, Linear, Module, RmsNorm, VarBuilder, embedding, linear_no_bias, rms_norm, }; use crate::{ models::{ - common::{GateUpDownMLP, QKNormAttention, eager_attention_forward}, + common::{GateUpDownMLP, QKNormAttention}, qwen3::config::Qwen3Config, }, - position_embed::rope::{RoPE, apply_rotary_pos_emb}, + position_embed::rope::RoPE, utils::tensor_utils::prepare_causal_attention_mask, }; -pub struct Qwen3Attention { - q_proj: Linear, - k_proj: Linear, - v_proj: Linear, - o_proj: Linear, - q_norm: RmsNorm, - k_norm: RmsNorm, - num_attention_heads: usize, - num_key_value_heads: usize, - num_kv_groups: usize, - head_dim: usize, - scaling: f64, - kv_cache: Option<(Tensor, Tensor)>, -} - -impl Qwen3Attention { - pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result { - let hidden_size = config.hidden_size; - let num_attention_heads = config.num_attention_heads; - let head_dim = config.head_dim; - let num_key_value_heads = config.num_key_value_heads; - let num_kv_groups = num_attention_heads / num_key_value_heads; - let scaling = 1f64 / f64::sqrt(head_dim as f64); - let q_proj = linear_b( - hidden_size, - num_attention_heads * head_dim, - config.attention_bias, - vb.pp("q_proj"), - )?; - let k_proj = linear_b( - hidden_size, - num_key_value_heads * head_dim, - config.attention_bias, - vb.pp("k_proj"), - )?; - let v_proj = linear_b( - hidden_size, - num_key_value_heads * head_dim, - config.attention_bias, - vb.pp("v_proj"), - )?; - let o_proj = linear_b( - num_attention_heads * head_dim, - hidden_size, - config.attention_bias, - vb.pp("o_proj"), - )?; - let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?; - let k_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("k_norm"))?; - Ok(Self { - q_proj, - k_proj, - v_proj, - o_proj, - q_norm, - k_norm, - num_attention_heads, - num_key_value_heads, - num_kv_groups, - head_dim, - scaling, - kv_cache: None, - }) - } - - pub 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)?.reshape(( - b_sz, - q_len, - self.num_attention_heads, - self.head_dim, - ))?; - let query_states = self.q_norm.forward(&query_states)?.transpose(1, 2)?; - let key_states = self.k_proj.forward(xs)?.reshape(( - b_sz, - q_len, - self.num_key_value_heads, - self.head_dim, - ))?; - let key_states = self.k_norm.forward(&key_states)?.transpose(1, 2)?; - let value_states = self.v_proj.forward(xs)?; - let value_states = value_states - .reshape((b_sz, q_len, self.num_key_value_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 attn_output = eager_attention_forward( - &query_states, - &key_states, - &value_states, - Some(self.num_kv_groups), - attention_mask, - self.scaling, - )?; - let attn_output = - attn_output.reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))?; - let attn_output = attn_output.apply(&self.o_proj)?; - Ok(attn_output) - } - - pub fn clear_kv_cache(&mut self) { - self.kv_cache = None - } -} - pub struct Qwen3DecoderLayer { - // self_attn: Qwen3Attention, self_attn: QKNormAttention, mlp: GateUpDownMLP, input_layernorm: RmsNorm, @@ -144,7 +22,6 @@ pub struct Qwen3DecoderLayer { impl Qwen3DecoderLayer { pub fn new(config: &Qwen3Config, vb: VarBuilder) -> Result { - // let self_attn = Qwen3Attention::new(config, vb.pp("self_attn"))?; let self_attn = QKNormAttention::new( vb.pp("self_attn"), config.hidden_size, diff --git a/src/models/qwen3_asr/generate.rs b/src/models/qwen3_asr/generate.rs index 742e4d3..efb8935 100644 --- a/src/models/qwen3_asr/generate.rs +++ b/src/models/qwen3_asr/generate.rs @@ -58,7 +58,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> { 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() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("qwen3-asr") .to_string(); diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs index 59dd5fa..8544e56 100644 --- a/src/models/qwen3vl/generate.rs +++ b/src/models/qwen3vl/generate.rs @@ -50,7 +50,7 @@ impl<'a> Qwen3VLGenerateModel<'a> { let generation_config: Qwen3GenerationConfig = serde_json::from_slice(&std::fs::read(generation_config_path)?)?; let model_name = std::path::Path::new(path) - .file_name() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("qwen3vl") .to_string(); diff --git a/src/models/rmbg2_0/generate.rs b/src/models/rmbg2_0/generate.rs index 4c44626..8366d71 100644 --- a/src/models/rmbg2_0/generate.rs +++ b/src/models/rmbg2_0/generate.rs @@ -42,7 +42,7 @@ impl RMBG2_0Model { 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() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("rmbg2.0") .to_string(); diff --git a/src/models/voxcpm/generate.rs b/src/models/voxcpm/generate.rs index 1cf1f21..eb82395 100644 --- a/src/models/voxcpm/generate.rs +++ b/src/models/voxcpm/generate.rs @@ -63,7 +63,7 @@ impl VoxCPMGenerate { }, }; let model_name = std::path::Path::new(path) - .file_name() + .file_name() .and_then(|s| s.to_str()) .unwrap_or("VoxCPM") .to_string();