// 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 #include #include #if defined(__i386__) || defined(__x86_64__) || defined(_M_IX86) || defined(_M_X64) #if defined(_MSC_VER) #include #else #include #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(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(a + 128), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(weights + 128), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(input + 128), _MM_HINT_T0); for (std::size_t k = 0; k < 64; k += 16) { const __m128i packed = _mm_loadu_si128(reinterpret_cast(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(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(*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(weights + 32), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(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(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(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(weights + 32), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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(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(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(a + 128), _MM_HINT_T0); _mm_prefetch(reinterpret_cast(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