use std::marker::PhantomData;
use cubecl::{
benchmark::{Benchmark, ProfileDuration, TimingMethod},
client::Client,
future,
prelude::*,
std::tensor::TensorHandle,
zspace::Shape,
};
use cubek_matmul::{
definition::MatmulElems,
definition::{MatmulPrecision, MatrixPrecision},
};
use cubek_std::InputBinding;
use cubek_test_utils::{RunSamples, TestInput};
use crate::{ConvolutionInputs, Strategy, eval::benchmarks::problem::Conv2dProblem, launch_ref};
type LhsG<MP> = <<MP as MatmulPrecision>::Lhs as MatrixPrecision>::Global;
type LhsS<MP> = <<MP as MatmulPrecision>::Lhs as MatrixPrecision>::Stage;
type RhsG<MP> = <<MP as MatmulPrecision>::Rhs as MatrixPrecision>::Global;
type AccG<MP> = <<MP as MatmulPrecision>::Acc as MatrixPrecision>::Global;
type AccR<MP> = <<MP as MatmulPrecision>::Acc as MatrixPrecision>::Register;
pub fn bench(
strategy: &Strategy,
problem: &Conv2dProblem,
num_samples: usize,
) -> Result<RunSamples, String> {
let device = cubecl::test_device();
let client = device.client();
let bench = Conv2dBench::<half::f16> {
problem: problem.clone(),
strategy: strategy.clone(),
device,
client,
samples: num_samples,
_phantom: PhantomData,
};
let durations = bench
.run(cubek_test_utils::timing_method(TimingMethod::System))
.map_err(|e| format!("benchmark failed: {e}"))?
.durations;
Ok(RunSamples::new(durations))
}
struct Conv2dBench<MP> {
problem: Conv2dProblem,
strategy: Strategy,
device: cubecl::Device,
client: Client,
samples: usize,
_phantom: PhantomData<MP>,
}
fn make_uniform_4d(client: &Client, shape: [usize; 4], dtype: ElemType, seed: u64) -> TensorHandle {
TestInput::builder(client.clone(), Shape::new(shape))
.dtype(dtype)
.uniform(seed, 0.0, 1.0)
.generate_without_host_data()
}
impl<MP: MatmulPrecision> Benchmark for Conv2dBench<MP> {
type Input = (TensorHandle, TensorHandle, TensorHandle);
type Output = ();
fn prepare(&self) -> Self::Input {
let client = self.device.client();
let input = make_uniform_4d(
&client,
self.problem.input_shape,
LhsG::<MP>::elem_type_native(),
0,
);
let weight = make_uniform_4d(
&client,
self.problem.weight_shape,
RhsG::<MP>::elem_type_native(),
1,
);
let bias = TestInput::builder(client.clone(), Shape::from(vec![self.problem.bias_shape]))
.dtype(AccG::<MP>::elem_type_native())
.layout(cubek_test_utils::StridedLayout::Explicit(vec![1]))
.uniform(2, 0.0, 1.0)
.generate_without_host_data();
(input, weight, bias)
}
fn execute(&self, (input, weight, bias): Self::Input) -> Result<(), String> {
let client = self.device.client();
let [n, _, h_in, w_in] = self.problem.input_shape;
let [c_out, _, k_h, k_w] = self.problem.weight_shape;
let [s_h, s_w] = self.problem.args.stride;
let [p_h, p_w] = self.problem.args.padding;
let [d_h, d_w] = self.problem.args.dilation;
let h_out = (h_in + 2 * p_h - d_h * (k_h - 1) - 1) / s_h + 1;
let w_out = (w_in + 2 * p_w - d_w * (k_w - 1) - 1) / s_w + 1;
let elems = MatmulElems::new_deprecated::<MP>();
let out: TensorHandle =
TensorHandle::empty(&client, vec![n, c_out, h_out, w_out], elems.acc_global);
launch_ref::<2>(
&self.strategy,
&self.client,
ConvolutionInputs::Forward {
input: InputBinding::Normal(input.binding(), elems.lhs_global),
weight: InputBinding::Normal(weight.binding(), elems.rhs_global),
bias: Some(InputBinding::Normal(bias.binding(), elems.acc_global)),
out: out.binding(),
},
self.problem.args.clone(),
elems,
)
.map_err(|it| format!("{it:?}"))?;
Ok(())
}
fn num_samples(&self) -> usize {
self.samples
}
fn name(&self) -> String {
let client = self.device.client();
format!(
"{}-conv2d-{}-{}-{}-{}",
client.name(),
LhsG::<MP>::elem_type_native(),
LhsS::<MP>::elem_type_native(),
AccR::<MP>::elem_type_native(),
AccG::<MP>::elem_type_native(),
)
.to_lowercase()
}
fn sync(&self) {
future::block_on(self.client.sync()).unwrap()
}
fn profile(&self, args: Self::Input) -> Result<ProfileDuration, String> {
cubek_test_utils::profile_launch(&self.client, "conv-bench", || self.execute(args))
}
}