AI Denoise's Apply switch becomes Method: Bilinear, Fast, Medium, Best, default Best, so an untouched raw writes nothing and develops through the mixture. `apply` is still read and never written: 0 is Bilinear, 1 keeps a network already chosen. - Best is the mixture of a flat and an edge expert with a learned gate; Medium and Fast are students distilled from it. 2.48 s, 0.79 s and 0.57 s for a 20 MP frame on TensorRT fp16. - Each network carries its own tile border (256 for the mixture, 192 for the students) through `dr_denoise::Shipped` and `TileNet::halo`. - The file is hashed once at open and each network keys its own cached result; Bilinear keeps the result in memory for the way back. - Each has an .a16w16 sibling for the Hexagon: 0.00 dB on the 6D gate, at most 0.11 dB with the noise scaled x0.5 to x4. - APK BUNDLED 19 -> 23; the PKGBUILD installs all three.
443 lines
16 KiB
Rust
443 lines
16 KiB
Rust
//! TRACES: FR-DEV-3g
|
||
//! A whole frame through a fixed-shape network, exactly (denoise.md §3.4).
|
||
//!
|
||
//! The network sees `TILE_IN`² photosites and its output is exact in the
|
||
//! central `TILE_IN − 2·HALO`: the halo is wider than its receptive field
|
||
//! (185 photosites, counted from the layers), so a tile's centre equals the
|
||
//! whole frame's at the same place. The frame is extended by reflection
|
||
//! about its edge photosites, which keeps every photosite's CFA colour, so
|
||
//! edge tiles see real context too.
|
||
//!
|
||
//! **Phase.** The network was trained on RGGB. A frame whose pattern starts
|
||
//! on another colour is read from one photosite up and/or left — the
|
||
//! reflection supplies that row or column — so its top-left is red, and the
|
||
//! output is read back from the same offset. Nothing is cropped.
|
||
|
||
use dr_decode::CfaPattern;
|
||
|
||
/// Photosites of context beyond a tile's kept centre, on every side, for a
|
||
/// single network; a mixture reaches further and says so through
|
||
/// [`TileNet::halo`].
|
||
pub const HALO: usize = 192;
|
||
|
||
/// A fixed-shape network: `mosaic` and `sigma`, `n×n` RGGB, in; `3×n×n`
|
||
/// planar linear camera RGB out.
|
||
///
|
||
/// The inputs are handed over, and the output is lent to `write` rather than
|
||
/// returned: a 1408² tile is 24 MB of output, and copying it out of the
|
||
/// runtime's buffer and back into the frame was a measurable share of a
|
||
/// frame's time.
|
||
pub trait TileNet {
|
||
/// The edge `n` of the square tile the network takes.
|
||
fn tile(&self) -> usize;
|
||
/// Photosites of context it needs past a tile's kept centre: at least
|
||
/// its receptive field. [`HALO`] unless the network says otherwise.
|
||
fn halo(&self) -> usize {
|
||
HALO
|
||
}
|
||
fn run(
|
||
&mut self,
|
||
mosaic: Vec<f32>,
|
||
sigma: Vec<f32>,
|
||
write: &mut dyn FnMut(&[f32]),
|
||
) -> Result<(), crate::DenoiseError>;
|
||
}
|
||
|
||
/// Index into `0..n` by reflection about the end photosites, any distance
|
||
/// out: …2 1 [0 1 2 … n−1] n−2 n−3…, period `2(n−1)`. Parity is kept, which
|
||
/// is what keeps a CFA colour.
|
||
#[inline]
|
||
pub fn reflect(i: isize, n: usize) -> usize {
|
||
if n == 1 {
|
||
return 0;
|
||
}
|
||
let p = 2 * (n as isize - 1);
|
||
let m = i.rem_euclid(p);
|
||
(if m < n as isize { m } else { p - m }) as usize
|
||
}
|
||
|
||
/// How far up and left to start reading so the first photosite is red.
|
||
pub fn rggb_offset(p: CfaPattern) -> Option<(usize, usize)> {
|
||
match p {
|
||
CfaPattern::Rggb => Some((0, 0)),
|
||
CfaPattern::Grbg => Some((0, 1)),
|
||
CfaPattern::Gbrg => Some((1, 0)),
|
||
CfaPattern::Bggr => Some((1, 1)),
|
||
_ => None,
|
||
}
|
||
}
|
||
|
||
/// Run `net` over an `h×w` mosaic given by `at(y, x)`, with σ from
|
||
/// `sigma(colour, value)`, and return `h×w` interleaved RGB.
|
||
///
|
||
/// `progress(done, total)` is called after each tile and stops the run by
|
||
/// returning `false`, in which case the result is `Ok(None)`.
|
||
#[allow(clippy::too_many_arguments)]
|
||
pub fn run_tiled(
|
||
net: &mut dyn TileNet,
|
||
h: usize,
|
||
w: usize,
|
||
pattern: CfaPattern,
|
||
at: &(dyn Fn(usize, usize) -> f32 + Sync),
|
||
sigma: &(dyn Fn(usize, f32) -> f32 + Sync),
|
||
progress: &mut dyn FnMut(usize, usize) -> bool,
|
||
) -> Result<Option<Vec<f32>>, crate::DenoiseError> {
|
||
let (dy, dx) = rggb_offset(pattern).ok_or_else(|| {
|
||
crate::DenoiseError::Unsupported(format!("{pattern:?} is not a Bayer pattern"))
|
||
})?;
|
||
let (n, halo) = (net.tile(), net.halo());
|
||
if n <= 2 * halo || !(n - 2 * halo).is_multiple_of(2) {
|
||
return Err(crate::DenoiseError::Model(format!(
|
||
"tile {n} leaves no even centre past a {halo} halo"
|
||
)));
|
||
}
|
||
let core = n - 2 * halo;
|
||
// In unified coordinates the frame spans u ∈ [dy, dy + h), v ∈ [dx, dx + w).
|
||
let (uh, uw) = (h + dy, w + dx);
|
||
let (ty, tx) = (uh.div_ceil(core), uw.div_ceil(core));
|
||
let total = ty * tx;
|
||
let origins: Vec<(usize, usize)> = (0..ty)
|
||
.flat_map(|i| (0..tx).map(move |j| (i * core, j * core)))
|
||
.collect();
|
||
let threads = std::thread::available_parallelism().map_or(1, |n| n.get());
|
||
|
||
// One tile's mosaic and σ, gathered on every core: rows are independent.
|
||
let gather = |u0: usize, v0: usize| {
|
||
let mut mos = vec![0.0f32; n * n];
|
||
let mut sig = vec![0.0f32; n * n];
|
||
let rows_per = n.div_ceil(threads).max(1);
|
||
std::thread::scope(|scope| {
|
||
for (chunk, (m, s)) in mos
|
||
.chunks_mut(rows_per * n)
|
||
.zip(sig.chunks_mut(rows_per * n))
|
||
.enumerate()
|
||
{
|
||
scope.spawn(move || {
|
||
for (i, (mrow, srow)) in m.chunks_mut(n).zip(s.chunks_mut(n)).enumerate() {
|
||
let r = chunk * rows_per + i;
|
||
// Unified row u = u0 + r − halo; frame row y = u − dy, reflected.
|
||
let u = u0 as isize + r as isize - halo as isize;
|
||
let y = reflect(u - dy as isize, h);
|
||
for c in 0..n {
|
||
let v = v0 as isize + c as isize - halo as isize;
|
||
let x = reflect(v - dx as isize, w);
|
||
let val = at(y, x);
|
||
mrow[c] = val;
|
||
// RGGB colour of the tile position (r, c).
|
||
srow[c] = sigma([[0, 1], [1, 2]][r & 1][c & 1], val);
|
||
}
|
||
}
|
||
});
|
||
}
|
||
});
|
||
(mos, sig)
|
||
};
|
||
|
||
// Pipelined: the next tile is gathered while the network runs this one,
|
||
// so the device does not wait on the CPU. A channel of one keeps at
|
||
// most two tiles' inputs alive.
|
||
let mut out = vec![0.0f32; h * w * 3];
|
||
let stop = std::sync::atomic::AtomicBool::new(false);
|
||
std::thread::scope(|scope| -> Result<Option<()>, crate::DenoiseError> {
|
||
let (tx_tiles, rx_tiles) = std::sync::mpsc::sync_channel(1);
|
||
let (origins, stop, gather) = (&origins, &stop, &gather);
|
||
scope.spawn(move || {
|
||
for &(u0, v0) in origins {
|
||
if stop.load(std::sync::atomic::Ordering::Relaxed) {
|
||
break;
|
||
}
|
||
if tx_tiles.send((u0, v0, gather(u0, v0))).is_err() {
|
||
break;
|
||
}
|
||
}
|
||
});
|
||
for k in 0..total {
|
||
let Ok((u0, v0, (mos, sig))) = rx_tiles.recv() else {
|
||
break;
|
||
};
|
||
let mut wrong = None;
|
||
let ran = net.run(mos, sig, &mut |rgb: &[f32]| {
|
||
if rgb.len() != 3 * n * n {
|
||
wrong = Some(rgb.len());
|
||
return;
|
||
}
|
||
// The tile's centre back into the frame: the frame rows it covers,
|
||
// split across cores (each row is written by one thread only).
|
||
let (y_lo, y_hi) = (
|
||
(u0 + dy.saturating_sub(u0)).max(dy) - dy,
|
||
(u0 + core).min(uh) - dy,
|
||
);
|
||
let (x_lo, x_hi) = ((v0.max(dx)) - dx, (v0 + core).min(uw) - dx);
|
||
if y_hi > y_lo && x_hi > x_lo {
|
||
let rows = &mut out[y_lo * w * 3..y_hi * w * 3];
|
||
let per = (y_hi - y_lo).div_ceil(threads).max(1);
|
||
std::thread::scope(|scope| {
|
||
for (chunk, block) in rows.chunks_mut(per * w * 3).enumerate() {
|
||
scope.spawn(move || {
|
||
for (i, row) in block.chunks_mut(w * 3).enumerate() {
|
||
let y = y_lo + chunk * per + i;
|
||
// Tile row of frame row y: u = y + dy = u0 + r − halo.
|
||
let r = y + dy + halo - u0;
|
||
for x in x_lo..x_hi {
|
||
let c = x + dx + halo - v0;
|
||
for ch in 0..3 {
|
||
row[x * 3 + ch] = rgb[ch * n * n + r * n + c];
|
||
}
|
||
}
|
||
}
|
||
});
|
||
}
|
||
});
|
||
}
|
||
});
|
||
if let Err(e) = ran {
|
||
stop.store(true, std::sync::atomic::Ordering::Relaxed);
|
||
return Err(e);
|
||
}
|
||
if let Some(len) = wrong {
|
||
stop.store(true, std::sync::atomic::Ordering::Relaxed);
|
||
return Err(crate::DenoiseError::Model(format!(
|
||
"network returned {len} values for a {n}² tile"
|
||
)));
|
||
}
|
||
if !progress(k + 1, total) {
|
||
stop.store(true, std::sync::atomic::Ordering::Relaxed);
|
||
// Drain so the producer is not left blocked on a full channel.
|
||
while rx_tiles.try_recv().is_ok() {}
|
||
return Ok(None);
|
||
}
|
||
}
|
||
Ok(Some(()))
|
||
})
|
||
.map(|done| done.map(|()| out))
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
#[test]
|
||
fn reflection_keeps_parity_any_distance_out() {
|
||
let n = 7;
|
||
for i in -40isize..40 {
|
||
let r = reflect(i, n);
|
||
assert!(r < n);
|
||
assert_eq!(
|
||
r % 2,
|
||
i.rem_euclid(2) as usize,
|
||
"index {i} reflected to {r}"
|
||
);
|
||
}
|
||
assert_eq!(reflect(-1, n), 1);
|
||
assert_eq!(reflect(7, n), 5);
|
||
}
|
||
|
||
/// A stand-in network with a known, finite reach: each output photosite
|
||
/// is its 2×2 quad's (R, mean G, B), averaged over the quads within
|
||
/// `reach` quads. Purely a function of the tile, like the real one.
|
||
struct BoxNet {
|
||
n: usize,
|
||
reach: usize,
|
||
}
|
||
|
||
impl TileNet for BoxNet {
|
||
fn tile(&self) -> usize {
|
||
self.n
|
||
}
|
||
fn run(
|
||
&mut self,
|
||
m: Vec<f32>,
|
||
_s: Vec<f32>,
|
||
write: &mut dyn FnMut(&[f32]),
|
||
) -> Result<(), crate::DenoiseError> {
|
||
let n = self.n;
|
||
let q = n / 2;
|
||
let quad = |qy: usize, qx: usize| {
|
||
let (y, x) = (2 * qy, 2 * qx);
|
||
[
|
||
m[y * n + x],
|
||
0.5 * (m[y * n + x + 1] + m[(y + 1) * n + x]),
|
||
m[(y + 1) * n + x + 1],
|
||
]
|
||
};
|
||
let mut out = vec![0.0; 3 * n * n];
|
||
for qy in 0..q {
|
||
for qx in 0..q {
|
||
let mut acc = [0.0f32; 3];
|
||
let mut cnt = 0.0;
|
||
for a in qy.saturating_sub(self.reach)..(qy + self.reach + 1).min(q) {
|
||
for b in qx.saturating_sub(self.reach)..(qx + self.reach + 1).min(q) {
|
||
let v = quad(a, b);
|
||
for c in 0..3 {
|
||
acc[c] += v[c];
|
||
}
|
||
cnt += 1.0;
|
||
}
|
||
}
|
||
for (dy, dx) in [(0, 0), (0, 1), (1, 0), (1, 1)] {
|
||
for c in 0..3 {
|
||
out[c * n * n + (2 * qy + dy) * n + 2 * qx + dx] = acc[c] / cnt;
|
||
}
|
||
}
|
||
}
|
||
}
|
||
write(&out);
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
/// The mosaic of a smooth colour field in `pattern`, read at (y, x).
|
||
fn field(pattern: CfaPattern) -> impl Fn(usize, usize) -> f32 {
|
||
move |y, x| {
|
||
let rgb = [0.2 + 0.0004 * x as f32, 0.5, 0.1 + 0.0003 * y as f32];
|
||
rgb[pattern.colour_at(x as u32, y as u32) as usize]
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn every_bayer_phase_comes_back_as_its_own_colours() {
|
||
// A frame of each pattern, its colours known: the network must see
|
||
// red where the frame's red photosites are, whatever the phase.
|
||
for p in [
|
||
CfaPattern::Rggb,
|
||
CfaPattern::Grbg,
|
||
CfaPattern::Gbrg,
|
||
CfaPattern::Bggr,
|
||
] {
|
||
let (h, w) = (300, 410);
|
||
let at = field(p);
|
||
let mut net = BoxNet {
|
||
n: 2 * HALO + 64,
|
||
reach: 0,
|
||
};
|
||
let out = run_tiled(&mut net, h, w, p, &at, &|_, _| 0.01, &mut |_, _| true)
|
||
.unwrap()
|
||
.unwrap();
|
||
for (y, x) in [(10, 10), (150, 201), (299, 409), (0, 0), (77, 333)] {
|
||
let o = &out[(y * w + x) * 3..(y * w + x) * 3 + 3];
|
||
let want = [0.2 + 0.0004 * x as f32, 0.5, 0.1 + 0.0003 * y as f32];
|
||
for c in 0..3 {
|
||
// Within the quad the binned value is at most a photosite away.
|
||
assert!(
|
||
(o[c] - want[c]).abs() < 0.0012,
|
||
"{p:?} at ({y},{x}) channel {c}: {} vs {}",
|
||
o[c],
|
||
want[c]
|
||
);
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn tiles_reproduce_one_pass_over_the_reflected_frame() {
|
||
// A network whose reach is inside the halo gives the same answer
|
||
// tiled small as in one tile covering everything.
|
||
let (h, w) = (230, 170);
|
||
for p in [CfaPattern::Rggb, CfaPattern::Bggr] {
|
||
let at = |y: usize, x: usize| ((y * 7919 + x * 104729) % 1000) as f32 / 1000.0;
|
||
let mut small = BoxNet {
|
||
n: 2 * HALO + 32,
|
||
reach: 20,
|
||
};
|
||
let mut big = BoxNet {
|
||
n: 2 * HALO + 256,
|
||
reach: 20,
|
||
};
|
||
let a = run_tiled(&mut small, h, w, p, &at, &|_, _| 0.0, &mut |_, _| true)
|
||
.unwrap()
|
||
.unwrap();
|
||
let b = run_tiled(&mut big, h, w, p, &at, &|_, _| 0.0, &mut |_, _| true)
|
||
.unwrap()
|
||
.unwrap();
|
||
let worst = a
|
||
.iter()
|
||
.zip(&b)
|
||
.map(|(x, y)| (x - y).abs())
|
||
.fold(0.0f32, f32::max);
|
||
assert!(worst < 1e-5, "{p:?}: tiled and whole differ by {worst}");
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn a_cancelled_run_returns_nothing() {
|
||
let mut net = BoxNet {
|
||
n: 2 * HALO + 32,
|
||
reach: 0,
|
||
};
|
||
let r = run_tiled(
|
||
&mut net,
|
||
100,
|
||
100,
|
||
CfaPattern::Rggb,
|
||
&|_, _| 0.5,
|
||
&|_, _| 0.0,
|
||
&mut |done, _| done < 2,
|
||
)
|
||
.unwrap();
|
||
assert!(r.is_none());
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod timing {
|
||
use super::*;
|
||
|
||
/// A network that answers instantly with an output of the right size,
|
||
/// so what is timed is the tiler alone: gathering each tile's mosaic and
|
||
/// σ, and writing its centre back.
|
||
struct Null(usize, Vec<f32>);
|
||
|
||
impl TileNet for Null {
|
||
fn tile(&self) -> usize {
|
||
self.0
|
||
}
|
||
fn run(
|
||
&mut self,
|
||
m: Vec<f32>,
|
||
_s: Vec<f32>,
|
||
write: &mut dyn FnMut(&[f32]),
|
||
) -> Result<(), crate::DenoiseError> {
|
||
// Stands for the runtime's own output buffer: allocated once.
|
||
if self.1.len() != 3 * m.len() {
|
||
self.1 = vec![m[0]; 3 * m.len()];
|
||
}
|
||
write(&self.1);
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
/// `cargo test --release -p dr-denoise tiler_overhead -- --ignored --nocapture`
|
||
#[test]
|
||
#[ignore]
|
||
fn tiler_overhead_on_a_6d_frame() {
|
||
let (h, w) = (3648, 5472);
|
||
let frame: Vec<f32> = (0..h * w).map(|i| (i % 977) as f32 / 977.0).collect();
|
||
let at = |y: usize, x: usize| frame[y * w + x];
|
||
let sigma = |_c: usize, v: f32| (0.001 * v + 1e-5).sqrt();
|
||
for n in [1408usize, 2048] {
|
||
let mut net = Null(n, Vec::new());
|
||
let t = std::time::Instant::now();
|
||
let mut tiles = 0;
|
||
run_tiled(
|
||
&mut net,
|
||
h,
|
||
w,
|
||
CfaPattern::Rggb,
|
||
&at,
|
||
&sigma,
|
||
&mut |_, total| {
|
||
tiles = total;
|
||
true
|
||
},
|
||
)
|
||
.unwrap();
|
||
let s = t.elapsed().as_secs_f64();
|
||
println!(
|
||
"tile {n}: {tiles} tiles, tiler alone {s:.2} s ({:.0} ms a tile)",
|
||
s / tiles as f64 * 1e3
|
||
);
|
||
}
|
||
}
|
||
}
|