use burn::tensor::{Tensor, backend::Backend};
pub(crate) const BROKEN_MATMUL_BOUNDARY: usize = 512;
const SAFE_M_SLAB: usize = 256;
pub(crate) fn safe_matmul<B: Backend, const D: usize>(
lhs: Tensor<B, D>,
rhs: Tensor<B, D>,
) -> Tensor<B, D> {
let dims = lhs.dims();
let (m, k) = (dims[D - 2], dims[D - 1]);
if m < BROKEN_MATMUL_BOUNDARY || k < BROKEN_MATMUL_BOUNDARY {
return lhs.matmul(rhs);
}
let mut slabs = Vec::with_capacity(m.div_ceil(SAFE_M_SLAB));
let mut offset = 0;
while offset < m {
let len = SAFE_M_SLAB.min(m - offset);
slabs.push(
lhs.clone()
.narrow(D - 2, offset, len)
.matmul(rhs.clone()),
);
offset += len;
}
Tensor::cat(slabs, D - 2)
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::TensorData;
type B = burn::backend::NdArray<f32>;
#[test]
fn slabbed_matches_single_matmul() {
let device = Default::default();
let (m, k, n) = (600, 576, 1536);
let a: Vec<f32> = (0..m * k).map(|i| ((i * 7 % 13) as f32 - 6.0) / 8.0).collect();
let b: Vec<f32> = (0..k * n).map(|i| ((i * 5 % 11) as f32 - 5.0) / 8.0).collect();
let lhs = Tensor::<B, 2>::from_data(TensorData::new(a, [m, k]), &device);
let rhs = Tensor::<B, 2>::from_data(TensorData::new(b, [k, n]), &device);
let single = lhs.clone().matmul(rhs.clone());
let slabbed = safe_matmul(lhs, rhs);
let single: Vec<f32> = single.into_data().to_vec().unwrap();
let slabbed: Vec<f32> = slabbed.into_data().to_vec().unwrap();
assert_eq!(single, slabbed, "M-slabbing must be exact");
}
#[test]
fn passthrough_below_boundary() {
let device = Default::default();
let (m, k, n) = (511, 64, 32);
let a: Vec<f32> = (0..m * k).map(|i| (i % 7) as f32).collect();
let b: Vec<f32> = (0..k * n).map(|i| (i % 5) as f32).collect();
let lhs = Tensor::<B, 2>::from_data(TensorData::new(a, [m, k]), &device);
let rhs = Tensor::<B, 2>::from_data(TensorData::new(b, [k, n]), &device);
let single: Vec<f32> = lhs.clone().matmul(rhs.clone()).into_data().to_vec().unwrap();
let wrapped: Vec<f32> = safe_matmul(lhs, rhs).into_data().to_vec().unwrap();
assert_eq!(single, wrapped);
}
}