diff --git a/Cargo.lock b/Cargo.lock index cd6ca1b..ca95210 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2445,9 +2445,8 @@ dependencies = [ [[package]] name = "openai_dive" -version = "1.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bac82ff995cdbe4e120fa87fe4d76bd452866dd7b0c9bcbdb562ba4be0f1786d" +version = "1.3.2" +source = "git+https://github.com/jhqxxx/openai-client.git#83363a83fbd73273e2a1474ca7fdfb1a899e82d4" dependencies = [ "bytes", "derive_builder", diff --git a/Cargo.toml b/Cargo.toml index 6587756..0e35108 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -21,7 +21,7 @@ base64 = "0.22.1" num = "0.4.3" minijinja = "2.12.0" tokenizers = "0.22.1" -openai_dive = { version = "1.3.0", features = ["stream"]} +openai_dive = { git = "https://github.com/jhqxxx/openai-client.git", version = "1.3.0", features = ["stream"]} uuid = { version = "1.18.1", features = ["v4"]} chrono = "0.4.42" rocket = "0.5.1" diff --git a/README.md b/README.md index 15a8f9f..97aea31 100644 --- a/README.md +++ b/README.md @@ -15,6 +15,7 @@ * Qwen2.5VL - 阿里通义千问 2.5 多模态大语言模型 * MiniCPM4 - 面壁智能 MiniCPM 系列语言模型 * VoxCPM - 面壁智能语音生成模型 +* Qwen3VL - 阿里通义千问 3 多模态大语言模型 ## 计划支持 我们持续扩展支持的模型列表,欢迎贡献! @@ -39,8 +40,8 @@ aha = { git = "https://github.com/jhqxxx/aha.git", features = ["cuda", "flash-at git clone https://github.com/jhqxxx/aha.git cd aha # 修改测试用例中模型路径 -# 运行 Qwen2.5VL 示例 -cargo test -F cuda qwen2_5vl_generate -- --nocapture +# 运行 Qwen3VL 示例 +cargo test -F cuda qwen3vl_generate -- --nocapture # 运行 MiniCPM4 示例 cargo test -F cuda minicpm_generate -- --nocapture @@ -91,6 +92,7 @@ fn main() -> Result<()> { │ │ ├── common │ │ ├── minicpm4 │ │ ├── qwen2_5vl +│ │ ├── qwen3vl │ │ ├── voxcpm │ │ └── mod.rs │ ├── position_embed @@ -121,6 +123,9 @@ fn main() -> Result<()> { 2. 提交新的 Issue,包含详细描述和复现步骤 ## 更新日志 +### v0.1.1 +* 添加 Qwen3VL 模型 + ### v0.1.0 * 初始版本发布 * 支持 Qwen2.5VL, MiniCPM4, VoxCPM 模型 diff --git a/assets/video/video_test.mp4 b/assets/video/video_test.mp4 new file mode 100644 index 0000000..47f5ae4 Binary files /dev/null and b/assets/video/video_test.mp4 differ diff --git a/src/lib.rs b/src/lib.rs index 3c90491..49d7fff 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -3,28 +3,3 @@ pub mod models; pub mod position_embed; pub mod tokenizer; pub mod utils; - -// pub enum ModelType { -// Qwen2_5VL, -// MiniCPM4, -// } - -// impl ModelType { -// pub fn init( -// model_type: ModelType, -// model_path: &str, -// device: Option<&Device>, -// dtype: Option, -// ) -> Result> { -// match model_type { -// ModelType::Qwen2_5VL => { -// let model = Qwen2_5VLGenerateModel::init(model_path, device, dtype)?; -// Ok(Box::new(model) as Box) -// }, -// ModelType::MiniCPM4 => { -// let model = MiniCPMGenerateModel::init(model_path, device, dtype)?; -// Ok(Box::new(model)as Box) -// } -// } -// } -// } diff --git a/src/models/common/mod.rs b/src/models/common/mod.rs index 383ff47..786cf0b 100644 --- a/src/models/common/mod.rs +++ b/src/models/common/mod.rs @@ -269,3 +269,54 @@ impl AttentionNobias { self.kv_cache = None } } + +pub fn eager_attention_forward( + query_states: &Tensor, + key_states: &Tensor, + value_states: &Tensor, + num_key_value_groups: Option, + attention_mask: Option<&Tensor>, + scaling: f64, +) -> Result { + let key_states = match num_key_value_groups { + Some(g) => repeat_kv(key_states.clone(), g)?.contiguous()?, + None => key_states.clone() + }; + let value_states = match num_key_value_groups { + Some(g) => repeat_kv(value_states.clone(), g)?.contiguous()?, + None => value_states.clone() + }; + let attn_output = { + #[cfg(not(feature = "flash-attn"))] + { + let attn_weights = query_states.matmul(&key_states.transpose(D::Minus2, D::Minus1)?)?; + let attn_weights = (attn_weights * scaling)?; + let attn_weights = match attention_mask { + None => attn_weights, + Some(mask) => attn_weights.broadcast_add(&mask.to_dtype(attn_weights.dtype())?)?, + }; + let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?; + attn_weights.matmul(&value_states)? + } + #[cfg(feature = "flash-attn")] + { + // use flash-attn, + // flash-attn shape: (bs, seq_len, num_head, head_dim) + let query_states = query_states.transpose(1, 2)?; + let key_states = key_states.transpose(1, 2)?; + let value_states = value_states.transpose(1, 2)?; + let attn_output = candle_flash_attn::flash_attn( + &query_states, + &key_states, + &value_states, + scaling as f32, + attention_mask.is_some(), + )? + .transpose(1, 2)?; + attn_output + } + }; + let attn_output = attn_output.transpose(1, 2)?.contiguous()?; + + Ok(attn_output) +} diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs index 231f176..512041d 100644 --- a/src/models/minicpm4/generate.rs +++ b/src/models/minicpm4/generate.rs @@ -53,7 +53,7 @@ impl<'a> MiniCPMGenerateModel<'a> { impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p); + let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); let mes_render = self.chat_template.apply_chat_template(&mes)?; let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; let mut seq_len = input_ids.dim(1)?; @@ -81,7 +81,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> { &mut self, mes: ChatCompletionParameters, ) -> Result>> { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p); + let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); let mes_render = self.chat_template.apply_chat_template(&mes)?; let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?; let mut seq_len = input_ids.dim(1)?; diff --git a/src/models/mod.rs b/src/models/mod.rs index e6e58c2..4ec507b 100644 --- a/src/models/mod.rs +++ b/src/models/mod.rs @@ -2,6 +2,7 @@ pub mod common; pub mod minicpm4; pub mod qwen2_5vl; pub mod voxcpm; +pub mod qwen3vl; use anyhow::Result; use openai_dive::v1::resources::chat::{ diff --git a/src/models/qwen2_5vl/generate.rs b/src/models/qwen2_5vl/generate.rs index b4be4aa..c189258 100644 --- a/src/models/qwen2_5vl/generate.rs +++ b/src/models/qwen2_5vl/generate.rs @@ -63,7 +63,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> { impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { fn generate(&mut self, mes: ChatCompletionParameters) -> Result { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p); + let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); let mes_render = self.chat_template.apply_chat_template(&mes)?; let input = self.pre_processor.process_info(&mes, &mes_render)?; let mut input_ids = self @@ -118,11 +118,12 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> { let response = build_completion_response(res, "qwen2.5vl"); Ok(response) } + fn generate_stream( &mut self, mes: ChatCompletionParameters, ) -> Result>> { - let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p); + let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None); let mes_render = self.chat_template.apply_chat_template(&mes)?; let input = self.pre_processor.process_info(&mes, &mes_render)?; let mut input_ids = self diff --git a/src/models/qwen2_5vl/model.rs b/src/models/qwen2_5vl/model.rs index 67bf406..27ba304 100644 --- a/src/models/qwen2_5vl/model.rs +++ b/src/models/qwen2_5vl/model.rs @@ -841,7 +841,7 @@ impl Qwen2_5VLTextModel { if seq_len <= 1 { None } else { - Some(&self.prepare_causal_attention_mask(b_size, seq_len, seqlen_offset)?) + Some(&self.prepare_causal_attention_mask(b_size, seq_len, 0)?) } }; for layer in self.layers.iter_mut() { diff --git a/src/models/qwen3vl/config.rs b/src/models/qwen3vl/config.rs new file mode 100644 index 0000000..57d241f --- /dev/null +++ b/src/models/qwen3vl/config.rs @@ -0,0 +1,89 @@ +use candle_nn::Activation; + + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Size { + pub longest_edge: usize, + pub shortest_edge: usize, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct PreprocessorConfig { + pub size: Size, + pub patch_size: usize, + pub temporal_patch_size: usize, + pub merge_size: usize, + pub image_mean: Vec, + pub image_std: Vec, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct RopeScaling { + pub rope_type: String, + pub mrope_section: Vec, + pub mrope_interleaved: bool, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3VLTextConfig { + pub attention_bias: bool, + pub attention_dropout: f32, + pub bos_token_id: usize, + pub dtype: String, + pub eos_token_id: usize, + pub head_dim: usize, + pub hidden_act: Activation, + pub hidden_size: usize, + pub initializer_range: f32, + pub intermediate_size: usize, + pub max_position_embeddings: usize, + pub num_attention_heads: usize, + pub num_hidden_layers: usize, + pub num_key_value_heads: usize, + 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, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3VLVisionConfig { + pub deepstack_visual_indexes: Vec, + pub depth: usize, + pub hidden_act: String, + pub hidden_size: usize, + pub in_channels: usize, + pub initializer_range: f32, + pub intermediate_size: usize, + pub num_heads: usize, + pub num_position_embeddings: usize, + pub out_hidden_size: usize, + pub patch_size: usize, + pub spatial_merge_size: usize, + pub temporal_patch_size: usize, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3VLConfig { + pub image_token_id: usize, + pub text_config: Qwen3VLTextConfig, + pub tie_word_embeddings: bool, + pub video_token_id: usize, + pub vision_config: Qwen3VLVisionConfig, + pub vision_end_token_id: usize, + pub vision_start_token_id: usize, +} + +#[derive(Debug, Clone, PartialEq, serde::Deserialize)] +pub struct Qwen3VLGenerationConfig { + pub bos_token_id: usize, + pub pad_token_id: usize, + pub do_sample: bool, + pub eos_token_id: Vec, + pub top_p: f32, + pub top_k: usize, + pub temperature: f32, + pub repetition_penalty: f32, +} \ No newline at end of file diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs new file mode 100644 index 0000000..122ec71 --- /dev/null +++ b/src/models/qwen3vl/generate.rs @@ -0,0 +1,199 @@ +use anyhow::{Result, anyhow}; +use candle_core::{DType, Device, Tensor}; +use candle_nn::VarBuilder; +use openai_dive::v1::resources::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse}; +use rocket::futures::Stream; +use rocket::async_stream::stream; + +use crate::{ + chat_template::ChatTemplate, + models::{GenerateModel, qwen3vl::{ + config::{Qwen3VLConfig, Qwen3VLGenerationConfig}, + model::Qwen3VLModel, + processor::Qwen3VLProcessor, + }}, + tokenizer::TokenizerModel, + utils::{ + build_completion_chunk_response, build_completion_response, find_type_files, get_device, get_dtype, get_logit_processor + }, +}; + +pub struct Qwen3VLGenerateModel<'a> { + chat_template: ChatTemplate<'a>, + tokenizer: TokenizerModel, + pre_processor: Qwen3VLProcessor, + qwen3_vl: Qwen3VLModel, + device: Device, + eos_token_id1: u32, + eos_token_id2: u32, + generation_config: Qwen3VLGenerationConfig, +} + +impl<'a> Qwen3VLGenerateModel<'a> { + pub fn init(path: &str, device: Option<&Device>, dtype: Option) -> Result { + let chat_template = ChatTemplate::init(path)?; + let tokenizer = TokenizerModel::init(path)?; + let config_path = path.to_string() + "/config.json"; + let cfg: Qwen3VLConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + let device = get_device(device); + let cfg_dtype = cfg.text_config.dtype.as_str(); + let dtype = get_dtype(dtype, cfg_dtype); + 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 = + serde_json::from_slice(&std::fs::read(generation_config_path)?)?; + Ok(Self { + chat_template, + tokenizer, + pre_processor, + qwen3_vl, + device, + eos_token_id1: generation_config.eos_token_id[0] as u32, + eos_token_id2: generation_config.eos_token_id[1] as u32, + generation_config, + }) + } + +} + +impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> { + fn generate(&mut self, mes: ChatCompletionParameters) -> Result { + let temperature = match mes.temperature { + None => self.generation_config.temperature, + Some(tem) => tem, + }; + let top_p = match mes.top_p { + None => self.generation_config.top_p, + Some(top_p) => top_p, + }; + let top_k = self.generation_config.top_k; + let mut logit_processor = get_logit_processor(Some(temperature), Some(top_p), Some(top_k)); + let mes_render = self.chat_template.apply_chat_template(&mes)?; + let input = self.pre_processor.process_info(&mes, &mes_render)?; + let mut input_ids = self + .tokenizer + .text_encode(input.replace_text.clone(), &self.device)?; + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let mut pixel_values = input.pixel_values.as_ref(); + let image_grid_thw = input.image_grid_thw.as_ref(); + let mut pixel_values_video = input.pixel_values_video.as_ref(); + let video_grid_thw = input.video_grid_thw.as_ref(); + let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; + let mut generate = Vec::new(); + let sample_len = mes.max_tokens.unwrap_or(1024); + for _ in 0..sample_len { + let logits = self.qwen3_vl.forward( + &input_ids, + pixel_values, + image_grid_thw, + pixel_values_video, + video_grid_thw, + Some(&cache_position), + seqlen_offset, + )?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + generate.push(next_token); + if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; + pixel_values = None; + pixel_values_video = None; + } + let res = self.tokenizer.token_decode(generate)?; + self.qwen3_vl.clear_kv_cache(); + let response = build_completion_response(res, "qwen3vl"); + Ok(response) + } + + fn generate_stream( + &mut self, + mes: ChatCompletionParameters, + ) -> Result>> { + let temperature = match mes.temperature { + None => self.generation_config.temperature, + Some(tem) => tem, + }; + let top_p = match mes.top_p { + None => self.generation_config.top_p, + Some(top_p) => top_p, + }; + let top_k = self.generation_config.top_k; + let mut logit_processor = get_logit_processor(Some(temperature), Some(top_p), Some(top_k)); + let mes_render = self.chat_template.apply_chat_template(&mes)?; + let input = self.pre_processor.process_info(&mes, &mes_render)?; + let mut input_ids = self + .tokenizer + .text_encode(input.replace_text.clone(), &self.device)?; + let mut seq_len = input_ids.dim(1)?; + let mut seqlen_offset = 0; + let pixel_values = input.pixel_values.clone(); + let image_grid_thw = input.image_grid_thw.clone(); + let pixel_values_video = input.pixel_values_video.clone(); + let video_grid_thw = input.video_grid_thw.clone(); + let mut cache_position = Tensor::arange(0u32, seq_len as u32, &self.device)?; + let sample_len = mes.max_tokens.unwrap_or(1024); + let stream = stream! { + let mut error_tokens = Vec::new(); + let mut pixel_values = pixel_values.as_ref(); + let image_grid_thw = image_grid_thw.as_ref(); + let mut pixel_values_video = pixel_values_video.as_ref(); + let video_grid_thw = video_grid_thw.as_ref(); + for _ in 0..sample_len { + let logits = self.qwen3_vl.forward( + &input_ids, + pixel_values, + image_grid_thw, + pixel_values_video, + video_grid_thw, + Some(&cache_position), + seqlen_offset, + )?; + let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?; + let next_token = logit_processor.sample(&logits)?; + let mut decode_ids = Vec::new(); + if !error_tokens.is_empty() { + decode_ids.extend_from_slice(&error_tokens); + } + decode_ids.push(next_token); + let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?; + if decoded_token.contains("�") { + error_tokens.push(next_token); + if error_tokens.len() > 3 { + error_tokens.clear(); + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; + pixel_values = None; + pixel_values_video = None; + continue; + } + error_tokens.clear(); + let chunk = build_completion_chunk_response(decoded_token, "qwen3vl", None, None); + yield Ok(chunk); + if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 { + break; + } + seqlen_offset += seq_len; + seq_len = 1; + input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?; + cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?; + pixel_values = None; + pixel_values_video = None; + } + self.qwen3_vl.clear_kv_cache(); + }; + Ok(stream) + } +} \ No newline at end of file diff --git a/src/models/qwen3vl/mod.rs b/src/models/qwen3vl/mod.rs new file mode 100644 index 0000000..75b27e4 --- /dev/null +++ b/src/models/qwen3vl/mod.rs @@ -0,0 +1,4 @@ +pub mod processor; +pub mod config; +pub mod model; +pub mod generate; \ No newline at end of file diff --git a/src/models/qwen3vl/model.rs b/src/models/qwen3vl/model.rs new file mode 100644 index 0000000..0f4eb33 --- /dev/null +++ b/src/models/qwen3vl/model.rs @@ -0,0 +1,1251 @@ +use anyhow::{Result, anyhow}; +use candle_core::{D, DType, IndexOp, Shape, Tensor}; +use candle_nn::{ + Activation, Embedding, Init, LayerNorm, LayerNormConfig, Linear, Module, RmsNorm, VarBuilder, + embedding, layer_norm, linear, linear_no_bias, rms_norm, +}; + +use crate::{ + models::{ + common::{MLPNoBias, eager_attention_forward}, + qwen3vl::config::{Qwen3VLConfig, Qwen3VLTextConfig, Qwen3VLVisionConfig}, + }, + position_embed::rope::{ + Qwen2_5VisionRotaryEmbedding, Qwen3VLTextRotaryEmbedding, apply_rotary_pos_emb, + apply_rotary_pos_emb_vision, + }, + utils::tensor_utils::{ + bitor_tensor, get_vision_next_indices, linspace, mask_index_add, masked_scatter_dim0, + nonzero_index, prepare_causal_attention_mask, prod_tensor_last_dim, split_tensor, + zero_index, + }, +}; + +pub struct Qwen3VLVisionMLP { + linear_fc1: Linear, + linear_fc2: Linear, + act_fn: Activation, +} + +impl Qwen3VLVisionMLP { + pub fn new(config: Qwen3VLVisionConfig, vb: VarBuilder) -> Result { + let hidden_size = config.hidden_size; + let intermediate_size = config.intermediate_size; + let linear_fc1 = linear(hidden_size, intermediate_size, vb.pp("linear_fc1"))?; + let linear_fc2 = linear(intermediate_size, hidden_size, vb.pp("linear_fc2"))?; + let act_fn = Activation::GeluPytorchTanh; + Ok(Self { + linear_fc1, + linear_fc2, + act_fn, + }) + } +} + +impl Module for Qwen3VLVisionMLP { + fn forward(&self, xs: &Tensor) -> candle_core::Result { + let xs = xs.apply(&self.linear_fc1)?.apply(&self.act_fn)?; + xs.apply(&self.linear_fc2) + } +} + +pub struct Qwen3VLVisionPatchEmbed { + conv3d_weight: Tensor, + conv3d_bias: Tensor, +} + +impl Qwen3VLVisionPatchEmbed { + pub fn new(cfg: &Qwen3VLVisionConfig, vb: VarBuilder) -> Result { + let patch_size = cfg.patch_size; + let temporal_patch_size = cfg.temporal_patch_size; + let in_channels = cfg.in_channels; + let embed_dim = cfg.hidden_size; + // conv3d weight key: visual.patch_embed.proj.weight, value: Tensor[dims 1024, 3, 2, 16, 16; bf16, cuda:0] + // (1024, 3, 2, 16, 16) -> (1024, 1536) -> (1536, 1024) + let conv3d_weight = vb + .get_with_hints( + ( + embed_dim, + in_channels, + temporal_patch_size, + patch_size, + patch_size, + ), + "proj.weight", + Init::Const(1.), + )? + .flatten(1, 4)? + .t()?; + // (1024) -> (1, 1024) + let conv3d_bias = vb + .get_with_hints((embed_dim,), "proj.bias", Init::Const(0.))? + .unsqueeze(0)?; + Ok(Self { + conv3d_weight, + conv3d_bias, + }) + } + + pub fn forward(&self, hidden_states: &Tensor) -> Result { + // hidden_states shape: (grid_t*grid_h*grid_w, c*temporal_patch_size*patch_size*patch_size) + // ((), 1536) matmul (1536, 1024) -> ((), 1024) + let hidden_states = hidden_states.matmul(&self.conv3d_weight)?; + let hidden_states = hidden_states.broadcast_add(&self.conv3d_bias)?; + Ok(hidden_states) + } +} + +pub struct Qwen3VLVisionPatchMerger { + hidden_size: usize, + use_postshuffle_norm: bool, + norm: LayerNorm, + linear_fc1: Linear, + act_fn: Activation, + linear_fc2: Linear, +} + +impl Qwen3VLVisionPatchMerger { + pub fn new( + config: &Qwen3VLVisionConfig, + vb: VarBuilder, + use_postshuffle_norm: bool, + ) -> Result { + let hidden_size = config.hidden_size * config.spatial_merge_size.pow(2); + let ln_config = LayerNormConfig { + eps: 1e-6, + remove_mean: true, // true for layernorm, false for RMSNorm + affine: true, // true for with bias, false for without bias + }; + let norm_size = if use_postshuffle_norm { + hidden_size + } else { + config.hidden_size + }; + let norm = layer_norm(norm_size, ln_config, vb.pp("norm"))?; + let linear_fc1 = linear(hidden_size, hidden_size, vb.pp("linear_fc1"))?; + let act_fn = Activation::Gelu; + let linear_fc2 = linear(hidden_size, config.out_hidden_size, vb.pp("linear_fc2"))?; + Ok(Self { + hidden_size, + use_postshuffle_norm, + norm, + linear_fc1, + act_fn, + linear_fc2, + }) + } + + pub fn forward(&self, xs: &Tensor) -> Result { + let xs = if self.use_postshuffle_norm { + xs.reshape(((), self.hidden_size))? + } else { + xs.clone() + }; + let xs = self.norm.forward(&xs)?.reshape(((), self.hidden_size))?; + let xs = self + .linear_fc2 + .forward(&self.act_fn.forward(&self.linear_fc1.forward(&xs)?)?)?; + Ok(xs) + } +} + +pub struct Qwen3VLVisionAttention { + num_heads: usize, + qkv: Linear, + proj: Linear, + scaling: f64, +} + +impl Qwen3VLVisionAttention { + pub fn new(config: Qwen3VLVisionConfig, vb: VarBuilder) -> Result { + let hidden_size = config.hidden_size; + let num_heads = config.num_heads; + let head_dim = hidden_size / num_heads; + let qkv = linear(hidden_size, hidden_size * 3, vb.pp("qkv"))?; + let proj = linear(hidden_size, hidden_size, vb.pp("proj"))?; + let scaling = 1.0 / (head_dim as f64).sqrt(); + + Ok(Self { + num_heads, + qkv, + proj, + scaling, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + cu_seqlens: &Tensor, + ) -> Result { + // xs: (seq_len, hidden_size) + let seq_length = xs.dim(0)?; + // (seq_len, hidden_size) -> (seq_len, hidden_size*3) + // -> (seq_len, 3, num_heads, head_dim) + // -> (3, seq_len, num_heads, head_dim) + let qkv_states = xs + .apply(&self.qkv)? + .reshape((seq_length, 3, self.num_heads, ()))? + .permute((1, 0, 2, 3))?; + // (seq_len, num_heads, head_dim) + let query_states = qkv_states.i(0)?.contiguous()?; + let key_states = qkv_states.i(1)?.contiguous()?; + let value_states = qkv_states.i(2)?.contiguous()?; + let (query_states, key_states) = + apply_rotary_pos_emb_vision(&query_states, &key_states, cos, sin)?; + // (seq_len, num_heads, head_dim) -> (num_heads, seq_len, head_dim) -> (1, num_heads, seq_len, head_dim) + let query_states = query_states.transpose(0, 1)?.unsqueeze(0)?.contiguous()?; + let key_states = key_states.transpose(0, 1)?.unsqueeze(0)?.contiguous()?; + let value_states = value_states.transpose(0, 1)?.unsqueeze(0)?.contiguous()?; + let cu_last_id = cu_seqlens.dim(0)? - 1; + let lengths = cu_seqlens.i(1..)?.sub(&cu_seqlens.i(..cu_last_id)?)?; + let chunks: Vec = lengths + .to_vec1::()? + .iter() + .map(|&x| x as usize) + .collect(); + let q_splits = split_tensor(&query_states, &chunks, 2)?; + let k_splits = split_tensor(&key_states, &chunks, 2)?; + let v_splits = split_tensor(&value_states, &chunks, 2)?; + + let mut attn_outputs = Vec::new(); + for (q, (k, v)) in q_splits.iter().zip(k_splits.iter().zip(v_splits.iter())) { + let output = eager_attention_forward(q, k, v, None, None, self.scaling)?; + attn_outputs.push(output); + } + let attn_output = Tensor::cat(&attn_outputs, 1)?; + let attn_output = attn_output.reshape((seq_length, ()))?.contiguous()?; + let attn_ouput = attn_output.apply(&self.proj)?; + Ok(attn_ouput) + } +} + +pub struct Qwen3VLVisionBlock { + norm1: LayerNorm, + norm2: LayerNorm, + attn: Qwen3VLVisionAttention, + mlp: Qwen3VLVisionMLP, +} + +impl Qwen3VLVisionBlock { + pub fn new(config: Qwen3VLVisionConfig, vb: VarBuilder) -> Result { + let ln_config = LayerNormConfig { + eps: 1e-6, + remove_mean: true, // true for layernorm, false for RMSNorm + affine: true, // true for with bias, false for without bias + }; + let norm1 = layer_norm(config.hidden_size, ln_config, vb.pp("norm1"))?; + let norm2 = layer_norm(config.hidden_size, ln_config, vb.pp("norm2"))?; + let attn = Qwen3VLVisionAttention::new(config.clone(), vb.pp("attn"))?; + let mlp = Qwen3VLVisionMLP::new(config, vb.pp("mlp"))?; + Ok(Self { + norm1, + norm2, + attn, + mlp, + }) + } + + pub fn forward( + &self, + xs: &Tensor, + cu_seqlens: &Tensor, + cos: &Tensor, + sin: &Tensor, + ) -> Result { + let residual = xs.clone(); + let xs = self.norm1.forward(xs)?; + let xs = self.attn.forward(&xs, cos, sin, cu_seqlens)?; + let xs = (residual + xs)?; + let residual = xs.clone(); + let xs = self.mlp.forward(&self.norm2.forward(&xs)?)?; + let xs = (residual + xs)?; + Ok(xs) + } +} + +pub struct Qwen3VLVisionModel { + spatial_merge_size: usize, + patch_embed: Qwen3VLVisionPatchEmbed, + pos_embed: Embedding, + num_grid_per_side: u32, + rotary_pos_emb: Qwen2_5VisionRotaryEmbedding, + blocks: Vec, + merger: Qwen3VLVisionPatchMerger, + deepstack_visual_indexes: Vec, + deepstack_merger_list: Vec, + dtype: DType, +} + +impl Qwen3VLVisionModel { + pub fn new(config: Qwen3VLVisionConfig, vb: VarBuilder) -> Result { + let spatial_merge_size = config.spatial_merge_size; + let patch_embed = Qwen3VLVisionPatchEmbed::new(&config, vb.pp("patch_embed"))?; + let pos_embed = embedding( + config.num_position_embeddings, + config.hidden_size, + vb.pp("pos_embed"), + )?; + let num_grid_per_side = (config.num_position_embeddings as f32).sqrt() as u32; + let head_dim = config.hidden_size / config.num_heads; + let rotary_pos_emb = Qwen2_5VisionRotaryEmbedding::new(head_dim / 2, None); + let mut blocks = Vec::new(); + let vb_blocks = vb.pp("blocks"); + for i in 0..config.depth { + let block = Qwen3VLVisionBlock::new(config.clone(), vb_blocks.pp(i))?; + blocks.push(block); + } + let merger = Qwen3VLVisionPatchMerger::new(&config, vb.pp("merger"), false)?; + let deepstack_visual_indexes = config.deepstack_visual_indexes.clone(); + let mut deepstack_merger_list = Vec::new(); + let vb_deepstack = vb.pp("deepstack_merger_list"); + for i in 0..deepstack_visual_indexes.len() { + let merger_i = Qwen3VLVisionPatchMerger::new(&config, vb_deepstack.pp(i), true)?; + deepstack_merger_list.push(merger_i); + } + Ok(Self { + spatial_merge_size, + patch_embed, + pos_embed, + num_grid_per_side, + rotary_pos_emb, + blocks, + merger, + deepstack_visual_indexes, + deepstack_merger_list, + dtype: vb.dtype(), + }) + } + + pub fn fast_pos_embed_interpolate(&self, grid_thw: &Tensor) -> Result { + let mut idx_list = vec![vec![]; 4]; + let mut weight_list = vec![vec![]; 4]; + let mut split_idx = vec![]; + for i in 0..grid_thw.dim(0)? { + let [_, h, w] = grid_thw.i(i)?.to_vec1::()?[..] else { + return Err(anyhow!(format!("grid_thw Expected exactly 3 elements"))); + }; + split_idx.push((h * w) as usize); + let num_grid_per_side = (self.num_grid_per_side - 1) as f32; + let h_idxs = linspace(0.0, num_grid_per_side, h as usize, grid_thw.device())?; + let w_idxs = linspace(0.0, num_grid_per_side, w as usize, grid_thw.device())?; + let h_idxs_floor = h_idxs.to_dtype(candle_core::DType::U32)?; + let w_idxs_floor = w_idxs.to_dtype(candle_core::DType::U32)?; + let h_idxs_ceil = h_idxs_floor + .affine(1.0, 1.0)? + .clamp(0u32, num_grid_per_side as u32)?; + let w_idxs_ceil = w_idxs_floor + .affine(1.0, 1.0)? + .clamp(0u32, num_grid_per_side as u32)?; + + let dh = h_idxs + .sub(&h_idxs_floor.to_dtype(h_idxs.dtype())?)? + .unsqueeze(D::Minus1)?; + let dw = w_idxs + .sub(&w_idxs_floor.to_dtype(h_idxs.dtype())?)? + .unsqueeze(0)?; + + let base_h = h_idxs_floor + .affine(self.num_grid_per_side as f64, 0.0)? + .unsqueeze(D::Minus1)?; + let base_h_ceil = h_idxs_ceil + .affine(self.num_grid_per_side as f64, 0.0)? + .unsqueeze(D::Minus1)?; + idx_list[0].extend_from_slice( + &base_h + .broadcast_add(&w_idxs_floor.unsqueeze(0)?)? + .flatten_all()? + .to_vec1::()?, + ); + idx_list[1].extend_from_slice( + &base_h + .broadcast_add(&w_idxs_ceil.unsqueeze(0)?)? + .flatten_all()? + .to_vec1::()?, + ); + idx_list[2].extend_from_slice( + &base_h_ceil + .broadcast_add(&w_idxs_floor.unsqueeze(0)?)? + .flatten_all()? + .to_vec1::()?, + ); + idx_list[3].extend_from_slice( + &base_h_ceil + .broadcast_add(&w_idxs_ceil.unsqueeze(0)?)? + .flatten_all()? + .to_vec1::()?, + ); + + let one_sub_dh = Tensor::ones_like(&dh)?.sub(&dh)?; + let one_sub_dw = Tensor::ones_like(&dw)?.sub(&dw)?; + + weight_list[0].extend_from_slice( + &one_sub_dh + .broadcast_mul(&one_sub_dw)? + .flatten_all()? + .to_vec1::()?, + ); + weight_list[1].extend_from_slice( + &one_sub_dh + .broadcast_mul(&dw)? + .flatten_all()? + .to_vec1::()?, + ); + weight_list[2].extend_from_slice( + &dh.broadcast_mul(&one_sub_dw)? + .flatten_all()? + .to_vec1::()?, + ); + weight_list[3] + .extend_from_slice(&dh.broadcast_mul(&dw)?.flatten_all()?.to_vec1::()?); + } + let idx_tensor = Tensor::new(idx_list, grid_thw.device())?; + let weight_tensor = Tensor::new(weight_list, grid_thw.device())?.to_dtype(self.dtype)?; + let pos_embeds = self + .pos_embed + .forward(&idx_tensor)? + .broadcast_mul(&weight_tensor.unsqueeze(D::Minus1)?)?; + let patch_pos_embeds = pos_embeds + .i(0)? + .add(&pos_embeds.i(1)?)? + .add(&pos_embeds.i(2)?)? + .add(&pos_embeds.i(3)?)?; + let mut patch_pos_embeds_permute = vec![]; + let patch_pos_embeds = split_tensor(&patch_pos_embeds, &split_idx, 0)?; + let merge_size = self.spatial_merge_size; + for i in 0..grid_thw.dim(0)? { + let [t, h, w] = grid_thw.i(i)?.to_vec1::()?[..] else { + return Err(anyhow!(format!("grid_thw Expected exactly 3 elements"))); + }; + let pos_embed = &patch_pos_embeds[i]; + let pos_emebd_last_dim = pos_embed.dim(D::Minus1)?; + let pos_embed = pos_embed.repeat((t as usize, 1))?; + let shape = Shape::from(vec![ + t as usize, + h as usize / merge_size, + merge_size, + w as usize / merge_size, + merge_size, + pos_emebd_last_dim, + ]); + let pos_embed = pos_embed + .reshape(shape)? + .permute((0, 1, 3, 2, 4, 5))? + .flatten(0, 4)?; + patch_pos_embeds_permute.push(pos_embed); + } + let patch_pos_embeds = Tensor::cat(&patch_pos_embeds_permute, 0)?; + Ok(patch_pos_embeds) + } + + pub fn rot_pos_emb(&self, grid_thw: &Tensor) -> Result { + let merge_size = self.spatial_merge_size; + + let max_hw = grid_thw.i((.., 1..))?.max_all()?.to_scalar::()?; + let freq_table = self + .rotary_pos_emb + .forward(max_hw as usize, grid_thw.device())?; + let mut pos_ids_vec = vec![]; + for i in 0..grid_thw.dim(0)? { + let [t, h, w] = grid_thw.i(i)?.to_vec1::()?[..] else { + return Err(anyhow!(format!("grid_thw Expected exactly 3 elements"))); + }; + let merged_h = h / merge_size as u32; + let merged_w = w / merge_size as u32; + let blocks_rows = Tensor::arange(0, merged_h, grid_thw.device())?; + let blocks_cols = Tensor::arange(0, merged_w, grid_thw.device())?; + let intra_row = Tensor::arange(0, merge_size as u32, grid_thw.device())?; + let intra_col = Tensor::arange(0, merge_size as u32, grid_thw.device())?; + + let row_idx = blocks_rows + .reshape(((), 1, 1, 1))? + .contiguous()? + .affine(merge_size as f64, 0.0)? + .broadcast_add(&intra_row.reshape((1, 1, (), 1))?.contiguous()?)?; + let col_idx = blocks_cols + .reshape((1, (), 1, 1))? + .contiguous()? + .affine(merge_size as f64, 0.0)? + .broadcast_add(&intra_col.reshape((1, 1, 1, ()))?.contiguous()?)?; + let row_idx = row_idx + .expand((merged_h as usize, merged_w as usize, merge_size, merge_size))? + .flatten_all()?; + let col_idx = col_idx + .expand((merged_h as usize, merged_w as usize, merge_size, merge_size))? + .flatten_all()?; + let mut coords = Tensor::stack(&[row_idx, col_idx], D::Minus1)?.contiguous()?; + if t > 1 { + coords = coords.repeat((t as usize, 1))?; + } + pos_ids_vec.push(coords); + } + let pos_ids = Tensor::cat(&pos_ids_vec, 0)?; + let pos_ids_h = pos_ids.i((.., 0))?.contiguous()?; + // 第二列是w维度的索引 + let pos_ids_w = pos_ids.i((.., 1))?.contiguous()?; + let rotary_pos_emb_h = freq_table.index_select(&pos_ids_h, 0)?; + let rotary_pos_emb_w = freq_table.index_select(&pos_ids_w, 0)?; + // 每个patch融合h索引和w索引两个的位置编码信息 + let rotary_pos_emb = Tensor::cat(&[rotary_pos_emb_h, rotary_pos_emb_w], 1)?.contiguous()?; + Ok(rotary_pos_emb) + } + + pub fn forward( + &self, + hidden_states: &Tensor, + grid_thw: &Tensor, + ) -> Result<(Tensor, Vec)> { + let hidden_states = self.patch_embed.forward(hidden_states)?; + let pos_embeds = self.fast_pos_embed_interpolate(grid_thw)?; + let hidden_states = hidden_states.broadcast_add(&pos_embeds)?; + let rotary_pos_emb = self.rot_pos_emb(grid_thw)?; + let seq_len = hidden_states.dim(0)?; + let mut hidden_states = hidden_states.reshape((seq_len, ()))?; + let rotary_pos_emb = rotary_pos_emb.reshape((seq_len, ()))?; + let emb = Tensor::cat(&[&rotary_pos_emb, &rotary_pos_emb], D::Minus1)?; + let cos = emb.cos()?; + let sin = emb.sin()?; + let cu_seqlens = grid_thw.i((.., 1))?.mul(&grid_thw.i((.., 2))?)?; + let grid_t = grid_thw.i((.., 0))?.to_vec1::()?; + let cu_seqlens_full = match cu_seqlens.rank() { + 1 => cu_seqlens.repeat(grid_t[0] as usize)?, + 2 => { + let mut cu_seqlens_repeat = Vec::new(); + for (index, t) in grid_t.iter().enumerate() { + cu_seqlens_repeat.push(cu_seqlens.i(index)?.repeat(*t as usize)?); + } + Tensor::cat(&cu_seqlens_repeat, 0)?.flatten_all()? + } + _ => { + return Err(anyhow!(format!("create cu_seqlens error"))); + } + }; + let cu_seqlens = cu_seqlens_full + .to_dtype(DType::F64)? + .cumsum(0)? + .to_dtype(DType::U32)? + .pad_with_zeros(D::Minus1, 1, 0)?; + let mut deepstack_feature_lists = vec![]; + for (layer_num, block) in self.blocks.iter().enumerate() { + hidden_states = block.forward(&hidden_states, &cu_seqlens, &cos, &sin)?; + if self.deepstack_visual_indexes.contains(&layer_num) { + if let Some(index) = self + .deepstack_visual_indexes + .iter() + .position(|&x| x == layer_num) + { + let deepstack_feature = + self.deepstack_merger_list[index].forward(&hidden_states)?; + deepstack_feature_lists.push(deepstack_feature); + } else { + println!("Value not found"); + } + } + } + hidden_states = self.merger.forward(&hidden_states)?; + Ok((hidden_states, deepstack_feature_lists)) + } +} + +pub struct Qwen3VLTextAttention { + 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, + hidden_size: usize, + scaling: f64, + kv_cache: Option<(Tensor, Tensor)>, +} + +impl Qwen3VLTextAttention { + pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result { + let hidden_size = config.hidden_size; + let num_attention_heads = config.num_attention_heads; + let head_dim = hidden_size / num_attention_heads; + 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, k_proj, v_proj, o_proj) = if config.attention_bias { + 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"))?; + (q_proj, k_proj, v_proj, o_proj) + } else { + let q_proj = + linear_no_bias(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?; + let k_proj = + 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"))?; + (q_proj, k_proj, v_proj, 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, + hidden_size, + 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.hidden_size))?; + 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 Qwen3VLTextDecoderLayer { + self_attn: Qwen3VLTextAttention, + mlp: MLPNoBias, + input_layernorm: RmsNorm, + post_attention_layernorm: RmsNorm, +} + +impl Qwen3VLTextDecoderLayer { + pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result { + let self_attn = Qwen3VLTextAttention::new(config.clone(), vb.pp("self_attn"))?; + let mlp = MLPNoBias::new( + vb.pp("mlp"), + config.hidden_size, + config.intermediate_size, + config.hidden_act, + )?; + let input_layernorm = rms_norm( + config.hidden_size, + config.rms_norm_eps, + vb.pp("input_layernorm"), + )?; + let post_attention_layernorm = rms_norm( + config.hidden_size, + config.rms_norm_eps, + vb.pp("post_attention_layernorm"), + )?; + Ok(Self { + self_attn, + mlp, + input_layernorm, + post_attention_layernorm, + }) + } + + pub fn forward( + &mut self, + xs: &Tensor, + cos: &Tensor, + sin: &Tensor, + attention_mask: Option<&Tensor>, + ) -> Result { + let residual = xs.clone(); + let xs = self.input_layernorm.forward(xs)?; + let xs = self.self_attn.forward(&xs, cos, sin, attention_mask)?; + let xs = residual.add(&xs)?; + let residual = xs.clone(); + let xs = self.post_attention_layernorm.forward(&xs)?; + let xs = self.mlp.forward(&xs)?; + let xs = residual.add(&xs)?; + Ok(xs) + } + + pub fn clear_kv_cache(&mut self) { + self.self_attn.clear_kv_cache(); + } +} + +pub struct Qwen3VLTextModel { + embed_tokens: Embedding, + layers: Vec, + norm: RmsNorm, + rotary_emb: Qwen3VLTextRotaryEmbedding, + mrope_section: Vec, +} + +impl Qwen3VLTextModel { + pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result { + let vocab_size = config.vocab_size; + let embed_tokens = embedding(vocab_size, config.hidden_size, vb.pp("embed_tokens"))?; + let mut layers = vec![]; + let vb_l = vb.pp("layers"); + for layer_idx in 0..config.num_hidden_layers { + let layer = Qwen3VLTextDecoderLayer::new(config.clone(), vb_l.pp(layer_idx))?; + 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 rotary_emb = Qwen3VLTextRotaryEmbedding::new(head_dim, config.rope_theta); + let mrope_section = config.rope_scaling.mrope_section.clone(); + Ok(Self { + embed_tokens, + layers, + norm, + rotary_emb, + mrope_section, + }) + } + + // fn deepstack_process( + // &self, + // xs: &Tensor, + // visual_pos_masks: &Tensor, + // visual_embeds: &Tensor, + // ) -> Result { + // let visual_nonzero_index = nonzero_index(&visual_pos_masks)?; + // let xs = xs.index_add(&visual_nonzero_index, visual_embeds, 0)?; + // Ok(xs) + // } + + pub fn forward( + &mut self, + inputs_embeds: &Tensor, + seqlen_offset: usize, + position_ids: Option<&Tensor>, + visual_pos_masks: Option<&Tensor>, + deepstack_visual_embeds: Option>, + ) -> Result { + let (b_size, seq_len, _) = inputs_embeds.dims3()?; + + let position_ids = match position_ids { + Some(ids) => ids.clone(), + None => Tensor::arange( + seqlen_offset as u32, + (seq_len + seqlen_offset) as u32, + inputs_embeds.device(), + )? + .unsqueeze(0)? + .unsqueeze(0)? + .broadcast_as((3, b_size, seq_len))?, + }; + let (cos, sin) = self.rotary_emb.forward( + &position_ids, + inputs_embeds.dtype(), + self.mrope_section.clone(), + )?; + let mut xs = inputs_embeds.clone(); + let attention_mask: Option<&Tensor> = { + if seq_len <= 1 { + None + } else { + Some(&prepare_causal_attention_mask( + b_size, + seq_len, + 0, + inputs_embeds.device(), + )?) + } + }; + for (layer_idx, layer) in self.layers.iter_mut().enumerate() { + xs = layer.forward(&xs, &cos, &sin, attention_mask)?; + if deepstack_visual_embeds.is_some() + && layer_idx < deepstack_visual_embeds.as_ref().unwrap().len() + { + xs = mask_index_add( + &xs.squeeze(0)?, + &visual_pos_masks.unwrap().squeeze(0)?, + &deepstack_visual_embeds.as_ref().unwrap()[layer_idx], + )? + .unsqueeze(0)?; + } + } + let xs = xs.apply(&self.norm)?; + Ok(xs) + } + + pub fn clear_kv_cache(&mut self) { + for layer in self.layers.iter_mut() { + layer.clear_kv_cache() + } + } +} + +pub struct Qwen3VLModel { + config: Qwen3VLConfig, + visual: Qwen3VLVisionModel, + language_model: Qwen3VLTextModel, + lm_head: Linear, + rope_deltas: Option, +} + +impl Qwen3VLModel { + pub fn new(config: Qwen3VLConfig, vb: VarBuilder) -> Result { + let config = config.clone(); + let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb.pp("visual"))?; + let language_model = + Qwen3VLTextModel::new(config.text_config.clone(), vb.pp("language_model"))?; + let lm_head = if config.tie_word_embeddings { + Linear::new(language_model.embed_tokens.embeddings().clone(), None) + } else { + linear_no_bias( + config.text_config.hidden_size, + config.text_config.vocab_size, + vb.pp("lm_head"), + )? + }; + Ok(Self { + config, + visual, + language_model, + lm_head, + rope_deltas: None, + }) + } + + fn get_vision_features( + &self, + pixel_values: &Tensor, + image_grid_thw: &Tensor, + ) -> Result<(Vec, Vec)> { + let (image_embeds, deepstack_image_embeds) = + self.visual.forward(pixel_values, image_grid_thw)?; + // torch.prod + let split_sizes: Vec = prod_tensor_last_dim(image_grid_thw)? + .to_vec1::()? + .iter() + .map(|&x| x as usize / self.visual.spatial_merge_size.pow(2)) + .collect(); + let image_embeds = split_tensor(&image_embeds, &split_sizes, 0)?; + Ok((image_embeds, deepstack_image_embeds)) + } + + fn get_placeholder_mask(&self, input_ids: &Tensor, is_image: bool) -> Result { + let special_token = if is_image { + Tensor::new(vec![self.config.image_token_id as u32], input_ids.device())? + } else { + Tensor::new(vec![self.config.video_token_id as u32], input_ids.device())? + }; + let special_mask = input_ids + .broadcast_eq(&special_token)? + .to_dtype(candle_core::DType::U32)?; + Ok(special_mask) + } + + fn get_rope_index( + &self, + input_ids: &Tensor, + image_grid_thw: Option<&Tensor>, + video_grid_thw: Option<&Tensor>, + mask: Option<&Tensor>, + ) -> Result<(Tensor, Tensor)> { + let video_grid_thw = match video_grid_thw { + Some(thw) => { + let grid_t = thw.i((.., 0))?.to_vec1::()?; + let mut v_thw_vec = Vec::new(); + for (index, t) in grid_t.iter().enumerate() { + let mut thw_i = thw.i(index)?.to_vec1::()?; + thw_i[0] = 1; + v_thw_vec.push( + Tensor::new(thw_i, thw.device())? + .repeat(*t as usize)? + .reshape((*t as usize, ()))?, + ); + } + Some(&Tensor::cat(&v_thw_vec, 0)?) + } + None => None, + }; + + let spatial_merge_size = self.config.vision_config.spatial_merge_size; + let image_token_id = self.config.image_token_id; + let video_token_id = self.config.video_token_id; + let vision_start_token_id = self.config.vision_start_token_id; + let mut mrope_position_deltas = vec![]; + if image_grid_thw.is_some() || video_grid_thw.is_some() { + let total_input_ids = input_ids.clone(); + let mask_ = mask + .cloned() + .unwrap_or(Tensor::ones_like(&total_input_ids)?) + .to_device(input_ids.device())?; + let mut position_ids = Tensor::ones( + (3, input_ids.dim(0)?, input_ids.dim(1)?), + input_ids.dtype(), + input_ids.device(), + )?; + let mut image_index = 0; + let mut video_index = 0; + + for i in 0..total_input_ids.dim(0)? { + let mut input_ids_i = total_input_ids.i(i)?; + let mask_i = mask_.i(i)?; + // 推理时, attention_mask如果是全1向量,取非0索引的操作没必要 + if mask_i.sum_all()?.to_scalar::()? != mask_i.dim(0)? as u32 { + let nonzero_idx = nonzero_index(&mask_i)?; + input_ids_i = input_ids_i.gather(&nonzero_idx, 0)?; + } + let mut text_start = 0; + let mut text_end = 0; + let mut thw = vec![]; + let mut llm_pos_ids_list: Vec = Vec::new(); + // vision start的下一个索引 + let vision_indices = + get_vision_next_indices(&input_ids_i, vision_start_token_id as u32); + + match vision_indices { + Ok(indeices) => { + let vision_tokens = input_ids_i.gather(&indeices, 0)?.to_vec1::()?; + let vision_indices_vec = indeices.to_vec1::()?; + for (j, &token) in vision_tokens.iter().enumerate() { + if token == image_token_id as u32 { + thw = image_grid_thw.unwrap().i(image_index)?.to_vec1::()?; + image_index += 1; + text_end = vision_indices_vec[j]; + } + if token == video_token_id as u32 { + thw = video_grid_thw.unwrap().i(video_index)?.to_vec1::()?; + text_end = vision_indices_vec[j]; + video_index += 1; + } + let llm_grid_t = thw[0]; + let llm_grid_h = thw[1] / spatial_merge_size as u32; + let llm_grid_w = thw[2] / spatial_merge_size as u32; + let text_len = text_end - text_start; + let start_idx = if !llm_pos_ids_list.is_empty() { + llm_pos_ids_list[llm_pos_ids_list.len() - 1] + .max_all()? + .to_scalar::()? + + 1 + } else { + 0 + }; + let pos_ids = Tensor::arange( + start_idx, + start_idx + text_len, + input_ids_i.device(), + )? + .unsqueeze(0)? + .broadcast_as((3usize, text_len as usize))?; + llm_pos_ids_list.push(pos_ids); + + let t_index = Tensor::arange( + start_idx + text_len, + start_idx + text_len + llm_grid_t, + input_ids_i.device(), + )? + .unsqueeze(D::Minus1)? + .broadcast_as(( + llm_grid_t as usize, + llm_grid_h as usize * llm_grid_w as usize, + ))? + .flatten_all()?; + let h_index = Tensor::arange( + start_idx + text_len, + start_idx + text_len + llm_grid_h, + input_ids_i.device(), + )? + .unsqueeze(0)? + .unsqueeze(D::Minus1)? + .broadcast_as(( + llm_grid_t as usize, + llm_grid_h as usize, + llm_grid_w as usize, + ))? + .flatten_all()?; + let w_index = Tensor::arange( + start_idx + text_len, + start_idx + text_len + llm_grid_w, + input_ids_i.device(), + )? + .unsqueeze(0)? + .unsqueeze(0)? + .broadcast_as(( + llm_grid_t as usize, + llm_grid_h as usize, + llm_grid_w as usize, + ))? + .flatten_all()?; + + let thw_index = Tensor::stack(&[t_index, h_index, w_index], 0)?; + llm_pos_ids_list.push(thw_index); + text_start = text_end + llm_grid_t * llm_grid_h * llm_grid_w; + } + } + Err(e) => { + println!("get vision_indices err: {}", e); + } + }; + if text_start < input_ids_i.dim(0)? as u32 { + let start_idx = if !llm_pos_ids_list.is_empty() { + llm_pos_ids_list[llm_pos_ids_list.len() - 1] + .max_all()? + .to_scalar::()? + + 1 + } else { + 0 + }; + let text_len = input_ids_i.dim(0)? as u32 - text_start; + let pos_ids = + Tensor::arange(start_idx, start_idx + text_len, input_ids_i.device())? + .unsqueeze(0)? + .broadcast_as((3usize, text_len as usize))?; + llm_pos_ids_list.push(pos_ids); + } + let llm_position = Tensor::cat(&llm_pos_ids_list, 1)?.reshape((3, 1, ()))?; + position_ids = position_ids + .slice_assign(&[(0..3), (i..i + 1), (0..input_ids.dim(1)?)], &llm_position)?; + let position_deltas = llm_position.max_all()?.to_scalar::()? as i64 + 1 + - input_ids_i.dim(0)? as i64; + mrope_position_deltas.push(position_deltas); + } + let mut mrope_position_deltas = Tensor::new(mrope_position_deltas, input_ids.device())?; + if mrope_position_deltas.rank() == 1 { + mrope_position_deltas = mrope_position_deltas.unsqueeze(0)?; + } + Ok((position_ids.contiguous()?, mrope_position_deltas)) + } else if let Some(mask) = mask { + let mut position_ids = mask + .to_dtype(candle_core::DType::F64)? + .cumsum(D::Minus1)? + .to_dtype(candle_core::DType::U32)? + .broadcast_sub(&Tensor::new(vec![1_u32], input_ids.device())?)?; + for i in 0..position_ids.dim(0)? { + let mut position_ids_i = position_ids.i(i)?; + let mask_i = mask.i(i)?; + // 如果有pad, 将填充位置置为1 + // 当bs>1, 可能存在不同序列长度,需要添加pad使seq_len长度一致 + if mask_i.sum_all()?.to_scalar::()? != mask_i.dim(0)? as u32 { + let zero_indices = zero_index(&mask_i)?; + let replace_1 = Tensor::ones( + zero_indices.dim(0)?, + candle_core::DType::U32, + input_ids.device(), + )?; + position_ids_i = position_ids_i + .scatter(&zero_indices, &replace_1, 0)? + .unsqueeze(0)?; + position_ids = position_ids + .slice_assign(&[(i..i + 1), (0..position_ids.dim(1)?)], &position_ids_i)?; + } + } + position_ids = position_ids + .unsqueeze(0)? + .broadcast_as((3, input_ids.dim(0)?, input_ids.dim(1)?))? + .contiguous()?; + let mut mrope_position_deltas = position_ids + .max(0)? + .max(D::Minus1)? + .broadcast_sub(&Tensor::new( + vec![mask.dim(D::Minus1)? as u32 - 1], + input_ids.device(), + )?)? + .contiguous()?; + if mrope_position_deltas.rank() == 1 { + mrope_position_deltas = mrope_position_deltas.unsqueeze(0)?; + } + Ok((position_ids, mrope_position_deltas)) + } else { + let position_ids = + Tensor::arange(0_u32, input_ids.dim(D::Minus1)? as u32, input_ids.device())? + .unsqueeze(0)? + .unsqueeze(0)? + .broadcast_as((3, input_ids.dim(0)?, input_ids.dim(D::Minus1)?))? + .contiguous()?; + let mrope_position_deltas = Tensor::zeros( + (input_ids.dim(0)?, 1), + input_ids.dtype(), + input_ids.device(), + )?; + Ok((position_ids, mrope_position_deltas)) + } + } + + pub fn forward( + &mut self, + input_ids: &Tensor, + pixel_values: Option<&Tensor>, + image_grid_thw: Option<&Tensor>, + pixel_values_video: Option<&Tensor>, + video_grid_thw: Option<&Tensor>, + cache_position: Option<&Tensor>, + seqlen_offset: usize, + ) -> Result { + let mut inputs_embeds = self.language_model.embed_tokens.forward(input_ids)?; + let mut image_mask = None; + let mut video_mask = None; + let mut deepstack_image_embeds = None; + let mut deepstack_video_embeds = None; + if let Some(pixel_values) = pixel_values + && let Some(image_grid_thw) = image_grid_thw + { + let (image_embeds, deepstack_img_embed) = + self.get_vision_features(pixel_values, image_grid_thw)?; + let image_embeds = Tensor::cat(&image_embeds, 0)?; + let vision_mask = self.get_placeholder_mask(&input_ids, true)?; + let n_image_tokens = vision_mask.sum_all()?.to_scalar::()?; + if n_image_tokens as usize != image_embeds.dim(0)? { + return Err(anyhow!(format!( + "n_image_token num: {} not equal to image_embed len: {}", + n_image_tokens, + image_embeds.dim(0)? + ))); + } + inputs_embeds = masked_scatter_dim0(&inputs_embeds, &image_embeds, &vision_mask)?; + image_mask = Some(vision_mask); + deepstack_image_embeds = Some(deepstack_img_embed); + } + if let Some(pixel_values_video) = pixel_values_video + && let Some(video_grid_thw) = video_grid_thw + { + let (video_embeds, deepstack_video_embed) = + self.get_vision_features(pixel_values_video, video_grid_thw)?; + let video_embeds = Tensor::cat(&video_embeds, 0)?; + let vision_mask = self.get_placeholder_mask(&input_ids, false)?; + let n_video_tokens = vision_mask.sum_all()?.to_scalar::()?; + if n_video_tokens as usize != video_embeds.dim(0)? { + return Err(anyhow!(format!( + "n_image_token num: {} not equal to image_embed len: {}", + n_video_tokens, + video_embeds.dim(0)? + ))); + } + inputs_embeds = masked_scatter_dim0(&inputs_embeds, &video_embeds, &vision_mask)?; + video_mask = Some(vision_mask); + deepstack_video_embeds = Some(deepstack_video_embed); + } + let mut visual_pos_mask = None; + let mut deepstack_visual_embeds = None; + if image_mask.is_some() && video_mask.is_some() { + let image_mask_ = image_mask.unwrap(); + let video_mask_ = video_mask.unwrap(); + let visual_mask = bitor_tensor(&image_mask_, &video_mask_)?; + let visual_none_zero_index = nonzero_index(&visual_mask)?; + let image_mask_joint = image_mask_.gather(&visual_none_zero_index, 0)?; + let image_nonzero_joint = nonzero_index(&image_mask_joint)?; + let video_mask_joint = video_mask_.gather(&visual_none_zero_index, 0)?; + let video_nonzero_joint = nonzero_index(&video_mask_joint)?; + let mut deepstack_embeds = vec![]; + let visual_len = visual_none_zero_index.dim(0)?; + for (img_embed, vid_embed) in deepstack_image_embeds + .unwrap() + .iter() + .zip(deepstack_video_embeds.unwrap().iter()) + { + let embed_joint = Tensor::zeros( + (visual_len, img_embed.dim(D::Minus1)?), + img_embed.dtype(), + img_embed.device(), + )?; + let embed_joint = embed_joint.index_add(&image_nonzero_joint, &img_embed, 0)?; + let embed_joint = embed_joint.index_add(&video_nonzero_joint, &vid_embed, 0)?; + deepstack_embeds.push(embed_joint); + } + visual_pos_mask = Some(visual_mask); + deepstack_visual_embeds = Some(deepstack_embeds); + } else if image_mask.is_some() { + visual_pos_mask = image_mask; + deepstack_visual_embeds = deepstack_image_embeds; + } else if video_mask.is_some() { + visual_pos_mask = video_mask; + deepstack_visual_embeds = deepstack_video_embeds; + } + + let position_ids; + let rope_deltas; + if (cache_position.is_some() && cache_position.unwrap().i(0)?.to_scalar::()? == 0) + || self.rope_deltas.is_none() + { + (position_ids, rope_deltas) = + self.get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)?; + self.rope_deltas = Some(rope_deltas); + } else { + let (bs, seq_len, _) = inputs_embeds.dims3()?; + let delta = if let Some(cache_position) = cache_position { + cache_position + .i(0)? + .to_dtype(self.rope_deltas.as_ref().unwrap().dtype())? + .broadcast_add(self.rope_deltas.as_ref().unwrap())? + .contiguous()? + .to_dtype(candle_core::DType::U32)? + } else { + Tensor::zeros(1, inputs_embeds.dtype(), inputs_embeds.device())? + }; + position_ids = Tensor::arange(0u32, seq_len as u32, input_ids.device())? + .unsqueeze(0)? + .broadcast_as((bs, seq_len))? + .broadcast_add(&delta)? + .unsqueeze(0)? + .broadcast_as((3, bs, seq_len))? + .contiguous()?; + } + let outputs = self.language_model.forward( + &inputs_embeds, + seqlen_offset, + Some(&position_ids), + visual_pos_mask.as_ref(), + deepstack_visual_embeds, + )?; + let seq_len = outputs.dim(1)?; + let hidden_state = outputs.narrow(1, seq_len - 1, 1)?; + let logits = self.lm_head.forward(&hidden_state)?; + Ok(logits) + } + + pub fn clear_kv_cache(&mut self) { + self.language_model.clear_kv_cache(); + } +} diff --git a/src/models/qwen3vl/processor.rs b/src/models/qwen3vl/processor.rs new file mode 100644 index 0000000..2e50526 --- /dev/null +++ b/src/models/qwen3vl/processor.rs @@ -0,0 +1,625 @@ +use std::collections::HashMap; + +use anyhow::{Result, anyhow}; +use candle_core::{DType, Device, IndexOp, Shape, Tensor}; +use ffmpeg_next as ffmpeg; +use image::DynamicImage; +use num::integer::lcm; +use openai_dive::v1::resources::chat::{ + ChatCompletionParameters, ChatMessage, ChatMessageContent, ChatMessageContentPart, +}; + +use crate::{ + models::qwen3vl::config::PreprocessorConfig, + utils::{ceil_by_factor, floor_by_factor, img_utils::get_image, round_by_factor}, +}; + +#[derive(Clone)] +pub struct VisionInput { + pub data: Tensor, + pub grid_thw: Tensor, +} + +#[derive(Clone)] +pub struct GeneralInput { + pub replace_text: String, + pub pixel_values: Option, + pub image_grid_thw: Option, + pub pixel_values_video: Option, + pub video_grid_thw: Option, +} + +#[allow(unused)] +#[derive(Debug, Clone)] +pub struct VideoMetadata { + total_num_frames: u32, + fps: f32, + width: u32, + height: u32, + duration: f32, + frame_indices: Vec, +} + +pub struct Qwen3VLProcessor { + img_process_cfg: PreprocessorConfig, + video_process_cfg: PreprocessorConfig, + device: Device, + dtype: DType, + image_token: String, + video_token: String, + vision_start_token: String, + vision_end_token: String, + fps: u32, + min_frames: u32, + max_frames: u32, +} + +impl Qwen3VLProcessor { + pub fn new(path: &str, device: &Device, dtype: DType) -> Result { + let path = path.to_string(); + assert!( + std::path::Path::new(&path).exists(), + "model path file not exists" + ); + let img_process_cfg_file = path.clone() + "/preprocessor_config.json"; + assert!( + std::path::Path::new(&img_process_cfg_file).exists(), + "preprocessor_config.json not exists in model path" + ); + let img_process_cfg: PreprocessorConfig = + serde_json::from_slice(&std::fs::read(img_process_cfg_file)?)?; + + let video_process_cfg_file = path.clone() + "/video_preprocessor_config.json"; + assert!( + std::path::Path::new(&video_process_cfg_file).exists(), + "video_preprocessor_config.json not exists in model path" + ); + let video_process_cfg: PreprocessorConfig = + serde_json::from_slice(&std::fs::read(video_process_cfg_file)?)?; + + let image_token = "<|image_pad|>".to_string(); + let video_token = "<|video_pad|>".to_string(); + let vision_start_token = "<|vision_start|>".to_string(); + let vision_end_token = "<|vision_end|>".to_string(); + Ok(Self { + img_process_cfg, + video_process_cfg, + device: device.clone(), + dtype, + image_token, + video_token, + vision_start_token, + vision_end_token, + fps: 2, + min_frames: 4, + max_frames: 768, + }) + } + + pub fn extract_vision_info( + &self, + mes: &ChatCompletionParameters, + ) -> Result>> { + let mut vision_map = HashMap::new(); + vision_map.insert("image".to_string(), Vec::new()); + vision_map.insert("video".to_string(), Vec::new()); + for chat_mes in mes.messages.clone() { + if let ChatMessage::User { content, .. } = chat_mes + && let ChatMessageContent::ContentPart(part_vec) = content + { + for part in part_vec { + if let ChatMessageContentPart::Image(img_part) = part { + let img_url = img_part.image_url; + vision_map.get_mut("image").unwrap().push(img_url.url); + } else if let ChatMessageContentPart::Video(video_part) = part { + let video_url = video_part.video_url; + vision_map.get_mut("video").unwrap().push(video_url.url); + } + } + } + } + Ok(vision_map) + } + + pub fn process_img( + &self, + img: &DynamicImage, + img_mean: &Tensor, + img_std: &Tensor, + ) -> Result { + let img_h = img.height(); + let img_w = img.width(); + // h,w resize成 28的倍数 + let (resize_h, resize_w) = img_smart_resize( + img_h, + img_w, + (self.img_process_cfg.patch_size * self.img_process_cfg.merge_size) as u32, + self.img_process_cfg.size.shortest_edge as u32, + self.img_process_cfg.size.longest_edge as u32, + None, + )?; + let img = img.resize_exact(resize_w, resize_h, image::imageops::FilterType::CatmullRom); + let img_vec = img.to_rgb8().into_raw(); + // (h, w, c) => (c, h, w) + let img_tensor = Tensor::from_slice( + &img_vec, + (resize_h as usize, resize_w as usize, 3), + &self.device, + )? + .permute((2, 0, 1))? + .to_dtype(self.dtype)?; + // 0-255 rescale to 0-1 + let img_tensor = img_tensor.affine(1.0 / 255.0, 0.)?; + // normalize + let img_tensor = img_tensor.broadcast_sub(img_mean)?.broadcast_div(img_std)?; + // (c, h, w) => (1, c, h, w) + let img_tensor = img_tensor.unsqueeze(0)?; + Ok(img_tensor) + } + + pub fn process_vision_tensor(&self, img_tensor: &Tensor) -> Result<(Tensor, Tensor)> { + let channel = img_tensor.dim(1)?; + let grid_t = img_tensor.dim(0)? / self.img_process_cfg.temporal_patch_size; + let grid_h = img_tensor.dim(2)? / self.img_process_cfg.patch_size; + let grid_w = img_tensor.dim(3)? / self.img_process_cfg.patch_size; + let shape = Shape::from(vec![ + grid_t, + self.img_process_cfg.temporal_patch_size, + channel, + grid_h / self.img_process_cfg.merge_size, + self.img_process_cfg.merge_size, + self.img_process_cfg.patch_size, + grid_w / self.img_process_cfg.merge_size, + self.img_process_cfg.merge_size, + self.img_process_cfg.patch_size, + ]); + let img_tensor = img_tensor.reshape(shape)?; + // shape to // grid_t, + // grid_h / merge_size, + // grid_w / merge_size, + // merge_size, + // merge_size, + // channel, + // temporal_patch_size, + // patch_size, + // patch_size, + let img_tensor = img_tensor.permute(vec![0, 3, 6, 4, 7, 2, 1, 5, 8])?; + let img_tensor = img_tensor + .reshape(( + grid_t * grid_h * grid_w, + channel + * self.img_process_cfg.temporal_patch_size + * self.img_process_cfg.patch_size + * self.img_process_cfg.patch_size, + ))? + .contiguous()?; + let grid_thw = Tensor::from_vec( + vec![grid_t as u32, grid_h as u32, grid_w as u32], + (1, 3), + &self.device, + )?; + Ok((img_tensor, grid_thw)) + } + + pub fn process_images( + &self, + imgs: Vec, + img_mean: &Tensor, + img_std: &Tensor, + ) -> Result { + let mut pixel_values_vec = Vec::new(); + let mut vision_grid_thws_vec = Vec::new(); + + for img in imgs { + let img_tensor = self.process_img(&img, img_mean, img_std)?; + let img_tensor = Tensor::cat(&[&img_tensor, &img_tensor], 0)?.contiguous()?; + let (img_tensor, grid_thw) = self.process_vision_tensor(&img_tensor)?; + pixel_values_vec.push(img_tensor); + vision_grid_thws_vec.push(grid_thw); + } + let pixel_values = Tensor::cat(&pixel_values_vec, 0)?; + let vision_grid_thws = Tensor::cat(&vision_grid_thws_vec, 0)?; + Ok(VisionInput { + data: pixel_values, + grid_thw: vision_grid_thws, + }) + } + + pub fn process_videos( + &self, + data: Vec, + img_mean: &Tensor, + img_std: &Tensor, + ) -> Result { + let mut pixel_values_vec = Vec::new(); + let mut vision_grid_thws_vec = Vec::new(); + for single_video in data { + // 0-255 rescale to 0-1 + let video_tensor = single_video.to_dtype(self.dtype)?.affine(1.0 / 255.0, 0.)?; + // normalize + let video_tensor = video_tensor + .broadcast_sub(img_mean)? + .broadcast_div(img_std)? + .contiguous()?; + let (video_tensor, video_grid_thw) = self.process_vision_tensor(&video_tensor)?; + pixel_values_vec.push(video_tensor); + vision_grid_thws_vec.push(video_grid_thw); + } + let pixel_values = Tensor::cat(&pixel_values_vec, 0)?.contiguous()?; + let vision_grid_thws = Tensor::cat(&vision_grid_thws_vec, 0)?.contiguous()?; + Ok(VisionInput { + data: pixel_values, + grid_thw: vision_grid_thws, + }) + } + + fn calculate_timestamps( + &self, + frames_indices: Vec, + fps: f32, + t_merge_size: usize, + ) -> Result> { + let indices = if frames_indices.len() % t_merge_size != 0 { + let mut frames_indices = frames_indices.clone(); + let last = frames_indices[frames_indices.len() - 1]; + let pad_len = t_merge_size - frames_indices.len() % t_merge_size; + for _ in 0..pad_len { + frames_indices.push(last); + } + frames_indices + } else { + frames_indices.clone() + }; + let timestamps: Vec = indices.iter().map(|&x| x as f32 / fps).collect(); + let mut stamps = Vec::new(); + for i in (0..timestamps.len()).step_by(t_merge_size) { + let stamp = (timestamps[i] + timestamps[i + t_merge_size - 1]) / 2.0; + stamps.push(stamp); + } + Ok(stamps) + } + + pub fn process_info( + &self, + messages: &ChatCompletionParameters, + text: &str, + ) -> Result { + let mut pixel_values = None; + let mut image_grid_thw = None; + let mut pixel_values_video = None; + let mut video_grid_thw: Option = None; + let mut video_metadata = None; + let vision_map = self.extract_vision_info(messages)?; + let img_mean = + Tensor::from_slice(&self.img_process_cfg.image_mean, (3, 1, 1), &self.device)? + .to_dtype(self.dtype)?; + let img_std = Tensor::from_slice(&self.img_process_cfg.image_std, (3, 1, 1), &self.device)? + .to_dtype(self.dtype)?; + for (key, vec) in vision_map { + // println!("key: {}, \nvalue: {:?}", key, vec); + if key.eq("image") { + let mut file_vec = Vec::new(); + for file in &vec { + let image = get_image(file); + match image { + Ok(img) => file_vec.push(img), + Err(e) => println!("get_image err: {:?}", e), + }; + } + if !file_vec.is_empty() { + let vision_input = self.process_images(file_vec, &img_mean, &img_std); + match vision_input { + Ok(img_input) => { + pixel_values = Some(img_input.data); + image_grid_thw = Some(img_input.grid_thw); + } + Err(e) => println!("img process_images err: {:?}", e), + }; + } + } + if key.eq("video") { + let mut file_vec = Vec::new(); + let mut video_infos = Vec::new(); + for file in &vec { + let video_data = get_video_data( + file, + self.video_process_cfg.patch_size as u32, + self.video_process_cfg.temporal_patch_size as u32, + self.video_process_cfg.merge_size as u32, + self.fps, + self.min_frames, + self.max_frames, + self.video_process_cfg.size.shortest_edge as u32, + self.video_process_cfg.size.longest_edge as u32, + &self.device, + ); + match video_data { + Ok((tensor, video_info)) => { + file_vec.push(tensor); + video_infos.push(video_info); + } + Err(e) => println!("get_video_data err: {:?}", e), + }; + } + if !file_vec.is_empty() { + let vision_input = self.process_videos(file_vec, &img_mean, &img_std); + match vision_input { + Ok(video_input) => { + pixel_values_video = Some(video_input.data); + video_grid_thw = Some(video_input.grid_thw); + video_metadata = Some(video_infos); + } + Err(e) => println!("video process_videos err: {:?}", e), + }; + } + } + } + let merge_length = self.img_process_cfg.merge_size.pow(2); + let mut text = text.to_string(); + if let Some(ref image_grid_thw) = image_grid_thw { + let mut index = 0; + while text.contains(&self.image_token) { + let grid_i = image_grid_thw.i(index)?; + let repeat_num = + grid_i.to_vec1::()?.iter().product::() as usize / merge_length; + let replace = "<|placeholder|>".repeat(repeat_num); + text = text.replacen(&self.image_token, &replace, 1); + index += 1; + } + text = text.replace("<|placeholder|>", &self.image_token); + } + if let Some(ref video_grid_thw) = video_grid_thw { + let mut index = 0; + while text.contains(&self.video_token) { + let grid_i = video_grid_thw.i(index)?; + let video_info = &video_metadata.as_ref().unwrap()[index]; + let curr_timestamp = self.calculate_timestamps( + video_info.frame_indices.clone(), + video_info.fps, + self.img_process_cfg.merge_size, + )?; + let mut video_placeholder = "".to_string(); + let [t, h, w] = grid_i.to_vec1::()?[..] else { + return Err(anyhow!(format!("grid_thw Expected exactly 3 elements"))); + }; + let frame_seqlen = h * w / merge_length as u32; + for frame_idx in 0..t { + let curr_time = curr_timestamp[frame_idx as usize]; + video_placeholder = video_placeholder + format!("<{:.1} seconds>", curr_time).as_str(); + video_placeholder = video_placeholder + self.vision_start_token.as_str(); + video_placeholder = video_placeholder + "<|placeholder|>".repeat(frame_seqlen as usize).as_str(); + video_placeholder = video_placeholder + self.vision_end_token.as_str(); + } + let three_token = format!("{}{}{}", self.vision_start_token, self.video_token, self.vision_end_token); + if text.contains(&three_token) { + text = text.replacen(&three_token, &video_placeholder, 1); + } else { + text = text.replacen(&self.video_token, &video_placeholder, 1); + } + index += 1; + } + text = text.replace("<|placeholder|>", &self.video_token); + } + let input = GeneralInput { + replace_text: text, + pixel_values, + image_grid_thw, + pixel_values_video, + video_grid_thw, + }; + Ok(input) + } +} + +pub fn img_smart_resize( + img_h: u32, + img_w: u32, + factor: u32, + min_pixels: u32, + max_pixels: u32, + video_ratio: Option, +) -> Result<(u32, u32)> { + if std::cmp::max(img_h, img_w) / std::cmp::min(img_h, img_w) > 200 { + return Err(anyhow!(format!( + "absolute aspect ratio mush be smaller than {}, got {}", + 200, + std::cmp::max(img_h, img_w) / std::cmp::min(img_h, img_w) + ))); + } + let mut image_factor = factor; + if let Some(ratio) = video_ratio { + image_factor = lcm(image_factor, ratio); + } + let mut h_bar = std::cmp::max(image_factor, round_by_factor(img_h, image_factor)); + let mut w_bar = std::cmp::max(image_factor, round_by_factor(img_w, image_factor)); + + if h_bar * w_bar > max_pixels { + let beta = ((img_h * img_w) as f32 / max_pixels as f32).sqrt(); + h_bar = floor_by_factor(img_h as f32 / beta, image_factor); + w_bar = floor_by_factor(img_w as f32 / beta, image_factor); + } else if h_bar * w_bar < min_pixels { + let beta = (min_pixels as f32 / (img_h * img_w) as f32).sqrt(); + h_bar = ceil_by_factor(img_h as f32 * beta, image_factor); + w_bar = ceil_by_factor(img_w as f32 * beta, image_factor); + } + Ok((h_bar, w_bar)) +} + +pub fn video_smart_resize( + num_frames: u32, + height: u32, + width: u32, + temporal_factor: u32, + factor: u32, + min_pixels: u32, + max_pixels: u32, + video_ratio: Option, +) -> Result<(u32, u32)> { + if num_frames < temporal_factor { + return Err(anyhow!(format!( + "{} must be larger than temporal_factor {}", + num_frames, temporal_factor + ))); + } + if height < factor || width < factor { + return Err(anyhow!(format!( + "height:{} or width:{} must be larger than factor:{}", + height, width, factor + ))); + } + if std::cmp::max(height, width) / std::cmp::min(height, width) > 200 { + return Err(anyhow!(format!( + "absolute aspect ratio mush be smaller than {}, got {}", + 200, + std::cmp::max(height, width) / std::cmp::min(height, width) + ))); + } + let mut image_factor = factor; + if let Some(ratio) = video_ratio { + image_factor = lcm(image_factor, ratio); + } + let mut h_bar = round_by_factor(height, image_factor); + let mut w_bar = round_by_factor(width, image_factor); + let t_bar = round_by_factor(num_frames, temporal_factor); + if t_bar * h_bar * w_bar > max_pixels { + let beta = ((num_frames * height * width) as f32 / max_pixels as f32).sqrt(); + h_bar = std::cmp::max( + image_factor, + floor_by_factor(height as f32 / beta, image_factor), + ); + w_bar = std::cmp::max( + image_factor, + floor_by_factor(width as f32 / beta, image_factor), + ); + } else if t_bar * h_bar * w_bar < min_pixels { + let beta = (min_pixels as f32 / (num_frames * height * width) as f32).sqrt(); + h_bar = ceil_by_factor(height as f32 * beta, image_factor); + w_bar = ceil_by_factor(width as f32 * beta, image_factor); + } + Ok((h_bar, w_bar)) +} + +pub fn get_video_data( + file: &String, + patch_size: u32, + temporal_patch_size: u32, + merge_size: u32, + fps: u32, + min_frames: u32, + max_frames: u32, + min_pixels: u32, + max_pixels: u32, + device: &Device, +) -> Result<(Tensor, VideoMetadata)> { + ffmpeg::init().map_err(|e| anyhow!(format!("Failed to initialize ffmpeg: {}", e)))?; + + let mut ictx = ffmpeg::format::input(&file) + .map_err(|e| anyhow!(format!("Failed to open video file: {}", e)))?; + let input = ictx + .streams() + .best(ffmpeg::media::Type::Video) + .ok_or_else(|| anyhow!(format!("No video stream found")))?; + let video_stream_index = input.index(); + let context_decoder = ffmpeg::codec::context::Context::from_parameters(input.parameters()) + .map_err(|e| anyhow!(format!("Failed to crate decoder context: {}", e)))?; + let mut decoder = context_decoder + .decoder() + .video() + .map_err(|e| anyhow!(format!("Failed to decoder video: {}", e)))?; + + let video_h = decoder.height(); + let video_w = decoder.width(); + let format = decoder.format(); + + let frames = input.frames(); + let rate = input.rate().0 as f32 / input.rate().1 as f32; + let duration = frames as f32 * 1.0 / rate; + // 1s取两帧 + let nframes = (frames as f32 / rate * fps as f32).round() as u32; + let nframes = std::cmp::min( + std::cmp::min(std::cmp::max(nframes, min_frames), max_frames), + frames as u32, + ); + let sample_interval = (frames as f32 / nframes as f32).round() as u32; + let mut frame_indices = Vec::new(); + let mut frame_id = 0_u32; + + // 图片帧使用scaler reshape的时候需要保证宽高是16的倍数,不然reshape出来的是损坏的图片 + // 所以计算resize的目标宽高时,需要用16和image_factor的最小公倍数 + let (resize_h, resize_w) = video_smart_resize( + nframes, + video_h, + video_w, + temporal_patch_size, + patch_size * merge_size, + min_pixels, + max_pixels, + Some(16), + )?; + let mut scaler = ffmpeg::software::scaling::context::Context::get( + format, + video_w, + video_h, + ffmpeg::format::Pixel::RGB24, + resize_w, + resize_h, + ffmpeg::software::scaling::flag::Flags::BILINEAR + | ffmpeg::software::scaling::flag::Flags::ACCURATE_RND, + ) + .map_err(|e| anyhow!(format!("Failed to crate scaler: {}", e)))?; + + let mut frames_vec = Vec::new(); + let mut receive_and_process_decoded_frames = + |decoder: &mut ffmpeg::decoder::Video| -> Result<()> { + let mut decoded = ffmpeg::frame::Video::empty(); + while decoder.receive_frame(&mut decoded).is_ok() { + if frame_id.is_multiple_of(sample_interval) { + frame_indices.push(frame_id); + let mut rgb_frame = ffmpeg::frame::Video::empty(); + scaler + .run(&decoded, &mut rgb_frame) + .map_err(|e| anyhow!(format!("Failed to scaler run decoded: {}", e)))?; + + // save_file(&rgb_frame, frame_id as usize); + let frame_data = rgb_frame.data(0); + let frame_tensor = Tensor::from_slice( + frame_data, + (resize_h as usize, resize_w as usize, 3), + device, + )? + .permute((2, 0, 1))?; + frames_vec.push(frame_tensor); + } + frame_id += 1; + } + Ok(()) + }; + + for (stream, packet) in ictx.packets() { + if stream.index() == video_stream_index { + decoder + .send_packet(&packet) + .map_err(|e| anyhow!(format!("Failed to send packet: {}", e)))?; + receive_and_process_decoded_frames(&mut decoder)?; + } + } + decoder + .send_eof() + .map_err(|e| anyhow!(format!("Failed to decoder.send_eof(): {}", e)))?; + receive_and_process_decoded_frames(&mut decoder)?; + + if frames_vec.is_empty() { + return Err(anyhow!("No frames extracted from video".to_string())); + } + // (t, c, h, w) + let frames_tensor = Tensor::stack(&frames_vec, 0)?.contiguous()?; + let video_info = VideoMetadata { + total_num_frames: frames as u32, + fps: rate, + width: video_w, + height: video_h, + duration, + frame_indices, + }; + Ok((frames_tensor, video_info)) +} diff --git a/src/position_embed/rope.rs b/src/position_embed/rope.rs index d4eba9f..2e9f4bb 100644 --- a/src/position_embed/rope.rs +++ b/src/position_embed/rope.rs @@ -64,7 +64,8 @@ pub fn apply_rotary_pos_emb_vision( // cos, sin -> (seq_len, head_dim) -> (seq_len, 1, head_dim) let cos = cos.unsqueeze(D::Minus2)?; let sin = sin.unsqueeze(D::Minus2)?; - + let cos = cos.to_dtype(q.dtype())?; + let sin = sin.to_dtype(q.dtype())?; let q_embed = q .broadcast_mul(&cos)? .add(&rotate_half(q)?.broadcast_mul(&sin)?)?; @@ -197,3 +198,75 @@ impl Qwen2_5VisionRotaryEmbedding { Ok(freqs) } } + +#[derive(Debug, Clone)] +pub struct Qwen3VLTextRotaryEmbedding { + inv_freq: Vec, +} + +impl Qwen3VLTextRotaryEmbedding { + pub fn new(dim: usize, theta_base: f32) -> Self { + let inv_freq = compute_default_rope_parameters(dim, theta_base); + Self { inv_freq } + } + + pub fn apply_interleaved_mrope( + &self, + freqs: &Tensor, + mrope_section: Vec, + ) -> Result { + let mut freqs_t = freqs.i(0)?.contiguous()?; //(3, bs, seq_len, head_dim //2) -> (bs, seq_len, head_dim //2) + + for dim in 1..3 { + let length = mrope_section[dim] * 3; + let idx = Tensor::arange_step(dim as u32, length as u32, 3, freqs.device())?; + let src = freqs.i(dim)?.contiguous()?; // (bs, seq_len, head_dim //2) + let src = src.index_select(&idx, D::Minus1)?.contiguous()?; + let idx = idx + .unsqueeze(0)? + .unsqueeze(0)? + .broadcast_as(src.shape())? + .contiguous()?; + freqs_t = freqs_t.scatter(&idx, &src, D::Minus1)?; + } + Ok(freqs_t) + } + pub fn forward( + &self, + position_ids: &Tensor, + dtype: DType, + mrope_section: Vec, + ) -> Result<(Tensor, Tensor)> { + // position_ids shape: (3, bs, position) -> (3, bs, 1, position) + let position_ids = if position_ids.rank() == 2 { + let (bs, len) = position_ids.dims2()?; + position_ids.unsqueeze(0)?.expand((3, bs, len))? + } else { + position_ids.clone() + }; + let position_ids_expanded = position_ids + .unsqueeze(D::Minus2)? + .to_dtype(DType::F32)? + .contiguous()?; + // inv_freq Vec -> Tensor(1, 1, head_dim / 2, 1) -> (3, bs, head_dim / 2, 1) + let inv_freq_expanded = Tensor::from_vec( + self.inv_freq.clone(), + (1, 1, self.inv_freq.len(), 1), + position_ids.device(), + )? + .broadcast_as((3, position_ids.dim(1)?, self.inv_freq.len(), 1))? + .to_dtype(DType::F32)? + .contiguous()?; + + // (3, bs, head_dim / 2, 1) matmul (3, bs, 1, position) + // -> (3, bs, head_dim / 2, seq_len) -> (3, bs, seq_len, head_dim / 2) + let freqs = inv_freq_expanded + .matmul(&position_ids_expanded)? + .transpose(2, 3)?; + let freqs = self.apply_interleaved_mrope(&freqs, mrope_section)?; + let emb = Tensor::cat(&[&freqs, &freqs], D::Minus1)?.contiguous()?; + let cos = emb.cos()?; + let sin = emb.sin()?; + Ok((cos.to_dtype(dtype)?, sin.to_dtype(dtype)?)) + } +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index ef2b9ee..a274ce0 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -5,7 +5,7 @@ pub mod video_utils; use anyhow::Result; use candle_core::{DType, Device}; -use candle_transformers::generation::LogitsProcessor; +use candle_transformers::generation::{LogitsProcessor, Sampling}; use openai_dive::v1::resources::{ chat::{ ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse, @@ -251,10 +251,33 @@ pub fn build_completion_chunk_response( response } -pub fn get_logit_processor(temperature: Option, top_p: Option) -> LogitsProcessor { - LogitsProcessor::new( - 34562, - temperature.map(|temp| temp as f64), - top_p.map(|tp| tp as f64), - ) +pub fn get_logit_processor( + temperature: Option, + top_p: Option, + top_k: Option, +) -> LogitsProcessor { + match top_k { + None => LogitsProcessor::new( + 34562, + temperature.map(|temp| temp as f64), + top_p.map(|tp| tp as f64), + ), + Some(k) => { + let sampling = match temperature { + None => Sampling::ArgMax, + Some(temperature) => match top_p { + None => Sampling::TopK { + k, + temperature: temperature as f64, + }, + Some(p) => Sampling::TopKThenTopP { + k, + p: p as f64, + temperature: temperature as f64, + }, + }, + }; + LogitsProcessor::from_sampling(34562, sampling) + } + } } diff --git a/src/utils/tensor_utils.rs b/src/utils/tensor_utils.rs index e616be0..8d9f9ae 100644 --- a/src/utils/tensor_utils.rs +++ b/src/utils/tensor_utils.rs @@ -42,7 +42,7 @@ pub fn repeat_kv(xs: Tensor, n_rep: usize) -> Result { } } -pub fn split(t: &Tensor, splits: &[usize], dim: D) -> Result> { +pub fn split_tensor(t: &Tensor, splits: &[usize], dim: D) -> Result> { let dim = dim.to_index(t.shape(), "split")?; let mut split_res = Vec::new(); let mut index = 0; @@ -241,3 +241,114 @@ pub fn linspace(start: f32, end: f32, steps: usize, device: &Device) -> Result Result { + assert!( + mask1.shape() == mask2.shape(), + " bitor_tensor two tensor shape mask be equal" + ); + let bitor = mask1.add(&mask2)?.ne(&Tensor::zeros_like(&mask1)?)?; + Ok(bitor) +} + +pub fn prod_tensor_last_dim(t: &Tensor) -> Result { + let prod = match t.rank() { + 0 => t.clone(), + 1 => { + let data_type = t.dtype(); + let t_prod = match data_type { + DType::U8 => { + let t_vec = t.to_vec1::()?; + let prod = t_vec.iter().product::(); + Tensor::from_slice(&vec![prod], 1, t.device())? + } + DType::U32 => { + let t_vec = t.to_vec1::()?; + let prod = t_vec.iter().product::(); + Tensor::from_slice(&vec![prod], 1, t.device())? + } + DType::I64 => { + let t_vec = t.to_vec1::()?; + let prod = t_vec.iter().product::(); + Tensor::from_slice(&vec![prod], 1, t.device())? + } + DType::F64 => { + let t_vec = t.to_vec1::()?; + let prod = t_vec.iter().product::(); + Tensor::from_slice(&vec![prod], 1, t.device())? + } + _ => { + let t_vec = t.to_vec1::()?; + let prod = t_vec.iter().product::(); + Tensor::from_slice(&vec![prod], 1, t.device())? + } + }; + t_prod + } + 2 => { + let data_type = t.dtype(); + let t_prod = match data_type { + DType::U8 => { + let t_vec = t.to_vec2::()?; + let mut prod_vec = vec![]; + for v in t_vec.iter() { + let prod = v.iter().product::(); + prod_vec.push(prod); + } + Tensor::new(prod_vec, t.device())? + } + DType::U32 => { + let t_vec = t.to_vec2::()?; + let mut prod_vec = vec![]; + for v in t_vec.iter() { + let prod = v.iter().product::(); + prod_vec.push(prod); + } + Tensor::new(prod_vec, t.device())? + } + DType::I64 => { + let t_vec = t.to_vec2::()?; + let mut prod_vec = vec![]; + for v in t_vec.iter() { + let prod = v.iter().product::(); + prod_vec.push(prod); + } + Tensor::new(prod_vec, t.device())? + } + DType::F64 => { + let t_vec = t.to_vec2::()?; + let mut prod_vec = vec![]; + for v in t_vec.iter() { + let prod = v.iter().product::(); + prod_vec.push(prod); + } + Tensor::new(prod_vec, t.device())? + } + _ => { + let t_vec = t.to_vec2::()?; + let mut prod_vec = vec![]; + for v in t_vec.iter() { + let prod = v.iter().product::(); + prod_vec.push(prod); + } + Tensor::new(prod_vec, t.device())? + } + }; + t_prod + } + _ => { + return Err(anyhow!(format!("can not action this dim"))); + } + }; + Ok(prod) +} + +pub fn mask_index_add( + original: &Tensor, + mask: &Tensor, + add: &Tensor, +) -> Result { + let visual_nonzero_index = nonzero_index(&mask)?; + let xs = original.index_add(&visual_nonzero_index, add, 0)?; + Ok(xs) +} diff --git a/tests/config_tests.rs b/tests/config_tests.rs index fce726e..7893df9 100644 --- a/tests/config_tests.rs +++ b/tests/config_tests.rs @@ -1,6 +1,5 @@ use aha::models::{ - minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig, - voxcpm::config::VoxCPMConfig, + minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig, qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig }; use anyhow::Result; @@ -34,3 +33,13 @@ fn voxcpm_config() -> Result<()> { println!("{:?}", config); Ok(()) } + +#[test] +fn qwen3vl_config() -> Result<()> { + // cargo test -F cuda qwen3vl_config -- --nocapture + let model_path = "/home/jhq/huggingface_model/Qwen/Qwen3-VL-4B-Instruct/"; + let config_path = model_path.to_string() + "/config.json"; + let config: Qwen3VLConfig = serde_json::from_slice(&std::fs::read(config_path)?)?; + println!("{:?}", config); + Ok(()) +} diff --git a/tests/messy_test.rs b/tests/messy_test.rs index 7a0d990..96940c8 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,13 +1,25 @@ -use aha::utils::audio_utils::load_audio_with_resample; +use aha::utils::{tensor_utils::bitor_tensor}; use anyhow::Result; +use candle_core::Tensor; #[test] fn messy_test() -> Result<()> { - let device = candle_core::Device::Cpu; - let wav_path = "./assets/audio/voice_01.wav"; - let audio_tensor = load_audio_with_resample(wav_path, device, Some(16000))?; - - println!("audio_tensor: {}", audio_tensor); + let device = &candle_core::Device::Cpu; + let image_mask = Tensor::new(vec![0u32, 0, 0, 1, 0, 1], device)?; + let video_mask = Tensor::new(vec![0u32, 1, 0, 1, 0, 1], device)?; + let visual_mask = bitor_tensor(&image_mask, &video_mask)?; + println!("visual_mask: {}", visual_mask); + // let x = Tensor::arange_step(0.0_f32, 5., 0.5, &device)?; + // let x_int = x.to_dtype(candle_core::DType::U32)?; + // println!("x: {}", x); + // println!("x_int: {}", x_int); + // let x_affine = x_int.affine(1.0, 1.0)?; + // println!("x_affine: {}", x_affine); + // let x_clamp = x_affine.clamp(0u32, 3u32)?; + // println!("x_clamp: {}", x_clamp); + // let wav_path = "./assets/audio/voice_01.wav"; + // let audio_tensor = load_audio_with_resample(wav_path, device, Some(16000))?; + // println!("audio_tensor: {}", audio_tensor); // let string = "你好啊".to_string(); // let vec_str: Vec= string.chars().map(|c| c.to_string()).collect(); // println!("vec_str: {:?}", vec_str); diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs new file mode 100644 index 0000000..db93581 --- /dev/null +++ b/tests/test_qwen3vl.rs @@ -0,0 +1,94 @@ +use std::{pin::pin, time::Instant}; + +use aha::models::{GenerateModel, qwen3vl::generate::Qwen3VLGenerateModel}; +use anyhow::Result; +use openai_dive::v1::resources::chat::ChatCompletionParameters; +use rocket::futures::StreamExt; + +#[test] +fn qwen3vl_generate() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3vl_generate -- --nocapture + + let model_path = "/home/jhq/huggingface_model/Qwen/Qwen3-VL-2B-Instruct/"; + + let message = r#" + { + "model": "qwen3vl", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "video", + "video_url": + { + "url": "./assets/video/video_test.mp4" + } + }, + { + "type": "text", + "text": "视频中发生了什么." + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut qwen3vl = Qwen3VLGenerateModel::init(model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let res = qwen3vl.generate(mes)?; + let i_duration = i_start.elapsed(); + println!("generate: \n {:?}", res); + println!("Time elapsed in generate is: {:?}", i_duration); + Ok(()) +} + +#[tokio::test] +async fn qwen3vl_stream() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda qwen3vl_stream -- --nocapture + + let model_path = "/home/jhq/huggingface_model/Qwen/Qwen3-VL-2B-Instruct/"; + + let message = r#" + { + "model": "qwen3vl", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image", + "image_url": + { + "url": "file://./assets/img/voxcpm.png" + } + }, + { + "type": "text", + "text": "描述这张图片" + } + ] + } + ] + } + "#; + let mes: ChatCompletionParameters = serde_json::from_str(message)?; + let i_start = Instant::now(); + let mut qwen3vl = Qwen3VLGenerateModel::init(model_path, None, None)?; + let i_duration = i_start.elapsed(); + println!("Time elapsed in load model is: {:?}", i_duration); + + let i_start = Instant::now(); + let mut stream = pin!(qwen3vl.generate_stream(mes)?); + while let Some(item) = stream.next().await { + println!("generate: \n {:?}", item); + } + let i_duration = i_start.elapsed(); + println!("Time elapsed in generate is: {:?}", i_duration); + Ok(()) +} diff --git a/tests/weight_test.rs b/tests/weight_test.rs index 14c04fa..d16e461 100644 --- a/tests/weight_test.rs +++ b/tests/weight_test.rs @@ -45,3 +45,19 @@ fn voxcpm_weight() -> Result<()> { ); Ok(()) } + +#[test] +fn qwen3vl_weight() -> Result<()> { + let model_path = "/home/jhq/huggingface_model/Qwen/Qwen3-VL-4B-Instruct/"; + let model_list = find_type_files(model_path, "safetensors")?; + + let device = Device::Cpu; + for m in &model_list { + let weights = safetensors::load(m, &device)?; + for (key, tensor) in weights.iter() { + println!("=== {} === {:?}", key, tensor.shape()); + } + } + println!("model_list: {:?}", model_list); + Ok(()) +} \ No newline at end of file