Spaces:
Sleeping
Sleeping
File size: 26,563 Bytes
28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 a9c1762 28a1a01 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 | // 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
|