File size: 24,872 Bytes
92dcc4e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
#include "kernels.hpp"

#include <algorithm>
#include <array>
#include <cmath>
#include <cstdlib>
#include <cstring>

#if (defined(CISM_HAVE_AVX2) || defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && defined(_MSC_VER)
#include <intrin.h>
#elif (defined(CISM_HAVE_AVX2) || defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && (defined(__GNUC__) || defined(__clang__))
#include <cpuid.h>
#endif

namespace cism {
// IEEE-754 binary16 conversion (round-to-nearest-even), portable: the AVX2
// TU uses vcvtph where available, but packing and scalar fallback need this
// everywhere (including ARM/SSE builds without F16C).
std::uint16_t fp32_to_fp16(float value) {
    std::uint32_t x;
    std::memcpy(&x, &value, 4);
    const std::uint16_t sign = static_cast<std::uint16_t>((x >> 16) & 0x8000u);
    const std::uint32_t ax = x & 0x7FFFFFFFu;
    const int e32 = static_cast<int>(ax >> 23);
    if (e32 == 255) return static_cast<std::uint16_t>(sign | 0x7BFFu);  // inf/NaN: saturate (weights are finite)
    const int e16 = e32 - 112;
    if (e16 >= 31) return static_cast<std::uint16_t>(sign | 0x7BFFu);  // overflow: max finite
    const std::uint32_t mant = ax & 0x7FFFFFu;
    if (e16 <= 0) {
        if (e16 < -10) return sign;  // underflow to zero
        // Subnormal: hidden 1 + RNE shift. shift = 126-E in [14,24].
        const int shift = 126 - e32;
        const std::uint32_t m = mant | 0x800000u;
        const std::uint32_t base = m >> shift;
        const std::uint32_t rbit = (m >> (shift - 1)) & 1u;
        const std::uint32_t sticky = m & ((shift == 1) ? 0u : ((1u << (shift - 1)) - 1u));
        const std::uint32_t half = base + ((rbit && (sticky || (base & 1u))) ? 1u : 0u);
        return static_cast<std::uint16_t>(sign | half);  // half <= 0x400: valid code
    }
    const std::uint32_t dropped = mant & 0x1FFFu;
    std::uint32_t half_mant = mant >> 13;
    if (dropped > 0x1000u || (dropped == 0x1000u && (half_mant & 1u))) {
        if (++half_mant == 0x400u) {
            // Mantissa carry: exact power of two (e16+1), except at e16==30
            // where it overflows to max finite (value was below 65504).
            if (e16 == 30) return static_cast<std::uint16_t>(sign | 0x7BFFu);
            return static_cast<std::uint16_t>(
                sign | (static_cast<std::uint32_t>(e16 + 1) << 10));
        }
    }
    return static_cast<std::uint16_t>(sign | (static_cast<std::uint32_t>(e16) << 10) | half_mant);
}

float fp16_to_fp32(std::uint16_t bits) {
    const std::uint32_t sign = (static_cast<std::uint32_t>(bits) & 0x8000u) << 16;
    const std::uint32_t exp = (bits >> 10) & 0x1Fu;
    const std::uint32_t mant = bits & 0x3FFu;
    std::uint32_t f;
    if (exp == 0) {
        if (mant == 0) {
            f = sign;
        } else {
            int e = -14;
            std::uint32_t m = mant;
            while (!(m & 0x400u)) {
                m <<= 1;
                --e;
            }
            m &= 0x3FFu;
            f = sign | (static_cast<std::uint32_t>(e + 127) << 23) | (m << 13);
        }
    } else if (exp == 31) {
        f = sign | 0x7F800000u | (mant << 13);
    } else {
        f = sign | ((exp + 112) << 23) | (mant << 13);
    }
    float out;
    std::memcpy(&out, &f, 4);
    return out;
}

const std::int8_t* fp4_element_lut() {
    // Raw E2M1 nibble (sign in bit 3) -> E2M1 value times two, exact in int8.
    static const std::int8_t lut[16] = {0, 1, 2, 3, 4, 6, 8, 12,
                                        0, -1, -2, -3, -4, -6, -8, -12};
    return lut;
}

float fp4_decode_scale(std::uint8_t bits) {
    // E4M3 (fn variant): subnormal mantissa/8 * 2^-6, normal (1+m/8) * 2^(e-7).
    const int exponent = (bits >> 3) & 15;
    const int mantissa = bits & 7;
    float value;
    if (exponent == 0) {
        value = static_cast<float>(mantissa) * (1.0f / 8.0f) * 0.015625f;
    } else {
        value = (1.0f + static_cast<float>(mantissa) / 8.0f) *
            std::ldexp(1.0f, exponent - 7);
    }
    return (bits & 128) ? -value : value;
}

const float* fp4_scale_lut() {
    // Decoded E4M3 scale, pre-halved: a half-value element times this entry
    // is exactly E2M1 * E4M3, so the vector FMA path needs no extra multiply.
    static const auto lut = []() {
        static std::array<float, 256> table;
        for (int bits = 0; bits < 256; ++bits) table[bits] = fp4_decode_scale(static_cast<std::uint8_t>(bits)) * 0.5f;
        return table;
    }();
    return lut.data();
}

std::uint8_t fp4_encode_scale(float value) {
    // Round-to-nearest-even into E4M3, saturating at the max normal (448),
    // matching torch.float8_e4m3fn for positive finite inputs.
    if (!(value > 0)) return 0;
    if (value >= 448.0f) return 0x7E;
    if (value < 0.015625f)  // smallest normal is 2^-6; below that, subnormals
        return static_cast<std::uint8_t>(std::nearbyint(value * 512.0f));
    int exponent = 0;
    std::frexp(value, &exponent);
    const int e = exponent - 1;                     // floor(log2(value))
    const float mantissa = std::ldexp(value, -e);   // [1, 2)
    int m = static_cast<int>(std::nearbyint((mantissa - 1.0f) * 8.0f));
    int eb = e + 7;
    if (m >= 8) { m = 0; ++eb; }
    if (eb > 15) return 0x7E;
    if (eb < 1)  // rounding landed below the normal range
        return static_cast<std::uint8_t>(std::nearbyint(value * 512.0f));
    return static_cast<std::uint8_t>((eb << 3) | m);
}

float dot_fp4_scalar(const std::uint8_t* weights, const float* input, const std::uint8_t* scales, std::size_t n) {
    const auto* elements = fp4_element_lut();
    const float* scale_lut = fp4_scale_lut();
    float sum = 0;
    for (std::size_t start = 0; start < n; start += 16) {
        const float scale = scale_lut[scales[start / 16]];
        float block_sum = 0;
        for (auto j = start; j < std::min(n, start + 16); ++j) {
            const int nibble = (weights[j / 2] >> (4 * (j % 2))) & 15;
            block_sum += static_cast<float>(elements[nibble]) * input[j];
        }
        sum += block_sum * scale;
    }
    return sum;
}
float dot_scalar(const float* a, const float* b, std::size_t n) {
    float result = 0;
    for (std::size_t i = 0; i < n; ++i) result += a[i] * b[i];
    return result;
}

float dot_int8_scalar(const std::int8_t* weights, const float* input, std::size_t n) {
    float sum = 0;
    for (std::size_t j = 0; j < n; ++j) sum += static_cast<float>(weights[j]) * input[j];
    return sum;
}

float dot_fp16_scalar(const std::uint16_t* weights, const float* input, std::size_t n) {
    float sum = 0;
    for (std::size_t j = 0; j < n; ++j) sum += fp16_to_fp32(weights[j]) * input[j];
    return sum;
}

float dot_int4_scalar(const std::uint8_t* weights, const float* input, const float* scales, std::size_t n) {
    // Split-half layout: byte j holds w[j] (low) and w[j+16] (high);
    // codes 0..15 map to (code-8), matching the packer exactly.
    float sum = 0;
    for (std::size_t start = 0; start < n; start += 32) {
        float block_sum = 0;
        const std::size_t m = std::min(n - start, static_cast<std::size_t>(32));
        for (std::size_t t = 0; t < m; ++t) {
            const int nibble = t < 16 ? (weights[t] & 15) : ((weights[t - 16] >> 4) & 15);
            block_sum += static_cast<float>(nibble - 8) * input[start + t];
        }
        weights += 16;
        sum += block_sum * scales[start / 32];
    }
    return sum;
}

// 5-arg adapters so the permuted kernel type keeps a working scalar
// fallback on non-AVX2 builds (the permuted buffer is simply unused).
float dot_int4_scalar5(const std::uint8_t* weights, const float* /*act_perm*/,

                       const float* input, const float* scales, std::size_t n) {
    return dot_int4_scalar(weights, input, scales, n);
}

float dot_fp4_scalar5(const std::uint8_t* weights, const float* /*act_perm*/,

                      const float* input, const std::uint8_t* scales, std::size_t n) {
    return dot_fp4_scalar(weights, input, scales, n);
}

// Canonical permuted-activation layout (scalar reference, bitwise stable):
// per full 32-chunk [c,c+32): OUT[c+k]=IN[c+2k], OUT[c+16+k]=IN[c+2k+1].
// Tail elements beyond (n/32)*32 are left undefined, matching the runtime
// permute_act32 contract. 64-wide inner blocking halves loop overhead with
// identical element order (two 32-blocks per iteration, in order).
void permute_act32_blocked(const float* input, std::size_t n, float* out) {
    std::size_t c = 0;
    for (; c + 64 <= n; c += 64) {
        for (std::size_t k = 0; k < 16; ++k) {
            out[c + k] = input[c + 2 * k];
            out[c + 16 + k] = input[c + 2 * k + 1];
        }
        for (std::size_t k = 0; k < 16; ++k) {
            out[c + 32 + k] = input[c + 32 + 2 * k];
            out[c + 48 + k] = input[c + 32 + 2 * k + 1];
        }
    }
    for (; c + 32 <= n; c += 32) {
        for (std::size_t k = 0; k < 16; ++k) {
            out[c + k] = input[c + 2 * k];
            out[c + 16 + k] = input[c + 2 * k + 1];
        }
    }
}

// Stable SiLU*up epilogue, bit-identical to Session::forward_tokens.
void silu_mul_scalar(const float* gate, const float* up, float* out, std::size_t n) {
    for (std::size_t i = 0; i < n; ++i) {
        const float value = gate[i];
        const float sigmoid = value >= 0 ? 1.0f / (1.0f + std::exp(-value)) :
            std::exp(value) / (1.0f + std::exp(value));
        out[i] = (value * sigmoid) * up[i];
    }
}

// Test-gated only: false unless CISM_FUSE_SILU=1. The runtime decode path
// never consults this (fusion stays off after the measured ~7% slowdown);
// it exists so experiments stay explicitly gated and PPL-checked.
bool silu_fusion_enabled() {
#ifdef _MSC_VER
#pragma warning(push)
#pragma warning(disable : 4996)
#endif
    const char* flag = std::getenv("CISM_FUSE_SILU");
#ifdef _MSC_VER
#pragma warning(pop)
#endif
    return flag != nullptr && flag[0] == '1' && flag[1] == '\0';
}

// Vector-exp activation blocks. Scalar loops use libm with the exact runtime
// formulas (bitwise reference); the AVX2 TU uses a ≤1-ULP polynomial exp.
// Dispatch resolved once per process (the CPU does not change under us).
namespace {
bool use_act_avx2() {
    static const bool cached =
#ifdef CISM_HAVE_AVX2
        has_avx2_cpu();
#else
        false;
#endif
    return cached;
}

inline float scalar_sigmoid(float x) {
    return x >= 0 ? 1.0f / (1.0f + std::exp(-x)) :
        std::exp(x) / (1.0f + std::exp(x));
}
}  // namespace

void act_exp(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
    if (use_act_avx2()) return act_exp_avx2(x, n);
#endif
    for (std::size_t i = 0; i < n; ++i) x[i] = std::exp(x[i]);
}

void act_silu(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
    if (use_act_avx2()) return act_silu_avx2(x, n);
#endif
    for (std::size_t i = 0; i < n; ++i) x[i] = x[i] * scalar_sigmoid(x[i]);
}

void act_silu_mul(const float* gate, const float* up, float* out, std::size_t n) {
#ifdef CISM_HAVE_AVX2
    if (use_act_avx2()) return act_silu_mul_avx2(gate, up, out, n);
#endif
    silu_mul_scalar(gate, up, out, n);
}

void act_sigmoid(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
    if (use_act_avx2()) return act_sigmoid_avx2(x, n);
#endif
    for (std::size_t i = 0; i < n; ++i) x[i] = scalar_sigmoid(x[i]);
}

void act_sigmoid_mul(float* o, const float* g, std::size_t n) {
#ifdef CISM_HAVE_AVX2
    if (use_act_avx2()) return act_sigmoid_mul_avx2(o, g, n);
#endif
    for (std::size_t i = 0; i < n; ++i) o[i] *= scalar_sigmoid(g[i]);
}

void act_silu_mul_plain(const float* gate, const float* up, float* out, std::size_t n) {
#ifdef CISM_HAVE_AVX2
    if (use_act_avx2()) return act_silu_mul_plain_avx2(gate, up, out, n);
#endif
    for (std::size_t i = 0; i < n; ++i) {
        const float value = gate[i];
        out[i] = (value * scalar_sigmoid(value)) * up[i];
    }
}

void act_gelu(float* x, std::size_t n) {
#ifdef CISM_HAVE_AVX2
    if (use_act_avx2()) return act_gelu_avx2(x, n);
#endif
    for (std::size_t i = 0; i < n; ++i) {
        const float value = x[i];
        x[i] = 0.5f * value * (1.0f + std::erf(value * 0.7071067811865475f));
    }
}

static bool has_avx2() {
#if defined(CISM_HAVE_AVX2) && defined(_MSC_VER)
    int regs[4];
    __cpuid(regs, 0);
    if (regs[0] < 7) return false;
    __cpuidex(regs, 1, 0);
    // Kernels use explicit FMA intrinsics, so require AVX, OSXSAVE, and FMA.
    constexpr int osxsave_avx_fma = (1 << 27) | (1 << 28) | (1 << 12);
    if ((regs[2] & osxsave_avx_fma) != osxsave_avx_fma) return false;
    if ((_xgetbv(0) & 6) != 6) return false;
    __cpuidex(regs, 7, 0);
    return (regs[1] & (1 << 5)) != 0;
#elif defined(CISM_HAVE_AVX2) && (defined(__GNUC__) || defined(__clang__))
    unsigned a, b, c, d;
    if (__get_cpuid_max(0, nullptr) < 7) return false;
    __cpuid_count(1, 0, a, b, c, d);
    constexpr unsigned osxsave_avx_fma = (1u << 27) | (1u << 28) | (1u << 12);
    if ((c & osxsave_avx_fma) != osxsave_avx_fma) return false;
    unsigned xcr0_low, xcr0_high;
    __asm__ volatile("xgetbv" : "=a"(xcr0_low), "=d"(xcr0_high) : "c"(0));
    if ((xcr0_low & 6) != 6) return false;
    __cpuid_count(7, 0, a, b, c, d);
    return (b & (1u << 5)) != 0;
#else
    return false;
#endif
}

// AVX-VNNI: CPUID leaf 7 subleaf 1 EAX[4] (Zen4/Zen5, Alder Lake+). VEX.256
// encoding needs only AVX OS state, so AVX2+FMA presence implies the state;
// still require has_avx2() so ancient/weird CPUs never take this branch.
static bool cpu_has_avx_vnni() {
#if (defined(CISM_HAVE_AVX_VNNI) || defined(CISM_HAVE_AVX512_VNNI)) && defined(_MSC_VER)
    if (!has_avx2()) return false;
    int regs[4];
    __cpuidex(regs, 7, 1);
    return (regs[0] & (1 << 4)) != 0;
#elif (defined(CISM_HAVE_AVX_VNNI) || defined(CISM_HAVE_AVX512_VNNI)) && (defined(__GNUC__) || defined(__clang__))
    if (!has_avx2()) return false;
    unsigned a, b, c, d;
    __cpuid_count(7, 1, a, b, c, d);
    return (a & (1u << 4)) != 0;
#else
    return false;
#endif
}

// AVX512-VNNI (EVEX encoding): leaf 7 subleaf 0 EBX[11] plus the AVX512
// foundation (F/BW/VL/DQ) and ZMM OS state (XCR0 opmask/ZMM_Hi256/Hi16_ZMM).
static bool cpu_has_avx512_vnni() {
#if defined(CISM_HAVE_AVX512_VNNI) && defined(_MSC_VER)
    if (!has_avx2()) return false;
    int regs[4];
    __cpuidex(regs, 7, 0);
    constexpr int want = (1 << 16) | (1 << 17) | (1 << 30) | (1 << 31) | (1 << 11);
    if ((regs[1] & want) != want) return false;
    return (_xgetbv(0) & 0xE0) == 0xE0;
#elif defined(CISM_HAVE_AVX512_VNNI) && (defined(__GNUC__) || defined(__clang__))
    if (!has_avx2()) return false;
    unsigned a, b, c, d;
    __cpuid_count(7, 0, a, b, c, d);
    constexpr unsigned want = (1u << 16) | (1u << 17) | (1u << 30) | (1u << 31) | (1u << 11);
    if ((b & want) != want) return false;
    unsigned lo, hi;
    __asm__ volatile("xgetbv" : "=a"(lo), "=d"(hi) : "c"(0));
    return (lo & 0xE0u) == 0xE0u;
#else
    return false;
#endif
}

bool has_avx2_cpu() { return has_avx2(); }

// F16C: leaf 1 ECX[29] plus the same AVX OS state as above (vcvtph needs
// YMM state). Gated on CISM_HAVE_F16C (compiler flag present).
static bool has_f16c() {
#if defined(CISM_HAVE_F16C) && defined(_MSC_VER)
    int regs[4];
    __cpuid(regs, 1);
    constexpr int osxsave_avx_fma_f16c = (1 << 27) | (1 << 28) | (1 << 12) | (1 << 29);
    if ((regs[2] & osxsave_avx_fma_f16c) != osxsave_avx_fma_f16c) return false;
    return (_xgetbv(0) & 6) == 6;
#elif defined(CISM_HAVE_F16C) && (defined(__GNUC__) || defined(__clang__))
    unsigned a, b, c, d;
    __cpuid(1, a, b, c, d);
    constexpr unsigned osxsave_avx_fma_f16c = (1u << 27) | (1u << 28) | (1u << 12) | (1u << 29);
    if ((c & osxsave_avx_fma_f16c) != osxsave_avx_fma_f16c) return false;
    unsigned xcr0_low, xcr0_high;
    __asm__ volatile("xgetbv" : "=a"(xcr0_low), "=d"(xcr0_high) : "c"(0));
    return (xcr0_low & 6) == 6;
#else
    return false;
#endif
}
bool has_f16c_cpu() { return has_f16c(); }

// SSE4.1: leaf 1 ECX[19]. SSSE3 (pshufb, needed by the int4/fp4 LUT path)
// is ECX[9]; require both so the nibble kernels never run without pshufb.
static bool has_sse41() {
#if (defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && defined(_MSC_VER)
    int regs[4];
    __cpuid(regs, 1);
    constexpr int want = (1 << 19) | (1 << 9);
    return (regs[2] & want) == want;
#elif (defined(CISM_HAVE_SSE41) || defined(CISM_HAVE_AVX)) && (defined(__GNUC__) || defined(__clang__))
    unsigned a, b, c, d;
    __cpuid(1, a, b, c, d);
    constexpr unsigned want = (1u << 19) | (1u << 9);
    return (c & want) == want;
#else
    return false;
#endif
}

// AVX (Sandy/Ivy/Bulldozer era): leaf 1 ECX[27,28] (OSXSAVE+AVX) plus XCR0
// SSE+AVX state. No FMA requirement (Sandy lacks it, Bulldozer has FMA4
// not FMA3); the AVX TU uses mul+add only.
static bool has_avx() {
#if defined(CISM_HAVE_AVX) && defined(_MSC_VER)
    if (!has_sse41()) return false;
    int regs[4];
    __cpuid(regs, 1);
    constexpr int want = (1 << 27) | (1 << 28);
    if ((regs[2] & want) != want) return false;
    return (_xgetbv(0) & 6) == 6;
#elif defined(CISM_HAVE_AVX) && (defined(__GNUC__) || defined(__clang__))
    if (!has_sse41()) return false;
    unsigned a, b, c, d;
    __cpuid(1, a, b, c, d);
    constexpr unsigned want = (1u << 27) | (1u << 28);
    if ((c & want) != want) return false;
    unsigned lo, hi;
    __asm__ volatile("xgetbv" : "=a"(lo), "=d"(hi) : "c"(0));
    return (lo & 6u) == 6u;
#else
    return false;
#endif
}

bool has_sse41_cpu() { return has_sse41(); }
bool has_avx_cpu() { return has_avx(); }

// The VNNI TU is compiled EVEX (MSVC /arch:AVX512) or VEX (GCC -mavxvnni);
// each encoding runs only on CPUs supporting exactly it.
bool has_avx_vnni_cpu() {
#if defined(CISM_HAVE_AVX512_VNNI) && !defined(CISM_HAVE_AVX_VNNI)
    (void)cpu_has_avx_vnni();
    return cpu_has_avx512_vnni();
#else
    return cpu_has_avx_vnni();
#endif
}

bool has_avx512_vnni_cpu() { return cpu_has_avx512_vnni(); }

bool has_neon_cpu() {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
    return true;
#else
    return false;
#endif
}

const char* kernel_variant() {
    if (has_neon_cpu()) return "neon";
    if (has_avx512_vnni_cpu()) return "avx512-vnni";
    if (has_avx_vnni_cpu()) return "avx-vnni";
    if (has_avx2_cpu()) return "avx2";
    if (has_avx_cpu()) return "avx";
    if (has_sse41_cpu()) return "sse41";
    return "scalar";
}

void quantize_row_i8(const float* input, std::size_t n, std::int8_t* values, float* scales) {
    // int8 sibling of runtime.cpp quantize_row: 127/absmax per 32-element
    // block, one fp32 dequant scale per block. Scalar reference for the VNNI
    // path; runs on every CPU (also the portable-kernels unit test target).
#ifdef CISM_HAVE_AVX2
    if (has_avx2()) {
        quantize_row_i8_avx2(input, n, values, scales);
        return;
    }
#endif
    for (std::size_t start = 0; start < n; start += 32) {
        const std::size_t end = std::min(start + 32, n);
        float absmax = 0.0f;
        for (std::size_t i = start; i < end; ++i) absmax = std::max(absmax, std::abs(input[i]));
        if (absmax == 0.0f) {
            for (std::size_t i = start; i < end; ++i) values[i] = 0;
            scales[start / 32] = 0.0f;
            continue;
        }
        const float norm = 127.0f / absmax;
        for (std::size_t i = start; i < end; ++i) {
            long q = std::lrint(static_cast<double>(input[i]) * norm);
            q = std::clamp<long>(q, -127, 127);
            values[i] = static_cast<std::int8_t>(q);
        }
        scales[start / 32] = absmax / 127.0f;
    }
}

Int8Q8VnniKernel int8_q8_vnni_kernel() {
    static const Int8Q8VnniKernel kernel = []() -> Int8Q8VnniKernel {
        if (has_avx_vnni_cpu()) return dot_int8_q8_vnni;
        return nullptr;
    }();
    return kernel;
}

DotKernel fp32_kernel() {
    static const DotKernel kernel = []() -> DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
        if (has_neon_cpu()) return dot_neon;
#endif
#if defined(CISM_HAVE_AVX512_VNNI)
        if (has_avx512_vnni_cpu()) return dot_fp32_avx512;
#endif
#ifdef CISM_HAVE_AVX2
        if (has_avx2()) return dot_avx2;
#endif
#ifdef CISM_HAVE_AVX
        if (has_avx()) return dot_avx_fp32;
#endif
#ifdef CISM_HAVE_SSE41
        if (has_sse41()) return dot_sse41;
#else
        (void)has_sse41();
        (void)has_avx();
#endif
        return dot_scalar;
    }();
    return kernel;
}

