use std::collections::HashMap;
use std::path::{Path, PathBuf};
use rayon::prelude::*;
use tch::{Device, Kind, Tensor};
use av_core::error::{AvError, AvResult};
use av_core::geometry::{letterbox, Aabb, Letterbox};
use av_tasks::augment::{scaled_dims, AugmentPlan};
pub struct SampleTensor {
pub x: Tensor,
pub boxes: Vec<[f32; 4]>,
pub labels: Vec<u32>,
}
impl Clone for SampleTensor {
fn clone(&self) -> Self {
Self {
x: self.x.copy(),
boxes: self.boxes.clone(),
labels: self.labels.clone(),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ResizeMode {
#[default]
Letterbox,
Stretch,
}
pub fn load_yolo_dir(
root: &Path,
split: &str,
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Vec<SampleTensor>> {
load_yolo_dir_with_mode(
root,
split,
img_size,
device,
ResizeMode::Letterbox,
imagenet_norm,
)
}
pub fn load_yolo_dir_with_mode(
root: &Path,
split: &str,
img_size: u32,
device: Device,
mode: ResizeMode,
imagenet_norm: bool,
) -> AvResult<Vec<SampleTensor>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
if !img_dir.is_dir() {
return Err(AvError::data(format!(
"数据集图片目录不存在: {}",
img_dir.display()
)));
}
let mut entries: Vec<PathBuf> = std::fs::read_dir(&img_dir)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| {
matches!(
p.extension().and_then(|e| e.to_str()),
Some("jpg") | Some("jpeg") | Some("png") | Some("bmp")
)
})
.collect();
entries.sort();
if entries.is_empty() {
return Err(AvError::data(format!(
"数据集图片目录为空: {}",
img_dir.display()
)));
}
let samples: Vec<SampleTensor> = entries
.par_iter()
.map(|img_path| -> AvResult<SampleTensor> {
let stem = img_path
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let rgb = image::open(img_path)
.map_err(|e| AvError::data(format!("读图失败 {}: {e}", img_path.display())))?
.to_rgb8();
let (ow, oh) = (rgb.width(), rgb.height());
let lb = match mode {
ResizeMode::Letterbox => Some(letterbox(ow, oh, img_size, img_size)),
ResizeMode::Stretch => None,
};
let (mut boxes, mut labels) = (Vec::new(), Vec::new());
if lbl_path.exists() {
let text = std::fs::read_to_string(&lbl_path)?;
let (px_boxes, px_labels) =
parse_yolo_label_text(&text, ow, oh, &lbl_path.display().to_string())?;
boxes = px_boxes
.iter()
.map(|g| match lb {
Some(lb) => {
let m = lb.map_box(Aabb::new(g[0], g[1], g[2], g[3]));
[m.x1, m.y1, m.x2, m.y2]
}
None => {
let (sx, sy) =
(img_size as f32 / ow as f32, img_size as f32 / oh as f32);
[g[0] * sx, g[1] * sy, g[2] * sx, g[3] * sy]
}
})
.collect();
labels = px_labels;
}
let x = rgb_to_input_tensor(&rgb, img_size, lb, device, imagenet_norm)?;
Ok(SampleTensor { x, boxes, labels })
})
.collect::<Result<Vec<_>, AvError>>()?;
Ok(samples)
}
fn is_image_entry_name(name: &str) -> bool {
match name.rsplit_once('.') {
Some((_, ext)) => matches!(ext, "jpg" | "jpeg" | "png" | "bmp"),
None => false,
}
}
fn label_entry_name(img_name: &str) -> Option<String> {
let rest = img_name.strip_prefix("images/")?;
let (stem, _) = rest.rsplit_once('.')?;
Some(format!("labels/{stem}.txt"))
}
pub(crate) fn parse_yolo_label_text(
text: &str,
ow: u32,
oh: u32,
lbl_disp: &str,
) -> AvResult<(Vec<[f32; 4]>, Vec<u32>)> {
let bad = || AvError::data(format!("标注解析失败: {lbl_disp}"));
let mut boxes = Vec::new();
let mut labels = Vec::new();
for line in text.lines() {
let mut it = line.split_whitespace();
let (Some(cls), Some(cx), Some(cy), Some(w), Some(h)) =
(it.next(), it.next(), it.next(), it.next(), it.next())
else {
continue;
};
let (cls, cx, cy, w, h): (f32, f32, f32, f32, f32) = (
cls.parse().map_err(|_| bad())?,
cx.parse().map_err(|_| bad())?,
cy.parse().map_err(|_| bad())?,
w.parse().map_err(|_| bad())?,
h.parse().map_err(|_| bad())?,
);
if !(cls.is_finite() && cx.is_finite() && cy.is_finite() && w.is_finite() && h.is_finite())
{
return Err(AvError::data(format!("标注含 NaN/Inf 值: {lbl_disp}")));
}
if cls < 0.0 {
return Err(AvError::data(format!("标注类别为负: {lbl_disp}")));
}
let (bw, bh) = (w * ow as f32, h * oh as f32);
let (bcx, bcy) = (cx * ow as f32, cy * oh as f32);
boxes.push([
bcx - bw / 2.0,
bcy - bh / 2.0,
bcx + bw / 2.0,
bcy + bh / 2.0,
]);
labels.push(cls as u32);
}
Ok((boxes, labels))
}
pub fn load_yolo_avpack(
pack: &Path,
split: &str,
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Vec<SampleTensor>> {
Ok(
load_yolo_avpack_named(pack, split, img_size, device, imagenet_norm)?
.into_iter()
.map(|(_, s)| s)
.collect(),
)
}
pub fn load_yolo_avpack_named(
pack: &Path,
split: &str,
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Vec<(String, SampleTensor)>> {
let prefix = format!("images/{split}/");
let reader = crate::avpack::AvPackReader::open(pack)?;
let mut names: Vec<String> = reader
.entries()
.iter()
.filter(|e| e.name.starts_with(&prefix) && is_image_entry_name(&e.name))
.map(|e| e.name.clone())
.collect();
names.sort();
if names.is_empty() {
return Err(AvError::data(format!(
"avpack 容器 {} 无 images/{split}/ 下图片(布局须为 images/<split> + labels/<split>)",
pack.display()
)));
}
let samples: Vec<(String, SampleTensor)> = names
.par_iter()
.map(|name| -> AvResult<(String, SampleTensor)> {
let bytes = reader.bytes(name)?;
let rgb = image::load_from_memory(bytes)
.map_err(|e| AvError::data(format!("读图失败(容器条目 {name}): {e}")))?
.to_rgb8();
let (ow, oh) = (rgb.width(), rgb.height());
let lb = letterbox(ow, oh, img_size, img_size);
let (mut boxes, mut labels) = (Vec::new(), Vec::new());
if let Some(lbl_name) = label_entry_name(name) {
if reader.entry(&lbl_name).is_some() {
let text = std::str::from_utf8(reader.bytes(&lbl_name)?)
.map_err(|_| AvError::data(format!("标注非 UTF-8: {lbl_name}")))?;
let (px_boxes, cls) = parse_yolo_label_text(text, ow, oh, &lbl_name)?;
boxes = px_boxes
.iter()
.map(|b| {
let m = lb.map_box(Aabb::new(b[0], b[1], b[2], b[3]));
[m.x1, m.y1, m.x2, m.y2]
})
.collect();
labels = cls;
}
}
let x = rgb_to_input_tensor(&rgb, img_size, Some(lb), device, imagenet_norm)?;
Ok((name.clone(), SampleTensor { x, boxes, labels }))
})
.collect::<Result<Vec<_>, AvError>>()?;
Ok(samples)
}
pub fn load_yolo_avpack_raw(pack: &Path, split: &str) -> AvResult<Vec<RawDetectSample>> {
let prefix = format!("images/{split}/");
let reader = crate::avpack::AvPackReader::open(pack)?;
let mut names: Vec<String> = reader
.entries()
.iter()
.filter(|e| e.name.starts_with(&prefix) && is_image_entry_name(&e.name))
.map(|e| e.name.clone())
.collect();
names.sort();
if names.is_empty() {
return Err(AvError::data(format!(
"avpack 容器 {} 无 images/{split}/ 下图片",
pack.display()
)));
}
let mut out = Vec::new();
for name in &names {
let bytes = reader.bytes(name)?;
let rgb = image::load_from_memory(bytes)
.map_err(|e| AvError::data(format!("读图失败(容器条目 {name}): {e}")))?
.to_rgb8();
let (ow, oh) = (rgb.width(), rgb.height());
let (mut boxes, mut labels) = (Vec::new(), Vec::new());
if let Some(lbl_name) = label_entry_name(name) {
if reader.entry(&lbl_name).is_some() {
let text = std::str::from_utf8(reader.bytes(&lbl_name)?)
.map_err(|_| AvError::data(format!("标注非 UTF-8: {lbl_name}")))?;
(boxes, labels) = parse_yolo_label_text(text, ow, oh, &lbl_name)?;
}
}
out.push(RawDetectSample {
w: ow,
h: oh,
rgb: rgb.into_raw(),
boxes,
labels,
});
}
Ok(out)
}
pub fn decode_image_tensor(
path: &Path,
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Tensor> {
decode_image_tensor_with_mode(path, img_size, device, ResizeMode::Letterbox, imagenet_norm)
}
pub fn decode_image_tensor_with_mode(
path: &Path,
img_size: u32,
device: Device,
mode: ResizeMode,
imagenet_norm: bool,
) -> AvResult<Tensor> {
let (x, _, _, _) = decode_image_with_meta(path, img_size, device, mode, imagenet_norm)?;
Ok(x)
}
pub fn decode_image_with_meta(
path: &Path,
img_size: u32,
device: Device,
mode: ResizeMode,
imagenet_norm: bool,
) -> AvResult<(Tensor, Option<Letterbox>, u32, u32)> {
let img = image::open(path)
.map_err(|e| AvError::data(format!("读图失败 {}: {e}", path.display())))?;
let rgb = img.to_rgb8();
let (ow, oh) = (rgb.width(), rgb.height());
let lb = match mode {
ResizeMode::Letterbox => Some(letterbox(ow, oh, img_size, img_size)),
ResizeMode::Stretch => None,
};
let x = rgb_to_input_tensor(&rgb, img_size, lb, device, imagenet_norm)?;
Ok((x, lb, ow, oh))
}
pub fn decode_rgb_with_meta(
rgb: &image::RgbImage,
img_size: u32,
device: Device,
mode: ResizeMode,
imagenet_norm: bool,
) -> AvResult<(Tensor, Option<Letterbox>)> {
let lb = match mode {
ResizeMode::Letterbox => Some(letterbox(rgb.width(), rgb.height(), img_size, img_size)),
ResizeMode::Stretch => None,
};
let x = rgb_to_input_tensor(rgb, img_size, lb, device, imagenet_norm)?;
Ok((x, lb))
}
thread_local! {
static FIR_RESIZER: std::cell::RefCell<fast_image_resize::Resizer> =
std::cell::RefCell::new(fast_image_resize::Resizer::new());
}
fn resize_rgb8(rgb: &image::RgbImage, tw: u32, th: u32) -> AvResult<image::RgbImage> {
use fast_image_resize as fir;
let src = fir::images::ImageRef::new(
rgb.width(),
rgb.height(),
rgb.as_raw(),
fir::PixelType::U8x3,
)
.map_err(|_| AvError::data("RGB8 缓冲长度与宽高不符"))?;
let mut dst = fir::images::Image::new(tw, th, fir::PixelType::U8x3);
let opts = fir::ResizeOptions::new()
.resize_alg(fir::ResizeAlg::Convolution(fir::FilterType::Bilinear));
FIR_RESIZER.with(|r| {
r.borrow_mut()
.resize(&src, &mut dst, Some(&opts))
.map_err(|e| AvError::data(format!("RGB 缩放失败: {e}")))
})?;
image::RgbImage::from_raw(tw, th, dst.into_vec())
.ok_or_else(|| AvError::data("RGB 缩放输出长度异常"))
}
fn rgb_to_input_tensor(
rgb: &image::RgbImage,
img_size: u32,
lb: Option<Letterbox>,
device: Device,
imagenet_norm: bool,
) -> AvResult<Tensor> {
let canvas: image::RgbImage = match lb {
Some(lb) => {
let nw = ((rgb.width() as f32 * lb.scale).round() as u32).clamp(1, img_size);
let nh = ((rgb.height() as f32 * lb.scale).round() as u32).clamp(1, img_size);
let resized = if (nw, nh) == (rgb.width(), rgb.height()) {
rgb.clone()
} else {
resize_rgb8(rgb, nw, nh)?
};
let mut c =
image::RgbImage::from_pixel(img_size, img_size, image::Rgb([114, 114, 114]));
image::imageops::overlay(
&mut c,
&resized,
lb.pad_left.round() as i64,
lb.pad_top.round() as i64,
);
c
}
None => resize_rgb8(rgb, img_size, img_size)?,
};
canvas_to_input_tensor(canvas, device, imagenet_norm)
}
fn canvas_to_input_tensor(
canvas: image::RgbImage,
device: Device,
imagenet_norm: bool,
) -> AvResult<Tensor> {
let (w, h) = (canvas.width() as usize, canvas.height() as usize);
let n = w * h;
const IMAGENET_MEAN: [f32; 3] = [0.485, 0.456, 0.406];
const IMAGENET_STD: [f32; 3] = [0.229, 0.224, 0.225];
let mut buf = vec![0f32; 3 * n];
for (i, px) in canvas.pixels().enumerate() {
let [r, g, b] = px.0;
if imagenet_norm {
buf[i] = (r as f32 / 255.0 - IMAGENET_MEAN[0]) / IMAGENET_STD[0];
buf[n + i] = (g as f32 / 255.0 - IMAGENET_MEAN[1]) / IMAGENET_STD[1];
buf[2 * n + i] = (b as f32 / 255.0 - IMAGENET_MEAN[2]) / IMAGENET_STD[2];
} else {
buf[i] = r as f32 / 255.0;
buf[n + i] = g as f32 / 255.0;
buf[2 * n + i] = b as f32 / 255.0;
}
}
Ok(Tensor::from_slice(&buf)
.to_kind(Kind::Float)
.to_device(device)
.reshape([3, h as i64, w as i64]))
}
pub fn stack_samples(samples: &[SampleTensor]) -> AvResult<Tensor> {
let xs: Vec<&Tensor> = samples.iter().map(|s| &s.x).collect();
Ok(Tensor::stack(&xs, 0))
}
#[derive(Debug)]
pub struct ClassifySample {
pub x: Tensor,
}
impl Clone for ClassifySample {
fn clone(&self) -> Self {
Self { x: self.x.copy() }
}
}
pub fn stack_classify(samples: &[ClassifySample]) -> AvResult<Tensor> {
let xs: Vec<&Tensor> = samples.iter().map(|s| &s.x).collect();
Ok(Tensor::stack(&xs, 0))
}
pub type ImageFolderData = (Vec<ClassifySample>, Vec<u32>, HashMap<String, u32>);
pub fn load_imagefolder(
root: &Path,
split: &str,
img_size: u32,
num_classes_from_dir: bool,
device: Device,
imagenet_norm: bool,
) -> AvResult<ImageFolderData> {
let _ = num_classes_from_dir; load_imagefolder_with_classes(root, split, img_size, None, device, imagenet_norm)
}
pub fn load_imagefolder_with_classes(
root: &Path,
split: &str,
img_size: u32,
classes: Option<&HashMap<String, u32>>,
device: Device,
imagenet_norm: bool,
) -> AvResult<ImageFolderData> {
let split_dir = root.join(split);
if !split_dir.is_dir() {
return Err(AvError::data(format!(
"ImageFolder split 目录不存在: {}",
split_dir.display()
)));
}
let mut wnids: Vec<String> = std::fs::read_dir(&split_dir)?
.filter_map(|e| e.ok())
.filter(|e| e.path().is_dir())
.filter_map(|e| e.file_name().into_string().ok())
.collect();
wnids.sort();
let class_map: HashMap<String, u32> = match classes {
Some(m) => {
for w in &wnids {
if !m.contains_key(w) {
return Err(AvError::data(format!(
"split {split} 含未知类别 {w}(不在给定词表中)"
)));
}
}
m.clone()
}
None => wnids
.iter()
.enumerate()
.map(|(i, w)| (w.clone(), i as u32))
.collect(),
};
if wnids.is_empty() {
return Err(AvError::data(format!(
"ImageFolder split 目录为空(无 wnid 子目录): {}",
split_dir.display()
)));
}
let mut jobs: Vec<(PathBuf, u32)> = Vec::new();
for wnid in &wnids {
let label = class_map[wnid];
let cls_dir = split_dir.join(wnid);
let mut files: Vec<PathBuf> = std::fs::read_dir(&cls_dir)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| {
matches!(
p.extension().and_then(|e| e.to_str()),
Some("jpg") | Some("jpeg") | Some("JPEG") | Some("png") | Some("bmp")
)
})
.collect();
files.sort();
jobs.extend(files.into_iter().map(|p| (p, label)));
}
if jobs.is_empty() {
return Err(AvError::data(format!(
"ImageFolder split 无图片: {}",
split_dir.display()
)));
}
let decoded: Vec<(ClassifySample, u32)> = jobs
.par_iter()
.map(|(img_path, label)| -> AvResult<(ClassifySample, u32)> {
let img = image::open(img_path)
.map_err(|e| AvError::data(format!("读图失败 {}: {e}", img_path.display())))?;
let x = rgb_to_input_tensor(&img.to_rgb8(), img_size, None, device, imagenet_norm)?;
Ok((ClassifySample { x }, *label))
})
.collect::<Result<Vec<_>, AvError>>()?;
let (samples, labels): (Vec<ClassifySample>, Vec<u32>) = decoded.into_iter().unzip();
if samples.is_empty() {
return Err(AvError::data(format!(
"ImageFolder split 无图片: {}",
split_dir.display()
)));
}
Ok((samples, labels, class_map))
}
fn parse_row_f32(tokens: &[&str], lbl_disp: &str) -> AvResult<Vec<f32>> {
tokens
.iter()
.map(|t| {
let v: f32 = t
.parse()
.map_err(|_| AvError::data(format!("标注解析失败: {lbl_disp}")))?;
if v.is_finite() {
Ok(v)
} else {
Err(AvError::data(format!("标注含 NaN/Inf 值: {lbl_disp}")))
}
})
.collect()
}
fn label_class(v: f32, lbl_disp: &str) -> AvResult<u32> {
if v >= 0.0 {
Ok(v as u32)
} else {
Err(AvError::data(format!("标注类别为负: {lbl_disp}")))
}
}
pub struct ObbSample {
pub x: Tensor,
pub boxes: Vec<[f32; 5]>,
pub labels: Vec<u32>,
}
impl Clone for ObbSample {
fn clone(&self) -> Self {
Self {
x: self.x.copy(),
boxes: self.boxes.clone(),
labels: self.labels.clone(),
}
}
}
pub fn load_dota_dir(
root: &Path,
split: &str,
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Vec<ObbSample>> {
use av_core::conventions::AngleDomain;
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
if !img_dir.is_dir() {
return Err(AvError::data(format!(
"数据集图片目录不存在: {}",
img_dir.display()
)));
}
let mut paths: Vec<PathBuf> = std::fs::read_dir(&img_dir)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| {
matches!(
p.extension().and_then(|e| e.to_str()),
Some("jpg") | Some("jpeg") | Some("png") | Some("bmp")
)
})
.collect();
paths.sort();
if paths.is_empty() {
return Err(AvError::data(format!(
"数据集图片目录为空: {}",
img_dir.display()
)));
}
let mut out = Vec::new();
for p in paths {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let img = image::open(&p).map_err(|e| AvError::data(format!("读图失败 {p:?}: {e}")))?;
let rgb = img.to_rgb8();
let (ow, oh) = (rgb.width() as f32, rgb.height() as f32);
let (x, lb) =
decode_rgb_with_meta(&rgb, img_size, device, ResizeMode::Letterbox, imagenet_norm)?;
let lb = lb.ok_or_else(|| AvError::data("dota 加载要求 letterbox 模式"))?;
let mut boxes = Vec::new();
let mut labels = Vec::new();
if lbl_path.exists() {
let disp = lbl_path.display().to_string();
for line in std::fs::read_to_string(&lbl_path)?.lines() {
let tokens: Vec<&str> = line.split_whitespace().collect();
if tokens.len() < 9 {
continue;
}
let vals = parse_row_f32(&tokens[..9], &disp)?;
let mut pts = [[0f32; 2]; 4];
for k in 0..4 {
pts[k] = [
vals[1 + 2 * k] * ow * lb.scale + lb.pad_left,
vals[2 + 2 * k] * oh * lb.scale + lb.pad_top,
];
}
let cx = (pts[0][0] + pts[2][0]) / 2.0;
let cy = (pts[0][1] + pts[2][1]) / 2.0;
let dx = pts[1][0] - pts[0][0];
let dy = pts[1][1] - pts[0][1];
let w = (dx * dx + dy * dy).sqrt();
let h = {
let ex = pts[3][0] - pts[0][0];
let ey = pts[3][1] - pts[0][1];
(ex * ex + ey * ey).sqrt()
};
if w < 1.0 || h < 1.0 {
continue;
}
let theta = dy.atan2(dx);
boxes.push([cx, cy, w, h, AngleDomain::Le90.normalize(theta)]);
labels.push(label_class(vals[0], &disp)?);
}
}
out.push(ObbSample { x, boxes, labels });
}
Ok(out)
}
pub fn stack_obb_samples(samples: &[ObbSample]) -> AvResult<Tensor> {
let xs: Vec<&Tensor> = samples.iter().map(|s| &s.x).collect();
Ok(Tensor::stack(&xs, 0))
}
pub struct SegSample {
pub x: Tensor,
pub masks: Vec<Vec<u8>>,
pub labels: Vec<u32>,
}
impl Clone for SegSample {
fn clone(&self) -> Self {
Self {
x: self.x.copy(),
masks: self.masks.clone(),
labels: self.labels.clone(),
}
}
}
pub fn rasterize_polygon(points: &[[f32; 2]], w: usize, h: usize) -> Vec<u8> {
let mut mask = vec![0u8; w * h];
let n = points.len();
if n < 3 || w == 0 || h == 0 {
return mask;
}
let min_y = points.iter().map(|p| p[1]).fold(f32::INFINITY, f32::min);
let max_y = points
.iter()
.map(|p| p[1])
.fold(f32::NEG_INFINITY, f32::max);
for y in 0..h {
let cy = y as f32 + 0.5;
if cy < min_y || cy > max_y {
continue;
}
let mut xs: Vec<f32> = Vec::new();
let mut j = n - 1;
for i in 0..n {
let (p, q) = (points[i], points[j]);
if (p[1] <= cy && q[1] > cy) || (q[1] <= cy && p[1] > cy) {
let t = (cy - p[1]) / (q[1] - p[1]);
xs.push(p[0] + t * (q[0] - p[0]));
}
j = i;
}
xs.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
for pair in xs.chunks(2) {
if pair.len() < 2 {
continue;
}
let x0 = ((pair[0] - 0.5).ceil() as isize).max(0);
let x1 = ((pair[1] - 0.5).ceil() as isize).min(w as isize);
for x in x0..x1 {
mask[y * w + x as usize] = 1;
}
}
}
mask
}
pub fn load_cocoseg_dir(
root: &Path,
split: &str,
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Vec<SegSample>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
if !img_dir.is_dir() {
return Err(AvError::data(format!(
"数据集图片目录不存在: {}",
img_dir.display()
)));
}
let mut paths: Vec<PathBuf> = std::fs::read_dir(&img_dir)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| {
matches!(
p.extension().and_then(|e| e.to_str()),
Some("jpg") | Some("jpeg") | Some("png") | Some("bmp")
)
})
.collect();
paths.sort();
if paths.is_empty() {
return Err(AvError::data(format!(
"数据集图片目录为空: {}",
img_dir.display()
)));
}
let mw = (img_size / 4) as usize;
let mh = (img_size / 4) as usize;
let mut out = Vec::new();
for p in paths {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let img = image::open(&p).map_err(|e| AvError::data(format!("读图失败 {p:?}: {e}")))?;
let rgb = img.to_rgb8();
let (ow, oh) = (rgb.width() as f32, rgb.height() as f32);
let (x, lb) =
decode_rgb_with_meta(&rgb, img_size, device, ResizeMode::Letterbox, imagenet_norm)?;
let lb = lb.ok_or_else(|| AvError::data("coco seg 加载要求 letterbox 模式"))?;
let mut masks = Vec::new();
let mut labels = Vec::new();
if lbl_path.exists() {
let disp = lbl_path.display().to_string();
for line in std::fs::read_to_string(&lbl_path)?.lines() {
let tokens: Vec<&str> = line.split_whitespace().collect();
if tokens.len() < 7 {
continue;
}
if !(tokens.len() - 1).is_multiple_of(2) {
return Err(AvError::data(format!(
"分割多边形坐标数为奇数(标注错位): {disp}"
)));
}
let vals = parse_row_f32(&tokens, &disp)?;
let n_pts = (vals.len() - 1) / 2;
let k = mw as f32 / img_size as f32;
let pts: Vec<[f32; 2]> = (0..n_pts)
.map(|i| {
[
(vals[1 + 2 * i] * ow * lb.scale + lb.pad_left) * k,
(vals[2 + 2 * i] * oh * lb.scale + lb.pad_top) * k,
]
})
.collect();
let mask = rasterize_polygon(&pts, mw, mh);
if mask.iter().all(|&v| v == 0) {
continue; }
masks.push(mask);
labels.push(label_class(vals[0], &disp)?);
}
}
out.push(SegSample { x, masks, labels });
}
Ok(out)
}
pub fn stack_seg_samples(samples: &[SegSample]) -> AvResult<Tensor> {
let xs: Vec<&Tensor> = samples.iter().map(|s| &s.x).collect();
Ok(Tensor::stack(&xs, 0))
}
pub struct KeypointSample {
pub x: Tensor,
pub boxes: Vec<[f32; 4]>,
pub kpts: Vec<Vec<[f32; 3]>>,
pub labels: Vec<u32>,
}
impl Clone for KeypointSample {
fn clone(&self) -> Self {
Self {
x: self.x.copy(),
boxes: self.boxes.clone(),
kpts: self.kpts.clone(),
labels: self.labels.clone(),
}
}
}
pub fn load_cocopose_dir(
root: &Path,
split: &str,
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Vec<KeypointSample>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
if !img_dir.is_dir() {
return Err(AvError::data(format!(
"数据集图片目录不存在: {}",
img_dir.display()
)));
}
let mut paths: Vec<PathBuf> = std::fs::read_dir(&img_dir)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| {
matches!(
p.extension().and_then(|e| e.to_str()),
Some("jpg") | Some("jpeg") | Some("png") | Some("bmp")
)
})
.collect();
paths.sort();
if paths.is_empty() {
return Err(AvError::data(format!(
"数据集图片目录为空: {}",
img_dir.display()
)));
}
let mut out = Vec::new();
for p in paths {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let img = image::open(&p).map_err(|e| AvError::data(format!("读图失败 {p:?}: {e}")))?;
let rgb = img.to_rgb8();
let (ow, oh) = (rgb.width() as f32, rgb.height() as f32);
let (x, lb) =
decode_rgb_with_meta(&rgb, img_size, device, ResizeMode::Letterbox, imagenet_norm)?;
let lb = lb.ok_or_else(|| AvError::data("cocopose 加载要求 letterbox 模式"))?;
let mut boxes = Vec::new();
let mut kpts = Vec::new();
let mut labels = Vec::new();
if lbl_path.exists() {
let disp = lbl_path.display().to_string();
for line in std::fs::read_to_string(&lbl_path)?.lines() {
let tokens: Vec<&str> = line.split_whitespace().collect();
if tokens.len() < 8 {
continue;
}
if !(tokens.len() - 5).is_multiple_of(3) {
return Err(AvError::data(format!(
"姿态行关键点字段非 3 的倍数(标注错位): {disp}"
)));
}
let vals = parse_row_f32(&tokens, &disp)?;
let n_k = (vals.len() - 5) / 3;
let kp: Vec<[f32; 3]> = (0..n_k)
.map(|j| {
[
vals[5 + 3 * j] * ow * lb.scale + lb.pad_left,
vals[6 + 3 * j] * oh * lb.scale + lb.pad_top,
vals[7 + 3 * j],
]
})
.collect();
boxes.push([
vals[1] * ow * lb.scale + lb.pad_left,
vals[2] * oh * lb.scale + lb.pad_top,
(vals[3] * ow * lb.scale).max(1e-3),
(vals[4] * oh * lb.scale).max(1e-3),
]);
kpts.push(kp);
labels.push(label_class(vals[0], &disp)?);
}
}
out.push(KeypointSample {
x,
boxes,
kpts,
labels,
});
}
Ok(out)
}
pub fn stack_kp_samples(samples: &[KeypointSample]) -> AvResult<Tensor> {
let xs: Vec<&Tensor> = samples.iter().map(|s| &s.x).collect();
Ok(Tensor::stack(&xs, 0))
}
pub struct RawDetectSample {
pub w: u32,
pub h: u32,
pub rgb: Vec<u8>,
pub boxes: Vec<[f32; 4]>,
pub labels: Vec<u32>,
}
pub struct RawKeypointSample {
pub w: u32,
pub h: u32,
pub rgb: Vec<u8>,
pub boxes: Vec<[f32; 4]>,
pub kpts: Vec<Vec<[f32; 3]>>,
pub labels: Vec<u32>,
}
pub struct RawSegSample {
pub w: u32,
pub h: u32,
pub rgb: Vec<u8>,
pub polys: Vec<Vec<[f32; 2]>>,
pub labels: Vec<u32>,
}
pub struct RawObbSample {
pub w: u32,
pub h: u32,
pub rgb: Vec<u8>,
pub corners: Vec<[[f32; 2]; 4]>,
pub labels: Vec<u32>,
}
fn list_image_files(img_dir: &Path) -> AvResult<Vec<PathBuf>> {
if !img_dir.is_dir() {
return Err(AvError::data(format!(
"数据集图片目录不存在: {}",
img_dir.display()
)));
}
let mut paths: Vec<PathBuf> = std::fs::read_dir(img_dir)?
.filter_map(|e| e.ok().map(|e| e.path()))
.filter(|p| {
matches!(
p.extension().and_then(|e| e.to_str()),
Some("jpg") | Some("jpeg") | Some("png") | Some("bmp")
)
})
.collect();
paths.sort();
if paths.is_empty() {
return Err(AvError::data(format!(
"数据集图片目录为空: {}",
img_dir.display()
)));
}
Ok(paths)
}
fn decode_raw_rgb(p: &Path) -> AvResult<(u32, u32, Vec<u8>)> {
let rgb = image::open(p)
.map_err(|e| AvError::data(format!("读图失败 {}: {e}", p.display())))?
.to_rgb8();
let (w, h) = (rgb.width(), rgb.height());
Ok((w, h, rgb.into_raw()))
}
pub fn load_yolo_dir_raw(root: &Path, split: &str) -> AvResult<Vec<RawDetectSample>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
let mut out = Vec::new();
for p in list_image_files(&img_dir)? {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let (ow, oh, rgb) = decode_raw_rgb(&p)?;
let (mut boxes, mut labels) = (Vec::new(), Vec::new());
if lbl_path.exists() {
let text = std::fs::read_to_string(&lbl_path)?;
let (px_boxes, px_labels) =
parse_yolo_label_text(&text, ow, oh, &lbl_path.display().to_string())?;
boxes = px_boxes;
labels = px_labels;
}
out.push(RawDetectSample {
w: ow,
h: oh,
rgb,
boxes,
labels,
});
}
Ok(out)
}
fn detect_content_tile(raw: &RawDetectSample, img_size: u32) -> AvResult<RawDetectSample> {
let s = (img_size as f32 / raw.w as f32).min(img_size as f32 / raw.h as f32);
let tw = ((raw.w as f32 * s).round() as u32).max(1);
let th = ((raw.h as f32 * s).round() as u32).max(1);
let img = image::RgbImage::from_raw(raw.w, raw.h, raw.rgb.clone())
.ok_or_else(|| AvError::data("RGB8 缓冲长度与宽高不符"))?;
let rgb = resize_rgb8(&img, tw, th)?.into_raw();
let boxes = raw
.boxes
.iter()
.map(|b| [b[0] * s, b[1] * s, b[2] * s, b[3] * s])
.collect();
Ok(RawDetectSample {
w: tw,
h: th,
rgb,
boxes,
labels: raw.labels.clone(),
})
}
pub fn build_detect_cache_from_dir(
root: &Path,
split: &str,
img_size: u32,
) -> AvResult<Vec<RawDetectSample>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
let tiles: Vec<RawDetectSample> = list_image_files(&img_dir)?
.into_par_iter()
.map(|p| -> AvResult<RawDetectSample> {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let (ow, oh, rgb) = decode_raw_rgb(&p)?;
let (boxes, labels) = if lbl_path.exists() {
let text = std::fs::read_to_string(&lbl_path)?;
let (px_boxes, px_labels) =
parse_yolo_label_text(&text, ow, oh, &lbl_path.display().to_string())?;
(px_boxes, px_labels)
} else {
(Vec::new(), Vec::new())
};
detect_content_tile(
&RawDetectSample {
w: ow,
h: oh,
rgb,
boxes,
labels,
},
img_size,
)
})
.collect::<AvResult<Vec<_>>>()?;
if tiles.is_empty() {
return Err(AvError::data(format!(
"贴片缓存为空: {}",
img_dir.display()
)));
}
Ok(tiles)
}
pub fn build_detect_tile_cache(
raws: Vec<RawDetectSample>,
img_size: u32,
) -> AvResult<Vec<RawDetectSample>> {
raws.into_par_iter()
.map(|raw| detect_content_tile(&raw, img_size))
.collect()
}
pub fn load_cocopose_dir_raw(root: &Path, split: &str) -> AvResult<Vec<RawKeypointSample>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
let mut out = Vec::new();
for p in list_image_files(&img_dir)? {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let (ow, oh, rgb) = decode_raw_rgb(&p)?;
let (mut boxes, mut kpts, mut labels) = (Vec::new(), Vec::new(), Vec::new());
if lbl_path.exists() {
let disp = lbl_path.display().to_string();
for line in std::fs::read_to_string(&lbl_path)?.lines() {
let tokens: Vec<&str> = line.split_whitespace().collect();
if tokens.len() < 8 {
continue;
}
if !(tokens.len() - 5).is_multiple_of(3) {
return Err(AvError::data(format!(
"姿态行关键点字段非 3 的倍数(标注错位): {disp}"
)));
}
let vals = parse_row_f32(&tokens, &disp)?;
let n_k = (vals.len() - 5) / 3;
let kp: Vec<[f32; 3]> = (0..n_k)
.map(|j| {
[
vals[5 + 3 * j] * ow as f32,
vals[6 + 3 * j] * oh as f32,
vals[7 + 3 * j],
]
})
.collect();
boxes.push([
vals[1] * ow as f32,
vals[2] * oh as f32,
vals[3] * ow as f32,
vals[4] * oh as f32,
]);
kpts.push(kp);
labels.push(label_class(vals[0], &disp)?);
}
}
out.push(RawKeypointSample {
w: ow,
h: oh,
rgb,
boxes,
kpts,
labels,
});
}
Ok(out)
}
pub fn load_cocoseg_dir_raw(root: &Path, split: &str) -> AvResult<Vec<RawSegSample>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
list_image_files(&img_dir)?
.par_iter()
.map(|p| -> AvResult<RawSegSample> {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let (ow, oh, rgb) = decode_raw_rgb(p)?;
let (mut polys, mut labels) = (Vec::new(), Vec::new());
if lbl_path.exists() {
let disp = lbl_path.display().to_string();
for line in std::fs::read_to_string(&lbl_path)?.lines() {
let tokens: Vec<&str> = line.split_whitespace().collect();
if tokens.len() < 7 {
continue;
}
if !(tokens.len() - 1).is_multiple_of(2) {
return Err(AvError::data(format!(
"分割多边形坐标数为奇数(标注错位): {disp}"
)));
}
let vals = parse_row_f32(&tokens, &disp)?;
let n_pts = (vals.len() - 1) / 2;
polys.push(
(0..n_pts)
.map(|i| [vals[1 + 2 * i] * ow as f32, vals[2 + 2 * i] * oh as f32])
.collect(),
);
labels.push(label_class(vals[0], &disp)?);
}
}
Ok(RawSegSample {
w: ow,
h: oh,
rgb,
polys,
labels,
})
})
.collect()
}
pub fn load_dota_dir_raw(root: &Path, split: &str) -> AvResult<Vec<RawObbSample>> {
let img_dir = root.join("images").join(split);
let lbl_dir = root.join("labels").join(split);
let mut out = Vec::new();
for p in list_image_files(&img_dir)? {
let stem = p
.file_stem()
.and_then(|s| s.to_str())
.ok_or_else(|| AvError::data("文件名非法"))?
.to_string();
let lbl_path = lbl_dir.join(format!("{stem}.txt"));
let (ow, oh, rgb) = decode_raw_rgb(&p)?;
let (mut corners, mut labels) = (Vec::new(), Vec::new());
if lbl_path.exists() {
let disp = lbl_path.display().to_string();
for line in std::fs::read_to_string(&lbl_path)?.lines() {
let tokens: Vec<&str> = line.split_whitespace().collect();
if tokens.len() < 9 {
continue;
}
let vals = parse_row_f32(&tokens[..9], &disp)?;
let mut pts = [[0f32; 2]; 4];
for k in 0..4 {
pts[k] = [vals[1 + 2 * k] * ow as f32, vals[2 + 2 * k] * oh as f32];
}
corners.push(pts);
labels.push(label_class(vals[0], &disp)?);
}
}
out.push(RawObbSample {
w: ow,
h: oh,
rgb,
corners,
labels,
});
}
Ok(out)
}
fn augment_rgb_image(w: u32, h: u32, rgb: &[u8], plan: &AugmentPlan) -> AvResult<image::RgbImage> {
let mut buf = rgb.to_vec();
if plan.flip {
av_tasks::augment::hflip_rgb(w as usize, h as usize, &mut buf);
}
av_tasks::augment::mul_rgb(&mut buf, plan.rgb_gains);
let mut img = image::RgbImage::from_raw(w, h, buf)
.ok_or_else(|| AvError::data("RGB8 缓冲长度与宽高不符"))?;
if plan.scale != 1.0 {
let (aw, ah) = scaled_dims(w, h, plan.scale);
if (aw, ah) != (w, h) {
img = resize_rgb8(&img, aw, ah)?;
}
}
Ok(img)
}
pub fn mosaic4_raw(items: [&RawDetectSample; 4]) -> AvResult<RawDetectSample> {
let (qw, qh) = (items[0].w, items[0].h);
let mut resized: [Option<Vec<u8>>; 4] = [None, None, None, None];
for (slot, it) in resized.iter_mut().zip(items.iter()) {
if (it.w, it.h) != (qw, qh) {
let img = image::RgbImage::from_raw(it.w, it.h, it.rgb.clone())
.ok_or_else(|| AvError::data("RGB8 缓冲长度与宽高不符"))?;
*slot = Some(resize_rgb8(&img, qw, qh)?.into_raw());
}
}
let mitems: [av_tasks::augment::MosaicItem<'_>; 4] = std::array::from_fn(|k| {
let it = items[k];
av_tasks::augment::MosaicItem {
rgb: match &resized[k] {
Some(buf) => buf.as_slice(),
None => &it.rgb,
},
src_w: it.w,
src_h: it.h,
boxes: &it.boxes,
labels: &it.labels,
}
});
let (rgb, boxes, labels) = av_tasks::augment::mosaic_compose(qw, qh, &mitems);
let (w, h) = av_tasks::augment::mosaic_canvas_dims(qw, qh);
Ok(RawDetectSample {
w,
h,
rgb,
boxes,
labels,
})
}
pub fn mixup_raw(a: &RawDetectSample, b: &RawDetectSample, lam: f32) -> AvResult<RawDetectSample> {
let brgb: Vec<u8> = if (a.w, a.h) == (b.w, b.h) {
b.rgb.clone()
} else {
let img = image::RgbImage::from_raw(b.w, b.h, b.rgb.clone())
.ok_or_else(|| AvError::data("RGB8 缓冲长度与宽高不符"))?;
resize_rgb8(&img, a.w, a.h)?.into_raw()
};
Ok(RawDetectSample {
w: a.w,
h: a.h,
rgb: av_tasks::augment::mixup_rgb(&a.rgb, &brgb, lam),
boxes: a
.boxes
.iter()
.copied()
.chain(b.boxes.iter().copied())
.collect(),
labels: a
.labels
.iter()
.copied()
.chain(b.labels.iter().copied())
.collect(),
})
}
pub fn encode_detect_sample(
raw: &RawDetectSample,
img_size: u32,
device: Device,
mode: ResizeMode,
plan: &AugmentPlan,
imagenet_norm: bool,
) -> AvResult<SampleTensor> {
let mut boxes = raw.boxes.clone();
for b in boxes.iter_mut() {
if plan.flip {
*b = av_tasks::augment::flip_box_xyxy(*b, raw.w as f32);
}
if plan.scale != 1.0 {
for v in b.iter_mut() {
*v *= plan.scale;
}
}
}
let (aw, ah) = scaled_dims(raw.w, raw.h, plan.scale);
let lb = match mode {
ResizeMode::Letterbox => Some(letterbox(aw, ah, img_size, img_size)),
ResizeMode::Stretch => None,
};
let mapped: Vec<[f32; 4]> = boxes
.iter()
.map(|b| match lb {
Some(lb) => {
let m = lb.map_box(Aabb::new(b[0], b[1], b[2], b[3]));
[m.x1, m.y1, m.x2, m.y2]
}
None => {
let (sx, sy) = (img_size as f32 / aw as f32, img_size as f32 / ah as f32);
[b[0] * sx, b[1] * sy, b[2] * sx, b[3] * sy]
}
})
.collect();
let img = augment_rgb_image(raw.w, raw.h, &raw.rgb, plan)?;
let x = rgb_to_input_tensor(&img, img_size, lb, device, imagenet_norm)?;
Ok(SampleTensor {
x,
boxes: mapped,
labels: raw.labels.clone(),
})
}
pub fn encode_keypoint_sample(
raw: &RawKeypointSample,
img_size: u32,
device: Device,
plan: &AugmentPlan,
imagenet_norm: bool,
) -> AvResult<KeypointSample> {
let (fw, s) = (raw.w as f32, plan.scale);
let mut boxes = raw.boxes.clone();
for b in boxes.iter_mut() {
if plan.flip {
b[0] = fw - b[0];
}
b[0] *= s;
b[1] *= s;
b[2] *= s;
b[3] *= s;
}
let mut kpts = raw.kpts.clone();
for g in kpts.iter_mut() {
for p in g.iter_mut() {
if plan.flip {
p[0] = fw - p[0];
}
p[0] *= s;
p[1] *= s;
}
}
if plan.flip {
av_tasks::augment::swap_coco17_keypoints(&mut kpts);
}
let (aw, ah) = scaled_dims(raw.w, raw.h, plan.scale);
let lb = letterbox(aw, ah, img_size, img_size);
let boxes: Vec<[f32; 4]> = boxes
.iter()
.map(|b| {
[
b[0] * lb.scale + lb.pad_left,
b[1] * lb.scale + lb.pad_top,
(b[2] * lb.scale).max(1e-3),
(b[3] * lb.scale).max(1e-3),
]
})
.collect();
let kpts: Vec<Vec<[f32; 3]>> = kpts
.iter()
.map(|g| {
g.iter()
.map(|p| {
[
p[0] * lb.scale + lb.pad_left,
p[1] * lb.scale + lb.pad_top,
p[2],
]
})
.collect()
})
.collect();
let img = augment_rgb_image(raw.w, raw.h, &raw.rgb, plan)?;
let x = rgb_to_input_tensor(&img, img_size, Some(lb), device, imagenet_norm)?;
Ok(KeypointSample {
x,
boxes,
kpts,
labels: raw.labels.clone(),
})
}
pub fn encode_seg_sample(
raw: &RawSegSample,
img_size: u32,
device: Device,
plan: &AugmentPlan,
imagenet_norm: bool,
) -> AvResult<SegSample> {
let (mw, mh) = ((img_size / 4) as usize, (img_size / 4) as usize);
let k = mw as f32 / img_size as f32;
let (fw, s) = (raw.w as f32, plan.scale);
let (aw, ah) = scaled_dims(raw.w, raw.h, plan.scale);
let lb = letterbox(aw, ah, img_size, img_size);
let (mut masks, mut labels) = (Vec::new(), Vec::new());
for (poly, label) in raw.polys.iter().zip(&raw.labels) {
let pts: Vec<[f32; 2]> = poly
.iter()
.map(|p| {
let (mut x, y) = (p[0], p[1]);
if plan.flip {
x = fw - x;
}
[
(x * s * lb.scale + lb.pad_left) * k,
(y * s * lb.scale + lb.pad_top) * k,
]
})
.collect();
let mask = rasterize_polygon(&pts, mw, mh);
if mask.iter().all(|&v| v == 0) {
continue; }
masks.push(mask);
labels.push(*label);
}
let img = augment_rgb_image(raw.w, raw.h, &raw.rgb, plan)?;
let x = rgb_to_input_tensor(&img, img_size, Some(lb), device, imagenet_norm)?;
Ok(SegSample { x, masks, labels })
}
#[derive(Debug)]
pub struct CachedSegSample {
pub w: u32,
pub h: u32,
pub polys: Vec<Vec<[f32; 2]>>,
pub labels: Vec<u32>,
pub content: Vec<u8>,
pub cw: u32,
pub ch: u32,
}
pub fn build_seg_cache_sample(raw: &RawSegSample, img_size: u32) -> AvResult<CachedSegSample> {
let img = image::RgbImage::from_raw(raw.w, raw.h, raw.rgb.to_vec())
.ok_or_else(|| AvError::data("RGB8 缓冲长度与宽高不符"))?;
let lb = letterbox(raw.w, raw.h, img_size, 1); let nw = ((raw.w as f32 * lb.scale).round() as u32).clamp(1, img_size);
let nh = ((raw.h as f32 * lb.scale).round() as u32).clamp(1, img_size);
let resized = resize_rgb8(&img, nw, nh)?;
Ok(CachedSegSample {
w: raw.w,
h: raw.h,
polys: raw.polys.clone(),
labels: raw.labels.clone(),
content: resized.into_raw(),
cw: nw,
ch: nh,
})
}
pub fn build_seg_cache(
raw: &[RawSegSample],
img_size: u32,
) -> AvResult<(Vec<CachedSegSample>, u64)> {
let mut bytes = 0u64;
let out = raw
.par_iter()
.map(|r| -> AvResult<CachedSegSample> {
let c = build_seg_cache_sample(r, img_size)?;
Ok(c)
})
.collect::<AvResult<Vec<_>>>()?;
for c in &out {
bytes += c.content.len() as u64;
}
Ok((out, bytes))
}
pub fn encode_seg_sample_cached(
c: &CachedSegSample,
img_size: u32,
device: Device,
plan: &AugmentPlan,
imagenet_norm: bool,
) -> AvResult<SegSample> {
let (masks, labels) = seg_masks_from_plan(c, img_size, plan);
let mut buf = c.content.clone();
if plan.flip {
av_tasks::augment::hflip_rgb(c.cw as usize, c.ch as usize, &mut buf);
}
av_tasks::augment::mul_rgb(&mut buf, plan.rgb_gains);
let (nw_t, nh_t, lb) = scaled_content_target(c, img_size, plan.scale);
let content = if plan.scale != 1.0 {
let img = image::RgbImage::from_raw(c.cw, c.ch, buf)
.ok_or_else(|| AvError::data("缓存贴片长度与宽高不符"))?;
resize_rgb8(&img, nw_t, nh_t)?
} else {
image::RgbImage::from_raw(c.cw, c.ch, buf)
.ok_or_else(|| AvError::data("缓存贴片长度与宽高不符"))?
};
let mut canvas = image::RgbImage::from_pixel(img_size, img_size, image::Rgb([114, 114, 114]));
image::imageops::overlay(
&mut canvas,
&content,
lb.pad_left.round() as i64,
lb.pad_top.round() as i64,
);
let x = canvas_to_input_tensor(canvas, device, imagenet_norm)?;
Ok(SegSample { x, masks, labels })
}
fn scaled_content_target(c: &CachedSegSample, img_size: u32, scale: f32) -> (u32, u32, Letterbox) {
let (aw, ah) = scaled_dims(c.w, c.h, scale);
let lb = letterbox(aw, ah, img_size, img_size);
let nw = ((aw as f32 * lb.scale).round() as u32).clamp(1, img_size);
let nh = ((ah as f32 * lb.scale).round() as u32).clamp(1, img_size);
(nw, nh, lb)
}
fn seg_masks_from_plan(
c: &CachedSegSample,
img_size: u32,
plan: &AugmentPlan,
) -> (Vec<Vec<u8>>, Vec<u32>) {
let (mw, mh) = ((img_size / 4) as usize, (img_size / 4) as usize);
let k = mw as f32 / img_size as f32;
let (fw, s) = (c.w as f32, plan.scale);
let (aw, ah) = scaled_dims(c.w, c.h, plan.scale);
let lb = letterbox(aw, ah, img_size, img_size);
let (mut masks, mut labels) = (Vec::new(), Vec::new());
for (poly, label) in c.polys.iter().zip(&c.labels) {
let pts: Vec<[f32; 2]> = poly
.iter()
.map(|p| {
let (mut x, y) = (p[0], p[1]);
if plan.flip {
x = fw - x;
}
[
(x * s * lb.scale + lb.pad_left) * k,
(y * s * lb.scale + lb.pad_top) * k,
]
})
.collect();
let mask = rasterize_polygon(&pts, mw, mh);
if mask.iter().all(|&v| v == 0) {
continue; }
masks.push(mask);
labels.push(*label);
}
(masks, labels)
}
pub fn encode_seg_batch_cached(
cache: &[CachedSegSample],
idx: &[usize],
plans: &[AugmentPlan],
img_size: u32,
device: Device,
imagenet_norm: bool,
) -> AvResult<Vec<SegSample>> {
idx.par_iter()
.zip(plans)
.map(|(&i, plan)| {
encode_seg_sample_cached(&cache[i], img_size, device, plan, imagenet_norm)
})
.collect()
}
pub fn build_seg_canvas_stack(
cache: &[CachedSegSample],
img_size: u32,
device: Device,
) -> AvResult<Tensor> {
let s = img_size as i64;
let gray = 114.0f32 / 255.0;
let mut bufs: Vec<Tensor> = Vec::with_capacity(cache.len());
for c in cache {
let lb = letterbox(c.w, c.h, img_size, img_size);
let (pl, pt) = (lb.pad_left.round() as i64, lb.pad_top.round() as i64);
let img = image::RgbImage::from_raw(c.cw, c.ch, c.content.clone())
.ok_or_else(|| AvError::data("缓存贴片长度与宽高不符"))?;
let content = canvas_to_input_tensor(img, device, false)?; let canvas = Tensor::full([3, s, s], gray as f64, (Kind::Float, device));
canvas
.narrow(2, pl, c.cw as i64)
.narrow(1, pt, c.ch as i64)
.copy_(&content);
bufs.push(canvas);
}
Ok(Tensor::stack(&bufs, 0))
}
pub fn encode_seg_sample_gpu(
stack: &Tensor,
i: u32,
c: &CachedSegSample,
img_size: u32,
plan: &AugmentPlan,
imagenet_norm: bool,
) -> AvResult<SegSample> {
let (masks, labels) = seg_masks_from_plan(c, img_size, plan);
let lb1 = letterbox(c.w, c.h, img_size, img_size);
let (pl1, pt1) = (lb1.pad_left.round() as i64, lb1.pad_top.round() as i64);
let mut x = stack.select(0, i as i64).copy(); if plan.flip {
let region = x.narrow(2, pl1, c.cw as i64).copy();
x.narrow(2, pl1, c.cw as i64).copy_(®ion.flip([2i64]));
}
if plan.rgb_gains != [1.0; 3] {
let gains = Tensor::from_slice(&plan.rgb_gains)
.to_kind(Kind::Float)
.to_device(x.device())
.reshape([3i64, 1, 1]);
let region = x.narrow(2, pl1, c.cw as i64).copy() * gains;
x.narrow(2, pl1, c.cw as i64).copy_(®ion.clamp(0.0, 1.0));
}
if plan.scale != 1.0 {
let (nw, nh, lb) = scaled_content_target(c, img_size, plan.scale);
let region = x
.narrow(2, pl1, c.cw as i64)
.narrow(1, pt1, c.ch as i64)
.copy();
let resized = region
.unsqueeze(0)
.upsample_bilinear2d([nh as i64, nw as i64], false, None, None)
.squeeze_dim(0);
let gray = 114.0f32 / 255.0;
let s = img_size as i64;
let canvas = Tensor::full([3, s, s], gray as f64, (Kind::Float, x.device()));
canvas
.narrow(2, lb.pad_left.round() as i64, nw as i64)
.narrow(1, lb.pad_top.round() as i64, nh as i64)
.copy_(&resized);
x = canvas;
}
if imagenet_norm {
const MEAN: [f32; 3] = [0.485, 0.456, 0.406];
const STD: [f32; 3] = [0.229, 0.224, 0.225];
let mean = Tensor::from_slice(&MEAN)
.to_kind(Kind::Float)
.to_device(x.device())
.reshape([3i64, 1, 1]);
let std = Tensor::from_slice(&STD)
.to_kind(Kind::Float)
.to_device(x.device())
.reshape([3i64, 1, 1]);
x = (x - mean) / std;
}
Ok(SegSample { x, masks, labels })
}
pub fn encode_obb_sample(
raw: &RawObbSample,
img_size: u32,
device: Device,
plan: &AugmentPlan,
imagenet_norm: bool,
) -> AvResult<ObbSample> {
use av_core::conventions::AngleDomain;
let (fw, s) = (raw.w as f32, plan.scale);
let (aw, ah) = scaled_dims(raw.w, raw.h, plan.scale);
let lb = letterbox(aw, ah, img_size, img_size);
let (mut boxes, mut labels) = (Vec::new(), Vec::new());
for (corners, label) in raw.corners.iter().zip(&raw.labels) {
let mut pts = [[0f32; 2]; 4];
for (dst, src) in pts.iter_mut().zip(corners) {
let (mut x, y) = (src[0], src[1]);
if plan.flip {
x = fw - x;
}
*dst = [
x * s * lb.scale + lb.pad_left,
y * s * lb.scale + lb.pad_top,
];
}
let cx = (pts[0][0] + pts[2][0]) / 2.0;
let cy = (pts[0][1] + pts[2][1]) / 2.0;
let dx = pts[1][0] - pts[0][0];
let dy = pts[1][1] - pts[0][1];
let w = (dx * dx + dy * dy).sqrt();
let h = {
let ex = pts[3][0] - pts[0][0];
let ey = pts[3][1] - pts[0][1];
(ex * ex + ey * ey).sqrt()
};
if w < 1.0 || h < 1.0 {
continue;
}
let theta = dy.atan2(dx);
boxes.push([cx, cy, w, h, AngleDomain::Le90.normalize(theta)]);
labels.push(*label);
}
let img = augment_rgb_image(raw.w, raw.h, &raw.rgb, plan)?;
let x = rgb_to_input_tensor(&img, img_size, Some(lb), device, imagenet_norm)?;
Ok(ObbSample { x, boxes, labels })
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn letterbox_label_mapping_roundtrip() {
let lb = letterbox(640, 480, 320, 320);
assert_eq!((lb.dst_w, lb.dst_h), (320, 320));
assert!((lb.scale - 0.5).abs() < 1e-6);
let g = Aabb::new(10.0, 20.0, 100.0, 80.0);
let m = lb.map_box(g);
assert!((m.x1 - 5.0).abs() < 1e-4, "x1={}", m.x1);
assert!((m.y1 - 50.0).abs() < 1e-4, "y1={}", m.y1); assert!((m.x2 - 50.0).abs() < 1e-4, "x2={}", m.x2);
assert!((m.y2 - 80.0).abs() < 1e-4, "y2={}", m.y2);
let back = lb.restore_box(m, 640, 480);
assert!((back.x1 - g.x1).abs() <= 1.0, "back={back:?}");
assert!((back.y1 - g.y1).abs() <= 1.0, "back={back:?}");
assert!((back.x2 - g.x2).abs() <= 1.0, "back={back:?}");
assert!((back.y2 - g.y2).abs() <= 1.0, "back={back:?}");
let lb2 = letterbox(427, 640, 320, 320);
assert_eq!((lb2.dst_w, lb2.dst_h), (320, 320));
let g2 = Aabb::new(3.0, 7.0, 424.0, 633.0);
let back2 = lb2.restore_box(lb2.map_box(g2), 427, 640);
assert!((back2.x1 - g2.x1).abs() <= 1.0, "back2={back2:?}");
assert!((back2.x2 - g2.x2).abs() <= 1.0, "back2={back2:?}");
assert!((back2.y2 - g2.y2).abs() <= 1.0, "back2={back2:?}");
}
#[test]
fn letterbox_decode_pads_gray_keeps_content() {
let dir = std::env::temp_dir().join(format!("av-ds-lb-test-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let p = dir.join("red.png");
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(&p)
.unwrap();
let x = decode_image_tensor_with_mode(&p, 32, Device::Cpu, ResizeMode::Letterbox, false)
.unwrap();
let px = |c: usize, y: usize, xx: usize| x.double_value(&[c as i64, y as i64, xx as i64]);
for c in 0..3 {
assert!(
(px(c, 0, 0) - 114.0 / 255.0).abs() < 1e-6,
"pad c{c}={}",
px(c, 0, 0)
);
}
assert!((px(0, 16, 16) - 1.0).abs() < 1e-6, "r={}", px(0, 16, 16));
assert!(px(1, 16, 16) < 1e-6, "g={}", px(1, 16, 16));
assert!(px(2, 16, 16) < 1e-6, "b={}", px(2, 16, 16));
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn stretch_mode_still_available() {
let dir = std::env::temp_dir().join(format!("av-ds-st-test-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let p = dir.join("red.png");
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(&p)
.unwrap();
let x =
decode_image_tensor_with_mode(&p, 32, Device::Cpu, ResizeMode::Stretch, false).unwrap();
let px = |c: usize, y: usize, xx: usize| x.double_value(&[c as i64, y as i64, xx as i64]);
assert!((px(0, 0, 0) - 1.0).abs() < 1e-6);
assert!(px(2, 31, 31) < 1e-6);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn imagenet_norm_maps_red_pixel_to_imagenet_domain() {
let dir = std::env::temp_dir().join(format!("av-ds-norm-test-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let p = dir.join("red.png");
image::RgbImage::from_pixel(8, 8, image::Rgb([255, 0, 0]))
.save(&p)
.unwrap();
let x01 =
decode_image_tensor_with_mode(&p, 8, Device::Cpu, ResizeMode::Stretch, false).unwrap();
assert!((x01.double_value(&[0, 0, 0]) - 1.0).abs() < 1e-6);
assert!(x01.double_value(&[1, 0, 0]).abs() < 1e-6);
let xn =
decode_image_tensor_with_mode(&p, 8, Device::Cpu, ResizeMode::Stretch, true).unwrap();
let expect = [
(1.0f32 - 0.485) / 0.229,
(0.0f32 - 0.456) / 0.224,
(0.0f32 - 0.406) / 0.225,
];
for (c, &e) in expect.iter().enumerate() {
assert!(
(xn.double_value(&[c as i64, 4, 4]) - e as f64).abs() < 1e-6,
"c{c}={}",
xn.double_value(&[c as i64, 4, 4])
);
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_yolo_dir_letterbox_maps_boxes_to_canvas() {
let dir = std::env::temp_dir().join(format!("av-ds-dir-test-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(lbl_dir.join("a.txt"), "0 0.5 0.5 1.0 1.0\n").unwrap();
let samples =
load_yolo_dir_with_mode(&dir, "train", 32, Device::Cpu, ResizeMode::Letterbox, false)
.unwrap();
assert_eq!(samples.len(), 1);
let b = samples[0].boxes[0];
assert!((b[0] - 0.0).abs() <= 1.0, "x1={}", b[0]);
assert!((b[1] - 8.0).abs() <= 1.0, "y1={}", b[1]);
assert!((b[2] - 32.0).abs() <= 1.0, "x2={}", b[2]);
assert!((b[3] - 24.0).abs() <= 1.0, "y2={}", b[3]);
assert_eq!(samples[0].labels, vec![0]);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_yolo_avpack_roundtrip_names_and_boxes() {
let dir = std::env::temp_dir().join(format!("av-ds-avpack-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
let (vimg, vlbl) = (dir.join("images/val"), dir.join("labels/val"));
for d in [&img_dir, &lbl_dir, &vimg, &vlbl] {
std::fs::create_dir_all(d).unwrap();
}
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(lbl_dir.join("a.txt"), "0 0.5 0.5 1.0 1.0\n").unwrap();
image::RgbImage::from_pixel(32, 32, image::Rgb([0, 0, 255]))
.save(img_dir.join("b.png"))
.unwrap();
std::fs::write(lbl_dir.join("b.txt"), "3 0.25 0.5 0.5 0.5\n").unwrap();
image::RgbImage::from_pixel(16, 16, image::Rgb([0, 255, 0]))
.save(vimg.join("c.png"))
.unwrap();
std::fs::write(vlbl.join("c.txt"), "1 0.5 0.5 0.5 0.5\n").unwrap();
let out = dir.join("ds.avpack");
let (count, _) = crate::avpack::pack_dir(&dir, &out).unwrap();
assert_eq!(count, 6, "a/b 图 + a/b 标注 + c 图 + c 标注");
let named = load_yolo_avpack_named(&out, "train", 32, Device::Cpu, false).unwrap();
assert_eq!(
named.iter().map(|(n, _)| n.as_str()).collect::<Vec<_>>(),
vec!["images/train/a.png", "images/train/b.png"],
"train split 只含 train 图,按名字排序"
);
let a = &named[0].1;
assert_eq!(a.labels, vec![0]);
assert!((a.boxes[0][0] - 0.0).abs() <= 1.0, "a x1={}", a.boxes[0][0]);
assert!((a.boxes[0][1] - 8.0).abs() <= 1.0, "a y1={}", a.boxes[0][1]);
assert!(
(a.boxes[0][2] - 32.0).abs() <= 1.0,
"a x2={}",
a.boxes[0][2]
);
assert!(
(a.boxes[0][3] - 24.0).abs() <= 1.0,
"a y2={}",
a.boxes[0][3]
);
let b = &named[1].1;
assert_eq!(b.labels, vec![3]);
for (got, want) in b.boxes[0].iter().zip([0.0f32, 8.0, 16.0, 24.0]) {
assert!((got - want).abs() < 1e-4, "b box {got} vs {want}");
}
let samples: Vec<SampleTensor> = named.into_iter().map(|(_, s)| s).collect();
assert_eq!(stack_samples(&samples).unwrap().size(), vec![2, 3, 32, 32]);
let val = load_yolo_avpack(&out, "val", 32, Device::Cpu, false).unwrap();
assert_eq!(val.len(), 1);
assert_eq!(val[0].labels, vec![1]);
let raw = load_yolo_avpack_raw(&out, "train").unwrap();
assert_eq!(raw.len(), 2);
assert_eq!(raw[0].boxes, vec![[0.0, 0.0, 64.0, 32.0]]);
assert_eq!(raw[1].boxes, vec![[0.0, 8.0, 16.0, 24.0]]);
let _ = std::fs::remove_dir_all(&dir);
let _ = std::fs::remove_file(&out);
}
#[test]
fn load_imagefolder_derives_sorted_class_ids() {
let dir = std::env::temp_dir().join(format!("av-ds-if-test-{}", std::process::id()));
let train = dir.join("train");
std::fs::create_dir_all(train.join("n02102040")).unwrap(); std::fs::create_dir_all(train.join("n01440764")).unwrap(); image::RgbImage::from_pixel(16, 8, image::Rgb([255, 0, 0]))
.save(train.join("n01440764").join("a.png"))
.unwrap();
image::RgbImage::from_pixel(8, 16, image::Rgb([0, 0, 255]))
.save(train.join("n02102040").join("b.png"))
.unwrap();
let (samples, labels, map) = load_imagefolder(&dir, "train", 32, true, Device::Cpu, false)
.expect("ImageFolder 应可加载");
assert_eq!(samples.len(), 2);
assert_eq!(labels, vec![0, 1]);
assert_eq!(map.len(), 2);
assert_eq!(map["n01440764"], 0);
assert_eq!(map["n02102040"], 1);
let px = |s: &ClassifySample, c: usize, y: usize, x: usize| {
s.x.double_value(&[c as i64, y as i64, x as i64])
};
assert!((px(&samples[0], 0, 0, 0) - 1.0).abs() < 1e-6, "应为红");
assert!((px(&samples[0], 1, 31, 31) - 0.0).abs() < 1e-6);
assert!((px(&samples[1], 2, 16, 16) - 1.0).abs() < 1e-6, "应为蓝");
assert!((px(&samples[1], 0, 0, 0) - 0.0).abs() < 1e-6);
let batch = stack_classify(&samples).unwrap();
assert_eq!(batch.size(), vec![2, 3, 32, 32]);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_imagefolder_reuses_train_class_map() {
let dir = std::env::temp_dir().join(format!("av-ds-ifmap-test-{}", std::process::id()));
let (tr, va) = (dir.join("train"), dir.join("val"));
for base in [&tr, &va] {
std::fs::create_dir_all(base.join("n01440764")).unwrap();
std::fs::create_dir_all(base.join("n02102040")).unwrap();
}
image::RgbImage::from_pixel(8, 8, image::Rgb([255, 0, 0]))
.save(tr.join("n01440764").join("a.png"))
.unwrap();
image::RgbImage::from_pixel(8, 8, image::Rgb([0, 255, 0]))
.save(tr.join("n02102040").join("b.png"))
.unwrap();
image::RgbImage::from_pixel(8, 8, image::Rgb([0, 0, 255]))
.save(va.join("n02102040").join("c.png"))
.unwrap();
let (_, _, train_map) =
load_imagefolder(&dir, "train", 16, true, Device::Cpu, false).unwrap();
let (_, val_labels, _) =
load_imagefolder_with_classes(&dir, "val", 16, Some(&train_map), Device::Cpu, false)
.unwrap();
assert_eq!(val_labels, vec![1], "val 复用 train 映射(n02102040 → 1)");
let mut partial = HashMap::new();
partial.insert("n01440764".to_string(), 0u32);
let err =
load_imagefolder_with_classes(&dir, "val", 16, Some(&partial), Device::Cpu, false)
.unwrap_err();
assert!(err.to_string().contains("n02102040"), "got: {err}");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_imagefolder_missing_split_errors() {
let dir = std::env::temp_dir().join(format!("av-ds-ifmiss-test-{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let err = load_imagefolder(&dir, "train", 16, true, Device::Cpu, false).unwrap_err();
assert!(err.to_string().contains("不存在"), "got: {err}");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn rasterize_polygon_rectangle_hand_computed() {
let pts = [[2.0, 2.0], [10.0, 2.0], [10.0, 8.0], [2.0, 8.0]];
let m = rasterize_polygon(&pts, 16, 16);
assert_eq!(m.iter().filter(|&&v| v == 1).count(), 48);
assert_eq!(m[2 * 16 + 2], 1, "(2,2) 应覆盖");
assert_eq!(m[7 * 16 + 9], 1, "(9,7) 应覆盖");
assert_eq!(m[2 * 16 + 10], 0, "(10,2) 中心 x=10.5 在界外");
assert_eq!(m[8 * 16 + 2], 0, "(2,8) 中心 y=8.5 在界外");
assert_eq!(m[16 + 2], 0);
assert_eq!(m[2 * 16 + 1], 0);
}
#[test]
fn rasterize_polygon_triangle_half_open() {
let pts = [[0.0, 0.0], [4.0, 0.0], [0.0, 4.0]];
let m = rasterize_polygon(&pts, 8, 8);
let cnt = m.iter().filter(|&&v| v == 1).count();
assert_eq!(cnt, 6);
assert_eq!(m[0], 1);
assert_eq!(m[2], 1); assert_eq!(m[8 + 1], 1); assert_eq!(m[2 * 8], 1); assert_eq!(m[3 * 8], 0); }
#[test]
fn rasterize_polygon_degenerate_inputs() {
assert!(rasterize_polygon(&[], 4, 4).iter().all(|&v| v == 0));
assert!(rasterize_polygon(&[[1.0, 1.0], [3.0, 3.0]], 4, 4)
.iter()
.all(|&v| v == 0));
let out = [
[100.0, 100.0],
[120.0, 100.0],
[120.0, 120.0],
[100.0, 120.0],
];
assert!(rasterize_polygon(&out, 4, 4).iter().all(|&v| v == 0));
let half = [[-4.0, -4.0], [4.0, -4.0], [4.0, 4.0], [-4.0, 4.0]];
let m = rasterize_polygon(&half, 4, 4);
assert_eq!(m.iter().filter(|&&v| v == 1).count(), 16, "整画布被覆盖");
}
#[test]
fn load_cocoseg_dir_rasterizes_and_skips_box_lines() {
let dir = std::env::temp_dir().join(format!("av-ds-cseg-test-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(
lbl_dir.join("a.txt"),
concat!(
"7 0.25 0.25 0.75 0.25 0.75 0.75 0.25 0.75\n",
"3 0.5 0.5 0.5 0.5\n",
"9 0.1 0.1 0.2 0.2 0.3 0.3 0.1 0.2 0.9 0.9 0.1 0.9\n",
),
)
.unwrap();
let samples = load_cocoseg_dir(&dir, "train", 32, Device::Cpu, false).unwrap();
assert_eq!(samples.len(), 1);
let s = &samples[0];
assert_eq!(s.labels, vec![7, 9]);
assert_eq!(s.masks.len(), 2);
let m0 = &s.masks[0];
assert_eq!(m0.len(), 8 * 8);
let cnt0 = m0.iter().filter(|&&v| v == 1).count();
assert_eq!(cnt0, 8, "掩码 (2,3)-(6,5) 应覆盖 4×2=8 像素,实际 {cnt0}");
assert_eq!(m0[3 * 8 + 3], 1, "中心点应覆盖");
assert_eq!(m0[0], 0, "画布角落(灰边)不应覆盖");
assert!(s.masks[1].contains(&1));
let x = stack_seg_samples(&samples).unwrap();
assert_eq!(x.size(), vec![1, 3, 32, 32]);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_cocopose_dir_maps_kpts_and_boxes_to_canvas() {
let dir = std::env::temp_dir().join(format!("av-ds-kp-test-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
image::RgbImage::from_pixel(64, 32, image::Rgb([0, 255, 0]))
.save(img_dir.join("b.png"))
.unwrap();
std::fs::write(
lbl_dir.join("a.txt"),
"0 0.5 0.5 0.5 0.5 0.25 0.25 2.0 0.5 0.5 0.0\n",
)
.unwrap();
std::fs::write(
lbl_dir.join("b.txt"),
"0 0.5 0.5 1.0 1.0 0.1 0.2 2.0 0.3 0.4 1.0 0.5 0.6 2.0\n",
)
.unwrap();
let samples = load_cocopose_dir(&dir, "train", 32, Device::Cpu, false).unwrap();
assert_eq!(samples.len(), 2);
let a = &samples[0];
assert_eq!(a.labels, vec![0]);
assert_eq!(a.boxes.len(), 1);
assert!((a.boxes[0][0] - 16.0).abs() < 1e-4, "cx={}", a.boxes[0][0]);
assert!((a.boxes[0][1] - 16.0).abs() < 1e-4, "cy={}", a.boxes[0][1]);
assert!((a.boxes[0][2] - 16.0).abs() < 1e-4, "w={}", a.boxes[0][2]);
assert!((a.boxes[0][3] - 8.0).abs() < 1e-4, "h={}", a.boxes[0][3]);
assert_eq!(a.kpts[0].len(), 2);
assert!(
(a.kpts[0][0][0] - 8.0).abs() < 1e-4,
"kx={}",
a.kpts[0][0][0]
);
assert!(
(a.kpts[0][0][1] - 12.0).abs() < 1e-4,
"ky={}",
a.kpts[0][0][1]
);
assert_eq!(a.kpts[0][0][2], 2.0);
assert!((a.kpts[0][1][0] - 16.0).abs() < 1e-4);
assert!((a.kpts[0][1][1] - 16.0).abs() < 1e-4);
assert_eq!(a.kpts[0][1][2], 0.0, "v=0 应原样保留");
let b = &samples[1];
assert_eq!(b.kpts[0].len(), 3, "K 应由行内 token 数推导");
assert!((b.kpts[0][0][0] - 3.2).abs() < 1e-4);
assert!((b.kpts[0][0][1] - 11.2).abs() < 1e-4);
assert_eq!(b.kpts[0][1][2], 1.0, "v=1(遮挡)应保留");
assert!((b.boxes[0][3] - 16.0).abs() < 1e-4);
let x = stack_kp_samples(&samples).unwrap();
assert_eq!(x.size(), vec![2, 3, 32, 32]);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn load_cocopose_dir_skips_malformed_lines() {
let dir = std::env::temp_dir().join(format!("av-ds-kpbad-test-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
image::RgbImage::from_pixel(32, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(
lbl_dir.join("a.txt"),
concat!(
"0 0.5 0.5 0.5 0.5 0.1 0.1\n", "\n", "0 0.5 0.5 0.5 0.5 0.2 0.2 2.0\n", ),
)
.unwrap();
let samples = load_cocopose_dir(&dir, "train", 32, Device::Cpu, false).unwrap();
assert_eq!(samples.len(), 1);
assert_eq!(samples[0].kpts.len(), 1);
assert_eq!(samples[0].kpts[0].len(), 1);
assert!((samples[0].kpts[0][0][0] - 6.4).abs() < 1e-4);
let _ = std::fs::remove_dir_all(&dir);
}
fn tensor_max_diff(a: &Tensor, b: &Tensor) -> f64 {
(a - b).abs().max().double_value(&[])
}
#[test]
fn encode_keypoint_none_matches_plain_loader() {
let dir = std::env::temp_dir().join(format!("av-ds-augkpn-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(
lbl_dir.join("a.txt"),
"0 0.5 0.5 0.5 0.5 0.25 0.25 2.0 0.5 0.5 0.0\n",
)
.unwrap();
let plain = load_cocopose_dir(&dir, "train", 32, Device::Cpu, false).unwrap();
let raw = load_cocopose_dir_raw(&dir, "train").unwrap();
assert_eq!(raw.len(), 1);
let enc =
encode_keypoint_sample(&raw[0], 32, Device::Cpu, &AugmentPlan::none(), false).unwrap();
assert_eq!(
tensor_max_diff(&plain[0].x, &enc.x),
0.0,
"none() 张量应逐位一致"
);
assert_eq!(plain[0].boxes, enc.boxes, "none() 框应逐位一致");
assert_eq!(plain[0].kpts, enc.kpts, "none() 关键点应逐位一致");
assert_eq!(plain[0].labels, enc.labels);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn encode_keypoint_flip_mirrors_and_swaps_coco17() {
use av_tasks::augment::COCO17_FLIP_SWAP;
let dir = std::env::temp_dir().join(format!("av-ds-augkpf-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
let mut img = image::RgbImage::new(32, 32);
for y in 0..32 {
for x in 0..32 {
img.put_pixel(
x,
y,
if x < 16 {
image::Rgb([255, 0, 0])
} else {
image::Rgb([0, 0, 255])
},
);
}
}
img.save(img_dir.join("a.png")).unwrap();
let mut line = String::from("0 0.5 0.5 0.5 0.5");
for i in 0..17 {
let v = if i % 3 == 0 { 0.0 } else { 2.0 };
line.push_str(&format!(" {} 0.5 {v}", (i as f32 + 1.0) / 19.0));
}
line.push('\n');
std::fs::write(lbl_dir.join("a.txt"), line).unwrap();
let plain = load_cocopose_dir(&dir, "train", 32, Device::Cpu, false).unwrap();
let raw = load_cocopose_dir_raw(&dir, "train").unwrap();
let plan = AugmentPlan {
flip: true,
..AugmentPlan::none()
};
let enc = encode_keypoint_sample(&raw[0], 32, Device::Cpu, &plan, false).unwrap();
let (p, e) = (&plain[0], &enc);
let px = |t: &Tensor, c: usize, y: usize, x: usize| {
t.double_value(&[c as i64, y as i64, x as i64])
};
assert!(px(&e.x, 0, 0, 0) < 1e-6, "翻转后左上应为蓝的 R=0");
assert!(
(px(&e.x, 2, 0, 0) - 1.0).abs() < 1e-6,
"翻转后左上应为蓝的 B=1"
);
assert!((px(&e.x, 0, 0, 31) - 1.0).abs() < 1e-6, "翻转后右上应为红");
assert!((e.boxes[0][0] - (32.0 - p.boxes[0][0])).abs() < 1e-4);
assert!((e.boxes[0][2] - p.boxes[0][2]).abs() < 1e-4);
for (i, &sw) in COCO17_FLIP_SWAP.iter().enumerate() {
let src = &p.kpts[0][sw];
assert!(
(e.kpts[0][i][0] - (32.0 - src[0])).abs() < 1e-4,
"kpt{i}: {} != 32−{}",
e.kpts[0][i][0],
src[0]
);
assert!((e.kpts[0][i][1] - src[1]).abs() < 1e-4, "kpt{i} y 不应变");
assert_eq!(e.kpts[0][i][2], src[2], "kpt{i} v 应随点换位");
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn encode_detect_flip_and_scale_boxes_hand_computed() {
let dir = std::env::temp_dir().join(format!("av-ds-augdet-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
image::RgbImage::from_pixel(64, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(lbl_dir.join("a.txt"), "3 0.75 0.25 0.25 0.25\n").unwrap();
let raw = load_yolo_dir_raw(&dir, "train").unwrap();
assert_eq!(raw[0].boxes, vec![[40.0, 4.0, 56.0, 12.0]], "原图像素 xyxy");
let plain =
load_yolo_dir_with_mode(&dir, "train", 32, Device::Cpu, ResizeMode::Letterbox, false)
.unwrap();
let enc = encode_detect_sample(
&raw[0],
32,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan::none(),
false,
)
.unwrap();
assert_eq!(tensor_max_diff(&plain[0].x, &enc.x), 0.0);
assert_eq!(plain[0].boxes, enc.boxes);
let fl = encode_detect_sample(
&raw[0],
32,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan {
flip: true,
..AugmentPlan::none()
},
false,
)
.unwrap();
let b = fl.boxes[0];
assert!(
(b[0] - 4.0).abs() < 1e-4 && (b[1] - 10.0).abs() < 1e-4,
"b={b:?}"
);
assert!(
(b[2] - 12.0).abs() < 1e-4 && (b[3] - 14.0).abs() < 1e-4,
"b={b:?}"
);
let sc = encode_detect_sample(
&raw[0],
32,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan {
scale: 0.5,
..AugmentPlan::none()
},
false,
)
.unwrap();
let b = sc.boxes[0];
assert!(
(b[0] - 20.0).abs() < 1e-4 && (b[1] - 10.0).abs() < 1e-4,
"b={b:?}"
);
assert!(
(b[2] - 28.0).abs() < 1e-4 && (b[3] - 14.0).abs() < 1e-4,
"b={b:?}"
);
assert_eq!(sc.x.size(), vec![3, 32, 32], "输出画布尺寸不变");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn encode_seg_none_matches_and_flip_mirrors_mask() {
let dir = std::env::temp_dir().join(format!("av-ds-augseg-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
image::RgbImage::from_pixel(32, 32, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(lbl_dir.join("a.txt"), "7 0.1 0.1 0.5 0.1 0.5 0.5 0.1 0.5\n").unwrap();
let plain = load_cocoseg_dir(&dir, "train", 32, Device::Cpu, false).unwrap();
let raw = load_cocoseg_dir_raw(&dir, "train").unwrap();
let none =
encode_seg_sample(&raw[0], 32, Device::Cpu, &AugmentPlan::none(), false).unwrap();
assert_eq!(plain[0].labels, none.labels);
assert_eq!(plain[0].masks, none.masks, "none() 掩码应逐位一致");
assert_eq!(tensor_max_diff(&plain[0].x, &none.x), 0.0);
let (mw, mh) = (8usize, 8usize);
assert_eq!(none.masks[0].iter().filter(|&&v| v == 1).count(), 9);
let fl = encode_seg_sample(
&raw[0],
32,
Device::Cpu,
&AugmentPlan {
flip: true,
..AugmentPlan::none()
},
false,
)
.unwrap();
assert_eq!(fl.labels, plain[0].labels, "翻转不应丢实例");
for y in 0..mh {
for x in 0..mw {
assert_eq!(
fl.masks[0][y * mw + x],
none.masks[0][y * mw + (mw - 1 - x)],
"翻转掩码应逐像素镜像 ({x},{y})"
);
}
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn encode_obb_none_matches_and_flip_negates_theta() {
use av_core::conventions::AngleDomain;
let dir = std::env::temp_dir().join(format!("av-ds-augobb-{}", std::process::id()));
let (img_dir, lbl_dir) = (dir.join("images/train"), dir.join("labels/train"));
std::fs::create_dir_all(&img_dir).unwrap();
std::fs::create_dir_all(&lbl_dir).unwrap();
let (th, cw, ch) = (30f32.to_radians(), 0.25f32, 0.125f32);
let (c, s) = (th.cos(), th.sin());
let corners: Vec<[f32; 2]> = [(-cw, -ch), (cw, -ch), (cw, ch), (-cw, ch)]
.iter()
.map(|&(dx, dy)| [0.5 + dx * c - dy * s, 0.5 + dx * s + dy * c])
.collect();
let line = format!(
"5 {}\n",
corners
.iter()
.map(|p| format!("{:.6} {:.6}", p[0], p[1]))
.collect::<Vec<_>>()
.join(" ")
);
image::RgbImage::from_pixel(64, 64, image::Rgb([255, 0, 0]))
.save(img_dir.join("a.png"))
.unwrap();
std::fs::write(lbl_dir.join("a.txt"), line).unwrap();
let plain = load_dota_dir(&dir, "train", 32, Device::Cpu, false).unwrap();
let raw = load_dota_dir_raw(&dir, "train").unwrap();
let none =
encode_obb_sample(&raw[0], 32, Device::Cpu, &AugmentPlan::none(), false).unwrap();
assert_eq!(plain[0].boxes.len(), 1);
assert_eq!(none.labels, plain[0].labels);
for (a, b) in plain[0].boxes.iter().zip(&none.boxes) {
for (va, vb) in a.iter().zip(b) {
assert!(
(va - vb).abs() < 1e-4,
"none() 应与 plain 一致: {va} vs {vb}"
);
}
}
assert_eq!(tensor_max_diff(&plain[0].x, &none.x), 0.0);
let fl = encode_obb_sample(
&raw[0],
32,
Device::Cpu,
&AugmentPlan {
flip: true,
..AugmentPlan::none()
},
false,
)
.unwrap();
let (p, e) = (&plain[0].boxes[0], &fl.boxes[0]);
assert!(
(e[0] - (32.0 - p[0])).abs() < 1e-4,
"cx 镜像: {} vs 32−{}",
e[0],
p[0]
);
assert!((e[1] - p[1]).abs() < 1e-4);
assert!(
(e[2] - p[2]).abs() < 1e-4 && (e[3] - p[3]).abs() < 1e-4,
"wh 不变"
);
let (pn, en) = (
AngleDomain::Le90.normalize(p[4]),
AngleDomain::Le90.normalize(e[4]),
);
assert!(
(en + pn).abs() < 1e-3 || ((en - pn).abs() < 1e-3 && (p[2] - p[3]).abs() < 1e-3),
"镜像应 θ→−θ(le90 归一化后):{pn} vs {en}"
);
let _ = std::fs::remove_dir_all(&dir);
}
fn solid_raw(
w: u32,
h: u32,
rgb: [u8; 3],
boxes: Vec<[f32; 4]>,
labels: Vec<u32>,
) -> RawDetectSample {
RawDetectSample {
w,
h,
rgb: vec![rgb; (w * h) as usize].into_iter().flatten().collect(),
boxes,
labels,
}
}
#[test]
fn mosaic4_raw_composes_quadrants_and_boxes() {
let red = solid_raw(2, 2, [255, 0, 0], vec![[0.0, 0.0, 2.0, 2.0]], vec![7]);
let green = solid_raw(2, 2, [0, 255, 0], vec![[0.0, 0.0, 1.0, 1.0]], vec![8]);
let blue = solid_raw(2, 2, [0, 0, 255], vec![[1.0, 1.0, 2.0, 2.0]], vec![9]);
let white = solid_raw(2, 2, [255, 255, 255], vec![[0.0, 1.0, 1.0, 2.0]], vec![10]);
let m = mosaic4_raw([&red, &green, &blue, &white]).unwrap();
assert_eq!((m.w, m.h), (4, 4));
assert_eq!(m.rgb.len(), 4 * 4 * 3);
let pixel = |x: usize, y: usize| &m.rgb[(y * 4 + x) * 3..(y * 4 + x) * 3 + 3];
assert_eq!(pixel(0, 0), &[255, 0, 0], "左上=锚点红");
assert_eq!(pixel(3, 0), &[0, 255, 0]);
assert_eq!(pixel(0, 3), &[0, 0, 255]);
assert_eq!(pixel(3, 3), &[255, 255, 255]);
assert_eq!(
m.boxes,
vec![
[0.0, 0.0, 2.0, 2.0],
[2.0, 0.0, 3.0, 1.0],
[1.0, 2.0 + 1.0, 2.0, 2.0 + 2.0],
[2.0, 2.0 + 1.0, 2.0 + 1.0, 2.0 + 2.0],
]
);
assert_eq!(m.labels, vec![7, 8, 9, 10]);
let rep = mosaic4_raw([&red, &red, &red, &red]).unwrap();
assert_eq!(
rep.boxes,
vec![
[0.0, 0.0, 2.0, 2.0],
[2.0, 0.0, 4.0, 2.0],
[0.0, 2.0, 2.0, 4.0],
[2.0, 2.0, 4.0, 4.0],
]
);
assert_eq!(rep.labels, vec![7; 4]);
assert!(rep.rgb.chunks(3).all(|px| px == [255, 0, 0]));
}
#[test]
fn mosaic4_raw_scales_boxes_with_pixels() {
let big = solid_raw(4, 4, [1, 2, 3], vec![[2.0, 2.0, 4.0, 4.0]], vec![3]);
let anchor = solid_raw(2, 2, [9, 9, 9], vec![], vec![]);
let m = mosaic4_raw([&anchor, &big, &big, &big]).unwrap();
assert_eq!((m.w, m.h), (4, 4));
assert_eq!(
m.boxes,
vec![
[3.0, 1.0, 4.0, 2.0],
[1.0, 3.0, 2.0, 4.0],
[3.0, 3.0, 4.0, 4.0]
]
);
assert_eq!(m.labels, vec![3, 3, 3]);
}
#[test]
fn mixup_raw_blends_and_merges_labels() {
let mut a = solid_raw(2, 1, [0, 0, 0], vec![[0.0, 0.0, 1.0, 1.0]], vec![1]);
a.rgb = vec![100, 0, 255, 0, 255, 0];
let mut b = solid_raw(2, 1, [0, 0, 0], vec![[0.0, 0.0, 2.0, 1.0]], vec![2]);
b.rgb = vec![200, 100, 0, 100, 0, 200];
let m = mixup_raw(&a, &b, 0.25).unwrap();
assert_eq!((m.w, m.h), (2, 1));
assert_eq!(&m.rgb[..3], &[175, 75, 64]);
assert_eq!(&m.rgb[3..], &[75, 64, 150]);
assert_eq!(m.boxes, vec![[0.0, 0.0, 1.0, 1.0], [0.0, 0.0, 2.0, 1.0]]);
assert_eq!(m.labels, vec![1, 2]);
let m1 = mixup_raw(&a, &b, 1.0).unwrap();
assert_eq!(m1.rgb, a.rgb);
let big = solid_raw(4, 2, [255, 255, 255], vec![[0.0, 0.0, 4.0, 2.0]], vec![5]);
let m2 = mixup_raw(&a, &big, 0.5).unwrap();
assert_eq!((m2.w, m2.h), (2, 1), "输出尺寸随 a");
assert_eq!(m2.labels, vec![1, 5], "并集含 partner 标签");
assert_eq!(m2.boxes.len(), 2);
}
}
#[cfg(all(test, feature = "torch"))]
fn synthetic_raw_for_test(w: u32, h: u32) -> RawSegSample {
let mut rgb = Vec::with_capacity((w * h * 3) as usize);
for y in 0..h {
for x in 0..w {
let fx = x as f32 / w as f32;
let fy = y as f32 / h as f32;
rgb.push((40.0 + 190.0 * fx) as u8);
rgb.push((60.0 + 160.0 * fy) as u8);
rgb.push((70.0 + 150.0 * (fx * 0.5 + fy * 0.5)) as u8);
}
}
let poly_rect = vec![
[0.1 * w as f32, 0.1 * h as f32],
[0.6 * w as f32, 0.12 * h as f32],
[0.62 * w as f32, 0.55 * h as f32],
[0.12 * w as f32, 0.5 * h as f32],
];
let poly_tri = vec![
[0.7 * w as f32, 0.6 * h as f32],
[0.95 * w as f32, 0.65 * h as f32],
[0.8 * w as f32, 0.92 * h as f32],
];
RawSegSample {
w,
h,
rgb,
polys: vec![poly_rect, poly_tri],
labels: vec![0, 1],
}
}
#[cfg(all(test, feature = "torch"))]
mod seg_cache_tests {
use super::*;
use av_tasks::augment::AugmentPlan;
fn synthetic_raw(w: u32, h: u32) -> RawSegSample {
super::synthetic_raw_for_test(w, h)
}
fn tensor_diff(a: &Tensor, b: &Tensor) -> (f32, f32) {
let d = (a - b).abs();
(
d.max().double_value(&[]) as f32,
d.mean(Kind::Float).double_value(&[]) as f32,
)
}
fn assert_masks_eq(a: &SegSample, b: &SegSample) {
assert_eq!(a.masks.len(), b.masks.len(), "实例数应一致");
for (ma, mb) in a.masks.iter().zip(&b.masks) {
assert_eq!(ma, mb, "掩码必须逐位一致");
}
assert_eq!(a.labels, b.labels);
}
#[test]
fn cached_none_plan_is_bit_exact() {
let raw = synthetic_raw(160, 128);
let img_size = 128;
let reference =
encode_seg_sample(&raw, img_size, Device::Cpu, &AugmentPlan::none(), false).unwrap();
let cached = build_seg_cache_sample(&raw, img_size).unwrap();
let fast =
encode_seg_sample_cached(&cached, img_size, Device::Cpu, &AugmentPlan::none(), false)
.unwrap();
let (dmax, _) = tensor_diff(&reference.x, &fast.x);
assert_eq!(dmax, 0.0, "none plan 必须逐位一致,实际最大差 {dmax}");
assert_masks_eq(&reference, &fast);
}
#[test]
fn cached_flip_plan_is_bit_exact() {
let raw = synthetic_raw(160, 128);
let img_size = 128;
let plan = AugmentPlan {
flip: true,
scale: 1.0,
rgb_gains: [1.0; 3],
};
let reference = encode_seg_sample(&raw, img_size, Device::Cpu, &plan, false).unwrap();
let cached = build_seg_cache_sample(&raw, img_size).unwrap();
let fast = encode_seg_sample_cached(&cached, img_size, Device::Cpu, &plan, false).unwrap();
let (dmax, _) = tensor_diff(&reference.x, &fast.x);
assert_eq!(dmax, 0.0, "flip 路径必须逐位一致,实际最大差 {dmax}");
assert_masks_eq(&reference, &fast);
}
#[test]
fn cached_gains_within_tolerance() {
let raw = synthetic_raw(160, 128);
let img_size = 128;
let plan = AugmentPlan {
flip: true,
scale: 1.0,
rgb_gains: [1.08, 0.92, 1.0],
};
let reference = encode_seg_sample(&raw, img_size, Device::Cpu, &plan, false).unwrap();
let cached = build_seg_cache_sample(&raw, img_size).unwrap();
let fast = encode_seg_sample_cached(&cached, img_size, Device::Cpu, &plan, false).unwrap();
let (dmax, dmean) = tensor_diff(&reference.x, &fast.x);
assert!(dmax <= 2.0 / 255.0, "增益路径最大差 {dmax} 超容差 2/255");
assert!(dmean < 0.2 / 255.0, "增益路径平均差 {dmean} 偏大");
assert_masks_eq(&reference, &fast);
}
#[test]
fn cached_scale_within_tolerance() {
let raw = synthetic_raw(160, 128);
let img_size = 128;
let plan = AugmentPlan {
flip: false,
scale: 1.15,
rgb_gains: [1.0; 3],
};
let reference = encode_seg_sample(&raw, img_size, Device::Cpu, &plan, false).unwrap();
let cached = build_seg_cache_sample(&raw, img_size).unwrap();
let fast = encode_seg_sample_cached(&cached, img_size, Device::Cpu, &plan, false).unwrap();
let (dmax, dmean) = tensor_diff(&reference.x, &fast.x);
assert!(dmax <= 4.0 / 255.0, "缩放路径最大差 {dmax} 超容差 4/255");
assert!(dmean < 1.0 / 255.0, "缩放路径平均差 {dmean} 偏大");
assert_masks_eq(&reference, &fast);
}
#[test]
fn cache_geometry_matches_letterbox() {
let raw = synthetic_raw(244, 204); let img_size = 128;
let c = build_seg_cache_sample(&raw, img_size).unwrap();
let lb = letterbox(raw.w, raw.h, img_size, 1);
let nw = ((raw.w as f32 * lb.scale).round() as u32).max(1);
let nh = ((raw.h as f32 * lb.scale).round() as u32).max(1);
assert_eq!((c.cw, c.ch), (nw, nh));
assert_eq!(c.content.len(), (nw * nh * 3) as usize);
}
#[test]
fn gpu_stack_matches_cpu_cached() {
if !matches!(Device::cuda_if_available(), Device::Cuda(_)) {
return; }
let device = Device::Cuda(0);
let raw = synthetic_raw(160, 128);
let img_size = 128;
let cached = build_seg_cache_sample(&raw, img_size).unwrap();
let stack =
build_seg_canvas_stack(std::slice::from_ref(&cached), img_size, device).unwrap();
assert_eq!(
stack.size(),
vec![1i64, 3, img_size as i64, img_size as i64],
"画布堆形状错误"
);
for (name, plan, tol_max, tol_mean) in [
(
"flip",
AugmentPlan {
flip: true,
scale: 1.0,
rgb_gains: [1.0; 3],
},
2.0f32,
0.2f32,
),
(
"gain",
AugmentPlan {
flip: false,
scale: 1.0,
rgb_gains: [1.06, 0.95, 1.0],
},
2.0,
0.2,
),
(
"scale",
AugmentPlan {
flip: false,
scale: 1.1,
rgb_gains: [1.0; 3],
},
8.0,
2.0,
),
] {
let cpu = encode_seg_sample_cached(&cached, img_size, Device::Cpu, &plan, false)
.unwrap()
.x
.to_device(device);
let gpu = encode_seg_sample_gpu(&stack, 0, &cached, img_size, &plan, false)
.unwrap()
.x;
let (dmax, dmean) = tensor_diff(&cpu, &gpu);
assert!(
dmax <= tol_max / 255.0 && dmean <= tol_mean / 255.0,
"{name}: GPU/CPU 差 max={dmax} mean={dmean} 超容差 ({tol_max},{tol_mean})/255"
);
}
}
}
#[test]
#[ignore = "基准:需要显式运行(--ignored --nocapture)"]
fn seg_cache_bench() {
use std::time::Instant;
let (w, h) = (2448u32, 2048u32);
let img_size = 640;
let n = 8;
let raws: Vec<RawSegSample> = (0..n)
.map(|k| {
let mut r = {
let mut s = synthetic_raw_for_test(w, h);
for v in s.rgb.iter_mut().step_by(97) {
*v = v.wrapping_add(k as u8 * 7);
}
s
};
r.polys = vec![vec![
[0.1 * w as f32, 0.1 * h as f32],
[0.6 * w as f32, 0.6 * h as f32],
]];
r.labels = vec![0];
r
})
.collect();
let plans = vec![
AugmentPlan {
flip: true,
scale: 1.05,
rgb_gains: [1.02, 0.98, 1.0]
};
n
];
let t0 = Instant::now();
let _old: Vec<_> = raws
.iter()
.zip(&plans)
.map(|(r, p)| encode_seg_sample(r, img_size, Device::Cpu, p, false).unwrap())
.collect();
let old_ms = t0.elapsed().as_millis() as f64 / n as f64;
let t1 = Instant::now();
let (cache, _bytes) = build_seg_cache(&raws, img_size).unwrap();
let build_ms = t1.elapsed().as_millis() as f64 / n as f64;
let t2 = Instant::now();
let idx: Vec<usize> = (0..n).collect();
let _new = encode_seg_batch_cached(&cache, &idx, &plans, img_size, Device::Cpu, false).unwrap();
let new_ms = t2.elapsed().as_millis() as f64 / n as f64;
println!("旧路径(全分辨率单线程): {old_ms:.1} ms/样本");
println!(
"新路径(缓存+rayon 并行): {new_ms:.1} ms/样本(缓存构建一次性 {build_ms:.1} ms/样本)"
);
println!("稳态加速比: {:.0}x", old_ms / new_ms);
assert!(old_ms > new_ms * 4.0, "新路径应显著快于旧路径");
}
#[cfg(test)]
mod detect_cache_tests {
use super::*;
fn tensor_diff(a: &Tensor, b: &Tensor) -> (f32, f32) {
let d = (a - b).abs();
(
d.max().double_value(&[]) as f32,
d.mean(Kind::Float).double_value(&[]) as f32,
)
}
fn synthetic_detect_raw(w: u32, h: u32) -> RawDetectSample {
let mut rgb = Vec::with_capacity((w * h * 3) as usize);
for y in 0..h {
for x in 0..w {
let fx = x as f32 / w as f32;
let fy = y as f32 / h as f32;
rgb.push((30.0 + 200.0 * fx) as u8);
rgb.push((50.0 + 180.0 * fy) as u8);
rgb.push((90.0 + 140.0 * (fx * 0.5 + fy * 0.5)) as u8);
}
}
RawDetectSample {
w,
h,
rgb,
boxes: vec![
[
0.1 * w as f32,
0.2 * h as f32,
0.4 * w as f32,
0.6 * h as f32,
],
[
0.55 * w as f32,
0.1 * h as f32,
0.9 * w as f32,
0.5 * h as f32,
],
],
labels: vec![1, 0],
}
}
fn boxes_of(s: &SampleTensor) -> Vec<[f32; 4]> {
s.boxes.clone()
}
#[test]
fn detect_cache_none_plan_is_bit_exact_landscape() {
let raw = synthetic_detect_raw(640, 426);
let img_size = 320;
let tile = detect_content_tile(&raw, img_size).unwrap();
assert_eq!(tile.w, img_size, "横图长边贴到 img_size");
assert_eq!(tile.h, ((426.0 / 640.0) * img_size as f32).round() as u32);
let reference = encode_detect_sample(
&raw,
img_size,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan::none(),
false,
)
.unwrap();
let fast = encode_detect_sample(
&tile,
img_size,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan::none(),
false,
)
.unwrap();
let (dmax, _) = tensor_diff(&reference.x, &fast.x);
assert_eq!(dmax, 0.0, "none plan 必须逐位一致,实际最大差 {dmax}");
assert_eq!(boxes_of(&reference), boxes_of(&fast), "框映射必须逐位一致");
assert_eq!(reference.labels, fast.labels);
}
#[test]
fn detect_cache_none_plan_is_bit_exact_portrait() {
let raw = synthetic_detect_raw(426, 640);
let img_size = 320;
let tile = detect_content_tile(&raw, img_size).unwrap();
assert_eq!(tile.h, img_size, "竖图长边贴到 img_size");
let reference = encode_detect_sample(
&raw,
img_size,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan::none(),
false,
)
.unwrap();
let fast = encode_detect_sample(
&tile,
img_size,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan::none(),
false,
)
.unwrap();
let (dmax, _) = tensor_diff(&reference.x, &fast.x);
assert_eq!(dmax, 0.0, "竖图 none plan 必须逐位一致,实际最大差 {dmax}");
assert_eq!(boxes_of(&reference), boxes_of(&fast));
}
#[test]
fn detect_cache_stretch_mode_is_bit_exact() {
let raw = synthetic_detect_raw(500, 333);
let img_size = 320;
let tile = detect_content_tile(&raw, img_size).unwrap();
let a = encode_detect_sample(
&raw,
img_size,
Device::Cpu,
ResizeMode::Stretch,
&AugmentPlan::none(),
false,
)
.unwrap();
let b = encode_detect_sample(
&tile,
img_size,
Device::Cpu,
ResizeMode::Stretch,
&AugmentPlan::none(),
false,
)
.unwrap();
for (ba, bb) in a.boxes.iter().zip(&b.boxes) {
for k in 0..4 {
assert!(
(ba[k] - bb[k]).abs() < 1.5,
"拉伸框几何偏差 ≤1.5px,实际 {}",
(ba[k] - bb[k]).abs()
);
}
}
}
#[test]
fn detect_tile_cache_batch_preserves_order_and_dims() {
let img_size = 160;
let raws = vec![
synthetic_detect_raw(320, 200),
synthetic_detect_raw(200, 320),
synthetic_detect_raw(160, 160),
];
let tiles = build_detect_tile_cache(raws, img_size).unwrap();
assert_eq!(tiles.len(), 3);
assert_eq!((tiles[0].w, tiles[0].h), (img_size, 100));
assert_eq!((tiles[1].w, tiles[1].h), (100, img_size));
assert_eq!((tiles[2].w, tiles[2].h), (img_size, img_size));
for t in &tiles {
assert!(t.boxes[0][0] < t.boxes[1][0]);
assert_eq!(t.labels, vec![1, 0]);
}
}
#[test]
fn detect_mosaic_on_tiles_matches_canvas_contract() {
let img_size = 160;
let tiles: Vec<RawDetectSample> = (0..4)
.map(|k| {
let mut r = synthetic_detect_raw(320 + k, 200 + 2 * k);
for v in r.rgb.iter_mut().step_by(97) {
*v = v.wrapping_add(k as u8 * 11);
}
detect_content_tile(&r, img_size).unwrap()
})
.collect();
let m = mosaic4_raw([&tiles[0], &tiles[1], &tiles[2], &tiles[3]]).unwrap();
assert_eq!((m.w, m.h), (2 * tiles[0].w, 2 * tiles[0].h));
assert_eq!(m.boxes.len(), 8, "四图框并集");
for b in &m.boxes {
assert!(b[0] >= 0.0 && b[1] >= 0.0 && b[2] <= m.w as f32 && b[3] <= m.h as f32);
}
let s = encode_detect_sample(
&m,
img_size,
Device::Cpu,
ResizeMode::Letterbox,
&AugmentPlan::none(),
false,
)
.unwrap();
assert_eq!(s.x.size(), [3, img_size as i64, img_size as i64]);
}
}