update model_name
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<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)]
|
||||
pub struct Qwen2_5VLTextModel {
|
||||
pub embed_tokens: candle_nn::Embedding,
|
||||
layers: Vec<Qwen2_5VLTextDecoderLayer>,
|
||||
layers: Vec<Qwen2DecoderLayer>,
|
||||
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"))?;
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user