add Qwen3VL model
This commit is contained in:
Generated
+2
-3
@@ -2445,9 +2445,8 @@ dependencies = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "openai_dive"
|
name = "openai_dive"
|
||||||
version = "1.3.0"
|
version = "1.3.2"
|
||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "git+https://github.com/jhqxxx/openai-client.git#83363a83fbd73273e2a1474ca7fdfb1a899e82d4"
|
||||||
checksum = "bac82ff995cdbe4e120fa87fe4d76bd452866dd7b0c9bcbdb562ba4be0f1786d"
|
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"bytes",
|
"bytes",
|
||||||
"derive_builder",
|
"derive_builder",
|
||||||
|
|||||||
+1
-1
@@ -21,7 +21,7 @@ base64 = "0.22.1"
|
|||||||
num = "0.4.3"
|
num = "0.4.3"
|
||||||
minijinja = "2.12.0"
|
minijinja = "2.12.0"
|
||||||
tokenizers = "0.22.1"
|
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"]}
|
uuid = { version = "1.18.1", features = ["v4"]}
|
||||||
chrono = "0.4.42"
|
chrono = "0.4.42"
|
||||||
rocket = "0.5.1"
|
rocket = "0.5.1"
|
||||||
|
|||||||
@@ -15,6 +15,7 @@
|
|||||||
* Qwen2.5VL - 阿里通义千问 2.5 多模态大语言模型
|
* Qwen2.5VL - 阿里通义千问 2.5 多模态大语言模型
|
||||||
* MiniCPM4 - 面壁智能 MiniCPM 系列语言模型
|
* MiniCPM4 - 面壁智能 MiniCPM 系列语言模型
|
||||||
* VoxCPM - 面壁智能语音生成模型
|
* 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
|
git clone https://github.com/jhqxxx/aha.git
|
||||||
cd aha
|
cd aha
|
||||||
# 修改测试用例中模型路径
|
# 修改测试用例中模型路径
|
||||||
# 运行 Qwen2.5VL 示例
|
# 运行 Qwen3VL 示例
|
||||||
cargo test -F cuda qwen2_5vl_generate -- --nocapture
|
cargo test -F cuda qwen3vl_generate -- --nocapture
|
||||||
|
|
||||||
# 运行 MiniCPM4 示例
|
# 运行 MiniCPM4 示例
|
||||||
cargo test -F cuda minicpm_generate -- --nocapture
|
cargo test -F cuda minicpm_generate -- --nocapture
|
||||||
@@ -91,6 +92,7 @@ fn main() -> Result<()> {
|
|||||||
│ │ ├── common
|
│ │ ├── common
|
||||||
│ │ ├── minicpm4
|
│ │ ├── minicpm4
|
||||||
│ │ ├── qwen2_5vl
|
│ │ ├── qwen2_5vl
|
||||||
|
│ │ ├── qwen3vl
|
||||||
│ │ ├── voxcpm
|
│ │ ├── voxcpm
|
||||||
│ │ └── mod.rs
|
│ │ └── mod.rs
|
||||||
│ ├── position_embed
|
│ ├── position_embed
|
||||||
@@ -121,6 +123,9 @@ fn main() -> Result<()> {
|
|||||||
2. 提交新的 Issue,包含详细描述和复现步骤
|
2. 提交新的 Issue,包含详细描述和复现步骤
|
||||||
|
|
||||||
## 更新日志
|
## 更新日志
|
||||||
|
### v0.1.1
|
||||||
|
* 添加 Qwen3VL 模型
|
||||||
|
|
||||||
### v0.1.0
|
### v0.1.0
|
||||||
* 初始版本发布
|
* 初始版本发布
|
||||||
* 支持 Qwen2.5VL, MiniCPM4, VoxCPM 模型
|
* 支持 Qwen2.5VL, MiniCPM4, VoxCPM 模型
|
||||||
|
|||||||
Binary file not shown.
-25
@@ -3,28 +3,3 @@ pub mod models;
|
|||||||
pub mod position_embed;
|
pub mod position_embed;
|
||||||
pub mod tokenizer;
|
pub mod tokenizer;
|
||||||
pub mod utils;
|
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<DType>,
|
|
||||||
// ) -> Result<Box<dyn GenerateModel>> {
|
|
||||||
// match model_type {
|
|
||||||
// ModelType::Qwen2_5VL => {
|
|
||||||
// let model = Qwen2_5VLGenerateModel::init(model_path, device, dtype)?;
|
|
||||||
// Ok(Box::new(model) as Box<dyn GenerateModel>)
|
|
||||||
// },
|
|
||||||
// ModelType::MiniCPM4 => {
|
|
||||||
// let model = MiniCPMGenerateModel::init(model_path, device, dtype)?;
|
|
||||||
// Ok(Box::new(model)as Box<dyn GenerateModel>)
|
|
||||||
// }
|
|
||||||
// }
|
|
||||||
// }
|
|
||||||
// }
|
|
||||||
|
|||||||
@@ -269,3 +269,54 @@ impl AttentionNobias {
|
|||||||
self.kv_cache = None
|
self.kv_cache = None
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn eager_attention_forward(
|
||||||
|
query_states: &Tensor,
|
||||||
|
key_states: &Tensor,
|
||||||
|
value_states: &Tensor,
|
||||||
|
num_key_value_groups: Option<usize>,
|
||||||
|
attention_mask: Option<&Tensor>,
|
||||||
|
scaling: f64,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -53,7 +53,7 @@ impl<'a> MiniCPMGenerateModel<'a> {
|
|||||||
|
|
||||||
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
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 mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let mut seq_len = input_ids.dim(1)?;
|
||||||
@@ -81,7 +81,7 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
|||||||
&mut self,
|
&mut self,
|
||||||
mes: ChatCompletionParameters,
|
mes: ChatCompletionParameters,
|
||||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||||
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 mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||||
let mut seq_len = input_ids.dim(1)?;
|
let mut seq_len = input_ids.dim(1)?;
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ pub mod common;
|
|||||||
pub mod minicpm4;
|
pub mod minicpm4;
|
||||||
pub mod qwen2_5vl;
|
pub mod qwen2_5vl;
|
||||||
pub mod voxcpm;
|
pub mod voxcpm;
|
||||||
|
pub mod qwen3vl;
|
||||||
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use openai_dive::v1::resources::chat::{
|
use openai_dive::v1::resources::chat::{
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ impl<'a> Qwen2_5VLGenerateModel<'a> {
|
|||||||
|
|
||||||
impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
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 mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||||
let mut input_ids = self
|
let mut input_ids = self
|
||||||
@@ -118,11 +118,12 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
|||||||
let response = build_completion_response(res, "qwen2.5vl");
|
let response = build_completion_response(res, "qwen2.5vl");
|
||||||
Ok(response)
|
Ok(response)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn generate_stream(
|
fn generate_stream(
|
||||||
&mut self,
|
&mut self,
|
||||||
mes: ChatCompletionParameters,
|
mes: ChatCompletionParameters,
|
||||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||||
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 mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||||
let mut input_ids = self
|
let mut input_ids = self
|
||||||
|
|||||||
@@ -841,7 +841,7 @@ impl Qwen2_5VLTextModel {
|
|||||||
if seq_len <= 1 {
|
if seq_len <= 1 {
|
||||||
None
|
None
|
||||||
} else {
|
} 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() {
|
for layer in self.layers.iter_mut() {
|
||||||
|
|||||||
@@ -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<f32>,
|
||||||
|
pub image_std: Vec<f32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||||
|
pub struct RopeScaling {
|
||||||
|
pub rope_type: String,
|
||||||
|
pub mrope_section: Vec<usize>,
|
||||||
|
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<usize>,
|
||||||
|
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<usize>,
|
||||||
|
pub top_p: f32,
|
||||||
|
pub top_k: usize,
|
||||||
|
pub temperature: f32,
|
||||||
|
pub repetition_penalty: f32,
|
||||||
|
}
|
||||||
@@ -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<DType>) -> Result<Self> {
|
||||||
|
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<ChatCompletionResponse> {
|
||||||
|
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<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,4 @@
|
|||||||
|
pub mod processor;
|
||||||
|
pub mod config;
|
||||||
|
pub mod model;
|
||||||
|
pub mod generate;
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -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<Tensor>,
|
||||||
|
pub image_grid_thw: Option<Tensor>,
|
||||||
|
pub pixel_values_video: Option<Tensor>,
|
||||||
|
pub video_grid_thw: Option<Tensor>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[allow(unused)]
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct VideoMetadata {
|
||||||
|
total_num_frames: u32,
|
||||||
|
fps: f32,
|
||||||
|
width: u32,
|
||||||
|
height: u32,
|
||||||
|
duration: f32,
|
||||||
|
frame_indices: Vec<u32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
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<Self> {
|
||||||
|
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<HashMap<String, Vec<String>>> {
|
||||||
|
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<Tensor> {
|
||||||
|
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<DynamicImage>,
|
||||||
|
img_mean: &Tensor,
|
||||||
|
img_std: &Tensor,
|
||||||
|
) -> Result<VisionInput> {
|
||||||
|
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<Tensor>,
|
||||||
|
img_mean: &Tensor,
|
||||||
|
img_std: &Tensor,
|
||||||
|
) -> Result<VisionInput> {
|
||||||
|
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<u32>,
|
||||||
|
fps: f32,
|
||||||
|
t_merge_size: usize,
|
||||||
|
) -> Result<Vec<f32>> {
|
||||||
|
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<f32> = 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<GeneralInput> {
|
||||||
|
let mut pixel_values = None;
|
||||||
|
let mut image_grid_thw = None;
|
||||||
|
let mut pixel_values_video = None;
|
||||||
|
let mut video_grid_thw: Option<Tensor> = 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::<u32>()?.iter().product::<u32>() 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::<u32>()?[..] 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<u32>,
|
||||||
|
) -> 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<u32>,
|
||||||
|
) -> 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))
|
||||||
|
}
|
||||||
@@ -64,7 +64,8 @@ pub fn apply_rotary_pos_emb_vision(
|
|||||||
// cos, sin -> (seq_len, head_dim) -> (seq_len, 1, head_dim)
|
// cos, sin -> (seq_len, head_dim) -> (seq_len, 1, head_dim)
|
||||||
let cos = cos.unsqueeze(D::Minus2)?;
|
let cos = cos.unsqueeze(D::Minus2)?;
|
||||||
let sin = sin.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
|
let q_embed = q
|
||||||
.broadcast_mul(&cos)?
|
.broadcast_mul(&cos)?
|
||||||
.add(&rotate_half(q)?.broadcast_mul(&sin)?)?;
|
.add(&rotate_half(q)?.broadcast_mul(&sin)?)?;
|
||||||
@@ -197,3 +198,75 @@ impl Qwen2_5VisionRotaryEmbedding {
|
|||||||
Ok(freqs)
|
Ok(freqs)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub struct Qwen3VLTextRotaryEmbedding {
|
||||||
|
inv_freq: Vec<f32>,
|
||||||
|
}
|
||||||
|
|
||||||
|
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<usize>,
|
||||||
|
) -> Result<Tensor> {
|
||||||
|
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<usize>,
|
||||||
|
) -> 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<f32> -> 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)?))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+30
-7
@@ -5,7 +5,7 @@ pub mod video_utils;
|
|||||||
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::{DType, Device};
|
use candle_core::{DType, Device};
|
||||||
use candle_transformers::generation::LogitsProcessor;
|
use candle_transformers::generation::{LogitsProcessor, Sampling};
|
||||||
use openai_dive::v1::resources::{
|
use openai_dive::v1::resources::{
|
||||||
chat::{
|
chat::{
|
||||||
ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
|
ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
|
||||||
@@ -251,10 +251,33 @@ pub fn build_completion_chunk_response(
|
|||||||
response
|
response
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn get_logit_processor(temperature: Option<f32>, top_p: Option<f32>) -> LogitsProcessor {
|
pub fn get_logit_processor(
|
||||||
LogitsProcessor::new(
|
temperature: Option<f32>,
|
||||||
34562,
|
top_p: Option<f32>,
|
||||||
temperature.map(|temp| temp as f64),
|
top_k: Option<usize>,
|
||||||
top_p.map(|tp| tp as f64),
|
) -> 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+112
-1
@@ -42,7 +42,7 @@ pub fn repeat_kv(xs: Tensor, n_rep: usize) -> Result<Tensor> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn split(t: &Tensor, splits: &[usize], dim: D) -> Result<Vec<Tensor>> {
|
pub fn split_tensor<D: Dim>(t: &Tensor, splits: &[usize], dim: D) -> Result<Vec<Tensor>> {
|
||||||
let dim = dim.to_index(t.shape(), "split")?;
|
let dim = dim.to_index(t.shape(), "split")?;
|
||||||
let mut split_res = Vec::new();
|
let mut split_res = Vec::new();
|
||||||
let mut index = 0;
|
let mut index = 0;
|
||||||
@@ -241,3 +241,114 @@ pub fn linspace(start: f32, end: f32, steps: usize, device: &Device) -> Result<T
|
|||||||
let t = Tensor::from_slice(&data, steps, device)?;
|
let t = Tensor::from_slice(&data, steps, device)?;
|
||||||
Ok(t)
|
Ok(t)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn bitor_tensor(mask1: &Tensor, mask2: &Tensor) -> Result<Tensor> {
|
||||||
|
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<Tensor> {
|
||||||
|
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::<u8>()?;
|
||||||
|
let prod = t_vec.iter().product::<u8>();
|
||||||
|
Tensor::from_slice(&vec![prod], 1, t.device())?
|
||||||
|
}
|
||||||
|
DType::U32 => {
|
||||||
|
let t_vec = t.to_vec1::<u32>()?;
|
||||||
|
let prod = t_vec.iter().product::<u32>();
|
||||||
|
Tensor::from_slice(&vec![prod], 1, t.device())?
|
||||||
|
}
|
||||||
|
DType::I64 => {
|
||||||
|
let t_vec = t.to_vec1::<i64>()?;
|
||||||
|
let prod = t_vec.iter().product::<i64>();
|
||||||
|
Tensor::from_slice(&vec![prod], 1, t.device())?
|
||||||
|
}
|
||||||
|
DType::F64 => {
|
||||||
|
let t_vec = t.to_vec1::<f64>()?;
|
||||||
|
let prod = t_vec.iter().product::<f64>();
|
||||||
|
Tensor::from_slice(&vec![prod], 1, t.device())?
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
let t_vec = t.to_vec1::<f32>()?;
|
||||||
|
let prod = t_vec.iter().product::<f32>();
|
||||||
|
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::<u8>()?;
|
||||||
|
let mut prod_vec = vec![];
|
||||||
|
for v in t_vec.iter() {
|
||||||
|
let prod = v.iter().product::<u8>();
|
||||||
|
prod_vec.push(prod);
|
||||||
|
}
|
||||||
|
Tensor::new(prod_vec, t.device())?
|
||||||
|
}
|
||||||
|
DType::U32 => {
|
||||||
|
let t_vec = t.to_vec2::<u32>()?;
|
||||||
|
let mut prod_vec = vec![];
|
||||||
|
for v in t_vec.iter() {
|
||||||
|
let prod = v.iter().product::<u32>();
|
||||||
|
prod_vec.push(prod);
|
||||||
|
}
|
||||||
|
Tensor::new(prod_vec, t.device())?
|
||||||
|
}
|
||||||
|
DType::I64 => {
|
||||||
|
let t_vec = t.to_vec2::<i64>()?;
|
||||||
|
let mut prod_vec = vec![];
|
||||||
|
for v in t_vec.iter() {
|
||||||
|
let prod = v.iter().product::<i64>();
|
||||||
|
prod_vec.push(prod);
|
||||||
|
}
|
||||||
|
Tensor::new(prod_vec, t.device())?
|
||||||
|
}
|
||||||
|
DType::F64 => {
|
||||||
|
let t_vec = t.to_vec2::<f64>()?;
|
||||||
|
let mut prod_vec = vec![];
|
||||||
|
for v in t_vec.iter() {
|
||||||
|
let prod = v.iter().product::<f64>();
|
||||||
|
prod_vec.push(prod);
|
||||||
|
}
|
||||||
|
Tensor::new(prod_vec, t.device())?
|
||||||
|
}
|
||||||
|
_ => {
|
||||||
|
let t_vec = t.to_vec2::<f32>()?;
|
||||||
|
let mut prod_vec = vec![];
|
||||||
|
for v in t_vec.iter() {
|
||||||
|
let prod = v.iter().product::<f32>();
|
||||||
|
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<Tensor> {
|
||||||
|
let visual_nonzero_index = nonzero_index(&mask)?;
|
||||||
|
let xs = original.index_add(&visual_nonzero_index, add, 0)?;
|
||||||
|
Ok(xs)
|
||||||
|
}
|
||||||
|
|||||||
+11
-2
@@ -1,6 +1,5 @@
|
|||||||
use aha::models::{
|
use aha::models::{
|
||||||
minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig,
|
minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig, qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig
|
||||||
voxcpm::config::VoxCPMConfig,
|
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
|
|
||||||
@@ -34,3 +33,13 @@ fn voxcpm_config() -> Result<()> {
|
|||||||
println!("{:?}", config);
|
println!("{:?}", config);
|
||||||
Ok(())
|
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(())
|
||||||
|
}
|
||||||
|
|||||||
+18
-6
@@ -1,13 +1,25 @@
|
|||||||
use aha::utils::audio_utils::load_audio_with_resample;
|
use aha::utils::{tensor_utils::bitor_tensor};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
|
use candle_core::Tensor;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn messy_test() -> Result<()> {
|
fn messy_test() -> Result<()> {
|
||||||
let device = candle_core::Device::Cpu;
|
let device = &candle_core::Device::Cpu;
|
||||||
let wav_path = "./assets/audio/voice_01.wav";
|
let image_mask = Tensor::new(vec![0u32, 0, 0, 1, 0, 1], device)?;
|
||||||
let audio_tensor = load_audio_with_resample(wav_path, device, Some(16000))?;
|
let video_mask = Tensor::new(vec![0u32, 1, 0, 1, 0, 1], device)?;
|
||||||
|
let visual_mask = bitor_tensor(&image_mask, &video_mask)?;
|
||||||
println!("audio_tensor: {}", audio_tensor);
|
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 string = "你好啊".to_string();
|
||||||
// let vec_str: Vec<String>= string.chars().map(|c| c.to_string()).collect();
|
// let vec_str: Vec<String>= string.chars().map(|c| c.to_string()).collect();
|
||||||
// println!("vec_str: {:?}", vec_str);
|
// println!("vec_str: {:?}", vec_str);
|
||||||
|
|||||||
@@ -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(())
|
||||||
|
}
|
||||||
@@ -45,3 +45,19 @@ fn voxcpm_weight() -> Result<()> {
|
|||||||
);
|
);
|
||||||
Ok(())
|
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(())
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user