merge pr/guobin211/12

This commit is contained in:
jhqxxx
2025-12-26 00:23:14 +08:00
6 changed files with 517 additions and 43 deletions
Generated
+1
View File
@@ -37,6 +37,7 @@ dependencies = [
"minijinja",
"modelscope",
"num",
"rayon",
"reqwest",
"rocket",
"serde",
+1
View File
@@ -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"]
+112 -30
View File
@@ -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<Vec<RgbaImage>> {
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<Result<RgbaImage>> = 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()
}
}
+8 -12
View File
@@ -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<DynamicImage> {
Err(anyhow!("get image from message failed".to_string()))
}
pub fn extract_image_url(mes: &ChatCompletionParameters) -> Result<Vec<String>> {
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<Vec<DynamicImage>> {
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)> {
-1
View File
@@ -1,4 +1,3 @@
use aha::utils::get_default_save_dir;
use anyhow::Result;
use candle_core::Tensor;
+395
View File
@@ -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(())
}