use super::Differentiable;
#[derive(Debug)]
pub struct GemmTask<'buffers, Element> {
a: &'buffers [Element],
b: &'buffers [Element],
m: usize,
n: usize,
k: usize,
a_strides: [usize; 2],
b_strides: [usize; 2],
}
impl<'buffers, Element> GemmTask<'buffers, Element> {
pub(crate) fn new(
a: &'buffers [Element],
a_strides: [usize; 2],
b: &'buffers [Element],
b_strides: [usize; 2],
m: usize,
k: usize,
n: usize,
) -> Self {
assert!(
m > 0 && k > 0 && n > 0,
"a gemm task needs non-empty extents"
);
let a_span = 1 + (m - 1) * a_strides[0] + (k - 1) * a_strides[1];
assert!(
a.len() >= a_span,
"the left operand slice does not span its {m} x {k} matrix"
);
let b_span = 1 + (k - 1) * b_strides[0] + (n - 1) * b_strides[1];
assert!(
b.len() >= b_span,
"the right operand slice does not span its {k} x {n} matrix"
);
Self {
a,
b,
m,
n,
k,
a_strides,
b_strides,
}
}
pub fn a(&self) -> &'buffers [Element] {
self.a
}
pub fn b(&self) -> &'buffers [Element] {
self.b
}
pub fn m(&self) -> usize {
self.m
}
pub fn n(&self) -> usize {
self.n
}
pub fn k(&self) -> usize {
self.k
}
pub fn a_strides(&self) -> [usize; 2] {
self.a_strides
}
pub fn b_strides(&self) -> [usize; 2] {
self.b_strides
}
}
pub(crate) fn multiply<Element: Differentiable>(task: &GemmTask<'_, Element>) -> Vec<Element> {
let mut accumulators = Vec::with_capacity(task.m * task.n);
for row in 0..task.m {
let a_row_start = row * task.a_strides[0];
let a_first = task.a[a_row_start].promote();
seed_row(&mut accumulators, &a_first, task);
let output = &mut accumulators[row * task.n..];
for step in 1..task.k {
let a_value = task.a[a_row_start + step * task.a_strides[1]].promote();
accumulate_row(output, &a_value, task, step);
}
}
accumulators.into_iter().map(Element::demote).collect()
}
fn seed_row<Element: Differentiable>(
accumulators: &mut Vec<Element::Accumulator>,
a_first: &Element::Accumulator,
task: &GemmTask<'_, Element>,
) {
if task.b_strides[1] == 1 {
let b_row = &task.b[..task.n];
accumulators.extend(
b_row
.iter()
.map(|b_element| a_first.clone() * b_element.promote()),
);
return;
}
accumulators.extend(
(0..task.n).map(|column| a_first.clone() * task.b[column * task.b_strides[1]].promote()),
);
}
fn accumulate_row<Element: Differentiable>(
output: &mut [Element::Accumulator],
a_value: &Element::Accumulator,
task: &GemmTask<'_, Element>,
step: usize,
) {
let b_row_start = step * task.b_strides[0];
if task.b_strides[1] == 1 {
let b_row = &task.b[b_row_start..b_row_start + task.n];
for (output_element, b_element) in output.iter_mut().zip(b_row) {
*output_element = output_element.clone() + a_value.clone() * b_element.promote();
}
return;
}
for (column, output_element) in output.iter_mut().enumerate() {
*output_element = output_element.clone()
+ a_value.clone() * task.b[b_row_start + column * task.b_strides[1]].promote();
}
}
#[cfg(test)]
#[path = "tests/gemm_tests.rs"]
mod tests;