#![allow(dead_code)]
use cubecl::prelude::*;
#[cube(launch_unchecked)]
fn add_one_kernel(input: &Array<f32>, output: &mut Array<f32>) {
if ABSOLUTE_POS < input.len() {
output[ABSOLUTE_POS] = input[ABSOLUTE_POS] + 1.0;
}
}
#[cube(launch_unchecked)]
fn smem_mirror_kernel(input: &Array<f32>, output: &mut Array<f32>) {
let mut staged = SharedMemory::<f32>::new(256usize);
let unit = UNIT_POS as usize;
let cdim = CUBE_DIM as usize;
if ABSOLUTE_POS < input.len() {
staged[unit] = input[ABSOLUTE_POS];
}
sync_cube();
let mirror = unit ^ (cdim - 1);
let src = CUBE_POS * cdim + mirror;
if ABSOLUTE_POS < output.len() && src < input.len() {
output[ABSOLUTE_POS] = staged[mirror];
}
}
#[cfg(test)]
mod tests {
use super::*;
use cubecl::wgpu::WgpuRuntime;
#[test]
fn add_one_runs_on_gpu() {
if crate::skip_no_gpu() {
return;
}
let device = Default::default();
let client = WgpuRuntime::client(&device);
let input = [1.0f32, 2.0, 3.0, 4.0];
let n = input.len();
let in_h = client.create_from_slice(f32::as_bytes(&input));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
add_one_kernel::launch_unchecked::<WgpuRuntime>(
&client,
CubeCount::Static(1, 1, 1),
CubeDim::new_1d(n as u32),
ArrayArg::from_raw_parts(in_h, n),
ArrayArg::from_raw_parts(out_h.clone(), n),
);
}
let bytes = client.read_one_unchecked(out_h);
let got = f32::from_bytes(&bytes);
assert_eq!(got, &[2.0, 3.0, 4.0, 5.0]);
}
#[test]
fn shared_memory_mirror_runs_on_gpu() {
if crate::skip_no_gpu() {
return;
}
let device = Default::default();
let client = WgpuRuntime::client(&device);
let n = 512usize; let input: Vec<f32> = (0..n).map(|i| i as f32).collect();
let in_h = client.create_from_slice(f32::as_bytes(&input));
let out_h = client.empty(n * core::mem::size_of::<f32>());
unsafe {
smem_mirror_kernel::launch_unchecked::<WgpuRuntime>(
&client,
CubeCount::Static(2, 1, 1),
CubeDim::new_1d(256),
ArrayArg::from_raw_parts(in_h, n),
ArrayArg::from_raw_parts(out_h.clone(), n),
);
}
let bytes = client.read_one_unchecked(out_h);
let got = f32::from_bytes(&bytes);
for cube in 0..2usize {
for j in 0..256usize {
let want = input[cube * 256 + (255 - j)];
assert_eq!(
got[cube * 256 + j],
want,
"cube {cube} lane {j}: shared-memory staging or barrier is broken"
);
}
}
}
}