use std::collections::VecDeque;
use cubecl::Runtime;
use cubecl::prelude::ComputeClient;
use crate::accelerate::Accelerator;
use crate::device::Device;
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;
#[derive(Debug, Clone, bon::Builder)]
pub struct DenoiserOptions {
#[builder(default = ChannelMode::Yuv)]
pub channel_mode: ChannelMode,
#[builder(default = DenoisingMode::Spacial)]
pub mode: DenoisingMode,
#[builder(default)]
pub algorithm: Algorithm,
#[builder(default = OutputFormat::F32)]
pub output_format: OutputFormat,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OutputFormat {
F32,
Wire { depth: Depth },
}
#[derive(Debug, Clone, PartialEq)]
pub enum FrameOutput {
F32(Vec<f32>),
Wire(Vec<u8>),
}
impl FrameOutput {
pub fn into_f32(self) -> Option<Vec<f32>> {
match self {
Self::F32(v) => Some(v),
Self::Wire(_) => None,
}
}
pub fn into_wire(self) -> Option<Vec<u8>> {
match self {
Self::Wire(v) => Some(v),
Self::F32(_) => None,
}
}
pub fn as_f32(&self) -> Option<&[f32]> {
match self {
Self::F32(v) => Some(v),
Self::Wire(_) => None,
}
}
pub fn as_wire(&self) -> Option<&[u8]> {
match self {
Self::Wire(v) => Some(v),
Self::F32(_) => None,
}
}
}
#[derive(Debug, Copy, Clone, PartialEq)]
pub enum Algorithm {
Nlmeans(NlmeansOptions),
NlmeansHq(NlmeansHqOptions),
Nl4d(Nl4dOptions),
}
impl Default for Algorithm {
fn default() -> Self {
Self::Nlmeans(NlmeansOptions::default())
}
}
#[derive(Debug, Copy, Clone, Default, PartialEq)]
pub struct NlmeansOptions {
pub prefilter: PrefilterMode,
pub motion_compensation: MotionCompensationMode,
pub tuning: NlmTuning,
}
#[derive(Debug, Copy, Clone, Default, PartialEq)]
pub struct NlmeansHqOptions {
pub nlm: NlmeansOptions,
pub hq: HqParams,
}
#[derive(Debug, Copy, Clone, PartialEq)]
pub struct Nl4dOptions {
pub motion: MotionSearch,
pub sigma: Option<f32>,
pub sigma_scale: f32,
pub thsad_scale: f32,
pub refine: u32,
pub spatial_radius: u32,
pub lambda_ht: Option<f32>,
pub lambda_ht_scale: f32,
pub c_min: f32,
pub mismatch_scale: f32,
pub confidence_variance: bool,
pub windowed_noise_estimation: 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,
lambda_ht: None,
lambda_ht_scale: 1.0,
c_min: defaults.c_min,
mismatch_scale: defaults.mismatch_scale,
confidence_variance: defaults.confidence_variance,
windowed_noise_estimation: false,
}
}
}
impl Nl4dOptions {
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()
}
}
}
pub fn nl4d_default_lambda_ht(channels: ChannelMode) -> f32 {
match channels {
ChannelMode::Luma | ChannelMode::Yuv => 5.3,
ChannelMode::Chroma => 4.2,
}
}
fn resolve_lambda_ht(opts: &Nl4dOptions, channels: ChannelMode) -> Result<f32, String> {
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)
}
#[derive(Debug, Copy, Clone, Default, PartialEq, Eq, strum_macros::EnumString)]
#[strum(ascii_case_insensitive)]
pub enum Preset {
Veryfast,
Fast,
#[default]
Base,
Slow,
Veryslow,
}
#[derive(Debug, Copy, Clone, PartialEq, Eq, strum_macros::EnumString)]
#[strum(ascii_case_insensitive)]
pub enum NlmeansVariant {
Fast,
Hq,
}
pub fn nlmeans_variant_for(preset: Preset) -> NlmeansVariant {
match preset {
Preset::Veryfast => NlmeansVariant::Fast,
Preset::Fast | Preset::Base | Preset::Slow | Preset::Veryslow => NlmeansVariant::Hq,
}
}
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,
}
}
pub fn nlmeans_search_radius_for(preset: Preset) -> u32 {
match preset {
Preset::Veryfast | Preset::Fast | Preset::Base => 2,
Preset::Slow | Preset::Veryslow => 4,
}
}
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,
}
}
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
},
}
}
#[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub enum DenoisingMode {
Spacial,
Temporal { radius: u32 },
}
#[derive(Debug, Copy, Clone, Default, PartialEq)]
pub struct NlmTuning {
pub search_radius: Option<u32>,
pub patch_radius: Option<u32>,
pub strength: Option<f32>,
pub self_weight: Option<f32>,
}
impl DenoiserOptions {
#[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),
Algorithm::Nl4d(opts) => NlmParams {
channels: self.channel_mode,
motion_compensation: opts.motion.into(),
temporal_radius,
hq: Some(opts.to_hq_params()),
..NlmParams::default()
},
}
}
fn nlm_params_for(&self, opts: NlmeansOptions, hq: Option<HqParams>, temporal_radius: u32) -> NlmParams {
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),
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum DenoiserError {
#[error("denoiser queue is full, collect the pending frame before pushing more")]
QueueFull,
#[error("denoiser failed earlier, reset the stream before using it again")]
Poisoned,
#[error("no accelerator from the priority list is available")]
NoAcceleratorAvailable,
#[error(transparent)]
Other(#[from] anyhow::Error),
}
enum Engine<R: Runtime> {
Nlm(Box<NlmDenoiser<R>>),
Nl4d(Box<Nl4dDenoiser<R>>),
}
impl<R: Runtime> Engine<R> {
fn is_nl4d(&self) -> bool {
matches!(self, Self::Nl4d(_))
}
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<Option<Pending<R>>, anyhow::Error> {
match self {
Self::Nlm(d) => d.denoise_submit(),
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 build_engine<R: Runtime>(
client: &ComputeClient<R>,
algorithm: &Algorithm,
params: NlmParams,
width: u32,
height: u32,
output_format: OutputFormat,
) -> Result<Engine<R>, DenoiserError> {
match algorithm {
Algorithm::Nl4d(opts) => {
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,
mismatch_scale: opts.mismatch_scale,
confidence_variance: opts.confidence_variance,
};
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<cubecl::cuda::CudaRuntime>),
#[cfg(feature = "rocm")]
Rocm(Engine<cubecl::hip::HipRuntime>),
#[cfg(any(feature = "vulkan", feature = "metal"))]
Wgpu(Engine<cubecl::wgpu::WgpuRuntime>),
}
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(),
}
}
#[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<cubecl::cuda::CudaRuntime>),
#[cfg(feature = "rocm")]
Rocm(Pending<cubecl::hip::HipRuntime>),
#[cfg(any(feature = "vulkan", feature = "metal"))]
Wgpu(Pending<cubecl::wgpu::WgpuRuntime>),
}
impl BackendPending {
fn wait(self) -> Result<FrameOutput, anyhow::Error> {
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(),
}
}
fn try_wait(self) -> Result<Result<FrameOutput, Self>, 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))),
},
}
}
}
pub const MAX_PENDING: usize = 2;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct WindowSpan {
pub behind: usize,
pub ahead: usize,
}
impl WindowSpan {
pub fn frame_count(&self) -> usize {
self.behind + 1 + self.ahead
}
}
pub struct Denoiser {
backend: Backend,
pending: VecDeque<BackendPending>,
accelerator: Accelerator,
width: u32,
height: u32,
temporal_radius: u32,
output_format: OutputFormat,
frames_pushed: u32,
poisoned: bool,
}
impl Denoiser {
pub fn create(
accelerators: &[Accelerator],
device: &Device,
width: u32,
height: u32,
options: DenoiserOptions,
) -> Result<Self, DenoiserError> {
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,
})
}
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
}
pub fn output_format(&self) -> OutputFormat {
self.output_format
}
pub fn window_span(&self) -> WindowSpan {
let radius = self.temporal_radius as usize;
let span = if self.backend.is_nl4d() {
2 * radius
} else {
radius
};
WindowSpan {
behind: span,
ahead: span,
}
}
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> {
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(())
}
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> {
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(())
}
pub fn push_frame_wire_priming(&mut self, planes: &[&[u8]], depth: Depth) -> Result<(), DenoiserError> {
if self.poisoned {
return Err(DenoiserError::Poisoned);
}
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(())
}
pub fn push_frame_priming(&mut self, frame: &[f32]) -> Result<(), DenoiserError> {
if self.poisoned {
return Err(DenoiserError::Poisoned);
}
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(())
}
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(),
}
}
pub fn recv_frame(&mut self) -> Result<Option<FrameOutput>, DenoiserError> {
if self.poisoned {
return Err(DenoiserError::Poisoned);
}
self.recv_frame_inner().inspect_err(|_| self.poisoned = true)
}
fn recv_frame_inner(&mut self) -> Result<Option<FrameOutput>, DenoiserError> {
let Some(pending) = self.pending.pop_front() else {
return Ok(None);
};
Ok(Some(pending.wait()?))
}
pub fn try_recv_frame(&mut self) -> Result<Option<FrameOutput>, 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<Option<FrameOutput>, 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)
},
}
}
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> {
while let Some(frame) = self.recv_frame_inner()? {
sink(frame);
}
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()))?,
}
self.frames_pushed = 0;
Ok(())
}
#[cfg(test)]
pub(crate) fn poison_for_test(&mut self) {
self.poisoned = true;
}
#[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<Backend, DenoiserError> {
match accel {
#[cfg(feature = "cuda")]
Accelerator::Cuda => {
let dev = device.to_cuda()?;
let client = <cubecl::cuda::CudaRuntime as Runtime>::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 = <cubecl::hip::HipRuntime as Runtime>::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 = <cubecl::wgpu::WgpuRuntime as Runtime>::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 = <cubecl::wgpu::WgpuRuntime as Runtime>::client(&dev);
Ok(Backend::Wgpu(build_engine(
&client,
algorithm,
params,
width,
height,
output_format,
)?))
},
#[cfg(docsrs)]
#[expect(
unreachable_patterns,
reason = "the arm only keeps the match exhaustive on docs.rs"
)]
_ => unreachable!(),
}
}
#[cfg(test)]
mod options_tests {
use super::*;
fn hq(hq: HqParams) -> Algorithm {
Algorithm::NlmeansHq(NlmeansHqOptions {
hq,
..NlmeansHqOptions::default()
})
}
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 - 5.3).abs() < f32::EPSILON);
assert!((chroma - 4.2).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 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 - 5.3).abs() < f32::EPSILON, "got {luma}");
assert!((chroma - 4.2).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}"
);
}
}
#[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() {
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() {
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);
assert_eq!(opts.confidence_variance, params.confidence_variance);
assert_eq!(opts.lambda_ht, None);
assert!((params.lambda_ht - nl4d_default_lambda_ht(ChannelMode::Yuv)).abs() < f32::EPSILON);
}
#[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"
);
}
#[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);
}
}
#[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));
}
#[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);
}
#[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<f32> {
vec![0.5f32; (w * h) as usize]
}
fn f32_out(out: FrameOutput) -> Vec<f32> {
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);
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);
}
#[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"),
}
}
#[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");
}
#[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");
}
#[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();
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);
d.push_frame(&frame(16, 16)).expect("push after drain failed");
}
#[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");
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<f32> {
vec![value; (w * h) as usize]
}
fn push_n_with_drain(d: &mut Denoiser, n: usize, value: f32, out: &mut Vec<Vec<f32>>) {
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);
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");
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"
);
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() {
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()
);
}
}
}