use taconite::cpu::par_rows;
fn bicubic(x: f64) -> f64 {
const A: f64 = -0.5;
let x = x.abs();
if x < 1.0 {
((A + 2.0) * x - (A + 3.0)) * x * x + 1.0
} else if x < 2.0 {
(((x - 5.0) * x + 8.0) * x - 4.0) * A
} else {
0.0
}
}
struct Axis {
taps: Vec<(usize, Vec<i32>)>,
precision: u32,
}
fn coeffs(in_size: usize, out_size: usize) -> Axis {
let scale = in_size as f64 / out_size as f64;
let filterscale = scale.max(1.0);
let support = 2.0 * filterscale;
let float: Vec<(usize, Vec<f64>)> = (0..out_size)
.map(|xx| {
let center = (xx as f64 + 0.5) * scale;
let xmin = ((center - support + 0.5) as i64).max(0) as usize;
let xmax = ((center + support + 0.5) as i64).min(in_size as i64) as usize;
let mut k: Vec<f64> = (xmin..xmax).map(|x| bicubic((x as f64 - center + 0.5) / filterscale)).collect();
let ww: f64 = k.iter().sum();
if ww != 0.0 {
k.iter_mut().for_each(|v| *v /= ww);
}
(xmin, k)
})
.collect();
let max_w = float.iter().flat_map(|(_, k)| k.iter().copied()).fold(0.0, f64::max);
let mut precision = 0;
while precision < 22 {
if (0.5 + max_w * (1u64 << (precision + 1)) as f64) as i64 >= 1 << 15 {
break;
}
precision += 1;
}
let s = (1u64 << precision) as f64;
let taps = float
.into_iter()
.map(|(xmin, k)| (xmin, k.iter().map(|&v| (v * s + if v < 0.0 { -0.5 } else { 0.5 }) as i32).collect()))
.collect();
Axis { taps, precision }
}
pub fn resize_rgb8(src: &[u8], w: usize, h: usize, ow: usize, oh: usize) -> Vec<u8> {
let pass = |ss: i64, p: u32| (ss >> p).clamp(0, 255) as u8;
let tmp = if ow == w {
src.to_vec()
} else {
let cx = coeffs(w, ow);
let mut tmp = vec![0u8; h * ow * 3];
par_rows(&mut tmp, ow * 3, |y0, piece| {
for (yi, row) in piece.chunks_mut(ow * 3).enumerate() {
let s = &src[(y0 + yi) * w * 3..][..w * 3];
for (xx, (xmin, k)) in cx.taps.iter().enumerate() {
for ch in 0..3 {
let mut ss = 1i64 << (cx.precision - 1);
for (i, &kv) in k.iter().enumerate() {
ss += s[(xmin + i) * 3 + ch] as i64 * kv as i64;
}
row[xx * 3 + ch] = pass(ss, cx.precision);
}
}
}
});
tmp
};
if oh == h {
return tmp;
}
let cy = coeffs(h, oh);
let mut out = vec![0u8; oh * ow * 3];
par_rows(&mut out, ow * 3, |y0, piece| {
for (yi, row) in piece.chunks_mut(ow * 3).enumerate() {
let (ymin, k) = &cy.taps[y0 + yi];
for (j, o) in row.iter_mut().enumerate() {
let mut ss = 1i64 << (cy.precision - 1);
for (i, &kv) in k.iter().enumerate() {
ss += tmp[(ymin + i) * ow * 3 + j] as i64 * kv as i64;
}
*o = pass(ss, cy.precision);
}
}
});
out
}
pub fn target_size(h: usize, w: usize, patch: usize, pool: usize, max_soft_tokens: usize) -> (usize, usize) {
let max_patches = max_soft_tokens * pool * pool;
let factor = ((max_patches * patch * patch) as f64 / (h * w) as f64).sqrt();
let side = pool * patch;
let mut th = (factor * h as f64 / side as f64).floor() as usize * side;
let mut tw = (factor * w as f64 / side as f64).floor() as usize * side;
let max_side = (max_patches / (pool * pool)) * side;
if th == 0 && tw > 0 {
th = side;
tw = ((w as f64 / h as f64).floor() as usize * side).min(max_side);
} else if tw == 0 && th > 0 {
tw = side;
th = ((h as f64 / w as f64).floor() as usize * side).min(max_side);
}
(th, tw)
}
pub struct Patches {
pub pixels: Vec<f32>,
pub positions: Vec<(usize, usize)>,
pub gh: usize,
pub gw: usize,
}
pub fn patches(rgb: &[u8], w: usize, h: usize, patch: usize, pool: usize, max_soft_tokens: usize) -> Option<Patches> {
let (th, tw) = target_size(h, w, patch, pool, max_soft_tokens);
if th == 0 || tw == 0 {
return None;
}
let img = if (th, tw) == (h, w) { rgb.to_vec() } else { resize_rgb8(rgb, w, h, tw, th) };
let (gh, gw) = (th / patch, tw / patch);
let pd = patch * patch * 3;
let mut pixels = vec![0f32; gh * gw * pd];
let mut positions = Vec::with_capacity(gh * gw);
for py in 0..gh {
for px in 0..gw {
let o = &mut pixels[(py * gw + px) * pd..][..pd];
for r in 0..patch {
for c in 0..patch {
let s = ((py * patch + r) * tw + px * patch + c) * 3;
for ch in 0..3 {
o[(r * patch + c) * 3 + ch] = img[s + ch] as f32 * (1.0f32 / 255.0);
}
}
}
positions.push((px, py));
}
}
Some(Patches { pixels, positions, gh, gw })
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn identity_pass_is_exact() {
let src: Vec<u8> = (0..6 * 5 * 3).map(|i| (i * 7 % 256) as u8).collect();
assert_eq!(resize_rgb8(&src, 6, 5, 6, 5), src);
}
#[test]
fn sizes_match_the_processor() {
assert_eq!(target_size(1764, 2646, 16, 3, 280), (39 * 16, 60 * 16));
assert_eq!(target_size(480, 640, 16, 3, 280), (42 * 16, 57 * 16));
assert_eq!(target_size(425, 640, 16, 3, 280), (39 * 16, 60 * 16));
}
}