update model_name

This commit is contained in:
jhqxxx
2026-03-30 20:33:48 +08:00
parent db2d3a1eb4
commit ccbe799b90
15 changed files with 106 additions and 175 deletions
+3 -2
View File
@@ -42,9 +42,10 @@ impl DeepseekOCRGenerateModel {
let device = &get_device(device); let device = &get_device(device);
let dtype = get_dtype(dtype, &cfg_dtype); let dtype = get_dtype(dtype, &cfg_dtype);
let model_name = std::path::Path::new(path) let model_name = std::path::Path::new(path)
.file_stem() // 获取文件名主干(不含扩展名) .file_name()
.and_then(|s| s.to_str()) .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() { let version = if model_name.contains("2") || cfg.vision_config.width.qwen2_0_5b.is_some() {
2usize 2usize
} else { } else {
+6 -1
View File
@@ -78,6 +78,11 @@ impl FunAsrNanoGenerateModel {
} }
let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device); let vb = VarBuilder::from_tensors(dict_to_hashmap, dtype, &device);
let fun_asr_nano = FunAsrNanoModel::new(vb, &cfg, &llm_cfg)?; 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 { Ok(Self {
tokenizer, tokenizer,
processor, processor,
@@ -87,7 +92,7 @@ impl FunAsrNanoGenerateModel {
eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config, generation_config,
model_name: "fun-asr-nano".to_string(), model_name,
}) })
} }
} }
+6 -1
View File
@@ -48,6 +48,11 @@ impl<'a> GlmAsrNanoGenerateModel<'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 glm_asr_nano = GlmAsrNanoModel::new(vb, cfg)?; 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 { Ok(Self {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -58,7 +63,7 @@ impl<'a> GlmAsrNanoGenerateModel<'a> {
eos_token_id1: 59246, eos_token_id1: 59246,
eos_token_id2: 59253, eos_token_id2: 59253,
eos_token_id3: 59255, eos_token_id3: 59255,
model_name: "glm-asr-nano".to_string(), model_name,
}) })
} }
} }
+6 -2
View File
@@ -57,7 +57,11 @@ impl GlmOcrGenerateModel {
let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: GlmOcrGenerationConfig = let generation_config: GlmOcrGenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?; 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 { Ok(Self {
// chat_template, // chat_template,
tokenizer, tokenizer,
@@ -65,7 +69,7 @@ impl GlmOcrGenerateModel {
model, model,
device, device,
eos_token_ids: generation_config.eos_token_id.clone(), 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_token_id: cfg.image_token_id,
image_start_token_id: cfg.image_start_token_id, image_start_token_id: cfg.image_start_token_id,
image_end_token_id: cfg.image_end_token_id, image_end_token_id: cfg.image_end_token_id,
+6 -1
View File
@@ -52,6 +52,11 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: HunyuanOCRGenerationConfig = let generation_config: HunyuanOCRGenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?; 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 { Ok(Self {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -61,7 +66,7 @@ impl<'a> HunyuanOCRGenerateModel<'a> {
eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config, generation_config,
model_name: "hunyuan_ocr".to_string(), model_name,
}) })
} }
} }
+6 -2
View File
@@ -40,7 +40,11 @@ impl<'a> MiniCPMGenerateModel<'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 minicpm = MiniCPMModel::new(vb, cfg)?; 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 { Ok(MiniCPMGenerateModel {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -48,7 +52,7 @@ impl<'a> MiniCPMGenerateModel<'a> {
device: device.clone(), device: device.clone(),
endoftext_id, endoftext_id,
im_end_id, im_end_id,
model_name: "minicpm4".to_string(), model_name,
}) })
} }
} }
+6 -2
View File
@@ -45,7 +45,11 @@ impl<'a> PaddleOCRVLGenerateModel<'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 paddleocr_vl = PaddleOCRVLModel::new(cfg.clone(), vb)?; 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 { Ok(PaddleOCRVLGenerateModel {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -54,7 +58,7 @@ impl<'a> PaddleOCRVLGenerateModel<'a> {
cfg, cfg,
device: device.clone(), device: device.clone(),
end_token_id, end_token_id,
model_name: "paddleocr_vl".to_string(), model_name,
}) })
} }
} }
+23
View File
@@ -1,5 +1,7 @@
use candle_nn::Activation; use candle_nn::Activation;
use crate::models::qwen2::Qwen2Config;
#[derive(Debug, Clone, PartialEq, serde::Deserialize)] #[derive(Debug, Clone, PartialEq, serde::Deserialize)]
pub struct VisionConfig { pub struct VisionConfig {
pub depth: usize, pub depth: usize,
@@ -53,6 +55,27 @@ pub struct Qwen2_5VLConfig {
pub vocab_size: usize, 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 struct VisionSetting {
pub image_factor: u32, pub image_factor: u32,
pub min_pixels: u32, pub min_pixels: u32,
+6 -2
View File
@@ -48,7 +48,11 @@ impl<'a> Qwen2_5VLGenerateModel<'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 qwen2_5_vl = Qwen2_5VLModel::new(cfg, vb)?; 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 { Ok(Qwen2_5VLGenerateModel {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -57,7 +61,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> {
device: device.clone(), device: device.clone(),
endoftext_id, endoftext_id,
im_end_id, im_end_id,
model_name: "qwen2.5vl".to_string(), model_name,
}) })
} }
} }
+4 -151
View File
@@ -4,8 +4,7 @@ use candle_nn::{Init, Linear, Module, RmsNorm, VarBuilder, linear, linear_no_bia
use crate::{ use crate::{
models::{ models::{
common::{GateUpDownMLP, eager_attention_forward}, common::{GateUpDownMLP, eager_attention_forward}, qwen2::Qwen2DecoderLayer, qwen2_5vl::config::{Qwen2_5VLConfig, RopeScaling}
qwen2_5vl::config::{Qwen2_5VLConfig, RopeScaling},
}, },
position_embed::rope::{ position_embed::rope::{
Qwen2_5VLTextRotaryEmbedding, Qwen2_5VisionRotaryEmbedding, apply_rotary_pos_emb, 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<Self> {
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<Tensor> {
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<Self> {
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<Tensor> {
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)] #[derive(Debug, Clone)]
pub struct Qwen2_5VLTextModel { pub struct Qwen2_5VLTextModel {
pub embed_tokens: candle_nn::Embedding, pub embed_tokens: candle_nn::Embedding,
layers: Vec<Qwen2_5VLTextDecoderLayer>, layers: Vec<Qwen2DecoderLayer>,
norm: RmsNorm, norm: RmsNorm,
rotary_emb: Qwen2_5VLTextRotaryEmbedding, rotary_emb: Qwen2_5VLTextRotaryEmbedding,
dtype: DType, dtype: DType,
@@ -677,9 +529,10 @@ impl Qwen2_5VLTextModel {
let head_dim = cfg.hidden_size / cfg.num_attention_heads; let head_dim = cfg.hidden_size / cfg.num_attention_heads;
let rotary_emb = Qwen2_5VLTextRotaryEmbedding::new(head_dim, cfg.rope_theta); let rotary_emb = Qwen2_5VLTextRotaryEmbedding::new(head_dim, cfg.rope_theta);
let mut layers = Vec::with_capacity(cfg.num_hidden_layers); let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
let qwen2cfg = cfg.to_qwen2cfg();
let vb_l = vb.pp("layers"); let vb_l = vb.pp("layers");
for layer_idx in 0..cfg.num_hidden_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) layers.push(layer)
} }
let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?; let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("norm"))?;
+6 -2
View File
@@ -42,7 +42,11 @@ impl<'a> Qwen3GenerateModel<'a> {
let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: Qwen3GenerationConfig = let generation_config: Qwen3GenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?; 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 { Ok(Qwen3GenerateModel {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -51,7 +55,7 @@ impl<'a> Qwen3GenerateModel<'a> {
eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config, generation_config,
model_name: "qwen3".to_string(), model_name,
}) })
} }
} }
+6 -2
View File
@@ -57,7 +57,11 @@ impl<'a> Qwen3AsrGenerateModel<'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 qwen3_asr = Qwen3ASRModel::new(vb, &cfg)?; 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 { Ok(Self {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -68,7 +72,7 @@ impl<'a> Qwen3AsrGenerateModel<'a> {
eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config, generation_config,
model_name: "qwen3-asr".to_string(), model_name,
}) })
} }
} }
+6 -1
View File
@@ -49,6 +49,11 @@ impl<'a> Qwen3VLGenerateModel<'a> {
let generation_config_path = path.to_string() + "/generation_config.json"; let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: Qwen3GenerationConfig = let generation_config: Qwen3GenerationConfig =
serde_json::from_slice(&std::fs::read(generation_config_path)?)?; 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 { Ok(Self {
chat_template, chat_template,
tokenizer, tokenizer,
@@ -58,7 +63,7 @@ impl<'a> Qwen3VLGenerateModel<'a> {
eos_token_id1: generation_config.eos_token_id[0] as u32, eos_token_id1: generation_config.eos_token_id[0] as u32,
eos_token_id2: generation_config.eos_token_id[1] as u32, eos_token_id2: generation_config.eos_token_id[1] as u32,
generation_config, generation_config,
model_name: "qwen3vl".to_string(), model_name,
}) })
} }
} }
+6 -1
View File
@@ -41,6 +41,11 @@ impl RMBG2_0Model {
Tensor::from_slice(&[0.485, 0.456, 0.406], (3, 1, 1), &device)?.to_dtype(dtype)?; Tensor::from_slice(&[0.485, 0.456, 0.406], (3, 1, 1), &device)?.to_dtype(dtype)?;
let img_std = let img_std =
Tensor::from_slice(&[0.229, 0.224, 0.225], (3, 1, 1), &device)?.to_dtype(dtype)?; 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 { Ok(Self {
model, model,
h: 1024, h: 1024,
@@ -49,7 +54,7 @@ impl RMBG2_0Model {
img_std, img_std,
device, device,
dtype, dtype,
model_name: "rmbg2.0".to_string(), model_name,
}) })
} }
+10 -5
View File
@@ -62,11 +62,16 @@ impl VoxCPMGenerate {
sample_rate: 16000, sample_rate: 16000,
}, },
}; };
let model_name = if audio_config.sample_rate == 16000 { let model_name = std::path::Path::new(path)
"VoxCPM".to_string() .file_name()
} else { .and_then(|s| s.to_str())
"VoxCPM1.5".to_string() .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( let audio_vae = AudioVAE::new(
vb_vae, vb_vae,
audio_config.encoder_dim, audio_config.encoder_dim,