//! TRACES: FR-DEV-3g //! A whole frame through a network, in tiles, exactly (denoise.md §3.4, §14). //! //! A tile's output is exact in its centre: past a halo wider than the //! network's receptive field (185 photosites for a single network, more for //! the mixture), 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. //! //! **Tile sizes.** A fixed-shape network takes one square ([`Sizes::Square`], //! 1408², of which Best keeps 896²). A network exported with any height and //! width ([`Sizes::Any`]) takes the frame whole when it is small enough, and //! otherwise the fewest equal tiles that are: [`plan`] picks the grid that //! computes the fewest photosites. If the first tile of a plan fails — a //! GPU out of memory — the limit is halved and the frame planned again. //! //! **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; /// The tiles a network takes. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Sizes { /// One square, `n` photosites a side. Square(usize), /// Any rectangle whose sides are multiples of `align`, at most `max` /// (rows, columns). Any { align: usize, max: (usize, usize) }, } /// A network: `mosaic` and `sigma`, `rows×cols` RGGB, in; `3×rows×cols` /// 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 a whole frame 300 MB, 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 tile sizes it takes. fn sizes(&self) -> Sizes; /// 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, rows: usize, cols: usize, mosaic: Vec, sigma: Vec, write: &mut dyn FnMut(&[f32]), ) -> Result<(), crate::DenoiseError>; } /// How a frame is cut: every tile `rows × cols` in, keeping its centre /// `core.0 × core.1` past the halo, on a `grid.0 × grid.1` grid. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct Plan { pub rows: usize, pub cols: usize, pub core: (usize, usize), pub grid: (usize, usize), } impl Plan { /// Photosites the network computes for the frame. pub fn work(&self) -> usize { self.grid.0 * self.grid.1 * self.rows * self.cols } } /// The tiles for an `uh × uw` frame (in the network's phase) at `sizes`, or /// `None` when no tile fits. /// /// Square tiles are today's grid. Any-size tiles are equal on each axis, so /// one call shape serves the frame — TensorRT's profile tunes for one, and /// the CUDA provider searches its algorithms once per shape — and the grid /// is the one with the least work: one tile whenever the frame and its /// halo fit under `max`. pub fn plan(uh: usize, uw: usize, halo: usize, sizes: Sizes) -> Option { match sizes { Sizes::Square(n) => { if n <= 2 * halo || !(n - 2 * halo).is_multiple_of(2) { return None; } let core = n - 2 * halo; Some(Plan { rows: n, cols: n, core: (core, core), grid: (uh.div_ceil(core), uw.div_ceil(core)), }) } Sizes::Any { align, max } => { // An even align keeps every tile origin on an even photosite, // so every tile starts on red. let align = align.max(2).next_multiple_of(2); let axis = |extent: usize, tiles: usize, limit: usize| { let size = (extent.div_ceil(tiles) + 2 * halo).next_multiple_of(align); let core = size.checked_sub(2 * halo)?; (size <= limit && core > 0 && core.is_multiple_of(2)).then_some((size, core)) }; let mut best: Option = None; for gy in 1..=16 { let Some((rows, cy)) = axis(uh, gy, max.0) else { continue; }; for gx in 1..=16 { let Some((cols, cx)) = axis(uw, gx, max.1) else { continue; }; let p = Plan { rows, cols, core: (cy, cx), grid: (uh.div_ceil(cy), uw.div_ceil(cx)), }; if best.is_none_or(|b| p.work() < b.work()) { best = Some(p); } } } best } } } /// 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)`. An any-size /// network whose first tile fails is planned again with tiles half that /// size, until a tile would keep no centre; then the failure is returned. #[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 halo = net.halo(); let (uh, uw) = (h + dy, w + dx); let mut sizes = net.sizes(); loop { let plan = plan(uh, uw, halo, sizes).ok_or_else(|| { crate::DenoiseError::Model(format!( "no tile of {sizes:?} keeps a centre past a {halo} halo" )) })?; match run_plan(net, plan, h, w, (dy, dx), halo, at, sigma, progress) { Err(Failed { error, first: true }) => { // The first call of a size is where a GPU runs out of // memory. Halve the larger kept centre of the tile that // failed — not the limit, which may be far above it, and not // the tile, half of which may be all halo — and plan again, // until no smaller tile keeps a centre. let Sizes::Any { align, .. } = sizes else { return Err(error); }; let (cr, cc) = plan.core; let smaller = if cr >= cc { (cr / 2 + 2 * halo, plan.cols) } else { (plan.rows, cc / 2 + 2 * halo) }; let next = Sizes::Any { align, max: smaller, }; if self::plan(uh, uw, halo, next).is_none() { return Err(error); } log::warn!( "learned denoise: a {}×{} tile failed ({error}); trying tiles up to {}×{}", plan.rows, plan.cols, smaller.0, smaller.1 ); sizes = next; } Err(Failed { error, .. }) => return Err(error), Ok(done) => return Ok(done), } } } /// A run that stopped on an error, and whether it was the plan's first call. struct Failed { error: crate::DenoiseError, first: bool, } #[allow(clippy::too_many_arguments)] fn run_plan( net: &mut dyn TileNet, plan: Plan, h: usize, w: usize, (dy, dx): (usize, usize), halo: usize, at: &(dyn Fn(usize, usize) -> f32 + Sync), sigma: &(dyn Fn(usize, f32) -> f32 + Sync), progress: &mut dyn FnMut(usize, usize) -> bool, ) -> Result>, Failed> { let Plan { rows: nr, cols: nc, core: (cr, cc), grid: (ty, tx), } = plan; // In unified coordinates the frame spans u ∈ [dy, dy + h), v ∈ [dx, dx + w). let (uh, uw) = (h + dy, w + dx); let total = ty * tx; let origins: Vec<(usize, usize)> = (0..ty) .flat_map(|i| (0..tx).map(move |j| (i * cr, j * cc))) .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; nr * nc]; let mut sig = vec![0.0f32; nr * nc]; let rows_per = nr.div_ceil(threads).max(1); std::thread::scope(|scope| { for (chunk, (m, s)) in mos .chunks_mut(rows_per * nc) .zip(sig.chunks_mut(rows_per * nc)) .enumerate() { scope.spawn(move || { for (i, (mrow, srow)) in m.chunks_mut(nc).zip(s.chunks_mut(nc)).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..nc { 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); let fail = |error, k: usize, stop: &std::sync::atomic::AtomicBool| { stop.store(true, std::sync::atomic::Ordering::Relaxed); Failed { error, first: k == 0, } }; std::thread::scope(|scope| -> Result, Failed> { 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(nr, nc, mos, sig, &mut |rgb: &[f32]| { if rgb.len() != 3 * nr * nc { 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.max(dy) - dy, (u0 + cr).min(uh) - dy); let (x_lo, x_hi) = (v0.max(dx) - dx, (v0 + cc).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 * nr * nc + r * nc + c]; } } } }); } }); } }); if let Err(e) = ran { while rx_tiles.try_recv().is_ok() {} return Err(fail(e, k, stop)); } if let Some(len) = wrong { while rx_tiles.try_recv().is_ok() {} return Err(fail( crate::DenoiseError::Model(format!( "network returned {len} values for a {nr}×{nc} tile" )), k, stop, )); } 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 { sizes: Sizes, reach: usize, /// Fails any call with more photosites than this, as a GPU out of /// memory does. fails_above: usize, calls: Vec<(usize, usize)>, } fn square(n: usize, reach: usize) -> BoxNet { BoxNet { sizes: Sizes::Square(n), reach, fails_above: usize::MAX, calls: Vec::new(), } } fn any(max: (usize, usize), reach: usize) -> BoxNet { BoxNet { sizes: Sizes::Any { align: 16, max }, reach, fails_above: usize::MAX, calls: Vec::new(), } } impl TileNet for BoxNet { fn sizes(&self) -> Sizes { self.sizes } fn run( &mut self, rows: usize, cols: usize, m: Vec, _s: Vec, write: &mut dyn FnMut(&[f32]), ) -> Result<(), crate::DenoiseError> { self.calls.push((rows, cols)); if rows * cols > self.fails_above { return Err(crate::DenoiseError::Model("out of memory".into())); } let (qr, qc) = (rows / 2, cols / 2); let quad = |qy: usize, qx: usize| { let (y, x) = (2 * qy, 2 * qx); [ m[y * cols + x], 0.5 * (m[y * cols + x + 1] + m[(y + 1) * cols + x]), m[(y + 1) * cols + x + 1], ] }; let plane = rows * cols; let mut out = vec![0.0; 3 * plane]; for qy in 0..qr { for qx in 0..qc { let mut acc = [0.0f32; 3]; let mut cnt = 0.0; for a in qy.saturating_sub(self.reach)..(qy + self.reach + 1).min(qr) { for b in qx.saturating_sub(self.reach)..(qx + self.reach + 1).min(qc) { 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 * plane + (2 * qy + dy) * cols + 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 = square(2 * HALO + 64, 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 = square(2 * HALO + 32, 20); let mut big = square(2 * HALO + 256, 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}"); } } /// The same frame through square tiles, one whole-frame call, a grid of /// any-size tiles, and a network that runs out of memory on the whole /// frame and is planned again: one answer. #[test] fn any_size_tiles_give_the_square_tiles_answer() { let (h, w) = (230, 170); for p in [ CfaPattern::Rggb, CfaPattern::Grbg, CfaPattern::Gbrg, CfaPattern::Bggr, ] { let at = |y: usize, x: usize| ((y * 7919 + x * 104729) % 1000) as f32 / 1000.0; let run = |net: &mut BoxNet| { run_tiled(net, h, w, p, &at, &|_, _| 0.0, &mut |_, _| true) .unwrap() .unwrap() }; // A reach of 6 quads is well inside the halo, and keeps a // debug-build test of four phases short. let want = run(&mut square(2 * HALO + 32, 6)); let mut whole = any((4096, 4096), 6); let got = run(&mut whole); assert_eq!(whole.calls.len(), 1, "the frame fits: one call"); assert_eq!(got, want, "{p:?}: whole frame"); let mut grid = any((2 * HALO + 96, 2 * HALO + 64), 6); let got = run(&mut grid); assert!(grid.calls.len() > 1); assert!( grid.calls.windows(2).all(|c| c[0] == c[1]), "one call shape" ); assert_eq!(got, want, "{p:?}: a grid of any-size tiles"); let mut tight = any((4096, 4096), 6); tight.fails_above = (2 * HALO + 200) * (2 * HALO + 200); let got = run(&mut tight); assert_eq!(got, want, "{p:?}: planned again after a failure"); assert!(tight.calls.len() > 2, "the whole frame failed, then tiles"); } } #[test] fn the_plan_is_one_tile_when_the_frame_fits_and_the_least_work_when_not() { // A 6D frame with Best's halo, under the whole-frame limit: one call. let one = plan( 3648, 5472, 256, Sizes::Any { align: 16, max: (4608, 6656), }, ) .unwrap(); assert_eq!(one.grid, (1, 1)); assert_eq!((one.rows, one.cols), (4160, 5984)); assert!(one.core.0 >= 3648 && one.core.1 >= 5472); // Too wide for one: the cheapest grid, every tile within the limit. let two = plan( 3648, 8192, 256, Sizes::Any { align: 16, max: (4608, 6656), }, ) .unwrap(); assert!(two.cols <= 6656 && two.rows <= 4608); assert_eq!(two.grid, (1, 2)); // And always less work than today's 1408 squares. let squares = plan(3648, 5472, 256, Sizes::Square(1408)).unwrap(); assert_eq!(squares.grid, (5, 7)); assert!(one.work() * 2 < squares.work()); // The whole-frame engine's limit on a 6 GB card: two tiles, each // within it, and still under half the work of the 1408 squares. let halves = plan( 3648, 5472, 256, Sizes::Any { align: 16, max: (4608, 3328), }, ) .unwrap(); assert_eq!(halves.grid, (1, 2)); assert_eq!((halves.rows, halves.cols), (4160, 3248)); assert!(halves.work() * 2 < squares.work()); // A limit no tile fits under. assert!(plan( 3648, 5472, 256, Sizes::Any { align: 16, max: (400, 400) } ) .is_none()); } #[test] fn a_cancelled_run_returns_nothing() { let mut net = square(2 * HALO + 32, 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 sizes(&self) -> Sizes { Sizes::Square(self.0) } fn run( &mut self, _rows: usize, _cols: usize, 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 ); } } }