use crate::error::TruenoError;
use super::compute::{gemm_blis, gemm_blis_with_prepacked_b};
use super::prepacked::PrepackedB;
#[cfg(feature = "parallel")]
use super::{MC, MR};
#[derive(Debug, Clone)]
pub struct HeijunkaScheduler {
pub num_threads: usize,
pub variance_threshold: f32,
}
impl Default for HeijunkaScheduler {
fn default() -> Self {
#[cfg(feature = "parallel")]
let threads = rayon::current_num_threads();
#[cfg(not(feature = "parallel"))]
let threads = 1;
Self {
num_threads: threads,
variance_threshold: 0.05, }
}
}
impl HeijunkaScheduler {
pub fn partition_m(&self, m: usize, mc: usize) -> Vec<std::ops::Range<usize>> {
let num_blocks = (m + mc - 1) / mc;
let blocks_per_thread = num_blocks / self.num_threads;
let remainder = num_blocks % self.num_threads;
let mut partitions = Vec::with_capacity(self.num_threads);
let mut start_block = 0;
for t in 0..self.num_threads {
let extra = if t < remainder { 1 } else { 0 };
let thread_blocks = blocks_per_thread + extra;
let start_row = start_block * mc;
let end_row = ((start_block + thread_blocks) * mc).min(m);
if start_row < end_row {
partitions.push(start_row..end_row);
}
start_block += thread_blocks;
}
partitions
}
}
#[cfg(feature = "parallel")]
pub(crate) fn gemm_should_run_serial(m: usize, n: usize, k: usize) -> bool {
let flops = m * n * k;
flops < 8_000_000 || (flops < 64_000_000 && n < 192)
}
#[cfg(feature = "parallel")]
pub fn gemm_blis_parallel(
m: usize,
n: usize,
k: usize,
a: &[f32],
b: &[f32],
c: &mut [f32],
) -> Result<(), TruenoError> {
use rayon::prelude::*;
contract_pre_amdahl_speedup!();
if a.len() != m * k || b.len() != k * n || c.len() != m * n {
return Err(TruenoError::InvalidInput("Dimension mismatch".to_string()));
}
let flops = m * n * k;
if gemm_should_run_serial(m, n, k) {
return gemm_blis(m, n, k, a, b, c, None);
}
let phys_cores = num_cpus::get_physical();
let max_threads = if flops < 64_000_000 {
2.min(phys_cores)
} else if flops < 512_000_000 {
4.min(phys_cores)
} else if flops < 4_000_000_000 {
8.min(phys_cores)
} else {
(phys_cores / 2).max(8).min(phys_cores)
};
let mut scheduler = HeijunkaScheduler::default();
scheduler.num_threads = scheduler.num_threads.min(max_threads);
let ps = if m <= MC { MR.max(m / scheduler.num_threads) } else { MC };
let partitions = scheduler.partition_m(m, ps);
let c_ptr = c.as_mut_ptr() as usize;
partitions.into_par_iter().for_each(|m_range| {
let m_local = m_range.len();
let m_start = m_range.start;
let a_local = &a[m_start * k..(m_start + m_local) * k];
let c_local = unsafe {
let ptr = c_ptr as *mut f32;
std::slice::from_raw_parts_mut(ptr.add(m_start * n), m_local * n)
};
let _ = gemm_blis(m_local, n, k, a_local, b, c_local, None);
});
Ok(())
}
#[cfg(feature = "parallel")]
pub fn gemm_blis_parallel_shared_b(
m: usize,
n: usize,
k: usize,
a: &[f32],
b: &[f32],
c: &mut [f32],
) -> Result<(), TruenoError> {
use rayon::prelude::*;
if a.len() != m * k || b.len() != k * n || c.len() != m * n {
return Err(TruenoError::InvalidInput("Dimension mismatch".to_string()));
}
let flops = m * n * k;
if !shared_b_path_available(flops) {
return gemm_blis(m, n, k, a, b, c, None);
}
let num_threads = shared_b_thread_count(flops).min(rayon::current_num_threads());
let blk = super::cache_topology::blocking_8x32();
let geo = SharedBGeometry {
mr: blk.mr, nr: blk.nr, mc: blk.mc.min(m),
nc: blk.nc.min(n),
kc: blk.kc,
};
let b_panels = geo.nc.div_ceil(geo.nr);
let mut packed_b = vec![0.0f32; b_panels * geo.nr * geo.kc];
let c_ptr = c.as_mut_ptr() as usize;
for jc in (0..n).step_by(geo.nc) {
let nc_block = geo.nc.min(n - jc);
for pc in (0..k).step_by(geo.kc) {
let kc_block = geo.kc.min(k - pc);
super::compute::pack_b_block_generic(
b,
n,
pc,
jc,
kc_block,
nc_block,
geo.nr,
&mut packed_b,
);
let block = SharedBBlock {
a,
k,
n,
c_ptr,
jc,
pc,
nc_block,
kc_block,
geo: &geo,
shared_b: &packed_b,
};
let m_per_thread = m.div_ceil(num_threads).div_ceil(geo.mr) * geo.mr;
(0..num_threads).into_par_iter().for_each(|tid| {
let ic_start = tid * m_per_thread;
if ic_start < m {
block.run_slice(ic_start, (ic_start + m_per_thread).min(m), m_per_thread);
}
});
}
}
Ok(())
}
#[cfg(feature = "parallel")]
fn shared_b_path_available(flops: usize) -> bool {
if flops < 8_000_000 {
return false;
}
#[cfg(target_arch = "x86_64")]
{
std::arch::is_x86_feature_detected!("avx512f")
}
#[cfg(not(target_arch = "x86_64"))]
{
false
}
}
#[cfg(feature = "parallel")]
fn shared_b_thread_count(flops: usize) -> usize {
let phys_cores = num_cpus::get_physical();
if flops < 64_000_000 {
2.min(phys_cores)
} else if flops < 512_000_000 {
4.min(phys_cores)
} else {
(phys_cores / 2).max(8).min(phys_cores)
}
}
#[cfg(feature = "parallel")]
struct SharedBGeometry {
mr: usize,
nr: usize,
mc: usize,
nc: usize,
kc: usize,
}
#[cfg(feature = "parallel")]
struct SharedBBlock<'a> {
a: &'a [f32],
k: usize,
n: usize,
c_ptr: usize,
jc: usize,
pc: usize,
nc_block: usize,
kc_block: usize,
geo: &'a SharedBGeometry,
shared_b: &'a [f32],
}
#[cfg(feature = "parallel")]
impl SharedBBlock<'_> {
fn run_slice(&self, ic_start: usize, ic_end: usize, m_per_thread: usize) {
thread_local! {
static TL_A: std::cell::RefCell<Vec<f32>> =
const { std::cell::RefCell::new(Vec::new()) };
}
TL_A.with(|tl| {
let geo = self.geo;
let needed = m_per_thread.div_ceil(geo.mr) * geo.mr * self.kc_block;
let mut packed_a = tl.borrow_mut();
if packed_a.len() < needed {
packed_a.resize(needed, 0.0);
}
for ic in (ic_start..ic_end).step_by(geo.mc) {
let mc_block = geo.mc.min(ic_end - ic);
super::packing::pack_a_block(
self.a,
self.k,
ic,
self.pc,
mc_block,
self.kc_block,
&mut packed_a,
);
self.run_panels(&packed_a, ic, mc_block);
}
});
}
fn run_panels(&self, packed_a: &[f32], ic: usize, mc_block: usize) {
let geo = self.geo;
let panels_n = self.nc_block.div_ceil(geo.nr);
for ir_panel in 0..mc_block.div_ceil(geo.mr) {
let ir = ir_panel * geo.mr;
let mr_block = geo.mr.min(mc_block - ir);
for jr_panel in 0..panels_n {
let jr = jr_panel * geo.nr;
let nr_block = geo.nr.min(self.nc_block - jr);
let a_panel = &packed_a[ir_panel * geo.mr * self.kc_block..];
let b_panel = &self.shared_b[jr_panel * geo.nr * self.kc_block..];
self.tile(a_panel, b_panel, ic + ir, self.jc + jr, mr_block, nr_block);
}
}
}
fn tile(
&self,
a_panel: &[f32],
b_panel: &[f32],
row: usize,
col: usize,
mr_block: usize,
nr_block: usize,
) {
#[cfg(target_arch = "x86_64")]
if mr_block == 8 && nr_block == 32 {
unsafe {
super::compute::avx512_microkernel_8x32_rowmajor(
self.kc_block,
a_panel.as_ptr(),
b_panel.as_ptr(),
(self.c_ptr as *mut f32).add(row * self.n + col),
self.n,
);
}
return;
}
self.scalar_tile(a_panel, b_panel, row, col, mr_block, nr_block);
}
fn scalar_tile(
&self,
a_panel: &[f32],
b_panel: &[f32],
row: usize,
col: usize,
mr_block: usize,
nr_block: usize,
) {
let geo = self.geo;
for ir_local in 0..mr_block {
for jr_local in 0..nr_block {
let mut sum = 0.0f32;
for p in 0..self.kc_block {
sum += a_panel[p * geo.mr + ir_local] * b_panel[p * geo.nr + jr_local];
}
unsafe {
let c = self.c_ptr as *mut f32;
*c.add((row + ir_local) * self.n + (col + jr_local)) += sum;
}
}
}
}
}
#[cfg(not(feature = "parallel"))]
pub fn gemm_blis_parallel(
m: usize,
n: usize,
k: usize,
a: &[f32],
b: &[f32],
c: &mut [f32],
) -> Result<(), TruenoError> {
gemm_blis(m, n, k, a, b, c, None)
}
#[cfg(feature = "parallel")]
pub fn gemm_blis_parallel_with_prepacked_b(
m: usize,
n: usize,
k: usize,
a: &[f32],
prepacked_b: &PrepackedB,
c: &mut [f32],
) -> Result<(), TruenoError> {
use rayon::prelude::*;
if a.len() != m * k || c.len() != m * n {
return Err(TruenoError::InvalidInput("Dimension mismatch".to_string()));
}
if prepacked_b.k != k || prepacked_b.n != n {
return Err(TruenoError::InvalidInput(format!(
"PrepackedB dimension mismatch: expected ({}, {}), got ({}, {})",
k, n, prepacked_b.k, prepacked_b.n
)));
}
if m * n * k < 1_000_000 {
return gemm_blis_with_prepacked_b(m, n, k, a, prepacked_b, c, None);
}
let scheduler = HeijunkaScheduler::default();
let partitions = scheduler.partition_m(m, MC);
let c_ptr = c.as_mut_ptr() as usize;
partitions.into_par_iter().for_each(|m_range| {
let m_local = m_range.len();
let m_start = m_range.start;
let a_local = &a[m_start * k..(m_start + m_local) * k];
let c_local = unsafe {
let ptr = c_ptr as *mut f32;
std::slice::from_raw_parts_mut(ptr.add(m_start * n), m_local * n)
};
let _ = gemm_blis_with_prepacked_b(m_local, n, k, a_local, prepacked_b, c_local, None);
});
Ok(())
}
#[cfg(not(feature = "parallel"))]
pub fn gemm_blis_parallel_with_prepacked_b(
m: usize,
n: usize,
k: usize,
a: &[f32],
prepacked_b: &PrepackedB,
c: &mut [f32],
) -> Result<(), TruenoError> {
gemm_blis_with_prepacked_b(m, n, k, a, prepacked_b, c, None)
}