diff --git a/.idea/av-denoise.iml b/.idea/av-denoise.iml index b294e21..c17e1b4 100644 --- a/.idea/av-denoise.iml +++ b/.idea/av-denoise.iml @@ -12,10 +12,13 @@ + + + diff --git a/Cargo.lock b/Cargo.lock index d516805..db6018b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -187,19 +187,25 @@ dependencies = [ "av-decoders", "av-denoise-core", "av-scenechange", + "bon", "bytesize", "clap", "crossbeam-channel", + "cubecl", + "etcetera", "ffms2-sys", "indicatif", "mimalloc", "pkg-config", "strum", "strum_macros", + "tempfile", + "thiserror", "tracing", "tracing-indicatif", "tracing-subscriber", "v_frame", + "wgpu", "y4m", ] @@ -211,11 +217,9 @@ dependencies = [ "bon", "clap", "cubecl", - "etcetera", "futures", "strum", "strum_macros", - "tempfile", "thiserror", "tracing", ] @@ -225,7 +229,7 @@ name = "av-denoise-vs" version = "0.5.0-alpha3" dependencies = [ "anyhow", - "av-denoise-core", + "av-denoise", "tracing", "tracing-subscriber", "vapoursynth", diff --git a/Cargo.toml b/Cargo.toml index 5fa9ade..34f5ec9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -20,6 +20,7 @@ strum_macros = "0.28" bon = "3" cubecl = { version = "0.10", features = ["std"] } av-denoise-core = { path = "av-denoise-core", version = "0.5.0-alpha3", default-features = false } +av-denoise = { path = "av-denoise", version = "0.5.0-alpha3", default-features = false } [profile.test] opt-level = 3 diff --git a/Justfile b/Justfile index 0a99e14..fad4048 100755 --- a/Justfile +++ b/Justfile @@ -66,6 +66,10 @@ build-static features=static_features: bench *ARGS: cargo bench -p av-denoise-core {{ARGS}} +# Benchmarks the host layer, end to end through `HostDenoiser` and `PlanarDenoiser`. +bench-host *ARGS: + cargo bench -p av-denoise --features vulkan {{ARGS}} + build-vs *ARGS: cargo build -p av-denoise-vs --release {{ARGS}} @@ -84,6 +88,7 @@ test-rust: cargo nextest run --release -p av-denoise --features vulkan,binary cargo nextest run --release -p av-denoise-vs --features vulkan cargo test --doc -p av-denoise-core --features vulkan + cargo test --doc -p av-denoise --features vulkan cargo check --workspace # Every Python test, including the ones that render on the GPU. diff --git a/README.md b/README.md index bfd49d9..1eea8df 100644 --- a/README.md +++ b/README.md @@ -171,8 +171,8 @@ without use. If the cache directory cannot be created, `av-denoise` logs a warning and carries on without a cache. -Library users can call `av_denoise::install_compilation_cache()` before `Denoiser::create` to get the same -behaviour in their own binary. It has to run before the first `Denoiser` exists, because building a CubeCL client +Library users can call `av_denoise::install_compilation_cache()` before `HostDenoiser::create` to get the same +behaviour in their own binary. It has to run before the first denoiser exists, because building a CubeCL client locks the global config. An embedder that wants to choose the cache directory itself can call `av_denoise::default_cache_dir()` to get the same default this crate uses, and `av_denoise::install_compilation_cache_at()` to install it, or any other directory, directly. diff --git a/av-denoise-core/Cargo.toml b/av-denoise-core/Cargo.toml index cc3129f..bf73faa 100644 --- a/av-denoise-core/Cargo.toml +++ b/av-denoise-core/Cargo.toml @@ -4,7 +4,7 @@ version.workspace = true edition.workspace = true license.workspace = true repository.workspace = true -description = "Core kernels and types for av-denoise (Do not use directly)" +description = "GPU denoising engines that read and write CubeCL handles" categories = ["multimedia::video"] keywords = ["cubecl", "denoise", "gpu"] readme = "README.md" @@ -28,14 +28,12 @@ tracing = "0.1" strum_macros = "0.28" strum = "0.28" bon = "3" -etcetera = "0.11" cubecl = { version = "0.10", features = ["std"] } [dev-dependencies] futures = "0.3" clap = { version = "4", features = ["derive"] } -tempfile = "3" [features] vulkan = ["cubecl/vulkan"] @@ -52,26 +50,14 @@ harness = false name = "bench_kernels" harness = false -[[bench]] -name = "denoise" -harness = false - [[bench]] name = "motion" harness = false -[[bench]] -name = "convert" -harness = false - [[bench]] name = "nl4d_ablation" harness = false -[[bench]] -name = "reseed" -harness = false - [[bench]] name = "mc_accuracy" harness = false diff --git a/av-denoise-core/README.md b/av-denoise-core/README.md index d3873e4..18c76c2 100644 --- a/av-denoise-core/README.md +++ b/av-denoise-core/README.md @@ -1,5 +1,44 @@ # av-denoise-core -The core library and kernels for the av-denoise system. +GPU denoising engines that read and write CubeCL plane handles. -You should use either the `av-denoise` library, binary or VapourSynth plugin, not this crate directly. +`Nlmeans` and `Nl4d` implement `Engine`. Push one frame of planes at a time. Each push returns how many +frames are ready, and every ready frame is written into planes you own with `Engine::emit_into` before the +next push. At the end of a stream, `Engine::finish` returns the tail count. + +`Nl4dOptions::default()` uses the base preset's `temporal_radius` and calibrated defaults for the rest. + +```rust,no_run +use av_denoise_core::{ChannelMode, DevicePlane, Engine, Geometry, Nl4d, Nl4dOptions, SampleFormat}; +use cubecl::Runtime; +use cubecl::wgpu::WgpuRuntime; + +# fn main() -> Result<(), av_denoise_core::Error> { +let device = ::Device::default(); +let client = WgpuRuntime::client(&device); +let geometry = Geometry { + width: 1920, + height: 1080, + channels: ChannelMode::Luma, + input: SampleFormat::U16 { depth: 10 }, + output: SampleFormat::U16 { depth: 10 }, +}; +let mut engine = Nl4d::new(&client, Nl4dOptions::default(), geometry)?; + +let input = client.empty(1920 * 1080 * 2); +let output = client.empty(1920 * 1080 * 2); +let input_planes = [DevicePlane::new(&input, 1920, 1080)]; +let output_planes = [DevicePlane::new(&output, 1920, 1080)]; + +let ready = engine.push(&input_planes)?; +for _ in 0..ready { + engine.emit_into(&output_planes)?; +} + +let tail = engine.finish()?; +for _ in 0..tail { + engine.emit_into(&output_planes)?; +} +# Ok(()) +# } +``` diff --git a/av-denoise-core/benches/bench_kernels.rs b/av-denoise-core/benches/bench_kernels.rs index 0c4b131..44bf077 100644 --- a/av-denoise-core/benches/bench_kernels.rs +++ b/av-denoise-core/benches/bench_kernels.rs @@ -1,8 +1,8 @@ -use av_denoise_core::Depth; -use cubecl::prelude::*; - mod kernels; +use av_denoise_core::bench_api::Device; +use clap::Parser; +use cubecl::prelude::*; use kernels::accumulate::AccumulateBench; use kernels::bilateral::BilateralBench; use kernels::collab_aggregate::{CollabNormaliseBench, CollabZeroAccumBench}; @@ -14,6 +14,7 @@ use kernels::distance::DistanceBench; use kernels::distance_pair::DistancePairBench; use kernels::distance_pair_ref::DistancePairRefBench; use kernels::distance_ref::DistanceRefBench; +use kernels::egress::{EgressBench, EgressFormat}; use kernels::finish::FinishBench; use kernels::fused_window::{ FusedPairWindowBench, @@ -24,6 +25,7 @@ use kernels::fused_window::{ use kernels::grain::{GRAIN_SIZES, GrainMeasureBench, GrainReducePartialsBench, GrainSaveVectorsBench}; use kernels::horizontal_sum::HSumBench; use kernels::horizontal_sum_pair::HSumPairBench; +use kernels::ingest::{IngestBench, IngestFormat}; use kernels::mc_block_match_coarse::BlockMatchCoarseBench; use kernels::mc_block_match_fine::BlockMatchFineBench; use kernels::mc_chain_compose::ChainComposeBench; @@ -32,9 +34,7 @@ use kernels::mc_downscale::DownscaleBench; use kernels::mc_warp::WarpBench; use kernels::mv_regularise::MvRegulariseBench; use kernels::noise_partial::NoisePartialBench; -use kernels::pack_wire::PackWireBench; use kernels::temporal_noise_stats::TemporalNoiseStatsBench; -use kernels::unpack_wire::UnpackWireBench; use kernels::vertical_weight::VWeightBench; use kernels::vweight_pair_accumulate::VWeightPairAccBench; use kernels::zero::ZeroBench; @@ -47,110 +47,133 @@ fn run_all(backend: &str, device: &R::Device) { println!("--- {backend} ---"); print_header(); - for &(ch, ch_name) in CHANNELS { + for &(channels, channel_name) in CHANNELS { run(CopyBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { - for depth in [Depth::Eight, Depth::Ten] { - run(PackWireBench { - client: client.clone(), - ch, - ch_name, - depth, - }); - } + + let ingest_cases = [ + (1920, 1080, 1, "luma", IngestFormat::U8), + (960, 540, 2, "chroma", IngestFormat::U16Ten), + (1920, 1080, 3, "yuv", IngestFormat::U16Ten), + (1920, 1080, 1, "luma", IngestFormat::F32), + ]; + for (width, height, channels, channel_name, format) in ingest_cases { + run(IngestBench { + client: client.clone(), + width, + height, + channels, + channel_name, + format, + }); } - for &(ch, ch_name) in CHANNELS { - for depth in [Depth::Eight, Depth::Ten] { - run(UnpackWireBench { - client: client.clone(), - ch, - ch_name, - depth, - }); - } + + let egress_cases = [ + (1920, 1080, 1, "luma", EgressFormat::U8), + (960, 540, 2, "chroma", EgressFormat::U16Ten), + (1920, 1080, 3, "yuv", EgressFormat::U16Ten), + (1920, 1080, 1, "luma", EgressFormat::F32), + ]; + for (width, height, channels, channel_name, format) in egress_cases { + run(EgressBench { + client: client.clone(), + width, + height, + channels, + channel_name, + format, + }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(ZeroBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + for &(channels, channel_name) in CHANNELS { run(DistWeightBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(DistWeightRefBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(FusedSingleWindowBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(FusedPairWindowBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(FusedSingleWindowRefBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(FusedPairWindowRefBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + for &(channels, channel_name) in CHANNELS { run(DistanceBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(DistanceRefBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(DistancePairBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(DistancePairRefBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } @@ -163,45 +186,44 @@ fn run_all(backend: &str, device: &R::Device) { run(VWeightBench { client: client.clone(), }); - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(VWeightPairAccBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + for &(channels, channel_name) in CHANNELS { run(AccumulateBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(FinishBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + for &(channels, channel_name) in CHANNELS { run(BilateralBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - // Immerkær noise estimate (stage 1 + stage 2 back-to-back). run(NoisePartialBench { client: client.clone(), }); - // Temporal-residual noise-stats kernel, the default `luma_fields` - // off variant every caller but nl4d's luma pass uses, and the on - // variant that pass will enable. + // `luma_fields` off as nlmeans runs it, and on as nl4d runs it with the noise map. run(TemporalNoiseStatsBench { client: client.clone(), luma_fields: false, @@ -211,9 +233,8 @@ fn run_all(backend: &str, device: &R::Device) { luma_fields: true, }); - // Motion-compensation kernels. Pyramid build and analyse are - // luma-only (ME doesn't look at chroma); warp runs per channel - // mode because its memory traffic scales with `stored_ch`. + // Pyramid build and block matching are luma-only since motion estimation ignores chroma. Warp + // runs per channel mode because its memory traffic scales with `stored_ch`. run(DownscaleBench { client: client.clone(), }); @@ -232,6 +253,7 @@ fn run_all(backend: &str, device: &R::Device) { run(MvRegulariseBench { client: client.clone(), }); + for &size in GRAIN_SIZES { run(GrainSaveVectorsBench { client: client.clone(), @@ -246,21 +268,23 @@ fn run_all(backend: &str, device: &R::Device) { size, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(CollabFusedBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, split_mv: false, noise_curve: false, strength_map: false, pooled: false, }); } + run(CollabFusedBench { client: client.clone(), - ch: 1, - ch_name: "luma", + channels: 1, + channel_name: "luma", split_mv: false, noise_curve: true, strength_map: false, @@ -268,8 +292,8 @@ fn run_all(backend: &str, device: &R::Device) { }); run(CollabFusedBench { client: client.clone(), - ch: 1, - ch_name: "luma", + channels: 1, + channel_name: "luma", split_mv: false, noise_curve: true, strength_map: true, @@ -277,43 +301,47 @@ fn run_all(backend: &str, device: &R::Device) { }); run(CollabFusedBench { client: client.clone(), - ch: 1, - ch_name: "luma", + channels: 1, + channel_name: "luma", split_mv: false, noise_curve: true, strength_map: true, pooled: true, }); - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(CollabFusedBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, split_mv: true, noise_curve: false, strength_map: false, pooled: false, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(CollabNormaliseBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(CollabZeroAccumBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(WarpBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } @@ -323,18 +351,16 @@ fn run_all(backend: &str, device: &R::Device) { #[derive(clap::Parser, Debug)] #[command(about = "NLMeans per-kernel benchmarks", long_about = None)] struct Cli { - /// GPU device to bind to. Format: `default`, `discrete[:N]`, - /// `integrated[:N]`, `virtual[:N]`, or `cpu`. + /// GPU device to bind to, one of `default`, `discrete[:N]`, `integrated[:N]`, `virtual[:N]` or `cpu`. #[arg(long, default_value = "default")] - device: av_denoise_core::Device, + device: Device, - /// Swallowed: cargo passes this when invoking the bench binary. + /// Swallowed, since cargo passes this when invoking the bench binary. #[arg(long, hide = true)] bench: bool, } fn main() { - use clap::Parser; let cli = Cli::parse(); println!("NLMeans Per-Kernel Benchmarks - 1920x1080 (TimingMethod::Device)"); diff --git a/av-denoise-core/benches/convert.rs b/av-denoise-core/benches/convert.rs deleted file mode 100644 index c8daf68..0000000 --- a/av-denoise-core/benches/convert.rs +++ /dev/null @@ -1,150 +0,0 @@ -use std::time::Instant; - -const W: usize = 1920; -const H: usize = 1080; -/// 4:2:0 sample count for one frame. -const SAMPLES: usize = W * H + 2 * ((W / 2) * (H / 2)); - -const WARMUP: usize = 5; -const ITERS: usize = 200; - -#[derive(clap::Parser, Debug)] -#[command(about = "Sample <-> f32 conversion benchmark", long_about = None)] -struct Cli { - /// Swallowed: cargo passes this when invoking the bench binary. - #[arg(long, hide = true)] - bench: bool, -} - -fn time(label: &str, mut f: impl FnMut() -> usize) { - for _ in 0..WARMUP { - std::hint::black_box(f()); - } - - let t = Instant::now(); - for _ in 0..ITERS { - std::hint::black_box(f()); - } - let per_ms = t.elapsed().as_secs_f64() / ITERS as f64 * 1000.0; - - println!("{label:<44} {per_ms:>8.3} ms/frame"); -} - -fn main() { - let _cli: Cli = clap::Parser::parse(); - - let normalized: Vec = (0..SAMPLES).map(|i| (i % 1024) as f32 / 1023.0).collect(); - - // Flat luma plane, the simplest read path. - let wire8: Vec = (0..SAMPLES).map(|i| (i % 256) as u8).collect(); - let wire10: Vec = (0..SAMPLES) - .flat_map(|i| ((i % 1024) as u16).to_le_bytes()) - .collect(); - - // Equal-length YUV444 planes for the interleaving path. - let yuv_pixels = SAMPLES / 3; - let plane8: Vec = (0..yuv_pixels).map(|i| (i % 256) as u8).collect(); - - println!("{SAMPLES} samples/frame (1080p 4:2:0), {ITERS} iters"); - - println!("-- output --"); - time("f32 -> 8-bit plane", || { - quantise_plane_narrow(&normalized, 255.0).len() - }); - time("f32 -> 10-bit plane", || { - quantise_plane_wide(&normalized, 1023.0).len() - }); - - println!("-- input --"); - time("8-bit plane -> f32", || read_plane_narrow(&wire8, 255.0).len()); - time("10-bit plane -> f32", || read_plane_wide(&wire10, 1023.0).len()); - - // The fused YUV444 path reads three planes per pixel by index. The - // reviewer of the converter task flagged that indexed reads may not - // vectorise as well as the zip form it replaced, because the bounds - // on the second and third planes cannot be proven. These two rows - // are what decides whether that concern is real. - println!("-- interleave (fused YUV444) --"); - time("8-bit YUV planes -> interleaved f32", || { - interleave_yuv_narrow(&plane8, &plane8, &plane8, 255.0).len() - }); - time("8-bit YUV planes -> interleaved f32 (sliced)", || { - interleave_yuv_narrow_sliced(&plane8, &plane8, &plane8, 255.0).len() - }); -} - -/// Mirrors `plane_to_f32`'s narrow arm. -fn read_plane_narrow(plane: &[u8], max: f32) -> Vec { - let out: Vec = (0..plane.len()).map(|i| plane[i] as f32 / max).collect(); - std::hint::black_box(&out); - out -} - -/// Mirrors `plane_to_f32`'s wide arm. -fn read_plane_wide(plane: &[u8], max: f32) -> Vec { - let samples = plane.len() / 2; - let out: Vec = (0..samples) - .map(|i| u16::from_le_bytes([plane[2 * i], plane[2 * i + 1]]) as f32 / max) - .collect(); - std::hint::black_box(&out); - out -} - -/// Mirrors `interleave_yuv_to_f32`'s narrow arm exactly as shipped. -fn interleave_yuv_narrow(y: &[u8], u: &[u8], v: &[u8], max: f32) -> Vec { - let pixels = y.len(); - let mut out = Vec::with_capacity(pixels * 3); - - for i in 0..pixels { - out.push(y[i] as f32 / max); - out.push(u[i] as f32 / max); - out.push(v[i] as f32 / max); - } - - std::hint::black_box(&out); - out -} - -/// The same loop with all three planes pre-sliced to `pixels`, which lets -/// the compiler drop the per-pixel bounds checks on `u` and `v`. -fn interleave_yuv_narrow_sliced(y: &[u8], u: &[u8], v: &[u8], max: f32) -> Vec { - let pixels = y.len(); - let (y, u, v) = (&y[..pixels], &u[..pixels], &v[..pixels]); - let mut out = Vec::with_capacity(pixels * 3); - - for i in 0..pixels { - out.push(y[i] as f32 / max); - out.push(u[i] as f32 / max); - out.push(v[i] as f32 / max); - } - - std::hint::black_box(&out); - out -} - -/// Mirrors `f32_to_plane`'s narrow arm. The indexed write matches -/// `Narrow::write` rather than a `zip`, because a bounds-check-free -/// iterator idiom would measure a loop the binary never runs. -fn quantise_plane_narrow(plane: &[f32], max: f32) -> Vec { - let mut out = vec![0u8; plane.len()]; - for (i, &v) in plane.iter().enumerate() { - out[i] = quantise(v, max) as u8; - } - std::hint::black_box(&out); - out -} - -/// Mirrors `f32_to_plane`'s wide arm, indexed to match `Wide::write`. -fn quantise_plane_wide(plane: &[f32], max: f32) -> Vec { - let mut out = vec![0u8; plane.len() * 2]; - for (i, &v) in plane.iter().enumerate() { - out[2 * i..2 * i + 2].copy_from_slice(&quantise(v, max).to_le_bytes()); - } - std::hint::black_box(&out); - out -} - -#[inline(always)] -fn quantise(v: f32, max: f32) -> u16 { - (v.clamp(0.0, 1.0) * max + 0.5) as u16 -} diff --git a/av-denoise-core/benches/kernels/accumulate.rs b/av-denoise-core/benches/kernels/accumulate.rs index e774ef5..f61128c 100644 --- a/av-denoise-core/benches/kernels/accumulate.rs +++ b/av-denoise-core/benches/kernels/accumulate.rs @@ -1,18 +1,18 @@ -use av_denoise_core::nlmeans::kernels::nlm_accumulate; +use av_denoise_core::bench_api::kernels::nlm_accumulate; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; use super::{ - H, + HEIGHT, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; @@ -28,8 +28,8 @@ pub struct AccumulateInput { pub struct AccumulateBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for AccumulateBench { @@ -37,15 +37,19 @@ impl Benchmark for AccumulateBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); + let weights_data = vec![0.5f32; pixels]; - let weights = self.client.create_from_slice(f32::as_bytes(&weights_data)); - let accum = self.client.empty(pixels * stored * size_of::()); + let weights_bytes = f32::as_bytes(&weights_data); + let weights = self.client.create_from_slice(weights_bytes); + let accum = self.client.empty(pixels * stored_ch * size_of::()); let weight_sum = self.client.empty(pixels * size_of::()); let max_weight = self.client.empty(pixels * size_of::()); + AccumulateInput { input, accum, @@ -57,16 +61,19 @@ impl Benchmark for AccumulateBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_accumulate::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.weights.clone(), pixels), ArrayArg::from_raw_parts(args.weights.clone(), pixels), @@ -75,20 +82,23 @@ impl Benchmark for AccumulateBench { 0u32, Q_X, Q_Y, - W, - H, + WIDTH, + HEIGHT, ); } + Ok(()) } fn name(&self) -> String { - format!("accumulate_1080p_{}", self.ch_name) + format!("accumulate_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/bilateral.rs b/av-denoise-core/benches/kernels/bilateral.rs index 3d06d50..82508a8 100644 --- a/av-denoise-core/benches/kernels/bilateral.rs +++ b/av-denoise-core/benches/kernels/bilateral.rs @@ -1,5 +1,5 @@ -use av_denoise_core::nlmeans::kernels::nlm_bilateral; -use av_denoise_core::nlmeans::prefilter::bilateral_radius; +use av_denoise_core::bench_api::kernels::nlm_bilateral; +use av_denoise_core::bench_api::prefilter::bilateral_radius; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; @@ -8,21 +8,21 @@ use super::{ BILATERAL_SIGMA_S, BLOCK_X, BLOCK_Y, - H, + HEIGHT, InputOutput, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; pub struct BilateralBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for BilateralBench { @@ -30,11 +30,13 @@ impl Benchmark for BilateralBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); - let output = self.client.empty(pixels * stored * size_of::()); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); + let output = self.client.empty(pixels * stored_ch * size_of::()); + InputOutput { input, output, @@ -43,38 +45,44 @@ impl Benchmark for BilateralBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; let radius = bilateral_radius(BILATERAL_SIGMA_S); + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_bilateral::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.output.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.output.clone(), pixels * stored_ch), 0u32, 1.0 / (2.0 * BILATERAL_SIGMA_S * BILATERAL_SIGMA_S), 1.0 / (2.0 * BILATERAL_SIGMA_R * BILATERAL_SIGMA_R), - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, radius, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { - format!("bilateral_1080p_{}", self.ch_name) + format!("bilateral_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/collab_aggregate.rs b/av-denoise-core/benches/kernels/collab_aggregate.rs index a91a1c4..12047b9 100644 --- a/av-denoise-core/benches/kernels/collab_aggregate.rs +++ b/av-denoise-core/benches/kernels/collab_aggregate.rs @@ -1,18 +1,22 @@ -use av_denoise_core::collab::kernels::aggregate::{collab_normalise, collab_zero_accum}; +use av_denoise_core::bench_api::collab::kernels::aggregate::{collab_normalise, collab_zero_accum}; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{BLOCK_X, BLOCK_Y, H, W, block_sync, stored_channels}; +use super::{BLOCK_X, BLOCK_Y, HEIGHT, WIDTH, block_sync, stored_channels}; -/// Divides the fixed-point accumulators the filters scattered into back -/// out to a finished 1080p frame plane, for each channel mode. Cost -/// scales with `stored_ch`, since one accumulator slot per channel is -/// read for every pixel. +/// The 65,535 workgroups per dimension GPU limit the library clamps to. +/// +/// The library's own constant is crate-private, so a bench target spells it out. +const MAX_GRID_1D: u32 = 65_535; + +/// Divides the fixed-point accumulators back out to a finished 1080p frame plane. +/// +/// Cost scales with `stored_ch`, since one accumulator slot per channel is read for every pixel. pub struct CollabNormaliseBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } #[derive(Clone)] @@ -27,48 +31,56 @@ impl Benchmark for CollabNormaliseBench { type Output = (); fn prepare(&self) -> Self::Input { - let stored = stored_channels(self.ch) as usize; - let pixels = (W * H) as usize; - - // Shaped like a real pass: a few dozen contributions per pixel, - // each already scaled into the accumulator's fixed point. - let accum_data: Vec = (0..pixels * stored).map(|i| (i % 97) as i32 * 8192).collect(); - let wsum_data: Vec = (0..pixels).map(|i| ((i % 31) + 20) as i32 * 8192).collect(); - - CollabNormaliseInput { - accum: self.client.create_from_slice(i32::as_bytes(&accum_data)), - wsum: self.client.create_from_slice(i32::as_bytes(&wsum_data)), - output: self.client.empty(pixels * stored * size_of::()), - } + let stored_ch = stored_channels(self.channels) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + + // Shaped like a real pass, with a few dozen contributions per pixel, each already scaled + // into the accumulator's fixed point. + let accum_data: Vec = (0..pixels * stored_ch) + .map(|index| (index % 97) as i32 * 8192) + .collect(); + let wsum_data: Vec = (0..pixels) + .map(|index| ((index % 31) + 20) as i32 * 8192) + .collect(); + + let accum_bytes = i32::as_bytes(&accum_data); + let accum = self.client.create_from_slice(accum_bytes); + let wsum_bytes = i32::as_bytes(&wsum_data); + let wsum = self.client.create_from_slice(wsum_bytes); + let output = self.client.empty(pixels * stored_ch * size_of::()); + + CollabNormaliseInput { accum, wsum, output } } fn execute(&self, args: Self::Input) -> Result<(), String> { - let stored = stored_channels(self.ch) as usize; - let pixels = (W * H) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let cubes_x = WIDTH.div_ceil(BLOCK_X); + let cubes_y = HEIGHT.div_ceil(BLOCK_Y); unsafe { collab_normalise::launch_unchecked::( &self.client, - CubeCount::new_2d(W.div_ceil(BLOCK_X), H.div_ceil(BLOCK_Y)), + CubeCount::new_2d(cubes_x, cubes_y), CubeDim::new_2d(BLOCK_X, BLOCK_Y), - stored, - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + stored_ch, + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.wsum.clone(), pixels), - ArrayArg::from_raw_parts(args.output.clone(), pixels * stored), - // A single-frame region. This bench measures one frame's - // worth of normalisation, not a cross-frame ring. + ArrayArg::from_raw_parts(args.output.clone(), pixels * stored_ch), + // One frame's region, since this bench measures a single frame's normalisation. 0u32, - W, - H, - self.ch, - stored as u32, + WIDTH, + HEIGHT, + self.channels, + stored_ch as u32, ); } + Ok(()) } fn name(&self) -> String { - format!("collab_normalise_1080p_{}", self.ch_name) + format!("collab_normalise_1080p_{}", self.channel_name) } fn sync(&self) { @@ -76,15 +88,15 @@ impl Benchmark for CollabNormaliseBench { } fn shapes(&self) -> Vec> { - vec![vec![W as usize, H as usize, self.ch as usize]] + vec![vec![WIDTH as usize, HEIGHT as usize, self.channels as usize]] } } /// Clears both accumulators, which runs once before every filter pass. pub struct CollabZeroAccumBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for CollabZeroAccumBench { @@ -92,46 +104,43 @@ impl Benchmark for CollabZeroAccumBench { type Output = (); fn prepare(&self) -> Self::Input { - let stored = stored_channels(self.ch) as usize; - let pixels = (W * H) as usize; - CollabNormaliseInput { - accum: self.client.empty(pixels * stored * size_of::()), - wsum: self.client.empty(pixels * size_of::()), - output: self.client.empty(size_of::()), - } + let stored_ch = stored_channels(self.channels) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let accum = self.client.empty(pixels * stored_ch * size_of::()); + let wsum = self.client.empty(pixels * size_of::()); + let output = self.client.empty(size_of::()); + + CollabNormaliseInput { accum, wsum, output } } fn execute(&self, args: Self::Input) -> Result<(), String> { - let stored = stored_channels(self.ch) as usize; - let pixels = (W * H) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let pixels = (WIDTH * HEIGHT) as usize; let dim = 256u32; - // The same 65,535-workgroups-per-dimension GPU limit the library - // clamps to, spelled out here because its own constant is - // crate-private and a bench is a separate crate target. - // `collab_zero_accum` strides, so the clamp still reaches every - // slot. - const MAX_GRID_1D: u32 = 65_535; - let grid = ((pixels * stored) as u32).div_ceil(dim).min(MAX_GRID_1D); + + // `collab_zero_accum` strides, so the clamped grid still reaches every slot. + let grid = ((pixels * stored_ch) as u32).div_ceil(dim).min(MAX_GRID_1D); unsafe { collab_zero_accum::launch_unchecked::( &self.client, CubeCount::new_1d(grid), CubeDim::new_1d(dim), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.wsum.clone(), pixels), // A single-frame region, as a single-frame caller passes. 0u32, pixels as u32, - stored as u32, + stored_ch as u32, grid * dim, ); } + Ok(()) } fn name(&self) -> String { - format!("collab_zero_accum_1080p_{}", self.ch_name) + format!("collab_zero_accum_1080p_{}", self.channel_name) } fn sync(&self) { @@ -139,6 +148,6 @@ impl Benchmark for CollabZeroAccumBench { } fn shapes(&self) -> Vec> { - vec![vec![W as usize, H as usize, self.ch as usize]] + vec![vec![WIDTH as usize, HEIGHT as usize, self.channels as usize]] } } diff --git a/av-denoise-core/benches/kernels/collab_fused.rs b/av-denoise-core/benches/kernels/collab_fused.rs index 9021d7a..97898a0 100644 --- a/av-denoise-core/benches/kernels/collab_fused.rs +++ b/av-denoise-core/benches/kernels/collab_fused.rs @@ -1,9 +1,13 @@ -use av_denoise_core::collab::geometry::{fused_cubes_x, ref_count, refs_along, strength_map_dims}; -use av_denoise_core::collab::kernels::aggregate::{cross_frame_accum_scale, kaiser_window, weight_scale}; -use av_denoise_core::collab::kernels::fused::{STRENGTH_MAP_LUMA, STRENGTH_MAP_OFF, collab_fused}; -use av_denoise_core::collab::kernels::transforms::dct_noise_profile; -use av_denoise_core::collab::{PATCH_SIZE, grid_frames, needs_warp_uniform_search}; -use av_denoise_core::nlmeans::NOISE_CURVE_BINS; +use av_denoise_core::bench_api::NOISE_CURVE_BINS; +use av_denoise_core::bench_api::collab::geometry::{fused_cubes_x, ref_count, refs_along, strength_map_dims}; +use av_denoise_core::bench_api::collab::kernels::aggregate::{ + cross_frame_accum_scale, + kaiser_window, + weight_scale, +}; +use av_denoise_core::bench_api::collab::kernels::fused::{STRENGTH_MAP_LUMA, STRENGTH_MAP_OFF, collab_fused}; +use av_denoise_core::bench_api::collab::kernels::transforms::dct_noise_profile; +use av_denoise_core::bench_api::collab::{PATCH_SIZE, grid_frames, needs_warp_uniform_search}; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; @@ -23,40 +27,29 @@ use super::nl4d_geometry::{ conf_stride, mv_stride, }; -use super::{H, W, block_sync, make_padded_frame, shapes_with_ch, stored_channels}; +use super::{HEIGHT, WIDTH, block_sync, make_padded_frame, shapes_with_channels, stored_channels}; -/// The fused collaborative kernel at the library's default search -/// geometry, over a 1080p frame ring. One 64-lane cube carries eight -/// reference patches, eight lanes each, so the grid is an eighth as wide -/// along x as the reference count. -/// -/// This is the whole collaborative stage in one launch, matching, -/// filtering and scatter. +/// The fused collaborative kernel at the library's default search geometry over a 1080p frame ring. /// -/// Confidence is uniformly 1.0, so no neighbour block is gated and every -/// candidate the kernel finds runs the full patch comparison. Gating a -/// block skips its comparisons entirely, so leaving it always open here -/// measures the worst case. A bench that gates freely would report a -/// time well under the real one. +/// This is the whole collaborative stage in one launch, matching, filtering and scatter. One 64-lane +/// cube carries eight reference patches of eight lanes each, so the grid is an eighth as wide along +/// x as the reference count. /// -/// `split_mv` picks which of the two motion fields the arm runs on, and -/// the two bracket the real cost of the covering-block search. +/// Confidence is uniformly 1.0, so no neighbour block is gated and every candidate runs the full +/// patch comparison. Gating a block skips its comparisons entirely, so this measures the worst case +/// and a bench that gates freely would report a time well under the real one. /// -/// `false` gives a zeroed field. Every block covering a patch then -/// predicts the same position, all four rectangles coincide, three of -/// them are dropped by the duplicate check and no extra pixel -/// comparison runs. That arm measures the duplicate check on its own. -/// -/// `true` gives each block a vector from its own grid parity, spaced -/// eight pixels apart, which is further than the refine window is wide. -/// The four rectangles covering a patch are then disjoint, nothing -/// deduplicates, and the neighbour search scores four times the -/// positions. That arm is the worst case, and a real motion field lands -/// between the two. +/// `split_mv` picks one of two motion fields that bracket the real cost of the covering-block +/// search. `false` gives a zeroed field. Every block covering a patch then predicts the same +/// position, the duplicate check drops three of the four rectangles and no extra pixel comparison +/// runs, so that arm measures the duplicate check on its own. `true` gives each block a vector from +/// its own grid parity, eight pixels apart, which is further than the refine window is wide. The +/// four rectangles are then disjoint, nothing deduplicates and the neighbour search scores four +/// times the positions. That arm is the worst case, and a real motion field lands between the two. pub struct CollabFusedBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, pub split_mv: bool, /// Launches with a stepped noise curve and `curve_valid = 1`. pub noise_curve: bool, @@ -69,21 +62,18 @@ pub struct CollabFusedBench { /// The luma pool ratio at the default lambda. The kernel's cost does not depend on its value. const POOL_RATIO: f32 = 2.2 / 3.78; -/// A curve that doubles the luma threshold in the darker half and halves it -/// in the brighter one. +/// A curve that doubles the luma threshold in the darker half and halves it in the brighter one. fn stepped_curve() -> [f32; NOISE_CURVE_BINS] { let mut curve = [2.0f32; NOISE_CURVE_BINS]; curve[NOISE_CURVE_BINS / 2..].fill(0.5); curve } -/// How far apart two neighbouring blocks' vectors sit in the split -/// field, in pixels. +/// How far apart two neighbouring blocks' vectors sit in the split field, in pixels. /// -/// `REFINE` is the rectangle's half-width, so two rectangles stay -/// disjoint once their centres are more than `2 * REFINE` apart. Eight -/// clears that with room and keeps every predicted position well inside -/// a 1080p frame. +/// `REFINE` is the rectangle's half-width, so two rectangles stay disjoint once their centres are +/// more than `2 * REFINE` apart. Eight clears that with room and keeps every predicted position +/// well inside a 1080p frame. const SPLIT_MV_SPACING: i32 = 8; #[derive(Clone)] @@ -94,8 +84,9 @@ pub struct CollabFusedInput { pub neighbour_slots: Handle, pub sigma: Handle, pub dct_profile: Handle, - /// The uniform aggregation window. The bench measures the kernel's - /// own cost, and a taper changes none of the work it does. + /// The uniform aggregation window. + /// + /// The bench measures the kernel's own cost, and a taper changes none of the work it does. pub kaiser: Handle, pub accum: Handle, pub wsum: Handle, @@ -112,65 +103,76 @@ impl Benchmark for CollabFusedBench { type Output = (); fn prepare(&self) -> Self::Input { - let stored_ch = stored_channels(self.ch); - let pixels = (W * H) as usize; + let stored_ch = stored_channels(self.channels); + let pixels = (WIDTH * HEIGHT) as usize; let frame_len = pixels * stored_ch as usize; let mut ring_data = Vec::new(); for _ in 0..N_FRAMES { - ring_data.extend(make_padded_frame(W, H, self.ch)); + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + ring_data.extend(frame); } - let ring = self.client.create_from_slice(f32::as_bytes(&ring_data)); - let blocks_x = W.div_ceil(BLK_STEP); - let blocks_y = H.div_ceil(BLK_STEP); + let ring_bytes = f32::as_bytes(&ring_data); + let ring = self.client.create_from_slice(ring_bytes); + + let blocks_x = WIDTH.div_ceil(BLK_STEP); + let blocks_y = HEIGHT.div_ceil(BLK_STEP); let align = self.client.properties().memory.alignment; - let mv_stride = mv_stride(blocks_x, blocks_y, align); - let conf_stride = conf_stride(blocks_x, blocks_y, align); + let neighbour_mv_stride = mv_stride(blocks_x, blocks_y, align); + let neighbour_conf_stride = conf_stride(blocks_x, blocks_y, align); - let mut mv_data = vec![0i32; (2 * RADIUS * mv_stride) as usize]; + let mut mv_data = vec![0i32; (2 * RADIUS * neighbour_mv_stride) as usize]; if self.split_mv { - for t in 0..2 * RADIUS { - for by in 0..blocks_y { - for bx in 0..blocks_x { - let base = (t * mv_stride + (by * blocks_x + bx) * 2) as usize; - mv_data[base] = (bx % 2) as i32 * SPLIT_MV_SPACING; - mv_data[base + 1] = (by % 2) as i32 * SPLIT_MV_SPACING; + for slot in 0..2 * RADIUS { + for block_y in 0..blocks_y { + for block_x in 0..blocks_x { + let base = (slot * neighbour_mv_stride + (block_y * blocks_x + block_x) * 2) as usize; + mv_data[base] = (block_x % 2) as i32 * SPLIT_MV_SPACING; + mv_data[base + 1] = (block_y % 2) as i32 * SPLIT_MV_SPACING; } } } } - let mv_field = self.client.create_from_slice(i32::as_bytes(&mv_data)); - let conf_data = vec![1.0f32; (2 * RADIUS * conf_stride) as usize]; - let confidence = self.client.create_from_slice(f32::as_bytes(&conf_data)); - let neighbour_slots = self.client.create_from_slice(u32::as_bytes(&NEIGHBOUR_SLOTS)); - // Sized for the stored lane count and filled for the logical - // ones, matching what `Nl4dDenoiser` uploads each pass. + let mv_bytes = i32::as_bytes(&mv_data); + let mv_field = self.client.create_from_slice(mv_bytes); + let conf_data = vec![1.0f32; (2 * RADIUS * neighbour_conf_stride) as usize]; + let conf_bytes = f32::as_bytes(&conf_data); + let confidence = self.client.create_from_slice(conf_bytes); + let slots_bytes = u32::as_bytes(&NEIGHBOUR_SLOTS); + let neighbour_slots = self.client.create_from_slice(slots_bytes); + + // Sized for the stored lane count and filled for the logical ones, matching what + // `Nl4dDenoiser` uploads each pass. let mut sigma_host = vec![0.0f32; stored_ch as usize]; - sigma_host[..self.ch as usize].fill(SIGMA); - let sigma = self.client.create_from_slice(f32::as_bytes(&sigma_host)); - let dct_profile = self - .client - .create_from_slice(f32::as_bytes(&dct_noise_profile(0.0))); - let kaiser = self.client.create_from_slice(f32::as_bytes(&kaiser_window(0.0))); + sigma_host[..self.channels as usize].fill(SIGMA); + let sigma_bytes = f32::as_bytes(&sigma_host); + let sigma = self.client.create_from_slice(sigma_bytes); + let profile_host = dct_noise_profile(0.0); + let profile_bytes = f32::as_bytes(&profile_host); + let dct_profile = self.client.create_from_slice(profile_bytes); + let kaiser_host = kaiser_window(0.0); + let kaiser_bytes = f32::as_bytes(&kaiser_host); + let kaiser = self.client.create_from_slice(kaiser_bytes); - // One accumulator region per ring slot, the same shape - // `Nl4dDenoiser` allocates, so the scatter crosses the same - // address range it does in the pipeline. + // One accumulator region per ring slot, the same shape `Nl4dDenoiser` allocates, so the + // scatter crosses the same address range it does in the pipeline. let accum = self .client .empty(frame_len * N_FRAMES as usize * size_of::()); let wsum = self.client.empty(pixels * N_FRAMES as usize * size_of::()); - let group_weight = self.client.empty(ref_count(W, H) * size_of::()); + let refs = ref_count(WIDTH, HEIGHT); + let group_weight = self.client.empty(refs * size_of::()); let curve_host = if self.noise_curve { stepped_curve() } else { [0.0f32; NOISE_CURVE_BINS] }; - let noise_curve = self.client.create_from_slice(f32::as_bytes(&curve_host)); + let curve_bytes = f32::as_bytes(&curve_host); + let noise_curve = self.client.create_from_slice(curve_bytes); - let (map_cols, map_rows) = strength_map_dims(W, H); + let (map_cols, map_rows) = strength_map_dims(WIDTH, HEIGHT); let map_len = (map_cols * map_rows) as usize; let map_host: Vec = if self.strength_map { (0..map_len) @@ -179,7 +181,8 @@ impl Benchmark for CollabFusedBench { } else { vec![1.0f32; map_len] }; - let strength_map = self.client.create_from_slice(f32::as_bytes(&map_host)); + let map_bytes = f32::as_bytes(&map_host); + let strength_map = self.client.create_from_slice(map_bytes); CollabFusedInput { ring, @@ -200,27 +203,35 @@ impl Benchmark for CollabFusedBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let stored_ch = stored_channels(self.ch); - let pixels = (W * H) as usize; + let stored_ch = stored_channels(self.channels); + let pixels = (WIDTH * HEIGHT) as usize; let frame_len = pixels * stored_ch as usize; - let refs = ref_count(W, H); - let refs_x = refs_along(W); - let refs_y = refs_along(H); + let refs = ref_count(WIDTH, HEIGHT); + let refs_x = refs_along(WIDTH); + let refs_y = refs_along(HEIGHT); - let blocks_x = W.div_ceil(BLK_STEP); - let blocks_y = H.div_ceil(BLK_STEP); + let blocks_x = WIDTH.div_ceil(BLK_STEP); + let blocks_y = HEIGHT.div_ceil(BLK_STEP); let align = self.client.properties().memory.alignment; - let mv_stride = mv_stride(blocks_x, blocks_y, align); - let conf_stride = conf_stride(blocks_x, blocks_y, align); + let neighbour_mv_stride = mv_stride(blocks_x, blocks_y, align); + let neighbour_conf_stride = conf_stride(blocks_x, blocks_y, align); - let (map_cols, map_rows) = strength_map_dims(W, H); + let (map_cols, map_rows) = strength_map_dims(WIDTH, HEIGHT); let map_mode = if self.strength_map { STRENGTH_MAP_LUMA } else { STRENGTH_MAP_OFF }; - let grid = CubeCount::new_2d(fused_cubes_x(W), refs_y); + let curve_valid = u32::from(self.noise_curve); + let dct_profile = dct_noise_profile(0.0); + let group_weight_scale = weight_scale(SIGMA, &dct_profile); + let accum_scale = cross_frame_accum_scale(SPATIAL_RADIUS, RADIUS); + let uniform_search = needs_warp_uniform_search(&self.client); + let frames_per_volume = grid_frames(RADIUS); + + let cubes_x = fused_cubes_x(WIDTH); + let grid = CubeCount::new_2d(cubes_x, refs_y); let dim = CubeDim::new_1d(64); unsafe { @@ -230,8 +241,11 @@ impl Benchmark for CollabFusedBench { dim, stored_ch as usize, ArrayArg::from_raw_parts(args.ring.clone(), args.ring_len), - ArrayArg::from_raw_parts(args.mv_field.clone(), (2 * RADIUS * mv_stride) as usize), - ArrayArg::from_raw_parts(args.confidence.clone(), (2 * RADIUS * conf_stride) as usize), + ArrayArg::from_raw_parts(args.mv_field.clone(), (2 * RADIUS * neighbour_mv_stride) as usize), + ArrayArg::from_raw_parts( + args.confidence.clone(), + (2 * RADIUS * neighbour_conf_stride) as usize, + ), ArrayArg::from_raw_parts(args.neighbour_slots.clone(), NEIGHBOUR_SLOTS.len()), ArrayArg::from_raw_parts(args.sigma.clone(), stored_ch as usize), ArrayArg::from_raw_parts(args.noise_curve.clone(), NOISE_CURVE_BINS), @@ -244,23 +258,23 @@ impl Benchmark for CollabFusedBench { CENTRE_SLOT, 0.0f32, LAMBDA_HT, - u32::from(self.noise_curve), + curve_valid, map_mode, - weight_scale(SIGMA, &dct_noise_profile(0.0)), - cross_frame_accum_scale(SPATIAL_RADIUS, RADIUS), - needs_warp_uniform_search(&self.client), + group_weight_scale, + accum_scale, + uniform_search, RADIUS, - grid_frames(RADIUS), + frames_per_volume, REFINE, - mv_stride, - conf_stride, + neighbour_mv_stride, + neighbour_conf_stride, BLK_STEP, BLKSIZE, blocks_x, blocks_y, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, K_MAX, stored_ch, SPATIAL_RADIUS, @@ -271,6 +285,7 @@ impl Benchmark for CollabFusedBench { self.pooled, ); } + Ok(()) } @@ -279,7 +294,10 @@ impl Benchmark for CollabFusedBench { let curve = if self.noise_curve { "_noise_curve" } else { "" }; let map = if self.strength_map { "_strength_map" } else { "" }; let pool = if self.pooled { "_pooled" } else { "" }; - format!("collab_fused_1080p_{}{field}{curve}{map}{pool}", self.ch_name) + format!( + "collab_fused_1080p_{}{field}{curve}{map}{pool}", + self.channel_name + ) } fn sync(&self) { @@ -287,6 +305,6 @@ impl Benchmark for CollabFusedBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/copy.rs b/av-denoise-core/benches/kernels/copy.rs index 0caa155..6804720 100644 --- a/av-denoise-core/benches/kernels/copy.rs +++ b/av-denoise-core/benches/kernels/copy.rs @@ -1,9 +1,18 @@ -use av_denoise_core::nlmeans::kernels::gpu_copy; +use av_denoise_core::bench_api::kernels::gpu_copy; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{BLOCK_1D, COPY_GRID_1D, H, W, block_sync, make_padded_frame, shapes_with_ch, stored_channels}; +use super::{ + BLOCK_1D, + COPY_GRID_1D, + HEIGHT, + WIDTH, + block_sync, + make_padded_frame, + shapes_with_channels, + stored_channels, +}; #[derive(Clone)] pub struct CopyInput { @@ -13,8 +22,8 @@ pub struct CopyInput { pub struct CopyBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for CopyBench { @@ -22,15 +31,19 @@ impl Benchmark for CopyBench { type Output = (); fn prepare(&self) -> Self::Input { - let frame = make_padded_frame(W, H, self.ch); - let src = self.client.create_from_slice(f32::as_bytes(&frame)); + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let src = self.client.create_from_slice(frame_bytes); let dst = self.client.empty(frame.len() * size_of::()); + CopyInput { src, dst } } fn execute(&self, args: Self::Input) -> Result<(), String> { - let len = (W * H) as usize * stored_channels(self.ch) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let len = (WIDTH * HEIGHT) as usize * stored_ch; let total_threads = COPY_GRID_1D * BLOCK_1D; + unsafe { gpu_copy::launch_unchecked::( &self.client, @@ -44,16 +57,19 @@ impl Benchmark for CopyBench { total_threads, ); } + Ok(()) } fn name(&self) -> String { - format!("gpu_copy_1080p_{}", self.ch_name) + format!("gpu_copy_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/dist_2d_weight.rs b/av-denoise-core/benches/kernels/dist_2d_weight.rs index 2544b9b..a7d4852 100644 --- a/av-denoise-core/benches/kernels/dist_2d_weight.rs +++ b/av-denoise-core/benches/kernels/dist_2d_weight.rs @@ -1,29 +1,29 @@ -use av_denoise_core::nlmeans::kernels::nlm_dist_2d_weight; +use av_denoise_core::bench_api::kernels::nlm_dist_2d_weight; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use super::{ BLOCK_X, BLOCK_Y, - H, + HEIGHT, InputOutput, PATCH_RADIUS, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, h2_inv_norm, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; pub struct DistWeightBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for DistWeightBench { @@ -31,10 +31,12 @@ impl Benchmark for DistWeightBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); let output = self.client.empty(pixels * size_of::()); + InputOutput { input, output, @@ -43,40 +45,47 @@ impl Benchmark for DistWeightBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_dist_2d_weight::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.output.clone(), pixels), 0u32, 0u32, Q_X, Q_Y, - h2_inv_norm(), + inv_norm, 0.0f32, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, PATCH_RADIUS, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { - format!("dist_2d_weight_1080p_{}", self.ch_name) + format!("dist_2d_weight_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/dist_2d_weight_ref.rs b/av-denoise-core/benches/kernels/dist_2d_weight_ref.rs index 52c6650..6dcb8e3 100644 --- a/av-denoise-core/benches/kernels/dist_2d_weight_ref.rs +++ b/av-denoise-core/benches/kernels/dist_2d_weight_ref.rs @@ -1,29 +1,29 @@ -use av_denoise_core::nlmeans::kernels::nlm_dist_2d_weight_ref; +use av_denoise_core::bench_api::kernels::nlm_dist_2d_weight_ref; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use super::{ BLOCK_X, BLOCK_Y, - H, + HEIGHT, InputOutput, PATCH_RADIUS, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, h2_inv_norm, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; pub struct DistWeightRefBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for DistWeightRefBench { @@ -31,10 +31,12 @@ impl Benchmark for DistWeightRefBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); let output = self.client.empty(pixels * size_of::()); + InputOutput { input, output, @@ -43,40 +45,47 @@ impl Benchmark for DistWeightRefBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_dist_2d_weight_ref::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.output.clone(), pixels), 0u32, 0u32, Q_X, Q_Y, - h2_inv_norm(), + inv_norm, 0.0f32, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, PATCH_RADIUS, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { - format!("dist_2d_weight_ref_1080p_{}", self.ch_name) + format!("dist_2d_weight_ref_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/distance.rs b/av-denoise-core/benches/kernels/distance.rs index 1bef8a2..ce0f852 100644 --- a/av-denoise-core/benches/kernels/distance.rs +++ b/av-denoise-core/benches/kernels/distance.rs @@ -1,18 +1,18 @@ -use av_denoise_core::nlmeans::kernels::nlm_distance; +use av_denoise_core::bench_api::kernels::nlm_distance; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; use super::{ - H, + HEIGHT, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; @@ -25,8 +25,8 @@ pub struct DistanceInput { pub struct DistanceBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for DistanceBench { @@ -34,10 +34,12 @@ impl Benchmark for DistanceBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); let dist = self.client.empty(pixels * size_of::()); + DistanceInput { input, dist, @@ -46,35 +48,41 @@ impl Benchmark for DistanceBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_distance::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.dist.clone(), pixels), 0u32, 0u32, Q_X, Q_Y, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, ); } + Ok(()) } fn name(&self) -> String { - format!("distance_1080p_{}", self.ch_name) + format!("distance_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/distance_pair.rs b/av-denoise-core/benches/kernels/distance_pair.rs index 0c1a9da..dddc3cf 100644 --- a/av-denoise-core/benches/kernels/distance_pair.rs +++ b/av-denoise-core/benches/kernels/distance_pair.rs @@ -1,18 +1,18 @@ -use av_denoise_core::nlmeans::kernels::nlm_distance_pair; +use av_denoise_core::bench_api::kernels::nlm_distance_pair; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; use super::{ - H, + HEIGHT, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; @@ -26,8 +26,8 @@ pub struct DistancePairInput { pub struct DistancePairBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for DistancePairBench { @@ -35,11 +35,13 @@ impl Benchmark for DistancePairBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); let dist_fwd = self.client.empty(pixels * size_of::()); let dist_bwd = self.client.empty(pixels * size_of::()); + DistancePairInput { input, dist_fwd, @@ -49,14 +51,17 @@ impl Benchmark for DistancePairBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_distance_pair::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.dist_fwd.clone(), pixels), ArrayArg::from_raw_parts(args.dist_bwd.clone(), pixels), @@ -65,21 +70,24 @@ impl Benchmark for DistancePairBench { 0u32, Q_X, Q_Y, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, ); } + Ok(()) } fn name(&self) -> String { - format!("distance_pair_1080p_{}", self.ch_name) + format!("distance_pair_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/distance_pair_ref.rs b/av-denoise-core/benches/kernels/distance_pair_ref.rs index 445b4c9..f75d7c1 100644 --- a/av-denoise-core/benches/kernels/distance_pair_ref.rs +++ b/av-denoise-core/benches/kernels/distance_pair_ref.rs @@ -1,25 +1,25 @@ -use av_denoise_core::nlmeans::kernels::nlm_distance_pair_ref; +use av_denoise_core::bench_api::kernels::nlm_distance_pair_ref; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use super::distance_pair::DistancePairInput; use super::{ - H, + HEIGHT, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; pub struct DistancePairRefBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for DistancePairRefBench { @@ -27,11 +27,13 @@ impl Benchmark for DistancePairRefBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); let dist_fwd = self.client.empty(pixels * size_of::()); let dist_bwd = self.client.empty(pixels * size_of::()); + DistancePairInput { input, dist_fwd, @@ -41,14 +43,17 @@ impl Benchmark for DistancePairRefBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_distance_pair_ref::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.dist_fwd.clone(), pixels), ArrayArg::from_raw_parts(args.dist_bwd.clone(), pixels), @@ -57,21 +62,24 @@ impl Benchmark for DistancePairRefBench { 0u32, Q_X, Q_Y, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, ); } + Ok(()) } fn name(&self) -> String { - format!("distance_pair_ref_1080p_{}", self.ch_name) + format!("distance_pair_ref_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/distance_ref.rs b/av-denoise-core/benches/kernels/distance_ref.rs index 3b03f51..fa9e993 100644 --- a/av-denoise-core/benches/kernels/distance_ref.rs +++ b/av-denoise-core/benches/kernels/distance_ref.rs @@ -1,25 +1,25 @@ -use av_denoise_core::nlmeans::kernels::nlm_distance_ref; +use av_denoise_core::bench_api::kernels::nlm_distance_ref; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use super::distance::DistanceInput; use super::{ - H, + HEIGHT, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; pub struct DistanceRefBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for DistanceRefBench { @@ -27,10 +27,12 @@ impl Benchmark for DistanceRefBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); let dist = self.client.empty(pixels * size_of::()); + DistanceInput { input, dist, @@ -39,35 +41,41 @@ impl Benchmark for DistanceRefBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_distance_ref::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.dist.clone(), pixels), 0u32, 0u32, Q_X, Q_Y, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, ); } + Ok(()) } fn name(&self) -> String { - format!("distance_ref_1080p_{}", self.ch_name) + format!("distance_ref_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/egress.rs b/av-denoise-core/benches/kernels/egress.rs new file mode 100644 index 0000000..7bc2825 --- /dev/null +++ b/av-denoise-core/benches/kernels/egress.rs @@ -0,0 +1,165 @@ +use av_denoise_core::bench_api::engine_kernels::{egress_f32, egress_words}; +use cubecl::benchmark::Benchmark; +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::{BLOCK_1D, block_sync, make_padded_frame, stored_channels}; + +/// How the planes under test are stored. +#[derive(Clone, Copy, Debug)] +pub enum EgressFormat { + U8, + U16Ten, + F32, +} + +impl EgressFormat { + fn samples_per_word(self) -> u32 { + match self { + EgressFormat::U8 => 4, + EgressFormat::U16Ten => 2, + EgressFormat::F32 => 1, + } + } + + fn max(self) -> f32 { + match self { + EgressFormat::U8 => 255.0, + EgressFormat::U16Ten => 1023.0, + EgressFormat::F32 => 1.0, + } + } +} + +#[derive(Clone)] +pub struct EgressInput { + frame: Handle, + planes: Vec, + placeholder: Handle, +} + +pub struct EgressBench { + pub client: ComputeClient, + pub width: u32, + pub height: u32, + pub channels: u32, + pub channel_name: &'static str, + pub format: EgressFormat, +} + +impl EgressBench { + fn pixels(&self) -> u32 { + self.width * self.height + } + + fn words(&self) -> u32 { + let samples_per_word = self.format.samples_per_word(); + self.pixels().div_ceil(samples_per_word) + } + + fn frame_len(&self) -> usize { + let stored_ch = stored_channels(self.channels); + self.pixels() as usize * stored_ch as usize + } +} + +impl Benchmark for EgressBench { + type Input = EgressInput; + type Output = (); + + fn prepare(&self) -> Self::Input { + let frame_data = make_padded_frame(self.width, self.height, self.channels); + let frame_bytes = f32::as_bytes(&frame_data); + let frame = self.client.create_from_slice(frame_bytes); + let plane_bytes = self.words() as usize * size_of::(); + let planes = (0..self.channels) + .map(|_| self.client.empty(plane_bytes)) + .collect(); + let placeholder = self.client.empty(size_of::()); + + EgressInput { + frame, + planes, + placeholder, + } + } + + fn execute(&self, args: Self::Input) -> Result<(), String> { + let pixels = self.pixels(); + let stored_ch = stored_channels(self.channels); + let frame_len = self.frame_len(); + let samples_per_word = self.format.samples_per_word(); + let max = self.format.max(); + let words = self.words(); + let plane_len = match self.format { + EgressFormat::F32 => pixels as usize, + _ => words as usize, + }; + let threads = match self.format { + EgressFormat::F32 => pixels, + _ => words, + }; + let groups = threads.div_ceil(BLOCK_1D).clamp(1, 65535); + let total_threads = groups * BLOCK_1D; + + // Planes past `channels` bind a distinct placeholder the kernel never writes. + let plane_0 = args.planes[0].clone(); + let plane_1 = args.planes.get(1).unwrap_or(&args.placeholder).clone(); + let plane_2 = args.planes.get(2).unwrap_or(&args.placeholder).clone(); + + unsafe { + match self.format { + EgressFormat::F32 => egress_f32::launch_unchecked::( + &self.client, + CubeCount::new_1d(groups), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(args.frame.clone(), frame_len), + ArrayArg::from_raw_parts(plane_0, plane_len), + ArrayArg::from_raw_parts(plane_1, plane_len), + ArrayArg::from_raw_parts(plane_2, plane_len), + pixels, + self.channels, + stored_ch, + total_threads, + ), + EgressFormat::U8 | EgressFormat::U16Ten => egress_words::launch_unchecked::( + &self.client, + CubeCount::new_1d(groups), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(args.frame.clone(), frame_len), + ArrayArg::from_raw_parts(plane_0, plane_len), + ArrayArg::from_raw_parts(plane_1, plane_len), + ArrayArg::from_raw_parts(plane_2, plane_len), + max, + pixels, + self.channels, + stored_ch, + samples_per_word, + words, + total_threads, + ), + } + } + + Ok(()) + } + + fn name(&self) -> String { + format!( + "egress_{}x{}_{:?}_{}", + self.width, self.height, self.format, self.channel_name + ) + } + + fn sync(&self) { + block_sync(&self.client); + } + + fn shapes(&self) -> Vec> { + vec![vec![ + self.width as usize, + self.height as usize, + self.channels as usize, + ]] + } +} diff --git a/av-denoise-core/benches/kernels/finish.rs b/av-denoise-core/benches/kernels/finish.rs index 5a90688..56408c5 100644 --- a/av-denoise-core/benches/kernels/finish.rs +++ b/av-denoise-core/benches/kernels/finish.rs @@ -1,16 +1,16 @@ -use av_denoise_core::nlmeans::kernels::nlm_finish; +use av_denoise_core::bench_api::kernels::nlm_finish; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; use super::{ - H, - W, + HEIGHT, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; @@ -26,8 +26,8 @@ pub struct FinishInput { pub struct FinishBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for FinishBench { @@ -35,17 +35,26 @@ impl Benchmark for FinishBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; - let frame = make_padded_frame(W, H, self.ch); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); - let accum_data = vec![0.25f32; pixels * stored]; - let accum = self.client.create_from_slice(f32::as_bytes(&accum_data)); - let ws_data = vec![1.0f32; pixels]; - let weight_sum = self.client.create_from_slice(f32::as_bytes(&ws_data)); - let mw_data = vec![0.8f32; pixels]; - let max_weight = self.client.create_from_slice(f32::as_bytes(&mw_data)); - let output = self.client.empty(pixels * stored * size_of::()); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); + + let accum_data = vec![0.25f32; pixels * stored_ch]; + let accum_bytes = f32::as_bytes(&accum_data); + let accum = self.client.create_from_slice(accum_bytes); + + let weight_sum_data = vec![1.0f32; pixels]; + let weight_sum_bytes = f32::as_bytes(&weight_sum_data); + let weight_sum = self.client.create_from_slice(weight_sum_bytes); + + let max_weight_data = vec![0.8f32; pixels]; + let max_weight_bytes = f32::as_bytes(&max_weight_data); + let max_weight = self.client.create_from_slice(max_weight_bytes); + + let output = self.client.empty(pixels * stored_ch * size_of::()); + FinishInput { input, output, @@ -57,37 +66,43 @@ impl Benchmark for FinishBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_finish::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.output.clone(), pixels * stored), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.output.clone(), pixels * stored_ch), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.max_weight.clone(), pixels), 0u32, 0u32, 1.0f32, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, ); } + Ok(()) } fn name(&self) -> String { - format!("finish_1080p_{}", self.ch_name) + format!("finish_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/fused_window.rs b/av-denoise-core/benches/kernels/fused_window.rs index ef514eb..961f3c5 100644 --- a/av-denoise-core/benches/kernels/fused_window.rs +++ b/av-denoise-core/benches/kernels/fused_window.rs @@ -1,4 +1,4 @@ -use av_denoise_core::nlmeans::kernels::{ +use av_denoise_core::bench_api::kernels::{ nlm_fused_pair_accumulate_window, nlm_fused_pair_accumulate_window_ref, nlm_fused_single_window, @@ -11,27 +11,28 @@ use cubecl::server::Handle; use super::{ BLOCK_X, BLOCK_Y, - H, + HEIGHT, PATCH_RADIUS, SEARCH_RADIUS, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, h2_inv_norm, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; -/// Zero-filled spatial-offset LUT for `SEARCH_RADIUS`. Zero everywhere -/// reproduces the old flat `noise_offset = 0.0` the single-window -/// benches measured before the LUT replaced that scalar, so the -/// timing stays comparable. +/// A zero-filled spatial-offset LUT for `SEARCH_RADIUS`. +/// +/// Zero everywhere applies no noise offset, matching the `0.0` `noise_offset` the pair-window rows pass. fn zero_spatial_offset_lut(client: &ComputeClient) -> (Handle, usize) { let side = (2 * SEARCH_RADIUS + 1) as usize; let lut = vec![0.0f32; side * side]; - let handle = client.create_from_slice(f32::as_bytes(&lut)); + let lut_bytes = f32::as_bytes(&lut); + let handle = client.create_from_slice(lut_bytes); + (handle, lut.len()) } @@ -60,16 +61,18 @@ pub struct WindowRefInput { frame_len: usize, } -fn prepare_window(client: &ComputeClient, ch: u32) -> WindowInput { - let pixels = (W * H) as usize; - let stored = stored_channels(ch) as usize; - let frame = make_padded_frame(W, H, ch); - let input = client.create_from_slice(f32::as_bytes(&frame)); - let accum = client.empty(pixels * stored * size_of::()); +fn prepare_window(client: &ComputeClient, channels: u32) -> WindowInput { + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(channels) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, channels); + let frame_bytes = f32::as_bytes(&frame); + let input = client.create_from_slice(frame_bytes); + let accum = client.empty(pixels * stored_ch * size_of::()); let weight_sum = client.empty(pixels * size_of::()); let max_weight = client.empty(pixels * size_of::()); let confidence_dummy = client.empty(size_of::()); let (spatial_offset_lut, spatial_offset_lut_len) = zero_spatial_offset_lut(client); + WindowInput { input, accum, @@ -82,17 +85,19 @@ fn prepare_window(client: &ComputeClient, ch: u32) -> WindowInput } } -fn prepare_window_ref(client: &ComputeClient, ch: u32) -> WindowRefInput { - let pixels = (W * H) as usize; - let stored = stored_channels(ch) as usize; - let frame = make_padded_frame(W, H, ch); - let input = client.create_from_slice(f32::as_bytes(&frame)); - let reference = client.create_from_slice(f32::as_bytes(&frame)); - let accum = client.empty(pixels * stored * size_of::()); +fn prepare_window_ref(client: &ComputeClient, channels: u32) -> WindowRefInput { + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(channels) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, channels); + let frame_bytes = f32::as_bytes(&frame); + let input = client.create_from_slice(frame_bytes); + let reference = client.create_from_slice(frame_bytes); + let accum = client.empty(pixels * stored_ch * size_of::()); let weight_sum = client.empty(pixels * size_of::()); let max_weight = client.empty(pixels * size_of::()); let confidence_dummy = client.empty(size_of::()); let (spatial_offset_lut, spatial_offset_lut_len) = zero_spatial_offset_lut(client); + WindowRefInput { input, reference, @@ -108,8 +113,8 @@ fn prepare_window_ref(client: &ComputeClient, ch: u32) -> WindowR pub struct FusedPairWindowBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for FusedPairWindowBench { @@ -117,20 +122,24 @@ impl Benchmark for FusedPairWindowBench { type Output = (); fn prepare(&self) -> Self::Input { - prepare_window(&self.client, self.ch) + prepare_window(&self.client, self.channels) } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_fused_pair_accumulate_window::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.max_weight.clone(), pixels), ArrayArg::from_raw_parts(args.confidence_dummy.clone(), 1), @@ -139,11 +148,11 @@ impl Benchmark for FusedPairWindowBench { 0u32, 0u32, 0u32, - h2_inv_norm(), + inv_norm, 0.0f32, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, PATCH_RADIUS, SEARCH_RADIUS, BLOCK_X, @@ -153,24 +162,27 @@ impl Benchmark for FusedPairWindowBench { 1u32, ); } + Ok(()) } fn name(&self) -> String { - format!("fused_pair_accumulate_window_1080p_{}", self.ch_name) + format!("fused_pair_accumulate_window_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } pub struct FusedSingleWindowBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for FusedSingleWindowBench { @@ -178,52 +190,59 @@ impl Benchmark for FusedSingleWindowBench { type Output = (); fn prepare(&self) -> Self::Input { - prepare_window(&self.client, self.ch) + prepare_window(&self.client, self.channels) } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_fused_single_window::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.max_weight.clone(), pixels), 0u32, - h2_inv_norm(), + inv_norm, ArrayArg::from_raw_parts(args.spatial_offset_lut.clone(), args.spatial_offset_lut_len), - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, PATCH_RADIUS, SEARCH_RADIUS, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { - format!("fused_single_window_1080p_{}", self.ch_name) + format!("fused_single_window_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } pub struct FusedPairWindowRefBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for FusedPairWindowRefBench { @@ -231,21 +250,25 @@ impl Benchmark for FusedPairWindowRefBench { type Output = (); fn prepare(&self) -> Self::Input { - prepare_window_ref(&self.client, self.ch) + prepare_window_ref(&self.client, self.channels) } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_fused_pair_accumulate_window_ref::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.reference.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.max_weight.clone(), pixels), ArrayArg::from_raw_parts(args.confidence_dummy.clone(), 1), @@ -254,11 +277,11 @@ impl Benchmark for FusedPairWindowRefBench { 0u32, 0u32, 0u32, - h2_inv_norm(), + inv_norm, 0.0f32, - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, PATCH_RADIUS, SEARCH_RADIUS, BLOCK_X, @@ -268,24 +291,27 @@ impl Benchmark for FusedPairWindowRefBench { 1u32, ); } + Ok(()) } fn name(&self) -> String { - format!("fused_pair_accumulate_window_ref_1080p_{}", self.ch_name) + format!("fused_pair_accumulate_window_ref_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } pub struct FusedSingleWindowRefBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for FusedSingleWindowRefBench { @@ -293,45 +319,52 @@ impl Benchmark for FusedSingleWindowRefBench { type Output = (); fn prepare(&self) -> Self::Input { - prepare_window_ref(&self.client, self.ch) + prepare_window_ref(&self.client, self.channels) } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_fused_single_window_ref::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), ArrayArg::from_raw_parts(args.reference.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.max_weight.clone(), pixels), 0u32, - h2_inv_norm(), + inv_norm, ArrayArg::from_raw_parts(args.spatial_offset_lut.clone(), args.spatial_offset_lut_len), - W, - H, - self.ch, + WIDTH, + HEIGHT, + self.channels, PATCH_RADIUS, SEARCH_RADIUS, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { - format!("fused_single_window_ref_1080p_{}", self.ch_name) + format!("fused_single_window_ref_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/grain.rs b/av-denoise-core/benches/kernels/grain.rs index 0a29852..f3992e9 100644 --- a/av-denoise-core/benches/kernels/grain.rs +++ b/av-denoise-core/benches/kernels/grain.rs @@ -1,9 +1,9 @@ -use av_denoise_core::nl4d::kernels::{grain_measure, grain_reduce_partials, grain_save_vectors}; +use av_denoise_core::bench_api::nl4d_kernels::{grain_measure, grain_reduce_partials, grain_save_vectors}; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_synthetic_frame}; +use super::{HEIGHT, WIDTH, block_sync, make_synthetic_frame}; const STEP: u32 = 8; const NEIGHBOURS: u32 = 4; @@ -44,8 +44,8 @@ impl GrainSize { pub const GRAIN_SIZES: &[GrainSize] = &[ GrainSize { - width: W, - height: H, + width: WIDTH, + height: HEIGHT, label: "1080p", }, GrainSize { @@ -77,8 +77,10 @@ impl Benchmark for GrainSaveVectorsBench { let blocks = self.size.blocks() as usize; let mv_host = vec![1i32; NEIGHBOURS as usize * blocks * 2]; let conf_host = vec![0.9f32; NEIGHBOURS as usize * blocks]; - let mv = self.client.create_from_slice(i32::as_bytes(&mv_host)); - let conf = self.client.create_from_slice(f32::as_bytes(&conf_host)); + let mv_bytes = i32::as_bytes(&mv_host); + let mv = self.client.create_from_slice(mv_bytes); + let conf_bytes = f32::as_bytes(&conf_host); + let conf = self.client.create_from_slice(conf_bytes); let saved_mv = self.client.empty(RING as usize * blocks * 2 * size_of::()); let saved_conf = self.client.empty(RING as usize * blocks * size_of::()); @@ -157,8 +159,8 @@ impl Benchmark for GrainMeasureBench { ring.push(frame[0]); // Flat outputs with a small per-pixel ripple, so the kept grain is not zero either. - let out_prev = vec![0.5f32; pixels]; - let out_t: Vec = (0..pixels) + let out_prev_host = vec![0.5f32; pixels]; + let out_t_host: Vec = (0..pixels) .map(|index| 0.5 + 0.001 * ((index % 7) as f32 - 3.0)) .collect(); let saved_mv_host = vec![0i32; 2 * blocks * 2]; @@ -166,15 +168,31 @@ impl Benchmark for GrainMeasureBench { let edges_host: Vec = (0..EDGE_COUNT).map(|index| index as f32 * 0.002).collect(); let hist_host = vec![0i32; HIST_TOTAL]; + let ring_bytes = f32::as_bytes(&ring); + let input = self.client.create_from_slice(ring_bytes); + let out_t_bytes = f32::as_bytes(&out_t_host); + let out_t = self.client.create_from_slice(out_t_bytes); + let out_prev_bytes = f32::as_bytes(&out_prev_host); + let out_prev = self.client.create_from_slice(out_prev_bytes); + let saved_mv_bytes = i32::as_bytes(&saved_mv_host); + let saved_mv = self.client.create_from_slice(saved_mv_bytes); + let saved_conf_bytes = f32::as_bytes(&saved_conf_host); + let saved_conf = self.client.create_from_slice(saved_conf_bytes); + let edges_bytes = f32::as_bytes(&edges_host); + let edges = self.client.create_from_slice(edges_bytes); + let hist_bytes = i32::as_bytes(&hist_host); + let hist = self.client.create_from_slice(hist_bytes); + let partials = self.client.empty(cells * PARTIAL_LANES * size_of::()); + GrainMeasureInput { - input: self.client.create_from_slice(f32::as_bytes(&ring)), - out_t: self.client.create_from_slice(f32::as_bytes(&out_t)), - out_prev: self.client.create_from_slice(f32::as_bytes(&out_prev)), - saved_mv: self.client.create_from_slice(i32::as_bytes(&saved_mv_host)), - saved_conf: self.client.create_from_slice(f32::as_bytes(&saved_conf_host)), - edges: self.client.create_from_slice(f32::as_bytes(&edges_host)), - hist: self.client.create_from_slice(i32::as_bytes(&hist_host)), - partials: self.client.empty(cells * PARTIAL_LANES * size_of::()), + input, + out_t, + out_prev, + saved_mv, + saved_conf, + edges, + hist, + partials, } } @@ -184,6 +202,8 @@ impl Benchmark for GrainMeasureBench { let pixels = self.size.pixels(); let blocks = self.size.blocks() as usize; let cells = self.size.cells(); + let blocks_x = width.div_ceil(STEP); + let blocks_y = height.div_ceil(STEP); unsafe { grain_measure::launch_unchecked::( @@ -208,8 +228,8 @@ impl Benchmark for GrainMeasureBench { width, height, 1u32, - width.div_ceil(STEP), - height.div_ceil(STEP), + blocks_x, + blocks_y, STEP, ); } @@ -253,10 +273,12 @@ impl Benchmark for GrainReducePartialsBench { let chunk_host = vec![0.0f32; STRENGTH_GROUPS * RECORD_LANES]; - GrainReducePartialsInput { - partials: self.client.create_from_slice(f32::as_bytes(&partials_host)), - chunk: self.client.create_from_slice(f32::as_bytes(&chunk_host)), - } + let partials_bytes = f32::as_bytes(&partials_host); + let partials = self.client.create_from_slice(partials_bytes); + let chunk_bytes = f32::as_bytes(&chunk_host); + let chunk = self.client.create_from_slice(chunk_bytes); + + GrainReducePartialsInput { partials, chunk } } fn execute(&self, args: Self::Input) -> Result<(), String> { diff --git a/av-denoise-core/benches/kernels/horizontal_sum.rs b/av-denoise-core/benches/kernels/horizontal_sum.rs index 0938bda..983a058 100644 --- a/av-denoise-core/benches/kernels/horizontal_sum.rs +++ b/av-denoise-core/benches/kernels/horizontal_sum.rs @@ -1,9 +1,9 @@ -use av_denoise_core::nlmeans::kernels::nlm_horizontal_sum; +use av_denoise_core::bench_api::kernels::nlm_horizontal_sum; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{BLOCK_X, BLOCK_Y, H, PATCH_RADIUS, W, block_sync, cube_count_2d, cube_dim_2d}; +use super::{BLOCK_X, BLOCK_Y, HEIGHT, PATCH_RADIUS, WIDTH, block_sync, cube_count_2d, cube_dim_2d}; #[derive(Clone)] pub struct HSumInput { @@ -20,39 +20,47 @@ impl Benchmark for HSumBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; + let pixels = (WIDTH * HEIGHT) as usize; let data = vec![0.5f32; pixels]; - let input = self.client.create_from_slice(f32::as_bytes(&data)); + let data_bytes = f32::as_bytes(&data); + let input = self.client.create_from_slice(data_bytes); let output = self.client.empty(pixels * size_of::()); + HSumInput { input, output } } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_horizontal_sum::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), + cube_count, + cube_dim, ArrayArg::from_raw_parts(args.input.clone(), pixels), ArrayArg::from_raw_parts(args.output.clone(), pixels), - W, - H, + WIDTH, + HEIGHT, PATCH_RADIUS, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { "horizontal_sum_1080p".to_string() } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - vec![vec![W as usize, H as usize]] + vec![vec![WIDTH as usize, HEIGHT as usize]] } } diff --git a/av-denoise-core/benches/kernels/horizontal_sum_pair.rs b/av-denoise-core/benches/kernels/horizontal_sum_pair.rs index 1701bbf..6fdcfbd 100644 --- a/av-denoise-core/benches/kernels/horizontal_sum_pair.rs +++ b/av-denoise-core/benches/kernels/horizontal_sum_pair.rs @@ -1,9 +1,9 @@ -use av_denoise_core::nlmeans::kernels::nlm_horizontal_sum_pair; +use av_denoise_core::bench_api::kernels::nlm_horizontal_sum_pair; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{BLOCK_X, BLOCK_Y, H, PATCH_RADIUS, W, block_sync, cube_count_2d, cube_dim_2d}; +use super::{BLOCK_X, BLOCK_Y, HEIGHT, PATCH_RADIUS, WIDTH, block_sync, cube_count_2d, cube_dim_2d}; #[derive(Clone)] pub struct HSumPairInput { @@ -22,12 +22,14 @@ impl Benchmark for HSumPairBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; + let pixels = (WIDTH * HEIGHT) as usize; let data = vec![0.5f32; pixels]; - let input_fwd = self.client.create_from_slice(f32::as_bytes(&data)); - let input_bwd = self.client.create_from_slice(f32::as_bytes(&data)); + let data_bytes = f32::as_bytes(&data); + let input_fwd = self.client.create_from_slice(data_bytes); + let input_bwd = self.client.create_from_slice(data_bytes); let output_fwd = self.client.empty(pixels * size_of::()); let output_bwd = self.client.empty(pixels * size_of::()); + HSumPairInput { input_fwd, input_bwd, @@ -37,33 +39,39 @@ impl Benchmark for HSumPairBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + unsafe { nlm_horizontal_sum_pair::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), + cube_count, + cube_dim, ArrayArg::from_raw_parts(args.input_fwd.clone(), pixels), ArrayArg::from_raw_parts(args.input_bwd.clone(), pixels), ArrayArg::from_raw_parts(args.output_fwd.clone(), pixels), ArrayArg::from_raw_parts(args.output_bwd.clone(), pixels), - W, - H, + WIDTH, + HEIGHT, PATCH_RADIUS, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { "horizontal_sum_pair_1080p".to_string() } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - vec![vec![W as usize, H as usize]] + vec![vec![WIDTH as usize, HEIGHT as usize]] } } diff --git a/av-denoise-core/benches/kernels/ingest.rs b/av-denoise-core/benches/kernels/ingest.rs new file mode 100644 index 0000000..3d88c13 --- /dev/null +++ b/av-denoise-core/benches/kernels/ingest.rs @@ -0,0 +1,152 @@ +use av_denoise_core::bench_api::engine_kernels::{ingest_f32, ingest_words}; +use cubecl::benchmark::Benchmark; +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::{BLOCK_1D, block_sync, stored_channels}; + +/// How the planes under test are stored. +#[derive(Clone, Copy, Debug)] +pub enum IngestFormat { + U8, + U16Ten, + F32, +} + +impl IngestFormat { + fn samples_per_word(self) -> u32 { + match self { + IngestFormat::U8 => 4, + IngestFormat::U16Ten => 2, + IngestFormat::F32 => 1, + } + } + + fn max(self) -> f32 { + match self { + IngestFormat::U8 => 255.0, + IngestFormat::U16Ten => 1023.0, + IngestFormat::F32 => 1.0, + } + } +} + +#[derive(Clone)] +pub struct IngestInput { + planes: Vec, + ring: Handle, +} + +pub struct IngestBench { + pub client: ComputeClient, + pub width: u32, + pub height: u32, + pub channels: u32, + pub channel_name: &'static str, + pub format: IngestFormat, +} + +impl IngestBench { + fn pixels(&self) -> u32 { + self.width * self.height + } + + fn words(&self) -> u32 { + let samples_per_word = self.format.samples_per_word(); + self.pixels().div_ceil(samples_per_word) + } + + fn ring_len(&self) -> usize { + let stored_ch = stored_channels(self.channels); + self.pixels() as usize * stored_ch as usize + } +} + +impl Benchmark for IngestBench { + type Input = IngestInput; + type Output = (); + + fn prepare(&self) -> Self::Input { + let plane_bytes = self.words() as usize * size_of::(); + let plane_data = vec![0x5Au8; plane_bytes]; + let planes = (0..self.channels) + .map(|_| self.client.create_from_slice(&plane_data)) + .collect(); + let ring_len = self.ring_len(); + let ring = self.client.empty(ring_len * size_of::()); + + IngestInput { planes, ring } + } + + fn execute(&self, args: Self::Input) -> Result<(), String> { + let pixels = self.pixels(); + let stored_ch = stored_channels(self.channels); + let groups = pixels.div_ceil(BLOCK_1D).min(65535); + let total_threads = groups * BLOCK_1D; + let words = self.words() as usize; + let ring_len = self.ring_len(); + let samples_per_word = self.format.samples_per_word(); + let max = self.format.max(); + + // Planes past `channels` are placeholders the kernel never reads. + let plane_0 = args.planes[0].clone(); + let plane_1 = args.planes.get(1).unwrap_or(&args.planes[0]).clone(); + let plane_2 = args.planes.get(2).unwrap_or(&args.planes[0]).clone(); + + unsafe { + match self.format { + IngestFormat::F32 => ingest_f32::launch_unchecked::( + &self.client, + CubeCount::new_1d(groups), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(plane_0, pixels as usize), + ArrayArg::from_raw_parts(plane_1, pixels as usize), + ArrayArg::from_raw_parts(plane_2, pixels as usize), + ArrayArg::from_raw_parts(args.ring.clone(), ring_len), + 0u32, + pixels, + self.channels, + stored_ch, + total_threads, + ), + IngestFormat::U8 | IngestFormat::U16Ten => ingest_words::launch_unchecked::( + &self.client, + CubeCount::new_1d(groups), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(plane_0, words), + ArrayArg::from_raw_parts(plane_1, words), + ArrayArg::from_raw_parts(plane_2, words), + ArrayArg::from_raw_parts(args.ring.clone(), ring_len), + max, + 0u32, + pixels, + self.channels, + stored_ch, + samples_per_word, + total_threads, + ), + } + } + + Ok(()) + } + + fn name(&self) -> String { + format!( + "ingest_{}x{}_{:?}_{}", + self.width, self.height, self.format, self.channel_name + ) + } + + fn sync(&self) { + block_sync(&self.client); + } + + fn shapes(&self) -> Vec> { + vec![vec![ + self.width as usize, + self.height as usize, + self.channels as usize, + ]] + } +} diff --git a/av-denoise-core/benches/kernels/mc_block_match_coarse.rs b/av-denoise-core/benches/kernels/mc_block_match_coarse.rs index 1ad90ca..17161b9 100644 --- a/av-denoise-core/benches/kernels/mc_block_match_coarse.rs +++ b/av-denoise-core/benches/kernels/mc_block_match_coarse.rs @@ -1,20 +1,20 @@ -use av_denoise_core::nlmeans::kernels::motion::nlm_mc_block_match_coarse; +use av_denoise_core::bench_api::kernels::motion::nlm_mc_block_match_coarse; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_synthetic_frame, shapes_with_ch}; +use super::{HEIGHT, WIDTH, block_sync, make_synthetic_frame, shapes_with_channels}; -/// Default MVTools-style block-matcher tuning at the coarse pyramid -/// level (`/2`). Mirrors the dispatcher's defaults so the bench -/// reflects steady-state cost when MC is enabled with no overrides. +// The library's default Mvtools block geometry, so the bench reflects the cost of motion +// compensation with no overrides. const FINE_BLKSIZE: u32 = 16; const FINE_STEP: u32 = 8; const SEARCH_RADIUS: u32 = 4; -/// Hierarchical coarse pass: one cube per coarse block, SAD search -/// over a `(2·r + 1)²` window on the `/2` luma pyramid level. Per-block -/// MV result is up-scaled to seed the fine pass. +/// The hierarchical coarse pass on the half-resolution luma pyramid level. +/// +/// One cube per coarse block runs a SAD search over a `(2·r + 1)²` window, and each block's +/// vector is scaled up to seed the fine pass. pub struct BlockMatchCoarseBench { pub client: ComputeClient, } @@ -31,15 +31,17 @@ impl Benchmark for BlockMatchCoarseBench { type Output = (); fn prepare(&self) -> Self::Input { - let coarse_w = W / 2; - let coarse_h = H / 2; - let centre_frame = make_synthetic_frame(coarse_w, coarse_h, 1); - let neighbour_frame = make_synthetic_frame(coarse_w, coarse_h, 1); - let centre = self.client.create_from_slice(f32::as_bytes(¢re_frame)); - let neighbour = self.client.create_from_slice(f32::as_bytes(&neighbour_frame)); + let coarse_width = WIDTH / 2; + let coarse_height = HEIGHT / 2; + let centre_frame = make_synthetic_frame(coarse_width, coarse_height, 1); + let neighbour_frame = make_synthetic_frame(coarse_width, coarse_height, 1); + let centre_bytes = f32::as_bytes(¢re_frame); + let centre = self.client.create_from_slice(centre_bytes); + let neighbour_bytes = f32::as_bytes(&neighbour_frame); + let neighbour = self.client.create_from_slice(neighbour_bytes); - let fine_blocks_x = W.div_ceil(FINE_STEP); - let fine_blocks_y = H.div_ceil(FINE_STEP); + let fine_blocks_x = WIDTH.div_ceil(FINE_STEP); + let fine_blocks_y = HEIGHT.div_ceil(FINE_STEP); let mv_field = self .client .empty((fine_blocks_x * fine_blocks_y * 2) as usize * size_of::()); @@ -52,17 +54,17 @@ impl Benchmark for BlockMatchCoarseBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let coarse_w = W / 2; - let coarse_h = H / 2; + let coarse_width = WIDTH / 2; + let coarse_height = HEIGHT / 2; let coarse_blksize = FINE_BLKSIZE / 2; let coarse_step = FINE_STEP / 2; let coarse_scale = 2u32; - let fine_blocks_x = W.div_ceil(FINE_STEP); - let fine_blocks_y = H.div_ceil(FINE_STEP); - let coarse_blocks_x = coarse_w.div_ceil(coarse_step); - let coarse_blocks_y = coarse_h.div_ceil(coarse_step); + let fine_blocks_x = WIDTH.div_ceil(FINE_STEP); + let fine_blocks_y = HEIGHT.div_ceil(FINE_STEP); + let coarse_blocks_x = coarse_width.div_ceil(coarse_step); + let coarse_blocks_y = coarse_height.div_ceil(coarse_step); - let level_len = (coarse_w * coarse_h) as usize; + let level_len = (coarse_width * coarse_height) as usize; let mv_len = (fine_blocks_x * fine_blocks_y * 2) as usize; let grid = CubeCount::new_2d(coarse_blocks_x, coarse_blocks_y); let dim = CubeDim::new_2d(8, 8); @@ -75,8 +77,8 @@ impl Benchmark for BlockMatchCoarseBench { ArrayArg::from_raw_parts(args.centre.clone(), level_len), ArrayArg::from_raw_parts(args.neighbour.clone(), level_len), ArrayArg::from_raw_parts(args.mv_field.clone(), mv_len), - coarse_w, - coarse_h, + coarse_width, + coarse_height, coarse_blksize, coarse_step, SEARCH_RADIUS, @@ -86,6 +88,7 @@ impl Benchmark for BlockMatchCoarseBench { FINE_STEP, ); } + Ok(()) } @@ -98,6 +101,6 @@ impl Benchmark for BlockMatchCoarseBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(1) + shapes_with_channels(1) } } diff --git a/av-denoise-core/benches/kernels/mc_block_match_fine.rs b/av-denoise-core/benches/kernels/mc_block_match_fine.rs index f34e578..0545997 100644 --- a/av-denoise-core/benches/kernels/mc_block_match_fine.rs +++ b/av-denoise-core/benches/kernels/mc_block_match_fine.rs @@ -1,25 +1,24 @@ -use av_denoise_core::nlmeans::kernels::motion::nlm_mc_block_match_fine; +use av_denoise_core::bench_api::kernels::motion::nlm_mc_block_match_fine; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_synthetic_frame, shapes_with_ch}; +use super::{HEIGHT, WIDTH, block_sync, make_synthetic_frame, shapes_with_channels}; const FINE_BLKSIZE: u32 = 16; const FINE_STEP: u32 = 8; const SEARCH_RADIUS: u32 = 4; -// Mirrors `motion::confidence::THSAD_PIXEL` / `thsad`, duplicated here -// since those host helpers are crate-internal and unreachable from the -// bench binary. +/// The library's per-pixel SAD threshold. +/// +/// The library's own constant and `thsad` helper are crate-private, so a bench target spells it out. const THSAD_PIXEL: f32 = 0.02; -/// Fine refinement pass at full resolution. Reads no seed (the bench -/// uses `use_seed = 0` so the kernel cost is bounded purely by the -/// `(2·r + 1)²` search window per block; fair across runs even if -/// the coarse bench tuning changes). `sad_noise_floor` is `0.0` (no -/// noise estimate available at this level) and `thsad` uses the -/// library's default `thsad_scale` of `1.0`. +/// The fine refinement pass at full resolution. +/// +/// It reads no seed, so the cost is bounded by the `(2·r + 1)²` search window per block and stays +/// fair even if the coarse bench's tuning changes. `sad_noise_floor` is 0.0 because no noise +/// estimate is available at this level, and `thsad` uses the library's default `thsad_scale` of 1.0. pub struct BlockMatchFineBench { pub client: ComputeClient, } @@ -37,13 +36,15 @@ impl Benchmark for BlockMatchFineBench { type Output = (); fn prepare(&self) -> Self::Input { - let centre_frame = make_synthetic_frame(W, H, 1); - let neighbour_frame = make_synthetic_frame(W, H, 1); - let centre = self.client.create_from_slice(f32::as_bytes(¢re_frame)); - let neighbour = self.client.create_from_slice(f32::as_bytes(&neighbour_frame)); + let centre_frame = make_synthetic_frame(WIDTH, HEIGHT, 1); + let neighbour_frame = make_synthetic_frame(WIDTH, HEIGHT, 1); + let centre_bytes = f32::as_bytes(¢re_frame); + let centre = self.client.create_from_slice(centre_bytes); + let neighbour_bytes = f32::as_bytes(&neighbour_frame); + let neighbour = self.client.create_from_slice(neighbour_bytes); - let blocks_x = W.div_ceil(FINE_STEP); - let blocks_y = H.div_ceil(FINE_STEP); + let blocks_x = WIDTH.div_ceil(FINE_STEP); + let blocks_y = HEIGHT.div_ceil(FINE_STEP); let mv_field = self .client .empty((blocks_x * blocks_y * 2) as usize * size_of::()); @@ -60,9 +61,9 @@ impl Benchmark for BlockMatchFineBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let level_len = (W * H) as usize; - let blocks_x = W.div_ceil(FINE_STEP); - let blocks_y = H.div_ceil(FINE_STEP); + let level_len = (WIDTH * HEIGHT) as usize; + let blocks_x = WIDTH.div_ceil(FINE_STEP); + let blocks_y = HEIGHT.div_ceil(FINE_STEP); let mv_len = (blocks_x * blocks_y * 2) as usize; let conf_len = (blocks_x * blocks_y) as usize; let thsad = (FINE_BLKSIZE * FINE_BLKSIZE) as f32 * THSAD_PIXEL; @@ -79,18 +80,19 @@ impl Benchmark for BlockMatchFineBench { ArrayArg::from_raw_parts(args.neighbour.clone(), level_len), ArrayArg::from_raw_parts(args.mv_field.clone(), mv_len), ArrayArg::from_raw_parts(args.confidence.clone(), conf_len), - true, // benchmark the full production-realistic cost, confidence included + true, // Includes the confidence write, the full production cost. 0.0, thsad, - W, - H, + WIDTH, + HEIGHT, FINE_BLKSIZE, FINE_STEP, SEARCH_RADIUS, - 0u32, // use_seed = 0; bench worst-case without coarse seed + 0u32, // `use_seed = 0`, the worst case without a coarse seed. blocks_x, ); } + Ok(()) } @@ -103,6 +105,6 @@ impl Benchmark for BlockMatchFineBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(1) + shapes_with_channels(1) } } diff --git a/av-denoise-core/benches/kernels/mc_chain_compose.rs b/av-denoise-core/benches/kernels/mc_chain_compose.rs index 1161ecc..0a45bef 100644 --- a/av-denoise-core/benches/kernels/mc_chain_compose.rs +++ b/av-denoise-core/benches/kernels/mc_chain_compose.rs @@ -1,28 +1,30 @@ -use av_denoise_core::nlmeans::kernels::motion::nlm_mc_chain_compose; +use av_denoise_core::bench_api::kernels::motion::nlm_mc_chain_compose; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W}; +use super::{HEIGHT, WIDTH, block_sync, shapes_with_channels}; const CHAIN_STEP: u32 = 8; const CHAIN_RADIUS: u32 = 4; const CHAIN_PAIR_RING_SLOTS: u32 = 2 * CHAIN_RADIUS; -/// Padded per-direction pair-ring stride in i32 elements, mirroring -/// `MotionCtx::pair_direction_stride` (not reachable from a bench -/// target, since it's `pub(crate)`). One `(dx, dy)` i32 pair per block, -/// rounded up to the GPU storage-buffer offset alignment (32 bytes). +/// The padded per-direction pair-ring stride in `i32` elements. +/// +/// It mirrors `MotionCtx::pair_direction_stride`, which is crate-private. Each block holds one +/// `(dx, dy)` pair, and the total is rounded up to a 32-byte storage-buffer offset alignment. fn padded_pair_direction_stride(blocks_x: u32, blocks_y: u32) -> u32 { let unpadded_bytes = (blocks_x as u64) * (blocks_y as u64) * 2 * size_of::() as u64; - (unpadded_bytes.next_multiple_of(32) / size_of::() as u64) as u32 + let padded_bytes = unpadded_bytes.next_multiple_of(32); + + (padded_bytes / size_of::() as u64) as u32 } -/// Chained motion-composition kernel at 1080p geometry, walking a -/// full `R = 4` chain (the deepest hop the default temporal radius -/// exercises). The pair ring holds synthetic (never-analysed) data; -/// the kernel's cost is dominated by the per-block hop walk, not the -/// values it reads, so the content doesn't need to be meaningful. +/// The chained motion-composition kernel at 1080p, walking a full `R = 4` chain. +/// +/// The pair ring holds synthetic data that was never analysed. The cost is dominated by the +/// per-block hop walk rather than the values it reads, so the content does not need to be +/// meaningful. pub struct ChainComposeBench { pub client: ComputeClient, } @@ -38,14 +40,15 @@ impl Benchmark for ChainComposeBench { type Output = (); fn prepare(&self) -> Self::Input { - let blocks_x = W.div_ceil(CHAIN_STEP); - let blocks_y = H.div_ceil(CHAIN_STEP); + let blocks_x = WIDTH.div_ceil(CHAIN_STEP); + let blocks_y = HEIGHT.div_ceil(CHAIN_STEP); let dir_len = padded_pair_direction_stride(blocks_x, blocks_y); let slot_len = 2 * dir_len; let pair_ring_len = CHAIN_PAIR_RING_SLOTS as usize * slot_len as usize; let pair_ring_data = vec![0i32; pair_ring_len]; - let pair_ring = self.client.create_from_slice(i32::as_bytes(&pair_ring_data)); + let pair_ring_bytes = i32::as_bytes(&pair_ring_data); + let pair_ring = self.client.create_from_slice(pair_ring_bytes); let mv_field = self .client .empty((blocks_x * blocks_y * 2) as usize * size_of::()); @@ -54,14 +57,13 @@ impl Benchmark for ChainComposeBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let blocks_x = W.div_ceil(CHAIN_STEP); - let blocks_y = H.div_ceil(CHAIN_STEP); + let blocks_x = WIDTH.div_ceil(CHAIN_STEP); + let blocks_y = HEIGHT.div_ceil(CHAIN_STEP); let dir_len = padded_pair_direction_stride(blocks_x, blocks_y); let slot_len = 2 * dir_len; let pair_ring_len = CHAIN_PAIR_RING_SLOTS as usize * slot_len as usize; - // One thread per output block, matching `run_chain_compose`'s - // production launch shape. + // One thread per output block, the same launch shape the library uses. let grid = CubeCount::new_2d(blocks_x, blocks_y); let dim = CubeDim::new_2d(1, 1); @@ -79,12 +81,13 @@ impl Benchmark for ChainComposeBench { dir_len, slot_len, CHAIN_STEP, - W, - H, + WIDTH, + HEIGHT, blocks_x, blocks_y, ); } + Ok(()) } @@ -93,10 +96,10 @@ impl Benchmark for ChainComposeBench { } fn sync(&self) { - super::block_sync(&self.client); + block_sync(&self.client); } fn shapes(&self) -> Vec> { - super::shapes_with_ch(1) + shapes_with_channels(1) } } diff --git a/av-denoise-core/benches/kernels/mc_confidence.rs b/av-denoise-core/benches/kernels/mc_confidence.rs index f76477d..8b45ec6 100644 --- a/av-denoise-core/benches/kernels/mc_confidence.rs +++ b/av-denoise-core/benches/kernels/mc_confidence.rs @@ -1,24 +1,23 @@ -use av_denoise_core::nlmeans::kernels::motion::nlm_mc_block_match_fine; -use av_denoise_core::nlmeans::motion::{DEFAULT_BLKSIZE, DEFAULT_OVERLAP}; +use av_denoise_core::bench_api::kernels::motion::nlm_mc_block_match_fine; +use av_denoise_core::bench_api::motion::{DEFAULT_BLKSIZE, DEFAULT_OVERLAP}; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_synthetic_frame, shapes_with_ch}; +use super::{HEIGHT, WIDTH, block_sync, make_synthetic_frame, shapes_with_channels}; const CONF_STEP: u32 = DEFAULT_BLKSIZE - DEFAULT_OVERLAP; const SEARCH_RADIUS: u32 = 0; -// Mirrors `motion::confidence::THSAD_PIXEL` / `thsad`, duplicated here -// since those host helpers are crate-internal and unreachable from the -// bench binary. +/// The library's per-pixel SAD threshold. +/// +/// The library's own constant and `thsad` helper are crate-private, so a bench target spells it out. const THSAD_PIXEL: f32 = 0.02; -/// The no-MC confidence pass. A single-candidate SAD (`search_radius = -/// 0`, no seed) at the library's default block geometry, the cost -/// profile of confidence weighting when motion compensation itself is -/// off. Cheaper than [`super::mc_block_match_fine::BlockMatchFineBench`], -/// which sweeps a real `(2·4 + 1)²` search window. +/// The confidence pass without motion compensation. +/// +/// It scores a single candidate with no seed and `search_radius = 0` at the library's default +/// block geometry, which is the cost of confidence weighting when motion compensation is off. pub struct McConfidenceBench { pub client: ComputeClient, } @@ -36,13 +35,15 @@ impl Benchmark for McConfidenceBench { type Output = (); fn prepare(&self) -> Self::Input { - let centre_frame = make_synthetic_frame(W, H, 1); - let neighbour_frame = make_synthetic_frame(W, H, 1); - let centre = self.client.create_from_slice(f32::as_bytes(¢re_frame)); - let neighbour = self.client.create_from_slice(f32::as_bytes(&neighbour_frame)); + let centre_frame = make_synthetic_frame(WIDTH, HEIGHT, 1); + let neighbour_frame = make_synthetic_frame(WIDTH, HEIGHT, 1); + let centre_bytes = f32::as_bytes(¢re_frame); + let centre = self.client.create_from_slice(centre_bytes); + let neighbour_bytes = f32::as_bytes(&neighbour_frame); + let neighbour = self.client.create_from_slice(neighbour_bytes); - let blocks_x = W.div_ceil(CONF_STEP); - let blocks_y = H.div_ceil(CONF_STEP); + let blocks_x = WIDTH.div_ceil(CONF_STEP); + let blocks_y = HEIGHT.div_ceil(CONF_STEP); let mv_scratch = self .client .empty((blocks_x * blocks_y * 2) as usize * size_of::()); @@ -59,9 +60,9 @@ impl Benchmark for McConfidenceBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let level_len = (W * H) as usize; - let blocks_x = W.div_ceil(CONF_STEP); - let blocks_y = H.div_ceil(CONF_STEP); + let level_len = (WIDTH * HEIGHT) as usize; + let blocks_x = WIDTH.div_ceil(CONF_STEP); + let blocks_y = HEIGHT.div_ceil(CONF_STEP); let mv_len = (blocks_x * blocks_y * 2) as usize; let conf_len = (blocks_x * blocks_y) as usize; let thsad = (DEFAULT_BLKSIZE * DEFAULT_BLKSIZE) as f32 * THSAD_PIXEL; @@ -78,18 +79,19 @@ impl Benchmark for McConfidenceBench { ArrayArg::from_raw_parts(args.neighbour.clone(), level_len), ArrayArg::from_raw_parts(args.mv_scratch.clone(), mv_len), ArrayArg::from_raw_parts(args.confidence.clone(), conf_len), - true, // this bench measures the confidence write itself + true, // The confidence write is what this bench measures. 0.0, thsad, - W, - H, + WIDTH, + HEIGHT, DEFAULT_BLKSIZE, CONF_STEP, SEARCH_RADIUS, - 0u32, // use_seed = 0 (no coarse pass in the no-MC path) + 0u32, // `use_seed = 0`, since there is no coarse pass without motion compensation. blocks_x, ); } + Ok(()) } @@ -102,6 +104,6 @@ impl Benchmark for McConfidenceBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(1) + shapes_with_channels(1) } } diff --git a/av-denoise-core/benches/kernels/mc_downscale.rs b/av-denoise-core/benches/kernels/mc_downscale.rs index c88b9e4..5793543 100644 --- a/av-denoise-core/benches/kernels/mc_downscale.rs +++ b/av-denoise-core/benches/kernels/mc_downscale.rs @@ -1,12 +1,13 @@ -use av_denoise_core::nlmeans::kernels::motion::nlm_mc_downscale; +use av_denoise_core::bench_api::kernels::motion::nlm_mc_downscale; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_synthetic_frame, shapes_with_ch}; +use super::{HEIGHT, WIDTH, block_sync, make_synthetic_frame, shapes_with_channels}; -/// 2x2 box downsample over a full-res luma frame into a `/2` slot. -/// Used to build the coarse pyramid level for motion estimation. +/// A 2x2 box downsample of a full-resolution luma frame into a half-resolution slot. +/// +/// It builds the coarse pyramid level for motion estimation. pub struct DownscaleBench { pub client: ComputeClient, } @@ -22,20 +23,26 @@ impl Benchmark for DownscaleBench { type Output = (); fn prepare(&self) -> Self::Input { - let src_frame = make_synthetic_frame(W, H, 1); - let src = self.client.create_from_slice(f32::as_bytes(&src_frame)); - let dst_w = W / 2; - let dst_h = H / 2; - let dst = self.client.empty((dst_w * dst_h) as usize * size_of::()); + let src_frame = make_synthetic_frame(WIDTH, HEIGHT, 1); + let src_bytes = f32::as_bytes(&src_frame); + let src = self.client.create_from_slice(src_bytes); + let dst_width = WIDTH / 2; + let dst_height = HEIGHT / 2; + let dst = self + .client + .empty((dst_width * dst_height) as usize * size_of::()); + DownscaleInput { src, dst } } fn execute(&self, args: Self::Input) -> Result<(), String> { - let dst_w = W / 2; - let dst_h = H / 2; + let dst_width = WIDTH / 2; + let dst_height = HEIGHT / 2; let block_x = 16u32; let block_y = 16u32; - let grid = CubeCount::new_2d(dst_w.div_ceil(block_x), dst_h.div_ceil(block_y)); + let cubes_x = dst_width.div_ceil(block_x); + let cubes_y = dst_height.div_ceil(block_y); + let grid = CubeCount::new_2d(cubes_x, cubes_y); let dim = CubeDim::new_2d(block_x, block_y); unsafe { @@ -43,16 +50,17 @@ impl Benchmark for DownscaleBench { &self.client, grid, dim, - ArrayArg::from_raw_parts(args.src.clone(), (W * H) as usize), - ArrayArg::from_raw_parts(args.dst.clone(), (dst_w * dst_h) as usize), + ArrayArg::from_raw_parts(args.src.clone(), (WIDTH * HEIGHT) as usize), + ArrayArg::from_raw_parts(args.dst.clone(), (dst_width * dst_height) as usize), 0u32, 0u32, - W, - H, - dst_w, - dst_h, + WIDTH, + HEIGHT, + dst_width, + dst_height, ); } + Ok(()) } @@ -65,6 +73,6 @@ impl Benchmark for DownscaleBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(1) + shapes_with_channels(1) } } diff --git a/av-denoise-core/benches/kernels/mc_warp.rs b/av-denoise-core/benches/kernels/mc_warp.rs index aa4d21b..3479550 100644 --- a/av-denoise-core/benches/kernels/mc_warp.rs +++ b/av-denoise-core/benches/kernels/mc_warp.rs @@ -1,19 +1,19 @@ -use av_denoise_core::nlmeans::kernels::motion::nlm_mc_warp; +use av_denoise_core::bench_api::kernels::motion::nlm_mc_warp; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_padded_frame, shapes_with_ch, stored_channels}; +use super::{HEIGHT, WIDTH, block_sync, make_padded_frame, shapes_with_channels, stored_channels}; const FINE_STEP: u32 = 8; -/// Apply the MV field to warp a neighbour into spatial alignment with -/// the centre. Cost is roughly one packed `Vector` read + -/// store per output pixel. +/// Warps a neighbour into spatial alignment with the centre using the motion field. +/// +/// Cost is roughly one packed `Vector` read and store per output pixel. pub struct WarpBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } #[derive(Clone)] @@ -30,12 +30,13 @@ impl Benchmark for WarpBench { type Output = (); fn prepare(&self) -> Self::Input { - let frame = make_padded_frame(W, H, self.ch); - let src = self.client.create_from_slice(f32::as_bytes(&frame)); + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); + let frame_bytes = f32::as_bytes(&frame); + let src = self.client.create_from_slice(frame_bytes); let dst = self.client.empty(frame.len() * size_of::()); - let blocks_x = W.div_ceil(FINE_STEP); - let blocks_y = H.div_ceil(FINE_STEP); + let blocks_x = WIDTH.div_ceil(FINE_STEP); + let blocks_y = HEIGHT.div_ceil(FINE_STEP); let mv_len = (blocks_x * blocks_y * 2) as usize; let mv_field = self.client.empty(mv_len * size_of::()); @@ -49,20 +50,22 @@ impl Benchmark for WarpBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let stored = stored_channels(self.ch) as usize; + let stored_ch = stored_channels(self.channels) as usize; let block_x = 16u32; let block_y = 16u32; - let grid = CubeCount::new_2d(W.div_ceil(block_x), H.div_ceil(block_y)); + let cubes_x = WIDTH.div_ceil(block_x); + let cubes_y = HEIGHT.div_ceil(block_y); + let grid = CubeCount::new_2d(cubes_x, cubes_y); let dim = CubeDim::new_2d(block_x, block_y); - let blocks_x = W.div_ceil(FINE_STEP); - let blocks_y = H.div_ceil(FINE_STEP); + let blocks_x = WIDTH.div_ceil(FINE_STEP); + let blocks_y = HEIGHT.div_ceil(FINE_STEP); unsafe { nlm_mc_warp::launch_unchecked::( &self.client, grid, dim, - stored, + stored_ch, ArrayArg::from_raw_parts(args.src.clone(), args.frame_len), ArrayArg::from_raw_parts(args.dst.clone(), args.frame_len), ArrayArg::from_raw_parts(args.mv_field.clone(), args.mv_len), @@ -71,15 +74,16 @@ impl Benchmark for WarpBench { FINE_STEP, blocks_x, blocks_y, - W, - H, + WIDTH, + HEIGHT, ); } + Ok(()) } fn name(&self) -> String { - format!("mc_warp_1080p_{}", self.ch_name) + format!("mc_warp_1080p_{}", self.channel_name) } fn sync(&self) { @@ -87,6 +91,6 @@ impl Benchmark for WarpBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/mod.rs b/av-denoise-core/benches/kernels/mod.rs index bcec798..49bea91 100644 --- a/av-denoise-core/benches/kernels/mod.rs +++ b/av-denoise-core/benches/kernels/mod.rs @@ -1,9 +1,3 @@ -use av_denoise_core::nlmeans::NlmParams; -pub use av_denoise_core::nlmeans::{BLOCK_X, BLOCK_Y}; -use cubecl::benchmark::{Benchmark, BenchmarkComputations, TimingMethod}; -use cubecl::prelude::*; -use cubecl::server::Handle; - pub mod accumulate; pub mod bilateral; pub mod collab_aggregate; @@ -15,11 +9,13 @@ pub mod distance; pub mod distance_pair; pub mod distance_pair_ref; pub mod distance_ref; +pub mod egress; pub mod finish; pub mod fused_window; pub mod grain; pub mod horizontal_sum; pub mod horizontal_sum_pair; +pub mod ingest; pub mod mc_block_match_coarse; pub mod mc_block_match_fine; pub mod mc_chain_compose; @@ -29,15 +25,19 @@ pub mod mc_warp; pub mod mv_regularise; pub mod nl4d_geometry; pub mod noise_partial; -pub mod pack_wire; pub mod temporal_noise_stats; -pub mod unpack_wire; pub mod vertical_weight; pub mod vweight_pair_accumulate; pub mod zero; -pub const W: u32 = 1920; -pub const H: u32 = 1080; +use av_denoise_core::bench_api::NlmParams; +pub use av_denoise_core::bench_api::{BLOCK_X, BLOCK_Y}; +use cubecl::benchmark::{Benchmark, BenchmarkComputations, TimingMethod}; +use cubecl::prelude::*; +use cubecl::server::Handle; + +pub const WIDTH: u32 = 1920; +pub const HEIGHT: u32 = 1080; pub const PATCH_RADIUS: u32 = 4; pub const SEARCH_RADIUS: u32 = 2; pub const Q_X: i32 = 1; @@ -47,24 +47,24 @@ pub const BILATERAL_SIGMA_R: f32 = 0.02; pub const BLOCK_1D: u32 = 256; pub const COPY_GRID_1D: u32 = 1024; -/// (logical channels, label). +/// The logical channel count and row label of each channel mode. pub const CHANNELS: &[(u32, &str)] = &[(1, "luma"), (2, "chroma"), (3, "yuv")]; -pub fn stored_channels(ch: u32) -> u32 { - match ch { +pub fn stored_channels(channels: u32) -> u32 { + match channels { 1 => 1, 2 => 2, _ => 4, } } -pub fn make_synthetic_frame(w: u32, h: u32, ch: u32) -> Vec { - let mut data = Vec::with_capacity((w * h * ch) as usize); - for y in 0..h { - for x in 0..w { +pub fn make_synthetic_frame(width: u32, height: u32, channels: u32) -> Vec { + let mut data = Vec::with_capacity((width * height * channels) as usize); + for y in 0..height { + for x in 0..width { let base = 0.5 + 0.2 * (x as f32 * 0.05).sin() * (y as f32 * 0.03).cos(); - for c in 0..ch { - let seed = (y * w + x) * ch + c; + for channel in 0..channels { + let seed = (y * width + x) * channels + channel; let hash = seed .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(340573321)); @@ -73,38 +73,46 @@ pub fn make_synthetic_frame(w: u32, h: u32, ch: u32) -> Vec { } } } + data } -/// Pad to next-pow2 lane count (matches `NlmDenoiser` internal storage). -pub fn make_padded_frame(w: u32, h: u32, ch: u32) -> Vec { - let stored = stored_channels(ch); - if stored == ch { - return make_synthetic_frame(w, h, ch); +/// A synthetic frame padded to the power-of-two lane count `NlmDenoiser` stores. +pub fn make_padded_frame(width: u32, height: u32, channels: u32) -> Vec { + let stored_ch = stored_channels(channels); + if stored_ch == channels { + return make_synthetic_frame(width, height, channels); } - let src = make_synthetic_frame(w, h, ch); - let mut data = vec![0.0f32; (w * h * stored) as usize]; - for i in 0..(w * h) as usize { - for c in 0..ch as usize { - data[i * stored as usize + c] = src[i * ch as usize + c]; + + let synthetic = make_synthetic_frame(width, height, channels); + let mut data = vec![0.0f32; (width * height * stored_ch) as usize]; + for i in 0..(width * height) as usize { + for channel in 0..channels as usize { + data[i * stored_ch as usize + channel] = synthetic[i * channels as usize + channel]; } } + data } -/// Welsch coefficient for the bench-default parameter set. Channel mode -/// is irrelevant here; `NlmParams::h2_inv_norm` only reads `patch_radius` -/// and `strength`. +/// The Welsch coefficient for the bench parameters. +/// +/// The channel mode stays at its default because the coefficient only depends on +/// `patch_radius` and `strength`. pub fn h2_inv_norm() -> f32 { - NlmParams { + let params = NlmParams { patch_radius: PATCH_RADIUS, ..NlmParams::default() - } - .h2_inv_norm() + }; + + params.h2_inv_norm() } pub fn cube_count_2d() -> CubeCount { - CubeCount::new_2d(W.div_ceil(BLOCK_X), H.div_ceil(BLOCK_Y)) + let cubes_x = WIDTH.div_ceil(BLOCK_X); + let cubes_y = HEIGHT.div_ceil(BLOCK_Y); + + CubeCount::new_2d(cubes_x, cubes_y) } pub fn cube_dim_2d() -> CubeDim { @@ -112,16 +120,15 @@ pub fn cube_dim_2d() -> CubeDim { } pub fn block_sync(client: &ComputeClient) { - cubecl::future::block_on(client.sync()).unwrap(); + let sync = client.sync(); + cubecl::future::block_on(sync).unwrap(); } -pub fn shapes_with_ch(ch: u32) -> Vec> { - vec![vec![W as usize, H as usize, ch as usize]] +pub fn shapes_with_channels(channels: u32) -> Vec> { + vec![vec![WIDTH as usize, HEIGHT as usize, channels as usize]] } -/// Shared input shape for kernels that take one framebuffer and write -/// one output of the same logical size (`dist_2d_weight`, its `_ref` -/// twin, and `bilateral`). +/// Buffers for a kernel that reads one frame and writes one output. #[derive(Clone)] pub struct InputOutput { pub input: Handle, @@ -136,24 +143,31 @@ pub fn print_header() { " {:5} {:>10} {:>10} {:>10} {:>10} {:>10}", "kernel", "samp", "mean", "median", "min", "max", "fps", ); - println!(" {}", "-".repeat(NAME_WIDTH + 6 + 12 * 5)); + + let rule = "-".repeat(NAME_WIDTH + 6 + 12 * 5); + println!(" {}", rule); } pub fn run(bench: B) { let name = bench.name(); match bench.run(TimingMethod::Device) { Ok(durations) => { - let c = BenchmarkComputations::new(&durations); - let mean_s = c.mean.as_secs_f64(); + let computations = BenchmarkComputations::new(&durations); + let mean_s = computations.mean.as_secs_f64(); let fps = if mean_s > 0.0 { 1.0 / mean_s } else { 0.0 }; + let mean = fmt_us(computations.mean); + let median = fmt_us(computations.median); + let min = fmt_us(computations.min); + let max = fmt_us(computations.max); + println!( " {:5} {:>10} {:>10} {:>10} {:>10} {:>10.2}", name, durations.durations.len(), - fmt_us(c.mean), - fmt_us(c.median), - fmt_us(c.min), - fmt_us(c.max), + mean, + median, + min, + max, fps, ); }, @@ -161,11 +175,11 @@ pub fn run(bench: B) { } } -fn fmt_us(d: core::time::Duration) -> String { - let us = d.as_secs_f64() * 1_000_000.0; - if us >= 1000.0 { - format!("{:.3} ms", us / 1000.0) +fn fmt_us(duration: core::time::Duration) -> String { + let micros = duration.as_secs_f64() * 1_000_000.0; + if micros >= 1000.0 { + format!("{:.3} ms", micros / 1000.0) } else { - format!("{us:.2} µs") + format!("{micros:.2} µs") } } diff --git a/av-denoise-core/benches/kernels/mv_regularise.rs b/av-denoise-core/benches/kernels/mv_regularise.rs index 33c17ce..82a3b7d 100644 --- a/av-denoise-core/benches/kernels/mv_regularise.rs +++ b/av-denoise-core/benches/kernels/mv_regularise.rs @@ -1,9 +1,9 @@ -use av_denoise_core::nl4d::kernels::nl4d_mv_regularise; +use av_denoise_core::bench_api::nl4d_kernels::nl4d_mv_regularise; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_synthetic_frame, shapes_with_ch}; +use super::{HEIGHT, WIDTH, block_sync, make_synthetic_frame, shapes_with_channels}; const BLKSIZE: u32 = 16; const STEP: u32 = 8; @@ -29,17 +29,23 @@ impl Benchmark for MvRegulariseBench { type Output = (); fn prepare(&self) -> Self::Input { - let blocks = (W.div_ceil(STEP) * H.div_ceil(STEP)) as usize; - let centre = self - .client - .create_from_slice(f32::as_bytes(&make_synthetic_frame(W, H, 1))); - let neighbour = self - .client - .create_from_slice(f32::as_bytes(&make_synthetic_frame(W, H, 1))); - let field: Vec = (0..2 * blocks).map(|i| (i % 7) as i32 - 3).collect(); - let mv_in = self.client.create_from_slice(i32::as_bytes(&field)); + let blocks_x = WIDTH.div_ceil(STEP); + let blocks_y = HEIGHT.div_ceil(STEP); + let blocks = (blocks_x * blocks_y) as usize; + + let centre_frame = make_synthetic_frame(WIDTH, HEIGHT, 1); + let centre_bytes = f32::as_bytes(¢re_frame); + let centre = self.client.create_from_slice(centre_bytes); + let neighbour_frame = make_synthetic_frame(WIDTH, HEIGHT, 1); + let neighbour_bytes = f32::as_bytes(&neighbour_frame); + let neighbour = self.client.create_from_slice(neighbour_bytes); + + let field: Vec = (0..2 * blocks).map(|index| (index % 7) as i32 - 3).collect(); + let field_bytes = i32::as_bytes(&field); + let mv_in = self.client.create_from_slice(field_bytes); let mv_out = self.client.empty(2 * blocks * size_of::()); let confidence = self.client.empty(blocks * size_of::()); + RegulariseInput { centre, neighbour, @@ -50,31 +56,35 @@ impl Benchmark for MvRegulariseBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let blocks_x = W.div_ceil(STEP); - let blocks_y = H.div_ceil(STEP); + let blocks_x = WIDTH.div_ceil(STEP); + let blocks_y = HEIGHT.div_ceil(STEP); let blocks = (blocks_x * blocks_y) as usize; let block_area = (BLKSIZE * BLKSIZE) as f32; + let lambda = FIELD_LAMBDA * block_area * THSAD_PIXEL; + let thsad = block_area * THSAD_PIXEL; + unsafe { nl4d_mv_regularise::launch_unchecked::( &self.client, CubeCount::new_2d(blocks_x, blocks_y), CubeDim::new_2d(8, 8), - ArrayArg::from_raw_parts(args.centre.clone(), (W * H) as usize), - ArrayArg::from_raw_parts(args.neighbour.clone(), (W * H) as usize), + ArrayArg::from_raw_parts(args.centre.clone(), (WIDTH * HEIGHT) as usize), + ArrayArg::from_raw_parts(args.neighbour.clone(), (WIDTH * HEIGHT) as usize), ArrayArg::from_raw_parts(args.mv_in.clone(), 2 * blocks), ArrayArg::from_raw_parts(args.mv_out.clone(), 2 * blocks), ArrayArg::from_raw_parts(args.confidence.clone(), blocks), - FIELD_LAMBDA * block_area * THSAD_PIXEL, + lambda, 0.0, - block_area * THSAD_PIXEL, - W, - H, + thsad, + WIDTH, + HEIGHT, BLKSIZE, STEP, blocks_x, blocks_y, ); } + Ok(()) } @@ -87,6 +97,6 @@ impl Benchmark for MvRegulariseBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(1) + shapes_with_channels(1) } } diff --git a/av-denoise-core/benches/kernels/nl4d_geometry.rs b/av-denoise-core/benches/kernels/nl4d_geometry.rs index adca878..fa19b10 100644 --- a/av-denoise-core/benches/kernels/nl4d_geometry.rs +++ b/av-denoise-core/benches/kernels/nl4d_geometry.rs @@ -1,10 +1,3 @@ -// The search geometry `nl4d` ships with, in one place. Every value here -// comes from `Nl4dParams::default()`. -// -// `benches/kernels/collab_fused.rs` builds its frame ring and member -// buffers from these, so one set of constants describes what the -// denoiser actually runs. - /// `Nl4dParams::default().temporal_radius`. pub const RADIUS: u32 = 2; /// `Nl4dParams::default().refine`. @@ -16,12 +9,11 @@ pub const K_MAX: u32 = 8; /// `Nl4dParams::default().lambda_ht`. pub const LAMBDA_HT: f32 = 4.158; -/// The motion field's block stride. Held at `collab::PATCH_SIZE` so a -/// block boundary lines up with a patch boundary. +/// The motion field's block stride. +/// +/// It is held at `collab::PATCH_SIZE` so a block boundary lines up with a patch boundary. pub const BLK_STEP: u32 = 8; -/// The library's own default motion block side length -/// (`MotionCompensationMode::Mvtools`'s `blksize`), distinct from -/// [`BLK_STEP`] above. +/// The library's default motion block side length, the `blksize` of `MotionCompensationMode::Mvtools`. pub const BLKSIZE: u32 = 16; /// Frames in the ring a pass reads. @@ -29,11 +21,11 @@ pub const N_FRAMES: u32 = 2 * RADIUS + 1; /// The physical ring slot a pass is centred on. pub const CENTRE_SLOT: u32 = RADIUS; -/// The centre slot is skipped, and physical slots run `0..N_FRAMES`, so -/// `NEIGHBOUR_SLOTS[t]` for the neighbour at temporal offset `k` is -/// `k + RADIUS`, laid out in the same `neighbour_idx_for_k` order -/// `crate::nlmeans::motion::chain` uses: ascending k on the negative -/// side first, then ascending k on the positive side. +/// The physical ring slot of each neighbour, skipping the centre. +/// +/// Slots run `0..N_FRAMES` and the neighbour at temporal offset `k` sits in slot `k + RADIUS`. +/// The order matches `neighbour_idx_for_k`, ascending `k` on the negative side first and then on +/// the positive side. pub const NEIGHBOUR_SLOTS: [u32; (2 * RADIUS) as usize] = [0, 1, 3, 4]; /// Sigma the hard-threshold bench filters at. @@ -41,25 +33,25 @@ pub const SIGMA: f32 = 0.02; /// The motion-field stride one neighbour occupies, in `i32` elements. /// -/// `MotionCtx` pads each neighbour's slice of the motion buffer up to -/// the runtime's buffer-binding alignment, and passes the padded -/// element count to the kernel as a `#[comptime]` stride. A rig that -/// passes the unpadded count compiles the kernel against a stride the -/// pipeline never uses. Pass `client.properties().memory.alignment` as -/// `align`. +/// `MotionCtx` pads each neighbour's slice of the motion buffer up to the runtime's buffer-binding +/// alignment and passes the padded count to the kernel as a `#[comptime]` stride. A rig that +/// passes the unpadded count compiles the kernel against a stride the pipeline never uses. Pass +/// `client.properties().memory.alignment` as `align`. pub fn mv_stride(blocks_x: u32, blocks_y: u32, align: u64) -> u32 { - padded_elems::(blocks_x as u64 * blocks_y as u64 * 2, align) + let elements = blocks_x as u64 * blocks_y as u64 * 2; + padded_elems::(elements, align) } -/// The confidence stride one neighbour occupies, in `f32` elements, -/// padded the way [`mv_stride`] describes. +/// The confidence stride one neighbour occupies, in `f32` elements, padded like [mv_stride]. pub fn conf_stride(blocks_x: u32, blocks_y: u32, align: u64) -> u32 { - padded_elems::(blocks_x as u64 * blocks_y as u64, align) + let elements = blocks_x as u64 * blocks_y as u64; + padded_elems::(elements, align) } -/// `elems` of `T` rounded up so they cover a whole number of `align` -/// byte boundaries. -fn padded_elems(elems: u64, align: u64) -> u32 { - let size = size_of::() as u64; - ((elems * size).next_multiple_of(align) / size) as u32 +/// `elements` of `T` rounded up to cover a whole number of `align`-byte boundaries. +fn padded_elems(elements: u64, align: u64) -> u32 { + let element_size = size_of::() as u64; + let padded_bytes = (elements * element_size).next_multiple_of(align); + + (padded_bytes / element_size) as u32 } diff --git a/av-denoise-core/benches/kernels/noise_partial.rs b/av-denoise-core/benches/kernels/noise_partial.rs index b91d78b..bc9da5e 100644 --- a/av-denoise-core/benches/kernels/noise_partial.rs +++ b/av-denoise-core/benches/kernels/noise_partial.rs @@ -1,4 +1,4 @@ -use av_denoise_core::nlmeans::kernels::{nlm_noise_partial, nlm_noise_reduce}; +use av_denoise_core::bench_api::kernels::{nlm_noise_partial, nlm_noise_reduce}; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; @@ -7,24 +7,24 @@ use super::{ BLOCK_1D, BLOCK_X, BLOCK_Y, - H, - W, + HEIGHT, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, make_padded_frame, - shapes_with_ch, + shapes_with_channels, }; /// Logical channel count for the bench's YUV storage frame. const NOISE_CHANNELS: u32 = 3; -/// Padded storage width for YUV (padded up to a vec4 lane). +/// Padded storage width for YUV, padded up to a vec4 lane. const NOISE_STORED_CH: u32 = 4; -/// Both stages of the Immerkær noise estimate, dispatched back-to-back -/// against a single 1080p YUV frame. `nlm_noise_partial` reduces every -/// `BLOCK_X × BLOCK_Y` cube down to one partial per channel lane, then -/// `nlm_noise_reduce` folds every partial into the frame-level total. +/// Both stages of the Immerkær noise estimate, run back to back on one 1080p YUV frame. +/// +/// `nlm_noise_partial` reduces every `BLOCK_X × BLOCK_Y` cube to one partial per channel lane, +/// then `nlm_noise_reduce` folds every partial into the frame-level total. pub struct NoisePartialBench { pub client: ComputeClient, } @@ -37,7 +37,7 @@ pub struct NoiseInput { } fn partials_len() -> usize { - (W.div_ceil(BLOCK_X) * H.div_ceil(BLOCK_Y) * 4) as usize + (WIDTH.div_ceil(BLOCK_X) * HEIGHT.div_ceil(BLOCK_Y) * 4) as usize } impl Benchmark for NoisePartialBench { @@ -45,10 +45,13 @@ impl Benchmark for NoisePartialBench { type Output = (); fn prepare(&self) -> Self::Input { - let frame = make_padded_frame(W, H, NOISE_CHANNELS); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); - let partials = self.client.empty(partials_len() * size_of::()); + let frame = make_padded_frame(WIDTH, HEIGHT, NOISE_CHANNELS); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); + let partial_lanes = partials_len(); + let partials = self.client.empty(partial_lanes * size_of::()); let results = self.client.empty(4 * size_of::()); + NoiseInput { input, partials, @@ -57,21 +60,23 @@ impl Benchmark for NoisePartialBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let total_input = (W * H * NOISE_STORED_CH) as usize; - let n_partials = partials_len(); - let num_partials = (n_partials / 4) as u32; + let total_input = (WIDTH * HEIGHT * NOISE_STORED_CH) as usize; + let partial_lanes = partials_len(); + let partial_count = (partial_lanes / 4) as u32; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); unsafe { nlm_noise_partial::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), + cube_count, + cube_dim, NOISE_STORED_CH as usize, ArrayArg::from_raw_parts(args.input.clone(), total_input), - ArrayArg::from_raw_parts(args.partials.clone(), n_partials), + ArrayArg::from_raw_parts(args.partials.clone(), partial_lanes), 0u32, - W, - H, + WIDTH, + HEIGHT, NOISE_CHANNELS, BLOCK_X, BLOCK_Y, @@ -83,10 +88,10 @@ impl Benchmark for NoisePartialBench { &self.client, CubeCount::new_1d(1), CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(args.partials.clone(), n_partials), + ArrayArg::from_raw_parts(args.partials.clone(), partial_lanes), ArrayArg::from_raw_parts(args.results.clone(), 4), 0u32, - num_partials, + partial_count, BLOCK_1D, ); } @@ -103,6 +108,6 @@ impl Benchmark for NoisePartialBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(NOISE_CHANNELS) + shapes_with_channels(NOISE_CHANNELS) } } diff --git a/av-denoise-core/benches/kernels/pack_wire.rs b/av-denoise-core/benches/kernels/pack_wire.rs deleted file mode 100644 index 8f53f3c..0000000 --- a/av-denoise-core/benches/kernels/pack_wire.rs +++ /dev/null @@ -1,82 +0,0 @@ -use av_denoise_core::Depth; -use av_denoise_core::nlmeans::kernels::gpu_pack_wire; -use cubecl::benchmark::Benchmark; -use cubecl::prelude::*; -use cubecl::server::Handle; - -use super::{BLOCK_1D, COPY_GRID_1D, H, W, block_sync, make_padded_frame, shapes_with_ch, stored_channels}; - -#[derive(Clone)] -pub struct PackWireInput { - src: Handle, - dst: Handle, -} - -pub struct PackWireBench { - pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, - /// Covers both wire codecs, 4 samples per word at 8-bit and 2 above it. - pub depth: Depth, -} - -impl PackWireBench { - fn samples(&self) -> u32 { - W * H * self.ch - } - - fn words(&self) -> u32 { - self.samples().div_ceil(self.depth.wire_pack().samples_per_word()) - } -} - -impl Benchmark for PackWireBench { - type Input = PackWireInput; - type Output = (); - - fn prepare(&self) -> Self::Input { - let frame = make_padded_frame(W, H, self.ch); - let src = self.client.create_from_slice(f32::as_bytes(&frame)); - let dst = self.client.empty(self.words() as usize * size_of::()); - PackWireInput { src, dst } - } - - fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = W * H; - let stored_ch = stored_channels(self.ch); - // One `Depth` yields both, so the bench cannot pair a scale with - // the wrong lane count the way hand-written literals could. - let pack = self.depth.wire_pack(); - let total_threads = COPY_GRID_1D * BLOCK_1D; - - unsafe { - gpu_pack_wire::launch_unchecked::( - &self.client, - CubeCount::new_1d(COPY_GRID_1D), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(args.src.clone(), (pixels * stored_ch) as usize), - ArrayArg::from_raw_parts(args.dst.clone(), self.words() as usize), - pack.max(), - pixels, - self.ch, - stored_ch, - self.ch, - false, - pack.samples_per_word(), - self.words(), - total_threads, - ); - } - Ok(()) - } - - fn name(&self) -> String { - format!("gpu_pack_wire_1080p_{:?}_{}", self.depth, self.ch_name) - } - fn sync(&self) { - block_sync(&self.client); - } - fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) - } -} diff --git a/av-denoise-core/benches/kernels/temporal_noise_stats.rs b/av-denoise-core/benches/kernels/temporal_noise_stats.rs index 3fc0dc9..1e33fde 100644 --- a/av-denoise-core/benches/kernels/temporal_noise_stats.rs +++ b/av-denoise-core/benches/kernels/temporal_noise_stats.rs @@ -1,23 +1,21 @@ -use av_denoise_core::nlmeans::kernels::nlm_temporal_noise_stats; +use av_denoise_core::bench_api::kernels::nlm_temporal_noise_stats; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{H, W, block_sync, make_padded_frame, shapes_with_ch}; +use super::{HEIGHT, WIDTH, block_sync, make_padded_frame, shapes_with_channels}; /// Logical channel count for the bench's YUV storage frame. const TEMPORAL_CHANNELS: u32 = 3; -/// Padded storage width for YUV (padded up to a vec4 lane). +/// Padded storage width for YUV, padded up to a vec4 lane. const TEMPORAL_STORED_CH: u32 = 4; /// Matches `nlmeans::noise::TEMPORAL_NOISE_BLOCK`. const TEMPORAL_BLOCK: u32 = 16; -/// The temporal-residual noise-stats kernel, diffing two 1080p YUV -/// ring slots against each other and reducing every `16 × 16` block -/// into its stats record. +/// The temporal-residual noise-stats kernel over two 1080p YUV ring slots. /// -/// `luma_fields` picks which of the kernel's two compiled variants this -/// row times, matching `nlm_temporal_noise_stats`'s own flag. +/// It diffs the slots and reduces every `16 × 16` block into its stats record. `luma_fields` +/// picks which of the kernel's two compiled variants this row times. pub struct TemporalNoiseStatsBench { pub client: ComputeClient, pub luma_fields: bool, @@ -30,16 +28,16 @@ pub struct TemporalNoiseStatsInput { } fn blocks() -> (u32, u32) { - (W.div_ceil(TEMPORAL_BLOCK), H.div_ceil(TEMPORAL_BLOCK)) + (WIDTH.div_ceil(TEMPORAL_BLOCK), HEIGHT.div_ceil(TEMPORAL_BLOCK)) } -/// Matches `nlmeans::noise::temporal_stats_record_len`. +/// The stats buffer length, one `nlmeans::noise::temporal_stats_record_len` record per block. /// -/// That is a sum and a sum of squares per stored channel, one lag-1 -/// total, and four quarter records of nine fields each. +/// A record is a sum and a sum of squares per stored channel, one lag-1 total, and four quarter +/// records of nine fields each. fn stats_len() -> usize { - let (bx, by) = blocks(); - (bx * by * (2 * TEMPORAL_STORED_CH + 37)) as usize + let (blocks_x, blocks_y) = blocks(); + (blocks_x * blocks_y * (2 * TEMPORAL_STORED_CH + 37)) as usize } impl Benchmark for TemporalNoiseStatsBench { @@ -47,18 +45,23 @@ impl Benchmark for TemporalNoiseStatsBench { type Output = (); fn prepare(&self) -> Self::Input { - // Two ring slots: a frame and a slightly perturbed copy, so the - // diff the kernel reduces isn't degenerately zero everywhere. - let frame = make_padded_frame(W, H, TEMPORAL_CHANNELS); + // Two ring slots, a frame and a slightly perturbed copy, so the diff the kernel reduces is + // not zero everywhere. + let frame = make_padded_frame(WIDTH, HEIGHT, TEMPORAL_CHANNELS); let mut ring = frame.clone(); - ring.extend(frame.iter().map(|&v| (v + 0.01).clamp(0.0, 1.0))); - let input = self.client.create_from_slice(f32::as_bytes(&ring)); - let stats = self.client.empty(stats_len() * size_of::()); + let perturbed = frame.iter().map(|&sample| (sample + 0.01).clamp(0.0, 1.0)); + ring.extend(perturbed); + + let ring_bytes = f32::as_bytes(&ring); + let input = self.client.create_from_slice(ring_bytes); + let total_stats = stats_len(); + let stats = self.client.empty(total_stats * size_of::()); + TemporalNoiseStatsInput { input, stats } } fn execute(&self, args: Self::Input) -> Result<(), String> { - let total_input = (2 * W * H * TEMPORAL_STORED_CH) as usize; + let total_input = (2 * WIDTH * HEIGHT * TEMPORAL_STORED_CH) as usize; let (blocks_x, blocks_y) = blocks(); let total_stats = stats_len(); @@ -72,8 +75,8 @@ impl Benchmark for TemporalNoiseStatsBench { ArrayArg::from_raw_parts(args.stats.clone(), total_stats), 1u32, 0u32, - W, - H, + WIDTH, + HEIGHT, TEMPORAL_STORED_CH, TEMPORAL_BLOCK, self.luma_fields, @@ -96,6 +99,6 @@ impl Benchmark for TemporalNoiseStatsBench { } fn shapes(&self) -> Vec> { - shapes_with_ch(TEMPORAL_CHANNELS) + shapes_with_channels(TEMPORAL_CHANNELS) } } diff --git a/av-denoise-core/benches/kernels/unpack_wire.rs b/av-denoise-core/benches/kernels/unpack_wire.rs deleted file mode 100644 index b35ba89..0000000 --- a/av-denoise-core/benches/kernels/unpack_wire.rs +++ /dev/null @@ -1,106 +0,0 @@ -use av_denoise_core::Depth; -use av_denoise_core::nlmeans::kernels::gpu_unpack_wire; -use cubecl::benchmark::Benchmark; -use cubecl::prelude::*; -use cubecl::server::Handle; - -use super::{BLOCK_1D, COPY_GRID_1D, H, W, block_sync, shapes_with_ch, stored_channels}; - -#[derive(Clone)] -pub struct UnpackWireInput { - src: Handle, - dst: Handle, -} - -pub struct UnpackWireBench { - pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, - /// Covers both wire codecs, 4 samples per word at 8-bit and 2 above it. - pub depth: Depth, -} - -impl UnpackWireBench { - fn samples(&self) -> u32 { - W * H * self.ch - } - - fn words(&self) -> u32 { - self.samples().div_ceil(self.depth.wire_pack().samples_per_word()) - } - - fn elements(&self) -> u32 { - W * H * stored_channels(self.ch) - } - - /// Wire bytes whose samples spread over the depth's whole range, so - /// the run measures a realistic spread of values rather than every - /// lane in a wave decoding to the same one. - fn wire(&self) -> Vec { - let bytes = self.depth.bytes_per_sample(); - let mask = (1u32 << self.depth.bits()) - 1; - let mut wire = vec![0u8; self.words() as usize * size_of::()]; - - // An xorshift rather than a multiply, so the samples are not all - // even and the fill reaches every value the depth can express. - for (i, chunk) in wire.chunks_exact_mut(bytes).enumerate() { - let mut hash = i as u32 ^ 0x9E37_79B9; - hash ^= hash << 13; - hash ^= hash >> 17; - hash ^= hash << 5; - chunk.copy_from_slice(&(hash & mask).to_le_bytes()[..bytes]); - } - - wire - } -} - -impl Benchmark for UnpackWireBench { - type Input = UnpackWireInput; - type Output = (); - - fn prepare(&self) -> Self::Input { - let wire = self.wire(); - let src = self.client.create_from_slice(&wire); - let dst = self.client.empty(self.elements() as usize * size_of::()); - UnpackWireInput { src, dst } - } - - fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = W * H; - let stored_ch = stored_channels(self.ch); - // One `Depth` yields both, so the bench cannot pair a scale with - // the wrong lane count the way hand-written literals could. - let pack = self.depth.wire_pack(); - let total_threads = COPY_GRID_1D * BLOCK_1D; - - unsafe { - gpu_unpack_wire::launch_unchecked::( - &self.client, - CubeCount::new_1d(COPY_GRID_1D), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(args.src.clone(), self.words() as usize), - ArrayArg::from_raw_parts(args.dst.clone(), self.elements() as usize), - pack.max(), - 0u32, - pixels, - self.ch, - stored_ch, - pack.samples_per_word(), - self.elements(), - total_threads, - ); - } - Ok(()) - } - - fn name(&self) -> String { - format!("gpu_unpack_wire_1080p_{:?}_{}", self.depth, self.ch_name) - } - fn sync(&self) { - block_sync(&self.client); - } - fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) - } -} diff --git a/av-denoise-core/benches/kernels/vertical_weight.rs b/av-denoise-core/benches/kernels/vertical_weight.rs index 32a0b0f..ed0e3d6 100644 --- a/av-denoise-core/benches/kernels/vertical_weight.rs +++ b/av-denoise-core/benches/kernels/vertical_weight.rs @@ -1,9 +1,19 @@ -use av_denoise_core::nlmeans::kernels::nlm_vertical_weight; +use av_denoise_core::bench_api::kernels::nlm_vertical_weight; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use super::horizontal_sum::HSumInput; -use super::{BLOCK_X, BLOCK_Y, H, PATCH_RADIUS, W, block_sync, cube_count_2d, cube_dim_2d, h2_inv_norm}; +use super::{ + BLOCK_X, + BLOCK_Y, + HEIGHT, + PATCH_RADIUS, + WIDTH, + block_sync, + cube_count_2d, + cube_dim_2d, + h2_inv_norm, +}; pub struct VWeightBench { pub client: ComputeClient, @@ -14,41 +24,50 @@ impl Benchmark for VWeightBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; + let pixels = (WIDTH * HEIGHT) as usize; let data = vec![0.5f32; pixels]; - let input = self.client.create_from_slice(f32::as_bytes(&data)); + let data_bytes = f32::as_bytes(&data); + let input = self.client.create_from_slice(data_bytes); let output = self.client.empty(pixels * size_of::()); + HSumInput { input, output } } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_vertical_weight::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), + cube_count, + cube_dim, ArrayArg::from_raw_parts(args.input.clone(), pixels), ArrayArg::from_raw_parts(args.output.clone(), pixels), - h2_inv_norm(), + inv_norm, 0.0f32, - W, - H, + WIDTH, + HEIGHT, PATCH_RADIUS, BLOCK_X, BLOCK_Y, ); } + Ok(()) } fn name(&self) -> String { "vertical_weight_1080p".to_string() } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - vec![vec![W as usize, H as usize]] + vec![vec![WIDTH as usize, HEIGHT as usize]] } } diff --git a/av-denoise-core/benches/kernels/vweight_pair_accumulate.rs b/av-denoise-core/benches/kernels/vweight_pair_accumulate.rs index e9eb31b..4ab40f3 100644 --- a/av-denoise-core/benches/kernels/vweight_pair_accumulate.rs +++ b/av-denoise-core/benches/kernels/vweight_pair_accumulate.rs @@ -1,4 +1,4 @@ -use av_denoise_core::nlmeans::kernels::nlm_vweight_pair_accumulate; +use av_denoise_core::bench_api::kernels::nlm_vweight_pair_accumulate; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; @@ -6,17 +6,17 @@ use cubecl::server::Handle; use super::{ BLOCK_X, BLOCK_Y, - H, + HEIGHT, PATCH_RADIUS, Q_X, Q_Y, - W, + WIDTH, block_sync, cube_count_2d, cube_dim_2d, h2_inv_norm, make_padded_frame, - shapes_with_ch, + shapes_with_channels, stored_channels, }; @@ -34,8 +34,8 @@ pub struct VWeightPairAccInput { pub struct VWeightPairAccBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for VWeightPairAccBench { @@ -43,17 +43,20 @@ impl Benchmark for VWeightPairAccBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; - let frame = make_padded_frame(W, H, self.ch); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let frame = make_padded_frame(WIDTH, HEIGHT, self.channels); let hsum = vec![0.5f32; pixels]; - let hsum_fwd = self.client.create_from_slice(f32::as_bytes(&hsum)); - let hsum_bwd = self.client.create_from_slice(f32::as_bytes(&hsum)); - let input = self.client.create_from_slice(f32::as_bytes(&frame)); - let accum = self.client.empty(pixels * stored * size_of::()); + let hsum_bytes = f32::as_bytes(&hsum); + let hsum_fwd = self.client.create_from_slice(hsum_bytes); + let hsum_bwd = self.client.create_from_slice(hsum_bytes); + let frame_bytes = f32::as_bytes(&frame); + let input = self.client.create_from_slice(frame_bytes); + let accum = self.client.empty(pixels * stored_ch * size_of::()); let weight_sum = self.client.empty(pixels * size_of::()); let max_weight = self.client.empty(pixels * size_of::()); let confidence_dummy = self.client.empty(size_of::()); + VWeightPairAccInput { hsum_fwd, hsum_bwd, @@ -67,18 +70,22 @@ impl Benchmark for VWeightPairAccBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let cube_count = cube_count_2d(); + let cube_dim = cube_dim_2d(); + let inv_norm = h2_inv_norm(); + unsafe { nlm_vweight_pair_accumulate::launch_unchecked::( &self.client, - cube_count_2d(), - cube_dim_2d(), - stored, + cube_count, + cube_dim, + stored_ch, ArrayArg::from_raw_parts(args.hsum_fwd.clone(), pixels), ArrayArg::from_raw_parts(args.hsum_bwd.clone(), pixels), ArrayArg::from_raw_parts(args.input.clone(), args.frame_len), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.max_weight.clone(), pixels), ArrayArg::from_raw_parts(args.confidence_dummy.clone(), 1), @@ -88,10 +95,10 @@ impl Benchmark for VWeightPairAccBench { 0u32, Q_X, Q_Y, - h2_inv_norm(), + inv_norm, 0.0f32, - W, - H, + WIDTH, + HEIGHT, PATCH_RADIUS, BLOCK_X, BLOCK_Y, @@ -100,16 +107,19 @@ impl Benchmark for VWeightPairAccBench { 1u32, ); } + Ok(()) } fn name(&self) -> String { - format!("vweight_pair_accumulate_1080p_{}", self.ch_name) + format!("vweight_pair_accumulate_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/kernels/zero.rs b/av-denoise-core/benches/kernels/zero.rs index 4e9275e..963fec5 100644 --- a/av-denoise-core/benches/kernels/zero.rs +++ b/av-denoise-core/benches/kernels/zero.rs @@ -1,9 +1,9 @@ -use av_denoise_core::nlmeans::kernels::gpu_zero_buffers; +use av_denoise_core::bench_api::kernels::gpu_zero_buffers; use cubecl::benchmark::Benchmark; use cubecl::prelude::*; use cubecl::server::Handle; -use super::{BLOCK_1D, COPY_GRID_1D, H, W, block_sync, shapes_with_ch, stored_channels}; +use super::{BLOCK_1D, COPY_GRID_1D, HEIGHT, WIDTH, block_sync, shapes_with_channels, stored_channels}; #[derive(Clone)] pub struct ZeroInput { @@ -14,8 +14,8 @@ pub struct ZeroInput { pub struct ZeroBench { pub client: ComputeClient, - pub ch: u32, - pub ch_name: &'static str, + pub channels: u32, + pub channel_name: &'static str, } impl Benchmark for ZeroBench { @@ -23,11 +23,12 @@ impl Benchmark for ZeroBench { type Output = (); fn prepare(&self) -> Self::Input { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; - let accum = self.client.empty(pixels * stored * size_of::()); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; + let accum = self.client.empty(pixels * stored_ch * size_of::()); let weight_sum = self.client.empty(pixels * size_of::()); let max_weight = self.client.empty(pixels * size_of::()); + ZeroInput { accum, weight_sum, @@ -36,32 +37,36 @@ impl Benchmark for ZeroBench { } fn execute(&self, args: Self::Input) -> Result<(), String> { - let pixels = (W * H) as usize; - let stored = stored_channels(self.ch) as usize; + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(self.channels) as usize; let total_threads = COPY_GRID_1D * BLOCK_1D; + unsafe { gpu_zero_buffers::launch_unchecked::( &self.client, CubeCount::new_1d(COPY_GRID_1D), CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored), + ArrayArg::from_raw_parts(args.accum.clone(), pixels * stored_ch), ArrayArg::from_raw_parts(args.weight_sum.clone(), pixels), ArrayArg::from_raw_parts(args.max_weight.clone(), pixels), - (pixels * stored) as u32, + (pixels * stored_ch) as u32, pixels as u32, total_threads, ); } + Ok(()) } fn name(&self) -> String { - format!("gpu_zero_buffers_1080p_{}", self.ch_name) + format!("gpu_zero_buffers_1080p_{}", self.channel_name) } + fn sync(&self) { block_sync(&self.client); } + fn shapes(&self) -> Vec> { - shapes_with_ch(self.ch) + shapes_with_channels(self.channels) } } diff --git a/av-denoise-core/benches/mc_accuracy.rs b/av-denoise-core/benches/mc_accuracy.rs index 0e76c5f..d34a52c 100644 --- a/av-denoise-core/benches/mc_accuracy.rs +++ b/av-denoise-core/benches/mc_accuracy.rs @@ -1,22 +1,23 @@ -//! Scores nl4d's motion field against synthetic clips with known -//! motion. Prints one table per arm. +//! Scores nl4d's motion field against synthetic clips with known motion //! -//! Run with `cargo bench -p av-denoise-core --bench mc_accuracy -- -//! --device discrete:1 --still brick=/path/to/brick.pgm --still -//! asterisk=/path/to/asterisk.pgm`. With no `--still` it runs on a -//! synthetic texture and says so. +//! Prints one table per arm. With no `--still` it runs on a synthetic texture and says so. +//! +//! ```text +//! cargo bench -p av-denoise-core --bench mc_accuracy -- \ +//! --device discrete:1 --still brick=/path/to/brick.pgm --still asterisk=/path/to/asterisk.pgm +//! ``` use std::path::PathBuf; -use av_denoise_core::nl4d::harness::{Clip, KindScore, MotionClass, Score, Still, score, synthesise}; -use av_denoise_core::nl4d::{Nl4dDenoiser, Nl4dParams}; -use av_denoise_core::nlmeans::{ChannelMode, MotionCompensationMode, NlmParams}; +use av_denoise_core::bench_api::harness::{Clip, KindScore, MotionClass, Score, Still, score, synthesise}; +use av_denoise_core::bench_api::{Device, HostIo, Nl4dDenoiser, Nl4dParams, NlmParams}; +use av_denoise_core::{ChannelMode, MotionCompensationMode}; +use clap::Parser; use cubecl::prelude::*; /// Grain levels on the 8-bit scale. const GRAIN: [f32; 3] = [2.0, 6.0, 12.0]; -/// A named still. struct NamedStill { name: String, still: Still, @@ -28,23 +29,25 @@ struct Arm { params: fn() -> Nl4dParams, } -/// `Nl4dParams::default` carries `ChannelMode::Yuv`, which expects -/// three interleaved planes per pushed frame. The harness only ever -/// synthesises a single luma plane, so every arm here switches to -/// `ChannelMode::Luma` instead. +/// The default parameters switched to `ChannelMode::Luma`. +/// +/// `Nl4dParams::default` carries `ChannelMode::Yuv`, which expects three interleaved planes per +/// pushed frame, and the harness only synthesises a single luma plane. fn baseline_params() -> Nl4dParams { + let nlm = NlmParams { + channels: ChannelMode::Luma, + ..Nl4dParams::default().nlm + }; + Nl4dParams { - nlm: NlmParams { - channels: ChannelMode::Luma, - ..Nl4dParams::default().nlm - }, + nlm, ..Nl4dParams::default() } } -/// `baseline_params` with `field_lambda` overridden, for context against -/// the shipped default. Building from `Nl4dParams::default()` directly -/// would panic with the three-channel default the harness cannot feed. +/// [baseline_params] with `field_lambda` at 0.5, for context against the shipped default. +/// +/// It builds on the luma baseline because the three-channel default would panic in the harness. fn with_lambda_0_5() -> Nl4dParams { Nl4dParams { field_lambda: 0.5, @@ -52,14 +55,16 @@ fn with_lambda_0_5() -> Nl4dParams { } } -/// `baseline_params` with the motion pyramid deepened to three levels, -/// to test whether the extra level earns its added kernel launch. +/// [baseline_params] with the motion pyramid deepened to three levels. +/// +/// It tests whether the extra level earns its added kernel launch. fn with_pyramid_3() -> Nl4dParams { - let mut p = baseline_params(); - if let MotionCompensationMode::Mvtools { pyramid_levels, .. } = &mut p.nlm.motion_compensation { + let mut params = baseline_params(); + if let MotionCompensationMode::Mvtools { pyramid_levels, .. } = &mut params.nlm.motion_compensation { *pyramid_levels = 3; } - p + + params } fn arms() -> Vec { @@ -83,54 +88,68 @@ fn parse_still(spec: &str) -> Result { let (name, path) = spec .split_once('=') .ok_or_else(|| format!("--still expects name=path, got {spec}"))?; - let bytes = std::fs::read(PathBuf::from(path)).map_err(|e| format!("{path}: {e}"))?; + let path_buf = PathBuf::from(path); + let bytes = std::fs::read(path_buf).map_err(|err| format!("{path}: {err}"))?; + let still = Still::from_pgm(&bytes)?; + Ok(NamedStill { name: name.to_string(), - still: Still::from_pgm(&bytes)?, + still, }) } fn run_clip(client: &ComputeClient, params: Nl4dParams, clip: &Clip) -> Score { let refine = params.refine; - let mut d = Nl4dDenoiser::::new(client, params, clip.width, clip.height).expect("construction failed"); + let mut denoiser = + Nl4dDenoiser::::new(client, params, clip.width, clip.height).expect("construction failed"); for frame in &clip.frames { - d.push_frame(frame); - let _ = d.denoise_submit().expect("denoise_submit failed"); + denoiser.push_frame(frame); + let _ = denoiser.denoise().expect("denoise failed"); } - let snap = d.motion_snapshot().expect("a pass ran once the window filled"); - score(clip, &snap, refine) + + let snapshot = denoiser + .motion_snapshot() + .expect("a pass ran once the window filled"); + + score(clip, &snapshot, refine) } -fn print_kind(label: &str, k: &KindScore) { - if k.patches == 0 { +fn print_kind(label: &str, kind: &KindScore) { + if kind.patches == 0 { return; } + + let corner_rate = 100.0 * kind.in_window_rate_corner(); + let covering_rate = 100.0 * kind.in_window_rate_covering(); + let epe_mean = kind.epe_mean(); + let epe_p95 = kind.epe_p95(); + let confidence = kind.confidence_median(); + println!( " {label:<9} {:>6} corner {:>5.1}% covering {:>5.1}% epe {:>5.2} / p95 {:>5.2} conf {:>4.2}", - k.patches, - 100.0 * k.in_window_rate_corner(), - 100.0 * k.in_window_rate_covering(), - k.epe_mean(), - k.epe_p95(), - k.confidence_median(), + kind.patches, corner_rate, covering_rate, epe_mean, epe_p95, confidence, ); } fn run_all(device: &R::Device, stills: &[NamedStill]) { let client = R::client(device); + for arm in arms() { println!(); println!("=== arm: {} ===", arm.name); + for still in stills { for class in MotionClass::ALL { for grain in GRAIN { let params = (arm.params)(); let clip = synthesise(&still.still, class, params.temporal_radius, grain / 255.0, 7); - let s = run_clip::(&client, params, &clip); - println!(" {:<10} {:<9} grain {grain:>4.0}", still.name, class.label()); - print_kind("plain", &s.plain); - print_kind("boundary", &s.boundary); - print_kind("occluded", &s.occluded); + let clip_score = run_clip::(&client, params, &clip); + let class_label = class.label(); + + println!(" {:<10} {:<9} grain {grain:>4.0}", still.name, class_label); + print_kind("plain", &clip_score.plain); + print_kind("boundary", &clip_score.boundary); + print_kind("occluded", &clip_score.occluded); } } } @@ -140,34 +159,35 @@ fn run_all(device: &R::Device, stills: &[NamedStill]) { #[derive(clap::Parser, Debug)] #[command(about = "Motion-field accuracy against synthetic known-motion clips", long_about = None)] struct Cli { - /// GPU device to bind to. Format: `default`, `discrete[:N]`, - /// `integrated[:N]`, `virtual[:N]`, or `cpu`. + /// GPU device to bind to, one of `default`, `discrete[:N]`, `integrated[:N]`, `virtual[:N]` or `cpu`. #[arg(long, default_value = "default")] - device: av_denoise_core::Device, + device: Device, /// A still to build clips from, as `name=path.pgm`. Repeatable. #[arg(long = "still")] stills: Vec, - /// Swallowed. Cargo passes this when invoking the bench binary. + /// Swallowed, since cargo passes this when invoking the bench binary. #[arg(long, hide = true)] bench: bool, } fn main() { - use clap::Parser; let cli = Cli::parse(); let stills: Vec = if cli.stills.is_empty() { println!("no --still given, running on a synthetic 256x256 texture"); - vec![NamedStill { + let still = Still::synthetic(256, 256); + let synthetic = NamedStill { name: "synthetic".to_string(), - still: Still::synthetic(256, 256), - }] + still, + }; + + vec![synthetic] } else { cli.stills .iter() - .map(|s| parse_still(s).unwrap_or_else(|e| panic!("{e}"))) + .map(|spec| parse_still(spec).unwrap_or_else(|err| panic!("{err}"))) .collect() }; diff --git a/av-denoise-core/benches/motion.rs b/av-denoise-core/benches/motion.rs index 0b04d74..9597712 100644 --- a/av-denoise-core/benches/motion.rs +++ b/av-denoise-core/benches/motion.rs @@ -1,17 +1,3 @@ -use std::hint::black_box; -use std::time::{Duration, Instant}; - -use av_denoise_core::nlmeans::{ - ChannelMode, - MotionCompensationMode, - MotionEstimation, - NlmDenoiser, - NlmParams, - Pending, - PrefilterMode, -}; -use cubecl::prelude::*; - #[expect( dead_code, reason = "the shared kernel module is included by several bench binaries, each of which uses \ @@ -20,23 +6,28 @@ use cubecl::prelude::*; #[path = "kernels/mod.rs"] mod kernels; +use std::hint::black_box; +use std::time::{Duration, Instant}; + +use av_denoise_core::bench_api::{Device, HostIo, NlmDenoiser, NlmParams, start_read, wait_read}; +use av_denoise_core::{ChannelMode, MotionCompensationMode, MotionEstimation, PrefilterMode}; +use clap::Parser; +use cubecl::prelude::*; use kernels::mc_block_match_coarse::BlockMatchCoarseBench; use kernels::mc_block_match_fine::BlockMatchFineBench; use kernels::mc_confidence::McConfidenceBench; use kernels::mc_downscale::DownscaleBench; use kernels::mc_warp::WarpBench; use kernels::mv_regularise::MvRegulariseBench; -use kernels::{CHANNELS, print_header, run}; +use kernels::{CHANNELS, make_synthetic_frame, print_header, run}; -const W: u32 = 1920; -const H: u32 = 1080; +const WIDTH: u32 = 1920; +const HEIGHT: u32 = 1080; const WARMUP_PIPELINE: usize = 2; const ITERS_PIPELINE: usize = 200; -fn make_synthetic_frame(w: u32, h: u32, ch: u32) -> Vec { - kernels::make_synthetic_frame(w, h, ch) -} +const TEMPORAL_RADIUS: u32 = 1; struct BenchResult { name: String, @@ -64,20 +55,22 @@ fn run_pipeline_bench( client: &ComputeClient, warmup: usize, iterations: usize, - mut f: impl FnMut(), + mut step: impl FnMut(), ) -> BenchResult { for _ in 0..warmup { - f(); - futures::executor::block_on(client.sync()).unwrap(); + step(); + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); } let mut times = Vec::with_capacity(iterations); - for _ in 0..iterations { let start = Instant::now(); - f(); - futures::executor::block_on(client.sync()).unwrap(); - times.push(start.elapsed()); + step(); + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); + let elapsed = start.elapsed(); + times.push(elapsed); } let total: Duration = times.iter().sum(); @@ -97,9 +90,11 @@ fn run_pipeline_bench( } } -const TEMPORAL_RADIUS: u32 = 1; - -fn temporal_params(radius: u32, channels: ChannelMode, mc: MotionCompensationMode) -> NlmParams { +fn temporal_params( + radius: u32, + channels: ChannelMode, + motion_compensation: MotionCompensationMode, +) -> NlmParams { NlmParams { temporal_radius: radius, search_radius: 2, @@ -108,7 +103,7 @@ fn temporal_params(radius: u32, channels: ChannelMode, mc: MotionCompensationMod self_weight: 1.0, channels, prefilter: PrefilterMode::None, - motion_compensation: mc, + motion_compensation, hq: None, } } @@ -123,9 +118,7 @@ fn mc_default() -> MotionCompensationMode { } } -/// Same block geometry as `mc_default`, but with `Chained` estimation -/// at the library's default refinement radius. Used by the direct vs -/// chained throughput comparison below. +/// [mc_default]'s block geometry with `Chained` estimation at the library's default refinement radius. fn mc_chained_default() -> MotionCompensationMode { MotionCompensationMode::Mvtools { blksize: 16, @@ -136,80 +129,87 @@ fn mc_chained_default() -> MotionCompensationMode { } } -/// Eager temporal pipeline (push → denoise → wait inline). Submits -/// one frame and blocks on its readback before the next push, so the -/// per-frame number is the full critical-path cost. +/// The eager temporal pipeline cost. +/// +/// Each frame is pushed, denoised and waited on inline before the next push, so the per-frame +/// number is the full critical-path cost. fn bench_eager( client: &ComputeClient, backend: &str, radius: u32, channels: ChannelMode, - ch_name: &str, - mc: MotionCompensationMode, + channel_name: &str, + motion_compensation: MotionCompensationMode, tag: &str, ) -> BenchResult { - let ch = channels.count(); - let params = temporal_params(radius, channels, mc); - let frame = make_synthetic_frame(W, H, ch); + let channel_count = channels.count(); + let params = temporal_params(radius, channels, motion_compensation); + let frame = make_synthetic_frame(WIDTH, HEIGHT, channel_count); let total_frames = 1 + 2 * params.temporal_radius as usize; - let name = format!("denoise_temporal{tag}_1080p_{ch_name}"); + let name = format!("denoise_temporal{tag}_1080p_{channel_name}"); - let mut denoiser = NlmDenoiser::::new(client, params, W, H); + let mut denoiser = NlmDenoiser::::new(client, params, WIDTH, HEIGHT); for _ in 0..total_frames - 1 { denoiser.push_frame(&frame); } - futures::executor::block_on(client.sync()).unwrap(); + + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); run_pipeline_bench(&name, backend, client, WARMUP_PIPELINE, ITERS_PIPELINE, || { denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser"); + let result = denoiser.denoise().unwrap().unwrap(); black_box(&result); }) } -/// Pipelined variant: kernels for frame N+1 are submitted before the -/// readback for frame N completes, so GPU and host overlap. Mirrors -/// the equivalent helper in `benches/nlmeans.rs`. +/// The pipelined temporal pipeline cost. +/// +/// Frame N+1's kernels are submitted before frame N's readback completes, so GPU and host work +/// overlap. fn bench_pipelined( client: &ComputeClient, backend: &str, radius: u32, channels: ChannelMode, - ch_name: &str, - mc: MotionCompensationMode, + channel_name: &str, + motion_compensation: MotionCompensationMode, tag: &str, ) -> BenchResult { - let ch = channels.count(); - let params = temporal_params(radius, channels, mc); - let frame = make_synthetic_frame(W, H, ch); + let channel_count = channels.count(); + let params = temporal_params(radius, channels, motion_compensation); + let frame = make_synthetic_frame(WIDTH, HEIGHT, channel_count); let total_frames = 1 + 2 * params.temporal_radius as usize; - let name = format!("denoise_temporal_pipelined{tag}_1080p_{ch_name}"); + let name = format!("denoise_temporal_pipelined{tag}_1080p_{channel_name}"); - let mut denoiser = NlmDenoiser::::new(client, params, W, H); + let mut denoiser = NlmDenoiser::::new(client, params, WIDTH, HEIGHT); for _ in 0..total_frames - 1 { denoiser.push_frame(&frame); } - futures::executor::block_on(client.sync()).unwrap(); + + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); denoiser.push_frame(&frame); - let mut in_flight: Option> = Some(denoiser.denoise_submit().unwrap().unwrap()); + let first = denoiser.denoise_submit_gpu().unwrap().unwrap(); + let first_read = start_read(client, first.handle); + let mut in_flight = Some(first_read); let result = run_pipeline_bench(&name, backend, client, WARMUP_PIPELINE, ITERS_PIPELINE, || { denoiser.push_frame(&frame); - let next = denoiser.denoise_submit().unwrap().unwrap(); - let output = in_flight.take().unwrap().wait().unwrap(); + let next = denoiser.denoise_submit_gpu().unwrap().unwrap(); + let next_read = start_read(client, next.handle); + + let previous = in_flight.take().unwrap(); + let output = wait_read(previous); black_box(&output); - in_flight = Some(next); + in_flight = Some(next_read); }); - if let Some(pending) = in_flight.take() { - let _ = pending.wait().unwrap(); + if let Some(read) = in_flight.take() { + let _ = wait_read(read); } + result } @@ -233,57 +233,73 @@ fn run_kernels(backend: &str, client: &ComputeClient) { run(MvRegulariseBench { client: client.clone(), }); - for &(ch, ch_name) in CHANNELS { + + for &(channels, channel_name) in CHANNELS { run(WarpBench { client: client.clone(), - ch, - ch_name, + channels, + channel_name, }); } } -/// For each channel mode, print the four temporal-pipeline rows -/// (with vs without MC, eager vs pipelined) adjacent so the cost -/// delta is visible at a glance. +/// Prints each channel mode's four temporal-pipeline rows side by side. +/// +/// The rows cover with and without motion compensation, eager and pipelined, so the cost delta is +/// visible at a glance. fn run_pipelines(backend: &str, client: &ComputeClient) { println!(); println!("--- {backend}: temporal pipeline (with vs without MC) ---"); - let variants: &[(MotionCompensationMode, &str)] = - &[(MotionCompensationMode::None, "_no_mc"), (mc_default(), "_mc")]; - let channels = [ + let motion_compensation = mc_default(); + let variants: &[(MotionCompensationMode, &str)] = &[ + (MotionCompensationMode::None, "_no_mc"), + (motion_compensation, "_mc"), + ]; + let channel_modes = [ ("luma", ChannelMode::Luma), ("chroma", ChannelMode::Chroma), ("yuv", ChannelMode::Yuv), ]; - for &(ch_name, mode) in &channels { - for &(mc, tag) in variants { - bench_eager::(client, backend, TEMPORAL_RADIUS, mode, ch_name, mc, tag).print(); - bench_pipelined::(client, backend, TEMPORAL_RADIUS, mode, ch_name, mc, tag).print(); + for &(channel_name, mode) in &channel_modes { + for &(variant, tag) in variants { + let eager = bench_eager::(client, backend, TEMPORAL_RADIUS, mode, channel_name, variant, tag); + eager.print(); + + let pipelined = + bench_pipelined::(client, backend, TEMPORAL_RADIUS, mode, channel_name, variant, tag); + pipelined.print(); } } + println!(); } -/// Direct vs chained MC estimation throughput at higher temporal -/// radii (r2, r4), luma only. The cost delta this comparison tracks -/// lives entirely in the MC dispatch path, not per-channel NLM -/// weighting, so extra channels would only add bench time without -/// adding signal. Prints eager and pipelined rows for each radius -/// and estimation strategy side by side, mirroring `run_pipelines`'s -/// "adjacent so the delta is visible at a glance" layout. +/// Direct and chained motion estimation throughput at radii 2 and 4, luma only. +/// +/// The cost delta lives entirely in the motion compensation path rather than per-channel weighting, +/// so extra channels would add bench time without adding signal. Eager and pipelined rows for each +/// radius and strategy print side by side. fn run_mc_estimation_comparison(backend: &str, client: &ComputeClient) { println!(); println!("--- {backend}: MC estimation comparison (direct vs chained, r2/r4) ---"); for &radius in &[2u32, 4u32] { - for (mc, label) in [(mc_default(), "direct"), (mc_chained_default(), "chained")] { + let direct = mc_default(); + let chained = mc_chained_default(); + for (variant, label) in [(direct, "direct"), (chained, "chained")] { let tag = format!("_mc_r{radius}_{label}"); - bench_eager::(client, backend, radius, ChannelMode::Luma, "luma", mc, &tag).print(); - bench_pipelined::(client, backend, radius, ChannelMode::Luma, "luma", mc, &tag).print(); + + let eager = bench_eager::(client, backend, radius, ChannelMode::Luma, "luma", variant, &tag); + eager.print(); + + let pipelined = + bench_pipelined::(client, backend, radius, ChannelMode::Luma, "luma", variant, &tag); + pipelined.print(); } } + println!(); } @@ -297,18 +313,16 @@ fn run_all(backend: &str, device: &R::Device) { #[derive(clap::Parser, Debug)] #[command(about = "Motion-compensation benches: per-kernel + end-to-end pipeline", long_about = None)] struct Cli { - /// GPU device to bind to. Format: `default`, `discrete[:N]`, - /// `integrated[:N]`, `virtual[:N]`, or `cpu`. + /// GPU device to bind to, one of `default`, `discrete[:N]`, `integrated[:N]`, `virtual[:N]` or `cpu`. #[arg(long, default_value = "default")] - device: av_denoise_core::Device, + device: Device, - /// Swallowed: cargo passes this when invoking the bench binary. + /// Swallowed, since cargo passes this when invoking the bench binary. #[arg(long, hide = true)] bench: bool, } fn main() { - use clap::Parser; let cli = Cli::parse(); println!("Motion-Compensation Benchmarks - 1920x1080"); diff --git a/av-denoise-core/benches/nl4d_ablation.rs b/av-denoise-core/benches/nl4d_ablation.rs index ab056ed..c8bcaa2 100644 --- a/av-denoise-core/benches/nl4d_ablation.rs +++ b/av-denoise-core/benches/nl4d_ablation.rs @@ -1,49 +1,48 @@ -//! Per-frame cost of the three collab kernels at the geometry the -//! pipeline actually runs them at. +//! Per-frame cost of the three collab kernels at the geometry the pipeline runs them at //! -//! A separate denoiser runs for luma and for chroma, and at 4:2:0 the -//! chroma planes are half-size on each axis. A bench that runs chroma at -//! full resolution would report four times the real work. Both planes are -//! measured here and summed, so the total is one frame's kernel cost. +//! Luma and chroma run in separate denoisers, and at 4:2:0 the chroma planes are half-size on each +//! axis. A bench that ran chroma at full resolution would report four times the real work, so both +//! planes are measured and summed into one frame's kernel cost. -use av_denoise_core::collab::geometry::{fused_cubes_x, ref_count, refs_along, strength_map_dims}; -use av_denoise_core::collab::kernels::aggregate::{ +use av_denoise_core::bench_api::collab::geometry::{fused_cubes_x, ref_count, refs_along, strength_map_dims}; +use av_denoise_core::bench_api::collab::kernels::aggregate::{ collab_normalise, collab_zero_accum, cross_frame_accum_scale, kaiser_window, weight_scale, }; -use av_denoise_core::collab::kernels::fused::{STRENGTH_MAP_OFF, collab_fused}; -use av_denoise_core::collab::kernels::transforms::dct_noise_profile; -use av_denoise_core::collab::{PATCH_SIZE, grid_frames, needs_warp_uniform_search}; -use av_denoise_core::nlmeans::{BLOCK_X, BLOCK_Y, NOISE_CURVE_BINS}; +use av_denoise_core::bench_api::collab::kernels::fused::{STRENGTH_MAP_OFF, collab_fused}; +use av_denoise_core::bench_api::collab::kernels::transforms::dct_noise_profile; +use av_denoise_core::bench_api::collab::{PATCH_SIZE, grid_frames, needs_warp_uniform_search}; +use av_denoise_core::bench_api::{BLOCK_X, BLOCK_Y, Device, NOISE_CURVE_BINS}; +use clap::Parser; use cubecl::benchmark::{Benchmark, BenchmarkComputations, TimingMethod}; use cubecl::prelude::*; use cubecl::server::Handle; #[derive(Clone, Copy)] -struct Geom { - w: u32, - h: u32, - ch: u32, - stored: u32, +struct PlaneGeometry { + width: u32, + height: u32, + channels: u32, + stored_channels: u32, label: &'static str, } -const PLANES: &[Geom] = &[ - Geom { - w: 1920, - h: 1080, - ch: 1, - stored: 1, +const PLANES: &[PlaneGeometry] = &[ + PlaneGeometry { + width: 1920, + height: 1080, + channels: 1, + stored_channels: 1, label: "luma 1920x1080 c1", }, - Geom { - w: 960, - h: 540, - ch: 2, - stored: 2, + PlaneGeometry { + width: 960, + height: 540, + channels: 2, + stored_channels: 2, label: "chroma 960x540 c2", }, ]; @@ -61,37 +60,38 @@ const SIGMA: f32 = 0.02; /// `Nl4dParams::default().lambda_ht`. const LAMBDA_HT: f32 = 4.158; -fn frame_data(g: Geom) -> Vec { - let mut data = Vec::with_capacity((g.w * g.h * g.stored) as usize); - for y in 0..g.h { - for x in 0..g.w { +fn frame_data(geometry: PlaneGeometry) -> Vec { + let mut data = Vec::with_capacity((geometry.width * geometry.height * geometry.stored_channels) as usize); + for y in 0..geometry.height { + for x in 0..geometry.width { let base = 0.5 + 0.2 * (x as f32 * 0.05).sin() * (y as f32 * 0.03).cos(); - for c in 0..g.stored { - let seed = (y * g.w + x) * g.stored + c; + for channel in 0..geometry.stored_channels { + let seed = (y * geometry.width + x) * geometry.stored_channels + channel; let hash = seed .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(340573321)); let noise = (hash as f32 / u32::MAX as f32 - 0.5) * 0.1; - data.push( - if c < g.ch { - (base + noise).clamp(0.0, 1.0) - } else { - 0.0 - }, - ); + let sample = if channel < geometry.channels { + (base + noise).clamp(0.0, 1.0) + } else { + 0.0 + }; + data.push(sample); } } } + data } fn block_sync(client: &ComputeClient) { - cubecl::future::block_on(client.sync()).unwrap(); + let sync = client.sync(); + cubecl::future::block_on(sync).unwrap(); } struct Rig { client: ComputeClient, - g: Geom, + geometry: PlaneGeometry, ring: Handle, ring_len: usize, mv_field: Handle, @@ -103,17 +103,17 @@ struct Rig { group_weight: Handle, sigma: Handle, dct_profile: Handle, - /// An all-zero noise curve, passed with `curve_valid = 0` where the - /// kernel never applies it. + /// An all-zero noise curve, passed with `curve_valid = 0` where the kernel never applies it. zero_curve: Handle, /// A strength map of ones, passed with the map off. unit_map: Handle, map_len: usize, /// The uniform aggregation window, which the `fused` row runs with. kaiser_off: Handle, - /// A `beta = 2` window, which the `fused_kaiser` row runs with. The - /// taper is two more loads and two more multiplies per scattered - /// pixel, and this is the row that says what they cost. + /// A `beta = 2` window, which the `fused_kaiser` row runs with. + /// + /// The kernel loads and applies the window taps for every scattered pixel in both rows, so this + /// row checks that the taper's values add no cost over the uniform window. kaiser_on: Handle, mv_len: usize, conf_len: usize, @@ -124,54 +124,87 @@ struct Rig { } impl Rig { - fn new(client: ComputeClient, g: Geom) -> Self { + fn new(client: ComputeClient, geometry: PlaneGeometry) -> Self { let mut ring_data = Vec::new(); for _ in 0..N_FRAMES { - ring_data.extend(frame_data(g)); + let frame = frame_data(geometry); + ring_data.extend(frame); } - let ring = client.create_from_slice(f32::as_bytes(&ring_data)); - - let blocks_x = g.w.div_ceil(BLK_STEP); - let blocks_y = g.h.div_ceil(BLK_STEP); - // `MotionCtx` pads each neighbour's slice of the motion and - // confidence buffers up to the runtime's binding alignment, and - // passes the padded element count as the kernel's stride. The - // rig pads the same way so the strides the kernels compile - // against here are the ones the pipeline compiles against. + + let ring_bytes = f32::as_bytes(&ring_data); + let ring = client.create_from_slice(ring_bytes); + + let blocks_x = geometry.width.div_ceil(BLK_STEP); + let blocks_y = geometry.height.div_ceil(BLK_STEP); + + // `MotionCtx` pads each neighbour's slice of the motion and confidence buffers up to the + // runtime's binding alignment and passes the padded count as the kernel's stride. The rig + // pads the same way so the kernels compile against the strides the pipeline uses. let align = client.properties().memory.alignment; let blocks = (blocks_x * blocks_y) as u64; let pad = |bytes: u64| bytes.next_multiple_of(align); - let mv_stride = (pad(blocks * 2 * size_of::() as u64) / size_of::() as u64) as u32; - let conf_stride = (pad(blocks * size_of::() as u64) / size_of::() as u64) as u32; + let mv_stride_bytes = pad(blocks * 2 * size_of::() as u64); + let mv_stride = (mv_stride_bytes / size_of::() as u64) as u32; + let conf_stride_bytes = pad(blocks * size_of::() as u64); + let conf_stride = (conf_stride_bytes / size_of::() as u64) as u32; let mv_len = (2 * RADIUS * mv_stride) as usize; let conf_len = (2 * RADIUS * conf_stride) as usize; - let refs = ref_count(g.w, g.h); - let pixels = (g.w * g.h) as usize; - let frame_len = pixels * g.stored as usize; + let refs = ref_count(geometry.width, geometry.height); + let pixels = (geometry.width * geometry.height) as usize; + let frame_len = pixels * geometry.stored_channels as usize; - let mut sigma_host = vec![0.0f32; g.stored as usize]; - sigma_host[..g.ch as usize].fill(SIGMA); + let mut sigma_host = vec![0.0f32; geometry.stored_channels as usize]; + sigma_host[..geometry.channels as usize].fill(SIGMA); - let (map_cols, map_rows) = strength_map_dims(g.w, g.h); + let (map_cols, map_rows) = strength_map_dims(geometry.width, geometry.height); let map_len = (map_cols * map_rows) as usize; - let unit_map = vec![1.0f32; map_len]; + let unit_map_host = vec![1.0f32; map_len]; + + let mv_host = vec![0i32; mv_len]; + let mv_bytes = i32::as_bytes(&mv_host); + let mv_field = client.create_from_slice(mv_bytes); + let conf_host = vec![1.0f32; conf_len]; + let conf_bytes = f32::as_bytes(&conf_host); + let confidence = client.create_from_slice(conf_bytes); + let slots_bytes = u32::as_bytes(&NEIGHBOUR_SLOTS); + let neighbour_slots = client.create_from_slice(slots_bytes); + let accum = client.empty(frame_len * N_FRAMES as usize * size_of::()); + let wsum = client.empty(pixels * N_FRAMES as usize * size_of::()); + let output = client.empty(frame_len * size_of::()); + let group_weight = client.empty(refs * size_of::()); + let sigma_bytes = f32::as_bytes(&sigma_host); + let sigma = client.create_from_slice(sigma_bytes); + let profile_host = dct_noise_profile(0.0); + let profile_bytes = f32::as_bytes(&profile_host); + let dct_profile = client.create_from_slice(profile_bytes); + let zero_curve_host = [0.0f32; NOISE_CURVE_BINS]; + let zero_curve_bytes = f32::as_bytes(&zero_curve_host); + let zero_curve = client.create_from_slice(zero_curve_bytes); + let unit_map_bytes = f32::as_bytes(&unit_map_host); + let unit_map = client.create_from_slice(unit_map_bytes); + let kaiser_off_host = kaiser_window(0.0); + let kaiser_off_bytes = f32::as_bytes(&kaiser_off_host); + let kaiser_off = client.create_from_slice(kaiser_off_bytes); + let kaiser_on_host = kaiser_window(2.0); + let kaiser_on_bytes = f32::as_bytes(&kaiser_on_host); + let kaiser_on = client.create_from_slice(kaiser_on_bytes); Self { - mv_field: client.create_from_slice(i32::as_bytes(&vec![0i32; mv_len])), - confidence: client.create_from_slice(f32::as_bytes(&vec![1.0f32; conf_len])), - neighbour_slots: client.create_from_slice(u32::as_bytes(&NEIGHBOUR_SLOTS)), - accum: client.empty(frame_len * N_FRAMES as usize * size_of::()), - wsum: client.empty(pixels * N_FRAMES as usize * size_of::()), - output: client.empty(frame_len * size_of::()), - group_weight: client.empty(refs * size_of::()), - sigma: client.create_from_slice(f32::as_bytes(&sigma_host)), - dct_profile: client.create_from_slice(f32::as_bytes(&dct_noise_profile(0.0))), - zero_curve: client.create_from_slice(f32::as_bytes(&[0.0f32; NOISE_CURVE_BINS])), - unit_map: client.create_from_slice(f32::as_bytes(&unit_map)), + mv_field, + confidence, + neighbour_slots, + accum, + wsum, + output, + group_weight, + sigma, + dct_profile, + zero_curve, + unit_map, map_len, - kaiser_off: client.create_from_slice(f32::as_bytes(&kaiser_window(0.0))), - kaiser_on: client.create_from_slice(f32::as_bytes(&kaiser_window(2.0))), + kaiser_off, + kaiser_on, ring_len: ring_data.len(), ring, mv_len, @@ -180,43 +213,51 @@ impl Rig { blocks_y, mv_stride, conf_stride, - g, + geometry, client, } } - /// The fused kernel, launched exactly as `Nl4dDenoiser` launches - /// it. Eight references share one 64-lane cube, so the grid is an - /// eighth as wide along x as the reference grid and the cube is 1D. - /// One row covers matching, filtering, and scatter together. + /// The fused kernel, launched exactly as `Nl4dDenoiser` launches it. + /// + /// Eight references share one 64-lane cube, so the grid is an eighth as wide along x as the + /// reference grid and the cube is 1D. One row covers matching, filtering and scatter together. fn fused(&self) { self.fused_with(&self.kaiser_off); } - /// [`Self::fused`] with the aggregation window on. + /// [Self::fused] with the aggregation window on. fn fused_kaiser(&self) { self.fused_with(&self.kaiser_on); } fn fused_with(&self, kaiser: &Handle) { - let g = self.g; - let refs = ref_count(g.w, g.h); - let refs_x = refs_along(g.w); - let pixels = (g.w * g.h) as usize; - let frame_len = pixels * g.stored as usize; - let (map_cols, map_rows) = strength_map_dims(g.w, g.h); + let geometry = self.geometry; + let refs = ref_count(geometry.width, geometry.height); + let refs_x = refs_along(geometry.width); + let pixels = (geometry.width * geometry.height) as usize; + let frame_len = pixels * geometry.stored_channels as usize; + let (map_cols, map_rows) = strength_map_dims(geometry.width, geometry.height); + + let cubes_x = fused_cubes_x(geometry.width); + let refs_y = refs_along(geometry.height); + let dct_profile = dct_noise_profile(0.0); + let group_weight_scale = weight_scale(SIGMA, &dct_profile); + let accum_scale = cross_frame_accum_scale(SPATIAL_RADIUS, RADIUS); + let uniform_search = needs_warp_uniform_search(&self.client); + let frames_per_volume = grid_frames(RADIUS); unsafe { collab_fused::launch_unchecked::( &self.client, - CubeCount::new_2d(fused_cubes_x(g.w), refs_along(g.h)), + CubeCount::new_2d(cubes_x, refs_y), CubeDim::new_1d(64), - g.stored as usize, + geometry.stored_channels as usize, ArrayArg::from_raw_parts(self.ring.clone(), self.ring_len), ArrayArg::from_raw_parts(self.mv_field.clone(), self.mv_len), ArrayArg::from_raw_parts(self.confidence.clone(), self.conf_len), ArrayArg::from_raw_parts(self.neighbour_slots.clone(), NEIGHBOUR_SLOTS.len()), - ArrayArg::from_raw_parts(self.sigma.clone(), g.stored as usize), + ArrayArg::from_raw_parts(self.sigma.clone(), geometry.stored_channels as usize), ArrayArg::from_raw_parts(self.zero_curve.clone(), NOISE_CURVE_BINS), ArrayArg::from_raw_parts(self.unit_map.clone(), self.map_len), ArrayArg::from_raw_parts(self.dct_profile.clone(), 8), @@ -229,11 +270,11 @@ impl Rig { LAMBDA_HT, 0u32, STRENGTH_MAP_OFF, - weight_scale(SIGMA, &dct_noise_profile(0.0)), - cross_frame_accum_scale(SPATIAL_RADIUS, RADIUS), - needs_warp_uniform_search(&self.client), + group_weight_scale, + accum_scale, + uniform_search, RADIUS, - grid_frames(RADIUS), + frames_per_volume, REFINE, self.mv_stride, self.conf_stride, @@ -241,11 +282,11 @@ impl Rig { BLKSIZE, self.blocks_x, self.blocks_y, - g.w, - g.h, - g.ch, + geometry.width, + geometry.height, + geometry.channels, K_MAX, - g.stored, + geometry.stored_channels, SPATIAL_RADIUS, refs_x, map_cols, @@ -257,33 +298,37 @@ impl Rig { } fn normalise(&self) { - let g = self.g; - let pixels = (g.w * g.h) as usize; - let frame_len = pixels * g.stored as usize; + let geometry = self.geometry; + let pixels = (geometry.width * geometry.height) as usize; + let frame_len = pixels * geometry.stored_channels as usize; + let cubes_x = geometry.width.div_ceil(BLOCK_X); + let cubes_y = geometry.height.div_ceil(BLOCK_Y); + unsafe { collab_normalise::launch_unchecked::( &self.client, - CubeCount::new_2d(g.w.div_ceil(BLOCK_X), g.h.div_ceil(BLOCK_Y)), + CubeCount::new_2d(cubes_x, cubes_y), CubeDim::new_2d(BLOCK_X, BLOCK_Y), - g.stored as usize, + geometry.stored_channels as usize, ArrayArg::from_raw_parts(self.accum.clone(), frame_len * N_FRAMES as usize), ArrayArg::from_raw_parts(self.wsum.clone(), pixels * N_FRAMES as usize), ArrayArg::from_raw_parts(self.output.clone(), frame_len), 0u32, - g.w, - g.h, - g.ch, - g.stored, + geometry.width, + geometry.height, + geometry.channels, + geometry.stored_channels, ); } } fn zero(&self) { - let g = self.g; - let pixels = (g.w * g.h) as usize; - let frame_len = pixels * g.stored as usize; + let geometry = self.geometry; + let pixels = (geometry.width * geometry.height) as usize; + let frame_len = pixels * geometry.stored_channels as usize; let dim = 256u32; let grid = (frame_len as u32).div_ceil(dim).min(65_535); + unsafe { collab_zero_accum::launch_unchecked::( &self.client, @@ -293,7 +338,7 @@ impl Rig { ArrayArg::from_raw_parts(self.wsum.clone(), pixels * N_FRAMES as usize), 0u32, pixels as u32, - g.stored, + geometry.stored_channels, grid * dim, ); } @@ -324,11 +369,12 @@ impl Benchmark for Arm<'_, R> { "normalise" => self.rig.normalise(), _ => self.rig.zero(), } + Ok(()) } fn name(&self) -> String { - format!("{:<15} {}", self.kernel, self.rig.g.label) + format!("{:<15} {}", self.kernel, self.rig.geometry.label) } fn sync(&self) { @@ -336,10 +382,11 @@ impl Benchmark for Arm<'_, R> { } fn shapes(&self) -> Vec> { + let geometry = self.rig.geometry; vec![vec![ - self.rig.g.w as usize, - self.rig.g.h as usize, - self.rig.g.ch as usize, + geometry.width as usize, + geometry.height as usize, + geometry.channels as usize, ]] } } @@ -347,29 +394,25 @@ impl Benchmark for Arm<'_, R> { #[derive(clap::Parser, Debug)] struct Cli { #[arg(long, default_value = "default")] - device: av_denoise_core::Device, + device: Device, #[arg(long, hide = true)] bench: bool, } fn main() { - use clap::Parser; let cli = Cli::parse(); #[cfg(feature = "vulkan")] { let device = cli.device.to_wgpu().expect("wgpu device conversion failed"); let client = cubecl::wgpu::WgpuRuntime::client(&device); + let alignment = client.properties().memory.alignment; println!("\ncollab kernels at real per-frame geometry, TimingMethod::Device"); println!(" device: {device:?}"); - println!( - " buffer alignment: {} bytes\n", - client.properties().memory.alignment - ); - - // (name, prime). A primed arm runs `fused` once before it is - // timed, so `normalise` reads real accumulator contents rather - // than an empty buffer. + println!(" buffer alignment: {} bytes\n", alignment); + + // (name, prime). A primed arm runs `fused` once before it is timed, so `normalise` reads + // real accumulator contents rather than an empty buffer. let kernels = [ ("zero_accum", false), ("fused", false), @@ -378,35 +421,37 @@ fn main() { ]; let mut totals = vec![0.0f64; kernels.len()]; - for g in PLANES { - let rig = Rig::::new(client.clone(), *g); - for (i, (k, prime)) in kernels.iter().enumerate() { + for plane in PLANES { + let rig = Rig::::new(client.clone(), *plane); + for (index, (kernel, prime)) in kernels.iter().enumerate() { let arm = Arm { rig: &rig, - kernel: k, + kernel, prime: *prime, }; let name = arm.name(); match arm.run(TimingMethod::Device) { - Ok(d) => { - let c = BenchmarkComputations::new(&d); - let ms = c.median.as_secs_f64() * 1000.0; - totals[i] += ms; - println!(" {name:<40} {ms:>8.3} ms"); + Ok(durations) => { + let computations = BenchmarkComputations::new(&durations); + let median_ms = computations.median.as_secs_f64() * 1000.0; + totals[index] += median_ms; + println!(" {name:<40} {median_ms:>8.3} ms"); }, - Err(e) => println!(" {name:<40} error: {e}"), + Err(err) => println!(" {name:<40} error: {err}"), } } + println!(); } println!(" --- per frame, both planes summed ---"); - let mut grand = 0.0; - for (i, (k, _)) in kernels.iter().enumerate() { - grand += totals[i]; - println!(" {:<40} {:>8.3} ms", *k, totals[i]); + let mut grand_total = 0.0; + for (index, (kernel, _)) in kernels.iter().enumerate() { + grand_total += totals[index]; + println!(" {:<40} {:>8.3} ms", *kernel, totals[index]); } - println!(" {:<40} {:>8.3} ms", "COLLAB TOTAL", grand); + + println!(" {:<40} {:>8.3} ms", "COLLAB TOTAL", grand_total); println!(); } diff --git a/av-denoise-core/benches/nlmeans.rs b/av-denoise-core/benches/nlmeans.rs index 5df0640..46aca28 100644 --- a/av-denoise-core/benches/nlmeans.rs +++ b/av-denoise-core/benches/nlmeans.rs @@ -1,21 +1,24 @@ use std::hint::black_box; use std::time::{Duration, Instant}; -use av_denoise_core::nlmeans::kernels::{nlm_accumulate, nlm_bilateral, nlm_dist_2d_weight, nlm_finish}; -use av_denoise_core::nlmeans::prefilter::bilateral_radius; -use av_denoise_core::nlmeans::{ +use av_denoise_core::bench_api::kernels::{nlm_accumulate, nlm_bilateral, nlm_dist_2d_weight, nlm_finish}; +use av_denoise_core::bench_api::prefilter::bilateral_radius; +use av_denoise_core::bench_api::{ BLOCK_X, BLOCK_Y, - ChannelMode, + Device, + HostIo, NlmDenoiser, NlmParams, - Pending, - PrefilterMode, + start_read, + wait_read, }; +use av_denoise_core::{ChannelMode, PrefilterMode}; +use clap::Parser; use cubecl::prelude::*; -const W: u32 = 1920; -const H: u32 = 1080; +const WIDTH: u32 = 1920; +const HEIGHT: u32 = 1080; const WARMUP_KERNEL: usize = 5; const ITERS_KERNEL: usize = 100; @@ -23,23 +26,38 @@ const ITERS_KERNEL: usize = 100; const WARMUP_PIPELINE: usize = 2; const ITERS_PIPELINE: usize = 500; -fn stored_channels(ch: u32) -> u32 { - match ch { +const BILATERAL_SIGMA_S: f32 = 3.0; +const BILATERAL_SIGMA_R: f32 = 0.02; + +const DENOISE_VARIANTS: &[(PrefilterMode, &str)] = &[ + (PrefilterMode::None, ""), + ( + PrefilterMode::Bilateral { + sigma_s: BILATERAL_SIGMA_S, + sigma_r: BILATERAL_SIGMA_R, + }, + "_rclip_bilateral", + ), + (PrefilterMode::NlmSpatial { strength_scale: 1.0 }, "_nlm_pilot"), +]; + +fn stored_channels(channels: u32) -> u32 { + match channels { 1 => 1, 2 => 2, - _ => 4, // YUV: 3 logical, 4 stored (vec3 -> vec4 padding) + _ => 4, // YUV has 3 logical channels stored as 4, padding vec3 to vec4. } } -fn make_synthetic_frame(w: u32, h: u32, ch: u32) -> Vec { - let mut data = Vec::with_capacity((w * h * ch) as usize); +fn make_synthetic_frame(width: u32, height: u32, channels: u32) -> Vec { + let mut data = Vec::with_capacity((width * height * channels) as usize); - for y in 0..h { - for x in 0..w { + for y in 0..height { + for x in 0..width { let base = 0.5 + 0.2 * (x as f32 * 0.05).sin() * (y as f32 * 0.03).cos(); - for c in 0..ch { - let seed = (y * w + x) * ch + c; + for channel in 0..channels { + let seed = (y * width + x) * channels + channel; let hash = seed .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(340573321)); @@ -52,19 +70,21 @@ fn make_synthetic_frame(w: u32, h: u32, ch: u32) -> Vec { data } -/// Pad to next-pow2 lane count (matches NlmDenoiser internal storage). -fn make_padded_frame(w: u32, h: u32, ch: u32) -> Vec { - let stored = stored_channels(ch); - if stored == ch { - return make_synthetic_frame(w, h, ch); +/// A synthetic frame padded to the power-of-two lane count `NlmDenoiser` stores. +fn make_padded_frame(width: u32, height: u32, channels: u32) -> Vec { + let stored_ch = stored_channels(channels); + if stored_ch == channels { + return make_synthetic_frame(width, height, channels); } - let src = make_synthetic_frame(w, h, ch); - let mut data = vec![0.0f32; (w * h * stored) as usize]; - for i in 0..(w * h) as usize { - for c in 0..ch as usize { - data[i * stored as usize + c] = src[i * ch as usize + c]; + + let synthetic = make_synthetic_frame(width, height, channels); + let mut data = vec![0.0f32; (width * height * stored_ch) as usize]; + for i in 0..(width * height) as usize { + for channel in 0..channels as usize { + data[i * stored_ch as usize + channel] = synthetic[i * channels as usize + channel]; } } + data } @@ -94,20 +114,22 @@ fn run_bench( client: &ComputeClient, warmup: usize, iterations: usize, - mut f: impl FnMut(), + mut step: impl FnMut(), ) -> BenchResult { for _ in 0..warmup { - f(); - futures::executor::block_on(client.sync()).unwrap(); + step(); + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); } let mut times = Vec::with_capacity(iterations); - for _ in 0..iterations { let start = Instant::now(); - f(); - futures::executor::block_on(client.sync()).unwrap(); - times.push(start.elapsed()); + step(); + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); + let elapsed = start.elapsed(); + times.push(elapsed); } let total: Duration = times.iter().sum(); @@ -128,39 +150,41 @@ fn run_bench( } } -fn div_ceil(a: u32, b: u32) -> u32 { - a.div_ceil(b) +fn div_ceil(value: u32, divisor: u32) -> u32 { + value.div_ceil(divisor) } fn bench_dist_2d_weight( client: &ComputeClient, backend: &str, - ch: u32, - ch_name: &str, + channels: u32, + channel_name: &str, ) -> BenchResult { - let pixels = (W * H) as usize; - let stored_ch = stored_channels(ch); - let frame = make_padded_frame(W, H, ch); - let input = client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(channels); + let frame = make_padded_frame(WIDTH, HEIGHT, channels); + let frame_bytes = f32::as_bytes(&frame); + let input = client.create_from_slice(frame_bytes); let output = client.empty(pixels * size_of::()); + let channel_mode = match channels { + 1 => ChannelMode::Luma, + 2 => ChannelMode::Chroma, + _ => ChannelMode::Yuv, + }; let params = NlmParams { patch_radius: 4, - channels: match ch { - 1 => ChannelMode::Luma, - 2 => ChannelMode::Chroma, - _ => ChannelMode::Yuv, - }, + channels: channel_mode, ..NlmParams::default() }; let h2_inv_norm = params.h2_inv_norm(); - let grid_x = div_ceil(W, BLOCK_X); - let grid_y = div_ceil(H, BLOCK_Y); + let grid_x = div_ceil(WIDTH, BLOCK_X); + let grid_y = div_ceil(HEIGHT, BLOCK_Y); let cube_count = CubeCount::new_2d(grid_x, grid_y); let cube_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); - let name = format!("dist_2d_weight_1080p_{ch_name}"); + let name = format!("dist_2d_weight_1080p_{channel_name}"); run_bench(&name, backend, client, WARMUP_KERNEL, ITERS_KERNEL, || unsafe { nlm_dist_2d_weight::launch_unchecked::( @@ -176,9 +200,9 @@ fn bench_dist_2d_weight( 0i32, h2_inv_norm, 0.0f32, - W, - H, - ch, + WIDTH, + HEIGHT, + channels, params.patch_radius, BLOCK_X, BLOCK_Y, @@ -189,27 +213,29 @@ fn bench_dist_2d_weight( fn bench_accumulate( client: &ComputeClient, backend: &str, - ch: u32, - ch_name: &str, + channels: u32, + channel_name: &str, ) -> BenchResult { - let pixels = (W * H) as usize; - let stored_ch = stored_channels(ch); - let frame = make_padded_frame(W, H, ch); - let input = client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(channels); + let frame = make_padded_frame(WIDTH, HEIGHT, channels); + let frame_bytes = f32::as_bytes(&frame); + let input = client.create_from_slice(frame_bytes); let weights_data = vec![0.5f32; pixels]; - let weights = client.create_from_slice(f32::as_bytes(&weights_data)); + let weights_bytes = f32::as_bytes(&weights_data); + let weights = client.create_from_slice(weights_bytes); let accum = client.empty(pixels * stored_ch as usize * size_of::()); let weight_sum = client.empty(pixels * size_of::()); let max_weight = client.empty(pixels * size_of::()); - let grid_x = div_ceil(W, BLOCK_X); - let grid_y = div_ceil(H, BLOCK_Y); + let grid_x = div_ceil(WIDTH, BLOCK_X); + let grid_y = div_ceil(HEIGHT, BLOCK_Y); let cube_count = CubeCount::new_2d(grid_x, grid_y); let cube_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); - let name = format!("accumulate_1080p_{ch_name}"); + let name = format!("accumulate_1080p_{channel_name}"); run_bench(&name, backend, client, WARMUP_KERNEL, ITERS_KERNEL, || unsafe { nlm_accumulate::launch_unchecked::( @@ -227,35 +253,44 @@ fn bench_accumulate( 0u32, 1i32, 0i32, - W, - H, + WIDTH, + HEIGHT, ); }) } -fn bench_finish(client: &ComputeClient, backend: &str, ch: u32, ch_name: &str) -> BenchResult { - let pixels = (W * H) as usize; - let stored_ch = stored_channels(ch); - let frame = make_padded_frame(W, H, ch); - let input = client.create_from_slice(f32::as_bytes(&frame)); +fn bench_finish( + client: &ComputeClient, + backend: &str, + channels: u32, + channel_name: &str, +) -> BenchResult { + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(channels); + let frame = make_padded_frame(WIDTH, HEIGHT, channels); + let frame_bytes = f32::as_bytes(&frame); + let input = client.create_from_slice(frame_bytes); let accum_data = vec![0.25f32; pixels * stored_ch as usize]; - let accum = client.create_from_slice(f32::as_bytes(&accum_data)); + let accum_bytes = f32::as_bytes(&accum_data); + let accum = client.create_from_slice(accum_bytes); - let ws_data = vec![1.0f32; pixels]; - let weight_sum = client.create_from_slice(f32::as_bytes(&ws_data)); + let weight_sum_data = vec![1.0f32; pixels]; + let weight_sum_bytes = f32::as_bytes(&weight_sum_data); + let weight_sum = client.create_from_slice(weight_sum_bytes); - let mw_data = vec![0.8f32; pixels]; - let max_weight = client.create_from_slice(f32::as_bytes(&mw_data)); + let max_weight_data = vec![0.8f32; pixels]; + let max_weight_bytes = f32::as_bytes(&max_weight_data); + let max_weight = client.create_from_slice(max_weight_bytes); let output = client.empty(pixels * stored_ch as usize * size_of::()); - let grid_x = div_ceil(W, BLOCK_X); - let grid_y = div_ceil(H, BLOCK_Y); + let grid_x = div_ceil(WIDTH, BLOCK_X); + let grid_y = div_ceil(HEIGHT, BLOCK_Y); let cube_count = CubeCount::new_2d(grid_x, grid_y); let cube_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); - let name = format!("finish_1080p_{ch_name}"); + let name = format!("finish_1080p_{channel_name}"); run_bench(&name, backend, client, WARMUP_KERNEL, ITERS_KERNEL, || unsafe { nlm_finish::launch_unchecked::( @@ -271,9 +306,9 @@ fn bench_finish(client: &ComputeClient, backend: &str, ch: u32, c 0u32, 0u32, 1.0f32, - W, - H, - ch, + WIDTH, + HEIGHT, + channels, ); }) } @@ -281,23 +316,24 @@ fn bench_finish(client: &ComputeClient, backend: &str, ch: u32, c fn bench_bilateral( client: &ComputeClient, backend: &str, - ch: u32, - ch_name: &str, + channels: u32, + channel_name: &str, ) -> BenchResult { - let pixels = (W * H) as usize; - let stored_ch = stored_channels(ch); - let frame = make_padded_frame(W, H, ch); - let input = client.create_from_slice(f32::as_bytes(&frame)); + let pixels = (WIDTH * HEIGHT) as usize; + let stored_ch = stored_channels(channels); + let frame = make_padded_frame(WIDTH, HEIGHT, channels); + let frame_bytes = f32::as_bytes(&frame); + let input = client.create_from_slice(frame_bytes); let output = client.empty(pixels * stored_ch as usize * size_of::()); let radius = bilateral_radius(BILATERAL_SIGMA_S); - let grid_x = div_ceil(W, BLOCK_X); - let grid_y = div_ceil(H, BLOCK_Y); + let grid_x = div_ceil(WIDTH, BLOCK_X); + let grid_y = div_ceil(HEIGHT, BLOCK_Y); let cube_count = CubeCount::new_2d(grid_x, grid_y); let cube_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); - let name = format!("bilateral_1080p_{ch_name}"); + let name = format!("bilateral_1080p_{channel_name}"); run_bench(&name, backend, client, WARMUP_KERNEL, ITERS_KERNEL, || unsafe { nlm_bilateral::launch_unchecked::( @@ -310,9 +346,9 @@ fn bench_bilateral( 0u32, 1.0 / (2.0 * BILATERAL_SIGMA_S * BILATERAL_SIGMA_S), 1.0 / (2.0 * BILATERAL_SIGMA_R * BILATERAL_SIGMA_R), - W, - H, - ch, + WIDTH, + HEIGHT, + channels, radius, BLOCK_X, BLOCK_Y, @@ -333,220 +369,201 @@ fn denoise_params(channels: ChannelMode, temporal_radius: u32, prefilter: Prefil } } -/// Push a frame (and, when needed, a matching reference) for the -/// configured prefilter mode. Used by the streaming pipeline benches so -/// the same push pattern works for `External` and non-`External` modes. -fn push_frame_for_prefilter( - denoiser: &mut NlmDenoiser, - frame: &[f32], - supply_reference: bool, -) { - if supply_reference { - denoiser.push_frame_with_reference(frame, frame); - } else { - denoiser.push_frame(frame); - } -} - -/// Steady-state streaming bench: every iteration pushes a fresh frame -/// (the real per-frame cost: upload plus optional prefilter) and then -/// calls the synchronous `denoise()` which waits for the readback. This -/// is the cost a caller pays if they push and wait in lockstep. +/// The steady-state spatial streaming cost. +/// +/// Every iteration pushes a fresh frame, the real per-frame upload and optional prefilter cost, then +/// calls the synchronous `denoise()` which waits for the readback. This is the cost a caller pays +/// when pushing and waiting in lockstep. fn bench_denoise_spatial( client: &ComputeClient, backend: &str, channels: ChannelMode, - ch_name: &str, + channel_name: &str, prefilter: PrefilterMode, tag: &str, ) -> BenchResult { - let ch = channels.count(); + let channel_count = channels.count(); let params = denoise_params(channels, 0, prefilter); - let frame = make_synthetic_frame(W, H, ch); - let supply_reference = matches!(prefilter, PrefilterMode::External); - let name = format!("denoise_spatial{tag}_1080p_{ch_name}"); + let frame = make_synthetic_frame(WIDTH, HEIGHT, channel_count); + let name = format!("denoise_spatial{tag}_1080p_{channel_name}"); - let mut denoiser = NlmDenoiser::::new(client, params, W, H); - futures::executor::block_on(client.sync()).unwrap(); + let mut denoiser = NlmDenoiser::::new(client, params, WIDTH, HEIGHT); + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); run_bench(&name, backend, client, WARMUP_PIPELINE, ITERS_PIPELINE, || { - push_frame_for_prefilter(&mut denoiser, &frame, supply_reference); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser"); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); black_box(&result); }) } -/// Steady-state temporal streaming bench. The window is pre-filled -/// outside the timer (a one-off cost in real usage), then every measured -/// iteration pushes one fresh frame and waits for that frame's denoise. +/// The steady-state temporal streaming cost. +/// +/// The window is pre-filled outside the timer because that is a one-off cost in real use. Every +/// measured iteration then pushes one fresh frame and waits for that frame's denoise. fn bench_denoise_temporal( client: &ComputeClient, backend: &str, channels: ChannelMode, - ch_name: &str, + channel_name: &str, prefilter: PrefilterMode, tag: &str, ) -> BenchResult { - let ch = channels.count(); + let channel_count = channels.count(); let params = denoise_params(channels, 1, prefilter); - let frame = make_synthetic_frame(W, H, ch); + let frame = make_synthetic_frame(WIDTH, HEIGHT, channel_count); let total_frames = 1 + 2 * params.temporal_radius as usize; - let supply_reference = matches!(prefilter, PrefilterMode::External); - let name = format!("denoise_temporal{tag}_1080p_{ch_name}"); + let name = format!("denoise_temporal{tag}_1080p_{channel_name}"); - let mut denoiser = NlmDenoiser::::new(client, params, W, H); + let mut denoiser = NlmDenoiser::::new(client, params, WIDTH, HEIGHT); for _ in 0..total_frames - 1 { - push_frame_for_prefilter(&mut denoiser, &frame, supply_reference); + denoiser.push_frame(&frame); } - futures::executor::block_on(client.sync()).unwrap(); + + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); run_bench(&name, backend, client, WARMUP_PIPELINE, ITERS_PIPELINE, || { - push_frame_for_prefilter(&mut denoiser, &frame, supply_reference); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser"); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); black_box(&result); }) } -/// Pipelined variant: each iteration pushes a fresh frame, submits its -/// denoise kernels (no wait), then blocks on the *previous* frame's -/// readback. With double-buffered output handles, frame N+1's kernels -/// run on the GPU while frame N's host readback is still in flight. +/// The pipelined temporal streaming cost. +/// +/// Each iteration pushes a fresh frame, submits its denoise kernels without waiting, then blocks on +/// the previous frame's readback. With double-buffered output handles, frame N+1's kernels run on +/// the GPU while frame N's host readback is still in flight. fn bench_denoise_temporal_pipelined( client: &ComputeClient, backend: &str, channels: ChannelMode, - ch_name: &str, + channel_name: &str, prefilter: PrefilterMode, tag: &str, ) -> BenchResult { - let ch = channels.count(); + let channel_count = channels.count(); let params = denoise_params(channels, 1, prefilter); - let frame = make_synthetic_frame(W, H, ch); + let frame = make_synthetic_frame(WIDTH, HEIGHT, channel_count); let total_frames = 1 + 2 * params.temporal_radius as usize; - let supply_reference = matches!(prefilter, PrefilterMode::External); - let name = format!("denoise_temporal_pipelined{tag}_1080p_{ch_name}"); + let name = format!("denoise_temporal_pipelined{tag}_1080p_{channel_name}"); - let mut denoiser = NlmDenoiser::::new(client, params, W, H); + let mut denoiser = NlmDenoiser::::new(client, params, WIDTH, HEIGHT); for _ in 0..total_frames - 1 { - push_frame_for_prefilter(&mut denoiser, &frame, supply_reference); + denoiser.push_frame(&frame); } - futures::executor::block_on(client.sync()).unwrap(); - // Prime the pipeline with one outstanding `Pending` so every measured - // iteration has previous work to wait on. - push_frame_for_prefilter(&mut denoiser, &frame, supply_reference); - let mut in_flight: Option> = Some(denoiser.denoise_submit().unwrap().unwrap()); + let sync = client.sync(); + futures::executor::block_on(sync).unwrap(); + + // Prime the pipeline with one outstanding readback so every measured iteration has previous + // work to wait on. + denoiser.push_frame(&frame); + let first = denoiser.denoise_submit_gpu().unwrap().unwrap(); + let first_read = start_read(client, first.handle); + let mut in_flight = Some(first_read); let result = run_bench(&name, backend, client, WARMUP_PIPELINE, ITERS_PIPELINE, || { - push_frame_for_prefilter(&mut denoiser, &frame, supply_reference); - let next = denoiser.denoise_submit().unwrap().unwrap(); - let output = in_flight.take().unwrap().wait().unwrap(); + denoiser.push_frame(&frame); + let next = denoiser.denoise_submit_gpu().unwrap().unwrap(); + let next_read = start_read(client, next.handle); + + let previous = in_flight.take().unwrap(); + let output = wait_read(previous); black_box(&output); - in_flight = Some(next); + in_flight = Some(next_read); }); - if let Some(pending) = in_flight.take() { - let _ = pending.wait().unwrap(); + if let Some(read) = in_flight.take() { + let _ = wait_read(read); } + result } -const BILATERAL_SIGMA_S: f32 = 3.0; -const BILATERAL_SIGMA_R: f32 = 0.02; - -const DENOISE_VARIANTS: &[(PrefilterMode, &str)] = &[ - (PrefilterMode::None, ""), - (PrefilterMode::External, "_rclip_external"), - ( - PrefilterMode::Bilateral { - sigma_s: BILATERAL_SIGMA_S, - sigma_r: BILATERAL_SIGMA_R, - }, - "_rclip_bilateral", - ), - (PrefilterMode::NlmSpatial { strength_scale: 1.0 }, "_nlm_pilot"), -]; - fn run_all_benches(backend: &str, device: &R::Device) { let client = R::client(device); println!("--- {backend} ---"); println!(); - let channels = [ + let channel_modes = [ (1u32, "luma", ChannelMode::Luma), (2, "chroma", ChannelMode::Chroma), (3, "yuv", ChannelMode::Yuv), ]; - for &(ch, ch_name, _) in &channels { - bench_dist_2d_weight::(&client, backend, ch, ch_name).print(); + for &(channels, channel_name, _) in &channel_modes { + let result = bench_dist_2d_weight::(&client, backend, channels, channel_name); + result.print(); } + println!(); - for &(ch, ch_name, _) in &channels { - bench_accumulate::(&client, backend, ch, ch_name).print(); + for &(channels, channel_name, _) in &channel_modes { + let result = bench_accumulate::(&client, backend, channels, channel_name); + result.print(); } + println!(); - for &(ch, ch_name, _) in &channels { - bench_finish::(&client, backend, ch, ch_name).print(); + for &(channels, channel_name, _) in &channel_modes { + let result = bench_finish::(&client, backend, channels, channel_name); + result.print(); } + println!(); - for &(ch, ch_name, _) in &channels { - bench_bilateral::(&client, backend, ch, ch_name).print(); + for &(channels, channel_name, _) in &channel_modes { + let result = bench_bilateral::(&client, backend, channels, channel_name); + result.print(); } + println!(); - // Group each channel mode's baseline and rclip variants together so - // before/after comparisons land on adjacent rows. - for &(_, ch_name, mode) in &channels { + // Group each channel mode's baseline and rclip variants together so before/after comparisons + // land on adjacent rows. + for &(_, channel_name, mode) in &channel_modes { for &(prefilter, tag) in DENOISE_VARIANTS { - bench_denoise_spatial::(&client, backend, mode, ch_name, prefilter, tag).print(); + let result = bench_denoise_spatial::(&client, backend, mode, channel_name, prefilter, tag); + result.print(); } } + println!(); - for &(_, ch_name, mode) in &channels { + for &(_, channel_name, mode) in &channel_modes { for &(prefilter, tag) in DENOISE_VARIANTS { - if matches!(prefilter, PrefilterMode::External) { - continue; - } - bench_denoise_temporal::(&client, backend, mode, ch_name, prefilter, tag).print(); - bench_denoise_temporal_pipelined::(&client, backend, mode, ch_name, prefilter, tag).print(); + let eager = bench_denoise_temporal::(&client, backend, mode, channel_name, prefilter, tag); + eager.print(); + + let pipelined = + bench_denoise_temporal_pipelined::(&client, backend, mode, channel_name, prefilter, tag); + pipelined.print(); } } + println!(); } -/// Bench-harness CLI. `cargo bench --bench nlmeans -- --device discrete:1` -/// selects the second discrete GPU. +/// Bench-harness CLI. +/// +/// `cargo bench --bench nlmeans -- --device discrete:1` selects the second discrete GPU. #[derive(clap::Parser, Debug)] #[command(about = "NLMeans benchmarks", long_about = None)] struct Cli { - /// GPU device to bind to. Format: `default`, `discrete[:N]`, - /// `integrated[:N]`, `virtual[:N]`, or `cpu`. + /// GPU device to bind to, one of `default`, `discrete[:N]`, `integrated[:N]`, `virtual[:N]` or `cpu`. #[arg(long, default_value = "default")] - device: av_denoise_core::Device, + device: Device, - /// Swallowed: cargo passes this when invoking the bench binary. + /// Swallowed, since cargo passes this when invoking the bench binary. #[arg(long, hide = true)] bench: bool, } fn main() { - use clap::Parser; let cli = Cli::parse(); println!("NLMeans Benchmarks - 1920x1080"); diff --git a/av-denoise-core/src/accelerate.rs b/av-denoise-core/src/accelerate.rs deleted file mode 100644 index b4e7928..0000000 --- a/av-denoise-core/src/accelerate.rs +++ /dev/null @@ -1,81 +0,0 @@ -//! The hardware backends kernels can run on. -//! -//! An [`Accelerator`] names a backend rather than a specific piece of -//! hardware. Which physical GPU it lands on is chosen separately, with -//! [`crate::Device`]. -//! -//! Only the backends whose crate feature is enabled exist at compile -//! time, so a build without the `cuda` feature has no `Accelerator::Cuda` -//! variant at all. -//! -//! [`Denoiser::create`](crate::Denoiser::create) takes a list of these -//! and uses the first one that starts successfully, which lets a program -//! prefer a fast backend and quietly fall back to a slower one. -//! -//! ```no_run -//! use av_denoise_core::accelerate::get_default_accelerators; -//! -//! // Every backend this build supports, in the order to try them. -//! let preferred = get_default_accelerators(); -//! # let _ = preferred; -//! ``` -//! -//! A list can also be written out by hand, such as -//! `vec![Accelerator::Cuda, Accelerator::Vulkan]` to prefer the vendor -//! backend and fall back to the portable one. -//! -//! Every accelerator here runs kernels on a GPU. There is no software -//! backend, because the collaborative filter aggregates through atomic -//! floating-point adds and cubecl's CPU runtime does not implement -//! atomics. [`crate::Device::Cpu`] still selects a software *device* -//! where the platform offers one, such as lavapipe under Vulkan. - -use strum_macros::{Display, EnumIter, EnumString, IntoStaticStr}; - -#[derive(Debug, Copy, Clone, Eq, PartialEq, IntoStaticStr, EnumString, EnumIter, Display)] -#[strum(serialize_all = "snake_case")] -/// A hardware backend that kernels can run on. -pub enum Accelerator { - #[cfg(any(feature = "cuda", docsrs))] - #[cfg_attr(docsrs, doc(cfg(feature = "cuda")))] - /// Runs kernels through the Nvidia CUDA backend. - /// - /// Nvidia GPUs only. - Cuda, - #[cfg(any(feature = "vulkan", docsrs))] - #[cfg_attr(docsrs, doc(cfg(feature = "vulkan")))] - /// Runs kernels through the wgpu Vulkan backend. - /// - /// This is the lightest and most portable option, because it works - /// on any platform and GPU that supports basic compute shaders. - Vulkan, - #[cfg(any(feature = "metal", docsrs))] - #[cfg_attr(docsrs, doc(cfg(feature = "metal")))] - /// Runs kernels through the wgpu Metal backend. - /// - /// This is the only option on Apple Silicon. - Metal, - #[cfg(any(feature = "rocm", docsrs))] - #[cfg_attr(docsrs, doc(cfg(feature = "rocm")))] - /// Runs kernels through the AMD ROCm backend. - /// - /// WARNING: ROCm is *not* the recommended backend for AMD GPUs, it is slower and often - /// plagued with issues from drivers, vulkan will almost certainly be faster and - /// less buggy. - /// - /// AMD GPUs only. - Rocm, -} - -/// Returns every accelerator this build enables, in the order to try -/// them. -pub fn get_default_accelerators() -> Vec { - use strum::IntoEnumIterator; - - let mut accelerator = Vec::new(); - for enabled in Accelerator::iter() { - accelerator.push(enabled); - } - - accelerator -} diff --git a/av-denoise-core/src/bench_api.rs b/av-denoise-core/src/bench_api.rs new file mode 100644 index 0000000..a7fd95e --- /dev/null +++ b/av-denoise-core/src/bench_api.rs @@ -0,0 +1,311 @@ +//! Internal items and host helpers shared by the benches and tests. Not a stable interface. + +pub mod collab { + pub use crate::collab::*; +} + +pub mod engine_kernels { + pub use crate::engine::kernels::*; +} + +pub mod harness { + pub use crate::nl4d::harness::*; +} + +pub mod kernels { + pub use crate::nlmeans::kernels::*; +} + +pub mod motion { + pub use crate::nlmeans::motion::*; +} + +pub mod nl4d_kernels { + pub use crate::nl4d::kernels::*; +} + +pub mod prefilter { + pub use crate::nlmeans::prefilter::*; +} + +use std::future::Future; +use std::pin::Pin; +use std::str::FromStr; + +use anyhow::Context; +use cubecl::bytes::Bytes; +use cubecl::prelude::*; +use cubecl::server::{Handle, ServerError}; + +use crate::engine::{DevicePlane, EgressSource, SampleFormat, egress}; +pub use crate::nl4d::denoiser::Nl4dDenoiser; +pub use crate::nl4d::params::Nl4dParams; +pub use crate::nl4d::snapshot::MotionSnapshot; +use crate::nlmeans::ChannelMode; +pub use crate::nlmeans::denoiser::{GpuOutput, NlmDenoiser}; +pub use crate::nlmeans::params::NlmParams; + +pub const BLOCK_X: u32 = crate::nlmeans::BLOCK_X; +pub const BLOCK_Y: u32 = crate::nlmeans::BLOCK_Y; +pub const NOISE_CURVE_BINS: usize = crate::nlmeans::NOISE_CURVE_BINS; + +/// A readback in flight, resolving to one buffer per handle read. +pub type ReadFuture = Pin, ServerError>> + Send>>; + +/// Starts reading `handle` back. Nothing is read until the future is first polled. +pub fn start_read(client: &ComputeClient, handle: Handle) -> ReadFuture { + let client = client.clone(); + let future = async move { client.read_async(vec![handle]).await }; + + Box::pin(future) +} + +/// Blocks on a readback and copies the frame out as f32. +pub fn wait_read(read: ReadFuture) -> Vec { + let buffers = cubecl::future::block_on(read).expect("readback failed"); + let samples = f32::from_bytes(&buffers[0]); + + samples.to_vec() +} + +/// Pushes interleaved host frames and reads denoised frames back as interleaved f32. +pub trait HostIo { + /// Uploads one frame of `width * height * channels` values and pushes it. + fn push_frame(&mut self, frame: &[f32]); + + /// Denoises the next frame, or returns `None` while the window is still filling. + fn denoise(&mut self) -> Result>, anyhow::Error>; + + /// Hands every frame the stream still holds to `sink`, then starts a fresh stream. + fn flush(&mut self, sink: impl FnMut(&[f32])) -> Result<(), anyhow::Error>; +} + +impl HostIo for NlmDenoiser { + fn push_frame(&mut self, frame: &[f32]) { + let (width, height, channels) = self.frame_shape(); + let handles = upload_frame(self.compute_client(), frame, width, height, channels); + let planes = device_planes(&handles, width, height); + + let pushed = self.push_planes(&planes, SampleFormat::F32); + pushed.expect("frame push failed"); + } + + fn denoise(&mut self) -> Result>, anyhow::Error> { + let Some(output) = self.denoise_submit_gpu()? else { + return Ok(None); + }; + + let frame = self.read_back(&output.handle)?; + Ok(Some(frame)) + } + + fn flush(&mut self, mut sink: impl FnMut(&[f32])) -> Result<(), anyhow::Error> { + let target = self.flush_target(); + + for _ in 0..target { + let output = loop { + let step = self.flush_step_gpu()?; + if let Some(output) = step { + break output; + } + }; + + let frame = self.read_back(&output.handle)?; + sink(&frame); + } + + self.reset_stream_state(); + + Ok(()) + } +} + +impl NlmDenoiser { + fn read_back(&self, frame: &Handle) -> Result, anyhow::Error> { + let shape = self.frame_shape(); + read_frame(self.compute_client(), self.placeholder(), frame, shape) + } +} + +impl HostIo for Nl4dDenoiser { + fn push_frame(&mut self, frame: &[f32]) { + let (width, height, channels) = self.frame_shape(); + let handles = upload_frame(self.compute_client(), frame, width, height, channels); + let planes = device_planes(&handles, width, height); + + let pushed = self.push_planes(&planes, SampleFormat::F32); + pushed.expect("frame push failed"); + } + + fn denoise(&mut self) -> Result>, anyhow::Error> { + let Some(region) = self.submit_passes()? else { + return Ok(None); + }; + + let handle = self.read_region(region); + let frame = self.read_back(&handle)?; + Ok(Some(frame)) + } + + fn flush(&mut self, mut sink: impl FnMut(&[f32])) -> Result<(), anyhow::Error> { + let regions = self.finish_passes()?; + + for region in regions { + let handle = self.read_region(region); + let frame = self.read_back(&handle)?; + sink(&frame); + } + + self.reset_stream(); + + Ok(()) + } +} + +impl Nl4dDenoiser { + fn read_back(&self, frame: &Handle) -> Result, anyhow::Error> { + let shape = self.frame_shape(); + read_frame(self.compute_client(), self.placeholder(), frame, shape) + } +} + +/// Splits an interleaved frame into one uploaded f32 plane per channel. +fn upload_frame( + client: &ComputeClient, + frame: &[f32], + width: u32, + height: u32, + channels: ChannelMode, +) -> Vec { + let channel_count = channels.count() as usize; + let expected = width as usize * height as usize * channel_count; + + assert_eq!( + frame.len(), + expected, + "frame size mismatch: expected {expected}, got {}", + frame.len() + ); + + let mut handles = Vec::with_capacity(channel_count); + + for channel in 0..channel_count { + let samples: Vec = frame + .iter() + .skip(channel) + .step_by(channel_count) + .copied() + .collect(); + let bytes = f32::as_bytes(&samples); + let handle = client.create_from_slice(bytes); + handles.push(handle); + } + + handles +} + +fn device_planes(handles: &[Handle], width: u32, height: u32) -> Vec> { + handles + .iter() + .map(|handle| DevicePlane::new(handle, width, height)) + .collect() +} + +/// Egresses an internal interleaved frame into f32 planes and reads them back interleaved. +fn read_frame( + client: &ComputeClient, + placeholder: &Handle, + frame: &Handle, + shape: (u32, u32, ChannelMode), +) -> Result, anyhow::Error> { + let (width, height, channels) = shape; + let pixels = width * height; + let channel_count = channels.count() as usize; + let plane_bytes = pixels as usize * size_of::(); + + let handles: Vec = (0..channel_count).map(|_| client.empty(plane_bytes)).collect(); + let planes = device_planes(&handles, width, height); + let source = EgressSource { + frame, + pixels, + channels: channels.count(), + stored_ch: channels.storage_count(), + }; + + egress(client, source, &planes, SampleFormat::F32, placeholder); + + let mut plane_samples = Vec::with_capacity(channel_count); + + for handle in handles { + let bytes = client.read_one(handle).context("plane readback failed")?; + let samples = f32::from_bytes(&bytes).to_vec(); + plane_samples.push(samples); + } + + let mut interleaved = Vec::with_capacity(pixels as usize * channel_count); + + for pixel in 0..pixels as usize { + for samples in &plane_samples { + interleaved.push(samples[pixel]); + } + } + + Ok(interleaved) +} + +/// A `--device` selector, such as `default` or `discrete:1`. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Device { + Default, + Discrete(usize), + Integrated(usize), + Virtual(usize), + Cpu, +} + +impl FromStr for Device { + type Err = String; + + fn from_str(selector: &str) -> Result { + let (kind, suffix) = match selector.split_once(':') { + Some((kind, index)) => (kind, Some(index)), + None => (selector, None), + }; + + if matches!(kind, "default" | "cpu") && suffix.is_some() { + return Err(format!("device kind '{kind}' takes no index, got '{selector}'")); + } + + let index_text = suffix.unwrap_or("0"); + let index = index_text.parse::(); + let index = index.map_err(|_| format!("invalid device index '{index_text}' in '{selector}'")); + + match kind { + "default" => Ok(Device::Default), + "cpu" => Ok(Device::Cpu), + "discrete" => Ok(Device::Discrete(index?)), + "integrated" => Ok(Device::Integrated(index?)), + "virtual" => Ok(Device::Virtual(index?)), + other => Err(format!( + "unknown device kind '{other}', expected default, discrete[:N], integrated[:N], virtual[:N], or cpu" + )), + } + } +} + +#[cfg(any(feature = "vulkan", feature = "metal"))] +impl Device { + pub fn to_wgpu(self) -> Result { + use cubecl::wgpu::WgpuDevice; + + let device = match self { + Device::Default => WgpuDevice::DefaultDevice, + Device::Discrete(index) => WgpuDevice::DiscreteGpu(index), + Device::Integrated(index) => WgpuDevice::IntegratedGpu(index), + Device::Virtual(index) => WgpuDevice::VirtualGpu(index), + Device::Cpu => WgpuDevice::Cpu, + }; + + Ok(device) + } +} diff --git a/av-denoise-core/src/collab/geometry.rs b/av-denoise-core/src/collab/geometry.rs index 8c30a8e..eba483f 100644 --- a/av-denoise-core/src/collab/geometry.rs +++ b/av-denoise-core/src/collab/geometry.rs @@ -1,28 +1,27 @@ +use super::{PATCH_SIZE, STEP}; + /// Number of reference patches along one axis. /// -/// `dim` must be at least `PATCH_SIZE`, which denoiser construction -/// validates. +/// `dim` must be at least `PATCH_SIZE`. pub fn refs_along(dim: u32) -> u32 { - (dim - super::PATCH_SIZE).div_ceil(super::STEP) + 1 + (dim - PATCH_SIZE).div_ceil(STEP) + 1 } -/// Cubes along x for [`crate::collab::kernels::fused::collab_fused`]. +/// Cubes along x for the `collab_fused` kernel. /// -/// That kernel gives each of its eight 8-lane groups one reference -/// patch, so a row of references needs an eighth as many cubes as -/// [`refs_along`] returns. The count rounds up, and the last cube of a -/// row runs dead groups for the references past the end. +/// Each cube runs eight 8-lane groups with one reference patch each. The count rounds up, so the +/// last cube of a row runs dead groups past the end. pub fn fused_cubes_x(width: u32) -> u32 { refs_along(width).div_ceil(8) } -/// Top-left pixel of reference index `i` along one axis. The last -/// reference clamps so its patch stays inside the frame. -pub fn ref_pos(i: u32, dim: u32) -> u32 { - (i * super::STEP).min(dim - super::PATCH_SIZE) +/// Top-left pixel of reference `index` along one axis. +/// +/// The last reference clamps so its patch stays inside the frame. +pub fn ref_pos(index: u32, dim: u32) -> u32 { + (index * STEP).min(dim - PATCH_SIZE) } -/// Total reference count for a frame. pub fn ref_count(width: u32, height: u32) -> usize { refs_along(width) as usize * refs_along(height) as usize } @@ -34,6 +33,7 @@ pub fn ref_count(width: u32, height: u32) -> usize { pub fn strength_map_dims(width: u32, height: u32) -> (u32, u32) { let cols = 2 * width.div_ceil(16); let rows = 2 * height.div_ceil(16); + (cols, rows) } @@ -44,15 +44,18 @@ mod tests { #[test] fn refs_cover_1080p_exactly() { // (1920 - 8) / 4 + 1 = 479, (1080 - 8) / 4 + 1 = 269. - assert_eq!(refs_along(1920), 479); - assert_eq!(refs_along(1080), 269); + let refs_x = refs_along(1920); + let refs_y = refs_along(1080); + assert_eq!(refs_x, 479); + assert_eq!(refs_y, 269); } #[test] fn fused_cubes_cover_every_reference() { - // 1920 gives 479 references, so the last of the 60 cubes runs - // seven live groups and one dead one. - assert_eq!(fused_cubes_x(1920), 60); + // 1920 gives 479 references, so the last of the 60 cubes runs seven live groups. + let full_hd_cubes = fused_cubes_x(1920); + assert_eq!(full_hd_cubes, 60); + for dim in [8u32, 9, 21, 64, 100, 104, 128, 1280, 1920, 3840] { let cubes = fused_cubes_x(dim); let refs = refs_along(dim); @@ -63,27 +66,29 @@ mod tests { #[test] fn last_ref_clamps_inside_the_frame() { - let w = 21; // not a multiple of STEP past PATCH_SIZE - let n = refs_along(w); - assert_eq!(ref_pos(n - 1, w), w - 8); - for i in 0..n { - assert!(ref_pos(i, w) + 8 <= w); + // Not a multiple of STEP past PATCH_SIZE. + let width = 21; + let ref_total = refs_along(width); + let last_pos = ref_pos(ref_total - 1, width); + assert_eq!(last_pos, width - 8); + + for i in 0..ref_total { + let pos = ref_pos(i, width); + assert!(pos + 8 <= width); } } #[test] fn every_pixel_is_covered_by_one_to_three_refs_per_axis() { - // Regular spacing gives 2 covering references per axis, and the - // single clamped edge gap can add a third. It never reaches a - // fourth, so the bound here is 1..=3, giving at most 9 (3 x 3) - // covering references per pixel in 2D. + // Regular spacing gives 2 covering references per axis, and the clamped edge gap can add a + // third. for dim in [8u32, 9, 16, 21, 64] { for x in 0..dim { - let n = refs_along(dim); - let covering = (0..n) + let ref_total = refs_along(dim); + let covering = (0..ref_total) .filter(|&i| { - let p = ref_pos(i, dim); - p <= x && x < p + 8 + let pos = ref_pos(i, dim); + pos <= x && x < pos + 8 }) .count(); assert!((1..=3).contains(&covering), "dim={dim} x={x} covering={covering}"); @@ -93,17 +98,22 @@ mod tests { #[test] fn ref_count_is_the_product_of_the_per_axis_counts() { - assert_eq!( - ref_count(1920, 1080), - refs_along(1920) as usize * refs_along(1080) as usize - ); + let count = ref_count(1920, 1080); + let refs_x = refs_along(1920); + let refs_y = refs_along(1080); + assert_eq!(count, refs_x as usize * refs_y as usize); } #[test] fn strength_map_dims_round_up_to_whole_blocks() { - assert_eq!(strength_map_dims(1920, 1080), (240, 136)); - assert_eq!(strength_map_dims(960, 540), (120, 68)); - assert_eq!(strength_map_dims(70, 54), (10, 8)); - assert_eq!(strength_map_dims(16, 16), (2, 2)); + let full_hd = strength_map_dims(1920, 1080); + let quarter_hd = strength_map_dims(960, 540); + let ragged = strength_map_dims(70, 54); + let two_blocks = strength_map_dims(16, 16); + + assert_eq!(full_hd, (240, 136)); + assert_eq!(quarter_hd, (120, 68)); + assert_eq!(ragged, (10, 8)); + assert_eq!(two_blocks, (2, 2)); } } diff --git a/av-denoise-core/src/collab/kernels/aggregate.rs b/av-denoise-core/src/collab/kernels/aggregate.rs index 49b7910..71ff969 100644 --- a/av-denoise-core/src/collab/kernels/aggregate.rs +++ b/av-denoise-core/src/collab/kernels/aggregate.rs @@ -1,68 +1,38 @@ use cubecl::prelude::*; -use crate::collab::kernels::transforms::RECIPROCAL_FLOOR; +use super::transforms::RECIPROCAL_FLOOR; use crate::collab::{MAX_K, PATCH_AREA, PATCH_SIZE, STEP}; /// The fixed-point scale a single-frame accumulator counts in. /// -/// Aggregation adds one weighted value per covering patch, one atomic add -/// each. Integer atomics are much faster than float atomics on most GPUs, -/// so the accumulators hold fixed-point integers. +/// The accumulators hold fixed-point integers because integer atomics are much faster than float +/// atomics on most GPUs. `2^19` is the largest power of two that keeps the worst case well inside +/// `i32`. One pass covers a pixel with at most 392 member patches (49 step-grid references times +/// `MAX_K`), each weighted by at most 1 and clamped to [ACCUM_CLAMP], so the accumulator peaks +/// near `1.03e9`, about half of `i32::MAX`. One unit is `1.9e-6`, well under the `2.4e-4` one +/// 12-bit code level spans. `wsum` peaks at the same place because [WEIGHT_GAIN] trades the +/// weight's smaller bound for exactly that much scale. /// -/// `2^19` is the largest power of two that keeps the worst case well -/// inside `i32`. One pass covers a pixel with at most 392 member patches -/// (49 references on the step grid, each with at most `MAX_K` members), -/// each weighted by at most 1 (see [`weight_scale`]) and clamped to -/// [`ACCUM_CLAMP`], so the accumulator peaks near `1.03e9`, about half of -/// `i32::MAX`. One unit is `1.9e-6`, well under the `2.4e-4` that one -/// 12-bit code level spans. `wsum` peaks at the same place, because -/// [`WEIGHT_GAIN`] trades a weight's smaller bound for exactly that much -/// more scale. -/// -/// A cross-frame ring collects several passes before it is read back and -/// needs [`cross_frame_accum_scale`] instead. [`scatter_patch`] takes the -/// scale as an argument so a single-frame caller keeps this one's full -/// precision. Either way the scale cancels in [`collab_normalise`], which -/// divides one accumulator by the other. +/// A cross-frame ring uses [cross_frame_accum_scale] instead. Either scale cancels in +/// `collab_normalise`, which divides one accumulator by the other. pub const ACCUM_SCALE: f32 = 524_288.0; -/// The headroom [`cross_frame_accum_scale`] leaves under `i32::MAX`, -/// matching the roughly two-fold margin [`ACCUM_SCALE`] leaves for the -/// single-pass case. +/// The headroom `cross_frame_accum_scale` leaves under `i32::MAX`, matching the two-fold margin +/// of [ACCUM_SCALE]. const CROSS_FRAME_SAFETY_FACTOR: f64 = 2.0; -/// The fixed-point scale [`crate::nl4d::Nl4dDenoiser`]'s cross-frame -/// accumulator ring counts in, in place of [`ACCUM_SCALE`]. -/// -/// One fixed constant will not do. A large `spatial_radius` or -/// `temporal_radius` pushes the worst-case accumulator value past -/// `i32::MAX`, and a scale small enough to survive that would throw away -/// precision at the radii people actually use. Deriving it per -/// configuration avoids both. -/// -/// The worst case is how many member patches can cover one pixel. A -/// member covering pixel `x` came from a reference whose own top-left is -/// within `(PATCH_SIZE - 1) + 2 * spatial_radius` pixels of `x`, and -/// references sit on the `STEP` grid, so along one axis at most -/// -/// ```text -/// refs_per_axis = ((PATCH_SIZE - 1) + 2 * spatial_radius) / STEP + 1 -/// ``` +/// The fixed-point scale for the cross-frame accumulator ring, in place of [ACCUM_SCALE]. /// -/// of them can reach it. Squaring for both axes and taking every one of -/// `MAX_K` members gives one pass's contribution count. A cross-frame -/// region collects that from `2 * temporal_radius + 1` passes in steady -/// state. A region can collect up to `4 * temporal_radius + 1`, the steady -/// `2 * temporal_radius + 1` plus `temporal_radius` head and -/// `temporal_radius` tail passes when a scene is short enough for one -/// frame to sit in both edge rings. Each -/// contribution is weighted by at most 1 (see [`weight_scale`]) and -/// clamped to [`ACCUM_CLAMP`]. The same figure sizes `wsum`, whose -/// contributions are bounded by [`WEIGHT_CLAMP`] instead but counted at -/// [`WEIGHT_GAIN`] times the scale. +/// Large radii push the worst case past `i32::MAX`, and one constant small enough for them would +/// waste precision at common radii, so the scale is derived per configuration. Along one axis at +/// most `((PATCH_SIZE - 1) + 2 * spatial_radius) / STEP + 1` step-grid references can reach a +/// pixel, and each brings `MAX_K` members. A region collects up to `4 * temporal_radius + 1` +/// passes, the steady `2 * temporal_radius + 1` plus `temporal_radius` head and tail passes when a +/// short scene puts one frame in both edge rings. /// -/// The result is the largest power of two that keeps that worst case -/// under `i32::MAX` divided by [`CROSS_FRAME_SAFETY_FACTOR`]. +/// Each contribution is weighted by at most 1 and clamped to [ACCUM_CLAMP], and `wsum` fits the +/// same budget through [WEIGHT_CLAMP] and [WEIGHT_GAIN]. The result is the largest power of two +/// that keeps the worst case under half of `i32::MAX`. pub fn cross_frame_accum_scale(spatial_radius: u32, temporal_radius: u32) -> f32 { let refs_per_axis = ((PATCH_SIZE - 1) + 2 * spatial_radius) / STEP + 1; let contribs_per_pass = refs_per_axis as f64 * refs_per_axis as f64 * MAX_K as f64; @@ -82,51 +52,24 @@ pub fn cross_frame_accum_scale(spatial_radius: u32, temporal_radius: u32) -> f32 scale as f32 } -/// The magnitude a single filtered value is clamped to before it enters -/// the accumulator. +/// The magnitude a filtered value is clamped to before it enters the accumulator. /// -/// The filter shrinks DCT coefficients of input already inside `[0, 1]`, -/// so a value this large means the filter has already gone wrong. The -/// clamp exists so that when it does, the accumulator saturates -/// gracefully rather than overflowing `i32` and wrapping into a wildly -/// wrong pixel. See [`ACCUM_SCALE`] for how this bound sizes the scale. +/// The filter shrinks coefficients of input already between 0 and 1, so a value this large means +/// the filter has gone wrong. The clamp makes the accumulator saturate instead of overflowing +/// `i32` into a wildly wrong pixel. pub const ACCUM_CLAMP: f32 = 5.0; -/// The constant every group weight is multiplied by before it reaches -/// the accumulator, so the weights land in a fixed-point-friendly band. -/// -/// A group's weight is `1 / sum(retained coefficient variance)`. That -/// sum is a multiple of the caller's `sigma^2`, so the weight's absolute -/// magnitude tracks `1 / sigma^2` and spans far too wide a range for a -/// fixed-point accumulator to hold directly. -/// -/// Scaling every weight by the same constant is free, because -/// aggregation computes `sum(w * x) / sum(w)` and any factor common to -/// every weight cancels exactly. This returns that constant. -/// -/// `sigma^2 * g_max^2` is the smallest retained sum any group can have. -/// The filter always keeps the group's DC coefficient, whose variance is -/// `sigma^2 * g[0]^2`, and `g[0]` is the profile's largest entry. -/// Dividing by it therefore bounds the normalised weight above by 1, -/// whatever `sigma` and whatever correlation shaping is in use, which is -/// what [`WEIGHT_CLAMP`] relies on. -/// -/// The bound below is `1/512`. Every member carries the plain `sigma^2`, -/// so a group of 512 coefficients sums to at most `512 * sigma^2 * g_max^2`, -/// which puts the normalised weight between `1/512` and 1. +/// The constant every group weight is multiplied by so it lands in a fixed-point-friendly band. /// -/// That range still does not fit `accum`'s fixed point with room to -/// spare, which is why `wsum` counts at [`WEIGHT_GAIN`] times its scale. -/// Without the gain, a group that keeps most of its coefficients rounds -/// away to nothing and takes its pixel with it. +/// A group weight is `1 / sum(retained coefficient variance)`, which tracks `1 / sigma^2` and +/// spans too wide a range for fixed point. Aggregation computes `sum(w * x) / sum(w)`, so a factor +/// common to every weight cancels exactly. /// -/// The [`RECIPROCAL_FLOOR`] fallback covers a `sigma` small enough that -/// `sigma^2 * g_max^2` falls under it, zero included. The filter builds -/// its weight with `safe_reciprocal(sum, RECIPROCAL_FLOOR)`, so below -/// that floor every weight saturates at `1 / RECIPROCAL_FLOOR` instead of -/// following the sum. Taking the larger of the two tracks whichever bound -/// the weight is actually against, so the upper bound of 1 holds either -/// way. +/// The filter always keeps the group DC, whose variance `sigma^2 * g_max^2` is the smallest +/// retained sum, so dividing by it bounds the weight above by 1, the bound [WEIGHT_CLAMP] relies +/// on. A group of 512 coefficients sums to at most `512 * sigma^2 * g_max^2`, so the weight lies +/// between `1/512` and 1. Below [RECIPROCAL_FLOOR] every weight saturates at +/// `1 / RECIPROCAL_FLOOR`, so taking the larger of the two keeps the bound of 1. pub fn weight_scale(sigma: f32, dct_profile: &[f32; 8]) -> f32 { let g_max = dct_profile.iter().copied().fold(0.0f32, f32::max); let norm = sigma * sigma * g_max * g_max; @@ -137,13 +80,10 @@ pub fn weight_scale(sigma: f32, dct_profile: &[f32; 8]) -> f32 { } } -/// The zeroth-order modified Bessel function of the first kind, from its own power series. +/// The zeroth-order modified Bessel function of the first kind, `sum_k ((x / 2)^k / k!)^2`. /// -/// `I0(x) = sum_k ((x / 2)^k / k!)^2`. -/// -/// The terms fall off by more than a factor of `x^2 / (4 k^2)` each time, so the series -/// is short and exact to `f64` well before the loop bound for every `beta` -/// [`kaiser_window`] accepts. +/// The series converges to `f64` precision well inside 32 terms for every `beta` that +/// [kaiser_window] accepts. fn bessel_i0(x: f64) -> f64 { let half = x / 2.0; let mut term = 1.0f64; @@ -155,130 +95,74 @@ fn bessel_i0(x: f64) -> f64 { sum } -/// The separable 8-tap Kaiser window [`scatter_patch`] tapers a patch's contribution with. -/// -/// A patch is one of many covering each pixel, and every one of them -/// made its own threshold decision. Weighting a patch's edge pixels less -/// than its centre blends those decisions together instead of letting -/// each patch's own decision reach its boundary at full strength, which -/// is what BM3D's aggregation window is for. The window is separable, so -/// eight taps cover the whole 8x8 patch: pixel `(i, j)` takes `w[i] * w[j]`. -/// -/// `w[i] = I0(beta * sqrt(1 - (2i / 7 - 1)^2)) / I0(beta)`, the standard -/// Kaiser window over 8 points, normalised so its peak is 1. Larger -/// `beta` tapers harder. BM3D uses 2.0. -/// -/// `beta = 0` returns all ones, because the numerator and denominator -/// are both `I0(0)`. That is the off switch, and it is exactly uniform -/// rather than nearly so, so a caller that wants no window gets the same -/// arithmetic the kernel did before this existed. +/// The separable 8-tap Kaiser window `scatter_patch` tapers a patch's contribution with. /// -/// Every tap is above zero and at most 1, so the window can only shrink a -/// contribution. [`WEIGHT_CLAMP`]'s bound of 1 on a weight and the -/// worst-case accumulator value both still hold, and neither -/// [`ACCUM_SCALE`] nor [`cross_frame_accum_scale`] needs rederiving. +/// Weighting a patch's edges less than its centre blends the threshold decisions of the many +/// patches covering a pixel, as BM3D's aggregation window does. Pixel `(i, j)` takes +/// `w[i] * w[j]`, with `w[i] = I0(beta * sqrt(1 - (2i / 7 - 1)^2)) / I0(beta)` so the peak is 1. +/// Larger `beta` tapers harder, BM3D uses 2.0, and `beta = 0` returns exactly all ones. /// -/// What the window does narrow is the other end of the range. The -/// smallest weight the fixed point has to resolve is scaled by the -/// smallest tap product, `w[0]^2`, which is `0.193` at `beta = 2`. The -/// weight floor lands around 640 units of `wsum` at the default geometry, -/// and around 124 after the corner's taper, so it still survives the -/// rounding [`WEIGHT_GAIN`] exists to keep it above. Only the geometry -/// moves this number, not match quality. +/// Every tap is above 0 and at most 1, so the accumulator bounds behind [ACCUM_SCALE] and +/// [cross_frame_accum_scale] still hold. The smallest weight shrinks by `w[0]^2`, `0.193` at +/// `beta = 2`, which leaves the weight floor around 124 units of `wsum` at the default geometry, +/// still above rounding. Only the geometry moves that floor, not match quality. pub fn kaiser_window(beta: f32) -> [f32; PATCH_SIZE as usize] { let denom = bessel_i0(beta as f64); let last = (PATCH_SIZE - 1) as f64; let mut window = [0.0f32; PATCH_SIZE as usize]; for (i, tap) in window.iter_mut().enumerate() { let position = 2.0 * i as f64 / last - 1.0; - *tap = (bessel_i0(beta as f64 * (1.0 - position * position).sqrt()) / denom) as f32; + let radius = (1.0 - position * position).sqrt(); + let numerator = bessel_i0(beta as f64 * radius); + *tap = (numerator / denom) as f32; } window } /// The magnitude a group weight is clamped to before it enters `wsum`. /// -/// A normalised weight is `weight_scale / sum`, and [`weight_scale`] -/// returns the smallest `sum` any group can have, so the weight never -/// exceeds 1. That is a fifth of the bound a filtered value needs, and -/// [`WEIGHT_GAIN`] is what turns the difference into resolution. +/// [weight_scale] bounds a normalised weight by 1, a fifth of [ACCUM_CLAMP], and [WEIGHT_GAIN] +/// turns that difference into resolution. pub const WEIGHT_CLAMP: f32 = 1.0; /// The extra fixed-point resolution `wsum` gets over `accum`. /// -/// Both accumulators are sized by the same worst case, how many -/// contributions can reach one pixel multiplied by the largest a single -/// contribution can be. A value is bounded by [`ACCUM_CLAMP`] and a -/// weight by the much smaller [`WEIGHT_CLAMP`], so counting the weight at -/// this multiple of the value's scale spends exactly the same `i32` -/// budget while resolving weights this many times finer. -/// -/// The resolution matters because every coefficient's variance is at -/// most `sigma^2 * g_max^2`, so a group of up to 512 of them can push -/// the normalised weight down to `1/512`. A weight that falls below half -/// a fixed-point unit contributes nothing at all to either accumulator, -/// and this is part of what keeps a poorly matched group above that -/// point. -/// -/// [`collab_normalise`] multiplies it back out, so it never reaches a -/// finished pixel. +/// A weight is bounded by [WEIGHT_CLAMP] and a value by the larger [ACCUM_CLAMP], so counting +/// weights at this multiple spends the same `i32` budget while resolving them this many times +/// finer. A weight can fall to `1/512`, and one below half a unit contributes nothing, so the gain +/// keeps poorly matched groups above that point. `collab_normalise` multiplies it back out. pub const WEIGHT_GAIN: f32 = ACCUM_CLAMP / WEIGHT_CLAMP; -/// Converts one weighted value into the accumulator's fixed point, at -/// `scale` ([`ACCUM_SCALE`] for a single-frame accumulator, or -/// [`cross_frame_accum_scale`]'s return value for a cross-frame ring). +/// Converts one weighted value into fixed point at `scale`. /// -/// Rounds rather than truncating. A filtered value is never negative, so -/// a toward-zero cast would bias every contribution the same way, and it -/// biases the value more than the weight it is divided by: at a weight of -/// three fixed-point units a value of 0.2 truncates to nothing while its -/// weight still counts three. The result is a weighted mean pulled toward -/// black, worst exactly where the weights are smallest. +/// `scale` is [ACCUM_SCALE] or a [cross_frame_accum_scale] result. It rounds rather than +/// truncates, because filtered values are never negative and truncation would pull the weighted +/// mean toward black where weights are smallest. #[cube] pub fn to_fixed(value: f32, scale: f32) -> i32 { let clamped = f32::clamp(value, -ACCUM_CLAMP, ACCUM_CLAMP); f32::round(clamped * scale) as i32 } -/// Converts one group weight into `wsum`'s fixed point, which counts at -/// [`WEIGHT_GAIN`] times the `scale` [`to_fixed`] uses. +/// Converts one group weight into `wsum`'s fixed point, at [WEIGHT_GAIN] times `scale`. /// -/// Rounds for the same reason [`to_fixed`] does. +/// It rounds for the same reason as `to_fixed`. #[cube] pub fn to_fixed_weight(weight: f32, scale: f32) -> i32 { let clamped = f32::clamp(weight, 0.0f32, WEIGHT_CLAMP); f32::round(clamped * scale * WEIGHT_GAIN) as i32 } -/// Adds one filtered patch to the accumulators at its own position, -/// inside whichever frame's region of the accumulators it belongs to. -/// -/// `value` is this thread's pixel of the patch, and `weight` the -/// normalised weight of the group the patch came from. Every thread in -/// the cube owns one of the patch's 64 pixels, so one call per member -/// scatters the whole patch. +/// Adds this thread's pixel of one filtered patch to the accumulators. /// -/// `kaiser` holds [`kaiser_window`]'s 8 taps, which taper the patch's -/// contribution toward its edges. A caller that wants no taper passes -/// eight ones, which [`kaiser_window`] returns at `beta = 0`. +/// Each of the cube's 64 threads owns the patch pixel `tid` picks, so one call per member scatters +/// the whole patch. `kaiser` holds [kaiser_window]'s 8 taps, and eight ones disable the taper. +/// `accum` and `wsum` hold one `frame_pixels` region per frame in ring-slot order, and +/// `frame_slot` picks this member's region. /// -/// `accum`/`wsum` hold one region per frame in a caller's window, each -/// `frame_pixels` (`width * height`) pixels wide, laid out back to back -/// in ring-slot order. `frame_slot` selects the region this member's own -/// frame owns. A single-frame caller passes `frame_slot = 0`, which folds -/// the frame offset away to nothing. -/// -/// `write_weight` adds the weight itself to `wsum`. Aggregation needs one -/// weight per covering patch, not one per channel, so only the pass over -/// the first channel sets this. -/// -/// `accum_scale` is the fixed-point scale [`to_fixed`] converts into, -/// [`ACCUM_SCALE`] for a single-frame accumulator or -/// [`cross_frame_accum_scale`]'s return value for a cross-frame ring. -/// Either cancels in [`collab_normalise`], so a caller picks whichever -/// matches how many passes can write into its accumulator. `wsum` counts -/// at [`WEIGHT_GAIN`] times that scale, which [`collab_normalise`] -/// multiplies back out. +/// `write_weight` adds the weight to `wsum`, and only the first channel's pass sets it because +/// aggregation needs one weight per patch. `accum_scale` is [ACCUM_SCALE] or a +/// [cross_frame_accum_scale] result, and `wsum` counts at [WEIGHT_GAIN] times it. #[cube] #[expect( clippy::too_many_arguments, @@ -305,10 +189,9 @@ pub fn scatter_patch( let col = tid % PATCH_SIZE; let local_pixel = (patch_y + row) * width + patch_x + col; let pixel = frame_slot * frame_pixels + local_pixel; - // The window multiplies the value and the weight by the same factor. - // `collab_normalise` divides one accumulator by the other, so it - // cancels wherever the coverage is uniform and reweights the blend - // where it is not, rather than shifting the pixel's level. + + // The window scales the value and the weight alike, so it cancels where coverage is uniform + // and only reweights the blend where it is not. let window = kaiser[row as usize] * kaiser[col as usize]; let weight = weight * window; Atomic::fetch_add( @@ -320,24 +203,13 @@ pub fn scatter_patch( } } -/// Clears both accumulators ahead of a filter pass, starting -/// `frame_offset` pixels into them. -/// -/// `accum` and `wsum` may hold more than one frame's worth of pixels, -/// laid out back to back in ring-slot order the way [`scatter_patch`] -/// addresses them, so `frame_offset` (in pixels, the same unit `pixels` -/// is) picks out which region this call zeroes. `frame_offset = 0` with -/// `pixels` covering the whole buffer zeroes it all in one call. +/// Zeroes `pixels` pixels of both accumulators, starting `frame_offset` pixels in. /// -/// The `accum` region is `pixels * stored_ch` slots wide and `wsum`'s is -/// `pixels` wide, so the loop is sized for the larger one and the weight -/// write is masked off past its end. -/// -/// Each thread steps forward by the whole grid's thread count rather than -/// owning one slot. A caller's dispatch grid is clamped to the GPU's -/// 65,535-workgroups-per-dimension limit, which a 4K or 8K frame can -/// exceed at this kernel's 256-thread block size, and striding is what -/// still reaches every slot past that clamp. +/// `client.empty` memory is not zeroed, so the accumulators must be cleared before a pass, and a +/// reused ring region must be cleared again before a later pass writes into it. The +/// `accum` region is `pixels * stored_ch` slots and `wsum`'s is `pixels`, so the weight write is +/// masked past its end. The loop is strided so the grid stays under the 65,535 workgroup limit, +/// which a 4K or 8K frame exceeds at 256 threads per block. #[cube(launch_unchecked)] pub fn collab_zero_accum( accum: &mut Array>, @@ -357,35 +229,15 @@ pub fn collab_zero_accum( } } -/// Turns one frame's region of the accumulators into a finished frame -/// plane. -/// -/// Each pixel's output is the weighted mean of every filtered patch that -/// covered it, which is `accum / wsum`. -/// -/// Both sides carry whatever fixed-point scale the caller's own -/// [`scatter_patch`] calls used, [`ACCUM_SCALE`] or a -/// [`cross_frame_accum_scale`] result, so that scale divides out and -/// never appears here. The weight sum counts at [`WEIGHT_GAIN`] times -/// that scale, which is the one factor that does not cancel, so the -/// ratio is multiplied by it. -/// -/// `accum`/`wsum` may hold more than one frame's worth of pixels, laid -/// out back to back in ring-slot order the way [`scatter_patch`] -/// addresses them, and `frame_offset` (in pixels) picks out which region -/// this call reads. `output` is always exactly one frame wide, so it is -/// indexed by the plain, offset-free pixel position. +/// Writes `accum / wsum` for one frame's region of the accumulators into `output`. /// -/// The weight sum is never zero. A group always contains its own -/// reference patch, and the references alone cover every pixel between -/// one and nine times over, since they sit on a grid of stride `STEP` and -/// are `PATCH_SIZE` wide. Coverage alone is not enough, because a weight -/// small enough to round to nothing would leave a covered pixel with an -/// empty weight sum. [`WEIGHT_GAIN`] keeps every group's weight above -/// that point. +/// Each pixel becomes the weighted mean of every filtered patch covering it. The fixed-point scale +/// cancels, so only [WEIGHT_GAIN] is multiplied back. `frame_offset` (in pixels) picks the region, +/// and `output` is one frame wide. /// -/// If the weight sum ever were to be zero anyway, the guard below returns -/// the accumulator untouched rather than a NaN or an infinity. +/// The weight sum is never zero, because the references cover every pixel between one and nine +/// times and [WEIGHT_GAIN] keeps every group weight from rounding to nothing. A zero sum returns +/// the accumulator untouched rather than a NaN. #[cube(launch_unchecked)] #[expect( clippy::too_many_arguments, @@ -409,85 +261,71 @@ pub fn collab_normalise( let local_pixel = y * width + x; let pixel = frame_offset + local_pixel; - let w = wsum[pixel as usize]; + let weight_total = wsum[pixel as usize]; - let mut out = Vector::::empty(); + let mut pixel_out = Vector::::empty(); #[unroll] - for c in 0..channels { - let a = accum[(pixel * stored_ch + c) as usize] as f32; - let mut v = a; - if w != 0i32 { - v = a * WEIGHT_GAIN / (w as f32); + for channel in 0..channels { + let accumulated = accum[(pixel * stored_ch + channel) as usize] as f32; + let mut value = accumulated; + if weight_total != 0i32 { + value = accumulated * WEIGHT_GAIN / (weight_total as f32); } - out[c as usize] = v; + pixel_out[channel as usize] = value; } - output[local_pixel as usize] = out; + output[local_pixel as usize] = pixel_out; } -// Ties the fixed-point headroom argument in `ACCUM_SCALE`'s docs to the -// constants it is actually derived from. A larger patch or a larger -// group both raise the number of contributions one pass can add to a -// pixel, and the scale has to come down to match either moving. +// Recheck ACCUM_SCALE's headroom if the patch area or group size moves, since both raise the +// number of contributions one pass adds to a pixel. const _: () = assert!( PATCH_AREA == 64 && MAX_K == 8, "recheck ACCUM_SCALE's headroom, the per-pass contribution bound moved" ); -// `cross_frame_accum_scale` needs no equivalent assertion. It re-derives -// its scale from `PATCH_SIZE`, `STEP`, and `MAX_K` on every call, so -// there is no baked-in number for a compile-time check to protect. Its -// arithmetic is guarded instead by a `debug_assert!` inside the function -// and by the exhaustive test below. - #[cfg(test)] mod tests { use super::*; use crate::nl4d::MAX_KAISER_BETA; - /// A filtered value is never negative, so a toward-zero cast biases - /// every contribution the same way. Rounding is what keeps the - /// weighted mean `collab_normalise` computes centred on the value the - /// filter actually produced. #[test] fn to_fixed_rounds_rather_than_truncating() { - assert_eq!(to_fixed(0.6, 1.0), 1); - assert_eq!(to_fixed(0.4, 1.0), 0); - assert_eq!(to_fixed(-0.6, 1.0), -1); + let rounded_up = to_fixed(0.6, 1.0); + let rounded_down = to_fixed(0.4, 1.0); + let negative = to_fixed(-0.6, 1.0); + + assert_eq!(rounded_up, 1); + assert_eq!(rounded_down, 0); + assert_eq!(negative, -1); } - /// A weight small enough to round away in `accum`'s fixed point still - /// reaches `wsum`, because `WEIGHT_GAIN` buys back exactly the - /// resolution the weight's smaller bound leaves unused. #[test] fn to_fixed_weight_resolves_finer_than_to_fixed() { let scale = 65_536.0f32; let weight = 0.3 / scale; - assert_eq!(to_fixed(weight, scale), 0); - assert_eq!(to_fixed_weight(weight, scale), 2); + let as_value = to_fixed(weight, scale); + let as_weight = to_fixed_weight(weight, scale); + + assert_eq!(as_value, 0); + assert_eq!(as_weight, 2); } - /// The weight's own bound, which is what lets `WEIGHT_GAIN` spend the - /// same `i32` budget as a value clamped to `ACCUM_CLAMP`. #[test] fn to_fixed_weight_clamps_at_one() { let scale = 1_024.0f32; - assert_eq!(to_fixed_weight(4.0, scale), to_fixed_weight(WEIGHT_CLAMP, scale),); - assert_eq!(to_fixed_weight(-1.0, scale), 0); + let above = to_fixed_weight(4.0, scale); + let at_clamp = to_fixed_weight(WEIGHT_CLAMP, scale); + let negative = to_fixed_weight(-1.0, scale); + + assert_eq!(above, at_clamp); + assert_eq!(negative, 0); } - // `Nl4dParams::validate` in `nl4d/params.rs` is what actually enforces - // these ranges; they are repeated here as plain numbers rather than - // imported, so this test does not depend on `nl4d` at all and keeps - // exercising the true worst case even if that module's ranges ever - // narrow. + // The ranges `Nl4dParams::validate` enforces, as plain numbers so the test keeps covering the + // true worst case if those ranges narrow. const SPATIAL_RADIUS_RANGE: std::ops::RangeInclusive = 1..=16; const TEMPORAL_RADIUS_RANGE: std::ops::RangeInclusive = 1..=8; - /// Recomputes the worst-case accumulator value the same way - /// [`cross_frame_accum_scale`]'s own `debug_assert!` does, for every - /// `(spatial_radius, temporal_radius)` pair the validated parameter - /// ranges allow, rather than only the pair a debug build happens to - /// exercise at runtime. #[test] fn every_spatial_and_temporal_radius_stays_under_the_safety_budget() { let budget = i32::MAX as f64 / CROSS_FRAME_SAFETY_FACTOR; @@ -517,10 +355,6 @@ mod tests { } } - /// Deriving the scale per configuration, rather than sizing one - /// constant for the widest configuration allowed, is what lets a - /// typical configuration keep more fixed-point precision. The - /// defaults should clear `2^15`. #[test] fn defaults_keep_at_least_a_2_15_scale() { let floor = 32_768.0f32; @@ -535,30 +369,36 @@ mod tests { #[test] fn kaiser_window_at_beta_zero_is_exactly_one_everywhere() { - assert_eq!(kaiser_window(0.0), [1.0f32; PATCH_SIZE as usize]); + let window = kaiser_window(0.0); + assert_eq!(window, [1.0f32; PATCH_SIZE as usize]); } - /// An even tap count puts the centre of the span between taps 3 and - /// 4, so those two are equal rather than one being above the other, - /// and the rise is checked up to that pair. + /// An even tap count puts the centre between taps 3 and 4, so the rise is checked up to that + /// pair. #[test] fn kaiser_window_is_symmetric_and_rises_to_the_centre() { for beta in [1.0f32, 2.0, 4.0, MAX_KAISER_BETA] { - let w = kaiser_window(beta); + let window = kaiser_window(beta); for i in 0..4 { assert!( - (w[i] - w[7 - i]).abs() < 1e-6, + (window[i] - window[7 - i]).abs() < 1e-6, "beta {beta}: tap {i} is {} and its mirror {}", - w[i], - w[7 - i], + window[i], + window[7 - i], ); + if i < 3 { - assert!(w[i + 1] > w[i], "beta {beta}: tap {} is not above tap {i}", i + 1,); + assert!( + window[i + 1] > window[i], + "beta {beta}: tap {} is not above tap {i}", + i + 1, + ); } } + assert!( - w.iter().all(|&t| t > 0.0 && t <= 1.0), - "beta {beta}: a tap is not above zero and at most 1, {w:?}", + window.iter().all(|&tap| tap > 0.0 && tap <= 1.0), + "beta {beta}: a tap is not above zero and at most 1, {window:?}", ); } } @@ -566,21 +406,21 @@ mod tests { #[test] fn kaiser_window_end_taps_are_the_bessel_ratio() { for beta in [1.0f32, 2.0, 4.0] { - let w = kaiser_window(beta); + let window = kaiser_window(beta); let expected = (1.0 / bessel_i0(beta as f64)) as f32; assert!( - (w[0] - expected).abs() < 1e-6, + (window[0] - expected).abs() < 1e-6, "beta {beta}: end tap {} against the ratio {expected}", - w[0], + window[0], ); - assert!((w[7] - expected).abs() < 1e-6); + assert!((window[7] - expected).abs() < 1e-6); } - // The figure the doc's margin arithmetic uses. - assert!((kaiser_window(2.0)[0] - 0.4388).abs() < 1e-3); + + // Pins the end tap behind the `w[0]^2 = 0.193` at `beta = 2` quoted in `kaiser_window`'s doc. + let beta_two_window = kaiser_window(2.0); + assert!((beta_two_window[0] - 0.4388).abs() < 1e-3); } - /// Ties the `392`-contribution figure in [`ACCUM_SCALE`]'s docs to - /// the model `cross_frame_accum_scale` is built on. #[test] fn contribution_model_reproduces_the_documented_392_at_the_default_spatial_radius() { let spatial_radius = 9u32; diff --git a/av-denoise-core/src/collab/kernels/fused/grid.rs b/av-denoise-core/src/collab/kernels/fused/grid.rs index 3089844..b84d69e 100644 --- a/av-denoise-core/src/collab/kernels/fused/grid.rs +++ b/av-denoise-core/src/collab/kernels/fused/grid.rs @@ -29,7 +29,7 @@ fn haar_strided_level_fwd( } } -/// The inverse of [haar_strided_level_fwd]. +/// The inverse of `haar_strided_level_fwd`. #[cube] fn haar_strided_level_inv( vals: &mut Array, @@ -159,7 +159,7 @@ pub(crate) fn grid_fwd(stack: &mut Array, #[comptime] grid_frames: u32) { } } -/// The inverse of [grid_fwd]. +/// The inverse of `grid_fwd`. #[cube] pub(crate) fn grid_inv(stack: &mut Array, #[comptime] grid_frames: u32) { let volumes = comptime!(MAX_K / grid_frames); @@ -189,7 +189,7 @@ pub(crate) fn grid_inv(stack: &mut Array, #[comptime] grid_frames: u32) { } } -/// Propagates each member's variance to the coefficient it lands on under [grid_fwd]. +/// Propagates each member's variance to the coefficient it lands on under `grid_fwd`. #[cube] pub(crate) fn grid_variance(v: &mut Array, #[comptime] grid_frames: u32) { let volumes = comptime!(MAX_K / grid_frames); @@ -275,14 +275,14 @@ pub(crate) fn grid_inv_host(column: &[f32; 8], grid_frames: u32) -> [f32; 8] { vals } -/// Propagates independent per-member variances through [grid_fwd_host]. +/// Propagates independent per-member variances through `grid_fwd_host`. /// /// Both outputs of a Haar pair carry the mean variance of its inputs. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] -pub(crate) fn grid_variance_host(v: &[f32; 8], grid_frames: u32) -> [f32; 8] { +pub(crate) fn grid_variance_host(variances: &[f32; 8], grid_frames: u32) -> [f32; 8] { let frames = grid_frames as usize; let volumes = 8 / frames; - let mut vals = *v; + let mut vals = *variances; let average = |vals: &mut [f32; 8], start: usize, stride: usize, len: usize| { let half = len / 2; let snapshot: Vec = (0..len).map(|k| vals[start + k * stride]).collect(); diff --git a/av-denoise-core/src/collab/kernels/fused/mod.rs b/av-denoise-core/src/collab/kernels/fused/mod.rs index 82d3c29..9fd6770 100644 --- a/av-denoise-core/src/collab/kernels/fused/mod.rs +++ b/av-denoise-core/src/collab/kernels/fused/mod.rs @@ -33,271 +33,94 @@ const NOISE_CURVE_SCALE_MIN: f32 = 0.33; /// The largest factor the noise curve may scale the luma threshold by. const NOISE_CURVE_SCALE_MAX: f32 = 3.0; -// The widest neighbour index this kernel ever packs is `2 * radius`, -// one past the last neighbour, and `radius` is capped at -// `MAX_TEMPORAL_RADIUS`. `pack_pos_t` gives `t` bits 26-31, so a value -// of 64 or more would silently overflow into nothing and corrupt the -// word. This ties the packer's field width to the radius ceiling that -// feeds it, so the bound is checked at compile time rather than -// assumed at the call site. +// The widest neighbour index packed is `2 * radius`, and `pack_pos_t`'s 6-bit `t` field silently +// corrupts the word at 64 or more. const _: () = assert!( 2 * MAX_TEMPORAL_RADIUS < 64, "pack_pos_t's 6-bit t field must hold every neighbour index collab_fused packs" ); -// A lane holds one 8-value column of each of `MAX_K` members, so the -// whole group fits `PATCH_AREA` slots only while the group size and the -// patch side are the same number. The stack transform's predicate -// ladder below also names the three levels 8, 4 and 2 outright. +// A lane holds one 8-value column of each of `MAX_K` members, so the group fits `PATCH_AREA` +// slots only while the group size and the patch side match. The stack transform's predicate +// ladder also names the levels 8, 4 and 2 outright. const _: () = assert!( MAX_K == PATCH_SIZE && MAX_K == 8, "collab_fused's per-lane group array and its three-level stack transform are written for \ MAX_K == PATCH_SIZE == 8" ); -// A candidate that never placed carries the distance `3.0e38`, written -// as a literal at each use below. A real distance is a sum of at most -// `PATCH_AREA` squared differences between values in `[0, 1]`, scaled -// by at most 3, so it never exceeds 192. `3.0e38` sits far above that -// and just below `f32::MAX`, so it always compares greater than a live -// candidate. The self-match takes `-1.0e38` at the other end, which -// sorts it below every real distance and pins it into slot 0. Both are -// literals rather than consts or `f32::INFINITY` because cubecl treats -// all of those as compile-time-only, and the shift-insert needs genuine -// mutable runtime variables. - -/// Groups each reference patch with the patches most similar to it, -/// filters the whole group jointly with a hard threshold in the -/// transform domain, and scatters every filtered member back into its -/// own frame. -/// -/// # Work decomposition -/// -/// One cube of 64 threads owns eight reference patches. Each 8-lane -/// group owns one of them, and lane `sub` of a group owns column `sub` -/// of every patch that group touches. That one mapping serves both -/// halves of the kernel. A candidate's 64 pixel differences are spread -/// eight ways during matching, and [`plane_ssd_reduce8`] folds the eight -/// column sums into the whole patch distance. A member's 64 filtered -/// pixels are spread the same eight ways during filtering, so both the -/// candidate reads and the scatter writes are coalesced. -/// -/// The reference patch's own column stays in registers for the whole -/// matching phase. Candidate pixels are read straight from global -/// memory. Neighbouring reference patches search heavily overlapping -/// windows at a step of 4, so the cache already serves those reads well -/// and a shared-memory tile would only cost occupancy. -/// -/// A row of references rarely divides into eights, so the last cube of -/// a row runs groups whose reference patch is past the end. A 1080p -/// frame has 479 references across, so this is a shipped path rather -/// than an edge case. Those groups stay live through the whole kernel, -/// working on a clamped copy of the last real reference, and are gated -/// only where they would write. -/// -/// # Barriers -/// -/// [`transpose8`] carries the only barrier inside the group-processing -/// loops. Every lane of the cube reaches it the same number of times, -/// because the transposes sit in fully unrolled loops with no run-time -/// condition around them. Nothing returns early, a dead group runs the -/// whole kernel, and the group size only ever gates which iterations do -/// arithmetic, never how many barriers a lane reaches. A workgroup -/// barrier reached by only part of the workgroup is undefined, so that -/// property is what the write gating and the clamped reference index -/// exist to preserve. -/// -/// The basis fill carries one more barrier, before either transform -/// runs. It is unconditional and sits before `live` is computed, so -/// every lane reaches it whatever the reference index later clamps to. -/// -/// # Search space -/// -/// The centre frame contributes the `spatial_radius` rectangle around -/// the reference patch, clipped to the frame. -/// -/// Each volume anchor then searches every neighbour frame. It contributes -/// one `refine` rectangle per motion block whose span contains the -/// anchor, each around the position that block's vector predicts the -/// anchor moved to, clipped the same way. A block grid at a step below -/// `blksize` gives several such blocks, so an anchor is searched wherever -/// any block covering it points. A position reached by more than one of -/// them is scored once, by the first rectangle that reaches it. -/// -/// Clipping each rectangle once keeps every candidate within it a -/// distinct position. Clamping each offset in turn would land several -/// offsets on the same edge position and let one physical patch count as -/// two. -/// -/// # Distance -/// -/// A candidate's distance is the channel-scaled sum of squared pixel -/// differences over the whole patch. -/// -/// # Confidence gate -/// -/// Every candidate stays in the running whatever its distance, so a -/// group fills wherever the search space is large enough. A covering -/// block whose confidence sits below `c_min` never runs the pixel -/// comparison, while the frame's other covering blocks still search. -/// The confidence comes from a motion block every lane of the group -/// shares, so the skip is uniform across the group. A volume left short -/// of frames by this gate makes the whole group fall back to the -/// single-frame group centred on the reference frame, so `c_min` can -/// change the output. -/// -/// # Selection -/// -/// The spatial search keeps the eight best centre-frame positions, one -/// per lane, ascending, through -/// [shift_insert8_gated](crate::collab::kernels::plane_ops::shift_insert8_gated). -/// A tie never displaces an incumbent, and the self-match scores a -/// sentinel below every real distance, which pins it into slot 0. -/// -/// The first `MAX_K / grid_frames` of those become volume anchors. Each -/// anchor keeps its best match in every neighbour frame, and the volume -/// keeps the `grid_frames - 1` best of those in ascending order. A -/// position an earlier volume already holds is skipped, so no patch -/// enters the group twice. -/// -/// # Members -/// -/// A member is a packed position. The neighbour it came from sits in the -/// bits above the coordinates, so the frame it was matched in is -/// recovered from the packed word when matching ends. Member -/// `s * grid_frames + t` is frame `t` of volume `s`, with the anchor at -/// `t = 0`. -/// -/// # Group size -/// -/// A group uses the grid when the spatial search held at least `MAX_K` -/// positions, `k_max` is `MAX_K`, and every volume filled all of its -/// frames. Otherwise it falls back to the single-frame group, the -/// spatial search's positions rounded down to a power of two and capped -/// at `k_max`. The decision is uniform across the group. -/// -/// # What the filter does -/// -/// For each active channel, every member's patch runs through a 2D DCT, -/// so each patch is described by 64 frequency coefficients instead of 64 -/// pixel values. A grid group then runs a Haar along time within each -/// volume and a Haar across the volumes, at each spatial position. A -/// fallback group runs a Haar across its stack instead. Content the group -/// agrees on collects into the low levels. A coefficient survives a hard -/// threshold when its magnitude reaches `lambda_ht` standard deviations -/// of its own propagated noise, where every member carries the plain -/// `sigma[c]^2`. Both transforms then invert. -/// -/// The spatial pass runs as a column DCT in registers, a transpose, and -/// a row DCT in registers, because a lane owns a column and the row pass -/// needs a row. The inverse runs the same three steps backwards, which -/// leaves the lane holding a column again in time for the scatter. -/// -/// Channel 0's threshold is scaled by the frame's noise curve at the -/// reference patch's mean luma, clamped to 0.33..=3. The group weight -/// keeps the plain sigma. With no curve the threshold is unchanged. -/// -/// A strength map scales thresholds further, one multiplier per 8x8 quarter. Each reference -/// patch takes the mean of the four quarters it overlaps. With `map_mode` at -/// [STRENGTH_MAP_LUMA] it scales channel 0's curve ratio before the clamp. With -/// [STRENGTH_MAP_ALL] it scales every channel's threshold when there is no curve. With a curve, -/// channel 0 keeps its curve threshold. [STRENGTH_MAP_OFF] leaves every threshold as it is. -/// -/// With `pooled` set, a coefficient is kept on the mean energy of itself and its four frequency -/// neighbours instead of its own, against `channel_lambda * pool_ratio`. See -/// [pooled_threshold](crate::collab::kernels::fused::pooled::pooled_threshold). -/// -/// The one coefficient that is both the group average and the patch's -/// spatial DC always survives the threshold, whatever its magnitude. A -/// group's mean brightness is signal, not something a noise threshold -/// should be able to zero out. -/// -/// # Group weight -/// -/// `group_weight` is `1 / sum(v_j)` over the coefficients the threshold -/// kept, computed from channel 0 only (luma dominates, and one weight -/// per group keeps aggregation simple downstream). When every member has -/// the same noise variance and the group keeps `n` coefficients this is -/// `1 / (sigma^2 * n)`, the usual inverse-variance weight, so a group -/// whose content agreed enough to keep more of its coefficients is -/// trusted more. Each lane sums the variance it retained over its own -/// eight positions and [`plane_ssd_reduce8`] folds the group's eight -/// partials together, which is why no shared array is needed for it. -/// -/// # Buffers -/// -/// `ring` is the frame ring, laid out one frame after another in -/// physical ring-slot order. `centre_slot` is the slot the pass is -/// centred on and `neighbour_slots` maps a packed neighbour index onto -/// its physical slot. -/// -/// `accum` and `wsum` hold one region per ring slot, the layout -/// [`scatter_patch`] addresses, so a member matched in a neighbour frame -/// scatters into that frame's own region rather than the centre's. -/// `accum_scale` is the fixed-point scale that scatter converts into. -/// -/// `group_weight` holds one weight per reference, and `sigma` one value -/// per stored channel. -/// -/// `noise_curve` holds `NOISE_CURVE_BINS` luma threshold ratios, each -/// sampled at the centre of an equal-width slice of the luma range. The -/// kernel interpolates between neighbouring centres. The curve only takes -/// effect when `curve_valid` is not 0. -/// -/// `strength_map` holds `map_cols * map_rows` multipliers, row-major, laid out by -/// [strength_map_dims](crate::collab::geometry::strength_map_dims). -/// -/// `kaiser` holds [`crate::collab::kernels::aggregate::kaiser_window`]'s 8 taps, which -/// taper each scattered patch toward its edges. Eight ones leave the aggregation uniform. -/// -/// `dct_profile` holds -/// [`crate::collab::kernels::transforms::dct_noise_profile`]'s 8 values. -/// Every member's coefficient variance at DCT position `(u, v)` scales -/// by `dct_profile[u] * dct_profile[v]` before the threshold reads it. -/// At `rho = 0` every entry is `1.0` and the multiply is a no-op. -/// -/// `grid_frames` is the frames per volume, from -/// [grid_frames](crate::collab::grid_frames). At 1 the grid compiles out -/// and every group is a single-frame one. -/// -/// `pool_ratio` scales each channel's lambda into the pooled threshold. It is only read with -/// `pooled` set. -/// -/// # Warp-uniform search -/// -/// `warp_uniform` decides how the spatial and trajectory searches are -/// walked. -/// -/// Both searches are group-scoped work: each 8-lane group owns one -/// reference patch, and every distance is completed by a shuffle across -/// just those eight lanes. Nothing in the algorithm needs the other -/// groups sharing a warp to keep step. -/// -/// The CUDA backend nevertheless lowers each of those shuffles to a -/// `__shfl_*_sync` naming the whole 32-lane warp. On Volta and later -/// such a shuffle waits for every lane it names, so a group still -/// searching blocks on groups that have already left the loop, and those -/// never come back. The clipped rectangles and the `c_min` skip both -/// give neighbouring groups different trip counts, so the warp -/// deadlocks and the launch never retires a frame. -/// -/// Setting `warp_uniform` walks fixed, comptime-sized rectangles -/// instead, in both searches, and masks every position the other walk -/// skips, whether clipped, gated, already scored or already claimed. -/// Every group in a warp then takes the same number of turns through the -/// same shuffles. A masked turn carries the same `3.0e38` an unfilled slot -/// holds, so it can never displace one. -/// -/// The candidates that do score, and the order they are offered in, are -/// exactly the ones the unset path visits, so both settings produce the -/// same group. Leave it unset on the wgpu backends, whose subgroup -/// operations reconverge on their own and which would only pay for the -/// dead turns. [`crate::collab::needs_warp_uniform_search`] is what -/// picks it per runtime. -/// -/// # Compilation cost -/// -/// The transforms unroll fully, which keeps the whole group in registers. +/// Groups each reference patch with its most similar patches, hard-thresholds the group in the +/// transform domain and scatters every member back into its own frame. +/// +/// A cube of 64 threads owns eight reference patches. Each 8-lane group owns one, and lane `sub` +/// owns column `sub` of every patch the group touches, so candidate reads and scatter writes are +/// coalesced and `plane_ssd_reduce8` completes each distance. Candidates are read straight from +/// global memory, because neighbouring references search overlapping windows the cache already +/// serves. +/// +/// Every lane must reach every barrier, since a barrier reached by only part of a workgroup is +/// undefined. The basis fill barrier is unconditional, and the `transpose8` barriers sit in fully +/// unrolled loops with no runtime condition around them. A group past the end of a row, which +/// every 1080p row has, works on a clamped copy of the last real reference and is gated only +/// where it writes. +/// +/// The centre frame contributes the `spatial_radius` rectangle around the reference, clipped to +/// the frame, and the best eight positions are kept. The self-match scores a sentinel below every +/// real distance, which pins it into slot 0. The first `MAX_K / grid_frames` positions become +/// volume anchors. In each neighbour frame an anchor searches one `refine` rectangle per covering +/// motion block, around where that block's vector moves it, and keeps its best match. Each +/// rectangle is clipped once so every candidate is a distinct position, and a position reached +/// twice or already in the group is scored once. A block whose confidence is below `c_min` is +/// skipped, uniformly across the group. +/// +/// A member is a `pack_pos_t` word, and member `s * grid_frames + t` is frame `t` of volume `s`. +/// The grid is used when the spatial search held at least `MAX_K` positions, `k_max` is `MAX_K` +/// and every volume filled. Otherwise the group falls back to the single-frame group, its +/// positions rounded down to a power of two and capped at `k_max`. A block skipped by `c_min` can +/// leave a volume short of frames, which makes the group fall back, so `c_min` can change the +/// output. +/// +/// For each channel, every member runs through a 2D DCT as a column pass, a transpose and a row +/// pass. A grid group then runs a Haar along time and across volumes, and a fallback group a Haar +/// across its stack. The transforms unroll fully, which keeps the whole group in registers. A +/// coefficient survives when its magnitude reaches `lambda_ht` standard deviations of its +/// propagated noise, with every member carrying `sigma[c]^2`. The group DC of the spatial DC +/// always survives, because a group's mean brightness is signal. Both transforms then invert. +/// +/// Channel 0's threshold is scaled by the noise curve at the reference's mean luma, interpolated +/// between bin centres and clamped to 0.33..=3, while the group weight keeps the plain sigma. The +/// strength map multiplier is the mean of the four 8x8 quarters the reference overlaps. +/// [STRENGTH_MAP_LUMA] scales channel 0's curve ratio before the clamp. [STRENGTH_MAP_ALL] scales +/// the chroma thresholds, and channel 0's too when there is no curve. [STRENGTH_MAP_OFF] leaves +/// every threshold alone. With `pooled` set, `pooled_threshold` keeps a coefficient on the mean +/// energy of itself and its four frequency neighbours, against `channel_lambda * pool_ratio`. +/// +/// `group_weight` gets `1 / sum(v_j)` over the kept coefficients, the inverse-variance weight, so +/// a group that keeps more coefficients is trusted more. It comes from channel 0 only, because +/// luma dominates and one weight per group keeps aggregation simple. `weight_scale` maps it into +/// the fixed-point band before the scatter. +/// +/// `ring` holds frames in ring-slot order. `centre_slot` is the slot the pass is centred on. +/// `neighbour_slots` maps a packed neighbour index to its slot. `accum` and `wsum` hold one region +/// per slot, so a neighbour-frame member scatters into its own frame's region. `accum_scale` is +/// the fixed-point scale of that scatter. `sigma` holds one value per stored channel. +/// `noise_curve` holds `NOISE_CURVE_BINS` ratios and is read only when `curve_valid` is not 0. +/// `strength_map` holds `map_cols * map_rows` row-major multipliers laid out by +/// [strength_map_dims](crate::collab::geometry::strength_map_dims). `kaiser` holds +/// [kaiser_window](crate::collab::kernels::aggregate::kaiser_window)'s taps. `dct_profile` holds +/// [dct_noise_profile](crate::collab::kernels::transforms::dct_noise_profile)'s values, and +/// coefficient `(u, v)`'s variance scales by `dct_profile[u] * dct_profile[v]`. `grid_frames` is +/// the frames per volume from [grid_frames](crate::collab::grid_frames), and 1 compiles the grid +/// out. +/// +/// `warp_uniform` walks both searches over fixed comptime rectangles and masks the positions the +/// clipped walk skips, so both settings score the same candidates in the same order. The clipped +/// rectangles and the `c_min` skip give neighbouring groups different trip counts. The CUDA +/// backend lowers each group-scoped shuffle to a `__shfl_*_sync` over the whole warp, so on Volta +/// and later those different trip counts deadlock the warp. A masked turn carries the +/// `3.0e38` an unfilled slot holds, so it never displaces a match. The wgpu backends reconverge on +/// their own, and [needs_warp_uniform_search](crate::collab::needs_warp_uniform_search) picks the +/// setting per runtime. #[cube(launch_unchecked)] #[expect( clippy::too_many_arguments, @@ -351,61 +174,55 @@ pub fn collab_fused( pool_ratio: f32, #[comptime] pooled: bool, ) { - let tid = UNIT_POS_X; - let grp = tid / 8u32; - let sub = tid % 8u32; + let thread_id = UNIT_POS_X; + let group = thread_id / 8u32; + let sub = thread_id % 8u32; let base = group_base(); let max_x = comptime!(width - PATCH_SIZE); let max_y = comptime!(height - PATCH_SIZE); - // The spatial basis, filled once and read by every lane for the rest - // of the kernel. It is 256 B against the transpose buffer's 2,080 B, - // and every lane reads all 64 of its entries, so keeping it shared - // costs nothing a per-lane copy would save. Shared memory is not - // what bounds this kernel's occupancy in any case, registers are. + // The basis is 256 B against the transpose buffer's 2,080 B, and registers rather than shared + // memory bound this kernel's occupancy, so it stays shared. let mut basis = SharedMemory::::new(PATCH_AREA as usize); - let mut tbuf = SharedMemory::::new(comptime!(8 * 65) as usize); - fill_dct8_basis(&mut basis, tid); + let mut transpose_buf = SharedMemory::::new(comptime!(8 * 65) as usize); + fill_dct8_basis(&mut basis, thread_id); sync_cube(); - // A dead group keeps working on the last real reference of the row - // so every read stays inside the frame and every lane reaches every - // barrier. `live` is what stops it writing. - let ref_x_index = CUBE_POS_X * 8u32 + grp; + // A dead group works on the last real reference of the row, so every read stays inside the + // frame and every lane reaches every barrier. `live` stops it writing. + let ref_x_index = CUBE_POS_X * 8u32 + group; let live = ref_x_index < refs_x; let ref_x_clamped = ref_x_index.min(refs_x - 1u32); - let rx = (ref_x_clamped * STEP).min(max_x); - let ry = (CUBE_POS_Y * STEP).min(max_y); + let ref_x = (ref_x_clamped * STEP).min(max_x); + let ref_y = (CUBE_POS_Y * STEP).min(max_y); - // Column `sub` of the reference patch, all channels, in registers - // for the whole search. + // Column `sub` of the reference patch, all channels, held in registers for the whole search. let mut current = Array::::new(comptime!(PATCH_SIZE * channels) as usize); #[unroll] for r in 0..PATCH_SIZE { - let px = read_line(ring, rx + sub, ry + r, centre_slot, width, height); + let pixel = read_line(ring, ref_x + sub, ref_y + r, centre_slot, width, height); #[unroll] for c in 0..channels { - current[(r * channels + c) as usize] = px[c as usize]; + current[(r * channels + c) as usize] = pixel[c as usize]; } } + // An unplaced candidate carries `3.0e38`, far above the largest real distance of 192 and below + // `f32::MAX`, and the self-match carries `-1.0e38`. They are literals because cubecl treats + // consts and `f32::INFINITY` as comptime-only, and the shift-insert needs runtime variables. let mut best_d = 3.0e38f32; let mut best_pos = 0u32; - // One scalar for the whole kernel, from the channel count. It - // multiplies the completed 64-pixel distance, not each squared - // difference. + // The channel scale multiplies the completed 64-pixel distance, not each squared difference. let scale = channel_scale(channels); - // The number of positions the spatial rectangle holds, which fixes the fallback group size - // below. let n_live = spatial_search( ring, ¤t, - rx, - ry, + ref_x, + ref_y, centre_slot, sub, base, @@ -459,10 +276,10 @@ pub fn collab_fused( let mut anchor = Array::::new(comptime!(PATCH_SIZE * channels) as usize); #[unroll] for r in 0..PATCH_SIZE { - let px = read_line(ring, anchor_x + sub, anchor_y + r, centre_slot, width, height); + let pixel = read_line(ring, anchor_x + sub, anchor_y + r, centre_slot, width, height); #[unroll] for c in 0..channels { - anchor[(r * channels + c) as usize] = px[c as usize]; + anchor[(r * channels + c) as usize] = pixel[c as usize]; } } @@ -496,8 +313,7 @@ pub fn collab_fused( ); } - // A grid needs a full spatial search for its anchors and every volume's last frame - // filled. The list is ascending, so a filled last slot means the whole volume is. + // The list is ascending, so a filled last slot means the whole volume is filled. use_grid = k_use == MAX_K; #[unroll] for volume in 0..volumes { @@ -512,34 +328,29 @@ pub fn collab_fused( k_use = select(use_grid, MAX_K, k_use); } - // The frame each member sits in, from its packed word, once before the channel loop. - // - // The frame is picked with [`select`] rather than a branch. A frame index that reaches - // [`read_line`] through a branch trips a bug in cubecl 0.10's global value numbering, which - // panics while compiling the shader and leaves the launch to do nothing at all. + // The frame is picked with `select` rather than a branch, because a frame index that reaches + // `read_line` through a branch panics cubecl 0.10's GVN pass and the launch silently does + // nothing. let mut member_slot = Array::::new(MAX_K as usize); #[unroll] for m in 0..MAX_K { let packed = member_pos[m as usize]; - let mt = unpack_t(packed); - // Clamped so the read below stays in range for a centre-frame member, whose value - // `select` then discards. The clamp lands on index 0, so it needs `neighbour_slots` to - // hold at least one entry. That is what every caller actually supplies, including - // `radius = 0` launches such as `Setup::spatial_only` and the standalone launch - // documented at `nl4d::tests::pipeline`, which still pass a one-element - // `neighbour_slots` even though there is no real neighbour to read. - let neighbour = u32::max(mt, 1u32) - 1u32; - member_slot[m as usize] = select(mt > 0u32, neighbour_slots[neighbour as usize], centre_slot); + let neighbour_field = unpack_t(packed); + // Clamped so the read stays in range for a centre-frame member, whose value `select` + // discards. This needs `neighbour_slots` to hold at least one entry, even at `radius = 0`. + let neighbour = u32::max(neighbour_field, 1u32) - 1u32; + member_slot[m as usize] = select( + neighbour_field > 0u32, + neighbour_slots[neighbour as usize], + centre_slot, + ); } - // The correlation profile is separable and the same for every - // member, so the lane's own half of it is read once. Lane `sub` - // ends up owning vertical frequency `sub` at every horizontal - // frequency, see the transform order below. + // Lane `sub` ends up owning vertical frequency `sub`, so its half of the separable profile is + // read once. let prof_sub = dct_profile[sub as usize]; - // The reference patch's mean luma picks its place on the frame's noise curve. Every lane - // reaches the reduction, so the group stays converged. + // Every lane reaches the luma reduction, so the group stays converged. let mut column_luma = 0.0f32; #[unroll] for r in 0..PATCH_SIZE { @@ -555,59 +366,51 @@ pub fn collab_fused( let lower_ratio = noise_curve[lower_bin as usize]; let upper_ratio = noise_curve[(lower_bin + 1u32) as usize]; let ratio = lower_ratio + (upper_ratio - lower_ratio) * fraction; - let map_scale = strength_map_scale(strength_map, rx, ry, map_cols, map_rows); + let map_scale = strength_map_scale(strength_map, ref_x, ref_y, map_cols, map_rows); let mapped_ratio = select(map_mode == STRENGTH_MAP_LUMA, ratio * map_scale, ratio); let curve_scale = f32::clamp(mapped_ratio, NOISE_CURVE_SCALE_MIN, NOISE_CURVE_SCALE_MAX); let other_lambda = select(map_mode == STRENGTH_MAP_ALL, lambda_ht * map_scale, lambda_ht); let luma_lambda = select(curve_valid != 0u32, lambda_ht * curve_scale, other_lambda); - // The group's normalised weight, computed from channel 0 and reused - // by every later channel's scatter. - let mut gw = 0.0f32; + // Computed from channel 0 and reused by every later channel's scatter. + let mut scaled_weight = 0.0f32; #[unroll] for c in 0..channels { let sigma_c = sigma[c as usize]; let base_sig2 = sigma_c * sigma_c; - // Column `sub` of every member, read out of the member's own - // frame. Lane `sub` holds `stack[m * 8 + r]` for member `m`, row - // `r`. + // Lane `sub` holds `stack[m * 8 + r]` for member `m`, row `r`, read from the member's own + // frame. let mut stack = Array::::new(PATCH_AREA as usize); - let mut v = Array::::new(MAX_K as usize); + let mut member_variance = Array::::new(MAX_K as usize); #[unroll] for m in 0..MAX_K { let packed = member_pos[m as usize]; - let mx = packed & 0x1FFFu32; - let my = (packed >> 13u32) & 0x1FFFu32; + let member_x = packed & 0x1FFFu32; + let member_y = (packed >> 13u32) & 0x1FFFu32; let src_slot = member_slot[m as usize]; - v[m as usize] = base_sig2; + member_variance[m as usize] = base_sig2; #[unroll] for r in 0..PATCH_SIZE { - let px = read_line(ring, mx + sub, my + r, src_slot, width, height); - stack[(m * PATCH_SIZE + r) as usize] = px[c as usize]; + let pixel = read_line(ring, member_x + sub, member_y + r, src_slot, width, height); + stack[(m * PATCH_SIZE + r) as usize] = pixel[c as usize]; } } - // The noise variance behind each member, propagated to a - // per-stack-level variance. The spatial profile is a constant - // factor across the stack axis and the ladder only averages, so - // it multiplies in at the threshold instead of here. + // The spatial profile is constant along the stack axis and the ladder only averages, so + // the profile multiplies in at the threshold instead. if comptime!(grid_frames > 1) { if use_grid { - grid_variance(&mut v, grid_frames); + grid_variance(&mut member_variance, grid_frames); } else { - stack_variance_ladder(&mut v, k_use); + stack_variance_ladder(&mut member_variance, k_use); } } else { - stack_variance_ladder(&mut v, k_use); + stack_variance_ladder(&mut member_variance, k_use); } - // 2D DCT forward, independently for each member's patch. The - // column pass runs over the rows the lane already holds, the - // transpose hands the lane a row, and the row pass runs over - // that. Lane `sub` comes out holding coefficient `(u = i, v = - // sub)` at slot `i`. + // Lane `sub` comes out holding coefficient `(u = i, v = sub)` at slot `i`. #[unroll] for m in 0..MAX_K { let mut line = Array::::new(PATCH_SIZE as usize); @@ -616,7 +419,7 @@ pub fn collab_fused( line[i as usize] = stack[(m * PATCH_SIZE + i) as usize]; } dct8_reg_fwd(&basis, &mut line); - transpose8(&mut tbuf, &mut line, sub, grp); + transpose8(&mut transpose_buf, &mut line, sub, group); dct8_reg_fwd(&basis, &mut line); #[unroll] for i in 0..PATCH_SIZE { @@ -624,9 +427,7 @@ pub fn collab_fused( } } - // Haar transform along the stack axis, at each of the lane's - // eight spatial positions. A lane owns every member at every - // position it holds, so nothing crosses lanes here. + // A lane owns every member at each of its positions, so nothing crosses lanes here. if comptime!(grid_frames > 1) { if use_grid { grid_fwd(&mut stack, grid_frames); @@ -637,9 +438,6 @@ pub fn collab_fused( stack_haar_fwd(&mut stack, k_use); } - // Hard threshold, and the group-DC exception described above. - // The lane's retained variance is summed here and folded across - // the group below. let channel_lambda = if comptime!(c == 0u32) { luma_lambda } else { @@ -650,7 +448,7 @@ pub fn collab_fused( let threshold = channel_lambda * pool_ratio; retained_v = pooled_threshold( &mut stack, - &v, + &member_variance, dct_profile, prof_sub, sub, @@ -665,16 +463,16 @@ pub fn collab_fused( #[unroll] for j in 0..MAX_K { if j < k_use { - let vj = v[j as usize] * factor; + let coeff_variance = member_variance[j as usize] * factor; let slot = (j * PATCH_SIZE + i) as usize; - let mut keep = f32::abs(stack[slot]) >= channel_lambda * f32::sqrt(vj); + let mut keep = f32::abs(stack[slot]) >= channel_lambda * f32::sqrt(coeff_variance); if comptime!(j == 0u32 && i == 0u32) { if sub == 0u32 { keep = true; } } if keep { - retained_v += vj; + retained_v += coeff_variance; } else { stack[slot] = 0.0f32; } @@ -683,32 +481,18 @@ pub fn collab_fused( } } - // The group weight has to be known before the scatter below, and - // only the first channel computes it, so the reduction runs here - // rather than after the inverse transforms. + // The scatter needs the weight, so the reduction runs before the inverse transforms. if comptime!(c == 0u32) { let sum = plane_ssd_reduce8(retained_v); - // `sum` adds non-negative variances, so it is never - // negative. `safe_reciprocal` checks for a non-finite sum - // explicitly rather than leaning on `f32::max` to discard - // one, so the weight is finite here whatever a given GPU - // does with NaN. - let w = safe_reciprocal(sum, RECIPROCAL_FLOOR); + let weight = safe_reciprocal(sum, RECIPROCAL_FLOOR); if live && sub == 0u32 { - group_weight[ref_idx as usize] = w; + group_weight[ref_idx as usize] = weight; } - // The accumulators count in fixed point, so the weight is - // scaled into the band `weight_scale` was built to put it - // in. Aggregation normalises by the weight sum, so scaling - // every weight by the same constant leaves the result - // exactly as it would have been. - gw = w * weight_scale; + scaled_weight = weight * weight_scale; } - // Haar inverse, back from stack coefficients to per-member DCT - // coefficients, then the spatial inverse in the opposite order - // to the forward pass. The lane holds a column again by the end - // of it, which is what makes the scatter below coalesced. + // The inverse runs in the opposite order, which leaves the lane holding a column again so + // the scatter is coalesced. if comptime!(grid_frames > 1) { if use_grid { grid_inv(&mut stack, grid_frames); @@ -727,7 +511,7 @@ pub fn collab_fused( line[i as usize] = stack[(m * PATCH_SIZE + i) as usize]; } dct8_reg_inv(&basis, &mut line); - transpose8(&mut tbuf, &mut line, sub, grp); + transpose8(&mut transpose_buf, &mut line, sub, group); dct8_reg_inv(&basis, &mut line); #[unroll] for i in 0..PATCH_SIZE { @@ -735,17 +519,14 @@ pub fn collab_fused( } } - // Every member of the group is written back, not just the - // reference patch, and each lands in its own frame's region of - // the accumulators. A neighbour-frame member therefore feeds the - // caller's cross-frame ring rather than being discarded once it - // has served the group's shared statistics. + // Every member is written back into its own frame's region, so neighbour-frame members + // feed the cross-frame ring. #[unroll] for m in 0..MAX_K { if live && m < k_use { let packed = member_pos[m as usize]; - let mx = packed & 0x1FFFu32; - let my = (packed >> 13u32) & 0x1FFFu32; + let member_x = packed & 0x1FFFu32; + let member_y = (packed >> 13u32) & 0x1FFFu32; let dst_slot = member_slot[m as usize]; #[unroll] for r in 0..PATCH_SIZE { @@ -754,9 +535,9 @@ pub fn collab_fused( wsum, kaiser, stack[(m * PATCH_SIZE + r) as usize], - gw, - mx, - my, + scaled_weight, + member_x, + member_y, r * PATCH_SIZE + sub, comptime!(c == 0u32), c, @@ -800,7 +581,7 @@ fn stack_haar_fwd(stack: &mut Array, k_use: u32) { } } -/// The inverse of [stack_haar_fwd]. +/// The inverse of `stack_haar_fwd`. #[cube] fn stack_haar_inv(stack: &mut Array, k_use: u32) { if k_use >= 2u32 { diff --git a/av-denoise-core/src/collab/kernels/fused/search.rs b/av-denoise-core/src/collab/kernels/fused/search.rs index 522431b..a121292 100644 --- a/av-denoise-core/src/collab/kernels/fused/search.rs +++ b/av-denoise-core/src/collab/kernels/fused/search.rs @@ -5,40 +5,28 @@ use crate::collab::kernels::group::{clamp_top_left, pack_pos_t}; use crate::collab::kernels::plane_ops::{plane_ssd_reduce8, shift_insert8, shift_insert8_gated}; use crate::nlmeans::kernels::helpers::read_line; -/// The lowest block index whose span contains the patch at `p` on one axis. +/// The lowest block index whose span contains the patch at `patch_start` on one axis. /// -/// Block `b` spans `b * step..b * step + blksize`, so the patch -/// `p..p + PATCH_SIZE` needs `b * step + blksize >= p + PATCH_SIZE`. -/// The highest such block is `p / step`, which the caller clamps to the -/// grid and uses as the low end's ceiling. -/// -/// This mirrors `covering_blocks` in the `mc_accuracy` bench's harness -/// module (`av-denoise-core/benches/harness/score.rs`), which the tests -/// below reproduce on the host to check the two stay in step. +/// Block `b` spans `b * step..b * step + blksize`, so it contains `patch_start..patch_start + PATCH_SIZE` +/// when `b * step + blksize >= patch_start + PATCH_SIZE`. The caller clamps the result to the highest +/// covering block, `patch_start / step`. It mirrors `covering_blocks` in `nl4d/harness/score.rs`. #[cube] -pub(crate) fn covering_lo(p: u32, #[comptime] blksize: u32, #[comptime] step: u32) -> u32 { - let past = u32::max(p + PATCH_SIZE, blksize) - blksize; - past.div_ceil(step) +pub(crate) fn covering_lo(patch_start: u32, #[comptime] blksize: u32, #[comptime] step: u32) -> u32 { + let overhang = u32::max(patch_start + PATCH_SIZE, blksize) - blksize; + overhang.div_ceil(step) } -/// The host mirror of [covering_lo], for tests that cannot launch a -/// kernel. +/// Host mirror of `covering_lo`. #[cfg(test)] -fn covering_lo_host(p: u32, blksize: u32, step: u32) -> u32 { - let past = u32::max(p + PATCH_SIZE, blksize) - blksize; - past.div_ceil(step) +fn covering_lo_host(patch_start: u32, blksize: u32, step: u32) -> u32 { + let overhang = u32::max(patch_start + PATCH_SIZE, blksize) - blksize; + overhang.div_ceil(step) } -/// The distance from the reference patch to the candidate whose -/// top-left pixel is `(x, y)` in frame `slot`. +/// The distance from the reference patch to the candidate with top-left `(x, y)` in frame `slot`. /// -/// Each lane holds one column of the reference patch and reads the -/// matching column of the candidate, so the eight per-lane partials -/// only become a whole-patch distance through -/// [plane_ssd_reduce8]. That reduction shuffles, so every lane of the -/// group has to reach it. Callers that end up discarding the result -/// still call this and drop the value afterwards rather than branching -/// around it. +/// Each lane sums its own column and `plane_ssd_reduce8` completes the distance with shuffles, so +/// every lane of the group must call this, even when the result is discarded. #[cube] pub(crate) fn candidate_distance( ring: &Array>, @@ -55,11 +43,11 @@ pub(crate) fn candidate_distance( let mut partial = 0.0f32; #[unroll] for r in 0..PATCH_SIZE { - let px = read_line(ring, x + sub, y + r, slot, width, height); + let pixel = read_line(ring, x + sub, y + r, slot, width, height); #[unroll] for c in 0..channels { - let d = current[(r * channels + c) as usize] - px[c as usize]; - partial += d * d; + let diff = current[(r * channels + c) as usize] - pixel[c as usize]; + partial += diff * diff; } } plane_ssd_reduce8(partial) * scale @@ -100,34 +88,27 @@ pub(crate) fn spatial_search( let s_top = clamp_top_left(ry as i32 - spatial_radius as i32, max_y); let s_bot = clamp_top_left(ry as i32 + spatial_radius as i32, max_y); - // The reference patch scores the lowest distance there is, which on - // textured content is enough to reach slot 0 on its own. On flat - // content every candidate scores that same distance, and - // `shift_insert8` leaves a tie with whichever candidate reached the - // slot first. A sentinel below every real distance pins the - // self-match whatever ties around it. + // The self-match takes a sentinel below every real distance, so it holds slot 0 even on flat + // content where every candidate ties. if warp_uniform { - // The clipped rectangle is never wider than the unclipped one, - // so walking the unclipped span covers every position the other - // path visits, in the same order, and the rest are masked. The - // span is comptime, so every group in the warp takes the same - // number of turns. + // The unclipped span covers every position the clipped walk visits, in the same order, and + // its comptime size gives every group in the warp the same number of turns. let span = comptime!(2 * spatial_radius + 1); for dy in 0..span { for dx in 0..span { let wanted_y = s_top + dy; let wanted_x = s_left + dx; let live_pos = wanted_x <= s_right && wanted_y <= s_bot; - // A masked turn still reads, so it is pinned to the last - // live position rather than left to run off the frame. - let cx = u32::min(wanted_x, s_right); - let cy = u32::min(wanted_y, s_bot); + + // A masked turn still reads, so it is pinned to the last live position. + let candidate_x = u32::min(wanted_x, s_right); + let candidate_y = u32::min(wanted_y, s_bot); let scored = candidate_distance( ring, current, - cx, - cy, + candidate_x, + candidate_y, centre_slot, sub, scale, @@ -135,31 +116,35 @@ pub(crate) fn spatial_search( height, channels, ); - // Only the branchless part of the insert is shared. The - // gated form tests a group-local distance before it - // shuffles, which is exactly the divergence this path - // exists to avoid. + + // The gated insert branches on a group-local distance before it shuffles, which is + // the divergence this path avoids. let mut dist = select(live_pos, scored, 3.0e38f32); - // A masked turn can land on the reference's own position - // once it has been pinned, so `live_pos` has to gate the - // sentinel too, or a dead turn would plant a second - // self-match in the group. - if live_pos && cx == rx && cy == ry { + + // A masked turn pinned onto the reference would otherwise plant a second + // self-match. + if live_pos && candidate_x == rx && candidate_y == ry { dist = -1.0e38f32; } - shift_insert8(best_d, best_pos, dist, pack_pos_t(cx, cy, 0u32), sub); + shift_insert8( + best_d, + best_pos, + dist, + pack_pos_t(candidate_x, candidate_y, 0u32), + sub, + ); } } } else { - let mut cy = s_top; - while cy <= s_bot { - let mut cx = s_left; - while cx <= s_right { + let mut candidate_y = s_top; + while candidate_y <= s_bot { + let mut candidate_x = s_left; + while candidate_x <= s_right { let mut dist = candidate_distance( ring, current, - cx, - cy, + candidate_x, + candidate_y, centre_slot, sub, scale, @@ -167,13 +152,20 @@ pub(crate) fn spatial_search( height, channels, ); - if cx == rx && cy == ry { + if candidate_x == rx && candidate_y == ry { dist = -1.0e38f32; } - shift_insert8_gated(best_d, best_pos, dist, pack_pos_t(cx, cy, 0u32), sub, base); - cx += 1u32; + shift_insert8_gated( + best_d, + best_pos, + dist, + pack_pos_t(candidate_x, candidate_y, 0u32), + sub, + base, + ); + candidate_x += 1u32; } - cy += 1u32; + candidate_y += 1u32; } } @@ -269,20 +261,20 @@ pub(crate) fn trajectory_search( let block_live = wanted_bx <= bx_hi && wanted_by <= by_hi; if warp_uniform { - let cbx = u32::min(wanted_bx, bx_hi); - let cby = u32::min(wanted_by, by_hi); - let block = cby * blocks_x + cbx; - let conf = confidence[(t * conf_stride + block) as usize]; - let block_scored = block_live && conf >= c_min; + let clamped_bx = u32::min(wanted_bx, bx_hi); + let clamped_by = u32::min(wanted_by, by_hi); + let block = clamped_by * blocks_x + clamped_bx; + let block_confidence = confidence[(t * conf_stride + block) as usize]; + let block_scored = block_live && block_confidence >= c_min; - let mv = (t * mv_stride + block * 2u32) as usize; - let px0 = anchor_x as i32 + mv_field[mv]; - let py0 = anchor_y as i32 + mv_field[mv + 1]; + let mv_index = (t * mv_stride + block * 2u32) as usize; + let predicted_x = anchor_x as i32 + mv_field[mv_index]; + let predicted_y = anchor_y as i32 + mv_field[mv_index + 1]; - let t_left = clamp_top_left(px0 - refine as i32, max_x); - let t_right = clamp_top_left(px0 + refine as i32, max_x); - let t_top = clamp_top_left(py0 - refine as i32, max_y); - let t_bot = clamp_top_left(py0 + refine as i32, max_y); + let t_left = clamp_top_left(predicted_x - refine as i32, max_x); + let t_right = clamp_top_left(predicted_x + refine as i32, max_x); + let t_top = clamp_top_left(predicted_y - refine as i32, max_y); + let t_bot = clamp_top_left(predicted_y + refine as i32, max_y); let span = comptime!(2 * refine + 1); for dy in 0..span { @@ -290,21 +282,22 @@ pub(crate) fn trajectory_search( let wanted_y = t_top + dy; let wanted_x = t_left + dx; let in_rect = wanted_x <= t_right && wanted_y <= t_bot; - let nx = u32::min(wanted_x, t_right); - let ny = u32::min(wanted_y, t_bot); - let packed = pack_pos_t(nx, ny, packed_t); + let candidate_x = u32::min(wanted_x, t_right); + let candidate_y = u32::min(wanted_y, t_bot); + let packed = pack_pos_t(candidate_x, candidate_y, packed_t); let mut skipped = false; #[unroll] for s in 0..max_rects { - if nx >= seen_left[s as usize] - && nx <= seen_right[s as usize] - && ny >= seen_top[s as usize] - && ny <= seen_bot[s as usize] + if candidate_x >= seen_left[s as usize] + && candidate_x <= seen_right[s as usize] + && candidate_y >= seen_top[s as usize] + && candidate_y <= seen_bot[s as usize] { skipped = true; } } + #[unroll] for k in 0..first { if member_pos[k as usize] == packed { @@ -314,7 +307,16 @@ pub(crate) fn trajectory_search( let live_pos = block_scored && in_rect && !skipped; let scored = candidate_distance( - ring, anchor, nx, ny, slot, sub, scale, width, height, channels, + ring, + anchor, + candidate_x, + candidate_y, + slot, + sub, + scale, + width, + height, + channels, ); let dist = select(live_pos, scored, 3.0e38f32); let better = dist < frame_d; @@ -330,34 +332,35 @@ pub(crate) fn trajectory_search( seen_bot[rect] = select(block_scored, t_bot, 0u32); } else if block_live { let block = wanted_by * blocks_x + wanted_bx; - let conf = confidence[(t * conf_stride + block) as usize]; - if conf >= c_min { - let mv = (t * mv_stride + block * 2u32) as usize; - let px0 = anchor_x as i32 + mv_field[mv]; - let py0 = anchor_y as i32 + mv_field[mv + 1]; - - let t_left = clamp_top_left(px0 - refine as i32, max_x); - let t_right = clamp_top_left(px0 + refine as i32, max_x); - let t_top = clamp_top_left(py0 - refine as i32, max_y); - let t_bot = clamp_top_left(py0 + refine as i32, max_y); - - let mut ny = t_top; - while ny <= t_bot { - let mut nx = t_left; - while nx <= t_right { - let packed = pack_pos_t(nx, ny, packed_t); + let block_confidence = confidence[(t * conf_stride + block) as usize]; + if block_confidence >= c_min { + let mv_index = (t * mv_stride + block * 2u32) as usize; + let predicted_x = anchor_x as i32 + mv_field[mv_index]; + let predicted_y = anchor_y as i32 + mv_field[mv_index + 1]; + + let t_left = clamp_top_left(predicted_x - refine as i32, max_x); + let t_right = clamp_top_left(predicted_x + refine as i32, max_x); + let t_top = clamp_top_left(predicted_y - refine as i32, max_y); + let t_bot = clamp_top_left(predicted_y + refine as i32, max_y); + + let mut candidate_y = t_top; + while candidate_y <= t_bot { + let mut candidate_x = t_left; + while candidate_x <= t_right { + let packed = pack_pos_t(candidate_x, candidate_y, packed_t); let mut skipped = false; #[unroll] for s in 0..max_rects { - if nx >= seen_left[s as usize] - && nx <= seen_right[s as usize] - && ny >= seen_top[s as usize] - && ny <= seen_bot[s as usize] + if candidate_x >= seen_left[s as usize] + && candidate_x <= seen_right[s as usize] + && candidate_y >= seen_top[s as usize] + && candidate_y <= seen_bot[s as usize] { skipped = true; } } + #[unroll] for k in 0..first { if member_pos[k as usize] == packed { @@ -367,7 +370,16 @@ pub(crate) fn trajectory_search( if !skipped { let dist = candidate_distance( - ring, anchor, nx, ny, slot, sub, scale, width, height, channels, + ring, + anchor, + candidate_x, + candidate_y, + slot, + sub, + scale, + width, + height, + channels, ); if dist < frame_d { frame_d = dist; @@ -375,9 +387,9 @@ pub(crate) fn trajectory_search( } } - nx += 1u32; + candidate_x += 1u32; } - ny += 1u32; + candidate_y += 1u32; } let rect = (iy * covers + ix) as usize; @@ -412,39 +424,19 @@ pub(crate) fn trajectory_search( #[cfg(test)] mod tests { use super::covering_lo_host; - - /// The host mirror of `covering_blocks` in the `mc_accuracy` bench's - /// harness (`benches/harness/score.rs`). Reproduced here, rather than - /// imported, because that module lives outside the crate as bench-only - /// code and cannot be a test dependency of the library. - /// - /// This pins the kernel's arithmetic against the harness's read of the - /// same geometry rather than launching a real kernel, so it catches the - /// two formulas drifting apart on paper but says nothing about whether - /// [super::covering_lo] compiles or runs correctly on a GPU; the - /// integration tests in `nl4d::tests` cover that by driving the whole - /// pipeline. - fn covering_blocks_host(p: u32, blksize: u32, step: u32, blocks: u32) -> (u32, u32) { - let hi = (p / step).min(blocks - 1); - let lo = if p + super::PATCH_SIZE <= blksize { - 0 - } else { - (p + super::PATCH_SIZE - blksize).div_ceil(step) - }; - (lo.min(hi), hi) - } + use crate::nl4d::harness::covering_blocks; #[test] fn covering_lo_matches_the_harness_across_a_range_of_geometries() { for (blksize, overlap) in [(16u32, 8u32), (16, 12), (32, 24), (8, 4), (16, 0)] { let step = blksize - overlap; let blocks = 8u32; - for p in (0..blocks * step).step_by(3) { - let (expect_lo, hi) = covering_blocks_host(p, blksize, step, blocks); - let got_lo = covering_lo_host(p, blksize, step).min(hi); + for patch_start in (0..blocks * step).step_by(3) { + let (expected_first, last) = covering_blocks(patch_start, blksize, step, blocks); + let got_first = covering_lo_host(patch_start, blksize, step).min(last); assert_eq!( - got_lo, expect_lo, - "blksize={blksize} step={step} p={p}: covering_lo disagrees with the \ + got_first, expected_first, + "blksize={blksize} step={step} p={patch_start}: covering_lo disagrees with the \ harness's covering_blocks" ); } diff --git a/av-denoise-core/src/collab/kernels/group.rs b/av-denoise-core/src/collab/kernels/group.rs index d28453f..4b57ff6 100644 --- a/av-denoise-core/src/collab/kernels/group.rs +++ b/av-denoise-core/src/collab/kernels/group.rs @@ -1,80 +1,62 @@ use cubecl::prelude::*; -/// Packs a patch's top-left position into one `u32`, x in the low 13 -/// bits and y in the next 13. +/// Packs a patch's top-left position into one `u32`, x in the low 13 bits and y in the next 13. /// -/// Both axes are 13 bits because a wider x would leave y too narrow to -/// cover a 4K-tall frame, and a position that overflows its field -/// corrupts the one beside it silently. -/// -/// Packing both coordinates into one word rather than carrying x and y -/// as a pair is what makes the top-K arrays and the duplicate check -/// single comparisons instead of two. +/// Both axes are 13 bits so y covers a 4K-tall frame, because a position that overflows its field +/// silently corrupts the other. One word makes the top-K arrays and the duplicate check single +/// comparisons. #[cube] pub(crate) fn pack_pos(x: u32, y: u32) -> u32 { (y << 13) | x } -/// The host-side mirror of [`pack_pos`], for building expected values in -/// tests without a GPU round trip. +/// Host mirror of `pack_pos`, for building expected values in tests. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] pub(crate) fn pack_pos_host(x: u32, y: u32) -> u32 { (y << 13) | x } -/// The host-side mirror of unpacking a position [`pack_pos`] produced. +/// Host mirror of unpacking a `pack_pos` word. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] pub(crate) fn unpack_pos_host(packed: u32) -> (u32, u32) { (packed & 0x1FFF, (packed >> 13) & 0x1FFF) } -/// Packs a candidate position and the neighbour it came from into one -/// word. -/// -/// The axes are 13 bits each rather than 16 and 11 because a 16-bit `x` -/// would leave `y` too narrow. `y` at 11 bits stops at 2047, which is -/// below a 2160-line frame, and a position that overflows its field -/// corrupts the one beside it silently. `x` takes bits 0-12 and `y` bits -/// 13-25, which leaves bits 26-31 for `t`, room for 63 neighbours -/// against the 16 a temporal radius of 8 asks for. -/// -/// `t` is 0 for a centre-frame position and `neighbour_index + 1` -/// otherwise, so a member's frame and its motion-block confidence are -/// both recoverable from the packed word alone. Keeping them here -/// rather than in a second array retires one register per slot in the -/// matching loop and one array in the filter stage. +/// Packs a candidate position and the neighbour it came from into one word. /// -/// The coordinates sit exactly where [`pack_pos`] puts them, so -/// [`unpack_pos_host`] reads a packed-with-`t` word correctly. +/// x takes bits 0-12, y bits 13-25 and `t` bits 26-31, which holds up to 63 neighbours. Thirteen +/// bits per axis keeps y above a 2160-line frame, because a position that overflows its field +/// silently corrupts the one beside it. `t` is 0 for a centre-frame position and +/// `neighbour_index + 1` otherwise, so a member's frame and motion-block confidence are recoverable +/// from the word alone. The coordinates sit where `pack_pos` puts them, so `unpack_pos_host` reads +/// this word too. #[cube] pub(crate) fn pack_pos_t(x: u32, y: u32, t: u32) -> u32 { (t << 26u32) | (y << 13u32) | x } -/// The neighbour field [`pack_pos_t`] wrote. +/// The neighbour field `pack_pos_t` wrote. #[cube] pub(crate) fn unpack_t(packed: u32) -> u32 { packed >> 26u32 } -/// The host-side mirror of [`pack_pos_t`]. +/// Host mirror of `pack_pos_t`. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] pub(crate) fn pack_pos_t_host(x: u32, y: u32, t: u32) -> u32 { (t << 26) | (y << 13) | x } -/// The host-side mirror of [`unpack_t`]. +/// Host mirror of `unpack_t`. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] pub(crate) fn unpack_t_host(packed: u32) -> u32 { packed >> 26 } -/// Clamps a candidate top-left coordinate to `[0, max_pos]`. +/// Clamps a candidate top-left coordinate to `0..=max_pos`. /// -/// Every candidate patch position on every axis goes through this before -/// anything reads from it, so a patch read at that position always starts -/// and ends inside the frame. No kernel that later consumes a group has -/// to clamp its own reads. +/// Every candidate position goes through this before anything reads it, so a patch read always +/// stays inside the frame and no later kernel has to clamp its own reads. #[cube] pub(crate) fn clamp_top_left(v: i32, max_pos: u32) -> u32 { let mut result = v; diff --git a/av-denoise-core/src/collab/kernels/plane_ops.rs b/av-denoise-core/src/collab/kernels/plane_ops.rs index a7b24f0..d2c8243 100644 --- a/av-denoise-core/src/collab/kernels/plane_ops.rs +++ b/av-denoise-core/src/collab/kernels/plane_ops.rs @@ -1,67 +1,46 @@ use cubecl::prelude::*; -/// The plane-local index of the first lane in the calling lane's -/// 8-lane group. +/// The plane-local index of the first lane in the calling lane's 8-lane group. /// -/// Every shuffle in this module addresses lanes relative to this, so a -/// group never reads a lane belonging to another group. Cubes are -/// launched 1-D, which makes the thread-to-plane mapping linear, so -/// `UNIT_POS_PLANE % 8` and `UNIT_POS_X % 8` agree at both wave32 and -/// wave64. +/// Every shuffle here addresses lanes relative to this, so a group never reads another group's +/// lane. Cubes launch 1-D, so `UNIT_POS_PLANE % 8` and `UNIT_POS_X % 8` agree at wave32 and wave64. #[cube] pub(crate) fn group_base() -> u32 { UNIT_POS_PLANE - UNIT_POS_PLANE % 8u32 } -/// The sum of `partial` across the calling lane's 8-lane group, -/// returned to every lane in it. +/// Sums `partial` across the calling lane's 8-lane group and returns the sum to every lane. /// -/// Three XOR shuffles fold 8 values into 8 copies of their sum. The -/// masks are 1, 2 and 4, all below 8, so a lane index XORed with any of -/// them stays inside its own group whatever the plane's width. -/// -/// Each lane holds the squared differences for one patch column, so -/// this is what completes a candidate's distance. +/// The XOR masks 1, 2 and 4 are all below 8, so the shuffles stay inside the group whatever the +/// plane width. #[cube] pub(crate) fn plane_ssd_reduce8(partial: f32) -> f32 { - let mut acc = partial; - acc += plane_shuffle_xor(acc, 1u32); - acc += plane_shuffle_xor(acc, 2u32); - acc += plane_shuffle_xor(acc, 4u32); - acc + let mut sum = partial; + sum += plane_shuffle_xor(sum, 1u32); + sum += plane_shuffle_xor(sum, 2u32); + sum += plane_shuffle_xor(sum, 4u32); + sum } /// Inserts one candidate into the group's sorted top-8. /// -/// The eight best candidates seen so far live one per lane, ascending, -/// slot 0 in the group's first lane. A new candidate shifts every slot -/// it beats one lane along and drops out the eighth. Nothing is stored -/// in shared memory and no second pass is needed. -/// -/// A candidate that ties an incumbent does not displace it, so the -/// first candidate seen at a given distance keeps its slot. Exactly -/// equal distances are common on flat content, so this is what fixes -/// which member a group keeps rather than leaving it to scheduling. +/// The eight best candidates live one per lane, ascending, with slot 0 in the group's first lane. +/// A new candidate shifts every slot it beats one lane along and drops the eighth. A tie never +/// displaces an incumbent, so on flat content the first candidate seen keeps its slot rather than +/// leaving it to scheduling. /// -/// `plane_shuffle_up` at the group's first lane returns a value from -/// the previous group. The `sub == 0` term discards it before it can be -/// used. -/// -/// Exactly two shuffles run, one per carried value. Whether the previous -/// lane also beat `d` is not shuffled, because it follows from `prev_d`. -/// The slots ascend, so the previous lane holds `prev_d` and beat `d` -/// exactly when `d < prev_d`. Each shuffle is an LDS crossbar operation -/// on every candidate, so deriving this rather than shuffling a flag is -/// worth the line of algebra. +/// `plane_shuffle_up` at the group's first lane reads the previous group, so the `sub == 0` term +/// discards it. The slots ascend, so the previous lane beat `distance` exactly when +/// `distance < prev_d`, which saves shuffling a flag. #[cube] -pub(crate) fn shift_insert8(best_d: &mut f32, best_pos: &mut u32, d: f32, packed: u32, sub: u32) { +pub(crate) fn shift_insert8(best_d: &mut f32, best_pos: &mut u32, distance: f32, packed: u32, sub: u32) { let prev_d = plane_shuffle_up(*best_d, 1u32); let prev_pos = plane_shuffle_up(*best_pos, 1u32); - if d < *best_d { - let first = sub == 0u32 || d >= prev_d; - if first { - *best_d = d; + if distance < *best_d { + let lands_here = sub == 0u32 || distance >= prev_d; + if lands_here { + *best_d = distance; *best_pos = packed; } else { *best_d = prev_d; @@ -70,59 +49,46 @@ pub(crate) fn shift_insert8(best_d: &mut f32, best_pos: &mut u32, d: f32, packed } } -/// [`shift_insert8`] with the shuffles skipped when the candidate -/// cannot place. -/// -/// The group's eighth-best distance sits in its last lane. A candidate -/// that does not beat it changes nothing, so the two shuffles the insert -/// costs are skipped. Every lane holds the same `d` and reads the same -/// broadcast, so the branch is uniform across the group and no lane sits -/// out a shuffle another lane takes. +/// `shift_insert8` with the shuffles skipped when the candidate cannot place. /// -/// The broadcast costs nothing. It fuses into the compare that sets the -/// execution mask, one `v_cmpx_gt_f32` with a `dpp8` modifier, while -/// each shuffle it skips is an LDS crossbar operation. That is why this -/// is the form the kernel uses rather than a later optimisation. -/// -/// This is a compute saving and never an admission decision. The eight -/// slots it produces are the same ones [`shift_insert8`] produces. +/// The group's eighth-best distance sits in its last lane, and a candidate that does not beat it +/// changes nothing. Every lane holds the same `distance` and reads the same broadcast, so the branch is +/// uniform across the group and no lane skips a shuffle another lane takes. The broadcast fuses +/// into the compare, while each skipped shuffle is an LDS crossbar operation. It keeps the same +/// eight slots as `shift_insert8`. #[cube] pub(crate) fn shift_insert8_gated( best_d: &mut f32, best_pos: &mut u32, - d: f32, + distance: f32, packed: u32, sub: u32, base: u32, ) { let worst = plane_shuffle(*best_d, base + 7u32); - if d < worst { - shift_insert8(best_d, best_pos, d, packed, sub); + if distance < worst { + shift_insert8(best_d, best_pos, distance, packed, sub); } } -/// Transposes one 8x8 block held one column per lane into one row per -/// lane, through `buf`. -/// -/// `v` holds the calling lane's 8 values on entry and its transposed 8 -/// on return. `slot` is the calling lane's group index, which picks the -/// group's own 65-float region of `buf`, and the stride is padded by -/// one past 64 so that eight lanes writing eight consecutive rows never -/// collide on a bank. +/// Transposes an 8x8 block held one column per lane into one row per lane, through `buf`. /// -/// The spatial row pass needs a row and a lane owns a column, so this -/// runs once before it and once after its inverse. +/// `v` holds the lane's 8 values on entry and its transposed 8 on return. `slot` picks the group's +/// own 65-float region of `buf`, padded one past 64 so eight lanes writing consecutive rows never +/// collide on a bank. It contains two `sync_cube()` barriers, so every lane of the cube must call +/// it. #[cube] pub(crate) fn transpose8(buf: &mut SharedMemory, v: &mut Array, sub: u32, slot: u32) { - let base = slot * 65u32; + let region_start = slot * 65u32; #[unroll] for i in 0..8u32 { - buf[(base + i * 8u32 + sub) as usize] = v[i as usize]; + buf[(region_start + i * 8u32 + sub) as usize] = v[i as usize]; } sync_cube(); + #[unroll] for i in 0..8u32 { - v[i as usize] = buf[(base + sub * 8u32 + i) as usize]; + v[i as usize] = buf[(region_start + sub * 8u32 + i) as usize]; } sync_cube(); } diff --git a/av-denoise-core/src/collab/kernels/transforms.rs b/av-denoise-core/src/collab/kernels/transforms.rs index 6079fa9..51284ed 100644 --- a/av-denoise-core/src/collab/kernels/transforms.rs +++ b/av-denoise-core/src/collab/kernels/transforms.rs @@ -2,154 +2,107 @@ use cubecl::prelude::*; use crate::collab::{MAX_K, PATCH_AREA, PATCH_SIZE}; -// The register-layout helpers below stride a group by `PATCH_SIZE`, -// because one lane holds one 8-value column of each of `MAX_K` members -// and the two counts happen to be the same number. A caller sizes the -// array it passes them as `PATCH_AREA`, which is only enough for the -// whole group while that holds. +// A lane holds one 8-value column of each of `MAX_K` members, so the Haar helpers stride a group +// by `PATCH_SIZE` and a `PATCH_AREA` array only holds the whole group while the two match. const _: () = assert!( MAX_K == PATCH_SIZE, "haar_reg_fwd_level and haar_reg_inv_level stride a group by PATCH_SIZE, which only holds \ the whole stack while MAX_K matches it" ); -/// The floor the filter passes to [`safe_reciprocal`] when it turns a -/// retained variance sum into a group weight. +/// The floor the filter passes to `safe_reciprocal` when turning a retained variance sum into a +/// group weight. /// -/// A group weight therefore never exceeds `1 / RECIPROCAL_FLOOR`, which -/// is the bound [`crate::collab::kernels::aggregate::weight_scale`] -/// needs to fit the weights into a fixed-point accumulator. +/// It caps a group weight at `1 / RECIPROCAL_FLOOR`, the bound +/// [weight_scale](crate::collab::kernels::aggregate::weight_scale) needs to fit the weights into a +/// fixed-point accumulator. pub const RECIPROCAL_FLOOR: f32 = 1e-12; -/// A reciprocal weight that never trusts a driver's `NaN`-`max` -/// behaviour. +/// Returns `1 / max(denom, floor)`, or 0 when `denom` is NaN or infinite. /// -/// Returns `1 / max(denom, floor)` for a finite denominator. A `denom` -/// that is `NaN` or infinite is caught first and returns `0`. -/// -/// The result only ever feeds a weight or a normalising divisor, never a -/// value that could stand in for real content, so zero is the right -/// fallback for both cases. Either means the reciprocal cannot be trusted -/// to mean anything. -/// -/// The explicit check matters because `f32::max(NaN, floor)` is not -/// guaranteed to discard the `NaN` on a GPU the way it does on the CPU. -/// SPIR-V's `FMax` leaves a `NaN` operand undefined and only `NMax` -/// promises to discard it, so leaning on `max` alone would make -/// correctness depend on which one a driver lowers `f32::max` to. +/// Zero is a safe fallback because the result only ever feeds a weight or a divisor. The explicit +/// check is needed because SPIR-V's `FMax` leaves a NaN operand undefined, so `f32::max` alone may +/// not discard it on a GPU. #[cube] pub(crate) fn safe_reciprocal(denom: f32, floor: f32) -> f32 { - let mut inv = 0.0f32; + let mut reciprocal = 0.0f32; if !denom.is_nan() && !denom.is_inf() { - inv = 1.0f32 / f32::max(denom, floor); + reciprocal = 1.0f32 / f32::max(denom, floor); } - inv + reciprocal } -/// Fills an 8x8 orthonormal DCT-II basis into shared memory, one entry -/// per thread. -/// -/// Entry `j*8+i` holds `c_j * cos(PI * (2i+1) * j / 16)`, with `c_0 = -/// 1/sqrt(8)` and `c_j = 0.5` for every other row. That scaling is what -/// makes the basis orthonormal, each row has unit norm and any two -/// distinct rows are orthogonal, so the same matrix run forwards is a -/// DCT and run transposed is its own inverse. +/// Fills an 8x8 orthonormal DCT-II basis into shared memory, one entry per thread. /// -/// Only the first 64 threads of the calling cube write an entry, callers -/// with more threads than that call this once per thread anyway. The -/// caller must call `sync_cube()` before reading `basis`. +/// Entry `j * 8 + i` holds `c_j * cos(PI * (2i + 1) * j / 16)`, with `c_0 = 1/sqrt(8)` and +/// `c_j = 0.5` otherwise. The basis is orthonormal, so its transpose is its inverse. Only the first +/// 64 threads write, and the caller must `sync_cube()` before reading `basis`. #[cube] pub(crate) fn fill_dct8_basis(basis: &mut SharedMemory, thread_id: u32) { if thread_id < PATCH_AREA { let i = thread_id % PATCH_SIZE; let j = thread_id / PATCH_SIZE; - let mut c = 0.5f32; + let mut row_scale = 0.5f32; if j == 0 { - c = 1.0f32 / f32::sqrt(8.0f32); + row_scale = 1.0f32 / f32::sqrt(8.0f32); } let angle = std::f32::consts::PI * (2.0f32 * i as f32 + 1.0f32) * j as f32 / 16.0f32; - basis[thread_id as usize] = c * f32::cos(angle); + basis[thread_id as usize] = row_scale * f32::cos(angle); } } -/// Runs one forward 8-point DCT over a line a lane already holds in -/// registers. +/// Runs a forward 8-point DCT over a line one lane holds in registers. /// -/// `line` holds the 8 values on entry and the 8 coefficients on return. -/// The basis stays in shared memory, because every lane reads all 64 of -/// its entries and a per-lane copy would cost 64 registers to save -/// nothing. -/// -/// The whole line is snapshotted before any output is written, so -/// writing an output cannot disturb an input a later output still needs. -/// A lane owns every value it touches, so no barrier is needed here or -/// around the call. +/// `line` holds 8 values on entry and 8 coefficients on return. The basis stays in shared memory +/// because a per-lane copy would cost 64 registers. The lane owns every value it touches, so no +/// barrier is needed. #[cube] pub(crate) fn dct8_reg_fwd(basis: &SharedMemory, line: &mut Array) { - let mut src = Array::::new(8usize); + let mut snapshot = Array::::new(8usize); #[unroll] for i in 0..PATCH_SIZE { - src[i as usize] = line[i as usize]; + snapshot[i as usize] = line[i as usize]; } + #[unroll] for j in 0..PATCH_SIZE { let mut sum = 0.0f32; #[unroll] for i in 0..PATCH_SIZE { - sum += basis[(j * PATCH_SIZE + i) as usize] * src[i as usize]; + sum += basis[(j * PATCH_SIZE + i) as usize] * snapshot[i as usize]; } line[j as usize] = sum; } } -/// The inverse of [`dct8_reg_fwd`], the transpose of the same basis. -/// -/// Because the basis is orthonormal its inverse is its own transpose, so -/// this reads `basis[j*8+i]` for output position `i` instead of output -/// position `j`, and sums over `j` instead of `i`, matching -/// [`dct8_reg_fwd`]'s calling convention otherwise. +/// The inverse of `dct8_reg_fwd`, using the transpose of the same basis. #[cube] pub(crate) fn dct8_reg_inv(basis: &SharedMemory, line: &mut Array) { - let mut src = Array::::new(8usize); + let mut snapshot = Array::::new(8usize); #[unroll] for j in 0..PATCH_SIZE { - src[j as usize] = line[j as usize]; + snapshot[j as usize] = line[j as usize]; } + #[unroll] for i in 0..PATCH_SIZE { let mut sum = 0.0f32; #[unroll] for j in 0..PATCH_SIZE { - sum += basis[(j * PATCH_SIZE + i) as usize] * src[j as usize]; + sum += basis[(j * PATCH_SIZE + i) as usize] * snapshot[j as usize]; } line[i as usize] = sum; } } -/// One level of the forward stack Haar, over a group a lane holds -/// entirely in registers. -/// -/// `stack` holds `MAX_K` members of 8 values each, member `k` at -/// `k * PATCH_SIZE + pos`, which is the slice of a group one lane owns -/// when each lane holds one column. This runs the butterfly at every one -/// of those 8 positions at once, so one call covers the whole level. -/// -/// Each butterfly is `(a, b) -> ((a+b)/sqrt(2), (a-b)/sqrt(2))`, which -/// has unit-norm rows and is its own inverse. A full multi-level -/// decomposition calls this once per level, over a length that halves -/// each time, with the coarsest approximation ending up at `k = 0` and -/// finer detail bands filling the rest in order. -/// -/// This takes the level's length as a `#[comptime]` argument, so every -/// index into `stack` is a compile-time constant. A register array -/// indexed by a runtime value is not a register array, it is scratch -/// memory, and the whole point of holding the group in registers is -/// lost the moment one dynamic index appears. The caller picks the -/// levels with a run of predicates on the group size rather than a -/// loop. +/// One level of the forward stack Haar over a group one lane holds in registers. /// -/// The level snapshots every value it reads before writing any output, -/// because the output range overlaps the input range. +/// `stack` holds `MAX_K` members of 8 values, member `k` at `k * PATCH_SIZE + pos`, and the +/// butterfly `(a, b) -> ((a + b) / sqrt(2), (a - b) / sqrt(2))` runs at all 8 positions. The +/// butterfly is orthonormal and its own inverse. A full +/// decomposition calls this once per level with a halving `len`, leaving the coarsest +/// approximation at `k = 0`. `len` is comptime so every index into `stack` is a constant, because +/// one runtime index turns the register array into scratch memory. #[cube] pub(crate) fn haar_reg_fwd_level(stack: &mut Array, #[comptime] len: u32) { let half = comptime!(len / 2); @@ -160,23 +113,21 @@ pub(crate) fn haar_reg_fwd_level(stack: &mut Array, #[comptime] len: u32) { for k in 0..len { snapshot[k as usize] = stack[(k * PATCH_SIZE + pos) as usize]; } + #[unroll] for p in 0..half { - let a = snapshot[(2u32 * p) as usize]; - let b = snapshot[(2u32 * p + 1u32) as usize]; - stack[(p * PATCH_SIZE + pos) as usize] = (a + b) * std::f32::consts::FRAC_1_SQRT_2; - stack[((half + p) * PATCH_SIZE + pos) as usize] = (a - b) * std::f32::consts::FRAC_1_SQRT_2; + let first = snapshot[(2u32 * p) as usize]; + let second = snapshot[(2u32 * p + 1u32) as usize]; + stack[(p * PATCH_SIZE + pos) as usize] = (first + second) * std::f32::consts::FRAC_1_SQRT_2; + stack[((half + p) * PATCH_SIZE + pos) as usize] = + (first - second) * std::f32::consts::FRAC_1_SQRT_2; } } } -/// One level of the inverse stack Haar, the mirror of -/// [`haar_reg_fwd_level`]. +/// One level of the inverse stack Haar, the mirror of `haar_reg_fwd_level`. /// -/// The butterfly is its own inverse, so this level combines the -/// approximation half with the detail half and writes the pair back -/// interleaved. A caller runs the levels in the opposite order to the -/// forward pass. +/// Levels run in the opposite order to the forward pass. #[cube] pub(crate) fn haar_reg_inv_level(stack: &mut Array, #[comptime] len: u32) { let half = comptime!(len / 2); @@ -187,29 +138,22 @@ pub(crate) fn haar_reg_inv_level(stack: &mut Array, #[comptime] len: u32) { for k in 0..len { snapshot[k as usize] = stack[(k * PATCH_SIZE + pos) as usize]; } + #[unroll] for p in 0..half { - let a = snapshot[p as usize]; - let b = snapshot[(half + p) as usize]; - stack[(2u32 * p * PATCH_SIZE + pos) as usize] = (a + b) * std::f32::consts::FRAC_1_SQRT_2; + let low = snapshot[p as usize]; + let high = snapshot[(half + p) as usize]; + stack[(2u32 * p * PATCH_SIZE + pos) as usize] = (low + high) * std::f32::consts::FRAC_1_SQRT_2; stack[((2u32 * p + 1u32) * PATCH_SIZE + pos) as usize] = - (a - b) * std::f32::consts::FRAC_1_SQRT_2; + (low - high) * std::f32::consts::FRAC_1_SQRT_2; } } } -/// One level of the variance propagation that shadows -/// [`haar_reg_fwd_level`]. -/// -/// A Haar butterfly's two outputs are each a sum of two independent -/// inputs scaled by `1/sqrt(2)`, so both land on the same variance, -/// `(va + vb) / 2`. Running this over the same levels, in the same -/// pairing order, as the signal turns `v[j]` into the variance of stack -/// coefficient `j`. +/// One level of the variance propagation that shadows `haar_reg_fwd_level`. /// -/// [`haar_variance_ladder`] is the host mirror this is checked against. -/// This takes the level length as a `#[comptime]` argument for the -/// reason [`haar_reg_fwd_level`] gives. +/// Both outputs of a butterfly carry `(va + vb) / 2`, so running this over the same levels turns +/// `v[j]` into the variance of stack coefficient `j`. #[cube] pub(crate) fn variance_reg_level(v: &mut Array, #[comptime] len: u32) { let half = comptime!(len / 2); @@ -218,37 +162,22 @@ pub(crate) fn variance_reg_level(v: &mut Array, #[comptime] len: u32) { for k in 0..len { snapshot[k as usize] = v[k as usize]; } + #[unroll] for p in 0..half { - let avg = (snapshot[(2u32 * p) as usize] + snapshot[(2u32 * p + 1u32) as usize]) * 0.5f32; - v[p as usize] = avg; - v[(half + p) as usize] = avg; + let mean = (snapshot[(2u32 * p) as usize] + snapshot[(2u32 * p + 1u32) as usize]) * 0.5f32; + v[p as usize] = mean; + v[(half + p) as usize] = mean; } } -/// The per-DCT-frequency multiplier a spatially correlated residual -/// leaves on top of an otherwise flat noise variance. -/// -/// Non-local means leaves a residual whose covariance falls off with -/// distance as `rho^d`, rather than the flat covariance a white-noise -/// model assumes. Projecting that through the same orthonormal 8-point -/// DCT basis [`fill_dct8_basis`] builds spreads a flat variance unevenly -/// across frequencies. Low frequencies pick up more of the noise power, -/// high frequencies less. +/// The per-DCT-frequency variance multiplier for a residual whose covariance falls off as `rho^d`. /// -/// `g(u) = sum_i sum_j B_u(i) * B_u(j) * rho^|i-j|`, where `B_u` is row -/// `u` of that basis. A patch's coefficient at `(u, v)` scales by -/// `g(u) * g(v)`, since the correlation model treats rows and columns -/// alike and the 2D DCT runs as two separable 1D passes. -/// -/// `sum_u g(u) = 8` for every `rho`, because the basis is orthonormal, so -/// this redistributes variance across frequencies rather than changing -/// the total and needs no separate normalisation. -/// -/// `rho <= 0` returns `[1.0; 8]` directly. The sum below is exactly 1.0 -/// there in exact arithmetic, but eight squared cosine terms in floating -/// point land a few bits off, and a caller with correlation shaping off -/// needs its result to be bit for bit what no shaping at all would give. +/// `g(u) = sum_i sum_j B_u(i) * B_u(j) * rho^|i-j|`, where `B_u` is row `u` of the basis +/// `fill_dct8_basis` builds, and a coefficient at `(u, v)` scales by `g(u) * g(v)`. The basis is +/// orthonormal, so `sum_u g(u) = 8` and the total variance is unchanged. `rho <= 0` returns exactly +/// `[1.0; 8]`, because the float sum lands a few bits off and an unshaped caller needs a bit-exact +/// result. pub fn dct_noise_profile(rho: f32) -> [f32; 8] { if rho <= 0.0 { return [1.0; 8]; @@ -257,59 +186,45 @@ pub fn dct_noise_profile(rho: f32) -> [f32; 8] { let rho = rho as f64; let mut basis = [[0.0f64; 8]; 8]; for (u, row) in basis.iter_mut().enumerate() { - let c = if u == 0 { 1.0 / 8.0f64.sqrt() } else { 0.5 }; + let row_scale = if u == 0 { 1.0 / 8.0f64.sqrt() } else { 0.5 }; for (i, entry) in row.iter_mut().enumerate() { let angle = std::f64::consts::PI * (2.0 * i as f64 + 1.0) * u as f64 / 16.0; - *entry = c * angle.cos(); + *entry = row_scale * angle.cos(); } } - let mut g = [0.0f32; 8]; - for (u, slot) in g.iter_mut().enumerate() { + let mut profile = [0.0f32; 8]; + for (u, slot) in profile.iter_mut().enumerate() { let mut sum = 0.0f64; - for (i, &bi) in basis[u].iter().enumerate() { - for (j, &bj) in basis[u].iter().enumerate() { - sum += bi * bj * rho.powi((i as i32 - j as i32).abs()); + for (i, &basis_i) in basis[u].iter().enumerate() { + for (j, &basis_j) in basis[u].iter().enumerate() { + let distance = (i as i32 - j as i32).abs(); + sum += basis_i * basis_j * rho.powi(distance); } } *slot = sum as f32; } - g + profile } -/// The host-side mirror of the variance propagation the stack Haar -/// applies to a per-coefficient noise variance instead of a signal. -/// -/// A Haar butterfly's two outputs are each a sum of two independent -/// values scaled by `1/sqrt(2)`, so if `va` and `vb` are the variances -/// of `a` and `b`, both outputs land on the same variance, `(va + -/// vb) / 2`. Squaring `1/sqrt(2)` gives the `1/2`, and both outputs get -/// the same value because both are a sum of the same two inputs, just -/// with one sign flipped. -/// -/// This runs the same multi-level recursion as [`haar_reg_fwd_level`], -/// over plain host `f32`, so a filter kernel can propagate a per-patch -/// sigma down to a per-coefficient sigma without a GPU round trip. -/// -/// Every caller is a test oracle, so this only builds under `cfg(test)` -/// with a GPU runtime feature enabled, matching the callers themselves. +/// Host mirror of the variance propagation `variance_reg_level` applies over `k_use` members. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] pub(crate) fn haar_variance_ladder(sig2: &[f32], k_use: u32) -> Vec { - let mut out = sig2.to_vec(); + let mut variances = sig2.to_vec(); let mut len = k_use; while len > 1 { let half = len / 2; - let snapshot = out[..len as usize].to_vec(); + let snapshot = variances[..len as usize].to_vec(); for p in 0..half { - let va = snapshot[(2 * p) as usize]; - let vb = snapshot[(2 * p + 1) as usize]; - let avg = (va + vb) / 2.0; - out[p as usize] = avg; - out[(half + p) as usize] = avg; + let first = snapshot[(2 * p) as usize]; + let second = snapshot[(2 * p + 1) as usize]; + let mean = (first + second) / 2.0; + variances[p as usize] = mean; + variances[(half + p) as usize] = mean; } len = half; } - out + variances } #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] @@ -318,24 +233,22 @@ mod tests { #[test] fn dct_noise_profile_rho_zero_is_uniform_identity() { - let g = dct_noise_profile(0.0); + let profile = dct_noise_profile(0.0); assert_eq!( - g, [1.0f32; 8], - "rho=0 must give exactly 1.0 at every frequency, got {g:?}" + profile, [1.0f32; 8], + "rho=0 must give exactly 1.0 at every frequency, got {profile:?}" ); - // A small negative rho is not a real correlation either, and - // must fall onto the same exact identity rather than running - // the summation with a negative base. - let g_neg = dct_noise_profile(-0.1); - assert_eq!(g_neg, [1.0f32; 8]); + // A negative rho is not a real correlation, so it takes the same exact identity. + let negative_profile = dct_noise_profile(-0.1); + assert_eq!(negative_profile, [1.0f32; 8]); } #[test] fn dct_noise_profile_sums_to_eight_across_a_range_of_rho() { for rho in [0.05f32, 0.3, 0.5, 0.67, 0.8, 0.85, 0.86, 0.95, 0.99] { - let g = dct_noise_profile(rho); - let sum: f32 = g.iter().sum(); + let profile = dct_noise_profile(rho); + let sum: f32 = profile.iter().sum(); assert!( (sum - 8.0).abs() < 1e-3, "rho={rho}: expected sum(g) == 8.0 (variance redistributed, not created or \ @@ -347,16 +260,16 @@ mod tests { #[test] fn dct_noise_profile_is_monotonically_decreasing_for_positive_rho() { for rho in [0.05f32, 0.3, 0.5, 0.67, 0.8, 0.85, 0.86, 0.95, 0.99] { - let g = dct_noise_profile(rho); + let profile = dct_noise_profile(rho); for u in 0..7 { assert!( - g[u] > g[u + 1], + profile[u] > profile[u + 1], "rho={rho}: expected g to strictly decrease with frequency (low frequencies \ carry more of a positively correlated residual's noise power), got \ g[{u}]={} <= g[{}]={}", - g[u], + profile[u], u + 1, - g[u + 1], + profile[u + 1], ); } } @@ -366,45 +279,36 @@ mod tests { fn uniform_variance_is_unchanged_by_the_ladder() { for k in [1u32, 2, 4, 8] { let sig2 = vec![0.3f32; k as usize]; - let out = haar_variance_ladder(&sig2, k); - for (idx, &v) in out.iter().enumerate() { - assert!((v - 0.3).abs() < 1e-6, "k={k} idx={idx}: got {v}"); + let variances = haar_variance_ladder(&sig2, k); + for (idx, &variance) in variances.iter().enumerate() { + assert!((variance - 0.3).abs() < 1e-6, "k={k} idx={idx}: got {variance}"); } } } #[test] fn two_element_ladder_averages_the_pair() { - let out = haar_variance_ladder(&[1.0, 0.0], 2); - assert_eq!(out.len(), 2); - assert!((out[0] - 0.5).abs() < 1e-6); - assert!((out[1] - 0.5).abs() < 1e-6); + let variances = haar_variance_ladder(&[1.0, 0.0], 2); + assert_eq!(variances.len(), 2); + assert!((variances[0] - 0.5).abs() < 1e-6); + assert!((variances[1] - 0.5).abs() < 1e-6); } #[test] fn k_use_of_one_is_the_identity() { - let out = haar_variance_ladder(&[0.7], 1); - assert_eq!(out, vec![0.7]); + let variances = haar_variance_ladder(&[0.7], 1); + assert_eq!(variances, vec![0.7]); } #[test] fn eight_element_ladder_matches_hand_computed_levels() { - // Level 1 (len=8, half=4) averages the pairs (0,1), (2,3), - // (4,5), (6,7), and writes each pair's average into both its - // approximation slot and its detail slot, giving - // [2, 2, 3, 2, 2, 2, 3, 2]. - // - // Level 2 (len=4, half=2) only touches positions 0..4, again - // writing each pair's average into both output slots, giving - // [2, 2.5, 2, 2.5] there and leaving 4..8 alone. - // - // Level 3 (len=2, half=1) only touches positions 0..2, folding - // that last pair down to [2.25, 2.25]. + // Level 1 averages each pair into both its slots, giving [2, 2, 3, 2, 2, 2, 3, 2]. Level 2 + // turns positions 0..4 into [2, 2.5, 2, 2.5], and level 3 folds 0..2 into [2.25, 2.25]. let sig2 = vec![1.0, 3.0, 2.0, 2.0, 5.0, 1.0, 4.0, 0.0]; - let out = haar_variance_ladder(&sig2, 8); + let variances = haar_variance_ladder(&sig2, 8); let expected = [2.25f32, 2.25, 2.0, 2.5, 2.0, 2.0, 3.0, 2.0]; - for (idx, (&got, &want)) in out.iter().zip(expected.iter()).enumerate() { + for (idx, (&got, &want)) in variances.iter().zip(expected.iter()).enumerate() { assert!((got - want).abs() < 1e-6, "idx={idx}: got {got} want {want}"); } } diff --git a/av-denoise-core/src/collab/mod.rs b/av-denoise-core/src/collab/mod.rs index d968ffb..17c45a0 100644 --- a/av-denoise-core/src/collab/mod.rs +++ b/av-denoise-core/src/collab/mod.rs @@ -1,68 +1,42 @@ -//! A BM3D-style collaborative filter core, used by the `nl4d` denoiser. -//! -//! Collaborative filtering cleans a frame by grouping similar patches -//! together and denoising the whole group at once, rather than pixel by -//! pixel. A reference patch collects a stack of the patches that match it -//! best, the stack is filtered as a unit, and the filtered results are -//! blended back onto the frame. -//! -//! [`geometry`] lays out the grid of reference patches a frame is covered -//! with, and sizes the buffers driven by that grid. [`kernels`] holds the -//! GPU code that does the grouping, filtering, and blending. - pub mod geometry; pub mod kernels; -use cubecl::prelude::*; - -/// Whether `R` needs [`kernels::fused::collab_fused`]'s warp-uniform -/// search. -/// -/// That kernel's searches are group-scoped: eight lanes share one -/// reference patch and complete every distance with a shuffle across -/// just those eight. The CUDA backend still lowers each shuffle to a -/// `__shfl_*_sync` over the whole 32-lane warp, which on Volta and later -/// waits for every lane it names. Groups that leave the search early -/// never come back to release the ones still in it, so the warp -/// deadlocks and the launch never retires a frame. -/// -/// The wgpu backends emit subgroup operations that reconverge on their -/// own, so they run the cheaper path that walks only the clipped -/// rectangles. See the `warp_uniform` argument for what the other path -/// changes. -pub fn needs_warp_uniform_search(client: &ComputeClient) -> bool { - R::name(client) == "cuda" -} - -// Every test in this tree runs against a real GPU runtime, see -// `tests::helpers::R`, so it only builds when a wgpu-backed feature is -// enabled. A cpu-only build skips it entirely. +// The tests run against a real GPU runtime, so they need a wgpu-backed feature. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] mod tests; +use cubecl::prelude::*; + /// Side length of a collaborative patch in pixels. pub const PATCH_SIZE: u32 = 8; -/// Pixels in one patch. pub const PATCH_AREA: u32 = PATCH_SIZE * PATCH_SIZE; /// Stride of the reference-patch grid. pub const STEP: u32 = 4; -/// Hard ceiling on the group size. Power of two, sized so a stack of K -/// 8x8 f32 patches stays small in shared memory. +/// Hard ceiling on the group size. +/// +/// A power of two, sized so a stack of 8x8 f32 patches stays small in shared memory. pub const MAX_K: u32 = 8; /// Hard ceiling on a cross-frame denoiser's temporal radius. /// -/// [`crate::nl4d::Nl4dParams::validate`] rejects a `temporal_radius` -/// outside `1..=MAX_TEMPORAL_RADIUS`, so this is also the largest radius -/// [`kernels::aggregate::cross_frame_accum_scale`] has to derive a scale -/// for. The widest cross-frame accumulator holds -/// `2 * MAX_TEMPORAL_RADIUS + 1` frames, which bounds how many passes can -/// write into one pixel before it is read back. +/// The widest cross-frame accumulator holds `2 * MAX_TEMPORAL_RADIUS + 1` frames, which bounds how +/// many passes can write into one pixel before it is read back. pub const MAX_TEMPORAL_RADIUS: u32 = 8; +/// Whether `R` needs the fused kernel's warp-uniform search. +/// +/// The fused kernel's searches are group-scoped. Eight lanes share one reference patch and finish +/// each distance with a shuffle across just those eight. CUDA lowers each shuffle to a +/// `__shfl_*_sync` over the whole 32-lane warp, which on Volta and later waits for every lane it +/// names. A group that leaves the search early never releases the ones still in it, so the warp +/// deadlocks. The wgpu backends reconverge on their own and run the cheaper clipped search. +pub fn needs_warp_uniform_search(client: &ComputeClient) -> bool { + R::name(client) == "cuda" +} + /// Frames per volume in a cross-frame group at `temporal_radius`. /// /// A group holds `MAX_K` patches as `MAX_K / grid_frames` volumes. Radius 0 has no neighbour to -/// follow, so it gives 1, which leaves every group a single-frame one. +/// follow, so every group is a single-frame one. pub fn grid_frames(temporal_radius: u32) -> u32 { match temporal_radius { 0 => 1, diff --git a/av-denoise-core/src/collab/tests/aggregate.rs b/av-denoise-core/src/collab/tests/aggregate.rs index 772e786..7c7a1ce 100644 --- a/av-denoise-core/src/collab/tests/aggregate.rs +++ b/av-denoise-core/src/collab/tests/aggregate.rs @@ -1,6 +1,6 @@ use cubecl::prelude::*; -use super::helpers::{R, make_client, noisy_field_over}; +use super::helpers::{R, make_client, noisy_flat_field}; use crate::collab::geometry::{fused_cubes_x, ref_count, refs_along, strength_map_dims}; use crate::collab::kernels::aggregate::{ ACCUM_SCALE, @@ -15,22 +15,29 @@ use crate::collab::kernels::transforms::dct_noise_profile; use crate::collab::{PATCH_SIZE, grid_frames, needs_warp_uniform_search}; use crate::nlmeans::{BLOCK_X, BLOCK_Y, NOISE_CURVE_BINS}; -/// Runs [`collab_normalise`] over hand-built accumulators. +/// Runs [collab_normalise] over hand-built accumulators. fn run_normalise(accum_host: &[i32], wsum_host: &[i32], width: u32, height: u32) -> Vec { let pixels = (width * height) as usize; assert_eq!(accum_host.len(), pixels); assert_eq!(wsum_host.len(), pixels); let client = make_client(); - let accum = client.create_from_slice(i32::as_bytes(accum_host)); - let wsum = client.create_from_slice(i32::as_bytes(wsum_host)); + let accum_bytes = i32::as_bytes(accum_host); + let wsum_bytes = i32::as_bytes(wsum_host); + let accum = client.create_from_slice(accum_bytes); + let wsum = client.create_from_slice(wsum_bytes); let output = client.empty(pixels * size_of::()); + let blocks_x = width.div_ceil(BLOCK_X); + let blocks_y = height.div_ceil(BLOCK_Y); + let grid = CubeCount::new_2d(blocks_x, blocks_y); + let dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); + unsafe { collab_normalise::launch_unchecked::( &client, - CubeCount::new_2d(width.div_ceil(BLOCK_X), height.div_ceil(BLOCK_Y)), - CubeDim::new_2d(BLOCK_X, BLOCK_Y), + grid, + dim, 1usize, ArrayArg::from_raw_parts(accum, pixels), ArrayArg::from_raw_parts(wsum, pixels), @@ -43,29 +50,28 @@ fn run_normalise(accum_host: &[i32], wsum_host: &[i32], width: u32, height: u32) ); } - let bytes = client.read_one(output).expect("normalise readback failed"); - f32::from_bytes(&bytes)[..pixels].to_vec() + let output_bytes = client.read_one(output).expect("normalise readback failed"); + + f32::from_bytes(&output_bytes)[..pixels].to_vec() } #[test] fn normalise_divides_one_accumulator_by_the_other() { - let (w, h) = (21u32, 16u32); - let pixels = (w * h) as usize; + let (width, height) = (21u32, 16u32); + let pixels = (width * height) as usize; - // Varied, non-constant fills, so a transposed index or a dropped - // pixel changes the answer rather than vanishing into a fixed point. + // Varied fills, so a transposed index or a dropped pixel changes the answer rather than vanishing + // into a fixed point. let accum: Vec = (0..pixels).map(|i| (i as i32 % 97) * 1000 - 4000).collect(); let wsum: Vec = (0..pixels).map(|i| (i as i32 % 13) + 1).collect(); - let got = run_normalise(&accum, &wsum, w, h); + let got = run_normalise(&accum, &wsum, width, height); for i in 0..pixels { - // `wsum` counts at `WEIGHT_GAIN` times `accum`'s scale, the one - // factor that does not cancel between the two. + // `wsum` counts at `WEIGHT_GAIN` times `accum`'s scale, the one factor that does not cancel. let want = accum[i] as f32 * WEIGHT_GAIN / wsum[i] as f32; - // Relative, because the ratios here run into the thousands and - // a single-precision divide is only good to about 1e-7 of the - // value either way. + // Relative, because the ratios run into the thousands and a single-precision divide is only + // good to about 1e-7 of the value. assert!( (got[i] - want).abs() <= want.abs() * 1e-6, "idx={i}: want {want} got {}", @@ -74,13 +80,10 @@ fn normalise_divides_one_accumulator_by_the_other() { } } -/// The fixed-point scale cancels, so a pixel whose accumulator and -/// weight sum were both built at the same scale reads back as the plain -/// ratio with no scale factor left in it. #[test] fn normalise_cancels_the_fixed_point_scale() { - let (w, h) = (16u32, 16u32); - let pixels = (w * h) as usize; + let (width, height) = (16u32, 16u32); + let pixels = (width * height) as usize; let value = 0.375f32; let weight = 0.25f32; @@ -89,38 +92,45 @@ fn normalise_cancels_the_fixed_point_scale() { let accum = vec![((value * weight * ACCUM_SCALE) as i32) * covering; pixels]; let wsum = vec![((weight * ACCUM_SCALE * WEIGHT_GAIN) as i32) * covering; pixels]; - let got = run_normalise(&accum, &wsum, w, h); - for (i, &v) in got.iter().enumerate() { - assert!((v - value).abs() < 1e-4, "idx={i}: want {value} got {v}"); + let got = run_normalise(&accum, &wsum, width, height); + + for (i, &pixel) in got.iter().enumerate() { + assert!((pixel - value).abs() < 1e-4, "idx={i}: want {value} got {pixel}"); } } #[test] fn a_zero_weight_sum_returns_the_accumulator_rather_than_a_nan() { - let (w, h) = (16u32, 16u32); - let pixels = (w * h) as usize; + let (width, height) = (16u32, 16u32); + let pixels = (width * height) as usize; let accum = vec![1234i32; pixels]; let wsum = vec![0i32; pixels]; - let got = run_normalise(&accum, &wsum, w, h); - for (i, &v) in got.iter().enumerate() { - assert!(v.is_finite(), "idx={i}: expected a finite value, got {v}"); - assert_eq!(v, 1234.0, "idx={i}"); + let got = run_normalise(&accum, &wsum, width, height); + + for (i, &pixel) in got.iter().enumerate() { + assert!(pixel.is_finite(), "idx={i}: expected a finite value, got {pixel}"); + assert_eq!(pixel, 1234.0, "idx={i}"); } } #[test] fn zero_accum_clears_both_buffers() { - let (w, h) = (16u32, 16u32); - let pixels = (w * h) as usize; + let (width, height) = (16u32, 16u32); + let pixels = (width * height) as usize; let client = make_client(); - let accum = client.create_from_slice(i32::as_bytes(&vec![42i32; pixels])); - let wsum = client.create_from_slice(i32::as_bytes(&vec![7i32; pixels])); + let filled_accum = vec![42i32; pixels]; + let filled_wsum = vec![7i32; pixels]; + let filled_accum_bytes = i32::as_bytes(&filled_accum); + let filled_wsum_bytes = i32::as_bytes(&filled_wsum); + let accum = client.create_from_slice(filled_accum_bytes); + let wsum = client.create_from_slice(filled_wsum_bytes); let dim = 256u32; let grid = (pixels as u32).div_ceil(dim); + unsafe { collab_zero_accum::launch_unchecked::( &client, @@ -135,23 +145,23 @@ fn zero_accum_clears_both_buffers() { ); } - let a = client.read_one(accum).expect("accum readback failed"); - let s = client.read_one(wsum).expect("wsum readback failed"); - assert!(i32::from_bytes(&a)[..pixels].iter().all(|&v| v == 0)); - assert!(i32::from_bytes(&s)[..pixels].iter().all(|&v| v == 0)); + let accum_bytes = client.read_one(accum).expect("accum readback failed"); + let wsum_bytes = client.read_one(wsum).expect("wsum readback failed"); + let accum_cleared = i32::from_bytes(&accum_bytes)[..pixels] + .iter() + .all(|&value| value == 0); + let wsum_cleared = i32::from_bytes(&wsum_bytes)[..pixels] + .iter() + .all(|&value| value == 0); + assert!(accum_cleared); + assert!(wsum_cleared); } -/// A buffer sized past the GPU's 65,535-workgroups-per-dimension -/// dispatch limit at the 256-thread block size the caller launches -/// this kernel with, and not a multiple of that block size either, so -/// the tail both needs the grid clamp and lands mid-block. +/// The buffer needs 65,626 workgroups of 256 threads, past the 65,535 dispatch limit, and is not a +/// multiple of 256, so the tail both needs the grid clamp and lands mid-block. /// -/// A one-thread-per-slot launch clamped to that limit stops short of -/// `pixels`, leaving the tail un-zeroed, which is exactly the silent -/// under-zeroing the task's clamp-without-striding trap describes. -/// `collab_zero_accum` is grid-strided so a clamped launch still walks -/// every slot in a second pass, this buffer needs a real 65,626-thread -/// unclamped grid, only 65,535 of which the dispatch actually starts. +/// A clamped one-thread-per-slot launch would stop short and leave the tail un-zeroed. +/// `collab_zero_accum` is grid-strided, so the clamped launch still walks every slot. #[test] fn zero_accum_clears_every_slot_of_a_buffer_past_the_grid_clamp() { const MAX_GRID_1D: u32 = 65_535; @@ -168,10 +178,15 @@ fn zero_accum_clears_every_slot_of_a_buffer_past_the_grid_clamp() { ); let client = make_client(); - let accum = client.create_from_slice(i32::as_bytes(&vec![42i32; pixels])); - let wsum = client.create_from_slice(i32::as_bytes(&vec![7i32; pixels])); + let filled_accum = vec![42i32; pixels]; + let filled_wsum = vec![7i32; pixels]; + let filled_accum_bytes = i32::as_bytes(&filled_accum); + let filled_wsum_bytes = i32::as_bytes(&filled_wsum); + let accum = client.create_from_slice(filled_accum_bytes); + let wsum = client.create_from_slice(filled_wsum_bytes); let grid = (pixels as u32).div_ceil(dim).min(MAX_GRID_1D); + unsafe { collab_zero_accum::launch_unchecked::( &client, @@ -186,28 +201,31 @@ fn zero_accum_clears_every_slot_of_a_buffer_past_the_grid_clamp() { ); } - let a = client.read_one(accum).expect("accum readback failed"); - let s = client.read_one(wsum).expect("wsum readback failed"); - let a = i32::from_bytes(&a); - let s = i32::from_bytes(&s); + let accum_bytes = client.read_one(accum).expect("accum readback failed"); + let wsum_bytes = client.read_one(wsum).expect("wsum readback failed"); + let accum_values = i32::from_bytes(&accum_bytes); + let wsum_values = i32::from_bytes(&wsum_bytes); + for i in 0..pixels { - assert_eq!(a[i], 0, "accum[{i}] left un-zeroed past the clamp point"); - assert_eq!(s[i], 0, "wsum[{i}] left un-zeroed past the clamp point"); + assert_eq!( + accum_values[i], 0, + "accum[{i}] left un-zeroed past the clamp point" + ); + assert_eq!(wsum_values[i], 0, "wsum[{i}] left un-zeroed past the clamp point"); } } -/// Groups, filters, and aggregates a frame end to end, returning the -/// finished plane and the weight sum behind it. +/// Groups, filters and aggregates a frame end to end, returning the finished plane and its weight sum. /// -/// The search runs at `radius = 0`, a one-frame ring with no -/// neighbours, so this covers the single-frame scatter path the -/// aggregation kernels are being checked on here. +/// The search runs at radius 0, a one-frame ring with no neighbours, so only the single-frame scatter +/// path runs. fn run_scatter_stage(frame: &[f32], width: u32, height: u32, sigma: f32) -> (Vec, Vec) { run_scatter_stage_windowed(frame, width, height, sigma, 0.0) } -/// [`run_scatter_stage`] with the aggregation window's `beta` chosen by -/// the caller. `0.0` is the uniform blend every other run here uses. +/// [run_scatter_stage] with the aggregation window's beta chosen by the caller. +/// +/// A beta of `0.0` is a uniform blend. fn run_scatter_stage_windowed( frame: &[f32], width: u32, @@ -221,27 +239,53 @@ fn run_scatter_stage_windowed( let k_max = 8u32; let pixels = (width * height) as usize; - let input = client.create_from_slice(f32::as_bytes(frame)); - let mv_dummy = client.create_from_slice(i32::as_bytes(&[0i32, 0i32])); - let conf_dummy = client.create_from_slice(f32::as_bytes(&[1.0f32])); - let slots_dummy = client.create_from_slice(u32::as_bytes(&[0u32])); + let frame_bytes = f32::as_bytes(frame); + let mv_dummy_bytes = i32::as_bytes(&[0i32, 0i32]); + let conf_dummy_bytes = f32::as_bytes(&[1.0f32]); + let slots_dummy_bytes = u32::as_bytes(&[0u32]); + let input = client.create_from_slice(frame_bytes); + let mv_dummy = client.create_from_slice(mv_dummy_bytes); + let conf_dummy = client.create_from_slice(conf_dummy_bytes); + let slots_dummy = client.create_from_slice(slots_dummy_bytes); let accum = client.empty(pixels * size_of::()); let wsum = client.empty(pixels * size_of::()); let group_weight = client.empty(refs * size_of::()); - let sigma_buf = client.create_from_slice(f32::as_bytes(&[sigma])); + + let sigma_values = [sigma]; + let sigma_bytes = f32::as_bytes(&sigma_values); + let sigma_buf = client.create_from_slice(sigma_bytes); let profile = dct_noise_profile(0.0); - let profile_buf = client.create_from_slice(f32::as_bytes(&profile)); - let kaiser_buf = client.create_from_slice(f32::as_bytes(&kaiser_window(kaiser_beta))); - let zero_curve = client.create_from_slice(f32::as_bytes(&[0.0f32; NOISE_CURVE_BINS])); + let profile_bytes = f32::as_bytes(&profile); + let profile_buf = client.create_from_slice(profile_bytes); + let kaiser = kaiser_window(kaiser_beta); + let kaiser_bytes = f32::as_bytes(&kaiser); + let kaiser_buf = client.create_from_slice(kaiser_bytes); + let zeroed_curve = [0.0f32; NOISE_CURVE_BINS]; + let curve_bytes = f32::as_bytes(&zeroed_curve); + let zero_curve = client.create_from_slice(curve_bytes); + let (map_cols, map_rows) = strength_map_dims(width, height); let map_len = (map_cols * map_rows) as usize; let unit_map = vec![1.0f32; map_len]; - let unit_map_buf = client.create_from_slice(f32::as_bytes(&unit_map)); + let unit_map_bytes = f32::as_bytes(&unit_map); + let unit_map_buf = client.create_from_slice(unit_map_bytes); let output = client.empty(pixels * size_of::()); let zero_dim = 256u32; + let zero_grid = (pixels as u32).div_ceil(zero_dim); + let cubes_x = fused_cubes_x(width); + let fused_grid = CubeCount::new_2d(cubes_x, refs_y); + let fused_dim = CubeDim::new_1d(64); + let scale = weight_scale(sigma, &profile); + let warp_uniform = needs_warp_uniform_search(&client); + let grid_frame_count = grid_frames(0); + let refs_x = refs_along(width); + let blocks_x = width.div_ceil(BLOCK_X); + let blocks_y = height.div_ceil(BLOCK_Y); + let normalise_grid = CubeCount::new_2d(blocks_x, blocks_y); + let normalise_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); + unsafe { - let zero_grid = (pixels as u32).div_ceil(zero_dim); collab_zero_accum::launch_unchecked::( &client, CubeCount::new_1d(zero_grid), @@ -255,8 +299,8 @@ fn run_scatter_stage_windowed( ); collab_fused::launch_unchecked::( &client, - CubeCount::new_2d(fused_cubes_x(width), refs_y), - CubeDim::new_1d(64), + fused_grid, + fused_dim, 1usize, ArrayArg::from_raw_parts(input.clone(), pixels), ArrayArg::from_raw_parts(mv_dummy, 2), @@ -275,11 +319,11 @@ fn run_scatter_stage_windowed( 2.7f32, 0u32, STRENGTH_MAP_OFF, - weight_scale(sigma, &profile), + scale, ACCUM_SCALE, - needs_warp_uniform_search(&client), + warp_uniform, 0u32, - grid_frames(0), + grid_frame_count, 0u32, 2u32, 1u32, @@ -293,7 +337,7 @@ fn run_scatter_stage_windowed( k_max, 1u32, 9u32, - refs_along(width), + refs_x, map_cols, map_rows, 0.0f32, @@ -301,8 +345,8 @@ fn run_scatter_stage_windowed( ); collab_normalise::launch_unchecked::( &client, - CubeCount::new_2d(width.div_ceil(BLOCK_X), height.div_ceil(BLOCK_Y)), - CubeDim::new_2d(BLOCK_X, BLOCK_Y), + normalise_grid, + normalise_dim, 1usize, ArrayArg::from_raw_parts(accum, pixels), ArrayArg::from_raw_parts(wsum.clone(), pixels), @@ -315,36 +359,27 @@ fn run_scatter_stage_windowed( ); } - let out = client.read_one(output).expect("output readback failed"); - let ws = client.read_one(wsum).expect("wsum readback failed"); - ( - f32::from_bytes(&out)[..pixels].to_vec(), - i32::from_bytes(&ws)[..pixels].to_vec(), - ) + let output_bytes = client.read_one(output).expect("output readback failed"); + let wsum_bytes = client.read_one(wsum).expect("wsum readback failed"); + let plane = f32::from_bytes(&output_bytes)[..pixels].to_vec(); + let weight_sum = i32::from_bytes(&wsum_bytes)[..pixels].to_vec(); + + (plane, weight_sum) } -/// The sharpest check the scatter has, and it does not depend on the -/// filter doing anything in particular. -/// -/// At `sigma = 0` the hard threshold keeps every coefficient, so each -/// member's filtered patch comes back as an exact copy of the input at -/// that member's own position. A member patch sitting at `q` contributes -/// its pixel `q + offset` to output pixel `q + offset`, so every single -/// contribution any pixel receives is that pixel's own input value, -/// whatever group it travelled through. The weighted mean of a set of -/// identical values is that value, so the whole scatter and normalise -/// path has to reproduce the input exactly. +/// At `sigma = 0` the hard threshold keeps every coefficient, so each member's filtered patch is an +/// exact copy of the input at that member's own position. /// -/// Any addressing mistake breaks this. A member written to the reference -/// patch's position instead of its own, a transposed `x`/`y`, or an -/// off-by-one in the pixel index all pull in a neighbouring pixel's -/// value and move the result. +/// Every contribution a pixel receives is then its own input value, and the weighted mean of +/// identical values is that value, so the scatter and normalise path must reproduce the input. A +/// member written to the reference's position, a transposed `x`/`y` or an off-by-one pixel index +/// all pull in a neighbouring pixel's value and move the result. #[test] fn scattering_every_member_at_zero_sigma_reproduces_the_input() { - let (w, h) = (48u32, 40u32); - let frame = noisy_field_over(w, h, 0.5, 0.05); + let (width, height) = (48u32, 40u32); + let frame = noisy_flat_field(width, height, 0.5, 0.05); - let (output, _) = run_scatter_stage(&frame, w, h, 0.0); + let (output, _) = run_scatter_stage(&frame, width, height, 0.0); for (idx, (&want, &have)) in frame.iter().zip(output.iter()).enumerate() { assert!( @@ -354,34 +389,28 @@ fn scattering_every_member_at_zero_sigma_reproduces_the_input() { } } -/// Proves the aggregation really covers every member and not just the -/// reference patch of each group. -/// -/// Reference patches alone sit on a grid of stride `STEP` and are -/// `PATCH_SIZE` wide, so they can cover any one pixel at most nine -/// times. Members are drawn from a window of radius 9 around their -/// reference, so once every member is written back an interior pixel -/// picks up far more contributions than that ceiling allows. +/// Reference patches sit on a stride-`STEP` grid and are `PATCH_SIZE` wide, so they cover any pixel +/// at most nine times. Members come from a radius-9 window around their reference, so writing every +/// member back gives an interior pixel far more contributions than that. /// -/// The weight sum is read rather than a contribution count, because -/// that is what aggregation actually divides by. Every group here -/// carries the same weight, since a flat noise field gives every group -/// the same retained variance, so the sum is proportional to the number -/// of covering patches. +/// The weight sum is read because aggregation divides by it. A flat noise field gives every group the +/// same retained variance and so the same weight, which makes the sum proportional to the number of +/// covering patches. #[test] fn every_member_reaches_the_weight_sum_not_only_the_reference_patch() { - let (w, h) = (64u32, 64u32); - let frame = noisy_field_over(w, h, 0.5, 0.02); + let (width, height) = (64u32, 64u32); + let frame = noisy_flat_field(width, height, 0.5, 0.02); - let (_, wsum) = run_scatter_stage(&frame, w, h, 0.02); + let (_, wsum) = run_scatter_stage(&frame, width, height, 0.02); // Away from the edges, where the search window is not truncated. let mut interior: Vec = Vec::new(); - for y in 16..h - 16 { - for x in 16..w - 16 { - interior.push(wsum[(y * w + x) as usize]); + for y in 16..height - 16 { + for x in 16..width - 16 { + interior.push(wsum[(y * width + x) as usize]); } } + assert!(!interior.is_empty()); let smallest = *interior.iter().min().expect("interior is non-empty"); @@ -390,32 +419,32 @@ fn every_member_reaches_the_weight_sum_not_only_the_reference_patch() { "every interior pixel must receive at least one contribution" ); - // One group's weight, taken as the largest single contribution any - // pixel could have received, bounds the count from above. Nine of - // them is the reference-only ceiling. - let per_patch = interior.iter().map(|&v| v as f64).fold(f64::INFINITY, f64::min); + // The smallest interior sum is at least one group's weight, so the spread is a lower bound on the + // largest contribution count. Nine is the reference-only ceiling. + let per_patch = interior + .iter() + .map(|&weight| weight as f64) + .fold(f64::INFINITY, f64::min); let biggest = *interior.iter().max().expect("interior is non-empty") as f64; + let spread = biggest / per_patch; assert!( - biggest / per_patch > 9.0, + spread > 9.0, "expected some interior pixel to collect more than the nine covering reference \ - patches a member-0-only writeback could manage, got a spread of {}", - biggest / per_patch, + patches a member-0-only writeback could manage, got a spread of {spread}", ); } -/// A window applied to the value but not to the weight would pull every -/// pixel toward zero, hardest at the patch edges where the taper is -/// deepest. Flat content is where that shows up exactly, because the -/// weighted mean of one value is that value however the weights fall. +/// A window applied to the value but not the weight would pull pixels toward zero, hardest at the +/// patch edges. Flat content shows that exactly, since the weighted mean of one value is that value. #[test] fn the_aggregation_window_leaves_flat_content_flat() { - let (w, h) = (64u32, 64u32); + let (width, height) = (64u32, 64u32); let level = 0.5f32; - let frame = vec![level; (w * h) as usize]; + let frame = vec![level; (width * height) as usize]; - let (out, wsum) = run_scatter_stage_windowed(&frame, w, h, 0.02, 2.0); + let (plane, wsum) = run_scatter_stage_windowed(&frame, width, height, 0.02, 2.0); - for (idx, (&got, &weight)) in out.iter().zip(wsum.iter()).enumerate() { + for (idx, (&got, &weight)) in plane.iter().zip(wsum.iter()).enumerate() { assert!(weight > 0, "pixel {idx} collected no weight at all"); assert!( (got - level).abs() < 1e-3, @@ -424,21 +453,19 @@ fn the_aggregation_window_leaves_flat_content_flat() { } } -/// The window enters only at the scatter, so the groups, the threshold -/// and every filtered value are identical across the two runs and the -/// difference between them is the taper alone. +/// The window enters only at the scatter, so the two runs differ by the taper alone. #[test] fn the_aggregation_window_reweights_the_blend() { - let (w, h) = (64u32, 64u32); - let frame = noisy_field_over(w, h, 0.5, 0.05); + let (width, height) = (64u32, 64u32); + let frame = noisy_flat_field(width, height, 0.5, 0.05); - let (uniform, _) = run_scatter_stage_windowed(&frame, w, h, 0.02, 0.0); - let (windowed, _) = run_scatter_stage_windowed(&frame, w, h, 0.02, 2.0); + let (uniform, _) = run_scatter_stage_windowed(&frame, width, height, 0.02, 0.0); + let (windowed, _) = run_scatter_stage_windowed(&frame, width, height, 0.02, 2.0); let moved = uniform .iter() .zip(windowed.iter()) - .filter(|(a, b)| (*a - *b).abs() > 1e-4) + .filter(|(uniform_value, windowed_value)| (*uniform_value - *windowed_value).abs() > 1e-4) .count(); assert!( moved > uniform.len() / 100, diff --git a/av-denoise-core/src/collab/tests/fused/behaviour.rs b/av-denoise-core/src/collab/tests/fused/behaviour.rs index c3019d1..33152ce 100644 --- a/av-denoise-core/src/collab/tests/fused/behaviour.rs +++ b/av-denoise-core/src/collab/tests/fused/behaviour.rs @@ -14,28 +14,23 @@ use crate::collab::geometry::refs_along; use crate::collab::kernels::transforms::dct_noise_profile; use crate::collab::tests::helpers::{deterministic_texture, plant_patch}; -/// At `sigma = 0` every threshold is zero, so nothing is discarded and -/// the transform chain must hand every member's own pixels back -/// unchanged. +/// At `sigma = 0` nothing is discarded, so every contribution a pixel receives is its own input +/// value and the weighted mean must reproduce it. /// -/// Every contribution any pixel receives is then that pixel's own input -/// value, whatever group carried it, and the weighted mean of a set of -/// identical values is that value. `k_max = 1` exercises the `k_use = 1` -/// case, where the stack transform is a no-op and only the 2D DCT round -/// trip runs. `k_max = 8` forces a full stack over content where every -/// position differs from every other, so all three Haar levels carry -/// non-trivial detail coefficients. +/// `k_max = 1` covers `k_use = 1`, where the stack transform is a no-op and only the 2D DCT round +/// trip runs. `k_max = 8` forces a full stack over content where every position differs, so all +/// three Haar levels carry non-trivial detail. #[test] fn zero_sigma_hands_every_member_back_unchanged() { - let (w, h) = (32u32, 32u32); - let frame = unique_frame(w, h); + let (width, height) = (32u32, 32u32); + let frame = unique_frame(width, height); for k_max in [1u32, 8] { - let mut s = Setup::spatial_only(frame.clone(), w, h); - s.k_max = k_max; - s.sigma = 0.0; - s.spatial_radius = 4; - let got = run_fused(&s); + let mut setup = Setup::spatial_only(frame.clone(), width, height); + setup.k_max = k_max; + setup.sigma = 0.0; + setup.spatial_radius = 4; + let got = run_fused(&setup); for (idx, &want) in frame.iter().enumerate() { assert!( @@ -51,26 +46,25 @@ fn zero_sigma_hands_every_member_back_unchanged() { } } -/// A covered pixel never ends with an empty weight sum, even when every neighbour holds content -/// unrelated to the centre. -/// -/// A group that reached the accumulators as nothing would leave such a pixel, which normalisation -/// can only render as black. +/// Every neighbour holds content unrelated to the centre. A group that reached the accumulators as +/// nothing would leave a covered pixel with an empty weight sum, which normalisation renders black. #[test] fn a_badly_matched_group_still_reaches_the_accumulators() { - let (w, h) = (32u32, 32u32); - let counts = reference_cover_counts(w, h); + let (width, height) = (32u32, 32u32); + let counts = reference_cover_counts(width, height); + + let mut setup = cross_frame_setup(width, height, 2); + setup.spatial_radius = 9; + setup.c_min = 0.0; - let mut s = cross_frame_setup(w, h, 2); - s.spatial_radius = 9; - s.c_min = 0.0; + let got = run_fused(&setup); + let base = setup.centre_slot as usize * setup.pixels(); - let got = run_fused(&s); - let base = s.centre_slot as usize * s.pixels(); for (idx, &count) in counts.iter().enumerate() { if count == 0 { continue; } + assert!( got.wsum[base + idx] > 0, "{count} references cover pixel {idx} and its weight sum is still {}", @@ -79,26 +73,26 @@ fn a_badly_matched_group_still_reaches_the_accumulators() { } } -/// The same invariant with the aggregation window on. -/// /// A patch corner is weighted by the square of the window's end tap, `0.193` at `beta = 2`, so the /// smallest weight the fixed point has to resolve drops about fivefold against the uniform case. #[test] fn a_windowed_badly_matched_group_still_reaches_the_accumulators() { - let (w, h) = (32u32, 32u32); - let counts = reference_cover_counts(w, h); + let (width, height) = (32u32, 32u32); + let counts = reference_cover_counts(width, height); + + let mut setup = cross_frame_setup(width, height, 2); + setup.spatial_radius = 9; + setup.c_min = 0.0; + setup.kaiser_beta = 2.0; - let mut s = cross_frame_setup(w, h, 2); - s.spatial_radius = 9; - s.c_min = 0.0; - s.kaiser_beta = 2.0; + let got = run_fused(&setup); + let base = setup.centre_slot as usize * setup.pixels(); - let got = run_fused(&s); - let base = s.centre_slot as usize * s.pixels(); for (idx, &count) in counts.iter().enumerate() { if count == 0 { continue; } + assert!( got.wsum[base + idx] > 0, "{count} references cover pixel {idx} and its weight sum is still {} with the \ @@ -108,27 +102,25 @@ fn a_windowed_badly_matched_group_still_reaches_the_accumulators() { } } -/// The reference patch is always the group's first member. +/// At `k_max = 1` a group scatters only slot 0, and `sigma = 0` gives every group the same weight, +/// so a pixel's weight counts the patches that covered it. /// -/// At `k_max = 1` a group holds exactly one member, so the only patch it -/// scatters is whichever position slot 0 ended up holding. `sigma = 0` -/// makes every group's weight the same constant, so the weight one pixel -/// accumulates counts the patches that covered it. That count must be -/// exactly the number of reference patches covering it, which only holds -/// if every group scattered its own reference position and nothing else. -/// A group that let a search result reach slot 0 would write somewhere -/// off the reference grid and leave the counts uneven. +/// That count matches the reference cover count only if every group scattered its own reference +/// position. A search result reaching slot 0 would write off the reference grid and leave the +/// counts uneven. #[test] fn the_reference_patch_is_always_the_first_member() { - let (w, h) = (32u32, 32u32); - let mut s = Setup::spatial_only(unique_frame(w, h), w, h); - s.k_max = 1; - s.sigma = 0.0; - let got = run_fused(&s); - - let counts = reference_cover_counts(w, h); + let (width, height) = (32u32, 32u32); + let frame = unique_frame(width, height); + let mut setup = Setup::spatial_only(frame, width, height); + setup.k_max = 1; + setup.sigma = 0.0; + let got = run_fused(&setup); + + let counts = reference_cover_counts(width, height); let unit = got.wsum[0] as i64 / counts[0]; assert!(unit > 0, "the per-patch weight increment must be positive"); + for (idx, &count) in counts.iter().enumerate() { assert_eq!( got.wsum[idx] as i64, @@ -140,48 +132,40 @@ fn the_reference_patch_is_always_the_first_member() { } } -/// The group size is the search space size rounded down to a power of -/// two, capped at `k_max`. -/// -/// At `spatial_radius = 1` the clipped rectangle holds 4 positions at a -/// corner reference, 6 at an edge one, and 9 in the interior. Rounding -/// therefore takes the edge references from 6 down to 4, and leaves the -/// interior ones at 8. Running the same frame at `k_max = 4` caps every -/// group at 4, so the two runs must agree exactly wherever rounding -/// already reached 4 and differ wherever it reached 8. +/// The group size is the search space size rounded down to a power of two, capped at `k_max`. /// -/// Clipping the rectangle once is what makes those counts right. Were -/// each offset clamped in turn instead, a corner would count nine -/// positions rather than four, several of them the same physical patch, -/// and the corner references would stop agreeing across the two runs. +/// At `spatial_radius = 1` the clipped rectangle holds 4 positions at a corner reference, 6 at an +/// edge one and 9 in the interior, which round to 4, 4 and 8. A `k_max = 4` run must therefore +/// match wherever rounding already reached 4 and differ wherever it reached 8. Clamping each offset +/// instead of clipping the rectangle would count nine positions at a corner, several of them the +/// same patch, and the corners would stop agreeing. #[test] fn group_size_rounds_down_to_a_power_of_two() { - let (w, h) = (64u32, 64u32); - let frame = unique_frame(w, h); + let (width, height) = (64u32, 64u32); + let frame = unique_frame(width, height); - let mut wide = Setup::spatial_only(frame.clone(), w, h); + let mut wide = Setup::spatial_only(frame.clone(), width, height); wide.spatial_radius = 1; - let mut narrow = Setup::spatial_only(frame, w, h); + let mut narrow = Setup::spatial_only(frame, width, height); narrow.spatial_radius = 1; narrow.k_max = 4; let wide = run_fused(&wide); let narrow = run_fused(&narrow); - let refs_x = refs_along(w); - let refs_y = refs_along(h); + let refs_x = refs_along(width); + let refs_y = refs_along(height); let mut interior_differed = 0usize; - for ry in 0..refs_y { - for rx in 0..refs_x { - let idx = (ry * refs_x + rx) as usize; - // A clipped axis contributes 2 positions instead of 3, so a - // reference is capped below 8 unless both of its axes are - // interior. - let clipped = rx == 0 || ry == 0 || rx == refs_x - 1 || ry == refs_y - 1; + for ref_y in 0..refs_y { + for ref_x in 0..refs_x { + let idx = (ref_y * refs_x + ref_x) as usize; + // A clipped axis contributes 2 positions instead of 3, so a reference is capped below 8 + // unless both of its axes are interior. + let clipped = ref_x == 0 || ref_y == 0 || ref_x == refs_x - 1 || ref_y == refs_y - 1; if clipped { assert_eq!( wide.group_weight[idx], narrow.group_weight[idx], - "reference ({rx}, {ry}) sees fewer than 8 positions, so both runs must \ + "reference ({ref_x}, {ref_y}) sees fewer than 8 positions, so both runs must \ round it to a group of 4" ); } else if wide.group_weight[idx] != narrow.group_weight[idx] { @@ -189,6 +173,7 @@ fn group_size_rounds_down_to_a_power_of_two() { } } } + let interior = ((refs_x - 2) * (refs_y - 2)) as usize; assert!( interior_differed * 2 > interior, @@ -197,49 +182,39 @@ fn group_size_rounds_down_to_a_power_of_two() { ); } -/// A group that finds a genuine twin agrees with itself, and a group -/// that does not carries far more detail into the threshold. -/// -/// One texture is planted twice over a flat background, at `(4, 4)` and -/// at `(16, 12)`. With `k_max = 2` the group at `(4, 4)` keeps the -/// self-match and exactly one other member, so the twin either is that -/// member or the matcher missed it. When it is, the two members are -/// pixel for pixel identical, the Haar difference across the pair is -/// exactly zero everywhere, and the threshold keeps nothing from that -/// level. When it is not, the second member is flat background against a -/// textured reference, the difference level carries the texture too, and -/// roughly twice as many coefficients survive, halving the weight. -/// -/// `lambda_ht` sits at 1.0 so the threshold keeps nearly every -/// coefficient it is offered, which is what makes the retained count -/// track the number of levels carrying content rather than the size of -/// the coefficients in them. +/// One texture is planted at `(4, 4)` and `(16, 12)` over a flat background, and `k_max = 2` keeps +/// the self-match plus one member. /// -/// The control run plants the same texture once, leaving nothing in the -/// window for the group to match. +/// When that member is the twin, the Haar difference across the pair is zero and the threshold +/// keeps nothing from that level. When it is flat background, the difference level carries the +/// texture too, roughly twice as many coefficients survive and the weight halves. `lambda_ht = 1.0` +/// keeps nearly every coefficient offered, so the retained count tracks how many levels carry +/// content. The control plants the texture once, leaving nothing to match. #[test] fn a_planted_twin_is_found() { - let (w, h) = (32u32, 32u32); + let (width, height) = (32u32, 32u32); let texture = deterministic_texture(7); - let mut twinned = vec![0.2f32; (w * h) as usize]; - plant_patch(&mut twinned, w, 4, 4, &texture); - plant_patch(&mut twinned, w, 16, 12, &texture); + let mut twinned = vec![0.2f32; (width * height) as usize]; + plant_patch(&mut twinned, width, 4, 4, &texture); + plant_patch(&mut twinned, width, 16, 12, &texture); - let mut alone = vec![0.2f32; (w * h) as usize]; - plant_patch(&mut alone, w, 4, 4, &texture); + let mut alone = vec![0.2f32; (width * height) as usize]; + plant_patch(&mut alone, width, 4, 4, &texture); let run = |frame: Vec| { - let mut s = Setup::spatial_only(frame, w, h); - s.spatial_radius = 12; - s.k_max = 2; - s.lambda_ht = 1.0; - run_fused(&s) + let mut setup = Setup::spatial_only(frame, width, height); + setup.spatial_radius = 12; + setup.k_max = 2; + setup.lambda_ht = 1.0; + run_fused(&setup) }; - let ref_idx = (4 / STEP + (4 / STEP) * refs_along(w)) as usize; - let with_twin = run(twinned).group_weight[ref_idx]; - let without_twin = run(alone).group_weight[ref_idx]; + let ref_idx = (4 / STEP + (4 / STEP) * refs_along(width)) as usize; + let twinned_run = run(twinned); + let alone_run = run(alone); + let with_twin = twinned_run.group_weight[ref_idx]; + let without_twin = alone_run.group_weight[ref_idx]; assert!( with_twin > without_twin * 1.5, @@ -248,52 +223,44 @@ fn a_planted_twin_is_found() { ); } -/// A neighbour whose motion-block confidence sits below `c_min` is -/// skipped outright, so no member ever comes from it and its region of -/// the accumulator ring stays untouched. -/// -/// The confidence field is uniform per neighbour here, so the skip is -/// the same decision for every group in the frame. Both neighbours hold -/// an exact copy, so neighbour 0 wins every tie. Gating it is what moves -/// every volume's match onto neighbour 1, and a slot that received even -/// one member would show a non-zero weight sum. +/// Confidence is uniform per neighbour, so the skip is the same decision for every group. Both +/// neighbours hold an exact copy and neighbour 0 wins ties, so gating it moves every match onto +/// neighbour 1, and a slot that received even one member would show a non-zero weight sum. #[test] fn a_gated_neighbour_receives_no_scatter() { - let (w, h) = (64u32, 64u32); - let mut s = three_frame_ring_with_a_planted_match(w, h); - // Neighbour 0 is ring slot 0 and neighbour 1 is ring slot 2, so this - // gates the first of the two. - let blocks = s.conf_stride as usize; - s.confidence[..blocks].fill(0.0); - s.confidence[blocks..].fill(1.0); - s.c_min = 0.5; - - let got = run_fused(&s); + let (width, height) = (64u32, 64u32); + let mut setup = three_frame_ring_with_a_planted_match(width, height); + // Neighbour 0 is ring slot 0 and neighbour 1 is ring slot 2, so this gates the first of the two. + let blocks = setup.conf_stride as usize; + setup.confidence[..blocks].fill(0.0); + setup.confidence[blocks..].fill(1.0); + setup.c_min = 0.5; + + let got = run_fused(&setup); + let gated_sum = got.frame_weight_sum(0); + let centre_sum = got.frame_weight_sum(1); + let ungated_sum = got.frame_weight_sum(2); assert_eq!( - got.frame_weight_sum(0), - 0, + gated_sum, 0, "the gated neighbour's slot must receive no scatter at all" ); - assert!(got.frame_weight_sum(1) > 0, "the centre slot received nothing"); - assert!( - got.frame_weight_sum(2) > 0, - "the ungated neighbour's slot received nothing" - ); + assert!(centre_sum > 0, "the centre slot received nothing"); + assert!(ungated_sum > 0, "the ungated neighbour's slot received nothing"); } #[test] fn noise_is_suppressed_on_a_flat_field() { - let (w, h) = (48u32, 48u32); + let (width, height) = (48u32, 48u32); let sigma = 0.04f32; - let s = flat_noise_setup(w, h, sigma); - let input_var = patch_pool_variance(&s.ring, w, h); - let got = run_fused(&s); - - // A run that wrote nothing would read a variance of zero and clear - // the bound below without filtering anything, so the output has to - // be shown to carry the field's own brightness first. - let output_mean: f64 = (0..got.accum.len()).map(|i| got.pixel(i)).sum::() / got.accum.len() as f64; + let setup = flat_noise_setup(width, height, sigma); + let input_var = patch_pool_variance(&setup.ring, width, height); + let got = run_fused(&setup); + + // A run that wrote nothing would read a variance of zero and clear the bound below without + // filtering, so the output must first be shown to keep the field's brightness. + let output_sum: f64 = (0..got.accum.len()).map(|i| got.pixel(i)).sum(); + let output_mean = output_sum / got.accum.len() as f64; assert!( (output_mean - 0.5).abs() < 0.01, "expected the filtered field to keep its 0.5 mean, got {output_mean}" @@ -310,45 +277,30 @@ fn noise_is_suppressed_on_a_flat_field() { #[test] fn group_weight_matches_uniform_theory() { - let (w, h) = (48u32, 48u32); + let (width, height) = (48u32, 48u32); let sigma = 0.04f32; - let s = flat_noise_setup(w, h, sigma); - let weights = run_fused(&s).group_weight; - - // With every member's variance equal to `sigma^2`, the ladder is a - // fixed point, as - // `transforms::tests::uniform_variance_is_unchanged_by_the_ladder` - // shows. Every coefficient the threshold could keep therefore also - // carries variance `sigma^2`, whatever level or spatial position it - // came from. - // - // `group_weight` is then exactly `1 / (sigma^2 * n_ret)`, so this - // backs out the mean retained count the run produced and checks two - // things about it. + let setup = flat_noise_setup(width, height, sigma); + let weights = run_fused(&setup).group_weight; + + // With every member's variance at `sigma^2` the ladder is a fixed point, so every coefficient + // the threshold could keep carries variance `sigma^2`. `group_weight` is then exactly + // `1 / (sigma^2 * n_ret)`, which backs out the mean retained count. // - // It must include at least the forced group DC. And a hard threshold - // at 2.7 standard deviations lets only about 0.7% of pure-noise - // coefficients through by chance, so out of the up to - // `k_max * PATCH_AREA - 1` coefficients besides the DC that a full - // 8-member group offers, the mean false-positive count should be - // small next to that ceiling rather than close to it. + // That count must include the forced group DC. A hard threshold at 2.7 standard deviations lets + // about 0.7% of pure-noise coefficients through, so out of the `k_max * PATCH_AREA - 1` + // non-DC coefficients a full group offers, false positives stay far below that ceiling. let sigma2 = sigma * sigma; - let mean_weight: f64 = weights.iter().map(|&w| w as f64).sum::() / weights.len() as f64; + let weight_total: f64 = weights.iter().map(|&weight| weight as f64).sum(); + let mean_weight = weight_total / weights.len() as f64; let mean_n_ret = 1.0 / (mean_weight * sigma2 as f64); let false_positive_rate = 0.007; // ~P(|Z| >= 2.7) for a standard normal, two-tailed let ceiling = (8 * 64 - 1) as f64; let expected_n_ret = 1.0 + ceiling * false_positive_rate; - // A run against the real kernel at this setup measures a mean - // retained count around 6 (close to `expected_n_ret`, ~4.5, and - // nowhere near a naive DC-only assumption of 1, which a 20% band - // around would reject this correct result outright). The lower - // bound below is what actually distinguishes a working threshold - // from two ways it could be broken: forced-DC-only (would measure - // exactly 1) and "threshold does nothing, keeps everything" (would - // measure close to `ceiling + 1`, an order of magnitude past the - // upper bound below). + // The kernel measures a mean retained count around 6 here, near `expected_n_ret` (about 4.5). + // The lower bound rejects a forced-DC-only threshold (exactly 1), and the upper bound rejects + // one that keeps everything (close to `ceiling + 1`). assert!( mean_n_ret > 2.0, "expected the mean retained count ({mean_n_ret}) to clearly exceed the forced-DC-\ @@ -363,35 +315,30 @@ fn group_weight_matches_uniform_theory() { ); } -/// `rho = 0` must leave the output bit for bit identical to what it -/// would be with no noise-shaping profile in the computation at all. -/// -/// This is checked two ways from the same noisy group, once through the -/// real `dct_noise_profile(0.0)` production path, and once through a -/// profile buffer built entirely by hand, `[1.0; 8]`, which is -/// mathematically the exact identity multiplier and so stands in for "no -/// profile logic at all" without needing a second copy of the kernel to -/// prove it against. +/// The production `dct_noise_profile(0.0)` path is compared against a hand-built `[1.0; 8]` +/// buffer, the exact identity multiplier, which stands in for no profile logic at all. #[test] fn dct_profile_rho_zero_matches_a_hand_built_all_ones_profile() { - let (w, h) = (48u32, 48u32); + let (width, height) = (48u32, 48u32); let sigma = 0.04f32; + let white_profile = dct_noise_profile(0.0); assert_eq!( - dct_noise_profile(0.0), - [1.0f32; 8], + white_profile, [1.0f32; 8], "dct_noise_profile(0.0) must be exactly [1.0; 8], the property this comparison relies on" ); - let produced = flat_noise_setup(w, h, sigma); - let mut hand_built = flat_noise_setup(w, h, sigma); + let produced = flat_noise_setup(width, height, sigma); + let mut hand_built = flat_noise_setup(width, height, sigma); hand_built.profile_override = Some([1.0f32; 8]); let produced = run_fused(&produced); let hand_built = run_fused(&hand_built); + let wrote_accum = produced.accum.iter().any(|&value| value != 0); + let wrote_weight = produced.group_weight.iter().any(|&weight| weight != 0.0); assert!( - produced.accum.iter().any(|&v| v != 0) || produced.group_weight.iter().any(|&w| w != 0.0), + wrote_accum || wrote_weight, "the kernel must actually have written output for this comparison to mean anything" ); assert_eq!( @@ -405,30 +352,25 @@ fn dct_profile_rho_zero_matches_a_hand_built_all_ones_profile() { ); } -/// Higher `rho` must retain more residual noise on a flat, noise-only -/// field than `rho = 0` does, at the same `lambda_ht`. +/// A positive `rho` moves variance from the high frequencies into the low ones, so a fixed +/// `lambda_ht` reaches a smaller threshold on most non-DC coefficients and more pure noise survives. /// -/// A positive `rho` moves variance out of the high frequencies and into -/// the low ones (`dct_noise_profile`'s own monotonic-decrease property), -/// so a fixed `lambda_ht` reaches a smaller threshold on most non-DC -/// coefficients than the white-noise assumption would, and more of the -/// pure noise sitting in those coefficients survives. This is the -/// documented, deliberate trade the shipped table's caveat describes. On -/// content where the true correlation is lower than the table assumes, -/// shaping under-shrinks rather than over-shrinks, trading a little -/// leftover noise for preserved detail. A flat, noise-only field -/// isolates that trade with nothing else going on. +/// This is the deliberate trade of correlation shaping. Where the true correlation is lower than +/// assumed it under-shrinks, leaving a little noise to preserve detail. A flat, noise-only field +/// isolates that trade. #[test] fn higher_rho_retains_more_noise_on_a_flat_field() { - let (w, h) = (48u32, 48u32); + let (width, height) = (48u32, 48u32); let sigma = 0.04f32; - let white = flat_noise_setup(w, h, sigma); - let mut shaped = flat_noise_setup(w, h, sigma); + let white = flat_noise_setup(width, height, sigma); + let mut shaped = flat_noise_setup(width, height, sigma); shaped.rho = 0.86; - let var_white = output_variance(&run_fused(&white)); - let var_shaped = output_variance(&run_fused(&shaped)); + let white_run = run_fused(&white); + let shaped_run = run_fused(&shaped); + let var_white = output_variance(&white_run); + let var_shaped = output_variance(&shaped_run); assert!( var_shaped > var_white * 1.05, @@ -437,24 +379,21 @@ fn higher_rho_retains_more_noise_on_a_flat_field() { ); } -/// At radius 1 each volume keeps one neighbour frame, and the first-listed neighbour wins a tie. -/// /// Both neighbours hold an exact copy of the centre, so every volume's two candidates tie at zero -/// and slot 0, neighbour 0, takes every match. A radius-1 group falling back to a single frame -/// would leave slot 0 empty as well. +/// and neighbour 0 takes every match. A radius-1 group falling back to a single frame would leave +/// slot 0 empty as well. #[test] fn radius_one_keeps_one_neighbour_per_volume() { - let s = three_frame_ring_with_a_planted_match(64, 64); - let got = run_fused(&s); - - assert!( - got.frame_weight_sum(0) > 0, - "neighbour 0 must hold every volume's frame" - ); - assert!(got.frame_weight_sum(1) > 0, "the centre slot received nothing"); + let setup = three_frame_ring_with_a_planted_match(64, 64); + let got = run_fused(&setup); + let first_sum = got.frame_weight_sum(0); + let centre_sum = got.frame_weight_sum(1); + let second_sum = got.frame_weight_sum(2); + + assert!(first_sum > 0, "neighbour 0 must hold every volume's frame"); + assert!(centre_sum > 0, "the centre slot received nothing"); assert_eq!( - got.frame_weight_sum(2), - 0, + second_sum, 0, "a 2x4 volume keeps one neighbour, so the tied second neighbour must receive nothing" ); } diff --git a/av-denoise-core/src/collab/tests/fused/mod.rs b/av-denoise-core/src/collab/tests/fused/mod.rs index eb71d66..4ca51d6 100644 --- a/av-denoise-core/src/collab/tests/fused/mod.rs +++ b/av-denoise-core/src/collab/tests/fused/mod.rs @@ -8,7 +8,7 @@ mod walks; use cubecl::prelude::*; use cubecl::server::Handle; -use super::helpers::{R, make_client, make_unique_frame, noisy_field_over}; +use super::helpers::{R, make_client, make_unique_frame, noisy_flat_field}; use crate::collab::geometry::{fused_cubes_x, ref_count, ref_pos, refs_along, strength_map_dims}; use crate::collab::kernels::aggregate::{WEIGHT_GAIN, cross_frame_accum_scale, kaiser_window, weight_scale}; use crate::collab::kernels::fused::{STRENGTH_MAP_OFF, collab_fused}; @@ -16,61 +16,48 @@ use crate::collab::kernels::transforms::dct_noise_profile; use crate::collab::{PATCH_SIZE, grid_frames, needs_warp_uniform_search}; use crate::nlmeans::{ChannelMode, NOISE_CURVE_BINS}; -/// The spatial search radius most runs below use. +/// The spatial search radius most runs use. /// -/// Large enough that a reference patch away from the frame edge scores -/// a 9x9 window, which is well past the eight members a group keeps, -/// and small enough that the whole sweep stays quick. It is a [`Setup`] -/// field rather than a constant so one test can narrow it far enough to -/// shrink a group below `k_max`. +/// A reference patch away from the frame edge scores a 9x9 window, well past the eight members a +/// group keeps, and the whole sweep stays quick. const SPATIAL_RADIUS: u32 = 4; -/// The group size most runs below use. The fused kernel carries one -/// member per lane of an 8-lane group, so this is the size it is built +/// The group size most runs use. +/// +/// The fused kernel carries one member per lane of an 8-lane group, so this is the size it is built /// for. const K_MAX: u32 = 8; -/// Motion-block side length. The kernel searches every block whose -/// `blksize` span contains a patch. +/// Motion-block side length. const BLKSIZE: u32 = 16; -/// Motion-block stride. It stays at `PATCH_SIZE` so a block boundary -/// lines up with a patch boundary. +/// Motion-block stride, equal to `PATCH_SIZE` so a block boundary lines up with a patch boundary. const BLK_STEP: u32 = 8; -/// The noise level the filter is told to shrink against. +/// The noise level the filter shrinks against. /// -/// Small enough against content in `[0, 1]` that the threshold keeps a -/// spread of coefficients rather than everything or nothing, so both -/// sides of the keep decision are exercised. +/// Against content between 0 and 1 the threshold keeps a spread of coefficients rather than +/// everything or nothing, so both sides of the keep decision are exercised. const SIGMA: f32 = 0.02; -/// A fixed hard-threshold multiplier, pinned independently of -/// `Nl4dParams::default().lambda_ht`. +/// A hard-threshold multiplier pinned independently of the shipped `lambda_ht` default. /// -/// Several tests in this file recorded their expected output at this -/// value, so it stays fixed even when the shipped default moves. +/// Several tests recorded their expected output at this value. const LAMBDA_HT: f32 = 5.3; -/// [`make_unique_frame`] rescaled into `[0, 1]`. +/// [make_unique_frame] rescaled to between 0 and 1. /// -/// That helper ramps to ten times the frame width, which suits a -/// matching test and breaks a filtering one. Everything downstream of -/// the match is defined over `[0, 1]`, the scatter clamps at -/// [`crate::collab::kernels::aggregate::ACCUM_CLAMP`], and a patch of -/// values in the hundreds both saturates that clamp and puts every -/// coefficient so far above the noise threshold that the threshold stops -/// being tested at all. Dividing by a constant leaves every 8x8 window -/// exactly as distinct as it was, so the tie-free property these runs -/// rely on is untouched. -pub(super) fn unique_frame(w: u32, h: u32) -> Vec { - let raw = make_unique_frame(w, h); +/// The raw ramp reaches ten times the frame width. Values in the hundreds saturate +/// [ACCUM_CLAMP](crate::collab::kernels::aggregate::ACCUM_CLAMP) and put every coefficient far above +/// the noise threshold, so the threshold would stop being tested. Dividing by a constant keeps every +/// 8x8 window exactly as distinct as before. +pub(super) fn unique_frame(width: u32, height: u32) -> Vec { + let raw = make_unique_frame(width, height); let peak = raw.iter().copied().fold(0.0f32, f32::max); - raw.into_iter().map(|v| v / peak).collect() + raw.into_iter().map(|value| value / peak).collect() } -/// Everything one launch of [`collab_fused`] takes, so a test reads as -/// the scenario it sets up rather than as an argument list. +/// Everything one launch of [collab_fused] takes. pub(super) struct Setup { pub(super) ring: Vec, pub(super) mv_field: Vec, @@ -90,18 +77,15 @@ pub(super) struct Setup { pub(super) k_max: u32, pub(super) sigma: f32, pub(super) lambda_ht: f32, - /// Residual correlation the noise profile is built for. `0.0` gives - /// the all-ones profile most runs use. + /// Residual correlation the noise profile is built for. `0.0` gives the all-ones profile. pub(super) rho: f32, - /// A profile buffer supplied outright, bypassing - /// [`dct_noise_profile`]. The weight scale still follows whatever - /// profile is in force. + /// A profile supplied outright, bypassing [dct_noise_profile]. + /// + /// The weight scale still follows whichever profile is in force. pub(super) profile_override: Option<[f32; 8]>, - /// The aggregation window's `beta`. `0.0`, what every run here uses - /// unless it says otherwise, is uniform aggregation. + /// The aggregation window's beta. `0.0` is uniform aggregation. pub(super) kaiser_beta: f32, - /// The frame's noise curve. `None` launches with `curve_valid = 0` - /// and a zeroed buffer. + /// The frame's noise curve. `None` launches with `curve_valid = 0` and a zeroed buffer. pub(super) noise_curve: Option<[f32; NOISE_CURVE_BINS]>, /// A strength map and the mode it applies in. `None` launches a unit map with the map off. pub(super) strength_map: Option<(Vec, u32)>, @@ -113,12 +97,13 @@ pub(super) struct Setup { } impl Setup { - /// A single-frame ring with no neighbours, which leaves every - /// candidate in the spatial window around the reference patch and - /// the motion, confidence, and neighbour-slot buffers as dummies - /// nothing reads. + /// A single-frame ring with no neighbours. + /// + /// Every candidate comes from the spatial window, and the motion, confidence and neighbour-slot + /// buffers are dummies nothing reads. pub(super) fn spatial_only(frame: Vec, width: u32, height: u32) -> Self { assert_eq!(frame.len(), (width * height) as usize); + Setup { ring: frame, mv_field: vec![0i32, 0i32], @@ -148,8 +133,7 @@ impl Setup { } } - /// Ring slots in this setup's frame ring, which is also how many - /// regions the accumulators carry. + /// Slots in the frame ring, which is also how many regions the accumulators carry. pub(super) fn frames(&self) -> u32 { let frame_len = self.width * self.height * self.stored_channels(); self.ring.len() as u32 / frame_len @@ -184,86 +168,78 @@ pub(super) struct Aggregated { } impl Aggregated { - /// One finished pixel, the weighted mean of every filtered patch - /// that covered it. - /// - /// This is what [`crate::collab::kernels::aggregate::collab_normalise`] - /// computes and what the caller actually sees, so a tolerance stated - /// against it is a tolerance in pixel values. Comparing the raw - /// accumulator instead would fail on a group-weight difference that - /// the division cancels out. + /// One finished pixel, the weighted mean of every filtered patch that covered it. /// - /// A pixel no member covered has a zero weight sum and reads zero. + /// This is what the caller sees, so a tolerance against it is in pixel values. Comparing the raw + /// accumulator instead would fail on a group-weight difference the division cancels. A pixel no + /// member covered reads zero. pub(super) fn pixel(&self, idx: usize) -> f64 { - let w = self.wsum[idx]; - if w == 0 { + let weight_sum = self.wsum[idx]; + if weight_sum == 0 { 0.0 } else { - // `wsum` counts at `WEIGHT_GAIN` times `accum`'s scale, the - // one factor that does not cancel between the two, exactly as - // `collab_normalise` multiplies it back out. - self.accum[idx] as f64 * WEIGHT_GAIN as f64 / w as f64 + // `wsum` counts at `WEIGHT_GAIN` times `accum`'s scale, the one factor that does not + // cancel, so it is multiplied back out as `collab_normalise` does. + self.accum[idx] as f64 * WEIGHT_GAIN as f64 / weight_sum as f64 } } - /// The total weight one ring slot's region received. A slot no - /// member scattered into reads exactly zero. + /// The total weight one ring slot's region received. pub(super) fn frame_weight_sum(&self, slot: usize) -> i64 { self.wsum[slot * self.pixels..(slot + 1) * self.pixels] .iter() - .map(|&v| v as i64) + .map(|&weight| weight as i64) .sum() } - /// A compact summary of the whole run, small enough to record as - /// literals and specific enough that a kernel writing nothing cannot - /// reproduce it. + /// A summary of the whole run, small enough to record as literals and specific enough that a + /// kernel writing nothing cannot reproduce it. fn digest(&self) -> Digest { - // Luma stores one channel per pixel across this file, so the two - // accumulators hold one entry each per pixel and share an index. + // Luma stores one channel per pixel, so the two accumulators share an index. assert_eq!(self.accum.len(), self.wsum.len()); - let n = self.accum.len(); + let len = self.accum.len(); let mut sum = 0.0f64; let mut sum_sq = 0.0f64; let mut covered = 0usize; - for idx in 0..n { - let v = self.pixel(idx); - sum += v; - sum_sq += v * v; + for idx in 0..len { + let pixel = self.pixel(idx); + sum += pixel; + sum_sq += pixel * pixel; if self.wsum[idx] != 0 { covered += 1; } } - let weight_mean = - self.group_weight.iter().map(|&w| w as f64).sum::() / self.group_weight.len() as f64; + + let weight_total = self.group_weight.iter().map(|&weight| weight as f64).sum::(); + let weight_mean = weight_total / self.group_weight.len() as f64; let mut probes = [0.0f64; PROBE_COUNT]; for (i, probe) in probes.iter_mut().enumerate() { - *probe = self.pixel(probe_index(i, n)); + let index = probe_index(i, len); + *probe = self.pixel(index); } Digest { covered, - pixel_mean: sum / n as f64, - pixel_rms: (sum_sq / n as f64).sqrt(), + pixel_mean: sum / len as f64, + pixel_rms: (sum_sq / len as f64).sqrt(), weight_mean, probes, } } } -/// How many individual pixels a [`Digest`] pins alongside its whole-run -/// statistics. +/// How many individual pixels a [Digest] pins alongside its whole-run statistics. const PROBE_COUNT: usize = 8; -/// The pixel a probe reads. The odd stride spreads the eight probes over -/// the buffer so no two land in one patch or one row. -pub(super) fn probe_index(i: usize, len: usize) -> usize { - (i * 7919 + 1013) % len +/// The pixel a probe reads. +/// +/// The odd stride spreads the eight probes over the buffer so no two land in one patch or one row. +pub(super) fn probe_index(probe: usize, len: usize) -> usize { + (probe * 7919 + 1013) % len } -/// One run's output, boiled down to numbers a test can carry as -/// literals. +/// One run's output, boiled down to numbers a test can carry as literals. pub(super) struct Digest { /// Pixels whose weight sum is non-zero. pub(super) covered: usize, @@ -273,67 +249,46 @@ pub(super) struct Digest { pub(super) pixel_rms: f64, /// Mean of the per-reference group weight. pub(super) weight_mean: f64, - /// Individual pixels at [`probe_index`] positions. + /// Individual pixels at [probe_index] positions. pub(super) probes: [f64; PROBE_COUNT], } /// How far a recorded whole-run statistic may move, relative. /// -/// Each of these sums thousands of values, so a single coefficient -/// falling the other side of the hard threshold moves one by around -/// `1e-8`. -/// -/// The literals below were recorded from an implementation that -/// truncated toward zero on the way into the accumulators, which biased -/// every contribution down by up to a fixed-point unit. -/// [`crate::collab::kernels::aggregate::to_fixed`] rounds instead, so the -/// values it produces sit about `1e-5` relative above the recorded ones. -/// That is the quantisation step itself moving, not the filter, and no -/// implementation can match across it more tightly than this. Re-recording -/// from the fused kernel would be worse than loosening, because these -/// literals are a second implementation's answer and matching the kernel -/// against itself would prove nothing. -/// -/// `2e-5` is still vanishingly small next to the difference a kernel that -/// stopped writing would produce. +/// Each statistic sums thousands of values, so one coefficient crossing the hard threshold moves it +/// by around `1e-8`. The recorded literals carry a truncate-toward-zero bias of up to one fixed-point +/// unit per contribution, while [to_fixed](crate::collab::kernels::aggregate::to_fixed) rounds, so +/// live values sit about `1e-5` relative above them. That is the quantisation step moving, not the +/// filter, and no implementation can match across it more tightly. Re-recording from the fused +/// kernel would prove nothing, since the literals are a second implementation's answer. `2e-5` is +/// still vanishingly small next to what a kernel that stopped writing would produce. const DIGEST_RELATIVE_TOLERANCE: f64 = 2.0e-5; /// How far a recorded probe pixel may move, absolute. /// -/// The hard threshold is a discontinuity, and a coefficient whose -/// magnitude sits within float rounding of `lambda_ht * sigma` can fall -/// either way. One such coefficient moves its group's reconstruction by -/// its own magnitude, and a probe reads one pixel rather than an -/// average, so this is the same `1e-3` (a quarter of an 8-bit code -/// level) the differential these literals were recorded from allowed. +/// A coefficient within float rounding of `lambda_ht * sigma` can fall either side of the hard +/// threshold and moves its group's reconstruction by its own magnitude. A probe reads one pixel +/// rather than an average, so it allows `1e-3`, a quarter of an 8-bit code level. const PROBE_TOLERANCE: f64 = 1.0e-3; -/// Checks a run against values recorded from a known-good -/// implementation. +/// Checks a run against values recorded from a second, known-good implementation. /// -/// Every expected value below was produced by -/// `collab_group_temporal` + `collab_filter_ht`, the two-kernel pair the -/// fused kernel replaces, on 2026-08-21, immediately before that pair -/// was deleted. The two agreed to `5e-9` on the whole-run statistics and -/// `5e-7` on the worst probe at the time of recording. -/// -/// Fixed literals rather than a second kernel is what keeps this -/// meaningful. A cubecl 0.10 compiler bug makes a failing shader -/// compile silently do nothing at all, leaving the buffers untouched, -/// and a test that compared the fused kernel against itself would have -/// compared zeros to zeros. Zeros do not match these. +/// The literals come from a two-kernel group-then-filter implementation, which agreed with the fused +/// kernel to `5e-9` on the whole-run statistics and `5e-7` on the worst probe when recorded. +/// Fixed literals keep this meaningful because a cubecl compiler bug can make a failing shader do +/// nothing at all, and a kernel compared against itself would then match zeros to zeros. pub(super) fn assert_matches_recorded(label: &str, got: &Aggregated, want: &Digest) { - let d = got.digest(); + let digest = got.digest(); assert_eq!( - d.covered, want.covered, + digest.covered, want.covered, "{label}: {} pixels carry weight, recorded {}", - d.covered, want.covered + digest.covered, want.covered ); for (name, have, expect) in [ - ("pixel_mean", d.pixel_mean, want.pixel_mean), - ("pixel_rms", d.pixel_rms, want.pixel_rms), - ("weight_mean", d.weight_mean, want.weight_mean), + ("pixel_mean", digest.pixel_mean, want.pixel_mean), + ("pixel_rms", digest.pixel_rms, want.pixel_rms), + ("weight_mean", digest.weight_mean, want.weight_mean), ] { let rel = (have - expect).abs() / expect.abs().max(1.0e-30); assert!( @@ -342,7 +297,7 @@ pub(super) fn assert_matches_recorded(label: &str, got: &Aggregated, want: &Dige ); } - for (i, (&have, &expect)) in d.probes.iter().zip(want.probes.iter()).enumerate() { + for (i, (&have, &expect)) in digest.probes.iter().zip(want.probes.iter()).enumerate() { assert!( (have - expect).abs() < PROBE_TOLERANCE, "{label}: probe {i} is {have}, recorded {expect}" @@ -370,35 +325,58 @@ pub(super) struct Buffers { refs_y: u32, } -pub(super) fn buffers(s: &Setup) -> Buffers { +pub(super) fn buffers(setup: &Setup) -> Buffers { let client = make_client(); - let refs_x = refs_along(s.width); - let refs_y = refs_along(s.height); - let refs = ref_count(s.width, s.height); - let frames = s.frames() as usize; - let stored_ch = s.stored_channels() as usize; - let accum_len = s.pixels() * stored_ch * frames; - let wsum_len = s.pixels() * frames; - - // Padding lanes past the live channels carry a zero sigma, as the - // denoiser uploads them. - let mut sigma = vec![0.0f32; stored_ch]; - let live_channels = s.channel_mode.count() as usize; - sigma[..live_channels].fill(s.sigma); + let refs_x = refs_along(setup.width); + let refs_y = refs_along(setup.height); + let refs = ref_count(setup.width, setup.height); + let frames = setup.frames() as usize; + let stored_channels = setup.stored_channels() as usize; + let accum_len = setup.pixels() * stored_channels * frames; + let wsum_len = setup.pixels() * frames; + + // Padding lanes past the live channels carry a zero sigma, as the denoiser uploads them. + let mut sigma = vec![0.0f32; stored_channels]; + let live_channels = setup.channel_mode.count() as usize; + sigma[..live_channels].fill(setup.sigma); + + let profile = setup.profile(); + let kaiser = kaiser_window(setup.kaiser_beta); + // Zeroed here rather than by `collab_zero_accum`, since the scatter is the only writer in these runs. + let zeroed_accum = vec![0i32; accum_len]; + let zeroed_wsum = vec![0i32; wsum_len]; + + let ring_bytes = f32::as_bytes(&setup.ring); + let mv_bytes = i32::as_bytes(&setup.mv_field); + let conf_bytes = f32::as_bytes(&setup.confidence); + let slots_bytes = u32::as_bytes(&setup.neighbour_slots); + let sigma_bytes = f32::as_bytes(&sigma); + let profile_bytes = f32::as_bytes(&profile); + let kaiser_bytes = f32::as_bytes(&kaiser); + let accum_bytes = i32::as_bytes(&zeroed_accum); + let wsum_bytes = i32::as_bytes(&zeroed_wsum); + let ring = client.create_from_slice(ring_bytes); + let mv_field = client.create_from_slice(mv_bytes); + let confidence = client.create_from_slice(conf_bytes); + let neighbour_slots = client.create_from_slice(slots_bytes); + let sigma = client.create_from_slice(sigma_bytes); + let dct_profile = client.create_from_slice(profile_bytes); + let kaiser = client.create_from_slice(kaiser_bytes); + let accum = client.create_from_slice(accum_bytes); + let wsum = client.create_from_slice(wsum_bytes); + let group_weight = client.empty(refs * size_of::()); Buffers { - ring: client.create_from_slice(f32::as_bytes(&s.ring)), - mv_field: client.create_from_slice(i32::as_bytes(&s.mv_field)), - confidence: client.create_from_slice(f32::as_bytes(&s.confidence)), - neighbour_slots: client.create_from_slice(u32::as_bytes(&s.neighbour_slots)), - sigma: client.create_from_slice(f32::as_bytes(&sigma)), - dct_profile: client.create_from_slice(f32::as_bytes(&s.profile())), - kaiser: client.create_from_slice(f32::as_bytes(&kaiser_window(s.kaiser_beta))), - // Zeroed here rather than by `collab_zero_accum`, since the - // scatter is the only thing writing them in these runs. - accum: client.create_from_slice(i32::as_bytes(&vec![0i32; accum_len])), - wsum: client.create_from_slice(i32::as_bytes(&vec![0i32; wsum_len])), - group_weight: client.empty(refs * size_of::()), + ring, + mv_field, + confidence, + neighbour_slots, + sigma, + dct_profile, + kaiser, + accum, + wsum, + group_weight, accum_len, wsum_len, refs, @@ -408,43 +386,47 @@ pub(super) fn buffers(s: &Setup) -> Buffers { } } -pub(super) fn read_back(b: Buffers, s: &Setup) -> Aggregated { - let accum = b.client.read_one(b.accum).expect("accum readback failed"); - let wsum = b.client.read_one(b.wsum).expect("wsum readback failed"); - let group_weight = b +pub(super) fn read_back(handles: Buffers, setup: &Setup) -> Aggregated { + let accum_bytes = handles .client - .read_one(b.group_weight) + .read_one(handles.accum) + .expect("accum readback failed"); + let wsum_bytes = handles + .client + .read_one(handles.wsum) + .expect("wsum readback failed"); + let weight_bytes = handles + .client + .read_one(handles.group_weight) .expect("group_weight readback failed"); Aggregated { - accum: i32::from_bytes(&accum)[..b.accum_len].to_vec(), - wsum: i32::from_bytes(&wsum)[..b.wsum_len].to_vec(), - group_weight: f32::from_bytes(&group_weight)[..b.refs].to_vec(), - pixels: s.pixels(), + accum: i32::from_bytes(&accum_bytes)[..handles.accum_len].to_vec(), + wsum: i32::from_bytes(&wsum_bytes)[..handles.wsum_len].to_vec(), + group_weight: f32::from_bytes(&weight_bytes)[..handles.refs].to_vec(), + pixels: setup.pixels(), } } -/// Launches [`collab_fused`] on its eight-references-per-cube grid and -/// reads back what it aggregated. +/// Launches [collab_fused] on its eight-references-per-cube grid and reads back what it aggregated. /// -/// The search walk is whichever one this runtime needs, so a plain run -/// covers whatever the shipping code would actually launch here. -pub(super) fn run_fused(s: &Setup) -> Aggregated { - run_fused_walk(s, None) +/// The search walk is whichever one this runtime needs, matching what the shipping code launches. +pub(super) fn run_fused(setup: &Setup) -> Aggregated { + run_fused_walk(setup, None) } -/// [`run_fused`] with the search walk pinned rather than taken from the -/// runtime, so one test can run both and compare them. -pub(super) fn run_fused_walk(s: &Setup, warp_uniform: Option) -> Aggregated { - let b = buffers(s); - let profile = s.profile(); - let curve = s.noise_curve.unwrap_or([0.0f32; NOISE_CURVE_BINS]); - let curve_buf = b.client.create_from_slice(f32::as_bytes(&curve)); - let curve_valid = u32::from(s.noise_curve.is_some()); +/// [run_fused] with the search walk pinned rather than taken from the runtime. +pub(super) fn run_fused_walk(setup: &Setup, warp_uniform: Option) -> Aggregated { + let handles = buffers(setup); + let profile = setup.profile(); + let curve = setup.noise_curve.unwrap_or([0.0f32; NOISE_CURVE_BINS]); + let curve_bytes = f32::as_bytes(&curve); + let curve_buf = handles.client.create_from_slice(curve_bytes); + let curve_valid = u32::from(setup.noise_curve.is_some()); - let (map_cols, map_rows) = strength_map_dims(s.width, s.height); + let (map_cols, map_rows) = strength_map_dims(setup.width, setup.height); let map_len = (map_cols * map_rows) as usize; - let (map_values, map_mode) = match &s.strength_map { + let (map_values, map_mode) = match &setup.strength_map { Some((values, mode)) => (values.clone(), *mode), None => (vec![1.0f32; map_len], STRENGTH_MAP_OFF), }; @@ -453,71 +435,77 @@ pub(super) fn run_fused_walk(s: &Setup, warp_uniform: Option) -> Aggregate map_len, "a strength map must cover the frame's quarter grid" ); - let map_buf = b.client.create_from_slice(f32::as_bytes(&map_values)); - let stored_ch = s.stored_channels(); + let map_bytes = f32::as_bytes(&map_values); + let map_buf = handles.client.create_from_slice(map_bytes); + let stored_channels = setup.stored_channels(); + + let cubes_x = fused_cubes_x(setup.width); + let grid = CubeCount::new_2d(cubes_x, handles.refs_y); + let dim = CubeDim::new_1d(64); + let scale = weight_scale(setup.sigma, &profile); + let accum_scale = setup.accum_scale(); + let warp_uniform = warp_uniform.unwrap_or_else(|| needs_warp_uniform_search(&handles.client)); + let grid_frame_count = grid_frames(setup.radius); + let pooled_ratio = setup.pooled.unwrap_or(0.0); unsafe { collab_fused::launch_unchecked::( - &b.client, - CubeCount::new_2d(fused_cubes_x(s.width), b.refs_y), - CubeDim::new_1d(64), - stored_ch as usize, - ArrayArg::from_raw_parts(b.ring.clone(), s.ring.len()), - ArrayArg::from_raw_parts(b.mv_field.clone(), s.mv_field.len()), - ArrayArg::from_raw_parts(b.confidence.clone(), s.confidence.len()), - ArrayArg::from_raw_parts(b.neighbour_slots.clone(), s.neighbour_slots.len()), - ArrayArg::from_raw_parts(b.sigma.clone(), stored_ch as usize), + &handles.client, + grid, + dim, + stored_channels as usize, + ArrayArg::from_raw_parts(handles.ring.clone(), setup.ring.len()), + ArrayArg::from_raw_parts(handles.mv_field.clone(), setup.mv_field.len()), + ArrayArg::from_raw_parts(handles.confidence.clone(), setup.confidence.len()), + ArrayArg::from_raw_parts(handles.neighbour_slots.clone(), setup.neighbour_slots.len()), + ArrayArg::from_raw_parts(handles.sigma.clone(), stored_channels as usize), ArrayArg::from_raw_parts(curve_buf, NOISE_CURVE_BINS), ArrayArg::from_raw_parts(map_buf, map_len), - ArrayArg::from_raw_parts(b.dct_profile.clone(), 8), - ArrayArg::from_raw_parts(b.kaiser.clone(), PATCH_SIZE as usize), - ArrayArg::from_raw_parts(b.accum.clone(), b.accum_len), - ArrayArg::from_raw_parts(b.wsum.clone(), b.wsum_len), - ArrayArg::from_raw_parts(b.group_weight.clone(), b.refs), - s.centre_slot, - s.c_min, - s.lambda_ht, + ArrayArg::from_raw_parts(handles.dct_profile.clone(), 8), + ArrayArg::from_raw_parts(handles.kaiser.clone(), PATCH_SIZE as usize), + ArrayArg::from_raw_parts(handles.accum.clone(), handles.accum_len), + ArrayArg::from_raw_parts(handles.wsum.clone(), handles.wsum_len), + ArrayArg::from_raw_parts(handles.group_weight.clone(), handles.refs), + setup.centre_slot, + setup.c_min, + setup.lambda_ht, curve_valid, map_mode, - weight_scale(s.sigma, &profile), - s.accum_scale(), - warp_uniform.unwrap_or_else(|| needs_warp_uniform_search(&b.client)), - s.radius, - grid_frames(s.radius), - s.refine, - s.mv_stride, - s.conf_stride, + scale, + accum_scale, + warp_uniform, + setup.radius, + grid_frame_count, + setup.refine, + setup.mv_stride, + setup.conf_stride, BLK_STEP, BLKSIZE, - s.blocks_x, - s.blocks_y, - s.width, - s.height, - s.channel_mode.count(), - s.k_max, - stored_ch, - s.spatial_radius, - b.refs_x, + setup.blocks_x, + setup.blocks_y, + setup.width, + setup.height, + setup.channel_mode.count(), + setup.k_max, + stored_channels, + setup.spatial_radius, + handles.refs_x, map_cols, map_rows, - s.pooled.unwrap_or(0.0), - s.pooled.is_some(), + pooled_ratio, + setup.pooled.is_some(), ); } - read_back(b, s) + read_back(handles, setup) } -/// A ring of `2 * radius + 1` frames of unique content, with a motion -/// field and a confidence field that both vary by block. -/// -/// The ring is laid out frame-major, exactly as `read_line` indexes it, -/// so one call to `make_unique_frame` over a `2 * radius + 1` times -/// taller image fills the whole ring with content no two 8x8 windows -/// share, across frames as well as within one. +/// A ring of `2 * radius + 1` frames of unique content, with motion and confidence fields that +/// vary by block. /// -/// The confidences run from below `c_min` to 1.0, so some blocks have -/// their whole window skipped and the rest are searched. +/// The ring is frame-major, so one [unique_frame] over a `2 * radius + 1` times taller image gives +/// content no two 8x8 windows share, across frames as well as within one. The confidences run from +/// below `c_min` to 1.0, so some blocks skip their whole window and the rest are searched. pub(super) fn cross_frame_setup(width: u32, height: u32, radius: u32) -> Setup { let frames = 2 * radius + 1; let blocks_x = width.div_ceil(BLK_STEP); @@ -529,17 +517,17 @@ pub(super) fn cross_frame_setup(width: u32, height: u32, radius: u32) -> Setup { let mut confidence = vec![0.0f32; (2 * radius * conf_stride) as usize]; for t in 0..(2 * radius) { for block in 0..conf_stride { - let mv = (t * mv_stride + block * 2) as usize; - // A spread of shifts in both signs, including some that push - // the refine window off the frame so the clip matters. - mv_field[mv] = (block % 11) as i32 - 5 + t as i32; - mv_field[mv + 1] = 4 - (block % 9) as i32 - t as i32; + let mv_index = (t * mv_stride + block * 2) as usize; + // A spread of shifts in both signs, some pushing the refine window off the frame so the + // clip matters. + mv_field[mv_index] = (block % 11) as i32 - 5 + t as i32; + mv_field[mv_index + 1] = 4 - (block % 9) as i32 - t as i32; confidence[(t * conf_stride + block) as usize] = ((block * 7 + t * 3) % 11) as f32 / 10.0; } } - // The centre sits in the middle of the ring, and the neighbours are - // the slots either side of it, nearest first. + // The centre sits in the middle of the ring, and the neighbours are the slots either side of it, + // nearest first. let centre_slot = radius; let mut neighbour_slots = Vec::new(); for t in 0..radius { @@ -547,8 +535,11 @@ pub(super) fn cross_frame_setup(width: u32, height: u32, radius: u32) -> Setup { neighbour_slots.push(radius + 1 + t); } + let ring = unique_frame(width, height * frames); + let placeholder = vec![0.0f32; (width * height) as usize]; + Setup { - ring: unique_frame(width, height * frames), + ring, mv_field, confidence, neighbour_slots, @@ -560,19 +551,17 @@ pub(super) fn cross_frame_setup(width: u32, height: u32, radius: u32) -> Setup { conf_stride, blocks_x, blocks_y, - ..Setup::spatial_only(vec![0.0f32; (width * height) as usize], width, height) + ..Setup::spatial_only(placeholder, width, height) } } -/// A three-frame ring whose neighbours hold an exact copy of the centre -/// frame, at the position the zero motion field predicts. +/// A three-frame ring whose neighbours hold an exact copy of the centre frame, at the position the +/// zero motion field predicts. /// -/// An exact copy scores distance zero, which every other candidate on -/// this content loses to. At radius 1 the group is four volumes of two -/// frames, and each volume's second frame is neighbour 0, which wins -/// every tie. `refine = 0` narrows each neighbour's rectangle to that -/// one predicted position, so there is nothing else in a neighbour for -/// a volume to pick instead. +/// An exact copy scores distance zero, which every other candidate loses to. At radius 1 the group +/// is four volumes of two frames, and each volume's second frame is neighbour 0, which wins every +/// tie. `refine = 0` narrows each neighbour to that one predicted position, so a volume has nothing +/// else to pick. pub(super) fn three_frame_ring_with_a_planted_match(width: u32, height: u32) -> Setup { let frame = unique_frame(width, height); let mut ring = Vec::with_capacity(frame.len() * 3); @@ -584,6 +573,7 @@ pub(super) fn three_frame_ring_with_a_planted_match(width: u32, height: u32) -> let blocks_y = height.div_ceil(BLK_STEP); let conf_stride = blocks_x * blocks_y; let mv_stride = conf_stride * 2; + let placeholder = vec![0.0f32; (width * height) as usize]; Setup { ring, @@ -597,7 +587,7 @@ pub(super) fn three_frame_ring_with_a_planted_match(width: u32, height: u32) -> conf_stride, blocks_x, blocks_y, - ..Setup::spatial_only(vec![0.0f32; (width * height) as usize], width, height) + ..Setup::spatial_only(placeholder, width, height) } } @@ -609,7 +599,7 @@ pub(super) fn three_frame_ring_with_a_planted_match(width: u32, height: u32) -> pub(super) fn five_frame_ring_with_jittered_copies(width: u32, height: u32) -> Setup { let radius = 2u32; let frame = unique_frame(width, height); - let jitter = noisy_field_over(width, height * 4, 0.5, 0.002); + let jitter = noisy_flat_field(width, height * 4, 0.5, 0.002); let pixels = (width * height) as usize; let mut ring = Vec::with_capacity(pixels * 5); @@ -632,6 +622,7 @@ pub(super) fn five_frame_ring_with_jittered_copies(width: u32, height: u32) -> S let blocks_y = height.div_ceil(BLK_STEP); let conf_stride = blocks_x * blocks_y; let mv_stride = conf_stride * 2; + let placeholder = vec![0.0f32; pixels]; Setup { ring, @@ -645,58 +636,59 @@ pub(super) fn five_frame_ring_with_jittered_copies(width: u32, height: u32) -> S conf_stride, blocks_x, blocks_y, - ..Setup::spatial_only(vec![0.0f32; pixels], width, height) + ..Setup::spatial_only(placeholder, width, height) } } -/// How many reference patches cover each pixel of a `width` by `height` -/// frame, on the same grid [`ref_pos`] lays out. +/// How many reference patches cover each pixel, on the same grid [ref_pos] lays out. pub(super) fn reference_cover_counts(width: u32, height: u32) -> Vec { let mut counts = vec![0i64; (width * height) as usize]; - for ry in 0..refs_along(height) { - for rx in 0..refs_along(width) { - let px = ref_pos(rx, width); - let py = ref_pos(ry, height); + for ref_y in 0..refs_along(height) { + for ref_x in 0..refs_along(width) { + let left = ref_pos(ref_x, width); + let top = ref_pos(ref_y, height); for row in 0..PATCH_SIZE { for col in 0..PATCH_SIZE { - counts[((py + row) * width + px + col) as usize] += 1; + counts[((top + row) * width + left + col) as usize] += 1; } } } } + counts } -pub(super) fn patch_pool_variance(frame: &[f32], w: u32, h: u32) -> f64 { +pub(super) fn patch_pool_variance(frame: &[f32], width: u32, height: u32) -> f64 { let mut pool: Vec = Vec::new(); - for ry in 0..refs_along(h) { - for rx in 0..refs_along(w) { - let px = ref_pos(rx, w); - let py = ref_pos(ry, h); + for ref_y in 0..refs_along(height) { + for ref_x in 0..refs_along(width) { + let left = ref_pos(ref_x, width); + let top = ref_pos(ref_y, height); for row in 0..PATCH_SIZE { for col in 0..PATCH_SIZE { - pool.push(frame[((py + row) * w + px + col) as usize] as f64); + pool.push(frame[((top + row) * width + left + col) as usize] as f64); } } } } + let mean = pool.iter().sum::() / pool.len() as f64; - pool.iter().map(|v| (v - mean).powi(2)).sum::() / pool.len() as f64 + pool.iter().map(|value| (value - mean).powi(2)).sum::() / pool.len() as f64 } -/// The variance of a run's finished pixels. pub(super) fn output_variance(got: &Aggregated) -> f64 { let values: Vec = (0..got.accum.len()).map(|i| got.pixel(i)).collect(); let mean = values.iter().sum::() / values.len() as f64; - values.iter().map(|v| (v - mean).powi(2)).sum::() / values.len() as f64 + values.iter().map(|value| (value - mean).powi(2)).sum::() / values.len() as f64 } -/// A flat field carrying nothing but noise, at the settings a real -/// caller would filter it with. -pub(super) fn flat_noise_setup(w: u32, h: u32, sigma: f32) -> Setup { - let mut s = Setup::spatial_only(noisy_field_over(w, h, 0.5, sigma), w, h); - s.spatial_radius = 9; - s.sigma = sigma; - s.lambda_ht = 2.7; - s +/// A flat field carrying nothing but noise, at the settings a real caller would filter it with. +pub(super) fn flat_noise_setup(width: u32, height: u32, sigma: f32) -> Setup { + let field = noisy_flat_field(width, height, 0.5, sigma); + let mut setup = Setup::spatial_only(field, width, height); + setup.spatial_radius = 9; + setup.sigma = sigma; + setup.lambda_ht = 2.7; + + setup } diff --git a/av-denoise-core/src/collab/tests/fused/noise_curve.rs b/av-denoise-core/src/collab/tests/fused/noise_curve.rs index 21fddcd..4b1f191 100644 --- a/av-denoise-core/src/collab/tests/fused/noise_curve.rs +++ b/av-denoise-core/src/collab/tests/fused/noise_curve.rs @@ -1,5 +1,5 @@ use super::{Aggregated, Setup, cross_frame_setup, run_fused}; -use crate::collab::tests::helpers::noisy_field_over; +use crate::collab::tests::helpers::noisy_flat_field; use crate::nlmeans::NOISE_CURVE_BINS; /// The side of the square frame the stepped-curve test filters. @@ -15,6 +15,7 @@ const CURVE_LAMBDA: f32 = 1.0; pub(super) fn stepped_curve() -> [f32; NOISE_CURVE_BINS] { let mut curve = [2.0f32; NOISE_CURVE_BINS]; curve[NOISE_CURVE_BINS / 2..].fill(0.5); + curve } @@ -32,9 +33,9 @@ pub(super) fn assert_identical(label: &str, got: &Aggregated, want: &Aggregated) } fn run_with_curve(curve: Option<[f32; NOISE_CURVE_BINS]>) -> Aggregated { - let mut s = cross_frame_setup(64, 64, 2); - s.noise_curve = curve; - run_fused(&s) + let mut setup = cross_frame_setup(64, 64, 2); + setup.noise_curve = curve; + run_fused(&setup) } /// Asserts every pixel in columns `x_start..x_end` of a single-frame run matches exactly. @@ -109,7 +110,7 @@ fn a_flat_curve_equals_scaling_lambda() { fn a_stepped_curve_thresholds_each_brightness_by_its_own_noise() { let side = STEP_FRAME_SIDE; let half = side / 2; - let noise = noisy_field_over(side, side, 0.0, 0.02); + let noise = noisy_flat_field(side, side, 0.0, 0.02); let mut frame = Vec::with_capacity(noise.len()); for (idx, sample) in noise.iter().enumerate() { let x = idx as u32 % side; @@ -119,7 +120,8 @@ fn a_stepped_curve_thresholds_each_brightness_by_its_own_noise() { let mut curved_setup = Setup::spatial_only(frame.clone(), side, side); curved_setup.lambda_ht = CURVE_LAMBDA; - curved_setup.noise_curve = Some(stepped_curve()); + let curve = stepped_curve(); + curved_setup.noise_curve = Some(curve); let curved = run_fused(&curved_setup); let plain = run_spatial_with_lambda(&frame, side, CURVE_LAMBDA); @@ -143,7 +145,7 @@ fn a_stepped_curve_thresholds_each_brightness_by_its_own_noise() { #[test] fn the_curve_is_sampled_at_bin_centres() { let side = STEP_FRAME_SIDE; - let frame = noisy_field_over(side, side, 0.25, 0.02); + let frame = noisy_flat_field(side, side, 0.25, 0.02); let mut curve = [0.33f32; NOISE_CURVE_BINS]; curve[3] = 2.0; diff --git a/av-denoise-core/src/collab/tests/fused/pooled.rs b/av-denoise-core/src/collab/tests/fused/pooled.rs index a84abd3..f066856 100644 --- a/av-denoise-core/src/collab/tests/fused/pooled.rs +++ b/av-denoise-core/src/collab/tests/fused/pooled.rs @@ -94,6 +94,7 @@ fn run_kernel( } } } + (result, retained) } @@ -152,6 +153,7 @@ fn reference( } } } + (result, retained) } @@ -174,6 +176,7 @@ fn seeded_group(seed: u32) -> Group { } } } + group } @@ -193,6 +196,7 @@ const UNIT: [f32; SIDE] = [1.0; SIDE]; fn the_kernel_matches_the_host_reference() { let profile = [1.3, 1.1, 1.0, 0.95, 0.9, 0.9, 0.9, 0.95]; let variances = [1.0, 1.0, 1.0, 1.0, 1.2, 1.2, 1.2, 1.2]; + for (seed, k_use) in [(1u32, 8u32), (2, 8), (3, 4), (4, 2)] { let group = seeded_group(seed); let (got, got_retained) = run_kernel(&group, &variances, &profile, k_use, 1.6, 2.7); @@ -327,24 +331,25 @@ fn both_walks_agree_with_pooling_on() { #[test] fn a_ragged_frame_with_pooling_on_completes_and_both_walks_agree() { - let mut setup = Setup::spatial_only(unique_frame(70, 54), 70, 54); + let frame = unique_frame(70, 54); + let mut setup = Setup::spatial_only(frame, 70, 54); setup.pooled = Some(RATIO); let divergent = run_fused_walk(&setup, Some(false)); let uniform = run_fused_walk(&setup, Some(true)); + let frame_sum = divergent.frame_weight_sum(0); assert_eq!(divergent.accum, uniform.accum); - assert!( - divergent.frame_weight_sum(0) > 0, - "the frame should receive weight" - ); + assert!(frame_sum > 0, "the frame should receive weight"); + let covered = (0..setup.pixels()).all(|idx| divergent.wsum[idx] > 0); assert!(covered, "every pixel of a ragged frame should be covered"); } #[test] fn a_small_fallback_group_with_pooling_on_stays_finite() { - let mut setup = Setup::spatial_only(unique_frame(64, 48), 64, 48); + let frame = unique_frame(64, 48); + let mut setup = Setup::spatial_only(frame, 64, 48); setup.k_max = 4; setup.pooled = Some(RATIO); @@ -360,13 +365,15 @@ fn a_small_fallback_group_with_pooling_on_stays_finite() { #[test] fn a_noiseless_flat_frame_passes_through_with_pooling_on() { - let mut setup = Setup::spatial_only(vec![0.4f32; 64 * 48], 64, 48); + let frame = vec![0.4f32; 64 * 48]; + let mut setup = Setup::spatial_only(frame, 64, 48); setup.sigma = 1.0e-6; setup.pooled = Some(RATIO); let pooled = run_fused(&setup); assert!(pooled.group_weight.iter().all(|weight| weight.is_finite())); + for idx in 0..setup.pixels() { let value = pooled.pixel(idx); assert!((value - 0.4).abs() < 1.0e-3, "pixel {idx} is {value}"); diff --git a/av-denoise-core/src/collab/tests/fused/recorded.rs b/av-denoise-core/src/collab/tests/fused/recorded.rs index 31f5eed..b05bd89 100644 --- a/av-denoise-core/src/collab/tests/fused/recorded.rs +++ b/av-denoise-core/src/collab/tests/fused/recorded.rs @@ -7,295 +7,257 @@ use super::{ run_fused, unique_frame, }; -use crate::collab::tests::helpers::noisy_field_over; +use crate::collab::tests::helpers::noisy_flat_field; -/// Content without ties, so nothing about the result depends on how the -/// insert breaks one. -/// -/// `make_unique_frame` is built so that any two distinct 8x8 windows -/// differ in most of their 64 pixels. +/// Content without ties, so nothing about the result depends on how the insert breaks one. #[test] fn fused_reproduces_recorded_output_on_unique_content() { - let (w, h) = (128u32, 96u32); - let s = Setup::spatial_only(unique_frame(w, h), w, h); - assert_matches_recorded( - "unique content", - &run_fused(&s), - &Digest { - covered: 12288, - pixel_mean: 0.500102660422, - pixel_rms: 0.577471243275, - weight_mean: 1250.000000000, - probes: [ - 0.917905456141, - 0.787302672863, - 0.650475382805, - 0.517024146186, - 0.386674649788, - 0.255359411240, - 0.120508321126, - 0.989773918601, - ], - }, - ); + let (width, height) = (128u32, 96u32); + let frame = unique_frame(width, height); + let setup = Setup::spatial_only(frame, width, height); + let got = run_fused(&setup); + let want = Digest { + covered: 12288, + pixel_mean: 0.500102660422, + pixel_rms: 0.577471243275, + weight_mean: 1250.000000000, + probes: [ + 0.917905456141, + 0.787302672863, + 0.650475382805, + 0.517024146186, + 0.386674649788, + 0.255359411240, + 0.120508321126, + 0.989773918601, + ], + }; + + assert_matches_recorded("unique content", &got, &want); } -/// The same content with its ramp turned on its side. -/// -/// `make_unique_frame` ramps along x, which makes a one-column shift far -/// costlier than a one-row shift, so every member a group keeps sits in -/// a narrow column band and the search rectangle's x extent never -/// decides anything. Transposing the frame moves that band onto the x -/// axis, so this is the run where the horizontal bounds are -/// load-bearing. +/// The unique frame ramps along x, so a one-column shift costs far more than a one-row shift and +/// the search rectangle's x extent never decides anything. Transposing it makes the horizontal +/// bounds load-bearing. #[test] fn fused_reproduces_recorded_output_on_a_transposed_ramp() { - let (w, h) = (128u32, 96u32); - let source = unique_frame(h, w); - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - frame[(y * w + x) as usize] = source[(x * h + y) as usize]; + let (width, height) = (128u32, 96u32); + let source = unique_frame(height, width); + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + frame[(y * width + x) as usize] = source[(x * height + y) as usize]; } } - let s = Setup::spatial_only(frame, w, h); - assert_matches_recorded( - "transposed ramp", - &run_fused(&s), - &Digest { - covered: 12288, - pixel_mean: 0.500026936557, - pixel_rms: 0.577366054447, - weight_mean: 1238.896681776, - probes: [ - 0.075242505755, - 0.724634047477, - 0.367116374354, - 0.014184951782, - 0.659262769363, - 0.307459000618, - 0.951007338131, - 0.589665272066, - ], - }, - ); + + let setup = Setup::spatial_only(frame, width, height); + let got = run_fused(&setup); + let want = Digest { + covered: 12288, + pixel_mean: 0.500026936557, + pixel_rms: 0.577366054447, + weight_mean: 1238.896681776, + probes: [ + 0.075242505755, + 0.724634047477, + 0.367116374354, + 0.014184951782, + 0.659262769363, + 0.307459000618, + 0.951007338131, + 0.589665272066, + ], + }; + + assert_matches_recorded("transposed ramp", &got, &want); } -/// A width whose reference count is not a multiple of 8 leaves the last -/// cube of each row partly out of range. +/// 104 pixels wide gives 25 reference patches at `STEP = 4`, so the fourth cube of each row runs one +/// live group and seven dead ones. /// -/// 104 pixels wide gives 25 reference patches at `STEP = 4`, so the -/// fourth cube runs one live group and seven dead ones. Those seven must -/// reach every barrier and write nothing. A dead group that scattered -/// would double the last reference's contribution, which shows up here -/// as a moved pixel rather than needing its own assertion. +/// The dead groups must reach every barrier and write nothing. One that scattered would double the +/// last reference's contribution and move a pixel. #[test] fn fused_reproduces_recorded_output_when_refs_are_not_a_multiple_of_eight() { - let (w, h) = (104u32, 96u32); - let s = Setup::spatial_only(unique_frame(w, h), w, h); - assert_matches_recorded( - "ragged reference row", - &run_fused(&s), - &Digest { - covered: 9984, - pixel_mean: 0.500121022858, - pixel_rms: 0.577469311374, - weight_mean: 1240.579711065, - probes: [ - 0.745903455294, - 0.891958951950, - 0.032306798299, - 0.177212221869, - 0.322843606131, - 0.468612211367, - 0.608005691977, - 0.754309082031, - ], - }, - ); + let (width, height) = (104u32, 96u32); + let frame = unique_frame(width, height); + let setup = Setup::spatial_only(frame, width, height); + let got = run_fused(&setup); + let want = Digest { + covered: 9984, + pixel_mean: 0.500121022858, + pixel_rms: 0.577469311374, + weight_mean: 1240.579711065, + probes: [ + 0.745903455294, + 0.891958951950, + 0.032306798299, + 0.177212221869, + 0.322843606131, + 0.468612211367, + 0.608005691977, + 0.754309082031, + ], + }; + + assert_matches_recorded("ragged reference row", &got, &want); } -/// The stack transform's shorter ladders, on a search space too small to -/// fill a group. +/// At `spatial_radius = 1` a corner reference sees 2x2 candidates and an edge one 2x3, both keeping +/// four, while an interior one sees 3x3 and fills to eight. /// -/// At `spatial_radius = 1` a corner reference sees a 2x2 rectangle, so -/// four candidates and a group of four, an edge reference sees 2x3 and -/// also keeps four, and an interior one sees 3x3 and fills to eight. -/// Every wider configuration fills every group to eight, so this is the -/// run where the 2- and 4-member ladders execute at all. +/// Every wider configuration fills every group to eight, so this is the run where the 2- and +/// 4-member ladders execute. #[test] fn fused_reproduces_recorded_output_on_a_short_search_space() { - let (w, h) = (64u32, 64u32); - let mut s = Setup::spatial_only(unique_frame(w, h), w, h); - s.spatial_radius = 1; - assert_matches_recorded( - "short search space", - &run_fused(&s), - &Digest { - covered: 4096, - pixel_mean: 0.500333883408, - pixel_rms: 0.577536371630, - weight_mean: 768.518540988, - probes: [ - 0.836218530965, - 0.571090123027, - 0.304096429037, - 0.033846737369, - 0.774293684286, - 0.508250150663, - 0.242787978384, - 0.977499961853, - ], - }, - ); + let (width, height) = (64u32, 64u32); + let frame = unique_frame(width, height); + let mut setup = Setup::spatial_only(frame, width, height); + setup.spatial_radius = 1; + let got = run_fused(&setup); + let want = Digest { + covered: 4096, + pixel_mean: 0.500333883408, + pixel_rms: 0.577536371630, + weight_mean: 768.518540988, + probes: [ + 0.836218530965, + 0.571090123027, + 0.304096429037, + 0.033846737369, + 0.774293684286, + 0.508250150663, + 0.242787978384, + 0.977499961853, + ], + }; + + assert_matches_recorded("short search space", &got, &want); } -/// The tie-break path, on content where a great many candidates score -/// the same distance. -/// -/// Noise over a flat field has no ramp to separate the candidates, so -/// this is the run where the self-match sentinel and the first-wins -/// insert decide the member set. +/// Noise over a flat field has no ramp to separate candidates, so the self-match sentinel and the +/// first-wins insert decide the member set. #[test] fn fused_reproduces_recorded_output_on_noise() { - let (w, h) = (64u32, 64u32); - let s = Setup::spatial_only(noisy_field_over(w, h, 0.5, 0.05), w, h); - assert_matches_recorded( - "noise", - &run_fused(&s), - &Digest { - covered: 4096, - pixel_mean: 0.500598531425, - pixel_rms: 0.500781913699, - weight_mean: 168.165750156, - probes: [ - 0.473047106911, - 0.525827771943, - 0.505852930189, - 0.501965226326, - 0.472216666744, - 0.498205827272, - 0.493595121410, - 0.515348414403, - ], - }, - ); + let (width, height) = (64u32, 64u32); + let frame = noisy_flat_field(width, height, 0.5, 0.05); + let setup = Setup::spatial_only(frame, width, height); + let got = run_fused(&setup); + let want = Digest { + covered: 4096, + pixel_mean: 0.500598531425, + pixel_rms: 0.500781913699, + weight_mean: 168.165750156, + probes: [ + 0.473047106911, + 0.525827771943, + 0.505852930189, + 0.501965226326, + 0.472216666744, + 0.498205827272, + 0.493595121410, + 0.515348414403, + ], + }; + + assert_matches_recorded("noise", &got, &want); } -/// A non-zero `rho`, where the correlation profile stops being all ones. +/// The kernel applies the correlation profile at the threshold, while the recording applied it to +/// each member's variance before the ladder. /// -/// The old filter multiplied the profile into each member's variance -/// before the variance ladder ran. The fused kernel multiplies it in at -/// the threshold instead. The ladder only averages and the profile is a -/// constant factor across the stack axis, so the two orders agree in -/// exact arithmetic, and this is the run that says so on a GPU. Every -/// other run here uses `dct_noise_profile(0.0)`, which is all ones and -/// cannot tell the two orders apart. `0.86` is the shipped table's high -/// end. +/// The ladder only averages and the profile is constant along the stack axis, so both orders agree +/// in exact arithmetic. Every other run uses the all-ones `rho = 0` profile, which cannot tell them +/// apart. `0.86` is the shipped table's high end. #[test] fn fused_reproduces_recorded_output_under_correlation_shaping() { - let (w, h) = (64u32, 64u32); - let mut s = Setup::spatial_only(noisy_field_over(w, h, 0.5, 0.05), w, h); - s.rho = 0.86; - assert_matches_recorded( - "correlation shaping", - &run_fused(&s), - &Digest { - covered: 4096, - pixel_mean: 0.500501375321, - pixel_rms: 0.502136468679, - weight_mean: 54.170715162, - probes: [ - 0.439152209001, - 0.572103197408, - 0.475816598569, - 0.509686441252, - 0.405882571403, - 0.498373582524, - 0.493889111273, - 0.547764034977, - ], - }, - ); + let (width, height) = (64u32, 64u32); + let frame = noisy_flat_field(width, height, 0.5, 0.05); + let mut setup = Setup::spatial_only(frame, width, height); + setup.rho = 0.86; + let got = run_fused(&setup); + let want = Digest { + covered: 4096, + pixel_mean: 0.500501375321, + pixel_rms: 0.502136468679, + weight_mean: 54.170715162, + probes: [ + 0.439152209001, + 0.572103197408, + 0.475816598569, + 0.509686441252, + 0.405882571403, + 0.498373582524, + 0.493889111273, + 0.547764034977, + ], + }; + + assert_matches_recorded("correlation shaping", &got, &want); } -/// The whole temporal path at once. It covers the `c_min` skip, the -/// volume grid with its single-frame fallback, and the scatter into each -/// member's own region of the accumulator ring. +/// Covers the `c_min` skip, the volume grid with its single-frame fallback, and the scatter into +/// each member's own region of the ring. /// -/// `cross_frame_setup` gives every block its own vector, so the search -/// reaches positions the corner block alone never pointed at, and its -/// confidences straddle `c_min`, so some groups build a grid and others -/// fall back. -/// -/// The digest comes from this kernel's own output, because no second -/// implementation exists for it. [assert_matches_recorded]'s warning -/// about comparing a kernel to itself is about a silently-broken shader -/// producing zeros, and this recording carries real, non-zero coverage. +/// Every block has its own vector, so the search reaches positions the corner block never pointed +/// at, and the confidences straddle `c_min`, so some groups build a grid and others fall back. No +/// second implementation exists, so the digest is this kernel's own output. It carries real, +/// non-zero coverage, which a silently broken shader cannot reproduce. #[test] fn fused_reproduces_recorded_output_across_frames() { - let s = cross_frame_setup(64, 64, 2); - assert_matches_recorded( - "cross frame", - &run_fused(&s), - &Digest { - covered: 16453, - pixel_mean: 0.441598762473, - pixel_rms: 0.546715666545, - weight_mean: 1045.608471951, - probes: [ - 0.837928771973, - 0.574595237938, - 0.300403234153, - 0.000000000000, - 0.775989927049, - 0.000000000000, - 0.237124125163, - 0.979660034180, - ], - }, - ); + let setup = cross_frame_setup(64, 64, 2); + let got = run_fused(&setup); + let want = Digest { + covered: 16453, + pixel_mean: 0.441598762473, + pixel_rms: 0.546715666545, + weight_mean: 1045.608471951, + probes: [ + 0.837928771973, + 0.574595237938, + 0.300403234153, + 0.000000000000, + 0.775989927049, + 0.000000000000, + 0.237124125163, + 0.979660034180, + ], + }; + + assert_matches_recorded("cross frame", &got, &want); } -/// A group with members in neighbour frames must scatter into those -/// frames' regions of the ring, not collapse onto the centre frame. +/// Each neighbour holds the centre plus its own jitter, so every volume keeps a different three of +/// them and every ring slot receives members somewhere. /// -/// Each of the four neighbours holds the centre plus its own jitter, so -/// every volume keeps a different three of them and every slot of the -/// five-frame ring receives members somewhere. The frame a member came -/// from is never written down between the match and the scatter, which -/// is what makes this easy to lose. -/// -/// The digest comes from this kernel's own output, because no second -/// implementation exists for it. +/// The frame a member came from is never written down between the match and the scatter, which +/// makes it easy to lose. No second implementation exists, so the digest is this kernel's own output. #[test] fn fused_scatters_into_every_member_frame() { - let s = five_frame_ring_with_jittered_copies(64, 64); - let got = run_fused(&s); + let setup = five_frame_ring_with_jittered_copies(64, 64); + let got = run_fused(&setup); + for slot in 0..5 { - assert!( - got.frame_weight_sum(slot) > 0, - "ring slot {slot} received nothing" - ); + let slot_sum = got.frame_weight_sum(slot); + assert!(slot_sum > 0, "ring slot {slot} received nothing"); } - assert_matches_recorded( - "planted cross-frame match", - &got, - &Digest { - covered: 19992, - pixel_mean: 0.488864379137, - pixel_rms: 0.570867709963, - weight_mean: 1233.333334961, - probes: [ - 0.836363474528, - 0.574619293213, - 0.300458908081, - 0.034561157227, - 0.774574279785, - 0.513134002686, - 0.238787333171, - 0.979087829590, - ], - }, - ); + + let want = Digest { + covered: 19992, + pixel_mean: 0.488864379137, + pixel_rms: 0.570867709963, + weight_mean: 1233.333334961, + probes: [ + 0.836363474528, + 0.574619293213, + 0.300458908081, + 0.034561157227, + 0.774574279785, + 0.513134002686, + 0.238787333171, + 0.979087829590, + ], + }; + + assert_matches_recorded("planted cross-frame match", &got, &want); } diff --git a/av-denoise-core/src/collab/tests/fused/strength_map.rs b/av-denoise-core/src/collab/tests/fused/strength_map.rs index 68d120d..33c4fb3 100644 --- a/av-denoise-core/src/collab/tests/fused/strength_map.rs +++ b/av-denoise-core/src/collab/tests/fused/strength_map.rs @@ -5,7 +5,7 @@ use super::{Aggregated, Setup, cross_frame_setup, run_fused, run_fused_walk}; use crate::collab::geometry::{ref_pos, refs_along, strength_map_dims}; use crate::collab::kernels::fused::strength_map::strength_map_scale; use crate::collab::kernels::fused::{STRENGTH_MAP_ALL, STRENGTH_MAP_LUMA}; -use crate::collab::tests::helpers::{R, make_client, noisy_field_over}; +use crate::collab::tests::helpers::{R, make_client, noisy_flat_field}; use crate::nlmeans::{ChannelMode, NOISE_CURVE_BINS}; const FRAME_SIDE: u32 = 64; @@ -25,15 +25,18 @@ fn map_scale_kernel( out[index as usize] = strength_map_scale(map, rx, ry, map_cols, map_rows); } -/// Every reference position of a `width` by `height` frame, as `(rx, ry)` pairs. +/// Every reference position of a `width` by `height` frame, as `(left, top)` pairs. fn reference_positions(width: u32, height: u32) -> Vec { let mut positions = Vec::new(); for ref_y in 0..refs_along(height) { for ref_x in 0..refs_along(width) { - positions.push(ref_pos(ref_x, width)); - positions.push(ref_pos(ref_y, height)); + let left = ref_pos(ref_x, width); + let top = ref_pos(ref_y, height); + positions.push(left); + positions.push(top); } } + positions } @@ -48,8 +51,10 @@ fn the_map_scale_is_the_mean_of_the_overlapped_quarters_on_a_ragged_frame() { let count = positions.len() / 2; let client = make_client(); - let map_buf = client.create_from_slice(f32::as_bytes(&map)); - let positions_buf = client.create_from_slice(u32::as_bytes(&positions)); + let map_bytes = f32::as_bytes(&map); + let positions_bytes = u32::as_bytes(&positions); + let map_buf = client.create_from_slice(map_bytes); + let positions_buf = client.create_from_slice(positions_bytes); let out_buf = client.empty(count * size_of::()); unsafe { @@ -69,19 +74,20 @@ fn the_map_scale_is_the_mean_of_the_overlapped_quarters_on_a_ragged_frame() { let got = f32::from_bytes(&out_bytes)[..count].to_vec(); let cols = map_cols as usize; + for (index, &scale) in got.iter().enumerate() { - let rx = positions[2 * index] as usize; - let ry = positions[2 * index + 1] as usize; - let col_lo = rx / 8; - let col_hi = rx.div_ceil(8).min(cols - 1); - let row_lo = ry / 8; - let row_hi = ry.div_ceil(8).min(map_rows as usize - 1); - let sum = map[row_lo * cols + col_lo] - + map[row_lo * cols + col_hi] - + map[row_hi * cols + col_lo] - + map[row_hi * cols + col_hi]; + let left = positions[2 * index] as usize; + let top = positions[2 * index + 1] as usize; + let first_col = left / 8; + let last_col = left.div_ceil(8).min(cols - 1); + let first_row = top / 8; + let last_row = top.div_ceil(8).min(map_rows as usize - 1); + let sum = map[first_row * cols + first_col] + + map[first_row * cols + last_col] + + map[last_row * cols + first_col] + + map[last_row * cols + last_col]; let want = sum / 4.0; - assert_eq!(scale, want, "reference at ({rx}, {ry})"); + assert_eq!(scale, want, "reference at ({left}, {top})"); } } @@ -160,7 +166,7 @@ fn the_luma_map_is_clamped_with_the_curve() { #[test] fn a_two_region_map_thresholds_each_region_by_its_own_multiplier() { let side = FRAME_SIDE; - let frame = noisy_field_over(side, side, 0.5, 0.02); + let frame = noisy_flat_field(side, side, 0.5, 0.02); let (cols, rows) = strength_map_dims(side, side); let mut map = Vec::with_capacity((cols * rows) as usize); for _ in 0..rows { @@ -188,8 +194,11 @@ fn a_two_region_map_thresholds_each_region_by_its_own_multiplier() { // from 40 on only by references reading 0.5, given the spatial radius of 4. assert_columns_identical("left region", &mapped, &left_scaled, side, 0, 24); assert_columns_identical("right region", &mapped, &right_scaled, side, 40, side); - assert!(columns_differ(&left_scaled, &plain, side, 0, 24)); - assert!(columns_differ(&right_scaled, &plain, side, 40, side)); + + let left_changed = columns_differ(&left_scaled, &plain, side, 0, 24); + let right_changed = columns_differ(&right_scaled, &plain, side, 40, side); + assert!(left_changed); + assert!(right_changed); } #[test] @@ -211,13 +220,13 @@ fn both_walks_agree_with_a_map_active() { /// A single-frame 3-channel ring, each channel carrying its own noise. fn three_channel_setup() -> Setup { let pixels = (FRAME_SIDE * FRAME_SIDE) as usize; - let noise = noisy_field_over(FRAME_SIDE, FRAME_SIDE * 3, 0.5, 0.02); - let stored_ch = ChannelMode::Yuv.storage_count() as usize; + let noise = noisy_flat_field(FRAME_SIDE, FRAME_SIDE * 3, 0.5, 0.02); + let stored_channels = ChannelMode::Yuv.storage_count() as usize; - let mut ring = vec![0.0f32; pixels * stored_ch]; + let mut ring = vec![0.0f32; pixels * stored_channels]; for pixel in 0..pixels { for channel in 0..3 { - ring[pixel * stored_ch + channel] = noise[channel * pixels + pixel]; + ring[pixel * stored_channels + channel] = noise[channel * pixels + pixel]; } } @@ -226,14 +235,15 @@ fn three_channel_setup() -> Setup { setup.ring = ring; setup.channel_mode = ChannelMode::Yuv; setup.lambda_ht = LAMBDA; + setup } /// Whether any accumulator of `channel` differs between two single-frame runs. fn channel_differs(first: &Aggregated, second: &Aggregated, channel: usize) -> bool { - let stored_ch = ChannelMode::Yuv.storage_count() as usize; - let first_values = first.accum.iter().skip(channel).step_by(stored_ch); - let second_values = second.accum.iter().skip(channel).step_by(stored_ch); + let stored_channels = ChannelMode::Yuv.storage_count() as usize; + let first_values = first.accum.iter().skip(channel).step_by(stored_channels); + let second_values = second.accum.iter().skip(channel).step_by(stored_channels); first_values .zip(second_values) .any(|(first_value, second_value)| first_value != second_value) @@ -258,6 +268,8 @@ fn a_luma_map_scales_only_channel_zero() { all_mapped.strength_map = Some((all_map, STRENGTH_MAP_ALL)); let all_channels = run_fused(&all_mapped); - assert!(channel_differs(&all_channels, &got, 1)); - assert!(channel_differs(&all_channels, &got, 2)); + let first_chroma_differs = channel_differs(&all_channels, &got, 1); + let second_chroma_differs = channel_differs(&all_channels, &got, 2); + assert!(first_chroma_differs); + assert!(second_chroma_differs); } diff --git a/av-denoise-core/src/collab/tests/fused/walks.rs b/av-denoise-core/src/collab/tests/fused/walks.rs index 9c72601..58da02f 100644 --- a/av-denoise-core/src/collab/tests/fused/walks.rs +++ b/av-denoise-core/src/collab/tests/fused/walks.rs @@ -1,19 +1,15 @@ use super::noise_curve::stepped_curve; use super::{Setup, cross_frame_setup, run_fused_walk, unique_frame}; -/// Asserts the two search walks aggregated the same thing, byte for -/// byte. +/// Asserts the two search walks aggregated the same thing, byte for byte. /// -/// Exact equality is the right bar rather than a tolerance. The -/// warp-uniform walk offers the same candidates, in the same order, and -/// scores them with the same arithmetic. Its extra turns carry the -/// `3.0e38` an unfilled slot already holds, which cannot displace a -/// slot, so they change nothing about the group that is retired. A -/// difference here means the masking let a dead position into a group -/// or dropped a live one, not that floating point drifted. -fn assert_walks_agree(label: &str, s: &Setup) { - let clipped = run_fused_walk(s, Some(false)); - let uniform = run_fused_walk(s, Some(true)); +/// The warp-uniform walk offers the same candidates in the same order with the same arithmetic. Its +/// extra turns carry the `3.0e38` an unfilled slot already holds, which cannot displace a slot. A +/// difference therefore means the masking let a dead position in or dropped a live one, not that +/// floating point drifted. +fn assert_walks_agree(label: &str, setup: &Setup) { + let clipped = run_fused_walk(setup, Some(false)); + let uniform = run_fused_walk(setup, Some(true)); assert_eq!( clipped.group_weight, uniform.group_weight, @@ -28,70 +24,55 @@ fn assert_walks_agree(label: &str, s: &Setup) { "{label}: the two search walks scattered different weights", ); assert!( - uniform.group_weight.iter().any(|w| *w > 0.0), + uniform.group_weight.iter().any(|weight| *weight > 0.0), "{label}: neither walk aggregated anything, so agreeing proves nothing", ); } -/// The warp-uniform walk is only correct if it is a pure change of -/// schedule, so this pins it against the walk it replaces on the spatial -/// search alone. -/// -/// The references along the left and top edges are the ones whose -/// clipped rectangle is narrower than the unclipped span, which is -/// exactly where the uniform walk takes turns the other one does not. -/// Those turns have to score nothing. +/// References along the left and top edges have a clipped rectangle narrower than the unclipped +/// span, which is where the uniform walk takes extra turns. Those turns must score nothing. #[test] fn warp_uniform_search_matches_the_clipped_search_on_the_spatial_pass() { - let (w, h) = (48u32, 48u32); - let s = Setup::spatial_only(unique_frame(w, h), w, h); + let (width, height) = (48u32, 48u32); + let frame = unique_frame(width, height); + let setup = Setup::spatial_only(frame, width, height); - assert_walks_agree("spatial", &s); + assert_walks_agree("spatial", &setup); } -/// The same equivalence across the temporal search, which is where the -/// two walks diverge most. -/// -/// [cross_frame_setup] is built for this: its motion vectors push some -/// refine windows off the frame so the clip matters, and its confidence -/// field straddles `c_min`, so blocks that one group scores its -/// neighbour skips. Under the clipped walk those are the two things that -/// give groups sharing a warp different trip counts, and under the -/// uniform walk they have to fold into the mask instead without moving -/// a single member. +/// [cross_frame_setup] pushes some refine windows off the frame and its confidences straddle `c_min`, +/// the two things that give groups sharing a warp different trip counts under the clipped walk. #[test] fn warp_uniform_search_matches_the_clipped_search_across_frames() { - let s = cross_frame_setup(64, 64, 2); + let setup = cross_frame_setup(64, 64, 2); - assert_walks_agree("cross frame", &s); + assert_walks_agree("cross frame", &setup); } -/// A gated neighbour is the one case where the uniform walk reads a -/// motion block the clipped walk never touches, so the `seen_*` slot it -/// leaves behind has to stay empty or it would hide positions a later -/// covering block still owes the search. +/// The uniform walk reads a gated motion block the clipped walk never touches, so the `seen_*` slot +/// it leaves must stay empty or it would hide positions a later covering block still owes the search. #[test] fn warp_uniform_search_matches_the_clipped_search_when_every_neighbour_is_gated() { - let mut s = cross_frame_setup(64, 64, 2); - // Above every confidence the setup plants, so no temporal block ever - // scores and every `seen_*` slot is one the uniform walk wrote. - s.c_min = 2.0; + let mut setup = cross_frame_setup(64, 64, 2); + // Above every planted confidence, so every `seen_*` slot is one the uniform walk wrote. + setup.c_min = 2.0; - assert_walks_agree("all gated", &s); + assert_walks_agree("all gated", &setup); } -/// The same equivalence at radius 1, where the grid is four volumes of two frames. +/// At radius 1 the grid is four volumes of two frames. #[test] fn warp_uniform_search_matches_the_clipped_search_at_radius_one() { - let s = cross_frame_setup(64, 64, 1); + let setup = cross_frame_setup(64, 64, 1); - assert_walks_agree("radius one", &s); + assert_walks_agree("radius one", &setup); } #[test] fn warp_uniform_search_matches_the_clipped_search_with_a_noise_curve() { - let mut s = cross_frame_setup(64, 64, 2); - s.noise_curve = Some(stepped_curve()); + let mut setup = cross_frame_setup(64, 64, 2); + let curve = stepped_curve(); + setup.noise_curve = Some(curve); - assert_walks_agree("noise curve", &s); + assert_walks_agree("noise curve", &setup); } diff --git a/av-denoise-core/src/collab/tests/grid.rs b/av-denoise-core/src/collab/tests/grid.rs index dbca2b6..b3e2573 100644 --- a/av-denoise-core/src/collab/tests/grid.rs +++ b/av-denoise-core/src/collab/tests/grid.rs @@ -38,24 +38,26 @@ fn grid_probe( #[cube(launch_unchecked)] fn grid_variance_probe(input: &Array, output: &mut Array, #[comptime] grid_frames: u32) { - let mut v = Array::::new(MAX_K as usize); + let mut variances = Array::::new(MAX_K as usize); #[unroll] for m in 0..MAX_K { - v[m as usize] = input[m as usize]; + variances[m as usize] = input[m as usize]; } - grid_variance(&mut v, grid_frames); + grid_variance(&mut variances, grid_frames); #[unroll] for m in 0..MAX_K { - output[m as usize] = v[m as usize]; + output[m as usize] = variances[m as usize]; } } fn run_grid(stack: &[f32], grid_frames: u32, inverse: bool) -> Vec { let client = make_client(); - let input = client.create_from_slice(f32::as_bytes(stack)); - let output = client.empty(std::mem::size_of_val(stack)); + let stack_bytes = f32::as_bytes(stack); + let input = client.create_from_slice(stack_bytes); + let stack_size = size_of_val(stack); + let output = client.empty(stack_size); unsafe { grid_probe::launch_unchecked::( @@ -69,13 +71,15 @@ fn run_grid(stack: &[f32], grid_frames: u32, inverse: bool) -> Vec { ); } - let bytes = client.read_one(output).expect("grid readback failed"); - f32::from_bytes(&bytes)[..stack.len()].to_vec() + let output_bytes = client.read_one(output).expect("grid readback failed"); + + f32::from_bytes(&output_bytes)[..stack.len()].to_vec() } -fn run_grid_variance(v: &[f32; 8], grid_frames: u32) -> Vec { +fn run_grid_variance(variances: &[f32; 8], grid_frames: u32) -> Vec { let client = make_client(); - let input = client.create_from_slice(f32::as_bytes(v)); + let variance_bytes = f32::as_bytes(variances); + let input = client.create_from_slice(variance_bytes); let output = client.empty(8 * size_of::()); unsafe { @@ -89,8 +93,9 @@ fn run_grid_variance(v: &[f32; 8], grid_frames: u32) -> Vec { ); } - let bytes = client.read_one(output).expect("grid variance readback failed"); - f32::from_bytes(&bytes)[..8].to_vec() + let output_bytes = client.read_one(output).expect("grid variance readback failed"); + + f32::from_bytes(&output_bytes)[..8].to_vec() } /// A stack of distinct, non-symmetric values so a wrong pairing or level order shows. @@ -100,13 +105,14 @@ fn distinct_stack() -> Vec { .collect() } -/// Member `m`'s value at spatial position `pos`, gathered into one column. +/// Every member's value at spatial position `pos`, gathered into one column. fn column(stack: &[f32], pos: u32) -> [f32; 8] { - let mut out = [0.0f32; 8]; - for (m, value) in out.iter_mut().enumerate() { + let mut values = [0.0f32; 8]; + for (m, value) in values.iter_mut().enumerate() { *value = stack[m * PATCH_SIZE as usize + pos as usize]; } - out + + values } #[test] @@ -117,8 +123,10 @@ fn gpu_grid_forward_matches_the_host_mirror() { let gpu = run_grid(&stack, grid_frames, false); for pos in 0..PATCH_SIZE { - let host = grid_fwd_host(&column(&stack, pos), grid_frames); + let input_column = column(&stack, pos); + let host = grid_fwd_host(&input_column, grid_frames); let device = column(&gpu, pos); + for m in 0..8 { assert!( (host[m] - device[m]).abs() < 1e-5, @@ -155,6 +163,7 @@ fn host_grid_inverse_undoes_the_host_forward() { for grid_frames in [2u32, 4] { let forward = grid_fwd_host(&input, grid_frames); let round_trip = grid_inv_host(&forward, grid_frames); + for m in 0..8 { assert!((input[m] - round_trip[m]).abs() < 1e-5, "T={grid_frames} m={m}"); } @@ -168,6 +177,7 @@ fn gpu_grid_variance_matches_the_host_mirror() { for grid_frames in [2u32, 4] { let host = grid_variance_host(&sig2, grid_frames); let gpu = run_grid_variance(&sig2, grid_frames); + for m in 0..8 { assert!( (host[m] - gpu[m]).abs() < 1e-5, @@ -191,8 +201,6 @@ fn grid_dc_variance_is_the_mean_member_variance() { } } -/// A volume whose frames are identical puts nothing in its temporal detail. -/// /// Coefficient `s * T + t` with `t > 0` is temporal detail after both passes. #[test] fn a_volume_constant_in_time_has_no_temporal_detail() { @@ -206,6 +214,7 @@ fn a_volume_constant_in_time_has_no_temporal_detail() { } let output = grid_fwd_host(&input, grid_frames); + for volume in 0..volumes { for frame in 1..grid_frames { let index = (volume * grid_frames + frame) as usize; diff --git a/av-denoise-core/src/collab/tests/group.rs b/av-denoise-core/src/collab/tests/group.rs index 9227b22..d9d98a3 100644 --- a/av-denoise-core/src/collab/tests/group.rs +++ b/av-denoise-core/src/collab/tests/group.rs @@ -11,10 +11,9 @@ use crate::collab::kernels::group::{ unpack_t_host, }; -/// Runs [`pack_pos`], [`pack_pos_t`], and [`clamp_top_left`] on the GPU, -/// one input per thread, so the host mirrors below are checked against -/// the kernels that actually consume them rather than against -/// themselves. +/// Runs [pack_pos], [pack_pos_t] and [clamp_top_left] on the GPU, one input per thread. +/// +/// This checks the host mirrors against the kernels that consume them rather than against themselves. #[cube(launch_unchecked)] fn group_helpers_kernel( xs: &Array, @@ -42,68 +41,68 @@ fn run_helpers( coords: &[i32], max_pos: &[u32], ) -> (Vec, Vec, Vec) { - let n = xs.len(); - assert_eq!(ys.len(), n); - assert_eq!(ts.len(), n); - assert_eq!(coords.len(), n); - assert_eq!(max_pos.len(), n); + let count = xs.len(); + assert_eq!(ys.len(), count); + assert_eq!(ts.len(), count); + assert_eq!(coords.len(), count); + assert_eq!(max_pos.len(), count); let client = make_client(); - let xs_buf = client.create_from_slice(u32::as_bytes(xs)); - let ys_buf = client.create_from_slice(u32::as_bytes(ys)); - let ts_buf = client.create_from_slice(u32::as_bytes(ts)); - let coords_buf = client.create_from_slice(i32::as_bytes(coords)); - let max_buf = client.create_from_slice(u32::as_bytes(max_pos)); - // These size the kernel's three output buffers, which hold one u32 - // per input coordinate. `size_of_val(xs)` reaches the same number - // but ties an output's size to an input's slice, which reads as if - // the buffers held `xs` itself. + let xs_bytes = u32::as_bytes(xs); + let ys_bytes = u32::as_bytes(ys); + let ts_bytes = u32::as_bytes(ts); + let coords_bytes = i32::as_bytes(coords); + let max_bytes = u32::as_bytes(max_pos); + let xs_buf = client.create_from_slice(xs_bytes); + let ys_buf = client.create_from_slice(ys_bytes); + let ts_buf = client.create_from_slice(ts_bytes); + let coords_buf = client.create_from_slice(coords_bytes); + let max_buf = client.create_from_slice(max_bytes); #[expect( clippy::manual_slice_size_calculation, - reason = "n is the element count these outputs hold, not xs's byte length" + reason = "count is the element count these outputs hold, not xs's byte length" )] - let packed_buf = client.empty(n * size_of::()); + let packed_buf = client.empty(count * size_of::()); #[expect( clippy::manual_slice_size_calculation, - reason = "n is the element count these outputs hold, not xs's byte length" + reason = "count is the element count these outputs hold, not xs's byte length" )] - let packed_t_buf = client.empty(n * size_of::()); + let packed_t_buf = client.empty(count * size_of::()); #[expect( clippy::manual_slice_size_calculation, - reason = "n is the element count these outputs hold, not xs's byte length" + reason = "count is the element count these outputs hold, not xs's byte length" )] - let clamped_buf = client.empty(n * size_of::()); + let clamped_buf = client.empty(count * size_of::()); unsafe { group_helpers_kernel::launch_unchecked::( &client, CubeCount::new_1d(1), CubeDim::new_1d(64), - ArrayArg::from_raw_parts(xs_buf, n), - ArrayArg::from_raw_parts(ys_buf, n), - ArrayArg::from_raw_parts(ts_buf, n), - ArrayArg::from_raw_parts(coords_buf, n), - ArrayArg::from_raw_parts(max_buf, n), - ArrayArg::from_raw_parts(packed_buf.clone(), n), - ArrayArg::from_raw_parts(packed_t_buf.clone(), n), - ArrayArg::from_raw_parts(clamped_buf.clone(), n), - n as u32, + ArrayArg::from_raw_parts(xs_buf, count), + ArrayArg::from_raw_parts(ys_buf, count), + ArrayArg::from_raw_parts(ts_buf, count), + ArrayArg::from_raw_parts(coords_buf, count), + ArrayArg::from_raw_parts(max_buf, count), + ArrayArg::from_raw_parts(packed_buf.clone(), count), + ArrayArg::from_raw_parts(packed_t_buf.clone(), count), + ArrayArg::from_raw_parts(clamped_buf.clone(), count), + count as u32, ); } - let packed = client.read_one(packed_buf).expect("packed readback failed"); - let packed_t = client.read_one(packed_t_buf).expect("packed_t readback failed"); - let clamped = client.read_one(clamped_buf).expect("clamped readback failed"); + let packed_bytes = client.read_one(packed_buf).expect("packed readback failed"); + let packed_t_bytes = client.read_one(packed_t_buf).expect("packed_t readback failed"); + let clamped_bytes = client.read_one(clamped_buf).expect("clamped readback failed"); + let packed = u32::from_bytes(&packed_bytes)[..count].to_vec(); + let packed_t = u32::from_bytes(&packed_t_bytes)[..count].to_vec(); + let clamped = u32::from_bytes(&clamped_bytes)[..count].to_vec(); - ( - u32::from_bytes(&packed)[..n].to_vec(), - u32::from_bytes(&packed_t)[..n].to_vec(), - u32::from_bytes(&clamped)[..n].to_vec(), - ) + (packed, packed_t, clamped) } -/// Positions spanning both halves of the packed word, including the -/// largest value each 13-bit field holds. +/// Positions spanning both halves of the packed word, including the largest value each 13-bit field +/// holds. const POSITIONS: &[(u32, u32)] = &[ (0, 0), (1, 0), @@ -117,18 +116,21 @@ const POSITIONS: &[(u32, u32)] = &[ #[test] fn packing_a_position_round_trips_through_the_host_mirror() { for &(x, y) in POSITIONS { - let (px, py) = unpack_pos_host(pack_pos_host(x, y)); - assert_eq!((px, py), (x, y), "({x}, {y}) did not survive the round trip"); + let packed = pack_pos_host(x, y); + let unpacked = unpack_pos_host(packed); + assert_eq!(unpacked, (x, y), "({x}, {y}) did not survive the round trip"); - // A neighbour index in the field above y must not leak into the - // coordinate unpack_pos_host reads. - let (px, py) = unpack_pos_host(pack_pos_t_host(x, y, 4)); - assert_eq!((px, py), (x, y), "t=4 leaked into the coordinates for ({x}, {y})"); + // A neighbour index in the field above y must not leak into the coordinates. + let packed_with_t = pack_pos_t_host(x, y, 4); + let unpacked_with_t = unpack_pos_host(packed_with_t); + assert_eq!( + unpacked_with_t, + (x, y), + "t=4 leaked into the coordinates for ({x}, {y})" + ); } } -/// The neighbour index rides in the bits above y without disturbing -/// either coordinate, so the existing unpack still reads them. #[test] fn pack_pos_t_round_trips_and_leaves_the_coordinates_readable() { for &(x, y, t) in &[ @@ -136,22 +138,24 @@ fn pack_pos_t_round_trips_and_leaves_the_coordinates_readable() { (1919, 1079, 0), (1912, 1072, 4), (7, 3, 1), - // All three fields at their maximum simultaneously, proving none - // bleeds into another at saturation. + // All three fields at their maximum at once, so none can bleed into another at saturation. (8191, 8191, 63), ] { let packed = pack_pos_t_host(x, y, t); - assert_eq!(unpack_pos_host(packed), (x, y), "coords for ({x},{y},{t})"); - assert_eq!(unpack_t_host(packed), t, "t for ({x},{y},{t})"); + let coords = unpack_pos_host(packed); + let unpacked_t = unpack_t_host(packed); + assert_eq!(coords, (x, y), "coords for ({x},{y},{t})"); + assert_eq!(unpacked_t, t, "t for ({x},{y},{t})"); } } -/// A centre-frame position packs a zero, which is what lets the filter -/// stage tell it apart from a motion-predicted member without a second -/// array. +/// A centre-frame position packs a zero, which lets the filter stage tell it apart from a +/// motion-predicted member without a second array. #[test] fn pack_pos_t_agrees_with_pack_pos_at_t_zero() { - assert_eq!(pack_pos_t_host(120, 400, 0), pack_pos_host(120, 400)); + let packed_t = pack_pos_t_host(120, 400, 0); + let packed = pack_pos_host(120, 400); + assert_eq!(packed_t, packed); } #[test] @@ -165,8 +169,9 @@ fn a_position_packs_x_low_and_y_high() { fn distinct_positions_pack_to_distinct_words() { let mut seen = std::collections::HashSet::new(); for &(x, y) in POSITIONS { + let packed = pack_pos_host(x, y); assert!( - seen.insert(pack_pos_host(x, y)), + seen.insert(packed), "({x}, {y}) collided with an earlier position" ); } @@ -176,26 +181,24 @@ fn distinct_positions_pack_to_distinct_words() { fn the_gpu_helpers_match_their_host_mirrors() { let xs: Vec = POSITIONS.iter().map(|&(x, _)| x).collect(); let ys: Vec = POSITIONS.iter().map(|&(_, y)| y).collect(); - // One t per position, covering the centre-frame value and a spread - // of neighbour indices. + // One t per position, covering the centre-frame value and a spread of neighbour indices. let ts: Vec = vec![0, 1, 4, 0, 2, 0, 3]; - // One coordinate per position, covering below the range, inside it, - // and past its top, against a max of 24 (a 32-wide frame's last - // legal 8x8 patch position). + // One coordinate per position, covering below the range, inside it and past its top, against a + // max of 24 (a 32-wide frame's last legal 8x8 patch position). let coords: Vec = vec![-9, -1, 0, 1, 24, 25, 4096]; let max_pos: Vec = vec![24; coords.len()]; let (packed, packed_t, clamped) = run_helpers(&xs, &ys, &ts, &coords, &max_pos); for (i, &(x, y)) in POSITIONS.iter().enumerate() { + let expected = pack_pos_host(x, y); + let expected_t = pack_pos_t_host(x, y, ts[i]); assert_eq!( - packed[i], - pack_pos_host(x, y), + packed[i], expected, "pack_pos disagreed with pack_pos_host at ({x}, {y})" ); assert_eq!( - packed_t[i], - pack_pos_t_host(x, y, ts[i]), + packed_t[i], expected_t, "pack_pos_t disagreed with pack_pos_t_host at ({x}, {y}, {})", ts[i] ); diff --git a/av-denoise-core/src/collab/tests/helpers.rs b/av-denoise-core/src/collab/tests/helpers.rs index 41ece8e..dcc3775 100644 --- a/av-denoise-core/src/collab/tests/helpers.rs +++ b/av-denoise-core/src/collab/tests/helpers.rs @@ -10,15 +10,12 @@ pub(super) fn make_client() -> ComputeClient { /// Adds independent pseudo-Gaussian noise to a flat `base` field. /// -/// Each sample sums four hash-derived uniforms in `[-0.5, 0.5]` -/// (Irwin-Hall) and rescales to the requested standard deviation, so the -/// same arguments always reproduce the same frame. Ported from -/// `src/nlmeans/tests/helpers.rs`, keeping only the flat-field case this -/// tree needs. -pub(super) fn noisy_field_over(w: u32, h: u32, base: f32, sigma: f32) -> Vec { +/// Each sample sums four hash-derived uniforms between -0.5 and 0.5 (Irwin-Hall) and rescales to +/// the requested standard deviation, so the same arguments always reproduce the same frame. +pub(super) fn noisy_flat_field(width: u32, height: u32, base: f32, sigma: f32) -> Vec { let unit_std = (1.0f32 / 3.0f32).sqrt(); - let mut frame = vec![0.0f32; (w * h) as usize]; - for idx in 0..(w * h) { + let mut frame = vec![0.0f32; (width * height) as usize]; + for idx in 0..(width * height) { let mut sum = 0.0f32; for k in 0..4u32 { let mut hash = (idx * 4 + k).wrapping_mul(2654435761).wrapping_add(0x9E3779B9); @@ -27,25 +24,23 @@ pub(super) fn noisy_field_over(w: u32, h: u32, base: f32, sigma: f32) -> Vec> 13; sum += (hash as f32 / u32::MAX as f32) - 0.5; } + frame[idx as usize] = base + (sum / unit_std) * sigma; } + frame } -/// A flat base value with a per-pixel hash offset that never repeats -/// across the frame. +/// A horizontal ramp plus a per-pixel hash offset, so no 8x8 window repeats anywhere in the frame. /// -/// Any two distinct 8x8 windows into this frame differ in most of their -/// 64 pixels, so a tiny admission threshold rejects every candidate but -/// the reference patch itself. The horizontal ramp on its own would -/// still leave two vertically-shifted patches identical, since the -/// frame would otherwise repeat down every row, so the hash term is -/// what actually makes every position unique. -pub(super) fn make_unique_frame(w: u32, h: u32) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let idx = y * w + x; +/// Any two distinct windows differ in most of their 64 pixels, so a tiny admission threshold rejects +/// every candidate but the reference patch itself. The ramp alone repeats down every row, so the hash +/// term is what makes vertically shifted patches differ. +pub(super) fn make_unique_frame(width: u32, height: u32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let idx = y * width + x; let mut hash = idx.wrapping_mul(2654435761).wrapping_add(0x9E3779B9); hash ^= hash >> 15; hash = hash.wrapping_mul(0x85EBCA6B); @@ -54,32 +49,32 @@ pub(super) fn make_unique_frame(w: u32, h: u32) -> Vec { frame[idx as usize] = x as f32 * 10.0 + offset * 10.0; } } + frame } -/// Writes an 8x8 patch into `frame` with its top-left corner at `(px, -/// py)`. -pub(super) fn plant_patch(frame: &mut [f32], w: u32, px: u32, py: u32, patch: &[f32; 64]) { +/// Writes an 8x8 patch into `frame` with its top-left corner at `(left, top)`. +pub(super) fn plant_patch(frame: &mut [f32], width: u32, left: u32, top: u32, patch: &[f32; 64]) { for row in 0..8u32 { for col in 0..8u32 { - let idx = (py + row) * w + (px + col); + let idx = (top + row) * width + (left + col); frame[idx as usize] = patch[(row * 8 + col) as usize]; } } } -/// A deterministic 8x8 texture with values well clear of the flat -/// backgrounds these tests plant it over. +/// A deterministic 8x8 texture with values well clear of the flat backgrounds the tests plant it over. pub(super) fn deterministic_texture(seed: u32) -> [f32; 64] { - let mut out = [0.0f32; 64]; - for (idx, v) in out.iter_mut().enumerate() { + let mut texture = [0.0f32; 64]; + for (idx, value) in texture.iter_mut().enumerate() { let mut hash = (idx as u32) .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(0x9E37_79B9)); hash ^= hash >> 15; hash = hash.wrapping_mul(0x85EBCA6B); hash ^= hash >> 13; - *v = 0.6 + (hash as f32 / u32::MAX as f32) * 0.3; + *value = 0.6 + (hash as f32 / u32::MAX as f32) * 0.3; } - out + + texture } diff --git a/av-denoise-core/src/collab/tests/plane_ops.rs b/av-denoise-core/src/collab/tests/plane_ops.rs index 29cd7b7..a40f581 100644 --- a/av-denoise-core/src/collab/tests/plane_ops.rs +++ b/av-denoise-core/src/collab/tests/plane_ops.rs @@ -16,51 +16,49 @@ fn ssd_reduce_kernel(input: &Array, out: &mut Array) { out[tid as usize] = plane_ssd_reduce8(input[tid as usize]); } -/// Runs [`plane_ssd_reduce8`] on the GPU, one lane per input, so the -/// host checks the code the kernels actually consume. +/// Runs [plane_ssd_reduce8] on the GPU, one lane per input. fn run_ssd_reduce(input: &[f32]) -> Vec { - let n = input.len(); + let count = input.len(); let client = make_client(); - let input_buf = client.create_from_slice(f32::as_bytes(input)); - // One output slot per input lane. `size_of_val(input)` reaches the - // same number but ties the output's size to the input's slice. + let input_bytes = f32::as_bytes(input); + let input_buf = client.create_from_slice(input_bytes); #[expect( clippy::manual_slice_size_calculation, - reason = "n is the element count this output holds, not the input's byte length" + reason = "count is the element count this output holds, not the input's byte length" )] - let out_buf = client.empty(n * size_of::()); + let out_buf = client.empty(count * size_of::()); unsafe { ssd_reduce_kernel::launch_unchecked::( &client, CubeCount::new_1d(1), - CubeDim::new_1d(n as u32), - ArrayArg::from_raw_parts(input_buf, n), - ArrayArg::from_raw_parts(out_buf.clone(), n), + CubeDim::new_1d(count as u32), + ArrayArg::from_raw_parts(input_buf, count), + ArrayArg::from_raw_parts(out_buf.clone(), count), ); } - let out = client.read_one(out_buf).expect("ssd reduce readback failed"); - f32::from_bytes(&out)[..n].to_vec() + let out_bytes = client.read_one(out_buf).expect("ssd reduce readback failed"); + + f32::from_bytes(&out_bytes)[..count].to_vec() } -/// Every lane of a group must come back holding that group's whole sum, -/// and groups must not bleed into each other. #[test] fn ssd_reduce8_sums_within_its_own_group() { - // Group 0 holds 1..=8 summing to 36, group 1 holds 100..=800 - // summing to 3600. Distinct magnitudes so a leak across the - // boundary cannot coincidentally match. + // Group 0 holds 1..=8 summing to 36, group 1 holds 100..=800 summing to 3600. The magnitudes + // differ so a leak across the boundary cannot coincidentally match. let input: Vec = (1..=8) - .map(|v| v as f32) - .chain((1..=8).map(|v| (v * 100) as f32)) + .map(|value| value as f32) + .chain((1..=8).map(|value| (value * 100) as f32)) .collect(); - let out = run_ssd_reduce(&input); - for (lane, &v) in out.iter().take(8).enumerate() { - assert_eq!(v, 36.0, "lane {lane} of group 0"); + let sums = run_ssd_reduce(&input); + + for (lane, &sum) in sums.iter().take(8).enumerate() { + assert_eq!(sum, 36.0, "lane {lane} of group 0"); } - for (lane, &v) in out.iter().enumerate().take(16).skip(8) { - assert_eq!(v, 3600.0, "lane {lane} of group 1"); + + for (lane, &sum) in sums.iter().enumerate().take(16).skip(8) { + assert_eq!(sum, 3600.0, "lane {lane} of group 1"); } } @@ -120,14 +118,16 @@ fn shift_insert_gated_kernel( out_p[tid as usize] = best_p; } -/// Runs [`shift_insert8`] over `dists`/`posns` on a single 8-lane group, -/// feeding every candidate through in order. +/// Runs [shift_insert8] over `dists`/`posns` on a single 8-lane group, feeding every candidate +/// through in order. fn run_shift_insert_with(dists: &[f32], posns: &[u32]) -> (Vec, Vec) { - let n = dists.len(); - assert_eq!(posns.len(), n); + let count = dists.len(); + assert_eq!(posns.len(), count); let client = make_client(); - let dists_buf = client.create_from_slice(f32::as_bytes(dists)); - let posns_buf = client.create_from_slice(u32::as_bytes(posns)); + let dists_bytes = f32::as_bytes(dists); + let posns_bytes = u32::as_bytes(posns); + let dists_buf = client.create_from_slice(dists_bytes); + let posns_buf = client.create_from_slice(posns_bytes); let out_d_buf = client.empty(MAX_K as usize * size_of::()); let out_p_buf = client.empty(MAX_K as usize * size_of::()); @@ -136,40 +136,42 @@ fn run_shift_insert_with(dists: &[f32], posns: &[u32]) -> (Vec, Vec) { &client, CubeCount::new_1d(1), CubeDim::new_1d(8), - ArrayArg::from_raw_parts(dists_buf, n), - ArrayArg::from_raw_parts(posns_buf, n), - n as u32, + ArrayArg::from_raw_parts(dists_buf, count), + ArrayArg::from_raw_parts(posns_buf, count), + count as u32, ArrayArg::from_raw_parts(out_d_buf.clone(), 8), ArrayArg::from_raw_parts(out_p_buf.clone(), 8), ); } - let out_d = client + let dist_bytes = client .read_one(out_d_buf) .expect("shift_insert dist readback failed"); - let out_p = client + let pos_bytes = client .read_one(out_p_buf) .expect("shift_insert pos readback failed"); - ( - f32::from_bytes(&out_d)[..8].to_vec(), - u32::from_bytes(&out_p)[..8].to_vec(), - ) + let best_dists = f32::from_bytes(&dist_bytes)[..8].to_vec(); + let best_posns = u32::from_bytes(&pos_bytes)[..8].to_vec(); + + (best_dists, best_posns) } -/// [`run_shift_insert_with`] with `posns` set to `(0..n)`. +/// [run_shift_insert_with] with `posns` set to `0..dists.len()`. fn run_shift_insert(dists: &[f32]) -> (Vec, Vec) { let posns: Vec = (0..dists.len() as u32).collect(); run_shift_insert_with(dists, &posns) } -/// Runs [`shift_insert8_gated`] over `dists`/`posns` on a single 8-lane -/// group, feeding every candidate through in order. +/// Runs [shift_insert8_gated] over `dists`/`posns` on a single 8-lane group, feeding every +/// candidate through in order. fn run_shift_insert_gated_with(dists: &[f32], posns: &[u32]) -> (Vec, Vec) { - let n = dists.len(); - assert_eq!(posns.len(), n); + let count = dists.len(); + assert_eq!(posns.len(), count); let client = make_client(); - let dists_buf = client.create_from_slice(f32::as_bytes(dists)); - let posns_buf = client.create_from_slice(u32::as_bytes(posns)); + let dists_bytes = f32::as_bytes(dists); + let posns_bytes = u32::as_bytes(posns); + let dists_buf = client.create_from_slice(dists_bytes); + let posns_buf = client.create_from_slice(posns_bytes); let out_d_buf = client.empty(MAX_K as usize * size_of::()); let out_p_buf = client.empty(MAX_K as usize * size_of::()); @@ -178,37 +180,34 @@ fn run_shift_insert_gated_with(dists: &[f32], posns: &[u32]) -> (Vec, Vec, posns: Vec, } -/// The inputs the four `shift_insert8` tests below use, so the gated -/// test in this module runs the identical inputs rather than a -/// duplicated copy of the literals. +/// The inputs of the four `shift_insert8` tests, shared so the gated test runs the identical inputs. fn insert_test_cases() -> Vec { vec![ InsertCase { @@ -228,13 +227,12 @@ fn insert_test_cases() -> Vec { }, InsertCase { name: "handles_a_descending_feed", - dists: (0..12).rev().map(|v| v as f32).collect(), + dists: (0..12).rev().map(|value| value as f32).collect(), posns: (0..12).collect(), }, ] } -/// The case from [`insert_test_cases`] with the given name. fn find_case(name: &str) -> InsertCase { insert_test_cases() .into_iter() @@ -242,8 +240,6 @@ fn find_case(name: &str) -> InsertCase { .unwrap_or_else(|| panic!("no insert_test_cases entry named {name}")) } -/// The eight lanes must hold the eight smallest distances in ascending -/// order, slot 0 the smallest. #[test] fn shift_insert8_matches_a_host_sort() { let case = find_case("matches_a_host_sort"); @@ -253,8 +249,7 @@ fn shift_insert8_matches_a_host_sort() { assert_eq!(&got_d[..8], &want[..8]); } -/// Ties must resolve to the candidate seen first, which is what fixes -/// which member a group keeps on flat content. +/// Ties resolve to the candidate seen first, which fixes which member a group keeps on flat content. #[test] fn shift_insert8_breaks_ties_toward_the_first_seen() { let case = find_case("breaks_ties_toward_the_first_seen"); @@ -262,20 +257,19 @@ fn shift_insert8_breaks_ties_toward_the_first_seen() { assert_eq!(&got_p[..8], &[0u32, 1, 2, 3, 4, 5, 6, 7]); } -/// Fewer candidates than slots leaves the tail at the sentinel rather -/// than at a stale or duplicated entry. #[test] fn shift_insert8_leaves_unfilled_slots_at_the_sentinel() { let case = find_case("leaves_unfilled_slots_at_the_sentinel"); let (got_d, _) = run_shift_insert(&case.dists); assert_eq!(&got_d[..3], &[1.0, 2.0, 3.0]); - for (slot, &d) in got_d.iter().enumerate().take(8).skip(3) { - assert!(d > 1.0e38, "slot {slot} should still be the sentinel"); + + for (slot, &dist) in got_d.iter().enumerate().take(8).skip(3) { + assert!(dist > 1.0e38, "slot {slot} should still be the sentinel"); } } -/// A strictly descending feed exercises the insert path on every single -/// candidate, which is the case the shift logic is easiest to get wrong on. +/// A strictly descending feed takes the insert path on every candidate, the case the shift logic +/// is easiest to get wrong on. #[test] fn shift_insert8_handles_a_descending_feed() { let case = find_case("handles_a_descending_feed"); @@ -283,9 +277,8 @@ fn shift_insert8_handles_a_descending_feed() { assert_eq!(&got_d[..8], &[0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0]); } -/// The gate is a compute saving, never an admission decision. It must -/// produce the same eight slots as the ungated insert on every input the -/// ungated one is tested on. +/// The gate is a compute saving, never an admission decision, so it must keep the same eight slots +/// as the ungated insert. #[test] fn shift_insert8_gated_matches_the_ungated_insert() { for case in insert_test_cases() { @@ -301,57 +294,49 @@ fn transpose_kernel(input: &Array, out: &mut Array) { let grp = tid / 8u32; let sub = tid % 8u32; let mut buf = SharedMemory::::new(8usize * 65usize); - let mut v = Array::::new(8usize); + let mut values = Array::::new(8usize); #[unroll] for i in 0..8u32 { - v[i as usize] = input[(tid * 8u32 + i) as usize]; + values[i as usize] = input[(tid * 8u32 + i) as usize]; } - transpose8(&mut buf, &mut v, sub, grp); + transpose8(&mut buf, &mut values, sub, grp); #[unroll] for i in 0..8u32 { - out[(tid * 8u32 + i) as usize] = v[i as usize]; + out[(tid * 8u32 + i) as usize] = values[i as usize]; } } -/// Runs [`transpose8`] over one cube of two 8-lane groups and reads back -/// both transposed blocks. -/// -/// Two groups is the point, because it catches cross-group bleed -/// through `buf`. +/// Runs [transpose8] over one cube of two 8-lane groups, so cross-group bleed through `buf` shows. fn run_transpose(input: &[f32]) -> Vec { - let n = input.len(); - assert_eq!(n, 128); + let count = input.len(); + assert_eq!(count, 128); let client = make_client(); - let input_buf = client.create_from_slice(f32::as_bytes(input)); - // One output slot per input value, sized from the count rather than - // from the input's own slice. + let input_bytes = f32::as_bytes(input); + let input_buf = client.create_from_slice(input_bytes); #[expect( clippy::manual_slice_size_calculation, - reason = "n is the element count this output holds, not the input's byte length" + reason = "count is the element count this output holds, not the input's byte length" )] - let out_buf = client.empty(n * size_of::()); + let out_buf = client.empty(count * size_of::()); unsafe { transpose_kernel::launch_unchecked::( &client, CubeCount::new_1d(1), CubeDim::new_1d(16), - ArrayArg::from_raw_parts(input_buf, n), - ArrayArg::from_raw_parts(out_buf.clone(), n), + ArrayArg::from_raw_parts(input_buf, count), + ArrayArg::from_raw_parts(out_buf.clone(), count), ); } - let out = client.read_one(out_buf).expect("transpose readback failed"); - f32::from_bytes(&out)[..n].to_vec() + let out_bytes = client.read_one(out_buf).expect("transpose readback failed"); + + f32::from_bytes(&out_bytes)[..count].to_vec() } -/// Element (row, col) of a group's block must come back at (col, row), -/// and two groups sharing a cube must not bleed into each other. -/// -/// Lane `sub` holds column `sub` on entry, so `input[sub * 8 + r]` is -/// element `(r, sub)`. After the transpose lane `sub` holds row `sub`, -/// so the same slot is element `(sub, r)`. Group 1's block is offset by -/// 1000 so a cross-group leak cannot coincidentally match. +/// Lane `sub` holds column `sub` on entry, so `input[sub * 8 + r]` is element `(r, sub)`. After the +/// transpose the same slot is element `(sub, r)`. Group 1 is offset by 1000 so a cross-group leak +/// cannot coincidentally match. #[test] fn transpose8_swaps_rows_and_columns_within_each_group() { let mut input = vec![0.0f32; 128]; @@ -363,12 +348,14 @@ fn transpose8_swaps_rows_and_columns_within_each_group() { } } } - let out = run_transpose(&input); + + let transposed = run_transpose(&input); + for grp in 0..2u32 { for sub in 0..8u32 { for r in 0..8u32 { let want = (sub * 8 + r) as f32 + grp as f32 * 1000.0; - let got = out[(grp * 64 + sub * 8 + r) as usize]; + let got = transposed[(grp * 64 + sub * 8 + r) as usize]; assert_eq!(got, want, "group {grp} lane {sub} slot {r}"); } } diff --git a/av-denoise-core/src/collab/tests/transforms.rs b/av-denoise-core/src/collab/tests/transforms.rs index aebf0cb..74e15f4 100644 --- a/av-denoise-core/src/collab/tests/transforms.rs +++ b/av-denoise-core/src/collab/tests/transforms.rs @@ -5,22 +5,15 @@ use crate::collab::MAX_K; use crate::collab::kernels::transforms::*; #[cube(launch_unchecked)] -fn safe_reciprocal_probe(denom: &Array, floor: &Array, out: &mut Array, n: u32) { +fn safe_reciprocal_probe(denom: &Array, floor: &Array, out: &mut Array, count: u32) { let tid = ABSOLUTE_POS_X; - if tid < n { + if tid < count { out[tid as usize] = safe_reciprocal(denom[tid as usize], floor[tid as usize]); } } -/// Pins `safe_reciprocal`'s own contract directly, rather than only -/// through whatever a caller's `f32::max` happens to do with a `NaN` on -/// this particular GPU backend. -/// -/// `collab_fused` and `collab_aggregate` both reach this same function -/// for their own weight and normalisation divisions, so a probe of the -/// function itself covers every call site at once, and does not depend -/// on a real denominator ever going non-finite in one of those larger -/// kernels to exercise the guard. +/// Probes `safe_reciprocal` itself, so the test checks the guard without relying on how a caller's +/// `f32::max` treats a `NaN` on this backend. #[test] fn safe_reciprocal_is_zero_for_a_non_finite_denominator_and_ordinary_otherwise() { let client = make_client(); @@ -34,93 +27,91 @@ fn safe_reciprocal_is_zero_for_a_non_finite_denominator_and_ordinary_otherwise() -3.0f32, ]; let floor = vec![1e-12f32; 6]; - let n = denom.len(); + let count = denom.len(); - let denom_buf = client.create_from_slice(f32::as_bytes(&denom)); - let floor_buf = client.create_from_slice(f32::as_bytes(&floor)); - let out_buf = client.empty(n * size_of::()); + let denom_bytes = f32::as_bytes(&denom); + let floor_bytes = f32::as_bytes(&floor); + let denom_buf = client.create_from_slice(denom_bytes); + let floor_buf = client.create_from_slice(floor_bytes); + let out_buf = client.empty(count * size_of::()); unsafe { safe_reciprocal_probe::launch_unchecked::( &client, CubeCount::new_1d(1), - CubeDim::new_1d(n as u32), - ArrayArg::from_raw_parts(denom_buf, n), - ArrayArg::from_raw_parts(floor_buf, n), - ArrayArg::from_raw_parts(out_buf.clone(), n), - n as u32, + CubeDim::new_1d(count as u32), + ArrayArg::from_raw_parts(denom_buf, count), + ArrayArg::from_raw_parts(floor_buf, count), + ArrayArg::from_raw_parts(out_buf.clone(), count), + count as u32, ); } - let bytes = client.read_one(out_buf).expect("safe_reciprocal readback failed"); - let out = f32::from_bytes(&bytes)[..n].to_vec(); + let out_bytes = client.read_one(out_buf).expect("safe_reciprocal readback failed"); + let reciprocals = f32::from_bytes(&out_bytes)[..count].to_vec(); assert_eq!( - out[0], 0.0, + reciprocals[0], 0.0, "a NaN denominator must yield exactly 0, got {}", - out[0] + reciprocals[0] ); assert_eq!( - out[1], 0.0, + reciprocals[1], 0.0, "a positive-infinite denominator must yield exactly 0, got {}", - out[1] + reciprocals[1] ); assert_eq!( - out[2], 0.0, + reciprocals[2], 0.0, "a negative-infinite denominator must yield exactly 0, got {}", - out[2] + reciprocals[2] ); assert_eq!( - out[3], 1e12, + reciprocals[3], 1e12, "an ordinary zero denominator floors to 1e12, got {}", - out[3] + reciprocals[3] ); assert!( - (out[4] - 0.2).abs() < 1e-6, + (reciprocals[4] - 0.2).abs() < 1e-6, "1 / max(5, 1e-12) should be 0.2, got {}", - out[4] + reciprocals[4] ); - // A negative but finite denominator is not a case any real caller - // hits (every caller here sums non-negative terms), but it still - // has to stay finite rather than flip the sign of the result. The - // floor outweighs it the same way it outweighs a legitimate zero. + // Callers only sum non-negative terms, but a negative finite denominator must still stay finite + // rather than flip the result's sign. The floor outweighs it the same way it outweighs a zero. assert_eq!( - out[5], 1e12, + reciprocals[5], 1e12, "a negative but finite denominator floors the same way zero does, got {}", - out[5] + reciprocals[5] ); } -/// Runs the same three-level variance ladder [`collab_fused`] runs, so -/// the host mirror is checked against the levels and pairing order the -/// shipped kernel actually applies. -/// -/// [`collab_fused`]: crate::collab::kernels::fused::collab_fused +/// Runs the same three-level variance ladder as +/// [collab_fused](crate::collab::kernels::fused::collab_fused), in the same pairing order. #[cube(launch_unchecked)] fn variance_ladder_kernel(input: &Array, k_use: u32, output: &mut Array) { - let mut v = Array::::new(MAX_K as usize); + let mut variances = Array::::new(MAX_K as usize); #[unroll] for k in 0..MAX_K { - v[k as usize] = input[k as usize]; + variances[k as usize] = input[k as usize]; } if k_use >= 8u32 { - variance_reg_level(&mut v, 8u32); + variance_reg_level(&mut variances, 8u32); } if k_use >= 4u32 { - variance_reg_level(&mut v, 4u32); + variance_reg_level(&mut variances, 4u32); } if k_use >= 2u32 { - variance_reg_level(&mut v, 2u32); + variance_reg_level(&mut variances, 2u32); } #[unroll] for k in 0..MAX_K { - output[k as usize] = v[k as usize]; + output[k as usize] = variances[k as usize]; } } fn run_variance_ladder(sig2: &[f32; 8], k_use: u32) -> Vec { let client = make_client(); - let input_buf = client.create_from_slice(f32::as_bytes(sig2)); + let input_bytes = f32::as_bytes(sig2); + let input_buf = client.create_from_slice(input_bytes); let output_buf = client.empty(8 * size_of::()); unsafe { @@ -134,18 +125,15 @@ fn run_variance_ladder(sig2: &[f32; 8], k_use: u32) -> Vec { ); } - let bytes = client + let output_bytes = client .read_one(output_buf) .expect("variance ladder readback failed"); - f32::from_bytes(&bytes)[..8].to_vec() + + f32::from_bytes(&output_bytes)[..8].to_vec() } -/// Pins the GPU ladder against the host mirror [`haar_variance_ladder`], -/// for every valid `k_use`, with non-uniform input variances. -/// -/// Uniform input would leave the ladder at a fixed point regardless of -/// pairing order, so it could not catch a level or pairing mismatch -/// between the two implementations. Only a non-uniform input can. +/// Uniform input sits at a fixed point under any pairing order, so only non-uniform input can catch +/// a level or pairing mismatch against [haar_variance_ladder]. #[test] fn gpu_variance_ladder_matches_the_host_mirror() { let sig2 = [0.7f32, 1.3, 0.2, 2.5, 0.05, 3.0, 1.1, 0.4]; diff --git a/av-denoise-core/src/denoiser.rs b/av-denoise-core/src/denoiser.rs deleted file mode 100644 index a06e461..0000000 --- a/av-denoise-core/src/denoiser.rs +++ /dev/null @@ -1,2433 +0,0 @@ -use std::collections::VecDeque; - -use cubecl::Runtime; -use cubecl::prelude::ComputeClient; - -use crate::accelerate::Accelerator; -use crate::device::Device; -use crate::nl4d::grain::GrainChunk; -use crate::nl4d::{Nl4dDenoiser, Nl4dParams}; -#[cfg(test)] -use crate::nlmeans::MotionEstimation; -use crate::nlmeans::{ - ChannelMode, - Depth, - HqParams, - MotionCompensationMode, - MotionSearch, - NlmDenoiser, - NlmParams, - Pending, - PrefilterMode, - TryWait, - hq_default_strength, - validate_dimensions, -}; -use crate::sniff::sniff_best_accelerator; - -/// How a [`Denoiser`] should be set up. -/// -/// Build one with `DenoiserOptions::builder()`. Every field has a -/// default, so only the parts you care about need naming. -/// -/// Only the settings every algorithm reads live here. Everything else -/// belongs to whichever [`Algorithm`] variant actually uses it. -#[derive(Debug, Clone, bon::Builder)] -pub struct DenoiserOptions { - /// Which channels of the frame to denoise. - #[builder(default = ChannelMode::Yuv)] - pub channel_mode: ChannelMode, - /// Whether to clean each frame on its own or across a temporal - /// window. - #[builder(default = DenoisingMode::Spacial)] - pub mode: DenoisingMode, - /// Which algorithm to run, along with the settings only that - /// algorithm reads. - #[builder(default)] - pub algorithm: Algorithm, - /// What format denoised frames come back in. - #[builder(default = OutputFormat::F32)] - pub output_format: OutputFormat, -} - -/// What a denoiser hands back from [`Denoiser::recv_frame`], -/// [`Denoiser::try_recv_frame`] and [`Denoiser::flush`]. -/// -/// The GPU only ever holds normalised `f32`. [`Depth`] is a wire concept, -/// so a denoiser that returns wire bytes is told its depth when it is -/// built rather than at each call. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum OutputFormat { - /// Normalised `f32`, one value per channel per pixel. - F32, - /// Wire bytes at this depth, quantised on the GPU. - Wire { depth: Depth }, -} - -/// One denoised frame, in whichever format its denoiser was built for. -#[derive(Debug, Clone, PartialEq)] -pub enum FrameOutput { - /// Normalised `f32`, `width * height * channels` values. - F32(Vec), - /// Wire bytes, interleaved except for a chroma pair, which is laid - /// out as U's whole region followed by V's. - Wire(Vec), -} - -impl FrameOutput { - /// The `f32` frame, or `None` if this came from a wire-mode denoiser. - pub fn into_f32(self) -> Option> { - match self { - Self::F32(v) => Some(v), - Self::Wire(_) => None, - } - } - - /// The wire bytes, or `None` if this came from an `f32`-mode denoiser. - pub fn into_wire(self) -> Option> { - match self { - Self::Wire(v) => Some(v), - Self::F32(_) => None, - } - } - - /// Borrows the `f32` frame, or `None` if this came from a wire-mode - /// denoiser. - pub fn as_f32(&self) -> Option<&[f32]> { - match self { - Self::F32(v) => Some(v), - Self::Wire(_) => None, - } - } - - /// Borrows the wire bytes, or `None` if this came from an `f32`-mode - /// denoiser. - pub fn as_wire(&self) -> Option<&[u8]> { - match self { - Self::Wire(v) => Some(v), - Self::F32(_) => None, - } - } -} - -/// Which denoising algorithm to run. -/// -/// Each variant carries its own settings, so a knob one algorithm has no -/// use for cannot be set on it. -#[derive(Debug, Copy, Clone, PartialEq)] -pub enum Algorithm { - /// The fast NLMeans path, with fixed weighting and no noise - /// measurement. - Nlmeans(NlmeansOptions), - /// NLMeans with its weighting matched to the measured noise level. - /// - /// This also uses a different default `strength`, one that adapts to - /// the temporal radius and the plane being denoised. See - /// [`crate::nlmeans::hq_default_strength`]. - NlmeansHq(NlmeansHqOptions), - /// Groups 8x8 patches across the motion-compensated temporal window - /// itself, rather than filtering with NLM first and grouping within - /// one frame afterward. - /// - /// No NLM weighting pass ever runs, so none of the NLM knobs appear - /// on [`Nl4dOptions`]. - Nl4d(Nl4dOptions), -} - -impl Default for Algorithm { - fn default() -> Self { - Self::Nlmeans(NlmeansOptions::default()) - } -} - -/// Settings for [`Algorithm::Nlmeans`]. -#[derive(Debug, Copy, Clone, Default, PartialEq)] -pub struct NlmeansOptions { - /// Which reference image the NLM weights are computed against. - /// - /// `None`, the default, compares patches on the noisy input - /// directly. Every other mode costs one extra GPU pass per frame. - pub prefilter: PrefilterMode, - /// Whether temporal denoising follows motion between frames. - /// - /// `None`, the default, turns motion compensation off. `Mvtools` - /// warps temporal neighbours into line with the centre frame before - /// the NLM weighting runs. - /// - /// Only has an effect when [`DenoiserOptions::mode`] is - /// `Temporal { .. }`. - pub motion_compensation: MotionCompensationMode, - /// Overrides for the NLM search radius, patch radius, strength, and - /// self-weight. - pub tuning: NlmTuning, -} - -/// Settings for [`Algorithm::NlmeansHq`]. -#[derive(Debug, Copy, Clone, Default, PartialEq)] -pub struct NlmeansHqOptions { - /// Everything the fast path takes, which HQ takes too. - pub nlm: NlmeansOptions, - /// The noise measurement and confidence weighting HQ adds on top. - pub hq: HqParams, -} - -/// Settings for [`Algorithm::Nl4d`]. -/// -/// nl4d runs the HQ front end only for its machinery, the frame ring, -/// the motion field, and the noise estimate. Nothing weights or averages -/// patches the NLM way, so the NLM knobs are absent here and the fields -/// below are the whole surface. -/// -/// The temporal radius comes from [`DenoiserOptions::mode`], which has to -/// be `Temporal { .. }`. Motion tracking is always on, because the -/// grouping kernel reads the motion field and confidence scores it -/// produces. -/// -/// `lambda_ht` has a per-plane default. `None` resolves through -/// [`nl4d_default_lambda_ht`] once the plane being denoised is known. -/// `lambda_ht_scale` then multiplies whichever value that resolves to. -/// -/// Every other default comes from [`Nl4dParams::default`]. -#[derive(Debug, Copy, Clone, PartialEq)] -pub struct Nl4dOptions { - /// How motion between frames is tracked. - pub motion: MotionSearch, - /// A fixed noise standard deviation in `[0, 1]` units, replacing the - /// automatic per-frame estimate. - /// - /// `None`, the default, measures the noise in each pushed frame and - /// smooths it over time. - pub sigma: Option, - /// A multiplier applied to the measured noise level before anything - /// reads it. Defaults to 1.0. - /// - /// This does nothing when `sigma` pins the noise level, because the - /// estimator never runs in that case. - pub sigma_scale: f32, - /// A multiplier on the per-block mismatch threshold, which sets how - /// much extra SAD a block tolerates before its confidence starts to - /// fall. Defaults to 1.0. - /// - /// Higher values tolerate larger mismatches. - pub thsad_scale: f32, - /// Half-width of the refine window searched around each neighbour - /// frame's motion-predicted position, in `1..=4`. Defaults to 2. - pub refine: u32, - /// Half-width of the spatial candidate window searched in the centre - /// frame, in `1..=16`. Defaults to 9. - pub spatial_radius: u32, - /// Hard-threshold multiplier on the propagated coefficient sigma. - /// Higher removes more noise and more fine detail. - /// - /// `None` resolves through [`nl4d_default_lambda_ht`], which returns - /// a different value for luma than for chroma. - pub lambda_ht: Option, - /// A multiplier applied to the resolved `lambda_ht`. Defaults to - /// 1.0. - /// - /// It scales an explicit `lambda_ht` and the calibrated per-plane - /// default alike, so one value moves both planes together. Has to - /// be finite and in `[0.1, 10.0]`. - pub lambda_ht_scale: f32, - /// The confidence floor below which a whole neighbour block is - /// skipped rather than scored, in `[0, 1)`. Defaults to 0.05. A - /// block below the floor is never scored, and a volume left short - /// of frames by the skip makes its group filter from the centre - /// frame alone. - pub c_min: f32, - /// The `beta` of the Kaiser window each filtered patch is tapered - /// with as it is aggregated. Defaults to 2.0. `0.0` is uniform - /// aggregation. See [`crate::nl4d::Nl4dParams::kaiser_beta`]. - pub kaiser_beta: f32, - /// Estimates noise fresh from each frame's own window instead of - /// smoothing it across the whole stream's history. Defaults to - /// `false`, matching every calibrated preset. - /// - /// `av-denoise-vs` turns this on unconditionally, because a - /// VapourSynth filter has to return the same pixels for a frame no - /// matter what order frames were requested in, and history-dependent - /// estimation breaks that guarantee under random access. See - /// [`HqParams::windowed_noise_estimation`]. - pub windowed_noise_estimation: bool, - /// See [`crate::nl4d::Nl4dParams::field_lambda`]. - pub field_lambda: f32, - /// See [crate::nl4d::Nl4dParams::noise_map]. - pub noise_map: bool, - /// See [crate::nl4d::Nl4dParams::flat_boost]. - pub flat_boost: f32, - /// See [crate::nl4d::Nl4dParams::chroma_flat_boost]. - pub chroma_flat_boost: f32, - /// See [crate::nl4d::Nl4dParams::shadow_soften]. - pub shadow_soften: f32, - /// See [crate::nl4d::Nl4dParams::flat_texture_cut]. - pub flat_texture_cut: f32, - /// See [crate::nl4d::Nl4dParams::pooled_threshold]. - pub pooled_threshold: bool, - /// See [crate::nl4d::Nl4dParams::grain_export]. - /// - /// Measured chunks stay on the GPU until drained, so callers drain after each flush. - pub grain_export: bool, -} - -impl Default for Nl4dOptions { - fn default() -> Self { - let defaults = Nl4dParams::default(); - let hq = HqParams::default(); - Self { - motion: MotionSearch::default(), - sigma: hq.sigma_override, - sigma_scale: hq.sigma_scale, - thsad_scale: hq.thsad_scale, - refine: defaults.refine, - spatial_radius: defaults.spatial_radius, - // Resolved per plane by `nl4d_default_lambda_ht` at - // construction time, once the plane being denoised is - // known. - lambda_ht: None, - lambda_ht_scale: 1.0, - c_min: defaults.c_min, - kaiser_beta: defaults.kaiser_beta, - windowed_noise_estimation: false, - field_lambda: defaults.field_lambda, - noise_map: defaults.noise_map, - flat_boost: defaults.flat_boost, - chroma_flat_boost: defaults.chroma_flat_boost, - shadow_soften: defaults.shadow_soften, - flat_texture_cut: defaults.flat_texture_cut, - pooled_threshold: defaults.pooled_threshold, - grain_export: defaults.grain_export, - } - } -} - -impl Nl4dOptions { - /// The front end's HQ parameters for this configuration. - /// - /// `temporal_confidence` is always on, because the grouping kernel - /// reads the confidence scores it produces. The two strength-related - /// switches keep their defaults, since nl4d never runs a weighting - /// pass for them to affect. - fn to_hq_params(self) -> HqParams { - HqParams { - sigma_override: self.sigma, - sigma_scale: self.sigma_scale, - thsad_scale: self.thsad_scale, - temporal_confidence: true, - windowed_noise_estimation: self.windowed_noise_estimation, - ..HqParams::default() - } - } -} - -/// The default `lambda_ht` for nl4d's hard-threshold stage, per plane. -/// -/// `lambda_ht` is how many standard deviations of estimated noise a -/// transform coefficient has to clear to survive. Raising it removes more -/// noise and more fine detail with it, so the value is a trade rather -/// than an optimum. -/// -/// Luma and chroma values are picked by eye from a ladder of renders against real -/// film grain, accepting more lost detail in exchange for less remaining -/// noise on heavy grain. The reason why we're going a bit heavier on high noise is because -/// the encoders end up reducing that detail _more_ than the denoiser does if -/// that extra entropy remains in and overall produces a worse final image. -/// -/// `ChannelMode::Yuv` reads the luma value, on the same "a fused pass is -/// dominated by luma" assumption [`hq_default_strength`] -/// makes for its own Yuv case. -/// -/// Luma and the fused Yuv mode use 4.158, and chroma uses 3.234. -pub fn nl4d_default_lambda_ht(channels: ChannelMode) -> f32 { - match channels { - ChannelMode::Luma | ChannelMode::Yuv => 4.158, - ChannelMode::Chroma => 3.234, - } -} - -/// The pooled threshold at each plane's default lambda, calibrated on real grain. -const NL4D_POOLED_THRESHOLD: f32 = 2.42; - -/// The ratio of nl4d's pooled threshold to its lambda for one plane. -/// -/// At the default lambda this gives the calibrated pooled threshold, and it scales with any -/// other lambda. -pub fn nl4d_pool_ratio(channels: ChannelMode) -> f32 { - NL4D_POOLED_THRESHOLD / nl4d_default_lambda_ht(channels) -} - -/// Resolves `Nl4dOptions.lambda_ht` for one plane, falling back to -/// [`nl4d_default_lambda_ht`] when the caller left it unset, then -/// applies `lambda_ht_scale`. -/// -/// The scale multiplies an explicit value and the calibrated default -/// alike, so it moves both planes together whether or not one of them -/// is pinned. -/// -/// The range check lives here rather than in [`Nl4dParams`], -/// which only ever sees the product. A scale of 0 would surface there as -/// a complaint about `lambda_ht`, naming a knob the caller never set. -fn resolve_lambda_ht(opts: &Nl4dOptions, channels: ChannelMode) -> Result { - if !(opts.lambda_ht_scale.is_finite() && (0.1..=10.0).contains(&opts.lambda_ht_scale)) { - return Err(format!( - "lambda_ht_scale must be finite and in [0.1, 10.0], got {}", - opts.lambda_ht_scale - )); - } - - let lambda_ht = opts.lambda_ht.unwrap_or_else(|| nl4d_default_lambda_ht(channels)); - - Ok(lambda_ht * opts.lambda_ht_scale) -} - -/// Speed vs quality dial. -/// -/// Each denoising family reads the same dial and fills in its own knobs -/// from it. For `nlmeans` that is [`nlmeans_variant_for`], -/// [`nlmeans_temporal_radius_for`], and [`nlmeans_search_radius_for`]. -/// For `nl4d` it is [`nl4d_temporal_radius_for`] and -/// [`nl4d_spatial_radius_for`]. -/// -/// Both front ends parse the same names from this one type, so a preset -/// resolves to the same dials everywhere it is used. -#[derive(Debug, Copy, Clone, Default, PartialEq, Eq, strum_macros::EnumString)] -#[strum(ascii_case_insensitive)] -pub enum Preset { - /// Fastest and lowest quality. - Veryfast, - /// One step up from `veryfast`. - Fast, - /// The default, favouring quality over speed. - #[default] - Base, - /// One step down from `veryslow`. - Slow, - /// Slowest and highest quality. - Veryslow, -} - -/// Which nlmeans implementation a preset, or an explicit choice, selects. -#[derive(Debug, Copy, Clone, PartialEq, Eq, strum_macros::EnumString)] -#[strum(ascii_case_insensitive)] -pub enum NlmeansVariant { - /// The fast path. Fixed weighting, no noise measurement. - Fast, - /// Quality focused. Calibrates its weighting to the noise level, - /// measured automatically per frame. - Hq, -} - -/// Which [`NlmeansVariant`] a preset runs. -pub fn nlmeans_variant_for(preset: Preset) -> NlmeansVariant { - match preset { - Preset::Veryfast => NlmeansVariant::Fast, - Preset::Fast | Preset::Base | Preset::Slow | Preset::Veryslow => NlmeansVariant::Hq, - } -} - -/// How many neighbouring frames on each side `nlmeans` looks at, at a -/// preset. -pub fn nlmeans_temporal_radius_for(preset: Preset) -> u32 { - match preset { - Preset::Veryfast => 0, - Preset::Fast => 1, - Preset::Base => 2, - Preset::Slow => 4, - Preset::Veryslow => 8, - } -} - -/// How far `nlmeans` looks for similar patches inside a frame, at a -/// preset. -pub fn nlmeans_search_radius_for(preset: Preset) -> u32 { - match preset { - Preset::Veryfast | Preset::Fast | Preset::Base => 2, - Preset::Slow | Preset::Veryslow => 4, - } -} - -/// How far the temporal window reaches at each preset, for `nl4d`. -/// -/// Unlike `nlmeans`, `veryfast` keeps a 1-frame window rather than -/// dropping to 0, because nl4d has nothing to do without neighbouring -/// frames to group against. -pub fn nl4d_temporal_radius_for(preset: Preset) -> u32 { - match preset { - Preset::Veryfast | Preset::Fast => 1, - Preset::Base => 2, - Preset::Slow => 4, - Preset::Veryslow => 8, - } -} - -/// How wide the centre frame's candidate search is at each preset, for -/// `nl4d`. -/// -/// `veryfast` shares its temporal radius with `fast`, so this is what -/// separates them. The window covers `(2 * radius + 1)^2` positions, so -/// 6 searches a little over half the candidates 9 does. -/// -/// Every preset from `fast` up uses the library default. Widening it -/// further at the slow end costs quadratically and has not been measured -/// to be worth it. -pub fn nl4d_spatial_radius_for(preset: Preset) -> u32 { - match preset { - Preset::Veryfast => 6, - Preset::Fast | Preset::Base | Preset::Slow | Preset::Veryslow => { - Nl4dOptions::default().spatial_radius - }, - } -} - -/// Whether a frame is cleaned on its own or alongside its neighbours. -#[derive(Debug, Copy, Clone, Eq, PartialEq)] -pub enum DenoisingMode { - /// Cleans each frame using only its own pixels. - Spacial, - /// Cleans each frame using a window of `2 * radius + 1` frames. - Temporal { radius: u32 }, -} - -/// NLM tuning knobs. -/// -/// Every field is optional. Whatever is left unset falls back to the -/// library default. -#[derive(Debug, Copy, Clone, Default, PartialEq)] -pub struct NlmTuning { - pub search_radius: Option, - pub patch_radius: Option, - pub strength: Option, - pub self_weight: Option, -} - -impl DenoiserOptions { - /// Turns this option set into the low-level [`NlmParams`] a backend - /// denoiser is built from. - /// - /// For nl4d this describes the front end only, since nl4d's own - /// grouping stage is configured from [`Nl4dOptions`] separately in - /// [`build_engine`]. - /// - /// Whichever default `strength` applies is folded in here. For the - /// HQ algorithm that comes from - /// [`crate::nlmeans::hq_default_strength`]. - /// - /// This is public so callers building per-plane options, and tests, - /// can read the resolved values without building a real `Denoiser`. - #[doc(hidden)] - pub fn to_nlm_params(&self) -> NlmParams { - let temporal_radius = match self.mode { - DenoisingMode::Spacial => 0, - DenoisingMode::Temporal { radius } => radius, - }; - - match self.algorithm { - Algorithm::Nlmeans(opts) => self.nlm_params_for(opts, None, temporal_radius), - Algorithm::NlmeansHq(opts) => self.nlm_params_for(opts.nlm, Some(opts.hq), temporal_radius), - // nl4d never runs a weighting pass, so `strength`, - // `search_radius`, `patch_radius`, and `self_weight` stay at - // their library defaults and no prefilter is built. - Algorithm::Nl4d(opts) => NlmParams { - channels: self.channel_mode, - motion_compensation: opts.motion.into(), - temporal_radius, - hq: Some(opts.to_hq_params()), - ..NlmParams::default() - }, - } - } - - /// [`Self::to_nlm_params`] for whichever of the two NLM algorithms - /// is running, with `hq` set only for the quality one. - fn nlm_params_for(&self, opts: NlmeansOptions, hq: Option, temporal_radius: u32) -> NlmParams { - // An explicit `strength` always wins, whether it came straight - // from `NlmTuning` or from a per-plane override the caller - // already folded in. - // - // Otherwise the default depends on `auto_strength`. With it on, - // HQ reads `strength` as a multiplier on the measured noise - // level, so it needs its own calibrated default rather than the - // fast path's absolute FFmpeg-style one. That calibrated default - // also varies with the temporal radius and with the plane - // `channel_mode` names, because each per-plane `Denoiser` - // carries its own channel mode. - // - // With auto-strength off, HQ reads `strength` as an absolute - // value just like the fast path, so it falls back to the same - // absolute default. - let strength = opts.tuning.strength.unwrap_or(match hq { - Some(hq) if hq.auto_strength => hq_default_strength(self.channel_mode, temporal_radius), - _ => NlmParams::default().strength, - }); - - let defaults = NlmParams::default(); - NlmParams { - channels: self.channel_mode, - prefilter: opts.prefilter, - motion_compensation: opts.motion_compensation, - temporal_radius, - hq, - strength, - search_radius: opts.tuning.search_radius.unwrap_or(defaults.search_radius), - patch_radius: opts.tuning.patch_radius.unwrap_or(defaults.patch_radius), - self_weight: opts.tuning.self_weight.unwrap_or(defaults.self_weight), - } - } -} - -/// Errors reported by the high-level [`Denoiser`]. -#[derive(Debug, thiserror::Error)] -pub enum DenoiserError { - /// An earlier denoised frame has not been collected yet, so pushing - /// again would overwrite it in the double-buffered output slot. - /// - /// Call [`Denoiser::recv_frame`] or [`Denoiser::try_recv_frame`], - /// then retry the same `push_frame` call. - #[error("denoiser queue is full, collect the pending frame before pushing more")] - QueueFull, - /// An earlier call failed, so how many frames are in flight is no - /// longer known and later output would not line up with its input. - /// - /// Call [`Denoiser::reset_stream`] to start a fresh stream, or drop - /// the denoiser. Either one settles a frame that - /// [`Denoiser::try_recv_frame`] has already polled, so it can block - /// for that readback's remaining latency. - #[error("denoiser failed earlier, reset the stream before using it again")] - Poisoned, - /// None of the accelerators in the priority list could be started. - #[error("no accelerator from the priority list is available")] - NoAcceleratorAvailable, - /// Anything else, wrapping the internal `anyhow` errors raised by - /// kernel dispatch and readback. - #[error(transparent)] - Other(#[from] anyhow::Error), -} - -/// Either denoiser a `Backend` runtime arm can hold. -/// -/// This keeps `Backend`'s own match arms at one line each. Without it, -/// adding a second denoiser type would multiply the runtime arms instead -/// of fanning out once here. -enum Engine { - Nlm(Box>), - Nl4d(Box>), -} - -impl Engine { - fn is_nl4d(&self) -> bool { - matches!(self, Self::Nl4d(_)) - } - - fn mark_continuation(&mut self) { - if let Self::Nl4d(d) = self { - d.mark_continuation(); - } - } - - fn push_frame(&mut self, frame: &[f32]) { - match self { - Self::Nlm(d) => d.push_frame(frame), - Self::Nl4d(d) => d.push_frame(frame), - } - } - - fn push_frame_wire(&mut self, planes: &[&[u8]], depth: Depth) { - match self { - Self::Nlm(d) => d.push_frame_wire(planes, depth), - Self::Nl4d(d) => d.push_frame_wire(planes, depth), - } - } - - fn denoise_submit(&mut self) -> Result>, anyhow::Error> { - match self { - Self::Nlm(d) => d.denoise_submit(), - // `Nl4dDenoiser::denoise_submit` already returns - // `DenoiserError` rather than `anyhow::Error`, so this leans - // on `DenoiserError`'s own `anyhow::Error` conversion instead - // of re-wrapping it. - Self::Nl4d(d) => d.denoise_submit().map_err(anyhow::Error::from), - } - } - - #[cfg(test)] - fn wire_outputs(&self) -> Option<&[cubecl::server::Handle; 2]> { - match self { - Self::Nlm(d) => d.wire_outputs_for_test(), - Self::Nl4d(d) => d.wire_outputs_for_test(), - } - } - - fn flush(&mut self, sink: impl FnMut(&FrameOutput)) -> Result<(), anyhow::Error> { - match self { - Self::Nlm(d) => d.flush(sink), - Self::Nl4d(d) => d.flush(sink).map_err(anyhow::Error::from), - } - } - - fn reset_stream(&mut self) { - match self { - Self::Nlm(d) => d.reset_stream_state(), - Self::Nl4d(d) => d.reset_stream(), - } - } - - fn drain_grain_chunks(&mut self) -> Result, DenoiserError> { - match self { - Self::Nlm(_) => Ok(Vec::new()), - Self::Nl4d(d) => d.drain_grain_chunks(), - } - } -} - -/// Builds whichever [`Engine`] `algorithm` calls for. -/// -/// `Algorithm::Nl4d` carries its own grouping tuning, which is not part -/// of `NlmParams`, so it is read from `algorithm` directly rather than -/// from `params`. This is also where an unset `lambda_ht` picks up its -/// calibrated per-plane default (`resolve_lambda_ht`), the same way -/// `to_nlm_params` resolves HQ's calibrated `strength`, since this is -/// the first point construction has both `opts` and `params.channels` -/// together. -fn build_engine( - client: &ComputeClient, - algorithm: &Algorithm, - params: NlmParams, - width: u32, - height: u32, - output_format: OutputFormat, -) -> Result, DenoiserError> { - match algorithm { - Algorithm::Nl4d(opts) => { - // nl4d groups patches across neighbouring frames, so there - // is nothing for it to do without a temporal window. - if params.temporal_radius == 0 { - return Err(DenoiserError::Other(anyhow::anyhow!( - "nl4d needs a temporal window, set DenoiserOptions::mode to \ - DenoisingMode::Temporal" - ))); - } - - let lambda_ht = resolve_lambda_ht(opts, params.channels) - .map_err(|e| DenoiserError::Other(anyhow::anyhow!(e)))?; - let nl4d_params = Nl4dParams { - temporal_radius: params.temporal_radius, - nlm: params, - refine: opts.refine, - spatial_radius: opts.spatial_radius, - lambda_ht, - c_min: opts.c_min, - kaiser_beta: opts.kaiser_beta, - field_lambda: opts.field_lambda, - noise_map: opts.noise_map, - flat_boost: opts.flat_boost, - chroma_flat_boost: opts.chroma_flat_boost, - shadow_soften: opts.shadow_soften, - flat_texture_cut: opts.flat_texture_cut, - pooled_threshold: opts.pooled_threshold, - grain_export: opts.grain_export, - }; - let denoiser = - Nl4dDenoiser::with_output_format(client, nl4d_params, width, height, output_format) - .map_err(|e| DenoiserError::Other(anyhow::anyhow!(e)))?; - Ok(Engine::Nl4d(Box::new(denoiser))) - }, - Algorithm::Nlmeans(_) | Algorithm::NlmeansHq(_) => Ok(Engine::Nlm(Box::new( - NlmDenoiser::with_output_format(client, params, width, height, output_format), - ))), - } -} - -enum Backend { - #[cfg(feature = "cuda")] - Cuda(Engine), - #[cfg(feature = "rocm")] - Rocm(Engine), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Wgpu(Engine), -} - -impl Backend { - fn is_nl4d(&self) -> bool { - match self { - #[cfg(feature = "cuda")] - Self::Cuda(e) => e.is_nl4d(), - #[cfg(feature = "rocm")] - Self::Rocm(e) => e.is_nl4d(), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Self::Wgpu(e) => e.is_nl4d(), - } - } - - fn mark_continuation(&mut self) { - match self { - #[cfg(feature = "cuda")] - Self::Cuda(e) => e.mark_continuation(), - #[cfg(feature = "rocm")] - Self::Rocm(e) => e.mark_continuation(), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Self::Wgpu(e) => e.mark_continuation(), - } - } - - #[cfg(test)] - fn wire_outputs(&self) -> Option<&[cubecl::server::Handle; 2]> { - match self { - #[cfg(feature = "cuda")] - Self::Cuda(e) => e.wire_outputs(), - #[cfg(feature = "rocm")] - Self::Rocm(e) => e.wire_outputs(), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Self::Wgpu(e) => e.wire_outputs(), - } - } -} - -enum BackendPending { - #[cfg(feature = "cuda")] - Cuda(Pending), - #[cfg(feature = "rocm")] - Rocm(Pending), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Wgpu(Pending), -} - -impl BackendPending { - fn wait(self) -> Result { - match self { - #[cfg(feature = "cuda")] - Self::Cuda(p) => p.wait(), - #[cfg(feature = "rocm")] - Self::Rocm(p) => p.wait(), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Self::Wgpu(p) => p.wait(), - } - } - - /// Polls the readback once. `Ok(Ok(frame))` is a landed frame, - /// `Ok(Err(self))` is a readback still in flight. - fn try_wait(self) -> Result, anyhow::Error> { - match self { - #[cfg(feature = "cuda")] - Self::Cuda(p) => match p.try_wait()? { - TryWait::Ready(frame) => Ok(Ok(frame)), - TryWait::NotReady(p) => Ok(Err(Self::Cuda(p))), - }, - #[cfg(feature = "rocm")] - Self::Rocm(p) => match p.try_wait()? { - TryWait::Ready(frame) => Ok(Ok(frame)), - TryWait::NotReady(p) => Ok(Err(Self::Rocm(p))), - }, - #[cfg(any(feature = "vulkan", feature = "metal"))] - Self::Wgpu(p) => match p.try_wait()? { - TryWait::Ready(frame) => Ok(Ok(frame)), - TryWait::NotReady(p) => Ok(Err(Self::Wgpu(p))), - }, - } - } -} - -/// How many readbacks the high-level [`Denoiser`] keeps in flight at -/// once. -/// -/// This has to match the backend's output-handle count, which is two. -/// Going past it would reuse the oldest pending frame's output handle -/// and quietly corrupt the results. -pub const MAX_PENDING: usize = 2; - -/// How many source frames a windowed operation needs behind and ahead -/// of its target frame, target frame itself not counted in either -/// number. -/// -/// `reseed` needs exactly `behind + 1 + ahead` frames, oldest first, -/// with the target frame sitting at index `behind`. This is what tells -/// a caller like `reseed` how wide a window to build, and it varies by -/// algorithm because nl4d's own cross-frame accumulator needs more -/// forward context than the NLM algorithms do. See -/// [`Denoiser::window_span`]. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct WindowSpan { - /// How many frames older than the target the window must include. - pub behind: usize, - /// How many frames newer than the target the window must include. - pub ahead: usize, - /// How the window is filled where it runs past a clip's ends. - pub edges: EdgePadding, -} - -/// How a windowed algorithm fills a window that runs past a clip's ends. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum EdgePadding { - /// Repeat the boundary frame, so every window has the same length. - Repeat, - /// Stop at the clip's ends and let the denoiser run off-centre passes there. - Shifted, -} - -impl WindowSpan { - /// The full window size this span describes, target frame - /// included: `behind + 1 + ahead`. - pub fn frame_count(&self) -> usize { - self.behind + 1 + self.ahead - } -} - -/// A stateful denoiser that cleans a stream of frames. -/// -/// Push frames in order with [`push_frame`](Self::push_frame) and -/// collect the cleaned ones with [`recv_frame`](Self::recv_frame) or -/// [`try_recv_frame`](Self::try_recv_frame). -/// -/// At the end of the stream call [`flush`](Self::flush) to drain -/// whatever temporal context is left. -/// -/// Input frames are `f32` values in `[0, 1]`, laid out as -/// `width * height * channels`. Output comes back as a [`FrameOutput`] -/// in whichever [`OutputFormat`] the options named. -/// -/// ```no_run -/// use av_denoise_core::accelerate::Accelerator; -/// use av_denoise_core::{ChannelMode, Denoiser, DenoiserOptions, DenoisingMode, Device}; -/// -/// # fn main() -> Result<(), Box> { -/// let options = DenoiserOptions::builder() -/// .channel_mode(ChannelMode::Luma) -/// .mode(DenoisingMode::Temporal { radius: 2 }) -/// .build(); -/// -/// let mut denoiser = Denoiser::create( -/// &[Accelerator::Vulkan], -/// &Device::Default, -/// 1920, -/// 1080, -/// options, -/// )?; -/// -/// let frames: Vec> = read_my_frames(); -/// let mut cleaned: Vec> = Vec::new(); -/// -/// for frame in &frames { -/// denoiser.push_frame(frame)?; -/// -/// // Temporal denoising runs a few frames behind the input, so -/// // there is not always one ready to collect. -/// if let Some(out) = denoiser.recv_frame()? { -/// cleaned.push(out.into_f32().expect("built for f32 output")); -/// } -/// } -/// -/// // Drain the frames still inside the temporal window. -/// denoiser.flush(|out| cleaned.push(out.into_f32().expect("built for f32 output")))?; -/// # Ok(()) -/// # } -/// # fn read_my_frames() -> Vec> { Vec::new() } -/// ``` -pub struct Denoiser { - backend: Backend, - pending: VecDeque, - accelerator: Accelerator, - width: u32, - height: u32, - temporal_radius: u32, - output_format: OutputFormat, - frames_pushed: u32, - /// Set once any call other than a `QueueFull` push has failed, so how - /// many frames are in flight is no longer known. - /// - /// Every entry point refuses to run while this is set. - /// [`Self::reset_stream`] clears it. - poisoned: bool, -} - -impl Denoiser { - /// Tries each accelerator in `accelerators` in order and builds a - /// denoiser on the first one that works. - /// - /// `device` picks a non-default device on the chosen runtime. - /// - /// # Thread stack size - /// - /// cubecl spawns its own per-device worker thread, named - /// `DS{U,D}-…`, and runs GPU kernel codegen on it. That thread gets - /// Rust's default stack, which is `RUST_MIN_STACK` or 2 MiB when - /// that is unset. - /// - /// The windowed NLM kernels unroll their body - /// `(2 * search_radius + 1)^2` times, so a `search_radius` of about - /// 5 or more can overflow the 2 MiB default and abort the process. - /// - /// Callers using a `search_radius` above 4 should call - /// [`crate::raise_codegen_stack_limit`] before any cubecl thread - /// spawns, usually right at the top of `main`. - pub fn create( - accelerators: &[Accelerator], - device: &Device, - width: u32, - height: u32, - options: DenoiserOptions, - ) -> Result { - let accelerator = - sniff_best_accelerator(accelerators, device).ok_or(DenoiserError::NoAcceleratorAvailable)?; - - let params = options.to_nlm_params(); - params.validate()?; - validate_dimensions(width, height)?; - - let temporal_radius = params.temporal_radius; - let backend = build_backend( - accelerator, - device, - &options.algorithm, - params, - width, - height, - options.output_format, - )?; - - Ok(Self { - backend, - pending: VecDeque::with_capacity(MAX_PENDING), - accelerator, - width, - height, - temporal_radius, - output_format: options.output_format, - frames_pushed: 0, - poisoned: false, - }) - } - - /// The accelerator [`sniff_best_accelerator`] picked. - pub fn selected_accelerator(&self) -> Accelerator { - self.accelerator - } - - /// The width passed at construction. - pub fn width(&self) -> u32 { - self.width - } - - /// The height passed at construction. - pub fn height(&self) -> u32 { - self.height - } - - /// The temporal radius the resolved parameters run at. - pub fn temporal_radius(&self) -> u32 { - self.temporal_radius - } - - /// The format every collected frame comes back in. - pub fn output_format(&self) -> OutputFormat { - self.output_format - } - - /// How many frames behind and ahead of a target frame this - /// denoiser needs pushed, in order, to produce that frame's - /// output through [`PlanarDenoiser::reseed`](crate::PlanarDenoiser::reseed). - /// - /// Both NLM algorithms only ever need their own `2 * radius + 1` - /// sliding window, symmetric around the target frame: - /// `WindowSpan { behind: radius, ahead: radius }`. - /// - /// nl4d's cross-frame accumulator scatters every pass's - /// contribution across the `2 * radius + 1` frames the pass - /// reaches, and a frame's own region only starts collecting once - /// the pass that first reaches it, the one centred `radius` frames - /// behind it, has run. That earliest pass is itself only real once - /// the front end's own window is full at that centre, which needs - /// `radius` more frames behind it again. So nl4d needs the target's - /// own `radius`-wide neighbourhood doubled on both sides: - /// `WindowSpan { behind: 2 * radius, ahead: 2 * radius }`. - /// - /// nl4d's windows stop at a clip's ends, while the NLM algorithms - /// repeat the boundary frame. - pub fn window_span(&self) -> WindowSpan { - let radius = self.temporal_radius as usize; - let is_nl4d = self.backend.is_nl4d(); - let span = if is_nl4d { 2 * radius } else { radius }; - let edges = if is_nl4d { - EdgePadding::Shifted - } else { - EdgePadding::Repeat - }; - - WindowSpan { - behind: span, - ahead: span, - edges, - } - } - - /// Uploads one frame into the temporal window. - /// - /// `frame` holds `width * height * channels` `f32` values in - /// `[0, 1]`. - /// - /// Once the window is full and the pipeline has room, this also - /// starts the kernels for the next denoised frame. - /// - /// Up to `MAX_PENDING` outputs can be in flight at once, so the GPU - /// runs one frame's kernels while the previous frame's readback is - /// still travelling. At that ceiling this returns - /// [`DenoiserError::QueueFull`], and the caller has to drain a frame - /// with [`Self::recv_frame`] before pushing more. - /// - /// Any other failure poisons the denoiser, so every further call - /// returns [`DenoiserError::Poisoned`] until [`Self::reset_stream`] - /// clears it. `QueueFull` does not poison, since it is the documented - /// retry signal above. - pub fn push_frame(&mut self, frame: &[f32]) -> Result<(), DenoiserError> { - if self.poisoned { - return Err(DenoiserError::Poisoned); - } - self.push_frame_inner(frame).inspect_err(|err| { - if !matches!(err, DenoiserError::QueueFull) { - self.poisoned = true; - } - }) - } - - fn push_frame_inner(&mut self, frame: &[f32]) -> Result<(), DenoiserError> { - // After `temporal_radius` real pushes the leading-edge mirror - // has primed the window, so the next push produces a pending - // frame. From then on every push takes a pending slot. - let window_full = self.frames_pushed > self.temporal_radius; - if window_full && self.pending.len() >= MAX_PENDING { - return Err(DenoiserError::QueueFull); - } - - match &mut self.backend { - #[cfg(feature = "cuda")] - Backend::Cuda(d) => { - d.push_frame(frame); - if let Some(p) = d.denoise_submit()? { - self.pending.push_back(BackendPending::Cuda(p)); - } - }, - #[cfg(feature = "rocm")] - Backend::Rocm(d) => { - d.push_frame(frame); - if let Some(p) = d.denoise_submit()? { - self.pending.push_back(BackendPending::Rocm(p)); - } - }, - #[cfg(any(feature = "vulkan", feature = "metal"))] - Backend::Wgpu(d) => { - d.push_frame(frame); - if let Some(p) = d.denoise_submit()? { - self.pending.push_back(BackendPending::Wgpu(p)); - } - }, - } - - self.frames_pushed = self.frames_pushed.saturating_add(1); - Ok(()) - } - - /// Uploads one frame held as wire bytes into the temporal window. - /// - /// `planes` holds one `width * height` plane per channel at `depth`, - /// which the GPU normalises and interleaves. The planes run Y, U, V - /// for a fused frame and U, V for a chroma pair. - /// - /// Queueing, poisoning, and the `QueueFull` retry signal work exactly - /// as they do for [`Self::push_frame`]. - pub fn push_frame_wire(&mut self, planes: &[&[u8]], depth: Depth) -> Result<(), DenoiserError> { - if self.poisoned { - return Err(DenoiserError::Poisoned); - } - self.push_frame_wire_inner(planes, depth).inspect_err(|err| { - if !matches!(err, DenoiserError::QueueFull) { - self.poisoned = true; - } - }) - } - - fn push_frame_wire_inner(&mut self, planes: &[&[u8]], depth: Depth) -> Result<(), DenoiserError> { - // The same window accounting `push_frame_inner` does. - let window_full = self.frames_pushed > self.temporal_radius; - if window_full && self.pending.len() >= MAX_PENDING { - return Err(DenoiserError::QueueFull); - } - - match &mut self.backend { - #[cfg(feature = "cuda")] - Backend::Cuda(d) => { - d.push_frame_wire(planes, depth); - if let Some(p) = d.denoise_submit()? { - self.pending.push_back(BackendPending::Cuda(p)); - } - }, - #[cfg(feature = "rocm")] - Backend::Rocm(d) => { - d.push_frame_wire(planes, depth); - if let Some(p) = d.denoise_submit()? { - self.pending.push_back(BackendPending::Rocm(p)); - } - }, - #[cfg(any(feature = "vulkan", feature = "metal"))] - Backend::Wgpu(d) => { - d.push_frame_wire(planes, depth); - if let Some(p) = d.denoise_submit()? { - self.pending.push_back(BackendPending::Wgpu(p)); - } - }, - } - - self.frames_pushed = self.frames_pushed.saturating_add(1); - Ok(()) - } - - /// Uploads one frame held as wire bytes into the temporal window - /// without starting a denoise. - /// - /// The wire counterpart of [`Self::push_frame_priming`]. - /// - /// A stream that starts with a priming push picks up mid-clip, so nl4d runs no head passes for it. - pub fn push_frame_wire_priming(&mut self, planes: &[&[u8]], depth: Depth) -> Result<(), DenoiserError> { - if self.poisoned { - return Err(DenoiserError::Poisoned); - } - - if self.frames_pushed == 0 { - self.backend.mark_continuation(); - } - - match &mut self.backend { - #[cfg(feature = "cuda")] - Backend::Cuda(d) => d.push_frame_wire(planes, depth), - #[cfg(feature = "rocm")] - Backend::Rocm(d) => d.push_frame_wire(planes, depth), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Backend::Wgpu(d) => d.push_frame_wire(planes, depth), - } - - self.frames_pushed = self.frames_pushed.saturating_add(1); - Ok(()) - } - - /// Uploads one frame into the temporal window without starting a - /// denoise. - /// - /// The ring advances exactly as it does for [`Self::push_frame`], so - /// the window still fills, but no kernels are submitted and no - /// output is queued. This is how a caller that can hand over a whole - /// window at once, rather than a strictly ordered stream, fills the - /// window in one go and lets only the last push in it submit. - /// - /// A stream that starts with a priming push picks up mid-clip, so nl4d runs no head passes for it. - /// - /// A failure elsewhere poisons the denoiser, so this refuses to run - /// until [`Self::reset_stream`] clears it. - pub fn push_frame_priming(&mut self, frame: &[f32]) -> Result<(), DenoiserError> { - if self.poisoned { - return Err(DenoiserError::Poisoned); - } - - if self.frames_pushed == 0 { - self.backend.mark_continuation(); - } - - match &mut self.backend { - #[cfg(feature = "cuda")] - Backend::Cuda(d) => d.push_frame(frame), - #[cfg(feature = "rocm")] - Backend::Rocm(d) => d.push_frame(frame), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Backend::Wgpu(d) => d.push_frame(frame), - } - - self.frames_pushed = self.frames_pushed.saturating_add(1); - Ok(()) - } - - /// Drops the current stream and returns to the state a fresh - /// denoiser starts in, keeping every GPU allocation. - /// - /// Frames still in flight are dropped without being read. A frame - /// that [`Self::try_recv_frame`] has already polled is the one - /// exception. Its readback is settled first, which blocks for the - /// rest of that readback's latency and fails the way a - /// [`Self::recv_frame`] on it would. At most one frame is ever in - /// that state. - /// - /// This also clears the poison an earlier failure left, so it is the - /// recovery path for [`DenoiserError::Poisoned`]. - pub fn reset_stream(&mut self) { - self.pending.clear(); - self.frames_pushed = 0; - self.poisoned = false; - - match &mut self.backend { - #[cfg(feature = "cuda")] - Backend::Cuda(d) => d.reset_stream(), - #[cfg(feature = "rocm")] - Backend::Rocm(d) => d.reset_stream(), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Backend::Wgpu(d) => d.reset_stream(), - } - } - - /// Blocks until the in-flight denoise finishes and returns the - /// cleaned frame. - /// - /// Returns `Ok(None)` when nothing is in flight, which happens while - /// the temporal window is still filling up. - /// - /// A failure poisons the denoiser, so every further call returns - /// [`DenoiserError::Poisoned`] until [`Self::reset_stream`] clears it. - pub fn recv_frame(&mut self) -> Result, DenoiserError> { - if self.poisoned { - return Err(DenoiserError::Poisoned); - } - self.recv_frame_inner().inspect_err(|_| self.poisoned = true) - } - - fn recv_frame_inner(&mut self) -> Result, DenoiserError> { - let Some(pending) = self.pending.pop_front() else { - return Ok(None); - }; - Ok(Some(pending.wait()?)) - } - - /// Polls the in-flight denoise once. - /// - /// Returns `Ok(None)` both when nothing is in flight and when the in-flight readback - /// has not landed yet, so `None` alone does not tell those two cases apart. - /// - /// A caller that needs the frame rather than just checking on it should - /// use [`Self::recv_frame`] instead. - /// - /// This only avoids blocking on the wgpu backends, meaning Vulkan and Metal. - /// On CUDA and ROCm the readback completes synchronously on its first poll, - /// so this call blocks until the readback lands there, the same as `recv_frame`. - /// - /// A poll that returns `None` for an in-flight frame commits the - /// denoiser to finishing that readback. Dropping the denoiser or - /// calling [`Self::reset_stream`] before it lands blocks until it - /// does. On the wgpu backends the first poll maps a staging buffer - /// that only the finished readback can unmap, and abandoning it - /// would break every later call on the device. - /// - /// A failure poisons the denoiser, so every further call returns - /// [`DenoiserError::Poisoned`] until [`Self::reset_stream`] clears it. - pub fn try_recv_frame(&mut self) -> Result, DenoiserError> { - if self.poisoned { - return Err(DenoiserError::Poisoned); - } - self.try_recv_frame_inner().inspect_err(|_| self.poisoned = true) - } - - fn try_recv_frame_inner(&mut self) -> Result, DenoiserError> { - let Some(pending) = self.pending.pop_front() else { - return Ok(None); - }; - - match pending.try_wait()? { - Ok(frame) => Ok(Some(frame)), - Err(pending) => { - self.pending.push_front(pending); - Ok(None) - }, - } - } - - /// Drains the in-flight frames and the trailing temporal tail, - /// handing each frame it produces to `sink`. - /// - /// The tail is padded by repeating the last pushed frame. - /// - /// On success the denoiser is ready for a fresh, unrelated stream of - /// the same size and parameters. Pushing again after a flush starts - /// a new temporal window from scratch, and flushing more than once - /// is fine. - /// - /// A failure poisons the denoiser, so every further call returns - /// [`DenoiserError::Poisoned`] until [`Self::reset_stream`] clears it. - pub fn flush(&mut self, sink: impl FnMut(FrameOutput)) -> Result<(), DenoiserError> { - if self.poisoned { - return Err(DenoiserError::Poisoned); - } - self.flush_inner(sink).inspect_err(|_| self.poisoned = true) - } - - fn flush_inner(&mut self, mut sink: impl FnMut(FrameOutput)) -> Result<(), DenoiserError> { - // Drain the whole pending pipeline, up to MAX_PENDING frames, - // before submitting the trailing-tail mirrors. This also leaves - // every output slot free, so the tail's own readbacks cannot be - // handed a slot a streaming readback is still reading. - while let Some(frame) = self.recv_frame_inner()? { - sink(frame); - } - - // The tail frames come back through each algorithm's own - // blocking readback, in the same format as every streaming - // frame, so they are quantised by the same pack kernel. - match &mut self.backend { - #[cfg(feature = "cuda")] - Backend::Cuda(d) => d.flush(|frame| sink(frame.clone()))?, - #[cfg(feature = "rocm")] - Backend::Rocm(d) => d.flush(|frame| sink(frame.clone()))?, - #[cfg(any(feature = "vulkan", feature = "metal"))] - Backend::Wgpu(d) => d.flush(|frame| sink(frame.clone()))?, - } - - // The backend has already reset its own stream indices. Reset - // the outer push counter too, so the next push re-arms the - // window-priming check at the top of `push_frame`. - self.frames_pushed = 0; - - Ok(()) - } - - /// Reads back the grain chunks measured since the last call, in frame order. - /// - /// Call it after [Self::flush]. Measured chunks stay on the GPU until drained, so callers - /// drain after each flush. Empty unless grain export is on. - pub fn drain_grain_chunks(&mut self) -> Result, DenoiserError> { - match &mut self.backend { - #[cfg(feature = "cuda")] - Backend::Cuda(d) => d.drain_grain_chunks(), - #[cfg(feature = "rocm")] - Backend::Rocm(d) => d.drain_grain_chunks(), - #[cfg(any(feature = "vulkan", feature = "metal"))] - Backend::Wgpu(d) => d.drain_grain_chunks(), - } - } - - /// Sets the poison flag directly, without going through a failing - /// call, so a test can check what a caller sees once the flag is - /// already set and what recovers it. - #[cfg(test)] - pub(crate) fn poison_for_test(&mut self) { - self.poisoned = true; - } - - /// The backend's packed-word output buffers, which only exist in - /// wire mode. - #[cfg(test)] - pub(crate) fn wire_outputs_for_test(&self) -> Option<&[cubecl::server::Handle; 2]> { - self.backend.wire_outputs() - } -} - -fn build_backend( - accel: Accelerator, - device: &Device, - algorithm: &Algorithm, - params: NlmParams, - width: u32, - height: u32, - output_format: OutputFormat, -) -> Result { - match accel { - #[cfg(feature = "cuda")] - Accelerator::Cuda => { - let dev = device.to_cuda()?; - let client = ::client(&dev); - Ok(Backend::Cuda(build_engine( - &client, - algorithm, - params, - width, - height, - output_format, - )?)) - }, - #[cfg(feature = "rocm")] - Accelerator::Rocm => { - let dev = device.to_amd()?; - let client = ::client(&dev); - Ok(Backend::Rocm(build_engine( - &client, - algorithm, - params, - width, - height, - output_format, - )?)) - }, - #[cfg(feature = "vulkan")] - Accelerator::Vulkan => { - let dev = device.to_wgpu()?; - let client = ::client(&dev); - Ok(Backend::Wgpu(build_engine( - &client, - algorithm, - params, - width, - height, - output_format, - )?)) - }, - #[cfg(feature = "metal")] - Accelerator::Metal => { - let dev = device.to_wgpu()?; - let client = ::client(&dev); - Ok(Backend::Wgpu(build_engine( - &client, - algorithm, - params, - width, - height, - output_format, - )?)) - }, - // Keeps the match exhaustive on docs.rs, where `cfg(docsrs)` - // widens the `Accelerator` enum to include variants whose - // backend feature is not enabled. Never reached at runtime. - #[cfg(docsrs)] - #[expect( - unreachable_patterns, - reason = "the arm only keeps the match exhaustive on docs.rs" - )] - _ => unreachable!(), - } -} - -#[cfg(test)] -mod options_tests { - use super::*; - - /// `Algorithm::NlmeansHq` with `hq` overridden and everything else - /// left at its default. - fn hq(hq: HqParams) -> Algorithm { - Algorithm::NlmeansHq(NlmeansHqOptions { - hq, - ..NlmeansHqOptions::default() - }) - } - - /// `Algorithm::Nlmeans` with `tuning` overridden. - fn fast_tuned(tuning: NlmTuning) -> Algorithm { - Algorithm::Nlmeans(NlmeansOptions { - tuning, - ..NlmeansOptions::default() - }) - } - - #[test] - fn nl4d_default_lambda_ht_differs_between_luma_and_chroma() { - let luma = nl4d_default_lambda_ht(ChannelMode::Luma); - let chroma = nl4d_default_lambda_ht(ChannelMode::Chroma); - - assert!((luma - 4.158).abs() < f32::EPSILON); - assert!((chroma - 3.234).abs() < f32::EPSILON); - assert!( - (chroma - luma).abs() > f32::EPSILON, - "the two planes should not resolve to the same default" - ); - } - - #[test] - fn nl4d_default_lambda_ht_yuv_reads_the_luma_value() { - let yuv = nl4d_default_lambda_ht(ChannelMode::Yuv); - let luma = nl4d_default_lambda_ht(ChannelMode::Luma); - - assert!((yuv - luma).abs() < f32::EPSILON); - } - - #[test] - fn nl4d_pool_ratio_gives_the_calibrated_threshold_at_each_default_lambda() { - for channels in [ChannelMode::Luma, ChannelMode::Yuv, ChannelMode::Chroma] { - let threshold = nl4d_pool_ratio(channels) * nl4d_default_lambda_ht(channels); - assert!( - (threshold - 2.42).abs() < 1.0e-6, - "{channels:?} gives {threshold}" - ); - } - - assert_eq!( - nl4d_pool_ratio(ChannelMode::Yuv), - nl4d_pool_ratio(ChannelMode::Luma) - ); - } - - #[test] - fn nl4d_options_default_to_pooling_on() { - assert!(Nl4dOptions::default().pooled_threshold); - } - - #[test] - fn resolve_lambda_ht_unset_uses_the_per_plane_default() { - let opts = Nl4dOptions::default(); - - let luma = resolve_lambda_ht(&opts, ChannelMode::Luma).expect("the default scale is in range"); - let chroma = resolve_lambda_ht(&opts, ChannelMode::Chroma).expect("the default scale is in range"); - - assert!((luma - 4.158).abs() < f32::EPSILON, "got {luma}"); - assert!((chroma - 3.234).abs() < f32::EPSILON, "got {chroma}"); - } - - #[test] - fn resolve_lambda_ht_explicit_value_overrides_every_plane() { - let opts = Nl4dOptions { - lambda_ht: Some(4.4), - ..Nl4dOptions::default() - }; - - for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] { - let got = resolve_lambda_ht(&opts, channels).expect("the default scale is in range"); - assert!( - (got - 4.4).abs() < f32::EPSILON, - "channels {channels:?} got {got}" - ); - } - } - - #[test] - fn resolve_lambda_ht_default_scale_leaves_the_value_alone() { - let opts = Nl4dOptions::default(); - - for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] { - let got = resolve_lambda_ht(&opts, channels).expect("the default scale is in range"); - let want = nl4d_default_lambda_ht(channels); - assert!( - (got - want).abs() < f32::EPSILON, - "channels {channels:?} got {got}" - ); - } - } - - #[test] - fn resolve_lambda_ht_scale_multiplies_the_per_plane_default() { - let opts = Nl4dOptions { - lambda_ht_scale: 1.1, - ..Nl4dOptions::default() - }; - - for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] { - let got = resolve_lambda_ht(&opts, channels).expect("1.1 is in range"); - let want = nl4d_default_lambda_ht(channels) * 1.1; - assert!( - (got - want).abs() < 1e-5, - "channels {channels:?} got {got}, want {want}" - ); - } - } - - /// The scale is not limited to the defaults. Pinning one plane and - /// scaling both is the combination this exists for. - #[test] - fn resolve_lambda_ht_scale_multiplies_an_explicit_value() { - let opts = Nl4dOptions { - lambda_ht: Some(5.0), - lambda_ht_scale: 0.9, - ..Nl4dOptions::default() - }; - - let got = resolve_lambda_ht(&opts, ChannelMode::Luma).expect("0.9 is in range"); - assert!((got - 4.5).abs() < 1e-5, "got {got}"); - } - - #[test] - fn resolve_lambda_ht_rejects_an_out_of_range_scale() { - for bad in [0.0, -1.0, 0.05, 10.5, f32::NAN, f32::INFINITY] { - let opts = Nl4dOptions { - lambda_ht_scale: bad, - ..Nl4dOptions::default() - }; - let err = resolve_lambda_ht(&opts, ChannelMode::Luma).unwrap_err(); - assert!( - err.contains("lambda_ht_scale"), - "lambda_ht_scale={bad} should be rejected, got {err}" - ); - } - } - - #[test] - fn the_default_algorithm_is_the_fast_nlmeans_path() { - let opts = DenoiserOptions::builder().build(); - assert_eq!(opts.algorithm, Algorithm::Nlmeans(NlmeansOptions::default())); - } - - #[test] - fn spatial_mode_maps_to_zero_temporal_radius() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Yuv) - .mode(DenoisingMode::Spacial) - .build(); - let params = opts.to_nlm_params(); - - assert_eq!(params.temporal_radius, 0); - assert_eq!(params.channels, ChannelMode::Yuv); - } - - #[test] - fn temporal_mode_propagates_radius() { - let opts = DenoiserOptions::builder() - .mode(DenoisingMode::Temporal { radius: 3 }) - .build(); - let params = opts.to_nlm_params(); - - assert_eq!(params.temporal_radius, 3); - } - - #[test] - fn prefilter_passthrough() { - let opts = DenoiserOptions::builder() - .algorithm(Algorithm::Nlmeans(NlmeansOptions { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: 0.02, - }, - ..NlmeansOptions::default() - })) - .build(); - let params = opts.to_nlm_params(); - - assert!(matches!(params.prefilter, PrefilterMode::Bilateral { .. })); - } - - #[test] - fn hq_unset_prefilter_defaults_to_none() { - let opts = DenoiserOptions::builder() - .algorithm(hq(HqParams::default())) - .build(); - let params = opts.to_nlm_params(); - - assert!(matches!(params.prefilter, PrefilterMode::None)); - } - - #[test] - fn fast_unset_prefilter_defaults_to_none() { - let opts = DenoiserOptions::builder() - .algorithm(Algorithm::Nlmeans(NlmeansOptions::default())) - .build(); - let params = opts.to_nlm_params(); - - assert!(matches!(params.prefilter, PrefilterMode::None)); - } - - #[test] - fn hq_unset_strength_defaults_to_hq_default_strength() { - // Default channel_mode is Yuv, default mode is Spacial (radius 0). - let opts = DenoiserOptions::builder() - .algorithm(hq(HqParams::default())) - .build(); - let params = opts.to_nlm_params(); - - let expected = hq_default_strength(ChannelMode::Yuv, 0); - assert!((params.strength - expected).abs() < f32::EPSILON); - } - - #[test] - fn hq_no_auto_strength_falls_back_to_the_legacy_absolute_default() { - // `effective_strength_with` only reads `strength` as a - // multiplier on the measured sigma when `auto_strength` is true. - // With it false, `strength` is an FFmpeg-style absolute value, - // so the fallback has to be the fast path's absolute default - // rather than a calibrated multiplier from - // `hq_default_strength`. - let opts = DenoiserOptions::builder() - .algorithm(hq(HqParams { - auto_strength: false, - ..HqParams::default() - })) - .build(); - let params = opts.to_nlm_params(); - - let expected = NlmParams::default().strength; - assert!( - (params.strength - expected).abs() < f32::EPSILON, - "expected the legacy absolute default {expected}, got {}, which looks like the \ - auto-strength multiplier table leaking through", - params.strength - ); - } - - #[test] - fn hq_luma_r4_uses_measured_table_value() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(DenoisingMode::Temporal { radius: 4 }) - .algorithm(hq(HqParams::default())) - .build(); - let params = opts.to_nlm_params(); - - assert!((params.strength - 0.35).abs() < f32::EPSILON); - } - - #[test] - fn hq_chroma_r4_uses_measured_table_value() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Chroma) - .mode(DenoisingMode::Temporal { radius: 4 }) - .algorithm(hq(HqParams::default())) - .build(); - let params = opts.to_nlm_params(); - - assert!((params.strength - 0.70).abs() < f32::EPSILON); - } - - #[test] - fn hq_yuv_r8_uses_measured_table_value() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Yuv) - .mode(DenoisingMode::Temporal { radius: 8 }) - .algorithm(hq(HqParams::default())) - .build(); - let params = opts.to_nlm_params(); - - assert!((params.strength - 0.30).abs() < f32::EPSILON); - } - - #[test] - fn hq_spacial_mode_uses_radius_zero_table_values() { - for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] { - let opts = DenoiserOptions::builder() - .channel_mode(channels) - .mode(DenoisingMode::Spacial) - .algorithm(hq(HqParams::default())) - .build(); - let params = opts.to_nlm_params(); - - let expected = hq_default_strength(channels, 0); - assert!( - (params.strength - expected).abs() < f32::EPSILON, - "for channels {channels:?} expected {expected}, got {}", - params.strength - ); - } - } - - #[test] - fn hq_explicit_strength_wins_over_the_table_for_every_plane() { - for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] { - let opts = DenoiserOptions::builder() - .channel_mode(channels) - .mode(DenoisingMode::Temporal { radius: 4 }) - .algorithm(Algorithm::NlmeansHq(NlmeansHqOptions { - nlm: NlmeansOptions { - tuning: NlmTuning { - strength: Some(0.99), - ..NlmTuning::default() - }, - ..NlmeansOptions::default() - }, - hq: HqParams::default(), - })) - .build(); - let params = opts.to_nlm_params(); - - assert!( - (params.strength - 0.99).abs() < f32::EPSILON, - "for channels {channels:?} the explicit strength was overridden by the table" - ); - } - } - - #[test] - fn fast_unset_strength_defaults_to_legacy_default() { - let opts = DenoiserOptions::builder() - .algorithm(Algorithm::Nlmeans(NlmeansOptions::default())) - .build(); - let params = opts.to_nlm_params(); - - assert!((params.strength - 1.2).abs() < f32::EPSILON); - } - - #[test] - fn nl4d_options_default_matches_nl4d_params_default() { - let opts = Nl4dOptions::default(); - let params = crate::nl4d::Nl4dParams::default(); - - assert_eq!(opts.refine, params.refine); - assert_eq!(opts.spatial_radius, params.spatial_radius); - assert!((opts.c_min - params.c_min).abs() < f32::EPSILON); - // The two `lambda_ht` fields hold different things, so they are - // not compared. `opts.lambda_ht` stays `None` and is deferred to - // `nl4d_default_lambda_ht` once the plane is known (see - // `resolve_lambda_ht_unset_uses_the_per_plane_default` above), - // while `params.lambda_ht` is a concrete default mirroring the - // Luma/Yuv value. - assert_eq!(opts.lambda_ht, None); - assert!((params.lambda_ht - nl4d_default_lambda_ht(ChannelMode::Yuv)).abs() < f32::EPSILON); - } - - /// nl4d takes its own noise and confidence knobs rather than a whole - /// [`HqParams`], so the three it does take have to reach the front - /// end and the rest have to arrive at their defaults. - #[test] - fn nl4d_builds_the_front_ends_hq_params_from_its_own_fields() { - let opts = DenoiserOptions::builder() - .mode(DenoisingMode::Temporal { radius: 2 }) - .algorithm(Algorithm::Nl4d(Nl4dOptions { - sigma: Some(0.02), - sigma_scale: 1.3, - thsad_scale: 0.8, - ..Nl4dOptions::default() - })) - .build(); - let params = opts.to_nlm_params(); - - let hq = params.hq.expect("nl4d always runs the hq front end"); - assert_eq!(hq.sigma_override, Some(0.02)); - assert!((hq.sigma_scale - 1.3).abs() < f32::EPSILON); - assert!((hq.thsad_scale - 0.8).abs() < f32::EPSILON); - assert!( - hq.temporal_confidence, - "the grouping kernel reads the confidence scores, so this cannot be off" - ); - } - - /// The temporal radius has one source now, `mode`, so nl4d cannot - /// disagree with the front end's ring about how wide the window is. - #[test] - fn nl4d_reads_its_temporal_radius_from_the_denoising_mode() { - for radius in [1u32, 4, 8] { - let opts = DenoiserOptions::builder() - .mode(DenoisingMode::Temporal { radius }) - .algorithm(Algorithm::Nl4d(Nl4dOptions::default())) - .build(); - - assert_eq!(opts.to_nlm_params().temporal_radius, radius); - } - } - - /// nl4d never runs an NLM weighting pass, so a prefilter would cost - /// a GPU pass per frame producing a reference image nothing reads. - #[test] - fn nl4d_never_builds_a_prefilter() { - let opts = DenoiserOptions::builder() - .mode(DenoisingMode::Temporal { radius: 2 }) - .algorithm(Algorithm::Nl4d(Nl4dOptions::default())) - .build(); - - assert!(matches!(opts.to_nlm_params().prefilter, PrefilterMode::None)); - } - - /// Nothing in the nl4d path reads `strength`, so it stays at the - /// library default rather than picking up HQ's calibrated table. - #[test] - fn nl4d_leaves_the_nlm_weighting_knobs_at_their_defaults() { - let defaults = NlmParams::default(); - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(DenoisingMode::Temporal { radius: 4 }) - .algorithm(Algorithm::Nl4d(Nl4dOptions::default())) - .build(); - let params = opts.to_nlm_params(); - - assert!((params.strength - defaults.strength).abs() < f32::EPSILON); - assert_eq!(params.search_radius, defaults.search_radius); - assert_eq!(params.patch_radius, defaults.patch_radius); - assert!((params.self_weight - defaults.self_weight).abs() < f32::EPSILON); - } - - /// nl4d always tracks motion, so its `MotionSearch` reaches the - /// front end as an active `Mvtools` mode. - #[test] - fn nl4d_motion_search_becomes_an_active_mvtools_mode() { - let opts = DenoiserOptions::builder() - .mode(DenoisingMode::Temporal { radius: 2 }) - .algorithm(Algorithm::Nl4d(Nl4dOptions { - motion: MotionSearch { - blksize: 32, - overlap: 16, - search_radius: 6, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - }, - ..Nl4dOptions::default() - })) - .build(); - let params = opts.to_nlm_params(); - - assert!(matches!( - params.motion_compensation, - MotionCompensationMode::Mvtools { - blksize: 32, - overlap: 16, - search_radius: 6, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - } - )); - } - - #[test] - fn nl4d_motion_search_defaults_match_the_front_ends_own_defaults() { - let opts = DenoiserOptions::builder() - .mode(DenoisingMode::Temporal { radius: 2 }) - .algorithm(Algorithm::Nl4d(Nl4dOptions::default())) - .build(); - let params = opts.to_nlm_params(); - - assert_eq!( - params.motion_compensation, - crate::nl4d::Nl4dParams::default().nlm.motion_compensation - ); - } - - #[test] - fn motion_compensation_passthrough() { - let opts = DenoiserOptions::builder() - .mode(DenoisingMode::Temporal { radius: 1 }) - .algorithm(Algorithm::Nlmeans(NlmeansOptions { - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 16, - overlap: 8, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }, - ..NlmeansOptions::default() - })) - .build(); - let params = opts.to_nlm_params(); - - assert!(matches!( - params.motion_compensation, - MotionCompensationMode::Mvtools { - blksize: 16, - overlap: 8, - search_radius: 4, - pyramid_levels: 2, - .. - } - )); - } - - #[test] - fn motion_compensation_defaults_to_none() { - let opts = DenoiserOptions::builder().build(); - let params = opts.to_nlm_params(); - assert!(matches!(params.motion_compensation, MotionCompensationMode::None)); - } - - #[test] - fn nlm_tuning_overrides_individual_fields() { - let defaults = NlmParams::default(); - let opts = DenoiserOptions::builder() - .algorithm(fast_tuned(NlmTuning { - search_radius: Some(7), - patch_radius: None, - strength: Some(2.5), - self_weight: None, - })) - .build(); - let params = opts.to_nlm_params(); - - assert_eq!(params.search_radius, 7); - assert_eq!(params.patch_radius, defaults.patch_radius); - assert!((params.strength - 2.5).abs() < f32::EPSILON); - assert!((params.self_weight - defaults.self_weight).abs() < f32::EPSILON); - } -} - -#[cfg(all(test, feature = "vulkan"))] -mod tests { - use super::*; - - fn opts(mode: DenoisingMode) -> DenoiserOptions { - DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(mode) - .build() - } - - fn frame(w: u32, h: u32) -> Vec { - vec![0.5f32; (w * h) as usize] - } - - fn f32_out(out: FrameOutput) -> Vec { - out.into_f32().expect("f32 output") - } - - #[test] - fn spatial_denoise_roundtrip() { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Spacial), - ) - .expect("denoiser construction failed"); - assert_eq!(d.selected_accelerator(), Accelerator::Vulkan); - - d.push_frame(&frame(16, 16)).expect("push failed"); - let out = f32_out(d.recv_frame().expect("recv failed").expect("no frame")); - assert_eq!(out.len(), 16 * 16); - } - - #[test] - fn nl4d_algorithm_round_trips_through_the_facade() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(DenoisingMode::Temporal { radius: 2 }) - .algorithm(Algorithm::Nl4d(Nl4dOptions::default())) - .build(); - let mut d = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts) - .expect("nl4d denoiser construction failed"); - assert_eq!(d.selected_accelerator(), Accelerator::Vulkan); - - // temporal_radius is 2, so a single push does not fill the - // window yet, the same convention every temporal algorithm here - // follows. - d.push_frame(&frame(16, 16)).expect("push failed"); - assert!(d.recv_frame().expect("recv failed").is_none()); - - let mut out = Vec::new(); - d.flush(|f| out.push(f32_out(f))).expect("flush failed"); - assert_eq!(out.len(), 1, "expected exactly one output for one pushed frame"); - assert_eq!(out[0].len(), 16 * 16); - } - - /// nl4d groups patches across neighbouring frames, so a spatial - /// mode leaves it nothing to do. The temporal radius has one source - /// now, `mode`, so this is the only way to ask for that. - #[test] - fn nl4d_rejects_a_spatial_denoising_mode() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(DenoisingMode::Spacial) - .algorithm(Algorithm::Nl4d(Nl4dOptions::default())) - .build(); - let result = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts); - - match result { - Err(DenoiserError::Other(e)) => assert!( - e.to_string().contains("temporal window"), - "unexpected error message: {e}" - ), - Err(other) => panic!("expected DenoiserError::Other, got {other:?}"), - Ok(_) => panic!("expected a rejection, got Ok"), - } - } - - /// Both NLM algorithms only need a symmetric `2r+1` window, so - /// `window_span` must report the same radius on both sides. - #[test] - fn window_span_is_symmetric_for_nlmeans() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(DenoisingMode::Temporal { radius: 3 }) - .algorithm(Algorithm::Nlmeans(NlmeansOptions::default())) - .build(); - let d = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts) - .expect("denoiser construction failed"); - - let span = d.window_span(); - assert_eq!(span.behind, 3, "behind should equal the temporal radius"); - assert_eq!(span.ahead, 3, "ahead should equal the temporal radius"); - assert_eq!(span.edges, EdgePadding::Repeat); - } - - /// nl4d's cross-frame accumulator needs the target's own `radius` - /// neighbourhood doubled on both sides, so both `behind` and - /// `ahead` must come out to `2 * radius`. - #[test] - fn window_span_is_doubled_on_both_sides_for_nl4d() { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(DenoisingMode::Temporal { radius: 3 }) - .algorithm(Algorithm::Nl4d(Nl4dOptions::default())) - .build(); - let d = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts) - .expect("nl4d denoiser construction failed"); - - let span = d.window_span(); - assert_eq!(span.behind, 6, "behind should equal 2 * the temporal radius"); - assert_eq!(span.ahead, 6, "ahead should equal 2 * the temporal radius"); - assert_eq!(span.edges, EdgePadding::Shifted); - } - - #[test] - fn invalid_params_surface_as_error() { - let bad = DenoiserOptions::builder() - .algorithm(Algorithm::Nlmeans(NlmeansOptions { - tuning: NlmTuning { - strength: Some(0.0), - ..NlmTuning::default() - }, - ..NlmeansOptions::default() - })) - .build(); - let result = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, bad); - - match result { - Err(DenoiserError::Other(_)) => {}, - Err(other) => panic!("expected DenoiserError::Other, got {other:?}"), - Ok(_) => panic!("expected validation error, got Ok"), - } - } - - #[test] - fn tiny_frame_dimensions_surface_as_error() { - let result = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 2, - 2, - opts(DenoisingMode::Spacial), - ); - - match result { - Err(DenoiserError::Other(e)) => { - assert!( - e.to_string().contains("supported minimum"), - "unexpected error message: {e}" - ); - }, - Err(other) => panic!("expected DenoiserError::Other, got {other:?}"), - Ok(_) => panic!("expected dimension validation error, got Ok"), - } - } - - #[test] - fn push_after_pending_returns_queue_full() { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Spacial), - ) - .unwrap(); - - // The pipeline is two deep because the output handles are - // double-buffered, so the first two pushes both submit. The - // third would overwrite the oldest pending frame's output slot, - // so it is rejected with QueueFull. - d.push_frame(&frame(16, 16)).unwrap(); - d.push_frame(&frame(16, 16)).unwrap(); - let err = d.push_frame(&frame(16, 16)).expect_err("expected QueueFull"); - assert!(matches!(err, DenoiserError::QueueFull)); - - let out = f32_out(d.recv_frame().unwrap().unwrap()); - assert_eq!(out.len(), 16 * 16); - - // After draining one slot the next push must succeed. - d.push_frame(&frame(16, 16)).expect("push after drain failed"); - } - - /// `QueueFull` is the documented retry signal, so it must not leave - /// the denoiser poisoned. - #[test] - fn queue_full_does_not_poison() { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Spacial), - ) - .unwrap(); - - d.push_frame(&frame(16, 16)).unwrap(); - d.push_frame(&frame(16, 16)).unwrap(); - let err = d.push_frame(&frame(16, 16)).expect_err("expected QueueFull"); - assert!(matches!(err, DenoiserError::QueueFull)); - assert!(!d.poisoned, "QueueFull must not poison the denoiser"); - - d.recv_frame().unwrap().expect("recv failed after QueueFull"); - - // The queue has room again, so the next push must succeed rather - // than being rejected as poisoned. - d.push_frame(&frame(16, 16)) - .expect("push after QueueFull drain should succeed, not poison"); - } - - #[test] - fn poisoned_denoiser_refuses_every_entry_point() { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Spacial), - ) - .unwrap(); - d.poisoned = true; - - assert!(matches!( - d.push_frame(&frame(16, 16)), - Err(DenoiserError::Poisoned) - )); - assert!(matches!(d.recv_frame(), Err(DenoiserError::Poisoned))); - assert!(matches!(d.try_recv_frame(), Err(DenoiserError::Poisoned))); - assert!(matches!(d.flush(|_| {}), Err(DenoiserError::Poisoned))); - } - - #[test] - fn reset_stream_clears_poison() { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Spacial), - ) - .unwrap(); - d.poisoned = true; - - d.reset_stream(); - assert!(!d.poisoned, "reset_stream must clear the poison flag"); - - d.push_frame(&frame(16, 16)) - .expect("push after reset_stream should succeed"); - } - - fn frame_filled(w: u32, h: u32, value: f32) -> Vec { - vec![value; (w * h) as usize] - } - - /// Pushes `n` frames of the given value, receiving along the way to - /// keep the in-flight pipeline below `MAX_PENDING`. - fn push_n_with_drain(d: &mut Denoiser, n: usize, value: f32, out: &mut Vec>) { - for _ in 0..n { - loop { - match d.push_frame(&frame_filled(16, 16, value)) { - Ok(()) => break, - Err(DenoiserError::QueueFull) => { - let f = d - .recv_frame() - .expect("recv ok") - .expect("queue full but recv yielded none"); - out.push(f32_out(f)); - }, - Err(e) => panic!("unexpected push error: {e:?}"), - } - } - } - } - - #[test] - fn flush_leaves_denoiser_reusable_spatial() { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Spacial), - ) - .unwrap(); - - let mut batch_a = Vec::new(); - push_n_with_drain(&mut d, 5, 0.25, &mut batch_a); - d.flush(|f| batch_a.push(f32_out(f))).expect("first flush failed"); - assert_eq!(batch_a.len(), 5); - - // After flush the pipeline must be empty. - assert!(d.recv_frame().unwrap().is_none()); - - let mut batch_b = Vec::new(); - push_n_with_drain(&mut d, 5, 0.75, &mut batch_b); - d.flush(|f| batch_b.push(f32_out(f))) - .expect("second flush failed"); - assert_eq!(batch_b.len(), 5); - - for v in batch_b.iter().flatten() { - assert!((v - 0.75).abs() < 0.1, "batch_b carried state from batch_a: {v}"); - } - for v in batch_a.iter().flatten() { - assert!((v - 0.25).abs() < 0.1, "batch_a value unexpectedly drifted: {v}"); - } - } - - #[test] - fn flush_leaves_denoiser_reusable_temporal() { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Temporal { radius: 1 }), - ) - .unwrap(); - - let mut batch_a = Vec::new(); - push_n_with_drain(&mut d, 5, 0.25, &mut batch_a); - d.flush(|f| batch_a.push(f32_out(f))).expect("first flush failed"); - assert_eq!(batch_a.len(), 5, "expected 5 frames from first batch"); - - // The temporal window must be empty after a flush, so the first - // push of the new stream should not produce a pending frame. - // With r=1 the window needs 3 frames before `denoise_submit` - // fires. - assert!(d.recv_frame().unwrap().is_none()); - d.push_frame(&frame_filled(16, 16, 0.75)).unwrap(); - assert!( - d.recv_frame().unwrap().is_none(), - "first push of new temporal stream should not produce output yet" - ); - - // Push 4 more frames (5 total in batch B) with drain. - let mut batch_b = Vec::new(); - push_n_with_drain(&mut d, 4, 0.75, &mut batch_b); - d.flush(|f| batch_b.push(f32_out(f))) - .expect("second flush failed"); - assert_eq!(batch_b.len(), 5, "expected 5 frames from second batch"); - - for v in batch_b.iter().flatten() { - assert!((v - 0.75).abs() < 0.1, "batch_b carried state from batch_a: {v}"); - } - } - - #[test] - fn flush_emits_exactly_n_outputs_for_small_n() { - // With temporal radius R=2 the window is 5 frames. Pushing fewer - // than R+1 frames means the window never fills while pushing, so - // flush must still emit one output per pushed frame rather than - // R+1 of them. - for n in 1..=5usize { - let mut d = Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 16, - 16, - opts(DenoisingMode::Temporal { radius: 2 }), - ) - .unwrap(); - - let mut out = Vec::new(); - push_n_with_drain(&mut d, n, 0.5, &mut out); - d.flush(|f| out.push(f32_out(f))).expect("flush failed"); - assert_eq!( - out.len(), - n, - "expected {n} outputs for {n} pushes, got {}", - out.len() - ); - } - } - - /// A `Pending` that was polled and then dropped has to settle its - /// readback first. On the wgpu backends the first poll maps a - /// staging buffer that only the finished readback unmaps, and a - /// still-mapped buffer back in the pool makes the next submit on - /// the device fail on cubecl's device thread, after which every - /// call on that device errors. - #[test] - fn dropping_a_polled_pending_frame_does_not_poison_the_device() { - let new = || { - Denoiser::create( - &[Accelerator::Vulkan], - &Device::Default, - 64, - 64, - opts(DenoisingMode::Spacial), - ) - .unwrap() - }; - - let mut d = new(); - d.push_frame(&frame(64, 64)).unwrap(); - // One poll starts the readback. Whether it lands on this poll - // depends on the GPU, and both outcomes must survive the drop. - let _ = d.try_recv_frame().unwrap(); - drop(d); - - // The staging pool is per device, so a fresh denoiser on the - // same device is handed the same buffers. - let mut d = new(); - for _ in 0..4 { - d.push_frame(&frame(64, 64)).unwrap(); - d.recv_frame() - .expect("readback after a dropped polled frame should not fail") - .expect("spatial mode emits one frame per push"); - } - d.flush(|_| {}).unwrap(); - } -} diff --git a/av-denoise-core/src/engine/io.rs b/av-denoise-core/src/engine/io.rs new file mode 100644 index 0000000..2d4dd50 --- /dev/null +++ b/av-denoise-core/src/engine/io.rs @@ -0,0 +1,168 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::kernels::{egress_f32, egress_words, ingest_f32, ingest_words}; +use super::{DevicePlane, SampleFormat}; +use crate::nlmeans::{BLOCK_1D, MAX_GRID_1D}; + +/// Where an ingest writes, as a slot of an interleaved ring. +pub(crate) struct IngestTarget<'a> { + pub ring: &'a Handle, + pub ring_len: usize, + pub offset: u32, + pub pixels: u32, + pub channels: u32, + pub stored_ch: u32, +} + +/// The handle bound for plane `index`, falling back to `placeholder` for planes the kernel never reads. +fn plane_or<'a>(planes: &[DevicePlane<'a>], index: usize, placeholder: &'a Handle) -> &'a Handle { + match planes.get(index) { + Some(plane) => plane.handle(), + None => placeholder, + } +} + +fn grid(pixels: u32) -> (CubeCount, u32) { + let groups = pixels.div_ceil(BLOCK_1D).clamp(1, MAX_GRID_1D); + let total_threads = groups * BLOCK_1D; + (CubeCount::new_1d(groups), total_threads) +} + +/// Queues the ingest of `planes` into one ring slot. +/// +/// The caller has validated `planes` against the engine's geometry. +pub(crate) fn ingest( + client: &ComputeClient, + planes: &[DevicePlane<'_>], + format: SampleFormat, + placeholder: &Handle, + target: IngestTarget<'_>, +) { + let (count, total_threads) = grid(target.pixels); + let plane_0 = plane_or(planes, 0, placeholder); + let plane_1 = plane_or(planes, 1, placeholder); + let plane_2 = plane_or(planes, 2, placeholder); + let ring = unsafe { ArrayArg::from_raw_parts(target.ring.clone(), target.ring_len) }; + + match format { + SampleFormat::F32 => { + let plane_len = target.pixels as usize; + + unsafe { + ingest_f32::launch_unchecked::( + client, + count, + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(plane_0.clone(), plane_len), + ArrayArg::from_raw_parts(plane_1.clone(), plane_len), + ArrayArg::from_raw_parts(plane_2.clone(), plane_len), + ring, + target.offset, + target.pixels, + target.channels, + target.stored_ch, + total_threads, + ); + } + }, + SampleFormat::U8 | SampleFormat::U16 { .. } => { + let samples_per_word = format.samples_per_word(); + let words = target.pixels.div_ceil(samples_per_word) as usize; + let max = format.max_value(); + + unsafe { + ingest_words::launch_unchecked::( + client, + count, + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(plane_0.clone(), words), + ArrayArg::from_raw_parts(plane_1.clone(), words), + ArrayArg::from_raw_parts(plane_2.clone(), words), + ring, + max, + target.offset, + target.pixels, + target.channels, + target.stored_ch, + samples_per_word, + total_threads, + ); + } + }, + } +} + +/// A finished interleaved f32 frame to write out. +pub(crate) struct EgressSource<'a> { + pub frame: &'a Handle, + pub pixels: u32, + pub channels: u32, + pub stored_ch: u32, +} + +/// Queues the write of `source` into `planes`. +/// +/// The caller has validated `planes` against the engine's geometry. +pub(crate) fn egress( + client: &ComputeClient, + source: EgressSource<'_>, + planes: &[DevicePlane<'_>], + format: SampleFormat, + placeholder: &Handle, +) { + let frame_len = (source.pixels * source.stored_ch) as usize; + let frame = unsafe { ArrayArg::from_raw_parts(source.frame.clone(), frame_len) }; + let plane_0 = plane_or(planes, 0, placeholder); + let plane_1 = plane_or(planes, 1, placeholder); + let plane_2 = plane_or(planes, 2, placeholder); + + match format { + SampleFormat::F32 => { + let (count, total_threads) = grid(source.pixels); + let plane_len = source.pixels as usize; + + unsafe { + egress_f32::launch_unchecked::( + client, + count, + CubeDim::new_1d(BLOCK_1D), + frame, + ArrayArg::from_raw_parts(plane_0.clone(), plane_len), + ArrayArg::from_raw_parts(plane_1.clone(), plane_len), + ArrayArg::from_raw_parts(plane_2.clone(), plane_len), + source.pixels, + source.channels, + source.stored_ch, + total_threads, + ); + } + }, + SampleFormat::U8 | SampleFormat::U16 { .. } => { + let samples_per_word = format.samples_per_word(); + let words = source.pixels.div_ceil(samples_per_word); + let (count, total_threads) = grid(words); + let max = format.max_value(); + let word_len = words as usize; + + unsafe { + egress_words::launch_unchecked::( + client, + count, + CubeDim::new_1d(BLOCK_1D), + frame, + ArrayArg::from_raw_parts(plane_0.clone(), word_len), + ArrayArg::from_raw_parts(plane_1.clone(), word_len), + ArrayArg::from_raw_parts(plane_2.clone(), word_len), + max, + source.pixels, + source.channels, + source.stored_ch, + samples_per_word, + words, + total_threads, + ); + } + }, + } +} diff --git a/av-denoise-core/src/engine/kernels.rs b/av-denoise-core/src/engine/kernels.rs new file mode 100644 index 0000000..478dfa7 --- /dev/null +++ b/av-denoise-core/src/engine/kernels.rs @@ -0,0 +1,203 @@ +use cubecl::prelude::*; + +/// Normalises word-packed planes into one interleaved ring slot. +/// +/// `max` is the largest sample code, and `offset` is the slot's first element in `ring`. Threads +/// stride over pixels by `total_threads` and write all `stored_ch` lanes, with padding lanes set to +/// zero. A plane is only read when `channels` covers it, so an unused plane can be a placeholder. +#[cube(launch_unchecked)] +#[expect( + clippy::too_many_arguments, + reason = "every argument is a binding or a comptime shape" +)] +pub fn ingest_words( + plane_0: &Array, + plane_1: &Array, + plane_2: &Array, + ring: &mut Array, + max: f32, + offset: u32, + #[comptime] pixels: u32, + #[comptime] channels: u32, + #[comptime] stored_ch: u32, + #[comptime] samples_per_word: u32, + #[comptime] total_threads: u32, +) { + let bits = comptime![32u32 / samples_per_word]; + let mask = comptime![(1u32 << (32u32 / samples_per_word)) - 1]; + + let mut pixel = ABSOLUTE_POS_X; + while pixel < pixels { + let word = (pixel / samples_per_word) as usize; + let shift = (pixel % samples_per_word) * bits; + let base = offset + pixel * stored_ch; + + let code_0 = (plane_0[word] >> shift) & mask; + ring[base as usize] = f32::cast_from(code_0) / max; + + if channels > 1 { + let code_1 = (plane_1[word] >> shift) & mask; + ring[(base + 1) as usize] = f32::cast_from(code_1) / max; + } + + if channels > 2 { + let code_2 = (plane_2[word] >> shift) & mask; + ring[(base + 2) as usize] = f32::cast_from(code_2) / max; + } + + #[unroll] + for lane in channels..stored_ch { + ring[(base + lane) as usize] = 0.0f32; + } + + pixel += total_threads; + } +} + +/// Copies f32 planes into one interleaved ring slot, with padding lanes set to zero. +/// +/// A plane is only read when `channels` covers it, so an unused plane can be a placeholder. +#[cube(launch_unchecked)] +#[expect( + clippy::too_many_arguments, + reason = "every argument is a binding or a comptime shape" +)] +pub fn ingest_f32( + plane_0: &Array, + plane_1: &Array, + plane_2: &Array, + ring: &mut Array, + offset: u32, + #[comptime] pixels: u32, + #[comptime] channels: u32, + #[comptime] stored_ch: u32, + #[comptime] total_threads: u32, +) { + let mut pixel = ABSOLUTE_POS_X; + while pixel < pixels { + let base = offset + pixel * stored_ch; + ring[base as usize] = plane_0[pixel as usize]; + + if channels > 1 { + ring[(base + 1) as usize] = plane_1[pixel as usize]; + } + + if channels > 2 { + ring[(base + 2) as usize] = plane_2[pixel as usize]; + } + + #[unroll] + for lane in channels..stored_ch { + ring[(base + lane) as usize] = 0.0f32; + } + + pixel += total_threads; + } +} + +/// Quantises an interleaved f32 frame into word-packed planes. +/// +/// Values are clamped to `0.0..=1.0` and rounded to the nearest code up to `max`. Each thread +/// packs and writes whole words, so no two threads share a word. Lanes past the last pixel read a +/// clamped index and are masked to zero, because a branch-derived index inside an unrolled loop +/// makes cubecl's GVN pass panic and the launch then silently writes nothing. A plane is only +/// written when `channels` covers it, so an unused plane can be a placeholder. +#[cube(launch_unchecked)] +#[expect( + clippy::too_many_arguments, + reason = "every argument is a binding or a comptime shape" +)] +pub fn egress_words( + frame: &Array, + plane_0: &mut Array, + plane_1: &mut Array, + plane_2: &mut Array, + max: f32, + #[comptime] pixels: u32, + #[comptime] channels: u32, + #[comptime] stored_ch: u32, + #[comptime] samples_per_word: u32, + #[comptime] words: u32, + #[comptime] total_threads: u32, +) { + let bits = comptime![32u32 / samples_per_word]; + + let mut word = ABSOLUTE_POS_X; + while word < words { + let mut packed_0 = 0u32; + let mut packed_1 = 0u32; + let mut packed_2 = 0u32; + + #[unroll] + for lane in 0..samples_per_word { + let pixel = word * samples_per_word + lane; + let clamped = u32::min(pixel, pixels - 1); + let in_range = pixel < pixels; + let base = clamped * stored_ch; + let shift = lane * bits; + + let value_0 = f32::clamp(frame[base as usize], 0.0, 1.0); + let code_0 = u32::cast_from(value_0 * max + 0.5); + packed_0 |= select(in_range, code_0, 0u32) << shift; + + if channels > 1 { + let value_1 = f32::clamp(frame[(base + 1) as usize], 0.0, 1.0); + let code_1 = u32::cast_from(value_1 * max + 0.5); + packed_1 |= select(in_range, code_1, 0u32) << shift; + } + + if channels > 2 { + let value_2 = f32::clamp(frame[(base + 2) as usize], 0.0, 1.0); + let code_2 = u32::cast_from(value_2 * max + 0.5); + packed_2 |= select(in_range, code_2, 0u32) << shift; + } + } + + plane_0[word as usize] = packed_0; + + if channels > 1 { + plane_1[word as usize] = packed_1; + } + + if channels > 2 { + plane_2[word as usize] = packed_2; + } + + word += total_threads; + } +} + +/// Copies an interleaved f32 frame into f32 planes, skipping padding lanes. +/// +/// A plane is only written when `channels` covers it, so an unused plane can be a placeholder. +#[cube(launch_unchecked)] +#[expect( + clippy::too_many_arguments, + reason = "every argument is a binding or a comptime shape" +)] +pub fn egress_f32( + frame: &Array, + plane_0: &mut Array, + plane_1: &mut Array, + plane_2: &mut Array, + #[comptime] pixels: u32, + #[comptime] channels: u32, + #[comptime] stored_ch: u32, + #[comptime] total_threads: u32, +) { + let mut pixel = ABSOLUTE_POS_X; + while pixel < pixels { + let base = pixel * stored_ch; + plane_0[pixel as usize] = frame[base as usize]; + + if channels > 1 { + plane_1[pixel as usize] = frame[(base + 1) as usize]; + } + + if channels > 2 { + plane_2[pixel as usize] = frame[(base + 2) as usize]; + } + + pixel += total_threads; + } +} diff --git a/av-denoise-core/src/engine/mod.rs b/av-denoise-core/src/engine/mod.rs new file mode 100644 index 0000000..2ddb531 --- /dev/null +++ b/av-denoise-core/src/engine/mod.rs @@ -0,0 +1,74 @@ +mod io; +pub(crate) mod kernels; +mod plane; + +#[cfg(test)] +mod tests; + +pub(crate) use self::io::{EgressSource, IngestTarget, egress, ingest}; +pub use self::plane::{DevicePlane, Geometry, SampleFormat}; +use crate::error::Error; +use crate::nl4d::grain::GrainChunk; + +/// How many frames before and after a target frame are needed to denoise it. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct WindowSpan { + pub behind: usize, + pub ahead: usize, + pub edges: EdgePadding, +} + +/// How a window is filled where it runs past a clip's ends. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum EdgePadding { + /// The boundary frame is repeated, so every window has the same length. + Repeat, + /// The window stops at the clip's ends and the engine runs off-centre passes there. + Shifted, +} + +impl WindowSpan { + /// The full window length, target frame included. + pub fn frame_count(&self) -> usize { + self.behind + 1 + self.ahead + } +} + +/// A stateful denoiser that reads and writes GPU planes. +/// +/// Push frames in order with [Engine::push]. It returns how many frames are ready, and every ready frame +/// must be written out with [Engine::emit_into] before the next push. At the end of a stream call +/// [Engine::finish] and emit the frames it reports. The next push then starts a new stream. +/// +/// Any GPU error leaves the engine refusing calls with [Error::NeedsReset] until [Engine::reset]. +pub trait Engine: Send { + /// Ingests one frame and returns how many frames are ready to emit. + fn push(&mut self, planes: &[DevicePlane<'_>]) -> Result; + + /// Ingests one frame as context only, producing no output from this call. + /// + /// A stream that starts with this picks up mid-clip instead of at a scene start. Context frames + /// still fill the window, so a later push or [Engine::finish] may emit them, and `finish` alone + /// reports at most [Engine::max_held_frames] of them. It returns [Error::ContextAfterPush] once the + /// stream has had a real push. + fn push_context(&mut self, planes: &[DevicePlane<'_>]) -> Result<(), Error>; + + /// Writes the oldest ready frame into `planes`. + fn emit_into(&mut self, planes: &[DevicePlane<'_>]) -> Result<(), Error>; + + /// Ends the stream and returns how many tail frames are ready to emit. + fn finish(&mut self) -> Result; + + /// Abandons the current stream, keeping every allocation. + fn reset(&mut self); + + fn window_span(&self) -> WindowSpan; + + /// The most frames the engine holds before it emits one. + fn max_held_frames(&self) -> usize; + + /// Reads back the film grain measured since the last call, in frame order. + fn drain_grain_chunks(&mut self) -> Result, Error> { + Ok(Vec::new()) + } +} diff --git a/av-denoise-core/src/engine/plane.rs b/av-denoise-core/src/engine/plane.rs new file mode 100644 index 0000000..faf1da5 --- /dev/null +++ b/av-denoise-core/src/engine/plane.rs @@ -0,0 +1,167 @@ +use cubecl::server::Handle; + +use crate::error::Error; +use crate::nlmeans::ChannelMode; + +/// How samples are stored in a plane. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum SampleFormat { + U8, + /// 16-bit words holding `depth` significant bits, between 9 and 16. + U16 { + depth: u8, + }, + /// Normalised samples between 0 and 1. + F32, +} + +impl SampleFormat { + pub(crate) fn validate(self) -> Result<(), Error> { + match self { + SampleFormat::U16 { depth } if !(9..=16).contains(&depth) => { + let message = format!("U16 depth must be between 9 and 16, got {depth}"); + Err(Error::InvalidGeometry(message)) + }, + _ => Ok(()), + } + } + + /// The largest sample value, which normalisation divides by. + pub(crate) fn max_value(self) -> f32 { + match self { + SampleFormat::U8 => 255.0, + SampleFormat::U16 { depth } => ((1u32 << depth) - 1) as f32, + SampleFormat::F32 => 1.0, + } + } + + pub(crate) fn samples_per_word(self) -> u32 { + match self { + SampleFormat::U8 => 4, + SampleFormat::U16 { .. } => 2, + SampleFormat::F32 => 1, + } + } + + /// Bytes a plane of `pixels` samples occupies, rounded up to whole words. + pub fn plane_bytes(self, pixels: u64) -> u64 { + let samples_per_word = self.samples_per_word() as u64; + let words = pixels.div_ceil(samples_per_word); + words * 4 + } +} + +/// One plane of samples on the GPU, tightly packed with a stride equal to its width. +/// +/// The handle must hold whole 4-byte words, which is `SampleFormat::plane_bytes(width * height)` bytes, +/// and [Engine::emit_into](crate::engine::Engine::emit_into) writes zeros into the padding past the last +/// sample. +#[derive(Debug, Clone, Copy)] +pub struct DevicePlane<'a> { + handle: &'a Handle, + width: u32, + height: u32, +} + +impl<'a> DevicePlane<'a> { + pub fn new(handle: &'a Handle, width: u32, height: u32) -> Self { + Self { + handle, + width, + height, + } + } + + pub fn handle(&self) -> &'a Handle { + self.handle + } + + pub fn width(&self) -> u32 { + self.width + } + + pub fn height(&self) -> u32 { + self.height + } +} + +/// The shape and sample formats an engine is built for. +/// +/// `width` and `height` are the dimensions of the planes the engine sees, so a chroma engine is built at +/// chroma size. Luma takes one plane, Chroma takes U and V, and Yuv takes Y, U and V at one size. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct Geometry { + pub width: u32, + pub height: u32, + pub channels: ChannelMode, + pub input: SampleFormat, + pub output: SampleFormat, +} + +impl Geometry { + pub(crate) fn validate(&self) -> Result<(), Error> { + self.input.validate()?; + self.output.validate()?; + Ok(()) + } + + /// The samples in one plane, in `u64` so no `u32` dimensions overflow it. + pub(crate) fn pixels(&self) -> u64 { + let width = u64::from(self.width); + let height = u64::from(self.height); + + width * height + } + + /// Rejects a geometry whose ring of `slots` frames stores more than `u32::MAX` elements. + /// + /// The kernels index the frame ring and its accumulators with `u32`, so a larger ring would wrap. + pub(crate) fn check_ring_fits(&self, slots: u64) -> Result<(), Error> { + let stored_ch = u64::from(self.channels.storage_count()); + let elements = self + .pixels() + .checked_mul(stored_ch) + .and_then(|frame_len| frame_len.checked_mul(slots)); + let fits = elements.is_some_and(|elements| elements <= u64::from(u32::MAX)); + if !fits { + let message = format!( + "a {slots} frame ring of {}x{} planes stores more than u32::MAX elements", + self.width, self.height + ); + return Err(Error::InvalidGeometry(message)); + } + + Ok(()) + } + + /// Checks plane count, dimensions, and that every handle covers its whole words. + pub(crate) fn check_planes(&self, planes: &[DevicePlane<'_>], format: SampleFormat) -> Result<(), Error> { + let expected = self.channels.count() as usize; + if planes.len() != expected { + let message = format!("expected {expected} planes, got {}", planes.len()); + return Err(Error::PlaneMismatch(message)); + } + + let pixels = self.pixels(); + let needed = format.plane_bytes(pixels); + + for (index, plane) in planes.iter().enumerate() { + let matches_size = plane.width == self.width && plane.height == self.height; + if !matches_size { + let message = format!( + "plane {index} is {}x{}, expected {}x{}", + plane.width, plane.height, self.width, self.height + ); + return Err(Error::PlaneMismatch(message)); + } + + let available = plane.handle.size_in_used(); + if available < needed { + let message = format!("plane {index} handle holds {available} bytes, needs {needed}"); + return Err(Error::PlaneMismatch(message)); + } + } + + Ok(()) + } +} diff --git a/av-denoise-core/src/engine/tests/kernels.rs b/av-denoise-core/src/engine/tests/kernels.rs new file mode 100644 index 0000000..a2c7a03 --- /dev/null +++ b/av-denoise-core/src/engine/tests/kernels.rs @@ -0,0 +1,418 @@ +use cubecl::prelude::*; +use cubecl::wgpu::WgpuRuntime; + +use crate::engine::{DevicePlane, EgressSource, IngestTarget, SampleFormat, egress, ingest}; + +type R = WgpuRuntime; + +fn client() -> ComputeClient { + let device = ::Device::default(); + R::client(&device) +} + +/// Sample codes where the first three pixels are always 0, `max` and `max - 1`. +fn codes(pixels: usize, channel: u16, max: u16) -> Vec { + let modulus = max as u32 + 1; + let mut codes: Vec = (0..pixels) + .map(|pixel| ((pixel as u32 * 37 + channel as u32 * 11) % modulus) as u16) + .collect(); + + let extremes = [0, max, max - 1]; + for (code, extreme) in codes.iter_mut().zip(extremes) { + *code = extreme; + } + + codes +} + +fn encode(codes: &[u16], format: SampleFormat) -> Vec { + let mut bytes: Vec = match format { + SampleFormat::U8 => codes.iter().map(|&code| code as u8).collect(), + SampleFormat::U16 { .. } => codes.iter().flat_map(|code| code.to_le_bytes()).collect(), + SampleFormat::F32 => unreachable!("encode is only for word formats"), + }; + let padded = bytes.len().div_ceil(4) * 4; + bytes.resize(padded, 0); + + bytes +} + +/// The GPU divide is not correctly rounded, so a sample can differ from the host by a few units in +/// the last place. +fn assert_close(actual: &[f32], expected: &[f32]) { + assert_eq!(actual.len(), expected.len()); + + for (index, (&actual, &expected)) in actual.iter().zip(expected).enumerate() { + let tolerance = expected.abs() * f32::EPSILON * 4.0; + let difference = (actual - expected).abs(); + assert!( + difference <= tolerance, + "sample {index}: got {actual}, expected {expected}" + ); + } +} + +fn run_ingest( + width: u32, + height: u32, + channels: u32, + stored_channels: u32, + format: SampleFormat, +) -> (Vec, Vec) { + let client = client(); + let pixels = (width * height) as usize; + let max = format.max_value(); + + let channel_codes: Vec> = (0..channels as u16) + .map(|channel| codes(pixels, channel, max as u16)) + .collect(); + let handles: Vec<_> = channel_codes + .iter() + .map(|codes| { + let bytes = encode(codes, format); + client.create_from_slice(&bytes) + }) + .collect(); + let planes: Vec<_> = handles + .iter() + .map(|handle| DevicePlane::new(handle, width, height)) + .collect(); + + // Two slots, writing into the second, so the offset is exercised. + let frame_len = pixels * stored_channels as usize; + let ring_values = vec![-1.0f32; frame_len * 2]; + let ring_bytes = f32::as_bytes(&ring_values); + let ring = client.create_from_slice(ring_bytes); + let placeholder = client.empty(4); + let target = IngestTarget { + ring: &ring, + ring_len: frame_len * 2, + offset: frame_len as u32, + pixels: pixels as u32, + channels, + stored_ch: stored_channels, + }; + + ingest(&client, &planes, format, &placeholder, target); + + let bytes = client.read_one(ring).expect("read ring"); + let actual = f32::from_bytes(&bytes)[frame_len..].to_vec(); + + let mut expected = vec![0.0f32; frame_len]; + for pixel in 0..pixels { + for channel in 0..channels as usize { + let code = channel_codes[channel][pixel]; + expected[pixel * stored_channels as usize + channel] = code as f32 / max; + } + } + + (actual, expected) +} + +#[test] +fn ingest_u8_luma_matches_the_host_oracle() { + let (actual, expected) = run_ingest(5, 3, 1, 1, SampleFormat::U8); + assert_close(&actual, &expected); +} + +#[test] +fn ingest_u8_odd_sized_chroma_matches_the_host_oracle() { + let (actual, expected) = run_ingest(3, 3, 2, 2, SampleFormat::U8); + assert_close(&actual, &expected); +} + +#[test] +fn ingest_u16_yuv_pads_the_fourth_lane_with_zero() { + let (actual, expected) = run_ingest(7, 5, 3, 4, SampleFormat::U16 { depth: 10 }); + assert_close(&actual, &expected); +} + +#[test] +fn ingest_u16_at_depth_16_matches_the_host_oracle() { + let (actual, expected) = run_ingest(4, 4, 1, 1, SampleFormat::U16 { depth: 16 }); + assert_close(&actual, &expected); +} + +#[test] +fn ingest_u8_across_many_cubes_matches_the_host_oracle() { + let (actual, expected) = run_ingest(300, 3, 3, 4, SampleFormat::U8); + assert_close(&actual, &expected); +} + +#[test] +fn ingest_u16_across_many_cubes_matches_the_host_oracle() { + let (actual, expected) = run_ingest(300, 3, 2, 2, SampleFormat::U16 { depth: 12 }); + assert_close(&actual, &expected); +} + +#[test] +fn ingest_u16_at_depth_16_across_many_cubes_matches_the_host_oracle() { + let (actual, expected) = run_ingest(300, 3, 1, 1, SampleFormat::U16 { depth: 16 }); + assert_close(&actual, &expected); +} + +#[test] +fn ingest_f32_copies_samples_across_many_cubes() { + run_f32_ingest(300, 3); +} + +#[test] +fn ingest_f32_copies_samples_and_zeroes_padding() { + run_f32_ingest(4, 3); +} + +fn run_f32_ingest(width: u32, height: u32) { + let client = client(); + let pixels = (width * height) as usize; + let channel_values: Vec> = (0..3) + .map(|channel| { + let mut values: Vec = (0..pixels) + .map(|pixel| (pixel + channel * 100) as f32 * 0.001) + .collect(); + values[..3].copy_from_slice(&[0.0, 1.0, 1.0 - f32::EPSILON]); + + values + }) + .collect(); + let handles: Vec<_> = channel_values + .iter() + .map(|values| { + let bytes = f32::as_bytes(values); + client.create_from_slice(bytes) + }) + .collect(); + let planes: Vec<_> = handles + .iter() + .map(|handle| DevicePlane::new(handle, width, height)) + .collect(); + let ring_values = vec![-1.0f32; pixels * 4]; + let ring_bytes = f32::as_bytes(&ring_values); + let ring = client.create_from_slice(ring_bytes); + let placeholder = client.empty(4); + let target = IngestTarget { + ring: &ring, + ring_len: pixels * 4, + offset: 0, + pixels: pixels as u32, + channels: 3, + stored_ch: 4, + }; + + ingest(&client, &planes, SampleFormat::F32, &placeholder, target); + + let bytes = client.read_one(ring).expect("read ring"); + let actual = f32::from_bytes(&bytes); + + for pixel in 0..pixels { + for channel in 0..3 { + assert_eq!(actual[pixel * 4 + channel], channel_values[channel][pixel]); + } + + assert_eq!(actual[pixel * 4 + 3], 0.0); + } +} + +/// An interleaved frame whose first pixels hit below zero, zero, one, above one and the top two codes. +/// +/// The padding lane holds a value no plane should ever receive. +fn egress_frame(pixels: usize, channels: u32, stored_channels: u32, format: SampleFormat) -> Vec { + let max = format.max_value(); + let extremes = [ + -0.5, + 0.0, + 1.0, + 1.5, + (max - 0.4) / max, + (max - 1.0) / max, + (max - 0.6) / max, + ]; + + let mut frame = vec![0.0f32; pixels * stored_channels as usize]; + for pixel in 0..pixels { + for lane in 0..stored_channels as usize { + let index = pixel * stored_channels as usize + lane; + let value = if lane >= channels as usize { + 7.0 + } else if pixel < extremes.len() { + extremes[(pixel + lane) % extremes.len()] + } else { + ((index * 7919) % 1000) as f32 / 900.0 - 0.05 + }; + frame[index] = value; + } + } + + frame +} + +fn bytes_per_sample(format: SampleFormat) -> usize { + match format { + SampleFormat::U8 => 1, + _ => 2, + } +} + +/// The host quantisation, a clamp then a round half up. +fn quantise_planes(frame: &[f32], channels: u32, stored_channels: u32, format: SampleFormat) -> Vec> { + let pixels = frame.len() / stored_channels as usize; + let max = format.max_value(); + + (0..channels as usize) + .map(|channel| { + let mut plane = Vec::new(); + for pixel in 0..pixels { + let value = frame[pixel * stored_channels as usize + channel].clamp(0.0, 1.0); + let code = (value * max + 0.5) as u32; + match format { + SampleFormat::U8 => plane.push(code as u8), + _ => plane.extend_from_slice(&(code as u16).to_le_bytes()), + } + } + + plane + }) + .collect() +} + +/// Runs `egress` and returns each plane's bytes up to its last pixel, ignoring the final word's padding. +fn run_egress( + frame: &[f32], + width: u32, + height: u32, + channels: u32, + stored_channels: u32, + format: SampleFormat, +) -> Vec> { + let client = client(); + let pixels = (width * height) as usize; + let frame_bytes = f32::as_bytes(frame); + let frame_handle = client.create_from_slice(frame_bytes); + let plane_bytes = format.plane_bytes(pixels as u64) as usize; + let outputs: Vec<_> = (0..channels) + .map(|_| { + let sentinel = vec![0xAAu8; plane_bytes]; + client.create_from_slice(&sentinel) + }) + .collect(); + let planes: Vec<_> = outputs + .iter() + .map(|handle| DevicePlane::new(handle, width, height)) + .collect(); + let placeholder = client.empty(4); + let source = EgressSource { + frame: &frame_handle, + pixels: pixels as u32, + channels, + stored_ch: stored_channels, + }; + + egress(&client, source, &planes, format, &placeholder); + + let length = pixels * bytes_per_sample(format); + outputs + .into_iter() + .map(|handle| { + let bytes = client.read_one(handle).expect("read plane"); + bytes[..length].to_vec() + }) + .collect() +} + +fn assert_egress_matches_oracle( + width: u32, + height: u32, + channels: u32, + stored_channels: u32, + format: SampleFormat, +) { + let pixels = (width * height) as usize; + let frame = egress_frame(pixels, channels, stored_channels, format); + let actual = run_egress(&frame, width, height, channels, stored_channels, format); + let expected = quantise_planes(&frame, channels, stored_channels, format); + assert_eq!(actual, expected); +} + +#[test] +fn egress_u8_luma_matches_the_pack_oracle() { + assert_egress_matches_oracle(5, 3, 1, 1, SampleFormat::U8); +} + +#[test] +fn egress_u8_odd_sized_chroma_matches_the_pack_oracle() { + assert_egress_matches_oracle(3, 3, 2, 2, SampleFormat::U8); +} + +#[test] +fn egress_u16_yuv_skips_the_padding_lane() { + assert_egress_matches_oracle(7, 5, 3, 4, SampleFormat::U16 { depth: 12 }); +} + +#[test] +fn egress_u16_at_depth_16_hits_the_top_code() { + assert_egress_matches_oracle(4, 4, 1, 1, SampleFormat::U16 { depth: 16 }); +} + +#[test] +fn egress_u8_across_many_cubes_matches_the_pack_oracle() { + assert_egress_matches_oracle(300, 9, 3, 4, SampleFormat::U8); +} + +#[test] +fn egress_u8_luma_across_many_cubes_matches_the_pack_oracle() { + assert_egress_matches_oracle(300, 9, 1, 1, SampleFormat::U8); +} + +#[test] +fn egress_u16_across_many_cubes_matches_the_pack_oracle() { + assert_egress_matches_oracle(300, 9, 2, 2, SampleFormat::U16 { depth: 10 }); +} + +#[test] +fn egress_u16_at_depth_16_across_many_cubes_matches_the_pack_oracle() { + assert_egress_matches_oracle(300, 9, 3, 4, SampleFormat::U16 { depth: 16 }); +} + +#[test] +fn egress_f32_writes_unclamped_samples() { + assert_f32_egress_copies(3, 2, 2, 2); +} + +#[test] +fn egress_f32_skips_the_padding_lane() { + assert_f32_egress_copies(7, 5, 3, 4); +} + +#[test] +fn egress_f32_across_many_cubes_copies_samples() { + assert_f32_egress_copies(300, 9, 3, 4); +} + +fn assert_f32_egress_copies(width: u32, height: u32, channels: u32, stored_channels: u32) { + let client = client(); + let pixels = (width * height) as usize; + let frame = egress_frame(pixels, channels, stored_channels, SampleFormat::F32); + let frame_bytes = f32::as_bytes(&frame); + let frame_handle = client.create_from_slice(frame_bytes); + let outputs: Vec<_> = (0..channels).map(|_| client.empty(pixels * 4)).collect(); + let planes: Vec<_> = outputs + .iter() + .map(|handle| DevicePlane::new(handle, width, height)) + .collect(); + let placeholder = client.empty(4); + let source = EgressSource { + frame: &frame_handle, + pixels: pixels as u32, + channels, + stored_ch: stored_channels, + }; + + egress(&client, source, &planes, SampleFormat::F32, &placeholder); + + for (channel, handle) in outputs.into_iter().enumerate() { + let bytes = client.read_one(handle).expect("read plane"); + let values = f32::from_bytes(&bytes); + for pixel in 0..pixels { + assert_eq!(values[pixel], frame[pixel * stored_channels as usize + channel]); + } + } +} diff --git a/av-denoise-core/src/engine/tests/mod.rs b/av-denoise-core/src/engine/tests/mod.rs new file mode 100644 index 0000000..d4b1d13 --- /dev/null +++ b/av-denoise-core/src/engine/tests/mod.rs @@ -0,0 +1,4 @@ +mod plane; + +#[cfg(any(feature = "vulkan", feature = "metal"))] +mod kernels; diff --git a/av-denoise-core/src/engine/tests/plane.rs b/av-denoise-core/src/engine/tests/plane.rs new file mode 100644 index 0000000..8b0ffd0 --- /dev/null +++ b/av-denoise-core/src/engine/tests/plane.rs @@ -0,0 +1,137 @@ +use crate::engine::{Geometry, SampleFormat}; +use crate::error::Error; +use crate::nlmeans::ChannelMode; + +#[test] +fn u16_depth_outside_9_to_16_is_rejected() { + for depth in [0, 8, 17] { + let format = SampleFormat::U16 { depth }; + let result = format.validate(); + assert!(matches!(result, Err(Error::InvalidGeometry(_)))); + } +} + +#[test] +fn u16_depth_inside_9_to_16_is_accepted() { + for depth in 9..=16 { + let format = SampleFormat::U16 { depth }; + assert!(format.validate().is_ok()); + } +} + +#[test] +fn max_value_matches_the_depth() { + assert_eq!(SampleFormat::U8.max_value(), 255.0); + assert_eq!(SampleFormat::U16 { depth: 10 }.max_value(), 1023.0); + assert_eq!(SampleFormat::U16 { depth: 16 }.max_value(), 65535.0); + assert_eq!(SampleFormat::F32.max_value(), 1.0); +} + +#[test] +fn plane_bytes_round_up_to_whole_words() { + assert_eq!(SampleFormat::U8.plane_bytes(9), 12); + assert_eq!(SampleFormat::U16 { depth: 10 }.plane_bytes(3), 8); + assert_eq!(SampleFormat::F32.plane_bytes(3), 12); +} + +fn sized_geometry(width: u32, height: u32, channels: ChannelMode) -> Geometry { + Geometry { + width, + height, + channels, + input: SampleFormat::F32, + output: SampleFormat::F32, + } +} + +#[test] +fn pixels_do_not_overflow_at_the_largest_dimensions() { + let geometry = sized_geometry(u32::MAX, u32::MAX, ChannelMode::Luma); + let expected = u64::from(u32::MAX) * u64::from(u32::MAX); + assert_eq!(geometry.pixels(), expected); +} + +#[test] +fn a_ring_of_exactly_u32_max_elements_fits() { + let geometry = sized_geometry(65_535, 65_537, ChannelMode::Luma); + let result = geometry.check_ring_fits(1); + assert!(result.is_ok()); +} + +#[test] +fn a_ring_one_frame_past_u32_max_elements_is_invalid_geometry() { + let geometry = sized_geometry(65_535, 65_537, ChannelMode::Luma); + let result = geometry.check_ring_fits(2); + assert!(matches!(result, Err(Error::InvalidGeometry(_)))); +} + +#[test] +fn a_ring_whose_size_overflows_u64_is_invalid_geometry() { + let geometry = sized_geometry(u32::MAX, u32::MAX, ChannelMode::Yuv); + let result = geometry.check_ring_fits(u64::MAX); + assert!(matches!(result, Err(Error::InvalidGeometry(_)))); +} + +#[cfg(any(feature = "vulkan", feature = "metal"))] +mod with_handles { + use cubecl::prelude::*; + use cubecl::wgpu::WgpuRuntime; + + use super::*; + use crate::engine::DevicePlane; + + fn client() -> ComputeClient { + let device = ::Device::default(); + WgpuRuntime::client(&device) + } + + fn geometry(channels: ChannelMode) -> Geometry { + Geometry { + width: 3, + height: 3, + channels, + input: SampleFormat::U8, + output: SampleFormat::U8, + } + } + + #[test] + fn rejects_the_wrong_plane_count() { + let client = client(); + let handle = client.empty(12); + let plane = DevicePlane::new(&handle, 3, 3); + let geometry = geometry(ChannelMode::Chroma); + let result = geometry.check_planes(&[plane], SampleFormat::U8); + assert!(matches!(result, Err(Error::PlaneMismatch(_)))); + } + + #[test] + fn rejects_mismatched_plane_dimensions() { + let client = client(); + let handle = client.empty(16); + let plane = DevicePlane::new(&handle, 4, 3); + let geometry = geometry(ChannelMode::Luma); + let result = geometry.check_planes(&[plane], SampleFormat::U8); + assert!(matches!(result, Err(Error::PlaneMismatch(_)))); + } + + #[test] + fn rejects_plane_handle_shorter_than_its_words() { + let client = client(); + let handle = client.empty(9); + let plane = DevicePlane::new(&handle, 3, 3); + let geometry = geometry(ChannelMode::Luma); + let result = geometry.check_planes(&[plane], SampleFormat::U8); + assert!(matches!(result, Err(Error::PlaneMismatch(_)))); + } + + #[test] + fn accepts_a_word_padded_plane() { + let client = client(); + let handle = client.empty(12); + let plane = DevicePlane::new(&handle, 3, 3); + let geometry = geometry(ChannelMode::Luma); + let result = geometry.check_planes(&[plane], SampleFormat::U8); + assert!(result.is_ok()); + } +} diff --git a/av-denoise-core/src/error.rs b/av-denoise-core/src/error.rs new file mode 100644 index 0000000..34e3475 --- /dev/null +++ b/av-denoise-core/src/error.rs @@ -0,0 +1,20 @@ +/// Errors returned by the denoising engines. +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("invalid options, {0}")] + InvalidOptions(String), + #[error("invalid geometry, {0}")] + InvalidGeometry(String), + #[error("plane mismatch, {0}")] + PlaneMismatch(String), + #[error("every ready frame must be emitted before the next push")] + OutputsPending, + #[error("no frame is ready to emit")] + NothingToEmit, + #[error("context frames must come before the first push of a stream")] + ContextAfterPush, + #[error("an earlier call failed, reset the engine before using it again")] + NeedsReset, + #[error(transparent)] + Gpu(#[from] anyhow::Error), +} diff --git a/av-denoise-core/src/frame/mod.rs b/av-denoise-core/src/frame/mod.rs deleted file mode 100644 index 05d0e0f..0000000 --- a/av-denoise-core/src/frame/mod.rs +++ /dev/null @@ -1,1559 +0,0 @@ -mod reseed; - -use std::collections::VecDeque; - -pub use self::reseed::ReseedWindow; -use crate::accelerate::Accelerator; -use crate::nl4d::grain::GrainChunk; -use crate::{ - Algorithm, - ChannelMode, - Denoiser, - DenoiserError, - DenoiserOptions, - DenoisingMode, - Depth, - Device, - FrameOutput, - Nl4dOptions, - NlmTuning, - NlmeansHqOptions, - NlmeansOptions, - OutputFormat, - WindowSpan, -}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum Subsampling { - Yuv420, - Yuv422, - Yuv444, -} - -impl Subsampling { - /// Halved axes round up, so an odd dimension keeps the extra sample, - /// matching what y4m and ffmpeg do. - pub fn chroma_dims(self, w: u32, h: u32) -> (u32, u32) { - match self { - Subsampling::Yuv420 => (w.div_ceil(2), h.div_ceil(2)), - Subsampling::Yuv422 => (w.div_ceil(2), h), - Subsampling::Yuv444 => (w, h), - } - } -} - -#[derive(Debug, Clone, Copy)] -pub struct FrameLayout { - pub width: u32, - pub height: u32, - pub subsampling: Subsampling, - pub depth: Depth, -} - -impl FrameLayout { - pub fn luma_pixels(&self) -> usize { - (self.width as usize) * (self.height as usize) - } - - pub fn chroma_dims(&self) -> (u32, u32) { - self.subsampling.chroma_dims(self.width, self.height) - } - - pub fn chroma_pixels(&self) -> usize { - let (w, h) = self.chroma_dims(); - (w as usize) * (h as usize) - } - - /// Wire size of the luma plane. - pub fn luma_bytes(&self) -> usize { - self.luma_pixels() * self.depth.bytes_per_sample() - } - - /// Wire size of one chroma plane. - pub fn chroma_bytes(&self) -> usize { - self.chroma_pixels() * self.depth.bytes_per_sample() - } - - /// A full black luma plane, used when no luma source is available. - pub fn black_luma_plane(&self) -> Vec { - fill_plane(self.luma_pixels(), 0, self.depth) - } - - /// A full neutral chroma plane, used when a source has no chroma. - pub fn neutral_chroma_plane(&self) -> Vec { - fill_plane(self.chroma_pixels(), self.depth.neutral_chroma(), self.depth) - } -} - -/// Builds a plane of `samples` copies of `value` in wire-byte form. -pub fn fill_plane(samples: usize, value: u16, depth: Depth) -> Vec { - match depth.bytes_per_sample() { - 1 => vec![value as u8; samples], - _ => { - let word = value.to_le_bytes(); - let mut out = Vec::with_capacity(samples * 2); - for _ in 0..samples { - out.extend_from_slice(&word); - } - out - }, - } -} - -/// A planar YUV frame holding little-endian wire bytes. -/// -/// Plane lengths come from [`FrameLayout`], so `y.len()` is -/// `layout.luma_bytes()` and both `u.len()` and `v.len()` are -/// `layout.chroma_bytes()`. -#[derive(Debug, Clone)] -pub struct Planes { - pub y: Vec, - pub u: Vec, - pub v: Vec, -} - -/// Which planes a caller wants cleaned, once `--channel-mode` (or the -/// equivalent host option) has been resolved. -/// -/// This is separate from the library's [`ChannelMode`] because this layer -/// may run more than one `Denoiser` in lockstep, one for luma and one for -/// chroma. It may also run a single fused three-channel denoiser instead. -/// Which of those applies depends on the caller's channel selection and -/// the source's chroma subsampling. -#[derive(Debug, Copy, Clone, PartialEq, Eq)] -pub enum ChannelIntent { - /// Denoise luma only. Chroma passes through. - Luma, - /// Denoise chroma only. Luma passes through. - Chroma, - /// Denoise both luma and chroma as two independent denoisers. - /// Chroma runs at the source's native subsampled resolution. - LumaChroma, - /// A single library `Denoiser` running the fused three-channel - /// kernel. Needs a YUV444 source, which is checked at ingest setup - /// time. - YuvFused, -} - -impl ChannelIntent { - /// Rejects the intent if the source's subsampling cannot support it. - pub fn validate_for_source(self, layout: FrameLayout) -> Result<(), anyhow::Error> { - match self { - ChannelIntent::YuvFused if layout.subsampling != Subsampling::Yuv444 => { - anyhow::bail!( - "--channel-mode yuv requires a YUV444 source, got {:?}. Convert the input first, for example with `ffmpeg -pix_fmt yuv444p`", - layout.subsampling - ); - }, - _ => Ok(()), - } - } -} - -/// The per-plane option set a caller resolves once and passes into -/// [`PlanarDenoiser::create`]. -#[derive(Debug, Clone)] -pub struct PlaneOptions { - pub accelerators: Vec, - pub device: Device, - pub intent: ChannelIntent, - pub mode: DenoisingMode, - /// Which denoising algorithm to run, along with the settings only - /// that algorithm reads. - pub algorithm: Algorithm, - /// Per-plane strength override for the luma denoiser. Takes - /// precedence over the algorithm's own `tuning.strength` when set. - /// Only has an effect on the two NLM algorithms. - pub luma_strength: Option, - /// Per-plane strength override for the chroma denoiser. Takes - /// precedence over the algorithm's own `tuning.strength` when set. - /// Only has an effect on the two NLM algorithms. - pub chroma_strength: Option, - /// Per-plane override for `lambda_ht`, luma. Takes precedence over - /// `algorithm`'s value when set, which itself falls back to a - /// calibrated per-plane default when nothing at all is set. Only - /// has an effect when `algorithm` is `Algorithm::Nl4d`, where it - /// pins the temporal grouping stage's hard threshold. - pub luma_lambda_ht: Option, - /// Per-plane override for `lambda_ht`, chroma. Takes precedence over - /// `algorithm`'s value when set, which itself falls back to a - /// calibrated per-plane default when nothing at all is set. Only - /// has an effect when `algorithm` is `Algorithm::Nl4d`, where it - /// pins the temporal grouping stage's hard threshold. - pub chroma_lambda_ht: Option, -} - -impl PlaneOptions { - /// Resolves `self.algorithm` for one plane, folding in the per-plane - /// overrides that apply to whichever algorithm `self.algorithm` is. - /// - /// For the two NLM algorithms that is `strength`. For `Nl4d` it is - /// `lambda_ht`, since nl4d has no NLM weighting pass for a strength - /// to affect. - /// - /// `Nl4d`'s `lambda_ht` stays `Option` all the way through - /// this method. When neither a per-plane flag nor the matching - /// shared flag was set, the result is `None`, deferred to - /// `nl4d_default_lambda_ht` at construction, once the plane being - /// denoised is known there too. That is what gives luma and chroma - /// different values when a caller passes no flags at all. - fn algorithm_for(&self, channels: ChannelMode) -> Algorithm { - let per_plane = |luma, chroma| match channels { - ChannelMode::Luma => luma, - ChannelMode::Chroma => chroma, - ChannelMode::Yuv => None, - }; - - match self.algorithm { - Algorithm::Nl4d(nl4d) => Algorithm::Nl4d(Nl4dOptions { - // Left unresolved when unset, since the calibrated - // default depends on the plane, which - // `nl4d_default_lambda_ht` resolves at construction. - lambda_ht: per_plane(self.luma_lambda_ht, self.chroma_lambda_ht).or(nl4d.lambda_ht), - grain_export: nl4d.grain_export && channels != ChannelMode::Chroma, - ..nl4d - }), - Algorithm::Nlmeans(nlm) => { - let strength = per_plane(self.luma_strength, self.chroma_strength); - Algorithm::Nlmeans(with_plane_strength(nlm, strength)) - }, - Algorithm::NlmeansHq(opts) => { - let strength = per_plane(self.luma_strength, self.chroma_strength); - Algorithm::NlmeansHq(NlmeansHqOptions { - nlm: with_plane_strength(opts.nlm, strength), - ..opts - }) - }, - } - } - - /// `depth` is the source's wire depth, which every denoiser - /// quantises to on the GPU. - fn denoiser_options(&self, channels: ChannelMode, depth: Depth) -> DenoiserOptions { - DenoiserOptions::builder() - .channel_mode(channels) - .mode(self.mode) - .algorithm(self.algorithm_for(channels)) - .output_format(OutputFormat::Wire { depth }) - .build() - } -} - -/// `nlm` with `strength` replaced by the per-plane override, when there -/// is one. An unset override leaves the shared value alone. -fn with_plane_strength(nlm: NlmeansOptions, strength: Option) -> NlmeansOptions { - match strength { - None => nlm, - Some(strength) => NlmeansOptions { - tuning: NlmTuning { - strength: Some(strength), - ..nlm.tuning - }, - ..nlm - }, - } -} - -/// Reads the result of a `PlanarDenoiser::push` call for the -/// push-then-drain-then-retry loop. -/// -/// `Ok(false)` means the push landed. `Ok(true)` means the queue was -/// full, so the caller should drain one output and push again. -/// -/// Any error other than `QueueFull` is passed on rather than discarded. -pub fn push_needs_retry(result: Result<(), DenoiserError>) -> Result { - match result { - Ok(()) => Ok(false), - Err(DenoiserError::QueueFull) => Ok(true), - Err(other) => Err(other.into()), - } -} - -/// Unwraps a denoised frame from one of the `Denoiser`s -/// [`PlanarDenoiser`] builds. -/// -/// Those are always built in [`crate::OutputFormat::Wire`], so the other -/// variant never reaches here. -fn expect_wire(out: FrameOutput) -> Vec { - out.into_wire() - .expect("PlanarDenoiser builds every Denoiser in wire output format") -} - -/// Splits a fused YUV444 wire frame into its three planes. -/// -/// The pack kernel leaves a three-channel frame interleaved, so this is -/// the byte-level counterpart of the host converter it replaced. That one -/// lives in `converter_tests` now, as the oracle this is checked against. -fn split_yuv_wire(wire: &[u8], depth: Depth) -> Planes { - let bytes = depth.bytes_per_sample(); - let pixels = wire.len() / (3 * bytes); - - let mut y = Vec::with_capacity(pixels * bytes); - let mut u = Vec::with_capacity(pixels * bytes); - let mut v = Vec::with_capacity(pixels * bytes); - - for pixel in wire.chunks_exact(3 * bytes) { - y.extend_from_slice(&pixel[..bytes]); - u.extend_from_slice(&pixel[bytes..2 * bytes]); - v.extend_from_slice(&pixel[2 * bytes..]); - } - - Planes { y, u, v } -} - -/// Splits a chroma wire frame into its U and V planes. -/// -/// The pack kernel writes U's whole region first and V's after it, so -/// each plane is one contiguous half of the buffer. -fn split_uv_wire(wire: &[u8]) -> (Vec, Vec) { - let (u, v) = wire.split_at(wire.len() / 2); - (u.to_vec(), v.to_vec()) -} - -/// The push a [`PlanarDenoiser`] runs against each enabled half, either -/// [`Denoiser::push_frame_wire`] or -/// [`Denoiser::push_frame_wire_priming`]. -type WirePush = fn(&mut Denoiser, &[&[u8]], Depth) -> Result<(), DenoiserError>; - -/// Wraps the luma and chroma `Denoiser` instances needed for one -/// subsampled YUV source. -/// -/// The caller pushes planar frames in and gets planar frames out. The -/// luma and chroma split is invisible from the outside. -pub struct PlanarDenoiser { - layout: FrameLayout, - luma: Option, - chroma: Option, - /// Set when the intent is `YuvFused`, in which case `luma` and - /// `chroma` are both unset. - yuv: Option, - // Source planes queued for passthrough when the matching denoiser is - // disabled. Only the disabled side's queue is ever filled. Entries - // are popped one per frame the enabled side emits, so temporal - // delays stay aligned. - luma_passthrough: VecDeque>, - chroma_passthrough: VecDeque<(Vec, Vec)>, - /// The temporal radius every owned denoiser runs at, resolved from - /// `opts.mode` at construction. - temporal_radius: u32, -} - -impl PlanarDenoiser { - pub fn create(opts: &PlaneOptions, layout: FrameLayout) -> Result { - let (chroma_w, chroma_h) = layout.chroma_dims(); - - if chroma_w == 0 || chroma_h == 0 { - anyhow::bail!( - "frame dimensions {}x{} are too small for subsampling {:?}", - layout.width, - layout.height, - layout.subsampling - ); - } - - opts.intent.validate_for_source(layout)?; - - let (denoise_luma, denoise_chroma, denoise_yuv) = match opts.intent { - ChannelIntent::Luma => (true, false, false), - ChannelIntent::Chroma => (false, true, false), - ChannelIntent::LumaChroma => (true, true, false), - ChannelIntent::YuvFused => (false, false, true), - }; - - let luma = denoise_luma - .then(|| { - Denoiser::create( - &opts.accelerators, - &opts.device, - layout.width, - layout.height, - opts.denoiser_options(ChannelMode::Luma, layout.depth), - ) - }) - .transpose()?; - - let chroma = denoise_chroma - .then(|| { - Denoiser::create( - &opts.accelerators, - &opts.device, - chroma_w, - chroma_h, - opts.denoiser_options(ChannelMode::Chroma, layout.depth), - ) - }) - .transpose()?; - - let yuv = denoise_yuv - .then(|| { - Denoiser::create( - &opts.accelerators, - &opts.device, - layout.width, - layout.height, - opts.denoiser_options(ChannelMode::Yuv, layout.depth), - ) - }) - .transpose()?; - - let temporal_radius = match opts.mode { - DenoisingMode::Spacial => 0, - DenoisingMode::Temporal { radius } => radius, - }; - - Ok(Self { - layout, - luma, - chroma, - yuv, - luma_passthrough: VecDeque::new(), - chroma_passthrough: VecDeque::new(), - temporal_radius, - }) - } - - /// The temporal radius the underlying denoisers run at. - pub fn temporal_radius(&self) -> u32 { - self.temporal_radius - } - - /// Pushes one planar frame. - /// - /// On `QueueFull` the caller should receive one frame and then retry - /// the whole call. Any other error is passed on unchanged. - /// - /// The denoiser push runs before either passthrough queue is - /// touched, so a retry replays the whole frame cleanly instead of - /// queueing the disabled side's plane twice. - /// - /// # Why a retry cannot duplicate a frame - /// - /// In `LumaChroma` mode `luma` and `chroma` are both real - /// `Denoiser`s with their own queues. A retry pushes again into - /// whichever half already succeeded, which would duplicate that - /// half's frame if the two could ever sit at different fill levels. - /// - /// They cannot. Both are built from the same `opts.mode`, so they - /// share a temporal radius and a `MAX_PENDING` ceiling. Every - /// successful push or receive moves both on by exactly one frame, - /// and a failed push moves neither, because the `QueueFull` check - /// runs before anything changes. - /// - /// So the two halves always enter this function with the same frame - /// count and the same pending depth, and the `QueueFull` check - /// inside `push_frame_wire` answers the same way for each. If the luma - /// push succeeds then the chroma push succeeds too, which makes the - /// duplicate unreachable. - pub fn push(&mut self, planes: &Planes) -> Result<(), DenoiserError> { - self.push_with(planes, Denoiser::push_frame_wire) - } - - /// Uploads one planar frame into the temporal window without starting - /// a denoise. - /// - /// Mirrors [`Self::push`], down to queueing the disabled side's - /// passthrough plane, but no output is ever produced for this call. - /// This is how the reseed paths fill a window's leading frames before - /// its real pushes start. - fn push_priming(&mut self, planes: &Planes) -> Result<(), DenoiserError> { - self.push_with(planes, Denoiser::push_frame_wire_priming) - } - - /// Shared body of [`Self::push`] and [`Self::push_priming`]. - /// - /// `push_frame` is [`Denoiser::push_frame_wire`] for a real push or - /// [`Denoiser::push_frame_wire_priming`] for a priming one, run - /// against whichever of `yuv`, `luma`, and `chroma` is enabled. - /// - /// The planes go over as wire bytes, so the normalisation and the - /// channel interleave both happen on the GPU. - fn push_with(&mut self, planes: &Planes, push_frame: WirePush) -> Result<(), DenoiserError> { - let depth = self.layout.depth; - - if let Some(d) = self.yuv.as_mut() { - push_frame(d, &[&planes.y, &planes.u, &planes.v], depth)?; - return Ok(()); - } - - if let Some(d) = self.luma.as_mut() { - push_frame(d, &[&planes.y], depth)?; - } - - if let Some(d) = self.chroma.as_mut() { - push_frame(d, &[&planes.u, &planes.v], depth)?; - } - - if self.luma.is_none() { - self.luma_passthrough.push_back(planes.y.clone()); - } - - if self.chroma.is_none() { - self.chroma_passthrough - .push_back((planes.u.clone(), planes.v.clone())); - } - - Ok(()) - } - - /// Blocks until each enabled half emits one frame, then reassembles - /// them into a planar frame. - /// - /// Returns `Ok(None)` if neither half had pending output. - pub fn recv(&mut self) -> Result, anyhow::Error> { - if let Some(d) = self.yuv.as_mut() { - return match d.recv_frame()? { - Some(packed) => Ok(Some(split_yuv_wire(&expect_wire(packed), self.layout.depth))), - None => Ok(None), - }; - } - - let luma_out = self - .luma - .as_mut() - .map(|d| d.recv_frame()) - .transpose()? - .flatten() - .map(expect_wire); - - let chroma_out = self - .chroma - .as_mut() - .map(|d| d.recv_frame()) - .transpose()? - .flatten() - .map(expect_wire); - - // A disabled side has no Denoiser to query. When the enabled side - // produced output, pop the matching source plane from the - // disabled side's passthrough queue instead. - let luma_passthrough = if self.luma.is_none() && chroma_out.is_some() { - self.luma_passthrough.pop_front() - } else { - None - }; - - let chroma_passthrough = if self.chroma.is_none() && luma_out.is_some() { - self.chroma_passthrough.pop_front() - } else { - None - }; - - if luma_out.is_none() && chroma_out.is_none() { - return Ok(None); - } - - let planes = self.assemble(luma_out, chroma_out, luma_passthrough, chroma_passthrough); - - Ok(Some(planes)) - } - - /// Reads back the luma grain chunks measured since the last call, in frame order. - /// - /// Measured chunks stay on the GPU until drained, so callers drain after each flush. - pub fn drain_grain_chunks(&mut self) -> Result, anyhow::Error> { - let source = self.luma.as_mut().or(self.yuv.as_mut()); - let Some(denoiser) = source else { - return Ok(Vec::new()); - }; - - let chunks = denoiser.drain_grain_chunks()?; - - Ok(chunks) - } - - /// Drains the temporal tail of both halves. - /// - /// `sink` is called once per emitted planar frame. - pub fn flush(&mut self, mut sink: impl FnMut(Planes)) -> Result<(), anyhow::Error> { - if let Some(d) = self.yuv.as_mut() { - let depth = self.layout.depth; - d.flush(|packed| sink(split_yuv_wire(&expect_wire(packed), depth)))?; - return Ok(()); - } - - let mut luma_buf: Vec> = Vec::new(); - let mut chroma_buf: Vec> = Vec::new(); - - if let Some(d) = self.luma.as_mut() { - d.flush(|v| luma_buf.push(expect_wire(v)))?; - } - - if let Some(d) = self.chroma.as_mut() { - d.flush(|v| chroma_buf.push(expect_wire(v)))?; - } - - // The two halves run in lockstep, so they flush the same number - // of frames. For each emitted frame the disabled side, if there - // is one, pops the matching source plane from its passthrough - // queue. - let count = luma_buf.len().max(chroma_buf.len()); - - for i in 0..count { - let y = if let Some(buf) = luma_buf.get_mut(i) { - std::mem::take(buf) - } else if let Some(src) = self.luma_passthrough.pop_front() { - src - } else { - self.layout.black_luma_plane() - }; - - let (u, v) = if let Some(packed) = chroma_buf.get(i) { - split_uv_wire(packed) - } else if let Some((src_u, src_v)) = self.chroma_passthrough.pop_front() { - (src_u, src_v) - } else { - ( - self.layout.neutral_chroma_plane(), - self.layout.neutral_chroma_plane(), - ) - }; - - sink(Planes { y, u, v }); - } - - if !self.luma_passthrough.is_empty() || !self.chroma_passthrough.is_empty() { - tracing::warn!( - luma_remaining = self.luma_passthrough.len(), - chroma_remaining = self.chroma_passthrough.len(), - "passthrough queue not fully drained after flush", - ); - self.luma_passthrough.clear(); - self.chroma_passthrough.clear(); - } - - Ok(()) - } - - /// The number of frames behind and ahead of a target frame a - /// [`Self::reseed`] window must supply, for whichever algorithm this - /// `PlanarDenoiser` runs. - /// - /// Every owned `Denoiser` was built from the same algorithm, so any - /// one of them answers for all of them. - pub fn window_span(&self) -> WindowSpan { - self.yuv - .as_ref() - .or(self.luma.as_ref()) - .or(self.chroma.as_ref()) - .expect("PlanarDenoiser always keeps at least one Denoiser") - .window_span() - } - - fn assemble( - &self, - luma: Option>, - chroma: Option>, - luma_passthrough: Option>, - chroma_passthrough: Option<(Vec, Vec)>, - ) -> Planes { - let y = match (luma, luma_passthrough) { - (Some(v), _) => v, - (None, Some(src)) => src, - (None, None) => self.layout.black_luma_plane(), - }; - - let (u, v) = match (chroma, chroma_passthrough) { - (Some(packed), _) => split_uv_wire(&packed), - (None, Some(src)) => src, - (None, None) => ( - self.layout.neutral_chroma_plane(), - self.layout.neutral_chroma_plane(), - ), - }; - - Planes { y, u, v } - } -} - -/// Reads and writes samples in one wire format. -/// -/// The implementor is chosen once per conversion, which keeps the -/// per-sample path free of depth branches. -trait SampleCodec { - const BYTES: usize; - - fn read(plane: &[u8], i: usize) -> u16; - fn write(plane: &mut [u8], i: usize, value: u16); -} - -/// One byte per sample. -struct Narrow; - -impl SampleCodec for Narrow { - const BYTES: usize = 1; - - #[inline(always)] - fn read(plane: &[u8], i: usize) -> u16 { - plane[i] as u16 - } - - #[inline(always)] - fn write(plane: &mut [u8], i: usize, value: u16) { - plane[i] = value as u8; - } -} - -/// Two bytes per sample, little-endian. -struct Wide; - -impl SampleCodec for Wide { - const BYTES: usize = 2; - - #[inline(always)] - fn read(plane: &[u8], i: usize) -> u16 { - u16::from_le_bytes([plane[2 * i], plane[2 * i + 1]]) - } - - #[inline(always)] - fn write(plane: &mut [u8], i: usize, value: u16) { - plane[2 * i..2 * i + 2].copy_from_slice(&value.to_le_bytes()); - } -} - -/// Quantises a normalised value to a native-depth sample. -#[inline(always)] -fn quantise(v: f32, max: f32) -> u16 { - (v.clamp(0.0, 1.0) * max + 0.5) as u16 -} - -/// Converts one wire-byte plane to normalised f32. -/// -/// `gpu_unpack_wire` does this on the device now. This host version is -/// the oracle that kernel is checked against. -pub fn plane_to_f32(plane: &[u8], depth: Depth) -> Vec { - let max = depth.max_value(); - - fn run(plane: &[u8], max: f32) -> Vec { - let samples = plane.len() / C::BYTES; - (0..samples).map(|i| C::read(plane, i) as f32 / max).collect() - } - - match depth.bytes_per_sample() { - 1 => run::(plane, max), - _ => run::(plane, max), - } -} - -/// Reverse of [`plane_to_f32`]. -pub fn f32_to_plane(plane: &[f32], depth: Depth) -> Vec { - let max = depth.max_value(); - - fn run(plane: &[f32], max: f32) -> Vec { - let mut out = vec![0u8; plane.len() * C::BYTES]; - for (i, &v) in plane.iter().enumerate() { - C::write(&mut out, i, quantise(v, max)); - } - out - } - - match depth.bytes_per_sample() { - 1 => run::(plane, max), - _ => run::(plane, max), - } -} - -/// Interleaves equal-length Y, U, and V planes from a YUV444 source into -/// `[Y0, U0, V0, Y1, U1, V1, ...]` as f32 in `[0, 1]`. -/// -/// This is the layout the library's fused three-channel kernel expects. -/// -/// `gpu_unpack_wire` does this on the device now. This host version is -/// the oracle that kernel is checked against. -pub fn interleave_yuv_to_f32(y: &[u8], u: &[u8], v: &[u8], depth: Depth) -> Vec { - debug_assert_eq!(y.len(), u.len()); - debug_assert_eq!(u.len(), v.len()); - - let max = depth.max_value(); - - fn run(y: &[u8], u: &[u8], v: &[u8], max: f32) -> Vec { - let pixels = y.len() / C::BYTES; - let mut out = Vec::with_capacity(pixels * 3); - - for i in 0..pixels { - out.push(C::read(y, i) as f32 / max); - out.push(C::read(u, i) as f32 / max); - out.push(C::read(v, i) as f32 / max); - } - - out - } - - match depth.bytes_per_sample() { - 1 => run::(y, u, v, max), - _ => run::(y, u, v, max), - } -} - -/// Interleaves separate U and V planes into `[U, V, U, V, ...]` as f32 -/// in `[0, 1]`. -/// -/// `gpu_unpack_wire` does this on the device now. This host version is -/// the oracle that kernel is checked against. -pub fn interleave_uv_to_f32(u: &[u8], v: &[u8], depth: Depth) -> Vec { - debug_assert_eq!(u.len(), v.len()); - - let max = depth.max_value(); - - fn run(u: &[u8], v: &[u8], max: f32) -> Vec { - let pixels = u.len() / C::BYTES; - let mut out = Vec::with_capacity(pixels * 2); - - for i in 0..pixels { - out.push(C::read(u, i) as f32 / max); - out.push(C::read(v, i) as f32 / max); - } - - out - } - - match depth.bytes_per_sample() { - 1 => run::(u, v, max), - _ => run::(u, v, max), - } -} - -/// Reverse of [`interleave_uv_to_f32`]. -pub fn unpack_uv_from_f32(packed: &[f32], chroma_pixels: usize, depth: Depth) -> (Vec, Vec) { - debug_assert_eq!(packed.len(), 2 * chroma_pixels); - - let max = depth.max_value(); - - fn run(packed: &[f32], chroma_pixels: usize, max: f32) -> (Vec, Vec) { - let mut u = vec![0u8; chroma_pixels * C::BYTES]; - let mut v = vec![0u8; chroma_pixels * C::BYTES]; - - for (i, chunk) in packed.as_chunks::<2>().0.iter().enumerate() { - C::write(&mut u, i, quantise(chunk[0], max)); - C::write(&mut v, i, quantise(chunk[1], max)); - } - - (u, v) - } - - match depth.bytes_per_sample() { - 1 => run::(packed, chroma_pixels, max), - _ => run::(packed, chroma_pixels, max), - } -} - -#[cfg(test)] -mod converter_tests { - use super::*; - - /// Reverse of [`interleave_yuv_to_f32`], and the oracle - /// `split_yuv_wire` is checked against. - /// - /// Production splits a fused YUV frame from the wire bytes the GPU - /// already quantised. This host version stays because cubecl can - /// compile a kernel to nothing without reporting an error, and a - /// kernel compared against itself compares zeros to zeros. - fn unpack_yuv_from_f32(packed: &[f32], pixels: usize, depth: Depth) -> Planes { - debug_assert_eq!(packed.len(), 3 * pixels); - - let max = depth.max_value(); - - fn run(packed: &[f32], pixels: usize, max: f32) -> Planes { - let mut y = vec![0u8; pixels * C::BYTES]; - let mut u = vec![0u8; pixels * C::BYTES]; - let mut v = vec![0u8; pixels * C::BYTES]; - - for (i, chunk) in packed.as_chunks::<3>().0.iter().enumerate() { - C::write(&mut y, i, quantise(chunk[0], max)); - C::write(&mut u, i, quantise(chunk[1], max)); - C::write(&mut v, i, quantise(chunk[2], max)); - } - - Planes { y, u, v } - } - - match depth.bytes_per_sample() { - 1 => run::(packed, pixels, max), - _ => run::(packed, pixels, max), - } - } - - /// Encodes native-depth samples into wire bytes, the inverse of what - /// the converters read. - fn wire(samples: &[u16], depth: Depth) -> Vec { - match depth.bytes_per_sample() { - 1 => samples.iter().map(|&s| s as u8).collect(), - _ => samples.iter().flat_map(|&s| s.to_le_bytes()).collect(), - } - } - - #[test] - fn plane_round_trips_boundary_codes_at_every_depth() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let max = depth.max_value() as u16; - let samples: Vec = vec![0, 1, 16, 64, 235, max / 2, max - 1, max] - .into_iter() - .filter(|&s| s <= max) - .collect(); - - let bytes = wire(&samples, depth); - let restored = f32_to_plane(&plane_to_f32(&bytes, depth), depth); - - assert_eq!(restored, bytes, "plane round trip failed at {depth:?}"); - } - } - - /// Samples above 8 bits are little-endian on the wire regardless of - /// host endianness. - #[test] - fn high_depth_samples_are_little_endian() { - // 1023 = 0x03FF -> [0xFF, 0x03] - let bytes = wire(&[1023, 0, 512], Depth::Ten); - assert_eq!(bytes, vec![0xFF, 0x03, 0x00, 0x00, 0x00, 0x02]); - - let f = plane_to_f32(&bytes, Depth::Ten); - assert!( - (f[0] - 1.0).abs() < 1e-6, - "0x03FF should normalize to 1.0, got {}", - f[0] - ); - assert_eq!(f[1], 0.0); - } - - #[test] - fn uv_interleave_round_trips_at_every_depth() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let max = depth.max_value() as u16; - let u_samples = vec![0, max / 4, max]; - let v_samples = vec![max, max / 2, 1]; - - let u_bytes = wire(&u_samples, depth); - let v_bytes = wire(&v_samples, depth); - - let packed = interleave_uv_to_f32(&u_bytes, &v_bytes, depth); - assert_eq!(packed.len(), 6, "packed UV length wrong at {depth:?}"); - - let (ru, rv) = unpack_uv_from_f32(&packed, 3, depth); - assert_eq!(ru, u_bytes, "U round trip failed at {depth:?}"); - assert_eq!(rv, v_bytes, "V round trip failed at {depth:?}"); - } - } - - #[test] - fn yuv_interleave_round_trips_at_every_depth() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let max = depth.max_value() as u16; - let y_samples = vec![0, max / 3, max]; - let u_samples = vec![max, 0, max / 2]; - let v_samples = vec![max / 4, max, 0]; - - let y_bytes = wire(&y_samples, depth); - let u_bytes = wire(&u_samples, depth); - let v_bytes = wire(&v_samples, depth); - - let packed = interleave_yuv_to_f32(&y_bytes, &u_bytes, &v_bytes, depth); - assert_eq!(packed.len(), 9, "packed YUV length wrong at {depth:?}"); - - let out = unpack_yuv_from_f32(&packed, 3, depth); - assert_eq!(out.y, y_bytes, "Y round trip failed at {depth:?}"); - assert_eq!(out.u, u_bytes, "U round trip failed at {depth:?}"); - assert_eq!(out.v, v_bytes, "V round trip failed at {depth:?}"); - } - } - - #[test] - fn split_uv_wire_matches_unpack_uv_from_f32() { - let u_src = [0.0, 0.25, 0.5, 1.0]; - let v_src = [1.0, 0.75, 0.5, 0.0]; - - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let packed: Vec = u_src.iter().zip(&v_src).flat_map(|(&u, &v)| [u, v]).collect(); - let (want_u, want_v) = unpack_uv_from_f32(&packed, 4, depth); - - // The kernel lays a chroma frame out as U's whole region - // followed by V's. - let wire: Vec = f32_to_plane(&u_src, depth) - .into_iter() - .chain(f32_to_plane(&v_src, depth)) - .collect(); - - let (u, v) = split_uv_wire(&wire); - assert_eq!(u, want_u, "U disagreed at {depth:?}"); - assert_eq!(v, want_v, "V disagreed at {depth:?}"); - } - } - - #[test] - fn split_yuv_wire_matches_unpack_yuv_from_f32() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let packed: Vec = (0..9).map(|i| i as f32 / 9.0).collect(); - - let want = unpack_yuv_from_f32(&packed, 3, depth); - let got = split_yuv_wire(&f32_to_plane(&packed, depth), depth); - - assert_eq!(got.y, want.y, "Y disagreed at {depth:?}"); - assert_eq!(got.u, want.u, "U disagreed at {depth:?}"); - assert_eq!(got.v, want.v, "V disagreed at {depth:?}"); - } - } - - #[test] - fn quantise_matches_the_clamping_form_including_nan() { - fn reference(v: f32, max: f32) -> u16 { - (v.clamp(0.0, 1.0) * max + 0.5) as u16 - } - - let max = 1023.0; - let cases = [ - -1.0, - -0.001, - 0.0, - 0.5, - 0.999, - 1.0, - 1.001, - 2.0, - f32::NAN, - f32::INFINITY, - f32::NEG_INFINITY, - ]; - - for v in cases { - assert_eq!(quantise(v, max), reference(v, max), "mismatch at {v}"); - } - } - - /// Limited-range codes normalise to matching values at every depth, - /// which is the property the whole design rests on. - /// - /// The match is within one 8-bit code level rather than exact. ITU - /// defines the limited-range endpoints as exact multiples, so 16 - /// becomes 64 and 235 becomes 940, but full scale is not a multiple, - /// because 255 becomes 1023. That leaves 235/255 and 940/1023 - /// differing by 0.0027, roughly 0.69 of an 8-bit step. - /// - /// Agreement below one step is the real property here. - /// - /// `normalized_scale_is_identical_across_depths` in - /// `src/nlmeans/mod.rs` pins the same property on the library's own - /// normalise helper. - #[test] - fn limited_range_codes_agree_across_depths() { - /// One 8-bit code level, the precision the endpoints agree to. - const TOL: f32 = 1.0 / 255.0; - - let eight = plane_to_f32(&wire(&[16, 235], Depth::Eight), Depth::Eight); - let ten = plane_to_f32(&wire(&[64, 940], Depth::Ten), Depth::Ten); - - for (a, b) in eight.iter().zip(ten.iter()) { - assert!((a - b).abs() < TOL, "8-bit {a} vs 10-bit {b}"); - } - } -} - -#[cfg(test)] -mod cli_options_tests { - use super::*; - use crate::nlmeans::NlmParams; - - /// A `PlaneOptions` with every field at a neutral default, so each test - /// only overrides what it cares about. - /// - /// `mode` and `algorithm` are the two fields every test below sets - /// for itself. - fn base_opts( - mode: DenoisingMode, - algorithm: Algorithm, - luma_strength: Option, - chroma_strength: Option, - ) -> PlaneOptions { - PlaneOptions { - accelerators: vec![], - device: Device::Default, - intent: ChannelIntent::LumaChroma, - mode, - algorithm, - luma_strength, - chroma_strength, - luma_lambda_ht: None, - chroma_lambda_ht: None, - } - } - - #[test] - fn luma_strength_alone_overrides_only_the_luma_plane() { - let opts = base_opts(DenoisingMode::Spacial, Algorithm::default(), Some(0.7), None); - - let luma = expect_nlmeans(opts.denoiser_options(ChannelMode::Luma, Depth::Eight).algorithm); - let chroma = expect_nlmeans(opts.denoiser_options(ChannelMode::Chroma, Depth::Eight).algorithm); - - assert!( - matches!(luma.tuning.strength, Some(s) if (s - 0.7).abs() < f32::EPSILON), - "expected luma tuning.strength = Some(0.7), got {:?}", - luma.tuning.strength - ); - assert_eq!( - chroma.tuning.strength, None, - "chroma plane should carry no override so the table default applies" - ); - } - - #[test] - fn both_per_plane_strengths_set_independently() { - let opts = base_opts(DenoisingMode::Spacial, Algorithm::default(), Some(0.7), Some(0.3)); - - let luma = expect_nlmeans(opts.denoiser_options(ChannelMode::Luma, Depth::Eight).algorithm); - let chroma = expect_nlmeans(opts.denoiser_options(ChannelMode::Chroma, Depth::Eight).algorithm); - - assert!( - matches!(luma.tuning.strength, Some(s) if (s - 0.7).abs() < f32::EPSILON), - "expected luma tuning.strength = Some(0.7), got {:?}", - luma.tuning.strength - ); - assert!( - matches!(chroma.tuning.strength, Some(s) if (s - 0.3).abs() < f32::EPSILON), - "expected chroma tuning.strength = Some(0.3), got {:?}", - chroma.tuning.strength - ); - } - - #[test] - fn no_overrides_hq_resolves_through_to_nlm_params_to_the_measured_tables() { - // Radius 4 in the measured tables is luma 0.35 and chroma - // 0.70 (see the table docs in `src/nlmeans/params.rs`). - let opts = base_opts( - DenoisingMode::Temporal { radius: 4 }, - Algorithm::NlmeansHq(NlmeansHqOptions::default()), - None, - None, - ); - - let luma_params: NlmParams = opts - .denoiser_options(ChannelMode::Luma, Depth::Eight) - .to_nlm_params(); - let chroma_params: NlmParams = opts - .denoiser_options(ChannelMode::Chroma, Depth::Eight) - .to_nlm_params(); - - assert!( - (luma_params.strength - 0.35).abs() < f32::EPSILON, - "expected luma strength 0.35 at r4, got {}", - luma_params.strength - ); - assert!( - (chroma_params.strength - 0.70).abs() < f32::EPSILON, - "expected chroma strength 0.70 at r4, got {}", - chroma_params.strength - ); - } - - /// A `PlaneOptions` running `Algorithm::Nl4d`, with every field at a - /// neutral default except the two per-plane `lambda_ht` overrides - /// under test. - fn nl4d_opts(luma_lambda_ht: Option, chroma_lambda_ht: Option) -> PlaneOptions { - PlaneOptions { - accelerators: vec![], - device: Device::Default, - intent: ChannelIntent::LumaChroma, - mode: DenoisingMode::Temporal { radius: 2 }, - algorithm: Algorithm::Nl4d(Nl4dOptions::default()), - luma_strength: None, - chroma_strength: None, - luma_lambda_ht, - chroma_lambda_ht, - } - } - - /// Unwraps an `Algorithm::Nlmeans`, panicking with the whole value - /// on any other variant. - fn expect_nlmeans(algorithm: Algorithm) -> NlmeansOptions { - match algorithm { - Algorithm::Nlmeans(n) => n, - other => panic!("expected Algorithm::Nlmeans, got {other:?}"), - } - } - - /// Unwraps an `Algorithm::Nl4d`, panicking with the whole value on - /// any other variant. - fn expect_nl4d(algorithm: Algorithm) -> Nl4dOptions { - match algorithm { - Algorithm::Nl4d(n) => n, - other => panic!("expected Algorithm::Nl4d, got {other:?}"), - } - } - - /// The routing property that matters most for a shared field: an - /// override aimed at one plane must never leak into the other - /// instance. `luma_lambda_ht` set alone must change nothing about - /// the chroma instance, and vice versa in the sibling test below. - #[test] - fn luma_lambda_ht_alone_overrides_only_the_luma_instance_for_nl4d() { - let opts = nl4d_opts(Some(4.0), None); - - let luma = expect_nl4d(opts.algorithm_for(ChannelMode::Luma)); - let chroma = expect_nl4d(opts.algorithm_for(ChannelMode::Chroma)); - - assert!((luma.lambda_ht.unwrap() - 4.0).abs() < f32::EPSILON); - assert_eq!( - chroma.lambda_ht, - Nl4dOptions::default().lambda_ht, - "chroma should stay unresolved here (None), deferred to its own per-plane \ - default at construction, got {:?}", - chroma.lambda_ht - ); - } - - #[test] - fn chroma_lambda_ht_alone_overrides_only_the_chroma_instance_for_nl4d() { - let opts = nl4d_opts(None, Some(4.0)); - - let luma = expect_nl4d(opts.algorithm_for(ChannelMode::Luma)); - let chroma = expect_nl4d(opts.algorithm_for(ChannelMode::Chroma)); - - assert_eq!( - luma.lambda_ht, - Nl4dOptions::default().lambda_ht, - "luma should stay unresolved here (None), deferred to its own per-plane \ - default at construction, got {:?}", - luma.lambda_ht - ); - assert!((chroma.lambda_ht.unwrap() - 4.0).abs() < f32::EPSILON); - } - - #[test] - fn both_planes_lambda_ht_set_independently_for_nl4d() { - let opts = nl4d_opts(Some(2.0), Some(3.5)); - - let luma = expect_nl4d(opts.algorithm_for(ChannelMode::Luma)); - let chroma = expect_nl4d(opts.algorithm_for(ChannelMode::Chroma)); - - assert!((luma.lambda_ht.unwrap() - 2.0).abs() < f32::EPSILON); - assert!((chroma.lambda_ht.unwrap() - 3.5).abs() < f32::EPSILON); - - // Every other field stays shared between the two instances even - // though lambda_ht diverges. - assert_eq!(luma.refine, chroma.refine); - assert_eq!(luma.spatial_radius, chroma.spatial_radius); - assert!((luma.c_min - chroma.c_min).abs() < f32::EPSILON); - } - - #[test] - fn unset_nl4d_overrides_resolve_to_different_lambda_ht_per_plane_end_to_end() { - let opts = nl4d_opts(None, None); - - let luma = expect_nl4d(opts.algorithm_for(ChannelMode::Luma)); - let chroma = expect_nl4d(opts.algorithm_for(ChannelMode::Chroma)); - - // Neither plane has anything set anywhere, so both stay - // unresolved at this layer... - assert_eq!(luma.lambda_ht, None); - assert_eq!(chroma.lambda_ht, None); - - // ...but resolving each through the same function construction - // uses (`nl4d_default_lambda_ht`, see `src/denoiser.rs`) gives - // luma and chroma different values, which is the whole point of - // a caller passing no flags at all getting both per-plane - // defaults. - let luma_default = crate::nl4d_default_lambda_ht(ChannelMode::Luma); - let chroma_default = crate::nl4d_default_lambda_ht(ChannelMode::Chroma); - assert!((luma_default - 4.158).abs() < f32::EPSILON); - assert!((chroma_default - 3.234).abs() < f32::EPSILON); - assert!((chroma_default - luma_default).abs() > f32::EPSILON); - } -} - -// Feature-gated because every test here builds its `PlaneOptions` from -// `chroma_only_opts`, which names the `Vulkan` accelerator variant. That -// variant only exists when the `vulkan` feature is enabled. -#[cfg(feature = "vulkan")] -#[cfg(test)] -mod passthrough_retry_tests { - use super::*; - use crate::accelerate::Accelerator; - use crate::{Algorithm, DenoisingMode}; - - /// Chroma-only intent, so `luma` is the disabled passthrough half and - /// `chroma` is the one that can report `QueueFull`. - /// - /// That is what drives the retry loop in `push_with_drain`. - fn chroma_only_opts() -> PlaneOptions { - PlaneOptions { - accelerators: vec![Accelerator::Vulkan], - device: Device::Default, - intent: ChannelIntent::Chroma, - mode: DenoisingMode::Spacial, - algorithm: Algorithm::default(), - luma_strength: None, - chroma_strength: None, - luma_lambda_ht: None, - chroma_lambda_ht: None, - } - } - - fn fake_planes(layout: FrameLayout) -> Planes { - Planes { - y: fill_plane(layout.luma_pixels(), layout.depth.neutral_chroma(), layout.depth), - u: layout.neutral_chroma_plane(), - v: layout.neutral_chroma_plane(), - } - } - - #[test] - fn queue_full_retry_does_not_double_queue_the_passthrough_plane() { - let layout = FrameLayout { - width: 16, - height: 16, - subsampling: Subsampling::Yuv420, - depth: Depth::Eight, - }; - let mut wd = - PlanarDenoiser::create(&chroma_only_opts(), layout).expect("denoiser construction failed"); - let planes = fake_planes(layout); - - // Spatial mode runs a depth-2 pipeline, so the first two pushes - // land directly. See `push_after_pending_returns_queue_full` in - // `src/denoiser.rs`. - wd.push(&planes).expect("first push should land"); - wd.push(&planes).expect("second push should land"); - - // Third push hits QueueFull on the chroma half. - let err = wd.push(&planes).expect_err("expected QueueFull"); - assert!( - matches!(err, DenoiserError::QueueFull), - "expected QueueFull, got {err:?}" - ); - - // Mirror the retry loop in `push_with_drain`. Drain one output, - // then retry the whole `push()` call for the same frame. - wd.recv().expect("recv after drain failed"); - wd.push(&planes).expect("retry push should land after drain"); - - // The chroma denoiser accepted three frames, two directly and - // one on the retry, and `recv` popped one back off. The disabled - // luma half's passthrough queue must track that one for one, and - // must not count the frame whose first attempt hit `QueueFull` - // twice. - assert_eq!( - wd.luma_passthrough.len(), - 2, - "expected exactly one passthrough entry per chroma frame actually accepted, got {}", - wd.luma_passthrough.len() - ); - } -} - -// Feature-gated because every test here builds its `PlaneOptions` from -// `luma_chroma_opts`, which names the `Vulkan` accelerator variant. That -// variant only exists when the `vulkan` feature is enabled. -#[cfg(feature = "vulkan")] -#[cfg(test)] -mod lumachroma_lockstep_tests { - use super::*; - use crate::accelerate::Accelerator; - use crate::{Algorithm, DenoisingMode}; - - /// Runs `luma` and `chroma` as two real `Denoiser`s in spatial mode. - /// - /// Spatial mode passes a uniform-valued plane through unchanged, as - /// the `uniform_*_passthrough` tests in `src/nlmeans/tests` show. The - /// test can therefore give each plane its own marker value and spot - /// the two halves drifting apart. - fn luma_chroma_opts() -> PlaneOptions { - PlaneOptions { - accelerators: vec![Accelerator::Vulkan], - device: Device::Default, - intent: ChannelIntent::LumaChroma, - mode: DenoisingMode::Spacial, - algorithm: Algorithm::default(), - luma_strength: None, - chroma_strength: None, - luma_lambda_ht: None, - chroma_lambda_ht: None, - } - } - - /// A uniform-valued frame whose luma and chroma planes each encode - /// `idx` with a different formula. - /// - /// If the round trip ever pairs luma from one push with chroma from - /// another, the two encodings disagree and the test catches it. - fn marked_planes(layout: FrameLayout, idx: u8) -> Planes { - let chroma_pixels = layout.chroma_pixels(); - let y_val = 10 + idx; - let uv_val = 200 - idx; - - Planes { - y: fill_plane(layout.luma_pixels(), y_val as u16, layout.depth), - u: fill_plane(chroma_pixels, uv_val as u16, layout.depth), - v: fill_plane(chroma_pixels, uv_val as u16, layout.depth), - } - } - - #[test] - fn queue_full_retries_never_desync_luma_and_chroma() { - let layout = FrameLayout { - width: 16, - height: 16, - subsampling: Subsampling::Yuv420, - depth: Depth::Eight, - }; - let mut wd = - PlanarDenoiser::create(&luma_chroma_opts(), layout).expect("denoiser construction failed"); - - // More pushes than the depth-2 pipeline holds, so this drives - // several `QueueFull`-then-retry cycles. - const N: u8 = 6; - let mut outputs: Vec = Vec::new(); - - for idx in 0..N { - let planes = marked_planes(layout, idx); - - // Mirror the retry loop in `push_with_drain` exactly, which - // is the sequence the CLI workers run. - if push_needs_retry(wd.push(&planes)).expect("push_needs_retry") { - if let Some(out) = wd.recv().expect("recv failed") { - outputs.push(out); - } - - wd.push(&planes).expect("retry push should land after drain"); - } - } - - wd.flush(|out| outputs.push(out)).expect("flush failed"); - - assert_eq!( - outputs.len(), - N as usize, - "expected exactly one output frame per input frame, got {}", - outputs.len() - ); - - for out in &outputs { - let y_val = out.y[0]; - let uv_val = out.u[0]; - let idx_from_y = y_val - 10; - let idx_from_uv = 200 - uv_val; - - assert_eq!( - idx_from_y, idx_from_uv, - "luma marker {y_val} (frame {idx_from_y}) and chroma marker {uv_val} \ - (frame {idx_from_uv}) disagree, so the luma and chroma pushes have drifted apart" - ); - } - } -} - -#[cfg(test)] -mod push_needs_retry_tests { - use super::*; - - #[test] - fn ok_means_no_retry() { - let outcome = push_needs_retry(Ok(())).expect("Ok(()) must not itself error"); - assert!(!outcome, "a landed push must not ask the caller to retry"); - } - - #[test] - fn queue_full_signals_retry() { - let outcome = - push_needs_retry(Err(DenoiserError::QueueFull)).expect("QueueFull must not itself error"); - assert!(outcome, "QueueFull must still trigger the retry-after-drain path"); - } - - #[test] - fn non_queue_full_errors_propagate_instead_of_being_swallowed() { - let synthetic = DenoiserError::Other(anyhow::anyhow!("synthetic readback failure")); - - let outcome = push_needs_retry(Err(synthetic)); - - assert!( - outcome.is_err(), - "a non-QueueFull push error must propagate instead of being silently treated as success" - ); - } -} - -#[cfg(test)] -mod layout_tests { - use super::*; - - fn layout(depth: Depth) -> FrameLayout { - FrameLayout { - width: 4, - height: 4, - subsampling: Subsampling::Yuv420, - depth, - } - } - - #[test] - fn byte_lengths_scale_with_depth() { - assert_eq!(layout(Depth::Eight).luma_bytes(), 16); - assert_eq!(layout(Depth::Ten).luma_bytes(), 32); - assert_eq!(layout(Depth::Eight).chroma_bytes(), 4); - assert_eq!(layout(Depth::Ten).chroma_bytes(), 8); - } - - #[test] - fn neutral_chroma_fill_is_correct_at_each_depth() { - let eight = layout(Depth::Eight).neutral_chroma_plane(); - assert_eq!(eight, vec![128u8; 4]); - - // 512 little-endian is [0x00, 0x02], repeated per sample. - let ten = layout(Depth::Ten).neutral_chroma_plane(); - assert_eq!(ten, vec![0x00, 0x02, 0x00, 0x02, 0x00, 0x02, 0x00, 0x02]); - - // 2048 little-endian is [0x00, 0x08]. - let twelve = layout(Depth::Twelve).neutral_chroma_plane(); - assert_eq!(twelve.len(), 8); - assert_eq!(&twelve[0..2], &[0x00, 0x08]); - } - - #[test] - fn black_luma_fill_is_zero_at_the_right_length() { - assert_eq!(layout(Depth::Eight).black_luma_plane(), vec![0u8; 16]); - assert_eq!(layout(Depth::Ten).black_luma_plane(), vec![0u8; 32]); - } -} - -#[cfg(test)] -mod chroma_dims_tests { - use super::*; - - #[test] - fn yuv420_even_dims_halve() { - assert_eq!(Subsampling::Yuv420.chroma_dims(1920, 1080), (960, 540)); - } - - #[test] - fn yuv420_odd_width_rounds_up() { - assert_eq!(Subsampling::Yuv420.chroma_dims(1919, 1080), (960, 540)); - } - - #[test] - fn yuv420_odd_height_rounds_up() { - assert_eq!(Subsampling::Yuv420.chroma_dims(1920, 1079), (960, 540)); - } - - #[test] - fn yuv420_odd_both_dims_round_up() { - assert_eq!(Subsampling::Yuv420.chroma_dims(1919, 1079), (960, 540)); - } - - #[test] - fn yuv422_even_width_halves() { - assert_eq!(Subsampling::Yuv422.chroma_dims(1920, 1080), (960, 1080)); - } - - #[test] - fn yuv422_odd_width_rounds_up() { - assert_eq!(Subsampling::Yuv422.chroma_dims(1919, 1080), (960, 1080)); - } - - #[test] - fn yuv444_passes_even_dims_through() { - assert_eq!(Subsampling::Yuv444.chroma_dims(1920, 1080), (1920, 1080)); - } - - #[test] - fn yuv444_passes_odd_dims_through() { - assert_eq!(Subsampling::Yuv444.chroma_dims(1919, 1079), (1919, 1079)); - } -} - -#[cfg(test)] -mod tests; diff --git a/av-denoise-core/src/frame/tests.rs b/av-denoise-core/src/frame/tests.rs deleted file mode 100644 index 8c8231b..0000000 --- a/av-denoise-core/src/frame/tests.rs +++ /dev/null @@ -1,1010 +0,0 @@ -//! Tests for [`super::PlanarDenoiser::reseed`]. - -use super::*; - -// Feature-gated because `test_plane_options` names the `Vulkan` -// accelerator variant, which only exists when the `vulkan` feature is -// enabled. -#[cfg(feature = "vulkan")] -mod reseed { - use super::*; - use crate::HqParams; - use crate::accelerate::Accelerator; - - fn layout() -> FrameLayout { - FrameLayout { - width: 64, - height: 64, - subsampling: Subsampling::Yuv420, - depth: Depth::Eight, - } - } - - /// A `PlaneOptions` running temporal nlmeans at radius `r`, denoising - /// both planes independently. - fn test_plane_options(r: u32) -> PlaneOptions { - PlaneOptions { - accelerators: vec![Accelerator::Vulkan], - device: Device::Default, - intent: ChannelIntent::LumaChroma, - mode: DenoisingMode::Temporal { radius: r }, - algorithm: Algorithm::default(), - luma_strength: None, - chroma_strength: None, - luma_lambda_ht: None, - chroma_lambda_ht: None, - } - } - - /// A `PlaneOptions` identical to [`test_plane_options`] except only - /// `intent` differs, for exercising a passthrough side. - fn test_plane_options_with_intent(r: u32, intent: ChannelIntent) -> PlaneOptions { - PlaneOptions { - intent, - ..test_plane_options(r) - } - } - - /// A small xorshift generator, deterministic across runs so the test - /// data does not vary between executions. - fn pseudo_random(mut x: u64) -> u64 { - x ^= x << 13; - x ^= x >> 7; - x ^= x << 17; - x - } - - /// One plane's bytes for frame `frame_idx`: a spatial ramp across the - /// plane, a per-frame offset, and a deterministic dither, all summed - /// and clamped so a temporal filter has real signal and real noise to - /// work with. - fn ramp_plane(pixels: usize, width: u32, frame_idx: usize, plane_seed: u64) -> Vec { - let width = width.max(1) as usize; - - (0..pixels) - .map(|i| { - let x = (i % width) as u32; - let y = (i / width) as u32; - let spatial = x.wrapping_add(y) % 120; - let frame_offset = (frame_idx as u32 * 7) % 60; - let seed = (i as u64) ^ (frame_idx as u64).wrapping_mul(0x9E3779B97F4A7C15) ^ plane_seed; - let dither = (pseudo_random(seed) % 16) as u32; - let value = 20 + spatial + frame_offset + dither; - value.min(235) as u8 - }) - .collect() - } - - /// Builds `count` `Planes` whose bytes vary per frame and per pixel, - /// so a temporal filter sees a non-degenerate signal. - fn ramp_clip(layout: &FrameLayout, count: usize) -> Vec { - let (chroma_w, _) = layout.chroma_dims(); - - (0..count) - .map(|frame_idx| Planes { - y: ramp_plane(layout.luma_pixels(), layout.width, frame_idx, 1), - u: ramp_plane(layout.chroma_pixels(), chroma_w, frame_idx, 2), - v: ramp_plane(layout.chroma_pixels(), chroma_w, frame_idx, 3), - }) - .collect() - } - - /// Renders `count` frames through the streaming path. - fn stream_all(opts: &PlaneOptions, frames: &[Planes]) -> Vec { - stream_all_with_layout(opts, layout(), frames) - } - - /// [`stream_all`] over a caller-chosen layout, for option sets that - /// need a source layout other than [`layout`], such as - /// `ChannelIntent::YuvFused`'s 4:4:4 requirement. - fn stream_all_with_layout( - opts: &PlaneOptions, - frame_layout: FrameLayout, - frames: &[Planes], - ) -> Vec { - let mut d = PlanarDenoiser::create(opts, frame_layout).unwrap(); - let mut out = Vec::new(); - for f in frames { - d.push(f).unwrap(); - if let Some(p) = d.recv().unwrap() { - out.push(p); - } - } - d.flush(|p| out.push(p)).unwrap(); - out - } - - fn window_of(frames: &[Planes], k: usize, r: usize) -> Vec { - (0..(2 * r + 1)) - .map(|i| { - let idx = (k + i).saturating_sub(r).min(frames.len() - 1); - frames[idx].clone() - }) - .collect() - } - - #[test] - fn reseed_matches_the_streaming_output_mid_clip() { - let opts = test_plane_options(2); - let frames = ramp_clip(&layout(), 12); - let streamed = stream_all(&opts, &frames); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let k = 6; - let got = d.reseed(&window_of(&frames, k, 2)).unwrap(); - - assert_eq!(got.y, streamed[k].y); - assert_eq!(got.u, streamed[k].u); - assert_eq!(got.v, streamed[k].v); - } - - #[test] - fn reseed_matches_the_streaming_output_at_both_clip_edges() { - let opts = test_plane_options(2); - let frames = ramp_clip(&layout(), 12); - let streamed = stream_all(&opts, &frames); - let last = frames.len() - 1; - - for k in [0usize, last] { - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let got = d.reseed(&window_of(&frames, k, 2)).unwrap(); - assert_eq!(got.y, streamed[k].y, "luma mismatch at k = {k}"); - assert_eq!(got.u, streamed[k].u, "u mismatch at k = {k}"); - assert_eq!(got.v, streamed[k].v, "v mismatch at k = {k}"); - } - } - - #[test] - fn reseed_recovers_a_half_poisoned_by_an_earlier_failure() { - let opts = test_plane_options(2); - let frames = ramp_clip(&layout(), 12); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - d.luma.as_mut().unwrap().poison_for_test(); - d.chroma.as_mut().unwrap().poison_for_test(); - - let got = d.reseed(&window_of(&frames, 6, 2)).unwrap(); - assert!(!got.y.is_empty()); - assert!(!got.u.is_empty()); - assert!(!got.v.is_empty()); - } - - /// Whether plain nlmeans's `reseed` stays order-independent under - /// repeated out-of-order reseeds on one long-lived denoiser, the - /// same stress the VapourSynth plugin's shuffled-access-order - /// harness test puts `avd.NLMeans` through. - /// - /// `Algorithm::Nlmeans`'s own doc comment says it runs with "no - /// noise measurement", and `NlmeansOptions` has no `hq` field for a - /// `PlaneOptions` built from it to carry, so `NlmParams::hq` stays - /// `None` and `fold_noise_estimate` never runs for it. There is no - /// stream-carried noise state for repeated `reseed` calls to - /// disagree about, so this is expected to hold without a - /// `windowed_noise_estimation` equivalent for nlmeans. This proves - /// that rather than assumes it, at a wider temporal radius and a - /// longer, more heavily shuffled clip than any other reseed test - /// here uses, so a history-dependent regression would have room to - /// show itself if one existed. - #[test] - fn nlmeans_repeated_out_of_order_reseeds_match_streaming() { - let opts = test_plane_options(4); - let frames = ramp_clip(&layout(), 24); - let streamed = stream_all(&opts, &frames); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - // Skews late, mirroring the plugin harness's shuffled order - // that first exposed the nl4d defect. - let order = [ - 18, 4, 23, 9, 12, 2, 20, 6, 15, 1, 22, 7, 17, 3, 11, 19, 0, 21, 8, 16, 5, 14, 10, 13, - ]; - - for &k in &order { - let got = d.reseed(&window_of(&frames, k, 4)).unwrap(); - assert_eq!(got.y, streamed[k].y, "luma mismatch at k = {k}"); - assert_eq!(got.u, streamed[k].u, "u mismatch at k = {k}"); - assert_eq!(got.v, streamed[k].v, "v mismatch at k = {k}"); - } - } - - #[test] - fn a_reseed_leaves_the_stream_positioned_for_the_next_frame() { - let opts = test_plane_options(2); - let frames = ramp_clip(&layout(), 12); - let streamed = stream_all(&opts, &frames); - let (k, r) = (6usize, 2usize); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - d.reseed(&window_of(&frames, k, r)).unwrap(); - d.push(&frames[k + 1 + r]).unwrap(); - let got = d.recv().unwrap().expect("frame k + 1"); - - assert_eq!(got.y, streamed[k + 1].y); - } - - #[test] - fn reseed_rejects_a_window_of_the_wrong_length() { - let opts = test_plane_options(2); - let frames = ramp_clip(&layout(), 12); - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - - let err = d.reseed(&frames[..3]).unwrap_err().to_string(); - assert!( - err.contains("5"), - "error should name the expected length, got {err}" - ); - } - - /// `ChannelIntent::Luma` leaves chroma disabled, so its planes travel - /// through the passthrough queue instead of a `Denoiser`. A reseed's - /// priming pushes queue one passthrough entry per window frame, and - /// this checks the entry `recv` pairs with the denoised centre is the - /// centre frame's own chroma, not a neighbour's. - #[test] - fn reseed_pairs_the_passthrough_plane_with_the_centre_frame() { - let opts = test_plane_options_with_intent(2, ChannelIntent::Luma); - let frames = ramp_clip(&layout(), 12); - let (k, r) = (6usize, 2usize); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let got = d.reseed(&window_of(&frames, k, r)).unwrap(); - - assert_eq!(got.u, frames[k].u, "u should pass through from the centre frame"); - assert_eq!(got.v, frames[k].v, "v should pass through from the centre frame"); - } - - /// The mirror of [`reseed_pairs_the_passthrough_plane_with_the_centre_frame`] - /// for `ChannelIntent::Chroma`, where luma is the disabled side. - #[test] - fn reseed_pairs_the_passthrough_luma_plane_with_the_centre_frame() { - let opts = test_plane_options_with_intent(2, ChannelIntent::Chroma); - let frames = ramp_clip(&layout(), 12); - let (k, r) = (6usize, 2usize); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let got = d.reseed(&window_of(&frames, k, r)).unwrap(); - - assert_eq!(got.y, frames[k].y, "y should pass through from the centre frame"); - } - - /// A single `reseed` call only checks the very first passthrough - /// entry `recv` pops. A leftover-count defect after the drop (an - /// extra or missing entry that still happens to leave the right one - /// at the front) would pass every single-shot test here and only - /// misalign the plane paired with the frame right after the centre, - /// once streaming resumes. - #[test] - fn reseed_then_streaming_keeps_the_passthrough_plane_aligned_on_the_next_frame() { - let opts = test_plane_options_with_intent(2, ChannelIntent::Luma); - let frames = ramp_clip(&layout(), 12); - let (k, r) = (6usize, 2usize); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - d.reseed(&window_of(&frames, k, r)).unwrap(); - d.push(&frames[k + 1 + r]).unwrap(); - let got = d.recv().unwrap().expect("frame k + 1"); - - assert_eq!(got.u, frames[k + 1].u, "u should pass through from frame k + 1"); - assert_eq!(got.v, frames[k + 1].v, "v should pass through from frame k + 1"); - } - - /// A `PlaneOptions` identical to [`test_plane_options`] except the - /// algorithm is `Nl4d`, which needs the wider window `reseed` has - /// to build for it instead of nlmeans's `2r+1` one. - /// - /// Pins `sigma` rather than leaving it on nl4d's automatic - /// per-frame estimate. That estimate is an exponential moving - /// average smoothed over every frame folded into it since the - /// stream last reset, so it carries genuine history from before - /// the window on a real, never-reset stream, history a windowed - /// `reseed` cannot supply and was never meant to reproduce. Pinning - /// it keeps these tests checking what `reseed`'s window shape and - /// pass sequence are actually responsible for, not that unrelated - /// warm-up behaviour. - fn nl4d_plane_options(r: u32) -> PlaneOptions { - PlaneOptions { - algorithm: Algorithm::Nl4d(Nl4dOptions { - sigma: Some(0.03), - ..Nl4dOptions::default() - }), - ..test_plane_options(r) - } - } - - /// The window a [`PlanarDenoiser::window_span`] of `span` needs for - /// target frame `k`, clamped at both clip ends exactly as - /// [`window_of`] clamps nlmeans's `2r+1` window. - /// - /// Reads the span from the accessor rather than hand-deriving it, - /// so this stays correct however the algorithm's own span is - /// shaped. - fn window_of_span(frames: &[Planes], k: usize, span: WindowSpan) -> Vec { - (0..span.frame_count()) - .map(|i| { - let idx = (k + i).saturating_sub(span.behind).min(frames.len() - 1); - frames[idx].clone() - }) - .collect() - } - - /// The shifted window around target frame `k`. It stops at the clip's - /// ends rather than repeating them, and returns the target's index in it. - fn shifted_window_of(frames: &[Planes], k: usize, span: WindowSpan) -> (Vec, ReseedWindowFlags) { - let first = k.saturating_sub(span.behind); - let last = (k + span.ahead).min(frames.len() - 1); - let window = frames[first..=last].to_vec(); - let flags = ReseedWindowFlags { - target: k - first, - at_clip_start: first == 0, - at_clip_end: last == frames.len() - 1, - }; - (window, flags) - } - - struct ReseedWindowFlags { - target: usize, - at_clip_start: bool, - at_clip_end: bool, - } - - fn reseed_shifted(denoiser: &mut PlanarDenoiser, frames: &[Planes], k: usize) -> Vec { - let span = denoiser.window_span(); - let (window, flags) = shifted_window_of(frames, k, span); - let request = ReseedWindow { - frames: &window, - target: flags.target, - at_clip_start: flags.at_clip_start, - at_clip_end: flags.at_clip_end, - }; - denoiser.reseed_window(request).unwrap() - } - - /// How many outputs a shifted reseed at `k` returns. That's the target - /// alone mid-clip, or the target through the clip's end. - fn got_len_for(denoiser: &PlanarDenoiser, clip_len: usize, k: usize) -> usize { - let span = denoiser.window_span(); - if k + span.ahead >= clip_len - 1 { - clip_len - k - } else { - 1 - } - } - - /// This is the test that would have caught the original defect: - /// `reseed` for a mid-clip frame under `Algorithm::Nl4d` must match - /// the streaming path's own output for that frame bit-for-bit, the - /// same property [`reseed_matches_the_streaming_output_mid_clip`] - /// checks for nlmeans. - #[test] - fn nl4d_reseed_matches_the_streaming_output_mid_clip() { - let opts = nl4d_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let streamed = stream_all(&opts, &frames); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let k = 8; - let span = d.window_span(); - let got = d.reseed(&window_of_span(&frames, k, span)).unwrap(); - - assert_eq!(got.y, streamed[k].y); - assert_eq!(got.u, streamed[k].u); - assert_eq!(got.v, streamed[k].v); - } - - /// The largest absolute per-sample difference between two same-sized - /// byte planes. - fn max_abs_diff(a: &[u8], b: &[u8]) -> i32 { - a.iter() - .zip(b.iter()) - .map(|(&x, &y)| (x as i32 - y as i32).abs()) - .max() - .unwrap_or(0) - } - - #[test] - fn nl4d_reseed_window_matches_streaming_at_every_frame() { - let option_sets = [ - nl4d_plane_options(2), - nl4d_windowed_plane_options(2), - nl4d_plane_options_with_intent(2, ChannelIntent::Luma), - nl4d_plane_options_with_intent(2, ChannelIntent::Chroma), - ]; - for opts in option_sets { - for clip_len in [3usize, 7, 12] { - let frames = ramp_clip(&layout(), clip_len); - let streamed = stream_all(&opts, &frames); - assert_eq!(streamed.len(), clip_len); - - for k in 0..clip_len { - let mut denoiser = PlanarDenoiser::create(&opts, layout()).unwrap(); - let got = reseed_shifted(&mut denoiser, &frames, k); - - let expected_len = got_len_for(&denoiser, clip_len, k); - assert_eq!(got.len(), expected_len, "len={clip_len} k={k}"); - - for (offset, planes) in got.iter().enumerate() { - let index = k + offset; - assert_eq!( - planes.y, streamed[index].y, - "len={clip_len} k={k} frame {index} luma" - ); - assert_eq!( - planes.u, streamed[index].u, - "len={clip_len} k={k} frame {index} u" - ); - assert_eq!( - planes.v, streamed[index].v, - "len={clip_len} k={k} frame {index} v" - ); - } - } - } - } - } - - /// [`ChannelIntent::YuvFused`] needs a 4:4:4 source, so this runs - /// the same match-streaming check as - /// [`nl4d_reseed_window_matches_streaming_at_every_frame`] over its - /// own 4:4:4 layout rather than sharing the 4:2:0 one. - #[test] - fn nl4d_reseed_window_matches_streaming_at_every_frame_in_yuv_fused_mode() { - let fused_layout = FrameLayout { - subsampling: Subsampling::Yuv444, - ..layout() - }; - let opts = nl4d_plane_options_with_intent(2, ChannelIntent::YuvFused); - for clip_len in [3usize, 7, 12] { - let frames = ramp_clip(&fused_layout, clip_len); - let streamed = stream_all_with_layout(&opts, fused_layout, &frames); - assert_eq!(streamed.len(), clip_len); - - for k in 0..clip_len { - let mut denoiser = PlanarDenoiser::create(&opts, fused_layout).unwrap(); - let got = reseed_shifted(&mut denoiser, &frames, k); - - let expected_len = got_len_for(&denoiser, clip_len, k); - assert_eq!(got.len(), expected_len, "len={clip_len} k={k}"); - - for (offset, planes) in got.iter().enumerate() { - let index = k + offset; - assert_eq!( - planes.y, streamed[index].y, - "len={clip_len} k={k} frame {index} luma" - ); - assert_eq!( - planes.u, streamed[index].u, - "len={clip_len} k={k} frame {index} u" - ); - assert_eq!( - planes.v, streamed[index].v, - "len={clip_len} k={k} frame {index} v" - ); - } - } - } - } - - #[test] - fn nl4d_reseed_window_pairs_passthrough_at_the_last_frame() { - let opts = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); - let frames = ramp_clip(&layout(), 16); - let last = frames.len() - 1; - let mut denoiser = PlanarDenoiser::create(&opts, layout()).unwrap(); - - let got = reseed_shifted(&mut denoiser, &frames, last); - - assert_eq!(got.len(), 1); - assert_eq!(got[0].u, frames[last].u); - assert_eq!(got[0].v, frames[last].v); - } - - #[test] - fn nl4d_reseed_window_pairs_passthrough_at_the_first_frame() { - let opts = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); - let frames = ramp_clip(&layout(), 16); - let mut denoiser = PlanarDenoiser::create(&opts, layout()).unwrap(); - - let got = reseed_shifted(&mut denoiser, &frames, 0); - - assert_eq!(got.len(), 1); - assert_eq!(got[0].u, frames[0].u); - assert_eq!(got[0].v, frames[0].v); - } - - #[test] - fn nl4d_reseed_window_then_streaming_continues_from_the_clip_start() { - let opts = nl4d_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let streamed = stream_all(&opts, &frames); - let mut denoiser = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = denoiser.window_span(); - - reseed_shifted(&mut denoiser, &frames, 1); - denoiser.push(&frames[2 + span.ahead]).unwrap(); - let next = denoiser.recv().unwrap().unwrap(); - - assert_eq!(next.y, streamed[2].y); - } - - /// After an nl4d `reseed`, ordinary sequential `push`/`recv` must - /// carry on producing the same frames the streaming path would - /// have, the nl4d mirror of - /// [`a_reseed_leaves_the_stream_positioned_for_the_next_frame`]. - #[test] - fn nl4d_reseed_then_streaming_continues_correctly() { - let opts = nl4d_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let streamed = stream_all(&opts, &frames); - let k = 8usize; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - d.reseed(&window_of_span(&frames, k, span)).unwrap(); - - // The next frame in source order after the reseed window's own - // last frame is `k + 1 + span.ahead`, the same relationship - // [`Denoise::render`]'s fast path in the VapourSynth plugin - // relies on for ordinary sequential continuation. - d.push(&frames[k + 1 + span.ahead]).unwrap(); - let got = d.recv().unwrap().expect("frame k + 1"); - - assert_eq!(got.y, streamed[k + 1].y); - assert_eq!(got.u, streamed[k + 1].u); - assert_eq!(got.v, streamed[k + 1].v); - } - - /// `reseed` rejects a window of the wrong length for nl4d too, and - /// the error names the wider length nl4d needs (`4r+1` at this - /// radius), not nlmeans's `2r+1`. - #[test] - fn nl4d_reseed_rejects_a_window_of_the_wrong_length() { - let opts = nl4d_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let expected = d.window_span().frame_count(); - - let err = d.reseed(&frames[..3]).unwrap_err().to_string(); - assert!( - err.contains(&expected.to_string()), - "error should name the expected length ({expected}), got {err}" - ); - } - - /// A `PlaneOptions` identical to [`nl4d_plane_options`] except only - /// `intent` differs, for exercising a passthrough side under nl4d's - /// wider window and multi-push pass sequence. - fn nl4d_plane_options_with_intent(r: u32, intent: ChannelIntent) -> PlaneOptions { - PlaneOptions { - intent, - ..nl4d_plane_options(r) - } - } - - /// The nl4d mirror of - /// [`reseed_pairs_the_passthrough_plane_with_the_centre_frame`]. - /// - /// nl4d's real-push loop drains after every emission, not only the - /// last, and each drain pops one passthrough entry. This checks - /// that walk still lands on the target's own entry rather than one - /// of the earlier, discarded regions' entries. - #[test] - fn nl4d_reseed_pairs_the_passthrough_plane_with_the_centre_frame() { - let opts = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); - let frames = ramp_clip(&layout(), 16); - let k = 8; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - let got = d.reseed(&window_of_span(&frames, k, span)).unwrap(); - - assert_eq!(got.u, frames[k].u, "u should pass through from the centre frame"); - assert_eq!(got.v, frames[k].v, "v should pass through from the centre frame"); - } - - /// The nl4d mirror of - /// [`reseed_pairs_the_passthrough_luma_plane_with_the_centre_frame`]. - #[test] - fn nl4d_reseed_pairs_the_passthrough_luma_plane_with_the_centre_frame() { - let opts = nl4d_plane_options_with_intent(2, ChannelIntent::Chroma); - let frames = ramp_clip(&layout(), 16); - let k = 8; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - let got = d.reseed(&window_of_span(&frames, k, span)).unwrap(); - - assert_eq!(got.y, frames[k].y, "y should pass through from the centre frame"); - } - - /// The nl4d mirror of - /// [`reseed_then_streaming_keeps_the_passthrough_plane_aligned_on_the_next_frame`]. - /// - /// A single-shot pairing test only checks the very first entry - /// `recv` pops after the drop. A leftover-count defect that still - /// happens to leave the right entry at the front would pass every - /// single-shot nl4d test above and only misalign the plane paired - /// with the frame right after the target, once streaming resumes. - #[test] - fn nl4d_reseed_then_streaming_keeps_the_passthrough_plane_aligned_on_the_next_frame() { - let opts = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); - let frames = ramp_clip(&layout(), 16); - let k = 8; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - d.reseed(&window_of_span(&frames, k, span)).unwrap(); - d.push(&frames[k + 1 + span.ahead]).unwrap(); - let got = d.recv().unwrap().expect("frame k + 1"); - - assert_eq!(got.u, frames[k + 1].u, "u should pass through from frame k + 1"); - assert_eq!(got.v, frames[k + 1].v, "v should pass through from frame k + 1"); - } - - /// A `PlaneOptions` identical to [`nl4d_plane_options`] except noise - /// estimation is window-local (`windowed_noise_estimation: true`) - /// and `sigma` is left on automatic estimation, the configuration - /// `av-denoise-vs` runs. - /// - /// Unlike [`nl4d_plane_options`], `sigma` is deliberately left - /// unpinned here: window-local estimation exists precisely so the - /// automatic estimate agrees between `reseed` and streaming, and a - /// pinned sigma would never have exercised that. - fn nl4d_windowed_plane_options(r: u32) -> PlaneOptions { - PlaneOptions { - algorithm: Algorithm::Nl4d(Nl4dOptions { - windowed_noise_estimation: true, - ..Nl4dOptions::default() - }), - ..test_plane_options(r) - } - } - - /// The property that would have caught the original random-access - /// bug directly: with window-local estimation on and `sigma` - /// automatic, a `reseed` for a mid-clip frame matches the streaming - /// path's own output for that frame bit-for-bit. The mirror of - /// [`nl4d_reseed_matches_the_streaming_output_mid_clip`], but with - /// the noise estimator actually exercised instead of sidestepped. - #[test] - fn nl4d_windowed_reseed_matches_the_streaming_output_mid_clip() { - let opts = nl4d_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let streamed = stream_all(&opts, &frames); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let k = 8; - let span = d.window_span(); - let got = d.reseed(&window_of_span(&frames, k, span)).unwrap(); - - assert_eq!(got.y, streamed[k].y); - assert_eq!(got.u, streamed[k].u); - assert_eq!(got.v, streamed[k].v); - } - - /// With window-local estimation on and `sigma` automatic, the fast - /// path and the reseed path must compute the same sigma for the - /// same window, so their outputs agree: a `reseed` at `k` followed - /// by an ordinary `push`/`recv` for `k + 1` must match a `reseed` - /// targeted directly at `k + 1` on a fresh denoiser. - /// - /// This is the property window-local estimation exists for. Without - /// it, the fast path keeps folding history the reseed path never - /// sees, so the two disagree even though both look at the same - /// window of real content. - #[test] - fn nl4d_windowed_fast_path_agrees_with_reseed_at_the_next_frame() { - let opts = nl4d_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let k = 8usize; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - d.reseed(&window_of_span(&frames, k, span)).unwrap(); - d.push(&frames[k + 1 + span.ahead]).unwrap(); - let via_fast_path = d.recv().unwrap().expect("frame k + 1"); - - let mut fresh = PlanarDenoiser::create(&opts, layout()).unwrap(); - let via_reseed = fresh.reseed(&window_of_span(&frames, k + 1, span)).unwrap(); - - assert_eq!(via_fast_path.y, via_reseed.y); - assert_eq!(via_fast_path.u, via_reseed.u); - assert_eq!(via_fast_path.v, via_reseed.v); - } - - /// A `PlaneOptions` identical to [`test_plane_options`] except the - /// algorithm is `NlmeansHq` with window-local estimation on and - /// `sigma` left on automatic estimation, the nlmeans mirror of - /// [`nl4d_windowed_plane_options`]. - fn nlmeans_hq_windowed_plane_options(r: u32) -> PlaneOptions { - PlaneOptions { - algorithm: Algorithm::NlmeansHq(NlmeansHqOptions { - nlm: NlmeansOptions::default(), - hq: HqParams { - windowed_noise_estimation: true, - ..HqParams::default() - }, - }), - ..test_plane_options(r) - } - } - - /// The nlmeans-hq mirror of - /// [`nl4d_windowed_reseed_matches_the_streaming_output_mid_clip`]. - /// - /// With window-local estimation on and `sigma` automatic, a `reseed` - /// for a mid-clip frame must match the streaming path's own output - /// for that frame bit-for-bit. No core test exercised HQ with - /// automatic sigma under reseed before this, which is how a - /// VapourSynth plugin filter that returns different pixels for the - /// same frame depending on request order shipped unnoticed. - #[test] - fn nlmeans_hq_windowed_reseed_matches_the_streaming_output_mid_clip() { - let opts = nlmeans_hq_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let streamed = stream_all(&opts, &frames); - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let k = 8; - let span = d.window_span(); - let got = d.reseed(&window_of_span(&frames, k, span)).unwrap(); - - assert_eq!(got.y, streamed[k].y); - assert_eq!(got.u, streamed[k].u); - assert_eq!(got.v, streamed[k].v); - } - - /// The clip-edge mirror of - /// [`nlmeans_hq_windowed_reseed_matches_the_streaming_output_mid_clip`]. - #[test] - fn nlmeans_hq_windowed_reseed_matches_the_streaming_output_at_both_clip_edges() { - let opts = nlmeans_hq_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let streamed = stream_all(&opts, &frames); - let last = frames.len() - 1; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - let got = d.reseed(&window_of_span(&frames, last, span)).unwrap(); - assert_eq!(got.y, streamed[last].y, "luma mismatch at the ahead edge"); - assert_eq!(got.u, streamed[last].u, "u mismatch at the ahead edge"); - assert_eq!(got.v, streamed[last].v, "v mismatch at the ahead edge"); - - const BEHIND_EDGE_TOLERANCE: i32 = 8; - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let got = d.reseed(&window_of_span(&frames, 0, span)).unwrap(); - let luma_diff = max_abs_diff(&got.y, &streamed[0].y); - assert!( - luma_diff <= BEHIND_EDGE_TOLERANCE, - "luma at the behind edge (k=0) drifted too far from streaming: max abs diff {luma_diff}" - ); - } - - /// The nlmeans-hq mirror of - /// [`nl4d_windowed_repeated_out_of_order_access_matches_streaming`]: - /// one long-lived `PlanarDenoiser` driven through the VapourSynth - /// plugin harness's exact shuffled access order with its hybrid - /// fast-path/`reseed` policy, every produced frame compared against - /// a true continuous stream. - /// - /// A single reseed, or a reseed followed by one push, both pass - /// under window-local estimation without exercising this, the same - /// way they did for nl4d: it takes a longer, repeatedly-reseeded run - /// to expose a carrier that survives `reset_stream_state` outside - /// the windowed gate. - #[test] - fn nlmeans_hq_windowed_repeated_out_of_order_access_matches_streaming() { - let opts = nlmeans_hq_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 14); - let streamed = stream_all(&opts, &frames); - let last = frames.len() - 1; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - let mut last_n: Option = None; - - // The VapourSynth plugin harness's exact shuffled order. - let order = [9usize, 0, 13, 4, 5, 6, 1, 12, 2, 11, 3, 10, 7, 8]; - const NEAR_START_TOLERANCE: i32 = 8; - - for &n in &order { - let fast = if last_n == Some(n.wrapping_sub(1)) && n > 0 { - let ahead = (n + span.ahead).min(last); - d.push(&frames[ahead]).unwrap(); - d.recv().unwrap() - } else { - None - }; - let got = match fast { - Some(out) => out, - None => d.reseed(&window_of_span(&frames, n, span)).unwrap(), - }; - last_n = Some(n); - - if n < span.behind { - let diff = max_abs_diff(&got.y, &streamed[n].y) - .max(max_abs_diff(&got.u, &streamed[n].u)) - .max(max_abs_diff(&got.v, &streamed[n].v)); - assert!( - diff <= NEAR_START_TOLERANCE, - "near-start frame n = {n} drifted too far from streaming: max abs diff {diff}" - ); - } else { - assert_eq!(got.y, streamed[n].y, "luma mismatch at n = {n}"); - assert_eq!(got.u, streamed[n].u, "u mismatch at n = {n}"); - assert_eq!(got.v, streamed[n].v, "v mismatch at n = {n}"); - } - } - } - - /// The nlmeans-hq mirror of - /// [`nl4d_windowed_fast_path_agrees_with_reseed_at_the_next_frame`]: - /// a `reseed` at `k` followed by an ordinary `push`/`recv` for - /// `k + 1` must match a `reseed` targeted directly at `k + 1` on a - /// fresh denoiser. - #[test] - fn nlmeans_hq_windowed_fast_path_agrees_with_reseed_at_the_next_frame() { - let opts = nlmeans_hq_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 16); - let k = 8usize; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - d.reseed(&window_of_span(&frames, k, span)).unwrap(); - d.push(&frames[k + 1 + span.ahead]).unwrap(); - let via_fast_path = d.recv().unwrap().expect("frame k + 1"); - - let mut fresh = PlanarDenoiser::create(&opts, layout()).unwrap(); - let via_reseed = fresh.reseed(&window_of_span(&frames, k + 1, span)).unwrap(); - - assert_eq!(via_fast_path.y, via_reseed.y); - assert_eq!(via_fast_path.u, via_reseed.u); - assert_eq!(via_fast_path.v, via_reseed.v); - } - - /// Reproduces the VapourSynth plugin harness's own `render` hybrid - /// policy exactly: frame 0 goes through `reseed`, and every - /// subsequent frame goes through the fast `push`/`recv` path, - /// falling back to `reseed` only when `recv` yields nothing. - fn render_sequence( - d: &mut PlanarDenoiser, - frames: &[Planes], - span: WindowSpan, - order: &[usize], - ) -> Vec { - let last = frames.len() - 1; - let mut last_n: Option = None; - let mut out = Vec::new(); - for &n in order { - let fast = if last_n == Some(n.wrapping_sub(1)) && n > 0 { - let ahead = (n + span.ahead).min(last); - d.push(&frames[ahead]).unwrap(); - d.recv().unwrap() - } else { - None - }; - let got = match fast { - Some(out) => out, - None => d.reseed(&window_of_span(frames, n, span)).unwrap(), - }; - last_n = Some(n); - out.push(got); - } - out - } - - /// Mirrors `a_sequential_run_after_a_seek_stays_correct_nlmeans` - /// exactly: a `reseed` at frame 11 of a 14-frame clip, then two - /// fast-path frames, compared against the same `render` hybrid - /// policy run straight through from frame 0. This is the reference - /// the VapourSynth harness actually uses, unlike - /// [`stream_all`], which is a true continuous stream with no - /// `reseed` in it at all. - #[test] - fn nlmeans_hq_windowed_sequential_run_after_a_seek_stays_correct() { - let opts = nlmeans_hq_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 14); - - let mut linear = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = linear.window_span(); - let linear_out = render_sequence(&mut linear, &frames, span, &(0..frames.len()).collect::>()); - - let mut seeked = PlanarDenoiser::create(&opts, layout()).unwrap(); - let seeked_out = render_sequence(&mut seeked, &frames, span, &[11, 12, 13]); - - for (i, n) in [12usize, 13].into_iter().enumerate() { - let got = &seeked_out[i + 1]; - let want = &linear_out[n]; - assert_eq!(got.y, want.y, "luma mismatch at n = {n}"); - assert_eq!(got.u, want.u, "u mismatch at n = {n}"); - assert_eq!(got.v, want.v, "v mismatch at n = {n}"); - } - } - - /// Diagnostic: same as - /// [`nlmeans_hq_windowed_sequential_run_after_a_seek_stays_correct`] - /// but at the VapourSynth harness's own clip size, 160x120. - #[test] - fn nlmeans_hq_windowed_sequential_run_after_a_seek_stays_correct_at_harness_size() { - let layout = FrameLayout { - width: 160, - height: 120, - subsampling: Subsampling::Yuv420, - depth: Depth::Eight, - }; - let opts = nlmeans_hq_windowed_plane_options(2); - let frames = ramp_clip(&layout, 14); - - let mut linear = PlanarDenoiser::create(&opts, layout).unwrap(); - let span = linear.window_span(); - let linear_out = render_sequence(&mut linear, &frames, span, &(0..frames.len()).collect::>()); - - let mut seeked = PlanarDenoiser::create(&opts, layout).unwrap(); - let seeked_out = render_sequence(&mut seeked, &frames, span, &[11, 12, 13]); - - for (i, n) in [12usize, 13].into_iter().enumerate() { - let got = &seeked_out[i + 1]; - let want = &linear_out[n]; - assert_eq!(got.y, want.y, "luma mismatch at n = {n}"); - assert_eq!(got.u, want.u, "u mismatch at n = {n}"); - assert_eq!(got.v, want.v, "v mismatch at n = {n}"); - } - } - - /// The property that actually reproduces the VapourSynth plugin - /// harness's `random_access_matches_sequential_access_nl4d` - /// end-to-end, at the core level: one long-lived `PlanarDenoiser` - /// driven through a shuffled access order with the plugin's own - /// hybrid fast-path/`reseed` policy, every produced frame compared - /// against a true continuous stream. - /// - /// A single reseed, or a reseed followed by one push, both pass - /// under window-local estimation without exercising this: it took - /// a longer, repeatedly-reseeded run to show that - /// `noise_estimator_temporal_only`'s "keep the last trustworthy - /// reading between folds" behaviour survives `reset_stream_state` - /// unwindowed even when every other chain is windowed, so a - /// `reseed`'s short real-push run can land on "no trustworthy - /// reading yet" while a true stream at the same frame is still - /// coasting on one from many frames back. - /// - /// Targets whose window covers either end of the clip reseed through - /// [PlanarDenoiser::reseed_window] with a shifted window. - #[test] - fn nl4d_windowed_repeated_out_of_order_access_matches_streaming() { - let opts = nl4d_windowed_plane_options(2); - let frames = ramp_clip(&layout(), 14); - let streamed = stream_all(&opts, &frames); - let last = frames.len() - 1; - - let mut d = PlanarDenoiser::create(&opts, layout()).unwrap(); - let span = d.window_span(); - let mut last_n: Option = None; - - // The VapourSynth plugin harness's exact shuffled order. - let order = [9usize, 0, 13, 4, 5, 6, 1, 12, 2, 11, 3, 10, 7, 8]; - - for &n in &order { - let fast = if last_n == Some(n.wrapping_sub(1)) && n > 0 { - let ahead = (n + span.ahead).min(last); - d.push(&frames[ahead]).unwrap(); - d.recv().unwrap() - } else { - None - }; - - let at_edge = n <= span.behind || n + span.ahead >= last; - let got = match fast { - Some(out) => out, - None if at_edge => { - let outputs = reseed_shifted(&mut d, &frames, n); - outputs[0].clone() - }, - None => d.reseed(&window_of_span(&frames, n, span)).unwrap(), - }; - last_n = Some(n); - - assert_eq!(got.y, streamed[n].y, "luma mismatch at n = {n}"); - assert_eq!(got.u, streamed[n].u, "u mismatch at n = {n}"); - assert_eq!(got.v, streamed[n].v, "v mismatch at n = {n}"); - } - } -} diff --git a/av-denoise-core/src/lib.rs b/av-denoise-core/src/lib.rs index 4d5d40d..861cee9 100644 --- a/av-denoise-core/src/lib.rs +++ b/av-denoise-core/src/lib.rs @@ -2,82 +2,46 @@ #![cfg_attr(docsrs, doc(auto_cfg))] #![doc = include_str!("../README.md")] -pub mod accelerate; -pub mod cache; -#[doc(hidden)] -pub mod collab; -mod denoiser; -pub mod device; -pub mod enumerate; -pub mod frame; #[doc(hidden)] +pub mod bench_api; +mod collab; +mod engine; +mod error; pub mod nl4d; -#[doc(hidden)] pub mod nlmeans; -mod probe; -pub mod sniff; -pub mod stack; -pub mod warmup; +mod options; -pub use cache::{ - COMPILATION_CACHE_ENV, - CacheError, - compilation_cache_dir, - default_cache_dir, - install_compilation_cache, - install_compilation_cache_at, - install_compilation_cache_once, -}; -pub use denoiser::{ - Algorithm, - Denoiser, - DenoiserError, - DenoiserOptions, - DenoisingMode, - EdgePadding, - FrameOutput, - MAX_PENDING, +pub use self::engine::{DevicePlane, EdgePadding, Engine, Geometry, SampleFormat, WindowSpan}; +pub use self::error::Error; +pub use self::nl4d::grain::{GrainChunk, SceneGrain, build_table}; +pub use self::nl4d::{ + Nl4d, Nl4dOptions, - NlmTuning, - NlmeansHqOptions, - NlmeansOptions, - NlmeansVariant, - OutputFormat, - Preset, - WindowSpan, nl4d_default_lambda_ht, nl4d_spatial_radius_for, nl4d_temporal_radius_for, - nlmeans_search_radius_for, - nlmeans_temporal_radius_for, - nlmeans_variant_for, }; -pub use device::Device; -pub use frame::{ - ChannelIntent, - FrameLayout, - PlanarDenoiser, - PlaneOptions, - Planes, - ReseedWindow, - Subsampling, - push_needs_retry, -}; -pub use nl4d::grain::{GrainChunk, SceneGrain, build_table}; -pub use nlmeans::{ +pub use self::nlmeans::{ ChannelMode, DEFAULT_PILOT_STRENGTH_SCALE, - Depth, + DenoisingMode, HqParams, MotionCompensationMode, MotionEstimation, MotionSearch, + NlmTuning, + Nlmeans, + NlmeansAlgorithm, + NlmeansHqOptions, + NlmeansOptions, + NlmeansVariant, PrefilterMode, - UnsupportedDepthError, - WirePack, - denormalize, - normalize, + nlmeans_search_radius_for, + nlmeans_temporal_radius_for, + nlmeans_variant_for, parse_prefilter, }; -pub use stack::{CODEGEN_STACK_BYTES, codegen_stack_is_sufficient, raise_codegen_stack_limit}; -pub use warmup::{WarmUp, kernel_key}; +pub use self::options::Preset; + +/// A hash of this crate's sources, which names a build's directory in the compiled kernel cache. +pub const KERNEL_HASH: &str = env!("AV_DENOISE_KERNEL_HASH"); diff --git a/av-denoise-core/src/nl4d/denoiser.rs b/av-denoise-core/src/nl4d/denoiser.rs index 7d836e6..5e44ea1 100644 --- a/av-denoise-core/src/nl4d/denoiser.rs +++ b/av-denoise-core/src/nl4d/denoiser.rs @@ -2,6 +2,7 @@ use cubecl::prelude::*; use cubecl::server::Handle; use super::grain::{GrainChunk, GrainExport, GrainGeometry}; +use super::nl4d_pool_ratio; use super::params::Nl4dParams; use super::regularise::run_regularise; use super::snapshot::{LastFields, MotionSnapshot, read_snapshot}; @@ -16,20 +17,18 @@ use crate::collab::kernels::aggregate::{ use crate::collab::kernels::fused::{STRENGTH_MAP_ALL, STRENGTH_MAP_LUMA, STRENGTH_MAP_OFF, collab_fused}; use crate::collab::kernels::transforms::dct_noise_profile; use crate::collab::{MAX_K, PATCH_SIZE, grid_frames, needs_warp_uniform_search}; -use crate::denoiser::{DenoiserError, FrameOutput, OutputFormat, nl4d_pool_ratio}; +use crate::engine::{DevicePlane, SampleFormat}; +use crate::nlmeans::denoiser::{BufferSize, front_buffer_sizes}; use crate::nlmeans::{ BLOCK_X, BLOCK_Y, ChannelMode, - Depth, MAX_GRID_1D, NOISE_CURVE_BINS, NlmDenoiser, - Pending, QuarterClasses, RingView, StrengthMapParams, - start_readback, }; /// Which accumulator regions a pass zeroes before scattering. @@ -43,6 +42,13 @@ enum AccumClear { Nothing, } +/// A finished accumulator region and the ring slot of the frame after it. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct CompletedRegion { + pub slot: u32, + pub next: Option, +} + /// How the current stream began. #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum StreamStart { @@ -52,29 +58,16 @@ enum StreamStart { Continuation, } -/// Groups similar 8x8 patches across a motion-compensated temporal -/// window and denoises each group jointly. +/// Groups similar 8x8 patches across a motion-compensated window and denoises each group jointly. /// -/// This drives the NLMeans front end, but only for its machinery, the -/// frame ring, the motion field, and the confidence scores built by -/// [`NlmDenoiser::submit_machinery`]. -/// No NLM weighting kernel ever runs. Instead, every submit hands the -/// noisy ring to [`collab_fused`], which groups patches by searching -/// both the centre frame spatially and each neighbour frame around -/// where motion compensation predicts a patch moved, shrinks each -/// group's coefficients in the transform domain, and scatters the -/// filtered members back into the accumulator ring. -/// [`collab_normalise`] then turns one region of that ring into a -/// finished frame. +/// The NLMeans front end supplies the frame ring, motion field and confidence scores, and no NLM +/// weighting runs. Each pass groups patches by searching the centre frame around each reference +/// and each neighbour frame around where motion predicts the patch moved. It shrinks each group's +/// coefficients in the transform domain and scatters the filtered members back into an +/// accumulator ring, which `collab_normalise` turns into finished frames. /// -/// Every pass scatters its filtered members into whichever frame each one -/// came from, not only the centre frame, so a frame's own output finishes -/// only once every pass that can reach it has run. -/// -/// Latency is `2 * temporal_radius` pushes, twice the front end's own -/// window depth. [`Self::denoise_submit`] returns `None` while the front -/// end's window is still filling, and [`Self::flush`] drains the frames -/// still held once the input stream ends. +/// A pass scatters each member into the frame it came from, so a frame finishes only once every +/// pass that can reach it has run. Latency is `2 * temporal_radius` pushes. pub struct Nl4dDenoiser { front: NlmDenoiser, width: u32, @@ -88,115 +81,90 @@ pub struct Nl4dDenoiser { lambda_ht: f32, c_min: f32, k_max: u32, - /// Whether [`collab_fused`] runs its warp-uniform search, decided - /// once from the runtime this denoiser was built on. See - /// [`needs_warp_uniform_search`]. + /// Whether `collab_fused` runs its warp-uniform search. warp_uniform: bool, - /// The fixed-point scale the cross-frame accumulator ring counts in, - /// from - /// [`crate::collab::kernels::aggregate::cross_frame_accum_scale`]. - /// - /// Both radii it derives from are fixed for the denoiser's lifetime, - /// so this is worked out once here rather than on every - /// [`Self::run_pass`] call. + /// The fixed-point scale the cross-frame accumulator ring counts in. accum_scale: f32, group_weight: Handle, sigma_buf: Handle, dct_profile_buf: Handle, - /// The aggregation window's 8 taps, built once from the caller's - /// `kaiser_beta`. Eight ones when that is 0. + /// The aggregation window's 8 taps, all ones when `kaiser_beta` is 0. kaiser_buf: Handle, - /// The correlation profile kept on the host too, so the weight - /// normalisation can be derived from it every submit without a - /// device readback. + /// The correlation profile kept on the host, so the weight normalisation needs no readback. dct_profile: [f32; 8], - /// Fixed-point accumulators the filter scatters into, one weighted - /// value per covering patch. + /// Fixed-point accumulators the filter scatters into, one region per slot of the frame ring. /// - /// These hold `1 + 2 * temporal_radius` frames' worth of pixels back - /// to back, one region per physical ring slot of the front end's own - /// frame ring. - /// - /// A pass contributes to every frame in the ring, so a frame's region - /// stays live across every pass run while it sits in the ring. See - /// [`Self::denoise_submit`] for when a region is read back. + /// A pass contributes to every frame in the ring, so a frame's region stays live for as long + /// as the frame sits in the ring. accum: Handle, wsum: Handle, - /// Two output buffers, alternated so one frame's kernels can overlap - /// the previous frame's readback. + /// Two output buffers, alternated so one frame's kernels can overlap the previous frame's + /// readback. outputs: [Handle; 2], next_output_slot: usize, - /// The format every readback this denoiser starts comes back in. - output_format: OutputFormat, - /// Packed-word destinations, one per entry of `outputs`, allocated - /// only in wire mode. - /// - /// These buffers rotate on the same slot counter, so each is free again exactly - /// when the `f32` slot it is packed from is free. - wire_outputs: Option<[Handle; 2]>, - /// How many passes [`Self::run_pass`] has run for the current stream. + /// Passes run for the current stream. /// - /// The stream's first pass zeroes the whole of `accum`/`wsum`, so a - /// zero here also means the ring may hold a previous stream's stale - /// contributions. It resets to 0 in [`Self::reset_stream`]. + /// At zero the accumulator ring may hold a previous stream's stale contributions, so the + /// stream's first pass zeroes all of it. passes_run: u32, - /// How the current stream began, set by [`Self::mark_continuation`]. stream_start: StreamStart, - /// The field buffers the last pass handed the fused kernel, for - /// [`Self::motion_snapshot`]. + /// The field buffers the last pass handed the fused kernel. last_fields: Option, - /// See [`Nl4dParams::field_lambda`]. field_lambda: f32, - /// The regularised motion field and confidence, laid out like the - /// front end's own, allocated only when `field_lambda > 0.0`. + /// The regularised motion field, allocated only when `field_lambda > 0.0`. reg_mv: Option, reg_conf: Option, - /// Whether passes apply the luma noise curve. - /// - /// Set when [Nl4dParams::noise_map](crate::nl4d::Nl4dParams::noise_map) - /// is on and the denoiser filters luma. + /// Whether passes apply the luma noise curve, set when the noise map is on and the denoiser + /// filters luma. apply_noise_map: bool, - /// The luma map's multipliers, set when `apply_noise_map` is and they are not all 1.0. + /// Set when the noise map applies and its multipliers are not all 1.0. luma_map: Option, - /// The flat boost a chroma denoiser applies, set when the noise map is on and the boost is not - /// 1.0. + /// The chroma flat boost, set when the noise map is on and the boost is not 1.0. chroma_map_boost: Option, /// A strength map of all 1.0, bound by every pass with no map to apply. unit_map: Handle, map_cols: u32, map_rows: u32, - /// The grain measurement state, present only when grain export is on. + /// Present only when grain export is on. grain: Option, } -impl Nl4dDenoiser { - /// Builds a new denoiser. - /// - /// Rejects an invalid `params` (see [`Nl4dParams::validate`]) and a - /// frame smaller than one collaborative patch on either axis. - pub fn new( - client: &ComputeClient, - params: Nl4dParams, - width: u32, - height: u32, - ) -> Result { - Self::with_output_format(client, params, width, height, OutputFormat::F32) +/// The element count of every geometry-sized buffer [Nl4dDenoiser::new] allocates beyond its frame rings. +/// +/// The front end is sized at the grouping radius, as `new` sets it. `accum` holds as many elements +/// as the frame ring, `wsum`, `outputs` and `group_weight` hold fewer, and `reg_mv` and `reg_conf` +/// match the front end's motion field and confidence. The grain's saved vectors span one more frame +/// than the motion field, so they are counted here. +pub(crate) fn buffer_sizes( + client: &ComputeClient, + params: &Nl4dParams, + width: u32, + height: u32, +) -> Vec { + let mut nlm_params = params.nlm.clone(); + nlm_params.temporal_radius = params.temporal_radius; + + let front = front_buffer_sizes(client, &nlm_params, width, height); + let mut buffers = front.buffers; + + let exports_grain = params.grain_export && nlm_params.channels != ChannelMode::Chroma; + if let (true, Some(blocks)) = (exports_grain, front.motion_blocks) { + let ring_frames = 1 + 2 * u64::from(params.temporal_radius); + let saved_mv = blocks.checked_mul(2 * ring_frames); + buffers.push(("grain saved vectors", saved_mv)); } - /// Builds a new denoiser whose readbacks come back in `output_format`. - /// - /// [`OutputFormat::Wire`] gives the denoiser a packed-word buffer - /// per output slot, so a readback quantises on the GPU and only the - /// wire bytes cross the bus. - /// - /// Rejects the same `params` and dimensions [`Self::new`] does. - pub fn with_output_format( + buffers +} + +impl Nl4dDenoiser { + /// Builds the denoiser, rejecting invalid `params` and frames smaller than one patch. + pub fn new( client: &ComputeClient, mut params: Nl4dParams, width: u32, height: u32, - output_format: OutputFormat, ) -> Result { params.validate()?; @@ -208,19 +176,13 @@ impl Nl4dDenoiser { )); } - // The front end's own temporal radius has to match the grouping - // radius exactly, since `submit_machinery`'s ring view walks the - // front end's own window, so this is forced here rather than - // trusted to the caller. + // The ring view `submit_machinery` returns walks the front end's own window, so its radius + // must match the grouping radius. params.nlm.temporal_radius = params.temporal_radius; - params.nlm.validate().map_err(|e| e.to_string())?; + params.nlm.validate().map_err(|error| error.to_string())?; - // The front end only supplies the ring, motion field, and - // confidence scores the collaborative stage reads. Its own - // buffers never leave the GPU, so it stays in `f32` whatever - // format this denoiser hands back. - let mut front = - NlmDenoiser::with_output_format(client, params.nlm.clone(), width, height, OutputFormat::F32); + let nlm_params = params.nlm.clone(); + let mut front = NlmDenoiser::new(client, nlm_params, width, height); let channels = params.nlm.channels; let apply_noise_map = params.noise_map && channels != ChannelMode::Chroma; @@ -245,26 +207,25 @@ impl Nl4dDenoiser { let frame_len = pixels * stored_ch as usize; let group_weight = client.empty(refs * size_of::()); - let sigma_buf = client.create_from_slice(f32::as_bytes(&vec![0.0f32; stored_ch as usize])); - // The correlation profile is purely spatial and this denoiser - // exposes no `rho` knob, so it is built once here from the - // white-noise default rather than every submit. + let sigma_host = vec![0.0f32; stored_ch as usize]; + let sigma_bytes = f32::as_bytes(&sigma_host); + let sigma_buf = client.create_from_slice(sigma_bytes); + + // The correlation profile is purely spatial and has no `rho` knob, so it is built once from + // the white-noise default. let dct_profile = dct_noise_profile(0.0); - let dct_profile_buf = client.create_from_slice(f32::as_bytes(&dct_profile)); - // The window depends only on `kaiser_beta`, which cannot change - // over a denoiser's life, so it is built here rather than every - // submit. - let kaiser_buf = client.create_from_slice(f32::as_bytes(&kaiser_window(params.kaiser_beta))); + let dct_profile_bytes = f32::as_bytes(&dct_profile); + let dct_profile_buf = client.create_from_slice(dct_profile_bytes); + let kaiser_taps = kaiser_window(params.kaiser_beta); + let kaiser_bytes = f32::as_bytes(&kaiser_taps); + let kaiser_buf = client.create_from_slice(kaiser_bytes); let (map_cols, map_rows) = strength_map_dims(width, height); let unit_map_host = vec![1.0f32; (map_cols * map_rows) as usize]; - let unit_map = client.create_from_slice(f32::as_bytes(&unit_map_host)); + let unit_map_bytes = f32::as_bytes(&unit_map_host); + let unit_map = client.create_from_slice(unit_map_bytes); - // One region per physical ring slot of the front end's own frame - // ring, `1 + 2 * temporal_radius` of them, see the `accum` field - // doc for why. The ring is zeroed in full by a stream's first pass - // rather than here, since `client.empty` gives no guarantee its memory - // starts zeroed. + // `client.empty` memory is not zeroed, so a stream's first pass zeroes the whole ring. let ring_frames = 1 + 2 * params.temporal_radius; let accum = client.empty(frame_len * ring_frames as usize * size_of::()); let wsum = client.empty(pixels * ring_frames as usize * size_of::()); @@ -272,41 +233,31 @@ impl Nl4dDenoiser { client.empty(frame_len * size_of::()), client.empty(frame_len * size_of::()), ]; - let wire_outputs = match output_format { - OutputFormat::F32 => None, - OutputFormat::Wire { depth } => { - let samples = pixels as u32 * channels.count(); - let words = samples.div_ceil(depth.wire_pack().samples_per_word()) as usize; - Some([ - client.empty(words * size_of::()), - client.empty(words * size_of::()), - ]) - }, - }; - // `motion_ctx()` panics without motion compensation, and - // `validate` above already requires it, so this is safe here. + // `motion_ctx()` panics without motion compensation, which `validate` already requires. let (reg_mv, reg_conf) = if params.field_lambda > 0.0 { - let mc = front.motion_ctx(); + let motion_ctx = front.motion_ctx(); let neighbours = 2 * params.temporal_radius as u64; - ( - Some(client.empty((neighbours * mc.mv_field_bytes_per_neighbour()) as usize)), - Some(client.empty((neighbours * mc.confidence_bytes_per_neighbour()) as usize)), - ) + let mv_bytes = neighbours * motion_ctx.mv_field_bytes_per_neighbour(); + let conf_bytes = neighbours * motion_ctx.confidence_bytes_per_neighbour(); + let reg_mv = client.empty(mv_bytes as usize); + let reg_conf = client.empty(conf_bytes as usize); + + (Some(reg_mv), Some(reg_conf)) } else { (None, None) }; let exports_grain = params.grain_export && channels != ChannelMode::Chroma; let grain = if exports_grain { - let mc = front.motion_ctx(); + let motion_ctx = front.motion_ctx(); let geometry = GrainGeometry { width, height, stored_ch, - blocks_x: mc.blocks_x, - blocks_y: mc.blocks_y, - step: mc.step, + blocks_x: motion_ctx.blocks_x, + blocks_y: motion_ctx.blocks_y, + step: motion_ctx.step, ring_frames, }; Some(GrainExport::new(client, geometry)) @@ -314,6 +265,9 @@ impl Nl4dDenoiser { None }; + let warp_uniform = needs_warp_uniform_search(client); + let accum_scale = cross_frame_accum_scale(params.spatial_radius, params.temporal_radius); + Ok(Self { front, width, @@ -326,8 +280,8 @@ impl Nl4dDenoiser { lambda_ht: params.lambda_ht, c_min: params.c_min, k_max, - warp_uniform: needs_warp_uniform_search(client), - accum_scale: cross_frame_accum_scale(params.spatial_radius, params.temporal_radius), + warp_uniform, + accum_scale, group_weight, sigma_buf, dct_profile_buf, @@ -337,8 +291,6 @@ impl Nl4dDenoiser { wsum, outputs, next_output_slot: 0, - output_format, - wire_outputs, passes_run: 0, stream_start: StreamStart::SceneStart, last_fields: None, @@ -355,49 +307,16 @@ impl Nl4dDenoiser { }) } - /// Pushes a new frame into the front end's ring buffer. + /// Runs the grouping passes a submit owes and returns the region they completed. /// - /// `frame` holds `width * height * channels` `f32` values in - /// `[0, 1]`, matching [`NlmDenoiser::push_frame`]. - pub fn push_frame(&mut self, frame: &[f32]) { - self.front.push_frame(frame); - } - - /// Pushes a new frame held as wire bytes into the front end's ring - /// buffer. - /// - /// `planes` holds one `width * height` plane per channel at `depth`, - /// matching [`NlmDenoiser::push_frame_wire`]. - pub fn push_frame_wire(&mut self, planes: &[&[u8]], depth: Depth) { - self.front.push_frame_wire(planes, depth); - } - - /// Runs one submit's worth of grouping, filtering, and aggregation, - /// and starts the readback. - /// - /// Returns `Ok(None)` while the front end's ring is still filling. + /// The push that fills a scene's ring runs head passes centred on the ring's first + /// `temporal_radius` frames, then the pass centred on its middle frame. Every later submit runs + /// one pass centred on the middle frame. /// - /// The push that fills a scene's ring runs head passes centred on the - /// ring's first `temporal_radius` frames, then the pass centred on its - /// middle frame. Every later submit runs one pass centred on the - /// ring's middle frame. - /// - /// Once more than `temporal_radius` passes have run, each pass centred - /// on the ring's middle completes the region `temporal_radius` frames - /// behind it, and that region is read back. A continuation stream - /// skips the head passes, so its first `temporal_radius` submits after - /// the ring fills return `Ok(None)`. Latency stays - /// `2 * temporal_radius` pushes for a scene start. - /// - /// There are two output slots, so at most two [`Pending`]s from this - /// denoiser may be outstanding at once. A third concurrent submit - /// reuses the oldest one's slot and silently corrupts it. - /// - /// The frame comes back in the [`OutputFormat`] this denoiser was - /// built with. [`OutputFormat::Wire`] quantises and packs the frame - /// on the GPU before the readback, so only the wire bytes cross the - /// bus. - pub fn denoise_submit(&mut self) -> Result>, DenoiserError> { + /// Once more than `temporal_radius` passes have run, each middle pass completes the region + /// `temporal_radius` frames behind it. A continuation stream skips the head passes, so its + /// first `temporal_radius` submits after the ring fills return `None`. + pub(crate) fn submit_passes(&mut self) -> Result, anyhow::Error> { if !self.front.window_ready() { return Ok(None); } @@ -428,14 +347,39 @@ impl Nl4dDenoiser { } let total_frames = 1 + 2 * radius; - let completed_slot = (view.centre_slot + total_frames - radius) % total_frames; - let (handle, slot) = self.normalise_region(completed_slot); - let next_slot = self.front.ring_slot(1); - self.measure_grain(completed_slot, Some(next_slot), slot); - let wire_dst = self.wire_outputs.as_ref().map(|outputs| &outputs[slot]); - let pending = self.start_readback(handle, wire_dst, self.output_format); - - Ok(Some(pending)) + let slot = (view.centre_slot + total_frames - radius) % total_frames; + let next = Some(self.front.ring_slot(1)); + + Ok(Some(CompletedRegion { slot, next })) + } + + /// Normalises a finished region into the next output slot and returns that slot's buffer. + pub(crate) fn read_region(&mut self, region: CompletedRegion) -> Handle { + let (handle, output_slot) = self.normalise_region(region.slot); + self.measure_grain(region.slot, region.next, output_slot); + + handle + } + + pub(crate) fn push_planes( + &mut self, + planes: &[DevicePlane<'_>], + format: SampleFormat, + ) -> Result<(), anyhow::Error> { + self.front.push_planes(planes, format) + } + + /// A 4-byte handle to bind for planes a kernel never reads. + pub(crate) fn placeholder(&self) -> &Handle { + self.front.placeholder() + } + + pub(crate) fn compute_client(&self) -> &ComputeClient { + self.front.compute_client() + } + + pub(crate) fn frame_shape(&self) -> (u32, u32, ChannelMode) { + (self.width, self.height, self.channels) } /// Marks the current stream as picking up mid-clip, so it runs no head passes. @@ -447,51 +391,31 @@ impl Nl4dDenoiser { } } - /// Runs the front end's motion and noise machinery for a pass centred on logical ring position `centre`. - fn machinery_at(&mut self, centre: u32) -> Result { + /// Runs the front end's motion and noise machinery for a pass centred on logical position `centre`. + fn machinery_at(&mut self, centre: u32) -> Result { let view = self.front.submit_machinery(centre)?; let view = view.expect("the ring is full whenever a pass runs"); - Ok(view) - } - /// Submits and waits for the result in one call. - /// - /// Prefer [`Self::denoise_submit`] when the caller can hold a frame - /// in flight. - /// - /// The frame comes back in the [`OutputFormat`] this denoiser was - /// built with. - pub fn denoise(&mut self) -> Result, DenoiserError> { - let Some(pending) = self.denoise_submit()? else { - return Ok(None); - }; - Ok(Some(pending.wait()?)) + Ok(view) } - /// Produces the frames still held at the end of a stream. - /// - /// A stream that filled its ring runs off-centre passes centred on its - /// last `temporal_radius` frames, then reads out the last - /// `2 * temporal_radius` frames' regions. A stream too short to fill - /// its ring pads the ring with copies of its last frame, runs every - /// real frame as a centre, then reads out every real frame. + /// Runs the passes a stream's end owes and returns the regions to read out, in emit order. /// - /// `sink` is called once per frame, in order, and the frame it - /// receives is only valid for that call. It arrives in the - /// [`OutputFormat`] this denoiser was built with, quantised by the - /// same pack kernel as every streaming frame. - pub fn flush(&mut self, mut sink: impl FnMut(&FrameOutput)) -> Result<(), DenoiserError> { + /// A stream that filled its ring runs tail passes centred on its last `temporal_radius` + /// frames, then reads out the last `2 * temporal_radius` regions. A shorter stream pads the + /// ring with copies of its last frame, runs every real frame as a centre and reads out every + /// real frame. + pub(crate) fn finish_passes(&mut self) -> Result, anyhow::Error> { let emit = self.flush_target() as u32; if emit == 0 { - self.reset_stream(); - return Ok(()); + return Ok(Vec::new()); } let radius = self.temporal_radius; let total_frames = 1 + 2 * radius; let short_stream = self.front.real_pushes() < total_frames as usize; if short_stream { - self.front.fill_ring_with_last_frame(); + self.front.fill_ring_with_last_frame()?; } let (centres, last_real) = if short_stream { @@ -500,11 +424,8 @@ impl Nl4dDenoiser { (radius + 1..2 * radius + 1, 2 * radius) }; - // A full ring with no pass run means a caller primed every slot - // through pushes alone and never called `denoise_submit`, so - // `accum`/`wsum` are still whatever the last stream, or nothing - // at all, left in them. The tail path's first pass has to clear - // the whole ring in that case, the same as a short stream does. + // A full ring with no pass run was primed through pushes alone, so its accumulators hold + // stale data and the first tail pass clears the whole ring. let mut clear = if short_stream || self.passes_run == 0 { AccumClear::WholeRing } else { @@ -517,34 +438,22 @@ impl Nl4dDenoiser { clear = AccumClear::Nothing; } - // Every output slot is free here. A caller reaches a flush only - // once its streaming readbacks have landed, and each readback - // below blocks before the next region reuses a slot. let first_region = last_real + 1 - emit; + let mut regions = Vec::with_capacity(emit as usize); for logical in first_region..=last_real { - let region_slot = self.front.ring_slot(logical); - let (handle, slot) = self.normalise_region(region_slot); - let next_slot = (logical < last_real).then(|| self.front.ring_slot(logical + 1)); - self.measure_grain(region_slot, next_slot, slot); - let wire_dst = self.wire_outputs.as_ref().map(|outputs| &outputs[slot]); - let pending = self.start_readback(handle, wire_dst, self.output_format); - let frame = pending.wait()?; - sink(&frame); + let slot = self.front.ring_slot(logical); + let next = (logical < last_real).then(|| self.front.ring_slot(logical + 1)); + regions.push(CompletedRegion { slot, next }); } - self.reset_stream(); - - Ok(()) + Ok(regions) } - /// Drops the current stream and returns to the state a fresh - /// denoiser starts in, keeping every GPU allocation. + /// Drops the current stream and returns to a fresh denoiser's state, keeping every allocation. /// - /// Clears the front end's own stream state plus the cross-frame - /// accumulator's pass counter and output slot, so a window primed - /// after this call never reads a previous window's stale - /// contributions out of the fixed-point `accum`/`wsum` ring. The next - /// stream starts a scene unless it is marked a continuation. + /// The pass counter resets, so the next stream's first pass clears the previous stream's + /// contributions from the accumulator ring. The next stream starts a scene unless it is marked + /// a continuation. pub fn reset_stream(&mut self) { self.front.reset_stream_state(); self.next_output_slot = 0; @@ -557,30 +466,30 @@ impl Nl4dDenoiser { } } - /// The motion field and confidence the last pass gave the fused - /// kernel, or `None` before any pass has run. + /// The motion field and confidence the last pass gave the fused kernel. /// - /// This is a synchronous readback for measurement tooling, not a - /// stable interface. + /// A synchronous readback for measurement tooling. It is not a stable interface. #[doc(hidden)] pub fn motion_snapshot(&self) -> Option { let fields = self.last_fields.as_ref()?; - let mc = self.front.motion_ctx(); - Some(read_snapshot( + let motion_ctx = self.front.motion_ctx(); + let snapshot = read_snapshot( self.front.compute_client(), fields, self.temporal_radius, - mc.blocks_x, - mc.blocks_y, - mc.step, - mc.blksize, - )) + motion_ctx.blocks_x, + motion_ctx.blocks_y, + motion_ctx.step, + motion_ctx.blksize, + ); + + Some(snapshot) } /// Reads back the grain chunks measured since the last call. Empty when export is off. /// /// Measured chunks stay on the GPU until drained, so callers drain after each flush. - pub fn drain_grain_chunks(&mut self) -> Result, DenoiserError> { + pub fn drain_grain_chunks(&mut self) -> Result, anyhow::Error> { let Some(grain) = self.grain.as_mut() else { return Ok(Vec::new()); }; @@ -601,18 +510,15 @@ impl Nl4dDenoiser { self.grain.as_ref().map_or(0, |grain| grain.measured_with_entry()) } - /// The front end this denoiser drives. #[cfg(test)] pub(crate) fn front_for_test(&self) -> &NlmDenoiser { &self.front } - /// How many tail frames [`Self::flush`] must emit for the stream - /// pushed so far. + /// How many tail frames [Self::finish_passes] must emit. /// - /// A stream that filled its ring holds `2 * temporal_radius` frames - /// whose regions are not yet read out. A shorter stream holds every - /// frame it pushed. + /// A stream that filled its ring holds `2 * temporal_radius` unread regions. A shorter stream + /// holds every frame it pushed. fn flush_target(&self) -> usize { let real_pushes = self.front.real_pushes(); if real_pushes == 0 { @@ -622,27 +528,15 @@ impl Nl4dDenoiser { } } - /// Runs the grouping, filtering, and aggregation kernels for one pass. - /// - /// The pass is centred on the physical ring slot `view.centre_slot`. - /// It groups the centre frame against every other frame in the ring, - /// and each filtered member scatters into the region of - /// `self.accum`/`self.wsum` for the frame it came from (see - /// [`collab_fused`]'s scatter). So one pass adds to every region in - /// the ring, not just the centre's. + /// Runs the grouping, filtering and aggregation kernels for one pass centred on `view`. /// - /// Before scattering, `clear` picks which regions to zero. - /// [`AccumClear::WholeRing`] zeroes every slot, since nothing has - /// cleared a new stream's ring. - /// [`AccumClear::NewestRegion`] does the same for the slot - /// `temporal_radius` ahead of the centre, which the newest frame has - /// just taken over. [`AccumClear::Nothing`] leaves every region as it - /// is, for an edge pass that sees no new frame. - fn run_pass(&mut self, view: &RingView, clear: AccumClear) -> Result<(), DenoiserError> { - // The frame-slot contract: `collab_fused`'s `centre_slot` and - // the ring view's own centre must be the same physical slot, or - // a member gets grouped against one frame and scattered as - // though it belonged to another. + /// Each filtered member scatters into the accumulator region of the frame it came from, so one + /// pass adds to every region in the ring. `clear` picks which regions are zeroed first. The + /// newest region is the slot `temporal_radius` ahead of the centre, which the newest frame has + /// just taken over. + fn run_pass(&mut self, view: &RingView, clear: AccumClear) -> Result<(), anyhow::Error> { + // The kernel's centre slot must be the ring view's centre, or a member is grouped against + // one frame and scattered into another. let centre_slot = view.centre_slot; let client = self.front.compute_client().clone(); @@ -653,9 +547,6 @@ impl Nl4dDenoiser { let frame_len = pixels * stored_ch as usize; let total_frames = 1 + 2 * self.temporal_radius; let ring_len = frame_len * total_frames as usize; - - // The accumulators' own ring, one region per physical slot of - // the frame ring above, the same `total_frames` count. let accum_ring_len = frame_len * total_frames as usize; let wsum_ring_len = pixels * total_frames as usize; @@ -663,17 +554,20 @@ impl Nl4dDenoiser { let mv_len = (neighbours * view.mv_stride) as usize; let conf_len = (neighbours * view.conf_stride) as usize; - let neighbour_slots_buf = client.create_from_slice(u32::as_bytes(&view.neighbour_slots)); + let neighbour_slot_bytes = u32::as_bytes(&view.neighbour_slots); + let neighbour_slots_buf = client.create_from_slice(neighbour_slot_bytes); let sigmas = self.front.current_sigmas_temporal_only(); let mut sigma_host = vec![0.0f32; stored_ch as usize]; sigma_host[..channels_count as usize].copy_from_slice(&sigmas[..channels_count as usize]); - self.sigma_buf = client.create_from_slice(f32::as_bytes(&sigma_host)); - let wnorm = weight_scale(sigma_host[0], &self.dct_profile); + let sigma_bytes = f32::as_bytes(&sigma_host); + self.sigma_buf = client.create_from_slice(sigma_bytes); + let weight_norm = weight_scale(sigma_host[0], &self.dct_profile); let curve_ratios = self.front.current_noise_curve().map(|curve| curve.ratios); let (ratios, curve_valid) = noise_curve_upload(curve_ratios, self.apply_noise_map); - let noise_curve_buf = client.create_from_slice(f32::as_bytes(&ratios)); + let ratio_bytes = f32::as_bytes(&ratios); + let noise_curve_buf = client.create_from_slice(ratio_bytes); let classes = self.front.current_quarter_classes(); if let Some(classes) = classes { @@ -687,7 +581,12 @@ impl Nl4dDenoiser { let map_upload = strength_map_upload(classes, curve_valid, self.luma_map, self.chroma_map_boost); let (strength_map_buf, map_mode) = match map_upload { - Some((multipliers, mode)) => (client.create_from_slice(f32::as_bytes(&multipliers)), mode), + Some((multipliers, mode)) => { + let multiplier_bytes = f32::as_bytes(&multipliers); + let strength_map_buf = client.create_from_slice(multiplier_bytes); + + (strength_map_buf, mode) + }, None => (self.unit_map.clone(), STRENGTH_MAP_OFF), }; let map_len = (self.map_cols * self.map_rows) as usize; @@ -696,31 +595,22 @@ impl Nl4dDenoiser { let refs_y = refs_along(self.height); let refs = ref_count(self.width, self.height); - // The kernel packs eight references into one 64-lane cube, so - // its grid is an eighth as wide as the reference grid along x. - let collab_grid = CubeCount::new_2d(fused_cubes_x(self.width), refs_y); + let collab_cubes_x = fused_cubes_x(self.width); + let collab_grid = CubeCount::new_2d(collab_cubes_x, refs_y); let collab_dim = CubeDim::new_1d(64); let zero_dim = 256u32; - // Sized for one frame's worth of the ring, and issued once per - // region the pass clears. - // - // Still clamped to the GPU's 65,535-workgroups-per-dimension - // limit, because one frame alone can exceed it. A 4:4:4 4K frame - // or an 8K luma plane both need more than that at 256 threads - // each. `collab_zero_accum` strides, so a clamped launch still - // reaches every slot in the frame. + // Clamped to the 65,535 workgroup limit, which a 4:4:4 4K frame or an 8K luma plane alone + // exceeds. `collab_zero_accum` strides, so a clamped launch still reaches every slot. let zero_workgroups_one_frame = (frame_len as u32).div_ceil(zero_dim).min(MAX_GRID_1D); let zero_grid_one_frame = CubeCount::new_1d(zero_workgroups_one_frame); let zero_total_threads_one_frame = zero_workgroups_one_frame * zero_dim; - let mc = self.front.motion_ctx(); - let blk_step = mc.step; - let blksize = mc.blksize; - let blocks_x = mc.blocks_x; - let blocks_y = mc.blocks_y; + let motion_ctx = self.front.motion_ctx(); + let blk_step = motion_ctx.step; + let blksize = motion_ctx.blksize; + let blocks_x = motion_ctx.blocks_x; + let blocks_y = motion_ctx.blocks_y; - // The physical slots whose regions this pass resets before - // scattering. let cleared_slots = match clear { AccumClear::WholeRing => 0..total_frames, AccumClear::NewestRegion => { @@ -732,24 +622,24 @@ impl Nl4dDenoiser { self.passes_run += 1; - // The field the fused kernel reads, the regularised one when the - // pass is on. let (mv_field, confidence) = match (self.reg_mv.as_ref(), self.reg_conf.as_ref()) { - (Some(mv), Some(conf)) => { + (Some(reg_mv), Some(reg_conf)) => { + let sad_noise_floor = self.front.sad_noise_floor_value(); + let thsad = self.front.thsad_value(); run_regularise::( &client, - mc, + motion_ctx, view, self.width, self.height, self.field_lambda, - self.front.sad_noise_floor_value(), - self.front.thsad_value(), - mv, - conf, - ) - .map_err(DenoiserError::Other)?; - (mv.clone(), conf.clone()) + sad_noise_floor, + thsad, + reg_mv, + reg_conf, + )?; + + (reg_mv.clone(), reg_conf.clone()) }, _ => (view.mv_field.clone(), view.confidence.clone()), }; @@ -762,15 +652,12 @@ impl Nl4dDenoiser { neighbours, }); + let frames_per_volume = grid_frames(self.temporal_radius); + unsafe { - // Clearing the whole ring in one dispatch would need - // `accum_ring_len.div_ceil(zero_dim)` workgroups, which grows - // with `total_frames`. At `temporal_radius = 4` a 1080p luma - // plane alone needs 72,900, already over the GPU's 65,535 - // limit. A rejected dispatch would leave the ring holding - // `client.empty`'s undefined memory instead of zero, which a - // fresh stream's first frames would then aggregate as though - // it were real. So each cleared region gets its own dispatch. + // One dispatch over the whole ring exceeds the 65,535 workgroup limit at + // `temporal_radius = 4` for a 1080p luma plane. A rejected dispatch leaves undefined + // memory that a fresh stream would aggregate as real, so each region gets its own. for slot in cleared_slots { collab_zero_accum::launch_unchecked::( &client, @@ -809,11 +696,11 @@ impl Nl4dDenoiser { self.lambda_ht, curve_valid, map_mode, - wnorm, + weight_norm, self.accum_scale, self.warp_uniform, self.temporal_radius, - grid_frames(self.temporal_radius), + frames_per_volume, self.refine, view.mv_stride, view.conf_stride, @@ -854,7 +741,8 @@ impl Nl4dDenoiser { // The motion field numbers neighbours in logical order and skips the centre, so the // frame at `centre + 1` is neighbour `centre`. let next_neighbour = (centre < last_real).then_some(centre); - grain.save_vectors(self.front.compute_client(), fields, centre_slot, next_neighbour); + let client = self.front.compute_client(); + grain.save_vectors(client, fields, centre_slot, next_neighbour); } /// Measures the grain of the frame just normalised into `output_slot`. @@ -868,10 +756,10 @@ impl Nl4dDenoiser { grain.measure(client, input, &self.outputs, slot_t, slot_next, output_slot); } - /// Normalises the accumulator region at physical slot `region_slot` into the next output buffer. + /// Normalises the accumulator region at `region_slot` into the next output buffer. /// - /// Returns the output buffer and its index. The region is left as it - /// was, and a later pass clears it before reuse. + /// Returns the output buffer and its index. The region is left as it was, and a later pass + /// clears it before reuse. fn normalise_region(&mut self, region_slot: u32) -> (Handle, usize) { let client = self.front.compute_client().clone(); let stored_ch = self.channels.storage_count(); @@ -880,7 +768,9 @@ impl Nl4dDenoiser { let total_frames = 1 + 2 * self.temporal_radius; let accum_ring_len = frame_len * total_frames as usize; let wsum_ring_len = pixels * total_frames as usize; - let agg_grid = CubeCount::new_2d(self.width.div_ceil(BLOCK_X), self.height.div_ceil(BLOCK_Y)); + let agg_cubes_x = self.width.div_ceil(BLOCK_X); + let agg_cubes_y = self.height.div_ceil(BLOCK_Y); + let agg_grid = CubeCount::new_2d(agg_cubes_x, agg_cubes_y); let agg_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); let slot = self.next_output_slot; @@ -905,30 +795,6 @@ impl Nl4dDenoiser { (self.outputs[slot].clone(), slot) } - - /// Starts an async readback of `handle`, wrapped in the same - /// [`Pending`] type [`NlmDenoiser`] returns. - /// - /// `wire_dst` is the packed-word buffer belonging to the slot - /// `handle` came from, and is `None` for an `f32` readback. - fn start_readback(&self, handle: Handle, wire_dst: Option<&Handle>, format: OutputFormat) -> Pending { - let pixels = (self.width * self.height) as usize; - start_readback( - self.front.compute_client(), - handle, - wire_dst, - self.channels.count(), - self.channels.storage_count(), - pixels, - format, - ) - } - - /// The packed-word destinations, which are `Some` only in wire mode. - #[cfg(test)] - pub(crate) fn wire_outputs_for_test(&self) -> Option<&[Handle; 2]> { - self.wire_outputs.as_ref() - } } /// The noise curve a pass uploads, and the flag that tells the kernel to apply it. diff --git a/av-denoise-core/src/nl4d/engine.rs b/av-denoise-core/src/nl4d/engine.rs new file mode 100644 index 0000000..5458562 --- /dev/null +++ b/av-denoise-core/src/nl4d/engine.rs @@ -0,0 +1,244 @@ +use std::collections::VecDeque; + +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::denoiser::{CompletedRegion, Nl4dDenoiser, buffer_sizes}; +use super::grain::GrainChunk; +use super::options::{Nl4dOptions, resolve_params}; +use crate::collab::PATCH_SIZE; +use crate::engine::{DevicePlane, EdgePadding, EgressSource, Engine, Geometry, WindowSpan, egress}; +use crate::error::Error; +use crate::nlmeans::denoiser::check_u32_indexable; + +/// Four-dimensional collaborative denoising over GPU planes. +pub struct Nl4d { + inner: Nl4dDenoiser, + client: ComputeClient, + geometry: Geometry, + temporal_radius: u32, + /// Regions finished but not yet emitted, oldest first. + pending: VecDeque, + /// Whether the regions in `pending` are a stream's tail. + finishing: bool, + /// Pushes since the last stream start, used to mark a context-led stream as a continuation. + pushes: usize, + /// Whether the current stream has had a real push, after which context frames are refused. + pushed: bool, + poisoned: bool, +} + +impl Nl4d { + /// Builds the engine, returning [Error::InvalidGeometry] or [Error::InvalidOptions] for unusable input. + pub fn new(client: &ComputeClient, options: Nl4dOptions, geometry: Geometry) -> Result { + if geometry.width < PATCH_SIZE || geometry.height < PATCH_SIZE { + let message = format!( + "frame dimensions {}x{} must be at least {PATCH_SIZE}x{PATCH_SIZE}", + geometry.width, geometry.height, + ); + return Err(Error::InvalidGeometry(message)); + } + + geometry.validate()?; + + let params = resolve_params(&options, geometry.channels)?; + + // Buffer sizing builds the motion block grid, which needs validated motion options. + let validated = params.validate(); + validated.map_err(Error::InvalidOptions)?; + let front_validated = params.nlm.validate(); + front_validated.map_err(|error| { + let message = error.to_string(); + Error::InvalidOptions(message) + })?; + + let ring_slots = 1 + 2 * u64::from(params.temporal_radius); + geometry.check_ring_fits(ring_slots)?; + + let sizes = buffer_sizes(client, ¶ms, geometry.width, geometry.height); + let indexable = check_u32_indexable(&sizes); + indexable.map_err(Error::InvalidGeometry)?; + + let inner = Nl4dDenoiser::new(client, params, geometry.width, geometry.height); + let inner = inner.map_err(Error::InvalidOptions)?; + + Ok(Self { + inner, + client: client.clone(), + geometry, + temporal_radius: options.temporal_radius, + pending: VecDeque::new(), + finishing: false, + pushes: 0, + pushed: false, + poisoned: false, + }) + } + + fn check_usable(&self) -> Result<(), Error> { + if self.poisoned { + return Err(Error::NeedsReset); + } + + Ok(()) + } + + fn check_nothing_pending(&self) -> Result<(), Error> { + if !self.pending.is_empty() { + return Err(Error::OutputsPending); + } + + Ok(()) + } + + fn check_no_push_yet(&self) -> Result<(), Error> { + if self.pushed { + return Err(Error::ContextAfterPush); + } + + Ok(()) + } + + /// Poisons the engine if `result` is an error. + fn guard(&mut self, result: Result) -> Result { + if result.is_err() { + self.poisoned = true; + } + + result.map_err(Error::Gpu) + } + + fn restart_stream(&mut self) { + self.inner.reset_stream(); + self.finishing = false; + self.pushes = 0; + self.pushed = false; + } + + fn write_out(&self, frame: &Handle, planes: &[DevicePlane<'_>]) { + let channels = self.geometry.channels; + // The constructor bounds the ring to `u32` elements, so one plane always fits. + let pixels = self.geometry.pixels() as u32; + let source = EgressSource { + frame, + pixels, + channels: channels.count(), + stored_ch: channels.storage_count(), + }; + let placeholder = self.inner.placeholder(); + + egress(&self.client, source, planes, self.geometry.output, placeholder); + } + + #[cfg(test)] + pub(crate) fn fail_through_guard_for_test(&mut self) -> Result<(), Error> { + let failure = Err(anyhow::anyhow!("forced failure")); + + self.guard(failure) + } +} + +impl Engine for Nl4d { + fn push(&mut self, planes: &[DevicePlane<'_>]) -> Result { + self.check_usable()?; + self.check_nothing_pending()?; + self.geometry.check_planes(planes, self.geometry.input)?; + + let pushed = self.inner.push_planes(planes, self.geometry.input); + self.guard(pushed)?; + + self.pushed = true; + + let submitted = self.inner.submit_passes(); + let region = self.guard(submitted)?; + if let Some(region) = region { + self.pending.push_back(region); + } + + self.pushes += 1; + + Ok(self.pending.len()) + } + + fn push_context(&mut self, planes: &[DevicePlane<'_>]) -> Result<(), Error> { + self.check_usable()?; + self.check_nothing_pending()?; + self.geometry.check_planes(planes, self.geometry.input)?; + self.check_no_push_yet()?; + + if self.pushes == 0 { + self.inner.mark_continuation(); + } + + let pushed = self.inner.push_planes(planes, self.geometry.input); + self.guard(pushed)?; + + self.pushes += 1; + + Ok(()) + } + + fn emit_into(&mut self, planes: &[DevicePlane<'_>]) -> Result<(), Error> { + self.check_usable()?; + self.geometry.check_planes(planes, self.geometry.output)?; + + let Some(region) = self.pending.pop_front() else { + return Err(Error::NothingToEmit); + }; + + let frame = self.inner.read_region(region); + self.write_out(&frame, planes); + + if self.finishing && self.pending.is_empty() { + self.restart_stream(); + } + + Ok(()) + } + + fn finish(&mut self) -> Result { + self.check_usable()?; + self.check_nothing_pending()?; + + let finished = self.inner.finish_passes(); + let regions = self.guard(finished)?; + + let count = regions.len(); + self.pending.extend(regions); + + if count == 0 { + self.restart_stream(); + } else { + self.finishing = true; + } + + Ok(count) + } + + fn reset(&mut self) { + self.pending.clear(); + self.poisoned = false; + self.restart_stream(); + } + + fn window_span(&self) -> WindowSpan { + let span = 2 * self.temporal_radius as usize; + + WindowSpan { + behind: span, + ahead: span, + edges: EdgePadding::Shifted, + } + } + + fn max_held_frames(&self) -> usize { + 2 * self.temporal_radius as usize + } + + fn drain_grain_chunks(&mut self) -> Result, Error> { + self.check_usable()?; + + let drained = self.inner.drain_grain_chunks(); + self.guard(drained) + } +} diff --git a/av-denoise-core/src/nl4d/grain/chunk.rs b/av-denoise-core/src/nl4d/grain/chunk.rs index b783c68..74640c4 100644 --- a/av-denoise-core/src/nl4d/grain/chunk.rs +++ b/av-denoise-core/src/nl4d/grain/chunk.rs @@ -1,14 +1,13 @@ use super::consts::{HIST_LEN, LAG_COUNT, STRENGTH_GROUPS}; -/// The grain statistics of up to [CHUNK_FRAMES](crate::nl4d::grain::consts::CHUNK_FRAMES) -/// consecutive completed frames of one scene. +/// The grain statistics of up to `CHUNK_FRAMES` consecutive completed frames of one scene. #[derive(Debug, Clone, PartialEq)] pub struct GrainChunk { /// How many completed frames the chunk covers. pub frames: u32, - /// Accepted source blocks, 16 luma bins by 64 std buckets. + /// Accepted source block counts, 16 luma bins by 64 std buckets. pub source_hist: Vec, - /// Accepted kept-grain blocks, in the same layout. + /// Accepted kept-grain block counts, in the same layout. pub kept_hist: Vec, /// The 46 source autocovariance sums of each strength group, in normalised units squared. pub autocov: Vec, diff --git a/av-denoise-core/src/nl4d/grain/consts.rs b/av-denoise-core/src/nl4d/grain/consts.rs index 837cd3a..f7416a8 100644 --- a/av-denoise-core/src/nl4d/grain/consts.rs +++ b/av-denoise-core/src/nl4d/grain/consts.rs @@ -4,7 +4,7 @@ pub(crate) const LUMA_BINS: usize = 16; pub(crate) const STD_BUCKETS: usize = 64; /// Counts in one histogram, luma bins by std buckets. pub(crate) const HIST_LEN: usize = LUMA_BINS * STD_BUCKETS; -/// The 46 autocovariance lags, half-plane, `dy` in `0..=3` and `dx` in `-6..=6`. +/// The 46 half-plane autocovariance lags, with `dy` in `0..=3` and `dx` in `-6..=6`. pub(crate) const LAG_COUNT: usize = 46; /// The lag sums plus the accepted pixel count. pub(crate) const AUTOCOV_LEN: usize = LAG_COUNT + 1; @@ -27,6 +27,9 @@ pub(crate) const LAGS: [(i32, i32); LAG_COUNT] = lag_table(); /// AV1's 24 causal lag-3 offsets in coefficient order, as `(dy, dx)`. pub(crate) const AR_OFFSETS: [(i32, i32); AR_COEFFS] = ar_offset_table(); +// A measured block needs a motion confidence of at least CONF_MIN, a denoised range below +// FLAT_RANGE, a denoised mean between LUMA_LOW and LUMA_HIGH, no pixel past CLIP_LOW or CLIP_HIGH, +// and a grain std above STD_MIN. pub(crate) const FLAT_RANGE: f32 = 5.0 / 255.0; pub(crate) const LUMA_LOW: f32 = 20.0 / 255.0; pub(crate) const LUMA_HIGH: f32 = 235.0 / 255.0; @@ -34,15 +37,26 @@ pub(crate) const CLIP_LOW: f32 = 4.0 / 255.0; pub(crate) const CLIP_HIGH: f32 = 251.0 / 255.0; pub(crate) const CONF_MIN: f32 = 0.8; pub(crate) const STD_MIN: f32 = 0.05 / 255.0; +/// The top std bucket edge. pub(crate) const STD_MAX: f32 = 32.0 / 255.0; +/// Fewest blocks for a luma bin to give a strength reading. +/// +/// A bin's kept reading counts as 0 below it. pub(crate) const MIN_BLOCKS_PER_BIN: u64 = 200; +/// Fewest readable luma bins for a segment to have its own strength. pub(crate) const MIN_POPULATED_BINS: usize = 3; +/// Fewest source pixels for a segment to solve its own texture. pub(crate) const MIN_AR_PIXELS: f64 = 50_000.0; +/// Most completed frames per [GrainChunk](crate::nl4d::grain::GrainChunk). pub(crate) const CHUNK_FRAMES: u32 = 24; +/// The median source std ratio, either way, past which a chunk starts a new segment. pub(crate) const DRIFT: f64 = 1.15; +/// Fewest source blocks for a chunk to start a new segment. pub(crate) const MIN_CHUNK_BLOCKS: u64 = 2_000; +/// Templates measured when calibrating a texture's std. pub(crate) const TEMPLATE_SEEDS: u32 = 24; +/// AV1 allows at most 14 luma scaling points. pub(crate) const MAX_POINTS: usize = 14; const fn lag_table() -> [(i32, i32); LAG_COUNT] { @@ -56,10 +70,13 @@ const fn lag_table() -> [(i32, i32); LAG_COUNT] { table[next] = (dy, dx); next += 1; } + dx += 1; } + dy += 1; } + table } @@ -74,9 +91,12 @@ const fn ar_offset_table() -> [(i32, i32); AR_COEFFS] { table[next] = (dy, dx); next += 1; } + dx += 1; } + dy += 1; } + table } diff --git a/av-denoise-core/src/nl4d/grain/export.rs b/av-denoise-core/src/nl4d/grain/export.rs index ed01d16..047ac4e 100644 --- a/av-denoise-core/src/nl4d/grain/export.rs +++ b/av-denoise-core/src/nl4d/grain/export.rs @@ -77,8 +77,8 @@ struct Completed { /// The grain measurement state of one nl4d denoiser. /// /// Each frame's vectors to the next frame are saved in a ring entry keyed by its physical ring -/// slot. Each completed frame adds into the newest chunk record, which stays on the GPU until -/// [Self::drain] reads every chunk back. +/// slot. A frame's vectors must be saved before the frame is measured. Each completed frame adds +/// into the newest chunk record, which stays on the GPU until [Self::drain] reads every chunk back. pub(crate) struct GrainExport { geometry: GrainGeometry, saved_mv: Handle, @@ -100,10 +100,14 @@ impl GrainExport { let saved_mv_host = vec![0i32; geometry.saved_mv_len()]; let saved_conf_host = vec![0.0f32; geometry.saved_conf_len()]; - let saved_mv = client.create_from_slice(i32::as_bytes(&saved_mv_host)); - let saved_conf = client.create_from_slice(f32::as_bytes(&saved_conf_host)); - let edges = client.create_from_slice(f32::as_bytes(&edges_host)); - let partials = client.empty(geometry.partials_len().max(1) * size_of::()); + let saved_mv_bytes = i32::as_bytes(&saved_mv_host); + let saved_mv = client.create_from_slice(saved_mv_bytes); + let saved_conf_bytes = f32::as_bytes(&saved_conf_host); + let saved_conf = client.create_from_slice(saved_conf_bytes); + let edges_bytes = f32::as_bytes(&edges_host); + let edges = client.create_from_slice(edges_bytes); + let partials_bytes = geometry.partials_len().max(1) * size_of::(); + let partials = client.empty(partials_bytes); Self { saved_valid: vec![false; geometry.ring_frames as usize], @@ -163,20 +167,20 @@ impl GrainExport { /// Measures the frame that just completed into the newest chunk. /// - /// `slot_next` is the ring slot of the next real frame, `None` for a stream's last frame. + /// `next_slot` is the ring slot of the next real frame, `None` for a stream's last frame. pub(crate) fn measure( &mut self, client: &ComputeClient, input: &Handle, outputs: &[Handle; 2], - slot_t: u32, - slot_next: Option, + frame_slot: u32, + next_slot: Option, output_slot: usize, ) { - let has_source = slot_next.is_some() && self.saved_valid[slot_t as usize]; + let has_source = next_slot.is_some() && self.saved_valid[frame_slot as usize]; #[cfg(test)] - if self.saved_valid[slot_t as usize] { + if self.saved_valid[frame_slot as usize] { self.measured_with_entry += 1; } @@ -184,8 +188,8 @@ impl GrainExport { .last_completed .filter(|previous| self.saved_valid[previous.ring_slot as usize]); let has_kept = kept_from.is_some(); - let next = slot_next.unwrap_or(slot_t); - let kept_entry = kept_from.map_or(slot_t, |previous| previous.ring_slot); + let next_or_own_slot = next_slot.unwrap_or(frame_slot); + let kept_entry = kept_from.map_or(frame_slot, |previous| previous.ring_slot); let prev_output = kept_from.map_or(output_slot, |previous| previous.output_slot); let (chunk_hist, chunk_autocov) = self.open_chunk(client); @@ -209,9 +213,9 @@ impl GrainExport { ArrayArg::from_raw_parts(self.edges.clone(), EDGES_LEN), ArrayArg::from_raw_parts(chunk_hist, 2 * HIST_LEN), ArrayArg::from_raw_parts(self.partials.clone(), geometry.partials_len()), - slot_t, - next, - slot_t, + frame_slot, + next_or_own_slot, + frame_slot, kept_entry, has_source as u32, has_kept as u32, @@ -240,7 +244,7 @@ impl GrainExport { } self.last_completed = Some(Completed { - ring_slot: slot_t, + ring_slot: frame_slot, output_slot, }); } @@ -250,8 +254,10 @@ impl GrainExport { if !self.chunk_open { let hist_host = vec![0i32; 2 * HIST_LEN]; let autocov_host = vec![0.0f32; GROUPED_AUTOCOV_LEN]; - let hist = client.create_from_slice(i32::as_bytes(&hist_host)); - let autocov = client.create_from_slice(f32::as_bytes(&autocov_host)); + let hist_bytes = i32::as_bytes(&hist_host); + let hist = client.create_from_slice(hist_bytes); + let autocov_bytes = f32::as_bytes(&autocov_host); + let autocov = client.create_from_slice(autocov_bytes); self.chunks.push(ChunkBuffers { hist, autocov, @@ -310,6 +316,7 @@ fn read_chunk( let mut chunk = GrainChunk::empty(); chunk.frames = buffer.frames; + for (target, &count) in chunk.source_hist.iter_mut().zip(source_counts) { *target = count as u32; } diff --git a/av-denoise-core/src/nl4d/grain/fit.rs b/av-denoise-core/src/nl4d/grain/fit.rs index 39299f0..738a262 100644 --- a/av-denoise-core/src/nl4d/grain/fit.rs +++ b/av-denoise-core/src/nl4d/grain/fit.rs @@ -24,6 +24,7 @@ pub(crate) fn bucket_edges() -> [f32; STD_BUCKETS + 1] { let fraction = index as f64 / STD_BUCKETS as f64; *edge = (STD_MIN as f64 * ratio.powf(fraction)) as f32; } + edges } @@ -36,6 +37,7 @@ pub(crate) fn bucket_of(std: f32, edges: &[f32]) -> usize { bucket = index; } } + bucket } @@ -205,6 +207,8 @@ fn solve(mut matrix: Vec<[f64; AR_COEFFS + 1]>) -> Option<[f64; AR_COEFFS]> { } /// Rounds the weights to AV1 integers with the finest shift that keeps each in a signed byte. +/// +/// AV1 allows an AR shift between 6 and 9. Weights that overflow at a shift of 6 are clamped. pub(crate) fn quantise_ar(coeffs: &[f64; AR_COEFFS]) -> ([i32; AR_COEFFS], u32) { for shift in [9u32, 8, 7, 6] { let scale = (1u32 << shift) as f64; @@ -222,13 +226,16 @@ pub(crate) fn quantise_ar(coeffs: &[f64; AR_COEFFS]) -> ([i32; AR_COEFFS], u32) /// Turns per-bin target sigmas into AV1 scaling points and the scaling shift. /// /// `points` holds `(luma, sigma)` with luma in 8-bit codes and sigma in normalised units. -/// `sigma_template` is the template std in 8-bit units. +/// `sigma_template` is the template std in 8-bit units. AV1 allows a scaling shift between 8 and +/// 11 and a byte per scaling value, so the finest shift that fits the largest value wins and +/// values past 255 are capped. pub(crate) fn scaling_points(points: &[(u8, f64)], sigma_template: f64) -> (Vec<(u8, u8)>, u32) { let raw: Vec = points .iter() .map(|&(_, sigma)| sigma * 255.0 / sigma_template) .collect(); let largest = raw.iter().copied().fold(0.0f64, f64::max); + let mut shift = 8u32; for candidate in [11u32, 10, 9, 8] { if largest * (1u32 << candidate) as f64 <= 255.0 { @@ -243,5 +250,6 @@ pub(crate) fn scaling_points(points: &[(u8, f64)], sigma_template: f64) -> (Vec< .zip(raw.iter()) .map(|(&(luma, _), &value)| (luma, (value * factor).round().min(255.0) as u8)) .collect(); + (scaled, shift) } diff --git a/av-denoise-core/src/nl4d/grain/segment.rs b/av-denoise-core/src/nl4d/grain/segment.rs index dc538b8..7e53124 100644 --- a/av-denoise-core/src/nl4d/grain/segment.rs +++ b/av-denoise-core/src/nl4d/grain/segment.rs @@ -50,6 +50,7 @@ struct BinStrength { kept: f64, } +/// Quantised AR weights and their shift. type Texture = ([i32; AR_COEFFS], u32); /// The texture band's lower edge, as a multiple of the median source std. @@ -79,16 +80,16 @@ fn overall_median(hist: &[u32], edges: &[f32]) -> Option { hist_median(&merged, edges) } -/// Splits a scene into segments, closing one when a chunk's grain drifts past `DRIFT`. +/// Splits a scene into segments, closing one when a chunk's grain drifts past [DRIFT]. pub(crate) fn segment_scene(scene: &SceneGrain) -> Vec { let edges = bucket_edges(); let mut segments: Vec = Vec::new(); let mut next_frame = scene.first_frame; for chunk in &scene.chunks { - let first = next_frame; - let last = first + chunk.frames as u64 - 1; - next_frame = last + 1; + let chunk_first = next_frame; + let chunk_last = chunk_first + chunk.frames as u64 - 1; + next_frame = chunk_last + 1; let starts_new = match segments.last() { None => true, @@ -96,15 +97,15 @@ pub(crate) fn segment_scene(scene: &SceneGrain) -> Vec { }; if starts_new { segments.push(Segment { - first_frame: first, - last_frame: last, + first_frame: chunk_first, + last_frame: chunk_last, stats: chunk.clone(), }); continue; } let open = segments.last_mut().expect("a segment is open"); - open.last_frame = last; + open.last_frame = chunk_last; open.stats.merge(chunk); } @@ -237,6 +238,7 @@ fn borrow(parts: &[Parts], index: usize, value: impl Fn(&Parts) -> Option) let Some(part) = parts.get(candidate) else { continue; }; + if same_scene_only && part.scene != scene { continue; } @@ -254,6 +256,8 @@ fn borrow(parts: &[Parts], index: usize, value: impl Fn(&Parts) -> Option) fn build_entry(part: &Parts, strength: &[BinStrength], texture: Texture) -> FittedEntry { let (coeffs, ar_shift) = texture; let (sigma_template, template_median) = template_stats(&coeffs, ar_shift); + // The measured stds are medians of mean-removed 8x8 blocks. The template's ratio of its + // whole-field std to its block median converts them to a whole-field std. let factor = sigma_template / template_median; let mut chosen: Vec = strength.to_vec(); @@ -271,6 +275,7 @@ fn build_entry(part: &Parts, strength: &[BinStrength], texture: Texture) -> Fitt .map(|entry| { let sigma_source = entry.source * factor; let sigma_kept = entry.kept * factor; + // Only the grain the denoiser removed is synthesised, so the kept grain comes off. let variance = (sigma_source * sigma_source - sigma_kept * sigma_kept).max(0.0); let centre = (entry.bin as f64 + 0.5) * 256.0 / LUMA_BINS as f64; (centre.round() as u8, variance.sqrt()) @@ -278,6 +283,7 @@ fn build_entry(part: &Parts, strength: &[BinStrength], texture: Texture) -> Fitt .collect(); let (points, scaling_shift) = scaling_points(&targets, sigma_template); + FittedEntry { first_frame: part.first_frame, last_frame: part.last_frame, diff --git a/av-denoise-core/src/nl4d/grain/table.rs b/av-denoise-core/src/nl4d/grain/table.rs index dca3d56..0b283a1 100644 --- a/av-denoise-core/src/nl4d/grain/table.rs +++ b/av-denoise-core/src/nl4d/grain/table.rs @@ -3,6 +3,7 @@ use std::fmt::Write; use super::consts::AR_COEFFS; use super::segment::{FittedEntry, SceneGrain, fit_scenes}; +/// Grain table timestamps count 100 ns ticks. const TICKS_PER_SECOND: u128 = 10_000_000; const SEED_BASE: u64 = 7391; const SEED_STEP: u64 = 1237; @@ -61,6 +62,7 @@ pub(crate) fn format_table(entries: &[FittedEntry], frame_rate: (u64, u64)) -> S let coeffs = join_values(&entry.ar_coeffs); writeln!(text, "E {start} {end} 1 {seed} 1").expect("writing to a String"); + // Lag 3, the AR and scaling shifts, and overlap on. writeln!( text, "\tp 3 {} 0 {} 0 1 128 192 256 128 192 256", @@ -68,6 +70,7 @@ pub(crate) fn format_table(entries: &[FittedEntry], frame_rate: (u64, u64)) -> S ) .expect("writing to a String"); writeln!(text, "\tsY {point_count} {points}").expect("writing to a String"); + // Chroma scales to zero everywhere, so only luma gets grain. text.push_str("\tsCb 2 0 0 255 0\n"); text.push_str("\tsCr 2 0 0 255 0\n"); writeln!(text, "\tcY {coeffs}").expect("writing to a String"); diff --git a/av-denoise-core/src/nl4d/grain/template.rs b/av-denoise-core/src/nl4d/grain/template.rs index 19d3026..7b9cc8e 100644 --- a/av-denoise-core/src/nl4d/grain/template.rs +++ b/av-denoise-core/src/nl4d/grain/template.rs @@ -33,10 +33,10 @@ fn round2(value: i32, shift: u32) -> i32 { /// Builds the 73x82 luma grain template an 8-bit AV1 decoder makes for these weights and seed. pub(crate) fn luma_template(coeffs: &[i32; AR_COEFFS], ar_shift: u32, seed: u16) -> Vec { - let mut rng = Av1Random::new(seed); + let mut random = Av1Random::new(seed); let mut grain = vec![0i32; TEMPLATE_ROWS * TEMPLATE_COLS]; for sample in grain.iter_mut() { - let index = rng.next(11) as usize; + let index = random.next(11) as usize; let gaussian = GAUSSIAN_SEQUENCE[index] as i32; *sample = round2(gaussian, 4); } @@ -65,8 +65,7 @@ pub(crate) fn template_seed(index: u32) -> u16 { /// The grain std of these weights' templates, and the median std of their 8x8 blocks. /// -/// Both read the 64x64 interior of [TEMPLATE_SEEDS](crate::nl4d::grain::consts::TEMPLATE_SEEDS) -/// templates. A block's std removes its own mean. +/// Both read the 64x64 interior of [TEMPLATE_SEEDS] templates. A block's std removes its own mean. pub(crate) fn template_stats(coeffs: &[i32; AR_COEFFS], ar_shift: u32) -> (f64, f64) { let capacity = TEMPLATE_SEEDS as usize * INTERIOR_SIZE * INTERIOR_SIZE; let mut samples = Vec::with_capacity(capacity); @@ -75,7 +74,8 @@ pub(crate) fn template_stats(coeffs: &[i32; AR_COEFFS], ar_shift: u32) -> (f64, let seed = template_seed(index); let template = luma_template(coeffs, ar_shift, seed); let interior = interior_of(&template); - block_stds.extend(block_stds_of(&interior)); + let interior_block_stds = block_stds_of(&interior); + block_stds.extend(interior_block_stds); samples.extend(interior); } @@ -89,8 +89,10 @@ fn interior_of(template: &[i32]) -> Vec { for y in INTERIOR_START..INTERIOR_START + INTERIOR_SIZE { let row_start = y * TEMPLATE_COLS + INTERIOR_START; let row = &template[row_start..row_start + INTERIOR_SIZE]; - interior.extend(row.iter().map(|&value| value as f64)); + let row_values = row.iter().map(|&value| value as f64); + interior.extend(row_values); } + interior } @@ -106,9 +108,11 @@ fn block_stds_of(interior: &[f64]) -> Vec { block.extend_from_slice(&interior[start..start + cell]); } - stds.push(sample_std(&block)); + let block_std = sample_std(&block); + stds.push(block_std); } } + stds } diff --git a/av-denoise-core/src/nl4d/grain/tests/export.rs b/av-denoise-core/src/nl4d/grain/tests/export.rs index 09029e7..201ce68 100644 --- a/av-denoise-core/src/nl4d/grain/tests/export.rs +++ b/av-denoise-core/src/nl4d/grain/tests/export.rs @@ -4,6 +4,7 @@ use cubecl::prelude::*; use cubecl::wgpu::WgpuRuntime; use super::synthetic::gaussian_field; +use crate::bench_api::HostIo; use crate::nl4d::grain::chunk::GrainChunk; use crate::nl4d::grain::consts::{CHUNK_FRAMES, HIST_LEN, LUMA_BINS, STD_BUCKETS, STRENGTH_GROUPS}; use crate::nl4d::grain::fit::{bucket_edges, hist_median}; @@ -40,24 +41,28 @@ fn make_client() -> ComputeClient { } fn params(grain_export: bool) -> Nl4dParams { + let motion_compensation = MotionCompensationMode::Mvtools { + blksize: 16, + overlap: 8, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Auto, + }; + let hq = HqParams::with_sigma(6.0 / 255.0); + let nlm = NlmParams { + temporal_radius: 2, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation, + hq: Some(hq), + }; + Nl4dParams { - nlm: NlmParams { - temporal_radius: 2, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 16, - overlap: 8, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Auto, - }, - hq: Some(HqParams::with_sigma(6.0 / 255.0)), - }, + nlm, temporal_radius: 2, grain_export, ..Nl4dParams::default() @@ -106,14 +111,17 @@ fn panning_frames(sigma: f32, count: u32) -> Vec> { fn denoise_stream(denoiser: &mut Nl4dDenoiser, frames: &[Vec], outputs: &mut Vec>) { for frame in frames { denoiser.push_frame(frame); - if let Some(pending) = denoiser.denoise_submit().expect("submit") { - let output = pending.wait().expect("readback"); - outputs.push(output.into_f32().expect("f32 output")); + + if let Some(output) = denoiser.denoise().expect("denoise") { + outputs.push(output); } } denoiser - .flush(|frame| outputs.push(frame.as_f32().expect("f32 output").to_vec())) + .flush(|frame| { + let output = frame.to_vec(); + outputs.push(output); + }) .expect("flush"); } @@ -135,6 +143,7 @@ fn run_sized( denoise_stream(&mut denoiser, frames, &mut outputs); let chunks = denoiser.drain_grain_chunks().expect("drain"); + (outputs, chunks, has_export) } @@ -171,6 +180,7 @@ fn lag_3_ratio(chunk: &GrainChunk) -> f64 { } assert!(zero_lag > 0.0); + lag_3 / zero_lag } @@ -181,7 +191,8 @@ fn frames_counted(chunks: &[GrainChunk]) -> u32 { #[test] fn export_off_allocates_nothing_and_drains_nothing() { let frames = grain_frames(SIGMA, 9); - let (_, chunks, has_export) = run(params(false), &frames); + let export_off = params(false); + let (_, chunks, has_export) = run(export_off, &frames); assert!(!has_export); assert!(chunks.is_empty()); @@ -190,8 +201,10 @@ fn export_off_allocates_nothing_and_drains_nothing() { #[test] fn export_does_not_change_the_output() { let frames = grain_frames(SIGMA, 9); - let (off, _, _) = run(params(false), &frames); - let (on, _, _) = run(params(true), &frames); + let export_off = params(false); + let export_on = params(true); + let (off, _, _) = run(export_off, &frames); + let (on, _, _) = run(export_on, &frames); assert_eq!(off, on); } @@ -199,21 +212,22 @@ fn export_does_not_change_the_output() { #[test] fn export_measures_the_source_grain() { let frames = grain_frames(SIGMA, 9); - let (_, chunks, has_export) = run(params(true), &frames); + let export_on = params(true); + let (_, chunks, has_export) = run(export_on, &frames); let merged = merged(&chunks); let median = median_of(&merged.source_hist); + let counted = frames_counted(&chunks); + let any_pixels = chunks + .iter() + .any(|chunk| chunk.pixels.iter().any(|&pixels| pixels > 0.0)); assert!(has_export); - assert_eq!(frames_counted(&chunks), 9); + assert_eq!(counted, 9); assert!( (median / SIGMA as f64 - 1.0).abs() < 0.1, "median {median} vs sigma {SIGMA}" ); - assert!( - chunks - .iter() - .any(|chunk| chunk.pixels.iter().any(|&pixels| pixels > 0.0)) - ); + assert!(any_pixels); } #[test] @@ -225,7 +239,8 @@ fn flicker_keeps_the_strength_and_a_short_texture() { } } - let (_, chunks, _) = run(params(true), &frames); + let export_on = params(true); + let (_, chunks, _) = run(export_on, &frames); let merged = merged(&chunks); let median = median_of(&merged.source_hist); let ratio = lag_3_ratio(&merged); @@ -240,7 +255,8 @@ fn flicker_keeps_the_strength_and_a_short_texture() { #[test] fn panning_content_measures_the_source_grain() { let frames = panning_frames(SIGMA, 9); - let (outputs, chunks, _) = run_sized(params(true), &frames, PAN_WIDTH, PAN_HEIGHT); + let export_on = params(true); + let (outputs, chunks, _) = run_sized(export_on, &frames, PAN_WIDTH, PAN_HEIGHT); let merged = merged(&chunks); let accepted: u32 = merged.source_hist.iter().sum(); let median = median_of(&merged.source_hist); @@ -256,7 +272,8 @@ fn panning_content_measures_the_source_grain() { #[test] fn kept_grain_is_weaker_than_source_grain() { let frames = grain_frames(SIGMA, 9); - let (_, chunks, _) = run(params(true), &frames); + let export_on = params(true); + let (_, chunks, _) = run(export_on, &frames); let merged = merged(&chunks); let kept_total: u32 = merged.kept_hist.iter().sum(); let source_total: u32 = merged.source_hist.iter().sum(); @@ -276,16 +293,18 @@ fn kept_grain_is_weaker_than_source_grain() { #[test] fn short_scene_measures_only_real_pairs() { let client = make_client(); - let mut denoiser = Nl4dDenoiser::::new(&client, params(true), WIDTH, HEIGHT).expect("construction"); + let export_on = params(true); + let mut denoiser = Nl4dDenoiser::::new(&client, export_on, WIDTH, HEIGHT).expect("construction"); let frames = grain_frames(SIGMA, 3); let mut outputs = Vec::new(); denoise_stream(&mut denoiser, &frames, &mut outputs); let chunks = denoiser.drain_grain_chunks().expect("drain"); + let counted = frames_counted(&chunks); assert_eq!(outputs.len(), 3); - assert_eq!(frames_counted(&chunks), 3); + assert_eq!(counted, 3); assert_eq!(denoiser.grain_measured_with_entry(), 2); } @@ -293,17 +312,19 @@ fn short_scene_measures_only_real_pairs() { fn fused_yuv_measures_the_luma_grain() { let mut yuv = params(true); yuv.nlm.channels = ChannelMode::Yuv; - let frames: Vec> = grain_frames(SIGMA, 9) + let luma_frames = grain_frames(SIGMA, 9); + let frames: Vec> = luma_frames .into_iter() .map(|frame| frame.iter().flat_map(|&luma| [luma, 0.5, 0.5]).collect()) .collect(); let (outputs, chunks, has_export) = run(yuv, &frames); let merged = merged(&chunks); let median = median_of(&merged.source_hist); + let counted = frames_counted(&chunks); assert!(has_export); assert_eq!(outputs.len(), 9); - assert_eq!(frames_counted(&chunks), 9); + assert_eq!(counted, 9); assert!( (median / SIGMA as f64 - 1.0).abs() < 0.1, "median {median} vs sigma {SIGMA}" @@ -314,7 +335,8 @@ fn fused_yuv_measures_the_luma_grain() { fn chroma_denoisers_never_export() { let mut chroma = params(true); chroma.nlm.channels = ChannelMode::Chroma; - let frames: Vec> = grain_frames(SIGMA, 9) + let plane_frames = grain_frames(SIGMA, 9); + let frames: Vec> = plane_frames .into_iter() .map(|frame| frame.iter().flat_map(|&sample| [sample, sample]).collect()) .collect(); @@ -327,7 +349,8 @@ fn chroma_denoisers_never_export() { #[test] fn chunks_close_when_full_and_at_each_stream_end() { let client = make_client(); - let mut denoiser = Nl4dDenoiser::::new(&client, params(true), WIDTH, HEIGHT).expect("construction"); + let export_on = params(true); + let mut denoiser = Nl4dDenoiser::::new(&client, export_on, WIDTH, HEIGHT).expect("construction"); let long_stream = grain_frames(SIGMA, 30); let short_stream = grain_frames(SIGMA, 5); let mut outputs = Vec::new(); diff --git a/av-denoise-core/src/nl4d/grain/tests/fit.rs b/av-denoise-core/src/nl4d/grain/tests/fit.rs index c78756e..1e9f104 100644 --- a/av-denoise-core/src/nl4d/grain/tests/fit.rs +++ b/av-denoise-core/src/nl4d/grain/tests/fit.rs @@ -39,6 +39,7 @@ fn scene_eight_weights() -> [f64; AR_COEFFS] { set_weight(&mut weights, (-2, 0), -0.148); set_weight(&mut weights, (-1, -1), -0.164); set_weight(&mut weights, (-1, 1), 0.117); + weights } @@ -57,10 +58,13 @@ fn bucket_edges_are_log_spaced_between_the_limits() { #[test] fn bucket_of_clamps_to_the_last_bucket() { let edges = bucket_edges(); + let lowest = bucket_of(STD_MIN, &edges); + let far_above = bucket_of(STD_MAX * 4.0, &edges); + let just_past_edge = bucket_of(edges[10] * 1.0001, &edges); - assert_eq!(bucket_of(STD_MIN, &edges), 0); - assert_eq!(bucket_of(STD_MAX * 4.0, &edges), STD_BUCKETS - 1); - assert_eq!(bucket_of(edges[10] * 1.0001, &edges), 10); + assert_eq!(lowest, 0); + assert_eq!(far_above, STD_BUCKETS - 1); + assert_eq!(just_past_edge, 10); } #[test] @@ -88,8 +92,9 @@ fn histogram_median_is_within_one_bucket_of_the_exact_median() { fn histogram_median_of_an_empty_histogram_is_none() { let edges = bucket_edges(); let counts = vec![0u32; STD_BUCKETS]; + let median = hist_median(&counts, &edges); - assert_eq!(hist_median(&counts, &edges), None); + assert_eq!(median, None); } #[test] @@ -107,8 +112,9 @@ fn yule_walker_recovers_known_weights() { #[test] fn yule_walker_rejects_an_empty_record() { let autocov = vec![0.0f64; LAG_COUNT + 1]; + let solved = yule_walker(&autocov); - assert!(yule_walker(&autocov).is_none()); + assert!(solved.is_none()); } #[test] @@ -141,7 +147,8 @@ fn calibration_recovers_a_known_sigma() { block.extend_from_slice(&field[start..start + 8]); } - block_stds.push(sample_std(&block)); + let block_std = sample_std(&block); + block_stds.push(block_std); } } @@ -201,12 +208,18 @@ fn correlation_of(record: &[f64]) -> Vec { #[test] fn measured_lanes_follow_the_lag_order() { for (lane, &(dy, dx)) in LAGS.iter().enumerate() { - assert_eq!(measured_lane(dy, dx), Some(lane)); - assert_eq!(measured_lane(-dy, -dx), Some(lane)); + let forward = measured_lane(dy, dx); + let mirrored = measured_lane(-dy, -dx); + + assert_eq!(forward, Some(lane)); + assert_eq!(mirrored, Some(lane)); } - assert_eq!(measured_lane(4, 0), None); - assert_eq!(measured_lane(0, 7), None); + let too_far_down = measured_lane(4, 0); + let too_far_right = measured_lane(0, 7); + + assert_eq!(too_far_down, None); + assert_eq!(too_far_right, None); } #[test] diff --git a/av-denoise-core/src/nl4d/grain/tests/kernels.rs b/av-denoise-core/src/nl4d/grain/tests/kernels.rs index e1ca75f..71f97a4 100644 --- a/av-denoise-core/src/nl4d/grain/tests/kernels.rs +++ b/av-denoise-core/src/nl4d/grain/tests/kernels.rs @@ -55,19 +55,25 @@ fn save_vectors_copies_one_neighbour_into_its_entry() { let saved_mv_zeros = vec![0i32; saved_mv_len]; let saved_conf_zeros = vec![0.0f32; saved_conf_len]; - let mv = client.create_from_slice(i32::as_bytes(&mv_host)); - let conf = client.create_from_slice(f32::as_bytes(&conf_host)); - let saved_mv = client.create_from_slice(i32::as_bytes(&saved_mv_zeros)); - let saved_conf = client.create_from_slice(f32::as_bytes(&saved_conf_zeros)); + let mv_bytes = i32::as_bytes(&mv_host); + let conf_bytes = f32::as_bytes(&conf_host); + let saved_mv_zero_bytes = i32::as_bytes(&saved_mv_zeros); + let saved_conf_zero_bytes = f32::as_bytes(&saved_conf_zeros); + let mv = client.create_from_slice(mv_bytes); + let conf = client.create_from_slice(conf_bytes); + let saved_mv = client.create_from_slice(saved_mv_zero_bytes); + let saved_conf = client.create_from_slice(saved_conf_zero_bytes); let neighbour = 2u32; let entry = 1u32; + let grid = CubeCount::new_1d(1); + let dim = CubeDim::new_1d(THREADS); unsafe { grain_save_vectors::launch_unchecked::( &client, - CubeCount::new_1d(1), - CubeDim::new_1d(THREADS), + grid, + dim, ArrayArg::from_raw_parts(mv, mv_host.len()), ArrayArg::from_raw_parts(conf, conf_host.len()), ArrayArg::from_raw_parts(saved_mv.clone(), saved_mv_len), @@ -234,6 +240,7 @@ fn with_kept_grain(mut scene: Scene, kept_shift: (i32, i32)) -> Scene { scene.out_t = add_noise(&moved_prev, KEPT_SIGMA, 4, shape); scene.out_prev = out_prev; scene.kept_mv.fill(kept_shift); + scene } @@ -281,20 +288,32 @@ fn run_measure(scene: &Scene, has_source: bool, has_kept: bool) -> (Vec, Ve let edges = bucket_edges(); let hist_zeros = vec![0i32; 2 * HIST_LEN]; - let input = client.create_from_slice(f32::as_bytes(&ring)); - let out_t = client.create_from_slice(f32::as_bytes(&scene.out_t)); - let out_prev = client.create_from_slice(f32::as_bytes(&scene.out_prev)); - let saved_mv = client.create_from_slice(i32::as_bytes(&saved_mv)); - let saved_conf = client.create_from_slice(f32::as_bytes(&saved_conf)); - let edges_buf = client.create_from_slice(f32::as_bytes(&edges)); - let hist = client.create_from_slice(i32::as_bytes(&hist_zeros)); + let ring_bytes = f32::as_bytes(&ring); + let out_t_bytes = f32::as_bytes(&scene.out_t); + let out_prev_bytes = f32::as_bytes(&scene.out_prev); + let saved_mv_bytes = i32::as_bytes(&saved_mv); + let saved_conf_bytes = f32::as_bytes(&saved_conf); + let edges_bytes = f32::as_bytes(&edges); + let hist_zero_bytes = i32::as_bytes(&hist_zeros); + let input = client.create_from_slice(ring_bytes); + let out_t = client.create_from_slice(out_t_bytes); + let out_prev = client.create_from_slice(out_prev_bytes); + let saved_mv = client.create_from_slice(saved_mv_bytes); + let saved_conf = client.create_from_slice(saved_conf_bytes); + let edges_buf = client.create_from_slice(edges_bytes); + let hist = client.create_from_slice(hist_zero_bytes); let partials = client.empty(cells * PARTIAL_LEN * size_of::()); + let grid = CubeCount::new_2d(shape.width / 8, shape.height / 8); + let dim = CubeDim::new_2d(8, 8); + let blocks_x = shape.blocks_x(); + let blocks_y = shape.blocks_y(); + unsafe { grain_measure::launch_unchecked::( &client, - CubeCount::new_2d(shape.width / 8, shape.height / 8), - CubeDim::new_2d(8, 8), + grid, + dim, 1usize, ArrayArg::from_raw_parts(input, 2 * pixels), ArrayArg::from_raw_parts(out_t, pixels), @@ -313,8 +332,8 @@ fn run_measure(scene: &Scene, has_source: bool, has_kept: bool) -> (Vec, Ve shape.width, shape.height, 1u32, - shape.blocks_x(), - shape.blocks_y(), + blocks_x, + blocks_y, shape.step, ); } @@ -380,9 +399,12 @@ fn assert_matches_mirror(scene: &Scene) -> Vec { let (gpu_hist, gpu_autocov) = run_measure(scene, true, true); let (host_hist, host_autocov) = mirror_of(scene, true, true); + let pixels = pixels_of(&gpu_autocov); + assert_eq!(gpu_hist, host_hist); assert_autocov_close(&gpu_autocov, &host_autocov); - assert!(pixels_of(&gpu_autocov) > 0.0); + assert!(pixels > 0.0); + gpu_hist } @@ -403,6 +425,7 @@ fn lag_3_ratio(autocov: &[f64]) -> f64 { } assert!(zero_lag > 0.0); + lag_3 / zero_lag } @@ -429,41 +452,48 @@ fn fill_cell_row(scene: &mut Scene, cell_y: u32, value: f32) { fn measure_matches_the_mirror_on_flat_grain() { let scene = flat_scene((0, 0), &[]); let hist = assert_matches_mirror(&scene); + let source = source_total(&hist); - assert!(source_total(&hist) > 0); + assert!(source > 0); } #[test] fn measure_matches_the_mirror_on_kept_grain() { let scene = kept_scene(PLAIN); let hist = assert_matches_mirror(&scene); + let source = source_total(&hist); + let kept = kept_total(&hist); - assert!(source_total(&hist) > 0); - assert!(kept_total(&hist) > 0); + assert!(source > 0); + assert!(kept > 0); } #[test] fn measure_matches_the_mirror_on_a_ragged_frame() { let scene = kept_scene(RAGGED); let hist = assert_matches_mirror(&scene); + let source = source_total(&hist); + let kept = kept_total(&hist); - assert!(source_total(&hist) > 0); - assert!(kept_total(&hist) > 0); + assert!(source > 0); + assert!(kept > 0); } #[test] fn flicker_leaves_the_record_unchanged() { let steady = flat_scene((0, 0), &[]); - let flickered = with_flicker(flat_scene((0, 0), &[])); + let unflickered = flat_scene((0, 0), &[]); + let flickered = with_flicker(unflickered); let (steady_hist, steady_autocov) = run_measure(&steady, true, false); let (flicker_hist, flicker_autocov) = run_measure(&flickered, true, false); let (host_hist, host_autocov) = mirror_of(&flickered, true, false); let steady_ratio = lag_3_ratio(&steady_autocov); let flicker_ratio = lag_3_ratio(&flicker_autocov); + let steady_source = source_total(&steady_hist); assert_eq!(flicker_hist, host_hist); assert_autocov_close(&flicker_autocov, &host_autocov); - assert!(source_total(&steady_hist) > 0); + assert!(steady_source > 0); assert_eq!(flicker_hist, steady_hist); assert!( (flicker_ratio - steady_ratio).abs() < 0.02, @@ -476,9 +506,11 @@ fn flicker_leaves_the_record_unchanged() { fn kept_grain_is_not_counted_without_a_kept_pair() { let scene = kept_scene(PLAIN); let (hist, _) = run_measure(&scene, true, false); + let source = source_total(&hist); + let kept = kept_total(&hist); - assert!(source_total(&hist) > 0); - assert_eq!(kept_total(&hist), 0); + assert!(source > 0); + assert_eq!(kept, 0); } #[test] @@ -488,9 +520,13 @@ fn a_low_kept_confidence_cell_is_rejected() { let block = PLAIN.block_at(3, 3); scene.kept_conf[block] = 0.5; let (after, _) = run_measure(&scene, true, true); + let kept_before = kept_total(&before); + let kept_after = kept_total(&after); + let source_before = source_total(&before); + let source_after = source_total(&after); - assert_eq!(kept_total(&after) + 1, kept_total(&before)); - assert_eq!(source_total(&after), source_total(&before)); + assert_eq!(kept_after + 1, kept_before); + assert_eq!(source_after, source_before); } #[test] @@ -523,8 +559,10 @@ fn a_textured_cell_is_rejected() { let textured = flat_scene((0, 0), &[(3, 3)]); let (plain_hist, _) = run_measure(&plain, true, false); let (textured_hist, _) = run_measure(&textured, true, false); + let plain_source = source_total(&plain_hist); + let textured_source = source_total(&textured_hist); - assert_eq!(source_total(&textured_hist) + 1, source_total(&plain_hist)); + assert_eq!(textured_source + 1, plain_source); } #[test] @@ -542,8 +580,10 @@ fn a_cell_without_grain_is_rejected() { let (plain_hist, _) = run_measure(&plain, true, false); let (empty_hist, _) = run_measure(&empty, true, false); + let plain_source = source_total(&plain_hist); + let empty_source = source_total(&empty_hist); - assert_eq!(source_total(&empty_hist) + 1, source_total(&plain_hist)); + assert_eq!(empty_source + 1, plain_source); } #[test] @@ -553,8 +593,10 @@ fn a_low_confidence_cell_is_rejected() { let block = PLAIN.block_at(3, 3); scene.source_conf[block] = 0.5; let (after, _) = run_measure(&scene, true, false); + let source_before = source_total(&before); + let source_after = source_total(&after); - assert_eq!(source_total(&after) + 1, source_total(&before)); + assert_eq!(source_after + 1, source_before); } #[test] @@ -577,12 +619,11 @@ fn clipped_blocks_are_rejected() { let (gpu_hist, _) = run_measure(&clipped, true, false); let (host_hist, _) = mirror_of(&clipped, true, false); let interior_per_row = PLAIN.cells_x() - 2; + let plain_source = source_total(&plain_hist); + let clipped_source = source_total(&gpu_hist); assert_eq!(gpu_hist, host_hist); - assert_eq!( - source_total(&gpu_hist) + 2 * interior_per_row, - source_total(&plain_hist) - ); + assert_eq!(clipped_source + 2 * interior_per_row, plain_source); } #[test] @@ -602,9 +643,11 @@ fn ten_bit_and_eight_bit_give_the_same_record() { let (eight_hist, eight_autocov) = run_measure(&eight, true, false); let (ten_hist, ten_autocov) = run_measure(&ten, true, false); + let eight_source = source_total(&eight_hist); + let ten_source = source_total(&ten_hist); - assert!(source_total(&eight_hist) > 0); - assert_eq!(source_total(&eight_hist), source_total(&ten_hist)); + assert!(eight_source > 0); + assert_eq!(eight_source, ten_source); assert_within_one_bucket(&eight_hist[..HIST_LEN], &ten_hist[..HIST_LEN]); assert_groups_close(&eight_autocov, &ten_autocov, 0.05); } @@ -709,14 +752,19 @@ fn reduce_adds_every_cell_into_its_group() { } let start: Vec = (0..GROUPED_AUTOCOV_LEN).map(|lane| lane as f32).collect(); - let partials = client.create_from_slice(f32::as_bytes(&partials_host)); - let chunk = client.create_from_slice(f32::as_bytes(&start)); + let partials_bytes = f32::as_bytes(&partials_host); + let start_bytes = f32::as_bytes(&start); + let partials = client.create_from_slice(partials_bytes); + let chunk = client.create_from_slice(start_bytes); + + let grid = CubeCount::new_1d(AUTOCOV_LEN as u32); + let dim = CubeDim::new_1d(REDUCE_THREADS); unsafe { grain_reduce_partials::launch_unchecked::( &client, - CubeCount::new_1d(AUTOCOV_LEN as u32), - CubeDim::new_1d(REDUCE_THREADS), + grid, + dim, ArrayArg::from_raw_parts(partials, partials_host.len()), ArrayArg::from_raw_parts(chunk.clone(), GROUPED_AUTOCOV_LEN), cells as u32, diff --git a/av-denoise-core/src/nl4d/grain/tests/mirror.rs b/av-denoise-core/src/nl4d/grain/tests/mirror.rs index f51a1eb..6a426cc 100644 --- a/av-denoise-core/src/nl4d/grain/tests/mirror.rs +++ b/av-denoise-core/src/nl4d/grain/tests/mirror.rs @@ -103,9 +103,9 @@ fn hist_slot(mean: f32, std: f32, edges: &[f32]) -> usize { /// Adds the cell's lag products to `record`, each pixel against its neighbour at every lag. /// /// `mean` is the cell's mean grain, taken off every pixel and neighbour first. -fn add_lag_sums(frame: &MirrorFrame, x0: u32, y0: u32, mean: f32, record: &mut [f64]) { - for y in y0..y0 + CELL { - for x in x0..x0 + CELL { +fn add_lag_sums(frame: &MirrorFrame, cell_left: u32, cell_top: u32, mean: f32, record: &mut [f64]) { + for y in cell_top..cell_top + CELL { + for x in cell_left..cell_left + CELL { let centre = (source_grain(frame, x, y) - mean) as f64; for (lane, &(dy, dx)) in LAGS.iter().enumerate() { @@ -129,20 +129,21 @@ pub(super) fn mirror_measure(frame: &MirrorFrame, edges: &[f32]) -> MirrorRecord for cell_y in 0..cells_y { for cell_x in 0..cells_x { - let x0 = cell_x * CELL; - let y0 = cell_y * CELL; + let cell_left = cell_x * CELL; + let cell_top = cell_y * CELL; let block = block_of(frame, cell_x, cell_y); let pixels = (CELL * CELL) as usize; let mut grain = Vec::with_capacity(pixels); let mut clean = Vec::with_capacity(pixels); let mut noisy = Vec::with_capacity(2 * pixels); let mut kept = Vec::with_capacity(pixels); - let mut prev = Vec::with_capacity(pixels); + let mut previous_values = Vec::with_capacity(pixels); - for y in y0..y0 + CELL { - for x in x0..x0 + CELL { + for y in cell_top..cell_top + CELL { + for x in cell_left..cell_left + CELL { let index = (y * frame.width + x) as usize; - grain.push(source_grain(frame, x, y)); + let pixel_grain = source_grain(frame, x, y); + grain.push(pixel_grain); clean.push(frame.out_t[index]); noisy.push(frame.source_t[index]); @@ -153,7 +154,7 @@ pub(super) fn mirror_measure(frame: &MirrorFrame, edges: &[f32]) -> MirrorRecord let warped = frame.out_t[kept_index]; let previous = frame.out_prev[index]; kept.push((warped - previous) * std::f32::consts::FRAC_1_SQRT_2); - prev.push(previous); + previous_values.push(previous); } } @@ -177,22 +178,22 @@ pub(super) fn mirror_measure(frame: &MirrorFrame, edges: &[f32]) -> MirrorRecord let group = bucket_of(grain_stats.std, edges) / BUCKETS_PER_GROUP; let record = &mut autocov[group * AUTOCOV_LEN..(group + 1) * AUTOCOV_LEN]; - add_lag_sums(frame, x0, y0, grain_stats.mean, record); + add_lag_sums(frame, cell_left, cell_top, grain_stats.mean, record); } let kept_stats = stats_of(&kept); - let prev_stats = stats_of(&prev); - let (prev_low, prev_high) = range_of(&prev); + let previous_stats = stats_of(&previous_values); + let (previous_low, previous_high) = range_of(&previous_values); let kept_ok = frame.has_kept && frame.kept_conf[block] >= CONF_MIN - && prev_high - prev_low < FLAT_RANGE - && prev_stats.mean > LUMA_LOW - && prev_stats.mean < LUMA_HIGH - && prev_low >= CLIP_LOW - && prev_high <= CLIP_HIGH + && previous_high - previous_low < FLAT_RANGE + && previous_stats.mean > LUMA_LOW + && previous_stats.mean < LUMA_HIGH + && previous_low >= CLIP_LOW + && previous_high <= CLIP_HIGH && kept_stats.std > STD_MIN; if kept_ok { - let slot = hist_slot(prev_stats.mean, kept_stats.std, edges); + let slot = hist_slot(previous_stats.mean, kept_stats.std, edges); hist[HIST_LEN + slot] += 1; } } diff --git a/av-denoise-core/src/nl4d/grain/tests/segment.rs b/av-denoise-core/src/nl4d/grain/tests/segment.rs index 0ed4405..b6cd01f 100644 --- a/av-denoise-core/src/nl4d/grain/tests/segment.rs +++ b/av-denoise-core/src/nl4d/grain/tests/segment.rs @@ -17,6 +17,7 @@ const EDGE_BUCKET: usize = 37; fn sparse_chunk() -> GrainChunk { let mut chunk = GrainChunk::empty(); chunk.frames = CHUNK_FRAMES; + chunk } @@ -35,6 +36,7 @@ fn narrow_chunk(std_codes: f32) -> GrainChunk { let record = grain_record(); add_group_record(&mut chunk, bucket / BUCKETS_PER_GROUP, &record); + chunk } @@ -85,16 +87,16 @@ fn edge_chunk() -> (GrainChunk, usize) { let edge_bucket = bucket_of(band_high as f32, &edges); assert_eq!(edge_bucket % BUCKETS_PER_GROUP, 0); + (chunk, edge_bucket / BUCKETS_PER_GROUP) } #[test] fn steady_scene_is_one_segment() { - let chunks = vec![ - chunk_at(2.0, 300, None), - chunk_at(2.05, 300, None), - chunk_at(1.98, 300, None), - ]; + let first = chunk_at(2.0, 300, None); + let second = chunk_at(2.05, 300, None); + let third = chunk_at(1.98, 300, None); + let chunks = vec![first, second, third]; let scene = scene_of(100, chunks); let segments = segment_scene(&scene); @@ -106,11 +108,10 @@ fn steady_scene_is_one_segment() { #[test] fn drift_splits_on_a_chunk_boundary() { - let chunks = vec![ - chunk_at(2.0, 300, None), - chunk_at(2.0, 300, None), - chunk_at(2.8, 300, None), - ]; + let first = chunk_at(2.0, 300, None); + let second = chunk_at(2.0, 300, None); + let drifted = chunk_at(2.8, 300, None); + let chunks = vec![first, second, drifted]; let scene = scene_of(0, chunks); let segments = segment_scene(&scene); @@ -121,7 +122,9 @@ fn drift_splits_on_a_chunk_boundary() { #[test] fn a_thin_chunk_joins_the_open_segment() { - let chunks = vec![chunk_at(2.0, 300, None), chunk_at(4.0, 10, None)]; + let full = chunk_at(2.0, 300, None); + let thin = chunk_at(4.0, 10, None); + let chunks = vec![full, thin]; let scene = scene_of(0, chunks); let segments = segment_scene(&scene); @@ -131,7 +134,9 @@ fn a_thin_chunk_joins_the_open_segment() { #[test] fn a_fitted_entry_has_points_and_weights() { - let scenes = vec![scene_of(0, vec![chunk_at(2.0, 300, None)])]; + let chunk = chunk_at(2.0, 300, None); + let scene = scene_of(0, vec![chunk]); + let scenes = vec![scene]; let entries = fit_scenes(&scenes); @@ -142,7 +147,9 @@ fn a_fitted_entry_has_points_and_weights() { #[test] fn kept_equal_to_source_gives_zero_strength() { - let scenes = vec![scene_of(0, vec![chunk_at(2.0, 300, Some(2.0))])]; + let chunk = chunk_at(2.0, 300, Some(2.0)); + let scene = scene_of(0, vec![chunk]); + let scenes = vec![scene]; let entries = fit_scenes(&scenes); @@ -151,7 +158,9 @@ fn kept_equal_to_source_gives_zero_strength() { #[test] fn kept_above_source_clamps_to_zero() { - let scenes = vec![scene_of(0, vec![chunk_at(2.0, 300, Some(3.0))])]; + let chunk = chunk_at(2.0, 300, Some(3.0)); + let scene = scene_of(0, vec![chunk]); + let scenes = vec![scene]; let entries = fit_scenes(&scenes); @@ -160,11 +169,14 @@ fn kept_above_source_clamps_to_zero() { #[test] fn a_thin_chunk_never_gets_its_own_entry() { - let second_chunks = vec![chunk_at(3.0, 300, None), sparse_chunk(), chunk_at(3.0, 300, None)]; - let scenes = vec![ - scene_of(0, vec![chunk_at(1.0, 300, None)]), - scene_of(24, second_chunks), - ]; + let first_chunk = chunk_at(1.0, 300, None); + let before_thin = chunk_at(3.0, 300, None); + let thin = sparse_chunk(); + let after_thin = chunk_at(3.0, 300, None); + let second_chunks = vec![before_thin, thin, after_thin]; + let first_scene = scene_of(0, vec![first_chunk]); + let second_scene = scene_of(24, second_chunks); + let scenes = vec![first_scene, second_scene]; let entries = fit_scenes(&scenes); @@ -177,10 +189,10 @@ fn a_thin_chunk_never_gets_its_own_entry() { fn missing_texture_borrows_from_the_nearest_segment() { let mut no_texture = chunk_at(2.0, 300, None); no_texture.pixels.fill(0.0); - let scenes = vec![ - scene_of(0, vec![chunk_at(2.0, 300, None)]), - scene_of(24, vec![no_texture]), - ]; + let textured = chunk_at(2.0, 300, None); + let first_scene = scene_of(0, vec![textured]); + let second_scene = scene_of(24, vec![no_texture]); + let scenes = vec![first_scene, second_scene]; let entries = fit_scenes(&scenes); @@ -190,10 +202,11 @@ fn missing_texture_borrows_from_the_nearest_segment() { #[test] fn sparse_scene_borrows_from_a_neighbouring_scene() { - let scenes = vec![ - scene_of(0, vec![chunk_at(2.0, 300, None)]), - scene_of(24, vec![sparse_chunk()]), - ]; + let dense = chunk_at(2.0, 300, None); + let sparse = sparse_chunk(); + let first_scene = scene_of(0, vec![dense]); + let second_scene = scene_of(24, vec![sparse]); + let scenes = vec![first_scene, second_scene]; let entries = fit_scenes(&scenes); @@ -203,7 +216,9 @@ fn sparse_scene_borrows_from_a_neighbouring_scene() { #[test] fn sparse_segment_with_no_donor_gets_no_entry() { - let scenes = vec![scene_of(0, vec![sparse_chunk()])]; + let sparse = sparse_chunk(); + let scene = scene_of(0, vec![sparse]); + let scenes = vec![scene]; let entries = fit_scenes(&scenes); @@ -225,11 +240,15 @@ fn an_outlier_group_leaves_the_texture_unchanged() { let mut mixed = grain_only.clone(); add_group_record(&mut mixed, grain_group, &coarse); - let clean_entries = fit_scenes(&[scene_of(0, vec![grain_only])]); - let outlier_entries = fit_scenes(&[scene_of(0, vec![with_outlier])]); - let mixed_entries = fit_scenes(&[scene_of(0, vec![mixed])]); + let clean_scene = scene_of(0, vec![grain_only]); + let outlier_scene = scene_of(0, vec![with_outlier]); + let mixed_scene = scene_of(0, vec![mixed]); + let clean_entries = fit_scenes(&[clean_scene]); + let outlier_entries = fit_scenes(&[outlier_scene]); + let mixed_entries = fit_scenes(&[mixed_scene]); + let grain = grain_record(); - assert!(coarse[0] > 10.0 * grain_record()[0]); + assert!(coarse[0] > 10.0 * grain[0]); assert!(clean_entries[0].ar_coeffs.iter().any(|&coeff| coeff != 0)); assert_eq!(outlier_entries[0].ar_coeffs, clean_entries[0].ar_coeffs); assert_ne!(mixed_entries[0].ar_coeffs, clean_entries[0].ar_coeffs); @@ -241,7 +260,8 @@ fn a_group_overlapping_the_band_edge_gives_the_texture() { let record = grain_record(); add_group_record(&mut chunk, edge_group, &record); - let entries = fit_scenes(&[scene_of(0, vec![chunk])]); + let scene = scene_of(0, vec![chunk]); + let entries = fit_scenes(&[scene]); assert_eq!(entries.len(), 1); assert!(entries[0].ar_coeffs.iter().any(|&coeff| coeff != 0)); @@ -253,18 +273,22 @@ fn a_group_past_the_band_gives_no_texture() { let record = grain_record(); add_group_record(&mut chunk, edge_group + 1, &record); - let entries = fit_scenes(&[scene_of(0, vec![chunk])]); + let scene = scene_of(0, vec![chunk]); + let entries = fit_scenes(&[scene]); assert!(entries.is_empty()); } #[test] fn a_sparse_segment_borrows_from_its_own_scene_first() { - let own_scene = vec![chunk_at(2.0, 300, None), narrow_chunk(3.0), narrow_chunk(4.0)]; - let scenes = vec![ - scene_of(0, own_scene), - scene_of(3 * CHUNK_FRAMES as u64, vec![chunk_at(1.0, 300, None)]), - ]; + let dense = chunk_at(2.0, 300, None); + let first_narrow = narrow_chunk(3.0); + let second_narrow = narrow_chunk(4.0); + let next_scene_chunk = chunk_at(1.0, 300, None); + let own_scene = vec![dense, first_narrow, second_narrow]; + let first_scene = scene_of(0, own_scene); + let second_scene = scene_of(3 * CHUNK_FRAMES as u64, vec![next_scene_chunk]); + let scenes = vec![first_scene, second_scene]; let entries = fit_scenes(&scenes); @@ -283,7 +307,8 @@ fn a_chunk_in_every_luma_bin_thins_to_the_point_limit() { chunk.source_hist[bin * STD_BUCKETS + bucket] = 300; } - let entries = fit_scenes(&[scene_of(0, vec![chunk])]); + let scene = scene_of(0, vec![chunk]); + let entries = fit_scenes(&[scene]); let points = &entries[0].points; let first_luma = points.first().expect("points").0; let last_luma = points.last().expect("points").0; diff --git a/av-denoise-core/src/nl4d/grain/tests/synthetic.rs b/av-denoise-core/src/nl4d/grain/tests/synthetic.rs index 04ed553..bf25984 100644 --- a/av-denoise-core/src/nl4d/grain/tests/synthetic.rs +++ b/av-denoise-core/src/nl4d/grain/tests/synthetic.rs @@ -40,6 +40,7 @@ pub(super) fn gaussian_field(width: usize, height: usize, seed: u64) -> Vec } field.truncate(width * height); + field } @@ -70,6 +71,7 @@ pub(super) fn ar_field(coeffs: &[f64; AR_COEFFS], width: usize, height: usize, s let start = y * full_width + warm; field.extend_from_slice(&grain[start..start + width]); } + field } @@ -90,6 +92,7 @@ pub(super) fn autocov_of(field: &[f64], width: usize, height: usize) -> Vec sums[LAG_COUNT] += 1.0; } } + sums } @@ -99,16 +102,17 @@ pub(super) fn cell_mean_removed_record(field: &[f64], width: usize, height: usiz let cell = CELL as usize; let mut sums = vec![0.0f64; LAG_COUNT + 1]; - for y0 in (0..height - 2 * cell).step_by(cell) { - for x0 in (cell..width - 2 * cell).step_by(cell) { + for cell_top in (0..height - 2 * cell).step_by(cell) { + for cell_left in (cell..width - 2 * cell).step_by(cell) { let mut total = 0.0; - for y in y0..y0 + cell { - total += field[y * width + x0..y * width + x0 + cell].iter().sum::(); + for y in cell_top..cell_top + cell { + let row = &field[y * width + cell_left..y * width + cell_left + cell]; + total += row.iter().sum::(); } let mean = total / (cell * cell) as f64; - for y in y0..y0 + cell { - for x in x0..x0 + cell { + for y in cell_top..cell_top + cell { + for x in cell_left..cell_left + cell { let centre = field[y * width + x] - mean; for (lane, &(dy, dx)) in LAGS.iter().enumerate() { let neighbour_y = (y as i32 + dy) as usize; @@ -136,6 +140,7 @@ pub(super) fn grain_record() -> Vec { weights[left] = 0.4; let field = ar_field(&weights, 300, 300, 5); + autocov_of(&field, 300, 300) } @@ -178,5 +183,6 @@ pub(super) fn chunk_at(std_codes: f32, blocks_per_bin: u32, kept_codes: Option Vec<&str> { #[test] fn boundaries_sit_half_a_frame_early() { - assert_eq!(boundary_ticks(0, NTSC), 0); - assert_eq!(boundary_ticks(24, NTSC), 9_801_458); - assert_eq!(boundary_ticks(1, NTSC), 208_542); + let frame_0_ticks = boundary_ticks(0, NTSC); + let frame_24_ticks = boundary_ticks(24, NTSC); + let frame_1_ticks = boundary_ticks(1, NTSC); + + assert_eq!(frame_0_ticks, 0); + assert_eq!(frame_24_ticks, 9_801_458); + assert_eq!(frame_1_ticks, 208_542); } #[test] fn entries_end_one_tick_before_the_next_start() { - let text = format_table(&[entry(0, 23), entry(24, 47)], NTSC); + let entries = [entry(0, 23), entry(24, 47)]; + let text = format_table(&entries, NTSC); let starts = entry_lines(&text); assert_eq!(starts[0], "E 0 9801457 1 7391 1"); @@ -42,7 +47,8 @@ fn entries_end_one_tick_before_the_next_start() { #[test] fn the_last_entry_ends_half_a_frame_after_its_last_frame() { - let text = format_table(&[entry(0, 23)], NTSC); + let entries = [entry(0, 23)]; + let text = format_table(&entries, NTSC); let first = text.lines().nth(1).expect("an entry line"); assert_eq!(first, "E 0 9801458 1 7391 1"); @@ -50,7 +56,8 @@ fn the_last_entry_ends_half_a_frame_after_its_last_frame() { #[test] fn an_entry_has_the_exact_format() { - let text = format_table(&[entry(0, 23)], NTSC); + let entries = [entry(0, 23)]; + let text = format_table(&entries, NTSC); let expected = "filmgrn1\n\ E 0 9801458 1 7391 1\n\ \tp 3 7 0 11 0 1 128 192 256 128 192 256\n\ @@ -66,7 +73,8 @@ fn an_entry_has_the_exact_format() { #[test] fn a_gap_between_entries_is_left_empty() { - let text = format_table(&[entry(0, 23), entry(48, 71)], NTSC); + let entries = [entry(0, 23), entry(48, 71)]; + let text = format_table(&entries, NTSC); let starts = entry_lines(&text); assert_eq!(starts[0], "E 0 9801458 1 7391 1"); @@ -91,24 +99,27 @@ fn no_entries_writes_header_only() { #[test] fn reruns_are_byte_identical() { - let first = format_table(&[entry(0, 23), entry(24, 47)], NTSC); - let second = format_table(&[entry(0, 23), entry(24, 47)], NTSC); + let first_entries = [entry(0, 23), entry(24, 47)]; + let second_entries = [entry(0, 23), entry(24, 47)]; + let first = format_table(&first_entries, NTSC); + let second = format_table(&second_entries, NTSC); assert_eq!(first, second); } #[test] fn scenes_are_sorted_before_fitting() { - let scenes = vec![ - SceneGrain { - first_frame: 24, - chunks: vec![chunk_at(2.0, 300, None)], - }, - SceneGrain { - first_frame: 0, - chunks: vec![chunk_at(2.0, 300, None)], - }, - ]; + let later_chunk = chunk_at(2.0, 300, None); + let earlier_chunk = chunk_at(2.0, 300, None); + let later = SceneGrain { + first_frame: 24, + chunks: vec![later_chunk], + }; + let earlier = SceneGrain { + first_frame: 0, + chunks: vec![earlier_chunk], + }; + let scenes = vec![later, earlier]; let text = build_table(&scenes, NTSC); let starts = entry_lines(&text); diff --git a/av-denoise-core/src/nl4d/grain/tests/template.rs b/av-denoise-core/src/nl4d/grain/tests/template.rs index ce58c11..5dbf87e 100644 --- a/av-denoise-core/src/nl4d/grain/tests/template.rs +++ b/av-denoise-core/src/nl4d/grain/tests/template.rs @@ -13,6 +13,7 @@ fn sums(samples: &[i32]) -> (i64, i64) { .iter() .map(|&sample| (sample as i64) * (sample as i64)) .sum(); + (sum, sum_sq) } @@ -41,8 +42,8 @@ fn lag_and_offset_tables_have_the_expected_order() { #[test] fn lfsr_matches_the_specification() { - let mut rng = Av1Random::new(1000); - let draws: Vec = (0..8).map(|_| rng.next(11)).collect(); + let mut random = Av1Random::new(1000); + let draws: Vec = (0..8).map(|_| random.next(11)).collect(); assert_eq!(draws, [1039, 519, 259, 129, 1088, 1568, 1808, 904]); } @@ -52,13 +53,16 @@ fn template_matches_the_reference_port() { let template = luma_template(&SCENE_EIGHT, 7, 1000); let (sum, sum_sq) = sums(&template); let row = |y: usize, x: usize| template[y * 82 + x..y * 82 + x + 8].to_vec(); + let top_row = row(3, 0); + let middle_row = row(40, 30); + let bottom_row = row(72, 70); assert_eq!(template.len(), 73 * 82); assert_eq!(sum, -2598); assert_eq!(sum_sq, 11_079_620); - assert_eq!(row(3, 0), [-29, 27, -56, -20, 7, 58, 21, 47]); - assert_eq!(row(40, 30), [-15, -67, -39, -9, 16, 45, 48, -1]); - assert_eq!(row(72, 70), [10, 4, 44, 60, -48, -10, 7, 37]); + assert_eq!(top_row, [-29, 27, -56, -20, 7, 58, 21, 47]); + assert_eq!(middle_row, [-15, -67, -39, -9, 16, 45, 48, -1]); + assert_eq!(bottom_row, [10, 4, 44, 60, -48, -10, 7, 37]); } #[test] @@ -87,6 +91,9 @@ fn template_stats_match_the_reference_port() { #[test] fn template_seeds_stay_sixteen_bit() { - assert_eq!(template_seed(0), 1000); - assert_eq!(template_seed(23), ((1000u32 + 7919 * 23) & 0xFFFF) as u16); + let first_seed = template_seed(0); + let later_seed = template_seed(23); + + assert_eq!(first_seed, 1000); + assert_eq!(later_seed, ((1000u32 + 7919 * 23) & 0xFFFF) as u16); } diff --git a/av-denoise-core/src/nl4d/harness/mod.rs b/av-denoise-core/src/nl4d/harness/mod.rs index bf18097..988d8d9 100644 --- a/av-denoise-core/src/nl4d/harness/mod.rs +++ b/av-denoise-core/src/nl4d/harness/mod.rs @@ -1,13 +1,8 @@ -//! Synthetic clips with known motion, and the scores that compare a -//! motion field against them. -//! -//! This exists for the `mc_accuracy` bench. It is not a stable -//! interface. - +// Shared with the mc_accuracy bench, not a stable interface. #![doc(hidden)] mod score; mod synth; -pub use score::{KindScore, Score, score}; -pub use synth::{Clip, MotionClass, Still, synthesise}; +pub use self::score::{KindScore, Score, covering_blocks, score}; +pub use self::synth::{Clip, MotionClass, Still, synthesise}; diff --git a/av-denoise-core/src/nl4d/harness/score.rs b/av-denoise-core/src/nl4d/harness/score.rs index 84f932d..c21231f 100644 --- a/av-denoise-core/src/nl4d/harness/score.rs +++ b/av-denoise-core/src/nl4d/harness/score.rs @@ -3,21 +3,20 @@ use crate::collab::PATCH_SIZE; use crate::collab::geometry::{ref_pos, refs_along}; use crate::nl4d::MotionSnapshot; -/// The inclusive range of blocks whose `b * step..b * step + blksize` -/// span contains the patch `p..p + PATCH_SIZE`, clamped to the grid. +/// The inclusive range of blocks whose span contains the patch starting at `patch_start`. /// -/// When `step == blksize` and the patch straddles a tile boundary, no -/// block fully contains it. The corner block is returned as the best -/// available in that case, because a consumer still needs something to -/// search. -pub fn covering_blocks(p: u32, blksize: u32, step: u32, blocks: u32) -> (u32, u32) { - let hi = (p / step).min(blocks - 1); - let lo = if p + PATCH_SIZE <= blksize { +/// Each block spans `b * step..b * step + blksize`, and the range is clamped to the grid. When +/// `step == blksize` and the patch straddles a tile boundary, no block contains it, so the corner +/// block is returned as the best available. +pub fn covering_blocks(patch_start: u32, blksize: u32, step: u32, blocks: u32) -> (u32, u32) { + let last_block = (patch_start / step).min(blocks - 1); + let first_block = if patch_start + PATCH_SIZE <= blksize { 0 } else { - (p + PATCH_SIZE - blksize).div_ceil(step) + (patch_start + PATCH_SIZE - blksize).div_ceil(step) }; - (lo.min(hi), hi) + + (first_block.min(last_block), last_block) } /// How a patch's ground truth classifies it. @@ -59,7 +58,7 @@ impl KindScore { if self.epe.is_empty() { 0.0 } else { - self.epe.iter().map(|&e| e as f64).sum::() / self.epe.len() as f64 + self.epe.iter().map(|&error| error as f64).sum::() / self.epe.len() as f64 } } @@ -72,14 +71,16 @@ impl KindScore { } } -fn percentile(values: &[f32], q: f64) -> f64 { +fn percentile(values: &[f32], quantile: f64) -> f64 { if values.is_empty() { return 0.0; } + let mut sorted = values.to_vec(); sorted.sort_by(|a, b| a.partial_cmp(b).expect("no NaN in scores")); - let idx = ((sorted.len() - 1) as f64 * q).round() as usize; - sorted[idx] as f64 + let index = ((sorted.len() - 1) as f64 * quantile).round() as usize; + + sorted[index] as f64 } /// The full score of one field against one clip. @@ -101,28 +102,29 @@ impl Score { } /// The largest-axis distance between a truth and an integer vector. -fn endpoint_error(truth: [f32; 2], v: [i32; 2]) -> f32 { - (truth[0] - v[0] as f32).abs().max((truth[1] - v[1] as f32).abs()) +fn endpoint_error(truth: [f32; 2], vector: [i32; 2]) -> f32 { + let error_x = (truth[0] - vector[0] as f32).abs(); + let error_y = (truth[1] - vector[1] as f32).abs(); + + error_x.max(error_y) } -/// Scores `snap` against `clip` over nl4d's reference grid and every -/// neighbour. +/// Scores `snapshot` against `clip` over nl4d's reference grid and every neighbour. /// -/// A patch is in window when its truth lies within `refine` pixels of -/// the vector on both axes. The corner reading uses the block whose -/// corner the patch sits on, the block nl4d reads today. The covering -/// reading takes the best of every block that covers the patch. -pub fn score(clip: &Clip, snap: &MotionSnapshot, refine: u32) -> Score { - let (w, h) = (clip.width, clip.height); - let mut out = Score::default(); +/// A patch is in window when its truth lies within `refine` pixels of the vector on both axes. +/// The corner reading uses the block the patch's corner sits on. The covering reading takes the +/// best of every block that covers the patch. +pub fn score(clip: &Clip, snapshot: &MotionSnapshot, refine: u32) -> Score { + let (width, height) = (clip.width, clip.height); + let mut result = Score::default(); assert_eq!( - snap.vectors.len(), - snap.confidence.len(), + snapshot.vectors.len(), + snapshot.confidence.len(), "vectors and confidence must carry the same neighbour count and convention" ); assert_eq!( - snap.vectors.len(), + snapshot.vectors.len(), clip.truth.len(), "the snapshot's neighbour count must match the clip's truth, which both index by \ `neighbour_idx_for_k`" @@ -130,30 +132,34 @@ pub fn score(clip: &Clip, snap: &MotionSnapshot, refine: u32) -> Score { for (t, truth) in clip.truth.iter().enumerate() { let occluded = &clip.occluded[t]; - for ry in 0..refs_along(h) { - for rx in 0..refs_along(w) { - let px = ref_pos(rx, w); - let py = ref_pos(ry, h); + for ref_y in 0..refs_along(height) { + for ref_x in 0..refs_along(width) { + let patch_x = ref_pos(ref_x, width); + let patch_y = ref_pos(ref_y, height); let mut sum = [0.0f32; 2]; let mut any_occluded = false; - for y in py..py + PATCH_SIZE { - for x in px..px + PATCH_SIZE { - let idx = (y * w + x) as usize; + for y in patch_y..patch_y + PATCH_SIZE { + for x in patch_x..patch_x + PATCH_SIZE { + let idx = (y * width + x) as usize; sum[0] += truth[idx][0]; sum[1] += truth[idx][1]; any_occluded |= occluded[idx]; } } + let area = (PATCH_SIZE * PATCH_SIZE) as f32; let mean = [sum[0] / area, sum[1] / area]; let mut spread = 0.0f32; - for y in py..py + PATCH_SIZE { - for x in px..px + PATCH_SIZE { - let d = truth[(y * w + x) as usize]; - spread = spread.max((d[0] - mean[0]).abs()).max((d[1] - mean[1]).abs()); + for y in patch_y..patch_y + PATCH_SIZE { + for x in patch_x..patch_x + PATCH_SIZE { + let displacement = truth[(y * width + x) as usize]; + let spread_x = (displacement[0] - mean[0]).abs(); + let spread_y = (displacement[1] - mean[1]).abs(); + spread = spread.max(spread_x).max(spread_y); } } + let kind = if any_occluded { PatchKind::Occluded } else if spread > 0.5 { @@ -162,35 +168,40 @@ pub fn score(clip: &Clip, snap: &MotionSnapshot, refine: u32) -> Score { PatchKind::Plain }; - let (bx_lo, bx_hi) = covering_blocks(px, snap.blksize, snap.step, snap.blocks_x); - let (by_lo, by_hi) = covering_blocks(py, snap.blksize, snap.step, snap.blocks_y); - let corner = (by_hi * snap.blocks_x + bx_hi) as usize; - let corner_v = snap.vectors[t][corner]; - let corner_err = endpoint_error(mean, corner_v); - - let mut best_err = corner_err; - for by in by_lo..=by_hi { - for bx in bx_lo..=bx_hi { - let v = snap.vectors[t][(by * snap.blocks_x + bx) as usize]; - best_err = best_err.min(endpoint_error(mean, v)); + let (first_block_x, last_block_x) = + covering_blocks(patch_x, snapshot.blksize, snapshot.step, snapshot.blocks_x); + let (first_block_y, last_block_y) = + covering_blocks(patch_y, snapshot.blksize, snapshot.step, snapshot.blocks_y); + let corner = (last_block_y * snapshot.blocks_x + last_block_x) as usize; + let corner_vector = snapshot.vectors[t][corner]; + let corner_error = endpoint_error(mean, corner_vector); + + let mut best_error = corner_error; + for by in first_block_y..=last_block_y { + for bx in first_block_x..=last_block_x { + let vector = snapshot.vectors[t][(by * snapshot.blocks_x + bx) as usize]; + let block_error = endpoint_error(mean, vector); + best_error = best_error.min(block_error); } } - let k = out.kind_mut(kind); - k.patches += 1; - if corner_err <= refine as f32 { - k.in_window_corner += 1; + let kind_score = result.kind_mut(kind); + kind_score.patches += 1; + if corner_error <= refine as f32 { + kind_score.in_window_corner += 1; } - if best_err <= refine as f32 { - k.in_window_covering += 1; + + if best_error <= refine as f32 { + kind_score.in_window_covering += 1; } - k.epe.push(corner_err); - k.confidence.push(snap.confidence[t][corner]); + + kind_score.epe.push(corner_error); + kind_score.confidence.push(snapshot.confidence[t][corner]); } } } - out + result } #[cfg(test)] @@ -198,32 +209,37 @@ mod tests { use super::*; use crate::nl4d::MotionSnapshot; - /// A 32x32 clip at radius 1 whose truth toward k = +1 is a uniform - /// `[3, 1]`, with nothing occluded. + /// A 32x32 clip at radius 1 whose truth toward k = +1 is a uniform `[3, 1]`, with nothing + /// occluded. fn uniform_clip() -> Clip { - let (w, h) = (32u32, 32u32); - let n = (w * h) as usize; + let (width, height) = (32u32, 32u32); + let pixels = (width * height) as usize; + Clip { - width: w, - height: h, + width, + height, radius: 1, - frames: vec![vec![0.5; n]; 3], - truth: vec![vec![[-3.0, -1.0]; n], vec![[3.0, 1.0]; n]], - occluded: vec![vec![false; n]; 2], + frames: vec![vec![0.5; pixels]; 3], + truth: vec![vec![[-3.0, -1.0]; pixels], vec![[3.0, 1.0]; pixels]], + occluded: vec![vec![false; pixels]; 2], } } /// One vector for every block of every neighbour. - fn uniform_snapshot(vx: i32, vy: i32) -> MotionSnapshot { + fn uniform_snapshot(vector_x: i32, vector_y: i32) -> MotionSnapshot { let (blocks_x, blocks_y) = (4u32, 4u32); let blocks = (blocks_x * blocks_y) as usize; + MotionSnapshot { blocks_x, blocks_y, step: 8, blksize: 16, offsets: vec![-1, 1], - vectors: vec![vec![[-vx, -vy]; blocks], vec![[vx, vy]; blocks]], + vectors: vec![ + vec![[-vector_x, -vector_y]; blocks], + vec![[vector_x, vector_y]; blocks], + ], confidence: vec![vec![0.9; blocks]; 2], } } @@ -231,108 +247,125 @@ mod tests { #[test] #[should_panic(expected = "neighbour count must match")] fn score_asserts_the_snapshots_neighbour_count_matches_the_clips_truth() { - // The clip carries truth for both neighbours (k = -1, +1) but - // the snapshot only carries one, so the two disagree on how - // many neighbours `neighbour_idx_for_k` indexes. - let mut snap = uniform_snapshot(3, 1); - snap.vectors.truncate(1); - snap.confidence.truncate(1); - score(&uniform_clip(), &snap, 2); + let mut snapshot = uniform_snapshot(3, 1); + snapshot.vectors.truncate(1); + snapshot.confidence.truncate(1); + + let clip = uniform_clip(); + score(&clip, &snapshot, 2); } #[test] #[should_panic(expected = "same neighbour count and convention")] fn score_asserts_vectors_and_confidence_carry_the_same_neighbour_count() { - let mut snap = uniform_snapshot(3, 1); - snap.confidence.pop(); - score(&uniform_clip(), &snap, 2); + let mut snapshot = uniform_snapshot(3, 1); + snapshot.confidence.pop(); + + let clip = uniform_clip(); + score(&clip, &snapshot, 2); } #[test] fn covering_blocks_for_the_default_geometry() { - // blksize 16, step 8: patch at 0 is covered by block 0 only, - // patch at 8 by blocks 0 and 1, patch at 16 by blocks 1 and 2. - assert_eq!(covering_blocks(0, 16, 8, 8), (0, 0)); - assert_eq!(covering_blocks(8, 16, 8, 8), (0, 1)); - assert_eq!(covering_blocks(16, 16, 8, 8), (1, 2)); + // At blksize 16 and step 8, the patch at 0 is covered by block 0 only, the patch at 8 by + // blocks 0 and 1, and the patch at 16 by blocks 1 and 2. + let at_zero = covering_blocks(0, 16, 8, 8); + let at_eight = covering_blocks(8, 16, 8, 8); + let at_sixteen = covering_blocks(16, 16, 8, 8); // step == blksize gives exactly one block. - assert_eq!(covering_blocks(24, 8, 8, 8), (3, 3)); + let step_equals_blksize = covering_blocks(24, 8, 8, 8); // The upper end clamps to the grid. - assert_eq!(covering_blocks(56, 16, 8, 7), (6, 6)); + let clamped = covering_blocks(56, 16, 8, 7); + + assert_eq!(at_zero, (0, 0)); + assert_eq!(at_eight, (0, 1)); + assert_eq!(at_sixteen, (1, 2)); + assert_eq!(step_equals_blksize, (3, 3)); + assert_eq!(clamped, (6, 6)); } #[test] fn a_straddling_patch_at_step_equal_blksize_falls_back_to_the_corner_block() { - // blksize 16, step 16: block 0 spans 0..16, block 1 spans - // 16..32. The patch at p = 10 spans 10..18, which no single - // block fully contains. No range is empty here, so the corner - // block (the one the patch's start pixel sits in) is returned - // as the best available search target, matching what the corner - // reading already reads today. - assert_eq!(covering_blocks(10, 16, 16, 8), (0, 0)); + // The patch spans 10..18, which neither block 0..16 nor block 16..32 fully contains. + let covering = covering_blocks(10, 16, 16, 8); + assert_eq!(covering, (0, 0)); } #[test] fn an_exact_field_scores_every_patch_in_window_with_zero_error() { - let s = score(&uniform_clip(), &uniform_snapshot(3, 1), 2); - assert!(s.plain.patches > 0); - assert_eq!(s.boundary.patches, 0); - assert_eq!(s.occluded.patches, 0); - assert_eq!(s.plain.in_window_rate_corner(), 1.0); - assert_eq!(s.plain.in_window_rate_covering(), 1.0); - assert_eq!(s.plain.epe_mean(), 0.0); - assert!((s.plain.confidence_median() - 0.9).abs() < 1e-6); + let clip = uniform_clip(); + let snapshot = uniform_snapshot(3, 1); + let scores = score(&clip, &snapshot, 2); + assert!(scores.plain.patches > 0); + assert_eq!(scores.boundary.patches, 0); + assert_eq!(scores.occluded.patches, 0); + assert_eq!(scores.plain.in_window_rate_corner(), 1.0); + assert_eq!(scores.plain.in_window_rate_covering(), 1.0); + assert_eq!(scores.plain.epe_mean(), 0.0); + assert!((scores.plain.confidence_median() - 0.9).abs() < 1e-6); } #[test] fn an_error_past_the_refine_window_scores_out_of_window() { - // Off by 3 on x, refine 2: out of window, endpoint error 3. - let s = score(&uniform_clip(), &uniform_snapshot(6, 1), 2); - assert_eq!(s.plain.in_window_rate_corner(), 0.0); - assert!((s.plain.epe_mean() - 3.0).abs() < 1e-6); - assert!((s.plain.epe_p95() - 3.0).abs() < 1e-6); - // Refine 3 admits it. - let s = score(&uniform_clip(), &uniform_snapshot(6, 1), 3); - assert_eq!(s.plain.in_window_rate_corner(), 1.0); + // Off by 3 on x, so refine 2 puts it out of window with an endpoint error of 3. + let clip = uniform_clip(); + let snapshot = uniform_snapshot(6, 1); + let scores = score(&clip, &snapshot, 2); + assert_eq!(scores.plain.in_window_rate_corner(), 0.0); + assert!((scores.plain.epe_mean() - 3.0).abs() < 1e-6); + assert!((scores.plain.epe_p95() - 3.0).abs() < 1e-6); + + let scores = score(&clip, &snapshot, 3); + assert_eq!(scores.plain.in_window_rate_corner(), 1.0); } #[test] fn the_covering_reading_takes_the_best_covering_block() { - // Corner blocks wrong, every other block right. Patches whose - // corner block is wrong but which another block covers still - // count in the covering reading. - let mut snap = uniform_snapshot(3, 1); - for by in 0..4u32 { - for bx in 0..4u32 { - if (bx + by) % 2 == 0 { - snap.vectors[1][(by * 4 + bx) as usize] = [30, 30]; + // Every other block is wrong, so patches whose corner block is wrong still count in the + // covering reading through another block. + let mut snapshot = uniform_snapshot(3, 1); + for block_y in 0..4u32 { + for block_x in 0..4u32 { + if (block_x + block_y) % 2 == 0 { + snapshot.vectors[1][(block_y * 4 + block_x) as usize] = [30, 30]; } } } - let s = score(&uniform_clip(), &snap, 2); - assert!(s.plain.in_window_rate_covering() > s.plain.in_window_rate_corner()); + + let clip = uniform_clip(); + let scores = score(&clip, &snapshot, 2); + let covering_rate = scores.plain.in_window_rate_covering(); + let corner_rate = scores.plain.in_window_rate_corner(); + assert!(covering_rate > corner_rate); } #[test] fn boundary_and_occluded_patches_are_classified_by_the_truth() { let mut clip = uniform_clip(); - let w = clip.width as usize; + let width = clip.width as usize; + // A vertical motion boundary at x = 16 toward k = +1. for y in 0..32usize { for x in 16..32usize { - clip.truth[1][y * w + x] = [0.0, 0.0]; + clip.truth[1][y * width + x] = [0.0, 0.0]; } } + // Pixel column 20 occluded toward k = +1. for y in 0..32usize { - clip.occluded[1][y * w + 20] = true; + clip.occluded[1][y * width + 20] = true; } - let s = score(&clip, &uniform_snapshot(3, 1), 2); + + let snapshot = uniform_snapshot(3, 1); + let scores = score(&clip, &snapshot, 2); assert!( - s.boundary.patches > 0, + scores.boundary.patches > 0, "patches straddling x = 16 are boundary patches" ); - assert!(s.occluded.patches > 0, "patches touching column 20 are occluded"); - assert!(s.plain.patches > 0); + assert!( + scores.occluded.patches > 0, + "patches touching column 20 are occluded" + ); + assert!(scores.plain.patches > 0); } } diff --git a/av-denoise-core/src/nl4d/harness/synth.rs b/av-denoise-core/src/nl4d/harness/synth.rs index 4305d6d..147cd65 100644 --- a/av-denoise-core/src/nl4d/harness/synth.rs +++ b/av-denoise-core/src/nl4d/harness/synth.rs @@ -1,6 +1,6 @@ use crate::nlmeans::motion::neighbour_idx_for_k; -/// A clean luma plane, values in `[0, 1]`. +/// A clean luma plane with values between 0 and 1. #[derive(Debug, Clone)] pub struct Still { pub width: u32, @@ -16,54 +16,64 @@ impl Still { if bytes.len() < 2 || &bytes[..2] != b"P5" { return Err("not a P5 pgm".to_string()); } + pos += 2; while fields.len() < 3 { while pos < bytes.len() && bytes[pos].is_ascii_whitespace() { pos += 1; } + if pos < bytes.len() && bytes[pos] == b'#' { while pos < bytes.len() && bytes[pos] != b'\n' { pos += 1; } continue; } + let start = pos; while pos < bytes.len() && bytes[pos].is_ascii_digit() { pos += 1; } + if start == pos { return Err("malformed pgm header".to_string()); } - let text = std::str::from_utf8(&bytes[start..pos]).map_err(|e| e.to_string())?; - fields.push(text.parse::().map_err(|e| e.to_string())?); + + let text = std::str::from_utf8(&bytes[start..pos]).map_err(|error| error.to_string())?; + let field = text.parse::().map_err(|error| error.to_string())?; + fields.push(field); } + // Exactly one whitespace byte separates maxval from the data. pos += 1; let (width, height, maxval) = (fields[0], fields[1], fields[2]); - // Widened to 64 bits so a header claiming huge dimensions cannot - // wrap back into a small `usize` on a 32-bit target and slip - // past the bounds check below. - let n64 = width as u64 * height as u64; + // Widened to 64 bits so huge header dimensions cannot wrap into a small `usize` on a 32-bit + // target and slip past the bounds check. + let sample_count_u64 = width as u64 * height as u64; let bytes_per_sample: u64 = if maxval > 255 { 2 } else { 1 }; let available = (bytes.len() - pos.min(bytes.len())) as u64; - if n64 * bytes_per_sample > available { + if sample_count_u64 * bytes_per_sample > available { return Err(format!( "pgm data truncated: header claims {width}x{height} at {bytes_per_sample} bytes/sample, \ only {available} bytes remain" )); } - let n = n64 as usize; + + let sample_count = sample_count_u64 as usize; let luma = if maxval > 255 { - let data = bytes.get(pos..pos + 2 * n).ok_or("pgm data truncated")?; + let data = bytes + .get(pos..pos + 2 * sample_count) + .ok_or("pgm data truncated")?; data.as_chunks::<2>() .0 .iter() - .map(|c| u16::from_be_bytes(*c) as f32 / maxval as f32) + .map(|chunk| u16::from_be_bytes(*chunk) as f32 / maxval as f32) .collect() } else { - let data = bytes.get(pos..pos + n).ok_or("pgm data truncated")?; - data.iter().map(|&v| v as f32 / maxval as f32).collect() + let data = bytes.get(pos..pos + sample_count).ok_or("pgm data truncated")?; + data.iter().map(|&value| value as f32 / maxval as f32).collect() }; + Ok(Still { width, height, luma }) } @@ -72,44 +82,50 @@ impl Still { let mut luma = vec![0.0f32; (width * height) as usize]; for y in 0..height { for x in 0..width { - let fx = x as f32 * 0.31; - let fy = y as f32 * 0.23; - let v = 0.5 + 0.2 * (fx.sin() * fy.cos()) + 0.1 * ((fx * 2.7).cos() + (fy * 3.1).sin()); - luma[(y * width + x) as usize] = v.clamp(0.0, 1.0); + let phase_x = x as f32 * 0.31; + let phase_y = y as f32 * 0.23; + let value = 0.5 + + 0.2 * (phase_x.sin() * phase_y.cos()) + + 0.1 * ((phase_x * 2.7).cos() + (phase_y * 3.1).sin()); + luma[(y * width + x) as usize] = value.clamp(0.0, 1.0); } } + Still { width, height, luma } } - /// Samples the still at a fractional position with a Lanczos-3 - /// kernel, clamping to the edge. - fn sample(&self, sx: f32, sy: f32) -> f32 { + /// Samples the still at a fractional position with a Lanczos-3 kernel, clamping to the edge. + fn sample(&self, sample_x: f32, sample_y: f32) -> f32 { const A: i32 = 3; + let lanczos = |t: f32| -> f32 { if t == 0.0 { 1.0 } else if t.abs() >= A as f32 { 0.0 } else { - let pt = std::f32::consts::PI * t; - (A as f32 * pt.sin() * (pt / A as f32).sin()) / (pt * pt) + let pi_t = std::f32::consts::PI * t; + (A as f32 * pi_t.sin() * (pi_t / A as f32).sin()) / (pi_t * pi_t) } }; - let x0 = sx.floor() as i32; - let y0 = sy.floor() as i32; - let mut acc = 0.0f32; - let mut wsum = 0.0f32; + + let x0 = sample_x.floor() as i32; + let y0 = sample_y.floor() as i32; + let mut weighted_sum = 0.0f32; + let mut weight_sum = 0.0f32; for j in (y0 - A + 1)..=(y0 + A) { - let wy = lanczos(sy - j as f32); - let yy = j.clamp(0, self.height as i32 - 1) as u32; + let weight_y = lanczos(sample_y - j as f32); + let clamped_y = j.clamp(0, self.height as i32 - 1) as u32; for i in (x0 - A + 1)..=(x0 + A) { - let w = wy * lanczos(sx - i as f32); - let xx = i.clamp(0, self.width as i32 - 1) as u32; - acc += w * self.luma[(yy * self.width + xx) as usize]; - wsum += w; + let weight_x = lanczos(sample_x - i as f32); + let weight = weight_y * weight_x; + let clamped_x = i.clamp(0, self.width as i32 - 1) as u32; + weighted_sum += weight * self.luma[(clamped_y * self.width + clamped_x) as usize]; + weight_sum += weight; } } - (acc / wsum).clamp(0.0, 1.0) + + (weighted_sum / weight_sum).clamp(0.0, 1.0) } } @@ -152,14 +168,16 @@ impl MotionClass { } } - /// Top-left corner and side of the cut-out rectangle in the centre - /// frame, a square a third of the shorter side, left of centre so - /// its rightward motion stays inside the frame. + /// Top-left corner and side of the cut-out square in the centre frame. + /// + /// The square is a third of the shorter side, left of centre so its rightward motion stays + /// inside the frame. pub fn cut_out_rect(width: u32, height: u32) -> (u32, u32, u32) { let side = (width.min(height) / 3).max(8); - let x0 = width / 4; - let y0 = (height - side) / 2; - (x0, y0, side) + let rect_x = width / 4; + let rect_y = (height - side) / 2; + + (rect_x, rect_y, side) } } @@ -171,27 +189,28 @@ pub struct Clip { pub radius: u32, /// `frames[i]` is the frame at offset `k = i - radius`. pub frames: Vec>, - /// `truth[t][pixel]` is where the centre frame's pixel lies in - /// neighbour `t`, as a displacement in pixels. + /// `truth[t][pixel]` is where the centre frame's pixel lies in neighbour `t`, as a + /// displacement in pixels. pub truth: Vec>, - /// `occluded[t][pixel]` is true when that pixel has no true match - /// in neighbour `t`. + /// `occluded[t][pixel]` is true when that pixel has no true match in neighbour `t`. pub occluded: Vec>, } -/// Where the centre frame's pixel `(x, y)` sits in the frame at offset -/// `k`, for the background of `class`. +/// Where the centre frame's background pixel `(x, y)` sits in the frame at offset `k`. fn background_displacement(class: MotionClass, k: i32, x: u32, y: u32, width: u32, height: u32) -> [f32; 2] { match class { MotionClass::IntegerPan | MotionClass::HalfPelPan => { - let v = class.velocity(); - [v[0] * k as f32, v[1] * k as f32] + let velocity = class.velocity(); + [velocity[0] * k as f32, velocity[1] * k as f32] }, MotionClass::Zoom => { - let s = ZOOM_PER_FRAME.powi(k); - let cx = width as f32 / 2.0; - let cy = height as f32 / 2.0; - [(x as f32 - cx) * (s - 1.0), (y as f32 - cy) * (s - 1.0)] + let scale = ZOOM_PER_FRAME.powi(k); + let centre_x = width as f32 / 2.0; + let centre_y = height as f32 / 2.0; + [ + (x as f32 - centre_x) * (scale - 1.0), + (y as f32 - centre_y) * (scale - 1.0), + ] }, MotionClass::CutOut => [0.0, 0.0], } @@ -200,98 +219,109 @@ fn background_displacement(class: MotionClass, k: i32, x: u32, y: u32, width: u3 /// Deterministic Gaussian grain from a hashed uniform pair. fn grain(idx: u32, seed: u32) -> f32 { let hash = |i: u32| -> f32 { - let mut h = i + let mut state = i .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(0x9E37_79B9)); - h ^= h >> 15; - h = h.wrapping_mul(0x85EB_CA6B); - h ^= h >> 13; - (h as f32 + 1.0) / (u32::MAX as f32 + 2.0) + state ^= state >> 15; + state = state.wrapping_mul(0x85EB_CA6B); + state ^= state >> 13; + (state as f32 + 1.0) / (u32::MAX as f32 + 2.0) }; + let u1 = hash(idx * 2); let u2 = hash(idx * 2 + 1); + (-2.0 * u1.ln()).sqrt() * (std::f32::consts::TAU * u2).cos() } -/// Builds the window of `2 * radius + 1` frames for `class`, with -/// Gaussian grain of `sigma` on every frame, and the ground truth toward -/// every neighbour. +/// Builds the window of `2 * radius + 1` frames for `class` and the ground truth toward every +/// neighbour. +/// +/// Every frame carries its own Gaussian grain of `sigma`. pub fn synthesise(still: &Still, class: MotionClass, radius: u32, sigma: f32, seed: u32) -> Clip { - let (w, h) = (still.width, still.height); - let n = (w * h) as usize; - let (cx0, cy0, side) = MotionClass::cut_out_rect(w, h); - let v = class.velocity(); + let (width, height) = (still.width, still.height); + let pixels = (width * height) as usize; + let (rect_x, rect_y, side) = MotionClass::cut_out_rect(width, height); + let velocity = class.velocity(); let in_rect_at = |x: f32, y: f32, k: i32| -> bool { - let ox = cx0 as f32 + v[0] * k as f32; - let oy = cy0 as f32 + v[1] * k as f32; - x >= ox && x < ox + side as f32 && y >= oy && y < oy + side as f32 + let origin_x = rect_x as f32 + velocity[0] * k as f32; + let origin_y = rect_y as f32 + velocity[1] * k as f32; + x >= origin_x && x < origin_x + side as f32 && y >= origin_y && y < origin_y + side as f32 }; let mut frames = Vec::with_capacity((2 * radius + 1) as usize); for i in 0..(2 * radius + 1) as i32 { let k = i - radius as i32; - let mut frame = vec![0.0f32; n]; - for y in 0..h { - for x in 0..w { - let idx = (y * w + x) as usize; + let mut frame = vec![0.0f32; pixels]; + for y in 0..height { + for x in 0..width { + let idx = (y * width + x) as usize; let value = if class == MotionClass::CutOut && in_rect_at(x as f32, y as f32, k) { // The rectangle's content, read from where it sat in the centre. - still.sample(x as f32 - v[0] * k as f32, y as f32 - v[1] * k as f32) + still.sample( + x as f32 - velocity[0] * k as f32, + y as f32 - velocity[1] * k as f32, + ) } else { - let d = background_displacement(class, k, x, y, w, h); - // The frame at k shows the centre's pixel p at p + d, so - // pixel (x, y) here comes from the centre's (x, y) - d. - // For a pan and a zoom the inverse map is exact. + let displacement = background_displacement(class, k, x, y, width, height); + // The frame at k shows the centre's pixel p at p + d, so pixel (x, y) here + // comes from the centre's (x, y) - d. For a pan and a zoom the inverse is exact. match class { MotionClass::Zoom => { - let s = ZOOM_PER_FRAME.powi(k); - let fx = w as f32 / 2.0 + (x as f32 - w as f32 / 2.0) / s; - let fy = h as f32 / 2.0 + (y as f32 - h as f32 / 2.0) / s; - still.sample(fx, fy) + let scale = ZOOM_PER_FRAME.powi(k); + let source_x = width as f32 / 2.0 + (x as f32 - width as f32 / 2.0) / scale; + let source_y = height as f32 / 2.0 + (y as f32 - height as f32 / 2.0) / scale; + still.sample(source_x, source_y) }, - _ => still.sample(x as f32 - d[0], y as f32 - d[1]), + _ => still.sample(x as f32 - displacement[0], y as f32 - displacement[1]), } }; let noise = if sigma > 0.0 { - sigma * grain(idx as u32, seed.wrapping_add(1000 * (i as u32 + 1))) + let frame_seed = seed.wrapping_add(1000 * (i as u32 + 1)); + sigma * grain(idx as u32, frame_seed) } else { 0.0 }; frame[idx] = (value + noise).clamp(0.0, 1.0); } } + frames.push(frame); } let neighbours = (2 * radius) as usize; - let mut truth = vec![vec![[0.0f32; 2]; n]; neighbours]; - let mut occluded = vec![vec![false; n]; neighbours]; + let mut truth = vec![vec![[0.0f32; 2]; pixels]; neighbours]; + let mut occluded = vec![vec![false; pixels]; neighbours]; for k in -(radius as i32)..=(radius as i32) { if k == 0 { continue; } + let t = neighbour_idx_for_k(radius, k) as usize; - for y in 0..h { - for x in 0..w { - let idx = (y * w + x) as usize; + for y in 0..height { + for x in 0..width { + let idx = (y * width + x) as usize; let foreground = class == MotionClass::CutOut && in_rect_at(x as f32, y as f32, 0); - let d = if foreground { - [v[0] * k as f32, v[1] * k as f32] + let displacement = if foreground { + [velocity[0] * k as f32, velocity[1] * k as f32] } else { - background_displacement(class, k, x, y, w, h) + background_displacement(class, k, x, y, width, height) }; - truth[t][idx] = d; + truth[t][idx] = displacement; + if class == MotionClass::CutOut && !foreground { - occluded[t][idx] = in_rect_at(x as f32 + d[0], y as f32 + d[1], k); + let moved_x = x as f32 + displacement[0]; + let moved_y = y as f32 + displacement[1]; + occluded[t][idx] = in_rect_at(moved_x, moved_y, k); } } } } Clip { - width: w, - height: h, + width, + height, radius, frames, truth, @@ -305,15 +335,16 @@ mod tests { fn ramp_still() -> Still { // Distinct values everywhere, so a shift is visible in any pixel. - let (w, h) = (64u32, 48u32); - let luma = (0..w * h) - .map(|i| ((i % w) as f32 * 0.9 / w as f32 + (i / w) as f32 * 0.1 / h as f32).clamp(0.0, 1.0)) + let (width, height) = (64u32, 48u32); + let luma = (0..width * height) + .map(|i| { + let ramp_x = (i % width) as f32 * 0.9 / width as f32; + let ramp_y = (i / width) as f32 * 0.1 / height as f32; + (ramp_x + ramp_y).clamp(0.0, 1.0) + }) .collect(); - Still { - width: w, - height: h, - luma, - } + + Still { width, height, luma } } #[test] @@ -321,21 +352,24 @@ mod tests { let still = ramp_still(); let clip = synthesise(&still, MotionClass::IntegerPan, 1, 0.0, 1); assert_eq!(clip.frames.len(), 3); - let (w, h) = (clip.width, clip.height); - let v = MotionClass::IntegerPan.velocity(); - // Frame k = +1 holds the still moved by +v. Check an interior pixel. + + let (width, height) = (clip.width, clip.height); + let velocity = MotionClass::IntegerPan.velocity(); + // Frame k = +1 holds the still moved by +velocity. Check an interior pixel. let (x, y) = (20u32, 20u32); - let moved = clip.frames[2][(y * w + x) as usize]; - let source = - still.luma[((y as i32 - v[1] as i32) as u32 * w + (x as i32 - v[0] as i32) as u32) as usize]; + let moved = clip.frames[2][(y * width + x) as usize]; + let source_y = (y as i32 - velocity[1] as i32) as u32; + let source_x = (x as i32 - velocity[0] as i32) as u32; + let source = still.luma[(source_y * width + source_x) as usize]; assert!( (moved - source).abs() < 1e-6, "an integer pan must copy pixels exactly" ); - // Truth toward k = +1 (t = 1 at radius 1) is +v everywhere. - for idx in 0..(w * h) as usize { - assert_eq!(clip.truth[1][idx], v); - assert_eq!(clip.truth[0][idx], [-v[0], -v[1]]); + + // Truth toward k = +1 (t = 1 at radius 1) is +velocity everywhere. + for idx in 0..(width * height) as usize { + assert_eq!(clip.truth[1][idx], velocity); + assert_eq!(clip.truth[0][idx], [-velocity[0], -velocity[1]]); assert!(!clip.occluded[1][idx]); } } @@ -344,27 +378,29 @@ mod tests { fn half_pel_pan_truth_has_a_half_pixel_component() { let still = ramp_still(); let clip = synthesise(&still, MotionClass::HalfPelPan, 1, 0.0, 1); - let v = MotionClass::HalfPelPan.velocity(); - assert!((v[0].fract().abs() - 0.5).abs() < 1e-6 || (v[1].fract().abs() - 0.5).abs() < 1e-6); - assert_eq!(clip.truth[1][100], v); + let velocity = MotionClass::HalfPelPan.velocity(); + let half_x = (velocity[0].fract().abs() - 0.5).abs() < 1e-6; + let half_y = (velocity[1].fract().abs() - 0.5).abs() < 1e-6; + assert!(half_x || half_y); + assert_eq!(clip.truth[1][100], velocity); } #[test] fn zoom_truth_grows_with_distance_from_the_centre() { let still = ramp_still(); let clip = synthesise(&still, MotionClass::Zoom, 1, 0.0, 1); - let (w, h) = (clip.width, clip.height); - let centre = ((h / 2) * w + w / 2) as usize; + let (width, height) = (clip.width, clip.height); + let centre = ((height / 2) * width + width / 2) as usize; let corner = 0usize; - let dc = clip.truth[1][centre]; - let dk = clip.truth[1][corner]; + let centre_shift = clip.truth[1][centre]; + let corner_shift = clip.truth[1][corner]; assert!( - dc[0].abs() < 0.01 && dc[1].abs() < 0.01, + centre_shift[0].abs() < 0.01 && centre_shift[1].abs() < 0.01, "the centre does not move under a zoom" ); assert!( - dk[0] < -0.1 && dk[1] < -0.1, - "the top-left corner moves outward, got {dk:?}" + corner_shift[0] < -0.1 && corner_shift[1] < -0.1, + "the top-left corner moves outward, got {corner_shift:?}" ); } @@ -372,17 +408,18 @@ mod tests { fn cut_out_marks_background_hidden_under_the_moved_rectangle() { let still = ramp_still(); let clip = synthesise(&still, MotionClass::CutOut, 1, 0.0, 1); - let w = clip.width; - let (x0, y0, side) = MotionClass::cut_out_rect(clip.width, clip.height); - let v = MotionClass::CutOut.velocity(); + let width = clip.width; + let (rect_x, rect_y, side) = MotionClass::cut_out_rect(clip.width, clip.height); + let velocity = MotionClass::CutOut.velocity(); + // A pixel inside the rectangle in the centre frame moves with it. - let inside = ((y0 + side / 2) * w + x0 + side / 2) as usize; - assert_eq!(clip.truth[1][inside], v); + let inside = ((rect_y + side / 2) * width + rect_x + side / 2) as usize; + assert_eq!(clip.truth[1][inside], velocity); assert!(!clip.occluded[1][inside]); - // A background pixel just to the right of the rectangle is - // covered once it moves right by v[0] pixels, so toward k = +1 - // it is occluded, and toward k = -1 it is not. - let just_right = ((y0 + side / 2) * w + x0 + side + 1) as usize; + + // A background pixel just right of the rectangle is covered once the rectangle moves + // right, so it is occluded toward k = +1 and not toward k = -1. + let just_right = ((rect_y + side / 2) * width + rect_x + side + 1) as usize; assert_eq!(clip.truth[1][just_right], [0.0, 0.0]); assert!(clip.occluded[1][just_right]); assert!(!clip.occluded[0][just_right]); @@ -393,14 +430,14 @@ mod tests { let still = Still::synthetic(128, 128); let clip = synthesise(&still, MotionClass::IntegerPan, 1, 0.0, 1); let noisy = synthesise(&still, MotionClass::IntegerPan, 1, 6.0 / 255.0, 1); - let n = clip.frames[1].len() as f32; - let var: f32 = clip.frames[1] + let pixels = clip.frames[1].len() as f32; + let variance: f32 = clip.frames[1] .iter() .zip(&noisy.frames[1]) - .map(|(a, b)| (a - b) * (a - b)) + .map(|(clean, grainy)| (clean - grainy) * (clean - grainy)) .sum::() - / n; - let sigma = var.sqrt(); + / pixels; + let sigma = variance.sqrt(); assert!( (sigma - 6.0 / 255.0).abs() < 0.1 * 6.0 / 255.0, "measured sigma {sigma}" @@ -413,35 +450,30 @@ mod tests { #[test] fn pgm_parses_8_and_16_bit_planes() { - let mut p8 = b"P5\n# comment\n2 2\n255\n".to_vec(); - p8.extend_from_slice(&[0, 128, 255, 64]); - let s = Still::from_pgm(&p8).expect("8-bit parse"); - assert_eq!((s.width, s.height), (2, 2)); - assert!((s.luma[1] - 128.0 / 255.0).abs() < 1e-6); - let mut p16 = b"P5 2 1 65535\n".to_vec(); - p16.extend_from_slice(&[0xFF, 0xFF, 0x00, 0x00]); - let s = Still::from_pgm(&p16).expect("16-bit parse"); - assert_eq!(s.luma, vec![1.0, 0.0]); + let mut pgm_8bit = b"P5\n# comment\n2 2\n255\n".to_vec(); + pgm_8bit.extend_from_slice(&[0, 128, 255, 64]); + let parsed = Still::from_pgm(&pgm_8bit).expect("8-bit parse"); + assert_eq!((parsed.width, parsed.height), (2, 2)); + assert!((parsed.luma[1] - 128.0 / 255.0).abs() < 1e-6); + + let mut pgm_16bit = b"P5 2 1 65535\n".to_vec(); + pgm_16bit.extend_from_slice(&[0xFF, 0xFF, 0x00, 0x00]); + let parsed = Still::from_pgm(&pgm_16bit).expect("16-bit parse"); + assert_eq!(parsed.luma, vec![1.0, 0.0]); } #[test] fn a_header_claiming_more_data_than_is_present_is_rejected() { - // The header claims a 100x100 plane, but only 4 sample bytes - // follow. Also exercises the overflow-safe path: `width * - // height` for a header this large already overflows `u32`. - let mut p8 = b"P5\n100 100\n255\n".to_vec(); - p8.extend_from_slice(&[0, 128, 255, 64]); - let err = Still::from_pgm(&p8).expect_err("truncated data must be rejected"); - assert!(err.contains("truncated"), "got {err}"); + let mut pgm = b"P5\n100 100\n255\n".to_vec(); + pgm.extend_from_slice(&[0, 128, 255, 64]); + let error = Still::from_pgm(&pgm).expect_err("truncated data must be rejected"); + assert!(error.contains("truncated"), "got {error}"); } #[test] fn a_header_with_dimensions_that_overflow_u32_is_rejected_not_wrapped() { - // `width * height` overflows `u32` here; widening to `u64` - // before multiplying must catch this as truncated data rather - // than wrapping to a small value that a short buffer satisfies. - let p8 = b"P5\n70000 70000\n255\n".to_vec(); - let err = Still::from_pgm(&p8).expect_err("an overflowing header must be rejected"); - assert!(err.contains("truncated"), "got {err}"); + let pgm = b"P5\n70000 70000\n255\n".to_vec(); + let error = Still::from_pgm(&pgm).expect_err("an overflowing header must be rejected"); + assert!(error.contains("truncated"), "got {error}"); } } diff --git a/av-denoise-core/src/nl4d/kernels/grain.rs b/av-denoise-core/src/nl4d/kernels/grain.rs index dd4777f..aa98b0c 100644 --- a/av-denoise-core/src/nl4d/kernels/grain.rs +++ b/av-denoise-core/src/nl4d/kernels/grain.rs @@ -19,20 +19,24 @@ use crate::nl4d::grain::consts::{ STRENGTH_GROUPS, }; +/// Threads in one measuring cube, one per pixel of a cell. const THREADS: u32 = 64; +/// The shared tile width, one cell plus [HALO_X] pixels each side. const TILE_W: u32 = 20; +/// The shared tile height, one cell plus the 3 lag rows below it. const TILE_H: u32 = 11; const TILE_LEN: u32 = TILE_W * TILE_H; +/// The largest horizontal lag. const HALO_X: u32 = 6; const LAG_LANES: u32 = 46; -/// Halving rounds that reduce `THREADS` values to one. +/// Halving rounds that reduce [THREADS] values to one. const REDUCE_ROUNDS: u32 = THREADS.ilog2(); -/// Halving rounds that reduce `REDUCE_THREADS` values to one. +/// Halving rounds that reduce [REDUCE_THREADS] values to one. const CHUNK_ROUNDS: u32 = REDUCE_THREADS.ilog2(); /// Shared lanes for sums of the grain, its square and the clean luma. const SUM_LANES: u32 = 3; -/// Copies one neighbour's motion vectors and confidence into a saved ring entry. +/// Copies one neighbour's motion vectors and confidence into saved ring entry `entry`. /// /// `mv_offset` and `conf_offset` are the neighbour's element offsets in the motion field and /// confidence. Each thread strides by `total_threads`, so a clamped grid still covers every block. @@ -61,7 +65,8 @@ pub fn grain_save_vectors( /// The lag lane of a half-plane offset, in the order of [LAGS](crate::nl4d::grain::consts::LAGS). /// -/// `dx_index` runs between 0 and 12 and is the horizontal offset plus [HALO_X]. +/// `dx_index` runs between 0 and 12 and is the horizontal offset plus [HALO_X]. Only half the plane +/// is measured because a lag and its mirror share one autocovariance. const fn lag_lane(dy: u32, dx_index: u32) -> u32 { if dy == 0 { dx_index - HALO_X @@ -93,6 +98,7 @@ fn clamp_shift(value: u32, shift: i32, #[comptime] limit: u32) -> u32 { clamped as u32 } +/// The motion block of a cell, clamped into the block grid. #[cube] fn block_for( cell_x: u32, @@ -157,7 +163,7 @@ fn luma_bin(mean: f32) -> u32 { u32::min(scaled, comptime!(LUMA_BINS as u32 - 1)) } -/// The sample standard deviation of `THREADS` values from their sum and sum of squares. +/// The sample standard deviation of [THREADS] values from their sum and sum of squares. #[cube] fn std_of(sum: f32, sum_sq: f32) -> f32 { let count = THREADS as f32; @@ -167,11 +173,23 @@ fn std_of(sum: f32, sum_sq: f32) -> f32 { /// Measures one completed frame's source grain and kept grain over one 8x8 cell. /// -/// Source grain is the motion-compensated difference of noisy frames `t` and `t + 1`, and kept -/// grain the same between outputs `t - 1` and `t`, both divided by sqrt 2. Each accepted cell adds -/// one count to its histogram. An accepted source cell writes its 46 lag sums, pixel count and -/// strength group to `partials`, and every other cell writes zeros there. The lag sums take the -/// cell's mean source grain off every pixel and halo neighbour first. +/// Launch one `CELL x CELL` cube per cell. Source grain is the motion-compensated difference of the +/// noisy ring slots `slot_t` and `slot_next`. Kept grain is the same between the denoised frames +/// `out_prev` at `t - 1` and `out_t` at `t`. Both are divided by sqrt 2 so they carry one frame's +/// grain std. `source_entry` and `kept_entry` pick each difference's saved vectors and confidence, +/// and a zero `has_source` or `has_kept` rejects every cell of that measurement. +/// +/// A source cell must sit off the left, right and bottom edges, so its lag halo never reads a +/// clamped pixel. It also needs a confident vector, a flat, mid-luma `out_t`, no source pixel of +/// `t` or `t + 1` near clipping, and a grain std above the floor. A kept cell passes the same +/// gates, apart from the edge gate, with flatness, luma and clipping all read from `out_prev`. +/// +/// Each accepted cell adds one count to its luma bin and std bucket of `hist`, the first +/// `HIST_LEN` counts for source grain and the next for kept grain. Every cell writes its partial, +/// the 46 lag sums, pixel count and strength group of its source grain, as zeros when rejected. +/// The lag sums take the cell's mean source grain off every pixel and halo neighbour first, so a +/// uniform brightness flicker between the two frames never reaches them. The strength group keeps +/// cells of outlying grain strength in a record of their own. #[cube(launch_unchecked)] #[expect( clippy::too_many_arguments, @@ -211,26 +229,26 @@ pub fn grain_measure( let blocks = comptime!(blocks_x * blocks_y); let local_x = UNIT_POS_X; let local_y = UNIT_POS_Y; - let tid = local_y * CELL + local_x; + let thread_id = local_y * CELL + local_x; let cell_x = CUBE_POS_X; let cell_y = CUBE_POS_Y; - let x0 = cell_x * CELL; - let y0 = cell_y * CELL; - let x = x0 + local_x; - let y = y0 + local_y; + let origin_x = cell_x * CELL; + let origin_y = cell_y * CELL; + let x = origin_x + local_x; + let y = origin_y + local_y; let block = block_for(cell_x, cell_y, blocks_x, blocks_y, step); let source_base = source_entry * blocks * 2; let kept_base = kept_entry * blocks * 2; let mut group = 0u32; // Source grain over the cell and its halo, every coordinate clamped into the frame. - let mut fill = tid; + let mut fill = thread_id; while fill < TILE_LEN { let tile_x = fill % TILE_W; let tile_y = fill / TILE_W; - let raw_x = x0 as i32 + tile_x as i32 - HALO_X as i32; + let raw_x = origin_x as i32 + tile_x as i32 - HALO_X as i32; let read_x = clamp_shift(0u32, raw_x, width); - let read_y = clamp_shift(y0 + tile_y, 0i32, height); + let read_y = clamp_shift(origin_y + tile_y, 0i32, height); tile[fill as usize] = source_grain_at( input, saved_mv, @@ -270,30 +288,30 @@ pub fn grain_measure( let next_y = clamp_shift(y, dy, height); let next = luma_at(input, next_x, next_y, slot_next, width, height); - sums[tid as usize] = own_grain; - sums[(THREADS + tid) as usize] = own_grain * own_grain; - sums[(2 * THREADS + tid) as usize] = clean; - lows[tid as usize] = clean; - lows[(THREADS + tid) as usize] = f32::min(current, next); - highs[tid as usize] = clean; - highs[(THREADS + tid) as usize] = f32::max(current, next); + sums[thread_id as usize] = own_grain; + sums[(THREADS + thread_id) as usize] = own_grain * own_grain; + sums[(2 * THREADS + thread_id) as usize] = clean; + lows[thread_id as usize] = clean; + lows[(THREADS + thread_id) as usize] = f32::min(current, next); + highs[thread_id as usize] = clean; + highs[(THREADS + thread_id) as usize] = f32::max(current, next); sync_cube(); #[unroll] for round in 0..REDUCE_ROUNDS { let stride = comptime!(THREADS >> (round + 1)); - if tid < stride { + if thread_id < stride { #[unroll] for lane in 0..SUM_LANES { - let here = (lane * THREADS + tid) as usize; - let there = (lane * THREADS + tid + stride) as usize; + let here = (lane * THREADS + thread_id) as usize; + let there = (lane * THREADS + thread_id + stride) as usize; sums[here] = sums[here] + sums[there]; } #[unroll] for lane in 0..2u32 { - let here = (lane * THREADS + tid) as usize; - let there = (lane * THREADS + tid + stride) as usize; + let here = (lane * THREADS + thread_id) as usize; + let there = (lane * THREADS + thread_id + stride) as usize; lows[here] = f32::min(lows[here], lows[there]); highs[here] = f32::max(highs[here], highs[there]); } @@ -302,16 +320,16 @@ pub fn grain_measure( sync_cube(); } - if tid == 0 { + if thread_id == 0 { let count = THREADS as f32; let grain_std = std_of(sums[0], sums[THREADS as usize]); let mean = sums[(2 * THREADS) as usize] / count; let range = highs[0] - lows[0]; let interior = cell_x >= 1 && cell_x + 2 <= cells_x && cell_y + 2 <= cells_y; - let conf = saved_conf[(source_entry * blocks + block) as usize]; - let ok = has_source != 0 + let confidence = saved_conf[(source_entry * blocks + block) as usize]; + let passes = has_source != 0 && interior - && conf >= CONF_MIN + && confidence >= CONF_MIN && range < FLAT_RANGE && mean > LUMA_LOW && mean < LUMA_HIGH @@ -320,7 +338,7 @@ pub fn grain_measure( && grain_std > STD_MIN; accepted[0] = 0.0f32; - if ok { + if passes { accepted[0] = 1.0f32; let bucket = bucket_for(grain_std, edges); let slot = luma_bin(mean) * STD_BUCKETS as u32 + bucket; @@ -344,7 +362,7 @@ pub fn grain_measure( let lane = comptime!(lag_lane(dy, dx_index)); let neighbour_index = (local_y + dy) * TILE_W + local_x + dx_index; let neighbour = tile[neighbour_index as usize] - grain_mean; - lags[(lane * THREADS + tid) as usize] = take * centre * neighbour; + lags[(lane * THREADS + thread_id) as usize] = take * centre * neighbour; } } } @@ -354,11 +372,11 @@ pub fn grain_measure( #[unroll] for round in 0..REDUCE_ROUNDS { let lag_stride = comptime!(THREADS >> (round + 1)); - if tid < lag_stride { + if thread_id < lag_stride { #[unroll] for index in 0..LAG_LANES { - let here = (index * THREADS + tid) as usize; - let there = (index * THREADS + tid + lag_stride) as usize; + let here = (index * THREADS + thread_id) as usize; + let there = (index * THREADS + thread_id + lag_stride) as usize; lags[here] = lags[here] + lags[there]; } } @@ -366,7 +384,7 @@ pub fn grain_measure( sync_cube(); } - if tid == 0 { + if thread_id == 0 { let cell_index = cell_y * cells_x + cell_x; let base = cell_index * PARTIAL_LEN as u32; @@ -388,26 +406,26 @@ pub fn grain_measure( let previous = out_prev[((y * width + x) * stored_ch) as usize]; let kept = (warped - previous) * std::f32::consts::FRAC_1_SQRT_2; - sums[tid as usize] = kept; - sums[(THREADS + tid) as usize] = kept * kept; - sums[(2 * THREADS + tid) as usize] = previous; - lows[tid as usize] = previous; - highs[tid as usize] = previous; + sums[thread_id as usize] = kept; + sums[(THREADS + thread_id) as usize] = kept * kept; + sums[(2 * THREADS + thread_id) as usize] = previous; + lows[thread_id as usize] = previous; + highs[thread_id as usize] = previous; sync_cube(); #[unroll] for round in 0..REDUCE_ROUNDS { let kept_stride = comptime!(THREADS >> (round + 1)); - if tid < kept_stride { + if thread_id < kept_stride { #[unroll] for index in 0..SUM_LANES { - let here = (index * THREADS + tid) as usize; - let there = (index * THREADS + tid + kept_stride) as usize; + let here = (index * THREADS + thread_id) as usize; + let there = (index * THREADS + thread_id + kept_stride) as usize; sums[here] = sums[here] + sums[there]; } - let here = tid as usize; - let there = (tid + kept_stride) as usize; + let here = thread_id as usize; + let there = (thread_id + kept_stride) as usize; lows[here] = f32::min(lows[here], lows[there]); highs[here] = f32::max(highs[here], highs[there]); } @@ -415,15 +433,15 @@ pub fn grain_measure( sync_cube(); } - if tid == 0 { + if thread_id == 0 { let count = THREADS as f32; let kept_std = std_of(sums[0], sums[THREADS as usize]); let mean = sums[(2 * THREADS) as usize] / count; let low = lows[0]; let high = highs[0]; - let conf = saved_conf[(kept_entry * blocks + block) as usize]; - let ok = has_kept != 0 - && conf >= CONF_MIN + let confidence = saved_conf[(kept_entry * blocks + block) as usize]; + let passes = has_kept != 0 + && confidence >= CONF_MIN && high - low < FLAT_RANGE && mean > LUMA_LOW && mean < LUMA_HIGH @@ -431,7 +449,7 @@ pub fn grain_measure( && high <= CLIP_HIGH && kept_std > STD_MIN; - if ok { + if passes { let bucket = bucket_for(kept_std, edges); let slot = HIST_LEN as u32 + luma_bin(mean) * STD_BUCKETS as u32 + bucket; Atomic::fetch_add(&hist[slot as usize], 1i32); @@ -439,27 +457,29 @@ pub fn grain_measure( } } -/// Adds every cell's partial for one lane into its strength group's chunk record, one cube per lane. +/// Adds one lane of every cell's partial into its strength group's record in `chunk`. /// -/// Each thread adds its cells into its own column of a per-group scratch, and each group's column -/// sums then reduce to one value. +/// Launch one cube of `REDUCE_THREADS` threads per lane, `AUTOCOV_LEN` cubes in all. `chunk` holds +/// one record per strength group and keeps its earlier sums. Each thread adds its cells into its +/// own column of a per-group scratch, and the columns then reduce to one value per group. A cell's +/// group is clamped to the last group, so a bad value never indexes past the scratch. #[cube(launch_unchecked)] pub fn grain_reduce_partials(partials: &Array, chunk: &mut Array, cells: u32) { let mut scratch = SharedMemory::::new((STRENGTH_GROUPS as u32 * REDUCE_THREADS) as usize); let lane = CUBE_POS_X; - let tid = UNIT_POS_X; + let thread_id = UNIT_POS_X; #[unroll] for group in 0..STRENGTH_GROUPS as u32 { - scratch[(group * REDUCE_THREADS + tid) as usize] = 0.0f32; + scratch[(group * REDUCE_THREADS + thread_id) as usize] = 0.0f32; } - let mut cell = tid; + let mut cell = thread_id; while cell < cells { let base = cell * PARTIAL_LEN as u32; let raw_group = u32::cast_from(partials[(base + AUTOCOV_LEN as u32) as usize]); let group = u32::min(raw_group, comptime!(STRENGTH_GROUPS as u32 - 1)); - let slot = (group * REDUCE_THREADS + tid) as usize; + let slot = (group * REDUCE_THREADS + thread_id) as usize; scratch[slot] = scratch[slot] + partials[(base + lane) as usize]; cell += REDUCE_THREADS; } @@ -469,11 +489,11 @@ pub fn grain_reduce_partials(partials: &Array, chunk: &mut Array, cell #[unroll] for round in 0..CHUNK_ROUNDS { let stride = comptime!(REDUCE_THREADS >> (round + 1)); - if tid < stride { + if thread_id < stride { #[unroll] for group in 0..STRENGTH_GROUPS as u32 { - let here = (group * REDUCE_THREADS + tid) as usize; - let there = (group * REDUCE_THREADS + tid + stride) as usize; + let here = (group * REDUCE_THREADS + thread_id) as usize; + let there = (group * REDUCE_THREADS + thread_id + stride) as usize; scratch[here] = scratch[here] + scratch[there]; } } @@ -481,8 +501,8 @@ pub fn grain_reduce_partials(partials: &Array, chunk: &mut Array, cell sync_cube(); } - if tid < STRENGTH_GROUPS as u32 { - let target = (tid * AUTOCOV_LEN as u32 + lane) as usize; - chunk[target] = chunk[target] + scratch[(tid * REDUCE_THREADS) as usize]; + if thread_id < STRENGTH_GROUPS as u32 { + let target = (thread_id * AUTOCOV_LEN as u32 + lane) as usize; + chunk[target] = chunk[target] + scratch[(thread_id * REDUCE_THREADS) as usize]; } } diff --git a/av-denoise-core/src/nl4d/kernels/mod.rs b/av-denoise-core/src/nl4d/kernels/mod.rs index 19ba181..f63222f 100644 --- a/av-denoise-core/src/nl4d/kernels/mod.rs +++ b/av-denoise-core/src/nl4d/kernels/mod.rs @@ -1,5 +1,3 @@ -//! GPU kernels that belong to nl4d alone. - #![doc(hidden)] mod grain; diff --git a/av-denoise-core/src/nl4d/kernels/regularise.rs b/av-denoise-core/src/nl4d/kernels/regularise.rs index 1cc1abf..4c57885 100644 --- a/av-denoise-core/src/nl4d/kernels/regularise.rs +++ b/av-denoise-core/src/nl4d/kernels/regularise.rs @@ -1,7 +1,7 @@ use cubecl::prelude::*; -/// How many vectors a block considers, its own, the neighbourhood -/// median, the four adjacent blocks' and zero. +/// How many vectors each block scores, its own, the neighbourhood median, its four adjacent +/// blocks' and zero. pub const REGULARISE_CANDIDATES: u32 = 7; /// The most neighbours a block has in its 3x3 neighbourhood. @@ -15,6 +15,7 @@ fn clamp_coord(value: i32, limit: i32) -> i32 { } else if value >= limit { result = limit - 1; } + result } @@ -24,46 +25,41 @@ fn abs_i32(value: i32) -> i32 { if value < 0 { result = -value; } + result } -/// Sorts the first `n` entries of `vals` in place and returns the lower -/// median. +/// Sorts the first `count` entries of `vals` in place and returns the lower median. #[cube] -fn median_of(vals: &mut Array, n: u32) -> i32 { +fn median_of(vals: &mut Array, count: u32) -> i32 { let mut i: u32 = 1; - while i < n { + while i < count { let key = vals[i as usize]; let mut j = i; while j > 0u32 && vals[(j - 1u32) as usize] > key { vals[j as usize] = vals[(j - 1u32) as usize]; j -= 1u32; } + vals[j as usize] = key; i += 1u32; } - vals[((n - 1u32) / 2u32) as usize] + + vals[((count - 1u32) / 2u32) as usize] } -/// Re-scores one block's motion vector against its neighbourhood and -/// writes the winner, with a fresh confidence, to the output field. +/// Re-scores each block's motion vector against its 3x3 neighbourhood median. /// -/// One cube handles one block of the field. Thread 0 gathers the 3x3 -/// neighbourhood's vectors from `mv_in`, takes their component-wise -/// median, and lays out the candidates in shared memory. The block's -/// own vector is candidate 0. Each of the next threads scores one -/// candidate by SAD over the block on the level-0 luma planes, plus -/// `lambda_pixel` times the candidate's distance from the median in -/// pixels. Thread 0 then picks the lowest cost, and a tie keeps the -/// earlier candidate, so the block's own vector wins every tie. +/// Launch one cube per block with at least `REGULARISE_CANDIDATES` threads. `centre` and +/// `neighbour` are full-resolution luma planes. Each candidate costs its SAD over the block plus +/// `lambda_pixel` times its distance in pixels from the component-wise median of the neighbours. +/// The lowest cost wins, and a tie keeps the earlier candidate, so the block's own vector wins +/// every tie. /// -/// The winner's confidence is derived from its SAD exactly as -/// `nlm_mc_block_match_fine` derives it, with the same -/// `sad_noise_floor` and `thsad`. +/// The winner goes to `mv_out` and a confidence from its SAD to `confidence_out`. The confidence +/// is 1 up to `sad_noise_floor` and falls to 0 at `thsad` past it. /// -/// `mv_in` and `mv_out` are separate buffers. Every block reads the -/// whole input before any block's output exists, so the result does -/// not depend on block order. +/// `mv_in` and `mv_out` must be separate buffers, so the result does not depend on block order. #[cube(launch_unchecked)] #[expect( clippy::too_many_arguments, @@ -85,153 +81,179 @@ pub fn nl4d_mv_regularise( #[comptime] blocks_x: u32, #[comptime] blocks_y: u32, ) { - let bx = CUBE_POS_X; - let by = CUBE_POS_Y; - let block = by * blocks_x + bx; + let block_col = CUBE_POS_X; + let block_row = CUBE_POS_Y; + let block = block_row * blocks_x + block_col; let local_x = UNIT_POS_X; let local_y = UNIT_POS_Y; let thread_id = local_y * CUBE_DIM_X + local_x; let block_pixels = comptime!(blksize * blksize); let mut centre_smem = SharedMemory::::new(block_pixels as usize); - let mut cand = SharedMemory::::new(comptime!(2 * REGULARISE_CANDIDATES) as usize); + let mut candidates = SharedMemory::::new(comptime!(2 * REGULARISE_CANDIDATES) as usize); let mut median = SharedMemory::::new(2usize); let mut sad_scratch = SharedMemory::::new(REGULARISE_CANDIDATES as usize); let mut cost = SharedMemory::::new(REGULARISE_CANDIDATES as usize); - let block_origin_x = bx as i32 * step as i32; - let block_origin_y = by as i32 * step as i32; + let block_origin_x = block_col as i32 * step as i32; + let block_origin_y = block_row as i32 * step as i32; // The centre tile, loaded once and shared by every candidate. - let mut py = local_y; - while py < blksize { - let mut px = local_x; - while px < blksize { - let cx = clamp_coord(block_origin_x + px as i32, width as i32); - let cy = clamp_coord(block_origin_y + py as i32, height as i32); - centre_smem[(py * blksize + px) as usize] = centre[(cy * width as i32 + cx) as usize]; - px += CUBE_DIM_X; + let mut pixel_y = local_y; + while pixel_y < blksize { + let mut pixel_x = local_x; + while pixel_x < blksize { + let centre_x = clamp_coord(block_origin_x + pixel_x as i32, width as i32); + let centre_y = clamp_coord(block_origin_y + pixel_y as i32, height as i32); + centre_smem[(pixel_y * blksize + pixel_x) as usize] = + centre[(centre_y * width as i32 + centre_x) as usize]; + pixel_x += CUBE_DIM_X; } - py += CUBE_DIM_Y; + + pixel_y += CUBE_DIM_Y; } if thread_id == 0u32 { - let mut xs = Array::::new(NEIGHBOURHOOD as usize); - let mut ys = Array::::new(NEIGHBOURHOOD as usize); - let mut n: u32 = 0; + let mut neighbour_xs = Array::::new(NEIGHBOURHOOD as usize); + let mut neighbour_ys = Array::::new(NEIGHBOURHOOD as usize); + let mut neighbour_count: u32 = 0; let mut dy: u32 = 0; while dy < 3u32 { let mut dx: u32 = 0; while dx < 3u32 { if dx != 1u32 || dy != 1u32 { - let nx = bx as i32 + dx as i32 - 1i32; - let ny = by as i32 + dy as i32 - 1i32; - if nx >= 0 && ny >= 0 && nx < blocks_x as i32 && ny < blocks_y as i32 { - let idx = ((ny as u32 * blocks_x + nx as u32) * 2u32) as usize; - xs[n as usize] = mv_in[idx]; - ys[n as usize] = mv_in[idx + 1]; - n += 1u32; + let neighbour_x = block_col as i32 + dx as i32 - 1i32; + let neighbour_y = block_row as i32 + dy as i32 - 1i32; + if neighbour_x >= 0 + && neighbour_y >= 0 + && neighbour_x < blocks_x as i32 + && neighbour_y < blocks_y as i32 + { + let mv_index = ((neighbour_y as u32 * blocks_x + neighbour_x as u32) * 2u32) as usize; + neighbour_xs[neighbour_count as usize] = mv_in[mv_index]; + neighbour_ys[neighbour_count as usize] = mv_in[mv_index + 1]; + neighbour_count += 1u32; } } + dx += 1u32; } + dy += 1u32; } + let own_x = mv_in[(block * 2u32) as usize]; let own_y = mv_in[(block * 2u32 + 1u32) as usize]; + // A block with no neighbours is its own median. - let mut mx = own_x; - let mut my = own_y; - if n > 0u32 { - mx = median_of(&mut xs, n); - my = median_of(&mut ys, n); + let mut median_x = own_x; + let mut median_y = own_y; + if neighbour_count > 0u32 { + median_x = median_of(&mut neighbour_xs, neighbour_count); + median_y = median_of(&mut neighbour_ys, neighbour_count); } - median[0] = mx; - median[1] = my; - - cand[0] = own_x; - cand[1] = own_y; - cand[2] = mx; - cand[3] = my; - // Left, right, up, down. Off-grid neighbours repeat the block's - // own vector, which the tie rule then discards. - let mut c: u32 = 2; + + median[0] = median_x; + median[1] = median_y; + + candidates[0] = own_x; + candidates[1] = own_y; + candidates[2] = median_x; + candidates[3] = median_y; + + // Left, right, up, down. Off-grid neighbours repeat the block's own vector, which the tie + // rule then discards. + let mut candidate: u32 = 2; let mut side: u32 = 0; while side < 4u32 { - let mut nx = bx as i32; - let mut ny = by as i32; + let mut neighbour_x = block_col as i32; + let mut neighbour_y = block_row as i32; if side == 0u32 { - nx -= 1; + neighbour_x -= 1; } else if side == 1u32 { - nx += 1; + neighbour_x += 1; } else if side == 2u32 { - ny -= 1; + neighbour_y -= 1; } else { - ny += 1; + neighbour_y += 1; } - let mut vx = own_x; - let mut vy = own_y; - if nx >= 0 && ny >= 0 && nx < blocks_x as i32 && ny < blocks_y as i32 { - let idx = ((ny as u32 * blocks_x + nx as u32) * 2u32) as usize; - vx = mv_in[idx]; - vy = mv_in[idx + 1]; + + let mut vector_x = own_x; + let mut vector_y = own_y; + if neighbour_x >= 0 + && neighbour_y >= 0 + && neighbour_x < blocks_x as i32 + && neighbour_y < blocks_y as i32 + { + let mv_index = ((neighbour_y as u32 * blocks_x + neighbour_x as u32) * 2u32) as usize; + vector_x = mv_in[mv_index]; + vector_y = mv_in[mv_index + 1]; } - cand[(c * 2u32) as usize] = vx; - cand[(c * 2u32 + 1u32) as usize] = vy; - c += 1u32; + + candidates[(candidate * 2u32) as usize] = vector_x; + candidates[(candidate * 2u32 + 1u32) as usize] = vector_y; + candidate += 1u32; side += 1u32; } - cand[(c * 2u32) as usize] = 0; - cand[(c * 2u32 + 1u32) as usize] = 0; + + candidates[(candidate * 2u32) as usize] = 0; + candidates[(candidate * 2u32 + 1u32) as usize] = 0; } + sync_cube(); if thread_id < REGULARISE_CANDIDATES { - let mvx = cand[(thread_id * 2u32) as usize]; - let mvy = cand[(thread_id * 2u32 + 1u32) as usize]; + let mvx = candidates[(thread_id * 2u32) as usize]; + let mvy = candidates[(thread_id * 2u32 + 1u32) as usize]; let mut sad: f32 = 0.0; for iy in 0..blksize { for ix in 0..blksize { - let cx = block_origin_x + ix as i32; - let cy = block_origin_y + iy as i32; + let centre_x = block_origin_x + ix as i32; + let centre_y = block_origin_y + iy as i32; let centre_val = centre_smem[(iy * blksize + ix) as usize]; - let nx = clamp_coord(cx + mvx, width as i32); - let ny = clamp_coord(cy + mvy, height as i32); - let diff = centre_val - neighbour[(ny * width as i32 + nx) as usize]; + let neighbour_x = clamp_coord(centre_x + mvx, width as i32); + let neighbour_y = clamp_coord(centre_y + mvy, height as i32); + let diff = centre_val - neighbour[(neighbour_y * width as i32 + neighbour_x) as usize]; let abs_diff = if diff < 0.0f32 { -diff } else { diff }; sad += abs_diff; } } + let deviation = abs_i32(mvx - median[0]) + abs_i32(mvy - median[1]); sad_scratch[thread_id as usize] = sad; cost[thread_id as usize] = sad + lambda_pixel * deviation as f32; } + sync_cube(); if thread_id == 0u32 { let mut best: u32 = 0; let mut best_cost = cost[0]; - let mut c: u32 = 1; - while c < REGULARISE_CANDIDATES { - if cost[c as usize] < best_cost { - best_cost = cost[c as usize]; - best = c; + let mut candidate: u32 = 1; + while candidate < REGULARISE_CANDIDATES { + if cost[candidate as usize] < best_cost { + best_cost = cost[candidate as usize]; + best = candidate; } - c += 1u32; + + candidate += 1u32; } - mv_out[(block * 2u32) as usize] = cand[(best * 2u32) as usize]; - mv_out[(block * 2u32 + 1u32) as usize] = cand[(best * 2u32 + 1u32) as usize]; + + mv_out[(block * 2u32) as usize] = candidates[(best * 2u32) as usize]; + mv_out[(block * 2u32 + 1u32) as usize] = candidates[(best * 2u32 + 1u32) as usize]; let mut excess = sad_scratch[best as usize] - sad_noise_floor; if excess < 0.0f32 { excess = 0.0f32; } + let thsad_sq = thsad * thsad; let excess_sq = excess * excess; let mut confidence = (thsad_sq - excess_sq) / (thsad_sq + excess_sq); if confidence < 0.0f32 { confidence = 0.0f32; } + confidence_out[block as usize] = confidence; } } diff --git a/av-denoise-core/src/nl4d/mod.rs b/av-denoise-core/src/nl4d/mod.rs index 141653c..e1c6f59 100644 --- a/av-denoise-core/src/nl4d/mod.rs +++ b/av-denoise-core/src/nl4d/mod.rs @@ -1,29 +1,40 @@ -//! nl4d groups patches across several noisy frames rather than within a -//! single one. +//! Collaborative denoising across a window of frames //! -//! [`crate::collab`] groups similar 8x8 patches within one frame and -//! denoises each group jointly. This module extends that search across a -//! motion-compensated window of frames. +//! Similar 8x8 patches are grouped and filtered together. Each group holds a few patches from the +//! centre frame, each followed through its neighbour frames along its motion vector. Grain differs +//! from frame to frame while texture stays, so a transform along time separates the two. //! -//! Each group is a few centre-frame patches, each followed through its -//! best-matching neighbour frames along its own motion vector. Patches -//! followed through time carry independent grain, so the transform along -//! time separates it from the texture they share. +//! - [Nl4d], the engine +//! - [Nl4dOptions] and the per-preset defaults it is built from +//! - film grain measurement for AV1 grain tables -mod denoiser; +pub(crate) mod denoiser; +mod engine; pub mod grain; -pub mod harness; -pub mod kernels; -mod params; +pub(crate) mod harness; +pub(crate) mod kernels; +mod options; +pub(crate) mod params; mod regularise; -mod snapshot; +pub(crate) mod snapshot; -// Every test in this tree runs against a real GPU runtime, see -// `tests::helpers::R`, so it only builds when a wgpu-backed feature is -// enabled. A cpu-only build skips it entirely. +// The tests run against a real GPU runtime, so they need a wgpu-backed feature. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] pub(crate) mod tests; -pub use self::denoiser::Nl4dDenoiser; -pub use self::params::{MAX_KAISER_BETA, Nl4dParams}; -pub use self::snapshot::MotionSnapshot; +#[cfg(all(test, any(feature = "vulkan", feature = "metal")))] +pub(crate) use self::denoiser::Nl4dDenoiser; +pub use self::engine::Nl4d; +pub(crate) use self::options::nl4d_pool_ratio; +#[cfg(all(test, any(feature = "vulkan", feature = "metal")))] +pub(crate) use self::options::resolve_params; +pub use self::options::{ + Nl4dOptions, + nl4d_default_lambda_ht, + nl4d_spatial_radius_for, + nl4d_temporal_radius_for, +}; +#[cfg(test)] +pub(crate) use self::params::MAX_KAISER_BETA; +pub(crate) use self::params::Nl4dParams; +pub(crate) use self::snapshot::MotionSnapshot; diff --git a/av-denoise-core/src/nl4d/options.rs b/av-denoise-core/src/nl4d/options.rs new file mode 100644 index 0000000..94f8736 --- /dev/null +++ b/av-denoise-core/src/nl4d/options.rs @@ -0,0 +1,238 @@ +use super::Nl4dParams; +use crate::error::Error; +use crate::nlmeans::{ChannelMode, HqParams, MotionSearch, NlmParams}; +use crate::options::Preset; + +/// Settings for [Nl4d](crate::Nl4d). +/// +/// The HQ front end runs only for its frame ring, motion field and noise estimate, so no NLM +/// weighting knobs appear here. Motion tracking is always on, because the grouping kernel reads +/// the motion field and confidence scores it produces. +#[derive(Debug, Copy, Clone, PartialEq)] +pub struct Nl4dOptions { + /// How motion between frames is tracked. + pub motion: MotionSearch, + /// A fixed noise standard deviation between 0 and 1, replacing the per-frame estimate. + /// + /// `None`, the default, measures the noise in each pushed frame and smooths it over time. + pub sigma: Option, + /// A multiplier on the measured noise level. Defaults to 1.0. + /// + /// It does nothing when `sigma` is set, because the estimator never runs. + pub sigma_scale: f32, + /// A multiplier on how much extra SAD a block tolerates before its confidence falls. Defaults + /// to 1.0. + pub thsad_scale: f32, + /// Half-width of the refine window around each neighbour frame's motion-predicted position, in + /// `1..=4`. Defaults to 2. + pub refine: u32, + /// Half-width of the spatial candidate window in the centre frame, in `1..=16`. Defaults to 9. + pub spatial_radius: u32, + /// Hard-threshold multiplier on the propagated coefficient sigma. + /// + /// Higher removes more noise and more fine detail. `None` resolves per plane through + /// [nl4d_default_lambda_ht]. + pub lambda_ht: Option, + /// A multiplier on the resolved `lambda_ht`, in `0.1..=10.0`. Defaults to 1.0. + /// + /// It scales an explicit `lambda_ht` and the per-plane default alike, so one value moves both + /// planes together. + pub lambda_ht_scale: f32, + /// The confidence floor below which a whole neighbour block is skipped, in `0.0..1.0`. Defaults + /// to 0.05. + /// + /// A volume left short of frames by the skip makes its group filter from the centre frame + /// alone. + pub c_min: f32, + /// The `beta` of the Kaiser window each filtered patch is tapered with as it is aggregated. + /// + /// Defaults to 2.0, and `0.0` is uniform aggregation. See + /// [Nl4dParams::kaiser_beta](crate::nl4d::Nl4dParams::kaiser_beta). + pub kaiser_beta: f32, + /// Estimates noise from each frame's own window instead of the whole stream's history. + /// + /// A VapourSynth filter must return the same pixels for a frame in any request order, and + /// history-dependent estimation breaks that under random access. Defaults to `false`, matching + /// every calibrated preset. + pub windowed_noise_estimation: bool, + /// See [Nl4dParams::field_lambda](crate::nl4d::Nl4dParams::field_lambda). + pub field_lambda: f32, + /// See [Nl4dParams::noise_map](crate::nl4d::Nl4dParams::noise_map). + pub noise_map: bool, + /// See [Nl4dParams::flat_boost](crate::nl4d::Nl4dParams::flat_boost). + pub flat_boost: f32, + /// See [Nl4dParams::chroma_flat_boost](crate::nl4d::Nl4dParams::chroma_flat_boost). + pub chroma_flat_boost: f32, + /// See [Nl4dParams::shadow_soften](crate::nl4d::Nl4dParams::shadow_soften). + pub shadow_soften: f32, + /// See [Nl4dParams::flat_texture_cut](crate::nl4d::Nl4dParams::flat_texture_cut). + pub flat_texture_cut: f32, + /// See [Nl4dParams::pooled_threshold](crate::nl4d::Nl4dParams::pooled_threshold). + pub pooled_threshold: bool, + /// Whether the denoiser measures the source's film grain for an AV1 grain table. + /// + /// Measured chunks stay on the GPU until drained, so callers drain after each flush. + pub grain_export: bool, + /// How many frames on each side of the centre frame the temporal window reaches, in `1..=8`. + pub temporal_radius: u32, +} + +impl Default for Nl4dOptions { + fn default() -> Self { + let defaults = Nl4dParams::default(); + let hq = HqParams::default(); + + Self { + motion: MotionSearch::default(), + sigma: hq.sigma_override, + sigma_scale: hq.sigma_scale, + thsad_scale: hq.thsad_scale, + refine: defaults.refine, + spatial_radius: defaults.spatial_radius, + // Resolved per plane at construction, once the plane is known. + lambda_ht: None, + lambda_ht_scale: 1.0, + c_min: defaults.c_min, + kaiser_beta: defaults.kaiser_beta, + windowed_noise_estimation: false, + field_lambda: defaults.field_lambda, + noise_map: defaults.noise_map, + flat_boost: defaults.flat_boost, + chroma_flat_boost: defaults.chroma_flat_boost, + shadow_soften: defaults.shadow_soften, + flat_texture_cut: defaults.flat_texture_cut, + pooled_threshold: defaults.pooled_threshold, + grain_export: defaults.grain_export, + temporal_radius: nl4d_temporal_radius_for(Preset::Base), + } + } +} + +/// The default `lambda_ht` for nl4d's hard-threshold stage, per plane. +/// +/// `lambda_ht` is how many standard deviations of estimated noise a transform coefficient must +/// clear to survive. Raising it removes more noise and more fine detail, so the value is a trade. +/// +/// The values were picked by eye on real film grain, accepting some lost detail for less remaining +/// noise on heavy grain. Encoders lose more detail to leftover grain than the denoiser does. +/// `ChannelMode::Yuv` uses the luma value, since a fused pass is dominated by luma. +pub fn nl4d_default_lambda_ht(channels: ChannelMode) -> f32 { + match channels { + ChannelMode::Luma | ChannelMode::Yuv => 4.158, + ChannelMode::Chroma => 3.234, + } +} + +/// The pooled threshold at each plane's default lambda, calibrated on real grain. +const NL4D_POOLED_THRESHOLD: f32 = 2.42; + +/// The ratio of nl4d's pooled threshold to its lambda for one plane. +/// +/// At the default lambda this gives the calibrated pooled threshold, and it scales with any +/// other lambda. +pub(crate) fn nl4d_pool_ratio(channels: ChannelMode) -> f32 { + NL4D_POOLED_THRESHOLD / nl4d_default_lambda_ht(channels) +} + +/// Resolves `lambda_ht` for one plane and applies `lambda_ht_scale`. +/// +/// The scale is range-checked here because [Nl4dParams] only sees the product, where a bad scale +/// would surface as a complaint about a `lambda_ht` the caller never set. +fn resolve_lambda_ht(opts: &Nl4dOptions, channels: ChannelMode) -> Result { + if !(opts.lambda_ht_scale.is_finite() && (0.1..=10.0).contains(&opts.lambda_ht_scale)) { + return Err(format!( + "lambda_ht_scale must be finite and in [0.1, 10.0], got {}", + opts.lambda_ht_scale + )); + } + + let lambda_ht = opts.lambda_ht.unwrap_or_else(|| nl4d_default_lambda_ht(channels)); + + Ok(lambda_ht * opts.lambda_ht_scale) +} + +/// How far the temporal window reaches at each preset. +/// +/// `veryfast` keeps a 1-frame window, because nl4d has nothing to group without neighbouring +/// frames. +pub fn nl4d_temporal_radius_for(preset: Preset) -> u32 { + match preset { + Preset::Veryfast | Preset::Fast => 1, + Preset::Base => 2, + Preset::Slow => 4, + Preset::Veryslow => 8, + } +} + +/// How wide the centre frame's candidate search is at each preset. +/// +/// `veryfast` shares its temporal radius with `fast`, so this is what separates them. The window +/// covers `(2 * radius + 1)^2` positions, so 6 searches a little under half the candidates 9 does. +/// Wider windows cost quadratically and have not been measured to help, so the slower presets keep +/// the default. +pub fn nl4d_spatial_radius_for(preset: Preset) -> u32 { + match preset { + Preset::Veryfast => 6, + Preset::Fast | Preset::Base | Preset::Slow | Preset::Veryslow => { + Nl4dOptions::default().spatial_radius + }, + } +} + +/// The front end's HQ parameters for these options. +/// +/// `temporal_confidence` is always on, because the grouping kernel reads the confidence scores. +fn hq_params(options: &Nl4dOptions) -> HqParams { + HqParams { + sigma_override: options.sigma, + sigma_scale: options.sigma_scale, + thsad_scale: options.thsad_scale, + temporal_confidence: true, + windowed_noise_estimation: options.windowed_noise_estimation, + ..HqParams::default() + } +} + +/// The front end parameters nl4d runs its motion and noise machinery with. +pub(crate) fn nlm_params(options: &Nl4dOptions, channels: ChannelMode) -> NlmParams { + let hq = hq_params(options); + + NlmParams { + channels, + motion_compensation: options.motion.into(), + temporal_radius: options.temporal_radius, + hq: Some(hq), + ..NlmParams::default() + } +} + +/// Resolves the options into an nl4d denoiser's parameters, with calibrated defaults for +/// `channels`. +pub(crate) fn resolve_params(options: &Nl4dOptions, channels: ChannelMode) -> Result { + if options.temporal_radius == 0 { + return Err(Error::InvalidOptions( + "nl4d needs a temporal radius of at least 1".to_string(), + )); + } + + let lambda_ht = resolve_lambda_ht(options, channels).map_err(Error::InvalidOptions)?; + let nlm = nlm_params(options, channels); + + Ok(Nl4dParams { + temporal_radius: options.temporal_radius, + nlm, + refine: options.refine, + spatial_radius: options.spatial_radius, + lambda_ht, + c_min: options.c_min, + kaiser_beta: options.kaiser_beta, + field_lambda: options.field_lambda, + noise_map: options.noise_map, + flat_boost: options.flat_boost, + chroma_flat_boost: options.chroma_flat_boost, + shadow_soften: options.shadow_soften, + flat_texture_cut: options.flat_texture_cut, + pooled_threshold: options.pooled_threshold, + grain_export: options.grain_export, + }) +} diff --git a/av-denoise-core/src/nl4d/params.rs b/av-denoise-core/src/nl4d/params.rs index d3c9a29..d3af157 100644 --- a/av-denoise-core/src/nl4d/params.rs +++ b/av-denoise-core/src/nl4d/params.rs @@ -1,89 +1,64 @@ +use crate::collab::MAX_TEMPORAL_RADIUS; use crate::nlmeans::{ChannelMode, HqParams, MotionCompensationMode, MotionEstimation, NlmParams}; -/// The largest [`Nl4dParams::kaiser_beta`] worth accepting. +/// The largest [Nl4dParams::kaiser_beta] worth accepting. /// -/// A Kaiser window's taps fall off faster the larger `beta` is. By 8 the -/// end tap is under a fiftieth of the centre, so a patch's edge pixels -/// contribute almost nothing and the step-4 grid is left covering each -/// pixel with a handful of centres rather than a blend. Past that the -/// window stops being a taper and starts being a mask, and the smallest -/// weights fall under what the fixed-point accumulators resolve. +/// By 8 the end tap is under a fiftieth of the centre, so a patch's edge pixels contribute almost +/// nothing and the step-4 grid covers each pixel with a handful of centres rather than a blend. +/// Past that the window acts as a mask, and the smallest weights fall under what the fixed-point +/// accumulators resolve. pub const MAX_KAISER_BETA: f32 = 8.0; /// The most motion blocks that may cover a reference patch on one axis. /// -/// A block grid at a step below `blksize` puts several blocks over one -/// patch, and -/// [`crate::collab::kernels::fused::collab_fused`] searches all of them. -/// It unrolls its per-neighbour duplicate-rectangle arrays over the -/// square of this bound, so the bound is what caps the shader's register -/// footprint. At 4 the arrays hold 16 rectangles and the shipped -/// geometry, `blksize = 16` at `overlap = 8`, uses 2. -/// -/// The step is `blksize - overlap`, so 4 admits an overlap of up to -/// three quarters of the block size. +/// The fused kernel searches every covering block and unrolls its duplicate-rectangle arrays over +/// the square of this bound, so the bound caps the shader's register footprint. The default +/// geometry, `blksize = 16` at `overlap = 8`, uses 2. The step is `blksize - overlap`, so 4 admits +/// an overlap of up to three quarters of the block size. pub const MAX_COVERING_BLOCKS: u32 = 4; -/// Tuning for [`super::Nl4dDenoiser`]. +/// Tuning for the nl4d denoiser. /// -/// `nlm` supplies the front end that builds the frame ring, the motion -/// field, and the confidence scores the temporal grouping reads. Its own -/// `temporal_radius` is overwritten at construction time from this -/// struct's own `temporal_radius`, so it does not need to be set by the -/// caller. +/// `nlm` configures the front end that builds the frame ring, motion field and confidence scores. +/// Its `temporal_radius` is overwritten from this struct's own at construction. #[derive(Debug, Clone)] pub struct Nl4dParams { - /// Machinery configuration for the front end. `hq` must be `Some` - /// with `temporal_confidence` on, and `motion_compensation` must be - /// active, because [`crate::nlmeans::NlmDenoiser::submit_machinery`] - /// only builds a ring view when both are on, and this denoiser is - /// built entirely on top of that call. + /// Configuration for the front end. + /// + /// `hq` must be `Some` with `temporal_confidence` on, and motion compensation must be active, + /// because the front end only builds a ring view when both are on. pub nlm: NlmParams, - /// How many frames on each side of the centre frame the temporal - /// search reaches into. In `1..=8`. + /// How many frames on each side of the centre frame the temporal search reaches, in `1..=8`. pub temporal_radius: u32, - /// Half-width of the refine window searched around each neighbour - /// frame's motion-predicted position. In `1..=4`. + /// Half-width of the refine window around each neighbour frame's motion-predicted position, in + /// `1..=4`. pub refine: u32, - /// Half-width of the spatial candidate window searched in the - /// centre frame. In `1..=16`. + /// Half-width of the spatial candidate window in the centre frame, in `1..=16`. pub spatial_radius: u32, /// Hard-threshold multiplier on the propagated coefficient sigma. - /// Higher shrinks more coefficients, so it removes more noise and - /// more fine detail. /// - /// Defaults to 4.158. Note that in reality luma and chroma want separately - /// tuned values. See [nl4d_default_lambda_ht](crate::nl4d_default_lambda_ht). + /// Higher removes more noise and more fine detail. Defaults to 4.158, though luma and chroma + /// want separately tuned values. See [nl4d_default_lambda_ht](crate::nl4d_default_lambda_ht). pub lambda_ht: f32, - /// The confidence floor below which a whole neighbour block is - /// skipped rather than scored, in `[0, 1)`. A block below the floor - /// is never scored, and a volume left short of frames by the skip - /// makes its group filter from the centre frame alone. - pub c_min: f32, - /// The `beta` of the Kaiser window each filtered patch is tapered - /// with as it is aggregated, in `0..=8`. + /// The confidence floor below which a whole neighbour block is skipped, in `0.0..1.0`. /// - /// A pixel is covered by many patches, each of which made its own - /// threshold decision. Tapering a patch toward its edges blends - /// those decisions rather than letting each reach its boundary at - /// full strength. Larger tapers harder. BM3D uses 2.0. + /// A volume left short of frames by the skip makes its group filter from the centre frame + /// alone. + pub c_min: f32, + /// The `beta` of the Kaiser window each filtered patch is tapered with as it is aggregated, in + /// `0.0..=8.0`. /// - /// Defaults to 2.0, BM3D's own value. `0.0` is exactly uniform - /// aggregation, which is what this did before the window existed. - /// See [`crate::collab::kernels::aggregate::kaiser_window`]. + /// A pixel is covered by many patches, each with its own threshold decision. Tapering each + /// patch toward its edges blends those decisions instead of letting each reach its boundary at + /// full strength. Defaults to 2.0, BM3D's own value, and `0.0` is uniform aggregation. pub kaiser_beta: f32, - /// The penalty on a block's vector deviating from its - /// neighbourhood's median, in the field regularisation pass. + /// The penalty on a block's vector deviating from its neighbourhood median, in the field + /// regularisation pass. /// - /// The pass re-scores each block's vector against the median of its - /// neighbours, the four adjacent blocks' vectors and zero, adding - /// this times the distance from the median, in pixels, scaled so - /// `1.0` weighs one pixel of deviation like a 5/255 per-pixel - /// mismatch. Defaults to `1.0`, calibrated with a `field_lambda` - /// sweep on the `mc_accuracy` bench. The pass gains most of its - /// accuracy by a moderate penalty and further increases add little, - /// so `1.0` sits inside that plateau rather than at its edge. `0.0` - /// skips the pass. + /// The median is taken over the four adjacent blocks' vectors and zero. The penalty is this + /// times the distance from the median in pixels, scaled so `1.0` weighs one pixel of deviation + /// like a 5/255 per-pixel mismatch. Defaults to `1.0`, calibrated on the `mc_accuracy` bench to + /// sit inside the plateau where larger values add little accuracy. `0.0` skips the pass. pub field_lambda: f32, /// Scales the luma threshold by how noisy each brightness level is in the current frame. /// @@ -160,9 +135,7 @@ impl Default for Nl4dParams { } impl Nl4dParams { - /// Rejects a configuration that would fail to launch, or that would - /// hit [`crate::nlmeans::NlmDenoiser::submit_machinery`]'s own - /// preconditions only once a real submit ran. + /// Rejects a configuration that would fail to launch or fail the front end's checks on submit. pub fn validate(&self) -> Result<(), String> { let Some(hq) = self.nlm.hq else { return Err( @@ -188,11 +161,8 @@ impl Nl4dParams { ); } - // Only checked once the geometry itself is sound. An overlap at - // or past blksize gives a step of 0, which `nlm.validate()` - // rejects on its own terms below with the real fault named. Left - // unguarded, that same case saturates the step to 1 here and - // reports a nonsensical covering-block count instead. + // An overlap at or past `blksize` is left to `nlm.validate()`, which names the real fault + // instead of a covering-block count computed from a saturated step. if let MotionCompensationMode::Mvtools { blksize, overlap, .. } = self.nlm.motion_compensation && overlap < blksize { @@ -208,11 +178,10 @@ impl Nl4dParams { } } - if !(1..=crate::collab::MAX_TEMPORAL_RADIUS).contains(&self.temporal_radius) { + if !(1..=MAX_TEMPORAL_RADIUS).contains(&self.temporal_radius) { return Err(format!( "temporal_radius={} must be in 1..={}", - self.temporal_radius, - crate::collab::MAX_TEMPORAL_RADIUS, + self.temporal_radius, MAX_TEMPORAL_RADIUS, )); } @@ -285,12 +254,14 @@ mod tests { #[test] fn validate_accepts_default() { - assert!(Nl4dParams::default().validate().is_ok()); + let params = Nl4dParams::default(); + assert!(params.validate().is_ok()); } #[test] fn the_noise_map_is_on_by_default() { - assert!(Nl4dParams::default().noise_map); + let params = Nl4dParams::default(); + assert!(params.noise_map); } #[test] @@ -326,15 +297,15 @@ mod tests { flat_boost: bad, ..Nl4dParams::default() }; - let err = flat.validate().expect_err("flat_boost out of range"); - assert!(err.contains("flat_boost"), "got {err}"); + let error = flat.validate().expect_err("flat_boost out of range"); + assert!(error.contains("flat_boost"), "got {error}"); let chroma = Nl4dParams { chroma_flat_boost: bad, ..Nl4dParams::default() }; - let err = chroma.validate().expect_err("chroma_flat_boost out of range"); - assert!(err.contains("chroma_flat_boost"), "got {err}"); + let error = chroma.validate().expect_err("chroma_flat_boost out of range"); + assert!(error.contains("chroma_flat_boost"), "got {error}"); } for bad in [0.05f32, 1.01, f32::NAN] { @@ -342,8 +313,8 @@ mod tests { shadow_soften: bad, ..Nl4dParams::default() }; - let err = params.validate().expect_err("shadow_soften out of range"); - assert!(err.contains("shadow_soften"), "got {err}"); + let error = params.validate().expect_err("shadow_soften out of range"); + assert!(error.contains("shadow_soften"), "got {error}"); } } @@ -365,32 +336,27 @@ mod tests { flat_texture_cut: bad, ..Nl4dParams::default() }; - let err = params.validate().expect_err("flat_texture_cut out of range"); - assert!(err.contains("flat_texture_cut"), "got {err}"); + let error = params.validate().expect_err("flat_texture_cut out of range"); + assert!(error.contains("flat_texture_cut"), "got {error}"); } } - /// A block geometry with `blksize / step` at or under - /// [`MAX_COVERING_BLOCKS`] is what the grouping kernel unrolls its - /// search over. - /// - /// The shipped geometry gives a step of 8 and so 2 covering blocks. - /// An overlap of three quarters of the block size gives a step of 4 - /// and exactly 4, the boundary. #[test] fn validate_accepts_block_geometries_up_to_the_covering_bound() { for (blksize, overlap, covers) in [(16u32, 8u32, 2u32), (16, 12, 4), (32, 24, 4), (8, 4, 2)] { + let motion_compensation = MotionCompensationMode::Mvtools { + blksize, + overlap, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Auto, + }; + let nlm = NlmParams { + motion_compensation, + ..Nl4dParams::default().nlm + }; let params = Nl4dParams { - nlm: NlmParams { - motion_compensation: MotionCompensationMode::Mvtools { - blksize, - overlap, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Auto, - }, - ..Nl4dParams::default().nlm - }, + nlm, ..Nl4dParams::default() }; assert!( @@ -400,54 +366,51 @@ mod tests { } } - /// Past the bound the kernel would unroll a far larger duplicate - /// check and hold far more rectangles in registers, so the - /// configuration is refused rather than compiled. #[test] fn validate_rejects_a_block_geometry_past_the_covering_bound() { for (blksize, overlap) in [(16u32, 13u32), (16, 14), (32, 31), (32, 25)] { + let motion_compensation = MotionCompensationMode::Mvtools { + blksize, + overlap, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Auto, + }; + let nlm = NlmParams { + motion_compensation, + ..Nl4dParams::default().nlm + }; let params = Nl4dParams { - nlm: NlmParams { - motion_compensation: MotionCompensationMode::Mvtools { - blksize, - overlap, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Auto, - }, - ..Nl4dParams::default().nlm - }, + nlm, ..Nl4dParams::default() }; - let err = params + let error = params .validate() .expect_err("a step this small should be rejected"); + let blksize_label = format!("blksize={blksize}"); + let overlap_label = format!("overlap={overlap}"); assert!( - err.contains(&format!("blksize={blksize}")) && err.contains(&format!("overlap={overlap}")), - "error should name the offending blksize and overlap, got {err}" + error.contains(&blksize_label) && error.contains(&overlap_label), + "error should name the offending blksize and overlap, got {error}" ); } } - /// An overlap equal to blksize gives a step of 0, which is really a - /// `nlm.validate()` fault, not a covering-block one. `Nl4dParams`'s - /// own check has to stay quiet about it, mirroring how construction - /// runs both validations in sequence, so the caller sees the overlap - /// constraint named rather than a nonsensical covering-block count - /// computed from a saturated step. #[test] fn overlap_equal_to_blksize_reports_the_overlap_constraint_not_covering_blocks() { + let motion_compensation = MotionCompensationMode::Mvtools { + blksize: 16, + overlap: 16, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Auto, + }; + let nlm = NlmParams { + motion_compensation, + ..Nl4dParams::default().nlm + }; let params = Nl4dParams { - nlm: NlmParams { - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 16, - overlap: 16, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Auto, - }, - ..Nl4dParams::default().nlm - }, + nlm, ..Nl4dParams::default() }; assert!( @@ -455,71 +418,71 @@ mod tests { "the covering-block check must not fire on a geometry nlm.validate() rejects on its \ own terms" ); - let err = params + + let error = params .nlm .validate() .expect_err("overlap == blksize must be rejected") .to_string(); assert!( - err.contains("overlap") && err.contains("blksize"), - "error should name the overlap constraint, got {err}" + error.contains("overlap") && error.contains("blksize"), + "error should name the overlap constraint, got {error}" ); assert!( - !err.contains("cover a patch"), - "error should not be the covering-block message, got {err}" + !error.contains("cover a patch"), + "error should not be the covering-block message, got {error}" ); } #[test] fn validate_rejects_missing_hq() { + let nlm = NlmParams { + hq: None, + ..Nl4dParams::default().nlm + }; let params = Nl4dParams { - nlm: NlmParams { - hq: None, - ..Nl4dParams::default().nlm - }, + nlm, ..Nl4dParams::default() }; - let err = params.validate().expect_err("expected rejection"); - assert!(err.contains("nlm.hq"), "error should name nlm.hq, got {err}"); + let error = params.validate().expect_err("expected rejection"); + assert!(error.contains("nlm.hq"), "error should name nlm.hq, got {error}"); } #[test] fn validate_rejects_inactive_motion_compensation() { + let nlm = NlmParams { + motion_compensation: MotionCompensationMode::None, + ..Nl4dParams::default().nlm + }; let params = Nl4dParams { - nlm: NlmParams { - motion_compensation: MotionCompensationMode::None, - ..Nl4dParams::default().nlm - }, + nlm, ..Nl4dParams::default() }; - let err = params.validate().expect_err("expected rejection"); + let error = params.validate().expect_err("expected rejection"); assert!( - err.contains("motion_compensation"), - "error should name nlm.motion_compensation, got {err}" + error.contains("motion_compensation"), + "error should name nlm.motion_compensation, got {error}" ); } - /// The latent precondition `submit_machinery` enforces at submit - /// time. Both motion compensation and the confidence buffer have to - /// be active, or that call returns an error. `validate` has to catch - /// a configuration that would hit that error before construction ever - /// gets that far. #[test] fn validate_rejects_missing_temporal_confidence() { + let hq = HqParams { + temporal_confidence: false, + ..HqParams::default() + }; + let nlm = NlmParams { + hq: Some(hq), + ..Nl4dParams::default().nlm + }; let params = Nl4dParams { - nlm: NlmParams { - hq: Some(HqParams { - temporal_confidence: false, - ..HqParams::default() - }), - ..Nl4dParams::default().nlm - }, + nlm, ..Nl4dParams::default() }; - let err = params.validate().expect_err("expected rejection"); + let error = params.validate().expect_err("expected rejection"); assert!( - err.contains("temporal_confidence"), - "error should name nlm.hq.temporal_confidence, got {err}" + error.contains("temporal_confidence"), + "error should name nlm.hq.temporal_confidence, got {error}" ); } @@ -605,12 +568,12 @@ mod tests { field_lambda: lambda, ..Nl4dParams::default() }; - let err = params + let error = params .validate() .expect_err("field_lambda={lambda} should be rejected"); assert!( - err.contains("field_lambda"), - "error should name field_lambda, got {err}" + error.contains("field_lambda"), + "error should name field_lambda, got {error}" ); } } diff --git a/av-denoise-core/src/nl4d/regularise.rs b/av-denoise-core/src/nl4d/regularise.rs index cb5edf2..0417913 100644 --- a/av-denoise-core/src/nl4d/regularise.rs +++ b/av-denoise-core/src/nl4d/regularise.rs @@ -12,16 +12,17 @@ use crate::nlmeans::motion::{ pyramid_slot_byte_offset, }; -/// Runs the field regularisation pass over every neighbour of `view`, -/// writing the result into `mv_out` and `conf_out`, which share the -/// front end's per-neighbour layout. +/// Runs the field regularisation pass over every neighbour of `view`. +/// +/// Each block's vector is re-scored against the median of its neighbourhood. The results land in +/// `mv_out` and `conf_out`, which share the front end's per-neighbour layout. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] pub(super) fn run_regularise( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, view: &RingView, width: u32, height: u32, @@ -31,37 +32,36 @@ pub(super) fn run_regularise( mv_out: &Handle, conf_out: &Handle, ) -> Result<(), anyhow::Error> { - let (fw, fh) = level_dims(width, height, 0); - let level_len = (fw * fh) as usize; - let blocks = (mc.blocks_x * mc.blocks_y) as usize; - let lambda_pixel = field_lambda * (mc.blksize * mc.blksize) as f32 * THSAD_PIXEL; - let centre = view.pyramid.clone().offset_start(pyramid_slot_byte_offset( + let (level_width, level_height) = level_dims(width, height, 0); + let level_len = (level_width * level_height) as usize; + let blocks = (motion_ctx.blocks_x * motion_ctx.blocks_y) as usize; + let block_area = (motion_ctx.blksize * motion_ctx.blksize) as f32; + let lambda_pixel = field_lambda * block_area * THSAD_PIXEL; + let centre_offset = pyramid_slot_byte_offset( width, height, view.frame_count, 0, view.centre_slot, - mc.align, - )); + motion_ctx.align, + ); + let centre = view.pyramid.clone().offset_start(centre_offset); for (t, &slot) in view.neighbour_slots.iter().enumerate() { let t = t as u32; - let neighbour = view.pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - view.frame_count, - 0, - slot, - mc.align, - )); - let mv_in = view.mv_field.clone().offset_start(mv_field_byte_offset(mc, t)); - let mv_dst = mv_out.clone().offset_start(mv_field_byte_offset(mc, t)); - let conf_dst = conf_out.clone().offset_start(confidence_byte_offset(mc, t)); + let neighbour_offset = + pyramid_slot_byte_offset(width, height, view.frame_count, 0, slot, motion_ctx.align); + let neighbour = view.pyramid.clone().offset_start(neighbour_offset); + let mv_offset = mv_field_byte_offset(motion_ctx, t); + let conf_offset = confidence_byte_offset(motion_ctx, t); + let mv_in = view.mv_field.clone().offset_start(mv_offset); + let mv_dst = mv_out.clone().offset_start(mv_offset); + let conf_dst = conf_out.clone().offset_start(conf_offset); unsafe { nl4d_mv_regularise::launch_unchecked::( client, - CubeCount::new_2d(mc.blocks_x, mc.blocks_y), + CubeCount::new_2d(motion_ctx.blocks_x, motion_ctx.blocks_y), CubeDim::new_2d(8, 8), ArrayArg::from_raw_parts(centre.clone(), level_len), ArrayArg::from_raw_parts(neighbour, level_len), @@ -71,12 +71,12 @@ pub(super) fn run_regularise( lambda_pixel, sad_noise_floor, thsad, - fw, - fh, - mc.blksize, - mc.step, - mc.blocks_x, - mc.blocks_y, + level_width, + level_height, + motion_ctx.blksize, + motion_ctx.step, + motion_ctx.blocks_x, + motion_ctx.blocks_y, ); } } diff --git a/av-denoise-core/src/nl4d/snapshot.rs b/av-denoise-core/src/nl4d/snapshot.rs index 68ce161..133e73b 100644 --- a/av-denoise-core/src/nl4d/snapshot.rs +++ b/av-denoise-core/src/nl4d/snapshot.rs @@ -1,15 +1,12 @@ use cubecl::prelude::*; use cubecl::server::Handle; -/// The motion field and confidence the last collaborative pass read, -/// copied back to the host. +/// The motion field and confidence the last collaborative pass read, copied back to the host. /// -/// `vectors[t][block]` is block `block`'s vector toward neighbour `t`, -/// in pixels, and `confidence[t][block]` that block's confidence in -/// `[0, 1]`. `offsets[t]` is neighbour `t`'s temporal offset from the -/// centre frame. Blocks run row-major over `blocks_x * blocks_y`, and -/// block `(bx, by)` covers `blksize` pixels starting at `bx * step` on -/// each axis. +/// `vectors[t][block]` is a block's vector toward neighbour `t` in pixels, and +/// `confidence[t][block]` its confidence between 0 and 1. `offsets[t]` is neighbour `t`'s temporal +/// offset from the centre frame. Blocks run row-major, and block `(bx, by)` covers `blksize` +/// pixels starting at `bx * step` and `by * step`. /// /// This exists for measurement tooling. It is not a stable interface. #[doc(hidden)] @@ -24,8 +21,7 @@ pub struct MotionSnapshot { pub confidence: Vec>, } -/// The device buffers one pass handed the fused kernel, kept so the -/// snapshot can read them back after the fact. +/// The device buffers one pass handed the fused kernel, kept for a later snapshot readback. pub(super) struct LastFields { pub mv_field: Handle, pub confidence: Handle, @@ -48,32 +44,33 @@ pub(super) fn read_snapshot( let mv_bytes = client .read_one(fields.mv_field.clone()) .expect("motion field readback failed"); - let mv = i32::from_bytes(&mv_bytes); + let mv_values = i32::from_bytes(&mv_bytes); let conf_bytes = client .read_one(fields.confidence.clone()) .expect("confidence readback failed"); - let conf = f32::from_bytes(&conf_bytes); + let conf_values = f32::from_bytes(&conf_bytes); let mut offsets = Vec::with_capacity(fields.neighbours as usize); let mut vectors = Vec::with_capacity(fields.neighbours as usize); let mut confidence = Vec::with_capacity(fields.neighbours as usize); for t in 0..fields.neighbours { - // Mirrors `neighbour_idx_for_k`, ascending k on the negative - // side first, then ascending k on the positive side. - let k = if t < radius { + // Matches `neighbour_idx_for_k`, negative offsets first, then positive, each ascending. + let offset = if t < radius { t as i32 - radius as i32 } else { t as i32 - radius as i32 + 1 }; - offsets.push(k); + offsets.push(offset); + let mv_base = (t * fields.mv_stride) as usize; - vectors.push( - (0..blocks) - .map(|b| [mv[mv_base + 2 * b], mv[mv_base + 2 * b + 1]]) - .collect(), - ); - let c_base = (t * fields.conf_stride) as usize; - confidence.push(conf[c_base..c_base + blocks].to_vec()); + let block_vectors = (0..blocks) + .map(|block| [mv_values[mv_base + 2 * block], mv_values[mv_base + 2 * block + 1]]) + .collect(); + vectors.push(block_vectors); + + let conf_base = (t * fields.conf_stride) as usize; + let block_confidence = conf_values[conf_base..conf_base + blocks].to_vec(); + confidence.push(block_confidence); } MotionSnapshot { diff --git a/av-denoise-core/src/nl4d/tests/edges.rs b/av-denoise-core/src/nl4d/tests/edges.rs index 5de76c5..56e4ffe 100644 --- a/av-denoise-core/src/nl4d/tests/edges.rs +++ b/av-denoise-core/src/nl4d/tests/edges.rs @@ -1,5 +1,7 @@ -use super::helpers::{R, SIGMA, make_client, noisy_copy_of, static_clip_params, textured_base}; +use super::helpers::{R, SIGMA, make_client, static_clip_params, textured_base}; +use crate::bench_api::HostIo; use crate::nl4d::Nl4dDenoiser; +use crate::nlmeans::tests::helpers::noisy_field_over; const SIZE: u32 = 64; @@ -14,22 +16,20 @@ fn run_stream(radius: u32, count: u32) -> (Vec>, Option) { let mut first_output_at = None; for seed in 0..count { - let frame = noisy_copy_of(&base, SIZE, SIZE, SIGMA, seed); + let frame = noisy_field_over(&base, SIZE, SIZE, SIGMA, seed); denoiser.push_frame(&frame); - let Some(pending) = denoiser.denoise_submit().expect("denoise_submit failed") else { + let Some(values) = denoiser.denoise().expect("denoise failed") else { continue; }; first_output_at.get_or_insert(seed + 1); - let frame = pending.wait().expect("readback failed"); - let values = frame.into_f32().expect("f32 output"); outputs.push(values); } denoiser .flush(|frame| { - let values = frame.as_f32().expect("f32 output").to_vec(); + let values = frame.to_vec(); outputs.push(values); }) .expect("flush failed"); @@ -61,11 +61,10 @@ fn first_output_push(denoiser: &mut Nl4dDenoiser, count: u32, first_seed: u32 let mut first_output_at = None; for push in 0..count { - let frame = noisy_copy_of(&base, SIZE, SIZE, SIGMA, first_seed + push); + let frame = noisy_field_over(&base, SIZE, SIZE, SIGMA, first_seed + push); denoiser.push_frame(&frame); - if let Some(pending) = denoiser.denoise_submit().expect("denoise_submit failed") { - pending.wait().expect("readback failed"); + if denoiser.denoise().expect("denoise failed").is_some() { first_output_at.get_or_insert(push + 1); } } @@ -149,10 +148,8 @@ fn edge_frames_are_denoised_as_strongly_as_mid_scene() { assert!(last <= 1.10 * middle, "last {last} vs middle {middle}"); } -/// A caller that primes a full ring with priming pushes and never -/// submits reaches `flush` with `passes_run == 0`, so the tail path's -/// accumulators were never cleared for this ring. That must not scatter -/// stale contributions into the output. +/// A ring primed without any submit reaches `flush` with `passes_run == 0` and uncleared +/// accumulators, which must not scatter stale contributions into the output. #[test] fn a_full_ring_primed_without_any_submit_flushes_without_black_output() { let radius = 2; @@ -163,28 +160,25 @@ fn a_full_ring_primed_without_any_submit_flushes_without_black_output() { let clean = base.clone(); denoiser.mark_continuation(); + for seed in 0..(2 * radius + 1) { - let frame = noisy_copy_of(&base, SIZE, SIZE, SIGMA, seed); + let frame = noisy_field_over(&base, SIZE, SIZE, SIGMA, seed); denoiser.push_frame(&frame); } let mut outputs = Vec::new(); denoiser .flush(|frame| { - let values = frame.as_f32().expect("f32 output").to_vec(); + let values = frame.to_vec(); outputs.push(values); }) .expect("flush failed"); assert_eq!(outputs.len(), 2 * radius as usize); - // Only two tail passes ever run here (no real pass warmed the ring - // first), so coverage per output frame is uneven and a couple of - // them sit above `SIGMA` rather than clearing it outright. The - // point of this test is that the fix stops the tail path from - // scattering into an uncleared accumulator, not that a never-warmed - // ring denoises as strongly as a normal stream, so the bound is - // looser than [residual_std] gets elsewhere in this file. + // Only two tail passes run, so coverage per frame is uneven and a couple of frames sit above + // `SIGMA`. This pins the stale scatter, not full-strength denoising, so the bound is looser than + // the `SIGMA` the other tests in this file use. for (index, output) in outputs.iter().enumerate() { let mean = output.iter().sum::() / output.len() as f32; assert!(mean > 0.1, "frame {index} came out black"); diff --git a/av-denoise-core/src/nl4d/tests/engine.rs b/av-denoise-core/src/nl4d/tests/engine.rs new file mode 100644 index 0000000..c30ebde --- /dev/null +++ b/av-denoise-core/src/nl4d/tests/engine.rs @@ -0,0 +1,707 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::helpers::noisy_frames; +use crate::bench_api::HostIo; +use crate::engine::{DevicePlane, EdgePadding, Engine, Geometry, SampleFormat}; +use crate::error::Error; +use crate::nl4d::grain::GrainChunk; +use crate::nl4d::{Nl4d, Nl4dDenoiser, Nl4dOptions, resolve_params}; +use crate::nlmeans::ChannelMode; +use crate::nlmeans::tests::helpers::{ + R, + make_client, + normalise_with_ingest, + read_interleaved, + upload_planes, +}; + +const WIDTH: u32 = 64; +const HEIGHT: u32 = 48; + +fn geometry(channels: ChannelMode) -> Geometry { + Geometry { + width: WIDTH, + height: HEIGHT, + channels, + input: SampleFormat::F32, + output: SampleFormat::F32, + } +} + +fn options(radius: u32, grain_export: bool) -> Nl4dOptions { + Nl4dOptions { + temporal_radius: radius, + grain_export, + ..Nl4dOptions::default() + } +} + +fn build_engine(client: &ComputeClient, radius: u32, channels: ChannelMode) -> Nl4d { + let options = options(radius, false); + let geometry = geometry(channels); + + Nl4d::new(client, options, geometry).expect("build") +} + +fn emit(engine: &mut Nl4d, client: &ComputeClient, channel_count: usize) -> Vec { + let pixels = (WIDTH * HEIGHT) as usize; + let outputs: Vec = (0..channel_count).map(|_| client.empty(pixels * 4)).collect(); + let planes: Vec<_> = outputs + .iter() + .map(|handle| DevicePlane::new(handle, WIDTH, HEIGHT)) + .collect(); + + engine.emit_into(&planes).expect("emit"); + + read_interleaved(client, &outputs) +} + +/// Pushes one frame and returns how many frames the engine reports ready. +fn push_frame(engine: &mut Nl4d, client: &ComputeClient, frame: &[f32], channel_count: usize) -> usize { + let handles = upload_planes(client, frame, channel_count); + let planes: Vec<_> = handles + .iter() + .map(|handle| DevicePlane::new(handle, WIDTH, HEIGHT)) + .collect(); + + engine.push(&planes).expect("push") +} + +fn push_context_frame(engine: &mut Nl4d, client: &ComputeClient, frame: &[f32], channel_count: usize) { + let handles = upload_planes(client, frame, channel_count); + let planes: Vec<_> = handles + .iter() + .map(|handle| DevicePlane::new(handle, WIDTH, HEIGHT)) + .collect(); + + engine.push_context(&planes).expect("push_context"); +} + +fn try_push_context(engine: &mut Nl4d, client: &ComputeClient, frame: &[f32]) -> Result<(), Error> { + let handles = upload_planes(client, frame, 1); + let planes = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + + engine.push_context(&planes) +} + +fn push_and_emit( + engine: &mut Nl4d, + client: &ComputeClient, + frame: &[f32], + channel_count: usize, + outputs: &mut Vec>, +) { + let ready = push_frame(engine, client, frame, channel_count); + for _ in 0..ready { + let output = emit(engine, client, channel_count); + outputs.push(output); + } +} + +fn finish_and_emit( + engine: &mut Nl4d, + client: &ComputeClient, + channel_count: usize, + outputs: &mut Vec>, +) { + let tail = engine.finish().expect("finish"); + for _ in 0..tail { + let output = emit(engine, client, channel_count); + outputs.push(output); + } +} + +/// Pushes every frame through `engine`, then the tail, returning each emitted frame. +fn drive( + engine: &mut Nl4d, + client: &ComputeClient, + frames: &[Vec], + channel_count: usize, +) -> Vec> { + let mut outputs = Vec::new(); + + for frame in frames { + push_and_emit(engine, client, frame, channel_count, &mut outputs); + } + + finish_and_emit(engine, client, channel_count, &mut outputs); + + outputs +} + +/// Pushes each input, then the tail, emitting every frame as `u8` planes. +fn drive_u8(engine: &mut Nl4d, client: &ComputeClient, inputs: &[Handle]) -> Vec> { + let pixels = (WIDTH * HEIGHT) as usize; + let mut outputs = Vec::new(); + + for input in inputs { + let planes = [DevicePlane::new(input, WIDTH, HEIGHT)]; + let ready = engine.push(&planes).expect("push"); + emit_u8(engine, client, ready, pixels, &mut outputs); + } + + let tail = engine.finish().expect("finish"); + emit_u8(engine, client, tail, pixels, &mut outputs); + + outputs +} + +/// Emits `frame_count` frames as `u8` planes into `outputs`. +fn emit_u8( + engine: &mut Nl4d, + client: &ComputeClient, + frame_count: usize, + pixels: usize, + outputs: &mut Vec>, +) { + for _ in 0..frame_count { + let output = client.empty(pixels); + let planes = [DevicePlane::new(&output, WIDTH, HEIGHT)]; + engine.emit_into(&planes).expect("emit"); + + let bytes = client.read_one(output).expect("read"); + outputs.push(bytes.to_vec()); + } +} + +struct Run { + frames: Vec>, + grain: Vec, +} + +/// Runs the engine, pushing the first `context` frames with `push_context`. +fn run_engine(options: Nl4dOptions, channels: ChannelMode, frames: &[Vec], context: usize) -> Run { + let client = make_client(); + let channel_count = channels.count() as usize; + let geometry = geometry(channels); + let mut engine = Nl4d::new(&client, options, geometry).expect("build"); + let mut outputs = Vec::new(); + + for (index, frame) in frames.iter().enumerate() { + if index < context { + push_context_frame(&mut engine, &client, frame, channel_count); + continue; + } + + push_and_emit(&mut engine, &client, frame, channel_count, &mut outputs); + } + + finish_and_emit(&mut engine, &client, channel_count, &mut outputs); + + // Drained through the trait object, the way the host layer reaches it. + let dyn_engine: &mut dyn Engine = &mut engine; + let grain = dyn_engine.drain_grain_chunks().expect("grain"); + + Run { + frames: outputs, + grain, + } +} + +/// Runs the `Nl4dDenoiser` oracle, marking a continuation and skipping submits for the first +/// `context` frames. +fn run_oracle(options: Nl4dOptions, channels: ChannelMode, frames: &[Vec], context: usize) -> Run { + let client = make_client(); + let params = resolve_params(&options, channels).expect("resolve"); + let mut denoiser = Nl4dDenoiser::new(&client, params, WIDTH, HEIGHT).expect("build"); + let mut outputs = Vec::new(); + + if context > 0 { + denoiser.mark_continuation(); + } + + for (index, frame) in frames.iter().enumerate() { + denoiser.push_frame(frame); + + if index < context { + continue; + } + + let output = denoiser.denoise().expect("denoise"); + if let Some(samples) = output { + outputs.push(samples); + } + } + + let collect = |output: &[f32]| { + let samples = output.to_vec(); + outputs.push(samples); + }; + denoiser.flush(collect).expect("flush"); + + let grain = denoiser.drain_grain_chunks().expect("grain"); + + Run { + frames: outputs, + grain, + } +} + +#[test] +fn luma_stream_matches_the_oracle_including_the_tail() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let actual_options = options(2, false); + let actual = run_engine(actual_options, ChannelMode::Luma, &frames, 0); + let expected_options = options(2, false); + let expected = run_oracle(expected_options, ChannelMode::Luma, &frames, 0); + assert_eq!(actual.frames.len(), 10); + assert_eq!(actual.frames, expected.frames); +} + +#[test] +fn chroma_stream_matches_the_oracle() { + let frames = noisy_frames(WIDTH, HEIGHT, 2, 8); + let actual_options = options(1, false); + let actual = run_engine(actual_options, ChannelMode::Chroma, &frames, 0); + let expected_options = options(1, false); + let expected = run_oracle(expected_options, ChannelMode::Chroma, &frames, 0); + assert_eq!(actual.frames.len(), 8); + assert_eq!(actual.frames, expected.frames); +} + +#[test] +fn yuv_stream_matches_the_oracle() { + let frames = noisy_frames(WIDTH, HEIGHT, 3, 8); + let actual_options = options(1, false); + let actual = run_engine(actual_options, ChannelMode::Yuv, &frames, 0); + let expected_options = options(1, false); + let expected = run_oracle(expected_options, ChannelMode::Yuv, &frames, 0); + assert_eq!(actual.frames, expected.frames); +} + +#[test] +fn a_short_stream_matches_the_oracle() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 3); + let actual_options = options(2, false); + let actual = run_engine(actual_options, ChannelMode::Luma, &frames, 0); + let expected_options = options(2, false); + let expected = run_oracle(expected_options, ChannelMode::Luma, &frames, 0); + assert_eq!(actual.frames.len(), 3); + assert_eq!(actual.frames, expected.frames); +} + +#[test] +fn a_context_led_stream_matches_the_oracle_continuation() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 12); + let actual_options = options(2, false); + let actual = run_engine(actual_options, ChannelMode::Luma, &frames, 4); + let expected_options = options(2, false); + let expected = run_oracle(expected_options, ChannelMode::Luma, &frames, 4); + assert_eq!(actual.frames, expected.frames); +} + +#[test] +fn grain_chunks_match_the_oracle() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let actual_options = options(2, true); + let actual = run_engine(actual_options, ChannelMode::Luma, &frames, 0); + let expected_options = options(2, true); + let expected = run_oracle(expected_options, ChannelMode::Luma, &frames, 0); + assert!(!actual.grain.is_empty()); + assert_eq!(actual.grain, expected.grain); +} + +#[test] +fn a_second_stream_after_finish_matches_a_fresh_engine() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + + drive(&mut engine, &client, &frames, 1); + let second = drive(&mut engine, &client, &frames, 1); + + let fresh_options = options(2, false); + let fresh = run_engine(fresh_options, ChannelMode::Luma, &frames, 0); + assert_eq!(second, fresh.frames); +} + +#[test] +fn a_context_led_second_stream_matches_a_fresh_engine() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 12); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + + drive(&mut engine, &client, &frames, 1); + + let mut second = Vec::new(); + for (index, frame) in frames.iter().enumerate() { + if index < 4 { + push_context_frame(&mut engine, &client, frame, 1); + continue; + } + + push_and_emit(&mut engine, &client, frame, 1, &mut second); + } + + finish_and_emit(&mut engine, &client, 1, &mut second); + + let fresh_options = options(2, false); + let fresh = run_engine(fresh_options, ChannelMode::Luma, &frames, 4); + assert_eq!(second, fresh.frames); +} + +#[test] +fn push_context_after_a_push_is_refused_and_changes_no_state() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + let mut outputs = Vec::new(); + + for (index, frame) in frames.iter().enumerate() { + push_and_emit(&mut engine, &client, frame, 1, &mut outputs); + + if index == 5 { + let refused = try_push_context(&mut engine, &client, &frames[9]); + assert!(matches!(refused, Err(Error::ContextAfterPush))); + } + } + + finish_and_emit(&mut engine, &client, 1, &mut outputs); + + let expected_options = options(2, false); + let expected = run_oracle(expected_options, ChannelMode::Luma, &frames, 0); + assert_eq!(outputs, expected.frames); +} + +#[test] +fn push_context_is_allowed_again_after_reset() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 2); + let client = make_client(); + let mut engine = build_engine(&client, 1, ChannelMode::Luma); + + push_frame(&mut engine, &client, &frames[0], 1); + engine.reset(); + + let context = try_push_context(&mut engine, &client, &frames[1]); + assert!(context.is_ok()); +} + +#[test] +fn push_context_is_allowed_again_after_the_tail_is_emitted() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 6); + let client = make_client(); + let mut engine = build_engine(&client, 1, ChannelMode::Luma); + + drive(&mut engine, &client, &frames, 1); + + let context = try_push_context(&mut engine, &client, &frames[0]); + assert!(context.is_ok()); +} + +/// Pushes frames until one is ready, leaving it unemitted, and returns the ready count and frames pushed. +fn engine_with_a_ready_frame(client: &ComputeClient, frames: &[Vec]) -> (Nl4d, usize, usize) { + let mut engine = build_engine(client, 1, ChannelMode::Luma); + let mut pushed = 0; + + loop { + let ready = push_frame(&mut engine, client, &frames[pushed], 1); + pushed += 1; + + if ready == 1 { + return (engine, ready, pushed); + } + } +} + +#[test] +fn push_before_emitting_returns_outputs_pending_and_keeps_the_frame() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let client = make_client(); + let (mut engine, _ready, pushed) = engine_with_a_ready_frame(&client, &frames); + + let handles = upload_planes(&client, &frames[pushed], 1); + let planes = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + let refused = engine.push(&planes); + assert!(matches!(refused, Err(Error::OutputsPending))); + + let output = emit(&mut engine, &client, 1); + let expected_options = options(1, false); + let expected = run_oracle(expected_options, ChannelMode::Luma, &frames, 0); + assert_eq!(output, expected.frames[0]); +} + +#[test] +fn finish_while_a_frame_is_ready_returns_outputs_pending() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let client = make_client(); + let (mut engine, _ready, _pushed) = engine_with_a_ready_frame(&client, &frames); + + let finished = engine.finish(); + assert!(matches!(finished, Err(Error::OutputsPending))); +} + +fn engine_owing_a_tail(client: &ComputeClient) -> (Nl4d, usize) { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 6); + let mut engine = build_engine(client, 1, ChannelMode::Luma); + let mut outputs = Vec::new(); + + for frame in &frames { + push_and_emit(&mut engine, client, frame, 1, &mut outputs); + } + + let tail = engine.finish().expect("finish"); + assert!(tail > 0); + + (engine, tail) +} + +#[test] +fn push_while_a_tail_is_owed_returns_outputs_pending() { + let client = make_client(); + let (mut engine, _tail) = engine_owing_a_tail(&client); + let frames = noisy_frames(WIDTH, HEIGHT, 1, 1); + let handles = upload_planes(&client, &frames[0], 1); + let planes = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + + let pushed = engine.push(&planes); + assert!(matches!(pushed, Err(Error::OutputsPending))); +} + +#[test] +fn finish_while_a_tail_is_owed_returns_outputs_pending() { + let client = make_client(); + let (mut engine, _tail) = engine_owing_a_tail(&client); + + let finished = engine.finish(); + assert!(matches!(finished, Err(Error::OutputsPending))); +} + +#[test] +fn emit_after_the_last_tail_frame_returns_nothing_to_emit() { + let client = make_client(); + let (mut engine, tail) = engine_owing_a_tail(&client); + + for _ in 0..tail { + emit(&mut engine, &client, 1); + } + + let output = client.empty((WIDTH * HEIGHT * 4) as usize); + let planes = [DevicePlane::new(&output, WIDTH, HEIGHT)]; + let emitted = engine.emit_into(&planes); + assert!(matches!(emitted, Err(Error::NothingToEmit))); +} + +#[test] +fn misuse_errors_do_not_poison() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 4); + let client = make_client(); + let mut engine = build_engine(&client, 1, ChannelMode::Luma); + let output = client.empty((WIDTH * HEIGHT * 4) as usize); + let planes = [DevicePlane::new(&output, WIDTH, HEIGHT)]; + + let nothing = engine.emit_into(&planes); + assert!(matches!(nothing, Err(Error::NothingToEmit))); + + let wrong_count = engine.push(&[]); + assert!(matches!(wrong_count, Err(Error::PlaneMismatch(_)))); + + let mut ready = 0; + for frame in &frames { + ready = push_frame(&mut engine, &client, frame, 1); + if ready == 1 { + break; + } + } + + assert_eq!(ready, 1); +} + +#[test] +fn a_plane_mismatch_mid_stream_changes_no_state() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + let mut outputs = Vec::new(); + + for (index, frame) in frames.iter().enumerate() { + push_and_emit(&mut engine, &client, frame, 1, &mut outputs); + + if index == 5 { + let wrong_count = engine.push(&[]); + assert!(matches!(wrong_count, Err(Error::PlaneMismatch(_)))); + + let wrong_context = engine.push_context(&[]); + assert!(matches!(wrong_context, Err(Error::PlaneMismatch(_)))); + } + } + + finish_and_emit(&mut engine, &client, 1, &mut outputs); + + let expected_options = options(2, false); + let expected = run_oracle(expected_options, ChannelMode::Luma, &frames, 0); + assert_eq!(outputs, expected.frames); +} + +#[test] +fn a_gpu_failure_poisons_until_reset() { + let client = make_client(); + let mut engine = build_engine(&client, 1, ChannelMode::Luma); + + let failure = engine.fail_through_guard_for_test(); + assert!(matches!(failure, Err(Error::Gpu(_)))); + + let frames = noisy_frames(WIDTH, HEIGHT, 1, 1); + let handles = upload_planes(&client, &frames[0], 1); + let input = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + let pushed = engine.push(&input); + assert!(matches!(pushed, Err(Error::NeedsReset))); + + engine.reset(); + let pushed_after_reset = engine.push(&input); + assert!(pushed_after_reset.is_ok()); +} + +#[test] +fn reset_mid_stream_with_a_ready_frame_matches_a_fresh_engine() { + let frames = noisy_frames(WIDTH, HEIGHT, 1, 10); + let client = make_client(); + let (mut engine, ready, _pushed) = engine_with_a_ready_frame(&client, &frames); + assert_eq!(ready, 1); + + engine.reset(); + let second = drive(&mut engine, &client, &frames, 1); + + let fresh_options = options(1, false); + let fresh = run_engine(fresh_options, ChannelMode::Luma, &frames, 0); + assert_eq!(second, fresh.frames); +} + +#[test] +fn u8_input_matches_f32_input_from_the_ingest_kernel() { + let client = make_client(); + let frames = noisy_frames(WIDTH, HEIGHT, 1, 4); + let codes: Vec> = frames + .iter() + .map(|frame| { + frame + .iter() + .map(|&value| (value.clamp(0.0, 1.0) * 255.0) as u8) + .collect() + }) + .collect(); + let u8_inputs: Vec = codes + .iter() + .map(|frame| client.create_from_slice(frame)) + .collect(); + let f32_inputs: Vec = codes + .iter() + .map(|frame| { + let normalised = normalise_with_ingest(&client, frame, WIDTH, HEIGHT); + let bytes = f32::as_bytes(&normalised); + client.create_from_slice(bytes) + }) + .collect(); + + let u8_geometry = Geometry { + input: SampleFormat::U8, + output: SampleFormat::U8, + width: WIDTH, + height: HEIGHT, + channels: ChannelMode::Luma, + }; + let f32_geometry = Geometry { + input: SampleFormat::F32, + ..u8_geometry + }; + let u8_options = options(1, false); + let f32_options = options(1, false); + let mut u8_engine = Nl4d::new(&client, u8_options, u8_geometry).expect("build u8"); + let mut f32_engine = Nl4d::new(&client, f32_options, f32_geometry).expect("build f32"); + + let u8_frames = drive_u8(&mut u8_engine, &client, &u8_inputs); + let f32_frames = drive_u8(&mut f32_engine, &client, &f32_inputs); + assert_eq!(u8_frames.len(), 4); + assert_eq!(u8_frames, f32_frames); +} + +#[test] +fn window_span_doubles_the_radius_with_shifted_edges() { + let client = make_client(); + let engine = build_engine(&client, 2, ChannelMode::Luma); + let span = engine.window_span(); + assert_eq!((span.behind, span.ahead), (4, 4)); + assert_eq!(span.edges, EdgePadding::Shifted); + + let held = engine.max_held_frames(); + assert_eq!(held, 4); +} + +#[test] +fn frames_smaller_than_a_patch_are_invalid_geometry() { + let client = make_client(); + let geometry = Geometry { + width: 4, + height: HEIGHT, + ..geometry(ChannelMode::Luma) + }; + + let engine_options = options(1, false); + let built = Nl4d::new(&client, engine_options, geometry); + assert!(matches!(built, Err(Error::InvalidGeometry(_)))); +} + +#[test] +fn a_plane_past_u32_pixels_is_invalid_geometry() { + let client = make_client(); + let oversized = Geometry { + width: 65_536, + height: 65_536, + ..geometry(ChannelMode::Luma) + }; + + let engine_options = options(1, false); + let built = Nl4d::new(&client, engine_options, oversized); + assert!(matches!(built, Err(Error::InvalidGeometry(_)))); +} + +/// Each 16384x16384 YUV frame stores 2^30 elements, so a five-frame ring passes `u32::MAX`. +#[test] +fn a_ring_past_u32_elements_is_invalid_geometry() { + let client = make_client(); + let oversized = Geometry { + width: 16_384, + height: 16_384, + ..geometry(ChannelMode::Yuv) + }; + + let engine_options = options(2, false); + let built = Nl4d::new(&client, engine_options, oversized); + assert!(matches!(built, Err(Error::InvalidGeometry(_)))); +} + +/// A 32768x24576 luma frame ring of five frames fits in `u32`, but its motion pyramid ring holds at +/// least 1.25 times as many elements and does not. +#[test] +fn a_motion_pyramid_past_u32_elements_is_invalid_geometry() { + let client = make_client(); + let oversized = Geometry { + width: 32_768, + height: 24_576, + ..geometry(ChannelMode::Luma) + }; + + let engine_options = options(2, false); + let built = Nl4d::new(&client, engine_options, oversized); + let Err(Error::InvalidGeometry(message)) = built else { + panic!("expected InvalidGeometry"); + }; + assert!(message.contains("pyramid"), "{message}"); +} + +#[test] +fn hd_and_4k_geometries_construct() { + let client = make_client(); + + for (width, height) in [(1920, 1080), (3840, 2160)] { + let sized = Geometry { + width, + height, + ..geometry(ChannelMode::Luma) + }; + + let engine_options = options(2, false); + let built = Nl4d::new(&client, engine_options, sized); + assert!(built.is_ok(), "{width}x{height} failed to construct"); + } +} diff --git a/av-denoise-core/src/nl4d/tests/grouping.rs b/av-denoise-core/src/nl4d/tests/grouping.rs index b6be051..b5328c6 100644 --- a/av-denoise-core/src/nl4d/tests/grouping.rs +++ b/av-denoise-core/src/nl4d/tests/grouping.rs @@ -17,32 +17,32 @@ use crate::collab::{PATCH_SIZE, STEP, grid_frames, needs_warp_uniform_search}; use crate::nlmeans::NOISE_CURVE_BINS; use crate::nlmeans::motion::neighbour_idx_for_k; -/// The motion block side length these fixtures score confidence -/// against, distinct from [`BLK_STEP`], which stays at `PATCH_SIZE` so -/// a block boundary lines up with a patch boundary. +/// The motion block side length these fixtures score confidence against. +/// +/// It differs from [BLK_STEP], which stays at `PATCH_SIZE` so a block boundary lines up with a +/// patch boundary. pub(super) const BLKSIZE: u32 = 16; const REFINE: u32 = 2; const K_MAX: u32 = 8; const SPATIAL_RADIUS: u32 = 4; +/// Where [twin_ring] plants the reference's spatial twin in the centre frame. +const TWIN_POS: (u32, u32) = (64, 48); + /// The knobs a run varies. Everything else follows the fixture. struct Knobs { c_min: f32, k_max: u32, sigma: f32, lambda_ht: f32, - /// Half-width of each neighbour's refine window, defaulting to the - /// module's [`REFINE`]. + /// Half-width of each neighbour's refine window. refine: u32, - /// Half-width of the centre frame's search window, defaulting to the - /// module's [SPATIAL_RADIUS]. + /// Half-width of the centre frame's search window. spatial_radius: u32, - /// The motion block side length, defaulting to the module's - /// [`BLKSIZE`]. At [`BLK_STEP`] exactly one block covers a patch. + /// The motion block side length. At [BLK_STEP] exactly one block covers a patch. blksize: u32, - /// Pins the search walk rather than taking it from the runtime. - /// `None`, the default, follows [`needs_warp_uniform_search`]. + /// Pins the search walk. `None` follows [needs_warp_uniform_search]. warp_uniform: Option, } @@ -61,7 +61,7 @@ impl Default for Knobs { } } -/// What one launch of [`collab_fused`] left behind. +/// What one launch of [collab_fused] left behind. struct FusedRun { wsum: Vec, group_weight: Vec, @@ -69,67 +69,92 @@ struct FusedRun { } impl FusedRun { - /// The total weight one ring slot's region received. A slot no - /// member scattered into reads exactly zero. + /// The total weight one ring slot's region received. fn frame_weight_sum(&self, slot: u32) -> i64 { let start = slot as usize * self.pixels; self.wsum[start..start + self.pixels] .iter() - .map(|&v| v as i64) + .map(|&weight| weight as i64) .sum() } - /// The total weight the whole ring received. Every group contributes - /// one patch of 64 pixels per member, so at a fixed per-group weight + /// The total weight the whole ring received. + /// + /// Every group contributes one patch of 64 pixels per member, so at a fixed per-group weight /// this counts members. fn total_weight(&self) -> i64 { - self.wsum.iter().map(|&v| v as i64).sum() + self.wsum.iter().map(|&weight| weight as i64).sum() } } -/// Launches [`collab_fused`] over a fixture, on the same one-cube-per- -/// eight-references grid `Nl4dDenoiser` uses, and reads back the -/// accumulator weights and the per-reference group weight. +/// Launches [collab_fused] over a fixture on the denoiser's eight-references-per-cube grid. /// -/// Luma always stores one channel per line, so the kernel's `Size` -/// selector is fixed at 1 here rather than threaded through as an -/// argument. -fn run_fused_over(fx: &RingFixture, k: Knobs) -> FusedRun { +/// Luma always stores one channel per line, so the kernel's `Size` selector is fixed at 1. +fn run_fused_over(fixture: &RingFixture, knobs: &Knobs) -> FusedRun { let client = make_client(); - let w = fx.width; - let h = fx.height; - let pixels = (w * h) as usize; - let frames = fx.ring.len() / pixels; - let refs = ref_count(w, h); - let refs_x = refs_along(w); + let width = fixture.width; + let height = fixture.height; + let pixels = (width * height) as usize; + let frames = fixture.ring.len() / pixels; + let refs = ref_count(width, height); + let refs_x = refs_along(width); let profile = dct_noise_profile(0.0); - - let ring_buf = client.create_from_slice(f32::as_bytes(&fx.ring)); - let mv_buf = client.create_from_slice(i32::as_bytes(&fx.mv_field)); - let conf_buf = client.create_from_slice(f32::as_bytes(&fx.confidence)); - let slots_buf = client.create_from_slice(u32::as_bytes(&fx.neighbour_slots)); - let sigma_buf = client.create_from_slice(f32::as_bytes(&[k.sigma])); - let profile_buf = client.create_from_slice(f32::as_bytes(&profile)); - let kaiser_buf = client.create_from_slice(f32::as_bytes(&kaiser_window(0.0))); - let zero_curve = client.create_from_slice(f32::as_bytes(&[0.0f32; NOISE_CURVE_BINS])); - let (map_cols, map_rows) = strength_map_dims(w, h); + let kaiser = kaiser_window(0.0); + let zeroed_curve = [0.0f32; NOISE_CURVE_BINS]; + let zeroed_accum = vec![0i32; pixels * frames]; + let zeroed_wsum = vec![0i32; pixels * frames]; + + let ring_bytes = f32::as_bytes(&fixture.ring); + let mv_bytes = i32::as_bytes(&fixture.mv_field); + let conf_bytes = f32::as_bytes(&fixture.confidence); + let slots_bytes = u32::as_bytes(&fixture.neighbour_slots); + let sigma = [knobs.sigma]; + let sigma_bytes = f32::as_bytes(&sigma); + let profile_bytes = f32::as_bytes(&profile); + let kaiser_bytes = f32::as_bytes(&kaiser); + let curve_bytes = f32::as_bytes(&zeroed_curve); + let ring_buf = client.create_from_slice(ring_bytes); + let mv_buf = client.create_from_slice(mv_bytes); + let conf_buf = client.create_from_slice(conf_bytes); + let slots_buf = client.create_from_slice(slots_bytes); + let sigma_buf = client.create_from_slice(sigma_bytes); + let profile_buf = client.create_from_slice(profile_bytes); + let kaiser_buf = client.create_from_slice(kaiser_bytes); + let zero_curve = client.create_from_slice(curve_bytes); + + let (map_cols, map_rows) = strength_map_dims(width, height); let map_len = (map_cols * map_rows) as usize; let unit_map = vec![1.0f32; map_len]; - let unit_map_buf = client.create_from_slice(f32::as_bytes(&unit_map)); - let accum = client.create_from_slice(i32::as_bytes(&vec![0i32; pixels * frames])); - let wsum = client.create_from_slice(i32::as_bytes(&vec![0i32; pixels * frames])); + let unit_map_bytes = f32::as_bytes(&unit_map); + let unit_map_buf = client.create_from_slice(unit_map_bytes); + + let accum_bytes = i32::as_bytes(&zeroed_accum); + let wsum_bytes = i32::as_bytes(&zeroed_wsum); + let accum = client.create_from_slice(accum_bytes); + let wsum = client.create_from_slice(wsum_bytes); let group_weight = client.empty(refs * size_of::()); + let cubes_x = fused_cubes_x(width); + let refs_y = refs_along(height); + let grid = CubeCount::new_2d(cubes_x, refs_y); + let dim = CubeDim::new_1d(64); + let scale = weight_scale(knobs.sigma, &profile); + let accum_scale = cross_frame_accum_scale(knobs.spatial_radius, fixture.radius); + let warp_uniform = knobs + .warp_uniform + .unwrap_or_else(|| needs_warp_uniform_search(&client)); + let grid_frame_count = grid_frames(fixture.radius); + unsafe { collab_fused::launch_unchecked::( &client, - CubeCount::new_2d(fused_cubes_x(w), refs_along(h)), - CubeDim::new_1d(64), + grid, + dim, 1usize, - ArrayArg::from_raw_parts(ring_buf, fx.ring.len()), - ArrayArg::from_raw_parts(mv_buf, fx.mv_field.len()), - ArrayArg::from_raw_parts(conf_buf, fx.confidence.len()), - ArrayArg::from_raw_parts(slots_buf, fx.neighbour_slots.len()), + ArrayArg::from_raw_parts(ring_buf, fixture.ring.len()), + ArrayArg::from_raw_parts(mv_buf, fixture.mv_field.len()), + ArrayArg::from_raw_parts(conf_buf, fixture.confidence.len()), + ArrayArg::from_raw_parts(slots_buf, fixture.neighbour_slots.len()), ArrayArg::from_raw_parts(sigma_buf, 1), ArrayArg::from_raw_parts(zero_curve, NOISE_CURVE_BINS), ArrayArg::from_raw_parts(unit_map_buf, map_len), @@ -138,30 +163,29 @@ fn run_fused_over(fx: &RingFixture, k: Knobs) -> FusedRun { ArrayArg::from_raw_parts(accum, pixels * frames), ArrayArg::from_raw_parts(wsum.clone(), pixels * frames), ArrayArg::from_raw_parts(group_weight.clone(), refs), - fx.centre_slot, - k.c_min, - k.lambda_ht, + fixture.centre_slot, + knobs.c_min, + knobs.lambda_ht, 0u32, STRENGTH_MAP_OFF, - weight_scale(k.sigma, &profile), - cross_frame_accum_scale(k.spatial_radius, fx.radius), - k.warp_uniform - .unwrap_or_else(|| needs_warp_uniform_search(&client)), - fx.radius, - grid_frames(fx.radius), - k.refine, - fx.mv_stride, - fx.conf_stride, + scale, + accum_scale, + warp_uniform, + fixture.radius, + grid_frame_count, + knobs.refine, + fixture.mv_stride, + fixture.conf_stride, BLK_STEP, - k.blksize, - fx.blocks_x, - fx.blocks_y, - w, - h, + knobs.blksize, + fixture.blocks_x, + fixture.blocks_y, + width, + height, 1u32, - k.k_max, + knobs.k_max, 1u32, - k.spatial_radius, + knobs.spatial_radius, refs_x, map_cols, map_rows, @@ -184,33 +208,27 @@ fn run_fused_over(fx: &RingFixture, k: Knobs) -> FusedRun { /// The temporal search looks where the motion field points. /// -/// `planted_ring` puts an exact copy of the reference patch in every -/// neighbour, shifted by `3 * k`, and seeds the motion field to predict -/// exactly that shift. A search that follows the prediction finds four -/// pixel-for-pixel copies of the reference patch, the whole group agrees, -/// and the Haar detail levels collapse to nothing, so the threshold keeps -/// very little and the group weight is high. -/// -/// The control zeroes the motion field, leaving every neighbour's refine -/// window over flat background instead. The copies still exist in the -/// ring, so this is a test of the prediction and not of whether the -/// content is reachable at all. +/// Every neighbour holds an exact copy of the reference patch shifted by `3 * k`, and the motion +/// field predicts that shift. Following it gives a group of exact copies whose Haar detail +/// collapses, so the group weight is high. The control zeroes the motion field while the copies +/// stay in the ring, so this tests the prediction and not whether the content is reachable. #[test] fn temporal_members_are_found_at_the_mv_prediction() { - let (w, h) = (96u32, 96u32); + let (width, height) = (96u32, 96u32); let radius = 2u32; let ref_pos = (64u32, 64u32); let patch = deterministic_texture(7); - let predicted = planted_ring(w, h, radius, ref_pos, 3, &patch, 0.2, |_| 1.0); - let mut blind = planted_ring(w, h, radius, ref_pos, 3, &patch, 0.2, |_| 1.0); + let predicted = planted_ring(width, height, radius, ref_pos, 3, &patch, 0.2, |_| 1.0); + let mut blind = planted_ring(width, height, radius, ref_pos, 3, &patch, 0.2, |_| 1.0); blind.mv_field.fill(0); - let refs_x = refs_along(w); + let refs_x = refs_along(width); let ref_idx = ((ref_pos.1 / STEP) * refs_x + (ref_pos.0 / STEP)) as usize; - let with_prediction = run_fused_over(&predicted, Knobs::default()).group_weight[ref_idx]; - let without = run_fused_over(&blind, Knobs::default()).group_weight[ref_idx]; + let knobs = Knobs::default(); + let with_prediction = run_fused_over(&predicted, &knobs).group_weight[ref_idx]; + let without = run_fused_over(&blind, &knobs).group_weight[ref_idx]; assert!( with_prediction > without * 1.5, @@ -220,28 +238,24 @@ fn temporal_members_are_found_at_the_mv_prediction() { ); } -/// A neighbour whose motion-block confidence sits below `c_min` is -/// skipped outright, so no member ever comes from it. +/// A neighbour whose motion-block confidence sits below `c_min` is skipped outright. /// -/// The confidence is uniform across every block of a neighbour's plane -/// here, so the skip is the same decision for every group in the frame -/// and that neighbour's whole region of the accumulator ring has to stay -/// exactly zero. The gated neighbour is k = -2, the first one searched, -/// which wins every tie on this fixture. Gating one of four neighbours -/// leaves every volume its three frames, so every other neighbour still -/// receives members. +/// The confidence is uniform per neighbour, so the gated neighbour's whole region of the ring +/// must stay exactly zero. The gated neighbour is k = -2, the first one searched, which wins every +/// tie on this fixture. Gating one of four neighbours leaves every volume its three frames, so +/// every other neighbour still receives members. #[test] fn low_confidence_neighbours_contribute_no_candidates() { - let (w, h) = (96u32, 96u32); + let (width, height) = (96u32, 96u32); let radius = 2u32; let ref_pos = (64u32, 64u32); let patch = deterministic_texture(11); - // Confidence 0.0 for k = -2, 1.0 for every other neighbour. - let fx = planted_ring(w, h, radius, ref_pos, 3, &patch, 0.2, |k| { + let fixture = planted_ring(width, height, radius, ref_pos, 3, &patch, 0.2, |k| { if k == -2 { 0.0 } else { 1.0 } }); - let run = run_fused_over(&fx, Knobs::default()); + let knobs = Knobs::default(); + let run = run_fused_over(&fixture, &knobs); for k in -(radius as i32)..=(radius as i32) { let slot = (k + radius as i32) as u32; @@ -266,164 +280,141 @@ fn low_confidence_neighbours_contribute_no_candidates() { fn a_volume_short_of_frames_sends_the_group_to_the_fallback() { let radius = 2u32; let patch = deterministic_texture(11); - let fx = planted_ring(96, 96, radius, (64, 64), 3, &patch, 0.2, |k| { + let fixture = planted_ring(96, 96, radius, (64, 64), 3, &patch, 0.2, |k| { if k > 0 { 0.0 } else { 1.0 } }); - let run = run_fused_over(&fx, Knobs::default()); + let knobs = Knobs::default(); + let run = run_fused_over(&fixture, &knobs); + let centre_weight = run.frame_weight_sum(fixture.centre_slot); + + assert!(centre_weight > 0, "the centre slot received nothing"); - assert!( - run.frame_weight_sum(fx.centre_slot) > 0, - "the centre slot received nothing" - ); for slot in 0..(2 * radius + 1) { - if slot == fx.centre_slot { + if slot == fixture.centre_slot { continue; } + + let weight = run.frame_weight_sum(slot); assert_eq!( - run.frame_weight_sum(slot), - 0, + weight, 0, "slot {slot} must receive nothing once every group falls back" ); } } -/// Every group fills to `k_max` however poor its candidates are, because -/// there is no admission gate. -/// -/// `noisy_ring` is built so no 8x8 window resembles any other, on any -/// frame, so every candidate is a bad match. `lambda_ht` is set high -/// enough that only the forced group DC survives the threshold, which -/// pins every group's retained variance at `sigma^2` and so every -/// group's weight at the same constant. The weight one member's patch -/// deposits is then the same fixed-point value everywhere, and the total -/// weight in the ring counts members outright. +/// Every group fills to `k_max` however poor its candidates are, because there is no admission +/// gate. /// -/// A run capped at `k_max = 1` holds every group to its self-match, so -/// the eight-member run has to deposit exactly eight times as much. An -/// admission gate anywhere would leave some group short and break the -/// ratio. +/// No 8x8 window of `noisy_ring` resembles any other, so every candidate is a bad match. A huge +/// `lambda_ht` keeps only the forced group DC, which pins every group's weight at the same +/// constant, so the total weight in the ring counts members. A run capped at `k_max = 1` holds +/// every group to its self-match, so the full run must deposit exactly eight times as much. #[test] fn no_admission_gate_means_the_group_always_fills() { - let (w, h) = (64u32, 64u32); + let (width, height) = (64u32, 64u32); let radius = 2u32; - let fx = noisy_ring(w, h, radius, 1.0); + let fixture = noisy_ring(width, height, radius, 1.0); - // The smallest search space any reference here sees is the 5x5 - // rectangle a corner clips to, so every group has at least eight - // positions to choose from and rounds up to a full stack. - let full = run_fused_over( - &fx, - Knobs { - lambda_ht: 1.0e6, - ..Knobs::default() - }, - ); - let single = run_fused_over( - &fx, - Knobs { - k_max: 1, - lambda_ht: 1.0e6, - ..Knobs::default() - }, - ); + // The smallest search space here is the 5x5 rectangle a corner clips to, so every group has + // at least eight positions to choose from. + let full_knobs = Knobs { + lambda_ht: 1.0e6, + ..Knobs::default() + }; + let single_knobs = Knobs { + k_max: 1, + lambda_ht: 1.0e6, + ..Knobs::default() + }; + let full = run_fused_over(&fixture, &full_knobs); + let single = run_fused_over(&fixture, &single_knobs); let one = single.total_weight(); assert!(one > 0, "the k_max = 1 run deposited no weight at all"); + + let full_weight = full.total_weight(); assert_eq!( - full.total_weight(), + full_weight, one * K_MAX as i64, "expected every group to carry {K_MAX} members, so {K_MAX}x the weight the \ one-member run deposited" ); } -/// Sets one block's vector toward neighbour `t`. -fn set_block_mv(fx: &mut RingFixture, t: u32, bx: u32, by: u32, mv: [i32; 2]) { - let block = by * fx.blocks_x + bx; - let base = (t * fx.mv_stride + block * 2) as usize; - fx.mv_field[base] = mv[0]; - fx.mv_field[base + 1] = mv[1]; +/// Sets the vector of block `(block_x, block_y)` toward neighbour index `neighbour`. +fn set_block_mv(fixture: &mut RingFixture, neighbour: u32, block_x: u32, block_y: u32, vector: [i32; 2]) { + let block = block_y * fixture.blocks_x + block_x; + let base = (neighbour * fixture.mv_stride + block * 2) as usize; + fixture.mv_field[base] = vector[0]; + fixture.mv_field[base + 1] = vector[1]; } -/// Writes an 8x8 patch into ring slot `slot` at `(px, py)`. -fn plant_in_slot(fx: &mut RingFixture, slot: u32, px: u32, py: u32, patch: &[f32; 64]) { - let pixels = (fx.width * fx.height) as usize; - let frame = &mut fx.ring[slot as usize * pixels..(slot as usize + 1) * pixels]; +/// Writes an 8x8 patch into ring slot `slot` with its top-left corner at `(x, y)`. +fn plant_in_slot(fixture: &mut RingFixture, slot: u32, x: u32, y: u32, patch: &[f32; 64]) { + let pixels = (fixture.width * fixture.height) as usize; + let frame = &mut fixture.ring[slot as usize * pixels..(slot as usize + 1) * pixels]; for row in 0..8u32 { for col in 0..8u32 { - frame[((py + row) * fx.width + px + col) as usize] = patch[(row * 8 + col) as usize]; + frame[((y + row) * fixture.width + x + col) as usize] = patch[(row * 8 + col) as usize]; } } } -/// Moves each neighbour's copy of the reference patch 20 pixels right, -/// leaving flat background where the reference sits, and points one -/// block's vector at the copy. +/// Moves each neighbour's copy of the reference patch 20 pixels right and points one block's +/// vector at it. /// -/// `planted_ring` at a zero shift puts a copy at the reference position -/// in every frame, so the copy there is erased first. Every block but -/// `(bx, by)` then holds the zeroed vector `planted_ring` left, which -/// points at flat background, so the copy is reachable only through -/// `(bx, by)`. +/// The copy at the reference position is erased first. Every other block keeps the zero vector, +/// which points at flat background, so the copy is reachable only through `block`. fn only_reachable_through( - fx: &mut RingFixture, + fixture: &mut RingFixture, ref_pos: (u32, u32), patch: &[f32; 64], - (bx, by): (u32, u32), + (block_x, block_y): (u32, u32), ) { let flat = [0.2f32; 64]; - for t in 0..fx.neighbour_slots.len() as u32 { - let slot = fx.neighbour_slots[t as usize]; - plant_in_slot(fx, slot, ref_pos.0, ref_pos.1, &flat); - plant_in_slot(fx, slot, ref_pos.0 + 20, ref_pos.1, patch); - set_block_mv(fx, t, 8, 8, [0, 0]); - set_block_mv(fx, t, bx, by, [20, 0]); + for neighbour in 0..fixture.neighbour_slots.len() as u32 { + let slot = fixture.neighbour_slots[neighbour as usize]; + plant_in_slot(fixture, slot, ref_pos.0, ref_pos.1, &flat); + plant_in_slot(fixture, slot, ref_pos.0 + 20, ref_pos.1, patch); + set_block_mv(fixture, neighbour, 8, 8, [0, 0]); + set_block_mv(fixture, neighbour, block_x, block_y, [20, 0]); } } -/// The corner block's vector points at flat background, and only a -/// neighbouring covering block's vector points at the planted copy. -/// -/// The reference at (64, 64) sits on the corner of block (8, 8) and is -/// also covered by blocks (7, 7), (8, 7) and (7, 8), since a 16-pixel -/// block at an 8-pixel step covers two patches per axis. A search that -/// reads only the corner block never sees the copy. +/// The corner block's vector points at flat background, and only a neighbouring covering block's +/// vector points at the planted copy. /// -/// Each of the three non-corner covering blocks is tried on its own, -/// `(7, 7)` diagonally, `(8, 7)` above and `(7, 8)` to the left, so a -/// kernel that read only the corner and the diagonal fails on two of -/// the three. +/// A 16-pixel block at an 8-pixel step covers two patches per axis, so the reference at (64, 64) +/// is covered by blocks (8, 8), (7, 7), (8, 7) and (7, 8). Each non-corner block is tried on its +/// own, so a kernel that read only the corner and the diagonal fails on two of the three. /// -/// The ring runs at radius 2, so the reference's volume keeps three of -/// the four copies the covering block reaches. A static twin at -/// [TWIN_POS] anchors the second volume on the same patch in every -/// frame, so that volume is identical in every run and the group weight -/// only moves with the reference's own volume. The control leaves every -/// block on the corner's zeroed vector, so no rectangle reaches the copy -/// however many blocks are read. +/// The ring runs at radius 2, so the reference's volume keeps three of the four copies the +/// covering block reaches. A static twin at [TWIN_POS] keeps the second volume identical in every +/// run, so the group weight only moves with the reference's own volume. The control leaves every +/// block on the zero vector, so no rectangle reaches the copy. #[test] fn a_covering_block_other_than_the_corner_finds_the_match() { - let (w, h) = (96u32, 96u32); + let (width, height) = (96u32, 96u32); let radius = 2u32; let ref_pos = (64u32, 64u32); let patch = deterministic_texture(13); - let refs_x = refs_along(w); + let refs_x = refs_along(width); let ref_idx = ((ref_pos.1 / STEP) * refs_x + (ref_pos.0 / STEP)) as usize; - // The same ring with every vector zeroed, so no block's rectangle - // reaches the copy however many blocks are read. - let mut corner_only = planted_ring(w, h, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); + let knobs = twin_knobs(); + + let mut corner_only = planted_ring(width, height, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); only_reachable_through(&mut corner_only, ref_pos, &patch, (8, 8)); plant_static_twins(&mut corner_only, &[TWIN_POS], &patch); corner_only.mv_field.fill(0); - let without = run_fused_over(&corner_only, twin_knobs()).group_weight[ref_idx]; + let without = run_fused_over(&corner_only, &knobs).group_weight[ref_idx]; for block in [(7u32, 7u32), (8, 7), (7, 8)] { - let mut fx = planted_ring(w, h, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); - only_reachable_through(&mut fx, ref_pos, &patch, block); - plant_static_twins(&mut fx, &[TWIN_POS], &patch); - let with_covering = run_fused_over(&fx, twin_knobs()).group_weight[ref_idx]; + let mut fixture = planted_ring(width, height, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); + only_reachable_through(&mut fixture, ref_pos, &patch, block); + plant_static_twins(&mut fixture, &[TWIN_POS], &patch); + let with_covering = run_fused_over(&fixture, &knobs).group_weight[ref_idx]; assert!( with_covering > without * 1.5, @@ -433,41 +424,38 @@ fn a_covering_block_other_than_the_corner_finds_the_match() { } } -/// Two covering blocks whose vectors differ by one pixel give -/// overlapping rectangles, and the reference's volume still finds the -/// copy they both reach. +/// Two covering blocks whose vectors differ by one pixel give overlapping rectangles, and the +/// reference's volume still finds the copy they both reach. /// -/// Block `(7, 7)` is visited first, so with a second vector its -/// rectangle reaches the copy and block `(8, 8)` then skips the overlap. -/// The reference's volume must hold the same copy either way. Three -/// static twins anchor the other three volumes on the same patch in -/// every frame, so the group weight only moves with the reference's own -/// volume, and two runs holding the same patches carry the same weight. -/// The control points no block at the copy, which shows the weight can -/// see the copy go missing. +/// Block `(7, 7)` is visited first, so with a second vector its rectangle reaches the copy and +/// block `(8, 8)` then skips the overlap. Three static twins keep the other volumes identical in +/// every run, so the group weight only moves with the reference's own volume. The control points +/// no block at the copy, which shows the weight can see the copy go missing. #[test] fn overlapping_covering_rectangles_still_find_the_match() { - let (w, h) = (96u32, 96u32); + let (width, height) = (96u32, 96u32); let radius = 1u32; let ref_pos = (64u32, 64u32); let patch = deterministic_texture(17); - let refs_x = refs_along(w); + let refs_x = refs_along(width); let ref_idx = ((ref_pos.1 / STEP) * refs_x + (ref_pos.0 / STEP)) as usize; let flat = [0.2f32; 64]; let build = |second_vector: Option<[i32; 2]>| { - let mut fx = planted_ring(w, h, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); - plant_static_twins(&mut fx, &[(64, 48), (48, 64), (48, 48)], &patch); - for t in 0..2u32 { - let slot = fx.neighbour_slots[t as usize]; - plant_in_slot(&mut fx, slot, ref_pos.0, ref_pos.1, &flat); - plant_in_slot(&mut fx, slot, ref_pos.0 + 20, ref_pos.1, &patch); - set_block_mv(&mut fx, t, 8, 8, [20, 0]); - if let Some(v) = second_vector { - set_block_mv(&mut fx, t, 7, 7, v); + let mut fixture = planted_ring(width, height, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); + plant_static_twins(&mut fixture, &[(64, 48), (48, 64), (48, 48)], &patch); + + for neighbour in 0..2u32 { + let slot = fixture.neighbour_slots[neighbour as usize]; + plant_in_slot(&mut fixture, slot, ref_pos.0, ref_pos.1, &flat); + plant_in_slot(&mut fixture, slot, ref_pos.0 + 20, ref_pos.1, &patch); + set_block_mv(&mut fixture, neighbour, 8, 8, [20, 0]); + if let Some(vector) = second_vector { + set_block_mv(&mut fixture, neighbour, 7, 7, vector); } } - fx + + fixture }; let mut unreachable = build(None); @@ -476,9 +464,10 @@ fn overlapping_covering_rectangles_still_find_the_match() { let one_covering_block = build(None); let two_covering_blocks = build(Some([21, 0])); - let one = run_fused_over(&one_covering_block, twin_knobs()).group_weight[ref_idx]; - let two = run_fused_over(&two_covering_blocks, twin_knobs()).group_weight[ref_idx]; - let none = run_fused_over(&unreachable, twin_knobs()).group_weight[ref_idx]; + let knobs = twin_knobs(); + let one = run_fused_over(&one_covering_block, &knobs).group_weight[ref_idx]; + let two = run_fused_over(&two_covering_blocks, &knobs).group_weight[ref_idx]; + let none = run_fused_over(&unreachable, &knobs).group_weight[ref_idx]; assert_eq!( two, one, @@ -496,41 +485,39 @@ fn overlapping_covering_rectangles_still_find_the_match() { /// /// Each becomes an exact spatial twin of a reference carrying the same patch, and its volume holds /// that patch in every frame as long as its covering blocks carry a zero vector. -fn plant_static_twins(fx: &mut RingFixture, positions: &[(u32, u32)], patch: &[f32; 64]) { - let frames = 2 * fx.radius + 1; +fn plant_static_twins(fixture: &mut RingFixture, positions: &[(u32, u32)], patch: &[f32; 64]) { + let frames = 2 * fixture.radius + 1; for slot in 0..frames { - for &(px, py) in positions { - plant_in_slot(fx, slot, px, py, patch); + for &(x, y) in positions { + plant_in_slot(fixture, slot, x, y, patch); } } } -/// With `blksize == step` exactly one block covers a patch, so a -/// neighbouring block's vector is never consulted. +/// With `blksize == step` exactly one block covers a patch, so a neighbouring block's vector is +/// never consulted. /// -/// The copy is reachable only through block `(7, 7)`, which covers the -/// patch at `blksize = 16` and does not at `blksize = 8`. +/// The copy is reachable only through block `(7, 7)`, which covers the patch at `blksize = 16` +/// and does not at `blksize = 8`. #[test] fn a_block_size_equal_to_the_step_reads_only_the_corner_block() { - let (w, h) = (96u32, 96u32); + let (width, height) = (96u32, 96u32); let radius = 2u32; let ref_pos = (64u32, 64u32); let patch = deterministic_texture(19); - let refs_x = refs_along(w); + let refs_x = refs_along(width); let ref_idx = ((ref_pos.1 / STEP) * refs_x + (ref_pos.0 / STEP)) as usize; - let mut fx = planted_ring(w, h, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); - only_reachable_through(&mut fx, ref_pos, &patch, (7, 7)); + let mut fixture = planted_ring(width, height, radius, ref_pos, 0, &patch, 0.2, |_| 1.0); + only_reachable_through(&mut fixture, ref_pos, &patch, (7, 7)); - let covering = run_fused_over(&fx, Knobs::default()).group_weight[ref_idx]; - let single = run_fused_over( - &fx, - Knobs { - blksize: BLK_STEP, - ..Knobs::default() - }, - ) - .group_weight[ref_idx]; + let covering_knobs = Knobs::default(); + let single_block_knobs = Knobs { + blksize: BLK_STEP, + ..Knobs::default() + }; + let covering = run_fused_over(&fixture, &covering_knobs).group_weight[ref_idx]; + let single = run_fused_over(&fixture, &single_block_knobs).group_weight[ref_idx]; assert!( covering > single * 1.5, @@ -539,52 +526,50 @@ fn a_block_size_equal_to_the_step_reads_only_the_corner_block() { ); } -/// Where [twin_ring] plants the reference's spatial twin in the centre frame. -const TWIN_POS: (u32, u32) = (64, 48); - /// A radius-2 ring whose centre frame holds the reference texture at (64, 64) and an exact twin /// at [TWIN_POS], so the spatial search anchors the second volume on the twin. /// -/// Every block moves the reference's copy by `(3k, 0)`, as `planted_ring` places it. The four -/// blocks covering the twin, `(7..=8, 5..=6)`, carry `twin_mv(k)` instead, and a copy of the twin -/// sits at `TWIN_POS + twin_copy(k)` in each neighbour. The two can differ, so a test can point -/// the twin's motion away from its copies. +/// Every block moves the reference's copy by `(3k, 0)`. The four blocks covering the twin, +/// `(7..=8, 5..=6)`, carry `twin_mv(k)` instead, and a copy of the twin sits at +/// `TWIN_POS + twin_copy(k)` in each neighbour. The two can differ, so a test can point the +/// twin's motion away from its copies. fn twin_ring( patch: &[f32; 64], twin_mv: impl Fn(i32) -> [i32; 2], twin_copy: impl Fn(i32) -> [i32; 2], ) -> RingFixture { let radius = 2u32; - let mut fx = planted_ring(96, 96, radius, (64, 64), 3, patch, 0.2, |_| 1.0); - let centre_slot = fx.centre_slot; - plant_in_slot(&mut fx, centre_slot, TWIN_POS.0, TWIN_POS.1, patch); + let mut fixture = planted_ring(96, 96, radius, (64, 64), 3, patch, 0.2, |_| 1.0); + let centre_slot = fixture.centre_slot; + plant_in_slot(&mut fixture, centre_slot, TWIN_POS.0, TWIN_POS.1, patch); for k in [-2i32, -1, 1, 2] { - let t = neighbour_idx_for_k(radius, k); - let slot = fx.neighbour_slots[t as usize]; + let neighbour = neighbour_idx_for_k(radius, k); + let slot = fixture.neighbour_slots[neighbour as usize]; - for by in 0..fx.blocks_y { - for bx in 0..fx.blocks_x { - set_block_mv(&mut fx, t, bx, by, [3 * k, 0]); + for block_y in 0..fixture.blocks_y { + for block_x in 0..fixture.blocks_x { + set_block_mv(&mut fixture, neighbour, block_x, block_y, [3 * k, 0]); } } - for by in 5..=6u32 { - for bx in 7..=8u32 { - set_block_mv(&mut fx, t, bx, by, twin_mv(k)); + let twin_vector = twin_mv(k); + for block_y in 5..=6u32 { + for block_x in 7..=8u32 { + set_block_mv(&mut fixture, neighbour, block_x, block_y, twin_vector); } } let [copy_dx, copy_dy] = twin_copy(k); let copy_x = (TWIN_POS.0 as i32 + copy_dx) as u32; let copy_y = (TWIN_POS.1 as i32 + copy_dy) as u32; - plant_in_slot(&mut fx, slot, copy_x, copy_y, patch); + plant_in_slot(&mut fixture, slot, copy_x, copy_y, patch); } - fx + fixture } -/// The knobs every twin-ring run shares, a spatial window wide enough to reach the twin. +/// A spatial window wide enough to reach the twin. fn twin_knobs() -> Knobs { Knobs { spatial_radius: 16, @@ -606,8 +591,9 @@ fn each_volume_follows_its_own_anchors_motion() { let own = twin_ring(&patch, |k| [0, 2 * k], |k| [0, 2 * k]); let borrowed = twin_ring(&patch, |k| [3 * k, 0], |k| [0, 2 * k]); - let with_own = run_fused_over(&own, twin_knobs()).group_weight[ref_idx]; - let with_borrowed = run_fused_over(&borrowed, twin_knobs()).group_weight[ref_idx]; + let knobs = twin_knobs(); + let with_own = run_fused_over(&own, &knobs).group_weight[ref_idx]; + let with_borrowed = run_fused_over(&borrowed, &knobs).group_weight[ref_idx]; assert!( with_own > with_borrowed * 1.5, @@ -621,9 +607,8 @@ fn each_volume_follows_its_own_anchors_motion() { /// Offsetting the copies of k = -2, the first neighbour searched and so the one that wins every /// tie, must leave the group untouched, because both volumes skip that frame for the three exact /// ones. A volume that kept its first three frames would hold the offset one instead. Two groups -/// holding the same patches carry the same weight, so the weight equals the run with no offset. -/// Offsetting k = -1 as well forces an offset frame into each volume, which moves the weight and -/// shows the comparison can see a change. +/// holding the same patches carry the same weight. Offsetting k = -1 as well forces an offset +/// frame into each volume, which shows the comparison can see a change. #[test] fn a_volume_keeps_its_best_frames() { let patch = deterministic_texture(29); @@ -637,9 +622,10 @@ fn a_volume_keeps_its_best_frames() { offset_copies(&mut two_offset, -2); offset_copies(&mut two_offset, -1); - let clean_weight = run_fused_over(&clean, twin_knobs()).group_weight[ref_idx]; - let one_weight = run_fused_over(&one_offset, twin_knobs()).group_weight[ref_idx]; - let two_weight = run_fused_over(&two_offset, twin_knobs()).group_weight[ref_idx]; + let knobs = twin_knobs(); + let clean_weight = run_fused_over(&clean, &knobs).group_weight[ref_idx]; + let one_weight = run_fused_over(&one_offset, &knobs).group_weight[ref_idx]; + let two_weight = run_fused_over(&two_offset, &knobs).group_weight[ref_idx]; assert_eq!( one_weight, clean_weight, @@ -653,11 +639,11 @@ fn a_volume_keeps_its_best_frames() { } /// Raises every texture pixel of neighbour `k`'s frame by 0.1, leaving the background alone. -fn offset_copies(fx: &mut RingFixture, k: i32) { - let pixels = (fx.width * fx.height) as usize; - let t = neighbour_idx_for_k(fx.radius, k); - let slot = fx.neighbour_slots[t as usize] as usize; - let frame = &mut fx.ring[slot * pixels..(slot + 1) * pixels]; +fn offset_copies(fixture: &mut RingFixture, k: i32) { + let pixels = (fixture.width * fixture.height) as usize; + let neighbour = neighbour_idx_for_k(fixture.radius, k); + let slot = fixture.neighbour_slots[neighbour as usize] as usize; + let frame = &mut fixture.ring[slot * pixels..(slot + 1) * pixels]; for value in frame.iter_mut() { if *value > 0.5 { *value += 0.1; @@ -668,8 +654,8 @@ fn offset_copies(fx: &mut RingFixture, k: i32) { /// A neighbour patch the first volume took is not reused by the second. /// /// The twin's blocks point at the reference's copies, which match the twin exactly. Reusing them -/// would make the twin's volume hold the same pixels as the control's, where the twin follows its -/// own exact copies, and the two weights would be equal. Skipping them leaves the twin's volume one +/// would give the twin's volume the same pixels as the control's, where the twin follows its own +/// exact copies, and the two weights would be equal. Skipping them leaves the twin's volume one /// unclaimed copy and two near misses, so its weight drops below the control's. #[test] fn a_position_claimed_by_an_earlier_volume_is_not_reused() { @@ -681,8 +667,9 @@ fn a_position_claimed_by_an_earlier_volume_is_not_reused() { let bait = twin_ring(&patch, |k| [3 * k, 16], |k| [3 * k, 16]); let control = twin_ring(&patch, |k| [0, 2 * k], |k| [0, 2 * k]); - let with_bait = run_fused_over(&bait, twin_knobs()).group_weight[ref_idx]; - let with_control = run_fused_over(&control, twin_knobs()).group_weight[ref_idx]; + let knobs = twin_knobs(); + let with_bait = run_fused_over(&bait, &knobs).group_weight[ref_idx]; + let with_control = run_fused_over(&control, &knobs).group_weight[ref_idx]; assert!( with_bait < with_control * 0.99, @@ -691,13 +678,10 @@ fn a_position_claimed_by_an_earlier_volume_is_not_reused() { ); } -/// The bait fixture above under both search walks, so the position skip -/// it relies on is not an artefact of whichever walk the runtime happens -/// to pick. +/// The claimed-position bait fixture under both search walks. /// -/// The bait's claimed positions are exactly where the two walks take -/// different turns to reach the same candidates, which is where a walk -/// that skipped the claim check differently would show up. +/// The claimed positions are exactly where the two walks take different turns to reach the same +/// candidates, so a walk that skipped the claim check differently would show up here. #[test] fn a_position_claimed_by_an_earlier_volume_agrees_across_search_walks() { let patch = deterministic_texture(31); @@ -712,8 +696,8 @@ fn a_position_claimed_by_an_earlier_volume_agrees_across_search_walks() { ..twin_knobs() }; - let clipped = run_fused_over(&bait, clipped_knobs); - let uniform = run_fused_over(&bait, uniform_knobs); + let clipped = run_fused_over(&bait, &clipped_knobs); + let uniform = run_fused_over(&bait, &uniform_knobs); assert_eq!( clipped.group_weight, uniform.group_weight, @@ -724,7 +708,7 @@ fn a_position_claimed_by_an_earlier_volume_agrees_across_search_walks() { "the two search walks scattered different weights" ); assert!( - uniform.group_weight.iter().any(|&w| w > 0.0), + uniform.group_weight.iter().any(|&weight| weight > 0.0), "neither walk aggregated anything, so agreeing proves nothing" ); } diff --git a/av-denoise-core/src/nl4d/tests/helpers.rs b/av-denoise-core/src/nl4d/tests/helpers.rs index 9e89624..94e9743 100644 --- a/av-denoise-core/src/nl4d/tests/helpers.rs +++ b/av-denoise-core/src/nl4d/tests/helpers.rs @@ -3,6 +3,7 @@ use cubecl::wgpu::WgpuRuntime; use crate::nl4d::Nl4dParams; use crate::nlmeans::motion::neighbour_idx_for_k; +use crate::nlmeans::tests::helpers::noisy_field_over; use crate::nlmeans::{ ChannelMode, HqParams, @@ -12,39 +13,48 @@ use crate::nlmeans::{ PrefilterMode, }; +pub(super) type R = WgpuRuntime; + pub(super) const SIGMA: f32 = 6.0 / 255.0; pub(super) const SPATIAL_RADIUS: u32 = 9; pub(super) const REFINE: u32 = 2; pub(super) const C_MIN: f32 = 0.05; pub(super) const LAMBDA_HT: f32 = 2.7; +/// The motion block step, equal to [PATCH_SIZE](crate::collab::PATCH_SIZE) so a block boundary +/// always lines up with a patch boundary. +pub(super) const BLK_STEP: u32 = 8; + /// Parameters for a still clip with sigma pinned to [SIGMA]. pub(super) fn static_clip_params(temporal_radius: u32) -> Nl4dParams { + let motion_compensation = MotionCompensationMode::Mvtools { + blksize: 16, + overlap: 8, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Auto, + }; + let hq = HqParams::with_sigma(SIGMA); + let nlm = NlmParams { + temporal_radius, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation, + hq: Some(hq), + }; + Nl4dParams { - nlm: NlmParams { - temporal_radius, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 16, - overlap: 8, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Auto, - }, - hq: Some(HqParams::with_sigma(SIGMA)), - }, + nlm, temporal_radius, refine: REFINE, spatial_radius: SPATIAL_RADIUS, lambda_ht: LAMBDA_HT, c_min: C_MIN, - // The shipped default, so these run the aggregation a real - // caller gets. + // The shipped default, so these run the aggregation a real caller gets. kaiser_beta: 2.0, field_lambda: 0.0, // No effect here, since sigma is pinned. @@ -55,91 +65,58 @@ pub(super) fn static_clip_params(temporal_radius: u32) -> Nl4dParams { // Off, so the pipeline tests keep the flat map their expectations were recorded against. flat_texture_cut: 1.0, // Off, so the pipeline tests keep the per-coefficient kernel their expectations were - // recorded against. `nl4d::tests::pooled` covers the pooled path. + // recorded against. pooled_threshold: false, grain_export: false, } } -/// A non-flat luma field, built from two out-of-phase sine waves rather -/// than noise, so it carries real spatial structure a denoiser can -/// either preserve or destroy. -pub(crate) fn textured_base(w: u32, h: u32) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let fx = x as f32 / w as f32; - let fy = y as f32 / h as f32; - let v = 0.5 - + 0.15 * (fx * 6.0 * std::f32::consts::PI).sin() * (fy * 4.0 * std::f32::consts::PI).cos(); - frame[(y * w + x) as usize] = v.clamp(0.05, 0.95); +/// A non-flat luma field built from two out-of-phase sine waves. +/// +/// It carries real spatial structure rather than noise, which a denoiser can either preserve or +/// destroy. +pub(crate) fn textured_base(width: u32, height: u32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let x_fraction = x as f32 / width as f32; + let y_fraction = y as f32 / height as f32; + let value = 0.5 + + 0.15 + * (x_fraction * 6.0 * std::f32::consts::PI).sin() + * (y_fraction * 4.0 * std::f32::consts::PI).cos(); + frame[(y * width + x) as usize] = value.clamp(0.05, 0.95); } } - frame -} -/// Adds independent pseudo-Gaussian noise to `base`, decorrelated across -/// `seed` so different seeds over the same base give independently -/// noisy copies of the same clean content. -pub(crate) fn noisy_copy_of(base: &[f32], w: u32, h: u32, sigma: f32, seed: u32) -> Vec { - let mut frame = vec![0.0f32; base.len()]; - for idx in 0..(w * h) { - let noise = unit_noise(idx, seed); - frame[idx as usize] = (base[idx as usize] + noise * sigma).clamp(0.0, 1.0); - } frame } -/// A pseudo-Gaussian sample with unit standard deviation for pixel `idx` -/// under `seed`. -pub(super) fn unit_noise(idx: u32, seed: u32) -> f32 { - let unit_std = (1.0f32 / 3.0f32).sqrt(); - let mut sum = 0.0f32; - for k in 0..4u32 { - let mut hash = (idx * 4 + k) - .wrapping_mul(2654435761) - .wrapping_add(seed.wrapping_mul(0x9E37_79B9).wrapping_add(k)); - hash ^= hash >> 15; - hash = hash.wrapping_mul(0x85EB_CA6B); - hash ^= hash >> 13; - sum += (hash as f32 / u32::MAX as f32) - 0.5; - } - sum / unit_std -} - /// PSNR between two equal-length planes, in dB. -pub(super) fn psnr(a: &[f32], b: &[f32]) -> f64 { - let mse: f64 = a +pub(super) fn psnr(output: &[f32], reference: &[f32]) -> f64 { + let mse: f64 = output .iter() - .zip(b.iter()) - .map(|(&x, &y)| (x as f64 - y as f64).powi(2)) + .zip(reference.iter()) + .map(|(&out, &expected)| (out as f64 - expected as f64).powi(2)) .sum::() - / a.len() as f64; + / output.len() as f64; if mse <= 0.0 { return f64::INFINITY; } + 10.0 * (1.0f64 / mse).log10() } -pub(super) type R = WgpuRuntime; - pub(super) fn make_client() -> ComputeClient { let device = ::Device::default(); R::client(&device) } -/// The block step and grid this tree's fixtures use, matching -/// [`crate::collab::PATCH_SIZE`] so a block boundary always lines up -/// with a patch boundary. -pub(super) const BLK_STEP: u32 = 8; - -/// A ring of `2 * radius + 1` frames, one physical slot per logical -/// temporal offset from `-radius` to `radius`, with the centre frame at -/// physical slot `radius`. +/// A ring of `2 * radius + 1` frames with the centre frame at physical slot `radius`. /// /// Every field is already shaped the way -/// [`crate::collab::kernels::fused::collab_fused`] expects to read it, -/// so a test only has to upload each `Vec` and launch. +/// [collab_fused](crate::collab::kernels::fused::collab_fused) reads it, so a test only uploads +/// each `Vec` and launches. pub(super) struct RingFixture { pub ring: Vec, pub mv_field: Vec, @@ -155,74 +132,71 @@ pub(super) struct RingFixture { pub height: u32, } -/// A deterministic 8x8 texture with values well clear of the flat -/// background these fixtures plant it over. +/// A deterministic 8x8 texture with values well clear of the flat background it is planted over. pub(super) fn deterministic_texture(seed: u32) -> [f32; 64] { - let mut out = [0.0f32; 64]; - for (idx, v) in out.iter_mut().enumerate() { + let mut texture = [0.0f32; 64]; + for (idx, value) in texture.iter_mut().enumerate() { let mut hash = (idx as u32) .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(0x9E37_79B9)); hash ^= hash >> 15; hash = hash.wrapping_mul(0x85EBCA6B); hash ^= hash >> 13; - *v = 0.6 + (hash as f32 / u32::MAX as f32) * 0.3; + *value = 0.6 + (hash as f32 / u32::MAX as f32) * 0.3; } - out + + texture } -/// Writes an 8x8 patch into `frame` with its top-left corner at `(px, -/// py)`. -fn plant_patch(frame: &mut [f32], w: u32, px: u32, py: u32, patch: &[f32; 64]) { +/// Writes an 8x8 patch into `frame` with its top-left corner at `(x, y)`. +fn plant_patch(frame: &mut [f32], width: u32, x: u32, y: u32, patch: &[f32; 64]) { for row in 0..8u32 { for col in 0..8u32 { - let idx = (py + row) * w + (px + col); + let idx = (y + row) * width + (x + col); frame[idx as usize] = patch[(row * 8 + col) as usize]; } } } -/// Builds a ring whose centre frame carries a distinctive 8x8 patch at -/// `ref_pos`, and whose neighbour frame for logical offset `k` carries -/// the same patch shifted by `shift_per_k * k` pixels on the x axis. +/// Builds a ring whose centre frame carries `patch` at `ref_pos`, and whose neighbour at logical +/// offset `k` carries it shifted by `shift_per_k * k` pixels along x. /// -/// The motion field is seeded to predict exactly that shift, at the -/// block covering `ref_pos`, so a correct search recovers the planted -/// patch through the motion prediction, not through luck. -/// -/// `conf` gives the per-neighbour confidence written into every block of -/// that neighbour's plane, keyed by the same logical offset `k` the -/// shift is keyed by. -#[expect(clippy::too_many_arguments)] +/// The motion field predicts exactly that shift at the block covering `ref_pos`, so a correct +/// search recovers the patch through the prediction, not through luck. `confidence_for(k)` is +/// written into every block of neighbour `k`'s confidence plane. +#[expect( + clippy::too_many_arguments, + reason = "the test helper takes the full set of parameters its cases vary" +)] pub(super) fn planted_ring( - w: u32, - h: u32, + width: u32, + height: u32, radius: u32, ref_pos: (u32, u32), shift_per_k: i32, patch: &[f32; 64], background: f32, - conf: impl Fn(i32) -> f32, + confidence_for: impl Fn(i32) -> f32, ) -> RingFixture { - let n_frames = 2 * radius + 1; + let frame_count = 2 * radius + 1; let centre_slot = radius; - let blocks_x = w.div_ceil(BLK_STEP); - let blocks_y = h.div_ceil(BLK_STEP); + let blocks_x = width.div_ceil(BLK_STEP); + let blocks_y = height.div_ceil(BLK_STEP); let mv_stride = blocks_x * blocks_y * 2; let conf_stride = blocks_x * blocks_y; - let (rx, ry) = ref_pos; - let bx = rx / BLK_STEP; - let by = ry / BLK_STEP; - let block = by * blocks_x + bx; + let (ref_x, ref_y) = ref_pos; + let block_x = ref_x / BLK_STEP; + let block_y = ref_y / BLK_STEP; + let block = block_y * blocks_x + block_x; - let mut ring = vec![0.0f32; (n_frames * w * h) as usize]; - for slot in 0..n_frames { + let mut ring = vec![0.0f32; (frame_count * width * height) as usize]; + for slot in 0..frame_count { let k = slot as i32 - radius as i32; - let frame = &mut ring[(slot * w * h) as usize..((slot + 1) * w * h) as usize]; + let frame = &mut ring[(slot * width * height) as usize..((slot + 1) * width * height) as usize]; frame.fill(background); - let px = (rx as i32 + shift_per_k * k) as u32; - plant_patch(frame, w, px, ry, patch); + let patch_x = (ref_x as i32 + shift_per_k * k) as u32; + plant_patch(frame, width, patch_x, ref_y, patch); } let mut mv_field = vec![0i32; (2 * radius * mv_stride) as usize]; @@ -232,16 +206,18 @@ pub(super) fn planted_ring( if k == 0 { continue; } - let t = neighbour_idx_for_k(radius, k); + + let neighbour = neighbour_idx_for_k(radius, k); let slot = (k + radius as i32) as u32; - neighbour_slots[t as usize] = slot; + neighbour_slots[neighbour as usize] = slot; - let mv_base = (t * mv_stride + block * 2) as usize; + let mv_base = (neighbour * mv_stride + block * 2) as usize; mv_field[mv_base] = shift_per_k * k; mv_field[mv_base + 1] = 0; - let c_base = t * conf_stride; - confidence[c_base as usize..(c_base + conf_stride) as usize].fill(conf(k)); + let conf_base = neighbour * conf_stride; + let neighbour_confidence = confidence_for(k); + confidence[conf_base as usize..(conf_base + conf_stride) as usize].fill(neighbour_confidence); } RingFixture { @@ -255,33 +231,30 @@ pub(super) fn planted_ring( blocks_y, mv_stride, conf_stride, - width: w, - height: h, + width, + height, } } -/// A ring of independent pseudo-random frames, with a zeroed motion -/// field and uniform confidence. +/// A ring of independent pseudo-random frames, with a zeroed motion field and uniform confidence. /// -/// No 8x8 window into this ring resembles any other, on any frame, so -/// every candidate a search finds is a poor match. It exists for the -/// no-admission-gate test, where the point is that the group still -/// fills to `k_max` despite that. -pub(super) fn noisy_ring(w: u32, h: u32, radius: u32, confidence_value: f32) -> RingFixture { - let n_frames = 2 * radius + 1; +/// No 8x8 window into this ring resembles any other, on any frame, so every candidate a search +/// finds is a poor match. +pub(super) fn noisy_ring(width: u32, height: u32, radius: u32, confidence_value: f32) -> RingFixture { + let frame_count = 2 * radius + 1; let centre_slot = radius; - let blocks_x = w.div_ceil(BLK_STEP); - let blocks_y = h.div_ceil(BLK_STEP); + let blocks_x = width.div_ceil(BLK_STEP); + let blocks_y = height.div_ceil(BLK_STEP); let mv_stride = blocks_x * blocks_y * 2; let conf_stride = blocks_x * blocks_y; - let mut ring = vec![0.0f32; (n_frames * w * h) as usize]; - for (idx, v) in ring.iter_mut().enumerate() { + let mut ring = vec![0.0f32; (frame_count * width * height) as usize]; + for (idx, value) in ring.iter_mut().enumerate() { let mut hash = (idx as u32).wrapping_mul(2654435761).wrapping_add(0x9E3779B9); hash ^= hash >> 15; hash = hash.wrapping_mul(0x85EBCA6B); hash ^= hash >> 13; - *v = hash as f32 / u32::MAX as f32; + *value = hash as f32 / u32::MAX as f32; } let mv_field = vec![0i32; (2 * radius * mv_stride) as usize]; @@ -291,8 +264,9 @@ pub(super) fn noisy_ring(w: u32, h: u32, radius: u32, confidence_value: f32) -> if k == 0 { continue; } - let t = neighbour_idx_for_k(radius, k); - neighbour_slots[t as usize] = (k + radius as i32) as u32; + + let neighbour = neighbour_idx_for_k(radius, k); + neighbour_slots[neighbour as usize] = (k + radius as i32) as u32; } RingFixture { @@ -306,7 +280,35 @@ pub(super) fn noisy_ring(w: u32, h: u32, radius: u32, confidence_value: f32) -> blocks_y, mv_stride, conf_stride, - width: w, - height: h, + width, + height, } } + +/// `count` interleaved frames of a drifting texture with independent noise per frame and channel. +pub(crate) fn noisy_frames(width: u32, height: u32, channels: u32, count: usize) -> Vec> { + let base = textured_base(width + count as u32, height); + let mut frames = Vec::with_capacity(count); + + for index in 0..count { + let mut frame = vec![0.0f32; (width * height * channels) as usize]; + + for channel in 0..channels { + let mut clean = Vec::with_capacity((width * height) as usize); + for row in 0..height { + let start = (row * (width + count as u32) + index as u32) as usize; + clean.extend_from_slice(&base[start..start + width as usize]); + } + + let seed = (index as u32) * 3 + channel; + let noisy = noisy_field_over(&clean, width, height, 0.03, seed); + for pixel in 0..(width * height) as usize { + frame[pixel * channels as usize + channel as usize] = noisy[pixel]; + } + } + + frames.push(frame); + } + + frames +} diff --git a/av-denoise-core/src/nl4d/tests/mod.rs b/av-denoise-core/src/nl4d/tests/mod.rs index 0fd5831..44264b6 100644 --- a/av-denoise-core/src/nl4d/tests/mod.rs +++ b/av-denoise-core/src/nl4d/tests/mod.rs @@ -1,8 +1,10 @@ pub(crate) mod helpers; mod edges; +mod engine; mod grouping; mod noise_map; +mod options; mod pipeline; mod pooled; mod regularise; diff --git a/av-denoise-core/src/nl4d/tests/noise_map.rs b/av-denoise-core/src/nl4d/tests/noise_map.rs index 0ce29b8..e47d89c 100644 --- a/av-denoise-core/src/nl4d/tests/noise_map.rs +++ b/av-denoise-core/src/nl4d/tests/noise_map.rs @@ -1,6 +1,8 @@ -use super::helpers::{R, make_client, unit_noise}; +use super::helpers::{R, make_client}; +use crate::bench_api::HostIo; use crate::nl4d::denoiser::noise_curve_upload; use crate::nl4d::{Nl4dDenoiser, Nl4dParams}; +use crate::nlmeans::tests::helpers::seeded_unit_gaussian; use crate::nlmeans::{ChannelMode, HqParams, NOISE_CURVE_BINS, NlmParams, PrefilterMode}; // Large enough that each luma bin the ramp crosses gathers the blocks a curve needs. At 320x240 @@ -11,8 +13,7 @@ const FRAMES: u32 = 12; const RAMP_TOP: f32 = 0.15; const RAMP_BOTTOM: f32 = 0.85; -/// A clip's denoised frames, and whether the front end built a noise -/// curve during any pass. +/// A clip's denoised frames, and whether the front end built a noise curve during any pass. struct ClipRun { outputs: Vec>, curve_seen: bool, @@ -28,8 +29,8 @@ fn noise_std_at(luma: f32) -> f32 { 0.004 + 0.02 * luma } -/// A static vertical brightness ramp with fresh noise per frame, in -/// `channels` interleaved planes that all carry the same ramp. +/// A static vertical brightness ramp with fresh noise per frame, in `channels` interleaved planes +/// that all carry the same ramp. fn ramp_clip(channels: u32) -> Vec> { let mut frames = Vec::new(); for frame_index in 0..FRAMES { @@ -43,13 +44,15 @@ fn ramp_clip(channels: u32) -> Vec> { let pixel = y * WIDTH + x; let sample = pixel * channels + channel; let seed = frame_index * channels + channel; - let noise = unit_noise(pixel, seed); + let noise = seeded_unit_gaussian(pixel, seed); frame[sample as usize] = (luma + noise * noise_std).clamp(0.0, 1.0); } } } + frames.push(frame); } + frames } @@ -59,13 +62,15 @@ fn ramp_params(channels: ChannelMode, noise_map: bool, sigma_scale: f32) -> Nl4d sigma_scale, ..HqParams::default() }; + let nlm = NlmParams { + channels, + prefilter: PrefilterMode::None, + hq: Some(hq), + ..defaults.nlm + }; + Nl4dParams { - nlm: NlmParams { - channels, - prefilter: PrefilterMode::None, - hq: Some(hq), - ..defaults.nlm - }, + nlm, noise_map, ..defaults } @@ -79,19 +84,17 @@ fn denoise_clip(params: Nl4dParams, frames: &[Vec]) -> ClipRun { let mut curve_seen = false; for frame in frames { denoiser.push_frame(frame); - let pending = denoiser.denoise_submit().expect("denoise_submit failed"); + let output = denoiser.denoise().expect("denoise failed"); curve_seen |= denoiser.front_for_test().current_noise_curve().is_some(); - if let Some(pending) = pending { - let output = pending.wait().expect("readback failed"); - let output_frame = output.into_f32().expect("f32 output"); + if let Some(output_frame) = output { outputs.push(output_frame); } } denoiser .flush(|frame| { - let output_frame = frame.as_f32().expect("f32 denoiser").to_vec(); + let output_frame = frame.to_vec(); outputs.push(output_frame); }) .expect("flush failed"); diff --git a/av-denoise-core/src/nl4d/tests/options.rs b/av-denoise-core/src/nl4d/tests/options.rs new file mode 100644 index 0000000..2b468f7 --- /dev/null +++ b/av-denoise-core/src/nl4d/tests/options.rs @@ -0,0 +1,293 @@ +use crate::error::Error; +use crate::nl4d::{ + Nl4dOptions, + Nl4dParams, + nl4d_default_lambda_ht, + nl4d_pool_ratio, + nl4d_spatial_radius_for, + nl4d_temporal_radius_for, + resolve_params, +}; +use crate::nlmeans::{ + ChannelMode, + MotionCompensationMode, + MotionEstimation, + MotionSearch, + NlmParams, + PrefilterMode, +}; +use crate::options::Preset; + +const ALL_CHANNELS: [ChannelMode; 3] = [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv]; + +fn lambda_ht_for(options: &Nl4dOptions, channels: ChannelMode) -> f32 { + let params = resolve_params(options, channels).expect("the options are valid"); + params.lambda_ht +} + +#[test] +fn nl4d_default_lambda_ht_differs_between_luma_and_chroma() { + let luma = nl4d_default_lambda_ht(ChannelMode::Luma); + let chroma = nl4d_default_lambda_ht(ChannelMode::Chroma); + + assert!((luma - 4.158).abs() < f32::EPSILON); + assert!((chroma - 3.234).abs() < f32::EPSILON); + assert!( + (chroma - luma).abs() > f32::EPSILON, + "the two planes should not resolve to the same default" + ); +} + +#[test] +fn nl4d_default_lambda_ht_yuv_reads_the_luma_value() { + let yuv = nl4d_default_lambda_ht(ChannelMode::Yuv); + let luma = nl4d_default_lambda_ht(ChannelMode::Luma); + + assert!((yuv - luma).abs() < f32::EPSILON); +} + +#[test] +fn nl4d_pool_ratio_gives_the_calibrated_threshold_at_each_default_lambda() { + for channels in ALL_CHANNELS { + let threshold = nl4d_pool_ratio(channels) * nl4d_default_lambda_ht(channels); + assert!( + (threshold - 2.42).abs() < 1.0e-6, + "{channels:?} gives {threshold}" + ); + } + + let yuv_ratio = nl4d_pool_ratio(ChannelMode::Yuv); + let luma_ratio = nl4d_pool_ratio(ChannelMode::Luma); + assert_eq!(yuv_ratio, luma_ratio); +} + +#[test] +fn nl4d_options_default_to_pooling_on() { + let options = Nl4dOptions::default(); + + assert!(options.pooled_threshold); +} + +#[test] +fn nl4d_options_default_temporal_radius_is_the_base_preset_radius() { + let options = Nl4dOptions::default(); + + assert_eq!(options.temporal_radius, nl4d_temporal_radius_for(Preset::Base)); + assert_eq!(options.temporal_radius, 2); +} + +#[test] +fn nl4d_spatial_radius_for_veryfast_is_narrower_than_the_default() { + let veryfast = nl4d_spatial_radius_for(Preset::Veryfast); + let base = nl4d_spatial_radius_for(Preset::Base); + let options = Nl4dOptions::default(); + + assert_eq!(veryfast, 6); + assert_eq!(base, options.spatial_radius); +} + +#[test] +fn nl4d_options_default_matches_nl4d_params_default() { + let options = Nl4dOptions::default(); + let params = Nl4dParams::default(); + let yuv_lambda_ht = nl4d_default_lambda_ht(ChannelMode::Yuv); + + assert_eq!(options.refine, params.refine); + assert_eq!(options.spatial_radius, params.spatial_radius); + assert!((options.c_min - params.c_min).abs() < f32::EPSILON); + assert_eq!(options.lambda_ht, None); + assert!((params.lambda_ht - yuv_lambda_ht).abs() < f32::EPSILON); +} + +#[test] +fn resolve_lambda_ht_unset_uses_the_per_plane_default() { + let options = Nl4dOptions::default(); + + let luma = lambda_ht_for(&options, ChannelMode::Luma); + let chroma = lambda_ht_for(&options, ChannelMode::Chroma); + + assert!((luma - 4.158).abs() < f32::EPSILON, "got {luma}"); + assert!((chroma - 3.234).abs() < f32::EPSILON, "got {chroma}"); +} + +#[test] +fn resolve_lambda_ht_explicit_value_overrides_every_plane() { + let options = Nl4dOptions { + lambda_ht: Some(4.4), + ..Nl4dOptions::default() + }; + + for channels in ALL_CHANNELS { + let got = lambda_ht_for(&options, channels); + assert!( + (got - 4.4).abs() < f32::EPSILON, + "channels {channels:?} got {got}" + ); + } +} + +#[test] +fn resolve_lambda_ht_default_scale_leaves_the_value_alone() { + let options = Nl4dOptions::default(); + + for channels in ALL_CHANNELS { + let got = lambda_ht_for(&options, channels); + let want = nl4d_default_lambda_ht(channels); + assert!( + (got - want).abs() < f32::EPSILON, + "channels {channels:?} got {got}" + ); + } +} + +#[test] +fn resolve_lambda_ht_scale_multiplies_the_per_plane_default() { + let options = Nl4dOptions { + lambda_ht_scale: 1.1, + ..Nl4dOptions::default() + }; + + for channels in ALL_CHANNELS { + let got = lambda_ht_for(&options, channels); + let want = nl4d_default_lambda_ht(channels) * 1.1; + assert!( + (got - want).abs() < 1e-5, + "channels {channels:?} got {got}, want {want}" + ); + } +} + +#[test] +fn resolve_lambda_ht_scale_multiplies_an_explicit_value() { + let options = Nl4dOptions { + lambda_ht: Some(4.0), + lambda_ht_scale: 1.5, + ..Nl4dOptions::default() + }; + + for channels in ALL_CHANNELS { + let got = lambda_ht_for(&options, channels); + assert!((got - 6.0).abs() < 1e-5, "channels {channels:?} got {got}"); + } +} + +#[test] +fn resolve_lambda_ht_rejects_an_out_of_range_scale() { + for bad in [0.0, -1.0, 0.05, 10.5, f32::NAN, f32::INFINITY] { + let options = Nl4dOptions { + lambda_ht_scale: bad, + ..Nl4dOptions::default() + }; + let result = resolve_params(&options, ChannelMode::Luma); + + let Err(Error::InvalidOptions(message)) = result else { + panic!("lambda_ht_scale={bad} should be rejected"); + }; + assert!( + message.contains("lambda_ht_scale"), + "lambda_ht_scale={bad} gave {message}" + ); + } +} + +#[test] +fn nl4d_builds_the_front_ends_hq_params_from_its_own_fields() { + let options = Nl4dOptions { + sigma: Some(0.02), + sigma_scale: 1.3, + thsad_scale: 0.8, + ..Nl4dOptions::default() + }; + let params = resolve_params(&options, ChannelMode::Yuv).expect("resolve"); + + let hq = params.nlm.hq.expect("nl4d always runs the hq front end"); + assert_eq!(hq.sigma_override, Some(0.02)); + assert!((hq.sigma_scale - 1.3).abs() < f32::EPSILON); + assert!((hq.thsad_scale - 0.8).abs() < f32::EPSILON); + assert!( + hq.temporal_confidence, + "the grouping kernel reads the confidence scores, so this cannot be off" + ); +} + +#[test] +fn nl4d_never_builds_a_prefilter() { + let options = Nl4dOptions::default(); + let params = resolve_params(&options, ChannelMode::Yuv).expect("resolve"); + + assert!(matches!(params.nlm.prefilter, PrefilterMode::None)); +} + +#[test] +fn nl4d_leaves_the_nlm_weighting_knobs_at_their_defaults() { + let defaults = NlmParams::default(); + let options = Nl4dOptions { + temporal_radius: 4, + ..Nl4dOptions::default() + }; + let params = resolve_params(&options, ChannelMode::Luma).expect("resolve"); + + assert!((params.nlm.strength - defaults.strength).abs() < f32::EPSILON); + assert_eq!(params.nlm.search_radius, defaults.search_radius); + assert_eq!(params.nlm.patch_radius, defaults.patch_radius); + assert!((params.nlm.self_weight - defaults.self_weight).abs() < f32::EPSILON); +} + +#[test] +fn nl4d_motion_search_becomes_an_active_mvtools_mode() { + let motion = MotionSearch { + blksize: 32, + overlap: 16, + search_radius: 6, + pyramid_levels: 1, + estimation: MotionEstimation::Direct, + }; + let options = Nl4dOptions { + motion, + ..Nl4dOptions::default() + }; + let params = resolve_params(&options, ChannelMode::Yuv).expect("resolve"); + + assert!(matches!( + params.nlm.motion_compensation, + MotionCompensationMode::Mvtools { + blksize: 32, + overlap: 16, + search_radius: 6, + pyramid_levels: 1, + estimation: MotionEstimation::Direct, + } + )); +} + +#[test] +fn nl4d_motion_search_defaults_match_the_front_ends_own_defaults() { + let options = Nl4dOptions::default(); + let params = resolve_params(&options, ChannelMode::Yuv).expect("resolve"); + let defaults = Nl4dParams::default(); + + assert_eq!(params.nlm.motion_compensation, defaults.nlm.motion_compensation); +} + +#[test] +fn a_zero_temporal_radius_is_rejected() { + let options = Nl4dOptions { + temporal_radius: 0, + ..Nl4dOptions::default() + }; + let result = resolve_params(&options, ChannelMode::Luma); + + assert!(matches!(result, Err(Error::InvalidOptions(_)))); +} + +#[test] +fn the_temporal_radius_reaches_both_the_front_end_and_the_grouping_stage() { + let options = Nl4dOptions { + temporal_radius: 3, + ..Nl4dOptions::default() + }; + let params = resolve_params(&options, ChannelMode::Luma).expect("resolve"); + + assert_eq!(params.temporal_radius, 3); + assert_eq!(params.nlm.temporal_radius, 3); +} diff --git a/av-denoise-core/src/nl4d/tests/pipeline.rs b/av-denoise-core/src/nl4d/tests/pipeline.rs index 54948f1..60eda85 100644 --- a/av-denoise-core/src/nl4d/tests/pipeline.rs +++ b/av-denoise-core/src/nl4d/tests/pipeline.rs @@ -8,11 +8,11 @@ use super::helpers::{ SIGMA, SPATIAL_RADIUS, make_client, - noisy_copy_of, psnr, static_clip_params, textured_base, }; +use crate::bench_api::HostIo; use crate::collab::geometry::{fused_cubes_x, ref_count, refs_along, strength_map_dims}; use crate::collab::kernels::aggregate::{ ACCUM_SCALE, @@ -24,44 +24,50 @@ use crate::collab::kernels::aggregate::{ }; use crate::collab::kernels::fused::{STRENGTH_MAP_OFF, collab_fused}; use crate::collab::kernels::transforms::dct_noise_profile; -use crate::collab::{MAX_K, PATCH_SIZE, grid_frames, needs_warp_uniform_search}; +use crate::collab::{MAX_K, MAX_TEMPORAL_RADIUS, PATCH_SIZE, grid_frames, needs_warp_uniform_search}; use crate::nl4d::{Nl4dDenoiser, Nl4dParams}; -use crate::nlmeans::{ChannelMode, NOISE_CURVE_BINS, NlmDenoiser, NlmParams}; +use crate::nlmeans::tests::helpers::noisy_field_over; +use crate::nlmeans::{BLOCK_X, BLOCK_Y, ChannelMode, NOISE_CURVE_BINS, NlmDenoiser, NlmParams}; -/// A static clip, camera and content both still, with independent -/// per-frame noise. Every emitted frame must come out well above the -/// noisy input's own PSNR against the clean base. #[test] fn denoises_a_static_noisy_clip() { let client = make_client(); - let (w, h) = (64u32, 64u32); + let (width, height) = (64u32, 64u32); let radius = 2u32; - let base = textured_base(w, h); - let n = 9usize; + let base = textured_base(width, height); + let frame_count = 9usize; - let noisy_frames: Vec> = (0..n as u32) - .map(|seed| noisy_copy_of(&base, w, h, SIGMA, seed)) + let noisy_frames: Vec> = (0..frame_count as u32) + .map(|seed| noisy_field_over(&base, width, height, SIGMA, seed)) .collect(); let params = static_clip_params(radius); - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); let mut outputs: Vec> = Vec::new(); for frame in &noisy_frames { - d.push_frame(frame); - if let Some(pending) = d.denoise_submit().expect("denoise_submit failed") { - let frame = pending.wait().expect("readback failed"); - outputs.push(frame.into_f32().expect("f32 output")); + denoiser.push_frame(frame); + if let Some(frame) = denoiser.denoise().expect("denoise failed") { + outputs.push(frame); } } - d.flush(|frame| outputs.push(frame.as_f32().expect("f32 denoiser").to_vec())) + + denoiser + .flush(|frame| { + let values = frame.to_vec(); + outputs.push(values); + }) .expect("flush failed"); - assert_eq!(outputs.len(), n, "expected one emitted frame per pushed frame"); + assert_eq!( + outputs.len(), + frame_count, + "expected one emitted frame per pushed frame" + ); - for (i, out) in outputs.iter().enumerate() { + for (i, output) in outputs.iter().enumerate() { let noisy_psnr = psnr(&noisy_frames[i], &base); - let out_psnr = psnr(out, &base); + let out_psnr = psnr(output, &base); assert!( out_psnr > noisy_psnr + 6.0, "frame {i}: expected at least a 6 dB PSNR improvement over the noisy input, got \ @@ -70,59 +76,59 @@ fn denoises_a_static_noisy_clip() { } } -/// `spatial_radius = 16` with `temporal_radius = MAX_TEMPORAL_RADIUS` (8) -/// is the widest configuration the parameter ranges allow, and so the one -/// whose cross-frame accumulator comes closest to overflowing `i32`. -/// -/// `cross_frame_accum_scale` sizes the fixed-point scale for it, so this -/// combination should denoise as cleanly as any other rather than -/// producing the non-finite or wildly out-of-range values an overflow -/// leaves behind. +/// The widest radii the parameter ranges allow come closest to overflowing the `i32` cross-frame +/// accumulator, so they must still denoise cleanly. #[test] fn denoises_at_the_widest_spatial_and_temporal_radius() { let client = make_client(); - let (w, h) = (64u32, 64u32); - let radius = crate::collab::MAX_TEMPORAL_RADIUS; - let base = textured_base(w, h); - let n = 3usize; + let (width, height) = (64u32, 64u32); + let radius = MAX_TEMPORAL_RADIUS; + let base = textured_base(width, height); + let frame_count = 3usize; - let noisy_frames: Vec> = (0..n as u32) - .map(|seed| noisy_copy_of(&base, w, h, SIGMA, seed)) + let noisy_frames: Vec> = (0..frame_count as u32) + .map(|seed| noisy_field_over(&base, width, height, SIGMA, seed)) .collect(); + let clip_params = static_clip_params(radius); let params = Nl4dParams { spatial_radius: 16, - ..static_clip_params(radius) + ..clip_params }; - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); let mut outputs: Vec> = Vec::new(); for frame in &noisy_frames { - d.push_frame(frame); - if let Some(pending) = d.denoise_submit().expect("denoise_submit failed") { - let frame = pending.wait().expect("readback failed"); - outputs.push(frame.into_f32().expect("f32 output")); + denoiser.push_frame(frame); + if let Some(frame) = denoiser.denoise().expect("denoise failed") { + outputs.push(frame); } } - d.flush(|frame| outputs.push(frame.as_f32().expect("f32 denoiser").to_vec())) + + denoiser + .flush(|frame| { + let values = frame.to_vec(); + outputs.push(values); + }) .expect("flush failed"); - assert_eq!(outputs.len(), n, "expected one emitted frame per pushed frame"); + assert_eq!( + outputs.len(), + frame_count, + "expected one emitted frame per pushed frame" + ); - for (i, out) in outputs.iter().enumerate() { - // An `i32` overflow wraps the fixed-point accumulator into a huge - // or negative value, which `collab_normalise` then divides - // through into non-finite or wildly out-of-range output. Checking - // finiteness first gives a clearer failure than letting a bad - // value fall through into the PSNR comparison below. + for (i, output) in outputs.iter().enumerate() { + // An overflow wraps the accumulator into non-finite output, so checking finiteness first + // gives a clearer failure than the PSNR comparison. assert!( - out.iter().all(|v| v.is_finite()), + output.iter().all(|value| value.is_finite()), "frame {i}: output contains non-finite values, a symptom of the accumulator \ overflow this test guards against" ); let noisy_psnr = psnr(&noisy_frames[i], &base); - let out_psnr = psnr(out, &base); + let out_psnr = psnr(output, &base); assert!( out_psnr > noisy_psnr, "frame {i}: expected a PSNR improvement over the noisy input at spatial_radius=16, \ @@ -131,77 +137,65 @@ fn denoises_at_the_widest_spatial_and_temporal_radius() { } } -/// Guards `run_pass`'s whole-ring accumulator zero against the GPU's -/// per-dimension dispatch limit. -/// -/// A single 1D dispatch has to stay at or under 65,535 workgroups on -/// every backend this project targets. Zeroing the whole cross-frame ring -/// in one dispatch would need `accum_ring_len` slots -/// (`width * height * stored_ch * (1 + 2 * temporal_radius)`) at 256 -/// threads per workgroup, which a large enough resolution and -/// `temporal_radius` pushes over that limit. -/// -/// A GPU that rejects an oversized dispatch leaves the ring holding -/// `client.empty`'s undefined memory instead of zero. Every later pass -/// then scatters real contributions on top of that, and -/// `collab_normalise` divides it through into wildly wrong output. -/// -/// `1024 * 1024` at `temporal_radius = MAX_TEMPORAL_RADIUS` (a -/// `1 + 2 * 8 = 17`-frame ring) needs `1024 * 1024 * 17 = 17,825,792` -/// accumulator elements, `69,632` workgroups at 256 threads each, -/// comfortably over the limit. This is also the exact failure mode a -/// real run hit at `temporal_radius = 4` on a 1080p input, `72,900` -/// workgroups for the luma plane alone. +/// Guards the whole-ring accumulator zero against the 65,535 workgroup limit on a single 1D +/// dispatch. /// -/// The frame count pushes the ring past a second full cycle -/// (`2 * total_frames + 3`), so this also stands as a regression guard -/// for slot reuse across more than one lap of the ring, not only the -/// pass-0 dispatch itself. +/// A rejected dispatch leaves the ring holding `client.empty` memory instead of zero, which +/// normalises into wildly wrong output. `1024 * 1024` at a 17-frame ring needs 69,632 workgroups +/// at 256 threads each. 1080p luma at `temporal_radius = 4` needs 72,900 workgroups. The frame +/// count runs the ring past a second full lap, so this also guards slot reuse. #[test] fn survives_a_ring_size_that_would_overflow_a_single_zero_dispatch() { let client = make_client(); - let (w, h) = (1024u32, 1024u32); - let radius = crate::collab::MAX_TEMPORAL_RADIUS; - let base = textured_base(w, h); + let (width, height) = (1024u32, 1024u32); + let radius = MAX_TEMPORAL_RADIUS; + let base = textured_base(width, height); let total_frames = 1 + 2 * radius; - let n = (2 * total_frames + 3) as usize; + let frame_count = (2 * total_frames + 3) as usize; - let noisy_frames: Vec> = (0..n as u32) - .map(|seed| noisy_copy_of(&base, w, h, SIGMA, seed)) + let noisy_frames: Vec> = (0..frame_count as u32) + .map(|seed| noisy_field_over(&base, width, height, SIGMA, seed)) .collect(); + // A narrow search keeps the test's time on the ring size under test. + let clip_params = static_clip_params(radius); let params = Nl4dParams { - // Cheaper than the module defaults so the test spends its time - // on the ring size under test, not on a wide spatial search - // this bug has nothing to do with. spatial_radius: 2, refine: 1, - ..static_clip_params(radius) + ..clip_params }; - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); let mut outputs: Vec> = Vec::new(); for frame in &noisy_frames { - d.push_frame(frame); - if let Some(pending) = d.denoise_submit().expect("denoise_submit failed") { - let frame = pending.wait().expect("readback failed"); - outputs.push(frame.into_f32().expect("f32 output")); + denoiser.push_frame(frame); + if let Some(frame) = denoiser.denoise().expect("denoise failed") { + outputs.push(frame); } } - d.flush(|frame| outputs.push(frame.as_f32().expect("f32 denoiser").to_vec())) + + denoiser + .flush(|frame| { + let values = frame.to_vec(); + outputs.push(values); + }) .expect("flush failed"); - assert_eq!(outputs.len(), n, "expected one emitted frame per pushed frame"); + assert_eq!( + outputs.len(), + frame_count, + "expected one emitted frame per pushed frame" + ); - for (i, out) in outputs.iter().enumerate() { + for (i, output) in outputs.iter().enumerate() { assert!( - out.iter().all(|v| v.is_finite()), + output.iter().all(|value| value.is_finite()), "frame {i}: output contains non-finite values, a symptom of the pass-0 dispatch \ this test guards against leaving the accumulator ring unzeroed" ); let noisy_psnr = psnr(&noisy_frames[i], &base); - let out_psnr = psnr(out, &base); + let out_psnr = psnr(output, &base); assert!( out_psnr > noisy_psnr, "frame {i}: expected a PSNR improvement over the noisy input, got noisy={noisy_psnr:.4} dB \ @@ -211,92 +205,62 @@ fn survives_a_ring_size_that_would_overflow_a_single_zero_dispatch() { } } -/// Guards the `centre_slot` contract `run_pass` depends on. The -/// slot the pass is centred on is what the reference patch is read from -/// and what an untouched member scatters back into, and nothing in the -/// type system pins it to the frame the caller means, so this test -/// plants content only one real frame carries and checks it survives in -/// that frame's own emitted output. -/// -/// `denoises_a_static_noisy_clip` is a weak canary for this specific -/// mismatch: every ring slot there holds the same base content, so -/// feeding the filter a valid but wrong slot would barely move its PSNR -/// (a neighbour frame denoises to essentially the same clean content). -/// Here one frame alone carries a strong, low-frequency marker block no -/// other frame has. If the pass ever centred on a different physical -/// ring slot, none of the marker's own patches would be a reference at -/// all, and the marker would be attenuated or absent from its frame's -/// own completed output. +/// Each pass must centre on the ring slot of the frame the caller means. /// -/// Emission lags `temporal_radius` passes behind the pass a frame is -/// the centre of (see [`Nl4dDenoiser::denoise_submit`]), so this -/// pushes `3 * radius + 1` frames, interleaving a `denoise_submit` after -/// every push the way a real caller does, and collects every emitted -/// output in order. Emitted output `k` is always real frame `k`'s own -/// completed region (see that same doc comment for why), so the marker, -/// planted on real frame `radius`, is checked against emitted output -/// `radius`. -/// -/// The marker is a big flat block, not fine detail, so ordinary -/// shrinkage cannot legitimately remove it, and the assertion checks a -/// wide margin rather than an exact value, so ordinary filtering noise -/// does not trip it. +/// Only one frame carries a large flat marker block. A pass centred on a different slot would +/// have none of the marker's patches as references, so the marker would fade from its frame's +/// output. A flat block cannot be removed by ordinary shrinkage, and the check uses a wide margin +/// so filtering noise does not trip it. #[test] fn output_carries_its_own_frames_marker_no_other_frame_has() { let client = make_client(); - let (w, h) = (64u32, 64u32); + let (width, height) = (64u32, 64u32); let radius = 2u32; - let base = textured_base(w, h); + let base = textured_base(width, height); - // `textured_base` never exceeds 0.65 (0.5 centre +/- 0.15 amplitude), - // so a marker at 0.92 is unambiguous against it even before adding - // noise or considering denoising error. + // `textured_base` never exceeds 0.65, so a marker at 0.92 is unambiguous against it. const MARKER: f32 = 0.92; const MARKER_X0: u32 = 24; const MARKER_Y0: u32 = 24; const MARKER_SIZE: u32 = 24; - // Read only the block's interior, away from its own edges, so patch - // boundary blending against the surrounding non-marker texture can't - // explain a low reading. + // Read only the block's interior, so blending at its edges cannot explain a low reading. const INTERIOR_MARGIN: u32 = 8; let mut marker_clean = base.clone(); for y in MARKER_Y0..MARKER_Y0 + MARKER_SIZE { for x in MARKER_X0..MARKER_X0 + MARKER_SIZE { - marker_clean[(y * w + x) as usize] = MARKER; + marker_clean[(y * width + x) as usize] = MARKER; } } - // The pass centred on real frame `radius` only finishes contributing - // to real frame `radius`'s own output `radius` passes later, and the - // pass centred on frame `f` runs once `f + radius + 1` real frames - // have been pushed, so `3 * radius + 1` real pushes are exactly - // enough to reach that emission without needing `flush` too. + // The pass centred on frame `f` runs once `f + radius + 1` frames are pushed, and frame + // `radius` emits `radius` passes after its own, so `3 * radius + 1` pushes reach it without + // `flush`. let marker_frame = radius; - let n_frames = 3 * radius + 1; - let frames: Vec> = (0..n_frames) + let frame_count = 3 * radius + 1; + let frames: Vec> = (0..frame_count) .map(|seed| { let content = if seed == marker_frame { &marker_clean } else { &base }; - noisy_copy_of(content, w, h, SIGMA, seed) + noisy_field_over(content, width, height, SIGMA, seed) }) .collect(); let params = static_clip_params(radius); - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); let mut outputs: Vec> = Vec::new(); for frame in &frames { - d.push_frame(frame); - if let Some(pending) = d.denoise_submit().expect("denoise_submit failed") { - let frame = pending.wait().expect("readback failed"); - outputs.push(frame.into_f32().expect("f32 output")); + denoiser.push_frame(frame); + if let Some(frame) = denoiser.denoise().expect("denoise failed") { + outputs.push(frame); } } - let out = outputs + // Outputs emit in frame order, so output `k` is real frame `k`'s own completed region. + let output = outputs .get(marker_frame as usize) .expect("enough frames were pushed for frame `radius`'s own output to have emitted"); @@ -304,10 +268,11 @@ fn output_carries_its_own_frames_marker_no_other_frame_has() { let mut count = 0usize; for y in (MARKER_Y0 + INTERIOR_MARGIN)..(MARKER_Y0 + MARKER_SIZE - INTERIOR_MARGIN) { for x in (MARKER_X0 + INTERIOR_MARGIN)..(MARKER_X0 + MARKER_SIZE - INTERIOR_MARGIN) { - sum += out[(y * w + x) as usize] as f64; + sum += output[(y * width + x) as usize] as f64; count += 1; } } + let mean = sum / count as f64; eprintln!("output_carries_its_own_frames_marker_no_other_frame_has: marker interior mean = {mean:.4}"); @@ -320,39 +285,39 @@ fn output_carries_its_own_frames_marker_no_other_frame_has() { ); } -/// `flush` must emit exactly as many frames as were pushed, whatever mix -/// of `denoise_submit` and `flush` produced them. #[test] fn flush_emits_exactly_the_pushed_frame_count() { let client = make_client(); - let (w, h) = (64u32, 64u32); + let (width, height) = (64u32, 64u32); let radius = 2u32; - let base = textured_base(w, h); - let n = 7u32; + let base = textured_base(width, height); + let frame_count = 7u32; let params = static_clip_params(radius); - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); let mut emitted = 0usize; - for seed in 0..n { - let frame = noisy_copy_of(&base, w, h, SIGMA, seed); - d.push_frame(&frame); - if d.denoise_submit().expect("denoise_submit failed").is_some() { + for seed in 0..frame_count { + let frame = noisy_field_over(&base, width, height, SIGMA, seed); + denoiser.push_frame(&frame); + if denoiser.denoise().expect("denoise failed").is_some() { emitted += 1; } } - d.flush(|_| emitted += 1).expect("flush failed"); - assert_eq!(emitted, n as usize, "expected exactly {n} emitted frames"); + denoiser.flush(|_| emitted += 1).expect("flush failed"); + + assert_eq!( + emitted, frame_count as usize, + "expected exactly {frame_count} emitted frames" + ); } -/// Launches the same collaborative and aggregation kernels -/// [`Nl4dDenoiser::run_pass`] runs, standalone, for a -/// single-frame ring at `radius = 0`. This is the "spatial-only" arm of -/// the hypothesis test below: identical grouping (no admission gate), -/// identical filter (hard threshold, same `lambda_ht`), identical noise -/// floor and `c_min`, with the only difference being that no temporal -/// candidates exist to search. +/// Runs the denoiser's collaborative and aggregation kernels standalone on a single-frame ring at +/// `radius = 0`. +/// +/// Grouping, filter, noise floor and `c_min` match the denoiser, so the only difference is that +/// no temporal candidates exist to search. #[expect( clippy::too_many_arguments, reason = "the test helper takes the full set of parameters its cases vary" @@ -360,8 +325,8 @@ fn flush_emits_exactly_the_pushed_frame_count() { fn run_spatial_only( client: &ComputeClient, noisy_centre: &[f32], - w: u32, - h: u32, + width: u32, + height: u32, spatial_radius: u32, refine: u32, c_min: f32, @@ -372,47 +337,62 @@ fn run_spatial_only( let k_max = MAX_K; let stored_ch = 1u32; let channels_count = 1u32; - let refs_x = refs_along(w); - let refs_y = refs_along(h); - let refs = ref_count(w, h); - let pixels = (w * h) as usize; + let refs_x = refs_along(width); + let refs_y = refs_along(height); + let refs = ref_count(width, height); + let pixels = (width * height) as usize; let frame_len = pixels; - let ring_buf = client.create_from_slice(f32::as_bytes(noisy_centre)); + let centre_bytes = f32::as_bytes(noisy_centre); + let ring_buf = client.create_from_slice(centre_bytes); let mv_dummy = client.empty(size_of::()); let conf_dummy = client.empty(size_of::()); let neighbour_slots_dummy = client.empty(size_of::()); let group_weight = client.empty(refs * size_of::()); - let sigma_buf = client.create_from_slice(f32::as_bytes(&[sigma])); + + let sigma_values = [sigma]; let dct_profile = dct_noise_profile(0.0); - let dct_profile_buf = client.create_from_slice(f32::as_bytes(&dct_profile)); - let kaiser_buf = client.create_from_slice(f32::as_bytes(&kaiser_window(0.0))); - let zero_curve = client.create_from_slice(f32::as_bytes(&[0.0f32; NOISE_CURVE_BINS])); - let (map_cols, map_rows) = strength_map_dims(w, h); + let kaiser = kaiser_window(0.0); + let zeroed_curve = [0.0f32; NOISE_CURVE_BINS]; + let sigma_bytes = f32::as_bytes(&sigma_values); + let sigma_buf = client.create_from_slice(sigma_bytes); + let dct_profile_bytes = f32::as_bytes(&dct_profile); + let dct_profile_buf = client.create_from_slice(dct_profile_bytes); + let kaiser_bytes = f32::as_bytes(&kaiser); + let kaiser_buf = client.create_from_slice(kaiser_bytes); + let curve_bytes = f32::as_bytes(&zeroed_curve); + let zero_curve = client.create_from_slice(curve_bytes); + + let (map_cols, map_rows) = strength_map_dims(width, height); let map_len = (map_cols * map_rows) as usize; let unit_map = vec![1.0f32; map_len]; - let unit_map_buf = client.create_from_slice(f32::as_bytes(&unit_map)); + let unit_map_bytes = f32::as_bytes(&unit_map); + let unit_map_buf = client.create_from_slice(unit_map_bytes); let accum = client.empty(frame_len * size_of::()); let wsum = client.empty(pixels * size_of::()); let output = client.empty(frame_len * size_of::()); - let agg_grid = CubeCount::new_2d( - w.div_ceil(crate::nlmeans::BLOCK_X), - h.div_ceil(crate::nlmeans::BLOCK_Y), - ); - let agg_dim = CubeDim::new_2d(crate::nlmeans::BLOCK_X, crate::nlmeans::BLOCK_Y); + let agg_cubes_x = width.div_ceil(BLOCK_X); + let agg_cubes_y = height.div_ceil(BLOCK_Y); + let agg_grid = CubeCount::new_2d(agg_cubes_x, agg_cubes_y); + let agg_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); let zero_dim = 256u32; let zero_workgroups = (frame_len as u32).div_ceil(zero_dim); let zero_grid = CubeCount::new_1d(zero_workgroups); + let zero_cube_dim = CubeDim::new_1d(zero_dim); - let wnorm = weight_scale(sigma, &dct_profile); + let fused_x = fused_cubes_x(width); + let fused_grid = CubeCount::new_2d(fused_x, refs_y); + let fused_dim = CubeDim::new_1d(64); + let weight_norm = weight_scale(sigma, &dct_profile); + let grid_frame_count = grid_frames(0); let centre_slot = 0u32; unsafe { collab_zero_accum::launch_unchecked::( client, zero_grid, - CubeDim::new_1d(zero_dim), + zero_cube_dim, ArrayArg::from_raw_parts(accum.clone(), frame_len), ArrayArg::from_raw_parts(wsum.clone(), pixels), 0u32, @@ -423,8 +403,8 @@ fn run_spatial_only( collab_fused::launch_unchecked::( client, - CubeCount::new_2d(fused_cubes_x(w), refs_y), - CubeDim::new_1d(64), + fused_grid, + fused_dim, stored_ch as usize, ArrayArg::from_raw_parts(ring_buf, noisy_centre.len()), ArrayArg::from_raw_parts(mv_dummy, 1), @@ -443,11 +423,11 @@ fn run_spatial_only( lambda_ht, 0u32, STRENGTH_MAP_OFF, - wnorm, + weight_norm, ACCUM_SCALE, warp_uniform, 0u32, - grid_frames(0), + grid_frame_count, refine, 1u32, 1u32, @@ -455,8 +435,8 @@ fn run_spatial_only( 8u32, 1u32, 1u32, - w, - h, + width, + height, channels_count, k_max, stored_ch, @@ -477,8 +457,8 @@ fn run_spatial_only( ArrayArg::from_raw_parts(wsum, pixels), ArrayArg::from_raw_parts(output.clone(), frame_len), 0u32, - w, - h, + width, + height, channels_count, stored_ch, ); @@ -488,60 +468,51 @@ fn run_spatial_only( f32::from_bytes(&bytes).to_vec() } -/// The hypothesis under test, as a unit assertion: on a static clip -/// where temporal candidates are near-duplicates, grouping across the -/// temporal window should cancel more grain than a spatial-only search -/// on the same frame ever could. -/// -/// Both arms run the exact same kernel, at the same `spatial_radius`, -/// `c_min`, `lambda_ht`, and fixed `sigma`, over the identical noisy -/// centre frame. The only difference is whether the search has a -/// temporal window to look in. +/// On a static clip, grouping across the temporal window cancels more grain than a spatial-only +/// search on the same frame. /// -/// `denoise_submit` is called after every push, exactly the way a real -/// caller drives this denoiser, and every emitted output is collected in -/// order. Emitted output `k` is real frame `k`'s own completed region -/// (see [`Nl4dDenoiser::denoise_submit`]'s scheduling), which is only -/// ready `radius` passes after the pass centred on frame `k` itself, so -/// `3 * radius + 1` frames are pushed. +/// Both arms run the same kernel with the same `spatial_radius`, `c_min`, `lambda_ht` and fixed +/// `sigma` over the same noisy centre frame. Frame `radius` emits `radius` passes after its own, +/// so `3 * radius + 1` frames are pushed. #[test] fn temporal_grouping_beats_spatial_only_on_a_static_clip() { let client = make_client(); - let (w, h) = (64u32, 64u32); + let (width, height) = (64u32, 64u32); let radius = 2u32; - let base = textured_base(w, h); + let base = textured_base(width, height); let noisy_frames: Vec> = (0..(3 * radius + 1)) - .map(|seed| noisy_copy_of(&base, w, h, SIGMA, seed)) + .map(|seed| noisy_field_over(&base, width, height, SIGMA, seed)) .collect(); let centre_index = radius as usize; let params = static_clip_params(radius); - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); let mut outputs: Vec> = Vec::new(); for frame in &noisy_frames { - d.push_frame(frame); - if let Some(pending) = d.denoise_submit().expect("denoise_submit failed") { - let frame = pending.wait().expect("readback failed"); - outputs.push(frame.into_f32().expect("f32 output")); + denoiser.push_frame(frame); + if let Some(frame) = denoiser.denoise().expect("denoise failed") { + outputs.push(frame); } } + let temporal_out = outputs .get(centre_index) .expect("enough frames were pushed for frame `radius`'s own output to have emitted") .clone(); + let warp_uniform = needs_warp_uniform_search(&client); let spatial_out = run_spatial_only( &client, &noisy_frames[centre_index], - w, - h, + width, + height, SPATIAL_RADIUS, REFINE, C_MIN, LAMBDA_HT, SIGMA, - needs_warp_uniform_search(&client), + warp_uniform, ); let temporal_psnr = psnr(&temporal_out, &base); @@ -560,66 +531,50 @@ fn temporal_grouping_beats_spatial_only_on_a_static_clip() { ); } -/// Isolates cross-frame aggregation's own contribution, apart from -/// temporal grouping's. -/// -/// Both arms here group across the identical radius-2 temporal window, -/// with the identical `lambda_ht`, and run the identical kernel. They -/// differ only in which filtered members reach the aggregate for real -/// frame `radius`'s own output. -/// -/// The cross-frame arm is `Nl4dDenoiser` itself, which keeps every member -/// that ever matched into that frame, whatever pass found it. +/// Isolates cross-frame aggregation's own contribution from temporal grouping's. /// -/// The centre-only arm runs the one pass centred on that frame and reads -/// back only the centre slot's own region of the accumulator ring. The -/// kernel scatters each member into the region of the frame it was -/// matched in, so that region holds exactly the members whose own frame -/// is the centre, which is what a single-pass centre-only design would -/// have produced. It is built by driving the same front end -/// (`NlmDenoiser::submit_machinery`) directly, so the ring, the motion -/// field, and the noise estimate are all the ones the real denoiser -/// used. +/// Both arms group across the same radius-2 window with the same `lambda_ht` and kernel. The +/// cross-frame arm is the denoiser itself, which keeps every member that matched into the judged +/// frame from any pass. The centre-only arm runs the one pass centred on that frame through the +/// same `NlmDenoiser` front end and reads back only the centre slot's region of the ring. Each +/// member scatters into the region of the frame it was matched in, so that region holds exactly +/// the members whose own frame is the centre. #[test] fn cross_frame_aggregation_beats_centre_only_at_the_same_lambda() { let client = make_client(); - let (w, h) = (64u32, 64u32); + let (width, height) = (64u32, 64u32); let radius = 2u32; - let base = textured_base(w, h); + let base = textured_base(width, height); - let n_frames = 3 * radius + 1; - let frames: Vec> = (0..n_frames) - .map(|seed| noisy_copy_of(&base, w, h, SIGMA, seed)) + let frame_count = 3 * radius + 1; + let frames: Vec> = (0..frame_count) + .map(|seed| noisy_field_over(&base, width, height, SIGMA, seed)) .collect(); let judged_frame = radius as usize; - // Cross-frame arm: the real denoiser, unmodified, driven the same - // way every other test above drives it. let params = static_clip_params(radius); - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); let mut outputs: Vec> = Vec::new(); for frame in &frames { - d.push_frame(frame); - if let Some(pending) = d.denoise_submit().expect("denoise_submit failed") { - let frame = pending.wait().expect("readback failed"); - outputs.push(frame.into_f32().expect("f32 output")); + denoiser.push_frame(frame); + if let Some(frame) = denoiser.denoise().expect("denoise failed") { + outputs.push(frame); } } + let cross_frame_out = outputs .get(judged_frame) .expect("enough frames were pushed for frame `radius`'s own output to have emitted") .clone(); - // Centre-only arm: the same front end, driven directly rather than - // through `Nl4dDenoiser`, so the pass centred on real frame `radius` - // can be aggregated on its own once it runs. - let mut nlm_params = static_clip_params(radius).nlm; + let clip_params = static_clip_params(radius); + let mut nlm_params = clip_params.nlm; nlm_params.temporal_radius = radius; - let mut front = NlmDenoiser::::new(&client, nlm_params, w, h); + let mut front = NlmDenoiser::::new(&client, nlm_params, width, height); - let pixels = (w * h) as usize; - let refs_x = refs_along(w); - let refs = ref_count(w, h); + let pixels = (width * height) as usize; + let refs_x = refs_along(width); + let refs = ref_count(width, height); let k_max = MAX_K; let total_frames = 1 + 2 * radius; @@ -627,9 +582,11 @@ fn cross_frame_aggregation_beats_centre_only_at_the_same_lambda() { let mut pass_index = 0u32; for frame in &frames { front.push_frame(frame); + let Some(view) = front.submit_machinery(radius).expect("submit_machinery failed") else { continue; }; + if pass_index != radius { pass_index += 1; continue; @@ -640,37 +597,61 @@ fn cross_frame_aggregation_beats_centre_only_at_the_same_lambda() { let neighbours = 2 * radius; let mv_len = (neighbours * view.mv_stride) as usize; let conf_len = (neighbours * view.conf_stride) as usize; - let neighbour_slots_buf = client.create_from_slice(u32::as_bytes(&view.neighbour_slots)); + let slots_bytes = u32::as_bytes(&view.neighbour_slots); + let neighbour_slots_buf = client.create_from_slice(slots_bytes); let sigmas = front.current_sigmas_temporal_only(); - let sigma_buf = client.create_from_slice(f32::as_bytes(&[sigmas[0]])); + let sigma = [sigmas[0]]; let profile = dct_noise_profile(0.0); - let profile_buf = client.create_from_slice(f32::as_bytes(&profile)); - let kaiser_buf = client.create_from_slice(f32::as_bytes(&kaiser_window(0.0))); - let zero_curve = client.create_from_slice(f32::as_bytes(&[0.0f32; NOISE_CURVE_BINS])); - let (map_cols, map_rows) = strength_map_dims(w, h); + let kaiser = kaiser_window(0.0); + let zeroed_curve = [0.0f32; NOISE_CURVE_BINS]; + let sigma_bytes = f32::as_bytes(&sigma); + let sigma_buf = client.create_from_slice(sigma_bytes); + let profile_bytes = f32::as_bytes(&profile); + let profile_buf = client.create_from_slice(profile_bytes); + let kaiser_bytes = f32::as_bytes(&kaiser); + let kaiser_buf = client.create_from_slice(kaiser_bytes); + let curve_bytes = f32::as_bytes(&zeroed_curve); + let zero_curve = client.create_from_slice(curve_bytes); + + let (map_cols, map_rows) = strength_map_dims(width, height); let map_len = (map_cols * map_rows) as usize; let unit_map = vec![1.0f32; map_len]; - let unit_map_buf = client.create_from_slice(f32::as_bytes(&unit_map)); - let wnorm = weight_scale(sigmas[0], &profile); + let unit_map_bytes = f32::as_bytes(&unit_map); + let unit_map_buf = client.create_from_slice(unit_map_bytes); + let weight_norm = weight_scale(sigmas[0], &profile); let accum_scale = cross_frame_accum_scale(SPATIAL_RADIUS, radius); let group_weight = client.empty(refs * size_of::()); - // The whole ring, because a member matched in a neighbour frame - // scatters into that frame's own region. Only the centre slot's - // region is read back below, which is exactly what makes this - // the centre-only arm. - let accum = client.create_from_slice(i32::as_bytes(&vec![0i32; pixels * total_frames as usize])); - let wsum = client.create_from_slice(i32::as_bytes(&vec![0i32; pixels * total_frames as usize])); + // The whole ring, because a member matched in a neighbour frame scatters into that + // frame's region. Reading back only the centre slot is what makes this the centre-only + // arm. + let zeroed_accum = vec![0i32; pixels * total_frames as usize]; + let zeroed_wsum = vec![0i32; pixels * total_frames as usize]; + let accum_bytes = i32::as_bytes(&zeroed_accum); + let accum = client.create_from_slice(accum_bytes); + let wsum_bytes = i32::as_bytes(&zeroed_wsum); + let wsum = client.create_from_slice(wsum_bytes); let output = client.empty(pixels * size_of::()); - let mc = front.motion_ctx(); + let motion = front.motion_ctx(); + + let fused_x = fused_cubes_x(width); + let refs_y = refs_along(height); + let fused_grid = CubeCount::new_2d(fused_x, refs_y); + let fused_dim = CubeDim::new_1d(64); + let warp_uniform = needs_warp_uniform_search(&client); + let grid_frame_count = grid_frames(radius); + let agg_cubes_x = width.div_ceil(BLOCK_X); + let agg_cubes_y = height.div_ceil(BLOCK_Y); + let agg_grid = CubeCount::new_2d(agg_cubes_x, agg_cubes_y); + let agg_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y); unsafe { collab_fused::launch_unchecked::( &client, - CubeCount::new_2d(fused_cubes_x(w), refs_along(h)), - CubeDim::new_1d(64), + fused_grid, + fused_dim, 1usize, ArrayArg::from_raw_parts(view.input.clone(), ring_len), ArrayArg::from_raw_parts(view.mv_field.clone(), mv_len.max(1)), @@ -689,20 +670,20 @@ fn cross_frame_aggregation_beats_centre_only_at_the_same_lambda() { LAMBDA_HT, 0u32, STRENGTH_MAP_OFF, - wnorm, + weight_norm, accum_scale, - needs_warp_uniform_search(&client), + warp_uniform, radius, - grid_frames(radius), + grid_frame_count, REFINE, view.mv_stride, view.conf_stride, - mc.step, - mc.blksize, - mc.blocks_x, - mc.blocks_y, - w, - h, + motion.step, + motion.blksize, + motion.blocks_x, + motion.blocks_y, + width, + height, 1u32, k_max, 1u32, @@ -716,29 +697,28 @@ fn cross_frame_aggregation_beats_centre_only_at_the_same_lambda() { collab_normalise::launch_unchecked::( &client, - CubeCount::new_2d( - w.div_ceil(crate::nlmeans::BLOCK_X), - h.div_ceil(crate::nlmeans::BLOCK_Y), - ), - CubeDim::new_2d(crate::nlmeans::BLOCK_X, crate::nlmeans::BLOCK_Y), + agg_grid, + agg_dim, 1usize, ArrayArg::from_raw_parts(accum, pixels * total_frames as usize), ArrayArg::from_raw_parts(wsum, pixels * total_frames as usize), ArrayArg::from_raw_parts(output.clone(), pixels), centre_slot * pixels as u32, - w, - h, + width, + height, 1u32, 1u32, ); } - let out = f32::from_bytes(&client.read_one(output).expect("readback failed"))[..pixels].to_vec(); + let output_bytes = client.read_one(output).expect("readback failed"); + let centre_only = f32::from_bytes(&output_bytes)[..pixels].to_vec(); assert!( - out.iter().all(|v| v.is_finite()), + centre_only.iter().all(|value| value.is_finite()), "the centre-only arm left a pixel with no contribution at all" ); - centre_only_out = Some(out); + + centre_only_out = Some(centre_only); break; } @@ -760,189 +740,177 @@ fn cross_frame_aggregation_beats_centre_only_at_the_same_lambda() { ); } -/// The snapshot reports the field the fused kernel was given. A clip -/// whose every frame is the previous one shifted right by 2 pixels must -/// report `[2 * k, 0]` toward the neighbour at offset `k`, at an -/// interior block, once the first pass has run. +/// A clip where every frame is the previous one shifted right by 2 pixels must report +/// `[2 * k, 0]` toward the neighbour at offset `k` once the first pass has run. #[test] fn motion_snapshot_reports_the_field_the_pass_used() { let client = make_client(); - let (w, h) = (96u32, 96u32); + let (width, height) = (96u32, 96u32); let radius = 2u32; - let base = textured_base(w, h); + let base = textured_base(width, height); let frames: Vec> = (0..5i32) .map(|k| { - let mut f = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let sx = (x as i32 - 2 * (k - 2)).clamp(0, w as i32 - 1) as u32; - f[(y * w + x) as usize] = base[(y * w + sx) as usize]; + let mut shifted = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let source_x = (x as i32 - 2 * (k - 2)).clamp(0, width as i32 - 1) as u32; + shifted[(y * width + x) as usize] = base[(y * width + source_x) as usize]; } } - f + + shifted }) .collect(); - let mut d = - Nl4dDenoiser::::new(&client, static_clip_params(radius), w, h).expect("construction failed"); - assert!(d.motion_snapshot().is_none(), "no pass has run yet"); + let params = static_clip_params(radius); + let mut denoiser = Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); + assert!(denoiser.motion_snapshot().is_none(), "no pass has run yet"); + for frame in &frames { - d.push_frame(frame); - let _ = d.denoise_submit().expect("denoise_submit failed"); + denoiser.push_frame(frame); + let _ = denoiser.denoise().expect("denoise failed"); } - let snap = d.motion_snapshot().expect("a pass has run"); - assert_eq!(snap.offsets, vec![-2, -1, 1, 2]); - assert_eq!(snap.step, 8); - assert_eq!(snap.blksize, 16); + let snapshot = denoiser.motion_snapshot().expect("a pass has run"); + assert_eq!(snapshot.offsets, vec![-2, -1, 1, 2]); + assert_eq!(snapshot.step, 8); + assert_eq!(snapshot.blksize, 16); + // Block (3, 3) covers pixels 24..40, well inside the frame. - let block = (3 * snap.blocks_x + 3) as usize; - for (t, &k) in snap.offsets.iter().enumerate() { + let block = (3 * snapshot.blocks_x + 3) as usize; + for (neighbour, &k) in snapshot.offsets.iter().enumerate() { assert_eq!( - snap.vectors[t][block], + snapshot.vectors[neighbour][block], [2 * k, 0], "neighbour k={k} should be tracked as a 2*k pixel shift" ); assert!( - snap.confidence[t][block] > 0.5, + snapshot.confidence[neighbour][block] > 0.5, "a clean shift must score confidently" ); } } -/// A panning clip with one flat block. The estimator ties on the flat -/// block and leaves it at the seed, so its vector differs from its -/// neighbours'. With `field_lambda` on, the pass pulls it to the -/// neighbourhood's vector, and the snapshot shows the regularised -/// field. With it off the field is the estimator's. +/// A panning clip with one flat block, which the estimator ties on and leaves at the seed. +/// +/// With `field_lambda` on, the snapshot shows the flat block pulled to its neighbours' vector. +/// With it off, the snapshot shows the estimator's field. #[test] fn field_regularisation_reaches_the_snapshot() { let client = make_client(); - let (w, h) = (128u32, 96u32); + let (width, height) = (128u32, 96u32); let radius = 1u32; - let mut base = textured_base(w, h); - // Flatten a 24x24 region centred on block (7, 5)'s own footprint, - // 52..76 x 36..60. Large enough to tie both the fine-level search and - // the coarse pyramid level for that one block, but small enough that - // its overlapping neighbours still see enough texture past the - // region's edge to estimate correctly, so only the centre block ties. + let mut base = textured_base(width, height); + + // Flatten 52..76 x 36..60 around block (7, 5). It is large enough to tie both the fine search + // and the coarse pyramid level for that block, but small enough that its overlapping + // neighbours still see texture, so only the centre block ties. for y in 36..60u32 { for x in 52..76u32 { - base[(y * w + x) as usize] = 0.5; + base[(y * width + x) as usize] = 0.5; } } + let frames: Vec> = (0..3i32) .map(|k| { - let mut f = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let sx = (x as i32 - 3 * (k - 1)).clamp(0, w as i32 - 1) as u32; - f[(y * w + x) as usize] = base[(y * w + sx) as usize]; + let mut shifted = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let source_x = (x as i32 - 3 * (k - 1)).clamp(0, width as i32 - 1) as u32; + shifted[(y * width + x) as usize] = base[(y * width + source_x) as usize]; } } - f + + shifted }) .collect(); let run = |lambda: f32| { + let clip_params = static_clip_params(radius); let params = Nl4dParams { field_lambda: lambda, - ..static_clip_params(radius) + ..clip_params }; - let mut d = Nl4dDenoiser::::new(&client, params, w, h).expect("construction failed"); + let mut denoiser = + Nl4dDenoiser::::new(&client, params, width, height).expect("construction failed"); for frame in &frames { - d.push_frame(frame); - let _ = d.denoise_submit().expect("denoise_submit failed"); + denoiser.push_frame(frame); + let _ = denoiser.denoise().expect("denoise failed"); } - d.motion_snapshot().expect("a pass ran") + + denoiser.motion_snapshot().expect("a pass ran") }; let off = run(0.0); let on = run(1.0); - // The block at (7, 5) spans pixels 56..72 x 40..56, inside the flat - // region on every frame. + + // The block at (7, 5) spans pixels 56..72 x 40..56, inside the flat region on every frame. let flat_block = (5 * off.blocks_x + 7) as usize; - let t_plus = 1usize; + let forward_neighbour = 1usize; assert_ne!( - off.vectors[t_plus][flat_block], + off.vectors[forward_neighbour][flat_block], [3, 0], "the flat block must not be tracked without help, or this test proves nothing" ); - assert_eq!(on.vectors[t_plus][flat_block], [3, 0]); + assert_eq!(on.vectors[forward_neighbour][flat_block], [3, 0]); + // A textured block is unchanged by the pass. let textured_block = (2 * off.blocks_x + 2) as usize; - assert_eq!(off.vectors[t_plus][textured_block], [3, 0]); - assert_eq!(on.vectors[t_plus][textured_block], [3, 0]); + assert_eq!(off.vectors[forward_neighbour][textured_block], [3, 0]); + assert_eq!(on.vectors[forward_neighbour][textured_block], [3, 0]); } -/// `Nl4dParams::default()`, changing only `channels`, which has to -/// switch to `Luma` because this file's helpers only ever synthesise a -/// single plane. Every other test in this file pins `field_lambda: -/// 0.0` so its recorded values stay stable, but the shipped default is -/// `1.0`, meaning a real caller always runs the field-regularisation -/// pass this configuration exercises. +/// The shipped default parameters, with `channels` switched to `Luma` because these fixtures only +/// synthesise a single plane. +/// +/// The other tests here pin `field_lambda` to 0.0 so their recorded values stay stable, while the +/// shipped default of 1.0 runs the field-regularisation pass. fn shipped_default_params() -> Nl4dParams { - let params = Nl4dParams { - nlm: NlmParams { - channels: ChannelMode::Luma, - ..Nl4dParams::default().nlm - }, - ..Nl4dParams::default() + let defaults = Nl4dParams::default(); + let nlm = NlmParams { + channels: ChannelMode::Luma, + ..defaults.nlm }; + let params = Nl4dParams { nlm, ..defaults }; assert_eq!( params.field_lambda, 1.0, "this helper exists to exercise the shipped default, not an override" ); + params } -/// Runs the pipeline at the true shipped defaults, in two phases. -/// -/// The first phase pushes a static, clean (noiseless) clip and checks -/// the field-regularisation pass leaves an interior block's vector at -/// exactly zero. A static clip's true motion is zero everywhere, so a -/// correctly regularised field has nothing to pull a well-textured -/// interior block's vector away from zero toward: every neighbouring -/// block's vector is also zero, so their median is zero, which is -/// already where the block sits. A field-regularisation dispatch that -/// reads the wrong pyramid slot, indexes the wrong neighbour, or -/// strides into a different block's data instead of its own would -/// instead pull in whatever mismatched vector sits there, and a real -/// motion vector would show up in a scene where nothing ever moved. -/// This phase carries no noise, because the estimator's own -/// noise-driven wobble would otherwise mask exactly the kind of small, -/// wrong-source displacement it exists to catch. +/// Runs the shipped defaults over a clean static clip, then a noisy one. /// -/// The second phase pushes a static, noisy clip through a fresh -/// denoiser at the same defaults and checks every emitted frame comes -/// out well above the noisy input's own PSNR, the same property -/// [`denoises_a_static_noisy_clip`] checks at `field_lambda: 0.0`. A -/// field-regularisation dispatch bug severe enough to corrupt the -/// motion field would feed the temporal grouping kernel the wrong -/// candidates and show up here as a smaller improvement, or none at -/// all. +/// A static clip's true motion is zero everywhere, so the regularised field must read exactly zero +/// at an interior block. A dispatch that read the wrong pyramid slot, neighbour or stride would +/// pull in a mismatched vector. The clean clip has no noise, because the estimator's noise-driven +/// wobble would mask that small displacement. The noisy clip must then gain at least 6 dB, since a +/// corrupt field would feed grouping the wrong candidates. #[test] fn shipped_defaults_denoise_a_static_clip_and_regularise_its_field_to_zero() { let client = make_client(); - let (w, h) = (96u32, 96u32); - let base = textured_base(w, h); - let radius = shipped_default_params().temporal_radius; - let n = (3 * radius + 1) as usize; - - // Phase 1: a clean, static clip, checking the regularised field. - let mut clean_d = Nl4dDenoiser::::new(&client, shipped_default_params(), w, h) + let (width, height) = (96u32, 96u32); + let base = textured_base(width, height); + let shipped_params = shipped_default_params(); + let radius = shipped_params.temporal_radius; + let frame_count = (3 * radius + 1) as usize; + + let clean_params = shipped_default_params(); + let mut clean_denoiser = Nl4dDenoiser::::new(&client, clean_params, width, height) .expect("construction failed for the clean phase"); - for _ in 0..n { - clean_d.push_frame(&base); - let _ = clean_d.denoise_submit().expect("denoise_submit failed"); + for _ in 0..frame_count { + clean_denoiser.push_frame(&base); + let _ = clean_denoiser.denoise().expect("denoise failed"); } - let snap = clean_d.motion_snapshot().expect("a pass ran"); - // Block (2, 2) spans pixels 16..32 on both axes, well inside the - // frame and away from any edge-clamping effects. - let interior_block = (2 * snap.blocks_x + 2) as usize; - for (t, &k) in snap.offsets.iter().enumerate() { + + let snapshot = clean_denoiser.motion_snapshot().expect("a pass ran"); + + // Block (2, 2) spans pixels 16..32 on both axes, away from any edge clamping. + let interior_block = (2 * snapshot.blocks_x + 2) as usize; + for (neighbour, &k) in snapshot.offsets.iter().enumerate() { assert_eq!( - snap.vectors[t][interior_block], + snapshot.vectors[neighbour][interior_block], [0, 0], "neighbour k={k}: a static, noiseless clip's regularised field must read exactly \ zero at an interior block; a nonzero vector here is what a wrong pyramid slot, \ @@ -950,28 +918,36 @@ fn shipped_defaults_denoise_a_static_clip_and_regularise_its_field_to_zero() { ); } - // Phase 2: a noisy version of the same clip, checking the output. - let frames: Vec> = (0..n as u32) - .map(|seed| noisy_copy_of(&base, w, h, SIGMA, seed)) + let frames: Vec> = (0..frame_count as u32) + .map(|seed| noisy_field_over(&base, width, height, SIGMA, seed)) .collect(); - let mut noisy_d = Nl4dDenoiser::::new(&client, shipped_default_params(), w, h) + let noisy_params = shipped_default_params(); + let mut noisy_denoiser = Nl4dDenoiser::::new(&client, noisy_params, width, height) .expect("construction failed for the noisy phase"); let mut outputs: Vec> = Vec::new(); for frame in &frames { - noisy_d.push_frame(frame); - if let Some(pending) = noisy_d.denoise_submit().expect("denoise_submit failed") { - let frame = pending.wait().expect("readback failed"); - outputs.push(frame.into_f32().expect("f32 output")); + noisy_denoiser.push_frame(frame); + if let Some(frame) = noisy_denoiser.denoise().expect("denoise failed") { + outputs.push(frame); } } - noisy_d - .flush(|frame| outputs.push(frame.as_f32().expect("f32 denoiser").to_vec())) + + noisy_denoiser + .flush(|frame| { + let values = frame.to_vec(); + outputs.push(values); + }) .expect("flush failed"); - assert_eq!(outputs.len(), n, "expected one emitted frame per pushed frame"); - for (i, out) in outputs.iter().enumerate() { + assert_eq!( + outputs.len(), + frame_count, + "expected one emitted frame per pushed frame" + ); + + for (i, output) in outputs.iter().enumerate() { let noisy_psnr = psnr(&frames[i], &base); - let out_psnr = psnr(out, &base); + let out_psnr = psnr(output, &base); assert!( out_psnr > noisy_psnr + 6.0, "frame {i}: expected at least a 6 dB PSNR improvement over the noisy input at the \ diff --git a/av-denoise-core/src/nl4d/tests/pooled.rs b/av-denoise-core/src/nl4d/tests/pooled.rs index eb6e920..47c19b5 100644 --- a/av-denoise-core/src/nl4d/tests/pooled.rs +++ b/av-denoise-core/src/nl4d/tests/pooled.rs @@ -1,5 +1,7 @@ -use super::helpers::{R, make_client, noisy_copy_of, static_clip_params}; +use super::helpers::{R, make_client, static_clip_params}; +use crate::bench_api::HostIo; use crate::nl4d::{Nl4dDenoiser, Nl4dParams}; +use crate::nlmeans::tests::helpers::noisy_field_over; const WIDTH: u32 = 96; const HEIGHT: u32 = 64; @@ -25,7 +27,7 @@ fn faint_line_clip() -> Vec> { } (0..FRAMES) - .map(|seed| noisy_copy_of(&clean, WIDTH, HEIGHT, GRAIN, seed)) + .map(|seed| noisy_field_over(&clean, WIDTH, HEIGHT, GRAIN, seed)) .collect() } @@ -36,20 +38,19 @@ fn denoise(params: Nl4dParams, frames: &[Vec]) -> Vec> { let mut outputs = Vec::new(); for frame in frames { denoiser.push_frame(frame); - let pending = denoiser.denoise_submit().expect("denoise_submit failed"); - if let Some(pending) = pending { - let output = pending.wait().expect("readback failed"); - let output_frame = output.into_f32().expect("f32 output"); + let output = denoiser.denoise().expect("denoise failed"); + if let Some(output_frame) = output { outputs.push(output_frame); } } denoiser .flush(|frame| { - let output_frame = frame.as_f32().expect("f32 denoiser").to_vec(); + let output_frame = frame.to_vec(); outputs.push(output_frame); }) .expect("flush failed"); + outputs } @@ -73,14 +74,17 @@ fn line_contrast(outputs: &[Vec]) -> f64 { } } } + on_sum / on_count - off_sum / off_count } fn params(pooled_threshold: bool, lambda_ht: f32) -> Nl4dParams { + let clip_params = static_clip_params(2); + Nl4dParams { pooled_threshold, lambda_ht, - ..static_clip_params(2) + ..clip_params } } @@ -88,8 +92,10 @@ fn params(pooled_threshold: bool, lambda_ht: f32) -> Nl4dParams { fn pooling_keeps_more_of_a_faint_line() { let frames = faint_line_clip(); - let pooled = denoise(params(true, 3.78), &frames); - let plain = denoise(params(false, 3.78), &frames); + let pooled_params = params(true, 3.78); + let plain_params = params(false, 3.78); + let pooled = denoise(pooled_params, &frames); + let plain = denoise(plain_params, &frames); let pooled_contrast = line_contrast(&pooled); let plain_contrast = line_contrast(&plain); @@ -103,8 +109,10 @@ fn pooling_keeps_more_of_a_faint_line() { fn lambda_still_changes_the_output_with_pooling_on() { let frames = faint_line_clip(); - let gentle = denoise(params(true, 3.0), &frames); - let strong = denoise(params(true, 4.5), &frames); + let gentle_params = params(true, 3.0); + let strong_params = params(true, 4.5); + let gentle = denoise(gentle_params, &frames); + let strong = denoise(strong_params, &frames); assert_ne!(gentle, strong); } diff --git a/av-denoise-core/src/nl4d/tests/regularise.rs b/av-denoise-core/src/nl4d/tests/regularise.rs index 8d9aafe..981df56 100644 --- a/av-denoise-core/src/nl4d/tests/regularise.rs +++ b/av-denoise-core/src/nl4d/tests/regularise.rs @@ -7,33 +7,39 @@ use crate::nlmeans::motion::THSAD_PIXEL; const BLKSIZE: u32 = 16; const STEP: u32 = 8; -/// One launch over a `blocks_x x blocks_y` grid, returning the output -/// field and confidence. +/// One launch over a `blocks_x x blocks_y` grid, returning the output field and confidence. fn run( - w: u32, - h: u32, + width: u32, + height: u32, centre: &[f32], neighbour: &[f32], mv_in: &[i32], lambda: f32, ) -> (Vec, Vec) { let client = make_client(); - let blocks_x = w.div_ceil(STEP); - let blocks_y = h.div_ceil(STEP); + let blocks_x = width.div_ceil(STEP); + let blocks_y = height.div_ceil(STEP); let blocks = (blocks_x * blocks_y) as usize; assert_eq!(mv_in.len(), 2 * blocks); - let centre_buf = client.create_from_slice(f32::as_bytes(centre)); - let neighbour_buf = client.create_from_slice(f32::as_bytes(neighbour)); - let mv_in_buf = client.create_from_slice(i32::as_bytes(mv_in)); + + let centre_bytes = f32::as_bytes(centre); + let neighbour_bytes = f32::as_bytes(neighbour); + let mv_in_bytes = i32::as_bytes(mv_in); + let centre_buf = client.create_from_slice(centre_bytes); + let neighbour_buf = client.create_from_slice(neighbour_bytes); + let mv_in_buf = client.create_from_slice(mv_in_bytes); let mv_out = client.empty(2 * blocks * size_of::()); let conf_out = client.empty(blocks * size_of::()); let thsad = (BLKSIZE * BLKSIZE) as f32 * THSAD_PIXEL; + let grid = CubeCount::new_2d(blocks_x, blocks_y); + let dim = CubeDim::new_2d(8, 8); + unsafe { nl4d_mv_regularise::launch_unchecked::( &client, - CubeCount::new_2d(blocks_x, blocks_y), - CubeDim::new_2d(8, 8), + grid, + dim, ArrayArg::from_raw_parts(centre_buf, centre.len()), ArrayArg::from_raw_parts(neighbour_buf, neighbour.len()), ArrayArg::from_raw_parts(mv_in_buf, 2 * blocks), @@ -42,8 +48,8 @@ fn run( lambda * (BLKSIZE * BLKSIZE) as f32 * THSAD_PIXEL, 0.0, thsad, - w, - h, + width, + height, BLKSIZE, STEP, blocks_x, @@ -51,139 +57,144 @@ fn run( ); } - let mv = i32::from_bytes(&client.read_one(mv_out).expect("mv readback"))[..2 * blocks].to_vec(); - let conf = f32::from_bytes(&client.read_one(conf_out).expect("conf readback"))[..blocks].to_vec(); - (mv, conf) + let mv_bytes = client.read_one(mv_out).expect("mv readback"); + let conf_bytes = client.read_one(conf_out).expect("conf readback"); + let field = i32::from_bytes(&mv_bytes)[..2 * blocks].to_vec(); + let confidence = f32::from_bytes(&conf_bytes)[..blocks].to_vec(); + + (field, confidence) } /// A frame with distinct values everywhere. -fn textured(w: u32, h: u32, seed: u32) -> Vec { - (0..w * h) - .map(|i| { - let mut x = i +fn textured(width: u32, height: u32, seed: u32) -> Vec { + (0..width * height) + .map(|index| { + let mut hash = index .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(0x9E37_79B9)); - x ^= x >> 15; - x = x.wrapping_mul(0x85EB_CA6B); - x ^= x >> 13; - 0.2 + 0.6 * (x as f32 / u32::MAX as f32) + hash ^= hash >> 15; + hash = hash.wrapping_mul(0x85EB_CA6B); + hash ^= hash >> 13; + 0.2 + 0.6 * (hash as f32 / u32::MAX as f32) }) .collect() } -/// `neighbour(x, y) = centre(x - dx, y - dy)`, so the true vector is -/// `(dx, dy)`. -fn shifted(centre: &[f32], w: u32, h: u32, dx: i32, dy: i32) -> Vec { - let mut out = vec![0.0f32; (w * h) as usize]; - for y in 0..h as i32 { - for x in 0..w as i32 { - let sx = (x - dx).clamp(0, w as i32 - 1) as u32; - let sy = (y - dy).clamp(0, h as i32 - 1) as u32; - out[(y as u32 * w + x as u32) as usize] = centre[(sy * w + sx) as usize]; +/// `neighbour(x, y) = centre(x - dx, y - dy)`, so the true vector is `(dx, dy)`. +fn shifted(centre: &[f32], width: u32, height: u32, dx: i32, dy: i32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height as i32 { + for x in 0..width as i32 { + let source_x = (x - dx).clamp(0, width as i32 - 1) as u32; + let source_y = (y - dy).clamp(0, height as i32 - 1) as u32; + frame[(y as u32 * width + x as u32) as usize] = centre[(source_y * width + source_x) as usize]; } } - out + + frame } -fn uniform_field(blocks: usize, v: [i32; 2]) -> Vec { - let mut f = Vec::with_capacity(2 * blocks); +fn uniform_field(blocks: usize, vector: [i32; 2]) -> Vec { + let mut field = Vec::with_capacity(2 * blocks); for _ in 0..blocks { - f.push(v[0]); - f.push(v[1]); + field.push(vector[0]); + field.push(vector[1]); } - f + + field } -/// A flat centre and neighbour score the same SAD at every candidate, -/// so an outlier vector in a smooth field moves to the neighbourhood's -/// median as soon as the penalty is positive. +/// A flat centre and neighbour score the same SAD at every candidate, so an outlier in a smooth +/// field moves to the neighbourhood's median as soon as the penalty is positive. #[test] fn an_outlier_in_a_flat_region_moves_to_the_median() { - let (w, h) = (64u32, 64u32); - let blocks_x = w.div_ceil(STEP); - let blocks = (blocks_x * h.div_ceil(STEP)) as usize; - let flat = vec![0.5f32; (w * h) as usize]; + let (width, height) = (64u32, 64u32); + let blocks_x = width.div_ceil(STEP); + let blocks_y = height.div_ceil(STEP); + let blocks = (blocks_x * blocks_y) as usize; + let flat = vec![0.5f32; (width * height) as usize]; let mut field = uniform_field(blocks, [3, 1]); let outlier = (4 * blocks_x + 4) as usize; field[2 * outlier] = -6; field[2 * outlier + 1] = 5; - let (out, _) = run(w, h, &flat, &flat, &field, 1.0); - assert_eq!([out[2 * outlier], out[2 * outlier + 1]], [3, 1]); + let (output, _) = run(width, height, &flat, &flat, &field, 1.0); + assert_eq!([output[2 * outlier], output[2 * outlier + 1]], [3, 1]); + // Every other block already sits on its median and stays put. - for b in 0..blocks { - if b != outlier { - assert_eq!([out[2 * b], out[2 * b + 1]], [3, 1], "block {b} moved"); + for block in 0..blocks { + if block != outlier { + assert_eq!( + [output[2 * block], output[2 * block + 1]], + [3, 1], + "block {block} moved" + ); } } } -/// `lambda = 0` is a plain re-score. On a clean shift with the field -/// already correct nothing moves, and on a flat region ties go to the -/// block's own vector, so the outlier stays. +/// `lambda = 0` is a plain re-score, and on a flat region ties go to the block's own vector, so +/// the outlier stays. #[test] fn a_zero_penalty_keeps_the_input_field_on_ties() { - let (w, h) = (64u32, 64u32); - let blocks_x = w.div_ceil(STEP); - let blocks = (blocks_x * h.div_ceil(STEP)) as usize; - let flat = vec![0.5f32; (w * h) as usize]; + let (width, height) = (64u32, 64u32); + let blocks_x = width.div_ceil(STEP); + let blocks_y = height.div_ceil(STEP); + let blocks = (blocks_x * blocks_y) as usize; + let flat = vec![0.5f32; (width * height) as usize]; let mut field = uniform_field(blocks, [3, 1]); let outlier = (4 * blocks_x + 4) as usize; field[2 * outlier] = -6; field[2 * outlier + 1] = 5; - let (out, _) = run(w, h, &flat, &flat, &field, 0.0); - assert_eq!(out, field); + let (output, _) = run(width, height, &flat, &flat, &field, 0.0); + assert_eq!(output, field); } -/// A block whose own vector matches far better than the median keeps -/// it, because its SAD margin exceeds the penalty. +/// A block whose own vector matches far better than the median keeps it, because its SAD margin +/// exceeds the penalty. #[test] fn a_true_boundary_block_keeps_its_vector_when_the_sad_margin_wins() { - let (w, h) = (64u32, 64u32); - let blocks_x = w.div_ceil(STEP); - let blocks = (blocks_x * h.div_ceil(STEP)) as usize; - let centre = textured(w, h, 1); - // The whole neighbour is the centre shifted by (2, 0). - let neighbour = shifted(¢re, w, h, 2, 0); - // The field says (0, 0) everywhere except one interior block that - // knows the truth. + let (width, height) = (64u32, 64u32); + let blocks_x = width.div_ceil(STEP); + let blocks_y = height.div_ceil(STEP); + let blocks = (blocks_x * blocks_y) as usize; + let centre = textured(width, height, 1); + let neighbour = shifted(¢re, width, height, 2, 0); + + // The field says (0, 0) everywhere except one interior block that knows the truth. let mut field = uniform_field(blocks, [0, 0]); let truthful = (4 * blocks_x + 4) as usize; field[2 * truthful] = 2; - let (out, conf) = run(w, h, ¢re, &neighbour, &field, 1.0); - assert_eq!([out[2 * truthful], out[2 * truthful + 1]], [2, 0]); + let (output, confidence) = run(width, height, ¢re, &neighbour, &field, 1.0); + assert_eq!([output[2 * truthful], output[2 * truthful + 1]], [2, 0]); assert!( - conf[truthful] > 0.9, + confidence[truthful] > 0.9, "an exact match scores a high confidence, got {}", - conf[truthful] + confidence[truthful] ); - // Its neighbours see (2, 0) among the adjacent candidates and take - // it, since textured content beats the penalty of one pixel. + + // Its neighbours take (2, 0) from the adjacent candidates, since textured content beats the + // penalty of one pixel. let right = truthful + 1; - assert_eq!([out[2 * right], out[2 * right + 1]], [2, 0]); + assert_eq!([output[2 * right], output[2 * right + 1]], [2, 0]); } -/// The median rule picks the lower of two middle values, not the upper, -/// and a corner block's three-member neighbourhood gets exercised with -/// genuinely different values along the way. +/// The median rule picks the lower of two middle values, and a corner block's three-member +/// neighbourhood resolves too. /// -/// The field is flat, so every candidate scores the same zero SAD and -/// only the penalty against the median decides the winner. The centre -/// block's eight neighbours split four and four between `-1` and `2`, -/// so the lower median is `-1` while the upper median would be `2`, and -/// the kernel must land on `-1`. Block `(0, 0)` is a grid corner with -/// only three neighbours, one of which is the centre block, so its -/// median comes from the mismatched values `-1, -1, 99` rather than -/// eight identical entries. +/// The field is flat, so only the penalty against the median decides the winner. The centre +/// block's eight neighbours split four and four between `-1` and `2`, so the lower median is `-1`. +/// Corner block `(0, 0)` has only three neighbours, so its median comes from `-1, -1, 99`. #[test] fn the_median_rule_picks_the_lower_of_two_middle_values() { - let (w, h) = (24u32, 24u32); - let blocks_x = w.div_ceil(STEP); - let blocks_y = h.div_ceil(STEP); + let (width, height) = (24u32, 24u32); + let blocks_x = width.div_ceil(STEP); + let blocks_y = height.div_ceil(STEP); assert_eq!((blocks_x, blocks_y), (3, 3)); - let flat = vec![0.5f32; (w * h) as usize]; + + let flat = vec![0.5f32; (width * height) as usize]; #[rustfmt::skip] let field: Vec = vec![ @@ -192,64 +203,69 @@ fn the_median_rule_picks_the_lower_of_two_middle_values() { 2, 0, 2, 0, 2, 0, ]; - let (out, _) = run(w, h, &flat, &flat, &field, 1.0); + let (output, _) = run(width, height, &flat, &flat, &field, 1.0); let centre = (blocks_x + 1) as usize; assert_eq!( - [out[2 * centre], out[2 * centre + 1]], + [output[2 * centre], output[2 * centre + 1]], [-1, 0], "the centre block must take the lower median, not the upper one" ); + let corner = 0usize; assert_eq!( - [out[2 * corner], out[2 * corner + 1]], + [output[2 * corner], output[2 * corner + 1]], [-1, 0], "the corner block's three-member median must resolve too" ); } -/// Confidence is recomputed for the winner, not copied from the input. #[test] fn confidence_follows_the_winning_vector() { - let (w, h) = (64u32, 64u32); - let blocks_x = w.div_ceil(STEP); - let blocks = (blocks_x * h.div_ceil(STEP)) as usize; - let centre = textured(w, h, 2); - let neighbour = shifted(¢re, w, h, 1, 1); + let (width, height) = (64u32, 64u32); + let blocks_x = width.div_ceil(STEP); + let blocks_y = height.div_ceil(STEP); + let blocks = (blocks_x * blocks_y) as usize; + let centre = textured(width, height, 2); + let neighbour = shifted(¢re, width, height, 1, 1); let field = uniform_field(blocks, [1, 1]); - let (_, conf) = run(w, h, ¢re, &neighbour, &field, 1.0); + let (_, confidence) = run(width, height, ¢re, &neighbour, &field, 1.0); let interior = (3 * blocks_x + 3) as usize; assert!( - conf[interior] > 0.99, + confidence[interior] > 0.99, "a perfect match must score ~1, got {}", - conf[interior] + confidence[interior] ); let wrong = uniform_field(blocks, [-3, -3]); - let (_, conf) = run(w, h, ¢re, &neighbour, &wrong, 0.0); + let (_, confidence) = run(width, height, ¢re, &neighbour, &wrong, 0.0); assert!( - conf[interior] < 0.5, + confidence[interior] < 0.5, "a wrong vector on texture must score low, got {}", - conf[interior] + confidence[interior] ); - // A block whose winner is not its own vector, so this test does not - // pass merely by reporting candidate 0's confidence unconditionally. - // `right` sits beside a block that knows the true shift, and its own - // vector is wrong, so it wins on its left neighbour's vector, which is - // candidate 2. - let boundary_centre = textured(w, h, 1); - let boundary_neighbour = shifted(&boundary_centre, w, h, 2, 0); + // `right` has a wrong vector and wins on its left neighbour's true one, candidate 2, so this + // does not pass by reporting candidate 0's confidence unconditionally. + let boundary_centre = textured(width, height, 1); + let boundary_neighbour = shifted(&boundary_centre, width, height, 2, 0); let mut boundary_field = uniform_field(blocks, [0, 0]); let truthful = (4 * blocks_x + 4) as usize; boundary_field[2 * truthful] = 2; let right = truthful + 1; - let (out, conf) = run(w, h, &boundary_centre, &boundary_neighbour, &boundary_field, 1.0); - assert_eq!([out[2 * right], out[2 * right + 1]], [2, 0]); + let (output, confidence) = run( + width, + height, + &boundary_centre, + &boundary_neighbour, + &boundary_field, + 1.0, + ); + assert_eq!([output[2 * right], output[2 * right + 1]], [2, 0]); assert!( - conf[right] > 0.9, + confidence[right] > 0.9, "right wins its neighbour's exact match (candidate 2), so confidence must be high, got {}", - conf[right] + confidence[right] ); } diff --git a/av-denoise-core/src/nl4d/tests/strength_map.rs b/av-denoise-core/src/nl4d/tests/strength_map.rs index fa37a1f..441ffbc 100644 --- a/av-denoise-core/src/nl4d/tests/strength_map.rs +++ b/av-denoise-core/src/nl4d/tests/strength_map.rs @@ -1,7 +1,9 @@ -use super::helpers::{R, make_client, unit_noise}; +use super::helpers::{R, make_client}; +use crate::bench_api::HostIo; use crate::collab::kernels::fused::{STRENGTH_MAP_ALL, STRENGTH_MAP_LUMA}; use crate::nl4d::denoiser::strength_map_upload; use crate::nl4d::{Nl4dDenoiser, Nl4dParams}; +use crate::nlmeans::tests::helpers::seeded_unit_gaussian; use crate::nlmeans::{ ChannelMode, HqParams, @@ -23,16 +25,16 @@ const FRAMES: u32 = 12; const SEAM_MARGIN: u32 = 48; fn two_quarters() -> QuarterClasses { - let classes = vec![ - Some(QuarterClass { - flat: true, - luma: 0.1, - }), - Some(QuarterClass { - flat: false, - luma: 0.1, - }), - ]; + let flat = QuarterClass { + flat: true, + luma: 0.1, + }; + let textured = QuarterClass { + flat: false, + luma: 0.1, + }; + let classes = vec![Some(flat), Some(textured)]; + QuarterClasses::from_classes(2, 1, classes) } @@ -60,8 +62,11 @@ fn the_chroma_map_uploads_without_a_curve() { fn nothing_uploads_without_classes_or_a_map() { let classes = two_quarters(); - assert_eq!(strength_map_upload(None, 1, Some(LUMA_MAP), None), None); - assert_eq!(strength_map_upload(Some(&classes), 1, None, None), None); + let without_classes = strength_map_upload(None, 1, Some(LUMA_MAP), None); + let without_map = strength_map_upload(Some(&classes), 1, None, None); + + assert_eq!(without_classes, None); + assert_eq!(without_map, None); } fn ramp_luma(y: u32) -> f32 { @@ -95,13 +100,15 @@ fn half_textured_clip(channels: u32) -> Vec> { let pixel = y * WIDTH + x; let sample = pixel * channels + channel; let seed = frame_index * channels + channel; - let noise = unit_noise(pixel, seed); + let noise = seeded_unit_gaussian(pixel, seed); frame[sample as usize] = (clean + noise * noise_std).clamp(0.0, 1.0); } } } + frames.push(frame); } + frames } @@ -112,13 +119,16 @@ fn clip_params( shadow_soften: f32, ) -> Nl4dParams { let defaults = Nl4dParams::default(); + let hq = HqParams::default(); + let nlm = NlmParams { + channels, + prefilter: PrefilterMode::None, + hq: Some(hq), + ..defaults.nlm + }; + Nl4dParams { - nlm: NlmParams { - channels, - prefilter: PrefilterMode::None, - hq: Some(HqParams::default()), - ..defaults.nlm - }, + nlm, flat_boost, chroma_flat_boost, shadow_soften, @@ -154,19 +164,17 @@ fn denoise_clip(params: Nl4dParams, frames: &[Vec]) -> ClipRun { let mut classes_seen = false; for frame in frames { denoiser.push_frame(frame); - let pending = denoiser.denoise_submit().expect("denoise_submit failed"); + let output = denoiser.denoise().expect("denoise failed"); classes_seen |= denoiser.front_for_test().current_quarter_classes().is_some(); - if let Some(pending) = pending { - let output = pending.wait().expect("readback failed"); - let output_frame = output.into_f32().expect("f32 output"); + if let Some(output_frame) = output { outputs.push(output_frame); } } denoiser .flush(|frame| { - let output_frame = frame.as_f32().expect("f32 denoiser").to_vec(); + let output_frame = frame.to_vec(); outputs.push(output_frame); }) .expect("flush failed"); @@ -231,8 +239,10 @@ fn dark_rows() -> std::ops::Range { fn the_defaults_differ_from_a_map_of_ones() { let frames = half_textured_clip(1); - let defaults = denoise_clip(default_params(ChannelMode::Luma), &frames); - let ones = denoise_clip(unit_params(ChannelMode::Luma), &frames); + let default_map_params = default_params(ChannelMode::Luma); + let unit_map_params = unit_params(ChannelMode::Luma); + let defaults = denoise_clip(default_map_params, &frames); + let ones = denoise_clip(unit_map_params, &frames); assert!(defaults.classes_seen, "the clip should form quarter classes"); assert_ne!(defaults.outputs, ones.outputs); @@ -242,12 +252,15 @@ fn the_defaults_differ_from_a_map_of_ones() { fn shadow_soften_removes_less_from_textured_darks() { let frames = half_textured_clip(1); - let softened = denoise_clip(clip_params(ChannelMode::Luma, 1.0, 1.0, 0.65), &frames); - let ones = denoise_clip(unit_params(ChannelMode::Luma), &frames); + let softened_params = clip_params(ChannelMode::Luma, 1.0, 1.0, 0.65); + let unit_map_params = unit_params(ChannelMode::Luma); + let softened = denoise_clip(softened_params, &frames); + let ones = denoise_clip(unit_map_params, &frames); let left = 0..WIDTH / 2 - SEAM_MARGIN; - let softened_std = removed_std(&frames, &softened.outputs, 1, left.clone(), dark_rows()); - let ones_std = removed_std(&frames, &ones.outputs, 1, left, dark_rows()); + let dark = dark_rows(); + let softened_std = removed_std(&frames, &softened.outputs, 1, left.clone(), dark.clone()); + let ones_std = removed_std(&frames, &ones.outputs, 1, left, dark); assert!(softened_std < ones_std, "{softened_std} vs {ones_std}"); } @@ -255,8 +268,10 @@ fn shadow_soften_removes_less_from_textured_darks() { fn flat_boost_removes_more_from_flat_grain() { let frames = half_textured_clip(1); - let boosted = denoise_clip(clip_params(ChannelMode::Luma, 1.5, 1.0, 1.0), &frames); - let ones = denoise_clip(unit_params(ChannelMode::Luma), &frames); + let boosted_params = clip_params(ChannelMode::Luma, 1.5, 1.0, 1.0); + let unit_map_params = unit_params(ChannelMode::Luma); + let boosted = denoise_clip(boosted_params, &frames); + let ones = denoise_clip(unit_map_params, &frames); let right = WIDTH / 2 + SEAM_MARGIN..WIDTH; let boosted_std = removed_std(&frames, &boosted.outputs, 1, right.clone(), 0..HEIGHT); @@ -282,10 +297,13 @@ fn the_noise_map_off_ignores_the_strength_map() { #[test] fn a_pinned_sigma_never_applies_a_map() { let frames = half_textured_clip(1); + let defaults_hq = HqParams::with_sigma(0.02); let mut defaults = default_params(ChannelMode::Luma); - defaults.nlm.hq = Some(HqParams::with_sigma(0.02)); + defaults.nlm.hq = Some(defaults_hq); + + let ones_hq = HqParams::with_sigma(0.02); let mut ones = unit_params(ChannelMode::Luma); - ones.nlm.hq = Some(HqParams::with_sigma(0.02)); + ones.nlm.hq = Some(ones_hq); let with_defaults = denoise_clip(defaults, &frames); let with_ones = denoise_clip(ones, &frames); @@ -298,8 +316,10 @@ fn a_pinned_sigma_never_applies_a_map() { fn a_chroma_denoiser_applies_its_flat_boost() { let frames = half_textured_clip(2); - let boosted = denoise_clip(default_params(ChannelMode::Chroma), &frames); - let ones = denoise_clip(unit_params(ChannelMode::Chroma), &frames); + let boosted_params = default_params(ChannelMode::Chroma); + let unit_map_params = unit_params(ChannelMode::Chroma); + let boosted = denoise_clip(boosted_params, &frames); + let ones = denoise_clip(unit_map_params, &frames); assert!( boosted.classes_seen, @@ -331,8 +351,10 @@ fn a_chroma_denoiser_with_the_noise_map_off_builds_nothing() { fn a_chroma_denoiser_boosts_its_second_channel() { let frames = half_textured_clip(2); - let boosted = denoise_clip(default_params(ChannelMode::Chroma), &frames); - let unboosted = denoise_clip(clip_params(ChannelMode::Chroma, 1.5, 1.0, 0.65), &frames); + let boosted_params = default_params(ChannelMode::Chroma); + let unboosted_params = clip_params(ChannelMode::Chroma, 1.5, 1.0, 0.65); + let boosted = denoise_clip(boosted_params, &frames); + let unboosted = denoise_clip(unboosted_params, &frames); let boosted_samples = channel_samples(&boosted.outputs, 2, 1); let unboosted_samples = channel_samples(&unboosted.outputs, 2, 1); @@ -346,21 +368,24 @@ fn a_chroma_denoiser_boosts_its_second_channel() { #[test] fn the_luma_denoiser_passes_the_texture_cut_to_its_front() { - let denoiser = build_nl4d(ChannelMode::Luma, Nl4dParams::default()); + let params = Nl4dParams::default(); + let denoiser = build_nl4d(ChannelMode::Luma, params); assert_eq!(denoiser.front_for_test().flat_texture_cut(), Some(0.21)); } #[test] fn the_fused_yuv_denoiser_passes_the_texture_cut_to_its_front() { - let denoiser = build_nl4d(ChannelMode::Yuv, Nl4dParams::default()); + let params = Nl4dParams::default(); + let denoiser = build_nl4d(ChannelMode::Yuv, params); assert_eq!(denoiser.front_for_test().flat_texture_cut(), Some(0.21)); } #[test] fn the_chroma_denoiser_never_gets_a_texture_cut() { - let denoiser = build_nl4d(ChannelMode::Chroma, Nl4dParams::default()); + let params = Nl4dParams::default(); + let denoiser = build_nl4d(ChannelMode::Chroma, params); assert_eq!(denoiser.front_for_test().flat_texture_cut(), None); } diff --git a/av-denoise-core/src/nlmeans/align.rs b/av-denoise-core/src/nlmeans/align.rs index 8bfc96e..4885484 100644 --- a/av-denoise-core/src/nlmeans/align.rs +++ b/av-denoise-core/src/nlmeans/align.rs @@ -1,31 +1,23 @@ use cubecl::prelude::*; -/// Byte alignment every buffer binding must start on, taken from the -/// runtime the denoiser is running against. +/// The byte alignment every per-slot buffer binding must start on, read from the runtime. /// -/// A GPU rejects a bind group whose buffer offset is not a multiple of -/// its `min_storage_buffer_offset_alignment`. Every buffer this crate -/// slices into per-slot regions therefore pads its slot stride up to -/// this value. -/// -/// Each backend reports its own figure, 32 bytes on the Vulkan adapters -/// we test against and up to 256 elsewhere, which is why the value is -/// read from the runtime rather than assumed. -/// -/// It is carried as its own type rather than a bare `u64` so it cannot -/// be swapped by mistake with the width, height, or frame-count -/// arguments it travels alongside. +/// A GPU rejects a bind group whose offset is not a multiple of its +/// `min_storage_buffer_offset_alignment`, so per-slot strides pad up to this value. Backends differ, +/// from 32 bytes on the tested Vulkan adapters up to 256 elsewhere, so it is read rather than +/// assumed. It is its own type so it cannot be swapped with the width, height or frame-count +/// arguments it travels with. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) struct StorageAlign(u64); impl StorageAlign { /// The alignment `client`'s runtime requires. /// - /// cubecl aligns every allocation it hands out to this same value, - /// so a slot offset that is a multiple of it always lands on a - /// boundary the backend accepts. + /// cubecl aligns every allocation to the same value, so a slot offset that is a multiple of it + /// always lands on a boundary the backend accepts. pub(crate) fn from_client(client: &ComputeClient) -> Self { - Self::new(client.properties().memory.alignment) + let alignment = client.properties().memory.alignment; + Self::new(alignment) } /// A fixed alignment, for tests that have no runtime to ask. @@ -37,18 +29,18 @@ impl StorageAlign { Self(bytes.max(1)) } - /// `bytes` rounded up to the next aligned boundary. pub(crate) fn pad_bytes(self, bytes: u64) -> u64 { bytes.next_multiple_of(self.0) } - /// A count of `T` rounded up so that many elements cover a whole - /// number of alignment boundaries. + /// [Self::pad_bytes], or `None` when the padded size overflows `u64`. + pub(crate) fn checked_pad_bytes(self, bytes: u64) -> Option { + bytes.checked_next_multiple_of(self.0) + } + + /// A count of `T` rounded up so the elements cover whole alignment boundaries. /// - /// Alignments are powers of two, so for any `T` whose size divides - /// the alignment this lands exactly on a boundary. For a larger `T` - /// the elements are already aligned, so the count comes back - /// unchanged. + /// A `T` larger than the alignment is already aligned, so its count comes back unchanged. pub(crate) fn pad_elems(self, elems: usize) -> usize { let per_boundary = (self.0 as usize).div_ceil(size_of::()).max(1); elems.next_multiple_of(per_boundary) @@ -80,8 +72,7 @@ mod tests { #[test] fn pad_elems_tracks_a_larger_alignment() { - // A 256-byte boundary holds 64 f32s, so the same element count - // pads eight times further than it does at 32 bytes. + // A 256-byte boundary holds 64 f32s. let align = StorageAlign::new(256); assert_eq!(align.pad_elems::(1), 64); assert_eq!(align.pad_elems::(64), 64); diff --git a/av-denoise-core/src/nlmeans/denoiser.rs b/av-denoise-core/src/nlmeans/denoiser.rs deleted file mode 100644 index 0646098..0000000 --- a/av-denoise-core/src/nlmeans/denoiser.rs +++ /dev/null @@ -1,2257 +0,0 @@ -use cubecl::prelude::*; -use cubecl::server::Handle; - -use super::align::StorageAlign; -use super::kernels::{gpu_copy, gpu_unpack_wire}; -use super::motion::{self, MotionCtx, MotionEstimation, build_pyramid_for_slot, run_pyramid_build}; -use super::noise::{ - EMA_ALPHA, - NoiseCtx, - NoiseCurve, - NoiseEstimator, - QuarterClasses, - TemporalNoiseReading, - TemporalNoiseSample, - TemporalStatsCtx, - build_spatial_offset_lut, - correlation_factor, - noise_partials_slot_stride_bytes, - partials_len, - read_temporal_stats_slot, - run_noise_estimate, - run_temporal_noise_stats, - sigma_block_p25_from_partials, - sigma_from_abs_sum, - temporal_noise_reading, - temporal_stats_buf_bytes, - zero_temporal_stats_slot, -}; -use super::params::{NlmParams, SEPARABLE_THRESHOLD, sigma_eff, validate_dimensions}; -use super::pending::{Pending, empty_output, start_readback}; -use super::prefilter::{PrefilterCtx, PrefilterMode, run_prefilter}; -use super::{BLOCK_1D, Depth, MAX_GRID_1D}; -use crate::denoiser::{DenoiserError, FrameOutput, OutputFormat}; - -/// A denoised frame that has finished its kernels but is still resident -/// on the GPU. -/// -/// [`NlmDenoiser::denoise_submit_gpu`] returns this instead of starting a -/// readback, for a caller that queues more GPU work against the frame -/// rather than pulling it back to the host straight away. -/// -/// `handle` points at one of the denoiser's two output slots, and it -/// stays valid only until that slot is reused. A denoiser only has two -/// output slots, so at most two outstanding `GpuOutput`s (or -/// [`Pending`]s, which are built from the same slots) may exist for one -/// denoiser at a time. Submitting a third before an earlier one is -/// consumed reuses its slot and silently corrupts it. -pub struct GpuOutput { - /// The GPU buffer holding the denoised frame. - pub handle: Handle, - /// Which of the denoiser's two output slots `handle` came from. - pub slot: usize, -} - -/// Handles and geometry a collaborative stage needs to read the ring. -/// -/// [`NlmDenoiser::submit_machinery`] builds this instead of running -/// any NLM denoising kernel, so a caller that wants the frame ring, the -/// motion fields, and the confidence scores without the NLM weighting -/// itself can read them straight from here. -/// -/// The handles are views into the denoiser's own buffers, valid until -/// the next push or the next machinery step reuses the slots they point -/// into. -pub(crate) struct RingView { - /// The whole input ring, indexable by physical frame slot. - pub input: Handle, - /// Chained motion fields, one per neighbour index. - pub mv_field: Handle, - /// Per-block confidence, one plane per neighbour index. - pub confidence: Handle, - /// Physical ring slot of the centre frame. - pub centre_slot: u32, - /// Physical ring slot per neighbour, in logical ring order around - /// the centre frame, skipping the centre itself. - /// - /// This stays a host `Vec` rather than a GPU buffer, because the - /// grouping kernel that consumes it indexes it per candidate on the - /// device, and it is the caller's job to upload it, not this one's. - pub neighbour_slots: Vec, - /// `i32` element stride between neighbours in `mv_field`. - pub mv_stride: u32, - /// `f32` element stride between neighbours in `confidence`. - pub conf_stride: u32, - /// The luma pyramid the motion estimator analysed, the reference - /// ring's when a prefilter is active and the input ring's - /// otherwise. Level 0 of slot `s` starts at - /// `pyramid_slot_byte_offset(width, height, frame_count, 0, s, align)`. - pub pyramid: Handle, - /// How many frames the ring holds. - pub frame_count: u32, -} - -/// The stateful NLMeans denoiser that owns the GPU buffers. -/// -/// It keeps a ring of frames in `input_buf`. Each `push_frame` uploads -/// one frame into the ring, and each `denoise` cleans the current centre -/// frame using the neighbours around it. -pub struct NlmDenoiser { - pub(super) client: ComputeClient, - pub(super) params: NlmParams, - pub(super) width: u32, - pub(super) height: u32, - /// The byte alignment every per-slot buffer view has to start on, - /// read from `client`'s runtime at construction. See - /// [`StorageAlign`]. - pub(super) align: StorageAlign, - - /// A running count of frames pushed. Taken modulo the window size it - /// gives the next physical slot in `input_buf` to overwrite. - pub(super) ring_head: usize, - /// How many frames are loaded so far, capped at the window size. - pub(super) frames_loaded: usize, - /// How many real pushes the current stream has seen, not counting - /// the duplicates the denoiser adds at either end. - /// - /// [`Self::reset_stream_state`] clears this. - pub(super) real_pushes: usize, - - /// The frame ring itself, one slot per frame in the window. - pub(super) input_buf: Handle, - /// The reference ring, shaped exactly like `input_buf`. - /// - /// It only exists when a prefilter is set, and it supplies the - /// distances the `_ref` kernels read. - pub(super) reference_buf: Option, - /// CPU scratch for repacking 3-channel YUV into 4 lanes. Empty when - /// no padding is needed. - pub(super) padding_scratch: Vec, - /// CPU scratch the wire push concatenates its planes into, so one - /// push costs one transfer whatever the channel mode. Reused, so it - /// allocates nothing after the first frame. - pub(super) upload_scratch: Vec, - /// The weighted-pixel accumulator, one entry per stored channel per - /// pixel. - pub(super) accum: Handle, - /// The total weight at each pixel. - pub(super) weight_sum: Handle, - /// The largest neighbour weight at each pixel. - pub(super) max_weight: Handle, - /// Weight scratch for the path that compares a frame against itself. - pub(super) weight_buf: Handle, - /// The raw forward distance on the separable path. - pub(super) raw_fwd: Handle, - /// The raw backward distance on the separable path. - pub(super) raw_bwd: Handle, - /// The forward row sums on the separable path. - pub(super) tmp_hsum: Handle, - /// The backward row sums on the separable path. - pub(super) tmp_hsum_bwd: Handle, - /// Two denoised output buffers, used in turn. - /// - /// A new submit writes into the next slot while the previous one may - /// still be reading back, which lets one frame's kernels overlap - /// with the frame before it. - pub(super) outputs: [Handle; 2], - /// Which output slot the next submit writes into. - pub(super) next_output_slot: usize, - /// The format every readback this denoiser starts comes back in. - pub(super) output_format: OutputFormat, - /// Packed-word destinations, one per entry of `outputs`, allocated - /// only in wire mode. They rotate on the same slot counter, so each - /// is free again exactly when the `f32` slot it is packed from is. - pub(super) wire_outputs: Option<[Handle; 2]>, - /// CPU scratch the blocking `denoise()` and `flush()` paths reuse - /// through `Pending::wait_into`, so they do not allocate per frame. - /// - /// It holds `output_format`'s variant, so `wait_into` keeps its - /// allocation rather than replacing it. - pub(super) output_scratch: FrameOutput, - - pub(super) h2_inv_norm: f32, - /// The distance floor the main pass subtracts before weighting. - /// - /// It is zero for `NlmSpatial`, because comparing one pilot output - /// against another no longer carries a noise floor, so subtracting - /// one would overweight patches that do not really match. - /// - /// Every other prefilter mode leaves it equal to - /// `input_noise_offset`. - pub(super) noise_offset: f32, - /// The distance floor for comparisons against noisy input pixels. - /// - /// The pilot pass always uses this value, because its own inputs - /// still carry the full noise floor even when `noise_offset` has - /// been zeroed for the main pass. - pub(super) input_noise_offset: f32, - pub use_separable: bool, - pub(super) use_reference: bool, - - /// How correlated the grain is between neighbouring pixels, - /// smoothed over time on the same schedule as the sigma estimate. - /// - /// It is `None` until the stream's first temporal sample, following - /// the same seeding rule as `NoiseEstimator`. That first sample sets - /// the value directly rather than blending up from an assumed zero, - /// so a stream does not spend its opening frames under-corrected. - /// - /// It only updates when a temporal sample exists, so the fast path - /// and a fixed `sigma_override` leave it `None` for the whole - /// stream. That reads as zero, meaning white noise and no - /// correction. - /// - /// With `hq.windowed_noise_estimation` set, a fold with no temporal - /// sample clears this to `None` instead of leaving it where it was, - /// the same reasoning `noise_estimator_temporal_only` documents: - /// coasting on an older window's reading is itself a form of - /// cross-window history. - pub(super) rho_smoothed: Option, - /// One noise-floor offset per search candidate, for the weighting - /// kernels that compare a frame against itself. - /// - /// The table covers the whole search window, laid out row-major. - /// - /// It is rebuilt from `noise_offset` and `rho_smoothed` on every - /// submit. See `Self::rebuild_spatial_offset_lut`. - /// - /// While `rho_smoothed` is unset every entry equals the flat - /// `noise_offset` scalar this table replaced. - pub(super) spatial_offset_lut: Handle, - - /// Scratch for the first stage of the noise estimate, one slot per - /// ring position. - /// - /// It only exists when the noise level is measured automatically, - /// meaning HQ is on and no fixed sigma was given. - /// - /// Giving each ring position its own slot keeps a frame's partials - /// intact between the push that queues them and the later fold that - /// reads them back once that frame reaches the centre. One shared - /// region would be overwritten long before then. - pub(super) noise_partials: Option, - /// The per-channel Immerkær totals for each ring slot, gated the - /// same way as `noise_partials`. - pub(super) noise_results: Option, - /// The temporal residual statistics for each ring slot, one record - /// per spatial block. - /// - /// Each record holds the sum and sum of squares per channel, plus - /// the lag-1 total that reveals correlated grain. - /// - /// It only exists when the noise level is measured automatically and - /// the temporal radius is at least 1, because with no neighbour - /// there is nothing to take a difference against. - pub(super) temporal_stats_buf: Option, - /// Smooths the median chain's raw per-frame estimate into a steady - /// per-channel sigma, which feeds `h2_inv_norm` and `sigma_y`. - /// - /// It never updates when the noise level is not measured - /// automatically. - pub(super) noise_estimator: NoiseEstimator, - /// Smooths the low chain's raw per-frame estimate into a steady - /// per-channel sigma, which feeds only `input_noise_offset` and - /// `noise_offset`. - /// - /// The low chain reads more cautiously than the median chain, - /// combining a lower-quartile spatial statistic with a - /// lower-quartile temporal one. Its consumers are the ones where - /// reading the noise too high destroys detail. - /// - /// It is inert under the same condition as `noise_estimator`. - pub(super) noise_estimator_low: NoiseEstimator, - /// Smooths the low chain's raw per-frame estimate a second time, - /// built the same way `noise_estimator_low` is except the temporal - /// reading never passes through `correlation_factor`. - /// - /// A consumer that squares this sigma into a shrinkage threshold - /// pays for any over-read twice over, so it can read this estimator - /// instead of `noise_estimator_low` to get the temporal reading on - /// its own, without a second correction stacked on top of it. - /// - /// It is inert under the same condition as `noise_estimator`. - pub(super) noise_estimator_low_unboosted: NoiseEstimator, - /// Smooths the temporal median on its own, with no maximum taken - /// against an Immerkær spatial reading and no correlation boost. - /// - /// It only updates on a fold that has a temporal sample trustworthy - /// enough for `aggregate_temporal_noise_stats` to produce one. A fold - /// with no such sample, whether because the temporal radius is zero - /// or because too little of the frame held still, leaves this - /// estimator exactly where it was. `current_sigmas_temporal_only` - /// treats "never updated" as "no trustworthy reading has arrived - /// yet" and falls back to `noise_estimator_low_unboosted` for that - /// case. - /// - /// With `hq.windowed_noise_estimation` set, a fold with no - /// trustworthy sample clears this estimator instead of leaving it - /// where it was: "keeps going between folds" is itself a form of - /// history carried past the current window, the same thing - /// window-local estimation exists to remove from the other chains. - /// - /// It is inert under the same condition as `noise_estimator`. - pub(super) noise_estimator_temporal_only: NoiseEstimator, - /// The latest centre frame's luma noise curve, or `None` when no - /// curve could be built. - /// - /// [Self::update_noise_estimate] updates it alongside the estimator - /// chains, and [Self::reset_stream_state] clears it. - pub(super) noise_curve: Option, - /// The latest centre frame's quarter classes, present exactly when `noise_curve` is. - pub(super) quarter_classes: Option, - - /// The motion-compensation geometry, present while motion - /// compensation is active. - pub(super) mc_ctx: Option, - /// The shifted input ring, shaped exactly like `input_buf`. - /// - /// The temporal kernels read their neighbours from here, and the - /// centre slot is a straight copy of `input_buf`. - pub(super) compensated_input_buf: Option, - /// The shifted reference ring, shaped like - /// `compensated_input_buf`, present when a prefilter is active. - pub(super) compensated_reference_buf: Option, - /// The motion field, with one slice per neighbour and two `i32` - /// components per block. - /// - /// The first half of the slices hold the neighbours behind the - /// centre frame, and the second half those ahead of it. - pub(super) mv_field_buf: Option, - /// The ring of adjacent-frame motion fields, indexed by slot, then - /// direction, then block. - /// - /// Direction 0 runs from the older frame to the newer one, and - /// direction 1 the other way. - /// - /// It only exists when `MotionEstimation::Chained` is active, and - /// the direct path never touches it. - /// - /// `motion::pair_ring_slot_count` explains why the slot count is - /// exactly enough, and `Self::pair_slot` shows how a frame's place - /// in the push sequence picks a slot. - pub(super) pair_ring_buf: Option, - /// The luma pyramid, indexed by level, then frame, then pixel. - pub(super) pyramid_input: Option, - /// The same pyramid built from the reference ring, present when a - /// prefilter is active. - pub(super) pyramid_reference: Option, - - /// The block geometry for the confidence pass that runs without - /// motion compensation. - /// - /// It only exists when confidence weighting is on and motion - /// compensation is off. When motion compensation is on, `mc_ctx` - /// supplies the geometry instead. - pub(super) confidence_ctx: Option, - /// The per-block match confidence, with one slice per neighbour - /// laid out like `mv_field_buf` but holding a single `f32` per - /// block. - /// - /// It exists whenever confidence weighting is on, whichever context - /// supplied the geometry. - pub(super) confidence_buf: Option, - /// A single-level luma pyramid ring feeding the confidence pass that - /// runs without motion compensation. It exists alongside - /// `confidence_ctx`. - pub(super) confidence_pyramid: Option, - /// Somewhere to throw away the motion vector that pass produces, - /// since nothing shifts by it without motion compensation. It exists - /// alongside `confidence_ctx`. - pub(super) confidence_mv_scratch: Option, - /// A small placeholder passed as the fine block-match kernel's - /// confidence argument when confidence weighting is off but motion - /// compensation still runs. - /// - /// The kernel drops the confidence write at compile time in that - /// case, so this buffer is never indexed and its size does not - /// matter. It is tiny, so unlike the buffers above it is always - /// allocated. - pub(super) confidence_dummy: Handle, - /// The smoothed sigma for channel 0, the plane motion estimation - /// treats as luma, which feeds the confidence noise floor. - /// - /// It is zero unless HQ is on. A fixed `sigma_override` sets it once - /// at construction, while automatic estimation refreshes it every - /// submit. - pub(super) sigma_y: f32, - - /// Whether `run_temporal_stats_for_slot` asks the stats kernel for - /// its four luma-only lanes. - /// - /// It defaults to off, and [Self::set_luma_noise_fields] is the - /// only way to change it. Leaving it off roughly halves the stats - /// kernel's cost at 1080p. - pub(super) luma_noise_fields: bool, - - /// The cut the luma flat map vetoes textured quarters at, or `None` for no veto. - pub(super) flat_texture_cut: Option, - - /// Whether the stream's edges run off-centre passes instead of copied padding. - /// - /// Set by nl4d. With it on, a stream gets no leading copies and a - /// centre with no temporal reading borrows the nearest one ahead. - pub(super) shifted_edges: bool, -} - -impl NlmDenoiser { - /// Builds a new denoiser. - /// - /// # Panics - /// - /// This panics if the parameters or the frame dimensions are - /// invalid. The high-level [`crate::Denoiser`] checks both first and - /// reports them as a `Result`, so most callers should use that - /// instead. - pub fn new(client: &ComputeClient, params: NlmParams, width: u32, height: u32) -> Self { - Self::with_output_format(client, params, width, height, OutputFormat::F32) - } - - /// Builds a new denoiser whose readbacks come back in - /// `output_format`. - /// - /// [`OutputFormat::Wire`] gives the denoiser a packed-word buffer - /// per output slot, so a readback quantises on the GPU and only the - /// wire bytes cross the bus. - /// - /// # Panics - /// - /// Panics under the same conditions as [`Self::new`]. - pub fn with_output_format( - client: &ComputeClient, - params: NlmParams, - width: u32, - height: u32, - output_format: OutputFormat, - ) -> Self { - params - .validate() - .expect("invalid NlmParams, call params.validate() first to get this as a Result"); - validate_dimensions(width, height) - .expect("unsupported frame dimensions, call validate_dimensions first to get this as a Result"); - - let align = StorageAlign::from_client(client); - let stored_ch = params.channels.storage_count(); - let total_frames = params.total_frames(); - let pixels = (width * height) as usize; - let frame_bytes = pixels * stored_ch as usize * size_of::(); - let scalar_bytes = pixels * size_of::(); - - let input_buf = client.empty(frame_bytes * total_frames as usize); - let reference_buf = if params.prefilter.needs_reference_buf() { - Some(client.empty(frame_bytes * total_frames as usize)) - } else { - None - }; - - let padding_scratch = if params.channels.count() != stored_ch { - vec![0.0f32; pixels * stored_ch as usize] - } else { - Vec::new() - }; - - let accum = client.empty(frame_bytes); - let weight_sum = client.empty(scalar_bytes); - let max_weight = client.empty(scalar_bytes); - let weight_buf = client.empty(scalar_bytes); - let raw_fwd = client.empty(scalar_bytes); - let raw_bwd = client.empty(scalar_bytes); - let tmp_hsum = client.empty(scalar_bytes); - let tmp_hsum_bwd = client.empty(scalar_bytes); - let outputs = [client.empty(frame_bytes), client.empty(frame_bytes)]; - let wire_outputs = match output_format { - OutputFormat::F32 => None, - OutputFormat::Wire { depth } => { - let samples = pixels as u32 * params.channels.count(); - let words = samples.div_ceil(depth.wire_pack().samples_per_word()) as usize; - Some([ - client.empty(words * size_of::()), - client.empty(words * size_of::()), - ]) - }, - }; - - let h2_inv_norm = params.h2_inv_norm(); - let input_noise_offset = params.noise_offset(); - // The pilot pass compares noisy input pixels, so it always - // keeps the full noise floor. Under `NlmSpatial` the main pass - // compares one pilot output against another, which no longer - // carries that floor, so subtracting it would overweight - // patches that do not really match. - let noise_offset = match params.prefilter { - PrefilterMode::NlmSpatial { .. } => 0.0, - _ => input_noise_offset, - }; - let output_scratch = empty_output(pixels, params.channels.count(), output_format); - let use_separable = params.patch_radius > SEPARABLE_THRESHOLD; - let use_reference = params.prefilter.needs_reference_buf(); - - // This stays unset until the first temporal sample lands, so - // the initial table matches the flat `noise_offset` scalar it - // replaces exactly. - let rho_smoothed: Option = None; - let spatial_offset_lut = client.create_from_slice(f32::as_bytes(&build_spatial_offset_lut( - params.search_radius, - 0.0, - noise_offset, - ))); - - // Automatic noise estimation only runs when HQ is on and the - // caller has not pinned a fixed sigma. The fast path and the - // fixed-sigma path allocate neither buffer and never launch the - // estimate kernels. - let auto_noise = params.hq.is_some_and(|hq| hq.sigma_override.is_none()); - let (noise_partials, noise_results) = if auto_noise { - let partials_ring_bytes = - noise_partials_slot_stride_bytes(width, height, align) * total_frames as u64; - let n_results = (total_frames * 4) as usize; - ( - Some(client.empty(partials_ring_bytes as usize)), - Some(client.empty(n_results * size_of::())), - ) - } else { - (None, None) - }; - - // The temporal residual estimator also needs a real neighbour - // to take a difference against, so it stays inert at a temporal - // radius of 0 even when automatic estimation is on. - let temporal_stats_buf = if auto_noise && params.temporal_radius >= 1 { - Some(client.empty(temporal_stats_buf_bytes( - width, - height, - stored_ch, - total_frames, - align, - ))) - } else { - None - }; - - // The motion-compensation buffers, allocated only when motion - // compensation is active and the temporal window reaches past - // the centre frame. A spatial-only pass never touches them. - let mc_ctx = if params.motion_compensation.is_active() && params.temporal_radius > 0 { - MotionCtx::new(params.motion_compensation, width, height, align) - } else { - None - }; - - let ( - compensated_input_buf, - compensated_reference_buf, - mv_field_buf, - pyramid_input, - pyramid_reference, - ) = if let Some(ctx) = mc_ctx.as_ref() { - let comp_in = client.empty(frame_bytes * total_frames as usize); - let comp_ref = if use_reference { - Some(client.empty(frame_bytes * total_frames as usize)) - } else { - None - }; - let neighbours = (2 * params.temporal_radius) as u64; - let mv_field = client.empty((neighbours * ctx.mv_field_bytes_per_neighbour()) as usize); - let pyramid_pixels = - motion::pyramid_pixels_per_frame(width, height, ctx.pyramid_levels, ctx.align); - let pyr_in_bytes = pyramid_pixels * total_frames as usize * size_of::(); - let pyr_in = client.empty(pyr_in_bytes); - let pyr_ref = if use_reference { - Some(client.empty(pyr_in_bytes)) - } else { - None - }; - (Some(comp_in), comp_ref, Some(mv_field), Some(pyr_in), pyr_ref) - } else { - (None, None, None, None, None) - }; - - // The pair ring is allocated only when `Chained` estimation is - // active, either because it was asked for or because `Auto` - // resolved to it at this temporal radius, and only on top of - // motion compensation already being on. The direct path never - // reads or writes it. - let is_chained = matches!( - params - .motion_compensation - .resolved_estimation(params.temporal_radius), - Some(MotionEstimation::Chained { .. }) - ); - let pair_ring_buf = if is_chained { - mc_ctx.as_ref().map(|ctx| { - let pair_ring_slots = motion::pair_ring_slot_count(params.temporal_radius) as u64; - let bytes = pair_ring_slots * ctx.pair_slot_bytes(); - client.empty(bytes as usize) - }) - } else { - None - }; - - // Confidence weighting, in either of its two forms, only runs - // when HQ has it enabled and the temporal window reaches past - // the centre frame. - // - // This check applies with motion compensation on as well. - // Without it, every submit would pay for the fine kernel's - // confidence write whether or not anything read the result. - let confidence_active = - params.hq.is_some_and(|hq| hq.temporal_confidence) && params.temporal_radius > 0; - - // Geometry for the confidence-only pass, needed only when - // motion compensation is not already supplying block geometry - // through its own analyse pass. - // - // This costs real extra work, because it needs its own luma - // pyramid ring and a block-match kernel per neighbour. - let confidence_only_active = confidence_active && mc_ctx.is_none(); - let confidence_ctx = confidence_only_active.then(|| MotionCtx::confidence_only(width, height, align)); - - // The confidence buffer uses whichever block geometry is - // available, but only when confidence weighting is on. - let confidence_geometry = if confidence_active { - mc_ctx.as_ref().or(confidence_ctx.as_ref()) - } else { - None - }; - let confidence_buf = confidence_geometry.map(|ctx| { - let neighbours = (2 * params.temporal_radius) as u64; - client.empty((neighbours * ctx.confidence_bytes_per_neighbour()) as usize) - }); - // Always allocated, tiny, and reused whenever the fine - // block-match kernel runs without writing confidence. - let confidence_dummy = client.empty(size_of::()); - - let (confidence_pyramid, confidence_mv_scratch) = if let Some(ctx) = confidence_ctx.as_ref() { - let pyramid_pixels = - motion::pyramid_pixels_per_frame(width, height, ctx.pyramid_levels, ctx.align); - let pyr_bytes = pyramid_pixels * total_frames as usize * size_of::(); - let mv_scratch_len = ctx.mv_slots_per_neighbour() * 2 * size_of::(); - (Some(client.empty(pyr_bytes)), Some(client.empty(mv_scratch_len))) - } else { - (None, None) - }; - - // A fixed `sigma_override` is the only source before the first - // estimate lands, and automatic estimation refreshes this every - // submit. See `update_noise_estimate`. - // - // The fast path leaves it at zero, which - // `motion::sad_noise_floor` turns into a zero floor, exactly - // what a caller with no estimate should get. - let sigma_y = params.hq.and_then(|hq| hq.sigma_override).unwrap_or(0.0); - - Self { - client: client.clone(), - params, - width, - height, - align, - ring_head: 0, - frames_loaded: 0, - real_pushes: 0, - input_buf, - reference_buf, - padding_scratch, - upload_scratch: Vec::new(), - accum, - weight_sum, - max_weight, - weight_buf, - raw_fwd, - raw_bwd, - tmp_hsum, - tmp_hsum_bwd, - outputs, - next_output_slot: 0, - output_format, - wire_outputs, - output_scratch, - h2_inv_norm, - noise_offset, - input_noise_offset, - use_separable, - use_reference, - rho_smoothed, - spatial_offset_lut, - noise_partials, - noise_results, - temporal_stats_buf, - noise_estimator: NoiseEstimator::default(), - noise_estimator_low: NoiseEstimator::default(), - noise_estimator_low_unboosted: NoiseEstimator::default(), - noise_estimator_temporal_only: NoiseEstimator::default(), - noise_curve: None, - quarter_classes: None, - mc_ctx, - compensated_input_buf, - compensated_reference_buf, - mv_field_buf, - pair_ring_buf, - pyramid_input, - pyramid_reference, - confidence_ctx, - confidence_buf, - confidence_pyramid, - confidence_mv_scratch, - confidence_dummy, - sigma_y, - luma_noise_fields: false, - flat_texture_cut: None, - shifted_edges: false, - } - } - - /// Pushes a new frame into the ring buffer. - /// - /// `frame` holds `width * height * channels` `f32` values in - /// `[0, 1]`. A 3-channel frame is repacked into 4 lanes through a - /// reused CPU scratch buffer. - /// - /// For `PrefilterMode::External` use - /// [`Self::push_frame_with_reference`] instead. - pub fn push_frame(&mut self, frame: &[f32]) { - assert!( - !matches!(self.params.prefilter, PrefilterMode::External), - "push_frame_with_reference is required when prefilter == External" - ); - - let slot = self.upload_into(&self.input_buf.clone(), frame); - self.run_post_upload_stages(slot); - } - - /// Pushes a new frame held as wire bytes, one slice per channel. - /// - /// Each plane holds `width * height` samples at `depth`, and the GPU - /// normalises them and interleaves the planes. `planes` runs Y, U, V - /// for a fused frame and U, V for a chroma pair. - /// - /// For `PrefilterMode::External` use - /// [`Self::push_frame_with_reference`] instead. - pub fn push_frame_wire(&mut self, planes: &[&[u8]], depth: Depth) { - assert!( - !matches!(self.params.prefilter, PrefilterMode::External), - "push_frame_with_reference is required when prefilter == External" - ); - - let slot = self.upload_wire_into(&self.input_buf.clone(), planes, depth); - self.run_post_upload_stages(slot); - } - - /// The work every push runs once its frame is in `slot`, from the - /// noise estimate through to the ring advance. - fn run_post_upload_stages(&mut self, slot: usize) { - self.run_noise_estimate_for_slot(slot as u32); - self.run_temporal_stats_for_slot(slot as u32); - self.seed_noise_estimate_if_first_frame(slot as u32); - - if let PrefilterMode::NlmSpatial { strength_scale } = self.params.prefilter { - self.run_nlm_spatial_pilot(slot as u32, strength_scale) - .expect("nlm spatial pilot dispatch failed"); - } else if self.params.prefilter.is_gpu_internal() { - self.run_prefilter_for_slot(slot); - } - - self.build_pyramids_for_slot(slot as u32); - self.build_confidence_pyramid_for_slot(slot as u32); - self.run_pair_analyse_for_slot(slot as u32); - - self.advance_ring(); - self.prime_leading_edge_if_first(); - } - - /// Pushes a new frame together with a reference image the caller - /// prefiltered itself. - /// - /// This is what `PrefilterMode::External` needs. Both slices hold - /// `width * height * channels` `f32` values in `[0, 1]`. - pub fn push_frame_with_reference(&mut self, frame: &[f32], reference: &[f32]) { - assert!( - matches!(self.params.prefilter, PrefilterMode::External), - "push_frame_with_reference requires prefilter == External" - ); - - let slot = self.upload_into(&self.input_buf.clone(), frame); - let reference_buf = self - .reference_buf - .as_ref() - .expect("reference buffer must exist for External prefilter") - .clone(); - self.upload_into_slot(&reference_buf, reference, slot); - - // The same order as `push_frame`. The noise estimate and its - // first-frame seed run before anything that could read the - // sigma. Building the pyramids only needs the reference upload - // just above, not the noise estimate, so the ordering does not - // change what either step sees. - self.run_noise_estimate_for_slot(slot as u32); - self.run_temporal_stats_for_slot(slot as u32); - self.seed_noise_estimate_if_first_frame(slot as u32); - - self.build_pyramids_for_slot(slot as u32); - self.build_confidence_pyramid_for_slot(slot as u32); - self.run_pair_analyse_for_slot(slot as u32); - - self.advance_ring(); - self.prime_leading_edge_if_first(); - } - - /// Uploads `frame` into the next ring slot of `dst` and returns the - /// physical slot it wrote. - fn upload_into(&mut self, dst: &Handle, frame: &[f32]) -> usize { - let total_frames = self.params.total_frames() as usize; - let slot = self.ring_head % total_frames; - self.upload_into_slot(dst, frame, slot); - slot - } - - fn upload_into_slot(&mut self, dst: &Handle, frame: &[f32], slot: usize) { - let channels = self.params.channels.count() as usize; - let stored_ch = self.params.channels.storage_count() as usize; - let pixels = self.width as usize * self.height as usize; - let expected = pixels * channels; - - assert_eq!( - frame.len(), - expected, - "frame size mismatch: expected {expected}, got {}", - frame.len() - ); - - let staging = if channels == stored_ch { - self.client.create_from_slice(f32::as_bytes(frame)) - } else { - for i in 0..pixels { - let dst_off = i * stored_ch; - let src_off = i * channels; - self.padding_scratch[dst_off..dst_off + channels] - .copy_from_slice(&frame[src_off..src_off + channels]); - } - self.client - .create_from_slice(f32::as_bytes(&self.padding_scratch)) - }; - - self.copy_frame_into_slot(dst, slot, &staging, 0, 1); - } - - /// Uploads one wire-byte frame into the next ring slot of `dst` and - /// returns the physical slot it wrote. - /// - /// The planes are concatenated into `upload_scratch` and uploaded as - /// one buffer, so a push costs one transfer whatever the channel - /// mode. The scratch is reused, so the concatenation allocates - /// nothing after the first frame. - fn upload_wire_into(&mut self, dst: &Handle, planes: &[&[u8]], depth: Depth) -> usize { - let total_frames = self.params.total_frames() as usize; - let slot = self.ring_head % total_frames; - self.upload_wire_into_slot(dst, planes, depth, slot); - slot - } - - fn upload_wire_into_slot(&mut self, dst: &Handle, planes: &[&[u8]], depth: Depth, slot: usize) { - let channels = self.params.channels.count(); - let stored_ch = self.params.channels.storage_count(); - let pixels = self.width * self.height; - let plane_bytes = pixels as usize * depth.bytes_per_sample(); - - assert_eq!( - planes.len(), - channels as usize, - "plane count mismatch: expected {channels}, got {}", - planes.len() - ); - - // Ten and Twelve share a byte width, so the plane-length check - // below cannot tell them apart. A wrong depth here divides by the - // wrong maximum and darkens the whole frame without failing - // anything else, so it is pinned against the depth this denoiser - // returns frames in. - if let OutputFormat::Wire { depth: out_depth } = self.output_format { - assert_eq!( - depth, out_depth, - "wire push depth {depth:?} does not match the denoiser's output depth {out_depth:?}" - ); - } - - self.upload_scratch.clear(); - for plane in planes { - assert_eq!( - plane.len(), - plane_bytes, - "plane size mismatch: expected {plane_bytes}, got {}", - plane.len() - ); - debug_assert!( - wire_samples_in_range(plane, depth), - "a sample is larger than {depth:?} can express" - ); - self.upload_scratch.extend_from_slice(plane); - } - - // The kernel reads whole words, so a plane that ends mid-word - // needs its last word backed by real storage. - let words = self.upload_scratch.len().div_ceil(size_of::()); - self.upload_scratch.resize(words * size_of::(), 0); - - let src = self.client.create_from_slice(&self.upload_scratch); - - let elements = pixels * stored_ch; - let total_frames = self.params.total_frames() as usize; - let grid = elements.div_ceil(BLOCK_1D).clamp(1, MAX_GRID_1D); - let total_threads = grid * BLOCK_1D; - - // One `wire_pack` for both, since a hand-paired maximum and lane - // width decode the wrong bits. - let pack = depth.wire_pack(); - - unsafe { - gpu_unpack_wire::launch_unchecked::( - &self.client, - CubeCount::new_1d(grid), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(src, words), - ArrayArg::from_raw_parts(dst.clone(), total_frames * elements as usize), - pack.max(), - slot as u32 * elements, - pixels, - channels, - stored_ch, - pack.samples_per_word(), - elements, - total_threads, - ) - }; - } - - fn run_prefilter_for_slot(&self, slot: usize) { - let reference_buf = self - .reference_buf - .as_ref() - .expect("reference buffer must exist for GPU prefilter"); - - let ctx = PrefilterCtx { - width: self.width, - height: self.height, - channels: self.params.channels.count(), - stored_ch: self.params.channels.storage_count(), - frame_count: self.params.total_frames(), - frame: slot as u32, - input_buf: &self.input_buf, - reference_buf, - }; - - run_prefilter::(self.params.prefilter, &self.client, &ctx).expect("prefilter dispatch failed"); - } - - /// Builds the motion-estimation pyramid for `slot` on the input - /// ring, and on the reference ring when there is one. - /// - /// This does nothing when motion compensation is off. - fn build_pyramids_for_slot(&self, slot: u32) { - let Some(ctx) = self.mc_ctx.as_ref() else { - return; - }; - - let stored_ch = self.params.channels.storage_count(); - let frame_count = self.params.total_frames(); - - if let Some(pyr) = self.pyramid_input.as_ref() { - build_pyramid_for_slot::( - &self.client, - ctx, - self.width, - self.height, - frame_count, - slot, - &self.input_buf, - pyr, - stored_ch, - ) - .expect("input pyramid build dispatch failed"); - } - - if let (Some(pyr_ref), Some(ref_buf)) = (self.pyramid_reference.as_ref(), self.reference_buf.as_ref()) - { - build_pyramid_for_slot::( - &self.client, - ctx, - self.width, - self.height, - frame_count, - slot, - ref_buf, - pyr_ref, - stored_ch, - ) - .expect("reference pyramid build dispatch failed"); - } - } - - /// Extracts the luma plane for `slot` into the pyramid the - /// confidence pass uses when motion compensation is off. - /// - /// This does nothing unless that pass is active. - /// - /// It always reads `input_buf`, even with a prefilter set. Comparing - /// the raw input keeps this path simple, rather than duplicating the - /// reference ring's pyramid. - /// - /// It calls `run_pyramid_build` directly rather than going through - /// [`Self::build_pyramids_for_slot`]. That helper only touches the - /// motion-compensation pyramids, and its context is never present at - /// the same time as this one, so it would return without building - /// anything. - fn build_confidence_pyramid_for_slot(&self, slot: u32) { - let (Some(ctx), Some(pyr)) = (self.confidence_ctx.as_ref(), self.confidence_pyramid.as_ref()) else { - return; - }; - - run_pyramid_build::( - &self.client, - ctx, - self.width, - self.height, - self.params.total_frames(), - slot, - &self.input_buf, - pyr, - self.params.channels.storage_count(), - ) - .expect("confidence pyramid build dispatch failed"); - } - - /// Whether `Chained` motion estimation is in use, either because it - /// was asked for or because `Auto` resolved to it at this temporal - /// radius. - /// - /// `resolved_estimation` is the one place that decision is made. - /// - /// This is separate from whether motion compensation itself is - /// active, which callers still have to check, because that also - /// needs a temporal radius above 0. - pub(super) fn is_chained(&self) -> bool { - matches!( - self.params - .motion_compensation - .resolved_estimation(self.params.temporal_radius), - Some(MotionEstimation::Chained { .. }) - ) - } - - /// Measures motion between the slot a push just wrote and the one - /// before it, storing both directions into the pair ring. - /// - /// This does nothing unless `Chained` estimation is active, and it - /// also does nothing for a stream's very first frame, which has no - /// older partner to pair against. - /// - /// Composition covers that first gap by reading the priming - /// duplicate's zero-filled pair instead. See - /// [`Self::zero_pair_slot_for_duplicate`]. - fn run_pair_analyse_for_slot(&self, newer_slot: u32) { - if self.ring_head == 0 { - return; - } - let Some(mc) = self.mc_ctx.as_ref() else { - return; - }; - if !self.is_chained() { - return; - } - let pair_ring = self - .pair_ring_buf - .as_ref() - .expect("pair_ring allocated when Chained is active"); - - // Match against the cleaner of the two buffers, exactly as - // `run_motion_compensation` does on the direct path. - let pyramid = self.pyramid_reference.as_ref().unwrap_or_else(|| { - self.pyramid_input - .as_ref() - .expect("pyramid_input allocated when mc_ctx is Some") - }); - - let total_frames = self.params.total_frames(); - let older_slot = (newer_slot + total_frames - 1) % total_frames; - let pair_slot = self.pair_slot(0); - - motion::run_pair_analyse::( - &self.client, - mc, - self.width, - self.height, - total_frames, - older_slot, - newer_slot, - pair_slot, - pyramid, - pair_ring, - &self.confidence_dummy, - ) - .expect("pair analyse dispatch failed"); - } - - /// Fills the pair-ring slot for a duplicated frame with zeroes, - /// which happens while priming a stream and during the - /// end-of-stream flush. - /// - /// This does nothing unless `Chained` estimation is active. - fn zero_pair_slot_for_duplicate(&self) { - let Some(mc) = self.mc_ctx.as_ref() else { - return; - }; - if !self.is_chained() { - return; - } - let pair_ring = self - .pair_ring_buf - .as_ref() - .expect("pair_ring allocated when Chained is active"); - let pair_slot = self.pair_slot(0); - motion::zero_pair_slot::(&self.client, mc, pair_ring, pair_slot); - } - - /// Queues the Immerkær noise estimate for `slot` on the input ring. - /// - /// This does nothing unless automatic noise estimation is active. - /// - /// The results are normally read back later, in - /// [`Self::denoise_submit`], once `slot` reaches the centre of the - /// temporal window. A stream's very first frame is read immediately - /// as well. See [`Self::seed_noise_estimate_if_first_frame`]. - fn run_noise_estimate_for_slot(&self, slot: u32) { - let (Some(partials_buf), Some(results_buf)) = - (self.noise_partials.as_ref(), self.noise_results.as_ref()) - else { - return; - }; - - let stride = noise_partials_slot_stride_bytes(self.width, self.height, self.align); - let partials_slot = partials_buf.clone().offset_start((slot as u64) * stride); - - let ctx = NoiseCtx { - width: self.width, - height: self.height, - channels: self.params.channels.count(), - stored_ch: self.params.channels.storage_count(), - frame_count: self.params.total_frames(), - frame: slot, - slot, - input_buf: &self.input_buf, - partials_buf: &partials_slot, - results_buf, - }; - - run_noise_estimate::(&self.client, &ctx).expect("noise estimate dispatch failed"); - } - - /// Queues the temporal residual statistics for `slot`, comparing it - /// against the slot immediately before it in the ring. - /// - /// This does nothing unless the temporal estimator is active. A - /// stream's very first frame has no predecessor, so its record is - /// zeroed instead. - /// - /// The centre slot's statistics are read back and combined later, in - /// [`Self::update_noise_estimate`]. - fn run_temporal_stats_for_slot(&self, slot: u32) { - let Some(stats_buf) = self.temporal_stats_buf.as_ref() else { - return; - }; - if self.ring_head == 0 { - self.zero_temporal_stats_for_slot(slot); - return; - } - - let total_frames = self.params.total_frames(); - let slot_prev = (slot + total_frames - 1) % total_frames; - - let ctx = TemporalStatsCtx { - width: self.width, - height: self.height, - stored_ch: self.params.channels.storage_count(), - frame_count: total_frames, - slot_new: slot, - slot_prev, - input_buf: &self.input_buf, - stats_buf, - align: self.align, - }; - - run_temporal_noise_stats::(&self.client, &ctx, self.luma_noise_fields) - .expect("temporal noise stats dispatch failed"); - } - - /// Turns the temporal-stats kernel's four luma-only lanes on or off. - /// - /// It defaults to off. - pub(crate) fn set_luma_noise_fields(&mut self, on: bool) { - self.luma_noise_fields = on; - } - - /// Sets the cut the luma flat map vetoes textured quarters at. - /// - /// It defaults to `None`, which leaves every flat quarter flat. - pub(crate) fn set_flat_texture_cut(&mut self, cut: Option) { - self.flat_texture_cut = cut; - } - - /// The latest frame's luma noise curve, or `None` when no curve is - /// available. - pub(crate) fn current_noise_curve(&self) -> Option { - self.noise_curve - } - - /// The latest frame's quarter classes, present exactly when - /// [Self::current_noise_curve] is. - pub(crate) fn current_quarter_classes(&self) -> Option<&QuarterClasses> { - self.quarter_classes.as_ref() - } - - #[cfg(test)] - pub(crate) fn flat_texture_cut(&self) -> Option { - self.flat_texture_cut - } - - /// Fills a duplicated slot's temporal-stats region with zeroes. - /// - /// A duplicate holds exactly the same pixels as the slot before it, - /// so measuring the difference would only ever produce an all-zero - /// record. Writing the zeroes is the cheaper way to the same answer. - /// - /// This does nothing unless the temporal estimator is active. - fn zero_temporal_stats_for_slot(&self, slot: u32) { - let Some(stats_buf) = self.temporal_stats_buf.as_ref() else { - return; - }; - zero_temporal_stats_slot::( - &self.client, - stats_buf, - self.width, - self.height, - self.params.channels.storage_count(), - slot, - self.align, - ); - } - - /// Reads the noise estimate once, for a stream's very first frame, - /// so push-time work has a real sigma to use. - /// - /// Automatic estimation normally refreshes the derived filter - /// parameters at submit time, in [`Self::update_noise_estimate`]. - /// But push-time GPU work that reads them, namely the NLM pilot, - /// runs before the first submit ever happens. - /// - /// Without this, that work would run on the absolute-strength - /// fallback chosen at construction for every frame up to the first - /// submit. One blocking read of the estimate this push just queued - /// fixes it from frame one onward. - /// - /// The first frame is spotted through `frames_loaded`, the same - /// counter [`Self::prime_leading_edge_if_first`] checks, but read - /// here before [`Self::advance_ring`] moves it on. It applies at - /// every temporal radius, not only when priming happens. - /// - /// The first submit folds the same frame's estimate in a second - /// time, which reproduces these values to within floating-point - /// rounding rather than exactly. - fn seed_noise_estimate_if_first_frame(&mut self, slot: u32) { - if self.frames_loaded != 0 { - return; - } - let Some(results_buf) = self.noise_results.as_ref() else { - return; - }; - - let bytes = self - .client - .read_one(results_buf.clone()) - .expect("noise-estimate seed readback failed"); - let data = f32::from_bytes(&bytes); - - // The stream's first frame has no predecessor, so its stats - // record is zeroed rather than measured. Seed from Immerkær - // alone. - let imm_low = self - .read_noise_partials_low(slot) - .expect("noise-partials seed readback failed"); - self.fold_noise_estimate(data, slot as usize, None, imm_low); - } - - /// Folds one ring slot's noise totals into both estimator chains and - /// recomputes everything derived from them. - /// - /// Those derived values are `h2_inv_norm`, `noise_offset`, - /// `input_noise_offset`, and `sigma_y`. - /// - /// [`Self::seed_noise_estimate_if_first_frame`] and - /// [`Self::update_noise_estimate`] both call this. They differ only - /// in how they obtain the readings and which slot they pass. - /// - /// # The estimator chains - /// - /// Each chain starts from an Immerkær reading and takes the larger - /// of that and a temporal-residual reading, when one is available. - /// - /// The temporal estimator sees correlated grain the Immerkær mask - /// reads too low, but a shot with little static content, because of - /// motion or a scene change, makes its reading unreliable. Taking - /// the larger value lets it raise an estimate but never lower one. - /// - /// The chains differ only in which statistic they read. The median - /// chain takes the frame-mean Immerkær total and the per-block - /// median of the temporal reading. The low chain takes Immerkær's - /// own lower-quartile block statistic and the temporal lower - /// quartile. - /// - /// `noise_offset` weighs patch distances by the square of the sigma, - /// so reading too high there scrubs fine texture. The low chain's - /// cautious statistics keep it from over-reading on shots where - /// texture leaks into the temporal residuals. - /// - /// The strength and the confidence floor stay on the median chain, - /// because that is what the dark-footage calibration validated. - /// - /// A third estimator, `noise_estimator_low_unboosted`, folds the - /// same lower-quartile temporal reading as the low chain but skips - /// the correlation boost described below. It exists for a consumer - /// that squares its sigma into a threshold, where a boost meant to - /// offset a spatial estimator's blind spot on correlated grain would - /// otherwise be applied a second time to a temporal reading that - /// already tracks that grain directly. - /// - /// A fourth estimator, `noise_estimator_temporal_only`, folds the - /// temporal median by itself, with neither the maximum against the - /// Immerkær spatial reading nor the correlation boost. It exists for - /// the same squaring consumer, for a stronger reason than the boost - /// alone: a spatial mask reads regularly repeating texture the same - /// way it reads noise, and taking the maximum against it lets that - /// misreading through no matter how accurate the temporal side is. - /// This estimator only folds a new value on a fold whose temporal - /// sample was trustworthy enough for `aggregate_temporal_noise_stats` - /// to produce, and otherwise keeps whatever it last held. - /// - /// # Grain correlation - /// - /// The temporal sample's correlation figure folds into - /// `rho_smoothed` once, whichever chain reads it. That value feeds - /// the spatial-offset table. - /// - /// A stream's first fold sets it directly rather than blending from - /// an assumed zero, the same convention `NoiseEstimator` uses for - /// its own first sample. - /// - /// It stays unset on the fast path and with a fixed sigma, because - /// no temporal sample ever arrives there. - /// - /// # Window-local estimation - /// - /// With `hq.windowed_noise_estimation` set, every chain and - /// `rho_smoothed` take this fold's own sample outright instead of - /// blending it into their running state, the same way each does for - /// its very first sample. `noise_estimator_temporal_only` and - /// `rho_smoothed` also stop keeping their last reading on a fold - /// that has no temporal sample of its own, clearing instead, since - /// that "keep going" behaviour is itself a form of cross-window - /// history. See - /// [`crate::nlmeans::HqParams::windowed_noise_estimation`]. - fn fold_noise_estimate( - &mut self, - data: &[f32], - slot: usize, - temporal: Option, - imm_low: [f32; 3], - ) { - let channels = self.params.channels.count() as usize; - let base = slot * 4; - - let mut raw = [0.0f32; 3]; - for (c, s) in raw.iter_mut().enumerate().take(channels) { - *s = sigma_from_abs_sum(data[base + c], self.width, self.height); - } - let mut raw_low = imm_low; - let mut raw_low_unboosted = imm_low; - let mut raw_temporal_only: Option<[f32; 3]> = None; - - // Resolved once per fold rather than re-read per estimator - // below, and `false` on every call the default configuration - // makes, since `hq.windowed_noise_estimation` is `false` unless - // a caller set it. See `HqParams::windowed_noise_estimation`. - let windowed = self.params.hq.is_some_and(|hq| hq.windowed_noise_estimation); - - if let Some(sample) = temporal { - let factor = correlation_factor(sample.rho); - for c in 0..channels { - raw[c] = raw[c].max(sample.sigma[c] * factor); - raw_low[c] = raw_low[c].max(sample.sigma_low[c] * factor); - raw_low_unboosted[c] = raw_low_unboosted[c].max(sample.sigma_low[c]); - } - raw_temporal_only = Some(sample.sigma); - self.rho_smoothed = Some( - if windowed { - sample.rho - } else { - match self.rho_smoothed { - None => sample.rho, - Some(prev) => EMA_ALPHA * sample.rho + (1.0 - EMA_ALPHA) * prev, - } - }, - ); - } else if windowed { - // Window-local estimation must not let an earlier push's - // correlation reading leak into a fold that has no temporal - // sample of its own, the same reasoning that gates - // `noise_estimator_temporal_only` above. Without this, - // `rho_smoothed` keeps whatever an earlier window last - // measured, so `spatial_offset_lut` would depend on how - // many pushes preceded this fold rather than only the - // current window's own content. - self.rho_smoothed = None; - } - - // The user's nudge on the measured noise level. It applies - // after the two readings are combined and before the smoothing - // step, so it scales the smoothed estimate and everything - // derived from it. - // - // This is only reached when no fixed sigma was given, so the HQ - // parameters are always present here. - let sigma_scale = self.params.hq.map_or(1.0, |hq| hq.sigma_scale); - for c in 0..channels { - raw[c] *= sigma_scale; - raw_low[c] *= sigma_scale; - raw_low_unboosted[c] *= sigma_scale; - } - if let Some(raw_t) = raw_temporal_only.as_mut() { - for s in raw_t.iter_mut().take(channels) { - *s *= sigma_scale; - } - } - - let updated = self.noise_estimator.update(&raw[..channels], windowed); - let mut smoothed = [0.0f32; 3]; - smoothed[..channels].copy_from_slice(updated); - - let updated_low = self.noise_estimator_low.update(&raw_low[..channels], windowed); - let mut smoothed_low = [0.0f32; 3]; - smoothed_low[..channels].copy_from_slice(updated_low); - - self.noise_estimator_low_unboosted - .update(&raw_low_unboosted[..channels], windowed); - - match raw_temporal_only { - Some(raw_t) => { - self.noise_estimator_temporal_only - .update(&raw_t[..channels], windowed); - }, - // Window-local estimation must not let an earlier push's - // trustworthy reading leak into a fold that has none of its - // own, or the target frame's result would depend on how - // many pushes preceded it in this call, the same history - // dependence window-local estimation exists to remove from - // the other chains. See `noise_estimator_temporal_only`. - None if windowed => self.noise_estimator_temporal_only.reset(), - None => {}, - } - - let eff = sigma_eff(&smoothed[..channels], self.params.channels); - self.h2_inv_norm = self.params.h2_inv_norm_with(Some(eff)); - self.input_noise_offset = self.params.noise_offset_with(Some(&smoothed_low[..channels])); - self.noise_offset = match self.params.prefilter { - PrefilterMode::NlmSpatial { .. } => 0.0, - _ => self.input_noise_offset, - }; - // Channel 0 is whatever motion estimation already treats as - // luma, as `nlm_mc_extract_luma` shows, so the confidence floor - // uses the median chain's estimate for that same plane. - self.sigma_y = smoothed[0]; - } - - fn advance_ring(&mut self) { - let total_frames = self.params.total_frames() as usize; - self.ring_head += 1; - if self.frames_loaded < total_frames { - self.frames_loaded += 1; - } - self.real_pushes += 1; - } - - /// Copies one frame from a slot of `src` into a slot of `dst`, - /// entirely on the GPU. - /// - /// `dst` has to use the same ring layout as `input_buf`. - /// `src_slots` is how many frames `src` holds, which is 1 for a - /// frame-sized staging buffer. - /// - /// Both handles are bound whole, and the kernel picks the slots - /// through its own offset arguments. Binding a slot directly would - /// need its byte offset to be a multiple of the GPU's - /// `min_storage_buffer_offset_alignment`, and a - /// `width * height * stored_ch` frame stride rarely lands on one. - fn copy_frame_into_slot( - &self, - dst: &Handle, - slot: usize, - src: &Handle, - src_slot: usize, - src_slots: usize, - ) { - let stored_ch = self.params.channels.storage_count(); - let frame_size = self.width * self.height * stored_ch; - let dst_slots = self.params.total_frames() as usize; - - let grid = frame_size.div_ceil(BLOCK_1D).min(MAX_GRID_1D); - let total_threads = grid * BLOCK_1D; - - unsafe { - gpu_copy::launch_unchecked::( - &self.client, - CubeCount::new_1d(grid), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(src.clone(), src_slots * frame_size as usize), - ArrayArg::from_raw_parts(dst.clone(), dst_slots * frame_size as usize), - src_slot as u32 * frame_size, - slot as u32 * frame_size, - frame_size, - total_threads, - ) - }; - } - - /// Copies the very first pushed frame into the leading ring slots, - /// so the temporal window starts out balanced rather than dropping - /// the opening frames. - /// - /// [`Self::flush`] does the same thing at the other end of the - /// stream. Does nothing with shifted edges on, since that mode - /// stops windows at the clip start instead of padding them. - fn prime_leading_edge_if_first(&mut self) { - if self.shifted_edges { - return; - } - - let r = self.params.temporal_radius as usize; - - if r == 0 || self.frames_loaded != 1 { - return; - } - - for _ in 0..r { - self.duplicate_last_frame(); - self.frames_loaded += 1; - } - } - - /// Copies the most recently pushed frame into the next ring slot. - /// - /// This runs at the end of a stream, keeping the window full as the - /// real frames ahead of the centre run out. - /// - /// Slots never overlap, so copying inside the same buffer is safe. - /// - /// The reference ring is copied in step when it exists, so the - /// weights are never computed from a stale slot. - pub(super) fn duplicate_last_frame(&mut self) { - let total_frames = self.params.total_frames() as usize; - let last_slot = (self.ring_head - 1) % total_frames; - let next_slot = self.ring_head % total_frames; - - let input_buf = self.input_buf.clone(); - self.copy_frame_into_slot(&input_buf, next_slot, &input_buf, last_slot, total_frames); - - // Skipped for `NlmSpatial`, because the pilot dispatch below - // rebuilds this slot's reference from scratch and would - // overwrite the copy straight away. - if !matches!(self.params.prefilter, PrefilterMode::NlmSpatial { .. }) - && let Some(reference_buf) = self.reference_buf.clone() - { - self.copy_frame_into_slot(&reference_buf, next_slot, &reference_buf, last_slot, total_frames); - } - - // Keep the pyramid and the noise estimate for the duplicated - // slot in step too, so a later denoise sees valid state at - // every ring slot it visits rather than whatever an older frame - // left behind at this position. - // - // The NLM pilot needs the same treatment, or the duplicated - // slot's reference would keep whatever an older frame last - // wrote there. - if let PrefilterMode::NlmSpatial { strength_scale } = self.params.prefilter { - self.run_nlm_spatial_pilot(next_slot as u32, strength_scale) - .expect("nlm spatial pilot dispatch failed"); - } - self.build_pyramids_for_slot(next_slot as u32); - self.build_confidence_pyramid_for_slot(next_slot as u32); - self.run_noise_estimate_for_slot(next_slot as u32); - self.zero_temporal_stats_for_slot(next_slot as u32); - // Runs before `ring_head` advances, so `pair_slot(0)` reads the - // same pre-advance `ring_head` as `run_pair_analyse_for_slot` - // (see `Self::pair_slot`). - self.zero_pair_slot_for_duplicate(); - - self.ring_head += 1; - } - - /// Queues the denoise kernels for the current window, without - /// reading the result back. - /// - /// Use this instead of [`Self::denoise_submit`] when the output has - /// more GPU work ahead of it, so the round trip to the host can be - /// skipped until the value the caller actually wants is ready. The - /// returned [`GpuOutput`] documents the lifetime the caller has to - /// respect. - /// - /// Returns `Ok(None)` while the temporal window is still filling. - pub fn denoise_submit_gpu(&mut self) -> Result, anyhow::Error> { - let total_frames = self.params.total_frames() as usize; - if self.frames_loaded < total_frames { - return Ok(None); - } - - if self.noise_results.is_some() { - self.update_noise_estimate(self.params.temporal_radius)?; - } - self.rebuild_spatial_offset_lut(); - - let slot = self.next_output_slot; - self.next_output_slot = (slot + 1) % self.outputs.len(); - - self.run_denoise_kernels(slot)?; - - Ok(Some(GpuOutput { - handle: self.outputs[slot].clone(), - slot, - })) - } - - /// Runs the per-submit machinery a collaborative stage builds on, - /// without launching any NLM denoising kernel. - /// - /// This refreshes the noise estimate the same way - /// [`Self::denoise_submit_gpu`] does, then runs the same per-neighbour - /// motion estimate [`Self::run_motion_compensation`] runs, minus the - /// `run_compensate` shift into the compensated buffers. The returned - /// [`RingView`] reads the unmodified input ring directly, so a - /// caller predicts where a patch moved from the motion field and - /// searches around that prediction itself, rather than reading - /// content a warp has already resampled. - /// - /// Returns `Ok(None)` while the temporal window is still filling, the - /// same condition [`Self::denoise_submit_gpu`] checks. - /// - /// The returned [`RingView`] is centred on logical ring position - /// `center_t`. - /// - /// # Errors - /// - /// Returns an error if the denoiser was not built with motion - /// compensation and temporal confidence both active, since a - /// [`RingView`] has nothing meaningful to hand back otherwise. - pub(crate) fn submit_machinery(&mut self, center_t: u32) -> Result, DenoiserError> { - debug_assert!( - center_t < self.params.total_frames(), - "center_t must be a logical ring position" - ); - - let total_frames = self.params.total_frames() as usize; - if self.frames_loaded < total_frames { - return Ok(None); - } - - if self.noise_results.is_some() { - self.update_noise_estimate(center_t)?; - } - - let neighbour_slots = self.run_motion_machinery(center_t)?; - - let mc = self.mc_ctx.as_ref().ok_or_else(|| { - DenoiserError::Other(anyhow::anyhow!( - "submit_machinery requires motion compensation to be active" - )) - })?; - let mv_field = self - .mv_field_buf - .as_ref() - .expect("mv_field allocated when mc_ctx is Some") - .clone(); - let confidence = self - .confidence_buf - .as_ref() - .ok_or_else(|| { - DenoiserError::Other(anyhow::anyhow!( - "submit_machinery requires HQ temporal confidence to be active" - )) - })? - .clone(); - let pyramid = self - .pyramid_reference - .as_ref() - .or(self.pyramid_input.as_ref()) - .expect("pyramid allocated when mc_ctx is Some") - .clone(); - - Ok(Some(RingView { - input: self.input_buf.clone(), - mv_field, - confidence, - centre_slot: self.phys_frame(center_t as i32), - neighbour_slots, - mv_stride: (mc.mv_field_bytes_per_neighbour() / size_of::() as u64) as u32, - conf_stride: (mc.confidence_bytes_per_neighbour() / size_of::() as u64) as u32, - pyramid, - frame_count: self.params.total_frames(), - })) - } - - /// The motion-compensation geometry the last [`Self::submit_machinery`] - /// call used. - /// - /// # Panics - /// - /// Panics if the denoiser was not built with motion compensation - /// active. Only call this on a denoiser [`Self::submit_machinery`] - /// has already returned `Some` for. - pub(crate) fn motion_ctx(&self) -> &MotionCtx { - self.mc_ctx - .as_ref() - .expect("motion_ctx called without motion compensation active") - } - - /// The SAD two noisy copies of one block show by chance, the floor - /// [`Self::submit_machinery`] scored confidence against. - /// - /// # Panics - /// - /// Panics under the same condition as [`Self::motion_ctx`]. - pub(crate) fn sad_noise_floor_value(&self) -> f32 { - let blksize = self.motion_ctx().blksize; - let sigma = crate::nlmeans::dispatch::mc_sad_noise_floor_sigma(self.params.prefilter, self.sigma_y); - motion::sad_noise_floor(blksize, sigma) - } - - /// The SAD threshold [`Self::submit_machinery`] scored confidence - /// against, the same one [`Self::sad_noise_floor_value`] is measured - /// past. - /// - /// # Panics - /// - /// Panics under the same condition as [`Self::motion_ctx`]. - pub(crate) fn thsad_value(&self) -> f32 { - let blksize = self.motion_ctx().blksize; - let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); - motion::thsad(blksize, thsad_scale) - } - - /// The compute client this denoiser dispatches kernels through, for - /// a collaborative stage that reads a [`RingView`]'s handles back or - /// launches its own kernels against them. - pub(crate) fn compute_client(&self) -> &cubecl::client::ComputeClient { - &self.client - } - - /// The whole input ring, indexable by physical frame slot. - pub(crate) fn input_ring(&self) -> &Handle { - &self.input_buf - } - - /// Queues the denoise kernels for the current window and starts the - /// readback. - /// - /// Returns a [`Pending`] whose `wait()` produces the denoised frame. - /// - /// There are two output handles, so a caller can keep two `Pending`s - /// in flight and let one frame's kernels overlap the previous - /// frame's readback. - /// - /// A third concurrent submit would reuse the oldest pending frame's - /// output handle and quietly corrupt the results. The high-level - /// [`crate::Denoiser`] holds callers to that limit through its - /// `MAX_PENDING` constant. - /// - /// The frame comes back in the [`OutputFormat`] this denoiser was - /// built with. [`OutputFormat::Wire`] quantises and packs the frame - /// on the GPU before the readback, so only the wire bytes cross the - /// bus. - /// - /// Returns `Ok(None)` while the temporal window is still filling. - pub fn denoise_submit(&mut self) -> Result>, anyhow::Error> { - let Some(output) = self.denoise_submit_gpu()? else { - return Ok(None); - }; - - // Start the readback right away, so the GPU-side copy is queued - // before the caller dispatches the next frame's kernels. - let pixels = (self.width * self.height) as usize; - Ok(Some(start_readback( - &self.client, - output.handle, - self.wire_outputs.as_ref().map(|w| &w[output.slot]), - self.params.channels.count(), - self.params.channels.storage_count(), - pixels, - self.output_format, - ))) - } - - /// The packed-word destinations, which are `Some` only in wire mode. - #[cfg(test)] - pub(crate) fn wire_outputs_for_test(&self) -> Option<&[Handle; 2]> { - self.wire_outputs.as_ref() - } - - /// The smoothed per-channel sigma estimate NLMeans is currently - /// filtering with. - /// - /// This is `sigma_override` broadcast to every channel when HQ - /// pinned a fixed sigma. Otherwise it is the median chain's smoothed - /// estimate once one has landed, and zeros before that first - /// estimate and on the fast path where no estimate ever runs. - pub fn current_sigmas(&self) -> [f32; 3] { - if let Some(sigma) = self.params.hq.and_then(|hq| hq.sigma_override) { - return [sigma; 3]; - } - - let channels = self.params.channels.count() as usize; - let mut sigmas = [0.0f32; 3]; - if let Some(smoothed) = self.noise_estimator.current() { - sigmas[..channels].copy_from_slice(&smoothed[..channels]); - } - sigmas - } - - /// The smoothed per-channel sigma estimate from the low chain, with - /// the correlation boost left out of its temporal reading. - /// - /// This is `sigma_override` broadcast to every channel when HQ - /// pinned a fixed sigma, the same as [`Self::current_sigmas`]. - /// Otherwise it is `noise_estimator_low_unboosted`'s smoothed - /// estimate once one has landed, and zeros before that first - /// estimate and on the fast path where no estimate ever runs. - /// - /// See the "The estimator chains" section of [`Self::fold_noise_estimate`] - /// for why a consumer would want this instead of - /// [`Self::current_sigmas`]. - pub fn current_sigmas_low_unboosted(&self) -> [f32; 3] { - if let Some(sigma) = self.params.hq.and_then(|hq| hq.sigma_override) { - return [sigma; 3]; - } - - let channels = self.params.channels.count() as usize; - let mut sigmas = [0.0f32; 3]; - if let Some(smoothed) = self.noise_estimator_low_unboosted.current() { - sigmas[..channels].copy_from_slice(&smoothed[..channels]); - } - sigmas - } - - /// The smoothed per-channel sigma estimate from the temporal median - /// alone, with no maximum taken against an Immerkær spatial reading - /// and no correlation boost. - /// - /// This is `sigma_override` broadcast to every channel when HQ - /// pinned a fixed sigma, the same as the other chains. Otherwise it - /// is `noise_estimator_temporal_only`'s smoothed estimate, once a - /// fold has arrived with a temporal sample trustworthy enough for - /// `aggregate_temporal_noise_stats` to produce one. - /// - /// Before that first trustworthy reading, this falls back to - /// [`Self::current_sigmas_low_unboosted`] instead of reading zero, - /// which would under-filter. Two situations reach that fallback: a - /// temporal radius of zero, where no temporal sample ever exists - /// because there is no neighbouring frame to difference against, and - /// any push where too little of the frame held still for - /// `aggregate_temporal_noise_stats` to trust its own reading. Once a - /// trustworthy reading has landed the smoothed estimate keeps going - /// on later folds that individually lack one, the same way the other - /// chains keep going between folds. - /// - /// See the "The estimator chains" section of [`Self::fold_noise_estimate`] - /// for why a consumer would want this instead of - /// [`Self::current_sigmas_low_unboosted`]. - pub fn current_sigmas_temporal_only(&self) -> [f32; 3] { - if let Some(sigma) = self.params.hq.and_then(|hq| hq.sigma_override) { - return [sigma; 3]; - } - - let channels = self.params.channels.count() as usize; - if let Some(smoothed) = self.noise_estimator_temporal_only.current() { - let mut sigmas = [0.0f32; 3]; - sigmas[..channels].copy_from_slice(&smoothed[..channels]); - return sigmas; - } - - self.current_sigmas_low_unboosted() - } - - /// Refreshes the derived filter parameters from the centre slot's - /// noise estimate. - /// - /// That estimate was queued several pushes ago, when this slot was - /// first written. See [`Self::run_noise_estimate_for_slot`]. - /// - /// The blocking read therefore lands on work the GPU has already - /// finished, rather than stalling the pipeline behind a fresh - /// dispatch. - fn update_noise_estimate(&mut self, center_t: u32) -> Result<(), anyhow::Error> { - debug_assert!( - center_t < self.params.total_frames(), - "center_t must be a logical ring position" - ); - - let results_buf = self - .noise_results - .as_ref() - .expect("noise_results allocated when auto noise is active") - .clone(); - - let bytes = self - .client - .read_one(results_buf) - .map_err(|e| anyhow::anyhow!("noise-estimate results readback failed: {e}"))?; - let data = f32::from_bytes(&bytes); - - let center_slot = self.phys_frame(center_t as i32) as usize; - - let reading = self.borrow_reading_ahead(center_t)?; - let imm_low = self.read_noise_partials_low(center_slot as u32)?; - - // The curve tracks the sample it was built alongside, so it only - // ever changes on a fold that produced a trustworthy sample. - // Under windowed estimation a fold with no sample clears it, the - // same way `noise_estimator_temporal_only` clears rather than - // coasts on an older window's reading. See - // `HqParams::windowed_noise_estimation`. - let windowed = self.params.hq.is_some_and(|hq| hq.windowed_noise_estimation); - match (&reading.sample, reading.curve, reading.classes) { - (Some(_), curve, classes) => { - self.noise_curve = curve; - self.quarter_classes = classes; - }, - (None, _, _) if windowed => { - self.noise_curve = None; - self.quarter_classes = None; - }, - (None, _, _) => {}, - } - - self.fold_noise_estimate(data, center_slot, reading.sample, imm_low); - - Ok(()) - } - - /// Reads one ring slot's noise partials back and reduces them to the - /// low chain's per-channel estimate. - /// - /// The shared ring handle is sliced by byte offset, so the transfer - /// only covers one slot rather than the whole ring. This matches - /// [`read_temporal_stats_slot`]. - fn read_noise_partials_low(&self, slot: u32) -> Result<[f32; 3], anyhow::Error> { - let partials_buf = self - .noise_partials - .as_ref() - .expect("noise_partials allocated when auto noise is active"); - - let slot_len_bytes = partials_len(self.width, self.height) as u64 * size_of::() as u64; - let stride = noise_partials_slot_stride_bytes(self.width, self.height, self.align); - let total_bytes = self.params.total_frames() as u64 * stride; - let start = (slot as u64) * stride; - let end_trim = total_bytes - start - slot_len_bytes; - - let sliced = partials_buf.clone().offset_start(start).offset_end(end_trim); - let bytes = self - .client - .read_one(sliced) - .map_err(|e| anyhow::anyhow!("noise partials readback failed: {e}"))?; - let data = f32::from_bytes(&bytes); - - Ok(sigma_block_p25_from_partials( - data, - self.params.channels.count(), - self.width, - self.height, - )) - } - - /// Reads the centre slot's temporal residual statistics back and - /// combines them into one reading. - /// - /// Both the sample and the curve are `None` when the temporal - /// estimator is inactive, which happens at a temporal radius of 0 or - /// with a fixed sigma, and also when the combining step itself - /// declines to produce a sample. See [temporal_noise_reading]. - /// - /// The curve is only built when `luma_noise_fields` is on, since it - /// needs the stats kernel's luma-only lanes. See - /// [Self::set_luma_noise_fields]. - pub(super) fn read_temporal_noise(&self, slot: u32) -> Result { - let Some(stats_buf) = self.temporal_stats_buf.as_ref() else { - return Ok(TemporalNoiseReading { - sample: None, - curve: None, - classes: None, - }); - }; - - let stored_ch = self.params.channels.storage_count(); - let channels = self.params.channels.count(); - let frame_count = self.params.total_frames(); - - let records = read_temporal_stats_slot::( - &self.client, - stats_buf, - self.width, - self.height, - stored_ch, - frame_count, - slot, - self.align, - )?; - - Ok(temporal_noise_reading( - &records, - channels, - stored_ch, - self.width, - self.height, - self.luma_noise_fields, - self.flat_texture_cut, - )) - } - - /// Rebuilds `spatial_offset_lut` from the current `noise_offset` and - /// `rho_smoothed`. - /// - /// This runs once per submit, after `noise_offset` has been - /// refreshed, so the table and the scalar it comes from never - /// disagree. - /// - /// The rebuild is cheap, covering at most 289 floats. - fn rebuild_spatial_offset_lut(&mut self) { - let lut = build_spatial_offset_lut( - self.params.search_radius, - self.rho_smoothed.unwrap_or(0.0), - self.noise_offset, - ); - self.spatial_offset_lut = self.client.create_from_slice(f32::as_bytes(&lut)); - } - - /// Submits the denoise and waits for it, all in one call. - /// - /// Prefer [`Self::denoise_submit`] when the caller can hold a frame - /// in flight, which lets one frame's kernels overlap the previous - /// frame's readback. - /// - /// Returns `Ok(None)` while not enough frames have been pushed. - /// - /// The frame comes back in the [`OutputFormat`] this denoiser was - /// built with. - /// - /// On success the return borrows a reusable internal buffer. Copy it - /// out if the data has to survive another call into the denoiser. - pub fn denoise(&mut self) -> Result, anyhow::Error> { - let Some(pending) = self.denoise_submit()? else { - return Ok(None); - }; - // The scratch already holds this denoiser's format, so - // `wait_into` refills it and it keeps its allocation. - pending.wait_into(&mut self.output_scratch)?; - Ok(Some(&self.output_scratch)) - } - - /// How many tail frames a call to [`Self::flush`] must emit for the - /// stream pushed so far. - /// - /// While frames were being pushed the backend produced one output - /// per push beyond the temporal radius, or none at all if the - /// stream was shorter than that. This is the remaining difference, - /// so a caller driving [`Self::flush_step_gpu`] directly knows how - /// many `Some` results to collect before the stream is fully - /// drained. - /// - /// It reads as zero for spatial mode, where there is no trailing - /// context to drain, and for a stream that has not pushed anything - /// yet. - pub(crate) fn flush_target(&self) -> usize { - let temporal_radius = self.params.temporal_radius as usize; - if temporal_radius == 0 || self.real_pushes == 0 { - 0 - } else { - self.real_pushes.min(temporal_radius) - } - } - - /// How many genuine `push_frame`/`push_frame_with_reference` calls - /// the current stream has seen, not counting the duplicates - /// [`Self::prime_leading_edge_if_first`] and [`Self::flush`] add at - /// either end. - /// - /// A collaborative stage built on top of [`Self::submit_machinery`] - /// reads this to size its own end-of-stream drain, the way - /// [`Self::flush_target`] sizes this front end's. - pub(crate) fn real_pushes(&self) -> usize { - self.real_pushes - } - - /// Runs one step of the end-of-stream drain. It duplicates the most - /// recently pushed frame forward and submits the window that - /// results. - /// - /// Returns `Ok(None)` while the very first duplicates are still - /// filling out a window that never reached its full size during - /// pushing. Every step after the window is full returns `Ok(Some)`, - /// so a caller has to stop on its own once it has collected - /// [`Self::flush_target`] outputs, not on seeing `None` again. - /// - /// [`Self::flush`] is a loop over this method. A caller that wants - /// the tail frames to stay on the GPU for further work, rather than - /// making a round trip through the host, can drive this directly - /// instead and check its own count against [`Self::flush_target`]. - /// - /// This assumes there is a frame to duplicate, which means - /// `flush_target() > 0`. Calling it on spatial mode, or before any - /// frame has been pushed, duplicates a frame that was never written. - pub(crate) fn flush_step_gpu(&mut self) -> Result, anyhow::Error> { - let total_frames = self.params.total_frames() as usize; - - self.duplicate_last_frame(); - if self.frames_loaded < total_frames { - self.frames_loaded += 1; - } - - self.denoise_submit_gpu() - } - - /// Produces the frames still held at the end of a stream. - /// - /// For the last few frames the temporal window is kept full by - /// repeating the final frame. - /// - /// `sink` is called once per frame produced, and the frame it - /// receives is only valid for that call. It arrives in the - /// [`OutputFormat`] this denoiser was built with, quantised by the - /// same pack kernel as every streaming frame. - pub fn flush(&mut self, mut sink: impl FnMut(&FrameOutput)) -> Result<(), anyhow::Error> { - let target = self.flush_target(); - let mut emitted = 0usize; - let pixels = (self.width * self.height) as usize; - - // Every output slot is free here. A caller reaches a flush only - // once its streaming readbacks have landed, and the readback - // below blocks, so no other readback is ever reading the slot - // this step is handed. - while emitted < target { - if let Some(output) = self.flush_step_gpu()? { - let pending = start_readback( - &self.client, - output.handle, - self.wire_outputs.as_ref().map(|w| &w[output.slot]), - self.params.channels.count(), - self.params.channels.storage_count(), - pixels, - self.output_format, - ); - pending.wait_into(&mut self.output_scratch)?; - sink(&self.output_scratch); - emitted += 1; - } - } - - // Leave the denoiser ready for a fresh stream of the same - // shape. The GPU buffers stay allocated and are overwritten one - // slot at a time as new frames arrive, and - // `prime_leading_edge_if_first` refills the leading edge as - // soon as the new stream's first frame lands. - self.reset_stream_state(); - - Ok(()) - } - - /// Resets the stream-tracking indices, so the next push starts a - /// fresh temporal stream. - /// - /// The GPU buffers are deliberately left alone. Like the pyramid and - /// noise-estimate buffers, the pair ring is always written before it - /// is read. - /// - /// A fresh stream's opening pushes overwrite every slot they touch - /// before anything reads it, so content from the previous stream is - /// never seen. - pub fn reset_stream_state(&mut self) { - self.ring_head = 0; - self.frames_loaded = 0; - self.next_output_slot = 0; - self.real_pushes = 0; - self.noise_estimator.reset(); - self.noise_estimator_low.reset(); - self.noise_estimator_low_unboosted.reset(); - self.noise_estimator_temporal_only.reset(); - self.noise_curve = None; - self.quarter_classes = None; - self.rho_smoothed = None; - } - - /// The physical slot holding the oldest frame in the window. - /// - /// This is only meaningful once a full window has been pushed. - pub(super) fn ring_start(&self) -> u32 { - let total_frames = self.params.total_frames() as usize; - (self.ring_head % total_frames) as u32 - } - - /// Turns a logical frame index within the window into its physical - /// slot inside `input_buf`. - pub(super) fn phys_frame(&self, logical: i32) -> u32 { - let total_frames = self.params.total_frames() as i32; - let wrapped = logical.rem_euclid(total_frames); - ((self.ring_start() as i32 + wrapped).rem_euclid(total_frames)) as u32 - } - - /// The pair-ring slot holding the gap between two neighbouring - /// frames in the window. - /// - /// It reduces `ring_head`, the running count of frames pushed - /// including duplicates, by the pair-ring size rather than by the - /// window size `Self::phys_frame` uses. - /// - /// Two callers reach the same slot for the same physical pair. - /// - /// At push time, with a gap index of 0 and `ring_head` still at the - /// value the frame just written was given, it returns the slot that - /// frame's pair with its predecessor belongs in. - /// - /// At compose time, with `ring_head` already advanced and the gap - /// index measured out from the window's centre, it returns the slot - /// an earlier push wrote. - /// - /// The two differ only in how far `ring_head` has moved since the - /// pair was created, and the gap index cancels exactly that much, so - /// the sum lands on the same slot either way. - pub(super) fn pair_slot(&self, gap_index: i32) -> u32 { - let radius = self.params.temporal_radius as i32; - debug_assert!( - radius > 0, - "pair ring is only meaningful when temporal_radius > 0" - ); - let n = 2 * radius; - ((self.ring_head as i32 + gap_index).rem_euclid(n)) as u32 - } -} - -/// True when every sample in `plane` fits the range `depth` expresses. -/// -/// `gpu_unpack_wire` is branch-free and divides every sample by the -/// depth's maximum, so a larger sample normalises above 1.0 and reaches -/// the filter as a value no clean frame can hold. Only 10 and 12-bit can -/// carry one, in the unused high bits of a 16-bit lane, so 8-bit is -/// always in range. -fn wire_samples_in_range(plane: &[u8], depth: Depth) -> bool { - if depth.bytes_per_sample() == 1 { - return true; - } - - let max = depth.max_value() as u32; - plane - .as_chunks::<2>() - .0 - .iter() - .all(|&s| u32::from(u16::from_le_bytes(s)) <= max) -} diff --git a/av-denoise-core/src/nlmeans/denoiser/machinery.rs b/av-denoise-core/src/nlmeans/denoiser/machinery.rs new file mode 100644 index 0000000..9766af1 --- /dev/null +++ b/av-denoise-core/src/nlmeans/denoiser/machinery.rs @@ -0,0 +1,146 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::NlmDenoiser; +use crate::nlmeans::dispatch::mc_sad_noise_floor_sigma; +use crate::nlmeans::motion::{self, MotionCtx}; +use crate::nlmeans::params::ChannelMode; + +/// Handles and geometry a collaborative stage needs to read the ring. +/// +/// The handles view the denoiser's own buffers and stay valid until the next push or machinery +/// step reuses their slots. +pub(crate) struct RingView { + /// The whole input ring, indexable by physical frame slot. + pub input: Handle, + /// Chained motion fields, one per neighbour. + pub mv_field: Handle, + /// Per-block confidence, one plane per neighbour. + pub confidence: Handle, + /// Physical ring slot of the centre frame. + pub centre_slot: u32, + /// Physical ring slot of each neighbour in logical order, skipping the centre. + pub neighbour_slots: Vec, + /// `i32` element stride between neighbours in `mv_field`. + pub mv_stride: u32, + /// `f32` element stride between neighbours in `confidence`. + pub conf_stride: u32, + /// The luma pyramid the motion estimator analysed. + /// + /// It is built from the reference ring when a prefilter is active. Level 0 of slot `s` starts + /// at `pyramid_slot_byte_offset(width, height, frame_count, 0, s, align)`. + pub pyramid: Handle, + pub frame_count: u32, +} + +impl NlmDenoiser { + /// Runs the per-submit noise and motion estimates without launching any NLM kernel. + /// + /// The view is centred on logical ring position `center_t` and reads the unshifted input + /// ring, so a caller searches around the predicted motion itself. Returns `Ok(None)` while + /// the window is still filling. + /// + /// # Errors + /// + /// Returns an error unless motion compensation and temporal confidence are both active. + pub(crate) fn submit_machinery(&mut self, center_t: u32) -> Result, anyhow::Error> { + debug_assert!( + center_t < self.params.total_frames(), + "center_t must be a logical ring position" + ); + + let total_frames = self.params.total_frames() as usize; + if self.frames_loaded < total_frames { + return Ok(None); + } + + if self.noise_results.is_some() { + self.update_noise_estimate(center_t)?; + } + + let neighbour_slots = self.run_motion_machinery(center_t)?; + + let motion_ctx = self + .mc_ctx + .as_ref() + .ok_or_else(|| anyhow::anyhow!("submit_machinery requires motion compensation to be active"))?; + let mv_field = self + .mv_field_buf + .as_ref() + .expect("mv_field allocated when mc_ctx is Some") + .clone(); + let confidence = self + .confidence_buf + .as_ref() + .ok_or_else(|| anyhow::anyhow!("submit_machinery requires HQ temporal confidence to be active"))? + .clone(); + let pyramid = self + .pyramid_reference + .as_ref() + .or(self.pyramid_input.as_ref()) + .expect("pyramid allocated when mc_ctx is Some") + .clone(); + + let centre_slot = self.phys_frame(center_t as i32); + let mv_stride = (motion_ctx.mv_field_bytes_per_neighbour() / size_of::() as u64) as u32; + let conf_stride = (motion_ctx.confidence_bytes_per_neighbour() / size_of::() as u64) as u32; + + let view = RingView { + input: self.input_buf.clone(), + mv_field, + confidence, + centre_slot, + neighbour_slots, + mv_stride, + conf_stride, + pyramid, + frame_count: self.params.total_frames(), + }; + Ok(Some(view)) + } + + /// The motion-compensation geometry. + /// + /// # Panics + /// + /// Panics unless motion compensation is active. + pub(crate) fn motion_ctx(&self) -> &MotionCtx { + self.mc_ctx + .as_ref() + .expect("motion_ctx called without motion compensation active") + } + + /// The SAD two noisy copies of one block show by chance, which confidence is scored against. + /// + /// # Panics + /// + /// Panics unless motion compensation is active. + pub(crate) fn sad_noise_floor_value(&self) -> f32 { + let blksize = self.motion_ctx().blksize; + let sigma = mc_sad_noise_floor_sigma(self.params.prefilter, self.sigma_y); + motion::sad_noise_floor(blksize, sigma) + } + + /// The SAD threshold that confidence is scored against. + /// + /// # Panics + /// + /// Panics unless motion compensation is active. + pub(crate) fn thsad_value(&self) -> f32 { + let blksize = self.motion_ctx().blksize; + let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); + motion::thsad(blksize, thsad_scale) + } + + pub(crate) fn compute_client(&self) -> &ComputeClient { + &self.client + } + + pub(crate) fn input_ring(&self) -> &Handle { + &self.input_buf + } + + pub(crate) fn frame_shape(&self) -> (u32, u32, ChannelMode) { + (self.width, self.height, self.params.channels) + } +} diff --git a/av-denoise-core/src/nlmeans/denoiser/mod.rs b/av-denoise-core/src/nlmeans/denoiser/mod.rs new file mode 100644 index 0000000..8648bd6 --- /dev/null +++ b/av-denoise-core/src/nlmeans/denoiser/mod.rs @@ -0,0 +1,507 @@ +mod machinery; +mod noise; +mod ring; +mod sizes; +mod stages; + +use cubecl::prelude::*; +use cubecl::server::Handle; + +pub(crate) use self::machinery::RingView; +pub(crate) use self::sizes::{BufferSize, check_u32_indexable, front_buffer_sizes}; +use super::align::StorageAlign; +use super::motion::{self, MotionCtx, MotionEstimation}; +use super::noise::{ + NoiseCurve, + NoiseEstimator, + QuarterClasses, + build_spatial_offset_lut, + noise_partials_slot_stride_bytes, + temporal_stats_buf_bytes, +}; +use super::params::{NlmParams, SEPARABLE_THRESHOLD, validate_dimensions}; +use super::prefilter::PrefilterMode; +use crate::engine::{DevicePlane, IngestTarget, SampleFormat, ingest}; + +/// A denoised frame still resident on the GPU. +/// +/// `handle` points at one of the denoiser's two output slots and stays valid until that slot is +/// reused, so holding more than two at once lets a later submit overwrite an earlier frame. +pub struct GpuOutput { + pub handle: Handle, +} + +/// The stateful NLMeans denoiser that owns the GPU buffers. +/// +/// Each push ingests one frame into a ring of frames, and each submit cleans the centre frame +/// using the neighbours around it. +pub struct NlmDenoiser { + pub(super) client: ComputeClient, + pub(super) params: NlmParams, + pub(super) width: u32, + pub(super) height: u32, + /// The byte alignment every per-slot buffer view starts on. + pub(super) align: StorageAlign, + + /// Frames pushed including duplicates, which modulo the window size picks the next slot. + pub(super) ring_head: usize, + /// Frames loaded so far, capped at the window size. + pub(super) frames_loaded: usize, + /// Real pushes in the current stream, not counting the duplicates added at either end. + pub(super) real_pushes: usize, + + /// The frame ring, one slot per frame in the window. + pub(super) input_buf: Handle, + /// The prefiltered ring the `_ref` distance kernels read, shaped like `input_buf`. + pub(super) reference_buf: Option, + /// A 4-byte handle bound for planes a kernel never reads. + pub(super) placeholder: Handle, + /// The weighted pixel sum, one entry per stored channel per pixel. + pub(super) accum: Handle, + pub(super) weight_sum: Handle, + pub(super) max_weight: Handle, + /// Weight scratch for the path that compares a frame against itself. + pub(super) weight_buf: Handle, + /// The raw forward distance on the separable path. + pub(super) raw_fwd: Handle, + /// The raw backward distance on the separable path. + pub(super) raw_bwd: Handle, + /// The forward row sums on the separable path. + pub(super) tmp_hsum: Handle, + /// The backward row sums on the separable path. + pub(super) tmp_hsum_bwd: Handle, + /// Two output buffers used in turn, so one frame's kernels overlap the previous readback. + pub(super) outputs: [Handle; 2], + pub(super) next_output_slot: usize, + + /// Scales patch distances into weights, one over the squared strength times the patch area. + pub(super) h2_inv_norm: f32, + /// The distance floor the main pass subtracts before weighting. + /// + /// It is zero under `NlmSpatial`, because two pilot outputs carry no noise floor and + /// subtracting one would overweight patches that do not really match. + pub(super) noise_offset: f32, + /// The distance floor for comparisons against noisy input pixels, which the pilot always uses. + pub(super) input_noise_offset: f32, + pub use_separable: bool, + pub(super) use_reference: bool, + + /// The smoothed grain correlation between neighbouring pixels. + /// + /// The first temporal sample sets it directly so the opening frames are not under-corrected. + /// `None` reads as white noise. + pub(super) rho_smoothed: Option, + /// One noise-floor offset per search candidate, row-major, for the self-comparison kernels. + pub(super) spatial_offset_lut: Handle, + + /// Per-ring-slot scratch for the first stage of the noise estimate. + /// + /// Each slot keeps a frame's partials intact until that frame reaches the centre and is + /// folded. + pub(super) noise_partials: Option, + /// The per-channel Immerkær totals for each ring slot. + pub(super) noise_results: Option, + /// Per-ring-slot temporal residual statistics, one record per spatial block. + pub(super) temporal_stats_buf: Option, + /// Smooths the median chain's estimate, which feeds `h2_inv_norm` and `sigma_y`. + pub(super) noise_estimator: NoiseEstimator, + /// Smooths the low chain's estimate, which feeds the noise offsets. + /// + /// The low chain reads lower-quartile statistics, because reading the noise too high there + /// destroys detail. + pub(super) noise_estimator_low: NoiseEstimator, + /// Smooths the low chain's estimate without the correlation boost. + /// + /// A consumer that squares sigma into a shrinkage threshold pays for an over-read twice, so + /// it reads this instead. + pub(super) noise_estimator_low_unboosted: NoiseEstimator, + /// Smooths the temporal median alone, with no spatial maximum and no correlation boost. + /// + /// It only updates on a fold with a trustworthy temporal sample. + pub(super) noise_estimator_temporal_only: NoiseEstimator, + /// The latest centre frame's luma noise curve, if one could be built. + pub(super) noise_curve: Option, + /// The latest centre frame's quarter classes, present exactly when `noise_curve` is. + pub(super) quarter_classes: Option, + + pub(super) mc_ctx: Option, + /// The motion-shifted input ring the temporal kernels read neighbours from. + pub(super) compensated_input_buf: Option, + /// The motion-shifted reference ring. + pub(super) compensated_reference_buf: Option, + /// Motion vectors, one slice per neighbour with two `i32` components per block. + /// + /// Neighbours behind the centre fill the first half of the slices. + pub(super) mv_field_buf: Option, + /// Adjacent-frame motion fields for chained estimation, by slot, then direction, then block. + /// + /// Direction 0 runs from the older frame to the newer one. + pub(super) pair_ring_buf: Option, + /// The luma pyramid, indexed by level, then frame, then pixel. + pub(super) pyramid_input: Option, + /// The luma pyramid built from the reference ring. + pub(super) pyramid_reference: Option, + + /// Block geometry for the confidence pass that runs without motion compensation. + pub(super) confidence_ctx: Option, + /// Per-block match confidence, laid out like `mv_field_buf` with one `f32` per block. + pub(super) confidence_buf: Option, + /// The single-level luma pyramid ring for the confidence-only pass. + pub(super) confidence_pyramid: Option, + /// Discarded motion vectors from the confidence-only pass. + pub(super) confidence_mv_scratch: Option, + /// The fine block-match kernel's confidence argument when confidence weighting is off. + /// + /// The kernel drops the confidence write at compile time in that case, so this buffer is + /// never indexed. + pub(super) confidence_dummy: Handle, + /// The smoothed luma sigma, which feeds the confidence noise floor. + pub(super) sigma_y: f32, + + /// Whether the temporal stats kernel runs its four luma-only lanes. + /// + /// Leaving them off roughly halves the kernel's cost at 1080p. + pub(super) luma_noise_fields: bool, + + /// The cut the luma flat map vetoes textured quarters at, or `None` for no veto. + pub(super) flat_texture_cut: Option, + + /// Whether the stream's edges run off-centre passes instead of copied padding. + /// + /// With it on, a stream gets no leading copies and a centre with no temporal reading + /// borrows the nearest one ahead. + pub(super) shifted_edges: bool, +} + +impl NlmDenoiser { + /// Builds a new denoiser. + /// + /// # Panics + /// + /// Panics if the parameters or the frame dimensions are invalid. [Nlmeans](crate::Nlmeans) + /// checks both and returns a `Result` instead. + pub fn new(client: &ComputeClient, params: NlmParams, width: u32, height: u32) -> Self { + params + .validate() + .expect("invalid NlmParams, call params.validate() first to get this as a Result"); + validate_dimensions(width, height) + .expect("unsupported frame dimensions, call validate_dimensions first to get this as a Result"); + + let align = StorageAlign::from_client(client); + let stored_ch = params.channels.storage_count(); + let total_frames = params.total_frames(); + let pixels = (width * height) as usize; + let frame_bytes = pixels * stored_ch as usize * size_of::(); + let scalar_bytes = pixels * size_of::(); + + let input_buf = client.empty(frame_bytes * total_frames as usize); + let reference_buf = if params.prefilter.needs_reference_buf() { + let reference = client.empty(frame_bytes * total_frames as usize); + Some(reference) + } else { + None + }; + + let placeholder = client.empty(4); + let accum = client.empty(frame_bytes); + let weight_sum = client.empty(scalar_bytes); + let max_weight = client.empty(scalar_bytes); + let weight_buf = client.empty(scalar_bytes); + let raw_fwd = client.empty(scalar_bytes); + let raw_bwd = client.empty(scalar_bytes); + let tmp_hsum = client.empty(scalar_bytes); + let tmp_hsum_bwd = client.empty(scalar_bytes); + let outputs = [client.empty(frame_bytes), client.empty(frame_bytes)]; + + let h2_inv_norm = params.h2_inv_norm(); + let input_noise_offset = params.noise_offset(); + // Under `NlmSpatial` the main pass compares pilot outputs, which carry no noise floor. + let noise_offset = match params.prefilter { + PrefilterMode::NlmSpatial { .. } => 0.0, + _ => input_noise_offset, + }; + let use_separable = params.patch_radius > SEPARABLE_THRESHOLD; + let use_reference = params.prefilter.needs_reference_buf(); + + let rho_smoothed: Option = None; + let initial_offsets = build_spatial_offset_lut(params.search_radius, 0.0, noise_offset); + let initial_offset_bytes = f32::as_bytes(&initial_offsets); + let spatial_offset_lut = client.create_from_slice(initial_offset_bytes); + + let auto_noise = params.hq.is_some_and(|hq| hq.sigma_override.is_none()); + let (noise_partials, noise_results) = if auto_noise { + let partials_ring_bytes = + noise_partials_slot_stride_bytes(width, height, align) * total_frames as u64; + let n_results = (total_frames * 4) as usize; + let partials = client.empty(partials_ring_bytes as usize); + let results = client.empty(n_results * size_of::()); + (Some(partials), Some(results)) + } else { + (None, None) + }; + + // The temporal estimator needs a neighbour frame to take a difference against. + let temporal_stats_buf = if auto_noise && params.temporal_radius >= 1 { + let stats_bytes = temporal_stats_buf_bytes(width, height, stored_ch, total_frames, align); + let stats = client.empty(stats_bytes); + Some(stats) + } else { + None + }; + + let mc_ctx = if params.motion_compensation.is_active() && params.temporal_radius > 0 { + MotionCtx::new(params.motion_compensation, width, height, align) + } else { + None + }; + + let ( + compensated_input_buf, + compensated_reference_buf, + mv_field_buf, + pyramid_input, + pyramid_reference, + ) = if let Some(motion_ctx) = mc_ctx.as_ref() { + let compensated_input = client.empty(frame_bytes * total_frames as usize); + let compensated_reference = if use_reference { + let reference = client.empty(frame_bytes * total_frames as usize); + Some(reference) + } else { + None + }; + + let neighbours = (2 * params.temporal_radius) as u64; + let mv_field_bytes = neighbours * motion_ctx.mv_field_bytes_per_neighbour(); + let mv_field = client.empty(mv_field_bytes as usize); + + let pyramid_pixels = + motion::pyramid_pixels_per_frame(width, height, motion_ctx.pyramid_levels, motion_ctx.align); + let pyramid_bytes = pyramid_pixels * total_frames as usize * size_of::(); + let input_pyramid = client.empty(pyramid_bytes); + let reference_pyramid = if use_reference { + let pyramid = client.empty(pyramid_bytes); + Some(pyramid) + } else { + None + }; + + ( + Some(compensated_input), + compensated_reference, + Some(mv_field), + Some(input_pyramid), + reference_pyramid, + ) + } else { + (None, None, None, None, None) + }; + + let estimation = params + .motion_compensation + .resolved_estimation(params.temporal_radius); + let is_chained = matches!(estimation, Some(MotionEstimation::Chained { .. })); + let pair_ring_buf = if is_chained { + mc_ctx.as_ref().map(|motion_ctx| { + let pair_ring_slots = motion::pair_ring_slot_count(params.temporal_radius) as u64; + let pair_ring_bytes = pair_ring_slots * motion_ctx.pair_slot_bytes(); + client.empty(pair_ring_bytes as usize) + }) + } else { + None + }; + + // This also gates the motion-compensated path, so no submit pays for an unread + // confidence write. + let confidence_active = + params.hq.is_some_and(|hq| hq.temporal_confidence) && params.temporal_radius > 0; + + let confidence_only_active = confidence_active && mc_ctx.is_none(); + let confidence_ctx = confidence_only_active.then(|| MotionCtx::confidence_only(width, height, align)); + + let confidence_geometry = if confidence_active { + mc_ctx.as_ref().or(confidence_ctx.as_ref()) + } else { + None + }; + let confidence_buf = confidence_geometry.map(|geometry| { + let neighbours = (2 * params.temporal_radius) as u64; + let confidence_bytes = neighbours * geometry.confidence_bytes_per_neighbour(); + client.empty(confidence_bytes as usize) + }); + + let confidence_dummy = client.empty(size_of::()); + + let (confidence_pyramid, confidence_mv_scratch) = if let Some(geometry) = confidence_ctx.as_ref() { + let pyramid_pixels = + motion::pyramid_pixels_per_frame(width, height, geometry.pyramid_levels, geometry.align); + let pyramid_bytes = pyramid_pixels * total_frames as usize * size_of::(); + let mv_scratch_len = geometry.mv_slots_per_neighbour() * 2 * size_of::(); + let pyramid = client.empty(pyramid_bytes); + let mv_scratch = client.empty(mv_scratch_len); + (Some(pyramid), Some(mv_scratch)) + } else { + (None, None) + }; + + let sigma_y = params.hq.and_then(|hq| hq.sigma_override).unwrap_or(0.0); + + Self { + client: client.clone(), + params, + width, + height, + align, + ring_head: 0, + frames_loaded: 0, + real_pushes: 0, + input_buf, + reference_buf, + placeholder, + accum, + weight_sum, + max_weight, + weight_buf, + raw_fwd, + raw_bwd, + tmp_hsum, + tmp_hsum_bwd, + outputs, + next_output_slot: 0, + h2_inv_norm, + noise_offset, + input_noise_offset, + use_separable, + use_reference, + rho_smoothed, + spatial_offset_lut, + noise_partials, + noise_results, + temporal_stats_buf, + noise_estimator: NoiseEstimator::default(), + noise_estimator_low: NoiseEstimator::default(), + noise_estimator_low_unboosted: NoiseEstimator::default(), + noise_estimator_temporal_only: NoiseEstimator::default(), + noise_curve: None, + quarter_classes: None, + mc_ctx, + compensated_input_buf, + compensated_reference_buf, + mv_field_buf, + pair_ring_buf, + pyramid_input, + pyramid_reference, + confidence_ctx, + confidence_buf, + confidence_pyramid, + confidence_mv_scratch, + confidence_dummy, + sigma_y, + luma_noise_fields: false, + flat_texture_cut: None, + shifted_edges: false, + } + } + + pub(crate) fn placeholder(&self) -> &Handle { + &self.placeholder + } + + /// Ingests `planes` into the next ring slot and runs every per-frame stage on it. + pub(crate) fn push_planes( + &mut self, + planes: &[DevicePlane<'_>], + format: SampleFormat, + ) -> Result<(), anyhow::Error> { + let total_frames = self.params.total_frames() as usize; + let slot = self.ring_head % total_frames; + let pixels = self.width * self.height; + let stored_ch = self.params.channels.storage_count(); + let frame_len = pixels * stored_ch; + let target = IngestTarget { + ring: &self.input_buf, + ring_len: total_frames * frame_len as usize, + offset: slot as u32 * frame_len, + pixels, + channels: self.params.channels.count(), + stored_ch, + }; + + ingest(&self.client, planes, format, &self.placeholder, target); + self.run_post_upload_stages(slot) + } + + /// Queues the denoise kernels for the current window without reading the result back. + /// + /// Returns `Ok(None)` while the window is still filling. + pub fn denoise_submit_gpu(&mut self) -> Result, anyhow::Error> { + let total_frames = self.params.total_frames() as usize; + if self.frames_loaded < total_frames { + return Ok(None); + } + + if self.noise_results.is_some() { + self.update_noise_estimate(self.params.temporal_radius)?; + } + + self.rebuild_spatial_offset_lut(); + + let slot = self.next_output_slot; + self.next_output_slot = (slot + 1) % self.outputs.len(); + + self.run_denoise_kernels(slot)?; + + let output = GpuOutput { + handle: self.outputs[slot].clone(), + }; + Ok(Some(output)) + } + + /// How many tail frames the end-of-stream drain must emit. + /// + /// It is zero in spatial mode and for a stream with no pushes. + pub(crate) fn flush_target(&self) -> usize { + let temporal_radius = self.params.temporal_radius as usize; + if temporal_radius == 0 || self.real_pushes == 0 { + 0 + } else { + self.real_pushes.min(temporal_radius) + } + } + + pub(crate) fn real_pushes(&self) -> usize { + self.real_pushes + } + + /// Runs one end-of-stream drain step, duplicating the last frame forward and submitting. + /// + /// Returns `Ok(None)` while the duplicates still fill a window that never filled during + /// pushing. Later steps always return `Some`, so callers stop after [Self::flush_target] + /// outputs. It needs at least one pushed frame and a temporal radius above 0. + pub(crate) fn flush_step_gpu(&mut self) -> Result, anyhow::Error> { + let total_frames = self.params.total_frames() as usize; + + self.duplicate_last_frame()?; + if self.frames_loaded < total_frames { + self.frames_loaded += 1; + } + + self.denoise_submit_gpu() + } + + /// Resets the stream state so the next push starts a fresh temporal stream. + /// + /// The GPU buffers are left alone, because a fresh stream writes every slot before reading it. + pub fn reset_stream_state(&mut self) { + self.ring_head = 0; + self.frames_loaded = 0; + self.next_output_slot = 0; + self.real_pushes = 0; + self.noise_estimator.reset(); + self.noise_estimator_low.reset(); + self.noise_estimator_low_unboosted.reset(); + self.noise_estimator_temporal_only.reset(); + self.noise_curve = None; + self.quarter_classes = None; + self.rho_smoothed = None; + } +} diff --git a/av-denoise-core/src/nlmeans/denoiser/noise.rs b/av-denoise-core/src/nlmeans/denoiser/noise.rs new file mode 100644 index 0000000..fb87bd2 --- /dev/null +++ b/av-denoise-core/src/nlmeans/denoiser/noise.rs @@ -0,0 +1,451 @@ +use anyhow::Context; +use cubecl::prelude::*; + +use super::NlmDenoiser; +use crate::nlmeans::noise::{ + EMA_ALPHA, + NoiseCtx, + NoiseCurve, + QuarterClasses, + TemporalNoiseReading, + TemporalNoiseSample, + TemporalStatsCtx, + build_spatial_offset_lut, + correlation_factor, + noise_partials_slot_stride_bytes, + partials_len, + read_temporal_stats_slot, + run_noise_estimate, + run_temporal_noise_stats, + sigma_block_p25_from_partials, + sigma_from_abs_sum, + temporal_noise_reading, + zero_temporal_stats_slot, +}; +use crate::nlmeans::params::sigma_eff; +use crate::nlmeans::prefilter::PrefilterMode; + +impl NlmDenoiser { + /// Queues the Immerkær noise estimate for `slot` when automatic estimation is active. + pub(super) fn run_noise_estimate_for_slot(&self, slot: u32) -> Result<(), anyhow::Error> { + let (Some(partials_buf), Some(results_buf)) = + (self.noise_partials.as_ref(), self.noise_results.as_ref()) + else { + return Ok(()); + }; + + let stride = noise_partials_slot_stride_bytes(self.width, self.height, self.align); + let partials_slot = partials_buf.clone().offset_start((slot as u64) * stride); + + let noise_ctx = NoiseCtx { + width: self.width, + height: self.height, + channels: self.params.channels.count(), + stored_ch: self.params.channels.storage_count(), + frame_count: self.params.total_frames(), + frame: slot, + slot, + input_buf: &self.input_buf, + partials_buf: &partials_slot, + results_buf, + }; + + run_noise_estimate::(&self.client, &noise_ctx).context("noise estimate dispatch failed") + } + + /// Queues the temporal residual statistics between `slot` and the slot before it. + /// + /// A stream's first frame has no predecessor, so its record is zeroed instead. + pub(super) fn run_temporal_stats_for_slot(&self, slot: u32) -> Result<(), anyhow::Error> { + let Some(stats_buf) = self.temporal_stats_buf.as_ref() else { + return Ok(()); + }; + + if self.ring_head == 0 { + self.zero_temporal_stats_for_slot(slot); + return Ok(()); + } + + let total_frames = self.params.total_frames(); + let slot_prev = (slot + total_frames - 1) % total_frames; + + let stats_ctx = TemporalStatsCtx { + width: self.width, + height: self.height, + stored_ch: self.params.channels.storage_count(), + frame_count: total_frames, + slot_new: slot, + slot_prev, + input_buf: &self.input_buf, + stats_buf, + align: self.align, + }; + + run_temporal_noise_stats::(&self.client, &stats_ctx, self.luma_noise_fields) + .context("temporal noise stats dispatch failed") + } + + pub(crate) fn set_luma_noise_fields(&mut self, on: bool) { + self.luma_noise_fields = on; + } + + pub(crate) fn set_flat_texture_cut(&mut self, cut: Option) { + self.flat_texture_cut = cut; + } + + pub(crate) fn current_noise_curve(&self) -> Option { + self.noise_curve + } + + pub(crate) fn current_quarter_classes(&self) -> Option<&QuarterClasses> { + self.quarter_classes.as_ref() + } + + #[cfg(test)] + pub(crate) fn flat_texture_cut(&self) -> Option { + self.flat_texture_cut + } + + /// Zeroes a duplicated slot's temporal stats record. + /// + /// A duplicate matches the slot before it, so zeroes are the measured answer at less cost. + pub(super) fn zero_temporal_stats_for_slot(&self, slot: u32) { + let Some(stats_buf) = self.temporal_stats_buf.as_ref() else { + return; + }; + + zero_temporal_stats_slot::( + &self.client, + stats_buf, + self.width, + self.height, + self.params.channels.storage_count(), + slot, + self.align, + ); + } + + /// Folds the first frame's noise estimate straight away so push-time work has a real sigma. + /// + /// Push-time stages such as the NLM pilot run before the first submit folds an estimate. + /// One blocking read here keeps them off the construction-time fallback strength. + /// + /// It must run after this push queues its estimate and before `advance_ring` moves + /// `frames_loaded`, which the first-frame check reads. The first submit folds the same frame + /// again, which is harmless because it reproduces these values to within rounding. + pub(super) fn seed_noise_estimate_if_first_frame(&mut self, slot: u32) -> Result<(), anyhow::Error> { + if self.frames_loaded != 0 { + return Ok(()); + } + + let Some(results_buf) = self.noise_results.as_ref() else { + return Ok(()); + }; + + let bytes = self + .client + .read_one(results_buf.clone()) + .context("noise-estimate seed readback failed")?; + let data = f32::from_bytes(&bytes); + + // The first frame has no temporal record, so Immerkær alone seeds it. + let immerkaer_low = self + .read_noise_partials_low(slot) + .context("noise-partials seed readback failed")?; + self.fold_noise_estimate(data, slot as usize, None, immerkaer_low); + + Ok(()) + } + + /// Folds one slot's noise readings into every estimator chain and refreshes the derived values. + /// + /// Each chain takes the larger of an Immerkær reading and a temporal reading, so the temporal + /// side can raise an estimate but an unreliable shot cannot lower it. The median chain drives + /// the strength. The low chain reads lower quartiles and drives the noise offsets, where + /// reading too high scrubs fine texture. + /// + /// The unboosted chain skips the correlation boost for consumers that square sigma into a + /// threshold. The temporal-only chain also skips the spatial maximum, because a spatial mask + /// reads repeating texture as noise. + /// + /// Under windowed estimation every chain takes this fold's sample outright, and a fold with no + /// temporal sample clears the temporal-only chain and `rho_smoothed`. Keeping an older + /// window's reading would make the result, `spatial_offset_lut` included, depend on how many + /// pushes came before rather than only on the current window. + fn fold_noise_estimate( + &mut self, + data: &[f32], + slot: usize, + temporal: Option, + immerkaer_low: [f32; 3], + ) { + let channels = self.params.channels.count() as usize; + let base = slot * 4; + + let mut raw = [0.0f32; 3]; + for (c, sigma) in raw.iter_mut().enumerate().take(channels) { + *sigma = sigma_from_abs_sum(data[base + c], self.width, self.height); + } + + let mut raw_low = immerkaer_low; + let mut raw_low_unboosted = immerkaer_low; + let mut raw_temporal_only: Option<[f32; 3]> = None; + + let windowed = self.params.hq.is_some_and(|hq| hq.windowed_noise_estimation); + + if let Some(sample) = temporal { + let factor = correlation_factor(sample.rho); + for c in 0..channels { + raw[c] = raw[c].max(sample.sigma[c] * factor); + raw_low[c] = raw_low[c].max(sample.sigma_low[c] * factor); + raw_low_unboosted[c] = raw_low_unboosted[c].max(sample.sigma_low[c]); + } + + raw_temporal_only = Some(sample.sigma); + + let rho = if windowed { + sample.rho + } else { + match self.rho_smoothed { + None => sample.rho, + Some(previous) => EMA_ALPHA * sample.rho + (1.0 - EMA_ALPHA) * previous, + } + }; + self.rho_smoothed = Some(rho); + } else if windowed { + // Window-local estimation drops an earlier window's correlation rather than keeping it. + self.rho_smoothed = None; + } + + // The user's scale applies after combining and before smoothing, so everything derived + // follows it. + let sigma_scale = self.params.hq.map_or(1.0, |hq| hq.sigma_scale); + for c in 0..channels { + raw[c] *= sigma_scale; + raw_low[c] *= sigma_scale; + raw_low_unboosted[c] *= sigma_scale; + } + + if let Some(raw_temporal) = raw_temporal_only.as_mut() { + for sigma in raw_temporal.iter_mut().take(channels) { + *sigma *= sigma_scale; + } + } + + let updated = self.noise_estimator.update(&raw[..channels], windowed); + let mut smoothed = [0.0f32; 3]; + smoothed[..channels].copy_from_slice(updated); + + let updated_low = self.noise_estimator_low.update(&raw_low[..channels], windowed); + let mut smoothed_low = [0.0f32; 3]; + smoothed_low[..channels].copy_from_slice(updated_low); + + self.noise_estimator_low_unboosted + .update(&raw_low_unboosted[..channels], windowed); + + match raw_temporal_only { + Some(raw_temporal) => { + self.noise_estimator_temporal_only + .update(&raw_temporal[..channels], windowed); + }, + // Window-local estimation clears rather than coasting on an older window's reading. + None if windowed => self.noise_estimator_temporal_only.reset(), + None => {}, + } + + let effective_sigma = sigma_eff(&smoothed[..channels], self.params.channels); + self.h2_inv_norm = self.params.h2_inv_norm_with(Some(effective_sigma)); + self.input_noise_offset = self.params.noise_offset_with(Some(&smoothed_low[..channels])); + self.noise_offset = match self.params.prefilter { + PrefilterMode::NlmSpatial { .. } => 0.0, + _ => self.input_noise_offset, + }; + + // Motion estimation treats channel 0 as luma. + self.sigma_y = smoothed[0]; + } + + /// Refreshes the derived filter values from the window centre's noise estimate. + /// + /// The estimate was queued when the slot was first written, so the blocking read lands on + /// finished work. + pub(super) fn update_noise_estimate(&mut self, center_t: u32) -> Result<(), anyhow::Error> { + debug_assert!( + center_t < self.params.total_frames(), + "center_t must be a logical ring position" + ); + + let results_buf = self + .noise_results + .as_ref() + .expect("noise_results allocated when auto noise is active") + .clone(); + + let bytes = self + .client + .read_one(results_buf) + .map_err(|error| anyhow::anyhow!("noise-estimate results readback failed: {error}"))?; + let data = f32::from_bytes(&bytes); + + let center_slot = self.phys_frame(center_t as i32) as usize; + + let reading = self.borrow_reading_ahead(center_t)?; + let immerkaer_low = self.read_noise_partials_low(center_slot as u32)?; + + // The curve only changes alongside a trustworthy sample. Windowed estimation clears it on + // a fold with none. + let windowed = self.params.hq.is_some_and(|hq| hq.windowed_noise_estimation); + match (&reading.sample, reading.curve, reading.classes) { + (Some(_), curve, classes) => { + self.noise_curve = curve; + self.quarter_classes = classes; + }, + (None, _, _) if windowed => { + self.noise_curve = None; + self.quarter_classes = None; + }, + (None, _, _) => {}, + } + + self.fold_noise_estimate(data, center_slot, reading.sample, immerkaer_low); + + Ok(()) + } + + /// Reads one slot's noise partials back and reduces them to the low chain's estimate. + /// + /// The ring handle is sliced to the slot, so the transfer skips the rest of the ring. + fn read_noise_partials_low(&self, slot: u32) -> Result<[f32; 3], anyhow::Error> { + let partials_buf = self + .noise_partials + .as_ref() + .expect("noise_partials allocated when auto noise is active"); + + let slot_len_bytes = partials_len(self.width, self.height) as u64 * size_of::() as u64; + let stride = noise_partials_slot_stride_bytes(self.width, self.height, self.align); + let total_bytes = self.params.total_frames() as u64 * stride; + let start = (slot as u64) * stride; + let end_trim = total_bytes - start - slot_len_bytes; + + let sliced = partials_buf.clone().offset_start(start).offset_end(end_trim); + let bytes = self + .client + .read_one(sliced) + .map_err(|error| anyhow::anyhow!("noise partials readback failed: {error}"))?; + let data = f32::from_bytes(&bytes); + + let sigmas = + sigma_block_p25_from_partials(data, self.params.channels.count(), self.width, self.height); + Ok(sigmas) + } + + /// Reads a slot's temporal residual statistics back and combines them into one reading. + /// + /// Every part is `None` when the temporal estimator is inactive or too little of the frame + /// held still. The curve also needs `luma_noise_fields` on. + pub(in crate::nlmeans) fn read_temporal_noise( + &self, + slot: u32, + ) -> Result { + let Some(stats_buf) = self.temporal_stats_buf.as_ref() else { + return Ok(TemporalNoiseReading { + sample: None, + curve: None, + classes: None, + }); + }; + + let stored_ch = self.params.channels.storage_count(); + let channels = self.params.channels.count(); + let frame_count = self.params.total_frames(); + + let records = read_temporal_stats_slot::( + &self.client, + stats_buf, + self.width, + self.height, + stored_ch, + frame_count, + slot, + self.align, + )?; + + let reading = temporal_noise_reading( + &records, + channels, + stored_ch, + self.width, + self.height, + self.luma_noise_fields, + self.flat_texture_cut, + ); + Ok(reading) + } + + /// Rebuilds `spatial_offset_lut` from `noise_offset` and `rho_smoothed`. + pub(super) fn rebuild_spatial_offset_lut(&mut self) { + let offsets = build_spatial_offset_lut( + self.params.search_radius, + self.rho_smoothed.unwrap_or(0.0), + self.noise_offset, + ); + let offset_bytes = f32::as_bytes(&offsets); + self.spatial_offset_lut = self.client.create_from_slice(offset_bytes); + } + + /// The median chain's smoothed per-channel sigma. + /// + /// A fixed `sigma_override` is returned for every channel. Before the first estimate, and + /// on the fast path, it is zero. + pub fn current_sigmas(&self) -> [f32; 3] { + if let Some(sigma) = self.params.hq.and_then(|hq| hq.sigma_override) { + return [sigma; 3]; + } + + let channels = self.params.channels.count() as usize; + let mut sigmas = [0.0f32; 3]; + if let Some(smoothed) = self.noise_estimator.current() { + sigmas[..channels].copy_from_slice(&smoothed[..channels]); + } + + sigmas + } + + /// The low chain's smoothed per-channel sigma without the correlation boost. + /// + /// A fixed `sigma_override` is returned for every channel. Before the first estimate, and + /// on the fast path, it is zero. + pub fn current_sigmas_low_unboosted(&self) -> [f32; 3] { + if let Some(sigma) = self.params.hq.and_then(|hq| hq.sigma_override) { + return [sigma; 3]; + } + + let channels = self.params.channels.count() as usize; + let mut sigmas = [0.0f32; 3]; + if let Some(smoothed) = self.noise_estimator_low_unboosted.current() { + sigmas[..channels].copy_from_slice(&smoothed[..channels]); + } + + sigmas + } + + /// The temporal median's smoothed per-channel sigma. + /// + /// A fixed `sigma_override` is returned for every channel. Until a trustworthy temporal + /// reading lands, this falls back to [Self::current_sigmas_low_unboosted], since zero would + /// under-filter. + pub fn current_sigmas_temporal_only(&self) -> [f32; 3] { + if let Some(sigma) = self.params.hq.and_then(|hq| hq.sigma_override) { + return [sigma; 3]; + } + + let channels = self.params.channels.count() as usize; + if let Some(smoothed) = self.noise_estimator_temporal_only.current() { + let mut sigmas = [0.0f32; 3]; + sigmas[..channels].copy_from_slice(&smoothed[..channels]); + return sigmas; + } + + self.current_sigmas_low_unboosted() + } +} diff --git a/av-denoise-core/src/nlmeans/denoiser/ring.rs b/av-denoise-core/src/nlmeans/denoiser/ring.rs new file mode 100644 index 0000000..fa7c2d3 --- /dev/null +++ b/av-denoise-core/src/nlmeans/denoiser/ring.rs @@ -0,0 +1,142 @@ +use anyhow::Context; +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::NlmDenoiser; +use crate::nlmeans::kernels::gpu_copy; +use crate::nlmeans::prefilter::PrefilterMode; +use crate::nlmeans::{BLOCK_1D, MAX_GRID_1D}; + +impl NlmDenoiser { + pub(super) fn advance_ring(&mut self) { + let total_frames = self.params.total_frames() as usize; + self.ring_head += 1; + if self.frames_loaded < total_frames { + self.frames_loaded += 1; + } + + self.real_pushes += 1; + } + + /// Copies frame `src_slot` of `src` into `slot` of `dst` on the GPU. + /// + /// `dst` uses the ring layout of `input_buf` and `src` holds `src_slots` frames. Both handles + /// are bound whole, because a slot's byte offset rarely meets the GPU's + /// `min_storage_buffer_offset_alignment`. + fn copy_frame_into_slot( + &self, + dst: &Handle, + slot: usize, + src: &Handle, + src_slot: usize, + src_slots: usize, + ) { + let stored_ch = self.params.channels.storage_count(); + let frame_size = self.width * self.height * stored_ch; + let dst_slots = self.params.total_frames() as usize; + + let grid = frame_size.div_ceil(BLOCK_1D).min(MAX_GRID_1D); + let total_threads = grid * BLOCK_1D; + + unsafe { + gpu_copy::launch_unchecked::( + &self.client, + CubeCount::new_1d(grid), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(src.clone(), src_slots * frame_size as usize), + ArrayArg::from_raw_parts(dst.clone(), dst_slots * frame_size as usize), + src_slot as u32 * frame_size, + slot as u32 * frame_size, + frame_size, + total_threads, + ) + }; + } + + /// Copies the first pushed frame into the leading ring slots so the window starts balanced. + /// + /// It does nothing with shifted edges on, since those windows stop at the clip start instead. + pub(super) fn prime_leading_edge_if_first(&mut self) -> Result<(), anyhow::Error> { + if self.shifted_edges { + return Ok(()); + } + + let temporal_radius = self.params.temporal_radius as usize; + + if temporal_radius == 0 || self.frames_loaded != 1 { + return Ok(()); + } + + for _ in 0..temporal_radius { + self.duplicate_last_frame()?; + self.frames_loaded += 1; + } + + Ok(()) + } + + /// Copies the last pushed frame into the next ring slot and refreshes that slot's state. + pub(in crate::nlmeans) fn duplicate_last_frame(&mut self) -> Result<(), anyhow::Error> { + let total_frames = self.params.total_frames() as usize; + let last_slot = (self.ring_head - 1) % total_frames; + let next_slot = self.ring_head % total_frames; + + // Slots never overlap, so copying within the same buffer is safe. + let input_buf = self.input_buf.clone(); + self.copy_frame_into_slot(&input_buf, next_slot, &input_buf, last_slot, total_frames); + + // `NlmSpatial` rebuilds this slot's reference with the pilot below. + if !matches!(self.params.prefilter, PrefilterMode::NlmSpatial { .. }) + && let Some(reference_buf) = self.reference_buf.clone() + { + self.copy_frame_into_slot(&reference_buf, next_slot, &reference_buf, last_slot, total_frames); + } + + // Every per-slot buffer is refreshed so a later pass never reads an older frame's state here. + if let PrefilterMode::NlmSpatial { strength_scale } = self.params.prefilter { + self.run_nlm_spatial_pilot(next_slot as u32, strength_scale) + .context("nlm spatial pilot dispatch failed")?; + } + + self.build_pyramids_for_slot(next_slot as u32)?; + self.build_confidence_pyramid_for_slot(next_slot as u32)?; + self.run_noise_estimate_for_slot(next_slot as u32)?; + self.zero_temporal_stats_for_slot(next_slot as u32); + // This runs before `ring_head` advances so `pair_slot(0)` matches the push-time slot. + self.zero_pair_slot_for_duplicate(); + + self.ring_head += 1; + + Ok(()) + } + + /// The physical slot holding the oldest frame in the window. + fn ring_start(&self) -> u32 { + let total_frames = self.params.total_frames() as usize; + (self.ring_head % total_frames) as u32 + } + + /// The physical `input_buf` slot of a logical frame index within the window. + pub(in crate::nlmeans) fn phys_frame(&self, logical: i32) -> u32 { + let total_frames = self.params.total_frames() as i32; + let wrapped = logical.rem_euclid(total_frames); + let start = self.ring_start() as i32; + ((start + wrapped).rem_euclid(total_frames)) as u32 + } + + /// The pair-ring slot holding the motion between two neighbouring frames. + /// + /// At push time the gap index is 0 and `ring_head` has not advanced yet. At compose time the + /// gap index is measured from the window centre, which cancels how far `ring_head` has moved + /// since, so both land on the same slot. + pub(in crate::nlmeans) fn pair_slot(&self, gap_index: i32) -> u32 { + let radius = self.params.temporal_radius as i32; + debug_assert!( + radius > 0, + "pair ring is only meaningful when temporal_radius > 0" + ); + + let pair_slots = 2 * radius; + ((self.ring_head as i32 + gap_index).rem_euclid(pair_slots)) as u32 + } +} diff --git a/av-denoise-core/src/nlmeans/denoiser/sizes.rs b/av-denoise-core/src/nlmeans/denoiser/sizes.rs new file mode 100644 index 0000000..7782338 --- /dev/null +++ b/av-denoise-core/src/nlmeans/denoiser/sizes.rs @@ -0,0 +1,163 @@ +use cubecl::prelude::*; + +use crate::nlmeans::align::StorageAlign; +use crate::nlmeans::motion::{self, MotionCtx, MotionEstimation}; +use crate::nlmeans::noise::{temporal_stats_blocks, temporal_stats_record_len}; +use crate::nlmeans::params::NlmParams; +use crate::nlmeans::{BLOCK_X, BLOCK_Y}; + +/// Bytes in one element of every buffer sized here, which all hold `f32`, `i32` or `u32`. +const ELEMENT_BYTES: u64 = 4; + +/// A buffer's name and its element count, `None` when the count overflows `u64`. +pub(crate) type BufferSize = (&'static str, Option); + +/// The geometry-sized buffers of an [NlmDenoiser](super::NlmDenoiser) beyond its frame rings. +pub(crate) struct FrontSizes { + pub buffers: Vec, + /// Blocks in one motion field, when motion compensation runs. + pub motion_blocks: Option, +} + +fn times(count: Option, factor: u64) -> Option { + count.and_then(|count| count.checked_mul(factor)) +} + +/// Elements in `slots` slots of `slot_bytes` bytes each. +fn slot_ring(slot_bytes: Option, slots: u64) -> Option { + let ring_bytes = times(slot_bytes, slots); + ring_bytes.map(|bytes| bytes / ELEMENT_BYTES) +} + +fn padded_slot_bytes(elements: Option, align: StorageAlign) -> Option { + let bytes = times(elements, ELEMENT_BYTES); + bytes.and_then(|bytes| align.checked_pad_bytes(bytes)) +} + +fn block_count(geometry: &MotionCtx) -> Option { + let blocks_x = u64::from(geometry.blocks_x); + let blocks_y = u64::from(geometry.blocks_y); + + blocks_x.checked_mul(blocks_y) +} + +/// Elements in a pyramid ring of `slots` frames, padded per level like the allocation. +fn pyramid_ring(width: u32, height: u32, levels: u32, align: StorageAlign, slots: u64) -> Option { + let per_frame = motion::pyramid_pixels_per_frame(width, height, levels, align); + let per_frame = u64::try_from(per_frame).ok(); + + times(per_frame, slots) +} + +/// The element count of every buffer [NlmDenoiser::new](super::NlmDenoiser::new) sizes from the frame, +/// other than the frame rings. +/// +/// Each count repeats its allocation's formula in checked `u64`. The geometry must have passed +/// [Geometry::check_ring_fits](crate::engine::Geometry::check_ring_fits) for the frame ring, which +/// keeps the pixel count within `u32` so the shared size helpers called here cannot overflow. +pub(crate) fn front_buffer_sizes( + client: &ComputeClient, + params: &NlmParams, + width: u32, + height: u32, +) -> FrontSizes { + let align = StorageAlign::from_client(client); + let stored_ch = params.channels.storage_count(); + let total_frames = u64::from(params.total_frames()); + let neighbours = 2 * u64::from(params.temporal_radius); + let mut buffers = Vec::new(); + + let auto_noise = params.hq.is_some_and(|hq| hq.sigma_override.is_none()); + if auto_noise { + let cubes_x = u64::from(width.div_ceil(BLOCK_X)); + let cubes_y = u64::from(height.div_ceil(BLOCK_Y)); + let partials = times(cubes_x.checked_mul(cubes_y), 4); + let slot_bytes = padded_slot_bytes(partials, align); + let ring = slot_ring(slot_bytes, total_frames); + buffers.push(("noise partials ring", ring)); + } + + if auto_noise && params.temporal_radius >= 1 { + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); + let blocks = u64::from(blocks_x).checked_mul(u64::from(blocks_y)); + let record_len = u64::from(temporal_stats_record_len(stored_ch)); + let slot_len = times(blocks, record_len); + let slot_bytes = padded_slot_bytes(slot_len, align); + let ring = slot_ring(slot_bytes, total_frames); + buffers.push(("temporal stats ring", ring)); + } + + let motion_active = params.motion_compensation.is_active() && params.temporal_radius > 0; + let mc_ctx = if motion_active { + MotionCtx::new(params.motion_compensation, width, height, align) + } else { + None + }; + + let motion_blocks = mc_ctx.as_ref().and_then(block_count); + + if let Some(motion_ctx) = mc_ctx.as_ref() { + let field_len = times(motion_blocks, 2); + let field_bytes = padded_slot_bytes(field_len, align); + let field = slot_ring(field_bytes, neighbours); + buffers.push(("motion field", field)); + + let pyramid = pyramid_ring(width, height, motion_ctx.pyramid_levels, align, total_frames); + buffers.push(("motion pyramid ring", pyramid)); + } + + let estimation = params + .motion_compensation + .resolved_estimation(params.temporal_radius); + let is_chained = matches!(estimation, Some(MotionEstimation::Chained { .. })); + if is_chained && mc_ctx.is_some() { + let direction_len = times(motion_blocks, 2); + let direction_bytes = padded_slot_bytes(direction_len, align); + let slot_bytes = times(direction_bytes, 2); + let pair_slots = u64::from(motion::pair_ring_slot_count(params.temporal_radius)); + let ring = slot_ring(slot_bytes, pair_slots); + buffers.push(("motion pair ring", ring)); + } + + let confidence_active = params.hq.is_some_and(|hq| hq.temporal_confidence) && params.temporal_radius > 0; + let confidence_only_active = confidence_active && mc_ctx.is_none(); + let confidence_ctx = confidence_only_active.then(|| MotionCtx::confidence_only(width, height, align)); + + let confidence_geometry = if confidence_active { + mc_ctx.as_ref().or(confidence_ctx.as_ref()) + } else { + None + }; + if let Some(geometry) = confidence_geometry { + let blocks = block_count(geometry); + let slot_bytes = padded_slot_bytes(blocks, align); + let confidence = slot_ring(slot_bytes, neighbours); + buffers.push(("confidence", confidence)); + } + + if let Some(geometry) = confidence_ctx.as_ref() { + let pyramid = pyramid_ring(width, height, geometry.pyramid_levels, align, total_frames); + buffers.push(("confidence pyramid ring", pyramid)); + + let blocks = block_count(geometry); + let scratch = times(blocks, 2); + buffers.push(("confidence vector scratch", scratch)); + } + + FrontSizes { + buffers, + motion_blocks, + } +} + +/// Rejects the first buffer that would hold more than `u32::MAX` elements, naming it. +pub(crate) fn check_u32_indexable(buffers: &[BufferSize]) -> Result<(), String> { + for (name, elements) in buffers { + let fits = elements.is_some_and(|elements| elements <= u64::from(u32::MAX)); + if !fits { + return Err(format!("the {name} would hold more than u32::MAX elements")); + } + } + + Ok(()) +} diff --git a/av-denoise-core/src/nlmeans/denoiser/stages.rs b/av-denoise-core/src/nlmeans/denoiser/stages.rs new file mode 100644 index 0000000..2926248 --- /dev/null +++ b/av-denoise-core/src/nlmeans/denoiser/stages.rs @@ -0,0 +1,197 @@ +use anyhow::Context; +use cubecl::prelude::*; + +use super::NlmDenoiser; +use crate::nlmeans::motion::{self, MotionEstimation, build_pyramid_for_slot, run_pyramid_build}; +use crate::nlmeans::prefilter::{PrefilterCtx, PrefilterMode, run_prefilter}; + +impl NlmDenoiser { + /// Runs every per-frame stage on a freshly pushed `slot`, through to the ring advance. + pub(super) fn run_post_upload_stages(&mut self, slot: usize) -> Result<(), anyhow::Error> { + self.run_noise_estimate_for_slot(slot as u32)?; + self.run_temporal_stats_for_slot(slot as u32)?; + self.seed_noise_estimate_if_first_frame(slot as u32)?; + + if let PrefilterMode::NlmSpatial { strength_scale } = self.params.prefilter { + self.run_nlm_spatial_pilot(slot as u32, strength_scale) + .context("nlm spatial pilot dispatch failed")?; + } else if self.params.prefilter.is_gpu_internal() { + self.run_prefilter_for_slot(slot)?; + } + + self.build_pyramids_for_slot(slot as u32)?; + self.build_confidence_pyramid_for_slot(slot as u32)?; + self.run_pair_analyse_for_slot(slot as u32)?; + + self.advance_ring(); + self.prime_leading_edge_if_first() + } + + fn run_prefilter_for_slot(&self, slot: usize) -> Result<(), anyhow::Error> { + let reference_buf = self + .reference_buf + .as_ref() + .expect("reference buffer must exist for GPU prefilter"); + + let prefilter_ctx = PrefilterCtx { + width: self.width, + height: self.height, + channels: self.params.channels.count(), + stored_ch: self.params.channels.storage_count(), + frame_count: self.params.total_frames(), + frame: slot as u32, + input_buf: &self.input_buf, + reference_buf, + }; + + run_prefilter::(self.params.prefilter, &self.client, &prefilter_ctx) + .context("prefilter dispatch failed") + } + + /// Builds the motion pyramids for `slot` on the input ring and, when present, the reference ring. + pub(super) fn build_pyramids_for_slot(&self, slot: u32) -> Result<(), anyhow::Error> { + let Some(motion_ctx) = self.mc_ctx.as_ref() else { + return Ok(()); + }; + + let stored_ch = self.params.channels.storage_count(); + let frame_count = self.params.total_frames(); + + if let Some(input_pyramid) = self.pyramid_input.as_ref() { + build_pyramid_for_slot::( + &self.client, + motion_ctx, + self.width, + self.height, + frame_count, + slot, + &self.input_buf, + input_pyramid, + stored_ch, + ) + .context("input pyramid build dispatch failed")?; + } + + if let (Some(reference_pyramid), Some(reference_buf)) = + (self.pyramid_reference.as_ref(), self.reference_buf.as_ref()) + { + build_pyramid_for_slot::( + &self.client, + motion_ctx, + self.width, + self.height, + frame_count, + slot, + reference_buf, + reference_pyramid, + stored_ch, + ) + .context("reference pyramid build dispatch failed")?; + } + + Ok(()) + } + + /// Builds the luma pyramid for `slot` that the confidence-only pass reads. + /// + /// It always reads `input_buf`, even with a prefilter, to avoid a second reference pyramid. + pub(super) fn build_confidence_pyramid_for_slot(&self, slot: u32) -> Result<(), anyhow::Error> { + let (Some(confidence_ctx), Some(pyramid)) = + (self.confidence_ctx.as_ref(), self.confidence_pyramid.as_ref()) + else { + return Ok(()); + }; + + run_pyramid_build::( + &self.client, + confidence_ctx, + self.width, + self.height, + self.params.total_frames(), + slot, + &self.input_buf, + pyramid, + self.params.channels.storage_count(), + ) + .context("confidence pyramid build dispatch failed") + } + + /// Whether `Chained` motion estimation is in use, either asked for or resolved from `Auto`. + /// + /// This ignores whether motion compensation is active, which also needs a temporal radius + /// above 0. + pub(in crate::nlmeans) fn is_chained(&self) -> bool { + let estimation = self + .params + .motion_compensation + .resolved_estimation(self.params.temporal_radius); + matches!(estimation, Some(MotionEstimation::Chained { .. })) + } + + /// Measures motion both ways between the slot just written and the one before it. + /// + /// The result goes into the pair ring. A stream's first frame has no older partner and is + /// skipped, so composition reads the priming duplicate's zeroed pair instead. + fn run_pair_analyse_for_slot(&self, newer_slot: u32) -> Result<(), anyhow::Error> { + if self.ring_head == 0 { + return Ok(()); + } + + let Some(motion_ctx) = self.mc_ctx.as_ref() else { + return Ok(()); + }; + + if !self.is_chained() { + return Ok(()); + } + + let pair_ring = self + .pair_ring_buf + .as_ref() + .expect("pair_ring allocated when Chained is active"); + + // Match against the cleaner of the two pyramids. + let pyramid = self.pyramid_reference.as_ref().unwrap_or_else(|| { + self.pyramid_input + .as_ref() + .expect("pyramid_input allocated when mc_ctx is Some") + }); + + let total_frames = self.params.total_frames(); + let older_slot = (newer_slot + total_frames - 1) % total_frames; + let pair_slot = self.pair_slot(0); + + motion::run_pair_analyse::( + &self.client, + motion_ctx, + self.width, + self.height, + total_frames, + older_slot, + newer_slot, + pair_slot, + pyramid, + pair_ring, + &self.confidence_dummy, + ) + .context("pair analyse dispatch failed") + } + + /// Zeroes the pair-ring slot of a duplicated frame. + pub(super) fn zero_pair_slot_for_duplicate(&self) { + let Some(motion_ctx) = self.mc_ctx.as_ref() else { + return; + }; + + if !self.is_chained() { + return; + } + + let pair_ring = self + .pair_ring_buf + .as_ref() + .expect("pair_ring allocated when Chained is active"); + let pair_slot = self.pair_slot(0); + motion::zero_pair_slot::(&self.client, motion_ctx, pair_ring, pair_slot); + } +} diff --git a/av-denoise-core/src/nlmeans/dispatch.rs b/av-denoise-core/src/nlmeans/dispatch.rs deleted file mode 100644 index 25963fc..0000000 --- a/av-denoise-core/src/nlmeans/dispatch.rs +++ /dev/null @@ -1,1431 +0,0 @@ -use cubecl::prelude::*; -use cubecl::server::Handle; - -use super::denoiser::NlmDenoiser; -use super::kernels::{ - gpu_copy, - gpu_zero_buffers, - nlm_accumulate, - nlm_distance, - nlm_distance_pair, - nlm_distance_pair_ref, - nlm_distance_ref, - nlm_finish, - nlm_fused_pair_accumulate_window, - nlm_fused_pair_accumulate_window_ref, - nlm_fused_single_window, - nlm_fused_single_window_ref, - nlm_horizontal_sum, - nlm_horizontal_sum_pair, - nlm_vertical_weight, - nlm_vweight_pair_accumulate, -}; -use super::motion::{ - self, - MotionCtx, - MotionEstimation, - confidence_byte_offset, - run_analyse, - run_compensate, - run_confidence_for_neighbour, - run_seeded_refine, -}; -use super::noise::{build_spatial_offset_lut, spatial_offset_factor, spatial_offset_lut_len}; -use super::prefilter::PrefilterMode; -use super::{BLOCK_1D, BLOCK_X, BLOCK_X_THIN, BLOCK_Y, BLOCK_Y_THIN, MAX_GRID_1D}; - -/// The sizes and grid shape one frame's work needs, bundled together so -/// the dispatch helpers do not each carry the same long argument list. -pub(super) struct LaunchCtx { - pub(super) total_frame_data: usize, - pub(super) frame_size: usize, - pub(super) pixels: usize, - pub(super) cube_count: CubeCount, - pub(super) cube_dim: CubeDim, - /// Alternate shape used by `nlm_accumulate`. See [`BLOCK_X_THIN`]. - pub(super) thin_cube_count: CubeCount, - pub(super) thin_cube_dim: CubeDim, -} - -/// Maps a nonzero temporal offset onto the neighbour index the analyse -/// and confidence passes use when filling their buffers. -/// -/// Those passes fill the negative offsets first, taking indices 0 up to -/// the radius minus 1, then the positive ones. This is the inverse of -/// that walk. -/// -/// See `NlmDenoiser::run_motion_compensation` and `run_confidence_pass`. -fn neighbour_idx_for_k(radius: u32, k: i32) -> u32 { - debug_assert_ne!(k, 0, "k=0 is the spatial pair, it has no neighbour index"); - debug_assert!( - k.unsigned_abs() <= radius, - "k={k} outside the temporal window ±{radius}" - ); - if k < 0 { - (k + radius as i32) as u32 - } else { - (radius as i32 - 1 + k) as u32 - } -} - -/// How much of the raw input's noise the `NlmSpatial` pilot's reference -/// image still carries. -/// -/// This was measured rather than derived, by sweeping 0, 0.25, 0.5, -/// 0.75, and 1.0 against the HQ variant with the NLM prefilter and -/// motion compensation, at light noise levels on `clean-1080p.mkv`. -/// -/// The sweep came back nearly flat, with at most 0.10 dB of spread -/// across the whole grid and both metrics agreeing on direction. Its -/// peak sat at 1.0, meaning no correction at all, which is the opposite -/// of the value chosen here. -/// -/// 0 is kept anyway. A floor large enough to swamp `thsad` leaves -/// confidence unable to tell a genuine mismatch apart from noise, no -/// matter whether this one clip happens to show that failure. Because -/// the grid is flat, choosing 0 costs nothing measurable either. -const NLM_SPATIAL_RESIDUAL_FRACTION: f32 = 0.0; - -/// How much of the raw input's noise the `Bilateral` prefilter's -/// reference image still carries. -/// -/// This was measured the same way as -/// [`NLM_SPATIAL_RESIDUAL_FRACTION`], using the repo's own calibrated -/// `sigma_s` and `sigma_r` pair with motion compensation on. -/// -/// Unlike the NLM pilot, this sweep landed cleanly on 0 at every noise -/// level tested, and the margin grew with the noise, from 0.16 dB at -/// the lightest level to 0.57 dB at the heaviest. So 0 here is a real -/// result rather than a judgement call. -/// -/// It matches [`NLM_SPATIAL_RESIDUAL_FRACTION`] by coincidence of two -/// different lines of reasoning, not because the two prefilters were -/// shown to share a value. They stay separate constants for that -/// reason. -const BILATERAL_RESIDUAL_FRACTION: f32 = 0.0; - -/// The sigma to hand [`motion::sad_noise_floor`] for the -/// motion-compensation block match. -/// -/// `run_motion_compensation` matches against the reference pyramid -/// whenever one exists, not the raw input pyramid. -/// [`motion::sad_noise_floor`] models the score two raw noisy copies -/// would produce, so the raw sigma only belongs there when the match -/// really runs on raw pixels. -/// -/// A prefilter that runs on the GPU, meaning the NLM pilot pass or the -/// bilateral blur, cleans the frame before the match ever sees it. The -/// raw floor therefore overstates the real one. -/// -/// The consequences are not subtle. With the NLM pilot, the default -/// block size, and a sigma of 0.02, the raw floor alone comes to about -/// 5.78 against a threshold of 5.12. That swamps confidence's whole -/// range and pins it at 1.0 everywhere, including genuinely occluded -/// blocks. -/// -/// # What is used instead -/// -/// Rather than asserting a model of how much noise each prefilter -/// leaves behind, which would only swap one unverified guess for -/// another, the sigma used is the raw sigma times a residual fraction -/// measured per prefilter by a quality sweep. See -/// [`NLM_SPATIAL_RESIDUAL_FRACTION`] and -/// [`BILATERAL_RESIDUAL_FRACTION`]. -/// -/// 0 means the raw floor contributes nothing, and 1 is the old -/// behaviour this replaced. Those are the sweep grid's two endpoints. -/// -/// An `External` reference comes from the caller with unknown noise, and -/// is not something this crate denoised, so it keeps the raw sigma just -/// as `PrefilterMode::None` does. -pub(super) fn mc_sad_noise_floor_sigma(prefilter: PrefilterMode, sigma_y: f32) -> f32 { - match prefilter { - PrefilterMode::NlmSpatial { .. } => sigma_y * NLM_SPATIAL_RESIDUAL_FRACTION, - PrefilterMode::Bilateral { .. } => sigma_y * BILATERAL_RESIDUAL_FRACTION, - PrefilterMode::External | PrefilterMode::None => sigma_y, - } -} - -/// The confidence arguments for one temporal pair dispatch. -/// -/// This carries whether confidence weighting is on, the forward and -/// backward per-block confidence views, and the block geometry the -/// kernel needs to map an output pixel onto its block. That mapping -/// mirrors the one `nlm_mc_warp` uses. -/// -/// When confidence is off, this holds the small placeholder buffer, a -/// false flag, and geometry that is never read. -struct ConfidenceArgs { - use_confidence: bool, - conf_fwd: ArrayArg, - conf_bwd: ArrayArg, - step: u32, - blocks_x: u32, - blocks_y: u32, -} - -impl NlmDenoiser { - fn input_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.input_buf.clone(), ctx.total_frame_data) } - } - - fn reference_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - let buf = self - .reference_buf - .as_ref() - .expect("reference buffer must exist when use_reference is set"); - unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) } - } - - /// The input array the temporal kernels read. - /// - /// With motion compensation active this is the compensated ring. - /// Otherwise it is the same as [`Self::input_arg`]. - fn input_arg_for_temporal(&self, ctx: &LaunchCtx) -> ArrayArg { - match self.compensated_input_buf.as_ref() { - Some(buf) => unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) }, - None => self.input_arg(ctx), - } - } - - /// The reference array the temporal `_ref` kernels read, following - /// the same rule as [`Self::input_arg_for_temporal`]. - fn reference_arg_for_temporal(&self, ctx: &LaunchCtx) -> ArrayArg { - match self.compensated_reference_buf.as_ref() { - Some(buf) => unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) }, - None => self.reference_arg(ctx), - } - } - - fn accum_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.accum.clone(), ctx.frame_size) } - } - - fn output_arg(&self, ctx: &LaunchCtx, slot: usize) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.outputs[slot].clone(), ctx.frame_size) } - } - - /// The whole reference ring, for a kernel that picks a slot itself. - /// - /// Binding one slot instead would need its byte offset to land on - /// one of the runtime's alignment boundaries, which a - /// `width * height * stored_ch` frame stride cannot promise. - fn reference_ring_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - let buf = self - .reference_buf - .as_ref() - .expect("reference buffer must exist for the nlm spatial pilot"); - unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) } - } - - fn weight_sum_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.weight_sum.clone(), ctx.pixels) } - } - - fn max_weight_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.max_weight.clone(), ctx.pixels) } - } - - fn weight_buf_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.weight_buf.clone(), ctx.pixels) } - } - - fn raw_fwd_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.raw_fwd.clone(), ctx.pixels) } - } - - fn raw_bwd_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.raw_bwd.clone(), ctx.pixels) } - } - - fn tmp_hsum_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.tmp_hsum.clone(), ctx.pixels) } - } - - fn tmp_hsum_bwd_arg(&self, ctx: &LaunchCtx) -> ArrayArg { - unsafe { ArrayArg::from_raw_parts(self.tmp_hsum_bwd.clone(), ctx.pixels) } - } - - /// A view of `spatial_offset_lut`, sized for this denoiser's search - /// radius. - /// - /// Only the spatial windowed kernels read it. Every other launch - /// uses the flat `noise_offset` scalar instead. - fn spatial_offset_lut_arg(&self) -> ArrayArg { - let len = spatial_offset_lut_len(self.params.search_radius); - unsafe { ArrayArg::from_raw_parts(self.spatial_offset_lut.clone(), len) } - } - - /// Builds the confidence arguments for one temporal pair, at an - /// offset that is never zero at any call site. - /// - /// The forward frame reads the neighbour at a positive offset, so - /// its confidence comes from that neighbour's slice. The backward - /// frame reads the negative offset, so its confidence comes from the - /// opposite slice. See `neighbour_idx_for_k`. - /// - /// Confidence weighting only runs when `confidence_buf` was - /// allocated, which `NlmDenoiser::new` decides, and when block - /// geometry exists, either from `mc_ctx` with motion compensation on - /// or from `confidence_ctx` without it. - /// - /// Otherwise this falls back to the one-element `confidence_dummy` - /// buffer with the flag off, so the kernel never reads it. - fn confidence_pair_args(&self, q_k: i32) -> ConfidenceArgs { - let geometry = self.mc_ctx.as_ref().or(self.confidence_ctx.as_ref()); - if let (Some(buf), Some(mc)) = (self.confidence_buf.as_ref(), geometry) { - let radius = self.params.temporal_radius; - let fwd_idx = neighbour_idx_for_k(radius, q_k); - let bwd_idx = neighbour_idx_for_k(radius, -q_k); - let conf_len = (mc.blocks_x * mc.blocks_y) as usize; - let fwd_handle = buf.clone().offset_start(confidence_byte_offset(mc, fwd_idx)); - let bwd_handle = buf.clone().offset_start(confidence_byte_offset(mc, bwd_idx)); - ConfidenceArgs { - use_confidence: true, - conf_fwd: unsafe { ArrayArg::from_raw_parts(fwd_handle, conf_len) }, - conf_bwd: unsafe { ArrayArg::from_raw_parts(bwd_handle, conf_len) }, - step: mc.step, - blocks_x: mc.blocks_x, - blocks_y: mc.blocks_y, - } - } else { - ConfidenceArgs { - use_confidence: false, - conf_fwd: unsafe { ArrayArg::from_raw_parts(self.confidence_dummy.clone(), 1) }, - conf_bwd: unsafe { ArrayArg::from_raw_parts(self.confidence_dummy.clone(), 1) }, - step: 1, - blocks_x: 1, - blocks_y: 1, - } - } - } - - /// The windowed fused step for a temporal neighbour. - /// - /// One launch covers every offset in the search window, keeping the - /// accumulator, weight sum, and max weight in registers throughout. - /// That collapses `(2 * search_radius + 1)^2` launches into one. - fn dispatch_fused_window_iter( - &self, - ctx: &LaunchCtx, - center_t: u32, - q_k: i32, - ) -> Result<(), anyhow::Error> { - let channels = self.params.channels.count(); - let _stored = self.params.channels.storage_count(); - let frame_t = self.phys_frame(center_t as i32); - let frame_fwd = self.phys_frame(center_t as i32 + q_k); - let frame_bwd = self.phys_frame(center_t as i32 - q_k); - let confidence = self.confidence_pair_args(q_k); - - if self.use_reference { - unsafe { - nlm_fused_pair_accumulate_window_ref::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg_for_temporal(ctx), - self.reference_arg_for_temporal(ctx), - self.accum_arg(ctx), - self.weight_sum_arg(ctx), - self.max_weight_arg(ctx), - confidence.conf_fwd, - confidence.conf_bwd, - confidence.use_confidence, - frame_t, - frame_fwd, - frame_bwd, - self.h2_inv_norm, - self.noise_offset, - self.width, - self.height, - channels, - self.params.patch_radius, - self.params.search_radius, - BLOCK_X, - BLOCK_Y, - confidence.step, - confidence.blocks_x, - confidence.blocks_y, - ); - } - } else { - unsafe { - nlm_fused_pair_accumulate_window::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg_for_temporal(ctx), - self.accum_arg(ctx), - self.weight_sum_arg(ctx), - self.max_weight_arg(ctx), - confidence.conf_fwd, - confidence.conf_bwd, - confidence.use_confidence, - frame_t, - frame_fwd, - frame_bwd, - self.h2_inv_norm, - self.noise_offset, - self.width, - self.height, - channels, - self.params.patch_radius, - self.params.search_radius, - BLOCK_X, - BLOCK_Y, - confidence.step, - confidence.blocks_x, - confidence.blocks_y, - ); - } - } - - Ok(()) - } - - /// The windowed fused step for the frame against itself. - /// - /// One launch covers every offset in the search window, walking it - /// in a single direction. Patch distance reads the same either way, - /// so a full window in one direction gives exactly the same total as - /// a half window walked in both. - fn dispatch_fused_single_window_iter(&self, ctx: &LaunchCtx, center_t: u32) -> Result<(), anyhow::Error> { - let channels = self.params.channels.count(); - let _stored = self.params.channels.storage_count(); - let frame_t = self.phys_frame(center_t as i32); - - if self.use_reference { - unsafe { - nlm_fused_single_window_ref::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg(ctx), - self.reference_arg(ctx), - self.accum_arg(ctx), - self.weight_sum_arg(ctx), - self.max_weight_arg(ctx), - frame_t, - self.h2_inv_norm, - self.spatial_offset_lut_arg(), - self.width, - self.height, - channels, - self.params.patch_radius, - self.params.search_radius, - BLOCK_X, - BLOCK_Y, - ); - } - } else { - unsafe { - nlm_fused_single_window::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg(ctx), - self.accum_arg(ctx), - self.weight_sum_arg(ctx), - self.max_weight_arg(ctx), - frame_t, - self.h2_inv_norm, - self.spatial_offset_lut_arg(), - self.width, - self.height, - channels, - self.params.patch_radius, - self.params.search_radius, - BLOCK_X, - BLOCK_Y, - ); - } - } - - Ok(()) - } - - /// The separable-path step for a temporal neighbour. - /// - /// It runs the paired distance, then the paired horizontal sums, - /// then one fused kernel that finishes the vertical sum, the - /// weights, and the accumulation together. - /// - /// That last kernel reads both horizontal-sum buffers itself, so no - /// weight buffer is ever written to global memory. - fn dispatch_separable_iter( - &self, - ctx: &LaunchCtx, - center_t: u32, - q_x: i32, - q_y: i32, - q_k: i32, - ) -> Result<(), anyhow::Error> { - let channels = self.params.channels.count(); - let frame_t = self.phys_frame(center_t as i32); - let frame_fwd = self.phys_frame(center_t as i32 + q_k); - let frame_bwd = self.phys_frame(center_t as i32 - q_k); - - if self.use_reference { - unsafe { - nlm_distance_pair_ref::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.reference_arg_for_temporal(ctx), - self.raw_fwd_arg(ctx), - self.raw_bwd_arg(ctx), - frame_t, - frame_fwd, - frame_bwd, - q_x, - q_y, - self.width, - self.height, - channels, - ); - } - } else { - unsafe { - nlm_distance_pair::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg_for_temporal(ctx), - self.raw_fwd_arg(ctx), - self.raw_bwd_arg(ctx), - frame_t, - frame_fwd, - frame_bwd, - q_x, - q_y, - self.width, - self.height, - channels, - ); - } - } - - unsafe { - nlm_horizontal_sum_pair::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.raw_fwd_arg(ctx), - self.raw_bwd_arg(ctx), - self.tmp_hsum_arg(ctx), - self.tmp_hsum_bwd_arg(ctx), - self.width, - self.height, - self.params.patch_radius, - BLOCK_X, - BLOCK_Y, - ); - } - - let confidence = self.confidence_pair_args(q_k); - unsafe { - nlm_vweight_pair_accumulate::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.tmp_hsum_arg(ctx), - self.tmp_hsum_bwd_arg(ctx), - self.input_arg_for_temporal(ctx), - self.accum_arg(ctx), - self.weight_sum_arg(ctx), - self.max_weight_arg(ctx), - confidence.conf_fwd, - confidence.conf_bwd, - confidence.use_confidence, - frame_fwd, - frame_bwd, - q_x, - q_y, - self.h2_inv_norm, - self.noise_offset, - self.width, - self.height, - self.params.patch_radius, - BLOCK_X, - BLOCK_Y, - confidence.step, - confidence.blocks_x, - confidence.blocks_y, - ); - } - - Ok(()) - } - - /// The separable-path step for the frame against itself. - /// - /// It runs the distance, the horizontal sums, the vertical sums and - /// weights, and then the accumulation. - /// - /// The weight map reads the same in either direction, so one buffer - /// serves both the forward and backward lookups. - fn dispatch_separable_iter_k0( - &self, - ctx: &LaunchCtx, - center_t: u32, - q_x: i32, - q_y: i32, - ) -> Result<(), anyhow::Error> { - let channels = self.params.channels.count(); - let frame_t = self.phys_frame(center_t as i32); - - if self.use_reference { - unsafe { - nlm_distance_ref::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.reference_arg(ctx), - self.raw_fwd_arg(ctx), - frame_t, - frame_t, - q_x, - q_y, - self.width, - self.height, - channels, - ); - } - } else { - unsafe { - nlm_distance::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg(ctx), - self.raw_fwd_arg(ctx), - frame_t, - frame_t, - q_x, - q_y, - self.width, - self.height, - channels, - ); - } - } - - unsafe { - nlm_horizontal_sum::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.raw_fwd_arg(ctx), - self.tmp_hsum_arg(ctx), - self.width, - self.height, - self.params.patch_radius, - BLOCK_X, - BLOCK_Y, - ); - } - - // The same correlation-adjusted offset the spatial windowed - // kernel's table holds for this candidate. It is computed - // directly rather than read back from the table, because this - // dispatch already knows the exact candidate offset the - // windowed kernel works out at compile time. - let offset = self.noise_offset * spatial_offset_factor(q_x, q_y, self.rho_smoothed.unwrap_or(0.0)); - unsafe { - nlm_vertical_weight::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.tmp_hsum_arg(ctx), - self.weight_buf_arg(ctx), - self.h2_inv_norm, - offset, - self.width, - self.height, - self.params.patch_radius, - BLOCK_X, - BLOCK_Y, - ); - } - - unsafe { - nlm_accumulate::launch_unchecked::( - &self.client, - ctx.thin_cube_count.clone(), - ctx.thin_cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg(ctx), - self.accum_arg(ctx), - self.weight_sum_arg(ctx), - self.weight_buf_arg(ctx), - self.weight_buf_arg(ctx), - self.max_weight_arg(ctx), - frame_t, - frame_t, - q_x, - q_y, - self.width, - self.height, - ); - } - - Ok(()) - } - - /// Works out how the neighbour at temporal offset `k` moved relative - /// to the centre frame, and writes the result into `mv_field` at - /// this neighbour's slot. - /// - /// Runs the chained compose-and-refine sequence when `Chained` - /// estimation is active, or a direct coarse-to-fine match otherwise. - /// `write_confidence` controls whether a per-block score also lands - /// in `confidence_arg`, the same as `run_analyse` and - /// `run_seeded_refine` document. - /// - /// Returns the neighbour's physical ring slot. This never shifts any - /// buffer, so both [`Self::run_motion_compensation`], which follows - /// it with `run_compensate`, and [`Self::run_motion_machinery`], - /// which does not, can share it without either duplicating the - /// estimate branch or paying for a shift neither one of them wants - /// in the other's place. - #[expect( - clippy::too_many_arguments, - reason = "the dispatch threads through every buffer and shape the kernel binds" - )] - fn run_motion_estimate( - &self, - mc: &MotionCtx, - analyse_pyramid: &Handle, - mv_field: &Handle, - confidence_arg: &Handle, - write_confidence: bool, - frame_count: u32, - centre_slot: u32, - center_t: u32, - k: i32, - neighbour_idx: u32, - sad_noise_floor: f32, - thsad: f32, - ) -> Result { - let neighbour_slot = self.phys_frame(center_t as i32 + k); - - if self.is_chained() { - self.run_chain_compose(center_t, k, neighbour_idx)?; - let refine_radius = match self - .params - .motion_compensation - .resolved_estimation(self.params.temporal_radius) - { - Some(MotionEstimation::Chained { refine_radius }) => refine_radius, - _ => unreachable!("is_chained() guarantees a resolved Chained estimation"), - }; - - run_seeded_refine::( - &self.client, - mc, - self.width, - self.height, - frame_count, - centre_slot, - neighbour_slot, - neighbour_idx, - refine_radius, - analyse_pyramid, - mv_field, - confidence_arg, - write_confidence, - sad_noise_floor, - thsad, - )?; - } else { - run_analyse::( - &self.client, - mc, - self.width, - self.height, - frame_count, - centre_slot, - neighbour_slot, - neighbour_idx, - analyse_pyramid, - mv_field, - confidence_arg, - write_confidence, - sad_noise_floor, - thsad, - )?; - } - - Ok(neighbour_slot) - } - - /// Runs the same per-neighbour motion estimate - /// [`Self::run_motion_compensation`] does, for every neighbour in - /// the temporal window, but shifts nothing into the compensated - /// buffers. - /// - /// Returns each neighbour's physical ring slot, in logical ring - /// order around `center_t`, skipping the centre itself. For - /// `center_t = temporal_radius` this runs from the furthest-behind - /// neighbour to the furthest-ahead, which is also the order - /// `NlmDenoiser::submit_machinery` hands back through - /// `RingView::neighbour_slots`. - /// - /// Returns an empty `Vec` when motion compensation is off, or when - /// there are no neighbours. - pub(super) fn run_motion_machinery(&self, center_t: u32) -> Result, anyhow::Error> { - let Some(mc) = self.mc_ctx.as_ref() else { - return Ok(Vec::new()); - }; - - let temporal_radius = self.params.temporal_radius; - if temporal_radius == 0 { - return Ok(Vec::new()); - } - - let frame_count = self.params.total_frames(); - let centre_slot = self.phys_frame(center_t as i32); - - let pyramid_input = self - .pyramid_input - .as_ref() - .expect("pyramid_input allocated when mc_ctx is Some"); - let mv_field = self - .mv_field_buf - .as_ref() - .expect("mv_field allocated when mc_ctx is Some"); - // Same placeholder convention `run_motion_compensation` uses: - // when confidence weighting is off, pass the small dummy buffer - // and tell the kernel not to write it. - let (confidence_arg, write_confidence): (&Handle, bool) = match self.confidence_buf.as_ref() { - Some(buf) => (buf, true), - None => (&self.confidence_dummy, false), - }; - let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); - let mc_sigma_y = mc_sad_noise_floor_sigma(self.params.prefilter, self.sigma_y); - let sad_noise_floor = motion::sad_noise_floor(mc.blksize, mc_sigma_y); - let thsad = motion::thsad(mc.blksize, thsad_scale); - - // Match against the cleaner of the two buffers, the same as - // `run_motion_compensation`. - let analyse_pyramid = self.pyramid_reference.as_ref().unwrap_or(pyramid_input); - - let mut neighbour_idx: u32 = 0; - let mut slots = Vec::with_capacity((frame_count - 1) as usize); - for logical in 0..frame_count { - if logical == center_t { - continue; - } - - let k = logical as i32 - center_t as i32; - let neighbour_slot = self.run_motion_estimate( - mc, - analyse_pyramid, - mv_field, - confidence_arg, - write_confidence, - frame_count, - centre_slot, - center_t, - k, - neighbour_idx, - sad_noise_floor, - thsad, - )?; - slots.push(neighbour_slot); - - neighbour_idx += 1; - } - - Ok(slots) - } - - /// Runs the motion-compensation sweep for one submit. - /// - /// It estimates the motion from the centre frame to each neighbour - /// and shifts them into the compensated buffers. - /// - /// The centre slot is copied through unchanged, so the temporal - /// kernels can read every slot the same way. - /// - /// This does nothing when motion compensation is off, or when there - /// are no neighbours. - fn run_motion_compensation(&self, center_t: u32) -> Result<(), anyhow::Error> { - let Some(mc) = self.mc_ctx.as_ref() else { - return Ok(()); - }; - - let temporal_radius = self.params.temporal_radius; - if temporal_radius == 0 { - return Ok(()); - } - - let frame_count = self.params.total_frames(); - let centre_slot = self.phys_frame(center_t as i32); - let stored_ch = self.params.channels.storage_count(); - - let pyramid_input = self - .pyramid_input - .as_ref() - .expect("pyramid_input allocated when mc_ctx is Some"); - let mv_field = self - .mv_field_buf - .as_ref() - .expect("mv_field allocated when mc_ctx is Some"); - let compensated_input = self - .compensated_input_buf - .as_ref() - .expect("compensated_input allocated when mc_ctx is Some"); - // `confidence_buf` only exists when confidence weighting is on, - // which `NlmDenoiser::new` decides. Motion compensation itself - // does not need it. - // - // The fine kernel still wants some buffer for its `confidence` - // argument, so when there is none, pass the small always-present - // dummy and tell the kernel not to write it. - let (confidence_arg, write_confidence): (&Handle, bool) = match self.confidence_buf.as_ref() { - Some(buf) => (buf, true), - None => (&self.confidence_dummy, false), - }; - let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); - let mc_sigma_y = mc_sad_noise_floor_sigma(self.params.prefilter, self.sigma_y); - let sad_noise_floor = motion::sad_noise_floor(mc.blksize, mc_sigma_y); - let thsad = motion::thsad(mc.blksize, thsad_scale); - - // The centre frame is copied straight through, so the temporal - // kernels can read every slot from the compensated buffer. - copy_frame_into_slot_handle::( - &self.client, - &self.input_buf, - compensated_input, - centre_slot as usize, - self.params.total_frames(), - self.width, - self.height, - stored_ch, - ); - if let (Some(ref_src), Some(ref_dst)) = ( - self.reference_buf.as_ref(), - self.compensated_reference_buf.as_ref(), - ) { - copy_frame_into_slot_handle::( - &self.client, - ref_src, - ref_dst, - centre_slot as usize, - self.params.total_frames(), - self.width, - self.height, - stored_ch, - ); - } - - // Match against the cleaner of the two buffers, which is the - // prefiltered reference pyramid whenever one exists. - let analyse_pyramid = self.pyramid_reference.as_ref().unwrap_or(pyramid_input); - - // One motion estimate and one shift per neighbour. They run in - // order from the furthest behind to the furthest ahead, skipping - // the centre, and their motion-field indices are contiguous so - // the field stays tight. - let radius = temporal_radius as i32; - let mut neighbour_idx: u32 = 0; - for k in -radius..=radius { - if k == 0 { - continue; - } - - let neighbour_slot = self.run_motion_estimate( - mc, - analyse_pyramid, - mv_field, - confidence_arg, - write_confidence, - frame_count, - centre_slot, - center_t, - k, - neighbour_idx, - sad_noise_floor, - thsad, - )?; - - run_compensate::( - &self.client, - mc, - self.params.channels.count(), - stored_ch, - self.width, - self.height, - frame_count, - neighbour_slot, - neighbour_idx, - &self.input_buf, - compensated_input, - mv_field, - )?; - - if let (Some(ref_src), Some(ref_dst)) = ( - self.reference_buf.as_ref(), - self.compensated_reference_buf.as_ref(), - ) { - run_compensate::( - &self.client, - mc, - self.params.channels.count(), - stored_ch, - self.width, - self.height, - frame_count, - neighbour_slot, - neighbour_idx, - ref_src, - ref_dst, - mv_field, - )?; - } - - neighbour_idx += 1; - } - - Ok(()) - } - - /// Scores each neighbour against the centre frame without searching - /// for motion, once per submit. - /// - /// Every block is matched where it stands, and the per-block score - /// goes into `confidence_buf`. - /// - /// This only runs when confidence is on and motion compensation is - /// off. `confidence_ctx` is `None` in every other case. - fn run_confidence_pass(&self, center_t: u32) -> Result<(), anyhow::Error> { - let Some(ctx) = self.confidence_ctx.as_ref() else { - return Ok(()); - }; - - // `confidence_ctx` is only ever built when the temporal radius - // is above 0, which `NlmDenoiser::new` sees to, so the radius - // cannot be zero here. - let temporal_radius = self.params.temporal_radius; - - let frame_count = self.params.total_frames(); - let centre_slot = self.phys_frame(center_t as i32); - - let luma_pyramid = self - .confidence_pyramid - .as_ref() - .expect("confidence_pyramid allocated when confidence_ctx is Some"); - let mv_scratch = self - .confidence_mv_scratch - .as_ref() - .expect("confidence_mv_scratch allocated when confidence_ctx is Some"); - let confidence_buf = self - .confidence_buf - .as_ref() - .expect("confidence_buf allocated when confidence_ctx is Some"); - - let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); - let sad_noise_floor = motion::sad_noise_floor(ctx.blksize, self.sigma_y); - let thsad = motion::thsad(ctx.blksize, thsad_scale); - - let radius = temporal_radius as i32; - let mut neighbour_idx: u32 = 0; - for k in -radius..=radius { - if k == 0 { - continue; - } - let neighbour_slot = self.phys_frame(center_t as i32 + k); - - run_confidence_for_neighbour::( - &self.client, - ctx, - self.width, - self.height, - frame_count, - centre_slot, - neighbour_slot, - neighbour_idx, - luma_pyramid, - mv_scratch, - confidence_buf, - sad_noise_floor, - thsad, - )?; - - neighbour_idx += 1; - } - - Ok(()) - } - - fn zero_accumulators(&self, ctx: &LaunchCtx) -> Result<(), anyhow::Error> { - let grid = (ctx.frame_size as u32).div_ceil(BLOCK_1D).min(MAX_GRID_1D); - let total_threads = grid * BLOCK_1D; - unsafe { - gpu_zero_buffers::launch_unchecked::( - &self.client, - CubeCount::new_1d(grid), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(self.accum.clone(), ctx.frame_size), - self.weight_sum_arg(ctx), - self.max_weight_arg(ctx), - ctx.frame_size as u32, - ctx.pixels as u32, - total_threads, - ); - } - - Ok(()) - } - - /// The shared `nlm_finish` launch, with the destination left to the - /// caller. - /// - /// `run_finish` writes to an output slot, while the NLM pilot writes - /// to a reference-ring slot instead. - fn run_finish_to( - &self, - ctx: &LaunchCtx, - center_frame: u32, - output_frame: u32, - output: ArrayArg, - ) -> Result<(), anyhow::Error> { - let channels = self.params.channels.count(); - unsafe { - nlm_finish::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg(ctx), - output, - ArrayArg::from_raw_parts(self.accum.clone(), ctx.frame_size), - self.weight_sum_arg(ctx), - self.max_weight_arg(ctx), - center_frame, - output_frame, - self.params.self_weight, - self.width, - self.height, - channels, - ); - } - - Ok(()) - } - - fn run_finish(&self, ctx: &LaunchCtx, center_t: u32, output_slot: usize) -> Result<(), anyhow::Error> { - self.run_finish_to( - ctx, - self.phys_frame(center_t as i32), - 0, - self.output_arg(ctx, output_slot), - ) - } - - /// Works out the launch shapes every per-frame dispatch shares, - /// covering both the main pass and the NLM pilot. - fn launch_ctx(&self) -> LaunchCtx { - let width = self.width; - let height = self.height; - let stored_ch = self.params.channels.storage_count(); - let total_frames = self.params.total_frames(); - let pixels = (width * height) as usize; - let frame_size = pixels * stored_ch as usize; - - LaunchCtx { - total_frame_data: frame_size * total_frames as usize, - frame_size, - pixels, - cube_count: CubeCount::new_2d(width.div_ceil(BLOCK_X), height.div_ceil(BLOCK_Y)), - cube_dim: CubeDim::new_2d(BLOCK_X, BLOCK_Y), - thin_cube_count: CubeCount::new_2d(width.div_ceil(BLOCK_X_THIN), height.div_ceil(BLOCK_Y_THIN)), - thin_cube_dim: CubeDim::new_2d(BLOCK_X_THIN, BLOCK_Y_THIN), - } - } - - /// Denoises a freshly pushed frame with the windowed spatial kernel - /// and stores the result in its reference-ring slot. - /// - /// This shares the frame-sized accumulators with the main pass. The - /// GPU queue runs in order and the main dispatch zeroes them again - /// before use, so sharing them is safe. - pub(super) fn run_nlm_spatial_pilot(&self, slot: u32, strength_scale: f32) -> Result<(), anyhow::Error> { - let ctx = self.launch_ctx(); - self.zero_accumulators(&ctx)?; - - let channels = self.params.channels.count(); - let pilot_h2 = self.h2_inv_norm / (strength_scale * strength_scale); - - // A flat table with no correlation adjustment. The pilot - // compares noisy input patches directly and keeps the full - // white-noise floor, which puts it outside the scope of that - // adjustment. The denoiser's `input_noise_offset` doc explains - // why. - // - // It is built fresh each call rather than cached, because - // `input_noise_offset` can change between pushes and this is a - // tiny one-off upload. - let pilot_lut = build_spatial_offset_lut(self.params.search_radius, 0.0, self.input_noise_offset); - let pilot_lut_handle = self.client.create_from_slice(f32::as_bytes(&pilot_lut)); - let pilot_lut_arg = unsafe { ArrayArg::::from_raw_parts(pilot_lut_handle, pilot_lut.len()) }; - - // Always read the noisy input here, never `reference_arg`. For - // `NlmSpatial` the reference buffer is the pilot's own output - // rather than an input to it, even though `use_reference` is - // true so the main pass picks the `_ref` kernels. - unsafe { - nlm_fused_single_window::launch_unchecked::( - &self.client, - ctx.cube_count.clone(), - ctx.cube_dim, - self.params.channels.storage_count() as usize, - self.input_arg(&ctx), - self.accum_arg(&ctx), - self.weight_sum_arg(&ctx), - self.max_weight_arg(&ctx), - slot, - pilot_h2, - pilot_lut_arg, - self.width, - self.height, - channels, - self.params.patch_radius, - self.params.search_radius, - BLOCK_X, - BLOCK_Y, - ); - } - - self.run_finish_to(&ctx, slot, slot, self.reference_ring_arg(&ctx)) - } - - pub(super) fn run_denoise_kernels(&mut self, output_slot: usize) -> Result<(), anyhow::Error> { - let temporal_radius = self.params.temporal_radius; - let search_radius = self.params.search_radius as i32; - - let ctx = self.launch_ctx(); - - let center_t = temporal_radius; - - // Motion compensation runs before any NLM dispatch, so the - // temporal kernels can read already-aligned neighbours from the - // compensated buffers. It does nothing when motion compensation - // is off. - self.run_motion_compensation(center_t)?; - // This does nothing unless the confidence pass without motion - // compensation is active. The two never both run, because - // `confidence_ctx` is only `Some` when motion compensation is - // off. - self.run_confidence_pass(center_t)?; - - self.zero_accumulators(&ctx)?; - let window_side = 2 * search_radius + 1; - let window_area = window_side * window_side; - - // Every temporal neighbour covers the full search window, so the - // plain fused path uses the windowed kernel, one launch per - // neighbour that loops over the offsets itself. - // - // The frame against itself still dispatches one offset at a - // time over half the window, because its weight map reads the - // same in either direction and the single-tile path is cheaper - // per offset. - // - // The reference-image and separable paths also go one offset at - // a time, until they gain windowed kernels of their own. - let k_start = -(temporal_radius as i32); - let use_windowed = !self.use_separable; - for q_k in k_start..=0 { - if use_windowed { - if q_k != 0 { - self.dispatch_fused_window_iter(&ctx, center_t, q_k)?; - } else { - self.dispatch_fused_single_window_iter(&ctx, center_t)?; - } - continue; - } - - for q_y in -search_radius..=search_radius { - for q_x in -search_radius..=search_radius { - let linear = q_k * window_area + q_y * window_side + q_x; - if linear >= 0 { - continue; - } - - if q_k == 0 { - self.dispatch_separable_iter_k0(&ctx, center_t, q_x, q_y)?; - } else { - self.dispatch_separable_iter(&ctx, center_t, q_x, q_y, q_k)?; - } - } - } - } - - self.run_finish(&ctx, center_t, output_slot)?; - - Ok(()) - } -} - -/// Copies one frame from a slot of `src` into the same slot of `dst`, -/// entirely on the GPU. -/// -/// Both buffers have to share the same ring layout. -/// -/// This is a free function rather than a method, so the -/// motion-compensation dispatcher can call it without borrowing the -/// denoiser again inside the per-submit method. -/// -/// Both rings are bound whole and the kernel picks the slot itself, for -/// the alignment reason [`NlmDenoiser::copy_frame_into_slot`] explains. -#[expect( - clippy::too_many_arguments, - reason = "the dispatch threads through every buffer and shape the kernel binds" -)] -fn copy_frame_into_slot_handle( - client: &ComputeClient, - src: &Handle, - dst: &Handle, - slot: usize, - frame_count: u32, - width: u32, - height: u32, - stored_ch: u32, -) { - let frame_size = width * height * stored_ch; - let ring_len = frame_count as usize * frame_size as usize; - let offset = slot as u32 * frame_size; - - let grid = frame_size.div_ceil(BLOCK_1D).min(MAX_GRID_1D); - let total_threads = grid * BLOCK_1D; - - unsafe { - gpu_copy::launch_unchecked::( - client, - CubeCount::new_1d(grid), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(src.clone(), ring_len), - ArrayArg::from_raw_parts(dst.clone(), ring_len), - offset, - offset, - frame_size, - total_threads, - ); - } -} - -#[cfg(test)] -mod tests { - use super::{ - BILATERAL_RESIDUAL_FRACTION, - NLM_SPATIAL_RESIDUAL_FRACTION, - PrefilterMode, - mc_sad_noise_floor_sigma, - neighbour_idx_for_k, - }; - - /// Walking every offset in order, skipping zero, and counting up - /// from 0 has to reproduce `neighbour_idx_for_k` exactly at every - /// radius. - /// - /// That is the same walk `run_motion_compensation` and - /// `run_confidence_pass` use to fill their buffers, so this pins - /// down the order those passes established. - #[test] - fn matches_the_sequential_fill_order() { - for radius in 1..=8u32 { - let mut expected = 0u32; - for k in -(radius as i32)..=(radius as i32) { - if k == 0 { - continue; - } - assert_eq!(neighbour_idx_for_k(radius, k), expected, "radius={radius} k={k}"); - expected += 1; - } - } - } - - /// The forward and backward slices must always land on different - /// indices, both in range. - /// - /// Getting that wrong would quietly apply one frame's confidence to - /// another frame's temporal weight. - #[test] - fn forward_and_backward_indices_are_distinct_and_in_range() { - for radius in 1..=8u32 { - for q_k in -(radius as i32)..0 { - let fwd = neighbour_idx_for_k(radius, q_k); - let bwd = neighbour_idx_for_k(radius, -q_k); - assert_ne!(fwd, bwd, "radius={radius} q_k={q_k}"); - assert!(fwd < 2 * radius, "radius={radius} q_k={q_k} fwd={fwd}"); - assert!(bwd < 2 * radius, "radius={radius} q_k={q_k} bwd={bwd}"); - } - } - } - - #[test] - fn radius_two_explicit_indices() { - assert_eq!(neighbour_idx_for_k(2, -2), 0); - assert_eq!(neighbour_idx_for_k(2, -1), 1); - assert_eq!(neighbour_idx_for_k(2, 1), 2); - assert_eq!(neighbour_idx_for_k(2, 2), 3); - } - - /// Pins the calibrated constant to a literal rather than to a value - /// worked out from the constant itself. - /// - /// That way a future recalibration fails this test instead of - /// quietly passing again. - #[test] - fn nlm_spatial_residual_fraction_is_calibrated_to_zero() { - assert_eq!(NLM_SPATIAL_RESIDUAL_FRACTION, 0.0); - } - - #[test] - fn bilateral_residual_fraction_is_calibrated_to_zero() { - assert_eq!(BILATERAL_RESIDUAL_FRACTION, 0.0); - } - - #[test] - fn mc_sad_noise_floor_sigma_scales_nlm_spatial_by_the_calibrated_fraction() { - let raw = 0.02f32; - assert_eq!( - mc_sad_noise_floor_sigma(PrefilterMode::NlmSpatial { strength_scale: 1.0 }, raw), - raw * NLM_SPATIAL_RESIDUAL_FRACTION - ); - } - - #[test] - fn mc_sad_noise_floor_sigma_scales_bilateral_by_the_calibrated_fraction() { - let raw = 0.02f32; - assert_eq!( - mc_sad_noise_floor_sigma( - PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: 0.02 - }, - raw - ), - raw * BILATERAL_RESIDUAL_FRACTION - ); - } - - #[test] - fn mc_sad_noise_floor_sigma_keeps_raw_sigma_for_none_and_external() { - let raw = 0.02f32; - assert_eq!(mc_sad_noise_floor_sigma(PrefilterMode::None, raw), raw); - assert_eq!(mc_sad_noise_floor_sigma(PrefilterMode::External, raw), raw); - } -} diff --git a/av-denoise-core/src/nlmeans/dispatch/mod.rs b/av-denoise-core/src/nlmeans/dispatch/mod.rs new file mode 100644 index 0000000..1082d22 --- /dev/null +++ b/av-denoise-core/src/nlmeans/dispatch/mod.rs @@ -0,0 +1,733 @@ +mod motion; + +use cubecl::prelude::*; + +pub(super) use self::motion::mc_sad_noise_floor_sigma; +#[cfg(all(test, any(feature = "vulkan", feature = "metal")))] +pub(super) use self::motion::{BILATERAL_RESIDUAL_FRACTION, NLM_SPATIAL_RESIDUAL_FRACTION}; +use super::denoiser::NlmDenoiser; +use super::kernels::{ + gpu_zero_buffers, + nlm_accumulate, + nlm_distance, + nlm_distance_pair, + nlm_distance_pair_ref, + nlm_distance_ref, + nlm_finish, + nlm_fused_pair_accumulate_window, + nlm_fused_pair_accumulate_window_ref, + nlm_fused_single_window, + nlm_fused_single_window_ref, + nlm_horizontal_sum, + nlm_horizontal_sum_pair, + nlm_vertical_weight, + nlm_vweight_pair_accumulate, +}; +use super::motion::{confidence_byte_offset, neighbour_idx_for_k}; +use super::noise::{build_spatial_offset_lut, spatial_offset_factor, spatial_offset_lut_len}; +use super::{BLOCK_1D, BLOCK_X, BLOCK_X_THIN, BLOCK_Y, BLOCK_Y_THIN, MAX_GRID_1D}; + +/// The sizes and launch shapes one frame's dispatches share. +pub(super) struct LaunchCtx { + pub(super) total_frame_data: usize, + pub(super) frame_size: usize, + pub(super) pixels: usize, + pub(super) cube_count: CubeCount, + pub(super) cube_dim: CubeDim, + /// The cube count for `nlm_accumulate`. See [BLOCK_X_THIN]. + pub(super) thin_cube_count: CubeCount, + pub(super) thin_cube_dim: CubeDim, +} + +/// The confidence arguments for one temporal pair dispatch. +/// +/// The block geometry must map an output pixel onto its block exactly as `nlm_mc_warp` does. With +/// confidence off this holds the placeholder buffer, a false flag and geometry the kernel never +/// reads. +struct ConfidenceArgs { + use_confidence: bool, + conf_fwd: ArrayArg, + conf_bwd: ArrayArg, + step: u32, + blocks_x: u32, + blocks_y: u32, +} + +impl NlmDenoiser { + fn input_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.input_buf.clone(), ctx.total_frame_data) } + } + + fn reference_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + let buf = self + .reference_buf + .as_ref() + .expect("reference buffer must exist when use_reference is set"); + unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) } + } + + /// The input ring the temporal kernels read, which is the compensated one under motion + /// compensation. + fn input_arg_for_temporal(&self, ctx: &LaunchCtx) -> ArrayArg { + match self.compensated_input_buf.as_ref() { + Some(buf) => unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) }, + None => self.input_arg(ctx), + } + } + + /// The reference ring the temporal `_ref` kernels read, which is the compensated one under + /// motion compensation. + fn reference_arg_for_temporal(&self, ctx: &LaunchCtx) -> ArrayArg { + match self.compensated_reference_buf.as_ref() { + Some(buf) => unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) }, + None => self.reference_arg(ctx), + } + } + + fn accum_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.accum.clone(), ctx.frame_size) } + } + + fn output_arg(&self, ctx: &LaunchCtx, slot: usize) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.outputs[slot].clone(), ctx.frame_size) } + } + + /// The whole reference ring, for a kernel that picks its slot itself. + /// + /// Binding one slot would need its byte offset to meet the runtime's alignment, which a + /// `width * height * stored_ch` frame stride cannot promise. + fn reference_ring_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + let buf = self + .reference_buf + .as_ref() + .expect("reference buffer must exist for the nlm spatial pilot"); + unsafe { ArrayArg::from_raw_parts(buf.clone(), ctx.total_frame_data) } + } + + fn weight_sum_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.weight_sum.clone(), ctx.pixels) } + } + + fn max_weight_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.max_weight.clone(), ctx.pixels) } + } + + fn weight_buf_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.weight_buf.clone(), ctx.pixels) } + } + + fn raw_fwd_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.raw_fwd.clone(), ctx.pixels) } + } + + fn raw_bwd_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.raw_bwd.clone(), ctx.pixels) } + } + + fn tmp_hsum_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.tmp_hsum.clone(), ctx.pixels) } + } + + fn tmp_hsum_bwd_arg(&self, ctx: &LaunchCtx) -> ArrayArg { + unsafe { ArrayArg::from_raw_parts(self.tmp_hsum_bwd.clone(), ctx.pixels) } + } + + /// The spatial offset table, sized for this denoiser's search radius. + fn spatial_offset_lut_arg(&self) -> ArrayArg { + let len = spatial_offset_lut_len(self.params.search_radius); + unsafe { ArrayArg::from_raw_parts(self.spatial_offset_lut.clone(), len) } + } + + /// Builds the confidence arguments for the temporal pair at the nonzero offset `q_k`. + /// + /// The forward frame reads the neighbour at `+q_k` and the backward frame the one at `-q_k`, + /// so each takes its confidence from that neighbour's slice. Weighting runs only when + /// `confidence_buf` exists and block geometry is available. Otherwise the flag is off and the + /// kernel never reads the placeholder. + fn confidence_pair_args(&self, q_k: i32) -> ConfidenceArgs { + let confidence_geometry = self.confidence_ctx.as_ref(); + let geometry = self.mc_ctx.as_ref().or(confidence_geometry); + if let (Some(confidence_buf), Some(geometry)) = (self.confidence_buf.as_ref(), geometry) { + let radius = self.params.temporal_radius; + let forward_idx = neighbour_idx_for_k(radius, q_k); + let backward_idx = neighbour_idx_for_k(radius, -q_k); + let conf_len = (geometry.blocks_x * geometry.blocks_y) as usize; + let forward_offset = confidence_byte_offset(geometry, forward_idx); + let backward_offset = confidence_byte_offset(geometry, backward_idx); + let forward_handle = confidence_buf.clone().offset_start(forward_offset); + let backward_handle = confidence_buf.clone().offset_start(backward_offset); + + ConfidenceArgs { + use_confidence: true, + conf_fwd: unsafe { ArrayArg::from_raw_parts(forward_handle, conf_len) }, + conf_bwd: unsafe { ArrayArg::from_raw_parts(backward_handle, conf_len) }, + step: geometry.step, + blocks_x: geometry.blocks_x, + blocks_y: geometry.blocks_y, + } + } else { + ConfidenceArgs { + use_confidence: false, + conf_fwd: unsafe { ArrayArg::from_raw_parts(self.confidence_dummy.clone(), 1) }, + conf_bwd: unsafe { ArrayArg::from_raw_parts(self.confidence_dummy.clone(), 1) }, + step: 1, + blocks_x: 1, + blocks_y: 1, + } + } + } + + /// The windowed fused step for a temporal neighbour. + /// + /// One launch covers the whole search window and keeps the accumulators in registers, in place + /// of `(2 * search_radius + 1)^2` launches. + fn dispatch_fused_window_iter( + &self, + ctx: &LaunchCtx, + center_t: u32, + q_k: i32, + ) -> Result<(), anyhow::Error> { + let channels = self.params.channels.count(); + let frame_t = self.phys_frame(center_t as i32); + let frame_fwd = self.phys_frame(center_t as i32 + q_k); + let frame_bwd = self.phys_frame(center_t as i32 - q_k); + let confidence = self.confidence_pair_args(q_k); + + if self.use_reference { + unsafe { + nlm_fused_pair_accumulate_window_ref::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg_for_temporal(ctx), + self.reference_arg_for_temporal(ctx), + self.accum_arg(ctx), + self.weight_sum_arg(ctx), + self.max_weight_arg(ctx), + confidence.conf_fwd, + confidence.conf_bwd, + confidence.use_confidence, + frame_t, + frame_fwd, + frame_bwd, + self.h2_inv_norm, + self.noise_offset, + self.width, + self.height, + channels, + self.params.patch_radius, + self.params.search_radius, + BLOCK_X, + BLOCK_Y, + confidence.step, + confidence.blocks_x, + confidence.blocks_y, + ); + } + } else { + unsafe { + nlm_fused_pair_accumulate_window::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg_for_temporal(ctx), + self.accum_arg(ctx), + self.weight_sum_arg(ctx), + self.max_weight_arg(ctx), + confidence.conf_fwd, + confidence.conf_bwd, + confidence.use_confidence, + frame_t, + frame_fwd, + frame_bwd, + self.h2_inv_norm, + self.noise_offset, + self.width, + self.height, + channels, + self.params.patch_radius, + self.params.search_radius, + BLOCK_X, + BLOCK_Y, + confidence.step, + confidence.blocks_x, + confidence.blocks_y, + ); + } + } + + Ok(()) + } + + /// The windowed fused step for the frame against itself. + /// + /// Patch distance is symmetric, so the full window walked in one direction gives the same total + /// as a half window walked both ways. + fn dispatch_fused_single_window_iter(&self, ctx: &LaunchCtx, center_t: u32) -> Result<(), anyhow::Error> { + let channels = self.params.channels.count(); + let frame_t = self.phys_frame(center_t as i32); + + if self.use_reference { + unsafe { + nlm_fused_single_window_ref::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg(ctx), + self.reference_arg(ctx), + self.accum_arg(ctx), + self.weight_sum_arg(ctx), + self.max_weight_arg(ctx), + frame_t, + self.h2_inv_norm, + self.spatial_offset_lut_arg(), + self.width, + self.height, + channels, + self.params.patch_radius, + self.params.search_radius, + BLOCK_X, + BLOCK_Y, + ); + } + } else { + unsafe { + nlm_fused_single_window::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg(ctx), + self.accum_arg(ctx), + self.weight_sum_arg(ctx), + self.max_weight_arg(ctx), + frame_t, + self.h2_inv_norm, + self.spatial_offset_lut_arg(), + self.width, + self.height, + channels, + self.params.patch_radius, + self.params.search_radius, + BLOCK_X, + BLOCK_Y, + ); + } + } + + Ok(()) + } + + /// The separable step for a temporal neighbour. + /// + /// The last kernel finishes the vertical sum, the weights and the accumulation from both + /// horizontal-sum buffers, so no weight buffer reaches global memory. + fn dispatch_separable_iter( + &self, + ctx: &LaunchCtx, + center_t: u32, + q_x: i32, + q_y: i32, + q_k: i32, + ) -> Result<(), anyhow::Error> { + let channels = self.params.channels.count(); + let frame_t = self.phys_frame(center_t as i32); + let frame_fwd = self.phys_frame(center_t as i32 + q_k); + let frame_bwd = self.phys_frame(center_t as i32 - q_k); + + if self.use_reference { + unsafe { + nlm_distance_pair_ref::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.reference_arg_for_temporal(ctx), + self.raw_fwd_arg(ctx), + self.raw_bwd_arg(ctx), + frame_t, + frame_fwd, + frame_bwd, + q_x, + q_y, + self.width, + self.height, + channels, + ); + } + } else { + unsafe { + nlm_distance_pair::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg_for_temporal(ctx), + self.raw_fwd_arg(ctx), + self.raw_bwd_arg(ctx), + frame_t, + frame_fwd, + frame_bwd, + q_x, + q_y, + self.width, + self.height, + channels, + ); + } + } + + unsafe { + nlm_horizontal_sum_pair::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.raw_fwd_arg(ctx), + self.raw_bwd_arg(ctx), + self.tmp_hsum_arg(ctx), + self.tmp_hsum_bwd_arg(ctx), + self.width, + self.height, + self.params.patch_radius, + BLOCK_X, + BLOCK_Y, + ); + } + + let confidence = self.confidence_pair_args(q_k); + unsafe { + nlm_vweight_pair_accumulate::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.tmp_hsum_arg(ctx), + self.tmp_hsum_bwd_arg(ctx), + self.input_arg_for_temporal(ctx), + self.accum_arg(ctx), + self.weight_sum_arg(ctx), + self.max_weight_arg(ctx), + confidence.conf_fwd, + confidence.conf_bwd, + confidence.use_confidence, + frame_fwd, + frame_bwd, + q_x, + q_y, + self.h2_inv_norm, + self.noise_offset, + self.width, + self.height, + self.params.patch_radius, + BLOCK_X, + BLOCK_Y, + confidence.step, + confidence.blocks_x, + confidence.blocks_y, + ); + } + + Ok(()) + } + + /// The separable step for the frame against itself. + /// + /// The weight map is symmetric, so one buffer serves both the forward and backward lookups. + fn dispatch_separable_iter_k0( + &self, + ctx: &LaunchCtx, + center_t: u32, + q_x: i32, + q_y: i32, + ) -> Result<(), anyhow::Error> { + let channels = self.params.channels.count(); + let frame_t = self.phys_frame(center_t as i32); + + if self.use_reference { + unsafe { + nlm_distance_ref::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.reference_arg(ctx), + self.raw_fwd_arg(ctx), + frame_t, + frame_t, + q_x, + q_y, + self.width, + self.height, + channels, + ); + } + } else { + unsafe { + nlm_distance::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg(ctx), + self.raw_fwd_arg(ctx), + frame_t, + frame_t, + q_x, + q_y, + self.width, + self.height, + channels, + ); + } + } + + unsafe { + nlm_horizontal_sum::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.raw_fwd_arg(ctx), + self.tmp_hsum_arg(ctx), + self.width, + self.height, + self.params.patch_radius, + BLOCK_X, + BLOCK_Y, + ); + } + + // The same correlation-adjusted offset the spatial table holds for this candidate, computed + // directly because the candidate is already known here. + let rho = self.rho_smoothed.unwrap_or(0.0); + let offset_factor = spatial_offset_factor(q_x, q_y, rho); + let offset = self.noise_offset * offset_factor; + unsafe { + nlm_vertical_weight::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.tmp_hsum_arg(ctx), + self.weight_buf_arg(ctx), + self.h2_inv_norm, + offset, + self.width, + self.height, + self.params.patch_radius, + BLOCK_X, + BLOCK_Y, + ); + } + + unsafe { + nlm_accumulate::launch_unchecked::( + &self.client, + ctx.thin_cube_count.clone(), + ctx.thin_cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg(ctx), + self.accum_arg(ctx), + self.weight_sum_arg(ctx), + self.weight_buf_arg(ctx), + self.weight_buf_arg(ctx), + self.max_weight_arg(ctx), + frame_t, + frame_t, + q_x, + q_y, + self.width, + self.height, + ); + } + + Ok(()) + } + + fn zero_accumulators(&self, ctx: &LaunchCtx) -> Result<(), anyhow::Error> { + let grid = (ctx.frame_size as u32).div_ceil(BLOCK_1D).min(MAX_GRID_1D); + let total_threads = grid * BLOCK_1D; + unsafe { + gpu_zero_buffers::launch_unchecked::( + &self.client, + CubeCount::new_1d(grid), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(self.accum.clone(), ctx.frame_size), + self.weight_sum_arg(ctx), + self.max_weight_arg(ctx), + ctx.frame_size as u32, + ctx.pixels as u32, + total_threads, + ); + } + + Ok(()) + } + + /// Launches `nlm_finish` into a destination the caller picks. + fn run_finish_to( + &self, + ctx: &LaunchCtx, + center_frame: u32, + output_frame: u32, + output: ArrayArg, + ) -> Result<(), anyhow::Error> { + let channels = self.params.channels.count(); + unsafe { + nlm_finish::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg(ctx), + output, + ArrayArg::from_raw_parts(self.accum.clone(), ctx.frame_size), + self.weight_sum_arg(ctx), + self.max_weight_arg(ctx), + center_frame, + output_frame, + self.params.self_weight, + self.width, + self.height, + channels, + ); + } + + Ok(()) + } + + fn run_finish(&self, ctx: &LaunchCtx, center_t: u32, output_slot: usize) -> Result<(), anyhow::Error> { + let center_frame = self.phys_frame(center_t as i32); + let output = self.output_arg(ctx, output_slot); + + self.run_finish_to(ctx, center_frame, 0, output) + } + + /// The launch shapes every per-frame dispatch shares, for the main pass and the NLM pilot. + fn launch_ctx(&self) -> LaunchCtx { + let width = self.width; + let height = self.height; + let stored_ch = self.params.channels.storage_count(); + let total_frames = self.params.total_frames(); + let pixels = (width * height) as usize; + let frame_size = pixels * stored_ch as usize; + + let cubes_x = width.div_ceil(BLOCK_X); + let cubes_y = height.div_ceil(BLOCK_Y); + let thin_cubes_x = width.div_ceil(BLOCK_X_THIN); + let thin_cubes_y = height.div_ceil(BLOCK_Y_THIN); + + LaunchCtx { + total_frame_data: frame_size * total_frames as usize, + frame_size, + pixels, + cube_count: CubeCount::new_2d(cubes_x, cubes_y), + cube_dim: CubeDim::new_2d(BLOCK_X, BLOCK_Y), + thin_cube_count: CubeCount::new_2d(thin_cubes_x, thin_cubes_y), + thin_cube_dim: CubeDim::new_2d(BLOCK_X_THIN, BLOCK_Y_THIN), + } + } + + /// Denoises a pushed frame with the windowed spatial kernel into its reference-ring slot. + /// + /// It shares the frame accumulators with the main pass. That is safe because the GPU queue runs + /// in order and the main pass zeroes them before use. + pub(super) fn run_nlm_spatial_pilot(&self, slot: u32, strength_scale: f32) -> Result<(), anyhow::Error> { + let ctx = self.launch_ctx(); + self.zero_accumulators(&ctx)?; + + let channels = self.params.channels.count(); + let pilot_h2 = self.h2_inv_norm / (strength_scale * strength_scale); + + // A flat table with no correlation adjustment, because the pilot compares noisy input + // patches and keeps the full white-noise floor. It is rebuilt each call since + // `input_noise_offset` can change between pushes. + let pilot_lut = build_spatial_offset_lut(self.params.search_radius, 0.0, self.input_noise_offset); + let pilot_lut_bytes = f32::as_bytes(&pilot_lut); + let pilot_lut_handle = self.client.create_from_slice(pilot_lut_bytes); + let pilot_lut_arg = unsafe { ArrayArg::::from_raw_parts(pilot_lut_handle, pilot_lut.len()) }; + + // Always the noisy input. For `NlmSpatial` the reference ring is this pass's output, even + // though `use_reference` is set so the main pass picks the `_ref` kernels. + unsafe { + nlm_fused_single_window::launch_unchecked::( + &self.client, + ctx.cube_count.clone(), + ctx.cube_dim, + self.params.channels.storage_count() as usize, + self.input_arg(&ctx), + self.accum_arg(&ctx), + self.weight_sum_arg(&ctx), + self.max_weight_arg(&ctx), + slot, + pilot_h2, + pilot_lut_arg, + self.width, + self.height, + channels, + self.params.patch_radius, + self.params.search_radius, + BLOCK_X, + BLOCK_Y, + ); + } + + let reference_ring = self.reference_ring_arg(&ctx); + self.run_finish_to(&ctx, slot, slot, reference_ring) + } + + pub(super) fn run_denoise_kernels(&mut self, output_slot: usize) -> Result<(), anyhow::Error> { + let temporal_radius = self.params.temporal_radius; + let search_radius = self.params.search_radius as i32; + + let ctx = self.launch_ctx(); + + let center_t = temporal_radius; + + // Both run before any NLM dispatch so the temporal kernels read aligned neighbours. + // `confidence_ctx` only exists with motion compensation off, so at most one does work. + self.run_motion_compensation(center_t)?; + self.run_confidence_pass(center_t)?; + + self.zero_accumulators(&ctx)?; + let window_side = 2 * search_radius + 1; + let window_area = window_side * window_side; + + // The windowed path makes one launch per temporal offset. The separable path launches per + // search offset, and `linear < 0` keeps half of the frame-against-itself window because + // its weights are symmetric. + let k_start = -(temporal_radius as i32); + let use_windowed = !self.use_separable; + for q_k in k_start..=0 { + if use_windowed { + if q_k != 0 { + self.dispatch_fused_window_iter(&ctx, center_t, q_k)?; + } else { + self.dispatch_fused_single_window_iter(&ctx, center_t)?; + } + + continue; + } + + for q_y in -search_radius..=search_radius { + for q_x in -search_radius..=search_radius { + let linear = q_k * window_area + q_y * window_side + q_x; + if linear >= 0 { + continue; + } + + if q_k == 0 { + self.dispatch_separable_iter_k0(&ctx, center_t, q_x, q_y)?; + } else { + self.dispatch_separable_iter(&ctx, center_t, q_x, q_y, q_k)?; + } + } + } + } + + self.run_finish(&ctx, center_t, output_slot)?; + + Ok(()) + } +} diff --git a/av-denoise-core/src/nlmeans/dispatch/motion.rs b/av-denoise-core/src/nlmeans/dispatch/motion.rs new file mode 100644 index 0000000..56e91a8 --- /dev/null +++ b/av-denoise-core/src/nlmeans/dispatch/motion.rs @@ -0,0 +1,427 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use crate::nlmeans::denoiser::NlmDenoiser; +use crate::nlmeans::kernels::gpu_copy; +use crate::nlmeans::motion::{ + self, + MotionCtx, + MotionEstimation, + run_analyse, + run_compensate, + run_confidence_for_neighbour, + run_seeded_refine, +}; +use crate::nlmeans::prefilter::PrefilterMode; +use crate::nlmeans::{BLOCK_1D, MAX_GRID_1D}; + +/// The share of the raw noise the `NlmSpatial` pilot's reference image still carries. +/// +/// A sweep from 0 to 1 on `clean-1080p.mkv` is flat to within 0.10 dB and peaks at 1.0. It is 0 +/// because a floor large enough to swamp `thsad` leaves confidence unable to tell a real mismatch +/// from noise, and the flat sweep makes that free. +pub(in crate::nlmeans) const NLM_SPATIAL_RESIDUAL_FRACTION: f32 = 0.0; + +/// The share of the raw noise the `Bilateral` prefilter's reference image still carries. +/// +/// A sweep with motion compensation on peaks at 0 at every noise level, by 0.16 dB at the lightest +/// and 0.57 dB at the heaviest. It stays separate from [NLM_SPATIAL_RESIDUAL_FRACTION] because the +/// two zeros come from different reasons. +pub(in crate::nlmeans) const BILATERAL_RESIDUAL_FRACTION: f32 = 0.0; + +/// The sigma to pass [motion::sad_noise_floor] for the motion-compensation block match. +/// +/// The match runs on the reference pyramid whenever one exists, and a GPU prefilter has already +/// cleaned it, so the raw floor overstates the real one. With the NLM pilot, the default block +/// size and a sigma of 0.02, the raw floor is about 5.78 against a threshold of 5.12, which pins +/// confidence at 1.0 even on occluded blocks. The raw sigma is therefore scaled by each +/// prefilter's measured residual fraction, and `PrefilterMode::None` keeps it whole. +pub(in crate::nlmeans) fn mc_sad_noise_floor_sigma(prefilter: PrefilterMode, sigma_y: f32) -> f32 { + match prefilter { + PrefilterMode::NlmSpatial { .. } => sigma_y * NLM_SPATIAL_RESIDUAL_FRACTION, + PrefilterMode::Bilateral { .. } => sigma_y * BILATERAL_RESIDUAL_FRACTION, + PrefilterMode::None => sigma_y, + } +} + +impl NlmDenoiser { + /// Estimates how the neighbour at temporal offset `k` moved and writes it to its `mv_field` + /// slot. + /// + /// `Chained` estimation composes the pair fields and refines the seed, and otherwise a direct + /// coarse-to-fine match runs. It shifts no buffer and returns the neighbour's physical slot. + #[expect( + clippy::too_many_arguments, + reason = "the dispatch threads through every buffer and shape the kernel binds" + )] + fn run_motion_estimate( + &self, + motion_ctx: &MotionCtx, + analyse_pyramid: &Handle, + mv_field: &Handle, + confidence_arg: &Handle, + write_confidence: bool, + frame_count: u32, + centre_slot: u32, + center_t: u32, + k: i32, + neighbour_idx: u32, + sad_noise_floor: f32, + thsad: f32, + ) -> Result { + let neighbour_slot = self.phys_frame(center_t as i32 + k); + + if self.is_chained() { + self.run_chain_compose(center_t, k, neighbour_idx)?; + let refine_radius = match self + .params + .motion_compensation + .resolved_estimation(self.params.temporal_radius) + { + Some(MotionEstimation::Chained { refine_radius }) => refine_radius, + _ => unreachable!("is_chained() guarantees a resolved Chained estimation"), + }; + + run_seeded_refine::( + &self.client, + motion_ctx, + self.width, + self.height, + frame_count, + centre_slot, + neighbour_slot, + neighbour_idx, + refine_radius, + analyse_pyramid, + mv_field, + confidence_arg, + write_confidence, + sad_noise_floor, + thsad, + )?; + } else { + run_analyse::( + &self.client, + motion_ctx, + self.width, + self.height, + frame_count, + centre_slot, + neighbour_slot, + neighbour_idx, + analyse_pyramid, + mv_field, + confidence_arg, + write_confidence, + sad_noise_floor, + thsad, + )?; + } + + Ok(neighbour_slot) + } + + /// Runs the motion estimate for every neighbour in the window without shifting any buffer. + /// + /// It returns each neighbour's physical slot in logical ring order, skipping the centre, or an + /// empty list when motion compensation is off or there are no neighbours. + pub(in crate::nlmeans) fn run_motion_machinery(&self, center_t: u32) -> Result, anyhow::Error> { + let Some(motion_ctx) = self.mc_ctx.as_ref() else { + return Ok(Vec::new()); + }; + + let temporal_radius = self.params.temporal_radius; + if temporal_radius == 0 { + return Ok(Vec::new()); + } + + let frame_count = self.params.total_frames(); + let centre_slot = self.phys_frame(center_t as i32); + + let pyramid_input = self + .pyramid_input + .as_ref() + .expect("pyramid_input allocated when mc_ctx is Some"); + let mv_field = self + .mv_field_buf + .as_ref() + .expect("mv_field allocated when mc_ctx is Some"); + let (confidence_arg, write_confidence): (&Handle, bool) = match self.confidence_buf.as_ref() { + Some(buf) => (buf, true), + None => (&self.confidence_dummy, false), + }; + let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); + let mc_sigma_y = mc_sad_noise_floor_sigma(self.params.prefilter, self.sigma_y); + let sad_noise_floor = motion::sad_noise_floor(motion_ctx.blksize, mc_sigma_y); + let thsad = motion::thsad(motion_ctx.blksize, thsad_scale); + + // Match against the reference pyramid when one exists, because it is cleaner. + let analyse_pyramid = self.pyramid_reference.as_ref().unwrap_or(pyramid_input); + + let mut neighbour_idx: u32 = 0; + let mut slots = Vec::with_capacity((frame_count - 1) as usize); + for logical in 0..frame_count { + if logical == center_t { + continue; + } + + let k = logical as i32 - center_t as i32; + let neighbour_slot = self.run_motion_estimate( + motion_ctx, + analyse_pyramid, + mv_field, + confidence_arg, + write_confidence, + frame_count, + centre_slot, + center_t, + k, + neighbour_idx, + sad_noise_floor, + thsad, + )?; + slots.push(neighbour_slot); + + neighbour_idx += 1; + } + + Ok(slots) + } + + /// Estimates each neighbour's motion and shifts it into the compensated rings. + /// + /// The centre slot is copied through unchanged so the temporal kernels read every slot the same + /// way. It does nothing when motion compensation is off or there are no neighbours. + pub(super) fn run_motion_compensation(&self, center_t: u32) -> Result<(), anyhow::Error> { + let Some(motion_ctx) = self.mc_ctx.as_ref() else { + return Ok(()); + }; + + let temporal_radius = self.params.temporal_radius; + if temporal_radius == 0 { + return Ok(()); + } + + let frame_count = self.params.total_frames(); + let centre_slot = self.phys_frame(center_t as i32); + let stored_ch = self.params.channels.storage_count(); + + let pyramid_input = self + .pyramid_input + .as_ref() + .expect("pyramid_input allocated when mc_ctx is Some"); + let mv_field = self + .mv_field_buf + .as_ref() + .expect("mv_field allocated when mc_ctx is Some"); + let compensated_input = self + .compensated_input_buf + .as_ref() + .expect("compensated_input allocated when mc_ctx is Some"); + // Without confidence weighting the fine kernel still needs a buffer, so it gets the dummy + // and is told not to write it. + let (confidence_arg, write_confidence): (&Handle, bool) = match self.confidence_buf.as_ref() { + Some(buf) => (buf, true), + None => (&self.confidence_dummy, false), + }; + let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); + let mc_sigma_y = mc_sad_noise_floor_sigma(self.params.prefilter, self.sigma_y); + let sad_noise_floor = motion::sad_noise_floor(motion_ctx.blksize, mc_sigma_y); + let thsad = motion::thsad(motion_ctx.blksize, thsad_scale); + + copy_frame_into_slot_handle::( + &self.client, + &self.input_buf, + compensated_input, + centre_slot as usize, + self.params.total_frames(), + self.width, + self.height, + stored_ch, + ); + if let (Some(ref_src), Some(ref_dst)) = ( + self.reference_buf.as_ref(), + self.compensated_reference_buf.as_ref(), + ) { + copy_frame_into_slot_handle::( + &self.client, + ref_src, + ref_dst, + centre_slot as usize, + self.params.total_frames(), + self.width, + self.height, + stored_ch, + ); + } + + // Match against the reference pyramid when one exists, because it is cleaner. + let analyse_pyramid = self.pyramid_reference.as_ref().unwrap_or(pyramid_input); + + // Neighbours run from the furthest behind to the furthest ahead, skipping the centre, so + // their motion-field indices stay contiguous. + let radius = temporal_radius as i32; + let mut neighbour_idx: u32 = 0; + for k in -radius..=radius { + if k == 0 { + continue; + } + + let neighbour_slot = self.run_motion_estimate( + motion_ctx, + analyse_pyramid, + mv_field, + confidence_arg, + write_confidence, + frame_count, + centre_slot, + center_t, + k, + neighbour_idx, + sad_noise_floor, + thsad, + )?; + + run_compensate::( + &self.client, + motion_ctx, + stored_ch, + self.width, + self.height, + frame_count, + neighbour_slot, + neighbour_idx, + &self.input_buf, + compensated_input, + mv_field, + )?; + + if let (Some(ref_src), Some(ref_dst)) = ( + self.reference_buf.as_ref(), + self.compensated_reference_buf.as_ref(), + ) { + run_compensate::( + &self.client, + motion_ctx, + stored_ch, + self.width, + self.height, + frame_count, + neighbour_slot, + neighbour_idx, + ref_src, + ref_dst, + mv_field, + )?; + } + + neighbour_idx += 1; + } + + Ok(()) + } + + /// Scores each neighbour against the centre frame with every block matched where it stands. + /// + /// It runs only when confidence is on and motion compensation is off, the one case where + /// `confidence_ctx` exists. + pub(super) fn run_confidence_pass(&self, center_t: u32) -> Result<(), anyhow::Error> { + let Some(ctx) = self.confidence_ctx.as_ref() else { + return Ok(()); + }; + + // `confidence_ctx` only exists for a temporal radius above 0. + let temporal_radius = self.params.temporal_radius; + + let frame_count = self.params.total_frames(); + let centre_slot = self.phys_frame(center_t as i32); + + let luma_pyramid = self + .confidence_pyramid + .as_ref() + .expect("confidence_pyramid allocated when confidence_ctx is Some"); + let mv_scratch = self + .confidence_mv_scratch + .as_ref() + .expect("confidence_mv_scratch allocated when confidence_ctx is Some"); + let confidence_buf = self + .confidence_buf + .as_ref() + .expect("confidence_buf allocated when confidence_ctx is Some"); + + let thsad_scale = self.params.hq.map_or(1.0, |hq| hq.thsad_scale); + let sad_noise_floor = motion::sad_noise_floor(ctx.blksize, self.sigma_y); + let thsad = motion::thsad(ctx.blksize, thsad_scale); + + let radius = temporal_radius as i32; + let mut neighbour_idx: u32 = 0; + for k in -radius..=radius { + if k == 0 { + continue; + } + + let neighbour_slot = self.phys_frame(center_t as i32 + k); + + run_confidence_for_neighbour::( + &self.client, + ctx, + self.width, + self.height, + frame_count, + centre_slot, + neighbour_slot, + neighbour_idx, + luma_pyramid, + mv_scratch, + confidence_buf, + sad_noise_floor, + thsad, + )?; + + neighbour_idx += 1; + } + + Ok(()) + } +} + +/// Copies one frame from a slot of `src` into the same slot of `dst`, which share a ring layout. +/// +/// Both rings are bound whole and the kernel picks the slot, because a slot's byte offset rarely +/// meets the GPU's `min_storage_buffer_offset_alignment`. It is a free function so the motion pass +/// can call it without re-borrowing the denoiser. +#[expect( + clippy::too_many_arguments, + reason = "the dispatch threads through every buffer and shape the kernel binds" +)] +fn copy_frame_into_slot_handle( + client: &ComputeClient, + src: &Handle, + dst: &Handle, + slot: usize, + frame_count: u32, + width: u32, + height: u32, + stored_ch: u32, +) { + let frame_size = width * height * stored_ch; + let ring_len = frame_count as usize * frame_size as usize; + let offset = slot as u32 * frame_size; + + let grid = frame_size.div_ceil(BLOCK_1D).min(MAX_GRID_1D); + let total_threads = grid * BLOCK_1D; + + unsafe { + gpu_copy::launch_unchecked::( + client, + CubeCount::new_1d(grid), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(src.clone(), ring_len), + ArrayArg::from_raw_parts(dst.clone(), ring_len), + offset, + offset, + frame_size, + total_threads, + ); + } +} diff --git a/av-denoise-core/src/nlmeans/edges.rs b/av-denoise-core/src/nlmeans/edges.rs index 9a8bf73..c0a6aee 100644 --- a/av-denoise-core/src/nlmeans/edges.rs +++ b/av-denoise-core/src/nlmeans/edges.rs @@ -14,7 +14,7 @@ impl NlmDenoiser { self.phys_frame(logical as i32) } - /// Turns the shifted-edge stream layout on or off. See the `shifted_edges` field. + /// Turns the shifted-edge stream layout on or off. pub(crate) fn set_shifted_edges(&mut self, on: bool) { self.shifted_edges = on; } @@ -22,18 +22,20 @@ impl NlmDenoiser { /// Copies the last pushed frame forward until every ring slot holds a frame. /// /// A stream shorter than the ring uses this at flush so its real frames can run as centres. - pub(crate) fn fill_ring_with_last_frame(&mut self) { + pub(crate) fn fill_ring_with_last_frame(&mut self) -> Result<(), anyhow::Error> { while !self.window_ready() { - self.duplicate_last_frame(); + self.duplicate_last_frame()?; self.frames_loaded += 1; } + + Ok(()) } /// The temporal reading a pass centred on `center_t` folds. /// - /// This is the centre's own reading. With shifted edges on and no - /// usable sample at the centre, it is the first usable reading at a - /// later ring position, or the centre's empty one when there is none. + /// It is the centre's own reading, unless shifted edges are on and the centre has no usable + /// sample. Then it is the first usable reading at a later ring position, or the centre's empty + /// one when there is none. pub(super) fn borrow_reading_ahead(&self, center_t: u32) -> Result { let centre_slot = self.ring_slot(center_t); let own = self.read_temporal_noise(centre_slot)?; diff --git a/av-denoise-core/src/nlmeans/engine.rs b/av-denoise-core/src/nlmeans/engine.rs new file mode 100644 index 0000000..9ceab8d --- /dev/null +++ b/av-denoise-core/src/nlmeans/engine.rs @@ -0,0 +1,222 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::denoiser::{GpuOutput, NlmDenoiser, check_u32_indexable, front_buffer_sizes}; +use super::options::{NlmeansAlgorithm, resolve_params}; +use super::params::validate_dimensions; +use crate::engine::{DevicePlane, EdgePadding, EgressSource, Engine, Geometry, WindowSpan, egress}; +use crate::error::Error; + +/// Non-local means over GPU planes. +pub struct Nlmeans { + front: NlmDenoiser, + geometry: Geometry, + /// The finished frame waiting for `emit_into`, if any. + ready: Option, + /// Tail frames still to produce after `finish`. + tail_remaining: usize, + /// Whether the current stream has had a real push, after which context frames are refused. + pushed: bool, + poisoned: bool, +} + +impl Nlmeans { + /// Builds the engine, returning [Error::InvalidGeometry] or [Error::InvalidOptions] for unusable input. + pub fn new( + client: &ComputeClient, + algorithm: NlmeansAlgorithm, + geometry: Geometry, + ) -> Result { + geometry.validate()?; + + let dimensions = validate_dimensions(geometry.width, geometry.height); + dimensions.map_err(|error| { + let message = error.to_string(); + Error::InvalidGeometry(message) + })?; + + let params = resolve_params(&algorithm, geometry.channels); + let validated = params.validate(); + validated.map_err(|error| { + let message = error.to_string(); + Error::InvalidOptions(message) + })?; + + let ring_slots = u64::from(params.total_frames()); + geometry.check_ring_fits(ring_slots)?; + + let sizes = front_buffer_sizes(client, ¶ms, geometry.width, geometry.height); + let indexable = check_u32_indexable(&sizes.buffers); + indexable.map_err(Error::InvalidGeometry)?; + + let front = NlmDenoiser::new(client, params, geometry.width, geometry.height); + + Ok(Self { + front, + geometry, + ready: None, + tail_remaining: 0, + pushed: false, + poisoned: false, + }) + } + + fn check_usable(&self) -> Result<(), Error> { + if self.poisoned { + return Err(Error::NeedsReset); + } + + Ok(()) + } + + fn check_nothing_pending(&self) -> Result<(), Error> { + if self.ready.is_some() || self.tail_remaining > 0 { + return Err(Error::OutputsPending); + } + + Ok(()) + } + + fn check_no_push_yet(&self) -> Result<(), Error> { + if self.pushed { + return Err(Error::ContextAfterPush); + } + + Ok(()) + } + + fn restart_stream(&mut self) { + self.front.reset_stream_state(); + self.pushed = false; + } + + /// Poisons the engine if `result` is an error. + fn guard(&mut self, result: Result) -> Result { + if result.is_err() { + self.poisoned = true; + } + + result.map_err(Error::Gpu) + } + + /// Steps the flush until it yields the next tail frame. + fn next_tail_frame(&mut self) -> Result { + loop { + let step = self.front.flush_step_gpu()?; + if let Some(output) = step { + return Ok(output); + } + } + } + + fn write_out(&self, frame: &Handle, planes: &[DevicePlane<'_>]) { + let channels = self.geometry.channels; + // The constructor bounds the ring to `u32` elements, so one plane always fits. + let pixels = self.geometry.pixels() as u32; + let source = EgressSource { + frame, + pixels, + channels: channels.count(), + stored_ch: channels.storage_count(), + }; + let client = self.front.compute_client(); + let placeholder = self.front.placeholder(); + + egress(client, source, planes, self.geometry.output, placeholder); + } + + #[cfg(test)] + pub(crate) fn fail_through_guard_for_test(&mut self) -> Result<(), Error> { + let failure = Err(anyhow::anyhow!("forced failure")); + + self.guard(failure) + } +} + +impl Engine for Nlmeans { + fn push(&mut self, planes: &[DevicePlane<'_>]) -> Result { + self.check_usable()?; + self.check_nothing_pending()?; + self.geometry.check_planes(planes, self.geometry.input)?; + + let pushed = self.front.push_planes(planes, self.geometry.input); + self.guard(pushed)?; + + self.pushed = true; + + let submitted = self.front.denoise_submit_gpu(); + let output = self.guard(submitted)?; + self.ready = output.map(|output| output.handle); + + let has_ready = self.ready.is_some(); + Ok(usize::from(has_ready)) + } + + fn push_context(&mut self, planes: &[DevicePlane<'_>]) -> Result<(), Error> { + self.check_usable()?; + self.check_nothing_pending()?; + self.geometry.check_planes(planes, self.geometry.input)?; + self.check_no_push_yet()?; + + let pushed = self.front.push_planes(planes, self.geometry.input); + self.guard(pushed) + } + + fn emit_into(&mut self, planes: &[DevicePlane<'_>]) -> Result<(), Error> { + self.check_usable()?; + self.geometry.check_planes(planes, self.geometry.output)?; + + if let Some(frame) = self.ready.take() { + self.write_out(&frame, planes); + return Ok(()); + } + + if self.tail_remaining == 0 { + return Err(Error::NothingToEmit); + } + + let step = self.next_tail_frame(); + let output = self.guard(step)?; + self.write_out(&output.handle, planes); + + self.tail_remaining -= 1; + if self.tail_remaining == 0 { + self.restart_stream(); + } + + Ok(()) + } + + fn finish(&mut self) -> Result { + self.check_usable()?; + self.check_nothing_pending()?; + + self.tail_remaining = self.front.flush_target(); + if self.tail_remaining == 0 { + self.restart_stream(); + } + + Ok(self.tail_remaining) + } + + fn reset(&mut self) { + self.ready = None; + self.tail_remaining = 0; + self.poisoned = false; + self.restart_stream(); + } + + fn window_span(&self) -> WindowSpan { + let radius = self.front.params.temporal_radius as usize; + + WindowSpan { + behind: radius, + ahead: radius, + edges: EdgePadding::Repeat, + } + } + + fn max_held_frames(&self) -> usize { + self.front.params.temporal_radius as usize + } +} diff --git a/av-denoise-core/src/nlmeans/kernels/accumulate.rs b/av-denoise-core/src/nlmeans/kernels/accumulate.rs index 671dd52..20c22ef 100644 --- a/av-denoise-core/src/nlmeans/kernels/accumulate.rs +++ b/av-denoise-core/src/nlmeans/kernels/accumulate.rs @@ -3,14 +3,11 @@ use cubecl::terminate; use super::helpers::{accumulate_pair, clamp_coord}; -/// Applies the forward and backward neighbour contributions at every -/// pixel from a weight map. +/// Adds the forward and backward neighbour contributions at every pixel from a weight map. /// -/// `weights_fwd` and `weights_bwd` may be the same buffer, which is what -/// the symmetric case at temporal offset 0 does. -/// -/// The backward lookup uses the clamped neighbour index so border pixels -/// still read a valid weight. +/// `weights_fwd` and `weights_bwd` may be the same buffer, as in the symmetric case at temporal +/// offset 0. The backward weight is read at the clamped neighbour so border pixels still read a +/// valid weight. #[cube(launch_unchecked)] pub fn nlm_accumulate( input: &Array>, @@ -47,18 +44,12 @@ pub fn nlm_accumulate( /// Turns the accumulated sums into the denoised output. /// -/// The result is `(original * m + acc) / (m + weight_sum)`, where `m` is -/// `wref * max_weight`. -/// -/// When the denominator is close to zero, meaning no usable match was -/// found anywhere in the search window, the original pixel is kept as it -/// is. +/// The result is `(original * self_weight + accum) / (self_weight + weight_sum)`, where +/// `self_weight` is `wref * max_weight`. When the denominator is close to zero, no usable match was +/// found and the original pixel is kept. /// -/// The centre pixel is read from `input`'s frame `center_frame` and the -/// result written to `output`'s frame `output_frame`. A caller writing -/// into a ring buffer slot therefore passes the whole ring, rather than -/// binding it at the slot's byte offset, which the GPU only accepts on -/// its own alignment boundaries. +/// `center_frame` and `output_frame` index into the whole bound rings, because a buffer can only +/// be bound at the GPU's offset alignment and a ring slot rarely lands on it. #[cube(launch_unchecked)] pub fn nlm_finish( input: &Array>, @@ -83,26 +74,27 @@ pub fn nlm_finish( let frame_idx = ((center_frame * height + y) * width + x) as usize; let output_idx = ((output_frame * height + y) * width + x) as usize; - let m = wref * max_weight[pixel_idx]; - let denominator = m + weight_sum[pixel_idx]; + let self_weight = wref * max_weight[pixel_idx]; + let denominator = self_weight + weight_sum[pixel_idx]; let original = input[frame_idx]; let accumulated = accum[pixel_idx]; - // `Vector::empty` zeroes its lanes, so the padding lanes a 3-channel - // frame gets in 4-lane storage stay 0 whichever branch runs below. + // `Vector::empty` zeroes its lanes, so the padding lanes of a 3-channel frame stay 0. let mut out = Vector::::empty(); if denominator > 1e-30f32 { let inv_denominator = 1.0f32 / denominator; + #[unroll] - for c in 0..channels { - out[c as usize] = (original[c as usize] * m + accumulated[c as usize]) * inv_denominator; + for channel in 0..channels { + out[channel as usize] = + (original[channel as usize] * self_weight + accumulated[channel as usize]) * inv_denominator; } } else { #[unroll] - for c in 0..channels { - out[c as usize] = original[c as usize]; + for channel in 0..channels { + out[channel as usize] = original[channel as usize]; } } diff --git a/av-denoise-core/src/nlmeans/kernels/bilateral.rs b/av-denoise-core/src/nlmeans/kernels/bilateral.rs index ab968eb..d3ca7bf 100644 --- a/av-denoise-core/src/nlmeans/kernels/bilateral.rs +++ b/av-denoise-core/src/nlmeans/kernels/bilateral.rs @@ -5,18 +5,13 @@ use super::helpers::{read_clamped_line, read_line}; /// A bilateral prefilter that blurs a frame without crossing edges. /// -/// Each neighbour's weight is the product of two Gaussians. One falls -/// off with distance in pixels, the other with difference in colour, so -/// a neighbour on the far side of an edge contributes almost nothing. +/// Each neighbour is weighted by +/// `exp(-(dx^2 + dy^2) * inv_two_sigma_s_sq - range_sq * inv_two_sigma_r_sq)`, so a neighbour +/// across an edge contributes almost nothing. The block caches a `(block + 2 * radius)^2` tile in +/// shared memory, one vector per pixel. /// -/// The block first loads a `(block + 2 * radius)^2` tile of source -/// pixels into shared memory, one vector per pixel, then each thread -/// convolves over its patch using -/// `w = exp(-(dx^2 + dy^2) * inv_two_sigma_s_sq - range_sq * inv_two_sigma_r_sq)`. -/// -/// The output keeps the input's channel layout, with padding lanes -/// copied straight through, so it can stand in for `input` in the `_ref` -/// distance kernels. +/// The output keeps the input's channel layout, so it can stand in as the reference image of the +/// `_ref` distance kernels. #[cube(launch_unchecked)] pub fn nlm_bilateral( input: &Array>, @@ -51,6 +46,7 @@ pub fn nlm_bilateral( let threads = block_x * block_y; let thread_id = local_y * block_x + local_x; let mut idx = thread_id; + if interior { while idx < tile_elems { let tile_x = idx % tile_width; @@ -83,7 +79,8 @@ pub fn nlm_bilateral( let patch_size = 2 * radius + 1; let mut weight_sum = 0.0f32; - let mut acc = Vector::::empty().fill(0.0f32); + let mut weighted_sum = Vector::::empty().fill(0.0f32); + for offset_y in 0..patch_size { for offset_x in 0..patch_size { let dy = offset_y as i32 - radius as i32; @@ -94,20 +91,22 @@ pub fn nlm_bilateral( let diff = neighbor - center; let mut range_sq = 0.0f32; + #[unroll] - for c in 0..channels { - range_sq += diff[c as usize] * diff[c as usize]; + for channel in 0..channels { + range_sq += diff[channel as usize] * diff[channel as usize]; } + let spatial = (dx * dx + dy * dy) as f32 * inv_two_sigma_s_sq; - let w = f32::exp(-(spatial + range_sq * inv_two_sigma_r_sq)); + let weight = f32::exp(-(spatial + range_sq * inv_two_sigma_r_sq)); - let line_w = Vector::::empty().fill(w); - acc += neighbor * line_w; - weight_sum += w; + let line_w = Vector::::empty().fill(weight); + weighted_sum += neighbor * line_w; + weight_sum += weight; } } - let inv = 1.0f32 / weight_sum; - let line_inv = Vector::::empty().fill(inv); - output[((frame * height + global_y) * width + global_x) as usize] = acc * line_inv; + let inv_weight_sum = 1.0f32 / weight_sum; + let line_inv = Vector::::empty().fill(inv_weight_sum); + output[((frame * height + global_y) * width + global_x) as usize] = weighted_sum * line_inv; } diff --git a/av-denoise-core/src/nlmeans/kernels/fused.rs b/av-denoise-core/src/nlmeans/kernels/fused.rs index 29906c4..86db33f 100644 --- a/av-denoise-core/src/nlmeans/kernels/fused.rs +++ b/av-denoise-core/src/nlmeans/kernels/fused.rs @@ -3,18 +3,11 @@ use cubecl::terminate; use super::helpers::{channel_scale, line_sum_sq, read_clamped_line, read_line, welsch_weight}; -/// Measures the distance, box-filters it over the patch, and turns the -/// result into a Welsch weight, all in one kernel. +/// Writes the Welsch weight of the patch distance between each pixel and its neighbour at `q`. /// -/// The block first loads a `(block + 2 * patch_radius)^2` tile of -/// per-pixel scaled distances into shared memory. Each thread then sums -/// its own `(2 * patch_radius + 1)^2` patch and applies the Welsch -/// kernel. -/// -/// An `interior` flag picks unclamped reads when the whole tile, and its -/// shifted twin, lie inside the image. Blocks near the border take the -/// clamped path instead. The flag is the same for every thread in the -/// block, so the branch costs nothing in divergence. +/// The block caches a `(block + 2 * patch_radius)^2` tile of per-pixel distances in shared memory. +/// Blocks whose tile and shifted tile lie inside the image skip clamping, and the branch never +/// diverges because the whole block takes the same side. #[cube(launch_unchecked)] pub fn nlm_dist_2d_weight( input: &Array>, @@ -105,6 +98,7 @@ pub fn nlm_dist_2d_weight( let center_tile_y = local_y + patch_radius; let patch_size = 2 * patch_radius + 1; let mut patch_sum = 0.0f32; + for offset_y in 0..patch_size { for offset_x in 0..patch_size { let smem_idx = ((center_tile_y - patch_radius + offset_y) * tile_width + center_tile_x @@ -119,12 +113,7 @@ pub fn nlm_dist_2d_weight( /// The reference-image version of `nlm_dist_2d_weight`. /// -/// Distances are read from `reference`, a prefiltered or externally -/// supplied image with the same layout as `input`. The weight output is -/// unchanged. -/// -/// This runs when a prefilter is active, so the weights are computed -/// from a cleaner image than the noisy input. +/// Distances are read from `reference`, a cleaner image with the same layout as the input. #[cube(launch_unchecked)] pub fn nlm_dist_2d_weight_ref( reference: &Array>, @@ -215,6 +204,7 @@ pub fn nlm_dist_2d_weight_ref( let center_tile_y = local_y + patch_radius; let patch_size = 2 * patch_radius + 1; let mut patch_sum = 0.0f32; + for offset_y in 0..patch_size { for offset_x in 0..patch_size { let smem_idx = ((center_tile_y - patch_radius + offset_y) * tile_width + center_tile_x @@ -227,42 +217,15 @@ pub fn nlm_dist_2d_weight_ref( output[(global_y * width + global_x) as usize] = welsch_weight(patch_sum, h2_inv_norm, noise_offset); } -/// Compares the centre frame against one pair of temporal neighbours, -/// covering the whole search window in a single launch. -/// -/// The kernel loops over every offset in the search window for one -/// temporal distance, keeping the running accumulator, weight sum, and -/// max weight in registers. Those are written to global memory once at -/// the end, which collapses `(2 * search_radius + 1)^2` launches into -/// one. -/// -/// # Caching the centre frame -/// -/// The centre frame is read once into a shared-memory tile of -/// `(block + 2 * patch_radius + 2 * search_radius)^2` pixels, big enough -/// to cover every neighbour offset the window can reach. -/// -/// The forward and backward comparisons both centre on that same patch, -/// so one cached tile serves both. Only the shifted neighbour pixels -/// come from global memory each iteration. +/// Accumulates one pair of temporal neighbours over the whole search window in a single launch. /// -/// That roughly halves the global read traffic compared with re-reading -/// the centre frame for every offset. +/// The running sums stay in registers and are added to `accum`, `weight_sum` and `max_weight` +/// once at the end. The centre frame is cached once in a +/// `(block + 2 * patch_radius + 2 * search_radius)^2` shared tile that both directions share. /// -/// The two distance tiles are reused across iterations, with a -/// `sync_cube` between them. -/// -/// # Confidence weighting -/// -/// When `use_confidence` is true, each weight is multiplied by its -/// block's confidence before it folds into the accumulators, using the -/// same pixel-to-block mapping `nlm_mc_warp` uses. -/// -/// The block index depends only on the pixel position, so it is the same -/// for every offset in the window. -/// -/// When `use_confidence` is false the lookup and the multiply are -/// dropped at compile time, and the confidence buffers are never read. +/// When `use_confidence` is set, each weight is scaled by its block's confidence from `conf_fwd` +/// and `conf_bwd`. `step`, `blocks_x` and `blocks_y` describe the motion block grid, which pixels +/// map onto the same way as in `nlm_mc_warp`. When it is unset the confidence buffers are never read. #[cube(launch_unchecked)] pub fn nlm_fused_pair_accumulate_window( input: &Array>, @@ -313,17 +276,16 @@ pub fn nlm_fused_pair_accumulate_window( let expanded_x0 = fwd_tile_x0 - search_radius as i32; let expanded_y0 = fwd_tile_y0 - search_radius as i32; - // Cache `frame_t` once across the expanded tile that covers every - // forward and shifted-backward center position. let mut idx = thread_id; while idx < expanded_elems { - let ex = idx % expanded_width; - let ey = idx / expanded_width; - let src_x = expanded_x0 + ex as i32; - let src_y = expanded_y0 + ey as i32; + let expanded_x = idx % expanded_width; + let expanded_y = idx / expanded_width; + let src_x = expanded_x0 + expanded_x as i32; + let src_y = expanded_y0 + expanded_y as i32; smem_center[idx as usize] = read_clamped_line(input, src_x, src_y, frame_t, width, height); idx += threads; } + sync_cube(); let mut accum_reg = Vector::::empty(); @@ -344,12 +306,8 @@ pub fn nlm_fused_pair_accumulate_window( let tile_x = idx % tile_width; let tile_y = idx / tile_width; - // Both the forward and backward centres sit at - // (tile_x + search_radius, tile_y + search_radius) in - // expanded-tile coordinates. The backward comparison is - // centred on the same output pixel as the forward one, - // mirroring it against `frame_bwd` at `-q` instead of - // `frame_fwd` at `+q`. + // The backward comparison mirrors the forward one against `frame_bwd` at `-q`, so + // both centre on the same cached patch. let center_idx = ((tile_y + search_radius) * expanded_width + (tile_x + search_radius)) as usize; let center = smem_center[center_idx]; @@ -384,6 +342,7 @@ pub fn nlm_fused_pair_accumulate_window( let patch_size = 2 * patch_radius + 1; let mut sum_fwd = 0.0f32; let mut sum_bwd = 0.0f32; + for offset_y in 0..patch_size { for offset_x in 0..patch_size { let smem_idx = ((center_tile_y - patch_radius + offset_y) * tile_width @@ -399,9 +358,9 @@ pub fn nlm_fused_pair_accumulate_window( let mut weight_bwd = welsch_weight(sum_bwd, h2_inv_norm, noise_offset); if use_confidence { - let bx = (global_x / step).min(blocks_x - 1); - let by = (global_y / step).min(blocks_y - 1); - let block_idx = (by * blocks_x + bx) as usize; + let block_col = (global_x / step).min(blocks_x - 1); + let block_row = (global_y / step).min(blocks_y - 1); + let block_idx = (block_row * blocks_x + block_col) as usize; weight_fwd *= conf_fwd[block_idx]; weight_bwd *= conf_bwd[block_idx]; } @@ -431,8 +390,7 @@ pub fn nlm_fused_pair_accumulate_window( max_weight_reg = f32::max(max_weight_reg, f32::max(weight_fwd, weight_bwd)); } - // Wait for every thread to finish reading the tiles before - // the next q overwrites them. + // The next offset overwrites the distance tiles. sync_cube(); } } @@ -447,30 +405,14 @@ pub fn nlm_fused_pair_accumulate_window( } } -/// Compares a frame against itself across the search window, which is -/// the spatial-only case. -/// -/// The structure matches `nlm_fused_pair_accumulate_window`, but this -/// kernel takes advantage of the weight map's symmetry. Patch distance -/// reads the same in either direction, so walking the full window one -/// way gives the same accumulator as the paired half-window version, -/// with one distance tile and one neighbour read per offset. +/// Accumulates a frame against itself over the whole search window, the spatial-only case. /// -/// The centre frame is cached in the expanded shared-memory tile, so -/// each offset only touches global memory for its shifted neighbour -/// pixel. +/// Patch distance is symmetric, so walking the full window one way gives the same result as the +/// paired half-window, with one distance tile and one neighbour read per offset. The zero offset is +/// skipped because `nlm_finish` adds it back through `wref * max_weight`. /// -/// The zero offset is skipped at compile time. `nlm_finish` folds that -/// self-contribution back in through its `wref * max_weight` term. -/// -/// # Spatial offset table -/// -/// `spatial_offset_lut` holds one noise-floor offset per candidate, laid -/// out row-major over the window. -/// -/// Nearby candidates share more of the grain's spatial correlation, so -/// their offset is reduced relative to distant ones. See -/// `noise::build_spatial_offset_lut`. +/// `spatial_offset_lut` holds one noise-floor offset per candidate, row-major over the window. +/// Nearby candidates share more of the grain's spatial correlation, so their offsets are smaller. #[cube(launch_unchecked)] pub fn nlm_fused_single_window( input: &Array>, @@ -514,13 +456,14 @@ pub fn nlm_fused_single_window( let mut idx = thread_id; while idx < expanded_elems { - let ex = idx % expanded_width; - let ey = idx / expanded_width; - let src_x = expanded_x0 + ex as i32; - let src_y = expanded_y0 + ey as i32; + let expanded_x = idx % expanded_width; + let expanded_y = idx / expanded_width; + let src_x = expanded_x0 + expanded_x as i32; + let src_y = expanded_y0 + expanded_y as i32; smem_center[idx as usize] = read_clamped_line(input, src_x, src_y, frame_t, width, height); idx += threads; } + sync_cube(); let mut accum_reg = Vector::::empty(); @@ -536,17 +479,12 @@ pub fn nlm_fused_single_window( let q_x = q_xi as i32 - search_radius as i32; let q_y = q_yi as i32 - search_radius as i32; if comptime!(q_x == 0 && q_y == 0) { - // Skip the zero offset. `nlm_finish` puts that - // contribution back through `wref * max_weight`. - // - // CubeCL has no `continue` yet. This does not become a - // branch in the kernel, because it is optimised out at - // compile time. + // An empty comptime branch stands in for `continue`, which cubecl lacks. } else { - let mut tidx = thread_id; - while tidx < tile_elems { - let tile_x = tidx % tile_width; - let tile_y = tidx / tile_width; + let mut tile_idx = thread_id; + while tile_idx < tile_elems { + let tile_x = tile_idx % tile_width; + let tile_y = tile_idx / tile_width; let center_idx = ((tile_y + search_radius) * expanded_width + (tile_x + search_radius)) as usize; let center = smem_center[center_idx]; @@ -558,9 +496,10 @@ pub fn nlm_fused_single_window( width, height, ); - smem_dist[tidx as usize] = line_sum_sq(center - neighbor, channels) * scale; - tidx += threads; + smem_dist[tile_idx as usize] = line_sum_sq(center - neighbor, channels) * scale; + tile_idx += threads; } + sync_cube(); if in_image { @@ -568,6 +507,7 @@ pub fn nlm_fused_single_window( let center_tile_y = local_y + patch_radius; let patch_size = 2 * patch_radius + 1; let mut patch_sum = 0.0f32; + for offset_y in 0..patch_size { for offset_x in 0..patch_size { let smem_idx = ((center_tile_y - patch_radius + offset_y) * tile_width @@ -577,6 +517,7 @@ pub fn nlm_fused_single_window( patch_sum += smem_dist[smem_idx]; } } + let lut_idx = (q_yi * window_side + q_xi) as usize; let offset = spatial_offset_lut[lut_idx]; let weight = welsch_weight(patch_sum, h2_inv_norm, offset); @@ -612,12 +553,7 @@ pub fn nlm_fused_single_window( /// The reference-image version of `nlm_fused_pair_accumulate_window`. /// -/// Distances, both the cached centre tile and the per-offset -/// neighbours, are read from `reference`. The pixels being accumulated -/// still come from `input`, so the original values reach `accum` while -/// the weights come from the cleaner reference frames. -/// -/// Confidence weighting works the same way as in the plain version. +/// Distances are read from `reference` while the accumulated pixels still come from `input`. #[cube(launch_unchecked)] pub fn nlm_fused_pair_accumulate_window_ref( input: &Array>, @@ -669,16 +605,16 @@ pub fn nlm_fused_pair_accumulate_window_ref( let expanded_x0 = fwd_tile_x0 - search_radius as i32; let expanded_y0 = fwd_tile_y0 - search_radius as i32; - // Cache `reference[frame_t]` once. let mut idx = thread_id; while idx < expanded_elems { - let ex = idx % expanded_width; - let ey = idx / expanded_width; - let src_x = expanded_x0 + ex as i32; - let src_y = expanded_y0 + ey as i32; + let expanded_x = idx % expanded_width; + let expanded_y = idx / expanded_width; + let src_x = expanded_x0 + expanded_x as i32; + let src_y = expanded_y0 + expanded_y as i32; smem_center[idx as usize] = read_clamped_line(reference, src_x, src_y, frame_t, width, height); idx += threads; } + sync_cube(); let mut accum_reg = Vector::::empty(); @@ -699,9 +635,6 @@ pub fn nlm_fused_pair_accumulate_window_ref( let tile_x = idx % tile_width; let tile_y = idx / tile_width; - // Both the forward and backward comparisons centre on - // the same patch of the centre frame. The plain - // variant's doc comment explains why. let center_idx = ((tile_y + search_radius) * expanded_width + (tile_x + search_radius)) as usize; let center = smem_center[center_idx]; @@ -736,6 +669,7 @@ pub fn nlm_fused_pair_accumulate_window_ref( let patch_size = 2 * patch_radius + 1; let mut sum_fwd = 0.0f32; let mut sum_bwd = 0.0f32; + for offset_y in 0..patch_size { for offset_x in 0..patch_size { let smem_idx = ((center_tile_y - patch_radius + offset_y) * tile_width @@ -751,14 +685,13 @@ pub fn nlm_fused_pair_accumulate_window_ref( let mut weight_bwd = welsch_weight(sum_bwd, h2_inv_norm, noise_offset); if use_confidence { - let bx = (global_x / step).min(blocks_x - 1); - let by = (global_y / step).min(blocks_y - 1); - let block_idx = (by * blocks_x + bx) as usize; + let block_col = (global_x / step).min(blocks_x - 1); + let block_row = (global_y / step).min(blocks_y - 1); + let block_idx = (block_row * blocks_x + block_col) as usize; weight_fwd *= conf_fwd[block_idx]; weight_bwd *= conf_bwd[block_idx]; } - // Pixel accumulation reads from `input`, not `reference`. let fwd_pixel = read_clamped_line( input, global_x as i32 + q_x, @@ -798,12 +731,7 @@ pub fn nlm_fused_pair_accumulate_window_ref( /// The reference-image version of `nlm_fused_single_window`. /// -/// Distances come from the reference image, both the cached centre and -/// the per-offset neighbours, while the pixels being accumulated come -/// from the input. -/// -/// `spatial_offset_lut` has the same layout as in -/// `nlm_fused_single_window`. +/// Distances are read from `reference` while the accumulated pixels still come from `input`. #[cube(launch_unchecked)] pub fn nlm_fused_single_window_ref( input: &Array>, @@ -848,13 +776,14 @@ pub fn nlm_fused_single_window_ref( let mut idx = thread_id; while idx < expanded_elems { - let ex = idx % expanded_width; - let ey = idx / expanded_width; - let src_x = expanded_x0 + ex as i32; - let src_y = expanded_y0 + ey as i32; + let expanded_x = idx % expanded_width; + let expanded_y = idx / expanded_width; + let src_x = expanded_x0 + expanded_x as i32; + let src_y = expanded_y0 + expanded_y as i32; smem_center[idx as usize] = read_clamped_line(reference, src_x, src_y, frame_t, width, height); idx += threads; } + sync_cube(); let mut accum_reg = Vector::::empty(); @@ -870,14 +799,12 @@ pub fn nlm_fused_single_window_ref( let q_x = q_xi as i32 - search_radius as i32; let q_y = q_yi as i32 - search_radius as i32; if comptime!(q_x == 0 && q_y == 0) { - // CubeCL has no `continue` yet. This does not become a - // branch in the kernel, because it is optimised out at - // compile time. + // An empty comptime branch stands in for `continue`, which cubecl lacks. } else { - let mut tidx = thread_id; - while tidx < tile_elems { - let tile_x = tidx % tile_width; - let tile_y = tidx / tile_width; + let mut tile_idx = thread_id; + while tile_idx < tile_elems { + let tile_x = tile_idx % tile_width; + let tile_y = tile_idx / tile_width; let center_idx = ((tile_y + search_radius) * expanded_width + (tile_x + search_radius)) as usize; let center = smem_center[center_idx]; @@ -889,9 +816,10 @@ pub fn nlm_fused_single_window_ref( width, height, ); - smem_dist[tidx as usize] = line_sum_sq(center - neighbor, channels) * scale; - tidx += threads; + smem_dist[tile_idx as usize] = line_sum_sq(center - neighbor, channels) * scale; + tile_idx += threads; } + sync_cube(); if in_image { @@ -899,6 +827,7 @@ pub fn nlm_fused_single_window_ref( let center_tile_y = local_y + patch_radius; let patch_size = 2 * patch_radius + 1; let mut patch_sum = 0.0f32; + for offset_y in 0..patch_size { for offset_x in 0..patch_size { let smem_idx = ((center_tile_y - patch_radius + offset_y) * tile_width @@ -908,6 +837,7 @@ pub fn nlm_fused_single_window_ref( patch_sum += smem_dist[smem_idx]; } } + let lut_idx = (q_yi * window_side + q_xi) as usize; let offset = spatial_offset_lut[lut_idx]; let weight = welsch_weight(patch_sum, h2_inv_norm, offset); diff --git a/av-denoise-core/src/nlmeans/kernels/helpers.rs b/av-denoise-core/src/nlmeans/kernels/helpers.rs index 956a87f..6a3f7e7 100644 --- a/av-denoise-core/src/nlmeans/kernels/helpers.rs +++ b/av-denoise-core/src/nlmeans/kernels/helpers.rs @@ -11,11 +11,9 @@ pub(super) fn clamp_coord(value: i32, #[comptime] limit: u32) -> u32 { result } -/// Reads the pixel at `(x, y)` in `frame`, clamped to the image edges on -/// both axes. +/// Reads the pixel at `(x, y)` in `frame`, clamped to the image edges. /// -/// The frame index is taken on trust, because callers always pass a -/// physical slot that holds loaded data. +/// Only the coordinates are clamped. `frame` must be a loaded physical slot. #[cube] pub(crate) fn read_clamped_line( buf: &Array>, @@ -33,8 +31,7 @@ pub(crate) fn read_clamped_line( /// The unchecked version of `read_clamped_line`. /// -/// The caller promises that `x` is inside `[0, width)` and `y` is inside -/// `[0, height)`. +/// `x` must be within `0..width` and `y` within `0..height`. #[cube] pub(crate) fn read_line( buf: &Array>, @@ -48,25 +45,22 @@ pub(crate) fn read_line( buf[idx as usize] } -/// Sums the squared differences across a vector's lanes. -/// -/// The loop unrolls fully at compile time, because `channels` is known -/// then. +/// Sums the squared differences across the first `channels` lanes. #[cube] pub(crate) fn line_sum_sq(diff: Vector, #[comptime] channels: u32) -> f32 { let mut sum = 0.0f32; + #[unroll] - for c in 0..channels { - sum += diff[c as usize] * diff[c as usize]; + for channel in 0..channels { + sum += diff[channel as usize] * diff[channel as usize]; } + sum } -/// The per-channel distance scale, which is 3 for luma, 1.5 for chroma, -/// and 1 for full YUV. +/// The per-channel distance scale, which is 3 for luma, 1.5 for chroma and 1 for full YUV. /// -/// Scaling this way lets all three channel modes share one -/// `h2_inv_norm`. +/// Scaling this way lets all three channel modes share one `h2_inv_norm`. #[cube] pub(crate) fn channel_scale(#[comptime] channels: u32) -> f32 { let mut scale = 1.0f32; @@ -80,26 +74,18 @@ pub(crate) fn channel_scale(#[comptime] channels: u32) -> f32 { /// The Welsch weight for a box-summed patch distance. /// -/// `noise_offset` is the distance two noisy copies of the same content -/// are expected to show. Subtracting it stops a good match being -/// penalised for the noise it carries. -/// -/// An offset of 0.0 gives exactly the plain weight, because the box sum -/// is never negative. +/// `noise_offset` is the distance two noisy copies of the same content are expected to show. +/// Subtracting it stops a good match being penalised for its noise. An offset of 0 gives the plain +/// weight, because the box sum is never negative. #[cube] pub(super) fn welsch_weight(sum: f32, h2_inv_norm: f32, noise_offset: f32) -> f32 { f32::exp(-f32::max(sum - noise_offset, 0.0) * h2_inv_norm) } -/// Adds the forward and backward neighbour contributions at the thread's -/// pixel. -/// -/// The forward neighbour sits at `(global + q, frame_fwd)` with -/// `weight_fwd`, and the backward one at `(global - q, frame_bwd)` with -/// `weight_bwd`. +/// Adds the forward and backward neighbour contributions at the thread's pixel. /// -/// One interior check per thread covers both reads, falling back to -/// clamped reads at the border. +/// The forward neighbour sits at `global + q` in `frame_fwd` and the backward one at `global - q` +/// in `frame_bwd`. Both reads fall back to clamped reads at the border. #[cube] pub(super) fn accumulate_pair( input: &Array>, @@ -148,8 +134,8 @@ pub(super) fn accumulate_pair( let line_w_fwd = Vector::::empty().fill(weight_fwd); let line_w_bwd = Vector::::empty().fill(weight_bwd); - let cur = accum[pixel_idx]; - accum[pixel_idx] = cur + fwd_pixel * line_w_fwd + bwd_pixel * line_w_bwd; + let cur_accum = accum[pixel_idx]; + accum[pixel_idx] = cur_accum + fwd_pixel * line_w_fwd + bwd_pixel * line_w_bwd; weight_sum[pixel_idx] += weight_fwd + weight_bwd; } diff --git a/av-denoise-core/src/nlmeans/kernels/memory.rs b/av-denoise-core/src/nlmeans/kernels/memory.rs index cd3853d..3a26163 100644 --- a/av-denoise-core/src/nlmeans/kernels/memory.rs +++ b/av-denoise-core/src/nlmeans/kernels/memory.rs @@ -1,16 +1,10 @@ use cubecl::prelude::*; -/// Copies `length` elements from `src[src_offset..]` into -/// `dst[dst_offset..]`, entirely on the GPU. +/// Copies `length` elements from `src[src_offset..]` into `dst[dst_offset..]`. /// -/// The loop is strided so the grid can stay under the 65,535 dispatch -/// limit. -/// -/// The offsets are kernel arguments rather than byte offsets on the -/// bound handles, which lets a caller address any slot of a ring buffer -/// whatever its stride. The GPU only accepts a buffer bound at a -/// multiple of its `min_storage_buffer_offset_alignment`, and a -/// `width * height * stored_ch` frame stride rarely lands on one. +/// The loop is strided so the grid stays under the 65,535 workgroup limit. The offsets are kernel +/// arguments because a buffer can only be bound at a multiple of +/// `min_storage_buffer_offset_alignment`, and a ring slot's stride rarely lands on one. #[cube(launch_unchecked)] pub fn gpu_copy( src: &Array, @@ -27,11 +21,10 @@ pub fn gpu_copy( } } -/// Zeroes `accum`, `weight_sum`, and `max_weight` in one dispatch. +/// Zeroes `accum`, `weight_sum` and `max_weight` in one dispatch. /// -/// The main loop covers all three up to `weight_len`, then a tail loop -/// finishes the channel-padded remainder of `accum`, which is always at -/// least as long as the other two. +/// `accum_len` must be at least `weight_len`, because a tail loop finishes the channel-padded +/// remainder of `accum`. #[cube(launch_unchecked)] pub fn gpu_zero_buffers( accum: &mut Array, @@ -55,150 +48,3 @@ pub fn gpu_zero_buffers( idx += total_threads; } } - -/// Quantises a denoised frame into wire bytes, packed into `u32` words. -/// -/// `src` holds `pixels * stored_ch` values. The lanes between `channels` -/// and `stored_ch` are padding and are skipped here, so the host reads -/// back only the samples it asked for. -/// -/// `outer` and `split_planes` decide how an output sample index maps -/// back into `src`. Interleaved output passes `outer = channels`, so the -/// quotient is the pixel and the remainder is the channel. Split output, -/// which is what a chroma pair needs, passes `outer = pixels` and gets -/// the reverse, laying each channel down as one contiguous region. -/// -/// `samples_per_word` is 4 for 8-bit wire and 2 for 10 and 12-bit, which -/// matches the host's `Narrow` and `Wide` codecs. -/// -/// A NaN sample does not come out as zero the way the host converter's -/// clamp does. The GPU clamp lowers to a min and max pair whose NaN -/// result is unspecified, and the float to integer cast is undefined on -/// NaN, so such a sample lands on an arbitrary byte. A denoised frame -/// holds no NaN, so nothing guards against it here. -/// -/// Every index is clamped rather than guarded by a branch. A -/// branch-derived index inside an unrolled loop makes cubecl's GVN pass -/// panic while compiling the shader, after which the launch silently -/// writes nothing. -/// -/// The loop is strided so the grid can stay under the dispatch limit. -#[cube(launch_unchecked)] -#[expect( - clippy::too_many_arguments, - reason = "every argument is a comptime shape the kernel specialises on" -)] -pub fn gpu_pack_wire( - src: &Array, - dst: &mut Array, - max: f32, - #[comptime] pixels: u32, - #[comptime] channels: u32, - #[comptime] stored_ch: u32, - #[comptime] outer: u32, - #[comptime] split_planes: bool, - #[comptime] samples_per_word: u32, - #[comptime] words: u32, - #[comptime] total_threads: u32, -) { - let samples = comptime![pixels * channels]; - let bits = comptime![32u32 / samples_per_word]; - - let mut word = ABSOLUTE_POS_X; - - while word < words { - let base = word * samples_per_word; - let mut acc = 0u32; - - #[unroll] - for lane in 0..samples_per_word { - let s = base + lane; - // Clamped, never branched. A lane past the last sample still - // reads a valid slot, and its value is dropped below. - let safe = u32::min(s, samples - 1); - - let a = safe / outer; - let b = safe % outer; - let src_idx = select(split_planes, b * stored_ch + a, a * stored_ch + b); - - let v = f32::clamp(src[src_idx as usize], 0.0, 1.0); - let q = u32::cast_from(v * max + 0.5); - - acc |= select(s < samples, q, 0u32) << (lane * bits); - } - - dst[word as usize] = acc; - word += total_threads; - } -} - -/// Normalises a wire-byte frame into one ring slot, interleaving its -/// planes and zeroing the padding lanes. -/// -/// `src` holds the planes concatenated in channel order, each `pixels` -/// samples long, packed `samples_per_word` to a `u32`. So a sample sits -/// at index `channel * pixels + pixel`. -/// -/// `src` reads whole words, so a plane whose sample count is not a whole -/// number of words needs its last word backed by real storage. The caller -/// pads the upload up to a multiple of four bytes, otherwise the kernel -/// reads past the allocation. -/// -/// `dst` is the ring buffer. This launch writes `elements` values from -/// `dst_offset`, which is the slot's own base. -/// -/// `max` is the depth's largest sample value, which each sample divides -/// by. The GPU divide is not correctly rounded, so a sample can land one -/// unit in the last place from what [`crate::frame::plane_to_f32`] -/// produces for the same byte. -/// -/// A channel at or above `channels` is a padding lane and comes out zero, -/// which is what keeps a `Yuv` frame's fourth lane from carrying a -/// previous frame's value. -/// -/// Every index is clamped rather than guarded by a branch. A -/// branch-derived index inside an unrolled loop makes cubecl's GVN pass -/// panic while compiling the shader, after which the launch silently -/// writes nothing. -/// -/// The loop is strided so the grid can stay under the dispatch limit. -#[cube(launch_unchecked)] -#[expect( - clippy::too_many_arguments, - reason = "every argument is a comptime shape the kernel specialises on" -)] -pub fn gpu_unpack_wire( - src: &Array, - dst: &mut Array, - max: f32, - dst_offset: u32, - #[comptime] pixels: u32, - #[comptime] channels: u32, - #[comptime] stored_ch: u32, - #[comptime] samples_per_word: u32, - #[comptime] elements: u32, - #[comptime] total_threads: u32, -) { - let bits = comptime![32u32 / samples_per_word]; - let mask = comptime![(1u32 << (32u32 / samples_per_word)) - 1]; - let wire_samples = comptime![pixels * channels]; - - let mut idx = ABSOLUTE_POS_X; - - while idx < elements { - let pixel = idx / stored_ch; - let ch = idx % stored_ch; - - // Clamped, never branched. A padding lane still reads a real - // sample and then discards it below. - let s = u32::min(ch * pixels + pixel, wire_samples - 1); - - let word = src[(s / samples_per_word) as usize]; - let sample = (word >> ((s % samples_per_word) * bits)) & mask; - - let v = f32::cast_from(sample) / max; - dst[(dst_offset + idx) as usize] = select(ch < channels, v, 0.0f32); - - idx += total_threads; - } -} diff --git a/av-denoise-core/src/nlmeans/kernels/mod.rs b/av-denoise-core/src/nlmeans/kernels/mod.rs index ad6f1e7..d0ce7ac 100644 --- a/av-denoise-core/src/nlmeans/kernels/mod.rs +++ b/av-denoise-core/src/nlmeans/kernels/mod.rs @@ -1,24 +1,3 @@ -//! The GPU kernels the denoiser launches. -//! -//! Everything here is cubecl code that runs on the device. The host side -//! decides which kernels to launch and with what arguments, which is the -//! `dispatch` module's job. -//! -//! A denoise pass measures how alike patches are, turns those distances -//! into weights, accumulates the weighted pixels, and then normalises -//! the result. -//! -//! `fused` does the first three steps in one kernel, which is the fast -//! path for small patches. `separable` splits the distance step into -//! horizontal and vertical passes so the cost stays linear as patches -//! grow. `accumulate` holds the shared accumulation and normalisation -//! steps. -//! -//! `bilateral` is the prefilter, `noise` measures the noise level, -//! [`motion`] tracks movement between frames, `memory` holds the copy, -//! zero, and wire packing and unpacking utilities, and `helpers` holds -//! the small pieces the kernels share. - mod accumulate; mod bilateral; mod fused; @@ -28,9 +7,9 @@ pub mod motion; mod noise; mod separable; -pub use accumulate::{nlm_accumulate, nlm_finish}; -pub use bilateral::nlm_bilateral; -pub use fused::{ +pub use self::accumulate::{nlm_accumulate, nlm_finish}; +pub use self::bilateral::nlm_bilateral; +pub use self::fused::{ nlm_dist_2d_weight, nlm_dist_2d_weight_ref, nlm_fused_pair_accumulate_window, @@ -38,9 +17,14 @@ pub use fused::{ nlm_fused_single_window, nlm_fused_single_window_ref, }; -pub use memory::{gpu_copy, gpu_pack_wire, gpu_unpack_wire, gpu_zero_buffers}; -pub use noise::{nlm_noise_partial, nlm_noise_reduce, nlm_temporal_noise_stats, nlm_temporal_stats_zero}; -pub use separable::{ +pub use self::memory::{gpu_copy, gpu_zero_buffers}; +pub use self::noise::{ + nlm_noise_partial, + nlm_noise_reduce, + nlm_temporal_noise_stats, + nlm_temporal_stats_zero, +}; +pub use self::separable::{ nlm_distance, nlm_distance_pair, nlm_distance_pair_ref, diff --git a/av-denoise-core/src/nlmeans/kernels/motion/block_match.rs b/av-denoise-core/src/nlmeans/kernels/motion/block_match.rs index 69307fd..28bd3ca 100644 --- a/av-denoise-core/src/nlmeans/kernels/motion/block_match.rs +++ b/av-denoise-core/src/nlmeans/kernels/motion/block_match.rs @@ -1,45 +1,15 @@ use cubecl::prelude::*; use cubecl::terminate; -/// Finds where each block of pixels moved to, by searching one level of -/// the luma pyramid for the offset with the smallest sum of absolute -/// differences. +/// Finds each block's motion on one pyramid level and seeds the fine grid with it. /// -/// One GPU block handles one image block. Its threads scan the -/// `(2 * search_radius + 1)^2` candidate offsets between them and -/// produce a single motion vector, in this level's pixels, at the -/// block's slot. +/// One GPU block handles one image block and picks the lowest SAD among the +/// `(2 * search_radius + 1)^2` candidates. `blksize` is capped at +/// [MAX_BLKSIZE](crate::nlmeans::motion::MAX_BLKSIZE) so the shared centre tile stays within 1024 +/// values. /// -/// # How the search is split up -/// -/// The block first loads its own `blksize x blksize` centre pixels into -/// shared memory once, splitting the rows across threads. -/// -/// Each thread then takes a strided share of the candidate offsets and -/// works out each one's full score on its own, reading the centre from -/// shared memory and the neighbour from global memory. It writes each of -/// its candidates' scores exactly once. -/// -/// Because every scratch slot has exactly one writer, no atomics and no -/// second reduction pass are needed. After a `sync_cube`, thread 0 walks -/// the scratch buffer and picks the winner. -/// -/// `blksize` is capped at [`MAX_BLKSIZE`], which is 32, so the shared -/// centre tile never exceeds 1024 values. -/// -/// # Seeding the fine pass -/// -/// `level_scale` rescales the motion vector into fine-level pixels. The -/// caller passes 2 raised to the coarse level. -/// -/// After finding its own winner, the block seeds every fine block whose -/// position falls inside its own source region. `step` and `fine_step` -/// are the coarse and fine grids' block spacings, which is what converts -/// between the two index spaces. -/// -/// The seeding code below spells out the exact formula. -/// -/// [`MAX_BLKSIZE`]: crate::nlmeans::motion::MAX_BLKSIZE +/// The winner is scaled by `level_scale`, 2 raised to the coarse level, and written to every fine +/// block inside this block's region. `step` and `fine_step` are the coarse and fine block spacings. #[cube(launch_unchecked)] pub fn nlm_mc_block_match_coarse( centre: &Array, @@ -55,11 +25,11 @@ pub fn nlm_mc_block_match_coarse( #[comptime] fine_blocks_y: u32, #[comptime] fine_step: u32, ) { - let bx = CUBE_POS_X; - let by = CUBE_POS_Y; + let block_col = CUBE_POS_X; + let block_row = CUBE_POS_Y; - let block_origin_x = bx as i32 * step as i32; - let block_origin_y = by as i32 * step as i32; + let block_origin_x = block_col as i32 * step as i32; + let block_origin_y = block_row as i32 * step as i32; let local_x = UNIT_POS_X; let local_y = UNIT_POS_Y; @@ -72,28 +42,22 @@ pub fn nlm_mc_block_match_coarse( let mut sad_scratch = SharedMemory::::new(candidates as usize); let mut centre_smem = SharedMemory::::new(block_pixels as usize); - // A cooperative load. Each thread claims a row-strided share of the - // block's pixels and caches the clamped centre value once, so every - // candidate below reads it from shared memory instead of fetching - // it again from global memory. - let mut py = local_y; - while py < blksize { - let mut px = local_x; - while px < blksize { - let cx_c = clamp_i32(block_origin_x + px as i32, level_width as i32); - let cy_c = clamp_i32(block_origin_y + py as i32, level_height as i32); - centre_smem[(py * blksize + px) as usize] = centre[(cy_c * level_width as i32 + cx_c) as usize]; - px += CUBE_DIM_X; + let mut pixel_y = local_y; + while pixel_y < blksize { + let mut pixel_x = local_x; + while pixel_x < blksize { + let clamped_x = clamp_i32(block_origin_x + pixel_x as i32, level_width as i32); + let clamped_y = clamp_i32(block_origin_y + pixel_y as i32, level_height as i32); + centre_smem[(pixel_y * blksize + pixel_x) as usize] = + centre[(clamped_y * level_width as i32 + clamped_x) as usize]; + pixel_x += CUBE_DIM_X; } - py += CUBE_DIM_Y; + + pixel_y += CUBE_DIM_Y; } + sync_cube(); - // Each thread takes a strided share of the candidates, works out - // each one's full score over the cached block, and writes its - // scratch slot exactly once. Splitting the work by candidate rather - // than by pixel means no two threads ever write the same slot, so - // no atomics and no reduce pass are needed. let mut candidate_idx = thread_id; while candidate_idx < candidates { let dy = candidate_idx / window_side; @@ -104,158 +68,106 @@ pub fn nlm_mc_block_match_coarse( let mut sad = 0.0f32; for iy in 0..blksize { for ix in 0..blksize { - let cx = block_origin_x + ix as i32; - let cy = block_origin_y + iy as i32; + let centre_x = block_origin_x + ix as i32; + let centre_y = block_origin_y + iy as i32; let centre_val = centre_smem[(iy * blksize + ix) as usize]; - let nx = clamp_i32(cx + mvx, level_width as i32); - let ny = clamp_i32(cy + mvy, level_height as i32); - let neighbour_val = neighbour[(ny * level_width as i32 + nx) as usize]; + let neighbour_x = clamp_i32(centre_x + mvx, level_width as i32); + let neighbour_y = clamp_i32(centre_y + mvy, level_height as i32); + let neighbour_val = neighbour[(neighbour_y * level_width as i32 + neighbour_x) as usize]; let diff = centre_val - neighbour_val; let abs_diff = if diff < 0.0f32 { -diff } else { diff }; sad += abs_diff; } } + sad_scratch[candidate_idx as usize] = sad; candidate_idx += threads; } + sync_cube(); if thread_id != 0 { terminate!(); } - // Walk the candidates and keep the lowest score. `window_side` is - // known at compile time, so this unrolls cleanly for small search - // radii. - // - // The starting value is deliberately huge so the first iteration - // always wins, which avoids a negative-initialiser pattern cubecl's - // macro does not lift cleanly. + // A huge start rather than a negative initialiser, which cubecl's macro does not lift cleanly. let mut best_sad = 1.0e30f32; let mut best_dx = 0i32; let mut best_dy = 0i32; + for dy in 0..window_side { for dx in 0..window_side { - let s = sad_scratch[(dy * window_side + dx) as usize]; - if s < best_sad { - best_sad = s; + let candidate_sad = sad_scratch[(dy * window_side + dx) as usize]; + if candidate_sad < best_sad { + best_sad = candidate_sad; best_dx = dx as i32 - search_radius as i32; best_dy = dy as i32 - search_radius as i32; } } } - // An exact tie resolves to the zero-motion candidate, rather than - // whichever candidate the scan above happened to reach first, which - // is the window corner. - // - // A block lying entirely inside a flat region scores the same at - // every candidate. Reporting motion there would seed the fine pass, - // and every block it covers, from a shifted position that no pixel - // comparison ever preferred. + // A tie resolves to zero motion, so a flat block never seeds the fine pass from a shifted + // position that no comparison preferred. let zero_sad = sad_scratch[(search_radius * window_side + search_radius) as usize]; if zero_sad <= best_sad { best_dx = 0i32; best_dy = 0i32; } - // Map the coarse block index into the fine block index space by - // position, rather than by simply doubling the index. - // - // Coarse block `bx` covers source pixels from `bx * step` up to - // `(bx + 1) * step` at the coarse level, which at the fine level - // means those bounds multiplied by `level_scale`. Dividing that - // span by `fine_step` gives the range of fine blocks this coarse - // block seeds. - // - // Floor-division tiling leaves every interior boundary touching, - // with no gap and no overlap. The two block counts are each their - // own ceil-division over a different width and a different step - // though, so they can round differently, and the coarse grid's - // nominal reach can fall just short of the fine grid's true edge. - // - // The last coarse block on each axis therefore extends its end to - // the fine grid's own edge, absorbing that remainder so every fine - // block still gets seeded exactly once. - // - // One formula covers all three cases, a genuinely coarser grid, two - // equal grids, and the ragged last block on either geometry, with - // no separate code path. - // - // How many fine blocks a coarse block seeds can vary from block to - // block, so this is a runtime `while` loop rather than an unrolled - // `for` loop. + // Coarse blocks map to fine blocks by position, `col * step * level_scale / fine_step`. The two + // block counts can round differently, so the last coarse block on each axis extends to the fine + // grid's edge and every fine block is seeded exactly once. The count varies per block, so the + // loops below are runtime `while` loops rather than unrolled. let mvx_fine = best_dx * level_scale as i32; let mvy_fine = best_dy * level_scale as i32; - let fbx_start = (bx * step * level_scale / fine_step).min(fine_blocks_x); + let fine_col_start = (block_col * step * level_scale / fine_step).min(fine_blocks_x); #[expect( clippy::useless_conversion, reason = "both branches have to expand to the same cubecl native type, which the \ conversion supplies" )] - let fbx_end = if bx == CUBE_COUNT_X - 1 { + let fine_col_end = if block_col == CUBE_COUNT_X - 1 { fine_blocks_x.into() } else { - ((bx + 1) * step * level_scale / fine_step).min(fine_blocks_x) + ((block_col + 1) * step * level_scale / fine_step).min(fine_blocks_x) }; - let fby_start = (by * step * level_scale / fine_step).min(fine_blocks_y); + let fine_row_start = (block_row * step * level_scale / fine_step).min(fine_blocks_y); #[expect( clippy::useless_conversion, reason = "both branches have to expand to the same cubecl native type, which the \ conversion supplies" )] - let fby_end = if by == CUBE_COUNT_Y - 1 { + let fine_row_end = if block_row == CUBE_COUNT_Y - 1 { fine_blocks_y.into() } else { - ((by + 1) * step * level_scale / fine_step).min(fine_blocks_y) + ((block_row + 1) * step * level_scale / fine_step).min(fine_blocks_y) }; - let mut fby = fby_start; - while fby < fby_end { - let mut fbx = fbx_start; - while fbx < fbx_end { - let idx = ((fby * fine_blocks_x + fbx) * 2) as usize; + let mut fine_row = fine_row_start; + while fine_row < fine_row_end { + let mut fine_col = fine_col_start; + while fine_col < fine_col_end { + let idx = ((fine_row * fine_blocks_x + fine_col) * 2) as usize; mv_field[idx] = mvx_fine; mv_field[idx + 1] = mvy_fine; - fbx += 1; + fine_col += 1; } - fby += 1; + + fine_row += 1; } } -/// Refines a motion estimate at full resolution. -/// -/// When `use_seed` is set, the block starts from the vector the coarse -/// pass left in `mv_field` and searches a small window around it. The -/// refined vector goes back into the same slot. -/// -/// The search itself works exactly like `nlm_mc_block_match_coarse`, -/// with the same cached centre tile and the same split by candidate. -/// -/// # Confidence -/// -/// When `write_confidence` is true, the block also writes a confidence -/// score derived from the winning score. That is what lets a poor match -/// suppress its own frame's contribution later on, instead of blurring -/// in content that does not belong. -/// -/// `sad_noise_floor` is the score two noisy copies of the same content -/// produce by chance. Subtracting it first stops a clean match being -/// penalised for the noise it carries. -/// -/// `thsad` is how far past that floor a block can go before its -/// confidence reaches zero. It has to be strictly positive whenever -/// `write_confidence` is true, or a perfect match divides zero by zero. +/// Refines each block's motion at full resolution and optionally scores its confidence. /// -/// A `search_radius` of 0 with no seed reduces the match to the single -/// unshifted candidate, which is useful for a confidence-only pass with -/// no motion search at all. +/// When `use_seed` is 1, the search window centres on the vector already in `mv_field`, and the +/// refined vector replaces it. A `search_radius` of 0 with no seed scores only the unshifted block. /// -/// When `write_confidence` is false, the whole confidence step is -/// dropped at compile time and costs nothing. Callers pass a small -/// placeholder buffer in that case, and its size never matters because -/// the kernel never reads it. +/// When `write_confidence` is set, each block also writes a confidence between 0 and 1 so a poor +/// match can suppress its frame. `sad_noise_floor` is the SAD two noisy copies of the same content +/// show, and `thsad` is how far past it the confidence reaches zero. `thsad` must be positive, or +/// a perfect match divides zero by zero. When it is unset `confidence` is never touched and can be +/// a placeholder. #[cube(launch_unchecked)] pub fn nlm_mc_block_match_fine( centre: &Array, @@ -273,15 +185,11 @@ pub fn nlm_mc_block_match_fine( use_seed: u32, #[comptime] blocks_x: u32, ) { - let bx = CUBE_POS_X; - let by = CUBE_POS_Y; + let block_col = CUBE_POS_X; + let block_row = CUBE_POS_Y; - let mv_slot = ((by * blocks_x + bx) * 2) as usize; + let mv_slot = ((block_row * blocks_x + block_col) * 2) as usize; - // Clippy flags these `.into()` calls as useless, but they are - // required. Both `if` branches have to produce the same cubecl - // `NativeExpand` type, and a bare `0i32` literal does not - // coerce inside the cube macro. #[expect( clippy::useless_conversion, reason = "both branches have to expand to the same cubecl native type, which the \ @@ -303,8 +211,8 @@ pub fn nlm_mc_block_match_fine( 0i32.into() }; - let block_origin_x = bx as i32 * step as i32; - let block_origin_y = by as i32 * step as i32; + let block_origin_x = block_col as i32 * step as i32; + let block_origin_y = block_row as i32 * step as i32; let local_x = UNIT_POS_X; let local_y = UNIT_POS_Y; @@ -317,28 +225,22 @@ pub fn nlm_mc_block_match_fine( let mut sad_scratch = SharedMemory::::new(candidates as usize); let mut centre_smem = SharedMemory::::new(block_pixels as usize); - // A cooperative load. Each thread claims a row-strided share of the - // block's pixels and caches the clamped centre value once, so every - // candidate below reads it from shared memory instead of fetching - // it again from global memory. - let mut py = local_y; - while py < blksize { - let mut px = local_x; - while px < blksize { - let cx_c = clamp_i32(block_origin_x + px as i32, width as i32); - let cy_c = clamp_i32(block_origin_y + py as i32, height as i32); - centre_smem[(py * blksize + px) as usize] = centre[(cy_c * width as i32 + cx_c) as usize]; - px += CUBE_DIM_X; + let mut pixel_y = local_y; + while pixel_y < blksize { + let mut pixel_x = local_x; + while pixel_x < blksize { + let clamped_x = clamp_i32(block_origin_x + pixel_x as i32, width as i32); + let clamped_y = clamp_i32(block_origin_y + pixel_y as i32, height as i32); + centre_smem[(pixel_y * blksize + pixel_x) as usize] = + centre[(clamped_y * width as i32 + clamped_x) as usize]; + pixel_x += CUBE_DIM_X; } - py += CUBE_DIM_Y; + + pixel_y += CUBE_DIM_Y; } + sync_cube(); - // Each thread takes a strided share of the candidates, works out - // each one's full score over the cached block, and writes its - // scratch slot exactly once. Splitting the work by candidate rather - // than by pixel means no two threads ever write the same slot, so - // no atomics and no reduce pass are needed. let mut candidate_idx = thread_id; while candidate_idx < candidates { let dy = candidate_idx / window_side; @@ -349,20 +251,22 @@ pub fn nlm_mc_block_match_fine( let mut sad = 0.0f32; for iy in 0..blksize { for ix in 0..blksize { - let cx = block_origin_x + ix as i32; - let cy = block_origin_y + iy as i32; + let centre_x = block_origin_x + ix as i32; + let centre_y = block_origin_y + iy as i32; let centre_val = centre_smem[(iy * blksize + ix) as usize]; - let nx = clamp_i32(cx + mvx, width as i32); - let ny = clamp_i32(cy + mvy, height as i32); - let neighbour_val = neighbour[(ny * width as i32 + nx) as usize]; + let neighbour_x = clamp_i32(centre_x + mvx, width as i32); + let neighbour_y = clamp_i32(centre_y + mvy, height as i32); + let neighbour_val = neighbour[(neighbour_y * width as i32 + neighbour_x) as usize]; let diff = centre_val - neighbour_val; let abs_diff = if diff < 0.0f32 { -diff } else { diff }; sad += abs_diff; } } + sad_scratch[candidate_idx as usize] = sad; candidate_idx += threads; } + sync_cube(); if thread_id != 0 { @@ -372,27 +276,20 @@ pub fn nlm_mc_block_match_fine( let mut best_sad = 1.0e30f32; let mut best_dx = seed_dx; let mut best_dy = seed_dy; + for dy in 0..window_side { for dx in 0..window_side { - let s = sad_scratch[(dy * window_side + dx) as usize]; - if s < best_sad { - best_sad = s; + let candidate_sad = sad_scratch[(dy * window_side + dx) as usize]; + if candidate_sad < best_sad { + best_sad = candidate_sad; best_dx = seed_dx + (dx as i32 - search_radius as i32); best_dy = seed_dy + (dy as i32 - search_radius as i32); } } } - // An exact tie resolves to the seed itself, adding no motion beyond - // whatever the coarse pass found, rather than to whichever - // candidate the scan above happened to reach first, which is the - // window corner. - // - // A block lying entirely inside a flat region scores the same at - // every candidate. Reporting motion there would warp in pixels no - // comparison ever preferred, and since the winning score is zero in - // that case, it would also write a perfect confidence for what may - // be a genuinely occluded block. + // A tie resolves to the seed. On a flat block every candidate ties at zero, and any other + // winner would warp in unpreferred pixels with a perfect confidence. let seed_sad = sad_scratch[(search_radius * window_side + search_radius) as usize]; if seed_sad <= best_sad { best_sad = seed_sad; @@ -408,13 +305,15 @@ pub fn nlm_mc_block_match_fine( if excess < 0.0f32 { excess = 0.0f32; } + let thsad_sq = thsad * thsad; let excess_sq = excess * excess; let mut confidence_val = (thsad_sq - excess_sq) / (thsad_sq + excess_sq); if confidence_val < 0.0f32 { confidence_val = 0.0f32; } - confidence[(by * blocks_x + bx) as usize] = confidence_val; + + confidence[(block_row * blocks_x + block_col) as usize] = confidence_val; } } diff --git a/av-denoise-core/src/nlmeans/kernels/motion/chain.rs b/av-denoise-core/src/nlmeans/kernels/motion/chain.rs index a721b91..208fefe 100644 --- a/av-denoise-core/src/nlmeans/kernels/motion/chain.rs +++ b/av-denoise-core/src/nlmeans/kernels/motion/chain.rs @@ -1,51 +1,19 @@ use cubecl::prelude::*; use cubecl::terminate; -/// Joins a run of adjacent-frame motion fields into one motion vector -/// per block, reaching from the centre frame to a distant neighbour. +/// Chains adjacent-frame motion fields into one vector per block, from the centre frame to a +/// distant neighbour. /// -/// Motion is only ever measured between neighbouring frames. To reach a -/// frame several steps away, this kernel follows the picture from one -/// frame to the next and adds up what it finds. +/// One thread walks one block for `steps` hops from `start_pair_slot`. Each hop reads the motion of +/// the block under the walking position and moves the position by it. Forward walks read +/// direction 0 and step to the next slot, backward walks read direction 1 and step to the previous +/// one. Consecutive hops use consecutive slots because the pair ring is keyed by the newer frame's +/// push order. Only the block lookup is clamped, so a chain can leave the frame without corrupting +/// later lookups. /// -/// The result goes into `mv_field` in the layout `nlm_mc_warp` and the -/// analyse kernels share. -/// -/// # The walk -/// -/// One thread handles one output block. -/// -/// Each hop finds the block nearest the walking position, clamped at the -/// edges exactly the way `nlm_mc_warp` maps a pixel to a block. It reads -/// that block's motion, adds it to the running total, and moves the -/// walking position by the same amount before the next hop. -/// -/// Only the block lookup clamps. The position itself keeps its true -/// value, so a long chain can wander outside the frame without -/// corrupting the lookups that follow. -/// -/// The walk starts at `start_pair_slot` and takes `steps` hops. Going -/// forward it reads direction 0 and moves to the next slot. Going -/// backward it reads direction 1 and moves to the previous one. -/// -/// Consecutive hops land on consecutive slots, because the pair ring is -/// keyed by the newer frame's place in the push sequence. See -/// `crate::nlmeans::motion::pair_ring_slot_count`. The caller therefore -/// only has to work out the first hop's slot. -/// -/// # Ring layout -/// -/// `pair_ring` holds every live adjacent-pair field, indexed by slot, -/// then direction, then block, then component. Direction 0 runs from the -/// older frame to the newer one, and direction 1 the other way. -/// -/// Each direction's slice is padded up to a 32-byte boundary, which is 8 -/// `i32` elements. -/// -/// `dir_len` and `slot_len` are those padded strides in elements, taken -/// from `MotionCtx::pair_direction_stride` and `pair_slot_stride`, not -/// the raw block count. The host writes at the padded offsets, so the -/// reads here have to use the same ones. +/// `pair_ring` is indexed by slot, direction, block, then component. Direction 0 runs from the +/// older frame to the newer one. `dir_len` and `slot_len` are the padded strides in elements, since +/// each direction is padded to a 32-byte boundary. #[cube(launch_unchecked)] pub fn nlm_mc_chain_compose( pair_ring: &Array, @@ -62,17 +30,17 @@ pub fn nlm_mc_chain_compose( #[comptime] blocks_x: u32, #[comptime] blocks_y: u32, ) { - let bx = ABSOLUTE_POS_X; - let by = ABSOLUTE_POS_Y; + let block_col = ABSOLUTE_POS_X; + let block_row = ABSOLUTE_POS_Y; - if bx >= blocks_x || by >= blocks_y { + if block_col >= blocks_x || block_row >= blocks_y { terminate!(); } let direction = comptime!(if forward { 0u32 } else { 1u32 }); - let mut pos_x = (bx * step + step / 2) as i32; - let mut pos_y = (by * step + step / 2) as i32; + let mut pos_x = (block_col * step + step / 2) as i32; + let mut pos_y = (block_row * step + step / 2) as i32; let mut acc_x = 0i32; let mut acc_y = 0i32; @@ -83,33 +51,29 @@ pub fn nlm_mc_chain_compose( (start_pair_slot + pair_ring_slots - i) % pair_ring_slots }; - let cx = clamp_i32(pos_x, width as i32) as u32; - let cy = clamp_i32(pos_y, height as i32) as u32; - let bxi = (cx / step).min(blocks_x - 1); - let byi = (cy / step).min(blocks_y - 1); + let clamped_x = clamp_i32(pos_x, width as i32) as u32; + let clamped_y = clamp_i32(pos_y, height as i32) as u32; + let hop_col = (clamped_x / step).min(blocks_x - 1); + let hop_row = (clamped_y / step).min(blocks_y - 1); - let base = slot * slot_len + direction * dir_len + (byi * blocks_x + bxi) * 2; - let fx = pair_ring[base as usize]; - let fy = pair_ring[(base + 1) as usize]; + let base = slot * slot_len + direction * dir_len + (hop_row * blocks_x + hop_col) * 2; + let hop_x = pair_ring[base as usize]; + let hop_y = pair_ring[(base + 1) as usize]; - acc_x += fx; - acc_y += fy; - pos_x += fx; - pos_y += fy; + acc_x += hop_x; + acc_y += hop_y; + pos_x += hop_x; + pos_y += hop_y; } - let out_idx = ((by * blocks_x + bx) * 2) as usize; + let out_idx = ((block_row * blocks_x + block_col) * 2) as usize; mv_field[out_idx] = acc_x; mv_field[out_idx + 1] = acc_y; } /// Fills both directions of one pair-ring slot with zeroes. /// -/// Duplicated ring slots, which appear while priming the stream and -/// again during the end-of-stream flush, hold the same content in both -/// frames. Their motion is zero by definition, so writing the zeroes is -/// both cheaper and exactly right compared with analysing identical -/// input. +/// A duplicated ring slot holds the same frame twice, so its motion is exactly zero. #[cube(launch_unchecked)] pub fn nlm_mc_pair_zero(dst: &mut Array, #[comptime] length: u32, #[comptime] total_threads: u32) { let mut idx = ABSOLUTE_POS_X; diff --git a/av-denoise-core/src/nlmeans/kernels/motion/downscale.rs b/av-denoise-core/src/nlmeans/kernels/motion/downscale.rs index bbd7907..c34b385 100644 --- a/av-denoise-core/src/nlmeans/kernels/motion/downscale.rs +++ b/av-denoise-core/src/nlmeans/kernels/motion/downscale.rs @@ -1,15 +1,10 @@ use cubecl::prelude::*; use cubecl::terminate; -/// Halves the size of a luma image by averaging each 2x2 group of -/// pixels. +/// Builds the next pyramid level by averaging each 2x2 group of luma pixels. /// -/// The source is one pyramid level, either full resolution or an -/// already-downscaled level, and the result goes into the next level's -/// slot. -/// -/// `src_frame` and `dst_frame` are slot indices inside the per-level -/// frame rings. +/// `src_frame` and `dst_frame` are slots in the per-level frame rings. An odd last row or column +/// repeats its edge pixel. #[cube(launch_unchecked)] pub fn nlm_mc_downscale( src: &Array, @@ -28,26 +23,22 @@ pub fn nlm_mc_downscale( terminate!(); } - let sx = x * 2; - let sy = y * 2; - let sx1 = if sx + 1 < src_width { sx + 1 } else { sx }; - let sy1 = if sy + 1 < src_height { sy + 1 } else { sy }; + let src_x = x * 2; + let src_y = y * 2; + let src_x1 = if src_x + 1 < src_width { src_x + 1 } else { src_x }; + let src_y1 = if src_y + 1 < src_height { src_y + 1 } else { src_y }; let src_base = src_frame * src_width * src_height; - let s00 = src[(src_base + sy * src_width + sx) as usize]; - let s10 = src[(src_base + sy * src_width + sx1) as usize]; - let s01 = src[(src_base + sy1 * src_width + sx) as usize]; - let s11 = src[(src_base + sy1 * src_width + sx1) as usize]; + let top_left = src[(src_base + src_y * src_width + src_x) as usize]; + let top_right = src[(src_base + src_y * src_width + src_x1) as usize]; + let bottom_left = src[(src_base + src_y1 * src_width + src_x) as usize]; + let bottom_right = src[(src_base + src_y1 * src_width + src_x1) as usize]; - let avg = (s00 + s10 + s01 + s11) * 0.25f32; + let avg = (top_left + top_right + bottom_left + bottom_right) * 0.25f32; dst[(dst_frame * dst_width * dst_height + y * dst_width + x) as usize] = avg; } -/// Copies the luma plane out of a packed input frame into a flat array, -/// which becomes level 0 of the pyramid. -/// -/// This lets the pyramid and analyse kernels work on luma alone, without -/// knowing anything about the channel layout. +/// Copies lane 0 of a packed frame into a flat luma array, which becomes level 0 of the pyramid. #[cube(launch_unchecked)] pub fn nlm_mc_extract_luma( src: &Array>, diff --git a/av-denoise-core/src/nlmeans/kernels/motion/mod.rs b/av-denoise-core/src/nlmeans/kernels/motion/mod.rs index 1ffc94f..d2a4498 100644 --- a/av-denoise-core/src/nlmeans/kernels/motion/mod.rs +++ b/av-denoise-core/src/nlmeans/kernels/motion/mod.rs @@ -1,24 +1,9 @@ -//! The GPU kernels that track motion between frames. -//! -//! Temporal denoising averages a pixel with the same position in nearby -//! frames. When the picture moves, that position holds different content -//! in each frame, which blurs moving edges. -//! -//! These kernels work out where each block of pixels moved to, then -//! shift the neighbouring frames back into line before the denoising -//! weights are computed. -//! -//! `downscale` builds the image pyramid, `block_match` searches each -//! level for the best match, `chain` joins adjacent-frame results into -//! one longer motion vector, and `warp` shifts a frame by the field that -//! comes out. - mod block_match; mod chain; mod downscale; mod warp; -pub use block_match::{nlm_mc_block_match_coarse, nlm_mc_block_match_fine}; -pub use chain::{nlm_mc_chain_compose, nlm_mc_pair_zero}; -pub use downscale::{nlm_mc_downscale, nlm_mc_extract_luma}; -pub use warp::nlm_mc_warp; +pub use self::block_match::{nlm_mc_block_match_coarse, nlm_mc_block_match_fine}; +pub use self::chain::{nlm_mc_chain_compose, nlm_mc_pair_zero}; +pub use self::downscale::{nlm_mc_downscale, nlm_mc_extract_luma}; +pub use self::warp::nlm_mc_warp; diff --git a/av-denoise-core/src/nlmeans/kernels/motion/warp.rs b/av-denoise-core/src/nlmeans/kernels/motion/warp.rs index 8759847..0cd3347 100644 --- a/av-denoise-core/src/nlmeans/kernels/motion/warp.rs +++ b/av-denoise-core/src/nlmeans/kernels/motion/warp.rs @@ -1,25 +1,11 @@ use cubecl::prelude::*; use cubecl::terminate; -/// Shifts a neighbour frame into line with the centre frame, using the -/// per-block motion field the analyse pass produced. +/// Shifts a neighbour frame into line with the centre frame using a per-block motion field. /// -/// Each output pixel works out which block it belongs to, reads that -/// block's motion vector, and takes its source pixel from the offset -/// position, clamped at the borders. Padding lanes are copied straight -/// through. -/// -/// # Overlapping blocks -/// -/// A pixel inside a block's interior takes that block's vector -/// directly. -/// -/// A pixel in the band where two adjacent blocks overlap, which happens -/// when the step is smaller than the block size, takes the vector of -/// whichever block is closest. Nothing is blended between them. -/// -/// That winner-takes-all rule is a simplification of MVTools' -/// raised-cosine blend. +/// Aligning the neighbours stops temporal averaging from blurring moving edges. Each pixel takes +/// the vector of block `pixel / step`, clamped to the grid, and reads its source pixel clamped at +/// the borders. Overlapping blocks are not blended. #[cube(launch_unchecked)] pub fn nlm_mc_warp( src: &Array>, @@ -40,17 +26,17 @@ pub fn nlm_mc_warp( terminate!(); } - let bx = (x / step).min(blocks_x - 1); - let by = (y / step).min(blocks_y - 1); + let block_col = (x / step).min(blocks_x - 1); + let block_row = (y / step).min(blocks_y - 1); - let mv_idx = ((by * blocks_x + bx) * 2) as usize; + let mv_idx = ((block_row * blocks_x + block_col) * 2) as usize; let mvx = mv_field[mv_idx]; let mvy = mv_field[mv_idx + 1]; - let sx = clamp_pos(x as i32 + mvx, width as i32); - let sy = clamp_pos(y as i32 + mvy, height as i32); + let src_x = clamp_pos(x as i32 + mvx, width as i32); + let src_y = clamp_pos(y as i32 + mvy, height as i32); - let src_idx = (src_frame * height + sy as u32) * width + sx as u32; + let src_idx = (src_frame * height + src_y as u32) * width + src_x as u32; let dst_idx = (dst_frame * height + y) * width + x; dst[dst_idx as usize] = src[src_idx as usize]; } diff --git a/av-denoise-core/src/nlmeans/kernels/noise.rs b/av-denoise-core/src/nlmeans/kernels/noise.rs index 813666e..104d9e0 100644 --- a/av-denoise-core/src/nlmeans/kernels/noise.rs +++ b/av-denoise-core/src/nlmeans/kernels/noise.rs @@ -18,16 +18,9 @@ use crate::nlmeans::noise::{ /// The per-block stage of the Immerkær noise estimate. /// -/// Every interior thread applies a 3x3 mask to its pixel, which cancels -/// out smooth content and leaves mostly noise. The block then sums the -/// absolute responses per channel into one partial total. -/// -/// Border pixels contribute zero, as do threads that land outside the -/// image because the grid overshoots on the last row or column of -/// blocks. -/// -/// Results are written as `partials[block_index * 4 + lane]`, with any -/// unused lane left at zero. +/// Each interior pixel applies a 3x3 mask that cancels smooth content and leaves mostly noise. The +/// block sums the absolute responses per channel into `partials[block_index * 4 + lane]`. Border +/// pixels, threads past the image and unused lanes contribute zero. #[cube(launch_unchecked)] pub fn nlm_noise_partial( input: &Array>, @@ -44,55 +37,58 @@ pub fn nlm_noise_partial( let x = ABSOLUTE_POS_X; let y = ABSOLUTE_POS_Y; - let tid = UNIT_POS_Y * block_x + UNIT_POS_X; + let thread_id = UNIT_POS_Y * block_x + UNIT_POS_X; let interior = x >= 1 && x < width - 1 && y >= 1 && y < height - 1; let mut response = Vector::::empty(); if interior { - let c = read_line(input, x, y, frame, width, height); - let l = read_line(input, x - 1, y, frame, width, height); - let r = read_line(input, x + 1, y, frame, width, height); - let u = read_line(input, x, y - 1, frame, width, height); - let d = read_line(input, x, y + 1, frame, width, height); - let ul = read_line(input, x - 1, y - 1, frame, width, height); - let ur = read_line(input, x + 1, y - 1, frame, width, height); - let dl = read_line(input, x - 1, y + 1, frame, width, height); - let dr = read_line(input, x + 1, y + 1, frame, width, height); + let centre = read_line(input, x, y, frame, width, height); + let left = read_line(input, x - 1, y, frame, width, height); + let right = read_line(input, x + 1, y, frame, width, height); + let up = read_line(input, x, y - 1, frame, width, height); + let down = read_line(input, x, y + 1, frame, width, height); + let up_left = read_line(input, x - 1, y - 1, frame, width, height); + let up_right = read_line(input, x + 1, y - 1, frame, width, height); + let down_left = read_line(input, x - 1, y + 1, frame, width, height); + let down_right = read_line(input, x + 1, y + 1, frame, width, height); let four = Vector::::empty().fill(4.0f32); let two = Vector::::empty().fill(2.0f32); - response = c * four - (l + r + u + d) * two + (ul + ur + dl + dr); + response = + centre * four - (left + right + up + down) * two + (up_left + up_right + down_left + down_right); } #[unroll] - for ch in 0..channels { - scratch[(tid * 4 + ch) as usize] = f32::abs(response[ch as usize]); + for channel in 0..channels { + scratch[(thread_id * 4 + channel) as usize] = f32::abs(response[channel as usize]); } + #[unroll] - for ch in channels..4u32 { - scratch[(tid * 4 + ch) as usize] = 0.0f32; + for channel in channels..4u32 { + scratch[(thread_id * 4 + channel) as usize] = 0.0f32; } sync_cube(); - if tid == 0 { + if thread_id == 0 { let cube_index = CUBE_POS_Y * CUBE_COUNT_X + CUBE_POS_X; + #[unroll] - for ch in 0..4u32 { + for channel in 0..4u32 { let mut sum = 0.0f32; for t in 0..threads { - sum += scratch[(t * 4 + ch) as usize]; + sum += scratch[(t * 4 + channel) as usize]; } - partials[(cube_index * 4 + ch) as usize] = sum; + + partials[(cube_index * 4 + channel) as usize] = sum; } } } /// The final stage of the Immerkær noise estimate. /// -/// A single block sums every partial into the per-channel totals for the -/// given ring slot. Each thread adds up a strided share of the partials, -/// then thread zero folds those shares together. +/// A single block of `block` threads sums every partial into the per-channel totals at +/// `results[slot * 4..]`. #[cube(launch_unchecked)] pub fn nlm_noise_reduce( partials: &Array, @@ -102,13 +98,13 @@ pub fn nlm_noise_reduce( #[comptime] block: u32, ) { let mut scratch = SharedMemory::::new(comptime!(block * 4) as usize); - let tid = UNIT_POS_X; + let thread_id = UNIT_POS_X; let mut sum0 = 0.0f32; let mut sum1 = 0.0f32; let mut sum2 = 0.0f32; let mut sum3 = 0.0f32; - let mut i = tid; + let mut i = thread_id; while i < num_partials { sum0 += partials[(i * 4) as usize]; sum1 += partials[(i * 4 + 1) as usize]; @@ -116,72 +112,54 @@ pub fn nlm_noise_reduce( sum3 += partials[(i * 4 + 3) as usize]; i += block; } - scratch[(tid * 4) as usize] = sum0; - scratch[(tid * 4 + 1) as usize] = sum1; - scratch[(tid * 4 + 2) as usize] = sum2; - scratch[(tid * 4 + 3) as usize] = sum3; + + scratch[(thread_id * 4) as usize] = sum0; + scratch[(thread_id * 4 + 1) as usize] = sum1; + scratch[(thread_id * 4 + 2) as usize] = sum2; + scratch[(thread_id * 4 + 3) as usize] = sum3; sync_cube(); - if tid == 0 { + if thread_id == 0 { #[unroll] - for ch in 0..4u32 { + for channel in 0..4u32 { let mut total = 0.0f32; for t in 0..block { - total += scratch[(t * 4 + ch) as usize]; + total += scratch[(t * 4 + channel) as usize]; } - results[(slot * 4 + ch) as usize] = total; + + results[(slot * 4 + channel) as usize] = total; } } } -/// Gathers temporal residual statistics for each spatial block, with one -/// GPU block per `block x block` region. +/// Writes one temporal residual record per `block x block` region of the frame. /// -/// For every pixel it computes the difference between the new slot and -/// the previous one, then reduces those differences into a single -/// record. The lag-1 product of neighbouring channel-0 differences is -/// what reveals grain correlated across nearby pixels. +/// The residual is the new slot minus the previous slot. Each record holds the per-channel `sum_d` +/// and `sum_d2`, then `sum_lag`, the lag-1 product of neighbouring channel-0 residuals, which +/// reveals grain correlated across nearby pixels. Pixels outside the frame contribute nothing, and +/// a lag pair never crosses a block boundary. /// -/// Each record also carries nine fields per 8x8 quarter of the block, in -/// top-left, top-right, bottom-left, bottom-right order. Every quarter -/// field reads channel 0 over the quarter's valid pixels. +/// Each record also carries nine channel-0 fields per 8x8 quarter, in top-left, top-right, +/// bottom-left, bottom-right order. /// /// - `sum_d` and `sum_d2` of the residual. /// - `luma_sum`, `luma_min` and `luma_max` of the new frame. -/// - `flatness`, the mean squared neighbour difference of a 4x4 grid. -/// Each grid cell averages a 2x2 group of pixels, and each pixel -/// averages the new and previous frames. A quarter smaller than 8x8 -/// writes `3.0e38` instead, so a flat gate downstream always rejects -/// it. -/// - `tensor_xx`, `tensor_yy` and `tensor_xy`, the structure tensor -/// sums of the temporal mean's 2x2 pixel gradients. Only windows lying -/// wholly inside the quarter and the frame count, 49 on a full quarter. -/// -/// With `luma_fields` off, the kernel skips every quarter tile, barrier -/// and reduction, and writes 0 to each quarter lane instead. +/// - `flatness`, the mean squared neighbour difference of a 4x4 grid of 2x2 cells over the mean of +/// both frames. A quarter smaller than 8x8 writes `3.0e38`, so a flat gate always rejects it. +/// - `tensor_xx`, `tensor_yy` and `tensor_xy`, the structure tensor sums of the temporal mean's +/// 2x2 gradients, from the 2x2 windows inside both the quarter and the frame, 49 on a full +/// quarter. /// -/// A block that runs past the frame edge uses only its in-frame part. -/// Pixels outside the frame contribute nothing, and a pair only forms -/// when its second pixel is still inside that part, so a pair never -/// crosses a block boundary. +/// With `luma_fields` off, every quarter lane is written as 0 and the quarter work is skipped. +/// `block` must be 16, because the quarter reductions assume 8x8 quarters and 64-slot segments. /// /// # Layout /// -/// Records go into `stats` one per block, at -/// `stats[block_index * (2 * stored_ch + 37) ..]`, laid out as every -/// `sum_d`, then every `sum_d2`, then `sum_lag`. Quarter `q` follows at -/// `2 * stored_ch + 1 + 9 * q`, with its fields at the offsets from -/// [QUARTER_SUM_D](crate::nlmeans::noise::QUARTER_SUM_D) to -/// [QUARTER_TENSOR_XY](crate::nlmeans::noise::QUARTER_TENSOR_XY). That -/// stride never depends on `luma_fields`. -/// -/// `stats` should already be sliced down to the new slot's own region of -/// the larger ring buffer. See `noise::run_temporal_noise_stats`. -/// -/// That is the same convention the motion-compensation kernels use for -/// their per-neighbour slices, and it means this kernel never needs to -/// know about the ring's other slots or the padding between them. +/// Records sit at `stats[block_index * (2 * stored_ch + 37)..]` as every `sum_d`, every `sum_d2`, +/// then `sum_lag`. Quarter `q` follows at `2 * stored_ch + 1 + 9 * q`, with its fields at the +/// `QUARTER_*` offsets in `nlmeans::noise`. The stride never depends on `luma_fields`. `stats` must +/// be bound to the new slot's own region of the ring. #[cube(launch_unchecked)] pub fn nlm_temporal_noise_stats( input: &Array>, @@ -194,12 +172,9 @@ pub fn nlm_temporal_noise_stats( #[comptime] block: u32, #[comptime] luma_fields: bool, ) { - // `record_len` is only ever used for `stats`' own output stride. The - // reduction scratch below is filled with `scratch_len` lanes per thread, - // so it is sized by `scratch_len` instead, or its unwritten tail would - // enter the reduction as uninitialised memory. The luma branch reuses it - // for three `threads`-long tensor runs, which fit as `scratch_len` is at - // least 3. + // The scratch is sized by `scratch_len`, not `record_len`, or its unwritten tail would enter + // the reduction as uninitialised memory. The luma branch reuses it for three `threads`-long + // tensor runs, which fit because `scratch_len` is at least 3. let quarter_lanes = comptime!(TEMPORAL_QUARTERS * TEMPORAL_QUARTER_FIELDS); let record_len = comptime!(2 * stored_ch + TEMPORAL_QUARTER_BASE + quarter_lanes); let scratch_len = comptime!(2 * stored_ch + 1); @@ -210,74 +185,67 @@ pub fn nlm_temporal_noise_stats( let local_x = UNIT_POS_X; let local_y = UNIT_POS_Y; - let tid = local_y * block + local_x; + let thread_id = local_y * block + local_x; let block_origin_x = CUBE_POS_X * block; let block_origin_y = CUBE_POS_Y * block; - let gx = block_origin_x + local_x; - let gy = block_origin_y + local_y; + let global_x = block_origin_x + local_x; + let global_y = block_origin_y + local_y; - let valid = gx < width && gy < height; + let valid = global_x < width && global_y < height; - // The in-block extent, truncated so a pair never reaches past this - // block's own slice of the frame. Ragged right and bottom edges use - // the truncated extent, the same way the block matcher's coarse - // kernel seeds its ragged last block from its position rather than - // from a fixed block size. + // The in-frame extent of this block, so a pair never reaches past its slice of the frame. let block_w = u32::min(block, width - block_origin_x); let block_h = u32::min(block, height - block_origin_y); - // These four stay at their harmless defaults, and are never read, - // when `luma_fields` is off. Declaring them either way costs a few - // registers at most, so only the real resource cost, the luma - // tiles below, is behind the flag. + // These are never read when `luma_fields` is off. Only the luma tiles below cost enough to sit + // behind the flag. let mut new_luma_raw = 0.0f32; let mut mean_luma_raw = 0.0f32; - // Sentinels far above and far below any normalised luma value - // (luma runs between 0 and 1), so an invalid pixel never wins the - // min or max reduction below. The negative one is built from a `mut` - // variable and a later compound assignment rather than a negative - // literal initializer. cubecl folds a negative literal used as a - // `let` initializer into a plain Rust constant. That constant cannot - // unify with the `f32::min`/`f32::max` calls below. + + // Sentinels outside the 0..=1 luma range, so an invalid pixel never wins the min or max + // reduction. cubecl folds a negative literal `let` initialiser into a plain Rust constant that + // cannot unify with `f32::min`/`f32::max`, so the negative one is built by a subtraction. let mut min_seed = 1.0e30f32; let mut max_seed = 1.0e30f32; max_seed -= 2.0e30f32; - let mut d = Vector::::empty(); + let mut residual = Vector::::empty(); if valid { - let c = read_line(input, gx, gy, slot_new, width, height); - let p = read_line(input, gx, gy, slot_prev, width, height); - d = c - p; + let new_pixel = read_line(input, global_x, global_y, slot_new, width, height); + let prev_pixel = read_line(input, global_x, global_y, slot_prev, width, height); + residual = new_pixel - prev_pixel; + if comptime!(luma_fields) { - new_luma_raw = c[0]; - mean_luma_raw = 0.5f32 * (c[0] + p[0]); - min_seed = c[0]; - max_seed = c[0]; + new_luma_raw = new_pixel[0]; + mean_luma_raw = 0.5f32 * (new_pixel[0] + prev_pixel[0]); + min_seed = new_pixel[0]; + max_seed = new_pixel[0]; } } - // The `.into()` calls are what let cubecl unify the two branches, - // because both arms have to expand to the same `NativeExpand`. - // Clippy cannot see that requirement. #[expect( clippy::useless_conversion, reason = "both branches have to expand to the same cubecl native type, which the \ conversion supplies" )] - let d0 = if valid { d[0] } else { 0.0f32.into() }; - d0_tile[tid as usize] = d0; + let d0 = if valid { residual[0] } else { 0.0f32.into() }; + d0_tile[thread_id as usize] = d0; #[unroll] - for ch in 0..stored_ch { + for channel in 0..stored_ch { #[expect( clippy::useless_conversion, reason = "both branches have to expand to the same cubecl native type, which the \ conversion supplies" )] - let v = if valid { d[ch as usize] } else { 0.0f32.into() }; - scratch[(tid * scratch_len + ch) as usize] = v; - scratch[(tid * scratch_len + stored_ch + ch) as usize] = v * v; + let lane_residual = if valid { + residual[channel as usize] + } else { + 0.0f32.into() + }; + scratch[(thread_id * scratch_len + channel) as usize] = lane_residual; + scratch[(thread_id * scratch_len + stored_ch + channel) as usize] = lane_residual * lane_residual; } sync_cube(); @@ -289,15 +257,15 @@ pub fn nlm_temporal_noise_stats( conversion supplies" )] let lag = if pair_valid { - d0 * d0_tile[(tid + 1) as usize] + d0 * d0_tile[(thread_id + 1) as usize] } else { 0.0f32.into() }; - scratch[(tid * scratch_len + 2 * stored_ch) as usize] = lag; + scratch[(thread_id * scratch_len + 2 * stored_ch) as usize] = lag; sync_cube(); - if tid == 0 { + if thread_id == 0 { let block_index = CUBE_POS_Y * CUBE_COUNT_X + CUBE_POS_X; let out_base = block_index * record_len; @@ -307,12 +275,11 @@ pub fn nlm_temporal_noise_stats( for t in 0..threads { total += scratch[(t * scratch_len + lane) as usize]; } + stats[(out_base + lane) as usize] = total; } - // With `luma_fields` off, the quarter lanes get an explicit 0 - // rather than whatever `stats` already held there, so a reader - // never sees stale data left over from an earlier frame's slot. + // The quarter lanes are written explicitly so a reader never sees an earlier frame's data. if comptime!(!luma_fields) { let quarters_start = out_base + 2 * stored_ch + TEMPORAL_QUARTER_BASE; @@ -323,14 +290,10 @@ pub fn nlm_temporal_noise_stats( } } - // Everything from here on computes the quarter lanes. With - // `luma_fields` off, none of it is compiled. `luma_fields` never - // varies within a launch, so every thread takes this branch - // identically and reaches every barrier inside it. + // The branch is comptime, so every thread reaches every barrier inside it. if comptime!(luma_fields) { - // Each quarter's 64 pixels sit in one contiguous 64-slot segment - // of the reduction tiles, so one tree reduction folds all four - // quarters at once. + // Each quarter's 64 pixels sit in one contiguous segment of the reduction tiles, so one tree + // reduction folds all four quarters at once. let quarter = (local_y / 8u32) * 2u32 + local_x / 8u32; let quarter_pos = (local_y % 8u32) * 8u32 + local_x % 8u32; let slot = quarter * 64u32 + quarter_pos; @@ -348,29 +311,26 @@ pub fn nlm_temporal_noise_stats( max_tile[slot as usize] = max_seed; residual_tile[slot as usize] = d0; residual_sq_tile[slot as usize] = d0 * d0; - // The lag pairs above are done with `d0_tile`, so it holds the - // temporal mean instead. A separate tile costs about 40% of this - // variant's throughput, because the extra shared memory lowers - // how many workgroups run at once. - d0_tile[tid as usize] = mean_luma_raw; + + // `d0_tile` is reused for the temporal mean, because a separate tile lowers occupancy enough + // to cost about 40% of this variant's throughput. + d0_tile[thread_id as usize] = mean_luma_raw; sync_cube(); - // Each thread whose 2x2 window lies inside its quarter and the - // frame adds one gradient of the temporal mean to the structure - // tensor. The scalar reduction above is done with `scratch`, so - // it holds the three sums, one `threads`-long run each. + // Each 2x2 window inside its quarter and the frame adds one gradient of the temporal mean to + // the structure tensor. `scratch` holds the three sums, one `threads`-long run each. let quarter_col = local_x % 8u32; let quarter_row = local_y % 8u32; let window_in_quarter = quarter_col < 7u32 && quarter_row < 7u32; - let window_in_frame = gx + 1u32 < width && gy + 1u32 < height; + let window_in_frame = global_x + 1u32 < width && global_y + 1u32 < height; let mut grad_x = 0.0f32; let mut grad_y = 0.0f32; if window_in_quarter && window_in_frame { - let top_left = d0_tile[tid as usize]; - let top_right = d0_tile[(tid + 1u32) as usize]; - let bottom_left = d0_tile[(tid + block) as usize]; - let bottom_right = d0_tile[(tid + block + 1u32) as usize]; + let top_left = d0_tile[thread_id as usize]; + let top_right = d0_tile[(thread_id + 1u32) as usize]; + let bottom_left = d0_tile[(thread_id + block) as usize]; + let bottom_right = d0_tile[(thread_id + block + 1u32) as usize]; grad_x = 0.5f32 * ((top_right + bottom_right) - (top_left + bottom_left)); grad_y = 0.5f32 * ((bottom_left + bottom_right) - (top_left + top_right)); } @@ -381,11 +341,9 @@ pub fn nlm_temporal_noise_stats( sync_cube(); - // Six halving rounds reduce each 64-slot segment to its first - // slot. `stride` is a per-round compile-time constant, because - // cubecl panics at JIT time on a `mut` seeded from a comptime - // value. The halving only guards which threads update a slot, so - // every thread reaches every `sync_cube()`. + // Six halving rounds reduce each 64-slot segment to its first slot. `stride` is a per-round + // comptime constant because cubecl panics at JIT time on a `mut` seeded from a comptime + // value. Every thread reaches every `sync_cube()`, the halving only guards the updates. #[unroll] for step in 0..6u32 { let stride = comptime!(32u32 >> step); @@ -406,15 +364,16 @@ pub fn nlm_temporal_noise_stats( scratch[yy_here] += scratch[yy_partner]; scratch[xy_here] += scratch[xy_partner]; } + sync_cube(); } - // 64 threads each average their own 2x2 group of `d0_tile` into - // one cell of an 8x8 grid, laid out row-major. - if tid < 64u32 { - let cell_y = tid / 8u32; - let cell_x = tid % 8u32; + // 64 threads each average a 2x2 group of `d0_tile` into one cell of a row-major 8x8 grid. + if thread_id < 64u32 { + let cell_y = thread_id / 8u32; + let cell_x = thread_id % 8u32; let mut cell = 0.0f32; + #[unroll] for dy in 0..2u32 { #[unroll] @@ -423,26 +382,26 @@ pub fn nlm_temporal_noise_stats( cell += d0_tile[idx as usize]; } } - cells[tid as usize] = cell / 4.0f32; + + cells[thread_id as usize] = cell / 4.0f32; } + sync_cube(); - // The same 64 threads each score their cell's right and down - // neighbour. A pair that would cross into another quarter is - // skipped, which leaves 24 pairs per quarter. The scores are - // stored quarter-major, 16 cells per quarter. - if tid < 64u32 { - let cell_y = tid / 8u32; - let cell_x = tid % 8u32; - let here = cells[tid as usize]; + // Each cell scores its right and down neighbour within its quarter, giving 24 pairs per + // quarter. The scores are stored quarter-major, 16 cells per quarter. + if thread_id < 64u32 { + let cell_y = thread_id / 8u32; + let cell_x = thread_id % 8u32; + let here = cells[thread_id as usize]; let mut local_energy = 0.0f32; if cell_x % 4u32 != 3u32 { - let right = cells[(tid + 1u32) as usize]; + let right = cells[(thread_id + 1u32) as usize]; local_energy += (here - right) * (here - right); } if cell_y % 4u32 != 3u32 { - let below = cells[(tid + 8u32) as usize]; + let below = cells[(thread_id + 8u32) as usize]; local_energy += (here - below) * (here - below); } @@ -450,28 +409,26 @@ pub fn nlm_temporal_noise_stats( let cell_pos = (cell_y % 4u32) * 4u32 + cell_x % 4u32; pair_tile[(cell_quarter * 16u32 + cell_pos) as usize] = local_energy; } + sync_cube(); - // Four halving rounds reduce each quarter's 16 scores to its - // first slot. + // Four halving rounds reduce each quarter's 16 scores to its first slot. #[unroll] for step in 0..4u32 { let energy_stride = comptime!(8u32 >> step); - if tid < 64u32 && tid % 16u32 < energy_stride { - pair_tile[tid as usize] += pair_tile[(tid + energy_stride) as usize]; + if thread_id < 64u32 && thread_id % 16u32 < energy_stride { + pair_tile[thread_id as usize] += pair_tile[(thread_id + energy_stride) as usize]; } + sync_cube(); } - // The first thread of each quarter writes that quarter's record. if quarter_pos == 0u32 { let quarter_x = quarter % 2u32; let quarter_y = quarter / 2u32; let full_width = block_w >= (quarter_x + 1u32) * 8u32; let full_height = block_h >= (quarter_y + 1u32) * 8u32; - // A quarter smaller than 8x8 keeps this sentinel, far above - // any real gradient energy, so a flat gate always rejects it. let mut flatness = 3.0e38f32; if full_width && full_height { flatness = pair_tile[(quarter * 16u32) as usize] / 24.0f32; @@ -495,10 +452,8 @@ pub fn nlm_temporal_noise_stats( /// Fills a slice of the temporal-stats ring with zeroes. /// -/// A duplicated ring slot holds exactly the same pixels as the one -/// before it, so measuring the difference would only ever produce an -/// all-zero record. Writing the zeroes directly is cheaper and gives the -/// aggregation step the same "nothing to measure here" signal. +/// A duplicated ring slot would only ever measure an all-zero record, so writing the zeroes is +/// cheaper and gives the same result. #[cube(launch_unchecked)] pub fn nlm_temporal_stats_zero( dst: &mut Array, diff --git a/av-denoise-core/src/nlmeans/kernels/separable.rs b/av-denoise-core/src/nlmeans/kernels/separable.rs index ccdc1e1..3db96d4 100644 --- a/av-denoise-core/src/nlmeans/kernels/separable.rs +++ b/av-denoise-core/src/nlmeans/kernels/separable.rs @@ -11,11 +11,10 @@ use super::helpers::{ welsch_weight, }; -/// Computes the per-pixel squared distance for both the forward and the -/// backward neighbour, filling both distance buffers in one pass. +/// Writes the per-pixel squared distances to the forward and backward neighbours in one pass. /// -/// The separable box filter and the fused weight-and-accumulate kernel -/// read those buffers afterwards. +/// The backward distance is stored at the backward neighbour's position, comparing `frame_bwd` at +/// the pixel against `frame_t` at the pixel plus `q`. #[cube(launch_unchecked)] pub fn nlm_distance_pair( input: &Array>, @@ -76,8 +75,7 @@ pub fn nlm_distance_pair( dist_bwd[pixel_idx] = line_sum_sq(bwd_center - bwd_neighbor, channels) * scale; } -/// Computes the per-pixel squared distance between each pixel and its -/// shifted neighbour, scaled for the channel mode in use. +/// Writes the channel-scaled squared distance between each pixel and its neighbour at `q`. #[cube(launch_unchecked)] pub fn nlm_distance( input: &Array>, @@ -119,12 +117,9 @@ pub fn nlm_distance( dist[(y * width + x) as usize] = line_sum_sq(center - neighbor, channels) * scale; } -/// Sums each row of a patch, which is the horizontal half of the -/// separable box filter. +/// The horizontal half of the separable box filter, summing `2 * patch_radius + 1` values per row. /// -/// The block loads a `(block_x + 2 * patch_radius) x block_y` tile into -/// shared memory, then each thread writes the sum of the -/// `2 * patch_radius + 1` values across its own row. +/// The block caches a `(block_x + 2 * patch_radius) x block_y` tile in shared memory. #[cube(launch_unchecked)] pub fn nlm_horizontal_sum( input: &Array, @@ -171,11 +166,11 @@ pub fn nlm_horizontal_sum( for offset_x in 0..patch_size { sum += smem[(smem_base + offset_x) as usize]; } + output[(global_y * width + global_x) as usize] = sum; } -/// Sums the row sums down each column, completing the patch total, then -/// turns it into a Welsch weight. +/// The vertical half of the separable box filter, writing the Welsch weight of each patch total. #[cube(launch_unchecked)] pub fn nlm_vertical_weight( input: &Array, @@ -222,13 +217,11 @@ pub fn nlm_vertical_weight( for offset_y in 0..patch_size { sum += smem[((local_y + offset_y) * block_x + local_x) as usize]; } + output[(global_y * width + global_x) as usize] = welsch_weight(sum, h2_inv_norm, noise_offset); } -/// The paired version of the horizontal box filter. -/// -/// The forward and backward passes share one tile load and one -/// `sync_cube`. +/// The paired version of `nlm_horizontal_sum`, sharing one tile pass and barrier for both buffers. #[cube(launch_unchecked)] pub fn nlm_horizontal_sum_pair( input_fwd: &Array, @@ -288,20 +281,14 @@ pub fn nlm_horizontal_sum_pair( output_bwd[out_idx] = sum_bwd; } -/// Finishes the separable paired path by summing down each column, -/// turning the result into Welsch weights, and accumulating them, all in -/// one kernel. +/// Finishes the paired separable path with the vertical sum, the Welsch weights and accumulation. /// -/// The backward tile is loaded from `hsum_bwd` at the block position -/// shifted the opposite way, so each thread's column sum lands directly -/// on the backward neighbour's weight. +/// The backward tile is read from `hsum_bwd` shifted by `-q`, so each column sum lands on the +/// backward neighbour's weight. /// -/// When `use_confidence` is true, each weight is multiplied by its -/// block's confidence before `accumulate_pair` folds it in, using the -/// same pixel-to-block mapping `nlm_mc_warp` uses. -/// -/// When `use_confidence` is false the lookup and the multiply are -/// dropped at compile time, and the confidence buffers are never read. +/// When `use_confidence` is set, each weight is scaled by its block's confidence from `conf_fwd` +/// and `conf_bwd`. `step`, `blocks_x` and `blocks_y` describe the motion block grid, which pixels +/// map onto the same way as in `nlm_mc_warp`. When it is unset the confidence buffers are never read. #[cube(launch_unchecked)] pub fn nlm_vweight_pair_accumulate( hsum_fwd: &Array, @@ -382,9 +369,9 @@ pub fn nlm_vweight_pair_accumulate( let mut weight_bwd = welsch_weight(sum_bwd, h2_inv_norm, noise_offset); if use_confidence { - let bx = (global_x / step).min(blocks_x - 1); - let by = (global_y / step).min(blocks_y - 1); - let block_idx = (by * blocks_x + bx) as usize; + let block_col = (global_x / step).min(blocks_x - 1); + let block_row = (global_y / step).min(blocks_y - 1); + let block_idx = (block_row * blocks_x + block_col) as usize; weight_fwd *= conf_fwd[block_idx]; weight_bwd *= conf_bwd[block_idx]; } @@ -395,10 +382,7 @@ pub fn nlm_vweight_pair_accumulate( ); } -/// The reference-image version of `nlm_distance`. -/// -/// Both the centre and the neighbour are read from `reference`. The -/// separable kernels downstream read the distance buffer unchanged. +/// The reference-image version of `nlm_distance`, reading both pixels from `reference`. #[cube(launch_unchecked)] pub fn nlm_distance_ref( reference: &Array>, @@ -440,10 +424,7 @@ pub fn nlm_distance_ref( dist[(y * width + x) as usize] = line_sum_sq(center - neighbor, channels) * scale; } -/// The reference-image version of `nlm_distance_pair`. -/// -/// Both the forward and the backward distances are read from -/// `reference`. +/// The reference-image version of `nlm_distance_pair`, reading every pixel from `reference`. #[cube(launch_unchecked)] pub fn nlm_distance_pair_ref( reference: &Array>, diff --git a/av-denoise-core/src/nlmeans/mod.rs b/av-denoise-core/src/nlmeans/mod.rs index 75f9345..06bcc7d 100644 --- a/av-denoise-core/src/nlmeans/mod.rs +++ b/av-denoise-core/src/nlmeans/mod.rs @@ -1,280 +1,68 @@ -//! The non-local means denoiser that sits behind [`crate::Denoiser`]. +//! Non-local means denoising on the GPU //! -//! Non-local means cleans a pixel by finding patches elsewhere that look -//! like the patch around it, then averaging them. Similar patches get a -//! large weight and dissimilar ones get almost none, so flat areas -//! smooth out while edges survive. +//! A pixel is averaged with pixels whose surrounding patches look alike, within one frame or across +//! a temporal window. The module provides: //! -//! The search can reach across neighbouring frames as well as within one -//! frame, which is what the temporal radius controls. -//! -//! # Layout -//! -//! `params` holds the tuning values and the calibrated defaults, and -//! [`NlmParams`] is the single struct everything else is built from. -//! -//! [`NlmDenoiser`] owns the GPU buffers and the frame ring, and -//! `dispatch` turns one set of parameters into the sequence of kernel -//! launches that produces a frame. -//! -//! [`kernels`] holds the GPU code itself. `noise` measures how noisy a -//! frame is, [`motion`] tracks movement between frames, and -//! [`prefilter`] builds the cleaner reference image that patches are -//! compared against. +//! - [Nlmeans], the engine +//! - the options and tuning parameters it is built from +//! - noise estimation, motion compensation and prefilters -pub mod kernels; -pub mod motion; -pub mod prefilter; +pub(crate) mod denoiser; +pub(crate) mod kernels; +pub(crate) mod motion; +pub(crate) mod params; +pub(crate) mod prefilter; mod align; -mod denoiser; mod dispatch; mod edges; +mod engine; mod noise; -mod params; -mod pending; +mod options; -// Every test in this tree runs against a real GPU runtime, see -// `tests::helpers::R`, so it only builds when a wgpu-backed feature is -// enabled. A cpu-only build skips it entirely, and the -// `cpu_smoke_tests` module in `src/denoiser.rs` covers that backend -// instead. +// The tests run against a real GPU runtime, so they need a wgpu-backed feature. #[cfg(all(test, any(feature = "vulkan", feature = "metal")))] -mod tests; - -pub(crate) use denoiser::RingView; -pub use denoiser::{GpuOutput, NlmDenoiser}; -pub use motion::{MotionCompensationMode, MotionEstimation, MotionSearch}; -pub use params::{ - ChannelMode, - HqParams, - MAX_PATCH_RADIUS, - MAX_SEARCH_RADIUS, - MAX_TEMPORAL_RADIUS, - MIN_FRAME_DIM, - NlmParams, - hq_default_strength, - validate_dimensions, -}; -pub(crate) use pending::start_readback; -pub use pending::{Pending, TryWait}; -pub use prefilter::{DEFAULT_PILOT_STRENGTH_SCALE, PrefilterMode, parse_prefilter}; +pub(crate) mod tests; -pub use self::noise::NOISE_CURVE_BINS; +pub(crate) use self::denoiser::{NlmDenoiser, RingView}; +pub use self::engine::Nlmeans; +pub use self::motion::{MotionCompensationMode, MotionEstimation, MotionSearch}; #[cfg(test)] pub(crate) use self::noise::QuarterClass; -pub(crate) use self::noise::{QuarterClasses, StrengthMapParams}; +pub(crate) use self::noise::{NOISE_CURVE_BINS, QuarterClasses, StrengthMapParams}; +#[cfg(all(test, any(feature = "vulkan", feature = "metal")))] +pub(crate) use self::options::resolve_params; +pub use self::options::{ + DenoisingMode, + NlmTuning, + NlmeansAlgorithm, + NlmeansHqOptions, + NlmeansOptions, + NlmeansVariant, + nlmeans_search_radius_for, + nlmeans_temporal_radius_for, + nlmeans_variant_for, +}; +pub use self::params::{ChannelMode, HqParams}; +#[cfg(all(test, any(feature = "vulkan", feature = "metal")))] +pub(crate) use self::params::{MAX_PATCH_RADIUS, MAX_SEARCH_RADIUS, MAX_TEMPORAL_RADIUS}; +pub(crate) use self::params::{NlmParams, hq_default_strength}; +pub use self::prefilter::{DEFAULT_PILOT_STRENGTH_SCALE, PrefilterMode, parse_prefilter}; /// Cube X dimension for tile-heavy fused/separable kernels. -pub const BLOCK_X: u32 = 32; +pub(crate) const BLOCK_X: u32 = 32; /// Cube Y dimension for tile-heavy fused/separable kernels. -pub const BLOCK_Y: u32 = 8; +pub(crate) const BLOCK_Y: u32 = 8; -/// Cube shape for the per-pixel `nlm_accumulate` kernel, which has no -/// shared-memory tile. +/// Cube shape for the per-pixel `nlm_accumulate` kernel, which has no shared-memory tile. /// -/// On RDNA-class GPUs this shape benchmarks 10 to 25% faster than the -/// tile-heavy default. The kernel waits on memory rather than compute, -/// so the extra threads hide the load latency. -pub const BLOCK_X_THIN: u32 = 32; -pub const BLOCK_Y_THIN: u32 = 16; +/// On RDNA-class GPUs it benchmarks 10 to 25% faster than the tile-heavy default, because the +/// kernel waits on memory and the extra threads hide the load latency. +pub(crate) const BLOCK_X_THIN: u32 = 32; +pub(crate) const BLOCK_Y_THIN: u32 = 16; -/// Largest 1D grid a dispatch may ask for, set by the WebGPU and Vulkan -/// limits. +/// Largest 1D grid a dispatch may ask for, set by the WebGPU and Vulkan limits. pub(crate) const MAX_GRID_1D: u32 = 65535; -/// Block size for 1D utility kernels (copy, zero). +/// Block size for the 1D copy and zero kernels. pub(crate) const BLOCK_1D: u32 = 256; - -/// Bit depth of a source's samples. -/// -/// Normalisation divides by [`Depth::max_value`], so a value in -/// normalised units means the same thing at every depth. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum Depth { - Eight, - Ten, - Twelve, -} - -/// Returned when a source declares a bit depth the denoiser does not handle. -#[derive(Debug, thiserror::Error)] -#[error("unsupported bit depth {0}, av-denoise supports 8, 10, and 12-bit")] -pub struct UnsupportedDepthError(pub usize); - -impl Depth { - /// Maps a declared bit depth onto a [`Depth`]. - pub fn from_bits(bits: usize) -> Result { - match bits { - 8 => Ok(Depth::Eight), - 10 => Ok(Depth::Ten), - 12 => Ok(Depth::Twelve), - other => Err(UnsupportedDepthError(other)), - } - } - - /// Bits per sample. - pub fn bits(self) -> usize { - match self { - Depth::Eight => 8, - Depth::Ten => 10, - Depth::Twelve => 12, - } - } - - /// Bytes each sample takes up on the wire. - /// - /// Depths above 8 use a little-endian 16-bit word. - pub fn bytes_per_sample(self) -> usize { - match self { - Depth::Eight => 1, - Depth::Ten | Depth::Twelve => 2, - } - } - - /// The largest sample value this depth can hold, which is also the - /// normalisation divisor. - pub fn max_value(self) -> f32 { - ((1u32 << self.bits()) - 1) as f32 - } - - /// The sample value that means neutral chroma at this depth. - pub fn neutral_chroma(self) -> u16 { - 1 << (self.bits() - 1) - } - - /// How [`crate::nlmeans::kernels::gpu_pack_wire`] packs samples at - /// this depth. - pub fn wire_pack(self) -> WirePack { - WirePack { - max: self.max_value(), - samples_per_word: 4 / self.bytes_per_sample() as u32, - } - } -} - -/// The quantisation scale and lane count `gpu_pack_wire` packs one -/// sample with. -/// -/// The kernel shifts a quantised sample by `lane * (32 / samples_per_word)` -/// and never masks it, so a `max` wider than the lane holds spills bits -/// into the neighbouring sample. Both values come from one [`Depth`] -/// through [`Depth::wire_pack`], which makes that pairing impossible to -/// get wrong. -#[derive(Debug, Clone, Copy, PartialEq)] -pub struct WirePack { - max: f32, - samples_per_word: u32, -} - -impl WirePack { - /// The scale a normalised value is quantised against. - pub fn max(self) -> f32 { - self.max - } - - /// How many samples share one `u32` word. - pub fn samples_per_word(self) -> u32 { - self.samples_per_word - } -} - -/// Scales native-depth samples into normalised `[0, 1]` f32. -pub fn normalize(input: &[u16], depth: Depth) -> Vec { - let max = depth.max_value(); - input.iter().map(|&v| v as f32 / max).collect() -} - -/// Reverse of [`normalize`]. -/// -/// Values outside `[0, 1]` are clamped, and `NaN` becomes 0. -pub fn denormalize(input: &[f32], depth: Depth) -> Vec { - let max = depth.max_value(); - input - .iter() - .map(|&v| (v * max).round().clamp(0.0, max) as u16) - .collect() -} - -#[cfg(test)] -mod depth_tests { - use super::*; - - #[test] - fn from_bits_accepts_supported_depths() { - assert_eq!(Depth::from_bits(8).unwrap(), Depth::Eight); - assert_eq!(Depth::from_bits(10).unwrap(), Depth::Ten); - assert_eq!(Depth::from_bits(12).unwrap(), Depth::Twelve); - } - - #[test] - fn from_bits_rejects_unsupported_depths() { - for bits in [0, 9, 14, 16] { - let err = Depth::from_bits(bits).expect_err("expected rejection"); - assert!( - err.to_string().contains(&bits.to_string()), - "error should name the depth, got {err}" - ); - } - } - - #[test] - fn depth_properties_match_the_format() { - assert_eq!(Depth::Eight.bytes_per_sample(), 1); - assert_eq!(Depth::Ten.bytes_per_sample(), 2); - assert_eq!(Depth::Twelve.bytes_per_sample(), 2); - - assert_eq!(Depth::Eight.max_value(), 255.0); - assert_eq!(Depth::Ten.max_value(), 1023.0); - assert_eq!(Depth::Twelve.max_value(), 4095.0); - - assert_eq!(Depth::Eight.neutral_chroma(), 128); - assert_eq!(Depth::Ten.neutral_chroma(), 512); - assert_eq!(Depth::Twelve.neutral_chroma(), 2048); - } - - /// Limited-range black and white land on matching normalised values - /// at every depth, which is what lets every calibrated constant in - /// the library stay depth-independent. - /// - /// The match is within one 8-bit code level rather than exact. ITU - /// defines the limited-range endpoints as exact multiples, so 235 - /// becomes 940 and then 3760, but full scale is not a multiple, - /// because 255 becomes 1023 and then 4095. - /// - /// That leaves 235/255 and 940/1023 differing by 0.0027, roughly - /// 0.69 of an 8-bit step. Agreement below one step is the real - /// property here. - #[test] - fn normalized_scale_is_identical_across_depths() { - /// One 8-bit code level, the precision the endpoints agree to. - const TOL: f32 = 1.0 / 255.0; - - let eight = normalize(&[16, 235], Depth::Eight); - let ten = normalize(&[64, 940], Depth::Ten); - let twelve = normalize(&[256, 3760], Depth::Twelve); - - for (a, b) in eight.iter().zip(ten.iter()) { - assert!((a - b).abs() < TOL, "8-bit {a} vs 10-bit {b}"); - } - for (a, b) in eight.iter().zip(twelve.iter()) { - assert!((a - b).abs() < TOL, "8-bit {a} vs 12-bit {b}"); - } - } - - #[test] - fn normalization_round_trips_at_every_depth() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let max = depth.max_value() as u16; - let original: Vec = vec![0, 1, 16, 64, 128, 235, max / 2, max - 1, max]; - let restored = denormalize(&normalize(&original, depth), depth); - assert_eq!(original, restored, "round trip failed at {depth:?}"); - } - } - - #[test] - fn denormalize_clamps_out_of_range_input() { - let out = denormalize(&[-0.5, 0.0, 1.0, 1.5], Depth::Ten); - assert_eq!(out, vec![0, 0, 1023, 1023]); - } -} diff --git a/av-denoise-core/src/nlmeans/motion/analyse.rs b/av-denoise-core/src/nlmeans/motion/analyse.rs index 2e1b4c1..1a36d66 100644 --- a/av-denoise-core/src/nlmeans/motion/analyse.rs +++ b/av-denoise-core/src/nlmeans/motion/analyse.rs @@ -7,52 +7,35 @@ use super::pyramid::{level_dims, pyramid_slot_byte_offset}; use crate::nlmeans::align::StorageAlign; use crate::nlmeans::kernels::motion::{nlm_mc_block_match_coarse, nlm_mc_block_match_fine}; -/// Where a neighbour's slice of the motion field starts. +/// Where a neighbour's motion-field slice starts. /// -/// The buffer is indexed by neighbour, then block, then component, with -/// two `i32` components per block. Each neighbour's slice is padded up -/// to an alignment boundary. See -/// [`MotionCtx::mv_field_bytes_per_neighbour`]. -pub(crate) fn mv_field_byte_offset(mc: &MotionCtx, neighbour_idx: u32) -> u64 { - (neighbour_idx as u64) * mc.mv_field_bytes_per_neighbour() +/// The field is indexed by neighbour, block and then component, with two `i32` per block. Each +/// slice is padded as [MotionCtx::mv_field_bytes_per_neighbour] describes. +pub(crate) fn mv_field_byte_offset(motion_ctx: &MotionCtx, neighbour_idx: u32) -> u64 { + (neighbour_idx as u64) * motion_ctx.mv_field_bytes_per_neighbour() } -/// Where a neighbour's slice of the confidence buffer starts. +/// Where a neighbour's confidence slice starts. /// -/// The layout mirrors the motion field, indexed by neighbour and then -/// block, but stores one `f32` per block rather than two `i32` -/// components. -/// -/// Each neighbour's slice is padded up to an alignment boundary. See -/// [`MotionCtx::confidence_bytes_per_neighbour`]. -pub(crate) fn confidence_byte_offset(mc: &MotionCtx, neighbour_idx: u32) -> u64 { - (neighbour_idx as u64) * mc.confidence_bytes_per_neighbour() +/// It mirrors the motion field's layout with one `f32` per block, padded as +/// [MotionCtx::confidence_bytes_per_neighbour] describes. +pub(crate) fn confidence_byte_offset(motion_ctx: &MotionCtx, neighbour_idx: u32) -> u64 { + (neighbour_idx as u64) * motion_ctx.confidence_bytes_per_neighbour() } -/// Works out how one neighbour frame moved relative to the centre -/// frame. -/// -/// A coarse pass runs on the smallest pyramid level, then a fine pass -/// refines its answer at full resolution. The result goes into -/// `mv_field` at the slot reserved for this neighbour. -/// -/// `sad_noise_floor` and `thsad` are the fine kernel's confidence -/// scalars. See -/// [`crate::nlmeans::kernels::motion::nlm_mc_block_match_fine`]. -/// -/// When `write_confidence` is true, a per-block confidence score also -/// goes into `confidence` at the matching slot. +/// Estimates how one neighbour frame moved relative to the centre frame. /// -/// When it is false, `confidence` is never indexed. Callers that do not -/// need the score can pass a small placeholder buffer and leave both -/// scalars at 0.0. +/// A coarse pass on the smallest pyramid level seeds a fine pass at full resolution, which writes +/// this neighbour's `mv_field` slot. With `write_confidence` set, a per-block score also lands in +/// its `confidence` slot. Otherwise `confidence` is never indexed, so a placeholder buffer works +/// and `sad_noise_floor` and `thsad` may stay at 0.0. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] pub(crate) fn run_analyse( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, width: u32, height: u32, frame_count: u32, @@ -66,56 +49,52 @@ pub(crate) fn run_analyse( sad_noise_floor: f32, thsad: f32, ) -> Result<(), anyhow::Error> { - let mv_offset = mv_field_byte_offset(mc, neighbour_idx); + let mv_offset = mv_field_byte_offset(motion_ctx, neighbour_idx); let mv_slot = mv_field.clone().offset_start(mv_offset); - let mv_slot_len = (mc.blocks_x as usize) * (mc.blocks_y as usize) * 2; + let mv_slot_len = (motion_ctx.blocks_x as usize) * (motion_ctx.blocks_y as usize) * 2; - // Only slice into `confidence` at its real per-neighbour offset - // when the kernel is going to write it. Otherwise `confidence` is a - // small placeholder buffer with no per-neighbour layout to offset - // into. + // A placeholder `confidence` has no per-neighbour layout to offset into. let (conf_slot, conf_slot_len) = if write_confidence { - let conf_offset = confidence_byte_offset(mc, neighbour_idx); + let conf_offset = confidence_byte_offset(motion_ctx, neighbour_idx); ( confidence.clone().offset_start(conf_offset), - (mc.blocks_x as usize) * (mc.blocks_y as usize), + (motion_ctx.blocks_x as usize) * (motion_ctx.blocks_y as usize), ) } else { (confidence.clone(), 1) }; - // The coarse pass, which only runs with more than one pyramid level. - if mc.pyramid_levels > 1 { - let coarse_level = mc.pyramid_levels - 1; - let (cw, ch) = level_dims(width, height, coarse_level); - let coarse_centre = pyramid.clone().offset_start(pyramid_slot_byte_offset( + if motion_ctx.pyramid_levels > 1 { + let coarse_level = motion_ctx.pyramid_levels - 1; + let (coarse_width, coarse_height) = level_dims(width, height, coarse_level); + let coarse_centre_offset = pyramid_slot_byte_offset( width, height, frame_count, coarse_level, centre_slot, - mc.align, - )); - let coarse_neighbour = pyramid.clone().offset_start(pyramid_slot_byte_offset( + motion_ctx.align, + ); + let coarse_centre = pyramid.clone().offset_start(coarse_centre_offset); + let coarse_neighbour_offset = pyramid_slot_byte_offset( width, height, frame_count, coarse_level, neighbour_slot, - mc.align, - )); - let level_len = (cw * ch) as usize; + motion_ctx.align, + ); + let coarse_neighbour = pyramid.clone().offset_start(coarse_neighbour_offset); + let level_len = (coarse_width * coarse_height) as usize; let coarse_scale = 1u32 << coarse_level; - // A coarse block covers the same content as a fine block scaled - // down by 2 raised to the coarse level. - let coarse_blksize = (mc.blksize / coarse_scale).max(2); - let coarse_step = (mc.step / coarse_scale).max(1); - let coarse_blocks_x = cw.div_ceil(coarse_step).max(1); - let coarse_blocks_y = ch.div_ceil(coarse_step).max(1); + // A coarse block covers the same content as a fine block scaled down by `coarse_scale`. + let coarse_blksize = (motion_ctx.blksize / coarse_scale).max(2); + let coarse_step = (motion_ctx.step / coarse_scale).max(1); + let coarse_blocks_x = coarse_width.div_ceil(coarse_step).max(1); + let coarse_blocks_y = coarse_height.div_ceil(coarse_step).max(1); let grid = CubeCount::new_2d(coarse_blocks_x, coarse_blocks_y); - // One block of threads per image block, sized to suit the 8x8 - // blocks a coarse level typically has. Those threads share the - // scoring work between them. + // One cube per image block, sized for the 8x8 blocks a coarse level typically has, with + // its threads sharing the scoring work. let dim = CubeDim::new_2d(8, 8); unsafe { @@ -126,46 +105,32 @@ pub(crate) fn run_analyse( ArrayArg::from_raw_parts(coarse_centre, level_len), ArrayArg::from_raw_parts(coarse_neighbour, level_len), ArrayArg::from_raw_parts(mv_slot.clone(), mv_slot_len), - cw, - ch, + coarse_width, + coarse_height, coarse_blksize, coarse_step, - mc.search_radius, + motion_ctx.search_radius, coarse_scale, - mc.blocks_x, - mc.blocks_y, - mc.step, + motion_ctx.blocks_x, + motion_ctx.blocks_y, + motion_ctx.step, ); } } else { - // With the pyramid disabled the fine pass has to start from a - // zero seed. There is no dedicated zeroing kernel for `i32` - // here, because the fine pass treats a missing seed as zero - // when there is only one pyramid level. + // A single level has no coarse seed, and the fine pass then treats the seed as zero. } - // The fine pass, which runs at full resolution. - let (fw, fh) = level_dims(width, height, 0); - let fine_centre = pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - 0, - centre_slot, - mc.align, - )); - let fine_neighbour = pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - 0, - neighbour_slot, - mc.align, - )); - let level_len = (fw * fh) as usize; - let grid = CubeCount::new_2d(mc.blocks_x, mc.blocks_y); + let (fine_width, fine_height) = level_dims(width, height, 0); + let fine_centre_offset = + pyramid_slot_byte_offset(width, height, frame_count, 0, centre_slot, motion_ctx.align); + let fine_centre = pyramid.clone().offset_start(fine_centre_offset); + let fine_neighbour_offset = + pyramid_slot_byte_offset(width, height, frame_count, 0, neighbour_slot, motion_ctx.align); + let fine_neighbour = pyramid.clone().offset_start(fine_neighbour_offset); + let level_len = (fine_width * fine_height) as usize; + let grid = CubeCount::new_2d(motion_ctx.blocks_x, motion_ctx.blocks_y); let dim = CubeDim::new_2d(8, 8); - let seeded = if mc.pyramid_levels > 1 { 1u32 } else { 0u32 }; + let seeded = if motion_ctx.pyramid_levels > 1 { 1u32 } else { 0u32 }; unsafe { nlm_mc_block_match_fine::launch_unchecked::( @@ -179,39 +144,31 @@ pub(crate) fn run_analyse( write_confidence, sad_noise_floor, thsad, - fw, - fh, - mc.blksize, - mc.step, - mc.search_radius, + fine_width, + fine_height, + motion_ctx.blksize, + motion_ctx.step, + motion_ctx.search_radius, seeded, - mc.blocks_x, + motion_ctx.blocks_x, ); } Ok(()) } -/// Cleans up the seed that chained motion estimation produced. -/// -/// The joined seed already sits in `mv_field` at this neighbour's slot. -/// This searches a small window around it and writes the corrected -/// vector back to the same place. +/// Refines the chained seed already sitting in this neighbour's `mv_field` slot. /// -/// Unlike [`run_analyse`] there is no coarse pass, because the joined -/// seed already carries the large movement. -/// -/// `refine_radius` is this pass's own search radius, set independently -/// of the direct path's `mc.search_radius`. Every other argument matches -/// the fine-pass call in `run_analyse`, including how confidence is -/// written. +/// It searches `refine_radius` around the seed and writes the corrected vector back. There is no +/// coarse pass, because the joined seed already carries the large movement. Confidence is written +/// the same way as in [run_analyse]. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] pub(crate) fn run_seeded_refine( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, width: u32, height: u32, frame_count: u32, @@ -226,39 +183,29 @@ pub(crate) fn run_seeded_refine( sad_noise_floor: f32, thsad: f32, ) -> Result<(), anyhow::Error> { - let mv_offset = mv_field_byte_offset(mc, neighbour_idx); + let mv_offset = mv_field_byte_offset(motion_ctx, neighbour_idx); let mv_slot = mv_field.clone().offset_start(mv_offset); - let mv_slot_len = (mc.blocks_x as usize) * (mc.blocks_y as usize) * 2; + let mv_slot_len = (motion_ctx.blocks_x as usize) * (motion_ctx.blocks_y as usize) * 2; let (conf_slot, conf_slot_len) = if write_confidence { - let conf_offset = confidence_byte_offset(mc, neighbour_idx); + let conf_offset = confidence_byte_offset(motion_ctx, neighbour_idx); ( confidence.clone().offset_start(conf_offset), - (mc.blocks_x as usize) * (mc.blocks_y as usize), + (motion_ctx.blocks_x as usize) * (motion_ctx.blocks_y as usize), ) } else { (confidence.clone(), 1) }; - let (fw, fh) = level_dims(width, height, 0); - let fine_centre = pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - 0, - centre_slot, - mc.align, - )); - let fine_neighbour = pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - 0, - neighbour_slot, - mc.align, - )); - let level_len = (fw * fh) as usize; - let grid = CubeCount::new_2d(mc.blocks_x, mc.blocks_y); + let (fine_width, fine_height) = level_dims(width, height, 0); + let fine_centre_offset = + pyramid_slot_byte_offset(width, height, frame_count, 0, centre_slot, motion_ctx.align); + let fine_centre = pyramid.clone().offset_start(fine_centre_offset); + let fine_neighbour_offset = + pyramid_slot_byte_offset(width, height, frame_count, 0, neighbour_slot, motion_ctx.align); + let fine_neighbour = pyramid.clone().offset_start(fine_neighbour_offset); + let level_len = (fine_width * fine_height) as usize; + let grid = CubeCount::new_2d(motion_ctx.blocks_x, motion_ctx.blocks_y); let dim = CubeDim::new_2d(8, 8); unsafe { @@ -273,13 +220,13 @@ pub(crate) fn run_seeded_refine( write_confidence, sad_noise_floor, thsad, - fw, - fh, - mc.blksize, - mc.step, + fine_width, + fine_height, + motion_ctx.blksize, + motion_ctx.step, refine_radius, 1u32, - mc.blocks_x, + motion_ctx.blocks_x, ); } @@ -296,131 +243,127 @@ mod tests { StorageAlign::new(32) } - fn mc(blksize: u32, overlap: u32) -> MotionCtx { - MotionCtx::new( - MotionCompensationMode::Mvtools { - blksize, - overlap, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }, - 64, - 64, - align(), - ) - .unwrap() + fn motion_ctx(blksize: u32, overlap: u32) -> MotionCtx { + let mode = MotionCompensationMode::Mvtools { + blksize, + overlap, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }; + let align = align(); + + MotionCtx::new(mode, 64, 64, align).unwrap() + } + + /// A 4x4 frame at a geometry that leaves exactly one block. + fn single_block_ctx() -> MotionCtx { + let mode = MotionCompensationMode::Mvtools { + blksize: 4, + overlap: 0, + search_radius: 1, + pyramid_levels: 1, + estimation: MotionEstimation::Direct, + }; + let align = align(); + + MotionCtx::new(mode, 4, 4, align).unwrap() } #[test] fn mv_field_offset_zero_for_first_neighbour() { - assert_eq!(mv_field_byte_offset(&mc(16, 8), 0), 0); + let ctx = motion_ctx(16, 8); + let offset = mv_field_byte_offset(&ctx, 0); + + assert_eq!(offset, 0); } #[test] fn mv_field_offset_advances_by_blocks() { - let m = mc(16, 8); - let per = (m.blocks_x as u64) * (m.blocks_y as u64) * 2 * 4; - assert_eq!(mv_field_byte_offset(&m, 3), 3 * per); + let ctx = motion_ctx(16, 8); + let stride = (ctx.blocks_x as u64) * (ctx.blocks_y as u64) * 2 * 4; + let offset = mv_field_byte_offset(&ctx, 3); + + assert_eq!(offset, 3 * stride); } #[test] fn confidence_offset_zero_for_first_neighbour() { - assert_eq!(confidence_byte_offset(&mc(16, 8), 0), 0); + let ctx = motion_ctx(16, 8); + let offset = confidence_byte_offset(&ctx, 0); + + assert_eq!(offset, 0); } #[test] fn confidence_offset_advances_by_blocks() { - let m = mc(16, 8); - let per = (m.blocks_x as u64) * (m.blocks_y as u64) * 4; - assert_eq!(confidence_byte_offset(&m, 3), 3 * per); + let ctx = motion_ctx(16, 8); + let stride = (ctx.blocks_x as u64) * (ctx.blocks_y as u64) * 4; + let offset = confidence_byte_offset(&ctx, 3); + + assert_eq!(offset, 3 * stride); } #[test] fn confidence_offset_is_one_component_not_two() { - // Confidence stores one `f32` per block and the motion field - // stores two `i32` components. Both are 4 bytes, so at the same - // block count the confidence stride should be exactly half the - // motion field's, as long as the unpadded stride already lands - // on a 32-byte boundary. This fixture's 64 blocks do. - let m = mc(16, 8); - assert_eq!(mv_field_byte_offset(&m, 1), 2 * confidence_byte_offset(&m, 1)); + // Both strides are 4 bytes per component, so confidence is exactly half the motion field + // while the unpadded stride is 32-byte aligned, which this fixture's 64 blocks are. + let ctx = motion_ctx(16, 8); + let mv_offset = mv_field_byte_offset(&ctx, 1); + let confidence_offset = confidence_byte_offset(&ctx, 1); + + assert_eq!(mv_offset, 2 * confidence_offset); } #[test] fn confidence_offset_pads_small_block_counts_to_32_bytes() { - // A 4x4 frame at this geometry has a single block, so the - // unpadded stride is only 4 bytes and would leave neighbour 1 - // at an offset that is not 32-aligned. - // - // wgpu rejects a bind-group offset that is not a multiple of - // its `min_storage_buffer_offset_alignment`, so the stride has - // to pad up to 32 bytes whatever the block count. - let m = MotionCtx::new( - MotionCompensationMode::Mvtools { - blksize: 4, - overlap: 0, - search_radius: 1, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - }, - 4, - 4, - align(), - ) - .unwrap(); + // One block gives an unpadded stride of 4 bytes, which would leave neighbour 1 off a + // 32-byte boundary. + let ctx = single_block_ctx(); assert_eq!( - m.blocks_x * m.blocks_y, + ctx.blocks_x * ctx.blocks_y, 1, "fixture should have exactly one block" ); - assert_eq!(confidence_byte_offset(&m, 0), 0); - assert_eq!(confidence_byte_offset(&m, 1), 32); - assert_eq!(confidence_byte_offset(&m, 2), 64); + + let first = confidence_byte_offset(&ctx, 0); + let second = confidence_byte_offset(&ctx, 1); + let third = confidence_byte_offset(&ctx, 2); + + assert_eq!(first, 0); + assert_eq!(second, 32); + assert_eq!(third, 64); } #[test] fn mv_field_offset_pads_small_block_counts_to_32_bytes() { - // The same fixture as - // `confidence_offset_pads_small_block_counts_to_32_bytes`, with - // one block. The unpadded motion-field stride is 8 bytes, which - // would leave neighbour 1 at an offset that is not 32-aligned. - let m = MotionCtx::new( - MotionCompensationMode::Mvtools { - blksize: 4, - overlap: 0, - search_radius: 1, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - }, - 4, - 4, - align(), - ) - .unwrap(); + // One block gives an unpadded stride of 8 bytes, which would leave neighbour 1 off a + // 32-byte boundary. + let ctx = single_block_ctx(); assert_eq!( - m.blocks_x * m.blocks_y, + ctx.blocks_x * ctx.blocks_y, 1, "fixture should have exactly one block" ); - assert_eq!(mv_field_byte_offset(&m, 0), 0); - assert_eq!(mv_field_byte_offset(&m, 1), 32); - assert_eq!(mv_field_byte_offset(&m, 2), 64); + + let first = mv_field_byte_offset(&ctx, 0); + let second = mv_field_byte_offset(&ctx, 1); + let third = mv_field_byte_offset(&ctx, 2); + + assert_eq!(first, 0); + assert_eq!(second, 32); + assert_eq!(third, 64); } #[test] fn mv_field_offset_pads_the_1080_square_odd_block_count_case() { - // A 1080x1080 frame at the library defaults gives 135x135 - // blocks, an odd count of 18,225. The unpadded stride of - // 145,800 bytes sits 8 past the preceding 32-byte boundary at - // 145,792, so it has to round up to 145,824 rather than leave - // neighbour 1 misaligned. - // - // The harness's usual 1920x1080 happens to land on an even - // block count at this geometry, so it never reaches this case. - let m = MotionCtx::new(MotionCompensationMode::mvtools_default(), 1080, 1080, align()).unwrap(); + // 1080x1080 at the defaults gives 135x135 blocks. The unpadded stride of 145,800 bytes sits + // 8 past a 32-byte boundary, so it must round up to 145,824. + let mode = MotionCompensationMode::mvtools_default(); + let align = align(); + let ctx = MotionCtx::new(mode, 1080, 1080, align).unwrap(); assert_eq!( - m.blocks_x * m.blocks_y, + ctx.blocks_x * ctx.blocks_y, 18225, "test premise: this geometry gives an odd block count" ); @@ -429,7 +372,11 @@ mod tests { 8, "test premise: the unpadded stride is not 32-aligned" ); - assert_eq!(mv_field_byte_offset(&m, 0), 0); - assert_eq!(mv_field_byte_offset(&m, 1), 145_824); + + let first = mv_field_byte_offset(&ctx, 0); + let second = mv_field_byte_offset(&ctx, 1); + + assert_eq!(first, 0); + assert_eq!(second, 145_824); } } diff --git a/av-denoise-core/src/nlmeans/motion/chain.rs b/av-denoise-core/src/nlmeans/motion/chain.rs index 597e149..7abadf1 100644 --- a/av-denoise-core/src/nlmeans/motion/chain.rs +++ b/av-denoise-core/src/nlmeans/motion/chain.rs @@ -9,42 +9,27 @@ use crate::nlmeans::{BLOCK_1D, MAX_GRID_1D}; /// Where one slot and direction of the pair ring starts. /// -/// The ring is indexed by slot, then direction, then block, then -/// component. Direction 0 runs from the older frame to the newer one, -/// and direction 1 the other way. Each direction's slice is padded up to -/// an alignment boundary. See [`MotionCtx::pair_direction_bytes`]. -/// -/// The slot is keyed by the newer frame's place in the push sequence, -/// reduced modulo the ring size. [`super::pair_ring_slot_count`] -/// explains why that sizing is safe. -/// -/// The `nlm_mc_chain_compose` kernel reads the whole ring as one array -/// and steps through it with the same padded strides, so any change here -/// has to stay in step with that kernel's own stride arguments. -pub(crate) fn pair_byte_offset(mc: &MotionCtx, pair_slot: u32, direction: u32) -> u64 { - (pair_slot as u64) * mc.pair_slot_bytes() + (direction as u64) * mc.pair_direction_bytes() +/// The ring is indexed by slot, direction, block and then component, and direction 0 runs from the +/// older frame to the newer one. Each direction is padded as [MotionCtx::pair_direction_bytes] +/// describes. A slot is keyed by the newer frame's push index modulo the ring size, which +/// [pair_ring_slot_count](crate::nlmeans::motion::pair_ring_slot_count) makes safe. +/// `nlm_mc_chain_compose` steps through the ring with the same padded strides, so the two must +/// change together. +pub(crate) fn pair_byte_offset(motion_ctx: &MotionCtx, pair_slot: u32, direction: u32) -> u64 { + (pair_slot as u64) * motion_ctx.pair_slot_bytes() + (direction as u64) * motion_ctx.pair_direction_bytes() } -/// Measures motion between one pushed frame and the one before it, in -/// both directions, and stores the result in the pair ring. -/// -/// `older_slot` and `newer_slot` are physical input-ring slots, the same -/// addressing `run_analyse` uses. +/// Measures motion between a pushed frame and the one before it, both ways, into the pair ring. /// -/// `pyramid` is the same one `run_motion_compensation` reads, which is -/// the reference ring when a prefilter is active and the raw input ring -/// otherwise. -/// -/// Nothing reads confidence at the pair level, so both directions turn -/// the fine kernel's confidence output off and pass `confidence_dummy` -/// as the placeholder target. +/// `older_slot` and `newer_slot` are physical input-ring slots. Nothing reads confidence at the pair +/// level, so both directions turn it off and target `confidence_dummy`. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] pub(crate) fn run_pair_analyse( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, width: u32, height: u32, frame_count: u32, @@ -55,10 +40,11 @@ pub(crate) fn run_pair_analyse( pair_ring: &Handle, confidence_dummy: &Handle, ) -> Result<(), anyhow::Error> { - let older_to_newer = pair_ring.clone().offset_start(pair_byte_offset(mc, pair_slot, 0)); + let older_to_newer_offset = pair_byte_offset(motion_ctx, pair_slot, 0); + let older_to_newer = pair_ring.clone().offset_start(older_to_newer_offset); run_analyse::( client, - mc, + motion_ctx, width, height, frame_count, @@ -73,10 +59,11 @@ pub(crate) fn run_pair_analyse( 1.0, )?; - let newer_to_older = pair_ring.clone().offset_start(pair_byte_offset(mc, pair_slot, 1)); + let newer_to_older_offset = pair_byte_offset(motion_ctx, pair_slot, 1); + let newer_to_older = pair_ring.clone().offset_start(newer_to_older_offset); run_analyse::( client, - mc, + motion_ctx, width, height, frame_count, @@ -96,33 +83,29 @@ pub(crate) fn run_pair_analyse( /// Fills both directions of one pair-ring slot with zeroes. /// -/// Duplicated ring slots, which appear while priming the stream and -/// again during the end-of-stream flush, hold the same content in both -/// frames, so their motion is zero by definition. -/// -/// Each direction gets its own dispatch at its own padded offset, -/// because the two directions no longer sit next to each other once -/// `pair_direction_bytes` pads between them. +/// Duplicated ring slots, which appear while priming and during the end-of-stream flush, hold the +/// same frame twice, so their motion is zero. Each direction gets its own dispatch because padding +/// separates the two. pub(crate) fn zero_pair_slot( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, pair_ring: &Handle, pair_slot: u32, ) { - let length = mc.pair_direction_len(); + let length = motion_ctx.pair_direction_len(); let grid = length.div_ceil(BLOCK_1D).min(MAX_GRID_1D); let total_threads = grid * BLOCK_1D; for direction in 0..2u32 { - let offset = pair_byte_offset(mc, pair_slot, direction); - let dst = pair_ring.clone().offset_start(offset); + let offset = pair_byte_offset(motion_ctx, pair_slot, direction); + let direction_slice = pair_ring.clone().offset_start(offset); unsafe { nlm_mc_pair_zero::launch_unchecked::( client, CubeCount::new_1d(grid), CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(dst, length as usize), + ArrayArg::from_raw_parts(direction_slice, length as usize), length, total_threads, ); @@ -130,21 +113,17 @@ pub(crate) fn zero_pair_slot( } } -/// Launches `nlm_mc_chain_compose` once, writing the joined motion field -/// into `mv_field` at this neighbour's slot. -/// -/// That is the same slot the direct path fills for the neighbour. +/// Launches `nlm_mc_chain_compose` into this neighbour's `mv_field` slot. /// -/// The padded per-direction and per-slot strides are passed through -/// explicitly, because the kernel reads the whole pair ring as one array -/// and has to step through it the same way `pair_byte_offset` does. +/// The padded strides are passed explicitly, because the kernel reads the whole pair ring as one +/// array and must step through it as [pair_byte_offset] does. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] fn dispatch_chain_compose( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, width: u32, height: u32, pair_ring_slots: u32, @@ -156,15 +135,12 @@ fn dispatch_chain_compose( mv_field: &Handle, neighbour_idx: u32, ) -> Result<(), anyhow::Error> { - let mv_slot = mv_field - .clone() - .offset_start(mv_field_byte_offset(mc, neighbour_idx)); - let mv_slot_len = (mc.blocks_x as usize) * (mc.blocks_y as usize) * 2; + let mv_offset = mv_field_byte_offset(motion_ctx, neighbour_idx); + let mv_slot = mv_field.clone().offset_start(mv_offset); + let mv_slot_len = (motion_ctx.blocks_x as usize) * (motion_ctx.blocks_y as usize) * 2; - // One thread per output block, so one block of threads per image - // block with a single thread in it. Unlike the scoring kernels - // there is no per-candidate work to spread across threads here. - let grid = CubeCount::new_2d(mc.blocks_x, mc.blocks_y); + // One single-thread cube per block, since there is no per-candidate work to share. + let grid = CubeCount::new_2d(motion_ctx.blocks_x, motion_ctx.blocks_y); let dim = CubeDim::new_2d(1, 1); unsafe { @@ -178,33 +154,30 @@ fn dispatch_chain_compose( forward, steps, pair_ring_slots, - mc.pair_direction_stride(), - mc.pair_slot_stride(), - mc.step, + motion_ctx.pair_direction_stride(), + motion_ctx.pair_slot_stride(), + motion_ctx.step, width, height, - mc.blocks_x, - mc.blocks_y, + motion_ctx.blocks_x, + motion_ctx.blocks_y, ); } Ok(()) } -/// Maps a nonzero temporal offset onto the neighbour index the analyse -/// and compose passes use inside the motion-field buffer. +/// Maps a nonzero temporal offset onto its motion-field neighbour index. /// -/// This mirrors `dispatch::neighbour_idx_for_k` exactly. It is -/// duplicated rather than shared because that helper is private to the -/// direct-path dispatch module, and `dispatch`'s own tests already cover -/// the mapping. -/// -/// It is `pub(crate)` so tests can look up the same neighbour index the -/// compose path writes to, rather than working out the formula -/// themselves. +/// Negative offsets take indices 0 up to the radius minus 1 and positive offsets follow, which is +/// the order the analyse, confidence and compose passes fill their buffers in. pub(crate) fn neighbour_idx_for_k(radius: u32, k: i32) -> u32 { - debug_assert_ne!(k, 0); - debug_assert!(k.unsigned_abs() <= radius); + debug_assert_ne!(k, 0, "k=0 is the spatial pair, it has no neighbour index"); + debug_assert!( + k.unsigned_abs() <= radius, + "k={k} outside the temporal window ±{radius}" + ); + if k < 0 { (k + radius as i32) as u32 } else { @@ -213,32 +186,23 @@ pub(crate) fn neighbour_idx_for_k(radius: u32, k: i32) -> u32 { } impl NlmDenoiser { - /// Joins the adjacent-frame fields into one motion field for the - /// neighbour at temporal offset `k`, writing it to field index - /// `neighbour_idx`. - /// - /// `k` must be nonzero and inside the ring, at most twice the - /// temporal radius either way. + /// Joins the adjacent-frame fields into one motion field for the neighbour at offset `k`. /// - /// The walk takes one hop per step out from the centre, following - /// the older-to-newer field for a positive `k` and the - /// newer-to-older field for a negative one. - /// - /// `center_t` is the centre frame's logical position in the ring. - /// - /// This does nothing unless `Chained` estimation is active. When it - /// is, `dispatch::run_motion_compensation` calls it once per - /// neighbour on every submit and then cleans up the seed with - /// `run_seeded_refine`. Tests call it directly as well. + /// `k` must be nonzero and land inside the + /// ring, at most twice the temporal radius either way. The walk takes one hop per step out from + /// the centre `center_t`, following the older-to-newer field for a positive `k` and the + /// newer-to-older field for a negative one. It does nothing unless `Chained` estimation is + /// active. pub(crate) fn run_chain_compose( &self, center_t: u32, k: i32, neighbour_idx: u32, ) -> Result<(), anyhow::Error> { - let Some(mc) = self.mc_ctx.as_ref() else { + let Some(motion_ctx) = self.mc_ctx.as_ref() else { return Ok(()); }; + if !self.is_chained() || k == 0 { return Ok(()); } @@ -273,11 +237,11 @@ impl NlmDenoiser { }; let start_pair_slot = self.pair_slot(start_gap); let pair_ring_slots = super::pair_ring_slot_count(radius); - let pair_ring_len = pair_ring_slots as usize * mc.pair_slot_stride() as usize; + let pair_ring_len = pair_ring_slots as usize * motion_ctx.pair_slot_stride() as usize; dispatch_chain_compose::( &self.client, - mc, + motion_ctx, self.width, self.height, pair_ring_slots, @@ -298,76 +262,75 @@ mod tests { use crate::nlmeans::align::StorageAlign; use crate::nlmeans::motion::{MotionCompensationMode, MotionEstimation}; + /// A 4-pixel-high frame cut into 4x4 blocks with no overlap. + fn small_block_ctx(width: u32) -> MotionCtx { + let mode = MotionCompensationMode::Mvtools { + blksize: 4, + overlap: 0, + search_radius: 1, + pyramid_levels: 1, + estimation: MotionEstimation::Direct, + }; + let align = StorageAlign::new(32); + + MotionCtx::new(mode, width, 4, align).unwrap() + } + #[test] fn neighbour_idx_for_k_matches_dispatch_convention() { - // The same walk `dispatch::neighbour_idx_for_k`'s own tests - // check. The negative offsets come first, taking indices 0 up - // to radius minus 1, then the positive ones follow. - assert_eq!(neighbour_idx_for_k(2, -2), 0); - assert_eq!(neighbour_idx_for_k(2, -1), 1); - assert_eq!(neighbour_idx_for_k(2, 1), 2); - assert_eq!(neighbour_idx_for_k(2, 2), 3); + // Negative offsets take indices 0 up to radius minus 1, then the positive ones follow. + let furthest_back = neighbour_idx_for_k(2, -2); + let nearest_back = neighbour_idx_for_k(2, -1); + let nearest_forward = neighbour_idx_for_k(2, 1); + let furthest_forward = neighbour_idx_for_k(2, 2); + + assert_eq!(furthest_back, 0); + assert_eq!(nearest_back, 1); + assert_eq!(nearest_forward, 2); + assert_eq!(furthest_forward, 3); } #[test] fn pair_byte_offset_pads_small_block_counts_to_32_bytes() { - // A 4x4 frame at this geometry has a single block, so the - // unpadded direction stride is only 8 bytes and would leave - // direction 1 at an offset that is not 32-aligned. - let m = MotionCtx::new( - MotionCompensationMode::Mvtools { - blksize: 4, - overlap: 0, - search_radius: 1, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - }, - 4, - 4, - StorageAlign::new(32), - ) - .unwrap(); + // One block gives an unpadded direction stride of 8 bytes, which would leave direction 1 + // off a 32-byte boundary. + let ctx = small_block_ctx(4); assert_eq!( - m.blocks_x * m.blocks_y, + ctx.blocks_x * ctx.blocks_y, 1, "fixture should have exactly one block" ); - assert_eq!(pair_byte_offset(&m, 0, 0), 0); - assert_eq!(pair_byte_offset(&m, 0, 1), 32); - assert_eq!(pair_byte_offset(&m, 1, 0), 64); - assert_eq!(pair_byte_offset(&m, 1, 1), 96); + + let slot_0_forward = pair_byte_offset(&ctx, 0, 0); + let slot_0_backward = pair_byte_offset(&ctx, 0, 1); + let slot_1_forward = pair_byte_offset(&ctx, 1, 0); + let slot_1_backward = pair_byte_offset(&ctx, 1, 1); + + assert_eq!(slot_0_forward, 0); + assert_eq!(slot_0_backward, 32); + assert_eq!(slot_1_forward, 64); + assert_eq!(slot_1_backward, 96); } #[test] fn pair_byte_offset_direction_one_pads_even_when_slot_base_is_aligned() { - // An 8x4 frame at this geometry has two blocks. The unpadded - // per-slot stride comes to 32 bytes, which is already aligned, - // so every slot's own base offset would be fine. - // - // The per-direction stride inside that slot is only 16 bytes - // though, so direction 1 still needs padding of its own even - // though direction 0's slot base never did. - let m = MotionCtx::new( - MotionCompensationMode::Mvtools { - blksize: 4, - overlap: 0, - search_radius: 1, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - }, - 8, - 4, - StorageAlign::new(32), - ) - .unwrap(); + // Two blocks give an aligned 32-byte unpadded slot stride, but the 16-byte direction stride + // inside it still needs padding of its own. + let ctx = small_block_ctx(8); assert_eq!( - m.blocks_x * m.blocks_y, + ctx.blocks_x * ctx.blocks_y, 2, "fixture should have exactly two blocks" ); - assert_eq!(pair_byte_offset(&m, 0, 0), 0); - assert_eq!(pair_byte_offset(&m, 0, 1), 32); - assert_eq!(pair_byte_offset(&m, 1, 0), 64); - assert_eq!(pair_byte_offset(&m, 1, 1), 96); + + let slot_0_forward = pair_byte_offset(&ctx, 0, 0); + let slot_0_backward = pair_byte_offset(&ctx, 0, 1); + let slot_1_forward = pair_byte_offset(&ctx, 1, 0); + let slot_1_backward = pair_byte_offset(&ctx, 1, 1); + + assert_eq!(slot_0_forward, 0); + assert_eq!(slot_0_backward, 32); + assert_eq!(slot_1_forward, 64); + assert_eq!(slot_1_backward, 96); } } diff --git a/av-denoise-core/src/nlmeans/motion/compensate.rs b/av-denoise-core/src/nlmeans/motion/compensate.rs index 84e4edc..a4c5f44 100644 --- a/av-denoise-core/src/nlmeans/motion/compensate.rs +++ b/av-denoise-core/src/nlmeans/motion/compensate.rs @@ -5,18 +5,14 @@ use super::MotionCtx; use super::analyse::mv_field_byte_offset; use crate::nlmeans::kernels::motion::nlm_mc_warp; -/// Shifts one neighbour frame into line with the centre frame. -/// -/// The result is a full frame written into `compensated` at the -/// neighbour's slot. +/// Shifts one neighbour frame into line with the centre frame, into its slot of `dst`. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] pub(crate) fn run_compensate( client: &ComputeClient, - mc: &MotionCtx, - channels: u32, + motion_ctx: &MotionCtx, stored_ch: u32, width: u32, height: u32, @@ -27,17 +23,17 @@ pub(crate) fn run_compensate( dst: &Handle, mv_field: &Handle, ) -> Result<(), anyhow::Error> { - let _ = channels; let block_x = 16u32; let block_y = 16u32; - let grid = CubeCount::new_2d(width.div_ceil(block_x), height.div_ceil(block_y)); + let cubes_x = width.div_ceil(block_x); + let cubes_y = height.div_ceil(block_y); + let grid = CubeCount::new_2d(cubes_x, cubes_y); let dim = CubeDim::new_2d(block_x, block_y); let total_pixels = (frame_count * height * width * stored_ch) as usize; - let mv_slice_len = (mc.blocks_x as usize) * (mc.blocks_y as usize) * 2; - let mv_slice = mv_field - .clone() - .offset_start(mv_field_byte_offset(mc, neighbour_idx)); + let mv_slice_len = (motion_ctx.blocks_x as usize) * (motion_ctx.blocks_y as usize) * 2; + let mv_offset = mv_field_byte_offset(motion_ctx, neighbour_idx); + let mv_slice = mv_field.clone().offset_start(mv_offset); unsafe { nlm_mc_warp::launch_unchecked::( @@ -50,9 +46,9 @@ pub(crate) fn run_compensate( ArrayArg::from_raw_parts(mv_slice, mv_slice_len), neighbour_slot, neighbour_slot, - mc.step, - mc.blocks_x, - mc.blocks_y, + motion_ctx.step, + motion_ctx.blocks_x, + motion_ctx.blocks_y, width, height, ); diff --git a/av-denoise-core/src/nlmeans/motion/confidence.rs b/av-denoise-core/src/nlmeans/motion/confidence.rs index d57999c..3587744 100644 --- a/av-denoise-core/src/nlmeans/motion/confidence.rs +++ b/av-denoise-core/src/nlmeans/motion/confidence.rs @@ -6,55 +6,34 @@ use super::analyse::confidence_byte_offset; use super::pyramid::{level_dims, pyramid_slot_byte_offset}; use crate::nlmeans::kernels::motion::nlm_mc_block_match_fine; -/// How far a pixel can be off, on average, before a block counts as no -/// longer matching. +/// The average per-pixel error at which a block stops counting as a match. /// -/// The value is in normalised luma units and works out at roughly 5 of -/// 255, which is the reference point MDegrain uses. -/// -/// [`thsad`] scales it by block area and by `thsad_scale` to get the -/// threshold a caller compares against. +/// It is in normalised luma units, roughly 5 of 255, which is the reference point MDegrain uses. pub(crate) const THSAD_PIXEL: f32 = 0.02; -/// The score two noisy copies of the same content are expected to -/// produce by chance. -/// -/// Subtracting this first means a block is only judged on how far its -/// content really differs, not on the noise it happens to carry. +/// The block SAD two noisy copies of the same content are expected to score by chance. /// -/// Each pixel's noisy-against-noisy absolute difference is the magnitude -/// of a zero-mean Gaussian with scale `sigma * sqrt(2)`, whose mean is -/// `2 * sigma / sqrt(pi)`. That is then summed over all `blksize^2` -/// pixels in the block. +/// Subtracting it first judges a block on its content rather than its noise. Each pixel's +/// noisy-against-noisy difference is the magnitude of a zero-mean Gaussian with scale +/// `sigma * sqrt(2)`, whose mean is `2 * sigma / sqrt(pi)`, summed over all `blksize^2` pixels. pub(crate) fn sad_noise_floor(blksize: u32, sigma_y: f32) -> f32 { let block_area = (blksize * blksize) as f32; block_area * 2.0 * sigma_y / std::f32::consts::PI.sqrt() } -/// How far past the noise floor a block can score before its confidence -/// reaches zero. +/// How far past the noise floor a block can score before its confidence reaches zero. /// -/// It is scaled by block area, so the value stays comparable across -/// block sizes. -/// -/// `thsad_scale` is the user-facing multiplier, exposed as -/// [`crate::nlmeans::HqParams::thsad_scale`]. +/// It scales with block area so it stays comparable across block sizes. `thsad_scale` is +/// [HqParams::thsad_scale](crate::nlmeans::HqParams::thsad_scale). pub(crate) fn thsad(blksize: u32, thsad_scale: f32) -> f32 { let block_area = (blksize * blksize) as f32; thsad_scale * block_area * THSAD_PIXEL } -/// Scores how well one neighbour matches the centre frame, without -/// searching for motion. -/// -/// This is what runs when confidence weighting is on but motion -/// compensation is off. Each block is scored exactly where it stands, at -/// level 0 of the pyramid, with no coarse seed and no search window. +/// Scores how well one neighbour matches the centre frame without a motion search. /// -/// The result goes into `confidence` at the slot reserved for this -/// neighbour. The motion vector the kernel also produces is thrown away -/// into `mv_scratch`, because with motion compensation off nothing would -/// use it. +/// Each block is scored where it stands at pyramid level 0, into this neighbour's `confidence` +/// slot. The motion vector the kernel also writes goes to `mv_scratch` and is discarded. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" @@ -74,24 +53,12 @@ pub(crate) fn run_confidence_for_neighbour( sad_noise_floor: f32, thsad: f32, ) -> Result<(), anyhow::Error> { - let (fw, fh) = level_dims(width, height, 0); - let centre = luma_pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - 0, - centre_slot, - ctx.align, - )); - let neighbour = luma_pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - 0, - neighbour_slot, - ctx.align, - )); - let level_len = (fw * fh) as usize; + let (fine_width, fine_height) = level_dims(width, height, 0); + let centre_offset = pyramid_slot_byte_offset(width, height, frame_count, 0, centre_slot, ctx.align); + let centre = luma_pyramid.clone().offset_start(centre_offset); + let neighbour_offset = pyramid_slot_byte_offset(width, height, frame_count, 0, neighbour_slot, ctx.align); + let neighbour = luma_pyramid.clone().offset_start(neighbour_offset); + let level_len = (fine_width * fine_height) as usize; let conf_offset = confidence_byte_offset(ctx, neighbour_idx); let conf_slot = confidence.clone().offset_start(conf_offset); @@ -110,11 +77,11 @@ pub(crate) fn run_confidence_for_neighbour( ArrayArg::from_raw_parts(neighbour, level_len), ArrayArg::from_raw_parts(mv_scratch.clone(), mv_slot_len), ArrayArg::from_raw_parts(conf_slot, conf_slot_len), - true, // the no-MC confidence pass always wants its output + true, sad_noise_floor, thsad, - fw, - fh, + fine_width, + fine_height, ctx.blksize, ctx.step, ctx.search_radius, @@ -140,7 +107,8 @@ mod tests { #[test] fn sad_noise_floor_zero_for_zero_sigma() { - assert_eq!(sad_noise_floor(16, 0.0), 0.0); + let floor = sad_noise_floor(16, 0.0); + assert_eq!(floor, 0.0); } #[test] @@ -152,6 +120,7 @@ mod tests { #[test] fn thsad_default_scale_is_positive() { - assert!(thsad(16, 1.0) > 0.0); + let threshold = thsad(16, 1.0); + assert!(threshold > 0.0); } } diff --git a/av-denoise-core/src/nlmeans/motion/mod.rs b/av-denoise-core/src/nlmeans/motion/mod.rs index 2f02888..33d5be5 100644 --- a/av-denoise-core/src/nlmeans/motion/mod.rs +++ b/av-denoise-core/src/nlmeans/motion/mod.rs @@ -1,71 +1,49 @@ -//! Following motion between frames so temporal denoising stays sharp. -//! -//! Temporal denoising averages a pixel with the same position in nearby -//! frames. When the camera or the content moves, that position holds -//! different content in each frame, and averaging it blurs the moving -//! parts. -//! -//! This module works out where each block of pixels moved to, then -//! shifts the neighbouring frames back into line with the current one -//! before the denoising weights are computed. -//! -//! # How a frame is tracked -//! -//! `pyramid` builds a stack of progressively smaller copies of the luma -//! plane. A search on a small copy finds large movements cheaply, and -//! the answer then seeds a short search at full resolution. -//! -//! `analyse` runs that search. `chain` handles distant neighbours by -//! measuring motion between adjacent frames and joining the results, -//! which reaches further than any single search window. -//! -//! `confidence` scores how well each block actually matched, so a block -//! that was occluded or changed can be held back rather than blurred in. -//! -//! `compensate` applies the finished motion field to a frame. - mod analyse; mod chain; mod compensate; mod confidence; mod pyramid; -pub(crate) use analyse::{confidence_byte_offset, mv_field_byte_offset, run_analyse, run_seeded_refine}; -#[cfg(all(test, any(feature = "vulkan", feature = "metal")))] -pub(crate) use chain::pair_byte_offset; -pub(crate) use chain::{neighbour_idx_for_k, run_pair_analyse, zero_pair_slot}; -pub(crate) use compensate::run_compensate; -pub(crate) use confidence::{THSAD_PIXEL, run_confidence_for_neighbour, sad_noise_floor, thsad}; use cubecl::prelude::*; use cubecl::server::Handle; -pub(crate) use pyramid::{level_dims, pyramid_pixels_per_frame, pyramid_slot_byte_offset, run_pyramid_build}; -use crate::nlmeans::align::StorageAlign; +pub(crate) use self::analyse::{ + confidence_byte_offset, + mv_field_byte_offset, + run_analyse, + run_seeded_refine, +}; +#[cfg(all(test, any(feature = "vulkan", feature = "metal")))] +pub(crate) use self::chain::pair_byte_offset; +pub(crate) use self::chain::{neighbour_idx_for_k, run_pair_analyse, zero_pair_slot}; +pub(crate) use self::compensate::run_compensate; +pub(crate) use self::confidence::{THSAD_PIXEL, run_confidence_for_neighbour, sad_noise_floor, thsad}; +pub(crate) use self::pyramid::{ + level_dims, + pyramid_pixels_per_frame, + pyramid_slot_byte_offset, + run_pyramid_build, +}; +use super::align::StorageAlign; /// The motion search's tuning, for a denoiser that always tracks motion. /// -/// [`MotionCompensationMode::Mvtools`] carries the same five values for -/// a denoiser that can also turn motion compensation off. +/// [MotionCompensationMode::Mvtools] carries the same five values. #[derive(Debug, Clone, Copy, PartialEq)] pub struct MotionSearch { - /// The side length of each motion-search block, in pixels at the - /// finest pyramid level. + /// The side length of each search block, in pixels at the finest pyramid level. pub blksize: u32, /// How many pixels neighbouring blocks overlap. /// - /// This has to be strictly below `blksize`, so the step between - /// blocks stays positive. + /// It must be strictly below `blksize` so the step between blocks stays positive. pub overlap: u32, /// The search radius in pixels at the finest pyramid level. /// - /// The coarse pass uses the same radius on a half-size image, so - /// its real reach is twice as far. + /// The coarse pass uses the same radius on a half-size image, so its reach is twice as far. pub search_radius: u32, - /// How many levels the pyramid has. + /// How many levels the pyramid has, up to [MAX_PYRAMID_LEVELS]. /// - /// `1` means a single full-resolution search. `2` adds a half-size - /// coarse pass that seeds the fine one. The maximum is - /// [`MAX_PYRAMID_LEVELS`]. + /// `1` searches at full resolution only, and `2` adds a half-size coarse pass that seeds it. pub pyramid_levels: u32, /// How motion toward each temporal neighbour is estimated. pub estimation: MotionEstimation, @@ -102,39 +80,26 @@ pub enum MotionCompensationMode { /// Motion compensation is off, and no extra buffers are allocated. #[default] None, - /// An estimator inspired by MVTools tracks each block, and the - /// neighbouring frames are shifted toward the centre frame at - /// denoise time. + /// An MVTools-style estimator tracks each block, and the neighbouring frames are shifted toward + /// the centre frame at denoise time. Mvtools { - /// The side length of each motion-search block, in pixels at - /// the finest pyramid level. + /// The side length of each search block, in pixels at the finest pyramid level. blksize: u32, /// How many pixels neighbouring blocks overlap. /// - /// This has to be strictly below `blksize`, so the step between - /// blocks stays positive. - /// - /// Anything above 0 leaves room for the raised-cosine blend in - /// the compensate step, which currently uses a - /// winner-takes-all rule instead. + /// It must be strictly below `blksize` so the step between blocks stays positive. overlap: u32, /// The search radius in pixels at the finest pyramid level. /// - /// The coarse pass uses the same radius on a half-size image, so - /// its real reach is twice as far. + /// The coarse pass uses the same radius on a half-size image, so its reach is twice as far. search_radius: u32, - /// How many levels the pyramid has. + /// How many levels the pyramid has, up to [MAX_PYRAMID_LEVELS]. /// - /// `1` means a single full-resolution search. `2` adds a - /// half-size coarse pass that seeds the fine one. The maximum is - /// [`MAX_PYRAMID_LEVELS`]. + /// `1` searches at full resolution only, and `2` adds a half-size coarse pass that seeds it. pyramid_levels: u32, /// How motion toward each temporal neighbour is estimated. /// - /// `Auto`, the default, picks a strategy from the temporal - /// radius and is what callers normally want. Naming `Direct` or - /// `Chained` is mostly useful for pinning one strategy in tests - /// and benches. + /// `Auto`, the default, picks a strategy from the temporal radius. estimation: MotionEstimation, }, } @@ -143,63 +108,45 @@ pub enum MotionCompensationMode { #[non_exhaustive] #[derive(Debug, Default, Clone, Copy, PartialEq)] pub enum MotionEstimation { - /// Picks `Direct` or `Chained` from the temporal radius when the - /// denoiser is built. - /// - /// [`MotionEstimation::resolve`] describes the rule and where it - /// came from. + /// Picks `Direct` or `Chained` from the temporal radius, as [MotionEstimation::resolve] does. #[default] Auto, - /// Matches every neighbour against the centre frame directly, at the - /// configured search radius. + /// Matches every neighbour against the centre frame directly at the configured search radius. /// - /// The cost grows with the temporal radius, because each neighbour - /// repeats the whole coarse and fine search. + /// Its cost grows with the temporal radius, because each neighbour repeats the coarse and fine + /// search. Direct, - /// Measures motion only between adjacent frames, once per pushed - /// frame. + /// Measures motion between adjacent frames once per push and joins it into a seed per + /// neighbour. /// - /// Those per-step vectors are then joined into a seed for each - /// neighbour, and a small seeded search cleans up whatever drift is - /// left. + /// A small seeded search then cleans up whatever drift is left. Chained { - /// The search radius for the seeded refinement pass, in pixels - /// at the finest pyramid level. + /// The seeded refinement's search radius, in pixels at the finest pyramid level. /// - /// It can be small, because the joined seed already carries most - /// of the real movement. + /// It can be small because the joined seed already carries most of the movement. refine_radius: u32, }, } -/// The default refinement radius for [`MotionEstimation::Chained`]. +/// The default refinement radius for [MotionEstimation::Chained]. pub const DEFAULT_REFINE_RADIUS: u32 = 2; -/// The temporal radius at which [`MotionEstimation::Auto`] switches from -/// `Direct` to `Chained`. -/// -/// Below this, `Direct` tracks slightly better, because the real motion -/// still fits inside its own search window. +/// The temporal radius at which [MotionEstimation::Auto] switches from `Direct` to `Chained`. /// -/// At or above it, `Chained` both stays inside its window and runs -/// faster, because its reach grows with the radius rather than being -/// capped by a fixed window. +/// Below it `Direct` tracks slightly better, because the motion still fits its search window. At or +/// above it `Chained` stays inside its window and runs faster, because its reach grows with the +/// radius. pub const CHAINED_RADIUS_THRESHOLD: u32 = 3; impl MotionEstimation { - /// Builds a `Chained` estimation with the library's default - /// refinement radius. + /// A `Chained` estimation with [DEFAULT_REFINE_RADIUS]. pub fn chained_default() -> Self { Self::Chained { refine_radius: DEFAULT_REFINE_RADIUS, } } - /// Resolves `Auto` against the temporal radius, always returning a - /// concrete `Direct` or `Chained`. - /// - /// `Direct` and `Chained` pass through unchanged whatever the - /// radius. [`CHAINED_RADIUS_THRESHOLD`] is where the switch happens. + /// Resolves `Auto` at [CHAINED_RADIUS_THRESHOLD], passing `Direct` and `Chained` through. pub fn resolve(self, temporal_radius: u32) -> Self { match self { Self::Auto if temporal_radius >= CHAINED_RADIUS_THRESHOLD => Self::chained_default(), @@ -224,49 +171,42 @@ impl MotionEstimation { } } -/// The default block size, matching MVTools and lining up well with the -/// patch sizes NLM typically uses. +/// The default block size, matching MVTools and the patch sizes NLM typically uses. pub const DEFAULT_BLKSIZE: u32 = 16; -/// The default block overlap, which is half the default block size. pub const DEFAULT_OVERLAP: u32 = 8; /// The default search radius at the finest level. /// -/// With a two-level pyramid this reaches motion of roughly 12 pixels at -/// full resolution. +/// With a two-level pyramid it reaches motion of roughly 12 pixels at full resolution. pub const DEFAULT_SEARCH_RADIUS: u32 = 4; /// The default number of pyramid levels. /// -/// Two levels give a single half-size coarse pass, which handles most -/// heavy-motion anime while keeping the number of kernel launches down. +/// One half-size coarse pass handles most heavy-motion anime while keeping kernel launches down. pub const DEFAULT_PYRAMID_LEVELS: u32 = 2; /// The hard ceiling on `pyramid_levels`. /// -/// Each extra level halves the resolution again and adds a kernel launch -/// per neighbour. Three is already more than 1080p content needs. +/// Each extra level halves the resolution again and adds a kernel launch per neighbour. Three is +/// already more than 1080p content needs. pub const MAX_PYRAMID_LEVELS: u32 = 3; /// The hard ceiling on `search_radius`. /// -/// The analyse kernel scores a `(2 * radius + 1)^2` window per block, so -/// the cost grows with the square of the radius. +/// The analyse kernel scores a `(2 * radius + 1)^2` window per block, so the cost grows with the +/// square of the radius. pub const MAX_SEARCH_RADIUS: u32 = 8; /// The hard ceiling on `blksize`. /// -/// Above this the per-block shared-memory tile grows uncomfortably large -/// on RDNA-class GPUs. +/// Above it the per-block shared-memory tile grows too large on RDNA-class GPUs. pub const MAX_BLKSIZE: u32 = 32; impl MotionCompensationMode { - /// Builds an `Mvtools` mode from the library defaults. + /// An `Mvtools` mode from the library defaults. /// - /// This pins `estimation` to `Direct` rather than the field's own - /// `Auto` default, so it never switches to `Chained` at larger - /// temporal radii the way an `Auto` configuration would. + /// It pins `estimation` to `Direct`, so it never switches to `Chained` at larger radii. pub fn mvtools_default() -> Self { Self::Mvtools { blksize: DEFAULT_BLKSIZE, @@ -277,23 +217,20 @@ impl MotionCompensationMode { } } - /// Whether motion compensation is active at all. pub(crate) fn is_active(self) -> bool { !matches!(self, Self::None) } - /// The estimation strategy this mode resolves to at - /// `temporal_radius`. + /// The concrete estimation strategy this mode resolves to at `temporal_radius`. /// - /// Returns `None` when the mode is not `Mvtools`, and never returns - /// `Auto`. See [`MotionEstimation::resolve`]. - /// - /// Every decision that depends on the strategy goes through here, - /// including pair-ring allocation, whether the push-time pair - /// analyse runs, and which branch the submit path takes. + /// It is `None` when the mode is not `Mvtools` and is never `Auto`. Every decision that depends + /// on the strategy goes through here, so they all agree. pub(crate) fn resolved_estimation(&self, temporal_radius: u32) -> Option { match *self { - Self::Mvtools { estimation, .. } => Some(estimation.resolve(temporal_radius)), + Self::Mvtools { estimation, .. } => { + let resolved = estimation.resolve(temporal_radius); + Some(resolved) + }, Self::None => None, } } @@ -316,27 +253,32 @@ impl MotionCompensationMode { "motion-compensation blksize={blksize} is too small, the minimum is 4 pixels per side" ); } + if blksize > MAX_BLKSIZE { anyhow::bail!( "motion-compensation blksize={blksize} exceeds the supported maximum of {MAX_BLKSIZE}" ); } + if blksize % 2 != 0 { anyhow::bail!( "motion-compensation blksize={blksize} must be even so the /2 coarse level is well-defined" ); } + if overlap >= blksize { anyhow::bail!( "motion-compensation overlap={overlap} must be strictly less than blksize, \ which is {blksize}, so the step between blocks stays positive" ); } + if search_radius == 0 || search_radius > MAX_SEARCH_RADIUS { anyhow::bail!( "motion-compensation search_radius={search_radius} must be in 1..={MAX_SEARCH_RADIUS}" ); } + if pyramid_levels == 0 || pyramid_levels > MAX_PYRAMID_LEVELS { anyhow::bail!( "motion-compensation pyramid_levels={pyramid_levels} must be in 1..={MAX_PYRAMID_LEVELS}" @@ -349,14 +291,9 @@ impl MotionCompensationMode { } } -/// The motion-compensation state a `NlmDenoiser` holds while motion -/// compensation is active. +/// The block geometry motion compensation runs with, worked out once when the denoiser is built. /// -/// It is worked out once at construction, so the hot dispatch path never -/// has to re-read the configuration enum. -/// -/// Only the fields the analyse and compensate dispatchers use live here. -/// The full configuration stays on [`MotionCompensationMode`]. +/// It keeps the hot dispatch path off the configuration enum. #[derive(Debug, Clone)] pub(crate) struct MotionCtx { pub blksize: u32, @@ -365,11 +302,8 @@ pub(crate) struct MotionCtx { pub pyramid_levels: u32, pub blocks_x: u32, pub blocks_y: u32, - /// The alignment every buffer this context slices per slot has to - /// respect, meaning the motion field, the confidence buffer, the - /// pair ring, and the pyramid. - /// - /// It is read from the runtime. See [`StorageAlign`]. + /// The alignment each per-slot slice of the motion field, confidence, pair ring and pyramid + /// must respect. pub align: StorageAlign, } @@ -401,91 +335,61 @@ impl MotionCtx { }) } - /// How many motion-field slots each neighbour needs, which is one - /// per block. pub fn mv_slots_per_neighbour(&self) -> usize { (self.blocks_x * self.blocks_y) as usize } - /// The padded per-neighbour motion-field stride in bytes. + /// The per-neighbour motion-field stride in bytes, two `i32` per block padded to the alignment. /// - /// Each block stores two `i32` components, and the total is rounded - /// up to the runtime's buffer-binding alignment. This is the same - /// convention [`Self::confidence_bytes_per_neighbour`] uses. - /// - /// The padding matters because wgpu rejects a bind-group offset that - /// is not a multiple of its `min_storage_buffer_offset_alignment`, - /// and an odd block count leaves the unpadded stride short of that - /// boundary. + /// wgpu rejects a bind-group offset that is not a multiple of its + /// `min_storage_buffer_offset_alignment`, and an odd block count leaves the unpadded stride + /// short of it. pub(crate) fn mv_field_bytes_per_neighbour(&self) -> u64 { let blocks = (self.blocks_x as u64) * (self.blocks_y as u64); - self.align.pad_bytes(blocks * 2 * size_of::() as u64) + let unpadded = blocks * 2 * size_of::() as u64; + self.align.pad_bytes(unpadded) } - /// The padded per-neighbour confidence-buffer stride in bytes. - /// - /// Each block stores one `f32`, rounded up to the runtime's - /// buffer-binding alignment the same way - /// [`Self::mv_field_bytes_per_neighbour`] rounds the motion field. + /// The per-neighbour confidence stride in bytes, one `f32` per block padded to the alignment. pub(crate) fn confidence_bytes_per_neighbour(&self) -> u64 { let blocks = (self.blocks_x as u64) * (self.blocks_y as u64); - self.align.pad_bytes(blocks * size_of::() as u64) + let unpadded = blocks * size_of::() as u64; + self.align.pad_bytes(unpadded) } - /// How many `i32` elements one direction of a pair-ring slot holds, - /// which is two per block. - /// - /// This is the unpadded count of the data itself. It is the - /// zero-fill length in `zero_pair_slot` and the input - /// [`Self::pair_direction_bytes`] pads. + /// The unpadded `i32` count of one pair-ring direction, two per block. pub(crate) fn pair_direction_len(&self) -> u32 { self.blocks_x * self.blocks_y * 2 } - /// The padded per-direction pair-ring stride in bytes, rounded up to - /// the runtime's buffer-binding alignment the same way - /// [`Self::confidence_bytes_per_neighbour`] rounds its own stride. + /// The per-direction pair-ring stride in bytes, padded to the alignment. /// - /// Both `pair_byte_offset`, which the host writes and zero-fills - /// at, and the chain-compose kernel's internal read stride use this - /// padded value, so a direction's data starts in the same place for - /// every reader and writer. + /// The host offsets and the chain-compose kernel's read stride both use it, so every reader and + /// writer finds a direction's data in the same place. pub(crate) fn pair_direction_bytes(&self) -> u64 { - self.align - .pad_bytes(self.pair_direction_len() as u64 * size_of::() as u64) + let unpadded = self.pair_direction_len() as u64 * size_of::() as u64; + self.align.pad_bytes(unpadded) } - /// The padded per-slot pair-ring stride in bytes, covering both - /// directions back to back. + /// The per-slot pair-ring stride in bytes, with both directions back to back. pub(crate) fn pair_slot_bytes(&self) -> u64 { 2 * self.pair_direction_bytes() } - /// The padded per-direction pair-ring stride in `i32` elements. - /// - /// The chain-compose kernel reads the whole pair ring as one array - /// and steps through it with this value, which matches - /// [`Self::pair_direction_bytes`] exactly. + /// [Self::pair_direction_bytes] in `i32` elements, which the chain-compose kernel steps by. pub(crate) fn pair_direction_stride(&self) -> u32 { (self.pair_direction_bytes() / size_of::() as u64) as u32 } - /// The padded per-slot pair-ring stride in `i32` elements, covering - /// both directions back to back. + /// [Self::pair_slot_bytes] in `i32` elements. pub(crate) fn pair_slot_stride(&self) -> u32 { 2 * self.pair_direction_stride() } - /// The block geometry for a confidence pass with no motion - /// compensation. - /// - /// It uses the library's default block size and overlap, one pyramid - /// level so there is no coarse pass, and a search radius of zero, so - /// each block is scored where it stands with no motion search at - /// all. + /// The block geometry for a confidence pass without motion compensation. /// - /// This is what runs when confidence weighting is on but no - /// `Mvtools` mode was configured to take geometry from. + /// It uses the default block size and overlap with one pyramid level and a search radius of 0, + /// so each block is scored where it stands. pub(crate) fn confidence_only(width: u32, height: u32, align: StorageAlign) -> Self { Self::new( MotionCompensationMode::Mvtools { @@ -503,36 +407,25 @@ impl MotionCtx { } } -/// How many slots the pair ring needs for a given temporal radius. +/// How many slots the pair ring needs for a temporal radius. /// -/// The pair ring stores one adjacent-frame motion field per gap between -/// consecutive frames in the temporal window. A window of -/// `2 * radius + 1` frames has exactly `2 * radius` gaps. -/// -/// A gap's field is only read while both of its frames are still inside -/// some window, which lasts exactly `2 * radius` frame pushes. -/// -/// Sizing the ring to match means a slot is reused precisely when its -/// old contents stop being needed, and never sooner. -/// `NlmDenoiser::pair_slot` works this out in full. +/// A window of `2 * radius + 1` frames has `2 * radius` gaps, each holding one adjacent-frame field. +/// A gap's field is read only while both its frames sit in some window, which lasts exactly +/// `2 * radius` pushes, so a slot is reused exactly when its old field stops being needed. pub(crate) fn pair_ring_slot_count(temporal_radius: u32) -> u32 { 2 * temporal_radius } -/// Builds the pyramid for the slot `push_frame` just uploaded. +/// Builds the pyramid for the slot a push just uploaded. /// -/// Level 0 luma is always extracted, and the smaller levels follow when -/// `pyramid_levels` is above 1. -/// -/// This is a thin wrapper around [`run_pyramid_build`], which already -/// handles both cases itself. +/// Level 0 luma is always extracted, and the smaller levels follow when `pyramid_levels` is above 1. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] pub(crate) fn build_pyramid_for_slot( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, width: u32, height: u32, frame_count: u32, @@ -543,7 +436,7 @@ pub(crate) fn build_pyramid_for_slot( ) -> Result<(), anyhow::Error> { run_pyramid_build::( client, - mc, + motion_ctx, width, height, frame_count, @@ -560,45 +453,45 @@ mod tests { #[test] fn none_is_inactive() { - let m = MotionCompensationMode::None; - assert!(!m.is_active()); - m.validate().unwrap(); + let mode = MotionCompensationMode::None; + assert!(!mode.is_active()); + mode.validate().unwrap(); } #[test] fn mvtools_default_is_active() { - let m = MotionCompensationMode::mvtools_default(); - assert!(m.is_active()); - m.validate().unwrap(); + let mode = MotionCompensationMode::mvtools_default(); + assert!(mode.is_active()); + mode.validate().unwrap(); } #[test] fn validate_rejects_tiny_blksize() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 2, overlap: 0, search_radius: 4, pyramid_levels: 2, estimation: MotionEstimation::Direct, }; - assert!(m.validate().is_err()); + assert!(mode.validate().is_err()); } #[test] fn validate_rejects_odd_blksize() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 9, overlap: 0, search_radius: 4, pyramid_levels: 2, estimation: MotionEstimation::Direct, }; - assert!(m.validate().is_err()); + assert!(mode.validate().is_err()); } #[test] fn validate_rejects_overlap_equal_to_blksize() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 16, search_radius: 4, @@ -606,112 +499,117 @@ mod tests { estimation: MotionEstimation::Direct, }; // An overlap equal to blksize would leave a step of 0. - assert!(m.validate().is_err()); + assert!(mode.validate().is_err()); } #[test] fn validate_accepts_half_overlap() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 8, search_radius: 4, pyramid_levels: 2, estimation: MotionEstimation::Direct, }; - m.validate().unwrap(); + mode.validate().unwrap(); } #[test] fn validate_rejects_zero_search_radius() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 4, search_radius: 0, pyramid_levels: 2, estimation: MotionEstimation::Direct, }; - assert!(m.validate().is_err()); + assert!(mode.validate().is_err()); } #[test] fn validate_rejects_zero_pyramid_levels() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 4, search_radius: 4, pyramid_levels: 0, estimation: MotionEstimation::Direct, }; - assert!(m.validate().is_err()); + assert!(mode.validate().is_err()); } #[test] fn chained_default_is_valid() { - let m = MotionCompensationMode::Mvtools { + let estimation = MotionEstimation::chained_default(); + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 8, search_radius: 4, pyramid_levels: 2, - estimation: MotionEstimation::chained_default(), + estimation, }; - m.validate().unwrap(); - assert_eq!( - m, - MotionCompensationMode::Mvtools { - blksize: 16, - overlap: 8, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Chained { - refine_radius: DEFAULT_REFINE_RADIUS - }, - } - ); + mode.validate().unwrap(); + + let expected_estimation = MotionEstimation::Chained { + refine_radius: DEFAULT_REFINE_RADIUS, + }; + let expected = MotionCompensationMode::Mvtools { + blksize: 16, + overlap: 8, + search_radius: 4, + pyramid_levels: 2, + estimation: expected_estimation, + }; + assert_eq!(mode, expected); } #[test] fn validate_rejects_zero_refine_radius() { - let m = MotionCompensationMode::Mvtools { + let estimation = MotionEstimation::Chained { refine_radius: 0 }; + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 8, search_radius: 4, pyramid_levels: 2, - estimation: MotionEstimation::Chained { refine_radius: 0 }, + estimation, }; - assert!(m.validate().is_err()); + assert!(mode.validate().is_err()); } #[test] fn validate_rejects_refine_radius_above_max() { - let m = MotionCompensationMode::Mvtools { + let estimation = MotionEstimation::Chained { + refine_radius: MAX_SEARCH_RADIUS + 1, + }; + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 8, search_radius: 4, pyramid_levels: 2, - estimation: MotionEstimation::Chained { - refine_radius: MAX_SEARCH_RADIUS + 1, - }, + estimation, }; - assert!(m.validate().is_err()); + assert!(mode.validate().is_err()); } #[test] fn validate_accepts_refine_radius_at_max() { - let m = MotionCompensationMode::Mvtools { + let estimation = MotionEstimation::Chained { + refine_radius: MAX_SEARCH_RADIUS, + }; + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 8, search_radius: 4, pyramid_levels: 2, - estimation: MotionEstimation::Chained { - refine_radius: MAX_SEARCH_RADIUS, - }, + estimation, }; - m.validate().unwrap(); + mode.validate().unwrap(); } #[test] fn motion_estimation_default_is_auto() { - assert_eq!(MotionEstimation::default(), MotionEstimation::Auto); + let estimation = MotionEstimation::default(); + assert_eq!(estimation, MotionEstimation::Auto); } #[test] @@ -722,14 +620,12 @@ mod tests { #[test] fn resolve_auto_at_and_above_threshold_gives_chained_default() { - assert_eq!( - MotionEstimation::Auto.resolve(CHAINED_RADIUS_THRESHOLD), - MotionEstimation::chained_default() - ); - assert_eq!( - MotionEstimation::Auto.resolve(8), - MotionEstimation::chained_default() - ); + let chained = MotionEstimation::chained_default(); + let at_threshold = MotionEstimation::Auto.resolve(CHAINED_RADIUS_THRESHOLD); + let above_threshold = MotionEstimation::Auto.resolve(8); + + assert_eq!(at_threshold, chained); + assert_eq!(above_threshold, chained); } #[test] @@ -749,14 +645,14 @@ mod tests { #[test] fn validate_accepts_auto() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 8, search_radius: 4, pyramid_levels: 2, estimation: MotionEstimation::Auto, }; - m.validate().unwrap(); + mode.validate().unwrap(); } #[test] @@ -766,24 +662,25 @@ mod tests { #[test] fn resolved_estimation_resolves_auto_from_the_mode() { - let m = MotionCompensationMode::Mvtools { + let mode = MotionCompensationMode::Mvtools { blksize: 16, overlap: 8, search_radius: 4, pyramid_levels: 2, estimation: MotionEstimation::Auto, }; - assert_eq!(m.resolved_estimation(1), Some(MotionEstimation::Direct)); - assert_eq!( - m.resolved_estimation(4), - Some(MotionEstimation::chained_default()) - ); + let chained = MotionEstimation::chained_default(); + assert_eq!(mode.resolved_estimation(1), Some(MotionEstimation::Direct)); + assert_eq!(mode.resolved_estimation(4), Some(chained)); } #[test] fn pair_ring_slot_count_is_double_radius() { - assert_eq!(pair_ring_slot_count(3), 6); - assert_eq!(pair_ring_slot_count(1), 2); + let radius_three = pair_ring_slot_count(3); + let radius_one = pair_ring_slot_count(1); + + assert_eq!(radius_three, 6); + assert_eq!(radius_one, 2); } #[test] @@ -795,7 +692,8 @@ mod tests { pyramid_levels: 2, estimation: MotionEstimation::Direct, }; - let ctx = MotionCtx::new(mode, 1920, 1080, StorageAlign::new(32)).unwrap(); + let align = StorageAlign::new(32); + let ctx = MotionCtx::new(mode, 1920, 1080, align).unwrap(); assert_eq!(ctx.step, 8); assert_eq!(ctx.blocks_x, 1920u32.div_ceil(8)); assert_eq!(ctx.blocks_y, 1080u32.div_ceil(8)); diff --git a/av-denoise-core/src/nlmeans/motion/pyramid.rs b/av-denoise-core/src/nlmeans/motion/pyramid.rs index 2ab897b..89b0b76 100644 --- a/av-denoise-core/src/nlmeans/motion/pyramid.rs +++ b/av-denoise-core/src/nlmeans/motion/pyramid.rs @@ -5,42 +5,27 @@ use super::MotionCtx; use crate::nlmeans::align::StorageAlign; use crate::nlmeans::kernels::motion::{nlm_mc_downscale, nlm_mc_extract_luma}; -/// How many luma pixels one frame takes up at `level`, padded up to a -/// whole number of alignment boundaries. +/// One frame's luma pixel count at `level`, padded to whole alignment boundaries. /// -/// Slot offsets are sums of whole slot strides, so padding the stride is -/// what keeps every offset aligned. -/// -/// wgpu rejects a bind-group offset that is not a multiple of its -/// `min_storage_buffer_offset_alignment`. A level whose pixel count does -/// not fill whole boundaries, such as a 180x137 chroma level, would -/// otherwise leave every odd slot short of one. -/// -/// Kernels only ever read a slot's leading `width * height` pixels, so -/// the padding is never touched. +/// Slot offsets are sums of whole strides, so padding the stride keeps every offset aligned. wgpu +/// rejects a bind-group offset that is not a multiple of its `min_storage_buffer_offset_alignment`, +/// which an unpadded level such as 180x137 would break on every odd slot. Kernels only read a +/// slot's leading `width * height` pixels, so the padding is never touched. fn level_slot_pixels(width: u32, height: u32, level: u32, align: StorageAlign) -> usize { - let (w, h) = level_dims(width, height, level); - align.pad_elems::((w as usize) * (h as usize)) + let (level_width, level_height) = level_dims(width, height, level); + let level_pixels = (level_width as usize) * (level_height as usize); + align.pad_elems::(level_pixels) } -/// How many luma pixels each frame takes up across every pyramid level. -/// -/// Level 0 contributes the full pixel count, and each level after that -/// halves both axes. -/// -/// Every level's contribution is padded to the alignment, which matches -/// the layout [`pyramid_slot_byte_offset`] addresses. +/// One frame's luma pixel count across every pyramid level, padded the way +/// [pyramid_slot_byte_offset] addresses it. pub fn pyramid_pixels_per_frame(width: u32, height: u32, levels: u32, align: StorageAlign) -> usize { (0..levels) .map(|level| level_slot_pixels(width, height, level, align)) .sum() } -/// Where a given level and frame slot starts inside the flat pyramid -/// buffer. -/// -/// The result is always a multiple of the alignment. See -/// [`level_slot_pixels`]. +/// Where a level and frame slot starts in the flat pyramid buffer, always on an alignment boundary. pub fn pyramid_slot_byte_offset( width: u32, height: u32, @@ -50,36 +35,37 @@ pub fn pyramid_slot_byte_offset( align: StorageAlign, ) -> u64 { let mut offset_pixels: usize = 0; - for l in 0..level { - offset_pixels += (frame_count as usize) * level_slot_pixels(width, height, l, align); + for lower_level in 0..level { + offset_pixels += (frame_count as usize) * level_slot_pixels(width, height, lower_level, align); } + offset_pixels += (frame as usize) * level_slot_pixels(width, height, level, align); (offset_pixels * size_of::()) as u64 } /// The pixel dimensions at `level`, where level 0 is full resolution. pub fn level_dims(width: u32, height: u32, level: u32) -> (u32, u32) { - let mut w = width; - let mut h = height; + let mut level_width = width; + let mut level_height = height; for _ in 0..level { - w = (w / 2).max(1); - h = (h / 2).max(1); + level_width = (level_width / 2).max(1); + level_height = (level_height / 2).max(1); } - (w, h) + + (level_width, level_height) } -/// Builds every pyramid level for the slot that was just uploaded, -/// starting from the packed full-resolution input. +/// Builds every pyramid level for the slot just uploaded. /// -/// Level 0 is the luma plane on its own. Each level after that is the -/// one before it at half size, averaged 2x2. +/// Level 0 is the luma plane alone, and each further level is the one before averaged 2x2 at half +/// size. #[expect( clippy::too_many_arguments, reason = "the dispatch threads through every buffer and shape the kernel binds" )] pub(crate) fn run_pyramid_build( client: &ComputeClient, - mc: &MotionCtx, + motion_ctx: &MotionCtx, width: u32, height: u32, frame_count: u32, @@ -88,7 +74,6 @@ pub(crate) fn run_pyramid_build( pyramid: &Handle, stored_ch: u32, ) -> Result<(), anyhow::Error> { - let _ = mc; extract_luma::( client, full_res, @@ -98,11 +83,22 @@ pub(crate) fn run_pyramid_build( height, frame_count, stored_ch, - mc.align, + motion_ctx.align, ); - for level in 1..mc.pyramid_levels { - downscale_level::(client, pyramid, slot, width, height, frame_count, level, mc.align); + + for level in 1..motion_ctx.pyramid_levels { + downscale_level::( + client, + pyramid, + slot, + width, + height, + frame_count, + level, + motion_ctx.align, + ); } + Ok(()) } @@ -123,17 +119,13 @@ fn extract_luma( ) { let block_x = 16u32; let block_y = 16u32; - let grid = CubeCount::new_2d(width.div_ceil(block_x), height.div_ceil(block_y)); + let cubes_x = width.div_ceil(block_x); + let cubes_y = height.div_ceil(block_y); + let grid = CubeCount::new_2d(cubes_x, cubes_y); let dim = CubeDim::new_2d(block_x, block_y); let full_len = (frame_count * height * width * stored_ch) as usize; - let level0_dst = pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - 0, - slot, - align, - )); + let level0_offset = pyramid_slot_byte_offset(width, height, frame_count, 0, slot, align); + let level0_dst = pyramid.clone().offset_start(level0_offset); let level0_len = (frame_count * height * width) as usize; unsafe { @@ -170,25 +162,15 @@ fn downscale_level( let (dst_w, dst_h) = level_dims(width, height, level); let block_x = 16u32; let block_y = 16u32; - let grid = CubeCount::new_2d(dst_w.div_ceil(block_x), dst_h.div_ceil(block_y)); + let cubes_x = dst_w.div_ceil(block_x); + let cubes_y = dst_h.div_ceil(block_y); + let grid = CubeCount::new_2d(cubes_x, cubes_y); let dim = CubeDim::new_2d(block_x, block_y); - let src = pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - level - 1, - slot, - align, - )); - let dst = pyramid.clone().offset_start(pyramid_slot_byte_offset( - width, - height, - frame_count, - level, - slot, - align, - )); + let src_offset = pyramid_slot_byte_offset(width, height, frame_count, level - 1, slot, align); + let src = pyramid.clone().offset_start(src_offset); + let dst_offset = pyramid_slot_byte_offset(width, height, frame_count, level, slot, align); + let dst = pyramid.clone().offset_start(dst_offset); let src_len = (src_w * src_h) as usize; let dst_len = (dst_w * dst_h) as usize; @@ -212,6 +194,7 @@ fn downscale_level( #[cfg(test)] mod tests { use super::*; + use crate::nlmeans::motion::MAX_PYRAMID_LEVELS; /// The alignment the Vulkan adapters these tests run on report. fn align() -> StorageAlign { @@ -220,43 +203,45 @@ mod tests { #[test] fn pyramid_pixels_single_level_matches_image() { - assert_eq!(pyramid_pixels_per_frame(64, 32, 1, align()), 64 * 32); + let align = align(); + let pixels = pyramid_pixels_per_frame(64, 32, 1, align); + assert_eq!(pixels, 64 * 32); } #[test] fn pyramid_pixels_two_levels_sums_levels() { - // Level 0 is 64x32, so 2048 pixels. Level 1 is 32x16, so 512. - // That gives 2560 in total. - assert_eq!(pyramid_pixels_per_frame(64, 32, 2, align()), 2048 + 512); + // Level 0 is 64x32, so 2048 pixels, and level 1 is 32x16, so 512. + let align = align(); + let pixels = pyramid_pixels_per_frame(64, 32, 2, align); + assert_eq!(pixels, 2048 + 512); } #[test] fn level_dims_halve() { - assert_eq!(level_dims(64, 32, 0), (64, 32)); - assert_eq!(level_dims(64, 32, 1), (32, 16)); - assert_eq!(level_dims(64, 32, 2), (16, 8)); + let full = level_dims(64, 32, 0); + let half = level_dims(64, 32, 1); + let quarter = level_dims(64, 32, 2); + + assert_eq!(full, (64, 32)); + assert_eq!(half, (32, 16)); + assert_eq!(quarter, (16, 8)); } #[test] fn slot_byte_offsets_respect_every_alignment_a_runtime_can_report() { - // A GPU rejects a bind-group offset that is not a multiple of - // `min_storage_buffer_offset_alignment`, which backends report - // anywhere from 4 to 256 bytes. - // - // Every dimension pair here has at least one level whose - // unpadded slot stride falls short. 360x274 is the chroma plane - // of a 720x548 frame, and its half-size level of 180x137 comes - // to 98,640 bytes, 16 short of a 32-byte boundary. + // Backends report alignments between 4 and 256 bytes. Each size has a level whose unpadded + // stride falls short, such as 180x137, the half-size level of a 720x548 frame's chroma + // plane, at 16 bytes short of a 32-byte boundary. for bytes in [4u64, 16, 32, 64, 256] { let align = StorageAlign::new(bytes); - for (w, h) in [(360, 274), (720, 548), (722, 546), (66, 66), (42, 28)] { - for level in 0..super::super::MAX_PYRAMID_LEVELS { + for (width, height) in [(360, 274), (720, 548), (722, 546), (66, 66), (42, 28)] { + for level in 0..MAX_PYRAMID_LEVELS { for frame in 0..5 { - let offset = pyramid_slot_byte_offset(w, h, 5, level, frame, align); + let offset = pyramid_slot_byte_offset(width, height, 5, level, frame, align); assert_eq!( offset % bytes, 0, - "align {bytes}: {w}x{h} level={level} frame={frame} lands at byte {offset}" + "align {bytes}: {width}x{height} level={level} frame={frame} lands at byte {offset}" ); } } @@ -266,20 +251,17 @@ mod tests { #[test] fn pixels_per_frame_covers_the_last_slot_of_every_level() { - // The allocation `pyramid_pixels_per_frame` sizes has to hold - // every slot `pyramid_slot_byte_offset` addresses, padding - // included, whatever alignment the runtime reports. - let (w, h, frames, levels) = (360u32, 274u32, 5u32, 3u32); + let (width, height, frames, levels) = (360u32, 274u32, 5u32, 3u32); for bytes in [4u64, 16, 32, 64, 256] { let align = StorageAlign::new(bytes); let total_bytes = - pyramid_pixels_per_frame(w, h, levels, align) * frames as usize * size_of::(); + pyramid_pixels_per_frame(width, height, levels, align) * frames as usize * size_of::(); for level in 0..levels { - let (lw, lh) = level_dims(w, h, level); - let last = pyramid_slot_byte_offset(w, h, frames, level, frames - 1, align) as usize; - let end = last + (lw * lh) as usize * size_of::(); + let (level_width, level_height) = level_dims(width, height, level); + let last = pyramid_slot_byte_offset(width, height, frames, level, frames - 1, align) as usize; + let end = last + (level_width * level_height) as usize * size_of::(); assert!( end <= total_bytes, "align {bytes}: level {level} slot {} ends at {end}, past the {total_bytes}-byte buffer", @@ -291,10 +273,9 @@ mod tests { #[test] fn slot_byte_offset_advances_past_full_levels() { - // 4 frames of a 64x32 image across 2 levels. Reaching level 1, - // frame 2 means skipping all of level 0, which is 8192 pixels, - // then two frames of level 1 at 512 pixels each, so 1024 more. - let bytes = pyramid_slot_byte_offset(64, 32, 4, 1, 2, align()); + // Level 1, frame 2 skips all of level 0 at 8192 pixels, then two level 1 frames at 512 each. + let align = align(); + let bytes = pyramid_slot_byte_offset(64, 32, 4, 1, 2, align); assert_eq!(bytes as usize, (8192 + 1024) * 4); } } diff --git a/av-denoise-core/src/nlmeans/noise/correlation.rs b/av-denoise-core/src/nlmeans/noise/correlation.rs new file mode 100644 index 0000000..ce05961 --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/correlation.rs @@ -0,0 +1,80 @@ +/// Correlation-correction points as `(rho, factor)`, sorted by `rho`. +/// +/// Calibrated on synthetic correlated-grain sweeps against the clean bench reference. Each factor +/// is how far the quality peak sits above the true sigma at that correlation, relative to the +/// white-noise optimum at the same sigma. At the heaviest correlation both quality metrics prefer +/// the raised value, and past the last point the table holds flat. +const CORRELATION_FACTOR_TABLE: [(f32, f32); 4] = [(0.0, 1.0), (0.3, 1.05), (0.5, 1.25), (0.65, 1.45)]; + +/// The factor that scales a measured temporal sigma up to allow for grain correlation. +pub(in crate::nlmeans) fn correlation_factor(rho: f32) -> f32 { + interpolate_table(&CORRELATION_FACTOR_TABLE, rho) +} + +/// Linearly interpolates a table of points sorted by `x`, clamping `x` to the table's range. +pub(super) fn interpolate_table(table: &[(f32, f32)], x: f32) -> f32 { + let x = x.clamp(table[0].0, table[table.len() - 1].0); + + for segment in table.windows(2) { + let (start_x, start_y) = segment[0]; + let (end_x, end_y) = segment[1]; + if x <= end_x { + if end_x == start_x { + return end_y; + } + + let fraction = (x - start_x) / (end_x - start_x); + return start_y + fraction * (end_y - start_y); + } + } + + table[table.len() - 1].1 +} + +/// The share of the noise floor that applies at a candidate offset under grain correlation. +/// +/// A nearby candidate shares part of its grain with the centre patch, so only part of the +/// white-noise floor is independent noise, and that part grows with distance. The centre returns 0 +/// because its true distance is zero. With no measured correlation every other offset returns 1, +/// which reproduces the flat white-noise floor. +pub(in crate::nlmeans) fn spatial_offset_factor(dx: i32, dy: i32, rho: f32) -> f32 { + if dx == 0 && dy == 0 { + return 0.0; + } + + if rho <= 0.0 { + return 1.0; + } + + let distance = ((dx * dx + dy * dy) as f32).sqrt(); + let log_rho = rho.ln(); + 1.0 - (distance * log_rho).exp() +} + +pub(in crate::nlmeans) fn spatial_offset_lut_len(search_radius: u32) -> usize { + let side = (2 * search_radius + 1) as usize; + side * side +} + +/// Builds the row-major noise-floor table for a search window. +/// +/// Each entry is the flat noise offset scaled by [spatial_offset_factor] at that candidate. It is +/// cheap enough to rebuild on every submit, at most 289 entries at the largest search radius. +pub(in crate::nlmeans) fn build_spatial_offset_lut( + search_radius: u32, + rho: f32, + noise_offset: f32, +) -> Vec { + let radius = search_radius as i32; + let side = (2 * search_radius + 1) as usize; + let mut lut = vec![0.0f32; side * side]; + + for dy in -radius..=radius { + for dx in -radius..=radius { + let index = ((dy + radius) as usize) * side + (dx + radius) as usize; + lut[index] = noise_offset * spatial_offset_factor(dx, dy, rho); + } + } + + lut +} diff --git a/av-denoise-core/src/nlmeans/noise/curve.rs b/av-denoise-core/src/nlmeans/noise/curve.rs index 437dfff..a15796e 100644 --- a/av-denoise-core/src/nlmeans/noise/curve.rs +++ b/av-denoise-core/src/nlmeans/noise/curve.rs @@ -1,4 +1,5 @@ -use super::{ +use super::stats::{median, sort_ascending}; +use super::temporal::{ AcceptedBlock, QUARTER_FLATNESS, QUARTER_LUMA_MAX, @@ -12,25 +13,30 @@ use super::{ TEMPORAL_QUARTER_FIELDS, TEMPORAL_QUARTER_SIZE, TEMPORAL_QUARTERS, - median, - sort_ascending, temporal_stats_record_len, }; /// How many luma bins the noise curve spans. pub const NOISE_CURVE_BINS: usize = 16; + /// The fewest quarters a bin needs before its median is trusted. const MIN_QUARTERS_PER_BIN: usize = 32; + /// The fewest populated bins a frame needs before it gets a curve. const MIN_POPULATED_BINS: usize = 3; + /// A quarter with a pixel at or below this luma may be clipped, so its noise reads low. pub(super) const CLIP_LOW: f32 = 4.0 / 255.0; + /// A quarter with a pixel at or above this luma may be clipped, so its noise reads low. pub(super) const CLIP_HIGH: f32 = 251.0 / 255.0; + /// How much texture a quarter may carry, as a fraction of the frame's noise variance. const FLAT_FACTOR: f32 = 0.5; + /// The flat gate never tightens below one code of squared gradient. const FLAT_FLOOR: f32 = (1.0 / 255.0) * (1.0 / 255.0); + /// How large a quarter's mean residual can be and still count as static. /// /// A quarter averages 64 pixels rather than a block's 256, so the noise in its mean is twice as @@ -158,9 +164,9 @@ fn static_quarters(records: &[f32], stored_ch: u32, accepted: &[AcceptedBlock]) for quarter_index in 0..TEMPORAL_QUARTERS { let offset_x = (quarter_index % 2) * TEMPORAL_QUARTER_SIZE; let offset_y = (quarter_index / 2) * TEMPORAL_QUARTER_SIZE; - let quarter_w = block.width.saturating_sub(offset_x).min(TEMPORAL_QUARTER_SIZE); - let quarter_h = block.height.saturating_sub(offset_y).min(TEMPORAL_QUARTER_SIZE); - let pixels = (quarter_w * quarter_h) as f32; + let quarter_width = block.width.saturating_sub(offset_x).min(TEMPORAL_QUARTER_SIZE); + let quarter_height = block.height.saturating_sub(offset_y).min(TEMPORAL_QUARTER_SIZE); + let pixels = (quarter_width * quarter_height) as f32; if pixels == 0.0 { continue; } diff --git a/av-denoise-core/src/nlmeans/noise/estimator.rs b/av-denoise-core/src/nlmeans/noise/estimator.rs new file mode 100644 index 0000000..80abfa2 --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/estimator.rs @@ -0,0 +1,45 @@ +/// How much weight the newest frame's estimate carries when smoothing. +pub(in crate::nlmeans) const EMA_ALPHA: f32 = 0.2; +/// The smallest smoothed sigma, 0.1 in 8-bit terms. +/// +/// An estimate near zero would send the derived strength to infinity. +pub(super) const SIGMA_FLOOR: f32 = 0.1 / 255.0; + +/// The smoothed per-channel noise level for one stream. +/// +/// Smoothing stops a single busy frame from spiking the strength, and the floor keeps near-clean +/// content on a usable normalisation factor. +#[derive(Debug, Default)] +pub(in crate::nlmeans) struct NoiseEstimator { + ema: Option>, +} + +impl NoiseEstimator { + /// Folds per-channel sigmas into the running estimate and returns the smoothed result. + /// + /// The first call, and every call with `windowed` set, takes the sample outright so a window's + /// reading never blends with frames outside it. Every element is floored at [SIGMA_FLOOR]. + pub(in crate::nlmeans) fn update(&mut self, sigmas: &[f32], windowed: bool) -> &[f32] { + match &mut self.ema { + Some(ema) if !windowed => { + for (smoothed, &sample) in ema.iter_mut().zip(sigmas.iter()) { + *smoothed = (EMA_ALPHA * sample + (1.0 - EMA_ALPHA) * *smoothed).max(SIGMA_FLOOR); + } + }, + _ => { + let floored: Vec = sigmas.iter().map(|&sigma| sigma.max(SIGMA_FLOOR)).collect(); + self.ema = Some(floored); + }, + } + + self.ema.as_deref().unwrap() + } + + pub(in crate::nlmeans) fn reset(&mut self) { + self.ema = None; + } + + pub(in crate::nlmeans) fn current(&self) -> Option<&[f32]> { + self.ema.as_deref() + } +} diff --git a/av-denoise-core/src/nlmeans/noise/mod.rs b/av-denoise-core/src/nlmeans/noise/mod.rs index 673aa1a..b91e289 100644 --- a/av-denoise-core/src/nlmeans/noise/mod.rs +++ b/av-denoise-core/src/nlmeans/noise/mod.rs @@ -1,2140 +1,65 @@ -//! Measuring how noisy a source is. -//! -//! The HQ variant matches its strength to the noise level, so it needs a -//! number for that level. This module produces one per frame. -//! -//! # Two ways of looking -//! -//! The Immerkær estimate runs a small mask over each frame that cancels -//! smooth content and leaves mostly noise. It is cheap and needs only -//! one frame, but it reads grain that is correlated between neighbouring -//! pixels too low, because such grain looks partly like content to the -//! mask. -//! -//! The temporal estimate compares a frame against the one before it. -//! Where nothing moved, whatever is left over is noise, and correlated -//! grain shows up in full. It needs static content to work, so motion -//! and scene changes make it unreliable. -//! -//! The two are combined by taking whichever reads higher, which lets the -//! temporal estimate correct an Immerkær under-read without letting an -//! unreliable one drag the estimate down. -//! -//! # Two chains -//! -//! The result feeds two separate smoothed estimates. -//! -//! The median chain reads typical noise and drives the filter strength. -//! -//! The low chain reads cautiously, using lower-quartile statistics, and -//! drives the distance floor. Reading that too high scrubs fine texture, -//! so it is deliberately the more conservative of the two. - +mod correlation; mod curve; +mod estimator; +mod spatial; +mod stats; mod strength_map; +mod temporal; +#[cfg(test)] +mod tests; -use cubecl::prelude::*; -use cubecl::server::Handle; - +pub(super) use self::correlation::{ + build_spatial_offset_lut, + correlation_factor, + spatial_offset_factor, + spatial_offset_lut_len, +}; pub use self::curve::NOISE_CURVE_BINS; pub(crate) use self::curve::NoiseCurve; +#[cfg(test)] pub(super) use self::curve::build_noise_curve; +pub(super) use self::estimator::{EMA_ALPHA, NoiseEstimator}; +pub(super) use self::spatial::{ + NoiseCtx, + noise_partials_slot_stride_bytes, + partials_len, + run_noise_estimate, + sigma_block_p25_from_partials, + sigma_from_abs_sum, +}; +#[cfg(test)] pub(super) use self::strength_map::classify_quarters; #[cfg(test)] pub(crate) use self::strength_map::{QuarterClass, QuarterTensor}; pub(crate) use self::strength_map::{QuarterClasses, StrengthMapParams}; -use super::align::StorageAlign; -use super::kernels::{ - nlm_noise_partial, - nlm_noise_reduce, - nlm_temporal_noise_stats, - nlm_temporal_stats_zero, +pub(super) use self::temporal::{ + QUARTER_FLATNESS, + QUARTER_LUMA_MAX, + QUARTER_LUMA_MIN, + QUARTER_LUMA_SUM, + QUARTER_SUM_D, + QUARTER_SUM_D2, + QUARTER_TENSOR_XX, + QUARTER_TENSOR_XY, + QUARTER_TENSOR_YY, + TEMPORAL_QUARTER_BASE, + TEMPORAL_QUARTER_FIELDS, + TEMPORAL_QUARTERS, + TemporalNoiseReading, + TemporalNoiseSample, + TemporalStatsCtx, + read_temporal_stats_slot, + run_temporal_noise_stats, + temporal_noise_reading, + temporal_stats_blocks, + temporal_stats_buf_bytes, + temporal_stats_record_len, + zero_temporal_stats_slot, }; -use super::{BLOCK_1D, BLOCK_X, BLOCK_Y, MAX_GRID_1D}; - -/// The inputs one Immerkær noise estimate needs. -/// -/// This lives only for the length of a single estimate call, which is -/// what makes the borrows on the denoiser's buffers sound. -pub(super) struct NoiseCtx<'a> { - pub width: u32, - pub height: u32, - pub channels: u32, - pub stored_ch: u32, - pub frame_count: u32, - pub frame: u32, - pub slot: u32, - pub input_buf: &'a Handle, - pub partials_buf: &'a Handle, - pub results_buf: &'a Handle, -} - -/// How many `f32` elements the first stage's partials buffer needs for -/// one frame. -/// -/// Each block covers a tile of the frame and contributes four lanes. -pub(super) fn partials_len(width: u32, height: u32) -> usize { - (width.div_ceil(BLOCK_X) * height.div_ceil(BLOCK_Y) * 4) as usize -} - -/// The byte stride between ring slots in the partials buffer, padded up -/// to the runtime's buffer-binding alignment. -/// -/// A small frame can leave [`partials_len`] short of a boundary, and -/// wgpu rejects a bind-group offset that is not a multiple of its -/// `min_storage_buffer_offset_alignment`. -/// -/// This matches [`temporal_stats_slot_stride_bytes`]. -pub(super) fn noise_partials_slot_stride_bytes(width: u32, height: u32, align: StorageAlign) -> u64 { - align.pad_bytes(partials_len(width, height) as u64 * size_of::() as u64) -} - -/// Runs both stages of the Immerkær noise estimate for one frame. -/// -/// The per-channel totals go into the results buffer at this frame's -/// slot. -/// -/// That buffer holds four values per ring slot, matching the input -/// ring's frame capacity. -pub(super) fn run_noise_estimate( - client: &ComputeClient, - ctx: &NoiseCtx<'_>, -) -> Result<(), anyhow::Error> { - let total_input = (ctx.frame_count * ctx.height * ctx.width * ctx.stored_ch) as usize; - let n_partials = partials_len(ctx.width, ctx.height); - let total_results = (ctx.frame_count * 4) as usize; - let stored_ch = ctx.stored_ch as usize; - - unsafe { - nlm_noise_partial::launch_unchecked::( - client, - CubeCount::new_2d(ctx.width.div_ceil(BLOCK_X), ctx.height.div_ceil(BLOCK_Y)), - CubeDim::new_2d(BLOCK_X, BLOCK_Y), - stored_ch, - ArrayArg::from_raw_parts(ctx.input_buf.clone(), total_input), - ArrayArg::from_raw_parts(ctx.partials_buf.clone(), n_partials), - ctx.frame, - ctx.width, - ctx.height, - ctx.channels, - BLOCK_X, - BLOCK_Y, - ); - } - - let num_partials = (n_partials / 4) as u32; - unsafe { - nlm_noise_reduce::launch_unchecked::( - client, - CubeCount::new_1d(1), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(ctx.partials_buf.clone(), n_partials), - ArrayArg::from_raw_parts(ctx.results_buf.clone(), total_results), - ctx.slot, - num_partials, - BLOCK_1D, - ); - } - - Ok(()) -} - -/// Turns the summed absolute mask responses into an Immerkær sigma. -/// -/// The interior area leaves out the one-pixel border the mask cannot -/// reach. -pub(super) fn sigma_from_abs_sum(abs_sum: f32, width: u32, height: u32) -> f32 { - let interior = ((width - 2) as f32) * ((height - 2) as f32); - (std::f32::consts::FRAC_PI_2).sqrt() * abs_sum / (6.0 * interior) -} - -/// The per-channel lower quartile of the per-block Immerkær sigmas. -/// -/// It reads one slot's partials directly, in the layout the first stage -/// wrote them. -/// -/// Each block's own sigma comes from the same formula -/// [`sigma_from_abs_sum`] uses, applied to however much of that block's -/// tile overlaps the frame's interior. -/// -/// A block with no interior overlap is skipped rather than diluting the -/// quartile with a spurious zero. -/// -/// Wherever noise is uneven across a frame, this quartile reads lower -/// than the frame-wide mean, which is the cautious estimate the low -/// chain wants. -/// -/// Channels past the active count stay at 0. -pub(super) fn sigma_block_p25_from_partials( - partials: &[f32], - channels: u32, - width: u32, - height: u32, -) -> [f32; 3] { - let cubes_x = width.div_ceil(BLOCK_X); - let cubes_y = height.div_ceil(BLOCK_Y); - let channels = channels as usize; - - let mut cube_sigmas: Vec> = vec![Vec::new(); channels]; - - for cy in 0..cubes_y { - let tile_y0 = cy * BLOCK_Y; - let tile_y1 = ((cy + 1) * BLOCK_Y).min(height); - let overlap_y0 = tile_y0.max(1); - let overlap_y1 = tile_y1.min(height - 1); - - for cx in 0..cubes_x { - let tile_x0 = cx * BLOCK_X; - let tile_x1 = ((cx + 1) * BLOCK_X).min(width); - let overlap_x0 = tile_x0.max(1); - let overlap_x1 = tile_x1.min(width - 1); - - if overlap_x1 <= overlap_x0 || overlap_y1 <= overlap_y0 { - continue; - } - - let area = ((overlap_x1 - overlap_x0) * (overlap_y1 - overlap_y0)) as f32; - let cube_index = (cy * cubes_x + cx) as usize; - let base = cube_index * 4; - - for (c, sigmas) in cube_sigmas.iter_mut().enumerate() { - let sum = partials[base + c]; - sigmas.push(std::f32::consts::FRAC_PI_2.sqrt() * sum / (6.0 * area)); - } - } - } - - let mut sigma_low = [0.0f32; 3]; - for (c, sigmas) in cube_sigmas.iter_mut().enumerate() { - if sigmas.is_empty() { - continue; - } - sort_ascending(sigmas); - sigma_low[c] = lower_quartile(sigmas); - } - sigma_low -} - -/// The spatial block size the temporal residual statistics use, with one -/// GPU block per square of this size. -pub(super) const TEMPORAL_NOISE_BLOCK: u32 = 16; - -/// How many `f32`s one block's stats record holds. -/// -/// That is a sum and a sum of squares per stored channel, one lag-1 -/// total, and one record for each of the block's four 8x8 quarters. -pub(super) fn temporal_stats_record_len(stored_ch: u32) -> u32 { - 2 * stored_ch + TEMPORAL_QUARTER_BASE + TEMPORAL_QUARTERS * TEMPORAL_QUARTER_FIELDS -} - -/// How many 8x8 quarters one block record holds. -pub(super) const TEMPORAL_QUARTERS: u32 = 4; -/// The side length of one quarter, in pixels. -pub(super) const TEMPORAL_QUARTER_SIZE: u32 = TEMPORAL_NOISE_BLOCK / 2; -/// Offset, past `2 * stored_ch`, of the first quarter record. -/// -/// Quarter `q` starts at `2 * stored_ch + 1 + 9 * q`, in top-left, -/// top-right, bottom-left, bottom-right order. -pub(super) const TEMPORAL_QUARTER_BASE: u32 = 1; -/// How many `f32`s one quarter record holds. -pub(super) const TEMPORAL_QUARTER_FIELDS: u32 = 9; -/// Offset, within a quarter record, of channel 0's summed residual. -pub(super) const QUARTER_SUM_D: u32 = 0; -/// Offset, within a quarter record, of channel 0's summed squared residual. -pub(super) const QUARTER_SUM_D2: u32 = 1; -/// Offset, within a quarter record, of the new frame's summed luma. -pub(super) const QUARTER_LUMA_SUM: u32 = 2; -/// Offset, within a quarter record, of the smoothed temporal mean's -/// gradient energy, or `3.0e38` on a quarter smaller than 8x8. -pub(super) const QUARTER_FLATNESS: u32 = 3; -/// Offset, within a quarter record, of the new frame's minimum luma. -pub(super) const QUARTER_LUMA_MIN: u32 = 4; -/// Offset, within a quarter record, of the new frame's maximum luma. -pub(super) const QUARTER_LUMA_MAX: u32 = 5; -/// Offset, within a quarter record, of the temporal mean's summed squared -/// horizontal gradient. -pub(super) const QUARTER_TENSOR_XX: u32 = 6; -/// Offset, within a quarter record, of the temporal mean's summed squared -/// vertical gradient. -pub(super) const QUARTER_TENSOR_YY: u32 = 7; -/// Offset, within a quarter record, of the temporal mean's summed -/// horizontal times vertical gradient. -pub(super) const QUARTER_TENSOR_XY: u32 = 8; - -/// The block grid covering a frame, laid out row-major. -/// -/// Ragged edges are truncated rather than padded, the same way the block -/// matcher handles its own ragged last block. -pub(super) fn temporal_stats_blocks(width: u32, height: u32) -> (u32, u32) { - ( - width.div_ceil(TEMPORAL_NOISE_BLOCK), - height.div_ceil(TEMPORAL_NOISE_BLOCK), - ) -} - -/// Number of `f32`s in one ring slot's stats region. -pub(super) fn temporal_stats_slot_len(width: u32, height: u32, stored_ch: u32) -> usize { - let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); - (blocks_x * blocks_y * temporal_stats_record_len(stored_ch)) as usize -} - -/// The byte stride between ring slots in the temporal-stats buffer, -/// padded up to the runtime's buffer-binding alignment. -/// -/// A small frame, or a single-channel mode, can leave -/// [`temporal_stats_slot_len`] short of a boundary, and wgpu rejects a -/// bind-group offset that is not a multiple of its -/// `min_storage_buffer_offset_alignment`. -/// -/// This matches `MotionCtx::confidence_bytes_per_neighbour`. -pub(super) fn temporal_stats_slot_stride_bytes( - width: u32, - height: u32, - stored_ch: u32, - align: StorageAlign, -) -> u64 { - align.pad_bytes(temporal_stats_slot_len(width, height, stored_ch) as u64 * size_of::() as u64) -} - -/// Total byte size of a `frame_count`-slot temporal-stats ring. -pub(super) fn temporal_stats_buf_bytes( - width: u32, - height: u32, - stored_ch: u32, - frame_count: u32, - align: StorageAlign, -) -> usize { - (temporal_stats_slot_stride_bytes(width, height, stored_ch, align) * frame_count as u64) as usize -} - -/// The inputs one temporal residual statistics dispatch needs, comparing -/// the new slot against the previous one on the input ring. -pub(super) struct TemporalStatsCtx<'a> { - pub width: u32, - pub height: u32, - pub stored_ch: u32, - pub frame_count: u32, - pub slot_new: u32, - pub slot_prev: u32, - pub input_buf: &'a Handle, - pub stats_buf: &'a Handle, - pub align: StorageAlign, -} - -/// Runs the temporal residual statistics kernel for the new slot, -/// writing one record per block into that slot's padded region. -/// -/// The kernel only ever addresses within its own slice, so it needs to -/// know nothing about the ring's other slots or the padding between -/// them. -/// -/// `luma_fields` gates the kernel's four luma-only lanes. With it off, -/// those lanes read 0 and the dispatch costs the same as before they -/// existed. See [nlm_temporal_noise_stats]. -pub(super) fn run_temporal_noise_stats( - client: &ComputeClient, - ctx: &TemporalStatsCtx<'_>, - luma_fields: bool, -) -> Result<(), anyhow::Error> { - let total_input = (ctx.frame_count * ctx.height * ctx.width * ctx.stored_ch) as usize; - let (blocks_x, blocks_y) = temporal_stats_blocks(ctx.width, ctx.height); - let slot_len = temporal_stats_slot_len(ctx.width, ctx.height, ctx.stored_ch); - let stride = temporal_stats_slot_stride_bytes(ctx.width, ctx.height, ctx.stored_ch, ctx.align); - let stats_slot = ctx.stats_buf.clone().offset_start((ctx.slot_new as u64) * stride); - - unsafe { - nlm_temporal_noise_stats::launch_unchecked::( - client, - CubeCount::new_2d(blocks_x, blocks_y), - CubeDim::new_2d(TEMPORAL_NOISE_BLOCK, TEMPORAL_NOISE_BLOCK), - ctx.stored_ch as usize, - ArrayArg::from_raw_parts(ctx.input_buf.clone(), total_input), - ArrayArg::from_raw_parts(stats_slot, slot_len), - ctx.slot_new, - ctx.slot_prev, - ctx.width, - ctx.height, - ctx.stored_ch, - TEMPORAL_NOISE_BLOCK, - luma_fields, - ); - } - - Ok(()) -} - -/// Fills one ring slot's temporal-stats region with zeroes. -/// -/// This runs when a slot is a copy of the one before it, which happens -/// while priming a stream and during the end-of-stream flush. -/// -/// The zeroes read as no static blocks with measurable noise, rather -/// than as a made-up reading of zero noise. -pub(super) fn zero_temporal_stats_slot( - client: &ComputeClient, - stats_buf: &Handle, - width: u32, - height: u32, - stored_ch: u32, - slot: u32, - align: StorageAlign, -) { - let slot_len = temporal_stats_slot_len(width, height, stored_ch) as u32; - let stride = temporal_stats_slot_stride_bytes(width, height, stored_ch, align); - let dst = stats_buf.clone().offset_start((slot as u64) * stride); - - let grid = slot_len.div_ceil(BLOCK_1D).min(MAX_GRID_1D); - let total_threads = grid * BLOCK_1D; - - unsafe { - nlm_temporal_stats_zero::launch_unchecked::( - client, - CubeCount::new_1d(grid), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(dst, slot_len as usize), - slot_len, - total_threads, - ); - } -} - -/// Reads exactly one ring slot's temporal-stats region back as owned -/// values. -/// -/// The shared ring handle is sliced by byte offset, so the transfer only -/// covers one slot rather than the whole ring. -#[expect( - clippy::too_many_arguments, - reason = "the dispatch threads through every buffer and shape the kernel binds" -)] -pub(super) fn read_temporal_stats_slot( - client: &ComputeClient, - stats_buf: &Handle, - width: u32, - height: u32, - stored_ch: u32, - frame_count: u32, - slot: u32, - align: StorageAlign, -) -> Result, anyhow::Error> { - let slot_len_bytes = temporal_stats_slot_len(width, height, stored_ch) as u64 * size_of::() as u64; - let stride = temporal_stats_slot_stride_bytes(width, height, stored_ch, align); - let total_bytes = frame_count as u64 * stride; - let start = (slot as u64) * stride; - let end_trim = total_bytes - start - slot_len_bytes; - - let sliced = stats_buf.clone().offset_start(start).offset_end(end_trim); - let bytes = client - .read_one(sliced) - .map_err(|e| anyhow::anyhow!("temporal noise stats readback failed: {e}"))?; - Ok(f32::from_bytes(&bytes).to_vec()) -} - -/// One centre slot's aggregated temporal-residual noise measurement. -#[derive(Debug, Clone, Copy, PartialEq)] -pub(super) struct TemporalNoiseSample { - /// The per-channel sigma, taken as the median over the static blocks - /// with measurable noise, in normalised units. Entries past the active - /// channel count stay 0. - pub sigma: [f32; 3], - /// The same per-channel sigma taken at the lower quartile instead of - /// the median, in normalised units. - /// - /// This reads more cautiously than `sigma`, for the consumers where - /// reading too high does more harm than reading too low. Entries - /// past the active channel count stay 0. - pub sigma_low: [f32; 3], - /// How correlated the grain is between neighbouring pixels. - /// - /// It is the median, over the static blocks with measurable noise, - /// of how strongly each residual matches the one beside it. - pub rho: f32, - /// What fraction of blocks counted as static with measurable noise. - pub static_fraction: f32, -} - -/// How large a block's mean residual can be and still count as static, -/// in normalised units. -/// -/// A block above this is treated as moving content rather than as noise. -const STATIC_GATE: f32 = 1.5 / 255.0; -/// The smallest block sigma that still counts as measurable noise. -/// -/// It gates the correlation median, the ceiling's reference set and the -/// measured sigma population behind `sigma` and `sigma_low`. -/// -/// Below this a block carries too little signal for its correlation -/// reading to mean anything. -const RHO_SIGMA_GATE: f32 = 0.3 / 255.0; -/// The smallest fraction of static blocks a sample needs to be trusted. -/// -/// Below this, motion, a scene change, or too little measurable noise -/// dominates the frame, and the Immerkær estimate is the only usable -/// reading. -const STATIC_FRACTION_MIN: f32 = 0.05; -/// How far above the surviving blocks' own lower quartile a block's -/// sigma may sit before it is treated as moving texture rather than -/// noise. -/// -/// # Why this is needed -/// -/// A block panning across texture can average out to nearly nothing -/// over its window, which clears [`STATIC_GATE`], while its variance is -/// entirely shifted texture rather than repeatable noise. That variance -/// runs several times higher than the rest of the frame's blocks. -/// -/// Real noise keeps a much narrower spread across a frame, even where -/// its magnitude genuinely varies, such as a dark region reading noisier -/// than a bright one. This threshold leaves room for that spread while -/// still catching the texture outliers. -/// -/// # Why the lower quartile -/// -/// The reference is the lower quartile rather than the median, so the -/// filter still works once texture makes up most of the surviving -/// blocks, as long as a genuinely static minority remains to anchor it. -/// -/// The quartile is computed only over blocks whose own sigma clears -/// [`RHO_SIGMA_GATE`]. Letterbox bars and other perfectly static regions -/// read a sigma of exactly 0, and leaving them in can drag the quartile -/// itself to 0, which would reject every block carrying real noise -/// rather than just the outliers. -/// -/// The measured population sits behind the same gate, so those regions -/// cannot drag the reported sigma down either. -/// -/// That exclusion only holds up while a genuinely low population remains -/// to anchor the quartile. [`aggregate_temporal_noise_stats`] covers -/// what happens when none does. -const SIGMA_OUTLIER_FACTOR: f32 = 5.0; - -/// One surviving block's per-channel stats, kept just long enough to -/// work out the reference the outlier check needs before deciding which -/// blocks are really static. -struct StaticGateCandidate { - index: usize, - width: u32, - height: u32, - sigmas: [f32; 3], - sigma_ch0: f32, - var_ch0: f32, - mean0: f32, - mean_lag: f32, - n_pairs: f32, -} - -/// One block [accepted_static_blocks] counted as static noise. -/// -/// It carries everything both [temporal_noise_reading] and -/// [build_noise_curve] need. `index` places it back in `records`, and -/// `width` and `height` give the extent of each of its quarters. -pub(super) struct AcceptedBlock { - /// This block's position in `records`, in units of one record. - pub(super) index: usize, - /// The block's in-frame width, in pixels. - pub(super) width: u32, - /// The block's in-frame height, in pixels. - pub(super) height: u32, - /// The per-channel sigma, as [temporal_noise_reading] - /// computes it. Entries past the active channel count stay 0. - pub(super) sigmas: [f32; 3], - /// Channel 0's mean residual over the block. - pub(super) mean0: f32, - /// Channel 0's mean lag-1 product over the block's adjacent pixel - /// pairs. - pub(super) mean_lag: f32, - /// Channel 0's residual variance over the block. - pub(super) var_ch0: f32, - /// How many adjacent pixel pairs the block holds. - pub(super) n_pairs: f32, -} - -/// Picks out the blocks a centre slot's records count as static noise. -/// -/// `records` holds exactly one slot's region, laid out block by block as -/// [`nlm_temporal_noise_stats`] documents. -/// -/// [temporal_noise_reading] builds both the scalar sample and the curve -/// from this same selection, so they cannot drift apart. -/// -/// # When this returns `None` -/// -/// The outlier check below had no way to validate its own ceiling. See -/// "When the ceiling cannot be trusted". -/// -/// # Deciding which blocks are static -/// -/// This takes two passes. -/// -/// The first checks each block's mean residual against [`STATIC_GATE`]. -/// A block whose average residual is near zero passes, but a block -/// panning across texture passes too, because displaced texture averages -/// toward zero over a block just as noise does. -/// -/// The second pass catches what the first let through. It rejects any -/// surviving block whose sigma sits far above the surviving population's -/// own lower quartile, by more than [`SIGMA_OUTLIER_FACTOR`]. A panning -/// block's variance comes from the texture it slid across, not from -/// noise shared with the rest of the frame. -/// -/// That quartile looks only at blocks clearing [`RHO_SIGMA_GATE`], so a -/// perfectly static region such as a letterbox bar cannot drag the -/// ceiling to 0 and reject every noisy block with it. -/// -/// A third filter drops any surviving block whose own sigma sits below -/// [`RHO_SIGMA_GATE`]. Such a block carries no measurable noise, so it -/// joins neither the reported sigma nor the static count. -/// -/// # When the ceiling cannot be trusted -/// -/// Excluding those low-sigma blocks has a cost of its own. If every -/// remaining block turns out to be texture, that population sets its own -/// ceiling and lets all of its members through, because a value never -/// exceeds a multiple of itself. -/// -/// Nothing in a single block's stats distinguishes texture from noise, -/// so the only check left is whether the surviving population shows any -/// internal spread. -/// -/// Real per-block sigma, measured over a finite sample, always varies a -/// little from block to block, even for physically uniform noise. So -/// once the low-sigma blocks are gone, a perfectly uniform remainder is -/// the signature of a self-selected outlier population with nothing to -/// anchor it. -/// -/// In that case this reports `None` rather than a confident sigma that -/// may be inflated by texture. -pub(super) fn accepted_static_blocks( - records: &[f32], - channels: u32, - stored_ch: u32, - width: u32, - height: u32, -) -> Option> { - let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); - let total_blocks = (blocks_x * blocks_y) as usize; - if total_blocks == 0 { - return None; - } - - let record_len = temporal_stats_record_len(stored_ch) as usize; - let channels = channels as usize; - let stored_ch = stored_ch as usize; - - let mut candidates = Vec::new(); - - for by in 0..blocks_y { - for bx in 0..blocks_x { - let block_index = (by * blocks_x + bx) as usize; - let rec = &records[block_index * record_len..(block_index + 1) * record_len]; - - let block_origin_x = bx * TEMPORAL_NOISE_BLOCK; - let block_origin_y = by * TEMPORAL_NOISE_BLOCK; - let block_w = TEMPORAL_NOISE_BLOCK.min(width - block_origin_x); - let block_h = TEMPORAL_NOISE_BLOCK.min(height - block_origin_y); - let n = (block_w * block_h) as f32; - let n_pairs = (block_h * block_w.saturating_sub(1)) as f32; - - let mean0 = rec[0] / n; - if mean0.abs() >= STATIC_GATE { - continue; - } - - let mut sigmas = [0.0f32; 3]; - let mut sigma_ch0 = 0.0f32; - let mut var_ch0 = 0.0f32; - for c in 0..channels { - let mean = rec[c] / n; - let var = (rec[stored_ch + c] / n - mean * mean).max(0.0); - let sigma_block = var.sqrt() / std::f32::consts::SQRT_2; - sigmas[c] = sigma_block; - if c == 0 { - sigma_ch0 = sigma_block; - var_ch0 = var; - } - } - - let mean_lag = if n_pairs > 0.0 { - rec[2 * stored_ch] / n_pairs - } else { - 0.0 - }; - - candidates.push(StaticGateCandidate { - index: block_index, - width: block_w, - height: block_h, - sigmas, - sigma_ch0, - var_ch0, - mean0, - mean_lag, - n_pairs, - }); - } - } - - // A perfectly static block, such as a letterbox bar or a duplicate - // frame, reads a sigma of 0 and clears STATIC_GATE without effort. - // Blocks like that can dominate the surviving population while - // carrying no measurable noise at all. - // - // Left in this reference set they drag the quartile toward 0, which - // then rejects every block with real noise instead of just the - // texture outliers the check exists for. - // - // Limiting the reference to blocks that themselves clear - // RHO_SIGMA_GATE keeps the ceiling anchored to blocks that could - // plausibly be noise. - let mut reference_sigma_ch0: Vec = candidates - .iter() - .map(|c| c.sigma_ch0) - .filter(|&sigma| sigma > RHO_SIGMA_GATE) - .collect(); - sort_ascending(&mut reference_sigma_ch0); - - // Dropping the low-sigma blocks throws away what they told us, and - // that has a cost. A value never exceeds a multiple of itself, so - // if every remaining block turns out to be a texture-panning block, - // that population sets its own ceiling and lets all of them - // through, which is exactly the failure the ceiling exists to - // prevent. - // - // Nothing available here tells texture and noise apart on its own, - // so this asks for corroboration instead. Some blocks must have - // been excluded, and the reference population left behind must show - // real internal spread, meaning its lowest and highest readings - // differ. - // - // Per-block sigma is measured over a finite sample, so even - // physically uniform noise varies a little from block to block. A - // reference population with no spread at all, sitting next to - // excluded low-sigma blocks, is the case this cannot resolve, so it - // reports None rather than a sigma that may be inflated by texture. - // - // A population where nothing was excluded, because there were no - // low-sigma blocks to begin with, skips this check and is trusted - // directly. - let candidates_were_excluded = candidates.len() > reference_sigma_ch0.len(); - let reference_has_spread = match (reference_sigma_ch0.first(), reference_sigma_ch0.last()) { - (Some(&lo), Some(&hi)) => hi > lo, - (None, _) | (_, None) => false, - }; - if candidates_were_excluded && !reference_has_spread { - return None; - } - - let sigma_ceiling = if reference_sigma_ch0.is_empty() { - 0.0 - } else { - lower_quartile(&reference_sigma_ch0) * SIGMA_OUTLIER_FACTOR - }; - - let mut accepted = Vec::new(); - - for candidate in candidates { - if candidate.sigma_ch0 > sigma_ceiling { - continue; - } - - // A perfectly static block, such as a letterbox bar or a - // duplicate frame, reads a sigma of 0 and clears STATIC_GATE - // without effort. It carries no measurable noise, so it says - // nothing about the frame's noise level. - // - // Left in, a population of them drags the lower quartile to 0, - // and nl4d reads that quartile as its whole noise estimate. - if candidate.sigma_ch0 <= RHO_SIGMA_GATE { - continue; - } - - accepted.push(AcceptedBlock { - index: candidate.index, - width: candidate.width, - height: candidate.height, - sigmas: candidate.sigmas, - mean0: candidate.mean0, - mean_lag: candidate.mean_lag, - var_ch0: candidate.var_ch0, - n_pairs: candidate.n_pairs, - }); - } - - Some(accepted) -} - -/// Combines one centre slot's per-block records into a single -/// [TemporalNoiseSample]. -/// -/// This is a thin wrapper over [temporal_noise_reading], kept for the -/// callers that only ever want the scalar sample. -/// -/// # When no sample is produced -/// -/// This returns `None` in three cases. -/// -/// [accepted_static_blocks] found no trustworthy way to set its own -/// outlier ceiling. -/// -/// Too few blocks counted as static, below [STATIC_FRACTION_MIN], so -/// motion dominates the frame and nothing here can be trusted. -/// -/// No static block carried measurable noise, so the correlation median -/// would be undefined. A zero-filled duplicate slot produces exactly -/// this. -/// -/// No production code calls this any more. `NlmDenoiser::read_temporal_noise` -/// calls [temporal_noise_reading] directly instead, and this stays -/// available only for the tests that pin it against the shared -/// selection. -#[cfg(test)] -pub(super) fn aggregate_temporal_noise_stats( - records: &[f32], - channels: u32, - stored_ch: u32, - width: u32, - height: u32, -) -> Option { - temporal_noise_reading(records, channels, stored_ch, width, height, false, None).sample -} - -/// One centre slot's temporal-noise reading, combining the scalar sample -/// with the curve built from the same accepted blocks. -pub(super) struct TemporalNoiseReading { - /// The scalar sample, as [temporal_noise_reading] computes it. - pub(super) sample: Option, - /// The per-frame luma noise curve, present only when `with_curve` was - /// set and a sample was produced. - pub(super) curve: Option, - /// The frame's quarter classes, built against `curve` and present exactly when it is. - pub(super) classes: Option, -} - -/// Builds one centre slot's [TemporalNoiseReading]. -/// -/// This calls [accepted_static_blocks] once and shares its result -/// between the scalar sample and the curve, so the two can never -/// disagree about which blocks count as static noise. -/// -/// `with_curve` gates the curve. It stays `None` whenever `with_curve` -/// is off or no sample was produced, since a frame too unreliable for a -/// scalar sample is too unreliable for a curve. -/// -/// `texture_cut` is the cut [classify_quarters] vetoes textured flat quarters at, or `None` for no -/// veto. -pub(super) fn temporal_noise_reading( - records: &[f32], - channels: u32, - stored_ch: u32, - width: u32, - height: u32, - with_curve: bool, - texture_cut: Option, -) -> TemporalNoiseReading { - let none = TemporalNoiseReading { - sample: None, - curve: None, - classes: None, - }; - - let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); - let total_blocks = (blocks_x * blocks_y) as usize; - - let Some(accepted) = accepted_static_blocks(records, channels, stored_ch, width, height) else { - return none; - }; - - let channels = channels as usize; - - let mut rho_samples = Vec::new(); - for block in &accepted { - if block.n_pairs > 0.0 { - let rho = (block.mean_lag - block.mean0 * block.mean0) / block.var_ch0; - rho_samples.push(rho.clamp(0.0, 1.0)); - } - } - - let static_fraction = accepted.len() as f32 / total_blocks as f32; - if static_fraction < STATIC_FRACTION_MIN || rho_samples.is_empty() { - return none; - } - - let mut static_sigmas: Vec> = vec![Vec::new(); channels]; - for block in &accepted { - for (c, sigmas) in static_sigmas.iter_mut().enumerate().take(channels) { - sigmas.push(block.sigmas[c]); - } - } - - let mut sigma = [0.0f32; 3]; - let mut sigma_low = [0.0f32; 3]; - for (c, sigmas) in static_sigmas.iter_mut().enumerate() { - sort_ascending(sigmas); - sigma[c] = median(sigmas); - sigma_low[c] = lower_quartile(sigmas); - } - sort_ascending(&mut rho_samples); - let rho = median(&rho_samples); - - let sample = TemporalNoiseSample { - sigma, - sigma_low, - rho, - static_fraction, - }; - - let curve = if with_curve { - build_noise_curve(records, stored_ch, &accepted, sample.sigma[0]) - } else { - None - }; - - let classes = curve - .as_ref() - .map(|curve| classify_quarters(records, stored_ch, width, height, curve, texture_cut)); - - TemporalNoiseReading { - sample: Some(sample), - curve, - classes, - } -} - -/// Sorts `values` into ascending order in place. -/// -/// Both [`median`] and [`lower_quartile`] need this first, so a caller -/// wanting both sorts once and passes the same slice to each. -pub(super) fn sort_ascending(values: &mut [f32]) { - values.sort_by(|a, b| a.partial_cmp(b).expect("noise stats are never NaN")); -} - -/// The median of an already-sorted slice, averaging the two middle -/// elements when the count is even. -/// -/// Callers only ever pass a non-empty slice. -pub(super) fn median(values: &[f32]) -> f32 { - let n = values.len(); - if n % 2 == 1 { - values[n / 2] - } else { - (values[n / 2 - 1] + values[n / 2]) / 2.0 - } -} - -/// The lower quartile of an already-sorted slice. -/// -/// It reads the value a quarter of the way along, interpolating between -/// the two neighbouring elements when that lands between them. -/// -/// A slice of one returns that element. Callers only ever pass a -/// non-empty slice. -fn lower_quartile(values: &[f32]) -> f32 { - let n = values.len(); - if n == 1 { - return values[0]; - } - let idx = 0.25 * (n - 1) as f32; - let lo = idx.floor() as usize; - let hi = idx.ceil() as usize; - if lo == hi { - values[lo] - } else { - let t = idx - lo as f32; - values[lo] + t * (values[hi] - values[lo]) - } -} - -/// The correlation-correction table, as points sorted by correlation. -/// -/// It was measured on synthetic correlated-grain sweeps against the -/// clean bench reference. -/// -/// Each factor is how far the quality peak sits above the true sigma at -/// that level of grain correlation, relative to the white-noise optimum -/// at the same sigma. -/// -/// At the heaviest correlation measured, both quality metrics prefer the -/// raised value. Past the last measured point the table holds flat. -const CORRELATION_FACTOR_TABLE: [(f32, f32); 4] = [(0.0, 1.0), (0.3, 1.05), (0.5, 1.25), (0.65, 1.45)]; - -/// The factor that turns a measured temporal sigma into an effective one -/// allowing for grain correlation. -/// -/// The effective sigma is the measured one multiplied by this. -pub(super) fn correlation_factor(rho: f32) -> f32 { - interpolate_table(&CORRELATION_FACTOR_TABLE, rho) -} - -/// Reads a value from a small table of points sorted by `x`, -/// interpolating between the two nearest entries. -/// -/// `x` is clamped to the table's own range first, so the result never -/// runs past the endpoints. -fn interpolate_table(table: &[(f32, f32)], x: f32) -> f32 { - let x = x.clamp(table[0].0, table[table.len() - 1].0); - - for pair in table.windows(2) { - let (x0, y0) = pair[0]; - let (x1, y1) = pair[1]; - if x <= x1 { - if x1 == x0 { - return y1; - } - let t = (x - x0) / (x1 - x0); - return y0 + t * (y1 - y0); - } - } - - table[table.len() - 1].1 -} - -/// How much of the noise floor still applies at a given candidate -/// offset, once grain correlation is taken into account. -/// -/// A nearby candidate shares some of its grain with the centre patch, so -/// only part of the white-noise floor is genuinely independent noise. -/// The further away the candidate, the less they share, and the more of -/// the floor applies. -/// -/// The centre itself always returns 0, because its true distance is zero -/// and none of the floor is independent noise there. -/// -/// With no measured correlation, which is also what an inert estimator -/// gives, every other offset returns 1, reproducing the flat white-noise -/// floor exactly. -pub(super) fn spatial_offset_factor(dx: i32, dy: i32, rho: f32) -> f32 { - if dx == 0 && dy == 0 { - return 0.0; - } - if rho <= 0.0 { - return 1.0; - } - let d = ((dx * dx + dy * dy) as f32).sqrt(); - 1.0 - (d * rho.ln()).exp() -} - -/// How many `f32`s a spatial-offset table needs at a given search -/// radius. -pub(super) fn spatial_offset_lut_len(search_radius: u32) -> usize { - let side = (2 * search_radius + 1) as usize; - side * side -} - -/// Builds the per-candidate noise-floor table for a search window, laid -/// out row-major. -/// -/// Each entry is the flat noise offset scaled by how much of it applies -/// at that candidate's distance. -/// -/// It is cheap enough to rebuild on every submit, reaching at most 289 -/// entries at the largest supported search radius. -pub(super) fn build_spatial_offset_lut(search_radius: u32, rho: f32, noise_offset: f32) -> Vec { - let r = search_radius as i32; - let side = (2 * search_radius + 1) as usize; - let mut lut = vec![0.0f32; side * side]; - for dy in -r..=r { - for dx in -r..=r { - let idx = ((dy + r) as usize) * side + (dx + r) as usize; - lut[idx] = noise_offset * spatial_offset_factor(dx, dy, rho); - } - } - lut -} - -/// How much weight the newest frame's estimate carries when smoothing. -/// -/// The sigma estimator below and the denoiser's own correlation -/// smoothing both use it. -pub(super) const EMA_ALPHA: f32 = 0.2; -/// The smallest smoothed sigma allowed, in normalised units, which works -/// out at 0.1 in 8-bit terms. -/// -/// An estimate near zero would send the derived strength to infinity. -const SIGMA_FLOOR: f32 = 0.1 / 255.0; - -/// The noise state for one stream. -/// -/// It smooths the per-frame estimates over time, so a single busy frame -/// cannot spike the strength, and applies a floor so near-clean content -/// keeps a usable normalisation factor. -#[derive(Debug, Default)] -pub(super) struct NoiseEstimator { - ema: Option>, -} - -impl NoiseEstimator { - /// Folds a new set of per-channel sigmas into the running estimate - /// and returns the smoothed result. - /// - /// The first call sets the state directly from the sample, because - /// there is no earlier estimate to blend with. So does every call - /// once `windowed` is set, which drops the running estimate - /// entirely and replaces it with this sample: the caller wants the - /// current window's own reading, not one blended with history from - /// frames outside it. - /// - /// Every element is floored at [`SIGMA_FLOOR`] either way. - pub(super) fn update(&mut self, sigmas: &[f32], windowed: bool) -> &[f32] { - match &mut self.ema { - Some(ema) if !windowed => { - for (e, &s) in ema.iter_mut().zip(sigmas.iter()) { - *e = (EMA_ALPHA * s + (1.0 - EMA_ALPHA) * *e).max(SIGMA_FLOOR); - } - }, - _ => { - self.ema = Some(sigmas.iter().map(|&s| s.max(SIGMA_FLOOR)).collect()); - }, - } - self.ema.as_deref().unwrap() - } - - /// Clears the running estimate. - /// - /// The next [`Self::update`] then starts from its own sample rather - /// than blending with stale state. - pub(super) fn reset(&mut self) { - self.ema = None; - } - - /// The current smoothed per-channel sigma, or `None` before the - /// first [`Self::update`] call. - pub(super) fn current(&self) -> Option<&[f32]> { - self.ema.as_deref() - } -} - #[cfg(test)] -mod tests { - use super::*; - - /// One single-channel block record with the given scalar lanes and - /// every quarter lane at 0. - fn scalar_record(sum_d: f32, sum_d2: f32, sum_lag: f32) -> Vec { - let record_len = temporal_stats_record_len(1) as usize; - let mut record = vec![0.0f32; record_len]; - record[0] = sum_d; - record[1] = sum_d2; - record[2] = sum_lag; - record - } - - #[test] - fn first_sample_initializes_state() { - let mut est = NoiseEstimator::default(); - let out = est.update(&[0.05, 0.02], false); - assert_eq!(out, &[0.05, 0.02]); - } - - #[test] - fn update_converges_toward_changed_level() { - let mut est = NoiseEstimator::default(); - est.update(&[0.02], false); - - let mut last = 0.0; - for _ in 0..200 { - last = est.update(&[0.10], false)[0]; - } - - assert!( - (last - 0.10).abs() < 1e-4, - "expected convergence close to 0.10, got {last}" - ); - } - - #[test] - fn update_floors_near_zero_samples() { - let mut est = NoiseEstimator::default(); - let out = est.update(&[0.0], false); - assert_eq!(out[0], SIGMA_FLOOR); - - let out = est.update(&[0.0], false); - assert_eq!(out[0], SIGMA_FLOOR); - } - - #[test] - fn reset_clears_state() { - let mut est = NoiseEstimator::default(); - est.update(&[0.10], false); - est.reset(); - - // After a reset the next update should start from its own - // sample rather than blending with the stale 0.10. - let out = est.update(&[0.02], false); - assert_eq!(out, &[0.02]); - } - - /// The property `windowed_noise_estimation` relies on: with - /// `windowed` set, every call ignores whatever the running estimate - /// held and returns exactly the sample just given, floored the same - /// way the very first sample is. - #[test] - fn windowed_update_ignores_prior_state() { - let mut est = NoiseEstimator::default(); - est.update(&[0.10], false); - est.update(&[0.20], false); - - let out = est.update(&[0.02], true); - assert_eq!(out, &[0.02]); - - // A second windowed call still ignores the first windowed - // call's own result, not just the pre-windowed history. - let out = est.update(&[0.0], true); - assert_eq!(out[0], SIGMA_FLOOR); - } - - #[test] - fn sigma_from_abs_sum_zero_for_zero_response() { - // No mask response anywhere in the frame has to estimate zero - // noise, whatever the frame size. - assert_eq!(sigma_from_abs_sum(0.0, 64, 64), 0.0); - } - - #[test] - fn partials_len_matches_cube_grid() { - assert_eq!(partials_len(32, 8), 4); // exactly one BLOCK_X x BLOCK_Y cube - assert_eq!(partials_len(33, 9), 16); // spills into a 2x2 cube grid - assert_eq!(partials_len(1920, 1080), 60 * 135 * 4); - } - - #[test] - fn noise_partials_slot_stride_bytes_pads_odd_cube_count() { - // A single block, whose 4 values come to 16 bytes and pad up to - // the next 32-byte multiple. - assert_eq!(noise_partials_slot_stride_bytes(32, 8, StorageAlign::new(32)), 32); - } - - #[test] - fn noise_partials_slot_stride_bytes_aligned_count_unchanged() { - // A 2x2 grid of blocks, whose 16 values come to 64 bytes. - // That is already a multiple of 32, so nothing is padded. - assert_eq!(noise_partials_slot_stride_bytes(33, 9, StorageAlign::new(32)), 64); - } - - /// A 70x20 frame grids into a ragged 3x3 layout of blocks, two full - /// rows and columns plus a partial one each way, so every block's - /// interior overlap is a different size. - /// - /// These are the hand-computed interior areas, row by row, summing - /// exactly to the frame's own interior area of 1224. - const RAGGED_CUBE_AREAS: [[f32; 3]; 3] = [[217.0, 224.0, 35.0], [248.0, 256.0, 40.0], [93.0, 96.0, 15.0]]; - - /// A uniform mask response gives every block the same sigma whatever - /// its area, because the area cancels out of the formula. - /// - /// The lower quartile of identical values is that same value, and it - /// has to match the frame-wide estimate over the same total - /// response. - #[test] - fn sigma_block_p25_from_partials_uniform_response_matches_frame_wide() { - let width = 70; - let height = 20; - let channels = 1; - let r = 0.02f32; - - let mut partials = vec![0.0f32; 3 * 3 * 4]; - let mut total_sum = 0.0f32; - for (cy, row) in RAGGED_CUBE_AREAS.iter().enumerate() { - for (cx, &area) in row.iter().enumerate() { - let sum = r * area; - partials[(cy * 3 + cx) * 4] = sum; - total_sum += sum; - } - } - - let sigma_low = sigma_block_p25_from_partials(&partials, channels, width, height); - let expected = sigma_from_abs_sum(total_sum, width, height); - - assert!( - (sigma_low[0] - expected).abs() < expected * 1e-4, - "uniform response should reproduce the frame-wide estimate {expected}, got {}", - sigma_low[0] - ); - } - - /// Nine blocks are given a shuffled set of sigmas from 1 to 9 in - /// 8-bit units, scaled through each block's own area, so sorting - /// them reproduces that run exactly. - /// - /// With nine values the lower quartile lands exactly on the third - /// smallest, with no interpolation. - #[test] - fn sigma_block_p25_from_partials_distinct_sums_pick_expected_cube() { - let width = 70; - let height = 20; - let channels = 1; - let sigma_targets_255 = [5.0f32, 2.0, 8.0, 1.0, 9.0, 3.0, 7.0, 4.0, 6.0]; - - let mut partials = vec![0.0f32; 3 * 3 * 4]; - for (cy, row) in RAGGED_CUBE_AREAS.iter().enumerate() { - for (cx, &area) in row.iter().enumerate() { - let i = cy * 3 + cx; - let sigma_target = sigma_targets_255[i] / 255.0; - let sum = sigma_target * 6.0 * area / std::f32::consts::FRAC_PI_2.sqrt(); - partials[i * 4] = sum; - } - } - - let sigma_low = sigma_block_p25_from_partials(&partials, channels, width, height); - let expected = 3.0 / 255.0; - assert!( - (sigma_low[0] - expected).abs() < 1e-4, - "expected the lower quartile to land on the third-smallest cube sigma {expected}, got {}", - sigma_low[0] - ); - } - - #[test] - fn temporal_stats_blocks_and_slot_len() { - assert_eq!(temporal_stats_blocks(32, 16), (2, 1)); - assert_eq!(temporal_stats_blocks(33, 17), (3, 2)); // ragged on both axes - assert_eq!(temporal_stats_record_len(1), 39); - assert_eq!(temporal_stats_record_len(4), 45); - assert_eq!(temporal_stats_slot_len(32, 16, 1), 78); // 2 blocks x record_len 39 - } - - #[test] - fn lower_quartile_odd_count_exact_index() { - // With 5 values the quartile lands on index 1 exactly, so no - // interpolation is needed. - assert_eq!(lower_quartile(&[1.0, 2.0, 3.0, 4.0, 5.0]), 2.0); - } - - #[test] - fn lower_quartile_even_count() { - // With 4 values the quartile lands at 0.75, between the first - // two. - let got = lower_quartile(&[1.0, 2.0, 3.0, 4.0]); - assert!((got - 1.75).abs() < 1e-6, "expected 1.75, got {got}"); - } - - #[test] - fn lower_quartile_interpolates_at_a_fractional_index() { - // With 3 values the quartile lands halfway between the first - // two. - let got = lower_quartile(&[10.0, 20.0, 30.0]); - assert!((got - 15.0).abs() < 1e-6, "expected 15.0, got {got}"); - } - - #[test] - fn lower_quartile_single_element_returns_it() { - assert_eq!(lower_quartile(&[42.0]), 42.0); - } - - #[test] - fn aggregate_only_static_block_contributes_to_sigma_and_rho() { - let width = 32; - let height = 16; - let stored_ch = 1; - let channels = 1; - let n = 256.0f32; - let n_pairs = 240.0f32; - - let sigma_target = 4.0 / 255.0; - let var = 2.0 * sigma_target * sigma_target; - let rho_target = 0.5f32; - let sum_d2_static = n * var; - let sum_lag_static = n_pairs * rho_target * var; - - let mean0_bad = 10.0 / 255.0; - let sum_d_bad = n * mean0_bad; - - let static_block = scalar_record(0.0, sum_d2_static, sum_lag_static); - let moving_block = scalar_record(sum_d_bad, 0.0, 0.0); - let records = [static_block, moving_block].concat(); - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("one of two blocks passes the static gate, above the 5% floor"); - - assert!((sample.static_fraction - 0.5).abs() < 1e-6); - assert!( - (sample.sigma[0] - sigma_target).abs() < 1e-4, - "sigma {} vs target {sigma_target}", - sample.sigma[0] - ); - assert!( - (sample.rho - rho_target).abs() < 1e-4, - "rho {} vs target {rho_target}", - sample.rho - ); - } - - /// Three perfectly static blocks, the letterbox shape, alongside five - /// carrying real noise. - /// - /// The static blocks are over a quarter of the population, so an - /// ungated lower quartile lands on zero. - #[test] - fn aggregate_excludes_perfectly_static_blocks_from_the_population() { - let width = 16 * 8; - let height = 16; - let stored_ch = 1; - let channels = 1; - let n = 256.0f32; - let n_pairs = 240.0f32; - let rho_target = 0.5f32; - - let sigmas_255 = [0.0f32, 0.0, 0.0, 2.0, 2.5, 3.0, 3.5, 4.0]; - - let mut records = Vec::new(); - for sigma_255 in sigmas_255.iter() { - let sigma = sigma_255 / 255.0; - let var = 2.0 * sigma * sigma; - let sum_d2 = n * var; - let sum_lag = n_pairs * rho_target * var; - let record = scalar_record(0.0, sum_d2, sum_lag); - records.extend(record); - } - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("five blocks clear the gate"); - - assert!( - (sample.static_fraction - 0.625).abs() < 1e-6, - "expected 5 of 8 blocks counted, got {}", - sample.static_fraction - ); - assert!( - (sample.sigma[0] - 3.0 / 255.0).abs() < 1e-4, - "expected median sigma 3/255 over [2,2.5,3,3.5,4], got {}", - sample.sigma[0] - ); - assert!( - (sample.sigma_low[0] - 2.5 / 255.0).abs() < 1e-4, - "expected lower-quartile sigma 2.5/255, not a zero dragged down by the static blocks, \ - got {}", - sample.sigma_low[0] - ); - } - - /// Five static blocks, each with its own sigma, four of which carry - /// real noise. - /// - /// That covers the odd-count median, the even-count median, and the - /// lower quartile in one pass. - #[test] - fn aggregate_median_over_multiple_static_blocks() { - let width = 16 * 5; - let height = 16; - let stored_ch = 1; - let channels = 1; - let n = 256.0f32; - let n_pairs = 240.0f32; - - // Block 0's sigma sits below RHO_SIGMA_GATE, so it carries no - // measurable noise and never enters the population. - let sigmas_255 = [0.1f32, 2.0, 3.0, 4.0, 5.0]; - let rhos = [0.0f32, 0.1, 0.3, 0.5, 0.7]; - - let mut records = Vec::new(); - for (sigma_255, rho) in sigmas_255.iter().zip(rhos.iter()) { - let sigma = sigma_255 / 255.0; - let var = 2.0 * sigma * sigma; - let sum_d2 = n * var; - let sum_lag = n_pairs * rho * var; - let record = scalar_record(0.0, sum_d2, sum_lag); - records.extend(record); - } - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("all five blocks are static"); - - assert!((sample.static_fraction - 0.8).abs() < 1e-6); - assert!( - (sample.sigma[0] - 3.5 / 255.0).abs() < 1e-4, - "expected median sigma 3.5/255 (middle of [2,3,4,5]), got {}", - sample.sigma[0] - ); - assert!( - (sample.sigma_low[0] - 2.75 / 255.0).abs() < 1e-4, - "expected lower-quartile sigma 2.75/255 (index 0.25*3=0.75 of [2,3,4,5]), got {}", - sample.sigma_low[0] - ); - assert!( - (sample.rho - 0.4).abs() < 1e-4, - "expected median rho 0.4 over the four blocks clearing the rho gate, got {}", - sample.rho - ); - } - - /// Only 1 of 25 blocks is static, which is below - /// `STATIC_FRACTION_MIN`. - /// - /// That single block carries perfectly valid noise otherwise, which - /// isolates the static-fraction floor from the other fallback. - #[test] - fn aggregate_below_static_floor_returns_none() { - let width = 16 * 5; - let height = 16 * 5; - let stored_ch = 1; - let channels = 1; - let n = 256.0f32; - - let record_len = temporal_stats_record_len(stored_ch) as usize; - let mut records = vec![0.0f32; 25 * record_len]; - let mean0_bad = 10.0 / 255.0; - for block in 1..25 { - records[block * record_len] = n * mean0_bad; - } - let sigma = 4.0 / 255.0; - let var = 2.0 * sigma * sigma; - records[1] = n * var; - records[2] = 240.0 * 0.5 * var; - - assert!( - aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height).is_none(), - "1 of 25 static blocks (4%) should fall back below the 5% floor" - ); - } - - /// A zero-filled duplicate slot has to fall back to Immerkær rather - /// than report a made-up sigma of zero. - /// - /// Every block passes the static check trivially, so this covers the - /// no-measurable-noise path rather than the static-fraction floor. - #[test] - fn aggregate_zeroed_slot_returns_none() { - let width = 32; - let height = 32; - let stored_ch = 1; - let channels = 1; - let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); - let record_len = temporal_stats_record_len(stored_ch) as usize; - let records = vec![0.0f32; (blocks_x * blocks_y) as usize * record_len]; - - assert!( - aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height).is_none(), - "a zero-filled duplicate slot's stats must fall back to Immerkær, not report sigma=0" - ); - } - - /// With YUV storage the three real channels sit in four lanes, so - /// each channel's sums have to be read at the right stride and the - /// unused padding lane must never affect the result. - #[test] - fn aggregate_multi_channel_layout_reads_correct_offsets() { - let width = 16; - let height = 16; - let stored_ch = 4; - let channels = 3; - let n = 256.0f32; - let n_pairs = 240.0f32; - - let sigmas_255 = [2.0f32, 4.0, 6.0]; - let rho_target = 0.6f32; - - let record_len = temporal_stats_record_len(stored_ch) as usize; - let mut record = vec![0.0f32; record_len]; - let mut var0 = 0.0f32; - for (c, sigma_255) in sigmas_255.iter().enumerate() { - let sigma = sigma_255 / 255.0; - let var = 2.0 * sigma * sigma; - record[stored_ch as usize + c] = n * var; - if c == 0 { - var0 = var; - } - } - record[2 * stored_ch as usize] = n_pairs * rho_target * var0; - - let sample = aggregate_temporal_noise_stats(&record, channels, stored_ch, width, height) - .expect("the single block is static with measurable channel-0 noise"); - - for (c, sigma_255) in sigmas_255.iter().enumerate() { - let expected = sigma_255 / 255.0; - assert!( - (sample.sigma[c] - expected).abs() < 1e-4, - "channel {c}: expected {expected}, got {}", - sample.sigma[c] - ); - } - assert!((sample.rho - rho_target).abs() < 1e-4); - } - - /// Builds one block's record from a chosen sigma and correlation, - /// with a mean residual of 0 so the block clears [`STATIC_GATE`] - /// without effort. - /// - /// The outlier tests below share this and only vary each block's - /// sigma. - fn zero_mean_block_record(sigma_255: f32, rho: f32) -> Vec { - let n = 256.0f32; - let n_pairs = 240.0f32; - let sigma = sigma_255 / 255.0; - let var = 2.0 * sigma * sigma; - scalar_record(0.0, n * var, n_pairs * rho * var) - } - - /// A 64x64 frame of 16 blocks where most stand in for panning - /// texture. - /// - /// Each has a mean residual of zero, clearing [`STATIC_GATE`] the - /// same way real noise does, but a sigma an order of magnitude above - /// the genuinely static minority. - /// - /// Their correlation reading is zero, so they look exactly like - /// white noise on that measure. That rules out a correlation-based - /// check as the fix. - /// - /// Without the outlier check the mean check lets every block - /// through, and the texture majority dominates the median, reading - /// their sigma instead of the real noise floor. - #[test] - fn aggregate_rejects_majority_zero_mean_texture_outliers() { - let width = 64; - let height = 64; - let stored_ch = 1; - let channels = 1; - - let background_sigma_255 = 2.0f32; - let background_rho = 0.1f32; - let texture_sigma_255 = 20.0f32; - - // 6 genuinely static blocks against 10 panning-texture ones. - // The texture blocks are the majority of the 16, so a plain - // median would read their level instead. - let mut records = Vec::new(); - for _ in 0..6 { - records.extend_from_slice(&zero_mean_block_record(background_sigma_255, background_rho)); - } - for _ in 0..10 { - records.extend_from_slice(&zero_mean_block_record(texture_sigma_255, 0.0)); - } - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("the static minority clears STATIC_FRACTION_MIN on its own"); - - let expected_sigma = background_sigma_255 / 255.0; - assert!( - (sample.sigma[0] - expected_sigma).abs() < 1e-4, - "expected the outlier gate to isolate the real noise floor {expected_sigma}, got {}", - sample.sigma[0] - ); - assert!( - (sample.static_fraction - 6.0 / 16.0).abs() < 1e-4, - "expected only the 6 background blocks to survive both gates, got static_fraction={}", - sample.static_fraction - ); - assert!( - (sample.rho - background_rho).abs() < 1e-4, - "expected rho to come from the surviving background blocks only, got {}", - sample.rho - ); - } - - /// The opposite failure. Every block here is genuinely static, but - /// the frame's real noise varies threefold across it, the way a dark - /// region often reads noisier than a bright one. - /// - /// None of that spread is texture, so the outlier check must not - /// drop any of it. - #[test] - fn aggregate_keeps_genuinely_static_blocks_despite_spatial_sigma_spread() { - let width = 64; - let height = 64; - let stored_ch = 1; - let channels = 1; - - let low_sigma_255 = 2.0f32; - let high_sigma_255 = 6.0f32; // 3x low_sigma_255, real spatial spread. - let rho = 0.1f32; - - let mut records = Vec::new(); - for _ in 0..8 { - records.extend_from_slice(&zero_mean_block_record(low_sigma_255, rho)); - } - for _ in 0..8 { - records.extend_from_slice(&zero_mean_block_record(high_sigma_255, rho)); - } - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("every block is static"); - - assert!( - (sample.static_fraction - 1.0).abs() < 1e-6, - "a real 3x spatial sigma spread must not trip the outlier gate, got static_fraction={}", - sample.static_fraction - ); - } - - /// Eight identical background blocks pin the surviving population's - /// lower quartile, whatever a ninth block's sigma turns out to be. - /// - /// With nine values the quartile lands exactly on the third - /// smallest, which is still one of the eight identical ones as long - /// as the ninth sorts above them. - /// - /// A 48x48 frame grids into exactly nine blocks, so that ninth block - /// is the only thing that varies. - fn outlier_factor_boundary_records( - background_sigma_255: f32, - background_rho: f32, - ninth_ratio: f32, - ) -> Vec { - let mut records = Vec::new(); - for _ in 0..8 { - records.extend_from_slice(&zero_mean_block_record(background_sigma_255, background_rho)); - } - records.extend_from_slice(&zero_mean_block_record(background_sigma_255 * ninth_ratio, 0.0)); - records - } - - /// The two ratios bracketing the outlier check's calibrated - /// boundary. - /// - /// They are written as literals rather than derived from - /// [`SIGMA_OUTLIER_FACTOR`], so the tests below assert the measured - /// calibration rather than the check's own arithmetic. Deriving them - /// would reduce to comparing the constant with itself, which proves - /// nothing about where it should sit. - /// - /// The measurement came from a scratchpad simulation during the - /// original investigation, which is not in this repo. A genuinely - /// static frame with a real noise spread up to fourfold survives, - /// and trimming starts at fivefold. - /// - /// So 5.0 was chosen with room above the real spread and margin - /// below where texture contamination is caught. - /// - /// Changing [`SIGMA_OUTLIER_FACTOR`] on purpose means updating these - /// two literals to match, or these tests will rightly fail. - const OUTLIER_FACTOR_SURVIVES_RATIO: f32 = 4.99; - const OUTLIER_FACTOR_REJECTS_RATIO: f32 = 5.01; - - #[test] - fn aggregate_outlier_factor_survives_just_under_threshold() { - let width = 48; - let height = 48; - let stored_ch = 1; - let channels = 1; - let background_sigma_255 = 2.0f32; - let background_rho = 0.1f32; - - let records = outlier_factor_boundary_records( - background_sigma_255, - background_rho, - OUTLIER_FACTOR_SURVIVES_RATIO, - ); - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("all 9 blocks clear the static-fraction floor"); - - assert!( - (sample.static_fraction - 1.0).abs() < 1e-6, - "a 9th block at {OUTLIER_FACTOR_SURVIVES_RATIO}x the reference must survive, \ - got static_fraction={}", - sample.static_fraction - ); - } - - #[test] - fn aggregate_outlier_factor_rejects_just_over_threshold() { - let width = 48; - let height = 48; - let stored_ch = 1; - let channels = 1; - let background_sigma_255 = 2.0f32; - let background_rho = 0.1f32; - - let records = outlier_factor_boundary_records( - background_sigma_255, - background_rho, - OUTLIER_FACTOR_REJECTS_RATIO, - ); - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("the 8 background blocks alone still clear the static-fraction floor"); - - assert!( - (sample.static_fraction - 8.0 / 9.0).abs() < 1e-6, - "a 9th block at {OUTLIER_FACTOR_REJECTS_RATIO}x the reference must be rejected, \ - got static_fraction={}", - sample.static_fraction - ); - } - - /// A 160x160 frame of 100 blocks, where 26 stand in for letterbox - /// bars with no residual at all and the other 74 carry real static - /// noise. - /// - /// Both groups clear [`STATIC_GATE`] without effort, the bars - /// because their residual is exactly zero and the noisy blocks - /// because theirs is centred on zero by construction, so all 100 - /// survive the first pass. - /// - /// # Why the noise levels vary - /// - /// The 74 noisy blocks cycle through five close sigma levels rather - /// than one repeated value, standing in for the sampling variance - /// any real measurement carries even under physically uniform noise. - /// - /// That spread is what lets them pass the anchor check the outlier - /// gate applies whenever it excludes low-sigma blocks. See - /// [`aggregate_temporal_noise_stats`]. - /// - /// A population with no internal spread, sitting next to excluded - /// low-sigma blocks, cannot be told apart from a self-selected - /// outlier run, and reports `None` instead. See - /// [`aggregate_returns_none_when_the_only_above_gate_population_is_texture`]. - /// - /// # Why the expected value is exact - /// - /// Splitting 74 blocks five ways gives counts of 15, 15, 15, 15, and - /// 14 across the five levels, with the remainder falling to the - /// earliest ones. - /// - /// The 26 zero blocks sit below [`RHO_SIGMA_GATE`], so they never - /// enter the population. The median of the remaining 74 lands - /// squarely inside the third level's run, so the expected sigma is - /// exactly that level rather than an approximation. - #[test] - fn aggregate_returns_correct_sigma_with_letterbox_zero_population() { - let width = 160; - let height = 160; - let stored_ch = 1; - let channels = 1; - - let background_rho = 0.2f32; - let sigma_levels_255 = [3.8f32, 3.9, 4.0, 4.1, 4.2]; - - let mut records = Vec::new(); - for _ in 0..26 { - records.extend_from_slice(&zero_mean_block_record(0.0, 0.0)); - } - for i in 0..74 { - let sigma_255 = sigma_levels_255[i % sigma_levels_255.len()]; - records.extend_from_slice(&zero_mean_block_record(sigma_255, background_rho)); - } - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("74 of 100 blocks carry real static noise, far above the 5% floor"); - - let expected_sigma = 4.0 / 255.0; - assert!( - (sample.sigma[0] - expected_sigma).abs() < 1e-4, - "expected the letterbox bars to leave the real noise floor near {expected_sigma} intact, got {}", - sample.sigma[0] - ); - assert!( - (sample.rho - background_rho).abs() < 1e-4, - "expected rho to come from the real-noise blocks, got {}", - sample.rho - ); - } - - /// The case the outlier exclusion opens up. - /// - /// The same 26 letterbox blocks, but all 74 of the others are - /// panning texture at one repeated sigma, with no genuine noise - /// anywhere above the gate. - /// - /// Every surviving candidate holds the same value, so the reference - /// population's lower quartile is that value too, and a value never - /// exceeds a multiple of itself. The ceiling therefore accepts every - /// one of its own outliers. - /// - /// A population that uniform, sitting next to 26 excluded zero-sigma - /// blocks, is exactly the case this cannot resolve. It has to report - /// `None` and let the caller fall back to the Immerkær reading, - /// rather than confidently report the texture level. - #[test] - fn aggregate_returns_none_when_the_only_above_gate_population_is_texture() { - let width = 160; - let height = 160; - let stored_ch = 1; - let channels = 1; - - let texture_sigma_255 = 20.0f32; - - let mut records = Vec::new(); - for _ in 0..26 { - records.extend_from_slice(&zero_mean_block_record(0.0, 0.0)); - } - for _ in 0..74 { - records.extend_from_slice(&zero_mean_block_record(texture_sigma_255, 0.0)); - } - - assert!( - aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height).is_none(), - "a homogeneous above-gate population with no genuine low anchor must fall back to \ - None rather than report the texture level as sigma" - ); - } - - /// The same case as - /// [`aggregate_rejects_majority_zero_mean_texture_outliers`], with a - /// letterbox-style zero-sigma population mixed in. - /// - /// The outlier check must still pick out the real noise floor from - /// the texture majority, and the zero blocks must not disturb it. The - /// zero blocks themselves carry no measurable noise, so they never - /// join the survivors. - #[test] - fn aggregate_rejects_texture_outliers_with_zero_population_present() { - let width = 64; - let height = 80; // 4 x 5 TEMPORAL_NOISE_BLOCK grid, 20 blocks. - let stored_ch = 1; - let channels = 1; - - let background_sigma_255 = 2.0f32; - let background_rho = 0.1f32; - let texture_sigma_255 = 20.0f32; - - let mut records = Vec::new(); - for _ in 0..6 { - records.extend_from_slice(&zero_mean_block_record(background_sigma_255, background_rho)); - } - for _ in 0..10 { - records.extend_from_slice(&zero_mean_block_record(texture_sigma_255, 0.0)); - } - for _ in 0..4 { - records.extend_from_slice(&zero_mean_block_record(0.0, 0.0)); - } - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("the static minority clears STATIC_FRACTION_MIN on its own"); - - let expected_sigma = background_sigma_255 / 255.0; - assert!( - (sample.sigma[0] - expected_sigma).abs() < 1e-4, - "expected the outlier gate to isolate the real noise floor {expected_sigma} despite \ - the zero population, got {}", - sample.sigma[0] - ); - assert!( - (sample.static_fraction - 6.0 / 20.0).abs() < 1e-4, - "expected only the 6 background blocks to survive, the zero blocks carry no \ - measurable noise, got static_fraction={}", - sample.static_fraction - ); - assert!( - (sample.rho - background_rho).abs() < 1e-4, - "expected rho to come from the surviving background blocks only, got {}", - sample.rho - ); - } - - /// The same case as - /// [`aggregate_keeps_genuinely_static_blocks_despite_spatial_sigma_spread`], - /// with a letterbox-style zero-sigma population mixed in. - /// - /// A real spread of noise across the frame must still survive the - /// outlier check, and the zero blocks must not trip it either. The - /// zero blocks themselves carry no measurable noise, so they never - /// join the survivors. - #[test] - fn aggregate_keeps_static_spread_with_zero_population_present() { - let width = 64; - let height = 80; // 4 x 5 TEMPORAL_NOISE_BLOCK grid, 20 blocks. - let stored_ch = 1; - let channels = 1; - - let low_sigma_255 = 2.0f32; - let high_sigma_255 = 6.0f32; // 3x low_sigma_255, real spatial spread. - let rho = 0.1f32; - - let mut records = Vec::new(); - for _ in 0..8 { - records.extend_from_slice(&zero_mean_block_record(low_sigma_255, rho)); - } - for _ in 0..8 { - records.extend_from_slice(&zero_mean_block_record(high_sigma_255, rho)); - } - for _ in 0..4 { - records.extend_from_slice(&zero_mean_block_record(0.0, 0.0)); - } - - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("every non-zero block is static"); - - assert!( - (sample.static_fraction - 0.8).abs() < 1e-6, - "a real 3x spatial sigma spread plus a zero population must not trip the outlier \ - gate, and the 4 zero blocks carry no measurable noise, got static_fraction={}", - sample.static_fraction - ); - } - - /// A block's correlation estimate can read above 1, because the - /// lag-1 total is averaged over adjacent pairs while the mean and - /// variance are averaged over every pixel. - /// - /// Each row here follows the pattern that maximises the ratio - /// between a row's lag-1 sum and its total variance, repeated down - /// every row of the block and scaled to clear both checks. - /// - /// This is a standalone construction rather than the usual test data - /// in this module, because it probes the formula's own bound instead - /// of a realistic noise scenario. - #[test] - fn aggregate_rho_estimate_stays_within_unit_range() { - let width = TEMPORAL_NOISE_BLOCK; - let height = TEMPORAL_NOISE_BLOCK; - let stored_ch = 1; - let channels = 1; - - let scale = 0.007f32; - let row: Vec = (1..=width) - .map(|i| scale * (i as f32 * std::f32::consts::PI / (width + 1) as f32).sin()) - .collect(); - - let n = (width * height) as f32; - let sum_d: f32 = row.iter().sum::() * height as f32; - let sum_d2: f32 = row.iter().map(|v| v * v).sum::() * height as f32; - let sum_lag: f32 = row.windows(2).map(|w| w[0] * w[1]).sum::() * height as f32; - - // Confirms the construction actually clears both gates this - // block needs to reach the rho computation at all, so the - // assertion below is testing the clamp and not a gate miss. - assert!( - (sum_d / n).abs() < STATIC_GATE, - "construction must clear the static gate" - ); - let var = sum_d2 / n - (sum_d / n) * (sum_d / n); - assert!( - (var.sqrt() / std::f32::consts::SQRT_2) > RHO_SIGMA_GATE, - "construction must clear the rho-sample sigma gate" - ); - - let records = scalar_record(sum_d, sum_d2, sum_lag); - let sample = aggregate_temporal_noise_stats(&records, channels, stored_ch, width, height) - .expect("the single block clears both gates"); - - assert!( - (0.0..=1.0).contains(&sample.rho), - "the mismatched-denominator estimate must stay within [0, 1], got {}", - sample.rho - ); - } - - #[test] - fn interpolate_table_linear_between_points() { - let table = [(0.0f32, 1.0f32), (0.5, 1.2), (1.0, 1.5)]; - assert!((interpolate_table(&table, 0.25) - 1.1).abs() < 1e-6); - assert!((interpolate_table(&table, 0.75) - 1.35).abs() < 1e-6); - assert_eq!(interpolate_table(&table, 0.0), 1.0); - assert_eq!(interpolate_table(&table, 1.0), 1.5); - } - - #[test] - fn interpolate_table_clamps_outside_range() { - let table = [(0.0f32, 1.0f32), (1.0, 2.0)]; - assert_eq!(interpolate_table(&table, -5.0), 1.0); - assert_eq!(interpolate_table(&table, 5.0), 2.0); - } - - #[test] - fn correlation_factor_matches_measured_table() { - assert_eq!(correlation_factor(0.0), 1.0); - assert!((correlation_factor(0.65) - 1.45).abs() < 1e-6); - // Clamped flat past the last measured point. - assert!((correlation_factor(0.9) - 1.45).abs() < 1e-6); - // White noise stays uncorrected. - assert_eq!(correlation_factor(-0.2), 1.0); - // Monotone non-decreasing across the measured range. - let mut last = 0.0; - for i in 0..=20 { - let f = correlation_factor(i as f32 / 20.0); - assert!(f >= last, "factor must not decrease, {f} < {last}"); - last = f; - } - } - - #[test] - fn spatial_offset_factor_rho_zero_is_white_identity() { - // With no correlation, every candidate but the centre keeps the - // full white-noise factor of 1. A negative value, which the - // aggregation never produces, has to behave the same way. - for rho in [0.0f32, -0.2] { - for dy in -3..=3 { - for dx in -3..=3 { - if dx == 0 && dy == 0 { - continue; - } - assert_eq!( - spatial_offset_factor(dx, dy, rho), - 1.0, - "dx={dx} dy={dy} rho={rho}" - ); - } - } - } - } - - #[test] - fn spatial_offset_factor_self_is_always_zero() { - for rho in [-0.2f32, 0.0, 0.3, 0.65, 1.0] { - assert_eq!(spatial_offset_factor(0, 0, rho), 0.0, "rho={rho}"); - } - } - - #[test] - fn spatial_offset_factor_monotone_nondecreasing_in_distance() { - let rho = 0.65f32; - let mut last = spatial_offset_factor(0, 0, rho); - for d in 1..=8 { - let f = spatial_offset_factor(d, 0, rho); - assert!(f >= last, "factor decreased at d={d}: {f} < {last}"); - last = f; - } - } - - #[test] - fn spatial_offset_factor_rho_0_65_shape() { - let rho = 0.65f32; - // One pixel away. - assert!((spatial_offset_factor(1, 0, rho) - (1.0 - rho)).abs() < 1e-6); - // Two pixels away along an axis. - assert!((spatial_offset_factor(2, 0, rho) - (1.0 - rho * rho)).abs() < 1e-6); - // A diagonal candidate sits sqrt(2) away, straight from the - // distance formula. - let d = 2.0f32.sqrt(); - let expected = 1.0 - (d * rho.ln()).exp(); - assert!((spatial_offset_factor(1, 1, rho) - expected).abs() < 1e-6); - // Every factor stays within [0, 1]. - for dy in -4..=4 { - for dx in -4..=4 { - let f = spatial_offset_factor(dx, dy, rho); - assert!((0.0..=1.0).contains(&f), "dx={dx} dy={dy} factor={f}"); - } - } - } - - #[test] - fn build_spatial_offset_lut_rho_zero_matches_flat_noise_offset() { - let search_radius = 3; - let noise_offset = 1.5f32; - let lut = build_spatial_offset_lut(search_radius, 0.0, noise_offset); - assert_eq!(lut.len(), spatial_offset_lut_len(search_radius)); - - let side = (2 * search_radius + 1) as usize; - let r = search_radius as i32; - for dy in -r..=r { - for dx in -r..=r { - let idx = ((dy + r) as usize) * side + (dx + r) as usize; - if dx == 0 && dy == 0 { - assert_eq!(lut[idx], 0.0, "self offset must be zero"); - } else { - assert_eq!(lut[idx], noise_offset, "dx={dx} dy={dy}"); - } - } - } - } - - #[test] - fn build_spatial_offset_lut_indexes_row_major_by_dy_then_dx() { - let search_radius = 2; - let lut = build_spatial_offset_lut(search_radius, 0.65, 10.0); - let side = (2 * search_radius + 1) as usize; - - // One step to the right lands at row 2, column 3 of the table, - // which is index 13. - let expected = 10.0 * spatial_offset_factor(1, 0, 0.65); - assert_eq!(lut[2 * side + 3], expected); - - // (dx=0, dy=-2) lands at row 0, column 2: index 0*5+2=2. - let expected = 10.0 * spatial_offset_factor(0, -2, 0.65); - assert_eq!(lut[2], expected); - - // Centre (dx=0, dy=0) is always zero. - assert_eq!(lut[2 * side + 2], 0.0); - } -} +pub(super) use self::temporal::{ + TEMPORAL_NOISE_BLOCK, + TEMPORAL_QUARTER_SIZE, + accepted_static_blocks, + aggregate_temporal_noise_stats, +}; diff --git a/av-denoise-core/src/nlmeans/noise/spatial.rs b/av-denoise-core/src/nlmeans/noise/spatial.rs new file mode 100644 index 0000000..9ae82d9 --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/spatial.rs @@ -0,0 +1,157 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::stats::{lower_quartile, sort_ascending}; +use crate::nlmeans::align::StorageAlign; +use crate::nlmeans::kernels::{nlm_noise_partial, nlm_noise_reduce}; +use crate::nlmeans::{BLOCK_1D, BLOCK_X, BLOCK_Y}; + +/// The inputs one Immerkær noise estimate needs. +pub(in crate::nlmeans) struct NoiseCtx<'a> { + pub width: u32, + pub height: u32, + pub channels: u32, + pub stored_ch: u32, + pub frame_count: u32, + pub frame: u32, + pub slot: u32, + pub input_buf: &'a Handle, + pub partials_buf: &'a Handle, + pub results_buf: &'a Handle, +} + +/// The `f32` length of one frame's first-stage partials, four lanes per block. +pub(in crate::nlmeans) fn partials_len(width: u32, height: u32) -> usize { + (width.div_ceil(BLOCK_X) * height.div_ceil(BLOCK_Y) * 4) as usize +} + +/// The byte stride between partials ring slots, padded to the buffer-binding alignment. +/// +/// wgpu rejects a bind-group offset that is not a multiple of its +/// `min_storage_buffer_offset_alignment`, and a small frame can leave [partials_len] short of one. +pub(in crate::nlmeans) fn noise_partials_slot_stride_bytes( + width: u32, + height: u32, + align: StorageAlign, +) -> u64 { + let partials_bytes = partials_len(width, height) as u64 * size_of::() as u64; + align.pad_bytes(partials_bytes) +} + +/// Runs both stages of the Immerkær noise estimate for one frame. +/// +/// The mask is cheap and needs only one frame, but it reads correlated grain low because the grain +/// looks partly like content to it. The per-channel totals land in this frame's slot of the results +/// buffer, which holds four values per slot of the input ring's frame capacity. +pub(in crate::nlmeans) fn run_noise_estimate( + client: &ComputeClient, + ctx: &NoiseCtx<'_>, +) -> Result<(), anyhow::Error> { + let total_input = (ctx.frame_count * ctx.height * ctx.width * ctx.stored_ch) as usize; + let partials_count = partials_len(ctx.width, ctx.height); + let total_results = (ctx.frame_count * 4) as usize; + let stored_ch = ctx.stored_ch as usize; + let cubes_x = ctx.width.div_ceil(BLOCK_X); + let cubes_y = ctx.height.div_ceil(BLOCK_Y); + + unsafe { + nlm_noise_partial::launch_unchecked::( + client, + CubeCount::new_2d(cubes_x, cubes_y), + CubeDim::new_2d(BLOCK_X, BLOCK_Y), + stored_ch, + ArrayArg::from_raw_parts(ctx.input_buf.clone(), total_input), + ArrayArg::from_raw_parts(ctx.partials_buf.clone(), partials_count), + ctx.frame, + ctx.width, + ctx.height, + ctx.channels, + BLOCK_X, + BLOCK_Y, + ); + } + + let block_count = (partials_count / 4) as u32; + + unsafe { + nlm_noise_reduce::launch_unchecked::( + client, + CubeCount::new_1d(1), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(ctx.partials_buf.clone(), partials_count), + ArrayArg::from_raw_parts(ctx.results_buf.clone(), total_results), + ctx.slot, + block_count, + BLOCK_1D, + ); + } + + Ok(()) +} + +/// Turns the summed absolute mask responses into an Immerkær sigma. +/// +/// The interior area leaves out the one-pixel border the mask cannot reach. +pub(in crate::nlmeans) fn sigma_from_abs_sum(abs_sum: f32, width: u32, height: u32) -> f32 { + let interior = ((width - 2) as f32) * ((height - 2) as f32); + (std::f32::consts::FRAC_PI_2).sqrt() * abs_sum / (6.0 * interior) +} + +/// The per-channel lower quartile of the per-block Immerkær sigmas in one slot's partials. +/// +/// Each block's sigma uses the [sigma_from_abs_sum] formula over its tile's overlap with the frame +/// interior. A block with no overlap is skipped so a spurious zero cannot dilute the quartile. +/// Where noise is uneven across a frame this reads lower than the frame-wide mean, which makes it +/// the cautious estimate. Channels past the active count stay 0. +pub(in crate::nlmeans) fn sigma_block_p25_from_partials( + partials: &[f32], + channels: u32, + width: u32, + height: u32, +) -> [f32; 3] { + let cubes_x = width.div_ceil(BLOCK_X); + let cubes_y = height.div_ceil(BLOCK_Y); + let channels = channels as usize; + + let mut cube_sigmas: Vec> = vec![Vec::new(); channels]; + + for cube_y in 0..cubes_y { + let tile_y0 = cube_y * BLOCK_Y; + let tile_y1 = ((cube_y + 1) * BLOCK_Y).min(height); + let overlap_y0 = tile_y0.max(1); + let overlap_y1 = tile_y1.min(height - 1); + + for cube_x in 0..cubes_x { + let tile_x0 = cube_x * BLOCK_X; + let tile_x1 = ((cube_x + 1) * BLOCK_X).min(width); + let overlap_x0 = tile_x0.max(1); + let overlap_x1 = tile_x1.min(width - 1); + + if overlap_x1 <= overlap_x0 || overlap_y1 <= overlap_y0 { + continue; + } + + let area = ((overlap_x1 - overlap_x0) * (overlap_y1 - overlap_y0)) as f32; + let cube_index = (cube_y * cubes_x + cube_x) as usize; + let base = cube_index * 4; + + for (channel, sigmas) in cube_sigmas.iter_mut().enumerate() { + let sum = partials[base + channel]; + let cube_sigma = std::f32::consts::FRAC_PI_2.sqrt() * sum / (6.0 * area); + sigmas.push(cube_sigma); + } + } + } + + let mut sigma_low = [0.0f32; 3]; + for (channel, sigmas) in cube_sigmas.iter_mut().enumerate() { + if sigmas.is_empty() { + continue; + } + + sort_ascending(sigmas); + sigma_low[channel] = lower_quartile(sigmas); + } + + sigma_low +} diff --git a/av-denoise-core/src/nlmeans/noise/stats.rs b/av-denoise-core/src/nlmeans/noise/stats.rs new file mode 100644 index 0000000..6cd115a --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/stats.rs @@ -0,0 +1,38 @@ +/// Sorts `values` into ascending order in place. +/// +/// [median] and [lower_quartile] both need sorted input, so a caller wanting both sorts once. +pub(super) fn sort_ascending(values: &mut [f32]) { + values.sort_by(|left, right| left.partial_cmp(right).expect("noise stats are never NaN")); +} + +/// The median of a sorted, non-empty slice. +/// +/// An even count averages the two middle elements. +pub(super) fn median(values: &[f32]) -> f32 { + let count = values.len(); + if count % 2 == 1 { + values[count / 2] + } else { + (values[count / 2 - 1] + values[count / 2]) / 2.0 + } +} + +/// The lower quartile of a sorted, non-empty slice. +/// +/// It interpolates between the two neighbouring elements when the quarter point falls between them. +pub(super) fn lower_quartile(values: &[f32]) -> f32 { + let count = values.len(); + if count == 1 { + return values[0]; + } + + let position = 0.25 * (count - 1) as f32; + let lower_index = position.floor() as usize; + let upper_index = position.ceil() as usize; + if lower_index == upper_index { + values[lower_index] + } else { + let fraction = position - lower_index as f32; + values[lower_index] + fraction * (values[upper_index] - values[lower_index]) + } +} diff --git a/av-denoise-core/src/nlmeans/noise/strength_map.rs b/av-denoise-core/src/nlmeans/noise/strength_map.rs index 709fc4f..8c68462 100644 --- a/av-denoise-core/src/nlmeans/noise/strength_map.rs +++ b/av-denoise-core/src/nlmeans/noise/strength_map.rs @@ -1,5 +1,5 @@ use super::curve::{CLIP_HIGH, CLIP_LOW, NoiseCurve, QUARTER_STATIC_GATE}; -use super::{ +use super::temporal::{ QUARTER_FLATNESS, QUARTER_LUMA_MAX, QUARTER_LUMA_MIN, @@ -25,10 +25,13 @@ use crate::collab::geometry::strength_map_dims; /// /// Grain alone reads about 0.35 to 0.45 of it, so faint lines under heavy grain land above the cut. const QUARTER_FLAT_FACTOR: f32 = 0.55; + /// A quarter noisier than this many times the curve at its luma is treated as motion, not grain. const MOTION_FACTOR: f32 = 2.5; + /// The luma at and below which [StrengthMapParams::shadow_soften] applies in full. const SHADOW_LOW: f32 = 128.0 / 255.0; + /// The luma at which the soften has faded back to 1.0. const SHADOW_HIGH: f32 = 160.0 / 255.0; @@ -96,8 +99,9 @@ impl QuarterClasses { .collect() } - /// One chroma threshold multiplier per quarter, row-major. Flat quarters get `flat_boost`, and - /// every other quarter gets 1.0. + /// One chroma threshold multiplier per quarter, row-major. + /// + /// Flat quarters get `flat_boost`, and every other quarter gets 1.0. pub(crate) fn chroma_multipliers(&self, flat_boost: f32) -> Vec { self.classes .iter() @@ -245,15 +249,15 @@ pub(in crate::nlmeans) fn classify_quarters( for block_x in 0..blocks_x { let block_index = (block_y * blocks_x + block_x) as usize; let record = &records[block_index * record_len..(block_index + 1) * record_len]; - let block_w = (width - block_x * TEMPORAL_NOISE_BLOCK).min(TEMPORAL_NOISE_BLOCK); - let block_h = (height - block_y * TEMPORAL_NOISE_BLOCK).min(TEMPORAL_NOISE_BLOCK); + let block_width = (width - block_x * TEMPORAL_NOISE_BLOCK).min(TEMPORAL_NOISE_BLOCK); + let block_height = (height - block_y * TEMPORAL_NOISE_BLOCK).min(TEMPORAL_NOISE_BLOCK); for quarter_index in 0..TEMPORAL_QUARTERS { let offset_x = (quarter_index % 2) * TEMPORAL_QUARTER_SIZE; let offset_y = (quarter_index / 2) * TEMPORAL_QUARTER_SIZE; - let quarter_w = block_w.saturating_sub(offset_x).min(TEMPORAL_QUARTER_SIZE); - let quarter_h = block_h.saturating_sub(offset_y).min(TEMPORAL_QUARTER_SIZE); - let pixels = (quarter_w * quarter_h) as f32; + let quarter_width = block_width.saturating_sub(offset_x).min(TEMPORAL_QUARTER_SIZE); + let quarter_height = block_height.saturating_sub(offset_y).min(TEMPORAL_QUARTER_SIZE); + let pixels = (quarter_width * quarter_height) as f32; if pixels == 0.0 { continue; } @@ -276,6 +280,7 @@ pub(in crate::nlmeans) fn classify_quarters( } let mut quarter_classes = QuarterClasses { cols, rows, classes }; + let active_cut = texture_cut.filter(|&cut| cut < 1.0); if let Some(cut) = active_cut { quarter_classes.veto_textured(&tensors, cut); diff --git a/av-denoise-core/src/nlmeans/noise/temporal.rs b/av-denoise-core/src/nlmeans/noise/temporal.rs new file mode 100644 index 0000000..4d89bfd --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/temporal.rs @@ -0,0 +1,529 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::curve::{NoiseCurve, build_noise_curve}; +use super::stats::{lower_quartile, median, sort_ascending}; +use super::strength_map::{QuarterClasses, classify_quarters}; +use crate::nlmeans::align::StorageAlign; +use crate::nlmeans::kernels::{nlm_temporal_noise_stats, nlm_temporal_stats_zero}; +use crate::nlmeans::{BLOCK_1D, MAX_GRID_1D}; + +/// The side of one temporal-stats block in pixels, with one GPU block per square. +pub(in crate::nlmeans) const TEMPORAL_NOISE_BLOCK: u32 = 16; +/// How many 8x8 quarters one block record holds. +pub(in crate::nlmeans) const TEMPORAL_QUARTERS: u32 = 4; +pub(in crate::nlmeans) const TEMPORAL_QUARTER_SIZE: u32 = TEMPORAL_NOISE_BLOCK / 2; +/// The offset of the first quarter record past `2 * stored_ch`. +/// +/// Quarter `q` starts at `2 * stored_ch + 1 + 9 * q`, in top-left, top-right, bottom-left, +/// bottom-right order. +pub(in crate::nlmeans) const TEMPORAL_QUARTER_BASE: u32 = 1; +/// How many `f32`s one quarter record holds. +pub(in crate::nlmeans) const TEMPORAL_QUARTER_FIELDS: u32 = 9; +/// The quarter-record offset of channel 0's summed residual. +pub(in crate::nlmeans) const QUARTER_SUM_D: u32 = 0; +/// The quarter-record offset of channel 0's summed squared residual. +pub(in crate::nlmeans) const QUARTER_SUM_D2: u32 = 1; +/// The quarter-record offset of the new frame's summed luma. +pub(in crate::nlmeans) const QUARTER_LUMA_SUM: u32 = 2; +/// The quarter-record offset of the smoothed temporal mean's gradient energy. +/// +/// A quarter smaller than 8x8 reads `3.0e38`. +pub(in crate::nlmeans) const QUARTER_FLATNESS: u32 = 3; +/// The quarter-record offset of the new frame's minimum luma. +pub(in crate::nlmeans) const QUARTER_LUMA_MIN: u32 = 4; +/// The quarter-record offset of the new frame's maximum luma. +pub(in crate::nlmeans) const QUARTER_LUMA_MAX: u32 = 5; +/// The quarter-record offset of the temporal mean's summed squared horizontal gradient. +pub(in crate::nlmeans) const QUARTER_TENSOR_XX: u32 = 6; +/// The quarter-record offset of the temporal mean's summed squared vertical gradient. +pub(in crate::nlmeans) const QUARTER_TENSOR_YY: u32 = 7; +/// The quarter-record offset of the temporal mean's summed horizontal times vertical gradient. +pub(in crate::nlmeans) const QUARTER_TENSOR_XY: u32 = 8; + +/// The largest mean block residual, in normalised units, that still counts as static. +/// +/// A block above this is moving content rather than noise. +pub(super) const STATIC_GATE: f32 = 1.5 / 255.0; +/// The smallest block sigma that counts as measurable noise. +/// +/// Below it a block carries too little signal for its correlation reading to mean anything. It +/// gates the correlation median, the outlier ceiling's reference set and the sigma population. +pub(super) const RHO_SIGMA_GATE: f32 = 0.3 / 255.0; +/// The smallest fraction of static blocks a sample needs to be trusted. +/// +/// Below it, motion, a scene change or too little measurable noise dominates the frame, and the +/// Immerkær estimate is the only usable reading. +const STATIC_FRACTION_MIN: f32 = 0.05; +/// How far above the surviving blocks' lower quartile a block's sigma may sit before it counts as +/// moving texture. +/// +/// A block panning across texture can average to nearly nothing and clear [STATIC_GATE], while +/// its variance is shifted texture running several times above the frame's other blocks. Real +/// noise keeps a much narrower spread, even where a dark region reads noisier than a bright one, +/// and this factor leaves room for that spread. +/// +/// The reference is the lower quartile rather than the median, so the check still works when +/// texture makes up most of the surviving blocks, as long as a static minority remains to anchor +/// it. The quartile only counts blocks above [RHO_SIGMA_GATE], because letterbox bars and other +/// perfectly static regions read a sigma of exactly 0 and would drag it to 0, rejecting every +/// block with real noise. +const SIGMA_OUTLIER_FACTOR: f32 = 5.0; + +/// How many `f32`s one block's stats record holds. +/// +/// That is a sum and a sum of squares per stored channel, one lag-1 total, and the four quarter +/// records. +pub(in crate::nlmeans) fn temporal_stats_record_len(stored_ch: u32) -> u32 { + 2 * stored_ch + TEMPORAL_QUARTER_BASE + TEMPORAL_QUARTERS * TEMPORAL_QUARTER_FIELDS +} + +/// The row-major block grid covering a frame, with ragged edges truncated rather than padded. +pub(in crate::nlmeans) fn temporal_stats_blocks(width: u32, height: u32) -> (u32, u32) { + ( + width.div_ceil(TEMPORAL_NOISE_BLOCK), + height.div_ceil(TEMPORAL_NOISE_BLOCK), + ) +} + +pub(super) fn temporal_stats_slot_len(width: u32, height: u32, stored_ch: u32) -> usize { + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); + (blocks_x * blocks_y * temporal_stats_record_len(stored_ch)) as usize +} + +/// The byte stride between temporal-stats ring slots, padded to the buffer-binding alignment. +/// +/// wgpu rejects a bind-group offset that is not a multiple of its +/// `min_storage_buffer_offset_alignment`, and a small frame or a single-channel mode can leave +/// [temporal_stats_slot_len] short of one. +fn temporal_stats_slot_stride_bytes(width: u32, height: u32, stored_ch: u32, align: StorageAlign) -> u64 { + let slot_bytes = temporal_stats_slot_len(width, height, stored_ch) as u64 * size_of::() as u64; + align.pad_bytes(slot_bytes) +} + +pub(in crate::nlmeans) fn temporal_stats_buf_bytes( + width: u32, + height: u32, + stored_ch: u32, + frame_count: u32, + align: StorageAlign, +) -> usize { + (temporal_stats_slot_stride_bytes(width, height, stored_ch, align) * frame_count as u64) as usize +} + +/// The inputs one temporal residual statistics dispatch needs, comparing `slot_new` against +/// `slot_prev`. +pub(in crate::nlmeans) struct TemporalStatsCtx<'a> { + pub width: u32, + pub height: u32, + pub stored_ch: u32, + pub frame_count: u32, + pub slot_new: u32, + pub slot_prev: u32, + pub input_buf: &'a Handle, + pub stats_buf: &'a Handle, + pub align: StorageAlign, +} + +/// Runs the temporal residual statistics kernel, writing one record per block into the new slot. +/// +/// Where nothing moved the residual is noise, so correlated grain shows in full, but motion and +/// scene changes make the reading unreliable. The kernel addresses only its own slot's slice, so +/// it needs nothing about the ring's other slots or the padding between them. With `luma_fields` +/// off, the kernel's four luma-only lanes read 0 at no extra cost. +pub(in crate::nlmeans) fn run_temporal_noise_stats( + client: &ComputeClient, + ctx: &TemporalStatsCtx<'_>, + luma_fields: bool, +) -> Result<(), anyhow::Error> { + let total_input = (ctx.frame_count * ctx.height * ctx.width * ctx.stored_ch) as usize; + let (blocks_x, blocks_y) = temporal_stats_blocks(ctx.width, ctx.height); + let slot_len = temporal_stats_slot_len(ctx.width, ctx.height, ctx.stored_ch); + let stride = temporal_stats_slot_stride_bytes(ctx.width, ctx.height, ctx.stored_ch, ctx.align); + let stats_slot = ctx.stats_buf.clone().offset_start((ctx.slot_new as u64) * stride); + + unsafe { + nlm_temporal_noise_stats::launch_unchecked::( + client, + CubeCount::new_2d(blocks_x, blocks_y), + CubeDim::new_2d(TEMPORAL_NOISE_BLOCK, TEMPORAL_NOISE_BLOCK), + ctx.stored_ch as usize, + ArrayArg::from_raw_parts(ctx.input_buf.clone(), total_input), + ArrayArg::from_raw_parts(stats_slot, slot_len), + ctx.slot_new, + ctx.slot_prev, + ctx.width, + ctx.height, + ctx.stored_ch, + TEMPORAL_NOISE_BLOCK, + luma_fields, + ); + } + + Ok(()) +} + +/// Fills one ring slot's temporal-stats region with zeroes. +/// +/// It runs on a slot copied from the one before it. The zeroes read as no static blocks with +/// measurable noise rather than as a reading of zero noise. +pub(in crate::nlmeans) fn zero_temporal_stats_slot( + client: &ComputeClient, + stats_buf: &Handle, + width: u32, + height: u32, + stored_ch: u32, + slot: u32, + align: StorageAlign, +) { + let slot_len = temporal_stats_slot_len(width, height, stored_ch) as u32; + let stride = temporal_stats_slot_stride_bytes(width, height, stored_ch, align); + let slot_region = stats_buf.clone().offset_start((slot as u64) * stride); + + let grid = slot_len.div_ceil(BLOCK_1D).min(MAX_GRID_1D); + let total_threads = grid * BLOCK_1D; + + unsafe { + nlm_temporal_stats_zero::launch_unchecked::( + client, + CubeCount::new_1d(grid), + CubeDim::new_1d(BLOCK_1D), + ArrayArg::from_raw_parts(slot_region, slot_len as usize), + slot_len, + total_threads, + ); + } +} + +/// Reads one ring slot's temporal-stats region back. +/// +/// The ring handle is sliced to the slot, so the transfer skips the rest of the ring. +#[expect( + clippy::too_many_arguments, + reason = "the readback needs the ring handle plus every shape value that locates one slot in it" +)] +pub(in crate::nlmeans) fn read_temporal_stats_slot( + client: &ComputeClient, + stats_buf: &Handle, + width: u32, + height: u32, + stored_ch: u32, + frame_count: u32, + slot: u32, + align: StorageAlign, +) -> Result, anyhow::Error> { + let slot_len_bytes = temporal_stats_slot_len(width, height, stored_ch) as u64 * size_of::() as u64; + let stride = temporal_stats_slot_stride_bytes(width, height, stored_ch, align); + let total_bytes = frame_count as u64 * stride; + let start = (slot as u64) * stride; + let end_trim = total_bytes - start - slot_len_bytes; + + let sliced = stats_buf.clone().offset_start(start).offset_end(end_trim); + let bytes = client + .read_one(sliced) + .map_err(|error| anyhow::anyhow!("temporal noise stats readback failed: {error}"))?; + let records = f32::from_bytes(&bytes).to_vec(); + + Ok(records) +} + +/// One centre slot's aggregated temporal-residual noise measurement. +#[derive(Debug, Clone, Copy, PartialEq)] +pub(in crate::nlmeans) struct TemporalNoiseSample { + /// The per-channel median sigma over the static blocks with measurable noise, in normalised + /// units. + /// + /// Entries past the active channel count stay 0. + pub sigma: [f32; 3], + /// The per-channel lower-quartile sigma in normalised units, for consumers where reading high + /// does more harm than reading low. + /// + /// Entries past the active channel count stay 0. + pub sigma_low: [f32; 3], + /// The median lag-1 correlation of the residuals, measuring how correlated the grain is + /// between neighbouring pixels. + pub rho: f32, + /// The fraction of blocks counted as static with measurable noise. + pub static_fraction: f32, +} + +/// A block that passed [STATIC_GATE], held until the outlier ceiling is known. +struct StaticGateCandidate { + index: usize, + width: u32, + height: u32, + sigmas: [f32; 3], + sigma_ch0: f32, + var_ch0: f32, + mean0: f32, + mean_lag: f32, + n_pairs: f32, +} + +/// One block [accepted_static_blocks] counted as static noise. +pub(in crate::nlmeans) struct AcceptedBlock { + /// The block's position in `records`, in units of one record. + pub(super) index: usize, + /// The in-frame width, which a ragged edge truncates. + pub(super) width: u32, + /// The in-frame height, which a ragged edge truncates. + pub(super) height: u32, + /// The per-channel sigma, with entries past the active channel count at 0. + pub(super) sigmas: [f32; 3], + /// Channel 0's mean residual. + pub(super) mean0: f32, + /// Channel 0's mean lag-1 product over adjacent pixel pairs. + pub(super) mean_lag: f32, + /// Channel 0's residual variance. + pub(super) var_ch0: f32, + /// How many adjacent pixel pairs the block holds. + pub(super) n_pairs: f32, +} + +/// Picks out the blocks in one slot's records that count as static noise. +/// +/// Three filters run in turn. [STATIC_GATE] keeps blocks whose mean residual is near zero, which a +/// block panning across texture also passes. [SIGMA_OUTLIER_FACTOR] then rejects blocks far above +/// the population's lower quartile, since a panning block's variance comes from texture rather +/// than shared noise. Last, blocks at or below [RHO_SIGMA_GATE] carry no measurable noise and are +/// dropped. +/// +/// It returns `None` when the outlier ceiling cannot be trusted. Excluding low-sigma blocks lets an +/// all-texture remainder set its own ceiling and pass it, because a value never exceeds a multiple +/// of itself. Real per-block sigma always varies a little over a finite sample, so a remainder with +/// no spread beside excluded blocks is reported as `None` rather than as a sigma texture may have +/// inflated. +pub(in crate::nlmeans) fn accepted_static_blocks( + records: &[f32], + channels: u32, + stored_ch: u32, + width: u32, + height: u32, +) -> Option> { + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); + let total_blocks = (blocks_x * blocks_y) as usize; + if total_blocks == 0 { + return None; + } + + let record_len = temporal_stats_record_len(stored_ch) as usize; + let channels = channels as usize; + let stored_ch = stored_ch as usize; + + let mut candidates = Vec::new(); + + for block_y in 0..blocks_y { + for block_x in 0..blocks_x { + let block_index = (block_y * blocks_x + block_x) as usize; + let record = &records[block_index * record_len..(block_index + 1) * record_len]; + + let block_origin_x = block_x * TEMPORAL_NOISE_BLOCK; + let block_origin_y = block_y * TEMPORAL_NOISE_BLOCK; + let block_width = TEMPORAL_NOISE_BLOCK.min(width - block_origin_x); + let block_height = TEMPORAL_NOISE_BLOCK.min(height - block_origin_y); + let pixel_count = (block_width * block_height) as f32; + let n_pairs = (block_height * block_width.saturating_sub(1)) as f32; + + let mean0 = record[0] / pixel_count; + if mean0.abs() >= STATIC_GATE { + continue; + } + + let mut sigmas = [0.0f32; 3]; + let mut sigma_ch0 = 0.0f32; + let mut var_ch0 = 0.0f32; + for channel in 0..channels { + let mean = record[channel] / pixel_count; + let variance = (record[stored_ch + channel] / pixel_count - mean * mean).max(0.0); + let sigma_block = variance.sqrt() / std::f32::consts::SQRT_2; + sigmas[channel] = sigma_block; + if channel == 0 { + sigma_ch0 = sigma_block; + var_ch0 = variance; + } + } + + let mean_lag = if n_pairs > 0.0 { + record[2 * stored_ch] / n_pairs + } else { + 0.0 + }; + + candidates.push(StaticGateCandidate { + index: block_index, + width: block_width, + height: block_height, + sigmas, + sigma_ch0, + var_ch0, + mean0, + mean_lag, + n_pairs, + }); + } + } + + // Perfectly static blocks such as letterbox bars read a sigma of 0. Left in the reference they + // drag the quartile to 0, which rejects every block with real noise. + let mut reference_sigma_ch0: Vec = candidates + .iter() + .map(|candidate| candidate.sigma_ch0) + .filter(|&sigma| sigma > RHO_SIGMA_GATE) + .collect(); + sort_ascending(&mut reference_sigma_ch0); + + // A reference with no spread beside excluded blocks may be all texture setting its own + // ceiling. With nothing excluded the population is trusted directly. + let candidates_were_excluded = candidates.len() > reference_sigma_ch0.len(); + let reference_has_spread = match (reference_sigma_ch0.first(), reference_sigma_ch0.last()) { + (Some(&lowest), Some(&highest)) => highest > lowest, + (None, _) | (_, None) => false, + }; + if candidates_were_excluded && !reference_has_spread { + return None; + } + + let sigma_ceiling = if reference_sigma_ch0.is_empty() { + 0.0 + } else { + lower_quartile(&reference_sigma_ch0) * SIGMA_OUTLIER_FACTOR + }; + + let mut accepted = Vec::new(); + + for candidate in candidates { + if candidate.sigma_ch0 > sigma_ceiling { + continue; + } + + // A zero-sigma block carries no measurable noise, and a population of them would drag + // the lower quartile to 0. + if candidate.sigma_ch0 <= RHO_SIGMA_GATE { + continue; + } + + accepted.push(AcceptedBlock { + index: candidate.index, + width: candidate.width, + height: candidate.height, + sigmas: candidate.sigmas, + mean0: candidate.mean0, + mean_lag: candidate.mean_lag, + var_ch0: candidate.var_ch0, + n_pairs: candidate.n_pairs, + }); + } + + Some(accepted) +} + +/// The scalar sample [temporal_noise_reading] builds, without a curve. +#[cfg(test)] +pub(in crate::nlmeans) fn aggregate_temporal_noise_stats( + records: &[f32], + channels: u32, + stored_ch: u32, + width: u32, + height: u32, +) -> Option { + temporal_noise_reading(records, channels, stored_ch, width, height, false, None).sample +} + +/// One centre slot's temporal-noise reading. +pub(in crate::nlmeans) struct TemporalNoiseReading { + pub(in crate::nlmeans) sample: Option, + /// The per-frame luma noise curve, present only when requested and a sample exists. + pub(in crate::nlmeans) curve: Option, + /// The frame's quarter classes, present exactly when `curve` is. + pub(in crate::nlmeans) classes: Option, +} + +/// Builds one centre slot's [TemporalNoiseReading] from a single [accepted_static_blocks] call. +/// +/// Sharing that selection keeps the scalar sample and the curve agreeing on which blocks are +/// static noise. The sample is `None` when the outlier ceiling cannot be trusted, when the static +/// fraction falls below [STATIC_FRACTION_MIN], or when no static block carries measurable noise, +/// as in a zero-filled duplicate slot. +/// +/// The curve needs `with_curve` and a sample, since a frame too unreliable for a scalar sample is +/// too unreliable for a curve. `texture_cut` is the cut [classify_quarters] vetoes textured flat +/// quarters at, or `None` for no veto. +pub(in crate::nlmeans) fn temporal_noise_reading( + records: &[f32], + channels: u32, + stored_ch: u32, + width: u32, + height: u32, + with_curve: bool, + texture_cut: Option, +) -> TemporalNoiseReading { + let none = TemporalNoiseReading { + sample: None, + curve: None, + classes: None, + }; + + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); + let total_blocks = (blocks_x * blocks_y) as usize; + + let Some(accepted) = accepted_static_blocks(records, channels, stored_ch, width, height) else { + return none; + }; + + let channels = channels as usize; + + let mut rho_samples = Vec::new(); + for block in &accepted { + if block.n_pairs > 0.0 { + let rho = (block.mean_lag - block.mean0 * block.mean0) / block.var_ch0; + let clamped_rho = rho.clamp(0.0, 1.0); + rho_samples.push(clamped_rho); + } + } + + let static_fraction = accepted.len() as f32 / total_blocks as f32; + if static_fraction < STATIC_FRACTION_MIN || rho_samples.is_empty() { + return none; + } + + let mut static_sigmas: Vec> = vec![Vec::new(); channels]; + for block in &accepted { + for (channel, sigmas) in static_sigmas.iter_mut().enumerate().take(channels) { + sigmas.push(block.sigmas[channel]); + } + } + + let mut sigma = [0.0f32; 3]; + let mut sigma_low = [0.0f32; 3]; + for (channel, sigmas) in static_sigmas.iter_mut().enumerate() { + sort_ascending(sigmas); + sigma[channel] = median(sigmas); + sigma_low[channel] = lower_quartile(sigmas); + } + + sort_ascending(&mut rho_samples); + let rho = median(&rho_samples); + + let sample = TemporalNoiseSample { + sigma, + sigma_low, + rho, + static_fraction, + }; + + let curve = if with_curve { + build_noise_curve(records, stored_ch, &accepted, sample.sigma[0]) + } else { + None + }; + + let classes = curve + .as_ref() + .map(|curve| classify_quarters(records, stored_ch, width, height, curve, texture_cut)); + + TemporalNoiseReading { + sample: Some(sample), + curve, + classes, + } +} diff --git a/av-denoise-core/src/nlmeans/noise/tests/correlation.rs b/av-denoise-core/src/nlmeans/noise/tests/correlation.rs new file mode 100644 index 0000000..02d7ffb --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/tests/correlation.rs @@ -0,0 +1,160 @@ +use crate::nlmeans::noise::correlation::{ + build_spatial_offset_lut, + correlation_factor, + interpolate_table, + spatial_offset_factor, + spatial_offset_lut_len, +}; + +#[test] +fn interpolate_table_linear_between_points() { + let table = [(0.0f32, 1.0f32), (0.5, 1.2), (1.0, 1.5)]; + let quarter = interpolate_table(&table, 0.25); + let three_quarters = interpolate_table(&table, 0.75); + let start = interpolate_table(&table, 0.0); + let end = interpolate_table(&table, 1.0); + + assert!((quarter - 1.1).abs() < 1e-6); + assert!((three_quarters - 1.35).abs() < 1e-6); + assert_eq!(start, 1.0); + assert_eq!(end, 1.5); +} + +#[test] +fn interpolate_table_clamps_outside_range() { + let table = [(0.0f32, 1.0f32), (1.0, 2.0)]; + let below = interpolate_table(&table, -5.0); + let above = interpolate_table(&table, 5.0); + + assert_eq!(below, 1.0); + assert_eq!(above, 2.0); +} + +#[test] +fn correlation_factor_matches_measured_table() { + let uncorrelated = correlation_factor(0.0); + let measured = correlation_factor(0.65); + let past_last_point = correlation_factor(0.9); + let white = correlation_factor(-0.2); + + assert_eq!(uncorrelated, 1.0); + assert!((measured - 1.45).abs() < 1e-6); + // Clamped flat past the last measured point. + assert!((past_last_point - 1.45).abs() < 1e-6); + // White noise stays uncorrected. + assert_eq!(white, 1.0); + + // Monotone non-decreasing across the measured range. + let mut last = 0.0; + for i in 0..=20 { + let factor = correlation_factor(i as f32 / 20.0); + assert!(factor >= last, "factor must not decrease, {factor} < {last}"); + last = factor; + } +} + +#[test] +fn spatial_offset_factor_rho_zero_is_white_identity() { + // A negative rho, which the aggregation never produces, behaves like zero. + for rho in [0.0f32, -0.2] { + for dy in -3..=3 { + for dx in -3..=3 { + if dx == 0 && dy == 0 { + continue; + } + + let factor = spatial_offset_factor(dx, dy, rho); + assert_eq!(factor, 1.0, "dx={dx} dy={dy} rho={rho}"); + } + } + } +} + +#[test] +fn spatial_offset_factor_self_is_always_zero() { + for rho in [-0.2f32, 0.0, 0.3, 0.65, 1.0] { + let factor = spatial_offset_factor(0, 0, rho); + assert_eq!(factor, 0.0, "rho={rho}"); + } +} + +#[test] +fn spatial_offset_factor_monotone_nondecreasing_in_distance() { + let rho = 0.65f32; + let mut last = spatial_offset_factor(0, 0, rho); + for distance in 1..=8 { + let factor = spatial_offset_factor(distance, 0, rho); + assert!( + factor >= last, + "factor decreased at d={distance}: {factor} < {last}" + ); + last = factor; + } +} + +#[test] +fn spatial_offset_factor_rho_0_65_shape() { + let rho = 0.65f32; + let one_pixel = spatial_offset_factor(1, 0, rho); + let two_pixels = spatial_offset_factor(2, 0, rho); + let diagonal = spatial_offset_factor(1, 1, rho); + + // One pixel away. + assert!((one_pixel - (1.0 - rho)).abs() < 1e-6); + // Two pixels away along an axis. + assert!((two_pixels - (1.0 - rho * rho)).abs() < 1e-6); + + // A diagonal candidate sits sqrt(2) away. + let distance = 2.0f32.sqrt(); + let log_rho = rho.ln(); + let expected = 1.0 - (distance * log_rho).exp(); + assert!((diagonal - expected).abs() < 1e-6); + + // Every factor stays within 0..=1. + for dy in -4..=4 { + for dx in -4..=4 { + let factor = spatial_offset_factor(dx, dy, rho); + assert!((0.0..=1.0).contains(&factor), "dx={dx} dy={dy} factor={factor}"); + } + } +} + +#[test] +fn build_spatial_offset_lut_rho_zero_matches_flat_noise_offset() { + let search_radius = 3; + let noise_offset = 1.5f32; + let lut = build_spatial_offset_lut(search_radius, 0.0, noise_offset); + let expected_len = spatial_offset_lut_len(search_radius); + assert_eq!(lut.len(), expected_len); + + let side = (2 * search_radius + 1) as usize; + let radius = search_radius as i32; + for dy in -radius..=radius { + for dx in -radius..=radius { + let index = ((dy + radius) as usize) * side + (dx + radius) as usize; + if dx == 0 && dy == 0 { + assert_eq!(lut[index], 0.0, "self offset must be zero"); + } else { + assert_eq!(lut[index], noise_offset, "dx={dx} dy={dy}"); + } + } + } +} + +#[test] +fn build_spatial_offset_lut_indexes_row_major_by_dy_then_dx() { + let search_radius = 2; + let lut = build_spatial_offset_lut(search_radius, 0.65, 10.0); + let side = (2 * search_radius + 1) as usize; + + // One step right lands at row 2, column 3, which is index 13. + let expected = 10.0 * spatial_offset_factor(1, 0, 0.65); + assert_eq!(lut[2 * side + 3], expected); + + // dx=0, dy=-2 lands at row 0, column 2, which is index 2. + let expected = 10.0 * spatial_offset_factor(0, -2, 0.65); + assert_eq!(lut[2], expected); + + // The centre is always zero. + assert_eq!(lut[2 * side + 2], 0.0); +} diff --git a/av-denoise-core/src/nlmeans/noise/tests/estimator.rs b/av-denoise-core/src/nlmeans/noise/tests/estimator.rs new file mode 100644 index 0000000..83fbe38 --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/tests/estimator.rs @@ -0,0 +1,58 @@ +use crate::nlmeans::noise::estimator::{NoiseEstimator, SIGMA_FLOOR}; + +#[test] +fn first_sample_initializes_state() { + let mut estimator = NoiseEstimator::default(); + let smoothed = estimator.update(&[0.05, 0.02], false); + assert_eq!(smoothed, &[0.05, 0.02]); +} + +#[test] +fn update_converges_toward_changed_level() { + let mut estimator = NoiseEstimator::default(); + estimator.update(&[0.02], false); + + let mut last = 0.0; + for _ in 0..200 { + last = estimator.update(&[0.10], false)[0]; + } + + assert!( + (last - 0.10).abs() < 1e-4, + "expected convergence close to 0.10, got {last}" + ); +} + +#[test] +fn update_floors_near_zero_samples() { + let mut estimator = NoiseEstimator::default(); + let smoothed = estimator.update(&[0.0], false); + assert_eq!(smoothed[0], SIGMA_FLOOR); + + let smoothed = estimator.update(&[0.0], false); + assert_eq!(smoothed[0], SIGMA_FLOOR); +} + +#[test] +fn reset_clears_state() { + let mut estimator = NoiseEstimator::default(); + estimator.update(&[0.10], false); + estimator.reset(); + + let smoothed = estimator.update(&[0.02], false); + assert_eq!(smoothed, &[0.02]); +} + +#[test] +fn windowed_update_ignores_prior_state() { + let mut estimator = NoiseEstimator::default(); + estimator.update(&[0.10], false); + estimator.update(&[0.20], false); + + let smoothed = estimator.update(&[0.02], true); + assert_eq!(smoothed, &[0.02]); + + // A second windowed call ignores the first windowed result too. + let smoothed = estimator.update(&[0.0], true); + assert_eq!(smoothed[0], SIGMA_FLOOR); +} diff --git a/av-denoise-core/src/nlmeans/noise/tests/mod.rs b/av-denoise-core/src/nlmeans/noise/tests/mod.rs new file mode 100644 index 0000000..f453322 --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/tests/mod.rs @@ -0,0 +1,5 @@ +mod correlation; +mod estimator; +mod spatial; +mod stats; +mod temporal; diff --git a/av-denoise-core/src/nlmeans/noise/tests/spatial.rs b/av-denoise-core/src/nlmeans/noise/tests/spatial.rs new file mode 100644 index 0000000..fc3445e --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/tests/spatial.rs @@ -0,0 +1,100 @@ +use crate::nlmeans::align::StorageAlign; +use crate::nlmeans::noise::spatial::{ + noise_partials_slot_stride_bytes, + partials_len, + sigma_block_p25_from_partials, + sigma_from_abs_sum, +}; + +/// Hand-computed interior areas of a 70x20 frame's ragged 3x3 block grid, summing to 1224. +const RAGGED_CUBE_AREAS: [[f32; 3]; 3] = [[217.0, 224.0, 35.0], [248.0, 256.0, 40.0], [93.0, 96.0, 15.0]]; + +#[test] +fn sigma_from_abs_sum_zero_for_zero_response() { + let sigma = sigma_from_abs_sum(0.0, 64, 64); + assert_eq!(sigma, 0.0); +} + +#[test] +fn partials_len_matches_cube_grid() { + let one_cube = partials_len(32, 8); + let spilled = partials_len(33, 9); + let full_hd = partials_len(1920, 1080); + + assert_eq!(one_cube, 4); // exactly one BLOCK_X x BLOCK_Y cube + assert_eq!(spilled, 16); // spills into a 2x2 cube grid + assert_eq!(full_hd, 60 * 135 * 4); +} + +#[test] +fn noise_partials_slot_stride_bytes_pads_odd_cube_count() { + // One block's 16 bytes pad up to the next 32-byte multiple. + let align = StorageAlign::new(32); + let stride = noise_partials_slot_stride_bytes(32, 8, align); + assert_eq!(stride, 32); +} + +#[test] +fn noise_partials_slot_stride_bytes_aligned_count_unchanged() { + // A 2x2 grid's 64 bytes are already a multiple of 32. + let align = StorageAlign::new(32); + let stride = noise_partials_slot_stride_bytes(33, 9, align); + assert_eq!(stride, 64); +} + +/// The area cancels out of each block's sigma, so a uniform response gives every block the same +/// sigma as the frame-wide estimate. +#[test] +fn sigma_block_p25_from_partials_uniform_response_matches_frame_wide() { + let width = 70; + let height = 20; + let channels = 1; + let response = 0.02f32; + + let mut partials = vec![0.0f32; 3 * 3 * 4]; + let mut total_sum = 0.0f32; + for (cube_y, row) in RAGGED_CUBE_AREAS.iter().enumerate() { + for (cube_x, &area) in row.iter().enumerate() { + let sum = response * area; + partials[(cube_y * 3 + cube_x) * 4] = sum; + total_sum += sum; + } + } + + let sigma_low = sigma_block_p25_from_partials(&partials, channels, width, height); + let expected = sigma_from_abs_sum(total_sum, width, height); + + assert!( + (sigma_low[0] - expected).abs() < expected * 1e-4, + "uniform response should reproduce the frame-wide estimate {expected}, got {}", + sigma_low[0] + ); +} + +/// Nine blocks get a shuffled 1..=9 run of 8-bit sigmas, so the lower quartile lands exactly on +/// the third smallest. +#[test] +fn sigma_block_p25_from_partials_distinct_sums_pick_expected_cube() { + let width = 70; + let height = 20; + let channels = 1; + let sigma_targets_255 = [5.0f32, 2.0, 8.0, 1.0, 9.0, 3.0, 7.0, 4.0, 6.0]; + + let mut partials = vec![0.0f32; 3 * 3 * 4]; + for (cube_y, row) in RAGGED_CUBE_AREAS.iter().enumerate() { + for (cube_x, &area) in row.iter().enumerate() { + let cube = cube_y * 3 + cube_x; + let sigma_target = sigma_targets_255[cube] / 255.0; + let sum = sigma_target * 6.0 * area / std::f32::consts::FRAC_PI_2.sqrt(); + partials[cube * 4] = sum; + } + } + + let sigma_low = sigma_block_p25_from_partials(&partials, channels, width, height); + let expected = 3.0 / 255.0; + assert!( + (sigma_low[0] - expected).abs() < 1e-4, + "expected the lower quartile to land on the third-smallest cube sigma {expected}, got {}", + sigma_low[0] + ); +} diff --git a/av-denoise-core/src/nlmeans/noise/tests/stats.rs b/av-denoise-core/src/nlmeans/noise/tests/stats.rs new file mode 100644 index 0000000..ebe70f5 --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/tests/stats.rs @@ -0,0 +1,28 @@ +use crate::nlmeans::noise::stats::lower_quartile; + +#[test] +fn lower_quartile_odd_count_exact_index() { + // With 5 values the quartile lands exactly on index 1. + let got = lower_quartile(&[1.0, 2.0, 3.0, 4.0, 5.0]); + assert_eq!(got, 2.0); +} + +#[test] +fn lower_quartile_even_count() { + // With 4 values the quartile lands at 0.75, between the first two. + let got = lower_quartile(&[1.0, 2.0, 3.0, 4.0]); + assert!((got - 1.75).abs() < 1e-6, "expected 1.75, got {got}"); +} + +#[test] +fn lower_quartile_interpolates_at_a_fractional_index() { + // With 3 values the quartile lands halfway between the first two. + let got = lower_quartile(&[10.0, 20.0, 30.0]); + assert!((got - 15.0).abs() < 1e-6, "expected 15.0, got {got}"); +} + +#[test] +fn lower_quartile_single_element_returns_it() { + let got = lower_quartile(&[42.0]); + assert_eq!(got, 42.0); +} diff --git a/av-denoise-core/src/nlmeans/noise/tests/temporal.rs b/av-denoise-core/src/nlmeans/noise/tests/temporal.rs new file mode 100644 index 0000000..2b7f278 --- /dev/null +++ b/av-denoise-core/src/nlmeans/noise/tests/temporal.rs @@ -0,0 +1,621 @@ +use crate::nlmeans::noise::temporal::{ + RHO_SIGMA_GATE, + STATIC_GATE, + TEMPORAL_NOISE_BLOCK, + aggregate_temporal_noise_stats, + temporal_stats_blocks, + temporal_stats_record_len, + temporal_stats_slot_len, +}; + +// The two ratios bracketing the outlier check's calibrated boundary. They are literals so the tests +// check the calibration rather than the constant against itself. A static frame with a real noise +// spread up to fourfold survives and trimming starts at fivefold, so a factor of 5.0 leaves room +// above real spread and catches texture. Changing `SIGMA_OUTLIER_FACTOR` means updating both. +const OUTLIER_FACTOR_SURVIVES_RATIO: f32 = 4.99; +const OUTLIER_FACTOR_REJECTS_RATIO: f32 = 5.01; + +/// One single-channel block record with the given scalar lanes and every quarter lane at 0. +fn scalar_record(sum_d: f32, sum_d2: f32, sum_lag: f32) -> Vec { + let record_len = temporal_stats_record_len(1) as usize; + let mut record = vec![0.0f32; record_len]; + record[0] = sum_d; + record[1] = sum_d2; + record[2] = sum_lag; + + record +} + +/// One full block's record with a zero mean residual, so it always clears [STATIC_GATE]. +fn zero_mean_block_record(sigma_255: f32, rho: f32) -> Vec { + let pixel_count = 256.0f32; + let n_pairs = 240.0f32; + let sigma = sigma_255 / 255.0; + let variance = 2.0 * sigma * sigma; + let sum_d2 = pixel_count * variance; + let sum_lag = n_pairs * rho * variance; + + scalar_record(0.0, sum_d2, sum_lag) +} + +/// A 48x48 frame of nine blocks, eight identical and a ninth at `ninth_ratio` times their sigma. +/// +/// The lower quartile of nine lands on the third smallest, which stays one of the eight identical +/// blocks while the ninth sorts above them. +fn outlier_factor_boundary_records( + background_sigma_255: f32, + background_rho: f32, + ninth_ratio: f32, +) -> Vec { + let mut records = Vec::new(); + for _ in 0..8 { + let record = zero_mean_block_record(background_sigma_255, background_rho); + records.extend_from_slice(&record); + } + + let ninth_sigma_255 = background_sigma_255 * ninth_ratio; + let ninth_record = zero_mean_block_record(ninth_sigma_255, 0.0); + records.extend_from_slice(&ninth_record); + + records +} + +#[test] +fn temporal_stats_blocks_and_slot_len() { + let exact_blocks = temporal_stats_blocks(32, 16); + let ragged_blocks = temporal_stats_blocks(33, 17); + let single_record_len = temporal_stats_record_len(1); + let padded_record_len = temporal_stats_record_len(4); + let slot_len = temporal_stats_slot_len(32, 16, 1); + + assert_eq!(exact_blocks, (2, 1)); + assert_eq!(ragged_blocks, (3, 2)); // ragged on both axes + assert_eq!(single_record_len, 39); + assert_eq!(padded_record_len, 45); + assert_eq!(slot_len, 78); // 2 blocks x record_len 39 +} + +#[test] +fn aggregate_only_static_block_contributes_to_sigma_and_rho() { + let width = 32; + let height = 16; + let stored_channels = 1; + let channels = 1; + let pixel_count = 256.0f32; + let n_pairs = 240.0f32; + + let sigma_target = 4.0 / 255.0; + let variance = 2.0 * sigma_target * sigma_target; + let rho_target = 0.5f32; + let sum_d2_static = pixel_count * variance; + let sum_lag_static = n_pairs * rho_target * variance; + + let mean0_bad = 10.0 / 255.0; + let sum_d_bad = pixel_count * mean0_bad; + + let static_block = scalar_record(0.0, sum_d2_static, sum_lag_static); + let moving_block = scalar_record(sum_d_bad, 0.0, 0.0); + let records = [static_block, moving_block].concat(); + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("one of two blocks passes the static gate, above the 5% floor"); + + assert!((sample.static_fraction - 0.5).abs() < 1e-6); + assert!( + (sample.sigma[0] - sigma_target).abs() < 1e-4, + "sigma {} vs target {sigma_target}", + sample.sigma[0] + ); + assert!( + (sample.rho - rho_target).abs() < 1e-4, + "rho {} vs target {rho_target}", + sample.rho + ); +} + +/// Three letterbox-style zero blocks are over a quarter of eight, so an ungated lower quartile +/// would land on zero. +#[test] +fn aggregate_excludes_perfectly_static_blocks_from_the_population() { + let width = 16 * 8; + let height = 16; + let stored_channels = 1; + let channels = 1; + let pixel_count = 256.0f32; + let n_pairs = 240.0f32; + let rho_target = 0.5f32; + + let sigmas_255 = [0.0f32, 0.0, 0.0, 2.0, 2.5, 3.0, 3.5, 4.0]; + + let mut records = Vec::new(); + for sigma_255 in sigmas_255.iter() { + let sigma = sigma_255 / 255.0; + let variance = 2.0 * sigma * sigma; + let sum_d2 = pixel_count * variance; + let sum_lag = n_pairs * rho_target * variance; + let record = scalar_record(0.0, sum_d2, sum_lag); + records.extend(record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("five blocks clear the gate"); + + assert!( + (sample.static_fraction - 0.625).abs() < 1e-6, + "expected 5 of 8 blocks counted, got {}", + sample.static_fraction + ); + assert!( + (sample.sigma[0] - 3.0 / 255.0).abs() < 1e-4, + "expected median sigma 3/255 over [2,2.5,3,3.5,4], got {}", + sample.sigma[0] + ); + assert!( + (sample.sigma_low[0] - 2.5 / 255.0).abs() < 1e-4, + "expected lower-quartile sigma 2.5/255, not a zero dragged down by the static blocks, \ + got {}", + sample.sigma_low[0] + ); +} + +#[test] +fn aggregate_median_over_multiple_static_blocks() { + let width = 16 * 5; + let height = 16; + let stored_channels = 1; + let channels = 1; + let pixel_count = 256.0f32; + let n_pairs = 240.0f32; + + // Block 0 sits below RHO_SIGMA_GATE, so it never enters the population. + let sigmas_255 = [0.1f32, 2.0, 3.0, 4.0, 5.0]; + let rhos = [0.0f32, 0.1, 0.3, 0.5, 0.7]; + + let mut records = Vec::new(); + for (sigma_255, rho) in sigmas_255.iter().zip(rhos.iter()) { + let sigma = sigma_255 / 255.0; + let variance = 2.0 * sigma * sigma; + let sum_d2 = pixel_count * variance; + let sum_lag = n_pairs * rho * variance; + let record = scalar_record(0.0, sum_d2, sum_lag); + records.extend(record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("all five blocks are static"); + + assert!((sample.static_fraction - 0.8).abs() < 1e-6); + assert!( + (sample.sigma[0] - 3.5 / 255.0).abs() < 1e-4, + "expected median sigma 3.5/255 (middle of [2,3,4,5]), got {}", + sample.sigma[0] + ); + assert!( + (sample.sigma_low[0] - 2.75 / 255.0).abs() < 1e-4, + "expected lower-quartile sigma 2.75/255 (index 0.25*3=0.75 of [2,3,4,5]), got {}", + sample.sigma_low[0] + ); + assert!( + (sample.rho - 0.4).abs() < 1e-4, + "expected median rho 0.4 over the four blocks clearing the rho gate, got {}", + sample.rho + ); +} + +/// The one static block carries valid noise, which isolates the static-fraction floor. +#[test] +fn aggregate_below_static_floor_returns_none() { + let width = 16 * 5; + let height = 16 * 5; + let stored_channels = 1; + let channels = 1; + let pixel_count = 256.0f32; + + let record_len = temporal_stats_record_len(stored_channels) as usize; + let mut records = vec![0.0f32; 25 * record_len]; + let mean0_bad = 10.0 / 255.0; + for block in 1..25 { + records[block * record_len] = pixel_count * mean0_bad; + } + + let sigma = 4.0 / 255.0; + let variance = 2.0 * sigma * sigma; + records[1] = pixel_count * variance; + records[2] = 240.0 * 0.5 * variance; + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height); + assert!( + sample.is_none(), + "1 of 25 static blocks (4%) should fall back below the 5% floor" + ); +} + +/// Every block passes the static check, so this reaches the no-measurable-noise path. +#[test] +fn aggregate_zeroed_slot_returns_none() { + let width = 32; + let height = 32; + let stored_channels = 1; + let channels = 1; + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); + let record_len = temporal_stats_record_len(stored_channels) as usize; + let records = vec![0.0f32; (blocks_x * blocks_y) as usize * record_len]; + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height); + assert!( + sample.is_none(), + "a zero-filled duplicate slot's stats must fall back to Immerkær, not report sigma=0" + ); +} + +/// YUV storage puts three channels in four lanes, and the padding lane must never affect the +/// result. +#[test] +fn aggregate_multi_channel_layout_reads_correct_offsets() { + let width = 16; + let height = 16; + let stored_channels = 4; + let channels = 3; + let pixel_count = 256.0f32; + let n_pairs = 240.0f32; + + let sigmas_255 = [2.0f32, 4.0, 6.0]; + let rho_target = 0.6f32; + + let record_len = temporal_stats_record_len(stored_channels) as usize; + let mut record = vec![0.0f32; record_len]; + let mut luma_variance = 0.0f32; + for (channel, sigma_255) in sigmas_255.iter().enumerate() { + let sigma = sigma_255 / 255.0; + let variance = 2.0 * sigma * sigma; + record[stored_channels as usize + channel] = pixel_count * variance; + if channel == 0 { + luma_variance = variance; + } + } + + record[2 * stored_channels as usize] = n_pairs * rho_target * luma_variance; + + let sample = aggregate_temporal_noise_stats(&record, channels, stored_channels, width, height) + .expect("the single block is static with measurable channel-0 noise"); + + for (channel, sigma_255) in sigmas_255.iter().enumerate() { + let expected = sigma_255 / 255.0; + assert!( + (sample.sigma[channel] - expected).abs() < 1e-4, + "channel {channel}: expected {expected}, got {}", + sample.sigma[channel] + ); + } + + assert!((sample.rho - rho_target).abs() < 1e-4); +} + +/// Ten of sixteen blocks stand in for panning texture, with a zero mean, a tenfold sigma and a +/// white-noise rho, so a correlation check cannot catch them. +#[test] +fn aggregate_rejects_majority_zero_mean_texture_outliers() { + let width = 64; + let height = 64; + let stored_channels = 1; + let channels = 1; + + let background_sigma_255 = 2.0f32; + let background_rho = 0.1f32; + let texture_sigma_255 = 20.0f32; + + let mut records = Vec::new(); + for _ in 0..6 { + let record = zero_mean_block_record(background_sigma_255, background_rho); + records.extend_from_slice(&record); + } + + for _ in 0..10 { + let record = zero_mean_block_record(texture_sigma_255, 0.0); + records.extend_from_slice(&record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("the static minority clears STATIC_FRACTION_MIN on its own"); + + let expected_sigma = background_sigma_255 / 255.0; + assert!( + (sample.sigma[0] - expected_sigma).abs() < 1e-4, + "expected the outlier gate to isolate the real noise floor {expected_sigma}, got {}", + sample.sigma[0] + ); + assert!( + (sample.static_fraction - 6.0 / 16.0).abs() < 1e-4, + "expected only the 6 background blocks to survive both gates, got static_fraction={}", + sample.static_fraction + ); + assert!( + (sample.rho - background_rho).abs() < 1e-4, + "expected rho to come from the surviving background blocks only, got {}", + sample.rho + ); +} + +/// Every block is static, with a threefold real noise spread like a dark region beside a bright one. +#[test] +fn aggregate_keeps_genuinely_static_blocks_despite_spatial_sigma_spread() { + let width = 64; + let height = 64; + let stored_channels = 1; + let channels = 1; + + let low_sigma_255 = 2.0f32; + let high_sigma_255 = 6.0f32; // 3x low_sigma_255, real spatial spread. + let rho = 0.1f32; + + let mut records = Vec::new(); + for _ in 0..8 { + let record = zero_mean_block_record(low_sigma_255, rho); + records.extend_from_slice(&record); + } + + for _ in 0..8 { + let record = zero_mean_block_record(high_sigma_255, rho); + records.extend_from_slice(&record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("every block is static"); + + assert!( + (sample.static_fraction - 1.0).abs() < 1e-6, + "a real 3x spatial sigma spread must not trip the outlier gate, got static_fraction={}", + sample.static_fraction + ); +} + +#[test] +fn aggregate_outlier_factor_survives_just_under_threshold() { + let width = 48; + let height = 48; + let stored_channels = 1; + let channels = 1; + let background_sigma_255 = 2.0f32; + let background_rho = 0.1f32; + + let records = outlier_factor_boundary_records( + background_sigma_255, + background_rho, + OUTLIER_FACTOR_SURVIVES_RATIO, + ); + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("all 9 blocks clear the static-fraction floor"); + + assert!( + (sample.static_fraction - 1.0).abs() < 1e-6, + "a 9th block at {OUTLIER_FACTOR_SURVIVES_RATIO}x the reference must survive, \ + got static_fraction={}", + sample.static_fraction + ); +} + +#[test] +fn aggregate_outlier_factor_rejects_just_over_threshold() { + let width = 48; + let height = 48; + let stored_channels = 1; + let channels = 1; + let background_sigma_255 = 2.0f32; + let background_rho = 0.1f32; + + let records = + outlier_factor_boundary_records(background_sigma_255, background_rho, OUTLIER_FACTOR_REJECTS_RATIO); + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("the 8 background blocks alone still clear the static-fraction floor"); + + assert!( + (sample.static_fraction - 8.0 / 9.0).abs() < 1e-6, + "a 9th block at {OUTLIER_FACTOR_REJECTS_RATIO}x the reference must be rejected, \ + got static_fraction={}", + sample.static_fraction + ); +} + +/// 26 of 100 blocks are zero letterbox bars and 74 carry static noise across five close levels. +/// +/// The spread stands in for real sampling variance, which lets the population pass the anchor +/// check. The 74 split as 15, 15, 15, 15 and 14, so the median lands exactly on the third level. +#[test] +fn aggregate_returns_correct_sigma_with_letterbox_zero_population() { + let width = 160; + let height = 160; + let stored_channels = 1; + let channels = 1; + + let background_rho = 0.2f32; + let sigma_levels_255 = [3.8f32, 3.9, 4.0, 4.1, 4.2]; + + let mut records = Vec::new(); + for _ in 0..26 { + let record = zero_mean_block_record(0.0, 0.0); + records.extend_from_slice(&record); + } + + for i in 0..74 { + let sigma_255 = sigma_levels_255[i % sigma_levels_255.len()]; + let record = zero_mean_block_record(sigma_255, background_rho); + records.extend_from_slice(&record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("74 of 100 blocks carry real static noise, far above the 5% floor"); + + let expected_sigma = 4.0 / 255.0; + assert!( + (sample.sigma[0] - expected_sigma).abs() < 1e-4, + "expected the letterbox bars to leave the real noise floor near {expected_sigma} intact, got {}", + sample.sigma[0] + ); + assert!( + (sample.rho - background_rho).abs() < 1e-4, + "expected rho to come from the real-noise blocks, got {}", + sample.rho + ); +} + +/// 26 zero letterbox blocks beside 74 texture blocks at one repeated sigma, which set and pass +/// their own ceiling. +#[test] +fn aggregate_returns_none_when_the_only_above_gate_population_is_texture() { + let width = 160; + let height = 160; + let stored_channels = 1; + let channels = 1; + + let texture_sigma_255 = 20.0f32; + + let mut records = Vec::new(); + for _ in 0..26 { + let record = zero_mean_block_record(0.0, 0.0); + records.extend_from_slice(&record); + } + + for _ in 0..74 { + let record = zero_mean_block_record(texture_sigma_255, 0.0); + records.extend_from_slice(&record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height); + assert!( + sample.is_none(), + "a homogeneous above-gate population with no genuine low anchor must fall back to \ + None rather than report the texture level as sigma" + ); +} + +#[test] +fn aggregate_rejects_texture_outliers_with_zero_population_present() { + let width = 64; + let height = 80; // 4 x 5 TEMPORAL_NOISE_BLOCK grid, 20 blocks. + let stored_channels = 1; + let channels = 1; + + let background_sigma_255 = 2.0f32; + let background_rho = 0.1f32; + let texture_sigma_255 = 20.0f32; + + let mut records = Vec::new(); + for _ in 0..6 { + let record = zero_mean_block_record(background_sigma_255, background_rho); + records.extend_from_slice(&record); + } + + for _ in 0..10 { + let record = zero_mean_block_record(texture_sigma_255, 0.0); + records.extend_from_slice(&record); + } + + for _ in 0..4 { + let record = zero_mean_block_record(0.0, 0.0); + records.extend_from_slice(&record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("the static minority clears STATIC_FRACTION_MIN on its own"); + + let expected_sigma = background_sigma_255 / 255.0; + assert!( + (sample.sigma[0] - expected_sigma).abs() < 1e-4, + "expected the outlier gate to isolate the real noise floor {expected_sigma} despite \ + the zero population, got {}", + sample.sigma[0] + ); + assert!( + (sample.static_fraction - 6.0 / 20.0).abs() < 1e-4, + "expected only the 6 background blocks to survive, the zero blocks carry no \ + measurable noise, got static_fraction={}", + sample.static_fraction + ); + assert!( + (sample.rho - background_rho).abs() < 1e-4, + "expected rho to come from the surviving background blocks only, got {}", + sample.rho + ); +} + +#[test] +fn aggregate_keeps_static_spread_with_zero_population_present() { + let width = 64; + let height = 80; // 4 x 5 TEMPORAL_NOISE_BLOCK grid, 20 blocks. + let stored_channels = 1; + let channels = 1; + + let low_sigma_255 = 2.0f32; + let high_sigma_255 = 6.0f32; // 3x low_sigma_255, real spatial spread. + let rho = 0.1f32; + + let mut records = Vec::new(); + for _ in 0..8 { + let record = zero_mean_block_record(low_sigma_255, rho); + records.extend_from_slice(&record); + } + + for _ in 0..8 { + let record = zero_mean_block_record(high_sigma_255, rho); + records.extend_from_slice(&record); + } + + for _ in 0..4 { + let record = zero_mean_block_record(0.0, 0.0); + records.extend_from_slice(&record); + } + + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("every non-zero block is static"); + + assert!( + (sample.static_fraction - 0.8).abs() < 1e-6, + "a real 3x spatial sigma spread plus a zero population must not trip the outlier \ + gate, and the 4 zero blocks carry no measurable noise, got static_fraction={}", + sample.static_fraction + ); +} + +/// A block's rho can read above 1, because the lag-1 total averages over adjacent pairs while the +/// variance averages over every pixel. Each row here follows the pattern that maximises that ratio. +#[test] +fn aggregate_rho_estimate_stays_within_unit_range() { + let width = TEMPORAL_NOISE_BLOCK; + let height = TEMPORAL_NOISE_BLOCK; + let stored_channels = 1; + let channels = 1; + + let scale = 0.007f32; + let row: Vec = (1..=width) + .map(|i| scale * (i as f32 * std::f32::consts::PI / (width + 1) as f32).sin()) + .collect(); + + let pixel_count = (width * height) as f32; + let sum_d: f32 = row.iter().sum::() * height as f32; + let sum_d2: f32 = row.iter().map(|value| value * value).sum::() * height as f32; + let sum_lag: f32 = row.windows(2).map(|pair| pair[0] * pair[1]).sum::() * height as f32; + + // Both gates must pass so the assertion below tests the clamp rather than a gate miss. + assert!( + (sum_d / pixel_count).abs() < STATIC_GATE, + "construction must clear the static gate" + ); + + let variance = sum_d2 / pixel_count - (sum_d / pixel_count) * (sum_d / pixel_count); + assert!( + (variance.sqrt() / std::f32::consts::SQRT_2) > RHO_SIGMA_GATE, + "construction must clear the rho-sample sigma gate" + ); + + let records = scalar_record(sum_d, sum_d2, sum_lag); + let sample = aggregate_temporal_noise_stats(&records, channels, stored_channels, width, height) + .expect("the single block clears both gates"); + + assert!( + (0.0..=1.0).contains(&sample.rho), + "the mismatched-denominator estimate must stay within [0, 1], got {}", + sample.rho + ); +} diff --git a/av-denoise-core/src/nlmeans/options.rs b/av-denoise-core/src/nlmeans/options.rs new file mode 100644 index 0000000..c926542 --- /dev/null +++ b/av-denoise-core/src/nlmeans/options.rs @@ -0,0 +1,137 @@ +use super::{ChannelMode, HqParams, MotionCompensationMode, NlmParams, PrefilterMode, hq_default_strength}; +use crate::options::Preset; + +/// Settings for the fast [NlmeansAlgorithm] path. +#[derive(Debug, Copy, Clone, Default, PartialEq)] +pub struct NlmeansOptions { + /// Which reference image the NLM weights are computed against. + /// + /// `None`, the default, compares patches on the noisy input. Every other mode costs one extra + /// GPU pass per frame. + pub prefilter: PrefilterMode, + /// Whether temporal denoising follows motion between frames. + /// + /// `None`, the default, turns it off. It only has an effect when `mode` is `Temporal { .. }`. + pub motion_compensation: MotionCompensationMode, + /// Overrides for the NLM search radius, patch radius, strength and self-weight. + pub tuning: NlmTuning, + /// Whether each frame is cleaned on its own or across a temporal window. + pub mode: DenoisingMode, +} + +/// Settings for the HQ [NlmeansAlgorithm] path. +#[derive(Debug, Copy, Clone, Default, PartialEq)] +pub struct NlmeansHqOptions { + /// Everything the fast path takes, which HQ takes too. + pub nlm: NlmeansOptions, + /// The noise measurement and confidence weighting HQ adds on top. + pub hq: HqParams, +} + +/// Which nlmeans implementation a preset, or an explicit choice, selects. +#[derive(Debug, Copy, Clone, PartialEq, Eq, strum_macros::EnumString)] +#[strum(ascii_case_insensitive)] +pub enum NlmeansVariant { + /// The fast path, with fixed weighting and no noise measurement. + Fast, + /// The quality path, which calibrates its weighting to the noise measured in each frame. + Hq, +} + +/// Which nlmeans variant to build, with its settings. +#[derive(Debug, Copy, Clone, PartialEq)] +pub enum NlmeansAlgorithm { + Fast(NlmeansOptions), + Hq(NlmeansHqOptions), +} + +impl NlmeansAlgorithm { + pub(crate) fn mode(&self) -> DenoisingMode { + match self { + NlmeansAlgorithm::Fast(options) => options.mode, + NlmeansAlgorithm::Hq(options) => options.nlm.mode, + } + } +} + +/// Which [NlmeansVariant] a preset runs. +pub fn nlmeans_variant_for(preset: Preset) -> NlmeansVariant { + match preset { + Preset::Veryfast => NlmeansVariant::Fast, + Preset::Fast | Preset::Base | Preset::Slow | Preset::Veryslow => NlmeansVariant::Hq, + } +} + +/// How many neighbouring frames on each side `nlmeans` looks at, at a preset. +pub fn nlmeans_temporal_radius_for(preset: Preset) -> u32 { + match preset { + Preset::Veryfast => 0, + Preset::Fast => 1, + Preset::Base => 2, + Preset::Slow => 4, + Preset::Veryslow => 8, + } +} + +/// How far `nlmeans` looks for similar patches inside a frame, at a preset. +pub fn nlmeans_search_radius_for(preset: Preset) -> u32 { + match preset { + Preset::Veryfast | Preset::Fast | Preset::Base => 2, + Preset::Slow | Preset::Veryslow => 4, + } +} + +/// Whether a frame is cleaned on its own or alongside its neighbours. +#[derive(Debug, Copy, Clone, Default, Eq, PartialEq)] +pub enum DenoisingMode { + /// Cleans each frame using only its own pixels. + #[default] + Spacial, + /// Cleans each frame using a window of `2 * radius + 1` frames. + Temporal { radius: u32 }, +} + +/// Optional NLM tuning overrides, each falling back to the library default when unset. +#[derive(Debug, Copy, Clone, Default, PartialEq)] +pub struct NlmTuning { + pub search_radius: Option, + pub patch_radius: Option, + pub strength: Option, + pub self_weight: Option, +} + +/// Resolves the options into the parameters a denoiser runs with. +/// +/// Calibrated defaults are filled in for `channels`, and an explicit `strength` always wins. +pub(crate) fn resolve_params(algorithm: &NlmeansAlgorithm, channels: ChannelMode) -> NlmParams { + let temporal_radius = match algorithm.mode() { + DenoisingMode::Spacial => 0, + DenoisingMode::Temporal { radius } => radius, + }; + + let (options, hq) = match *algorithm { + NlmeansAlgorithm::Fast(options) => (options, None), + NlmeansAlgorithm::Hq(options) => (options.nlm, Some(options.hq)), + }; + + let defaults = NlmParams::default(); + // With `auto_strength` on, HQ reads `strength` as a multiplier on the measured noise, so it + // needs its own calibrated default. + let default_strength = match hq { + Some(hq) if hq.auto_strength => hq_default_strength(channels, temporal_radius), + _ => defaults.strength, + }; + let strength = options.tuning.strength.unwrap_or(default_strength); + + NlmParams { + channels, + prefilter: options.prefilter, + motion_compensation: options.motion_compensation, + temporal_radius, + hq, + strength, + search_radius: options.tuning.search_radius.unwrap_or(defaults.search_radius), + patch_radius: options.tuning.patch_radius.unwrap_or(defaults.patch_radius), + self_weight: options.tuning.self_weight.unwrap_or(defaults.self_weight), + } +} diff --git a/av-denoise-core/src/nlmeans/params.rs b/av-denoise-core/src/nlmeans/params.rs index 35460ce..8cff601 100644 --- a/av-denoise-core/src/nlmeans/params.rs +++ b/av-denoise-core/src/nlmeans/params.rs @@ -1,87 +1,45 @@ use super::{MotionCompensationMode, PrefilterMode, prefilter}; -/// The reference value patch distances are normalised against, matching -/// FFmpeg's nlmeans at 255 squared. +/// The value patch distances are normalised against, 255 squared as in FFmpeg's nlmeans. /// -/// Distances here are measured in `[0, 1]` units, so this constant folds -/// in the scale-up back to 8-bit terms. +/// Distances are measured between 0 and 1, so this folds the scale back up to 8-bit terms. pub(super) const NLM_NORM: f32 = 255.0 * 255.0; -/// A scaling factor inherited from FFmpeg's nlmeans, kept so our -/// `strength` parameter means the same thing theirs does. +/// A scale factor from FFmpeg's nlmeans, kept so `strength` means the same as theirs. pub(super) const NLM_LEGACY: f32 = 3.0; -/// The measured HQ default `strength` for luma at each temporal radius, -/// indexed by `temporal_radius.min(8)`. -/// -/// `hq_default_strength` reads this for `ChannelMode::Luma` and -/// `ChannelMode::Yuv`. +/// The measured HQ default luma `strength`, indexed by `temporal_radius.min(8)`. const HQ_DEFAULT_STRENGTH_LUMA: [f32; 9] = [0.45, 0.45, 0.42, 0.42, 0.35, 0.35, 0.35, 0.30, 0.30]; -/// The measured HQ default `strength` for chroma at each temporal -/// radius, indexed by `temporal_radius.min(8)`. -/// -/// `hq_default_strength` reads this for `ChannelMode::Chroma`. +/// The measured HQ default chroma `strength`, indexed by `temporal_radius.min(8)`. const HQ_DEFAULT_STRENGTH_CHROMA: [f32; 9] = [1.00, 0.85, 0.70, 0.70, 0.70, 0.70, 0.70, 0.70, 0.70]; -/// The calibrated default `strength` multiplier for `nlmeans-hq`'s -/// auto-strength mode. -/// -/// The answer depends on which plane is being denoised and how far the -/// temporal window reaches. -/// -/// # Where the numbers come from -/// -/// Every entry is a measured value rather than a fitted curve. Both -/// tables come from quality-harness sweeps that score a grid of -/// strengths at three noise levels per radius. -/// -/// The luma sweep covers each radius from 0 to 8 directly. -/// -/// The chroma sweep pins luma at the value already chosen for that -/// radius, so the chroma numbers stay clean, and covers radii 0, 1, 2, -/// 4, and 8 with a bracketed peak at each. -/// -/// At each measured radius the chosen value is the one whose worst XPSNR -/// gain across the tested noise levels is highest, so it holds up at -/// whichever noise level is hardest to serve. -/// -/// # Shape of the tables -/// -/// Luma never rises with radius, because a wider temporal window already -/// gathers more samples to average over. -/// -/// Chroma falls the same way out to radius 2 and then holds flat at -/// 0.70. Radii 3, 5, 6, and 7 sit on that measured plateau rather than -/// being swept directly. +/// The calibrated default `strength` multiplier for HQ auto-strength. /// -/// `ChannelMode::Yuv` reads the luma table, on the assumption that a -/// fused pass is dominated by luma. That mode was not part of the sweep, -/// so this is an assumption rather than a measurement. +/// Every entry is measured, from quality sweeps over three noise levels per radius that keep the +/// strength whose worst XPSNR gain is highest. Luma is swept at every radius from 0 to 8 and never +/// rises with radius, because a wider window already gathers more samples. Chroma is swept at +/// radii 0, 1, 2, 4 and 8 with luma pinned, and holds flat at 0.70 from radius 2, so radii 3, 5, 6 +/// and 7 sit on that plateau. /// -/// # Clamping -/// -/// `temporal_radius` is clamped to the last table index, which matches -/// [`MAX_TEMPORAL_RADIUS`]. That is only a safety net, because -/// [`NlmParams::validate`] already rejects anything larger. +/// `ChannelMode::Yuv` reads the luma table, which is an assumption because the sweeps do not +/// cover that mode. Clamping the radius to [MAX_TEMPORAL_RADIUS] is a safety net, since +/// [NlmParams::validate] already rejects anything larger. pub fn hq_default_strength(channels: ChannelMode, temporal_radius: u32) -> f32 { - let idx = temporal_radius.min(MAX_TEMPORAL_RADIUS) as usize; + let index = temporal_radius.min(MAX_TEMPORAL_RADIUS) as usize; match channels { - ChannelMode::Luma | ChannelMode::Yuv => HQ_DEFAULT_STRENGTH_LUMA[idx], - ChannelMode::Chroma => HQ_DEFAULT_STRENGTH_CHROMA[idx], + ChannelMode::Luma | ChannelMode::Yuv => HQ_DEFAULT_STRENGTH_LUMA[index], + ChannelMode::Chroma => HQ_DEFAULT_STRENGTH_CHROMA[index], } } -/// The smallest frame side length the denoiser supports. +/// The smallest supported frame side. /// -/// The Immerkær noise estimate only reads interior pixels, because its -/// 3x3 mask cannot reach the one-pixel border. A frame under 3 pixels -/// across has no interior at all, which leaves the estimate undefined. +/// The Immerkær estimate's 3x3 mask only reads interior pixels, and a frame under 3 pixels across +/// has none. pub const MIN_FRAME_DIM: u32 = 3; /// Rejects frame dimensions the kernels cannot handle. -/// -/// Both denoiser constructors call this before allocating any buffer. pub fn validate_dimensions(width: u32, height: u32) -> Result<(), anyhow::Error> { if width < MIN_FRAME_DIM || height < MIN_FRAME_DIM { anyhow::bail!( @@ -90,62 +48,44 @@ pub fn validate_dimensions(width: u32, height: u32) -> Result<(), anyhow::Error> least one interior pixel" ); } + Ok(()) } -/// The patch radius above which the dispatcher switches to the separable -/// path, so per-pixel cost stays linear in `patch_radius`. +/// The patch radius above which the separable path runs, keeping per-pixel cost linear in +/// `patch_radius`. pub(super) const SEPARABLE_THRESHOLD: u32 = 8; /// The hard ceiling on `patch_radius`. /// -/// The fused kernels load a `(block + 2 * patch_radius)^2` tile into -/// shared memory, and anything larger runs out of it on RDNA-class GPUs. +/// The fused kernels load a `(block + 2 * patch_radius)^2` tile into shared memory, and anything +/// larger runs out of it on RDNA-class GPUs. pub const MAX_PATCH_RADIUS: u32 = 16; /// The hard ceiling on `search_radius`. /// -/// Shared memory is not the limit here. The windowed kernel's tile is -/// `(block + 2 * patch_radius + 2 * search_radius)^2 * stored_ch * 4` -/// bytes, which stays comfortably inside hardware limits at every -/// supported size. -/// -/// The real cost is the kernel's fully unrolled -/// `(2 * search_radius + 1)^2` window loop. Both its compiled size and -/// how long it takes to generate grow with the radius. See the -/// stack-size note in `.cargo/config.toml`. -/// -/// The per-offset dispatch path, used when `patch_radius` forces the -/// separable fallback, is limited by the same constant, which keeps its -/// `(2 * search_radius + 1)^2` launches per temporal offset reasonable. +/// The windowed kernel's tile fits in shared memory at every supported size. The limit is its fully +/// unrolled `(2 * search_radius + 1)^2` window loop, whose compiled size and build time grow with +/// the radius. It also bounds the separable path's `(2 * search_radius + 1)^2` launches per +/// temporal offset. pub const MAX_SEARCH_RADIUS: u32 = 8; /// The hard ceiling on `temporal_radius`. /// -/// The ring buffer holds `2 * radius + 1` frames, so device memory grows -/// with it. At radius 16, 1080p YUV would need roughly 540 MB for the -/// input alone. +/// The ring holds `2 * radius + 1` frames, so device memory grows with it. At radius 16, 1080p YUV +/// would need roughly 540 MB for the input alone. pub const MAX_TEMPORAL_RADIUS: u32 = 8; -/// The hard ceiling on the radius the bilateral prefilter derives from -/// `sigma_s`, where `radius = ceil(2 * sigma_s).max(1)`. See -/// `prefilter::bilateral_radius`. -/// -/// `nlm_bilateral` loads a `(32 + 2r) x (8 + 2r)` tile of -/// `Vector` into shared memory. `N` reaches 4 for YUV storage at -/// 4 bytes per `f32`, so the tile costs `16 * (32 + 2r) * (8 + 2r)` -/// bytes. +/// The hard ceiling on the bilateral prefilter's radius, `ceil(2 * sigma_s).max(1)`. /// -/// RDNA-class hardware gives 64 KiB of shared memory to work with. At -/// `r = 22` the tile is 63,232 bytes, or 61.75 KiB, which fits. At -/// `r = 23` it is 67,392 bytes, or 65.8 KiB, which does not. -/// -/// So 22 is the largest radius that fits, and `bilateral_radius` reaches -/// it at `sigma_s = 11.0`. +/// `nlm_bilateral` loads a `(32 + 2r) x (8 + 2r)` tile of up to four `f32` lanes into shared +/// memory, which costs `16 * (32 + 2r) * (8 + 2r)` bytes. Against the 64 KiB RDNA-class hardware +/// offers, `r = 22` fits at 63,232 bytes and `r = 23` does not at 67,392. `bilateral_radius` +/// reaches 22 at `sigma_s = 11.0`. pub const MAX_BILATERAL_RADIUS: u32 = 22; -#[derive(Debug, Clone, Copy, PartialEq, Eq)] /// Which channels of a frame the denoiser works on. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ChannelMode { /// The single brightness channel, with distances scaled by 3.0. Luma, @@ -167,9 +107,8 @@ impl ChannelMode { /// How many channels each pixel occupies in GPU storage. /// - /// This is padded up to the next supported vector width so kernels - /// can read whole `Line` values at once. Backends only support - /// power-of-two widths, so YUV pads from 3 up to 4. + /// Kernels read whole `Line` values and backends only support power-of-two widths, so YUV + /// pads from 3 to 4. pub fn storage_count(self) -> u32 { match self { ChannelMode::Luma => 1, @@ -181,65 +120,41 @@ impl ChannelMode { /// Parameters for the quality-focused `nlmeans-hq` variant. /// -/// The measured noise level drives both the effective strength and the -/// distance floor, so the weighting adapts to how noisy the source -/// really is. +/// The measured noise level drives both the effective strength and the distance floor. #[derive(Debug, Clone, Copy, PartialEq)] pub struct HqParams { - /// Reads `strength` as a multiplier on the noise level, making the - /// effective FFmpeg-style strength `strength * sigma_eff * 255`. + /// Reads `strength` as a multiplier on the noise level. Defaults to true. /// - /// Defaults to true. + /// The FFmpeg-style strength becomes `strength * sigma_eff * 255`. pub auto_strength: bool, - /// Subtracts the expected noise floor from patch distances before - /// weighting, so a match is not penalised for the noise it carries. + /// Subtracts the expected noise floor from patch distances before weighting. Defaults to true. /// - /// Defaults to true. + /// A match is then not penalised for the noise it carries. pub noise_floor: bool, - /// A fixed noise standard deviation in `[0, 1]` units, replacing the - /// automatic per-frame estimate. + /// A fixed noise sigma between 0 and 1 in place of the automatic estimate. /// - /// `None`, the default, measures the noise in each pushed frame and - /// smooths it over time. `Some` applies one fixed value to every - /// frame instead. - /// - /// The CLI takes this in 8-bit units through `--hq-sigma` and - /// divides by 255. + /// `None`, the default, measures each pushed frame and smooths the result over time. pub sigma_override: Option, - /// Weights each temporal neighbour by how well it block-matches the - /// centre frame, so occlusion or a change of content collapses its - /// contribution rather than blurring it in. - /// - /// Only has an effect when `temporal_radius` is above 0. Defaults to + /// Weights each temporal neighbour by how well it block-matches the centre frame. Defaults to /// true. + /// + /// Occlusion or a change of content then collapses a neighbour's contribution rather than + /// blurring it in. It only has an effect when `temporal_radius` is above 0. pub temporal_confidence: bool, - /// A multiplier on the per-pixel mismatch threshold, which sets how - /// much extra SAD a block tolerates before its confidence starts to - /// fall. + /// A multiplier on the per-pixel mismatch threshold. Defaults to 1.0. /// - /// Higher values tolerate larger mismatches. Defaults to 1.0. + /// Higher values let a block carry more extra SAD before its confidence starts to fall. pub thsad_scale: f32, - /// A multiplier applied to each channel's measured sigma before it - /// folds into the running estimate. - /// - /// `1.0`, the default, keeps the measurement as it is. This does - /// nothing when `sigma_override` is set, because the estimator never - /// runs in that case. + /// A multiplier on each channel's measured sigma before it joins the running estimate. Defaults + /// to 1.0. /// - /// The CLI takes this as `--hq-sigma-scale`. + /// It has no effect with `sigma_override` set, because the estimator never runs. pub sigma_scale: f32, - /// Estimates noise fresh from each submit's own window, instead of - /// smoothing it with an exponential average carried forward from - /// every earlier frame the stream has folded. + /// Estimates noise from the current window alone rather than an EMA over every earlier frame. /// - /// `false`, the default, keeps the temporal EMA every calibrated - /// preset assumes. `true` makes the automatic estimate depend only - /// on the frames currently in the window, so a - /// [`crate::frame::PlanarDenoiser::reseed`] targeting frame `n` and - /// a continuous stream that reaches frame `n` compute the same - /// sigma, regardless of how each got there. This does nothing when - /// `sigma_override` pins a fixed value, because the estimator never - /// runs in that case. + /// `false`, the default, keeps the EMA every calibrated preset assumes. `true` makes a reseed + /// at frame `n` and a continuous stream reaching frame `n` compute the same sigma. It has no + /// effect with `sigma_override` set. pub windowed_noise_estimation: bool, } @@ -258,8 +173,7 @@ impl Default for HqParams { } impl HqParams { - /// The HQ defaults with a fixed noise level in `[0, 1]` units, which - /// skips automatic estimation. + /// The HQ defaults with a fixed noise sigma between 0 and 1, which skips automatic estimation. pub fn with_sigma(sigma: f32) -> Self { Self { sigma_override: Some(sigma), @@ -273,39 +187,28 @@ impl HqParams { pub struct NlmParams { /// How many frames on each side of the current one to look at. /// - /// 0 means each frame is cleaned on its own. Anything higher uses a - /// window of `2 * radius + 1` frames. + /// 0 cleans each frame on its own. pub temporal_radius: u32, - /// Half the width of the search window, which covers - /// `(2 * radius + 1)^2` pixels. Defaults to 2. + /// Half the width of the search window. Defaults to 2. pub search_radius: u32, - /// Half the width of a compared patch, which covers - /// `(2 * radius + 1)^2` pixels. Defaults to 4. + /// Half the width of a compared patch. Defaults to 4. pub patch_radius: u32, /// How hard to filter. Higher values smooth more. Defaults to 1.2. pub strength: f32, - /// How much weight the centre pixel gets in the average. + /// How much weight the centre pixel gets in the average. Defaults to 1.0. /// - /// Defaults to 1.0. Set it to 0 for pure NLM, where the centre pixel - /// only counts through the patches that match it. + /// 0 gives pure NLM, where the centre pixel only counts through the patches that match it. pub self_weight: f32, - /// Which channels to process. pub channels: ChannelMode, - /// Which image the patch distances are measured against. + /// The image patch distances are measured on. Defaults to `None`. /// - /// Defaults to `None`. When set, the weights come from a prefiltered - /// or externally supplied image while the pixels being averaged - /// still come from the original input. + /// The averaged pixels always come from the original input. pub prefilter: PrefilterMode, - /// Whether temporal denoising follows motion between frames. - /// - /// Defaults to `None`. Set to `Mvtools`, each submit warps the - /// temporal neighbours into line with the centre frame before the - /// NLM weighting runs. + /// Whether temporal denoising follows motion between frames. Defaults to `None`. /// - /// Only has an effect when `temporal_radius` is above 0. + /// It only has an effect when `temporal_radius` is above 0. pub motion_compensation: MotionCompensationMode, - /// The quality-mode parameters. `None` runs the fast path unchanged. + /// The quality-mode parameters, or `None` for the fast path. pub hq: Option, } @@ -326,13 +229,10 @@ impl Default for NlmParams { } impl NlmParams { - /// The FFmpeg-style strength the weighting actually uses. + /// The FFmpeg-style strength the weighting uses. /// - /// With HQ auto-strength the user's value multiplies the noise - /// level, so one setting follows sources of different noisiness. - /// - /// `sigma_eff` is the scale-weighted RMS of the per-channel noise - /// estimates. + /// With HQ auto-strength, `strength` multiplies `sigma_eff`, so one setting follows sources of + /// different noisiness. pub(super) fn effective_strength_with(&self, sigma_eff: Option) -> f32 { match (self.hq, sigma_eff) { (Some(hq), Some(sigma)) if hq.auto_strength => self.strength * sigma * 255.0, @@ -340,57 +240,43 @@ impl NlmParams { } } - /// `h2_inv_norm` for a noise estimate given here, ignoring whatever - /// `self.hq.sigma_override` holds. - /// - /// The denoiser calls this each submit to refresh the value from a - /// freshly measured sigma. + /// `h2_inv_norm` for the given noise estimate, ignoring `sigma_override`. pub fn h2_inv_norm_with(&self, sigma_eff: Option) -> f32 { - let s_size = (2 * self.patch_radius + 1) * (2 * self.patch_radius + 1); - let s = self.effective_strength_with(sigma_eff); - NLM_NORM / (NLM_LEGACY * s * s * s_size as f32) + let patch_area = (2 * self.patch_radius + 1) * (2 * self.patch_radius + 1); + let strength = self.effective_strength_with(sigma_eff); + NLM_NORM / (NLM_LEGACY * strength * strength * patch_area as f32) } - /// `h2_inv_norm` using `self.hq.sigma_override` as the noise - /// estimate, or no estimate at all on the fast path. - /// - /// HQ denoisers that estimate noise automatically call - /// [`Self::h2_inv_norm_with`] each submit instead. + /// `h2_inv_norm` from `sigma_override`, or with no estimate on the fast path. pub fn h2_inv_norm(&self) -> f32 { - self.h2_inv_norm_with(self.hq.and_then(|hq| hq.sigma_override)) + let sigma_override = self.hq.and_then(|hq| hq.sigma_override); + self.h2_inv_norm_with(sigma_override) } - /// The patch distance two noisy copies of the same content are - /// expected to show, for a given set of per-channel sigmas. + /// The patch distance two noisy copies of the same content are expected to show. /// - /// Each active channel contributes `2 * channel_scale * sigma^2`, - /// summed over all `(2 * patch_radius + 1)^2` taps. - /// - /// Returns 0 when the HQ noise floor is off, or when there is no - /// estimate to apply. + /// Each active channel adds `2 * channel_scale * sigma^2` per tap over all + /// `(2 * patch_radius + 1)^2` taps. It is 0 with the HQ noise floor off or no estimate. pub(super) fn noise_offset_with(&self, sigmas: Option<&[f32]>) -> f32 { match (self.hq, sigmas) { (Some(hq), Some(sigmas)) if hq.noise_floor => { - let s_size = (2 * self.patch_radius + 1) * (2 * self.patch_radius + 1); + let patch_area = (2 * self.patch_radius + 1) * (2 * self.patch_radius + 1); let scale = channel_scale(self.channels); let count = self.channels.count() as usize; - let sum_sq: f32 = sigmas.iter().take(count).map(|&s| s * s).sum(); - 2.0 * scale * sum_sq * s_size as f32 + let sum_sq: f32 = sigmas.iter().take(count).map(|&sigma| sigma * sigma).sum(); + 2.0 * scale * sum_sq * patch_area as f32 }, _ => 0.0, } } - /// `noise_offset` with `self.hq.sigma_override` applied to every - /// active channel, or no estimate at all on the fast path. - /// - /// HQ denoisers that estimate noise automatically call - /// [`Self::noise_offset_with`] each submit instead. + /// `noise_offset` with `sigma_override` on every active channel, or 0 without one. pub(super) fn noise_offset(&self) -> f32 { match self.hq.and_then(|hq| hq.sigma_override) { Some(sigma) => { let sigmas = [sigma; 3]; - self.noise_offset_with(Some(&sigmas[..self.channels.count() as usize])) + let active = &sigmas[..self.channels.count() as usize]; + self.noise_offset_with(Some(active)) }, None => 0.0, } @@ -400,12 +286,8 @@ impl NlmParams { 1 + 2 * self.temporal_radius } - /// Rejects parameter combinations that would fail to launch, by - /// running the kernels past their shared-memory or register limits, - /// or that would produce meaningless output. - /// - /// `NlmDenoiser::new` calls this for you. Callers building params by - /// hand can call it directly to see errors before construction. + /// Rejects parameters that would push a kernel past its shared-memory or register limits, or + /// produce meaningless output. pub fn validate(&self) -> Result<(), anyhow::Error> { if self.patch_radius > MAX_PATCH_RADIUS { anyhow::bail!( @@ -489,6 +371,7 @@ impl NlmParams { sigma_s, ); } + if !sigma_r.is_finite() || sigma_r <= 0.0 { anyhow::bail!( "bilateral prefilter sigma_r must be finite and greater than 0, got \ @@ -497,16 +380,9 @@ impl NlmParams { sigma_r, ); } - // sigma_s decides the shared-memory tile radius through - // `prefilter::bilateral_radius`. Checking that derived - // radius, rather than working out an equivalent sigma_s - // threshold here, keeps this in step if the formula ever - // changes. - // - // A very large sigma_s can also overflow the radius inside - // the tile-size expression. This check catches that too, - // because an overflowed radius always lands far past the - // maximum. + + // Checking the derived radius keeps this in step with its formula, and also catches a + // sigma_s large enough to overflow the tile-size expression. let bilateral_radius = prefilter::bilateral_radius(sigma_s); if bilateral_radius > MAX_BILATERAL_RADIUS { anyhow::bail!( @@ -519,16 +395,11 @@ impl NlmParams { MAX_BILATERAL_RADIUS, ); } - // A sigma can be finite and positive yet small enough that - // `sigma * sigma` underflows to 0.0 in f32, which happens - // below roughly 3.8e-20. That makes the reciprocal - // normalisation factor the kernel uses infinite. - // - // Checking the same derived factor `run_bilateral` computes - // for the launch catches this wherever the underflow - // threshold actually falls, without picking a sigma cutoff - // by hand or repeating the expression here. - if !prefilter::inv_two_sigma_sq(sigma_s).is_finite() { + + // A positive sigma below roughly 3.8e-20 squares to 0.0 in f32, so the launch's own + // factor is checked rather than a hand-picked sigma cutoff. + let inv_two_sigma_s_sq = prefilter::inv_two_sigma_sq(sigma_s); + if !inv_two_sigma_s_sq.is_finite() { anyhow::bail!( "bilateral prefilter sigma_s is too small, got {}. Squaring it \ underflows to 0 in f32, which makes the spatial-weight \ @@ -536,7 +407,9 @@ impl NlmParams { sigma_s, ); } - if !prefilter::inv_two_sigma_sq(sigma_r).is_finite() { + + let inv_two_sigma_r_sq = prefilter::inv_two_sigma_sq(sigma_r); + if !inv_two_sigma_r_sq.is_finite() { anyhow::bail!( "bilateral prefilter sigma_r is too small, got {}. Squaring it \ underflows to 0 in f32, which makes the range-weight normalisation \ @@ -553,6 +426,7 @@ impl NlmParams { strength_scale, ); } + if self.patch_radius > SEPARABLE_THRESHOLD { anyhow::bail!( "the nlm pilot uses the windowed spatial kernel, which supports \ @@ -569,11 +443,7 @@ impl NlmParams { } } -/// The per-channel distance scale for a channel mode, which is 3 for -/// luma, 1.5 for chroma, and 1 for full YUV. -/// -/// This matches the `channel_scale` the weighting kernels use on the -/// GPU, and it is the same for every channel within a given mode. +/// The per-channel distance scale, matching the `channel_scale` the weighting kernels use. pub(super) fn channel_scale(channels: ChannelMode) -> f32 { match channels { ChannelMode::Luma => 3.0, @@ -582,605 +452,12 @@ pub(super) fn channel_scale(channels: ChannelMode) -> f32 { } } -/// The scale-weighted RMS of the per-channel noise estimates, over the -/// channels a mode actually uses. -/// -/// Because `channel_scale` is the same for every channel in a mode, the -/// weighting cancels out and this is really just a plain RMS. +/// The scale-weighted RMS of the active channels' noise estimates. /// -/// Anything in `sigmas` past the mode's channel count is ignored. +/// The scale is the same for every channel in a mode, so this is a plain RMS. Entries past the +/// mode's channel count are ignored. pub(super) fn sigma_eff(sigmas: &[f32], channels: ChannelMode) -> f32 { let count = channels.count() as usize; - let sum_sq: f32 = sigmas.iter().take(count).map(|&s| s * s).sum(); + let sum_sq: f32 = sigmas.iter().take(count).map(|&sigma| sigma * sigma).sum(); (sum_sq / count as f32).sqrt() } - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn noise_offset_scales_with_sigma_and_patch_size() { - let sigma = 4.0 / 255.0; - let params = NlmParams { - patch_radius: 4, - hq: Some(HqParams::with_sigma(sigma)), - ..NlmParams::default() - }; - - let expected = 6.0 * sigma * sigma * 81.0; - assert!( - (params.noise_offset() - expected).abs() < 1e-6, - "expected {expected}, got {}", - params.noise_offset() - ); - } - - #[test] - fn noise_offset_zero_without_noise_floor() { - let params = NlmParams { - hq: Some(HqParams { - auto_strength: true, - noise_floor: false, - sigma_override: Some(4.0 / 255.0), - temporal_confidence: true, - thsad_scale: 1.0, - sigma_scale: 1.0, - windowed_noise_estimation: false, - }), - ..NlmParams::default() - }; - - assert_eq!(params.noise_offset(), 0.0); - } - - #[test] - fn noise_offset_zero_without_hq() { - let params = NlmParams::default(); - assert_eq!(params.noise_offset(), 0.0); - } - - #[test] - fn h2_inv_norm_with_auto_strength_matches_hand_computed() { - let sigma = 8.0 / 255.0; - let params = NlmParams { - strength: 1.0, - hq: Some(HqParams::with_sigma(sigma)), - ..NlmParams::default() - }; - - let s_size = (2 * params.patch_radius + 1) * (2 * params.patch_radius + 1); - let effective_strength = 1.0 * sigma * 255.0; - let expected = NLM_NORM / (NLM_LEGACY * effective_strength * effective_strength * s_size as f32); - - assert!( - (params.h2_inv_norm() - expected).abs() < 1e-6, - "expected {expected}, got {}", - params.h2_inv_norm() - ); - } - - #[test] - fn validate_rejects_zero_hq_sigma() { - let params = NlmParams { - hq: Some(HqParams::with_sigma(0.0)), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_rejects_hq_sigma_above_one() { - let params = NlmParams { - hq: Some(HqParams::with_sigma(1.5)), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_rejects_nan_hq_sigma() { - let params = NlmParams { - hq: Some(HqParams::with_sigma(f32::NAN)), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_rejects_zero_thsad_scale() { - let params = NlmParams { - hq: Some(HqParams { - thsad_scale: 0.0, - ..HqParams::default() - }), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_rejects_negative_thsad_scale() { - let params = NlmParams { - hq: Some(HqParams { - thsad_scale: -1.0, - ..HqParams::default() - }), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_rejects_nan_thsad_scale() { - let params = NlmParams { - hq: Some(HqParams { - thsad_scale: f32::NAN, - ..HqParams::default() - }), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_accepts_default_thsad_scale() { - let params = NlmParams { - hq: Some(HqParams::default()), - ..NlmParams::default() - }; - assert!(params.validate().is_ok()); - } - - #[test] - fn hq_params_default_sigma_scale_is_one() { - assert_eq!(HqParams::default().sigma_scale, 1.0); - } - - #[test] - fn validate_rejects_sigma_scale_below_the_minimum() { - let params = NlmParams { - hq: Some(HqParams { - sigma_scale: 0.05, - ..HqParams::default() - }), - ..NlmParams::default() - }; - let err = params.validate().expect_err("0.05 is below the 0.1 minimum"); - assert!( - err.to_string().contains("hq sigma_scale"), - "error should name the field, got {err}" - ); - } - - #[test] - fn validate_rejects_sigma_scale_above_the_maximum() { - let params = NlmParams { - hq: Some(HqParams { - sigma_scale: 10.5, - ..HqParams::default() - }), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_rejects_nan_sigma_scale() { - let params = NlmParams { - hq: Some(HqParams { - sigma_scale: f32::NAN, - ..HqParams::default() - }), - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_accepts_sigma_scale_at_the_bounds() { - let low = NlmParams { - hq: Some(HqParams { - sigma_scale: 0.1, - ..HqParams::default() - }), - ..NlmParams::default() - }; - assert!(low.validate().is_ok()); - - let high = NlmParams { - hq: Some(HqParams { - sigma_scale: 10.0, - ..HqParams::default() - }), - ..NlmParams::default() - }; - assert!(high.validate().is_ok()); - } - - #[test] - fn noise_offset_with_handles_distinct_per_channel_sigmas() { - let sigma_u = 4.0 / 255.0; - let sigma_v = 10.0 / 255.0; - let params = NlmParams { - patch_radius: 4, - channels: ChannelMode::Chroma, - hq: Some(HqParams { - auto_strength: true, - noise_floor: true, - sigma_override: None, - temporal_confidence: true, - thsad_scale: 1.0, - sigma_scale: 1.0, - windowed_noise_estimation: false, - }), - ..NlmParams::default() - }; - - let s_size = (2 * params.patch_radius + 1) * (2 * params.patch_radius + 1); - // The chroma scale of 1.5 applies per channel, and each channel - // keeps its own sigma rather than sharing one. - let expected = 2.0 * 1.5 * (sigma_u * sigma_u + sigma_v * sigma_v) * s_size as f32; - - let got = params.noise_offset_with(Some(&[sigma_u, sigma_v])); - assert!((got - expected).abs() < 1e-9, "expected {expected}, got {got}"); - } - - #[test] - fn sigma_eff_is_rms_over_active_channels() { - let sigmas = [3.0 / 255.0, 4.0 / 255.0]; - let got = sigma_eff(&sigmas, ChannelMode::Chroma); - let expected = ((sigmas[0] * sigmas[0] + sigmas[1] * sigmas[1]) / 2.0).sqrt(); - assert!((got - expected).abs() < 1e-9, "expected {expected}, got {got}"); - } - - #[test] - fn validate_rejects_non_positive_pilot_strength_scale() { - let zero = NlmParams { - prefilter: PrefilterMode::NlmSpatial { strength_scale: 0.0 }, - ..NlmParams::default() - }; - assert!(zero.validate().is_err()); - - let nan = NlmParams { - prefilter: PrefilterMode::NlmSpatial { - strength_scale: f32::NAN, - }, - ..NlmParams::default() - }; - assert!(nan.validate().is_err()); - } - - #[test] - fn validate_rejects_pilot_with_patch_radius_above_separable_threshold() { - let params = NlmParams { - prefilter: PrefilterMode::NlmSpatial { strength_scale: 1.0 }, - patch_radius: SEPARABLE_THRESHOLD + 1, - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn validate_accepts_pilot_within_limits() { - let params = NlmParams { - prefilter: PrefilterMode::NlmSpatial { strength_scale: 1.0 }, - patch_radius: SEPARABLE_THRESHOLD, - ..NlmParams::default() - }; - assert!(params.validate().is_ok()); - } - - #[test] - fn validate_rejects_non_positive_bilateral_sigma_r() { - // A sigma_r of 0 makes inv_two_sigma_r_sq infinite. The centre - // tap's range_sq is 0, and 0 times infinity is NaN, which - // poisons every pixel of the reference image. - let params = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: 0.0, - }, - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - - let negative = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: -0.02, - }, - ..NlmParams::default() - }; - assert!(negative.validate().is_err()); - - let nan = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: f32::NAN, - }, - ..NlmParams::default() - }; - assert!(nan.validate().is_err()); - - let inf = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: f32::INFINITY, - }, - ..NlmParams::default() - }; - assert!(inf.validate().is_err()); - } - - #[test] - fn validate_rejects_non_positive_bilateral_sigma_s() { - // A sigma_s of 0 makes inv_two_sigma_s_sq infinite. The centre - // tap's spatial_dist_sq is 0, so the spatial term is poisoned by - // the same 0 times infinity NaN. - let params = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 0.0, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - - let negative = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: -3.0, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(negative.validate().is_err()); - - let nan = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: f32::NAN, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(nan.validate().is_err()); - - let inf = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: f32::INFINITY, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(inf.validate().is_err()); - } - - #[test] - fn validate_accepts_positive_finite_bilateral_sigmas() { - let params = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(params.validate().is_ok()); - } - - #[test] - fn validate_accepts_a_small_positive_bilateral_sigma_at_the_boundary() { - // Pins the guard to `<= 0.0` rather than `< 0.0`. A value that - // is small but strictly positive, and far enough from the f32 - // underflow cliff that squaring it stays a normal float, has to - // be accepted for either field on its own. - // - // 1e-6 squares to 1e-12, nowhere near the smallest normal f32 of - // about 1.18e-38, so `inv_two_sigma_sq` stays finite here. - let safe_small = 1e-6_f32; - assert!( - (safe_small * safe_small).is_normal(), - "the test value itself must not underflow" - ); - - let small_sigma_s = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: safe_small, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(small_sigma_s.validate().is_ok()); - - let small_sigma_r = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: safe_small, - }, - ..NlmParams::default() - }; - assert!(small_sigma_r.validate().is_ok()); - } - - #[test] - fn validate_rejects_a_subnormal_bilateral_sigma_that_underflows_on_squaring() { - // `f32::MIN_POSITIVE`, the smallest normal positive f32 at about - // 1.1754944e-38, is finite and above 0, so a guard that only - // checked the raw value let it through. - // - // Squaring it underflows to exactly 0.0 in f32, because its true - // square of about 1.38e-76 is far below the smallest subnormal - // of about 1.4e-45. `inv_two_sigma_sq` then divides by zero and - // returns infinity. - // - // That is the same NaN poisoning this validation exists to - // prevent, reached through a sigma other than exactly 0.0. - let sq = f32::MIN_POSITIVE * f32::MIN_POSITIVE; - assert_eq!(sq, 0.0, "this test assumes MIN_POSITIVE underflows on squaring"); - let inv = prefilter::inv_two_sigma_sq(f32::MIN_POSITIVE); - assert!( - !inv.is_finite(), - "this test assumes the derived factor is infinite here" - ); - - let sigma_s = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: f32::MIN_POSITIVE, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!( - sigma_s.validate().is_err(), - "a subnormal sigma_s that underflows to an infinite normalisation factor must be rejected" - ); - - let sigma_r = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 3.0, - sigma_r: f32::MIN_POSITIVE, - }, - ..NlmParams::default() - }; - assert!( - sigma_r.validate().is_err(), - "a subnormal sigma_r that underflows to an infinite normalisation factor must be rejected" - ); - } - - /// A `sigma_s` of 16.0 gives a bilateral radius of 32, which is the - /// worked example in `prefilter.rs` and well past the maximum of 22. - #[test] - fn validate_rejects_bilateral_sigma_s_above_the_smem_ceiling() { - let params = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 16.0, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - let err = params.validate().expect_err("radius 32 exceeds the 22 ceiling"); - assert!( - err.to_string().contains("sigma_s"), - "error should name the field, got {err}" - ); - } - - /// A `sigma_s` of 1e9 overflows the tile-size arithmetic if it ever - /// reaches the kernel launch. - /// - /// Validation has to reject it long before that, through the same - /// radius check any other oversized `sigma_s` hits. - #[test] - fn validate_rejects_extreme_bilateral_sigma_s() { - let params = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 1e9, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - /// The boundary pair for [`MAX_BILATERAL_RADIUS`], written as - /// literal `sigma_s` values rather than derived from the constant. - /// - /// A `sigma_s` of 11.0 gives a radius of 22, right at the ceiling, - /// so it is accepted. A `sigma_s` of 11.01 gives 23, one past it, so - /// it is rejected. - #[test] - fn validate_accepts_bilateral_sigma_s_at_the_smem_ceiling() { - let params = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 11.0, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(params.validate().is_ok()); - } - - #[test] - fn validate_rejects_bilateral_sigma_s_just_above_the_smem_ceiling() { - let params = NlmParams { - prefilter: PrefilterMode::Bilateral { - sigma_s: 11.01, - sigma_r: 0.02, - }, - ..NlmParams::default() - }; - assert!(params.validate().is_err()); - } - - #[test] - fn sigma_eff_ignores_channels_past_the_mode_count() { - // Luma only reads the first element, even when handed extra - // chroma samples. - let sigmas = [6.0 / 255.0, 100.0 / 255.0, 200.0 / 255.0]; - let got = sigma_eff(&sigmas, ChannelMode::Luma); - assert!( - (got - sigmas[0]).abs() < 1e-9, - "expected {}, got {got}", - sigmas[0] - ); - } - - #[test] - fn hq_default_strength_matches_the_measured_luma_table() { - const EXPECTED: [f32; 9] = [0.45, 0.45, 0.42, 0.42, 0.35, 0.35, 0.35, 0.30, 0.30]; - for (radius, &expected) in EXPECTED.iter().enumerate() { - let got = hq_default_strength(ChannelMode::Luma, radius as u32); - assert!( - (got - expected).abs() < f32::EPSILON, - "at radius {radius} expected {expected}, got {got}" - ); - } - } - - #[test] - fn hq_default_strength_matches_the_measured_chroma_table() { - const EXPECTED: [f32; 9] = [1.00, 0.85, 0.70, 0.70, 0.70, 0.70, 0.70, 0.70, 0.70]; - for (radius, &expected) in EXPECTED.iter().enumerate() { - let got = hq_default_strength(ChannelMode::Chroma, radius as u32); - assert!( - (got - expected).abs() < f32::EPSILON, - "at radius {radius} expected {expected}, got {got}" - ); - } - } - - #[test] - fn hq_default_strength_yuv_reads_the_luma_table() { - for radius in 0..=8u32 { - let yuv = hq_default_strength(ChannelMode::Yuv, radius); - let luma = hq_default_strength(ChannelMode::Luma, radius); - assert!( - (yuv - luma).abs() < f32::EPSILON, - "at radius {radius} yuv is {yuv} but luma is {luma}" - ); - } - } - - #[test] - fn validate_dimensions_rejects_frames_below_the_minimum() { - assert!(validate_dimensions(2, 64).is_err()); - assert!(validate_dimensions(64, 2).is_err()); - assert!(validate_dimensions(0, 0).is_err()); - } - - #[test] - fn validate_dimensions_accepts_the_minimum() { - assert!(validate_dimensions(MIN_FRAME_DIM, MIN_FRAME_DIM).is_ok()); - assert!(validate_dimensions(1920, 1080).is_ok()); - } - - #[test] - fn hq_default_strength_clamps_radius_above_the_table() { - let at_max = hq_default_strength(ChannelMode::Luma, MAX_TEMPORAL_RADIUS); - let above_max = hq_default_strength(ChannelMode::Luma, MAX_TEMPORAL_RADIUS + 5); - assert!( - (at_max - above_max).abs() < f32::EPSILON, - "expected clamping to hold the last table entry, got {at_max} vs {above_max}" - ); - } -} diff --git a/av-denoise-core/src/nlmeans/pending.rs b/av-denoise-core/src/nlmeans/pending.rs deleted file mode 100644 index 467643c..0000000 --- a/av-denoise-core/src/nlmeans/pending.rs +++ /dev/null @@ -1,363 +0,0 @@ -use std::future::Future; -use std::marker::PhantomData; -use std::pin::Pin; -use std::task::{Context, Poll, Waker}; - -use cubecl::bytes::Bytes; -use cubecl::client::ComputeClient; -use cubecl::prelude::*; -use cubecl::server::{Handle, ServerError}; - -use super::kernels::gpu_pack_wire; -use super::{BLOCK_1D, Depth, MAX_GRID_1D}; -use crate::denoiser::{FrameOutput, OutputFormat}; - -pub(crate) type ReadFuture = Pin, ServerError>> + Send>>; - -/// A denoise that is still in flight. -/// -/// The kernels are queued, the GPU may still be working on them, and the -/// readback to the host has not finished. -/// -/// The readback only starts on the first poll, so a `Pending` that was -/// never polled costs nothing to drop. One that [`Self::try_wait`] has -/// polled owns a mapped staging buffer on the wgpu backends until it -/// lands, so dropping it blocks until the readback finishes. See the -/// `Drop` impl. -pub struct Pending { - pub(super) fut: ReadFuture, - /// Set once `fut` has been polled and cleared again once it has - /// produced its result. See the `Drop` impl for why this matters. - polled: bool, - pub(super) channels: u32, - pub(super) stored_ch: u32, - pub(super) pixels: usize, - pub(super) format: OutputFormat, - pub(super) _marker: PhantomData, -} - -impl Pending { - /// Wraps an in-flight readback future into a `Pending`. - /// - /// `fut` is the future returned by reading back a single output - /// buffer. `channels` is how many channels the caller wants out, - /// `stored_ch` is how many are actually laid out per pixel in that - /// buffer, and `pixels` is the frame's pixel count. `format` is what - /// the buffer holds, which is `f32` samples with padding lanes for - /// [`OutputFormat::F32`] and already-packed wire bytes for - /// [`OutputFormat::Wire`]. - /// - /// In `F32` mode `wait` and `wait_into` strip the padding lanes - /// between `channels` and `stored_ch`. In `Wire` mode the pack kernel - /// has already done that on the GPU. - pub(crate) fn new( - fut: ReadFuture, - channels: u32, - stored_ch: u32, - pixels: usize, - format: OutputFormat, - ) -> Self { - Self { - fut, - polled: false, - channels, - stored_ch, - pixels, - format, - _marker: PhantomData, - } - } - - /// Blocks until the readback finishes and returns the denoised frame - /// in a fresh buffer of this `Pending`'s output format. - /// - /// In `F32` mode the YUV padding lanes are stripped, so the buffer - /// holds exactly `pixels * channels` values. - pub fn wait(self) -> Result { - let mut out = empty_output(self.pixels, self.channels, self.format); - self.wait_into(&mut out)?; - Ok(out) - } - - /// Blocks until the readback finishes and writes the result into `dst`, - /// which is cleared first. - /// - /// `dst` keeps its allocation when it already holds this `Pending`'s - /// output format, so a caller can reuse one buffer when running frame - /// after frame. - pub fn wait_into(mut self, dst: &mut FrameOutput) -> Result<(), anyhow::Error> { - let (pixels, channels, stored_ch, format) = (self.pixels, self.channels, self.stored_ch, self.format); - let result = cubecl::future::block_on(self.fut.as_mut()); - // The future has produced its result either way, so there is - // nothing left for `Drop` to settle. - self.polled = false; - let bytes = result?.remove(0); - unpack_into(&bytes, pixels, channels, stored_ch, format, dst); - Ok(()) - } - - /// Polls the readback once. - /// - /// `TryWait::NotReady` hands the same `Pending` back, so a caller - /// that gets it can only poll again by calling `try_wait` on that - /// returned value. There is no way to poll a future that has already - /// produced its result. The poll uses a no-op waker, so nothing ever - /// wakes a caller when the readback lands. A caller that wants the - /// frame has to keep calling `try_wait` again itself, on whatever - /// `NotReady` the previous call returned. - /// - /// This only avoids blocking on the wgpu backends, meaning Vulkan and Metal, - /// where readiness is external state a discarded wakeup does not lose. - /// - /// On CUDA and ROCm the readback future's first poll runs a blocking driver wait - /// internally, so `try_wait` blocks for the full kernel and readback latency there - /// and `NotReady` is never actually returned. - /// - /// The first poll starts the readback, and from then on the - /// `Pending` is committed to it. Dropping a `NotReady` instead of - /// polling it again blocks until the readback lands. See the `Drop` - /// impl. - pub fn try_wait(mut self) -> Result, anyhow::Error> { - let waker = Waker::noop(); - let mut cx = Context::from_waker(waker); - self.polled = true; - let poll = self.fut.as_mut().poll(&mut cx); - if poll.is_ready() { - self.polled = false; - } - match poll { - Poll::Ready(Ok(mut bytes)) => { - let bytes = bytes.remove(0); - let mut out = empty_output(self.pixels, self.channels, self.format); - unpack_into( - &bytes, - self.pixels, - self.channels, - self.stored_ch, - self.format, - &mut out, - ); - Ok(TryWait::Ready(out)) - }, - Poll::Ready(Err(e)) => Err(e.into()), - Poll::Pending => Ok(TryWait::NotReady(self)), - } - } -} - -impl Drop for Pending { - /// Settles a readback that was polled but never landed. - /// - /// On the wgpu backends the first poll reserves a staging buffer, - /// records the copy into it and maps it. The buffer is only unmapped - /// when the `Bytes` the future resolves to is dropped. Dropping the - /// future before that hands the still-mapped buffer back to the - /// device's staging pool, and the next submit that touches it fails - /// with "buffer is still mapped" on cubecl's device thread, which - /// takes every later call on that device down with it. - /// - /// Blocking on the future here resolves it to its `Bytes`, which are - /// dropped straight away and unmap the buffer. An unpolled future - /// has reserved nothing, so it is left alone. A future that already - /// produced its result cannot be polled again, so it is skipped too. - fn drop(&mut self) { - if !self.polled || std::thread::panicking() { - return; - } - let _ = cubecl::future::block_on(self.fut.as_mut()); - } -} - -/// An empty buffer of `format`, sized for one frame. -pub(super) fn empty_output(pixels: usize, channels: u32, format: OutputFormat) -> FrameOutput { - let samples = pixels * channels as usize; - match format { - OutputFormat::F32 => FrameOutput::F32(Vec::with_capacity(samples)), - OutputFormat::Wire { depth } => { - FrameOutput::Wire(Vec::with_capacity(samples * depth.bytes_per_sample())) - }, - } -} - -/// Turns a raw readback buffer into `dst`. -/// -/// This is the shared step behind `wait_into` and `try_wait`, so the -/// blocking and non-blocking paths cannot drift apart. A `dst` that -/// already holds the right variant keeps its allocation, and one that -/// does not is replaced. -fn unpack_into( - bytes: &Bytes, - pixels: usize, - channels: u32, - stored_ch: u32, - format: OutputFormat, - dst: &mut FrameOutput, -) { - let samples = pixels * channels as usize; - match (format, dst) { - (OutputFormat::F32, FrameOutput::F32(out)) => { - unpack_bytes_into(bytes, pixels, channels, stored_ch, out); - }, - (OutputFormat::Wire { depth }, FrameOutput::Wire(out)) => { - unpack_wire_into(bytes, samples * depth.bytes_per_sample(), out); - }, - (OutputFormat::F32, dst) => { - let mut out = Vec::new(); - unpack_bytes_into(bytes, pixels, channels, stored_ch, &mut out); - *dst = FrameOutput::F32(out); - }, - (OutputFormat::Wire { depth }, dst) => { - let mut out = Vec::new(); - unpack_wire_into(bytes, samples * depth.bytes_per_sample(), &mut out); - *dst = FrameOutput::Wire(out); - }, - } -} - -/// The outcome of polling a [`Pending`] without blocking. -pub enum TryWait { - /// The readback landed. This is the denoised frame. - Ready(FrameOutput), - /// The readback has not landed. This is the same `Pending`, still in - /// flight and now committed to finishing. Dropping it blocks until - /// the readback lands rather than abandoning it. - NotReady(Pending), -} - -/// Starts the readback of a denoised frame sitting in `handle`, in -/// whichever format the caller asked for. -/// -/// In `Wire` mode this first queues [`gpu_pack_wire`] against -/// `wire_dst` and reads that back instead, so the quantised bytes cross -/// the bus rather than four times as many `f32`s. `wire_dst` is the -/// caller's own word buffer for the slot `handle` came from. -/// -/// No sync sits between the pack launch and the read. `read_async` -/// submits its copy descriptors on the same stream as the launches ahead -/// of it, so it is already the sync point and the pack is just one more -/// queued operation. -/// -/// Both algorithms build their `Pending` through here, so their readback -/// paths cannot drift apart. -pub(crate) fn start_readback( - client: &ComputeClient, - handle: Handle, - wire_dst: Option<&Handle>, - channels: u32, - stored_ch: u32, - pixels: usize, - format: OutputFormat, -) -> Pending { - let handle = match (format, wire_dst) { - (OutputFormat::F32, _) => handle, - (OutputFormat::Wire { depth }, Some(dst)) => { - pack_wire(client, &handle, dst, channels, stored_ch, pixels, depth); - dst.clone() - }, - (OutputFormat::Wire { .. }, None) => { - unreachable!("a wire-mode denoiser allocates its wire buffers at construction") - }, - }; - - // The future is wrapped in an `async move` that owns a cloned - // `ComputeClient`, which is cheap because it shares its internals. - // That owned client lives inside the future, so the future is - // genuinely `'static` and the `Pending` can outlive the denoiser - // without any lifetime tricks. - let client = client.clone(); - let fut = Box::pin(async move { client.read_async(vec![handle]).await }); - - Pending::new(fut, channels, stored_ch, pixels, format) -} - -/// Queues [`gpu_pack_wire`] over `src`, writing the packed words into `dst`. -fn pack_wire( - client: &ComputeClient, - src: &Handle, - dst: &Handle, - channels: u32, - stored_ch: u32, - pixels: usize, - depth: Depth, -) { - let pack = depth.wire_pack(); - let samples = pixels as u32 * channels; - let words = samples.div_ceil(pack.samples_per_word()); - - let split_planes = wire_splits_planes(channels); - let outer = if split_planes { pixels as u32 } else { channels }; - - let grid = words.div_ceil(BLOCK_1D).clamp(1, MAX_GRID_1D); - let total_threads = grid * BLOCK_1D; - - unsafe { - gpu_pack_wire::launch_unchecked::( - client, - CubeCount::new_1d(grid), - CubeDim::new_1d(BLOCK_1D), - ArrayArg::from_raw_parts(src.clone(), pixels * stored_ch as usize), - ArrayArg::from_raw_parts(dst.clone(), words as usize), - pack.max(), - pixels as u32, - channels, - stored_ch, - outer, - split_planes, - pack.samples_per_word(), - words, - total_threads, - ); - } -} - -/// Whether a frame with this many channels goes out as one contiguous -/// region per channel rather than interleaved. -/// -/// Only the chroma pair splits. A luma frame has nothing to split, and a -/// packed YUV frame stays interleaved because that is the layout its -/// consumer already reads. -pub(crate) fn wire_splits_planes(channels: u32) -> bool { - channels == 2 -} - -/// Copies the first `len` bytes of a readback buffer into `dst`, which is -/// cleared first. -/// -/// The pack kernel writes whole `u32` words, so the buffer can run up to -/// three bytes past the frame. Everything the caller wants is already -/// quantised and free of padding lanes. -fn unpack_wire_into(bytes: &Bytes, len: usize, dst: &mut Vec) { - dst.clear(); - dst.extend_from_slice(&bytes[..len]); -} - -/// Unpacks a raw readback buffer straight into `dst`, stripping the padding lanes -/// between `channels` and `stored_ch`. -fn unpack_bytes_into(bytes: &Bytes, pixels: usize, channels: u32, stored_ch: u32, dst: &mut Vec) { - let data = f32::from_bytes(bytes); - unpack_frame(data, pixels, channels as usize, stored_ch as usize, dst); -} - -/// Copies `channels` values out of every pixel in `data` into `dst`, -/// skipping any padding lanes `stored_ch` added beyond that. -/// -/// `data` holds `pixels * stored_ch` values. When `channels` and -/// `stored_ch` are equal this is a plain copy. `dst` is cleared first. -pub(super) fn unpack_frame( - data: &[f32], - pixels: usize, - channels: usize, - stored_ch: usize, - dst: &mut Vec, -) { - dst.clear(); - if channels == stored_ch { - dst.extend_from_slice(data); - } else { - dst.reserve(pixels * channels); - for pixel in 0..pixels { - let src = pixel * stored_ch; - dst.extend_from_slice(&data[src..src + channels]); - } - } -} diff --git a/av-denoise-core/src/nlmeans/prefilter/bilateral.rs b/av-denoise-core/src/nlmeans/prefilter/bilateral.rs index 70efd89..2508018 100644 --- a/av-denoise-core/src/nlmeans/prefilter/bilateral.rs +++ b/av-denoise-core/src/nlmeans/prefilter/bilateral.rs @@ -4,28 +4,19 @@ use super::PrefilterCtx; use crate::nlmeans::kernels::nlm_bilateral; use crate::nlmeans::{BLOCK_X, BLOCK_Y}; -/// The kernel radius derived from `sigma_s`. +/// The kernel radius for `sigma_s`. /// -/// Stopping at two sigma covers over 95% of the Gaussian's mass, and it -/// keeps shared memory and register use bounded. +/// Two sigma covers over 95% of the Gaussian's mass while keeping shared memory and register use +/// bounded. pub fn bilateral_radius(sigma_s: f32) -> u32 { ((2.0 * sigma_s).ceil() as u32).max(1) } -/// The normalisation factor `1 / (2 * sigma^2)` for the bilateral -/// Gaussian. +/// The bilateral Gaussian's normalisation factor `1 / (2 * sigma^2)`. /// -/// The spatial and range terms share this because they use the same -/// Gaussian shape. -/// -/// It is computed on the host and passed to the kernel as a plain `f32`, -/// so this is the only place the value is derived from a sigma. The -/// kernel only ever multiplies by it. -/// -/// Squaring a very small but positive sigma can underflow to 0.0 in -/// `f32`, which makes this factor infinite even though the sigma itself -/// was finite and positive. Code validating user input should check this -/// value rather than only the sign and finiteness of the sigma. +/// The spatial and range terms share it, and the kernel only multiplies by it. A tiny positive +/// sigma can square to 0.0 in `f32` and make it infinite, so validation checks this value rather +/// than the sigma alone. pub(crate) fn inv_two_sigma_sq(sigma: f32) -> f32 { 1.0 / (2.0 * sigma * sigma) } @@ -43,10 +34,13 @@ pub(super) fn run_bilateral( let inv_two_sigma_s_sq = inv_two_sigma_sq(sigma_s); let inv_two_sigma_r_sq = inv_two_sigma_sq(sigma_r); + let cubes_x = ctx.width.div_ceil(BLOCK_X); + let cubes_y = ctx.height.div_ceil(BLOCK_Y); + unsafe { nlm_bilateral::launch_unchecked::( client, - CubeCount::new_2d(ctx.width.div_ceil(BLOCK_X), ctx.height.div_ceil(BLOCK_Y)), + CubeCount::new_2d(cubes_x, cubes_y), CubeDim::new_2d(BLOCK_X, BLOCK_Y), stored_ch, ArrayArg::from_raw_parts(ctx.input_buf.clone(), total), diff --git a/av-denoise-core/src/nlmeans/prefilter/mod.rs b/av-denoise-core/src/nlmeans/prefilter/mod.rs index 226a492..b2db904 100644 --- a/av-denoise-core/src/nlmeans/prefilter/mod.rs +++ b/av-denoise-core/src/nlmeans/prefilter/mod.rs @@ -1,45 +1,34 @@ mod bilateral; -pub use bilateral::bilateral_radius; -pub(crate) use bilateral::inv_two_sigma_sq; use cubecl::prelude::*; use cubecl::server::Handle; -/// How the reference image for each frame is produced. -/// -/// NLM compares patches to decide how much two pixels look alike. Doing -/// that on a noisy image means comparing noise as well as content, so a -/// cleaner reference image can give better weights. +pub use self::bilateral::bilateral_radius; +pub(crate) use self::bilateral::inv_two_sigma_sq; + +/// How each frame's reference image is produced. /// -/// The pixels being averaged always come from the original input. Only -/// the weights change. +/// Comparing patches on a noisy image compares the noise too, so a cleaner reference gives better +/// weights. The averaged pixels always come from the original input. #[non_exhaustive] #[derive(Debug, Default, Clone, Copy, PartialEq)] pub enum PrefilterMode { - /// No reference image, so patches are compared on the noisy input. - /// This costs nothing extra. + /// No reference image, so patches are compared on the noisy input at no extra cost. #[default] None, - /// The caller supplies the reference frame through - /// [`super::NlmDenoiser::push_frame_with_reference`]. - External, /// A quick bilateral blur run on the GPU at push time. Bilateral { sigma_s: f32, sigma_r: f32 }, - /// A spatial NLM pilot pass. - /// - /// Each frame is denoised with the windowed spatial kernel at push - /// time and the result is kept as the reference image. + /// A spatial NLM pilot pass at push time, kept as the reference image. NlmSpatial { /// How much of the main pass strength the pilot pass uses. strength_scale: f32, }, } -/// The measured default strength for the pilot pass, as a multiplier on -/// the main pass strength. +/// The default pilot strength, as a multiplier on the main pass strength. /// -/// A calibration sweep across noise levels put the XPSNR plateau for -/// `PrefilterMode::NlmSpatial` at this value. +/// A calibration sweep across noise levels puts the XPSNR plateau for `PrefilterMode::NlmSpatial` at +/// this value. pub const DEFAULT_PILOT_STRENGTH_SCALE: f32 = 0.4; impl PrefilterMode { @@ -48,39 +37,29 @@ impl PrefilterMode { !matches!(self, Self::None) } - /// Whether this mode builds its reference on the GPU during - /// `push_frame`, rather than taking one from the caller. + /// Whether this mode builds its reference on the GPU during a push. pub(crate) fn is_gpu_internal(self) -> bool { matches!(self, Self::Bilateral { .. } | Self::NlmSpatial { .. }) } } -/// Parses a `--prefilter`-style string into a [`PrefilterMode`]. -/// -/// `"none"` or an empty string means [`PrefilterMode::None`]. -/// -/// `"nlm"` or `"nlm:"` builds -/// [`PrefilterMode::NlmSpatial`]. `strength_scale` multiplies the main -/// pass strength for the pilot pass. Bare `"nlm"` uses -/// [`DEFAULT_PILOT_STRENGTH_SCALE`]. +/// Parses a `--prefilter`-style string into a [PrefilterMode]. /// -/// `"bilateral:,"` builds [`PrefilterMode::Bilateral`]. -/// -/// This never produces [`PrefilterMode::External`], since that mode has -/// no string form: it requires the caller to supply a reference frame -/// through [`super::NlmDenoiser::push_frame_with_reference`]. -pub fn parse_prefilter(s: &str) -> Result { - if s == "none" || s.is_empty() { +/// `none` or an empty string gives [PrefilterMode::None]. `nlm` or `nlm:` gives +/// [PrefilterMode::NlmSpatial], with bare `nlm` at [DEFAULT_PILOT_STRENGTH_SCALE]. +/// `bilateral:,` gives [PrefilterMode::Bilateral]. +pub fn parse_prefilter(value: &str) -> Result { + if value == "none" || value.is_empty() { return Ok(PrefilterMode::None); } - if s == "nlm" { + if value == "nlm" { return Ok(PrefilterMode::NlmSpatial { strength_scale: DEFAULT_PILOT_STRENGTH_SCALE, }); } - if let Some(rest) = s.strip_prefix("nlm:") { + if let Some(rest) = value.strip_prefix("nlm:") { let strength_scale: f32 = rest .trim() .parse() @@ -89,7 +68,7 @@ pub fn parse_prefilter(s: &str) -> Result { return Ok(PrefilterMode::NlmSpatial { strength_scale }); } - if let Some(rest) = s.strip_prefix("bilateral:") { + if let Some(rest) = value.strip_prefix("bilateral:") { let parts: Vec<&str> = rest.split(',').collect(); if parts.len() != 2 { @@ -103,14 +82,13 @@ pub fn parse_prefilter(s: &str) -> Result { } anyhow::bail!( - "unknown prefilter '{s}', expected `none`, `nlm[:]`, or `bilateral:,`" + "unknown prefilter '{value}', expected `none`, `nlm[:]`, or `bilateral:,`" ) } /// The inputs one prefilter dispatch needs. /// -/// This lives only for the length of a single `push_frame`, which is -/// what makes the borrows on the denoiser's buffers sound. +/// It lives for one push, which makes the borrows on the denoiser's buffers sound. pub(crate) struct PrefilterCtx<'a> { pub width: u32, pub height: u32, @@ -122,20 +100,16 @@ pub(crate) struct PrefilterCtx<'a> { pub reference_buf: &'a Handle, } -/// Runs the GPU prefilter for the frame that was uploaded last. -/// -/// `None` and `External` do nothing here. +/// Runs the GPU prefilter for the frame uploaded last. pub(crate) fn run_prefilter( mode: PrefilterMode, client: &ComputeClient, ctx: &PrefilterCtx<'_>, ) -> Result<(), anyhow::Error> { match mode { - PrefilterMode::None | PrefilterMode::External => Ok(()), - // The pilot needs the full accumulator context, meaning accum, - // weight_sum, max_weight, and h2_inv_norm, which `PrefilterCtx` - // does not carry. `NlmDenoiser::run_nlm_spatial_pilot` - // dispatches it directly instead of coming through here. + PrefilterMode::None => Ok(()), + // The pilot needs the accumulators and `h2_inv_norm`, which this context lacks, so the + // denoiser dispatches it directly. PrefilterMode::NlmSpatial { .. } => Ok(()), PrefilterMode::Bilateral { sigma_s, sigma_r } => { bilateral::run_bilateral::(client, ctx, sigma_s, sigma_r) @@ -153,43 +127,45 @@ mod tests { assert!(!PrefilterMode::None.is_gpu_internal()); } - #[test] - fn external_needs_buffer_but_not_gpu() { - assert!(PrefilterMode::External.needs_reference_buf()); - assert!(!PrefilterMode::External.is_gpu_internal()); - } - #[test] fn bilateral_is_gpu_internal() { - let m = PrefilterMode::Bilateral { + let mode = PrefilterMode::Bilateral { sigma_s: 3.0, sigma_r: 0.02, }; - assert!(m.needs_reference_buf()); - assert!(m.is_gpu_internal()); + assert!(mode.needs_reference_buf()); + assert!(mode.is_gpu_internal()); } #[test] fn nlm_spatial_is_gpu_internal() { - let m = PrefilterMode::NlmSpatial { strength_scale: 1.0 }; + let mode = PrefilterMode::NlmSpatial { strength_scale: 1.0 }; - assert!(m.needs_reference_buf()); - assert!(m.is_gpu_internal()); + assert!(mode.needs_reference_buf()); + assert!(mode.is_gpu_internal()); } #[test] fn bilateral_radius_truncates_at_two_sigma() { - assert_eq!(bilateral_radius(0.1), 1); - assert_eq!(bilateral_radius(1.0), 2); - assert_eq!(bilateral_radius(3.0), 6); - assert_eq!(bilateral_radius(3.5), 7); + let tiny = bilateral_radius(0.1); + let unit = bilateral_radius(1.0); + let wide = bilateral_radius(3.0); + let fractional = bilateral_radius(3.5); + + assert_eq!(tiny, 1); + assert_eq!(unit, 2); + assert_eq!(wide, 6); + assert_eq!(fractional, 7); } #[test] fn none_and_empty_prefilters_parse() { - assert!(matches!(parse_prefilter("none").unwrap(), PrefilterMode::None)); - assert!(matches!(parse_prefilter("").unwrap(), PrefilterMode::None)); + let none = parse_prefilter("none").unwrap(); + let empty = parse_prefilter("").unwrap(); + + assert!(matches!(none, PrefilterMode::None)); + assert!(matches!(empty, PrefilterMode::None)); } #[test] @@ -225,31 +201,13 @@ mod tests { #[test] fn malformed_nlm_scale_is_rejected() { - let err = parse_prefilter("nlm:x").expect_err("expected parse failure"); - assert!(err.to_string().contains("nlm")); + let error = parse_prefilter("nlm:x").expect_err("expected parse failure"); + assert!(error.to_string().contains("nlm")); } #[test] fn unknown_prefilter_is_rejected() { - assert!(parse_prefilter("garbage").is_err()); - } - - #[test] - fn external_cannot_be_parsed_from_a_string() { - for s in [ - "external", - "none", - "nlm", - "nlm:0.5", - "bilateral:3.0,0.02", - "garbage", - ] { - if let Ok(mode) = parse_prefilter(s) { - assert!( - !matches!(mode, PrefilterMode::External), - "parse_prefilter must never produce External, got it from '{s}'" - ); - } - } + let result = parse_prefilter("garbage"); + assert!(result.is_err()); } } diff --git a/av-denoise-core/src/nlmeans/tests/alignment.rs b/av-denoise-core/src/nlmeans/tests/alignment.rs index 94ebb57..e2dcb11 100644 --- a/av-denoise-core/src/nlmeans/tests/alignment.rs +++ b/av-denoise-core/src/nlmeans/tests/alignment.rs @@ -1,11 +1,5 @@ -//! Frame sizes whose per-slot byte strides do not land on the GPU's -//! buffer-offset alignment, which is 32 bytes on the adapters these -//! tests run against. -//! -//! A buffer view bound at such an offset is rejected outright, so these -//! sizes used to abort the whole pipeline instead of denoising. - use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; fn luma_params(motion_compensation: MotionCompensationMode) -> NlmParams { @@ -22,47 +16,42 @@ fn luma_params(motion_compensation: MotionCompensationMode) -> NlmParams { } } -/// Pushes five uniform frames and denoises the centre one, asserting -/// the result is the same uniform value it went in as. -fn denoise_uniform(w: u32, h: u32, params: NlmParams) { +/// Pushes five uniform frames and asserts the centre one denoises to the same uniform value. +/// +/// The sizes passed in give per-slot byte strides off the 32-byte buffer-offset alignment of the +/// test adapters, and a view bound at such an offset is rejected outright. +fn denoise_uniform(width: u32, height: u32, params: NlmParams) { let client = make_client(); - let frame = make_uniform_frame(w, h, 1, 0.5); + let frame = make_uniform_frame(width, height, 1, 0.5); - let mut d = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); for _ in 0..5 { - d.push_frame(&frame); + denoiser.push_frame(&frame); } - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - assert_eq!(result.len(), (w * h) as usize); - for (i, &v) in result.iter().enumerate() { + let result = denoiser.denoise().unwrap().unwrap(); + + assert_eq!(result.len(), (width * height) as usize); + for (i, &value) in result.iter().enumerate() { assert!( - (v - 0.5).abs() < 1e-3, - "pixel {i}: expected 0.5 (uniform input passthrough), got {v}" + (value - 0.5).abs() < 1e-3, + "pixel {i}: expected 0.5 (uniform input passthrough), got {value}" ); } } #[test] fn denoises_when_the_frame_ring_slot_stride_is_unaligned() { - // A 34x34 luma slot is 4624 bytes, which is a multiple of 16 but - // not of 32, so every odd ring slot starts 16 bytes short of an - // alignment boundary. - denoise_uniform(34, 34, luma_params(MotionCompensationMode::None)); + // A 34x34 luma slot is 4624 bytes, a multiple of 16 but not 32, so every odd ring slot starts + // 16 bytes short of a boundary. + let params = luma_params(MotionCompensationMode::None); + denoise_uniform(34, 34, params); } #[test] fn denoises_the_nlm_spatial_pilot_when_the_reference_ring_slot_stride_is_unaligned() { - // The pilot writes its own output into the reference ring, so at - // 34x34 it targets a slot starting 4624 bytes in, 16 short of an - // alignment boundary. Only a temporal radius above 0 ever reaches - // a slot past the first. + // The pilot writes into the reference ring, so at 34x34 it targets a slot 4624 bytes in, 16 + // short of a boundary. Only a temporal radius above 0 reaches a slot past the first. let mut params = luma_params(MotionCompensationMode::None); params.prefilter = PrefilterMode::NlmSpatial { strength_scale: DEFAULT_PILOT_STRENGTH_SCALE, @@ -72,21 +61,16 @@ fn denoises_the_nlm_spatial_pilot_when_the_reference_ring_slot_stride_is_unalign #[test] fn denoises_with_motion_compensation_when_the_pyramid_slot_stride_is_unaligned() { - // A 42x28 luma slot is 4704 bytes, a clean 32-byte multiple, so - // only the motion pyramid is at stake here. Its half-size level is - // 21x14, which comes to 1176 bytes, 24 short of a boundary. - // - // That is the same shape as the 720x548 chroma plane that first hit - // this. - denoise_uniform( - 42, - 28, - luma_params(MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }), - ); + // A 42x28 luma slot is 4704 bytes, a clean 32-byte multiple, so only the motion pyramid is at + // stake. Its 21x14 half-size level is 1176 bytes, 24 short of a boundary. This matches the + // shape of a 720x548 chroma plane. + let motion_compensation = MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }; + let params = luma_params(motion_compensation); + denoise_uniform(42, 28, params); } diff --git a/av-denoise-core/src/nlmeans/tests/confidence.rs b/av-denoise-core/src/nlmeans/tests/confidence.rs index dccf051..b2b3cbb 100644 --- a/av-denoise-core/src/nlmeans/tests/confidence.rs +++ b/av-denoise-core/src/nlmeans/tests/confidence.rs @@ -1,6 +1,7 @@ use cubecl::prelude::*; use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::kernels::motion::nlm_mc_block_match_fine; use crate::nlmeans::motion::{ MotionCompensationMode, @@ -13,24 +14,24 @@ use crate::nlmeans::motion::{ }; use crate::nlmeans::*; -/// Lays `frames` out as a single-level pyramid buffer, each frame at -/// the slot stride `run_pyramid_build` writes to. Slots are padded up -/// to the GPU's storage-buffer offset alignment, so a frame does not -/// necessarily start right where the one before it ended. +/// Lays `frames` out as a single-level pyramid at the slot stride `run_pyramid_build` writes. +/// +/// Slots are padded up to the storage-buffer offset alignment, so a frame does not always start +/// where the one before it ended. fn pack_single_level_pyramid(frames: &[&[f32]], width: u32, height: u32, align: StorageAlign) -> Vec { let stride = pyramid_pixels_per_frame(width, height, 1, align); let mut data = vec![0.0f32; frames.len() * stride]; for (slot, frame) in frames.iter().enumerate() { data[slot * stride..slot * stride + frame.len()].copy_from_slice(frame); } + data } -/// Runs the fine block-match kernel over a single `blksize × blksize` -/// block (one cube covers the whole frame, no seed, no search window) -/// and returns the resulting confidence score. Isolates the confidence -/// expression from the search behaviour already covered by -/// `tests::motion_compensation`. +/// Runs the fine block-match kernel over one `blksize x blksize` block and returns its confidence. +/// +/// One cube covers the whole frame with no seed and no search window. With `blksize = 1` this +/// isolates the confidence expression (floor subtraction, threshold, clamp) from the SAD reduction. fn run_fine_confidence( blksize: u32, centre: &[f32], @@ -43,8 +44,10 @@ fn run_fine_confidence( assert_eq!(centre.len(), level_len); assert_eq!(neighbour.len(), level_len); - let centre_buf = client.create_from_slice(f32::as_bytes(centre)); - let neighbour_buf = client.create_from_slice(f32::as_bytes(neighbour)); + let centre_bytes = f32::as_bytes(centre); + let neighbour_bytes = f32::as_bytes(neighbour); + let centre_buf = client.create_from_slice(centre_bytes); + let neighbour_buf = client.create_from_slice(neighbour_bytes); let mv_field = client.empty(2 * size_of::()); let confidence = client.empty(size_of::()); @@ -77,116 +80,85 @@ fn run_fine_confidence( f32::from_bytes(&bytes)[0] } -// NOTE on `blksize` below. Tests (a)-(e) use `blksize = 1` (a -// single-pixel block) rather than a realistic multi-pixel block size. -// This isolates the confidence *expression* (floor subtraction, -// threshold, clamp) from the SAD reduction itself, which has its own -// dedicated coverage. A prior version of this SAD reduction had every -// thread in the cube accumulate into the *same* `SharedMemory` scratch -// slot via a plain `+=` with no atomics and no per-thread -// partial/reduce split, a data race whenever more than one thread -// contributed to a block. That's now fixed (see `block_match.rs`'s -// candidate-parallel reduction and `tests::motion_compensation`'s -// `exact_sad_*`/`argmin_*` tests, which exercise it directly at a -// realistic multi-pixel `blksize`). `blksize = 1` here is kept for -// isolation, not as a workaround. - -/// Zero mismatch, with a floor present. Excess clamps to zero, so -/// confidence must be exactly 1 regardless of how large the floor is. #[test] fn confidence_perfect_match_is_exactly_one() { let sigma = 0.1; let floor = sad_noise_floor(1, sigma); - let th = thsad(1, 1.0); - let confidence = run_fine_confidence(1, &[0.5], &[0.5], floor, th); + let threshold = thsad(1, 1.0); + let confidence = run_fine_confidence(1, &[0.5], &[0.5], floor, threshold); assert_eq!(confidence, 1.0, "zero mismatch must give exactly full confidence"); } -/// A per-pixel diff at or below the noise floor must not cost any -/// confidence. The excess clamps to zero either way. #[test] fn confidence_diff_within_noise_floor_is_exactly_one() { let sigma = 0.1; let floor = sad_noise_floor(1, sigma); - let th = thsad(1, 1.0); + let threshold = thsad(1, 1.0); // Half the floor, comfortably inside the "this is just noise" region. let diff = floor * 0.5; - let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, th); + let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, threshold); assert_eq!( confidence, 1.0, "a diff under the noise floor must give exactly full confidence" ); } -/// A mismatch whose excess (over the floor) reaches `thsad` must -/// collapse confidence to exactly zero, matching the documented -/// "0 at `S ≥ thsad`" behaviour. #[test] fn confidence_excess_at_thsad_is_exactly_zero() { let sigma = 0.1; let floor = sad_noise_floor(1, sigma); - let th = thsad(1, 1.0); - let diff = floor + th; - let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, th); + let threshold = thsad(1, 1.0); + let diff = floor + threshold; + let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, threshold); assert!( confidence < 1e-4, "excess reaching thsad must collapse confidence to ~zero, got {confidence}" ); } -/// A mismatch far beyond `thsad` must also collapse to zero (the -/// clamp, not just a small-but-positive value). #[test] fn confidence_excess_beyond_thsad_is_exactly_zero() { let sigma = 0.1; let floor = sad_noise_floor(1, sigma); - let th = thsad(1, 1.0); - let diff = floor + 10.0 * th; - let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, th); + let threshold = thsad(1, 1.0); + let diff = floor + 10.0 * threshold; + let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, threshold); assert_eq!( confidence, 0.0, "a gross mismatch must collapse confidence to exactly zero" ); } -/// Confidence must decrease monotonically (MDegrain-style) as the -/// excess over the floor grows from 0 to `thsad`. #[test] fn confidence_decreases_as_excess_grows() { let sigma = 0.1; let floor = sad_noise_floor(1, sigma); - let th = thsad(1, 1.0); + let threshold = thsad(1, 1.0); let excess_fractions = [0.0f32, 0.2, 0.4, 0.6, 0.8, 1.0]; - let mut prev = f32::INFINITY; - for &frac in &excess_fractions { - let diff = floor + frac * th; - let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, th); + let mut previous = f32::INFINITY; + for &fraction in &excess_fractions { + let diff = floor + fraction * threshold; + let confidence = run_fine_confidence(1, &[0.5], &[0.5 + diff], floor, threshold); assert!( - confidence <= prev + 1e-6, + confidence <= previous + 1e-6, "confidence should be non-increasing as excess grows: excess={}·thsad gave {confidence}, \ - previous was {prev}", - frac, + previous was {previous}", + fraction, ); - prev = confidence; + previous = confidence; } + assert!( - prev < 1e-4, - "excess reaching thsad must land at ~zero confidence, got {prev}" + previous < 1e-4, + "excess reaching thsad must land at ~zero confidence, got {previous}" ); } -/// Realistic-blksize matched case. Two independently-noisy copies of -/// the same clean content at `blksize = 16` (the library default), -/// exercising the full multi-thread SAD reduction rather than the -/// single-pixel isolation above. `sad_noise_floor` is calibrated so -/// the SAD two noisy copies produce by chance sits at the floor on -/// average, so confidence should stay high (the floor "absorbs" the -/// noise) rather than collapsing just because real content is noisy. +/// `sad_noise_floor` is calibrated so the chance SAD of two noisy copies sits at the floor on average. /// -/// This only demonstrates the formula discriminating when `best_sad` -/// is the true SAD. A racy reduction that undercounts it lets the -/// assertion pass without the floor doing any work. +/// The check only means something when `best_sad` is the full SAD. A racy reduction that +/// undercounts it passes without the floor doing any work. #[test] fn confidence_matched_noisy_content_at_blksize_16_is_near_one() { let blksize = 16; @@ -195,8 +167,8 @@ fn confidence_matched_noisy_content_at_blksize_16_is_near_one() { let neighbour = noisy_copy(blksize, 0.5, sigma, 31); let floor = sad_noise_floor(blksize, sigma); - let th = thsad(blksize, 1.0); - let confidence = run_fine_confidence(blksize, ¢re, &neighbour, floor, th); + let threshold = thsad(blksize, 1.0); + let confidence = run_fine_confidence(blksize, ¢re, &neighbour, floor, threshold); assert!( confidence > 0.9, "two independently-noisy copies of the same content at blksize=16 \ @@ -204,14 +176,8 @@ fn confidence_matched_noisy_content_at_blksize_16_is_near_one() { ); } -/// Realistic-blksize mismatched case. Same noise characteristics but a -/// clearly different underlying signal (a different base level), so -/// the true per-pixel diff is far larger than noise alone. Confidence -/// must collapse toward 0. An undercounted `best_sad` keeps -/// confidence artificially high instead. At this blksize/cube-dim -/// combination a racy reduction undercounts by roughly two orders of -/// magnitude (`best_sad ≈ 2.4` instead of `≈ 153.6` for a uniform -/// 0.6 mismatch), hiding all but the most extreme real mismatches. +/// A racy SAD reduction undercounts `best_sad` 64-fold at this block and cube size (about 1.6 +/// instead of 102.4), which keeps confidence falsely high. #[test] fn confidence_mismatched_block_at_blksize_16_is_near_zero() { let blksize = 16; @@ -220,8 +186,8 @@ fn confidence_mismatched_block_at_blksize_16_is_near_zero() { let neighbour = noisy_copy(blksize, 0.9, sigma, 41); let floor = sad_noise_floor(blksize, sigma); - let th = thsad(blksize, 1.0); - let confidence = run_fine_confidence(blksize, ¢re, &neighbour, floor, th); + let threshold = thsad(blksize, 1.0); + let confidence = run_fine_confidence(blksize, ¢re, &neighbour, floor, threshold); assert!( confidence < 0.1, "a block with a genuinely different base level (0.5 vs 0.9) should \ @@ -229,51 +195,36 @@ fn confidence_mismatched_block_at_blksize_16_is_near_zero() { ); } -/// Covers the bug where the confidence noise floor was sized from the -/// raw input sigma rather than the prefiltered one. -/// -/// `mc_sad_noise_floor_sigma` in `dispatch.rs` is unit tested there, -/// being a plain host-side function. This test exercises what it feeds -/// into instead. +/// Pins the bug where the confidence floor was sized from the raw input sigma instead of the +/// prefiltered one. /// -/// It builds two block pairs at the small residual noise level a cleaned -/// reference carries, well below a typical raw sigma. One pair genuinely -/// matches, holding the same content with independent residual noise. -/// The other is genuinely occluded, with a modest but real shift in -/// base level. -/// -/// Both run under two floors. The raw input sigma is the old behaviour, -/// far too large for this content, and a floor sized to the residual -/// noise is what the fix produces. -/// -/// The raw floor swamps the threshold and pins even the occluded pair -/// near full confidence, making it indistinguishable from the matched -/// one, which reproduces the bug. -/// -/// The residual-scale floor keeps the matched pair confident while -/// letting the occluded pair drop well below. +/// Both pairs carry the small residual noise of a cleaned reference. Under the raw-sigma floor the +/// occluded pair reads as confident as the matched one. Under a residual-scale floor only the +/// matched pair stays confident. #[test] fn confidence_discriminates_occluded_from_matched_at_prefilter_scale_noise() { let blksize = 16; - // Representative of a prefiltered reference's small residual noise - // (an order of magnitude below a typical raw σ_y). + // A prefiltered reference's residual noise, an order of magnitude below a typical raw sigma. let residual_sigma = 0.002f32; - // Representative raw input sigma, the pre-fix floor's input. let raw_sigma = 0.02f32; - let th = thsad(blksize, 1.0); + let threshold = thsad(blksize, 1.0); let matched_centre = noisy_copy(blksize, 0.5, residual_sigma, 50); let matched_neighbour = noisy_copy(blksize, 0.5, residual_sigma, 51); - // A modest but genuine base-level shift, small next to a typical - // raw floor but large next to the residual noise itself. + // A base-level shift that is small next to a raw floor but large next to the residual noise. let occluded_centre = noisy_copy(blksize, 0.5, residual_sigma, 52); let occluded_neighbour = noisy_copy(blksize, 0.518, residual_sigma, 53); let raw_floor = sad_noise_floor(blksize, raw_sigma); let matched_conf_raw_floor = - run_fine_confidence(blksize, &matched_centre, &matched_neighbour, raw_floor, th); - let occluded_conf_raw_floor = - run_fine_confidence(blksize, &occluded_centre, &occluded_neighbour, raw_floor, th); + run_fine_confidence(blksize, &matched_centre, &matched_neighbour, raw_floor, threshold); + let occluded_conf_raw_floor = run_fine_confidence( + blksize, + &occluded_centre, + &occluded_neighbour, + raw_floor, + threshold, + ); assert!( matched_conf_raw_floor > 0.99, "matched pair under the raw-sigma floor should already read as confident, \ @@ -287,10 +238,20 @@ fn confidence_discriminates_occluded_from_matched_at_prefilter_scale_noise() { ); let residual_floor = sad_noise_floor(blksize, residual_sigma); - let matched_conf_residual_floor = - run_fine_confidence(blksize, &matched_centre, &matched_neighbour, residual_floor, th); - let occluded_conf_residual_floor = - run_fine_confidence(blksize, &occluded_centre, &occluded_neighbour, residual_floor, th); + let matched_conf_residual_floor = run_fine_confidence( + blksize, + &matched_centre, + &matched_neighbour, + residual_floor, + threshold, + ); + let occluded_conf_residual_floor = run_fine_confidence( + blksize, + &occluded_centre, + &occluded_neighbour, + residual_floor, + threshold, + ); assert!( matched_conf_residual_floor > 0.9, "matched pair should stay confident under the residual-scale floor too, not penalised \ @@ -303,12 +264,7 @@ fn confidence_discriminates_occluded_from_matched_at_prefilter_scale_noise() { ); } -/// `run_analyse` (the with-MC path) must fill the confidence buffer -/// with real, in-range values, not leave it at whatever sentinel the -/// caller pre-seeded it with. Pyramid data is supplied directly rather -/// than built via `run_pyramid_build`, a single-level pyramid being -/// just its two frames laid out at the pyramid's own slot stride (see -/// [`pack_single_level_pyramid`]). +/// The pyramid is packed by hand as one level instead of built by `run_pyramid_build`. #[test] fn run_analyse_fills_confidence_buffer() { let client = make_client(); @@ -316,31 +272,30 @@ fn run_analyse_fills_confidence_buffer() { let height = 16; let frame_count = 2; - let mc = MotionCtx::new( - MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - }, - width, - height, - test_align(), - ) - .unwrap(); + let mode = MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 1, + estimation: MotionEstimation::Direct, + }; + let align = test_align(); + let mc = MotionCtx::new(mode, width, height, align).unwrap(); let frame0 = noisy_copy(width, 0.5, 4.0 / 255.0, 10); let frame1 = noisy_copy(width, 0.5, 4.0 / 255.0, 11); - let pyramid_data = pack_single_level_pyramid(&[&frame0, &frame1], width, height, test_align()); - let pyramid = client.create_from_slice(f32::as_bytes(&pyramid_data)); + let pyramid_data = pack_single_level_pyramid(&[&frame0, &frame1], width, height, align); + let pyramid_bytes = f32::as_bytes(&pyramid_data); + let pyramid = client.create_from_slice(pyramid_bytes); - let mv_field = client.empty(mc.mv_slots_per_neighbour() * 2 * size_of::()); + let mv_field_bytes = mc.mv_slots_per_neighbour() * 2 * size_of::(); + let mv_field = client.empty(mv_field_bytes); let sentinel = vec![-1.0f32; mc.mv_slots_per_neighbour()]; - let confidence = client.create_from_slice(f32::as_bytes(&sentinel)); + let sentinel_bytes = f32::as_bytes(&sentinel); + let confidence = client.create_from_slice(sentinel_bytes); let floor = sad_noise_floor(mc.blksize, 4.0 / 255.0); - let th = thsad(mc.blksize, 1.0); + let threshold = thsad(mc.blksize, 1.0); run_analyse::( &client, @@ -356,26 +311,26 @@ fn run_analyse_fills_confidence_buffer() { &confidence, true, floor, - th, + threshold, ) .expect("run_analyse dispatch failed"); let bytes = client.read_one(confidence).expect("confidence readback failed"); let data = f32::from_bytes(&bytes); assert_eq!(data.len(), mc.mv_slots_per_neighbour()); - for (i, &v) in data.iter().enumerate() { - assert!(v.is_finite(), "block {i}: non-finite confidence {v}"); - assert!((0.0..=1.0).contains(&v), "block {i}: out-of-range confidence {v}"); + for (i, &value) in data.iter().enumerate() { + assert!(value.is_finite(), "block {i}: non-finite confidence {value}"); + assert!( + (0.0..=1.0).contains(&value), + "block {i}: out-of-range confidence {value}" + ); assert_ne!( - v, -1.0, + value, -1.0, "block {i}: confidence left at the sentinel, kernel didn't write it" ); } } -/// `run_confidence_for_neighbour` (the no-MC path) must fill the -/// confidence buffer the same way, using the confidence-only geometry -/// instead of a real `Mvtools` configuration. #[test] fn run_confidence_for_neighbour_fills_confidence_buffer() { let client = make_client(); @@ -383,19 +338,23 @@ fn run_confidence_for_neighbour_fills_confidence_buffer() { let height = 16; let frame_count = 2; - let ctx = MotionCtx::confidence_only(width, height, test_align()); + let align = test_align(); + let ctx = MotionCtx::confidence_only(width, height, align); let frame0 = noisy_copy(width, 0.5, 4.0 / 255.0, 20); let frame1 = noisy_copy(width, 0.5, 4.0 / 255.0, 21); - let pyramid_data = pack_single_level_pyramid(&[&frame0, &frame1], width, height, test_align()); - let pyramid = client.create_from_slice(f32::as_bytes(&pyramid_data)); + let pyramid_data = pack_single_level_pyramid(&[&frame0, &frame1], width, height, align); + let pyramid_bytes = f32::as_bytes(&pyramid_data); + let pyramid = client.create_from_slice(pyramid_bytes); - let mv_scratch = client.empty(ctx.mv_slots_per_neighbour() * 2 * size_of::()); + let mv_scratch_bytes = ctx.mv_slots_per_neighbour() * 2 * size_of::(); + let mv_scratch = client.empty(mv_scratch_bytes); let sentinel = vec![-1.0f32; ctx.mv_slots_per_neighbour()]; - let confidence = client.create_from_slice(f32::as_bytes(&sentinel)); + let sentinel_bytes = f32::as_bytes(&sentinel); + let confidence = client.create_from_slice(sentinel_bytes); let floor = sad_noise_floor(ctx.blksize, 4.0 / 255.0); - let th = thsad(ctx.blksize, 1.0); + let threshold = thsad(ctx.blksize, 1.0); run_confidence_for_neighbour::( &client, @@ -410,18 +369,21 @@ fn run_confidence_for_neighbour_fills_confidence_buffer() { &mv_scratch, &confidence, floor, - th, + threshold, ) .expect("run_confidence_for_neighbour dispatch failed"); let bytes = client.read_one(confidence).expect("confidence readback failed"); let data = f32::from_bytes(&bytes); assert_eq!(data.len(), ctx.mv_slots_per_neighbour()); - for (i, &v) in data.iter().enumerate() { - assert!(v.is_finite(), "block {i}: non-finite confidence {v}"); - assert!((0.0..=1.0).contains(&v), "block {i}: out-of-range confidence {v}"); + for (i, &value) in data.iter().enumerate() { + assert!(value.is_finite(), "block {i}: non-finite confidence {value}"); + assert!( + (0.0..=1.0).contains(&value), + "block {i}: out-of-range confidence {value}" + ); assert_ne!( - v, -1.0, + value, -1.0, "block {i}: confidence left at the sentinel, kernel didn't write it" ); } @@ -430,9 +392,9 @@ fn run_confidence_for_neighbour_fills_confidence_buffer() { #[test] fn confidence_buf_filled_without_motion_compensation() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 1, 0.5); let params = NlmParams { temporal_radius: 1, @@ -454,43 +416,47 @@ fn confidence_buf_filled_without_motion_compensation() { }), }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - d.denoise().unwrap(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.denoise().unwrap(); assert!( - d.mc_ctx.is_none(), + denoiser.mc_ctx.is_none(), "this test exercises the no-MC confidence path" ); - let ctx = d + let ctx = denoiser .confidence_ctx .as_ref() .expect("confidence_ctx must be allocated"); - let handle = d + let handle = denoiser .confidence_buf .as_ref() .expect("confidence_buf must be allocated") .clone(); - let bytes = d.client.read_one(handle).expect("confidence readback failed"); + let bytes = denoiser + .client + .read_one(handle) + .expect("confidence readback failed"); let data = f32::from_bytes(&bytes); assert_eq!(data.len(), 2 * ctx.mv_slots_per_neighbour()); - for (i, &v) in data.iter().enumerate() { - assert!(v.is_finite(), "block {i}: non-finite confidence {v}"); - assert!((0.0..=1.0).contains(&v), "block {i}: out-of-range confidence {v}"); + for (i, &value) in data.iter().enumerate() { + assert!(value.is_finite(), "block {i}: non-finite confidence {value}"); + assert!( + (0.0..=1.0).contains(&value), + "block {i}: out-of-range confidence {value}" + ); } } -/// Same wiring check with motion compensation active. The analyse -/// fine pass must fill `confidence_buf` using `mc_ctx`'s geometry. #[test] fn confidence_buf_filled_with_motion_compensation() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 1, 0.5); let params = NlmParams { temporal_radius: 1, @@ -518,51 +484,49 @@ fn confidence_buf_filled_with_motion_compensation() { }), }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - d.denoise().unwrap(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.denoise().unwrap(); - let mc = d + let mc = denoiser .mc_ctx .as_ref() .expect("mc_ctx must be allocated when MC is active"); assert!( - d.confidence_ctx.is_none(), + denoiser.confidence_ctx.is_none(), "confidence_ctx is only for the no-MC path; MC-active reuses mc_ctx" ); - let handle = d + let handle = denoiser .confidence_buf .as_ref() .expect("confidence_buf must be allocated") .clone(); - let bytes = d.client.read_one(handle).expect("confidence readback failed"); + let bytes = denoiser + .client + .read_one(handle) + .expect("confidence readback failed"); let data = f32::from_bytes(&bytes); assert_eq!(data.len(), 2 * mc.mv_slots_per_neighbour()); - for (i, &v) in data.iter().enumerate() { - assert!(v.is_finite(), "block {i}: non-finite confidence {v}"); - assert!((0.0..=1.0).contains(&v), "block {i}: out-of-range confidence {v}"); + for (i, &value) in data.iter().enumerate() { + assert!(value.is_finite(), "block {i}: non-finite confidence {value}"); + assert!( + (0.0..=1.0).contains(&value), + "block {i}: out-of-range confidence {value}" + ); } } -/// Motion compensation active with HQ off, which is how the fast path -/// usually runs it. -/// -/// Confidence weighting needs HQ with the flag on, so with HQ off no -/// confidence buffer should be allocated even though motion -/// compensation is active. -/// -/// The fine block-match kernel still runs, because it always has to -/// solve for the motion vector, but it writes no confidence and takes a -/// placeholder buffer. +/// The fine block-match kernel still runs for the motion vectors but takes a placeholder confidence +/// buffer. #[test] fn confidence_buf_absent_with_motion_compensation_and_no_hq() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 1, 0.5); let params = NlmParams { temporal_radius: 1, @@ -582,35 +546,29 @@ fn confidence_buf_absent_with_motion_compensation_and_no_hq() { hq: None, }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - d.denoise().unwrap(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.denoise().unwrap(); assert!( - d.mc_ctx.is_some(), + denoiser.mc_ctx.is_some(), "this test exercises the MC-active, HQ-off path" ); assert!( - d.confidence_buf.is_none(), + denoiser.confidence_buf.is_none(), "confidence_buf must stay absent without HQ, even with MC active" ); } -/// Motion compensation active and HQ on, but temporal confidence turned -/// off. -/// -/// Confidence weighting follows the flag whether or not motion -/// compensation already supplies block geometry, so nothing should be -/// allocated here either. That matches the test further down for the -/// same case without motion compensation. +/// Confidence weighting follows the flag even when motion compensation already supplies block geometry. #[test] fn confidence_buf_absent_with_motion_compensation_when_temporal_confidence_disabled() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 1, 0.5); let params = NlmParams { temporal_radius: 1, @@ -633,22 +591,22 @@ fn confidence_buf_absent_with_motion_compensation_when_temporal_confidence_disab }), }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - d.denoise().unwrap(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.denoise().unwrap(); assert!( - d.mc_ctx.is_some(), + denoiser.mc_ctx.is_some(), "this test exercises the MC-active path with confidence explicitly disabled" ); assert!( - d.confidence_ctx.is_none(), + denoiser.confidence_ctx.is_none(), "confidence_ctx is only for the no-MC path" ); assert!( - d.confidence_buf.is_none(), + denoiser.confidence_buf.is_none(), "confidence_buf must stay absent when temporal_confidence is off, even with MC active" ); } @@ -657,19 +615,14 @@ fn confidence_buf_absent_with_motion_compensation_when_temporal_confidence_disab fn confidence_buf_absent_without_mc_or_hq() { let client = make_client(); let params = NlmParams::default(); - let d = NlmDenoiser::::new(&client, params, 16, 16); + let denoiser = NlmDenoiser::::new(&client, params, 16, 16); - assert!(d.confidence_ctx.is_none()); - assert!(d.confidence_buf.is_none()); - assert!(d.confidence_pyramid.is_none()); - assert!(d.confidence_mv_scratch.is_none()); + assert!(denoiser.confidence_ctx.is_none()); + assert!(denoiser.confidence_buf.is_none()); + assert!(denoiser.confidence_pyramid.is_none()); + assert!(denoiser.confidence_mv_scratch.is_none()); } -/// HQ with temporal confidence off and no motion compensation should -/// also allocate nothing. -/// -/// Without motion compensation the confidence pass is the only thing -/// that would allocate the buffer, and it follows the same flag. #[test] fn confidence_buf_absent_when_temporal_confidence_disabled() { let client = make_client(); @@ -681,8 +634,8 @@ fn confidence_buf_absent_when_temporal_confidence_disabled() { }), ..NlmParams::default() }; - let d = NlmDenoiser::::new(&client, params, 16, 16); + let denoiser = NlmDenoiser::::new(&client, params, 16, 16); - assert!(d.confidence_ctx.is_none()); - assert!(d.confidence_buf.is_none()); + assert!(denoiser.confidence_ctx.is_none()); + assert!(denoiser.confidence_buf.is_none()); } diff --git a/av-denoise-core/src/nlmeans/tests/dispatch.rs b/av-denoise-core/src/nlmeans/tests/dispatch.rs new file mode 100644 index 0000000..d473766 --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/dispatch.rs @@ -0,0 +1,86 @@ +use crate::nlmeans::dispatch::{ + BILATERAL_RESIDUAL_FRACTION, + NLM_SPATIAL_RESIDUAL_FRACTION, + mc_sad_noise_floor_sigma, +}; +use crate::nlmeans::motion::neighbour_idx_for_k; +use crate::nlmeans::prefilter::PrefilterMode; + +#[test] +fn matches_the_sequential_fill_order() { + for radius in 1..=8u32 { + let mut expected = 0u32; + for k in -(radius as i32)..=(radius as i32) { + if k == 0 { + continue; + } + + let idx = neighbour_idx_for_k(radius, k); + assert_eq!(idx, expected, "radius={radius} k={k}"); + expected += 1; + } + } +} + +/// Distinct indices keep one frame's confidence off another frame's temporal weight. +#[test] +fn forward_and_backward_indices_are_distinct_and_in_range() { + for radius in 1..=8u32 { + for q_k in -(radius as i32)..0 { + let forward = neighbour_idx_for_k(radius, q_k); + let backward = neighbour_idx_for_k(radius, -q_k); + assert_ne!(forward, backward, "radius={radius} q_k={q_k}"); + assert!(forward < 2 * radius, "radius={radius} q_k={q_k} fwd={forward}"); + assert!(backward < 2 * radius, "radius={radius} q_k={q_k} bwd={backward}"); + } + } +} + +#[test] +fn radius_two_explicit_indices() { + let back_two = neighbour_idx_for_k(2, -2); + let back_one = neighbour_idx_for_k(2, -1); + let forward_one = neighbour_idx_for_k(2, 1); + let forward_two = neighbour_idx_for_k(2, 2); + assert_eq!(back_two, 0); + assert_eq!(back_one, 1); + assert_eq!(forward_one, 2); + assert_eq!(forward_two, 3); +} + +/// A literal, so a recalibration fails here instead of passing against itself. +#[test] +fn nlm_spatial_residual_fraction_is_calibrated_to_zero() { + assert_eq!(NLM_SPATIAL_RESIDUAL_FRACTION, 0.0); +} + +#[test] +fn bilateral_residual_fraction_is_calibrated_to_zero() { + assert_eq!(BILATERAL_RESIDUAL_FRACTION, 0.0); +} + +#[test] +fn mc_sad_noise_floor_sigma_scales_nlm_spatial_by_the_calibrated_fraction() { + let raw = 0.02f32; + let prefilter = PrefilterMode::NlmSpatial { strength_scale: 1.0 }; + let floor = mc_sad_noise_floor_sigma(prefilter, raw); + assert_eq!(floor, raw * NLM_SPATIAL_RESIDUAL_FRACTION); +} + +#[test] +fn mc_sad_noise_floor_sigma_scales_bilateral_by_the_calibrated_fraction() { + let raw = 0.02f32; + let prefilter = PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: 0.02, + }; + let floor = mc_sad_noise_floor_sigma(prefilter, raw); + assert_eq!(floor, raw * BILATERAL_RESIDUAL_FRACTION); +} + +#[test] +fn mc_sad_noise_floor_sigma_keeps_raw_sigma_for_none() { + let raw = 0.02f32; + let floor = mc_sad_noise_floor_sigma(PrefilterMode::None, raw); + assert_eq!(floor, raw); +} diff --git a/av-denoise-core/src/nlmeans/tests/edges.rs b/av-denoise-core/src/nlmeans/tests/edges.rs index 842231e..a206e07 100644 --- a/av-denoise-core/src/nlmeans/tests/edges.rs +++ b/av-denoise-core/src/nlmeans/tests/edges.rs @@ -1,7 +1,8 @@ use cubecl::prelude::*; use super::helpers::*; -use crate::nl4d::tests::helpers::{noisy_copy_of, textured_base}; +use crate::bench_api::HostIo; +use crate::nl4d::tests::helpers::textured_base; use crate::nlmeans::*; const RADIUS: u32 = 2; @@ -32,7 +33,8 @@ fn edge_params(windowed: bool) -> NlmParams { } fn shifted_denoiser(client: &ComputeClient, windowed: bool) -> NlmDenoiser { - let mut denoiser = NlmDenoiser::::new(client, edge_params(windowed), SIZE, SIZE); + let params = edge_params(windowed); + let mut denoiser = NlmDenoiser::::new(client, params, SIZE, SIZE); denoiser.set_shifted_edges(true); denoiser.set_luma_noise_fields(true); denoiser @@ -41,7 +43,7 @@ fn shifted_denoiser(client: &ComputeClient, windowed: bool) -> NlmDenoiser fn push_grain(denoiser: &mut NlmDenoiser, count: u32) { let base = textured_base(SIZE, SIZE); for seed in 0..count { - let frame = noisy_copy_of(&base, SIZE, SIZE, GRAIN, seed); + let frame = noisy_field_over(&base, SIZE, SIZE, GRAIN, seed); denoiser.push_frame(&frame); } } @@ -85,7 +87,7 @@ fn centre_zero_falls_back_when_every_reading_ahead_is_rejected() { let client = make_client(); let mut denoiser = shifted_denoiser(&client, true); let base = textured_base(SIZE, SIZE); - let frozen = noisy_copy_of(&base, SIZE, SIZE, GRAIN, 0); + let frozen = noisy_field_over(&base, SIZE, SIZE, GRAIN, 0); for _ in 0..(2 * RADIUS + 1) { denoiser.push_frame(&frozen); } @@ -103,13 +105,13 @@ fn a_second_stream_never_reads_the_first_streams_f0_stats() { let client = make_client(); let total_frames = 2 * RADIUS + 1; let mut denoiser = shifted_denoiser(&client, true); - // One more than the ring, so the first stream writes a real record - // into the slot the second stream's f0 lands in. + // One more than the ring, so the first stream writes a real record into the slot the second + // stream's f0 lands in. push_grain(&mut denoiser, total_frames + 1); denoiser.reset_stream_state(); let base = textured_base(SIZE, SIZE); - let frozen = noisy_copy_of(&base, SIZE, SIZE, GRAIN, 0); + let frozen = noisy_field_over(&base, SIZE, SIZE, GRAIN, 0); for _ in 0..total_frames { denoiser.push_frame(&frozen); } @@ -127,7 +129,7 @@ fn fill_ring_with_last_frame_makes_a_short_stream_ready() { let mut denoiser = shifted_denoiser(&client, false); push_grain(&mut denoiser, 2); - denoiser.fill_ring_with_last_frame(); + denoiser.fill_ring_with_last_frame().expect("fill ring"); assert!(denoiser.window_ready()); assert_eq!(denoiser.real_pushes(), 2); @@ -136,7 +138,8 @@ fn fill_ring_with_last_frame_makes_a_short_stream_ready() { #[test] fn nlm_mode_keeps_the_leading_copies() { let client = make_client(); - let mut denoiser = NlmDenoiser::::new(&client, edge_params(false), SIZE, SIZE); + let params = edge_params(false); + let mut denoiser = NlmDenoiser::::new(&client, params, SIZE, SIZE); push_grain(&mut denoiser, 1); diff --git a/av-denoise-core/src/nlmeans/tests/engine.rs b/av-denoise-core/src/nlmeans/tests/engine.rs new file mode 100644 index 0000000..3bc1c3f --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/engine.rs @@ -0,0 +1,629 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::helpers::{ + R, + make_client, + make_noisy_gaussian_frame, + normalise_with_ingest, + read_interleaved, + upload_planes, +}; +use crate::bench_api::HostIo; +use crate::engine::{DevicePlane, Engine, Geometry, SampleFormat}; +use crate::error::Error; +use crate::nlmeans::{ + ChannelMode, + DenoisingMode, + MotionCompensationMode, + NlmDenoiser, + Nlmeans, + NlmeansAlgorithm, + NlmeansHqOptions, + NlmeansOptions, + resolve_params, +}; + +const WIDTH: u32 = 48; +const HEIGHT: u32 = 40; + +fn geometry(channels: ChannelMode) -> Geometry { + Geometry { + width: WIDTH, + height: HEIGHT, + channels, + input: SampleFormat::F32, + output: SampleFormat::F32, + } +} + +fn hq(radius: u32) -> NlmeansAlgorithm { + let mode = match radius { + 0 => DenoisingMode::Spacial, + radius => DenoisingMode::Temporal { radius }, + }; + let nlm_options = NlmeansOptions { + mode, + ..NlmeansOptions::default() + }; + + NlmeansAlgorithm::Hq(NlmeansHqOptions { + nlm: nlm_options, + ..NlmeansHqOptions::default() + }) +} + +fn frames(count: usize, channels: u32) -> Vec> { + (0..count) + .map(|index| { + let base = 0.3 + index as f32 * 0.01; + make_noisy_gaussian_frame(WIDTH, HEIGHT, channels, base, &[0.03]) + }) + .collect() +} + +fn emit(engine: &mut Nlmeans, client: &ComputeClient, channel_count: usize) -> Vec { + let pixels = (WIDTH * HEIGHT) as usize; + let outputs: Vec = (0..channel_count).map(|_| client.empty(pixels * 4)).collect(); + let planes: Vec<_> = outputs + .iter() + .map(|handle| DevicePlane::new(handle, WIDTH, HEIGHT)) + .collect(); + + engine.emit_into(&planes).expect("emit"); + + read_interleaved(client, &outputs) +} + +/// Pushes one frame and returns how many frames the engine reports ready. +fn push_frame( + engine: &mut Nlmeans, + client: &ComputeClient, + frame: &[f32], + channel_count: usize, +) -> usize { + let handles = upload_planes(client, frame, channel_count); + let planes: Vec<_> = handles + .iter() + .map(|handle| DevicePlane::new(handle, WIDTH, HEIGHT)) + .collect(); + + engine.push(&planes).expect("push") +} + +fn try_push_context(engine: &mut Nlmeans, client: &ComputeClient, frame: &[f32]) -> Result<(), Error> { + let handles = upload_planes(client, frame, 1); + let planes = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + + engine.push_context(&planes) +} + +fn push_and_emit( + engine: &mut Nlmeans, + client: &ComputeClient, + frame: &[f32], + channel_count: usize, + outputs: &mut Vec>, +) { + let ready = push_frame(engine, client, frame, channel_count); + for _ in 0..ready { + let output = emit(engine, client, channel_count); + outputs.push(output); + } +} + +fn finish_and_emit( + engine: &mut Nlmeans, + client: &ComputeClient, + channel_count: usize, + outputs: &mut Vec>, +) { + let tail = engine.finish().expect("finish"); + for _ in 0..tail { + let output = emit(engine, client, channel_count); + outputs.push(output); + } +} + +/// Pushes every frame through `engine`, then the tail, returning each emitted frame. +fn drive( + engine: &mut Nlmeans, + client: &ComputeClient, + frames: &[Vec], + channel_count: usize, +) -> Vec> { + let mut outputs = Vec::new(); + + for frame in frames { + push_and_emit(engine, client, frame, channel_count, &mut outputs); + } + + finish_and_emit(engine, client, channel_count, &mut outputs); + + outputs +} + +fn build_engine(client: &ComputeClient, radius: u32, channels: ChannelMode) -> Nlmeans { + let algorithm = hq(radius); + let geometry = geometry(channels); + + Nlmeans::new(client, algorithm, geometry).expect("build") +} + +fn run_engine(radius: u32, channels: ChannelMode, frames: &[Vec]) -> Vec> { + let client = make_client(); + let channel_count = channels.count() as usize; + let mut engine = build_engine(&client, radius, channels); + + drive(&mut engine, &client, frames, channel_count) +} + +fn run_oracle(radius: u32, channels: ChannelMode, frames: &[Vec]) -> Vec> { + let client = make_client(); + let algorithm = hq(radius); + let params = resolve_params(&algorithm, channels); + let mut denoiser = NlmDenoiser::new(&client, params, WIDTH, HEIGHT); + let mut outputs = Vec::new(); + + for frame in frames { + denoiser.push_frame(frame); + + let output = denoiser.denoise().expect("denoise"); + if let Some(samples) = output { + outputs.push(samples); + } + } + + let collect = |output: &[f32]| { + let samples = output.to_vec(); + outputs.push(samples); + }; + denoiser.flush(collect).expect("flush"); + + outputs +} + +#[test] +fn spatial_luma_matches_the_oracle() { + let frames = frames(4, 1); + let actual = run_engine(0, ChannelMode::Luma, &frames); + let expected = run_oracle(0, ChannelMode::Luma, &frames); + assert_eq!(actual, expected); +} + +#[test] +fn temporal_chroma_matches_the_oracle_including_the_tail() { + let frames = frames(8, 2); + let actual = run_engine(2, ChannelMode::Chroma, &frames); + let expected = run_oracle(2, ChannelMode::Chroma, &frames); + assert_eq!(actual.len(), 8); + assert_eq!(actual, expected); +} + +#[test] +fn temporal_yuv_matches_the_oracle() { + let frames = frames(6, 3); + let actual = run_engine(1, ChannelMode::Yuv, &frames); + let expected = run_oracle(1, ChannelMode::Yuv, &frames); + assert_eq!(actual, expected); +} + +#[test] +fn a_stream_shorter_than_the_window_matches_the_oracle() { + let frames = frames(2, 1); + let actual = run_engine(2, ChannelMode::Luma, &frames); + let expected = run_oracle(2, ChannelMode::Luma, &frames); + assert_eq!(actual.len(), 2); + assert_eq!(actual, expected); +} + +#[test] +fn a_second_stream_after_finish_matches_a_fresh_engine() { + let frames = frames(6, 1); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + + drive(&mut engine, &client, &frames, 1); + let second = drive(&mut engine, &client, &frames, 1); + + let fresh = run_engine(2, ChannelMode::Luma, &frames); + assert_eq!(second, fresh); +} + +#[test] +fn push_context_then_one_push_matches_the_streaming_centre() { + let radius = 2u32; + let context = 2 * radius as usize; + let window = context + 1; + let frames = frames(window, 1); + let client = make_client(); + let mut engine = build_engine(&client, radius, ChannelMode::Luma); + + for frame in &frames[..context] { + let handles = upload_planes(&client, frame, 1); + let planes = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + engine.push_context(&planes).expect("push context"); + } + + let ready = push_frame(&mut engine, &client, &frames[context], 1); + assert_eq!(ready, 1); + + let actual = emit(&mut engine, &client, 1); + let expected = run_oracle(radius, ChannelMode::Luma, &frames); + assert_eq!(actual, expected[radius as usize]); +} + +#[test] +fn push_context_after_a_push_is_refused_and_changes_no_state() { + let frames = frames(6, 1); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + let mut outputs = Vec::new(); + + for (index, frame) in frames.iter().enumerate() { + push_and_emit(&mut engine, &client, frame, 1, &mut outputs); + + if index == 2 { + let refused = try_push_context(&mut engine, &client, &frames[5]); + assert!(matches!(refused, Err(Error::ContextAfterPush))); + } + } + + finish_and_emit(&mut engine, &client, 1, &mut outputs); + + let expected = run_oracle(2, ChannelMode::Luma, &frames); + assert_eq!(outputs, expected); +} + +#[test] +fn push_context_is_allowed_again_after_reset() { + let frames = frames(2, 1); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + + push_frame(&mut engine, &client, &frames[0], 1); + engine.reset(); + + let context = try_push_context(&mut engine, &client, &frames[1]); + assert!(context.is_ok()); +} + +#[test] +fn push_context_is_allowed_again_after_the_tail_is_emitted() { + let frames = frames(4, 1); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + + drive(&mut engine, &client, &frames, 1); + + let context = try_push_context(&mut engine, &client, &frames[0]); + assert!(context.is_ok()); +} + +#[test] +fn push_context_is_allowed_again_after_a_finish_with_no_tail() { + let frames = frames(2, 1); + let client = make_client(); + let mut engine = build_engine(&client, 0, ChannelMode::Luma); + + drive(&mut engine, &client, &frames, 1); + + let context = try_push_context(&mut engine, &client, &frames[0]); + assert!(context.is_ok()); +} + +#[test] +fn push_before_emitting_returns_outputs_pending_and_keeps_the_frame() { + let frames = frames(3, 1); + let client = make_client(); + let mut engine = build_engine(&client, 0, ChannelMode::Luma); + let handles = upload_planes(&client, &frames[0], 1); + let planes = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + + let ready = engine.push(&planes).expect("push"); + assert_eq!(ready, 1); + + let second = engine.push(&planes); + assert!(matches!(second, Err(Error::OutputsPending))); + + let output = emit(&mut engine, &client, 1); + let expected = run_oracle(0, ChannelMode::Luma, &frames[..1]); + assert_eq!(output, expected[0]); +} + +#[test] +fn finish_while_a_frame_is_ready_returns_outputs_pending() { + let frames = frames(1, 1); + let client = make_client(); + let mut engine = build_engine(&client, 0, ChannelMode::Luma); + + let ready = push_frame(&mut engine, &client, &frames[0], 1); + assert_eq!(ready, 1); + + let finished = engine.finish(); + assert!(matches!(finished, Err(Error::OutputsPending))); +} + +fn engine_owing_a_tail(client: &ComputeClient) -> (Nlmeans, usize) { + let frames = frames(4, 1); + let mut engine = build_engine(client, 1, ChannelMode::Luma); + let mut outputs = Vec::new(); + + for frame in &frames { + push_and_emit(&mut engine, client, frame, 1, &mut outputs); + } + + let tail = engine.finish().expect("finish"); + assert!(tail > 0); + + (engine, tail) +} + +#[test] +fn push_while_a_tail_is_owed_returns_outputs_pending() { + let client = make_client(); + let (mut engine, _tail) = engine_owing_a_tail(&client); + let frame = frames(1, 1); + let handles = upload_planes(&client, &frame[0], 1); + let planes = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + + let pushed = engine.push(&planes); + assert!(matches!(pushed, Err(Error::OutputsPending))); +} + +#[test] +fn emit_after_the_last_tail_frame_returns_nothing_to_emit() { + let client = make_client(); + let (mut engine, tail) = engine_owing_a_tail(&client); + + for _ in 0..tail { + emit(&mut engine, &client, 1); + } + + let output = client.empty((WIDTH * HEIGHT * 4) as usize); + let planes = [DevicePlane::new(&output, WIDTH, HEIGHT)]; + let emitted = engine.emit_into(&planes); + assert!(matches!(emitted, Err(Error::NothingToEmit))); +} + +#[test] +fn misuse_errors_do_not_poison() { + let client = make_client(); + let mut engine = build_engine(&client, 0, ChannelMode::Luma); + let output = client.empty((WIDTH * HEIGHT * 4) as usize); + let planes = [DevicePlane::new(&output, WIDTH, HEIGHT)]; + + let nothing = engine.emit_into(&planes); + assert!(matches!(nothing, Err(Error::NothingToEmit))); + + let wrong_count = engine.push(&[]); + assert!(matches!(wrong_count, Err(Error::PlaneMismatch(_)))); + + let frame = frames(1, 1); + let ready = push_frame(&mut engine, &client, &frame[0], 1); + assert_eq!(ready, 1); +} + +#[test] +fn a_plane_mismatch_mid_stream_changes_no_state() { + let frames = frames(6, 1); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + let mut outputs = Vec::new(); + + for (index, frame) in frames.iter().enumerate() { + push_and_emit(&mut engine, &client, frame, 1, &mut outputs); + + if index == 3 { + let wrong_count = engine.push(&[]); + assert!(matches!(wrong_count, Err(Error::PlaneMismatch(_)))); + } + } + + finish_and_emit(&mut engine, &client, 1, &mut outputs); + + let expected = run_oracle(2, ChannelMode::Luma, &frames); + assert_eq!(outputs, expected); +} + +#[test] +fn a_gpu_failure_poisons_until_reset() { + let client = make_client(); + let mut engine = build_engine(&client, 0, ChannelMode::Luma); + + let failure = engine.fail_through_guard_for_test(); + assert!(matches!(failure, Err(Error::Gpu(_)))); + + let frame = frames(1, 1); + let handles = upload_planes(&client, &frame[0], 1); + let input = [DevicePlane::new(&handles[0], WIDTH, HEIGHT)]; + let pushed = engine.push(&input); + assert!(matches!(pushed, Err(Error::NeedsReset))); + + engine.reset(); + let ready = engine.push(&input).expect("push after reset"); + assert_eq!(ready, 1); +} + +#[test] +fn reset_mid_stream_with_a_ready_frame_matches_a_fresh_engine() { + let frames = frames(6, 1); + let client = make_client(); + let mut engine = build_engine(&client, 2, ChannelMode::Luma); + + let mut ready = 0; + for frame in &frames[..3] { + ready = push_frame(&mut engine, &client, frame, 1); + } + + assert_eq!(ready, 1); + + engine.reset(); + let second = drive(&mut engine, &client, &frames, 1); + + let fresh = run_engine(2, ChannelMode::Luma, &frames); + assert_eq!(second, fresh); +} + +/// Pushes each input, then the tail, emitting every frame as `u8` planes. +fn drive_u8(engine: &mut Nlmeans, client: &ComputeClient, inputs: &[Handle]) -> Vec> { + let pixels = (WIDTH * HEIGHT) as usize; + let mut outputs = Vec::new(); + + for input in inputs { + let planes = [DevicePlane::new(input, WIDTH, HEIGHT)]; + let ready = engine.push(&planes).expect("push"); + emit_u8(engine, client, ready, pixels, &mut outputs); + } + + let tail = engine.finish().expect("finish"); + emit_u8(engine, client, tail, pixels, &mut outputs); + + outputs +} + +fn emit_u8( + engine: &mut Nlmeans, + client: &ComputeClient, + frame_count: usize, + pixels: usize, + outputs: &mut Vec>, +) { + for _ in 0..frame_count { + let output = client.empty(pixels); + let planes = [DevicePlane::new(&output, WIDTH, HEIGHT)]; + engine.emit_into(&planes).expect("emit"); + + let bytes = client.read_one(output).expect("read"); + outputs.push(bytes.to_vec()); + } +} + +#[test] +fn u8_input_matches_f32_input_from_the_ingest_kernel() { + let client = make_client(); + let codes: Vec> = (0..6) + .map(|frame_index| { + (0..WIDTH * HEIGHT) + .map(|index| ((index + frame_index * 7) % 251) as u8) + .collect() + }) + .collect(); + let u8_inputs: Vec = codes + .iter() + .map(|frame| client.create_from_slice(frame)) + .collect(); + let f32_inputs: Vec = codes + .iter() + .map(|frame| { + let normalised = normalise_with_ingest(&client, frame, WIDTH, HEIGHT); + let bytes = f32::as_bytes(&normalised); + client.create_from_slice(bytes) + }) + .collect(); + + let u8_geometry = Geometry { + width: WIDTH, + height: HEIGHT, + channels: ChannelMode::Luma, + input: SampleFormat::U8, + output: SampleFormat::U8, + }; + let f32_geometry = Geometry { + input: SampleFormat::F32, + ..u8_geometry + }; + let u8_algorithm = hq(2); + let f32_algorithm = hq(2); + let mut u8_engine = Nlmeans::new(&client, u8_algorithm, u8_geometry).expect("build u8"); + let mut f32_engine = Nlmeans::new(&client, f32_algorithm, f32_geometry).expect("build f32"); + + let u8_frames = drive_u8(&mut u8_engine, &client, &u8_inputs); + let f32_frames = drive_u8(&mut f32_engine, &client, &f32_inputs); + assert_eq!(u8_frames.len(), 6); + assert_eq!(u8_frames, f32_frames); +} + +#[test] +fn a_plane_past_u32_pixels_is_invalid_geometry() { + let client = make_client(); + let oversized = Geometry { + width: 65_536, + height: 65_536, + ..geometry(ChannelMode::Luma) + }; + + let algorithm = hq(0); + let built = Nlmeans::new(&client, algorithm, oversized); + assert!(matches!(built, Err(Error::InvalidGeometry(_)))); +} + +/// Each 16384x16384 YUV frame stores 2^30 elements, so a five-frame ring passes `u32::MAX`. +#[test] +fn a_ring_past_u32_elements_is_invalid_geometry() { + let client = make_client(); + let oversized = Geometry { + width: 16_384, + height: 16_384, + ..geometry(ChannelMode::Yuv) + }; + + let algorithm = hq(2); + let built = Nlmeans::new(&client, algorithm, oversized); + assert!(matches!(built, Err(Error::InvalidGeometry(_)))); +} + +/// A 32768x24576 luma frame ring of five frames fits in `u32`, but its two-level pyramid ring holds +/// 1.25 times as many elements and does not. +#[test] +fn a_motion_pyramid_past_u32_elements_is_invalid_geometry() { + let client = make_client(); + let oversized = Geometry { + width: 32_768, + height: 24_576, + ..geometry(ChannelMode::Luma) + }; + let nlm_options = NlmeansOptions { + mode: DenoisingMode::Temporal { radius: 2 }, + motion_compensation: MotionCompensationMode::mvtools_default(), + ..NlmeansOptions::default() + }; + let algorithm = NlmeansAlgorithm::Hq(NlmeansHqOptions { + nlm: nlm_options, + ..NlmeansHqOptions::default() + }); + + let built = Nlmeans::new(&client, algorithm, oversized); + let Err(Error::InvalidGeometry(message)) = built else { + panic!("expected InvalidGeometry"); + }; + assert!(message.contains("pyramid"), "{message}"); +} + +#[test] +fn hd_and_4k_geometries_construct() { + let client = make_client(); + + for (width, height) in [(1920, 1080), (3840, 2160)] { + let sized = Geometry { + width, + height, + ..geometry(ChannelMode::Luma) + }; + + let algorithm = hq(2); + let built = Nlmeans::new(&client, algorithm, sized); + assert!(built.is_ok(), "{width}x{height} failed to construct"); + } +} + +#[test] +fn window_span_and_held_frames_follow_the_radius() { + let client = make_client(); + let engine = build_engine(&client, 3, ChannelMode::Luma); + let span = engine.window_span(); + assert_eq!((span.behind, span.ahead), (3, 3)); + + let held = engine.max_held_frames(); + assert_eq!(held, 3); +} diff --git a/av-denoise-core/src/nlmeans/tests/gpu_submit.rs b/av-denoise-core/src/nlmeans/tests/gpu_submit.rs index b28405a..037151e 100644 --- a/av-denoise-core/src/nlmeans/tests/gpu_submit.rs +++ b/av-denoise-core/src/nlmeans/tests/gpu_submit.rs @@ -1,13 +1,7 @@ -//! `denoise_submit_gpu` and `flush_step_gpu` skip the readback that -//! `denoise_submit` and `flush` start automatically, handing back the -//! raw GPU handle instead. -//! -//! These tests pin that the GPU-resident path produces exactly the same -//! frames, in the same order, as the readback path it is built from. - use cubecl::prelude::*; use super::helpers::*; +use crate::bench_api::{GpuOutput, HostIo}; use crate::nlmeans::*; fn temporal_params(radius: u32) -> NlmParams { @@ -24,13 +18,12 @@ fn temporal_params(radius: u32) -> NlmParams { } } -/// Distinct frames so a mixed-up slot or a stale readback shows up as a -/// mismatch rather than being masked by every frame looking the same. -fn distinct_frames(w: u32, h: u32, count: usize) -> Vec> { +/// Distinct frames, so a mixed-up slot or a stale readback shows up as a mismatch. +fn distinct_frames(width: u32, height: u32, count: usize) -> Vec> { (0..count) .map(|i| { - let noise_val = 0.2 + (i as f32) * 0.05; - make_frame_with_noisy_region(w, h, 1, 0.5, w / 2, h / 2, 3, noise_val) + let noise_value = 0.2 + (i as f32) * 0.05; + make_frame_with_noisy_region(width, height, 1, 0.5, width / 2, height / 2, 3, noise_value) }) .collect() } @@ -38,51 +31,47 @@ fn distinct_frames(w: u32, h: u32, count: usize) -> Vec> { #[test] fn submit_gpu_matches_submit() { let client = make_client(); - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut via_readback = NlmDenoiser::::new(&client, temporal_params(1), w, h); - let mut via_gpu = NlmDenoiser::::new(&client, temporal_params(1), w, h); + let readback_params = temporal_params(1); + let gpu_params = temporal_params(1); + let mut via_readback = NlmDenoiser::::new(&client, readback_params, width, height); + let mut via_gpu = NlmDenoiser::::new(&client, gpu_params, width, height); - for frame in distinct_frames(w, h, 5) { + let frames = distinct_frames(width, height, 5); + for frame in frames { via_readback.push_frame(&frame); via_gpu.push_frame(&frame); - let expected = via_readback.denoise_submit().expect("submit failed"); + let expected = via_readback.denoise().expect("denoise failed"); let actual = via_gpu.denoise_submit_gpu().expect("submit_gpu failed"); match (expected, actual) { (None, None) => {}, - (Some(pending), Some(output)) => { - let expected_frame = pending - .wait() - .expect("wait failed") - .into_f32() - .expect("f32 output"); - + (Some(expected_frame), Some(output)) => { let bytes = client.read_one(output.handle).expect("gpu readback failed"); let actual_frame = f32::from_bytes(&bytes); assert_eq!(actual_frame.len(), expected_frame.len(), "frame length mismatch"); - for (i, (a, b)) in expected_frame.iter().zip(actual_frame.iter()).enumerate() { + + let pairs = expected_frame.iter().zip(actual_frame.iter()).enumerate(); + for (i, (readback_value, gpu_value)) in pairs { assert!( - (a - b).abs() < 1e-6, - "pixel {i}: denoise_submit gave {a}, denoise_submit_gpu gave {b}" + (readback_value - gpu_value).abs() < 1e-6, + "pixel {i}: denoise gave {readback_value}, denoise_submit_gpu gave {gpu_value}" ); } }, - (a, b) => panic!( - "denoise_submit and denoise_submit_gpu disagreed on readiness: {} vs {}", - a.is_some(), - b.is_some() + (readback, gpu) => panic!( + "denoise and denoise_submit_gpu disagreed on readiness: {} vs {}", + readback.is_some(), + gpu.is_some() ), } } } -/// Reads a `GpuOutput` fully back to the host, the same way `flush` -/// reads it back internally, so a test can compare it against a frame -/// produced through the plain readback path. fn read_output(client: &ComputeClient, output: GpuOutput) -> Vec { let bytes = client.read_one(output.handle).expect("gpu readback failed"); f32::from_bytes(&bytes).to_vec() @@ -94,30 +83,38 @@ fn assert_frames_match(expected: &[Vec], actual: &[Vec]) { actual.len(), "flush and flush_step_gpu produced different frame counts" ); - for (frame_idx, (e, a)) in expected.iter().zip(actual.iter()).enumerate() { - assert_eq!(e.len(), a.len(), "frame {frame_idx}: length mismatch"); - for (i, (ev, av)) in e.iter().zip(a.iter()).enumerate() { + + for (frame_idx, (expected_frame, actual_frame)) in expected.iter().zip(actual.iter()).enumerate() { + assert_eq!( + expected_frame.len(), + actual_frame.len(), + "frame {frame_idx}: length mismatch" + ); + + let pairs = expected_frame.iter().zip(actual_frame.iter()).enumerate(); + for (i, (expected_value, actual_value)) in pairs { assert!( - (ev - av).abs() < 1e-6, - "frame {frame_idx}, pixel {i}: flush gave {ev}, flush_step_gpu gave {av}" + (expected_value - actual_value).abs() < 1e-6, + "frame {frame_idx}, pixel {i}: flush gave {expected_value}, flush_step_gpu gave {actual_value}" ); } } } -/// Runs `pushes` frames through two denoisers built from the same -/// parameters, drains one with `flush` and the other by hand-driving -/// `flush_step_gpu` against `flush_target`, and checks the two tails -/// agree frame for frame. +/// Drains one denoiser with `flush` and a twin by hand-driving `flush_step_gpu`, then checks the two +/// tails agree frame for frame. fn compare_flush_paths(radius: u32, pushes: usize) { let client = make_client(); - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut via_flush = NlmDenoiser::::new(&client, temporal_params(radius), w, h); - let mut via_step = NlmDenoiser::::new(&client, temporal_params(radius), w, h); + let flush_params = temporal_params(radius); + let step_params = temporal_params(radius); + let mut via_flush = NlmDenoiser::::new(&client, flush_params, width, height); + let mut via_step = NlmDenoiser::::new(&client, step_params, width, height); - for frame in distinct_frames(w, h, pushes) { + let frames = distinct_frames(width, height, pushes); + for frame in frames { via_flush.push_frame(&frame); let _ = via_flush.denoise().expect("denoise failed"); @@ -126,17 +123,21 @@ fn compare_flush_paths(radius: u32, pushes: usize) { } let mut expected = Vec::new(); - via_flush - .flush(|frame| expected.push(frame.as_f32().expect("f32 denoiser").to_vec())) - .expect("flush failed"); + let collect = |frame: &[f32]| { + let samples = frame.to_vec(); + expected.push(samples); + }; + via_flush.flush(collect).expect("flush failed"); let target = via_step.flush_target(); let mut actual = Vec::new(); while actual.len() < target { if let Some(output) = via_step.flush_step_gpu().expect("flush_step_gpu failed") { - actual.push(read_output(&client, output)); + let frame = read_output(&client, output); + actual.push(frame); } } + assert!( via_step.flush_step_gpu().is_ok(), "a further flush_step_gpu call past the target should still succeed" @@ -152,44 +153,38 @@ fn compare_flush_paths(radius: u32, pushes: usize) { #[test] fn flush_step_gpu_emits_the_same_count_and_frames_short_stream() { - // Radius 2, one push: the window never fills during pushing, so the - // whole drain happens inside the padding phase of the flush. + // One push at radius 2 never fills the window, so the whole drain runs in the padding phase. compare_flush_paths(2, 1); } #[test] fn flush_step_gpu_emits_the_same_count_and_frames_long_stream() { - // Radius 2, five pushes: the window fills while pushing, so the - // drain runs entirely through the trailing-tail phase of the flush. + // Five pushes at radius 2 fill the window, so the whole drain runs in the trailing-tail phase. compare_flush_paths(2, 5); } -/// Radius 2, two pushes. Neither of the two tests above exercises a -/// single drain that needs both behaviours. The short-stream case stops -/// as soon as the window finishes filling, and the long-stream case -/// never fills the window during the drain at all, since it was already -/// full before the drain began. Here the window is exactly one frame -/// short when the drain starts, and the target is two frames, so the -/// first duplicate both finishes filling the window and produces the -/// first output, then a second duplicate runs with the window already -/// full and produces the second. One drain call has to switch from one -/// behaviour to the other partway through, which is what this test -/// pins down. +/// Pins a single drain that switches from the padding phase to the trailing-tail phase. +/// +/// With radius 2 and two pushes the window is one frame short when the drain starts, and the target +/// is two frames. The first step fills the window and emits, then the second runs on a full window. #[test] fn flush_step_gpu_emits_the_same_count_and_frames_mixed_phase_stream() { let radius = 2; let pushes = 2; let client = make_client(); - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut via_flush = NlmDenoiser::::new(&client, temporal_params(radius), w, h); - let mut via_step = NlmDenoiser::::new(&client, temporal_params(radius), w, h); + let flush_params = temporal_params(radius); + let step_params = temporal_params(radius); + let mut via_flush = NlmDenoiser::::new(&client, flush_params, width, height); + let mut via_step = NlmDenoiser::::new(&client, step_params, width, height); + let frames = distinct_frames(width, height, pushes); let mut during_pushes_flush = 0usize; let mut during_pushes_step = 0usize; - for frame in distinct_frames(w, h, pushes) { + for frame in frames { via_flush.push_frame(&frame); if via_flush.denoise().expect("denoise failed").is_some() { during_pushes_flush += 1; @@ -204,6 +199,7 @@ fn flush_step_gpu_emits_the_same_count_and_frames_mixed_phase_stream() { during_pushes_step += 1; } } + assert_eq!( during_pushes_flush, 0, "radius {radius} pushes {pushes}: the window should still be filling, so nothing \ @@ -211,11 +207,7 @@ fn flush_step_gpu_emits_the_same_count_and_frames_mixed_phase_stream() { ); assert_eq!(during_pushes_step, during_pushes_flush); - // Confirm the drain actually straddles both phases rather than just - // assuming it from the chosen parameters. The window has to still be - // short when the drain starts, and filling it in has to leave at - // least one more output for the trailing-tail phase to supply, or - // every output would come from a single phase after all. + // Checks the drain really straddles both phases instead of assuming it from the parameters. let total_frames = via_step.params.total_frames() as usize; let target = via_step.flush_target(); let frames_short = total_frames - via_step.frames_loaded; @@ -233,14 +225,17 @@ fn flush_step_gpu_emits_the_same_count_and_frames_mixed_phase_stream() { ); let mut expected = Vec::new(); - via_flush - .flush(|frame| expected.push(frame.as_f32().expect("f32 denoiser").to_vec())) - .expect("flush failed"); + let collect = |frame: &[f32]| { + let samples = frame.to_vec(); + expected.push(samples); + }; + via_flush.flush(collect).expect("flush failed"); let mut actual = Vec::new(); while actual.len() < target { if let Some(output) = via_step.flush_step_gpu().expect("flush_step_gpu failed") { - actual.push(read_output(&client, output)); + let frame = read_output(&client, output); + actual.push(frame); } } @@ -262,11 +257,12 @@ fn flush_step_gpu_emits_the_same_count_and_frames_mixed_phase_stream() { #[test] fn current_sigmas_reads_zero_on_the_fast_path() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.5, 8, 8, 3, 0.9); + let width = 16; + let height = 16; + let frame = make_frame_with_noisy_region(width, height, 1, 0.5, 8, 8, 3, 0.9); - let mut denoiser = NlmDenoiser::::new(&client, temporal_params(0), w, h); + let params = temporal_params(0); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); assert_eq!(denoiser.current_sigmas(), [0.0, 0.0, 0.0]); denoiser.push_frame(&frame); @@ -281,9 +277,9 @@ fn current_sigmas_reads_zero_on_the_fast_path() { #[test] fn current_sigmas_broadcasts_a_pinned_sigma_override() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.5, 8, 8, 3, 0.9); + let width = 16; + let height = 16; + let frame = make_frame_with_noisy_region(width, height, 1, 0.5, 8, 8, 3, 0.9); let params = NlmParams { hq: Some(HqParams { @@ -297,7 +293,7 @@ fn current_sigmas_broadcasts_a_pinned_sigma_override() { }), ..temporal_params(0) }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); // A pinned sigma reads back immediately, before any frame is pushed. assert_eq!(denoiser.current_sigmas(), [6.0 / 255.0; 3]); @@ -310,9 +306,9 @@ fn current_sigmas_broadcasts_a_pinned_sigma_override() { #[test] fn current_sigmas_matches_the_median_estimator_once_it_folds() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.5, 8, 8, 3, 0.9); + let width = 16; + let height = 16; + let frame = make_frame_with_noisy_region(width, height, 1, 0.5, 8, 8, 3, 0.9); let params = NlmParams { hq: Some(HqParams { @@ -326,7 +322,7 @@ fn current_sigmas_matches_the_median_estimator_once_it_folds() { }), ..temporal_params(0) }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); let _ = denoiser.denoise().expect("denoise failed"); diff --git a/av-denoise-core/src/nlmeans/tests/helpers.rs b/av-denoise-core/src/nlmeans/tests/helpers.rs index a906318..1b87451 100644 --- a/av-denoise-core/src/nlmeans/tests/helpers.rs +++ b/av-denoise-core/src/nlmeans/tests/helpers.rs @@ -1,53 +1,60 @@ use cubecl::prelude::*; +use cubecl::server::Handle; use cubecl::wgpu::WgpuRuntime; +use crate::engine::{DevicePlane, IngestTarget, SampleFormat, ingest}; pub(super) use crate::nlmeans::align::StorageAlign; -#[cfg(feature = "vulkan")] -use crate::{ChannelMode, Denoiser, DenoiserOptions, DenoisingMode, accelerate::Accelerator, device::Device}; +use crate::nlmeans::{ + ChannelMode, + DenoisingMode, + NlmDenoiser, + NlmeansAlgorithm, + NlmeansOptions, + resolve_params, +}; -pub(super) type R = WgpuRuntime; +pub(crate) type R = WgpuRuntime; -pub(super) fn make_client() -> ComputeClient { +pub(crate) fn make_client() -> ComputeClient { let device = ::Device::default(); R::client(&device) } -/// Buffer-binding alignment the test runtime reports, the same value -/// `NlmDenoiser` lays its per-slot buffers out against. +/// Buffer-binding alignment the test runtime reports, which `NlmDenoiser` lays its slots out against. pub(super) fn test_align() -> StorageAlign { - StorageAlign::from_client(&make_client()) + let client = make_client(); + StorageAlign::from_client(&client) } -pub(super) fn make_uniform_frame(w: u32, h: u32, ch: u32, val: f32) -> Vec { - vec![val; (w * h * ch) as usize] +pub(super) fn make_uniform_frame(width: u32, height: u32, channels: u32, value: f32) -> Vec { + vec![value; (width * height * channels) as usize] } -/// Creates a frame with a patch of noise (not just a single pixel) -/// so that NLMeans has matching noisy patches to work with. +/// Creates a frame with a square of noise so NLMeans has matching noisy patches to work with. #[expect( clippy::too_many_arguments, reason = "the test helper takes the full set of parameters its cases vary" )] pub(super) fn make_frame_with_noisy_region( - w: u32, - h: u32, - ch: u32, + width: u32, + height: u32, + channels: u32, base: f32, - cx: u32, - cy: u32, + centre_x: u32, + centre_y: u32, radius: u32, - noise_val: f32, + noise_value: f32, ) -> Vec { - let mut frame = vec![base; (w * h * ch) as usize]; + let mut frame = vec![base; (width * height * channels) as usize]; for dy in 0..=radius * 2 { for dx in 0..=radius * 2 { - let x = cx + dx - radius; - let y = cy + dy - radius; + let x = centre_x + dx - radius; + let y = centre_y + dy - radius; - if x < w && y < h { - for c in 0..ch { - frame[((y * w + x) * ch + c) as usize] = noise_val; + if x < width && y < height { + for channel in 0..channels { + frame[((y * width + x) * channels + channel) as usize] = noise_value; } } } @@ -56,18 +63,23 @@ pub(super) fn make_frame_with_noisy_region( frame } -/// Flat base value plus deterministic pseudo-Gaussian noise, densely -/// packed as `pixels * ch` (no channel padding). Each sample sums four -/// hash-derived uniforms in `[-0.5, 0.5]` (Irwin-Hall) and rescales to -/// the requested per-channel standard deviation, so the same arguments -/// always reproduce the same frame. `sigmas` is indexed per channel and -/// wraps if shorter than `ch`. -pub(super) fn make_noisy_gaussian_frame(w: u32, h: u32, ch: u32, base: f32, sigmas: &[f32]) -> Vec { - let mut frame = vec![0.0f32; (w * h * ch) as usize]; +/// Flat base value plus deterministic pseudo-Gaussian noise, densely packed as `pixels * channels`. +/// +/// Each sample sums four hash-derived uniforms between -0.5 and 0.5 (Irwin-Hall) and rescales to +/// the per-channel standard deviation. `sigmas` is indexed per channel and wraps if shorter than +/// `channels`. +pub(super) fn make_noisy_gaussian_frame( + width: u32, + height: u32, + channels: u32, + base: f32, + sigmas: &[f32], +) -> Vec { + let mut frame = vec![0.0f32; (width * height * channels) as usize]; // Sum of 4 independent Uniform(-0.5, 0.5) samples has variance 4/12 = 1/3. let unit_std = (1.0f32 / 3.0f32).sqrt(); - for idx in 0..(w * h * ch) { + for idx in 0..(width * height * channels) { let mut sum = 0.0f32; for k in 0..4u32 { let mut hash = (idx * 4 + k).wrapping_mul(2654435761).wrapping_add(0x9E3779B9); @@ -76,221 +88,243 @@ pub(super) fn make_noisy_gaussian_frame(w: u32, h: u32, ch: u32, base: f32, sigm hash ^= hash >> 13; sum += (hash as f32 / u32::MAX as f32) - 0.5; } - let c = (idx % ch) as usize; - let sigma = sigmas[c % sigmas.len()]; + + let channel = (idx % channels) as usize; + let sigma = sigmas[channel % sigmas.len()]; frame[idx as usize] = (base + (sum / unit_std) * sigma).clamp(0.0, 1.0); } frame } -/// Adds independent pseudo-Gaussian noise to an arbitrary `w * h` -/// clean field, clamping once after. Each sample sums four -/// hash-derived uniforms in `[-0.5, 0.5]` (Irwin-Hall) and rescales to -/// the requested standard deviation, decorrelated across `seed` values -/// so two calls with different `seed`s over the *same* `clean` field -/// produce two independently noisy copies of it. -pub(super) fn noisy_field_over(clean: &[f32], w: u32, h: u32, sigma: f32, seed: u32) -> Vec { +/// A unit-variance pseudo-Gaussian sample at `idx`, decorrelated across `seed` values. +/// +/// It sums four hash-derived uniforms between -0.5 and 0.5 (Irwin-Hall), whose variance is 1/3. +pub(crate) fn seeded_unit_gaussian(idx: u32, seed: u32) -> f32 { let unit_std = (1.0f32 / 3.0f32).sqrt(); - let mut frame = vec![0.0f32; (w * h) as usize]; - for idx in 0..(w * h) { - let mut sum = 0.0f32; - for k in 0..4u32 { - let mut hash = (idx * 4 + k) - .wrapping_mul(2654435761) - .wrapping_add(seed.wrapping_mul(0x9E37_79B9).wrapping_add(k)); - hash ^= hash >> 15; - hash = hash.wrapping_mul(0x85EB_CA6B); - hash ^= hash >> 13; - sum += (hash as f32 / u32::MAX as f32) - 0.5; - } - frame[idx as usize] = (clean[idx as usize] + (sum / unit_std) * sigma).clamp(0.0, 1.0); + + let mut sum = 0.0f32; + for k in 0..4u32 { + let mut hash = (idx * 4 + k) + .wrapping_mul(2654435761) + .wrapping_add(seed.wrapping_mul(0x9E37_79B9).wrapping_add(k)); + hash ^= hash >> 15; + hash = hash.wrapping_mul(0x85EB_CA6B); + hash ^= hash >> 13; + sum += (hash as f32 / u32::MAX as f32) - 0.5; + } + + sum / unit_std +} + +/// Adds independent pseudo-Gaussian noise to a `width * height` clean field, clamping once after. +/// +/// Two calls with different `seed`s over the same `clean` field produce two independently noisy +/// copies of it. +pub(crate) fn noisy_field_over(clean: &[f32], width: u32, height: u32, sigma: f32, seed: u32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for idx in 0..(width * height) { + let sample = seeded_unit_gaussian(idx, seed); + frame[idx as usize] = (clean[idx as usize] + sample * sigma).clamp(0.0, 1.0); } + frame } -/// Independent pseudo-Gaussian noise sample at `(idx, seed)` over a -/// flat `base` field, decorrelated across `seed` values so two calls -/// with the same `size`/`base`/`sigma` but different `seed`s produce -/// two independently-noisy copies of the same clean signal. A thin -/// wrapper over [`noisy_field_over`] for the common flat-field case. +/// A `size * size` flat `base` field with [noisy_field_over] noise added. pub(super) fn noisy_copy(size: u32, base: f32, sigma: f32, seed: u32) -> Vec { - noisy_field_over(&vec![base; (size * size) as usize], size, size, sigma, seed) + let clean = vec![base; (size * size) as usize]; + noisy_field_over(&clean, size, size, sigma, seed) } -/// Builds a frame of grain that is correlated between neighbouring -/// pixels. -/// -/// An independent white noise field is blurred horizontally with a -/// `[0.25, 0.5, 0.25]` kernel, clamped at the edges, then added to -/// `base` and clamped once more. -/// -/// The blur is identical every call, so two calls with different seeds -/// are independent of each other while each carries the blur's own -/// spatial correlation. +/// Builds a frame of grain correlated between horizontal neighbours by a `[0.25, 0.5, 0.25]` blur. /// -/// That works out at a lag-1 correlation of two thirds for any input -/// distribution, with the variance scaled by 0.375, the sum of the taps -/// squared. -pub(super) fn correlated_noisy_frame(w: u32, h: u32, base: f32, sigma_pre: f32, seed: u32) -> Vec { - let unit_std = (1.0f32 / 3.0f32).sqrt(); - let mut raw = vec![0.0f32; (w * h) as usize]; - for idx in 0..(w * h) { - let mut sum = 0.0f32; - for k in 0..4u32 { - let mut hash = (idx * 4 + k) - .wrapping_mul(2654435761) - .wrapping_add(seed.wrapping_mul(0x9E37_79B9).wrapping_add(k)); - hash ^= hash >> 15; - hash = hash.wrapping_mul(0x85EB_CA6B); - hash ^= hash >> 13; - sum += (hash as f32 / u32::MAX as f32) - 0.5; - } - raw[idx as usize] = (sum / unit_std) * sigma_pre; - } - - let mut out = vec![0.0f32; raw.len()]; - for y in 0..h { - for x in 0..w { - let xl = x.saturating_sub(1); - let xr = (x + 1).min(w - 1); - let l = raw[(y * w + xl) as usize]; - let c = raw[(y * w + x) as usize]; - let r = raw[(y * w + xr) as usize]; - let blurred = 0.25 * l + 0.5 * c + 0.25 * r; - out[(y * w + x) as usize] = (base + blurred).clamp(0.0, 1.0); - } - } - out +/// The blur gives a lag-1 correlation of two thirds for any input distribution and scales the +/// variance by 0.375, the sum of the taps squared. +pub(super) fn correlated_noisy_frame( + width: u32, + height: u32, + base: f32, + sigma_pre: f32, + seed: u32, +) -> Vec { + correlated_noisy_frame_with_tap(width, height, base, sigma_pre, seed, 0.25) } -/// Builds grain correlated the same way [`correlated_noisy_frame`] does, -/// with a tunable horizontal blur tap `a` (kernel `[a, 1 - 2a, a]`) -/// instead of that function's fixed `0.25`. +/// Builds grain correlated by a horizontal `[tap, 1 - 2 * tap, tap]` blur, clamped at the edges. /// -/// For taps `(a, b, a)` applied to unit-variance white noise, the -/// lag-1 correlation along x works out to `2*a*b / (2*a^2 + b^2)`. -/// `a = 0.25` reduces to `correlated_noisy_frame`'s two thirds; `a = -/// 0.125` gives a predicted correlation of about `0.316`, an -/// intermediate value between the uncorrelated and two-thirds cases. -/// Both are confirmed empirically, not just derived, in -/// `residual_correlation.rs`. +/// For taps `(tap, 1 - 2 * tap, tap)` over unit-variance white noise, the lag-1 correlation along x +/// is `2 * tap * (1 - 2 * tap) / (2 * tap^2 + (1 - 2 * tap)^2)`. A tap of `0.125` gives a +/// correlation of about `0.316`. pub(super) fn correlated_noisy_frame_with_tap( - w: u32, - h: u32, + width: u32, + height: u32, base: f32, sigma_pre: f32, seed: u32, - a: f32, + tap: f32, ) -> Vec { - let b = 1.0 - 2.0 * a; - let unit_std = (1.0f32 / 3.0f32).sqrt(); - let mut raw = vec![0.0f32; (w * h) as usize]; - for idx in 0..(w * h) { - let mut sum = 0.0f32; - for k in 0..4u32 { - let mut hash = (idx * 4 + k) - .wrapping_mul(2654435761) - .wrapping_add(seed.wrapping_mul(0x9E37_79B9).wrapping_add(k)); - hash ^= hash >> 15; - hash = hash.wrapping_mul(0x85EB_CA6B); - hash ^= hash >> 13; - sum += (hash as f32 / u32::MAX as f32) - 0.5; - } - raw[idx as usize] = (sum / unit_std) * sigma_pre; + let centre_tap = 1.0 - 2.0 * tap; + + let mut raw = vec![0.0f32; (width * height) as usize]; + for idx in 0..(width * height) { + let sample = seeded_unit_gaussian(idx, seed); + raw[idx as usize] = sample * sigma_pre; } let mut out = vec![0.0f32; raw.len()]; - for y in 0..h { - for x in 0..w { - let xl = x.saturating_sub(1); - let xr = (x + 1).min(w - 1); - let l = raw[(y * w + xl) as usize]; - let c = raw[(y * w + x) as usize]; - let r = raw[(y * w + xr) as usize]; - let blurred = a * l + b * c + a * r; - out[(y * w + x) as usize] = (base + blurred).clamp(0.0, 1.0); + for y in 0..height { + for x in 0..width { + let left_x = x.saturating_sub(1); + let right_x = (x + 1).min(width - 1); + let left = raw[(y * width + left_x) as usize]; + let centre = raw[(y * width + x) as usize]; + let right = raw[(y * width + right_x) as usize]; + let blurred = tap * left + centre_tap * centre + tap * right; + out[(y * width + x) as usize] = (base + blurred).clamp(0.0, 1.0); } } + out } -/// A deterministic frame with real spatial structure at more than one -/// scale, rather than a flat field or a smooth gradient. +/// A deterministic frame with spatial structure at more than one scale. /// -/// Two out-of-phase sine waves plus a finer third one give NLMeans -/// patches with genuinely varying content to match against, so a test -/// built over this frame exercises the same weight spread real footage -/// produces instead of the uniform, always-maximal weights a flat frame -/// hands every candidate. -pub(super) fn make_textured_frame(w: u32, h: u32) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let fx = x as f32 / w as f32; - let fy = y as f32 / h as f32; - let v = 0.5 - + 0.2 * (fx * 8.0 * std::f32::consts::PI).sin() * (fy * 6.0 * std::f32::consts::PI).cos() - + 0.1 * (fx * 20.0 * std::f32::consts::PI).sin(); - frame[(y * w + x) as usize] = v.clamp(0.05, 0.95); +/// Two out-of-phase sine waves plus a finer third one give NLMeans patches with varying content, so +/// the weights spread the way they do on real footage instead of all being maximal. +pub(super) fn make_textured_frame(width: u32, height: u32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let norm_x = x as f32 / width as f32; + let norm_y = y as f32 / height as f32; + let value = 0.5 + + 0.2 + * (norm_x * 8.0 * std::f32::consts::PI).sin() + * (norm_y * 6.0 * std::f32::consts::PI).cos() + + 0.1 * (norm_x * 20.0 * std::f32::consts::PI).sin(); + frame[(y * width + x) as usize] = value.clamp(0.05, 0.95); } } + frame } -/// Smooth horizontal luma gradient from `lo` to `hi` inclusive, -/// replicated down every row. -pub(super) fn make_gradient_frame(w: u32, h: u32, lo: f32, hi: f32) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let t = x as f32 / (w - 1).max(1) as f32; - frame[(y * w + x) as usize] = lo + (hi - lo) * t; +/// Horizontal luma gradient from `low` to `high` inclusive, repeated down every row. +pub(super) fn make_gradient_frame(width: u32, height: u32, low: f32, high: f32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let fraction = x as f32 / (width - 1).max(1) as f32; + frame[(y * width + x) as usize] = low + (high - low) * fraction; } } + frame } -/// Expands a densely packed `pixels * ch` frame into the padded -/// `pixels * stored_ch` GPU storage layout used by `Vector`-typed -/// kernel buffers (extra lanes zeroed). Mirrors the padding -/// `NlmDenoiser::upload_into_slot` applies internally. -/// Builds a top-level [`Denoiser`] at the given temporal radius over -/// Luma, with every other option left at its default, on the Vulkan -/// accelerator. -#[cfg(feature = "vulkan")] -pub(super) fn test_denoiser(radius: u32, w: u32, h: u32) -> Denoiser { - let opts = DenoiserOptions::builder() - .channel_mode(ChannelMode::Luma) - .mode(DenoisingMode::Temporal { radius }) - .build(); - Denoiser::create(&[Accelerator::Vulkan], &Device::Default, w, h, opts) - .expect("denoiser construction failed") +/// Builds a Luma [NlmDenoiser] at the given temporal radius with default options. +pub(super) fn test_denoiser(radius: u32, width: u32, height: u32) -> NlmDenoiser { + let options = NlmeansOptions { + mode: DenoisingMode::Temporal { radius }, + ..NlmeansOptions::default() + }; + let algorithm = NlmeansAlgorithm::Fast(options); + let params = resolve_params(&algorithm, ChannelMode::Luma); + let client = make_client(); + + NlmDenoiser::new(&client, params, width, height) } -/// A deterministic frame whose pixels ramp from `0.2` to `0.8` across -/// the row and shift a little with `i`, so a sequence built from -/// increasing `i` gives every frame in the window distinct content. -pub(super) fn ramp_frame(w: u32, h: u32, i: usize) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let t = (x as f32 + y as f32 * w as f32) / (w * h) as f32; - frame[(y * w + x) as usize] = (0.2 + 0.6 * t + i as f32 * 0.01).clamp(0.0, 1.0); +/// A frame that ramps from `0.2` to `0.8` and shifts a little with `frame_index`. +/// +/// Each `frame_index` gives distinct content. +pub(super) fn ramp_frame(width: u32, height: u32, frame_index: usize) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let fraction = (x as f32 + y as f32 * width as f32) / (width * height) as f32; + frame[(y * width + x) as usize] = + (0.2 + 0.6 * fraction + frame_index as f32 * 0.01).clamp(0.0, 1.0); } } + frame } -pub(super) fn pad_channels(dense: &[f32], pixels: usize, ch: u32, stored_ch: u32) -> Vec { - if ch == stored_ch { +/// Pads a dense `pixels * channels` frame to the `pixels * stored_ch` layout the ingest kernel writes. +/// +/// The extra lanes are zeroed. +pub(super) fn pad_channels(dense: &[f32], pixels: usize, channels: u32, stored_ch: u32) -> Vec { + if channels == stored_ch { return dense.to_vec(); } - let ch = ch as usize; + + let channels = channels as usize; let stored_ch = stored_ch as usize; let mut out = vec![0.0f32; pixels * stored_ch]; - for p in 0..pixels { - out[p * stored_ch..p * stored_ch + ch].copy_from_slice(&dense[p * ch..p * ch + ch]); + for pixel in 0..pixels { + let source = &dense[pixel * channels..pixel * channels + channels]; + out[pixel * stored_ch..pixel * stored_ch + channels].copy_from_slice(source); } + out } + +/// Splits an interleaved host frame into one uploaded plane per channel. +pub(crate) fn upload_planes(client: &ComputeClient, frame: &[f32], channels: usize) -> Vec { + (0..channels) + .map(|channel| { + let plane: Vec = frame.iter().skip(channel).step_by(channels).copied().collect(); + let bytes = f32::as_bytes(&plane); + client.create_from_slice(bytes) + }) + .collect() +} + +pub(crate) fn read_interleaved(client: &ComputeClient, planes: &[Handle]) -> Vec { + let channels: Vec> = planes + .iter() + .map(|handle| { + let bytes = client.read_one(handle.clone()).expect("read plane"); + f32::from_bytes(&bytes).to_vec() + }) + .collect(); + + let pixels = channels[0].len(); + let mut interleaved = Vec::with_capacity(pixels * channels.len()); + for pixel in 0..pixels { + for channel in &channels { + interleaved.push(channel[pixel]); + } + } + + interleaved +} + +/// The f32 samples the real ingest kernel produces for one plane of `u8` codes. +pub(crate) fn normalise_with_ingest( + client: &ComputeClient, + codes: &[u8], + width: u32, + height: u32, +) -> Vec { + let pixels = width * height; + let input = client.create_from_slice(codes); + let planes = [DevicePlane::new(&input, width, height)]; + let placeholder = client.create_from_slice(&[0u8; 4]); + let scratch = client.empty(pixels as usize * 4); + let target = IngestTarget { + ring: &scratch, + ring_len: pixels as usize, + offset: 0, + pixels, + channels: 1, + stored_ch: 1, + }; + + ingest(client, &planes, SampleFormat::U8, &placeholder, target); + + let bytes = client.read_one(scratch).expect("read ingest"); + f32::from_bytes(&bytes).to_vec() +} diff --git a/av-denoise-core/src/nlmeans/tests/hq.rs b/av-denoise-core/src/nlmeans/tests/hq.rs index 5d32c10..3235777 100644 --- a/av-denoise-core/src/nlmeans/tests/hq.rs +++ b/av-denoise-core/src/nlmeans/tests/hq.rs @@ -1,8 +1,7 @@ use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; -/// Shared baseline for the fast path. Each test overrides just the -/// field it's exercising via struct-update syntax. fn base_params() -> NlmParams { NlmParams { temporal_radius: 0, @@ -17,25 +16,58 @@ fn base_params() -> NlmParams { } } -/// With both HQ features off, `effective_strength` and `noise_offset` -/// degenerate to exactly what the fast path already computes, so the -/// two denoisers must agree bit-for-bit. +/// Pushes five frames with a moving noisy square, then flushes. +/// +/// Asserts every output is finite and in range, and that each pushed frame produces one output. +fn run_temporal_smoke(params: NlmParams, width: u32, height: u32) { + let client = make_client(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + + let frames: Vec> = (0..5) + .map(|i| make_frame_with_noisy_region(width, height, 1, 0.5, 6 + i, 8, 2, 0.8)) + .collect(); + + let mut emitted = 0usize; + let check = |frame: &[f32]| { + for (i, &value) in frame.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (0.0..=1.0).contains(&value), + "pixel {i}: out-of-range output {value}" + ); + } + }; + + for frame in &frames { + denoiser.push_frame(frame); + if let Some(result) = denoiser.denoise().unwrap() { + check(&result); + emitted += 1; + } + } + + denoiser + .flush(|frame| { + check(frame); + emitted += 1; + }) + .unwrap(); + + assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); +} + +/// With both features off, `effective_strength` and `noise_offset` reduce to the fast path's values. #[test] fn hq_disabled_features_match_fast_mode() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.5, 8, 8, 3, 0.9); + let width = 16; + let height = 16; + let frame = make_frame_with_noisy_region(width, height, 1, 0.5, 8, 8, 3, 0.9); - let mut fast = NlmDenoiser::::new(&client, base_params(), w, h); + let fast_params = base_params(); + let mut fast = NlmDenoiser::::new(&client, fast_params, width, height); fast.push_frame(&frame); - let fast_out = fast - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let fast_out = fast.denoise().unwrap().unwrap(); let hq_params = NlmParams { hq: Some(HqParams { @@ -49,15 +81,9 @@ fn hq_disabled_features_match_fast_mode() { }), ..base_params() }; - let mut hq = NlmDenoiser::::new(&client, hq_params, w, h); + let mut hq = NlmDenoiser::::new(&client, hq_params, width, height); hq.push_frame(&frame); - let hq_out = hq - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let hq_out = hq.denoise().unwrap().unwrap(); assert_eq!( fast_out, hq_out, @@ -65,37 +91,22 @@ fn hq_disabled_features_match_fast_mode() { ); } -/// Turning on the noise floor shifts every patch distance by a -/// nonzero offset. Neighbours whose distance falls below that offset -/// get clamped to full weight instead of the fast path's decayed -/// weight. That changes their contribution to the weighted average -/// relative to neighbours that stay above the offset, so the two -/// outputs must differ somewhere while staying finite and in range. +/// The sigma sits far above realistic noise because the solid block gives only a few discrete +/// patch distances. /// -/// `sigma` here is set well above the CLI's "heavy noise" guidance. -/// The synthetic frame is a solid block edge rather than real -/// per-pixel noise, so patch distances only take a few discrete -/// values, either zero or a multiple of one mismatched tap's -/// contribution, instead of the small continuum real sensor noise -/// would produce. A small, realistic sigma would sit below every -/// nonzero distance and never clamp anything. The larger sigma exists -/// purely to land the offset between two of those discrete steps. +/// A realistic sigma would sit below every nonzero distance and clamp nothing. This one lands the +/// offset between two of those steps. #[test] fn hq_noise_floor_changes_output() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.5, 8, 8, 3, 0.9); + let width = 16; + let height = 16; + let frame = make_frame_with_noisy_region(width, height, 1, 0.5, 8, 8, 3, 0.9); - let mut fast = NlmDenoiser::::new(&client, base_params(), w, h); + let fast_params = base_params(); + let mut fast = NlmDenoiser::::new(&client, fast_params, width, height); fast.push_frame(&frame); - let fast_out = fast - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let fast_out = fast.denoise().unwrap().unwrap(); let hq_params = NlmParams { hq: Some(HqParams { @@ -109,21 +120,18 @@ fn hq_noise_floor_changes_output() { }), ..base_params() }; - let mut hq = NlmDenoiser::::new(&client, hq_params, w, h); + let mut hq = NlmDenoiser::::new(&client, hq_params, width, height); hq.push_frame(&frame); - let hq_out = hq - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let hq_out = hq.denoise().unwrap().unwrap(); let mut max_diff = 0.0f32; - for (i, (&f, &q)) in fast_out.iter().zip(hq_out.iter()).enumerate() { - assert!(q.is_finite(), "pixel {i}: non-finite HQ output {q}"); - assert!((0.0..=1.0).contains(&q), "pixel {i}: out-of-range HQ output {q}"); - max_diff = max_diff.max((f - q).abs()); + for (i, (&fast_value, &hq_value)) in fast_out.iter().zip(hq_out.iter()).enumerate() { + assert!(hq_value.is_finite(), "pixel {i}: non-finite HQ output {hq_value}"); + assert!( + (0.0..=1.0).contains(&hq_value), + "pixel {i}: out-of-range HQ output {hq_value}" + ); + max_diff = max_diff.max((fast_value - hq_value).abs()); } assert!( @@ -132,95 +140,46 @@ fn hq_noise_floor_changes_output() { ); } -/// Mirrors `spatial::uniform_image_passthrough`. -/// -/// Every patch distance is zero on a flat frame, so both HQ features do -/// nothing and the output has to come back unchanged. #[test] fn hq_uniform_input_passthrough() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 16; + let height = 16; + let frame = make_uniform_frame(width, height, 1, 0.5); let params = NlmParams { hq: Some(HqParams::with_sigma(8.0 / 255.0)), ..base_params() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - for (i, &v) in result.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-5, "pixel {i}: expected 0.5, got {v}"); + let result = denoiser.denoise().unwrap().unwrap(); + + for (i, &value) in result.iter().enumerate() { + assert!((value - 0.5).abs() < 1e-5, "pixel {i}: expected 0.5, got {value}"); } } #[test] fn hq_temporal_smoke() { - let client = make_client(); - let w = 16; - let h = 16; - let params = NlmParams { temporal_radius: 1, hq: Some(HqParams::with_sigma(6.0 / 255.0)), ..base_params() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); - - let frames: Vec> = (0..5) - .map(|i| make_frame_with_noisy_region(w, h, 1, 0.5, 6 + i, 8, 2, 0.8)) - .collect(); - - let mut emitted = 0usize; - let check = |frame: &[f32]| { - for (i, &v) in frame.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((0.0..=1.0).contains(&v), "pixel {i}: out-of-range output {v}"); - } - }; - - for frame in &frames { - denoiser.push_frame(frame); - if let Some(result) = denoiser.denoise().unwrap() { - check(result.as_f32().expect("f32 denoiser")); - emitted += 1; - } - } - - denoiser - .flush(|frame| { - check(frame.as_f32().expect("f32 denoiser")); - emitted += 1; - }) - .unwrap(); - - assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); + run_temporal_smoke(params, 16, 16); } -/// `sigma_override: None` measures the noise level from the pushed -/// frame instead of requiring a caller-supplied value, and the -/// measured sigma must still drive real denoising. -/// -/// Uses per-pixel Gaussian noise rather than a solid noisy block. The -/// Immerkær estimator responds to genuine high-frequency variation. A -/// single flat block only disturbs its boundary ring, so it reads back -/// close to the noise floor and barely denoises anything. +/// Uses per-pixel Gaussian noise because the Immerkær estimator reads a solid block close to the +/// noise floor, since only its boundary ring varies. #[test] fn hq_auto_sigma_denoises() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[8.0 / 255.0]); + let width = 32; + let height = 32; + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[8.0 / 255.0]); let params = NlmParams { hq: Some(HqParams { @@ -235,15 +194,9 @@ fn hq_auto_sigma_denoises() { ..base_params() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); let mut max_diff = 0.0f32; for (i, (&input, &output)) in frame.iter().zip(result.iter()).enumerate() { @@ -261,14 +214,8 @@ fn hq_auto_sigma_denoises() { ); } -/// Mirrors `hq_temporal_smoke` but with the noise level measured -/// automatically instead of supplied up front. #[test] fn hq_auto_sigma_temporal_smoke() { - let client = make_client(); - let w = 16; - let h = 16; - let params = NlmParams { temporal_radius: 1, hq: Some(HqParams { @@ -283,50 +230,21 @@ fn hq_auto_sigma_temporal_smoke() { ..base_params() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); - - let frames: Vec> = (0..5) - .map(|i| make_frame_with_noisy_region(w, h, 1, 0.5, 6 + i, 8, 2, 0.8)) - .collect(); - - let mut emitted = 0usize; - let check = |frame: &[f32]| { - for (i, &v) in frame.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((0.0..=1.0).contains(&v), "pixel {i}: out-of-range output {v}"); - } - }; - - for frame in &frames { - denoiser.push_frame(frame); - if let Some(result) = denoiser.denoise().unwrap() { - check(result.as_f32().expect("f32 denoiser")); - emitted += 1; - } - } - - denoiser - .flush(|frame| { - check(frame.as_f32().expect("f32 denoiser")); - emitted += 1; - }) - .unwrap(); - - assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); + run_temporal_smoke(params, 16, 16); } #[test] fn hq_override_skips_estimation() { let client = make_client(); - let w = 16; - let h = 16; + let width = 16; + let height = 16; let params = NlmParams { hq: Some(HqParams::with_sigma(8.0 / 255.0)), ..base_params() }; - let denoiser = NlmDenoiser::::new(&client, params, w, h); + let denoiser = NlmDenoiser::::new(&client, params, width, height); assert!( denoiser.noise_partials.is_none(), @@ -338,20 +256,13 @@ fn hq_override_skips_estimation() { ); } -/// `reset_stream_state` (called by `flush`) must clear the noise -/// estimator's EMA so a new stream doesn't inherit the previous -/// stream's noise level. Verified by observable behaviour. Pushing a -/// low-noise frame right after a reset must derive exactly the same -/// `h2_inv_norm` / `noise_offset` as a brand-new denoiser that only -/// ever saw that frame, instead of a value blended with the earlier -/// high-noise estimate. #[test] fn hq_reset_clears_noise_state() { let client = make_client(); - let w = 16; - let h = 16; - let noisy = make_frame_with_noisy_region(w, h, 1, 0.5, 8, 8, 3, 0.9); - let low = make_uniform_frame(w, h, 1, 0.5); + let width = 16; + let height = 16; + let noisy = make_frame_with_noisy_region(width, height, 1, 0.5, 8, 8, 3, 0.9); + let low = make_uniform_frame(width, height, 1, 0.5); let params = NlmParams { hq: Some(HqParams { @@ -366,7 +277,7 @@ fn hq_reset_clears_noise_state() { ..base_params() }; - let mut denoiser = NlmDenoiser::::new(&client, params.clone(), w, h); + let mut denoiser = NlmDenoiser::::new(&client, params.clone(), width, height); denoiser.push_frame(&noisy); denoiser.denoise().unwrap(); @@ -374,7 +285,7 @@ fn hq_reset_clears_noise_state() { denoiser.push_frame(&low); denoiser.denoise().unwrap(); - let mut fresh = NlmDenoiser::::new(&client, params, w, h); + let mut fresh = NlmDenoiser::::new(&client, params, width, height); fresh.push_frame(&low); fresh.denoise().unwrap(); @@ -388,16 +299,8 @@ fn hq_reset_clears_noise_state() { ); } -/// HQ auto-σ combined with the nlm-spatial pilot over a short temporal -/// sequence. Mirrors `hq_auto_sigma_temporal_smoke` but with the pilot -/// enabled. Every produced frame must stay finite and in range, and the -/// pipeline must still emit exactly one output per pushed frame. #[test] fn hq_pilot_temporal_end_to_end() { - let client = make_client(); - let w = 16; - let h = 16; - let params = NlmParams { temporal_radius: 1, prefilter: PrefilterMode::NlmSpatial { strength_scale: 1.0 }, @@ -413,52 +316,17 @@ fn hq_pilot_temporal_end_to_end() { ..base_params() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); - - let frames: Vec> = (0..5) - .map(|i| make_frame_with_noisy_region(w, h, 1, 0.5, 6 + i, 8, 2, 0.8)) - .collect(); - - let mut emitted = 0usize; - let check = |frame: &[f32]| { - for (i, &v) in frame.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((0.0..=1.0).contains(&v), "pixel {i}: out-of-range output {v}"); - } - }; - - for frame in &frames { - denoiser.push_frame(frame); - if let Some(result) = denoiser.denoise().unwrap() { - check(result.as_f32().expect("f32 denoiser")); - emitted += 1; - } - } - - denoiser - .flush(|frame| { - check(frame.as_f32().expect("f32 denoiser")); - emitted += 1; - }) - .unwrap(); - - assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); + run_temporal_smoke(params, 16, 16); } -/// The nlm-spatial pilot changes what the main pass reads as its -/// distance signal, so HQ with the pilot enabled must diverge from -/// plain HQ (`prefilter: None`) on the same input. -/// -/// Uses per-pixel Gaussian noise rather than a solid noisy block (as -/// `hq_noise_floor_changes_output` explains, a flat block only -/// disturbs its boundary ring, leaving patch distances elsewhere -/// identical whether or not the pilot ran). +/// Uses per-pixel Gaussian noise because a solid block only changes patch distances around its +/// boundary ring, whether or not the pilot ran. #[test] fn hq_pilot_differs_from_unguided() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[10.0 / 255.0]); + let width = 32; + let height = 32; + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[10.0 / 255.0]); let hq_params = |prefilter: PrefilterMode| NlmParams { prefilter, @@ -474,39 +342,28 @@ fn hq_pilot_differs_from_unguided() { ..base_params() }; - let mut unguided = NlmDenoiser::::new(&client, hq_params(PrefilterMode::None), w, h); + let unguided_params = hq_params(PrefilterMode::None); + let mut unguided = NlmDenoiser::::new(&client, unguided_params, width, height); unguided.push_frame(&frame); - let unguided_out = unguided - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - let mut piloted = NlmDenoiser::::new( - &client, - hq_params(PrefilterMode::NlmSpatial { strength_scale: 1.0 }), - w, - h, - ); + let unguided_out = unguided.denoise().unwrap().unwrap(); + + let piloted_params = hq_params(PrefilterMode::NlmSpatial { strength_scale: 1.0 }); + let mut piloted = NlmDenoiser::::new(&client, piloted_params, width, height); piloted.push_frame(&frame); - let piloted_out = piloted - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let piloted_out = piloted.denoise().unwrap().unwrap(); let mut max_diff = 0.0f32; - for (i, (&a, &b)) in unguided_out.iter().zip(piloted_out.iter()).enumerate() { - assert!(b.is_finite(), "pixel {i}: non-finite piloted output {b}"); + let pairs = unguided_out.iter().zip(piloted_out.iter()).enumerate(); + for (i, (&unguided_value, &piloted_value)) in pairs { assert!( - (0.0..=1.0).contains(&b), - "pixel {i}: out-of-range piloted output {b}" + piloted_value.is_finite(), + "pixel {i}: non-finite piloted output {piloted_value}" ); - max_diff = max_diff.max((a - b).abs()); + assert!( + (0.0..=1.0).contains(&piloted_value), + "pixel {i}: out-of-range piloted output {piloted_value}" + ); + max_diff = max_diff.max((unguided_value - piloted_value).abs()); } assert!( @@ -515,12 +372,11 @@ fn hq_pilot_differs_from_unguided() { ); } -/// HQ temporal params tuned so a uniformly mismatched neighbour sits -/// well past `thsad` (confidence collapses to ~0) while the plain NLM -/// patch weight for that same neighbour stays significant on its own. -/// `strength`/`patch_radius` are chosen so the two thresholds don't -/// coincide, otherwise confidence toggling wouldn't be separable from -/// the intrinsic Welsch suppression. +/// HQ temporal params where a uniformly mismatched neighbour sits well past `thsad` while its plain +/// NLM weight stays significant. +/// +/// `strength` and `patch_radius` keep the two thresholds apart, so confidence is separable from the +/// Welsch suppression. fn temporal_conf_params(temporal_confidence: bool) -> NlmParams { NlmParams { temporal_radius: 1, @@ -543,81 +399,56 @@ fn temporal_conf_params(temporal_confidence: bool) -> NlmParams { } } -/// Confidence weighting has to suppress a mismatched neighbour while -/// leaving a matching one alone. -/// -/// The centre frame and the following neighbour are flat at 0.5. The -/// preceding neighbour is flat at 0.55, mismatched everywhere by 0.05, -/// which is well past the default threshold at the library's block -/// size. -/// -/// Without confidence, plain NLM weighting still gives that neighbour a -/// noticeable weight, because patch distances stay small at this -/// strength, and the output drifts away from 0.5. +/// The previous neighbour is flat at 0.55, a 0.05 mismatch well past the default threshold. /// -/// With confidence, its block-level mismatch collapses that weight to -/// almost nothing and the output stays at 0.5. -/// -/// This is the same style of comparison -/// `temporal_asymmetric_frames_correct_weights` uses for the fast path, -/// measured here against confidence being off rather than against a -/// hand-computed target. +/// Without confidence, plain NLM still gives it a noticeable weight at this strength and the output +/// drifts from 0.5. #[test] fn hq_temporal_confidence_suppresses_mismatched_neighbour() { let client = make_client(); - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let prev = make_uniform_frame(w, h, 1, 0.55); - let center = make_uniform_frame(w, h, 1, 0.5); - let next = make_uniform_frame(w, h, 1, 0.5); + let previous = make_uniform_frame(width, height, 1, 0.55); + let centre = make_uniform_frame(width, height, 1, 0.5); + let next = make_uniform_frame(width, height, 1, 0.5); let run = |temporal_confidence: bool| { - let mut d = NlmDenoiser::::new(&client, temporal_conf_params(temporal_confidence), w, h); - d.push_frame(&prev); - d.push_frame(¢er); - d.push_frame(&next); - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() + let params = temporal_conf_params(temporal_confidence); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&previous); + denoiser.push_frame(¢re); + denoiser.push_frame(&next); + denoiser.denoise().unwrap().unwrap() }; let off = run(false); let on = run(true); - let off_dev = (off[(8 * w + 8) as usize] - 0.5).abs(); - let on_dev = (on[(8 * w + 8) as usize] - 0.5).abs(); + let off_deviation = (off[(8 * width + 8) as usize] - 0.5).abs(); + let on_deviation = (on[(8 * width + 8) as usize] - 0.5).abs(); assert!( - off_dev > 5e-3, + off_deviation > 5e-3, "without confidence weighting the mismatched neighbour should pull the \ - output measurably away from 0.5, got deviation {off_dev}" + output measurably away from 0.5, got deviation {off_deviation}" ); assert!( - on_dev < off_dev * 0.5, + on_deviation < off_deviation * 0.5, "confidence weighting should suppress the mismatched neighbour's \ - contribution: off deviation {off_dev}, on deviation {on_dev}" + contribution: off deviation {off_deviation}, on deviation {on_deviation}" ); } -/// HQ with `temporal_confidence` set to `false` must ignore the -/// confidence machinery entirely, not just apply a weak version of it. -/// `thsad_scale` only ever feeds the confidence threshold (see -/// `HqParams::thsad_scale`), so if the disabled kernel path genuinely -/// never reads the confidence buffer, sweeping it must leave the -/// output bitwise unchanged. This is the observable form of "compiles -/// to the same code as before confidence consumption existed" for a -/// config the plan requires to match prior HQ output exactly. +/// `thsad_scale` only feeds the confidence threshold, so with confidence off a sweep must leave the +/// output bitwise unchanged. #[test] fn hq_temporal_confidence_disabled_ignores_thsad_scale() { let client = make_client(); - let w = 16; - let h = 16; + let width = 16; + let height = 16; let frames: Vec> = (0..3) - .map(|i| make_frame_with_noisy_region(w, h, 1, 0.5, 6 + i, 8, 2, 0.8)) + .map(|i| make_frame_with_noisy_region(width, height, 1, 0.5, 6 + i, 8, 2, 0.8)) .collect(); let run = |thsad_scale: f32| { @@ -640,16 +471,13 @@ fn hq_temporal_confidence_disabled_ignores_thsad_scale() { windowed_noise_estimation: false, }), }; - let mut d = NlmDenoiser::::new(&client, params, w, h); + + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); for frame in &frames { - d.push_frame(frame); + denoiser.push_frame(frame); } - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() + + denoiser.denoise().unwrap().unwrap() }; let base = run(1.0); @@ -661,17 +489,8 @@ fn hq_temporal_confidence_disabled_ignores_thsad_scale() { ); } -/// End-to-end smoke test covering HQ temporal denoising with motion -/// compensation and (default-on) confidence weighting together, over a -/// short synthetic sequence. Every produced frame must stay finite and -/// in range, mirroring `hq_auto_sigma_temporal_smoke` with MC layered -/// on top. #[test] fn hq_temporal_mc_confidence_smoke() { - let client = make_client(); - let w = 32; - let h = 32; - let params = NlmParams { temporal_radius: 1, motion_compensation: MotionCompensationMode::Mvtools { @@ -685,51 +504,19 @@ fn hq_temporal_mc_confidence_smoke() { ..base_params() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); - - let frames: Vec> = (0..5) - .map(|i| make_frame_with_noisy_region(w, h, 1, 0.5, 6 + i, 8, 2, 0.8)) - .collect(); - - let mut emitted = 0usize; - let check = |frame: &[f32]| { - for (i, &v) in frame.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((0.0..=1.0).contains(&v), "pixel {i}: out-of-range output {v}"); - } - }; - - for frame in &frames { - denoiser.push_frame(frame); - if let Some(result) = denoiser.denoise().unwrap() { - check(result.as_f32().expect("f32 denoiser")); - emitted += 1; - } - } - - denoiser - .flush(|frame| { - check(frame.as_f32().expect("f32 denoiser")); - emitted += 1; - }) - .unwrap(); - - assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); + run_temporal_smoke(params, 32, 32); } -/// `sigma_scale` multiplies each channel's raw sigma before it folds -/// into the running EMA, so doubling it must exactly double the folded -/// estimator state (the EMA's first sample sets its state directly, -/// with no nonlinearity to round off) and quadruple `noise_offset` -/// (quadratic in sigma). Checking both consumers from a single fold -/// proves the multiply sits before the blend feeds either of them, -/// rather than one of them picking it up incidentally. +/// The EMA's first sample sets its state directly, so doubling `sigma_scale` exactly doubles the +/// folded estimate and quadruples `noise_offset`. +/// +/// Checking both consumers from one fold shows the multiply sits before the blend feeds either. #[test] fn hq_sigma_scale_multiplies_the_folded_estimate() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[8.0 / 255.0]); + let width = 32; + let height = 32; + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[8.0 / 255.0]); let run = |sigma_scale: f32| { let params = NlmParams { @@ -744,14 +531,16 @@ fn hq_sigma_scale_multiplies_the_folded_estimate() { }), ..base_params() }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.denoise().unwrap(); - let folded = d + + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.denoise().unwrap(); + + let folded = denoiser .noise_estimator .current() .expect("estimator should hold a value after one push")[0]; - (folded, d.noise_offset) + (folded, denoiser.noise_offset) }; let (folded_1x, offset_1x) = run(1.0); diff --git a/av-denoise-core/src/nlmeans/tests/machinery.rs b/av-denoise-core/src/nlmeans/tests/machinery.rs index e62d9e8..e401bd8 100644 --- a/av-denoise-core/src/nlmeans/tests/machinery.rs +++ b/av-denoise-core/src/nlmeans/tests/machinery.rs @@ -1,43 +1,34 @@ -//! `submit_machinery` runs the NLM denoiser's ring, motion, and -//! confidence machinery without launching any NLM denoising kernel, so -//! a separate collaborative stage can read the same ring, motion -//! fields, and confidence scores the NLM path builds. -//! -//! These tests pin that the returned [`RingView`] geometry and content -//! line up with what a real, non-trivial motion sequence produces. - use cubecl::prelude::*; use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::motion::neighbour_idx_for_k; use crate::nlmeans::*; const RADIUS: u32 = 2; const SIZE: u32 = 128; -/// A single `size x size` noisy world pattern, read at a shifting -/// horizontal offset and clamped at the edges, so a sequence built from -/// increasing `shift` values translates the same content one pixel to -/// the right per frame. +// The motion module's default block size and overlap, giving the `step = blksize - overlap = 8` +// geometry the assertions assume. +const DEFAULT_BLKSIZE_FOR_TEST: u32 = 16; +const DEFAULT_OVERLAP_FOR_TEST: u32 = 8; + +/// A noisy world read at a horizontal offset of `shift` and clamped at the edges. /// -/// A flat gradient (varying only along x) was tried first and rejected. -/// Its SAD is identical at every vertical candidate offset, the same -/// degeneracy `block_match.rs`'s tie-break comments describe for a -/// uniform block, and the fine pass's tie-break resolves that to -/// whatever the coarse pass seeded rather than to zero, letting a wrong -/// vertical offset slip through the "within 1 px" assertion below -/// undetected (confirmed empirically while writing this test). Dense 2D -/// noise gives every candidate a distinct score, so the block match has -/// a genuine, unambiguous minimum at the planted shift. +/// Increasing `shift` by one moves the content one pixel right. Dense 2D noise gives every candidate +/// offset a distinct SAD, so the block match has one clear minimum at the planted shift. Content that +/// varies only along x ties every vertical offset, and the tie-break can then let a wrong vertical +/// offset pass the 1 px checks. fn translating_frame(size: u32, shift: i32) -> Vec { let world = noisy_copy(size, 0.5, 0.2, 777); let mut frame = vec![0.0f32; (size * size) as usize]; for y in 0..size { for x in 0..size { - let sx = (x as i32 - shift).clamp(0, size as i32 - 1) as u32; - frame[(y * size + x) as usize] = world[(y * size + sx) as usize]; + let source_x = (x as i32 - shift).clamp(0, size as i32 - 1) as u32; + frame[(y * size + x) as usize] = world[(y * size + source_x) as usize]; } } + frame } @@ -61,40 +52,32 @@ fn machinery_params() -> NlmParams { } } -// Mirrors `motion::DEFAULT_BLKSIZE`/`DEFAULT_OVERLAP`, spelled out locally -// so this file does not need to reach into the `motion` module just for -// two constants already fixed by the geometry the assertions below -// assume (`step = blksize - overlap = 8`). -const DEFAULT_BLKSIZE_FOR_TEST: u32 = 16; -const DEFAULT_OVERLAP_FOR_TEST: u32 = 8; - -/// Pushes `2 * RADIUS + 1` frames of a one-pixel-per-frame translating -/// world, exactly filling the temporal window, and returns the built -/// denoiser. +/// Pushes `2 * RADIUS + 1` frames of a world translating one pixel per frame, exactly filling the +/// window. fn push_translating_sequence(client: &ComputeClient) -> NlmDenoiser { - let mut d = NlmDenoiser::::new(client, machinery_params(), SIZE, SIZE); + let params = machinery_params(); + let mut denoiser = NlmDenoiser::::new(client, params, SIZE, SIZE); let total_frames = 2 * RADIUS + 1; - for n in 0..total_frames { - let frame = translating_frame(SIZE, n as i32); - d.push_frame(&frame); + for frame_index in 0..total_frames { + let frame = translating_frame(SIZE, frame_index as i32); + denoiser.push_frame(&frame); } - d + + denoiser } #[test] fn submit_machinery_reports_ring_view_with_correct_motion_and_confidence() { let client = make_client(); - let mut d = push_translating_sequence(&client); + let mut denoiser = push_translating_sequence(&client); - let view = d + let view = denoiser .submit_machinery(RADIUS) .expect("submit_machinery dispatch failed") .expect("window is exactly full, submit_machinery should report Some"); - // The centre slot differs from every neighbour slot. With exactly - // `2 * RADIUS + 1` pushes into a same-sized ring, every frame landed - // in its own distinct physical slot, so this also confirms the ring - // never doubled a slot up. + // With exactly `2 * RADIUS + 1` pushes into a ring of that size, every frame lands in its own slot, + // so this also confirms the ring never doubled a slot up. for &slot in &view.neighbour_slots { assert_ne!( slot, view.centre_slot, @@ -108,38 +91,38 @@ fn submit_machinery_reports_ring_view_with_correct_motion_and_confidence() { "one neighbour slot per non-zero k in -RADIUS..=RADIUS" ); - let mc = d.motion_ctx(); - let bx = (64 / mc.step).min(mc.blocks_x - 1); - let by = (64 / mc.step).min(mc.blocks_y - 1); + let motion = denoiser.motion_ctx(); + let block_x = (64 / motion.step).min(motion.blocks_x - 1); + let block_y = (64 / motion.step).min(motion.blocks_y - 1); - let nidx = neighbour_idx_for_k(RADIUS, 1); - let mv_idx = (nidx * view.mv_stride + (by * mc.blocks_x + bx) * 2) as usize; - let mv_bytes = d + let neighbour_idx = neighbour_idx_for_k(RADIUS, 1); + let mv_idx = (neighbour_idx * view.mv_stride + (block_y * motion.blocks_x + block_x) * 2) as usize; + let mv_field = view.mv_field.clone(); + let mv_bytes = denoiser .compute_client() - .read_one(view.mv_field.clone()) + .read_one(mv_field) .expect("mv_field readback failed"); - let mv = i32::from_bytes(&mv_bytes); + let motion_vectors = i32::from_bytes(&mv_bytes); - // The sequence translates by exactly one pixel per frame, so the - // immediate forward neighbour (k = 1) moved by exactly (1, 0) - // relative to the centre. + // The world shifts one pixel per frame, so the forward neighbour at k = 1 moved by exactly (1, 0). assert!( - (mv[mv_idx] - 1).abs() <= 1, + (motion_vectors[mv_idx] - 1).abs() <= 1, "expected mv.x within 1px of the planted shift of 1, got {}", - mv[mv_idx] + motion_vectors[mv_idx] ); assert!( - mv[mv_idx + 1].abs() <= 1, + motion_vectors[mv_idx + 1].abs() <= 1, "expected mv.y within 1px of the planted shift of 0, got {}", - mv[mv_idx + 1] + motion_vectors[mv_idx + 1] ); - let conf_idx = (nidx * view.conf_stride + (by * mc.blocks_x + bx)) as usize; - let conf_bytes = d + let confidence_idx = (neighbour_idx * view.conf_stride + (block_y * motion.blocks_x + block_x)) as usize; + let confidence_field = view.confidence.clone(); + let confidence_bytes = denoiser .compute_client() - .read_one(view.confidence.clone()) + .read_one(confidence_field) .expect("confidence readback failed"); - let confidence = f32::from_bytes(&conf_bytes)[conf_idx]; + let confidence = f32::from_bytes(&confidence_bytes)[confidence_idx]; assert!( confidence.is_finite() && (0.0..=1.0).contains(&confidence), @@ -154,107 +137,107 @@ fn submit_machinery_reports_ring_view_with_correct_motion_and_confidence() { #[test] fn submit_machinery_at_centre_zero_lists_every_later_slot() { let client = make_client(); - let mut d = push_translating_sequence(&client); + let mut denoiser = push_translating_sequence(&client); let total_frames = 2 * RADIUS + 1; - let view = d + let view = denoiser .submit_machinery(0) .expect("submit_machinery dispatch failed") .expect("window is exactly full"); - let expected: Vec = (1..total_frames).map(|logical| d.ring_slot(logical)).collect(); - assert_eq!(view.centre_slot, d.ring_slot(0)); + let expected: Vec = (1..total_frames) + .map(|logical| denoiser.ring_slot(logical)) + .collect(); + let centre_slot = denoiser.ring_slot(0); + assert_eq!(view.centre_slot, centre_slot); assert_eq!(view.neighbour_slots, expected); } #[test] fn submit_machinery_at_the_last_slot_lists_every_earlier_slot() { let client = make_client(); - let mut d = push_translating_sequence(&client); + let mut denoiser = push_translating_sequence(&client); let last = 2 * RADIUS; - let view = d + let view = denoiser .submit_machinery(last) .expect("submit_machinery dispatch failed") .expect("window is exactly full"); - let expected: Vec = (0..last).map(|logical| d.ring_slot(logical)).collect(); - assert_eq!(view.centre_slot, d.ring_slot(last)); + let expected: Vec = (0..last).map(|logical| denoiser.ring_slot(logical)).collect(); + let centre_slot = denoiser.ring_slot(last); + assert_eq!(view.centre_slot, centre_slot); assert_eq!(view.neighbour_slots, expected); } -/// From centre 0 the translating world moves one pixel right per frame, -/// so the neighbour at logical `2 * RADIUS` sits `2 * RADIUS` frames -/// ahead. That neighbour is the last one submitted, landing at field -/// index `2 * RADIUS - 1`. +/// From centre 0 the neighbour at logical `2 * RADIUS` sits `2 * RADIUS` pixels to the right. +/// +/// It is the last neighbour submitted, so it lands at field index `2 * RADIUS - 1`. #[test] fn submit_machinery_at_centre_zero_finds_motion_at_the_far_offset() { let client = make_client(); - let mut d = push_translating_sequence(&client); + let mut denoiser = push_translating_sequence(&client); let far = 2 * RADIUS; - let view = d + let view = denoiser .submit_machinery(0) .expect("submit_machinery dispatch failed") .expect("window is exactly full"); - let mc = d.motion_ctx(); - let bx = (64 / mc.step).min(mc.blocks_x - 1); - let by = (64 / mc.step).min(mc.blocks_y - 1); + let motion = denoiser.motion_ctx(); + let block_x = (64 / motion.step).min(motion.blocks_x - 1); + let block_y = (64 / motion.step).min(motion.blocks_y - 1); let neighbour_idx = far - 1; - let mv_idx = (neighbour_idx * view.mv_stride + (by * mc.blocks_x + bx) * 2) as usize; - let mv_bytes = d + let mv_idx = (neighbour_idx * view.mv_stride + (block_y * motion.blocks_x + block_x) * 2) as usize; + let mv_field = view.mv_field.clone(); + let mv_bytes = denoiser .compute_client() - .read_one(view.mv_field.clone()) + .read_one(mv_field) .expect("mv_field readback failed"); - let mv = i32::from_bytes(&mv_bytes); + let motion_vectors = i32::from_bytes(&mv_bytes); assert!( - (mv[mv_idx] - far as i32).abs() <= 1, + (motion_vectors[mv_idx] - far as i32).abs() <= 1, "expected mv.x within 1px of the planted shift of {far}, got {}", - mv[mv_idx] + motion_vectors[mv_idx] ); } -/// The ring view carries the pyramid the estimator analysed and the -/// window size, and the front end reports the SAD noise floor it scored -/// confidence with. #[test] fn ring_view_exposes_the_analysed_pyramid_and_the_noise_floor() { let client = make_client(); - let mut d = push_translating_sequence(&client); - let view = d + let mut denoiser = push_translating_sequence(&client); + let view = denoiser .submit_machinery(RADIUS) .expect("submit_machinery dispatch failed") .expect("window is exactly full, submit_machinery should report Some"); let frames = machinery_params().total_frames(); assert_eq!(view.frame_count, frames); - // The pyramid holds every level of every slot, so it is at least - // one full-resolution luma plane per slot. - let bytes = client - .read_one(view.pyramid.clone()) - .expect("pyramid readback failed"); + + // The pyramid holds every level of every slot, so it is at least one full-resolution luma plane + // per slot. + let pyramid = view.pyramid.clone(); + let bytes = client.read_one(pyramid).expect("pyramid readback failed"); let plane = bytes.len() / (frames as usize * size_of::()); assert!(plane > 0, "the pyramid must hold at least one plane per slot"); - assert!(d.sad_noise_floor_value() >= 0.0); + assert!(denoiser.sad_noise_floor_value() >= 0.0); } -/// A window that has not filled yet reports `None`, the same convention -/// `denoise_submit_gpu` uses. #[test] fn submit_machinery_none_while_window_is_filling() { let client = make_client(); - let mut d = NlmDenoiser::::new(&client, machinery_params(), SIZE, SIZE); + let params = machinery_params(); + let mut denoiser = NlmDenoiser::::new(&client, params, SIZE, SIZE); // Fewer than `2 * RADIUS + 1` pushes, so the window never fills. - for n in 0..RADIUS { - let frame = translating_frame(SIZE, n as i32); - d.push_frame(&frame); + for frame_index in 0..RADIUS { + let frame = translating_frame(SIZE, frame_index as i32); + denoiser.push_frame(&frame); } - let result = d + let result = denoiser .submit_machinery(RADIUS) .expect("submit_machinery dispatch failed"); assert!( @@ -263,89 +246,32 @@ fn submit_machinery_none_while_window_is_filling() { ); } -/// Priming a whole window with [`Denoiser::push_frame_priming`], then -/// submitting only the last real push, must produce the same frame the -/// streaming path emits for the window's centre. Filling the window this -/// way is what a caller with random-order access to a fixed window, such -/// as a VapourSynth plugin, needs to reseed on every frame request -/// instead of pushing one frame at a time in order. #[cfg(feature = "vulkan")] #[test] fn priming_pushes_then_one_submit_matches_the_streaming_centre() { - let r = 2u32; - let window: Vec> = (0..(2 * r + 1) as usize).map(|i| ramp_frame(64, 64, i)).collect(); - - let mut windowed = test_denoiser(r, 64, 64); - for frame in &window[..(2 * r) as usize] { - windowed.push_frame_priming(frame).unwrap(); - } - windowed.push_frame(&window[(2 * r) as usize]).unwrap(); - let got = windowed.recv_frame().unwrap().expect("one frame"); - - let mut streamed = test_denoiser(r, 64, 64); - let mut emitted = Vec::new(); - for frame in &window { - streamed.push_frame(frame).unwrap(); - if let Some(out) = streamed.recv_frame().unwrap() { - emitted.push(out); - } + let radius = 2u32; + let window: Vec> = (0..(2 * radius + 1) as usize) + .map(|i| ramp_frame(64, 64, i)) + .collect(); + + let mut windowed = test_denoiser(radius, 64, 64); + for frame in &window[..(2 * radius) as usize] { + windowed.push_frame(frame); } - assert_eq!(emitted.len(), (r + 1) as usize); - assert_eq!(got, emitted[r as usize]); -} - -#[cfg(feature = "vulkan")] -#[test] -fn try_recv_frame_returns_none_when_nothing_is_in_flight() { - let mut d = test_denoiser(2, 64, 64); - assert_eq!(d.try_recv_frame().unwrap(), None); -} + windowed.push_frame(&window[(2 * radius) as usize]); + let got = windowed.denoise().unwrap().expect("one frame"); -#[cfg(feature = "vulkan")] -#[test] -fn try_recv_frame_observes_a_landed_readback_within_a_bounded_poll() { - // A poll count is the wrong proxy for the wall-clock interval this - // test needs to cover (cold pipeline compile plus dispatch plus - // readback), since a faster CPU makes each poll cheaper and so - // needs *more* of them for the same GPU latency. A deadline covers - // both a slow GPU and a fast CPU the same way. - const DEADLINE: std::time::Duration = std::time::Duration::from_secs(30); - - let r = 2u32; - // `temporal_radius + 1` pushes is exactly enough to prime the window - // and submit one denoise, leaving exactly one readback in flight. - let window: Vec> = (0..(r + 1) as usize).map(|i| ramp_frame(64, 64, i)).collect(); - - let mut polled = test_denoiser(r, 64, 64); + let mut streamed = test_denoiser(radius, 64, 64); + let mut emitted = Vec::new(); for frame in &window { - polled.push_frame(frame).unwrap(); - } + streamed.push_frame(frame); - let start = std::time::Instant::now(); - let mut got = None; - let mut polls = 0; - while start.elapsed() < DEADLINE { - polls += 1; - if let Some(frame) = polled.try_recv_frame().unwrap() { - got = Some(frame); - break; + if let Some(output) = streamed.denoise().unwrap() { + emitted.push(output); } } - let got = got.unwrap_or_else(|| panic!("readback never landed within {DEADLINE:?} ({polls} polls)")); - eprintln!( - "try_recv_frame landed after {polls} poll(s), {:?}", - start.elapsed() - ); - - let mut blocking = test_denoiser(r, 64, 64); - for frame in &window { - blocking.push_frame(frame).unwrap(); - } - let expected = blocking - .recv_frame() - .unwrap() - .expect("blocking denoiser should have a frame ready"); - assert_eq!(got, expected); + assert_eq!(emitted.len(), (radius + 1) as usize); + assert_eq!(got, emitted[radius as usize]); } diff --git a/av-denoise-core/src/nlmeans/tests/mod.rs b/av-denoise-core/src/nlmeans/tests/mod.rs index 9fce434..2fa1205 100644 --- a/av-denoise-core/src/nlmeans/tests/mod.rs +++ b/av-denoise-core/src/nlmeans/tests/mod.rs @@ -1,26 +1,27 @@ -mod helpers; +pub(crate) mod helpers; mod alignment; mod confidence; +mod dispatch; mod edges; +mod engine; mod gpu_submit; mod hq; mod machinery; mod motion_compensation; mod noise; mod noise_curve; -mod pack_wire; -mod pending_drop; -mod pending_outlives; +mod options; +mod params; mod prefilter; mod residual_correlation; mod separable; +mod sizes; mod spatial; mod spatial_offset; mod split_sigma; mod strength_map; mod temporal; mod temporal_noise; -mod unpack_wire; mod util; mod validation; diff --git a/av-denoise-core/src/nlmeans/tests/motion_compensation.rs b/av-denoise-core/src/nlmeans/tests/motion_compensation.rs deleted file mode 100644 index f077550..0000000 --- a/av-denoise-core/src/nlmeans/tests/motion_compensation.rs +++ /dev/null @@ -1,1744 +0,0 @@ -use cubecl::prelude::*; -use cubecl::server::Handle; - -use super::helpers::*; -use crate::nlmeans::kernels::motion::{nlm_mc_block_match_coarse, nlm_mc_block_match_fine}; -use crate::nlmeans::motion::{ - CHAINED_RADIUS_THRESHOLD, - DEFAULT_BLKSIZE, - DEFAULT_OVERLAP, - DEFAULT_PYRAMID_LEVELS, - DEFAULT_SEARCH_RADIUS, - MotionCtx, - mv_field_byte_offset, - neighbour_idx_for_k, - pair_byte_offset, -}; -use crate::nlmeans::*; - -/// Build a frame with a constant background and a bright square at -/// `(square_x, square_y)`. Used to simulate translating content -/// across the temporal window. -fn frame_with_square( - w: u32, - h: u32, - background: f32, - square_x: u32, - square_y: u32, - square_size: u32, - square_val: f32, -) -> Vec { - let mut frame = vec![background; (w * h) as usize]; - for dy in 0..square_size { - for dx in 0..square_size { - let x = square_x + dx; - let y = square_y + dy; - if x < w && y < h { - frame[(y * w + x) as usize] = square_val; - } - } - } - frame -} - -/// Launches `nlm_mc_block_match_fine` directly over a single block -/// covering the whole `blksize × blksize` buffer (one cube, `blocks_x -/// = blocks_y = 1`, `use_seed = 0`), returning the winning MV and -/// confidence score. Exercises the kernel's SAD reduction and argmin -/// directly, without `run_analyse`'s pyramid/geometry plumbing. -fn run_fine_block_match_single_block( - blksize: u32, - search_radius: u32, - centre: &[f32], - neighbour: &[f32], - sad_noise_floor: f32, - thsad: f32, -) -> (i32, i32, f32) { - let client = make_client(); - let level_len = (blksize * blksize) as usize; - assert_eq!(centre.len(), level_len); - assert_eq!(neighbour.len(), level_len); - - let centre_buf = client.create_from_slice(f32::as_bytes(centre)); - let neighbour_buf = client.create_from_slice(f32::as_bytes(neighbour)); - let mv_field = client.empty(2 * size_of::()); - let confidence = client.empty(size_of::()); - - let grid = CubeCount::new_2d(1, 1); - let dim = CubeDim::new_2d(8, 8); - - unsafe { - nlm_mc_block_match_fine::launch_unchecked::( - &client, - grid, - dim, - ArrayArg::from_raw_parts(centre_buf, level_len), - ArrayArg::from_raw_parts(neighbour_buf, level_len), - ArrayArg::from_raw_parts(mv_field.clone(), 2), - ArrayArg::from_raw_parts(confidence.clone(), 1), - true, - sad_noise_floor, - thsad, - blksize, - blksize, - blksize, - blksize, - search_radius, - 0u32, - 1, - ); - } - - let mv_bytes = client.read_one(mv_field).expect("mv readback failed"); - let mv = i32::from_bytes(&mv_bytes); - let conf_bytes = client.read_one(confidence).expect("confidence readback failed"); - let confidence = f32::from_bytes(&conf_bytes)[0]; - (mv[0], mv[1], confidence) -} - -/// Recovers the fine kernel's raw `best_sad` from its confidence -/// output by inverting the confidence formula (`confidence = (thsad² - -/// S²) / (thsad² + S²)` inverts to `S = thsad · sqrt((1 - confidence) / -/// (1 + confidence))`). Requires `sad_noise_floor = 0.0` (so `excess == -/// best_sad`) and a `thsad` comfortably larger than the expected SAD, -/// so the confidence value stays well clear of both the `≈ 1` corner -/// (catastrophic cancellation computing `1 - confidence`) and the `= 0` -/// clamp corner (all precision lost). -fn recover_sad_from_confidence(confidence: f32, thsad: f32) -> f32 { - thsad * ((1.0 - confidence) / (1.0 + confidence)).sqrt() -} - -/// Exact-SAD test. A uniform `|Δ| = d` mismatch between centre and -/// neighbour across the whole block, at zero MV (`search_radius = 0`, -/// a single candidate), must give `best_sad = blksize² · d` exactly -/// (within a tight f32 tolerance). This is the ground truth the -/// block-match kernel's SAD reduction is supposed to compute. -/// -/// Every pixel's contribution must land in the candidate's SAD sum -/// exactly once. A racy shared-memory accumulation (multiple threads -/// `+=`-ing one candidate slot without atomics) loses most -/// contributions and undercounts the SAD by orders of magnitude. -#[test] -fn block_match_fine_exact_sad_uniform_mismatch() { - let blksize = 16u32; - let d = 0.1f32; - let centre = vec![0.25f32; (blksize * blksize) as usize]; - let neighbour = vec![0.25f32 + d; (blksize * blksize) as usize]; - - let expected_sad = (blksize * blksize) as f32 * d; - // Comfortably above the expected SAD so the confidence readout - // avoids both precision corners (see `recover_sad_from_confidence`). - let thsad = 3.0 * expected_sad; - - let (_, _, confidence) = run_fine_block_match_single_block(blksize, 0, ¢re, &neighbour, 0.0, thsad); - let measured_sad = recover_sad_from_confidence(confidence, thsad); - - assert!( - (measured_sad - expected_sad).abs() < expected_sad * 0.01, - "uniform |Δ|={d} over a {blksize}x{blksize} block should give best_sad \ - = {expected_sad} (blksize²·d), measured {measured_sad} (confidence={confidence})", - ); -} - -/// Clamps `val - delta` into `[0, limit)`. Used to build a "clean -/// shift" neighbour frame below whose out-of-range edge pixels use the -/// same clamp-to-edge convention the kernel itself applies (`clamp_i32` -/// in `block_match.rs`), even though the test's block/search geometry -/// never actually reaches those edges. -fn shift_clamped(val: i32, delta: i32, limit: i32) -> i32 { - (val - delta).clamp(0, limit - 1) -} - -/// Argmin correctness. The neighbour frame is a clean `(+2, +1)` pixel -/// shift of the centre frame's content (`neighbour(x, y) = -/// centre(x - 2, y - 1)`), so the true best match sits exactly at -/// `mv = (2, 1)`. Content is a deterministic pseudo-random pattern -/// (`noisy_copy`, reused here purely as "content-rich, no ties" filler) -/// rather than a flat value, so every other candidate in the search -/// window gives a strictly larger SAD and the argmin is unambiguous. -/// -/// Runs the fine kernel over a `3×3` block grid at the library's -/// default blksize/search radius so the winning block (the centre one, -/// away from the frame edges) sees the same clamped addressing the -/// production dispatch path uses, but reads back only that one block's -/// MV. The argmin is only meaningful when each candidate's SAD -/// reflects its true cost. Corrupted per-candidate sums make the -/// winner quasi-arbitrary. -/// -/// Confidence is turned off and given a small placeholder buffer, since -/// this test only checks the winning motion vector. That also covers the -/// path where the confidence write is dropped at compile time. -#[test] -fn block_match_fine_argmin_finds_clean_shift() { - let w = 64u32; - let h = 64u32; - let blksize = DEFAULT_BLKSIZE; - let step = blksize; - let search_radius = DEFAULT_SEARCH_RADIUS; - let blocks_x = 3u32; - let blocks_y = 3u32; - - let centre = noisy_copy(w, 0.5, 0.2, 123); - let mut neighbour = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let sx = shift_clamped(x as i32, 2, w as i32) as u32; - let sy = shift_clamped(y as i32, 1, h as i32) as u32; - neighbour[(y * w + x) as usize] = centre[(sy * w + sx) as usize]; - } - } - - let client = make_client(); - let level_len = (w * h) as usize; - let centre_buf = client.create_from_slice(f32::as_bytes(¢re)); - let neighbour_buf = client.create_from_slice(f32::as_bytes(&neighbour)); - let mv_len = (blocks_x * blocks_y * 2) as usize; - let mv_field = client.empty(mv_len * size_of::()); - // Confidence is turned off below, so this placeholder is never - // indexed whatever its size. - let confidence = client.empty(size_of::()); - - let grid = CubeCount::new_2d(blocks_x, blocks_y); - let dim = CubeDim::new_2d(8, 8); - - unsafe { - nlm_mc_block_match_fine::launch_unchecked::( - &client, - grid, - dim, - ArrayArg::from_raw_parts(centre_buf, level_len), - ArrayArg::from_raw_parts(neighbour_buf, level_len), - ArrayArg::from_raw_parts(mv_field.clone(), mv_len), - ArrayArg::from_raw_parts(confidence, 1), - false, - 0.0, - 1.0, - w, - h, - blksize, - step, - search_radius, - 0u32, - blocks_x, - ); - } - - let bytes = client.read_one(mv_field).expect("mv readback failed"); - let mv = i32::from_bytes(&bytes); - // Middle block (bx=1, by=1). Its content and search window sit - // `blksize` pixels away from every frame edge, well clear of the - // `search_radius + shift` margin, so no clamped addressing is hit. - let (mid_bx, mid_by) = (1u32, 1u32); - let idx = ((mid_by * blocks_x + mid_bx) * 2) as usize; - assert_eq!( - (mv[idx], mv[idx + 1]), - (2, 1), - "a clean (+2, +1) shift of the centre content should give exactly \ - MV=(2, 1) at default blksize={blksize}/search_radius={search_radius}, got ({}, {})", - mv[idx], - mv[idx + 1], - ); -} - -/// Launches `nlm_mc_block_match_coarse` directly over a single coarse -/// block covering the whole `blksize × blksize` buffer (one cube, -/// `fine_blocks_x = fine_blocks_y = 1`, `fine_step = blksize`, so the -/// coarse block seeds exactly that one fine block), returning the -/// coarse MV it writes into `mv_field`. `level_scale` is fixed at `1` -/// so the returned MV is the raw coarse-level offset, unscaled. -fn run_coarse_block_match_single_block( - blksize: u32, - search_radius: u32, - centre: &[f32], - neighbour: &[f32], -) -> (i32, i32) { - let client = make_client(); - let level_len = (blksize * blksize) as usize; - assert_eq!(centre.len(), level_len); - assert_eq!(neighbour.len(), level_len); - - let centre_buf = client.create_from_slice(f32::as_bytes(centre)); - let neighbour_buf = client.create_from_slice(f32::as_bytes(neighbour)); - let mv_field = client.empty(2 * size_of::()); - - let grid = CubeCount::new_2d(1, 1); - let dim = CubeDim::new_2d(8, 8); - - unsafe { - nlm_mc_block_match_coarse::launch_unchecked::( - &client, - grid, - dim, - ArrayArg::from_raw_parts(centre_buf, level_len), - ArrayArg::from_raw_parts(neighbour_buf, level_len), - ArrayArg::from_raw_parts(mv_field.clone(), 2), - blksize, - blksize, - blksize, - blksize, - search_radius, - 1, - 1, - 1, - blksize, - ); - } - - let mv_bytes = client.read_one(mv_field).expect("mv readback failed"); - let mv = i32::from_bytes(&mv_bytes); - (mv[0], mv[1]) -} - -/// SAD tie-break regression test (fine pass). A block lying entirely -/// inside a flat region has the same value everywhere, including at -/// every edge-clamped read the search window's candidates touch, so -/// every candidate's SAD is exactly `0.0`, an exact tie. The argmin -/// must resolve that tie to the zero-motion seed (here `(0, 0)`, since -/// `use_seed = 0`), not the window's `(-search_radius, -search_radius)` -/// corner, the first candidate the raster scan reaches. -#[test] -fn block_match_fine_flat_region_tie_resolves_to_zero_motion() { - let blksize = 16u32; - let search_radius = 4u32; - let value = 0.5f32; - let centre = vec![value; (blksize * blksize) as usize]; - let neighbour = vec![value; (blksize * blksize) as usize]; - - let (mvx, mvy, confidence) = - run_fine_block_match_single_block(blksize, search_radius, ¢re, &neighbour, 0.0, 1.0); - - assert_eq!( - (mvx, mvy), - (0, 0), - "a flat region gives an exact SAD tie at every candidate, which must \ - resolve to the zero-motion seed, not the window corner \ - (-{search_radius}, -{search_radius}); got ({mvx}, {mvy})", - ); - // `best_sad == 0` here regardless of which candidate wins the tie, - // so confidence is `1.0` either way. The MV assertion above is - // what actually distinguishes the fix. Asserted anyway so a future - // change to the confidence formula that breaks the exact-zero case - // shows up here too. - assert_eq!( - confidence, 1.0, - "an exact SAD=0 match should give full confidence" - ); -} - -/// SAD tie-break regression test (coarse pass). Same premise as -/// `block_match_fine_flat_region_tie_resolves_to_zero_motion`, but for -/// `nlm_mc_block_match_coarse`. A corner-favouring tie here seeds every -/// fine block the coarse block covers from the wrong position, so -/// under `Chained` estimation the bias compounds across pyramid -/// levels. -#[test] -fn block_match_coarse_flat_region_tie_resolves_to_zero_motion() { - let blksize = 16u32; - let search_radius = 4u32; - let value = 0.5f32; - let centre = vec![value; (blksize * blksize) as usize]; - let neighbour = vec![value; (blksize * blksize) as usize]; - - let (mvx, mvy) = run_coarse_block_match_single_block(blksize, search_radius, ¢re, &neighbour); - - assert_eq!( - (mvx, mvy), - (0, 0), - "a flat region gives an exact SAD tie at every candidate, which the \ - coarse pass must resolve to the zero-motion candidate, not the window \ - corner (-{search_radius}, -{search_radius}); got ({mvx}, {mvy})", - ); -} - -#[test] -fn motion_compensation_uniform_passthrough() { - let client = make_client(); - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 1, 0.5); - - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }, - hq: None, - }; - - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - assert_eq!(result.len(), (w * h) as usize); - for (i, &v) in result.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!( - (v - 0.5).abs() < 1e-3, - "pixel {i}: expected 0.5 (uniform input passthrough), got {v}" - ); - } -} - -/// Motion compensation and the bilateral prefilter together. -/// -/// The reference ring has to be shifted along with the input, and the -/// output has to stay finite and in range. -#[test] -fn motion_compensation_with_bilateral_finite() { - let client = make_client(); - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 1, 0.5); - - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::Bilateral { - sigma_s: 1.0, - sigma_r: 0.1, - }, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }, - hq: None, - }; - - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - assert_eq!(result.len(), (w * h) as usize); - for (i, &v) in result.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((-0.01..=1.01).contains(&v), "pixel {i}: out-of-range output {v}"); - } -} - -/// A three-frame sequence where a bright square moves diagonally by two -/// pixels each frame. -/// -/// With motion compensation on, the denoise has to run end to end -/// without crashing, produce finite output that stays in range, and -/// leave the square exactly where the centre frame put it. -/// -/// That last point is what catches misaligned temporal blending, whether -/// from a neighbour shifted wrongly or from a bad centre-frame copy into -/// the compensated buffer. -#[test] -fn motion_compensation_translating_square_preserves_centre() { - let client = make_client(); - let w = 32u32; - let h = 32u32; - let bg = 0.3; - let sq_val = 0.8; - let sq_size = 4u32; - - // Translate the square diagonally by 2 px per frame so the - // temporal kernel without MC would see misaligned content at the - // same (x, y) across frames. Centre frame's square sits at (14, 14). - let f0 = frame_with_square(w, h, bg, 12, 12, sq_size, sq_val); - let f1 = frame_with_square(w, h, bg, 14, 14, sq_size, sq_val); - let f2 = frame_with_square(w, h, bg, 16, 16, sq_size, sq_val); - - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }, - hq: None, - }; - - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&f0); - d.push_frame(&f1); - d.push_frame(&f2); - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - assert_eq!(result.len(), (w * h) as usize); - for (i, &v) in result.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((-0.01..=1.01).contains(&v), "pixel {i}: out-of-range output {v}"); - } - - // The centre of the centre-frame's square (15, 15) should still - // look brightly square-like. We allow significant tolerance: - // temporal blending will pull it toward the background even with - // MC because the warping is integer-pixel and the square's edges - // may not align perfectly across all neighbours. The assertion - // is just that the centre pixel hasn't been smeared *below* the - // halfway point between bg and sq_val. - let halfway = (bg + sq_val) * 0.5; - let centre_val = result[(15 * w + 15) as usize]; - assert!( - centre_val > halfway, - "centre of moving square should remain above halfway between bg ({bg}) \ - and sq_val ({sq_val}) (= {halfway}), got {centre_val}", - ); - - // Conversely, the background a few pixels away from the centre - // frame's square must stay near `bg`. If MC mis-warped neighbour - // squares into the wrong position, the background would brighten. - let bg_val = result[(2 * w + 2) as usize]; - assert!( - (bg_val - bg).abs() < 0.05, - "background pixel (2, 2) should stay near {bg}, got {bg_val} \ - (MC may be warping neighbour squares into the background region)", - ); -} - -/// Guards the motion field's binding offsets against misalignment. -/// -/// A 1080x1080 frame at the library defaults gives 135x135 blocks, an -/// odd count of 18,225, whose unpadded per-neighbour stride is not a -/// 32-byte multiple. -/// -/// wgpu rejects a binding offset that is not a multiple of its -/// `min_storage_buffer_offset_alignment`, so the second neighbour's -/// dispatch fails outright at this exact size unless the stride is -/// padded. See `mv_field_byte_offset`. -/// -/// The 1920x1080 size every other test uses happens to land on an even -/// block count here and never reaches the bug, which is why this needs -/// its own size. -#[test] -fn motion_compensation_1080_square_odd_block_count_dispatch_succeeds() { - let client = make_client(); - let w = 1080u32; - let h = 1080u32; - let frame = make_uniform_frame(w, h, 1, 0.5); - - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::mvtools_default(), - hq: None, - }; - let mc = MotionCtx::new(params.motion_compensation, w, h, test_align()).unwrap(); - assert_eq!( - mc.blocks_x * mc.blocks_y, - 18225, - "test premise: this geometry gives an odd block count" - ); - - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - assert_eq!(result.len(), (w * h) as usize); - for (i, &v) in result.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!( - (v - 0.5).abs() < 1e-3, - "pixel {i}: expected 0.5 (uniform input passthrough), got {v}" - ); - } -} - -/// Guards the pair ring's binding offsets against misalignment. -/// -/// This uses the same odd-block-count geometry as -/// `motion_compensation_1080_square_odd_block_count_dispatch_succeeds` -/// above, but pins `Chained` estimation so the pair ring is actually -/// allocated and used. -/// -/// The defaults pin `Direct`, whose dispatch never touches the pair ring -/// at all, so the test above only ever binds into the motion field. -/// -/// Every other `Chained` test in this file uses a small frame with 256 -/// blocks, an even count, so none of them reach an odd one either. -/// -/// Pushing enough real frames past the priming window covers both the -/// zeroed duplicate slots and the real hops, all of which bind into the -/// pair ring. -#[test] -fn motion_compensation_1080_square_odd_block_count_chained_dispatch_succeeds() { - let client = make_client(); - let w = 1080u32; - let h = 1080u32; - let radius = 2u32; - let frames: Vec> = (0..8) - .map(|i| make_frame_with_noisy_region(w, h, 1, 0.5, 200 + i * 4, 200, 8, 0.8)) - .collect(); - - let params = NlmParams { - temporal_radius: radius, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: DEFAULT_BLKSIZE, - overlap: DEFAULT_OVERLAP, - search_radius: DEFAULT_SEARCH_RADIUS, - pyramid_levels: DEFAULT_PYRAMID_LEVELS, - estimation: MotionEstimation::chained_default(), - }, - hq: None, - }; - let mc = MotionCtx::new(params.motion_compensation, w, h, test_align()).unwrap(); - assert_eq!( - mc.blocks_x * mc.blocks_y, - 18225, - "test premise: this geometry gives an odd block count" - ); - - let mut d = NlmDenoiser::::new(&client, params, w, h); - assert!( - d.pair_ring_buf.is_some(), - "test premise: Chained estimation must allocate the pair ring" - ); - - let check = |frame: &[f32]| { - for (i, &v) in frame.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((-0.01..=1.01).contains(&v), "pixel {i}: out-of-range output {v}"); - } - }; - - let mut emitted = 0usize; - for frame in &frames { - d.push_frame(frame); - if let Some(result) = d.denoise().unwrap() { - check(result.as_f32().expect("f32 denoiser")); - emitted += 1; - } - } - d.flush(|frame| { - check(frame.as_f32().expect("f32 denoiser")); - emitted += 1; - }) - .unwrap(); - - assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); -} - -// --- Coarse-seeding tiling and pyramid level-0 extraction guards. -// Both defects live in the shared coarse+fine machinery every -// estimation mode funnels through (`run_analyse`), so they're -// exercised here through a real `Direct`-estimation `NlmDenoiser`, -// reading `mv_field_buf` back directly (the same pattern -// `composed_centre_mv` above already establishes for the Chained -// tests). - -/// Builds a `w x h` frame from two independently-rich `half`-wide -/// halves (`left`, `right`), each optionally shifted locally within -/// its own half by `left_shift`/`right_shift` fine-level pixels -/// (`0` for the unshifted base frame). Local, per-half clamped -/// shifting (via `shift_clamped`) keeps each half's content a -/// self-contained translation, so a block deep inside one half never -/// legitimately depends on the other half's content. -fn split_half_frame( - w: u32, - h: u32, - half: u32, - left: &[f32], - right: &[f32], - left_shift: i32, - right_shift: i32, -) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let idx = (y * w + x) as usize; - if x < half { - let lx = shift_clamped(x as i32, left_shift, half as i32) as u32; - frame[idx] = left[(y * half + lx) as usize]; - } else { - let rx = shift_clamped((x - half) as i32, right_shift, half as i32) as u32; - frame[idx] = right[(y * half + rx) as usize]; - } - } - } - frame -} - -/// Pushes `base` twice then `neighbour` once through a fresh `Direct` -/// `NlmDenoiser` at `mode`, runs one `denoise()`, and reads back the -/// forward (`k = 1`) neighbour's MV field. Mirrors the push order -/// every other test in this file relies on (`push, push, push`, single -/// `denoise()`, centre = the second push, see -/// `motion_compensation_translating_square_preserves_centre`'s own -/// comment for the derivation), so `neighbour` (the third push) is -/// what the returned MV field is matched against. -fn direct_mv_field_for_forward_neighbour( - mode: MotionCompensationMode, - w: u32, - h: u32, - base: &[f32], - neighbour: &[f32], -) -> Vec { - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: mode, - hq: None, - }; - - let client = make_client(); - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(base); - d.push_frame(base); - d.push_frame(neighbour); - d.denoise().unwrap(); - - let mc = MotionCtx::new(mode, w, h, test_align()).unwrap(); - let neighbour_idx = neighbour_idx_for_k(1, 1); - let mv_field = d - .mv_field_buf - .as_ref() - .expect("mv_field allocated when mc_ctx is Some"); - let offset = mv_field_byte_offset(&mc, neighbour_idx); - let sliced = mv_field.clone().offset_start(offset); - let bytes = d.client.read_one(sliced).expect("mv readback failed"); - i32::from_bytes(&bytes).to_vec() -} - -/// Equal-grid seeding guard. At `blksize=8, overlap=4, -/// pyramid_levels=2` the coarse and fine grids come out to the *same* -/// block count (`coarse_blocks_x == blocks_x`, verified below -/// algebraically rather than assumed), so seeding must map by -/// position, not by index doubling. A fixed-doubling map -/// (`fine_bx_origin = bx * level_scale`) lands a coarse block's -/// displacement at fine index `2 * bx` instead of `bx`, seeding -/// roughly half the fine grid from a coarse block spatially tied to a -/// *different* region. -/// -/// The content is a left half and a right half, each independently -/// detailed, and each moving by a different amount between the centre -/// frame and the one after it. -/// -/// Both shifts are multiples of the pyramid scale, so the coarse level -/// sees an exact half-size version of each. -/// -/// The blocks under test sit deep inside their own half, well clear of -/// both the boundary between the halves and the frame edges. A -/// correctly-seeded block therefore finds its own half's motion, while -/// one seeded from the wrong half lands outside the fine pass's small -/// search window and cannot recover. -/// -/// At this geometry a left-half block always gets seeded from another -/// left-half block, which carries the same motion, so a seeding defect -/// only shows up on the right half. -#[test] -fn coarse_seeding_handles_equal_grids() { - let w = 128u32; - let h = 64u32; - let half = 64u32; - - let mode = MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 3, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }; - let mc = MotionCtx::new(mode, w, h, test_align()).unwrap(); - - // Recompute `run_analyse`'s coarse-grid formula directly (mirrors - // `analyse.rs`'s own derivation) rather than trusting the brief's - // worked numbers, confirming this geometry actually lands on - // "equal grids" before relying on it. - let coarse_scale = 1u32 << (mc.pyramid_levels - 1); - let coarse_step = (mc.step / coarse_scale).max(1); - let cw = w / coarse_scale; - let coarse_blocks_x = cw.div_ceil(coarse_step).max(1); - assert_eq!( - coarse_blocks_x, mc.blocks_x, - "test premise: this geometry must give equal coarse/fine grids" - ); - - let left = noisy_copy(half, 0.5, 0.2, 201); - let right = noisy_copy(half, 0.5, 0.2, 202); - let left_shift = 4i32; - let right_shift = -4i32; - - let base = split_half_frame(w, h, half, &left, &right, 0, 0); - let shifted = split_half_frame(w, h, half, &left, &right, left_shift, right_shift); - - let data = direct_mv_field_for_forward_neighbour(mode, w, h, &base, &shifted); - - // Deep inside each half, clear of the boundary at x=64 and the - // frame edges by well over `search_radius + blksize` (= 11). - let by = 4u32; - let bx_left = 8u32; - let bx_right = 24u32; - let idx_left = ((by * mc.blocks_x + bx_left) * 2) as usize; - let idx_right = ((by * mc.blocks_x + bx_right) * 2) as usize; - - assert_eq!( - (data[idx_left], data[idx_left + 1]), - (left_shift, 0), - "left-half block should recover the left half's own motion ({left_shift}, 0), got ({}, {})", - data[idx_left], - data[idx_left + 1], - ); - assert_eq!( - (data[idx_right], data[idx_right + 1]), - (right_shift, 0), - "right-half block should recover the right half's own motion ({right_shift}, 0), got \ - ({}, {}); a wrong value here means it was seeded from the wrong (left-half) coarse block", - data[idx_right], - data[idx_right + 1], - ); -} - -/// Half-grid seeding coverage at the genuine `fine_blocks_x/y == -/// 2 * coarse_blocks_x/y` geometry. `mc.step == 1` reaches this ratio -/// through the `.max(1)` floor clamp in `run_analyse`'s `coarse_step` -/// derivation (`coarse_step = (1 / 2).max(1) = 1`), verified -/// algebraically below rather than assumed. A coarse block's region -/// really does correspond to exactly `level_scale` fine blocks here, -/// so the position-based tiling formula must reduce to plain index -/// doubling, both for uniform motion and for the left/right split -/// (different motion per half) the equal-grid test above exercises. -#[test] -fn coarse_seeding_still_correct_at_half_grid() { - let w = 48u32; - let h = 16u32; - let half = 24u32; - - let mode = MotionCompensationMode::Mvtools { - blksize: 4, - overlap: 3, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }; - let mc = MotionCtx::new(mode, w, h, test_align()).unwrap(); - - let coarse_scale = 1u32 << (mc.pyramid_levels - 1); - let coarse_step = (mc.step / coarse_scale).max(1); - let cw = w / coarse_scale; - let ch = h / coarse_scale; - let coarse_blocks_x = cw.div_ceil(coarse_step).max(1); - let coarse_blocks_y = ch.div_ceil(coarse_step).max(1); - assert_eq!(mc.step, 1, "test premise: step must floor-clamp coarse_step to 1"); - assert_eq!( - mc.blocks_x, - 2 * coarse_blocks_x, - "test premise: this geometry must give a genuine 2:1 fine:coarse ratio in x" - ); - assert_eq!( - mc.blocks_y, - 2 * coarse_blocks_y, - "test premise: this geometry must give a genuine 2:1 fine:coarse ratio in y" - ); - - let left = noisy_copy(half, 0.5, 0.2, 301); - let right = noisy_copy(half, 0.5, 0.2, 302); - // Deep inside each half, clear of the boundary at x=24 and the - // frame edges by well over `search_radius + blksize` (= 6). - let by = 6u32; - let bx_left = 10u32; - let bx_right = 32u32; - - let mv_at = |left_shift: i32, right_shift: i32| -> ((i32, i32), (i32, i32)) { - let base = split_half_frame(w, h, half, &left, &right, 0, 0); - let shifted = split_half_frame(w, h, half, &left, &right, left_shift, right_shift); - let data = direct_mv_field_for_forward_neighbour(mode, w, h, &base, &shifted); - let idx_left = ((by * mc.blocks_x + bx_left) * 2) as usize; - let idx_right = ((by * mc.blocks_x + bx_right) * 2) as usize; - ( - (data[idx_left], data[idx_left + 1]), - (data[idx_right], data[idx_right + 1]), - ) - }; - - // Uniform sub-case. Both halves share one motion vector, so any - // coarse block seeds any fine block correctly and seeding defects - // stay invisible. - let (uni_left, uni_right) = mv_at(2, 2); - assert_eq!(uni_left, (2, 0), "uniform motion: left block got {uni_left:?}"); - assert_eq!(uni_right, (2, 0), "uniform motion: right block got {uni_right:?}"); - - // Varying sub-case. Left and right halves move oppositely, so - // each fine block only recovers its own half's motion when the - // tiling formula seeds it from the coarse block covering the same - // region. - let (var_left, var_right) = mv_at(2, -2); - assert_eq!(var_left, (2, 0), "varying motion: left block got {var_left:?}"); - assert_eq!( - var_right, - (-2, 0), - "varying motion: right block got {var_right:?}" - ); -} - -/// Pyramid level-0 extraction guard. `build_pyramid_for_slot` must -/// extract level-0 luma even when `pyramid_levels == 1`, not skip the -/// whole pyramid build. This is a real `Mvtools` MC config (not -/// `MotionCtx::confidence_only`, a separate path that doesn't share -/// this build step), so the fine pass runs unseeded straight off -/// level 0. If level 0 is never written, the recovered MV has no -/// reason to match the known shift below. -#[test] -fn pyramid_level0_extracted_at_one_level() { - let w = 64u32; - let h = 64u32; - let dx = 2i32; - let dy = 1i32; - - let mode = MotionCompensationMode::Mvtools { - blksize: DEFAULT_BLKSIZE, - overlap: DEFAULT_OVERLAP, - search_radius: DEFAULT_SEARCH_RADIUS, - pyramid_levels: 1, - estimation: MotionEstimation::Direct, - }; - let mc = MotionCtx::new(mode, w, h, test_align()).unwrap(); - - let world = noisy_copy(w, 0.5, 0.2, 77); - let shifted = frame_shifted_by(&world, w, h, dx, dy); - - let data = direct_mv_field_for_forward_neighbour(mode, w, h, &world, &shifted); - - // Interior block, well clear of the frame edges. - let bx = mc.blocks_x / 2; - let by = mc.blocks_y / 2; - let idx = ((by * mc.blocks_x + bx) * 2) as usize; - - assert_eq!( - (data[idx], data[idx + 1]), - (dx, dy), - "a clean ({dx}, {dy}) shift with pyramid_levels=1 should give exactly that MV at an \ - interior block once level-0 luma is actually extracted, got ({}, {})", - data[idx], - data[idx + 1], - ); -} - -/// Guards the seeding of the ragged blocks at a frame's trailing edge. -/// -/// The coarse and fine block counts each round up over a different width -/// and a different step, so they can round differently even where the -/// ratio between the two grids is otherwise consistent. -/// -/// At the library defaults, certain frame sizes leave the coarse grid -/// exactly one fine block short of the frame's true edge. The last -/// coarse block on each axis therefore has to extend its reach to -/// absorb the remainder. -/// -/// A trailing column or row no coarse block writes quietly keeps -/// whatever the motion field already holds, which is zero from a fresh -/// buffer here but a stale frame's motion in production. -/// -/// The premises are verified below rather than assumed. -/// -/// # Why the motion is uniform and negative -/// -/// The motion here is one uniform diagonal shift, not the two-part -/// split `coarse_seeding_handles_equal_grids` uses above. -/// -/// This defect leaves a block never seeded at all rather than seeded -/// from the wrong place, and that other test already covers the wrong -/// place. So what separates seeded from unseeded here is a shift large -/// enough to escape the fine-only search window while still small -/// enough to be reachable through a correct seed. -/// -/// The shift is negative on both axes on purpose. A block on the -/// trailing edge has only one genuinely valid column and row to read, -/// because the rest of its tile clamps to that same edge pixel. -/// -/// The search at that position can only tell apart offsets that pull -/// content in from the interior. Offsets reaching further past the edge -/// clamp to the same neighbour pixel for every candidate and contribute -/// nothing that distinguishes them. -/// -/// A positive shift at a trailing edge would be unrecoverable by any -/// search, seeded or not, and would not isolate this bug. -#[test] -fn coarse_seeding_covers_ragged_last_block() { - let shift = -6i32; - let mode = MotionCompensationMode::Mvtools { - blksize: DEFAULT_BLKSIZE, - overlap: DEFAULT_OVERLAP, - search_radius: 4, - pyramid_levels: 2, - estimation: MotionEstimation::Direct, - }; - - // `w_gap` (57) is congruent to 1 modulo `step` (8), giving a - // coarse grid exactly one block short of the fine grid on that - // axis (matching the exact arithmetic the review's own worked - // example used, `width = 601`). `h_nice` (64) is an exact multiple - // of `step`, giving an ordinary equal coarse/fine grid on the - // other axis (no gap there). Each case below uses one frame ragged - // on a single axis, rather than one frame ragged on both, because - // two dimensions both congruent to 1 modulo an even step are both - // odd, and an odd-by-odd pixel count can never be a multiple of - // the ring buffers' own 32-byte frame-stride alignment - // requirement. - let w_gap = 57u32; - let h_nice = 64u32; - - // Builds the MV field for a `w x h` frame under a uniform diagonal - // `(shift, shift)` translation and returns it alongside the - // `MotionCtx` used to index it. - let build = |w: u32, h: u32| -> (MotionCtx, Vec) { - let mc = MotionCtx::new(mode, w, h, test_align()).unwrap(); - let world = make_noisy_gaussian_frame(w, h, 1, 0.5, &[0.2]); - let shifted = frame_shifted_by(&world, w, h, shift, shift); - let data = direct_mv_field_for_forward_neighbour(mode, w, h, &world, &shifted); - (mc, data) - }; - let at = |mc: &MotionCtx, data: &[i32], bx: u32, by: u32| -> (i32, i32) { - let idx = ((by * mc.blocks_x + bx) * 2) as usize; - (data[idx], data[idx + 1]) - }; - // Test premise, verified rather than assumed. Mirrors - // `run_analyse`'s own coarse-grid formula (see the equal-grid test - // above for the same derivation style) to confirm this geometry - // actually lands on the ragged, one-short-of-the-fine-grid case on - // exactly the named axis, and an ordinary equal grid on the other. - let assert_ragged_on = |mc: &MotionCtx, w: u32, h: u32, ragged_axis_is_x: bool| { - let coarse_scale = 1u32 << (mc.pyramid_levels - 1); - let coarse_step = (mc.step / coarse_scale).max(1); - let coarse_blocks_x = (w / coarse_scale).div_ceil(coarse_step).max(1); - let coarse_blocks_y = (h / coarse_scale).div_ceil(coarse_step).max(1); - if ragged_axis_is_x { - assert_eq!(w % mc.step, 1, "test premise: width must be step*k + 1"); - assert_eq!( - coarse_blocks_x, - mc.blocks_x - 1, - "test premise: ragged coarse grid, one block short in x" - ); - assert_eq!( - coarse_blocks_y, mc.blocks_y, - "test premise: y axis is an ordinary equal grid here" - ); - } else { - assert_eq!(h % mc.step, 1, "test premise: height must be step*k + 1"); - assert_eq!( - coarse_blocks_y, - mc.blocks_y - 1, - "test premise: ragged coarse grid, one block short in y" - ); - assert_eq!( - coarse_blocks_x, mc.blocks_x, - "test premise: x axis is an ordinary equal grid here" - ); - } - }; - - // X-axis case, last column at a non-edge row. - let (mc_x, data_x) = build(w_gap, h_nice); - assert_ragged_on(&mc_x, w_gap, h_nice, true); - let mid_bx_x = mc_x.blocks_x / 2; - let mid_by_x = mc_x.blocks_y / 2; - assert_eq!( - at(&mc_x, &data_x, mid_bx_x, mid_by_x), - (shift, shift), - "interior control block (x-axis case) should recover ({shift}, {shift})" - ); - assert_eq!( - at(&mc_x, &data_x, mc_x.blocks_x - 1, mid_by_x), - (shift, shift), - "last-column block (x-axis coverage gap) should recover ({shift}, {shift})" - ); - - // Y-axis case, last row at a non-edge column (the same frame, - // transposed). - let (mc_y, data_y) = build(h_nice, w_gap); - assert_ragged_on(&mc_y, h_nice, w_gap, false); - let mid_bx_y = mc_y.blocks_x / 2; - let mid_by_y = mc_y.blocks_y / 2; - assert_eq!( - at(&mc_y, &data_y, mid_bx_y, mid_by_y), - (shift, shift), - "interior control block (y-axis case) should recover ({shift}, {shift})" - ); - assert_eq!( - at(&mc_y, &data_y, mid_bx_y, mc_y.blocks_y - 1), - (shift, shift), - "last-row block (y-axis coverage gap) should recover ({shift}, {shift})" - ); -} - -// --- Chained motion estimation, pair ring + composition kernel --- -// -// These tests exercise `NlmDenoiser::run_chain_compose` and the -// push-time pair analyse directly. Nothing in the submit path calls -// either yet, so every test drives them explicitly. - -const CHAIN_TEST_RADIUS: u32 = 2; -const CHAIN_TEST_SIZE: u32 = 64; - -fn chained_params(refine_radius: u32) -> NlmParams { - NlmParams { - temporal_radius: CHAIN_TEST_RADIUS, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - // Matches the MC geometry the other tests in this file already - // exercise through the real push pipeline (blksize=8, overlap=4, - // pyramid_levels=2). - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Chained { refine_radius }, - }, - hq: None, - } -} - -/// Builds a `w × h` frame that reads `world` shifted by `(dx, dy)` -/// pixels, clamped to `world`'s own edges, i.e. `frame(x, y) = world(x -/// - dx, y - dy)`. A sequence built from the same `world` with `dx = n -/// * v` for increasing `n` gives adjacent frames that differ by -/// exactly `(v, v)` everywhere except right at the frame edges. -fn frame_shifted_by(world: &[f32], w: u32, h: u32, dx: i32, dy: i32) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let sx = shift_clamped(x as i32, dx, w as i32) as u32; - let sy = shift_clamped(y as i32, dy, h as i32) as u32; - frame[(y * w + x) as usize] = world[(sy * w + sx) as usize]; - } - } - frame -} - -/// Same idea as `frame_shifted_by`, but wraps at the frame edges -/// (`rem_euclid`) instead of clamping. `world`'s content is spatially -/// unstructured (see `noisy_copy`), so a wrapped shift keeps every -/// pixel position fully translation-invariant, with no degenerate -/// clamped border region anywhere in the frame regardless of how large -/// `dx`/`dy` grow. Used by the k4 alignment test below, which pushes -/// many frames with a steadily growing absolute shift and needs the -/// whole frame to stay a clean, uniform translation of `world`. -fn frame_shifted_wrapped(world: &[f32], w: u32, h: u32, dx: i32, dy: i32) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let sx = (x as i32 - dx).rem_euclid(w as i32) as u32; - let sy = (y as i32 - dy).rem_euclid(h as i32) as u32; - frame[(y * w + x) as usize] = world[(sy * w + sx) as usize]; - } - } - frame -} - -/// Pushes a constant-velocity sequence (`v` pixels/frame on both axes, -/// built from one rich pattern) through a `Chained` denoiser. Pushes -/// `1 + 3 * radius + 2` real frames past the stream's first one. `1 + -/// 3 * radius` is enough for every one of the `2 * radius` pair-ring -/// slots to be overwritten by real analyse past the initial priming -/// duplicates (see `NlmDenoiser::pair_slot`'s doc comment for the -/// `2 * radius`-push lifetime this relies on), plus 2 frames of margin. -/// The chosen frame size and a centre-block readout keep every match -/// clear of both the frame edges and the accumulated drift. -fn push_constant_velocity(client: &ComputeClient, radius: u32, v: i32) -> NlmDenoiser { - let w = CHAIN_TEST_SIZE; - let h = CHAIN_TEST_SIZE; - let world = noisy_copy(w, 0.5, 0.2, 99); - - let mut d = NlmDenoiser::::new(client, chained_params(2), w, h); - - let real_pushes = 1 + 3 * radius as i32 + 2; - for n in 0..real_pushes { - let frame = frame_shifted_by(&world, w, h, n * v, n * v); - d.push_frame(&frame); - } - d -} - -/// Runs `run_chain_compose` for neighbour offset `k` and reads back the -/// composed MV at the block nearest the frame's centre. -fn composed_centre_mv(d: &NlmDenoiser, center_t: u32, k: i32, neighbour_idx: u32) -> (i32, i32) { - d.run_chain_compose(center_t, k, neighbour_idx) - .expect("chain compose dispatch failed"); - - let mc = MotionCtx::new(d.params.motion_compensation, d.width, d.height, d.align).unwrap(); - let mv_field = d - .mv_field_buf - .as_ref() - .expect("mv_field allocated when mc_ctx is Some"); - let offset = mv_field_byte_offset(&mc, neighbour_idx); - let sliced = mv_field.clone().offset_start(offset); - let bytes = d.client.read_one(sliced).expect("mv readback failed"); - let data = i32::from_bytes(&bytes); - - let bx = mc.blocks_x / 2; - let by = mc.blocks_y / 2; - let idx = ((by * mc.blocks_x + bx) * 2) as usize; - (data[idx], data[idx + 1]) -} - -/// Asserts that the `radius` pair-ring writes starting at -/// `ring_head_before` (the `ring_head` value the first of those writes -/// saw, pre-advance) are all zero in both directions. Used to check -/// duplicated slots (priming or flush), whose pair is zero motion by -/// definition. -fn assert_pair_ring_zero_from(d: &NlmDenoiser, ring_head_before: i32, radius: u32) { - let mc = MotionCtx::new(d.params.motion_compensation, d.width, d.height, d.align).unwrap(); - let pair_ring = d - .pair_ring_buf - .as_ref() - .expect("pair_ring allocated when Chained is active"); - let pair_ring_slots = 2 * radius as i32; - let dir_len = mc.pair_direction_len() as usize; - - for i in 0..radius as i32 { - let slot = (ring_head_before + i).rem_euclid(pair_ring_slots) as u32; - for direction in 0..2u32 { - let offset = pair_byte_offset(&mc, slot, direction); - let sliced = pair_ring.clone().offset_start(offset); - let bytes = d.client.read_one(sliced).expect("pair ring readback failed"); - let data = i32::from_bytes(&bytes); - assert!( - data[..dir_len].iter().all(|&v| v == 0), - "duplicate pair slot {slot} direction {direction} should be zero-filled, got {:?}", - &data[..dir_len], - ); - } - } -} - -#[test] -fn chain_compose_zero_motion_gives_zero_mv() { - let client = make_client(); - let radius = CHAIN_TEST_RADIUS; - let d = push_constant_velocity(&client, radius, 0); - - for k in 1..=radius as i32 { - let forward_idx = neighbour_idx_for_k(radius, k); - assert_eq!( - composed_centre_mv(&d, radius, k, forward_idx), - (0, 0), - "forward k={k} should compose to zero motion on a static sequence" - ); - let backward_idx = neighbour_idx_for_k(radius, -k); - assert_eq!( - composed_centre_mv(&d, radius, -k, backward_idx), - (0, 0), - "backward k={k} should compose to zero motion on a static sequence" - ); - } -} - -/// Constant-velocity sequence. The forward-composed MV to neighbour k -/// must equal exactly `k * v` for every k up to the temporal radius. -/// -/// Uses `v = 2` rather than `1`, since the block-match geometry here -/// runs a 2-level pyramid, and a `v = 1` fine-level shift is only half -/// a pixel at the coarse level, ambiguous enough that the coarse pass -/// can lock onto the wrong candidate for either pair direction -/// (confirmed by hand-checking the pair ring directly during -/// debugging). `v = 2` gives a clean one-pixel shift at both levels. -#[test] -fn chain_compose_constant_velocity_matches_k_times_v() { - let client = make_client(); - let radius = CHAIN_TEST_RADIUS; - let v = 2; - let d = push_constant_velocity(&client, radius, v); - - for k in 1..=radius as i32 { - let forward_idx = neighbour_idx_for_k(radius, k); - assert_eq!( - composed_centre_mv(&d, radius, k, forward_idx), - (k * v, k * v), - "forward k={k} should compose to exactly k*v = ({}, {})", - k * v, - k * v - ); - } -} - -/// From centre 0, a chained walk out to `k = 2 * radius` reaches every -/// pair the constant-velocity fixture ever wrote, the far end of the -/// ring rather than the temporal radius. The composed MV must still -/// equal `k * v`, this time landing in the last neighbour slot, -/// `2 * radius - 1`. -#[test] -fn chain_compose_reaches_twice_the_radius_from_the_ring_start() { - let client = make_client(); - let radius = CHAIN_TEST_RADIUS; - let v = 2; - let d = push_constant_velocity(&client, radius, v); - let far = 2 * radius as i32; - - let composed = composed_centre_mv(&d, 0, far, 2 * radius - 1); - - assert_eq!(composed, (far * v, far * v)); -} - -/// Same constant-velocity sequence walked backward. The composed MV to -/// neighbour -k must equal `-k * v`, the mirror image of the forward -/// case, since the backward pass reads the newer→older field at every -/// hop instead of older→newer. Uses `v = 2` for the same reason as -/// `chain_compose_constant_velocity_matches_k_times_v`. -#[test] -fn chain_compose_backward_direction_matches_negative_k_times_v() { - let client = make_client(); - let radius = CHAIN_TEST_RADIUS; - let v = 2; - let d = push_constant_velocity(&client, radius, v); - - for k in 1..=radius as i32 { - let backward_idx = neighbour_idx_for_k(radius, -k); - assert_eq!( - composed_centre_mv(&d, radius, -k, backward_idx), - (-k * v, -k * v), - "backward k={k} should compose to exactly -k*v = ({}, {})", - -k * v, - -k * v - ); - } -} - -/// Duplicated ring slots (stream priming and end-of-stream flush) get -/// a zero-filled pair field rather than a real analyse. Pushes real, -/// nonzero-motion frames in between so the check isn't trivially true -/// from a freshly-allocated, still-zeroed buffer. -#[test] -fn chain_compose_duplicated_slot_pairs_are_zero_filled() { - let client = make_client(); - let radius = CHAIN_TEST_RADIUS; - let w = CHAIN_TEST_SIZE; - let h = CHAIN_TEST_SIZE; - let world = noisy_copy(w, 0.5, 0.2, 7); - - let mut d = NlmDenoiser::::new(&client, chained_params(2), w, h); - - // During priming the very first push has no older partner (`ring_head == - // 0`), so the analyse skip leaves the window's very first gap to - // be filled entirely by the `radius` auto-primed duplicates that - // follow. Each of those duplicates' own pair write happens with - // `ring_head` starting at 1 (right after the real push's - // `advance_ring`). - d.push_frame(&world); - assert_pair_ring_zero_from(&d, 1, radius); - - // Overwrite every pair slot with real, nonzero motion before - // checking the flush path, so its zero-fill isn't indistinguishable - // from an untouched buffer. - for n in 1..=(3 * radius) { - let frame = frame_shifted_by(&world, w, h, n as i32, n as i32); - d.push_frame(&frame); - } - - let ring_head_before_flush = d.ring_head as i32; - d.flush(|_| {}).expect("flush failed"); - assert_pair_ring_zero_from(&d, ring_head_before_flush, radius); -} - -// --- Chained motion estimation, submit-path wiring (dispatch.rs) --- -// -// The tests above drive `run_chain_compose` directly. These exercise the -// real per-submit dispatch branch (`run_motion_compensation` in -// `dispatch.rs`), which composes, refines, and warps automatically on -// every `push_frame` + `denoise`/`flush` call once `estimation` is -// `Chained`. - -/// Builds an HQ config with `Chained` motion estimation at the given -/// temporal radius and refinement radius. Auto noise estimation and -/// temporal confidence stay on, mirroring `hq_temporal_mc_confidence_smoke` -/// in `tests/hq.rs` but with the chained estimator instead of direct. -fn chained_hq_params(radius: u32, refine_radius: u32) -> NlmParams { - NlmParams { - temporal_radius: radius, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Chained { refine_radius }, - }, - hq: Some(HqParams { - auto_strength: true, - noise_floor: true, - sigma_override: None, - temporal_confidence: true, - thsad_scale: 1.0, - sigma_scale: 1.0, - windowed_noise_estimation: false, - }), - } -} - -/// End-to-end smoke test. A `Chained`-estimation denoiser must produce -/// finite, `[0, 1]` output for every pushed frame, whether from pushes -/// or the trailing flush, at the given temporal radius. Exercises the -/// full `push_frame` → `denoise` → `flush` pipeline, so the compose and -/// seeded-refine dispatch branch in `run_motion_compensation` actually -/// runs (not just `run_chain_compose` in isolation). -fn chained_end_to_end_finite(radius: u32) { - let client = make_client(); - let w = 32u32; - let h = 32u32; - - let mut denoiser = NlmDenoiser::::new(&client, chained_hq_params(radius, 2), w, h); - - let frames: Vec> = (0..8) - .map(|i| make_frame_with_noisy_region(w, h, 1, 0.5, 6 + i, 8, 2, 0.8)) - .collect(); - - let mut emitted = 0usize; - let check = |frame: &[f32]| { - for (i, &v) in frame.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((0.0..=1.0).contains(&v), "pixel {i}: out-of-range output {v}"); - } - }; - - for frame in &frames { - denoiser.push_frame(frame); - if let Some(result) = denoiser.denoise().unwrap() { - check(result.as_f32().expect("f32 denoiser")); - emitted += 1; - } - } - - denoiser - .flush(|frame| { - check(frame.as_f32().expect("f32 denoiser")); - emitted += 1; - }) - .unwrap(); - - assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); -} - -#[test] -fn chained_end_to_end_finite_r2() { - chained_end_to_end_finite(2); -} - -#[test] -fn chained_end_to_end_finite_r4() { - chained_end_to_end_finite(4); -} - -/// `MotionCompensationMode::mvtools_default()` sets `estimation: Direct` -/// (see its own doc comment). Building the same configuration two ways -/// the codebase supports it, the convenience constructor and an -/// explicit `Mvtools` struct literal, must give bit-identical output for -/// identical input, since dispatch behaviour depends only on the -/// resulting value. Guards against the new chained dispatch branch in -/// `run_motion_compensation` accidentally changing what the Direct -/// branch does. -#[test] -fn direct_estimation_default_and_explicit_construction_match_bit_for_bit() { - let client = make_client(); - let w = 32u32; - let h = 32u32; - let frame = make_frame_with_noisy_region(w, h, 1, 0.5, 16, 16, 4, 0.8); - - let run = |mc: MotionCompensationMode| { - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: mc, - hq: None, - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.push_frame(&frame); - d.push_frame(&frame); - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() - }; - - let via_default = run(MotionCompensationMode::mvtools_default()); - let via_explicit = run(MotionCompensationMode::Mvtools { - blksize: DEFAULT_BLKSIZE, - overlap: DEFAULT_OVERLAP, - search_radius: DEFAULT_SEARCH_RADIUS, - pyramid_levels: DEFAULT_PYRAMID_LEVELS, - estimation: MotionEstimation::Direct, - }); - - assert_eq!( - via_default, via_explicit, - "Direct estimation must give the same output regardless of which \ - constructor produced the MotionCompensationMode value" - ); -} - -// --- Auto estimation, resolution allocates the pair ring exactly like -// an explicit Chained request would --- - -fn auto_params(temporal_radius: u32) -> NlmParams { - NlmParams { - temporal_radius, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 2, - pyramid_levels: 2, - estimation: MotionEstimation::Auto, - }, - hq: None, - } -} - -#[test] -fn auto_estimation_at_high_radius_allocates_pair_ring() { - let client = make_client(); - let params = auto_params(CHAINED_RADIUS_THRESHOLD); - - let d = NlmDenoiser::::new(&client, params, 32, 32); - - assert!( - d.pair_ring_buf.is_some(), - "Auto at radius {CHAINED_RADIUS_THRESHOLD} (>= CHAINED_RADIUS_THRESHOLD) should \ - resolve to Chained and allocate the pair ring" - ); -} - -#[test] -fn auto_estimation_at_low_radius_does_not_allocate_pair_ring() { - let client = make_client(); - let params = auto_params(CHAINED_RADIUS_THRESHOLD - 1); - - let d = NlmDenoiser::::new(&client, params, 32, 32); - - assert!( - d.pair_ring_buf.is_none(), - "Auto at radius {} (< CHAINED_RADIUS_THRESHOLD) should resolve to Direct \ - and skip the pair ring", - CHAINED_RADIUS_THRESHOLD - 1 - ); -} - -// --- The key regression test, chained beats direct once k*v exceeds -// direct's search window --- - -/// Temporal radius for the k4 alignment tests below. Chosen so k=4 is -/// reachable (needs `radius >= 4`). -const K4_RADIUS: u32 = 4; -/// Frame side length. Comfortably larger than any candidate offset the -/// coarse+fine search (or its seeded refine) can reach, so a wrapped -/// shift is never ambiguous with its aliased counterpart on the other -/// side of the torus. -const K4_SIZE: u32 = 128; -/// Per-frame diagonal shift in pixels. At this geometry (`mc.search_radius -/// = 4`, `pyramid_levels = 2`), the coarse pass doubles reach to 8px -/// (half-res search radius 4, scaled by the /2 level) and the fine pass -/// adds another 4px around that seed, so direct's reachable window is -/// 12px total. `k = 1` motion (this value) sits comfortably inside that -/// window. `k = 4` motion (`4 * K4_V = 16` px) exceeds it. -const K4_V: i32 = 4; - -fn k4_params(estimation: MotionEstimation) -> NlmParams { - NlmParams { - temporal_radius: K4_RADIUS, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::Mvtools { - blksize: 8, - overlap: 4, - search_radius: 4, - pyramid_levels: 2, - estimation, - }, - hq: None, - } -} - -/// Reads back frame `slot` from a ring-buffer handle laid out -/// `[total_frames][h][w][stored_ch]` f32, matching `input_buf` / -/// `compensated_input_buf`'s shared layout. -fn read_frame_slot( - client: &ComputeClient, - buf: &Handle, - slot: u32, - w: u32, - h: u32, - stored_ch: u32, -) -> Vec { - let frame_size = (w * h * stored_ch) as usize; - let byte_offset = (slot as u64) * (frame_size as u64) * (size_of::() as u64); - let sliced = buf.clone().offset_start(byte_offset); - let bytes = client.read_one(sliced).expect("frame readback failed"); - f32::from_bytes(&bytes)[..frame_size].to_vec() -} - -/// Pushes a diagonal constant-velocity sequence (`K4_V` px/frame, -/// wrapped at the edges via `frame_shifted_wrapped` so the whole frame -/// stays a clean, uniform translation of one rich pattern with no -/// degenerate border region) through a `k4_params`-geometry denoiser, -/// then returns the mean absolute residual between the centre frame and -/// the forward `k`-neighbour's *compensated* (warped) copy, averaged -/// over the whole frame. -/// -/// Pushes `2 * K4_RADIUS + 4` real frames before reading back, clearing -/// the stream's leading-edge priming duplicates (which need only `> -/// 2 * radius + 1` real pushes, see `NlmDenoiser::pair_slot`'s doc -/// comment) with a small margin to spare. -fn k4_compensated_residual(estimation: MotionEstimation, k: i32) -> f32 { - let client = make_client(); - let w = K4_SIZE; - let h = K4_SIZE; - let world = noisy_copy(w, 0.5, 0.2, 55); - - let mut d = NlmDenoiser::::new(&client, k4_params(estimation), w, h); - - let real_pushes = 2 * K4_RADIUS as i32 + 4; - for n in 0..real_pushes { - let frame = frame_shifted_wrapped(&world, w, h, n * K4_V, n * K4_V); - d.push_frame(&frame); - } - d.denoise().unwrap(); - - let radius = d.params.temporal_radius; - let stored_ch = d.params.channels.storage_count(); - let centre_slot = d.phys_frame(radius as i32); - let neighbour_slot = d.phys_frame(radius as i32 + k); - let compensated = d - .compensated_input_buf - .as_ref() - .expect("compensated buf allocated when MC is active"); - - let centre_frame = read_frame_slot(&d.client, &d.input_buf, centre_slot, w, h, stored_ch); - let warped = read_frame_slot(&d.client, compensated, neighbour_slot, w, h, stored_ch); - - let mut sum = 0.0f32; - let mut count = 0u32; - for y in 0..h { - for x in 0..w { - let idx = (y * w + x) as usize; - sum += (centre_frame[idx] - warped[idx]).abs(); - count += 1; - } - } - sum / count as f32 -} - -/// THE key regression test for chained motion estimation. At `k = 4` -/// the true displacement is `4 * K4_V = 16` px, beyond direct's -/// reachable window (`≈ 12` px at this geometry), while chained -/// composes four exact per-step vectors and only needs a `refine_radius -/// = 2` correction to land on the true offset. A real, wide-margin -/// assertion, not just a smoke check, since this comparison is the -/// entire reason `Chained` estimation exists. -#[test] -fn chained_beats_direct_at_k4_beyond_direct_window() { - let direct_residual = k4_compensated_residual(MotionEstimation::Direct, 4); - let chained_residual = k4_compensated_residual(MotionEstimation::Chained { refine_radius: 2 }, 4); - - assert!( - direct_residual > 0.02, - "expected direct's k=4 match to show a real misalignment residual \ - (window reach ≈12px, true motion 16px), got {direct_residual}" - ); - assert!( - chained_residual < direct_residual * 0.5, - "chained's composed+refined k=4 alignment should beat direct's by a \ - wide margin: chained={chained_residual}, direct={direct_residual}" - ); -} - -/// Companion sanity check. At `k = 1` the true displacement (`K4_V` = 4 -/// px) sits comfortably inside direct's reach, so direct should already -/// align cleanly there. This isolates the k=4 test's effect to the -/// window-size failure mode rather than a blanket "chained is always -/// better" claim. -#[test] -fn direct_already_aligns_at_k1_inside_its_window() { - let direct_residual = k4_compensated_residual(MotionEstimation::Direct, 1); - assert!( - direct_residual < 0.02, - "direct should align cleanly at k=1 (motion {K4_V}px, well inside its ~12px reach), got {direct_residual}" - ); -} diff --git a/av-denoise-core/src/nlmeans/tests/motion_compensation/block_match.rs b/av-denoise-core/src/nlmeans/tests/motion_compensation/block_match.rs new file mode 100644 index 0000000..b3841f5 --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/motion_compensation/block_match.rs @@ -0,0 +1,262 @@ +use cubecl::prelude::*; + +use super::frame_shifted_by; +use crate::nlmeans::kernels::motion::{nlm_mc_block_match_coarse, nlm_mc_block_match_fine}; +use crate::nlmeans::motion::{DEFAULT_BLKSIZE, DEFAULT_SEARCH_RADIUS}; +use crate::nlmeans::tests::helpers::*; + +/// Runs the fine block-match kernel over one block covering the whole buffer. +/// +/// It returns the winning motion vector and its confidence, with no seed and no pyramid. +fn run_fine_block_match_single_block( + blksize: u32, + search_radius: u32, + centre: &[f32], + neighbour: &[f32], + sad_noise_floor: f32, + thsad: f32, +) -> (i32, i32, f32) { + let client = make_client(); + let level_len = (blksize * blksize) as usize; + assert_eq!(centre.len(), level_len); + assert_eq!(neighbour.len(), level_len); + + let centre_bytes = f32::as_bytes(centre); + let neighbour_bytes = f32::as_bytes(neighbour); + let centre_buf = client.create_from_slice(centre_bytes); + let neighbour_buf = client.create_from_slice(neighbour_bytes); + let mv_field = client.empty(2 * size_of::()); + let confidence = client.empty(size_of::()); + + let grid = CubeCount::new_2d(1, 1); + let dim = CubeDim::new_2d(8, 8); + + unsafe { + nlm_mc_block_match_fine::launch_unchecked::( + &client, + grid, + dim, + ArrayArg::from_raw_parts(centre_buf, level_len), + ArrayArg::from_raw_parts(neighbour_buf, level_len), + ArrayArg::from_raw_parts(mv_field.clone(), 2), + ArrayArg::from_raw_parts(confidence.clone(), 1), + true, + sad_noise_floor, + thsad, + blksize, + blksize, + blksize, + blksize, + search_radius, + 0u32, + 1, + ); + } + + let mv_bytes = client.read_one(mv_field).expect("mv readback failed"); + let mv = i32::from_bytes(&mv_bytes); + let confidence_bytes = client.read_one(confidence).expect("confidence readback failed"); + let confidence = f32::from_bytes(&confidence_bytes)[0]; + (mv[0], mv[1], confidence) +} + +/// Recovers the fine kernel's best SAD by inverting its confidence formula. +/// +/// `confidence = (thsad² - S²) / (thsad² + S²)` inverts to `S = thsad · sqrt((1 - confidence) / +/// (1 + confidence))`. It needs `sad_noise_floor = 0.0` and a `thsad` well above the expected SAD, +/// which keeps the confidence clear of cancellation near 1 and of the clamp at 0. +fn recover_sad_from_confidence(confidence: f32, thsad: f32) -> f32 { + thsad * ((1.0 - confidence) / (1.0 + confidence)).sqrt() +} + +/// A racy shared-memory SAD reduction drops contributions and undercounts by orders of magnitude. +#[test] +fn block_match_fine_exact_sad_uniform_mismatch() { + let blksize = 16u32; + let mismatch = 0.1f32; + let centre = vec![0.25f32; (blksize * blksize) as usize]; + let neighbour = vec![0.25f32 + mismatch; (blksize * blksize) as usize]; + + let expected_sad = (blksize * blksize) as f32 * mismatch; + // Well above the expected SAD so the confidence stays clear of both precision corners. + let thsad = 3.0 * expected_sad; + + let (_, _, confidence) = run_fine_block_match_single_block(blksize, 0, ¢re, &neighbour, 0.0, thsad); + let measured_sad = recover_sad_from_confidence(confidence, thsad); + + assert!( + (measured_sad - expected_sad).abs() < expected_sad * 0.01, + "uniform |Δ|={mismatch} over a {blksize}x{blksize} block should give best_sad \ + = {expected_sad} (blksize²·d), measured {measured_sad} (confidence={confidence})", + ); +} + +/// The neighbour is the centre shifted by `(+2, +1)`, so the argmin must land exactly there. +/// +/// Rich content keeps every other candidate's SAD strictly larger. A 3x3 grid at the library +/// defaults gives the middle block the same clamped addressing the production dispatch uses. +/// Confidence is turned off, which also covers the path where its write is compiled out. +#[test] +fn block_match_fine_argmin_finds_clean_shift() { + let width = 64u32; + let height = 64u32; + let blksize = DEFAULT_BLKSIZE; + let step = blksize; + let search_radius = DEFAULT_SEARCH_RADIUS; + let blocks_x = 3u32; + let blocks_y = 3u32; + + let centre = noisy_copy(width, 0.5, 0.2, 123); + let neighbour = frame_shifted_by(¢re, width, height, 2, 1); + + let client = make_client(); + let level_len = (width * height) as usize; + let centre_bytes = f32::as_bytes(¢re); + let neighbour_bytes = f32::as_bytes(&neighbour); + let centre_buf = client.create_from_slice(centre_bytes); + let neighbour_buf = client.create_from_slice(neighbour_bytes); + let mv_len = (blocks_x * blocks_y * 2) as usize; + let mv_field = client.empty(mv_len * size_of::()); + // Confidence is off below, so this placeholder is never indexed. + let confidence = client.empty(size_of::()); + + let grid = CubeCount::new_2d(blocks_x, blocks_y); + let dim = CubeDim::new_2d(8, 8); + + unsafe { + nlm_mc_block_match_fine::launch_unchecked::( + &client, + grid, + dim, + ArrayArg::from_raw_parts(centre_buf, level_len), + ArrayArg::from_raw_parts(neighbour_buf, level_len), + ArrayArg::from_raw_parts(mv_field.clone(), mv_len), + ArrayArg::from_raw_parts(confidence, 1), + false, + 0.0, + 1.0, + width, + height, + blksize, + step, + search_radius, + 0u32, + blocks_x, + ); + } + + let bytes = client.read_one(mv_field).expect("mv readback failed"); + let mv = i32::from_bytes(&bytes); + + // The middle block and its search window sit `blksize` pixels from every edge, so no clamped + // read is hit. + let middle_block_x = 1u32; + let middle_block_y = 1u32; + let mv_index = ((middle_block_y * blocks_x + middle_block_x) * 2) as usize; + assert_eq!( + (mv[mv_index], mv[mv_index + 1]), + (2, 1), + "a clean (+2, +1) shift of the centre content should give exactly \ + MV=(2, 1) at default blksize={blksize}/search_radius={search_radius}, got ({}, {})", + mv[mv_index], + mv[mv_index + 1], + ); +} + +/// Runs the coarse block-match kernel over one block covering the whole buffer. +/// +/// The block seeds exactly one fine block at a level scale of 1, so the returned vector is the raw +/// coarse offset. +fn run_coarse_block_match_single_block( + blksize: u32, + search_radius: u32, + centre: &[f32], + neighbour: &[f32], +) -> (i32, i32) { + let client = make_client(); + let level_len = (blksize * blksize) as usize; + assert_eq!(centre.len(), level_len); + assert_eq!(neighbour.len(), level_len); + + let centre_bytes = f32::as_bytes(centre); + let neighbour_bytes = f32::as_bytes(neighbour); + let centre_buf = client.create_from_slice(centre_bytes); + let neighbour_buf = client.create_from_slice(neighbour_bytes); + let mv_field = client.empty(2 * size_of::()); + + let grid = CubeCount::new_2d(1, 1); + let dim = CubeDim::new_2d(8, 8); + + unsafe { + nlm_mc_block_match_coarse::launch_unchecked::( + &client, + grid, + dim, + ArrayArg::from_raw_parts(centre_buf, level_len), + ArrayArg::from_raw_parts(neighbour_buf, level_len), + ArrayArg::from_raw_parts(mv_field.clone(), 2), + blksize, + blksize, + blksize, + blksize, + search_radius, + 1, + 1, + 1, + blksize, + ); + } + + let mv_bytes = client.read_one(mv_field).expect("mv readback failed"); + let mv = i32::from_bytes(&mv_bytes); + (mv[0], mv[1]) +} + +/// Every candidate ties at a SAD of 0, and the raster scan reaches the window corner first. +#[test] +fn block_match_fine_flat_region_tie_resolves_to_zero_motion() { + let blksize = 16u32; + let search_radius = 4u32; + let value = 0.5f32; + let centre = vec![value; (blksize * blksize) as usize]; + let neighbour = vec![value; (blksize * blksize) as usize]; + + let (mv_x, mv_y, confidence) = + run_fine_block_match_single_block(blksize, search_radius, ¢re, &neighbour, 0.0, 1.0); + + assert_eq!( + (mv_x, mv_y), + (0, 0), + "a flat region gives an exact SAD tie at every candidate, which must \ + resolve to the zero-motion seed, not the window corner \ + (-{search_radius}, -{search_radius}); got ({mv_x}, {mv_y})", + ); + + // Confidence is 1.0 whichever candidate wins the tie, so only the vector above pins the + // tie-break. This guards the exact-zero case of the confidence formula. + assert_eq!( + confidence, 1.0, + "an exact SAD=0 match should give full confidence" + ); +} + +/// A corner-biased tie here mis-seeds every fine block under it, which compounds across pyramid +/// levels under `Chained` estimation. +#[test] +fn block_match_coarse_flat_region_tie_resolves_to_zero_motion() { + let blksize = 16u32; + let search_radius = 4u32; + let value = 0.5f32; + let centre = vec![value; (blksize * blksize) as usize]; + let neighbour = vec![value; (blksize * blksize) as usize]; + + let (mv_x, mv_y) = run_coarse_block_match_single_block(blksize, search_radius, ¢re, &neighbour); + + assert_eq!( + (mv_x, mv_y), + (0, 0), + "a flat region gives an exact SAD tie at every candidate, which the \ + coarse pass must resolve to the zero-motion candidate, not the window \ + corner (-{search_radius}, -{search_radius}); got ({mv_x}, {mv_y})", + ); +} diff --git a/av-denoise-core/src/nlmeans/tests/motion_compensation/chain.rs b/av-denoise-core/src/nlmeans/tests/motion_compensation/chain.rs new file mode 100644 index 0000000..b03110e --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/motion_compensation/chain.rs @@ -0,0 +1,436 @@ +use cubecl::prelude::*; +use cubecl::server::Handle; + +use super::frame_shifted_by; +use crate::bench_api::HostIo; +use crate::nlmeans::motion::{ + CHAINED_RADIUS_THRESHOLD, + MotionCtx, + mv_field_byte_offset, + neighbour_idx_for_k, + pair_byte_offset, +}; +use crate::nlmeans::tests::helpers::*; +use crate::nlmeans::*; + +const CHAIN_TEST_RADIUS: u32 = 2; +const CHAIN_TEST_SIZE: u32 = 64; + +/// Temporal radius large enough to reach `k = 4`. +const K4_RADIUS: u32 = 4; +/// Frame side, larger than any offset the search can reach, so a wrapped shift is never confused +/// with its alias on the other side. +const K4_SIZE: u32 = 128; +/// Diagonal shift per frame in pixels. +/// +/// Direct reaches about 12 px at this geometry, 8 px from the coarse pass plus 4 px from the fine +/// pass. So `k = 1` sits inside its reach and `k = 4` (16 px) does not. +const K4_V: i32 = 4; + +fn chained_params(refine_radius: u32) -> NlmParams { + NlmParams { + temporal_radius: CHAIN_TEST_RADIUS, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Chained { refine_radius }, + }, + hq: None, + } +} + +/// Like [frame_shifted_by](crate::nlmeans::tests::motion_compensation::frame_shifted_by) but wraps +/// at the edges, so a growing shift never leaves a clamped border. +fn frame_shifted_wrapped(world: &[f32], width: u32, height: u32, shift_x: i32, shift_y: i32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let source_x = (x as i32 - shift_x).rem_euclid(width as i32) as u32; + let source_y = (y as i32 - shift_y).rem_euclid(height as i32) as u32; + frame[(y * width + x) as usize] = world[(source_y * width + source_x) as usize]; + } + } + frame +} + +/// Pushes a sequence moving `velocity` pixels per frame on both axes through a `Chained` denoiser. +/// +/// A pair-ring slot lives for `2 * radius` pushes (see +/// [pair_ring_slot_count](crate::nlmeans::motion::pair_ring_slot_count)), so `1 + 3 * radius` +/// pushes replace every priming duplicate with a real analyse. Two more frames give margin. +fn push_constant_velocity(client: &ComputeClient, radius: u32, velocity: i32) -> NlmDenoiser { + let width = CHAIN_TEST_SIZE; + let height = CHAIN_TEST_SIZE; + let world = noisy_copy(width, 0.5, 0.2, 99); + + let params = chained_params(2); + let mut denoiser = NlmDenoiser::::new(client, params, width, height); + + let real_pushes = 1 + 3 * radius as i32 + 2; + for frame_number in 0..real_pushes { + let shift = frame_number * velocity; + let frame = frame_shifted_by(&world, width, height, shift, shift); + denoiser.push_frame(&frame); + } + denoiser +} + +/// Runs chain compose for offset `k` and reads back the composed vector at the centre block. +fn composed_centre_mv(denoiser: &NlmDenoiser, center_t: u32, k: i32, neighbour_idx: u32) -> (i32, i32) { + denoiser + .run_chain_compose(center_t, k, neighbour_idx) + .expect("chain compose dispatch failed"); + + let motion_ctx = MotionCtx::new( + denoiser.params.motion_compensation, + denoiser.width, + denoiser.height, + denoiser.align, + ) + .unwrap(); + let mv_field = denoiser + .mv_field_buf + .as_ref() + .expect("mv_field allocated when mc_ctx is Some"); + let offset = mv_field_byte_offset(&motion_ctx, neighbour_idx); + let sliced = mv_field.clone().offset_start(offset); + let bytes = denoiser.client.read_one(sliced).expect("mv readback failed"); + let data = i32::from_bytes(&bytes); + + let block_x = motion_ctx.blocks_x / 2; + let block_y = motion_ctx.blocks_y / 2; + let mv_index = ((block_y * motion_ctx.blocks_x + block_x) * 2) as usize; + (data[mv_index], data[mv_index + 1]) +} + +/// Asserts the `radius` pair-ring writes starting at `ring_head_before` are zero in both directions. +/// +/// `ring_head_before` is the head the first write saw, before it advanced. Duplicated slots hold +/// zero motion by definition. +fn assert_pair_ring_zero_from(denoiser: &NlmDenoiser, ring_head_before: i32, radius: u32) { + let motion_ctx = MotionCtx::new( + denoiser.params.motion_compensation, + denoiser.width, + denoiser.height, + denoiser.align, + ) + .unwrap(); + let pair_ring = denoiser + .pair_ring_buf + .as_ref() + .expect("pair_ring allocated when Chained is active"); + let pair_ring_slots = 2 * radius as i32; + let direction_len = motion_ctx.pair_direction_len() as usize; + + for i in 0..radius as i32 { + let slot = (ring_head_before + i).rem_euclid(pair_ring_slots) as u32; + for direction in 0..2u32 { + let offset = pair_byte_offset(&motion_ctx, slot, direction); + let sliced = pair_ring.clone().offset_start(offset); + let bytes = denoiser + .client + .read_one(sliced) + .expect("pair ring readback failed"); + let data = i32::from_bytes(&bytes); + assert!( + data[..direction_len].iter().all(|&value| value == 0), + "duplicate pair slot {slot} direction {direction} should be zero-filled, got {:?}", + &data[..direction_len], + ); + } + } +} + +#[test] +fn chain_compose_zero_motion_gives_zero_mv() { + let client = make_client(); + let radius = CHAIN_TEST_RADIUS; + let denoiser = push_constant_velocity(&client, radius, 0); + + for k in 1..=radius as i32 { + let forward_idx = neighbour_idx_for_k(radius, k); + let forward_mv = composed_centre_mv(&denoiser, radius, k, forward_idx); + assert_eq!( + forward_mv, + (0, 0), + "forward k={k} should compose to zero motion on a static sequence" + ); + + let backward_idx = neighbour_idx_for_k(radius, -k); + let backward_mv = composed_centre_mv(&denoiser, radius, -k, backward_idx); + assert_eq!( + backward_mv, + (0, 0), + "backward k={k} should compose to zero motion on a static sequence" + ); + } +} + +/// A velocity of 1 is half a pixel at the coarse level, ambiguous enough for the coarse pass to +/// lock onto the wrong candidate, so this uses 2. +#[test] +fn chain_compose_constant_velocity_matches_k_times_v() { + let client = make_client(); + let radius = CHAIN_TEST_RADIUS; + let velocity = 2; + let denoiser = push_constant_velocity(&client, radius, velocity); + + for k in 1..=radius as i32 { + let forward_idx = neighbour_idx_for_k(radius, k); + let forward_mv = composed_centre_mv(&denoiser, radius, k, forward_idx); + assert_eq!( + forward_mv, + (k * velocity, k * velocity), + "forward k={k} should compose to exactly k*v = ({}, {})", + k * velocity, + k * velocity + ); + } +} + +/// From centre 0 the walk reaches `k = 2 * radius`, the far end of the ring, and lands in the last +/// neighbour slot. +#[test] +fn chain_compose_reaches_twice_the_radius_from_the_ring_start() { + let client = make_client(); + let radius = CHAIN_TEST_RADIUS; + let velocity = 2; + let denoiser = push_constant_velocity(&client, radius, velocity); + let far = 2 * radius as i32; + + let composed = composed_centre_mv(&denoiser, 0, far, 2 * radius - 1); + + assert_eq!(composed, (far * velocity, far * velocity)); +} + +#[test] +fn chain_compose_backward_direction_matches_negative_k_times_v() { + let client = make_client(); + let radius = CHAIN_TEST_RADIUS; + let velocity = 2; + let denoiser = push_constant_velocity(&client, radius, velocity); + + for k in 1..=radius as i32 { + let backward_idx = neighbour_idx_for_k(radius, -k); + let backward_mv = composed_centre_mv(&denoiser, radius, -k, backward_idx); + assert_eq!( + backward_mv, + (-k * velocity, -k * velocity), + "backward k={k} should compose to exactly -k*v = ({}, {})", + -k * velocity, + -k * velocity + ); + } +} + +/// Moving frames are pushed between priming and flush, so the check cannot pass on a still-zeroed +/// buffer. +#[test] +fn chain_compose_duplicated_slot_pairs_are_zero_filled() { + let client = make_client(); + let radius = CHAIN_TEST_RADIUS; + let width = CHAIN_TEST_SIZE; + let height = CHAIN_TEST_SIZE; + let world = noisy_copy(width, 0.5, 0.2, 7); + + let params = chained_params(2); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + + // The first push has no older partner, so the `radius` primed duplicates that follow fill the + // first gap. Their pair writes start with `ring_head` at 1. + denoiser.push_frame(&world); + assert_pair_ring_zero_from(&denoiser, 1, radius); + + // Real motion in every pair slot makes the flush's zero-fill distinguishable. + for frame_number in 1..=(3 * radius) { + let shift = frame_number as i32; + let frame = frame_shifted_by(&world, width, height, shift, shift); + denoiser.push_frame(&frame); + } + + let ring_head_before_flush = denoiser.ring_head as i32; + denoiser.flush(|_| {}).expect("flush failed"); + assert_pair_ring_zero_from(&denoiser, ring_head_before_flush, radius); +} + +fn auto_params(temporal_radius: u32) -> NlmParams { + NlmParams { + temporal_radius, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Auto, + }, + hq: None, + } +} + +#[test] +fn auto_estimation_at_high_radius_allocates_pair_ring() { + let client = make_client(); + let params = auto_params(CHAINED_RADIUS_THRESHOLD); + + let denoiser = NlmDenoiser::::new(&client, params, 32, 32); + + assert!( + denoiser.pair_ring_buf.is_some(), + "Auto at radius {CHAINED_RADIUS_THRESHOLD} (>= CHAINED_RADIUS_THRESHOLD) should \ + resolve to Chained and allocate the pair ring" + ); +} + +#[test] +fn auto_estimation_at_low_radius_does_not_allocate_pair_ring() { + let client = make_client(); + let params = auto_params(CHAINED_RADIUS_THRESHOLD - 1); + + let denoiser = NlmDenoiser::::new(&client, params, 32, 32); + + assert!( + denoiser.pair_ring_buf.is_none(), + "Auto at radius {} (< CHAINED_RADIUS_THRESHOLD) should resolve to Direct \ + and skip the pair ring", + CHAINED_RADIUS_THRESHOLD - 1 + ); +} + +fn k4_params(estimation: MotionEstimation) -> NlmParams { + NlmParams { + temporal_radius: K4_RADIUS, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 4, + pyramid_levels: 2, + estimation, + }, + hq: None, + } +} + +/// Reads frame `slot` from a ring buffer of `height * width * stored_channels` f32 frames. +fn read_frame_slot( + client: &ComputeClient, + buf: &Handle, + slot: u32, + width: u32, + height: u32, + stored_channels: u32, +) -> Vec { + let frame_size = (width * height * stored_channels) as usize; + let byte_offset = (slot as u64) * (frame_size as u64) * (size_of::() as u64); + let sliced = buf.clone().offset_start(byte_offset); + let bytes = client.read_one(sliced).expect("frame readback failed"); + f32::from_bytes(&bytes)[..frame_size].to_vec() +} + +/// Returns the mean absolute residual between the centre frame and the forward `k` neighbour's +/// warped copy, over a wrapped sequence moving `K4_V` pixels per frame. +/// +/// `2 * K4_RADIUS + 4` pushes clear the leading priming duplicates, which need more than +/// `2 * radius + 1` real pushes, with a small margin. +fn k4_compensated_residual(estimation: MotionEstimation, k: i32) -> f32 { + let client = make_client(); + let width = K4_SIZE; + let height = K4_SIZE; + let world = noisy_copy(width, 0.5, 0.2, 55); + + let params = k4_params(estimation); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + + let real_pushes = 2 * K4_RADIUS as i32 + 4; + for frame_number in 0..real_pushes { + let shift = frame_number * K4_V; + let frame = frame_shifted_wrapped(&world, width, height, shift, shift); + denoiser.push_frame(&frame); + } + denoiser.denoise().unwrap(); + + let radius = denoiser.params.temporal_radius; + let stored_channels = denoiser.params.channels.storage_count(); + let centre_slot = denoiser.phys_frame(radius as i32); + let neighbour_slot = denoiser.phys_frame(radius as i32 + k); + let compensated = denoiser + .compensated_input_buf + .as_ref() + .expect("compensated buf allocated when MC is active"); + + let centre_frame = read_frame_slot( + &denoiser.client, + &denoiser.input_buf, + centre_slot, + width, + height, + stored_channels, + ); + let warped = read_frame_slot( + &denoiser.client, + compensated, + neighbour_slot, + width, + height, + stored_channels, + ); + + let mut sum = 0.0f32; + let mut count = 0u32; + for y in 0..height { + for x in 0..width { + let index = (y * width + x) as usize; + sum += (centre_frame[index] - warped[index]).abs(); + count += 1; + } + } + sum / count as f32 +} + +/// At `k = 4` the 16 px motion is beyond direct's 12 px reach, while chained composes four exact +/// steps and only needs a refine radius of 2. +#[test] +fn chained_beats_direct_at_k4_beyond_direct_window() { + let direct_residual = k4_compensated_residual(MotionEstimation::Direct, 4); + let chained_residual = k4_compensated_residual(MotionEstimation::Chained { refine_radius: 2 }, 4); + + assert!( + direct_residual > 0.02, + "expected direct's k=4 match to show a real misalignment residual \ + (window reach ≈12px, true motion 16px), got {direct_residual}" + ); + assert!( + chained_residual < direct_residual * 0.5, + "chained's composed+refined k=4 alignment should beat direct's by a \ + wide margin: chained={chained_residual}, direct={direct_residual}" + ); +} + +/// Pins the `k = 4` result on the window size rather than chained always winning. +#[test] +fn direct_already_aligns_at_k1_inside_its_window() { + let direct_residual = k4_compensated_residual(MotionEstimation::Direct, 1); + assert!( + direct_residual < 0.02, + "direct should align cleanly at k=1 (motion {K4_V}px, well inside its ~12px reach), got {direct_residual}" + ); +} diff --git a/av-denoise-core/src/nlmeans/tests/motion_compensation/end_to_end.rs b/av-denoise-core/src/nlmeans/tests/motion_compensation/end_to_end.rs new file mode 100644 index 0000000..0efef8b --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/motion_compensation/end_to_end.rs @@ -0,0 +1,433 @@ +use crate::bench_api::HostIo; +use crate::nlmeans::motion::{ + DEFAULT_BLKSIZE, + DEFAULT_OVERLAP, + DEFAULT_PYRAMID_LEVELS, + DEFAULT_SEARCH_RADIUS, + MotionCtx, +}; +use crate::nlmeans::tests::helpers::*; +use crate::nlmeans::*; + +/// Builds a flat frame with a bright square at `(square_x, square_y)`. +fn frame_with_square( + width: u32, + height: u32, + background: f32, + square_x: u32, + square_y: u32, + square_size: u32, + square_value: f32, +) -> Vec { + let mut frame = vec![background; (width * height) as usize]; + for offset_y in 0..square_size { + for offset_x in 0..square_size { + let x = square_x + offset_x; + let y = square_y + offset_y; + if x < width && y < height { + frame[(y * width + x) as usize] = square_value; + } + } + } + frame +} + +#[test] +fn motion_compensation_uniform_passthrough() { + let client = make_client(); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 1, 0.5); + + let params = NlmParams { + temporal_radius: 1, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }, + hq: None, + }; + + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); + + assert_eq!(result.len(), (width * height) as usize); + for (i, &value) in result.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (value - 0.5).abs() < 1e-3, + "pixel {i}: expected 0.5 (uniform input passthrough), got {value}" + ); + } +} + +#[test] +fn motion_compensation_with_bilateral_finite() { + let client = make_client(); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 1, 0.5); + + let params = NlmParams { + temporal_radius: 1, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::Bilateral { + sigma_s: 1.0, + sigma_r: 0.1, + }, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }, + hq: None, + }; + + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); + + assert_eq!(result.len(), (width * height) as usize); + for (i, &value) in result.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (-0.01..=1.01).contains(&value), + "pixel {i}: out-of-range output {value}" + ); + } +} + +#[test] +fn motion_compensation_translating_square_preserves_centre() { + let client = make_client(); + let width = 32u32; + let height = 32u32; + let background = 0.3; + let square_value = 0.8; + let square_size = 4u32; + + // The square moves 2 px diagonally per frame, so without motion compensation the temporal + // kernel would see misaligned content. The centre frame's square sits at (14, 14). + let previous_frame = frame_with_square(width, height, background, 12, 12, square_size, square_value); + let centre_frame = frame_with_square(width, height, background, 14, 14, square_size, square_value); + let next_frame = frame_with_square(width, height, background, 16, 16, square_size, square_value); + + let params = NlmParams { + temporal_radius: 1, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }, + hq: None, + }; + + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&previous_frame); + denoiser.push_frame(¢re_frame); + denoiser.push_frame(&next_frame); + let result = denoiser.denoise().unwrap().unwrap(); + + assert_eq!(result.len(), (width * height) as usize); + for (i, &value) in result.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (-0.01..=1.01).contains(&value), + "pixel {i}: out-of-range output {value}" + ); + } + + // Integer-pixel warping can soften the square's edges, so the square's centre pixel only has + // to stay above halfway between the background and the square. + let halfway = (background + square_value) * 0.5; + let centre_value = result[(15 * width + 15) as usize]; + assert!( + centre_value > halfway, + "centre of moving square should remain above halfway between bg ({background}) \ + and sq_val ({square_value}) (= {halfway}), got {centre_value}", + ); + + // A neighbour square warped to the wrong place would brighten the background. + let background_value = result[(2 * width + 2) as usize]; + assert!( + (background_value - background).abs() < 0.05, + "background pixel (2, 2) should stay near {background}, got {background_value} \ + (MC may be warping neighbour squares into the background region)", + ); +} + +/// A 1080x1080 frame at the defaults has 135x135 blocks, an odd count whose unpadded +/// per-neighbour stride is not a 32-byte multiple. +/// +/// wgpu rejects a binding offset that is not a multiple of `min_storage_buffer_offset_alignment`, +/// so the second neighbour's dispatch fails unless +/// [mv_field_byte_offset](crate::nlmeans::motion::mv_field_byte_offset) pads the stride. 1920x1080 +/// gives an even block count and never reaches this bug, which is why the test needs its own size. +#[test] +fn motion_compensation_1080_square_odd_block_count_dispatch_succeeds() { + let client = make_client(); + let width = 1080u32; + let height = 1080u32; + let frame = make_uniform_frame(width, height, 1, 0.5); + + let params = NlmParams { + temporal_radius: 1, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::mvtools_default(), + hq: None, + }; + let align = test_align(); + let motion_ctx = MotionCtx::new(params.motion_compensation, width, height, align).unwrap(); + assert_eq!( + motion_ctx.blocks_x * motion_ctx.blocks_y, + 18225, + "test premise: this geometry gives an odd block count" + ); + + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); + + assert_eq!(result.len(), (width * height) as usize); + for (i, &value) in result.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (value - 0.5).abs() < 1e-3, + "pixel {i}: expected 0.5 (uniform input passthrough), got {value}" + ); + } +} + +/// The same odd block count under `Chained` estimation, which binds into the pair ring as well. +/// The other `Chained` tests all have an even block count. +/// +/// Pushing frames past the priming window binds both the zeroed duplicate slots and the real hops. +#[test] +fn motion_compensation_1080_square_odd_block_count_chained_dispatch_succeeds() { + let client = make_client(); + let width = 1080u32; + let height = 1080u32; + let radius = 2u32; + let frames: Vec> = (0..8) + .map(|i| make_frame_with_noisy_region(width, height, 1, 0.5, 200 + i * 4, 200, 8, 0.8)) + .collect(); + + let params = NlmParams { + temporal_radius: radius, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: DEFAULT_BLKSIZE, + overlap: DEFAULT_OVERLAP, + search_radius: DEFAULT_SEARCH_RADIUS, + pyramid_levels: DEFAULT_PYRAMID_LEVELS, + estimation: MotionEstimation::chained_default(), + }, + hq: None, + }; + let align = test_align(); + let motion_ctx = MotionCtx::new(params.motion_compensation, width, height, align).unwrap(); + assert_eq!( + motion_ctx.blocks_x * motion_ctx.blocks_y, + 18225, + "test premise: this geometry gives an odd block count" + ); + + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + assert!( + denoiser.pair_ring_buf.is_some(), + "test premise: Chained estimation must allocate the pair ring" + ); + + let check = |frame: &[f32]| { + for (i, &value) in frame.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (-0.01..=1.01).contains(&value), + "pixel {i}: out-of-range output {value}" + ); + } + }; + + let mut emitted = 0usize; + for frame in &frames { + denoiser.push_frame(frame); + if let Some(result) = denoiser.denoise().unwrap() { + check(&result); + emitted += 1; + } + } + + denoiser + .flush(|frame| { + check(frame); + emitted += 1; + }) + .unwrap(); + + assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); +} + +/// HQ parameters with auto noise estimation, temporal confidence and `Chained` estimation. +fn chained_hq_params(radius: u32, refine_radius: u32) -> NlmParams { + NlmParams { + temporal_radius: radius, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Chained { refine_radius }, + }, + hq: Some(HqParams { + auto_strength: true, + noise_floor: true, + sigma_override: None, + temporal_confidence: true, + thsad_scale: 1.0, + sigma_scale: 1.0, + windowed_noise_estimation: false, + }), + } +} + +/// Checks a `Chained` denoiser emits finite output in `0.0..=1.0` for every pushed frame. +fn chained_end_to_end_finite(radius: u32) { + let client = make_client(); + let width = 32u32; + let height = 32u32; + + let params = chained_hq_params(radius, 2); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + + let frames: Vec> = (0..8) + .map(|i| make_frame_with_noisy_region(width, height, 1, 0.5, 6 + i, 8, 2, 0.8)) + .collect(); + + let mut emitted = 0usize; + let check = |frame: &[f32]| { + for (i, &value) in frame.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (0.0..=1.0).contains(&value), + "pixel {i}: out-of-range output {value}" + ); + } + }; + + for frame in &frames { + denoiser.push_frame(frame); + if let Some(result) = denoiser.denoise().unwrap() { + check(&result); + emitted += 1; + } + } + + denoiser + .flush(|frame| { + check(frame); + emitted += 1; + }) + .unwrap(); + + assert_eq!(emitted, frames.len(), "expected one output per pushed frame"); +} + +#[test] +fn chained_end_to_end_finite_r2() { + chained_end_to_end_finite(2); +} + +#[test] +fn chained_end_to_end_finite_r4() { + chained_end_to_end_finite(4); +} + +#[test] +fn direct_estimation_default_and_explicit_construction_match_bit_for_bit() { + let client = make_client(); + let width = 32u32; + let height = 32u32; + let frame = make_frame_with_noisy_region(width, height, 1, 0.5, 16, 16, 4, 0.8); + + let run = |motion_compensation: MotionCompensationMode| { + let params = NlmParams { + temporal_radius: 1, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation, + hq: None, + }; + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.push_frame(&frame); + denoiser.denoise().unwrap().unwrap() + }; + + let default_mode = MotionCompensationMode::mvtools_default(); + let via_default = run(default_mode); + let explicit_mode = MotionCompensationMode::Mvtools { + blksize: DEFAULT_BLKSIZE, + overlap: DEFAULT_OVERLAP, + search_radius: DEFAULT_SEARCH_RADIUS, + pyramid_levels: DEFAULT_PYRAMID_LEVELS, + estimation: MotionEstimation::Direct, + }; + let via_explicit = run(explicit_mode); + + assert_eq!( + via_default, via_explicit, + "Direct estimation must give the same output regardless of which \ + constructor produced the MotionCompensationMode value" + ); +} diff --git a/av-denoise-core/src/nlmeans/tests/motion_compensation/mod.rs b/av-denoise-core/src/nlmeans/tests/motion_compensation/mod.rs new file mode 100644 index 0000000..c2d8a2c --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/motion_compensation/mod.rs @@ -0,0 +1,22 @@ +mod block_match; +mod chain; +mod end_to_end; +mod seeding; + +/// Clamps `value - delta` into `0..limit`, matching the kernel's clamp-to-edge reads. +fn shift_clamped(value: i32, delta: i32, limit: i32) -> i32 { + (value - delta).clamp(0, limit - 1) +} + +/// Builds a frame that reads `world` shifted by `(shift_x, shift_y)` pixels, clamped to its edges. +fn frame_shifted_by(world: &[f32], width: u32, height: u32, shift_x: i32, shift_y: i32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let source_x = shift_clamped(x as i32, shift_x, width as i32) as u32; + let source_y = shift_clamped(y as i32, shift_y, height as i32) as u32; + frame[(y * width + x) as usize] = world[(source_y * width + source_x) as usize]; + } + } + frame +} diff --git a/av-denoise-core/src/nlmeans/tests/motion_compensation/seeding.rs b/av-denoise-core/src/nlmeans/tests/motion_compensation/seeding.rs new file mode 100644 index 0000000..27513e9 --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/motion_compensation/seeding.rs @@ -0,0 +1,396 @@ +use cubecl::prelude::*; + +use super::{frame_shifted_by, shift_clamped}; +use crate::bench_api::HostIo; +use crate::nlmeans::motion::{ + DEFAULT_BLKSIZE, + DEFAULT_OVERLAP, + DEFAULT_SEARCH_RADIUS, + MotionCtx, + mv_field_byte_offset, + neighbour_idx_for_k, +}; +use crate::nlmeans::tests::helpers::*; +use crate::nlmeans::*; + +/// Builds a frame from two independent halves, each shifted within its own half. +/// +/// Clamping per half keeps a block deep inside one half from ever depending on the other. +fn split_half_frame( + width: u32, + height: u32, + half: u32, + left: &[f32], + right: &[f32], + left_shift: i32, + right_shift: i32, +) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let index = (y * width + x) as usize; + if x < half { + let left_x = shift_clamped(x as i32, left_shift, half as i32) as u32; + frame[index] = left[(y * half + left_x) as usize]; + } else { + let right_x = shift_clamped((x - half) as i32, right_shift, half as i32) as u32; + frame[index] = right[(y * half + right_x) as usize]; + } + } + } + frame +} + +/// Pushes `base` twice then `neighbour` through a `Direct` denoiser and reads back the forward +/// neighbour's motion field. +/// +/// The second push is the centre frame, so the field matches `base` against `neighbour`. +fn direct_mv_field_for_forward_neighbour( + mode: MotionCompensationMode, + width: u32, + height: u32, + base: &[f32], + neighbour: &[f32], +) -> Vec { + let params = NlmParams { + temporal_radius: 1, + search_radius: 2, + patch_radius: 2, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: mode, + hq: None, + }; + + let client = make_client(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(base); + denoiser.push_frame(base); + denoiser.push_frame(neighbour); + denoiser.denoise().unwrap(); + + let align = test_align(); + let motion_ctx = MotionCtx::new(mode, width, height, align).unwrap(); + let neighbour_idx = neighbour_idx_for_k(1, 1); + let mv_field = denoiser + .mv_field_buf + .as_ref() + .expect("mv_field allocated when mc_ctx is Some"); + let offset = mv_field_byte_offset(&motion_ctx, neighbour_idx); + let sliced = mv_field.clone().offset_start(offset); + let bytes = denoiser.client.read_one(sliced).expect("mv readback failed"); + i32::from_bytes(&bytes).to_vec() +} + +/// At this geometry the coarse and fine grids have the same block count, so seeding must map by +/// position rather than by doubling the index. +/// +/// Each half moves by its own amount, and a block seeded from the wrong half lands outside the +/// fine search window. Both shifts are multiples of the pyramid scale, so the coarse level sees an +/// exact half-size copy. Left-half blocks always seed from left-half blocks here, so only the right +/// half shows the defect. +#[test] +fn coarse_seeding_handles_equal_grids() { + let width = 128u32; + let height = 64u32; + let half = 64u32; + + let mode = MotionCompensationMode::Mvtools { + blksize: 8, + overlap: 4, + search_radius: 3, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }; + let align = test_align(); + let motion_ctx = MotionCtx::new(mode, width, height, align).unwrap(); + + // Recomputes the coarse grid width to confirm this geometry gives equal grids. + let coarse_scale = 1u32 << (motion_ctx.pyramid_levels - 1); + let coarse_step = (motion_ctx.step / coarse_scale).max(1); + let coarse_width = width / coarse_scale; + let coarse_blocks_x = coarse_width.div_ceil(coarse_step).max(1); + assert_eq!( + coarse_blocks_x, motion_ctx.blocks_x, + "test premise: this geometry must give equal coarse/fine grids" + ); + + let left = noisy_copy(half, 0.5, 0.2, 201); + let right = noisy_copy(half, 0.5, 0.2, 202); + let left_shift = 4i32; + let right_shift = -4i32; + + let base = split_half_frame(width, height, half, &left, &right, 0, 0); + let shifted = split_half_frame(width, height, half, &left, &right, left_shift, right_shift); + + let data = direct_mv_field_for_forward_neighbour(mode, width, height, &base, &shifted); + + // Deep inside each half, clear of the boundary at x=64 and the frame edges by well over + // `search_radius + blksize` (11). + let block_y = 4u32; + let left_block_x = 8u32; + let right_block_x = 24u32; + let left_index = ((block_y * motion_ctx.blocks_x + left_block_x) * 2) as usize; + let right_index = ((block_y * motion_ctx.blocks_x + right_block_x) * 2) as usize; + + assert_eq!( + (data[left_index], data[left_index + 1]), + (left_shift, 0), + "left-half block should recover the left half's own motion ({left_shift}, 0), got ({}, {})", + data[left_index], + data[left_index + 1], + ); + assert_eq!( + (data[right_index], data[right_index + 1]), + (right_shift, 0), + "right-half block should recover the right half's own motion ({right_shift}, 0), got \ + ({}, {}); a wrong value here means it was seeded from the wrong (left-half) coarse block", + data[right_index], + data[right_index + 1], + ); +} + +/// At a step of 1 the coarse step floors to 1, giving a genuine 2:1 fine to coarse grid. +/// +/// Position-based seeding must then reduce to plain index doubling, both for uniform motion and +/// for halves moving in opposite directions. +#[test] +fn coarse_seeding_still_correct_at_half_grid() { + let width = 48u32; + let height = 16u32; + let half = 24u32; + + let mode = MotionCompensationMode::Mvtools { + blksize: 4, + overlap: 3, + search_radius: 2, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }; + let align = test_align(); + let motion_ctx = MotionCtx::new(mode, width, height, align).unwrap(); + + let coarse_scale = 1u32 << (motion_ctx.pyramid_levels - 1); + let coarse_step = (motion_ctx.step / coarse_scale).max(1); + let coarse_width = width / coarse_scale; + let coarse_height = height / coarse_scale; + let coarse_blocks_x = coarse_width.div_ceil(coarse_step).max(1); + let coarse_blocks_y = coarse_height.div_ceil(coarse_step).max(1); + assert_eq!( + motion_ctx.step, 1, + "test premise: step must floor-clamp coarse_step to 1" + ); + assert_eq!( + motion_ctx.blocks_x, + 2 * coarse_blocks_x, + "test premise: this geometry must give a genuine 2:1 fine:coarse ratio in x" + ); + assert_eq!( + motion_ctx.blocks_y, + 2 * coarse_blocks_y, + "test premise: this geometry must give a genuine 2:1 fine:coarse ratio in y" + ); + + let left = noisy_copy(half, 0.5, 0.2, 301); + let right = noisy_copy(half, 0.5, 0.2, 302); + + // Deep inside each half, clear of the boundary at x=24 and the frame edges by well over + // `search_radius + blksize` (6). + let block_y = 6u32; + let left_block_x = 10u32; + let right_block_x = 32u32; + + let mv_at = |left_shift: i32, right_shift: i32| -> ((i32, i32), (i32, i32)) { + let base = split_half_frame(width, height, half, &left, &right, 0, 0); + let shifted = split_half_frame(width, height, half, &left, &right, left_shift, right_shift); + let data = direct_mv_field_for_forward_neighbour(mode, width, height, &base, &shifted); + let left_index = ((block_y * motion_ctx.blocks_x + left_block_x) * 2) as usize; + let right_index = ((block_y * motion_ctx.blocks_x + right_block_x) * 2) as usize; + ( + (data[left_index], data[left_index + 1]), + (data[right_index], data[right_index + 1]), + ) + }; + + // Both halves share one vector here, so any coarse block seeds any fine block correctly. + let (uniform_left, uniform_right) = mv_at(2, 2); + assert_eq!( + uniform_left, + (2, 0), + "uniform motion: left block got {uniform_left:?}" + ); + assert_eq!( + uniform_right, + (2, 0), + "uniform motion: right block got {uniform_right:?}" + ); + + // The halves move oppositely, so a fine block only recovers its own half's motion when it is + // seeded from the coarse block covering the same region. + let (varying_left, varying_right) = mv_at(2, -2); + assert_eq!( + varying_left, + (2, 0), + "varying motion: left block got {varying_left:?}" + ); + assert_eq!( + varying_right, + (-2, 0), + "varying motion: right block got {varying_right:?}" + ); +} + +/// With a single pyramid level the fine pass runs unseeded straight off level 0, so a skipped +/// level 0 extraction leaves the recovered vector unrelated to the shift. +#[test] +fn pyramid_level0_extracted_at_one_level() { + let width = 64u32; + let height = 64u32; + let shift_x = 2i32; + let shift_y = 1i32; + + let mode = MotionCompensationMode::Mvtools { + blksize: DEFAULT_BLKSIZE, + overlap: DEFAULT_OVERLAP, + search_radius: DEFAULT_SEARCH_RADIUS, + pyramid_levels: 1, + estimation: MotionEstimation::Direct, + }; + let align = test_align(); + let motion_ctx = MotionCtx::new(mode, width, height, align).unwrap(); + + let world = noisy_copy(width, 0.5, 0.2, 77); + let shifted = frame_shifted_by(&world, width, height, shift_x, shift_y); + + let data = direct_mv_field_for_forward_neighbour(mode, width, height, &world, &shifted); + + // An interior block, well clear of the frame edges. + let block_x = motion_ctx.blocks_x / 2; + let block_y = motion_ctx.blocks_y / 2; + let mv_index = ((block_y * motion_ctx.blocks_x + block_x) * 2) as usize; + + assert_eq!( + (data[mv_index], data[mv_index + 1]), + (shift_x, shift_y), + "a clean ({shift_x}, {shift_y}) shift with pyramid_levels=1 should give exactly that MV at an \ + interior block once level-0 luma is actually extracted, got ({}, {})", + data[mv_index], + data[mv_index + 1], + ); +} + +/// The coarse and fine block counts round up over different widths and steps, so at some sizes +/// the coarse grid ends one fine block short and its last block must reach the trailing edge. +/// +/// An unseeded trailing block keeps whatever the motion field already held. The shift escapes the +/// unseeded search window but is reachable through a correct seed. It is negative because a +/// trailing-edge block can only tell apart offsets that pull content in from the interior, so a +/// positive shift there is unrecoverable by any search. +#[test] +fn coarse_seeding_covers_ragged_last_block() { + let shift = -6i32; + let mode = MotionCompensationMode::Mvtools { + blksize: DEFAULT_BLKSIZE, + overlap: DEFAULT_OVERLAP, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }; + + // A ragged side of 57 is one more than a multiple of the step (8), so the coarse grid ends one + // block short on that axis. The even side of 64 is an exact multiple. Each case is ragged on + // one axis only, because two odd sides can never give a pixel count that meets the ring + // buffers' 32-byte frame-stride alignment. + let ragged_side = 57u32; + let even_side = 64u32; + + let build_mv_field = |width: u32, height: u32| -> (MotionCtx, Vec) { + let align = test_align(); + let motion_ctx = MotionCtx::new(mode, width, height, align).unwrap(); + let world = make_noisy_gaussian_frame(width, height, 1, 0.5, &[0.2]); + let shifted = frame_shifted_by(&world, width, height, shift, shift); + let data = direct_mv_field_for_forward_neighbour(mode, width, height, &world, &shifted); + (motion_ctx, data) + }; + let mv_at = |motion_ctx: &MotionCtx, data: &[i32], block_x: u32, block_y: u32| -> (i32, i32) { + let mv_index = ((block_y * motion_ctx.blocks_x + block_x) * 2) as usize; + (data[mv_index], data[mv_index + 1]) + }; + // Recomputes the coarse grid to confirm the named axis is ragged and the other is equal. + let assert_ragged_on = |motion_ctx: &MotionCtx, width: u32, height: u32, ragged_axis_is_x: bool| { + let coarse_scale = 1u32 << (motion_ctx.pyramid_levels - 1); + let coarse_step = (motion_ctx.step / coarse_scale).max(1); + let coarse_blocks_x = (width / coarse_scale).div_ceil(coarse_step).max(1); + let coarse_blocks_y = (height / coarse_scale).div_ceil(coarse_step).max(1); + + if ragged_axis_is_x { + assert_eq!( + width % motion_ctx.step, + 1, + "test premise: width must be step*k + 1" + ); + assert_eq!( + coarse_blocks_x, + motion_ctx.blocks_x - 1, + "test premise: ragged coarse grid, one block short in x" + ); + assert_eq!( + coarse_blocks_y, motion_ctx.blocks_y, + "test premise: y axis is an ordinary equal grid here" + ); + } else { + assert_eq!( + height % motion_ctx.step, + 1, + "test premise: height must be step*k + 1" + ); + assert_eq!( + coarse_blocks_y, + motion_ctx.blocks_y - 1, + "test premise: ragged coarse grid, one block short in y" + ); + assert_eq!( + coarse_blocks_x, motion_ctx.blocks_x, + "test premise: x axis is an ordinary equal grid here" + ); + } + }; + + // X-axis case, the last column at a non-edge row. + let (x_case_ctx, x_case_field) = build_mv_field(ragged_side, even_side); + assert_ragged_on(&x_case_ctx, ragged_side, even_side, true); + let middle_x = x_case_ctx.blocks_x / 2; + let middle_y = x_case_ctx.blocks_y / 2; + let interior_mv = mv_at(&x_case_ctx, &x_case_field, middle_x, middle_y); + assert_eq!( + interior_mv, + (shift, shift), + "interior control block (x-axis case) should recover ({shift}, {shift})" + ); + let last_column_mv = mv_at(&x_case_ctx, &x_case_field, x_case_ctx.blocks_x - 1, middle_y); + assert_eq!( + last_column_mv, + (shift, shift), + "last-column block (x-axis coverage gap) should recover ({shift}, {shift})" + ); + + // Y-axis case, the last row at a non-edge column, using the transposed frame size. + let (y_case_ctx, y_case_field) = build_mv_field(even_side, ragged_side); + assert_ragged_on(&y_case_ctx, even_side, ragged_side, false); + let middle_x = y_case_ctx.blocks_x / 2; + let middle_y = y_case_ctx.blocks_y / 2; + let interior_mv = mv_at(&y_case_ctx, &y_case_field, middle_x, middle_y); + assert_eq!( + interior_mv, + (shift, shift), + "interior control block (y-axis case) should recover ({shift}, {shift})" + ); + let last_row_mv = mv_at(&y_case_ctx, &y_case_field, middle_x, y_case_ctx.blocks_y - 1); + assert_eq!( + last_row_mv, + (shift, shift), + "last-row block (y-axis coverage gap) should recover ({shift}, {shift})" + ); +} diff --git a/av-denoise-core/src/nlmeans/tests/noise.rs b/av-denoise-core/src/nlmeans/tests/noise.rs index b552f09..d8b852d 100644 --- a/av-denoise-core/src/nlmeans/tests/noise.rs +++ b/av-denoise-core/src/nlmeans/tests/noise.rs @@ -2,24 +2,26 @@ use cubecl::prelude::*; use super::helpers::*; use crate::nlmeans::noise::{NoiseCtx, partials_len, run_noise_estimate, sigma_from_abs_sum}; -use crate::nlmeans::{Depth, denormalize, normalize}; -/// Uploads `dense` (packed `pixels * ch`) as the padded GPU storage -/// layout and runs both noise-estimate stages for a single frame in a -/// one-slot ring. Returns the four raw per-lane absolute-sum totals. -fn estimate_abs_sums(w: u32, h: u32, ch: u32, stored_ch: u32, dense: &[f32]) -> [f32; 4] { +/// Runs both noise-estimate stages over one frame in a one-slot ring. +/// +/// `dense` is packed as `pixels * channels` and is padded to `stored_ch` on upload. Returns the four +/// raw per-lane absolute-sum totals. +fn estimate_abs_sums(width: u32, height: u32, channels: u32, stored_ch: u32, dense: &[f32]) -> [f32; 4] { let client = make_client(); - let pixels = (w * h) as usize; - let padded = pad_channels(dense, pixels, ch, stored_ch); + let pixels = (width * height) as usize; + let padded = pad_channels(dense, pixels, channels, stored_ch); - let input_buf = client.create_from_slice(f32::as_bytes(&padded)); - let partials_buf = client.empty(partials_len(w, h) * size_of::()); + let input_bytes = f32::as_bytes(&padded); + let partials_bytes = partials_len(width, height) * size_of::(); + let input_buf = client.create_from_slice(input_bytes); + let partials_buf = client.empty(partials_bytes); let results_buf = client.empty(4 * size_of::()); let ctx = NoiseCtx { - width: w, - height: h, - channels: ch, + width, + height, + channels, stored_ch, frame_count: 1, frame: 0, @@ -36,15 +38,20 @@ fn estimate_abs_sums(w: u32, h: u32, ch: u32, stored_ch: u32, dense: &[f32]) -> [data[0], data[1], data[2], data[3]] } +/// The luma sigma measured over `frame`. +fn estimate_luma_sigma(width: u32, height: u32, frame: &[f32]) -> f32 { + let sums = estimate_abs_sums(width, height, 1, 1, frame); + sigma_from_abs_sum(sums[0], width, height) +} + #[test] fn noise_estimate_recovers_known_sigma() { - let w = 256; - let h = 256; + let width = 256; + let height = 256; let true_sigma = 8.0 / 255.0; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[true_sigma]); - let sums = estimate_abs_sums(w, h, 1, 1, &frame); - let estimated = sigma_from_abs_sum(sums[0], w, h); + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[true_sigma]); + let estimated = estimate_luma_sigma(width, height, &frame); let rel_err = (estimated - true_sigma).abs() / true_sigma; assert!( @@ -53,16 +60,13 @@ fn noise_estimate_recovers_known_sigma() { ); } -/// A perfectly uniform frame has zero mask response everywhere, so the -/// estimate must land below the noise floor. #[test] fn noise_estimate_zero_on_uniform() { - let w = 128; - let h = 128; + let width = 128; + let height = 128; - let frame = make_uniform_frame(w, h, 1, 0.5); - let sums = estimate_abs_sums(w, h, 1, 1, &frame); - let estimated = sigma_from_abs_sum(sums[0], w, h); + let frame = make_uniform_frame(width, height, 1, 0.5); + let estimated = estimate_luma_sigma(width, height, &frame); assert!( estimated < 0.2 / 255.0, @@ -70,22 +74,21 @@ fn noise_estimate_zero_on_uniform() { ); } -/// YUV storage. Distinct per-channel sigmas must be recovered -/// independently and the unused 4th (padding) lane must return exactly -/// zero (stage 1 zeroes it explicitly for every thread). +/// The unused 4th lane must read exactly zero because stage 1 zeroes it for every thread. #[test] fn noise_estimate_per_channel() { - let w = 256; - let h = 256; + let width = 256; + let height = 256; let true_sigma_y = 8.0 / 255.0; let true_sigma_uv = 2.0 / 255.0; - let frame = make_noisy_gaussian_frame(w, h, 3, 0.5, &[true_sigma_y, true_sigma_uv, true_sigma_uv]); - let sums = estimate_abs_sums(w, h, 3, 4, &frame); + let sigmas = [true_sigma_y, true_sigma_uv, true_sigma_uv]; + let frame = make_noisy_gaussian_frame(width, height, 3, 0.5, &sigmas); + let sums = estimate_abs_sums(width, height, 3, 4, &frame); - let estimated_y = sigma_from_abs_sum(sums[0], w, h); - let estimated_u = sigma_from_abs_sum(sums[1], w, h); - let estimated_v = sigma_from_abs_sum(sums[2], w, h); + let estimated_y = sigma_from_abs_sum(sums[0], width, height); + let estimated_u = sigma_from_abs_sum(sums[1], width, height); + let estimated_v = sigma_from_abs_sum(sums[2], width, height); let rel_err_y = (estimated_y - true_sigma_y).abs() / true_sigma_y; let rel_err_u = (estimated_u - true_sigma_uv).abs() / true_sigma_uv; @@ -106,29 +109,23 @@ fn noise_estimate_per_channel() { assert_eq!(sums[3], 0.0, "padding lane must be exactly zero, got {}", sums[3]); } -/// Smooth linear gradient plus noise. The mask is orthogonal to affine -/// content so this mainly documents that gradient content doesn't blow -/// up the estimate, bounding the known content bias rather than hiding -/// it. +/// The mask is orthogonal to affine content, so this bounds the known content bias on a gradient. #[test] fn noise_estimate_gradient_bias_bounded() { - let w = 256; - let h = 256; + let width = 256; + let height = 256; let true_sigma = 6.0 / 255.0; - // Noise is generated around a mid-gray base (safely away from the - // clamp bounds for this sigma) and re-centred to zero-mean before - // being layered onto the gradient, so the final clamp never clips - // and doesn't disturb the gradient's linearity. - let gradient = make_gradient_frame(w, h, 0.2, 0.8); - let noise_frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[true_sigma]); + // Noise is built around mid-grey, clear of the clamp bounds, then re-centred to zero mean before + // it is layered onto the gradient. The final clamp then never clips the gradient. + let gradient = make_gradient_frame(width, height, 0.2, 0.8); + let noise_frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[true_sigma]); let frame: Vec = gradient .iter() .zip(noise_frame.iter()) - .map(|(&g, &n)| (g + (n - 0.5)).clamp(0.0, 1.0)) + .map(|(&gradient_value, &noise_value)| (gradient_value + (noise_value - 0.5)).clamp(0.0, 1.0)) .collect(); - let sums = estimate_abs_sums(w, h, 1, 1, &frame); - let estimated = sigma_from_abs_sum(sums[0], w, h); + let estimated = estimate_luma_sigma(width, height, &frame); let rel_err = (estimated - true_sigma).abs() / true_sigma; assert!( @@ -137,29 +134,31 @@ fn noise_estimate_gradient_bias_bounded() { ); } -/// Quantises a normalised frame through a given depth and back, the -/// round trip a real source of that depth goes through. -fn requantise(frame: &[f32], depth: Depth) -> Vec { - normalize(&denormalize(frame, depth), depth) +/// Quantises a normalised frame to `bits` and back, as a real source of that depth would be. +fn requantise(frame: &[f32], bits: u32) -> Vec { + let max = ((1u32 << bits) - 1) as f32; + + frame + .iter() + .map(|&value| (value * max).round().clamp(0.0, max) / max) + .collect() } -/// The same content at 8-bit and at 10-bit must yield the same measured -/// sigma. Normalising by `(1 << bits) - 1` holds the normalised scale -/// fixed across depths, which is what lets every calibrated constant in -/// the library stay depth-independent. +/// Normalising by `(1 << bits) - 1` holds the scale fixed across depths, so calibrated constants +/// stay depth-independent. #[test] fn sigma_estimate_agrees_across_bit_depths() { - let w = 256; - let h = 256; + let width = 256; + let height = 256; let true_sigma = 8.0 / 255.0; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[true_sigma]); + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[true_sigma]); - let eight = requantise(&frame, Depth::Eight); - let ten = requantise(&frame, Depth::Ten); + let eight = requantise(&frame, 8); + let ten = requantise(&frame, 10); - let sigma_eight = sigma_from_abs_sum(estimate_abs_sums(w, h, 1, 1, &eight)[0], w, h); - let sigma_ten = sigma_from_abs_sum(estimate_abs_sums(w, h, 1, 1, &ten)[0], w, h); + let sigma_eight = estimate_luma_sigma(width, height, &eight); + let sigma_ten = estimate_luma_sigma(width, height, &ten); let rel_diff = (sigma_eight - sigma_ten).abs() / sigma_eight; assert!( @@ -167,31 +166,28 @@ fn sigma_estimate_agrees_across_bit_depths() { "8-bit sigma {sigma_eight} vs 10-bit sigma {sigma_ten} (rel diff {rel_diff:.4})" ); - // Both must still recover the true sigma, not merely agree with - // each other on a wrong answer. - for (label, s) in [("8-bit", sigma_eight), ("10-bit", sigma_ten)] { - let err = (s - true_sigma).abs() / true_sigma; - assert!(err <= 0.20, "{label} sigma {s} vs true {true_sigma}"); + // Both must still recover the true sigma, not just agree on a wrong answer. + for (label, sigma) in [("8-bit", sigma_eight), ("10-bit", sigma_ten)] { + let err = (sigma - true_sigma).abs() / true_sigma; + assert!(err <= 0.20, "{label} sigma {sigma} vs true {true_sigma}"); } } -/// Grain too fine for 8-bit to represent survives 10-bit quantisation. +/// The sigma is half an 8-bit step, which 10-bit resolves across two of its own steps. /// -/// The chosen sigma is half of one 8-bit step (1/255), so 8-bit cannot -/// carry it, while 10-bit resolves it across two of its own steps -/// (1/1023). It also sits 5x above `SIGMA_FLOOR` (0.1/255), so a pass -/// shows the floor is not clipping a legitimate fine measurement. +/// It sits 5x above `SIGMA_FLOOR` (0.1/255), so a pass also shows the floor does not clip a real +/// fine measurement. #[test] fn fine_grain_survives_ten_bit_quantisation() { - let w = 256; - let h = 256; + let width = 256; + let height = 256; let true_sigma = 0.5 / 255.0; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[true_sigma]); - let ten = requantise(&frame, Depth::Ten); - let eight = requantise(&frame, Depth::Eight); + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[true_sigma]); + let ten = requantise(&frame, 10); + let eight = requantise(&frame, 8); - let estimated = sigma_from_abs_sum(estimate_abs_sums(w, h, 1, 1, &ten)[0], w, h); + let estimated = estimate_luma_sigma(width, height, &ten); let err = (estimated - true_sigma).abs() / true_sigma; assert!( @@ -199,11 +195,9 @@ fn fine_grain_survives_ten_bit_quantisation() { "10-bit estimate {estimated} vs true {true_sigma} (rel err {err:.3})" ); - // The same grain through 8-bit picks up quantisation noise worth - // more than half its own amplitude, so the estimate inflates. That - // gap is the point of the depth, and it must be visible here or the - // test above is measuring nothing 8-bit could not already do. - let estimated_eight = sigma_from_abs_sum(estimate_abs_sums(w, h, 1, 1, &eight)[0], w, h); + // Through 8-bit the grain picks up quantisation noise worth more than half its amplitude, so the + // estimate inflates. Without that gap the 10-bit check above proves nothing. + let estimated_eight = estimate_luma_sigma(width, height, &eight); assert!( estimated_eight > estimated, diff --git a/av-denoise-core/src/nlmeans/tests/noise_curve.rs b/av-denoise-core/src/nlmeans/tests/noise_curve.rs index cc80ea7..2600edd 100644 --- a/av-denoise-core/src/nlmeans/tests/noise_curve.rs +++ b/av-denoise-core/src/nlmeans/tests/noise_curve.rs @@ -1,4 +1,5 @@ use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::noise::{ NoiseCurve, QUARTER_FLATNESS, @@ -83,10 +84,9 @@ pub(super) fn write_quarter( record[base + QUARTER_TENSOR_XY as usize] = quarter.tensor_xy; } -/// Records for a single row of full 16x16 blocks, luma only (`stored_ch` 1). +/// Records for a single row of full 16x16 luma-only blocks. /// -/// Every four quarters form one block, in top-left, top-right, -/// bottom-left, bottom-right order. +/// Every four quarters form one block, in top-left, top-right, bottom-left, bottom-right order. pub(super) fn synthetic_records(quarters: &[SyntheticQuarter]) -> Vec { assert_eq!(quarters.len() % 4, 0, "quarters must fill whole blocks"); @@ -126,8 +126,9 @@ fn curve_for(quarters: &[SyntheticQuarter]) -> Option { build_noise_curve(&records, stored_ch, &accepted, sample.sigma[0]) } -/// Gives each block the residual means `+m`, `+m`, `-m`, `-m`, so every -/// block's mean stays 0 while its 16x16 sigma reads above its quarters'. +/// Gives each block the residual means `+m`, `+m`, `-m`, `-m`. +/// +/// Every block's mean stays 0 while its 16x16 sigma reads above its quarters'. fn with_cancelling_means(quarters: &mut [SyntheticQuarter], mean: f32) { for (index, quarter) in quarters.iter_mut().enumerate() { quarter.mean_residual = if index % 4 < 2 { mean } else { -mean }; @@ -147,14 +148,14 @@ fn reading_sample_equals_the_scalar_aggregation() { let mut records = synthetic_records(&good_quarters); - // Two blocks with a large mean residual, so the static gate rejects - // them before either function ever reads their sigma. + // Two blocks with a large mean residual, so the static gate rejects them before either function + // reads their sigma. let record_len = temporal_stats_record_len(stored_ch) as usize; - let n = 256.0f32; - let mean0_bad = 10.0 / 255.0; + let block_pixels = 256.0f32; + let bad_mean = 10.0 / 255.0; for _ in 0..2 { let mut bad = vec![0.0f32; record_len]; - bad[0] = n * mean0_bad; + bad[0] = block_pixels * bad_mean; records.extend(bad); } @@ -170,25 +171,27 @@ fn reading_sample_equals_the_scalar_aggregation() { #[test] fn curve_uses_the_median_per_bin_and_normalises() { let mut quarters = quarters_at(0.15, 0.02, 160); - quarters.extend(quarters_at(0.35, 0.01, 160)); - quarters.extend(quarters_at(0.6, 0.005, 160)); + let middle_band = quarters_at(0.35, 0.01, 160); + let bright_band = quarters_at(0.6, 0.005, 160); + quarters.extend(middle_band); + quarters.extend(bright_band); let curve = curve_for(&quarters).expect("three populated bins"); - let tol = 1e-6; - assert!((curve.ratios[2] - 2.0).abs() < tol, "{:?}", curve.ratios); - assert!((curve.ratios[5] - 1.0).abs() < tol, "{:?}", curve.ratios); - assert!((curve.ratios[9] - 0.5).abs() < tol, "{:?}", curve.ratios); + let tolerance = 1e-6; + assert!((curve.ratios[2] - 2.0).abs() < tolerance, "{:?}", curve.ratios); + assert!((curve.ratios[5] - 1.0).abs() < tolerance, "{:?}", curve.ratios); + assert!((curve.ratios[9] - 0.5).abs() < tolerance, "{:?}", curve.ratios); // The flat low and high ends, held at the nearest populated bin. - assert!((curve.ratios[0] - 2.0).abs() < tol); - assert!((curve.ratios[1] - 2.0).abs() < tol); - assert!((curve.ratios[15] - 0.5).abs() < tol); + assert!((curve.ratios[0] - 2.0).abs() < tolerance); + assert!((curve.ratios[1] - 2.0).abs() < tolerance); + assert!((curve.ratios[15] - 0.5).abs() < tolerance); // Linear interpolation between the populated bins. let expected = 2.0 + (1.0 - 2.0) / 3.0; assert!( - (curve.ratios[3] - expected).abs() < tol, + (curve.ratios[3] - expected).abs() < tolerance, "{} vs {expected}", curve.ratios[3] ); @@ -196,64 +199,72 @@ fn curve_uses_the_median_per_bin_and_normalises() { #[test] fn each_quarter_is_binned_by_its_own_mean_luma() { - // Every block mixes one quarter from each bin, so binning by the - // block's mean luma would put them all in one bin. + // Every block mixes one quarter from each bin, so binning by the block's mean luma would put them + // all in one bin. let mut quarters = Vec::new(); for _ in 0..40 { - quarters.push(quarter_at(0.15, 0.02)); - quarters.push(quarter_at(0.35, 0.01)); - quarters.push(quarter_at(0.6, 0.005)); - quarters.push(quarter_at(0.9, 0.01)); + let dark = quarter_at(0.15, 0.02); + let middle = quarter_at(0.35, 0.01); + let bright = quarter_at(0.6, 0.005); + let brightest = quarter_at(0.9, 0.01); + quarters.push(dark); + quarters.push(middle); + quarters.push(bright); + quarters.push(brightest); } let curve = curve_for(&quarters).expect("four populated bins"); - let tol = 1e-6; - assert!((curve.ratios[2] - 2.0).abs() < tol, "{:?}", curve.ratios); - assert!((curve.ratios[5] - 1.0).abs() < tol, "{:?}", curve.ratios); - assert!((curve.ratios[9] - 0.5).abs() < tol, "{:?}", curve.ratios); - assert!((curve.ratios[14] - 1.0).abs() < tol, "{:?}", curve.ratios); + let tolerance = 1e-6; + assert!((curve.ratios[2] - 2.0).abs() < tolerance, "{:?}", curve.ratios); + assert!((curve.ratios[5] - 1.0).abs() < tolerance, "{:?}", curve.ratios); + assert!((curve.ratios[9] - 0.5).abs() < tolerance, "{:?}", curve.ratios); + assert!((curve.ratios[14] - 1.0).abs() < tolerance, "{:?}", curve.ratios); } #[test] fn a_bin_with_fewer_than_32_quarters_is_empty() { - let base = |third_bin: usize| { + let build = |third_bin: usize| { let mut quarters = quarters_at(0.15, 0.02, 40); - quarters.extend(quarters_at(0.35, 0.02, 40)); - quarters.extend(quarters_at(0.6, 0.02, third_bin)); + let middle_band = quarters_at(0.35, 0.02, 40); + let bright_band = quarters_at(0.6, 0.02, third_bin); + quarters.extend(middle_band); + quarters.extend(bright_band); + // Pads the frame to whole blocks with a bin that never fills. let padding = 4 - third_bin % 4; - quarters.extend(quarters_at(0.95, 0.02, padding)); + let padding_band = quarters_at(0.95, 0.02, padding); + quarters.extend(padding_band); quarters }; - let short = base(31); - assert!( - curve_for(&short).is_none(), - "only 2 of the 3 bins reach 32 quarters" - ); + let short = build(31); + let short_curve = curve_for(&short); + assert!(short_curve.is_none(), "only 2 of the 3 bins reach 32 quarters"); - let enough = base(32); - assert!(curve_for(&enough).is_some(), "32 quarters fill the third bin"); + let enough = build(32); + let enough_curve = curve_for(&enough); + assert!(enough_curve.is_some(), "32 quarters fill the third bin"); } #[test] fn fewer_than_three_bins_gives_none() { let mut quarters = quarters_at(0.15, 0.02, 40); - quarters.extend(quarters_at(0.35, 0.02, 40)); + let middle_band = quarters_at(0.35, 0.02, 40); + quarters.extend(middle_band); - assert!( - curve_for(&quarters).is_none(), - "only 2 bins populated, below the minimum of 3" - ); + let curve = curve_for(&quarters); + assert!(curve.is_none(), "only 2 bins populated, below the minimum of 3"); } -/// Two full bins plus a third bin that reaches 32 quarters only if all -/// five `odd_one` quarters count. Each odd quarter shares its block with -/// three plain ones, so its parent block is still accepted. +/// Two full bins plus a third that reaches 32 quarters only if all five `odd_one` quarters count. +/// +/// Each odd quarter shares its block with three plain ones, so its parent block is still accepted. fn frame_with_odd_quarters(odd_one: SyntheticQuarter) -> Vec { let mut quarters = quarters_at(0.15, 0.02, 40); - quarters.extend(quarters_at(0.35, 0.02, 40)); + let middle_band = quarters_at(0.35, 0.02, 40); + quarters.extend(middle_band); + for index in 0..36 { let replaced = index % 4 == 3 && index < 20; let quarter = if replaced { odd_one } else { quarter_at(0.6, 0.02) }; @@ -263,19 +274,17 @@ fn frame_with_odd_quarters(odd_one: SyntheticQuarter) -> Vec { quarters } -/// Asserts the frame from [frame_with_odd_quarters] forms a curve with -/// plain quarters, and none with `gated` ones, which shows the gate kept -/// them out. +/// Asserts [frame_with_odd_quarters] forms a curve with plain quarters and none with `gated` ones. fn assert_gate_keeps_the_bin_empty(gated: SyntheticQuarter) { - let plain = frame_with_odd_quarters(quarter_at(0.6, 0.02)); - assert!( - curve_for(&plain).is_some(), - "the ungated frame should form a curve" - ); + let plain_quarter = quarter_at(0.6, 0.02); + let plain = frame_with_odd_quarters(plain_quarter); + let plain_curve = curve_for(&plain); + assert!(plain_curve.is_some(), "the ungated frame should form a curve"); let with_gated = frame_with_odd_quarters(gated); + let gated_curve = curve_for(&with_gated); assert!( - curve_for(&with_gated).is_none(), + gated_curve.is_none(), "the gated quarters must not let the third bin reach the 32-quarter minimum" ); } @@ -319,8 +328,9 @@ fn the_quarter_static_gate_passes_a_mean_between_the_block_and_quarter_gates() { }; let quarters = frame_with_odd_quarters(slow); + let curve = curve_for(&quarters); assert!( - curve_for(&quarters).is_some(), + curve.is_some(), "a 2.5 code mean is static for a 64-pixel quarter" ); } @@ -344,31 +354,37 @@ fn the_rho_sigma_gate_rejects_a_noiseless_quarter() { fn quarters_of_a_rejected_block_never_count() { let build = |third_sigma: f32| { let mut quarters = quarters_at(0.15, 0.01, 40); - quarters.extend(quarters_at(0.35, 0.01, 40)); - quarters.extend(quarters_at(0.6, third_sigma, 40)); + let middle_band = quarters_at(0.35, 0.01, 40); + let bright_band = quarters_at(0.6, third_sigma, 40); + quarters.extend(middle_band); + quarters.extend(bright_band); quarters }; // Below 5 times the blocks' lower quartile, so these blocks stay. let kept = build(0.04); - assert!(curve_for(&kept).is_some()); + let kept_curve = curve_for(&kept); + assert!(kept_curve.is_some()); - // Above it, so the outlier check rejects the whole block, even - // though each quarter alone would pass the quarter gates. + // Above it, so the outlier check rejects the whole block even though each quarter alone would + // pass the quarter gates. let rejected = build(0.2); - assert!(curve_for(&rejected).is_none()); + let rejected_curve = curve_for(&rejected); + assert!(rejected_curve.is_none()); } #[test] fn the_curve_normalises_by_the_quarter_median() { let sigma = 0.01; let mut quarters = quarters_at(0.15, sigma, 40); - quarters.extend(quarters_at(0.35, sigma, 40)); - quarters.extend(quarters_at(0.6, sigma, 40)); + let middle_band = quarters_at(0.35, sigma, 40); + let bright_band = quarters_at(0.6, sigma, 40); + quarters.extend(middle_band); + quarters.extend(bright_band); with_cancelling_means(&mut quarters, 0.005); - // The 16x16 median reads above the quarters' own sigma, so - // normalising by it would put every ratio below 1. + // The 16x16 median reads above the quarters' own sigma, so normalising by it would put every + // ratio below 1. let stored_ch = 1; let channels = 1; let (width, height) = frame_dims(quarters.len()); @@ -384,11 +400,13 @@ fn the_curve_normalises_by_the_quarter_median() { #[test] fn the_flat_gate_uses_the_block_median() { - // The quarters read sigma 0.01, which would set the flat limit at - // 5.0e-5. Their blocks read about 0.0106, which sets it at 5.625e-5. + // The quarters read sigma 0.01, which would set the flat limit at 5.0e-5. Their blocks read about + // 0.0106, which sets it at 5.625e-5. let build = |flatness: f32| { let mut quarters = quarters_at(0.15, 0.01, 40); - quarters.extend(quarters_at(0.35, 0.01, 40)); + let middle_band = quarters_at(0.35, 0.01, 40); + quarters.extend(middle_band); + for _ in 0..40 { let quarter = SyntheticQuarter { flatness, @@ -402,19 +420,21 @@ fn the_flat_gate_uses_the_block_median() { }; let between = build(5.3e-5); + let between_curve = curve_for(&between); assert!( - curve_for(&between).is_some(), + between_curve.is_some(), "above the quarter limit but below the block limit" ); let textured = build(5.8e-5); - assert!(curve_for(&textured).is_none(), "above the block limit"); + let textured_curve = curve_for(&textured); + assert!(textured_curve.is_none(), "above the block limit"); } #[test] fn a_ragged_parent_reads_its_partial_quarters_with_their_own_pixel_count() { - // Each row holds a full block and a 12 pixel wide ragged one. The - // ragged block's right quarters are 4x8, so they hold 32 pixels. + // Each row holds a full block and a 12 pixel wide ragged one. The ragged block's right quarters + // are 4x8, so they hold 32 pixels. let rows = 32u32; let width = 28; let height = 16 * rows; @@ -423,8 +443,8 @@ fn a_ragged_parent_reads_its_partial_quarters_with_their_own_pixel_count() { let record_len = temporal_stats_record_len(stored_ch) as usize; let mut records = vec![0.0f32; 2 * rows as usize * record_len]; - // The full blocks fill three bins at sigma 0.01. Every other quarter - // reads 0.02, and textured or ragged ones never reach a bin. + // The full blocks fill three bins at sigma 0.01. Every other quarter reads 0.02, and textured or + // ragged ones never reach a bin. let textured = SyntheticQuarter { flatness: 1.0, ..quarter_at(0.5, 0.02) @@ -469,9 +489,8 @@ fn a_ragged_parent_reads_its_partial_quarters_with_their_own_pixel_count() { let curve = build_noise_curve(&records, stored_ch, &accepted, sample.sigma[0]).expect("three populated bins"); - // 96 quarters read 0.01 and 160 read 0.02, so the quarter median is - // 0.02 only when the 64 partial quarters read their sigma over 32 - // pixels. Over 64 they would read 0.0141, and skipped the median + // 96 quarters read 0.01 and 160 read 0.02, so the quarter median is 0.02 only when the 64 partial + // quarters read their sigma over 32 pixels. Over 64 they would read 0.0141, and skipped the median // would fall to 0.015. for ratio in curve.ratios { assert!((ratio - 0.5).abs() < 1e-5, "{:?}", curve.ratios); @@ -481,13 +500,16 @@ fn a_ragged_parent_reads_its_partial_quarters_with_their_own_pixel_count() { #[test] fn letterbox_bars_do_not_reach_the_curve() { let mut base = quarters_at(0.15, 0.02, 160); - base.extend(quarters_at(0.35, 0.01, 160)); - base.extend(quarters_at(0.6, 0.005, 160)); + let middle_band = quarters_at(0.35, 0.01, 160); + let bright_band = quarters_at(0.6, 0.005, 160); + base.extend(middle_band); + base.extend(bright_band); let base_curve = curve_for(&base).expect("three populated bins"); let mut with_bars = base.clone(); - with_bars.extend(quarters_at(0.0, 0.0, 800)); + let bars = quarters_at(0.0, 0.0, 800); + with_bars.extend(bars); let bar_curve = curve_for(&with_bars).expect("three populated bins"); @@ -496,19 +518,21 @@ fn letterbox_bars_do_not_reach_the_curve() { #[test] fn no_passing_quarter_gives_none() { - // Every block's mean cancels to 0, so each block is accepted, but - // every quarter alone fails the static gate. + // Every block's mean cancels to 0, so each block is accepted, but every quarter alone fails the + // static gate. let mut quarters = quarters_at(0.15, 0.01, 40); - quarters.extend(quarters_at(0.35, 0.01, 40)); - quarters.extend(quarters_at(0.6, 0.01, 40)); + let middle_band = quarters_at(0.35, 0.01, 40); + let bright_band = quarters_at(0.6, 0.01, 40); + quarters.extend(middle_band); + quarters.extend(bright_band); with_cancelling_means(&mut quarters, 3.0 / 255.0); - assert!(curve_for(&quarters).is_none()); + let curve = curve_for(&quarters); + assert!(curve.is_none()); } -/// A brightness ramp with static noise in each band, so a curve forms -/// with at least three populated bins. -pub(super) fn ramp_frame(width: u32, height: u32, seed: u32) -> Vec { +/// Three brightness bands with static noise, so a curve forms with at least three populated bins. +pub(super) fn banded_noisy_frame(width: u32, height: u32, seed: u32) -> Vec { let band_height = height / 3; let mut clean = vec![0.0f32; (width * height) as usize]; for y in 0..height { @@ -521,6 +545,7 @@ pub(super) fn ramp_frame(width: u32, height: u32, seed: u32) -> Vec { clean[(y * width + x) as usize] = luma; } } + noisy_field_over(&clean, width, height, 0.02, seed) } @@ -555,7 +580,7 @@ fn reset_stream_state_clears_the_curve() { let mut curve_seen = false; for i in 0..12u32 { - let frame = ramp_frame(width, height, 100 + i); + let frame = banded_noisy_frame(width, height, 100 + i); denoiser.push_frame(&frame); let _ = denoiser.denoise().unwrap(); if denoiser.current_noise_curve().is_some() { @@ -563,6 +588,7 @@ fn reset_stream_state_clears_the_curve() { break; } } + assert!(curve_seen, "expected a curve to form over the brightness ramp"); denoiser.reset_stream_state(); diff --git a/av-denoise-core/src/nlmeans/tests/options.rs b/av-denoise-core/src/nlmeans/tests/options.rs new file mode 100644 index 0000000..77276c0 --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/options.rs @@ -0,0 +1,237 @@ +use crate::nlmeans::{ + ChannelMode, + DenoisingMode, + HqParams, + MotionCompensationMode, + MotionEstimation, + NlmParams, + NlmTuning, + NlmeansAlgorithm, + NlmeansHqOptions, + NlmeansOptions, + PrefilterMode, + hq_default_strength, + resolve_params, +}; + +fn fast(options: NlmeansOptions) -> NlmeansAlgorithm { + NlmeansAlgorithm::Fast(options) +} + +fn hq_with(hq: HqParams, mode: DenoisingMode) -> NlmeansAlgorithm { + let nlm_options = NlmeansOptions { + mode, + ..NlmeansOptions::default() + }; + + NlmeansAlgorithm::Hq(NlmeansHqOptions { nlm: nlm_options, hq }) +} + +#[test] +fn spatial_mode_maps_to_zero_temporal_radius() { + let options = NlmeansOptions { + mode: DenoisingMode::Spacial, + ..NlmeansOptions::default() + }; + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert_eq!(params.temporal_radius, 0); + assert_eq!(params.channels, ChannelMode::Yuv); +} + +#[test] +fn temporal_mode_propagates_radius() { + let options = NlmeansOptions { + mode: DenoisingMode::Temporal { radius: 3 }, + ..NlmeansOptions::default() + }; + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Luma); + + assert_eq!(params.temporal_radius, 3); +} + +#[test] +fn prefilter_passthrough() { + let options = NlmeansOptions { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: 0.02, + }, + ..NlmeansOptions::default() + }; + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert!(matches!(params.prefilter, PrefilterMode::Bilateral { .. })); +} + +#[test] +fn hq_unset_prefilter_defaults_to_none() { + let hq = HqParams::default(); + let algorithm = hq_with(hq, DenoisingMode::Spacial); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert!(matches!(params.prefilter, PrefilterMode::None)); +} + +#[test] +fn fast_unset_prefilter_defaults_to_none() { + let options = NlmeansOptions::default(); + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert!(matches!(params.prefilter, PrefilterMode::None)); +} + +#[test] +fn hq_unset_strength_defaults_to_hq_default_strength() { + let hq = HqParams::default(); + let algorithm = hq_with(hq, DenoisingMode::Spacial); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + let expected = hq_default_strength(ChannelMode::Yuv, 0); + assert!((params.strength - expected).abs() < f32::EPSILON); +} + +#[test] +fn hq_no_auto_strength_falls_back_to_the_legacy_absolute_default() { + let hq = HqParams { + auto_strength: false, + ..HqParams::default() + }; + let algorithm = hq_with(hq, DenoisingMode::Spacial); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + let expected = NlmParams::default().strength; + assert!((params.strength - expected).abs() < f32::EPSILON); +} + +#[test] +fn hq_luma_r4_uses_measured_table_value() { + let mode = DenoisingMode::Temporal { radius: 4 }; + let hq = HqParams::default(); + let algorithm = hq_with(hq, mode); + let params = resolve_params(&algorithm, ChannelMode::Luma); + + assert!((params.strength - 0.35).abs() < f32::EPSILON); +} + +#[test] +fn hq_chroma_r4_uses_measured_table_value() { + let mode = DenoisingMode::Temporal { radius: 4 }; + let hq = HqParams::default(); + let algorithm = hq_with(hq, mode); + let params = resolve_params(&algorithm, ChannelMode::Chroma); + + assert!((params.strength - 0.70).abs() < f32::EPSILON); +} + +#[test] +fn hq_yuv_r8_uses_measured_table_value() { + let mode = DenoisingMode::Temporal { radius: 8 }; + let hq = HqParams::default(); + let algorithm = hq_with(hq, mode); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert!((params.strength - 0.30).abs() < f32::EPSILON); +} + +#[test] +fn hq_spacial_mode_uses_radius_zero_table_values() { + for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] { + let hq = HqParams::default(); + let algorithm = hq_with(hq, DenoisingMode::Spacial); + let params = resolve_params(&algorithm, channels); + + let expected = hq_default_strength(channels, 0); + assert!((params.strength - expected).abs() < f32::EPSILON); + } +} + +#[test] +fn hq_explicit_strength_wins_over_the_table_for_every_plane() { + for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] { + let nlm_options = NlmeansOptions { + mode: DenoisingMode::Temporal { radius: 4 }, + tuning: NlmTuning { + strength: Some(0.99), + ..NlmTuning::default() + }, + ..NlmeansOptions::default() + }; + let hq = HqParams::default(); + let algorithm = NlmeansAlgorithm::Hq(NlmeansHqOptions { nlm: nlm_options, hq }); + let params = resolve_params(&algorithm, channels); + + assert!((params.strength - 0.99).abs() < f32::EPSILON); + } +} + +#[test] +fn fast_unset_strength_defaults_to_legacy_default() { + let options = NlmeansOptions::default(); + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert!((params.strength - 1.2).abs() < f32::EPSILON); +} + +#[test] +fn motion_compensation_passthrough() { + let options = NlmeansOptions { + mode: DenoisingMode::Temporal { radius: 1 }, + motion_compensation: MotionCompensationMode::Mvtools { + blksize: 16, + overlap: 8, + search_radius: 4, + pyramid_levels: 2, + estimation: MotionEstimation::Direct, + }, + ..NlmeansOptions::default() + }; + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert!(matches!( + params.motion_compensation, + MotionCompensationMode::Mvtools { + blksize: 16, + overlap: 8, + search_radius: 4, + pyramid_levels: 2, + .. + } + )); +} + +#[test] +fn motion_compensation_defaults_to_none() { + let options = NlmeansOptions::default(); + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert!(matches!(params.motion_compensation, MotionCompensationMode::None)); +} + +#[test] +fn nlm_tuning_overrides_individual_fields() { + let defaults = NlmParams::default(); + let options = NlmeansOptions { + tuning: NlmTuning { + search_radius: Some(7), + patch_radius: None, + strength: Some(2.5), + self_weight: None, + }, + ..NlmeansOptions::default() + }; + let algorithm = fast(options); + let params = resolve_params(&algorithm, ChannelMode::Yuv); + + assert_eq!(params.search_radius, 7); + assert_eq!(params.patch_radius, defaults.patch_radius); + assert!((params.strength - 2.5).abs() < f32::EPSILON); + assert!((params.self_weight - defaults.self_weight).abs() < f32::EPSILON); +} diff --git a/av-denoise-core/src/nlmeans/tests/pack_wire.rs b/av-denoise-core/src/nlmeans/tests/pack_wire.rs deleted file mode 100644 index 873e873..0000000 --- a/av-denoise-core/src/nlmeans/tests/pack_wire.rs +++ /dev/null @@ -1,569 +0,0 @@ -use cubecl::prelude::*; - -use super::helpers::{R, make_client}; -#[cfg(feature = "vulkan")] -use super::helpers::{ramp_frame, test_denoiser}; -use crate::Depth; -use crate::nlmeans::kernels::gpu_pack_wire; -#[cfg(feature = "vulkan")] -use crate::{ - ChannelMode, - Denoiser, - DenoiserOptions, - DenoisingMode, - OutputFormat, - accelerate::Accelerator, - device::Device, -}; - -const BLOCK: u32 = 256; - -/// Runs the kernel and returns the wire bytes it produced, trimmed to the -/// plane's exact byte length. -fn pack( - src: &[f32], - pixels: u32, - channels: u32, - stored_ch: u32, - depth: Depth, - split_planes: bool, -) -> Vec { - let samples_per_word = 4 / depth.bytes_per_sample() as u32; - let words = (pixels * channels).div_ceil(samples_per_word); - let grid = words.div_ceil(BLOCK).max(1); - - pack_with_grid(src, pixels, channels, stored_ch, depth, split_planes, grid, BLOCK) -} - -/// The same, with the launch geometry chosen by the caller so a test can -/// force a grid smaller than the word count. -#[expect( - clippy::too_many_arguments, - reason = "mirrors the kernel's own argument list, plus the launch geometry" -)] -fn pack_with_grid( - src: &[f32], - pixels: u32, - channels: u32, - stored_ch: u32, - depth: Depth, - split_planes: bool, - grid: u32, - block: u32, -) -> Vec { - let client = make_client(); - - let samples = pixels * channels; - let bytes = samples as usize * depth.bytes_per_sample(); - let samples_per_word = 4 / depth.bytes_per_sample() as u32; - let words = samples.div_ceil(samples_per_word); - let outer = if split_planes { pixels } else { channels }; - - let total_threads = grid * block; - - let src_buf = client.create_from_slice(f32::as_bytes(src)); - let dst_buf = client.empty(words as usize * size_of::()); - - unsafe { - gpu_pack_wire::launch_unchecked::( - &client, - CubeCount::new_1d(grid), - CubeDim::new_1d(block), - ArrayArg::from_raw_parts(src_buf, src.len()), - ArrayArg::from_raw_parts(dst_buf.clone(), words as usize), - depth.max_value(), - pixels, - channels, - stored_ch, - outer, - split_planes, - samples_per_word, - words, - total_threads, - ); - } - - let out = client.read_one(dst_buf).expect("pack readback failed"); - out[..bytes].to_vec() -} - -/// The only check that can catch a silently-dead kernel, since a kernel -/// compared against itself compares zeros to zeros. -#[test] -fn packed_bytes_match_the_host_converter_at_every_depth() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let pixels = 64u32; - // A ramp over the full range, including both ends and values that - // land either side of a rounding boundary. - let src: Vec = (0..pixels).map(|i| i as f32 / (pixels - 1) as f32).collect(); - - let got = pack(&src, pixels, 1, 1, depth, false); - let want = crate::frame::f32_to_plane(&src, depth); - - assert_eq!( - got, want, - "depth {depth:?} must match the host converter byte for byte" - ); - } -} - -#[test] -fn padding_lanes_are_skipped() { - let pixels = 16u32; - let channels = 3u32; - let stored_ch = 4u32; - - // Every padding lane holds 1.0, which would be visible as 0xFF if it - // ever reached the output. - let mut src = vec![1.0f32; (pixels * stored_ch) as usize]; - let wanted: Vec = (0..pixels * channels).map(|i| i as f32 / 255.0).collect(); - for p in 0..pixels as usize { - for c in 0..channels as usize { - src[p * stored_ch as usize + c] = wanted[p * channels as usize + c]; - } - } - - let got = pack(&src, pixels, channels, stored_ch, Depth::Eight, false); - let want = crate::frame::f32_to_plane(&wanted, Depth::Eight); - - assert_eq!(got, want); -} - -#[test] -fn a_sample_count_that_is_not_a_whole_number_of_words_writes_its_tail() { - // 13 samples at 8-bit is three whole words plus one byte. - let pixels = 13u32; - let src: Vec = (0..pixels).map(|i| i as f32 / (pixels - 1) as f32).collect(); - - let got = pack(&src, pixels, 1, 1, Depth::Eight, false); - let want = crate::frame::f32_to_plane(&src, Depth::Eight); - - assert_eq!(got.len(), 13); - assert_eq!(got, want); -} - -#[test] -fn split_planes_writes_each_channel_as_one_region() { - let pixels = 32u32; - let channels = 2u32; - - let src: Vec = (0..pixels * channels).map(|i| i as f32 / 255.0).collect(); - - let got = pack(&src, pixels, channels, channels, Depth::Eight, true); - - let (u, v) = crate::frame::unpack_uv_from_f32(&src, pixels as usize, Depth::Eight); - let want: Vec = u.into_iter().chain(v).collect(); - - assert_eq!(got, want); -} - -/// Forces a grid far smaller than the word count so every thread runs the -/// strided loop several times. Without this the loop body runs at most -/// once and `word += total_threads` is never exercised. -#[test] -fn the_strided_loop_covers_every_word_when_the_grid_is_smaller_than_the_frame() { - let pixels = 512u32; - let src: Vec = (0..pixels).map(|i| (i % 256) as f32 / 255.0).collect(); - - // 128 words against 64 threads, so each thread walks two words. - let got = pack_with_grid(&src, pixels, 1, 1, Depth::Eight, false, 1, 64); - let want = crate::frame::f32_to_plane(&src, Depth::Eight); - - assert_eq!(got, want); -} - -#[test] -fn split_planes_at_ten_bit_writes_each_channel_as_one_region() { - let pixels = 32u32; - let channels = 2u32; - - let src: Vec = (0..pixels * channels) - .map(|i| i as f32 / (pixels * channels - 1) as f32) - .collect(); - - let got = pack(&src, pixels, channels, channels, Depth::Ten, true); - - let (u, v) = crate::frame::unpack_uv_from_f32(&src, pixels as usize, Depth::Ten); - let want: Vec = u.into_iter().chain(v).collect(); - - assert_eq!(got, want); -} - -#[test] -fn padding_lanes_are_skipped_at_ten_bit() { - let pixels = 16u32; - let channels = 3u32; - let stored_ch = 4u32; - - let mut src = vec![1.0f32; (pixels * stored_ch) as usize]; - let wanted: Vec = (0..pixels * channels) - .map(|i| i as f32 / (pixels * channels - 1) as f32) - .collect(); - for p in 0..pixels as usize { - for c in 0..channels as usize { - src[p * stored_ch as usize + c] = wanted[p * channels as usize + c]; - } - } - - let got = pack(&src, pixels, channels, stored_ch, Depth::Ten, false); - let want = crate::frame::f32_to_plane(&wanted, Depth::Ten); - - assert_eq!(got, want); -} - -/// Builds a top-level [`Denoiser`] at temporal radius 1 over `mode`, -/// collecting frames in `format`. -#[cfg(feature = "vulkan")] -pub(super) fn denoiser(mode: ChannelMode, format: OutputFormat, w: u32, h: u32) -> Denoiser { - let opts = DenoiserOptions::builder() - .channel_mode(mode) - .mode(DenoisingMode::Temporal { radius: 1 }) - .output_format(format) - .build(); - - Denoiser::create(&[Accelerator::Vulkan], &Device::Default, w, h, opts) - .expect("denoiser construction failed") -} - -/// Pushes `count` deterministic frames into `d`, one per index. -#[cfg(feature = "vulkan")] -pub(super) fn push_ramp(d: &mut Denoiser, w: u32, h: u32, count: usize) { - for i in 0..count { - d.push_frame(&ramp_frame(w, h, i)).expect("push failed"); - } -} - -/// The differential test the pack kernel lives or dies on. A kernel that -/// silently compiled to nothing returns zeros, which `f32_to_plane` of a -/// real frame never does. -#[cfg(feature = "vulkan")] -#[test] -fn wire_mode_output_matches_the_f32_path() { - let (w, h) = (16u32, 16u32); - - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let mut f32_side = test_denoiser(1, w, h); - let mut wire_side = denoiser(ChannelMode::Luma, OutputFormat::Wire { depth }, w, h); - - push_ramp(&mut f32_side, w, h, 3); - push_ramp(&mut wire_side, w, h, 3); - - let want = f32_side - .recv_frame() - .expect("f32 recv failed") - .expect("a frame is ready") - .into_f32() - .expect("an f32 denoiser returns f32"); - - let got = wire_side - .recv_frame() - .expect("wire recv failed") - .expect("a frame is ready") - .into_wire() - .expect("a wire denoiser returns wire bytes"); - - assert_eq!(got, crate::frame::f32_to_plane(&want, depth), "depth {depth:?}"); - } -} - -/// A chroma pair goes out as U's whole region followed by V's, not -/// interleaved, so it matches what a planar consumer writes. -#[cfg(feature = "vulkan")] -#[test] -fn wire_mode_chroma_matches_unpack_uv_from_f32() { - let (w, h) = (16u32, 16u32); - let depth = Depth::Ten; - - let mut f32_side = denoiser(ChannelMode::Chroma, OutputFormat::F32, w, h); - let mut wire_side = denoiser(ChannelMode::Chroma, OutputFormat::Wire { depth }, w, h); - - // A chroma frame holds two channels per pixel, which `ramp_frame` - // covers by producing twice as many values. - push_ramp(&mut f32_side, w, h * 2, 3); - push_ramp(&mut wire_side, w, h * 2, 3); - - let want = f32_side - .recv_frame() - .expect("f32 recv failed") - .expect("a frame is ready") - .into_f32() - .expect("an f32 denoiser returns f32"); - - let got = wire_side - .recv_frame() - .expect("wire recv failed") - .expect("a frame is ready") - .into_wire() - .expect("a wire denoiser returns wire bytes"); - - let (u, v) = crate::frame::unpack_uv_from_f32(&want, (w * h) as usize, depth); - let expected: Vec = u.into_iter().chain(v).collect(); - - assert_eq!(got, expected); -} - -/// 13x3 luma is 39 samples, which is nine 8-bit words plus three bytes, -/// so the kernel's last word is a partial one. -#[cfg(feature = "vulkan")] -#[test] -fn wire_mode_handles_a_plane_that_is_not_a_whole_number_of_words() { - let (w, h) = (13u32, 3u32); - let depth = Depth::Eight; - - let mut f32_side = test_denoiser(1, w, h); - let mut wire_side = denoiser(ChannelMode::Luma, OutputFormat::Wire { depth }, w, h); - - push_ramp(&mut f32_side, w, h, 3); - push_ramp(&mut wire_side, w, h, 3); - - let want = f32_side - .recv_frame() - .expect("f32 recv failed") - .expect("a frame is ready") - .into_f32() - .expect("an f32 denoiser returns f32"); - - let got = wire_side - .recv_frame() - .expect("wire recv failed") - .expect("a frame is ready") - .into_wire() - .expect("a wire denoiser returns wire bytes"); - - assert_eq!(got.len(), 39); - assert_eq!(got, crate::frame::f32_to_plane(&want, depth)); -} - -/// The flush tail must come back quantised by the same kernel as every -/// other frame, not by a host converter that could drift from it. -#[cfg(feature = "vulkan")] -#[test] -fn a_wire_flush_matches_an_f32_flush_at_every_depth() { - let (w, h) = (16u32, 16u32); - - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let mut f32_side = denoiser(ChannelMode::Luma, OutputFormat::F32, w, h); - let mut wire_side = denoiser(ChannelMode::Luma, OutputFormat::Wire { depth }, w, h); - - push_ramp(&mut f32_side, w, h, 3); - push_ramp(&mut wire_side, w, h, 3); - - let mut want = Vec::new(); - f32_side - .flush(|out| { - want.push(crate::frame::f32_to_plane( - out.as_f32().expect("f32 denoiser flushes f32"), - depth, - )) - }) - .expect("f32 flush failed"); - - let mut got = Vec::new(); - wire_side - .flush(|out| got.push(out.as_wire().expect("wire denoiser flushes wire").to_vec())) - .expect("wire flush failed"); - - assert!(!got.is_empty(), "flush emitted nothing at depth {depth:?}"); - assert_eq!(got, want, "depth {depth:?}"); - } -} - -/// A flushed chroma pair splits into U's whole region followed by V's, -/// the same way a streaming one does. -#[cfg(feature = "vulkan")] -#[test] -fn a_wire_chroma_flush_matches_unpack_uv_from_f32() { - let (w, h) = (16u32, 16u32); - let depth = Depth::Ten; - - let mut f32_side = denoiser(ChannelMode::Chroma, OutputFormat::F32, w, h); - let mut wire_side = denoiser(ChannelMode::Chroma, OutputFormat::Wire { depth }, w, h); - - // A chroma frame holds two channels per pixel, which `ramp_frame` - // covers by producing twice as many values. - push_ramp(&mut f32_side, w, h * 2, 3); - push_ramp(&mut wire_side, w, h * 2, 3); - - let mut want = Vec::new(); - f32_side - .flush(|out| { - let (u, v) = crate::frame::unpack_uv_from_f32( - out.as_f32().expect("f32 denoiser flushes f32"), - (w * h) as usize, - depth, - ); - want.push(u.into_iter().chain(v).collect::>()); - }) - .expect("f32 flush failed"); - - let mut got = Vec::new(); - wire_side - .flush(|out| got.push(out.as_wire().expect("wire denoiser flushes wire").to_vec())) - .expect("wire flush failed"); - - assert!(!got.is_empty(), "flush emitted nothing"); - assert_eq!(got, want); -} - -/// A flush takes the same wire slots streaming does, so one that begins -/// while a submitted readback is still in flight would be handed a slot -/// that readback is still reading. -/// -/// The push after the single drain leaves a readback outstanding at the -/// moment `flush` is called, which draining to empty first would hide. -#[cfg(feature = "vulkan")] -#[test] -fn flushing_with_a_readback_in_flight_matches_the_f32_path() { - let (w, h) = (16u32, 16u32); - let depth = Depth::Ten; - - let mut f32_side = denoiser(ChannelMode::Luma, OutputFormat::F32, w, h); - let mut wire_side = denoiser(ChannelMode::Luma, OutputFormat::Wire { depth }, w, h); - - let mut want = Vec::new(); - let mut got = Vec::new(); - - push_ramp(&mut f32_side, w, h, 3); - push_ramp(&mut wire_side, w, h, 3); - - want.push(crate::frame::f32_to_plane( - &f32_side - .recv_frame() - .expect("f32 recv failed") - .expect("a frame is ready") - .into_f32() - .expect("an f32 denoiser returns f32"), - depth, - )); - got.push( - wire_side - .recv_frame() - .expect("wire recv failed") - .expect("a frame is ready") - .into_wire() - .expect("a wire denoiser returns wire bytes"), - ); - - let frame = ramp_frame(w, h, 3); - f32_side.push_frame(&frame).expect("f32 push failed"); - wire_side.push_frame(&frame).expect("wire push failed"); - - f32_side - .flush(|out| { - want.push(crate::frame::f32_to_plane( - out.as_f32().expect("f32 denoiser flushes f32"), - depth, - )) - }) - .expect("f32 flush failed"); - wire_side - .flush(|out| got.push(out.as_wire().expect("wire denoiser flushes wire").to_vec())) - .expect("wire flush failed"); - - assert_eq!(got.len(), 4, "four pushes at radius 1 emit four frames"); - assert_eq!(got, want); -} - -/// A wire buffer handed out while its own readback is still reading it -/// would let one frame's bytes appear in another's. -/// -/// Both pushes land before either drain, so two readbacks are in flight -/// at once and the second occupies the slot the first has not released. -/// Draining after every push would leave one readback live at a time, -/// where a wrong slot index cannot be observed at all. -#[cfg(feature = "vulkan")] -#[test] -fn reusing_the_wire_slots_keeps_every_frame_distinct() { - let (w, h) = (16u32, 16u32); - let depth = Depth::Eight; - - let mut f32_side = denoiser(ChannelMode::Luma, OutputFormat::F32, w, h); - let mut wire_side = denoiser(ChannelMode::Luma, OutputFormat::Wire { depth }, w, h); - - let mut got = Vec::new(); - let mut want = Vec::new(); - - for pair in 0..3 { - for k in 0..2 { - let frame = ramp_frame(w, h, pair * 2 + k); - f32_side.push_frame(&frame).expect("f32 push failed"); - wire_side.push_frame(&frame).expect("wire push failed"); - } - - while let Some(out) = f32_side.recv_frame().expect("f32 recv failed") { - want.push(crate::frame::f32_to_plane( - &out.into_f32().expect("f32 denoiser returns f32"), - depth, - )); - } - - while let Some(out) = wire_side.recv_frame().expect("wire recv failed") { - got.push(out.into_wire().expect("wire denoiser returns wire bytes")); - } - } - - assert_eq!(got.len(), 5, "six pushes at radius 1 emit five frames"); - assert_eq!(got, want); -} - -#[cfg(feature = "vulkan")] -#[test] -fn an_f32_denoiser_allocates_no_wire_buffers() { - let d = denoiser(ChannelMode::Luma, OutputFormat::F32, 16, 16); - assert!(d.wire_outputs_for_test().is_none()); -} - -#[cfg(feature = "vulkan")] -#[test] -fn try_recv_frame_in_wire_mode_returns_none_when_nothing_is_in_flight() { - let mut d = denoiser( - ChannelMode::Luma, - OutputFormat::Wire { depth: Depth::Eight }, - 64, - 64, - ); - assert_eq!(d.try_recv_frame().unwrap(), None); -} - -/// Wire mode must not quietly become blocking, and the bytes it polls -/// out must be the ones the blocking path returns. -#[cfg(feature = "vulkan")] -#[test] -fn try_wait_still_reports_not_ready_without_blocking_in_wire_mode() { - // A poll count is the wrong proxy for the wall-clock interval this - // test needs to cover (cold pipeline compile plus dispatch plus - // readback), since a faster CPU makes each poll cheaper and so - // needs *more* of them for the same GPU latency. A deadline covers - // both a slow GPU and a fast CPU the same way. - const DEADLINE: std::time::Duration = std::time::Duration::from_secs(30); - - let (w, h) = (64u32, 64u32); - let format = OutputFormat::Wire { depth: Depth::Eight }; - - // Two pushes at radius 1 prime the window and submit one denoise, - // leaving exactly one readback in flight. - let mut polled = denoiser(ChannelMode::Luma, format, w, h); - push_ramp(&mut polled, w, h, 2); - - let start = std::time::Instant::now(); - let mut got = None; - let mut polls = 0; - while start.elapsed() < DEADLINE { - polls += 1; - if let Some(frame) = polled.try_recv_frame().unwrap() { - got = Some(frame); - break; - } - } - let got = got.unwrap_or_else(|| panic!("readback never landed within {DEADLINE:?} ({polls} polls)")); - - let mut blocking = denoiser(ChannelMode::Luma, format, w, h); - push_ramp(&mut blocking, w, h, 2); - let expected = blocking - .recv_frame() - .unwrap() - .expect("blocking denoiser should have a frame ready"); - - assert!(matches!(got, crate::FrameOutput::Wire(_))); - assert_eq!(got, expected); -} diff --git a/av-denoise-core/src/nlmeans/tests/params.rs b/av-denoise-core/src/nlmeans/tests/params.rs new file mode 100644 index 0000000..c0c3acc --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/params.rs @@ -0,0 +1,571 @@ +use crate::nlmeans::params::{ + ChannelMode, + HqParams, + MAX_TEMPORAL_RADIUS, + MIN_FRAME_DIM, + NLM_LEGACY, + NLM_NORM, + NlmParams, + SEPARABLE_THRESHOLD, + hq_default_strength, + sigma_eff, + validate_dimensions, +}; +use crate::nlmeans::prefilter::{self, PrefilterMode}; + +#[test] +fn noise_offset_scales_with_sigma_and_patch_size() { + let sigma = 4.0 / 255.0; + let params = NlmParams { + patch_radius: 4, + hq: Some(HqParams::with_sigma(sigma)), + ..NlmParams::default() + }; + + let expected = 6.0 * sigma * sigma * 81.0; + let got = params.noise_offset(); + assert!((got - expected).abs() < 1e-6, "expected {expected}, got {got}"); +} + +#[test] +fn noise_offset_zero_without_noise_floor() { + let params = NlmParams { + hq: Some(HqParams { + auto_strength: true, + noise_floor: false, + sigma_override: Some(4.0 / 255.0), + temporal_confidence: true, + thsad_scale: 1.0, + sigma_scale: 1.0, + windowed_noise_estimation: false, + }), + ..NlmParams::default() + }; + + assert_eq!(params.noise_offset(), 0.0); +} + +#[test] +fn noise_offset_zero_without_hq() { + let params = NlmParams::default(); + assert_eq!(params.noise_offset(), 0.0); +} + +#[test] +fn h2_inv_norm_with_auto_strength_matches_hand_computed() { + let sigma = 8.0 / 255.0; + let params = NlmParams { + strength: 1.0, + hq: Some(HqParams::with_sigma(sigma)), + ..NlmParams::default() + }; + + let patch_area = (2 * params.patch_radius + 1) * (2 * params.patch_radius + 1); + let effective_strength = 1.0 * sigma * 255.0; + let expected = NLM_NORM / (NLM_LEGACY * effective_strength * effective_strength * patch_area as f32); + + let got = params.h2_inv_norm(); + assert!((got - expected).abs() < 1e-6, "expected {expected}, got {got}"); +} + +#[test] +fn validate_rejects_zero_hq_sigma() { + let params = NlmParams { + hq: Some(HqParams::with_sigma(0.0)), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_rejects_hq_sigma_above_one() { + let params = NlmParams { + hq: Some(HqParams::with_sigma(1.5)), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_rejects_nan_hq_sigma() { + let params = NlmParams { + hq: Some(HqParams::with_sigma(f32::NAN)), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_rejects_zero_thsad_scale() { + let params = NlmParams { + hq: Some(HqParams { + thsad_scale: 0.0, + ..HqParams::default() + }), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_rejects_negative_thsad_scale() { + let params = NlmParams { + hq: Some(HqParams { + thsad_scale: -1.0, + ..HqParams::default() + }), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_rejects_nan_thsad_scale() { + let params = NlmParams { + hq: Some(HqParams { + thsad_scale: f32::NAN, + ..HqParams::default() + }), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_accepts_default_thsad_scale() { + let params = NlmParams { + hq: Some(HqParams::default()), + ..NlmParams::default() + }; + assert!(params.validate().is_ok()); +} + +#[test] +fn hq_params_default_sigma_scale_is_one() { + let defaults = HqParams::default(); + assert_eq!(defaults.sigma_scale, 1.0); +} + +#[test] +fn validate_rejects_sigma_scale_below_the_minimum() { + let params = NlmParams { + hq: Some(HqParams { + sigma_scale: 0.05, + ..HqParams::default() + }), + ..NlmParams::default() + }; + let err = params.validate().expect_err("0.05 is below the 0.1 minimum"); + assert!( + err.to_string().contains("hq sigma_scale"), + "error should name the field, got {err}" + ); +} + +#[test] +fn validate_rejects_sigma_scale_above_the_maximum() { + let params = NlmParams { + hq: Some(HqParams { + sigma_scale: 10.5, + ..HqParams::default() + }), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_rejects_nan_sigma_scale() { + let params = NlmParams { + hq: Some(HqParams { + sigma_scale: f32::NAN, + ..HqParams::default() + }), + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_accepts_sigma_scale_at_the_bounds() { + let low = NlmParams { + hq: Some(HqParams { + sigma_scale: 0.1, + ..HqParams::default() + }), + ..NlmParams::default() + }; + assert!(low.validate().is_ok()); + + let high = NlmParams { + hq: Some(HqParams { + sigma_scale: 10.0, + ..HqParams::default() + }), + ..NlmParams::default() + }; + assert!(high.validate().is_ok()); +} + +#[test] +fn noise_offset_with_handles_distinct_per_channel_sigmas() { + let sigma_u = 4.0 / 255.0; + let sigma_v = 10.0 / 255.0; + let params = NlmParams { + patch_radius: 4, + channels: ChannelMode::Chroma, + hq: Some(HqParams { + auto_strength: true, + noise_floor: true, + sigma_override: None, + temporal_confidence: true, + thsad_scale: 1.0, + sigma_scale: 1.0, + windowed_noise_estimation: false, + }), + ..NlmParams::default() + }; + + let patch_area = (2 * params.patch_radius + 1) * (2 * params.patch_radius + 1); + // The chroma scale of 1.5 applies per channel, and each channel keeps its own sigma. + let expected = 2.0 * 1.5 * (sigma_u * sigma_u + sigma_v * sigma_v) * patch_area as f32; + + let got = params.noise_offset_with(Some(&[sigma_u, sigma_v])); + assert!((got - expected).abs() < 1e-9, "expected {expected}, got {got}"); +} + +#[test] +fn sigma_eff_is_rms_over_active_channels() { + let sigmas = [3.0 / 255.0, 4.0 / 255.0]; + let got = sigma_eff(&sigmas, ChannelMode::Chroma); + let expected = ((sigmas[0] * sigmas[0] + sigmas[1] * sigmas[1]) / 2.0).sqrt(); + assert!((got - expected).abs() < 1e-9, "expected {expected}, got {got}"); +} + +#[test] +fn validate_rejects_non_positive_pilot_strength_scale() { + let zero = NlmParams { + prefilter: PrefilterMode::NlmSpatial { strength_scale: 0.0 }, + ..NlmParams::default() + }; + assert!(zero.validate().is_err()); + + let nan = NlmParams { + prefilter: PrefilterMode::NlmSpatial { + strength_scale: f32::NAN, + }, + ..NlmParams::default() + }; + assert!(nan.validate().is_err()); +} + +#[test] +fn validate_rejects_pilot_with_patch_radius_above_separable_threshold() { + let params = NlmParams { + prefilter: PrefilterMode::NlmSpatial { strength_scale: 1.0 }, + patch_radius: SEPARABLE_THRESHOLD + 1, + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn validate_accepts_pilot_within_limits() { + let params = NlmParams { + prefilter: PrefilterMode::NlmSpatial { strength_scale: 1.0 }, + patch_radius: SEPARABLE_THRESHOLD, + ..NlmParams::default() + }; + assert!(params.validate().is_ok()); +} + +#[test] +fn validate_rejects_non_positive_bilateral_sigma_r() { + // The centre tap's range distance is 0, and 0 times an infinite factor is NaN, which poisons + // every pixel of the reference image. + let params = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: 0.0, + }, + ..NlmParams::default() + }; + assert!(params.validate().is_err()); + + let negative = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: -0.02, + }, + ..NlmParams::default() + }; + assert!(negative.validate().is_err()); + + let nan = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: f32::NAN, + }, + ..NlmParams::default() + }; + assert!(nan.validate().is_err()); + + let infinite = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: f32::INFINITY, + }, + ..NlmParams::default() + }; + assert!(infinite.validate().is_err()); +} + +#[test] +fn validate_rejects_non_positive_bilateral_sigma_s() { + // The centre tap's spatial distance is 0, so an infinite factor poisons it with the same NaN. + let params = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 0.0, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(params.validate().is_err()); + + let negative = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: -3.0, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(negative.validate().is_err()); + + let nan = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: f32::NAN, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(nan.validate().is_err()); + + let infinite = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: f32::INFINITY, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(infinite.validate().is_err()); +} + +#[test] +fn validate_accepts_positive_finite_bilateral_sigmas() { + let params = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(params.validate().is_ok()); +} + +/// Pins the guard to `<= 0.0`, with a value whose square stays a normal float. +#[test] +fn validate_accepts_a_small_positive_bilateral_sigma_at_the_boundary() { + let safe_small = 1e-6_f32; + assert!( + (safe_small * safe_small).is_normal(), + "the test value itself must not underflow" + ); + + let small_sigma_s = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: safe_small, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(small_sigma_s.validate().is_ok()); + + let small_sigma_r = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: safe_small, + }, + ..NlmParams::default() + }; + assert!(small_sigma_r.validate().is_ok()); +} + +#[test] +fn validate_rejects_a_subnormal_bilateral_sigma_that_underflows_on_squaring() { + // `f32::MIN_POSITIVE` is finite and above 0, but its square underflows to exactly 0.0, which + // makes the derived factor infinite. + let squared = f32::MIN_POSITIVE * f32::MIN_POSITIVE; + assert_eq!( + squared, 0.0, + "this test assumes MIN_POSITIVE underflows on squaring" + ); + let inverse = prefilter::inv_two_sigma_sq(f32::MIN_POSITIVE); + assert!( + !inverse.is_finite(), + "this test assumes the derived factor is infinite here" + ); + + let sigma_s = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: f32::MIN_POSITIVE, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!( + sigma_s.validate().is_err(), + "a subnormal sigma_s that underflows to an infinite normalisation factor must be rejected" + ); + + let sigma_r = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 3.0, + sigma_r: f32::MIN_POSITIVE, + }, + ..NlmParams::default() + }; + assert!( + sigma_r.validate().is_err(), + "a subnormal sigma_r that underflows to an infinite normalisation factor must be rejected" + ); +} + +#[test] +fn validate_rejects_bilateral_sigma_s_above_the_smem_ceiling() { + let params = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 16.0, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + let err = params.validate().expect_err("radius 32 exceeds the 22 ceiling"); + assert!( + err.to_string().contains("sigma_s"), + "error should name the field, got {err}" + ); +} + +/// A `sigma_s` of 1e9 would overflow the tile-size arithmetic at launch. +#[test] +fn validate_rejects_extreme_bilateral_sigma_s() { + let params = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 1e9, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +/// A `sigma_s` of 11.0 gives a radius of 22, exactly the ceiling. +#[test] +fn validate_accepts_bilateral_sigma_s_at_the_smem_ceiling() { + let params = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 11.0, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(params.validate().is_ok()); +} + +/// A `sigma_s` of 11.01 gives a radius of 23, one past the ceiling. +#[test] +fn validate_rejects_bilateral_sigma_s_just_above_the_smem_ceiling() { + let params = NlmParams { + prefilter: PrefilterMode::Bilateral { + sigma_s: 11.01, + sigma_r: 0.02, + }, + ..NlmParams::default() + }; + assert!(params.validate().is_err()); +} + +#[test] +fn sigma_eff_ignores_channels_past_the_mode_count() { + let sigmas = [6.0 / 255.0, 100.0 / 255.0, 200.0 / 255.0]; + let got = sigma_eff(&sigmas, ChannelMode::Luma); + assert!( + (got - sigmas[0]).abs() < 1e-9, + "expected {}, got {got}", + sigmas[0] + ); +} + +#[test] +fn hq_default_strength_matches_the_measured_luma_table() { + const EXPECTED: [f32; 9] = [0.45, 0.45, 0.42, 0.42, 0.35, 0.35, 0.35, 0.30, 0.30]; + for (radius, &expected) in EXPECTED.iter().enumerate() { + let got = hq_default_strength(ChannelMode::Luma, radius as u32); + assert!( + (got - expected).abs() < f32::EPSILON, + "at radius {radius} expected {expected}, got {got}" + ); + } +} + +#[test] +fn hq_default_strength_matches_the_measured_chroma_table() { + const EXPECTED: [f32; 9] = [1.00, 0.85, 0.70, 0.70, 0.70, 0.70, 0.70, 0.70, 0.70]; + for (radius, &expected) in EXPECTED.iter().enumerate() { + let got = hq_default_strength(ChannelMode::Chroma, radius as u32); + assert!( + (got - expected).abs() < f32::EPSILON, + "at radius {radius} expected {expected}, got {got}" + ); + } +} + +#[test] +fn hq_default_strength_yuv_reads_the_luma_table() { + for radius in 0..=8u32 { + let yuv = hq_default_strength(ChannelMode::Yuv, radius); + let luma = hq_default_strength(ChannelMode::Luma, radius); + assert!( + (yuv - luma).abs() < f32::EPSILON, + "at radius {radius} yuv is {yuv} but luma is {luma}" + ); + } +} + +#[test] +fn validate_dimensions_rejects_frames_below_the_minimum() { + let small_width = validate_dimensions(2, 64); + let small_height = validate_dimensions(64, 2); + let empty = validate_dimensions(0, 0); + assert!(small_width.is_err()); + assert!(small_height.is_err()); + assert!(empty.is_err()); +} + +#[test] +fn validate_dimensions_accepts_the_minimum() { + let minimum = validate_dimensions(MIN_FRAME_DIM, MIN_FRAME_DIM); + let full_hd = validate_dimensions(1920, 1080); + assert!(minimum.is_ok()); + assert!(full_hd.is_ok()); +} + +#[test] +fn hq_default_strength_clamps_radius_above_the_table() { + let at_max = hq_default_strength(ChannelMode::Luma, MAX_TEMPORAL_RADIUS); + let above_max = hq_default_strength(ChannelMode::Luma, MAX_TEMPORAL_RADIUS + 5); + assert!( + (at_max - above_max).abs() < f32::EPSILON, + "expected clamping to hold the last table entry, got {at_max} vs {above_max}" + ); +} diff --git a/av-denoise-core/src/nlmeans/tests/pending_drop.rs b/av-denoise-core/src/nlmeans/tests/pending_drop.rs deleted file mode 100644 index 0e679ce..0000000 --- a/av-denoise-core/src/nlmeans/tests/pending_drop.rs +++ /dev/null @@ -1,64 +0,0 @@ -//! Dropping a `Pending` that was polled but has not landed must settle its readback -//! rather than abandon it. -//! -//! On the wgpu backends the first poll maps a staging buffer that only the finished -//! readback unmaps. -//! -//! A buffer handed back to the device's staging pool while still mapped makes the next -//! submit that touches it fail on cubecl's device thread, and every later call on that -//! device then errors. - -use super::helpers::*; -use crate::nlmeans::*; - -/// Large enough that the GPU cannot have finished by the time the -/// first poll runs, a few microseconds after submit. The readback is -/// large too, so its staging page is not one a smaller read would pick. -/// The radii stay small because a wide kernel is what costs codegen -/// stack, and this test is about the readback, not the kernel. -const SIZE: u32 = 2048; - -fn params() -> NlmParams { - NlmParams { - temporal_radius: 0, - search_radius: 3, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::None, - hq: None, - } -} - -fn submit(client: &cubecl::client::ComputeClient, frame: &[f32]) -> Pending { - let mut denoiser = NlmDenoiser::::new(client, params(), SIZE, SIZE); - denoiser.push_frame(frame); - denoiser - .denoise_submit() - .expect("submit failed") - .expect("spatial mode submits one readback per push") -} - -#[test] -fn dropping_a_polled_pending_settles_its_readback() { - let client = make_client(); - let frame = make_uniform_frame(SIZE, SIZE, 1, 0.5); - - let pending = submit(&client, &frame); - let not_ready = match pending.try_wait().expect("poll failed") { - TryWait::NotReady(pending) => pending, - TryWait::Ready(_) => { - panic!("the first poll landed, so the drop path cannot be exercised at this size") - }, - }; - drop(not_ready); - - let out = submit(&client, &frame) - .wait() - .expect("a readback after a dropped polled Pending must still work") - .into_f32() - .expect("f32 output"); - assert_eq!(out.len(), (SIZE * SIZE) as usize); -} diff --git a/av-denoise-core/src/nlmeans/tests/pending_outlives.rs b/av-denoise-core/src/nlmeans/tests/pending_outlives.rs deleted file mode 100644 index 4548e3e..0000000 --- a/av-denoise-core/src/nlmeans/tests/pending_outlives.rs +++ /dev/null @@ -1,54 +0,0 @@ -//! A `Pending` has to keep working after the denoiser that produced it -//! is dropped. -//! -//! The readback future is built inside an `async move` that owns a -//! cloned `ComputeClient`, which is what makes that safe. This test pins -//! that arrangement down. -//! -//! If a future cubecl release changes `read_async` so its future borrows -//! from the client, or if a refactor here brings back a -//! lifetime-erasing `transmute`, this test reaches freed memory. Miri -//! and ASan will catch it, and a plain build may fail too. - -use super::helpers::*; -use crate::nlmeans::*; - -#[test] -fn pending_survives_denoiser_drop() { - let client = make_client(); - let params = NlmParams { - temporal_radius: 0, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::None, - hq: None, - }; - - let w = 16; - let h = 16; - let frame = make_uniform_frame(w, h, 1, 0.5); - - let pending = { - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); - denoiser.push_frame(&frame); - denoiser - .denoise_submit() - .expect("submit failed") - .expect("expected a pending readback for spatial mode") - // `denoiser` is dropped here, and `pending` has to stay valid. - }; - - let out = pending - .wait() - .expect("wait failed") - .into_f32() - .expect("f32 output"); - assert_eq!(out.len(), (w * h) as usize); - for (i, &v) in out.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-5, "pixel {i}: expected 0.5, got {v}"); - } -} diff --git a/av-denoise-core/src/nlmeans/tests/prefilter.rs b/av-denoise-core/src/nlmeans/tests/prefilter.rs index 16647ed..912195b 100644 --- a/av-denoise-core/src/nlmeans/tests/prefilter.rs +++ b/av-denoise-core/src/nlmeans/tests/prefilter.rs @@ -1,11 +1,12 @@ use cubecl::prelude::*; use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; -/// Reads back a single-slot reference ring buffer (`temporal_radius: 0`, -/// so the ring holds exactly one frame and the whole handle is that -/// frame, no byte-offset slicing needed). +/// Reads back the reference ring, which holds exactly one frame at `temporal_radius: 0`. +/// +/// The single slot is the whole buffer, so the handle needs no byte-offset slicing. fn read_single_slot_reference(denoiser: &NlmDenoiser) -> Vec { let handle = denoiser .reference_buf @@ -19,260 +20,38 @@ fn read_single_slot_reference(denoiser: &NlmDenoiser) -> Vec { f32::from_bytes(&bytes).to_vec() } -/// Mean absolute horizontal + vertical neighbour difference, a simple -/// roughness proxy for single-channel dense (`stored_ch == 1`) frames. -/// Lower means smoother. -fn mean_abs_neighbour_diff(frame: &[f32], w: u32, h: u32) -> f32 { - let w = w as usize; - let h = h as usize; +/// Mean absolute difference to the right and lower neighbours of a single-channel frame. +/// +/// It is a simple roughness proxy, where lower means smoother. +fn mean_abs_neighbour_diff(frame: &[f32], width: u32, height: u32) -> f32 { + let width = width as usize; + let height = height as usize; let mut sum = 0.0f32; let mut count = 0usize; - for y in 0..h { - for x in 0..w { - let v = frame[y * w + x]; - if x + 1 < w { - sum += (frame[y * w + x + 1] - v).abs(); + for y in 0..height { + for x in 0..width { + let centre = frame[y * width + x]; + if x + 1 < width { + sum += (frame[y * width + x + 1] - centre).abs(); count += 1; } - if y + 1 < h { - sum += (frame[(y + 1) * w + x] - v).abs(); + + if y + 1 < height { + sum += (frame[(y + 1) * width + x] - centre).abs(); count += 1; } } } - sum / count as f32 -} - -#[test] -fn external_reference_equals_input_matches_baseline() { - let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.3, 8, 8, 2, 0.7); - - let baseline = { - let params = NlmParams { - temporal_radius: 0, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::None, - hq: None, - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() - }; - - let with_ref = { - let params = NlmParams { - temporal_radius: 0, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::External, - motion_compensation: MotionCompensationMode::None, - hq: None, - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame_with_reference(&frame, &frame); - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() - }; - - assert_eq!(baseline.len(), with_ref.len()); - for (i, (a, b)) in baseline.iter().zip(with_ref.iter()).enumerate() { - assert!((a - b).abs() < 1e-5, "pixel {i}: baseline={a}, with_ref={b}"); - } -} - -/// `push_frame` seeds the noise estimator from the stream's first -/// frame so any push-time GPU work that reads σ never sees the -/// absolute-strength fallback (`seed_noise_estimate_if_first_frame`'s -/// doc says this applies for every `temporal_radius`). -/// `push_frame_with_reference` queues the same estimate but must reach -/// the same seed, not leave the estimator unset until the first -/// `denoise_submit` runs. -#[test] -fn push_frame_with_reference_seeds_noise_estimate_on_first_frame() { - let client = make_client(); - let w = 32; - let h = 32; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[8.0 / 255.0]); - - let params = NlmParams { - temporal_radius: 0, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::External, - motion_compensation: MotionCompensationMode::None, - hq: Some(HqParams { - auto_strength: true, - noise_floor: true, - sigma_override: None, - temporal_confidence: true, - thsad_scale: 1.0, - sigma_scale: 1.0, - windowed_noise_estimation: false, - }), - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame_with_reference(&frame, &frame); - assert!( - d.noise_estimator.current().is_some(), - "push_frame_with_reference must seed the noise estimator on the stream's first \ - frame, the same as push_frame" - ); -} - -/// Separable path (patch_radius > 2) variant of the identity check. -#[test] -fn external_reference_separable_matches_baseline() { - let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.3, 8, 8, 2, 0.7); - - let baseline = { - let params = NlmParams { - temporal_radius: 0, - search_radius: 2, - patch_radius: 4, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::None, - hq: None, - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() - }; - - let with_ref = { - let params = NlmParams { - temporal_radius: 0, - search_radius: 2, - patch_radius: 4, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::External, - motion_compensation: MotionCompensationMode::None, - hq: None, - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame_with_reference(&frame, &frame); - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() - }; - - for (i, (a, b)) in baseline.iter().zip(with_ref.iter()).enumerate() { - assert!((a - b).abs() < 1e-4, "pixel {i}: baseline={a}, with_ref={b}"); - } -} - -#[test] -fn external_reference_temporal_matches_baseline() { - let client = make_client(); - let w = 16; - let h = 16; - let frames = [ - make_frame_with_noisy_region(w, h, 1, 0.3, 8, 8, 2, 0.7), - make_frame_with_noisy_region(w, h, 1, 0.3, 7, 8, 2, 0.65), - make_frame_with_noisy_region(w, h, 1, 0.3, 9, 8, 2, 0.75), - ]; - - let baseline = { - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::None, - hq: None, - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - for f in &frames { - d.push_frame(f); - } - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() - }; - - let with_ref = { - let params = NlmParams { - temporal_radius: 1, - search_radius: 2, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::External, - motion_compensation: MotionCompensationMode::None, - hq: None, - }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - for f in &frames { - d.push_frame_with_reference(f, f); - } - d.denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec() - }; - - for (i, (a, b)) in baseline.iter().zip(with_ref.iter()).enumerate() { - assert!((a - b).abs() < 1e-5, "pixel {i}: baseline={a}, with_ref={b}"); - } + sum / count as f32 } -/// Bilateral prefilter on a uniform image must reproduce the uniform -/// value exactly (weights sum to anything, but the weighted average of -/// identical values is itself). #[test] fn bilateral_uniform_image_passthrough() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 16; + let height = 16; + let frame = make_uniform_frame(width, height, 1, 0.5); let params = NlmParams { temporal_radius: 0, @@ -289,30 +68,21 @@ fn bilateral_uniform_image_passthrough() { hq: None, }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - for (i, &v) in result.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-4, "pixel {i}: expected 0.5, got {v}"); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); + + for (i, &value) in result.iter().enumerate() { + assert!((value - 0.5).abs() < 1e-4, "pixel {i}: expected 0.5, got {value}"); } } -/// Bilateral smoke test on noisy input. Verifies the kernel produces -/// finite, in-range outputs (we trust the kernel correctness from the -/// uniform-image and identity tests). #[test] fn bilateral_noisy_image_finite() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_frame_with_noisy_region(w, h, 1, 0.4, 8, 8, 3, 0.8); + let width = 16; + let height = 16; + let frame = make_frame_with_noisy_region(width, height, 1, 0.4, 8, 8, 3, 0.8); let params = NlmParams { temporal_radius: 0, @@ -329,19 +99,16 @@ fn bilateral_noisy_image_finite() { hq: None, }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - for (i, &v) in result.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((-0.01..=1.01).contains(&v), "pixel {i}: out-of-range output {v}"); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); + + for (i, &value) in result.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (-0.01..=1.01).contains(&value), + "pixel {i}: out-of-range output {value}" + ); } } @@ -362,16 +129,18 @@ fn nlm_spatial_params() -> NlmParams { #[test] fn nlm_spatial_pilot_fills_reference() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[6.0 / 255.0]); + let width = 16; + let height = 16; + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[6.0 / 255.0]); - let mut d = NlmDenoiser::::new(&client, nlm_spatial_params(), w, h); - d.push_frame(&frame); + let params = nlm_spatial_params(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); - let reference = read_single_slot_reference(&d); + let reference = read_single_slot_reference(&denoiser); assert_eq!(reference.len(), frame.len()); + let mut differs = false; for (i, (&input, &pilot)) in frame.iter().zip(reference.iter()).enumerate() { assert!(pilot.is_finite(), "pixel {i}: non-finite pilot output {pilot}"); @@ -379,27 +148,30 @@ fn nlm_spatial_pilot_fills_reference() { (0.0..=1.0).contains(&pilot), "pixel {i}: out-of-range pilot output {pilot}" ); + if (input - pilot).abs() > 1e-6 { differs = true; } } + assert!(differs, "pilot output must differ from the noisy input somewhere"); } #[test] fn nlm_spatial_pilot_smooths() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[10.0 / 255.0]); + let width = 32; + let height = 32; + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[10.0 / 255.0]); - let mut d = NlmDenoiser::::new(&client, nlm_spatial_params(), w, h); - d.push_frame(&frame); + let params = nlm_spatial_params(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); - let reference = read_single_slot_reference(&d); + let reference = read_single_slot_reference(&denoiser); - let input_roughness = mean_abs_neighbour_diff(&frame, w, h); - let pilot_roughness = mean_abs_neighbour_diff(&reference, w, h); + let input_roughness = mean_abs_neighbour_diff(&frame, width, height); + let pilot_roughness = mean_abs_neighbour_diff(&reference, width, height); assert!( pilot_roughness < input_roughness, @@ -407,34 +179,30 @@ fn nlm_spatial_pilot_smooths() { ); } -/// A uniform frame has zero patch distance everywhere, so the pilot's -/// weighted average reproduces the input value exactly. #[test] fn nlm_spatial_pilot_uniform_passthrough() { let client = make_client(); - let w = 16; - let h = 16; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 16; + let height = 16; + let frame = make_uniform_frame(width, height, 1, 0.5); - let mut d = NlmDenoiser::::new(&client, nlm_spatial_params(), w, h); - d.push_frame(&frame); + let params = nlm_spatial_params(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); - let reference = read_single_slot_reference(&d); + let reference = read_single_slot_reference(&denoiser); - for (i, &v) in reference.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-4, "pixel {i}: expected 0.5, got {v}"); + for (i, &value) in reference.iter().enumerate() { + assert!((value - 0.5).abs() < 1e-4, "pixel {i}: expected 0.5, got {value}"); } } -/// The main-pass offset must be zeroed for `NlmSpatial` (pilot-vs-pilot -/// distances no longer carry the noise floor) while the pilot-facing -/// input offset keeps the full value the HQ noise-floor math would -/// otherwise apply to the main pass too. +/// Pilot-vs-pilot distances carry no noise floor, so only the pilot-facing input offset keeps it. #[test] fn nlm_spatial_zeros_main_offset_but_keeps_input_offset() { let client = make_client(); - let w = 16; - let h = 16; + let width = 16; + let height = 16; let sigma = 8.0 / 255.0; let params = NlmParams { @@ -449,7 +217,7 @@ fn nlm_spatial_zeros_main_offset_but_keeps_input_offset() { "test setup: expected a nonzero noise floor" ); - let denoiser = NlmDenoiser::::new(&client, params, w, h); + let denoiser = NlmDenoiser::::new(&client, params, width, height); assert_eq!( denoiser.noise_offset, 0.0, diff --git a/av-denoise-core/src/nlmeans/tests/residual_correlation.rs b/av-denoise-core/src/nlmeans/tests/residual_correlation.rs deleted file mode 100644 index 3533c82..0000000 --- a/av-denoise-core/src/nlmeans/tests/residual_correlation.rs +++ /dev/null @@ -1,1303 +0,0 @@ -use super::helpers::*; -use crate::nlmeans::*; - -/// Pearson correlation between a field and its neighbour one pixel to -/// the right, skipping the last column so no clamped edge pixel is ever -/// paired with itself. -fn lag1_horizontal(field: &[f32], w: u32, h: u32) -> f64 { - let mut sum_a = 0.0f64; - let mut sum_b = 0.0f64; - let mut sum_ab = 0.0f64; - let mut sum_aa = 0.0f64; - let mut sum_bb = 0.0f64; - let mut n = 0.0f64; - for y in 0..h { - for x in 0..(w - 1) { - let a = field[(y * w + x) as usize] as f64; - let b = field[(y * w + x + 1) as usize] as f64; - sum_a += a; - sum_b += b; - sum_ab += a * b; - sum_aa += a * a; - sum_bb += b * b; - n += 1.0; - } - } - let mean_a = sum_a / n; - let mean_b = sum_b / n; - let cov = sum_ab / n - mean_a * mean_b; - let var_a = sum_aa / n - mean_a * mean_a; - let var_b = sum_bb / n - mean_b * mean_b; - cov / (var_a.sqrt() * var_b.sqrt()) -} - -/// Same as [`lag1_horizontal`] but along y, skipping the last row. -fn lag1_vertical(field: &[f32], w: u32, h: u32) -> f64 { - let mut sum_a = 0.0f64; - let mut sum_b = 0.0f64; - let mut sum_ab = 0.0f64; - let mut sum_aa = 0.0f64; - let mut sum_bb = 0.0f64; - let mut n = 0.0f64; - for y in 0..(h - 1) { - for x in 0..w { - let a = field[(y * w + x) as usize] as f64; - let b = field[((y + 1) * w + x) as usize] as f64; - sum_a += a; - sum_b += b; - sum_ab += a * b; - sum_aa += a * a; - sum_bb += b * b; - n += 1.0; - } - } - let mean_a = sum_a / n; - let mean_b = sum_b / n; - let cov = sum_ab / n - mean_a * mean_b; - let var_a = sum_aa / n - mean_a * mean_a; - let var_b = sum_bb / n - mean_b * mean_b; - cov / (var_a.sqrt() * var_b.sqrt()) -} - -/// The actual standard deviation a `[a, 1 - 2a, a]` horizontal blur of -/// unit-variance white noise leaves behind, scaled by `sigma_pre`. -/// `a = 0` (no blur) reduces to `sigma_pre` itself, matching -/// [`noisy_field_over`]'s plain injection. -fn tap_sigma(sigma_pre: f32, a: f32) -> f32 { - let b = 1.0 - 2.0 * a; - sigma_pre * (2.0 * a * a + b * b).sqrt() -} - -fn std_dev(field: &[f32]) -> f64 { - let n = field.len() as f64; - let mean: f64 = field.iter().map(|&v| v as f64).sum::() / n; - let var: f64 = field.iter().map(|&v| (v as f64 - mean).powi(2)).sum::() / n; - var.sqrt() -} - -/// One measured row: an input grain correlation run through the NLM -/// front end at a given search/temporal radius, reporting what comes -/// out the other side. -struct Measurement { - rho_out_h: f64, - rho_out_v: f64, - sigma_ratio: f64, -} - -/// Pushes `2 * temporal_radius + 1` independently-seeded noisy frames -/// (via `make_noise(seed)`) through the front end and reads back the -/// output emitted on the final push. -/// -/// That final push is the one whose center frame is fully real on both -/// sides (no leading-mirror duplicate anywhere in its window), so the -/// residual it produces reflects genuine averaging across distinct -/// noise realisations, not a partially-duplicated one. -/// -/// Compares the output against the flat `clean` reference to get the -/// residual, and the middle pushed frame (the one the output is -/// centered on) against the same reference to get the pre-denoise -/// input noise, so the sigma ratio and the input frame's own -/// correlation are measured on the exact same noise realisation the -/// output derives from. -#[expect( - clippy::too_many_arguments, - reason = "the test helper takes the full set of parameters its cases vary" -)] -fn measure( - client: &cubecl::prelude::ComputeClient, - w: u32, - h: u32, - base: f32, - sigma: f32, - make_noise: impl Fn(u32) -> Vec, - search_radius: u32, - temporal_radius: u32, -) -> Measurement { - let params = NlmParams { - temporal_radius, - search_radius, - patch_radius: 2, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::None, - hq: Some(HqParams::with_sigma(sigma)), - }; - let mut denoiser = NlmDenoiser::::new(client, params, w, h); - - let n_push = 2 * temporal_radius + 1; - let mut center_noisy: Option> = None; - let mut output: Option> = None; - for i in 0..n_push { - let frame = make_noise(100 + i); - if i == temporal_radius { - center_noisy = Some(frame.clone()); - } - denoiser.push_frame(&frame); - let result = denoiser.denoise().unwrap(); - if i == n_push - 1 { - output = result.map(|o| o.as_f32().expect("f32 denoiser").to_vec()); - } - } - - let output = output.expect("a fully real window must emit on its final push"); - let center_noisy = center_noisy.expect("center frame must have been pushed"); - - let clean = vec![base; (w * h) as usize]; - let residual: Vec = output.iter().zip(clean.iter()).map(|(&o, &c)| o - c).collect(); - let input_noise: Vec = center_noisy - .iter() - .zip(clean.iter()) - .map(|(&o, &c)| o - c) - .collect(); - - Measurement { - rho_out_h: lag1_horizontal(&residual, w, h), - rho_out_v: lag1_vertical(&residual, w, h), - sigma_ratio: std_dev(&residual) / std_dev(&input_noise), - } -} - -/// Measures how the NLM front end's own residual noise correlation -/// compares to the correlation of the grain it was fed, across a -/// sweep of input correlations and window shapes. -/// -/// The collaborative second stage that follows NLM assumes white -/// (uncorrelated) noise when it shrinks DCT coefficients. It never sees -/// the input grain, only NLM's residual, so the question a noise-shaping -/// fix needs answered is not "what is the input grain's correlation" but -/// "what correlation does NLM's own output residual actually carry." -/// -/// NLM is a local weighted average, and neighbouring output pixels draw -/// on heavily overlapping windows of input pixels, so the mechanism -/// predicts the residual comes out *more* correlated than the input, -/// purely from the window overlap, independent of whether the input -/// noise carried any correlation of its own. A search radius of 0 -/// removes that overlap (each pixel then only ever draws on its own -/// position across frames), so it isolates the mechanism: if residual -/// correlation collapses toward the input's own correlation at search -/// radius 0, that confirms the window overlap, not the input grain, is -/// what's inducing it. -/// -/// Three input correlations are covered: uncorrelated grain -/// ([`noisy_field_over`]), an intermediate tap -/// ([`correlated_noisy_frame_with_tap`] at `a = 0.125`, predicted -/// `rho ~= 0.316`), and [`correlated_noisy_frame`]'s fixed two thirds. -/// All three are horizontal-only blurs, so the input correlation is -/// anisotropic by construction (`rho_v ~= 0` for every one of them); -/// both axes are measured on the residual regardless, since NLM's own -/// window is isotropic and could induce vertical correlation the input -/// never had. -#[test] -fn nlm_residual_correlation_exceeds_input_correlation_and_tracks_window_overlap() { - let client = make_client(); - let w = 160; - let h = 160; - let base = 0.5f32; - let sigma_pre = 0.06f32; - - struct Case { - label: &'static str, - rho_in_h: f64, - rho_in_v: f64, - } - let cases = [ - Case { - label: "rho_in=0.00", - rho_in_h: 0.0, - rho_in_v: 0.0, - }, - Case { - label: "rho_in=0.32", - rho_in_h: 0.316, - rho_in_v: 0.0, - }, - Case { - label: "rho_in=0.67", - rho_in_h: 2.0 / 3.0, - rho_in_v: 0.0, - }, - ]; - - let configs = [(2u32, 0u32), (2u32, 2u32), (0u32, 0u32), (0u32, 2u32)]; - - eprintln!( - "residual correlation sweep (w={w} h={h} sigma_pre={sigma_pre}):\n\ - {:<12} {:>6} {:>6} | {:>6} {:>6} | {:>8} {:>8} | {:>8}", - "input", "R_s", "R_t", "rho_in_h", "rho_in_v", "rho_out_h", "rho_out_v", "sig_ratio" - ); - - // Collect every row so the assertions below can reason about the - // whole table at once instead of one measurement in isolation. - struct Row { - label: &'static str, - search_radius: u32, - temporal_radius: u32, - rho_in_h: f64, - m: Measurement, - } - let mut rows = Vec::new(); - - for case in &cases { - for &(search_radius, temporal_radius) in &configs { - let m = match case.label { - "rho_in=0.00" => { - let clean = vec![base; (w * h) as usize]; - measure( - &client, - w, - h, - base, - tap_sigma(sigma_pre, 0.0), - |seed| noisy_field_over(&clean, w, h, sigma_pre, seed), - search_radius, - temporal_radius, - ) - }, - "rho_in=0.32" => measure( - &client, - w, - h, - base, - tap_sigma(sigma_pre, 0.125), - |seed| correlated_noisy_frame_with_tap(w, h, base, sigma_pre, seed, 0.125), - search_radius, - temporal_radius, - ), - _ => measure( - &client, - w, - h, - base, - tap_sigma(sigma_pre, 0.25), - |seed| correlated_noisy_frame(w, h, base, sigma_pre, seed), - search_radius, - temporal_radius, - ), - }; - - eprintln!( - "{:<12} {:>6} {:>6} | {:>8.4} {:>8.4} | {:>8.4} {:>8.4} | {:>8.4}", - case.label, - search_radius, - temporal_radius, - case.rho_in_h, - case.rho_in_v, - m.rho_out_h, - m.rho_out_v, - m.sigma_ratio - ); - - rows.push(Row { - label: case.label, - search_radius, - temporal_radius, - rho_in_h: case.rho_in_h, - m, - }); - } - } - - // The core mechanism claim: at search_radius=2, the residual's - // horizontal correlation must exceed the input's, for every input - // correlation tested including the uncorrelated one. Window overlap - // alone induces correlation that was never in the input. - for row in &rows { - if row.search_radius == 2 { - assert!( - row.m.rho_out_h > row.rho_in_h, - "{} at search_radius=2 temporal_radius={}: residual rho_h={:.4} did not exceed \ - input rho_h={:.4}, expected window overlap to raise it", - row.label, - row.temporal_radius, - row.m.rho_out_h, - row.rho_in_h - ); - } - } - - // The isolating claim: at search_radius=0, no spatial window overlap - // exists (each pixel only ever draws on its own position across - // frames), so residual rho_h should track the input's own - // correlation closely rather than staying inflated. The tolerance - // (0.08 absolute) covers the small, real shift temporal averaging at - // temporal_radius=2 introduces (measured up to ~0.025 here), not - // measurement noise, since every generator is seeded and - // deterministic. - for row in &rows { - if row.search_radius == 0 { - assert!( - (row.m.rho_out_h - row.rho_in_h).abs() < 0.08, - "{} at search_radius=0 temporal_radius={}: residual rho_h={:.4} should track the \ - input's rho_h={:.4} within 0.08 once spatial window overlap is removed", - row.label, - row.temporal_radius, - row.m.rho_out_h, - row.rho_in_h - ); - } - } - - // The residual's vertical correlation must stay small everywhere the - // input's vertical correlation was zero and no spatial window is in - // play (search_radius=0): there is no mixing mechanism at play there - // to manufacture it. - for row in &rows { - if row.search_radius == 0 { - assert!( - row.m.rho_out_v.abs() < 0.05, - "{} at search_radius=0 temporal_radius={}: residual rho_v={:.4} should stay near \ - zero, the input never carried vertical correlation and there is no spatial \ - window to manufacture it", - row.label, - row.temporal_radius, - row.m.rho_out_v - ); - } - } - - // The surprising finding this test exists to pin down: NLM's own - // search window is isotropic (a square), and once it is in play - // (search_radius=2) it dominates over the input's own anisotropy. - // Even though every input generator here is horizontal-only - // (rho_in_v is always 0), the residual's vertical correlation still - // comes out substantial, well above what the input ever carried, - // because the window mixes vertical neighbours' noise together too. - // A separable noise-shaping fix that assumes the same correlation - // profile on both axes is closer to the truth here than one that - // trusted the input's own anisotropy. - for row in &rows { - if row.search_radius == 2 { - assert!( - row.m.rho_out_v > 0.3, - "{} at search_radius=2 temporal_radius={}: residual rho_v={:.4} did not rise \ - well above the input's zero vertical correlation, expected the isotropic \ - window to induce substantial vertical correlation regardless", - row.label, - row.temporal_radius, - row.m.rho_out_v - ); - } - } -} - -// --------------------------------------------------------------------- -// Follow-up: does the flat-content curve above hold at every search -// radius that actually ships, and does it survive on textured content -// where NLM's weights are far less uniform than on a flat field? -// --------------------------------------------------------------------- - -/// Bounded variant of [`lag1_horizontal`], restricted to the half-open -/// rectangle `[x0, x1) x [y0, y1)`. Lets a single frame's correlation be -/// measured separately over sub-regions, a flat tile and a textured -/// tile side by side, rather than only over the whole field at once. -fn lag1_horizontal_rect(field: &[f32], w: u32, x0: u32, y0: u32, x1: u32, y1: u32) -> f64 { - let mut sum_a = 0.0f64; - let mut sum_b = 0.0f64; - let mut sum_ab = 0.0f64; - let mut sum_aa = 0.0f64; - let mut sum_bb = 0.0f64; - let mut n = 0.0f64; - for y in y0..y1 { - for x in x0..(x1 - 1) { - let a = field[(y * w + x) as usize] as f64; - let b = field[(y * w + x + 1) as usize] as f64; - sum_a += a; - sum_b += b; - sum_ab += a * b; - sum_aa += a * a; - sum_bb += b * b; - n += 1.0; - } - } - let mean_a = sum_a / n; - let mean_b = sum_b / n; - let cov = sum_ab / n - mean_a * mean_b; - let var_a = sum_aa / n - mean_a * mean_a; - let var_b = sum_bb / n - mean_b * mean_b; - cov / (var_a.sqrt() * var_b.sqrt()) -} - -/// Same as [`lag1_horizontal_rect`] but along y. -fn lag1_vertical_rect(field: &[f32], w: u32, x0: u32, y0: u32, x1: u32, y1: u32) -> f64 { - let mut sum_a = 0.0f64; - let mut sum_b = 0.0f64; - let mut sum_ab = 0.0f64; - let mut sum_aa = 0.0f64; - let mut sum_bb = 0.0f64; - let mut n = 0.0f64; - for y in y0..(y1 - 1) { - for x in x0..x1 { - let a = field[(y * w + x) as usize] as f64; - let b = field[((y + 1) * w + x) as usize] as f64; - sum_a += a; - sum_b += b; - sum_ab += a * b; - sum_aa += a * a; - sum_bb += b * b; - n += 1.0; - } - } - let mean_a = sum_a / n; - let mean_b = sum_b / n; - let cov = sum_ab / n - mean_a * mean_b; - let var_a = sum_aa / n - mean_a * mean_a; - let var_b = sum_bb / n - mean_b * mean_b; - cov / (var_a.sqrt() * var_b.sqrt()) -} - -/// Standard deviation restricted to the half-open rectangle -/// `[x0, x1) x [y0, y1)`. -fn std_dev_rect(field: &[f32], w: u32, x0: u32, y0: u32, x1: u32, y1: u32) -> f64 { - let n = ((x1 - x0) * (y1 - y0)) as f64; - let mut sum = 0.0f64; - for y in y0..y1 { - for x in x0..x1 { - sum += field[(y * w + x) as usize] as f64; - } - } - let mean = sum / n; - let mut var = 0.0f64; - for y in y0..y1 { - for x in x0..x1 { - let v = field[(y * w + x) as usize] as f64; - var += (v - mean).powi(2); - } - } - (var / n).sqrt() -} - -/// Runs the NLM front end once over `2 * temporal_radius + 1` pushed -/// frames of `make_noise(clean, seed)`, seeded from `seed_base`, and -/// returns the output emitted on the final push (the one whose center -/// frame has a fully real window on both sides) together with the noisy -/// frame that produced that center. A generalisation of the sweep -/// above's inlined push loop, parameterised over `clean` and -/// `patch_radius` so it can drive both a flat and a textured reference. -#[expect( - clippy::too_many_arguments, - reason = "the test helper takes the full set of parameters its cases vary" -)] -fn run_front_end( - client: &cubecl::prelude::ComputeClient, - w: u32, - h: u32, - clean: &[f32], - sigma: f32, - make_noise: impl Fn(&[f32], u32) -> Vec, - search_radius: u32, - temporal_radius: u32, - patch_radius: u32, - seed_base: u32, -) -> (Vec, Vec) { - let params = NlmParams { - temporal_radius, - search_radius, - patch_radius, - strength: 1.2, - self_weight: 1.0, - channels: ChannelMode::Luma, - prefilter: PrefilterMode::None, - motion_compensation: MotionCompensationMode::None, - hq: Some(HqParams::with_sigma(sigma)), - }; - let mut denoiser = NlmDenoiser::::new(client, params, w, h); - - let n_push = 2 * temporal_radius + 1; - let mut center_noisy: Option> = None; - let mut output: Option> = None; - for i in 0..n_push { - let frame = make_noise(clean, seed_base + i); - if i == temporal_radius { - center_noisy = Some(frame.clone()); - } - denoiser.push_frame(&frame); - let result = denoiser.denoise().unwrap(); - if i == n_push - 1 { - output = result.map(|o| o.as_f32().expect("f32 denoiser").to_vec()); - } - } - - ( - output.expect("a fully real window must emit on its final push"), - center_noisy.expect("center frame must have been pushed"), - ) -} - -/// One measured row of residual correlation, common to both the -/// flat-content and the differenced textured-content measurements -/// below. -struct Sample { - rho_out_h: f64, - rho_out_v: f64, - sigma_ratio: f64, -} - -/// Measures residual correlation on content whose clean reference is -/// flat, where `output - clean` is already an unbiased noise-only -/// residual: a flat field has no structure for NLM to respond to, so -/// nothing but noise-driven averaging can separate the output from the -/// clean value it started from. Same measurement `measure` above makes, -/// generalised over an arbitrary flat `clean` and `patch_radius`. -#[expect( - clippy::too_many_arguments, - reason = "the test helper takes the full set of parameters its cases vary" -)] -fn measure_flat( - client: &cubecl::prelude::ComputeClient, - w: u32, - h: u32, - clean: &[f32], - sigma: f32, - make_noise: impl Fn(&[f32], u32) -> Vec, - search_radius: u32, - temporal_radius: u32, - patch_radius: u32, -) -> Sample { - let (output, center_noisy) = run_front_end( - client, - w, - h, - clean, - sigma, - make_noise, - search_radius, - temporal_radius, - patch_radius, - 100, - ); - let residual: Vec = output.iter().zip(clean.iter()).map(|(&o, &c)| o - c).collect(); - let input_noise: Vec = center_noisy - .iter() - .zip(clean.iter()) - .map(|(&o, &c)| o - c) - .collect(); - Sample { - rho_out_h: lag1_horizontal(&residual, w, h), - rho_out_v: lag1_vertical(&residual, w, h), - sigma_ratio: std_dev(&residual) / std_dev(&input_noise), - } -} - -/// Measures residual correlation on content with real structure, where -/// `output - clean` is not a clean noise residual on its own: NLM's own -/// weighted average responds to the texture's edges and gradients too, -/// and that structural response is a deterministic function of the -/// clean image, not noise. Subtracting a fixed clean reference would -/// fold that deterministic bias into what is supposed to be a -/// noise-only correlation measurement. -/// -/// Instead this runs the front end twice over the same clean field, -/// with two independently-seeded noise realisations (`seed_base` 100 -/// and 500, far enough apart that even the widest temporal window used -/// anywhere in this file, `2 * 4 + 1 = 9` pushed frames, cannot make the -/// two seed ranges overlap), and differences the two outputs. -/// -/// For two identically-distributed, independent copies `A` and `B` of -/// the same random field, `Corr(A - B) == Corr(A)` exactly: the -/// deterministic structural component the two outputs would otherwise -/// share cancels out of the difference, and both the variance and the -/// covariance of `A - B` double relative to `A` alone in the same -/// proportion (independence removes any cross term), so their ratio, -/// the correlation, is unchanged. `A - B`'s own lag-1 correlation is -/// already the answer, no separate bias correction needed. Its standard -/// deviation is `sqrt(2)` times a single realisation's, so `sigma_ratio` -/// divides that back out before comparing against the input noise's own -/// standard deviation. -#[expect( - clippy::too_many_arguments, - reason = "the test helper takes the full set of parameters its cases vary" -)] -fn measure_diff( - client: &cubecl::prelude::ComputeClient, - w: u32, - h: u32, - clean: &[f32], - sigma: f32, - make_noise: impl Fn(&[f32], u32) -> Vec, - search_radius: u32, - temporal_radius: u32, - patch_radius: u32, -) -> Sample { - let (output_a, noisy_a) = run_front_end( - client, - w, - h, - clean, - sigma, - &make_noise, - search_radius, - temporal_radius, - patch_radius, - 100, - ); - let (output_b, _noisy_b) = run_front_end( - client, - w, - h, - clean, - sigma, - &make_noise, - search_radius, - temporal_radius, - patch_radius, - 500, - ); - let diff: Vec = output_a - .iter() - .zip(output_b.iter()) - .map(|(&a, &b)| a - b) - .collect(); - let input_noise: Vec = noisy_a.iter().zip(clean.iter()).map(|(&o, &c)| o - c).collect(); - let sqrt2 = 2.0f64.sqrt(); - Sample { - rho_out_h: lag1_horizontal(&diff, w, h), - rho_out_v: lag1_vertical(&diff, w, h), - sigma_ratio: (std_dev(&diff) / sqrt2) / std_dev(&input_noise), - } -} - -/// Runs [`measure_diff`] but reports correlation separately over two -/// sub-rectangles of the frame instead of the whole field, so a single -/// frame containing both flat and textured regions can be checked for -/// whether its residual correlation actually differs between them. -#[expect( - clippy::too_many_arguments, - reason = "the test helper takes the full set of parameters its cases vary" -)] -fn measure_diff_two_regions( - client: &cubecl::prelude::ComputeClient, - w: u32, - h: u32, - clean: &[f32], - sigma: f32, - make_noise: impl Fn(&[f32], u32) -> Vec, - search_radius: u32, - temporal_radius: u32, - patch_radius: u32, - region_a: (u32, u32, u32, u32), - region_b: (u32, u32, u32, u32), -) -> (Sample, Sample) { - let (output_a, noisy_a) = run_front_end( - client, - w, - h, - clean, - sigma, - &make_noise, - search_radius, - temporal_radius, - patch_radius, - 100, - ); - let (output_b, _noisy_b) = run_front_end( - client, - w, - h, - clean, - sigma, - &make_noise, - search_radius, - temporal_radius, - patch_radius, - 500, - ); - let diff: Vec = output_a - .iter() - .zip(output_b.iter()) - .map(|(&a, &b)| a - b) - .collect(); - let input_noise: Vec = noisy_a.iter().zip(clean.iter()).map(|(&o, &c)| o - c).collect(); - let sqrt2 = 2.0f64.sqrt(); - - let sample_for = |(x0, y0, x1, y1): (u32, u32, u32, u32)| Sample { - rho_out_h: lag1_horizontal_rect(&diff, w, x0, y0, x1, y1), - rho_out_v: lag1_vertical_rect(&diff, w, x0, y0, x1, y1), - sigma_ratio: (std_dev_rect(&diff, w, x0, y0, x1, y1) / sqrt2) - / std_dev_rect(&input_noise, w, x0, y0, x1, y1), - }; - - (sample_for(region_a), sample_for(region_b)) -} - -/// Primary sweep for the search-radius-to-residual-correlation table: -/// every search radius that actually ships (`1..=4`; the presets use 2 -/// and 4), measured on both flat and textured content, uncorrelated -/// input noise throughout. -/// -/// `task-D1-residual-rho-report.md` flagged its flat-content-only -/// numbers as a likely upper bound: on flat content every candidate in -/// NLM's search window carries an equal weight, but on real footage the -/// Welsch weighting suppresses candidates whose patch doesn't actually -/// match, so less averaging happens and the induced correlation should -/// come out lower. This sweep checks that directly instead of trusting -/// the extrapolation, by running the identical measurement on -/// [`make_textured_frame`] alongside the flat field, at every shipping -/// search radius. -/// -/// Textured content cannot reuse the flat measurement's `output - clean` -/// trick (see [`measure_diff`]'s doc comment for why), so it is measured -/// by differencing two independently-seeded denoised outputs of the same -/// clean field instead; [`Sample::sigma_ratio`] on that side is already -/// rescaled back to a single-realisation footing. -#[test] -fn nlm_residual_correlation_search_radius_sweep_flat_vs_textured() { - let client = make_client(); - let w = 160; - let h = 160; - let base = 0.5f32; - let sigma_pre = 0.06f32; - let patch_radius = 2; - - let flat_clean = vec![base; (w * h) as usize]; - let textured_clean = make_textured_frame(w, h); - - struct Row { - search_radius: u32, - flat: Sample, - textured: Sample, - } - let mut rows = Vec::new(); - - eprintln!( - "search radius sweep, flat vs textured (w={w} h={h} sigma_pre={sigma_pre} \ - patch_radius={patch_radius}, uncorrelated input):\n\ - {:>3} | {:>8} {:>8} {:>8} | {:>8} {:>8} {:>8}", - "R_s", "flat_h", "flat_v", "flat_sig", "tex_h", "tex_v", "tex_sig" - ); - - for &search_radius in &[1u32, 2, 3, 4] { - let flat = measure_flat( - &client, - w, - h, - &flat_clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 0, - patch_radius, - ); - let textured = measure_diff( - &client, - w, - h, - &textured_clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 0, - patch_radius, - ); - - eprintln!( - "{:>3} | {:>8.4} {:>8.4} {:>8.4} | {:>8.4} {:>8.4} {:>8.4}", - search_radius, - flat.rho_out_h, - flat.rho_out_v, - flat.sigma_ratio, - textured.rho_out_h, - textured.rho_out_v, - textured.sigma_ratio - ); - - rows.push(Row { - search_radius, - flat, - textured, - }); - } - - // Validity guard: every configuration must show real smoothing - // before its correlation number means anything. A sigma_ratio near - // 1.0 means the front end barely denoised at all, the failure mode - // task-D1 hit with an uncalibrated strength. - for row in &rows { - assert!( - row.flat.sigma_ratio < 0.9, - "search_radius={}: flat sigma_ratio={:.4} too close to 1.0, no real smoothing \ - happened, this configuration's correlation number is not valid data", - row.search_radius, - row.flat.sigma_ratio - ); - assert!( - row.textured.sigma_ratio < 0.9, - "search_radius={}: textured sigma_ratio={:.4} too close to 1.0, no real smoothing \ - happened, this configuration's correlation number is not valid data", - row.search_radius, - row.textured.sigma_ratio - ); - } - - // The core regression guard: window overlap induces real - // correlation on flat content at every shipping search radius, not - // just the radius-2 case task-D1 covered. - for row in &rows { - assert!( - row.flat.rho_out_h > 0.3, - "search_radius={}: flat rho_out_h={:.4} did not show substantial window-induced \ - correlation", - row.search_radius, - row.flat.rho_out_h - ); - } - - // The headline finding this sweep exists to check: contrary to - // task-D1's "probably an upper bound" caveat, the sinusoidal texture - // in make_textured_frame tracks the flat field closely at every - // radius (measured max gap 0.0234 at search_radius=4). The 0.08 - // tolerance leaves comfortable margin above that while still - // catching a real divergence between the two content types. - for row in &rows { - assert!( - (row.flat.rho_out_h - row.textured.rho_out_h).abs() < 0.08, - "search_radius={}: flat rho_out_h={:.4} and textured rho_out_h={:.4} diverge by more \ - than the measured tolerance, flat and textured no longer agree", - row.search_radius, - row.flat.rho_out_h, - row.textured.rho_out_h - ); - } - - // Residual correlation should not fall as the window widens: a - // bigger search radius only ever adds more overlapping candidates, - // never fewer, so the induced correlation should rise or plateau, - // not reverse. - for pair in rows.windows(2) { - assert!( - pair[1].flat.rho_out_h >= pair[0].flat.rho_out_h - 1e-6, - "flat rho_out_h dropped from search_radius={} ({:.4}) to search_radius={} ({:.4})", - pair[0].search_radius, - pair[0].flat.rho_out_h, - pair[1].search_radius, - pair[1].flat.rho_out_h - ); - } -} - -/// The table this file's primary sweep produces at the shipped default -/// `patch_radius`, across every search radius the presets can select -/// (`0..=4`). -/// -/// The primary sweep above measured at `patch_radius=2`, following -/// `task-D1-residual-rho-report.md`'s convention. The library default is -/// `patch_radius=4` -/// ([`crate::nlmeans::NlmParams::patch_radius`]'s doc comment), and this -/// file's own secondary-checks test found that shift is not small: at -/// `search_radius=2`, `patch_radius=4` measured `rho_out_h=0.7770` -/// against `patch_radius=2`'s `0.7006`, a 0.077 gap. A table built at -/// the wrong patch radius does not describe the configuration that -/// actually runs, so this is the sweep to bake into the codebase, not -/// the `patch_radius=2` one above. -/// -/// `search_radius=0` is included here (the earlier sweep started at 1), -/// since a complete table needs every radius the presets can reach, and -/// `search_radius=0` is the isolating case: no spatial window overlap, -/// so residual correlation should collapse to near zero regardless of -/// `patch_radius` (a larger patch still only compares a single fixed -/// candidate against itself when there is no window to search). -#[test] -fn nlm_residual_correlation_search_radius_sweep_at_shipped_patch_radius() { - let client = make_client(); - let w = 160; - let h = 160; - let base = 0.5f32; - let sigma_pre = 0.06f32; - let patch_radius = 4; - - let flat_clean = vec![base; (w * h) as usize]; - let textured_clean = make_textured_frame(w, h); - - struct Row { - search_radius: u32, - flat: Sample, - textured: Sample, - } - let mut rows = Vec::new(); - - eprintln!( - "search radius sweep at shipped patch_radius={patch_radius}, flat vs textured \ - (w={w} h={h} sigma_pre={sigma_pre}, uncorrelated input):\n\ - {:>3} | {:>8} {:>8} {:>8} | {:>8} {:>8} {:>8}", - "R_s", "flat_h", "flat_v", "flat_sig", "tex_h", "tex_v", "tex_sig" - ); - - for &search_radius in &[0u32, 1, 2, 3, 4] { - let flat = measure_flat( - &client, - w, - h, - &flat_clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 0, - patch_radius, - ); - let textured = measure_diff( - &client, - w, - h, - &textured_clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 0, - patch_radius, - ); - - eprintln!( - "{:>3} | {:>8.4} {:>8.4} {:>8.4} | {:>8.4} {:>8.4} {:>8.4}", - search_radius, - flat.rho_out_h, - flat.rho_out_v, - flat.sigma_ratio, - textured.rho_out_h, - textured.rho_out_v, - textured.sigma_ratio - ); - - rows.push(Row { - search_radius, - flat, - textured, - }); - } - - // Validity guard. search_radius=0 is exempt: with no spatial window - // at all, HQ's calibrated weighting has essentially nothing to - // average over, so almost no smoothing happens and sigma_ratio sits - // near 1.0 by construction, not because the measurement is broken. - // That is exactly D1's isolating case, kept here for a complete - // table rather than trusted as a valid correlation measurement. - for row in &rows { - if row.search_radius == 0 { - continue; - } - assert!( - row.flat.sigma_ratio < 0.9, - "search_radius={}: flat sigma_ratio={:.4} too close to 1.0, no real smoothing \ - happened, this configuration's correlation number is not valid data", - row.search_radius, - row.flat.sigma_ratio - ); - assert!( - row.textured.sigma_ratio < 0.9, - "search_radius={}: textured sigma_ratio={:.4} too close to 1.0, no real smoothing \ - happened, this configuration's correlation number is not valid data", - row.search_radius, - row.textured.sigma_ratio - ); - } - - // The isolating case: search_radius=0 has no window to overlap, so - // residual correlation should stay near the uncorrelated input's own - // ~0, on both content types, mirroring D1's search_radius=0 finding. - for row in &rows { - if row.search_radius == 0 { - assert!( - row.flat.rho_out_h.abs() < 0.1, - "search_radius=0: flat rho_out_h={:.4} should stay near zero, there is no \ - spatial window to manufacture correlation", - row.flat.rho_out_h - ); - } - } - - // The core regression guard at every radius with a real window: flat - // content shows substantial window-induced correlation. - for row in &rows { - if row.search_radius >= 1 { - assert!( - row.flat.rho_out_h > 0.3, - "search_radius={}: flat rho_out_h={:.4} did not show substantial window-induced \ - correlation", - row.search_radius, - row.flat.rho_out_h - ); - } - } - - // Every row here should sit at or above its patch_radius=2 - // counterpart from the primary sweep: a larger patch radius makes - // the per-candidate weight estimate less noisy (more pixels - // averaged into the distance), which lets the window's own - // near-uniform weighting come through more cleanly rather than - // being scattered by measurement noise, so it raises the induced - // correlation rather than lowering it (the secondary-checks test - // above already pins this down at search_radius=2 alone; this - // extends the same direction check across the radii newly measured - // here). - for row in &rows { - if row.search_radius == 0 { - continue; - } - assert!( - row.flat.rho_out_h > 0.6, - "search_radius={}: patch_radius=4 flat rho_out_h={:.4} unexpectedly low, expected it \ - to sit above the patch_radius=2 sweep's own values at this radius", - row.search_radius, - row.flat.rho_out_h - ); - } - - // Residual correlation should not fall as the window widens. - for pair in rows.windows(2) { - assert!( - pair[1].flat.rho_out_h >= pair[0].flat.rho_out_h - 1e-6, - "flat rho_out_h dropped from search_radius={} ({:.4}) to search_radius={} ({:.4})", - pair[0].search_radius, - pair[0].flat.rho_out_h, - pair[1].search_radius, - pair[1].flat.rho_out_h - ); - } -} - -/// Secondary checks, fewer points than the primary sweep, on what else -/// the residual correlation depends on besides `search_radius`. All -/// three vary one parameter at a time away from a shared baseline -/// (`search_radius=2`, `patch_radius=2`, `temporal_radius=0`, -/// uncorrelated input, flat content), reusing that baseline's own row -/// from the sweep above rather than re-measuring it. -#[test] -fn nlm_residual_correlation_patch_radius_temporal_radius_and_input_correlation() { - let client = make_client(); - let w = 160; - let h = 160; - let base = 0.5f32; - let sigma_pre = 0.06f32; - let flat_clean = vec![base; (w * h) as usize]; - let search_radius = 2; - - let baseline = measure_flat( - &client, - w, - h, - &flat_clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 0, - 2, - ); - - let patch4 = measure_flat( - &client, - w, - h, - &flat_clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 0, - 4, - ); - - let temporal2 = measure_flat( - &client, - w, - h, - &flat_clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 2, - 2, - ); - - let rho_in_h = 2.0 / 3.0; - let corr_input = measure_flat( - &client, - w, - h, - &flat_clean, - sigma_pre, - |_clean, seed| correlated_noisy_frame(w, h, base, sigma_pre, seed), - search_radius, - 0, - 2, - ); - - eprintln!( - "secondary checks at search_radius=2, flat content (w={w} h={h} sigma_pre={sigma_pre}):\n\ - {:<32} {:>8} {:>8} {:>8}", - "config", "rho_h", "rho_v", "sig_ratio" - ); - for (label, s) in [ - ("baseline patch_radius=2 R_t=0 rho_in=0", &baseline), - ("patch_radius=4", &patch4), - ("temporal_radius=2", &temporal2), - ("input rho_h=0.67", &corr_input), - ] { - eprintln!( - "{:<32} {:>8.4} {:>8.4} {:>8.4}", - label, s.rho_out_h, s.rho_out_v, s.sigma_ratio - ); - } - - // Validity guard on every configuration measured here. - for (label, s) in [ - ("baseline", &baseline), - ("patch_radius=4", &patch4), - ("temporal_radius=2", &temporal2), - ("input rho_h=0.67", &corr_input), - ] { - assert!( - s.sigma_ratio < 0.9, - "{label}: sigma_ratio={:.4} too close to 1.0, no real smoothing happened, this \ - configuration's correlation number is not valid data", - s.sigma_ratio - ); - } - - // task-D1 measured temporal averaging diluting the window-induced - // correlation slightly (~0.02-0.05 at search_radius=2). Confirm the - // direction holds with patch_radius=2 held fixed here too. - assert!( - temporal2.rho_out_h < baseline.rho_out_h, - "temporal_radius=2 rho_out_h={:.4} should be below temporal_radius=0's {:.4}, temporal \ - averaging is expected to dilute the spatial window's contribution", - temporal2.rho_out_h, - baseline.rho_out_h - ); - - // Correlated input should raise the residual's correlation further - // above the uncorrelated baseline (task-D1's core finding, that the - // window imposes a floor and the input's own correlation lifts it - // further within the remaining headroom below 1), and it must not - // exceed 1 either. - assert!( - corr_input.rho_out_h > baseline.rho_out_h, - "correlated input (rho_in_h={rho_in_h:.4}) rho_out_h={:.4} should exceed the uncorrelated \ - baseline's {:.4}", - corr_input.rho_out_h, - baseline.rho_out_h - ); - assert!( - corr_input.rho_out_h < 1.0, - "correlated input rho_out_h={:.4} must stay below 1.0", - corr_input.rho_out_h - ); -} - -/// A single clean frame with a flat region on the left and -/// [`make_textured_frame`]'s own sine formula on the right, split at -/// `split_x`. Lets one NLM run be checked for whether its residual -/// correlation actually differs between the two regions within the -/// same frame, rather than only across separately-generated frames. -fn make_flat_and_textured_frame(w: u32, h: u32, split_x: u32) -> Vec { - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let v = if x < split_x { - 0.5 - } else { - let fx = x as f32 / w as f32; - let fy = y as f32 / h as f32; - let raw = 0.5 - + 0.2 * (fx * 8.0 * std::f32::consts::PI).sin() * (fy * 6.0 * std::f32::consts::PI).cos() - + 0.1 * (fx * 20.0 * std::f32::consts::PI).sin(); - raw.clamp(0.05, 0.95) - }; - frame[(y * w + x) as usize] = v; - } - } - frame -} - -/// Whether residual correlation differs between a flat region and a -/// textured region within a *single* frame, rather than across two -/// separately-generated frames. -/// -/// The sweep above already runs flat and textured content as two -/// entirely separate frames, which cannot rule out that some other -/// difference between the two setups (frame content statistics beyond -/// just flat-vs-textured, boundary handling, whatever) is responsible -/// for how closely they tracked. Building one frame with both region -/// types side by side and measuring each region's own correlation -/// removes that confound: if a spatially-varying second-stage fix ever -/// turns out to be necessary, it is because the *same* frame carries -/// two different residual-correlation profiles at once, not because two -/// different test frames happened to differ. -/// -/// The two analysis tiles sit 15 pixels from the seam and from every -/// frame edge, comfortably outside `patch_radius + search_radius = 4`'s -/// reach, so neither NLM's own output near a tile boundary nor the -/// correlation measurement within a tile can be contaminated by the -/// other region. -#[test] -fn nlm_residual_correlation_within_a_single_frame_flat_vs_textured_regions() { - let client = make_client(); - let w = 200; - let h = 120; - let sigma_pre = 0.06f32; - let search_radius = 2; - let patch_radius = 2; - let split_x = 100; - - let clean = make_flat_and_textured_frame(w, h, split_x); - let flat_region = (15u32, 15u32, 85u32, 105u32); - let textured_region = (115u32, 15u32, 185u32, 105u32); - - let (flat, textured) = measure_diff_two_regions( - &client, - w, - h, - &clean, - sigma_pre, - |clean, seed| noisy_field_over(clean, w, h, sigma_pre, seed), - search_radius, - 0, - patch_radius, - flat_region, - textured_region, - ); - - eprintln!( - "within-frame flat vs textured region (w={w} h={h} sigma_pre={sigma_pre} \ - search_radius={search_radius} patch_radius={patch_radius}):\n\ - {:<10} {:>8} {:>8} {:>8}", - "region", "rho_h", "rho_v", "sig_ratio" - ); - for (label, s) in [("flat", &flat), ("textured", &textured)] { - eprintln!( - "{:<10} {:>8.4} {:>8.4} {:>8.4}", - label, s.rho_out_h, s.rho_out_v, s.sigma_ratio - ); - } - eprintln!( - "within-frame gap: rho_h {:.4}, rho_v {:.4}", - (flat.rho_out_h - textured.rho_out_h).abs(), - (flat.rho_out_v - textured.rho_out_v).abs() - ); - - // Validity guard on both regions independently. - assert!( - flat.sigma_ratio < 0.9, - "flat region sigma_ratio={:.4} too close to 1.0, no real smoothing happened", - flat.sigma_ratio - ); - assert!( - textured.sigma_ratio < 0.9, - "textured region sigma_ratio={:.4} too close to 1.0, no real smoothing happened", - textured.sigma_ratio - ); - - // The finding this test exists to pin down: within one frame, the - // two regions' residual correlations stay close, matching the - // separate-frames sweep above rather than contradicting it. The - // 0.1 tolerance is generous relative to the 0.08 used for the - // separate-frame comparison, since each tile here covers far fewer - // pixels (70x90) than a full 160x160 frame, so sampling noise in - // the correlation estimate itself is larger. - assert!( - (flat.rho_out_h - textured.rho_out_h).abs() < 0.1, - "within one frame, flat region rho_out_h={:.4} and textured region rho_out_h={:.4} \ - diverge by more than the tolerance; a single per-frame correlation profile would not \ - be structurally sound if this fails", - flat.rho_out_h, - textured.rho_out_h - ); -} diff --git a/av-denoise-core/src/nlmeans/tests/residual_correlation/mod.rs b/av-denoise-core/src/nlmeans/tests/residual_correlation/mod.rs new file mode 100644 index 0000000..476aa17 --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/residual_correlation/mod.rs @@ -0,0 +1,312 @@ +mod parameters; +mod regions; +mod search_radius; + +use cubecl::prelude::ComputeClient; + +use super::helpers::*; +use crate::bench_api::HostIo; +use crate::nlmeans::*; + +/// Residual correlation and smoothing measured for one configuration. +struct Sample { + rho_out_h: f64, + rho_out_v: f64, + /// The residual's standard deviation over the input noise's. + sigma_ratio: f64, +} + +/// Pearson correlation between each pixel and its right neighbour, within `left..right` by +/// `top..bottom`. +/// +/// The last column is skipped so no pixel is paired with a clamped edge copy of itself. +fn lag1_horizontal_rect(field: &[f32], width: u32, left: u32, top: u32, right: u32, bottom: u32) -> f64 { + let mut sum_current = 0.0f64; + let mut sum_next = 0.0f64; + let mut sum_product = 0.0f64; + let mut sum_current_sq = 0.0f64; + let mut sum_next_sq = 0.0f64; + let mut count = 0.0f64; + for y in top..bottom { + for x in left..(right - 1) { + let current = field[(y * width + x) as usize] as f64; + let next = field[(y * width + x + 1) as usize] as f64; + sum_current += current; + sum_next += next; + sum_product += current * next; + sum_current_sq += current * current; + sum_next_sq += next * next; + count += 1.0; + } + } + + let mean_current = sum_current / count; + let mean_next = sum_next / count; + let covariance = sum_product / count - mean_current * mean_next; + let variance_current = sum_current_sq / count - mean_current * mean_current; + let variance_next = sum_next_sq / count - mean_next * mean_next; + covariance / (variance_current.sqrt() * variance_next.sqrt()) +} + +/// Pearson correlation between each pixel and the one below it, within `left..right` by +/// `top..bottom`, skipping the last row. +fn lag1_vertical_rect(field: &[f32], width: u32, left: u32, top: u32, right: u32, bottom: u32) -> f64 { + let mut sum_current = 0.0f64; + let mut sum_next = 0.0f64; + let mut sum_product = 0.0f64; + let mut sum_current_sq = 0.0f64; + let mut sum_next_sq = 0.0f64; + let mut count = 0.0f64; + for y in top..(bottom - 1) { + for x in left..right { + let current = field[(y * width + x) as usize] as f64; + let next = field[((y + 1) * width + x) as usize] as f64; + sum_current += current; + sum_next += next; + sum_product += current * next; + sum_current_sq += current * current; + sum_next_sq += next * next; + count += 1.0; + } + } + + let mean_current = sum_current / count; + let mean_next = sum_next / count; + let covariance = sum_product / count - mean_current * mean_next; + let variance_current = sum_current_sq / count - mean_current * mean_current; + let variance_next = sum_next_sq / count - mean_next * mean_next; + covariance / (variance_current.sqrt() * variance_next.sqrt()) +} + +fn lag1_horizontal(field: &[f32], width: u32, height: u32) -> f64 { + lag1_horizontal_rect(field, width, 0, 0, width, height) +} + +fn lag1_vertical(field: &[f32], width: u32, height: u32) -> f64 { + lag1_vertical_rect(field, width, 0, 0, width, height) +} + +fn std_dev(field: &[f32]) -> f64 { + let count = field.len() as f64; + let mean: f64 = field.iter().map(|&value| value as f64).sum::() / count; + let variance: f64 = field + .iter() + .map(|&value| (value as f64 - mean).powi(2)) + .sum::() + / count; + variance.sqrt() +} + +fn std_dev_rect(field: &[f32], width: u32, left: u32, top: u32, right: u32, bottom: u32) -> f64 { + let count = ((right - left) * (bottom - top)) as f64; + + let mut sum = 0.0f64; + for y in top..bottom { + for x in left..right { + sum += field[(y * width + x) as usize] as f64; + } + } + + let mean = sum / count; + + let mut variance = 0.0f64; + for y in top..bottom { + for x in left..right { + let value = field[(y * width + x) as usize] as f64; + variance += (value - mean).powi(2); + } + } + + (variance / count).sqrt() +} + +fn subtract(minuend: &[f32], subtrahend: &[f32]) -> Vec { + minuend + .iter() + .zip(subtrahend.iter()) + .map(|(&left, &right)| left - right) + .collect() +} + +/// Pushes `2 * temporal_radius + 1` noisy frames seeded from `seed_base` through the front end. +/// +/// It returns the output of the final push, the first whose window holds no priming duplicate, +/// together with the noisy centre frame that output is built on. +#[expect( + clippy::too_many_arguments, + reason = "the test helper takes the full set of parameters its cases vary" +)] +fn run_front_end( + client: &ComputeClient, + width: u32, + height: u32, + clean: &[f32], + sigma: f32, + make_noise: impl Fn(&[f32], u32) -> Vec, + search_radius: u32, + temporal_radius: u32, + patch_radius: u32, + seed_base: u32, +) -> (Vec, Vec) { + let params = NlmParams { + temporal_radius, + search_radius, + patch_radius, + strength: 1.2, + self_weight: 1.0, + channels: ChannelMode::Luma, + prefilter: PrefilterMode::None, + motion_compensation: MotionCompensationMode::None, + hq: Some(HqParams::with_sigma(sigma)), + }; + let mut denoiser = NlmDenoiser::::new(client, params, width, height); + + let push_count = 2 * temporal_radius + 1; + let mut center_noisy: Option> = None; + let mut output: Option> = None; + for i in 0..push_count { + let frame = make_noise(clean, seed_base + i); + if i == temporal_radius { + center_noisy = Some(frame.clone()); + } + + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap(); + if i == push_count - 1 { + output = result; + } + } + + ( + output.expect("a fully real window must emit on its final push"), + center_noisy.expect("center frame must have been pushed"), + ) +} + +/// Measures residual correlation against a flat clean reference. +/// +/// A flat field gives NLM no structure to respond to, so `output - clean` is a noise-only residual. +#[expect( + clippy::too_many_arguments, + reason = "the test helper takes the full set of parameters its cases vary" +)] +fn measure_flat( + client: &ComputeClient, + width: u32, + height: u32, + clean: &[f32], + sigma: f32, + make_noise: impl Fn(&[f32], u32) -> Vec, + search_radius: u32, + temporal_radius: u32, + patch_radius: u32, +) -> Sample { + let (output, center_noisy) = run_front_end( + client, + width, + height, + clean, + sigma, + make_noise, + search_radius, + temporal_radius, + patch_radius, + 100, + ); + let residual = subtract(&output, clean); + let input_noise = subtract(¢er_noisy, clean); + + Sample { + rho_out_h: lag1_horizontal(&residual, width, height), + rho_out_v: lag1_vertical(&residual, width, height), + sigma_ratio: std_dev(&residual) / std_dev(&input_noise), + } +} + +/// Denoises two independent noise realisations of `clean` and returns their difference, plus the +/// first realisation's input noise. +/// +/// On textured content `output - clean` carries NLM's deterministic response to structure, which +/// cancels in the difference. Two independent, identically distributed outputs differ with the +/// same correlation and `sqrt(2)` times the standard deviation. The seed bases 100 and 500 never +/// overlap for windows of up to 9 frames. +#[expect( + clippy::too_many_arguments, + reason = "the test helper takes the full set of parameters its cases vary" +)] +fn denoised_difference( + client: &ComputeClient, + width: u32, + height: u32, + clean: &[f32], + sigma: f32, + make_noise: impl Fn(&[f32], u32) -> Vec, + search_radius: u32, + temporal_radius: u32, + patch_radius: u32, +) -> (Vec, Vec) { + let (output_a, noisy_a) = run_front_end( + client, + width, + height, + clean, + sigma, + &make_noise, + search_radius, + temporal_radius, + patch_radius, + 100, + ); + let (output_b, _noisy_b) = run_front_end( + client, + width, + height, + clean, + sigma, + &make_noise, + search_radius, + temporal_radius, + patch_radius, + 500, + ); + + let difference = subtract(&output_a, &output_b); + let input_noise = subtract(&noisy_a, clean); + (difference, input_noise) +} + +/// Measures residual correlation by differencing two denoised realisations of `clean`. +#[expect( + clippy::too_many_arguments, + reason = "the test helper takes the full set of parameters its cases vary" +)] +fn measure_diff( + client: &ComputeClient, + width: u32, + height: u32, + clean: &[f32], + sigma: f32, + make_noise: impl Fn(&[f32], u32) -> Vec, + search_radius: u32, + temporal_radius: u32, + patch_radius: u32, +) -> Sample { + let (difference, input_noise) = denoised_difference( + client, + width, + height, + clean, + sigma, + make_noise, + search_radius, + temporal_radius, + patch_radius, + ); + let sqrt2 = 2.0f64.sqrt(); + + Sample { + rho_out_h: lag1_horizontal(&difference, width, height), + rho_out_v: lag1_vertical(&difference, width, height), + sigma_ratio: (std_dev(&difference) / sqrt2) / std_dev(&input_noise), + } +} diff --git a/av-denoise-core/src/nlmeans/tests/residual_correlation/parameters.rs b/av-denoise-core/src/nlmeans/tests/residual_correlation/parameters.rs new file mode 100644 index 0000000..9838b5e --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/residual_correlation/parameters.rs @@ -0,0 +1,315 @@ +use super::{Sample, measure_flat}; +use crate::nlmeans::tests::helpers::*; + +/// The standard deviation an `a, 1 - 2a, a` horizontal blur leaves on white noise of `sigma_pre`. +fn tap_sigma(sigma_pre: f32, tap: f32) -> f32 { + let centre_tap = 1.0 - 2.0 * tap; + sigma_pre * (2.0 * tap * tap + centre_tap * centre_tap).sqrt() +} + +/// Overlapping search windows make the residual more correlated than the input grain. +/// +/// A search radius of 0 removes the overlap and isolates that mechanism. Every input is a +/// horizontal-only blur, so its vertical correlation is near 0. +#[test] +fn nlm_residual_correlation_exceeds_input_correlation_and_tracks_window_overlap() { + let client = make_client(); + let width = 160; + let height = 160; + let base = 0.5f32; + let sigma_pre = 0.06f32; + + struct Case { + label: &'static str, + rho_in_h: f64, + rho_in_v: f64, + } + let cases = [ + Case { + label: "rho_in=0.00", + rho_in_h: 0.0, + rho_in_v: 0.0, + }, + Case { + label: "rho_in=0.32", + rho_in_h: 0.316, + rho_in_v: 0.0, + }, + Case { + label: "rho_in=0.67", + rho_in_h: 2.0 / 3.0, + rho_in_v: 0.0, + }, + ]; + + let configs = [(2u32, 0u32), (2u32, 2u32), (0u32, 0u32), (0u32, 2u32)]; + + eprintln!( + "residual correlation sweep (w={width} h={height} sigma_pre={sigma_pre}):\n\ + {:<12} {:>6} {:>6} | {:>6} {:>6} | {:>8} {:>8} | {:>8}", + "input", "R_s", "R_t", "rho_in_h", "rho_in_v", "rho_out_h", "rho_out_v", "sig_ratio" + ); + + struct Row { + label: &'static str, + search_radius: u32, + temporal_radius: u32, + rho_in_h: f64, + measurement: Sample, + } + let mut rows = Vec::new(); + + let clean = vec![base; (width * height) as usize]; + for case in &cases { + for &(search_radius, temporal_radius) in &configs { + let measurement = match case.label { + "rho_in=0.00" => { + let sigma = tap_sigma(sigma_pre, 0.0); + measure_flat( + &client, + width, + height, + &clean, + sigma, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + temporal_radius, + 2, + ) + }, + "rho_in=0.32" => { + let sigma = tap_sigma(sigma_pre, 0.125); + measure_flat( + &client, + width, + height, + &clean, + sigma, + |_clean, seed| { + correlated_noisy_frame_with_tap(width, height, base, sigma_pre, seed, 0.125) + }, + search_radius, + temporal_radius, + 2, + ) + }, + _ => { + let sigma = tap_sigma(sigma_pre, 0.25); + measure_flat( + &client, + width, + height, + &clean, + sigma, + |_clean, seed| correlated_noisy_frame(width, height, base, sigma_pre, seed), + search_radius, + temporal_radius, + 2, + ) + }, + }; + + eprintln!( + "{:<12} {:>6} {:>6} | {:>8.4} {:>8.4} | {:>8.4} {:>8.4} | {:>8.4}", + case.label, + search_radius, + temporal_radius, + case.rho_in_h, + case.rho_in_v, + measurement.rho_out_h, + measurement.rho_out_v, + measurement.sigma_ratio + ); + + rows.push(Row { + label: case.label, + search_radius, + temporal_radius, + rho_in_h: case.rho_in_h, + measurement, + }); + } + } + + // At search radius 2, window overlap alone raises the residual's horizontal correlation above + // the input's. + for row in &rows { + if row.search_radius == 2 { + assert!( + row.measurement.rho_out_h > row.rho_in_h, + "{} at search_radius=2 temporal_radius={}: residual rho_h={:.4} did not exceed \ + input rho_h={:.4}, expected window overlap to raise it", + row.label, + row.temporal_radius, + row.measurement.rho_out_h, + row.rho_in_h + ); + } + } + + // At search radius 0 there is no window overlap, so the residual tracks the input. The 0.08 + // tolerance covers the real shift of up to about 0.025 that temporal radius 2 adds, not + // measurement noise, since every generator is seeded. + for row in &rows { + if row.search_radius == 0 { + assert!( + (row.measurement.rho_out_h - row.rho_in_h).abs() < 0.08, + "{} at search_radius=0 temporal_radius={}: residual rho_h={:.4} should track the \ + input's rho_h={:.4} within 0.08 once spatial window overlap is removed", + row.label, + row.temporal_radius, + row.measurement.rho_out_h, + row.rho_in_h + ); + } + } + + // With no spatial window, nothing can create vertical correlation the input never had. + for row in &rows { + if row.search_radius == 0 { + assert!( + row.measurement.rho_out_v.abs() < 0.05, + "{} at search_radius=0 temporal_radius={}: residual rho_v={:.4} should stay near \ + zero, the input never carried vertical correlation and there is no spatial \ + window to manufacture it", + row.label, + row.temporal_radius, + row.measurement.rho_out_v + ); + } + } + + // The square search window mixes vertical neighbours too, so it creates strong vertical + // correlation even from horizontal-only grain. + for row in &rows { + if row.search_radius == 2 { + assert!( + row.measurement.rho_out_v > 0.3, + "{} at search_radius=2 temporal_radius={}: residual rho_v={:.4} did not rise \ + well above the input's zero vertical correlation, expected the isotropic \ + window to induce substantial vertical correlation regardless", + row.label, + row.temporal_radius, + row.measurement.rho_out_v + ); + } + } +} + +/// Varies patch radius, temporal radius and input correlation one at a time from a baseline of +/// search radius 2, patch radius 2, temporal radius 0 and uncorrelated flat input. +#[test] +fn nlm_residual_correlation_patch_radius_temporal_radius_and_input_correlation() { + let client = make_client(); + let width = 160; + let height = 160; + let base = 0.5f32; + let sigma_pre = 0.06f32; + let flat_clean = vec![base; (width * height) as usize]; + let search_radius = 2; + + let baseline = measure_flat( + &client, + width, + height, + &flat_clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 0, + 2, + ); + + let patch_radius_4 = measure_flat( + &client, + width, + height, + &flat_clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 0, + 4, + ); + + let temporal_radius_2 = measure_flat( + &client, + width, + height, + &flat_clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 2, + 2, + ); + + let rho_in_h = 2.0 / 3.0; + let correlated_input = measure_flat( + &client, + width, + height, + &flat_clean, + sigma_pre, + |_clean, seed| correlated_noisy_frame(width, height, base, sigma_pre, seed), + search_radius, + 0, + 2, + ); + + eprintln!( + "secondary checks at search_radius=2, flat content (w={width} h={height} sigma_pre={sigma_pre}):\n\ + {:<32} {:>8} {:>8} {:>8}", + "config", "rho_h", "rho_v", "sig_ratio" + ); + for (label, sample) in [ + ("baseline patch_radius=2 R_t=0 rho_in=0", &baseline), + ("patch_radius=4", &patch_radius_4), + ("temporal_radius=2", &temporal_radius_2), + ("input rho_h=0.67", &correlated_input), + ] { + eprintln!( + "{:<32} {:>8.4} {:>8.4} {:>8.4}", + label, sample.rho_out_h, sample.rho_out_v, sample.sigma_ratio + ); + } + + for (label, sample) in [ + ("baseline", &baseline), + ("patch_radius=4", &patch_radius_4), + ("temporal_radius=2", &temporal_radius_2), + ("input rho_h=0.67", &correlated_input), + ] { + assert!( + sample.sigma_ratio < 0.9, + "{label}: sigma_ratio={:.4} too close to 1.0, no real smoothing happened, this \ + configuration's correlation number is not valid data", + sample.sigma_ratio + ); + } + + // Temporal averaging dilutes the window-induced correlation slightly, by about 0.02 to 0.05 at + // search radius 2. + assert!( + temporal_radius_2.rho_out_h < baseline.rho_out_h, + "temporal_radius=2 rho_out_h={:.4} should be below temporal_radius=0's {:.4}, temporal \ + averaging is expected to dilute the spatial window's contribution", + temporal_radius_2.rho_out_h, + baseline.rho_out_h + ); + + // The window sets a floor and correlated input lifts the residual further, within the headroom + // below 1. + assert!( + correlated_input.rho_out_h > baseline.rho_out_h, + "correlated input (rho_in_h={rho_in_h:.4}) rho_out_h={:.4} should exceed the uncorrelated \ + baseline's {:.4}", + correlated_input.rho_out_h, + baseline.rho_out_h + ); + assert!( + correlated_input.rho_out_h < 1.0, + "correlated input rho_out_h={:.4} must stay below 1.0", + correlated_input.rho_out_h + ); +} diff --git a/av-denoise-core/src/nlmeans/tests/residual_correlation/regions.rs b/av-denoise-core/src/nlmeans/tests/residual_correlation/regions.rs new file mode 100644 index 0000000..96b450a --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/residual_correlation/regions.rs @@ -0,0 +1,151 @@ +use cubecl::prelude::ComputeClient; + +use super::{Sample, denoised_difference, lag1_horizontal_rect, lag1_vertical_rect, std_dev_rect}; +use crate::nlmeans::tests::helpers::*; + +/// A rectangle as `(left, top, right, bottom)`, covering `left..right` by `top..bottom`. +type Region = (u32, u32, u32, u32); + +/// Builds a frame that is flat left of `split_x` and carries +/// [make_textured_frame](crate::nlmeans::tests::helpers::make_textured_frame)'s sine pattern to +/// the right. +fn make_flat_and_textured_frame(width: u32, height: u32, split_x: u32) -> Vec { + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let value = if x < split_x { + 0.5 + } else { + let fraction_x = x as f32 / width as f32; + let fraction_y = y as f32 / height as f32; + let raw = 0.5 + + 0.2 + * (fraction_x * 8.0 * std::f32::consts::PI).sin() + * (fraction_y * 6.0 * std::f32::consts::PI).cos() + + 0.1 * (fraction_x * 20.0 * std::f32::consts::PI).sin(); + raw.clamp(0.05, 0.95) + }; + frame[(y * width + x) as usize] = value; + } + } + frame +} + +/// Measures residual correlation by differencing two denoised realisations, separately over two +/// regions of the frame. +#[expect( + clippy::too_many_arguments, + reason = "the test helper takes the full set of parameters its cases vary" +)] +fn measure_diff_two_regions( + client: &ComputeClient, + width: u32, + height: u32, + clean: &[f32], + sigma: f32, + make_noise: impl Fn(&[f32], u32) -> Vec, + search_radius: u32, + temporal_radius: u32, + patch_radius: u32, + region_a: Region, + region_b: Region, +) -> (Sample, Sample) { + let (difference, input_noise) = denoised_difference( + client, + width, + height, + clean, + sigma, + make_noise, + search_radius, + temporal_radius, + patch_radius, + ); + let sqrt2 = 2.0f64.sqrt(); + + let sample_for = |(left, top, right, bottom): Region| { + let difference_std = std_dev_rect(&difference, width, left, top, right, bottom); + let input_std = std_dev_rect(&input_noise, width, left, top, right, bottom); + Sample { + rho_out_h: lag1_horizontal_rect(&difference, width, left, top, right, bottom), + rho_out_v: lag1_vertical_rect(&difference, width, left, top, right, bottom), + sigma_ratio: (difference_std / sqrt2) / input_std, + } + }; + + (sample_for(region_a), sample_for(region_b)) +} + +/// Separate frames cannot rule out other differences between two setups, so this places flat and +/// textured regions side by side in one frame. +/// +/// Both tiles sit 15 pixels from the seam and every edge, outside the reach of +/// `patch_radius + search_radius` (4). +#[test] +fn nlm_residual_correlation_within_a_single_frame_flat_vs_textured_regions() { + let client = make_client(); + let width = 200; + let height = 120; + let sigma_pre = 0.06f32; + let search_radius = 2; + let patch_radius = 2; + let split_x = 100; + + let clean = make_flat_and_textured_frame(width, height, split_x); + let flat_region = (15u32, 15u32, 85u32, 105u32); + let textured_region = (115u32, 15u32, 185u32, 105u32); + + let (flat, textured) = measure_diff_two_regions( + &client, + width, + height, + &clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 0, + patch_radius, + flat_region, + textured_region, + ); + + eprintln!( + "within-frame flat vs textured region (w={width} h={height} sigma_pre={sigma_pre} \ + search_radius={search_radius} patch_radius={patch_radius}):\n\ + {:<10} {:>8} {:>8} {:>8}", + "region", "rho_h", "rho_v", "sig_ratio" + ); + for (label, sample) in [("flat", &flat), ("textured", &textured)] { + eprintln!( + "{:<10} {:>8.4} {:>8.4} {:>8.4}", + label, sample.rho_out_h, sample.rho_out_v, sample.sigma_ratio + ); + } + eprintln!( + "within-frame gap: rho_h {:.4}, rho_v {:.4}", + (flat.rho_out_h - textured.rho_out_h).abs(), + (flat.rho_out_v - textured.rho_out_v).abs() + ); + + assert!( + flat.sigma_ratio < 0.9, + "flat region sigma_ratio={:.4} too close to 1.0, no real smoothing happened", + flat.sigma_ratio + ); + assert!( + textured.sigma_ratio < 0.9, + "textured region sigma_ratio={:.4} too close to 1.0, no real smoothing happened", + textured.sigma_ratio + ); + + // Each 70x90 tile has far fewer pixels than a 160x160 frame, so sampling noise is larger and + // the tolerance is looser than the 0.08 used across separate frames. + assert!( + (flat.rho_out_h - textured.rho_out_h).abs() < 0.1, + "within one frame, flat region rho_out_h={:.4} and textured region rho_out_h={:.4} \ + diverge by more than the tolerance; a single per-frame correlation profile would not \ + be structurally sound if this fails", + flat.rho_out_h, + textured.rho_out_h + ); +} diff --git a/av-denoise-core/src/nlmeans/tests/residual_correlation/search_radius.rs b/av-denoise-core/src/nlmeans/tests/residual_correlation/search_radius.rs new file mode 100644 index 0000000..d63ba69 --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/residual_correlation/search_radius.rs @@ -0,0 +1,276 @@ +use super::{Sample, measure_diff, measure_flat}; +use crate::nlmeans::tests::helpers::*; + +/// On texture the Welsch weighting suppresses candidates whose patch does not match, which could +/// lower the induced correlation, so textured content is measured alongside flat. +#[test] +fn nlm_residual_correlation_search_radius_sweep_flat_vs_textured() { + let client = make_client(); + let width = 160; + let height = 160; + let base = 0.5f32; + let sigma_pre = 0.06f32; + let patch_radius = 2; + + let flat_clean = vec![base; (width * height) as usize]; + let textured_clean = make_textured_frame(width, height); + + struct Row { + search_radius: u32, + flat: Sample, + textured: Sample, + } + let mut rows = Vec::new(); + + eprintln!( + "search radius sweep, flat vs textured (w={width} h={height} sigma_pre={sigma_pre} \ + patch_radius={patch_radius}, uncorrelated input):\n\ + {:>3} | {:>8} {:>8} {:>8} | {:>8} {:>8} {:>8}", + "R_s", "flat_h", "flat_v", "flat_sig", "tex_h", "tex_v", "tex_sig" + ); + + for &search_radius in &[1u32, 2, 3, 4] { + let flat = measure_flat( + &client, + width, + height, + &flat_clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 0, + patch_radius, + ); + let textured = measure_diff( + &client, + width, + height, + &textured_clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 0, + patch_radius, + ); + + eprintln!( + "{:>3} | {:>8.4} {:>8.4} {:>8.4} | {:>8.4} {:>8.4} {:>8.4}", + search_radius, + flat.rho_out_h, + flat.rho_out_v, + flat.sigma_ratio, + textured.rho_out_h, + textured.rho_out_v, + textured.sigma_ratio + ); + + rows.push(Row { + search_radius, + flat, + textured, + }); + } + + // A sigma ratio near 1.0 means almost no smoothing happened, so the correlation is not valid + // data. + for row in &rows { + assert!( + row.flat.sigma_ratio < 0.9, + "search_radius={}: flat sigma_ratio={:.4} too close to 1.0, no real smoothing \ + happened, this configuration's correlation number is not valid data", + row.search_radius, + row.flat.sigma_ratio + ); + assert!( + row.textured.sigma_ratio < 0.9, + "search_radius={}: textured sigma_ratio={:.4} too close to 1.0, no real smoothing \ + happened, this configuration's correlation number is not valid data", + row.search_radius, + row.textured.sigma_ratio + ); + } + + // Window overlap induces real correlation on flat content at every shipping search radius. + for row in &rows { + assert!( + row.flat.rho_out_h > 0.3, + "search_radius={}: flat rho_out_h={:.4} did not show substantial window-induced \ + correlation", + row.search_radius, + row.flat.rho_out_h + ); + } + + // The sinusoidal texture tracks the flat field closely at every radius, with a measured + // maximum gap of 0.0234 at search radius 4. The 0.08 tolerance leaves margin above that. + for row in &rows { + assert!( + (row.flat.rho_out_h - row.textured.rho_out_h).abs() < 0.08, + "search_radius={}: flat rho_out_h={:.4} and textured rho_out_h={:.4} diverge by more \ + than the measured tolerance, flat and textured no longer agree", + row.search_radius, + row.flat.rho_out_h, + row.textured.rho_out_h + ); + } + + // A wider window only adds overlapping candidates, so the induced correlation rises or + // plateaus. + for pair in rows.windows(2) { + assert!( + pair[1].flat.rho_out_h >= pair[0].flat.rho_out_h - 1e-6, + "flat rho_out_h dropped from search_radius={} ({:.4}) to search_radius={} ({:.4})", + pair[0].search_radius, + pair[0].flat.rho_out_h, + pair[1].search_radius, + pair[1].flat.rho_out_h + ); + } +} + +/// The patch radius shifts the result noticeably, with 0.7770 at the default of 4 against 0.7006 +/// at 2 for search radius 2. Search radius 0 is included as the case with no window overlap. +#[test] +fn nlm_residual_correlation_search_radius_sweep_at_shipped_patch_radius() { + let client = make_client(); + let width = 160; + let height = 160; + let base = 0.5f32; + let sigma_pre = 0.06f32; + let patch_radius = 4; + + let flat_clean = vec![base; (width * height) as usize]; + let textured_clean = make_textured_frame(width, height); + + struct Row { + search_radius: u32, + flat: Sample, + textured: Sample, + } + let mut rows = Vec::new(); + + eprintln!( + "search radius sweep at shipped patch_radius={patch_radius}, flat vs textured \ + (w={width} h={height} sigma_pre={sigma_pre}, uncorrelated input):\n\ + {:>3} | {:>8} {:>8} {:>8} | {:>8} {:>8} {:>8}", + "R_s", "flat_h", "flat_v", "flat_sig", "tex_h", "tex_v", "tex_sig" + ); + + for &search_radius in &[0u32, 1, 2, 3, 4] { + let flat = measure_flat( + &client, + width, + height, + &flat_clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 0, + patch_radius, + ); + let textured = measure_diff( + &client, + width, + height, + &textured_clean, + sigma_pre, + |clean, seed| noisy_field_over(clean, width, height, sigma_pre, seed), + search_radius, + 0, + patch_radius, + ); + + eprintln!( + "{:>3} | {:>8.4} {:>8.4} {:>8.4} | {:>8.4} {:>8.4} {:>8.4}", + search_radius, + flat.rho_out_h, + flat.rho_out_v, + flat.sigma_ratio, + textured.rho_out_h, + textured.rho_out_v, + textured.sigma_ratio + ); + + rows.push(Row { + search_radius, + flat, + textured, + }); + } + + // Search radius 0 is exempt, because with no spatial window almost no smoothing happens and + // its sigma ratio sits near 1.0 by construction. + for row in &rows { + if row.search_radius == 0 { + continue; + } + + assert!( + row.flat.sigma_ratio < 0.9, + "search_radius={}: flat sigma_ratio={:.4} too close to 1.0, no real smoothing \ + happened, this configuration's correlation number is not valid data", + row.search_radius, + row.flat.sigma_ratio + ); + assert!( + row.textured.sigma_ratio < 0.9, + "search_radius={}: textured sigma_ratio={:.4} too close to 1.0, no real smoothing \ + happened, this configuration's correlation number is not valid data", + row.search_radius, + row.textured.sigma_ratio + ); + } + + // With no window to overlap, the residual stays near the uncorrelated input's 0. + for row in &rows { + if row.search_radius == 0 { + assert!( + row.flat.rho_out_h.abs() < 0.1, + "search_radius=0: flat rho_out_h={:.4} should stay near zero, there is no \ + spatial window to manufacture correlation", + row.flat.rho_out_h + ); + } + } + + // Flat content shows substantial window-induced correlation at every radius with a window. + for row in &rows { + if row.search_radius >= 1 { + assert!( + row.flat.rho_out_h > 0.3, + "search_radius={}: flat rho_out_h={:.4} did not show substantial window-induced \ + correlation", + row.search_radius, + row.flat.rho_out_h + ); + } + } + + // A larger patch makes each candidate's weight less noisy, which lets the window's + // near-uniform weighting through and raises the correlation above the patch radius 2 values. + for row in &rows { + if row.search_radius == 0 { + continue; + } + + assert!( + row.flat.rho_out_h > 0.6, + "search_radius={}: patch_radius=4 flat rho_out_h={:.4} unexpectedly low, expected it \ + to sit above the patch_radius=2 sweep's own values at this radius", + row.search_radius, + row.flat.rho_out_h + ); + } + + // Residual correlation should not fall as the window widens. + for pair in rows.windows(2) { + assert!( + pair[1].flat.rho_out_h >= pair[0].flat.rho_out_h - 1e-6, + "flat rho_out_h dropped from search_radius={} ({:.4}) to search_radius={} ({:.4})", + pair[0].search_radius, + pair[0].flat.rho_out_h, + pair[1].search_radius, + pair[1].flat.rho_out_h + ); + } +} diff --git a/av-denoise-core/src/nlmeans/tests/separable.rs b/av-denoise-core/src/nlmeans/tests/separable.rs index a4d3174..ba33fea 100644 --- a/av-denoise-core/src/nlmeans/tests/separable.rs +++ b/av-denoise-core/src/nlmeans/tests/separable.rs @@ -1,4 +1,5 @@ use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; #[test] @@ -16,26 +17,20 @@ fn separable_uniform_passthrough() { hq: None, }; - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 1, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); assert!(denoiser.use_separable, "should use separable for patch_radius=9"); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - for (i, &v) in result.iter().enumerate() { + for (i, &value) in result.iter().enumerate() { assert!( - (v - 0.5).abs() < 1e-4, - "separable: pixel {i}: expected 0.5, got {v}" + (value - 0.5).abs() < 1e-4, + "separable: pixel {i}: expected 0.5, got {value}" ); } } @@ -55,27 +50,21 @@ fn separable_yuv_passthrough() { hq: None, }; - let w = 32; - let h = 32; - let frame = make_uniform_frame(w, h, 3, 0.5); + let width = 32; + let height = 32; + let frame = make_uniform_frame(width, height, 3, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); assert!(denoiser.use_separable); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - assert_eq!(result.len(), (w * h * 3) as usize); + let result = denoiser.denoise().unwrap().unwrap(); + assert_eq!(result.len(), (width * height * 3) as usize); - for (i, &v) in result.iter().enumerate() { + for (i, &value) in result.iter().enumerate() { assert!( - (v - 0.5).abs() < 1e-4, - "separable yuv: pixel {i}: expected 0.5, got {v}" + (value - 0.5).abs() < 1e-4, + "separable yuv: pixel {i}: expected 0.5, got {value}" ); } } @@ -95,33 +84,27 @@ fn separable_symmetry_preserved() { hq: None, }; - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut frame = vec![0.5f32; (w * h) as usize]; - for y in 0..h { - for x in 0..(w / 2) { - let val = 0.3 + 0.4 * (x as f32 / w as f32); - frame[(y * w + x) as usize] = val; - frame[(y * w + (w - 1 - x)) as usize] = val; + let mut frame = vec![0.5f32; (width * height) as usize]; + for y in 0..height { + for x in 0..(width / 2) { + let value = 0.3 + 0.4 * (x as f32 / width as f32); + frame[(y * width + x) as usize] = value; + frame[(y * width + (width - 1 - x)) as usize] = value; } } - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - for y in 0..h { - for x in 0..(w / 2) { - let left = result[(y * w + x) as usize]; - let right = result[(y * w + (w - 1 - x)) as usize]; + let result = denoiser.denoise().unwrap().unwrap(); + + for y in 0..height { + for x in 0..(width / 2) { + let left = result[(y * width + x) as usize]; + let right = result[(y * width + (width - 1 - x)) as usize]; assert!( (left - right).abs() < 1e-4, "separable symmetry broken at ({x},{y}): \ diff --git a/av-denoise-core/src/nlmeans/tests/sizes.rs b/av-denoise-core/src/nlmeans/tests/sizes.rs new file mode 100644 index 0000000..282c0af --- /dev/null +++ b/av-denoise-core/src/nlmeans/tests/sizes.rs @@ -0,0 +1,90 @@ +use cubecl::server::Handle; + +use super::helpers::{R, make_client}; +use crate::nlmeans::denoiser::{BufferSize, front_buffer_sizes}; +use crate::nlmeans::*; + +// Odd dimensions leave ragged block grids and unaligned slots, so every padding rule is exercised. +const WIDTH: u32 = 101; +const HEIGHT: u32 = 67; + +fn allocated_elements(handle: Option<&Handle>) -> Option { + handle.map(|handle| handle.size_in_used() / 4) +} + +/// The front end's computed count for `name`, which must appear exactly once. +fn computed(buffers: &[BufferSize], name: &str) -> Option { + let mut matching = buffers.iter().filter(|(buffer, _)| *buffer == name); + let (_, elements) = matching.next()?; + assert!(matching.next().is_none(), "{name} listed twice"); + + *elements +} + +fn assert_sizes_match_allocations(params: NlmParams) { + let client = make_client(); + let sizes = front_buffer_sizes(&client, ¶ms, WIDTH, HEIGHT); + let denoiser = NlmDenoiser::::new(&client, params, WIDTH, HEIGHT); + + let pairs = [ + ("noise partials ring", denoiser.noise_partials.as_ref()), + ("temporal stats ring", denoiser.temporal_stats_buf.as_ref()), + ("motion field", denoiser.mv_field_buf.as_ref()), + ("motion pyramid ring", denoiser.pyramid_input.as_ref()), + ("motion pair ring", denoiser.pair_ring_buf.as_ref()), + ("confidence", denoiser.confidence_buf.as_ref()), + ("confidence pyramid ring", denoiser.confidence_pyramid.as_ref()), + ( + "confidence vector scratch", + denoiser.confidence_mv_scratch.as_ref(), + ), + ]; + for (name, handle) in pairs { + let allocated = allocated_elements(handle); + let expected = computed(&sizes.buffers, name); + assert_eq!(expected, allocated, "{name}"); + } + + for (name, _) in &sizes.buffers { + let known = pairs.iter().any(|(pair_name, _)| pair_name == name); + assert!(known, "{name} has no allocation to compare against"); + } +} + +fn hq_params(temporal_radius: u32, motion_compensation: MotionCompensationMode) -> NlmParams { + NlmParams { + temporal_radius, + motion_compensation, + hq: Some(HqParams::default()), + ..NlmParams::default() + } +} + +#[test] +fn chained_motion_sizes_match_the_allocations() { + let motion = MotionCompensationMode::Mvtools { + blksize: 16, + overlap: 8, + search_radius: 4, + pyramid_levels: 3, + estimation: MotionEstimation::chained_default(), + }; + let params = hq_params(3, motion); + + assert_sizes_match_allocations(params); +} + +#[test] +fn direct_motion_sizes_match_the_allocations() { + let motion = MotionCompensationMode::mvtools_default(); + let params = hq_params(1, motion); + + assert_sizes_match_allocations(params); +} + +#[test] +fn confidence_only_sizes_match_the_allocations() { + let params = hq_params(2, MotionCompensationMode::None); + + assert_sizes_match_allocations(params); +} diff --git a/av-denoise-core/src/nlmeans/tests/spatial.rs b/av-denoise-core/src/nlmeans/tests/spatial.rs index 45b6f0d..6f65c2c 100644 --- a/av-denoise-core/src/nlmeans/tests/spatial.rs +++ b/av-denoise-core/src/nlmeans/tests/spatial.rs @@ -1,4 +1,5 @@ use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; #[test] @@ -16,23 +17,17 @@ fn uniform_image_passthrough() { hq: None, }; - let w = 16; - let h = 16; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 16; + let height = 16; + let frame = make_uniform_frame(width, height, 1, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - for (i, &v) in result.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-5, "pixel {i}: expected 0.5, got {v}"); + for (i, &value) in result.iter().enumerate() { + assert!((value - 0.5).abs() < 1e-5, "pixel {i}: expected 0.5, got {value}"); } } @@ -51,24 +46,18 @@ fn uniform_yuv_passthrough() { hq: None, }; - let w = 16; - let h = 16; - let frame = make_uniform_frame(w, h, 3, 0.5); + let width = 16; + let height = 16; + let frame = make_uniform_frame(width, height, 3, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - assert_eq!(result.len(), (w * h * 3) as usize); - - for (i, &v) in result.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-5, "pixel {i}: expected 0.5, got {v}"); + let result = denoiser.denoise().unwrap().unwrap(); + assert_eq!(result.len(), (width * height * 3) as usize); + + for (i, &value) in result.iter().enumerate() { + assert!((value - 0.5).abs() < 1e-5, "pixel {i}: expected 0.5, got {value}"); } } @@ -87,24 +76,21 @@ fn uniform_chroma_passthrough() { hq: None, }; - let w = 16; - let h = 16; - let frame = make_uniform_frame(w, h, 2, 0.5); + let width = 16; + let height = 16; + let frame = make_uniform_frame(width, height, 2, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - assert_eq!(result.len(), (w * h * 2) as usize); - - for (i, &v) in result.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-5, "pixel {i}: expected ~0.5, got {v}"); + let result = denoiser.denoise().unwrap().unwrap(); + assert_eq!(result.len(), (width * height * 2) as usize); + + for (i, &value) in result.iter().enumerate() { + assert!( + (value - 0.5).abs() < 1e-5, + "pixel {i}: expected ~0.5, got {value}" + ); } } @@ -123,24 +109,18 @@ fn noisy_region_suppressed() { hq: None, }; - let w = 32; - let h = 32; - let mut frame = vec![0.5f32; (w * h) as usize]; - frame[(16 * w + 16) as usize] = 0.8; + let width = 32; + let height = 32; + let mut frame = vec![0.5f32; (width * height) as usize]; + frame[(16 * width + 16) as usize] = 0.8; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - let noisy_idx = (16 * w + 16) as usize; - let denoised = result[noisy_idx]; + let noisy_index = (16 * width + 16) as usize; + let denoised = result[noisy_index]; assert!( denoised < 0.8, @@ -163,29 +143,23 @@ fn high_strength_smooths_heavily() { hq: None, }; - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut frame = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - let val = if y % 2 == 0 { 0.3 } else { 0.7 }; - for x in 0..w { - frame[(y * w + x) as usize] = val; + let mut frame = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + let row_value = if y % 2 == 0 { 0.3 } else { 0.7 }; + for x in 0..width { + frame[(y * width + x) as usize] = row_value; } } - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - let center = result[(8 * w + 8) as usize]; + let center = result[(8 * width + 8) as usize]; assert!( (center - 0.5).abs() < 0.15, "high strength should smooth toward ~0.5, got {center}" @@ -207,24 +181,18 @@ fn low_strength_preserves_original() { hq: None, }; - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut frame = vec![0.5f32; (w * h) as usize]; - frame[(8 * w + 8) as usize] = 0.8; + let mut frame = vec![0.5f32; (width * height) as usize]; + frame[(8 * width + 8) as usize] = 0.8; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - let pixel = result[(8 * w + 8) as usize]; + let pixel = result[(8 * width + 8) as usize]; assert!( (pixel - 0.8).abs() < 0.05, "low strength should preserve original ~0.8, got {pixel}" @@ -246,24 +214,21 @@ fn self_weight_zero_uniform() { hq: None, }; - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let frame = make_uniform_frame(w, h, 1, 0.5); + let frame = make_uniform_frame(width, height, 1, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - for (i, &v) in result.iter().enumerate() { - assert!((v - 0.5).abs() < 1e-5, "pixel {i}: expected ~0.5, got {v}"); + for (i, &value) in result.iter().enumerate() { + assert!( + (value - 0.5).abs() < 1e-5, + "pixel {i}: expected ~0.5, got {value}" + ); } } @@ -275,11 +240,11 @@ fn spatial_only_no_delay() { ..NlmParams::default() }; - let w = 8; - let h = 8; - let frame = make_uniform_frame(w, h, 3, 0.5); + let width = 8; + let height = 8; + let frame = make_uniform_frame(width, height, 3, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); let result = denoiser.denoise().unwrap(); @@ -301,33 +266,27 @@ fn symmetry_preserved() { hq: None, }; - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut frame = vec![0.5f32; (w * h) as usize]; - for y in 0..h { - for x in 0..(w / 2) { - let val = 0.3 + 0.4 * (x as f32 / w as f32); - frame[(y * w + x) as usize] = val; - frame[(y * w + (w - 1 - x)) as usize] = val; + let mut frame = vec![0.5f32; (width * height) as usize]; + for y in 0..height { + for x in 0..(width / 2) { + let value = 0.3 + 0.4 * (x as f32 / width as f32); + frame[(y * width + x) as usize] = value; + frame[(y * width + (width - 1 - x)) as usize] = value; } } - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - for y in 0..h { - for x in 0..(w / 2) { - let left = result[(y * w + x) as usize]; - let right = result[(y * w + (w - 1 - x)) as usize]; + let result = denoiser.denoise().unwrap().unwrap(); + + for y in 0..height { + for x in 0..(width / 2) { + let left = result[(y * width + x) as usize]; + let right = result[(y * width + (width - 1 - x)) as usize]; assert!( (left - right).abs() < 1e-5, "symmetry broken at ({x},{y}): \ @@ -352,20 +311,14 @@ fn clamp_to_edge_no_darkening() { hq: None, }; - let w = 8; - let h = 8; - let frame = make_uniform_frame(w, h, 1, 0.7); + let width = 8; + let height = 8; + let frame = make_uniform_frame(width, height, 1, 0.7); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); let corner = result[0]; assert!( diff --git a/av-denoise-core/src/nlmeans/tests/spatial_offset.rs b/av-denoise-core/src/nlmeans/tests/spatial_offset.rs index 4dafadd..291cd16 100644 --- a/av-denoise-core/src/nlmeans/tests/spatial_offset.rs +++ b/av-denoise-core/src/nlmeans/tests/spatial_offset.rs @@ -1,11 +1,10 @@ use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; -/// The shared parameters for these tests, running spatially only with a -/// fixed sigma so the noise offset is nonzero and steady. +/// Spatial-only parameters with a fixed sigma, so the noise offset is nonzero and steady. /// -/// A fixed sigma also keeps automatic estimation inactive, so nothing -/// but the test itself ever touches the correlation state. +/// A fixed sigma also keeps automatic estimation off, so only the test touches the correlation state. fn attenuation_params(sigma: f32) -> NlmParams { NlmParams { temporal_radius: 0, @@ -28,57 +27,43 @@ fn attenuation_params(sigma: f32) -> NlmParams { } } -/// Setting the correlation state directly, as the estimator would after -/// folding a temporal sample, has to change the spatial weighting on -/// correlated content. -/// -/// Reducing the noise-floor offset for nearby candidates changes which -/// patch distances reach full weight, which changes the weighted -/// average. +/// Sets the correlation state directly, as the estimator would after folding a temporal sample. #[test] fn rho_attenuation_changes_spatial_weighting_on_correlated_content() { let client = make_client(); - let w = 48; - let h = 48; + let width = 48; + let height = 48; let sigma_marginal = 20.0 / 255.0; let sigma_pre = sigma_marginal / 0.375f32.sqrt(); - let frame = correlated_noisy_frame(w, h, 0.5, sigma_pre, 7); + let frame = correlated_noisy_frame(width, height, 0.5, sigma_pre, 7); let params = attenuation_params(sigma_marginal); - let mut rho_zero = NlmDenoiser::::new(&client, params.clone(), w, h); + let mut rho_zero = NlmDenoiser::::new(&client, params.clone(), width, height); rho_zero.push_frame(&frame); - let rho_zero_out = rho_zero - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let rho_zero_out = rho_zero.denoise().unwrap().unwrap(); - let mut rho_high = NlmDenoiser::::new(&client, params, w, h); + let mut rho_high = NlmDenoiser::::new(&client, params, width, height); rho_high.rho_smoothed = Some(0.65); rho_high.push_frame(&frame); - let rho_high_out = rho_high - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let rho_high_out = rho_high.denoise().unwrap().unwrap(); let mut max_diff = 0.0f32; - for (i, (&a, &b)) in rho_zero_out.iter().zip(rho_high_out.iter()).enumerate() { - assert!(a.is_finite() && b.is_finite(), "pixel {i}: non-finite output"); + let pairs = rho_zero_out.iter().zip(rho_high_out.iter()).enumerate(); + for (i, (&zero_value, &high_value)) in pairs { + assert!( + zero_value.is_finite() && high_value.is_finite(), + "pixel {i}: non-finite output" + ); assert!( - (0.0..=1.0).contains(&a), - "pixel {i}: rho=0 output out of range: {a}" + (0.0..=1.0).contains(&zero_value), + "pixel {i}: rho=0 output out of range: {zero_value}" ); assert!( - (0.0..=1.0).contains(&b), - "pixel {i}: rho=0.65 output out of range: {b}" + (0.0..=1.0).contains(&high_value), + "pixel {i}: rho=0.65 output out of range: {high_value}" ); - max_diff = max_diff.max((a - b).abs()); + max_diff = max_diff.max((zero_value - high_value).abs()); } assert!( @@ -87,65 +72,42 @@ fn rho_attenuation_changes_spatial_weighting_on_correlated_content() { ); } -/// The windowed and separable paths must agree once correlation is -/// taken into account. -/// -/// The windowed kernel reads each candidate's offset from the table, -/// while the separable path computes the same value on the host for -/// each dispatched candidate. Both have to arrive at the same factor. -/// -/// The comparison covers interior pixels only, keeping a margin clear of -/// every clamped read either path makes. -/// -/// Forcing the separable path mirrors the cross-check in -/// `temporal::windowed_vs_separable_psnr`. Unlike that one, this -/// compares pixels directly rather than through PSNR. +/// The windowed kernel reads each candidate's offset from the table, while the separable path +/// computes it on the host. /// -/// The two paths already handle clamped borders differently, which this -/// change did not touch, so the test also asserts the same agreement -/// with no correlation. That separates the existing difference from -/// anything the table wiring could have introduced. +/// Only interior pixels are compared, clear of every clamped read. The two paths already differ on +/// clamped borders, so rho 0 is checked too. That separates the border difference from anything +/// the offset table could introduce. #[test] fn windowed_and_separable_agree_under_rho_attenuation() { let client = make_client(); - let w = 48; - let h = 48; + let width = 48; + let height = 48; let sigma_marginal = 20.0 / 255.0; let sigma_pre = sigma_marginal / 0.375f32.sqrt(); let margin = 7usize; // search_radius (4) + patch_radius (3) - let frame = correlated_noisy_frame(w, h, 0.5, sigma_pre, 11); + let frame = correlated_noisy_frame(width, height, 0.5, sigma_pre, 11); let params = attenuation_params(sigma_marginal); for rho in [0.0f32, 0.65] { - let mut windowed = NlmDenoiser::::new(&client, params.clone(), w, h); + let mut windowed = NlmDenoiser::::new(&client, params.clone(), width, height); windowed.rho_smoothed = Some(rho); windowed.push_frame(&frame); - let windowed_out = windowed - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let windowed_out = windowed.denoise().unwrap().unwrap(); - let mut separable = NlmDenoiser::::new(&client, params.clone(), w, h); + let mut separable = NlmDenoiser::::new(&client, params.clone(), width, height); separable.use_separable = true; separable.rho_smoothed = Some(rho); separable.push_frame(&frame); - let separable_out = separable - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let separable_out = separable.denoise().unwrap().unwrap(); let mut max_diff_interior = 0.0f32; - for y in margin..(h as usize - margin) { - for x in margin..(w as usize - margin) { - let idx = y * w as usize + x; - max_diff_interior = max_diff_interior.max((windowed_out[idx] - separable_out[idx]).abs()); + for y in margin..(height as usize - margin) { + for x in margin..(width as usize - margin) { + let idx = y * width as usize + x; + let diff = (windowed_out[idx] - separable_out[idx]).abs(); + max_diff_interior = max_diff_interior.max(diff); } } diff --git a/av-denoise-core/src/nlmeans/tests/split_sigma.rs b/av-denoise-core/src/nlmeans/tests/split_sigma.rs index 1f21ed2..faca1f2 100644 --- a/av-denoise-core/src/nlmeans/tests/split_sigma.rs +++ b/av-denoise-core/src/nlmeans/tests/split_sigma.rs @@ -1,10 +1,11 @@ use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; -/// Shared HQ auto-estimation params, k=0 only, so the temporal chain -/// stays inert (`temporal_stats_buf` needs `temporal_radius >= 1`) and -/// the two Immerkær spatial statistics (frame mean vs block p25) are -/// the only thing that can tell the median and low chains apart. +/// HQ auto-estimation params at k=0. +/// +/// The temporal chain needs `temporal_radius >= 1`, so at k=0 only the two spatial statistics (frame +/// mean and block p25) can tell the median and low chains apart. fn auto_params_k0() -> NlmParams { NlmParams { temporal_radius: 0, @@ -27,36 +28,34 @@ fn auto_params_k0() -> NlmParams { } } -/// A frame whose top half carries `sigma_a` noise and whose bottom -/// half carries `sigma_b`, both generated by `make_noisy_gaussian_frame` -/// at the same base level and spliced row-wise. -fn block_heterogeneous_frame(w: u32, h: u32, base: f32, sigma_a: f32, sigma_b: f32) -> Vec { - let top = make_noisy_gaussian_frame(w, h, 1, base, &[sigma_a]); - let bottom = make_noisy_gaussian_frame(w, h, 1, base, &[sigma_b]); - let row_len = w as usize; - let half = (h / 2) as usize; - let mut out = top; - for row in half..h as usize { +/// A frame whose top half carries `sigma_a` noise and whose bottom half carries `sigma_b`, both at +/// the same base level. +fn block_heterogeneous_frame(width: u32, height: u32, base: f32, sigma_a: f32, sigma_b: f32) -> Vec { + let top = make_noisy_gaussian_frame(width, height, 1, base, &[sigma_a]); + let bottom = make_noisy_gaussian_frame(width, height, 1, base, &[sigma_b]); + let row_len = width as usize; + let half = (height / 2) as usize; + let mut frame = top; + for row in half..height as usize { let start = row * row_len; - out[start..start + row_len].copy_from_slice(&bottom[start..start + row_len]); + frame[start..start + row_len].copy_from_slice(&bottom[start..start + row_len]); } - out + + frame } -/// Spatially uniform synthetic noise. The low chain's block-p25 sits -/// close to the median chain's frame mean (uniform noise means p25 is -/// nearly the median), so the derived `noise_offset` must stay close -/// to what the pre-split single-chain design would have produced from -/// the median chain's own smoothed sigma. +/// Uniform noise puts the block p25 close to the frame mean, so the low chain stays close to the +/// median chain and `noise_offset` close to the median-based offset. #[test] fn uniform_noise_low_chain_matches_median_chain() { let client = make_client(); - let w = 256; - let h = 256; + let width = 256; + let height = 256; let sigma = 8.0 / 255.0; - let frame = make_noisy_gaussian_frame(w, h, 1, 0.5, &[sigma]); + let frame = make_noisy_gaussian_frame(width, height, 1, 0.5, &[sigma]); - let mut denoiser = NlmDenoiser::::new(&client, auto_params_k0(), w, h); + let params = auto_params_k0(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); denoiser.denoise().unwrap(); @@ -66,37 +65,35 @@ fn uniform_noise_low_chain_matches_median_chain() { .current() .expect("seeded on first push")[0]; - let rel_err = (low - median).abs() / median; + let relative_error = (low - median).abs() / median; assert!( - rel_err <= 0.15, - "uniform noise: low chain {low} should sit close to median chain {median} (rel err {rel_err:.3})" + relative_error <= 0.15, + "uniform noise: low chain {low} should sit close to median chain {median} (rel err {relative_error:.3})" ); let expected_offset = denoiser.params.noise_offset_with(Some(&[median])); - let offset_rel_err = (denoiser.noise_offset - expected_offset).abs() / expected_offset; + let offset_relative_error = (denoiser.noise_offset - expected_offset).abs() / expected_offset; assert!( - offset_rel_err <= 0.3, + offset_relative_error <= 0.3, "uniform noise: noise_offset {} should stay close to the pre-split median-based offset \ - {expected_offset} (rel err {offset_rel_err:.3})", + {expected_offset} (rel err {offset_relative_error:.3})", denoiser.noise_offset ); } -/// Block-heterogeneous noise (half the frame at a low sigma, half at a -/// much higher sigma). The low chain must read below the median chain, -/// so `noise_offset` tracks the conservative statistic strictly below -/// what the median chain's own smoothed sigma would produce, while -/// `h2_inv_norm` (strength) keeps tracking the median chain exactly. +/// Half the frame sits at a low sigma and half at a much higher one, so the low chain reads below +/// the median chain. #[test] fn split_noise_offset_tracks_low_chain_strength_tracks_median() { let client = make_client(); - let w = 256; - let h = 256; + let width = 256; + let height = 256; let sigma_a = 2.0 / 255.0; let sigma_b = 20.0 / 255.0; - let frame = block_heterogeneous_frame(w, h, 0.5, sigma_a, sigma_b); + let frame = block_heterogeneous_frame(width, height, 0.5, sigma_a, sigma_b); - let mut denoiser = NlmDenoiser::::new(&client, auto_params_k0(), w, h); + let params = auto_params_k0(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); denoiser.denoise().unwrap(); @@ -125,30 +122,17 @@ fn split_noise_offset_tracks_low_chain_strength_tracks_median() { ); } -/// Ring isolation. A slot's own stage-1 partials must survive untouched -/// from the push that queues them until that slot reaches the centre -/// and gets folded, even though several other slots get pushed (and -/// dispatch their own noise estimate) in between. +/// A slot's stage-1 partials must survive from its push until it reaches the centre and is folded. /// -/// `temporal_radius = 2` (a 5-slot window) primes the leading edge with -/// two duplicates of frame 0 (low sigma), then two further high-sigma -/// pushes fill the window and trigger the first denoise. Its centre -/// slot is one of those low-sigma duplicates, untouched since priming. -/// A single shared (non-ring) partials scratch buffer would instead -/// have been overwritten by the most recent high-sigma dispatch by the -/// time it's read back, so the low chain's estimate would jump toward -/// the high sigma instead of staying low. -/// -/// One more high-sigma push then rotates the centre onto the first -/// real high-sigma frame, and the low chain must rise to reflect it, -/// confirming the ring keeps serving each slot's own fresh data as it -/// becomes centre rather than latching onto whichever slot happened to -/// isolate correctly first. +/// At `temporal_radius = 2` the leading edge is primed with copies of the low-sigma frame 0, and two +/// high-sigma pushes then fill the window. The first centre is a low-sigma copy, so a shared scratch +/// buffer would hold the latest high-sigma partials by then and the low estimate would jump. One +/// more high-sigma push centres the first high-sigma frame, and the low chain must rise with it. #[test] fn partials_ring_isolates_slots_between_push_and_fold() { let client = make_client(); - let w = 128; - let h = 128; + let width = 128; + let height = 128; let sigma_low = 2.0 / 255.0; let sigma_high = 30.0 / 255.0; @@ -156,12 +140,12 @@ fn partials_ring_isolates_slots_between_push_and_fold() { temporal_radius: 2, ..auto_params_k0() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); - let frame_low = make_noisy_gaussian_frame(w, h, 1, 0.5, &[sigma_low]); - let frame_high_1 = make_noisy_gaussian_frame(w, h, 1, 0.5, &[sigma_high]); - let frame_high_2 = make_noisy_gaussian_frame(w, h, 1, 0.5, &[sigma_high]); - let frame_high_3 = make_noisy_gaussian_frame(w, h, 1, 0.5, &[sigma_high]); + let frame_low = make_noisy_gaussian_frame(width, height, 1, 0.5, &[sigma_low]); + let frame_high_1 = make_noisy_gaussian_frame(width, height, 1, 0.5, &[sigma_high]); + let frame_high_2 = make_noisy_gaussian_frame(width, height, 1, 0.5, &[sigma_high]); + let frame_high_3 = make_noisy_gaussian_frame(width, height, 1, 0.5, &[sigma_high]); denoiser.push_frame(&frame_low); assert!(denoiser.denoise().unwrap().is_none(), "window not full yet"); @@ -191,34 +175,33 @@ fn partials_ring_isolates_slots_between_push_and_fold() { ); } -/// A frame whose top `split_rows` rows carry `sigma_top` noise and whose -/// remaining rows carry `sigma_bottom`, over a flat mid-grey field. +/// A flat mid-grey frame whose top `split_rows` rows carry `sigma_top` noise and the rest `sigma_bottom`. fn row_split_noisy_frame( - w: u32, - h: u32, + width: u32, + height: u32, split_rows: u32, sigma_top: f32, sigma_bottom: f32, seed: u32, ) -> Vec { - let clean = vec![0.5f32; (w * h) as usize]; - let top = noisy_field_over(&clean, w, h, sigma_top, seed); - let bottom = noisy_field_over(&clean, w, h, sigma_bottom, seed); + let clean = vec![0.5f32; (width * height) as usize]; + let top = noisy_field_over(&clean, width, height, sigma_top, seed); + let bottom = noisy_field_over(&clean, width, height, sigma_bottom, seed); - let split = (split_rows * w) as usize; + let split = (split_rows * width) as usize; let mut frame = bottom; frame[..split].copy_from_slice(&top[..split]); + frame } -/// Three eighths of the blocks carry a low sigma and the rest a higher -/// one, so the lower quartile of block sigmas lands on the low group -/// while the median lands on the high group. +/// Three eighths of the blocks carry a low sigma and the rest a higher one, so the lower quartile of +/// block sigmas lands on the low group while the median lands on the high group. #[test] fn temporal_only_chain_reads_the_median_block_sigma() { let client = make_client(); - let w = 128; - let h = 128; + let width = 128; + let height = 128; let sigma_low = 3.0 / 255.0; let sigma_high = 9.0 / 255.0; @@ -226,10 +209,10 @@ fn temporal_only_chain_reads_the_median_block_sigma() { temporal_radius: 2, ..auto_params_k0() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); for seed in 0..8 { - let frame = row_split_noisy_frame(w, h, 48, sigma_low, sigma_high, 300 + seed); + let frame = row_split_noisy_frame(width, height, 48, sigma_low, sigma_high, 300 + seed); denoiser.push_frame(&frame); let _ = denoiser.denoise().unwrap(); } @@ -239,10 +222,10 @@ fn temporal_only_chain_reads_the_median_block_sigma() { .current() .expect("a temporal sample should have folded by now")[0]; - let rel_err = (temporal_only - sigma_high).abs() / sigma_high; + let relative_error = (temporal_only - sigma_high).abs() / sigma_high; assert!( - rel_err <= 0.15, + relative_error <= 0.15, "temporal-only chain {temporal_only} should read the median block sigma near {sigma_high} \ - (rel err {rel_err:.3})" + (rel err {relative_error:.3})" ); } diff --git a/av-denoise-core/src/nlmeans/tests/strength_map.rs b/av-denoise-core/src/nlmeans/tests/strength_map.rs index 7fe2072..bd17a72 100644 --- a/av-denoise-core/src/nlmeans/tests/strength_map.rs +++ b/av-denoise-core/src/nlmeans/tests/strength_map.rs @@ -1,13 +1,14 @@ use super::helpers::{R, make_client}; use super::noise_curve::{ SyntheticQuarter, + banded_noisy_frame, frame_dims, quarter_at, quarters_at, - ramp_frame, synthetic_records, write_quarter, }; +use crate::bench_api::HostIo; use crate::collab::geometry::strength_map_dims; use crate::nlmeans::noise::{ NOISE_CURVE_BINS, @@ -24,6 +25,17 @@ use crate::nlmeans::{ChannelMode, HqParams, MotionCompensationMode, NlmDenoiser, const SIGMA: f32 = 0.02; +const ORIENTED: QuarterTensor = QuarterTensor { + xx: 1.0, + yy: 0.1, + xy: 0.0, +}; +const ISOTROPIC: QuarterTensor = QuarterTensor { + xx: 1.0, + yy: 1.0, + xy: 0.0, +}; + /// A flat curve predicting [SIGMA] at every luma. fn flat_curve() -> NoiseCurve { NoiseCurve { @@ -32,21 +44,25 @@ fn flat_curve() -> NoiseCurve { } } -/// Classes a single block whose four quarters are all `quarter`. -fn classify_one(quarter: SyntheticQuarter) -> QuarterClasses { +/// Classifies a single block whose four quarters are all `quarter`. +fn classify_block(quarter: SyntheticQuarter, cut: Option) -> QuarterClasses { let quarters = vec![quarter; 4]; let (width, height) = frame_dims(quarters.len()); let records = synthetic_records(&quarters); let curve = flat_curve(); - classify_quarters(&records, 1, width, height, &curve, None) + classify_quarters(&records, 1, width, height, &curve, cut) } -fn is_flat(quarter: SyntheticQuarter) -> bool { - let classes = classify_one(quarter); +fn all_flat(classes: &QuarterClasses) -> bool { let multipliers = classes.chroma_multipliers(2.0); multipliers.iter().all(|&multiplier| multiplier == 2.0) } +fn is_flat(quarter: SyntheticQuarter) -> bool { + let classes = classify_block(quarter, None); + all_flat(&classes) +} + fn with_flatness(fraction_of_variance: f32) -> SyntheticQuarter { SyntheticQuarter { flatness: fraction_of_variance * SIGMA * SIGMA, @@ -59,20 +75,10 @@ fn flat_grid(cols: usize, rows: usize) -> QuarterClasses { flat: true, luma: 0.2, }); - QuarterClasses::from_classes(cols, rows, vec![class; cols * rows]) + let classes = vec![class; cols * rows]; + QuarterClasses::from_classes(cols, rows, classes) } -const ORIENTED: QuarterTensor = QuarterTensor { - xx: 1.0, - yy: 0.1, - xy: 0.0, -}; -const ISOTROPIC: QuarterTensor = QuarterTensor { - xx: 1.0, - yy: 1.0, - xy: 0.0, -}; - /// A flat block whose four quarters carry a strongly horizontal structure tensor. fn oriented_flat_block(cut: Option) -> QuarterClasses { let quarter = SyntheticQuarter { @@ -80,27 +86,24 @@ fn oriented_flat_block(cut: Option) -> QuarterClasses { tensor_yy: 0.1, ..with_flatness(0.4) }; - let quarters = vec![quarter; 4]; - let (width, height) = frame_dims(quarters.len()); - let records = synthetic_records(&quarters); - let curve = flat_curve(); - classify_quarters(&records, 1, width, height, &curve, cut) -} - -fn all_flat(classes: &QuarterClasses) -> bool { - let multipliers = classes.chroma_multipliers(2.0); - multipliers.iter().all(|&multiplier| multiplier == 2.0) + classify_block(quarter, cut) } #[test] fn a_flat_noisy_quarter_is_flat() { - assert!(is_flat(with_flatness(0.4))); + let quarter = with_flatness(0.4); + let flat = is_flat(quarter); + assert!(flat); } #[test] fn the_flat_cut_sits_at_0_55_of_the_quarters_own_variance() { - assert!(is_flat(with_flatness(0.54))); - assert!(!is_flat(with_flatness(0.56))); + let below_cut = with_flatness(0.54); + let above_cut = with_flatness(0.56); + let below_cut_flat = is_flat(below_cut); + let above_cut_flat = is_flat(above_cut); + assert!(below_cut_flat); + assert!(!above_cut_flat); } #[test] @@ -109,13 +112,18 @@ fn a_moving_quarter_is_not_flat() { mean_residual: 3.5 / 255.0, ..quarter_at(0.3, SIGMA) }; - assert!(!is_flat(moving)); + let flat = is_flat(moving); + assert!(!flat); } #[test] fn a_quarter_far_noisier_than_the_curve_is_not_flat() { - assert!(is_flat(quarter_at(0.3, 2.4 * SIGMA))); - assert!(!is_flat(quarter_at(0.3, 2.6 * SIGMA))); + let slightly_noisier = quarter_at(0.3, 2.4 * SIGMA); + let far_noisier = quarter_at(0.3, 2.6 * SIGMA); + let slightly_noisier_flat = is_flat(slightly_noisier); + let far_noisier_flat = is_flat(far_noisier); + assert!(slightly_noisier_flat); + assert!(!far_noisier_flat); } #[test] @@ -128,8 +136,10 @@ fn clipped_quarters_are_not_flat() { luma_max: 253.0 / 255.0, ..quarter_at(0.3, SIGMA) }; - assert!(!is_flat(clipped_low)); - assert!(!is_flat(clipped_high)); + let clipped_low_flat = is_flat(clipped_low); + let clipped_high_flat = is_flat(clipped_high); + assert!(!clipped_low_flat); + assert!(!clipped_high_flat); } #[test] @@ -138,12 +148,15 @@ fn a_ragged_quarter_is_not_flat() { flatness: 3.0e38, ..quarter_at(0.3, SIGMA) }; - assert!(!is_flat(ragged)); + let flat = is_flat(ragged); + assert!(!flat); } #[test] fn a_noiseless_quarter_is_not_flat() { - assert!(!is_flat(quarter_at(0.3, 0.0))); + let noiseless = quarter_at(0.3, 0.0); + let flat = is_flat(noiseless); + assert!(!flat); } #[test] @@ -190,7 +203,8 @@ fn flat_quarters_take_the_boost_at_any_luma() { ]; let quarters = QuarterClasses::from_classes(2, 1, classes); - assert_eq!(quarters.luma_multipliers(params), vec![1.5, 1.5]); + let multipliers = quarters.luma_multipliers(params); + assert_eq!(multipliers, vec![1.5, 1.5]); } #[test] @@ -208,7 +222,8 @@ fn chroma_multipliers_never_soften() { ]; let quarters = QuarterClasses::from_classes(3, 1, classes); - assert_eq!(quarters.chroma_multipliers(1.5), vec![1.5, 1.0, 1.0]); + let multipliers = quarters.chroma_multipliers(1.5); + assert_eq!(multipliers, vec![1.5, 1.0, 1.0]); } #[test] @@ -231,7 +246,9 @@ fn unit_params_give_a_map_of_ones() { let quarters = QuarterClasses::from_classes(3, 1, classes); assert!(params.is_identity()); - assert_eq!(quarters.luma_multipliers(params), vec![1.0, 1.0, 1.0]); + + let multipliers = quarters.luma_multipliers(params); + assert_eq!(multipliers, vec![1.0, 1.0, 1.0]); } #[test] @@ -264,6 +281,7 @@ fn a_quarter_past_the_frame_edge_gets_one() { } else { [32.0, 0.0, 32.0, 0.0] }; + for (quarter_index, &quarter_pixels) in pixels.iter().enumerate() { if quarter_pixels > 0.0 { let quarter = quarter_at(0.1, SIGMA); @@ -271,6 +289,7 @@ fn a_quarter_past_the_frame_edge_gets_one() { } } } + let curve = flat_curve(); let params = StrengthMapParams { flat_boost: 1.5, @@ -289,8 +308,10 @@ fn a_quarter_past_the_frame_edge_gets_one() { #[test] fn a_reading_carries_classes_exactly_when_it_carries_a_curve() { let mut quarters = quarters_at(0.15, 0.02, 160); - quarters.extend(quarters_at(0.35, 0.01, 160)); - quarters.extend(quarters_at(0.6, 0.005, 160)); + let middle_band = quarters_at(0.35, 0.01, 160); + let bright_band = quarters_at(0.6, 0.005, 160); + quarters.extend(middle_band); + quarters.extend(bright_band); let (width, height) = frame_dims(quarters.len()); let records = synthetic_records(&quarters); @@ -333,7 +354,7 @@ fn the_front_end_keeps_classes_beside_the_curve_and_resets_both() { let mut classes_seen = false; for i in 0..12u32 { - let frame = ramp_frame(width, height, 100 + i); + let frame = banded_noisy_frame(width, height, 100 + i); denoiser.push_frame(&frame); let _ = denoiser.denoise().unwrap(); @@ -342,6 +363,7 @@ fn the_front_end_keeps_classes_beside_the_curve_and_resets_both() { assert_eq!(curve_present, classes_present, "push {i}"); classes_seen |= classes_present; } + assert!(classes_seen, "expected classes to form over the brightness ramp"); denoiser.reset_stream_state(); @@ -369,7 +391,9 @@ fn an_isotropic_neighbourhood_stays_flat() { let counts = classes.veto_textured(&tensors, 0.2); assert_eq!((counts.flat, counts.vetoed), (9, 0)); - assert!(all_flat(&classes)); + + let flat = all_flat(&classes); + assert!(flat); } #[test] @@ -391,7 +415,8 @@ fn the_veto_skips_quarters_without_a_class() { flat: true, luma: 0.2, }); - let mut classes = QuarterClasses::from_classes(2, 1, vec![flat, None]); + let row = vec![flat, None]; + let mut classes = QuarterClasses::from_classes(2, 1, row); let huge_oriented = QuarterTensor { xx: 100.0, yy: 0.0, @@ -414,7 +439,8 @@ fn a_vetoed_dark_quarter_takes_the_shadow_soften() { classes.veto_textured(&[ORIENTED], 0.5); - assert_eq!(classes.luma_multipliers(params), vec![0.65]); + let multipliers = classes.luma_multipliers(params); + assert_eq!(multipliers, vec![0.65]); } #[test] @@ -449,12 +475,14 @@ fn classify_quarters_vetoes_an_oriented_flat_block() { fn classify_quarters_without_a_cut_leaves_an_oriented_flat_block_flat() { let classes = oriented_flat_block(None); - assert!(all_flat(&classes)); + let flat = all_flat(&classes); + assert!(flat); } #[test] fn a_cut_of_one_leaves_an_oriented_flat_block_flat() { let classes = oriented_flat_block(Some(1.0)); - assert!(all_flat(&classes)); + let flat = all_flat(&classes); + assert!(flat); } diff --git a/av-denoise-core/src/nlmeans/tests/temporal.rs b/av-denoise-core/src/nlmeans/tests/temporal.rs index 78526d2..0ccce62 100644 --- a/av-denoise-core/src/nlmeans/tests/temporal.rs +++ b/av-denoise-core/src/nlmeans/tests/temporal.rs @@ -1,4 +1,7 @@ +use cubecl::prelude::ComputeClient; + use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; #[test] @@ -11,15 +14,13 @@ fn temporal_requires_full_window() { ..NlmParams::default() }; - let w = 8; - let h = 8; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 8; + let height = 8; + let frame = make_uniform_frame(width, height, 1, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); - // Leading-edge mirror fills R past slots from the very first push, so the - // window only needs R+1 real pushes (= 2 for radius 1) before the first - // submit produces output. + // The leading-edge mirror fills the R past slots, so the window needs only R+1 real pushes. denoiser.push_frame(&frame); assert!( denoiser.denoise().unwrap().is_none(), @@ -49,28 +50,22 @@ fn temporal_denoise_uniform() { hq: None, }; - let w = 8; - let h = 8; + let width = 8; + let height = 8; - let frame = make_uniform_frame(w, h, 1, 0.5); + let frame = make_uniform_frame(width, height, 1, 0.5); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame); denoiser.push_frame(&frame); denoiser.push_frame(&frame); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - for (i, &v) in result.iter().enumerate() { + for (i, &value) in result.iter().enumerate() { assert!( - (v - 0.5).abs() < 1e-4, - "temporal uniform: pixel {i} expected ~0.5, got {v}" + (value - 0.5).abs() < 1e-4, + "temporal uniform: pixel {i} expected ~0.5, got {value}" ); } } @@ -90,29 +85,23 @@ fn temporal_with_noisy_center_frame() { hq: None, }; - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let clean = make_uniform_frame(w, h, 1, 0.5); - let noisy = make_frame_with_noisy_region(w, h, 1, 0.5, 8, 8, 1, 0.8); + let clean = make_uniform_frame(width, height, 1, 0.5); + let noisy = make_frame_with_noisy_region(width, height, 1, 0.5, 8, 8, 1, 0.8); - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&clean); denoiser.push_frame(&noisy); denoiser.push_frame(&clean); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - let center_val = result[(8 * w + 8) as usize]; + let center_value = result[(8 * width + 8) as usize]; assert!( - center_val < 0.8, - "temporal denoising should suppress noise, got {center_val}" + center_value < 0.8, + "temporal denoising should suppress noise, got {center_value}" ); } @@ -131,37 +120,31 @@ fn temporal_asymmetric_frames_correct_weights() { hq: None, }; - let w = 16; - let h = 16; + let width = 16; + let height = 16; - let mut frame0 = vec![0.5f32; (w * h) as usize]; + let mut frame0 = vec![0.5f32; (width * height) as usize]; for y in 6..10 { for x in 6..10 { - frame0[(y * w + x) as usize] = 0.9; + frame0[(y * width + x) as usize] = 0.9; } } - let frame1 = vec![0.5f32; (w * h) as usize]; - let frame2 = vec![0.5f32; (w * h) as usize]; + let frame1 = vec![0.5f32; (width * height) as usize]; + let frame2 = vec![0.5f32; (width * height) as usize]; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); denoiser.push_frame(&frame0); denoiser.push_frame(&frame1); denoiser.push_frame(&frame2); - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let result = denoiser.denoise().unwrap().unwrap(); - let center_val = result[(8 * w + 8) as usize]; + let center_value = result[(8 * width + 8) as usize]; assert!( - (center_val - 0.5).abs() < 0.1, + (center_value - 0.5).abs() < 0.1, "temporal asymmetric: center should stay near 0.5 \ - (past frame de-weighted), got {center_val}" + (past frame de-weighted), got {center_value}" ); } @@ -175,21 +158,23 @@ fn flush_produces_remaining_frames() { ..NlmParams::default() }; - let w = 8; - let h = 8; + let width = 8; + let height = 8; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); for _ in 0..4 { - let frame = make_uniform_frame(w, h, 1, 0.5); + let frame = make_uniform_frame(width, height, 1, 0.5); denoiser.push_frame(&frame); let _ = denoiser.denoise().unwrap(); } let mut remaining: Vec> = Vec::new(); - denoiser - .flush(|frame| remaining.push(frame.as_f32().expect("f32 denoiser").to_vec())) - .unwrap(); + let collect = |frame: &[f32]| { + let samples = frame.to_vec(); + remaining.push(samples); + }; + denoiser.flush(collect).unwrap(); assert_eq!( remaining.len(), 1, @@ -197,19 +182,16 @@ fn flush_produces_remaining_frames() { ); for frame in &remaining { - assert_eq!(frame.len(), (w * h) as usize); + assert_eq!(frame.len(), (width * height) as usize); } } -/// `N` pushes at temporal radius `R` must produce exactly `N` total emissions -/// (during pushes + flush). Regression guard against the old bug where the -/// leading `R` logical frames were silently dropped (every scene lost its -/// first frame under `--temporal-radius >= 1`). +/// Pins the bug where the leading `R` frames of every scene were dropped. #[test] fn temporal_push_flush_frame_count_matches() { let client = make_client(); - let w = 8; - let h = 8; + let width = 8; + let height = 8; for radius in 1..=2 { let params = NlmParams { @@ -218,15 +200,14 @@ fn temporal_push_flush_frame_count_matches() { prefilter: PrefilterMode::None, ..NlmParams::default() }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); const PUSHES: usize = 10; let mut during_pushes = 0usize; for i in 0..PUSHES { - // Distinct frames so the kernel can't accidentally satisfy a - // count check by mis-pairing duplicate buffers. + // Distinct frames stop mis-paired duplicate buffers from satisfying the count. let value = 0.1 + (i as f32) * 0.05; - let frame = make_uniform_frame(w, h, 1, value); + let frame = make_uniform_frame(width, height, 1, value); denoiser.push_frame(&frame); if denoiser.denoise().unwrap().is_some() { during_pushes += 1; @@ -244,116 +225,84 @@ fn temporal_push_flush_frame_count_matches() { } } -/// Deterministic per-frame noisy copy of `base`, decorrelated across -/// `seed`. Same Irwin-Hall hash as `noisy_copy`, generalised to a -/// non-uniform base image instead of a flat value. -fn noisy_copy_of(base: &[f32], seed: u32, sigma: f32) -> Vec { - let unit_std = (1.0f32 / 3.0f32).sqrt(); - base.iter() - .enumerate() - .map(|(idx, &b)| { - let idx = idx as u32; - let mut sum = 0.0f32; - for k in 0..4u32 { - let mut hash = (idx * 4 + k) - .wrapping_mul(2654435761) - .wrapping_add(seed.wrapping_mul(0x9E37_79B9).wrapping_add(k)); - hash ^= hash >> 15; - hash = hash.wrapping_mul(0x85EB_CA6B); - hash ^= hash >> 13; - sum += (hash as f32 / u32::MAX as f32) - 0.5; - } - (b + (sum / unit_std) * sigma).clamp(0.0, 1.0) - }) - .collect() -} - fn psnr(reference: &[f32], test: &[f32]) -> f64 { let mse: f64 = reference .iter() .zip(test.iter()) - .map(|(&r, &t)| { - let d = (r as f64) - (t as f64); - d * d + .map(|(&reference_value, &test_value)| { + let difference = (reference_value as f64) - (test_value as f64); + difference * difference }) .sum::() / reference.len() as f64; + if mse <= 1e-20 { return 999.0; } + 10.0 * (1.0f64 / mse).log10() } -/// Structured content for the search-radius regression tests below. -/// Combines a gradient (a smooth region for NLM to average) with a -/// block of a different value (an edge NLM should preserve rather -/// than blur across). -fn structured_base(w: u32, h: u32) -> Vec { - let mut base = make_gradient_frame(w, h, 0.2, 0.8); - let bx0 = w / 3; - let by0 = h / 3; - for y in by0..by0 * 2 { - for x in bx0..bx0 * 2 { - base[(y * w + x) as usize] = 0.15; +/// A gradient (a smooth region to average) with a block of another value (an edge to preserve). +fn structured_base(width: u32, height: u32) -> Vec { + let mut base = make_gradient_frame(width, height, 0.2, 0.8); + let block_left = width / 3; + let block_top = height / 3; + for y in block_top..block_top * 2 { + for x in block_left..block_left * 2 { + base[(y * width + x) as usize] = 0.15; } } + base } -/// Runs `params` through the windowed (default) dispatch and again -/// through the separable dispatch (forced via the public -/// `use_separable` flag, an independently-implemented path that -/// doesn't share the windowed pair kernel's code), denoising `frames` -/// of noisy copies of `base` both times. Returns `(windowed_psnr, -/// separable_psnr)` against `base`. +/// Denoises `frames` through the windowed and separable dispatches, returning each PSNR against `base`. +/// +/// The separable path shares no code with the windowed pair kernel, so it acts as an independent +/// reference. fn windowed_vs_separable_psnr( - client: &cubecl::prelude::ComputeClient, + client: &ComputeClient, params: &NlmParams, - w: u32, - h: u32, + width: u32, + height: u32, base: &[f32], frames: &[Vec], ) -> (f64, f64) { - let mut windowed = NlmDenoiser::::new(client, params.clone(), w, h); + let windowed_params = params.clone(); + let mut windowed = NlmDenoiser::::new(client, windowed_params, width, height); for frame in frames { windowed.push_frame(frame); } - let windowed_result = windowed - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - let mut separable = NlmDenoiser::::new(client, params.clone(), w, h); + + let windowed_result = windowed.denoise().unwrap().unwrap(); + + let separable_params = params.clone(); + let mut separable = NlmDenoiser::::new(client, separable_params, width, height); separable.use_separable = true; for frame in frames { separable.push_frame(frame); } - let separable_result = separable - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - (psnr(base, &windowed_result), psnr(base, &separable_result)) + + let separable_result = separable.denoise().unwrap().unwrap(); + + let windowed_psnr = psnr(base, &windowed_result); + let separable_psnr = psnr(base, &separable_result); + + (windowed_psnr, separable_psnr) } -/// The backward temporal weight in `nlm_fused_pair_accumulate_window[_ref]` -/// must be measured against the same centre patch as the value it -/// multiplies. A weight measured against a shifted patch instead grows -/// wrong with the search offset, so this checks the windowed dispatch -/// against the independent separable dispatch at a search radius large -/// enough to expose a shift. +/// The windowed kernel's backward temporal weight must use the same centre patch as the value it +/// multiplies. +/// +/// A weight measured against a shifted patch grows wrong with the search offset, which these search +/// radii are large enough to expose. #[test] fn temporal_windowed_matches_separable_at_search_5_and_6() { let client = make_client(); - let w = 128; - let h = 128; - let base = structured_base(w, h); + let width = 128; + let height = 128; + let base = structured_base(width, height); for search_radius in [5u32, 6] { let params = NlmParams { @@ -369,10 +318,12 @@ fn temporal_windowed_matches_separable_at_search_5_and_6() { }; let sigma = 16.0 / 255.0; - let frames: Vec> = (0..9).map(|i| noisy_copy_of(&base, i, sigma)).collect(); + let frames: Vec> = (0..9) + .map(|seed| noisy_field_over(&base, width, height, sigma, seed)) + .collect(); let (windowed_psnr, separable_psnr) = - windowed_vs_separable_psnr(&client, ¶ms, w, h, &base, &frames); + windowed_vs_separable_psnr(&client, ¶ms, width, height, &base, &frames); assert!( (windowed_psnr - separable_psnr).abs() < 1.5, @@ -382,17 +333,13 @@ fn temporal_windowed_matches_separable_at_search_5_and_6() { } } -/// Same check as [`temporal_windowed_matches_separable_at_search_5_and_6`] -/// for `nlm_fused_pair_accumulate_window_ref`, the variant that reads -/// patch distances from a prefiltered reference clip instead of the raw -/// input. A prefilter is active so both the windowed and separable -/// dispatches route through their `_ref` kernels. +/// A prefilter is active so both dispatches read patch distances from the prefiltered reference. #[test] fn temporal_windowed_ref_matches_separable_ref_at_search_5_and_6() { let client = make_client(); - let w = 128; - let h = 128; - let base = structured_base(w, h); + let width = 128; + let height = 128; + let base = structured_base(width, height); for search_radius in [5u32, 6] { let params = NlmParams { @@ -411,10 +358,12 @@ fn temporal_windowed_ref_matches_separable_ref_at_search_5_and_6() { }; let sigma = 16.0 / 255.0; - let frames: Vec> = (0..9).map(|i| noisy_copy_of(&base, i, sigma)).collect(); + let frames: Vec> = (0..9) + .map(|seed| noisy_field_over(&base, width, height, sigma, seed)) + .collect(); let (windowed_psnr, separable_psnr) = - windowed_vs_separable_psnr(&client, ¶ms, w, h, &base, &frames); + windowed_vs_separable_psnr(&client, ¶ms, width, height, &base, &frames); assert!( (windowed_psnr - separable_psnr).abs() < 1.5, @@ -424,21 +373,15 @@ fn temporal_windowed_ref_matches_separable_ref_at_search_5_and_6() { } } -/// Same check as [`temporal_windowed_matches_separable_at_search_5_and_6`] -/// at the maximum supported search radius. Ignored by default. The -/// windowed kernel's fully unrolled window loop at this size overflows a -/// debug build's codegen stack even at the stack size -/// `.cargo/config.toml` sets (the spatial windowed kernel hits the same -/// limit). Release builds compile it fine. Run with -/// `cargo test --release -- --ignored -/// temporal_windowed_matches_separable_at_the_search_ceiling`. +/// A debug build overflows its codegen stack on the fully unrolled window loop at this radius, even +/// at the stack size `.cargo/config.toml` sets. #[test] #[ignore = "debug build codegen overflows the stack at search_radius=8, run with --release"] fn temporal_windowed_matches_separable_at_the_search_ceiling() { let client = make_client(); - let w = 128; - let h = 128; - let base = structured_base(w, h); + let width = 128; + let height = 128; + let base = structured_base(width, height); let params = NlmParams { temporal_radius: 4, @@ -453,9 +396,12 @@ fn temporal_windowed_matches_separable_at_the_search_ceiling() { }; let sigma = 16.0 / 255.0; - let frames: Vec> = (0..9).map(|i| noisy_copy_of(&base, i, sigma)).collect(); + let frames: Vec> = (0..9) + .map(|seed| noisy_field_over(&base, width, height, sigma, seed)) + .collect(); - let (windowed_psnr, separable_psnr) = windowed_vs_separable_psnr(&client, ¶ms, w, h, &base, &frames); + let (windowed_psnr, separable_psnr) = + windowed_vs_separable_psnr(&client, ¶ms, width, height, &base, &frames); assert!( (windowed_psnr - separable_psnr).abs() < 1.5, @@ -464,18 +410,15 @@ fn temporal_windowed_matches_separable_at_the_search_ceiling() { ); } -/// Uniform-content sanity check at the same search radii as -/// [`temporal_windowed_matches_separable_at_search_5_and_6`]. Uniform -/// input makes every patch distance zero regardless of which pixel a -/// kernel reads, so this cannot catch a mis-centred weight, but it does -/// catch a kernel reading or writing outside its intended memory region, -/// which would pull in unrelated data and break uniformity even here. +/// Uniform input zeroes every patch distance, so this cannot catch a mis-centred weight. +/// +/// It does catch a kernel reading or writing outside its region, which breaks the uniformity. #[test] fn temporal_uniform_passthrough_search_5_and_6() { let client = make_client(); - let w = 64; - let h = 64; - let frame = make_uniform_frame(w, h, 1, 0.5); + let width = 64; + let height = 64; + let frame = make_uniform_frame(width, height, 1, 0.5); for search_radius in [5u32, 6] { let params = NlmParams { @@ -490,22 +433,17 @@ fn temporal_uniform_passthrough_search_5_and_6() { hq: None, }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); for _ in 0..5 { denoiser.push_frame(&frame); } - let result = denoiser - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); - - for (i, &v) in result.iter().enumerate() { + + let result = denoiser.denoise().unwrap().unwrap(); + + for (i, &value) in result.iter().enumerate() { assert!( - (v - 0.5).abs() < 1e-3, - "search_radius={search_radius}: pixel {i} expected ~0.5, got {v}" + (value - 0.5).abs() < 1e-3, + "search_radius={search_radius}: pixel {i} expected ~0.5, got {value}" ); } } diff --git a/av-denoise-core/src/nlmeans/tests/temporal_noise.rs b/av-denoise-core/src/nlmeans/tests/temporal_noise.rs index e96aa57..7609134 100644 --- a/av-denoise-core/src/nlmeans/tests/temporal_noise.rs +++ b/av-denoise-core/src/nlmeans/tests/temporal_noise.rs @@ -1,6 +1,7 @@ use cubecl::prelude::*; use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::noise::{ NoiseCtx, QUARTER_FLATNESS, @@ -31,53 +32,52 @@ use crate::nlmeans::noise::{ }; use crate::nlmeans::*; -/// Uploads `prev`/`new` (densely packed `pixels * stored_ch`, no -/// channel padding needed by these tests) as ring slots 0 and 1, runs -/// the temporal-residual stats kernel diffing slot 1 against slot 0, -/// and returns slot 1's stats region. +/// Runs the temporal-residual stats kernel on `previous` and `new` as ring slots 0 and 1. /// -/// The stats buffer starts out full of garbage, so a lane the kernel -/// never writes shows up as a mismatch. +/// Both frames are packed as `pixels * stored_ch`, and slot 1's stats are returned. The stats +/// buffer starts full of garbage, so a lane the kernel never writes shows up as a mismatch. fn run_temporal_stats( - w: u32, - h: u32, + width: u32, + height: u32, stored_ch: u32, - prev: &[f32], + previous: &[f32], new: &[f32], luma_fields: bool, ) -> Vec { - run_temporal_stats_prefilled(w, h, stored_ch, prev, new, luma_fields, -7.0) + run_temporal_stats_prefilled(width, height, stored_ch, previous, new, luma_fields, -7.0) } /// [run_temporal_stats] over a stats buffer first filled with `fill`. fn run_temporal_stats_prefilled( - w: u32, - h: u32, + width: u32, + height: u32, stored_ch: u32, - prev: &[f32], + previous: &[f32], new: &[f32], luma_fields: bool, fill: f32, ) -> Vec { let client = make_client(); let frame_count = 2u32; - let frame_len = (w * h * stored_ch) as usize; - assert_eq!(prev.len(), frame_len); + let frame_len = (width * height * stored_ch) as usize; + assert_eq!(previous.len(), frame_len); assert_eq!(new.len(), frame_len); let mut ring = vec![0.0f32; frame_len * frame_count as usize]; - ring[..frame_len].copy_from_slice(prev); + ring[..frame_len].copy_from_slice(previous); ring[frame_len..].copy_from_slice(new); - let input_buf = client.create_from_slice(f32::as_bytes(&ring)); + let ring_bytes = f32::as_bytes(&ring); + let input_buf = client.create_from_slice(ring_bytes); let align = test_align(); - let stats_bytes = temporal_stats_buf_bytes(w, h, stored_ch, frame_count, align); + let stats_bytes = temporal_stats_buf_bytes(width, height, stored_ch, frame_count, align); let stats_fill = vec![fill; stats_bytes / size_of::()]; - let stats_buf = client.create_from_slice(f32::as_bytes(&stats_fill)); + let stats_fill_bytes = f32::as_bytes(&stats_fill); + let stats_buf = client.create_from_slice(stats_fill_bytes); let ctx = TemporalStatsCtx { - width: w, - height: h, + width, + height, stored_ch, frame_count, slot_new: 1, @@ -88,64 +88,75 @@ fn run_temporal_stats_prefilled( }; run_temporal_noise_stats::(&client, &ctx, luma_fields).expect("temporal noise stats dispatch failed"); - read_temporal_stats_slot::(&client, &stats_buf, w, h, stored_ch, frame_count, 1, align) - .expect("readback failed") + read_temporal_stats_slot::( + &client, + &stats_buf, + width, + height, + stored_ch, + frame_count, + 1, + align, + ) + .expect("readback failed") } -/// CPU oracle mirroring `nlm_temporal_noise_stats`'s block geometry -/// and summation independently of the kernel, so the kernel-unit -/// tests cross-check the GPU output against a from-scratch -/// implementation rather than a hand-derived closed form. +/// CPU oracle for `nlm_temporal_noise_stats`, written from scratch rather than from a closed form. /// -/// Only valid for single-channel (`stored_ch == 1`) `new`/`prev` frames, -/// because the quarter records always read channel 0 at a stride of -/// `stored_ch`, matching `sch` below only in that case. Every caller in -/// this file passes `stored_ch == 1`. -fn reference_temporal_stats(w: u32, h: u32, stored_ch: u32, prev: &[f32], new: &[f32]) -> Vec { - let sch = stored_ch as usize; - let (blocks_x, blocks_y) = temporal_stats_blocks(w, h); +/// Only valid for `stored_ch == 1`, because [quarter_fields_host] indexes the frames as single-channel. +fn reference_temporal_stats( + width: u32, + height: u32, + stored_ch: u32, + previous: &[f32], + new: &[f32], +) -> Vec { + let channel_stride = stored_ch as usize; + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); let record_len = temporal_stats_record_len(stored_ch) as usize; let mut out = vec![0.0f32; (blocks_x * blocks_y) as usize * record_len]; - for by in 0..blocks_y { - for bx in 0..blocks_x { - let ox = bx * TEMPORAL_NOISE_BLOCK; - let oy = by * TEMPORAL_NOISE_BLOCK; - let bw = TEMPORAL_NOISE_BLOCK.min(w - ox); - let bh = TEMPORAL_NOISE_BLOCK.min(h - oy); + for block_y in 0..blocks_y { + for block_x in 0..blocks_x { + let origin_x = block_x * TEMPORAL_NOISE_BLOCK; + let origin_y = block_y * TEMPORAL_NOISE_BLOCK; + let block_width = TEMPORAL_NOISE_BLOCK.min(width - origin_x); + let block_height = TEMPORAL_NOISE_BLOCK.min(height - origin_y); - let mut sum_d = vec![0.0f32; sch]; - let mut sum_d2 = vec![0.0f32; sch]; + let mut sum_d = vec![0.0f32; channel_stride]; + let mut sum_d2 = vec![0.0f32; channel_stride]; let mut sum_lag = 0.0f32; - for ly in 0..bh { - let y = oy + ly; - let mut d0_row = Vec::with_capacity(bw as usize); - for lx in 0..bw { - let x = ox + lx; - let idx = (y * w + x) as usize; - for c in 0..sch { - let d = new[idx * sch + c] - prev[idx * sch + c]; - sum_d[c] += d; - sum_d2[c] += d * d; - if c == 0 { - d0_row.push(d); + for local_y in 0..block_height { + let y = origin_y + local_y; + let mut luma_residual_row = Vec::with_capacity(block_width as usize); + for local_x in 0..block_width { + let x = origin_x + local_x; + let idx = (y * width + x) as usize; + for channel in 0..channel_stride { + let sample = idx * channel_stride + channel; + let residual = new[sample] - previous[sample]; + sum_d[channel] += residual; + sum_d2[channel] += residual * residual; + if channel == 0 { + luma_residual_row.push(residual); } } } - for lx in 0..(bw as usize).saturating_sub(1) { - sum_lag += d0_row[lx] * d0_row[lx + 1]; + + for local_x in 0..(block_width as usize).saturating_sub(1) { + sum_lag += luma_residual_row[local_x] * luma_residual_row[local_x + 1]; } } - let block_index = (by * blocks_x + bx) as usize; + let block_index = (block_y * blocks_x + block_x) as usize; let base = block_index * record_len; - out[base..base + sch].copy_from_slice(&sum_d); - out[base + sch..base + 2 * sch].copy_from_slice(&sum_d2); - out[base + 2 * sch] = sum_lag; + out[base..base + channel_stride].copy_from_slice(&sum_d); + out[base + channel_stride..base + 2 * channel_stride].copy_from_slice(&sum_d2); + out[base + 2 * channel_stride] = sum_lag; for quarter in 0..TEMPORAL_QUARTERS { - let fields = quarter_fields_host(new, prev, w, h, bx, by, quarter); + let fields = quarter_fields_host(new, previous, width, height, block_x, block_y, quarter); let quarter_start = base + quarter_offset(stored_ch, quarter); let quarter_end = quarter_start + TEMPORAL_QUARTER_FIELDS as usize; out[quarter_start..quarter_end].copy_from_slice(&fields); @@ -161,35 +172,33 @@ fn quarter_offset(stored_ch: u32, quarter: u32) -> usize { (2 * stored_ch + TEMPORAL_QUARTER_BASE + quarter * TEMPORAL_QUARTER_FIELDS) as usize } -/// The nine fields one quarter's record carries, computed on the host. +/// The nine fields of one quarter's record, computed on the host from single-channel frames. /// -/// `new` and `prev` are densely packed single-channel frames, matching -/// every caller in this file. An empty quarter keeps the kernel's min -/// and max seeds. +/// An empty quarter keeps the kernel's min and max seeds. fn quarter_fields_host( new: &[f32], - prev: &[f32], + previous: &[f32], width: u32, height: u32, - bx: u32, - by: u32, + block_x: u32, + block_y: u32, quarter: u32, ) -> [f32; 9] { let size = TEMPORAL_QUARTER_SIZE; - let origin_x = bx * TEMPORAL_NOISE_BLOCK + (quarter % 2) * size; - let origin_y = by * TEMPORAL_NOISE_BLOCK + (quarter / 2) * size; - let quarter_w = size.min(width.saturating_sub(origin_x)); - let quarter_h = size.min(height.saturating_sub(origin_y)); + let origin_x = block_x * TEMPORAL_NOISE_BLOCK + (quarter % 2) * size; + let origin_y = block_y * TEMPORAL_NOISE_BLOCK + (quarter / 2) * size; + let quarter_width = size.min(width.saturating_sub(origin_x)); + let quarter_height = size.min(height.saturating_sub(origin_y)); let mut sum_d = 0.0f32; let mut sum_d2 = 0.0f32; let mut luma_sum = 0.0f32; let mut luma_min = 1.0e30f32; let mut luma_max = -1.0e30f32; - for y in 0..quarter_h { - for x in 0..quarter_w { + for y in 0..quarter_height { + for x in 0..quarter_width { let idx = ((origin_y + y) * width + origin_x + x) as usize; - let residual = new[idx] - prev[idx]; + let residual = new[idx] - previous[idx]; sum_d += residual; sum_d2 += residual * residual; luma_sum += new[idx]; @@ -198,37 +207,41 @@ fn quarter_fields_host( } } - let full = quarter_w == size && quarter_h == size; + let full = quarter_width == size && quarter_height == size; let flatness = if full { let mut cells = [0.0f32; 16]; - for cy in 0..4u32 { - for cx in 0..4u32 { + for cell_y in 0..4u32 { + for cell_x in 0..4u32 { let mut cell = 0.0f32; for dy in 0..2u32 { for dx in 0..2u32 { - let idx = ((origin_y + 2 * cy + dy) * width + origin_x + 2 * cx + dx) as usize; - cell += 0.5 * (new[idx] + prev[idx]); + let row = origin_y + 2 * cell_y + dy; + let column = origin_x + 2 * cell_x + dx; + let idx = (row * width + column) as usize; + cell += 0.5 * (new[idx] + previous[idx]); } } - cells[(cy * 4 + cx) as usize] = cell / 4.0; + + cells[(cell_y * 4 + cell_x) as usize] = cell / 4.0; } } let mut energy = 0.0f32; - for cy in 0..4usize { - for cx in 0..4usize { - let here = cells[cy * 4 + cx]; - if cx + 1 < 4 { - let right = cells[cy * 4 + cx + 1]; + for cell_y in 0..4usize { + for cell_x in 0..4usize { + let here = cells[cell_y * 4 + cell_x]; + if cell_x + 1 < 4 { + let right = cells[cell_y * 4 + cell_x + 1]; energy += (here - right) * (here - right); } - if cy + 1 < 4 { - let below = cells[(cy + 1) * 4 + cx]; + if cell_y + 1 < 4 { + let below = cells[(cell_y + 1) * 4 + cell_x]; energy += (here - below) * (here - below); } } } + energy / 24.0 } else { 3.0e38 @@ -236,14 +249,14 @@ fn quarter_fields_host( let mean_at = |x: u32, y: u32| { let idx = ((origin_y + y) * width + origin_x + x) as usize; - 0.5 * (new[idx] + prev[idx]) + 0.5 * (new[idx] + previous[idx]) }; let mut tensor_xx = 0.0f32; let mut tensor_yy = 0.0f32; let mut tensor_xy = 0.0f32; - for y in 0..quarter_h.saturating_sub(1) { - for x in 0..quarter_w.saturating_sub(1) { + for y in 0..quarter_height.saturating_sub(1) { + for x in 0..quarter_width.saturating_sub(1) { let top_left = mean_at(x, y); let top_right = mean_at(x + 1, y); let bottom_left = mean_at(x, y + 1); @@ -269,173 +282,165 @@ fn quarter_fields_host( fields } -fn assert_close(actual: &[f32], expected: &[f32], tol: f32, msg: &str) { - assert_eq!(actual.len(), expected.len(), "{msg}: length mismatch"); - for (i, (&a, &e)) in actual.iter().zip(expected.iter()).enumerate() { +fn assert_close(actual: &[f32], expected: &[f32], tolerance: f32, message: &str) { + assert_eq!(actual.len(), expected.len(), "{message}: length mismatch"); + for (i, (&actual_value, &expected_value)) in actual.iter().zip(expected.iter()).enumerate() { assert!( - (a - e).abs() <= tol, - "{msg}: index {i}: got {a}, expected {e} (tol {tol})" + (actual_value - expected_value).abs() <= tolerance, + "{message}: index {i}: got {actual_value}, expected {expected_value} (tol {tolerance})" ); } } -/// A constant difference over a frame that grids into exactly four full -/// blocks, with no ragged edges. -/// -/// Every block's sums then reduce to a closed form, so the expected -/// values can be written out by hand. +/// The frame grids into exactly four full blocks, so every block's sums have a closed form. #[test] fn kernel_uniform_diff_exact_sums() { - let w = 32; - let h = 32; - let k = 0.1f32; - let prev = vec![0.0f32; (w * h) as usize]; - let new = vec![k; (w * h) as usize]; - - let n = 256.0f32; - let n_pairs = 240.0f32; - let expected_block = [n * k, n * k * k, n_pairs * k * k]; - - let got = run_temporal_stats(w, h, 1, &prev, &new, true); - let oracle = reference_temporal_stats(w, h, 1, &prev, &new); + let width = 32; + let height = 32; + let difference = 0.1f32; + let previous = vec![0.0f32; (width * height) as usize]; + let new = vec![difference; (width * height) as usize]; + + let block_pixels = 256.0f32; + let lag_pairs = 240.0f32; + let expected_block = [ + block_pixels * difference, + block_pixels * difference * difference, + lag_pairs * difference * difference, + ]; + + let got = run_temporal_stats(width, height, 1, &previous, &new, true); + let oracle = reference_temporal_stats(width, height, 1, &previous, &new); let record_len = temporal_stats_record_len(1) as usize; for block in 0..4 { - let rec = &got[block * record_len..block * record_len + 3]; - assert_close(rec, &expected_block, 1e-4, &format!("block {block}")); + let record = &got[block * record_len..block * record_len + 3]; + let label = format!("block {block}"); + assert_close(record, &expected_block, 1e-4, &label); } + assert_close(&got, &oracle, 1e-4, "kernel vs CPU oracle"); } -/// A horizontal ramp difference over a grid of six full blocks. -/// -/// Nothing is asserted in closed form here. The kernel is checked -/// against the independent CPU reference instead, which covers the -/// varying per-pixel content the uniform test cannot reach. +/// Checked against the CPU oracle rather than a closed form, covering per-pixel variation the uniform +/// test cannot reach. #[test] fn kernel_gradient_diff_exact_sums() { - let w = 48; - let h = 32; - let mut prev = vec![0.0f32; (w * h) as usize]; - let mut new = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - new[(y * w + x) as usize] = x as f32; + let width = 48; + let height = 32; + let mut previous = vec![0.0f32; (width * height) as usize]; + let mut new = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + new[(y * width + x) as usize] = x as f32; } } - // Keep `prev` at zero so `d = new`. - prev.fill(0.0); - let got = run_temporal_stats(w, h, 1, &prev, &new, true); - let oracle = reference_temporal_stats(w, h, 1, &prev, &new); + // `previous` stays zero so the residual equals `new`. + previous.fill(0.0); + + let got = run_temporal_stats(width, height, 1, &previous, &new, true); + let oracle = reference_temporal_stats(width, height, 1, &previous, &new); assert_close(&got, &oracle, 1e-2, "kernel vs CPU oracle (gradient diff)"); } -/// A 33x17 frame is ragged on both axes, leaving a last column one -/// pixel wide and a last row one pixel tall. +/// A 33x17 frame leaves a last column one pixel wide and a last row one pixel tall. /// -/// A uniform difference still has a closed form per block, just with -/// each block's own truncated counts. -/// -/// This asserts those directly rather than only against the CPU -/// reference. +/// A uniform difference still has a closed form per block, using each block's truncated counts. #[test] fn kernel_ragged_block_dims() { - let w = 33; - let h = 17; - let k = 0.2f32; - let prev = vec![0.0f32; (w * h) as usize]; - let new = vec![k; (w * h) as usize]; + let width = 33; + let height = 17; + let difference = 0.2f32; + let previous = vec![0.0f32; (width * height) as usize]; + let new = vec![difference; (width * height) as usize]; - let (blocks_x, blocks_y) = temporal_stats_blocks(w, h); + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); assert_eq!((blocks_x, blocks_y), (3, 2)); - let got = run_temporal_stats(w, h, 1, &prev, &new, true); - let oracle = reference_temporal_stats(w, h, 1, &prev, &new); - // `luma_sum` reduces on the device through a shared-memory tree, - // rather than the oracle's plain sequential loop. Summing the same - // 256 copies of a value like 0.2, which float32 cannot hold exactly, - // in a different order picks up its own small rounding error, so - // this needs a little more slack than the other lanes. + let got = run_temporal_stats(width, height, 1, &previous, &new, true); + let oracle = reference_temporal_stats(width, height, 1, &previous, &new); + // `luma_sum` reduces through a shared-memory tree on the device but a sequential loop in the + // oracle. Summing 256 copies of a value f32 cannot hold exactly in a different order needs more + // slack than the other lanes. assert_close(&got, &oracle, 5e-4, "kernel vs CPU oracle (ragged)"); - for by in 0..blocks_y { - for bx in 0..blocks_x { - let bw = TEMPORAL_NOISE_BLOCK.min(w - bx * TEMPORAL_NOISE_BLOCK); - let bh = TEMPORAL_NOISE_BLOCK.min(h - by * TEMPORAL_NOISE_BLOCK); - let n = (bw * bh) as f32; - let n_pairs = (bh * bw.saturating_sub(1)) as f32; - let expected = [n * k, n * k * k, n_pairs * k * k]; - - let block = (by * blocks_x + bx) as usize; + for block_y in 0..blocks_y { + for block_x in 0..blocks_x { + let block_width = TEMPORAL_NOISE_BLOCK.min(width - block_x * TEMPORAL_NOISE_BLOCK); + let block_height = TEMPORAL_NOISE_BLOCK.min(height - block_y * TEMPORAL_NOISE_BLOCK); + let block_pixels = (block_width * block_height) as f32; + let lag_pairs = (block_height * block_width.saturating_sub(1)) as f32; + let expected = [ + block_pixels * difference, + block_pixels * difference * difference, + lag_pairs * difference * difference, + ]; + + let block = (block_y * blocks_x + block_x) as usize; let record_len = temporal_stats_record_len(1) as usize; - let rec = &got[block * record_len..block * record_len + 3]; - // A looser tolerance than the reference comparison above. - // This is a hand-derived closed form summing up to 256 - // copies of a value f32 cannot hold exactly, so it picks up - // a little more floating-point slack than comparing two - // sums that walk the same block in the same order. - assert_close( - rec, - &expected, - 5e-3, - &format!("ragged block ({bx},{by}), dims {bw}x{bh}"), - ); + let record = &got[block * record_len..block * record_len + 3]; + // Looser than the oracle comparison, because the kernel accumulates up to 256 copies of a + // value f32 cannot hold exactly while the closed form multiplies once. + let label = format!("ragged block ({block_x},{block_y}), dims {block_width}x{block_height}"); + assert_close(record, &expected, 5e-3, &label); } } } -fn assert_relative(actual: f32, expected: f32, tol: f32, msg: &str) { +fn assert_relative(actual: f32, expected: f32, tolerance: f32, message: &str) { let rel_err = (actual - expected).abs() / expected.abs().max(1e-8); assert!( - rel_err <= tol, - "{msg}: got {actual}, expected {expected} (rel err {rel_err})" + rel_err <= tolerance, + "{message}: got {actual}, expected {expected} (rel err {rel_err})" ); } /// A textured frame with a small, pixel-varying residual. -fn textured_pair(w: u32, h: u32) -> (Vec, Vec) { - let mut new = vec![0.0f32; (w * h) as usize]; - let mut prev = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let idx = (y * w + x) as usize; +fn textured_pair(width: u32, height: u32) -> (Vec, Vec) { + let mut new = vec![0.0f32; (width * height) as usize]; + let mut previous = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let idx = (y * width + x) as usize; let value = 0.2 + 0.6 * ((x * 7 + y * 13) % 17) as f32 / 17.0; new[idx] = value; let offset = 0.02 * (((x + y) % 3) as f32 - 1.0); - prev[idx] = (value + offset).clamp(0.0, 1.0); + previous[idx] = (value + offset).clamp(0.0, 1.0); } } - (prev, new) + + (previous, new) } -fn assert_field_close(actual: f32, expected: f32, msg: &str) { - let tol = 1e-4 * expected.abs().max(1.0); +fn assert_field_close(actual: f32, expected: f32, message: &str) { + let tolerance = 1e-4 * expected.abs().max(1.0); assert!( - (actual - expected).abs() <= tol, - "{msg}: got {actual}, expected {expected} (tol {tol})" + (actual - expected).abs() <= tolerance, + "{message}: got {actual}, expected {expected} (tol {tolerance})" ); } -/// Checks every quarter record of a single-channel run against the host -/// mirror, and returns how many quarters carried real flatness. -fn assert_quarters_match_mirror(w: u32, h: u32, prev: &[f32], new: &[f32]) -> usize { +/// Checks every quarter record of a single-channel run against the host mirror. +/// +/// Returns how many quarters carried real flatness. +fn assert_quarters_match_mirror(width: u32, height: u32, previous: &[f32], new: &[f32]) -> usize { let stored_ch = 1u32; - let got = run_temporal_stats(w, h, stored_ch, prev, new, true); + let got = run_temporal_stats(width, height, stored_ch, previous, new, true); let record_len = temporal_stats_record_len(stored_ch) as usize; - let (blocks_x, blocks_y) = temporal_stats_blocks(w, h); + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); let mut real_flatness = 0; - for by in 0..blocks_y { - for bx in 0..blocks_x { - let block_index = (by * blocks_x + bx) as usize; + for block_y in 0..blocks_y { + for block_x in 0..blocks_x { + let block_index = (block_y * blocks_x + block_x) as usize; let record = &got[block_index * record_len..(block_index + 1) * record_len]; for quarter in 0..TEMPORAL_QUARTERS { - let expected = quarter_fields_host(new, prev, w, h, bx, by, quarter); + let expected = quarter_fields_host(new, previous, width, height, block_x, block_y, quarter); let offset = quarter_offset(stored_ch, quarter); let fields = &record[offset..offset + TEMPORAL_QUARTER_FIELDS as usize]; - let label = format!("{w}x{h} block ({bx},{by}) quarter {quarter}"); + let label = format!("{width}x{height} block ({block_x},{block_y}) quarter {quarter}"); let exact_fields = [ QUARTER_SUM_D, @@ -447,7 +452,8 @@ fn assert_quarters_match_mirror(w: u32, h: u32, prev: &[f32], new: &[f32]) -> us ]; for field in exact_fields { let index = field as usize; - assert_field_close(fields[index], expected[index], &format!("{label} field {field}")); + let field_label = format!("{label} field {field}"); + assert_field_close(fields[index], expected[index], &field_label); } let min_index = QUARTER_LUMA_MIN as usize; @@ -461,12 +467,8 @@ fn assert_quarters_match_mirror(w: u32, h: u32, prev: &[f32], new: &[f32]) -> us if expected_flatness == 3.0e38 { assert_eq!(got_flatness, 3.0e38, "{label}: flatness sentinel"); } else { - assert_relative( - got_flatness, - expected_flatness, - 1e-4, - &format!("{label} flatness"), - ); + let flatness_label = format!("{label} flatness"); + assert_relative(got_flatness, expected_flatness, 1e-4, &flatness_label); real_flatness += 1; } } @@ -474,16 +476,12 @@ fn assert_quarters_match_mirror(w: u32, h: u32, prev: &[f32], new: &[f32]) -> us } // The scalar lanes stay exactly as the 16x16 oracle computes them. - let oracle = reference_temporal_stats(w, h, stored_ch, prev, new); + let oracle = reference_temporal_stats(width, height, stored_ch, previous, new); for block in 0..(blocks_x * blocks_y) as usize { - let got_rec = &got[block * record_len..block * record_len + 3]; - let oracle_rec = &oracle[block * record_len..block * record_len + 3]; - assert_close( - got_rec, - oracle_rec, - 1e-4, - &format!("block {block} sum_d/sum_d2/lag"), - ); + let got_record = &got[block * record_len..block * record_len + 3]; + let oracle_record = &oracle[block * record_len..block * record_len + 3]; + let label = format!("block {block} sum_d/sum_d2/lag"); + assert_close(got_record, oracle_record, 1e-4, &label); } real_flatness @@ -491,11 +489,11 @@ fn assert_quarters_match_mirror(w: u32, h: u32, prev: &[f32], new: &[f32]) -> us #[test] fn quarter_records_match_the_host_mirror_on_texture() { - let w = 48u32; - let h = 32u32; - let (prev, new) = textured_pair(w, h); + let width = 48u32; + let height = 32u32; + let (previous, new) = textured_pair(width, height); - let real_flatness = assert_quarters_match_mirror(w, h, &prev, &new); + let real_flatness = assert_quarters_match_mirror(width, height, &previous, &new); assert_eq!(real_flatness, 24, "every quarter of six full blocks is full"); } @@ -503,54 +501,51 @@ fn quarter_records_match_the_host_mirror_on_texture() { #[test] fn quarter_records_match_the_host_mirror_on_flat_noise() { let size = 32u32; - let prev = noisy_copy(size, 0.5, 0.01, 1); + let previous = noisy_copy(size, 0.5, 0.01, 1); let new = noisy_copy(size, 0.5, 0.01, 2); - let real_flatness = assert_quarters_match_mirror(size, size, &prev, &new); + let real_flatness = assert_quarters_match_mirror(size, size, &previous, &new); assert_eq!(real_flatness, 16); } -/// A 45x29 frame leaves the last column 13 pixels wide and the last row -/// 13 pixels tall, so each ragged block holds full and partial quarters. +/// A 45x29 frame leaves the last column 13 pixels wide and the last row 13 pixels tall, so each +/// ragged block holds full and partial quarters. #[test] fn quarter_records_match_the_host_mirror_on_ragged_edges() { - let w = 45u32; - let h = 29u32; - let (prev, new) = textured_pair(w, h); + let width = 45u32; + let height = 29u32; + let (previous, new) = textured_pair(width, height); - let real_flatness = assert_quarters_match_mirror(w, h, &prev, &new); + let real_flatness = assert_quarters_match_mirror(width, height, &previous, &new); - // Block (0,0) has 4 full quarters, block (1,0) 4, block (2,0) 2, - // block (0,1) 2, block (1,1) 2 and block (2,1) 1. + // Block (0,0) has 4 full quarters, block (1,0) 4, block (2,0) 2, block (0,1) 2, block (1,1) 2 + // and block (2,1) 1. assert_eq!(real_flatness, 15); } -/// A 40x24 frame leaves the last column and row exactly 8 pixels, so -/// some quarters lie entirely outside the frame. +/// A 40x24 frame leaves the last column and row exactly 8 pixels, so some quarters lie entirely +/// outside the frame. #[test] fn quarter_records_match_the_host_mirror_with_empty_quarters() { - let w = 40u32; - let h = 24u32; - let (prev, new) = textured_pair(w, h); + let width = 40u32; + let height = 24u32; + let (previous, new) = textured_pair(width, height); - let real_flatness = assert_quarters_match_mirror(w, h, &prev, &new); + let real_flatness = assert_quarters_match_mirror(width, height, &previous, &new); assert_eq!(real_flatness, 4 + 4 + 2 + 2 + 2 + 1); } -/// A full block of a constant base plus independent noise in each frame -/// has to read low flatness in every quarter, showing the -/// smoothed-temporal-mean gate does not mistake the noise for texture. #[test] fn flat_noisy_block_reads_low_flatness_in_every_quarter() { let size = 16u32; let sigma = 0.01f32; let stored_ch = 1u32; - let prev = noisy_copy(size, 0.5, sigma, 1); + let previous = noisy_copy(size, 0.5, sigma, 1); let new = noisy_copy(size, 0.5, sigma, 2); - let got = run_temporal_stats(size, size, stored_ch, &prev, &new, true); + let got = run_temporal_stats(size, size, stored_ch, &previous, &new, true); for quarter in 0..TEMPORAL_QUARTERS { let offset = quarter_offset(stored_ch, quarter); @@ -562,29 +557,27 @@ fn flat_noisy_block_reads_low_flatness_in_every_quarter() { } } -/// With `luma_fields` off, the 36 quarter lanes read 0 even over a -/// buffer full of garbage, and every other lane matches the on run bit -/// for bit. +/// The off run starts from a buffer of garbage, so a zero quarter lane was written by the kernel. #[test] fn luma_fields_off_leaves_the_quarter_lanes_zero_and_the_rest_unchanged() { - let w = 45u32; - let h = 29u32; + let width = 45u32; + let height = 29u32; let stored_ch = 1u32; - let (prev, new) = textured_pair(w, h); + let (previous, new) = textured_pair(width, height); - let with_luma = run_temporal_stats(w, h, stored_ch, &prev, &new, true); - let without_luma = run_temporal_stats_prefilled(w, h, stored_ch, &prev, &new, false, 123.0); + let with_luma = run_temporal_stats(width, height, stored_ch, &previous, &new, true); + let without_luma = run_temporal_stats_prefilled(width, height, stored_ch, &previous, &new, false, 123.0); let record_len = temporal_stats_record_len(stored_ch) as usize; let scalar_len = 2 * stored_ch as usize + 1; - let (blocks_x, blocks_y) = temporal_stats_blocks(w, h); + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); for block in 0..(blocks_x * blocks_y) as usize { let base = block * record_len; - let with_rec = &with_luma[base..base + scalar_len]; - let without_rec = &without_luma[base..base + scalar_len]; + let with_record = &with_luma[base..base + scalar_len]; + let without_record = &without_luma[base..base + scalar_len]; assert_eq!( - with_rec, without_rec, + with_record, without_record, "block {block}: sum_d/sum_d2/lag must not depend on luma_fields" ); @@ -602,10 +595,10 @@ fn white_noise_pair_recovers_known_sigma() { let size = 256; let true_sigma = 8.0 / 255.0; - let prev = noisy_copy(size, 0.5, true_sigma, 1); + let previous = noisy_copy(size, 0.5, true_sigma, 1); let new = noisy_copy(size, 0.5, true_sigma, 2); - let records = run_temporal_stats(size, size, 1, &prev, &new, true); + let records = run_temporal_stats(size, size, 1, &previous, &new, true); let sample = aggregate_temporal_noise_stats(&records, 1, 1, size, size) .expect("a static white-noise pair should clear the static-block floor"); @@ -617,31 +610,23 @@ fn white_noise_pair_recovers_known_sigma() { ); } -/// Grain that is correlated between neighbouring pixels, made by -/// blurring a white noise field horizontally. -/// -/// The sigma is still recoverable within 10%, because the temporal -/// measurement does not care about spatial correlation. +/// The temporal measurement ignores spatial correlation, so the sigma stays within 10%. /// -/// The correlation reading has to clearly register the blur. That is the -/// whole point of this estimator next to Immerkær, which reads -/// correlated grain low. See -/// `hq_temporal_folds_correlated_grain_above_immerkaer_alone`. +/// The correlation reading must still register the blur, which Immerkær reads low. #[test] fn correlated_noise_pair_recovers_marginal_sigma_and_rho() { - let w = 256; - let h = 256; + let width = 256; + let height = 256; let sigma_marginal = 8.0 / 255.0; - // The blur scales the variance by 0.375, the sum of its squared - // weights, so the sigma before the blur has to be raised to land - // the blurred field on the target. + // The blur scales the variance by 0.375, the sum of its squared taps, so the sigma before the + // blur is raised to land the blurred field on the target. let sigma_pre = sigma_marginal / 0.375f32.sqrt(); - let prev = correlated_noisy_frame(w, h, 0.5, sigma_pre, 11); - let new = correlated_noisy_frame(w, h, 0.5, sigma_pre, 12); + let previous = correlated_noisy_frame(width, height, 0.5, sigma_pre, 11); + let new = correlated_noisy_frame(width, height, 0.5, sigma_pre, 12); - let records = run_temporal_stats(w, h, 1, &prev, &new, true); - let sample = aggregate_temporal_noise_stats(&records, 1, 1, w, h) + let records = run_temporal_stats(width, height, 1, &previous, &new, true); + let sample = aggregate_temporal_noise_stats(&records, 1, 1, width, height) .expect("a static correlated-noise pair should clear the static-block floor"); let rel_err = (sample.sigma[0] - sigma_marginal).abs() / sigma_marginal; @@ -657,34 +642,32 @@ fn correlated_noise_pair_recovers_marginal_sigma_and_rho() { ); } -/// A horizontal ramp shifted by a few pixels stands in for motion. -/// -/// The steady per-pixel offset that introduces overwhelms the static -/// check almost everywhere, so most blocks, if not all of them, should -/// come out as non-static. +/// A ramp shifted by a few pixels stands in for motion, and its steady per-pixel offset fails the +/// static check almost everywhere. #[test] fn moving_content_pair_mostly_non_static() { - let w = 64; - let h = 64; + let width = 64; + let height = 64; let shift = 4u32; - let mut prev = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - prev[(y * w + x) as usize] = x as f32 / w as f32; + let mut previous = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + previous[(y * width + x) as usize] = x as f32 / width as f32; } } - let mut new = vec![0.0f32; (w * h) as usize]; - for y in 0..h { - for x in 0..w { - let xs = (x + shift).min(w - 1); - new[(y * w + x) as usize] = prev[(y * w + xs) as usize]; + + let mut new = vec![0.0f32; (width * height) as usize]; + for y in 0..height { + for x in 0..width { + let source_x = (x + shift).min(width - 1); + new[(y * width + x) as usize] = previous[(y * width + source_x) as usize]; } } - let records = run_temporal_stats(w, h, 1, &prev, &new, true); - let sample = aggregate_temporal_noise_stats(&records, 1, 1, w, h); - let static_fraction = sample.map(|s| s.static_fraction).unwrap_or(0.0); + let records = run_temporal_stats(width, height, 1, &previous, &new, true); + let sample = aggregate_temporal_noise_stats(&records, 1, 1, width, height); + let static_fraction = sample.map(|sample| sample.static_fraction).unwrap_or(0.0); assert!( static_fraction < 0.5, @@ -692,40 +675,33 @@ fn moving_content_pair_mostly_non_static() { ); } -/// The whole pipeline, fed a synthetic stream of correlated grain. -/// -/// The denoiser's estimator has to settle within 25% of the true sigma -/// once the correlation correction is applied. +/// Immerkær alone reads several times lower on the same correlated grain. /// -/// Immerkær alone reads several times lower on the same content, which -/// is the exact problem this estimator exists to fix. -/// -/// The blur's lag-1 correlation works out at two thirds, from the same -/// coefficients that give the 0.375 variance scale. +/// The blur's lag-1 correlation is two thirds, from the same taps that give the 0.375 variance scale. #[test] fn hq_temporal_folds_correlated_grain_above_immerkaer_alone() { let client = make_client(); - let w = 128; - let h = 128; + let width = 128; + let height = 128; let sigma_marginal = 8.0 / 255.0; let sigma_pre = sigma_marginal / 0.375f32.sqrt(); let base = 0.5f32; - let n_frames = 14; - let frames: Vec> = (0..n_frames) - .map(|i| correlated_noisy_frame(w, h, base, sigma_pre, 100 + i as u32)) + let frame_count = 14; + let frames: Vec> = (0..frame_count) + .map(|i| correlated_noisy_frame(width, height, base, sigma_pre, 100 + i as u32)) .collect(); - // What the old (Immerkær-only) estimator would read on this same - // correlated content, computed directly rather than through the - // denoiser. + // What Immerkær alone reads on this content, computed directly rather than through the denoiser. let immerkaer_only = { - let input_buf = client.create_from_slice(f32::as_bytes(&frames[0])); - let partials_buf = client.empty(partials_len(w, h) * size_of::()); + let input_bytes = f32::as_bytes(&frames[0]); + let partials_bytes = partials_len(width, height) * size_of::(); + let input_buf = client.create_from_slice(input_bytes); + let partials_buf = client.empty(partials_bytes); let results_buf = client.empty(4 * size_of::()); let ctx = NoiseCtx { - width: w, - height: h, + width, + height, channels: 1, stored_ch: 1, frame_count: 1, @@ -736,10 +712,12 @@ fn hq_temporal_folds_correlated_grain_above_immerkaer_alone() { results_buf: &results_buf, }; run_noise_estimate::(&client, &ctx).expect("immerkaer dispatch failed"); + let bytes = client.read_one(results_buf).expect("immerkaer readback failed"); let data = f32::from_bytes(&bytes); - sigma_from_abs_sum(data[0], w, h) + sigma_from_abs_sum(data[0], width, height) }; + assert!( immerkaer_only < sigma_marginal * 0.5, "expected Immerkær alone to read well below the marginal truth {sigma_marginal} \ @@ -766,7 +744,7 @@ fn hq_temporal_folds_correlated_grain_above_immerkaer_alone() { }), }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); for frame in &frames { denoiser.push_frame(frame); let _ = denoiser.denoise().unwrap(); @@ -786,19 +764,13 @@ fn hq_temporal_folds_correlated_grain_above_immerkaer_alone() { ); } -/// `rho_smoothed`'s first update must seed directly from the first -/// temporal sample rather than blend it against an assumed 0. On -/// correlated grain with true rho 2/3, blending from 0 with -/// `EMA_ALPHA = 0.2` would read about 0.13 after the very first -/// sample, an 80% relative error, while seeding directly should land -/// within the same 25% tolerance -/// `hq_temporal_folds_correlated_grain_above_immerkaer_alone` uses for -/// the folded sigma on the same content. +/// On grain with true rho 2/3, blending from 0 with `EMA_ALPHA = 0.2` would read about 0.13 after +/// the first sample, an 80% relative error. The 25% tolerance separates seeding from blending. #[test] fn rho_smoothed_seeds_from_first_sample_not_from_zero() { let client = make_client(); - let w = 128; - let h = 128; + let width = 128; + let height = 128; let sigma_marginal = 8.0 / 255.0; let sigma_pre = sigma_marginal / 0.375f32.sqrt(); let base = 0.5f32; @@ -824,10 +796,10 @@ fn rho_smoothed_seeds_from_first_sample_not_from_zero() { }), }; - let mut denoiser = NlmDenoiser::::new(&client, params, w, h); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); let mut first_seeded_rho = None; - for i in 0..(w.min(h)) { - let frame = correlated_noisy_frame(w, h, base, sigma_pre, 200 + i); + for i in 0..(width.min(height)) { + let frame = correlated_noisy_frame(width, height, base, sigma_pre, 200 + i); denoiser.push_frame(&frame); let _ = denoiser.denoise().unwrap(); if let Some(rho) = denoiser.rho_smoothed { @@ -847,56 +819,45 @@ fn rho_smoothed_seeds_from_first_sample_not_from_zero() { ); } -/// A frame split into a static top half and a panning bottom half. -/// -/// The top is flat content with independent measurement noise. The -/// bottom is fine texture panning a few pixels between frames, with the -/// same noise on top. -/// -/// The panning residual averages close to zero over a block, so the mean -/// check alone lets nearly the whole frame through. -/// -/// That is the case `moving_content_pair_mostly_non_static` does not -/// cover, because its ramp leaves a block-mean residual far above the -/// check rather than near zero. +/// The top half is static flat content and the bottom half is fine texture panning a few pixels, both +/// under the same noise. /// -/// Only the top half's variance is really noise, so the aggregation has -/// to recover close to the true sigma and reject most of the panning -/// half. +/// The panning residual averages near zero over a block, so the mean check alone lets nearly the +/// whole frame through. Only the top half is really noise, so the aggregation must reject most of +/// the bottom half. #[test] fn moving_texture_pair_near_zero_mean_residual_recovers_true_sigma() { - let w = 128; - let h = 128; + let width = 128; + let height = 128; let true_sigma = 2.0 / 255.0; let texture_sigma_pre = 24.0 / 255.0 / 0.375f32.sqrt(); let shift = 3u32; let base = 0.5f32; - // One texture field for the bottom half, independent of the - // measurement noise added below. - let raw_texture = correlated_noisy_frame(w, h / 2, base, texture_sigma_pre, 1); - - let mut clean_prev = vec![base; (w * h) as usize]; - let mut clean_new = vec![base; (w * h) as usize]; - for y in 0..(h / 2) { - for x in 0..w { - let prev_val = raw_texture[(y * w + x) as usize]; - let xs = (x + shift).min(w - 1); - let new_val = raw_texture[(y * w + xs) as usize]; - let out_y = h / 2 + y; - clean_prev[(out_y * w + x) as usize] = prev_val; - clean_new[(out_y * w + x) as usize] = new_val; + // One texture field for the bottom half, independent of the measurement noise added below. + let raw_texture = correlated_noisy_frame(width, height / 2, base, texture_sigma_pre, 1); + + let mut clean_previous = vec![base; (width * height) as usize]; + let mut clean_new = vec![base; (width * height) as usize]; + for y in 0..(height / 2) { + for x in 0..width { + let previous_value = raw_texture[(y * width + x) as usize]; + let source_x = (x + shift).min(width - 1); + let new_value = raw_texture[(y * width + source_x) as usize]; + let out_y = height / 2 + y; + clean_previous[(out_y * width + x) as usize] = previous_value; + clean_new[(out_y * width + x) as usize] = new_value; } } - let prev = noisy_field_over(&clean_prev, w, h, true_sigma, 11); - let new = noisy_field_over(&clean_new, w, h, true_sigma, 12); + let previous = noisy_field_over(&clean_previous, width, height, true_sigma, 11); + let new = noisy_field_over(&clean_new, width, height, true_sigma, 12); - let records = run_temporal_stats(w, h, 1, &prev, &new, true); - let sample = aggregate_temporal_noise_stats(&records, 1, 1, w, h) + let records = run_temporal_stats(width, height, 1, &previous, &new, true); + let sample = aggregate_temporal_noise_stats(&records, 1, 1, width, height) .expect("the static top half should clear STATIC_FRACTION_MIN on its own"); - let (blocks_x, blocks_y) = temporal_stats_blocks(w, h); + let (blocks_x, blocks_y) = temporal_stats_blocks(width, height); let total_blocks = blocks_x * blocks_y; let top_half_blocks = total_blocks / 2; assert!( diff --git a/av-denoise-core/src/nlmeans/tests/unpack_wire.rs b/av-denoise-core/src/nlmeans/tests/unpack_wire.rs deleted file mode 100644 index d930224..0000000 --- a/av-denoise-core/src/nlmeans/tests/unpack_wire.rs +++ /dev/null @@ -1,368 +0,0 @@ -use cubecl::prelude::*; - -#[cfg(feature = "vulkan")] -use super::helpers::ramp_frame; -use super::helpers::{R, make_client}; -#[cfg(feature = "vulkan")] -use super::pack_wire::denoiser; -use crate::Depth; -use crate::nlmeans::kernels::gpu_unpack_wire; -#[cfg(feature = "vulkan")] -use crate::{ChannelMode, OutputFormat}; - -const BLOCK: u32 = 256; - -/// Fills every destination slot before the launch, so a slot the kernel -/// never writes is visible instead of reading as a plausible zero. -const SENTINEL: f32 = 1.0; - -/// Runs the kernel over `wire` and returns the `f32` frame it wrote, -/// including any padding lanes. -fn unpack(wire: &[u8], pixels: u32, channels: u32, stored_ch: u32, depth: Depth) -> Vec { - let elements = pixels * stored_ch; - let out = launch(wire, pixels, channels, stored_ch, depth, None, 0, elements); - out[..elements as usize].to_vec() -} - -/// The same, with the grid forced so a test can make every thread walk -/// the strided loop more than once. -fn unpack_with_grid( - wire: &[u8], - pixels: u32, - channels: u32, - stored_ch: u32, - depth: Depth, - grid: u32, -) -> Vec { - let elements = pixels * stored_ch; - let out = launch(wire, pixels, channels, stored_ch, depth, Some(grid), 0, elements); - out[..elements as usize].to_vec() -} - -/// Runs one launch and returns the whole destination buffer, sentinel -/// slots included. -#[expect( - clippy::too_many_arguments, - reason = "mirrors the kernel's own argument list, plus the launch geometry" -)] -fn launch( - wire: &[u8], - pixels: u32, - channels: u32, - stored_ch: u32, - depth: Depth, - grid_override: Option, - dst_offset: u32, - dst_len: u32, -) -> Vec { - let client = make_client(); - - let pack = depth.wire_pack(); - let elements = pixels * stored_ch; - let grid = grid_override.unwrap_or_else(|| elements.div_ceil(BLOCK).max(1)); - let total_threads = grid * BLOCK; - - // The kernel reads whole words, so a plane that ends mid-word needs - // its last word backed by real storage. - let mut padded = wire.to_vec(); - padded.resize(wire.len().div_ceil(4) * 4, 0); - - let seed = vec![SENTINEL; dst_len as usize]; - - let src = client.create_from_slice(&padded); - let dst = client.create_from_slice(f32::as_bytes(&seed)); - - unsafe { - gpu_unpack_wire::launch_unchecked::( - &client, - CubeCount::new_1d(grid), - CubeDim::new_1d(BLOCK), - ArrayArg::from_raw_parts(src, padded.len() / 4), - ArrayArg::from_raw_parts(dst.clone(), dst_len as usize), - pack.max(), - dst_offset, - pixels, - channels, - stored_ch, - pack.samples_per_word(), - elements, - total_threads, - ); - } - - let out = client.read_one(dst).expect("unpack readback failed"); - f32::from_bytes(&out)[..dst_len as usize].to_vec() -} - -/// Turns normalised samples into the wire bytes the kernel reads. -fn wire_bytes(samples: &[f32], depth: Depth) -> Vec { - crate::frame::f32_to_plane(samples, depth) -} - -/// The GPU divide is not correctly rounded, so a normalised sample can -/// land one representable step either side of the host's value. -fn within_one_ulp(a: f32, b: f32) -> bool { - (a.to_bits() as i64 - b.to_bits() as i64).abs() <= 1 -} - -/// Compares whole frames, reporting the first sample that drifts further -/// than the divide can account for. -fn assert_frame(got: &[f32], want: &[f32], what: &str) { - assert_eq!(got.len(), want.len(), "length mismatch, {what}"); - for (i, (&g, &w)) in got.iter().zip(want).enumerate() { - assert!(within_one_ulp(g, w), "{what}, sample {i}, got {g} want {w}"); - } -} - -/// A ramp over the whole range, so every test covers both ends. -fn ramp(pixels: u32) -> Vec { - (0..pixels).map(|i| i as f32 / (pixels - 1) as f32).collect() -} - -/// The only check that can catch a silently-dead kernel, since a kernel -/// compared against itself compares zeros to zeros. -#[test] -fn luma_matches_the_host_converter_at_every_depth() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let pixels = 64u32; - let wire = wire_bytes(&ramp(pixels), depth); - - let got = unpack(&wire, pixels, 1, 1, depth); - let want = crate::frame::plane_to_f32(&wire, depth); - - assert_frame(&got, &want, &format!("depth {depth:?}")); - } -} - -#[test] -fn chroma_matches_interleave_uv_from_the_host_at_every_depth() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let pixels = 32u32; - let u = ramp(pixels); - let v: Vec = u.iter().map(|s| 1.0 - s).collect(); - - let u_wire = wire_bytes(&u, depth); - let v_wire = wire_bytes(&v, depth); - let mut wire = u_wire.clone(); - wire.extend_from_slice(&v_wire); - - let got = unpack(&wire, pixels, 2, 2, depth); - let want = crate::frame::interleave_uv_to_f32(&u_wire, &v_wire, depth); - - assert_frame(&got, &want, &format!("depth {depth:?}")); - } -} - -/// The padding lane starts at [`SENTINEL`], so a kernel that leaves lane -/// 3 alone fails here rather than passing on a zeroed allocation. -#[test] -fn yuv_matches_the_host_and_zeroes_the_padding_lane() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let pixels = 16u32; - let make = |off: f32| -> Vec { - (0..pixels) - .map(|i| ((i as f32 / (pixels - 1) as f32) * 0.5 + off).clamp(0.0, 1.0)) - .collect() - }; - let (y, u, v) = (make(0.0), make(0.25), make(0.5)); - let (yw, uw, vw) = ( - wire_bytes(&y, depth), - wire_bytes(&u, depth), - wire_bytes(&v, depth), - ); - - let mut wire = yw.clone(); - wire.extend_from_slice(&uw); - wire.extend_from_slice(&vw); - - // Yuv stores 4 lanes per pixel and fills 3. - let got = unpack(&wire, pixels, 3, 4, depth); - let want = crate::frame::interleave_yuv_to_f32(&yw, &uw, &vw, depth); - - for p in 0..pixels as usize { - for c in 0..3usize { - let (g, w) = (got[p * 4 + c], want[p * 3 + c]); - assert!( - within_one_ulp(g, w), - "depth {depth:?} pixel {p} channel {c}, got {g} want {w}" - ); - } - assert_eq!(got[p * 4 + 3], 0.0, "the padding lane must be zero"); - } - } -} - -/// A ring slot other than the first. Without this the kernel writes every -/// frame over slot 0 and every other test still passes. -#[test] -fn a_non_zero_dst_offset_writes_its_own_slot_and_leaves_the_others_alone() { - for depth in [Depth::Eight, Depth::Ten] { - let pixels = 64u32; - let wire = wire_bytes(&ramp(pixels), depth); - - let got = launch(&wire, pixels, 1, 1, depth, None, pixels, pixels * 2); - let want = crate::frame::plane_to_f32(&wire, depth); - - assert_frame(&got[pixels as usize..], &want, &format!("depth {depth:?}")); - assert_eq!( - &got[..pixels as usize], - &vec![SENTINEL; pixels as usize][..], - "depth {depth:?}, the slot below must be untouched" - ); - } -} - -#[test] -fn a_sample_count_that_is_not_a_whole_number_of_words_reads_its_tail() { - // 13 samples at 8-bit is three whole words plus one byte, and at 10 - // and 12-bit it is six whole words plus two bytes. - let pixels = 13u32; - - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let wire = wire_bytes(&ramp(pixels), depth); - - let got = unpack(&wire, pixels, 1, 1, depth); - let want = crate::frame::plane_to_f32(&wire, depth); - - assert_frame(&got, &want, &format!("depth {depth:?}")); - } -} - -/// Forces a grid far smaller than the element count so every thread runs -/// the strided loop several times. Without this the loop body runs at most -/// once and `idx += total_threads` is never exercised. -#[test] -fn the_strided_loop_covers_every_element_when_the_grid_is_small() { - let pixels = 1024u32; - - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let samples: Vec = (0..pixels).map(|i| (i % 251) as f32 / 250.0).collect(); - let wire = wire_bytes(&samples, depth); - - let got = unpack_with_grid(&wire, pixels, 1, 1, depth, 1); - let want = crate::frame::plane_to_f32(&wire, depth); - - assert_frame(&got, &want, &format!("depth {depth:?}")); - } -} - -/// One frame's worth of wire planes, plus the densely interleaved `f32` -/// frame holding exactly the same quantised values. -/// -/// Both sides start from the wire bytes, so the comparison turns on the -/// kernel rather than on host rounding. -#[cfg(feature = "vulkan")] -fn wire_and_f32_frame(w: u32, h: u32, channels: usize, i: usize, depth: Depth) -> (Vec>, Vec) { - let pixels = (w * h) as usize; - - let wire: Vec> = (0..channels) - .map(|c| crate::frame::f32_to_plane(&ramp_frame(w, h, i * channels + c), depth)) - .collect(); - - let normalised: Vec> = wire - .iter() - .map(|plane| crate::frame::plane_to_f32(plane, depth)) - .collect(); - - let mut dense = Vec::with_capacity(pixels * channels); - for p in 0..pixels { - for plane in &normalised { - dense.push(plane[p]); - } - } - - (wire, dense) -} - -/// The quantised codes a frame lands on at `depth`. -/// -/// The comparison runs in codes rather than raw `f32`, since one code is -/// the smallest difference the output can carry. -#[cfg(feature = "vulkan")] -fn codes(frame: &[f32], depth: Depth) -> Vec { - let wire = crate::frame::f32_to_plane(frame, depth); - match depth.bytes_per_sample() { - 1 => wire.iter().map(|&b| u32::from(b)).collect(), - _ => wire - .as_chunks::<2>() - .0 - .iter() - .map(|&s| u32::from(u16::from_le_bytes(s))) - .collect(), - } -} - -/// The differential test the whole push path rests on. A kernel that -/// silently compiled to nothing writes zeros, which the `f32` push of a -/// real frame never produces. -/// -/// The GPU divide is not correctly rounded, so a normalised sample can -/// land one representable step either side of the host's value. That -/// moves a denoised sample by at most one code and never further. -#[cfg(feature = "vulkan")] -#[test] -fn a_wire_push_denoises_within_one_code_of_an_f32_push() { - let (w, h) = (16u32, 16u32); - let modes = [ - (ChannelMode::Luma, 1usize), - (ChannelMode::Chroma, 2), - (ChannelMode::Yuv, 3), - ]; - - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - for (mode, channels) in modes { - let mut f32_side = denoiser(mode, OutputFormat::F32, w, h); - let mut wire_side = denoiser(mode, OutputFormat::F32, w, h); - - for i in 0..3 { - let (wire, dense) = wire_and_f32_frame(w, h, channels, i, depth); - let planes: Vec<&[u8]> = wire.iter().map(Vec::as_slice).collect(); - - f32_side.push_frame(&dense).expect("f32 push failed"); - wire_side - .push_frame_wire(&planes, depth) - .expect("wire push failed"); - } - - let want = f32_side - .recv_frame() - .expect("f32 recv failed") - .expect("a frame is ready") - .into_f32() - .expect("an f32 denoiser returns f32"); - - let got = wire_side - .recv_frame() - .expect("wire recv failed") - .expect("a frame is ready") - .into_f32() - .expect("an f32 denoiser returns f32"); - - let want = codes(&want, depth); - let got = codes(&got, depth); - - assert_eq!(got.len(), want.len(), "depth {depth:?} mode {mode:?}"); - for (i, (&g, &w)) in got.iter().zip(&want).enumerate() { - let drift = g.abs_diff(w); - assert!( - drift <= 1, - "depth {depth:?} mode {mode:?}, sample {i} moved {drift} codes, \ - got {g} want {w}" - ); - } - } - } -} - -#[cfg(feature = "vulkan")] -#[test] -#[should_panic(expected = "plane count mismatch")] -fn a_plane_count_that_disagrees_with_the_channel_mode_is_rejected() { - let (w, h) = (16u32, 16u32); - let depth = Depth::Eight; - - let mut d = denoiser(ChannelMode::Luma, OutputFormat::F32, w, h); - let plane = crate::frame::f32_to_plane(&ramp_frame(w, h, 0), depth); - - let _ = d.push_frame_wire(&[&plane, &plane], depth); -} diff --git a/av-denoise-core/src/nlmeans/tests/util.rs b/av-denoise-core/src/nlmeans/tests/util.rs index d15e8cf..43c2043 100644 --- a/av-denoise-core/src/nlmeans/tests/util.rs +++ b/av-denoise-core/src/nlmeans/tests/util.rs @@ -1,27 +1,15 @@ use super::helpers::*; +use crate::bench_api::HostIo; use crate::nlmeans::*; -#[test] -fn normalization_round_trips_across_the_full_range() { - for depth in [Depth::Eight, Depth::Ten, Depth::Twelve] { - let max = depth.max_value() as u16; - let original: Vec = (0..=max).collect(); - let restored = denormalize(&normalize(&original, depth), depth); - - assert_eq!(original, restored, "round trip failed at {depth:?}"); - } -} - -/// Drives most weights toward zero by combining a near-maximum search -/// radius with extremely low strength (large `h2_inv_norm`), and on -/// noisy content so denominators land near the underflow guard in -/// `nlm_finish`. The output must contain no `inf`/`nan` regardless. +/// Very low strength on noisy content drives most weights toward zero, so the weight sums land near +/// the underflow guard in `nlm_finish`. #[test] fn extreme_params_produce_finite_output() { let client = make_client(); - let w = 32; - let h = 32; - let frame = make_frame_with_noisy_region(w, h, 1, 0.1, 16, 16, 5, 0.9); + let width = 32; + let height = 32; + let frame = make_frame_with_noisy_region(width, height, 1, 0.1, 16, 16, 5, 0.9); let params = NlmParams { temporal_radius: 0, @@ -35,18 +23,15 @@ fn extreme_params_produce_finite_output() { hq: None, }; - let mut d = NlmDenoiser::::new(&client, params, w, h); - d.push_frame(&frame); - let result = d - .denoise() - .unwrap() - .unwrap() - .as_f32() - .expect("f32 denoiser") - .to_vec(); + let mut denoiser = NlmDenoiser::::new(&client, params, width, height); + denoiser.push_frame(&frame); + let result = denoiser.denoise().unwrap().unwrap(); - for (i, &v) in result.iter().enumerate() { - assert!(v.is_finite(), "pixel {i}: non-finite output {v}"); - assert!((-0.01..=1.01).contains(&v), "pixel {i}: out-of-range output {v}"); + for (i, &value) in result.iter().enumerate() { + assert!(value.is_finite(), "pixel {i}: non-finite output {value}"); + assert!( + (-0.01..=1.01).contains(&value), + "pixel {i}: out-of-range output {value}" + ); } } diff --git a/av-denoise-core/src/options.rs b/av-denoise-core/src/options.rs new file mode 100644 index 0000000..6ae2208 --- /dev/null +++ b/av-denoise-core/src/options.rs @@ -0,0 +1,27 @@ +/// Speed vs quality dial. +/// +/// Each denoising family reads the same dial and fills in its own knobs +/// from it. For `nlmeans` that is +/// [nlmeans_variant_for](crate::nlmeans_variant_for), +/// [nlmeans_temporal_radius_for](crate::nlmeans_temporal_radius_for), and +/// [nlmeans_search_radius_for](crate::nlmeans_search_radius_for). +/// For `nl4d` it is [nl4d_temporal_radius_for](crate::nl4d_temporal_radius_for) and +/// [nl4d_spatial_radius_for](crate::nl4d_spatial_radius_for). +/// +/// Both front ends parse the same names from this one type, so a preset +/// resolves to the same dials everywhere it is used. +#[derive(Debug, Copy, Clone, Default, PartialEq, Eq, strum_macros::EnumString)] +#[strum(ascii_case_insensitive)] +pub enum Preset { + /// Fastest and lowest quality. + Veryfast, + /// One step up from `veryfast`. + Fast, + /// The default, favouring quality over speed. + #[default] + Base, + /// One step down from `veryslow`. + Slow, + /// Slowest and highest quality. + Veryslow, +} diff --git a/av-denoise-core/src/probe.rs b/av-denoise-core/src/probe.rs deleted file mode 100644 index 753de30..0000000 --- a/av-denoise-core/src/probe.rs +++ /dev/null @@ -1,126 +0,0 @@ -//! Opening a backend's client without taking the process down with it. -//! -//! A build can enable a backend whose driver libraries are not -//! installed. Some backends do not report that as an error. The CUDA -//! runtime loads `libcuda` dynamically on its own worker thread and -//! panics there when the load fails, and the panic reaches the caller -//! as a second panic when cubecl unwraps the dead worker's channel. -//! -//! [`open_client`] runs that work under [`catch_unwind`], so a missing -//! driver reads as "this backend is not available here" rather than as -//! a crash. That is what makes a single binary with `cuda`, `rocm`, and -//! `vulkan` all enabled usable on a machine that has only one of them. - -use std::panic::{self, AssertUnwindSafe}; -use std::sync::Mutex; - -use cubecl::client::ComputeClient; -use cubecl::prelude::*; - -use crate::accelerate::Accelerator; - -/// Backends already reported as unavailable, and the lock guarding the -/// panic hook. -/// -/// The hook is process-wide, so two probes running at once would race to -/// restore each other's. Holding this for the length of a probe keeps -/// them in single file, and the list inside it keeps a backend from -/// warning again every time it is probed. -static PROBED: Mutex> = Mutex::new(Vec::new()); - -/// Opens a client for `accelerator` on `device`, or reports that the -/// backend cannot run here. -/// -/// The client is synchronised before it is handed back. cubecl kernels -/// are fully asynchronous, so a successful `sync()` is what proves the -/// backend works, and no test kernel is needed. -pub(crate) fn open_client( - accelerator: Accelerator, - device: &R::Device, -) -> Option> { - let mut probed = PROBED.lock().unwrap_or_else(|err| err.into_inner()); - - let opened = quiet_panics(|| { - let client = R::client(device); - cubecl::future::block_on(client.sync()).map(|()| client) - }); - - match opened { - Ok(Ok(client)) => Some(client), - Ok(Err(err)) => { - tracing::debug!(err = ?err, "could not use the {accelerator} runtime"); - None - }, - Err(_) => { - // Only the first probe of a backend says anything. A denoise - // run probes once per denoiser it builds, and a missing - // driver is worth one line, not one per scene. - if !probed.contains(&accelerator) { - probed.push(accelerator); - tracing::warn!( - "the {accelerator} backend is enabled but did not start, its driver libraries are probably missing" - ); - } - None - }, - } -} - -/// Runs `f`, turning a panic into an `Err` and routing the panic message -/// to the debug log rather than to stderr. -/// -/// A failing backend prints its own panic from its worker thread before -/// the caller ever sees one, so the hook is quietened for as long as `f` -/// runs and put back afterwards. -/// -/// The hook is process-wide. Callers hold [`PROBED`] across this so two -/// probes cannot race to restore each other's. -fn quiet_panics(f: impl FnOnce() -> T) -> std::thread::Result { - let previous = panic::take_hook(); - panic::set_hook(Box::new(|info| tracing::debug!("{info}"))); - let out = panic::catch_unwind(AssertUnwindSafe(f)); - panic::set_hook(previous); - out -} - -#[cfg(test)] -mod tests { - use std::sync::Arc; - use std::sync::atomic::{AtomicBool, Ordering}; - - use super::*; - - #[test] - fn a_panic_inside_becomes_an_error() { - let _probed = PROBED.lock().unwrap_or_else(|err| err.into_inner()); - - assert!(quiet_panics(|| panic!("the backend fell over")).is_err()); - assert_eq!(quiet_panics(|| 7).unwrap(), 7); - } - - /// The hook has to come back however the probe ended, or every later - /// panic in the process reports at debug level. - /// - /// Holding [`PROBED`], the way [`open_client`] does, keeps a probe on - /// another thread from swapping the hook mid-test. - #[test] - fn the_panic_hook_is_restored() { - let _probed = PROBED.lock().unwrap_or_else(|err| err.into_inner()); - - let marker = Arc::new(AtomicBool::new(false)); - let flag = marker.clone(); - panic::set_hook(Box::new(move |_| flag.store(true, Ordering::SeqCst))); - - let _ = quiet_panics(|| panic!("swallowed by the quiet hook")); - assert!( - !marker.load(Ordering::SeqCst), - "the quiet hook did not replace the installed one", - ); - - let _ = panic::catch_unwind(|| panic!("seen by the restored hook")); - let restored = marker.load(Ordering::SeqCst); - let _ = panic::take_hook(); - - assert!(restored, "the probe left its own panic hook installed"); - } -} diff --git a/av-denoise-core/src/sniff.rs b/av-denoise-core/src/sniff.rs deleted file mode 100644 index 895cb3d..0000000 --- a/av-denoise-core/src/sniff.rs +++ /dev/null @@ -1,97 +0,0 @@ -//! Picking a backend that actually works on this machine. -//! -//! A build can enable several accelerators, but only some of them will -//! run on any given machine. A CUDA build on a machine with no Nvidia -//! driver, for instance, compiles fine and then fails at runtime. -//! -//! [`sniff_best_accelerator`] takes a list in order of preference and -//! returns the first one that really starts on the device the caller -//! asked for, so callers can list a fast backend followed by a safe -//! fallback. -//! -//! The probe runs on that device rather than on the backend's default -//! one. Opening a client is what proves a backend works, and opening it -//! on a card the caller did not choose both tests the wrong hardware and -//! pays that card's first-time driver initialisation. -//! -//! [`Denoiser::create`](crate::Denoiser::create) calls this for you, so -//! reach for it directly only when you want to know the answer without -//! building a denoiser. -//! -//! ```no_run -//! use av_denoise_core::Device; -//! use av_denoise_core::accelerate::get_default_accelerators; -//! use av_denoise_core::sniff::sniff_best_accelerator; -//! -//! match sniff_best_accelerator(&get_default_accelerators(), &Device::Default) { -//! Some(accelerator) => println!("running on {accelerator}"), -//! None => println!("no usable backend on this machine"), -//! } -//! ``` - -use cubecl::prelude::*; - -use crate::accelerate::Accelerator; -use crate::device::Device; -use crate::probe::open_client; - -/// Tries each accelerator in turn and returns the first one whose client -/// can be built and synchronised on `device`. -/// -/// cubecl kernels are fully asynchronous, so a successful -/// `client.sync()` is enough to prove the backend works. No test kernel -/// is needed. -/// -/// An accelerator that cannot express `device` at all is treated as -/// unavailable and the search moves on. That is the same answer the -/// caller would get from trying to build on it, one step earlier. So is -/// a backend whose driver libraries are missing, however loudly it -/// fails, see [`probe`](crate::probe). -pub fn sniff_best_accelerator(enable: &[Accelerator], device: &Device) -> Option { - for accelerator in enable { - let is_enabled = match accelerator { - #[cfg(feature = "cuda")] - Accelerator::Cuda => match device.to_cuda() { - Ok(dev) => probe_runtime::(*accelerator, &dev), - Err(_) => false, - }, - #[cfg(feature = "rocm")] - Accelerator::Rocm => match device.to_amd() { - Ok(dev) => probe_runtime::(*accelerator, &dev), - Err(_) => false, - }, - #[cfg(feature = "vulkan")] - Accelerator::Vulkan => match device.to_wgpu() { - Ok(dev) => probe_runtime::(*accelerator, &dev), - Err(_) => false, - }, - #[cfg(feature = "metal")] - Accelerator::Metal => match device.to_wgpu() { - Ok(dev) => probe_runtime::(*accelerator, &dev), - Err(_) => false, - }, - // docs.rs widens the `Accelerator` variants behind - // `cfg(docsrs)` so they all appear in the rendered enum, even - // when the matching backend feature is off. This arm keeps - // the match exhaustive there and is never reached at - // runtime. - #[cfg(docsrs)] - #[expect( - unreachable_patterns, - reason = "the arm only keeps the match exhaustive on docs.rs" - )] - _ => unreachable!(), - }; - - if is_enabled { - return Some(*accelerator); - } - } - - None -} - -/// Reports whether `accelerator` can really run on `device`. -fn probe_runtime(accelerator: Accelerator, device: &R::Device) -> bool { - open_client::(accelerator, device).is_some() -} diff --git a/av-denoise-core/src/stack.rs b/av-denoise-core/src/stack.rs deleted file mode 100644 index 38d232a..0000000 --- a/av-denoise-core/src/stack.rs +++ /dev/null @@ -1,76 +0,0 @@ -use std::ffi::OsStr; - -/// Stack bytes the kernel codegen thread needs. -pub const CODEGEN_STACK_BYTES: usize = 16 << 20; - -/// Raises `RUST_MIN_STACK` to [`CODEGEN_STACK_BYTES`] when it is unset. -/// -/// No-op if the variable is already set. -/// -/// # Safety -/// -/// The same safety rules as [std::env::set_var] applies. -pub unsafe fn raise_codegen_stack_limit() { - if std::env::var_os("RUST_MIN_STACK").is_none() { - // SAFETY: forwarded from this function's own precondition. - unsafe { std::env::set_var("RUST_MIN_STACK", CODEGEN_STACK_BYTES.to_string()) }; - } -} - -/// Whether the process's `RUST_MIN_STACK` is large enough for codegen. -pub fn codegen_stack_is_sufficient() -> bool { - limit_is_sufficient(std::env::var_os("RUST_MIN_STACK").as_deref()) -} - -/// Parses a raw `RUST_MIN_STACK` value and checks it against [`CODEGEN_STACK_BYTES`]. -/// -/// Unset is insufficient. A value that fails to parse is also insufficient. -fn limit_is_sufficient(raw: Option<&OsStr>) -> bool { - raw.and_then(|v| v.to_str()) - .and_then(|v| v.parse::().ok()) - .is_some_and(|bytes| bytes >= CODEGEN_STACK_BYTES) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn unset_is_insufficient() { - assert!(!limit_is_sufficient(None)); - } - - #[test] - fn zero_is_insufficient() { - assert!(!limit_is_sufficient(Some(OsStr::new("0")))); - } - - #[test] - fn one_below_the_limit_is_insufficient() { - let value = (CODEGEN_STACK_BYTES - 1).to_string(); - assert!(!limit_is_sufficient(Some(OsStr::new(&value)))); - } - - #[test] - fn exactly_the_limit_is_sufficient() { - let value = CODEGEN_STACK_BYTES.to_string(); - assert!(limit_is_sufficient(Some(OsStr::new(&value)))); - } - - #[test] - fn above_the_limit_is_sufficient() { - let value = (CODEGEN_STACK_BYTES * 2).to_string(); - assert!(limit_is_sufficient(Some(OsStr::new(&value)))); - } - - #[test] - fn surrounding_whitespace_is_insufficient() { - let value = format!(" {CODEGEN_STACK_BYTES} "); - assert!(!limit_is_sufficient(Some(OsStr::new(&value)))); - } - - #[test] - fn non_numeric_is_insufficient() { - assert!(!limit_is_sufficient(Some(OsStr::new("not-a-number")))); - } -} diff --git a/av-denoise-core/src/warmup.rs b/av-denoise-core/src/warmup.rs deleted file mode 100644 index ecb6c7c..0000000 --- a/av-denoise-core/src/warmup.rs +++ /dev/null @@ -1,412 +0,0 @@ -//! Lets one process fill a cold kernel cache while the others wait. -//! -//! CubeCL compiles a kernel the first time it is dispatched, and writes -//! the result to the cache [`install_compilation_cache`] points it at. -//! The table of contents that cache reads is a snapshot taken when the -//! GPU client is built, so a process that starts while a second process -//! is still compiling sees an empty table and shares nothing with it. -//! Every process in that group pays the full compilation cost and -//! appends its own copy of the same kernels. -//! -//! [`WarmUp`] closes that window with a lock file next to the cache. -//! The first process in holds it while it compiles, the rest block, and -//! by the time they build their own client the cache has everything they -//! need. Once a run finishes it leaves a stamp file behind, and later -//! processes see the stamp and skip the lock entirely, so the cost is -//! paid once rather than on every chunk. -//! -//! [`install_compilation_cache`]: crate::install_compilation_cache - -use std::collections::HashSet; -use std::fs::{File, OpenOptions, TryLockError}; -use std::path::{Path, PathBuf}; -use std::sync::{LazyLock, Mutex}; -use std::time::{Duration, Instant}; - -use cubecl::hash::StableHasher; - -use crate::cache::compilation_cache_dir; -use crate::frame::{FrameLayout, PlaneOptions}; - -/// How long a process waits for the one ahead of it before giving up and -/// compiling for itself. -/// -/// Compiling this crate's kernels takes about ten seconds on a quiet -/// machine, and a machine running an encode is not quiet, so the limit -/// is generous. Waiting longer than this is worse than duplicating the -/// work, because the encoder above has nothing to do until a frame -/// arrives. -const WAIT_LIMIT: Duration = Duration::from_secs(180); - -/// How often a waiting process retries the lock. -const POLL_INTERVAL: Duration = Duration::from_millis(100); - -/// The keys this process already holds a place for. -static CLAIMED_KEYS: LazyLock>> = LazyLock::new(|| Mutex::new(HashSet::new())); - -/// Records that this process is taking the place for `key`, and reports -/// whether it was free to take. -fn claim_key(key: u128) -> bool { - CLAIMED_KEYS - .lock() - .expect("warm-up key mutex poisoned") - .insert(key) -} - -/// Gives `key` back, so a later filter in this process can queue for it -/// again. -fn release_key(key: u128) { - CLAIMED_KEYS - .lock() - .expect("warm-up key mutex poisoned") - .remove(&key); -} - -/// Identifies the set of kernels a denoiser compiles. -/// -/// Radii, channel mode, depth and algorithm are all baked into the -/// kernels at compile time, so a change to any of them produces a -/// different set and a warm cache for one set says nothing about -/// another. Processes that would compile different kernels have no -/// reason to wait for each other, and a stamp left by one must not -/// convince the other that its own kernels are cached. -/// -/// The key is taken from the `Debug` rendering of both inputs rather -/// than from a hand-written list of the fields that reach a kernel. -/// A hand-written list silently stops covering a field the day someone -/// adds one, and the failure that follows is a process trusting a stamp -/// for kernels it has never compiled. Reading everything keeps the key -/// correct without maintenance. -/// -/// Including fields that only ever reach the GPU as runtime arguments, -/// such as strength, makes the key finer than it strictly needs to be. -/// Two runs that differ only in strength warm up separately instead of -/// sharing. That costs one extra warm-up and is the safe direction to -/// err in. -/// -/// This crate's version goes into the key as well, because a release can -/// upgrade CubeCL, which files its own cache under the CubeCL version. -/// Without the version a stamp written before an upgrade would tell every -/// process after it that a cache emptied by that upgrade is warm, and the -/// whole first wave would compile at once again with nothing said about -/// it. A rebuild of the kernel sources needs nothing here, because the -/// stamps live in that build's own cache directory. -pub fn kernel_key(options: &PlaneOptions, layout: FrameLayout) -> u128 { - StableHasher::hash_one(&format!("{}|{options:?}|{layout:?}", env!("CARGO_PKG_VERSION"))) -} - -/// A held place in the queue to fill a cold cache. -/// -/// Obtained from [`WarmUp::begin`] and given up with [`WarmUp::finish`] -/// once the kernels are compiled. Dropping one without calling `finish` -/// releases the lock without leaving a stamp, so a run that failed part -/// way through does not convince the next process that the cache is -/// warm. -#[derive(Debug)] -pub struct WarmUp { - lock: File, - stamp: PathBuf, - key: u128, -} - -impl WarmUp { - /// Takes a place in the queue for the kernels `key` identifies. - /// - /// Returns `Some` while holding the lock, and the caller compiles - /// under it. Returns `None` when there is nothing to wait for, which - /// covers a cache that is already warm for these kernels, a cache - /// that is turned off, and a lock that could not be taken in - /// [`WAIT_LIMIT`]. In every one of those the caller carries on and - /// compiles as it always did. - /// - /// Blocks for as long as the process ahead takes to compile, so it - /// belongs on the path that builds a denoiser rather than on the - /// path that renders a frame. - pub fn begin(key: u128) -> Option { - Self::begin_in(compilation_cache_dir()?, key, WAIT_LIMIT) - } - - /// [`WarmUp::begin`] against an explicit directory and wait limit. - /// - /// Split out so tests can drive the queue without the process-wide - /// cache directory, and without waiting the full [`WAIT_LIMIT`] to - /// see what a contended lock does. - fn begin_in(dir: &Path, key: u128, wait_limit: Duration) -> Option { - // A file lock is held by the process rather than by the handle - // that took it, so a second filter in this process asking for - // the same kernels would wait out `wait_limit` on a lock its own - // process already holds. One script can easily build two - // filters, so take the place at most once per key per process - // and let the second caller carry on. - if !claim_key(key) { - return None; - } - - let held = Self::acquire(dir, key, wait_limit); - - if held.is_none() { - release_key(key); - } - - held - } - - /// [`WarmUp::begin_in`] without the in-process bookkeeping, which its - /// caller takes care of on both the success and the failure path. - fn acquire(dir: &Path, key: u128, wait_limit: Duration) -> Option { - let stamp = dir.join(format!("warm-{key:032x}.stamp")); - - if stamp.exists() { - return None; - } - - let lock = open_lock_file(&dir.join(format!("warm-{key:032x}.lock")))?; - - if !wait_for_lock(&lock, wait_limit) { - return None; - } - - // Whoever held the lock has finished compiling by now, so the - // stamp answers differently than it did above. Checking it again - // is what turns the queue into a single warm-up rather than one - // per waiting process. - if stamp.exists() { - let _ = lock.unlock(); - return None; - } - - tracing::debug!(?stamp, "compiling kernels for a cold cache"); - Some(Self { lock, stamp, key }) - } - - /// Records that the kernels are compiled and lets the next process - /// through. - pub fn finish(self) { - if let Err(err) = std::fs::write(&self.stamp, b"") { - // The next process reads a missing stamp as a cold cache and - // compiles again, which is slower but still correct. - tracing::debug!(stamp = ?self.stamp, %err, "cannot write the kernel warm-up stamp"); - } - } -} - -impl Drop for WarmUp { - fn drop(&mut self) { - let _ = self.lock.unlock(); - release_key(self.key); - } -} - -/// Opens the lock file, creating it when this is the first process to -/// ask for these kernels. -/// -/// A directory that cannot be written is reported and then ignored, the -/// same way an uncreatable cache directory is. Denoising works without -/// the queue, it just compiles more than once. -fn open_lock_file(path: &Path) -> Option { - match OpenOptions::new() - .create(true) - .read(true) - .write(true) - .truncate(false) - .open(path) - { - Ok(file) => Some(file), - Err(err) => { - tracing::debug!(?path, %err, "cannot open the kernel warm-up lock, compiling unqueued"); - None - }, - } -} - -/// Blocks until the lock is held, giving up after `wait_limit`. -/// -/// Returns whether the lock is held. The lock is a real advisory file -/// lock rather than a file whose presence means "taken", so the -/// operating system releases it when Av1an kills a worker mid-compile -/// and the next process in line wakes up straight away. -fn wait_for_lock(lock: &File, wait_limit: Duration) -> bool { - let start = Instant::now(); - - loop { - match lock.try_lock() { - Ok(()) => return true, - Err(TryLockError::WouldBlock) => {}, - Err(TryLockError::Error(err)) => { - tracing::debug!(%err, "cannot take the kernel warm-up lock, compiling unqueued"); - return false; - }, - } - - if start.elapsed() >= wait_limit { - tracing::warn!( - "waited {:?} for another process to compile kernels, compiling for ourselves", - wait_limit, - ); - return false; - } - - std::thread::sleep(POLL_INTERVAL); - } -} - -#[cfg(test)] -mod tests { - use super::*; - - /// A key no test shares with another, so the lock and stamp files of - /// one test cannot be seen by the next. - fn key(n: u128) -> u128 { - 0xa5a5_0000_0000_0000_0000_0000_0000_0000 + n - } - - /// Short enough that a contended lock fails the test quickly rather - /// than holding it up for the three real minutes. - const BRIEFLY: Duration = Duration::from_millis(50); - - /// The lock is taken here rather than through a second `begin_in`, - /// because a file lock is held per process and a second `begin_in` - /// would be answered by the in-process registry instead of by the - /// lock this is about. - #[test] - fn a_lock_held_elsewhere_keeps_this_process_out() { - let dir = tempfile::tempdir().unwrap(); - let path = dir.path().join(format!("warm-{:032x}.lock", key(1))); - let elsewhere = open_lock_file(&path).unwrap(); - elsewhere.lock().unwrap(); - - assert!( - WarmUp::begin_in(dir.path(), key(1), BRIEFLY).is_none(), - "a caller gives up rather than compiling alongside the process ahead", - ); - } - - #[test] - fn one_process_takes_one_place_per_key() { - let dir = tempfile::tempdir().unwrap(); - - let first = WarmUp::begin_in(dir.path(), key(6), BRIEFLY); - assert!(first.is_some(), "the first caller compiles"); - - assert!( - WarmUp::begin_in(dir.path(), key(6), BRIEFLY).is_none(), - "the second caller carries on rather than waiting for itself", - ); - } - - #[test] - fn a_released_place_can_be_taken_again() { - let dir = tempfile::tempdir().unwrap(); - - drop(WarmUp::begin_in(dir.path(), key(7), BRIEFLY)); - - assert!( - WarmUp::begin_in(dir.path(), key(7), BRIEFLY).is_some(), - "the key is free again once the place is given up", - ); - } - - #[test] - fn a_finished_warm_up_lets_the_next_process_straight_through() { - let dir = tempfile::tempdir().unwrap(); - - WarmUp::begin_in(dir.path(), key(2), BRIEFLY).unwrap().finish(); - - assert!( - WarmUp::begin_in(dir.path(), key(2), BRIEFLY).is_none(), - "a warm cache needs no queue", - ); - } - - #[test] - fn an_abandoned_warm_up_leaves_the_cache_cold() { - let dir = tempfile::tempdir().unwrap(); - - drop(WarmUp::begin_in(dir.path(), key(3), BRIEFLY)); - - assert!( - WarmUp::begin_in(dir.path(), key(3), BRIEFLY).is_some(), - "no stamp means the kernels still need compiling", - ); - } - - /// A `PlaneOptions` with no accelerator named, so that this test - /// module builds whichever backend feature is on. - fn options() -> PlaneOptions { - PlaneOptions { - accelerators: Vec::new(), - device: crate::Device::Default, - intent: crate::ChannelIntent::LumaChroma, - mode: crate::DenoisingMode::Temporal { radius: 2 }, - algorithm: crate::Algorithm::default(), - luma_strength: None, - chroma_strength: None, - luma_lambda_ht: None, - chroma_lambda_ht: None, - } - } - - fn layout() -> FrameLayout { - FrameLayout { - width: 1920, - height: 1080, - subsampling: crate::Subsampling::Yuv420, - depth: crate::Depth::Eight, - } - } - - #[test] - fn the_same_settings_give_the_same_key() { - assert_eq!(kernel_key(&options(), layout()), kernel_key(&options(), layout())); - } - - #[test] - fn a_different_depth_gives_a_different_key() { - let ten_bit = FrameLayout { - depth: crate::Depth::Ten, - ..layout() - }; - - assert_ne!(kernel_key(&options(), layout()), kernel_key(&options(), ten_bit)); - } - - #[test] - fn a_different_radius_gives_a_different_key() { - let wider = PlaneOptions { - mode: crate::DenoisingMode::Temporal { radius: 3 }, - ..options() - }; - - assert_ne!(kernel_key(&options(), layout()), kernel_key(&wider, layout())); - } - - #[test] - fn grain_export_gives_a_different_key() { - let exporting = PlaneOptions { - algorithm: crate::Algorithm::Nl4d(crate::Nl4dOptions { - grain_export: true, - ..crate::Nl4dOptions::default() - }), - ..options() - }; - let plain = PlaneOptions { - algorithm: crate::Algorithm::Nl4d(crate::Nl4dOptions::default()), - ..options() - }; - - assert_ne!(kernel_key(&plain, layout()), kernel_key(&exporting, layout())); - } - - #[test] - fn different_kernels_do_not_wait_for_each_other() { - let dir = tempfile::tempdir().unwrap(); - - let first = WarmUp::begin_in(dir.path(), key(4), BRIEFLY); - let second = WarmUp::begin_in(dir.path(), key(5), BRIEFLY); - - assert!( - first.is_some() && second.is_some(), - "separate keys queue separately" - ); - } -} diff --git a/av-denoise-vs/Cargo.toml b/av-denoise-vs/Cargo.toml index b8e150a..8a76803 100644 --- a/av-denoise-vs/Cargo.toml +++ b/av-denoise-vs/Cargo.toml @@ -10,15 +10,15 @@ description = "VapourSynth core low-level plugin for av-denoise providing fast, crate-type = ["cdylib", "rlib"] [dependencies] -av-denoise-core.workspace = true +av-denoise.workspace = true anyhow.workspace = true tracing.workspace = true vapoursynth = "0.5.6" tracing-subscriber = { version = "0.3", features = ["env-filter"] } [features] -vulkan = ["av-denoise-core/vulkan"] -metal = ["av-denoise-core/metal"] -cuda = ["av-denoise-core/cuda"] -rocm = ["av-denoise-core/rocm"] +vulkan = ["av-denoise/vulkan"] +metal = ["av-denoise/metal"] +cuda = ["av-denoise/cuda"] +rocm = ["av-denoise/rocm"] default = ["vulkan"] diff --git a/av-denoise-vs/src/filter.rs b/av-denoise-vs/src/filter.rs index 206c1d4..c85d131 100644 --- a/av-denoise-vs/src/filter.rs +++ b/av-denoise-vs/src/filter.rs @@ -1,14 +1,8 @@ -//! The `avd.NLMeans` and `avd.NL4D` filters. -//! -//! Both share one `Denoise` filter type and one GPU pipeline underneath. -//! They differ only in the [`AlgorithmKind`] their creation function -//! passes to [`plane_options_from`]. - use std::collections::HashMap; use std::sync::Mutex; use anyhow::{Error, Result, anyhow}; -use av_denoise_core::{EdgePadding, FrameLayout, PlanarDenoiser, Planes, ReseedWindow, WarmUp, WindowSpan}; +use av_denoise::{EdgePadding, FrameLayout, PlanarDenoiser, Planes, ReseedWindow, WarmUp, WindowSpan}; use vapoursynth::core::CoreRef; use vapoursynth::plugins::{Filter, FrameContext}; use vapoursynth::prelude::{API, FrameRef, FrameRefMut, Node, Property}; @@ -20,37 +14,25 @@ use crate::{init_logging, pin_plugin_library}; /// The running pipeline and the output frame it last produced. /// -/// VapourSynth may call `get_frame` from several threads, so the -/// pipeline sits behind a mutex and requests queue on it. One pipeline -/// is enough because the GPU is the bottleneck. +/// VapourSynth may call `get_frame` from several threads, so the pipeline sits behind a mutex and +/// requests queue on it. One pipeline is enough because the GPU is the bottleneck. struct State { denoiser: PlanarDenoiser, - /// The output frame index last served. - last: Option, - /// The cold-cache queue place this filter holds, until the first - /// frame proves the kernels are compiled and cached. - /// - /// CubeCL compiles a kernel when it is first dispatched rather than - /// when the denoiser is built, so the place has to be held across a - /// frame and cannot be given up at the end of creation. + last_served: Option, + /// The cold-cache queue place this filter holds until its first frame is rendered. /// - /// A process that builds the filter and then renders nothing keeps - /// the place until it exits, and other workers wait out the queue's - /// limit before compiling for themselves. That needs a long lived - /// process holding a node it never pulls a frame from, which is rare - /// enough to accept. + /// CubeCL compiles a kernel on its first dispatch rather than when the denoiser is built, so the + /// place is held across the first frame. A process that builds the filter and never pulls a frame + /// keeps the place until it exits, making other workers wait out the queue's limit before compiling + /// for themselves, which is rare enough to accept. warm_up: Option, /// Outputs from the clip's final flush that no request has taken yet. tail: Option, } impl State { - /// Gives up this filter's place in the cold-cache queue, now that a - /// frame has been through and every kernel it needs is compiled and - /// written to the cache. - /// - /// Does nothing after the first call, and nothing at all when the - /// filter never took a place. + /// Gives up this filter's place in the cold-cache queue once a frame has compiled and cached every + /// kernel it needs. fn finish_warm_up(&mut self) { if let Some(warm_up) = self.warm_up.take() { warm_up.finish(); @@ -58,32 +40,26 @@ impl State { } } -/// A denoising filter backed by one [`PlanarDenoiser`] pipeline. +/// A denoising filter backed by one [PlanarDenoiser] pipeline. /// -/// `avd.NLMeans` and `avd.NL4D` both build one of these, differing only -/// in the algorithm baked into `state.denoiser` at creation. +/// `avd.NLMeans` and `avd.NL4D` both build one, differing only in the algorithm the denoiser runs. pub struct Denoise<'core> { source: Node<'core>, layout: FrameLayout, - /// How many source frames a window at output frame `n` needs behind - /// and ahead of `n`, read once from - /// [`PlanarDenoiser::window_span`] at creation. nlmeans and nl4d - /// report different spans, so this is asked for rather than - /// assumed symmetric. + /// How many source frames a window needs behind and ahead of its output frame. + /// + /// nlmeans and nl4d report different spans, so this comes from [PlanarDenoiser::window_span] + /// rather than being assumed symmetric. span: WindowSpan, - /// The source clip's frame count, read once at creation. source_len: usize, state: Mutex, } impl<'core> Denoise<'core> { - /// Builds the shared parts of a `avd.NLMeans` or `avd.NL4D` filter. + /// Builds an `avd.NLMeans` or `avd.NL4D` filter. /// - /// Raises the stack limit first, since a `PlanarDenoiser` is created - /// below and cubecl only spawns its codegen thread once that - /// happens. Rejects the source's format and resolution before - /// touching the GPU, and rejects a variable-resolution source, which - /// `FrameLayout` has no way to represent. + /// Rejects the source's format and resolution before touching the GPU, including a + /// variable-resolution source, which [FrameLayout] has no way to represent. pub(crate) fn create( _api: API, _core: CoreRef<'core>, @@ -94,19 +70,18 @@ impl<'core> Denoise<'core> { init_logging(); pin_plugin_library(); - // `export_vapoursynth_plugin!` expands to the whole body of the - // plugin's entry point, so there is no earlier hook of ours to - // raise the stack limit in. - // SAFETY: best-effort mutation at the earliest hook this plugin - // gets. The host may already have other threads touching the - // environment, so this cannot guarantee exclusive access, but the - // alternative is a hard abort during codegen. - unsafe { av_denoise_core::raise_codegen_stack_limit() }; + // cubecl spawns its codegen thread when the denoiser below is created, so the stack limit is + // raised first. `export_vapoursynth_plugin!` owns the plugin's entry point, so there is no + // earlier hook of ours to do it in. + // SAFETY: best-effort mutation at the earliest hook this plugin gets. The host may already + // have other threads touching the environment, so this cannot guarantee exclusive access, but + // the alternative is a hard abort during codegen. + unsafe { av_denoise::raise_codegen_stack_limit() }; let info = source.info(); let (width, height) = match info.resolution { - Property::Constant(res) => (res.width as u32, res.height as u32), + Property::Constant(resolution) => (resolution.width as u32, resolution.height as u32), Property::Variable => { anyhow::bail!("clips with variable resolution are not supported"); }, @@ -124,54 +99,57 @@ impl<'core> Denoise<'core> { let layout = layout_from_format(raw_format, width, height)?; let plane_options = plane_options_from(raw, algorithm_kind, layout)?; - // Av1an runs one of these per chunk, so without a cache every - // chunk pays the ten seconds it takes to compile the kernels. - // The queue below keeps the first wave of chunks from all paying - // it at once. - av_denoise_core::install_compilation_cache_once(); - let warm_up = WarmUp::begin(av_denoise_core::kernel_key(&plane_options, layout)); + // Av1an runs one of these per chunk, so without a cache every chunk pays the ten seconds it + // takes to compile the kernels. The queue keeps the first wave of chunks from all paying it at + // once. + av_denoise::install_compilation_cache_once(); + let cache_key = av_denoise::kernel_key(&plane_options, layout); + let warm_up = WarmUp::begin(cache_key); let denoiser = PlanarDenoiser::create(&plane_options, layout)?; let span = denoiser.window_span(); + let state = State { + denoiser, + last_served: None, + warm_up, + tail: None, + }; + Ok(Self { source, layout, span, source_len: info.num_frames, - state: Mutex::new(State { - denoiser, - last: None, - warm_up, - tail: None, - }), + state: Mutex::new(state), }) } - /// The ordered source indices `reseed` or `reseed_window` needs for output frame `n`. + /// The ordered source indices `reseed` or `reseed_window` needs for an output frame. /// - /// `EdgePadding::Repeat` repeats the boundary frame so `reseed` sees the exact - /// window length it expects. `EdgePadding::Shifted` stops at either end of the - /// clip instead, with no repeats, for `reseed_window`. - fn window(&self, n: usize) -> Vec { + /// `EdgePadding::Repeat` repeats the boundary frame so `reseed` sees the exact window length it + /// expects. `EdgePadding::Shifted` stops at either end of the clip instead, with no repeats, for + /// `reseed_window`. + fn window(&self, output_index: usize) -> Vec { let last_frame = self.source_len - 1; match self.span.edges { - EdgePadding::Repeat => window_indices(n, self.span.behind, self.span.ahead, last_frame), + EdgePadding::Repeat => { + window_indices(output_index, self.span.behind, self.span.ahead, last_frame) + }, EdgePadding::Shifted => { - let range = shifted_window_range(n, self.span.behind, self.span.ahead, last_frame); + let range = shifted_window_range(output_index, self.span.behind, self.span.ahead, last_frame); range.collect() }, } } - /// The source indices output frame `n` needs, deduplicated so each - /// one is requested and fetched from VapourSynth only once. + /// The source indices an output frame needs, deduplicated so each one is requested and fetched + /// from VapourSynth only once. /// - /// Only `EdgePadding::Repeat` windows can repeat an index, at either end of the - /// clip. This is sorted since [`Self::window`] is already non-decreasing, so - /// sorting is a no-op kept for clarity. - fn unique_window(&self, n: usize) -> Vec { - let mut indices = self.window(n); + /// Only `EdgePadding::Repeat` windows can repeat an index, at either end of the clip. + /// [Self::window] is already non-decreasing, so the sort is a no-op kept for clarity. + fn unique_window(&self, output_index: usize) -> Vec { + let mut indices = self.window(output_index); indices.sort_unstable(); indices.dedup(); indices @@ -179,104 +157,108 @@ impl<'core> Denoise<'core> { /// Renders one output frame, applying the hybrid fast/rebuild policy. /// - /// A sequential request, straight after the last frame produced, pushes one - /// frame through the running stream. Under shifted edges, the request that - /// reaches the clip's end instead flushes the stream once, and the frames - /// after it are served from the tail cache. Anything else, including frame - /// 0, abandons the stream and rebuilds it from an explicit window, which - /// costs more but is correct from any starting point. - fn render(&self, n: usize, fetch: impl Fn(usize) -> Result) -> Result { + /// A sequential request, straight after the last frame produced, pushes one frame through the + /// running stream. Under shifted edges, the request that reaches the clip's end flushes the stream + /// once instead, and the frames after it are served from the tail cache. Anything else, including + /// frame 0, rebuilds the stream from an explicit window, which costs more but is correct from any + /// starting point. + fn render( + &self, + output_index: usize, + fetch: impl Fn(usize) -> Result, + ) -> Result { let mut state = self.state.lock().expect("denoiser mutex poisoned"); let last_frame = self.source_len - 1; let shifted = self.span.edges == EdgePadding::Shifted; - // Read the anchor, then clear it before anything touches the pipeline. - // - // Every path below either reaches a `state.last = Some(n)` or leaves through `?`, - // so an error out of `fetch`, `push`, `recv`, `flush`, or `reseed` can never leave - // the anchor claiming a position the stream has moved past. - let sequential = state.last == Some(n.wrapping_sub(1)) && n > 0; - state.last = None; - - let cached = state.tail.as_mut().and_then(|tail| tail.take(n)); - if let Some(out) = cached { - state.last = Some(n); - return Ok(out); + // The anchor is cleared before anything touches the pipeline. Every path below either sets it + // again or leaves through `?`, so an error can never leave it claiming a position the stream + // has moved past. + let sequential = state.last_served == Some(output_index.wrapping_sub(1)) && output_index > 0; + state.last_served = None; + + let cached = state.tail.as_mut().and_then(|tail| tail.take(output_index)); + if let Some(denoised) = cached { + state.last_served = Some(output_index); + return Ok(denoised); } state.tail = None; - let ahead = n + self.span.ahead; + let window_end = output_index + self.span.ahead; - // This request is the first whose window would run past the clip's last - // frame. The previous request pushed that last frame into a live stream, - // and no tail cache covers this index, so flushing the stream yields this - // frame and every later one. - if sequential && shifted && ahead == last_frame + 1 { + // This request is the first whose window would run past the clip's last frame. The previous + // request pushed that last frame into a live stream, and no tail cache covers this index, so + // flushing the stream yields this frame and every later one. + if sequential && shifted && window_end == last_frame + 1 { let mut outputs = Vec::new(); state.denoiser.flush(|planes| outputs.push(planes))?; let mut outputs = outputs.into_iter(); - let out = outputs + let denoised = outputs .next() .ok_or_else(|| anyhow!("the clip's final flush produced no frame"))?; let rest: Vec = outputs.collect(); - state.tail = Some(TailCache::new(n + 1, rest)); - state.last = Some(n); + let tail = TailCache::new(output_index + 1, rest); + state.tail = Some(tail); + state.last_served = Some(output_index); state.finish_warm_up(); - return Ok(out); + return Ok(denoised); } - if sequential && (!shifted || ahead <= last_frame) { - let next = ahead.min(last_frame); - let frame = fetch(next)?; + if sequential && (!shifted || window_end <= last_frame) { + let next_index = window_end.min(last_frame); + let frame = fetch(next_index)?; state.denoiser.push(&frame)?; - if let Some(out) = state.denoiser.recv()? { - state.last = Some(n); + if let Some(denoised) = state.denoiser.recv()? { + state.last_served = Some(output_index); state.finish_warm_up(); - return Ok(out); + return Ok(denoised); } + // The stream did not yield, so fall through and rebuild. } - let indices = self.window(n); + let indices = self.window(output_index); let window: Vec = indices .iter() .map(|&index| fetch(index)) .collect::>()?; - let out = if shifted { - let first = indices[0]; + let denoised = if shifted { + let first_index = indices[0]; let last_index = *indices .last() .ok_or_else(|| anyhow!("window produced no frame indices"))?; let request = ReseedWindow { frames: &window, - target: n - first, - at_clip_start: first == 0, + target: output_index - first_index, + at_clip_start: first_index == 0, at_clip_end: last_index == last_frame, }; + let mut outputs = state.denoiser.reseed_window(request)?.into_iter(); - let out = outputs + let denoised = outputs .next() .ok_or_else(|| anyhow!("reseed produced no frame"))?; let rest: Vec = outputs.collect(); if !rest.is_empty() { - state.tail = Some(TailCache::new(n + 1, rest)); + let tail = TailCache::new(output_index + 1, rest); + state.tail = Some(tail); } - out + + denoised } else { state.denoiser.reseed(&window)? }; - state.last = Some(n); + state.last_served = Some(output_index); state.finish_warm_up(); - Ok(out) + Ok(denoised) } } -/// Packs one source frame's three planes into a [`Planes`], dropping -/// each plane's row padding. +/// Packs one source frame's three planes into a [Planes], dropping each plane's row padding. fn pack_frame(frame: &FrameRef, depth_bytes: usize) -> Planes { let pack = |plane: usize| -> Vec { let stride = frame.stride(plane); @@ -288,24 +270,28 @@ fn pack_frame(frame: &FrameRef, depth_bytes: usize) -> Planes { pack_plane(data, stride, width_bytes, height) }; + let y_plane = pack(0); + let u_plane = pack(1); + let v_plane = pack(2); + Planes { - y: pack(0), - u: pack(1), - v: pack(2), + y: y_plane, + u: u_plane, + v: v_plane, } } -/// Writes a denoised [`Planes`] into a freshly allocated output frame. +/// Writes a denoised [Planes] into a freshly allocated output frame. fn unpack_into_frame(frame: &mut FrameRefMut, planes: &Planes, depth_bytes: usize) { let sources = [&planes.y, &planes.u, &planes.v]; - for (plane, src) in sources.into_iter().enumerate() { + for (plane, packed) in sources.into_iter().enumerate() { let stride = frame.stride(plane); let height = frame.height(plane); let width_bytes = frame.width(plane) * depth_bytes; // SAFETY: `stride * height` is exactly the byte range VapourSynth // allocated for this plane. let data = unsafe { std::slice::from_raw_parts_mut(frame.data_ptr_mut(plane), stride * height) }; - unpack_plane_into(data, stride, width_bytes, height, src); + unpack_plane_into(data, stride, width_bytes, height, packed); } } @@ -319,11 +305,12 @@ impl<'core> Filter<'core> for Denoise<'core> { _api: API, _core: CoreRef<'core>, context: FrameContext, - n: usize, + output_index: usize, ) -> Result>, Error> { - for idx in self.unique_window(n) { - self.source.request_frame_filter(context, idx); + for index in self.unique_window(output_index) { + self.source.request_frame_filter(context, index); } + Ok(None) } @@ -332,29 +319,32 @@ impl<'core> Filter<'core> for Denoise<'core> { _api: API, core: CoreRef<'core>, context: FrameContext, - n: usize, + output_index: usize, ) -> Result, Error> { let mut frames: HashMap> = HashMap::new(); - for idx in self.unique_window(n) { + for index in self.unique_window(output_index) { let frame = self .source - .get_frame_filter(context, idx) - .ok_or_else(|| anyhow!("couldn't get source frame {idx}"))?; - frames.insert(idx, frame); + .get_frame_filter(context, index) + .ok_or_else(|| anyhow!("couldn't get source frame {index}"))?; + frames.insert(index, frame); } let depth_bytes = self.layout.depth.bytes_per_sample(); - let fetch = |idx: usize| -> Result { + let fetch = |index: usize| -> Result { let frame = frames - .get(&idx) + .get(&index) .expect("get_frame_initial requested the same window as get_frame"); - Ok(pack_frame(frame, depth_bytes)) + let packed = pack_frame(frame, depth_bytes); + Ok(packed) }; - let planes = self.render(n, fetch)?; + let planes = self.render(output_index, fetch)?; - let prop_src = frames.get(&n).expect("the window always includes n"); - let format = prop_src.format(); + let props_source = frames + .get(&output_index) + .expect("the window always includes output_index"); + let format = props_source.format(); let resolution = Resolution { width: self.layout.width as usize, height: self.layout.height as usize, @@ -363,9 +353,10 @@ impl<'core> Filter<'core> for Denoise<'core> { // SAFETY: the frame's plane data starts uninitialized, but // `unpack_into_frame` below writes every byte of every plane // before the frame is returned to VapourSynth. - let mut out = unsafe { FrameRefMut::new_uninitialized(core, Some(prop_src), format, resolution) }; - unpack_into_frame(&mut out, &planes, depth_bytes); + let mut output_frame = + unsafe { FrameRefMut::new_uninitialized(core, Some(props_source), format, resolution) }; + unpack_into_frame(&mut output_frame, &planes, depth_bytes); - Ok(out.into()) + Ok(output_frame.into()) } } diff --git a/av-denoise-vs/src/frames.rs b/av-denoise-vs/src/frames.rs index 0e65f17..c97bb68 100644 --- a/av-denoise-vs/src/frames.rs +++ b/av-denoise-vs/src/frames.rs @@ -1,32 +1,24 @@ -//! Converts between VapourSynth's strided plane buffers and the tightly -//! packed rows core's converters accept, plus the frame window core -//! windowed algorithms request around each output frame, including the -//! shifted window nl4d uses at a clip's edges. - use std::ops::RangeInclusive; -use av_denoise_core::Planes; +use av_denoise::Planes; -/// Copies `src`, a plane with row stride `stride` bytes, into a new -/// tightly packed buffer of `width_bytes * height` bytes. +/// Copies a plane with a row stride of `stride` bytes into a tightly packed buffer. /// -/// `width_bytes` is `width * bytes_per_sample`, not a pixel count, so -/// this works the same at any bit depth. Passing a pixel count here -/// packs the wrong number of bytes per row. +/// `width_bytes` is `width * bytes_per_sample`, not a pixel count, so this works the same at any bit +/// depth. Passing a pixel count here packs the wrong number of bytes per row. pub fn pack_plane(src: &[u8], stride: usize, width_bytes: usize, height: usize) -> Vec { let mut packed = Vec::with_capacity(width_bytes * height); for row in src.chunks(stride).take(height) { packed.extend_from_slice(&row[..width_bytes]); } + packed } -/// Writes a tightly packed plane, `src`, back into `dst`, a strided -/// buffer with row stride `stride` bytes. The reverse of [`pack_plane`]. +/// Writes the tightly packed plane `src` back into `dst`, which has a row stride of `stride` bytes. /// -/// `width_bytes` is `width * bytes_per_sample`, not a pixel count, so -/// this works the same at any bit depth. `dst`'s padding bytes, if any, -/// are left untouched. +/// The reverse of [pack_plane]. `width_bytes` is `width * bytes_per_sample`, not a pixel count. +/// Padding bytes in `dst` are left untouched. pub fn unpack_plane_into(dst: &mut [u8], stride: usize, width_bytes: usize, height: usize, src: &[u8]) { for (y, row) in dst.chunks_mut(stride).take(height).enumerate() { let packed_row = &src[y * width_bytes..(y + 1) * width_bytes]; @@ -34,37 +26,27 @@ pub fn unpack_plane_into(dst: &mut [u8], stride: usize, width_bytes: usize, heig } } -/// The `behind + 1 + ahead` source frame indices for the window around -/// output frame `n`, `behind` older and `ahead` newer, clamped so -/// nothing runs off either end of a clip whose last valid index is -/// `last_frame`. -/// -/// `behind` and `ahead` come from the denoiser's own -/// [`av_denoise_core::PlanarDenoiser::window_span`], so this stays -/// correct for whichever algorithm the denoiser is running rather than -/// assuming every algorithm needs the same symmetric window. +/// The `behind + 1 + ahead` source indices around an output frame, clamped to `0..=last_frame`. /// -/// Under repeated edges, frame requests and window builds both call this, so -/// the two always agree on which frames a window at `n` pulls in. -pub fn window_indices(n: usize, behind: usize, ahead: usize, last_frame: usize) -> Vec { +/// Under repeated edges, frame requests and window builds both call this, so the two always agree on +/// which frames a window pulls in. +pub fn window_indices(output_index: usize, behind: usize, ahead: usize, last_frame: usize) -> Vec { (0..=behind + ahead) - .map(|i| (n + i).saturating_sub(behind).min(last_frame)) + .map(|i| (output_index + i).saturating_sub(behind).min(last_frame)) .collect() } -/// The source frame range for output frame `n` under shifted edges. +/// The source frame range around an output frame under shifted edges. /// -/// It is `behind` older and `ahead` newer frames, cut short at either -/// end of a clip whose last valid index is `last_frame` rather than -/// repeating a boundary frame. +/// The range is cut short at either end of the clip rather than repeating a boundary frame. pub fn shifted_window_range( - n: usize, + output_index: usize, behind: usize, ahead: usize, last_frame: usize, ) -> RangeInclusive { - let first = n.saturating_sub(behind); - let last = (n + ahead).min(last_frame); + let first = output_index.saturating_sub(behind); + let last = (output_index + ahead).min(last_frame); first..=last } @@ -82,8 +64,8 @@ impl TailCache { Self { first, frames } } - pub fn take(&mut self, n: usize) -> Option { - let index = n.checked_sub(self.first)?; + pub fn take(&mut self, output_index: usize) -> Option { + let index = output_index.checked_sub(self.first)?; self.frames.get_mut(index)?.take() } } diff --git a/av-denoise-vs/src/lib.rs b/av-denoise-vs/src/lib.rs index 9cd57a4..e91bb5b 100644 --- a/av-denoise-vs/src/lib.rs +++ b/av-denoise-vs/src/lib.rs @@ -11,13 +11,12 @@ use vapoursynth::plugins::{Filter, FilterArgument, Metadata}; use vapoursynth::prelude::{API, Node}; use vapoursynth::{export_vapoursynth_plugin, make_filter_function}; -use crate::filter::Denoise; -use crate::params::{AlgorithmKind, RawParams}; +use self::filter::Denoise; +use self::params::{AlgorithmKind, RawParams}; /// Installs the tracing subscriber that writes the plugin's logs to stderr. /// -/// `RUST_LOG` picks what is printed, and without it the plugin logs at `warn` so -/// an ordinary render stays quiet. +/// `RUST_LOG` picks what is printed. Without it the plugin logs at `warn` so an ordinary render stays quiet. fn init_logging() { let filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("warn")); @@ -29,28 +28,20 @@ fn init_logging() { /// Keeps this plugin's library mapped for the rest of the process. /// -/// VapourSynth unloads every plugin library when it frees a core, and -/// vspipe frees its core right before exiting. The GPU runtime this -/// plugin builds spawns a device thread per accelerator, plus a -/// polling thread per stream on the wgpu backends, and those threads -/// run for the rest of the process. The device thread never blocks. It -/// spins, yields, then sleeps briefly, over and over. On Windows its -/// first wake after `FreeLibrary` returns into unmapped code and the -/// process dies with an access violation, after every frame was -/// already written. The polling thread parks or waits in the driver, -/// and dies the same way once anything wakes it. +/// VapourSynth unloads every plugin library when it frees a core, and vspipe frees its core right +/// before exiting. The GPU runtime runs a device thread per accelerator, plus a polling thread per +/// stream on the wgpu backends, until the process exits. The device thread never blocks. It spins, +/// yields and sleeps in a loop, so on Windows its first wake after `FreeLibrary` returns into unmapped +/// code and the process dies with an access violation after every frame was written. The polling +/// thread parks or waits in the driver, and dies the same way once anything wakes it. /// -/// Pinning the module makes the unload a no-op, so the threads stay -/// valid until process exit terminates them. On Linux the loader -/// already refuses to unload a library that registered thread-local -/// destructors, which is what happens as soon as this plugin's threads -/// start, so nothing needs doing there. macOS is not covered and has -/// not been tested. +/// Pinning the module makes the unload a no-op. Linux needs nothing, since its loader refuses to +/// unload a library that registered thread-local destructors, which these threads do as soon as they +/// start. macOS is not covered and has not been tested. /// -/// This runs once, on the first filter creation. That is before any -/// device thread exists, since only a filter builds a denoiser. The -/// plugin's init function runs earlier, but the export macro owns its -/// body and this plugin has no code of its own in it. +/// This runs once, on the first filter creation, which is before any device thread exists since only +/// a filter builds a denoiser. The plugin's init function runs earlier, but the export macro owns its +/// body. fn pin_plugin_library() { static PIN: std::sync::Once = std::sync::Once::new(); PIN.call_once(|| { @@ -75,34 +66,33 @@ fn pin_plugin_library_windows() { let mut module: *mut c_void = std::ptr::null_mut(); // SAFETY: `address` is a code address inside this library, which is // what `FROM_ADDRESS` asks for, and `module` is a valid out pointer. - let ok = unsafe { + let pinned = unsafe { GetModuleHandleExW( GET_MODULE_HANDLE_EX_FLAG_PIN | GET_MODULE_HANDLE_EX_FLAG_FROM_ADDRESS, address, &mut module, ) }; - if ok == 0 { + if pinned == 0 { tracing::warn!("could not pin the plugin library, the process may crash at exit"); } } -/// Reads one optional UTF-8 script argument, naming `field` in the error -/// when the bytes are not valid UTF-8. +/// Reads an optional UTF-8 script argument, naming `field` in the error when it is not valid UTF-8. fn opt_string(bytes: Option<&[u8]>, field: &str) -> Result, Error> { - bytes - .map(|b| String::from_utf8(b.to_vec()).map_err(|_| anyhow::anyhow!("{field} must be valid UTF-8"))) - .transpose() + let Some(bytes) = bytes else { + return Ok(None); + }; + + let owned = bytes.to_vec(); + let text = String::from_utf8(owned).map_err(|_| anyhow::anyhow!("{field} must be valid UTF-8"))?; + Ok(Some(text)) } -/// Reads the optional `accelerators` script argument, a comma-separated -/// list of accelerator names, into the `Vec` [`RawParams`] -/// wants. +/// Reads the optional `accelerators` script argument, a comma-separated list of accelerator names. /// -/// VapourSynth script arguments have no native string array type that -/// fits cleanly into `make_filter_function!`'s generated argument -/// string, so this reuses the plain `data` type and splits it, matching -/// how `channel_mode` and `device` already take a single string. +/// VapourSynth has no string array argument type that fits `make_filter_function!`'s generated +/// argument string, so this takes the plain `data` type and splits it. fn opt_accelerators(bytes: Option<&[u8]>) -> Result>, Error> { let Some(joined) = opt_string(bytes, "accelerators")? else { return Ok(None); @@ -111,7 +101,7 @@ fn opt_accelerators(bytes: Option<&[u8]>) -> Result>, Error> let names: Vec = joined .split(',') .map(str::trim) - .filter(|s| !s.is_empty()) + .filter(|name| !name.is_empty()) .map(str::to_string) .collect(); @@ -122,13 +112,14 @@ fn opt_accelerators(bytes: Option<&[u8]>) -> Result>, Error> Ok(Some(names)) } -/// Reads an optional on/off script argument. fn opt_bool(value: Option) -> Option { - value.map(|v| v != 0) + value.map(|flag| flag != 0) } -/// Builds a [`RawParams`] from a filter function's raw script arguments. -#[expect(clippy::too_many_arguments)] +#[expect( + clippy::too_many_arguments, + reason = "takes one parameter per optional argument across both VapourSynth filters" +)] fn raw_params( strength: Option, variant: Option<&[u8]>, @@ -192,7 +183,10 @@ fn raw_params( make_filter_function! { NlmeansFunction, "NLMeans" - #[expect(clippy::too_many_arguments)] + #[expect( + clippy::too_many_arguments, + reason = "each parameter is a VapourSynth filter argument, so they cannot be grouped" + )] fn create_nlmeans<'core>( api: API, core: CoreRef<'core>, @@ -243,19 +237,23 @@ make_filter_function! { None, )?; let filter = Denoise::create(api, core, clip, AlgorithmKind::Nlmeans, &raw)?; - Ok(Some(Box::new(filter))) + let boxed: Box + 'core> = Box::new(filter); + Ok(Some(boxed)) } } make_filter_function! { Nl4dFunction, "NL4D" - /// Estimates its automatic noise level fresh from each frame's own - /// temporal window, rather than smoothing it across the whole - /// stream, so a frame denoises to the same pixels no matter what - /// order VapourSynth requests frames in. Passing `sigma` pins the - /// noise level and skips that estimator entirely. - #[expect(clippy::too_many_arguments)] + /// Creates an `avd.NL4D` filter. + /// + /// The automatic noise level is estimated from each frame's own temporal window, so a frame + /// denoises to the same pixels in any order VapourSynth requests frames. Passing `sigma` pins the + /// noise level and skips the estimator. + #[expect( + clippy::too_many_arguments, + reason = "each parameter is a VapourSynth filter argument, so they cannot be grouped" + )] fn create_nl4d<'core>( api: API, core: CoreRef<'core>, @@ -312,7 +310,8 @@ make_filter_function! { pooled_threshold, )?; let filter = Denoise::create(api, core, clip, AlgorithmKind::Nl4d, &raw)?; - Ok(Some(Box::new(filter))) + let boxed: Box + 'core> = Box::new(filter); + Ok(Some(boxed)) } } diff --git a/av-denoise-vs/src/params.rs b/av-denoise-vs/src/params.rs index 3d1c0f6..db6aa7c 100644 --- a/av-denoise-vs/src/params.rs +++ b/av-denoise-vs/src/params.rs @@ -1,15 +1,5 @@ -//! Turns a VapourSynth clip's format and a filter's script arguments -//! into the option types `av-denoise-core` denoises with. -//! -//! Everything here is a pure function over plain values, with no -//! VapourSynth core and no GPU, so the whole accept/reject matrix is -//! unit-testable. [`Format`](vapoursynth::format::Format) itself cannot -//! be built outside a running core, so [`layout_from_format`] takes a -//! [`RawFormat`] of the plain fields it needs instead. The caller in -//! `filter.rs` does the short extraction from a real `Format`. - -use av_denoise_core::accelerate::{Accelerator, get_default_accelerators}; -use av_denoise_core::{ +use av_denoise::accelerate::{Accelerator, get_default_accelerators}; +use av_denoise::{ Algorithm, ChannelIntent, DenoisingMode, @@ -37,14 +27,10 @@ use av_denoise_core::{ }; use vapoursynth::format::{ColorFamily, SampleType}; -/// The handful of format fields [`layout_from_format`] actually reads. +/// The format fields [layout_from_format] reads. /// -/// The real caller is `vapoursynth::format::Format`, which wraps a -/// pointer only a running VapourSynth core can hand out, so it cannot be -/// built in a unit test. A caller with a real `Format` builds one of -/// these from `format.sample_type()`, `format.bits_per_sample()`, -/// `format.sub_sampling_w()`, `format.sub_sampling_h()`, and -/// `format.color_family()`. +/// `vapoursynth::format::Format` wraps a pointer only a running VapourSynth core can hand out, so it +/// cannot be built in a unit test. #[derive(Debug, Clone, Copy)] pub struct RawFormat { pub sample_type: SampleType, @@ -54,21 +40,12 @@ pub struct RawFormat { pub color_family: ColorFamily, } -/// Validates a clip's format and turns it into a [`FrameLayout`]. +/// Validates a clip's format and turns it into a [FrameLayout]. /// -/// Accepts integer YUV420, YUV422, and YUV444 sources at 8, 10, or -/// 12-bit. Rejects RGB, since the denoiser's channel distance weights -/// are calibrated for YUV. Rejects float sample types and any other -/// chroma subsampling. -/// -/// Rejects GRAY too. Core's [`Subsampling`] has no "no chroma" variant, -/// so a GRAY source would have to be represented as YUV444, which makes -/// [`av_denoise_core::frame::FrameLayout::chroma_dims`] report -/// full-resolution chroma planes that do not exist. The filter would -/// then have to fabricate and push full-size neutral chroma every -/// frame, four times the real data volume of true 4:2:0 chroma, purely -/// to work around a gap in the geometry type. GRAY is out of scope -/// until core can represent a source with no chroma planes at all. +/// Accepts integer YUV420, YUV422 and YUV444 sources at 8, 10 or 12 bits. RGB is rejected because +/// the denoiser's channel distance weights are calibrated for YUV. GRAY is rejected because +/// [Subsampling] has no "no chroma" variant, and representing it as YUV444 would push full-size +/// neutral chroma every frame, four times the data volume of true 4:2:0 chroma. pub fn layout_from_format(format: RawFormat, width: u32, height: u32) -> Result { match format.color_family { ColorFamily::YUV => {}, @@ -96,9 +73,9 @@ pub fn layout_from_format(format: RawFormat, width: u32, height: u32) -> Result< (0, 0) => Subsampling::Yuv444, (1, 0) => Subsampling::Yuv422, (1, 1) => Subsampling::Yuv420, - (w, h) => { + (subsampling_w, subsampling_h) => { anyhow::bail!( - "unsupported chroma subsampling (subsampling_w={w}, subsampling_h={h}), av-denoise-vs accepts YUV420, YUV422, and YUV444" + "unsupported chroma subsampling (subsampling_w={subsampling_w}, subsampling_h={subsampling_h}), av-denoise-vs accepts YUV420, YUV422, and YUV444" ); }, }; @@ -112,19 +89,12 @@ pub fn layout_from_format(format: RawFormat, width: u32, height: u32) -> Result< } /// Which denoising algorithm a filter function runs. -/// -/// `avd.Nlmeans` builds [`AlgorithmKind::Nlmeans`], `avd.Nl4d` builds -/// [`AlgorithmKind::Nl4d`]. This has no `Hq` variant since the -/// VapourSynth plugin does not expose the HQ variant separately, it -/// takes a plain algorithm choice per filter function. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum AlgorithmKind { Nlmeans, Nl4d, } -/// The name a [`NlmeansVariant`] parses back from, used in error -/// messages. fn variant_name(variant: NlmeansVariant) -> &'static str { match variant { NlmeansVariant::Fast => "fast", @@ -132,32 +102,24 @@ fn variant_name(variant: NlmeansVariant) -> &'static str { } } -/// Resolves an explicit `variant` string into an [`NlmeansVariant`]. -/// -/// Uses [`av_denoise_core`]'s own parser, the same one the CLI's -/// `--variant` flag resolves through, so a name accepted on the CLI is -/// accepted here too. +/// Parses a `variant` name with the parser the CLI's `--variant` flag uses, so both accept the same +/// names. fn parse_variant(raw: &str) -> Result { raw.parse::() .map_err(|_| anyhow::anyhow!("unknown variant '{raw}', expected one of fast, hq")) } -/// Resolves an explicit `preset` string into a [`Preset`]. -/// -/// Uses [`av_denoise_core`]'s own parser, the same one the CLI's -/// `--preset` flag resolves through, so a name accepted on the CLI is -/// accepted here too. +/// Parses a `preset` name with the parser the CLI's `--preset` flag uses, so both accept the same +/// names. fn parse_preset(raw: &str) -> Result { raw.parse::().map_err(|_| { anyhow::anyhow!("unknown preset '{raw}', expected one of veryfast, fast, base, slow, veryslow") }) } -/// The raw script arguments a filter function receives, before they are -/// validated and folded into a [`PlaneOptions`]. +/// The script arguments a filter function receives, before they are validated. /// -/// Every field is optional. An unset field falls back to the library's -/// own default for whichever algorithm is being built. +/// An unset field falls back to the library's own default for the algorithm being built. #[derive(Debug, Clone, Default)] pub struct RawParams { pub strength: Option, @@ -189,14 +151,11 @@ pub struct RawParams { pub pooled_threshold: Option, } -/// Turns a nonnegative script integer into a `u32`, naming `field` in -/// the error when it is negative. fn nonnegative(value: i64, field: &str) -> Result { u32::try_from(value).map_err(|_| anyhow::anyhow!("{field} must not be negative, got {value}")) } -/// Resolves an explicit `channel_mode` string into a [`ChannelIntent`], -/// rejecting anything the source can't support. +/// Parses a `channel_mode` name, rejecting a mode the source cannot support. fn parse_channel_mode(raw: &str, layout: FrameLayout) -> Result { let intent = match raw.to_ascii_lowercase().as_str() { "luma" => ChannelIntent::Luma, @@ -214,42 +173,15 @@ fn parse_channel_mode(raw: &str, layout: FrameLayout) -> Result Result { - // Resolved once and read by both algorithms below, exactly like the - // CLI's own `--preset`. An explicit `variant`, `temporal_radius`, or - // `search_radius` overrides whatever the preset would have picked, - // matching `NlmeansArgs::resolve_preset`'s precedence. + // An explicit `variant`, `temporal_radius` or `search_radius` overrides what the preset picks, + // matching the CLI's precedence. let preset = match &raw.preset { None => Preset::default(), - Some(p) => parse_preset(p)?, + Some(name) => parse_preset(name)?, }; - // Only `Nlmeans` reads `variant` at all, so `Nl4d` never parses it, - // it is rejected as a mismatched parameter below instead if set. + // Only nlmeans reads `variant`, so nl4d never parses it and a set value fails as a mismatch below. let variant = match algorithm_kind { AlgorithmKind::Nlmeans => match raw.variant.as_deref() { None => nlmeans_variant_for(preset), - Some(v) => parse_variant(v)?, + Some(name) => parse_variant(name)?, }, AlgorithmKind::Nl4d => NlmeansVariant::Hq, }; @@ -357,7 +282,7 @@ pub fn plane_options_from( if let Some(radius) = raw.search_radius && radius > 4 - && !av_denoise_core::codegen_stack_is_sufficient() + && !av_denoise::codegen_stack_is_sufficient() { anyhow::bail!( "search_radius {radius} needs a raised stack, but RUST_MIN_STACK is not set. Values above 4 overflow the default 2 MiB stack during kernel codegen" @@ -371,28 +296,25 @@ pub fn plane_options_from( let device = match &raw.device { None => Device::default(), - Some(s) => s + Some(name) => name .parse() - .map_err(|e| anyhow::anyhow!("invalid device '{s}': {e}"))?, + .map_err(|error| anyhow::anyhow!("invalid device '{name}': {error}"))?, }; let accelerators = match &raw.accelerators { None => get_default_accelerators(), Some(names) => names .iter() - .map(|s| { - s.parse::() - .map_err(|e| anyhow::anyhow!("invalid accelerator '{s}': {e}")) + .map(|name| { + name.parse::() + .map_err(|error| anyhow::anyhow!("invalid accelerator '{name}': {error}")) }) .collect::, _>>()?, }; - // Both algorithms resolve their preset-driven temporal radius the - // same way the CLI does: an explicit `temporal_radius` overrides - // whatever the preset picks. nl4d groups patches across a temporal - // window and has no spatial-only mode, but no preset ever resolves - // it to 0, so the `radius == 0` arm below only actually triggers - // for `nlmeans`, whose `veryfast` preset does. + // An explicit `temporal_radius` overrides the preset's, the same way the CLI resolves it. nl4d has + // no spatial-only mode, and no nl4d preset resolves to 0, so only the nlmeans `veryfast` preset + // picks the spatial mode by default. let preset_temporal_radius = match algorithm_kind { AlgorithmKind::Nlmeans => nlmeans_temporal_radius_for(preset), AlgorithmKind::Nl4d => nl4d_temporal_radius_for(preset), @@ -411,128 +333,112 @@ pub fn plane_options_from( let algorithm = match algorithm_kind { AlgorithmKind::Nlmeans => { + let search_radius = match raw.search_radius { + None => nlmeans_search_radius_for(preset), + Some(radius) => nonnegative(radius, "search_radius")?, + }; + let tuning = NlmTuning { - search_radius: Some(match raw.search_radius { - None => nlmeans_search_radius_for(preset), - Some(r) => nonnegative(r, "search_radius")?, - }), + search_radius: Some(search_radius), patch_radius: raw .patch_radius - .map(|r| nonnegative(r, "patch_radius")) + .map(|radius| nonnegative(radius, "patch_radius")) .transpose()?, - strength: raw.strength.map(|v| v as f32), + strength: raw.strength.map(|value| value as f32), ..NlmTuning::default() }; + let motion_search = MotionSearch::default(); let motion_compensation = match raw.motion_compensation { - Some(true) => MotionCompensationMode::from(MotionSearch::default()), + Some(true) => MotionCompensationMode::from(motion_search), Some(false) | None => MotionCompensationMode::None, }; let prefilter = match &raw.prefilter { None => PrefilterMode::None, - Some(s) => { - let mode = parse_prefilter(s)?; - // `parse_prefilter`'s string grammar has no form - // that produces `External`, but the check stays - // here as a boundary guard rather than trusting - // that invariant silently: `External` needs a - // reference frame supplied through - // `push_frame_with_reference`, which this plugin - // has no way to call. - if matches!(mode, PrefilterMode::External) { - anyhow::bail!( - "prefilter 'external' is not supported by av-denoise-vs, which has no way to supply a reference frame" - ); - } - mode - }, + Some(spec) => parse_prefilter(spec)?, + }; + + let nlm = NlmeansOptions { + prefilter, + motion_compensation, + tuning, + mode, }; match variant { - NlmeansVariant::Fast => Algorithm::Nlmeans(NlmeansOptions { - prefilter, - motion_compensation, - tuning, - }), - NlmeansVariant::Hq => Algorithm::NlmeansHq(NlmeansHqOptions { - nlm: NlmeansOptions { - prefilter, - motion_compensation, - tuning, - }, - hq: HqParams { - sigma_override: raw.sigma.map(|v| v as f32), + NlmeansVariant::Fast => Algorithm::Nlmeans(nlm), + NlmeansVariant::Hq => { + let hq = HqParams { + sigma_override: raw.sigma.map(|value| value as f32), sigma_scale: raw .sigma_scale - .map(|v| v as f32) + .map(|value| value as f32) .unwrap_or_else(|| HqParams::default().sigma_scale), - // A VapourSynth filter has to return the same - // pixels for a frame no matter what order - // frames were requested in, and history- - // dependent estimation breaks that guarantee - // under random access. See `Nl4dOptions`'s own - // `windowed_noise_estimation` field for the - // same reasoning applied to nl4d. + // A VapourSynth filter has to return the same pixels for a frame in any + // request order, which history-dependent estimation breaks under random access. windowed_noise_estimation: true, ..HqParams::default() - }, - }), + }; + + let options = NlmeansHqOptions { nlm, hq }; + Algorithm::NlmeansHq(options) + }, } }, - AlgorithmKind::Nl4d => Algorithm::Nl4d(Nl4dOptions { - // A VapourSynth filter has to return the same pixels for a - // frame no matter what order frames were requested in. - // window-local estimation computes sigma from only the - // frames in the current window, so the fast path and a - // `reseed` after random access agree by construction. There - // is no reason to expose the stream-history-dependent - // temporal EMA here at all. - windowed_noise_estimation: true, - sigma: raw.sigma.map(|v| v as f32), - sigma_scale: raw - .sigma_scale - .map(|v| v as f32) - .unwrap_or_else(|| Nl4dOptions::default().sigma_scale), - lambda_ht: raw.lambda_ht.map(|v| v as f32), - lambda_ht_scale: raw - .lambda_ht_scale - .map(|v| v as f32) - .unwrap_or_else(|| Nl4dOptions::default().lambda_ht_scale), - spatial_radius: match raw.spatial_radius { - Some(r) => nonnegative(r, "spatial_radius")?, - None => nl4d_spatial_radius_for(preset), - }, - refine: match raw.refine { - Some(r) => nonnegative(r, "refine")?, - None => Nl4dOptions::default().refine, - }, - noise_map: match raw.noise_map { - Some(enabled) => enabled, - None => Nl4dOptions::default().noise_map, - }, - flat_boost: raw - .flat_boost - .map(|v| v as f32) - .unwrap_or_else(|| Nl4dOptions::default().flat_boost), - chroma_flat_boost: raw - .chroma_flat_boost - .map(|v| v as f32) - .unwrap_or_else(|| Nl4dOptions::default().chroma_flat_boost), - shadow_soften: raw - .shadow_soften - .map(|v| v as f32) - .unwrap_or_else(|| Nl4dOptions::default().shadow_soften), - flat_texture_cut: raw - .flat_texture_cut - .map(|v| v as f32) - .unwrap_or_else(|| Nl4dOptions::default().flat_texture_cut), - pooled_threshold: match raw.pooled_threshold { - Some(enabled) => enabled, - None => Nl4dOptions::default().pooled_threshold, - }, - ..Nl4dOptions::default() - }), + AlgorithmKind::Nl4d => { + let options = Nl4dOptions { + // A VapourSynth filter has to return the same pixels for a frame in any request + // order. Window-local estimation reads sigma from only the current window, so the + // fast path and a reseed after random access agree by construction. + windowed_noise_estimation: true, + sigma: raw.sigma.map(|value| value as f32), + sigma_scale: raw + .sigma_scale + .map(|value| value as f32) + .unwrap_or_else(|| Nl4dOptions::default().sigma_scale), + lambda_ht: raw.lambda_ht.map(|value| value as f32), + lambda_ht_scale: raw + .lambda_ht_scale + .map(|value| value as f32) + .unwrap_or_else(|| Nl4dOptions::default().lambda_ht_scale), + spatial_radius: match raw.spatial_radius { + Some(radius) => nonnegative(radius, "spatial_radius")?, + None => nl4d_spatial_radius_for(preset), + }, + refine: match raw.refine { + Some(radius) => nonnegative(radius, "refine")?, + None => Nl4dOptions::default().refine, + }, + noise_map: match raw.noise_map { + Some(enabled) => enabled, + None => Nl4dOptions::default().noise_map, + }, + flat_boost: raw + .flat_boost + .map(|value| value as f32) + .unwrap_or_else(|| Nl4dOptions::default().flat_boost), + chroma_flat_boost: raw + .chroma_flat_boost + .map(|value| value as f32) + .unwrap_or_else(|| Nl4dOptions::default().chroma_flat_boost), + shadow_soften: raw + .shadow_soften + .map(|value| value as f32) + .unwrap_or_else(|| Nl4dOptions::default().shadow_soften), + flat_texture_cut: raw + .flat_texture_cut + .map(|value| value as f32) + .unwrap_or_else(|| Nl4dOptions::default().flat_texture_cut), + pooled_threshold: match raw.pooled_threshold { + Some(enabled) => enabled, + None => Nl4dOptions::default().pooled_threshold, + }, + ..Nl4dOptions::default() + }; + + Algorithm::Nl4d(options) + }, }; Ok(PlaneOptions { @@ -541,9 +447,9 @@ pub fn plane_options_from( intent, mode, algorithm, - luma_strength: raw.luma_strength.map(|v| v as f32), - chroma_strength: raw.chroma_strength.map(|v| v as f32), - luma_lambda_ht: raw.luma_lambda_ht.map(|v| v as f32), - chroma_lambda_ht: raw.chroma_lambda_ht.map(|v| v as f32), + luma_strength: raw.luma_strength.map(|value| value as f32), + chroma_strength: raw.chroma_strength.map(|value| value as f32), + luma_lambda_ht: raw.luma_lambda_ht.map(|value| value as f32), + chroma_lambda_ht: raw.chroma_lambda_ht.map(|value| value as f32), }) } diff --git a/av-denoise-vs/tests/frames.rs b/av-denoise-vs/tests/frames.rs index d54a6b7..74b335e 100644 --- a/av-denoise-vs/tests/frames.rs +++ b/av-denoise-vs/tests/frames.rs @@ -1,49 +1,53 @@ -use av_denoise_core::Planes; +use av_denoise::Planes; use av_denoise_vs::frames::{TailCache, pack_plane, shifted_window_range, unpack_plane_into, window_indices}; -/// A strided buffer whose padding bytes are all 0xAA, so a bug that -/// reads padding is visible rather than silently plausible. +/// A strided buffer whose padding bytes are all 0xAA, so a bug that reads padding is visible rather +/// than silently plausible. fn strided(width_bytes: usize, height: usize, stride: usize) -> Vec { - let mut buf = vec![0xAAu8; stride * height]; + let mut buffer = vec![0xAAu8; stride * height]; for y in 0..height { for x in 0..width_bytes { - buf[y * stride + x] = (y * width_bytes + x) as u8; + buffer[y * stride + x] = (y * width_bytes + x) as u8; } } - buf + + buffer } #[test] fn packing_drops_the_row_padding() { - let (w, h, stride) = (5usize, 3usize, 8usize); - let packed = pack_plane(&strided(w, h, stride), stride, w, h); + let (width, height, stride) = (5usize, 3usize, 8usize); + let src = strided(width, height, stride); + let packed = pack_plane(&src, stride, width, height); - assert_eq!(packed.len(), w * h); + assert_eq!(packed.len(), width * height); assert!(!packed.contains(&0xAA), "padding leaked into the packed buffer"); - for y in 0..h { - for x in 0..w { - assert_eq!(packed[y * w + x], (y * w + x) as u8); + for y in 0..height { + for x in 0..width { + assert_eq!(packed[y * width + x], (y * width + x) as u8); } } } #[test] fn packing_round_trips_through_unpacking() { - let (w, h, stride) = (5usize, 3usize, 8usize); - let src = strided(w, h, stride); - let packed = pack_plane(&src, stride, w, h); + let (width, height, stride) = (5usize, 3usize, 8usize); + let src = strided(width, height, stride); + let packed = pack_plane(&src, stride, width, height); - let mut dst = vec![0xAAu8; stride * h]; - unpack_plane_into(&mut dst, stride, w, h, &packed); + let mut dst = vec![0xAAu8; stride * height]; + unpack_plane_into(&mut dst, stride, width, height, &packed); assert_eq!(dst, src, "round trip changed the buffer"); } #[test] fn an_unpadded_stride_is_a_straight_copy() { - let (w, h) = (5usize, 3usize); - let src = strided(w, h, w); - assert_eq!(pack_plane(&src, w, w, h), src); + let (width, height) = (5usize, 3usize); + let src = strided(width, height, width); + let packed = pack_plane(&src, width, width, height); + + assert_eq!(packed, src); } #[test] @@ -65,8 +69,7 @@ fn window_indices_are_untouched_mid_clip() { #[test] fn a_short_clip_clamps_from_both_ends_at_once() { - // Three frames, span two on each side, so every index saturates - // somewhere. + // Three frames with a span of two on each side, so every index saturates somewhere. assert_eq!(window_indices(1, 2, 2, 2), vec![0, 0, 1, 2, 2]); } @@ -75,24 +78,17 @@ fn a_zero_span_window_is_the_frame_itself() { assert_eq!(window_indices(4, 0, 0, 11), vec![4]); } -/// nl4d's window is wider ahead of the target than behind it, so -/// `window_indices` must honour `behind` and `ahead` independently -/// rather than assuming a symmetric window, the property that made -/// `window_indices` exist in the first place. +/// nl4d's window is wider ahead of the target than behind it. #[test] fn an_asymmetric_span_is_not_forced_symmetric() { assert_eq!(window_indices(5, 2, 4, 11), vec![3, 4, 5, 6, 7, 8, 9]); } -/// The asymmetric case clamps at the clip start too: `behind` alone -/// saturates while `ahead` still reaches forward normally. #[test] fn an_asymmetric_span_clamps_at_the_clip_start() { assert_eq!(window_indices(1, 2, 4, 11), vec![0, 0, 1, 2, 3, 4, 5]); } -/// The asymmetric case clamps at the clip end too: `ahead` saturates -/// while `behind` still reaches back normally. #[test] fn an_asymmetric_span_clamps_at_the_clip_end() { assert_eq!(window_indices(9, 2, 4, 11), vec![7, 8, 9, 10, 11, 11, 11]); @@ -100,10 +96,12 @@ fn an_asymmetric_span_clamps_at_the_clip_end() { #[test] fn a_two_byte_depth_packs_by_bytes_not_samples() { - // 4 samples of 16-bit data is 8 bytes per row, stride 12. - let (width_bytes, h, stride) = (8usize, 2usize, 12usize); - let packed = pack_plane(&strided(width_bytes, h, stride), stride, width_bytes, h); - assert_eq!(packed.len(), width_bytes * h); + // 4 samples of 16-bit data is 8 bytes per row, with a stride of 12. + let (width_bytes, height, stride) = (8usize, 2usize, 12usize); + let src = strided(width_bytes, height, stride); + let packed = pack_plane(&src, stride, width_bytes, height); + + assert_eq!(packed.len(), width_bytes * height); assert!(!packed.contains(&0xAA)); } @@ -139,7 +137,10 @@ fn planes_marked(mark: u8) -> Planes { #[test] fn the_tail_cache_serves_frames_in_order_once() { - let frames = vec![planes_marked(7), planes_marked(8), planes_marked(9)]; + let marked_7 = planes_marked(7); + let marked_8 = planes_marked(8); + let marked_9 = planes_marked(9); + let frames = vec![marked_7, marked_8, marked_9]; let mut cache = TailCache::new(7, frames); let frame_8 = cache.take(8).map(|planes| planes.y); diff --git a/av-denoise-vs/tests/params.rs b/av-denoise-vs/tests/params.rs index 4fe4cf4..25820fe 100644 --- a/av-denoise-vs/tests/params.rs +++ b/av-denoise-vs/tests/params.rs @@ -1,11 +1,12 @@ -use av_denoise_core::frame::Subsampling; -use av_denoise_core::{ +use av_denoise::{ DenoisingMode, Depth, + FrameLayout, MotionCompensationMode, NlmeansVariant, PrefilterMode, Preset, + Subsampling, nl4d_spatial_radius_for, nl4d_temporal_radius_for, nlmeans_search_radius_for, @@ -55,14 +56,21 @@ fn test_format_gray(bits_per_sample: u8) -> RawFormat { } } +fn yuv420_layout() -> FrameLayout { + let format = test_format_yuv(1, 1, 8); + layout_from_format(format, 160, 120).unwrap() +} + #[test] fn accepts_the_three_supported_subsamplings_at_eight_bit() { - for (fmt, expected) in [ + let cases = [ (test_format_yuv(1, 1, 8), Subsampling::Yuv420), (test_format_yuv(1, 0, 8), Subsampling::Yuv422), (test_format_yuv(0, 0, 8), Subsampling::Yuv444), - ] { - let layout = layout_from_format(fmt, 160, 120).unwrap(); + ]; + for (format, expected) in cases { + let layout = layout_from_format(format, 160, 120).unwrap(); + assert_eq!(layout.subsampling, expected); assert_eq!(layout.depth, Depth::Eight); } @@ -70,9 +78,9 @@ fn accepts_the_three_supported_subsamplings_at_eight_bit() { #[test] fn rejects_rgb_with_a_clear_message() { - let err = layout_from_format(test_format_rgb(8), 160, 120) - .unwrap_err() - .to_string(); + let format = test_format_rgb(8); + let err = layout_from_format(format, 160, 120).unwrap_err().to_string(); + assert!(err.to_lowercase().contains("rgb"), "got {err}"); assert!( err.to_lowercase().contains("yuv"), @@ -82,17 +90,17 @@ fn rejects_rgb_with_a_clear_message() { #[test] fn rejects_float_clips_with_a_clear_message() { - let err = layout_from_format(test_format_yuv_float(), 160, 120) - .unwrap_err() - .to_string(); + let format = test_format_yuv_float(); + let err = layout_from_format(format, 160, 120).unwrap_err().to_string(); + assert!(err.to_lowercase().contains("float"), "got {err}"); } #[test] fn rejects_gray_with_a_clear_message() { - let err = layout_from_format(test_format_gray(8), 160, 120) - .unwrap_err() - .to_string(); + let format = test_format_gray(8); + let err = layout_from_format(format, 160, 120).unwrap_err().to_string(); + assert!(err.to_lowercase().contains("gray"), "got {err}"); assert!( err.to_lowercase().contains("yuv"), @@ -106,10 +114,11 @@ fn strength_is_rejected_for_nl4d() { strength: Some(0.5), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("strength"), "got {err}"); assert!( err.to_lowercase().contains("nl4d"), @@ -123,10 +132,11 @@ fn luma_lambda_ht_is_rejected_for_nlmeans() { luma_lambda_ht: Some(1.2), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("lambda_ht"), "got {err}"); assert!( err.to_lowercase().contains("nlmeans"), @@ -140,10 +150,11 @@ fn patch_radius_is_rejected_for_nl4d() { patch_radius: Some(3), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("patch_radius"), "got {err}"); assert!( err.to_lowercase().contains("nl4d"), @@ -157,10 +168,11 @@ fn search_radius_is_rejected_for_nl4d() { search_radius: Some(3), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("search_radius"), "got {err}"); assert!( err.to_lowercase().contains("nl4d"), @@ -172,14 +184,16 @@ fn search_radius_is_rejected_for_nl4d() { fn the_nl4d_mismatch_error_wins_over_the_stack_guard() { // SAFETY: single-threaded test, no denoiser thread exists yet. unsafe { std::env::remove_var("RUST_MIN_STACK") }; + let raw = RawParams { search_radius: Some(6), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap_err() .to_string(); + assert!( err.to_lowercase().contains("nl4d"), "search_radius on nl4d should fail as a mismatched parameter before the stack guard runs, got {err}" @@ -197,8 +211,9 @@ fn per_plane_overrides_reach_plane_options() { chroma_strength: Some(0.3), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + assert_eq!(opts.luma_strength, Some(0.6)); assert_eq!(opts.chroma_strength, Some(0.3)); } @@ -209,10 +224,11 @@ fn sigma_reaches_nl4d_options() { sigma: Some(6.0), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nl4d, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::Nl4d(nl4d) => { + av_denoise::Algorithm::Nl4d(nl4d) => { assert_eq!(nl4d.sigma, Some(6.0_f32)); }, other => panic!("expected Nl4d algorithm, got {other:?}"), @@ -226,10 +242,11 @@ fn sigma_is_rejected_for_nlmeans_fast() { variant: Some("fast".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("sigma"), "got {err}"); assert!( err.to_lowercase().contains("fast"), @@ -244,10 +261,11 @@ fn sigma_reaches_hq_sigma_override_under_variant_hq() { variant: Some("hq".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => { + av_denoise::Algorithm::NlmeansHq(hq) => { assert_eq!(hq.hq.sigma_override, Some(6.0_f32)); }, other => panic!("expected NlmeansHq algorithm, got {other:?}"), @@ -260,10 +278,11 @@ fn variant_hq_produces_the_hq_algorithm() { variant: Some("hq".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + assert!( - matches!(opts.algorithm, av_denoise_core::Algorithm::NlmeansHq(_)), + matches!(opts.algorithm, av_denoise::Algorithm::NlmeansHq(_)), "got {:?}", opts.algorithm ); @@ -275,10 +294,11 @@ fn variant_fast_produces_the_fast_algorithm() { variant: Some("fast".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + assert!( - matches!(opts.algorithm, av_denoise_core::Algorithm::Nlmeans(_)), + matches!(opts.algorithm, av_denoise::Algorithm::Nlmeans(_)), "got {:?}", opts.algorithm ); @@ -287,10 +307,11 @@ fn variant_fast_produces_the_fast_algorithm() { #[test] fn no_variant_given_defaults_to_the_hq_algorithm() { let raw = RawParams::default(); - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + assert!( - matches!(opts.algorithm, av_denoise_core::Algorithm::NlmeansHq(_)), + matches!(opts.algorithm, av_denoise::Algorithm::NlmeansHq(_)), "got {:?}", opts.algorithm ); @@ -302,10 +323,11 @@ fn unrecognised_variant_errors_clearly() { variant: Some("turbo".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.contains("turbo"), "got {err}"); assert!( err.to_lowercase().contains("fast") && err.to_lowercase().contains("hq"), @@ -316,8 +338,9 @@ fn unrecognised_variant_errors_clearly() { #[test] fn nlmeans_default_temporal_radius_is_two() { let raw = RawParams::default(); - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + assert_eq!(opts.mode, DenoisingMode::Temporal { radius: 2 }); } @@ -327,8 +350,9 @@ fn nlmeans_explicit_zero_temporal_radius_stays_spacial() { temporal_radius: Some(0), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + assert_eq!(opts.mode, DenoisingMode::Spacial); } @@ -338,10 +362,11 @@ fn the_hq_arm_sets_windowed_noise_estimation() { variant: Some("hq".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => { + av_denoise::Algorithm::NlmeansHq(hq) => { assert!(hq.hq.windowed_noise_estimation); }, other => panic!("expected NlmeansHq algorithm, got {other:?}"), @@ -352,14 +377,16 @@ fn the_hq_arm_sets_windowed_noise_estimation() { fn a_large_search_radius_is_rejected_when_the_stack_is_not_raised() { // SAFETY: single-threaded test, no denoiser thread exists yet. unsafe { std::env::remove_var("RUST_MIN_STACK") }; + let raw = RawParams { search_radius: Some(6), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.contains("RUST_MIN_STACK"), "got {err}"); } @@ -372,10 +399,11 @@ fn sigma_scale_reaches_hq_params_under_variant_hq() { sigma_scale: Some(1.5), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => { + av_denoise::Algorithm::NlmeansHq(hq) => { assert!((hq.hq.sigma_scale - 1.5).abs() < f32::EPSILON); }, other => panic!("expected NlmeansHq algorithm, got {other:?}"), @@ -388,10 +416,11 @@ fn sigma_scale_reaches_nl4d_options() { sigma_scale: Some(1.5), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nl4d, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::Nl4d(nl4d) => { + av_denoise::Algorithm::Nl4d(nl4d) => { assert!((nl4d.sigma_scale - 1.5).abs() < f32::EPSILON); }, other => panic!("expected Nl4d algorithm, got {other:?}"), @@ -405,10 +434,11 @@ fn sigma_scale_is_rejected_for_nlmeans_fast() { variant: Some("fast".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("sigma_scale"), "got {err}"); assert!( err.to_lowercase().contains("fast"), @@ -425,10 +455,11 @@ fn prefilter_none_and_empty_parse() { prefilter: Some(raw_prefilter.to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => { + av_denoise::Algorithm::NlmeansHq(hq) => { assert!(matches!(hq.nlm.prefilter, PrefilterMode::None)); }, other => panic!("expected NlmeansHq algorithm, got {other:?}"), @@ -442,10 +473,11 @@ fn prefilter_bilateral_parses() { prefilter: Some("bilateral:3.0,0.02".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => match hq.nlm.prefilter { + av_denoise::Algorithm::NlmeansHq(hq) => match hq.nlm.prefilter { PrefilterMode::Bilateral { sigma_s, sigma_r } => { assert!((sigma_s - 3.0).abs() < f32::EPSILON); assert!((sigma_r - 0.02).abs() < f32::EPSILON); @@ -462,10 +494,11 @@ fn prefilter_nlm_variants_parse() { prefilter: Some("nlm:0.8".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => match hq.nlm.prefilter { + av_denoise::Algorithm::NlmeansHq(hq) => match hq.nlm.prefilter { PrefilterMode::NlmSpatial { strength_scale } => { assert!((strength_scale - 0.8).abs() < f32::EPSILON); }, @@ -481,27 +514,25 @@ fn prefilter_unknown_string_errors_clearly() { prefilter: Some("garbage".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("prefilter"), "got {err}"); } -/// `parse_prefilter`'s string grammar has no form that produces -/// `External`, and the boundary in `plane_options_from` rejects it if it -/// ever did, since this plugin has no way to supply the reference frame -/// `External` needs. #[test] fn prefilter_external_cannot_be_produced() { let raw = RawParams { prefilter: Some("external".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!( err.to_lowercase().contains("prefilter"), "'external' has no string form, so it should fail as an unknown prefilter, got {err}" @@ -514,10 +545,11 @@ fn prefilter_is_rejected_for_nl4d() { prefilter: Some("none".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("prefilter"), "got {err}"); assert!(err.to_lowercase().contains("nl4d"), "got {err}"); } @@ -532,7 +564,7 @@ fn preset_resolves_the_same_dials_as_core_for_nlmeans() { preset: Some(name.to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); let want_variant = nlmeans_variant_for(preset); @@ -546,25 +578,27 @@ fn preset_resolves_the_same_dials_as_core_for_nlmeans() { radius: want_temporal_radius, } }; + assert_eq!( opts.mode, want_mode, "preset {name} resolved the wrong temporal radius" ); - let got_variant_is_hq = matches!(opts.algorithm, av_denoise_core::Algorithm::NlmeansHq(_)); + let got_variant_is_hq = matches!(opts.algorithm, av_denoise::Algorithm::NlmeansHq(_)); + assert_eq!( got_variant_is_hq, want_variant == NlmeansVariant::Hq, "preset {name} resolved the wrong variant" ); - if let av_denoise_core::Algorithm::NlmeansHq(hq) = opts.algorithm { + if let av_denoise::Algorithm::NlmeansHq(hq) = opts.algorithm { assert_eq!( hq.nlm.tuning.search_radius, Some(want_search_radius), "preset {name} resolved the wrong search radius" ); - } else if let av_denoise_core::Algorithm::Nlmeans(nlm) = opts.algorithm { + } else if let av_denoise::Algorithm::Nlmeans(nlm) = opts.algorithm { assert_eq!( nlm.tuning.search_radius, Some(want_search_radius), @@ -582,7 +616,7 @@ fn preset_resolves_the_same_dials_as_core_for_nl4d() { preset: Some(name.to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nl4d, layout).unwrap(); let want_temporal_radius = nl4d_temporal_radius_for(preset); @@ -597,7 +631,7 @@ fn preset_resolves_the_same_dials_as_core_for_nl4d() { ); match opts.algorithm { - av_denoise_core::Algorithm::Nl4d(nl4d) => { + av_denoise::Algorithm::Nl4d(nl4d) => { assert_eq!( nl4d.spatial_radius, want_spatial_radius, "preset {name} resolved the wrong spatial radius" @@ -611,12 +645,14 @@ fn preset_resolves_the_same_dials_as_core_for_nl4d() { #[test] fn unset_preset_defaults_to_base() { let raw = RawParams::default(); - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nl4d, layout).unwrap(); + let base_temporal_radius = nl4d_temporal_radius_for(Preset::Base); + assert_eq!( opts.mode, DenoisingMode::Temporal { - radius: nl4d_temporal_radius_for(Preset::Base) + radius: base_temporal_radius } ); } @@ -627,10 +663,11 @@ fn unrecognised_preset_errors_clearly() { preset: Some("turbo".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.contains("turbo"), "got {err}"); } @@ -641,24 +678,25 @@ fn explicit_temporal_radius_overrides_the_preset() { temporal_radius: Some(3), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nl4d, layout).unwrap(); + assert_eq!(opts.mode, DenoisingMode::Temporal { radius: 3 }); } #[test] fn explicit_variant_overrides_the_preset() { let raw = RawParams { - // `base` resolves to the `hq` variant; an explicit `variant` - // must win anyway. + // `base` resolves to the `hq` variant, so the explicit `variant` must win. preset: Some("base".to_string()), variant: Some("fast".to_string()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + assert!( - matches!(opts.algorithm, av_denoise_core::Algorithm::Nlmeans(_)), + matches!(opts.algorithm, av_denoise::Algorithm::Nlmeans(_)), "got {:?}", opts.algorithm ); @@ -669,10 +707,11 @@ fn explicit_variant_overrides_the_preset() { #[test] fn motion_compensation_defaults_to_false() { let raw = RawParams::default(); - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => { + av_denoise::Algorithm::NlmeansHq(hq) => { assert!(matches!(hq.nlm.motion_compensation, MotionCompensationMode::None)); }, other => panic!("expected NlmeansHq algorithm, got {other:?}"), @@ -685,10 +724,11 @@ fn motion_compensation_true_turns_it_on() { motion_compensation: Some(true), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => { + av_denoise::Algorithm::NlmeansHq(hq) => { assert!(matches!( hq.nlm.motion_compensation, MotionCompensationMode::Mvtools { .. } @@ -704,10 +744,11 @@ fn motion_compensation_false_stays_off() { motion_compensation: Some(false), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::NlmeansHq(hq) => { + av_denoise::Algorithm::NlmeansHq(hq) => { assert!(matches!(hq.nlm.motion_compensation, MotionCompensationMode::None)); }, other => panic!("expected NlmeansHq algorithm, got {other:?}"), @@ -720,10 +761,11 @@ fn motion_compensation_is_rejected_for_nl4d() { motion_compensation: Some(true), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap_err() .to_string(); + assert!(err.to_lowercase().contains("motion_compensation"), "got {err}"); assert!(err.to_lowercase().contains("nl4d"), "got {err}"); } @@ -734,10 +776,11 @@ fn lambda_ht_scale_reaches_nl4d_options() { lambda_ht_scale: Some(1.15), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nl4d, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::Nl4d(nl4d) => { + av_denoise::Algorithm::Nl4d(nl4d) => { assert!((nl4d.lambda_ht_scale - 1.15).abs() < 1e-6) }, other => panic!("expected Nl4d, got {other:?}"), @@ -750,10 +793,11 @@ fn shared_lambda_ht_reaches_nl4d_options() { lambda_ht: Some(4.6), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let opts = plane_options_from(&raw, AlgorithmKind::Nl4d, layout).unwrap(); + match opts.algorithm { - av_denoise_core::Algorithm::Nl4d(nl4d) => assert_eq!(nl4d.lambda_ht, Some(4.6)), + av_denoise::Algorithm::Nl4d(nl4d) => assert_eq!(nl4d.lambda_ht, Some(4.6)), other => panic!("expected Nl4d, got {other:?}"), } } @@ -765,12 +809,13 @@ fn spatial_radius_overrides_the_preset() { preset: Some("base".into()), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => nl4d.spatial_radius, + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => nl4d.spatial_radius, other => panic!("expected Nl4d, got {other:?}"), } }; @@ -780,12 +825,13 @@ fn spatial_radius_overrides_the_preset() { spatial_radius: Some((from_preset + 1) as i64), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => assert_eq!(nl4d.spatial_radius, from_preset + 1), + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => assert_eq!(nl4d.spatial_radius, from_preset + 1), other => panic!("expected Nl4d, got {other:?}"), } } @@ -796,12 +842,13 @@ fn refine_reaches_nl4d_options() { refine: Some(3), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => assert_eq!(nl4d.refine, 3), + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => assert_eq!(nl4d.refine, 3), other => panic!("expected Nl4d, got {other:?}"), } } @@ -809,12 +856,13 @@ fn refine_reaches_nl4d_options() { #[test] fn noise_map_defaults_to_on() { let raw = RawParams::default(); - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => assert!(nl4d.noise_map), + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => assert!(nl4d.noise_map), other => panic!("expected Nl4d, got {other:?}"), } } @@ -825,12 +873,13 @@ fn noise_map_false_turns_it_off() { noise_map: Some(false), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => assert!(!nl4d.noise_map), + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => assert!(!nl4d.noise_map), other => panic!("expected Nl4d, got {other:?}"), } } @@ -841,22 +890,24 @@ fn noise_map_is_rejected_for_nlmeans() { noise_map: Some(true), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.contains("noise_map"), "got {err}"); } #[test] fn pooled_threshold_defaults_to_on() { let raw = RawParams::default(); - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => assert!(nl4d.pooled_threshold), + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => assert!(nl4d.pooled_threshold), other => panic!("expected Nl4d, got {other:?}"), } } @@ -867,12 +918,13 @@ fn pooled_threshold_false_turns_it_off() { pooled_threshold: Some(false), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => assert!(!nl4d.pooled_threshold), + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => assert!(!nl4d.pooled_threshold), other => panic!("expected Nl4d, got {other:?}"), } } @@ -883,17 +935,18 @@ fn pooled_threshold_is_rejected_for_nlmeans() { pooled_threshold: Some(true), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.contains("pooled_threshold"), "got {err}"); } #[test] fn the_four_nl4d_dials_are_rejected_for_nlmeans() { - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - for (name, raw) in [ + let layout = yuv420_layout(); + let cases = [ ( "lambda_ht_scale", RawParams { @@ -922,10 +975,12 @@ fn the_four_nl4d_dials_are_rejected_for_nlmeans() { ..RawParams::default() }, ), - ] { + ]; + for (name, raw) in cases { let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.contains(name), "error should name {name}, got {err}"); } } @@ -939,12 +994,13 @@ fn strength_map_params_reach_nl4d_options() { flat_texture_cut: Some(0.3), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => { + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => { assert_eq!(nl4d.flat_boost, 2.0); assert_eq!(nl4d.chroma_flat_boost, 1.2); assert_eq!(nl4d.shadow_soften, 0.8); @@ -957,13 +1013,14 @@ fn strength_map_params_reach_nl4d_options() { #[test] fn strength_map_params_default_to_the_library_values() { let raw = RawParams::default(); - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); - let defaults = av_denoise_core::Nl4dOptions::default(); - match plane_options_from(&raw, AlgorithmKind::Nl4d, layout) + let layout = yuv420_layout(); + let defaults = av_denoise::Nl4dOptions::default(); + let algorithm = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap() - .algorithm - { - av_denoise_core::Algorithm::Nl4d(nl4d) => { + .algorithm; + + match algorithm { + av_denoise::Algorithm::Nl4d(nl4d) => { assert_eq!(nl4d.flat_boost, defaults.flat_boost); assert_eq!(nl4d.chroma_flat_boost, defaults.chroma_flat_boost); assert_eq!(nl4d.shadow_soften, defaults.shadow_soften); @@ -975,7 +1032,7 @@ fn strength_map_params_default_to_the_library_values() { #[test] fn strength_map_params_are_rejected_for_nlmeans() { - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let cases = [ ( "flat_boost", @@ -1010,6 +1067,7 @@ fn strength_map_params_are_rejected_for_nlmeans() { let err = plane_options_from(&raw, AlgorithmKind::Nlmeans, layout) .unwrap_err() .to_string(); + assert!(err.contains(name), "{name}: got {err}"); } } @@ -1020,9 +1078,10 @@ fn a_negative_spatial_radius_is_rejected() { spatial_radius: Some(-1), ..RawParams::default() }; - let layout = layout_from_format(test_format_yuv(1, 1, 8), 160, 120).unwrap(); + let layout = yuv420_layout(); let err = plane_options_from(&raw, AlgorithmKind::Nl4d, layout) .unwrap_err() .to_string(); + assert!(err.contains("spatial_radius"), "got {err}"); } diff --git a/av-denoise-vs/tests/vs_harness.py b/av-denoise-vs/tests/vs_harness.py index ed5aeb9..52ce20d 100644 --- a/av-denoise-vs/tests/vs_harness.py +++ b/av-denoise-vs/tests/vs_harness.py @@ -309,7 +309,7 @@ def _parity_source_filter(): PARITY_TEMPORAL_RADIUS = 2 # nl4d's WindowSpan is `{behind: 2 * radius, ahead: 2 * radius}` -# (av-denoise-core/src/denoiser.rs, `PlanarDenoiser::window_span`), so a +# (av-denoise/src/planar/mod.rs, `PlanarDenoiser::window_span`), so a # window reaches `2 * radius` frames behind its centre. Those are the # only output frames close enough to the clip's start for the two front # ends' leading-edge padding to differ: `reseed` (what the plugin's @@ -324,7 +324,7 @@ def _parity_source_filter(): # The bound the leading edge frames are allowed to drift within, the # same value and reasoning as `BEHIND_EDGE_TOLERANCE` in -# av-denoise-core/src/frame/tests.rs: full-range luma codes span 255, +# av-denoise/src/planar/tests/mod.rs: full-range luma codes span 255, # and the extra duplicated history moves the result by at most a # handful of 8-bit codes. PARITY_LEADING_EDGE_TOLERANCE = 8 diff --git a/av-denoise/Cargo.toml b/av-denoise/Cargo.toml index 1755532..a9a2a0f 100644 --- a/av-denoise/Cargo.toml +++ b/av-denoise/Cargo.toml @@ -32,6 +32,10 @@ required-features = ["binary"] [dependencies] av-denoise-core.workspace = true anyhow.workspace = true +bon.workspace = true +cubecl.workspace = true +etcetera = "0.11" +thiserror.workspace = true tracing.workspace = true strum.workspace = true strum_macros.workspace = true @@ -56,14 +60,19 @@ indicatif = { version = "0.18", optional = true, default-features = false } tracing-indicatif = { version = "0.3", optional = true } tracing-subscriber = { version = "0.3", optional = true, features = ["env-filter"] } +[dev-dependencies] +clap = { version = "4", features = ["derive"] } +tempfile = "3" +wgpu = "29" + [build-dependencies] pkg-config = { version = "0.3.34", optional = true } [features] -vulkan = ["av-denoise-core/vulkan"] -metal = ["av-denoise-core/metal"] -cuda = ["av-denoise-core/cuda"] -rocm = ["av-denoise-core/rocm"] +vulkan = ["av-denoise-core/vulkan", "cubecl/vulkan"] +metal = ["av-denoise-core/metal", "cubecl/metal"] +cuda = ["av-denoise-core/cuda", "cubecl/cuda"] +rocm = ["av-denoise-core/rocm", "cubecl/hip"] binary = [ "bytesize", "clap", @@ -81,3 +90,11 @@ binary = [ # Links ffms2 and FFmpeg statically from the vcpkg tree. See `just build-static`. static-ffms2 = ["binary", "av-decoders/ffms2_static", "dep:pkg-config"] default = ["vulkan"] + +[[bench]] +name = "denoise" +harness = false + +[[bench]] +name = "reseed" +harness = false diff --git a/av-denoise-core/benches/denoise.rs b/av-denoise/benches/denoise.rs similarity index 59% rename from av-denoise-core/benches/denoise.rs rename to av-denoise/benches/denoise.rs index bcb7c0b..edc099c 100644 --- a/av-denoise-core/benches/denoise.rs +++ b/av-denoise/benches/denoise.rs @@ -1,20 +1,22 @@ use std::time::{Duration, Instant}; -use av_denoise_core::accelerate::Accelerator; -use av_denoise_core::{ +use av_denoise::accelerate::Accelerator; +use av_denoise::{ Algorithm, ChannelMode, - Denoiser, + DenoiserError, DenoiserOptions, DenoisingMode, Device, + HostDenoiser, MotionCompensationMode, NlmeansOptions, PrefilterMode, }; +use clap::Parser; -const W: u32 = 1920; -const H: u32 = 1080; +const WIDTH: u32 = 1920; +const HEIGHT: u32 = 1080; const WARMUP: usize = 5; const ITERS: usize = 100; @@ -23,39 +25,42 @@ const BILATERAL_SIGMA_S: f32 = 3.0; const BILATERAL_SIGMA_R: f32 = 0.02; #[derive(clap::Parser, Debug)] -#[command(about = "End-to-end Denoiser benchmark", long_about = None)] +#[command(about = "End-to-end HostDenoiser benchmark", long_about = None)] struct Cli { - /// GPU device to bind to. Format: `default`, `discrete[:N]`, - /// `integrated[:N]`, `virtual[:N]`, or `cpu`. + /// GPU device to bind to, one of `default`, `discrete[:N]`, `integrated[:N]`, `virtual[:N]` or `cpu`. #[arg(long, default_value = "default")] device: Device, - /// Accelerator priority list (comma-delimited). Defaults to all - /// compiled-in accelerators. - #[arg(long, value_delimiter = ',', default_values_t = av_denoise_core::accelerate::get_default_accelerators())] + /// Accelerator priority list, comma-delimited. Defaults to all compiled-in accelerators. + #[arg(long, value_delimiter = ',', default_values_t = av_denoise::accelerate::get_default_accelerators())] accelerators: Vec, - /// Swallowed: cargo passes this when invoking the bench binary. + /// Swallowed, since cargo passes this when invoking the bench binary. #[arg(long, hide = true)] bench: bool, } -fn make_synthetic_frame(w: u32, h: u32, ch: u32) -> Vec { - let mut data = Vec::with_capacity((w * h * ch) as usize); - for y in 0..h { - for x in 0..w { +/// One 8-bit plane per channel, a smooth pattern plus hashed noise. +fn make_synthetic_planes(width: u32, height: u32, channels: u32) -> Vec> { + let mut planes = vec![Vec::with_capacity((width * height) as usize); channels as usize]; + + for y in 0..height { + for x in 0..width { let base = 0.5 + 0.2 * (x as f32 * 0.05).sin() * (y as f32 * 0.03).cos(); - for c in 0..ch { - let seed = (y * w + x) * ch + c; + + for (channel, plane) in planes.iter_mut().enumerate() { + let seed = (y * width + x) * channels + channel as u32; let hash = seed .wrapping_mul(2654435761) .wrapping_add(seed.wrapping_mul(340573321)); let noise = (hash as f32 / u32::MAX as f32 - 0.5) * 0.1; - data.push((base + noise).clamp(0.0, 1.0)); + let value = (base + noise).clamp(0.0, 1.0); + plane.push((value * 255.0 + 0.5) as u8); } } } - data + + planes } struct BenchResult { @@ -88,11 +93,12 @@ fn options(channel_mode: ChannelMode, mode: DenoisingMode, algorithm: Algorithm) /// The fast NLM path with a prefilter and a motion-compensation mode. fn nlm(prefilter: PrefilterMode, motion_compensation: MotionCompensationMode) -> Algorithm { - Algorithm::Nlmeans(NlmeansOptions { + let nlmeans_options = NlmeansOptions { prefilter, motion_compensation, ..NlmeansOptions::default() - }) + }; + Algorithm::Nlmeans(nlmeans_options) } fn bench_push_recv( @@ -103,48 +109,46 @@ fn bench_push_recv( mode: DenoisingMode, algorithm: Algorithm, ) -> Result { - let ch = channel_mode.count(); - let frame = make_synthetic_frame(W, H, ch); + let channels = channel_mode.count(); + let planes = make_synthetic_planes(WIDTH, HEIGHT, channels); + let frame: Vec<&[u8]> = planes.iter().map(Vec::as_slice).collect(); - let mut denoiser = Denoiser::create(accelerators, device, W, H, options(channel_mode, mode, algorithm))?; + let denoiser_options = options(channel_mode, mode, algorithm); + let mut denoiser = HostDenoiser::create(accelerators, device, WIDTH, HEIGHT, denoiser_options)?; let accelerator = denoiser.selected_accelerator(); - // Fill the temporal window so subsequent push/recv steady-state - // lines up. NLMeans mirrors the first pushed frame into the leading - // `R` ring slots, so it emits early and pushing `window - 1` frames - // can trip `QueueFull` at radius ≥ 2. nl4d fills its ring with real - // frames instead. Use a defensive push that drains a pending if the - // queue is full, then drain everything before steady-state. + // Fill the temporal window so the steady-state push/recv lines up. NLMeans mirrors the first + // frame into the leading ring slots and emits early, so pushing `window - 1` frames can hit + // `QueueFull` at radius 2 or more. A full queue drains one frame before the push is retried. let temporal_radius = match mode { DenoisingMode::Spacial => 0, DenoisingMode::Temporal { radius } => radius, }; let window = 2 * temporal_radius + 1; for _ in 0..window.saturating_sub(1) { - if let Err(av_denoise_core::DenoiserError::QueueFull) = denoiser.push_frame(&frame) { - let _ = denoiser.recv_frame()?; - denoiser.push_frame(&frame)?; + if let Err(DenoiserError::QueueFull) = denoiser.push(&frame) { + let _ = denoiser.recv()?; + denoiser.push(&frame)?; } } - while denoiser.recv_frame()?.is_some() {} + + while denoiser.recv()?.is_some() {} for _ in 0..WARMUP { - denoiser.push_frame(&frame)?; - let _ = denoiser.recv_frame()?; + denoiser.push(&frame)?; + let _ = denoiser.recv()?; } let mut times = Vec::with_capacity(ITERS); for _ in 0..ITERS { let start = Instant::now(); - denoiser.push_frame(&frame)?; - let _out = denoiser.recv_frame()?; + denoiser.push(&frame)?; + let _received = denoiser.recv()?; times.push(start.elapsed()); } - // Drain the temporal tail so every pushed frame is accounted for - // before the denoiser drops. An unpolled `Pending` has started no - // readback and is free to drop, so this is bookkeeping rather than - // a safety requirement. + // Drain the temporal tail so every pushed frame is accounted for. An unpolled `Pending` is free + // to drop, so this is bookkeeping rather than a safety requirement. denoiser.flush(|_| {})?; let total: Duration = times.iter().sum(); @@ -166,12 +170,11 @@ fn bench_push_recv( fn main() { // SAFETY: single-threaded at entry, no race possible. - unsafe { av_denoise_core::raise_codegen_stack_limit() }; + unsafe { av_denoise::raise_codegen_stack_limit() }; - use clap::Parser; let cli = Cli::parse(); - println!("Denoiser E2E Benchmarks - {W}×{H}"); + println!("HostDenoiser E2E Benchmarks - {WIDTH}×{HEIGHT}"); println!(" warmup={WARMUP}, timed={ITERS}"); println!(" device: {:?}", cli.device); println!(" accelerators: {:?}", cli.accelerators); @@ -181,102 +184,104 @@ fn main() { sigma_s: BILATERAL_SIGMA_S, sigma_r: BILATERAL_SIGMA_R, }; - let mc = MotionCompensationMode::mvtools_default(); + let motion_compensation = MotionCompensationMode::mvtools_default(); - // Side-by-side ordering: each temporal config is followed by its - // motion-compensation variant so the cost delta from `--motion-compensation` - // is visible on adjacent rows. + let plain = nlm(PrefilterMode::None, MotionCompensationMode::None); + let plain_mc = nlm(PrefilterMode::None, motion_compensation); + let bilateral_only = nlm(bilateral, MotionCompensationMode::None); + let bilateral_mc = nlm(bilateral, motion_compensation); + + // Each temporal config is followed by its motion-compensation variant, so the cost of + // `--motion-compensation` shows on adjacent rows. let configs: &[(&str, ChannelMode, DenoisingMode, Algorithm)] = &[ - ( - "spatial_luma", - ChannelMode::Luma, - DenoisingMode::Spacial, - nlm(PrefilterMode::None, MotionCompensationMode::None), - ), + ("spatial_luma", ChannelMode::Luma, DenoisingMode::Spacial, plain), ( "spatial_chroma", ChannelMode::Chroma, DenoisingMode::Spacial, - nlm(PrefilterMode::None, MotionCompensationMode::None), - ), - ( - "spatial_yuv", - ChannelMode::Yuv, - DenoisingMode::Spacial, - nlm(PrefilterMode::None, MotionCompensationMode::None), + plain, ), + ("spatial_yuv", ChannelMode::Yuv, DenoisingMode::Spacial, plain), ( "temporal_r1_yuv", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 1 }, - nlm(PrefilterMode::None, MotionCompensationMode::None), + plain, ), ( "temporal_r1_yuv+mc", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 1 }, - nlm(PrefilterMode::None, mc), + plain_mc, ), ( "temporal_r2_yuv", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 2 }, - nlm(PrefilterMode::None, MotionCompensationMode::None), + plain, ), ( "temporal_r2_yuv+mc", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 2 }, - nlm(PrefilterMode::None, mc), + plain_mc, ), ( "spatial_luma+bilateral", ChannelMode::Luma, DenoisingMode::Spacial, - nlm(bilateral, MotionCompensationMode::None), + bilateral_only, ), ( "spatial_chroma+bilateral", ChannelMode::Chroma, DenoisingMode::Spacial, - nlm(bilateral, MotionCompensationMode::None), + bilateral_only, ), ( "spatial_yuv+bilateral", ChannelMode::Yuv, DenoisingMode::Spacial, - nlm(bilateral, MotionCompensationMode::None), + bilateral_only, ), ( "temporal_r1_yuv+bilateral", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 1 }, - nlm(bilateral, MotionCompensationMode::None), + bilateral_only, ), ( "temporal_r1_yuv+bilateral+mc", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 1 }, - nlm(bilateral, mc), + bilateral_mc, ), ( "temporal_r2_yuv+bilateral", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 2 }, - nlm(bilateral, MotionCompensationMode::None), + bilateral_only, ), ( "temporal_r2_yuv+bilateral+mc", ChannelMode::Yuv, DenoisingMode::Temporal { radius: 2 }, - nlm(bilateral, mc), + bilateral_mc, ), ]; - for (name, ch, mode, algorithm) in configs { - match bench_push_recv(name, &cli.accelerators, &cli.device, *ch, *mode, *algorithm) { + for (name, channel_mode, mode, algorithm) in configs { + let outcome = bench_push_recv( + name, + &cli.accelerators, + &cli.device, + *channel_mode, + *mode, + *algorithm, + ); + match outcome { Ok(result) => result.print(), - Err(err) => eprintln!("[{name}] failed: {err:?}"), + Err(error) => eprintln!("[{name}] failed: {error:?}"), } } } diff --git a/av-denoise-core/benches/reseed.rs b/av-denoise/benches/reseed.rs similarity index 58% rename from av-denoise-core/benches/reseed.rs rename to av-denoise/benches/reseed.rs index 98c9856..abc40a7 100644 --- a/av-denoise-core/benches/reseed.rs +++ b/av-denoise/benches/reseed.rs @@ -1,25 +1,26 @@ use std::time::{Duration, Instant}; -use av_denoise_core::accelerate::Accelerator; -use av_denoise_core::{ +use av_denoise::accelerate::Accelerator; +use av_denoise::{ Algorithm, ChannelIntent, ChannelMode, - Denoiser, DenoiserOptions, DenoisingMode, Depth, Device, FrameLayout, + HostDenoiser, PlanarDenoiser, PlaneOptions, Planes, Subsampling, push_needs_retry, }; +use clap::Parser; -const W: u32 = 1920; -const H: u32 = 1080; +const WIDTH: u32 = 1920; +const HEIGHT: u32 = 1080; const RADIUS: u32 = 2; const WARMUP: usize = 5; @@ -28,25 +29,23 @@ const ITERS: usize = 100; #[derive(clap::Parser, Debug)] #[command(about = "Cost of a reseed relative to a sequential frame", long_about = None)] struct Cli { - /// GPU device to bind to. Format: `default`, `discrete[:N]`, - /// `integrated[:N]`, `virtual[:N]`, or `cpu`. + /// GPU device to bind to, one of `default`, `discrete[:N]`, `integrated[:N]`, `virtual[:N]` or `cpu`. #[arg(long, default_value = "default")] device: Device, - /// Accelerator priority list (comma-delimited). Defaults to all - /// compiled-in accelerators. - #[arg(long, value_delimiter = ',', default_values_t = av_denoise_core::accelerate::get_default_accelerators())] + /// Accelerator priority list, comma-delimited. Defaults to all compiled-in accelerators. + #[arg(long, value_delimiter = ',', default_values_t = av_denoise::accelerate::get_default_accelerators())] accelerators: Vec, - /// Swallowed: cargo passes this when invoking the bench binary. + /// Swallowed, since cargo passes this when invoking the bench binary. #[arg(long, hide = true)] bench: bool, } fn layout() -> FrameLayout { FrameLayout { - width: W, - height: H, + width: WIDTH, + height: HEIGHT, subsampling: Subsampling::Yuv420, depth: Depth::Eight, } @@ -66,19 +65,19 @@ fn plane_options(accelerators: &[Accelerator], device: &Device) -> PlaneOptions } } -/// A small xorshift generator, deterministic across runs so the synthetic -/// clip does not vary between executions. -fn pseudo_random(mut x: u64) -> u64 { - x ^= x << 13; - x ^= x >> 7; - x ^= x << 17; - x +/// A small xorshift generator, so the synthetic clip is the same on every run. +fn pseudo_random(mut state: u64) -> u64 { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + state } -/// One plane's wire bytes for frame `frame_idx`: a spatial ramp across the -/// plane plus a per-frame offset and a deterministic dither, so a temporal -/// filter sees real signal and real noise to work with. -fn ramp_plane(pixels: usize, width: u32, frame_idx: usize, plane_seed: u64) -> Vec { +/// One plane's wire bytes for frame `frame_index`. +/// +/// A spatial ramp plus a per-frame offset and a deterministic dither give a temporal filter real +/// signal and real noise to work with. +fn ramp_plane(pixels: usize, width: u32, frame_index: usize, plane_seed: u64) -> Vec { let width = width.max(1) as usize; (0..pixels) @@ -86,8 +85,8 @@ fn ramp_plane(pixels: usize, width: u32, frame_idx: usize, plane_seed: u64) -> V let x = (i % width) as u32; let y = (i / width) as u32; let spatial = x.wrapping_add(y) % 120; - let frame_offset = (frame_idx as u32 * 7) % 60; - let seed = (i as u64) ^ (frame_idx as u64).wrapping_mul(0x9E3779B97F4A7C15) ^ plane_seed; + let frame_offset = (frame_index as u32 * 7) % 60; + let seed = (i as u64) ^ (frame_index as u64).wrapping_mul(0x9E3779B97F4A7C15) ^ plane_seed; let dither = (pseudo_random(seed) % 16) as u32; let value = 20 + spatial + frame_offset + dither; value.min(235) as u8 @@ -95,31 +94,35 @@ fn ramp_plane(pixels: usize, width: u32, frame_idx: usize, plane_seed: u64) -> V .collect() } -fn make_planes(layout: &FrameLayout, frame_idx: usize) -> Planes { - let (chroma_w, _) = layout.chroma_dims(); +fn make_planes(layout: &FrameLayout, frame_index: usize) -> Planes { + let (chroma_width, _) = layout.chroma_dims(); + let y_plane = ramp_plane(layout.luma_pixels(), layout.width, frame_index, 1); + let u_plane = ramp_plane(layout.chroma_pixels(), chroma_width, frame_index, 2); + let v_plane = ramp_plane(layout.chroma_pixels(), chroma_width, frame_index, 3); + Planes { - y: ramp_plane(layout.luma_pixels(), layout.width, frame_idx, 1), - u: ramp_plane(layout.chroma_pixels(), chroma_w, frame_idx, 2), - v: ramp_plane(layout.chroma_pixels(), chroma_w, frame_idx, 3), + y: y_plane, + u: u_plane, + v: v_plane, } } -/// `count` frames, each with distinct content, for building sliding -/// windows out of without re-generating a window's frames per call. +/// `count` frames with distinct content, so sliding windows are cut from it without regenerating +/// frames per call. fn make_clip(layout: &FrameLayout, count: usize) -> Vec { (0..count).map(|i| make_planes(layout, i)).collect() } -/// The accelerator a real denoiser would pick for `accelerators` and -/// `device`, read from a throwaway probe since `PlanarDenoiser` may own -/// up to three inner denoisers and exposes no single accelerator getter. +/// The accelerator a real denoiser would pick, read from a throwaway probe. +/// +/// `PlanarDenoiser` may own up to three inner denoisers and exposes no single accelerator getter. fn selected_accelerator(accelerators: &[Accelerator], device: &Device) -> Result { - let opts = DenoiserOptions::builder() + let probe_options = DenoiserOptions::builder() .channel_mode(ChannelMode::Luma) .mode(DenoisingMode::Spacial) .algorithm(Algorithm::default()) .build(); - let probe = Denoiser::create(accelerators, device, 4, 4, opts)?; + let probe = HostDenoiser::create(accelerators, device, 4, 4, probe_options)?; Ok(probe.selected_accelerator()) } @@ -164,59 +167,60 @@ fn summarise(name: &str, accelerator: Accelerator, times: &[Duration]) -> BenchR } } -/// Times `push` plus `recv` per output frame once the window is primed and -/// the stream is in steady state. +/// Times `push` plus `recv` per output frame once the window is primed and the stream is steady. fn bench_sequential(accelerators: &[Accelerator], device: &Device) -> Result { let layout = layout(); - let opts = plane_options(accelerators, device); - let mut denoiser = PlanarDenoiser::create(&opts, layout)?; + let options = plane_options(accelerators, device); + let mut denoiser = PlanarDenoiser::create(&options, layout)?; let accelerator = selected_accelerator(accelerators, device)?; let radius = denoiser.temporal_radius(); let window = 2 * radius + 1; let clip = make_clip(&layout, window as usize + WARMUP + ITERS); - // Prime the window outside the timed region so steady-state push/recv - // lines up: one recv for every push from here on. + // Priming happens outside the timed region, so from here on every push has one recv. for frame in &clip[..window.saturating_sub(1) as usize] { - if push_needs_retry(denoiser.push(frame))? { + let pushed = denoiser.push(frame); + if push_needs_retry(pushed)? { let _ = denoiser.recv()?; denoiser.push(frame)?; } } + while denoiser.recv()?.is_some() {} - let mut idx = window.saturating_sub(1) as usize; + let mut next_frame = window.saturating_sub(1) as usize; for _ in 0..WARMUP { - denoiser.push(&clip[idx])?; + denoiser.push(&clip[next_frame])?; let _ = denoiser.recv()?; - idx += 1; + next_frame += 1; } let mut times = Vec::with_capacity(ITERS); for _ in 0..ITERS { let start = Instant::now(); - denoiser.push(&clip[idx])?; - let _out = denoiser.recv()?; + denoiser.push(&clip[next_frame])?; + let _received = denoiser.recv()?; times.push(start.elapsed()); - idx += 1; + next_frame += 1; } - // Drain the trailing temporal frames before the denoiser drops, so no - // pending readback dies in flight with its GPU buffer still mapped. + // Drain the trailing temporal frames so every pushed frame is accounted for. An unpolled + // `Pending` is free to drop, so this is bookkeeping rather than a safety requirement. denoiser.flush(|_| {})?; - Ok(summarise("sequential", accelerator, ×)) + let result = summarise("sequential", accelerator, ×); + Ok(result) } -/// Times one `reseed` call over a fresh window per iteration. The -/// denoiser is built once, before timing starts, and every window is -/// built ahead of time too, so only `reseed` itself is on the clock. +/// Times one `reseed` call over a fresh window per iteration. +/// +/// The denoiser and every window are built before timing starts, so only `reseed` is on the clock. fn bench_reseed(accelerators: &[Accelerator], device: &Device) -> Result { let layout = layout(); - let opts = plane_options(accelerators, device); - let mut denoiser = PlanarDenoiser::create(&opts, layout)?; + let options = plane_options(accelerators, device); + let mut denoiser = PlanarDenoiser::create(&options, layout)?; let accelerator = selected_accelerator(accelerators, device)?; let radius = denoiser.temporal_radius() as usize; @@ -233,41 +237,43 @@ fn bench_reseed(accelerators: &[Accelerator], device: &Device) -> Result result, - Err(err) => { - eprintln!("[sequential] failed: {err:?}"); + Err(error) => { + eprintln!("[sequential] failed: {error:?}"); return; }, }; sequential.print(); - let reseed = match bench_reseed(&cli.accelerators, &cli.device) { + let reseed_outcome = bench_reseed(&cli.accelerators, &cli.device); + let reseed = match reseed_outcome { Ok(result) => result, - Err(err) => { - eprintln!("[reseed] failed: {err:?}"); + Err(error) => { + eprintln!("[reseed] failed: {error:?}"); return; }, }; diff --git a/av-denoise/src/backend/accelerate.rs b/av-denoise/src/backend/accelerate.rs new file mode 100644 index 0000000..45ab63b --- /dev/null +++ b/av-denoise/src/backend/accelerate.rs @@ -0,0 +1,60 @@ +//! The hardware backends kernels can run on + +use strum_macros::{Display, EnumIter, EnumString, IntoStaticStr}; + +/// A hardware backend that kernels can run on. +/// +/// It names a backend rather than a specific GPU, which is chosen separately with +/// [Device](crate::Device). Only the backends whose crate feature is enabled exist. +/// +/// Every accelerator runs kernels on a GPU. There is no software backend, because the collaborative +/// filter aggregates through atomic floating-point adds and cubecl's CPU runtime does not implement +/// atomics. [Device::Cpu](crate::Device::Cpu) still selects a software device where the platform +/// offers one, such as lavapipe under Vulkan. +#[derive(Debug, Copy, Clone, Eq, PartialEq, IntoStaticStr, EnumString, EnumIter, Display)] +#[strum(serialize_all = "snake_case")] +pub enum Accelerator { + #[cfg(any(feature = "cuda", docsrs))] + #[cfg_attr(docsrs, doc(cfg(feature = "cuda")))] + /// Runs kernels through the Nvidia CUDA backend, on Nvidia GPUs only. + Cuda, + #[cfg(any(feature = "vulkan", docsrs))] + #[cfg_attr(docsrs, doc(cfg(feature = "vulkan")))] + /// Runs kernels through the wgpu Vulkan backend. + /// + /// This is the lightest and most portable option, because it works on any platform and GPU that + /// supports basic compute shaders. + Vulkan, + #[cfg(any(feature = "metal", docsrs))] + #[cfg_attr(docsrs, doc(cfg(feature = "metal")))] + /// Runs kernels through the wgpu Metal backend. + /// + /// This is the only option on Apple Silicon. + Metal, + #[cfg(any(feature = "rocm", docsrs))] + #[cfg_attr(docsrs, doc(cfg(feature = "rocm")))] + /// Runs kernels through the AMD ROCm backend, on AMD GPUs only. + /// + /// ROCm is not the recommended backend for AMD GPUs. It is slower and often hit by driver issues, + /// so Vulkan is almost always faster and less buggy. + Rocm, +} + +/// Returns every accelerator this build enables, in the order to try them. +/// +/// ```no_run +/// use av_denoise::accelerate::get_default_accelerators; +/// +/// let preferred = get_default_accelerators(); +/// # let _ = preferred; +/// ``` +pub fn get_default_accelerators() -> Vec { + use strum::IntoEnumIterator; + + let mut accelerators = Vec::new(); + for enabled in Accelerator::iter() { + accelerators.push(enabled); + } + + accelerators +} diff --git a/av-denoise-core/src/device.rs b/av-denoise/src/backend/device.rs similarity index 64% rename from av-denoise-core/src/device.rs rename to av-denoise/src/backend/device.rs index f032761..c386731 100644 --- a/av-denoise-core/src/device.rs +++ b/av-denoise/src/backend/device.rs @@ -1,40 +1,29 @@ -//! Which physical device to run on. -//! -//! A [`Device`] picks the hardware inside a backend, while -//! [`Accelerator`](crate::accelerate::Accelerator) picks the backend -//! itself. The two are chosen separately because most backends expose -//! more than one device. -//! -//! [`Device`] implements `FromStr`, so it can be taken straight from a -//! command-line flag or a config file. -//! -//! ``` -//! use av_denoise_core::Device; -//! -//! // Let the backend decide. -//! assert_eq!("default".parse::().unwrap(), Device::Default); -//! -//! // Or name the second discrete GPU in the machine. -//! assert_eq!( -//! "discrete:1".parse::().unwrap(), -//! Device::Discrete { index: 1 }, -//! ); -//! ``` +//! Which physical device to run on use std::fmt; use std::str::FromStr; -/// Where to run the compute. +/// Which physical device inside a backend runs the compute. /// -/// Each variant maps onto the concrete `Device` type of whichever cubecl -/// runtime was selected. +/// The backend itself is picked separately with [Accelerator](crate::accelerate::Accelerator), +/// because most backends expose more than one device. Each variant maps onto the concrete `Device` +/// type of whichever cubecl runtime was selected. /// -/// Not every variant makes sense on every backend. `Integrated` and -/// `Virtual` are wgpu-only, and `Cpu` does nothing on the `cpu` runtime -/// while selecting `WgpuDevice::Cpu` on wgpu. +/// `Integrated`, `Virtual` and `Cpu` are wgpu-only. Asking for a variant a runtime cannot honour +/// returns an error from the matching `to_*` conversion. /// -/// Asking for a variant a runtime cannot honour returns an error from -/// the matching `to_*` conversion. +/// ``` +/// use av_denoise::Device; +/// +/// // Let the backend decide. +/// assert_eq!("default".parse::().unwrap(), Device::Default); +/// +/// // Or name the second discrete GPU in the machine. +/// assert_eq!( +/// "discrete:1".parse::().unwrap(), +/// Device::Discrete { index: 1 }, +/// ); +/// ``` #[derive(Debug, Default, Clone, PartialEq, Eq)] pub enum Device { /// Backend-chosen default device. @@ -42,56 +31,58 @@ pub enum Device { Default, /// Discrete GPU at ordinal `index`. /// - /// Maps to `CudaDevice { index }`, `AmdDevice { index }`, or - /// `WgpuDevice::DiscreteGpu(index)`. + /// Maps to `CudaDevice { index }`, `AmdDevice { index }`, or `WgpuDevice::DiscreteGpu(index)`. Discrete { index: usize }, /// Integrated GPU at ordinal `index`. wgpu-only. Integrated { index: usize }, /// Virtual GPU at ordinal `index`. wgpu-only. Virtual { index: usize }, - /// The software device. Valid on the `cpu` runtime, and on wgpu - /// where it picks the lavapipe or software adapter. + /// The software device, which picks the lavapipe or software adapter on wgpu. Cpu, } impl FromStr for Device { type Err = String; - /// Accepts the same spellings as the bench CLI. + /// Parses a device selector. /// /// - `default`, which takes no index /// - `discrete[:N]`, `integrated[:N]`, and `virtual[:N]`, where `N` defaults to 0 /// - `cpu`, which takes no index - fn from_str(s: &str) -> Result { - let (kind, suffix) = match s.split_once(':') { - Some((kind, idx)) => (kind, Some(idx)), - None => (s, None), + fn from_str(selector: &str) -> Result { + let (kind, suffix) = match selector.split_once(':') { + Some((kind, index_text)) => (kind, Some(index_text)), + None => (selector, None), }; if matches!(kind, "default" | "cpu") && suffix.is_some() { return Err(format!( - "device kind '{kind}' takes no index, got '{s}'. Only discrete, integrated, and virtual take an index" + "device kind '{kind}' takes no index, got '{selector}'. Only discrete, integrated, and virtual take an index" )); } - let parse_index = |idx: &str| -> Result { - idx.parse() - .map_err(|_| format!("invalid device index '{idx}' in '{s}'")) + let parse_index = |index_text: &str| -> Result { + index_text + .parse() + .map_err(|_| format!("invalid device index '{index_text}' in '{selector}'")) }; - let idx = suffix.unwrap_or("0"); + let index_text = suffix.unwrap_or("0"); match kind { "default" => Ok(Device::Default), "cpu" => Ok(Device::Cpu), - "discrete" => Ok(Device::Discrete { - index: parse_index(idx)?, - }), - "integrated" => Ok(Device::Integrated { - index: parse_index(idx)?, - }), - "virtual" => Ok(Device::Virtual { - index: parse_index(idx)?, - }), + "discrete" => { + let index = parse_index(index_text)?; + Ok(Device::Discrete { index }) + }, + "integrated" => { + let index = parse_index(index_text)?; + Ok(Device::Integrated { index }) + }, + "virtual" => { + let index = parse_index(index_text)?; + Ok(Device::Virtual { index }) + }, other => Err(format!( "unknown device kind '{other}', expected default, discrete[:N], integrated[:N], virtual[:N], or cpu" )), @@ -100,15 +91,14 @@ impl FromStr for Device { } impl fmt::Display for Device { - /// Writes the selector spelling [`FromStr`] accepts, so a device - /// prints as `discrete:1` rather than as its enum variant. - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + /// Writes the selector spelling [FromStr] accepts, such as `discrete:1`. + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { match self { - Device::Default => f.write_str("default"), - Device::Discrete { index } => write!(f, "discrete:{index}"), - Device::Integrated { index } => write!(f, "integrated:{index}"), - Device::Virtual { index } => write!(f, "virtual:{index}"), - Device::Cpu => f.write_str("cpu"), + Device::Default => formatter.write_str("default"), + Device::Discrete { index } => write!(formatter, "discrete:{index}"), + Device::Integrated { index } => write!(formatter, "integrated:{index}"), + Device::Virtual { index } => write!(formatter, "virtual:{index}"), + Device::Cpu => formatter.write_str("cpu"), } } } @@ -143,6 +133,7 @@ impl Device { impl Device { pub fn to_wgpu(&self) -> Result { use cubecl::wgpu::WgpuDevice; + Ok(match self { Device::Default => WgpuDevice::DefaultDevice, Device::Discrete { index } => WgpuDevice::DiscreteGpu(*index), @@ -206,10 +197,11 @@ mod tests { #[test] fn rejected_index_error_names_the_kind() { - let err = "default:1".parse::().unwrap_err(); - assert!(err.contains("default"), "{err}"); - let err = "cpu:2".parse::().unwrap_err(); - assert!(err.contains("cpu"), "{err}"); + let default_error = "default:1".parse::().unwrap_err(); + assert!(default_error.contains("default"), "{default_error}"); + + let cpu_error = "cpu:2".parse::().unwrap_err(); + assert!(cpu_error.contains("cpu"), "{cpu_error}"); } #[test] diff --git a/av-denoise-core/src/enumerate.rs b/av-denoise/src/backend/enumerate.rs similarity index 63% rename from av-denoise-core/src/enumerate.rs rename to av-denoise/src/backend/enumerate.rs index 2a17138..9d2af0d 100644 --- a/av-denoise-core/src/enumerate.rs +++ b/av-denoise/src/backend/enumerate.rs @@ -1,9 +1,11 @@ +//! Listing the devices each backend can see + use cubecl::device::DeviceId; use cubecl::prelude::*; -use crate::accelerate::Accelerator; -use crate::device::Device; -use crate::probe::open_client; +use super::accelerate::Accelerator; +use super::device::Device; +use super::probe::open_client; /// What one backend reports about this machine. #[derive(Debug, Clone, PartialEq, Eq)] @@ -12,9 +14,8 @@ pub struct BackendDevices { pub accelerator: Accelerator, /// Whether the backend started at all. /// - /// A build can enable a backend the machine has no driver for, and - /// that backend reports no devices because it never ran, not - /// because the machine has no hardware. + /// A backend the machine has no driver for reports no devices because it never ran, not because + /// the machine has no hardware. pub available: bool, /// The devices the backend can see, in the order it lists them. /// @@ -24,36 +25,34 @@ pub struct BackendDevices { /// Asks each backend in `enable` which devices it can see. /// -/// Backends are reported in the order given, including the ones that -/// could not start, so a caller can tell "no such hardware" apart from -/// "no such driver". +/// Backends are reported in the order given, including the ones that could not start, so a caller +/// can tell "no such hardware" apart from "no such driver". pub fn enumerate_devices(enable: &[Accelerator]) -> Vec { enable .iter() .map(|accelerator| match accelerator { #[cfg(feature = "cuda")] Accelerator::Cuda => match Device::Default.to_cuda() { - Ok(dev) => query_runtime::(*accelerator, &dev), + Ok(cuda_device) => query_runtime::(*accelerator, &cuda_device), Err(_) => unavailable(*accelerator), }, #[cfg(feature = "rocm")] Accelerator::Rocm => match Device::Default.to_amd() { - Ok(dev) => query_runtime::(*accelerator, &dev), + Ok(amd_device) => query_runtime::(*accelerator, &amd_device), Err(_) => unavailable(*accelerator), }, #[cfg(feature = "vulkan")] Accelerator::Vulkan => match Device::Default.to_wgpu() { - Ok(dev) => query_runtime::(*accelerator, &dev), + Ok(wgpu_device) => query_runtime::(*accelerator, &wgpu_device), Err(_) => unavailable(*accelerator), }, #[cfg(feature = "metal")] Accelerator::Metal => match Device::Default.to_wgpu() { - Ok(dev) => query_runtime::(*accelerator, &dev), + Ok(wgpu_device) => query_runtime::(*accelerator, &wgpu_device), Err(_) => unavailable(*accelerator), }, - // Keeps the match exhaustive on docs.rs, where `cfg(docsrs)` - // widens the `Accelerator` enum to include variants whose - // backend feature is not enabled. Never reached at runtime. + // Keeps the match exhaustive on docs.rs, where `cfg(docsrs)` widens the `Accelerator` enum to + // include variants whose backend feature is not enabled. Never reached at runtime. #[cfg(docsrs)] #[expect( unreachable_patterns, @@ -74,24 +73,20 @@ fn unavailable(accelerator: Accelerator) -> BackendDevices { /// Opens a client on `device` and lists what that backend can see. /// -/// A backend that cannot open a client at all, because its driver -/// libraries are missing, is reported as unavailable rather than -/// allowed to take the process down. See [`probe`](crate::probe). +/// A backend that cannot open a client at all, because its driver libraries are missing, is reported +/// as unavailable rather than allowed to take the process down. fn query_runtime(accelerator: Accelerator, device: &R::Device) -> BackendDevices { let Some(client) = open_client::(accelerator, device) else { return unavailable(accelerator); }; - // Type ids 0 to 3 are the device kinds `Device` can name. Anything - // else the runtime reports is hardware this tool cannot select. - // - // Not every backend filters by the type id it is given. ROCm and - // CUDA report their whole device list for each one, so the same - // device comes back on every pass and is kept only once. + // Type ids 0 to 3 are the device kinds `Device` can name, and anything else is hardware this tool + // cannot select. ROCm and CUDA report their whole device list for every type id, so a device that + // comes back on several passes is kept only once. let mut devices: Vec = Vec::new(); for type_id in 0..=3 { - for id in client.enumerate_devices(type_id) { - if let Some(device) = to_device(id) + for device_id in client.enumerate_devices(type_id) { + if let Some(device) = to_device(device_id) && !devices.contains(&device) { devices.push(device); @@ -108,8 +103,8 @@ fn query_runtime(accelerator: Accelerator, device: &R::Device) -> Ba /// Maps a cubecl device id onto the selector that names it. /// -/// The type ids come from cubecl's own ordering of device kinds. -/// Backends that report a kind this tool cannot select return `None`. +/// The type ids come from cubecl's own ordering of device kinds. A kind this tool cannot select +/// returns `None`. fn to_device(id: DeviceId) -> Option { let index = id.index_id as usize; match id.type_id { @@ -127,26 +122,35 @@ mod tests { #[test] fn device_kinds_map_from_type_ids() { - assert_eq!( - to_device(DeviceId::new(0, 1)), - Some(Device::Discrete { index: 1 }), - ); - assert_eq!( - to_device(DeviceId::new(1, 0)), - Some(Device::Integrated { index: 0 }), - ); - assert_eq!(to_device(DeviceId::new(2, 2)), Some(Device::Virtual { index: 2 }),); - assert_eq!(to_device(DeviceId::new(3, 0)), Some(Device::Cpu)); + let discrete_id = DeviceId::new(0, 1); + let integrated_id = DeviceId::new(1, 0); + let virtual_id = DeviceId::new(2, 2); + let cpu_id = DeviceId::new(3, 0); + + let discrete = to_device(discrete_id); + let integrated = to_device(integrated_id); + let virtual_gpu = to_device(virtual_id); + let cpu = to_device(cpu_id); + + assert_eq!(discrete, Some(Device::Discrete { index: 1 })); + assert_eq!(integrated, Some(Device::Integrated { index: 0 })); + assert_eq!(virtual_gpu, Some(Device::Virtual { index: 2 })); + assert_eq!(cpu, Some(Device::Cpu)); } #[test] fn unknown_type_ids_are_skipped() { - assert_eq!(to_device(DeviceId::new(4, 0)), None); + let unknown_id = DeviceId::new(4, 0); + let unknown = to_device(unknown_id); + + assert_eq!(unknown, None); } #[test] fn no_backends_lists_nothing() { - assert!(enumerate_devices(&[]).is_empty()); + let reported = enumerate_devices(&[]); + + assert!(reported.is_empty()); } #[cfg(feature = "vulkan")] diff --git a/av-denoise/src/backend/mod.rs b/av-denoise/src/backend/mod.rs new file mode 100644 index 0000000..59df825 --- /dev/null +++ b/av-denoise/src/backend/mod.rs @@ -0,0 +1,105 @@ +pub mod accelerate; +pub mod device; +pub mod enumerate; +mod probe; +pub mod sniff; + +use av_denoise_core::{Engine, Geometry, Nl4d, Nl4dOptions, Nlmeans, NlmeansAlgorithm}; +use cubecl::Runtime; +use cubecl::prelude::ComputeClient; + +use self::accelerate::Accelerator; +pub use self::device::Device; +use self::sniff::sniff_best_accelerator; +use crate::host::DenoiserError; +use crate::host::io::{ClientIo, PlaneIo}; + +/// Which engine to build, and for what planes. +#[derive(Debug, Clone, Copy)] +pub enum EngineSpec { + Nlmeans { + algorithm: NlmeansAlgorithm, + geometry: Geometry, + }, + Nl4d { + options: Nl4dOptions, + geometry: Geometry, + }, +} + +/// An engine and the plane I/O for the runtime it was built on. +pub(crate) struct BuiltEngine { + pub engine: Box, + pub io: Box, + pub accelerator: Accelerator, +} + +fn build_on( + client: ComputeClient, + spec: EngineSpec, + accelerator: Accelerator, +) -> Result { + let engine: Box = match spec { + EngineSpec::Nlmeans { algorithm, geometry } => { + let engine = Nlmeans::new(&client, algorithm, geometry)?; + Box::new(engine) + }, + EngineSpec::Nl4d { options, geometry } => { + let engine = Nl4d::new(&client, options, geometry)?; + Box::new(engine) + }, + }; + + let io = ClientIo::new(client); + + Ok(BuiltEngine { + engine, + io: Box::new(io), + accelerator, + }) +} + +/// Builds `spec` on the first accelerator in `accelerators` that works. +pub(crate) fn build_engine( + accelerators: &[Accelerator], + device: &Device, + spec: EngineSpec, +) -> Result { + let accelerator = sniff_best_accelerator(accelerators, device); + let accelerator = accelerator.ok_or(DenoiserError::NoAcceleratorAvailable)?; + + match accelerator { + #[cfg(feature = "cuda")] + Accelerator::Cuda => { + let cuda_device = device.to_cuda()?; + let client = ::client(&cuda_device); + build_on(client, spec, accelerator) + }, + #[cfg(feature = "rocm")] + Accelerator::Rocm => { + let amd_device = device.to_amd()?; + let client = ::client(&amd_device); + build_on(client, spec, accelerator) + }, + #[cfg(feature = "vulkan")] + Accelerator::Vulkan => { + let wgpu_device = device.to_wgpu()?; + let client = ::client(&wgpu_device); + build_on(client, spec, accelerator) + }, + #[cfg(feature = "metal")] + Accelerator::Metal => { + let wgpu_device = device.to_wgpu()?; + let client = ::client(&wgpu_device); + build_on(client, spec, accelerator) + }, + // Keeps the match exhaustive on docs.rs, where `cfg(docsrs)` widens the `Accelerator` enum to + // include variants whose backend feature is not enabled. Never reached at runtime. + #[cfg(docsrs)] + #[expect( + unreachable_patterns, + reason = "the arm only keeps the match exhaustive on docs.rs" + )] + _ => unreachable!(), + } +} diff --git a/av-denoise/src/backend/probe.rs b/av-denoise/src/backend/probe.rs new file mode 100644 index 0000000..0aa79b1 --- /dev/null +++ b/av-denoise/src/backend/probe.rs @@ -0,0 +1,116 @@ +use std::panic::{self, AssertUnwindSafe, PanicHookInfo}; +use std::sync::Mutex; + +use cubecl::client::ComputeClient; +use cubecl::prelude::*; + +use super::accelerate::Accelerator; + +/// Backends already reported as unavailable, and the lock guarding the panic hook. +/// +/// The hook is process-wide, so two probes running at once would race to restore each other's. +/// Holding this for the length of a probe keeps them in single file, and the list inside it stops a +/// backend from warning again every time it is probed. +static PROBED: Mutex> = Mutex::new(Vec::new()); + +/// Opens a client for `accelerator` on `device`, or reports that the backend cannot run here. +/// +/// The client is synchronised before it is handed back. cubecl kernels are fully asynchronous, so a +/// successful `sync()` proves the backend works and no test kernel is needed. +/// +/// Some backends report missing driver libraries by panicking. The CUDA runtime loads `libcuda` on +/// its own worker thread and panics there, which reaches the caller as a second panic when cubecl +/// unwraps the dead worker's channel. Catching that lets one binary with several backends enabled +/// run on a machine that has only one of them. +pub(crate) fn open_client( + accelerator: Accelerator, + device: &R::Device, +) -> Option> { + let mut probed = PROBED.lock().unwrap_or_else(|err| err.into_inner()); + + let opened = quiet_panics(|| { + let client = R::client(device); + let synced = client.sync(); + cubecl::future::block_on(synced).map(|()| client) + }); + + match opened { + Ok(Ok(client)) => Some(client), + Ok(Err(err)) => { + tracing::debug!(err = ?err, "could not use the {accelerator} runtime"); + None + }, + Err(_) => { + // A denoise run probes once per denoiser it builds, and a missing driver is worth one + // line rather than one per scene. + if !probed.contains(&accelerator) { + probed.push(accelerator); + tracing::warn!( + "the {accelerator} backend is enabled but did not start, its driver libraries are probably missing" + ); + } + None + }, + } +} + +/// Runs `work`, turning a panic into an `Err` and routing the panic message to the debug log. +/// +/// A failing backend prints its own panic from its worker thread before the caller sees one, so the +/// hook is quietened while `work` runs and put back afterwards. Callers hold [PROBED] across this so +/// two probes cannot race to restore each other's hook. +fn quiet_panics(work: impl FnOnce() -> T) -> std::thread::Result { + let previous = panic::take_hook(); + let quiet_hook = Box::new(|info: &PanicHookInfo<'_>| tracing::debug!("{info}")); + panic::set_hook(quiet_hook); + + let unwind_safe_work = AssertUnwindSafe(work); + let result = panic::catch_unwind(unwind_safe_work); + panic::set_hook(previous); + result +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::sync::atomic::{AtomicBool, Ordering}; + + use super::*; + + #[test] + fn a_panic_inside_becomes_an_error() { + let _probed = PROBED.lock().unwrap_or_else(|err| err.into_inner()); + + let panicked = quiet_panics(|| panic!("the backend fell over")); + let returned = quiet_panics(|| 7); + + assert!(panicked.is_err()); + assert_eq!(returned.unwrap(), 7); + } + + /// The hook has to come back however the probe ended, or every later panic reports at debug level. + /// + /// Holding [PROBED], as [open_client] does, keeps a probe on another thread from swapping the hook + /// mid-test. + #[test] + fn the_panic_hook_is_restored() { + let _probed = PROBED.lock().unwrap_or_else(|err| err.into_inner()); + + let marker = Arc::new(AtomicBool::new(false)); + let flag = marker.clone(); + let marking_hook = Box::new(move |_: &PanicHookInfo<'_>| flag.store(true, Ordering::SeqCst)); + panic::set_hook(marking_hook); + + let _ = quiet_panics(|| panic!("swallowed by the quiet hook")); + assert!( + !marker.load(Ordering::SeqCst), + "the quiet hook did not replace the installed one", + ); + + let _ = panic::catch_unwind(|| panic!("seen by the restored hook")); + let restored = marker.load(Ordering::SeqCst); + let _ = panic::take_hook(); + + assert!(restored, "the probe left its own panic hook installed"); + } +} diff --git a/av-denoise/src/backend/sniff.rs b/av-denoise/src/backend/sniff.rs new file mode 100644 index 0000000..d83eb7f --- /dev/null +++ b/av-denoise/src/backend/sniff.rs @@ -0,0 +1,75 @@ +//! Picking a backend that actually works on this machine + +use cubecl::prelude::*; + +use super::accelerate::Accelerator; +use super::device::Device; +use super::probe::open_client; + +/// Returns the first accelerator, in order of preference, whose client can be built and synchronised +/// on `device`. +/// +/// cubecl kernels are fully asynchronous, so a successful `client.sync()` is enough to prove the +/// backend works and no test kernel is needed. The probe runs on `device` rather than the backend's +/// default, because opening a client on another card tests the wrong hardware and pays that card's +/// first-time driver initialisation. +/// +/// An accelerator that cannot express `device` at all is treated as unavailable, which is the answer +/// building on it would give one step earlier. So is a backend whose driver libraries are missing, +/// however loudly it fails. +/// +/// ```no_run +/// use av_denoise::Device; +/// use av_denoise::accelerate::get_default_accelerators; +/// use av_denoise::sniff::sniff_best_accelerator; +/// +/// let preferred = get_default_accelerators(); +/// match sniff_best_accelerator(&preferred, &Device::Default) { +/// Some(accelerator) => println!("running on {accelerator}"), +/// None => println!("no usable backend on this machine"), +/// } +/// ``` +pub fn sniff_best_accelerator(enable: &[Accelerator], device: &Device) -> Option { + for accelerator in enable { + let is_enabled = match accelerator { + #[cfg(feature = "cuda")] + Accelerator::Cuda => match device.to_cuda() { + Ok(cuda_device) => probe_runtime::(*accelerator, &cuda_device), + Err(_) => false, + }, + #[cfg(feature = "rocm")] + Accelerator::Rocm => match device.to_amd() { + Ok(amd_device) => probe_runtime::(*accelerator, &amd_device), + Err(_) => false, + }, + #[cfg(feature = "vulkan")] + Accelerator::Vulkan => match device.to_wgpu() { + Ok(wgpu_device) => probe_runtime::(*accelerator, &wgpu_device), + Err(_) => false, + }, + #[cfg(feature = "metal")] + Accelerator::Metal => match device.to_wgpu() { + Ok(wgpu_device) => probe_runtime::(*accelerator, &wgpu_device), + Err(_) => false, + }, + // Keeps the match exhaustive on docs.rs, where `cfg(docsrs)` widens the `Accelerator` enum to + // include variants whose backend feature is not enabled. Never reached at runtime. + #[cfg(docsrs)] + #[expect( + unreachable_patterns, + reason = "the arm only keeps the match exhaustive on docs.rs" + )] + _ => unreachable!(), + }; + + if is_enabled { + return Some(*accelerator); + } + } + + None +} + +fn probe_runtime(accelerator: Accelerator, device: &R::Device) -> bool { + open_client::(accelerator, device).is_some() +} diff --git a/av-denoise/src/bin/cli/common.rs b/av-denoise/src/bin/cli/common.rs index 80b0112..bafad80 100644 --- a/av-denoise/src/bin/cli/common.rs +++ b/av-denoise/src/bin/cli/common.rs @@ -1,7 +1,6 @@ use super::InputSource; -/// The flags every denoising family takes, whatever it does with the -/// frames once it has them. +/// Flags shared by every denoising subcommand. #[derive(Debug, Clone, clap::Args)] pub struct CommonArgs { /// Where to read frames from. @@ -21,7 +20,7 @@ pub struct CommonArgs { /// /// The source's bit depth is detected automatically. 8, 10, and /// 12-bit sources are supported and the output keeps the source's - /// depth. Other depths are rejected with a clear error message. + /// depth. Other depths are rejected with an error. #[arg(short, long)] pub input: InputSource, @@ -30,8 +29,7 @@ pub struct CommonArgs { /// Each worker uses its own GPU memory for the frame ring /// buffer, so higher values trade GPU memory for throughput. /// - /// `1` is valid and useful for debugging. Defaults to 2 when - /// unset. + /// `1` is valid and useful for debugging. Defaults to 2 when unset. #[arg(short = 'W', long)] pub workers: Option, diff --git a/av-denoise/src/bin/cli/input.rs b/av-denoise/src/bin/cli/input.rs index c9dd718..1c8774b 100644 --- a/av-denoise/src/bin/cli/input.rs +++ b/av-denoise/src/bin/cli/input.rs @@ -18,11 +18,13 @@ pub enum InputSource { impl InputSource { /// Opens the stream this source names. /// - /// Only the piped variants are readable here. A path is opened - /// with ffms2 instead. + /// Only the piped variants are readable as a stream. A path is an error. pub fn open_reader(&self) -> Result, anyhow::Error> { match self { - InputSource::Stdin => Ok(Box::new(std::io::stdin().lock())), + InputSource::Stdin => { + let stdin = std::io::stdin().lock(); + Ok(Box::new(stdin)) + }, InputSource::Fd(fd) => open_fd(*fd), InputSource::File(path) => anyhow::bail!( "`{}` is a file path and is opened with ffms2, not read as a stream", @@ -40,38 +42,39 @@ impl FromStr for InputSource { /// - `-` and `pipe:0` are standard input /// - `pipe:N` for `N` of 3 or above is an inherited descriptor /// - anything else is a path on disk - fn from_str(s: &str) -> Result { - if s == "-" { + fn from_str(raw: &str) -> Result { + if raw == "-" { return Ok(InputSource::Stdin); } - if let Some(rest) = s.strip_prefix("pipe:") { + if let Some(rest) = raw.strip_prefix("pipe:") { let fd: u32 = rest .parse() - .map_err(|_| format!("pipe: expects a file descriptor number (got `{s}`)"))?; + .map_err(|_| format!("pipe: expects a file descriptor number (got `{raw}`)"))?; return match fd { 0 => Ok(InputSource::Stdin), 1 => Err("pipe:1 is this process's stdout, which carries the denoised y4m".to_string()), 2 => Err("pipe:2 is this process's stderr, which carries log output".to_string()), - n => Ok(InputSource::Fd(n)), + inherited => Ok(InputSource::Fd(inherited)), }; } - if s.is_empty() { + if raw.is_empty() { return Err("expected a file path, `-`, or `pipe:N`".to_string()); } - Ok(InputSource::File(PathBuf::from(s))) + let path = PathBuf::from(raw); + Ok(InputSource::File(path)) } } impl fmt::Display for InputSource { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { match self { - InputSource::Stdin => f.write_str("stdin"), - InputSource::Fd(fd) => write!(f, "pipe:{fd}"), - InputSource::File(path) => write!(f, "{}", path.display()), + InputSource::Stdin => formatter.write_str("stdin"), + InputSource::Fd(fd) => write!(formatter, "pipe:{fd}"), + InputSource::File(path) => write!(formatter, "{}", path.display()), } } } @@ -81,7 +84,7 @@ impl fmt::Display for InputSource { fn open_fd(fd: u32) -> Result, anyhow::Error> { let path = format!("/dev/fd/{fd}"); let file = std::fs::File::open(&path) - .map_err(|e| anyhow::anyhow!("--input pipe:{fd} could not open {path}: {e}"))?; + .map_err(|error| anyhow::anyhow!("--input pipe:{fd} could not open {path}: {error}"))?; Ok(Box::new(file)) } @@ -99,70 +102,86 @@ mod tests { use super::*; - fn parse(s: &str) -> Result { - s.parse() + fn parse(raw: &str) -> Result { + raw.parse() } #[test] fn dash_is_stdin() { - assert_eq!(parse("-").unwrap(), InputSource::Stdin); + let source = parse("-").unwrap(); + + assert_eq!(source, InputSource::Stdin); } #[test] fn pipe_zero_is_stdin() { - assert_eq!(parse("pipe:0").unwrap(), InputSource::Stdin); + let source = parse("pipe:0").unwrap(); + + assert_eq!(source, InputSource::Stdin); } #[test] fn pipe_three_is_an_inherited_descriptor() { - assert_eq!(parse("pipe:3").unwrap(), InputSource::Fd(3)); + let source = parse("pipe:3").unwrap(); + + assert_eq!(source, InputSource::Fd(3)); } #[test] fn our_own_output_descriptors_are_rejected() { - assert!(parse("pipe:1").unwrap_err().contains("stdout")); - assert!(parse("pipe:2").unwrap_err().contains("stderr")); + let stdout_error = parse("pipe:1").unwrap_err(); + let stderr_error = parse("pipe:2").unwrap_err(); + + assert!(stdout_error.contains("stdout")); + assert!(stderr_error.contains("stderr")); } #[test] fn non_numeric_descriptor_is_rejected() { - assert!(parse("pipe:x").unwrap_err().contains("file descriptor")); + let error = parse("pipe:x").unwrap_err(); + + assert!(error.contains("file descriptor")); } #[test] fn anything_else_is_a_path() { - assert_eq!( - parse("noisy.mkv").unwrap(), - InputSource::File(PathBuf::from("noisy.mkv")) - ); - assert_eq!(parse("./-").unwrap(), InputSource::File(PathBuf::from("./-"))); - assert_eq!( - parse("./pipe:3").unwrap(), - InputSource::File(PathBuf::from("./pipe:3")) - ); + let plain = parse("noisy.mkv").unwrap(); + let dash = parse("./-").unwrap(); + let pipe_lookalike = parse("./pipe:3").unwrap(); + + let plain_path = PathBuf::from("noisy.mkv"); + let dash_path = PathBuf::from("./-"); + let pipe_lookalike_path = PathBuf::from("./pipe:3"); + + assert_eq!(plain, InputSource::File(plain_path)); + assert_eq!(dash, InputSource::File(dash_path)); + assert_eq!(pipe_lookalike, InputSource::File(pipe_lookalike_path)); } #[test] fn empty_is_rejected() { - assert!(parse("").is_err()); + let result = parse(""); + + assert!(result.is_err()); } #[test] fn display_round_trips_the_typed_spelling() { + let path = PathBuf::from("noisy.mkv"); + let file = InputSource::File(path); + assert_eq!(InputSource::Stdin.to_string(), "stdin"); assert_eq!(InputSource::Fd(3).to_string(), "pipe:3"); - assert_eq!( - InputSource::File(PathBuf::from("noisy.mkv")).to_string(), - "noisy.mkv" - ); + assert_eq!(file.to_string(), "noisy.mkv"); } /// `/dev/fd/N` reopens whatever the descriptor points at, so a temp - /// file stands in for the inherited pipe a harness would hand us. + /// file stands in for an inherited pipe. #[cfg(unix)] #[test] fn open_reader_reads_an_inherited_descriptor() { - let path = std::env::temp_dir().join(format!("av-denoise-fd-{}.bin", std::process::id())); + let file_name = format!("av-denoise-fd-{}.bin", std::process::id()); + let path = std::env::temp_dir().join(file_name); let mut file = std::fs::File::create(&path).expect("temp file should create"); file.write_all(b"YUV4MPEG2 frames").expect("payload should write"); @@ -176,24 +195,29 @@ mod tests { .open_reader() .expect("inherited descriptor should open"); - let mut got = String::new(); - reader.read_to_string(&mut got).expect("payload should read back"); + let mut read_back = String::new(); + reader + .read_to_string(&mut read_back) + .expect("payload should read back"); drop(reader); drop(file); let _ = std::fs::remove_file(&path); - assert_eq!(got, "YUV4MPEG2 frames"); + assert_eq!(read_back, "YUV4MPEG2 frames"); } #[test] fn open_reader_rejects_a_path() { + let path = PathBuf::from("noisy.mkv"); + let source = InputSource::File(path); + // `Box` isn't `Debug`, so `expect_err` can't be used here. - let err = match InputSource::File(PathBuf::from("noisy.mkv")).open_reader() { + let error = match source.open_reader() { Ok(_) => panic!("paths are opened with ffms2, not read as a stream"), - Err(e) => e, + Err(error) => error, }; - assert!(err.to_string().contains("noisy.mkv")); + assert!(error.to_string().contains("noisy.mkv")); } } diff --git a/av-denoise/src/bin/cli/list_devices.rs b/av-denoise/src/bin/cli/list_devices.rs index 27ee3fd..705f576 100644 --- a/av-denoise/src/bin/cli/list_devices.rs +++ b/av-denoise/src/bin/cli/list_devices.rs @@ -2,30 +2,27 @@ use av_denoise::Device; use av_denoise::accelerate::Accelerator; use av_denoise::enumerate::{BackendDevices, enumerate_devices}; -/// Lists the devices every backend in `enable` can see. +/// Lists the devices every backend in `enable` can see, one row per device. /// -/// One row per device, naming the backends that offer it, so the -/// selector printed in the first column can be passed straight to -/// `--device`. +/// The first column of each row can be passed straight to `--device`. pub fn run_list_devices(enable: &[Accelerator]) -> String { - format_table(&enumerate_devices(enable)) + let reported = enumerate_devices(enable); + format_table(&reported) } /// Renders what the backends reported as a table. /// -/// Kept apart from the enumeration so it can be tested without a GPU. -/// -/// The table opens with a blank line. Graphics drivers write their own -/// notices to stderr while the backends start, and the header is hard to -/// pick out of that without a gap in front of it. +/// It is separate from the enumeration so it can be tested without a GPU. The table opens +/// with a blank line because graphics drivers write their own notices to stderr while the +/// backends start, and the header is hard to pick out without a gap. fn format_table(reported: &[BackendDevices]) -> String { let rows = build_rows(reported); - let unavailable: Vec<&BackendDevices> = reported.iter().filter(|b| !b.available).collect(); + let unavailable: Vec<&BackendDevices> = reported.iter().filter(|backend| !backend.available).collect(); - let mut out = String::from("\n"); + let mut table = String::from("\n"); if rows.is_empty() { - out.push_str("No usable devices found.\n"); + table.push_str("No usable devices found.\n"); } else { let width = rows .iter() @@ -34,21 +31,27 @@ fn format_table(reported: &[BackendDevices]) -> String { .max() .unwrap_or(0); - out.push_str(&format!("{: = unavailable.iter().map(|b| b.accelerator.to_string()).collect(); - out.push_str(&format!( - "\nEnabled but unavailable on this machine: {}\n", - names.join(", "), - )); + let names: Vec = unavailable + .iter() + .map(|backend| backend.accelerator.to_string()) + .collect(); + let name_list = names.join(", "); + let footer = format!("\nEnabled but unavailable on this machine: {name_list}\n"); + table.push_str(&footer); } - out + table } /// Turns per-backend device lists into per-device backend lists. @@ -59,7 +62,7 @@ fn format_table(reported: &[BackendDevices]) -> String { /// works on a backend that runs at all. /// /// The ordinals are per-backend, so two backends sharing a row do not -/// necessarily share a physical card. +/// always share a physical card. fn build_rows(reported: &[BackendDevices]) -> Vec<(String, Vec)> { let mut rows: Vec<(Device, Vec)> = Vec::new(); @@ -72,14 +75,17 @@ fn build_rows(reported: &[BackendDevices]) -> Vec<(String, Vec)> { } }; - for backend in reported.iter().filter(|b| b.available) { + let available = reported.iter().filter(|backend| backend.available); + for backend in available { record(Device::Default, backend.accelerator); + for device in &backend.devices { record(device.clone(), backend.accelerator); } } rows.sort_by_key(|(device, _)| sort_key(device)); + rows.into_iter() .map(|(device, backends)| (device.to_string(), backends)) .collect() @@ -116,12 +122,12 @@ mod tests { /// is named once per row however often it reports that device. #[test] fn rows_merge_backends_offering_the_same_device() { - let reported = vec![ - backend(Accelerator::Vulkan, vec![Device::Discrete { index: 0 }]), - backend(Accelerator::Vulkan, vec![Device::Discrete { index: 0 }]), - ]; + let first = backend(Accelerator::Vulkan, vec![Device::Discrete { index: 0 }]); + let second = backend(Accelerator::Vulkan, vec![Device::Discrete { index: 0 }]); + let reported = vec![first, second]; let rows = build_rows(&reported); + assert_eq!( rows, vec![ @@ -133,17 +139,18 @@ mod tests { #[test] fn rows_are_sorted_by_kind_then_ordinal() { - let reported = vec![backend( - Accelerator::Vulkan, - vec![ - Device::Cpu, - Device::Discrete { index: 1 }, - Device::Integrated { index: 0 }, - Device::Discrete { index: 0 }, - ], - )]; + let devices = vec![ + Device::Cpu, + Device::Discrete { index: 1 }, + Device::Integrated { index: 0 }, + Device::Discrete { index: 0 }, + ]; + let vulkan = backend(Accelerator::Vulkan, devices); + let reported = vec![vulkan]; + + let rows = build_rows(&reported); + let names: Vec = rows.into_iter().map(|(device, _)| device).collect(); - let names: Vec = build_rows(&reported).into_iter().map(|(d, _)| d).collect(); assert_eq!( names, ["default", "discrete:0", "discrete:1", "integrated:0", "cpu"], @@ -158,28 +165,33 @@ mod tests { devices: Vec::new(), }]; - assert!(build_rows(&reported).is_empty()); + let rows = build_rows(&reported); + + assert!(rows.is_empty()); } #[test] fn the_table_opens_with_a_blank_line() { - let reported = vec![backend(Accelerator::Vulkan, vec![Device::Cpu])]; + let vulkan = backend(Accelerator::Vulkan, vec![Device::Cpu]); + let reported = vec![vulkan]; + + let table = format_table(&reported); assert!( - format_table(&reported).starts_with('\n'), + table.starts_with('\n'), "driver notices on stderr run straight into the header", ); } #[test] fn table_pads_the_device_column() { - let reported = vec![backend( - Accelerator::Vulkan, - vec![Device::Integrated { index: 0 }], - )]; + let vulkan = backend(Accelerator::Vulkan, vec![Device::Integrated { index: 0 }]); + let reported = vec![vulkan]; + + let table = format_table(&reported); assert_eq!( - format_table(&reported), + table, "\nDEVICE BACKENDS\ndefault vulkan\nintegrated:0 vulkan\n", ); } @@ -192,14 +204,18 @@ mod tests { devices: Vec::new(), }]; + let table = format_table(&reported); + assert_eq!( - format_table(&reported), + table, "\nNo usable devices found.\n\nEnabled but unavailable on this machine: vulkan\n", ); } #[test] fn nothing_enabled_says_so() { - assert_eq!(format_table(&[]), "\nNo usable devices found.\n"); + let table = format_table(&[]); + + assert_eq!(table, "\nNo usable devices found.\n"); } } diff --git a/av-denoise/src/bin/cli/mod.rs b/av-denoise/src/bin/cli/mod.rs index f5d7714..da9dfcf 100644 --- a/av-denoise/src/bin/cli/mod.rs +++ b/av-denoise/src/bin/cli/mod.rs @@ -19,15 +19,11 @@ pub use self::motion::MotionArgs; pub use self::nl4d::Nl4dArgs; pub use self::nlmeans::NlmeansArgs; -/// The options `main` runs a denoising pass with. -/// -/// `planes` is the library-side option set that drives `PlanarDenoiser`. -/// `progress` is a CLI concern the library layer has no use for, so it -/// stays here rather than on `PlaneOptions`. +/// The options a denoising run takes. #[derive(Debug, Clone)] pub struct RunOptions { pub planes: PlaneOptions, - /// Draws the denoising progress bar. + /// Whether to draw the denoising progress bar. pub progress: bool, /// Where to write the AV1 film grain table, when one was asked for. pub grain_table: Option, @@ -45,6 +41,7 @@ pub enum CliChannelMode { Yuv, } +/// Turns the `--channel-mode` list into the planes to denoise. pub fn resolve_channel_intent(modes: &[CliChannelMode]) -> Result { if modes.is_empty() { anyhow::bail!("--channel-mode must contain at least one value"); @@ -57,21 +54,26 @@ pub fn resolve_channel_intent(modes: &[CliChannelMode]) -> Result 1 || chroma_count > 1 || yuv_count > 1 { anyhow::bail!("--channel-mode entries must be unique"); } - Ok(match (has_yuv, has_luma, has_chroma) { + let intent = match (has_yuv, has_luma, has_chroma) { (true, _, _) => ChannelIntent::YuvFused, (false, true, true) => ChannelIntent::LumaChroma, (false, true, false) => ChannelIntent::Luma, (false, false, true) => ChannelIntent::Chroma, (false, false, false) => unreachable!("empty list rejected above"), - }) + }; + + Ok(intent) } #[derive(Debug, Parser)] @@ -80,10 +82,9 @@ pub struct Args { /// Speed vs quality dial. /// /// `veryfast` is the fastest and lowest-quality end of the dial. - /// For `nlmeans` it runs the `fast` variant with no temporal window - /// and matches this tool's original default behavior. For `nl4d` it - /// keeps a 1-frame temporal window, which that algorithm needs, and - /// narrows the spatial search instead. + /// For `nlmeans` it runs the `fast` variant with no temporal window. + /// For `nl4d` it keeps a 1-frame temporal window, which that + /// algorithm needs, and narrows the spatial search instead. /// /// Going up the list widens the temporal window, from a 1-frame /// radius at `fast` to an 8-frame radius at `veryslow`. `fast` and @@ -174,16 +175,15 @@ pub enum Command { /// temporal window. Nlmeans(NlmeansArgs), - /// Denoise by grouping matching patches across several noisy frames - /// directly, rather than filtering with non-local means first. + /// Denoise by grouping matching patches across several noisy frames. /// - /// `nl4d` measures the noise level, tracks motion, and scores how - /// well each neighbour frame matches, the same way `nlmeans hq` - /// does. No NLM weighting pass ever runs. Instead, patches are - /// grouped straight out of the noisy frames, searching both the - /// centre frame spatially and each neighbour frame around where a - /// patch is predicted to have moved, then each group's coefficients - /// are shrunk jointly. + /// Patches are grouped straight out of the noisy frames, and no + /// non-local means pass runs. `nl4d` measures the noise level, + /// tracks motion, and scores how well each neighbour frame matches, + /// the same way `nlmeans hq` does. Each group searches the centre + /// frame spatially and each neighbour frame around where a patch is + /// predicted to have moved, then the group's coefficients are shrunk + /// jointly. /// /// Motion tracking is always on, and every preset keeps a temporal /// window, which this algorithm needs. diff --git a/av-denoise/src/bin/cli/motion.rs b/av-denoise/src/bin/cli/motion.rs index e690ace..470ca87 100644 --- a/av-denoise/src/bin/cli/motion.rs +++ b/av-denoise/src/bin/cli/motion.rs @@ -1,15 +1,6 @@ use av_denoise::MotionSearch; -/// How motion between frames is tracked, for the families that track it. -/// -/// When the camera or content moves between frames, the brightness at -/// the same `(x, y)` is different content in each frame. Motion tracking -/// looks at where each block of pixels moved, so a temporal pass lines -/// neighbouring frames up instead of blurring moving edges. -/// -/// Whether tracking runs at all is the family's own business. -/// [`super::NlmeansArgs`] takes a `--motion-compensation` switch for it. -/// [`super::Nl4dArgs`] always tracks motion. +/// Motion-search flags for the subcommands that track motion between frames. #[derive(Debug, Clone, Default, clap::Args)] pub struct MotionArgs { /// Size of each motion-search block, in pixels. Must be even. @@ -34,9 +25,7 @@ pub struct MotionArgs { /// /// The coarse pyramid pass reaches further (search radius times /// 2 for a 2-level pyramid), so for typical content the default - /// is fine. - /// - /// Raise it for very fast motion. + /// is fine. Raise it for very fast motion. /// /// Defaults to 4 when unset. #[arg(long)] @@ -48,9 +37,8 @@ pub struct MotionArgs { /// large motion). /// /// `2` (default) does a coarse pass on a half-size image first, - /// then refines at full resolution. - /// - /// This handles much larger motion at modest extra cost. + /// then refines at full resolution. This handles much larger motion + /// at modest extra cost. /// /// Defaults to 2 when unset. #[arg(long)] @@ -59,9 +47,6 @@ pub struct MotionArgs { impl MotionArgs { /// Whether any of these flags was given. - /// - /// A family that can leave motion tracking off uses this to warn - /// when the flags would go nowhere. pub fn any_set(&self) -> bool { self.mc_blksize.is_some() || self.mc_overlap.is_some() @@ -69,10 +54,10 @@ impl MotionArgs { || self.mc_pyramid_levels.is_some() } - /// These flags as the library's [`MotionSearch`], with the library - /// default for whatever was left unset. + /// Converts these flags to a [MotionSearch], with the library default for each unset flag. pub fn to_motion_search(&self) -> MotionSearch { let defaults = MotionSearch::default(); + MotionSearch { blksize: self.mc_blksize.unwrap_or(defaults.blksize), overlap: self.mc_overlap.unwrap_or(defaults.overlap), diff --git a/av-denoise/src/bin/cli/nl4d.rs b/av-denoise/src/bin/cli/nl4d.rs index a9e215c..0d9aaa4 100644 --- a/av-denoise/src/bin/cli/nl4d.rs +++ b/av-denoise/src/bin/cli/nl4d.rs @@ -12,13 +12,9 @@ use av_denoise::{ use super::{Args, CommonArgs, MotionArgs, Preset, RunOptions, resolve_channel_intent}; -/// Flags for `nl4d`, which groups 8x8 patches across the temporal window -/// itself, rather than filtering with `nlmeans` first and grouping -/// within one frame afterward. +/// Flags for the `nl4d` subcommand. /// -/// nl4d measures the noise level and tracks motion the way `nlmeans hq` -/// does, but never weights or averages patches the NLM way, so none of -/// the NLM knobs appear here. +/// nl4d never weights or averages patches the NLM way, so none of the NLM flags appear here. #[derive(Debug, Clone, clap::Args)] pub struct Nl4dArgs { #[command(flatten)] @@ -109,7 +105,7 @@ pub struct Nl4dArgs { /// noise. `3` is subtle grain, `6` is clearly visible grain, `12` /// and up is heavy noise. /// - /// Always expressed on an 8-bit 0-255 scale, no matter the + /// Always expressed on an 8-bit scale from 0 to 255, no matter the /// source's actual bit depth. #[arg(long)] pub sigma: Option, @@ -168,8 +164,6 @@ pub struct Nl4dArgs { /// Turns off the pooled threshold, which judges each frequency together with its neighbours /// so faint texture survives. - /// - /// With this flag and `--flat-boost 1.5`, the output matches earlier releases. #[arg(long)] pub no_pooled_threshold: bool, @@ -206,10 +200,9 @@ pub struct Nl4dArgs { /// Estimates noise from a local window instead of a temporal EMA /// over stream history. /// - /// Experimental measurement switch for comparing the two estimators - /// on real footage. Not a committed public interface. Off by - /// default, which keeps the temporal EMA every calibrated preset - /// assumes. + /// An experimental switch for comparing the two estimators on real + /// footage. It is not a stable interface. Off by default, which + /// keeps the temporal EMA every calibrated preset assumes. #[arg(long, hide = true)] pub windowed_noise_estimation: bool, @@ -228,13 +221,11 @@ pub struct Nl4dArgs { } impl Nl4dArgs { - /// How far the temporal window reaches, from the explicit flag or - /// the active `--preset`. + /// How far the temporal window reaches, from the explicit flag or the active `--preset`. /// - /// nl4d groups patches across neighbouring frames, so a window of 0 - /// leaves it nothing to do. No preset resolves that way, so only an - /// explicit `--temporal-radius 0` reaches the guard, and it is - /// rejected here rather than left to fail deep inside construction. + /// nl4d groups patches across neighbouring frames, so a radius of 0 leaves it nothing to do. + /// Only an explicit `--temporal-radius 0` can ask for one, and it is rejected here rather + /// than left to fail deep inside construction. fn temporal_radius(&self, preset: Preset) -> Result { let radius = self .temporal_radius @@ -250,8 +241,7 @@ impl Nl4dArgs { Ok(radius) } - /// Turns the parsed flags plus the shared globals into the options - /// the ingest pipeline takes. + /// Builds the run options from these flags and the global flags. pub fn build_options(&self, globals: &Args) -> Result { let defaults = Nl4dOptions::default(); let intent = resolve_channel_intent(&globals.channel_mode)?; @@ -268,71 +258,70 @@ impl Nl4dArgs { ); } - // Check the raw 8-bit value here so an out-of-range `--sigma` - // reports the number the user typed. The library re-validates - // the same bound after the /255 normalisation, but its message - // speaks in [0, 1] units. + // Checked in 8-bit units so the error reports the number the user typed. The library's + // own check runs after the /255 normalisation, in 0..=1 units. if let Some(sigma) = self.sigma && (!sigma.is_finite() || sigma <= 0.0 || sigma > 255.0) { anyhow::bail!("--sigma must be a finite value in (0, 255] 8-bit units (got {sigma})"); } - if self.sigma.is_some() && self.sigma_scale.is_some_and(|v| v != 1.0) { + if self.sigma.is_some() && self.sigma_scale.is_some_and(|scale| scale != 1.0) { tracing::warn!("--sigma-scale has no effect when --sigma pins the noise level"); } self.warn_on_dead_per_plane_flags(intent); + let temporal_radius = self.temporal_radius(globals.preset)?; + + let spatial_radius = self + .spatial_radius + .unwrap_or_else(|| nl4d_spatial_radius_for(globals.preset)); + let nl4d_options = Nl4dOptions { + motion: self.motion.to_motion_search(), + temporal_radius, + sigma: self.sigma.map(|sigma| sigma / 255.0), + sigma_scale: self.sigma_scale.unwrap_or(defaults.sigma_scale), + thsad_scale: self.thsad_scale.unwrap_or(defaults.thsad_scale), + refine: self.refine.unwrap_or(defaults.refine), + spatial_radius, + // Left unset here because the default depends on the plane being denoised, which + // is not known yet. The per-plane overrides and defaults apply once it is. + lambda_ht: self.lambda_ht, + lambda_ht_scale: self.lambda_ht_scale.unwrap_or(defaults.lambda_ht_scale), + c_min: self.c_min.unwrap_or(defaults.c_min), + kaiser_beta: self.kaiser_beta.unwrap_or(defaults.kaiser_beta), + field_lambda: self.field_lambda.unwrap_or(defaults.field_lambda), + // Off unless the hidden flag asks, because every calibrated preset assumes the + // temporal EMA. Window-local estimation gives random-access determinism. + windowed_noise_estimation: self.windowed_noise_estimation, + noise_map: defaults.noise_map && !self.no_noise_map, + flat_boost: self.flat_boost.unwrap_or(defaults.flat_boost), + chroma_flat_boost: self.chroma_flat_boost.unwrap_or(defaults.chroma_flat_boost), + shadow_soften: self.shadow_soften.unwrap_or(defaults.shadow_soften), + flat_texture_cut: self.flat_texture_cut.unwrap_or(defaults.flat_texture_cut), + pooled_threshold: defaults.pooled_threshold && !self.no_pooled_threshold, + grain_export: exports_grain, + }; + let mode = DenoisingMode::Temporal { + radius: temporal_radius, + }; + + let planes = PlaneOptions { + accelerators: globals.accelerators.clone(), + device: globals.device.clone(), + intent, + mode, + algorithm: Algorithm::Nl4d(nl4d_options), + // nl4d has no NLM weighting pass for a strength to apply to. + luma_strength: None, + chroma_strength: None, + luma_lambda_ht: self.luma_lambda_ht, + chroma_lambda_ht: self.chroma_lambda_ht, + }; + Ok(RunOptions { - planes: PlaneOptions { - accelerators: globals.accelerators.clone(), - device: globals.device.clone(), - intent, - mode: DenoisingMode::Temporal { - radius: self.temporal_radius(globals.preset)?, - }, - algorithm: Algorithm::Nl4d(Nl4dOptions { - motion: self.motion.to_motion_search(), - sigma: self.sigma.map(|s| s / 255.0), - sigma_scale: self.sigma_scale.unwrap_or(defaults.sigma_scale), - thsad_scale: self.thsad_scale.unwrap_or(defaults.thsad_scale), - refine: self.refine.unwrap_or(defaults.refine), - spatial_radius: self - .spatial_radius - .unwrap_or_else(|| nl4d_spatial_radius_for(globals.preset)), - // Left unresolved when unset, rather than picked from - // `defaults` here, because the default depends on which - // plane is being denoised and that is not known yet. - // `PlaneOptions::algorithm_for` (av-denoise-core/src/frame/mod.rs) - // applies `--luma-`/`--chroma-lambda-ht` on top of this - // once the plane is known, and construction fills in the - // per-plane default for whatever is still unset. - lambda_ht: self.lambda_ht, - lambda_ht_scale: self.lambda_ht_scale.unwrap_or(defaults.lambda_ht_scale), - c_min: self.c_min.unwrap_or(defaults.c_min), - kaiser_beta: self.kaiser_beta.unwrap_or(defaults.kaiser_beta), - field_lambda: self.field_lambda.unwrap_or(defaults.field_lambda), - // The CLI keeps the temporal EMA every calibrated - // preset assumes by default. Only `av-denoise-vs` - // needs window-local estimation, for random-access - // determinism. `--windowed-noise-estimation` exists - // to measure the difference on real footage. - windowed_noise_estimation: self.windowed_noise_estimation, - noise_map: defaults.noise_map && !self.no_noise_map, - flat_boost: self.flat_boost.unwrap_or(defaults.flat_boost), - chroma_flat_boost: self.chroma_flat_boost.unwrap_or(defaults.chroma_flat_boost), - shadow_soften: self.shadow_soften.unwrap_or(defaults.shadow_soften), - flat_texture_cut: self.flat_texture_cut.unwrap_or(defaults.flat_texture_cut), - pooled_threshold: defaults.pooled_threshold && !self.no_pooled_threshold, - grain_export: exports_grain, - }), - // nl4d has no NLM weighting pass for a strength to apply to. - luma_strength: None, - chroma_strength: None, - luma_lambda_ht: self.luma_lambda_ht, - chroma_lambda_ht: self.chroma_lambda_ht, - }, + planes, progress: globals.progress, grain_table: self.unstable_export_av1_fgs.clone(), }) @@ -362,11 +351,10 @@ impl Nl4dArgs { mod tests { use clap::Parser; - use super::super::Command; use super::*; + use crate::cli::Command; - /// Parses a full argv into the `nl4d` subcommand's args plus the - /// globals they resolve against. + /// Parses a `nl4d` argv reading stdin, with `extra` after the subcommand. fn parse(extra: &[&str]) -> (Args, Nl4dArgs) { let mut argv = vec!["av-denoise", "nl4d", "-i", "-"]; argv.extend_from_slice(extra); @@ -380,8 +368,7 @@ mod tests { (args, nl4d) } - /// Parses an argv that clap itself is expected to reject, which is - /// what an `nlmeans`-only flag should now be. + /// Parses a `nl4d` argv that clap is expected to reject. fn parse_err(extra: &[&str]) -> clap::Error { let mut argv = vec!["av-denoise", "nl4d", "-i", "-"]; argv.extend_from_slice(extra); @@ -389,8 +376,6 @@ mod tests { Args::try_parse_from(argv).expect_err("expected clap to reject this argv") } - /// Unwraps a `RunOptions`'s `Algorithm::Nl4d`, panicking with the - /// whole value on any other variant. fn expect_nl4d(opts: &RunOptions) -> Nl4dOptions { match opts.planes.algorithm { Algorithm::Nl4d(nl4d) => nl4d, @@ -401,6 +386,7 @@ mod tests { #[test] fn nl4d_subcommand_parses_with_no_extra_flags() { let (_, nl4d) = parse(&[]); + assert_eq!(nl4d.temporal_radius, None); assert_eq!(nl4d.refine, None); assert_eq!(nl4d.spatial_radius, None); @@ -453,80 +439,94 @@ mod tests { fn lambda_ht_scale_flows_into_the_nl4d_algorithm() { let (args, nl4d) = parse(&["--lambda-ht-scale", "1.1"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); - assert!((expect_nl4d(&opts).lambda_ht_scale - 1.1).abs() < f32::EPSILON); + assert!((nl4d_options.lambda_ht_scale - 1.1).abs() < f32::EPSILON); } #[test] fn unset_lambda_ht_scale_resolves_to_the_library_default() { let (args, nl4d) = parse(&[]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); let defaults = Nl4dOptions::default(); - assert!((expect_nl4d(&opts).lambda_ht_scale - defaults.lambda_ht_scale).abs() < f32::EPSILON); + assert!((nl4d_options.lambda_ht_scale - defaults.lambda_ht_scale).abs() < f32::EPSILON); } #[test] fn field_lambda_flows_into_the_nl4d_algorithm() { let (args, nl4d) = parse(&["--field-lambda", "0.7"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); assert!((nl4d.field_lambda.unwrap() - 0.7).abs() < f32::EPSILON); - assert!((expect_nl4d(&opts).field_lambda - 0.7).abs() < f32::EPSILON); + assert!((nl4d_options.field_lambda - 0.7).abs() < f32::EPSILON); } #[test] fn unset_field_lambda_resolves_to_the_library_default() { let (args, nl4d) = parse(&[]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); let defaults = Nl4dOptions::default(); assert_eq!(nl4d.field_lambda, None); - assert!((expect_nl4d(&opts).field_lambda - defaults.field_lambda).abs() < f32::EPSILON); + assert!((nl4d_options.field_lambda - defaults.field_lambda).abs() < f32::EPSILON); } #[test] fn noise_map_defaults_to_on() { let (args, nl4d) = parse(&[]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - assert!(expect_nl4d(&opts).noise_map); + let nl4d_options = expect_nl4d(&opts); + + assert!(nl4d_options.noise_map); } #[test] fn no_noise_map_turns_it_off() { let (args, nl4d) = parse(&["--no-noise-map"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - assert!(!expect_nl4d(&opts).noise_map); + let nl4d_options = expect_nl4d(&opts); + + assert!(!nl4d_options.noise_map); } #[test] fn there_is_no_positive_noise_map_flag() { - let err = parse_err(&["--noise-map"]); - assert_eq!(err.kind(), clap::error::ErrorKind::UnknownArgument); + let error = parse_err(&["--noise-map"]); + + assert_eq!(error.kind(), clap::error::ErrorKind::UnknownArgument); } #[test] fn pooled_threshold_defaults_to_on() { let (args, nl4d) = parse(&[]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - assert!(expect_nl4d(&opts).pooled_threshold); + let nl4d_options = expect_nl4d(&opts); + + assert!(nl4d_options.pooled_threshold); } #[test] fn no_pooled_threshold_turns_it_off() { let (args, nl4d) = parse(&["--no-pooled-threshold"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - assert!(!expect_nl4d(&opts).pooled_threshold); + let nl4d_options = expect_nl4d(&opts); + + assert!(!nl4d_options.pooled_threshold); } #[test] fn there_is_no_positive_pooled_threshold_flag() { - let err = parse_err(&["--pooled-threshold"]); - assert_eq!(err.kind(), clap::error::ErrorKind::UnknownArgument); + let error = parse_err(&["--pooled-threshold"]); + + assert_eq!(error.kind(), clap::error::ErrorKind::UnknownArgument); } /// nl4d never runs an NLM weighting pass, so the flags that only - /// configure one are gone rather than silently ignored. + /// configure one are rejected rather than silently ignored. #[test] fn nlmeans_only_flags_are_rejected() { let valued = [ @@ -550,20 +550,22 @@ mod tests { ]; for flag in valued { - let err = parse_err(&[flag, "1"]); + let error = parse_err(&[flag, "1"]); + assert_eq!( - err.kind(), + error.kind(), clap::error::ErrorKind::UnknownArgument, - "{flag} should not exist under nl4d, got {err}" + "{flag} should not exist under nl4d, got {error}" ); } for flag in switches { - let err = parse_err(&[flag]); + let error = parse_err(&[flag]); + assert_eq!( - err.kind(), + error.kind(), clap::error::ErrorKind::UnknownArgument, - "{flag} should not exist under nl4d, got {err}" + "{flag} should not exist under nl4d, got {error}" ); } } @@ -576,17 +578,16 @@ mod tests { "--luma-mismatch-scale=1.0", "--chroma-mismatch-scale=1.0", ] { - let err = parse_err(&[flag]); + let error = parse_err(&[flag]); + assert_eq!( - err.kind(), + error.kind(), clap::error::ErrorKind::UnknownArgument, - "{flag} must no longer parse, got {err}" + "{flag} must no longer parse, got {error}" ); } } - /// The motion-search knobs stay, because nl4d reads the motion field - /// they shape. #[test] fn motion_flags_flow_into_the_motion_search() { let (args, nl4d) = parse(&[ @@ -608,15 +609,14 @@ mod tests { assert_eq!(motion.pyramid_levels, 1); } - /// Only an explicit flag can ask for a window nl4d cannot use. No - /// preset resolves to 0. #[test] fn an_explicit_temporal_radius_of_zero_is_rejected() { let (args, nl4d) = parse(&["--temporal-radius", "0"]); - let err = nl4d + let error = nl4d .build_options(&args) .expect_err("radius 0 must be rejected under nl4d"); - assert!(err.to_string().contains("--temporal-radius"), "got {err}"); + + assert!(error.to_string().contains("--temporal-radius"), "got {error}"); } #[test] @@ -633,7 +633,8 @@ mod tests { let (args, nl4d) = parse(&["--preset", preset]); let opts = nl4d .build_options(&args) - .unwrap_or_else(|e| panic!("preset {preset} should resolve: {e}")); + .unwrap_or_else(|error| panic!("preset {preset} should resolve: {error}")); + let nl4d_options = expect_nl4d(&opts); assert_eq!( opts.planes.mode, @@ -641,8 +642,7 @@ mod tests { "preset {preset} resolved to the wrong temporal radius" ); assert_eq!( - expect_nl4d(&opts).spatial_radius, - spatial_radius, + nl4d_options.spatial_radius, spatial_radius, "preset {preset} resolved to the wrong spatial radius" ); } @@ -653,20 +653,23 @@ mod tests { #[test] fn veryfast_searches_fewer_candidates_than_fast() { let (fast_args, fast) = parse(&["--preset", "fast"]); - let (vf_args, vf) = parse(&["--preset", "veryfast"]); + let (veryfast_args, veryfast) = parse(&["--preset", "veryfast"]); - let fast = expect_nl4d(&fast.build_options(&fast_args).expect("fast should resolve")); - let vf = expect_nl4d(&vf.build_options(&vf_args).expect("veryfast should resolve")); + let fast_opts = fast.build_options(&fast_args).expect("fast should resolve"); + let veryfast_opts = veryfast + .build_options(&veryfast_args) + .expect("veryfast should resolve"); + let fast_nl4d = expect_nl4d(&fast_opts); + let veryfast_nl4d = expect_nl4d(&veryfast_opts); assert!( - vf.spatial_radius < fast.spatial_radius, + veryfast_nl4d.spatial_radius < fast_nl4d.spatial_radius, "veryfast ({}) should search a narrower window than fast ({})", - vf.spatial_radius, - fast.spatial_radius + veryfast_nl4d.spatial_radius, + fast_nl4d.spatial_radius ); } - /// An explicit flag outranks the preset on both dials. #[test] fn explicit_radii_outrank_the_preset() { let (args, nl4d) = parse(&[ @@ -678,9 +681,10 @@ mod tests { "12", ]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); assert_eq!(opts.planes.mode, DenoisingMode::Temporal { radius: 4 }); - assert_eq!(expect_nl4d(&opts).spatial_radius, 12); + assert_eq!(nl4d_options.spatial_radius, 12); } #[test] @@ -689,7 +693,7 @@ mod tests { let opts = nl4d .build_options(&args) .expect("default preset should resolve for nl4d"); - let nl4d_opts = expect_nl4d(&opts); + let nl4d_options = expect_nl4d(&opts); let defaults = Nl4dOptions::default(); assert_eq!( @@ -697,40 +701,38 @@ mod tests { DenoisingMode::Temporal { radius: 2 }, "base preset resolves to temporal radius 2" ); - assert_eq!(nl4d_opts.motion, defaults.motion); - assert_eq!(nl4d_opts.refine, defaults.refine); - assert_eq!(nl4d_opts.spatial_radius, defaults.spatial_radius); - assert_eq!(nl4d_opts.sigma, defaults.sigma); - assert!((nl4d_opts.sigma_scale - defaults.sigma_scale).abs() < f32::EPSILON); - assert!((nl4d_opts.thsad_scale - defaults.thsad_scale).abs() < f32::EPSILON); + assert_eq!(nl4d_options.motion, defaults.motion); + assert_eq!(nl4d_options.refine, defaults.refine); + assert_eq!(nl4d_options.spatial_radius, defaults.spatial_radius); + assert_eq!(nl4d_options.sigma, defaults.sigma); + assert!((nl4d_options.sigma_scale - defaults.sigma_scale).abs() < f32::EPSILON); + assert!((nl4d_options.thsad_scale - defaults.thsad_scale).abs() < f32::EPSILON); assert_eq!( - nl4d_opts.lambda_ht, defaults.lambda_ht, + nl4d_options.lambda_ht, defaults.lambda_ht, "unset --lambda-ht should stay None here, resolved later per plane" ); - assert!((nl4d_opts.c_min - defaults.c_min).abs() < f32::EPSILON); + assert!((nl4d_options.c_min - defaults.c_min).abs() < f32::EPSILON); } #[test] fn explicit_grouping_flags_override_the_library_defaults() { let (args, nl4d) = parse(&["--refine", "4", "--spatial-radius", "6", "--c-min", "0.2"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - let nl4d_opts = expect_nl4d(&opts); + let nl4d_options = expect_nl4d(&opts); - assert_eq!(nl4d_opts.refine, 4); - assert_eq!(nl4d_opts.spatial_radius, 6); - assert!((nl4d_opts.c_min - 0.2).abs() < f32::EPSILON); + assert_eq!(nl4d_options.refine, 4); + assert_eq!(nl4d_options.spatial_radius, 6); + assert!((nl4d_options.c_min - 0.2).abs() < f32::EPSILON); } - /// `--sigma` is typed in 8-bit units and the library takes it - /// normalised, so `build_options` is where the /255 happens. #[test] fn sigma_is_normalised_out_of_eight_bit_units() { let (args, nl4d) = parse(&["--sigma", "6"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); + + let sigma = nl4d_options.sigma.expect("--sigma should reach the options"); - let sigma = expect_nl4d(&opts) - .sigma - .expect("--sigma should reach the options"); assert!( (sigma - 6.0 / 255.0).abs() < f32::EPSILON, "expected 6/255, got {sigma}" @@ -740,33 +742,31 @@ mod tests { #[test] fn out_of_range_sigma_is_rejected() { let (args, nl4d) = parse(&["--sigma", "300"]); - let err = nl4d.build_options(&args).expect_err("300 is out of range"); + let error = nl4d.build_options(&args).expect_err("300 is out of range"); - assert!(err.to_string().contains("--sigma"), "got {err}"); + assert!(error.to_string().contains("--sigma"), "got {error}"); } #[test] fn sigma_scale_and_thsad_scale_flow_into_the_nl4d_algorithm() { let (args, nl4d) = parse(&["--sigma-scale", "1.2", "--thsad-scale", "0.7"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - let nl4d_opts = expect_nl4d(&opts); + let nl4d_options = expect_nl4d(&opts); - assert!((nl4d_opts.sigma_scale - 1.2).abs() < f32::EPSILON); - assert!((nl4d_opts.thsad_scale - 0.7).abs() < f32::EPSILON); + assert!((nl4d_options.sigma_scale - 1.2).abs() < f32::EPSILON); + assert!((nl4d_options.thsad_scale - 0.7).abs() < f32::EPSILON); } #[test] fn per_plane_lambda_ht_flags_parse_to_the_typed_values() { let (_, nl4d) = parse(&["--luma-lambda-ht", "2.0", "--chroma-lambda-ht", "3.5"]); + assert!((nl4d.luma_lambda_ht.unwrap() - 2.0).abs() < f32::EPSILON); assert!((nl4d.chroma_lambda_ht.unwrap() - 3.5).abs() < f32::EPSILON); } - /// `build_options` carries the two per-plane `lambda_ht` overrides - /// straight through onto `PlaneOptions`, unresolved. - /// `PlaneOptions::algorithm_for` (`av-denoise-core/src/frame/mod.rs`) - /// is what actually resolves them per plane, so this test only checks - /// the flow into `PlaneOptions`. + /// The per-plane overrides reach `PlaneOptions` unresolved, because + /// they are resolved per plane later, so this only checks that flow. #[test] fn luma_lambda_ht_alone_flows_into_cli_options_luma_field_only() { let (args, nl4d) = parse(&["--luma-lambda-ht", "2.0"]); @@ -816,61 +816,67 @@ mod tests { "0.3", ]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - let algorithm = expect_nl4d(&opts); + let nl4d_options = expect_nl4d(&opts); - assert_eq!(algorithm.flat_boost, 2.0); - assert_eq!(algorithm.chroma_flat_boost, 1.2); - assert_eq!(algorithm.shadow_soften, 0.8); - assert_eq!(algorithm.flat_texture_cut, 0.3); + assert_eq!(nl4d_options.flat_boost, 2.0); + assert_eq!(nl4d_options.chroma_flat_boost, 1.2); + assert_eq!(nl4d_options.shadow_soften, 0.8); + assert_eq!(nl4d_options.flat_texture_cut, 0.3); } #[test] fn unset_strength_map_flags_resolve_to_the_library_defaults() { let (args, nl4d) = parse(&[]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); - let algorithm = expect_nl4d(&opts); + let nl4d_options = expect_nl4d(&opts); let defaults = Nl4dOptions::default(); - assert_eq!(algorithm.flat_boost, defaults.flat_boost); - assert_eq!(algorithm.chroma_flat_boost, defaults.chroma_flat_boost); - assert_eq!(algorithm.shadow_soften, defaults.shadow_soften); - assert_eq!(algorithm.flat_texture_cut, defaults.flat_texture_cut); + assert_eq!(nl4d_options.flat_boost, defaults.flat_boost); + assert_eq!(nl4d_options.chroma_flat_boost, defaults.chroma_flat_boost); + assert_eq!(nl4d_options.shadow_soften, defaults.shadow_soften); + assert_eq!(nl4d_options.flat_texture_cut, defaults.flat_texture_cut); } #[test] fn export_flag_sets_the_table_and_turns_export_on() { let (args, nl4d) = parse(&["--unstable-export-av1-fgs", "out.tbl"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); let expected = std::path::Path::new("out.tbl"); assert_eq!(opts.grain_table.as_deref(), Some(expected)); - assert!(expect_nl4d(&opts).grain_export); + assert!(nl4d_options.grain_export); } #[test] fn export_works_with_a_fixed_sigma() { let (args, nl4d) = parse(&["--sigma", "4", "--unstable-export-av1-fgs", "x"]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); - assert!(expect_nl4d(&opts).grain_export); + assert!(nl4d_options.grain_export); } #[test] fn export_is_off_without_the_flag() { let (args, nl4d) = parse(&[]); let opts = nl4d.build_options(&args).expect("build_options should succeed"); + let nl4d_options = expect_nl4d(&opts); assert_eq!(opts.grain_table, None); - assert!(!expect_nl4d(&opts).grain_export); + assert!(!nl4d_options.grain_export); } #[test] fn export_rejects_chroma_only() { let (args, nl4d) = parse(&["--channel-mode", "chroma", "--unstable-export-av1-fgs", "out.tbl"]); - let err = nl4d + let error = nl4d .build_options(&args) .expect_err("chroma only cannot export luma grain"); - assert!(err.to_string().contains("--unstable-export-av1-fgs"), "got {err}"); + assert!( + error.to_string().contains("--unstable-export-av1-fgs"), + "got {error}" + ); } } diff --git a/av-denoise/src/bin/cli/nlmeans.rs b/av-denoise/src/bin/cli/nlmeans.rs index 73137fe..37ef647 100644 --- a/av-denoise/src/bin/cli/nlmeans.rs +++ b/av-denoise/src/bin/cli/nlmeans.rs @@ -46,7 +46,7 @@ pub struct NlmeansArgs { /// blur first, then compares patches against that cleaner image. /// /// `sigma_s` is the spatial blur radius in pixels, greater than 0 - /// and at most 11.0 (anything beyond this is insane.) + /// and at most 11.0. /// /// `sigma_r` is the colour-similarity threshold, greater than 0. /// Values above `0` and up to `1` are typical for normalised @@ -87,11 +87,11 @@ pub struct NlmeansArgs { #[arg(long)] pub search_radius: Option, - /// Size of each patch being compared. The patch is - /// `(2*patch_radius + 1)` pixels square. + /// Radius of each patch being compared. /// - /// Larger patches preserve fine structure better but cost more - /// GPU memory. Library default is 4. + /// The patch is `(2*patch_radius + 1)` pixels square. Larger patches + /// preserve fine structure better but cost more GPU memory. Library + /// default is 4. #[arg(long)] pub patch_radius: Option, @@ -151,7 +151,7 @@ pub struct NlmeansArgs { /// noise. `3` is subtle grain, `6` is clearly visible grain, `12` /// and up is heavy noise. /// - /// Always expressed on an 8-bit 0-255 scale, no matter the + /// Always expressed on an 8-bit scale from 0 to 255, no matter the /// source's actual bit depth. #[arg(long)] pub hq_sigma: Option, @@ -227,11 +227,10 @@ pub struct NlmeansArgs { /// Estimates noise from a local window instead of a temporal EMA /// over stream history. /// - /// Experimental measurement switch for comparing the two estimators - /// on real footage. Not a committed public interface. Off by - /// default, which keeps the temporal EMA every calibrated preset - /// assumes. Only applies to `--variant hq`, since `fast` never - /// measures noise. + /// An experimental switch for comparing the two estimators on real + /// footage. It is not a stable interface. Off by default, which + /// keeps the temporal EMA every calibrated preset assumes. Only + /// applies to `--variant hq`, since `fast` never measures noise. #[arg(long, hide = true)] pub windowed_noise_estimation: bool, @@ -262,17 +261,16 @@ impl NlmeansArgs { } } - /// Builds the library's [`av_denoise::Algorithm`] from the resolved - /// preset and the flags the chosen variant reads. + /// Builds the library's [Algorithm](av_denoise::Algorithm) from the resolved preset and + /// the flags the chosen variant reads. /// - /// Flags that are set but do nothing for this configuration are - /// reported as warnings. + /// Flags that are set but do nothing for this configuration are reported as warnings. pub fn resolve_algorithm( &self, resolved: ResolvedPreset, nlm: NlmeansOptions, ) -> Result { - let sigma_scale_is_set = self.hq_sigma_scale.is_some_and(|v| v != 1.0); + let sigma_scale_is_set = self.hq_sigma_scale.is_some_and(|scale| scale != 1.0); match resolved.variant { Variant::Fast => { @@ -292,13 +290,12 @@ impl NlmeansArgs { { tracing::warn!("--hq-* options are ignored unless --variant hq is selected"); } + Ok(av_denoise::Algorithm::Nlmeans(nlm)) }, Variant::Hq => { - // Check the raw 8-bit value here so an out-of-range - // `--hq-sigma` reports the number the user typed. The - // library re-validates the same bound after the /255 - // normalisation, but its message speaks in [0, 1] units. + // Checked in 8-bit units so the error reports the number the user typed. The + // library's own check runs after the /255 normalisation, in 0..=1 units. if let Some(sigma) = self.hq_sigma && (!sigma.is_finite() || sigma <= 0.0 || sigma > 255.0) { @@ -309,30 +306,25 @@ impl NlmeansArgs { tracing::warn!("--hq-sigma-scale has no effect when --hq-sigma pins the noise level"); } - Ok(av_denoise::Algorithm::NlmeansHq(NlmeansHqOptions { - nlm, - hq: av_denoise::HqParams { - auto_strength: !self.hq_no_auto_strength, - noise_floor: !self.hq_no_noise_floor, - sigma_override: self.hq_sigma.map(|s| s / 255.0), - temporal_confidence: !self.hq_no_temporal_confidence, - thsad_scale: self.hq_thsad_scale.unwrap_or(1.0), - sigma_scale: self.hq_sigma_scale.unwrap_or(1.0), - // The CLI keeps the temporal EMA every - // calibrated preset assumes by default. Only - // `av-denoise-vs` needs window-local estimation, - // for random-access determinism. - // `--windowed-noise-estimation` exists to - // measure the difference on real footage. - windowed_noise_estimation: self.windowed_noise_estimation, - }, - })) + let hq = av_denoise::HqParams { + auto_strength: !self.hq_no_auto_strength, + noise_floor: !self.hq_no_noise_floor, + sigma_override: self.hq_sigma.map(|sigma| sigma / 255.0), + temporal_confidence: !self.hq_no_temporal_confidence, + thsad_scale: self.hq_thsad_scale.unwrap_or(1.0), + sigma_scale: self.hq_sigma_scale.unwrap_or(1.0), + // Off unless the hidden flag asks, because every calibrated preset assumes + // the temporal EMA. Window-local estimation gives random-access determinism. + windowed_noise_estimation: self.windowed_noise_estimation, + }; + let options = NlmeansHqOptions { nlm, hq }; + + Ok(av_denoise::Algorithm::NlmeansHq(options)) }, } } - /// Turns the parsed flags plus the shared globals into the options - /// the ingest pipeline takes. + /// Builds the run options from these flags and the global flags. pub fn build_options(&self, globals: &Args) -> Result { let resolved = self.resolve_preset(globals.preset); @@ -359,41 +351,47 @@ impl NlmeansArgs { the spatial path doesn't use temporal neighbours" ); } + self.motion.to_motion_search().into() } else { if self.motion.any_set() { tracing::warn!("--mc-* options are ignored unless --motion-compensation is set"); } + MotionCompensationMode::None }; - // search_radius always has a resolved value (explicit flag or the - // active preset), so it's always carried into the tuning override. + // The search radius always resolves, from the flag or the preset, so it always + // overrides the library tuning. + let tuning = NlmTuning { + search_radius: Some(resolved.search_radius), + patch_radius: self.patch_radius, + strength: self.strength, + self_weight: self.self_weight, + }; let nlm = NlmeansOptions { prefilter, motion_compensation, - tuning: NlmTuning { - search_radius: Some(resolved.search_radius), - patch_radius: self.patch_radius, - strength: self.strength, - self_weight: self.self_weight, - }, + tuning, + mode, + }; + + let algorithm = self.resolve_algorithm(resolved, nlm)?; + let planes = PlaneOptions { + accelerators: globals.accelerators.clone(), + device: globals.device.clone(), + intent, + mode, + algorithm, + luma_strength: self.luma_strength, + chroma_strength: self.chroma_strength, + // `nlmeans` has no grouping stage for a threshold to apply to. + luma_lambda_ht: None, + chroma_lambda_ht: None, }; Ok(RunOptions { - planes: PlaneOptions { - accelerators: globals.accelerators.clone(), - device: globals.device.clone(), - intent, - mode, - algorithm: self.resolve_algorithm(resolved, nlm)?, - luma_strength: self.luma_strength, - chroma_strength: self.chroma_strength, - // `nlmeans` has no grouping stage, so these stay unset - // here. `Nl4dArgs::build_options` fills them in afterwards. - luma_lambda_ht: None, - chroma_lambda_ht: None, - }, + planes, progress: globals.progress, grain_table: None, }) @@ -405,14 +403,12 @@ mod tests { use av_denoise::DEFAULT_PILOT_STRENGTH_SCALE; use clap::Parser; - use super::super::{Args, CliChannelMode, Command, InputSource, Preset}; use super::*; + use crate::cli::{Args, CliChannelMode, Command, InputSource, Preset}; - /// Parses a full argv into the `nlmeans` subcommand's args plus the - /// globals they resolve against. + /// Parses a `nlmeans` argv reading stdin, with `extra` after the subcommand. /// - /// `extra` follows the subcommand because the family-owned flags - /// only exist there. The globals parse in that position too. + /// `extra` follows the subcommand because the subcommand's own flags only exist there. fn parse(extra: &[&str]) -> (Args, NlmeansArgs) { let mut argv = vec!["av-denoise", "nlmeans", "-i", "-"]; argv.extend_from_slice(extra); @@ -426,8 +422,7 @@ mod tests { (args, nlm) } - /// Parses an argv whose `-i`/`-W` values are supplied by the caller, - /// unlike [`parse`] which pins the input to stdin. + /// Parses a `nlmeans` argv whose `-i`/`-W` values come from `extra`. fn parse_input(extra: &[&str]) -> NlmeansArgs { let mut argv = vec!["av-denoise", "nlmeans"]; argv.extend_from_slice(extra); @@ -447,6 +442,7 @@ mod tests { #[test] fn default_preset_is_base() { let resolved = resolve(&[]); + assert_eq!(resolved.variant, Variant::Hq); assert_eq!(resolved.temporal_radius, 2); assert_eq!(resolved.search_radius, 2); @@ -455,6 +451,7 @@ mod tests { #[test] fn veryfast_matches_former_defaults() { let resolved = resolve(&["--preset", "veryfast"]); + assert_eq!(resolved.variant, Variant::Fast); assert_eq!(resolved.temporal_radius, 0); assert_eq!(resolved.search_radius, 2); @@ -463,6 +460,7 @@ mod tests { #[test] fn fast_grid_row() { let resolved = resolve(&["--preset", "fast"]); + assert_eq!(resolved.variant, Variant::Hq); assert_eq!(resolved.temporal_radius, 1); assert_eq!(resolved.search_radius, 2); @@ -471,6 +469,7 @@ mod tests { #[test] fn base_grid_row() { let resolved = resolve(&["--preset", "base"]); + assert_eq!(resolved.variant, Variant::Hq); assert_eq!(resolved.temporal_radius, 2); assert_eq!(resolved.search_radius, 2); @@ -479,6 +478,7 @@ mod tests { #[test] fn slow_grid_row() { let resolved = resolve(&["--preset", "slow"]); + assert_eq!(resolved.variant, Variant::Hq); assert_eq!(resolved.temporal_radius, 4); assert_eq!(resolved.search_radius, 4); @@ -487,6 +487,7 @@ mod tests { #[test] fn veryslow_grid_row() { let resolved = resolve(&["--preset", "veryslow"]); + assert_eq!(resolved.variant, Variant::Hq); assert_eq!(resolved.temporal_radius, 8); assert_eq!(resolved.search_radius, 4); @@ -495,6 +496,7 @@ mod tests { #[test] fn explicit_variant_overrides_the_preset() { let resolved = resolve(&["--preset", "base", "--variant", "fast"]); + assert_eq!(resolved.variant, Variant::Fast); assert_eq!(resolved.temporal_radius, 2); assert_eq!(resolved.search_radius, 2); @@ -503,6 +505,7 @@ mod tests { #[test] fn explicit_temporal_radius_overrides_the_preset() { let resolved = resolve(&["--preset", "base", "--temporal-radius", "6"]); + assert_eq!(resolved.variant, Variant::Hq); assert_eq!(resolved.temporal_radius, 6); assert_eq!(resolved.search_radius, 2); @@ -511,6 +514,7 @@ mod tests { #[test] fn explicit_search_radius_overrides_the_preset() { let resolved = resolve(&["--preset", "veryslow", "--search-radius", "1"]); + assert_eq!(resolved.variant, Variant::Hq); assert_eq!(resolved.temporal_radius, 8); assert_eq!(resolved.search_radius, 1); @@ -519,32 +523,19 @@ mod tests { #[test] fn globals_parse_before_the_subcommand() { let args = Args::parse_from(["av-denoise", "--preset", "slow", "nlmeans", "-i", "-"]); + assert!(matches!(args.preset, Preset::Slow)); } - /// The five tests below check that every global flag is accepted - /// before the `nlmeans` token. That is the shape the `Justfile` - /// uses, passing `-A vulkan,cpu` ahead of the subcommand. - /// - /// Despite the name they do not guard `global = true`. Clap accepts - /// `Args`'s own fields in the leading position either way, because - /// `global` only starts to matter once the subcommand token has - /// been read. Dropping `global = true` from `accelerators`, - /// `device`, `channel_mode`, and `progress` in turn left all five - /// tests passing. - /// - /// What `global = true` really controls is the trailing position, - /// after `nlmeans`. The `*_parses_after_the_subcommand` tests - /// further down cover that. - /// - /// The repeated `vulkan` is deliberate. What this pins is that the - /// comma delimiter splits the value into a list, and `Vulkan` is - /// the only accelerator variant a default build is guaranteed to - /// have, so naming it twice tests the split without depending on a - /// second backend feature being enabled. - /// - /// Feature-gated because it names the `Vulkan` accelerator variant, - /// which only exists when its feature is enabled. + // The five tests below check that every global flag is accepted before the `nlmeans` + // token, the shape the `Justfile` uses for `-A vulkan,cpu`. They pass with or without + // `global = true`, because clap accepts `Args`'s own fields in the leading position either + // way. `global = true` controls the trailing position, which the + // `*_parses_after_the_subcommand` tests cover. + + /// The repeated `vulkan` tests the comma split without needing a second backend feature, + /// since `Vulkan` is the only variant a default build is guaranteed to have. Gated because + /// that variant only exists with the `vulkan` feature. #[cfg(feature = "vulkan")] #[test] fn accelerators_are_accepted_before_the_subcommand() { @@ -556,6 +547,7 @@ mod tests { "-i", "-", ]); + assert_eq!( args.accelerators, vec![ @@ -571,6 +563,7 @@ mod tests { #[test] fn short_accelerators_flag_is_accepted_before_the_subcommand() { let args = Args::parse_from(["av-denoise", "-A", "vulkan", "nlmeans", "-i", "-"]); + assert_eq!( args.accelerators, vec![av_denoise::accelerate::Accelerator::Vulkan] @@ -580,44 +573,45 @@ mod tests { #[test] fn device_is_accepted_before_the_subcommand() { let args = Args::parse_from(["av-denoise", "--device", "cpu", "nlmeans", "-i", "-"]); + assert!(matches!(args.device, av_denoise::Device::Cpu)); } #[test] fn channel_mode_is_accepted_before_the_subcommand() { let args = Args::parse_from(["av-denoise", "--channel-mode", "chroma", "nlmeans", "-i", "-"]); + assert_eq!(args.channel_mode, vec![CliChannelMode::Chroma]); } #[test] fn progress_flag_is_accepted_before_the_subcommand() { let args = Args::parse_from(["av-denoise", "--progress", "nlmeans", "-i", "-"]); + assert!(args.progress); } - /// The other half of the rule. `--strength` is owned by the - /// `nlmeans` subcommand rather than being global, so it must be - /// rejected before `nlmeans` even though the globals above are - /// accepted there. + /// `--strength` belongs to the `nlmeans` subcommand rather than + /// being global, so it must be rejected before `nlmeans` even + /// though the globals above are accepted there. #[test] fn a_subcommand_owned_flag_is_rejected_before_the_subcommand() { - let err = Args::try_parse_from(["av-denoise", "--strength", "1.2", "nlmeans", "-i", "-"]) + let error = Args::try_parse_from(["av-denoise", "--strength", "1.2", "nlmeans", "-i", "-"]) .expect_err("--strength is subcommand-owned, not global"); - assert!(err.to_string().contains("strength"), "got {err}"); + + assert!(error.to_string().contains("strength"), "got {error}"); } - /// This is the test that really depends on `--accelerators` being - /// `global = true`. Without it, a flag placed after `nlmeans` is - /// rejected as unknown to the subcommand's own parser. The - /// `denoise-file` recipe in the `Justfile` passes user flags in - /// exactly this position. - /// - /// Feature-gated because it names the `Vulkan` accelerator variant, - /// which only exists when its feature is enabled. + /// Depends on `--accelerators` being `global = true`. Without it, a + /// flag after `nlmeans` is rejected as unknown to the subcommand, + /// and the `denoise-file` recipe in the `Justfile` passes user flags + /// in exactly this position. Gated because it names the `Vulkan` + /// accelerator variant, which only exists with its feature. #[cfg(feature = "vulkan")] #[test] fn accelerators_parses_after_the_subcommand() { let (args, _) = parse(&["--accelerators", "vulkan"]); + assert_eq!( args.accelerators, vec![av_denoise::accelerate::Accelerator::Vulkan] @@ -625,25 +619,27 @@ mod tests { } /// Same `global = true` dependency as - /// [`accelerators_parses_after_the_subcommand`], for `--device`. + /// [accelerators_parses_after_the_subcommand], for `--device`. #[test] fn device_parses_after_the_subcommand() { let (args, _) = parse(&["--device", "cpu"]); + assert!(matches!(args.device, av_denoise::Device::Cpu)); } /// Same `global = true` dependency as - /// [`accelerators_parses_after_the_subcommand`], for - /// `--channel-mode`. + /// [accelerators_parses_after_the_subcommand], for `--channel-mode`. #[test] fn channel_mode_parses_after_the_subcommand() { let (args, _) = parse(&["--channel-mode", "chroma"]); + assert_eq!(args.channel_mode, vec![CliChannelMode::Chroma]); } #[test] fn channel_mode_defaults_to_luma_and_chroma() { let (args, _) = parse(&[]); + assert_eq!( args.channel_mode, vec![CliChannelMode::Luma, CliChannelMode::Chroma] @@ -654,6 +650,7 @@ mod tests { fn default_channel_mode_resolves_to_the_luma_chroma_intent() { let (args, _) = parse(&[]); let intent = resolve_channel_intent(&args.channel_mode).expect("default should resolve"); + assert!( matches!(intent, av_denoise::ChannelIntent::LumaChroma), "expected LumaChroma, got {intent:?}", @@ -663,12 +660,14 @@ mod tests { #[test] fn variant_flag_is_case_insensitive() { let (_, nlm) = parse(&["--variant", "HQ"]); + assert_eq!(nlm.variant, Some(Variant::Hq)); } #[test] fn hq_sigma_scale_parses_to_the_typed_value() { let (_, nlm) = parse(&["--hq-sigma-scale", "2.0"]); + assert_eq!(nlm.hq_sigma_scale, Some(2.0)); } @@ -676,8 +675,9 @@ mod tests { fn unset_hq_sigma_scale_resolves_to_the_library_default_of_one() { let (args, nlm) = parse(&["--variant", "hq"]); let resolved = nlm.resolve_preset(args.preset); + let nlm_options = NlmeansOptions::default(); let algorithm = nlm - .resolve_algorithm(resolved, NlmeansOptions::default()) + .resolve_algorithm(resolved, nlm_options) .expect("resolution should succeed"); match algorithm { @@ -690,8 +690,9 @@ mod tests { fn hq_flags_on_the_fast_variant_still_resolve_to_nlmeans() { let (args, nlm) = parse(&["--variant", "fast", "--hq-sigma-scale", "2.0"]); let resolved = nlm.resolve_preset(args.preset); + let nlm_options = NlmeansOptions::default(); let algorithm = nlm - .resolve_algorithm(resolved, NlmeansOptions::default()) + .resolve_algorithm(resolved, nlm_options) .expect("resolution should succeed"); assert!(matches!(algorithm, av_denoise::Algorithm::Nlmeans(_))); @@ -701,14 +702,18 @@ mod tests { fn hq_sigma_and_sigma_scale_both_resolve() { let (args, nlm) = parse(&["--variant", "hq", "--hq-sigma", "6", "--hq-sigma-scale", "2.0"]); let resolved = nlm.resolve_preset(args.preset); + let nlm_options = NlmeansOptions::default(); let algorithm = nlm - .resolve_algorithm(resolved, NlmeansOptions::default()) + .resolve_algorithm(resolved, nlm_options) .expect("resolution should succeed"); match algorithm { av_denoise::Algorithm::NlmeansHq(opts) => { assert!( - matches!(opts.hq.sigma_override, Some(s) if (s - 6.0 / 255.0).abs() < f32::EPSILON), + matches!( + opts.hq.sigma_override, + Some(sigma) if (sigma - 6.0 / 255.0).abs() < f32::EPSILON + ), "expected sigma_override Some(6/255), got {:?}", opts.hq.sigma_override, ); @@ -722,34 +727,41 @@ mod tests { fn out_of_range_hq_sigma_is_rejected() { let (args, nlm) = parse(&["--variant", "hq", "--hq-sigma", "300"]); let resolved = nlm.resolve_preset(args.preset); - let err = nlm - .resolve_algorithm(resolved, NlmeansOptions::default()) + let nlm_options = NlmeansOptions::default(); + let error = nlm + .resolve_algorithm(resolved, nlm_options) .expect_err("300 is out of range"); - assert!(err.to_string().contains("--hq-sigma"), "got {err}"); + assert!(error.to_string().contains("--hq-sigma"), "got {error}"); } #[test] fn progress_defaults_to_false() { let (args, _) = parse(&[]); + assert!(!args.progress); } #[test] fn progress_flag_sets_it_true() { let (args, _) = parse(&["--progress"]); + assert!(args.progress); } #[test] fn none_and_empty_prefilters() { - assert!(matches!(parse_prefilter("none").unwrap(), PrefilterMode::None)); - assert!(matches!(parse_prefilter("").unwrap(), PrefilterMode::None)); + let none = parse_prefilter("none").unwrap(); + let empty = parse_prefilter("").unwrap(); + + assert!(matches!(none, PrefilterMode::None)); + assert!(matches!(empty, PrefilterMode::None)); } #[test] fn bilateral_with_values() { let mode = parse_prefilter("bilateral:3.0,0.02").unwrap(); + assert!(matches!( mode, PrefilterMode::Bilateral { @@ -762,6 +774,7 @@ mod tests { #[test] fn bare_nlm_uses_default_strength_scale() { let mode = parse_prefilter("nlm").unwrap(); + assert!(matches!( mode, PrefilterMode::NlmSpatial { strength_scale } @@ -772,6 +785,7 @@ mod tests { #[test] fn nlm_with_explicit_strength_scale() { let mode = parse_prefilter("nlm:0.8").unwrap(); + assert!(matches!( mode, PrefilterMode::NlmSpatial { strength_scale } if (strength_scale - 0.8).abs() < f32::EPSILON @@ -780,85 +794,99 @@ mod tests { #[test] fn malformed_nlm_scale_is_rejected() { - let err = parse_prefilter("nlm:x").expect_err("expected parse failure"); - assert!(err.to_string().contains("nlm")); + let error = parse_prefilter("nlm:x").expect_err("expected parse failure"); + + assert!(error.to_string().contains("nlm")); } #[test] fn unknown_prefilter_is_rejected() { - assert!(parse_prefilter("garbage").is_err()); + let result = parse_prefilter("garbage"); + + assert!(result.is_err()); } #[test] fn the_old_subcommands_are_rejected() { - assert!(Args::try_parse_from(["av-denoise", "stdin"]).is_err()); - assert!(Args::try_parse_from(["av-denoise", "file", "-i", "noisy.mkv"]).is_err()); + let stdin_result = Args::try_parse_from(["av-denoise", "stdin"]); + let file_result = Args::try_parse_from(["av-denoise", "file", "-i", "noisy.mkv"]); + + assert!(stdin_result.is_err()); + assert!(file_result.is_err()); } #[test] fn a_path_parses_as_a_file_source() { let nlm = parse_input(&["--input", "noisy.mkv"]); - assert_eq!( - nlm.common.input, - InputSource::File(std::path::PathBuf::from("noisy.mkv")) - ); + let path = std::path::PathBuf::from("noisy.mkv"); + + assert_eq!(nlm.common.input, InputSource::File(path)); } #[test] fn a_dash_parses_as_stdin() { let nlm = parse_input(&["-i", "-"]); + assert_eq!(nlm.common.input, InputSource::Stdin); } #[test] fn a_pipe_parses_as_a_descriptor() { let nlm = parse_input(&["-i", "pipe:3"]); + assert_eq!(nlm.common.input, InputSource::Fd(3)); } #[test] fn an_unreadable_descriptor_fails_to_parse() { - let err = Args::try_parse_from(["av-denoise", "nlmeans", "-i", "pipe:1"]) + let error = Args::try_parse_from(["av-denoise", "nlmeans", "-i", "pipe:1"]) .expect_err("pipe:1 is our own stdout"); - assert!(err.to_string().contains("stdout"), "got {err}"); + + assert!(error.to_string().contains("stdout"), "got {error}"); } #[test] fn workers_is_unset_by_default() { let nlm = parse_input(&["-i", "noisy.mkv"]); + assert_eq!(nlm.common.workers, None); } #[test] fn workers_carries_the_typed_value() { let nlm = parse_input(&["-i", "noisy.mkv", "--workers", "4"]); + assert_eq!(nlm.common.workers, Some(4)); } #[test] fn frame_budget_is_unset_by_default() { let nlm = parse_input(&["-i", "noisy.mkv"]); + assert_eq!(nlm.common.frame_budget, None); } #[test] fn frame_budget_reads_decimal_and_binary_units_apart() { let decimal = parse_input(&["-i", "noisy.mkv", "--frame-budget", "8GB"]); + assert_eq!(decimal.common.frame_budget, Some(8_000_000_000)); let binary = parse_input(&["-i", "noisy.mkv", "--frame-budget", "8GiB"]); + assert_eq!(binary.common.frame_budget, Some(8_589_934_592)); } #[test] fn frame_budget_reads_a_bare_number_as_bytes() { let nlm = parse_input(&["-i", "noisy.mkv", "--frame-budget", "512"]); + assert_eq!(nlm.common.frame_budget, Some(512)); } #[test] fn a_malformed_frame_budget_is_rejected() { - let err = Args::try_parse_from([ + let error = Args::try_parse_from([ "av-denoise", "nlmeans", "-i", @@ -869,14 +897,14 @@ mod tests { .expect_err("banana is not a size"); assert!( - err.to_string().contains("banana"), - "error should name the value: {err}" + error.to_string().contains("banana"), + "error should name the value: {error}" ); } #[test] fn nlmeans_has_no_export_flag() { - let err = Args::try_parse_from([ + let error = Args::try_parse_from([ "av-denoise", "nlmeans", "-i", @@ -886,7 +914,7 @@ mod tests { ]) .expect_err("only nl4d exports grain"); - assert_eq!(err.kind(), clap::error::ErrorKind::UnknownArgument); + assert_eq!(error.kind(), clap::error::ErrorKind::UnknownArgument); } #[test] diff --git a/av-denoise/src/bin/frame_index.rs b/av-denoise/src/bin/frame_index.rs index b63c267..40f18af 100644 --- a/av-denoise/src/bin/frame_index.rs +++ b/av-denoise/src/bin/frame_index.rs @@ -8,21 +8,18 @@ use ffms2_sys::{FFMS_GetFrameInfo, FFMS_GetNumFrames, FFMS_GetTrackFromVideo}; pub struct IndexEntry { /// Presentation timestamp, in the video track's time base. pub pts: i64, - /// Whether ffms2 marked this entry as a keyframe. pub keyframe: bool, } /// Reads the video index behind `decoder`. /// -/// Returns `None` when `decoder` is not backed by ffms2, or when ffms2 declines to describe the -/// track, when this happens caller then keeps every frame. +/// Returns `None` when `decoder` is not backed by ffms2 or when ffms2 cannot describe the track. pub fn read_index(decoder: &mut Decoder) -> Option> { let total = decoder.get_video_details().total_frames?; let source = decoder.get_ffms2_impl()?.video_source; - // SAFETY: a live `Ffms2Decoder` holds a non-null video source, and - // the track belongs to that source rather than to us, so it stays - // valid for as long as the decoder does. + // SAFETY: a live `Ffms2Decoder` holds a non-null video source, and the track belongs to that + // source, so it stays valid for as long as the decoder does. let track = unsafe { FFMS_GetTrackFromVideo(source) }; if track.is_null() { @@ -36,24 +33,21 @@ pub fn read_index(decoder: &mut Decoder) -> Option> { return None; } - // `FFMS_GetFrameInfo` indexes the track's entries without checking the - // bound, so the track's own count is what the read has to stay under. - // The video properties describe the same track but are reported - // separately, and a disagreement must not turn into a read past the end. + // `FFMS_GetFrameInfo` does not check its bound, so the read stays under the track's own count + // in case the separately reported video properties disagree with it. let total = total.min(track_frames as usize); let mut index = Vec::with_capacity(total); for i in 0..total { - // SAFETY: `track` is non-null, and `i` stays below the entry - // count the track itself reports. - let info = unsafe { FFMS_GetFrameInfo(track, i as i32) }; + // SAFETY: `track` is non-null, and `i` stays below the entry count the track reports. + let info_ptr = unsafe { FFMS_GetFrameInfo(track, i as i32) }; - if info.is_null() { + if info_ptr.is_null() { return None; } // SAFETY: checked non-null just above. - let info = unsafe { &*info }; + let info = unsafe { &*info_ptr }; index.push(IndexEntry { pts: info.PTS, @@ -72,23 +66,22 @@ pub fn phantom_indices(index: &[IndexEntry]) -> BTreeSet { return phantom; } - // Nothing before the first keyframe can be decoded, so ffms2 - // answers those positions with a repeat of the keyframe. - let lead = index.iter().position(|e| e.keyframe).unwrap_or(0); + // Nothing before the first keyframe can be decoded, so ffms2 answers those positions with a + // repeat of the keyframe. + let first_keyframe = index.iter().position(|entry| entry.keyframe).unwrap_or(0); - // One leading picture is the ordinary case. More than that means the - // index marks no keyframe for a stretch of the file, so say how much - // is going rather than shortening the output quietly. - if lead > 1 { + // One leading picture is the ordinary case. More means the index marks no keyframe for a + // stretch of the file, so the drop is reported rather than shortening the output quietly. + if first_keyframe > 1 { tracing::warn!( - dropped = lead, - "the index marks no keyframe until entry {lead}, dropping every entry before it", + dropped = first_keyframe, + "the index marks no keyframe until entry {first_keyframe}, dropping every entry before it", ); } - phantom.extend(0..lead); + phantom.extend(0..first_keyframe); - let gaps: Vec = (lead + 1..index.len()) + let gaps: Vec = (first_keyframe + 1..index.len()) .map(|i| index[i].pts.saturating_sub(index[i - 1].pts)) .collect(); @@ -96,24 +89,23 @@ pub fn phantom_indices(index: &[IndexEntry]) -> BTreeSet { return phantom; } - // The median gap is the clip's real frame spacing. A handful of - // phantom entries cannot move it, however far apart they sit. - let mut sorted = gaps.clone(); - sorted.sort_unstable(); - let median = sorted[sorted.len() / 2]; + // The median gap is the clip's real frame spacing. A handful of phantom entries cannot move + // it, however far apart they sit. + let mut sorted_gaps = gaps.clone(); + sorted_gaps.sort_unstable(); + let median = sorted_gaps[sorted_gaps.len() / 2]; - // Variable frame rate pacing puts real frames closer together than the - // median, which is the same signature a phantom leaves. Telling the two - // apart by timing only works on a clip that is otherwise regular, so a - // clip that is not keeps every entry. - let regular = gaps + // Variable frame rate pacing puts real frames closer together than the median, which is the + // same signature a phantom leaves. Timing only tells them apart on an otherwise regular clip, + // so an irregular clip keeps every entry. + let regular_gaps = gaps .iter() .filter(|&&gap| gap.saturating_sub(median).saturating_abs().saturating_mul(4) <= median) .count(); - if regular * 10 < gaps.len() * 9 { + if regular_gaps * 10 < gaps.len() * 9 { tracing::debug!( - regular, + regular = regular_gaps, gaps = gaps.len(), "frame spacing is too irregular to tell phantom entries from variable frame rate pacing", ); @@ -121,14 +113,12 @@ pub fn phantom_indices(index: &[IndexEntry]) -> BTreeSet { return phantom; } - // A phantom shares a timeline slot with the frame that follows it, landing just ahead of - // that frame's timestamp. The entry to drop is therefore the earlier of the pair. - // Doubling the gap rather than halving the median keeps the comparison exact, - // which matters when a clip's time base is close enough to its frame rate that the median - // gap is a single unit. + // A phantom lands just ahead of the frame that follows it, so the earlier entry of a + // too-close pair is dropped. Doubling the gap rather than halving the median keeps the + // comparison exact when the median gap is a single time base unit. for (offset, &gap) in gaps.iter().enumerate() { if gap.saturating_mul(2) < median { - phantom.insert(lead + offset); + phantom.insert(first_keyframe + offset); } } @@ -139,7 +129,7 @@ pub fn phantom_indices(index: &[IndexEntry]) -> BTreeSet { mod tests { use super::*; - /// Builds an index at a steady 42-unit spacing, with the first entry marked as the keyframe. + /// Builds an index at a steady 42-unit spacing with the first entry as the keyframe. fn regular(count: usize) -> Vec { (0..count) .map(|i| IndexEntry { @@ -151,12 +141,17 @@ mod tests { #[test] fn a_regular_index_has_no_phantoms() { - assert!(phantom_indices(®ular(20)).is_empty()); + let index = regular(20); + let phantom = phantom_indices(&index); + + assert!(phantom.is_empty()); } #[test] fn an_empty_index_has_no_phantoms() { - assert!(phantom_indices(&[]).is_empty()); + let phantom = phantom_indices(&[]); + + assert!(phantom.is_empty()); } #[test] @@ -165,12 +160,15 @@ mod tests { pts: 0, keyframe: false, }]; - index.extend((1..20).map(|i| IndexEntry { + let following = (1..20).map(|i| IndexEntry { pts: 41 + (i as i64 - 1) * 42, keyframe: i == 1, - })); + }); + index.extend(following); - assert_eq!(phantom_indices(&index), BTreeSet::from([0])); + let phantom = phantom_indices(&index); + + assert_eq!(phantom, BTreeSet::from([0])); } #[test] @@ -197,18 +195,21 @@ mod tests { keyframe: false, }, ]; - index.extend((5..60).map(|i| IndexEntry { + let following = (5..60).map(|i| IndexEntry { pts: 125 + (i as i64 - 5) * 42, keyframe: false, - })); + }); + index.extend(following); - assert_eq!(phantom_indices(&index), BTreeSet::from([1, 3])); + let phantom = phantom_indices(&index); + + assert_eq!(phantom, BTreeSet::from([1, 3])); } #[test] fn a_time_base_as_tight_as_the_frame_rate_still_finds_a_phantom() { - // Every real frame is one unit apart, so a phantom shows up as a - // repeated timestamp rather than as a fraction of a wider gap. + // Every real frame is one unit apart, so a phantom shows up as a repeated timestamp rather + // than as a fraction of a wider gap. let mut index: Vec = (0..40) .map(|i| IndexEntry { pts: i as i64, @@ -218,13 +219,15 @@ mod tests { index[4].pts = index[3].pts; - assert_eq!(phantom_indices(&index), BTreeSet::from([3])); + let phantom = phantom_indices(&index); + + assert_eq!(phantom, BTreeSet::from([3])); } #[test] fn variable_frame_rate_pacing_keeps_every_entry() { - // A third of these frames arrive at a third of the median spacing. - // They are real, and the timing rule cannot tell them from phantoms. + // A third of these frames arrive at a third of the median spacing. They are real, and the + // timing rule cannot tell them from phantoms. let mut pts = 0; let index: Vec = (0..30) .map(|i| { @@ -238,11 +241,16 @@ mod tests { }) .collect(); - assert!(phantom_indices(&index).is_empty()); + let phantom = phantom_indices(&index); + + assert!(phantom.is_empty()); } #[test] fn repeated_pictures_at_regular_spacing_are_kept() { - assert!(phantom_indices(®ular(2159)).is_empty()); + let index = regular(2159); + let phantom = phantom_indices(&index); + + assert!(phantom.is_empty()); } } diff --git a/av-denoise/src/bin/main.rs b/av-denoise/src/bin/main.rs index f70830f..9a25009 100644 --- a/av-denoise/src/bin/main.rs +++ b/av-denoise/src/bin/main.rs @@ -1,5 +1,4 @@ -use clap::Parser; -use tracing_subscriber::EnvFilter; +//! The `av-denoise` command line tool. mod cli; mod frame_index; @@ -8,7 +7,10 @@ mod progress; mod warm_start; mod y4m_format; -use cli::{Args, Command, InputSource, RunOptions, run_list_devices}; +use clap::Parser; +use tracing_subscriber::EnvFilter; + +use self::cli::{Args, Command, InputSource, RunOptions, run_list_devices}; #[global_allocator] static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc; @@ -16,12 +18,10 @@ static GLOBAL: mimalloc::MiMalloc = mimalloc::MiMalloc; /// Scene workers used when `--workers` is not given. const DEFAULT_WORKERS: usize = 2; -/// Frame budget in bytes used when `--frame-budget` is not given. (1GB) +/// Frame budget in bytes used when `--frame-budget` is not given, 1 GiB. const DEFAULT_FRAME_BUDGET_BYTES: u64 = 1 << 30; -/// Routes an input to the scene-parallel pipeline. -/// -/// A path opens with ffms2. A pipe reads a y4m stream. +/// Runs the scene-parallel pipeline on an input, with defaults for unset worker and budget flags. fn run_input( opts: &RunOptions, input: &InputSource, @@ -38,43 +38,48 @@ fn run_input( fn main() -> anyhow::Result<()> { // SAFETY: still single-threaded, no other thread can race the env mutation. - unsafe { av_denoise_core::raise_codegen_stack_limit() }; + unsafe { av_denoise::raise_codegen_stack_limit() }; let args = Args::parse(); if std::env::var("RUST_LOG").is_err() { - // `list-devices` prints a table and nothing else, and the - // backends chatter at info level while they start up, so it - // starts quieter than a denoising run. - let default = match args.command { + // `list-devices` prints only a table, and the backends log at info level while they + // start, so it runs quieter than a denoising run. + let default_filter = match args.command { Command::ListDevices => "warn", _ => "info", }; - unsafe { std::env::set_var("RUST_LOG", default) }; + + // SAFETY: still single-threaded, no other thread can race the env mutation. + unsafe { std::env::set_var("RUST_LOG", default_filter) }; } + let env_filter = EnvFilter::from_default_env(); + let log_writer = progress::tracing_writer(); tracing_subscriber::fmt() - .with_env_filter(EnvFilter::from_default_env()) - .with_writer(progress::tracing_writer()) + .with_env_filter(env_filter) + .with_writer(log_writer) .init(); - // Listing devices compiles no kernels, so it runs before the cache - // is installed and skips it entirely. + // Listing devices compiles no kernels, so it skips the kernel cache. if matches!(args.command, Command::ListDevices) { - print!("{}", run_list_devices(&args.accelerators)); + let table = run_list_devices(&args.accelerators); + print!("{table}"); return Ok(()); } - // Point CubeCL at a kernel cache. This has to run before - // Denoiser::create, because the first CubeCL client locks the global - // config the moment it is built. + // The first CubeCL client locks the global config, so the cache must be installed before + // any denoiser is created. match av_denoise::install_compilation_cache() { Ok(Some(path)) => tracing::info!(?path, "caching compiled kernels"), Ok(None) => tracing::info!( "kernel caching is off, every run recompiles. Unset {} to turn it back on.", av_denoise::COMPILATION_CACHE_ENV, ), - Err(err) => return Err(anyhow::Error::new(err).context("unable to install the kernel cache")), + Err(error) => { + let error = anyhow::Error::new(error).context("unable to install the kernel cache"); + return Err(error); + }, } let (opts, input, workers, frame_budget) = match &args.command { diff --git a/av-denoise/src/bin/pipeline/convert.rs b/av-denoise/src/bin/pipeline/convert.rs index 4d9dc62..47511d3 100644 --- a/av-denoise/src/bin/pipeline/convert.rs +++ b/av-denoise/src/bin/pipeline/convert.rs @@ -1,8 +1,10 @@ use std::sync::Arc; use av_denoise::{FrameLayout, Planes, Subsampling}; +use v_frame::chroma::ChromaSubsampling; use v_frame::frame::Frame; use v_frame::pixel::Pixel; +use v_frame::plane::Plane; /// A decoded frame at whichever sample width the source uses. pub enum DecodedFrame { @@ -54,97 +56,100 @@ impl SourcePixel for u16 { } } -/// Checks each plane's byte length against the layout, failing with an error naming which plane -/// is wrong, the length found and the length expected. +/// Checks each plane's byte length against the layout. +/// +/// The error names the wrong plane, the length found and the length expected. pub fn check_plane_lens(planes: &Planes, layout: FrameLayout) -> Result<(), anyhow::Error> { - for (name, got, expected) in [ + for (plane_name, actual, expected) in [ ("y", planes.y.len(), layout.luma_bytes()), ("u", planes.u.len(), layout.chroma_bytes()), ("v", planes.v.len(), layout.chroma_bytes()), ] { - if got != expected { - anyhow::bail!("{name} plane is {got} bytes, expected {expected} from the frame layout"); + if actual != expected { + anyhow::bail!("{plane_name} plane is {actual} bytes, expected {expected} from the frame layout"); } } Ok(()) } -pub fn planes_from_v_frame_u8( - frame: &v_frame::frame::Frame, - layout: FrameLayout, -) -> Result { - let y = collect_plane_u8(&frame.y_plane); - let u = frame +pub fn planes_from_v_frame_u8(frame: &Frame, layout: FrameLayout) -> Result { + let y_plane = collect_plane_u8(&frame.y_plane); + let u_plane = frame .u_plane .as_ref() .map(collect_plane_u8) .unwrap_or_else(|| layout.neutral_chroma_plane()); - let v = frame + let v_plane = frame .v_plane .as_ref() .map(collect_plane_u8) .unwrap_or_else(|| layout.neutral_chroma_plane()); - let planes = Planes { y, u, v }; + let planes = Planes { + y: y_plane, + u: u_plane, + v: v_plane, + }; check_plane_lens(&planes, layout)?; Ok(planes) } -pub fn planes_from_v_frame_u16( - frame: &v_frame::frame::Frame, - layout: FrameLayout, -) -> Result { - let y = collect_plane_u16(&frame.y_plane); - let u = frame +pub fn planes_from_v_frame_u16(frame: &Frame, layout: FrameLayout) -> Result { + let y_plane = collect_plane_u16(&frame.y_plane); + let u_plane = frame .u_plane .as_ref() .map(collect_plane_u16) .unwrap_or_else(|| layout.neutral_chroma_plane()); - let v = frame + let v_plane = frame .v_plane .as_ref() .map(collect_plane_u16) .unwrap_or_else(|| layout.neutral_chroma_plane()); - let planes = Planes { y, u, v }; + let planes = Planes { + y: y_plane, + u: u_plane, + v: v_plane, + }; check_plane_lens(&planes, layout)?; Ok(planes) } -pub fn collect_plane_u8(plane: &v_frame::plane::Plane) -> Vec { +pub fn collect_plane_u8(plane: &Plane) -> Vec { let width = plane.width().get(); let height = plane.height().get(); - let mut out = Vec::with_capacity(width * height); + let mut bytes = Vec::with_capacity(width * height); for row in plane.rows() { - out.extend_from_slice(row); + bytes.extend_from_slice(row); } - out + bytes } -pub fn collect_plane_u16(plane: &v_frame::plane::Plane) -> Vec { +/// Serialises a plane to little-endian bytes, two per sample. +pub fn collect_plane_u16(plane: &Plane) -> Vec { let width = plane.width().get(); let height = plane.height().get(); - let mut out = Vec::with_capacity(width * height * 2); + let mut bytes = Vec::with_capacity(width * height * 2); for row in plane.rows() { for &sample in row { - out.extend_from_slice(&sample.to_le_bytes()); + let sample_bytes = sample.to_le_bytes(); + bytes.extend_from_slice(&sample_bytes); } } - out + bytes } pub fn subsampling_from_av_decoders( - chroma_subsampling: v_frame::chroma::ChromaSubsampling, + chroma_subsampling: ChromaSubsampling, ) -> Result { - use v_frame::chroma::ChromaSubsampling; - match chroma_subsampling { ChromaSubsampling::Yuv420 => Ok(Subsampling::Yuv420), ChromaSubsampling::Yuv422 => Ok(Subsampling::Yuv422), diff --git a/av-denoise/src/bin/pipeline/coordinator.rs b/av-denoise/src/bin/pipeline/coordinator.rs index f2b2d72..f61c1b9 100644 --- a/av-denoise/src/bin/pipeline/coordinator.rs +++ b/av-denoise/src/bin/pipeline/coordinator.rs @@ -10,6 +10,7 @@ use super::source::SourceInfo; use crate::progress::{self, denoise_progress_bar}; use crate::y4m_format::subsampling_to_y4m; +/// A denoised frame tagged with its position in the output. pub struct OutputMsg { pub global_idx: u64, pub planes: Planes, @@ -17,18 +18,19 @@ pub struct OutputMsg { pub fn spawn_coordinator( info: SourceInfo, - rx: crossbeam_channel::Receiver, + outputs: crossbeam_channel::Receiver, staged: crossbeam_channel::Receiver, visible: bool, permits: crossbeam_channel::Sender<()>, output: W, ) -> thread::JoinHandle> { - thread::spawn(move || run_coordinator(info, rx, staged, visible, permits, output)) + thread::spawn(move || run_coordinator(info, outputs, staged, visible, permits, output)) } +/// Writes the y4m header, then every frame the workers send, in order. pub fn run_coordinator( info: SourceInfo, - rx: crossbeam_channel::Receiver, + outputs: crossbeam_channel::Receiver, staged: crossbeam_channel::Receiver, visible: bool, permits: crossbeam_channel::Sender<()>, @@ -45,69 +47,65 @@ pub fn run_coordinator( builder = builder.with_pixel_aspect(pixel_aspect); } - // Forwards the source's `X` params, `XCOLORRANGE=` being the common one. for extension in info.vendor_extensions { builder = builder.append_vendor_extension(extension); } let mut encoder = builder.write_header(output)?; - // Counts frames written to the output, which lags the frames read by - // the depth of the worker pipelines. Emitted frames are the honest - // measure of progress, because the count stalls whenever whatever - // consumes our stdout stops reading. - let pb = denoise_progress_bar(info.estimated_frames, visible); + // Counts frames written rather than read, which lags by the depth of the worker pipelines. + // Written frames are the honest measure because the count stalls whenever the consumer of + // stdout stops reading. + let progress_bar = denoise_progress_bar(info.estimated_frames, visible); - // The first frame only lands once a worker has compiled its - // kernels, which takes seconds. A steady tick draws the bar right - // away and keeps its elapsed time moving until then. - pb.enable_steady_tick(Duration::from_millis(250)); + // The first frame only lands once a worker has compiled its kernels, which takes seconds. A + // steady tick draws the bar right away and keeps its elapsed time moving until then. + progress_bar.enable_steady_tick(Duration::from_millis(250)); - let result = emit_frames(&mut encoder, &rx, &staged, &pb, &permits); + let result = emit_frames(&mut encoder, &outputs, &staged, &progress_bar, &permits); - progress::finish(&pb); + progress::finish(&progress_bar); result } -/// Reorders worker output by frame index and writes it out, updating `pb` as frames land. +/// Reorders worker output by frame index and writes it out, updating `progress_bar` as frames land. /// /// Runs until every worker has hung up, then checks the frames written against the count the /// dispatcher staged. No count means the dispatcher failed and reports its own error. pub fn emit_frames( encoder: &mut y4m::Encoder, - rx: &crossbeam_channel::Receiver, + outputs: &crossbeam_channel::Receiver, staged: &crossbeam_channel::Receiver, - pb: &ProgressBar, + progress_bar: &ProgressBar, permits: &crossbeam_channel::Sender<()>, ) -> Result<(), anyhow::Error> { let mut pending: BTreeMap = BTreeMap::new(); let mut next_emit: u64 = 0; - while let Ok(msg) = rx.recv() { - pending.insert(msg.global_idx, msg.planes); + while let Ok(message) = outputs.recv() { + pending.insert(message.global_idx, message.planes); while let Some(planes) = pending.remove(&next_emit) { let frame = Y4mFrame::new([&planes.y, &planes.u, &planes.v], None); encoder.write_frame(&frame)?; next_emit += 1; - // Returning the permit is what lets the decoder run further - // ahead. The send never blocks, because permits held plus - // permits waiting is always the channel's capacity. The - // result is discarded because it fails once the dispatcher - // has already errored out and dropped its receiver. + // Returning the permit lets the decoder run further ahead. The send never blocks + // because permits held plus permits waiting always equal the channel's capacity. + // Its result is discarded because it fails once the dispatcher has errored out and + // dropped its receiver. let _ = permits.send(()); } - pb.set_position(next_emit); + progress_bar.set_position(next_emit); } let Ok(total) = staged.recv() else { return Ok(()); }; - pb.set_length(total); + progress_bar.set_length(total); if next_emit != total { anyhow::bail!( diff --git a/av-denoise/src/bin/pipeline/decode.rs b/av-denoise/src/bin/pipeline/decode.rs index 21397a5..7512bc4 100644 --- a/av-denoise/src/bin/pipeline/decode.rs +++ b/av-denoise/src/bin/pipeline/decode.rs @@ -105,6 +105,7 @@ fn run_decode_thread( phantom, info, } = opened; + let result = match info.layout.depth { Depth::Eight => pump_decoder::(decoder, &phantom, &start), Depth::Ten | Depth::Twelve => pump_decoder::(decoder, &phantom, &start), @@ -130,7 +131,10 @@ fn read_frame(decoder: &mut Decoder) -> Option, match decoder.read_video_frame::() { Ok(frame) => Some(Ok(frame)), Err(DecoderError::EndOfFile) => None, - Err(err) => Some(Err(err.into())), + Err(err) => { + let error = anyhow::Error::from(err); + Some(Err(error)) + }, } } @@ -141,7 +145,7 @@ pub fn pump_frames( frames: I, phantom: &BTreeSet, permits: &crossbeam_channel::Receiver<()>, - out: &crossbeam_channel::Sender, + decoded_tx: &crossbeam_channel::Sender, ) -> Result<(), anyhow::Error> where T: SourcePixel, @@ -162,7 +166,7 @@ where let shared = Arc::new(frame); let decoded = T::into_decoded(shared); - if out.send(Ok(decoded)).is_err() { + if decoded_tx.send(Ok(decoded)).is_err() { return Ok(()); } } diff --git a/av-denoise/src/bin/pipeline/grain_table.rs b/av-denoise/src/bin/pipeline/grain_table.rs index 06ac201..2961eea 100644 --- a/av-denoise/src/bin/pipeline/grain_table.rs +++ b/av-denoise/src/bin/pipeline/grain_table.rs @@ -17,9 +17,9 @@ pub struct TablePath { impl TablePath { /// Creates the temporary file beside `path`. pub fn create(path: &Path) -> Result { - let mut temporary = path.as_os_str().to_owned(); - temporary.push(".tmp"); - let temporary = PathBuf::from(temporary); + let mut temporary_name = path.as_os_str().to_owned(); + temporary_name.push(".tmp"); + let temporary = PathBuf::from(temporary_name); File::create(&temporary) .with_context(|| format!("unable to create the grain table at {}", path.display()))?; @@ -48,6 +48,7 @@ impl TablePath { .with_context(|| format!("unable to move the grain table to {}", self.path.display()))?; self.written = true; + Ok(()) } } diff --git a/av-denoise/src/bin/pipeline/mod.rs b/av-denoise/src/bin/pipeline/mod.rs index 1bd98e7..75f966c 100644 --- a/av-denoise/src/bin/pipeline/mod.rs +++ b/av-denoise/src/bin/pipeline/mod.rs @@ -37,6 +37,8 @@ pub fn run( let visible = denoise_bar_visible(opts.progress, is_terminal); let owned_input = input.clone(); let opener = move || source::open_source(&owned_input); + let output = stdout(); + let grain_table = opts.grain_table.clone(); run_with( &opts.planes, @@ -44,8 +46,8 @@ pub fn run( workers, frame_budget_bytes, visible, - stdout(), - opts.grain_table.clone(), + output, + grain_table, ) } @@ -101,8 +103,9 @@ where ); let (staged_tx, staged_rx) = crossbeam_channel::bounded::(1); - let (job_tx, worker_handles, out_rx) = spawn_workers(planes, layout, workers); - let coordinator = spawn_coordinator(info.clone(), out_rx, staged_rx, visible, give, output); + let (job_tx, worker_handles, output_rx) = spawn_workers(planes, layout, workers); + let coordinator_info = info.clone(); + let coordinator = spawn_coordinator(coordinator_info, output_rx, staged_rx, visible, give, output); let frames = decode_thread.start(take); let dispatched = match layout.depth { diff --git a/av-denoise/src/bin/pipeline/scenes.rs b/av-denoise/src/bin/pipeline/scenes.rs index 7afd46b..febf8d6 100644 --- a/av-denoise/src/bin/pipeline/scenes.rs +++ b/av-denoise/src/bin/pipeline/scenes.rs @@ -79,6 +79,7 @@ impl SceneSplitter { } self.queue.clear(); + released } diff --git a/av-denoise/src/bin/pipeline/source.rs b/av-denoise/src/bin/pipeline/source.rs index d1b2390..a57a41f 100644 --- a/av-denoise/src/bin/pipeline/source.rs +++ b/av-denoise/src/bin/pipeline/source.rs @@ -64,6 +64,7 @@ pub fn open_file(path: &Path) -> Result { let colour = read_file_colour(&mut decoder); let vendor_extensions: Vec = colour.range.into_iter().collect(); + log_forwarded_colour(colour.pixel_aspect, &vendor_extensions); let info = SourceInfo { @@ -87,9 +88,10 @@ fn log_forwarded_colour(pixel_aspect: Option, vendor_extensions: &[y } let aspect = pixel_aspect.map(|ratio| (ratio.num, ratio.den)); - let range = vendor_extensions - .first() - .map(|extension| String::from_utf8_lossy(extension.value()).into_owned()); + let range = vendor_extensions.first().map(|extension| { + let value = extension.value(); + String::from_utf8_lossy(value).into_owned() + }); tracing::info!( pixel_aspect = ?aspect, @@ -115,7 +117,9 @@ pub fn color_range_extension(range: i32) -> Option { _ => return None, }; - y4m::VendorExtensionString::new(tag.to_vec()).ok() + let value = tag.to_vec(); + + y4m::VendorExtensionString::new(value).ok() } /// Turns an ffms2 sample aspect ratio into a y4m pixel aspect. @@ -146,14 +150,14 @@ pub fn read_file_colour(decoder: &mut Decoder) -> FileColour { // SAFETY: a live `Ffms2Decoder` holds a non-null video source, and the properties it // returns belong to that source, so they stay valid while the decoder does. - let properties = unsafe { FFMS_GetVideoProperties(ffms2.video_source) }; + let properties_ptr = unsafe { FFMS_GetVideoProperties(ffms2.video_source) }; - if properties.is_null() { + if properties_ptr.is_null() { return empty; } // SAFETY: checked non-null just above. - let properties = unsafe { &*properties }; + let properties = unsafe { &*properties_ptr }; let range = color_range_extension(properties.ColorRange); let pixel_aspect = pixel_aspect_from_sar(properties.SARNum, properties.SARDen); diff --git a/av-denoise/src/bin/pipeline/stage.rs b/av-denoise/src/bin/pipeline/stage.rs index 4c708b1..78c6cdb 100644 --- a/av-denoise/src/bin/pipeline/stage.rs +++ b/av-denoise/src/bin/pipeline/stage.rs @@ -11,14 +11,14 @@ pub const IN_TRANSIT_FRAMES: usize = LOOKAHEAD_DISTANCE + 2 + PREFETCH_FRAMES + /// Frames in flight this run allows. pub fn frame_permits(budget_bytes: u64, frame_bytes: usize, workers: usize, radius: u32) -> usize { - // A worker emits nothing until `push` first returns QueueFull, which - // takes `radius + MAX_PENDING + 1` pushes, and the frames upstream of - // the workers hold permits too. Fewer permits than that and the - // dispatcher waits on a permit nothing can release. + // A worker emits nothing until `push` first returns QueueFull, which takes + // `radius + MAX_PENDING + 1` pushes, and the frames upstream of the workers hold permits too. + // Fewer permits than that and the dispatcher waits on a permit nothing can release. let per_worker = radius as usize + av_denoise::MAX_PENDING + 2; let floor = workers * per_worker + IN_TRANSIT_FRAMES; + let afforded = frames_afforded(budget_bytes, frame_bytes); - frames_afforded(budget_bytes, frame_bytes).max(floor) + afforded.max(floor) } /// Frames the budget pays for at this frame size. @@ -30,8 +30,8 @@ pub fn frames_afforded(budget_bytes: u64, frame_bytes: usize) -> usize { /// Renders a byte count in the decimal units `--frame-budget` accepts. /// -/// The space `ByteSize` puts before the unit goes, so the result is one -/// shell argument a caller can paste straight back into the flag. +/// The space `ByteSize` puts before the unit is removed, so the result is one shell argument that +/// pastes straight back into the flag. pub fn size_string(bytes: u64) -> String { bytesize::ByteSize::b(bytes) .display() @@ -42,22 +42,23 @@ pub fn size_string(bytes: u64) -> String { /// Rounds a byte count up to the precision [size_string] prints at. /// -/// The rendering keeps one decimal place, so a raw minimum can round -/// down to a size that still fails the budget check. +/// The rendering keeps one decimal place, so a raw minimum can round down to a size that still +/// fails the budget check. pub fn suggested_budget(bytes: u64) -> String { let unit = [bytesize::GB, bytesize::MB, bytesize::KB] .into_iter() .find(|&unit| bytes >= unit) .unwrap_or(1); let step = (unit / 10).max(1); + let rounded = bytes.div_ceil(step) * step; - size_string(bytes.div_ceil(step) * step) + size_string(rounded) } /// Frames in flight this run allows, refusing a budget below the floor. /// -/// A budget the floor has to raise serialises the pipeline, so it fails -/// here rather than running on with too few frames in flight. +/// A budget the floor has to raise serialises the pipeline, so it fails here rather than running +/// on with too few frames in flight. pub fn checked_frame_permits( budget_bytes: u64, frame_bytes: usize, @@ -68,13 +69,14 @@ pub fn checked_frame_permits( let permits = frame_permits(budget_bytes, frame_bytes, workers, radius); if permits > afforded { - let suggestion = suggested_budget(permits as u64 * frame_bytes as u64); + let required_bytes = permits as u64 * frame_bytes as u64; + let suggestion = suggested_budget(required_bytes); + let budget = size_string(budget_bytes); anyhow::bail!( "--frame-budget {budget} affords {afforded} frames at {frame_bytes} bytes per frame, \ but {workers} workers at temporal radius {radius} need at least {permits}. Pass at \ least --frame-budget {suggestion}.", - budget = size_string(budget_bytes), ); } @@ -83,9 +85,8 @@ pub fn checked_frame_permits( /// Builds a counting semaphore holding `count` permits. /// -/// Returns the giving end and the taking end. The coordinator holds the -/// giving end, so if it dies the dispatcher's next take fails instead of -/// blocking forever. +/// Returns the giving end and the taking end. The coordinator holds the giving end, so if it dies +/// the dispatcher's next take fails instead of blocking forever. pub fn frame_permit_channel( count: usize, ) -> (crossbeam_channel::Sender<()>, crossbeam_channel::Receiver<()>) { @@ -180,8 +181,7 @@ pub struct StagedFrame { /// One scene, offered to whichever worker is free. /// -/// `frames` closes when the scene has no more frames, which is how a -/// worker knows to flush. +/// `frames` closes when the scene has no more frames, which is how a worker knows to flush. pub struct SceneJob { pub scene_idx: u32, pub frames: crossbeam_channel::Receiver, diff --git a/av-denoise/src/bin/pipeline/tests/convert.rs b/av-denoise/src/bin/pipeline/tests/convert.rs index 53e8d81..ad4a630 100644 --- a/av-denoise/src/bin/pipeline/tests/convert.rs +++ b/av-denoise/src/bin/pipeline/tests/convert.rs @@ -1,24 +1,25 @@ +use std::num::{NonZeroU8, NonZeroUsize}; +use std::sync::Arc; + use av_denoise::{Depth, FrameLayout, Subsampling}; +use v_frame::chroma::ChromaSubsampling; +use v_frame::frame::{Frame, FrameBuilder}; -use crate::pipeline::convert::{collect_plane_u16, planes_from_v_frame_u8, planes_from_v_frame_u16}; +use crate::pipeline::convert::{ + SourcePixel, + collect_plane_u16, + planes_from_v_frame_u8, + planes_from_v_frame_u16, +}; -/// A 10-bit v_frame plane serialises to little-endian wire bytes at -/// twice the sample count. #[test] fn collect_plane_u16_writes_little_endian_bytes() { - use std::num::{NonZeroU8, NonZeroUsize}; - - use v_frame::chroma::ChromaSubsampling; - use v_frame::frame::{Frame, FrameBuilder}; - - let mut frame: Frame = FrameBuilder::new( - NonZeroUsize::new(2).expect("width is non-zero"), - NonZeroUsize::new(2).expect("height is non-zero"), - ChromaSubsampling::Yuv420, - NonZeroU8::new(10).expect("depth is non-zero"), - ) - .build() - .expect("a 2x2 10-bit frame builds"); + let width = NonZeroUsize::new(2).expect("width is non-zero"); + let height = NonZeroUsize::new(2).expect("height is non-zero"); + let bit_depth = NonZeroU8::new(10).expect("depth is non-zero"); + let mut frame: Frame = FrameBuilder::new(width, height, ChromaSubsampling::Yuv420, bit_depth) + .build() + .expect("a 2x2 10-bit frame builds"); frame .y_plane @@ -37,25 +38,18 @@ fn collect_plane_u16_writes_little_endian_bytes() { #[test] fn planes_from_v_frame_u8_matching_layout_succeeds() { - use std::num::{NonZeroU8, NonZeroUsize}; - - use v_frame::chroma::ChromaSubsampling; - use v_frame::frame::FrameBuilder; - let layout = FrameLayout { width: 2, height: 2, subsampling: Subsampling::Yuv420, depth: Depth::Eight, }; - let frame: v_frame::frame::Frame = FrameBuilder::new( - NonZeroUsize::new(2).expect("width is non-zero"), - NonZeroUsize::new(2).expect("height is non-zero"), - ChromaSubsampling::Yuv420, - NonZeroU8::new(8).expect("depth is non-zero"), - ) - .build() - .expect("a 2x2 8-bit frame builds"); + let width = NonZeroUsize::new(2).expect("width is non-zero"); + let height = NonZeroUsize::new(2).expect("height is non-zero"); + let bit_depth = NonZeroU8::new(8).expect("depth is non-zero"); + let frame: Frame = FrameBuilder::new(width, height, ChromaSubsampling::Yuv420, bit_depth) + .build() + .expect("a 2x2 8-bit frame builds"); let planes = planes_from_v_frame_u8(&frame, layout).expect("matching layout should not error"); @@ -66,25 +60,18 @@ fn planes_from_v_frame_u8_matching_layout_succeeds() { #[test] fn planes_from_v_frame_u16_matching_layout_succeeds() { - use std::num::{NonZeroU8, NonZeroUsize}; - - use v_frame::chroma::ChromaSubsampling; - use v_frame::frame::FrameBuilder; - let layout = FrameLayout { width: 2, height: 2, subsampling: Subsampling::Yuv420, depth: Depth::Ten, }; - let frame: v_frame::frame::Frame = FrameBuilder::new( - NonZeroUsize::new(2).expect("width is non-zero"), - NonZeroUsize::new(2).expect("height is non-zero"), - ChromaSubsampling::Yuv420, - NonZeroU8::new(10).expect("depth is non-zero"), - ) - .build() - .expect("a 2x2 10-bit frame builds"); + let width = NonZeroUsize::new(2).expect("width is non-zero"); + let height = NonZeroUsize::new(2).expect("height is non-zero"); + let bit_depth = NonZeroU8::new(10).expect("depth is non-zero"); + let frame: Frame = FrameBuilder::new(width, height, ChromaSubsampling::Yuv420, bit_depth) + .build() + .expect("a 2x2 10-bit frame builds"); let planes = planes_from_v_frame_u16(&frame, layout).expect("matching layout should not error"); @@ -95,19 +82,12 @@ fn planes_from_v_frame_u16_matching_layout_succeeds() { #[test] fn planes_from_v_frame_u8_mismatched_layout_errors() { - use std::num::{NonZeroU8, NonZeroUsize}; - - use v_frame::chroma::ChromaSubsampling; - use v_frame::frame::FrameBuilder; - - let frame: v_frame::frame::Frame = FrameBuilder::new( - NonZeroUsize::new(2).expect("width is non-zero"), - NonZeroUsize::new(2).expect("height is non-zero"), - ChromaSubsampling::Yuv420, - NonZeroU8::new(8).expect("depth is non-zero"), - ) - .build() - .expect("a 2x2 8-bit frame builds"); + let width = NonZeroUsize::new(2).expect("width is non-zero"); + let height = NonZeroUsize::new(2).expect("height is non-zero"); + let bit_depth = NonZeroU8::new(8).expect("depth is non-zero"); + let frame: Frame = FrameBuilder::new(width, height, ChromaSubsampling::Yuv420, bit_depth) + .build() + .expect("a 2x2 8-bit frame builds"); let layout = FrameLayout { width: 4, @@ -117,38 +97,30 @@ fn planes_from_v_frame_u8_mismatched_layout_errors() { }; let err = planes_from_v_frame_u8(&frame, layout).expect_err("a smaller frame should not pass"); - let msg = err.to_string(); + let message = err.to_string(); - assert!(msg.contains('y'), "error should name the plane: {msg}"); + assert!(message.contains('y'), "error should name the plane: {message}"); assert!( - msg.contains('4'), - "error should name the 2x2 plane's length (4): {msg}" + message.contains('4'), + "error should name the 2x2 plane's length (4): {message}" ); assert!( - msg.contains("16"), - "error should name the layout's expected length (16): {msg}" + message.contains("16"), + "error should name the layout's expected length (16): {message}" ); } #[test] fn a_decoded_frame_only_unwraps_to_its_own_depth() { - use std::num::{NonZeroU8, NonZeroUsize}; - use std::sync::Arc; - - use v_frame::chroma::ChromaSubsampling; - use v_frame::frame::FrameBuilder; - - use crate::pipeline::convert::SourcePixel; - - let frame: v_frame::frame::Frame = FrameBuilder::new( - NonZeroUsize::new(2).expect("width is non-zero"), - NonZeroUsize::new(2).expect("height is non-zero"), - ChromaSubsampling::Yuv420, - NonZeroU8::new(8).expect("depth is non-zero"), - ) - .build() - .expect("a 2x2 8-bit frame builds"); - let decoded = u8::into_decoded(Arc::new(frame)); - - assert!(u16::from_decoded(decoded).is_none()); + let width = NonZeroUsize::new(2).expect("width is non-zero"); + let height = NonZeroUsize::new(2).expect("height is non-zero"); + let bit_depth = NonZeroU8::new(8).expect("depth is non-zero"); + let frame: Frame = FrameBuilder::new(width, height, ChromaSubsampling::Yuv420, bit_depth) + .build() + .expect("a 2x2 8-bit frame builds"); + let shared = Arc::new(frame); + let decoded = u8::into_decoded(shared); + let unwrapped = u16::from_decoded(decoded); + + assert!(unwrapped.is_none()); } diff --git a/av-denoise/src/bin/pipeline/tests/coordinator.rs b/av-denoise/src/bin/pipeline/tests/coordinator.rs index 4a9925d..d85512a 100644 --- a/av-denoise/src/bin/pipeline/tests/coordinator.rs +++ b/av-denoise/src/bin/pipeline/tests/coordinator.rs @@ -5,21 +5,19 @@ use crate::pipeline::coordinator::{OutputMsg, emit_frames}; use crate::pipeline::stage::frame_permit_channel; use crate::y4m_format::subsampling_to_y4m; -fn encoder(buf: &mut Vec) -> y4m::Encoder<&mut Vec> { +fn encoder(buffer: &mut Vec) -> y4m::Encoder<&mut Vec> { let layout = tiny_layout(); + let frame_rate = y4m::Ratio::new(30, 1); + let colorspace = subsampling_to_y4m(layout.subsampling, layout.depth); - y4m::encode( - layout.width as usize, - layout.height as usize, - y4m::Ratio::new(30, 1), - ) - .with_colorspace(subsampling_to_y4m(layout.subsampling, layout.depth)) - .write_header(buf) - .expect("header write failed") + y4m::encode(layout.width as usize, layout.height as usize, frame_rate) + .with_colorspace(colorspace) + .write_header(buffer) + .expect("header write failed") } fn send_frames(indices: &[u64]) -> crossbeam_channel::Receiver { - let (tx, rx) = crossbeam_channel::unbounded::(); + let (output_tx, output_rx) = crossbeam_channel::unbounded::(); let layout = tiny_layout(); let planes = tiny_planes(layout); @@ -28,27 +26,27 @@ fn send_frames(indices: &[u64]) -> crossbeam_channel::Receiver { global_idx, planes: planes.clone(), }; - tx.send(message).expect("the receiver is alive"); + output_tx.send(message).expect("the receiver is alive"); } - rx + output_rx } fn staged_count(count: Option) -> crossbeam_channel::Receiver { - let (tx, rx) = crossbeam_channel::bounded(1); + let (staged_tx, staged_rx) = crossbeam_channel::bounded(1); if let Some(count) = count { - tx.send(count).expect("the channel has room"); + staged_tx.send(count).expect("the channel has room"); } - rx + staged_rx } #[test] fn emit_frames_errors_when_a_frame_index_is_lost() { - let mut buf = Vec::new(); - let mut encoder = encoder(&mut buf); - let rx = send_frames(&[0, 2]); + let mut buffer = Vec::new(); + let mut encoder = encoder(&mut buffer); + let outputs = send_frames(&[0, 2]); let staged = staged_count(Some(3)); let (give, take) = frame_permit_channel(4); take.recv().expect("a permit is available"); @@ -56,21 +54,21 @@ fn emit_frames_errors_when_a_frame_index_is_lost() { let progress = ProgressBar::hidden(); - let err = - emit_frames(&mut encoder, &rx, &staged, &progress, &give).expect_err("expected a lost-frame error"); - let msg = err.to_string(); + let err = emit_frames(&mut encoder, &outputs, &staged, &progress, &give) + .expect_err("expected a lost-frame error"); + let message = err.to_string(); assert!( - msg.contains('1') && msg.contains('3'), - "error should name 1 written vs 3 staged: {msg}" + message.contains('1') && message.contains('3'), + "error should name 1 written vs 3 staged: {message}" ); } #[test] fn emit_frames_finishes_when_every_staged_frame_is_written() { - let mut buf = Vec::new(); - let mut encoder = encoder(&mut buf); - let rx = send_frames(&[1, 0, 2]); + let mut buffer = Vec::new(); + let mut encoder = encoder(&mut buffer); + let outputs = send_frames(&[1, 0, 2]); let staged = staged_count(Some(3)); let (give, take) = frame_permit_channel(4); @@ -80,22 +78,22 @@ fn emit_frames_finishes_when_every_staged_frame_is_written() { let progress = ProgressBar::hidden(); - emit_frames(&mut encoder, &rx, &staged, &progress, &give).expect("three of three frames written"); + emit_frames(&mut encoder, &outputs, &staged, &progress, &give).expect("three of three frames written"); assert_eq!(take.len(), 4, "every permit came back"); } #[test] fn emit_frames_leaves_the_error_to_a_dispatcher_that_sent_no_count() { - let mut buf = Vec::new(); - let mut encoder = encoder(&mut buf); - let rx = send_frames(&[0]); + let mut buffer = Vec::new(); + let mut encoder = encoder(&mut buffer); + let outputs = send_frames(&[0]); let staged = staged_count(None); let (give, take) = frame_permit_channel(2); take.recv().expect("a permit is available"); let progress = ProgressBar::hidden(); - emit_frames(&mut encoder, &rx, &staged, &progress, &give) + emit_frames(&mut encoder, &outputs, &staged, &progress, &give) .expect("the dispatcher reports its own failure"); } diff --git a/av-denoise/src/bin/pipeline/tests/decode.rs b/av-denoise/src/bin/pipeline/tests/decode.rs index 62b0c44..4349edd 100644 --- a/av-denoise/src/bin/pipeline/tests/decode.rs +++ b/av-denoise/src/bin/pipeline/tests/decode.rs @@ -1,85 +1,87 @@ use std::collections::BTreeSet; -use std::io::{Cursor, Read}; use std::num::{NonZeroU8, NonZeroUsize}; use v_frame::chroma::ChromaSubsampling; use v_frame::frame::{Frame, FrameBuilder}; -use super::y4m_clip; +use super::{y4m_clip, y4m_reader}; use crate::pipeline::convert::SourcePixel; use crate::pipeline::decode::{DecodeThread, FrameMsg, pump_frames}; use crate::pipeline::source::open_y4m; use crate::pipeline::stage::frame_permit_channel; fn tiny_frame() -> Frame { - FrameBuilder::new( - NonZeroUsize::new(2).expect("width is non-zero"), - NonZeroUsize::new(2).expect("height is non-zero"), - ChromaSubsampling::Yuv420, - NonZeroU8::new(8).expect("depth is non-zero"), - ) - .build() - .expect("a 2x2 8-bit frame builds") + let width = NonZeroUsize::new(2).expect("width is non-zero"); + let height = NonZeroUsize::new(2).expect("height is non-zero"); + let bit_depth = NonZeroU8::new(8).expect("depth is non-zero"); + + FrameBuilder::new(width, height, ChromaSubsampling::Yuv420, bit_depth) + .build() + .expect("a 2x2 8-bit frame builds") } fn frames(count: usize) -> impl Iterator, anyhow::Error>> { - (0..count).map(|_| Ok(tiny_frame())) + (0..count).map(|_| { + let frame = tiny_frame(); + Ok(frame) + }) } #[test] fn phantom_frames_are_skipped_and_take_no_permit() { let (_give, take) = frame_permit_channel(8); - let (out_tx, out_rx) = crossbeam_channel::unbounded::(); + let (decoded_tx, decoded_rx) = crossbeam_channel::unbounded::(); let phantom = BTreeSet::from([1, 3]); let decoded = frames(6); - pump_frames(decoded, &phantom, &take, &out_tx).expect("pumping should succeed"); - drop(out_tx); + pump_frames(decoded, &phantom, &take, &decoded_tx).expect("pumping should succeed"); + drop(decoded_tx); - assert_eq!(out_rx.iter().count(), 4); + assert_eq!(decoded_rx.iter().count(), 4); assert_eq!(take.len(), 4, "only the four sent frames hold a permit"); } #[test] fn a_read_error_stops_pumping_after_earlier_frames() { let (_give, take) = frame_permit_channel(8); - let (out_tx, out_rx) = crossbeam_channel::unbounded::(); - let corrupt = Err(anyhow::anyhow!("corrupt packet")); + let (decoded_tx, decoded_rx) = crossbeam_channel::unbounded::(); + let decode_error = anyhow::anyhow!("corrupt packet"); + let corrupt = Err(decode_error); let corrupt_frame = std::iter::once(corrupt); let leading = frames(2); let failing = leading.chain(corrupt_frame); let phantom = BTreeSet::new(); - let result = pump_frames(failing, &phantom, &take, &out_tx); + let result = pump_frames(failing, &phantom, &take, &decoded_tx); assert!(result.is_err()); - assert_eq!(out_rx.len(), 2, "frames before the error are still sent"); + assert_eq!(decoded_rx.len(), 2, "frames before the error are still sent"); } #[test] fn pumping_stops_quietly_when_the_splitter_hangs_up() { let (_give, take) = frame_permit_channel(8); - let (out_tx, out_rx) = crossbeam_channel::bounded::(1); - drop(out_rx); + let (decoded_tx, decoded_rx) = crossbeam_channel::bounded::(1); + drop(decoded_rx); let decoded = frames(4); let phantom = BTreeSet::new(); - pump_frames(decoded, &phantom, &take, &out_tx).expect("a closed channel is not an error"); + pump_frames(decoded, &phantom, &take, &decoded_tx).expect("a closed channel is not an error"); } #[test] fn pumping_fails_when_the_coordinator_drops_every_permit() { let (give, take) = frame_permit_channel(1); drop(give); - let (out_tx, _out_rx) = crossbeam_channel::unbounded::(); + let (decoded_tx, _decoded_rx) = crossbeam_channel::unbounded::(); let decoded = frames(4); let phantom = BTreeSet::new(); - let result = pump_frames(decoded, &phantom, &take, &out_tx); + let result = pump_frames(decoded, &phantom, &take, &decoded_tx); assert!(result.is_err(), "a second frame cannot get a permit"); } @@ -88,7 +90,7 @@ fn pumping_fails_when_the_coordinator_drops_every_permit() { fn the_thread_reports_the_source_then_streams_every_frame() { let bytes = y4m_clip(3); let (thread, info) = DecodeThread::spawn(move || { - let reader: Box = Box::new(Cursor::new(bytes)); + let reader = y4m_reader(bytes); open_y4m(reader) }) .expect("the clip opens"); @@ -105,13 +107,18 @@ fn the_thread_reports_the_source_then_streams_every_frame() { for message in received { let decoded = message.expect("every frame decodes"); - assert!(u8::from_decoded(decoded).is_some()); + let unwrapped = u8::from_decoded(decoded); + + assert!(unwrapped.is_some()); } } #[test] fn an_open_failure_is_returned_from_spawn() { - let result = DecodeThread::spawn(|| Err(anyhow::anyhow!("no such file"))); + let result = DecodeThread::spawn(|| { + let open_error = anyhow::anyhow!("no such file"); + Err(open_error) + }); let err = match result { Ok(_) => panic!("the open failed, so spawn must fail"), @@ -125,7 +132,7 @@ fn an_open_failure_is_returned_from_spawn() { fn dropping_an_unstarted_thread_lets_it_exit() { let bytes = y4m_clip(1); let (thread, _info) = DecodeThread::spawn(move || { - let reader: Box = Box::new(Cursor::new(bytes)); + let reader = y4m_reader(bytes); open_y4m(reader) }) .expect("the clip opens"); diff --git a/av-denoise/src/bin/pipeline/tests/dispatch.rs b/av-denoise/src/bin/pipeline/tests/dispatch.rs index 3f03294..8b70342 100644 --- a/av-denoise/src/bin/pipeline/tests/dispatch.rs +++ b/av-denoise/src/bin/pipeline/tests/dispatch.rs @@ -12,7 +12,8 @@ use crate::pipeline::stage::SceneJob; #[test] fn dispatch_fails_on_a_forwarded_decode_error() { let bytes = y4m_clip(3); - let reader: Box = Box::new(Cursor::new(bytes)); + let cursor = Cursor::new(bytes); + let reader: Box = Box::new(cursor); let mut opened = open_y4m(reader).expect("the clip opens"); let (frames_tx, frames_rx) = crossbeam_channel::unbounded::(); @@ -26,9 +27,8 @@ fn dispatch_fails_on_a_forwarded_decode_error() { frames_tx.send(Ok(decoded)).expect("the receiver is alive"); } - frames_tx - .send(Err(anyhow::anyhow!("corrupt packet"))) - .expect("the receiver is alive"); + let decode_error = anyhow::anyhow!("corrupt packet"); + frames_tx.send(Err(decode_error)).expect("the receiver is alive"); drop(frames_tx); let (job_tx, job_rx) = crossbeam_channel::bounded::(0); diff --git a/av-denoise/src/bin/pipeline/tests/grain_table.rs b/av-denoise/src/bin/pipeline/tests/grain_table.rs index f23b21d..3baebf3 100644 --- a/av-denoise/src/bin/pipeline/tests/grain_table.rs +++ b/av-denoise/src/bin/pipeline/tests/grain_table.rs @@ -54,15 +54,19 @@ fn nl4d_opts(grain_export: bool) -> PlaneOptions { fn run(bytes: Vec, workers: usize, table: Option) -> Result, anyhow::Error> { let output = SharedBuffer::default(); - let options = nl4d_opts(table.is_some()); + let writer = output.clone(); + let grain_export = table.is_some(); + let options = nl4d_opts(grain_export); let opener = move || { - let reader: Box = Box::new(Cursor::new(bytes)); + let cursor = Cursor::new(bytes); + let reader: Box = Box::new(cursor); open_y4m(reader) }; - run_with(&options, opener, workers, 1 << 30, false, output.clone(), table)?; + run_with(&options, opener, workers, 1 << 30, false, writer, table)?; let written = output.0.lock().expect("buffer lock").clone(); + Ok(written) } @@ -70,22 +74,27 @@ fn run(bytes: Vec, workers: usize, table: Option) -> Result fn a_run_with_the_flag_writes_a_table() { let dir = TestDir::new("writes"); let path = dir.path().join("grain.tbl"); + let clip = grainy_clip(40); - run(grainy_clip(40), 2, Some(path.clone())).expect("the run succeeds"); + run(clip, 2, Some(path.clone())).expect("the run succeeds"); let text = std::fs::read_to_string(&path).expect("the table exists"); + let temporary = dir.path().join("grain.tbl.tmp"); + assert!(text.starts_with("filmgrn1\n")); assert!(text.contains("\nE "), "no entry was fitted, got {text}"); - assert!(!dir.path().join("grain.tbl.tmp").exists()); + assert!(!temporary.exists()); } #[test] fn the_flag_does_not_change_the_output() { let dir = TestDir::new("output"); let path = dir.path().join("grain.tbl"); + let plain_clip = grainy_clip(40); + let export_clip = grainy_clip(40); - let without = run(grainy_clip(40), 1, None).expect("plain run"); - let with = run(grainy_clip(40), 1, Some(path)).expect("export run"); + let without = run(plain_clip, 1, None).expect("plain run"); + let with = run(export_clip, 1, Some(path)).expect("export run"); assert_eq!(without, with); } @@ -95,13 +104,16 @@ fn worker_count_does_not_change_the_table() { let dir = TestDir::new("workers"); let single = dir.path().join("single.tbl"); let multi = dir.path().join("multi.tbl"); + let single_clip = grainy_clip(60); + let multi_clip = grainy_clip(60); - run(grainy_clip(60), 1, Some(single.clone())).expect("single worker"); - run(grainy_clip(60), 3, Some(multi.clone())).expect("three workers"); + run(single_clip, 1, Some(single.clone())).expect("single worker"); + run(multi_clip, 3, Some(multi.clone())).expect("three workers"); let single_text = std::fs::read_to_string(single).expect("single table"); let multi_text = std::fs::read_to_string(multi).expect("multi table"); let entries = single_text.matches("\nE ").count(); + assert!(entries >= 2, "expected an entry per scene, got {single_text}"); assert_eq!(single_text, multi_text); } @@ -110,19 +122,22 @@ fn worker_count_does_not_change_the_table() { fn failed_run_leaves_no_table() { let dir = TestDir::new("failed"); let path = dir.path().join("grain.tbl"); + let clip = multi_scene_clip(0); - let result = run(multi_scene_clip(0), 1, Some(path.clone())); + let result = run(clip, 1, Some(path.clone())); + let temporary = dir.path().join("grain.tbl.tmp"); assert!(result.is_err()); assert!(!path.exists()); - assert!(!dir.path().join("grain.tbl.tmp").exists()); + assert!(!temporary.exists()); } #[test] fn an_unwritable_path_fails_before_decoding() { let path = PathBuf::from("/nonexistent-directory/grain.tbl"); + let clip = multi_scene_clip(10); - let err = run(multi_scene_clip(10), 1, Some(path)).expect_err("cannot write there"); + let err = run(clip, 1, Some(path)).expect_err("cannot write there"); assert!(err.to_string().contains("grain table"), "got {err}"); } diff --git a/av-denoise/src/bin/pipeline/tests/mod.rs b/av-denoise/src/bin/pipeline/tests/mod.rs index f13b533..80c0f0c 100644 --- a/av-denoise/src/bin/pipeline/tests/mod.rs +++ b/av-denoise/src/bin/pipeline/tests/mod.rs @@ -9,11 +9,15 @@ mod source; mod stage; mod worker; +use std::io::{Cursor, Read}; #[cfg(feature = "vulkan")] use std::sync::{Arc, Mutex}; -use av_denoise::frame::fill_plane; -use av_denoise::{Depth, FrameLayout, Planes, Subsampling}; +#[cfg(feature = "vulkan")] +use av_denoise::accelerate::Accelerator; +#[cfg(feature = "vulkan")] +use av_denoise::{Algorithm, ChannelIntent, DenoisingMode, Device, PlaneOptions}; +use av_denoise::{Depth, FrameLayout, Planes, Subsampling, fill_plane}; /// A writer the test can read back after the coordinator thread drops it. #[cfg(feature = "vulkan")] @@ -32,9 +36,29 @@ impl std::io::Write for SharedBuffer { } } +pub(super) fn y4m_reader(bytes: Vec) -> Box { + let cursor = Cursor::new(bytes); + + Box::new(cursor) +} + +#[cfg(feature = "vulkan")] +pub(super) fn temporal_opts() -> PlaneOptions { + PlaneOptions { + accelerators: vec![Accelerator::Vulkan], + device: Device::Default, + intent: ChannelIntent::LumaChroma, + mode: DenoisingMode::Temporal { radius: 1 }, + algorithm: Algorithm::default(), + luma_strength: None, + chroma_strength: None, + luma_lambda_ht: None, + chroma_lambda_ht: None, + } +} + pub fn tiny_layout() -> FrameLayout { - // 4:2:0 chroma at this size is 4x4, clearing the denoiser's 3x3 - // minimum frame dimension. + // 4:2:0 chroma at this size is 4x4, clearing the denoiser's 3x3 minimum frame dimension. FrameLayout { width: 8, height: 8, @@ -46,9 +70,10 @@ pub fn tiny_layout() -> FrameLayout { pub fn tiny_planes(layout: FrameLayout) -> Planes { let luma_pixels = layout.luma_pixels(); let neutral = layout.depth.neutral_chroma(); + let luma = fill_plane(luma_pixels, neutral, layout.depth); Planes { - y: fill_plane(luma_pixels, neutral, layout.depth), + y: luma, u: layout.neutral_chroma_plane(), v: layout.neutral_chroma_plane(), } @@ -57,7 +82,8 @@ pub fn tiny_planes(layout: FrameLayout) -> Planes { /// A 4x4 8-bit 4:2:0 y4m clip holding `frames` flat frames. pub fn y4m_clip(frames: usize) -> Vec { let mut bytes = Vec::new(); - let mut encoder = y4m::encode(4, 4, y4m::Ratio::new(25, 1)) + let frame_rate = y4m::Ratio::new(25, 1); + let mut encoder = y4m::encode(4, 4, frame_rate) .with_colorspace(y4m::Colorspace::C420) .write_header(&mut bytes) .expect("header should write"); @@ -81,16 +107,16 @@ pub const SCENE_CLIP_SIZE: usize = 64; /// A tiny xorshift so the clip is the same on every run. pub fn pattern(seed: u32, len: usize) -> Vec { let mut state = seed.max(1); - let mut out = Vec::with_capacity(len); + let mut samples = Vec::with_capacity(len); for _ in 0..len { state ^= state << 13; state ^= state >> 17; state ^= state << 5; - out.push((state & 0xff) as u8); + samples.push((state & 0xff) as u8); } - out + samples } /// A 64x64 8-bit 4:2:0 clip of `frames` frames, with `XCOLORRANGE=LIMITED` in the header. @@ -99,9 +125,10 @@ pub fn pattern(seed: u32, len: usize) -> Vec { /// hard cut. pub fn multi_scene_clip(frames: usize) -> Vec { let mut bytes = Vec::new(); - let range = - y4m::VendorExtensionString::new(b"COLORRANGE=LIMITED".to_vec()).expect("the extension has no spaces"); - let mut encoder = y4m::encode(SCENE_CLIP_SIZE, SCENE_CLIP_SIZE, y4m::Ratio::new(25, 1)) + let range_tag = b"COLORRANGE=LIMITED".to_vec(); + let range = y4m::VendorExtensionString::new(range_tag).expect("the extension has no spaces"); + let frame_rate = y4m::Ratio::new(25, 1); + let mut encoder = y4m::encode(SCENE_CLIP_SIZE, SCENE_CLIP_SIZE, frame_rate) .with_colorspace(y4m::Colorspace::C420) .append_vendor_extension(range) .write_header(&mut bytes) @@ -138,7 +165,8 @@ pub fn grainy_clip(frames: usize) -> Vec { let width = 192; let height = 128; let mut bytes = Vec::new(); - let mut encoder = y4m::encode(width, height, y4m::Ratio::new(25, 1)) + let frame_rate = y4m::Ratio::new(25, 1); + let mut encoder = y4m::encode(width, height, frame_rate) .with_colorspace(y4m::Colorspace::C420) .write_header(&mut bytes) .expect("header should write"); diff --git a/av-denoise/src/bin/pipeline/tests/run.rs b/av-denoise/src/bin/pipeline/tests/run.rs index 35c9041..19af416 100644 --- a/av-denoise/src/bin/pipeline/tests/run.rs +++ b/av-denoise/src/bin/pipeline/tests/run.rs @@ -2,39 +2,25 @@ use std::io::{Cursor, Read}; -use av_denoise::accelerate::Accelerator; -use av_denoise::{Algorithm, ChannelIntent, DenoisingMode, Device, PlaneOptions}; - -use super::{SCENE_CLIP_SIZE, SharedBuffer, multi_scene_clip}; +use super::{SCENE_CLIP_SIZE, SharedBuffer, multi_scene_clip, temporal_opts}; use crate::pipeline::run_with; use crate::pipeline::source::open_y4m; use crate::pipeline::stage::frame_permits; -fn temporal_opts() -> PlaneOptions { - PlaneOptions { - accelerators: vec![Accelerator::Vulkan], - device: Device::Default, - intent: ChannelIntent::LumaChroma, - mode: DenoisingMode::Temporal { radius: 1 }, - algorithm: Algorithm::default(), - luma_strength: None, - chroma_strength: None, - luma_lambda_ht: None, - chroma_lambda_ht: None, - } -} - fn run_over(bytes: Vec, workers: usize, budget: u64) -> Result, anyhow::Error> { let output = SharedBuffer::default(); + let writer = output.clone(); let options = temporal_opts(); let opener = move || { - let reader: Box = Box::new(Cursor::new(bytes)); + let cursor = Cursor::new(bytes); + let reader: Box = Box::new(cursor); open_y4m(reader) }; - run_with(&options, opener, workers, budget, false, output.clone(), None)?; + run_with(&options, opener, workers, budget, false, writer, None)?; let written = output.0.lock().expect("buffer lock").clone(); + Ok(written) } @@ -53,8 +39,9 @@ fn frame_count(y4m_bytes: &[u8]) -> usize { fn a_pipe_round_trips_every_frame_and_its_colour_range() { let clip = multi_scene_clip(30); let output = run_over(clip, 2, 1 << 30).expect("the run succeeds"); + let frames = frame_count(&output); - assert_eq!(frame_count(&output), 30); + assert_eq!(frames, 30); assert!(output.windows(17).any(|window| window == b"XCOLORRANGE=LIMIT")); } @@ -62,8 +49,9 @@ fn a_pipe_round_trips_every_frame_and_its_colour_range() { fn a_one_frame_pipe_round_trips() { let clip = multi_scene_clip(1); let output = run_over(clip, 2, 1 << 30).expect("the run succeeds"); + let frames = frame_count(&output); - assert_eq!(frame_count(&output), 1); + assert_eq!(frames, 1); } #[test] @@ -83,6 +71,7 @@ fn a_run_at_the_permit_floor_finishes() { let clip = multi_scene_clip(40); let output = run_over(clip, 2, budget).expect("the floor must not deadlock"); + let frames = frame_count(&output); - assert_eq!(frame_count(&output), 40); + assert_eq!(frames, 40); } diff --git a/av-denoise/src/bin/pipeline/tests/scenes.rs b/av-denoise/src/bin/pipeline/tests/scenes.rs index c905fe8..5ad8ccc 100644 --- a/av-denoise/src/bin/pipeline/tests/scenes.rs +++ b/av-denoise/src/bin/pipeline/tests/scenes.rs @@ -13,7 +13,8 @@ const HEIGHT: usize = 128; /// Builds an 8-bit 4:2:0 y4m clip where each entry of `scene_lengths` is one scene. fn clip(scene_lengths: &[usize]) -> Vec { let mut bytes = Vec::new(); - let mut encoder = y4m::encode(WIDTH, HEIGHT, y4m::Ratio::new(24, 1)) + let frame_rate = y4m::Ratio::new(24, 1); + let mut encoder = y4m::encode(WIDTH, HEIGHT, frame_rate) .with_colorspace(y4m::Colorspace::C420) .write_header(&mut bytes) .expect("header should write"); @@ -38,7 +39,9 @@ fn clip(scene_lengths: &[usize]) -> Vec { } fn decoder_over(bytes: &[u8]) -> Decoder { - let reader: Box = Box::new(Cursor::new(bytes.to_vec())); + let owned = bytes.to_vec(); + let cursor = Cursor::new(owned); + let reader: Box = Box::new(cursor); let y4m_decoder = y4m::decode(reader).expect("y4m header should parse"); Decoder::from_decoder_impl(DecoderImpl::Y4m(y4m_decoder)).expect("decoder should open") @@ -62,12 +65,14 @@ fn split(bytes: &[u8]) -> Vec> { let mut decided = Vec::new(); while let Ok(frame) = decoder.read_video_frame::() { - let released = splitter.push(Arc::new(frame)); + let shared = Arc::new(frame); + let released = splitter.push(shared); decided.extend(released); } let tail = splitter.finish(); decided.extend(tail); + decided } @@ -85,9 +90,10 @@ fn the_pipeline_clip_cuts_at_every_pattern_switch() { let bytes = multi_scene_clip(40); let decided = split(&bytes); let expected_starts: Vec = (0..40).step_by(SCENE_LENGTH).collect(); + let starts = scene_starts(&decided); assert_eq!(decided.len(), 40); - assert_eq!(scene_starts(&decided), expected_starts); + assert_eq!(starts, expected_starts); } #[test] @@ -101,9 +107,10 @@ fn cuts_match_the_whole_clip_pass_over_several_scenes() { ); let decided = split(&bytes); + let starts = scene_starts(&decided); assert_eq!(decided.len(), expected_count); - assert_eq!(scene_starts(&decided), expected_starts); + assert_eq!(starts, expected_starts); } #[test] @@ -111,9 +118,10 @@ fn cuts_match_the_whole_clip_pass_at_the_look_ahead_length() { let bytes = clip(&[LOOKAHEAD_DISTANCE + 1]); let (expected_starts, expected_count) = reference(&bytes); let decided = split(&bytes); + let starts = scene_starts(&decided); assert_eq!(decided.len(), expected_count); - assert_eq!(scene_starts(&decided), expected_starts); + assert_eq!(starts, expected_starts); } #[test] @@ -121,9 +129,10 @@ fn cuts_match_the_whole_clip_pass_on_two_frames() { let bytes = clip(&[1, 1]); let (expected_starts, expected_count) = reference(&bytes); let decided = split(&bytes); + let starts = scene_starts(&decided); assert_eq!(decided.len(), expected_count); - assert_eq!(scene_starts(&decided), expected_starts); + assert_eq!(starts, expected_starts); } #[test] @@ -145,10 +154,11 @@ fn every_frame_is_released_once_in_order() { let mut released = Vec::new(); while let Ok(frame) = decoder.read_video_frame::() { - let frame = Arc::new(frame); - pushed.push(Arc::clone(&frame)); + let shared = Arc::new(frame); + let pushed_frame = Arc::clone(&shared); + pushed.push(pushed_frame); - let decided = splitter.push(frame); + let decided = splitter.push(shared); released.extend(decided); } @@ -158,7 +168,9 @@ fn every_frame_is_released_once_in_order() { assert_eq!(released.len(), pushed.len()); for (original, decided) in pushed.iter().zip(&released) { - assert!(Arc::ptr_eq(original, &decided.frame)); + let same_frame = Arc::ptr_eq(original, &decided.frame); + + assert!(same_frame); } } @@ -179,7 +191,8 @@ fn nothing_is_released_before_the_look_ahead_fills() { for pushed in 0..LOOKAHEAD_DISTANCE { let frame = decoder.read_video_frame::().expect("the clip has 20 frames"); - let released = splitter.push(Arc::new(frame)); + let shared = Arc::new(frame); + let released = splitter.push(shared); assert!( released.is_empty(), @@ -188,7 +201,8 @@ fn nothing_is_released_before_the_look_ahead_fills() { } let frame = decoder.read_video_frame::().expect("the clip has 20 frames"); - let released = splitter.push(Arc::new(frame)); + let shared = Arc::new(frame); + let released = splitter.push(shared); assert_eq!( released.len(), diff --git a/av-denoise/src/bin/pipeline/tests/source.rs b/av-denoise/src/bin/pipeline/tests/source.rs index dda5402..6ae56f1 100644 --- a/av-denoise/src/bin/pipeline/tests/source.rs +++ b/av-denoise/src/bin/pipeline/tests/source.rs @@ -1,19 +1,20 @@ -use std::io::{Cursor, Read}; - use av_denoise::{Depth, Subsampling}; +use super::y4m_reader; use crate::pipeline::convert::SourcePixel; use crate::pipeline::source::{color_range_extension, open_y4m, pixel_aspect_from_sar}; fn y4m_bytes(colorspace: y4m::Colorspace, extension: Option<&str>, frames: usize) -> Vec { let mut bytes = Vec::new(); - let mut builder = y4m::encode(4, 4, y4m::Ratio::new(25, 1)) + let frame_rate = y4m::Ratio::new(25, 1); + let pixel_aspect = y4m::Ratio::new(4, 3); + let mut builder = y4m::encode(4, 4, frame_rate) .with_colorspace(colorspace) - .with_pixel_aspect(y4m::Ratio::new(4, 3)); + .with_pixel_aspect(pixel_aspect); if let Some(extension) = extension { - let vendor = y4m::VendorExtensionString::new(extension.as_bytes().to_vec()) - .expect("the extension has no spaces"); + let value = extension.as_bytes().to_vec(); + let vendor = y4m::VendorExtensionString::new(value).expect("the extension has no spaces"); builder = builder.append_vendor_extension(vendor); } @@ -34,14 +35,10 @@ fn y4m_bytes(colorspace: y4m::Colorspace, extension: Option<&str>, frames: usize bytes } -fn reader(bytes: Vec) -> Box { - Box::new(Cursor::new(bytes)) -} - #[test] fn a_pipe_keeps_its_vendor_extensions_and_pixel_aspect() { let bytes = y4m_bytes(y4m::Colorspace::C420, Some("COLORRANGE=LIMITED"), 1); - let source = reader(bytes); + let source = y4m_reader(bytes); let opened = open_y4m(source).expect("a 4:2:0 pipe opens"); let extensions: Vec<&[u8]> = opened @@ -59,7 +56,7 @@ fn a_pipe_keeps_its_vendor_extensions_and_pixel_aspect() { #[test] fn a_ten_bit_pipe_reports_its_layout() { let bytes = y4m_bytes(y4m::Colorspace::C422p10, None, 1); - let source = reader(bytes); + let source = y4m_reader(bytes); let opened = open_y4m(source).expect("a 10-bit 4:2:2 pipe opens"); assert_eq!(opened.info.layout.width, 4); @@ -70,7 +67,7 @@ fn a_ten_bit_pipe_reports_its_layout() { #[test] fn ten_bit_stream_round_trips_header_and_plane_sizes() { let bytes = y4m_bytes(y4m::Colorspace::C420p10, None, 2); - let source = reader(bytes); + let source = y4m_reader(bytes); let mut opened = open_y4m(source).expect("a 10-bit 4:2:0 pipe opens"); let layout = opened.info.layout; @@ -91,7 +88,7 @@ fn ten_bit_stream_round_trips_header_and_plane_sizes() { #[test] fn a_pipe_has_no_phantom_frames_or_frame_estimate() { let bytes = y4m_bytes(y4m::Colorspace::C420, None, 3); - let source = reader(bytes); + let source = y4m_reader(bytes); let opened = open_y4m(source).expect("a 4:2:0 pipe opens"); assert!(opened.phantom.is_empty()); @@ -101,13 +98,15 @@ fn a_pipe_has_no_phantom_frames_or_frame_estimate() { #[test] fn a_mono_pipe_is_rejected_without_panicking() { let mut bytes = Vec::new(); - y4m::encode(4, 4, y4m::Ratio::new(25, 1)) + let frame_rate = y4m::Ratio::new(25, 1); + y4m::encode(4, 4, frame_rate) .with_colorspace(y4m::Colorspace::Cmono) .write_header(&mut bytes) .expect("header should write"); - let source = reader(bytes); - let err = match open_y4m(source) { + let source = y4m_reader(bytes); + let result = open_y4m(source); + let err = match result { Ok(_) => panic!("a mono pipe must be rejected"), Err(err) => err, }; @@ -132,10 +131,9 @@ fn a_full_range_maps_to_the_full_tag() { #[test] fn an_unspecified_or_unknown_range_adds_no_tag() { for range in [0, 3, -1] { - assert!( - color_range_extension(range).is_none(), - "range {range} must add no tag" - ); + let extension = color_range_extension(range); + + assert!(extension.is_none(), "range {range} must add no tag"); } } diff --git a/av-denoise/src/bin/pipeline/tests/stage.rs b/av-denoise/src/bin/pipeline/tests/stage.rs index 48e77c9..a4aaf77 100644 --- a/av-denoise/src/bin/pipeline/tests/stage.rs +++ b/av-denoise/src/bin/pipeline/tests/stage.rs @@ -2,6 +2,8 @@ use std::thread; use std::time::Duration; use super::{tiny_layout, tiny_planes}; +use crate::pipeline::decode::PREFETCH_FRAMES; +use crate::pipeline::scenes::LOOKAHEAD_DISTANCE; use crate::pipeline::stage::{IN_TRANSIT_FRAMES, SceneJob, Stager, checked_frame_permits, frame_permits}; /// Stages one frame per flag, starting a new scene wherever a flag is set. @@ -11,10 +13,13 @@ fn stage_all(flags: &[bool], jobs: &crossbeam_channel::Sender) -> Resu let mut stager = Stager::new(jobs); for &starts_scene in flags { - stager.stage(planes.clone(), starts_scene)?; + let frame_planes = planes.clone(); + stager.stage(frame_planes, starts_scene)?; } - Ok(stager.finish()) + let staged = stager.finish(); + + Ok(staged) } fn flags(scene_lengths: &[usize]) -> Vec { @@ -24,22 +29,21 @@ fn flags(scene_lengths: &[usize]) -> Vec { .collect() } -/// Drains every job the stager offers, returning each scene index with -/// the frame indices that scene carried. +/// Drains every job the stager offers, returning each scene index with its frame indices. /// -/// Runs on its own thread because the scene queue is a rendezvous, so the -/// stager blocks until someone claims each job. -fn collect_jobs(rx: crossbeam_channel::Receiver) -> thread::JoinHandle)>> { +/// Runs on its own thread because the scene queue is a rendezvous, so the stager blocks until +/// someone claims each job. +fn collect_jobs(job_rx: crossbeam_channel::Receiver) -> thread::JoinHandle)>> { thread::spawn(move || { - let mut out = Vec::new(); + let mut jobs = Vec::new(); - while let Ok(job) = rx.recv() { - let idx = job.scene_idx; - let frames = job.frames.iter().map(|f| f.global_idx).collect(); - out.push((idx, frames)); + while let Ok(job) = job_rx.recv() { + let scene_idx = job.scene_idx; + let frames = job.frames.iter().map(|frame| frame.global_idx).collect(); + jobs.push((scene_idx, frames)); } - out + jobs }) } @@ -79,19 +83,32 @@ fn a_scene_job_channel_closes_when_its_scene_ends() { ); } -/// A worker that claims a scene and dies without draining it must surface -/// as an error rather than hanging the stager. #[test] fn staging_fails_rather_than_hanging_when_the_pool_dies() { let (job_tx, job_rx) = crossbeam_channel::bounded::(0); - let pool = thread::spawn(move || drop(job_rx.recv())); + let (die_tx, die_rx) = crossbeam_channel::bounded::<()>(0); + let pool = thread::spawn(move || { + let claimed = job_rx.recv(); + let _ = die_rx.recv(); + drop(claimed); + }); - let scene_starts = flags(&[10]); + let layout = tiny_layout(); + let planes = tiny_planes(layout); + let mut stager = Stager::new(&job_tx); - let err = stage_all(&scene_starts, &job_tx).expect_err("staging must not hang"); + let first_planes = planes.clone(); + stager + .stage(first_planes, true) + .expect("the pool claims the first scene"); + // The scene's channel is unbounded, so a send only fails once the claimed job is dropped. + // The pool holds the job until told to die, so the drop lands between the two frames. + drop(die_tx); pool.join().expect("pool panicked"); + let err = stager.stage(planes, false).expect_err("staging must not hang"); + assert!( err.to_string().contains("disconnect"), "error should name the disconnect: {err}" @@ -114,7 +131,9 @@ fn a_leading_frame_without_a_scene_flag_still_opens_a_scene() { #[test] fn frame_permits_follows_the_budget_when_it_clears_the_floor() { // A 1080p 8-bit 4:2:0 frame is 3,110,400 bytes, so 1 GiB affords 345. - assert_eq!(frame_permits(1 << 30, 3_110_400, 4, 0), 345); + let permits = frame_permits(1 << 30, 3_110_400, 4, 0); + + assert_eq!(permits, 345); } #[test] @@ -122,7 +141,9 @@ fn frame_permits_applies_the_floor_when_the_budget_is_too_small() { let floor = 4 * (av_denoise::MAX_PENDING + 2) + IN_TRANSIT_FRAMES; // A 4K 10-bit frame is 24,883,200 bytes, so 1 MiB affords none. - assert_eq!(frame_permits(1 << 20, 24_883_200, 4, 0), floor); + let permits = frame_permits(1 << 20, 24_883_200, 4, 0); + + assert_eq!(permits, floor); } #[test] @@ -130,14 +151,18 @@ fn a_budget_below_the_floor_is_rejected() { // A 4K 10-bit frame is 24,883,200 bytes, so 1 MB affords none. let err = checked_frame_permits(1_000_000, 24_883_200, 4, 8) .expect_err("1 MB cannot feed 4 workers at radius 8"); - let msg = err.to_string(); + let message = err.to_string(); let floor = 4 * (8 + av_denoise::MAX_PENDING + 2) + IN_TRANSIT_FRAMES; + let floor_text = format!("at least {floor}"); - assert!(msg.contains("affords 0 frames"), "got {msg}"); - assert!(msg.contains(&format!("at least {floor}")), "got {msg}"); + assert!(message.contains("affords 0 frames"), "got {message}"); + assert!(message.contains(&floor_text), "got {message}"); // 60 frames at 24,883,200 bytes is 1,492,992,000, rounded up to 1.5GB. - assert!(msg.contains("Pass at least --frame-budget 1.5GB"), "got {msg}"); + assert!( + message.contains("Pass at least --frame-budget 1.5GB"), + "got {message}" + ); } #[test] @@ -153,8 +178,8 @@ fn the_floor_covers_a_workers_first_output_at_every_radius() { for radius in [0u32, 1, 4, 8] { let permits = frame_permits(1, 199_065_600, 1, radius); - // Pushes a worker needs before `push` first returns QueueFull, - // which is the first point it can emit and return a permit. + // Pushes a worker needs before `push` first returns QueueFull, which is the first point + // it can emit and return a permit. let first_output = radius as usize + av_denoise::MAX_PENDING + 1; assert!( @@ -169,18 +194,13 @@ fn the_floor_covers_the_look_ahead_and_prefetch() { // A budget far too small for any real frame, so the floor decides. let permits = frame_permits(1, 199_065_600, 1, 0); let first_output = av_denoise::MAX_PENDING + 1; - let held_upstream = - crate::pipeline::scenes::LOOKAHEAD_DISTANCE + 2 + crate::pipeline::decode::PREFETCH_FRAMES + 1; + let held_upstream = LOOKAHEAD_DISTANCE + 2 + PREFETCH_FRAMES + 1; assert!(permits >= first_output + held_upstream, "got {permits}"); } -/// A worker that claims a scene and stops reading it must not stop -/// later scenes being offered to anyone else. -/// -/// Scene 0 holds ten frames that nobody drains. Bounding each scene's -/// channel would block the stager inside scene 0 and starve every idle -/// worker behind it, which is the stall this pins. +/// Scene 0 holds ten frames that nobody drains. Bounding each scene's channel would block the +/// stager inside scene 0 and starve every idle worker behind it, which is the stall this pins. #[test] fn a_backlogged_scene_does_not_stop_later_scenes_being_offered() { let (job_tx, job_rx) = crossbeam_channel::bounded::(0); @@ -188,13 +208,12 @@ fn a_backlogged_scene_does_not_stop_later_scenes_being_offered() { let consumer = thread::spawn(move || { let first = job_rx.recv().expect("scene 0 is offered"); - // Held rather than discarded. Dropping a job closes its frame - // channel, and the stager is still filling scene 1's. + // Held rather than discarded, because dropping a job closes its frame channel and the + // stager is still filling scene 1's. let second = job_rx.recv_timeout(Duration::from_secs(5)).ok(); let offered_while_backlogged = second.is_some(); - // Drain everything either way, so a failing run finishes and - // reports instead of hanging. + // Drain everything either way, so a failing run finishes and reports instead of hanging. for _ in first.frames.iter() {} while job_rx.recv().is_ok() {} drop(second); @@ -207,8 +226,10 @@ fn a_backlogged_scene_does_not_stop_later_scenes_being_offered() { stage_all(&scene_starts, &job_tx).expect("staging should not stall"); drop(job_tx); + let offered = consumer.join().expect("consumer panicked"); + assert!( - consumer.join().expect("consumer panicked"), - "scene 1 must be offered while scene 0 is still backlogged", + offered, + "scene 1 must be offered while scene 0 is still backlogged" ); } diff --git a/av-denoise/src/bin/pipeline/tests/worker.rs b/av-denoise/src/bin/pipeline/tests/worker.rs index 2027b72..f394843 100644 --- a/av-denoise/src/bin/pipeline/tests/worker.rs +++ b/av-denoise/src/bin/pipeline/tests/worker.rs @@ -1,59 +1,37 @@ -// `temporal_opts` and the one test that uses it are the only things -// naming `Accelerator::Vulkan`, `Algorithm`, `Device`, and -// `ChannelIntent`. Their imports are gated the same way to keep -// cpu-only builds free of unused-import warnings. +// The test, its helper and their imports are gated because the `Vulkan` accelerator variant only +// exists with the `vulkan` feature, and gating the imports keeps cpu-only builds free of warnings. #[cfg(feature = "vulkan")] -use av_denoise::accelerate::Accelerator; +use std::collections::VecDeque; + #[cfg(feature = "vulkan")] -use av_denoise::{Algorithm, ChannelIntent, DenoisingMode, Device, PlanarDenoiser, PlaneOptions}; +use av_denoise::PlanarDenoiser; #[cfg(feature = "vulkan")] -use super::{tiny_layout, tiny_planes}; +use super::{temporal_opts, tiny_layout, tiny_planes}; #[cfg(feature = "vulkan")] use crate::pipeline::coordinator::OutputMsg; #[cfg(feature = "vulkan")] use crate::pipeline::worker::flush_worker; -/// Gated because it names the `Vulkan` accelerator variant, which only -/// exists when the `vulkan` feature is enabled. -#[cfg(feature = "vulkan")] -fn temporal_opts() -> PlaneOptions { - PlaneOptions { - accelerators: vec![Accelerator::Vulkan], - device: Device::Default, - intent: ChannelIntent::LumaChroma, - mode: DenoisingMode::Temporal { radius: 1 }, - algorithm: Algorithm::default(), - luma_strength: None, - chroma_strength: None, - luma_lambda_ht: None, - chroma_lambda_ht: None, - } -} - -/// Gated because it depends on `temporal_opts`, which names the -/// `Vulkan` accelerator variant and only builds when the `vulkan` -/// feature is enabled. #[cfg(feature = "vulkan")] #[test] fn flush_worker_errors_when_coordinator_has_disconnected() { let layout = tiny_layout(); let options = temporal_opts(); - let mut wd = PlanarDenoiser::create(&options, layout).expect("denoiser construction failed"); + let mut denoiser = PlanarDenoiser::create(&options, layout).expect("denoiser construction failed"); let planes = tiny_planes(layout); - // One push into a temporal window leaves a trailing tail that - // `flush` will pad and emit. - wd.push(&planes).expect("push failed"); + // One push into a temporal window leaves a trailing tail that `flush` pads and emits. + denoiser.push(&planes).expect("push failed"); - let mut pending: std::collections::VecDeque = std::collections::VecDeque::new(); + let mut pending: VecDeque = VecDeque::new(); pending.push_back(0); - let (tx, rx) = crossbeam_channel::unbounded::(); - drop(rx); + let (output_tx, output_rx) = crossbeam_channel::unbounded::(); + drop(output_rx); let mut warm_up = None; - let err = flush_worker(&mut wd, &mut warm_up, &mut pending, &tx) + let err = flush_worker(&mut denoiser, &mut warm_up, &mut pending, &output_tx) .expect_err("expected the coordinator disconnect to surface as an error"); assert!( diff --git a/av-denoise/src/bin/pipeline/worker.rs b/av-denoise/src/bin/pipeline/worker.rs index 3075c30..3ee916a 100644 --- a/av-denoise/src/bin/pipeline/worker.rs +++ b/av-denoise/src/bin/pipeline/worker.rs @@ -1,3 +1,4 @@ +use std::collections::VecDeque; use std::thread; use av_denoise::{FrameLayout, PlanarDenoiser, PlaneOptions, Planes, SceneGrain, WarmUp, push_needs_retry}; @@ -10,11 +11,9 @@ pub type WorkerJoin = thread::JoinHandle, anyhow::Error>> /// Spawns `workers` worker threads over one shared scene queue. /// -/// The queue is a rendezvous, so a scene is only offered when a worker is -/// free and at most `workers` scenes are ever in flight. -/// -/// Returns the queue's sender, their join handles, and the shared output -/// channel they emit denoised frames on. +/// The queue is a rendezvous, so a scene is only offered when a worker is free and at most +/// `workers` scenes are ever in flight. Returns the queue's sender, the join handles and the +/// shared output channel the workers emit denoised frames on. pub fn spawn_workers( opts: &PlaneOptions, layout: FrameLayout, @@ -25,42 +24,40 @@ pub fn spawn_workers( crossbeam_channel::Receiver, ) { let (job_tx, job_rx) = crossbeam_channel::bounded::(0); - let (out_tx, out_rx) = crossbeam_channel::unbounded::(); + let (output_tx, output_rx) = crossbeam_channel::unbounded::(); let mut worker_handles: Vec = Vec::with_capacity(workers); for worker_id in 0..workers { let opts = opts.clone(); - let out_tx = out_tx.clone(); + let output_tx = output_tx.clone(); let job_rx = job_rx.clone(); - worker_handles.push(thread::spawn(move || { - run_worker(worker_id, opts, layout, job_rx, out_tx) - })); + let handle = thread::spawn(move || run_worker(worker_id, opts, layout, job_rx, output_tx)); + worker_handles.push(handle); } - // Drop the original sender so the channel closes once every worker - // clone has terminated. - drop(out_tx); + // Dropping the original sender lets the channel close once every worker's clone is gone. + drop(output_tx); - (job_tx, worker_handles, out_rx) + (job_tx, worker_handles, output_rx) } +/// Denoises every scene this worker claims and returns the grain it measured. pub fn run_worker( worker_id: usize, opts: PlaneOptions, layout: FrameLayout, jobs: crossbeam_channel::Receiver, - tx: crossbeam_channel::Sender, + output_tx: crossbeam_channel::Sender, ) -> Result, anyhow::Error> { let mut scenes = Vec::new(); let mut denoiser_slot: Option = None; - // The cold-cache queue place this worker's denoiser holds, until its - // first output frame proves the kernels are compiled and cached. + + // Held until the denoiser's first output frame proves its kernels are compiled and cached. let mut warm_up: Option = None; while let Ok(job) = jobs.recv() { - // Built on the first claimed scene, so a worker that never claims - // one never compiles. + // Built on the first claimed scene, so a worker that never claims one never compiles. if denoiser_slot.is_none() { let (denoiser, place) = create_denoiser(&opts, layout)?; denoiser_slot = Some(denoiser); @@ -73,16 +70,13 @@ pub fn run_worker( tracing::debug!(worker_id, scene_idx = job.scene_idx, "worker started scene"); - // Indices of pushed-but-not-yet-emitted frames, in push order. - let mut pending: std::collections::VecDeque = Default::default(); + // Indices of pushed but not yet emitted frames, in push order. + let mut pending: VecDeque = VecDeque::new(); let mut first_frame = None; - // Nothing is received straight after the push. - // `push_with_drain` handles backpressure through QueueFull - // when the 2-deep pending pipeline fills, and `flush_worker` - // drains the tail below. Receiving after every push would clamp - // the pipeline back to depth 1 and put the GPU readback in the - // critical path of the next push. + // Nothing is received straight after a push, because that would clamp the 2-deep pending + // pipeline to depth 1 and put the GPU readback in the next push's critical path. + // `push_with_drain` drains on QueueFull and `flush_worker` drains the tail instead. for frame in job.frames { first_frame.get_or_insert(frame.global_idx); push_with_drain( @@ -91,13 +85,13 @@ pub fn run_worker( &mut pending, frame.global_idx, &frame.planes, - &tx, + &output_tx, )?; } - // Reuse the PlanarDenoiser across scenes. Flushing here ensures - // no temporal window spans two of them. - flush_worker(denoiser, &mut warm_up, &mut pending, &tx)?; + // The denoiser is reused across scenes, so flushing here stops a temporal window + // spanning two of them. + flush_worker(denoiser, &mut warm_up, &mut pending, &output_tx)?; let chunks = denoiser.drain_grain_chunks()?; if let Some(first_frame) = first_frame @@ -110,23 +104,25 @@ pub fn run_worker( Ok(scenes) } -/// Push one frame, draining any pending output first if the queue is full. +/// Pushes one frame, emitting the oldest pending output and retrying when the queue is full. pub fn push_with_drain( denoiser: &mut PlanarDenoiser, warm_up: &mut Option, - pending: &mut std::collections::VecDeque, + pending: &mut VecDeque, global_idx: u64, planes: &Planes, - tx: &crossbeam_channel::Sender, + output_tx: &crossbeam_channel::Sender, ) -> Result<(), anyhow::Error> { pending.push_back(global_idx); - if push_needs_retry(denoiser.push(planes))? { - if let Some(out) = denoiser.recv()? { + let push_result = denoiser.push(planes); + + if push_needs_retry(push_result)? { + if let Some(denoised) = denoiser.recv()? { let oldest_idx = pending .pop_front() .expect("pending has at least one entry on QueueFull recv"); - send_output(tx, oldest_idx, out)?; + send_output(output_tx, oldest_idx, denoised)?; finish_warm_up(warm_up); } @@ -137,33 +133,36 @@ pub fn push_with_drain( } pub fn send_output( - tx: &crossbeam_channel::Sender, + output_tx: &crossbeam_channel::Sender, global_idx: u64, planes: Planes, ) -> Result<(), anyhow::Error> { - tx.send(OutputMsg { global_idx, planes }) + output_tx + .send(OutputMsg { global_idx, planes }) .map_err(|_| anyhow::anyhow!("coordinator disconnected")) } +/// Flushes the denoiser, emitting every remaining frame against its pending index. pub fn flush_worker( denoiser: &mut PlanarDenoiser, warm_up: &mut Option, - pending: &mut std::collections::VecDeque, - tx: &crossbeam_channel::Sender, + pending: &mut VecDeque, + output_tx: &crossbeam_channel::Sender, ) -> Result<(), anyhow::Error> { let mut disconnected = false; - denoiser.flush(|out| { + denoiser.flush(|denoised| { if disconnected { return; } if let Some(global_idx) = pending.pop_front() { - let msg = OutputMsg { + let message = OutputMsg { global_idx, - planes: out, + planes: denoised, }; - let did_send = tx.send(msg).is_ok(); + let did_send = output_tx.send(message).is_ok(); + if did_send { finish_warm_up(warm_up); } else { diff --git a/av-denoise/src/bin/progress.rs b/av-denoise/src/bin/progress.rs index fdca16c..307c513 100644 --- a/av-denoise/src/bin/progress.rs +++ b/av-denoise/src/bin/progress.rs @@ -4,81 +4,67 @@ use indicatif::{MultiProgress, ProgressBar, ProgressStyle}; use tracing_indicatif::IndicatifWriter; use tracing_indicatif::writer::Stderr; -/// Bar style shared by every phase. const BAR_TEMPLATE: &str = "{msg} [{bar:40}] {pos}/{len} ({per_sec}, eta {eta})"; -/// The process-wide bar container. static PROGRESS: OnceLock = OnceLock::new(); /// The `MultiProgress` every bar is registered on. /// -/// Bars go through this rather than drawing to stderr directly. That -/// lets tracing output routed through [`tracing_writer`] suspend them -/// while it writes. -/// -/// `MultiProgress::new` draws to stderr and reports itself hidden when -/// stderr is not a terminal, so a redirected run emits nothing. +/// Bars draw through it so tracing output from [tracing_writer] can suspend them while it writes. +/// It reports itself hidden when stderr is not a terminal, so a redirected run emits nothing. fn multi() -> &'static MultiProgress { PROGRESS.get_or_init(MultiProgress::new) } -/// The writer to hand to the tracing subscriber. +/// The writer for the tracing subscriber. /// -/// Each write is wrapped in `MultiProgress::suspend`, so log lines land -/// above an intact bar instead of overwriting it. +/// Each write is wrapped in `MultiProgress::suspend`, so log lines land above an intact bar +/// instead of overwriting it. pub fn tracing_writer() -> IndicatifWriter { - IndicatifWriter::new(multi().clone()) + let progress = multi().clone(); + IndicatifWriter::new(progress) } /// Whether the denoising progress bar should be drawn. /// -/// That bar is opt-in. It shows only when `progress` is set and the -/// target stream is a terminal. -/// -/// It runs for the whole encode, alongside whatever the consumer of our -/// output is printing, so leaving it off by default keeps a piped run -/// readable. -/// -/// The terminal check is a parameter rather than read directly here so -/// this stays unit-testable without a real tty. +/// The bar is opt-in because it runs for the whole encode alongside whatever the consumer of the +/// output prints, and leaving it off keeps a piped run readable. The terminal check is a parameter +/// so this stays testable without a real tty. pub fn denoise_bar_visible(progress: bool, stream_is_terminal: bool) -> bool { progress && stream_is_terminal } -/// Builds a bar registered on the shared [`multi`]. +/// Builds a bar registered on the shared [multi]. /// -/// Returns a hidden bar when `visible` is false, so callers can drive it -/// unconditionally without branching on visibility. +/// Returns a hidden bar when `visible` is false, so callers can drive it unconditionally. fn bar(total_frames: Option, message: &str, visible: bool) -> ProgressBar { if !visible { return ProgressBar::hidden(); } - let pb = match total_frames { - Some(n) => ProgressBar::new(n as u64), + let progress_bar = match total_frames { + Some(total) => ProgressBar::new(total as u64), None => ProgressBar::no_length(), }; if let Ok(style) = ProgressStyle::with_template(BAR_TEMPLATE) { - pb.set_style(style); + progress_bar.set_style(style); } - pb.set_message(message.to_owned()); + progress_bar.set_message(message.to_owned()); - multi().add(pb) + multi().add(progress_bar) } -/// Builds the denoising progress bar, tracking frames written to the -/// output. +/// Builds the denoising progress bar, which tracks frames written to the output. pub fn denoise_progress_bar(total_frames: Option, visible: bool) -> ProgressBar { bar(total_frames, "denoising", visible) } -/// Clears a finished bar and drops it from the shared [`multi`], so it -/// leaves no blank line behind. -pub fn finish(pb: &ProgressBar) { - pb.finish_and_clear(); - multi().remove(pb); +/// Clears a finished bar and removes it from [multi] so it leaves no blank line behind. +pub fn finish(progress_bar: &ProgressBar) { + progress_bar.finish_and_clear(); + multi().remove(progress_bar); } #[cfg(test)] @@ -87,48 +73,66 @@ mod tests { #[test] fn denoise_bar_shown_when_requested_and_terminal() { - assert!(denoise_bar_visible(true, true)); + let visible = denoise_bar_visible(true, true); + + assert!(visible); } #[test] fn denoise_bar_hidden_when_not_requested() { - assert!(!denoise_bar_visible(false, true)); + let visible = denoise_bar_visible(false, true); + + assert!(!visible); } #[test] fn denoise_bar_hidden_when_stream_is_not_a_terminal() { - assert!(!denoise_bar_visible(true, false)); + let visible = denoise_bar_visible(true, false); + + assert!(!visible); } #[test] fn denoise_bar_hidden_when_neither_requested_nor_a_terminal() { - assert!(!denoise_bar_visible(false, false)); + let visible = denoise_bar_visible(false, false); + + assert!(!visible); } #[test] fn bar_template_parses() { - assert!(ProgressStyle::with_template(BAR_TEMPLATE).is_ok()); + let style = ProgressStyle::with_template(BAR_TEMPLATE); + + assert!(style.is_ok()); } #[test] fn denoise_bar_hidden_when_not_visible() { - assert!(denoise_progress_bar(Some(10), false).is_hidden()); + let progress_bar = denoise_progress_bar(Some(10), false); + + assert!(progress_bar.is_hidden()); } #[test] fn denoise_bar_uses_total_as_length() { - assert_eq!(denoise_progress_bar(Some(10), true).length(), Some(10)); + let progress_bar = denoise_progress_bar(Some(10), true); + + assert_eq!(progress_bar.length(), Some(10)); } #[test] fn denoise_bar_without_total_has_no_length() { - assert_eq!(denoise_progress_bar(None, true).length(), None); + let progress_bar = denoise_progress_bar(None, true); + + assert_eq!(progress_bar.length(), None); } #[test] fn finish_marks_the_bar_done() { - let pb = denoise_progress_bar(Some(10), true); - finish(&pb); - assert!(pb.is_finished()); + let progress_bar = denoise_progress_bar(Some(10), true); + + finish(&progress_bar); + + assert!(progress_bar.is_finished()); } } diff --git a/av-denoise/src/bin/warm_start.rs b/av-denoise/src/bin/warm_start.rs index 1dd760e..14b1ca9 100644 --- a/av-denoise/src/bin/warm_start.rs +++ b/av-denoise/src/bin/warm_start.rs @@ -1,26 +1,24 @@ use av_denoise::{FrameLayout, PlanarDenoiser, PlaneOptions, WarmUp, kernel_key}; -/// The only place either CLI mode builds a [`PlanarDenoiser`]. +/// Builds a [PlanarDenoiser] after taking a place in the cross-process warm-up queue. /// -/// Takes a place in the cross-process warm-up queue first, so concurrent -/// `av-denoise` processes do not each compile into a cold cache. The -/// returned place, if any, is the caller's to finish once this -/// denoiser has produced its first output frame — see [`WarmUp`] for -/// why it cannot be finished any earlier than that. +/// The queue stops concurrent `av-denoise` processes each compiling into a cold cache. The caller +/// finishes the returned place once the denoiser has produced its first output frame, and [WarmUp] +/// explains why it cannot be finished any earlier. pub fn create_denoiser( opts: &PlaneOptions, layout: FrameLayout, ) -> Result<(PlanarDenoiser, Option), anyhow::Error> { - let warm_up = WarmUp::begin(kernel_key(opts, layout)); + let key = kernel_key(opts, layout); + let warm_up = WarmUp::begin(key); let denoiser = PlanarDenoiser::create(opts, layout)?; Ok((denoiser, warm_up)) } -/// Gives up a cold-cache queue place after a frame has proven the -/// kernels it names are compiled and cached. Does nothing once already -/// finished, or if no place was taken. Mirrors -/// `av-denoise-vs`'s `State::finish_warm_up`. +/// Gives up a cold-cache queue place once a frame has proven its kernels are compiled and cached. +/// +/// Does nothing when no place is held. pub fn finish_warm_up(warm_up: &mut Option) { if let Some(warm_up) = warm_up.take() { warm_up.finish(); @@ -46,8 +44,9 @@ mod tests { luma_lambda_ht: None, chroma_lambda_ht: None, }; - // Zero width collapses the 4:2:0 chroma plane to nothing, which - // `PlanarDenoiser::create` rejects before touching the GPU. + + // Zero width collapses the 4:2:0 chroma plane to nothing, which `PlanarDenoiser::create` + // rejects before touching the GPU. let layout = FrameLayout { width: 0, height: 0, diff --git a/av-denoise/src/bin/y4m_format.rs b/av-denoise/src/bin/y4m_format.rs index cd4361e..d5b2550 100644 --- a/av-denoise/src/bin/y4m_format.rs +++ b/av-denoise/src/bin/y4m_format.rs @@ -1,9 +1,8 @@ use av_denoise::{Depth, Subsampling}; -/// Maps our [`Subsampling`] and [`Depth`] onto the [`y4m::Colorspace`] -/// used to read the input and write the output header. -pub fn subsampling_to_y4m(s: Subsampling, depth: Depth) -> y4m::Colorspace { - match (s, depth) { +/// Maps a [Subsampling] and [Depth] onto the matching [y4m::Colorspace]. +pub fn subsampling_to_y4m(subsampling: Subsampling, depth: Depth) -> y4m::Colorspace { + match (subsampling, depth) { (Subsampling::Yuv420, Depth::Eight) => y4m::Colorspace::C420, (Subsampling::Yuv420, Depth::Ten) => y4m::Colorspace::C420p10, (Subsampling::Yuv420, Depth::Twelve) => y4m::Colorspace::C420p12, @@ -16,13 +15,11 @@ pub fn subsampling_to_y4m(s: Subsampling, depth: Depth) -> y4m::Colorspace { } } -/// Maps a [`y4m::Colorspace`] back onto our [`Subsampling`] and -/// [`Depth`]. +/// Maps a [y4m::Colorspace] back onto a [Subsampling] and [Depth]. /// -/// Grayscale and any other unsupported colorspace are rejected with an -/// error naming what is required instead. -pub fn subsampling_from_y4m(c: y4m::Colorspace) -> Result<(Subsampling, Depth), anyhow::Error> { - let sub = match c { +/// Grayscale and other unsupported colorspaces are rejected with an error naming what is required. +pub fn subsampling_from_y4m(colorspace: y4m::Colorspace) -> Result<(Subsampling, Depth), anyhow::Error> { + let subsampling = match colorspace { y4m::Colorspace::C420 | y4m::Colorspace::C420jpeg | y4m::Colorspace::C420paldv @@ -34,28 +31,25 @@ pub fn subsampling_from_y4m(c: y4m::Colorspace) -> Result<(Subsampling, Depth), other => anyhow::bail!("unsupported y4m colorspace {other:?}, need 4:2:0, 4:2:2, or 4:4:4"), }; - let depth = Depth::from_bits(c.get_bit_depth())?; + let bit_depth = colorspace.get_bit_depth(); + let depth = Depth::from_bits(bit_depth)?; - Ok((sub, depth)) + Ok((subsampling, depth)) } -/// Pulls the `X`-prefixed vendor extension params out of a decoded y4m -/// header's raw params bytes, `XCOLORRANGE=LIMITED` being the common one. +/// Pulls the `X`-prefixed vendor extension params, such as `XCOLORRANGE=LIMITED`, out of a y4m header. /// -/// The leading `X` is stripped so the result can go straight to -/// [`y4m::EncoderBuilder::append_vendor_extension`], which adds the `X` -/// back when it writes the output header. -/// -/// This is how whatever colorspace and range tags the source declared -/// reach the output instead of being dropped. -/// -/// A token that [`y4m::VendorExtensionString`] rejects, which means one -/// containing a space, is skipped rather than failing the run. +/// The leading `X` is stripped because [y4m::EncoderBuilder::append_vendor_extension] adds it back. +/// A token that [y4m::VendorExtensionString] rejects, one containing a space, is skipped rather +/// than failing the run. pub fn y4m_vendor_extensions(raw_params: &[u8]) -> Vec { raw_params - .split(|&b| b == b' ') - .filter(|tok| tok.first() == Some(&b'X')) - .filter_map(|tok| y4m::VendorExtensionString::new(tok[1..].to_vec()).ok()) + .split(|&byte| byte == b' ') + .filter(|token| token.first() == Some(&b'X')) + .filter_map(|token| { + let value = token[1..].to_vec(); + y4m::VendorExtensionString::new(value).ok() + }) .collect() } @@ -77,47 +71,52 @@ mod colorspace_tests { (Subsampling::Yuv444, Depth::Twelve), ]; - for (sub, depth) in combos { - let cs = subsampling_to_y4m(sub, depth); - let (rsub, rdepth) = subsampling_from_y4m(cs).expect("should map back"); + for (subsampling, depth) in combos { + let colorspace = subsampling_to_y4m(subsampling, depth); + let (mapped_subsampling, mapped_depth) = + subsampling_from_y4m(colorspace).expect("should map back"); - assert_eq!(rsub, sub, "subsampling lost for {cs:?}"); - assert_eq!(rdepth, depth, "depth lost for {cs:?}"); + assert_eq!( + mapped_subsampling, subsampling, + "subsampling lost for {colorspace:?}" + ); + assert_eq!(mapped_depth, depth, "depth lost for {colorspace:?}"); } } #[test] fn ten_bit_420_maps_to_c420p10() { - // `y4m::Colorspace` derives only `Debug, Clone, Copy`, not - // `PartialEq`, so `assert_eq!` won't compile here. - assert!(matches!( - subsampling_to_y4m(Subsampling::Yuv420, Depth::Ten), - y4m::Colorspace::C420p10 - )); + let colorspace = subsampling_to_y4m(Subsampling::Yuv420, Depth::Ten); + + // `y4m::Colorspace` does not derive `PartialEq`, so `assert_eq!` won't compile here. + assert!(matches!(colorspace, y4m::Colorspace::C420p10)); } #[test] fn eight_bit_420_variants_all_map_to_yuv420_eight() { - for cs in [ + for colorspace in [ y4m::Colorspace::C420, y4m::Colorspace::C420jpeg, y4m::Colorspace::C420paldv, y4m::Colorspace::C420mpeg2, ] { - let (sub, depth) = subsampling_from_y4m(cs).expect("should map"); - assert_eq!(sub, Subsampling::Yuv420); + let (subsampling, depth) = subsampling_from_y4m(colorspace).expect("should map"); + + assert_eq!(subsampling, Subsampling::Yuv420); assert_eq!(depth, Depth::Eight); } } #[test] fn grayscale_colorspaces_are_rejected_with_a_clear_message() { - for cs in [y4m::Colorspace::Cmono, y4m::Colorspace::Cmono12] { - let err = subsampling_from_y4m(cs).expect_err("grayscale should be rejected"); - let msg = err.to_string(); + for colorspace in [y4m::Colorspace::Cmono, y4m::Colorspace::Cmono12] { + let err = subsampling_from_y4m(colorspace).expect_err("grayscale should be rejected"); + let message = err.to_string(); + let colorspace_name = format!("{colorspace:?}"); + assert!( - msg.contains(&format!("{cs:?}")), - "error should name the offending colorspace, got {msg}" + message.contains(&colorspace_name), + "error should name the offending colorspace, got {message}" ); } } diff --git a/av-denoise-core/src/cache.rs b/av-denoise/src/cache.rs similarity index 58% rename from av-denoise-core/src/cache.rs rename to av-denoise/src/cache.rs index 37f9f70..6cbbd94 100644 --- a/av-denoise-core/src/cache.rs +++ b/av-denoise/src/cache.rs @@ -1,37 +1,22 @@ -//! Where CubeCL keeps its compiled kernels. +//! Where CubeCL keeps its compiled kernels //! -//! Compiling this crate's kernels takes about ten seconds, and CubeCL -//! caches nothing on its own. Its cache setting defaults to `None`, so -//! every run recompiles from scratch unless something points it at a -//! directory. +//! Compiling the kernels was measured at about ten seconds, and CubeCL caches nothing by default. +//! [install_compilation_cache] points it at [default_cache_dir], or at the directory named by the +//! `AV_DENOISE_COMPILATION_CACHE` environment variable, so that cost is paid once per machine rather +//! than once per run. Setting that variable to `off` disables caching, +//! which benchmarks want because a warm cache hides the compile cost a first run pays. //! -//! [`install_compilation_cache`] points it at one. By default, that is -//! [`default_cache_dir`], `av-denoise` inside the platform's cache -//! directory, which turns the ten seconds into a cost paid once per -//! machine rather than once per run. +//! Each build of the kernels gets its own subdirectory, and an install removes other builds' +//! subdirectories once they have gone unused for a week. //! -//! The `AV_DENOISE_COMPILATION_CACHE` environment variable overrides the -//! location, which is what CI runs and containers use to put the cache -//! on a mounted volume. Setting it to `off` disables caching entirely, -//! which is what benchmarking wants, because a warm cache hides the -//! compilation cost that a first run pays. -//! -//! The cache root holds one subdirectory per build of this crate's -//! sources, named by a hash of them. A build only reads kernels it -//! compiled itself, and installing removes other builds' subdirectories -//! once they have gone unused for a week. -//! -//! [`install_compilation_cache`] has to run before the first -//! [`Denoiser`](crate::Denoiser) is created, because building a CubeCL -//! client locks the global config. -//! -//! If using `av-denoise` as a library you may want to specify the cache directory -//! path yourself via [`install_compilation_cache_at`] instead. +//! The cache has to be installed before the first [HostDenoiser](crate::HostDenoiser) is created, +//! because building a CubeCL client locks the global config. Library callers can pick the directory +//! themselves with [install_compilation_cache_at]. //! //! ```no_run //! # fn main() -> Result<(), Box> { //! // Call this at the top of `main`, before any denoiser exists. -//! match av_denoise_core::install_compilation_cache()? { +//! match av_denoise::install_compilation_cache()? { //! Some(path) => println!("caching compiled kernels in {}", path.display()), //! None => println!("kernel caching is off, every run recompiles"), //! } @@ -49,40 +34,34 @@ use cubecl::config::cache::CacheConfig; use cubecl::config::{CubeClRuntimeConfig, RuntimeConfig}; use etcetera::base_strategy::{BaseStrategy, choose_base_strategy}; -/// The environment variable that overrides where compiled kernels are -/// cached, or turns caching off. +/// The environment variable that overrides where compiled kernels are cached, or turns caching off. pub const COMPILATION_CACHE_ENV: &str = "AV_DENOISE_COMPILATION_CACHE"; /// Where compiled kernels are cached, once an install has settled it. static CACHE_DIR: OnceLock> = OnceLock::new(); -/// The directory name this crate uses inside the user's cache directory. const CACHE_DIR_NAME: &str = "av-denoise"; -/// A hash of this crate's sources, naming this build's subdirectory of -/// the cache root. +/// A hash of the kernel sources that names this build's subdirectory of the cache root. /// -/// CubeCL keys cached kernels by their signature rather than their body, -/// so each build of the kernels needs a directory of its own. -pub(crate) const KERNEL_HASH: &str = env!("AV_DENOISE_KERNEL_HASH"); +/// CubeCL keys cached kernels by their signature rather than their body, so each build of the +/// kernels needs a directory of its own. +pub(crate) const KERNEL_HASH: &str = av_denoise_core::KERNEL_HASH; -/// How long another build's subdirectory goes unused before an install -/// removes it. +/// How long another build's subdirectory can go unused before an install removes it. const STALE_BUILD_AGE: Duration = Duration::from_secs(7 * 24 * 60 * 60); -/// The values of [`COMPILATION_CACHE_ENV`] that turn caching off. +/// The values of [COMPILATION_CACHE_ENV] that turn caching off, compared without regard to case. /// -/// Compared without regard to case. `off` is the documented spelling and -/// the others are here so that a reasonable guess does not silently -/// create a directory named `0`. +/// `off` is the documented spelling. The others stop a reasonable guess from silently creating a +/// directory named `0`. const DISABLE_WORDS: [&str; 4] = ["off", "0", "false", "none"]; /// Something went wrong installing the kernel cache. #[derive(Debug, thiserror::Error)] pub enum CacheError { - /// The CubeCL global config was already set up before this helper - /// ran, so the cache directory can no longer be installed. - #[error("CubeCL global config already initialized. Install the cache before any Denoiser::create")] + /// The CubeCL global config was set up before the cache could be installed. + #[error("CubeCL global config already initialized. Install the cache before any HostDenoiser::create")] AlreadyInitialised, /// The cache directory does not exist and could not be created. #[error("cannot create the kernel cache directory {path}", path = path.display())] @@ -95,32 +74,30 @@ pub enum CacheError { /// Where compiled kernels go. /// -/// [`Disabled`](CacheLocation::Disabled) is reachable only when -/// [`COMPILATION_CACHE_ENV`] names one of [`DISABLE_WORDS`]. Every -/// platform has a default directory, so nothing else produces it. +/// [Disabled](CacheLocation::Disabled) is reachable only when [COMPILATION_CACHE_ENV] names one of +/// [DISABLE_WORDS]. Every platform has a default directory, so nothing else produces it. #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) enum CacheLocation { - /// Nothing is cached and every run recompiles. Disabled, - /// Compiled kernels are written under this directory. Dir(PathBuf), } /// The directory compiled kernels are cached in when nothing overrides it. pub fn default_cache_dir() -> PathBuf { - let platform_cache = choose_base_strategy().ok().map(|s| s.cache_dir()); + let platform_cache = choose_base_strategy().ok().map(|strategy| strategy.cache_dir()); if platform_cache.is_none() { tracing::warn!("no platform cache directory available, falling back to the temporary directory"); } - resolve_default_dir(platform_cache, std::env::temp_dir()) + + let temp_dir = std::env::temp_dir(); + resolve_default_dir(platform_cache, temp_dir) } -fn resolve_default_dir(platform_cache: Option, temp: PathBuf) -> PathBuf { - platform_cache.unwrap_or(temp).join(CACHE_DIR_NAME) +fn resolve_default_dir(platform_cache: Option, temp_dir: PathBuf) -> PathBuf { + platform_cache.unwrap_or(temp_dir).join(CACHE_DIR_NAME) } -/// Decides where compiled kernels go, from a default and the environment -/// override alone. +/// Decides where compiled kernels go from a default and the environment override alone. pub(crate) fn resolve_cache_location( env: Option<&OsStr>, default: impl FnOnce() -> PathBuf, @@ -129,27 +106,29 @@ pub(crate) fn resolve_cache_location( let text = raw.to_string_lossy(); let trimmed = text.trim(); if !trimmed.is_empty() { - if DISABLE_WORDS.iter().any(|w| trimmed.eq_ignore_ascii_case(w)) { + if DISABLE_WORDS + .iter() + .any(|word| trimmed.eq_ignore_ascii_case(word)) + { return CacheLocation::Disabled; } - // `raw.to_str()` fails only when the value is not UTF-8, in - // which case it cannot be trimmed portably, so it is used - // unchanged rather than dropped. + + // A value that is not UTF-8 cannot be trimmed portably, so it is used unchanged. let dir = match raw.to_str() { - Some(s) => PathBuf::from(s.trim()), + Some(utf8) => PathBuf::from(utf8.trim()), None => PathBuf::from(raw), }; return CacheLocation::Dir(dir); } } - CacheLocation::Dir(default()) + let dir = default(); + CacheLocation::Dir(dir) } -/// Points CubeCL's compilation and autotune caches at this build's -/// subdirectory of `dir`, creating it if it does not exist. +/// Points CubeCL's compilation and autotune caches at this build's subdirectory of `dir`. /// -/// This is the entry point for a caller using this crate directly. +/// The subdirectory is created if it does not exist. pub fn install_compilation_cache_at(dir: &Path) -> Result<(), CacheError> { let build_dir = build_cache_dir(dir); if let Err(source) = std::fs::create_dir_all(&build_dir) { @@ -166,15 +145,13 @@ pub fn install_compilation_cache_at(dir: &Path) -> Result<(), CacheError> { Ok(()) } -/// Installs the cache at [`default_cache_dir`], or the directory -/// [`COMPILATION_CACHE_ENV`] names, unless that variable turns caching off. +/// Installs the cache at [default_cache_dir], or the directory [COMPILATION_CACHE_ENV] names. /// -/// Returns this build's subdirectory of that root, or `Ok(None)` only -/// when the variable disables caching. +/// Returns this build's subdirectory of that root, or `Ok(None)` when the variable disables caching. /// -/// A directory that cannot be created is reported through `tracing` logs and then -/// ignored, because denoising works without a cache. A caller that wants a -/// creation failure reported should use [`install_compilation_cache_at`]. +/// A directory that cannot be created is logged through `tracing` and then ignored, because denoising +/// works without a cache. A caller that wants that failure reported should use +/// [install_compilation_cache_at]. pub fn install_compilation_cache() -> Result, CacheError> { let env_override = std::env::var_os(COMPILATION_CACHE_ENV); let location = resolve_cache_location(env_override.as_deref(), default_cache_dir); @@ -196,11 +173,8 @@ pub fn install_compilation_cache() -> Result, CacheError> { set_runtime_config(&path)?; tidy_cache_root(&root, &path); - // Only an install that reached the config has a directory worth - // recording. A directory that could not be created has already - // answered `Ok(None)` above, and latching that would leave - // `compilation_cache_dir` saying `None` for the rest of the process - // even if a later call succeeded. + // Only a successful install is recorded, so a failed one never latches `None` for the rest + // of the process. let _ = CACHE_DIR.set(Some(path.clone())); Ok(Some(path)) @@ -208,16 +182,12 @@ pub fn install_compilation_cache() -> Result, CacheError> { /// Points CubeCL at a cache the first time it runs, and reports where. /// -/// Written for callers that are not a `main`, such as the VapourSynth -/// plugin, where filter creation is the earliest hook there is and runs -/// once per filter rather than once per process. +/// This suits callers without a `main`, such as the VapourSynth plugin, where filter creation is +/// the earliest hook and runs once per filter. /// -/// A failure to install is reported through `tracing` and then ignored, -/// because a plugin that refuses to denoise is worse than one that -/// recompiles. Failing here means something else configured CubeCL -/// first, which may well have pointed it at a cache of its own. What is -/// lost is knowing where that cache is, which is why this answers `None` -/// rather than guessing. +/// A failure to install is logged through `tracing` and then ignored, because a plugin that refuses +/// to denoise is worse than one that recompiles. Failing means something else configured CubeCL +/// first, possibly with a cache of its own, so this answers `None` rather than guessing where it is. pub fn install_compilation_cache_once() -> Option<&'static Path> { static ONCE: Once = Once::new(); @@ -235,43 +205,32 @@ pub fn install_compilation_cache_once() -> Option<&'static Path> { /// This build's directory of cached kernels. /// -/// `None` until an install succeeds, and `None` for good when caching is -/// off. Callers that want to sit alongside the cache, such as -/// [`WarmUp`](crate::WarmUp), have nowhere to put their own files until -/// this answers. +/// `None` until an install succeeds, and `None` for good when caching is off. pub fn compilation_cache_dir() -> Option<&'static Path> { CACHE_DIR.get()?.as_deref() } -/// Points the CubeCL global config at `path`. -/// -/// Split out because both entry points build the same config once they -/// have settled on a directory that exists. fn set_runtime_config(path: &Path) -> Result<(), CacheError> { - let mut cfg = CubeClRuntimeConfig::from_current_dir().override_from_env(); - cfg.compilation.cache = Some(CacheConfig::File(path.to_path_buf())); - cfg.autotune.cache = CacheConfig::File(path.to_path_buf()); - - // `RuntimeConfig::set` panics if the singleton is already set up. - // Catching that turns an abort into a typed error for the caller. - // - // CubeCL does not expose a fallible version of this call, so the - // panic is the only signal available. - std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { - CubeClRuntimeConfig::set(cfg); - })) - .map_err(|_| CacheError::AlreadyInitialised) + let cache_dir = path.to_path_buf(); + let mut config = CubeClRuntimeConfig::from_current_dir().override_from_env(); + config.compilation.cache = Some(CacheConfig::File(cache_dir.clone())); + config.autotune.cache = CacheConfig::File(cache_dir); + + // `RuntimeConfig::set` panics if the singleton is already set up, and CubeCL has no fallible + // version, so the panic is caught and turned into a typed error. + let install = std::panic::AssertUnwindSafe(|| { + CubeClRuntimeConfig::set(config); + }); + std::panic::catch_unwind(install).map_err(|_| CacheError::AlreadyInitialised) } fn build_cache_dir(root: &Path) -> PathBuf { root.join(KERNEL_HASH) } -/// Marks this build's subdirectory as used and removes stale ones from -/// other builds. +/// Marks this build's subdirectory as used and removes stale ones from other builds. /// -/// Both steps are best effort, because a cache that is tidied late still -/// works. +/// Both steps are best effort, because a cache that is tidied late still works. fn tidy_cache_root(root: &Path, build_dir: &Path) { let now = SystemTime::now(); @@ -282,19 +241,19 @@ fn tidy_cache_root(root: &Path, build_dir: &Path) { prune_stale_builds(root, KERNEL_HASH, now); } -/// Removes subdirectories of `root` left by other builds that have gone -/// unused for longer than [STALE_BUILD_AGE](crate::cache::STALE_BUILD_AGE). +/// Removes other builds' subdirectories of `root` that have gone unused for longer than +/// [STALE_BUILD_AGE]. /// -/// Only directories named like a build hash are candidates, so anything -/// else in the root is left alone. Every error is ignored. +/// Only directories named like a build hash are candidates, so anything else in the root is left +/// alone. Every error is ignored. fn prune_stale_builds(root: &Path, current_hash: &str, now: SystemTime) { let Ok(entries) = std::fs::read_dir(root) else { return; }; for entry in entries.flatten() { - let name = entry.file_name(); - let Some(name) = name.to_str() else { + let file_name = entry.file_name(); + let Some(name) = file_name.to_str() else { continue; }; @@ -302,8 +261,8 @@ fn prune_stale_builds(root: &Path, current_hash: &str, now: SystemTime) { continue; } - // `DirEntry::metadata` does not follow symlinks, so a link is - // never mistaken for a build directory. + // `DirEntry::metadata` does not follow symlinks, so a link is never mistaken for a build + // directory. let Ok(metadata) = entry.metadata() else { continue; }; @@ -313,7 +272,8 @@ fn prune_stale_builds(root: &Path, current_hash: &str, now: SystemTime) { let age = now.duration_since(modified).unwrap_or_default(); if metadata.is_dir() && age > STALE_BUILD_AGE { - let _ = std::fs::remove_dir_all(entry.path()); + let stale_path = entry.path(); + let _ = std::fs::remove_dir_all(stale_path); } } } @@ -340,102 +300,108 @@ mod tests { #[test] fn with_nothing_set_the_default_wins() { + let location = resolve(None, "/home/u/.cache/av-denoise"); + assert_eq!( - resolve(None, "/home/u/.cache/av-denoise"), + location, CacheLocation::Dir(PathBuf::from("/home/u/.cache/av-denoise")), ); } #[test] fn an_explicit_path_overrides_the_default() { - assert_eq!( - resolve(Some("/mnt/cache"), "/home/u/.cache/av-denoise"), - CacheLocation::Dir(PathBuf::from("/mnt/cache")), - ); + let location = resolve(Some("/mnt/cache"), "/home/u/.cache/av-denoise"); + + assert_eq!(location, CacheLocation::Dir(PathBuf::from("/mnt/cache"))); } #[test] fn the_disable_words_turn_caching_off() { for word in ["off", "OFF", "Off", "0", "false", "FALSE", "none", " off "] { - assert_eq!( - resolve(Some(word), "/home/u/.cache/av-denoise"), - CacheLocation::Disabled, - "{word} should disable caching", - ); + let location = resolve(Some(word), "/home/u/.cache/av-denoise"); + assert_eq!(location, CacheLocation::Disabled, "{word} should disable caching"); } } #[test] fn a_path_containing_a_disable_word_is_still_a_path() { - assert_eq!( - resolve(Some("/tmp/offsite"), "/home/u/.cache/av-denoise"), - CacheLocation::Dir(PathBuf::from("/tmp/offsite")), - ); + let location = resolve(Some("/tmp/offsite"), "/home/u/.cache/av-denoise"); + + assert_eq!(location, CacheLocation::Dir(PathBuf::from("/tmp/offsite"))); } #[test] fn an_empty_variable_takes_the_default() { + let empty = resolve(Some(""), "/home/u/.cache/av-denoise"); + let blank = resolve(Some(" "), "/home/u/.cache/av-denoise"); + assert_eq!( - resolve(Some(""), "/home/u/.cache/av-denoise"), + empty, CacheLocation::Dir(PathBuf::from("/home/u/.cache/av-denoise")), ); assert_eq!( - resolve(Some(" "), "/home/u/.cache/av-denoise"), + blank, CacheLocation::Dir(PathBuf::from("/home/u/.cache/av-denoise")), ); } #[test] fn a_padded_explicit_path_is_trimmed() { - assert_eq!( - resolve(Some(" /mnt/cache "), "/home/u/.cache/av-denoise"), - CacheLocation::Dir(PathBuf::from("/mnt/cache")), - ); + let location = resolve(Some(" /mnt/cache "), "/home/u/.cache/av-denoise"); + + assert_eq!(location, CacheLocation::Dir(PathBuf::from("/mnt/cache"))); } - /// A non-UTF-8 value cannot be trimmed portably, so it is used unchanged. #[cfg(unix)] #[test] fn a_non_utf8_override_is_carried_through_untrimmed() { - use std::ffi::OsString; use std::os::unix::ffi::OsStringExt; - // `0xFF` is not valid UTF-8 in any position, so `bytes` never - // decodes to a `&str`. + // `0xFF` is not valid UTF-8 in any position, so `bytes` never decodes to a `&str`. let bytes = vec![b'/', b'm', b'n', b't', b'/', 0xFF, b'x']; let env = OsString::from_vec(bytes.clone()); - assert_eq!( - resolve_cache_location(Some(&env), || PathBuf::from("/home/u/.cache/av-denoise")), - CacheLocation::Dir(PathBuf::from(OsString::from_vec(bytes))), - ); + let expected_os = OsString::from_vec(bytes); + let expected = PathBuf::from(expected_os); + + let location = resolve_cache_location(Some(&env), || PathBuf::from("/home/u/.cache/av-denoise")); + + assert_eq!(location, CacheLocation::Dir(expected)); } #[test] fn resolve_default_dir_joins_the_platform_cache_directory() { - assert_eq!( - resolve_default_dir(Some(PathBuf::from("/home/u/.cache")), PathBuf::from("/tmp"),), - PathBuf::from("/home/u/.cache/av-denoise"), - ); + let platform_cache = Some(PathBuf::from("/home/u/.cache")); + let temp_dir = PathBuf::from("/tmp"); + + let resolved = resolve_default_dir(platform_cache, temp_dir); + + assert_eq!(resolved, PathBuf::from("/home/u/.cache/av-denoise")); } #[test] fn resolve_default_dir_falls_back_to_the_temporary_directory() { - assert_eq!( - resolve_default_dir(None, PathBuf::from("/tmp")), - PathBuf::from("/tmp/av-denoise"), - ); + let temp_dir = PathBuf::from("/tmp"); + + let resolved = resolve_default_dir(None, temp_dir); + + assert_eq!(resolved, PathBuf::from("/tmp/av-denoise")); } #[test] fn the_build_cache_dir_is_the_root_joined_with_the_kernel_hash() { let root = PathBuf::from("/home/u/.cache/av-denoise"); let expected = root.join(KERNEL_HASH); - assert_eq!(build_cache_dir(&root), expected); + + let build_dir = build_cache_dir(&root); + + assert_eq!(build_dir, expected); } #[test] fn the_kernel_hash_is_sixteen_lowercase_hex_chars() { - assert!(is_build_hash(KERNEL_HASH), "{KERNEL_HASH} is not a build hash"); + let looks_like_hash = is_build_hash(KERNEL_HASH); + + assert!(looks_like_hash, "{KERNEL_HASH} is not a build hash"); } #[test] @@ -447,7 +413,8 @@ mod tests { "0123456789abcdef0", "warm-0123456789ab", ] { - assert!(!is_build_hash(name), "{name} should not look like a build hash"); + let looks_like_hash = is_build_hash(name); + assert!(!looks_like_hash, "{name} should not look like a build hash"); } } @@ -495,7 +462,9 @@ mod tests { fn pruning_a_missing_root_does_nothing() { let root = tempfile::tempdir().expect("temp dir"); let missing = root.path().join("missing"); - prune_stale_builds(&missing, KERNEL_HASH, SystemTime::now()); + let now = SystemTime::now(); + + prune_stale_builds(&missing, KERNEL_HASH, now); } #[cfg(unix)] diff --git a/av-denoise/src/host/depth.rs b/av-denoise/src/host/depth.rs new file mode 100644 index 0000000..42bbd00 --- /dev/null +++ b/av-denoise/src/host/depth.rs @@ -0,0 +1,113 @@ +use av_denoise_core::SampleFormat; + +/// Bit depth of a source's samples. +/// +/// Normalisation divides by [Depth::max_value], so a value in normalised units means the same thing +/// at every depth. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Depth { + Eight, + Ten, + Twelve, +} + +/// Returned when a source declares a bit depth the denoiser does not handle. +#[derive(Debug, thiserror::Error)] +#[error("unsupported bit depth {0}, av-denoise supports 8, 10, and 12-bit")] +pub struct UnsupportedDepthError(pub usize); + +impl Depth { + /// Maps a declared bit depth onto a [Depth]. + pub fn from_bits(bits: usize) -> Result { + match bits { + 8 => Ok(Depth::Eight), + 10 => Ok(Depth::Ten), + 12 => Ok(Depth::Twelve), + other => Err(UnsupportedDepthError(other)), + } + } + + /// Bits per sample. + pub fn bits(self) -> usize { + match self { + Depth::Eight => 8, + Depth::Ten => 10, + Depth::Twelve => 12, + } + } + + /// Bytes each sample takes up on the wire. + /// + /// Depths above 8 use a little-endian 16-bit word. + pub fn bytes_per_sample(self) -> usize { + match self { + Depth::Eight => 1, + Depth::Ten | Depth::Twelve => 2, + } + } + + /// The largest sample value this depth can hold, which is also the normalisation divisor. + pub fn max_value(self) -> f32 { + ((1u32 << self.bits()) - 1) as f32 + } + + /// The sample value that means neutral chroma at this depth. + pub fn neutral_chroma(self) -> u16 { + 1 << (self.bits() - 1) + } + + /// The plane format an engine reads and writes at this depth. + pub(crate) fn sample_format(self) -> SampleFormat { + match self { + Depth::Eight => SampleFormat::U8, + Depth::Ten => SampleFormat::U16 { depth: 10 }, + Depth::Twelve => SampleFormat::U16 { depth: 12 }, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn from_bits_accepts_supported_depths() { + let eight_bit = Depth::from_bits(8).unwrap(); + let ten_bit = Depth::from_bits(10).unwrap(); + let twelve_bit = Depth::from_bits(12).unwrap(); + + assert_eq!(eight_bit, Depth::Eight); + assert_eq!(ten_bit, Depth::Ten); + assert_eq!(twelve_bit, Depth::Twelve); + } + + #[test] + fn from_bits_rejects_unsupported_depths() { + for bits in [0, 7, 9, 11, 16] { + let result = Depth::from_bits(bits); + assert!(result.is_err(), "{bits} bits should be rejected"); + } + } + + #[test] + fn depth_properties_match_the_format() { + assert_eq!(Depth::Eight.bytes_per_sample(), 1); + assert_eq!(Depth::Ten.bytes_per_sample(), 2); + assert_eq!(Depth::Twelve.bytes_per_sample(), 2); + + assert_eq!(Depth::Eight.max_value(), 255.0); + assert_eq!(Depth::Ten.max_value(), 1023.0); + assert_eq!(Depth::Twelve.max_value(), 4095.0); + + assert_eq!(Depth::Eight.neutral_chroma(), 128); + assert_eq!(Depth::Ten.neutral_chroma(), 512); + assert_eq!(Depth::Twelve.neutral_chroma(), 2048); + } + + #[test] + fn sample_format_matches_the_depth() { + assert_eq!(Depth::Eight.sample_format(), SampleFormat::U8); + assert_eq!(Depth::Ten.sample_format(), SampleFormat::U16 { depth: 10 }); + assert_eq!(Depth::Twelve.sample_format(), SampleFormat::U16 { depth: 12 }); + } +} diff --git a/av-denoise/src/host/io.rs b/av-denoise/src/host/io.rs new file mode 100644 index 0000000..fcb66a8 --- /dev/null +++ b/av-denoise/src/host/io.rs @@ -0,0 +1,50 @@ +use std::future::Future; +use std::pin::Pin; + +use cubecl::Runtime; +use cubecl::bytes::Bytes; +use cubecl::prelude::ComputeClient; +use cubecl::server::{Handle, ServerError}; + +/// A readback in flight, resolving to one buffer per handle read. +pub(crate) type ReadFuture = Pin, ServerError>> + Send>>; + +/// Uploads and reads back planes on whichever runtime an engine was built on. +pub(crate) trait PlaneIo: Send { + fn upload(&self, bytes: &[u8]) -> Handle; + + fn allocate(&self, bytes: usize) -> Handle; + + /// Starts reading `handles` back to the host. + /// + /// Nothing is read until the returned future is first polled. + fn read(&self, handles: Vec) -> ReadFuture; +} + +pub(crate) struct ClientIo { + client: ComputeClient, +} + +impl ClientIo { + pub(crate) fn new(client: ComputeClient) -> Self { + Self { client } + } +} + +impl PlaneIo for ClientIo { + fn upload(&self, bytes: &[u8]) -> Handle { + self.client.create_from_slice(bytes) + } + + fn allocate(&self, bytes: usize) -> Handle { + self.client.empty(bytes) + } + + fn read(&self, handles: Vec) -> ReadFuture { + // The future owns its own client, so it stays valid after the denoiser that started it is dropped. + let client = self.client.clone(); + let future = async move { client.read_async(handles).await }; + + Box::pin(future) + } +} diff --git a/av-denoise/src/host/mod.rs b/av-denoise/src/host/mod.rs new file mode 100644 index 0000000..26daa4d --- /dev/null +++ b/av-denoise/src/host/mod.rs @@ -0,0 +1,454 @@ +mod depth; +pub(crate) mod io; +mod options; +mod pending; + +#[cfg(test)] +mod tests; + +use std::collections::VecDeque; + +use av_denoise_core::{DevicePlane, Engine, GrainChunk, WindowSpan}; +use cubecl::server::Handle; + +pub use self::depth::{Depth, UnsupportedDepthError}; +use self::io::PlaneIo; +pub use self::options::{Algorithm, DenoiserOptions}; +pub(crate) use self::pending::{Pending, TryWait}; +use crate::backend::accelerate::Accelerator; +use crate::backend::{Device, build_engine}; + +/// How many readbacks a [HostDenoiser] keeps in flight at once. +pub const MAX_PENDING: usize = 2; + +/// How many output plane sets a [HostDenoiser] rotates through. +const OUTPUT_SETS: usize = 2; + +/// Each frame in flight reads its own output set, so a set is never rewritten before its readback lands. +const _: () = assert!(OUTPUT_SETS >= MAX_PENDING); + +/// Errors reported by a [HostDenoiser]. +#[derive(Debug, thiserror::Error)] +pub enum DenoiserError { + /// [MAX_PENDING] frames are waiting to be collected. + /// + /// Call [HostDenoiser::recv] or [HostDenoiser::try_recv], then retry the same push. + #[error("denoiser queue is full, collect the pending frame before pushing more")] + QueueFull, + /// An earlier failure means later output would not line up with its input. + /// + /// Call [HostDenoiser::reset_stream] to start a fresh stream, or drop the denoiser. + #[error("denoiser failed earlier, reset the stream before using it again")] + Poisoned, + /// None of the accelerators in the priority list could be started. + #[error("no accelerator from the priority list is available")] + NoAcceleratorAvailable, + /// The engine rejected its options, its planes or a call. + #[error(transparent)] + Engine(#[from] av_denoise_core::Error), + /// Anything else, such as a failed readback. + #[error(transparent)] + Other(#[from] anyhow::Error), +} + +impl DenoiserError { + /// Whether this error leaves the stream unusable until [HostDenoiser::reset_stream]. + /// + /// Rejected planes are caught before any state changes and never reach this check. A full queue + /// and a context frame sent after a push change no state either. Any other failure may leave the + /// host partly advanced, so it poisons. + fn poisons_stream(&self) -> bool { + !matches!( + self, + DenoiserError::QueueFull + | DenoiserError::Poisoned + | DenoiserError::Engine(av_denoise_core::Error::ContextAfterPush) + ) + } +} + +/// A stateful denoiser that cleans a stream of frames held as wire bytes. +/// +/// Push frames in order with [Self::push] and collect the cleaned ones with [Self::recv] or +/// [Self::try_recv]. At the end of the stream call [Self::flush] to drain the temporal tail. +/// +/// Each frame is one byte buffer per plane. Luma takes Y, chroma takes U and V, and YUV takes all +/// three at one size. Every output frame comes back in the same layout. +/// +/// ```no_run +/// use av_denoise::accelerate::Accelerator; +/// use av_denoise::{ChannelMode, DenoiserOptions, DenoisingMode, Device, HostDenoiser}; +/// +/// # fn main() -> Result<(), Box> { +/// let options = DenoiserOptions::builder() +/// .channel_mode(ChannelMode::Luma) +/// .mode(DenoisingMode::Temporal { radius: 2 }) +/// .build(); +/// +/// let mut denoiser = HostDenoiser::create(&[Accelerator::Vulkan], &Device::Default, 1920, 1080, options)?; +/// +/// let frame = vec![128u8; 1920 * 1080]; +/// let mut cleaned = Vec::new(); +/// +/// for _ in 0..10 { +/// denoiser.push(&[&frame])?; +/// +/// if let Some(planes) = denoiser.recv()? { +/// cleaned.push(planes); +/// } +/// } +/// +/// denoiser.flush(|planes| cleaned.push(planes))?; +/// # Ok(()) +/// # } +/// ``` +pub struct HostDenoiser { + engine: Box, + io: Box, + accelerator: Accelerator, + width: u32, + height: u32, + temporal_radius: u32, + plane_count: usize, + /// The wire byte length of one plane. + plane_length: usize, + /// The bytes one device plane takes up, rounded up to whole words. + device_plane_bytes: usize, + /// Output plane sets, allocated on first use. + outputs: Vec>, + next_output: usize, + pending: VecDeque, + /// Set once a call fails with an error that [DenoiserError::poisons_stream]. + poisoned: bool, +} + +impl HostDenoiser { + /// Builds a denoiser on the first accelerator in `accelerators` that works. + /// + /// `device` picks a non-default device on the chosen runtime. + /// + /// # Thread stack size + /// + /// cubecl runs kernel codegen on its own worker thread, which gets Rust's default stack. A + /// `search_radius` above 4 can overflow it, so callers using one should call + /// [raise_codegen_stack_limit](crate::raise_codegen_stack_limit) at the top of `main`. + pub fn create( + accelerators: &[Accelerator], + device: &Device, + width: u32, + height: u32, + options: DenoiserOptions, + ) -> Result { + let temporal_radius = options.temporal_radius(); + let is_nl4d = matches!(options.algorithm, Algorithm::Nl4d(_)); + + if is_nl4d && temporal_radius == 0 { + let error = anyhow::anyhow!( + "nl4d needs a temporal window, set DenoiserOptions::mode to DenoisingMode::Temporal" + ); + return Err(DenoiserError::Other(error)); + } + + let spec = options.algorithm.engine_spec(&options, width, height); + let built = build_engine(accelerators, device, spec)?; + + let format = options.depth.sample_format(); + let pixels = width as usize * height as usize; + let device_plane_bytes = format.plane_bytes(pixels as u64) as usize; + let plane_length = pixels * options.depth.bytes_per_sample(); + let plane_count = options.channel_mode.count() as usize; + + Ok(Self { + engine: built.engine, + io: built.io, + accelerator: built.accelerator, + width, + height, + temporal_radius, + plane_count, + plane_length, + device_plane_bytes, + outputs: Vec::with_capacity(OUTPUT_SETS), + next_output: 0, + pending: VecDeque::with_capacity(MAX_PENDING), + poisoned: false, + }) + } + + pub fn selected_accelerator(&self) -> Accelerator { + self.accelerator + } + + pub fn width(&self) -> u32 { + self.width + } + + pub fn height(&self) -> u32 { + self.height + } + + pub fn temporal_radius(&self) -> u32 { + self.temporal_radius + } + + /// How many frames behind and ahead of a target frame are needed to denoise it. + pub fn window_span(&self) -> WindowSpan { + self.engine.window_span() + } + + /// Uploads one frame and starts denoising any frame it completes. + /// + /// Returns [DenoiserError::QueueFull] once [MAX_PENDING] frames wait to be collected, which does + /// not poison the denoiser. A rejected plane leaves it usable too. Any other failure poisons it + /// until [Self::reset_stream]. + pub fn push(&mut self, planes: &[&[u8]]) -> Result<(), DenoiserError> { + if self.poisoned { + return Err(DenoiserError::Poisoned); + } + + if self.pending.len() >= MAX_PENDING { + return Err(DenoiserError::QueueFull); + } + + self.check_planes(planes)?; + + let result = self.push_inner(planes); + self.poison_on_error(result) + } + + fn push_inner(&mut self, planes: &[&[u8]]) -> Result<(), DenoiserError> { + let handles = self.upload(planes); + let device_planes = device_planes(&handles, self.width, self.height); + let ready = self.engine.push(&device_planes)?; + + // Emitting past the free output sets would overwrite a set whose readback has not landed. + let in_flight = self.pending.len(); + if in_flight + ready > OUTPUT_SETS { + let error = anyhow::anyhow!( + "push readied {ready} frames with {in_flight} in flight, past the {OUTPUT_SETS} output sets" + ); + return Err(DenoiserError::Other(error)); + } + + for _ in 0..ready { + let slot = self.next_output; + self.next_output = (slot + 1) % OUTPUT_SETS; + + let pending = self.emit_into_slot(slot)?; + self.pending.push_back(pending); + } + + Ok(()) + } + + /// Uploads one frame as context only, without producing output. + /// + /// A stream that starts with this picks up mid-clip, so nl4d runs no head passes for it. + pub fn push_priming(&mut self, planes: &[&[u8]]) -> Result<(), DenoiserError> { + if self.poisoned { + return Err(DenoiserError::Poisoned); + } + + self.check_planes(planes)?; + + let result = self.push_priming_inner(planes); + self.poison_on_error(result) + } + + fn push_priming_inner(&mut self, planes: &[&[u8]]) -> Result<(), DenoiserError> { + let handles = self.upload(planes); + let device_planes = device_planes(&handles, self.width, self.height); + self.engine.push_context(&device_planes)?; + + Ok(()) + } + + /// Blocks until the oldest frame in flight lands and returns it, one buffer per plane. + /// + /// Returns `Ok(None)` when nothing is in flight. + pub fn recv(&mut self) -> Result>>, DenoiserError> { + if self.poisoned { + return Err(DenoiserError::Poisoned); + } + + let result = self.recv_inner(); + self.poison_on_error(result) + } + + fn recv_inner(&mut self) -> Result>>, DenoiserError> { + let Some(pending) = self.pending.pop_front() else { + return Ok(None); + }; + + let planes = pending.wait()?; + Ok(Some(planes)) + } + + /// Polls the oldest frame in flight once. + /// + /// Returns `Ok(None)` both when nothing is in flight and when the readback has not landed. Only + /// the wgpu backends avoid blocking here. + /// + /// Once polled, a frame's readback has to finish. Dropping the denoiser or calling + /// [Self::reset_stream] before it lands blocks until it does. + pub fn try_recv(&mut self) -> Result>>, DenoiserError> { + if self.poisoned { + return Err(DenoiserError::Poisoned); + } + + let result = self.try_recv_inner(); + self.poison_on_error(result) + } + + fn try_recv_inner(&mut self) -> Result>>, DenoiserError> { + let Some(pending) = self.pending.pop_front() else { + return Ok(None); + }; + + match pending.try_wait()? { + TryWait::Ready(planes) => Ok(Some(planes)), + TryWait::NotReady(pending) => { + self.pending.push_front(pending); + Ok(None) + }, + } + } + + /// Drains the frames in flight and the temporal tail, handing each frame to `sink`. + /// + /// The denoiser is ready for a fresh stream afterwards. + pub fn flush(&mut self, sink: impl FnMut(Vec>)) -> Result<(), DenoiserError> { + if self.poisoned { + return Err(DenoiserError::Poisoned); + } + + let result = self.flush_inner(sink); + self.poison_on_error(result) + } + + fn flush_inner(&mut self, mut sink: impl FnMut(Vec>)) -> Result<(), DenoiserError> { + while let Some(planes) = self.recv_inner()? { + sink(planes); + } + + let tail = self.engine.finish()?; + + for _ in 0..tail { + let pending = self.emit_into_slot(0)?; + let planes = pending.wait()?; + sink(planes); + } + + self.next_output = 0; + + Ok(()) + } + + /// Abandons the current stream, keeping every GPU allocation, and clears any poison. + /// + /// Frames in flight are dropped unread, except one [Self::try_recv] already polled, which is + /// settled first. + pub fn reset_stream(&mut self) { + self.pending.clear(); + self.engine.reset(); + self.next_output = 0; + self.poisoned = false; + } + + /// Reads back the grain chunks measured since the last call, in frame order. + /// + /// Empty unless grain export is on. + pub fn drain_grain_chunks(&mut self) -> Result, DenoiserError> { + if self.poisoned { + return Err(DenoiserError::Poisoned); + } + + let drained = self.engine.drain_grain_chunks(); + let drained = drained.map_err(DenoiserError::from); + + self.poison_on_error(drained) + } + + #[cfg(test)] + pub(crate) fn poison_for_test(&mut self) { + self.poisoned = true; + } + + fn poison_on_error(&mut self, result: Result) -> Result { + if let Err(error) = &result + && error.poisons_stream() + { + self.poisoned = true; + } + + result + } + + /// Rejects a frame with the wrong number of planes or a plane of the wrong length. + fn check_planes(&self, planes: &[&[u8]]) -> Result<(), DenoiserError> { + if planes.len() != self.plane_count { + let message = format!("expected {} planes, got {}", self.plane_count, planes.len()); + let error = av_denoise_core::Error::PlaneMismatch(message); + return Err(DenoiserError::Engine(error)); + } + + for (index, plane) in planes.iter().enumerate() { + if plane.len() != self.plane_length { + let message = format!( + "plane {index} holds {} bytes, expected {}", + plane.len(), + self.plane_length + ); + let error = av_denoise_core::Error::PlaneMismatch(message); + return Err(DenoiserError::Engine(error)); + } + } + + Ok(()) + } + + /// Uploads each plane, padded to whole words. + fn upload(&self, planes: &[&[u8]]) -> Vec { + let mut handles = Vec::with_capacity(planes.len()); + + for plane in planes { + let handle = if plane.len() == self.device_plane_bytes { + self.io.upload(plane) + } else { + let mut padded = plane.to_vec(); + padded.resize(self.device_plane_bytes, 0); + self.io.upload(&padded) + }; + handles.push(handle); + } + + handles + } + + /// Writes the oldest ready frame into output set `slot` and starts reading it back. + fn emit_into_slot(&mut self, slot: usize) -> Result { + while self.outputs.len() <= slot { + let set = (0..self.plane_count) + .map(|_| self.io.allocate(self.device_plane_bytes)) + .collect(); + self.outputs.push(set); + } + + let handles = self.outputs[slot].clone(); + let device_planes = device_planes(&handles, self.width, self.height); + self.engine.emit_into(&device_planes)?; + + let future = self.io.read(handles); + let plane_lengths = vec![self.plane_length; self.plane_count]; + let pending = Pending::new(future, plane_lengths); + + Ok(pending) + } +} + +fn device_planes(handles: &[Handle], width: u32, height: u32) -> Vec> { + handles + .iter() + .map(|handle| DevicePlane::new(handle, width, height)) + .collect() +} diff --git a/av-denoise/src/host/options.rs b/av-denoise/src/host/options.rs new file mode 100644 index 0000000..fc09357 --- /dev/null +++ b/av-denoise/src/host/options.rs @@ -0,0 +1,113 @@ +use av_denoise_core::{ + ChannelMode, + DenoisingMode, + Geometry, + Nl4dOptions, + NlmeansAlgorithm, + NlmeansHqOptions, + NlmeansOptions, +}; + +use super::Depth; +use crate::backend::EngineSpec; + +/// How a [HostDenoiser](crate::HostDenoiser) should be set up. +/// +/// Build one with `DenoiserOptions::builder()`. Every field has a default, so only the parts you care +/// about need naming. +#[derive(Debug, Clone, bon::Builder)] +pub struct DenoiserOptions { + /// Which channels of the frame to denoise. + #[builder(default = ChannelMode::Yuv)] + pub channel_mode: ChannelMode, + /// Whether to clean each frame on its own or across a temporal window. + /// + /// This wins over any mode or temporal radius set inside `algorithm`. + #[builder(default = DenoisingMode::Spacial)] + pub mode: DenoisingMode, + /// Which algorithm to run, along with the settings only that algorithm reads. + #[builder(default)] + pub algorithm: Algorithm, + /// The bit depth of the wire bytes going in and coming out. + #[builder(default = Depth::Eight)] + pub depth: Depth, +} + +impl DenoiserOptions { + /// The temporal radius `mode` asks for. + pub(crate) fn temporal_radius(&self) -> u32 { + match self.mode { + DenoisingMode::Spacial => 0, + DenoisingMode::Temporal { radius } => radius, + } + } +} + +/// Which denoising algorithm to run. +/// +/// Each variant carries its own settings, so a knob one algorithm has no use for cannot be set on it. +#[derive(Debug, Copy, Clone, PartialEq)] +pub enum Algorithm { + /// The fast NLMeans path, with fixed weighting and no noise measurement. + Nlmeans(NlmeansOptions), + /// NLMeans with its weighting matched to the measured noise level. + NlmeansHq(NlmeansHqOptions), + /// Groups 8x8 patches across the motion-compensated temporal window. + Nl4d(Nl4dOptions), +} + +impl Default for Algorithm { + fn default() -> Self { + let options = NlmeansOptions::default(); + Self::Nlmeans(options) + } +} + +impl Algorithm { + /// The engine to build for `options`, at `width` by `height`. + /// + /// `options.mode` is copied into the algorithm's own options, so it wins over whatever they hold. + pub(crate) fn engine_spec(&self, options: &DenoiserOptions, width: u32, height: u32) -> EngineSpec { + let format = options.depth.sample_format(); + let geometry = Geometry { + width, + height, + channels: options.channel_mode, + input: format, + output: format, + }; + + match *self { + Algorithm::Nlmeans(nlm) => { + let nlm = NlmeansOptions { + mode: options.mode, + ..nlm + }; + let algorithm = NlmeansAlgorithm::Fast(nlm); + + EngineSpec::Nlmeans { algorithm, geometry } + }, + Algorithm::NlmeansHq(hq) => { + let nlm = NlmeansOptions { + mode: options.mode, + ..hq.nlm + }; + let hq = NlmeansHqOptions { nlm, ..hq }; + let algorithm = NlmeansAlgorithm::Hq(hq); + + EngineSpec::Nlmeans { algorithm, geometry } + }, + Algorithm::Nl4d(nl4d) => { + let nl4d = Nl4dOptions { + temporal_radius: options.temporal_radius(), + ..nl4d + }; + + EngineSpec::Nl4d { + options: nl4d, + geometry, + } + }, + } + } +} diff --git a/av-denoise/src/host/pending.rs b/av-denoise/src/host/pending.rs new file mode 100644 index 0000000..2299a2c --- /dev/null +++ b/av-denoise/src/host/pending.rs @@ -0,0 +1,95 @@ +use std::task::{Context, Poll, Waker}; + +use cubecl::bytes::Bytes; + +use super::io::ReadFuture; + +/// A denoised frame whose readback has not finished. +/// +/// The readback starts on the first poll. Dropping a `Pending` that was polled but has not landed +/// blocks until it lands. +pub struct Pending { + future: ReadFuture, + /// Set while `future` has been polled but has not produced its result. + polled: bool, + /// The byte length of each plane, which trims the word padding off each buffer. + plane_lengths: Vec, +} + +/// The outcome of polling a [Pending] without blocking. +pub enum TryWait { + /// The readback landed. This is the frame, one buffer per plane. + Ready(Vec>), + /// The readback has not landed. Dropping it blocks until it does. + NotReady(Pending), +} + +impl Pending { + pub(crate) fn new(future: ReadFuture, plane_lengths: Vec) -> Self { + Self { + future, + polled: false, + plane_lengths, + } + } + + /// Blocks until the readback finishes and returns one buffer per plane. + pub fn wait(mut self) -> Result>, anyhow::Error> { + let future = self.future.as_mut(); + let result = cubecl::future::block_on(future); + self.polled = false; + + let buffers = result?; + let planes = trim_planes(buffers, &self.plane_lengths); + + Ok(planes) + } + + /// Polls the readback once. + /// + /// Only the wgpu backends avoid blocking here. On CUDA and ROCm the first poll waits for the + /// whole readback. + pub fn try_wait(mut self) -> Result { + let waker = Waker::noop(); + let mut context = Context::from_waker(waker); + + self.polled = true; + let poll = self.future.as_mut().poll(&mut context); + if poll.is_ready() { + self.polled = false; + } + + match poll { + Poll::Ready(Ok(buffers)) => { + let planes = trim_planes(buffers, &self.plane_lengths); + Ok(TryWait::Ready(planes)) + }, + Poll::Ready(Err(error)) => Err(error.into()), + Poll::Pending => Ok(TryWait::NotReady(self)), + } + } +} + +impl Drop for Pending { + /// Settles a readback that was polled but never landed. + /// + /// On the wgpu backends the first poll maps a staging buffer, and only dropping the finished + /// readback's bytes unmaps it. A still-mapped buffer back in the device's pool makes the next + /// submit on that device fail. Blocking here finishes the readback and drops its bytes. + fn drop(&mut self) { + if !self.polled || std::thread::panicking() { + return; + } + + let future = self.future.as_mut(); + let _ = cubecl::future::block_on(future); + } +} + +fn trim_planes(buffers: Vec, plane_lengths: &[usize]) -> Vec> { + buffers + .iter() + .zip(plane_lengths) + .map(|(buffer, &length)| buffer[..length].to_vec()) + .collect() +} diff --git a/av-denoise/src/host/tests/denoiser.rs b/av-denoise/src/host/tests/denoiser.rs new file mode 100644 index 0000000..8967bd4 --- /dev/null +++ b/av-denoise/src/host/tests/denoiser.rs @@ -0,0 +1,724 @@ +use av_denoise_core::{ + ChannelMode, + DenoisingMode, + DevicePlane, + EdgePadding, + Engine, + Nl4dOptions, + NlmTuning, + NlmeansOptions, + WindowSpan, +}; + +use crate::backend::Device; +use crate::backend::accelerate::Accelerator; +use crate::host::{Algorithm, DenoiserError, DenoiserOptions, HostDenoiser, OUTPUT_SETS}; + +fn luma_options(mode: DenoisingMode) -> DenoiserOptions { + DenoiserOptions::builder() + .channel_mode(ChannelMode::Luma) + .mode(mode) + .build() +} + +fn create(width: u32, height: u32, options: DenoiserOptions) -> Result { + HostDenoiser::create(&[Accelerator::Vulkan], &Device::Default, width, height, options) +} + +fn luma_denoiser(mode: DenoisingMode) -> HostDenoiser { + let options = luma_options(mode); + create(16, 16, options).expect("denoiser construction failed") +} + +fn frame(width: u32, height: u32) -> Vec { + vec![128u8; (width * height) as usize] +} + +fn frame_filled(value: u8) -> Vec { + vec![value; 16 * 16] +} + +fn luma_plane(mut planes: Vec>) -> Vec { + assert_eq!(planes.len(), 1, "a luma frame has one plane"); + planes.remove(0) +} + +/// Flushes `denoiser`, collecting each frame's luma plane into `outputs`. +fn flush_luma(denoiser: &mut HostDenoiser, outputs: &mut Vec>) -> Result<(), DenoiserError> { + denoiser.flush(|planes| { + let plane = luma_plane(planes); + outputs.push(plane); + }) +} + +/// Pushes `count` frames of `value`, receiving whenever the queue is full. +fn push_n_with_drain(denoiser: &mut HostDenoiser, count: usize, value: u8, outputs: &mut Vec>) { + let plane = frame_filled(value); + + for _ in 0..count { + loop { + match denoiser.push(&[&plane]) { + Ok(()) => break, + Err(DenoiserError::QueueFull) => { + let received = denoiser.recv().expect("recv ok"); + let planes = received.expect("queue full but recv yielded none"); + let plane = luma_plane(planes); + outputs.push(plane); + }, + Err(error) => panic!("unexpected push error: {error:?}"), + } + } + } +} + +#[test] +fn spatial_denoise_roundtrip() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + assert_eq!(denoiser.selected_accelerator(), Accelerator::Vulkan); + + let plane = frame(16, 16); + denoiser.push(&[&plane]).expect("push failed"); + + let received = denoiser.recv().expect("recv failed"); + let planes = received.expect("no frame"); + let output_plane = luma_plane(planes); + + assert_eq!(output_plane.len(), 16 * 16); +} + +#[test] +fn a_plane_that_is_not_a_whole_number_of_words_round_trips() { + let options = luma_options(DenoisingMode::Spacial); + let mut denoiser = create(13, 9, options).expect("denoiser construction failed"); + + let plane = frame(13, 9); + denoiser.push(&[&plane]).expect("push failed"); + + let received = denoiser.recv().expect("recv failed"); + let planes = received.expect("no frame"); + let output_plane = luma_plane(planes); + + assert_eq!(output_plane.len(), 13 * 9); +} + +#[test] +fn a_chroma_frame_keeps_its_u_and_v_order() { + let options = DenoiserOptions::builder() + .channel_mode(ChannelMode::Chroma) + .mode(DenoisingMode::Spacial) + .build(); + let mut denoiser = create(16, 16, options).expect("denoiser construction failed"); + + let u_plane = frame_filled(60); + let v_plane = frame_filled(190); + denoiser.push(&[&u_plane, &v_plane]).expect("push failed"); + + let received = denoiser.recv().expect("recv failed"); + let planes = received.expect("no frame"); + assert_eq!(planes.len(), 2); + + for &sample in &planes[0] { + assert!( + sample.abs_diff(60) <= 2, + "U plane sample {sample}, expected about 60" + ); + } + + for &sample in &planes[1] { + assert!( + sample.abs_diff(190) <= 2, + "V plane sample {sample}, expected about 190" + ); + } +} + +#[test] +fn a_wrong_plane_count_is_rejected_before_uploading() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + + let plane = frame(16, 16); + let result = denoiser.push(&[&plane, &plane]); + + let is_plane_mismatch = matches!( + result, + Err(DenoiserError::Engine(av_denoise_core::Error::PlaneMismatch(_))) + ); + assert!(is_plane_mismatch, "got {result:?}"); +} + +#[test] +fn a_plane_of_the_wrong_length_is_rejected() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + + let plane = vec![128u8; 16 * 15]; + let result = denoiser.push(&[&plane]); + + let is_plane_mismatch = matches!( + result, + Err(DenoiserError::Engine(av_denoise_core::Error::PlaneMismatch(_))) + ); + assert!(is_plane_mismatch, "got {result:?}"); +} + +/// An engine that returns scripted outcomes and writes nothing. +/// +/// `push_outcome` receives how many pushes came before it. +struct ScriptedEngine { + push_outcome: fn(usize) -> Result, + emit_outcome: fn() -> Result<(), av_denoise_core::Error>, + finish_outcome: fn() -> Result, + pushes: usize, +} + +impl ScriptedEngine { + fn pushing(push_outcome: fn(usize) -> Result) -> Self { + Self { + push_outcome, + emit_outcome: || Ok(()), + finish_outcome: || Ok(0), + pushes: 0, + } + } +} + +impl Engine for ScriptedEngine { + fn push(&mut self, _planes: &[DevicePlane<'_>]) -> Result { + let earlier_pushes = self.pushes; + self.pushes += 1; + + (self.push_outcome)(earlier_pushes) + } + + fn push_context(&mut self, _planes: &[DevicePlane<'_>]) -> Result<(), av_denoise_core::Error> { + Ok(()) + } + + fn emit_into(&mut self, _planes: &[DevicePlane<'_>]) -> Result<(), av_denoise_core::Error> { + (self.emit_outcome)() + } + + fn finish(&mut self) -> Result { + (self.finish_outcome)() + } + + fn reset(&mut self) {} + + fn window_span(&self) -> WindowSpan { + WindowSpan { + behind: 0, + ahead: 0, + edges: EdgePadding::Repeat, + } + } + + fn max_held_frames(&self) -> usize { + 0 + } +} + +#[test] +fn a_push_readying_more_frames_than_output_sets_fails_instead_of_overwriting() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + let engine = ScriptedEngine::pushing(|_| Ok(OUTPUT_SETS + 1)); + denoiser.engine = Box::new(engine); + + let plane = frame(16, 16); + let result = denoiser.push(&[&plane]); + + assert!(matches!(result, Err(DenoiserError::Other(_))), "got {result:?}"); + assert!( + denoiser.pending.is_empty(), + "no output set should have been emitted" + ); + + let pushed = denoiser.push(&[&plane]); + assert!(matches!(pushed, Err(DenoiserError::Poisoned))); +} + +#[test] +fn a_push_filling_every_output_set_with_one_in_flight_fails_instead_of_overwriting() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + let engine = ScriptedEngine::pushing(|earlier_pushes| match earlier_pushes { + 0 => Ok(1), + _ => Ok(OUTPUT_SETS), + }); + denoiser.engine = Box::new(engine); + + let plane = frame(16, 16); + denoiser.push(&[&plane]).expect("first push should land"); + assert_eq!(denoiser.pending.len(), 1); + + let result = denoiser.push(&[&plane]); + + assert!(matches!(result, Err(DenoiserError::Other(_))), "got {result:?}"); + assert_eq!( + denoiser.pending.len(), + 1, + "no output set should have been emitted" + ); +} + +#[test] +fn a_failed_emit_during_push_poisons_the_denoiser() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + let engine = ScriptedEngine { + emit_outcome: || Err(av_denoise_core::Error::NothingToEmit), + ..ScriptedEngine::pushing(|_| Ok(1)) + }; + denoiser.engine = Box::new(engine); + + let plane = frame(16, 16); + let result = denoiser.push(&[&plane]); + + let is_nothing_to_emit = matches!( + result, + Err(DenoiserError::Engine(av_denoise_core::Error::NothingToEmit)) + ); + assert!(is_nothing_to_emit, "got {result:?}"); + + let pushed = denoiser.push(&[&plane]); + assert!(matches!(pushed, Err(DenoiserError::Poisoned))); +} + +#[test] +fn a_failed_finish_during_flush_poisons_the_denoiser() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + let engine = ScriptedEngine { + finish_outcome: || Err(av_denoise_core::Error::OutputsPending), + ..ScriptedEngine::pushing(|_| Ok(0)) + }; + denoiser.engine = Box::new(engine); + + let flushed = denoiser.flush(|_| {}); + + let is_outputs_pending = matches!( + flushed, + Err(DenoiserError::Engine(av_denoise_core::Error::OutputsPending)) + ); + assert!(is_outputs_pending, "got {flushed:?}"); + + let plane = frame(16, 16); + let pushed = denoiser.push(&[&plane]); + assert!(matches!(pushed, Err(DenoiserError::Poisoned))); +} + +#[test] +fn nl4d_algorithm_round_trips_through_the_host() { + let algorithm = Algorithm::Nl4d(Nl4dOptions::default()); + let options = DenoiserOptions::builder() + .channel_mode(ChannelMode::Luma) + .mode(DenoisingMode::Temporal { radius: 2 }) + .algorithm(algorithm) + .build(); + let mut denoiser = create(16, 16, options).expect("nl4d denoiser construction failed"); + assert_eq!(denoiser.selected_accelerator(), Accelerator::Vulkan); + + let plane = frame(16, 16); + denoiser.push(&[&plane]).expect("push failed"); + + let received = denoiser.recv().expect("recv failed"); + assert!(received.is_none()); + + let mut outputs = Vec::new(); + flush_luma(&mut denoiser, &mut outputs).expect("flush failed"); + + assert_eq!( + outputs.len(), + 1, + "expected exactly one output for one pushed frame" + ); + assert_eq!(outputs[0].len(), 16 * 16); +} + +/// nl4d groups patches across neighbouring frames, so a spatial mode leaves it nothing to do. +#[test] +fn nl4d_rejects_a_spatial_denoising_mode() { + let algorithm = Algorithm::Nl4d(Nl4dOptions::default()); + let options = DenoiserOptions::builder() + .channel_mode(ChannelMode::Luma) + .mode(DenoisingMode::Spacial) + .algorithm(algorithm) + .build(); + let result = create(16, 16, options); + + match result { + Err(DenoiserError::Other(error)) => assert!( + error.to_string().contains("temporal window"), + "unexpected error message: {error}" + ), + Err(other) => panic!("expected DenoiserError::Other, got {other:?}"), + Ok(_) => panic!("expected a rejection, got Ok"), + } +} + +#[test] +fn window_span_is_symmetric_for_nlmeans() { + let algorithm = Algorithm::Nlmeans(NlmeansOptions::default()); + let options = DenoiserOptions::builder() + .channel_mode(ChannelMode::Luma) + .mode(DenoisingMode::Temporal { radius: 3 }) + .algorithm(algorithm) + .build(); + let denoiser = create(16, 16, options).expect("denoiser construction failed"); + + let span = denoiser.window_span(); + + assert_eq!(span.behind, 3, "behind should equal the temporal radius"); + assert_eq!(span.ahead, 3, "ahead should equal the temporal radius"); + assert_eq!(span.edges, EdgePadding::Repeat); +} + +#[test] +fn window_span_is_doubled_on_both_sides_for_nl4d() { + let algorithm = Algorithm::Nl4d(Nl4dOptions::default()); + let options = DenoiserOptions::builder() + .channel_mode(ChannelMode::Luma) + .mode(DenoisingMode::Temporal { radius: 3 }) + .algorithm(algorithm) + .build(); + let denoiser = create(16, 16, options).expect("nl4d denoiser construction failed"); + + let span = denoiser.window_span(); + + assert_eq!(span.behind, 6, "behind should equal 2 * the temporal radius"); + assert_eq!(span.ahead, 6, "ahead should equal 2 * the temporal radius"); + assert_eq!(span.edges, EdgePadding::Shifted); +} + +#[test] +fn invalid_params_surface_as_error() { + let tuning = NlmTuning { + strength: Some(0.0), + ..NlmTuning::default() + }; + let nlmeans_options = NlmeansOptions { + tuning, + ..NlmeansOptions::default() + }; + let algorithm = Algorithm::Nlmeans(nlmeans_options); + let options = DenoiserOptions::builder().algorithm(algorithm).build(); + let result = create(16, 16, options); + + match result { + Err(DenoiserError::Engine(av_denoise_core::Error::InvalidOptions(_))) => {}, + Err(other) => panic!("expected an invalid options error, got {other:?}"), + Ok(_) => panic!("expected validation error, got Ok"), + } +} + +#[test] +fn tiny_frame_dimensions_surface_as_error() { + let options = luma_options(DenoisingMode::Spacial); + let result = create(2, 2, options); + + match result { + Err(DenoiserError::Engine(error)) => assert!( + error.to_string().contains("supported minimum"), + "unexpected error message: {error}" + ), + Err(other) => panic!("expected an engine error, got {other:?}"), + Ok(_) => panic!("expected dimension validation error, got Ok"), + } +} + +#[test] +fn push_after_pending_returns_queue_full() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + let plane = frame(16, 16); + + denoiser.push(&[&plane]).unwrap(); + denoiser.push(&[&plane]).unwrap(); + let error = denoiser.push(&[&plane]).expect_err("expected QueueFull"); + assert!(matches!(error, DenoiserError::QueueFull)); + + let received = denoiser.recv().unwrap(); + let planes = received.unwrap(); + let output_plane = luma_plane(planes); + + assert_eq!(output_plane.len(), 16 * 16); + + denoiser.push(&[&plane]).expect("push after drain failed"); +} + +#[test] +fn queue_full_does_not_poison() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + let plane = frame(16, 16); + + denoiser.push(&[&plane]).unwrap(); + denoiser.push(&[&plane]).unwrap(); + let error = denoiser.push(&[&plane]).expect_err("expected QueueFull"); + assert!(matches!(error, DenoiserError::QueueFull)); + assert!(!denoiser.poisoned, "QueueFull must not poison the denoiser"); + + let received = denoiser.recv().unwrap(); + received.expect("recv failed after QueueFull"); + + denoiser + .push(&[&plane]) + .expect("push after QueueFull drain should succeed, not poison"); +} + +#[test] +fn poisoned_denoiser_refuses_every_entry_point() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + denoiser.poisoned = true; + + let plane = frame(16, 16); + + let pushed = denoiser.push(&[&plane]); + assert!(matches!(pushed, Err(DenoiserError::Poisoned))); + + let primed = denoiser.push_priming(&[&plane]); + assert!(matches!(primed, Err(DenoiserError::Poisoned))); + + let received = denoiser.recv(); + assert!(matches!(received, Err(DenoiserError::Poisoned))); + + let polled = denoiser.try_recv(); + assert!(matches!(polled, Err(DenoiserError::Poisoned))); + + let flushed = denoiser.flush(|_| {}); + assert!(matches!(flushed, Err(DenoiserError::Poisoned))); + + let drained = denoiser.drain_grain_chunks(); + assert!(matches!(drained, Err(DenoiserError::Poisoned))); +} + +#[test] +fn a_gpu_failure_during_push_poisons_the_denoiser() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + let engine = ScriptedEngine::pushing(|_| { + let failure = anyhow::anyhow!("synthetic dispatch failure"); + Err(av_denoise_core::Error::Gpu(failure)) + }); + denoiser.engine = Box::new(engine); + + let plane = frame(16, 16); + let result = denoiser.push(&[&plane]); + assert!(result.is_err()); + + let pushed = denoiser.push(&[&plane]); + assert!(matches!(pushed, Err(DenoiserError::Poisoned))); +} + +#[test] +fn a_rejected_plane_does_not_poison_the_denoiser() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + + let short_plane = vec![128u8; 4]; + let result = denoiser.push(&[&short_plane]); + assert!(result.is_err()); + + let plane = frame(16, 16); + denoiser + .push(&[&plane]) + .expect("a push after a rejected plane should land"); +} + +#[test] +fn priming_after_a_push_is_rejected_without_poisoning() { + let mut denoiser = luma_denoiser(DenoisingMode::Temporal { radius: 1 }); + let plane = frame(16, 16); + + denoiser.push(&[&plane]).expect("first push should land"); + + let primed = denoiser.push_priming(&[&plane]); + let is_context_after_push = matches!( + primed, + Err(DenoiserError::Engine(av_denoise_core::Error::ContextAfterPush)) + ); + assert!(is_context_after_push, "got {primed:?}"); + + denoiser + .push(&[&plane]) + .expect("a push after a rejected priming frame should land"); +} + +#[test] +fn reset_stream_clears_poison() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + denoiser.poisoned = true; + + denoiser.reset_stream(); + assert!(!denoiser.poisoned, "reset_stream must clear the poison flag"); + + let plane = frame(16, 16); + denoiser + .push(&[&plane]) + .expect("push after reset_stream should succeed"); +} + +#[test] +fn flush_leaves_denoiser_reusable_spatial() { + let mut denoiser = luma_denoiser(DenoisingMode::Spacial); + + let mut batch_a = Vec::new(); + push_n_with_drain(&mut denoiser, 5, 64, &mut batch_a); + flush_luma(&mut denoiser, &mut batch_a).expect("first flush failed"); + assert_eq!(batch_a.len(), 5); + + let received = denoiser.recv().unwrap(); + assert!(received.is_none()); + + let mut batch_b = Vec::new(); + push_n_with_drain(&mut denoiser, 5, 191, &mut batch_b); + flush_luma(&mut denoiser, &mut batch_b).expect("second flush failed"); + assert_eq!(batch_b.len(), 5); + + for &sample in batch_b.iter().flatten() { + assert!( + sample.abs_diff(191) < 25, + "batch_b carried state from batch_a: {sample}" + ); + } + + for &sample in batch_a.iter().flatten() { + assert!( + sample.abs_diff(64) < 25, + "batch_a value unexpectedly drifted: {sample}" + ); + } +} + +#[test] +fn flush_leaves_denoiser_reusable_temporal() { + let mut denoiser = luma_denoiser(DenoisingMode::Temporal { radius: 1 }); + + let mut batch_a = Vec::new(); + push_n_with_drain(&mut denoiser, 5, 64, &mut batch_a); + flush_luma(&mut denoiser, &mut batch_a).expect("first flush failed"); + assert_eq!(batch_a.len(), 5, "expected 5 frames from first batch"); + + // With radius 1 the window needs more than one push before anything is ready. + let received = denoiser.recv().unwrap(); + assert!(received.is_none()); + + let plane = frame_filled(191); + denoiser.push(&[&plane]).unwrap(); + + let received = denoiser.recv().unwrap(); + assert!( + received.is_none(), + "first push of new temporal stream should not produce output yet" + ); + + let mut batch_b = Vec::new(); + push_n_with_drain(&mut denoiser, 4, 191, &mut batch_b); + flush_luma(&mut denoiser, &mut batch_b).expect("second flush failed"); + assert_eq!(batch_b.len(), 5, "expected 5 frames from second batch"); + + for &sample in batch_b.iter().flatten() { + assert!( + sample.abs_diff(191) < 25, + "batch_b carried state from batch_a: {sample}" + ); + } +} + +#[test] +fn flush_emits_exactly_n_outputs_for_small_n() { + for count in 1..=5usize { + let mut denoiser = luma_denoiser(DenoisingMode::Temporal { radius: 2 }); + + let mut outputs = Vec::new(); + push_n_with_drain(&mut denoiser, count, 128, &mut outputs); + flush_luma(&mut denoiser, &mut outputs).expect("flush failed"); + + assert_eq!( + outputs.len(), + count, + "expected {count} outputs for {count} pushes, got {}", + outputs.len() + ); + } +} + +#[test] +fn dropping_a_polled_pending_frame_does_not_poison_the_device() { + let create_denoiser = || { + let options = luma_options(DenoisingMode::Spacial); + create(64, 64, options).unwrap() + }; + let plane = frame(64, 64); + + let mut denoiser = create_denoiser(); + denoiser.push(&[&plane]).unwrap(); + + // Whether the readback lands on this poll depends on the GPU, and both outcomes must survive the drop. + let _ = denoiser.try_recv().unwrap(); + drop(denoiser); + + // The staging pool is per device, so a fresh denoiser on the same device is handed the same buffers. + let mut denoiser = create_denoiser(); + for _ in 0..4 { + denoiser.push(&[&plane]).unwrap(); + + let received = denoiser + .recv() + .expect("readback after a dropped polled frame should not fail"); + received.expect("spatial mode emits one frame per push"); + } + + denoiser.flush(|_| {}).unwrap(); +} + +#[test] +fn try_recv_observes_a_landed_readback_within_a_bounded_poll() { + // A deadline covers both a slow GPU and a fast CPU, where a poll count would not. + const DEADLINE: std::time::Duration = std::time::Duration::from_secs(30); + + let radius = 2u32; + let create_denoiser = || { + let options = luma_options(DenoisingMode::Temporal { radius }); + create(64, 64, options).unwrap() + }; + + // `radius + 1` pushes fill the window and leave exactly one readback in flight. + let window: Vec> = (0..=radius as usize) + .map(|index| { + (0..64 * 64) + .map(|pixel| ((pixel * 7 + index * 13) % 256) as u8) + .collect() + }) + .collect(); + + let mut polled = create_denoiser(); + for plane in &window { + polled.push(&[plane]).unwrap(); + } + + let start = std::time::Instant::now(); + let mut landed = None; + let mut polls = 0; + while start.elapsed() < DEADLINE { + polls += 1; + + if let Some(planes) = polled.try_recv().unwrap() { + landed = Some(planes); + break; + } + } + + let Some(landed) = landed else { + panic!("readback never landed within {DEADLINE:?} ({polls} polls)"); + }; + + let mut blocking = create_denoiser(); + for plane in &window { + blocking.push(&[plane]).unwrap(); + } + + let received = blocking.recv().unwrap(); + let expected = received.expect("blocking denoiser should have a frame ready"); + + assert_eq!(landed, expected); +} + +#[test] +fn try_recv_returns_none_when_nothing_is_in_flight() { + let mut denoiser = luma_denoiser(DenoisingMode::Temporal { radius: 2 }); + let polled = denoiser.try_recv().unwrap(); + + assert_eq!(polled, None); +} diff --git a/av-denoise/src/host/tests/mod.rs b/av-denoise/src/host/tests/mod.rs new file mode 100644 index 0000000..e0f0f61 --- /dev/null +++ b/av-denoise/src/host/tests/mod.rs @@ -0,0 +1,5 @@ +#[cfg(feature = "vulkan")] +mod denoiser; +mod options; +#[cfg(feature = "vulkan")] +mod pending; diff --git a/av-denoise/src/host/tests/options.rs b/av-denoise/src/host/tests/options.rs new file mode 100644 index 0000000..cd1fe25 --- /dev/null +++ b/av-denoise/src/host/tests/options.rs @@ -0,0 +1,131 @@ +use av_denoise_core::{ + ChannelMode, + DenoisingMode, + Nl4dOptions, + NlmeansAlgorithm, + NlmeansHqOptions, + NlmeansOptions, + SampleFormat, +}; + +use crate::backend::EngineSpec; +use crate::host::{Algorithm, DenoiserOptions, Depth}; + +fn spec_for(options: &DenoiserOptions) -> EngineSpec { + options.algorithm.engine_spec(options, 32, 16) +} + +fn expect_nlmeans(spec: EngineSpec) -> NlmeansAlgorithm { + match spec { + EngineSpec::Nlmeans { algorithm, .. } => algorithm, + other => panic!("expected an nlmeans spec, got {other:?}"), + } +} + +fn expect_nl4d(spec: EngineSpec) -> Nl4dOptions { + match spec { + EngineSpec::Nl4d { options, .. } => options, + other => panic!("expected an nl4d spec, got {other:?}"), + } +} + +#[test] +fn the_default_algorithm_is_the_fast_nlmeans_path() { + let options = DenoiserOptions::builder().build(); + let expected = Algorithm::Nlmeans(NlmeansOptions::default()); + + assert_eq!(options.algorithm, expected); +} + +#[test] +fn the_geometry_carries_the_size_planes_and_depth() { + let options = DenoiserOptions::builder() + .channel_mode(ChannelMode::Chroma) + .depth(Depth::Ten) + .build(); + + let EngineSpec::Nlmeans { geometry, .. } = spec_for(&options) else { + panic!("expected an nlmeans spec"); + }; + + assert_eq!(geometry.width, 32); + assert_eq!(geometry.height, 16); + assert_eq!(geometry.channels, ChannelMode::Chroma); + assert_eq!(geometry.input, SampleFormat::U16 { depth: 10 }); + assert_eq!(geometry.output, SampleFormat::U16 { depth: 10 }); +} + +#[test] +fn the_denoising_mode_wins_over_the_fast_options_mode() { + let fast_options = NlmeansOptions { + mode: DenoisingMode::Temporal { radius: 5 }, + ..NlmeansOptions::default() + }; + let algorithm = Algorithm::Nlmeans(fast_options); + let options = DenoiserOptions::builder() + .mode(DenoisingMode::Spacial) + .algorithm(algorithm) + .build(); + + let spec = spec_for(&options); + let NlmeansAlgorithm::Fast(fast) = expect_nlmeans(spec) else { + panic!("expected the fast variant"); + }; + + assert_eq!(fast.mode, DenoisingMode::Spacial); +} + +#[test] +fn the_denoising_mode_wins_over_the_hq_options_mode() { + let algorithm = Algorithm::NlmeansHq(NlmeansHqOptions::default()); + let options = DenoiserOptions::builder() + .mode(DenoisingMode::Temporal { radius: 3 }) + .algorithm(algorithm) + .build(); + + let spec = spec_for(&options); + let NlmeansAlgorithm::Hq(hq) = expect_nlmeans(spec) else { + panic!("expected the hq variant"); + }; + + assert_eq!(hq.nlm.mode, DenoisingMode::Temporal { radius: 3 }); +} + +#[test] +fn the_denoising_mode_sets_the_nl4d_temporal_radius() { + for radius in [1u32, 4, 8] { + let nl4d_options = Nl4dOptions { + temporal_radius: 2, + ..Nl4dOptions::default() + }; + let algorithm = Algorithm::Nl4d(nl4d_options); + let options = DenoiserOptions::builder() + .mode(DenoisingMode::Temporal { radius }) + .algorithm(algorithm) + .build(); + + let spec = spec_for(&options); + let nl4d = expect_nl4d(spec); + + assert_eq!(nl4d.temporal_radius, radius); + } +} + +#[test] +fn other_algorithm_options_pass_through_untouched() { + let nl4d = Nl4dOptions { + sigma: Some(0.02), + refine: 3, + ..Nl4dOptions::default() + }; + let options = DenoiserOptions::builder() + .mode(DenoisingMode::Temporal { radius: 2 }) + .algorithm(Algorithm::Nl4d(nl4d)) + .build(); + + let spec = spec_for(&options); + let resolved = expect_nl4d(spec); + + assert_eq!(resolved.sigma, Some(0.02)); + assert_eq!(resolved.refine, 3); +} diff --git a/av-denoise/src/host/tests/pending.rs b/av-denoise/src/host/tests/pending.rs new file mode 100644 index 0000000..820866e --- /dev/null +++ b/av-denoise/src/host/tests/pending.rs @@ -0,0 +1,77 @@ +use av_denoise_core::{ChannelMode, DenoisingMode, NlmTuning, NlmeansOptions}; + +use crate::backend::Device; +use crate::backend::accelerate::Accelerator; +use crate::host::{Algorithm, DenoiserOptions, HostDenoiser, Pending, TryWait}; + +/// A spatial luma denoiser with small radii, since wide kernels cost codegen stack. +fn spatial_luma(size: u32) -> HostDenoiser { + let tuning = NlmTuning { + search_radius: Some(3), + patch_radius: Some(2), + ..NlmTuning::default() + }; + let nlmeans_options = NlmeansOptions { + tuning, + ..NlmeansOptions::default() + }; + let algorithm = Algorithm::Nlmeans(nlmeans_options); + let options = DenoiserOptions::builder() + .channel_mode(ChannelMode::Luma) + .mode(DenoisingMode::Spacial) + .algorithm(algorithm) + .build(); + + HostDenoiser::create(&[Accelerator::Vulkan], &Device::Default, size, size, options) + .expect("denoiser construction failed") +} + +/// Pushes one uniform frame and takes its readback out of the denoiser. +fn submit(denoiser: &mut HostDenoiser, size: u32) -> Pending { + let plane = vec![128u8; (size * size) as usize]; + denoiser.push(&[&plane]).expect("push failed"); + + let pending = denoiser.pending.pop_front(); + pending.expect("spatial mode submits one readback per push") +} + +#[test] +fn pending_survives_denoiser_drop() { + let pending = { + let mut denoiser = spatial_luma(16); + submit(&mut denoiser, 16) + }; + + let planes = pending.wait().expect("wait failed"); + + assert_eq!(planes.len(), 1); + assert_eq!(planes[0].len(), 16 * 16); + + for (index, &sample) in planes[0].iter().enumerate() { + assert_eq!(sample, 128, "pixel {index}"); + } +} + +/// Large enough that the GPU cannot have finished by the time the first poll runs. +const LARGE_SIZE: u32 = 2048; + +#[test] +fn dropping_a_polled_pending_settles_its_readback() { + let mut denoiser = spatial_luma(LARGE_SIZE); + + let pending = submit(&mut denoiser, LARGE_SIZE); + let not_ready = match pending.try_wait().expect("poll failed") { + TryWait::NotReady(pending) => pending, + TryWait::Ready(_) => { + panic!("the first poll landed, so the drop path cannot be exercised at this size") + }, + }; + drop(not_ready); + + let pending = submit(&mut denoiser, LARGE_SIZE); + let planes = pending + .wait() + .expect("a readback after a dropped polled Pending must still work"); + + assert_eq!(planes[0].len(), (LARGE_SIZE * LARGE_SIZE) as usize); +} diff --git a/av-denoise/src/lib.rs b/av-denoise/src/lib.rs index 97c81d3..a3cd6b5 100644 --- a/av-denoise/src/lib.rs +++ b/av-denoise/src/lib.rs @@ -2,4 +2,80 @@ #![cfg_attr(docsrs, doc(auto_cfg))] #![doc = include_str!("../README.md")] -pub use av_denoise_core::*; +mod backend; +pub mod cache; +mod host; +mod planar; +pub mod stack; +pub mod warmup; + +pub use av_denoise_core::{ + ChannelMode, + DEFAULT_PILOT_STRENGTH_SCALE, + DenoisingMode, + DevicePlane, + EdgePadding, + Engine, + Error as EngineError, + Geometry, + GrainChunk, + HqParams, + KERNEL_HASH, + MotionCompensationMode, + MotionEstimation, + MotionSearch, + Nl4d, + Nl4dOptions, + NlmTuning, + Nlmeans, + NlmeansAlgorithm, + NlmeansHqOptions, + NlmeansOptions, + NlmeansVariant, + PrefilterMode, + Preset, + SampleFormat, + SceneGrain, + WindowSpan, + build_table, + nl4d_default_lambda_ht, + nl4d_spatial_radius_for, + nl4d_temporal_radius_for, + nlmeans_search_radius_for, + nlmeans_temporal_radius_for, + nlmeans_variant_for, + parse_prefilter, +}; + +pub use self::backend::{Device, accelerate, device, enumerate, sniff}; +pub use self::cache::{ + COMPILATION_CACHE_ENV, + CacheError, + compilation_cache_dir, + default_cache_dir, + install_compilation_cache, + install_compilation_cache_at, + install_compilation_cache_once, +}; +pub use self::host::{ + Algorithm, + DenoiserError, + DenoiserOptions, + Depth, + HostDenoiser, + MAX_PENDING, + UnsupportedDepthError, +}; +pub use self::planar::{ + ChannelIntent, + FrameLayout, + PlanarDenoiser, + PlaneOptions, + Planes, + ReseedWindow, + Subsampling, + fill_plane, + push_needs_retry, +}; +pub use self::stack::{CODEGEN_STACK_BYTES, codegen_stack_is_sufficient, raise_codegen_stack_limit}; +pub use self::warmup::{WarmUp, kernel_key}; diff --git a/av-denoise/src/planar/mod.rs b/av-denoise/src/planar/mod.rs new file mode 100644 index 0000000..2162f5b --- /dev/null +++ b/av-denoise/src/planar/mod.rs @@ -0,0 +1,651 @@ +mod reseed; + +#[cfg(test)] +mod tests; + +use std::collections::VecDeque; + +use av_denoise_core::{ + ChannelMode, + DenoisingMode, + GrainChunk, + Nl4dOptions, + NlmTuning, + NlmeansHqOptions, + NlmeansOptions, + WindowSpan, +}; + +pub use self::reseed::ReseedWindow; +use crate::backend::Device; +use crate::backend::accelerate::Accelerator; +use crate::host::{Algorithm, DenoiserError, DenoiserOptions, Depth, HostDenoiser}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Subsampling { + Yuv420, + Yuv422, + Yuv444, +} + +impl Subsampling { + /// The chroma plane size for a `width` by `height` frame. + /// + /// Halved axes round up, so an odd dimension keeps the extra sample, matching what y4m and ffmpeg do. + pub fn chroma_dims(self, width: u32, height: u32) -> (u32, u32) { + match self { + Subsampling::Yuv420 => (width.div_ceil(2), height.div_ceil(2)), + Subsampling::Yuv422 => (width.div_ceil(2), height), + Subsampling::Yuv444 => (width, height), + } + } +} + +#[derive(Debug, Clone, Copy)] +pub struct FrameLayout { + pub width: u32, + pub height: u32, + pub subsampling: Subsampling, + pub depth: Depth, +} + +impl FrameLayout { + pub fn luma_pixels(&self) -> usize { + (self.width as usize) * (self.height as usize) + } + + pub fn chroma_dims(&self) -> (u32, u32) { + self.subsampling.chroma_dims(self.width, self.height) + } + + pub fn chroma_pixels(&self) -> usize { + let (chroma_width, chroma_height) = self.chroma_dims(); + (chroma_width as usize) * (chroma_height as usize) + } + + /// Wire size of the luma plane. + pub fn luma_bytes(&self) -> usize { + self.luma_pixels() * self.depth.bytes_per_sample() + } + + /// Wire size of one chroma plane. + pub fn chroma_bytes(&self) -> usize { + self.chroma_pixels() * self.depth.bytes_per_sample() + } + + /// A full black luma plane, used when no luma source is available. + pub fn black_luma_plane(&self) -> Vec { + let samples = self.luma_pixels(); + fill_plane(samples, 0, self.depth) + } + + /// A full neutral chroma plane, used when a source has no chroma. + pub fn neutral_chroma_plane(&self) -> Vec { + let samples = self.chroma_pixels(); + let neutral = self.depth.neutral_chroma(); + fill_plane(samples, neutral, self.depth) + } +} + +/// Builds a plane of `samples` copies of `value` in wire-byte form. +pub fn fill_plane(samples: usize, value: u16, depth: Depth) -> Vec { + match depth.bytes_per_sample() { + 1 => vec![value as u8; samples], + _ => { + let word = value.to_le_bytes(); + let mut plane = Vec::with_capacity(samples * 2); + for _ in 0..samples { + plane.extend_from_slice(&word); + } + + plane + }, + } +} + +/// A planar YUV frame holding little-endian wire bytes. +/// +/// Plane lengths come from [FrameLayout], so `y.len()` is `layout.luma_bytes()` and both `u.len()` +/// and `v.len()` are `layout.chroma_bytes()`. +#[derive(Debug, Clone)] +pub struct Planes { + pub y: Vec, + pub u: Vec, + pub v: Vec, +} + +/// Which planes a caller wants cleaned. +/// +/// This is separate from [ChannelMode] because a [PlanarDenoiser] may run two `HostDenoiser`s in +/// lockstep, one for luma and one for chroma, or a single fused three-channel one. Which applies +/// depends on the caller's channel selection and the source's chroma subsampling. +#[derive(Debug, Copy, Clone, PartialEq, Eq)] +pub enum ChannelIntent { + /// Denoise luma only. Chroma passes through. + Luma, + /// Denoise chroma only. Luma passes through. + Chroma, + /// Denoise luma and chroma as two independent denoisers. + /// + /// Chroma runs at the source's native subsampled resolution. + LumaChroma, + /// A single `HostDenoiser` running the fused three-channel kernel, which needs a YUV444 source. + YuvFused, +} + +impl ChannelIntent { + /// Rejects the intent if the source's subsampling cannot support it. + pub fn validate_for_source(self, layout: FrameLayout) -> Result<(), anyhow::Error> { + match self { + ChannelIntent::YuvFused if layout.subsampling != Subsampling::Yuv444 => { + anyhow::bail!( + "--channel-mode yuv requires a YUV444 source, got {:?}. Convert the input first, for example with `ffmpeg -pix_fmt yuv444p`", + layout.subsampling + ); + }, + _ => Ok(()), + } + } +} + +/// The per-plane option set a caller resolves once and passes into [PlanarDenoiser::create]. +#[derive(Debug, Clone)] +pub struct PlaneOptions { + pub accelerators: Vec, + pub device: Device, + pub intent: ChannelIntent, + /// Whether to clean each frame on its own or across a temporal window. + /// + /// This wins over `NlmeansOptions.mode` and `Nl4dOptions.temporal_radius` inside `algorithm`. + pub mode: DenoisingMode, + /// Which denoising algorithm to run, along with the settings only that algorithm reads. + pub algorithm: Algorithm, + /// Strength override for the luma denoiser, which wins over the algorithm's `tuning.strength`. + /// + /// Only the two NLM algorithms read it. + pub luma_strength: Option, + /// Strength override for the chroma denoiser, which wins over the algorithm's `tuning.strength`. + /// + /// Only the two NLM algorithms read it. + pub chroma_strength: Option, + /// Luma override for nl4d's `lambda_ht`, the temporal grouping stage's hard threshold. + /// + /// It wins over `algorithm`'s value, which itself falls back to a calibrated per-plane default + /// when neither is set. + pub luma_lambda_ht: Option, + /// Chroma override for nl4d's `lambda_ht`, the temporal grouping stage's hard threshold. + /// + /// It wins over `algorithm`'s value, which itself falls back to a calibrated per-plane default + /// when neither is set. + pub chroma_lambda_ht: Option, +} + +impl PlaneOptions { + /// Resolves `self.algorithm` for one plane, folding in that plane's overrides. + /// + /// The NLM algorithms take a strength override. Nl4d takes a `lambda_ht` override instead, because + /// it has no NLM weighting pass for a strength to affect. + /// + /// An unset nl4d `lambda_ht` stays `None`, so construction resolves it with + /// [nl4d_default_lambda_ht](crate::nl4d_default_lambda_ht) once the plane is known. That is what + /// gives luma and chroma different values when a caller sets nothing. + fn algorithm_for(&self, channels: ChannelMode) -> Algorithm { + let per_plane = |luma, chroma| match channels { + ChannelMode::Luma => luma, + ChannelMode::Chroma => chroma, + ChannelMode::Yuv => None, + }; + + match self.algorithm { + Algorithm::Nl4d(nl4d) => { + let plane_lambda_ht = per_plane(self.luma_lambda_ht, self.chroma_lambda_ht); + let options = Nl4dOptions { + lambda_ht: plane_lambda_ht.or(nl4d.lambda_ht), + grain_export: nl4d.grain_export && channels != ChannelMode::Chroma, + ..nl4d + }; + + Algorithm::Nl4d(options) + }, + Algorithm::Nlmeans(nlm) => { + let strength = per_plane(self.luma_strength, self.chroma_strength); + let options = with_plane_strength(nlm, strength); + + Algorithm::Nlmeans(options) + }, + Algorithm::NlmeansHq(hq) => { + let strength = per_plane(self.luma_strength, self.chroma_strength); + let nlm = with_plane_strength(hq.nlm, strength); + let options = NlmeansHqOptions { nlm, ..hq }; + + Algorithm::NlmeansHq(options) + }, + } + } + + /// The options for one plane's denoiser at the source's wire `depth`. + /// + /// Every denoiser quantises to `depth` on the GPU. + fn denoiser_options(&self, channels: ChannelMode, depth: Depth) -> DenoiserOptions { + let algorithm = self.algorithm_for(channels); + + DenoiserOptions::builder() + .channel_mode(channels) + .mode(self.mode) + .algorithm(algorithm) + .depth(depth) + .build() + } +} + +/// Replaces `nlm`'s strength with the per-plane override, when there is one. +fn with_plane_strength(nlm: NlmeansOptions, strength: Option) -> NlmeansOptions { + match strength { + None => nlm, + Some(strength) => { + let tuning = NlmTuning { + strength: Some(strength), + ..nlm.tuning + }; + + NlmeansOptions { tuning, ..nlm } + }, + } +} + +/// Reads the result of a [PlanarDenoiser::push] for the push, drain and retry loop. +/// +/// `Ok(false)` means the push landed. `Ok(true)` means the queue was full, so the caller should drain +/// one output and push again. Any error other than `QueueFull` is passed on rather than discarded. +pub fn push_needs_retry(result: Result<(), DenoiserError>) -> Result { + match result { + Ok(()) => Ok(false), + Err(DenoiserError::QueueFull) => Ok(true), + Err(other) => Err(other.into()), + } +} + +/// Turns one denoised frame's per-plane buffers into `N` planes. +fn into_array(planes: Vec>) -> [Vec; N] { + let count = planes.len(); + let array = planes.try_into(); + array.unwrap_or_else(|_| panic!("expected {N} planes from a HostDenoiser, got {count}")) +} + +fn into_yuv(planes: Vec>) -> Planes { + let [y, u, v] = into_array(planes); + Planes { y, u, v } +} + +fn into_uv(planes: Vec>) -> (Vec, Vec) { + let [u_plane, v_plane] = into_array(planes); + (u_plane, v_plane) +} + +fn into_luma(planes: Vec>) -> Vec { + let [y_plane] = into_array(planes); + y_plane +} + +/// The push a [PlanarDenoiser] runs against each enabled half, either [HostDenoiser::push] or +/// [HostDenoiser::push_priming]. +type WirePush = fn(&mut HostDenoiser, &[&[u8]]) -> Result<(), DenoiserError>; + +/// Wraps the luma and chroma `HostDenoiser` instances needed for one subsampled YUV source. +/// +/// The caller pushes planar frames in and gets planar frames out. The luma and chroma split is +/// invisible from the outside. +pub struct PlanarDenoiser { + layout: FrameLayout, + luma: Option, + chroma: Option, + /// Set when the intent is `YuvFused`, in which case `luma` and `chroma` are both unset. + yuv: Option, + // Source planes queued for the disabled side, if there is one. One entry is popped per frame the + // enabled side emits, so temporal delays stay aligned. + luma_passthrough: VecDeque>, + chroma_passthrough: VecDeque<(Vec, Vec)>, + temporal_radius: u32, +} + +impl PlanarDenoiser { + pub fn create(options: &PlaneOptions, layout: FrameLayout) -> Result { + let (chroma_width, chroma_height) = layout.chroma_dims(); + + if chroma_width == 0 || chroma_height == 0 { + anyhow::bail!( + "frame dimensions {}x{} are too small for subsampling {:?}", + layout.width, + layout.height, + layout.subsampling + ); + } + + options.intent.validate_for_source(layout)?; + + let (denoise_luma, denoise_chroma, denoise_yuv) = match options.intent { + ChannelIntent::Luma => (true, false, false), + ChannelIntent::Chroma => (false, true, false), + ChannelIntent::LumaChroma => (true, true, false), + ChannelIntent::YuvFused => (false, false, true), + }; + + let luma = denoise_luma + .then(|| { + let luma_options = options.denoiser_options(ChannelMode::Luma, layout.depth); + HostDenoiser::create( + &options.accelerators, + &options.device, + layout.width, + layout.height, + luma_options, + ) + }) + .transpose()?; + + let chroma = denoise_chroma + .then(|| { + let chroma_options = options.denoiser_options(ChannelMode::Chroma, layout.depth); + HostDenoiser::create( + &options.accelerators, + &options.device, + chroma_width, + chroma_height, + chroma_options, + ) + }) + .transpose()?; + + let yuv = denoise_yuv + .then(|| { + let yuv_options = options.denoiser_options(ChannelMode::Yuv, layout.depth); + HostDenoiser::create( + &options.accelerators, + &options.device, + layout.width, + layout.height, + yuv_options, + ) + }) + .transpose()?; + + let temporal_radius = match options.mode { + DenoisingMode::Spacial => 0, + DenoisingMode::Temporal { radius } => radius, + }; + + Ok(Self { + layout, + luma, + chroma, + yuv, + luma_passthrough: VecDeque::new(), + chroma_passthrough: VecDeque::new(), + temporal_radius, + }) + } + + /// The temporal radius the underlying denoisers run at. + pub fn temporal_radius(&self) -> u32 { + self.temporal_radius + } + + /// Pushes one planar frame. + /// + /// On `QueueFull` the caller should receive one frame and then retry the whole call. Any other + /// error is passed on unchanged. The denoiser push runs before either passthrough queue is touched, + /// so a retry replays the frame cleanly instead of queueing the disabled side's plane twice. + /// + /// # Why a retry cannot duplicate a frame + /// + /// In `LumaChroma` mode a retry pushes again into whichever half already succeeded, which would + /// duplicate that half's frame if the two halves could sit at different fill levels. They cannot. + /// Both share a temporal radius and a [MAX_PENDING](crate::MAX_PENDING) ceiling, and every + /// successful push or receive moves both on by exactly one frame. A failed push moves neither, + /// because the `QueueFull` check runs before anything changes. So if the luma push succeeds then + /// the chroma push succeeds too. + pub fn push(&mut self, planes: &Planes) -> Result<(), DenoiserError> { + self.push_with(planes, HostDenoiser::push) + } + + /// Uploads one planar frame into the temporal window without starting a denoise. + /// + /// It queues the disabled side's passthrough plane like [Self::push], but never produces output. + fn push_priming(&mut self, planes: &Planes) -> Result<(), DenoiserError> { + self.push_with(planes, HostDenoiser::push_priming) + } + + /// Runs `push_frame` against whichever of `yuv`, `luma` and `chroma` is enabled. + /// + /// The planes go over as wire bytes, so the normalisation and the channel interleave both happen + /// on the GPU. + fn push_with(&mut self, planes: &Planes, push_frame: WirePush) -> Result<(), DenoiserError> { + self.check_planes(planes)?; + + if let Some(denoiser) = self.yuv.as_mut() { + push_frame(denoiser, &[&planes.y, &planes.u, &planes.v])?; + return Ok(()); + } + + if let Some(denoiser) = self.luma.as_mut() { + push_frame(denoiser, &[&planes.y])?; + } + + if let Some(denoiser) = self.chroma.as_mut() { + push_frame(denoiser, &[&planes.u, &planes.v])?; + } + + if self.luma.is_none() { + self.luma_passthrough.push_back(planes.y.clone()); + } + + if self.chroma.is_none() { + let chroma_pair = (planes.u.clone(), planes.v.clone()); + self.chroma_passthrough.push_back(chroma_pair); + } + + Ok(()) + } + + /// Rejects a frame whose plane lengths do not match the layout. + /// + /// A half that rejects a plane stays usable, so checking every plane before either half is pushed keeps + /// luma and chroma from drifting a frame apart. + fn check_planes(&self, planes: &Planes) -> Result<(), DenoiserError> { + let luma_bytes = self.layout.luma_bytes(); + let chroma_bytes = self.layout.chroma_bytes(); + let expected = [ + ("Y", planes.y.len(), luma_bytes), + ("U", planes.u.len(), chroma_bytes), + ("V", planes.v.len(), chroma_bytes), + ]; + + for (name, length, expected_length) in expected { + if length != expected_length { + let message = format!("plane {name} holds {length} bytes, expected {expected_length}"); + let error = av_denoise_core::Error::PlaneMismatch(message); + return Err(DenoiserError::Engine(error)); + } + } + + Ok(()) + } + + /// Blocks until each enabled half emits one frame, then reassembles them into a planar frame. + /// + /// Returns `Ok(None)` if neither half had pending output. + pub fn recv(&mut self) -> Result, anyhow::Error> { + if let Some(denoiser) = self.yuv.as_mut() { + let received = denoiser.recv()?; + let planes = received.map(into_yuv); + return Ok(planes); + } + + let luma_out = self + .luma + .as_mut() + .map(|denoiser| denoiser.recv()) + .transpose()? + .flatten() + .map(into_luma); + + let chroma_out = self + .chroma + .as_mut() + .map(|denoiser| denoiser.recv()) + .transpose()? + .flatten() + .map(into_uv); + + // A disabled side has no HostDenoiser to query, so it pops its matching source plane when the + // enabled side produced output. + let luma_passthrough = if self.luma.is_none() && chroma_out.is_some() { + self.luma_passthrough.pop_front() + } else { + None + }; + + let chroma_passthrough = if self.chroma.is_none() && luma_out.is_some() { + self.chroma_passthrough.pop_front() + } else { + None + }; + + if luma_out.is_none() && chroma_out.is_none() { + return Ok(None); + } + + let planes = self.assemble(luma_out, chroma_out, luma_passthrough, chroma_passthrough); + + Ok(Some(planes)) + } + + /// Reads back the luma grain chunks measured since the last call, in frame order. + /// + /// Measured chunks stay on the GPU until drained, so callers drain after each flush. + pub fn drain_grain_chunks(&mut self) -> Result, anyhow::Error> { + let source = self.luma.as_mut().or(self.yuv.as_mut()); + let Some(denoiser) = source else { + return Ok(Vec::new()); + }; + + let chunks = denoiser.drain_grain_chunks()?; + + Ok(chunks) + } + + /// Drains the temporal tail of both halves. + /// + /// `sink` is called once per emitted planar frame. + pub fn flush(&mut self, mut sink: impl FnMut(Planes)) -> Result<(), anyhow::Error> { + if let Some(denoiser) = self.yuv.as_mut() { + denoiser.flush(|planes| { + let frame = into_yuv(planes); + sink(frame); + })?; + return Ok(()); + } + + let mut luma_frames: Vec> = Vec::new(); + let mut chroma_frames: Vec<(Vec, Vec)> = Vec::new(); + + if let Some(denoiser) = self.luma.as_mut() { + denoiser.flush(|planes| { + let luma = into_luma(planes); + luma_frames.push(luma); + })?; + } + + if let Some(denoiser) = self.chroma.as_mut() { + denoiser.flush(|planes| { + let chroma = into_uv(planes); + chroma_frames.push(chroma); + })?; + } + + // The two halves run in lockstep, so they flush the same number of frames. For each frame the + // disabled side, if there is one, pops the matching source plane from its passthrough queue. + let count = luma_frames.len().max(chroma_frames.len()); + + for i in 0..count { + let y_plane = if let Some(frame) = luma_frames.get_mut(i) { + std::mem::take(frame) + } else if let Some(source) = self.luma_passthrough.pop_front() { + source + } else { + self.layout.black_luma_plane() + }; + + let (u_plane, v_plane) = if let Some(pair) = chroma_frames.get_mut(i) { + std::mem::take(pair) + } else if let Some((source_u, source_v)) = self.chroma_passthrough.pop_front() { + (source_u, source_v) + } else { + ( + self.layout.neutral_chroma_plane(), + self.layout.neutral_chroma_plane(), + ) + }; + + sink(Planes { + y: y_plane, + u: u_plane, + v: v_plane, + }); + } + + if !self.luma_passthrough.is_empty() || !self.chroma_passthrough.is_empty() { + tracing::warn!( + luma_remaining = self.luma_passthrough.len(), + chroma_remaining = self.chroma_passthrough.len(), + "passthrough queue not fully drained after flush", + ); + self.luma_passthrough.clear(); + self.chroma_passthrough.clear(); + } + + Ok(()) + } + + /// The number of frames behind and ahead of a target frame a [Self::reseed] window must supply. + /// + /// Every owned `HostDenoiser` was built from the same algorithm, so any one of them answers for all + /// of them. + pub fn window_span(&self) -> WindowSpan { + self.yuv + .as_ref() + .or(self.luma.as_ref()) + .or(self.chroma.as_ref()) + .expect("PlanarDenoiser always keeps at least one HostDenoiser") + .window_span() + } + + fn assemble( + &self, + luma: Option>, + chroma: Option<(Vec, Vec)>, + luma_passthrough: Option>, + chroma_passthrough: Option<(Vec, Vec)>, + ) -> Planes { + let y_plane = match (luma, luma_passthrough) { + (Some(plane), _) => plane, + (None, Some(source)) => source, + (None, None) => self.layout.black_luma_plane(), + }; + + let (u_plane, v_plane) = match (chroma, chroma_passthrough) { + (Some(pair), _) => pair, + (None, Some(source)) => source, + (None, None) => ( + self.layout.neutral_chroma_plane(), + self.layout.neutral_chroma_plane(), + ), + }; + + Planes { + y: y_plane, + u: u_plane, + v: v_plane, + } + } +} diff --git a/av-denoise-core/src/frame/reseed.rs b/av-denoise/src/planar/reseed.rs similarity index 90% rename from av-denoise-core/src/frame/reseed.rs rename to av-denoise/src/planar/reseed.rs index c212ceb..9d32438 100644 --- a/av-denoise-core/src/frame/reseed.rs +++ b/av-denoise/src/planar/reseed.rs @@ -1,6 +1,6 @@ use std::collections::VecDeque; -use super::{PlanarDenoiser, Planes}; +use super::{DenoiserError, PlanarDenoiser, Planes}; use crate::EdgePadding; /// An explicit window of source frames for [PlanarDenoiser::reseed_window]. @@ -64,6 +64,8 @@ impl PlanarDenoiser { return target_output.ok_or_else(|| anyhow::anyhow!("reseed produced no frame, this is a bug")); } + self.check_window(window)?; + self.luma_passthrough.clear(); self.chroma_passthrough.clear(); self.reset_streams(); @@ -86,8 +88,8 @@ impl PlanarDenoiser { let mut result = None; for planes in tail { self.push(planes)?; - if let Some(out) = self.recv()? { - result = Some(out); + if let Some(denoised) = self.recv()? { + result = Some(denoised); } } @@ -138,6 +140,8 @@ impl PlanarDenoiser { anyhow::bail!("reseed window does not match the span {span:?} for its edge flags"); } + self.check_window(frames)?; + self.luma_passthrough.clear(); self.chroma_passthrough.clear(); self.reset_streams(); @@ -158,8 +162,8 @@ impl PlanarDenoiser { let mut outputs = Vec::new(); for planes in real_frames { self.push(planes)?; - if let Some(out) = self.recv()? { - outputs.push(out); + if let Some(denoised) = self.recv()? { + outputs.push(denoised); } } @@ -167,7 +171,8 @@ impl PlanarDenoiser { let target_output = outputs.pop(); let target_output = target_output .ok_or_else(|| anyhow::anyhow!("a full window produced no frame, this is a bug"))?; - return Ok(vec![target_output]); + let target_only = vec![target_output]; + return Ok(target_only); } let tail_count = frames.len().min(2 * radius); @@ -184,6 +189,15 @@ impl PlanarDenoiser { Ok(target_onwards) } + /// Checks every frame of a window, so a bad plane is rejected before the running stream is reset. + fn check_window(&self, frames: &[Planes]) -> Result<(), DenoiserError> { + for planes in frames { + self.check_planes(planes)?; + } + + Ok(()) + } + /// Abandons the running stream on every enabled half. fn reset_streams(&mut self) { let halves = [self.yuv.as_mut(), self.luma.as_mut(), self.chroma.as_mut()]; @@ -204,8 +218,8 @@ impl PlanarDenoiser { } if self.chroma.is_none() { - self.chroma_passthrough - .push_back((planes.u.clone(), planes.v.clone())); + let chroma_pair = (planes.u.clone(), planes.v.clone()); + self.chroma_passthrough.push_back(chroma_pair); } } } diff --git a/av-denoise/src/planar/tests/chroma_dims.rs b/av-denoise/src/planar/tests/chroma_dims.rs new file mode 100644 index 0000000..83e8b14 --- /dev/null +++ b/av-denoise/src/planar/tests/chroma_dims.rs @@ -0,0 +1,41 @@ +use super::*; + +#[test] +fn yuv420_even_dims_halve() { + assert_eq!(Subsampling::Yuv420.chroma_dims(1920, 1080), (960, 540)); +} + +#[test] +fn yuv420_odd_width_rounds_up() { + assert_eq!(Subsampling::Yuv420.chroma_dims(1919, 1080), (960, 540)); +} + +#[test] +fn yuv420_odd_height_rounds_up() { + assert_eq!(Subsampling::Yuv420.chroma_dims(1920, 1079), (960, 540)); +} + +#[test] +fn yuv420_odd_both_dims_round_up() { + assert_eq!(Subsampling::Yuv420.chroma_dims(1919, 1079), (960, 540)); +} + +#[test] +fn yuv422_even_width_halves() { + assert_eq!(Subsampling::Yuv422.chroma_dims(1920, 1080), (960, 1080)); +} + +#[test] +fn yuv422_odd_width_rounds_up() { + assert_eq!(Subsampling::Yuv422.chroma_dims(1919, 1080), (960, 1080)); +} + +#[test] +fn yuv444_passes_even_dims_through() { + assert_eq!(Subsampling::Yuv444.chroma_dims(1920, 1080), (1920, 1080)); +} + +#[test] +fn yuv444_passes_odd_dims_through() { + assert_eq!(Subsampling::Yuv444.chroma_dims(1919, 1079), (1919, 1079)); +} diff --git a/av-denoise/src/planar/tests/cli_options.rs b/av-denoise/src/planar/tests/cli_options.rs new file mode 100644 index 0000000..044aab6 --- /dev/null +++ b/av-denoise/src/planar/tests/cli_options.rs @@ -0,0 +1,207 @@ +use av_denoise_core::NlmeansAlgorithm; + +use super::*; +use crate::backend::EngineSpec; + +/// A `PlaneOptions` with every field other than the four arguments at a neutral default. +fn base_options( + mode: DenoisingMode, + algorithm: Algorithm, + luma_strength: Option, + chroma_strength: Option, +) -> PlaneOptions { + PlaneOptions { + accelerators: vec![], + device: Device::Default, + intent: ChannelIntent::LumaChroma, + mode, + algorithm, + luma_strength, + chroma_strength, + luma_lambda_ht: None, + chroma_lambda_ht: None, + } +} + +#[test] +fn luma_strength_alone_overrides_only_the_luma_plane() { + let algorithm = Algorithm::default(); + let plane_options = base_options(DenoisingMode::Spacial, algorithm, Some(0.7), None); + + let luma_options = plane_options.denoiser_options(ChannelMode::Luma, Depth::Eight); + let chroma_options = plane_options.denoiser_options(ChannelMode::Chroma, Depth::Eight); + let luma = expect_nlmeans(luma_options.algorithm); + let chroma = expect_nlmeans(chroma_options.algorithm); + + assert!( + matches!(luma.tuning.strength, Some(strength) if (strength - 0.7).abs() < f32::EPSILON), + "expected luma tuning.strength = Some(0.7), got {:?}", + luma.tuning.strength + ); + assert_eq!( + chroma.tuning.strength, None, + "chroma plane should carry no override so the table default applies" + ); +} + +#[test] +fn both_per_plane_strengths_set_independently() { + let algorithm = Algorithm::default(); + let plane_options = base_options(DenoisingMode::Spacial, algorithm, Some(0.7), Some(0.3)); + + let luma_options = plane_options.denoiser_options(ChannelMode::Luma, Depth::Eight); + let chroma_options = plane_options.denoiser_options(ChannelMode::Chroma, Depth::Eight); + let luma = expect_nlmeans(luma_options.algorithm); + let chroma = expect_nlmeans(chroma_options.algorithm); + + assert!( + matches!(luma.tuning.strength, Some(strength) if (strength - 0.7).abs() < f32::EPSILON), + "expected luma tuning.strength = Some(0.7), got {:?}", + luma.tuning.strength + ); + assert!( + matches!(chroma.tuning.strength, Some(strength) if (strength - 0.3).abs() < f32::EPSILON), + "expected chroma tuning.strength = Some(0.3), got {:?}", + chroma.tuning.strength + ); +} + +#[test] +fn no_overrides_hq_leaves_strength_to_the_per_plane_table() { + let hq_options = NlmeansHqOptions::default(); + let plane_options = base_options( + DenoisingMode::Temporal { radius: 4 }, + Algorithm::NlmeansHq(hq_options), + None, + None, + ); + + for channels in [ChannelMode::Luma, ChannelMode::Chroma] { + let options = plane_options.denoiser_options(channels, Depth::Eight); + let spec = options.algorithm.engine_spec(&options, 16, 16); + + let EngineSpec::Nlmeans { + algorithm: NlmeansAlgorithm::Hq(hq), + geometry, + } = spec + else { + panic!("expected an HQ nlmeans spec for {channels:?}, got {spec:?}"); + }; + + assert_eq!(geometry.channels, channels); + assert_eq!(hq.nlm.mode, DenoisingMode::Temporal { radius: 4 }); + assert_eq!( + hq.nlm.tuning.strength, None, + "{channels:?} should use the calibrated table" + ); + } +} + +/// A `PlaneOptions` running `Algorithm::Nl4d` with only the two `lambda_ht` overrides set. +fn nl4d_options(luma_lambda_ht: Option, chroma_lambda_ht: Option) -> PlaneOptions { + let default_nl4d = Nl4dOptions::default(); + + PlaneOptions { + accelerators: vec![], + device: Device::Default, + intent: ChannelIntent::LumaChroma, + mode: DenoisingMode::Temporal { radius: 2 }, + algorithm: Algorithm::Nl4d(default_nl4d), + luma_strength: None, + chroma_strength: None, + luma_lambda_ht, + chroma_lambda_ht, + } +} + +fn expect_nlmeans(algorithm: Algorithm) -> NlmeansOptions { + match algorithm { + Algorithm::Nlmeans(options) => options, + other => panic!("expected Algorithm::Nlmeans, got {other:?}"), + } +} + +fn expect_nl4d(algorithm: Algorithm) -> Nl4dOptions { + match algorithm { + Algorithm::Nl4d(options) => options, + other => panic!("expected Algorithm::Nl4d, got {other:?}"), + } +} + +#[test] +fn luma_lambda_ht_alone_overrides_only_the_luma_instance_for_nl4d() { + let plane_options = nl4d_options(Some(4.0), None); + let default_lambda_ht = Nl4dOptions::default().lambda_ht; + + let luma_algorithm = plane_options.algorithm_for(ChannelMode::Luma); + let chroma_algorithm = plane_options.algorithm_for(ChannelMode::Chroma); + let luma = expect_nl4d(luma_algorithm); + let chroma = expect_nl4d(chroma_algorithm); + + assert!((luma.lambda_ht.unwrap() - 4.0).abs() < f32::EPSILON); + assert_eq!( + chroma.lambda_ht, default_lambda_ht, + "chroma should stay unresolved here (None), deferred to its own per-plane \ + default at construction, got {:?}", + chroma.lambda_ht + ); +} + +#[test] +fn chroma_lambda_ht_alone_overrides_only_the_chroma_instance_for_nl4d() { + let plane_options = nl4d_options(None, Some(4.0)); + let default_lambda_ht = Nl4dOptions::default().lambda_ht; + + let luma_algorithm = plane_options.algorithm_for(ChannelMode::Luma); + let chroma_algorithm = plane_options.algorithm_for(ChannelMode::Chroma); + let luma = expect_nl4d(luma_algorithm); + let chroma = expect_nl4d(chroma_algorithm); + + assert_eq!( + luma.lambda_ht, default_lambda_ht, + "luma should stay unresolved here (None), deferred to its own per-plane \ + default at construction, got {:?}", + luma.lambda_ht + ); + assert!((chroma.lambda_ht.unwrap() - 4.0).abs() < f32::EPSILON); +} + +#[test] +fn both_planes_lambda_ht_set_independently_for_nl4d() { + let plane_options = nl4d_options(Some(2.0), Some(3.5)); + + let luma_algorithm = plane_options.algorithm_for(ChannelMode::Luma); + let chroma_algorithm = plane_options.algorithm_for(ChannelMode::Chroma); + let luma = expect_nl4d(luma_algorithm); + let chroma = expect_nl4d(chroma_algorithm); + + assert!((luma.lambda_ht.unwrap() - 2.0).abs() < f32::EPSILON); + assert!((chroma.lambda_ht.unwrap() - 3.5).abs() < f32::EPSILON); + + // Every other field stays shared between the two instances even though lambda_ht diverges. + assert_eq!(luma.refine, chroma.refine); + assert_eq!(luma.spatial_radius, chroma.spatial_radius); + assert!((luma.c_min - chroma.c_min).abs() < f32::EPSILON); +} + +#[test] +fn unset_nl4d_overrides_resolve_to_different_lambda_ht_per_plane_end_to_end() { + let plane_options = nl4d_options(None, None); + + let luma_algorithm = plane_options.algorithm_for(ChannelMode::Luma); + let chroma_algorithm = plane_options.algorithm_for(ChannelMode::Chroma); + let luma = expect_nl4d(luma_algorithm); + let chroma = expect_nl4d(chroma_algorithm); + + // Neither plane has anything set, so both stay unresolved at this layer. + assert_eq!(luma.lambda_ht, None); + assert_eq!(chroma.lambda_ht, None); + + // Construction resolves each through `nl4d_default_lambda_ht`, which gives luma and chroma + // different values. + let luma_default = crate::nl4d_default_lambda_ht(ChannelMode::Luma); + let chroma_default = crate::nl4d_default_lambda_ht(ChannelMode::Chroma); + assert!((luma_default - 4.158).abs() < f32::EPSILON); + assert!((chroma_default - 3.234).abs() < f32::EPSILON); + assert!((chroma_default - luma_default).abs() > f32::EPSILON); +} diff --git a/av-denoise/src/planar/tests/layout.rs b/av-denoise/src/planar/tests/layout.rs new file mode 100644 index 0000000..fcdf705 --- /dev/null +++ b/av-denoise/src/planar/tests/layout.rs @@ -0,0 +1,45 @@ +use super::*; + +fn layout(depth: Depth) -> FrameLayout { + FrameLayout { + width: 4, + height: 4, + subsampling: Subsampling::Yuv420, + depth, + } +} + +#[test] +fn byte_lengths_scale_with_depth() { + let eight_bit = layout(Depth::Eight); + let ten_bit = layout(Depth::Ten); + + assert_eq!(eight_bit.luma_bytes(), 16); + assert_eq!(ten_bit.luma_bytes(), 32); + assert_eq!(eight_bit.chroma_bytes(), 4); + assert_eq!(ten_bit.chroma_bytes(), 8); +} + +#[test] +fn neutral_chroma_fill_is_correct_at_each_depth() { + let eight = layout(Depth::Eight).neutral_chroma_plane(); + assert_eq!(eight, vec![128u8; 4]); + + // 512 little-endian is [0x00, 0x02], repeated per sample. + let ten = layout(Depth::Ten).neutral_chroma_plane(); + assert_eq!(ten, vec![0x00, 0x02, 0x00, 0x02, 0x00, 0x02, 0x00, 0x02]); + + // 2048 little-endian is [0x00, 0x08]. + let twelve = layout(Depth::Twelve).neutral_chroma_plane(); + assert_eq!(twelve.len(), 8); + assert_eq!(&twelve[0..2], &[0x00, 0x08]); +} + +#[test] +fn black_luma_fill_is_zero_at_the_right_length() { + let eight = layout(Depth::Eight).black_luma_plane(); + let ten = layout(Depth::Ten).black_luma_plane(); + + assert_eq!(eight, vec![0u8; 16]); + assert_eq!(ten, vec![0u8; 32]); +} diff --git a/av-denoise/src/planar/tests/lockstep.rs b/av-denoise/src/planar/tests/lockstep.rs new file mode 100644 index 0000000..c011a25 --- /dev/null +++ b/av-denoise/src/planar/tests/lockstep.rs @@ -0,0 +1,188 @@ +use super::*; +use crate::accelerate::Accelerator; +use crate::{Algorithm, DenoisingMode}; + +/// Runs `luma` and `chroma` as two real `HostDenoiser`s in spatial mode. +/// +/// Spatial mode passes a uniform-valued plane through unchanged, so each plane can carry its own +/// marker value and the two halves drifting apart shows up. +fn luma_chroma_options() -> PlaneOptions { + PlaneOptions { + accelerators: vec![Accelerator::Vulkan], + device: Device::Default, + intent: ChannelIntent::LumaChroma, + mode: DenoisingMode::Spacial, + algorithm: Algorithm::default(), + luma_strength: None, + chroma_strength: None, + luma_lambda_ht: None, + chroma_lambda_ht: None, + } +} + +/// A uniform-valued frame whose luma and chroma planes encode `frame_index` with different formulas. +/// +/// Pairing luma from one push with chroma from another makes the two encodings disagree. +fn marked_planes(layout: FrameLayout, frame_index: u8) -> Planes { + let luma_pixels = layout.luma_pixels(); + let chroma_pixels = layout.chroma_pixels(); + let luma_marker = 10 + frame_index; + let chroma_marker = 200 - frame_index; + + Planes { + y: fill_plane(luma_pixels, luma_marker as u16, layout.depth), + u: fill_plane(chroma_pixels, chroma_marker as u16, layout.depth), + v: fill_plane(chroma_pixels, chroma_marker as u16, layout.depth), + } +} + +#[test] +fn distinct_u_and_v_planes_come_back_in_order() { + let layout = FrameLayout { + width: 16, + height: 16, + subsampling: Subsampling::Yuv420, + depth: Depth::Eight, + }; + let options = luma_chroma_options(); + let mut denoiser = PlanarDenoiser::create(&options, layout).expect("denoiser construction failed"); + + let luma_pixels = layout.luma_pixels(); + let chroma_pixels = layout.chroma_pixels(); + let planes = Planes { + y: fill_plane(luma_pixels, 100, layout.depth), + u: fill_plane(chroma_pixels, 60, layout.depth), + v: fill_plane(chroma_pixels, 190, layout.depth), + }; + denoiser.push(&planes).expect("push failed"); + + let received = denoiser.recv().expect("recv failed"); + let denoised = received.expect("spatial mode emits one frame per push"); + + for &sample in &denoised.u { + assert!(sample.abs_diff(60) <= 2, "U sample {sample}, expected about 60"); + } + + for &sample in &denoised.v { + assert!(sample.abs_diff(190) <= 2, "V sample {sample}, expected about 190"); + } +} + +#[test] +fn queue_full_retries_never_desync_luma_and_chroma() { + let layout = FrameLayout { + width: 16, + height: 16, + subsampling: Subsampling::Yuv420, + depth: Depth::Eight, + }; + let options = luma_chroma_options(); + let mut denoiser = PlanarDenoiser::create(&options, layout).expect("denoiser construction failed"); + + // More pushes than the depth-2 pipeline holds, so this drives several `QueueFull` retries. + const FRAME_COUNT: u8 = 6; + let mut outputs: Vec = Vec::new(); + + for frame_index in 0..FRAME_COUNT { + let planes = marked_planes(layout, frame_index); + + // The push, drain and retry sequence a streaming caller runs. + let pushed = denoiser.push(&planes); + let needs_retry = push_needs_retry(pushed).expect("push_needs_retry"); + if needs_retry { + if let Some(denoised) = denoiser.recv().expect("recv failed") { + outputs.push(denoised); + } + + denoiser + .push(&planes) + .expect("retry push should land after drain"); + } + } + + denoiser + .flush(|denoised| outputs.push(denoised)) + .expect("flush failed"); + + assert_eq!( + outputs.len(), + FRAME_COUNT as usize, + "expected exactly one output frame per input frame, got {}", + outputs.len() + ); + + for (position, denoised) in outputs.iter().enumerate() { + let luma_marker = denoised.y[0]; + let chroma_marker = denoised.u[0]; + let index_from_luma = luma_marker - 10; + let index_from_chroma = 200 - chroma_marker; + + assert_eq!( + index_from_luma, index_from_chroma, + "luma marker {luma_marker} (frame {index_from_luma}) and chroma marker {chroma_marker} \ + (frame {index_from_chroma}) disagree, so the luma and chroma pushes have drifted apart" + ); + assert_eq!( + index_from_luma as usize, position, + "output {position} carries frame {index_from_luma}, so frames came back out of order" + ); + } +} + +#[test] +fn a_wrong_length_u_plane_is_rejected_without_advancing_either_half() { + let layout = FrameLayout { + width: 16, + height: 16, + subsampling: Subsampling::Yuv420, + depth: Depth::Eight, + }; + let options = luma_chroma_options(); + let mut denoiser = PlanarDenoiser::create(&options, layout).expect("denoiser construction failed"); + + let mut short_u = marked_planes(layout, 0); + short_u.u.pop(); + + let rejected = denoiser.push(&short_u); + let is_plane_mismatch = matches!( + rejected, + Err(DenoiserError::Engine(av_denoise_core::Error::PlaneMismatch(_))) + ); + assert!(is_plane_mismatch, "got {rejected:?}"); + + const FRAME_COUNT: u8 = 4; + let mut outputs: Vec = Vec::new(); + + for frame_index in 0..FRAME_COUNT { + let planes = marked_planes(layout, frame_index); + + let pushed = denoiser.push(&planes); + let needs_retry = push_needs_retry(pushed).expect("push_needs_retry"); + if needs_retry { + if let Some(denoised) = denoiser.recv().expect("recv failed") { + outputs.push(denoised); + } + + denoiser + .push(&planes) + .expect("retry push should land after drain"); + } + } + + denoiser + .flush(|denoised| outputs.push(denoised)) + .expect("flush failed"); + + assert_eq!(outputs.len(), FRAME_COUNT as usize); + + for (position, denoised) in outputs.iter().enumerate() { + let index_from_luma = denoised.y[0] - 10; + let index_from_chroma = 200 - denoised.u[0]; + + assert_eq!(index_from_luma as usize, position, "luma came back out of order"); + assert_eq!( + index_from_chroma as usize, position, + "chroma came back out of order" + ); + } +} diff --git a/av-denoise/src/planar/tests/mod.rs b/av-denoise/src/planar/tests/mod.rs new file mode 100644 index 0000000..3f40f46 --- /dev/null +++ b/av-denoise/src/planar/tests/mod.rs @@ -0,0 +1,1119 @@ +mod chroma_dims; +mod cli_options; +mod layout; +// Gated on `vulkan` because `luma_chroma_options` names the `Vulkan` accelerator variant. +#[cfg(feature = "vulkan")] +mod lockstep; +// Gated on `vulkan` because `chroma_only_options` names the `Vulkan` accelerator variant. +#[cfg(feature = "vulkan")] +mod passthrough_retry; +mod push_retry; + +use super::*; + +// The tests name the `Vulkan` accelerator, which only exists with the `vulkan` feature. +#[cfg(feature = "vulkan")] +mod reseed { + use super::*; + use crate::HqParams; + use crate::accelerate::Accelerator; + + fn layout() -> FrameLayout { + FrameLayout { + width: 64, + height: 64, + subsampling: Subsampling::Yuv420, + depth: Depth::Eight, + } + } + + /// Temporal nlmeans at `radius`, denoising both planes independently. + fn test_plane_options(radius: u32) -> PlaneOptions { + PlaneOptions { + accelerators: vec![Accelerator::Vulkan], + device: Device::Default, + intent: ChannelIntent::LumaChroma, + mode: DenoisingMode::Temporal { radius }, + algorithm: Algorithm::default(), + luma_strength: None, + chroma_strength: None, + luma_lambda_ht: None, + chroma_lambda_ht: None, + } + } + + fn test_plane_options_with_intent(radius: u32, intent: ChannelIntent) -> PlaneOptions { + PlaneOptions { + intent, + ..test_plane_options(radius) + } + } + + /// A small xorshift generator, so the test data is the same on every run. + fn pseudo_random(mut state: u64) -> u64 { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + state + } + + /// One plane's bytes for frame `frame_index`. + /// + /// A spatial ramp, a per-frame offset and a deterministic dither are summed and clamped, so a + /// temporal filter has real signal and real noise to work with. + fn ramp_plane(pixels: usize, width: u32, frame_index: usize, plane_seed: u64) -> Vec { + let width = width.max(1) as usize; + + (0..pixels) + .map(|i| { + let x = (i % width) as u32; + let y = (i / width) as u32; + let spatial = x.wrapping_add(y) % 120; + let frame_offset = (frame_index as u32 * 7) % 60; + let seed = (i as u64) ^ (frame_index as u64).wrapping_mul(0x9E3779B97F4A7C15) ^ plane_seed; + let dither = (pseudo_random(seed) % 16) as u32; + let value = 20 + spatial + frame_offset + dither; + value.min(235) as u8 + }) + .collect() + } + + /// `count` frames whose bytes vary per frame and per pixel, so a temporal filter sees a + /// non-degenerate signal. + fn ramp_clip(layout: &FrameLayout, count: usize) -> Vec { + let (chroma_width, _) = layout.chroma_dims(); + + (0..count) + .map(|frame_index| { + let y_plane = ramp_plane(layout.luma_pixels(), layout.width, frame_index, 1); + let u_plane = ramp_plane(layout.chroma_pixels(), chroma_width, frame_index, 2); + let v_plane = ramp_plane(layout.chroma_pixels(), chroma_width, frame_index, 3); + + Planes { + y: y_plane, + u: u_plane, + v: v_plane, + } + }) + .collect() + } + + /// Renders every frame through the streaming path. + fn stream_all(options: &PlaneOptions, frames: &[Planes]) -> Vec { + let frame_layout = layout(); + stream_all_with_layout(options, frame_layout, frames) + } + + fn stream_all_with_layout( + options: &PlaneOptions, + frame_layout: FrameLayout, + frames: &[Planes], + ) -> Vec { + let mut denoiser = PlanarDenoiser::create(options, frame_layout).unwrap(); + let mut outputs = Vec::new(); + for frame in frames { + denoiser.push(frame).unwrap(); + if let Some(planes) = denoiser.recv().unwrap() { + outputs.push(planes); + } + } + + denoiser.flush(|planes| outputs.push(planes)).unwrap(); + outputs + } + + fn window_of(frames: &[Planes], target: usize, radius: usize) -> Vec { + (0..(2 * radius + 1)) + .map(|i| { + let index = (target + i).saturating_sub(radius).min(frames.len() - 1); + frames[index].clone() + }) + .collect() + } + + #[test] + fn reseed_matches_the_streaming_output_mid_clip() { + let frame_layout = layout(); + let options = test_plane_options(2); + let frames = ramp_clip(&frame_layout, 12); + let streamed = stream_all(&options, &frames); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let target_frame = 6; + let window = window_of(&frames, target_frame, 2); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!(got.y, streamed[target_frame].y); + assert_eq!(got.u, streamed[target_frame].u); + assert_eq!(got.v, streamed[target_frame].v); + } + + #[test] + fn reseed_matches_the_streaming_output_at_both_clip_edges() { + let frame_layout = layout(); + let options = test_plane_options(2); + let frames = ramp_clip(&frame_layout, 12); + let streamed = stream_all(&options, &frames); + let last = frames.len() - 1; + + for target_frame in [0usize, last] { + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let window = window_of(&frames, target_frame, 2); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!( + got.y, streamed[target_frame].y, + "luma mismatch at k = {target_frame}" + ); + assert_eq!( + got.u, streamed[target_frame].u, + "u mismatch at k = {target_frame}" + ); + assert_eq!( + got.v, streamed[target_frame].v, + "v mismatch at k = {target_frame}" + ); + } + } + + #[test] + fn reseed_recovers_a_half_poisoned_by_an_earlier_failure() { + let frame_layout = layout(); + let options = test_plane_options(2); + let frames = ramp_clip(&frame_layout, 12); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + denoiser.luma.as_mut().unwrap().poison_for_test(); + denoiser.chroma.as_mut().unwrap().poison_for_test(); + + let window = window_of(&frames, 6, 2); + let got = denoiser.reseed(&window).unwrap(); + + assert!(!got.y.is_empty()); + assert!(!got.u.is_empty()); + assert!(!got.v.is_empty()); + } + + /// Plain nlmeans carries no noise state between frames, so repeated out-of-order reseeds on one + /// denoiser must match streaming. + /// + /// The radius is wider and the clip longer and more shuffled than in any other reseed test here, + /// so a history-dependent regression has room to show. + #[test] + fn nlmeans_repeated_out_of_order_reseeds_match_streaming() { + let frame_layout = layout(); + let options = test_plane_options(4); + let frames = ramp_clip(&frame_layout, 24); + let streamed = stream_all(&options, &frames); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + + // Skews late, like the VapourSynth plugin harness's shuffled order. + let order = [ + 18, 4, 23, 9, 12, 2, 20, 6, 15, 1, 22, 7, 17, 3, 11, 19, 0, 21, 8, 16, 5, 14, 10, 13, + ]; + + for &target_frame in &order { + let window = window_of(&frames, target_frame, 4); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!( + got.y, streamed[target_frame].y, + "luma mismatch at k = {target_frame}" + ); + assert_eq!( + got.u, streamed[target_frame].u, + "u mismatch at k = {target_frame}" + ); + assert_eq!( + got.v, streamed[target_frame].v, + "v mismatch at k = {target_frame}" + ); + } + } + + #[test] + fn a_reseed_leaves_the_stream_positioned_for_the_next_frame() { + let frame_layout = layout(); + let options = test_plane_options(2); + let frames = ramp_clip(&frame_layout, 12); + let streamed = stream_all(&options, &frames); + let (target_frame, radius) = (6usize, 2usize); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let window = window_of(&frames, target_frame, radius); + denoiser.reseed(&window).unwrap(); + denoiser.push(&frames[target_frame + 1 + radius]).unwrap(); + let got = denoiser.recv().unwrap().expect("frame k + 1"); + + assert_eq!(got.y, streamed[target_frame + 1].y); + } + + #[test] + fn reseed_rejects_a_window_of_the_wrong_length() { + let frame_layout = layout(); + let options = test_plane_options(2); + let frames = ramp_clip(&frame_layout, 12); + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + + let error = denoiser.reseed(&frames[..3]).unwrap_err().to_string(); + + assert!( + error.contains("5"), + "error should name the expected length, got {error}" + ); + } + + /// Reseeds to `good_target`, tries a reseed whose window has one short U plane, then pushes the + /// next frame and returns its output. + /// + /// The bad reseed must be rejected, so the output matches a denoiser that never saw it. + fn next_frame_after_a_rejected_reseed( + options: &PlaneOptions, + frames: &[Planes], + good_target: usize, + bad_target: usize, + ) -> (Planes, Planes) { + let frame_layout = layout(); + let mut control = PlanarDenoiser::create(options, frame_layout).unwrap(); + let mut tested = PlanarDenoiser::create(options, frame_layout).unwrap(); + let span = control.window_span(); + + let good_window = window_of_span(frames, good_target, span); + let mut bad_window = window_of_span(frames, bad_target, span); + let middle = bad_window.len() / 2; + bad_window[middle].u.pop(); + + control.reseed(&good_window).unwrap(); + tested.reseed(&good_window).unwrap(); + + let rejected = tested.reseed(&bad_window); + assert!( + rejected.is_err(), + "a window with a short U plane should be rejected" + ); + + let next_frame = &frames[good_target + 1 + span.ahead]; + control.push(next_frame).unwrap(); + tested.push(next_frame).unwrap(); + + let expected = control.recv().unwrap().expect("control frame"); + let got = tested.recv().unwrap().expect("tested frame"); + (expected, got) + } + + #[test] + fn a_rejected_reseed_leaves_the_running_stream_untouched() { + let frame_layout = layout(); + let options = test_plane_options(2); + let frames = ramp_clip(&frame_layout, 12); + + let (expected, got) = next_frame_after_a_rejected_reseed(&options, &frames, 6, 3); + + assert_eq!(got.y, expected.y); + assert_eq!(got.u, expected.u); + assert_eq!(got.v, expected.v); + } + + #[test] + fn a_rejected_nl4d_reseed_window_leaves_the_running_stream_untouched() { + let frame_layout = layout(); + let options = nl4d_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + + let (expected, got) = next_frame_after_a_rejected_reseed(&options, &frames, 6, 9); + + assert_eq!(got.y, expected.y); + assert_eq!(got.u, expected.u); + assert_eq!(got.v, expected.v); + } + + /// `ChannelIntent::Luma` sends chroma through the passthrough queue, and a reseed queues one + /// entry per window frame. + /// + /// The entry paired with the denoised centre must be the centre frame's own chroma, not a + /// neighbour's. + #[test] + fn reseed_pairs_the_passthrough_plane_with_the_centre_frame() { + let frame_layout = layout(); + let options = test_plane_options_with_intent(2, ChannelIntent::Luma); + let frames = ramp_clip(&frame_layout, 12); + let (target_frame, radius) = (6usize, 2usize); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let window = window_of(&frames, target_frame, radius); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!( + got.u, frames[target_frame].u, + "u should pass through from the centre frame" + ); + assert_eq!( + got.v, frames[target_frame].v, + "v should pass through from the centre frame" + ); + } + + #[test] + fn reseed_pairs_the_passthrough_luma_plane_with_the_centre_frame() { + let frame_layout = layout(); + let options = test_plane_options_with_intent(2, ChannelIntent::Chroma); + let frames = ramp_clip(&frame_layout, 12); + let (target_frame, radius) = (6usize, 2usize); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let window = window_of(&frames, target_frame, radius); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!( + got.y, frames[target_frame].y, + "y should pass through from the centre frame" + ); + } + + /// A single reseed only checks the first passthrough entry `recv` pops. + /// + /// An extra or missing entry that still leaves the right one at the front only misaligns the + /// plane paired with the next frame, once streaming resumes. + #[test] + fn reseed_then_streaming_keeps_the_passthrough_plane_aligned_on_the_next_frame() { + let frame_layout = layout(); + let options = test_plane_options_with_intent(2, ChannelIntent::Luma); + let frames = ramp_clip(&frame_layout, 12); + let (target_frame, radius) = (6usize, 2usize); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let window = window_of(&frames, target_frame, radius); + denoiser.reseed(&window).unwrap(); + denoiser.push(&frames[target_frame + 1 + radius]).unwrap(); + let got = denoiser.recv().unwrap().expect("frame k + 1"); + + assert_eq!( + got.u, + frames[target_frame + 1].u, + "u should pass through from frame k + 1" + ); + assert_eq!( + got.v, + frames[target_frame + 1].v, + "v should pass through from frame k + 1" + ); + } + + /// Temporal nl4d at `radius` with a pinned `sigma`. + /// + /// The automatic estimate is a moving average over every frame since the stream last reset, + /// history a windowed `reseed` cannot supply. Pinning it keeps these tests on what the window + /// shape and pass sequence decide. + fn nl4d_plane_options(radius: u32) -> PlaneOptions { + let nl4d_options = Nl4dOptions { + sigma: Some(0.03), + ..Nl4dOptions::default() + }; + + PlaneOptions { + algorithm: Algorithm::Nl4d(nl4d_options), + ..test_plane_options(radius) + } + } + + /// The window `span` needs for frame `target`, clamped at both clip ends like `window_of`. + fn window_of_span(frames: &[Planes], target: usize, span: WindowSpan) -> Vec { + (0..span.frame_count()) + .map(|i| { + let index = (target + i).saturating_sub(span.behind).min(frames.len() - 1); + frames[index].clone() + }) + .collect() + } + + /// The shifted window around frame `target`. + /// + /// It stops at the clip's ends rather than repeating them, and returns the target's index in it. + fn shifted_window_of( + frames: &[Planes], + target: usize, + span: WindowSpan, + ) -> (Vec, ReseedWindowFlags) { + let first = target.saturating_sub(span.behind); + let last = (target + span.ahead).min(frames.len() - 1); + let window = frames[first..=last].to_vec(); + let flags = ReseedWindowFlags { + target: target - first, + at_clip_start: first == 0, + at_clip_end: last == frames.len() - 1, + }; + (window, flags) + } + + struct ReseedWindowFlags { + target: usize, + at_clip_start: bool, + at_clip_end: bool, + } + + fn reseed_shifted(denoiser: &mut PlanarDenoiser, frames: &[Planes], target: usize) -> Vec { + let span = denoiser.window_span(); + let (window, flags) = shifted_window_of(frames, target, span); + let request = ReseedWindow { + frames: &window, + target: flags.target, + at_clip_start: flags.at_clip_start, + at_clip_end: flags.at_clip_end, + }; + denoiser.reseed_window(request).unwrap() + } + + /// How many outputs a shifted reseed at `target` returns. + /// + /// That is the target alone mid-clip, or the target through the clip's end. + fn shifted_output_count(denoiser: &PlanarDenoiser, clip_len: usize, target: usize) -> usize { + let span = denoiser.window_span(); + if target + span.ahead >= clip_len - 1 { + clip_len - target + } else { + 1 + } + } + + #[test] + fn nl4d_reseed_matches_the_streaming_output_mid_clip() { + let frame_layout = layout(); + let options = nl4d_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let streamed = stream_all(&options, &frames); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let target_frame = 8; + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!(got.y, streamed[target_frame].y); + assert_eq!(got.u, streamed[target_frame].u); + assert_eq!(got.v, streamed[target_frame].v); + } + + fn max_abs_diff(left: &[u8], right: &[u8]) -> i32 { + left.iter() + .zip(right.iter()) + .map(|(&left_sample, &right_sample)| (left_sample as i32 - right_sample as i32).abs()) + .max() + .unwrap_or(0) + } + + #[test] + fn nl4d_reseed_window_matches_streaming_at_every_frame() { + let frame_layout = layout(); + let option_sets = [ + nl4d_plane_options(2), + nl4d_windowed_plane_options(2), + nl4d_plane_options_with_intent(2, ChannelIntent::Luma), + nl4d_plane_options_with_intent(2, ChannelIntent::Chroma), + ]; + + for options in option_sets { + for clip_len in [3usize, 7, 12] { + let frames = ramp_clip(&frame_layout, clip_len); + let streamed = stream_all(&options, &frames); + assert_eq!(streamed.len(), clip_len); + + for target_frame in 0..clip_len { + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let got = reseed_shifted(&mut denoiser, &frames, target_frame); + + let expected_len = shifted_output_count(&denoiser, clip_len, target_frame); + assert_eq!(got.len(), expected_len, "len={clip_len} k={target_frame}"); + + for (offset, planes) in got.iter().enumerate() { + let index = target_frame + offset; + assert_eq!( + planes.y, streamed[index].y, + "len={clip_len} k={target_frame} frame {index} luma" + ); + assert_eq!( + planes.u, streamed[index].u, + "len={clip_len} k={target_frame} frame {index} u" + ); + assert_eq!( + planes.v, streamed[index].v, + "len={clip_len} k={target_frame} frame {index} v" + ); + } + } + } + } + } + + /// `ChannelIntent::YuvFused` needs a 4:4:4 source, so it runs over its own layout. + #[test] + fn nl4d_reseed_window_matches_streaming_at_every_frame_in_yuv_fused_mode() { + let fused_layout = FrameLayout { + subsampling: Subsampling::Yuv444, + ..layout() + }; + let options = nl4d_plane_options_with_intent(2, ChannelIntent::YuvFused); + + for clip_len in [3usize, 7, 12] { + let frames = ramp_clip(&fused_layout, clip_len); + let streamed = stream_all_with_layout(&options, fused_layout, &frames); + assert_eq!(streamed.len(), clip_len); + + for target_frame in 0..clip_len { + let mut denoiser = PlanarDenoiser::create(&options, fused_layout).unwrap(); + let got = reseed_shifted(&mut denoiser, &frames, target_frame); + + let expected_len = shifted_output_count(&denoiser, clip_len, target_frame); + assert_eq!(got.len(), expected_len, "len={clip_len} k={target_frame}"); + + for (offset, planes) in got.iter().enumerate() { + let index = target_frame + offset; + assert_eq!( + planes.y, streamed[index].y, + "len={clip_len} k={target_frame} frame {index} luma" + ); + assert_eq!( + planes.u, streamed[index].u, + "len={clip_len} k={target_frame} frame {index} u" + ); + assert_eq!( + planes.v, streamed[index].v, + "len={clip_len} k={target_frame} frame {index} v" + ); + } + } + } + } + + #[test] + fn nl4d_reseed_window_pairs_passthrough_at_the_last_frame() { + let frame_layout = layout(); + let options = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); + let frames = ramp_clip(&frame_layout, 16); + let last = frames.len() - 1; + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + + let got = reseed_shifted(&mut denoiser, &frames, last); + + assert_eq!(got.len(), 1); + assert_eq!(got[0].u, frames[last].u); + assert_eq!(got[0].v, frames[last].v); + } + + #[test] + fn nl4d_reseed_window_pairs_passthrough_at_the_first_frame() { + let frame_layout = layout(); + let options = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); + let frames = ramp_clip(&frame_layout, 16); + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + + let got = reseed_shifted(&mut denoiser, &frames, 0); + + assert_eq!(got.len(), 1); + assert_eq!(got[0].u, frames[0].u); + assert_eq!(got[0].v, frames[0].v); + } + + #[test] + fn nl4d_reseed_window_then_streaming_continues_from_the_clip_start() { + let frame_layout = layout(); + let options = nl4d_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let streamed = stream_all(&options, &frames); + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + + reseed_shifted(&mut denoiser, &frames, 1); + denoiser.push(&frames[2 + span.ahead]).unwrap(); + let next = denoiser.recv().unwrap().unwrap(); + + assert_eq!(next.y, streamed[2].y); + } + + #[test] + fn nl4d_reseed_then_streaming_continues_correctly() { + let frame_layout = layout(); + let options = nl4d_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let streamed = stream_all(&options, &frames); + let target_frame = 8usize; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + denoiser.reseed(&window).unwrap(); + + // The next frame in source order after the reseed window's last frame is + // `target_frame + 1 + span.ahead`. + denoiser.push(&frames[target_frame + 1 + span.ahead]).unwrap(); + let got = denoiser.recv().unwrap().expect("frame k + 1"); + + assert_eq!(got.y, streamed[target_frame + 1].y); + assert_eq!(got.u, streamed[target_frame + 1].u); + assert_eq!(got.v, streamed[target_frame + 1].v); + } + + /// The error names nl4d's wider `4r+1` window length, not nlmeans's `2r+1`. + #[test] + fn nl4d_reseed_rejects_a_window_of_the_wrong_length() { + let frame_layout = layout(); + let options = nl4d_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let expected = denoiser.window_span().frame_count(); + let expected_text = expected.to_string(); + + let error = denoiser.reseed(&frames[..3]).unwrap_err().to_string(); + + assert!( + error.contains(&expected_text), + "error should name the expected length ({expected}), got {error}" + ); + } + + fn nl4d_plane_options_with_intent(radius: u32, intent: ChannelIntent) -> PlaneOptions { + PlaneOptions { + intent, + ..nl4d_plane_options(radius) + } + } + + /// nl4d drains after every emission during a reseed, not only the last, and each drain pops one + /// passthrough entry. + /// + /// The walk must still land on the target's own entry rather than an earlier, discarded one. + #[test] + fn nl4d_reseed_pairs_the_passthrough_plane_with_the_centre_frame() { + let frame_layout = layout(); + let options = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); + let frames = ramp_clip(&frame_layout, 16); + let target_frame = 8; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!( + got.u, frames[target_frame].u, + "u should pass through from the centre frame" + ); + assert_eq!( + got.v, frames[target_frame].v, + "v should pass through from the centre frame" + ); + } + + #[test] + fn nl4d_reseed_pairs_the_passthrough_luma_plane_with_the_centre_frame() { + let frame_layout = layout(); + let options = nl4d_plane_options_with_intent(2, ChannelIntent::Chroma); + let frames = ramp_clip(&frame_layout, 16); + let target_frame = 8; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!( + got.y, frames[target_frame].y, + "y should pass through from the centre frame" + ); + } + + /// A single-shot pairing test only checks the first entry `recv` pops after the drop. + /// + /// An extra or missing entry that still leaves the right one at the front only misaligns the + /// plane paired with the frame after the target, once streaming resumes. + #[test] + fn nl4d_reseed_then_streaming_keeps_the_passthrough_plane_aligned_on_the_next_frame() { + let frame_layout = layout(); + let options = nl4d_plane_options_with_intent(2, ChannelIntent::Luma); + let frames = ramp_clip(&frame_layout, 16); + let target_frame = 8; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + denoiser.reseed(&window).unwrap(); + denoiser.push(&frames[target_frame + 1 + span.ahead]).unwrap(); + let got = denoiser.recv().unwrap().expect("frame k + 1"); + + assert_eq!( + got.u, + frames[target_frame + 1].u, + "u should pass through from frame k + 1" + ); + assert_eq!( + got.v, + frames[target_frame + 1].v, + "v should pass through from frame k + 1" + ); + } + + /// Temporal nl4d at `radius` with window-local noise estimation and an automatic `sigma`, as + /// `av-denoise-vs` runs it. + /// + /// `sigma` stays unpinned because window-local estimation exists so the automatic estimate + /// agrees between `reseed` and streaming. + fn nl4d_windowed_plane_options(radius: u32) -> PlaneOptions { + let nl4d_options = Nl4dOptions { + windowed_noise_estimation: true, + ..Nl4dOptions::default() + }; + + PlaneOptions { + algorithm: Algorithm::Nl4d(nl4d_options), + ..test_plane_options(radius) + } + } + + /// Like `nl4d_reseed_matches_the_streaming_output_mid_clip`, with the noise estimator running + /// instead of pinned. + #[test] + fn nl4d_windowed_reseed_matches_the_streaming_output_mid_clip() { + let frame_layout = layout(); + let options = nl4d_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let streamed = stream_all(&options, &frames); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let target_frame = 8; + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!(got.y, streamed[target_frame].y); + assert_eq!(got.u, streamed[target_frame].u); + assert_eq!(got.v, streamed[target_frame].v); + } + + /// With window-local estimation, a `reseed` at a frame then a `push`/`recv` for the next frame must + /// match a `reseed` at the next frame on a fresh denoiser. + /// + /// Without it the fast path folds history the reseed path never sees, so the two disagree on the + /// same window of content. + #[test] + fn nl4d_windowed_fast_path_agrees_with_reseed_at_the_next_frame() { + let frame_layout = layout(); + let options = nl4d_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let target_frame = 8usize; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + denoiser.reseed(&window).unwrap(); + denoiser.push(&frames[target_frame + 1 + span.ahead]).unwrap(); + let via_fast_path = denoiser.recv().unwrap().expect("frame k + 1"); + + let mut fresh = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let next_window = window_of_span(&frames, target_frame + 1, span); + let via_reseed = fresh.reseed(&next_window).unwrap(); + + assert_eq!(via_fast_path.y, via_reseed.y); + assert_eq!(via_fast_path.u, via_reseed.u); + assert_eq!(via_fast_path.v, via_reseed.v); + } + + /// Temporal nlmeans HQ at `radius` with window-local noise estimation and an automatic `sigma`. + fn nlmeans_hq_windowed_plane_options(radius: u32) -> PlaneOptions { + let hq = HqParams { + windowed_noise_estimation: true, + ..HqParams::default() + }; + let hq_options = NlmeansHqOptions { + nlm: NlmeansOptions::default(), + hq, + }; + + PlaneOptions { + algorithm: Algorithm::NlmeansHq(hq_options), + ..test_plane_options(radius) + } + } + + /// Pins a VapourSynth plugin bug where HQ with an automatic `sigma` returned different pixels for + /// the same frame depending on request order. + #[test] + fn nlmeans_hq_windowed_reseed_matches_the_streaming_output_mid_clip() { + let frame_layout = layout(); + let options = nlmeans_hq_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let streamed = stream_all(&options, &frames); + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let target_frame = 8; + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + let got = denoiser.reseed(&window).unwrap(); + + assert_eq!(got.y, streamed[target_frame].y); + assert_eq!(got.u, streamed[target_frame].u); + assert_eq!(got.v, streamed[target_frame].v); + } + + #[test] + fn nlmeans_hq_windowed_reseed_matches_the_streaming_output_at_both_clip_edges() { + let frame_layout = layout(); + let options = nlmeans_hq_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let streamed = stream_all(&options, &frames); + let last = frames.len() - 1; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let last_window = window_of_span(&frames, last, span); + let got = denoiser.reseed(&last_window).unwrap(); + + assert_eq!(got.y, streamed[last].y, "luma mismatch at the ahead edge"); + assert_eq!(got.u, streamed[last].u, "u mismatch at the ahead edge"); + assert_eq!(got.v, streamed[last].v, "v mismatch at the ahead edge"); + + const BEHIND_EDGE_TOLERANCE: i32 = 8; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let first_window = window_of_span(&frames, 0, span); + let got = denoiser.reseed(&first_window).unwrap(); + let luma_diff = max_abs_diff(&got.y, &streamed[0].y); + + assert!( + luma_diff <= BEHIND_EDGE_TOLERANCE, + "luma at the behind edge (k=0) drifted too far from streaming: max abs diff {luma_diff}" + ); + } + + /// Drives one denoiser through the VapourSynth plugin harness's shuffled order with its hybrid + /// fast-path and `reseed` policy, comparing every frame with a true stream. + /// + /// A single reseed, or a reseed then one push, passes under window-local estimation without + /// exercising this. It takes a longer, repeatedly reseeded run to expose state that window-local + /// estimation fails to clear. + #[test] + fn nlmeans_hq_windowed_repeated_out_of_order_access_matches_streaming() { + let frame_layout = layout(); + let options = nlmeans_hq_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 14); + let streamed = stream_all(&options, &frames); + let last = frames.len() - 1; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let mut previous_index: Option = None; + + // The VapourSynth plugin harness's exact shuffled order. + let order = [9usize, 0, 13, 4, 5, 6, 1, 12, 2, 11, 3, 10, 7, 8]; + const NEAR_START_TOLERANCE: i32 = 8; + + for &frame_index in &order { + let fast_output = if previous_index == Some(frame_index.wrapping_sub(1)) && frame_index > 0 { + let ahead = (frame_index + span.ahead).min(last); + denoiser.push(&frames[ahead]).unwrap(); + denoiser.recv().unwrap() + } else { + None + }; + let got = match fast_output { + Some(planes) => planes, + None => { + let window = window_of_span(&frames, frame_index, span); + denoiser.reseed(&window).unwrap() + }, + }; + previous_index = Some(frame_index); + + if frame_index < span.behind { + let luma_diff = max_abs_diff(&got.y, &streamed[frame_index].y); + let u_diff = max_abs_diff(&got.u, &streamed[frame_index].u); + let v_diff = max_abs_diff(&got.v, &streamed[frame_index].v); + let max_diff = luma_diff.max(u_diff).max(v_diff); + + assert!( + max_diff <= NEAR_START_TOLERANCE, + "near-start frame n = {frame_index} drifted too far from streaming: max abs diff {max_diff}" + ); + } else { + assert_eq!( + got.y, streamed[frame_index].y, + "luma mismatch at n = {frame_index}" + ); + assert_eq!(got.u, streamed[frame_index].u, "u mismatch at n = {frame_index}"); + assert_eq!(got.v, streamed[frame_index].v, "v mismatch at n = {frame_index}"); + } + } + } + + #[test] + fn nlmeans_hq_windowed_fast_path_agrees_with_reseed_at_the_next_frame() { + let frame_layout = layout(); + let options = nlmeans_hq_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 16); + let target_frame = 8usize; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let window = window_of_span(&frames, target_frame, span); + denoiser.reseed(&window).unwrap(); + denoiser.push(&frames[target_frame + 1 + span.ahead]).unwrap(); + let via_fast_path = denoiser.recv().unwrap().expect("frame k + 1"); + + let mut fresh = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let next_window = window_of_span(&frames, target_frame + 1, span); + let via_reseed = fresh.reseed(&next_window).unwrap(); + + assert_eq!(via_fast_path.y, via_reseed.y); + assert_eq!(via_fast_path.u, via_reseed.u); + assert_eq!(via_fast_path.v, via_reseed.v); + } + + /// Renders `order` with the VapourSynth plugin harness's hybrid `render` policy. + /// + /// A frame that directly follows the previous request takes the fast `push`/`recv` path, and any + /// other frame, or one where `recv` yields nothing, goes through `reseed`. + fn render_sequence( + denoiser: &mut PlanarDenoiser, + frames: &[Planes], + span: WindowSpan, + order: &[usize], + ) -> Vec { + let last = frames.len() - 1; + let mut previous_index: Option = None; + let mut outputs = Vec::new(); + for &frame_index in order { + let fast_output = if previous_index == Some(frame_index.wrapping_sub(1)) && frame_index > 0 { + let ahead = (frame_index + span.ahead).min(last); + denoiser.push(&frames[ahead]).unwrap(); + denoiser.recv().unwrap() + } else { + None + }; + let got = match fast_output { + Some(planes) => planes, + None => { + let window = window_of_span(frames, frame_index, span); + denoiser.reseed(&window).unwrap() + }, + }; + previous_index = Some(frame_index); + outputs.push(got); + } + + outputs + } + + /// Mirrors the plugin's `a_sequential_run_after_a_seek_stays_correct_nlmeans`. + /// + /// After a reseed at frame 11 of a 14-frame clip, the two fast-path frames that follow are + /// compared with the same policy run from frame 0. That is the reference the VapourSynth harness + /// uses, rather than the true continuous stream `stream_all` produces. + #[test] + fn nlmeans_hq_windowed_sequential_run_after_a_seek_stays_correct() { + let frame_layout = layout(); + let options = nlmeans_hq_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 14); + + let mut linear = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = linear.window_span(); + let linear_order: Vec = (0..frames.len()).collect(); + let linear_out = render_sequence(&mut linear, &frames, span, &linear_order); + + let mut seeked = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let seeked_out = render_sequence(&mut seeked, &frames, span, &[11, 12, 13]); + + for (i, frame_index) in [12usize, 13].into_iter().enumerate() { + let got = &seeked_out[i + 1]; + let expected = &linear_out[frame_index]; + + assert_eq!(got.y, expected.y, "luma mismatch at n = {frame_index}"); + assert_eq!(got.u, expected.u, "u mismatch at n = {frame_index}"); + assert_eq!(got.v, expected.v, "v mismatch at n = {frame_index}"); + } + } + + /// The same check at the VapourSynth harness's clip size, 160x120. + #[test] + fn nlmeans_hq_windowed_sequential_run_after_a_seek_stays_correct_at_harness_size() { + let harness_layout = FrameLayout { + width: 160, + height: 120, + subsampling: Subsampling::Yuv420, + depth: Depth::Eight, + }; + let options = nlmeans_hq_windowed_plane_options(2); + let frames = ramp_clip(&harness_layout, 14); + + let mut linear = PlanarDenoiser::create(&options, harness_layout).unwrap(); + let span = linear.window_span(); + let linear_order: Vec = (0..frames.len()).collect(); + let linear_out = render_sequence(&mut linear, &frames, span, &linear_order); + + let mut seeked = PlanarDenoiser::create(&options, harness_layout).unwrap(); + let seeked_out = render_sequence(&mut seeked, &frames, span, &[11, 12, 13]); + + for (i, frame_index) in [12usize, 13].into_iter().enumerate() { + let got = &seeked_out[i + 1]; + let expected = &linear_out[frame_index]; + + assert_eq!(got.y, expected.y, "luma mismatch at n = {frame_index}"); + assert_eq!(got.u, expected.u, "u mismatch at n = {frame_index}"); + assert_eq!(got.v, expected.v, "v mismatch at n = {frame_index}"); + } + } + + /// Drives one denoiser through a shuffled order with the VapourSynth plugin's hybrid fast-path and + /// `reseed` policy, comparing every frame with a true stream. + /// + /// It reproduces the plugin's `random_access_matches_sequential_access_nl4d` at the core level. + /// It pins a defect where, under window-local estimation, the temporal-only noise estimator kept + /// its last trustworthy reading on folds without one, unlike every other chain. A reseed starts + /// from `reset_stream_state`, so its short run could find no reading while a true stream still + /// coasted on one from many frames back. Targets whose window covers either clip end reseed + /// through `reseed_window` with a shifted window. + #[test] + fn nl4d_windowed_repeated_out_of_order_access_matches_streaming() { + let frame_layout = layout(); + let options = nl4d_windowed_plane_options(2); + let frames = ramp_clip(&frame_layout, 14); + let streamed = stream_all(&options, &frames); + let last = frames.len() - 1; + + let mut denoiser = PlanarDenoiser::create(&options, frame_layout).unwrap(); + let span = denoiser.window_span(); + let mut previous_index: Option = None; + + // The VapourSynth plugin harness's exact shuffled order. + let order = [9usize, 0, 13, 4, 5, 6, 1, 12, 2, 11, 3, 10, 7, 8]; + + for &frame_index in &order { + let fast_output = if previous_index == Some(frame_index.wrapping_sub(1)) && frame_index > 0 { + let ahead = (frame_index + span.ahead).min(last); + denoiser.push(&frames[ahead]).unwrap(); + denoiser.recv().unwrap() + } else { + None + }; + + let at_edge = frame_index <= span.behind || frame_index + span.ahead >= last; + let got = match fast_output { + Some(planes) => planes, + None if at_edge => { + let outputs = reseed_shifted(&mut denoiser, &frames, frame_index); + outputs[0].clone() + }, + None => { + let window = window_of_span(&frames, frame_index, span); + denoiser.reseed(&window).unwrap() + }, + }; + previous_index = Some(frame_index); + + assert_eq!( + got.y, streamed[frame_index].y, + "luma mismatch at n = {frame_index}" + ); + assert_eq!(got.u, streamed[frame_index].u, "u mismatch at n = {frame_index}"); + assert_eq!(got.v, streamed[frame_index].v, "v mismatch at n = {frame_index}"); + } + } +} diff --git a/av-denoise/src/planar/tests/passthrough_retry.rs b/av-denoise/src/planar/tests/passthrough_retry.rs new file mode 100644 index 0000000..64d263e --- /dev/null +++ b/av-denoise/src/planar/tests/passthrough_retry.rs @@ -0,0 +1,69 @@ +use super::*; +use crate::accelerate::Accelerator; +use crate::{Algorithm, DenoisingMode}; + +/// Chroma-only intent, so `luma` is the disabled passthrough half and `chroma` is the one that can +/// report `QueueFull`. +fn chroma_only_options() -> PlaneOptions { + PlaneOptions { + accelerators: vec![Accelerator::Vulkan], + device: Device::Default, + intent: ChannelIntent::Chroma, + mode: DenoisingMode::Spacial, + algorithm: Algorithm::default(), + luma_strength: None, + chroma_strength: None, + luma_lambda_ht: None, + chroma_lambda_ht: None, + } +} + +fn fake_planes(layout: FrameLayout) -> Planes { + let luma_pixels = layout.luma_pixels(); + let neutral = layout.depth.neutral_chroma(); + + Planes { + y: fill_plane(luma_pixels, neutral, layout.depth), + u: layout.neutral_chroma_plane(), + v: layout.neutral_chroma_plane(), + } +} + +#[test] +fn queue_full_retry_does_not_double_queue_the_passthrough_plane() { + let layout = FrameLayout { + width: 16, + height: 16, + subsampling: Subsampling::Yuv420, + depth: Depth::Eight, + }; + let options = chroma_only_options(); + let mut denoiser = PlanarDenoiser::create(&options, layout).expect("denoiser construction failed"); + let planes = fake_planes(layout); + + // Spatial mode runs a depth-2 pipeline, so the first two pushes land directly. + denoiser.push(&planes).expect("first push should land"); + denoiser.push(&planes).expect("second push should land"); + + // The third push hits QueueFull on the chroma half. + let err = denoiser.push(&planes).expect_err("expected QueueFull"); + assert!( + matches!(err, DenoiserError::QueueFull), + "expected QueueFull, got {err:?}" + ); + + // Drain one output, then retry the whole push for the same frame. + denoiser.recv().expect("recv after drain failed"); + denoiser + .push(&planes) + .expect("retry push should land after drain"); + + // Chroma accepted three frames and `recv` popped one, so the luma passthrough queue must hold + // two and must not count the frame whose first attempt hit `QueueFull` twice. + assert_eq!( + denoiser.luma_passthrough.len(), + 2, + "expected exactly one passthrough entry per chroma frame actually accepted, got {}", + denoiser.luma_passthrough.len() + ); +} diff --git a/av-denoise/src/planar/tests/push_retry.rs b/av-denoise/src/planar/tests/push_retry.rs new file mode 100644 index 0000000..5ab61d1 --- /dev/null +++ b/av-denoise/src/planar/tests/push_retry.rs @@ -0,0 +1,26 @@ +use super::*; + +#[test] +fn ok_means_no_retry() { + let outcome = push_needs_retry(Ok(())).expect("Ok(()) must not itself error"); + assert!(!outcome, "a landed push must not ask the caller to retry"); +} + +#[test] +fn queue_full_signals_retry() { + let outcome = push_needs_retry(Err(DenoiserError::QueueFull)).expect("QueueFull must not itself error"); + assert!(outcome, "QueueFull must still trigger the retry-after-drain path"); +} + +#[test] +fn non_queue_full_errors_propagate_instead_of_being_swallowed() { + let cause = anyhow::anyhow!("synthetic readback failure"); + let synthetic = DenoiserError::Other(cause); + + let outcome = push_needs_retry(Err(synthetic)); + + assert!( + outcome.is_err(), + "a non-QueueFull push error must propagate instead of being silently treated as success" + ); +} diff --git a/av-denoise/src/stack.rs b/av-denoise/src/stack.rs new file mode 100644 index 0000000..6d89153 --- /dev/null +++ b/av-denoise/src/stack.rs @@ -0,0 +1,99 @@ +use std::ffi::OsStr; + +/// Stack bytes the kernel codegen thread needs. +pub const CODEGEN_STACK_BYTES: usize = 16 << 20; + +/// Raises `RUST_MIN_STACK` to [CODEGEN_STACK_BYTES] when it is unset. +/// +/// # Safety +/// +/// The same safety rules as [std::env::set_var] apply. +pub unsafe fn raise_codegen_stack_limit() { + if std::env::var_os("RUST_MIN_STACK").is_none() { + let stack_bytes = CODEGEN_STACK_BYTES.to_string(); + + // SAFETY: forwarded from this function's own precondition. + unsafe { std::env::set_var("RUST_MIN_STACK", stack_bytes) }; + } +} + +/// Whether the process's `RUST_MIN_STACK` is large enough for codegen. +pub fn codegen_stack_is_sufficient() -> bool { + let raw = std::env::var_os("RUST_MIN_STACK"); + limit_is_sufficient(raw.as_deref()) +} + +/// Checks a raw `RUST_MIN_STACK` value against [CODEGEN_STACK_BYTES]. +/// +/// A value that is unset or fails to parse is insufficient. +fn limit_is_sufficient(raw: Option<&OsStr>) -> bool { + raw.and_then(|value| value.to_str()) + .and_then(|text| text.parse::().ok()) + .is_some_and(|bytes| bytes >= CODEGEN_STACK_BYTES) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn unset_is_insufficient() { + assert!(!limit_is_sufficient(None)); + } + + #[test] + fn zero_is_insufficient() { + let raw = OsStr::new("0"); + let sufficient = limit_is_sufficient(Some(raw)); + + assert!(!sufficient); + } + + #[test] + fn one_below_the_limit_is_insufficient() { + let value = (CODEGEN_STACK_BYTES - 1).to_string(); + + let raw = OsStr::new(&value); + let sufficient = limit_is_sufficient(Some(raw)); + + assert!(!sufficient); + } + + #[test] + fn exactly_the_limit_is_sufficient() { + let value = CODEGEN_STACK_BYTES.to_string(); + + let raw = OsStr::new(&value); + let sufficient = limit_is_sufficient(Some(raw)); + + assert!(sufficient); + } + + #[test] + fn above_the_limit_is_sufficient() { + let value = (CODEGEN_STACK_BYTES * 2).to_string(); + + let raw = OsStr::new(&value); + let sufficient = limit_is_sufficient(Some(raw)); + + assert!(sufficient); + } + + #[test] + fn surrounding_whitespace_is_insufficient() { + let value = format!(" {CODEGEN_STACK_BYTES} "); + + let raw = OsStr::new(&value); + let sufficient = limit_is_sufficient(Some(raw)); + + assert!(!sufficient); + } + + #[test] + fn non_numeric_is_insufficient() { + let raw = OsStr::new("not-a-number"); + let sufficient = limit_is_sufficient(Some(raw)); + + assert!(!sufficient); + } +} diff --git a/av-denoise/src/warmup.rs b/av-denoise/src/warmup.rs new file mode 100644 index 0000000..ae4eea8 --- /dev/null +++ b/av-denoise/src/warmup.rs @@ -0,0 +1,402 @@ +//! Lets one process fill a cold kernel cache while the others wait +//! +//! CubeCL compiles a kernel the first time it is dispatched and writes it to the cache that +//! [install_compilation_cache](crate::install_compilation_cache) points it at. The cache's table of +//! contents is a snapshot taken when the GPU client is built, so a process that starts while another +//! is still compiling shares nothing with it, pays the full compilation cost again and writes its +//! own copy of the same kernels. +//! +//! [WarmUp] holds a lock file next to the cache while the first process compiles, and the rest +//! block on it, so each waiting process builds its own client after the wait, when the cache +//! already holds every kernel. A finished run leaves a stamp file behind, so later processes skip +//! the lock entirely. + +use std::collections::HashSet; +use std::fs::{File, OpenOptions, TryLockError}; +use std::path::{Path, PathBuf}; +use std::sync::{LazyLock, Mutex}; +use std::time::{Duration, Instant}; + +use cubecl::hash::StableHasher; + +use crate::cache::compilation_cache_dir; +use crate::planar::{FrameLayout, PlaneOptions}; + +/// How long a process waits for the one ahead of it before compiling for itself. +/// +/// Compiling the kernels was measured at about ten seconds on a quiet machine, and a machine running +/// an encode is not quiet, so the limit is generous. Waiting longer is worse than duplicating the +/// work, because the encoder has nothing to do until a frame arrives. +const WAIT_LIMIT: Duration = Duration::from_secs(180); + +const POLL_INTERVAL: Duration = Duration::from_millis(100); + +/// The keys this process already holds a place for. +static CLAIMED_KEYS: LazyLock>> = LazyLock::new(|| Mutex::new(HashSet::new())); + +/// Records that this process is taking the place for `key`, and reports whether it was free. +fn claim_key(key: u128) -> bool { + CLAIMED_KEYS + .lock() + .expect("warm-up key mutex poisoned") + .insert(key) +} + +/// Gives `key` back, so a later filter in this process can queue for it again. +fn release_key(key: u128) { + CLAIMED_KEYS + .lock() + .expect("warm-up key mutex poisoned") + .remove(&key); +} + +/// Identifies the set of kernels a denoiser compiles. +/// +/// Radii, channel mode, depth and algorithm are baked into the kernels at compile time, so processes +/// compiling different sets never wait for each other and a stamp for one set never vouches for another. +/// +/// The key hashes the `Debug` rendering of both inputs rather than a hand-written field list, which +/// would silently stop covering a newly added field and let a process trust a stamp for kernels it +/// never compiled. Fields that only reach the GPU at runtime, such as strength, make the key finer +/// than needed, which costs an extra warm-up and is the safe direction to err in. +/// +/// The crate version is part of the key because a release can upgrade CubeCL, which files its cache +/// under the CubeCL version, and a stamp from before the upgrade would call the emptied cache warm. +/// A rebuild of the kernel sources needs nothing here, because the stamps live in that build's own +/// cache directory. +pub fn kernel_key(options: &PlaneOptions, layout: FrameLayout) -> u128 { + let version = env!("CARGO_PKG_VERSION"); + let rendered = format!("{version}|{options:?}|{layout:?}"); + StableHasher::hash_one(&rendered) +} + +/// A held place in the queue to fill a cold cache. +/// +/// Obtained from [WarmUp::begin] and given up with [WarmUp::finish] once the kernels are compiled. +/// Dropping one without calling `finish` releases the lock without leaving a stamp, so a run that +/// failed part way through does not convince the next process that the cache is warm. +#[derive(Debug)] +pub struct WarmUp { + lock: File, + stamp: PathBuf, + key: u128, +} + +impl WarmUp { + /// Takes a place in the queue for the kernels `key` identifies. + /// + /// Returns `Some` while holding the lock, and the caller compiles under it. Returns `None` when + /// the cache is already warm for these kernels, caching is off, or the lock could not be taken + /// within the three minute wait limit. In every one of those the caller compiles as usual. + /// + /// Blocks for as long as the process ahead takes to compile, so it belongs on the path that + /// builds a denoiser rather than the path that renders a frame. + pub fn begin(key: u128) -> Option { + let dir = compilation_cache_dir()?; + Self::begin_in(dir, key, WAIT_LIMIT) + } + + /// [WarmUp::begin] against an explicit directory and wait limit. + fn begin_in(dir: &Path, key: u128, wait_limit: Duration) -> Option { + // A file lock is held by the process rather than the handle, so a second filter in this + // process asking for the same kernels would wait out `wait_limit` on its own lock. One script + // can easily build two filters, so the place is taken at most once per key per process and + // the second caller carries on. + if !claim_key(key) { + return None; + } + + let held = Self::acquire(dir, key, wait_limit); + + if held.is_none() { + release_key(key); + } + + held + } + + /// [WarmUp::begin_in] without the in-process bookkeeping, which its caller handles. + fn acquire(dir: &Path, key: u128, wait_limit: Duration) -> Option { + let stamp_name = format!("warm-{key:032x}.stamp"); + let stamp = dir.join(stamp_name); + + if stamp.exists() { + return None; + } + + let lock_name = format!("warm-{key:032x}.lock"); + let lock_path = dir.join(lock_name); + let lock = open_lock_file(&lock_path)?; + + if !wait_for_lock(&lock, wait_limit) { + return None; + } + + // Whoever held the lock has finished compiling by now, and checking the stamp again is what + // turns the queue into a single warm-up rather than one per waiting process. + if stamp.exists() { + let _ = lock.unlock(); + return None; + } + + tracing::debug!(?stamp, "compiling kernels for a cold cache"); + Some(Self { lock, stamp, key }) + } + + /// Records that the kernels are compiled and lets the next process through. + pub fn finish(self) { + if let Err(err) = std::fs::write(&self.stamp, b"") { + // A missing stamp reads as a cold cache, which is slower but still correct. + tracing::debug!(stamp = ?self.stamp, %err, "cannot write the kernel warm-up stamp"); + } + } +} + +impl Drop for WarmUp { + fn drop(&mut self) { + let _ = self.lock.unlock(); + release_key(self.key); + } +} + +/// Opens the lock file, creating it when this is the first process to ask for these kernels. +/// +/// A directory that cannot be written is logged and then ignored, because denoising works without +/// the queue and only compiles more than once. +fn open_lock_file(path: &Path) -> Option { + let opened = OpenOptions::new() + .create(true) + .read(true) + .write(true) + .truncate(false) + .open(path); + + match opened { + Ok(file) => Some(file), + Err(err) => { + tracing::debug!(?path, %err, "cannot open the kernel warm-up lock, compiling unqueued"); + None + }, + } +} + +/// Blocks until the lock is held, giving up after `wait_limit`. +/// +/// Returns whether the lock is held. It is a real advisory file lock rather than a file whose presence +/// means "taken", so the operating system releases it when Av1an kills a worker mid-compile and the +/// next process in line wakes up straight away. +fn wait_for_lock(lock: &File, wait_limit: Duration) -> bool { + let start = Instant::now(); + + loop { + match lock.try_lock() { + Ok(()) => return true, + Err(TryLockError::WouldBlock) => {}, + Err(TryLockError::Error(err)) => { + tracing::debug!(%err, "cannot take the kernel warm-up lock, compiling unqueued"); + return false; + }, + } + + if start.elapsed() >= wait_limit { + tracing::warn!( + "waited {:?} for another process to compile kernels, compiling for ourselves", + wait_limit, + ); + return false; + } + + std::thread::sleep(POLL_INTERVAL); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + /// A key no other test shares, so one test's lock and stamp files are never seen by the next. + fn key(index: u128) -> u128 { + 0xa5a5_0000_0000_0000_0000_0000_0000_0000 + index + } + + /// Short enough that a contended lock fails the test quickly rather than after three minutes. + const BRIEFLY: Duration = Duration::from_millis(50); + + /// The lock is taken directly because a file lock is held per process, and a second `begin_in` + /// would be answered by the in-process registry instead of the lock. + #[test] + fn a_lock_held_elsewhere_keeps_this_process_out() { + let dir = tempfile::tempdir().unwrap(); + let lock_key = key(1); + let lock_name = format!("warm-{lock_key:032x}.lock"); + let path = dir.path().join(lock_name); + let elsewhere = open_lock_file(&path).unwrap(); + elsewhere.lock().unwrap(); + + let place = WarmUp::begin_in(dir.path(), lock_key, BRIEFLY); + + assert!( + place.is_none(), + "a caller gives up rather than compiling alongside the process ahead", + ); + } + + #[test] + fn one_process_takes_one_place_per_key() { + let dir = tempfile::tempdir().unwrap(); + let shared_key = key(6); + + let first = WarmUp::begin_in(dir.path(), shared_key, BRIEFLY); + assert!(first.is_some(), "the first caller compiles"); + + let second = WarmUp::begin_in(dir.path(), shared_key, BRIEFLY); + assert!( + second.is_none(), + "the second caller carries on rather than waiting for itself", + ); + } + + #[test] + fn a_released_place_can_be_taken_again() { + let dir = tempfile::tempdir().unwrap(); + let place_key = key(7); + + let released = WarmUp::begin_in(dir.path(), place_key, BRIEFLY); + drop(released); + + let retaken = WarmUp::begin_in(dir.path(), place_key, BRIEFLY); + assert!( + retaken.is_some(), + "the key is free again once the place is given up", + ); + } + + #[test] + fn a_finished_warm_up_lets_the_next_process_straight_through() { + let dir = tempfile::tempdir().unwrap(); + let place_key = key(2); + + WarmUp::begin_in(dir.path(), place_key, BRIEFLY).unwrap().finish(); + + let next = WarmUp::begin_in(dir.path(), place_key, BRIEFLY); + assert!(next.is_none(), "a warm cache needs no queue"); + } + + #[test] + fn an_abandoned_warm_up_leaves_the_cache_cold() { + let dir = tempfile::tempdir().unwrap(); + let place_key = key(3); + + let abandoned = WarmUp::begin_in(dir.path(), place_key, BRIEFLY); + drop(abandoned); + + let next = WarmUp::begin_in(dir.path(), place_key, BRIEFLY); + assert!(next.is_some(), "no stamp means the kernels still need compiling"); + } + + /// A `PlaneOptions` with no accelerator named, so this module builds with any backend feature. + fn options() -> PlaneOptions { + PlaneOptions { + accelerators: Vec::new(), + device: crate::Device::Default, + intent: crate::ChannelIntent::LumaChroma, + mode: crate::DenoisingMode::Temporal { radius: 2 }, + algorithm: crate::Algorithm::default(), + luma_strength: None, + chroma_strength: None, + luma_lambda_ht: None, + chroma_lambda_ht: None, + } + } + + fn layout() -> FrameLayout { + FrameLayout { + width: 1920, + height: 1080, + subsampling: crate::Subsampling::Yuv420, + depth: crate::Depth::Eight, + } + } + + #[test] + fn the_same_settings_give_the_same_key() { + let first_options = options(); + let first_layout = layout(); + let first = kernel_key(&first_options, first_layout); + + let second_options = options(); + let second_layout = layout(); + let second = kernel_key(&second_options, second_layout); + + assert_eq!(first, second); + } + + #[test] + fn a_different_depth_gives_a_different_key() { + let ten_bit = FrameLayout { + depth: crate::Depth::Ten, + ..layout() + }; + + let test_options = options(); + let eight_bit = layout(); + let eight_bit_key = kernel_key(&test_options, eight_bit); + let ten_bit_key = kernel_key(&test_options, ten_bit); + + assert_ne!(eight_bit_key, ten_bit_key); + } + + #[test] + fn a_different_radius_gives_a_different_key() { + let wider = PlaneOptions { + mode: crate::DenoisingMode::Temporal { radius: 3 }, + ..options() + }; + + let default_options = options(); + let test_layout = layout(); + let default_key = kernel_key(&default_options, test_layout); + let wider_key = kernel_key(&wider, test_layout); + + assert_ne!(default_key, wider_key); + } + + #[test] + fn grain_export_gives_a_different_key() { + let exporting_options = crate::Nl4dOptions { + grain_export: true, + ..crate::Nl4dOptions::default() + }; + let exporting = PlaneOptions { + algorithm: crate::Algorithm::Nl4d(exporting_options), + ..options() + }; + + let plain_options = crate::Nl4dOptions::default(); + let plain = PlaneOptions { + algorithm: crate::Algorithm::Nl4d(plain_options), + ..options() + }; + + let test_layout = layout(); + let plain_key = kernel_key(&plain, test_layout); + let exporting_key = kernel_key(&exporting, test_layout); + + assert_ne!(plain_key, exporting_key); + } + + #[test] + fn different_kernels_do_not_wait_for_each_other() { + let dir = tempfile::tempdir().unwrap(); + let first_key = key(4); + let second_key = key(5); + + let first = WarmUp::begin_in(dir.path(), first_key, BRIEFLY); + let second = WarmUp::begin_in(dir.path(), second_key, BRIEFLY); + + assert!( + first.is_some() && second.is_some(), + "separate keys queue separately" + ); + } +} diff --git a/av-denoise/tests/fixtures/parity.txt b/av-denoise/tests/fixtures/parity.txt new file mode 100644 index 0000000..0ecad29 --- /dev/null +++ b/av-denoise/tests/fixtures/parity.txt @@ -0,0 +1,383 @@ +# adapters: AMD Radeon AI PRO R9700 (RADV GFX1201)|radv|Mesa 26.2.2-arch3.2; AMD Radeon AI PRO R9700 (RADV GFX1201)|radv|Mesa 26.2.2-arch3.2; AMD Ryzen 9 9950X 16-Core Processor (RADV RAPHAEL_MENDOCINO)|radv|Mesa 26.2.2-arch3.2 +nl4d_chroma_10 f0.u 161f5b580e5064dd +nl4d_chroma_10 f0.v adb1904c0b9328ad +nl4d_chroma_10 f0.y 11cc4683a0611f99 +nl4d_chroma_10 f1.u b6bc81c1fd1b76c6 +nl4d_chroma_10 f1.v 81fd9ae313e869bf +nl4d_chroma_10 f1.y 6314ed1c36536d09 +nl4d_chroma_10 f10.u d639b63f54ba446b +nl4d_chroma_10 f10.v c6eef6f8cdc530fa +nl4d_chroma_10 f10.y f4dfaad0a380a7e3 +nl4d_chroma_10 f11.u 1e4e184845cec46b +nl4d_chroma_10 f11.v b777d39c49928832 +nl4d_chroma_10 f11.y 060b898058902cdb +nl4d_chroma_10 f2.u 892e38df5acf117c +nl4d_chroma_10 f2.v 48bfe26b7067e1c6 +nl4d_chroma_10 f2.y 586dac4893bb0b5e +nl4d_chroma_10 f3.u 7d6b84fca3d88fc3 +nl4d_chroma_10 f3.v f3648a7728998344 +nl4d_chroma_10 f3.y 8fcba87f5350cf5f +nl4d_chroma_10 f4.u 5ac068c597a36295 +nl4d_chroma_10 f4.v 3b0a83f09e65ea97 +nl4d_chroma_10 f4.y f41ce69ba683fe34 +nl4d_chroma_10 f5.u 80276fb5dffbac0d +nl4d_chroma_10 f5.v f11ed94bd284009c +nl4d_chroma_10 f5.y b33c4217a2a5567c +nl4d_chroma_10 f6.u 2809beb330d8bd0b +nl4d_chroma_10 f6.v 5308179884af107d +nl4d_chroma_10 f6.y f95163970c4a3e42 +nl4d_chroma_10 f7.u 759ecf183a5e888b +nl4d_chroma_10 f7.v 9de3eb0b738acb1e +nl4d_chroma_10 f7.y e6e619a8fbfc4f74 +nl4d_chroma_10 f8.u df5b506de4bae36d +nl4d_chroma_10 f8.v 22c6d1c96831b818 +nl4d_chroma_10 f8.y df7b3b736ff10a2e +nl4d_chroma_10 f9.u 628fd606b926170e +nl4d_chroma_10 f9.v 83cb1401d0dbf987 +nl4d_chroma_10 f9.y 8a2dac912021485b +nl4d_grain_8 f0.u de3a29958a0cfdee +nl4d_grain_8 f0.v e06919110df6aaf8 +nl4d_grain_8 f0.y 3a91ab99fe317322 +nl4d_grain_8 f1.u 56c2b42b5e3ba76e +nl4d_grain_8 f1.v 21bc3e02b9454632 +nl4d_grain_8 f1.y c0be61fd7efa1a16 +nl4d_grain_8 f10.u c35b9c01d13ff06f +nl4d_grain_8 f10.v 00154a8154a2a331 +nl4d_grain_8 f10.y ff92d6ba20092197 +nl4d_grain_8 f11.u 63ad2397626689c9 +nl4d_grain_8 f11.v 9fe1a2d3acdf7357 +nl4d_grain_8 f11.y 07aa59f519bf235b +nl4d_grain_8 f2.u aa48d2b79a5ac418 +nl4d_grain_8 f2.v 153ee290cdf20dfb +nl4d_grain_8 f2.y cd361ce5085a5753 +nl4d_grain_8 f3.u 7d28bce3b73ed591 +nl4d_grain_8 f3.v e7e89de054de038a +nl4d_grain_8 f3.y 7fd336e1f6483583 +nl4d_grain_8 f4.u 7c2eb1ec251768d4 +nl4d_grain_8 f4.v f1fcc4442b36dc57 +nl4d_grain_8 f4.y 46b294be5f075502 +nl4d_grain_8 f5.u a5bcebb6b3e397ed +nl4d_grain_8 f5.v f85f6833a2f52f67 +nl4d_grain_8 f5.y 3e83b20dd67d84d2 +nl4d_grain_8 f6.u fccd90249eca907f +nl4d_grain_8 f6.v 1a7358cf1f6b0497 +nl4d_grain_8 f6.y 5c79d91832a4f020 +nl4d_grain_8 f7.u 99191803fddb2ffc +nl4d_grain_8 f7.v 2c253716e5210740 +nl4d_grain_8 f7.y 9a175b94728f8631 +nl4d_grain_8 f8.u 76720368694b9dda +nl4d_grain_8 f8.v d1fc3836b5ad3c64 +nl4d_grain_8 f8.y 96ab849789b00c0e +nl4d_grain_8 f9.u 31699de899f8a1f6 +nl4d_grain_8 f9.v a3080004132f2010 +nl4d_grain_8 f9.y dec85fdcf53748e6 +nl4d_grain_8 grain cfe6c57268d13d4f +nl4d_luma_8 f0.u de3a29958a0cfdee +nl4d_luma_8 f0.v e06919110df6aaf8 +nl4d_luma_8 f0.y 3a91ab99fe317322 +nl4d_luma_8 f1.u 56c2b42b5e3ba76e +nl4d_luma_8 f1.v 21bc3e02b9454632 +nl4d_luma_8 f1.y c0be61fd7efa1a16 +nl4d_luma_8 f10.u c35b9c01d13ff06f +nl4d_luma_8 f10.v 00154a8154a2a331 +nl4d_luma_8 f10.y ff92d6ba20092197 +nl4d_luma_8 f11.u 63ad2397626689c9 +nl4d_luma_8 f11.v 9fe1a2d3acdf7357 +nl4d_luma_8 f11.y 07aa59f519bf235b +nl4d_luma_8 f2.u aa48d2b79a5ac418 +nl4d_luma_8 f2.v 153ee290cdf20dfb +nl4d_luma_8 f2.y cd361ce5085a5753 +nl4d_luma_8 f3.u 7d28bce3b73ed591 +nl4d_luma_8 f3.v e7e89de054de038a +nl4d_luma_8 f3.y 7fd336e1f6483583 +nl4d_luma_8 f4.u 7c2eb1ec251768d4 +nl4d_luma_8 f4.v f1fcc4442b36dc57 +nl4d_luma_8 f4.y 46b294be5f075502 +nl4d_luma_8 f5.u a5bcebb6b3e397ed +nl4d_luma_8 f5.v f85f6833a2f52f67 +nl4d_luma_8 f5.y 3e83b20dd67d84d2 +nl4d_luma_8 f6.u fccd90249eca907f +nl4d_luma_8 f6.v 1a7358cf1f6b0497 +nl4d_luma_8 f6.y 5c79d91832a4f020 +nl4d_luma_8 f7.u 99191803fddb2ffc +nl4d_luma_8 f7.v 2c253716e5210740 +nl4d_luma_8 f7.y 9a175b94728f8631 +nl4d_luma_8 f8.u 76720368694b9dda +nl4d_luma_8 f8.v d1fc3836b5ad3c64 +nl4d_luma_8 f8.y 96ab849789b00c0e +nl4d_luma_8 f9.u 31699de899f8a1f6 +nl4d_luma_8 f9.v a3080004132f2010 +nl4d_luma_8 f9.y dec85fdcf53748e6 +nl4d_lumachroma_10 f0.u 161f5b580e5064dd +nl4d_lumachroma_10 f0.v adb1904c0b9328ad +nl4d_lumachroma_10 f0.y 85405eca45bceebf +nl4d_lumachroma_10 f1.u b6bc81c1fd1b76c6 +nl4d_lumachroma_10 f1.v 81fd9ae313e869bf +nl4d_lumachroma_10 f1.y dced2157e6b329cb +nl4d_lumachroma_10 f10.u d639b63f54ba446b +nl4d_lumachroma_10 f10.v c6eef6f8cdc530fa +nl4d_lumachroma_10 f10.y 494b16347a0ae43e +nl4d_lumachroma_10 f11.u 1e4e184845cec46b +nl4d_lumachroma_10 f11.v b777d39c49928832 +nl4d_lumachroma_10 f11.y 4c5e34cf02d74d32 +nl4d_lumachroma_10 f2.u 892e38df5acf117c +nl4d_lumachroma_10 f2.v 48bfe26b7067e1c6 +nl4d_lumachroma_10 f2.y 1fb402cbd5b25f54 +nl4d_lumachroma_10 f3.u 7d6b84fca3d88fc3 +nl4d_lumachroma_10 f3.v f3648a7728998344 +nl4d_lumachroma_10 f3.y 31e4694b95ebd8df +nl4d_lumachroma_10 f4.u 5ac068c597a36295 +nl4d_lumachroma_10 f4.v 3b0a83f09e65ea97 +nl4d_lumachroma_10 f4.y 6c0cb4387420c742 +nl4d_lumachroma_10 f5.u 80276fb5dffbac0d +nl4d_lumachroma_10 f5.v f11ed94bd284009c +nl4d_lumachroma_10 f5.y 99e5ee0e991e62a3 +nl4d_lumachroma_10 f6.u 2809beb330d8bd0b +nl4d_lumachroma_10 f6.v 5308179884af107d +nl4d_lumachroma_10 f6.y 43e53e07d2e1e767 +nl4d_lumachroma_10 f7.u 759ecf183a5e888b +nl4d_lumachroma_10 f7.v 9de3eb0b738acb1e +nl4d_lumachroma_10 f7.y a06857772a909dd9 +nl4d_lumachroma_10 f8.u df5b506de4bae36d +nl4d_lumachroma_10 f8.v 22c6d1c96831b818 +nl4d_lumachroma_10 f8.y 4d90b3731f3d012e +nl4d_lumachroma_10 f9.u 628fd606b926170e +nl4d_lumachroma_10 f9.v 83cb1401d0dbf987 +nl4d_lumachroma_10 f9.y 76f67fb89ac73f9e +nl4d_reseed_8 f0.u a7a8985a94219195 +nl4d_reseed_8 f0.v a6ba3cab66a1e14e +nl4d_reseed_8 f0.y 8bbd866b5e366bde +nl4d_reseed_8 f1.u 490aab2f8de9432f +nl4d_reseed_8 f1.v 388329719d64308b +nl4d_reseed_8 f1.y 8dbdced79806242c +nl4d_reseed_8 f10.u 3c544d4c5c5e0466 +nl4d_reseed_8 f10.v b884225f1f4f4014 +nl4d_reseed_8 f10.y 83c019c921163a96 +nl4d_reseed_8 f11.u 34f0a9dbdfb2dd65 +nl4d_reseed_8 f11.v dd98fa5a0490bd91 +nl4d_reseed_8 f11.y 40edf847a8356b20 +nl4d_reseed_8 f2.u 96814423f16645e0 +nl4d_reseed_8 f2.v 028ca41cf77399ae +nl4d_reseed_8 f2.y 78c5dcd96ad51433 +nl4d_reseed_8 f3.u 5f47139ddfa8fd51 +nl4d_reseed_8 f3.v 71a028f4ba8242e6 +nl4d_reseed_8 f3.y 4e54d1c85b8a496f +nl4d_reseed_8 f4.u 2cc80c12eef33583 +nl4d_reseed_8 f4.v 8a38e8cae02f28be +nl4d_reseed_8 f4.y a6a418d5ec33631a +nl4d_reseed_8 f5.u fc70d7ff20bb8613 +nl4d_reseed_8 f5.v 59a431607b2d5a3b +nl4d_reseed_8 f5.y e6d031ad55009594 +nl4d_reseed_8 f6.u 1ec35b96284e672d +nl4d_reseed_8 f6.v ef76521cc139512d +nl4d_reseed_8 f6.y 68775ea2469988bb +nl4d_reseed_8 f7.u 53c46090ad950f0e +nl4d_reseed_8 f7.v 2d482d3d3ce38c75 +nl4d_reseed_8 f7.y a3e1b13d1e8932d3 +nl4d_reseed_8 f8.u 30720e9c31605ac8 +nl4d_reseed_8 f8.v 33d60bd12ea0fcc1 +nl4d_reseed_8 f8.y c848ba40cc59342f +nl4d_reseed_8 f9.u e4e83c24a5c1ed93 +nl4d_reseed_8 f9.v 464e4fc8b9db8559 +nl4d_reseed_8 f9.y 1f72af395684b2d2 +nl4d_reseed_8 reseed.u fc70d7ff20bb8613 +nl4d_reseed_8 reseed.v 59a431607b2d5a3b +nl4d_reseed_8 reseed.y e6d031ad55009594 +nl4d_short_8 f0.u de3a29958a0cfdee +nl4d_short_8 f0.v e06919110df6aaf8 +nl4d_short_8 f0.y a88df6ca31fce47d +nl4d_short_8 f1.u 56c2b42b5e3ba76e +nl4d_short_8 f1.v 21bc3e02b9454632 +nl4d_short_8 f1.y 86ce99c97794c164 +nl4d_short_8 f2.u aa48d2b79a5ac418 +nl4d_short_8 f2.v 153ee290cdf20dfb +nl4d_short_8 f2.y 144fbe62f6663f5e +nl4d_yuv_8 f0.u b877cf9c924be2ee +nl4d_yuv_8 f0.v 97fcc69c38d153d8 +nl4d_yuv_8 f0.y e034fe39b18571a6 +nl4d_yuv_8 f1.u 84c2e194f3f06fd6 +nl4d_yuv_8 f1.v 673319d76e8e6866 +nl4d_yuv_8 f1.y e17a75dc27502ff9 +nl4d_yuv_8 f10.u ef4b41583f16a75f +nl4d_yuv_8 f10.v 5c69fbf8c61a30e4 +nl4d_yuv_8 f10.y 0fca5ca8a9d95cd8 +nl4d_yuv_8 f11.u 93d9f8d0d17cfa21 +nl4d_yuv_8 f11.v 9238766e56b03100 +nl4d_yuv_8 f11.y 4b6062220072ba83 +nl4d_yuv_8 f2.u f4c0e5cbf398ba99 +nl4d_yuv_8 f2.v 4d502d607e8e82e1 +nl4d_yuv_8 f2.y 547a96d5f4a4b5e0 +nl4d_yuv_8 f3.u 1394c61ca4f39cbb +nl4d_yuv_8 f3.v 3ddcf4cfbf8aeb64 +nl4d_yuv_8 f3.y 9dd1c7a91a4b3873 +nl4d_yuv_8 f4.u b905ec5d1f3a2bf7 +nl4d_yuv_8 f4.v 60a5e4facd97c9ea +nl4d_yuv_8 f4.y c16d07d842105902 +nl4d_yuv_8 f5.u 35ff99dbf3619bd2 +nl4d_yuv_8 f5.v ec6f050fe64921f4 +nl4d_yuv_8 f5.y 7131076a689de95e +nl4d_yuv_8 f6.u e789300d31ff8d6e +nl4d_yuv_8 f6.v f3c045b93eae731a +nl4d_yuv_8 f6.y 9a7aa069be7e901f +nl4d_yuv_8 f7.u ee2d1bd9f3929862 +nl4d_yuv_8 f7.v e24246e3ee863c98 +nl4d_yuv_8 f7.y 04690bbd76f45ea6 +nl4d_yuv_8 f8.u 0d884bda5d1f3439 +nl4d_yuv_8 f8.v bac2240ca0518c14 +nl4d_yuv_8 f8.y a3cea1562490cac4 +nl4d_yuv_8 f9.u 0fa9087726a56dcb +nl4d_yuv_8 f9.v 10fc1181945814c1 +nl4d_yuv_8 f9.y 210fe7c99806f7e6 +nlm_fast_spatial_8 f0.u 4617b7de6002d675 +nlm_fast_spatial_8 f0.v b6a491ada7b1c0d4 +nlm_fast_spatial_8 f0.y 5576512fe4885820 +nlm_fast_spatial_8 f1.u 724f90b6044216b7 +nlm_fast_spatial_8 f1.v 0ba069a1487ec38f +nlm_fast_spatial_8 f1.y b1c39b8fdfc8e1f5 +nlm_fast_spatial_8 f2.u fc8cfe52fedd0875 +nlm_fast_spatial_8 f2.v e76290349add76e5 +nlm_fast_spatial_8 f2.y d9276d4dd9ca3760 +nlm_fast_spatial_8 f3.u 1718a5d78eefa2fc +nlm_fast_spatial_8 f3.v 66584f27ad6aff9d +nlm_fast_spatial_8 f3.y 6426a45c9448164a +nlm_fast_spatial_8 f4.u 33f6baef6090221f +nlm_fast_spatial_8 f4.v d60e7b685b2ffe29 +nlm_fast_spatial_8 f4.y cfffef291d7fc390 +nlm_fast_spatial_8 f5.u d88118a851062faf +nlm_fast_spatial_8 f5.v 0b3d22b81b2f57c4 +nlm_fast_spatial_8 f5.y 14e4cc9c282972d7 +nlm_fast_temporal_8 f0.u de3a29958a0cfdee +nlm_fast_temporal_8 f0.v e06919110df6aaf8 +nlm_fast_temporal_8 f0.y fcb26120eb5d94db +nlm_fast_temporal_8 f1.u fd62c278ce48fa70 +nlm_fast_temporal_8 f1.v c18c4b057d07f972 +nlm_fast_temporal_8 f1.y d779ac64aa959b10 +nlm_fast_temporal_8 f2.u 90d97a80ce199581 +nlm_fast_temporal_8 f2.v 7a4d6aa2b4d2ea6d +nlm_fast_temporal_8 f2.y f8b56f929264b910 +nlm_fast_temporal_8 f3.u eada4ac4effff9d1 +nlm_fast_temporal_8 f3.v 32de1fc4111b0a6f +nlm_fast_temporal_8 f3.y 6753f1fb81a84ae3 +nlm_fast_temporal_8 f4.u a75076e37736553d +nlm_fast_temporal_8 f4.v 0069ec52ab66fd3c +nlm_fast_temporal_8 f4.y b8c1a06c065769f7 +nlm_fast_temporal_8 f5.u 352d8ab8c3cd0f9d +nlm_fast_temporal_8 f5.v e6ab0ab322a449ad +nlm_fast_temporal_8 f5.y 25007fc91c11d289 +nlm_fast_temporal_8 f6.u d0aa912deafc17ae +nlm_fast_temporal_8 f6.v 23b1bb13fa24c92f +nlm_fast_temporal_8 f6.y 32978b4b98581379 +nlm_fast_temporal_8 f7.u 94cd33fec9855e38 +nlm_fast_temporal_8 f7.v a703b87e3f207542 +nlm_fast_temporal_8 f7.y 80ece6d9cdb674de +nlm_fast_temporal_8 f8.u c3d29e4b003342a6 +nlm_fast_temporal_8 f8.v 80f8d3d551b5624a +nlm_fast_temporal_8 f8.y a0cdacda49786ee5 +nlm_fast_temporal_8 f9.u 31699de899f8a1f6 +nlm_fast_temporal_8 f9.v a3080004132f2010 +nlm_fast_temporal_8 f9.y c90fd0fd2a204614 +nlm_hq_reseed_8 f0.u 04753832ff160c19 +nlm_hq_reseed_8 f0.v c4cdcd5f252d9180 +nlm_hq_reseed_8 f0.y c199cdf4b6d63943 +nlm_hq_reseed_8 f1.u c1f7b4e11b04975b +nlm_hq_reseed_8 f1.v e47d574c6ea6337e +nlm_hq_reseed_8 f1.y 60a4598b837cdda0 +nlm_hq_reseed_8 f10.u cf9ab4585f69d1e6 +nlm_hq_reseed_8 f10.v cc72471bc094d40a +nlm_hq_reseed_8 f10.y c3e01289efd5ee85 +nlm_hq_reseed_8 f11.u 9fea58fe56eedb22 +nlm_hq_reseed_8 f11.v 048f6872f3a8139b +nlm_hq_reseed_8 f11.y 5bcef01b2a68b941 +nlm_hq_reseed_8 f2.u a111b54fa94204ad +nlm_hq_reseed_8 f2.v ff5a79bea66348c8 +nlm_hq_reseed_8 f2.y 3857fa203c1b56d5 +nlm_hq_reseed_8 f3.u 048567c2e213b3ae +nlm_hq_reseed_8 f3.v 3cff0cf2af268acf +nlm_hq_reseed_8 f3.y f24929570c0a5c42 +nlm_hq_reseed_8 f4.u 4e6e862eec5e45f1 +nlm_hq_reseed_8 f4.v a06e57a303351550 +nlm_hq_reseed_8 f4.y 05aa09ffc7de5ac2 +nlm_hq_reseed_8 f5.u 57b8a591bde2dda2 +nlm_hq_reseed_8 f5.v 98012faf98f1ea85 +nlm_hq_reseed_8 f5.y e9db1c6a8f13683f +nlm_hq_reseed_8 f6.u 1c8209e74c683edb +nlm_hq_reseed_8 f6.v 0da09d2d0446d498 +nlm_hq_reseed_8 f6.y 1c7b1711d7da120c +nlm_hq_reseed_8 f7.u 38a2d4331d256efc +nlm_hq_reseed_8 f7.v 92a589cd380f3385 +nlm_hq_reseed_8 f7.y 7614ae699a520d71 +nlm_hq_reseed_8 f8.u 3949798131611062 +nlm_hq_reseed_8 f8.v ce53158d3f27eee1 +nlm_hq_reseed_8 f8.y bda3e0c7f5ba36e6 +nlm_hq_reseed_8 f9.u be730d32af5ce9b9 +nlm_hq_reseed_8 f9.v 060f4becf320d8b5 +nlm_hq_reseed_8 f9.y 2932a9f48ef1da7a +nlm_hq_reseed_8 reseed.u ee5a033eff1cec62 +nlm_hq_reseed_8 reseed.v 091466ca4c54dd44 +nlm_hq_reseed_8 reseed.y c1191ab5bcb8b59a +nlm_hq_short_8 f0.u de3a29958a0cfdee +nlm_hq_short_8 f0.v e06919110df6aaf8 +nlm_hq_short_8 f0.y 18173befc5c8a643 +nlm_hq_short_8 f1.u 56c2b42b5e3ba76e +nlm_hq_short_8 f1.v 21bc3e02b9454632 +nlm_hq_short_8 f1.y 92fca4a0c774d537 +nlm_hq_temporal_10 f0.u 32d98959d8b47ef9 +nlm_hq_temporal_10 f0.v efe1ac26c4b48580 +nlm_hq_temporal_10 f0.y 403175fad9ca03bc +nlm_hq_temporal_10 f1.u 02541d255613ac65 +nlm_hq_temporal_10 f1.v 3535b70313d44148 +nlm_hq_temporal_10 f1.y 5e5b049b3aa70e9b +nlm_hq_temporal_10 f2.u 68a2fa57eef9993f +nlm_hq_temporal_10 f2.v 4fd30237b406f7c6 +nlm_hq_temporal_10 f2.y 5f884426f82d6c3d +nlm_hq_temporal_10 f3.u 1037960377095d4e +nlm_hq_temporal_10 f3.v 31cbefd6eea66b89 +nlm_hq_temporal_10 f3.y f220ee7455a8281e +nlm_hq_temporal_10 f4.u 2e29d4c058d324b5 +nlm_hq_temporal_10 f4.v fbd3547ce5e3e8aa +nlm_hq_temporal_10 f4.y 0ad46d104bc93bcf +nlm_hq_temporal_10 f5.u 4452545706606b26 +nlm_hq_temporal_10 f5.v 9e419603e4b4874d +nlm_hq_temporal_10 f5.y 637fee770dd99737 +nlm_hq_temporal_10 f6.u c7ec78d52c88cb3d +nlm_hq_temporal_10 f6.v 2f247ce505843049 +nlm_hq_temporal_10 f6.y 08599c0c5c975af3 +nlm_hq_temporal_10 f7.u f00de96b084ac96d +nlm_hq_temporal_10 f7.v 39b8fb0faf3b24f6 +nlm_hq_temporal_10 f7.y f39a68a6ac679a5e +nlm_hq_temporal_10 f8.u 45cd0f0da37b7402 +nlm_hq_temporal_10 f8.v 41729f2bb18ef3bd +nlm_hq_temporal_10 f8.y 0ff51c2c41ff00e7 +nlm_hq_temporal_10 f9.u 3ab7a1b00a1951a9 +nlm_hq_temporal_10 f9.v 3dbc3defe23b0c27 +nlm_hq_temporal_10 f9.y e10c3015fa8750b6 +nlm_hq_yuv_8 f0.u f970724b17f24e86 +nlm_hq_yuv_8 f0.v 15d15d6b817104fb +nlm_hq_yuv_8 f0.y 4c0f956cc3ba66b7 +nlm_hq_yuv_8 f1.u 3f90a59f584adb7f +nlm_hq_yuv_8 f1.v 271769c2ff81ac00 +nlm_hq_yuv_8 f1.y 8d0c2026167508f7 +nlm_hq_yuv_8 f2.u 2ceb47a4d646a500 +nlm_hq_yuv_8 f2.v 7b5529870e936dc5 +nlm_hq_yuv_8 f2.y 17e33d0daed22553 +nlm_hq_yuv_8 f3.u c6db31e60fd4f0c2 +nlm_hq_yuv_8 f3.v 34d575a3b0b00a6c +nlm_hq_yuv_8 f3.y d179e0244d43d93e +nlm_hq_yuv_8 f4.u 9e4dc2dc78d11d27 +nlm_hq_yuv_8 f4.v 2fe6433a3093ec65 +nlm_hq_yuv_8 f4.y 86a568874c9beff9 +nlm_hq_yuv_8 f5.u 03571ed1600f8e73 +nlm_hq_yuv_8 f5.v fc26a667bfd079fc +nlm_hq_yuv_8 f5.y 1651ff29203647c4 +nlm_hq_yuv_8 f6.u 6227b6335493fde0 +nlm_hq_yuv_8 f6.v cb12d1968b1cdb2d +nlm_hq_yuv_8 f6.y 08f5aab713b6f6b5 +nlm_hq_yuv_8 f7.u 015c5765d390e62c +nlm_hq_yuv_8 f7.v 03b3b7f76910c833 +nlm_hq_yuv_8 f7.y 97d1a658749b6dd4 +nlm_hq_yuv_8 f8.u 3d07fb9423b4358a +nlm_hq_yuv_8 f8.v 8936a47d748a06cc +nlm_hq_yuv_8 f8.y 9dec3321bf7baa7d +nlm_hq_yuv_8 f9.u 21674dce3917607e +nlm_hq_yuv_8 f9.v 153ac3f2416ee668 +nlm_hq_yuv_8 f9.y 911b7a061aa3a328 diff --git a/av-denoise/tests/parity/clips.rs b/av-denoise/tests/parity/clips.rs new file mode 100644 index 0000000..2205bff --- /dev/null +++ b/av-denoise/tests/parity/clips.rs @@ -0,0 +1,56 @@ +use av_denoise::{Depth, FrameLayout, Planes}; + +/// A moving sine texture plus hashed noise, so motion search and the noise estimator both have work. +pub fn clip(layout: FrameLayout, frame_count: usize) -> Vec { + let (chroma_width, chroma_height) = layout.chroma_dims(); + let mut frames = Vec::with_capacity(frame_count); + + for index in 0..frame_count { + let y_plane = plane(layout.width, layout.height, layout.depth, index, 0); + let u_plane = plane(chroma_width, chroma_height, layout.depth, index, 1); + let v_plane = plane(chroma_width, chroma_height, layout.depth, index, 2); + frames.push(Planes { + y: y_plane, + u: u_plane, + v: v_plane, + }); + } + + frames +} + +fn plane(width: u32, height: u32, depth: Depth, frame: usize, channel: u32) -> Vec { + let max_code = depth.max_value(); + let shift = frame as f32 * 1.5; + let mut samples = Vec::with_capacity((width * height) as usize); + + for row in 0..height { + for column in 0..width { + let phase_x = (column as f32 + shift) / width as f32 * std::f32::consts::TAU * 3.0; + let phase_y = row as f32 / height as f32 * std::f32::consts::TAU * 2.0; + let texture = 0.5 + 0.2 * phase_x.sin() * phase_y.cos(); + let noise = hashed_noise(column, row, frame as u32, channel) * 0.03; + let value = (texture + noise).clamp(0.0, 1.0); + let code = (value * max_code).round() as u16; + samples.push(code); + } + } + + encode(&samples, depth) +} + +fn hashed_noise(column: u32, row: u32, frame: u32, channel: u32) -> f32 { + let mut hash = column.wrapping_mul(0x9E37_79B9) ^ row.wrapping_mul(0x85EB_CA6B); + hash ^= frame.wrapping_mul(0xC2B2_AE35) ^ channel.wrapping_mul(0x27D4_EB2F); + hash ^= hash >> 15; + hash = hash.wrapping_mul(0x2C1B_3C6D); + hash ^= hash >> 12; + hash as f32 / u32::MAX as f32 - 0.5 +} + +fn encode(samples: &[u16], depth: Depth) -> Vec { + match depth.bytes_per_sample() { + 1 => samples.iter().map(|&sample| sample as u8).collect(), + _ => samples.iter().flat_map(|sample| sample.to_le_bytes()).collect(), + } +} diff --git a/av-denoise/tests/parity/configs.rs b/av-denoise/tests/parity/configs.rs new file mode 100644 index 0000000..65b8c6c --- /dev/null +++ b/av-denoise/tests/parity/configs.rs @@ -0,0 +1,222 @@ +use av_denoise::accelerate::Accelerator; +use av_denoise::{ + Algorithm, + ChannelIntent, + DenoisingMode, + Depth, + Device, + FrameLayout, + Nl4dOptions, + NlmeansHqOptions, + NlmeansOptions, + PlaneOptions, + Subsampling, +}; + +pub struct ParityConfig { + pub name: &'static str, + pub options: PlaneOptions, + pub layout: FrameLayout, + pub frames: usize, + pub reseed: bool, +} + +const WIDTH: u32 = 160; +const HEIGHT: u32 = 128; + +fn layout(subsampling: Subsampling, depth: Depth) -> FrameLayout { + FrameLayout { + width: WIDTH, + height: HEIGHT, + subsampling, + depth, + } +} + +fn options(intent: ChannelIntent, mode: DenoisingMode, algorithm: Algorithm) -> PlaneOptions { + PlaneOptions { + accelerators: vec![Accelerator::Vulkan], + device: Device::Default, + intent, + mode, + algorithm, + luma_strength: None, + chroma_strength: None, + luma_lambda_ht: None, + chroma_lambda_ht: None, + } +} + +pub fn all() -> Vec { + let temporal = DenoisingMode::Temporal { radius: 2 }; + let fast = Algorithm::Nlmeans(NlmeansOptions::default()); + let hq = Algorithm::NlmeansHq(NlmeansHqOptions::default()); + let nl4d = Algorithm::Nl4d(Nl4dOptions::default()); + let grain_options = Nl4dOptions { + grain_export: true, + ..Nl4dOptions::default() + }; + let nl4d_grain = Algorithm::Nl4d(grain_options); + let windowed_options = Nl4dOptions { + windowed_noise_estimation: true, + ..Nl4dOptions::default() + }; + let nl4d_windowed = Algorithm::Nl4d(windowed_options); + let yuv420_8 = layout(Subsampling::Yuv420, Depth::Eight); + let yuv420_10 = layout(Subsampling::Yuv420, Depth::Ten); + let yuv444_8 = layout(Subsampling::Yuv444, Depth::Eight); + + let nlm_fast_spatial_8 = config( + "nlm_fast_spatial_8", + ChannelIntent::LumaChroma, + DenoisingMode::Spacial, + fast, + yuv420_8, + 6, + false, + ); + let nlm_fast_temporal_8 = config( + "nlm_fast_temporal_8", + ChannelIntent::LumaChroma, + temporal, + fast, + yuv420_8, + 10, + false, + ); + let nlm_hq_temporal_10 = config( + "nlm_hq_temporal_10", + ChannelIntent::LumaChroma, + temporal, + hq, + yuv420_10, + 10, + false, + ); + let nlm_hq_short_8 = config( + "nlm_hq_short_8", + ChannelIntent::Luma, + temporal, + hq, + yuv420_8, + 2, + false, + ); + let nlm_hq_yuv_8 = config( + "nlm_hq_yuv_8", + ChannelIntent::YuvFused, + temporal, + hq, + yuv444_8, + 10, + false, + ); + let nl4d_luma_8 = config( + "nl4d_luma_8", + ChannelIntent::Luma, + temporal, + nl4d, + yuv420_8, + 12, + false, + ); + let nl4d_chroma_10 = config( + "nl4d_chroma_10", + ChannelIntent::Chroma, + temporal, + nl4d, + yuv420_10, + 12, + false, + ); + let nl4d_lumachroma_10 = config( + "nl4d_lumachroma_10", + ChannelIntent::LumaChroma, + temporal, + nl4d, + yuv420_10, + 12, + false, + ); + let nl4d_yuv_8 = config( + "nl4d_yuv_8", + ChannelIntent::YuvFused, + temporal, + nl4d, + yuv444_8, + 12, + false, + ); + let nl4d_short_8 = config( + "nl4d_short_8", + ChannelIntent::Luma, + temporal, + nl4d, + yuv420_8, + 3, + false, + ); + let nl4d_grain_8 = config( + "nl4d_grain_8", + ChannelIntent::Luma, + temporal, + nl4d_grain, + yuv420_8, + 12, + false, + ); + let nl4d_reseed_8 = config( + "nl4d_reseed_8", + ChannelIntent::LumaChroma, + temporal, + nl4d_windowed, + yuv420_8, + 12, + true, + ); + let nlm_hq_reseed_8 = config( + "nlm_hq_reseed_8", + ChannelIntent::LumaChroma, + temporal, + hq, + yuv420_8, + 12, + true, + ); + + vec![ + nlm_fast_spatial_8, + nlm_fast_temporal_8, + nlm_hq_temporal_10, + nlm_hq_short_8, + nlm_hq_yuv_8, + nl4d_luma_8, + nl4d_chroma_10, + nl4d_lumachroma_10, + nl4d_yuv_8, + nl4d_short_8, + nl4d_grain_8, + nl4d_reseed_8, + nlm_hq_reseed_8, + ] +} + +fn config( + name: &'static str, + intent: ChannelIntent, + mode: DenoisingMode, + algorithm: Algorithm, + layout: FrameLayout, + frames: usize, + reseed: bool, +) -> ParityConfig { + let plane_options = options(intent, mode, algorithm); + + ParityConfig { + name, + options: plane_options, + layout, + frames, + reseed, + } +} diff --git a/av-denoise/tests/parity/fixture.rs b/av-denoise/tests/parity/fixture.rs new file mode 100644 index 0000000..9b88797 --- /dev/null +++ b/av-denoise/tests/parity/fixture.rs @@ -0,0 +1,105 @@ +use std::collections::BTreeMap; +use std::path::PathBuf; + +use wgpu::{Adapter, Backends, Instance, InstanceDescriptor}; + +pub type Entries = BTreeMap<(String, String), String>; + +const ADAPTERS_PREFIX: &str = "# adapters: "; + +/// The recorded adapter fingerprint and hashes. +pub struct Fixture { + pub adapters: String, + pub entries: Entries, +} + +/// FNV-1a 64, chosen because its output never changes between Rust releases. +pub fn hash(bytes: &[u8]) -> u64 { + let mut state = 0xcbf2_9ce4_8422_2325u64; + + for &byte in bytes { + state ^= byte as u64; + state = state.wrapping_mul(0x0000_0100_0000_01b3); + } + + state +} + +pub fn path() -> PathBuf { + let manifest_dir = env!("CARGO_MANIFEST_DIR"); + PathBuf::from(manifest_dir).join("tests/fixtures/parity.txt") +} + +/// Names and drivers of every Vulkan adapter, sorted and joined with `; `. +pub fn adapter_fingerprint() -> String { + let mut descriptor = InstanceDescriptor::new_without_display_handle(); + descriptor.backends = Backends::VULKAN; + let instance = Instance::new(descriptor); + let enumeration = instance.enumerate_adapters(Backends::VULKAN); + let adapters = cubecl::future::block_on(enumeration); + + let mut described: Vec = adapters.iter().map(describe).collect(); + described.sort(); + described.join("; ") +} + +/// Panics when the recorded adapters differ from this machine's. +pub fn check_adapters(recorded: &str) { + let current = adapter_fingerprint(); + if recorded == current { + return; + } + + panic!( + "parity fixture was recorded on different adapters\n recorded: {recorded}\n current: {current}\n \ + re-record with AVD_PARITY_RECORD=1 on this machine before starting a refactor, never mid-refactor" + ); +} + +/// Reads the fixture. +/// +/// Every line after the adapters header is `