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