use crate::DatumType;
use crate::block_quant::*;
use crate::isa::Arch;
use crate::isa::Isa;
use crate::isa::IsaSet;
use crate::mmm::{Query, Suitable};
use crate::mmm_tiers::MmmTier;
use crate::pack::PackedFormat;
#[cfg(any(tract_avx512vnni, tract_avxvnni, tract_amx_int8))]
use crate::pack::PackedI8K4;
#[cfg(tract_amx_int8)]
use super::amx::PackedAmxA;
#[cfg(tract_amx_bf16)]
use super::amx_bf16::{PackedAmxBf16A, PackedBf16K2};
#[cfg(tract_avx512vnni)]
use super::fma_width::has_dual_avx512_fma;
#[cfg(tract_avx512vnni)]
const AVX512VNNI_WIDE_TILE: fn() -> isize = || if has_dual_avx512_fma() { 50 } else { -1 };
#[cfg(tract_amx_bf16)]
const AMX_BF16_OPT_IN: fn() -> isize =
|| if crate::knobs::TRACT_AMX_BF16.get() { 100 } else { crate::isa::NEVER_PREFERRED };
use super::*;
#[derive(Clone, Copy)]
struct KernelChoice {
mr: usize,
nr: usize,
scale: f32,
name: fn() -> &'static str,
}
fn tile_util(d: usize, tile: usize) -> f32 {
if d == 0 {
return 1.0;
}
let batches = d.div_ceil(tile);
d as f32 / (batches * tile) as f32
}
fn pick_mmm(choices: &[KernelChoice], m: Option<usize>, n: Option<usize>) -> &'static str {
let key = |c: &KernelChoice| -> (f32, i32, i32) {
let m_u = m.map(|m| tile_util(m, c.mr)).unwrap_or(1.0);
let n_u = n.map(|n| tile_util(n, c.nr)).unwrap_or(1.0);
let m_b = m.map(|m| m.div_ceil(c.mr)).unwrap_or(1) as i32;
let n_b = n.map(|n| n.div_ceil(c.nr)).unwrap_or(1) as i32;
(c.scale * m_u * n_u, -(m_b * n_b), c.nr as i32)
};
let best = choices
.iter()
.max_by(|a, b| key(a).partial_cmp(&key(b)).unwrap())
.expect("non-empty kernel pool");
(best.name)()
}
MMMExternKernel!(x86_64; avx_mmm_f32_8x8 <f32>(8, 8)@(256,4) isa(X86_64Avx));
MMMExternKernel!(x86_64; avx_mmm_f32_16x5<f32>(16,5)@(256,4) isa(X86_64Avx));
MMMExternKernel!(x86_64; avx_mmm_f32_16x6<f32>(16,6)@(256,4) isa(X86_64Avx));
MMMExternKernel!(x86_64; avx_mmm_f32_24x4<f32>(24,4)@(256,4) isa(X86_64Avx));
MMMExternKernel!(x86_64; avx_mmm_f32_32x3<f32>(32,3)@(256,4) isa(X86_64Avx));
MMMExternKernel!(x86_64; avx_mmm_f32_40x2<f32>(40,2)@(256,4) isa(X86_64Avx));
MMMExternKernel!(x86_64; avx_mmm_f32_64x1<f32>(64,1)@(256,4) isa(X86_64Avx));
const FMA_F32_PEER: fn() -> isize = || crate::isa::peer_of(Isa::X86_64Fma, Isa::X86_64Avx512f);
MMMExternKernel!(x86_64; fma_mmm_f32_8x8 <f32>(8, 8)@(256,4) isa(X86_64Avx, X86_64Fma) boost(FMA_F32_PEER));
MMMExternKernel!(x86_64; fma_mmm_f32_16x6<f32>(16,6)@(256,4) isa(X86_64Avx, X86_64Fma) boost(FMA_F32_PEER));
MMMExternKernel!(x86_64; fma_mmm_f32_16x5<f32>(16,5)@(256,4) isa(X86_64Avx, X86_64Fma) boost(FMA_F32_PEER));
MMMExternKernel!(x86_64; fma_mmm_f32_24x4<f32>(24,4)@(256,4) isa(X86_64Avx, X86_64Fma) boost(FMA_F32_PEER));
MMMExternKernel!(x86_64; fma_mmm_f32_40x2<f32>(40,2)@(256,4) isa(X86_64Avx, X86_64Fma) boost(FMA_F32_PEER));
MMMExternKernel!(x86_64; fma_mmm_f32_64x1<f32>(64,1)@(256,4) isa(X86_64Avx, X86_64Fma) boost(FMA_F32_PEER));
pub fn pq40_r32() -> PackedBlockQuantFormat {
PackedBlockQuantFormat::new(&Q4_0, 32, 16, false)
}
pub fn pq20t_r32() -> PackedBlockQuantFormat {
PackedBlockQuantFormat::new(&Q2_0_T, 32, 0, false)
}
MMMExternKernel! { x86_64; fma_mmm_f32_32x1<f32>(32,1)@(256,4) isa(X86_64Avx, X86_64Fma, X86_64F16c)
packing[1] = q40f32 => |k| k.with_packing_a(pq40_r32());
packing[2] = q40f16 => |k| k.with_packing(pq40_r32(), f16::packing(1));
packing[3] = f16f16 => |k| k.with_packing(f16::packing(32), f16::packing(1));
packing[4] = f16f32 => |k| k.with_packing(f16::packing(32), f32::packing(1));
packing[5] = f32f16 => |k| k.with_packing(f32::packing(32), f16::packing(1));
boost(FMA_F32_PEER)
store(f16)
}
MMMExternKernel!(x86_64; fma_mmm_f32_32x3<f32>(32,3)@(256,4) isa(X86_64Avx, X86_64Fma)
packing[1] = f32f16 => |k| k.with_packing(f32::packing(32).align(256), f16::packing(3));
packing[2] = f16f32 => |k| k.with_packing(f16::packing(32).align(256), f32::packing(3));
packing[3] = f16f16 => |k| k.with_packing(f16::packing(32).align(256), f16::packing(3));
boost(FMA_F32_PEER)
store(f16)
);
MMMExternKernel!(x86_64; avx512_mmm_f32_128x1<f32>(128, 1)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_16x1 <f32>( 16, 1)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_16x12<f32>( 16,12)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_16x8 <f32>( 16, 8)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_32x6 <f32>( 32, 6)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_32x5 <f32>( 32, 5)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_48x4 <f32>( 48, 4)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_64x3 <f32>( 64, 3)@(512,4) isa(X86_64Avx512f));
MMMExternKernel!(x86_64; avx512_mmm_f32_80x2 <f32>( 80, 2)@(512,4) isa(X86_64Avx512f));
MMMExternKernel! { x86_64; avx_mmm_i32_8x4<i32>(8,4)@(256,4) isa(X86_64Avx)
packing[1] = i8i8 => |k| k.with_packing(PackedFormat::new(DatumType::I8, 8, 256), PackedFormat::new(DatumType::I8, 4, 4));
store(i8)
}
MMMExternKernel! { x86_64; avx2_mmm_i32_8x8<i32>(8,8)@(256,4) isa(X86_64Avx2)
packing[1] = i8i8 => |k| k.with_packing(PackedFormat::new(DatumType::I8, 8, 256), PackedFormat::new(DatumType::I8, 8, 4));
store(i8)
}
#[cfg(tract_avx512vnni)]
MMMExternKernel! { x86_64; avx512vnni_mmm_i32_8x8<i32>(8,8)@(256,4) isa(X86_64Avx512f, X86_64Avx512Vnni)
packing[1] = i8i8 => |k| k.with_packing(PackedI8K4::new(8), PackedI8K4::new(8));
store(i8)
}
#[cfg(tract_avx512vnni)]
MMMExternKernel! { x86_64; avx512vnni_mmm_i32_16x16<i32>(16,16)@(64,4) isa(X86_64Avx512f, X86_64Avx512Vnni)
packing[1] = i8i8 => |k| k.with_packing(PackedI8K4::new(16), PackedI8K4::new(16));
boost(AVX512VNNI_WIDE_TILE)
store(i8)
}
#[cfg(tract_avxvnni)]
MMMExternKernel! { x86_64; avxvnni_mmm_i32_8x8<i32>(8,8)@(256,4) isa(X86_64Avx2, X86_64AvxVnni)
packing[1] = i8i8 => |k| k.with_packing(PackedI8K4::new(8), PackedI8K4::new(8));
store(i8)
}
#[cfg(tract_amx_int8)]
MMMExternKernel! { x86_64; avx512amx_mmm_i32_8x8<i32>(8,8)@(64,4) isa(X86_64Avx512f, X86_64AmxInt8)
packing[1] = i8i8 => |k| k.with_packing(PackedAmxA::new(8), PackedI8K4::new(8));
store(i8)
}
#[cfg(tract_amx_int8)]
MMMExternKernel! { x86_64; avx512amx_mmm_i32_16x16<i32>(16,16)@(64,4) isa(X86_64Avx512f, X86_64AmxInt8)
packing[1] = i8i8 => |k| k.with_packing(PackedAmxA::new(16), PackedI8K4::new(16));
boost(|| 100)
store(i8)
}
#[cfg(tract_amx_bf16)]
MMMExternKernel! { x86_64; avx512amx_mmm_f32_16x16<f32>(16,16)@(64,4) isa(X86_64Avx512f, X86_64AmxBf16)
packing[1] = f32f32_bf16 => |k| k.with_packing(PackedAmxBf16A::new(16), PackedBf16K2::new(16));
boost(AMX_BF16_OPT_IN)
}
#[cfg(tract_avx512vnni)]
fn avx512vnni_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
_suitable: &[Suitable],
) -> Option<&'static str> {
match (dt, query.n) {
(DatumType::I32, Some(1)) => None,
(DatumType::I32, _) if !has_dual_avx512_fma() => Some(avx512vnni_mmm_i32_8x8.name.as_str()),
(DatumType::I32, _) => {
let big = |o: Option<usize>, t: usize| o.is_none_or(|v| v >= t);
Some(if big(query.m, 16) && big(query.n, 16) {
avx512vnni_mmm_i32_16x16.name.as_str()
} else {
avx512vnni_mmm_i32_8x8.name.as_str()
})
}
_ => None,
}
}
#[cfg(tract_avx512vnni)]
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 6,
name: "avx512vnni",
applies: |isa| {
isa.has(Isa::X86_64Avx2) && isa.has(Isa::X86_64Fma) && isa.has(Isa::X86_64Avx512f)
&& isa.has(Isa::X86_64Avx512Vnni)
},
preferred: avx512vnni_preferred,
}
}
#[cfg(tract_avxvnni)]
fn avxvnni_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
_suitable: &[Suitable],
) -> Option<&'static str> {
match (dt, query.n) {
(DatumType::I32, Some(1)) => None,
(DatumType::I32, _) => Some(avxvnni_mmm_i32_8x8.name.as_str()),
_ => None,
}
}
#[cfg(tract_avxvnni)]
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 3,
name: "avxvnni",
applies: |isa| isa.has(Isa::X86_64Avx2) && isa.has(Isa::X86_64AvxVnni),
preferred: avxvnni_preferred,
}
}
#[cfg(tract_amx_bf16)]
fn amx_bf16_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
_suitable: &[Suitable],
) -> Option<&'static str> {
if dt != DatumType::F32 || query.n == Some(1) {
return None;
}
let big = |o: Option<usize>, t: usize| o.is_none_or(|v| v >= t);
if big(query.m, 16) && big(query.n, 16) && big(query.k, 32) {
Some(avx512amx_mmm_f32_16x16.name.as_str())
} else {
None
}
}
#[cfg(tract_amx_bf16)]
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 8,
name: "avx512amx-bf16",
applies: |isa| {
crate::knobs::TRACT_AMX_BF16.get()
&& isa.has(Isa::X86_64Avx2)
&& isa.has(Isa::X86_64Fma)
&& isa.has(Isa::X86_64Avx512f)
&& isa.has(Isa::X86_64AmxBf16)
},
preferred: amx_bf16_preferred,
}
}
#[cfg(tract_amx_int8)]
fn amx_int8_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
_suitable: &[Suitable],
) -> Option<&'static str> {
match (dt, query.n) {
(DatumType::I32, Some(1)) => None,
(DatumType::I32, _) => {
let big = |o: Option<usize>, t: usize| o.is_none_or(|v| v >= t);
Some(if big(query.m, 16) && big(query.n, 16) && big(query.k, 64) {
avx512amx_mmm_i32_16x16.name.as_str()
} else {
avx512amx_mmm_i32_8x8.name.as_str()
})
}
_ => None,
}
}
#[cfg(tract_amx_int8)]
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 7,
name: "avx512amx-int8",
applies: |isa| {
isa.has(Isa::X86_64Avx2)
&& isa.has(Isa::X86_64Fma)
&& isa.has(Isa::X86_64Avx512f)
&& isa.has(Isa::X86_64Avx512Vnni)
&& isa.has(Isa::X86_64AmxInt8)
},
preferred: amx_int8_preferred,
}
}
fn avx2_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
_suitable: &[Suitable],
) -> Option<&'static str> {
match (dt, query.n) {
(DatumType::I32, Some(1)) => None,
(DatumType::I32, _) => Some(&mmm::avx2_mmm_i32_8x8.name.as_str()),
_ => None,
}
}
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 2,
name: "avx2",
applies: |isa| isa.has(Isa::X86_64Avx2),
preferred: avx2_preferred,
}
}
const AVX_CHOICES: &[KernelChoice] = &[
KernelChoice { mr: 16, nr: 6, scale: 1.0, name: || avx_mmm_f32_16x6.name.as_str() },
KernelChoice { mr: 16, nr: 5, scale: 0.98, name: || avx_mmm_f32_16x5.name.as_str() },
KernelChoice { mr: 24, nr: 4, scale: 0.95, name: || avx_mmm_f32_24x4.name.as_str() },
KernelChoice { mr: 32, nr: 3, scale: 0.93, name: || avx_mmm_f32_32x3.name.as_str() },
KernelChoice { mr: 40, nr: 2, scale: 0.90, name: || avx_mmm_f32_40x2.name.as_str() },
KernelChoice { mr: 8, nr: 8, scale: 0.80, name: || avx_mmm_f32_8x8.name.as_str() },
];
fn avx_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
_suitable: &[Suitable],
) -> Option<&'static str> {
match (dt, query.n) {
(DatumType::I32, Some(1)) => None,
(DatumType::I32, _) => Some(avx_mmm_i32_8x4.name.as_str()),
(DatumType::F32, _) => match query.n {
None => Some(avx_mmm_f32_16x6.name.as_str()),
Some(1) => Some(avx_mmm_f32_64x1.name.as_str()),
Some(2) => Some(avx_mmm_f32_40x2.name.as_str()),
Some(3) => Some(avx_mmm_f32_32x3.name.as_str()),
Some(4) => Some(avx_mmm_f32_24x4.name.as_str()),
Some(5) => Some(avx_mmm_f32_16x5.name.as_str()),
Some(6) => Some(avx_mmm_f32_16x6.name.as_str()),
Some(8) => Some(avx_mmm_f32_8x8.name.as_str()),
Some(_) => Some(pick_mmm(AVX_CHOICES, query.m, query.n)),
},
_ => None,
}
}
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 1,
name: "avx",
applies: |isa| isa.has(Isa::X86_64Avx),
preferred: avx_preferred,
}
}
const FMA_CHOICES: &[KernelChoice] = &[
KernelChoice { mr: 8, nr: 8, scale: 44.0 / 60.0, name: || fma_mmm_f32_8x8.name.as_str() },
KernelChoice { mr: 16, nr: 6, scale: 54.0 / 60.0, name: || fma_mmm_f32_16x6.name.as_str() },
KernelChoice { mr: 16, nr: 5, scale: 54.0 / 60.0, name: || fma_mmm_f32_16x5.name.as_str() },
KernelChoice { mr: 24, nr: 4, scale: 54.0 / 60.0, name: || fma_mmm_f32_24x4.name.as_str() },
KernelChoice { mr: 32, nr: 3, scale: 54.0 / 60.0, name: || fma_mmm_f32_32x3.name.as_str() },
KernelChoice { mr: 40, nr: 2, scale: 54.0 / 60.0, name: || fma_mmm_f32_40x2.name.as_str() },
];
fn fma_mmm_f32(suitable: &[Suitable], query: &Query) -> Option<&'static str> {
match super::vendor() {
super::Vendor::Intel => {
super::intel_fma_linear::linear_model().preferred(suitable, query.m, query.k, query.n)
}
super::Vendor::Amd => {
super::amd_fma_linear::linear_model().preferred(suitable, query.m, query.k, query.n)
}
super::Vendor::Other => match query.n {
None => Some(fma_mmm_f32_16x6.name.as_str()),
Some(1) => unreachable!("n == 1 answered above"),
Some(2) => Some(fma_mmm_f32_40x2.name.as_str()),
Some(3) => Some(fma_mmm_f32_32x3.name.as_str()),
Some(4) => Some(fma_mmm_f32_24x4.name.as_str()),
Some(5) => Some(fma_mmm_f32_16x5.name.as_str()),
Some(6) => Some(fma_mmm_f32_16x6.name.as_str()),
Some(8) => Some(fma_mmm_f32_8x8.name.as_str()),
Some(_) => Some(pick_mmm(FMA_CHOICES, query.m, query.n)),
},
}
}
fn fma_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
suitable: &[Suitable],
) -> Option<&'static str> {
match (dt, query.n) {
(DatumType::F32, Some(1)) => Some(fma_mmm_f32_64x1.name.as_str()),
(DatumType::F32, _) => fma_mmm_f32(suitable, query),
_ => None,
}
}
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 4,
name: "fma",
applies: |isa| isa.has(Isa::X86_64Avx2) && isa.has(Isa::X86_64Fma),
preferred: fma_preferred,
}
}
const X86_F32_CHOICES: &[KernelChoice] = &[
KernelChoice { mr: 16, nr: 12, scale: 1.000, name: || avx512_mmm_f32_16x12.name.as_str() },
KernelChoice { mr: 16, nr: 8, scale: 0.995, name: || avx512_mmm_f32_16x8.name.as_str() },
KernelChoice { mr: 32, nr: 5, scale: 0.992, name: || avx512_mmm_f32_32x5.name.as_str() },
KernelChoice { mr: 32, nr: 6, scale: 0.990, name: || avx512_mmm_f32_32x6.name.as_str() },
KernelChoice { mr: 48, nr: 4, scale: 0.978, name: || avx512_mmm_f32_48x4.name.as_str() },
KernelChoice { mr: 16, nr: 6, scale: 0.964, name: || fma_mmm_f32_16x6.name.as_str() },
KernelChoice { mr: 24, nr: 4, scale: 0.948, name: || fma_mmm_f32_24x4.name.as_str() },
KernelChoice { mr: 16, nr: 5, scale: 0.935, name: || fma_mmm_f32_16x5.name.as_str() },
KernelChoice { mr: 32, nr: 3, scale: 0.919, name: || fma_mmm_f32_32x3.name.as_str() },
KernelChoice { mr: 64, nr: 3, scale: 0.895, name: || avx512_mmm_f32_64x3.name.as_str() },
KernelChoice { mr: 40, nr: 2, scale: 0.842, name: || fma_mmm_f32_40x2.name.as_str() },
KernelChoice { mr: 8, nr: 8, scale: 0.788, name: || fma_mmm_f32_8x8.name.as_str() },
KernelChoice { mr: 80, nr: 2, scale: 0.766, name: || avx512_mmm_f32_80x2.name.as_str() },
KernelChoice { mr: 128, nr: 1, scale: 0.378, name: || avx512_mmm_f32_128x1.name.as_str() },
];
fn avx512_mmv_f32(suitable: &[Suitable], query: &Query) -> Option<&'static str> {
match super::vendor() {
super::Vendor::Intel => match query.m {
Some(m) if m < 128 => super::intel_avx512_mmv_linear::linear_model().preferred(
suitable,
Some(m),
query.k,
Some(1),
),
_ => Some(avx512_mmm_f32_128x1.name.as_str()),
},
_ => match query.m {
Some(m) if m < 31 => Some(avx512_mmm_f32_16x1.name.as_str()),
_ => Some(avx512_mmm_f32_128x1.name.as_str()),
},
}
}
fn avx512_mmm_f32(suitable: &[Suitable], query: &Query) -> Option<&'static str> {
match super::vendor() {
super::Vendor::Intel => super::intel_avx512_linear::linear_model()
.preferred(suitable, query.m, query.k, query.n),
super::Vendor::Amd => {
super::amd_avx512_linear::linear_model().preferred(suitable, query.m, query.k, query.n)
}
super::Vendor::Other => {
if let Some(1) = query.n {
unreachable!("n == 1 answered above");
}
Some(pick_mmm(X86_F32_CHOICES, query.m, query.n))
}
}
}
fn avx512f_preferred(
_isa: &IsaSet,
dt: DatumType,
query: &Query,
suitable: &[Suitable],
) -> Option<&'static str> {
match (dt, query.n) {
(DatumType::F32, Some(1)) => avx512_mmv_f32(suitable, query),
(DatumType::F32, _) => avx512_mmm_f32(suitable, query),
_ => None,
}
}
inventory::submit! {
MmmTier {
arch: Some(Arch::X86_64),
precedence: 5,
name: "avx512f",
applies: |isa| isa.has(Isa::X86_64Avx2) && isa.has(Isa::X86_64Fma) && isa.has(Isa::X86_64Avx512f),
preferred: avx512f_preferred,
}
}
#[cfg(all(test, target_arch = "x86_64"))]
mod tests {
use super::*;
use crate::frame::mmm::{AsInputValue, FusedSpec};
use tract_data::internal::*;
#[test]
fn avx512_128x1_add_unicast_with_strided_c() -> TractResult<()> {
if !is_x86_feature_detected!("avx512f") {
return Ok(());
}
let (m, k_each, n) = (1000usize, 256usize, 13usize);
let a0: Vec<f32> = (0..m * k_each).map(|i| ((i % 17) as f32 - 8.0) / 16.0).collect();
let a1: Vec<f32> = (0..m * k_each).map(|i| ((i % 19) as f32 - 9.0) / 18.0).collect();
let b0: Vec<f32> = (0..k_each * n).map(|i| ((i % 13) as f32 - 6.0) / 13.0).collect();
let b1: Vec<f32> = (0..k_each * n).map(|i| ((i % 11) as f32 - 5.0) / 10.0).collect();
let mut expected = vec![0.0f32; m * n];
for r in 0..m {
for c in 0..n {
let mut acc = 0.0f32;
for kk in 0..k_each {
acc += a0[r * k_each + kk] * b0[kk * n + c];
acc += a1[r * k_each + kk] * b1[kk * n + c];
}
expected[r * n + c] = acc;
}
}
let ker = avx512_mmm_f32_128x1.mmm();
let (pack_a, pack_b) = &ker.packings()[0];
let pack_one =
|buf: Vec<f32>, rows, cols, m_axis, k_axis, pack: &dyn crate::mmm::MMMInputFormat| {
let t =
tract_ndarray::Array2::from_shape_vec((rows, cols), buf).unwrap().into_tensor();
pack.prepare_one(&t, k_axis, m_axis).unwrap()
};
let pa0 = pack_one(a0, m, k_each, 0, 1, &**pack_a);
let pa1 = pack_one(a1, m, k_each, 0, 1, &**pack_a);
let pb0 = pack_one(b0, k_each, n, 1, 0, &**pack_b);
let pb1 = pack_one(b1, k_each, n, 1, 0, &**pack_b);
let spatial = 13usize;
let mut c_backing = Tensor::zero::<f32>(&[m, spatial, n])?;
let c_spec = unsafe { ker.c_from_data_and_strides(4, (spatial * n) as isize, 1) };
unsafe {
let c_view = c_backing.view_mut();
let c = c_spec.wrap(&c_view);
let ops: TVec<FusedSpec> = tvec!(
FusedSpec::AddMatMul {
a: AsInputValue::Borrowed(&*pa0),
b: AsInputValue::Borrowed(&*pb0),
packing: 0,
},
FusedSpec::Store(c),
);
ker.run(m, n, &ops)?;
}
unsafe {
let c_view = c_backing.view_mut();
let c_for_unicast = c_spec.wrap(&c_view);
let c_for_store = c_spec.wrap(&c_view);
let ops: TVec<FusedSpec> = tvec!(
FusedSpec::AddMatMul {
a: AsInputValue::Borrowed(&*pa1),
b: AsInputValue::Borrowed(&*pb1),
packing: 0,
},
FusedSpec::AddUnicast(c_for_unicast),
FusedSpec::Store(c_for_store),
);
ker.run(m, n, &ops)?;
}
let c_slice = c_backing.to_plain_array_view::<f32>()?;
let mut max_err = 0.0f32;
let mut wrong_cells = 0;
for r in 0..m {
for cc in 0..n {
let got = c_slice[[r, 0, cc]];
let exp = expected[r * n + cc];
let e = (got - exp).abs();
if e > 1e-3 {
wrong_cells += 1;
}
max_err = max_err.max(e);
}
}
assert!(
max_err < 1e-3,
"avx512_mmm_f32_128x1 wrong output at squeezenet shape: \
max_err={max_err}, {wrong_cells}/{} cells off",
m * n,
);
Ok(())
}
}