// AVX2-or-better family: everything dispatched on dot_avx2 stays valid when
// the zmm kernel leads (same ISA superset, same per-row results contract).
static bool use_avx2_kernels() {
    const auto selected = fp32_kernel();
    if (selected == dot_avx2) return true;
#ifdef CISM_HAVE_AVX512_VNNI
    if (selected == dot_fp32_avx512) return true;
#endif
    return false;
}

const char* kernel_name() {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
    if (fp32_kernel() == dot_neon) return "neon";
#endif
#ifdef CISM_HAVE_AVX
    if (fp32_kernel() == dot_avx_fp32) return "avx";
#endif
#ifdef CISM_HAVE_SSE41
    if (fp32_kernel() == dot_sse41) return "sse41";
#endif
#ifdef CISM_HAVE_AVX512_VNNI
    if (fp32_kernel() == dot_fp32_avx512) return "avx512";
#endif
    return fp32_kernel() == dot_scalar ? "scalar" : "avx2";
}

Int8DotKernel int8_kernel() {
    static const Int8DotKernel kernel = []() -> Int8DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
        if (fp32_kernel() == dot_neon) return dot_int8_neon;
#endif
#if defined(CISM_HAVE_AVX512_VNNI)
        // zmm int8 measured SLOWER than AVX2 on Xeon 8375C 2-vCPU (87 vs 99
        // tok/s: 512-bit throttling eats the width gain), so it is opt-in
        // only (CISM_ENABLE_AVX512_INT8=1) for experimentation/other tin.
        if (has_avx512_vnni_cpu()) {
            const char* flag = std::getenv("CISM_ENABLE_AVX512_INT8");
            if (flag != nullptr && flag[0] == '1' && flag[1] == '\0')
                return dot_int8_avx512;
        }
#endif
#ifdef CISM_HAVE_AVX2
        if (use_avx2_kernels()) return dot_int8_avx2;
#endif
#ifdef CISM_HAVE_SSE41
        if (has_sse41()) return dot_int8_sse41;
#endif
        return dot_int8_scalar;
    }();
    return kernel;
}

