add qwen3vl 4b, 8b, 32b

This commit is contained in:
Ziyi
2025-11-12 03:20:55 -10:00
parent d15cd45315
commit b5cfcb4684
8 changed files with 49 additions and 17 deletions
+24
View File
@@ -21,8 +21,16 @@ pub enum WhichModel {
MiniCPM4_0_5B,
#[value(name = "qwen2.5vl-3b")]
Qwen2_5vl3B,
#[value(name = "qwen2.5vl-7b")]
Qwen2_5vl7B,
#[value(name = "qwen3vl-2b")]
Qwen3vl2B,
#[value(name = "qwen3vl-4b")]
Qwen3vl4B,
#[value(name = "qwen3vl-8b")]
Qwen3vl8B,
#[value(name = "qwen3vl-32b")]
Qwen3vl32B,
}
pub trait GenerateModel {
@@ -84,9 +92,25 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result<ModelInstance<'_
let model = Qwen2_5VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen2_5VL(model)
}
WhichModel::Qwen2_5vl7B => {
let model = Qwen2_5VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen2_5VL(model)
}
WhichModel::Qwen3vl2B => {
let model = Qwen3VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen3VL(model)
}
WhichModel::Qwen3vl4B => {
let model = Qwen3VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen3VL(model)
}
WhichModel::Qwen3vl8B => {
let model = Qwen3VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen3VL(model)
}
WhichModel::Qwen3vl32B => {
let model = Qwen3VLGenerateModel::init(path, None, None)?;
ModelInstance::Qwen3VL(model)
}
};
Ok(model)
-1
View File
@@ -42,7 +42,6 @@ pub struct Qwen3VLTextConfig {
pub rms_norm_eps: f64,
pub rope_scaling: RopeScaling,
pub rope_theta: f32,
pub tie_word_embeddings: bool,
pub use_cache: bool,
pub vocab_size: usize,
}
-1
View File
@@ -47,7 +47,6 @@ impl<'a> Qwen3VLGenerateModel<'a> {
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
let model_list = find_type_files(path, "safetensors")?;
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
let vb = vb.pp("model");
let qwen3_vl = Qwen3VLModel::new(cfg, vb)?;
let generation_config_path = path.to_string() + "/generation_config.json";
let generation_config: Qwen3VLGenerationConfig =
+8 -7
View File
@@ -568,7 +568,7 @@ impl Qwen3VLTextAttention {
pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result<Self> {
let hidden_size = config.hidden_size;
let num_attention_heads = config.num_attention_heads;
let head_dim = hidden_size / 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);
@@ -576,7 +576,7 @@ impl Qwen3VLTextAttention {
let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj = linear(hidden_size, hidden_size, vb.pp("o_proj"))?;
let o_proj = linear(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
} else {
let q_proj =
@@ -585,7 +585,7 @@ impl Qwen3VLTextAttention {
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
let v_proj =
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
let o_proj = linear_no_bias(hidden_size, hidden_size, vb.pp("o_proj"))?;
let o_proj = linear_no_bias(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
(q_proj, k_proj, v_proj, o_proj)
};
let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
@@ -652,7 +652,7 @@ impl Qwen3VLTextAttention {
attention_mask,
self.scaling,
)?;
let attn_output = attn_output.reshape((b_sz, q_len, self.hidden_size))?;
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)
}
@@ -738,7 +738,7 @@ impl Qwen3VLTextModel {
layers.push(layer)
}
let norm = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("norm"))?;
let head_dim = config.hidden_size / config.num_attention_heads;
let head_dim = config.head_dim;
let rotary_emb = Qwen3VLTextRotaryEmbedding::new(head_dim, config.rope_theta);
let mrope_section = config.rope_scaling.mrope_section.clone();
Ok(Self {
@@ -823,10 +823,11 @@ pub struct Qwen3VLModel {
impl Qwen3VLModel {
pub fn new(config: Qwen3VLConfig, vb: VarBuilder) -> Result<Self> {
let vb_m = vb.pp("model");
let config = config.clone();
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb.pp("visual"))?;
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
let language_model =
Qwen3VLTextModel::new(config.text_config.clone(), vb.pp("language_model"))?;
Qwen3VLTextModel::new(config.text_config.clone(), vb_m.pp("language_model"))?;
let lm_head = if config.tie_word_embeddings {
Linear::new(language_model.embed_tokens.embeddings().clone(), None)
} else {