use av_core::config::AugmentCfg;
use crate::rng::XorShift;
pub const COCO17_FLIP_SWAP: [usize; 17] =
[0, 2, 1, 4, 3, 6, 5, 8, 7, 10, 9, 12, 11, 14, 13, 16, 15];
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct AugmentPlan {
pub flip: bool,
pub scale: f32,
pub rgb_gains: [f32; 3],
}
impl AugmentPlan {
pub fn none() -> Self {
Self {
flip: false,
scale: 1.0,
rgb_gains: [1.0; 3],
}
}
pub fn is_none(&self) -> bool {
!self.flip && self.scale == 1.0 && self.rgb_gains == [1.0; 3]
}
}
pub fn has_strength(cfg: &AugmentCfg) -> bool {
cfg.mosaic > 0.0
|| cfg.mixup > 0.0
|| cfg.flip > 0.0
|| cfg.hsv.iter().any(|&g| g > 0.0)
|| cfg.scale_jitter.is_some()
}
pub fn strong_aug_active(cfg: &AugmentCfg, epoch: u32, total_epochs: u32) -> bool {
cfg.close_last_epochs == 0 || epoch <= total_epochs.saturating_sub(cfg.close_last_epochs)
}
pub fn draw_plan(cfg: &AugmentCfg, rng: &mut XorShift) -> AugmentPlan {
AugmentPlan {
flip: cfg.flip > 0.0 && rng.next_f32() < cfg.flip,
scale: match cfg.scale_jitter {
Some([lo, hi]) => rng.next_range(lo, hi),
None => 1.0,
},
rgb_gains: [
1.0 + rng.next_range(-cfg.hsv[0], cfg.hsv[0]),
1.0 + rng.next_range(-cfg.hsv[1], cfg.hsv[1]),
1.0 + rng.next_range(-cfg.hsv[2], cfg.hsv[2]),
],
}
}
pub fn hflip_rgb(w: usize, h: usize, rgb: &mut [u8]) {
for y in 0..h {
let row = &mut rgb[y * w * 3..(y + 1) * w * 3];
for x in 0..w / 2 {
let (a, b) = (x * 3, (w - 1 - x) * 3);
row.swap(a, b);
row.swap(a + 1, b + 1);
row.swap(a + 2, b + 2);
}
}
}
pub fn mul_rgb(rgb: &mut [u8], gains: [f32; 3]) {
if gains == [1.0, 1.0, 1.0] {
return; }
let mut lut = [[0u8; 256]; 3];
for (c, gc) in gains.iter().enumerate() {
for (v, slot) in lut[c].iter_mut().enumerate() {
*slot = (v as f32 * gc).round().clamp(0.0, 255.0) as u8;
}
}
for px in rgb.chunks_exact_mut(3) {
px[0] = lut[0][px[0] as usize];
px[1] = lut[1][px[1] as usize];
px[2] = lut[2][px[2] as usize];
}
}
pub fn flip_x(x: f32, w: f32) -> f32 {
w - x
}
pub fn flip_box_xyxy(b: [f32; 4], w: f32) -> [f32; 4] {
[w - b[2], b[1], w - b[0], b[3]]
}
pub fn flip_box_cxcywh(b: [f32; 4], w: f32) -> [f32; 4] {
[w - b[0], b[1], b[2], b[3]]
}
pub fn flip_keypoints(kpts: &mut [Vec<[f32; 3]>], w: f32) {
for g in kpts.iter_mut() {
for p in g.iter_mut() {
p[0] = w - p[0];
}
}
swap_coco17_keypoints(kpts);
}
pub fn swap_coco17_keypoints(kpts: &mut [Vec<[f32; 3]>]) {
for g in kpts.iter_mut() {
if g.len() == 17 {
let old = g.clone();
for (i, &src) in COCO17_FLIP_SWAP.iter().enumerate() {
g[i] = old[src];
}
}
}
}
pub fn flip_polygon(poly: &mut [[f32; 2]], w: f32) {
for p in poly.iter_mut() {
p[0] = w - p[0];
}
}
pub fn scale_point(p: &mut [f32; 2], s: f32) {
p[0] *= s;
p[1] *= s;
}
pub fn scaled_dims(w: u32, h: u32, s: f32) -> (u32, u32) {
(
((w as f32 * s).round() as u32).max(1),
((h as f32 * s).round() as u32).max(1),
)
}
pub fn mosaic_canvas_dims(w: u32, h: u32) -> (u32, u32) {
(w * 2, h * 2)
}
pub fn mosaic_paste_quadrant(
canvas: &mut [u8],
canvas_w: usize,
src: &[u8],
qw: usize,
qh: usize,
x0: usize,
y0: usize,
) {
debug_assert_eq!(src.len(), qw * qh * 3, "源缓冲必须是象限尺寸");
for ry in 0..qh {
let dst = ((y0 + ry) * canvas_w + x0) * 3;
let src_off = ry * qw * 3;
canvas[dst..dst + qw * 3].copy_from_slice(&src[src_off..src_off + qw * 3]);
}
}
pub fn mosaic_map_box(
b: [f32; 4],
sw: f32,
sh: f32,
qw: f32,
qh: f32,
x0: f32,
y0: f32,
) -> Option<[f32; 4]> {
let (sx, sy) = (qw / sw, qh / sh);
let x1 = (b[0] * sx + x0).clamp(x0, x0 + qw);
let y1 = (b[1] * sy + y0).clamp(y0, y0 + qh);
let x2 = (b[2] * sx + x0).clamp(x0, x0 + qw);
let y2 = (b[3] * sy + y0).clamp(y0, y0 + qh);
if x2 - x1 <= 0.0 || y2 - y1 <= 0.0 {
None
} else {
Some([x1, y1, x2, y2])
}
}
pub struct MosaicItem<'a> {
pub rgb: &'a [u8],
pub src_w: u32,
pub src_h: u32,
pub boxes: &'a [[f32; 4]],
pub labels: &'a [u32],
}
pub fn mosaic_compose(
qw: u32,
qh: u32,
items: &[MosaicItem<'_>; 4],
) -> (Vec<u8>, Vec<[f32; 4]>, Vec<u32>) {
let (qw, qh) = (qw as usize, qh as usize);
let (cw, ch) = (qw * 2, qh * 2);
let mut canvas = vec![0u8; cw * ch * 3];
let origins = [(0usize, 0usize), (qw, 0), (0, qh), (qw, qh)];
let (mut boxes, mut labels) = (Vec::new(), Vec::new());
for (k, item) in items.iter().enumerate() {
let (x0, y0) = origins[k];
mosaic_paste_quadrant(&mut canvas, cw, item.rgb, qw, qh, x0, y0);
for (b, &l) in item.boxes.iter().zip(item.labels) {
if let Some(mb) = mosaic_map_box(
*b,
item.src_w as f32,
item.src_h as f32,
qw as f32,
qh as f32,
x0 as f32,
y0 as f32,
) {
boxes.push(mb);
labels.push(l);
}
}
}
(canvas, boxes, labels)
}
pub const MIXUP_BETA_ALPHA: f32 = 0.2;
pub fn beta_symmetric(rng: &mut XorShift, alpha: f32) -> f32 {
if alpha <= 0.0 {
return 0.5;
}
let g1 = gamma_one(rng, alpha);
let g2 = gamma_one(rng, alpha);
let s = g1 + g2;
if s <= 0.0 {
0.5
} else {
(g1 / s).clamp(f32::EPSILON, 1.0 - f32::EPSILON)
}
}
fn gamma_one(rng: &mut XorShift, shape: f32) -> f32 {
let (boost, a) = if shape < 1.0 {
let u = rng.next_f32().max(1e-12);
(u.powf(1.0 / shape), shape + 1.0)
} else {
(1.0, shape)
};
let d = a - 1.0 / 3.0;
let c = (1.0 / (9.0 * d)).sqrt();
loop {
let u1 = rng.next_f32().max(1e-12);
let u2 = rng.next_f32();
let z = (-2.0 * u1.ln()).sqrt() * (std::f32::consts::TAU * u2).cos();
let v = (1.0 + c * z).powi(3);
if v <= 0.0 {
continue;
}
let u = rng.next_f32();
let z2 = z * z;
if u < 1.0 - 0.0331 * z2 * z2 || u.ln() < 0.5 * z2 + d * (1.0 - v + v.ln()) {
return boost * d * v;
}
}
}
pub fn mixup_rgb(a: &[u8], b: &[u8], lam: f32) -> Vec<u8> {
debug_assert_eq!(a.len(), b.len(), "mixup 两缓冲必须等长");
let w = 1.0 - lam;
a.iter()
.zip(b.iter())
.map(|(&x, &y)| (x as f32 * lam + y as f32 * w).round().clamp(0.0, 255.0) as u8)
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct CompositeDraw {
pub mosaic: bool,
pub mixup: bool,
pub mixup_lam: f32,
}
impl CompositeDraw {
pub fn none() -> Self {
Self {
mosaic: false,
mixup: false,
mixup_lam: 0.5,
}
}
}
pub fn draw_composite(cfg: &AugmentCfg, rng: &mut XorShift) -> CompositeDraw {
let mosaic = cfg.mosaic > 0.0 && rng.next_f32() < cfg.mosaic;
let mixup = cfg.mixup > 0.0 && rng.next_f32() < cfg.mixup;
let mixup_lam = if mixup {
beta_symmetric(rng, MIXUP_BETA_ALPHA)
} else {
0.5
};
CompositeDraw {
mosaic,
mixup,
mixup_lam,
}
}
#[cfg(test)]
mod tests {
use super::*;
use av_core::config::AugmentCfg;
#[test]
fn flip_box_xyxy_hand_computed() {
let b = flip_box_xyxy([10.0, 20.0, 30.0, 40.0], 100.0);
assert_eq!(b, [70.0, 20.0, 90.0, 40.0]);
let b2 = flip_box_xyxy([0.0, 0.0, 50.5, 8.0], 50.5);
assert_eq!(b2, [0.0, 0.0, 50.5, 8.0]);
let round = flip_box_xyxy(flip_box_xyxy([3.25, 4.0, 17.75, 9.5], 32.0), 32.0);
assert_eq!(round, [3.25, 4.0, 17.75, 9.5]);
}
#[test]
fn flip_box_cxcywh_hand_computed() {
let b = flip_box_cxcywh([24.0, 16.0, 8.0, 4.0], 64.0);
assert_eq!(b, [40.0, 16.0, 8.0, 4.0]);
}
#[test]
fn flip_keypoints_swaps_coco17_pairs() {
let mut g: Vec<[f32; 3]> = (0..17)
.map(|i| {
[
10.0 + i as f32,
20.0 + i as f32,
if i == 3 { 0.0 } else { 2.0 },
]
})
.collect();
g[1] = [10.0, 21.0, 2.0]; g[2] = [30.0, 22.0, 2.0]; g[15] = [50.0, 35.0, 1.0]; g[16] = [60.0, 36.0, 0.0];
let mut insts = vec![g];
flip_keypoints(&mut insts, 100.0);
let g = &insts[0];
assert_eq!(g[0][0], 100.0 - 10.0, "鼻子 x 应镜像");
assert_eq!(g[0][1], 20.0);
assert_eq!(g[1][0], 70.0, "新[1] 应为旧[2] 镜像");
assert_eq!(g[1][1], 22.0);
assert_eq!(g[1][2], 2.0, "v 随点换位");
assert_eq!(g[2][0], 90.0);
assert_eq!(g[2][1], 21.0);
assert_eq!(g[15][0], 40.0, "100-60");
assert_eq!(g[15][2], 0.0);
assert_eq!(g[16][0], 50.0);
assert_eq!(g[16][2], 1.0);
assert_eq!(g[7][0], 82.0, "100-18");
assert_eq!(g[7][1], 28.0);
}
#[test]
fn flip_keypoints_non_17_no_swap() {
let mut insts = vec![vec![[10.0, 0.0, 2.0], [20.0, 1.0, 1.0], [30.0, 2.0, 2.0]]];
flip_keypoints(&mut insts, 40.0);
let g = &insts[0];
assert_eq!(g[0], [30.0, 0.0, 2.0]);
assert_eq!(g[1], [20.0, 1.0, 1.0]);
assert_eq!(g[2], [10.0, 2.0, 2.0]);
}
#[test]
fn flip_polygon_hand_computed() {
let mut poly = [[2.0, 2.0], [10.0, 2.0], [10.0, 8.0], [2.0, 8.0]];
flip_polygon(&mut poly, 12.0);
assert_eq!(poly, [[10.0, 2.0], [2.0, 2.0], [2.0, 8.0], [10.0, 8.0]]);
}
#[test]
fn scaled_dims_and_points() {
assert_eq!(scaled_dims(640, 480, 1.1), (704, 528));
assert_eq!(scaled_dims(640, 480, 0.9), (576, 432));
assert_eq!(scaled_dims(10, 10, 1.0), (10, 10));
assert_eq!(scaled_dims(3, 3, 0.05), (1, 1), "不缩到 0");
let mut p = [100.0, 50.0];
scale_point(&mut p, 0.5);
assert_eq!(p, [50.0, 25.0]);
}
#[test]
fn hflip_rgb_two_pixels() {
let mut rgb = vec![255, 0, 0, 0, 0, 255];
hflip_rgb(2, 1, &mut rgb);
assert_eq!(rgb, vec![0, 0, 255, 255, 0, 0]);
}
#[test]
fn hflip_rgb_rows_independent() {
let mut rgb = Vec::new();
for y in 0..2 {
for x in 0..3 {
rgb.extend_from_slice(&[x + 1, y + 1, 0]);
}
}
hflip_rgb(3, 2, &mut rgb);
for (i, px) in rgb.chunks(3).enumerate() {
let (x, y) = ((2 - i % 3) as u8, (i / 3) as u8); assert_eq!(px, &[x + 1, y + 1, 0], "像素 {i}");
}
}
#[test]
fn mul_rgb_gain_and_clamp() {
let mut rgb = vec![200, 100, 10, 0, 255, 128];
mul_rgb(&mut rgb, [2.0, 0.5, 1.0]);
assert_eq!(rgb[0], 255, "上截断");
assert_eq!(rgb[1], 50);
assert_eq!(rgb[2], 10);
assert_eq!(rgb[3], 0, "下截断");
assert_eq!(rgb[4], 128, "127.5 四舍五入到 128");
assert_eq!(rgb[5], 128);
let mut same = vec![7, 128, 255];
mul_rgb(&mut same, [1.0, 1.0, 1.0]);
assert_eq!(same, vec![7, 128, 255], "gain=1 应逐位恒等");
}
#[test]
fn draw_plan_deterministic_and_in_range() {
let cfg = AugmentCfg {
flip: 1.0,
hsv: [0.1, 0.2, 0.3],
scale_jitter: Some([0.9, 1.1]),
..AugmentCfg::default()
};
let mut a = XorShift::new(42);
let mut b = XorShift::new(42);
for _ in 0..50 {
let pa = draw_plan(&cfg, &mut a);
let pb = draw_plan(&cfg, &mut b);
assert_eq!(pa, pb, "同种子同序列应可复现");
assert!(pa.flip, "p=1 必翻");
assert!(
(0.9..1.1).contains(&pa.scale),
"scale 应在 [0.9, 1.1): {pa:?}"
);
for (c, (&gain, &g)) in pa.rgb_gains.iter().zip(&cfg.hsv).enumerate() {
assert!(
(1.0 - g..1.0 + g).contains(&gain),
"通道 {c} 增益应在 [1−{g}, 1+{g}): {gain}"
);
}
}
let mut r = XorShift::new(7);
let off = AugmentCfg::default();
for _ in 0..50 {
let p = draw_plan(&off, &mut r);
assert!(!p.flip);
assert_eq!(p.scale, 1.0);
assert_eq!(p.rgb_gains, [1.0; 3]);
assert!(p.is_none());
}
}
#[test]
fn strong_aug_active_close_last_semantics() {
let mut cfg = AugmentCfg {
close_last_epochs: 20,
..AugmentCfg::default()
};
assert!(strong_aug_active(&cfg, 1, 200));
assert!(strong_aug_active(&cfg, 180, 200));
assert!(!strong_aug_active(&cfg, 181, 200));
assert!(!strong_aug_active(&cfg, 200, 200));
cfg.close_last_epochs = 0;
assert!(strong_aug_active(&cfg, 200, 200));
cfg.close_last_epochs = 300;
assert!(!strong_aug_active(&cfg, 1, 200));
}
#[test]
fn has_strength_matches_config_axes() {
assert!(!has_strength(&AugmentCfg::default()));
assert!(has_strength(&AugmentCfg {
flip: 0.5,
..AugmentCfg::default()
}));
assert!(has_strength(&AugmentCfg {
hsv: [0.0, 0.2, 0.0],
..AugmentCfg::default()
}));
assert!(has_strength(&AugmentCfg {
scale_jitter: Some([0.9, 1.1]),
..AugmentCfg::default()
}));
assert!(has_strength(&AugmentCfg {
mosaic: 1.0,
..AugmentCfg::default()
}));
assert!(has_strength(&AugmentCfg {
mixup: 0.2,
..AugmentCfg::default()
}));
}
#[test]
fn mosaic_canvas_dims_hand_computed() {
assert_eq!(mosaic_canvas_dims(2, 3), (4, 6));
assert_eq!(mosaic_canvas_dims(1, 1), (2, 2));
}
#[test]
fn mosaic_map_box_hand_computed() {
assert_eq!(
mosaic_map_box([0.5, 0.5, 1.5, 1.5], 2.0, 2.0, 2.0, 2.0, 0.0, 0.0),
Some([0.5, 0.5, 1.5, 1.5])
);
assert_eq!(
mosaic_map_box([0.0, 0.0, 1.0, 1.0], 2.0, 2.0, 2.0, 2.0, 2.0, 0.0),
Some([2.0, 0.0, 3.0, 1.0])
);
assert_eq!(
mosaic_map_box([2.0, 2.0, 4.0, 4.0], 4.0, 4.0, 2.0, 2.0, 0.0, 2.0),
Some([1.0, 3.0, 2.0, 4.0])
);
assert_eq!(
mosaic_map_box([-1.0, 0.0, 3.0, 1.0], 2.0, 2.0, 2.0, 2.0, 0.0, 0.0),
Some([0.0, 0.0, 2.0, 1.0])
);
assert_eq!(
mosaic_map_box([4.0, 4.0, 6.0, 6.0], 2.0, 2.0, 2.0, 2.0, 0.0, 0.0),
None
);
assert_eq!(
mosaic_map_box([1.0, 1.0, 1.0, 2.0], 2.0, 2.0, 2.0, 2.0, 0.0, 0.0),
None
);
}
#[test]
fn mosaic_compose_hand_computed() {
let px = |r, g, b| {
vec![[r, g, b]; 4]
.into_iter()
.flatten()
.collect::<Vec<u8>>()
};
let red = px(255, 0, 0);
let green = px(0, 255, 0);
let blue = px(0, 0, 255);
let white = px(255, 255, 255);
let items = [
MosaicItem {
rgb: &red,
src_w: 2,
src_h: 2,
boxes: &[[0.0, 0.0, 2.0, 2.0]],
labels: &[7],
},
MosaicItem {
rgb: &green,
src_w: 2,
src_h: 2,
boxes: &[[0.0, 0.0, 1.0, 1.0]],
labels: &[8],
},
MosaicItem {
rgb: &blue,
src_w: 2,
src_h: 2,
boxes: &[[1.0, 1.0, 2.0, 2.0]],
labels: &[9],
},
MosaicItem {
rgb: &white,
src_w: 2,
src_h: 2,
boxes: &[[0.0, 1.0, 1.0, 2.0]],
labels: &[10],
},
];
let (canvas, boxes, labels) = mosaic_compose(2, 2, &items);
assert_eq!(canvas.len(), 4 * 4 * 3);
let pixel = |x: usize, y: usize| &canvas[(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!(
boxes,
vec![
[0.0, 0.0, 2.0, 2.0], [2.0, 0.0, 3.0, 1.0], [1.0, 3.0, 2.0, 4.0], [2.0, 3.0, 3.0, 4.0], ]
);
assert_eq!(labels, vec![7, 8, 9, 10]);
}
#[test]
fn mosaic_compose_scaled_source_maps_boxes() {
let rgb2 = vec![9u8; 2 * 2 * 3];
let items = [
MosaicItem {
rgb: &rgb2,
src_w: 4,
src_h: 4,
boxes: &[[2.0, 2.0, 4.0, 4.0]],
labels: &[1],
},
MosaicItem {
rgb: &rgb2,
src_w: 4,
src_h: 4,
boxes: &[],
labels: &[],
},
MosaicItem {
rgb: &rgb2,
src_w: 4,
src_h: 4,
boxes: &[],
labels: &[],
},
MosaicItem {
rgb: &rgb2,
src_w: 4,
src_h: 4,
boxes: &[],
labels: &[],
},
];
let (_, boxes, labels) = mosaic_compose(2, 2, &items);
assert_eq!(boxes, vec![[1.0, 1.0, 2.0, 2.0]]);
assert_eq!(labels, vec![1]);
}
#[test]
fn mixup_rgb_lambda_hand_computed() {
let out = mixup_rgb(&[100, 0, 255], &[200, 100, 0], 0.25);
assert_eq!(out, vec![175, 75, 64]);
let a = [1u8, 127, 253];
let b = [200u8, 60, 9];
assert_eq!(mixup_rgb(&a, &b, 1.0), a.to_vec(), "λ=1 应逐位等于 a");
assert_eq!(mixup_rgb(&a, &b, 0.0), b.to_vec(), "λ=0 应逐位等于 b");
let out = mixup_rgb(&[255, 0], &[255, 0], 0.5);
assert_eq!(out, vec![255, 0]);
}
#[test]
fn beta_symmetric_deterministic_and_centered() {
assert_eq!(beta_symmetric(&mut XorShift::new(1), 0.0), 0.5, "α≤0 → 0.5");
assert_eq!(beta_symmetric(&mut XorShift::new(1), -1.0), 0.5);
let mut a = XorShift::new(42);
let mut b = XorShift::new(42);
let n = 4000;
let mut sum = 0f32;
for _ in 0..n {
let x = beta_symmetric(&mut a, 0.2);
let y = beta_symmetric(&mut b, 0.2);
assert_eq!(x, y, "同种子应可复现");
assert!(x > 0.0 && x < 1.0, "λ 应在开区间 (0,1): {x}");
sum += x;
}
let mean = sum / n as f32;
assert!(
(mean - 0.5).abs() < 0.05,
"Beta(0.2,0.2) 均值应≈0.5,得到 {mean}"
);
}
#[test]
fn draw_composite_zero_prob_consumes_nothing() {
let cfg = AugmentCfg::default();
let mut r1 = XorShift::new(123);
let mut r2 = XorShift::new(123);
for _ in 0..50 {
let d = draw_composite(&cfg, &mut r1);
assert_eq!(d, CompositeDraw::none());
assert_eq!(r1.next_f32(), r2.next_f32(), "RNG 状态不应被推动");
}
}
#[test]
fn draw_composite_coins_and_lambda() {
let mut cfg = AugmentCfg {
mosaic: 1.0,
..AugmentCfg::default()
};
let mut r = XorShift::new(9);
for _ in 0..50 {
let d = draw_composite(&cfg, &mut r);
assert!(d.mosaic && !d.mixup && d.mixup_lam == 0.5);
}
cfg.mixup = 1.0;
let mut r = XorShift::new(5);
let mut first = None;
for i in 0..50 {
let d = draw_composite(&cfg, &mut r);
assert!(d.mosaic && d.mixup);
assert!(d.mixup_lam > 0.0 && d.mixup_lam < 1.0);
if i == 0 {
first = Some(d);
}
}
let mut r2 = XorShift::new(5);
assert_eq!(draw_composite(&cfg, &mut r2), first.unwrap());
}
}