use cubecl::{prelude::*, std::tensor::TensorHandle, zspace::Shape};
use cubek_convolution::{
ConvolutionArgs, DirectTensors, eval::cpu_reference::ConvSpec, launch_direct,
};
use cubek_test_utils::{HostData, HostDataType, TestInput};
#[derive(Debug, Clone, Copy)]
struct Case {
b: usize,
c_in: usize,
c_out: usize,
groups: usize,
size: usize,
kh: usize,
kw: usize,
stride: usize,
dilation: usize,
padding: usize,
bias: bool,
}
impl Case {
fn ramp(n: usize, period: usize) -> Vec<f32> {
(0..n).map(|i| ((i % period) as f32) - 1.0).collect()
}
fn check(&self, dtype: ElemType) -> Result<(), String> {
let client = cubecl::test_device().client();
let args = ConvolutionArgs::<2> {
stride: [self.stride; 2],
padding: [self.padding; 2],
dilation: [self.dilation; 2],
};
let spec = ConvSpec {
batches: self.b,
in_h: self.size,
in_w: self.size,
channels: self.c_in,
out_channels: self.c_out,
args: args.clone(),
kernel_size: [self.kh, self.kw],
};
let in_shape = [self.b, self.size, self.size, self.c_in];
let w_shape = [self.c_out, self.kh, self.kw, self.c_in / self.groups];
let out_shape = [self.b, spec.out_h(), spec.out_w(), self.c_out];
let (input, input_host) = TestInput::builder(client.clone(), Shape::new(in_shape))
.dtype(dtype)
.custom(Self::ramp(in_shape.iter().product(), 4))
.generate_with_f32_host_data();
let (weight, weight_host) = TestInput::builder(client.clone(), Shape::new(w_shape))
.dtype(dtype)
.custom(Self::ramp(w_shape.iter().product(), 3))
.generate_with_f32_host_data();
let (bias, bias_host) = self
.bias
.then(|| {
TestInput::builder(client.clone(), Shape::new([self.c_out]))
.dtype(dtype)
.custom(Self::ramp(self.c_out, 5))
.generate_with_f32_host_data()
})
.unzip();
let out: TensorHandle = TestInput::builder(client.clone(), Shape::new(out_shape))
.dtype(dtype)
.custom(vec![4096.0; out_shape.iter().product::<usize>()])
.generate_without_host_data();
launch_direct::<2>(
&client,
DirectTensors {
input: input.binding(),
weight: weight.binding(),
bias: bias.map(|bias| bias.binding()),
out: out.clone().binding(),
},
args,
self.groups,
dtype,
)
.map_err(|e| format!("setup: {e:?}"))?;
let got = HostData::from_tensor_handle(&client, out, HostDataType::F32);
let want = spec.cpu_reference(&input_host, &weight_host, bias_host.as_ref());
for b in 0..self.b {
for oh in 0..out_shape[1] {
for ow in 0..out_shape[2] {
for oc in 0..self.c_out {
let (g, w) = (
got.get_f32(&[b, oh, ow, oc]),
want.get_f32(&[b, oh, ow, oc]),
);
if g != w {
return Err(format!("({b},{oh},{ow},{oc}) got {g} want {w}"));
}
}
}
}
}
Ok(())
}
}
const CHANNELS: [(usize, usize, usize); 11] = [
(3, 8, 1),
(16, 16, 1),
(32, 32, 1),
(64, 64, 1),
(64, 12, 1),
(64, 6, 1),
(64, 5, 1),
(128, 32, 2),
(32, 64, 2),
(192, 48, 3),
(104, 26, 1),
];
const GEOMETRIES: [(usize, usize, usize, usize, usize); 6] = [
(3, 3, 1, 1, 1),
(3, 3, 2, 1, 1),
(3, 3, 1, 2, 2),
(3, 3, 1, 1, 0),
(1, 1, 1, 1, 0),
(1, 3, 1, 1, 1),
];
fn check_grid(dtype: ElemType) {
let mut bad = Vec::new();
for (c_in, c_out, groups) in CHANNELS {
for (kh, kw, stride, dilation, padding) in GEOMETRIES {
for bias in [false, true] {
let case = Case {
b: 2,
c_in,
c_out,
groups,
size: 6,
kh,
kw,
stride,
dilation,
padding,
bias,
};
if let Err(e) = case.check(dtype) {
bad.push(format!("{case:?}: {e}"));
}
}
}
}
assert!(bad.is_empty(), "failing cases:\n{}", bad.join("\n"));
}
#[test]
fn direct_f32_matches_the_reference() {
check_grid(f32::elem_type_native());
}
#[test]
fn direct_f16_matches_the_reference() {
check_grid(half::f16::elem_type_native());
}