use cubecl::prelude::*;
use crate::collab::{MAX_K, PATCH_AREA, PATCH_SIZE};
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"
);
pub const RECIPROCAL_FLOOR: f32 = 1e-12;
#[cube]
pub(crate) fn safe_reciprocal(denom: f32, floor: f32) -> f32 {
let mut inv = 0.0f32;
if !denom.is_nan() && !denom.is_inf() {
inv = 1.0f32 / f32::max(denom, floor);
}
inv
}
#[cube]
pub(crate) fn fill_dct8_basis(basis: &mut SharedMemory<f32>, 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;
if j == 0 {
c = 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);
}
}
#[cube]
pub(crate) fn dct8_reg_fwd(basis: &SharedMemory<f32>, line: &mut Array<f32>) {
let mut src = Array::<f32>::new(8usize);
#[unroll]
for i in 0..PATCH_SIZE {
src[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];
}
line[j as usize] = sum;
}
}
#[cube]
pub(crate) fn dct8_reg_inv(basis: &SharedMemory<f32>, line: &mut Array<f32>) {
let mut src = Array::<f32>::new(8usize);
#[unroll]
for j in 0..PATCH_SIZE {
src[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];
}
line[i as usize] = sum;
}
}
#[cube]
pub(crate) fn haar_reg_fwd_level(stack: &mut Array<f32>, #[comptime] len: u32) {
let half = comptime!(len / 2);
#[unroll]
for pos in 0..PATCH_SIZE {
let mut snapshot = Array::<f32>::new(MAX_K as usize);
#[unroll]
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;
}
}
}
#[cube]
pub(crate) fn haar_reg_inv_level(stack: &mut Array<f32>, #[comptime] len: u32) {
let half = comptime!(len / 2);
#[unroll]
for pos in 0..PATCH_SIZE {
let mut snapshot = Array::<f32>::new(MAX_K as usize);
#[unroll]
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;
stack[((2u32 * p + 1u32) * PATCH_SIZE + pos) as usize] =
(a - b) * std::f32::consts::FRAC_1_SQRT_2;
}
}
}
#[cube]
pub(crate) fn variance_reg_level(v: &mut Array<f32>, #[comptime] len: u32) {
let half = comptime!(len / 2);
let mut snapshot = Array::<f32>::new(MAX_K as usize);
#[unroll]
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;
}
}
pub fn dct_noise_profile(rho: f32) -> [f32; 8] {
if rho <= 0.0 {
return [1.0; 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 };
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();
}
}
let mut g = [0.0f32; 8];
for (u, slot) in g.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());
}
}
*slot = sum as f32;
}
g
}
#[cfg(all(test, any(feature = "vulkan", feature = "metal")))]
pub(crate) fn haar_variance_ladder(sig2: &[f32], k_use: u32) -> Vec<f32> {
let mut out = sig2.to_vec();
let mut len = k_use;
while len > 1 {
let half = len / 2;
let snapshot = out[..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;
}
len = half;
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dct_noise_profile_rho_zero_is_uniform_identity() {
let g = dct_noise_profile(0.0);
assert_eq!(
g, [1.0f32; 8],
"rho=0 must give exactly 1.0 at every frequency, got {g:?}"
);
let g_neg = dct_noise_profile(-0.1);
assert_eq!(g_neg, [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();
assert!(
(sum - 8.0).abs() < 1e-3,
"rho={rho}: expected sum(g) == 8.0 (variance redistributed, not created or \
destroyed), got {sum}"
);
}
}
#[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);
for u in 0..7 {
assert!(
g[u] > g[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],
u + 1,
g[u + 1],
);
}
}
}
#[test]
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}");
}
}
}
#[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);
}
#[test]
fn k_use_of_one_is_the_identity() {
let out = haar_variance_ladder(&[0.7], 1);
assert_eq!(out, vec![0.7]);
}
#[test]
fn eight_element_ladder_matches_hand_computed_levels() {
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 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() {
assert!((got - want).abs() < 1e-6, "idx={idx}: got {got} want {want}");
}
}
}