Fp16DotKernel fp16_kernel() {
    static const Fp16DotKernel kernel = []() -> Fp16DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
        if (fp32_kernel() == dot_neon) return dot_fp16_scalar;
#endif
#ifdef CISM_HAVE_F16C
        if (has_f16c_cpu()) return dot_fp16_avx2;
#endif
        return dot_fp16_scalar;
    }();
    return kernel;
}

Int4DotKernel int4_kernel() {
    static const Int4DotKernel kernel = []() -> Int4DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
        if (fp32_kernel() == dot_neon) return dot_int4_neon;
#endif
#ifdef CISM_HAVE_AVX2
        if (use_avx2_kernels()) return dot_int4_avx2;
#endif
#ifdef CISM_HAVE_SSE41
        if (has_sse41()) return dot_int4_sse41;
#endif
        return dot_int4_scalar5;
    }();
    return kernel;
}

Int8Q8DotKernel int8_q8_kernel() {
    static const Int8Q8DotKernel kernel = []() -> Int8Q8DotKernel {
#ifdef CISM_HAVE_AVX2
        if (use_avx2_kernels()) return dot_int8_q8_avx2;
#endif
        return nullptr;
    }();
    return kernel;
}

Int4Q8DotKernel int4_q8_kernel() {
    static const Int4Q8DotKernel kernel = []() -> Int4Q8DotKernel {
#ifdef CISM_HAVE_AVX2
        if (use_avx2_kernels()) return dot_int4_q8_avx2;
#endif
        return nullptr;
    }();
    return kernel;
}

Fp4Q8DotKernel fp4_q8_kernel() {
    static const Fp4Q8DotKernel kernel = []() -> Fp4Q8DotKernel {
#ifdef CISM_HAVE_AVX2
        if (use_avx2_kernels()) return dot_fp4_q8_avx2;
#endif
        return nullptr;
    }();
    return kernel;
}

Fp4DotKernel fp4_kernel() {
    static const Fp4DotKernel kernel = []() -> Fp4DotKernel {
#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
        if (fp32_kernel() == dot_neon) return dot_fp4_neon;
#endif
#ifdef CISM_HAVE_AVX2
        if (use_avx2_kernels()) return dot_fp4_avx2;
#endif
#ifdef CISM_HAVE_SSE41
        if (has_sse41()) return dot_fp4_sse41;
#endif
        return dot_fp4_scalar5;
    }();
    return kernel;
}
}  // namespace cism