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.
730 lines
26 KiB
Rust
730 lines
26 KiB
Rust
//! 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
|
||
);
|
||
}
|
||
}
|
||
}
|