Denoise a whole frame in one call where the GPU takes any size

A fixed 1408 tile is exact only in its centre, and Best keeps 896 of
every 1408 it computes: 2.47 photosites of work for each one kept. The
tiler now takes a network of any size as well as a square one, and
plans the frame as the fewest equal tiles under the rung's limit --
one tile, the whole frame and its reflected border, whenever it fits.
If the first call of a plan fails, as a GPU out of memory does, the
kept centre is halved and the frame planned again.

Each shipped network names its any-size sibling (mosaic-best.onnx
beside mosaic-best-1408.onnx). OnnxNet::open takes it where the engine
runs whole frames and the file is installed, and the 1408 tiles
otherwise; open_tiled forces the tiles, and denoise_raw's DR_PLAN=tiles
uses it to compare. The cache key stays on the fixed model: the output
is the same network's. Tests hold any-size tiles, a grid of them and a
plan rebuilt after a failure to the square tiles' answer in every Bayer
phase.
This commit is contained in:
2026-10-06 21:47:19 -04:00
parent 56f4180347
commit 9cba420fd5
6 changed files with 461 additions and 108 deletions
+348 -76
View File
@@ -1,12 +1,18 @@
//! TRACES: FR-DEV-3g
//! A whole frame through a fixed-shape network, exactly (denoise.md §3.4).
//! A whole frame through a network, in tiles, exactly (denoise.md §3.4, §14).
//!
//! 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.
//! 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
@@ -20,16 +26,26 @@ use dr_decode::CfaPattern;
/// [`TileNet::halo`].
pub const HALO: usize = 192;
/// A fixed-shape network: `mosaic` and `sigma`, `n×n` RGGB, in; `3×n×n`
/// 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 copying it out of the
/// runtime's buffer and back into the frame was a measurable share of a
/// frame's time.
/// 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 edge `n` of the square tile the network takes.
fn tile(&self) -> usize;
/// 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 {
@@ -37,12 +53,87 @@ pub trait TileNet {
}
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.
@@ -71,7 +162,9 @@ pub fn rggb_offset(p: CfaPattern) -> Option<(usize, usize)> {
/// `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)`.
/// 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,
@@ -85,40 +178,103 @@ pub fn run_tiled(
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 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),
}
}
let core = n - 2 * halo;
}
/// 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 (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)))
.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; n * n];
let mut sig = vec![0.0f32; n * n];
let rows_per = n.div_ceil(threads).max(1);
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 * n)
.zip(sig.chunks_mut(rows_per * n))
.chunks_mut(rows_per * nc)
.zip(sig.chunks_mut(rows_per * nc))
.enumerate()
{
scope.spawn(move || {
for (i, (mrow, srow)) in m.chunks_mut(n).zip(s.chunks_mut(n)).enumerate() {
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..n {
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);
@@ -138,7 +294,14 @@ pub fn run_tiled(
// 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<Option<()>, crate::DenoiseError> {
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 || {
@@ -156,18 +319,15 @@ pub fn run_tiled(
break;
};
let mut wrong = None;
let ran = net.run(mos, sig, &mut |rgb: &[f32]| {
if rgb.len() != 3 * n * n {
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 + 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);
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);
@@ -181,7 +341,7 @@ pub fn run_tiled(
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];
row[x * 3 + ch] = rgb[ch * nr * nc + r * nc + c];
}
}
}
@@ -191,14 +351,18 @@ pub fn run_tiled(
}
});
if let Err(e) = ran {
stop.store(true, std::sync::atomic::Ordering::Relaxed);
return Err(e);
while rx_tiles.try_recv().is_ok() {}
return Err(fail(e, k, stop));
}
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"
)));
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);
@@ -236,37 +400,65 @@ mod tests {
/// 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,
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 tile(&self) -> usize {
self.n
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> {
let n = self.n;
let q = n / 2;
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 * n + x],
0.5 * (m[y * n + x + 1] + m[(y + 1) * n + x]),
m[(y + 1) * n + x + 1],
m[y * cols + x],
0.5 * (m[y * cols + x + 1] + m[(y + 1) * cols + x]),
m[(y + 1) * cols + x + 1],
]
};
let mut out = vec![0.0; 3 * n * n];
for qy in 0..q {
for qx in 0..q {
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(q) {
for b in qx.saturating_sub(self.reach)..(qx + self.reach + 1).min(q) {
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];
@@ -276,7 +468,7 @@ mod tests {
}
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;
out[c * plane + (2 * qy + dy) * cols + 2 * qx + dx] = acc[c] / cnt;
}
}
}
@@ -306,10 +498,7 @@ mod tests {
] {
let (h, w) = (300, 410);
let at = field(p);
let mut net = BoxNet {
n: 2 * HALO + 64,
reach: 0,
};
let mut net = square(2 * HALO + 64, 0);
let out = run_tiled(&mut net, h, w, p, &at, &|_, _| 0.01, &mut |_, _| true)
.unwrap()
.unwrap();
@@ -336,14 +525,8 @@ mod tests {
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 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();
@@ -359,12 +542,99 @@ mod tests {
}
}
/// 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());
// 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 = BoxNet {
n: 2 * HALO + 32,
reach: 0,
};
let mut net = square(2 * HALO + 32, 0);
let r = run_tiled(
&mut net,
100,
@@ -389,11 +659,13 @@ mod timing {
struct Null(usize, Vec<f32>);
impl TileNet for Null {
fn tile(&self) -> usize {
self.0
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]),