use core::mem::MaybeUninit;
use crate::cache;
use crate::kernel::epilogue::{Epilogue, Identity};
use crate::kernel::{AlphaStatus, BetaStatus, KernelFamily, SCRATCH_LEN};
use crate::parallel::{self, Parallelism, Ptr};
use crate::scalar::Scalar;
use crate::simd::{KernelSimd, SimdOps};
use crate::tuning;
use crate::workspace::{Regions, Workspace};
#[inline]
pub(crate) fn alpha_status<T: crate::scalar::Scalar>(a: T) -> AlphaStatus {
if a == T::ONE {
AlphaStatus::One
} else {
AlphaStatus::Other
}
}
#[inline]
pub(crate) fn beta_status<T: crate::scalar::Scalar>(b: T) -> BetaStatus {
if b == T::ZERO {
BetaStatus::Zero
} else if b == T::ONE {
BetaStatus::One
} else {
BetaStatus::Other
}
}
#[inline]
fn do_pack_lhs(rsa: isize, mc_eff: usize, mr: usize, want_pack: bool) -> bool {
rsa != 1 || !mc_eff.is_multiple_of(mr) || want_pack
}
#[inline(always)]
fn prefetch_c_tile<T>(c: *const T, rsc: isize, csc: isize, mr_eff: usize, nr_eff: usize) {
#[cfg(target_arch = "x86_64")]
unsafe {
use core::arch::x86_64::{_MM_HINT_T0, _mm_prefetch};
let esz = core::mem::size_of::<T>();
if rsc == 1 {
let col_bytes = mr_eff * esz;
for j in 0..nr_eff {
let col = c.offset(j as isize * csc) as *const i8;
let mut off = 0;
while off < col_bytes {
_mm_prefetch::<_MM_HINT_T0>(col.add(off));
off += 64;
}
}
} else if csc == 1 {
let row_bytes = nr_eff * esz;
for i in 0..mr_eff {
let row = c.offset(i as isize * rsc) as *const i8;
let mut off = 0;
while off < row_bytes {
_mm_prefetch::<_MM_HINT_T0>(row.add(off));
off += 64;
}
}
}
}
#[cfg(not(target_arch = "x86_64"))]
{
let _ = (c, rsc, csc, mr_eff, nr_eff);
}
}
#[allow(clippy::too_many_arguments)]
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_inner::<Fam, S, MR_REG, NR, Identity>(
simd, m, k, n, alpha, a, rsa, csa, b, rsb, csb, beta, c, rsc, csc, par, ws, None,
&Identity,
)
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_epilogue<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 {
run_inner::<Fam, S, MR_REG, NR, E>(
simd, m, k, n, alpha, a, rsa, csa, b, rsb, csb, beta, c, rsc, csc, par, ws, None, epi,
)
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_packed_rhs<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,
packed_b: *const Fam::Rhs,
kc: usize,
nc: usize,
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_inner::<Fam, S, MR_REG, NR, Identity>(
simd,
m,
k,
n,
alpha,
a,
rsa,
csa,
packed_b,
0,
0,
beta,
c,
rsc,
csc,
par,
ws,
Some((packed_b, kc, nc)),
&Identity,
)
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_packed_rhs_epilogue<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,
packed_b: *const Fam::Rhs,
kc: usize,
nc: usize,
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 {
run_inner::<Fam, S, MR_REG, NR, E>(
simd,
m,
k,
n,
alpha,
a,
rsa,
csa,
packed_b,
0,
0,
beta,
c,
rsc,
csc,
par,
ws,
Some((packed_b, kc, nc)),
epi,
)
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn pack_rhs_full<Fam: KernelFamily>(
dst: *mut Fam::Rhs,
b: *const Fam::Rhs,
rsb: isize,
csb: isize,
k: usize,
n: usize,
kc: usize,
nc: usize,
nr: usize,
) {
unsafe {
let mut d = dst;
let mut jc = 0;
while jc < n {
let nc_eff = core::cmp::min(nc, n - jc);
let n_nt = nc_eff.div_ceil(nr);
let mut pc = 0;
while pc < k {
let kc_eff = core::cmp::min(kc, k - pc);
let kc_eff_pad = kc_eff.next_multiple_of(Fam::DEPTH_MULTIPLE);
for jt in 0..n_nt {
let col = jc + jt * nr;
let nr_eff = core::cmp::min(nr, nc_eff - jt * nr);
let src = b.offset(pc as isize * rsb + col as isize * csb);
Fam::pack_rhs(d, src, rsb, csb, kc_eff, nr_eff, nr);
d = d.add(kc_eff_pad * nr);
}
pc += kc;
}
jc += nc;
}
}
}
#[allow(clippy::too_many_arguments)]
unsafe fn run_inner<Fam, S, const MR_REG: usize, const NR: usize, E>(
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,
packed: Option<(*const Fam::Rhs, usize, usize)>,
epi: &E,
) where
Fam: KernelFamily,
S: KernelSimd<Fam::Lhs, Fam::Rhs, Fam::Acc, Fam::Out>,
E: Epilogue<Fam>,
{
let epi = *epi;
unsafe {
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_eq!(
core::mem::size_of::<Fam::Lhs>(),
core::mem::size_of::<Fam::Rhs>()
);
let sizeof_lhs = core::mem::size_of::<Fam::Lhs>().max(1);
let blk = cache::topology().blocking(mr, nr, sizeof_lhs, m, n, k);
let mc = blk.mc.next_multiple_of(mr).max(mr);
let (kc, nc) = match packed {
Some((_, pkc, pnc)) => (pkc.max(1), pnc.next_multiple_of(nr).max(nr)),
None => {
let kc = if Fam::OUT_IS_ACC {
blk.kc.next_multiple_of(Fam::DEPTH_MULTIPLE)
} else {
k
};
(kc.max(1), blk.nc.next_multiple_of(nr).max(nr))
}
};
let q_depth = Fam::DEPTH_MULTIPLE;
let kc_pad_block = kc.next_multiple_of(q_depth);
assert!(
q_depth == 1 || packed.is_none() || kc >= k,
"prepacked-RHS with DEPTH_MULTIPLE > 1 requires a single depth slice (kc >= k)"
);
let n_nt_max = nc.div_ceil(nr);
let n_jobs_max = m.div_ceil(mc) * n_nt_max;
let mnk = m.saturating_mul(n).saturating_mul(k);
let n_threads = par.resolve(mnk, n_jobs_max);
const PAR_JOBS_PER_WORKER: usize = 8;
let mc = if n_threads > 1 && n_jobs_max < n_threads * PAR_JOBS_PER_WORKER {
let n_mc_target = (n_threads * PAR_JOBS_PER_WORKER)
.div_ceil(n_nt_max)
.min(m.div_ceil(mr));
mc.min(m.div_ceil(n_mc_target).next_multiple_of(mr).max(mr))
} else {
mc
};
let n_mc = m.div_ceil(mc); let n_jobs_max = n_mc * n_nt_max;
let jobs_per_worker = n_jobs_max.div_ceil(n_threads.max(1));
let reuse_cols = jobs_per_worker.min(n_nt_max) * nr;
let pack_stride = cache::lhs_pack_stride_bytes();
let lhs_step_bytes = csa
.unsigned_abs()
.saturating_mul(core::mem::size_of::<Fam::Lhs>());
let strided_lhs = rsa == 1
&& lhs_step_bytes >= pack_stride
&& lhs_step_bytes.saturating_mul(kc.min(k)) >= cache::lhs_pack_span_bytes()
&& n_nt_max >= tuning::lhs_pack_reuse();
let want_pack_lhs =
reuse_cols > tuning::lhs_pack_threshold() || strided_lhs || Fam::FORCE_PACK_LHS;
const SHARED_A_MIN_THREADS: usize = 16;
let shared_a = n_threads > 1
&& (rsa != 1 || want_pack_lhs)
&& (mnk >= tuning::shared_lhs_mnk() || n_threads >= SHARED_A_MIN_THREADS);
let prepacked = packed.is_some();
let pack_b = !prepacked && (m > tuning::rhs_pack_threshold() || Fam::FORCE_PACK_RHS);
let ws_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 prefetch_c = ws_bytes > cache::prefetch_ws_bytes();
let need_a_pack = rsa != 1 || !m.is_multiple_of(mr) || want_pack_lhs;
let a_full_region = mc
.next_multiple_of(mr)
.checked_mul(kc_pad_block)
.unwrap_or_else(|| {
panic!("gemmkit: GEMM {m}x{k}x{n} is too large; the LHS pack region size overflows usize")
});
let b_full_elems = nc
.next_multiple_of(nr)
.checked_mul(kc_pad_block)
.unwrap_or_else(|| {
panic!("gemmkit: GEMM {m}x{k}x{n} is too large; the RHS pack region size overflows usize")
});
let (a_per_region, a_regions) = if need_a_pack {
(a_full_region, if shared_a { n_mc } else { n_threads })
} else {
(0, 0)
};
let b_elems = if pack_b { b_full_elems } else { 0 };
let regions: Regions<Fam::Lhs> = if need_a_pack || pack_b {
ws.regions::<Fam::Lhs>(a_per_region, a_regions, b_elems)
} else {
Regions {
a_base: core::ptr::null_mut(),
a_stride: 0,
b_base: core::ptr::null_mut(),
}
};
let a_base = Ptr(regions.a_base);
let a_stride = regions.a_stride;
let b_base = match packed {
Some((pb, _, _)) => Ptr(pb as *mut Fam::Rhs),
None => Ptr(regions.b_base as *mut Fam::Rhs),
};
let a = Ptr(a as *mut Fam::Lhs);
let b = Ptr(b as *mut Fam::Rhs);
let c = Ptr(c);
let ash = alpha_status(alpha);
let mut jc_off = 0usize;
let mut jc = 0;
while jc < n {
let nc_eff = core::cmp::min(nc, n - jc);
let n_nt = nc_eff.div_ceil(nr);
let n_jobs = n_mc * n_nt;
let grain = if shared_a {
parallel::job_grain(n_jobs, n_threads)
} else if (rsa != 1 || want_pack_lhs) && n_mc >= n_threads {
parallel::packed_block_grain(n_nt, n_mc, n_threads)
} else {
parallel::job_grain(n_jobs, n_threads)
};
let mut pc = 0;
while pc < k {
let kc_eff = core::cmp::min(kc, k - pc);
let kc_eff_pad = kc_eff.next_multiple_of(q_depth);
let first = pc == 0;
let last = pc + kc_eff >= k;
let beta_eff = if first { beta } else { Fam::Acc::ONE };
let bst = if first {
beta_status(beta)
} else {
BetaStatus::One
};
if pack_b {
let bcur = parallel::JobCursor::new(n_nt, parallel::job_grain(n_nt, n_threads));
parallel::for_each_worker(n_threads, |_tid| {
let (b, b_base) = (b, b_base);
while let Some((s, e)) = bcur.next_chunk() {
for jt in s..e {
let col = jc + jt * nr;
let nr_eff = core::cmp::min(nr, nc_eff - jt * nr);
let dst = b_base.0.add(jt * kc_eff_pad * nr);
let src = b.0.offset(pc as isize * rsb + col as isize * csb)
as *const Fam::Rhs;
Fam::pack_rhs(dst, src, rsb, csb, kc_eff, nr_eff, nr);
}
}
});
}
if shared_a {
let acur = parallel::JobCursor::new(n_mc, parallel::job_grain(n_mc, n_threads));
parallel::for_each_worker(n_threads, |_tid| {
let (a, a_base) = (a, a_base);
while let Some((s, e)) = acur.next_chunk() {
for ic_idx in s..e {
let ic = ic_idx * mc;
let mc_eff = core::cmp::min(mc, m - ic);
let dst = a_base.0.add(ic_idx * a_stride);
let src = a.0.offset(ic as isize * rsa + pc as isize * csa)
as *const Fam::Lhs;
Fam::pack_lhs(dst, src, rsa, csa, mc_eff, kc_eff, mr);
}
}
});
}
let cur = parallel::JobCursor::new(n_jobs, grain);
parallel::for_each_worker(n_threads, |tid| {
let (a, b, c, a_base, b_base, epi) = (a, b, c, a_base, b_base, epi);
let a_buf = if shared_a {
core::ptr::null_mut::<Fam::Lhs>()
} else {
a_base.0.add(tid * a_stride)
};
let mut scratch = [const { MaybeUninit::<Fam::Acc>::uninit() }; SCRATCH_LEN];
let scratch_ptr = scratch.as_mut_ptr() as *mut Fam::Acc;
let mut cur_ic = usize::MAX;
let mut a_panel_base: *const Fam::Lhs = core::ptr::null();
let mut a_cs: isize = 0;
let mut packed_a = false;
while let Some((start, end)) = cur.next_chunk() {
for q in start..end {
let ic_idx = q / n_nt;
let jt = q % n_nt;
let ic = ic_idx * mc;
let mc_eff = core::cmp::min(mc, m - ic);
if ic_idx != cur_ic {
cur_ic = ic_idx;
if shared_a {
a_panel_base =
a_base.0.add(ic_idx * a_stride) as *const Fam::Lhs;
a_cs = mr as isize;
packed_a = true;
} else {
let src = a.0.offset(ic as isize * rsa + pc as isize * csa)
as *const Fam::Lhs;
if do_pack_lhs(rsa, mc_eff, mr, want_pack_lhs) {
Fam::pack_lhs(a_buf, src, rsa, csa, mc_eff, kc_eff, mr);
a_panel_base = a_buf as *const Fam::Lhs;
a_cs = mr as isize;
packed_a = true;
} else {
a_panel_base = src;
a_cs = csa;
packed_a = false;
}
}
}
let col = jc + jt * nr;
let nr_eff = core::cmp::min(nr, nc_eff - jt * nr);
let (bpan, b_rs_k, b_cs_k) = if prepacked {
(
b_base.0.add(jc_off + nr * (n_nt * pc + jt * kc_eff_pad))
as *const Fam::Rhs,
nr as isize,
1,
)
} else if pack_b {
(
b_base.0.add(jt * kc_eff_pad * nr) as *const Fam::Rhs,
nr as isize,
1,
)
} else {
(
b.0.offset(pc as isize * rsb + col as isize * csb)
as *const Fam::Rhs,
rsb,
csb,
)
};
simd.vectorize(|| {
let mut ir = 0;
while ir < mc_eff {
let mr_eff = core::cmp::min(mr, mc_eff - ir);
let apan = if packed_a {
a_panel_base.add((ir / mr) * mr * kc_eff_pad)
} else {
a_panel_base.offset(ir as isize * rsa)
};
let cptr =
c.0.offset((ic + ir) as isize * rsc + col as isize * csc);
if prefetch_c {
prefetch_c_tile(cptr, rsc, csc, mr_eff, nr_eff);
}
Fam::microkernel_epi::<S, E, MR_REG, NR>(
simd,
kc_eff,
alpha,
beta_eff,
ash,
bst,
apan,
a_cs,
bpan,
b_rs_k,
b_cs_k,
cptr,
rsc,
csc,
mr_eff,
nr_eff,
ic + ir,
col,
last,
&epi,
scratch_ptr,
);
ir += mr;
}
});
}
}
});
pc += kc;
}
jc_off += n_nt * nr * k.next_multiple_of(q_depth);
jc += nc;
}
}
}