use cubecl::prelude::*;
use cubecl::server::Handle;
use super::params::Nl4dParams;
use crate::collab::geometry::{fused_cubes_x, ref_count, refs_along};
use crate::collab::kernels::aggregate::{
collab_normalise,
collab_zero_accum,
cross_frame_accum_scale,
weight_scale,
};
use crate::collab::kernels::fused::collab_fused;
use crate::collab::kernels::transforms::dct_noise_profile;
use crate::collab::{MAX_K, PATCH_SIZE};
use crate::denoiser::DenoiserError;
use crate::nlmeans::{BLOCK_X, BLOCK_Y, ChannelMode, MAX_GRID_1D, NlmDenoiser, Pending, RingView};
pub struct Nl4dDenoiser<R: Runtime> {
front: NlmDenoiser<R>,
width: u32,
height: u32,
channels: ChannelMode,
temporal_radius: u32,
refine: u32,
spatial_radius: u32,
lambda_ht: f32,
c_min: f32,
mismatch_scale: f32,
confidence_variance: bool,
k_max: u32,
accum_scale: f32,
group_weight: Handle,
sigma_buf: Handle,
dct_profile_buf: Handle,
dct_profile: [f32; 8],
accum: Handle,
wsum: Handle,
outputs: [Handle; 2],
next_output_slot: usize,
passes_run: u32,
}
impl<R: Runtime> Nl4dDenoiser<R> {
pub fn new(
client: &ComputeClient<R>,
mut params: Nl4dParams,
width: u32,
height: u32,
) -> Result<Self, String> {
params.validate()?;
if width < PATCH_SIZE || height < PATCH_SIZE {
return Err(format!(
"frame dimensions {width}x{height} must be at least {p}x{p} for the \
collaborative filter's patch grid",
p = PATCH_SIZE,
));
}
params.nlm.temporal_radius = params.temporal_radius;
params.nlm.validate().map_err(|e| e.to_string())?;
let front = NlmDenoiser::new(client, params.nlm.clone(), width, height);
let channels = params.nlm.channels;
let stored_ch = channels.storage_count();
let k_max = MAX_K;
let refs = ref_count(width, height);
let pixels = (width * height) as usize;
let frame_len = pixels * stored_ch as usize;
let group_weight = client.empty(refs * size_of::<f32>());
let sigma_buf = client.create_from_slice(f32::as_bytes(&vec![0.0f32; stored_ch as usize]));
let dct_profile = dct_noise_profile(0.0);
let dct_profile_buf = client.create_from_slice(f32::as_bytes(&dct_profile));
let ring_frames = 1 + 2 * params.temporal_radius;
let accum = client.empty(frame_len * ring_frames as usize * size_of::<i32>());
let wsum = client.empty(pixels * ring_frames as usize * size_of::<i32>());
let outputs = [
client.empty(frame_len * size_of::<f32>()),
client.empty(frame_len * size_of::<f32>()),
];
Ok(Self {
front,
width,
height,
channels,
temporal_radius: params.temporal_radius,
refine: params.refine,
spatial_radius: params.spatial_radius,
lambda_ht: params.lambda_ht,
c_min: params.c_min,
mismatch_scale: params.mismatch_scale,
confidence_variance: params.confidence_variance,
k_max,
accum_scale: cross_frame_accum_scale(params.spatial_radius, params.temporal_radius),
group_weight,
sigma_buf,
dct_profile_buf,
dct_profile,
accum,
wsum,
outputs,
next_output_slot: 0,
passes_run: 0,
})
}
pub fn push_frame(&mut self, frame: &[f32]) {
self.front.push_frame(frame);
}
pub fn denoise_submit(&mut self) -> Result<Option<Pending<R>>, DenoiserError> {
let Some(view) = self.front.submit_machinery()? else {
return Ok(None);
};
let Some(handle) = self.run_collab_stage(&view)? else {
return Ok(None);
};
Ok(Some(self.start_readback(handle)))
}
pub fn denoise(&mut self) -> Result<Option<Vec<f32>>, DenoiserError> {
let Some(pending) = self.denoise_submit()? else {
return Ok(None);
};
Ok(Some(pending.wait()?))
}
pub fn flush(&mut self, mut sink: impl FnMut(&[f32])) -> Result<(), DenoiserError> {
let target = self.flush_target();
let mut emitted = 0usize;
while emitted < target {
if let Some(view) = self.front.flush_step_machinery()?
&& let Some(handle) = self.run_collab_stage(&view)?
{
let pending = self.start_readback(handle);
let frame = pending.wait()?;
sink(&frame);
emitted += 1;
}
}
self.front.reset_stream_state();
self.next_output_slot = 0;
self.passes_run = 0;
Ok(())
}
fn flush_target(&self) -> usize {
let real_pushes = self.front.real_pushes();
if real_pushes == 0 {
0
} else {
real_pushes.min(2 * self.temporal_radius as usize)
}
}
fn run_collab_stage(&mut self, view: &RingView) -> Result<Option<Handle>, DenoiserError> {
let centre_slot = view.centre_slot;
let client = self.front.compute_client().clone();
let stored_ch = self.channels.storage_count();
let channels_count = self.channels.count();
let pixels = (self.width * self.height) as usize;
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;
let accum_ring_len = frame_len * total_frames as usize;
let wsum_ring_len = pixels * total_frames as usize;
let neighbours = 2 * self.temporal_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 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 refs_x = refs_along(self.width);
let refs_y = refs_along(self.height);
let refs = ref_count(self.width, self.height);
let collab_grid = CubeCount::new_2d(fused_cubes_x(self.width), refs_y);
let collab_dim = CubeDim::new_1d(64);
let agg_grid = CubeCount::new_2d(self.width.div_ceil(BLOCK_X), self.height.div_ceil(BLOCK_Y));
let agg_dim = CubeDim::new_2d(BLOCK_X, BLOCK_Y);
let zero_dim = 256u32;
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 mismatch_thsad = self.front.thsad_value() * self.mismatch_scale;
let newest_slot = (centre_slot + self.temporal_radius) % total_frames;
let completed_slot = (centre_slot + total_frames - self.temporal_radius) % total_frames;
let pass_index = self.passes_run;
self.passes_run += 1;
unsafe {
if pass_index == 0 {
for slot in 0..total_frames {
collab_zero_accum::launch_unchecked::<R>(
&client,
zero_grid_one_frame.clone(),
CubeDim::new_1d(zero_dim),
ArrayArg::from_raw_parts(self.accum.clone(), accum_ring_len),
ArrayArg::from_raw_parts(self.wsum.clone(), wsum_ring_len),
slot * pixels as u32,
pixels as u32,
stored_ch,
zero_total_threads_one_frame,
);
}
} else {
collab_zero_accum::launch_unchecked::<R>(
&client,
zero_grid_one_frame,
CubeDim::new_1d(zero_dim),
ArrayArg::from_raw_parts(self.accum.clone(), accum_ring_len),
ArrayArg::from_raw_parts(self.wsum.clone(), wsum_ring_len),
newest_slot * pixels as u32,
pixels as u32,
stored_ch,
zero_total_threads_one_frame,
);
}
collab_fused::launch_unchecked::<R>(
&client,
collab_grid,
collab_dim,
stored_ch as usize,
ArrayArg::from_raw_parts(view.input.clone(), ring_len),
ArrayArg::from_raw_parts(view.mv_field.clone(), mv_len.max(1)),
ArrayArg::from_raw_parts(view.confidence.clone(), conf_len.max(1)),
ArrayArg::from_raw_parts(neighbour_slots_buf, view.neighbour_slots.len().max(1)),
ArrayArg::from_raw_parts(self.sigma_buf.clone(), stored_ch as usize),
ArrayArg::from_raw_parts(self.dct_profile_buf.clone(), 8),
ArrayArg::from_raw_parts(self.accum.clone(), accum_ring_len),
ArrayArg::from_raw_parts(self.wsum.clone(), wsum_ring_len),
ArrayArg::from_raw_parts(self.group_weight.clone(), refs),
centre_slot,
0.0f32,
self.c_min,
mismatch_thsad,
self.lambda_ht,
wnorm,
self.accum_scale,
self.confidence_variance,
self.temporal_radius,
self.refine,
view.mv_stride,
view.conf_stride,
blk_step,
blksize,
blocks_x,
blocks_y,
self.width,
self.height,
channels_count,
self.k_max,
stored_ch,
self.spatial_radius,
refs_x,
);
}
if pass_index < self.temporal_radius {
return Ok(None);
}
let slot = self.next_output_slot;
self.next_output_slot = (slot + 1) % self.outputs.len();
unsafe {
collab_normalise::launch_unchecked::<R>(
&client,
agg_grid,
agg_dim,
stored_ch as usize,
ArrayArg::from_raw_parts(self.accum.clone(), accum_ring_len),
ArrayArg::from_raw_parts(self.wsum.clone(), wsum_ring_len),
ArrayArg::from_raw_parts(self.outputs[slot].clone(), frame_len),
completed_slot * pixels as u32,
self.width,
self.height,
channels_count,
stored_ch,
);
}
Ok(Some(self.outputs[slot].clone()))
}
fn start_readback(&self, handle: Handle) -> Pending<R> {
let client = self.front.compute_client().clone();
let fut = Box::pin(async move { client.read_async(vec![handle]).await });
let pixels = (self.width * self.height) as usize;
Pending::new(fut, self.channels.count(), self.channels.storage_count(), pixels)
}
}