use crate::kernel::FloatGemm;
#[cfg(feature = "half")]
use crate::kernel::MixedGemm;
use crate::kernel::epilogue::Epilogue;
use crate::parallel::{self, JobCursor, Parallelism, Ptr};
use crate::scalar::Float;
#[cfg(feature = "half")]
use crate::scalar::NarrowFloat;
#[cfg(any(feature = "half", feature = "int8"))]
use crate::simd::KernelSimd;
use crate::simd::SimdOps;
use crate::workspace::Workspace;
const MT: usize = 4;
const NT: usize = 4;
fn tile_sweep(
m: usize,
n: usize,
bytes: usize,
par: Parallelism,
body: impl Fn(usize, usize) + Copy + Send + Sync,
) {
let n_row_tiles = m.div_ceil(MT);
let n_col_tiles = n.div_ceil(NT);
let n_tiles = n_row_tiles * n_col_tiles;
let n_threads = par.resolve_bandwidth(bytes, n_tiles);
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);
}
});
}
#[inline]
fn packed_line_stride<T>(k: usize) -> usize {
let lane = (64 / core::mem::size_of::<T>().max(1)).max(1);
let lines = k.div_ceil(lane).max(1);
let odd_lines = if lines.is_multiple_of(2) {
lines + 1
} else {
lines
};
odd_lines * lane
}
#[inline]
#[allow(clippy::too_many_arguments)]
unsafe fn pack_k_contiguous<T: Copy>(
dst: *mut T,
src: *const T,
lead: usize,
t_begin: usize,
t_end: usize,
dst_stride: usize,
lead_stride: isize,
depth_stride: isize,
) {
unsafe {
let tile = crate::tuning::pack_transpose_tile();
let mut t0 = t_begin;
while t0 < t_end {
let te = core::cmp::min(t0.saturating_add(tile), t_end);
for t in t0..te {
let col = src.offset(t as isize * depth_stride);
for l in 0..lead {
*dst.add(l * dst_stride + t) = *col.offset(l as isize * lead_stride);
}
}
t0 = te;
}
}
}
#[allow(clippy::too_many_arguments)]
unsafe fn pack_k_contiguous_par<T: Copy>(
dst: *mut T,
src: *const T,
lead: usize,
k: usize,
dst_stride: usize,
lead_stride: isize,
depth_stride: isize,
par: Parallelism,
) -> usize {
let bytes = lead
.saturating_mul(k)
.saturating_mul(2)
.saturating_mul(core::mem::size_of::<T>());
let n_threads = par.resolve_bandwidth(bytes, k);
if n_threads <= 1 {
unsafe {
pack_k_contiguous::<T>(dst, src, lead, 0, k, dst_stride, lead_stride, depth_stride)
};
return 1;
}
let n_chunks = k.min(
n_threads
.saturating_mul(crate::tuning::parallel_oversample())
.max(1),
);
let chunk = k.div_ceil(n_chunks.max(1));
let (dst, src) = (Ptr(dst), Ptr(src as *mut T));
let cur = JobCursor::new(n_chunks, 1);
parallel::for_each_worker(n_threads, |_tid| {
let (dst, src) = (dst, src);
while let Some((s, e)) = cur.next_chunk() {
let t_begin = s * chunk;
let t_end = core::cmp::min(e * chunk, k);
unsafe {
pack_k_contiguous::<T>(
dst.0,
src.0 as *const T,
lead,
t_begin,
t_end,
dst_stride,
lead_stride,
depth_stride,
)
};
}
});
n_threads
}
#[allow(clippy::too_many_arguments)]
unsafe fn prepack_operands<T: Copy>(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
a: *const T,
rsa: isize,
csa: isize,
b: *const T,
rsb: isize,
csb: isize,
par: Parallelism,
) -> (*const T, isize, isize, *const T, isize, isize) {
let pack_a = csa != 1;
let pack_b = rsb != 1;
if !pack_a && !pack_b {
return (a, rsa, csa, b, rsb, csb);
}
let stride = packed_line_stride::<T>(k);
unsafe {
let a_elems = if pack_a { m.saturating_mul(stride) } else { 0 };
let b_elems = if pack_b { n.saturating_mul(stride) } else { 0 };
let r = ws.regions::<T>(a_elems, 1, b_elems);
let (mut a, mut rsa, mut csa) = (a, rsa, csa);
let (mut b, mut rsb, mut csb) = (b, rsb, csb);
if pack_a {
pack_k_contiguous_par::<T>(r.a_base, a, m, k, stride, rsa, csa, par);
a = r.a_base;
rsa = stride as isize;
csa = 1;
}
if pack_b {
pack_k_contiguous_par::<T>(r.b_base, b, n, k, stride, csb, rsb, par);
b = r.b_base;
rsb = 1;
csb = stride as isize;
}
(a, rsa, csa, b, rsb, csb)
}
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_epi<T, S, E>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
ws: &mut Workspace,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
b: *const T,
rsb: isize,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
epi: &E,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
let epi = *epi;
unsafe {
let (a, rsa, csa, b, rsb, csb) =
prepack_operands::<T>(ws, m, k, n, a, rsa, csa, b, rsb, csb, par);
debug_assert!(
csa == 1 && rsb == 1,
"small_mn kernel requires A rows / B cols unit-stride along k"
);
let n_row_tiles = m.div_ceil(MT);
let sizeof = core::mem::size_of::<T>();
let bytes = m
.saturating_mul(k)
.saturating_add(k.saturating_mul(n))
.saturating_add(m.saturating_mul(n))
.saturating_mul(sizeof);
let a = Ptr(a as *mut T);
let b = Ptr(b as *mut T);
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 T;
let b = b.0 as *const T;
let c = c.0;
simd.vectorize(|| {
for q in q_start..q_end {
let it = q % n_row_tiles;
let jt = q / n_row_tiles;
let i0 = it * MT;
let j0 = jt * NT;
let mi = core::cmp::min(MT, m - i0);
let nj = core::cmp::min(NT, n - j0);
if mi == MT && nj == NT {
full_tile::<T, S, E, MT, NT>(
simd, k, i0, j0, alpha, a, rsa, b, csb, beta, c, rsc, csc, &epi,
);
} else {
for cc in 0..nj {
for ir in 0..mi {
cell_dot::<T, S, E>(
simd,
k,
i0 + ir,
j0 + cc,
alpha,
a,
rsa,
b,
csb,
beta,
c,
rsc,
csc,
&epi,
);
}
}
}
}
});
};
tile_sweep(m, n, bytes, par, body);
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn full_tile<T, S, E, const MT: usize, const NT: usize>(
simd: S,
k: usize,
i0: usize,
j0: usize,
alpha: T,
a: *const T,
rsa: isize,
b: *const T,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
epi: &E,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
let rows: [*const T; MT] = core::array::from_fn(|r| a.offset((i0 + r) as isize * rsa));
let cols: [*const T; NT] = core::array::from_fn(|cc| b.offset((j0 + cc) as isize * csb));
let mut acc = [[simd.zero(); MT]; NT];
let mut kk = 0;
while kk + lanes <= k {
let av: [S::Reg; MT] = core::array::from_fn(|r| simd.loadu(rows[r].add(kk)));
for cc in 0..NT {
let bv = simd.loadu(cols[cc].add(kk));
for r in 0..MT {
acc[cc][r] = simd.mul_add(av[r], bv, acc[cc][r]);
}
}
kk += lanes;
}
for cc in 0..NT {
for r in 0..MT {
let mut dot = simd.reduce_sum(acc[cc][r]);
let mut t = kk;
while t < k {
dot = (*rows[r].add(t)).mul_add(*cols[cc].add(t), dot);
t += 1;
}
let cp = c.offset((i0 + r) as isize * rsc + (j0 + cc) as isize * csc);
let ov = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*cp
} else {
beta * *cp
};
let out = alpha.mul_add(dot, ov);
*cp = if E::IS_IDENTITY {
out
} else {
epi.apply(out, i0 + r, j0 + cc)
};
}
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn cell_dot<T, S, E>(
simd: S,
k: usize,
i: usize,
j: usize,
alpha: T,
a: *const T,
rsa: isize,
b: *const T,
csb: isize,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
epi: &E,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
unsafe {
let row = a.offset(i as isize * rsa); let col = b.offset(j as isize * csb); let dot = super::dot_contiguous::<T, S>(simd, k, row, col);
let cp = c.offset(i as isize * rsc + j as isize * csc);
let ov = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*cp
} else {
beta * *cp
};
let out = alpha.mul_add(dot, ov);
*cp = if E::IS_IDENTITY {
out
} else {
epi.apply(out, i, j)
};
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_mixed_epi<N, S, E>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
ws: &mut Workspace,
alpha: f32,
a: *const N,
rsa: isize,
csa: isize,
b: *const N,
rsb: isize,
csb: isize,
beta: f32,
c: *mut N,
rsc: isize,
csc: isize,
epi: &E,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
E: Epilogue<MixedGemm<N>>,
{
let epi = *epi;
unsafe {
let (a, rsa, csa, b, rsb, csb) =
prepack_operands::<N>(ws, m, k, n, a, rsa, csa, b, rsb, csb, par);
debug_assert!(
csa == 1 && rsb == 1,
"small_mn kernel requires A rows / B cols unit-stride along k"
);
let n_row_tiles = m.div_ceil(MT);
let sizeof = core::mem::size_of::<N>();
let bytes = m
.saturating_mul(k)
.saturating_add(k.saturating_mul(n))
.saturating_add(m.saturating_mul(n))
.saturating_mul(sizeof);
let a = Ptr(a as *mut N);
let b = Ptr(b as *mut N);
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 N;
let b = b.0 as *const N;
let c = c.0;
simd.vectorize(|| {
for q in q_start..q_end {
let it = q % n_row_tiles;
let jt = q / n_row_tiles;
let i0 = it * MT;
let j0 = jt * NT;
let mi = core::cmp::min(MT, m - i0);
let nj = core::cmp::min(NT, n - j0);
if mi == MT && nj == NT {
full_tile_mixed::<N, S, E, MT, NT>(
simd, k, i0, j0, alpha, a, rsa, b, csb, beta, c, rsc, csc, &epi,
);
} else {
for cc in 0..nj {
for ir in 0..mi {
cell_dot_mixed::<N, S, E>(
simd,
k,
i0 + ir,
j0 + cc,
alpha,
a,
rsa,
b,
csb,
beta,
c,
rsc,
csc,
&epi,
);
}
}
}
}
});
};
tile_sweep(m, n, bytes, par, body);
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn full_tile_mixed<N, S, E, const MT: usize, const NT: usize>(
simd: S,
k: usize,
i0: usize,
j0: usize,
alpha: f32,
a: *const N,
rsa: isize,
b: *const N,
csb: isize,
beta: f32,
c: *mut N,
rsc: isize,
csc: isize,
epi: &E,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
E: Epilogue<MixedGemm<N>>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let rows: [*const N; MT] = core::array::from_fn(|r| a.offset((i0 + r) as isize * rsa));
let cols: [*const N; NT] = core::array::from_fn(|cc| b.offset((j0 + cc) as isize * csb));
let mut acc: [[<S as SimdOps<f32>>::Reg; MT]; NT] = [[simd.zero(); MT]; NT];
let mut kk = 0;
while kk + lanes <= k {
let av: [<S as SimdOps<f32>>::Reg; MT] =
core::array::from_fn(|r| simd.load_lhs(rows[r].add(kk)));
for cc in 0..NT {
let bv = simd.load_lhs(cols[cc].add(kk));
for r in 0..MT {
acc[cc][r] = simd.mul_add(av[r], bv, acc[cc][r]);
}
}
kk += lanes;
}
for cc in 0..NT {
for r in 0..MT {
let mut dot = simd.reduce_sum(acc[cc][r]);
let mut t = kk;
while t < k {
dot += (*rows[r].add(t)).widen() * (*cols[cc].add(t)).widen();
t += 1;
}
let cp = c.offset((i0 + r) as isize * rsc + (j0 + cc) as isize * csc);
let ov = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*cp).widen()
} else {
beta * (*cp).widen()
};
let out = alpha * dot + ov;
*cp = if E::IS_IDENTITY {
N::narrow(out)
} else {
epi.apply(out, i0 + r, j0 + cc)
};
}
}
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn cell_dot_mixed<N, S, E>(
simd: S,
k: usize,
i: usize,
j: usize,
alpha: f32,
a: *const N,
rsa: isize,
b: *const N,
csb: isize,
beta: f32,
c: *mut N,
rsc: isize,
csc: isize,
epi: &E,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
E: Epilogue<MixedGemm<N>>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let row = a.offset(i as isize * rsa);
let col = b.offset(j as isize * csb);
let mut acc = simd.zero();
let mut kk = 0;
while kk + lanes <= k {
acc = simd.mul_add(simd.load_lhs(row.add(kk)), simd.load_lhs(col.add(kk)), acc);
kk += lanes;
}
let mut dot = simd.reduce_sum(acc);
while kk < k {
dot += (*row.add(kk)).widen() * (*col.add(kk)).widen();
kk += 1;
}
let cp = c.offset(i as isize * rsc + j as isize * csc);
let ov = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*cp).widen()
} else {
beta * (*cp).widen()
};
let out = alpha * dot + ov;
*cp = if E::IS_IDENTITY {
N::narrow(out)
} else {
epi.apply(out, i, j)
};
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_int<S>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
ws: &mut Workspace,
alpha: i32,
a: *const i8,
rsa: isize,
csa: isize,
b: *const i8,
rsb: isize,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let (a, rsa, csa, b, rsb, csb) =
prepack_operands::<i8>(ws, m, k, n, a, rsa, csa, b, rsb, csb, par);
debug_assert!(
csa == 1 && rsb == 1,
"small_mn kernel requires A rows / B cols unit-stride along k"
);
let n_row_tiles = m.div_ceil(MT);
let bytes = m
.saturating_mul(k)
.saturating_add(k.saturating_mul(n))
.saturating_mul(core::mem::size_of::<i8>())
.saturating_add(
m.saturating_mul(n)
.saturating_mul(core::mem::size_of::<i32>()),
);
let a = Ptr(a as *mut i8);
let b = Ptr(b as *mut i8);
let c = Ptr(c);
let body = move |q_start: usize, q_end: usize| {
let (a, b, c) = (a, b, c);
let a = a.0 as *const i8;
let b = b.0 as *const i8;
let c = c.0;
simd.vectorize(|| {
for q in q_start..q_end {
let it = q % n_row_tiles;
let jt = q / n_row_tiles;
let i0 = it * MT;
let j0 = jt * NT;
let mi = core::cmp::min(MT, m - i0);
let nj = core::cmp::min(NT, n - j0);
if mi == MT && nj == NT {
full_tile_int::<S, MT, NT>(
simd, k, i0, j0, alpha, a, rsa, b, csb, beta, c, rsc, csc,
);
} else {
for cc in 0..nj {
for ir in 0..mi {
cell_dot_int::<S>(
simd,
k,
i0 + ir,
j0 + cc,
alpha,
a,
rsa,
b,
csb,
beta,
c,
rsc,
csc,
);
}
}
}
}
});
};
tile_sweep(m, n, bytes, par, body);
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn full_tile_int<S, const MT: usize, const NT: usize>(
simd: S,
k: usize,
i0: usize,
j0: usize,
alpha: i32,
a: *const i8,
rsa: isize,
b: *const i8,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let lanes = <S as SimdOps<i32>>::LANES;
let rows: [*const i8; MT] = core::array::from_fn(|r| a.offset((i0 + r) as isize * rsa));
let cols: [*const i8; NT] = core::array::from_fn(|cc| b.offset((j0 + cc) as isize * csb));
let mut acc: [[<S as SimdOps<i32>>::Reg; MT]; NT] = [[simd.zero(); MT]; NT];
let mut kk = 0;
while kk + lanes <= k {
let av: [<S as SimdOps<i32>>::Reg; MT] = core::array::from_fn(|r| {
<S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, rows[r].add(kk))
});
for cc in 0..NT {
let bv = <S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, cols[cc].add(kk));
for r in 0..MT {
acc[cc][r] = simd.mul_add(av[r], bv, acc[cc][r]);
}
}
kk += lanes;
}
for cc in 0..NT {
for r in 0..MT {
let mut dot = simd.reduce_sum(acc[cc][r]);
let mut t = kk;
while t < k {
dot = dot.wrapping_add(
(*rows[r].add(t) as i32).wrapping_mul(*cols[cc].add(t) as i32),
);
t += 1;
}
let cp = c.offset((i0 + r) as isize * rsc + (j0 + cc) as isize * csc);
let ov = if beta == 0 {
0
} else if beta == 1 {
*cp
} else {
beta.wrapping_mul(*cp)
};
*cp = alpha.wrapping_mul(dot).wrapping_add(ov);
}
}
}
}
#[cfg(feature = "int8")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn cell_dot_int<S>(
simd: S,
k: usize,
i: usize,
j: usize,
alpha: i32,
a: *const i8,
rsa: isize,
b: *const i8,
csb: isize,
beta: i32,
c: *mut i32,
rsc: isize,
csc: isize,
) where
S: KernelSimd<i8, i8, i32, i32>,
{
unsafe {
let lanes = <S as SimdOps<i32>>::LANES;
let row = a.offset(i as isize * rsa);
let col = b.offset(j as isize * csb);
let mut acc = simd.zero();
let mut kk = 0;
while kk + lanes <= k {
acc = simd.mul_add(
<S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, row.add(kk)),
<S as KernelSimd<i8, i8, i32, i32>>::load_lhs(simd, col.add(kk)),
acc,
);
kk += lanes;
}
let mut dot = simd.reduce_sum(acc);
while kk < k {
dot = dot.wrapping_add((*row.add(kk) as i32).wrapping_mul(*col.add(kk) as i32));
kk += 1;
}
let cp = c.offset(i as isize * rsc + j as isize * csc);
let ov = if beta == 0 {
0
} else if beta == 1 {
*cp
} else {
beta.wrapping_mul(*cp)
};
*cp = alpha.wrapping_mul(dot).wrapping_add(ov);
}
}
#[cfg(test)]
mod tests {
use super::{pack_k_contiguous, packed_line_stride};
fn fixture(lead: usize, k: usize) -> (Vec<f32>, usize, Vec<f32>) {
let depth_stride = lead as isize; let src: Vec<f32> = (0..lead * k).map(|i| i as f32 * 0.25 - 3.0).collect();
let stride = packed_line_stride::<f32>(k);
let mut whole = vec![f32::NAN; lead * stride];
unsafe {
pack_k_contiguous::<f32>(
whole.as_mut_ptr(),
src.as_ptr(),
lead,
0,
k,
stride,
1,
depth_stride,
)
};
(src, stride, whole)
}
#[test]
fn a_huge_transpose_tile_leaves_a_shifted_range_correct() {
let (lead, k) = (3usize, 8usize);
let (src, stride, whole) = fixture(lead, k);
let prev = crate::tuning::pack_transpose_tile();
crate::tuning::set_pack_transpose_tile(usize::MAX);
let mut got = vec![f32::NAN; lead * stride];
unsafe {
pack_k_contiguous::<f32>(
got.as_mut_ptr(),
src.as_ptr(),
lead,
k / 2,
k,
stride,
1,
lead as isize,
)
};
crate::tuning::set_pack_transpose_tile(prev);
for l in 0..lead {
let (a, b) = (&got[l * stride..][..k], &whole[l * stride..][..k]);
assert_eq!(a[k / 2..], b[k / 2..], "line {l}: shifted range");
assert!(a[..k / 2].iter().all(|v| v.is_nan()), "line {l}: wrote low");
}
}
#[test]
fn depth_ranges_compose_to_the_whole_copy() {
for &(lead, k) in &[(1usize, 33usize), (4, 37), (5, 64), (16, 100), (3, 7)] {
let (src, stride, whole) = fixture(lead, k);
let splits: [&[usize]; 4] = [
&[0, k],
&[0, 1, k],
&[0, 1, 7, 20.min(k), k],
&[0, 16.min(k), 32.min(k), 48.min(k), k],
];
for cuts in splits {
let mut got = vec![f32::NAN; lead * stride];
for w in cuts.windows(2) {
let (t0, t1) = (w[0], w[1]);
if t0 >= t1 {
continue; }
unsafe {
pack_k_contiguous::<f32>(
got.as_mut_ptr(),
src.as_ptr(),
lead,
t0,
t1,
stride,
1,
lead as isize,
)
};
}
for l in 0..lead {
let (a, b) = (&got[l * stride..][..k], &whole[l * stride..][..k]);
assert_eq!(a, b, "lead={lead} k={k} cuts={cuts:?} line {l}");
}
}
}
}
#[test]
#[cfg(feature = "parallel")]
fn parallel_copy_matches_the_serial_one() {
use super::pack_k_contiguous_par;
use crate::parallel::Parallelism;
let lead = 4usize;
let floor = crate::cache::gemv_parallel_floor_bytes();
let k_want = floor / (lead * 2 * 4) * 2 + 64;
let k = k_want.clamp(64, 1 << 20);
let (src, stride, whole) = fixture(lead, k);
let mut forked = false;
for par in [
Parallelism::Serial,
Parallelism::Rayon(0),
Parallelism::Rayon(4),
] {
let mut got = vec![f32::NAN; lead * stride];
let width = unsafe {
pack_k_contiguous_par::<f32>(
got.as_mut_ptr(),
src.as_ptr(),
lead,
k,
stride,
1,
lead as isize,
par,
)
};
forked |= width > 1;
for l in 0..lead {
let (a, b) = (&got[l * stride..][..k], &whole[l * stride..][..k]);
assert_eq!(a, b, "{par:?} width={width} k={k} line {l}");
}
}
let cores = std::thread::available_parallelism().map_or(1, |n| n.get());
assert!(
forked || cores <= 1 || k != k_want,
"no arm forked: the parallel copy is going untested (k={k} floor={floor})"
);
}
}