From dc4b5930112a20999d16d331270b892af632ca5e Mon Sep 17 00:00:00 2001 From: jhqxxx Date: Tue, 9 Dec 2025 15:52:26 +0800 Subject: [PATCH] add ffmpeg feature --- Cargo.toml | 3 ++- README.md | 36 +++++++++++++++++++++++---- src/models/qwen2_5vl/processor.rs | 8 +++++- src/models/qwen3vl/processor.rs | 9 ++++++- src/utils/video_utils.rs | 24 +++++++++--------- tests/test_gelab_zero.rs | 41 +++++++++++++++++++++++++++++++ tests/test_qwen3vl.rs | 10 +++----- 7 files changed, 105 insertions(+), 26 deletions(-) create mode 100644 tests/test_gelab_zero.rs diff --git a/Cargo.toml b/Cargo.toml index b8da2b8..5f4f683 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -14,7 +14,7 @@ candle-flash-attn = { git = "https://github.com/huggingface/candle.git", version serde = "1.0.226" serde_json = "1.0.145" anyhow = "1.0.100" -ffmpeg-next = "8.0.0" +ffmpeg-next = { version = "8.0.0", optional = true } image = "0.25.8" reqwest = { version = "0.12.23", features = ["blocking"] } base64 = "0.22.1" @@ -34,6 +34,7 @@ dirs = "6.0.0" [features] flash-attn=["candle-flash-attn"] cuda=["candle-nn/cuda", "candle-core/cuda", "candle-transformers/cuda"] +ffmpeg=["ffmpeg-next"] [lints.clippy] needless_range_loop = "allow" diff --git a/README.md b/README.md index 255ba3c..739067b 100644 --- a/README.md +++ b/README.md @@ -26,13 +26,39 @@ ⭐ 如果这个项目对你有帮助,请给我们一个 Star! ## 环境依赖 -1. ffmpeg: -* ubuntu/WSL +* 启用ffmpeg的feature时: + * ubuntu/WSL + ```bash + sudo apt-get update + sudo apt-get install -y clang pkg-config ffmpeg libavutil-dev libavcodec-dev libavformat-dev libavfilter-dev libavdevice-dev libswresample-dev libswscale-dev + ``` + * windows参考: https://github.com/zmwangx/rust-ffmpeg/wiki/Notes-on-building + +## 功能特性 +项目提供了几个可选的功能特性,您可以根据需要启用它们: +* flash-attn: 启用 Flash Attention 支持以提升模型推理性能: ```bash -sudo apt-get update -sudo apt-get install -y clang pkg-config ffmpeg libavutil-dev libavcodec-dev libavformat-dev libavfilter-dev libavdevice-dev libswresample-dev libswscale-dev +cargo build --features flash-attn +``` + +* cuda: 为 candle 核心组件启用 CUDA 支持,实现 GPU 加速计算: +```bash +cargo build --features cuda +``` + +* ffmpeg: 启用 FFmpeg 支持,提供多媒体处理功能: +```bash +cargo build --features ffmpeg +``` +* 组合使用功能特性 + +```bash +# 同时启用 CUDA 和 Flash Attention 以获得最佳性能 +cargo build --features "cuda,flash-attn" + +# 启用所有功能特性 +cargo build --features "cuda,flash-attn,ffmpeg" ``` -* windows参考: https://github.com/zmwangx/rust-ffmpeg/wiki/Notes-on-building ## 安装及使用 diff --git a/src/models/qwen2_5vl/processor.rs b/src/models/qwen2_5vl/processor.rs index c4ccf10..e15b696 100644 --- a/src/models/qwen2_5vl/processor.rs +++ b/src/models/qwen2_5vl/processor.rs @@ -5,6 +5,7 @@ use aha_openai_dive::v1::resources::chat::{ }; use anyhow::{Result, anyhow}; use candle_core::{DType, Device, IndexOp, Shape, Tensor}; +#[cfg(feature = "ffmpeg")] use ffmpeg_next as ffmpeg; use image::DynamicImage; use num::integer::lcm; @@ -33,6 +34,7 @@ pub struct GeneralInput { pub second_per_grid_ts: Option>, } +#[allow(unused)] pub struct Qwen2_5VLProcessor { vision_setting: VisionSetting, device: Device, @@ -213,6 +215,7 @@ impl Qwen2_5VLProcessor { }) } + #[allow(unused_mut)] pub fn process_info( &self, messages: &ChatCompletionParameters, @@ -221,7 +224,7 @@ impl Qwen2_5VLProcessor { 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 video_grid_thw: Option = None; let mut second_per_grid_ts = None; let vision_map = self.extract_vision_info(messages)?; let img_mean = @@ -251,6 +254,7 @@ impl Qwen2_5VLProcessor { }; } } + #[cfg(feature = "ffmpeg")] if key.eq("video") { let mut file_vec = Vec::new(); for file in &vec { @@ -294,6 +298,7 @@ impl Qwen2_5VLProcessor { } text = text.replace("<|placeholder|>", &self.image_token); } + #[cfg(feature = "ffmpeg")] if let Some(ref video_grid_thw) = video_grid_thw { let mut index = 0; while text.contains(&self.video_token) { @@ -359,6 +364,7 @@ pub fn smart_resize( Ok((h_bar, w_bar)) } +#[cfg(feature = "ffmpeg")] pub fn get_video_data( file: &String, vision_setting: &VisionSetting, diff --git a/src/models/qwen3vl/processor.rs b/src/models/qwen3vl/processor.rs index 3125c6c..1b73058 100644 --- a/src/models/qwen3vl/processor.rs +++ b/src/models/qwen3vl/processor.rs @@ -5,6 +5,7 @@ use aha_openai_dive::v1::resources::chat::{ }; use anyhow::{Result, anyhow}; use candle_core::{DType, Device, IndexOp, Shape, Tensor}; +#[cfg(feature = "ffmpeg")] use ffmpeg_next as ffmpeg; use image::DynamicImage; use num::integer::lcm; @@ -44,6 +45,7 @@ pub struct VideoMetadata { frame_indices: Vec, } +#[allow(unused)] pub struct Qwen3VLProcessor { img_process_cfg: PreprocessorConfig, video_process_cfg: PreprocessorConfig, @@ -256,6 +258,7 @@ impl Qwen3VLProcessor { }) } + #[allow(unused)] fn calculate_timestamps( &self, frames_indices: Vec, @@ -282,6 +285,7 @@ impl Qwen3VLProcessor { Ok(stamps) } + #[allow(unused)] pub fn process_info( &self, messages: &ChatCompletionParameters, @@ -291,7 +295,7 @@ impl Qwen3VLProcessor { let mut image_grid_thw = None; let mut pixel_values_video = None; let mut video_grid_thw: Option = None; - let mut video_metadata = None; + let mut video_metadata: Option> = 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)? @@ -320,6 +324,7 @@ impl Qwen3VLProcessor { }; } } + #[cfg(feature = "ffmpeg")] if key.eq("video") { let mut file_vec = Vec::new(); let mut video_infos = Vec::new(); @@ -371,6 +376,7 @@ impl Qwen3VLProcessor { } text = text.replace("<|placeholder|>", &self.image_token); } + #[cfg(feature = "ffmpeg")] if let Some(ref video_grid_thw) = video_grid_thw { let mut index = 0; while text.contains(&self.video_token) { @@ -471,6 +477,7 @@ pub fn video_smart_resize( Ok((h_bar, w_bar)) } +#[cfg(feature = "ffmpeg")] pub fn get_video_data( file: &String, patch_size: u32, diff --git a/src/utils/video_utils.rs b/src/utils/video_utils.rs index b6efeee..0545cb0 100644 --- a/src/utils/video_utils.rs +++ b/src/utils/video_utils.rs @@ -1,14 +1,14 @@ -use std::{fs::File, io::Write}; +// use std::{fs::File, io::Write}; -use ffmpeg_next as ffmpeg; +// use ffmpeg_next as ffmpeg; -#[allow(unused)] -fn save_file( - frame: &ffmpeg::frame::Video, - index: usize, -) -> std::result::Result<(), std::io::Error> { - let mut file = File::create(format!("frame{}.ppm", index))?; - file.write_all(format!("P6\n{} {}\n255\n", frame.width(), frame.height()).as_bytes())?; - file.write_all(frame.data(0))?; - Ok(()) -} +// #[allow(unused)] +// fn save_file( +// frame: &ffmpeg::frame::Video, +// index: usize, +// ) -> std::result::Result<(), std::io::Error> { +// let mut file = File::create(format!("frame{}.ppm", index))?; +// file.write_all(format!("P6\n{} {}\n255\n", frame.width(), frame.height()).as_bytes())?; +// file.write_all(frame.data(0))?; +// Ok(()) +// } diff --git a/tests/test_gelab_zero.rs b/tests/test_gelab_zero.rs new file mode 100644 index 0000000..29db2d7 --- /dev/null +++ b/tests/test_gelab_zero.rs @@ -0,0 +1,41 @@ +use std::time::Instant; + +use aha::models::{GenerateModel, qwen3vl::generate::Qwen3VLGenerateModel}; +use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; +use anyhow::Result; + +#[test] +fn gelab_zero_generate() -> Result<()> { + // test with cuda: RUST_BACKTRACE=1 cargo test -F cuda gelab_zero_generate -r -- --nocapture + + let model_path = "/home/jhq/huggingface_model/stepfun-ai/GELab-Zero-4B-preview"; + + let message = r#" + { + "model": "gelab-zero", + "messages": [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Hello, GELab-Zero!" + } + ] + } + ] + } + "#; + 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(()) +} diff --git a/tests/test_qwen3vl.rs b/tests/test_qwen3vl.rs index 31c20ce..02ed4bc 100644 --- a/tests/test_qwen3vl.rs +++ b/tests/test_qwen3vl.rs @@ -19,23 +19,21 @@ fn qwen3vl_generate() -> Result<()> { "role": "user", "content": [ { - "type": "image", - "image_url": + "type": "video", + "video_url": { - "url": "file://./assets/img/ocr_test1.png" + "url": "./assets/video/video_test.mp4" } }, { "type": "text", - "text": "请分析图片并提取所有可见文本内容,按从左到右、从上到下的布局,返回纯文本" + "text": "视频里发生了什么" } ] } ] } "#; - // ./assets/video/video_test.mp4 - let mes: ChatCompletionParameters = serde_json::from_str(message)?; let i_start = Instant::now(); let mut qwen3vl = Qwen3VLGenerateModel::init(model_path, None, None)?;