update model_name
This commit is contained in:
@@ -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 {
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,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"))?;
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user