diff --git a/core/dr-denoise/src/onnx.rs b/core/dr-denoise/src/onnx.rs index ed9a20d..5b47052 100644 --- a/core/dr-denoise/src/onnx.rs +++ b/core/dr-denoise/src/onnx.rs @@ -44,15 +44,21 @@ impl TileNet for OnnxNet { self.tile } - fn run(&mut self, mosaic: &[f32], sigma: &[f32]) -> Result, DenoiseError> { + fn run( + &mut self, + mosaic: Vec, + sigma: Vec, + write: &mut dyn FnMut(&[f32]), + ) -> Result<(), DenoiseError> { let n = self.tile; let shape = ndarray::IxDyn(&[1, 1, n, n]); + // The vectors become the tensors: no copy on the way in. let m = ort::value::Tensor::from_array( - ndarray::Array::from_shape_vec(shape.clone(), mosaic.to_vec()) + ndarray::Array::from_shape_vec(shape.clone(), mosaic) .map_err(|e| DenoiseError::Model(e.to_string()))?, )?; let s = ort::value::Tensor::from_array( - ndarray::Array::from_shape_vec(shape, sigma.to_vec()) + ndarray::Array::from_shape_vec(shape, sigma) .map_err(|e| DenoiseError::Model(e.to_string()))?, )?; let acquired = self.model.acquire()?; @@ -65,6 +71,8 @@ impl TileNet for OnnxNet { "output is {dims:?}, expected [1, 3, {n}, {n}]" ))); } - Ok(data.to_vec()) + // And none on the way out: the frame is written from the runtime's buffer. + write(data); + Ok(()) } } diff --git a/core/dr-denoise/src/tile.rs b/core/dr-denoise/src/tile.rs index cc14ad0..7057bd2 100644 --- a/core/dr-denoise/src/tile.rs +++ b/core/dr-denoise/src/tile.rs @@ -20,10 +20,20 @@ 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; - fn run(&mut self, mosaic: &[f32], sigma: &[f32]) -> Result, crate::DenoiseError>; + 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 @@ -61,8 +71,8 @@ pub fn run_tiled( h: usize, w: usize, pattern: CfaPattern, - at: &dyn Fn(usize, usize) -> f32, - sigma: &dyn Fn(usize, f32) -> f32, + 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(|| { @@ -79,58 +89,120 @@ pub fn run_tiled( 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 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 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; } - 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 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); + } } - if !progress(k + 1, total) { - return Ok(None); - } - } - Ok(Some(out)) + Ok(Some(())) + }) + .map(|done| done.map(|()| out)) } #[cfg(test)] @@ -165,7 +237,12 @@ mod tests { fn tile(&self) -> usize { self.n } - fn run(&mut self, m: &[f32], _s: &[f32]) -> Result, crate::DenoiseError> { + 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| { @@ -197,7 +274,8 @@ mod tests { } } } - Ok(out) + write(&out); + Ok(()) } } @@ -293,3 +371,65 @@ mod tests { 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 + ); + } + } +}