use crate::kernel::FloatGemm;
use crate::kernel::epilogue::Epilogue;
use crate::parallel::{self, JobCursor, Parallelism, Ptr};
use crate::scalar::Float;
#[cfg(feature = "half")]
use crate::scalar::NarrowFloat;
#[cfg(feature = "half")]
use crate::simd::KernelSimd;
use crate::simd::SimdOps;
const MB_REG: usize = 8;
#[inline]
fn row_sweep(
rows: usize,
block: usize,
n_threads: usize,
body: impl Fn(usize, usize) + Copy + Send + Sync,
) {
if n_threads <= 1 {
body(0, rows);
return;
}
let n_blocks = rows.div_ceil(block);
let cur = JobCursor::new(n_blocks, parallel::job_grain(n_blocks, n_threads));
parallel::for_each_worker(n_threads, |_tid| {
while let Some((bs, be)) = cur.next_chunk() {
let row_start = bs * block;
let row_end = core::cmp::min(be * block, rows);
body(row_start, row_end);
}
});
}
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_typed_epi<T, S, E>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
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>>,
{
unsafe {
if n == 1 {
core_epi::<T, S, E>(
simd, m, k, par, alpha, a, rsa, csa, b, rsb, beta, c, rsc, false, epi,
);
} else {
core_epi::<T, S, E>(
simd, n, k, par, alpha, b, csb, rsb, a, csa, beta, c, csc, true, epi,
);
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline]
unsafe fn core_epi<T, S, E>(
simd: S,
rows: usize,
k: usize,
par: Parallelism,
alpha: T,
mat: *const T,
mat_rs: isize,
mat_cs: isize,
vec: *const T,
vec_s: isize,
beta: T,
out: *mut T,
out_s: isize,
swap_rc: bool,
epi: &E,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
E: Epilogue<FloatGemm<T>>,
{
let epi = *epi;
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
let sizeof = core::mem::size_of::<T>();
let dot_legal = mat_cs == 1 && vec_s == 1;
let axpy = mat_rs == 1 && out_s == 1 && !axpy_yields_to_dot(rows, lanes, dot_legal);
let output_block = axpy && output_register_block(rows, sizeof, k);
let dot = !axpy && dot_legal;
let bytes_touched = (rows.saturating_mul(k) + k + rows).saturating_mul(sizeof);
let n_threads = if axpy && axpy_row_split_loses(rows) {
1
} else {
par.resolve_bandwidth(bytes_touched, rows)
};
let block = if output_block { MB_REG * lanes } else { lanes }.max(1);
let mat = Ptr(mat as *mut T);
let vec = Ptr(vec as *mut T);
let out = Ptr(out);
let body = move |row_start: usize, row_end: usize| {
let (mat, vec, out, epi) = (mat, vec, out, epi);
let mat = mat.0 as *const T;
let vec = vec.0 as *const T;
let out = out.0;
simd.vectorize(|| {
if output_block {
axpy_regblocked::<T, S>(
simd, row_start, row_end, k, alpha, mat, mat_cs, vec, vec_s, beta, out,
);
} else if axpy {
axpy_plain::<T, S>(
simd, row_start, row_end, k, alpha, mat, mat_cs, vec, vec_s, beta, out,
);
} else if dot {
dot_rows::<T, S>(
simd, row_start, row_end, k, alpha, mat, mat_rs, vec, beta, out, out_s,
);
} else {
strided_rows::<T, S>(
simd, row_start, row_end, k, alpha, mat, mat_rs, mat_cs, vec, vec_s, beta,
out, out_s,
);
}
if !E::IS_IDENTITY {
for i in row_start..row_end {
let op = out.offset(i as isize * out_s);
let (r, c) = if swap_rc { (0, i) } else { (i, 0) };
*op = epi.apply(*op, r, c);
}
}
});
};
row_sweep(rows, block, n_threads, body);
}
}
#[inline]
fn axpy_yields_to_dot(rows: usize, lanes: usize, dot_legal: bool) -> bool {
dot_legal && rows < lanes
}
#[inline]
fn axpy_row_split_loses(rows: usize) -> bool {
let floor = crate::tuning::gemv_axpy_par_min_rows();
floor != 0 && rows < floor
}
#[inline]
fn output_register_block(rows: usize, sizeof: usize, k: usize) -> bool {
k <= crate::tuning::k_stream_max()
&& rows.saturating_mul(sizeof) > crate::cache::gemv_regblock_engage_bytes()
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn axpy_regblocked<T, S>(
simd: S,
s: usize,
e: usize,
k: usize,
alpha: T,
mat: *const T,
mat_cs: isize,
vec: *const T,
vec_s: isize,
beta: T,
out: *mut T,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
{
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
let mb = MB_REG * lanes;
let mut i = s;
while i + mb <= e {
let mut acc = [simd.zero(); MB_REG];
if beta == T::ONE {
for (r, a) in acc.iter_mut().enumerate() {
*a = simd.loadu(out.add(i + r * lanes));
}
} else if beta != T::ZERO {
let bv = simd.splat(beta);
for (r, a) in acc.iter_mut().enumerate() {
*a = simd.mul(simd.loadu(out.add(i + r * lanes)), bv);
}
}
for kk in 0..k {
let sv = simd.splat(alpha * *vec.offset(kk as isize * vec_s));
let col = mat.offset(kk as isize * mat_cs).add(i);
for (r, a) in acc.iter_mut().enumerate() {
*a = simd.mul_add(simd.loadu(col.add(r * lanes)), sv, *a);
}
}
for (r, a) in acc.iter().enumerate() {
simd.storeu(out.add(i + r * lanes), *a);
}
i += mb;
}
while i + lanes <= e {
let mut acc = if beta == T::ONE {
simd.loadu(out.add(i))
} else if beta == T::ZERO {
simd.zero()
} else {
simd.mul(simd.loadu(out.add(i)), simd.splat(beta))
};
for kk in 0..k {
let sv = simd.splat(alpha * *vec.offset(kk as isize * vec_s));
acc = simd.mul_add(simd.loadu(mat.offset(kk as isize * mat_cs).add(i)), sv, acc);
}
simd.storeu(out.add(i), acc);
i += lanes;
}
while i < e {
let op = out.add(i);
let mut acc = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*op
} else {
beta * *op
};
for kk in 0..k {
let s = alpha * *vec.offset(kk as isize * vec_s);
acc = s.mul_add(*mat.offset(kk as isize * mat_cs).add(i), acc);
}
*op = acc;
i += 1;
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn axpy_plain<T, S>(
simd: S,
s: usize,
e: usize,
k: usize,
alpha: T,
mat: *const T,
mat_cs: isize,
vec: *const T,
vec_s: isize,
beta: T,
out: *mut T,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
{
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
for i in s..e {
let op = out.add(i);
if beta == T::ZERO {
*op = T::ZERO;
} else if beta != T::ONE {
*op = beta * *op;
}
}
const KB: usize = 4;
let mut kk = 0;
while kk + KB <= k {
let scal: [T; KB] =
core::array::from_fn(|j| alpha * *vec.offset((kk + j) as isize * vec_s));
let sv: [S::Reg; KB] = core::array::from_fn(|j| simd.splat(scal[j]));
let col: [*const T; KB] =
core::array::from_fn(|j| mat.offset((kk + j) as isize * mat_cs));
let mut i = s;
while i + lanes <= e {
let mut ov = simd.loadu(out.add(i));
for j in 0..KB {
ov = simd.mul_add(simd.loadu(col[j].add(i)), sv[j], ov);
}
simd.storeu(out.add(i), ov);
i += lanes;
}
while i < e {
let op = out.add(i);
let mut o = *op;
for j in 0..KB {
o = scal[j].mul_add(*col[j].add(i), o);
}
*op = o;
i += 1;
}
kk += KB;
}
while kk < k {
let scal = alpha * *vec.offset(kk as isize * vec_s);
let sv = simd.splat(scal);
let col = mat.offset(kk as isize * mat_cs);
let mut i = s;
while i + lanes <= e {
let mv = simd.loadu(col.add(i));
let ov = simd.loadu(out.add(i));
simd.storeu(out.add(i), simd.mul_add(mv, sv, ov));
i += lanes;
}
while i < e {
let op = out.add(i);
*op = scal.mul_add(*col.add(i), *op);
i += 1;
}
kk += 1;
}
}
}
const DOT_RB: usize = 4;
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn dot_rows<T, S>(
simd: S,
s: usize,
e: usize,
k: usize,
alpha: T,
mat: *const T,
mat_rs: isize,
vec: *const T,
beta: T,
out: *mut T,
out_s: isize,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
{
unsafe {
let lanes = <S as SimdOps<T>>::LANES;
let mut i = s;
while i + DOT_RB <= e {
let rows: [*const T; DOT_RB] =
core::array::from_fn(|r| mat.offset((i + r) as isize * mat_rs));
let mut acc = [simd.zero(); DOT_RB];
let mut kk = 0;
while kk + lanes <= k {
let v = simd.loadu(vec.add(kk));
for r in 0..DOT_RB {
acc[r] = simd.mul_add(simd.loadu(rows[r].add(kk)), v, acc[r]);
}
kk += lanes;
}
let mut dots: [T; DOT_RB] = core::array::from_fn(|r| simd.reduce_sum(acc[r]));
while kk < k {
let y = *vec.add(kk);
for r in 0..DOT_RB {
dots[r] = (*rows[r].add(kk)).mul_add(y, dots[r]);
}
kk += 1;
}
for (r, dot) in dots.into_iter().enumerate() {
let op = out.offset((i + r) as isize * out_s);
let ov = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*op
} else {
beta * *op
};
*op = alpha.mul_add(dot, ov);
}
i += DOT_RB;
}
while i < e {
let row = mat.offset(i as isize * mat_rs);
let dot = super::dot_contiguous::<T, S>(simd, k, row, vec);
let op = out.offset(i as isize * out_s);
let ov = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*op
} else {
beta * *op
};
*op = alpha.mul_add(dot, ov);
i += 1;
}
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn strided_rows<T, S>(
_simd: S,
s: usize,
e: usize,
k: usize,
alpha: T,
mat: *const T,
mat_rs: isize,
mat_cs: isize,
vec: *const T,
vec_s: isize,
beta: T,
out: *mut T,
out_s: isize,
) where
T: Float<Acc = T>,
S: SimdOps<T>,
{
unsafe {
for i in s..e {
let mut dot = T::ZERO;
for kk in 0..k {
dot = (*mat.offset(i as isize * mat_rs + kk as isize * mat_cs))
.mul_add(*vec.offset(kk as isize * vec_s), dot);
}
let op = out.offset(i as isize * out_s);
let ov = if beta == T::ZERO {
T::ZERO
} else if beta == T::ONE {
*op
} else {
beta * *op
};
*op = alpha.mul_add(dot, ov);
}
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn run_mixed<N, S>(
simd: S,
m: usize,
k: usize,
n: usize,
par: Parallelism,
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,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
{
unsafe {
if n == 1 {
core_mixed::<N, S>(simd, m, k, par, alpha, a, rsa, csa, b, rsb, beta, c, rsc);
} else {
core_mixed::<N, S>(simd, n, k, par, alpha, b, csb, rsb, a, csa, beta, c, csc);
}
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline]
unsafe fn core_mixed<N, S>(
simd: S,
rows: usize,
k: usize,
par: Parallelism,
alpha: f32,
mat: *const N,
mat_rs: isize,
mat_cs: isize,
vec: *const N,
vec_s: isize,
beta: f32,
out: *mut N,
out_s: isize,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let sizeof = core::mem::size_of::<N>();
let dot_legal = mat_cs == 1 && vec_s == 1;
let axpy = mat_rs == 1 && out_s == 1 && !axpy_yields_to_dot(rows, lanes, dot_legal);
let dot = !axpy && dot_legal;
let bytes_touched = (rows.saturating_mul(k) + k + rows).saturating_mul(sizeof);
let n_threads = par.resolve_bandwidth(bytes_touched, rows);
let block = if axpy { MB_REG * lanes } else { lanes }.max(1);
let mat = Ptr(mat as *mut N);
let vec = Ptr(vec as *mut N);
let out = Ptr(out);
let body = move |row_start: usize, row_end: usize| {
let (mat, vec, out) = (mat, vec, out);
let mat = mat.0 as *const N;
let vec = vec.0 as *const N;
let out = out.0;
simd.vectorize(|| {
if axpy {
axpy_mixed::<N, S>(
simd, row_start, row_end, k, alpha, mat, mat_cs, vec, vec_s, beta, out,
);
} else if dot {
dot_rows_mixed::<N, S>(
simd, row_start, row_end, k, alpha, mat, mat_rs, vec, beta, out, out_s,
);
} else {
strided_rows_mixed::<N, S>(
simd, row_start, row_end, k, alpha, mat, mat_rs, mat_cs, vec, vec_s, beta,
out, out_s,
);
}
});
};
row_sweep(rows, block, n_threads, body);
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn axpy_mixed<N, S>(
simd: S,
s: usize,
e: usize,
k: usize,
alpha: f32,
mat: *const N,
mat_cs: isize,
vec: *const N,
vec_s: isize,
beta: f32,
out: *mut N,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let mb = MB_REG * lanes;
let mut i = s;
while i + mb <= e {
let mut acc: [<S as SimdOps<f32>>::Reg; MB_REG] = [simd.zero(); MB_REG];
if beta == 1.0 {
for (r, a) in acc.iter_mut().enumerate() {
*a = simd.load_out(out.add(i + r * lanes));
}
} else if beta != 0.0 {
let bv = simd.splat(beta);
for (r, a) in acc.iter_mut().enumerate() {
*a = simd.mul(simd.load_out(out.add(i + r * lanes)), bv);
}
}
for kk in 0..k {
let sv = simd.splat(alpha * (*vec.offset(kk as isize * vec_s)).widen());
let col = mat.offset(kk as isize * mat_cs).add(i);
for (r, a) in acc.iter_mut().enumerate() {
*a = simd.mul_add(simd.load_lhs(col.add(r * lanes)), sv, *a);
}
}
for (r, a) in acc.iter().enumerate() {
simd.store_out(out.add(i + r * lanes), *a);
}
i += mb;
}
while i + lanes <= e {
let mut acc = if beta == 1.0 {
simd.load_out(out.add(i))
} else if beta == 0.0 {
simd.zero()
} else {
simd.mul(simd.load_out(out.add(i)), simd.splat(beta))
};
for kk in 0..k {
let sv = simd.splat(alpha * (*vec.offset(kk as isize * vec_s)).widen());
acc = simd.mul_add(
simd.load_lhs(mat.offset(kk as isize * mat_cs).add(i)),
sv,
acc,
);
}
simd.store_out(out.add(i), acc);
i += lanes;
}
while i < e {
let op = out.add(i);
let mut acc: f32 = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*op).widen()
} else {
beta * (*op).widen()
};
for kk in 0..k {
let sv = alpha * (*vec.offset(kk as isize * vec_s)).widen();
acc += sv * (*mat.offset(kk as isize * mat_cs).add(i)).widen();
}
*op = N::narrow(acc);
i += 1;
}
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn dot_rows_mixed<N, S>(
simd: S,
s: usize,
e: usize,
k: usize,
alpha: f32,
mat: *const N,
mat_rs: isize,
vec: *const N,
beta: f32,
out: *mut N,
out_s: isize,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let mut i = s;
while i + DOT_RB <= e {
let rows: [*const N; DOT_RB] =
core::array::from_fn(|r| mat.offset((i + r) as isize * mat_rs));
let mut acc: [<S as SimdOps<f32>>::Reg; DOT_RB] = [simd.zero(); DOT_RB];
let mut kk = 0;
while kk + lanes <= k {
let v = simd.load_lhs(vec.add(kk));
for r in 0..DOT_RB {
acc[r] = simd.mul_add(simd.load_lhs(rows[r].add(kk)), v, acc[r]);
}
kk += lanes;
}
let mut dots: [f32; DOT_RB] = core::array::from_fn(|r| simd.reduce_sum(acc[r]));
while kk < k {
let y = (*vec.add(kk)).widen();
for r in 0..DOT_RB {
dots[r] += (*rows[r].add(kk)).widen() * y;
}
kk += 1;
}
for (r, dot) in dots.into_iter().enumerate() {
let op = out.offset((i + r) as isize * out_s);
let ov = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*op).widen()
} else {
beta * (*op).widen()
};
*op = N::narrow(alpha * dot + ov);
}
i += DOT_RB;
}
while i < e {
let row = mat.offset(i as isize * mat_rs);
let dot = dot_contiguous_mixed::<N, S>(simd, k, row, vec);
let op = out.offset(i as isize * out_s);
let ov = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*op).widen()
} else {
beta * (*op).widen()
};
*op = N::narrow(alpha * dot + ov);
i += 1;
}
}
}
#[cfg(feature = "half")]
#[inline(always)]
unsafe fn dot_contiguous_mixed<N, S>(simd: S, k: usize, x: *const N, y: *const N) -> f32
where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
{
unsafe {
let lanes = <S as SimdOps<f32>>::LANES;
let mut acc = simd.zero();
let mut kk = 0;
while kk + lanes <= k {
acc = simd.mul_add(simd.load_lhs(x.add(kk)), simd.load_lhs(y.add(kk)), acc);
kk += lanes;
}
let mut dot = simd.reduce_sum(acc);
while kk < k {
dot += (*x.add(kk)).widen() * (*y.add(kk)).widen();
kk += 1;
}
dot
}
}
#[cfg(feature = "half")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn strided_rows_mixed<N, S>(
_simd: S,
s: usize,
e: usize,
k: usize,
alpha: f32,
mat: *const N,
mat_rs: isize,
mat_cs: isize,
vec: *const N,
vec_s: isize,
beta: f32,
out: *mut N,
out_s: isize,
) where
N: NarrowFloat,
S: KernelSimd<N, N, f32, N>,
{
unsafe {
for i in s..e {
let mut dot: f32 = 0.0;
for kk in 0..k {
dot += (*mat.offset(i as isize * mat_rs + kk as isize * mat_cs)).widen()
* (*vec.offset(kk as isize * vec_s)).widen();
}
let op = out.offset(i as isize * out_s);
let ov = if beta == 0.0 {
0.0
} else if beta == 1.0 {
(*op).widen()
} else {
beta * (*op).widen()
};
*op = N::narrow(alpha * dot + ov);
}
}
}
#[cfg(test)]
mod tests {
use super::{DOT_RB, MB_REG, axpy_row_split_loses, axpy_yields_to_dot};
use crate::simd::{ScalarTok, SimdOps};
#[test]
fn short_sweeps_yield_to_the_dot_form() {
for &lanes in &[1usize, 2, 4, 8, 16] {
for rows in 0..=(2 * lanes) {
assert!(
!axpy_yields_to_dot(rows, lanes, false),
"rows={rows} lanes={lanes}: nothing to yield to when the dot strides fail"
);
assert_eq!(
axpy_yields_to_dot(rows, lanes, true),
rows < lanes,
"rows={rows} lanes={lanes}"
);
}
}
assert!(!axpy_yields_to_dot(1, 1, true));
}
#[test]
fn axpy_row_split_floor_follows_its_knob() {
let prev = crate::tuning::gemv_axpy_par_min_rows();
crate::tuning::set_gemv_axpy_par_min_rows(0);
for rows in [0usize, 1, 16, 4096, usize::MAX] {
assert!(
!axpy_row_split_loses(rows),
"a 0 floor must never gate (rows={rows})"
);
}
for floor in [1usize, 16, 4096, 16384] {
crate::tuning::set_gemv_axpy_par_min_rows(floor);
for rows in [0usize, 1, 15, 16, 4095, 4096, 16383, 16384, usize::MAX] {
assert_eq!(
axpy_row_split_loses(rows),
rows < floor,
"floor={floor} rows={rows}"
);
}
}
crate::tuning::set_gemv_axpy_par_min_rows(prev);
}
macro_rules! axpy_regblock_check {
($fn:ident, $t:ty, $tol:expr) => {
fn $fn<S: SimdOps<$t>>(simd: S, label: &str) {
let lanes = <S as SimdOps<$t>>::LANES;
let rows = MB_REG * lanes + lanes + lanes.saturating_sub(1);
let k = 37usize;
let mat: Vec<$t> = (0..rows * k)
.map(|i| (((i as u64 * 1103515245 + 12345) % 251) as $t) * 0.008 - 1.0)
.collect();
let vec: Vec<$t> = (0..k)
.map(|i| (((i as u64 * 2654435761) % 193) as $t) * 0.01 - 0.9)
.collect();
let out0: Vec<$t> = (0..rows)
.map(|i| (((i as u64 * 40503) % 131) as $t) * 0.05 - 3.0)
.collect();
for &(alpha, beta) in &[
(1.3 as $t, 0.0 as $t),
(0.7 as $t, 1.0 as $t),
(1.1 as $t, 2.5 as $t),
] {
let mut out = out0.clone();
unsafe {
simd.vectorize(|| {
super::axpy_regblocked::<$t, S>(
simd,
0,
rows,
k,
alpha,
mat.as_ptr(),
rows as isize,
vec.as_ptr(),
1,
beta,
out.as_mut_ptr(),
);
});
}
for i in 0..rows {
let mut acc = if beta == 0.0 {
0.0 as $t
} else {
beta * out0[i]
};
for kk in 0..k {
acc += mat[kk * rows + i] * (alpha * vec[kk]);
}
let tol = $tol * (1.0 as $t + acc.abs());
assert!(
(out[i] - acc).abs() <= tol,
"{label} lanes={lanes} beta={beta} row {i}: got {} want {}",
out[i],
acc
);
}
}
}
};
}
axpy_regblock_check!(check_f32, f32, 1e-4);
axpy_regblock_check!(check_f64, f64, 1e-10);
#[test]
fn axpy_regblocked_spans_all_regimes() {
check_f32(ScalarTok, "scalar/f32");
check_f64(ScalarTok, "scalar/f64");
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
use crate::simd::{Avx512F, Fma};
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
check_f32(Fma, "fma/f32");
check_f64(Fma, "fma/f64");
}
if is_x86_feature_detected!("avx512f") {
check_f32(Avx512F, "avx512f/f32");
check_f64(Avx512F, "avx512f/f64");
}
}
#[cfg(target_arch = "aarch64")]
{
check_f32(crate::simd::Neon, "neon/f32");
check_f64(crate::simd::Neon, "neon/f64");
}
}
macro_rules! dot_rows_bit_identity_check {
($fn:ident, $t:ty) => {
fn $fn<S: SimdOps<$t>>(simd: S, label: &str) {
let lanes = <S as SimdOps<$t>>::LANES;
let rows = DOT_RB * 2 + 3;
let k = lanes * 5 + 3;
let mat: Vec<$t> = (0..rows * k)
.map(|i| (((i as u64 * 1103515245 + 12345) % 251) as $t) * 0.008 - 1.0)
.collect();
let vec: Vec<$t> = (0..k)
.map(|i| (((i as u64 * 2654435761) % 193) as $t) * 0.01 - 0.9)
.collect();
let out0: Vec<$t> = (0..rows)
.map(|i| (((i as u64 * 40503) % 131) as $t) * 0.05 - 3.0)
.collect();
for &(alpha, beta) in &[
(1.3 as $t, 0.0 as $t),
(0.7 as $t, 1.0 as $t),
(1.1 as $t, 2.5 as $t),
] {
let mut out = out0.clone();
let mut refr = out0.clone();
unsafe {
simd.vectorize(|| {
super::dot_rows::<$t, S>(
simd,
0,
rows,
k,
alpha,
mat.as_ptr(),
k as isize,
vec.as_ptr(),
beta,
out.as_mut_ptr(),
1,
);
for i in 0..rows {
let row = mat.as_ptr().add(i * k);
let dot = crate::special::dot_contiguous::<$t, S>(
simd,
k,
row,
vec.as_ptr(),
);
let ov = if beta == 0.0 as $t {
0.0 as $t
} else if beta == 1.0 as $t {
refr[i]
} else {
beta * refr[i]
};
refr[i] = alpha * dot + ov;
}
});
}
for i in 0..rows {
assert_eq!(
out[i].to_bits(),
refr[i].to_bits(),
"{label} lanes={lanes} beta={beta} row {i}: blocked {} vs ref {}",
out[i],
refr[i]
);
}
}
}
};
}
dot_rows_bit_identity_check!(dot_check_f32, f32);
dot_rows_bit_identity_check!(dot_check_f64, f64);
#[test]
fn dot_rows_bit_identical() {
dot_check_f32(ScalarTok, "scalar/f32");
dot_check_f64(ScalarTok, "scalar/f64");
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
use crate::simd::{Avx512F, Fma};
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
dot_check_f32(Fma, "fma/f32");
dot_check_f64(Fma, "fma/f64");
}
if is_x86_feature_detected!("avx512f") {
dot_check_f32(Avx512F, "avx512f/f32");
dot_check_f64(Avx512F, "avx512f/f64");
}
}
#[cfg(target_arch = "aarch64")]
{
dot_check_f32(crate::simd::Neon, "neon/f32");
dot_check_f64(crate::simd::Neon, "neon/f64");
}
}
}