use core::mem::MaybeUninit;
use crate::driver::{self, alpha_status, beta_status};
use crate::kernel::epilogue::{Epilogue, Identity};
use crate::kernel::{KernelFamily, MAX_MR, SCRATCH_LEN};
use crate::parallel::{self, JobCursor, Parallelism, Ptr};
use crate::simd::{KernelSimd, SimdOps};
use crate::workspace::Workspace;
const SMALL_K_MAX: usize = 32;
#[allow(clippy::too_many_arguments)]
#[cfg_attr(not(feature = "int8"), allow(dead_code))]
pub unsafe fn run<Fam, S, const MR_REG: usize, const NR: usize>(
simd: S,
m: usize,
k: usize,
n: usize,
alpha: Fam::Acc,
a: *const Fam::Lhs,
rsa: isize,
csa: isize,
b: *const Fam::Rhs,
rsb: isize,
csb: isize,
beta: Fam::Acc,
c: *mut Fam::Out,
rsc: isize,
csc: isize,
par: Parallelism,
ws: &mut Workspace,
) where
Fam: KernelFamily,
S: KernelSimd<Fam::Lhs, Fam::Rhs, Fam::Acc, Fam::Out>,
{
unsafe {
run_epi::<Fam, S, Identity, MR_REG, NR>(
simd, m, k, n, alpha, a, rsa, csa, b, rsb, csb, beta, c, rsc, csc, &Identity, par, ws,
)
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_epi<Fam, S, E, const MR_REG: usize, const NR: usize>(
simd: S,
m: usize,
k: usize,
n: usize,
alpha: Fam::Acc,
a: *const Fam::Lhs,
rsa: isize,
csa: isize,
b: *const Fam::Rhs,
rsb: isize,
csb: isize,
beta: Fam::Acc,
c: *mut Fam::Out,
rsc: isize,
csc: isize,
epi: &E,
par: Parallelism,
ws: &mut Workspace,
) where
Fam: KernelFamily,
S: KernelSimd<Fam::Lhs, Fam::Rhs, Fam::Acc, Fam::Out>,
E: Epilogue<Fam>,
{
unsafe {
if Fam::FORCE_PACK_LHS || Fam::FORCE_PACK_RHS || rsa != 1 || k > SMALL_K_MAX {
driver::run_epilogue::<Fam, S, E, MR_REG, NR>(
simd, m, k, n, alpha, a, rsa, csa, b, rsb, csb, beta, c, rsc, csc, epi, par, ws,
);
return;
}
let epi = *epi;
let lanes = <S as SimdOps<Fam::Acc>>::LANES;
let mr = MR_REG * lanes;
let nr = NR;
debug_assert!(mr * nr <= SCRATCH_LEN, "microtile exceeds scratch capacity");
debug_assert!(mr <= MAX_MR, "microtile rows exceed MAX_MR");
let n_row_tiles = m.div_ceil(mr);
let n_col_tiles = n.div_ceil(nr);
let n_tiles = n_row_tiles * n_col_tiles;
let last_ic = (n_row_tiles - 1) * mr;
let bottom_eff = m - last_ic; let bottom_partial = bottom_eff < mr;
let ash = alpha_status(alpha);
let bst = beta_status(beta);
let bytes = m
.saturating_mul(k)
.saturating_mul(core::mem::size_of::<Fam::Lhs>())
.saturating_add(
k.saturating_mul(n)
.saturating_mul(core::mem::size_of::<Fam::Rhs>()),
)
.saturating_add(
m.saturating_mul(n)
.saturating_mul(core::mem::size_of::<Fam::Out>()),
);
let n_threads = par.resolve_bandwidth(bytes, n_tiles);
let a = Ptr(a as *mut Fam::Lhs);
let b = Ptr(b as *mut Fam::Rhs);
let c = Ptr(c);
let body = move |q_start: usize, q_end: usize| {
let (a, b, c, epi) = (a, b, c, epi);
let a = a.0 as *const Fam::Lhs;
let b = b.0 as *const Fam::Rhs;
let c = c.0;
simd.vectorize(|| {
let mut scratch = [const { MaybeUninit::<Fam::Acc>::uninit() }; SCRATCH_LEN];
let scratch_ptr = scratch.as_mut_ptr() as *mut Fam::Acc;
let mut pad = [const { MaybeUninit::<Fam::Lhs>::uninit() }; MAX_MR * SMALL_K_MAX];
let pad_base = if bottom_partial {
Fam::pack_lhs(
pad.as_mut_ptr() as *mut Fam::Lhs,
a.offset(last_ic as isize * rsa),
rsa,
csa,
bottom_eff,
k,
mr,
);
pad.as_ptr() as *const Fam::Lhs
} else {
core::ptr::null()
};
for q in q_start..q_end {
let jt = q / n_row_tiles;
let it = q % n_row_tiles;
let ic = it * mr;
let jc = jt * nr;
let mr_eff = core::cmp::min(mr, m - ic);
let nr_eff = core::cmp::min(nr, n - jc);
let (apan, a_cs) = if mr_eff < mr {
(pad_base, mr as isize)
} else {
(a.offset(ic as isize * rsa), csa)
};
let bpan = b.offset(jc as isize * csb);
let cptr = c.offset(ic as isize * rsc + jc as isize * csc);
Fam::microkernel_epi::<S, E, MR_REG, NR>(
simd,
k,
alpha,
beta,
ash,
bst,
apan,
a_cs,
bpan,
rsb,
csb,
cptr,
rsc,
csc,
mr_eff,
nr_eff,
ic,
jc,
true,
&epi,
scratch_ptr,
);
}
});
};
if n_threads <= 1 {
body(0, n_tiles);
return;
}
let cur = JobCursor::new(n_tiles, parallel::job_grain(n_tiles, n_threads));
parallel::for_each_worker(n_threads, |_tid| {
while let Some((s, e)) = cur.next_chunk() {
body(s, e);
}
});
}
}