//! 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, sigma: Vec, 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>, 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, 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, _s: Vec, 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); impl TileNet for Null { fn tile(&self) -> usize { self.0 } fn run( &mut self, m: Vec, _s: Vec, 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 = (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 ); } } }