#![allow(unsafe_op_in_unsafe_fn)]
use std::mem::MaybeUninit;
use std::ops::Range;
use super::{GemmOutT, Im2Col};
use rten_tensor::Matrix;
pub mod generic;
mod simd_generic;
#[cfg(target_arch = "aarch64")]
pub mod aarch64;
#[cfg(target_arch = "x86_64")]
pub mod x86_64;
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
pub mod wasm;
#[derive(Clone, Copy)]
pub enum Lhs<'a, T> {
Packed(&'a [u8]),
Unpacked {
data: *const T,
len: usize,
row_stride: usize,
},
}
#[derive(Clone, Debug, PartialEq)]
pub struct PackedLayout {
size: usize,
align: usize,
panel_stride: usize,
pub must_pack: bool,
}
impl PackedLayout {
pub fn new(size: usize, align: usize, panel_stride: usize) -> PackedLayout {
debug_assert_eq!(size % align, 0);
debug_assert_eq!(size % panel_stride, 0);
PackedLayout {
size,
align,
panel_stride,
must_pack: false,
}
}
pub fn size(&self) -> usize {
self.size
}
pub fn align(&self) -> usize {
self.align
}
pub fn panel_stride(&self) -> usize {
self.panel_stride
}
}
#[derive(Debug)]
pub struct QuantParams<'a, T> {
pub zero_point: &'a [T],
}
impl<T> Clone for QuantParams<'_, T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for QuantParams<'_, T> {}
pub unsafe trait Kernel<LhsT, RhsT, OutT>: Sync {
fn new() -> Option<Self>
where
Self: Sized;
fn mr(&self) -> usize;
fn nr(&self) -> usize;
fn name(&self) -> &'static str;
fn may_saturate(&self) -> bool {
false
}
fn im2col_row_count_step(&self) -> usize {
1
}
fn im2col_col_count_step(&self) -> usize {
self.nr()
}
fn packed_a_layout(
&self,
a: Matrix<LhsT>,
rows: usize,
cols: usize,
quant: Option<QuantParams<LhsT>>,
) -> PackedLayout;
fn pack_a_block(
&self,
out: &mut [MaybeUninit<u8>],
a: Matrix<LhsT>,
rows: Range<usize>,
cols: Range<usize>,
quant: Option<QuantParams<LhsT>>,
);
fn packed_b_layout(
&self,
rows: usize,
cols: usize,
quant: Option<QuantParams<RhsT>>,
) -> PackedLayout;
fn pack_b_block(
&self,
out: &mut [MaybeUninit<u8>],
b: Matrix<RhsT>,
rows: Range<usize>,
cols: Range<usize>,
quant: Option<QuantParams<RhsT>>,
);
fn pack_im2col(
&self,
out: &mut [MaybeUninit<u8>],
image: &Im2Col<RhsT>,
rows: Range<usize>,
cols: Range<usize>,
zero_point: Option<RhsT>,
);
unsafe fn kernel(
&self,
tile_ptr: *mut OutT,
tile_row_stride: usize,
a: Lhs<LhsT>,
b: &[u8],
used_rows: usize,
used_cols: usize,
depth: usize,
alpha: f32,
beta: OutT,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
);
fn gemv_kernel(
&self,
out: MatVecOutput<OutT>,
a: &[LhsT],
b: Matrix<RhsT>,
alpha: f32,
a_quant: Option<QuantParams<LhsT>>,
b_quant: Option<QuantParams<RhsT>>,
);
}
pub struct TempTile<T: GemmOutT, const MR: usize, const NR: usize> {
data: [[MaybeUninit<T>; NR]; MR],
}
impl<T: GemmOutT, const MR: usize, const NR: usize> TempTile<T, MR, NR> {
pub fn new() -> Self {
TempTile {
data: [[MaybeUninit::<T>::uninit(); NR]; MR],
}
}
pub fn as_mut_ptr(&mut self) -> *mut MaybeUninit<T> {
self.data.as_mut_ptr() as *mut MaybeUninit<T>
}
pub unsafe fn accumulate_into(
&self,
dest: *mut MaybeUninit<T>,
n_rows: usize,
n_cols: usize,
row_stride: usize,
beta: T,
) {
if beta != T::zero() {
for i in 0..n_rows {
for j in 0..n_cols {
unsafe {
let out_el = dest.add(row_stride * i + j);
let tmp = (*out_el).assume_init();
out_el.write(MaybeUninit::new(
beta * tmp + self.data.get_unchecked(i).get_unchecked(j).assume_init(),
));
}
}
}
} else {
for i in 0..n_rows {
for j in 0..n_cols {
unsafe {
let out_el = dest.add(row_stride * i + j);
out_el.write(MaybeUninit::new(
self.data.get_unchecked(i).get_unchecked(j).assume_init(),
));
}
}
}
}
}
}
unsafe trait Int8DotProduct {
type X8;
type I32;
fn supports_indexed_dot_product() -> bool {
false
}
fn dot_product(self, a: Self::X8, b: Self::X8, c: Self::I32) -> Self::I32;
#[allow(unused_variables)]
fn indexed_dot_product<const IDX: u32>(
self,
a: Self::X8,
b: Self::X8,
c: Self::I32,
) -> Self::I32
where
Self: Sized,
{
unimplemented!("indexed_dot_product not supported")
}
}
pub struct MatVecOutput<'a, T, B = T> {
data: &'a mut [MaybeUninit<T>],
beta: B,
}
impl<'a, T, B: Copy + Default + PartialEq> MatVecOutput<'a, T, B> {
pub fn from_uninit_slice(data: &'a mut [MaybeUninit<T>]) -> Self {
MatVecOutput {
data,
beta: B::default(),
}
}
pub fn from_slice(data: &'a mut [T], beta: B) -> Self {
let data = unsafe { std::mem::transmute::<&'a mut [T], &'a mut [MaybeUninit<T>]>(data) };
MatVecOutput { data, beta }
}
pub fn slice_mut(&mut self, range: Range<usize>) -> MatVecOutput<'_, T, B> {
MatVecOutput {
data: &mut self.data[range],
beta: self.beta,
}
}
pub fn as_bool_beta(&mut self) -> MatVecOutput<'_, T, bool> {
MatVecOutput {
data: self.data,
beta: self.beta != B::default(),
}
}
}
#[cfg(test)]
mod tests {
use rten_tensor::prelude::*;
use rten_tensor::rng::XorShiftRng;
use rten_tensor::{AssumeInit, MatrixLayout, NdTensor, RandomSource};
use super::{Kernel, Lhs};
use crate::PackingBuffer;
fn run_kernel_bench<LhsT: Clone + Default, RhsT: Clone + Default, OutT: Clone + Default>(
kernel: &dyn Kernel<LhsT, RhsT, OutT>,
) where
XorShiftRng: RandomSource<LhsT> + RandomSource<RhsT>,
{
let mut output = NdTensor::zeros([kernel.mr(), kernel.nr()]);
let target_size = 24 * 1024;
let m = kernel.mr();
let n = kernel.nr();
let k = target_size / (m * size_of::<LhsT>() + n * size_of::<RhsT>());
let mut rng = XorShiftRng::new(1234);
let a = NdTensor::rand([m, k], &mut rng);
let b = NdTensor::rand([k, n], &mut rng);
let a_layout = kernel.packed_a_layout(a.view(), a.rows(), a.cols(), None);
let b_layout = kernel.packed_b_layout(b.rows(), b.cols(), None);
let mut packed_a_buf = PackingBuffer::new();
let packed_a = packed_a_buf.alloc(a_layout.size(), a_layout.align());
kernel.pack_a_block(packed_a, a.view(), 0..a.rows(), 0..a.cols(), None);
let packed_a = unsafe { packed_a.assume_init() };
let mut packed_b_buf = PackingBuffer::new();
let packed_b = packed_b_buf.alloc(b_layout.size(), b_layout.align());
kernel.pack_b_block(packed_b, b.view(), 0..b.rows(), 0..b.cols(), None);
let packed_b = unsafe { packed_b.assume_init() };
let n_iters = 1_000_000;
let start = std::time::Instant::now();
for _ in 0..n_iters {
unsafe {
let out_row_stride = output.stride(0);
kernel.kernel(
output.data_mut().unwrap().as_mut_ptr(),
out_row_stride,
Lhs::Packed(packed_a),
packed_b,
kernel.mr(),
kernel.nr(),
k,
1.,
OutT::default(), None, None, );
}
output.fill(OutT::default());
}
let elapsed = start.elapsed().as_secs_f64();
let n_ops: u64 = 2 * m as u64 * n as u64 * k as u64 * n_iters as u64;
let gflops = (n_ops as f64 / 1e9) / elapsed;
println!(
"{}: {} iters in {:.4}s = {:.2} GFLOPS",
kernel.name(),
n_iters,
elapsed,
gflops
);
}
struct KernelBench<LhsT, RhsT, OutT> {
kernels: Vec<Box<dyn Kernel<LhsT, RhsT, OutT>>>,
}
impl<LhsT, RhsT, OutT> KernelBench<LhsT, RhsT, OutT> {
fn new() -> Self {
Self {
kernels: Vec::new(),
}
}
fn add<K: Kernel<LhsT, RhsT, OutT> + 'static>(&mut self) {
if let Some(kernel) = K::new() {
self.kernels.push(Box::new(kernel))
}
}
fn run_bench(&self)
where
LhsT: Clone + Default,
RhsT: Clone + Default,
OutT: Clone + Default,
XorShiftRng: RandomSource<RhsT> + RandomSource<LhsT>,
{
let filter = std::env::var("RTEN_BENCH_FILTER").ok();
for kernel in &self.kernels {
let filter_match = filter
.as_ref()
.map(|f| kernel.name().contains(f))
.unwrap_or(true);
if !filter_match {
continue;
}
run_kernel_bench(kernel.as_ref())
}
}
}
#[test]
#[ignore]
fn bench_kernel_f32() {
let mut kernels = KernelBench::<f32, f32, f32>::new();
kernels.add::<super::generic::GenericKernel>();
#[cfg(target_arch = "aarch64")]
{
kernels.add::<super::aarch64::ArmNeonKernel>();
}
#[cfg(target_arch = "x86_64")]
{
kernels.add::<super::x86_64::FmaKernel>();
kernels.add::<super::x86_64::Avx512Kernel>();
}
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
{
kernels.add::<super::wasm::WasmKernel>();
}
kernels.run_bench();
}
#[test]
#[ignore]
fn bench_kernel_int8() {
let mut kernels = KernelBench::<u8, i8, i32>::new();
kernels.add::<super::generic::GenericKernel>();
#[cfg(target_arch = "aarch64")]
{
kernels.add::<super::aarch64::ArmInt8MlalKernel>();
kernels.add::<super::aarch64::ArmInt8DotKernel>();
kernels.add::<super::aarch64::ArmInt8MMKernel>();
}
#[cfg(target_arch = "x86_64")]
{
kernels.add::<super::x86_64::Avx2Int8Kernel>();
kernels.add::<super::x86_64::Avx512Int8Kernel>();
}
#[cfg(target_arch = "wasm32")]
#[cfg(target_feature = "simd128")]
{
kernels.add::<super::wasm::WasmInt8Kernel>();
}
kernels.run_bench();
}
}