use cubecl::prelude::Scalar;
use cubek_matmul::{
definition::{MatmulElems, MatmulGlobalElems},
definition::{MatmulPrecision, MatrixPrecision},
};
use cubek_test_utils::{HostData, Progress};
use crate::{
Strategy,
eval::{
benchmarks::problem::Conv2dProblem,
cpu_reference::{ConvSpec, cpu_reference_result, strategy_result},
},
};
type LhsG<MP> = <<MP as MatmulPrecision>::Lhs as MatrixPrecision>::Global;
type RhsG<MP> = <<MP as MatmulPrecision>::Rhs as MatrixPrecision>::Global;
type AccG<MP> = <<MP as MatmulPrecision>::Acc as MatrixPrecision>::Global;
pub struct Conv2dCorrectness;
impl cubek_test_utils::Correctness for Conv2dCorrectness {
type Problem = Conv2dProblem;
type Strategy = Strategy;
fn kernel_result(
&self,
strategy: &Strategy,
problem: &Conv2dProblem,
seeds: &[u64],
) -> Result<HostData, String> {
let device = cubecl::test_device();
let client = device.client();
let (spec, dtypes) = build_spec_and_dtypes::<half::f16>(problem);
strategy_result(client, spec, strategy.clone(), dtypes, seeds[0], seeds[1])
}
fn reference_result(
&self,
problem: &Conv2dProblem,
seeds: &[u64],
progress: Option<&Progress>,
) -> Result<HostData, String> {
let device = cubecl::test_device();
let client = device.client();
let (spec, dtypes) = build_spec_and_dtypes::<half::f16>(problem);
cpu_reference_result(client, spec, dtypes, seeds[0], seeds[1], progress)
}
}
fn build_spec_and_dtypes<MP: MatmulPrecision>(p: &Conv2dProblem) -> (ConvSpec, MatmulElems) {
let [n, c_in, h_in, w_in] = p.input_shape;
let [c_out, _, k_h, k_w] = p.weight_shape;
let dtypes = MatmulElems::from_globals(&MatmulGlobalElems {
lhs: LhsG::<MP>::elem_type_native(),
rhs: RhsG::<MP>::elem_type_native(),
out: AccG::<MP>::elem_type_native(),
});
let spec = ConvSpec {
batches: n,
in_h: h_in,
in_w: w_in,
channels: c_in,
out_channels: c_out,
args: p.args.clone(),
kernel_size: [k_h, k_w],
};
(spec, dtypes)
}