From f35631157ab8aae687e1cea1fee75bafbbf14509 Mon Sep 17 00:00:00 2001 From: michaelbguo Date: Wed, 24 Dec 2025 21:42:05 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E4=BC=98=E5=8C=96RMBG2.0=E5=9B=BE?= =?UTF-8?q?=E5=83=8F=E5=A4=84=E7=90=86=E6=80=A7=E8=83=BD=EF=BC=8C=E6=94=AF?= =?UTF-8?q?=E6=8C=81=E5=B9=B6=E8=A1=8C=E6=89=B9=E9=87=8F=E5=A4=84=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 1 + Cargo.toml | 1 + src/models/rmbg2_0/generate.rs | 132 ++++++++--- src/utils/img_utils.rs | 24 +- tests/test_rmbg2_0_perf.rs | 396 +++++++++++++++++++++++++++++++++ 5 files changed, 511 insertions(+), 43 deletions(-) create mode 100644 tests/test_rmbg2_0_perf.rs diff --git a/Cargo.lock b/Cargo.lock index 484227d..6545371 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 0b77857..707ac2f 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 1cff361..d355d7c 100644 --- a/src/models/rmbg2_0/generate.rs +++ b/src/models/rmbg2_0/generate.rs @@ -1,8 +1,9 @@ use aha_openai_dive::v1::resources::chat::ChatCompletionParameters; use anyhow::Result; -use candle_core::{DType, Device, Tensor}; +use candle_core::{DType, Device, IndexOp, Tensor}; use candle_nn::VarBuilder; -use image::{Rgba, RgbaImage}; +use image::RgbaImage; +use rayon::prelude::*; use crate::{ models::rmbg2_0::model::BiRefNet, @@ -44,37 +45,106 @@ impl RMBG2_0 { }) } + pub fn h(&self) -> u32 { + self.h + } + + pub fn w(&self) -> u32 { + self.w + } + + pub fn img_mean(&self) -> &Tensor { + &self.img_mean + } + + pub fn img_std(&self) -> &Tensor { + &self.img_std + } + + pub fn device(&self) -> &Device { + &self.device + } + + pub fn dtype(&self) -> DType { + self.dtype + } + + pub fn model(&self) -> &BiRefNet { + &self.model + } + pub fn generate(&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 + let batch_tensor = Tensor::stack(&tensors, 0)?; + let batch_output = self.model.forward(&batch_tensor)?; + + // 并行后处理:生成 RGBA 图像 + let results: Vec> = meta + .into_par_iter() + .enumerate() + .map(|(i, (img, height, width))| { + let rmbg_tensor = batch_output.i(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..0aee160 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,30 @@ 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 { let mut img_vec = Vec::new(); - for chat_mes in mes.messages.clone() { + // 使用引用避免 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.clone()); } } } } - 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/test_rmbg2_0_perf.rs b/tests/test_rmbg2_0_perf.rs new file mode 100644 index 0000000..8793f3c --- /dev/null +++ b/tests/test_rmbg2_0_perf.rs @@ -0,0 +1,396 @@ +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 = 50; + + 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 = 20; + + 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(()) +}