test1111111 / native /kernels_sse.cpp
spitfire4794's picture
Fresh Surjo-50m bench on current code (all 4 quants) + HW inventory
a9c1762
Raw History Blame Contribute Delete
26.6 kB
// native/kernels_sse.cpp — SSE4.1 (+ AVX float) fallback kernels for pre-AVX2 x86.
//
// Purpose: give Nehalem/Sandy/Bulldozer-era CPUs a SIMD path instead of
// falling all the way to scalar. Haswell and newer stay on the AVX2 path
// (dot_avx2 / dot_int8_avx2 / ...); dispatch (kernels.cpp, NOT owned here)
// keeps preferring AVX2 whenever has_avx2_cpu() is true.
//
// CPU coverage (after the integrator wires dispatch):
// +--------------------------------+-------------------------------+-------------------------------------------
// | CPU family | ISA available | CISM path (intended) |
// +--------------------------------+-------------------------------+-------------------------------------------
// | Intel Nehalem/Westmere | SSE4.2 (superset of | dot_sse41 / dot_int8_sse41 / |
// | (2008-2010), old Xeons, | SSE4.1 + SSSE3) | dot_int4_sse41 / dot_fp4_sse41 |
// | Silvermont/Goldmont Atoms | | |
// +--------------------------------+-------------------------------+-------------------------------------------
// | Intel Sandy Bridge / | AVX (256-bit float only, | fp32: dot_avx_fp32 (AVX, mul+add) |
// | Ivy Bridge (2011-2012) | no integer AVX2, no FMA) | integers: SSE4.1 kernels above |
// +--------------------------------+-------------------------------+-------------------------------------------
// | AMD Bulldozer / Piledriver / | SSE4.2 + AVX + FMA4 + XOP | SAME as Sandy: SSE4.1 + AVX fp32 |
// | Steamroller (2011-2014) | NOTE: FMA4 only, NO FMA3. | Do NOT emit FMA3 (_mm256_fmadd_ps / |
// | | This TU uses mul+add only, | _mm_fmadd_ps) and do NOT use XOP/FMA4 |
// | | so it is safe on both Intel | intrinsics — mul+add runs everywhere. |
// | | and AMD of this era. | |
// +--------------------------------+-------------------------------+-------------------------------------------
// | Intel Haswell+ (2013-), | AVX2 + FMA3 | Stays on dot_avx2 / dot_int8_avx2 / ... |
// | AMD Excavator/Zen+ | | SSE TU unused (dispatch prefers AVX2). |
// +--------------------------------+-------------------------------+-------------------------------------------
//
// Conventions (mirror kernels_avx2.cpp):
// - 64-wide inner blocking, 4 independent accumulators, T0 prefetches,
// scalar tail. SIMD order may differ from scalar.
// - _mm_dp_ps is deliberately avoided (slow on Conroe/Nehalem-era cores);
// all reductions use mul+add.
// - int4/fp4 reuse the permuted-activation layout from runtime.cpp /
// kernels_avx2.cpp: per full 32-chunk [c,c+32), PERM[c+k]=ORIG[c+2k]
// (evens) and PERM[c+16+k]=ORIG[c+2k+1] (odds) for k=0..15. Low nibbles
// (evens) dot PERM[i..i+15], high nibbles (odds) dot PERM[i+16..i+31).
// Nibble decode uses _mm_shuffle_epi8 (pshufb — SSSE3, always present on
// any SSE4.1 CPU) with the int4/fp4 LUTs from kernels.hpp.
// - Numeric changes from SIMD reordering stay PPL-gated (see VALIDATION.md
// noise note), same rule as every other SIMD path.
// - Float FMA is NOT assumed anywhere in this TU (Sandy/Ivy lack it,
// Bulldozer has FMA4 not FMA3): use _mm_mul_ps + _mm_add_ps and
// _mm256_mul_ps + _mm256_add_ps only. No _mm_fmadd_* / _mm256_fmadd_*.
// - Integer kernels stay SSE4.1 (__m128 only) even when AVX is available:
// Sandy/Ivy have no integer _mm256 AVX2 ops, so no _mm256 integer
// intrinsic appears outside the AVX-guarded fp32 function.
// - Standalone compile: C++20 + SSE4.1, no AVX2 intrinsics outside the
// AVX-guarded dot_avx_fp32. A compiler invoked WITHOUT -msse4.1/-mavx
// (or without /arch:AVX on MSVC) still builds: every function falls back
// to the scalar reference from kernels.hpp (do not redeclare scalars
// here — include kernels.hpp).
//
// Integrator wiring (SHOW ONLY — do not apply here) is listed in the
// delivery message, not in this file.
#include "kernels.hpp"
#include <cstddef>
#include <cstdint>
#include <cstring>
#if defined(__i386__) || defined(__x86_64__) || defined(_M_IX86) || defined(_M_X64)
#if defined(_MSC_VER)
#include <intrin.h>
#else
#include <immintrin.h>
#endif
#endif
// ---- ISA availability ------------------------------------------------------
// CISM_SSE41_OK: __m128 float + SSSE3 shuffle + SSE4.1 int8->int32 cvt are
// usable in this TU. On MSVC, SSE4.1 intrinsics need no /arch flag (codegen
// flag only gates AVX+), so any x86 MSVC build can emit them; runtime gating
// (CPUID in kernels.cpp) still protects pre-SSE4.1 CPUs. On GCC/Clang the
// -msse4.1 (or -mavx/-mavx2, which imply it) flag is required, hence the
// __SSE4_1__ check. CISM_HAVE_SSE41 is also honored so the future CMake
// per-file flag (-msse4.1 / default on MSVC) can force-enable explicitly.
// CISM_AVX_OK: 256-bit FLOAT ops only (_mm256_loadu_ps/_mm256_mul_ps/
// _mm256_add_ps/cast/extract). No integer _mm256 ops here (those are AVX2).
#if !defined(CISM_SSE41_OK)
#if defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_SSE4_1) || defined(__SSE4_1__) || \
defined(__AVX__) || defined(__AVX2__) || \
(defined(_MSC_VER) && (defined(_M_X64) || defined(_M_IX86)))
#define CISM_SSE41_OK 1
#endif
#endif
#if !defined(CISM_AVX_OK)
#if defined(CISM_HAVE_AVX) || defined(__AVX__) || defined(__AVX2__) || \
(defined(_MSC_VER) && (defined(__AVX__) || defined(__AVX2__)))
#define CISM_AVX_OK 1
#endif
#endif
namespace cism {
#ifdef CISM_SSE41_OK
namespace {
// Horizontal sum of 4 float lanes, fixed order, deterministic.
inline float sse_reduce(__m128 value) {
__m128 sum = _mm_add_ps(value, _mm_movehl_ps(value, value));
sum = _mm_add_ss(sum, _mm_shuffle_ps(sum, sum, _MM_SHUFFLE(1, 1, 1, 1)));
return _mm_cvtss_f32(sum);
}
// Widen the low 4 int8 lanes to 4 floats (SSE4.1 _mm_cvtepi8_epi32).
inline __m128 sse_cvt8x4(__m128i v) {
return _mm_cvtepi32_ps(_mm_cvtepi8_epi32(v));
}
// Widen 4 int8 weights at an arbitrary (possibly unaligned) address.
// Used for the 4-wide vector tail; main loops use 16-byte loads + shifts.
inline __m128 sse_int8_4(const std::int8_t* weights) {
std::int32_t packed = 0;
std::memcpy(&packed, weights, sizeof(packed));
return sse_cvt8x4(_mm_cvtsi32_si128(static_cast<int>(packed)));
}
} // namespace
#endif
// ---- fp32 ------------------------------------------------------------------
float dot_sse41(const float* a, const float* b, std::size_t n) {
#ifdef CISM_SSE41_OK
// 64-wide inner blocking, same 4 accumulators in order (bitwise identical
// to the 16-wide loop below); T0 prefetches help DRAM-streaming spill
// sizes and are free for cache-resident rows (never fault).
// mul+add only: no _mm_dp_ps (slow), no FMA (not on Sandy/Bulldozer-FMA4).
__m128 sum0 = _mm_setzero_ps(), sum1 = _mm_setzero_ps();
__m128 sum2 = _mm_setzero_ps(), sum3 = _mm_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, a += 64, b += 64) {
_mm_prefetch(reinterpret_cast<const char*>(a + 128), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(b + 128), _MM_HINT_T0);
sum0 = _mm_add_ps(sum0, _mm_mul_ps(_mm_loadu_ps(a), _mm_loadu_ps(b)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(_mm_loadu_ps(a + 4), _mm_loadu_ps(b + 4)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(_mm_loadu_ps(a + 8), _mm_loadu_ps(b + 8)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(_mm_loadu_ps(a + 12), _mm_loadu_ps(b + 12)));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(_mm_loadu_ps(a + 16), _mm_loadu_ps(b + 16)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(_mm_loadu_ps(a + 20), _mm_loadu_ps(b + 20)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(_mm_loadu_ps(a + 24), _mm_loadu_ps(b + 24)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(_mm_loadu_ps(a + 28), _mm_loadu_ps(b + 28)));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(_mm_loadu_ps(a + 32), _mm_loadu_ps(b + 32)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(_mm_loadu_ps(a + 36), _mm_loadu_ps(b + 36)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(_mm_loadu_ps(a + 40), _mm_loadu_ps(b + 40)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(_mm_loadu_ps(a + 44), _mm_loadu_ps(b + 44)));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(_mm_loadu_ps(a + 48), _mm_loadu_ps(b + 48)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(_mm_loadu_ps(a + 52), _mm_loadu_ps(b + 52)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(_mm_loadu_ps(a + 56), _mm_loadu_ps(b + 56)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(_mm_loadu_ps(a + 60), _mm_loadu_ps(b + 60)));
}
for (; i + 16 <= n; i += 16, a += 16, b += 16) {
sum0 = _mm_add_ps(sum0, _mm_mul_ps(_mm_loadu_ps(a), _mm_loadu_ps(b)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(_mm_loadu_ps(a + 4), _mm_loadu_ps(b + 4)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(_mm_loadu_ps(a + 8), _mm_loadu_ps(b + 8)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(_mm_loadu_ps(a + 12), _mm_loadu_ps(b + 12)));
}
__m128 sum = _mm_add_ps(_mm_add_ps(sum0, sum1), _mm_add_ps(sum2, sum3));
for (; i + 4 <= n; i += 4, a += 4, b += 4)
sum = _mm_add_ps(sum, _mm_mul_ps(_mm_loadu_ps(a), _mm_loadu_ps(b)));
float result = sse_reduce(sum);
for (; i < n; ++i, ++a, ++b) result += *a * *b;
return result;
#else
return dot_scalar(a, b, n);
#endif
}
// ---- int8 x fp32 ------------------------------------------------------------
float dot_int8_sse41(const std::int8_t* weights, const float* input, std::size_t n) {
#ifdef CISM_SSE41_OK
// Four independent mul+add chains; 64-wide outer blocking (four 16-groups
// per iteration, same 4 accumulators in order) halves loop overhead with
// bitwise-identical accumulation vs the 16-wide loop. T0 prefetches mirror
// the AVX2 int8 kernel (weights+128 bytes, input+128 floats = 512 B).
__m128 sum0 = _mm_setzero_ps(), sum1 = _mm_setzero_ps();
__m128 sum2 = _mm_setzero_ps(), sum3 = _mm_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, weights += 64, input += 64) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 128), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(input + 128), _MM_HINT_T0);
for (std::size_t k = 0; k < 64; k += 16) {
const __m128i packed =
_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + k));
const __m128 w0 = sse_cvt8x4(packed);
const __m128 w1 = sse_cvt8x4(_mm_srli_si128(packed, 4));
const __m128 w2 = sse_cvt8x4(_mm_srli_si128(packed, 8));
const __m128 w3 = sse_cvt8x4(_mm_srli_si128(packed, 12));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w0, _mm_loadu_ps(input + k)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w1, _mm_loadu_ps(input + k + 4)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w2, _mm_loadu_ps(input + k + 8)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w3, _mm_loadu_ps(input + k + 12)));
}
}
for (; i + 16 <= n; i += 16, weights += 16, input += 16) {
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128 w0 = sse_cvt8x4(packed);
const __m128 w1 = sse_cvt8x4(_mm_srli_si128(packed, 4));
const __m128 w2 = sse_cvt8x4(_mm_srli_si128(packed, 8));
const __m128 w3 = sse_cvt8x4(_mm_srli_si128(packed, 12));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w0, _mm_loadu_ps(input)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w1, _mm_loadu_ps(input + 4)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w2, _mm_loadu_ps(input + 8)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w3, _mm_loadu_ps(input + 12)));
}
__m128 sum = _mm_add_ps(_mm_add_ps(sum0, sum1), _mm_add_ps(sum2, sum3));
for (; i + 4 <= n; i += 4, weights += 4, input += 4)
sum = _mm_add_ps(sum, _mm_mul_ps(sse_int8_4(weights), _mm_loadu_ps(input)));
float result = sse_reduce(sum);
for (; i < n; ++i, ++weights, ++input)
result += static_cast<float>(*weights) * *input;
return result;
#else
return dot_int8_scalar(weights, input, n);
#endif
}
// ---- int4 (split-half nibbles, linear activations) ---------------------------
float dot_int4_sse41(const std::uint8_t* weights, const float* act_perm, const float* act_orig,
const float* scales, std::size_t n) {
#ifdef CISM_SSE41_OK
(void)act_perm; // split-half dots against linear acts; no permute.
// Byte j holds w[j] (low) and w[j+16] (high): sub-8 (SSSE3-free zone
// uses plain integer sub, no LUT) then widen+scale, same order as AVX2.
// One fp32 scale per 32-block, broadcast with _mm_set1_ps; mul+add only.
// 64-wide blocking = two 32-blocks per iteration, same 4 accumulators.
const __m128i mask = _mm_set1_epi8(15);
const __m128i eight = _mm_set1_epi8(8);
__m128 sum0 = _mm_setzero_ps(), sum1 = _mm_setzero_ps();
__m128 sum2 = _mm_setzero_ps(), sum3 = _mm_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, weights += 32, act_orig += 64) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 32), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(act_orig + 64), _MM_HINT_T0);
for (std::size_t b = 0; b < 2; ++b) {
const std::size_t off_w = b * 16;
const std::size_t off_a = b * 32;
const std::size_t blk = i / 32 + b;
const __m128i packed =
_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + off_w));
const __m128i lo = _mm_sub_epi8(_mm_and_si128(packed, mask), eight);
const __m128i hi = _mm_sub_epi8(_mm_and_si128(_mm_srli_epi16(packed, 4), mask), eight);
const __m128 scale = _mm_set1_ps(scales[blk]);
const __m128 w0 = _mm_mul_ps(sse_cvt8x4(lo), scale);
const __m128 w1 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 4)), scale);
const __m128 w2 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 8)), scale);
const __m128 w3 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 12)), scale);
const __m128 w4 = _mm_mul_ps(sse_cvt8x4(hi), scale);
const __m128 w5 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 4)), scale);
const __m128 w6 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 8)), scale);
const __m128 w7 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 12)), scale);
const float* ap = act_orig + off_a;
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w0, _mm_loadu_ps(ap)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w1, _mm_loadu_ps(ap + 4)));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w2, _mm_loadu_ps(ap + 8)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w3, _mm_loadu_ps(ap + 12)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w4, _mm_loadu_ps(ap + 16)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w5, _mm_loadu_ps(ap + 20)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w6, _mm_loadu_ps(ap + 24)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w7, _mm_loadu_ps(ap + 28)));
}
}
for (; i + 32 <= n; i += 32, weights += 16, act_orig += 32) {
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128i lo = _mm_sub_epi8(_mm_and_si128(packed, mask), eight);
const __m128i hi = _mm_sub_epi8(_mm_and_si128(_mm_srli_epi16(packed, 4), mask), eight);
const __m128 scale = _mm_set1_ps(scales[i / 32]);
const __m128 w0 = _mm_mul_ps(sse_cvt8x4(lo), scale);
const __m128 w1 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 4)), scale);
const __m128 w2 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 8)), scale);
const __m128 w3 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 12)), scale);
const __m128 w4 = _mm_mul_ps(sse_cvt8x4(hi), scale);
const __m128 w5 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 4)), scale);
const __m128 w6 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 8)), scale);
const __m128 w7 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 12)), scale);
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w0, _mm_loadu_ps(act_orig)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w1, _mm_loadu_ps(act_orig + 4)));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w2, _mm_loadu_ps(act_orig + 8)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w3, _mm_loadu_ps(act_orig + 12)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w4, _mm_loadu_ps(act_orig + 16)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w5, _mm_loadu_ps(act_orig + 20)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w6, _mm_loadu_ps(act_orig + 24)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w7, _mm_loadu_ps(act_orig + 28)));
}
float result = sse_reduce(_mm_add_ps(_mm_add_ps(sum0, sum1), _mm_add_ps(sum2, sum3)));
if (i < n) {
// Partial tail block: split-half block-relative indexing via the
// scalar nibble code. i is a multiple of 32; weights points at the
// current 16-byte block (uniform stride), act_orig at element i.
result += dot_int4_scalar(weights, act_orig, scales + i / 32, n - i);
}
return result;
#else
(void)act_perm;
return dot_int4_scalar(weights, act_orig, scales, n);
#endif
}
// ---- fp4 (E2M1 elements, E4M3 scales, permuted activations) ------------------
float dot_fp4_sse41(const std::uint8_t* weights, const float* act_perm, const float* act_orig,
const std::uint8_t* scales, std::size_t n) {
#ifdef CISM_SSE41_OK
// Same layout/scale mapping as dot_fp4_avx2: lo[0..7]+hi[0..7] are row
// elements 0-15 (scale0), lo[8..15]+hi[8..15] are 16-31 (scale1) — the
// 16-element scale boundary cuts across the even/odd split. Elements are
// half-values (exact int8) via pshufb; the halved E4M3 table entry makes
// half_value * table == E2M1 * E4M3 with no extra multiply. mul+add only.
const __m128i mask = _mm_set1_epi8(15);
const __m128i lut = _mm_loadu_si128(reinterpret_cast<const __m128i*>(fp4_element_lut()));
const float* scale_lut = fp4_scale_lut();
__m128 sum0 = _mm_setzero_ps(), sum1 = _mm_setzero_ps();
__m128 sum2 = _mm_setzero_ps(), sum3 = _mm_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, weights += 32, act_perm += 64) {
_mm_prefetch(reinterpret_cast<const char*>(weights + 32), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(act_perm + 64), _MM_HINT_T0);
for (std::size_t b = 0; b < 2; ++b) {
const std::size_t off_w = b * 16;
const std::size_t off_a = b * 32;
const std::size_t sblk = i / 16 + b * 2;
const __m128i packed =
_mm_loadu_si128(reinterpret_cast<const __m128i*>(weights + off_w));
const __m128i low = _mm_and_si128(packed, mask);
const __m128i high = _mm_and_si128(_mm_srli_epi16(packed, 4), mask);
const __m128i lo = _mm_shuffle_epi8(lut, low);
const __m128i hi = _mm_shuffle_epi8(lut, high);
const __m128 scale0 = _mm_set1_ps(scale_lut[scales[sblk]]);
const __m128 scale1 = _mm_set1_ps(scale_lut[scales[sblk + 1]]);
const __m128 w0 = _mm_mul_ps(sse_cvt8x4(lo), scale0);
const __m128 w1 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 4)), scale0);
const __m128 w2 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 8)), scale1);
const __m128 w3 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 12)), scale1);
const __m128 w4 = _mm_mul_ps(sse_cvt8x4(hi), scale0);
const __m128 w5 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 4)), scale0);
const __m128 w6 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 8)), scale1);
const __m128 w7 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 12)), scale1);
const float* ap = act_perm + off_a;
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w0, _mm_loadu_ps(ap)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w1, _mm_loadu_ps(ap + 4)));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w2, _mm_loadu_ps(ap + 8)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w3, _mm_loadu_ps(ap + 12)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w4, _mm_loadu_ps(ap + 16)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w5, _mm_loadu_ps(ap + 20)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w6, _mm_loadu_ps(ap + 24)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w7, _mm_loadu_ps(ap + 28)));
}
}
for (; i + 32 <= n; i += 32, weights += 16, act_perm += 32) {
const __m128i packed = _mm_loadu_si128(reinterpret_cast<const __m128i*>(weights));
const __m128i low = _mm_and_si128(packed, mask);
const __m128i high = _mm_and_si128(_mm_srli_epi16(packed, 4), mask);
const __m128i lo = _mm_shuffle_epi8(lut, low);
const __m128i hi = _mm_shuffle_epi8(lut, high);
const std::size_t sblk = i / 16;
const __m128 scale0 = _mm_set1_ps(scale_lut[scales[sblk]]);
const __m128 scale1 = _mm_set1_ps(scale_lut[scales[sblk + 1]]);
const __m128 w0 = _mm_mul_ps(sse_cvt8x4(lo), scale0);
const __m128 w1 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 4)), scale0);
const __m128 w2 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 8)), scale1);
const __m128 w3 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(lo, 12)), scale1);
const __m128 w4 = _mm_mul_ps(sse_cvt8x4(hi), scale0);
const __m128 w5 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 4)), scale0);
const __m128 w6 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 8)), scale1);
const __m128 w7 = _mm_mul_ps(sse_cvt8x4(_mm_srli_si128(hi, 12)), scale1);
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w0, _mm_loadu_ps(act_perm)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w1, _mm_loadu_ps(act_perm + 4)));
sum0 = _mm_add_ps(sum0, _mm_mul_ps(w2, _mm_loadu_ps(act_perm + 8)));
sum1 = _mm_add_ps(sum1, _mm_mul_ps(w3, _mm_loadu_ps(act_perm + 12)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w4, _mm_loadu_ps(act_perm + 16)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w5, _mm_loadu_ps(act_perm + 20)));
sum2 = _mm_add_ps(sum2, _mm_mul_ps(w6, _mm_loadu_ps(act_perm + 24)));
sum3 = _mm_add_ps(sum3, _mm_mul_ps(w7, _mm_loadu_ps(act_perm + 28)));
}
float result = sse_reduce(_mm_add_ps(_mm_add_ps(sum0, sum1), _mm_add_ps(sum2, sum3)));
if (i < n) {
// Same tail contract as int4: PERM tail undefined, ORIGINAL-order
// row with the scalar nibble code. i is a multiple of 32, so the
// slice (weights+i/2, act_orig+i, scales+i/16) preserves parity and
// reads only to the logical end, including odd row widths.
result += dot_fp4_scalar(weights + i / 2, act_orig + i, scales + i / 16, n - i);
}
return result;
#else
(void)act_perm;
return dot_fp4_scalar(weights, act_orig, scales, n);
#endif
}
// ---- AVX (Sandy/Ivy Bridge) fp32-only fast path -----------------------------
// Uses 256-bit FLOAT ops only (_mm256_loadu_ps/_mm256_mul_ps/_mm256_add_ps +
// cast/extract for the reduce). No integer _mm256 ops (those are AVX2), no
// _mm256_fmadd_ps (no FMA on Sandy/Ivy; Bulldozer has FMA4, not FMA3), no
// _mm256_broadcast_ss (kept to plain loads/mul/add so /arch:AVX is enough).
// Integer kernels intentionally stay SSE4.1 on these CPUs.
float dot_avx_fp32(const float* a, const float* b, std::size_t n) {
#ifdef CISM_AVX_OK
__m256 sum0 = _mm256_setzero_ps(), sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps(), sum3 = _mm256_setzero_ps();
std::size_t i = 0;
for (; i + 64 <= n; i += 64, a += 64, b += 64) {
_mm_prefetch(reinterpret_cast<const char*>(a + 128), _MM_HINT_T0);
_mm_prefetch(reinterpret_cast<const char*>(b + 128), _MM_HINT_T0);
sum0 = _mm256_add_ps(sum0, _mm256_mul_ps(_mm256_loadu_ps(a), _mm256_loadu_ps(b)));
sum1 = _mm256_add_ps(sum1,
_mm256_mul_ps(_mm256_loadu_ps(a + 8), _mm256_loadu_ps(b + 8)));
sum2 = _mm256_add_ps(sum2,
_mm256_mul_ps(_mm256_loadu_ps(a + 16), _mm256_loadu_ps(b + 16)));
sum3 = _mm256_add_ps(sum3,
_mm256_mul_ps(_mm256_loadu_ps(a + 24), _mm256_loadu_ps(b + 24)));
sum0 = _mm256_add_ps(sum0,
_mm256_mul_ps(_mm256_loadu_ps(a + 32), _mm256_loadu_ps(b + 32)));
sum1 = _mm256_add_ps(sum1,
_mm256_mul_ps(_mm256_loadu_ps(a + 40), _mm256_loadu_ps(b + 40)));
sum2 = _mm256_add_ps(sum2,
_mm256_mul_ps(_mm256_loadu_ps(a + 48), _mm256_loadu_ps(b + 48)));
sum3 = _mm256_add_ps(sum3,
_mm256_mul_ps(_mm256_loadu_ps(a + 56), _mm256_loadu_ps(b + 56)));
}
for (; i + 32 <= n; i += 32, a += 32, b += 32) {
sum0 = _mm256_add_ps(sum0, _mm256_mul_ps(_mm256_loadu_ps(a), _mm256_loadu_ps(b)));
sum1 = _mm256_add_ps(sum1,
_mm256_mul_ps(_mm256_loadu_ps(a + 8), _mm256_loadu_ps(b + 8)));
sum2 = _mm256_add_ps(sum2,
_mm256_mul_ps(_mm256_loadu_ps(a + 16), _mm256_loadu_ps(b + 16)));
sum3 = _mm256_add_ps(sum3,
_mm256_mul_ps(_mm256_loadu_ps(a + 24), _mm256_loadu_ps(b + 24)));
}
__m256 sum = _mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3));
for (; i + 8 <= n; i += 8, a += 8, b += 8)
sum = _mm256_add_ps(sum, _mm256_mul_ps(_mm256_loadu_ps(a), _mm256_loadu_ps(b)));
// AVX-only reduce: 256 -> 2x128 -> scalar (no AVX2 integer ops).
__m128 lanes = _mm_add_ps(_mm256_castps256_ps128(sum), _mm256_extractf128_ps(sum, 1));
lanes = _mm_add_ps(lanes, _mm_movehl_ps(lanes, lanes));
lanes = _mm_add_ss(lanes, _mm_shuffle_ps(lanes, lanes, _MM_SHUFFLE(1, 1, 1, 1)));
float result = _mm_cvtss_f32(lanes);
for (; i < n; ++i, ++a, ++b) result += *a * *b;
return result;
#else
// Compiler without AVX (or ARM build): prefer the SSE4.1 path when it is
// available so Sandy-era shape is preserved; otherwise scalar. Keeps the
// symbol linkable everywhere; dispatch never selects it without AVX.
#ifdef CISM_SSE41_OK
return dot_sse41(a, b, n);
#else
return dot_scalar(a, b, n);
#endif
#endif
}
} // namespace cism