update code format
This commit is contained in:
@@ -280,11 +280,11 @@ pub fn eager_attention_forward(
|
|||||||
) -> Result<Tensor> {
|
) -> Result<Tensor> {
|
||||||
let key_states = match num_key_value_groups {
|
let key_states = match num_key_value_groups {
|
||||||
Some(g) => repeat_kv(key_states.clone(), g)?.contiguous()?,
|
Some(g) => repeat_kv(key_states.clone(), g)?.contiguous()?,
|
||||||
None => key_states.clone()
|
None => key_states.clone(),
|
||||||
};
|
};
|
||||||
let value_states = match num_key_value_groups {
|
let value_states = match num_key_value_groups {
|
||||||
Some(g) => repeat_kv(value_states.clone(), g)?.contiguous()?,
|
Some(g) => repeat_kv(value_states.clone(), g)?.contiguous()?,
|
||||||
None => value_states.clone()
|
None => value_states.clone(),
|
||||||
};
|
};
|
||||||
let attn_output = {
|
let attn_output = {
|
||||||
#[cfg(not(feature = "flash-attn"))]
|
#[cfg(not(feature = "flash-attn"))]
|
||||||
|
|||||||
+1
-1
@@ -1,8 +1,8 @@
|
|||||||
pub mod common;
|
pub mod common;
|
||||||
pub mod minicpm4;
|
pub mod minicpm4;
|
||||||
pub mod qwen2_5vl;
|
pub mod qwen2_5vl;
|
||||||
pub mod voxcpm;
|
|
||||||
pub mod qwen3vl;
|
pub mod qwen3vl;
|
||||||
|
pub mod voxcpm;
|
||||||
|
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use openai_dive::v1::resources::chat::{
|
use openai_dive::v1::resources::chat::{
|
||||||
|
|||||||
@@ -1,6 +1,5 @@
|
|||||||
use candle_nn::Activation;
|
use candle_nn::Activation;
|
||||||
|
|
||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||||
pub struct Size {
|
pub struct Size {
|
||||||
pub longest_edge: usize,
|
pub longest_edge: usize,
|
||||||
|
|||||||
@@ -1,20 +1,26 @@
|
|||||||
use anyhow::{Result, anyhow};
|
use anyhow::{Result, anyhow};
|
||||||
use candle_core::{DType, Device, Tensor};
|
use candle_core::{DType, Device, Tensor};
|
||||||
use candle_nn::VarBuilder;
|
use candle_nn::VarBuilder;
|
||||||
use openai_dive::v1::resources::chat::{ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse};
|
use openai_dive::v1::resources::chat::{
|
||||||
use rocket::futures::Stream;
|
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||||
|
};
|
||||||
use rocket::async_stream::stream;
|
use rocket::async_stream::stream;
|
||||||
|
use rocket::futures::Stream;
|
||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
chat_template::ChatTemplate,
|
chat_template::ChatTemplate,
|
||||||
models::{GenerateModel, qwen3vl::{
|
models::{
|
||||||
config::{Qwen3VLConfig, Qwen3VLGenerationConfig},
|
GenerateModel,
|
||||||
model::Qwen3VLModel,
|
qwen3vl::{
|
||||||
processor::Qwen3VLProcessor,
|
config::{Qwen3VLConfig, Qwen3VLGenerationConfig},
|
||||||
}},
|
model::Qwen3VLModel,
|
||||||
|
processor::Qwen3VLProcessor,
|
||||||
|
},
|
||||||
|
},
|
||||||
tokenizer::TokenizerModel,
|
tokenizer::TokenizerModel,
|
||||||
utils::{
|
utils::{
|
||||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device, get_dtype, get_logit_processor
|
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||||
|
get_dtype, get_logit_processor,
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -57,7 +63,6 @@ impl<'a> Qwen3VLGenerateModel<'a> {
|
|||||||
generation_config,
|
generation_config,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
pub mod processor;
|
|
||||||
pub mod config;
|
pub mod config;
|
||||||
pub mod model;
|
|
||||||
pub mod generate;
|
pub mod generate;
|
||||||
|
pub mod model;
|
||||||
|
pub mod processor;
|
||||||
|
|||||||
@@ -1135,7 +1135,7 @@ impl Qwen3VLModel {
|
|||||||
let (image_embeds, deepstack_img_embed) =
|
let (image_embeds, deepstack_img_embed) =
|
||||||
self.get_vision_features(pixel_values, image_grid_thw)?;
|
self.get_vision_features(pixel_values, image_grid_thw)?;
|
||||||
let image_embeds = Tensor::cat(&image_embeds, 0)?;
|
let image_embeds = Tensor::cat(&image_embeds, 0)?;
|
||||||
let vision_mask = self.get_placeholder_mask(&input_ids, true)?;
|
let vision_mask = self.get_placeholder_mask(input_ids, true)?;
|
||||||
let n_image_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
|
let n_image_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
|
||||||
if n_image_tokens as usize != image_embeds.dim(0)? {
|
if n_image_tokens as usize != image_embeds.dim(0)? {
|
||||||
return Err(anyhow!(format!(
|
return Err(anyhow!(format!(
|
||||||
@@ -1154,7 +1154,7 @@ impl Qwen3VLModel {
|
|||||||
let (video_embeds, deepstack_video_embed) =
|
let (video_embeds, deepstack_video_embed) =
|
||||||
self.get_vision_features(pixel_values_video, video_grid_thw)?;
|
self.get_vision_features(pixel_values_video, video_grid_thw)?;
|
||||||
let video_embeds = Tensor::cat(&video_embeds, 0)?;
|
let video_embeds = Tensor::cat(&video_embeds, 0)?;
|
||||||
let vision_mask = self.get_placeholder_mask(&input_ids, false)?;
|
let vision_mask = self.get_placeholder_mask(input_ids, false)?;
|
||||||
let n_video_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
|
let n_video_tokens = vision_mask.sum_all()?.to_scalar::<u32>()?;
|
||||||
if n_video_tokens as usize != video_embeds.dim(0)? {
|
if n_video_tokens as usize != video_embeds.dim(0)? {
|
||||||
return Err(anyhow!(format!(
|
return Err(anyhow!(format!(
|
||||||
@@ -1190,8 +1190,8 @@ impl Qwen3VLModel {
|
|||||||
img_embed.dtype(),
|
img_embed.dtype(),
|
||||||
img_embed.device(),
|
img_embed.device(),
|
||||||
)?;
|
)?;
|
||||||
let embed_joint = embed_joint.index_add(&image_nonzero_joint, &img_embed, 0)?;
|
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)?;
|
let embed_joint = embed_joint.index_add(&video_nonzero_joint, vid_embed, 0)?;
|
||||||
deepstack_embeds.push(embed_joint);
|
deepstack_embeds.push(embed_joint);
|
||||||
}
|
}
|
||||||
visual_pos_mask = Some(visual_mask);
|
visual_pos_mask = Some(visual_mask);
|
||||||
|
|||||||
@@ -259,7 +259,7 @@ impl Qwen3VLProcessor {
|
|||||||
fps: f32,
|
fps: f32,
|
||||||
t_merge_size: usize,
|
t_merge_size: usize,
|
||||||
) -> Result<Vec<f32>> {
|
) -> Result<Vec<f32>> {
|
||||||
let indices = if frames_indices.len() % t_merge_size != 0 {
|
let indices = if !frames_indices.len().is_multiple_of(t_merge_size) {
|
||||||
let mut frames_indices = frames_indices.clone();
|
let mut frames_indices = frames_indices.clone();
|
||||||
let last = frames_indices[frames_indices.len() - 1];
|
let last = frames_indices[frames_indices.len() - 1];
|
||||||
let pad_len = t_merge_size - frames_indices.len() % t_merge_size;
|
let pad_len = t_merge_size - frames_indices.len() % t_merge_size;
|
||||||
@@ -385,12 +385,15 @@ impl Qwen3VLProcessor {
|
|||||||
let frame_seqlen = h * w / merge_length as u32;
|
let frame_seqlen = h * w / merge_length as u32;
|
||||||
for frame_idx in 0..t {
|
for frame_idx in 0..t {
|
||||||
let curr_time = curr_timestamp[frame_idx as usize];
|
let curr_time = curr_timestamp[frame_idx as usize];
|
||||||
video_placeholder = video_placeholder + format!("<{:.1} seconds>", curr_time).as_str();
|
video_placeholder += format!("<{:.1} seconds>", curr_time).as_str();
|
||||||
video_placeholder = video_placeholder + self.vision_start_token.as_str();
|
video_placeholder += self.vision_start_token.as_str();
|
||||||
video_placeholder = video_placeholder + "<|placeholder|>".repeat(frame_seqlen as usize).as_str();
|
video_placeholder += "<|placeholder|>".repeat(frame_seqlen as usize).as_str();
|
||||||
video_placeholder = video_placeholder + self.vision_end_token.as_str();
|
video_placeholder += self.vision_end_token.as_str();
|
||||||
}
|
}
|
||||||
let three_token = format!("{}{}{}", self.vision_start_token, self.video_token, self.vision_end_token);
|
let three_token = format!(
|
||||||
|
"{}{}{}",
|
||||||
|
self.vision_start_token, self.video_token, self.vision_end_token
|
||||||
|
);
|
||||||
if text.contains(&three_token) {
|
if text.contains(&three_token) {
|
||||||
text = text.replacen(&three_token, &video_placeholder, 1);
|
text = text.replacen(&three_token, &video_placeholder, 1);
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
+12
-18
@@ -247,7 +247,7 @@ pub fn bitor_tensor(mask1: &Tensor, mask2: &Tensor) -> Result<Tensor> {
|
|||||||
mask1.shape() == mask2.shape(),
|
mask1.shape() == mask2.shape(),
|
||||||
" bitor_tensor two tensor shape mask be equal"
|
" bitor_tensor two tensor shape mask be equal"
|
||||||
);
|
);
|
||||||
let bitor = mask1.add(&mask2)?.ne(&Tensor::zeros_like(&mask1)?)?;
|
let bitor = mask1.add(mask2)?.ne(&Tensor::zeros_like(mask1)?)?;
|
||||||
Ok(bitor)
|
Ok(bitor)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -256,38 +256,37 @@ pub fn prod_tensor_last_dim(t: &Tensor) -> Result<Tensor> {
|
|||||||
0 => t.clone(),
|
0 => t.clone(),
|
||||||
1 => {
|
1 => {
|
||||||
let data_type = t.dtype();
|
let data_type = t.dtype();
|
||||||
let t_prod = match data_type {
|
match data_type {
|
||||||
DType::U8 => {
|
DType::U8 => {
|
||||||
let t_vec = t.to_vec1::<u8>()?;
|
let t_vec = t.to_vec1::<u8>()?;
|
||||||
let prod = t_vec.iter().product::<u8>();
|
let prod = t_vec.iter().product::<u8>();
|
||||||
Tensor::from_slice(&vec![prod], 1, t.device())?
|
Tensor::from_slice(&[prod], 1, t.device())?
|
||||||
}
|
}
|
||||||
DType::U32 => {
|
DType::U32 => {
|
||||||
let t_vec = t.to_vec1::<u32>()?;
|
let t_vec = t.to_vec1::<u32>()?;
|
||||||
let prod = t_vec.iter().product::<u32>();
|
let prod = t_vec.iter().product::<u32>();
|
||||||
Tensor::from_slice(&vec![prod], 1, t.device())?
|
Tensor::from_slice(&[prod], 1, t.device())?
|
||||||
}
|
}
|
||||||
DType::I64 => {
|
DType::I64 => {
|
||||||
let t_vec = t.to_vec1::<i64>()?;
|
let t_vec = t.to_vec1::<i64>()?;
|
||||||
let prod = t_vec.iter().product::<i64>();
|
let prod = t_vec.iter().product::<i64>();
|
||||||
Tensor::from_slice(&vec![prod], 1, t.device())?
|
Tensor::from_slice(&[prod], 1, t.device())?
|
||||||
}
|
}
|
||||||
DType::F64 => {
|
DType::F64 => {
|
||||||
let t_vec = t.to_vec1::<f64>()?;
|
let t_vec = t.to_vec1::<f64>()?;
|
||||||
let prod = t_vec.iter().product::<f64>();
|
let prod = t_vec.iter().product::<f64>();
|
||||||
Tensor::from_slice(&vec![prod], 1, t.device())?
|
Tensor::from_slice(&[prod], 1, t.device())?
|
||||||
}
|
}
|
||||||
_ => {
|
_ => {
|
||||||
let t_vec = t.to_vec1::<f32>()?;
|
let t_vec = t.to_vec1::<f32>()?;
|
||||||
let prod = t_vec.iter().product::<f32>();
|
let prod = t_vec.iter().product::<f32>();
|
||||||
Tensor::from_slice(&vec![prod], 1, t.device())?
|
Tensor::from_slice(&[prod], 1, t.device())?
|
||||||
}
|
}
|
||||||
};
|
}
|
||||||
t_prod
|
|
||||||
}
|
}
|
||||||
2 => {
|
2 => {
|
||||||
let data_type = t.dtype();
|
let data_type = t.dtype();
|
||||||
let t_prod = match data_type {
|
match data_type {
|
||||||
DType::U8 => {
|
DType::U8 => {
|
||||||
let t_vec = t.to_vec2::<u8>()?;
|
let t_vec = t.to_vec2::<u8>()?;
|
||||||
let mut prod_vec = vec![];
|
let mut prod_vec = vec![];
|
||||||
@@ -333,8 +332,7 @@ pub fn prod_tensor_last_dim(t: &Tensor) -> Result<Tensor> {
|
|||||||
}
|
}
|
||||||
Tensor::new(prod_vec, t.device())?
|
Tensor::new(prod_vec, t.device())?
|
||||||
}
|
}
|
||||||
};
|
}
|
||||||
t_prod
|
|
||||||
}
|
}
|
||||||
_ => {
|
_ => {
|
||||||
return Err(anyhow!(format!("can not action this dim")));
|
return Err(anyhow!(format!("can not action this dim")));
|
||||||
@@ -343,12 +341,8 @@ pub fn prod_tensor_last_dim(t: &Tensor) -> Result<Tensor> {
|
|||||||
Ok(prod)
|
Ok(prod)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn mask_index_add(
|
pub fn mask_index_add(original: &Tensor, mask: &Tensor, add: &Tensor) -> Result<Tensor> {
|
||||||
original: &Tensor,
|
let visual_nonzero_index = nonzero_index(mask)?;
|
||||||
mask: &Tensor,
|
|
||||||
add: &Tensor,
|
|
||||||
) -> Result<Tensor> {
|
|
||||||
let visual_nonzero_index = nonzero_index(&mask)?;
|
|
||||||
let xs = original.index_add(&visual_nonzero_index, add, 0)?;
|
let xs = original.index_add(&visual_nonzero_index, add, 0)?;
|
||||||
Ok(xs)
|
Ok(xs)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
use aha::models::{
|
use aha::models::{
|
||||||
minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig, qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig
|
minicpm4::config::MiniCPM4Config, qwen2_5vl::config::Qwen2_5VLConfig,
|
||||||
|
qwen3vl::config::Qwen3VLConfig, voxcpm::config::VoxCPMConfig,
|
||||||
};
|
};
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -1,4 +1,4 @@
|
|||||||
use aha::utils::{tensor_utils::bitor_tensor};
|
use aha::utils::tensor_utils::bitor_tensor;
|
||||||
use anyhow::Result;
|
use anyhow::Result;
|
||||||
use candle_core::Tensor;
|
use candle_core::Tensor;
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user