//! 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. pub const HALO: usize = 192; /// A fixed-shape network: `mosaic` and `sigma`, `n×n` RGGB, in; `3×n×n` /// planar linear camera RGB out. pub trait TileNet { /// The edge `n` of the square tile the network takes. fn tile(&self) -> usize; fn run(&mut self, mosaic: &[f32], sigma: &[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, sigma: &dyn Fn(usize, f32) -> f32, 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 = net.tile(); 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 mut out = vec![0.0f32; h * w * 3]; let mut mos = vec![0.0f32; n * n]; let mut sig = vec![0.0f32; n * n]; // RGGB colour of unified position (u, v). let colour = |u: usize, v: usize| [[0, 1], [1, 2]][u & 1][v & 1]; for (k, (i, j)) in (0..ty) .flat_map(|i| (0..tx).map(move |j| (i, j))) .enumerate() { let (u0, v0) = (i * core, j * core); for r in 0..n { // 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); mos[r * n + c] = val; sig[r * n + c] = sigma(colour(r, c), val); } } let rgb = net.run(&mos, &sig)?; if rgb.len() != 3 * n * n { return Err(crate::DenoiseError::Model(format!( "network returned {} values for a {n}² tile", rgb.len() ))); } for r in HALO..HALO + core { let u = u0 + r - HALO; if u < dy || u >= uh { continue; } let y = u - dy; for c in HALO..HALO + core { let v = v0 + c - HALO; if v < dx || v >= uw { continue; } let x = v - dx; let o = (y * w + x) * 3; for ch in 0..3 { out[o + ch] = rgb[ch * n * n + r * n + c]; } } } if !progress(k + 1, total) { return Ok(None); } } Ok(Some(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: &[f32], _s: &[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; } } } } Ok(out) } } /// 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()); } }