Spaces:
Sleeping
Sleeping
Download native/kernels_sse.cpp from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 26.6 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_sse.cpp
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/native/kernels_sse.cpp
-
curl -L -o kernels_sse.cpp https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_sse.cpp
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. | |
| // ---- 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). | |
| namespace cism { | |
| 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 | |
| // ---- fp32 ------------------------------------------------------------------ | |
| float dot_sse41(const float* a, const float* b, std::size_t n) { | |
| // 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; | |
| return dot_scalar(a, b, n); | |
| } | |
| // ---- int8 x fp32 ------------------------------------------------------------ | |
| float dot_int8_sse41(const std::int8_t* weights, const float* input, std::size_t n) { | |
| // 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; | |
| return dot_int8_scalar(weights, input, n); | |
| } | |
| // ---- 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) { | |
| (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; | |
| (void)act_perm; | |
| return dot_int4_scalar(weights, act_orig, scales, n); | |
| } | |
| // ---- 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) { | |
| // 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; | |
| (void)act_perm; | |
| return dot_fp4_scalar(weights, act_orig, scales, n); | |
| } | |
| // ---- 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) { | |
| __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; | |
| // 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. | |
| return dot_sse41(a, b, n); | |
| return dot_scalar(a, b, n); | |
| } | |
| } // namespace cism | |