Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 26 additions & 0 deletions linalg/src/x86_64_fma.rs
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ pub mod panel_extract;
pub mod rms_norm;
pub mod softmax;

const AVX: fn() -> bool = || is_x86_feature_detected!("avx");
const AVX2: fn() -> bool = || is_x86_feature_detected!("avx2");
const FMA: fn() -> bool = || is_x86_feature_detected!("fma");
const AVX512F: fn() -> bool = || is_x86_feature_detected!("avx512f");
Expand All @@ -70,6 +71,11 @@ const AVX512VNNI: fn() -> bool = || is_x86_feature_detected!("avx512vnni");
tanh_impl!(f32, fma_tanh_f32, 8, 8, is_x86_feature_detected!("fma"));
sigmoid_impl!(f32, fma_sigmoid_f32, 8, 8, is_x86_feature_detected!("fma"));

// AVX-without-FMA ports of the fma kernels above (each vfmadd132ps expanded
// to an in-place vmulps+vaddps pair) for CPUs outside the fma tier.
tanh_impl!(f32, avx_tanh_f32, 8, 8, is_x86_feature_detected!("avx"));
sigmoid_impl!(f32, avx_sigmoid_f32, 8, 8, is_x86_feature_detected!("avx"));

// AVX-512 (zmm, 16-wide) variants. The assembly lives in x86_64/avx512/; the
// main loop handles 64 lanes (4 zmm) per iteration with a 16-lane tail, so
// nr()=16 (any multiple of 16 is safe).
Expand All @@ -78,6 +84,21 @@ sigmoid_impl!(f32, avx512_sigmoid_f32, 16, 16, is_x86_feature_detected!("avx512f

fn plug_avx2(_ops: &mut Ops) {}

/// Element-wise kernels for AVX-capable CPUs outside the fma tier: the
/// mul_by_scalar / max / min asm is plain AVX, and sigmoid / tanh have
/// dedicated mul+add ports. softmax uses fma asm and keeps its generic
/// fallback on this tier.
fn plug_avx(ops: &mut Ops) {
ops.sigmoid_f32 = Box::new(|| avx_sigmoid_f32::ew());
ops.tanh_f32 = Box::new(|| avx_tanh_f32::ew());

ops.mul_by_scalar_f32 = Box::new(|| by_scalar::x86_64_avx_f32_mul_by_scalar_32n::ew());
ops.max_f32 = Box::new(|| max::x86_64_fma_max_f32_32n::red());
ops.min_f32 = Box::new(|| min::x86_64_fma_min_f32_32n::red());

log::info!("sigmoid_f32, tanh_f32, mul_by_scalar_f32, max_f32, min_f32: x86_64/avx activated");
}

fn plug_fma(ops: &mut Ops) {
panel_extract::plug(ops);

Expand Down Expand Up @@ -146,6 +167,11 @@ fn plug_avx512f(ops: &mut Ops) {

pub fn plug(ops: &mut Ops) {
mmm::plug(ops);
if is_x86_feature_detected!("avx")
&& !(is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma"))
{
plug_avx(ops);
}
if is_x86_feature_detected!("avx2") {
plug_avx2(ops);
if is_x86_feature_detected!("fma") {
Expand Down
6 changes: 4 additions & 2 deletions linalg/src/x86_64_fma/by_scalar.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,9 @@ unsafe fn x86_64_avx_f32_mul_by_scalar_32n_run(buf: &mut [f32], scalar: f32) {
let len = buf.len();
let ptr = buf.as_ptr();
std::arch::asm!("
vbroadcastss ymm0, xmm0
// reg-source vbroadcastss needs avx2; this kernel must stay avx-safe
vpermilps xmm0, xmm0, 0
vinsertf128 ymm0, ymm0, xmm0, 1
2:
vmovaps ymm4, [{ptr}]
vmovaps ymm5, [{ptr} + 32]
Expand Down Expand Up @@ -48,7 +50,7 @@ unsafe fn x86_64_avx_f32_mul_by_scalar_32n_run(buf: &mut [f32], scalar: f32) {
pub mod test_x86_64_avx_f32_mul_by_scalar_32n {
use super::*;
by_scalar_frame_tests!(
is_x86_feature_detected!("avx2"),
is_x86_feature_detected!("avx"),
f32,
x86_64_avx_f32_mul_by_scalar_32n,
|a, b| a * b
Expand Down
6 changes: 4 additions & 2 deletions linalg/src/x86_64_fma/max.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,9 @@ unsafe fn x86_64_fma_max_f32_32n_run(buf: &[f32]) -> f32 {
let ptr = buf.as_ptr();
let mut acc = f32::MIN;
std::arch::asm!("
vbroadcastss ymm0, xmm0
// reg-source vbroadcastss needs avx2; this kernel must stay avx-safe
vpermilps xmm0, xmm0, 0
vinsertf128 ymm0, ymm0, xmm0, 1
vmovaps ymm1, ymm0
vmovaps ymm2, ymm0
vmovaps ymm3, ymm0
Expand Down Expand Up @@ -63,7 +65,7 @@ unsafe fn x86_64_fma_max_f32_32n_run(buf: &[f32]) -> f32 {
#[cfg(test)]
mod test_x86_64_fma_max_f32_32n {
use super::*;
crate::max_frame_tests!(is_x86_feature_detected!("avx2"), f32, x86_64_fma_max_f32_32n);
crate::max_frame_tests!(is_x86_feature_detected!("avx"), f32, x86_64_fma_max_f32_32n);
}

// AVX-512 version: processes 64 f32 per loop iteration (4 zmm registers of 16
Expand Down
6 changes: 4 additions & 2 deletions linalg/src/x86_64_fma/min.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,9 @@ unsafe fn x86_64_fma_min_f32_32n_run(buf: &[f32]) -> f32 {
let ptr = buf.as_ptr();
let mut acc = f32::MAX;
std::arch::asm!("
vbroadcastss ymm0, xmm0
// reg-source vbroadcastss needs avx2; this kernel must stay avx-safe
vpermilps xmm0, xmm0, 0
vinsertf128 ymm0, ymm0, xmm0, 1
vmovaps ymm1, ymm0
vmovaps ymm2, ymm0
vmovaps ymm3, ymm0
Expand Down Expand Up @@ -63,5 +65,5 @@ unsafe fn x86_64_fma_min_f32_32n_run(buf: &[f32]) -> f32 {
#[cfg(test)]
mod test_x86_64_fma_min_f32_32n {
use super::*;
crate::min_frame_tests!(is_x86_feature_detected!("avx2"), f32, x86_64_fma_min_f32_32n);
crate::min_frame_tests!(is_x86_feature_detected!("avx"), f32, x86_64_fma_min_f32_32n);
}
59 changes: 59 additions & 0 deletions linalg/src/x86_64_fma/mmm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,17 @@ fn pick_mmm(candidates: &[KernelChoice], m: Option<usize>, n: Option<usize>) ->
(best.ctor)()
}

// AVX-without-FMA f32 tier for pre-Haswell CPUs (Sandy Bridge / Ivy Bridge):
// same tile geometries as their fma_ siblings but the inner loops use
// vmulps+vaddps, and add_unicast avoids the avx2-only vgatherdps.
MMMExternKernel!(avx_mmm_f32_8x8 <f32>(8, 8)@(256,4) where(AVX) quality(ManuallyOptimized));
MMMExternKernel!(avx_mmm_f32_16x5<f32>(16,5)@(256,4) where(AVX) quality(ManuallyOptimized));
MMMExternKernel!(avx_mmm_f32_16x6<f32>(16,6)@(256,4) where(AVX) quality(ManuallyOptimized));
MMMExternKernel!(avx_mmm_f32_24x4<f32>(24,4)@(256,4) where(AVX) quality(ManuallyOptimized));
MMMExternKernel!(avx_mmm_f32_32x3<f32>(32,3)@(256,4) where(AVX) quality(ManuallyOptimized));
MMMExternKernel!(avx_mmm_f32_40x2<f32>(40,2)@(256,4) where(AVX) quality(ManuallyOptimized));
MMMExternKernel!(avx_mmm_f32_64x1<f32>(64,1)@(256,4) where(AVX) quality(ManuallyOptimized));

MMMExternKernel!(fma_mmm_f32_8x8 <f32>(8, 8)@(256,4) where(FMA) quality(ManuallyOptimized));
MMMExternKernel!(fma_mmm_f32_16x6<f32>(16,6)@(256,4) where(FMA) quality(ManuallyOptimized));
MMMExternKernel!(fma_mmm_f32_16x5<f32>(16,5)@(256,4) where(FMA) quality(ManuallyOptimized));
Expand Down Expand Up @@ -222,6 +233,14 @@ MMMExternKernel! { avx512amx_mmm_f32_16x16<f32>(16,16)@(64,4) where(AVX512AMX_BF
}

pub fn plug(ops: &mut Ops) {
// The fma f32 tier below needs avx2 (vgatherdps) on top of fma; whenever it
// can't plug, cover every avx-capable CPU (Sandy/Ivy Bridge without fma,
// AMD Bulldozer-family with fma but no avx2) with the mul+add tier.
if is_x86_feature_detected!("avx")
&& !(is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma"))
{
plug_avx(ops);
}
if is_x86_feature_detected!("avx2") {
plug_avx2(ops);
// AVX-VNNI runs on AVX2-only Atom-class cores (Alder Lake-E, Sierra
Expand Down Expand Up @@ -381,6 +400,46 @@ pub fn plug_avx2(ops: &mut Ops) {
log::info!("qmmm_i32: x86_64/avx2 activated");
}

/// f32 kernels for AVX-capable CPUs that can't run the fma tier (Sandy/Ivy
/// Bridge without fma; AMD Bulldozer-family with fma but no avx2). Never
/// active alongside plug_fma: these kernels replace the generic fallback,
/// not the fma_ ones.
pub fn plug_avx(ops: &mut Ops) {
ops.mmm_impls.extend([
avx_mmm_f32_8x8.mmm(),
avx_mmm_f32_16x5.mmm(),
avx_mmm_f32_16x6.mmm(),
avx_mmm_f32_24x4.mmm(),
avx_mmm_f32_32x3.mmm(),
avx_mmm_f32_40x2.mmm(),
avx_mmm_f32_64x1.mmm(), // mmv candidate (nr==1; excluded from n>=2 picks)
]);

ops.mmv_f32 = Box::new(|_, _| avx_mmm_f32_64x1.mmm());

const AVX_CHOICES: &[KernelChoice] = &[
KernelChoice { mr: 16, nr: 6, scale: 1.0, ctor: || avx_mmm_f32_16x6.mmm() },
KernelChoice { mr: 16, nr: 5, scale: 0.98, ctor: || avx_mmm_f32_16x5.mmm() },
KernelChoice { mr: 24, nr: 4, scale: 0.95, ctor: || avx_mmm_f32_24x4.mmm() },
KernelChoice { mr: 32, nr: 3, scale: 0.93, ctor: || avx_mmm_f32_32x3.mmm() },
KernelChoice { mr: 40, nr: 2, scale: 0.90, ctor: || avx_mmm_f32_40x2.mmm() },
KernelChoice { mr: 8, nr: 8, scale: 0.80, ctor: || avx_mmm_f32_8x8.mmm() },
];
ops.mmm_f32 = Box::new(|m, _, n| match n {
None => avx_mmm_f32_16x6.mmm(),
Some(1) => avx_mmm_f32_64x1.mmm(),
Some(2) => avx_mmm_f32_40x2.mmm(),
Some(3) => avx_mmm_f32_32x3.mmm(),
Some(4) => avx_mmm_f32_24x4.mmm(),
Some(5) => avx_mmm_f32_16x5.mmm(),
Some(6) => avx_mmm_f32_16x6.mmm(),
Some(8) => avx_mmm_f32_8x8.mmm(),
Some(_) => pick_mmm(AVX_CHOICES, m, n),
});

log::info!("mmm_f32, mmv_f32: x86_64/avx (no fma) activated");
}

pub fn plug_fma(ops: &mut Ops) {
ops.mmm_impls.extend([
fma_mmm_f32_8x8.mmm(),
Expand Down
37 changes: 37 additions & 0 deletions linalg/x86_64/fma/2x5/packed_packed_loop1/mul-add.S.raw
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
// Tile size: 2x5, AVX without FMA (vmulps+vaddps)
// Accumulators: ymm0-9
// Row regs: ymm10-11
// Col reg: ymm14; mul scratch: ymm15

vmovaps ymm10, [rax]
vmovaps ymm11, [rax + 32]

vbroadcastss ymm14, dword ptr [rcx]
vmulps ymm15, ymm10, ymm14
vaddps ymm0, ymm0, ymm15
vmulps ymm15, ymm11, ymm14
vaddps ymm1, ymm1, ymm15

vbroadcastss ymm14, dword ptr [rcx + 4]
vmulps ymm15, ymm10, ymm14
vaddps ymm2, ymm2, ymm15
vmulps ymm15, ymm11, ymm14
vaddps ymm3, ymm3, ymm15

vbroadcastss ymm14, dword ptr [rcx + 8]
vmulps ymm15, ymm10, ymm14
vaddps ymm4, ymm4, ymm15
vmulps ymm15, ymm11, ymm14
vaddps ymm5, ymm5, ymm15

vbroadcastss ymm14, dword ptr [rcx + 12]
vmulps ymm15, ymm10, ymm14
vaddps ymm6, ymm6, ymm15
vmulps ymm15, ymm11, ymm14
vaddps ymm7, ymm7, ymm15

vbroadcastss ymm14, dword ptr [rcx + 16]
vmulps ymm15, ymm10, ymm14
vaddps ymm8, ymm8, ymm15
vmulps ymm15, ymm11, ymm14
vaddps ymm9, ymm9, ymm15
46 changes: 46 additions & 0 deletions linalg/x86_64/fma/2x6/packed_packed_loop1/mul-add.S.raw
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
// Tile size: 2x6, AVX without FMA (vmulps+vaddps)
// Accumulators: ymm0-11
// Row regs: ymm12-13
// Col reg: ymm14; mul scratch: ymm15

vmovaps ymm12, [rax]
vmovaps ymm13, [rax + 32]

vbroadcastss ymm14, dword ptr [rcx]
vmulps ymm15, ymm12, ymm14
vaddps ymm0, ymm0, ymm15
vmulps ymm15, ymm13, ymm14
vaddps ymm1, ymm1, ymm15

vbroadcastss ymm14, dword ptr [rcx + 4]
vmulps ymm15, ymm12, ymm14
vaddps ymm2, ymm2, ymm15
vmulps ymm15, ymm13, ymm14
vaddps ymm3, ymm3, ymm15

vbroadcastss ymm14, dword ptr [rcx + 8]
vmulps ymm15, ymm12, ymm14
vaddps ymm4, ymm4, ymm15
vmulps ymm15, ymm13, ymm14
vaddps ymm5, ymm5, ymm15

vbroadcastss ymm14, dword ptr [rcx + 12]
vmulps ymm15, ymm12, ymm14
vaddps ymm6, ymm6, ymm15
vmulps ymm15, ymm13, ymm14
vaddps ymm7, ymm7, ymm15

vbroadcastss ymm14, dword ptr [rcx + 16]
vmulps ymm15, ymm12, ymm14
vaddps ymm8, ymm8, ymm15
vmulps ymm15, ymm13, ymm14
vaddps ymm9, ymm9, ymm15

vbroadcastss ymm14, dword ptr [rcx + 20]
vmulps ymm15, ymm12, ymm14
vaddps ymm10, ymm10, ymm15
vmulps ymm15, ymm13, ymm14
vaddps ymm11, ymm11, ymm15

add rax, 64
add rcx, 24
48 changes: 48 additions & 0 deletions linalg/x86_64/fma/3x4/packed_packed_loop1/mul-add.S.raw
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
// Tile size: 3x4, AVX without FMA (vmulps+vaddps)
// Accumulators: ymm0-11
// Row regs: ymm12-14
// Col broadcast doubles as mul scratch: ymm15 (re-broadcast per row)

vmovaps ymm12, [rax]
vmovaps ymm13, [rax + 32]
vmovaps ymm14, [rax + 64]

vbroadcastss ymm15, dword ptr [rcx + 0]
vmulps ymm15, ymm15, ymm12
vaddps ymm0, ymm0, ymm15
vbroadcastss ymm15, dword ptr [rcx + 0]
vmulps ymm15, ymm15, ymm13
vaddps ymm1, ymm1, ymm15
vbroadcastss ymm15, dword ptr [rcx + 0]
vmulps ymm15, ymm15, ymm14
vaddps ymm2, ymm2, ymm15

vbroadcastss ymm15, dword ptr [rcx + 4]
vmulps ymm15, ymm15, ymm12
vaddps ymm3, ymm3, ymm15
vbroadcastss ymm15, dword ptr [rcx + 4]
vmulps ymm15, ymm15, ymm13
vaddps ymm4, ymm4, ymm15
vbroadcastss ymm15, dword ptr [rcx + 4]
vmulps ymm15, ymm15, ymm14
vaddps ymm5, ymm5, ymm15

vbroadcastss ymm15, dword ptr [rcx + 8]
vmulps ymm15, ymm15, ymm12
vaddps ymm6, ymm6, ymm15
vbroadcastss ymm15, dword ptr [rcx + 8]
vmulps ymm15, ymm15, ymm13
vaddps ymm7, ymm7, ymm15
vbroadcastss ymm15, dword ptr [rcx + 8]
vmulps ymm15, ymm15, ymm14
vaddps ymm8, ymm8, ymm15

vbroadcastss ymm15, dword ptr [rcx + 12]
vmulps ymm15, ymm15, ymm12
vaddps ymm9, ymm9, ymm15
vbroadcastss ymm15, dword ptr [rcx + 12]
vmulps ymm15, ymm15, ymm13
vaddps ymm10, ymm10, ymm15
vbroadcastss ymm15, dword ptr [rcx + 12]
vmulps ymm15, ymm15, ymm14
vaddps ymm11, ymm11, ymm15
55 changes: 55 additions & 0 deletions linalg/x86_64/fma/4x3/packed_packed_loop1/mul-add.S.raw
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
// Tile size: 4x3, AVX without FMA (vmulps+vaddps)
// Accumulators: ymm0-11
// Row reg: ymm12
// Col broadcasts double as mul scratch: ymm13-15 (re-broadcast per row)

vmovaps ymm12, [rax]

vbroadcastss ymm13, dword ptr [rcx + 0]
vmulps ymm13, ymm13, ymm12
vaddps ymm0, ymm0, ymm13
vbroadcastss ymm14, dword ptr [rcx + 4]
vmulps ymm14, ymm14, ymm12
vaddps ymm4, ymm4, ymm14
vbroadcastss ymm15, dword ptr [rcx + 8]
vmulps ymm15, ymm15, ymm12
vaddps ymm8, ymm8, ymm15

vmovaps ymm12, [rax + 32]

vbroadcastss ymm13, dword ptr [rcx + 0]
vmulps ymm13, ymm13, ymm12
vaddps ymm1, ymm1, ymm13
vbroadcastss ymm14, dword ptr [rcx + 4]
vmulps ymm14, ymm14, ymm12
vaddps ymm5, ymm5, ymm14
vbroadcastss ymm15, dword ptr [rcx + 8]
vmulps ymm15, ymm15, ymm12
vaddps ymm9, ymm9, ymm15

vmovaps ymm12, [rax + 64]

vbroadcastss ymm13, dword ptr [rcx + 0]
vmulps ymm13, ymm13, ymm12
vaddps ymm2, ymm2, ymm13
vbroadcastss ymm14, dword ptr [rcx + 4]
vmulps ymm14, ymm14, ymm12
vaddps ymm6, ymm6, ymm14
vbroadcastss ymm15, dword ptr [rcx + 8]
vmulps ymm15, ymm15, ymm12
vaddps ymm10, ymm10, ymm15

vmovaps ymm12, [rax + 96]

vbroadcastss ymm13, dword ptr [rcx + 0]
vmulps ymm13, ymm13, ymm12
vaddps ymm3, ymm3, ymm13
vbroadcastss ymm14, dword ptr [rcx + 4]
vmulps ymm14, ymm14, ymm12
vaddps ymm7, ymm7, ymm14
vbroadcastss ymm15, dword ptr [rcx + 8]
vmulps ymm15, ymm15, ymm12
vaddps ymm11, ymm11, ymm15

add rcx, 12
add rax, 128
Loading