diff --git a/Cargo.lock b/Cargo.lock index 44c4e7a..05f2c36 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -37,6 +37,7 @@ dependencies = [ "minijinja", "modelscope", "num", + "rayon", "reqwest", "rocket", "serde", diff --git a/Cargo.toml b/Cargo.toml index d864a72..752fa9f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -31,6 +31,7 @@ clap = { version = "4.5.51", features = ["derive"] } modelscope = "0.1.0" dirs = "6.0.0" url = "2.5.7" +rayon = "1.10" [features] flash-attn=["candle-flash-attn"] diff --git a/src/models/rmbg2_0/generate.rs b/src/models/rmbg2_0/generate.rs index 57e0a5e..961b90c 100644 --- a/src/models/rmbg2_0/generate.rs +++ b/src/models/rmbg2_0/generate.rs @@ -7,7 +7,8 @@ use anyhow::Result; use base64::{Engine, prelude::BASE64_STANDARD}; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; -use image::{Rgba, RgbaImage}; +use image::RgbaImage; +use rayon::prelude::*; use rocket::futures::{Stream, stream}; use crate::{ @@ -52,38 +53,119 @@ impl RMBG2_0Model { }) } + #[cfg(test)] + pub fn h(&self) -> u32 { + self.h + } + + #[cfg(test)] + pub fn w(&self) -> u32 { + self.w + } + + #[cfg(test)] + pub fn img_mean(&self) -> &Tensor { + &self.img_mean + } + + #[cfg(test)] + pub fn img_std(&self) -> &Tensor { + &self.img_std + } + + #[cfg(test)] + pub fn device(&self) -> &Device { + &self.device + } + #[cfg(test)] + pub fn dtype(&self) -> DType { + self.dtype + } + + #[cfg(test)] + pub fn model(&self) -> &BiRefNet { + &self.model + } pub fn inference(&self, mes: ChatCompletionParameters) -> Result> { let imgs = extract_images(&mes)?; - let mut rmbg_png = vec![]; - for img in imgs { - let height = img.height(); - let width = img.width(); - let img_tensor = img_transform_with_resize( - &img, - self.h, - self.w, - &self.img_mean, - &self.img_std, - &self.device, - self.dtype, - )? - .unsqueeze(0)?; - let rmbg_img = self.model.forward(&img_tensor)?.squeeze(0)?; - let alpha_img = float_tensor_to_dynamic_image(&rmbg_img)?; - let alpha_img = - alpha_img.resize_exact(width, height, image::imageops::FilterType::CatmullRom); - let alpha_gray = alpha_img.to_luma8(); - let mut rgba_img = RgbaImage::new(width, height); - - // 遍历像素并组合 - for (x, y, pixel) in img.to_rgb8().enumerate_pixels() { - let alpha_value = alpha_gray.get_pixel(x, y).0[0]; - let rgba_pixel = Rgba([pixel.0[0], pixel.0[1], pixel.0[2], alpha_value]); - rgba_img.put_pixel(x, y, rgba_pixel); - } - rmbg_png.push(rgba_img); + if imgs.is_empty() { + return Ok(vec![]); } - Ok(rmbg_png) + + // 并行预处理:提取原始尺寸和转换为 tensor + let preprocessed: Vec<_> = imgs + .par_iter() + .map(|img| { + let height = img.height(); + let width = img.width(); + let tensor = img_transform_with_resize( + img, + self.h, + self.w, + &self.img_mean, + &self.img_std, + &self.device, + self.dtype, + ); + (img.clone(), height, width, tensor) + }) + .collect(); + + // 检查预处理是否有错误 + let mut tensors = Vec::with_capacity(preprocessed.len()); + let mut meta: Vec<_> = Vec::with_capacity(preprocessed.len()); + for (img, height, width, tensor_result) in preprocessed { + let tensor = tensor_result?; + tensors.push(tensor); + meta.push((img, height, width)); + } + + // 批量推理:将所有图片合并为一个 batch + // to guobin211: 感谢你贡献的代码,不过现在模型中可变形卷积的实现只支持batch_size=1,所以推理还是用的循环QaQ + // let batch_tensor = Tensor::stack(&tensors, 0)?; + // let batch_output = self.model.forward(&batch_tensor)?; + let mut batch_output = vec![]; + for img_tensor in tensors { + let output = self.model.forward(&img_tensor.unsqueeze(0)?)?.squeeze(0)?; + batch_output.push(output); + } + + // 并行后处理:生成 RGBA 图像 + let results: Vec> = meta + .into_par_iter() + .enumerate() + .map(|(i, (img, height, width))| { + // let rmbg_tensor = batch_output.i(i)?; + let rmbg_tensor = &batch_output[i]; + let alpha_img = float_tensor_to_dynamic_image(rmbg_tensor)?; + let alpha_img = + alpha_img.resize_exact(width, height, image::imageops::FilterType::CatmullRom); + let alpha_gray = alpha_img.to_luma8(); + + let rgb_img = img.to_rgb8(); + let rgb_raw = rgb_img.as_raw(); + let alpha_raw = alpha_gray.as_raw(); + let pixel_count = (width * height) as usize; + let mut rgba_raw = vec![0u8; pixel_count * 4]; + + // 并行分块写入 + rgba_raw + .par_chunks_mut(4) + .enumerate() + .for_each(|(idx, chunk)| { + let src = idx * 3; + chunk[0] = rgb_raw[src]; + chunk[1] = rgb_raw[src + 1]; + chunk[2] = rgb_raw[src + 2]; + chunk[3] = alpha_raw[idx]; + }); + + RgbaImage::from_raw(width, height, rgba_raw) + .ok_or_else(|| anyhow::anyhow!("Failed to create RGBA image")) + }) + .collect(); + + results.into_iter().collect() } } diff --git a/src/utils/img_utils.rs b/src/utils/img_utils.rs index 1dba19b..7120ab0 100644 --- a/src/utils/img_utils.rs +++ b/src/utils/img_utils.rs @@ -8,6 +8,7 @@ use anyhow::{Result, anyhow}; use base64::{Engine, engine::general_purpose}; use candle_core::{DType, Device, Tensor}; use image::{DynamicImage, ImageBuffer, ImageReader, Rgb, RgbImage, imageops}; +use rayon::prelude::*; use crate::utils::{ceil_by_factor, floor_by_factor, round_by_factor}; @@ -78,31 +79,26 @@ pub fn get_image(file: &str) -> Result { Err(anyhow!("get image from message failed".to_string())) } -pub fn extract_image_url(mes: &ChatCompletionParameters) -> Result> { +pub fn extract_image_url(mes: &ChatCompletionParameters) -> Vec<&String> { let mut img_vec = Vec::new(); - for chat_mes in mes.messages.clone() { + for chat_mes in &mes.messages { 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; - img_vec.push(img_url.url); + img_vec.push(&img_part.image_url.url); } } } } - Ok(img_vec) + img_vec } pub fn extract_images(mes: &ChatCompletionParameters) -> Result> { - let img_url_vec = extract_image_url(mes)?; - let mut img_vec = Vec::new(); - for url in img_url_vec { - let img = get_image(&url)?; - img_vec.push(img); - } - Ok(img_vec) + let img_url_vec = extract_image_url(mes); + // 并行下载图片 + img_url_vec.par_iter().map(|url| get_image(url)).collect() } pub fn generate_target_ratios_sorted(min_num: u32, max_num: u32) -> Vec<(u32, u32)> { diff --git a/tests/messy_test.rs b/tests/messy_test.rs index 9a0a5cc..066c450 100644 --- a/tests/messy_test.rs +++ b/tests/messy_test.rs @@ -1,4 +1,3 @@ -use aha::utils::get_default_save_dir; use anyhow::Result; use candle_core::Tensor; diff --git a/tests/test_rmbg2_0_perf.rs b/tests/test_rmbg2_0_perf.rs new file mode 100644 index 0000000..f8f6a4d --- /dev/null +++ b/tests/test_rmbg2_0_perf.rs @@ -0,0 +1,395 @@ +use std::time::Instant; + +use anyhow::Result; +use image::{ImageReader, Rgba, RgbaImage}; +use rayon::prelude::*; + +/// 测试像素组合性能对比 +#[test] +fn test_pixel_combine_performance() -> Result<()> { + // cargo test test_pixel_combine_performance -r -- --nocapture + let img_path = "./assets/img/gougou.jpg"; + let img = ImageReader::open(img_path)?.decode()?; + // 缩小图片以加快测试 + let img = img.resize(1024, 1024, image::imageops::FilterType::Nearest); + let width = img.width(); + let height = img.height(); + let rgb_img = img.to_rgb8(); + + // 模拟 alpha 通道 + let alpha_gray = + image::GrayImage::from_fn(width, height, |x, y| image::Luma([((x + y) % 256) as u8])); + + let iterations = 10; + + println!("=== 像素组合性能测试 ==="); + println!("图片尺寸: {}x{}", width, height); + println!(); + + // 旧方法:逐像素操作 + let start = Instant::now(); + for _ in 0..iterations { + let mut rgba_img = RgbaImage::new(width, height); + for (x, y, pixel) in rgb_img.enumerate_pixels() { + let alpha_value = alpha_gray.get_pixel(x, y).0[0]; + let rgba_pixel = Rgba([pixel.0[0], pixel.0[1], pixel.0[2], alpha_value]); + rgba_img.put_pixel(x, y, rgba_pixel); + } + std::hint::black_box(&rgba_img); + } + let old_duration = start.elapsed(); + println!( + "旧方法(逐像素): {:?}, 平均: {:?}", + old_duration, + old_duration / iterations + ); + + // 新方法:串行索引赋值 + let start = Instant::now(); + for _ in 0..iterations { + let rgb_raw = rgb_img.as_raw(); + let alpha_raw = alpha_gray.as_raw(); + let pixel_count = (width * height) as usize; + let mut rgba_raw = vec![0u8; pixel_count * 4]; + + for i in 0..pixel_count { + let dst = i * 4; + let src = i * 3; + rgba_raw[dst] = rgb_raw[src]; + rgba_raw[dst + 1] = rgb_raw[src + 1]; + rgba_raw[dst + 2] = rgb_raw[src + 2]; + rgba_raw[dst + 3] = alpha_raw[i]; + } + + let rgba_img = RgbaImage::from_raw(width, height, rgba_raw).unwrap(); + std::hint::black_box(&rgba_img); + } + let serial_duration = start.elapsed(); + println!( + "新方法(串行索引): {:?}, 平均: {:?}", + serial_duration, + serial_duration / iterations + ); + + // 新方法:并行分块写入 + let start = Instant::now(); + for _ in 0..iterations { + let rgb_raw = rgb_img.as_raw(); + let alpha_raw = alpha_gray.as_raw(); + let pixel_count = (width * height) as usize; + let mut rgba_raw = vec![0u8; pixel_count * 4]; + + rgba_raw + .par_chunks_mut(4) + .enumerate() + .for_each(|(i, chunk)| { + let src = i * 3; + chunk[0] = rgb_raw[src]; + chunk[1] = rgb_raw[src + 1]; + chunk[2] = rgb_raw[src + 2]; + chunk[3] = alpha_raw[i]; + }); + + let rgba_img = RgbaImage::from_raw(width, height, rgba_raw).unwrap(); + std::hint::black_box(&rgba_img); + } + let parallel_duration = start.elapsed(); + println!( + "新方法(并行): {:?}, 平均: {:?}", + parallel_duration, + parallel_duration / iterations + ); + + let speedup_serial = old_duration.as_secs_f64() / serial_duration.as_secs_f64(); + let speedup_parallel = old_duration.as_secs_f64() / parallel_duration.as_secs_f64(); + println!(); + println!("串行索引 vs 逐像素: {:.2}x", speedup_serial); + println!("并行 vs 逐像素: {:.2}x", speedup_parallel); + + Ok(()) +} + +/// 测试图片 resize 性能对比(串行 vs 并行) +#[test] +fn test_image_resize_parallel_vs_serial() -> Result<()> { + // cargo test test_image_resize_parallel_vs_serial -r -- --nocapture + let img_path = "./assets/img/gougou.jpg"; + let img = ImageReader::open(img_path)?.decode()?; + // 缩小原图以加快测试 + let img = img.resize(2048, 2048, image::imageops::FilterType::Nearest); + + let num_images = 4; + let imgs: Vec<_> = (0..num_images).map(|_| img.clone()).collect(); + let target_h = 1024u32; + let target_w = 1024u32; + + let iterations = 10; + + println!("=== 图片 Resize 性能测试 ==="); + println!("图片数量: {}", num_images); + println!("原始尺寸: {}x{}", img.width(), img.height()); + println!("目标尺寸: {}x{}", target_w, target_h); + println!(); + + // 串行 resize + let start = Instant::now(); + for _ in 0..iterations { + let mut results = Vec::with_capacity(num_images); + for img in &imgs { + let resized = + img.resize_exact(target_w, target_h, image::imageops::FilterType::CatmullRom); + results.push(resized); + } + std::hint::black_box(&results); + } + let serial_duration = start.elapsed(); + println!( + "串行 resize: {:?}, 平均: {:?}", + serial_duration, + serial_duration / iterations + ); + + // 并行 resize + let start = Instant::now(); + for _ in 0..iterations { + let results: Vec<_> = imgs + .par_iter() + .map(|img| { + img.resize_exact(target_w, target_h, image::imageops::FilterType::CatmullRom) + }) + .collect(); + std::hint::black_box(&results); + } + let parallel_duration = start.elapsed(); + println!( + "并行 resize: {:?}, 平均: {:?}", + parallel_duration, + parallel_duration / iterations + ); + + let speedup = serial_duration.as_secs_f64() / parallel_duration.as_secs_f64(); + println!(); + println!("并行 vs 串行: {:.2}x", speedup); + + Ok(()) +} + +/// 测试后处理阶段并行 vs 串行性能(纯图像操作) +#[test] +fn test_postprocess_parallel_vs_serial() -> Result<()> { + // cargo test test_postprocess_parallel_vs_serial -r -- --nocapture + let img_path = "./assets/img/gougou.jpg"; + let img = ImageReader::open(img_path)?.decode()?; + // 缩小图片以加快测试 + let img = img.resize(1024, 1024, image::imageops::FilterType::Nearest); + let width = img.width(); + let height = img.height(); + + // 模拟多张图片的后处理数据 + let num_images = 4; + let rgb_imgs: Vec<_> = (0..num_images).map(|_| img.to_rgb8()).collect(); + let alpha_grays: Vec<_> = (0..num_images) + .map(|_| { + image::GrayImage::from_fn(width, height, |x, y| image::Luma([((x + y) % 256) as u8])) + }) + .collect(); + + let iterations = 10; + + println!("=== 后处理阶段性能测试(纯图像操作)==="); + println!("图片数量: {}", num_images); + println!("图片尺寸: {}x{}", width, height); + println!(); + + // 串行后处理(for-in 循环) + let start = Instant::now(); + for _ in 0..iterations { + let mut results = Vec::with_capacity(num_images); + for i in 0..num_images { + let rgb_raw = rgb_imgs[i].as_raw(); + let alpha_raw = alpha_grays[i].as_raw(); + let pixel_count = (width * height) as usize; + let mut rgba_raw = vec![0u8; pixel_count * 4]; + + for j in 0..pixel_count { + let dst = j * 4; + let src = j * 3; + rgba_raw[dst] = rgb_raw[src]; + rgba_raw[dst + 1] = rgb_raw[src + 1]; + rgba_raw[dst + 2] = rgb_raw[src + 2]; + rgba_raw[dst + 3] = alpha_raw[j]; + } + + let rgba_img = RgbaImage::from_raw(width, height, rgba_raw).unwrap(); + results.push(rgba_img); + } + std::hint::black_box(&results); + } + let serial_duration = start.elapsed(); + println!( + "串行后处理: {:?}, 平均: {:?}", + serial_duration, + serial_duration / iterations + ); + + // 并行后处理(外层并行 + 内层并行) + let start = Instant::now(); + for _ in 0..iterations { + let results: Vec<_> = (0..num_images) + .into_par_iter() + .map(|i| { + let rgb_raw = rgb_imgs[i].as_raw(); + let alpha_raw = alpha_grays[i].as_raw(); + let pixel_count = (width * height) as usize; + let mut rgba_raw = vec![0u8; pixel_count * 4]; + + rgba_raw + .par_chunks_mut(4) + .enumerate() + .for_each(|(j, chunk)| { + let src = j * 3; + chunk[0] = rgb_raw[src]; + chunk[1] = rgb_raw[src + 1]; + chunk[2] = rgb_raw[src + 2]; + chunk[3] = alpha_raw[j]; + }); + + RgbaImage::from_raw(width, height, rgba_raw).unwrap() + }) + .collect(); + std::hint::black_box(&results); + } + let parallel_duration = start.elapsed(); + println!( + "并行后处理: {:?}, 平均: {:?}", + parallel_duration, + parallel_duration / iterations + ); + + let speedup = serial_duration.as_secs_f64() / parallel_duration.as_secs_f64(); + println!(); + println!("后处理性能提升: {:.2}x", speedup); + + Ok(()) +} + +/// 测试完整图像处理流程(resize + 像素合并)串行 vs 并行 +#[test] +fn test_full_image_pipeline_parallel_vs_serial() -> Result<()> { + // cargo test test_full_image_pipeline_parallel_vs_serial -r -- --nocapture + let img_path = "./assets/img/gougou.jpg"; + let img = ImageReader::open(img_path)?.decode()?; + // 缩小图片以加快测试 + let img = img.resize(2048, 2048, image::imageops::FilterType::Nearest); + let orig_width = img.width(); + let orig_height = img.height(); + + let num_images = 4; + let imgs: Vec<_> = (0..num_images).map(|_| img.clone()).collect(); + let target_h = 1024u32; + let target_w = 1024u32; + + // 模拟 alpha 蒙版 + let alpha_grays: Vec<_> = (0..num_images) + .map(|_| { + image::GrayImage::from_fn(orig_width, orig_height, |x, y| { + image::Luma([((x + y) % 256) as u8]) + }) + }) + .collect(); + + let iterations = 10; + + println!("=== 完整图像处理流程性能测试 ==="); + println!("图片数量: {}", num_images); + println!("原始尺寸: {}x{}", orig_width, orig_height); + println!("处理尺寸: {}x{}", target_w, target_h); + println!(); + + // 串行处理流程 + let start = Instant::now(); + for _ in 0..iterations { + let mut results = Vec::with_capacity(num_images); + for i in 0..num_images { + // 预处理:resize 到模型输入尺寸 + let _resized = + imgs[i].resize_exact(target_w, target_h, image::imageops::FilterType::CatmullRom); + + // 后处理:合并 RGB 和 alpha + let rgb_img = imgs[i].to_rgb8(); + let rgb_raw = rgb_img.as_raw(); + let alpha_raw = alpha_grays[i].as_raw(); + let pixel_count = (orig_width * orig_height) as usize; + let mut rgba_raw = vec![0u8; pixel_count * 4]; + + for j in 0..pixel_count { + let dst = j * 4; + let src = j * 3; + rgba_raw[dst] = rgb_raw[src]; + rgba_raw[dst + 1] = rgb_raw[src + 1]; + rgba_raw[dst + 2] = rgb_raw[src + 2]; + rgba_raw[dst + 3] = alpha_raw[j]; + } + + let rgba_img = RgbaImage::from_raw(orig_width, orig_height, rgba_raw).unwrap(); + results.push(rgba_img); + } + std::hint::black_box(&results); + } + let serial_duration = start.elapsed(); + println!( + "串行流程: {:?}, 平均: {:?}", + serial_duration, + serial_duration / iterations + ); + + // 并行处理流程 + let start = Instant::now(); + for _ in 0..iterations { + // 并行预处理 + let _resized: Vec<_> = imgs + .par_iter() + .map(|img| { + img.resize_exact(target_w, target_h, image::imageops::FilterType::CatmullRom) + }) + .collect(); + + // 并行后处理 + let results: Vec<_> = (0..num_images) + .into_par_iter() + .map(|i| { + let rgb_img = imgs[i].to_rgb8(); + let rgb_raw = rgb_img.as_raw(); + let alpha_raw = alpha_grays[i].as_raw(); + let pixel_count = (orig_width * orig_height) as usize; + let mut rgba_raw = vec![0u8; pixel_count * 4]; + + rgba_raw + .par_chunks_mut(4) + .enumerate() + .for_each(|(j, chunk)| { + let src = j * 3; + chunk[0] = rgb_raw[src]; + chunk[1] = rgb_raw[src + 1]; + chunk[2] = rgb_raw[src + 2]; + chunk[3] = alpha_raw[j]; + }); + + RgbaImage::from_raw(orig_width, orig_height, rgba_raw).unwrap() + }) + .collect(); + std::hint::black_box(&results); + } + let parallel_duration = start.elapsed(); + println!( + "并行流程: {:?}, 平均: {:?}", + parallel_duration, + parallel_duration / iterations + ); + + let speedup = serial_duration.as_secs_f64() / parallel_duration.as_secs_f64(); + println!(); + println!("完整流程性能提升: {:.2}x", speedup); + + Ok(()) +}