Files
DarkRoom/core/dr-denoise/src/tile.rs
T
dtourolle 1e1aa1442b Size the whole-frame profile for a 6 GB card
TensorRT plans its memory for the profile's largest shape, and up to a
whole 6D frame with Best's border (4608 x 6656) it asked for 4.9-5.9 GB
and would not build on the RTX 3050. The profile now ends at 4608 x 3328
(15 MP), tuned for 4160 x 3248, and the tiler cuts a 6D frame into two
such tiles: 27 MP of work for 20 MP kept, against 49 MP in 1408 tiles.
The engine's directory names the profile, so a later range never loads
an engine built for this one.
2026-10-06 23:14:10 -04:00

730 lines
26 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 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<f32>,
sigma: Vec<f32>,
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<Plan> {
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<Plan> = 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<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 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<Option<Vec<f32>>, 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<Option<()>, 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<f32>,
_s: Vec<f32>,
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<f32>);
impl TileNet for Null {
fn sizes(&self) -> Sizes {
Sizes::Square(self.0)
}
fn run(
&mut self,
_rows: usize,
_cols: usize,
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
);
}
}
}