use cubecl::prelude::*;
#[cube(launch_unchecked)]
pub fn gpu_copy(
src: &Array<f32>,
dst: &mut Array<f32>,
src_offset: u32,
dst_offset: u32,
#[comptime] length: u32,
#[comptime] total_threads: u32,
) {
let mut idx = ABSOLUTE_POS_X;
while idx < length {
dst[(dst_offset + idx) as usize] = src[(src_offset + idx) as usize];
idx += total_threads;
}
}
#[cube(launch_unchecked)]
pub fn gpu_zero_buffers(
accum: &mut Array<f32>,
weight_sum: &mut Array<f32>,
max_weight: &mut Array<f32>,
#[comptime] accum_len: u32,
#[comptime] weight_len: u32,
#[comptime] total_threads: u32,
) {
let mut idx = ABSOLUTE_POS_X;
while idx < weight_len {
accum[idx as usize] = 0.0f32;
weight_sum[idx as usize] = 0.0f32;
max_weight[idx as usize] = 0.0f32;
idx += total_threads;
}
while idx < accum_len {
accum[idx as usize] = 0.0f32;
idx += total_threads;
}
}
#[cube(launch_unchecked)]
#[expect(
clippy::too_many_arguments,
reason = "every argument is a comptime shape the kernel specialises on"
)]
pub fn gpu_pack_wire(
src: &Array<f32>,
dst: &mut Array<u32>,
max: f32,
#[comptime] pixels: u32,
#[comptime] channels: u32,
#[comptime] stored_ch: u32,
#[comptime] outer: u32,
#[comptime] split_planes: bool,
#[comptime] samples_per_word: u32,
#[comptime] words: u32,
#[comptime] total_threads: u32,
) {
let samples = comptime![pixels * channels];
let bits = comptime![32u32 / samples_per_word];
let mut word = ABSOLUTE_POS_X;
while word < words {
let base = word * samples_per_word;
let mut acc = 0u32;
#[unroll]
for lane in 0..samples_per_word {
let s = base + lane;
let safe = u32::min(s, samples - 1);
let a = safe / outer;
let b = safe % outer;
let src_idx = select(split_planes, b * stored_ch + a, a * stored_ch + b);
let v = f32::clamp(src[src_idx as usize], 0.0, 1.0);
let q = u32::cast_from(v * max + 0.5);
acc |= select(s < samples, q, 0u32) << (lane * bits);
}
dst[word as usize] = acc;
word += total_threads;
}
}
#[cube(launch_unchecked)]
#[expect(
clippy::too_many_arguments,
reason = "every argument is a comptime shape the kernel specialises on"
)]
pub fn gpu_unpack_wire(
src: &Array<u32>,
dst: &mut Array<f32>,
max: f32,
dst_offset: u32,
#[comptime] pixels: u32,
#[comptime] channels: u32,
#[comptime] stored_ch: u32,
#[comptime] samples_per_word: u32,
#[comptime] elements: u32,
#[comptime] total_threads: u32,
) {
let bits = comptime![32u32 / samples_per_word];
let mask = comptime![(1u32 << (32u32 / samples_per_word)) - 1];
let wire_samples = comptime![pixels * channels];
let mut idx = ABSOLUTE_POS_X;
while idx < elements {
let pixel = idx / stored_ch;
let ch = idx % stored_ch;
let s = u32::min(ch * pixels + pixel, wire_samples - 1);
let word = src[(s / samples_per_word) as usize];
let sample = (word >> ((s % samples_per_word) * bits)) & mask;
let v = f32::cast_from(sample) / max;
dst[(dst_offset + idx) as usize] = select(ch < channels, v, 0.0f32);
idx += total_threads;
}
}