add qwen2.5vl model
This commit is contained in:
@@ -0,0 +1,97 @@
|
||||
use candle_nn::Activation;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct VisionConfig {
|
||||
pub depth: usize,
|
||||
pub hidden_act: Activation,
|
||||
pub hidden_size: usize,
|
||||
pub intermediate_size: usize,
|
||||
pub num_heads: usize,
|
||||
pub in_chans: usize,
|
||||
pub out_hidden_size: usize,
|
||||
pub patch_size: usize,
|
||||
pub spatial_merge_size: usize,
|
||||
pub spatial_patch_size: usize,
|
||||
pub window_size: usize,
|
||||
pub fullatt_block_indexes: Vec<usize>,
|
||||
pub tokens_per_second: usize,
|
||||
pub temporal_patch_size: usize,
|
||||
}
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct RopeScaling {
|
||||
pub r#type: String,
|
||||
pub mrope_section: Vec<usize>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct Config {
|
||||
pub attention_dropout: f32,
|
||||
pub bos_token_id: usize,
|
||||
pub eos_token_id: usize,
|
||||
pub vision_start_token_id: usize,
|
||||
pub vision_end_token_id: usize,
|
||||
pub vision_token_id: usize,
|
||||
pub image_token_id: usize,
|
||||
pub video_token_id: usize,
|
||||
pub hidden_act: Activation,
|
||||
pub hidden_size: usize,
|
||||
pub initializer_range: f32,
|
||||
pub intermediate_size: usize,
|
||||
pub max_position_embeddings: usize,
|
||||
pub max_window_layers: usize,
|
||||
pub num_attention_heads: usize,
|
||||
pub num_hidden_layers: usize,
|
||||
pub num_key_value_heads: usize,
|
||||
pub rms_norm_eps: f64,
|
||||
pub rope_theta: f32,
|
||||
pub sliding_window: usize,
|
||||
pub tie_word_embeddings: bool,
|
||||
pub torch_dtype: String,
|
||||
pub use_sliding_window: bool,
|
||||
pub vision_config: VisionConfig,
|
||||
pub rope_scaling: RopeScaling,
|
||||
pub vocab_size: usize,
|
||||
}
|
||||
|
||||
pub struct VisionSetting {
|
||||
pub image_factor: u32,
|
||||
pub min_pixels: u32,
|
||||
pub max_pixels: u32,
|
||||
pub max_ratio: u32,
|
||||
pub temporal_patch_size: usize,
|
||||
pub patch_size: usize,
|
||||
pub merge_size: usize,
|
||||
pub video_min_pixels: u32,
|
||||
pub video_max_pixels: u32,
|
||||
pub video_total_pixels: u32,
|
||||
pub frame_factor: u32,
|
||||
pub fps: f32,
|
||||
pub fps_min_frames: u32,
|
||||
pub fps_max_frames: u32,
|
||||
pub image_mean: Vec<f32>,
|
||||
pub image_std: Vec<f32>,
|
||||
}
|
||||
|
||||
impl VisionSetting {
|
||||
pub fn default() -> Self {
|
||||
Self {
|
||||
image_factor: 28,
|
||||
min_pixels: 4 * 28 * 28,
|
||||
max_pixels: 16384 * 28 * 28,
|
||||
// max_pixels: 1000 * 28 * 28,
|
||||
max_ratio: 200,
|
||||
temporal_patch_size: 2,
|
||||
patch_size: 14,
|
||||
merge_size: 2,
|
||||
video_min_pixels: 128 * 28 * 28,
|
||||
video_max_pixels: 768 * 28 * 28,
|
||||
video_total_pixels: 24576 * 28 * 28,
|
||||
frame_factor: 2,
|
||||
fps: 2.0,
|
||||
fps_min_frames: 4,
|
||||
fps_max_frames: 768,
|
||||
image_mean: vec![0.48145466_f32, 0.4578275, 0.40821073],
|
||||
image_std: vec![0.26862954, 0.26130258, 0.27577711],
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
use crate::models::qwen2_5vl::config::Config;
|
||||
use crate::utils::utils::{find_safetensors_files, get_device, get_dtype};
|
||||
use crate::{
|
||||
chat_template::chat_template::ChatTemplate,
|
||||
models::{
|
||||
GenerateModel,
|
||||
qwen2_5vl::{model::Qwen2_5VLModel, processor::Qwen2_5VLProcessor},
|
||||
},
|
||||
tokenizer::tokenizer::TokenizerModel,
|
||||
};
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{D, DType, Device, IndexOp, Tensor};
|
||||
use candle_nn::VarBuilder;
|
||||
use candle_transformers::generation::LogitsProcessor;
|
||||
use openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
|
||||
pub struct Qwen2_5VLGenerateModel<'a> {
|
||||
chat_template: ChatTemplate<'a>,
|
||||
tokenizer: TokenizerModel,
|
||||
pre_processor: Qwen2_5VLProcessor,
|
||||
qwen2_5_vl: Qwen2_5VLModel,
|
||||
device: Device,
|
||||
dtype: DType,
|
||||
}
|
||||
|
||||
impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
||||
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: Config = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
||||
let device = &get_device(device);
|
||||
let cfg_dtype = cfg.torch_dtype.as_str();
|
||||
let dtype = get_dtype(dtype, cfg_dtype);
|
||||
let pre_processor = Qwen2_5VLProcessor::new(device, dtype)?;
|
||||
|
||||
let model_list = find_safetensors_files(&path)?;
|
||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||
let qwen2_5_vl = Qwen2_5VLModel::new(cfg, vb)?;
|
||||
Ok(Qwen2_5VLGenerateModel {
|
||||
chat_template,
|
||||
tokenizer,
|
||||
pre_processor,
|
||||
qwen2_5_vl,
|
||||
device: device.clone(),
|
||||
dtype: dtype,
|
||||
})
|
||||
}
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<String> {
|
||||
let temperature = match mes.temperature {
|
||||
Some(temp) => Some(temp as f64),
|
||||
None => None,
|
||||
};
|
||||
let top_p = match mes.top_p {
|
||||
Some(tp) => Some(tp as f64),
|
||||
None => None,
|
||||
};
|
||||
let mut logit_processor = LogitsProcessor::new(34562, temperature, top_p);
|
||||
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 end_of_text_id = self.qwen2_5_vl.cfg.bos_token_id as u32;
|
||||
let im_end_id = self.qwen2_5_vl.cfg.eos_token_id as u32;
|
||||
let mut pixel_values = if input.pixel_values.is_some() {
|
||||
Some(&input.pixel_values.unwrap().clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let image_grid_thw = if input.image_grid_thw.is_some() {
|
||||
Some(&input.image_grid_thw.unwrap().clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let mut pixel_values_video = if input.pixel_values_video.is_some() {
|
||||
Some(&input.pixel_values_video.unwrap().clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let video_grid_thw = if input.video_grid_thw.is_some() {
|
||||
Some(&input.video_grid_thw.unwrap().clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let second_per_grid_ts = if input.second_per_grid_ts.is_some() {
|
||||
Some(input.second_per_grid_ts.unwrap().clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let mut mask = Tensor::ones_like(&input_ids)?;
|
||||
let mut cache_position = Tensor::ones_like(&input_ids.i(0)?)?
|
||||
.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())?)?;
|
||||
|
||||
let mut generate = Vec::new();
|
||||
let sample_len = match mes.max_tokens {
|
||||
Some(max) => max,
|
||||
None => 512,
|
||||
};
|
||||
for _ in 0..sample_len {
|
||||
let logits = self.qwen2_5_vl.forward(
|
||||
&input_ids,
|
||||
pixel_values,
|
||||
image_grid_thw,
|
||||
pixel_values_video,
|
||||
video_grid_thw,
|
||||
&mask,
|
||||
Some(&cache_position),
|
||||
seqlen_offset,
|
||||
second_per_grid_ts.clone(),
|
||||
)?;
|
||||
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 == end_of_text_id || next_token == im_end_id {
|
||||
break;
|
||||
}
|
||||
seqlen_offset += seq_len;
|
||||
seq_len = 1;
|
||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||
let appendd_mask = Tensor::ones((1, 1), mask.dtype(), &self.device)?;
|
||||
mask = Tensor::cat(&[mask, appendd_mask], 1)?;
|
||||
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.qwen2_5_vl.clear_kv_cache();
|
||||
Ok(res)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
pub mod config;
|
||||
pub mod model;
|
||||
pub mod processor;
|
||||
pub mod generate;
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,462 @@
|
||||
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::qwen2_5vl::config::VisionSetting,
|
||||
utils::{
|
||||
img_utils::get_image,
|
||||
utils::{ceil_by_factor, floor_by_factor, 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>,
|
||||
pub second_per_grid_ts: Option<Vec<f32>>,
|
||||
}
|
||||
|
||||
pub struct Qwen2_5VLProcessor {
|
||||
vision_setting: VisionSetting,
|
||||
device: Device,
|
||||
dtype: DType,
|
||||
image_token: String,
|
||||
video_token: String,
|
||||
}
|
||||
|
||||
impl Qwen2_5VLProcessor {
|
||||
pub fn new(device: &Device, dtype: DType) -> Result<Self> {
|
||||
let vision_setting = VisionSetting::default();
|
||||
let image_token = "<|image_pad|>".to_string();
|
||||
let video_token = "<|video_pad|>".to_string();
|
||||
Ok(Self {
|
||||
vision_setting,
|
||||
device: device.clone(),
|
||||
dtype,
|
||||
image_token,
|
||||
video_token,
|
||||
})
|
||||
}
|
||||
|
||||
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() {
|
||||
match chat_mes {
|
||||
ChatMessage::User { content, name } => match content {
|
||||
ChatMessageContent::ContentPart(part_vec) => {
|
||||
for part in part_vec {
|
||||
match part {
|
||||
ChatMessageContentPart::Image(img_part) => {
|
||||
let img_url = img_part.image_url;
|
||||
vision_map.get_mut("image").unwrap().push(img_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) = smart_resize(img_h, img_w, &self.vision_setting, true, 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.vision_setting.temporal_patch_size;
|
||||
let grid_h = img_tensor.dim(2)? / self.vision_setting.patch_size;
|
||||
let grid_w = img_tensor.dim(3)? / self.vision_setting.patch_size;
|
||||
let shape = Shape::from(vec![
|
||||
grid_t,
|
||||
self.vision_setting.temporal_patch_size,
|
||||
channel,
|
||||
grid_h / self.vision_setting.merge_size,
|
||||
self.vision_setting.merge_size,
|
||||
self.vision_setting.patch_size,
|
||||
grid_w / self.vision_setting.merge_size,
|
||||
self.vision_setting.merge_size,
|
||||
self.vision_setting.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.vision_setting.temporal_patch_size
|
||||
* self.vision_setting.patch_size
|
||||
* self.vision_setting.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,
|
||||
})
|
||||
}
|
||||
|
||||
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 = None;
|
||||
let mut second_per_grid_ts = None;
|
||||
let vision_map = self.extract_vision_info(messages)?;
|
||||
let img_mean =
|
||||
Tensor::from_slice(&self.vision_setting.image_mean, (3, 1, 1), &self.device)?
|
||||
.to_dtype(self.dtype)?;
|
||||
let img_std = Tensor::from_slice(&self.vision_setting.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.len() > 0 {
|
||||
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();
|
||||
for file in &vec {
|
||||
let video_data = get_video_data(file, &self.vision_setting, &self.device);
|
||||
match video_data {
|
||||
Ok(tensor) => file_vec.push(tensor),
|
||||
Err(e) => println!("get_video_data err: {:?}", e),
|
||||
};
|
||||
}
|
||||
if file_vec.len() > 0 {
|
||||
let vision_input = self.process_videos(file_vec, &img_mean, &img_std);
|
||||
match vision_input {
|
||||
Ok(video_input) => {
|
||||
let video_num = video_input.grid_thw.dim(0)?;
|
||||
pixel_values_video = Some(video_input.data);
|
||||
video_grid_thw = Some(video_input.grid_thw);
|
||||
let second_per_grid = vec![
|
||||
self.vision_setting.temporal_patch_size
|
||||
as f32
|
||||
/ self.vision_setting.fps;
|
||||
video_num
|
||||
];
|
||||
second_per_grid_ts = Some(second_per_grid);
|
||||
}
|
||||
Err(e) => println!("video process_videos err: {:?}", e),
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
let merge_length = self.vision_setting.merge_size.pow(2);
|
||||
let mut text = text.to_string();
|
||||
if image_grid_thw.is_some() {
|
||||
let mut index = 0;
|
||||
while text.contains(&self.image_token) {
|
||||
let grid_i = image_grid_thw.as_ref().unwrap().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 video_grid_thw.is_some() {
|
||||
let mut index = 0;
|
||||
while text.contains(&self.video_token) {
|
||||
let grid_i = video_grid_thw.as_ref().unwrap().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.video_token, &replace, 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,
|
||||
second_per_grid_ts,
|
||||
};
|
||||
Ok(input)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn smart_resize(
|
||||
img_h: u32,
|
||||
img_w: u32,
|
||||
vision_setting: &VisionSetting,
|
||||
is_img: bool,
|
||||
video_ratio: Option<u32>,
|
||||
) -> Result<(u32, u32)> {
|
||||
if std::cmp::max(img_h, img_w) / std::cmp::min(img_h, img_w) > vision_setting.max_ratio {
|
||||
return Err(anyhow!(format!(
|
||||
"absolute aspect ratio mush be smaller than {}, got {}",
|
||||
vision_setting.max_ratio,
|
||||
std::cmp::max(img_h, img_w) / std::cmp::min(img_h, img_w)
|
||||
)));
|
||||
}
|
||||
let mut image_factor = vision_setting.image_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));
|
||||
let mut max_pixels = 0u32;
|
||||
let mut min_pixels = 0u32;
|
||||
if is_img {
|
||||
min_pixels = vision_setting.min_pixels;
|
||||
max_pixels = vision_setting.max_pixels;
|
||||
} else {
|
||||
min_pixels = vision_setting.video_min_pixels;
|
||||
max_pixels = vision_setting.video_max_pixels;
|
||||
}
|
||||
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 get_video_data(
|
||||
file: &String,
|
||||
vision_setting: &VisionSetting,
|
||||
device: &Device,
|
||||
) -> Result<Tensor> {
|
||||
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).round() as u32;
|
||||
// 1s取两帧
|
||||
let min_frames = ceil_by_factor(
|
||||
vision_setting.fps_min_frames as f32,
|
||||
vision_setting.frame_factor,
|
||||
);
|
||||
let max_frames = floor_by_factor(
|
||||
vision_setting.fps_max_frames as f32,
|
||||
vision_setting.frame_factor,
|
||||
);
|
||||
let nframes = (frames as f32 / rate as f32 * vision_setting.fps) as u32;
|
||||
let nframes = std::cmp::min(std::cmp::max(nframes, min_frames), max_frames);
|
||||
let nframes = round_by_factor(nframes, vision_setting.frame_factor);
|
||||
let sample_interval = (frames as f32 / nframes as f32).round() as u32;
|
||||
let mut frame_id = 0_u32;
|
||||
|
||||
// 图片帧使用scaler reshape的时候需要保证宽高是16的倍数,不然reshape出来的是损坏的图片
|
||||
// 所以计算resize的目标宽高时,需要用16和image_factor的最小公倍数
|
||||
let (resize_h, resize_w) = smart_resize(video_h, video_w, vision_setting, false, 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 % sample_interval == 0 {
|
||||
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()?;
|
||||
Ok(frames_tensor)
|
||||
}
|
||||
Reference in New Issue
Block a user