1use std::collections::VecDeque;
2
3use cubecl::Runtime;
4use cubecl::prelude::ComputeClient;
5
6use crate::accelerate::Accelerator;
7use crate::device::Device;
8use crate::nl4d::{Nl4dDenoiser, Nl4dParams};
9#[cfg(test)]
10use crate::nlmeans::MotionEstimation;
11use crate::nlmeans::{
12 ChannelMode,
13 HqParams,
14 MotionCompensationMode,
15 MotionSearch,
16 NlmDenoiser,
17 NlmParams,
18 Pending,
19 PrefilterMode,
20 hq_default_strength,
21 validate_dimensions,
22};
23use crate::sniff::sniff_best_accelerator;
24
25#[derive(Debug, Clone, bon::Builder)]
33pub struct DenoiserOptions {
34 #[builder(default = ChannelMode::Yuv)]
36 pub channel_mode: ChannelMode,
37 #[builder(default = DenoisingMode::Spacial)]
40 pub mode: DenoisingMode,
41 #[builder(default)]
44 pub algorithm: Algorithm,
45}
46
47#[derive(Debug, Copy, Clone, PartialEq)]
52pub enum Algorithm {
53 Nlmeans(NlmeansOptions),
56 NlmeansHq(NlmeansHqOptions),
62 Nl4d(Nl4dOptions),
69}
70
71impl Default for Algorithm {
72 fn default() -> Self {
73 Self::Nlmeans(NlmeansOptions::default())
74 }
75}
76
77#[derive(Debug, Copy, Clone, Default, PartialEq)]
79pub struct NlmeansOptions {
80 pub prefilter: PrefilterMode,
85 pub motion_compensation: MotionCompensationMode,
94 pub tuning: NlmTuning,
97}
98
99#[derive(Debug, Copy, Clone, Default, PartialEq)]
101pub struct NlmeansHqOptions {
102 pub nlm: NlmeansOptions,
104 pub hq: HqParams,
106}
107
108#[derive(Debug, Copy, Clone, PartialEq)]
126pub struct Nl4dOptions {
127 pub motion: MotionSearch,
129 pub sigma: Option<f32>,
135 pub sigma_scale: f32,
141 pub thsad_scale: f32,
147 pub refine: u32,
150 pub spatial_radius: u32,
153 pub lambda_ht: Option<f32>,
159 pub lambda_ht_scale: f32,
166 pub c_min: f32,
171 pub mismatch_scale: f32,
178 pub confidence_variance: bool,
182 pub windowed_noise_estimation: bool,
192}
193
194impl Default for Nl4dOptions {
195 fn default() -> Self {
196 let defaults = Nl4dParams::default();
197 let hq = HqParams::default();
198 Self {
199 motion: MotionSearch::default(),
200 sigma: hq.sigma_override,
201 sigma_scale: hq.sigma_scale,
202 thsad_scale: hq.thsad_scale,
203 refine: defaults.refine,
204 spatial_radius: defaults.spatial_radius,
205 lambda_ht: None,
209 lambda_ht_scale: 1.0,
210 c_min: defaults.c_min,
211 mismatch_scale: defaults.mismatch_scale,
212 confidence_variance: defaults.confidence_variance,
213 windowed_noise_estimation: false,
214 }
215 }
216}
217
218impl Nl4dOptions {
219 fn to_hq_params(self) -> HqParams {
226 HqParams {
227 sigma_override: self.sigma,
228 sigma_scale: self.sigma_scale,
229 thsad_scale: self.thsad_scale,
230 temporal_confidence: true,
231 windowed_noise_estimation: self.windowed_noise_estimation,
232 ..HqParams::default()
233 }
234 }
235}
236
237pub fn nl4d_default_lambda_ht(channels: ChannelMode) -> f32 {
256 match channels {
257 ChannelMode::Luma | ChannelMode::Yuv => 5.3,
258 ChannelMode::Chroma => 4.2,
259 }
260}
261
262fn resolve_lambda_ht(opts: &Nl4dOptions, channels: ChannelMode) -> Result<f32, String> {
274 if !(opts.lambda_ht_scale.is_finite() && (0.1..=10.0).contains(&opts.lambda_ht_scale)) {
275 return Err(format!(
276 "lambda_ht_scale must be finite and in [0.1, 10.0], got {}",
277 opts.lambda_ht_scale
278 ));
279 }
280
281 let lambda_ht = opts.lambda_ht.unwrap_or_else(|| nl4d_default_lambda_ht(channels));
282
283 Ok(lambda_ht * opts.lambda_ht_scale)
284}
285
286#[derive(Debug, Copy, Clone, Default, PartialEq, Eq, strum_macros::EnumString)]
297#[strum(ascii_case_insensitive)]
298pub enum Preset {
299 Veryfast,
301 Fast,
303 #[default]
305 Base,
306 Slow,
308 Veryslow,
310}
311
312#[derive(Debug, Copy, Clone, PartialEq, Eq, strum_macros::EnumString)]
314#[strum(ascii_case_insensitive)]
315pub enum NlmeansVariant {
316 Fast,
318 Hq,
321}
322
323pub fn nlmeans_variant_for(preset: Preset) -> NlmeansVariant {
325 match preset {
326 Preset::Veryfast => NlmeansVariant::Fast,
327 Preset::Fast | Preset::Base | Preset::Slow | Preset::Veryslow => NlmeansVariant::Hq,
328 }
329}
330
331pub fn nlmeans_temporal_radius_for(preset: Preset) -> u32 {
334 match preset {
335 Preset::Veryfast => 0,
336 Preset::Fast => 1,
337 Preset::Base => 2,
338 Preset::Slow => 4,
339 Preset::Veryslow => 8,
340 }
341}
342
343pub fn nlmeans_search_radius_for(preset: Preset) -> u32 {
346 match preset {
347 Preset::Veryfast | Preset::Fast | Preset::Base => 2,
348 Preset::Slow | Preset::Veryslow => 4,
349 }
350}
351
352pub fn nl4d_temporal_radius_for(preset: Preset) -> u32 {
358 match preset {
359 Preset::Veryfast | Preset::Fast => 1,
360 Preset::Base => 2,
361 Preset::Slow => 4,
362 Preset::Veryslow => 8,
363 }
364}
365
366pub fn nl4d_spatial_radius_for(preset: Preset) -> u32 {
377 match preset {
378 Preset::Veryfast => 6,
379 Preset::Fast | Preset::Base | Preset::Slow | Preset::Veryslow => {
380 Nl4dOptions::default().spatial_radius
381 },
382 }
383}
384
385#[derive(Debug, Copy, Clone, Eq, PartialEq)]
387pub enum DenoisingMode {
388 Spacial,
390 Temporal { radius: u32 },
392}
393
394#[derive(Debug, Copy, Clone, Default, PartialEq)]
399pub struct NlmTuning {
400 pub search_radius: Option<u32>,
401 pub patch_radius: Option<u32>,
402 pub strength: Option<f32>,
403 pub self_weight: Option<f32>,
404}
405
406impl DenoiserOptions {
407 #[doc(hidden)]
421 pub fn to_nlm_params(&self) -> NlmParams {
422 let temporal_radius = match self.mode {
423 DenoisingMode::Spacial => 0,
424 DenoisingMode::Temporal { radius } => radius,
425 };
426
427 match self.algorithm {
428 Algorithm::Nlmeans(opts) => self.nlm_params_for(opts, None, temporal_radius),
429 Algorithm::NlmeansHq(opts) => self.nlm_params_for(opts.nlm, Some(opts.hq), temporal_radius),
430 Algorithm::Nl4d(opts) => NlmParams {
434 channels: self.channel_mode,
435 motion_compensation: opts.motion.into(),
436 temporal_radius,
437 hq: Some(opts.to_hq_params()),
438 ..NlmParams::default()
439 },
440 }
441 }
442
443 fn nlm_params_for(&self, opts: NlmeansOptions, hq: Option<HqParams>, temporal_radius: u32) -> NlmParams {
446 let strength = opts.tuning.strength.unwrap_or(match hq {
462 Some(hq) if hq.auto_strength => hq_default_strength(self.channel_mode, temporal_radius),
463 _ => NlmParams::default().strength,
464 });
465
466 let defaults = NlmParams::default();
467 NlmParams {
468 channels: self.channel_mode,
469 prefilter: opts.prefilter,
470 motion_compensation: opts.motion_compensation,
471 temporal_radius,
472 hq,
473 strength,
474 search_radius: opts.tuning.search_radius.unwrap_or(defaults.search_radius),
475 patch_radius: opts.tuning.patch_radius.unwrap_or(defaults.patch_radius),
476 self_weight: opts.tuning.self_weight.unwrap_or(defaults.self_weight),
477 }
478 }
479}
480
481#[derive(Debug, thiserror::Error)]
483pub enum DenoiserError {
484 #[error("denoiser queue is full, collect the pending frame before pushing more")]
490 QueueFull,
491 #[error("no accelerator from the priority list is available")]
493 NoAcceleratorAvailable,
494 #[error(transparent)]
497 Other(#[from] anyhow::Error),
498}
499
500enum Engine<R: Runtime> {
506 Nlm(Box<NlmDenoiser<R>>),
507 Nl4d(Box<Nl4dDenoiser<R>>),
508}
509
510impl<R: Runtime> Engine<R> {
511 fn is_nl4d(&self) -> bool {
512 matches!(self, Self::Nl4d(_))
513 }
514
515 fn push_frame(&mut self, frame: &[f32]) {
516 match self {
517 Self::Nlm(d) => d.push_frame(frame),
518 Self::Nl4d(d) => d.push_frame(frame),
519 }
520 }
521
522 fn denoise_submit(&mut self) -> Result<Option<Pending<R>>, anyhow::Error> {
523 match self {
524 Self::Nlm(d) => d.denoise_submit(),
525 Self::Nl4d(d) => d.denoise_submit().map_err(anyhow::Error::from),
530 }
531 }
532
533 fn flush(&mut self, sink: impl FnMut(&[f32])) -> Result<(), anyhow::Error> {
534 match self {
535 Self::Nlm(d) => d.flush(sink),
536 Self::Nl4d(d) => d.flush(sink).map_err(anyhow::Error::from),
537 }
538 }
539
540 fn reset_stream(&mut self) {
541 match self {
542 Self::Nlm(d) => d.reset_stream_state(),
543 Self::Nl4d(d) => d.reset_stream(),
544 }
545 }
546}
547
548fn build_engine<R: Runtime>(
558 client: &ComputeClient<R>,
559 algorithm: &Algorithm,
560 params: NlmParams,
561 width: u32,
562 height: u32,
563) -> Result<Engine<R>, DenoiserError> {
564 match algorithm {
565 Algorithm::Nl4d(opts) => {
566 if params.temporal_radius == 0 {
569 return Err(DenoiserError::Other(anyhow::anyhow!(
570 "nl4d needs a temporal window, set DenoiserOptions::mode to \
571 DenoisingMode::Temporal"
572 )));
573 }
574
575 let lambda_ht = resolve_lambda_ht(opts, params.channels)
576 .map_err(|e| DenoiserError::Other(anyhow::anyhow!(e)))?;
577 let nl4d_params = Nl4dParams {
578 temporal_radius: params.temporal_radius,
579 nlm: params,
580 refine: opts.refine,
581 spatial_radius: opts.spatial_radius,
582 lambda_ht,
583 c_min: opts.c_min,
584 mismatch_scale: opts.mismatch_scale,
585 confidence_variance: opts.confidence_variance,
586 };
587 let denoiser = Nl4dDenoiser::new(client, nl4d_params, width, height)
588 .map_err(|e| DenoiserError::Other(anyhow::anyhow!(e)))?;
589 Ok(Engine::Nl4d(Box::new(denoiser)))
590 },
591 Algorithm::Nlmeans(_) | Algorithm::NlmeansHq(_) => Ok(Engine::Nlm(Box::new(NlmDenoiser::new(
592 client, params, width, height,
593 )))),
594 }
595}
596
597enum Backend {
598 #[cfg(feature = "cuda")]
599 Cuda(Engine<cubecl::cuda::CudaRuntime>),
600 #[cfg(feature = "rocm")]
601 Rocm(Engine<cubecl::hip::HipRuntime>),
602 #[cfg(any(feature = "vulkan", feature = "metal"))]
603 Wgpu(Engine<cubecl::wgpu::WgpuRuntime>),
604}
605
606impl Backend {
607 fn is_nl4d(&self) -> bool {
608 match self {
609 #[cfg(feature = "cuda")]
610 Self::Cuda(e) => e.is_nl4d(),
611 #[cfg(feature = "rocm")]
612 Self::Rocm(e) => e.is_nl4d(),
613 #[cfg(any(feature = "vulkan", feature = "metal"))]
614 Self::Wgpu(e) => e.is_nl4d(),
615 }
616 }
617}
618
619enum BackendPending {
620 #[cfg(feature = "cuda")]
621 Cuda(Pending<cubecl::cuda::CudaRuntime>),
622 #[cfg(feature = "rocm")]
623 Rocm(Pending<cubecl::hip::HipRuntime>),
624 #[cfg(any(feature = "vulkan", feature = "metal"))]
625 Wgpu(Pending<cubecl::wgpu::WgpuRuntime>),
626}
627
628impl BackendPending {
629 fn wait(self) -> Result<Vec<f32>, anyhow::Error> {
630 match self {
631 #[cfg(feature = "cuda")]
632 Self::Cuda(p) => p.wait(),
633 #[cfg(feature = "rocm")]
634 Self::Rocm(p) => p.wait(),
635 #[cfg(any(feature = "vulkan", feature = "metal"))]
636 Self::Wgpu(p) => p.wait(),
637 }
638 }
639}
640
641pub const MAX_PENDING: usize = 2;
648
649#[derive(Debug, Clone, Copy, PartialEq, Eq)]
660pub struct WindowSpan {
661 pub behind: usize,
663 pub ahead: usize,
665}
666
667impl WindowSpan {
668 pub fn frame_count(&self) -> usize {
671 self.behind + 1 + self.ahead
672 }
673}
674
675pub struct Denoiser {
725 backend: Backend,
726 pending: VecDeque<BackendPending>,
727 accelerator: Accelerator,
728 width: u32,
729 height: u32,
730 channels: u32,
731 temporal_radius: u32,
732 frames_pushed: u32,
733}
734
735impl Denoiser {
736 pub fn create(
763 accelerators: &[Accelerator],
764 device: &Device,
765 width: u32,
766 height: u32,
767 options: DenoiserOptions,
768 ) -> Result<Self, DenoiserError> {
769 let accelerator =
770 sniff_best_accelerator(accelerators, device).ok_or(DenoiserError::NoAcceleratorAvailable)?;
771
772 let params = options.to_nlm_params();
773 params.validate()?;
774 validate_dimensions(width, height)?;
775
776 let channels = params.channels.count();
777 let temporal_radius = params.temporal_radius;
778 let backend = build_backend(accelerator, device, &options.algorithm, params, width, height)?;
779
780 Ok(Self {
781 backend,
782 pending: VecDeque::with_capacity(MAX_PENDING),
783 accelerator,
784 width,
785 height,
786 channels,
787 temporal_radius,
788 frames_pushed: 0,
789 })
790 }
791
792 pub fn selected_accelerator(&self) -> Accelerator {
794 self.accelerator
795 }
796
797 pub fn width(&self) -> u32 {
799 self.width
800 }
801
802 pub fn height(&self) -> u32 {
804 self.height
805 }
806
807 pub fn temporal_radius(&self) -> u32 {
809 self.temporal_radius
810 }
811
812 pub fn window_span(&self) -> WindowSpan {
830 let radius = self.temporal_radius as usize;
831 let span = if self.backend.is_nl4d() {
832 2 * radius
833 } else {
834 radius
835 };
836 WindowSpan {
837 behind: span,
838 ahead: span,
839 }
840 }
841
842 pub fn push_frame(&mut self, frame: &[f32]) -> Result<(), DenoiserError> {
856 let window_full = self.frames_pushed > self.temporal_radius;
860 if window_full && self.pending.len() >= MAX_PENDING {
861 return Err(DenoiserError::QueueFull);
862 }
863
864 match &mut self.backend {
865 #[cfg(feature = "cuda")]
866 Backend::Cuda(d) => {
867 d.push_frame(frame);
868 if let Some(p) = d.denoise_submit()? {
869 self.pending.push_back(BackendPending::Cuda(p));
870 }
871 },
872 #[cfg(feature = "rocm")]
873 Backend::Rocm(d) => {
874 d.push_frame(frame);
875 if let Some(p) = d.denoise_submit()? {
876 self.pending.push_back(BackendPending::Rocm(p));
877 }
878 },
879 #[cfg(any(feature = "vulkan", feature = "metal"))]
880 Backend::Wgpu(d) => {
881 d.push_frame(frame);
882 if let Some(p) = d.denoise_submit()? {
883 self.pending.push_back(BackendPending::Wgpu(p));
884 }
885 },
886 }
887
888 self.frames_pushed = self.frames_pushed.saturating_add(1);
889 Ok(())
890 }
891
892 pub fn push_frame_priming(&mut self, frame: &[f32]) -> Result<(), DenoiserError> {
901 match &mut self.backend {
902 #[cfg(feature = "cuda")]
903 Backend::Cuda(d) => d.push_frame(frame),
904 #[cfg(feature = "rocm")]
905 Backend::Rocm(d) => d.push_frame(frame),
906 #[cfg(any(feature = "vulkan", feature = "metal"))]
907 Backend::Wgpu(d) => d.push_frame(frame),
908 }
909
910 self.frames_pushed = self.frames_pushed.saturating_add(1);
911 Ok(())
912 }
913
914 pub fn reset_stream(&mut self) {
919 self.pending.clear();
920 self.frames_pushed = 0;
921
922 match &mut self.backend {
923 #[cfg(feature = "cuda")]
924 Backend::Cuda(d) => d.reset_stream(),
925 #[cfg(feature = "rocm")]
926 Backend::Rocm(d) => d.reset_stream(),
927 #[cfg(any(feature = "vulkan", feature = "metal"))]
928 Backend::Wgpu(d) => d.reset_stream(),
929 }
930 }
931
932 pub fn recv_frame(&mut self) -> Result<Option<Vec<f32>>, DenoiserError> {
938 let Some(pending) = self.pending.pop_front() else {
939 return Ok(None);
940 };
941 Ok(Some(pending.wait()?))
942 }
943
944 pub fn try_recv_frame(&mut self) -> Result<Option<Vec<f32>>, DenoiserError> {
950 self.recv_frame()
951 }
952
953 pub fn flush(&mut self, mut sink: impl FnMut(Vec<f32>)) -> Result<(), DenoiserError> {
966 while let Some(frame) = self.recv_frame()? {
969 sink(frame);
970 }
971
972 let pixels = (self.width * self.height) as usize;
973 let channels = self.channels as usize;
974 let scratch_cap = pixels * channels;
975
976 match &mut self.backend {
977 #[cfg(feature = "cuda")]
978 Backend::Cuda(d) => d.flush(|slice| {
979 let mut v = Vec::with_capacity(scratch_cap);
980 v.extend_from_slice(slice);
981 sink(v);
982 })?,
983 #[cfg(feature = "rocm")]
984 Backend::Rocm(d) => d.flush(|slice| {
985 let mut v = Vec::with_capacity(scratch_cap);
986 v.extend_from_slice(slice);
987 sink(v);
988 })?,
989 #[cfg(any(feature = "vulkan", feature = "metal"))]
990 Backend::Wgpu(d) => d.flush(|slice| {
991 let mut v = Vec::with_capacity(scratch_cap);
992 v.extend_from_slice(slice);
993 sink(v);
994 })?,
995 }
996
997 self.frames_pushed = 0;
1001
1002 Ok(())
1003 }
1004}
1005
1006fn build_backend(
1007 accel: Accelerator,
1008 device: &Device,
1009 algorithm: &Algorithm,
1010 params: NlmParams,
1011 width: u32,
1012 height: u32,
1013) -> Result<Backend, DenoiserError> {
1014 match accel {
1015 #[cfg(feature = "cuda")]
1016 Accelerator::Cuda => {
1017 let dev = device.to_cuda()?;
1018 let client = <cubecl::cuda::CudaRuntime as Runtime>::client(&dev);
1019 Ok(Backend::Cuda(build_engine(
1020 &client, algorithm, params, width, height,
1021 )?))
1022 },
1023 #[cfg(feature = "rocm")]
1024 Accelerator::Rocm => {
1025 let dev = device.to_amd()?;
1026 let client = <cubecl::hip::HipRuntime as Runtime>::client(&dev);
1027 Ok(Backend::Rocm(build_engine(
1028 &client, algorithm, params, width, height,
1029 )?))
1030 },
1031 #[cfg(feature = "vulkan")]
1032 Accelerator::Vulkan => {
1033 let dev = device.to_wgpu()?;
1034 let client = <cubecl::wgpu::WgpuRuntime as Runtime>::client(&dev);
1035 Ok(Backend::Wgpu(build_engine(
1036 &client, algorithm, params, width, height,
1037 )?))
1038 },
1039 #[cfg(feature = "metal")]
1040 Accelerator::Metal => {
1041 let dev = device.to_wgpu()?;
1042 let client = <cubecl::wgpu::WgpuRuntime as Runtime>::client(&dev);
1043 Ok(Backend::Wgpu(build_engine(
1044 &client, algorithm, params, width, height,
1045 )?))
1046 },
1047 #[cfg(docsrs)]
1051 #[allow(unreachable_patterns)]
1052 _ => unreachable!(),
1053 }
1054}
1055
1056#[cfg(test)]
1057mod options_tests {
1058 use super::*;
1059
1060 fn hq(hq: HqParams) -> Algorithm {
1063 Algorithm::NlmeansHq(NlmeansHqOptions {
1064 hq,
1065 ..NlmeansHqOptions::default()
1066 })
1067 }
1068
1069 fn fast_tuned(tuning: NlmTuning) -> Algorithm {
1071 Algorithm::Nlmeans(NlmeansOptions {
1072 tuning,
1073 ..NlmeansOptions::default()
1074 })
1075 }
1076
1077 #[test]
1078 fn nl4d_default_lambda_ht_differs_between_luma_and_chroma() {
1079 let luma = nl4d_default_lambda_ht(ChannelMode::Luma);
1080 let chroma = nl4d_default_lambda_ht(ChannelMode::Chroma);
1081
1082 assert!((luma - 5.3).abs() < f32::EPSILON);
1083 assert!((chroma - 4.2).abs() < f32::EPSILON);
1084 assert!(
1085 (chroma - luma).abs() > f32::EPSILON,
1086 "the two planes should not resolve to the same default"
1087 );
1088 }
1089
1090 #[test]
1091 fn nl4d_default_lambda_ht_yuv_reads_the_luma_value() {
1092 let yuv = nl4d_default_lambda_ht(ChannelMode::Yuv);
1093 let luma = nl4d_default_lambda_ht(ChannelMode::Luma);
1094
1095 assert!((yuv - luma).abs() < f32::EPSILON);
1096 }
1097
1098 #[test]
1099 fn resolve_lambda_ht_unset_uses_the_per_plane_default() {
1100 let opts = Nl4dOptions::default();
1101
1102 let luma = resolve_lambda_ht(&opts, ChannelMode::Luma).expect("the default scale is in range");
1103 let chroma = resolve_lambda_ht(&opts, ChannelMode::Chroma).expect("the default scale is in range");
1104
1105 assert!((luma - 5.3).abs() < f32::EPSILON, "got {luma}");
1106 assert!((chroma - 4.2).abs() < f32::EPSILON, "got {chroma}");
1107 }
1108
1109 #[test]
1110 fn resolve_lambda_ht_explicit_value_overrides_every_plane() {
1111 let opts = Nl4dOptions {
1112 lambda_ht: Some(4.4),
1113 ..Nl4dOptions::default()
1114 };
1115
1116 for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] {
1117 let got = resolve_lambda_ht(&opts, channels).expect("the default scale is in range");
1118 assert!(
1119 (got - 4.4).abs() < f32::EPSILON,
1120 "channels {channels:?} got {got}"
1121 );
1122 }
1123 }
1124
1125 #[test]
1126 fn resolve_lambda_ht_default_scale_leaves_the_value_alone() {
1127 let opts = Nl4dOptions::default();
1128
1129 for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] {
1130 let got = resolve_lambda_ht(&opts, channels).expect("the default scale is in range");
1131 let want = nl4d_default_lambda_ht(channels);
1132 assert!(
1133 (got - want).abs() < f32::EPSILON,
1134 "channels {channels:?} got {got}"
1135 );
1136 }
1137 }
1138
1139 #[test]
1140 fn resolve_lambda_ht_scale_multiplies_the_per_plane_default() {
1141 let opts = Nl4dOptions {
1142 lambda_ht_scale: 1.1,
1143 ..Nl4dOptions::default()
1144 };
1145
1146 for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] {
1147 let got = resolve_lambda_ht(&opts, channels).expect("1.1 is in range");
1148 let want = nl4d_default_lambda_ht(channels) * 1.1;
1149 assert!(
1150 (got - want).abs() < 1e-5,
1151 "channels {channels:?} got {got}, want {want}"
1152 );
1153 }
1154 }
1155
1156 #[test]
1159 fn resolve_lambda_ht_scale_multiplies_an_explicit_value() {
1160 let opts = Nl4dOptions {
1161 lambda_ht: Some(5.0),
1162 lambda_ht_scale: 0.9,
1163 ..Nl4dOptions::default()
1164 };
1165
1166 let got = resolve_lambda_ht(&opts, ChannelMode::Luma).expect("0.9 is in range");
1167 assert!((got - 4.5).abs() < 1e-5, "got {got}");
1168 }
1169
1170 #[test]
1171 fn resolve_lambda_ht_rejects_an_out_of_range_scale() {
1172 for bad in [0.0, -1.0, 0.05, 10.5, f32::NAN, f32::INFINITY] {
1173 let opts = Nl4dOptions {
1174 lambda_ht_scale: bad,
1175 ..Nl4dOptions::default()
1176 };
1177 let err = resolve_lambda_ht(&opts, ChannelMode::Luma).unwrap_err();
1178 assert!(
1179 err.contains("lambda_ht_scale"),
1180 "lambda_ht_scale={bad} should be rejected, got {err}"
1181 );
1182 }
1183 }
1184
1185 #[test]
1186 fn the_default_algorithm_is_the_fast_nlmeans_path() {
1187 let opts = DenoiserOptions::builder().build();
1188 assert_eq!(opts.algorithm, Algorithm::Nlmeans(NlmeansOptions::default()));
1189 }
1190
1191 #[test]
1192 fn spatial_mode_maps_to_zero_temporal_radius() {
1193 let opts = DenoiserOptions::builder()
1194 .channel_mode(ChannelMode::Yuv)
1195 .mode(DenoisingMode::Spacial)
1196 .build();
1197 let params = opts.to_nlm_params();
1198
1199 assert_eq!(params.temporal_radius, 0);
1200 assert_eq!(params.channels, ChannelMode::Yuv);
1201 }
1202
1203 #[test]
1204 fn temporal_mode_propagates_radius() {
1205 let opts = DenoiserOptions::builder()
1206 .mode(DenoisingMode::Temporal { radius: 3 })
1207 .build();
1208 let params = opts.to_nlm_params();
1209
1210 assert_eq!(params.temporal_radius, 3);
1211 }
1212
1213 #[test]
1214 fn prefilter_passthrough() {
1215 let opts = DenoiserOptions::builder()
1216 .algorithm(Algorithm::Nlmeans(NlmeansOptions {
1217 prefilter: PrefilterMode::Bilateral {
1218 sigma_s: 3.0,
1219 sigma_r: 0.02,
1220 },
1221 ..NlmeansOptions::default()
1222 }))
1223 .build();
1224 let params = opts.to_nlm_params();
1225
1226 assert!(matches!(params.prefilter, PrefilterMode::Bilateral { .. }));
1227 }
1228
1229 #[test]
1230 fn hq_unset_prefilter_defaults_to_none() {
1231 let opts = DenoiserOptions::builder()
1232 .algorithm(hq(HqParams::default()))
1233 .build();
1234 let params = opts.to_nlm_params();
1235
1236 assert!(matches!(params.prefilter, PrefilterMode::None));
1237 }
1238
1239 #[test]
1240 fn fast_unset_prefilter_defaults_to_none() {
1241 let opts = DenoiserOptions::builder()
1242 .algorithm(Algorithm::Nlmeans(NlmeansOptions::default()))
1243 .build();
1244 let params = opts.to_nlm_params();
1245
1246 assert!(matches!(params.prefilter, PrefilterMode::None));
1247 }
1248
1249 #[test]
1250 fn hq_unset_strength_defaults_to_hq_default_strength() {
1251 let opts = DenoiserOptions::builder()
1253 .algorithm(hq(HqParams::default()))
1254 .build();
1255 let params = opts.to_nlm_params();
1256
1257 let expected = hq_default_strength(ChannelMode::Yuv, 0);
1258 assert!((params.strength - expected).abs() < f32::EPSILON);
1259 }
1260
1261 #[test]
1262 fn hq_no_auto_strength_falls_back_to_the_legacy_absolute_default() {
1263 let opts = DenoiserOptions::builder()
1270 .algorithm(hq(HqParams {
1271 auto_strength: false,
1272 ..HqParams::default()
1273 }))
1274 .build();
1275 let params = opts.to_nlm_params();
1276
1277 let expected = NlmParams::default().strength;
1278 assert!(
1279 (params.strength - expected).abs() < f32::EPSILON,
1280 "expected the legacy absolute default {expected}, got {}, which looks like the \
1281 auto-strength multiplier table leaking through",
1282 params.strength
1283 );
1284 }
1285
1286 #[test]
1287 fn hq_luma_r4_uses_measured_table_value() {
1288 let opts = DenoiserOptions::builder()
1289 .channel_mode(ChannelMode::Luma)
1290 .mode(DenoisingMode::Temporal { radius: 4 })
1291 .algorithm(hq(HqParams::default()))
1292 .build();
1293 let params = opts.to_nlm_params();
1294
1295 assert!((params.strength - 0.35).abs() < f32::EPSILON);
1296 }
1297
1298 #[test]
1299 fn hq_chroma_r4_uses_measured_table_value() {
1300 let opts = DenoiserOptions::builder()
1301 .channel_mode(ChannelMode::Chroma)
1302 .mode(DenoisingMode::Temporal { radius: 4 })
1303 .algorithm(hq(HqParams::default()))
1304 .build();
1305 let params = opts.to_nlm_params();
1306
1307 assert!((params.strength - 0.70).abs() < f32::EPSILON);
1308 }
1309
1310 #[test]
1311 fn hq_yuv_r8_uses_measured_table_value() {
1312 let opts = DenoiserOptions::builder()
1313 .channel_mode(ChannelMode::Yuv)
1314 .mode(DenoisingMode::Temporal { radius: 8 })
1315 .algorithm(hq(HqParams::default()))
1316 .build();
1317 let params = opts.to_nlm_params();
1318
1319 assert!((params.strength - 0.30).abs() < f32::EPSILON);
1320 }
1321
1322 #[test]
1323 fn hq_spacial_mode_uses_radius_zero_table_values() {
1324 for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] {
1325 let opts = DenoiserOptions::builder()
1326 .channel_mode(channels)
1327 .mode(DenoisingMode::Spacial)
1328 .algorithm(hq(HqParams::default()))
1329 .build();
1330 let params = opts.to_nlm_params();
1331
1332 let expected = hq_default_strength(channels, 0);
1333 assert!(
1334 (params.strength - expected).abs() < f32::EPSILON,
1335 "for channels {channels:?} expected {expected}, got {}",
1336 params.strength
1337 );
1338 }
1339 }
1340
1341 #[test]
1342 fn hq_explicit_strength_wins_over_the_table_for_every_plane() {
1343 for channels in [ChannelMode::Luma, ChannelMode::Chroma, ChannelMode::Yuv] {
1344 let opts = DenoiserOptions::builder()
1345 .channel_mode(channels)
1346 .mode(DenoisingMode::Temporal { radius: 4 })
1347 .algorithm(Algorithm::NlmeansHq(NlmeansHqOptions {
1348 nlm: NlmeansOptions {
1349 tuning: NlmTuning {
1350 strength: Some(0.99),
1351 ..NlmTuning::default()
1352 },
1353 ..NlmeansOptions::default()
1354 },
1355 hq: HqParams::default(),
1356 }))
1357 .build();
1358 let params = opts.to_nlm_params();
1359
1360 assert!(
1361 (params.strength - 0.99).abs() < f32::EPSILON,
1362 "for channels {channels:?} the explicit strength was overridden by the table"
1363 );
1364 }
1365 }
1366
1367 #[test]
1368 fn fast_unset_strength_defaults_to_legacy_default() {
1369 let opts = DenoiserOptions::builder()
1370 .algorithm(Algorithm::Nlmeans(NlmeansOptions::default()))
1371 .build();
1372 let params = opts.to_nlm_params();
1373
1374 assert!((params.strength - 1.2).abs() < f32::EPSILON);
1375 }
1376
1377 #[test]
1378 fn nl4d_options_default_matches_nl4d_params_default() {
1379 let opts = Nl4dOptions::default();
1380 let params = crate::nl4d::Nl4dParams::default();
1381
1382 assert_eq!(opts.refine, params.refine);
1383 assert_eq!(opts.spatial_radius, params.spatial_radius);
1384 assert!((opts.c_min - params.c_min).abs() < f32::EPSILON);
1385 assert_eq!(opts.confidence_variance, params.confidence_variance);
1386 assert_eq!(opts.lambda_ht, None);
1393 assert!((params.lambda_ht - nl4d_default_lambda_ht(ChannelMode::Yuv)).abs() < f32::EPSILON);
1394 }
1395
1396 #[test]
1400 fn nl4d_builds_the_front_ends_hq_params_from_its_own_fields() {
1401 let opts = DenoiserOptions::builder()
1402 .mode(DenoisingMode::Temporal { radius: 2 })
1403 .algorithm(Algorithm::Nl4d(Nl4dOptions {
1404 sigma: Some(0.02),
1405 sigma_scale: 1.3,
1406 thsad_scale: 0.8,
1407 ..Nl4dOptions::default()
1408 }))
1409 .build();
1410 let params = opts.to_nlm_params();
1411
1412 let hq = params.hq.expect("nl4d always runs the hq front end");
1413 assert_eq!(hq.sigma_override, Some(0.02));
1414 assert!((hq.sigma_scale - 1.3).abs() < f32::EPSILON);
1415 assert!((hq.thsad_scale - 0.8).abs() < f32::EPSILON);
1416 assert!(
1417 hq.temporal_confidence,
1418 "the grouping kernel reads the confidence scores, so this cannot be off"
1419 );
1420 }
1421
1422 #[test]
1425 fn nl4d_reads_its_temporal_radius_from_the_denoising_mode() {
1426 for radius in [1u32, 4, 8] {
1427 let opts = DenoiserOptions::builder()
1428 .mode(DenoisingMode::Temporal { radius })
1429 .algorithm(Algorithm::Nl4d(Nl4dOptions::default()))
1430 .build();
1431
1432 assert_eq!(opts.to_nlm_params().temporal_radius, radius);
1433 }
1434 }
1435
1436 #[test]
1439 fn nl4d_never_builds_a_prefilter() {
1440 let opts = DenoiserOptions::builder()
1441 .mode(DenoisingMode::Temporal { radius: 2 })
1442 .algorithm(Algorithm::Nl4d(Nl4dOptions::default()))
1443 .build();
1444
1445 assert!(matches!(opts.to_nlm_params().prefilter, PrefilterMode::None));
1446 }
1447
1448 #[test]
1451 fn nl4d_leaves_the_nlm_weighting_knobs_at_their_defaults() {
1452 let defaults = NlmParams::default();
1453 let opts = DenoiserOptions::builder()
1454 .channel_mode(ChannelMode::Luma)
1455 .mode(DenoisingMode::Temporal { radius: 4 })
1456 .algorithm(Algorithm::Nl4d(Nl4dOptions::default()))
1457 .build();
1458 let params = opts.to_nlm_params();
1459
1460 assert!((params.strength - defaults.strength).abs() < f32::EPSILON);
1461 assert_eq!(params.search_radius, defaults.search_radius);
1462 assert_eq!(params.patch_radius, defaults.patch_radius);
1463 assert!((params.self_weight - defaults.self_weight).abs() < f32::EPSILON);
1464 }
1465
1466 #[test]
1469 fn nl4d_motion_search_becomes_an_active_mvtools_mode() {
1470 let opts = DenoiserOptions::builder()
1471 .mode(DenoisingMode::Temporal { radius: 2 })
1472 .algorithm(Algorithm::Nl4d(Nl4dOptions {
1473 motion: MotionSearch {
1474 blksize: 32,
1475 overlap: 16,
1476 search_radius: 6,
1477 pyramid_levels: 1,
1478 estimation: MotionEstimation::Direct,
1479 },
1480 ..Nl4dOptions::default()
1481 }))
1482 .build();
1483 let params = opts.to_nlm_params();
1484
1485 assert!(matches!(
1486 params.motion_compensation,
1487 MotionCompensationMode::Mvtools {
1488 blksize: 32,
1489 overlap: 16,
1490 search_radius: 6,
1491 pyramid_levels: 1,
1492 estimation: MotionEstimation::Direct,
1493 }
1494 ));
1495 }
1496
1497 #[test]
1498 fn nl4d_motion_search_defaults_match_the_front_ends_own_defaults() {
1499 let opts = DenoiserOptions::builder()
1500 .mode(DenoisingMode::Temporal { radius: 2 })
1501 .algorithm(Algorithm::Nl4d(Nl4dOptions::default()))
1502 .build();
1503 let params = opts.to_nlm_params();
1504
1505 assert_eq!(
1506 params.motion_compensation,
1507 crate::nl4d::Nl4dParams::default().nlm.motion_compensation
1508 );
1509 }
1510
1511 #[test]
1512 fn motion_compensation_passthrough() {
1513 let opts = DenoiserOptions::builder()
1514 .mode(DenoisingMode::Temporal { radius: 1 })
1515 .algorithm(Algorithm::Nlmeans(NlmeansOptions {
1516 motion_compensation: MotionCompensationMode::Mvtools {
1517 blksize: 16,
1518 overlap: 8,
1519 search_radius: 4,
1520 pyramid_levels: 2,
1521 estimation: MotionEstimation::Direct,
1522 },
1523 ..NlmeansOptions::default()
1524 }))
1525 .build();
1526 let params = opts.to_nlm_params();
1527
1528 assert!(matches!(
1529 params.motion_compensation,
1530 MotionCompensationMode::Mvtools {
1531 blksize: 16,
1532 overlap: 8,
1533 search_radius: 4,
1534 pyramid_levels: 2,
1535 ..
1536 }
1537 ));
1538 }
1539
1540 #[test]
1541 fn motion_compensation_defaults_to_none() {
1542 let opts = DenoiserOptions::builder().build();
1543 let params = opts.to_nlm_params();
1544 assert!(matches!(params.motion_compensation, MotionCompensationMode::None));
1545 }
1546
1547 #[test]
1548 fn nlm_tuning_overrides_individual_fields() {
1549 let defaults = NlmParams::default();
1550 let opts = DenoiserOptions::builder()
1551 .algorithm(fast_tuned(NlmTuning {
1552 search_radius: Some(7),
1553 patch_radius: None,
1554 strength: Some(2.5),
1555 self_weight: None,
1556 }))
1557 .build();
1558 let params = opts.to_nlm_params();
1559
1560 assert_eq!(params.search_radius, 7);
1561 assert_eq!(params.patch_radius, defaults.patch_radius);
1562 assert!((params.strength - 2.5).abs() < f32::EPSILON);
1563 assert!((params.self_weight - defaults.self_weight).abs() < f32::EPSILON);
1564 }
1565}
1566
1567#[cfg(all(test, feature = "vulkan"))]
1568mod tests {
1569 use super::*;
1570
1571 fn opts(mode: DenoisingMode) -> DenoiserOptions {
1572 DenoiserOptions::builder()
1573 .channel_mode(ChannelMode::Luma)
1574 .mode(mode)
1575 .build()
1576 }
1577
1578 fn frame(w: u32, h: u32) -> Vec<f32> {
1579 vec![0.5f32; (w * h) as usize]
1580 }
1581
1582 #[test]
1583 fn spatial_denoise_roundtrip() {
1584 let mut d = Denoiser::create(
1585 &[Accelerator::Vulkan],
1586 &Device::Default,
1587 16,
1588 16,
1589 opts(DenoisingMode::Spacial),
1590 )
1591 .expect("denoiser construction failed");
1592 assert_eq!(d.selected_accelerator(), Accelerator::Vulkan);
1593
1594 d.push_frame(&frame(16, 16)).expect("push failed");
1595 let out = d.recv_frame().expect("recv failed").expect("no frame");
1596 assert_eq!(out.len(), 16 * 16);
1597 }
1598
1599 #[test]
1600 fn nl4d_algorithm_round_trips_through_the_facade() {
1601 let opts = DenoiserOptions::builder()
1602 .channel_mode(ChannelMode::Luma)
1603 .mode(DenoisingMode::Temporal { radius: 2 })
1604 .algorithm(Algorithm::Nl4d(Nl4dOptions::default()))
1605 .build();
1606 let mut d = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts)
1607 .expect("nl4d denoiser construction failed");
1608 assert_eq!(d.selected_accelerator(), Accelerator::Vulkan);
1609
1610 d.push_frame(&frame(16, 16)).expect("push failed");
1614 assert!(d.recv_frame().expect("recv failed").is_none());
1615
1616 let mut out = Vec::new();
1617 d.flush(|f| out.push(f)).expect("flush failed");
1618 assert_eq!(out.len(), 1, "expected exactly one output for one pushed frame");
1619 assert_eq!(out[0].len(), 16 * 16);
1620 }
1621
1622 #[test]
1626 fn nl4d_rejects_a_spatial_denoising_mode() {
1627 let opts = DenoiserOptions::builder()
1628 .channel_mode(ChannelMode::Luma)
1629 .mode(DenoisingMode::Spacial)
1630 .algorithm(Algorithm::Nl4d(Nl4dOptions::default()))
1631 .build();
1632 let result = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts);
1633
1634 match result {
1635 Err(DenoiserError::Other(e)) => assert!(
1636 e.to_string().contains("temporal window"),
1637 "unexpected error message: {e}"
1638 ),
1639 Err(other) => panic!("expected DenoiserError::Other, got {other:?}"),
1640 Ok(_) => panic!("expected a rejection, got Ok"),
1641 }
1642 }
1643
1644 #[test]
1647 fn window_span_is_symmetric_for_nlmeans() {
1648 let opts = DenoiserOptions::builder()
1649 .channel_mode(ChannelMode::Luma)
1650 .mode(DenoisingMode::Temporal { radius: 3 })
1651 .algorithm(Algorithm::Nlmeans(NlmeansOptions::default()))
1652 .build();
1653 let d = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts)
1654 .expect("denoiser construction failed");
1655
1656 let span = d.window_span();
1657 assert_eq!(span.behind, 3, "behind should equal the temporal radius");
1658 assert_eq!(span.ahead, 3, "ahead should equal the temporal radius");
1659 }
1660
1661 #[test]
1665 fn window_span_is_doubled_on_both_sides_for_nl4d() {
1666 let opts = DenoiserOptions::builder()
1667 .channel_mode(ChannelMode::Luma)
1668 .mode(DenoisingMode::Temporal { radius: 3 })
1669 .algorithm(Algorithm::Nl4d(Nl4dOptions::default()))
1670 .build();
1671 let d = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, opts)
1672 .expect("nl4d denoiser construction failed");
1673
1674 let span = d.window_span();
1675 assert_eq!(span.behind, 6, "behind should equal 2 * the temporal radius");
1676 assert_eq!(span.ahead, 6, "ahead should equal 2 * the temporal radius");
1677 }
1678
1679 #[test]
1680 fn invalid_params_surface_as_error() {
1681 let bad = DenoiserOptions::builder()
1682 .algorithm(Algorithm::Nlmeans(NlmeansOptions {
1683 tuning: NlmTuning {
1684 strength: Some(0.0),
1685 ..NlmTuning::default()
1686 },
1687 ..NlmeansOptions::default()
1688 }))
1689 .build();
1690 let result = Denoiser::create(&[Accelerator::Vulkan], &Device::Default, 16, 16, bad);
1691
1692 match result {
1693 Err(DenoiserError::Other(_)) => {},
1694 Err(other) => panic!("expected DenoiserError::Other, got {other:?}"),
1695 Ok(_) => panic!("expected validation error, got Ok"),
1696 }
1697 }
1698
1699 #[test]
1700 fn tiny_frame_dimensions_surface_as_error() {
1701 let result = Denoiser::create(
1702 &[Accelerator::Vulkan],
1703 &Device::Default,
1704 2,
1705 2,
1706 opts(DenoisingMode::Spacial),
1707 );
1708
1709 match result {
1710 Err(DenoiserError::Other(e)) => {
1711 assert!(
1712 e.to_string().contains("supported minimum"),
1713 "unexpected error message: {e}"
1714 );
1715 },
1716 Err(other) => panic!("expected DenoiserError::Other, got {other:?}"),
1717 Ok(_) => panic!("expected dimension validation error, got Ok"),
1718 }
1719 }
1720
1721 #[test]
1722 fn push_after_pending_returns_queue_full() {
1723 let mut d = Denoiser::create(
1724 &[Accelerator::Vulkan],
1725 &Device::Default,
1726 16,
1727 16,
1728 opts(DenoisingMode::Spacial),
1729 )
1730 .unwrap();
1731
1732 d.push_frame(&frame(16, 16)).unwrap();
1737 d.push_frame(&frame(16, 16)).unwrap();
1738 let err = d.push_frame(&frame(16, 16)).expect_err("expected QueueFull");
1739 assert!(matches!(err, DenoiserError::QueueFull));
1740
1741 let out = d.recv_frame().unwrap().unwrap();
1742 assert_eq!(out.len(), 16 * 16);
1743
1744 d.push_frame(&frame(16, 16)).expect("push after drain failed");
1746 }
1747
1748 fn frame_filled(w: u32, h: u32, value: f32) -> Vec<f32> {
1749 vec![value; (w * h) as usize]
1750 }
1751
1752 fn push_n_with_drain(d: &mut Denoiser, n: usize, value: f32, out: &mut Vec<Vec<f32>>) {
1755 for _ in 0..n {
1756 loop {
1757 match d.push_frame(&frame_filled(16, 16, value)) {
1758 Ok(()) => break,
1759 Err(DenoiserError::QueueFull) => {
1760 let f = d
1761 .recv_frame()
1762 .expect("recv ok")
1763 .expect("queue full but recv yielded none");
1764 out.push(f);
1765 },
1766 Err(e) => panic!("unexpected push error: {e:?}"),
1767 }
1768 }
1769 }
1770 }
1771
1772 #[test]
1773 fn flush_leaves_denoiser_reusable_spatial() {
1774 let mut d = Denoiser::create(
1775 &[Accelerator::Vulkan],
1776 &Device::Default,
1777 16,
1778 16,
1779 opts(DenoisingMode::Spacial),
1780 )
1781 .unwrap();
1782
1783 let mut batch_a = Vec::new();
1784 push_n_with_drain(&mut d, 5, 0.25, &mut batch_a);
1785 d.flush(|f| batch_a.push(f)).expect("first flush failed");
1786 assert_eq!(batch_a.len(), 5);
1787
1788 assert!(d.recv_frame().unwrap().is_none());
1790
1791 let mut batch_b = Vec::new();
1792 push_n_with_drain(&mut d, 5, 0.75, &mut batch_b);
1793 d.flush(|f| batch_b.push(f)).expect("second flush failed");
1794 assert_eq!(batch_b.len(), 5);
1795
1796 for v in batch_b.iter().flatten() {
1797 assert!((v - 0.75).abs() < 0.1, "batch_b carried state from batch_a: {v}");
1798 }
1799 for v in batch_a.iter().flatten() {
1800 assert!((v - 0.25).abs() < 0.1, "batch_a value unexpectedly drifted: {v}");
1801 }
1802 }
1803
1804 #[test]
1805 fn flush_leaves_denoiser_reusable_temporal() {
1806 let mut d = Denoiser::create(
1807 &[Accelerator::Vulkan],
1808 &Device::Default,
1809 16,
1810 16,
1811 opts(DenoisingMode::Temporal { radius: 1 }),
1812 )
1813 .unwrap();
1814
1815 let mut batch_a = Vec::new();
1816 push_n_with_drain(&mut d, 5, 0.25, &mut batch_a);
1817 d.flush(|f| batch_a.push(f)).expect("first flush failed");
1818 assert_eq!(batch_a.len(), 5, "expected 5 frames from first batch");
1819
1820 assert!(d.recv_frame().unwrap().is_none());
1825 d.push_frame(&frame_filled(16, 16, 0.75)).unwrap();
1826 assert!(
1827 d.recv_frame().unwrap().is_none(),
1828 "first push of new temporal stream should not produce output yet"
1829 );
1830
1831 let mut batch_b = Vec::new();
1833 push_n_with_drain(&mut d, 4, 0.75, &mut batch_b);
1834 d.flush(|f| batch_b.push(f)).expect("second flush failed");
1835 assert_eq!(batch_b.len(), 5, "expected 5 frames from second batch");
1836
1837 for v in batch_b.iter().flatten() {
1838 assert!((v - 0.75).abs() < 0.1, "batch_b carried state from batch_a: {v}");
1839 }
1840 }
1841
1842 #[test]
1843 fn flush_emits_exactly_n_outputs_for_small_n() {
1844 for n in 1..=5usize {
1849 let mut d = Denoiser::create(
1850 &[Accelerator::Vulkan],
1851 &Device::Default,
1852 16,
1853 16,
1854 opts(DenoisingMode::Temporal { radius: 2 }),
1855 )
1856 .unwrap();
1857
1858 let mut out = Vec::new();
1859 push_n_with_drain(&mut d, n, 0.5, &mut out);
1860 d.flush(|f| out.push(f)).expect("flush failed");
1861 assert_eq!(
1862 out.len(),
1863 n,
1864 "expected {n} outputs for {n} pushes, got {}",
1865 out.len()
1866 );
1867 }
1868 }
1869}