use crate::tensor::{Metadata, Shape, Strides};
const ROW_TILE: usize = 32;
#[derive(Debug, PartialEq, Eq, Clone, Copy)]
pub enum MatmulTransformAction {
Keep,
MergeBatches {
rows: usize,
},
}
impl MatmulTransformAction {
pub fn apply(&self, meta: &mut Metadata) {
let rows = match self {
MatmulTransformAction::Keep => return,
MatmulTransformAction::MergeBatches { rows } => *rows,
};
let rank = meta.rank();
let stride_rows = merged_row_stride(meta.shape(), meta.strides())
.expect("The action requires batch-contiguous rows");
for i in 0..rank - 2 {
meta.shape_mut()[i] = 1;
meta.strides_mut()[i] = rows * stride_rows;
}
meta.shape_mut()[rank - 2] = rows;
meta.strides_mut()[rank - 2] = stride_rows;
}
}
#[derive(Debug, PartialEq, Eq, Clone, Copy)]
pub struct MatmulTransformAnalysis {
batches: usize,
rows: usize,
cols: usize,
mergeable: bool,
}
impl MatmulTransformAnalysis {
pub fn from_shapes(lhs: &Shape, rhs: &Shape) -> Self {
Self::new(lhs, rhs, true)
}
pub fn from_metadata(lhs: &Metadata, rhs: &Metadata, out: &Metadata) -> Self {
let mergeable = merged_row_stride(lhs.shape(), lhs.strides()).is_some()
&& merged_row_stride(out.shape(), out.strides()).is_some();
Self::new(lhs.shape(), rhs.shape(), mergeable)
}
fn new(lhs: &Shape, rhs: &Shape, rows_contiguous: bool) -> Self {
let rank = lhs.num_dims();
let rank_rhs = rhs.num_dims();
let batches = lhs[..rank - 2].iter().product();
let rows = lhs[rank - 2];
let cols = rhs[rank_rhs - 1];
let rhs_shared = rhs[..rank_rhs - 2].iter().all(|&dim| dim == 1);
Self {
batches,
rows,
cols,
mergeable: rhs_shared && rows_contiguous,
}
}
}
#[derive(Debug, Default, Clone, Copy)]
pub enum MatmulTransformPolicy {
#[default]
BetterTiling,
Never,
}
impl MatmulTransformPolicy {
pub fn action(&self, analysis: &MatmulTransformAnalysis) -> MatmulTransformAction {
match self {
MatmulTransformPolicy::Never => MatmulTransformAction::Keep,
MatmulTransformPolicy::BetterTiling => {
if !analysis.mergeable || analysis.batches == 1 {
return MatmulTransformAction::Keep;
}
if analysis.rows.is_multiple_of(ROW_TILE) {
return MatmulTransformAction::Keep;
}
let rows = analysis.batches * analysis.rows;
if !squarer(rows, analysis.rows, analysis.cols) {
return MatmulTransformAction::Keep;
}
MatmulTransformAction::MergeBatches { rows }
}
}
}
}
fn squarer(rows_new: usize, rows: usize, cols: usize) -> bool {
rows_new.min(cols) * rows.max(cols) > rows.min(cols) * rows_new.max(cols)
}
fn merged_row_stride(shape: &Shape, strides: &Strides) -> Option<usize> {
let rank = shape.num_dims();
let mut chained: Option<(usize, usize)> = None;
for i in (0..rank - 1).rev() {
if shape[i] == 1 {
continue;
}
match chained {
None => chained = Some((strides[i], strides[i] * shape[i])),
Some((row_stride, expected)) => {
if strides[i] != expected {
return None;
}
chained = Some((row_stride, strides[i] * shape[i]));
}
}
}
match chained {
Some((row_stride, _)) => Some(row_stride),
None => Some(shape[rank - 1] * strides[rank - 1]),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::shape;
fn analysis(lhs: &[usize], rhs: &[usize]) -> MatmulTransformAnalysis {
MatmulTransformAnalysis::from_shapes(&Shape::from(lhs), &Shape::from(rhs))
}
fn action(analysis: &MatmulTransformAnalysis) -> MatmulTransformAction {
MatmulTransformPolicy::default().action(analysis)
}
#[test]
fn merges_batched_vec_mat() {
let analysis = analysis(&[16, 1, 4096], &[4096, 14336]);
assert_eq!(
action(&analysis),
MatmulTransformAction::MergeBatches { rows: 16 }
);
}
#[test]
fn keeps_row_tiled_batches() {
let analysis = analysis(&[8, 512, 4096], &[4096, 14336]);
assert_eq!(action(&analysis), MatmulTransformAction::Keep);
}
#[test]
fn keeps_single_batch() {
let analysis = analysis(&[1, 1, 4096], &[4096, 14336]);
assert_eq!(action(&analysis), MatmulTransformAction::Keep);
}
#[test]
fn keeps_matrix() {
let analysis = analysis(&[16, 4096], &[4096, 14336]);
assert_eq!(action(&analysis), MatmulTransformAction::Keep);
}
#[test]
fn keeps_batched_rhs() {
let analysis = analysis(&[8, 1, 64], &[8, 64, 32]);
assert_eq!(action(&analysis), MatmulTransformAction::Keep);
}
#[test]
fn merges_with_broadcast_rhs_rank() {
let analysis = analysis(&[8, 1, 64], &[1, 64, 32]);
assert_eq!(
action(&analysis),
MatmulTransformAction::MergeBatches { rows: 8 }
);
}
#[test]
fn keeps_when_merge_leaves_square() {
let analysis = analysis(&[4, 1000, 64], &[64, 64]);
assert_eq!(action(&analysis), MatmulTransformAction::Keep);
}
#[test]
fn merges_multiple_batch_dims() {
let analysis = analysis(&[2, 8, 1, 4096], &[4096, 14336]);
assert_eq!(
action(&analysis),
MatmulTransformAction::MergeBatches { rows: 16 }
);
}
#[test]
fn never_policy_keeps() {
let analysis = analysis(&[16, 1, 4096], &[4096, 14336]);
assert_eq!(
MatmulTransformPolicy::Never.action(&analysis),
MatmulTransformAction::Keep
);
}
#[test]
fn metadata_merges_pitched_rows() {
let lhs = Metadata::new(shape![16, 1, 4096], crate::strides![8192, 4096, 1]);
let rhs = Metadata::new(shape![1, 4096, 14336], crate::strides![0, 14336, 1]);
let out = Metadata::new(shape![16, 1, 14336], crate::strides![14336, 14336, 1]);
let analysis = MatmulTransformAnalysis::from_metadata(&lhs, &rhs, &out);
assert_eq!(
action(&analysis),
MatmulTransformAction::MergeBatches { rows: 16 }
);
}
#[test]
fn metadata_keeps_batch_holes() {
let lhs = Metadata::new(shape![4, 3, 8], crate::strides![48, 8, 1]);
let rhs = Metadata::new(shape![1, 8, 32], crate::strides![0, 32, 1]);
let out = Metadata::new(shape![4, 3, 32], crate::strides![96, 32, 1]);
let analysis = MatmulTransformAnalysis::from_metadata(&lhs, &rhs, &out);
assert_eq!(action(&analysis), MatmulTransformAction::Keep);
}
#[test]
fn metadata_merges_contiguous() {
let lhs = Metadata::new(shape![16, 1, 4096], crate::strides![4096, 4096, 1]);
let rhs = Metadata::new(shape![1, 4096, 14336], crate::strides![0, 14336, 1]);
let out = Metadata::new(shape![16, 1, 14336], crate::strides![14336, 14336, 1]);
let analysis = MatmulTransformAnalysis::from_metadata(&lhs, &rhs, &out);
assert_eq!(
action(&analysis),
MatmulTransformAction::MergeBatches { rows: 16 }
);
}
#[test]
fn metadata_ignores_unit_dim_strides() {
let lhs = Metadata::new(shape![16, 1, 4096], crate::strides![4096, 12345, 1]);
let rhs = Metadata::new(shape![1, 4096, 14336], crate::strides![0, 14336, 1]);
let out = Metadata::new(shape![16, 1, 14336], crate::strides![14336, 99999, 1]);
let analysis = MatmulTransformAnalysis::from_metadata(&lhs, &rhs, &out);
assert_eq!(
action(&analysis),
MatmulTransformAction::MergeBatches { rows: 16 }
);
}
#[test]
fn apply_folds_batches_in_place() {
let mut lhs = Metadata::new(shape![16, 1, 4096], crate::strides![4096, 4096, 1]);
MatmulTransformAction::MergeBatches { rows: 16 }.apply(&mut lhs);
assert_eq!(lhs.shape(), &shape![1, 16, 4096]);
assert_eq!(lhs.strides()[1], 4096);
assert_eq!(lhs.strides()[2], 1);
}
#[test]
fn apply_keep_is_noop() {
let mut lhs = Metadata::new(shape![16, 1, 4096], crate::strides![4096, 4096, 1]);
MatmulTransformAction::Keep.apply(&mut lhs);
assert_eq!(lhs.shape(), &shape![16, 1, 4096]);
}
}