use cubecl::{
client::Client,
ir::ElemType,
throughput::{ThroughputKey, ThroughputMode},
tune::Work,
};
use cubek_matmul::definition::{MatmulCost, MatmulGlobalElems};
#[derive(Debug, Clone)]
pub struct Conv2dCost {
pub batch: usize,
pub channels_in: usize,
pub spatial_in: [usize; 2],
pub channels_out: usize,
pub kernel: [usize; 2],
pub spatial_out: [usize; 2],
pub bias_elems: usize,
pub elems: MatmulGlobalElems,
}
impl Conv2dCost {
fn gemm(&self) -> MatmulCost {
let [h_out, w_out] = self.spatial_out;
let [k_h, k_w] = self.kernel;
MatmulCost {
batches: 1,
m: self.batch * h_out * w_out,
n: self.channels_out,
k: self.channels_in * k_h * k_w,
elems: self.elems.clone(),
}
}
pub fn work(&self) -> Work {
let (read, written) = self.traffic();
Work {
compute_ops: self.compute_ops(),
bytes: read + written,
}
}
pub fn compute_ops(&self) -> usize {
self.gemm().compute_ops()
}
pub fn traffic(&self) -> (usize, usize) {
let [h_in, w_in] = self.spatial_in;
let [h_out, w_out] = self.spatial_out;
let [k_h, k_w] = self.kernel;
let input = self.batch * self.channels_in * h_in * w_in;
let filter = self.channels_out * self.channels_in * k_h * k_w;
let output = self.batch * self.channels_out * h_out * w_out;
(
input * self.elems.lhs.size()
+ filter * self.elems.rhs.size()
+ self.bias_elems * self.elems.out.size(),
output * self.elems.out.size(),
)
}
pub fn compute_key(&self, client: &Client) -> ThroughputKey {
self.gemm().compute_key(client)
}
}
#[derive(Debug, Clone, Copy)]
pub struct DepthwiseCost {
pub batch: usize,
pub channels: usize,
pub size: usize,
pub out_size: usize,
pub kernel: usize,
pub dtype: ElemType,
}
impl DepthwiseCost {
pub fn work(&self) -> Work {
let (read, written) = self.traffic();
Work {
compute_ops: self.compute_ops(),
bytes: read + written,
}
}
pub fn compute_ops(&self) -> usize {
let taps = self.kernel * self.kernel;
self.batch * self.out_size * self.out_size * self.channels * (2 * taps).saturating_sub(1)
}
pub fn traffic(&self) -> (usize, usize) {
let maps = self.batch * self.channels;
let filter = self.channels * self.kernel * self.kernel;
let size = self.dtype.size();
(
(maps * self.size * self.size + filter) * size,
maps * self.out_size * self.out_size * size,
)
}
pub fn compute_key(&self) -> ThroughputKey {
ThroughputKey {
mode: ThroughputMode::ComputeDirect { dtype: self.dtype },
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use cubecl::ir::{ElemType, FloatKind};
fn cost() -> Conv2dCost {
let f16 = ElemType::Float(FloatKind::F16);
Conv2dCost {
batch: 2,
channels_in: 3,
spatial_in: [8, 8],
channels_out: 4,
kernel: [3, 3],
spatial_out: [6, 6],
bias_elems: 0,
elems: MatmulGlobalElems {
lhs: f16,
rhs: f16,
out: f16,
},
}
}
#[test]
fn contracts_over_the_filter_window_for_every_output_pixel() {
let outputs = 2 * 6 * 6 * 4;
let taps = 3 * 3 * 3;
assert_eq!(cost().compute_ops(), outputs * (2 * taps - 1));
}
#[test]
fn moves_the_maps_and_the_filter_once_each() {
let (read, written) = cost().traffic();
assert_eq!(read, (2 * 3 * 8 * 8 + 4 * 3 * 3 * 3) * 2);
assert_eq!(written, 2 * 4 * 6 * 6 * 2);
}
#[test]
fn a_bias_is_read_alongside_the_filter() {
let mut with_bias = cost();
with_bias.bias_elems = 4;
assert_eq!(with_bias.traffic().0, cost().traffic().0 + 4 * 2);
}
}