use crate::{GemmTask, MapOperation};
use super::operand::{Operand, classify};
const ROW_MAJOR: i32 = 101;
const NO_TRANSPOSE: i32 = 111;
const TRANSPOSE: i32 = 112;
fn transpose(operand: &Operand) -> i32 {
if operand.transposed {
TRANSPOSE
} else {
NO_TRANSPOSE
}
}
const FLOP_THRESHOLD: usize = 1 << 13;
const MAP_THRESHOLD: usize = 1 << 7;
#[link(name = "Accelerate", kind = "framework")]
unsafe extern "C" {
fn cblas_sgemm(
order: i32,
transpose_a: i32,
transpose_b: i32,
m: i32,
n: i32,
k: i32,
alpha: f32,
a: *const f32,
leading_a: i32,
b: *const f32,
leading_b: i32,
beta: f32,
c: *mut f32,
leading_c: i32,
);
fn cblas_dgemm(
order: i32,
transpose_a: i32,
transpose_b: i32,
m: i32,
n: i32,
k: i32,
alpha: f64,
a: *const f64,
leading_a: i32,
b: *const f64,
leading_b: i32,
beta: f64,
c: *mut f64,
leading_c: i32,
);
fn vvexpf(mapped: *mut f32, elements: *const f32, count: *const i32);
fn vvlogf(mapped: *mut f32, elements: *const f32, count: *const i32);
fn vvsqrtf(mapped: *mut f32, elements: *const f32, count: *const i32);
fn vvtanhf(mapped: *mut f32, elements: *const f32, count: *const i32);
fn vvexp(mapped: *mut f64, elements: *const f64, count: *const i32);
fn vvlog(mapped: *mut f64, elements: *const f64, count: *const i32);
fn vvsqrt(mapped: *mut f64, elements: *const f64, count: *const i32);
fn vvtanh(mapped: *mut f64, elements: *const f64, count: *const i32);
}
pub(super) fn gemm_f32(task: &GemmTask<'_, f32>) -> Option<Vec<f32>> {
if flops(task.m(), task.n(), task.k()) < FLOP_THRESHOLD {
return None;
}
executed_f32(task)
}
pub(super) fn gemm_f64(task: &GemmTask<'_, f64>) -> Option<Vec<f64>> {
if flops(task.m(), task.n(), task.k()) < FLOP_THRESHOLD {
return None;
}
executed_f64(task)
}
fn flops(m: usize, n: usize, k: usize) -> usize {
2usize.saturating_mul(m).saturating_mul(n).saturating_mul(k)
}
pub(super) fn executed_f32(task: &GemmTask<'_, f32>) -> Option<Vec<f32>> {
let a = classify(task.a_strides(), task.m(), task.k())?;
let b = classify(task.b_strides(), task.k(), task.n())?;
let m = i32::try_from(task.m()).ok()?;
let n = i32::try_from(task.n()).ok()?;
let k = i32::try_from(task.k()).ok()?;
let mut product = vec![0.0_f32; task.m() * task.n()];
unsafe {
cblas_sgemm(
ROW_MAJOR,
transpose(&a),
transpose(&b),
m,
n,
k,
1.0,
task.a().as_ptr(),
a.leading,
task.b().as_ptr(),
b.leading,
0.0,
product.as_mut_ptr(),
n,
);
}
Some(product)
}
pub(super) fn map_f32(operation: MapOperation, elements: &[f32]) -> Option<Vec<f32>> {
if elements.len() < MAP_THRESHOLD {
return None;
}
let count = i32::try_from(elements.len()).ok()?;
let mut mapped = vec![0.0_f32; elements.len()];
unsafe {
match operation {
MapOperation::Exp => vvexpf(mapped.as_mut_ptr(), elements.as_ptr(), &count),
MapOperation::Ln => vvlogf(mapped.as_mut_ptr(), elements.as_ptr(), &count),
MapOperation::Sqrt => vvsqrtf(mapped.as_mut_ptr(), elements.as_ptr(), &count),
MapOperation::Tanh => vvtanhf(mapped.as_mut_ptr(), elements.as_ptr(), &count),
}
}
Some(mapped)
}
pub(super) fn map_f64(operation: MapOperation, elements: &[f64]) -> Option<Vec<f64>> {
if elements.len() < MAP_THRESHOLD {
return None;
}
let count = i32::try_from(elements.len()).ok()?;
let mut mapped = vec![0.0_f64; elements.len()];
unsafe {
match operation {
MapOperation::Exp => vvexp(mapped.as_mut_ptr(), elements.as_ptr(), &count),
MapOperation::Ln => vvlog(mapped.as_mut_ptr(), elements.as_ptr(), &count),
MapOperation::Sqrt => vvsqrt(mapped.as_mut_ptr(), elements.as_ptr(), &count),
MapOperation::Tanh => vvtanh(mapped.as_mut_ptr(), elements.as_ptr(), &count),
}
}
Some(mapped)
}
pub(super) fn executed_f64(task: &GemmTask<'_, f64>) -> Option<Vec<f64>> {
let a = classify(task.a_strides(), task.m(), task.k())?;
let b = classify(task.b_strides(), task.k(), task.n())?;
let m = i32::try_from(task.m()).ok()?;
let n = i32::try_from(task.n()).ok()?;
let k = i32::try_from(task.k()).ok()?;
let mut product = vec![0.0_f64; task.m() * task.n()];
unsafe {
cblas_dgemm(
ROW_MAJOR,
transpose(&a),
transpose(&b),
m,
n,
k,
1.0,
task.a().as_ptr(),
a.leading,
task.b().as_ptr(),
b.leading,
0.0,
product.as_mut_ptr(),
n,
);
}
Some(product)
}
#[cfg(test)]
#[path = "tests/accelerate_tests.rs"]
mod tests;