File size: 44,291 Bytes
28a1a01
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a9c1762
 
 
 
40d0cd2
 
 
a9c1762
40d0cd2
 
a9c1762
 
 
 
 
28a1a01
a9c1762
28a1a01
 
 
 
 
 
 
40d0cd2
28a1a01
40d0cd2
28a1a01
 
 
 
40d0cd2
28a1a01
40d0cd2
28a1a01
 
 
 
 
 
 
 
 
40d0cd2
28a1a01
 
 
 
40d0cd2
28a1a01
40d0cd2
28a1a01
40d0cd2
28a1a01
40d0cd2
28a1a01
 
 
 
40d0cd2
28a1a01
 
 
 
 
 
 
a9c1762
 
 
 
28a1a01
 
 
 
 
 
 
 
 
 
 
40d0cd2
28a1a01
 
 
 
 
40d0cd2
28a1a01
40d0cd2
28a1a01
40d0cd2
28a1a01
40d0cd2
28a1a01
 
 
 
40d0cd2
28a1a01
 
 
 
a9c1762
 
 
 
28a1a01
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
40d0cd2
 
 
 
 
 
 
 
 
 
28a1a01
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a9c1762
40d0cd2
 
28a1a01
 
 
a9c1762
 
 
 
 
28a1a01
a9c1762
40d0cd2
 
 
 
28a1a01
a9c1762
40d0cd2
 
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
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
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
// ARM64 NEON kernels for phones and ARM laptops/desktops (Android, iOS,
// Apple Silicon, Raspberry Pi). Same blocking shapes and semantics as the
// AVX2 kernels (64-wide inner blocking, 4 accumulators, permuted int4/fp4
// activations); SIMD order may differ from scalar, so numeric changes stay
// PPL-gated like every other SIMD path. Non-ARM builds compile scalar
// wrappers so the symbols always link (dispatch never selects them there).
//
// What runs where (this TU only; kernels.cpp dispatch is untouched):
// - Baseline NEON everywhere: ARMv8.0 FP+SIMD only (FMLA, tbl, uzp). No
//   dotprod / i8mm required. This is what kernels.cpp selects on any ARM64.
// - SDOT (ARMv8.2+dotprod, vdotq_s32): Snapdragon 888 (Kryo 680 = X1/A78)
//   has ASIMDDP in most configs. Probed at runtime INSIDE this TU via
//   getauxval(AT_HWCAP) & HWCAP_ASIMDDP on Linux/Android; Apple Silicon
//   always has it (compile note, no getauxval there). SDOT only accelerates
//   int8xint8 (Q8 activations) - it does NOT help the fp32-act paths below,
//   so dot_int8/int4/fp4_neon keep their widening MLA loops; the Q8
//   SDOT/widening pairs live here as dot_int8_q8_*_impl,
//   dot_int4_q8_*_impl, dot_fp4_q8_*_impl (+ *_neon_internal dispatch) for
//   future wiring (kernels.cpp currently exposes no NEON Q8 kernel, so the
//   fp32-act paths stay the live ones; integer dots are exact, so SDOT vs
//   widening is results-neutral, unlike fp32 reorder which stays PPL-gated).
//   int4/fp4 Q8 decode nibbles via vtbl, vzip even/odd halves back to
//   sequential int8, then SDOT into int32 with a float scale epilogue.
// - i8mm (ARMv8.6 usmmla): Snapdragon 888 does NOT have it. Commented stub
//   only, guarded by __ARM_FEATURE_MATMUL_INT8, never selected (would need
//   AT_HWCAP2/HWCAP2_I8MM which the 888 lacks). Baseline build unaffected.
//
// Blocking (AVX2 standard, adjusted for 128-bit vectors):
// - dot_fp32: 64-wide (16 FMA) + 32-wide (8 FMA, bitwise-identical unroll of
//   two 16-groups) + 16-wide (4 FMA = AVX2's 32-wide, same 4 vectors) + 4-wide
//   + scalar. 4 independent accumulators throughout.
// - dot_int8: 64 + 32 + 8 + scalar (8-weight groups, same element counts as
//   AVX2). Prefetch only in the 64-wide loop, like AVX2.
// - dot_int4/fp4: 64-wide (two 32-blocks) + 32-wide + scalar tail that reads
//   only to the logical end from the ORIGINAL-order pointer. int4 scales are
//   per-32 (one scale covers both halves); fp4 scales are per-16 and the
//   16-element boundary CUTS ACROSS the even/odd split (same as AVX2), so each
//   32-block needs 4 scaled quarters, not 2 scaled halves.
// - Prefetches mirror AVX2 element distances (fp32 +128 floats = 512 B,
//   int8 +128, int4/fp4 weights +32 B / act_perm +64 floats). On the 888
//   (64 B lines, 32-64 KB L1D, ~512 KB L2) +128 floats is 8 lines / 512 B
//   ahead: helps DRAM-streaming spill sizes, free for L1-resident rows,
//   never faults (skipped when n < 64).
//
// Termux owner (S21 FE): measure, don't guess. See the validation list at
// the bottom of this file's commit message / task return: PPL parity
// (fp32 vs int8 vs hybrid-int4/fp4), tokens/s single-thread, HWCAP check
// (getauxval ASIMDDP present?), and the fp4 scale-boundary numpy check.

#include "kernels.hpp"

#include <cstddef>
#include <cstdint>

#if defined(CISM_HAVE_NEON) && defined(__aarch64__)
#include <arm_neon.h>

#if defined(__linux__)
#include <sys/auxv.h>
// HWCAP_ASIMDDP lives in <asm/hwcap.h> on glibc but <sys/auxv.h> already
// defines it on Bionic (Android/Termux). Try the asm header when present,
// otherwise fall back to the architectural bit number below.
#if defined(__has_include)
#if __has_include(<asm/hwcap.h>)
#include <asm/hwcap.h>
#endif
#endif
#endif

namespace cism {

namespace neon_detail {
// Architectural HWCAP bit for ASIMDDP on AArch64 Linux (bit 20). Used only
// when the libc headers did not already define HWCAP_ASIMDDP.
#if defined(__linux__) && defined(AT_HWCAP) && defined(HWCAP_ASIMDDP)
constexpr unsigned long kHwcapDotprod = HWCAP_ASIMDDP;
#elif defined(__linux__) && defined(AT_HWCAP)
constexpr unsigned long kHwcapDotprod = (1UL << 20);
#else
constexpr unsigned long kHwcapDotprod = 0UL;
#endif

// Runtime dotprod probe, TU-internal only (kernels.cpp is untouched).
// - Apple Silicon (M1+): always has dotprod, no getauxval needed.
// - Linux/Android (Termux on S21 FE): getauxval(AT_HWCAP) & ASIMDDP.
// - Other ARM64 OSes: compile-time feature only, else false (safe).
[[maybe_unused]] inline bool has_dotprod_runtime() {
#if defined(__APPLE__) && defined(__aarch64__)
    (void)kHwcapDotprod;
    return true;  // All Apple Silicon ships ARMv8.2+dotprod or later.
#elif defined(__linux__) && defined(__aarch64__) && defined(AT_HWCAP)
    return (getauxval(AT_HWCAP) & kHwcapDotprod) != 0UL;
#elif defined(__ARM_FEATURE_DOTPROD)
    return true;  // TU built with +dotprod: encoding is always legal.
#else
    return false;
#endif
}
}  // namespace neon_detail

float dot_neon(const float* a, const float* b, std::size_t n) {
    float32x4_t sum0 = vdupq_n_f32(0), sum1 = vdupq_n_f32(0);
    float32x4_t sum2 = vdupq_n_f32(0), sum3 = vdupq_n_f32(0);
    std::size_t i = 0;
    for (; i + 64 <= n; i += 64, a += 64, b += 64) {
        // +128 floats = 512 B ahead = AVX2's kPrefetchDistance; 8 lines on
        // the 888, sensible for L1D 32-64 KB, free when cache-resident.
        __builtin_prefetch(a + 128, 0, 3);
        __builtin_prefetch(b + 128, 0, 3);
        sum0 = vfmaq_f32(sum0, vld1q_f32(a), vld1q_f32(b));
        sum1 = vfmaq_f32(sum1, vld1q_f32(a + 4), vld1q_f32(b + 4));
        sum2 = vfmaq_f32(sum2, vld1q_f32(a + 8), vld1q_f32(b + 8));
        sum3 = vfmaq_f32(sum3, vld1q_f32(a + 12), vld1q_f32(b + 12));
        sum0 = vfmaq_f32(sum0, vld1q_f32(a + 16), vld1q_f32(b + 16));
        sum1 = vfmaq_f32(sum1, vld1q_f32(a + 20), vld1q_f32(b + 20));
        sum2 = vfmaq_f32(sum2, vld1q_f32(a + 24), vld1q_f32(b + 24));
        sum3 = vfmaq_f32(sum3, vld1q_f32(a + 28), vld1q_f32(b + 28));
        sum0 = vfmaq_f32(sum0, vld1q_f32(a + 32), vld1q_f32(b + 32));
        sum1 = vfmaq_f32(sum1, vld1q_f32(a + 36), vld1q_f32(b + 36));
        sum2 = vfmaq_f32(sum2, vld1q_f32(a + 40), vld1q_f32(b + 40));
        sum3 = vfmaq_f32(sum3, vld1q_f32(a + 44), vld1q_f32(b + 44));
        sum0 = vfmaq_f32(sum0, vld1q_f32(a + 48), vld1q_f32(b + 48));
        sum1 = vfmaq_f32(sum1, vld1q_f32(a + 52), vld1q_f32(b + 52));
        sum2 = vfmaq_f32(sum2, vld1q_f32(a + 56), vld1q_f32(b + 56));
        sum3 = vfmaq_f32(sum3, vld1q_f32(a + 60), vld1q_f32(b + 60));
    }
    // 32-wide: literal AVX2 element count (8 vectors), same 4 accumulators in
    // order - bitwise identical to two 16-wide iterations, halves loop
    // overhead for 32 <= remainder < 64. No prefetch (matches AVX2: prefetches
    // live only in the 64-wide loop).
    for (; i + 32 <= n; i += 32, a += 32, b += 32) {
        sum0 = vfmaq_f32(sum0, vld1q_f32(a), vld1q_f32(b));
        sum1 = vfmaq_f32(sum1, vld1q_f32(a + 4), vld1q_f32(b + 4));
        sum2 = vfmaq_f32(sum2, vld1q_f32(a + 8), vld1q_f32(b + 8));
        sum3 = vfmaq_f32(sum3, vld1q_f32(a + 12), vld1q_f32(b + 12));
        sum0 = vfmaq_f32(sum0, vld1q_f32(a + 16), vld1q_f32(b + 16));
        sum1 = vfmaq_f32(sum1, vld1q_f32(a + 20), vld1q_f32(b + 20));
        sum2 = vfmaq_f32(sum2, vld1q_f32(a + 24), vld1q_f32(b + 24));
        sum3 = vfmaq_f32(sum3, vld1q_f32(a + 28), vld1q_f32(b + 28));
    }
    // 16-wide = AVX2's 32-wide in vectors (4 FMA chains, one vector each).
    for (; i + 16 <= n; i += 16, a += 16, b += 16) {
        sum0 = vfmaq_f32(sum0, vld1q_f32(a), vld1q_f32(b));
        sum1 = vfmaq_f32(sum1, vld1q_f32(a + 4), vld1q_f32(b + 4));
        sum2 = vfmaq_f32(sum2, vld1q_f32(a + 8), vld1q_f32(b + 8));
        sum3 = vfmaq_f32(sum3, vld1q_f32(a + 12), vld1q_f32(b + 12));
    }
    float32x4_t sum = vaddq_f32(vaddq_f32(sum0, sum1), vaddq_f32(sum2, sum3));
    float result = vaddvq_f32(sum);
    for (; i + 4 <= n; i += 4, a += 4, b += 4)
        result += vaddvq_f32(vmulq_f32(vld1q_f32(a), vld1q_f32(b)));
    for (; i < n; ++i, ++a, ++b) result += *a * *b;
    return result;
}

// Widen 8 int8 weights to float32x4 pairs via s8->s16->s32->f32.
inline void int8_to_f32x4(const std::int8_t* w, float32x4_t& lo, float32x4_t& hi) {
    int8x8_t v = vld1_s8(w);
    int16x8_t w16 = vmovl_s8(v);
    lo = vcvtq_f32_s32(vmovl_s16(vget_low_s16(w16)));
    hi = vcvtq_f32_s32(vmovl_s16(vget_high_s16(w16)));
}

// SDOT decision (documented here because dot_int8_neon is the natural place
// to look): vdotq_s32 computes int8xint8 -> int32, so it does NOT accelerate
// this int8-weights x fp32-activations kernel - the weights still need the
// s8->f32 widen above before they can FMA against float acts. SDOT only helps
// the int8-ACTIVATION (Q8) path, which has no NEON kernel wired in kernels.cpp
// yet. Hence dot_int8_neon keeps the widening MLA loop below unconditionally,
// and the SDOT machinery lives in the TU-internal Q8 helpers that follow
// (dot_int8_q8_*_impl + dot_int8_q8_neon_internal) with a wiring note.
float dot_int8_neon(const std::int8_t* weights, const float* input, std::size_t n) {
    float32x4_t sum0 = vdupq_n_f32(0), sum1 = vdupq_n_f32(0);
    float32x4_t sum2 = vdupq_n_f32(0), sum3 = vdupq_n_f32(0);
    std::size_t i = 0;
    for (; i + 64 <= n; i += 64, weights += 64, input += 64) {
        __builtin_prefetch(weights + 128, 0, 3);
        __builtin_prefetch(input + 128, 0, 3);
        for (int k = 0; k < 8; ++k) {
            float32x4_t wlo, whi;
            int8_to_f32x4(weights + 8 * k, wlo, whi);
            float32x4_t* acc = k % 4 == 0 ? &sum0 : k % 4 == 1 ? &sum1 : k % 4 == 2 ? &sum2 : &sum3;
            *acc = vfmaq_f32(*acc, wlo, vld1q_f32(input + 8 * k));
            *acc = vfmaq_f32(*acc, whi, vld1q_f32(input + 8 * k + 4));
        }
    }
    for (; i + 32 <= n; i += 32, weights += 32, input += 32) {
        for (int k = 0; k < 4; ++k) {
            float32x4_t wlo, whi;
            int8_to_f32x4(weights + 8 * k, wlo, whi);
            float32x4_t* acc = k == 0 ? &sum0 : k == 1 ? &sum1 : k == 2 ? &sum2 : &sum3;
            *acc = vfmaq_f32(*acc, wlo, vld1q_f32(input + 8 * k));
            *acc = vfmaq_f32(*acc, whi, vld1q_f32(input + 8 * k + 4));
        }
    }
    float32x4_t sum = vaddq_f32(vaddq_f32(sum0, sum1), vaddq_f32(sum2, sum3));
    float result = vaddvq_f32(sum);
    for (; i + 8 <= n; i += 8, weights += 8, input += 8) {
        float32x4_t wlo, whi;
        int8_to_f32x4(weights, wlo, whi);
        result += vaddvq_f32(vmulq_f32(wlo, vld1q_f32(input)));
        result += vaddvq_f32(vmulq_f32(whi, vld1q_f32(input + 4)));
    }
    for (; i < n; ++i, ++weights, ++input)
        result += static_cast<float>(*weights) * *input;
    return result;
}

// ---- TU-internal int8-activation (Q8) helpers: SDOT vs widening ----------
// Future-wiring only: kernels.cpp exposes no NEON Q8 kernel today (int8_q8 /
// vnni selectors return nullptr on ARM), so nothing outside this TU calls
// these. They exist so the S21 FE SDOT decision is implemented and reviewable
// here without touching dispatch. Integer dots are exact, and both impls share
// the same 128/32/tail blocking with the same float-accumulation order, so
// SDOT vs widening is results-neutral (no PPL gate needed between them; the
// Q8-vs-fp32-act choice itself stays PPL-gated like AVX2's Q8 path).
// Signature mirrors the x86 VNNI contract: int8 weights, int8 acts quantized
// per-32 (127/absmax, e.g. quantize_row_i8), one fp32 dequant scale per
// 32-block; the caller multiplies the per-row weight scale outside, exactly
// like Matrix::multiply does for dot_int8_q8_avx2 / dot_int8_q8_vnni.

// Exact dot of 16 int8 pairs via widening MLA (runs on every ARM64).
[[maybe_unused]] static inline std::int32_t dot16_s8_widening(const std::int8_t* w,
                                                              const std::int8_t* a) {
    int8x16_t wv = vld1q_s8(w);
    int8x16_t av = vld1q_s8(a);
    int16x8_t lo = vmull_s8(vget_low_s8(wv), vget_low_s8(av));
    int16x8_t hi = vmull_s8(vget_high_s8(wv), vget_high_s8(av));
    int32x4_t acc = vpadalq_s16(vdupq_n_s32(0), lo);
    acc = vpadalq_s16(acc, hi);
    return vaddvq_s32(acc);
}

[[maybe_unused]] static float dot_int8_q8_widening_impl(const std::int8_t* weights,
                                                        const std::int8_t* act,
                                                        const float* act_scales, std::size_t n) {
    // 128-wide inner blocking (four 32-groups, one float acc each) + 32-wide
    // + scalar tail to the logical end. Four independent float chains hide
    // convert/multiply latency, same shape as dot_int8_q8_avx2's 128-wide.
    float acc0 = 0, acc1 = 0, acc2 = 0, acc3 = 0;
    std::size_t i = 0;
    for (; i + 128 <= n; i += 128, weights += 128, act += 128) {
        __builtin_prefetch(weights + 256, 0, 3);
        __builtin_prefetch(act + 256, 0, 3);
        acc0 += static_cast<float>(dot16_s8_widening(weights, act) +
                                   dot16_s8_widening(weights + 16, act + 16)) *
                act_scales[i / 32];
        acc1 += static_cast<float>(dot16_s8_widening(weights + 32, act + 32) +
                                   dot16_s8_widening(weights + 48, act + 48)) *
                act_scales[i / 32 + 1];
        acc2 += static_cast<float>(dot16_s8_widening(weights + 64, act + 64) +
                                   dot16_s8_widening(weights + 80, act + 80)) *
                act_scales[i / 32 + 2];
        acc3 += static_cast<float>(dot16_s8_widening(weights + 96, act + 96) +
                                   dot16_s8_widening(weights + 112, act + 112)) *
                act_scales[i / 32 + 3];
    }
    float result = (acc0 + acc1) + (acc2 + acc3);
    for (; i + 32 <= n; i += 32, weights += 32, act += 32)
        result += static_cast<float>(dot16_s8_widening(weights, act) +
                                     dot16_s8_widening(weights + 16, act + 16)) *
                  act_scales[i / 32];
    if (i < n) {
        // Partial tail block, read only to the row's logical end.
        float block_sum = 0;
        const float scale = act_scales[i / 32];
        for (std::size_t j = 0; i < n; ++i, ++j)
            block_sum += static_cast<float>(weights[j]) * static_cast<float>(act[j]);
        result += block_sum * scale;
    }
    return result;
}

#ifdef __ARM_FEATURE_DOTPROD
// Exact dot of 16 int8 pairs via SDOT (needs ARMv8.2+dotprod encoding).
[[maybe_unused]] static inline std::int32_t dot16_s8_sdot(const std::int8_t* w, const std::int8_t* a) {
    int32x4_t acc = vdupq_n_s32(0);
    acc = vdotq_s32(acc, vld1q_s8(w), vld1q_s8(a));
    return vaddvq_s32(acc);
}

[[maybe_unused]] static float dot_int8_q8_sdot_impl(const std::int8_t* weights, const std::int8_t* act,
                                                    const float* act_scales, std::size_t n) {
    // Identical blocking/order to the widening impl above (only the 16-dot
    // primitive differs), so the two are bitwise identical end to end.
    float acc0 = 0, acc1 = 0, acc2 = 0, acc3 = 0;
    std::size_t i = 0;
    for (; i + 128 <= n; i += 128, weights += 128, act += 128) {
        __builtin_prefetch(weights + 256, 0, 3);
        __builtin_prefetch(act + 256, 0, 3);
        acc0 += static_cast<float>(dot16_s8_sdot(weights, act) +
                                   dot16_s8_sdot(weights + 16, act + 16)) *
                act_scales[i / 32];
        acc1 += static_cast<float>(dot16_s8_sdot(weights + 32, act + 32) +
                                   dot16_s8_sdot(weights + 48, act + 48)) *
                act_scales[i / 32 + 1];
        acc2 += static_cast<float>(dot16_s8_sdot(weights + 64, act + 64) +
                                   dot16_s8_sdot(weights + 80, act + 80)) *
                act_scales[i / 32 + 2];
        acc3 += static_cast<float>(dot16_s8_sdot(weights + 96, act + 96) +
                                   dot16_s8_sdot(weights + 112, act + 112)) *
                act_scales[i / 32 + 3];
    }
    float result = (acc0 + acc1) + (acc2 + acc3);
    for (; i + 32 <= n; i += 32, weights += 32, act += 32)
        result += static_cast<float>(dot16_s8_sdot(weights, act) +
                                     dot16_s8_sdot(weights + 16, act + 16)) *
                  act_scales[i / 32];
    if (i < n) {
        float block_sum = 0;
        const float scale = act_scales[i / 32];
        for (std::size_t j = 0; i < n; ++i, ++j)
            block_sum += static_cast<float>(weights[j]) * static_cast<float>(act[j]);
        result += block_sum * scale;
    }
    return result;
}
#endif

// TU-internal dispatch: SDOT when the CPU reports ASIMDDP, else widening.
// Wiring note for a future kernels.cpp change (OUT OF SCOPE here): expose a
// NEON int8_q8 selector returning this function once the runtime quantizes
// shared activations to int8 per-32 (quantize_row_i8 ABI) and multiplies the
// per-row weight scale at the call site, mirroring the AVX2 Q8/VNNI branches
// in Matrix::multiply/gemm. Until then this stays internal-only so dispatch
// behavior is byte-for-byte unchanged on every phone.
[[maybe_unused]] static float dot_int8_q8_neon_internal(const std::int8_t* weights,
                                                        const std::int8_t* act,
                                                        const float* act_scales, std::size_t n) {
    if (neon_detail::has_dotprod_runtime()) {
#ifdef __ARM_FEATURE_DOTPROD
        return dot_int8_q8_sdot_impl(weights, act, act_scales, n);
#else
        // CPU has SDOT but this TU was built without the +dotprod encoding
        // (baseline flags per CMakeLists), so vdotq_s32 is not compiled in.
        // Fall through to widening; rebuild with -march=armv8.2-a+dotprod
        // (or an SDOT-only sub-TU) to unlock the fast path. Correctness first.
#endif
    }
    return dot_int8_q8_widening_impl(weights, act, act_scales, n);
}

// Forward declaration: defined below beside the live int4/fp4 kernels
// (vtbl decode shared by the live fp32-act path and the Q8 helpers here).
inline int8x16_t nibbles_to_s8(const std::uint8_t* packed, int low, const int8x16_t& lut);

// ---- TU-internal int4/fp4 Q8 helpers: vtbl decode routed through SDOT -----
// Future-wiring only (same status as the int8 Q8 pair above): kernels.cpp
// exposes no NEON Q8 kernels, so nothing outside this TU calls these and no
// kernels.hpp signature changes. The live fp32-activation int4/fp4 kernels
// keep widen+FMLA because SDOT needs int8xint8 inputs; THESE helpers take
// int8 activations (VNNI-style ABI: quantize_row_i8 per-32, one fp32 dequant
// scale per 32-block; int4 weight scale per-32 fp32, fp4 weight scale per-16
// E4M3 byte, caller-side like the AVX2 Q8 branches), so the vtbl-decoded
// int8 weights dot via SDOT into int32 with a float scale epilogue.
// Widening fallback is results-neutral (integer dots exact; shared blocking
// and float accumulation order), guarded by __ARM_FEATURE_DOTPROD so the
// baseline build (no +dotprod flags per CMakeLists) never sees the encoding.

// Dot of two int8x16 vectors via widening MLA (baseline ARMv8.0).
[[maybe_unused]] static inline std::int32_t dot16_vec_widening(int8x16_t wv, int8x16_t av) {
    int16x8_t lo = vmull_s8(vget_low_s8(wv), vget_low_s8(av));
    int16x8_t hi = vmull_s8(vget_high_s8(wv), vget_high_s8(av));
    int32x4_t acc = vpadalq_s16(vdupq_n_s32(0), lo);
    acc = vpadalq_s16(acc, hi);
    return vaddvq_s32(acc);
}

#ifdef __ARM_FEATURE_DOTPROD
// Same 16-dot via one SDOT (needs the ARMv8.2+dotprod encoding).
[[maybe_unused]] static inline std::int32_t dot16_vec_sdot(int8x16_t wv, int8x16_t av) {
    int32x4_t acc = vdupq_n_s32(0);
    acc = vdotq_s32(acc, wv, av);
    return vaddvq_s32(acc);
}
#endif

// Decode 32 packed split-half int4 nibbles (16 bytes) to sequential int8
// halves. Low nibbles are already w[0..15], high nibbles w[16..31] (both
// linear), so no vzip re-interleave is needed: s0/s1 dot directly against
// linear int8 acts. int4-only; fp4 keeps adjacent packing + nibbles_seq32.
// Sub-8 (not the LUT): int4 codes are offset-binary (value = code-8) while
// the legacy LUT is two's-complement — decoding via LUT is silently wrong.
[[maybe_unused]] static inline void nibbles_seq32_split(const std::uint8_t* packed,
                                                        int8x16_t& s0, int8x16_t& s1) {
    s0 = nibbles_sub8(packed, 1);
    s1 = nibbles_sub8(packed, 0);
}

// Decode 32 packed ADJACENT nibbles to sequential int8 halves (fp4 path:
// low nibbles are evens, high nibbles odds; vzip re-interleaves to
// s0 = elements 0..15, s1 = elements 16..31).
[[maybe_unused]] static inline void nibbles_seq32(const std::uint8_t* packed, const int8x16_t& lut,
                                                  int8x16_t& s0, int8x16_t& s1) {
    int8x16_t lo = nibbles_to_s8(packed, 1, lut);
    int8x16_t hi = nibbles_to_s8(packed, 0, lut);
    s0 = vzip1q_s8(lo, hi);
    s1 = vzip2q_s8(lo, hi);
}

// One 32-weight int4 block -> int32 dot against 32 int8 acts.
[[maybe_unused]] static inline std::int32_t int4_block_dot_wide(const std::uint8_t* w, const std::int8_t* a) {
    int8x16_t s0, s1;
    nibbles_seq32_split(w, s0, s1);
    return dot16_vec_widening(s0, vld1q_s8(a)) + dot16_vec_widening(s1, vld1q_s8(a + 16));
}

#ifdef __ARM_FEATURE_DOTPROD
[[maybe_unused]] static inline std::int32_t int4_block_dot_sdot(const std::uint8_t* w, const std::int8_t* a) {
    int8x16_t s0, s1;
    nibbles_seq32_split(w, s0, s1);
    return dot16_vec_sdot(s0, vld1q_s8(a)) + dot16_vec_sdot(s1, vld1q_s8(a + 16));
}
#endif

[[maybe_unused]] static float dot_int4_q8_widening_impl(const std::uint8_t* weights, const std::int8_t* act,
                                                        const float* wscales, const float* act_scales,
                                                        std::size_t n) {
    // Mirrors dot_int4_q8_avx2: 128-wide (four 32-blocks, 4 float chains) +
    // 32-wide + scalar tail to the logical end. Per-32 weight scale times
    // per-32 act scale in the block epilogue. Sub-8 decode needs no LUT.
    std::size_t i = 0;
    for (; i + 128 <= n; i += 128, weights += 64, act += 128) {
        __builtin_prefetch(weights + 128, 0, 3);
        __builtin_prefetch(act + 256, 0, 3);
        acc0 += static_cast<float>(int4_block_dot_wide(weights, act)) *
                wscales[i / 32] * act_scales[i / 32];
        acc1 += static_cast<float>(int4_block_dot_wide(weights + 16, act + 32)) *
                wscales[i / 32 + 1] * act_scales[i / 32 + 1];
        acc2 += static_cast<float>(int4_block_dot_wide(weights + 32, act + 64)) *
                wscales[i / 32 + 2] * act_scales[i / 32 + 2];
        acc3 += static_cast<float>(int4_block_dot_wide(weights + 48, act + 96)) *
                wscales[i / 32 + 3] * act_scales[i / 32 + 3];
    }
    float result = (acc0 + acc1) + (acc2 + acc3);
    for (; i + 32 <= n; i += 32, weights += 16, act += 32)
        result += static_cast<float>(int4_block_dot_wide(weights, act)) *
                  wscales[i / 32] * act_scales[i / 32];
    if (i < n) {
        // Partial tail block, read only to the row's logical end (i % 32 == 0
        // here, so scales[i/32] is the exact next per-32 scale and (i%2)==0
        // keeps nibble parity aligned with the scalar reference).
        const float scale = wscales[i / 32] * act_scales[i / 32];
        float block_sum = 0;
        for (std::size_t j = 0; i < n; ++i, ++j) {
            const int nibble = j < 16 ? (weights[j] & 15) : ((weights[j - 16] >> 4) & 15);
            block_sum += static_cast<float>(nibble - 8) * static_cast<float>(act[j]);
        }
        result += block_sum * scale;
    }
    return result;
}

#ifdef __ARM_FEATURE_DOTPROD
[[maybe_unused]] static float dot_int4_q8_sdot_impl(const std::uint8_t* weights, const std::int8_t* act,
                                                    const float* wscales, const float* act_scales,
                                                    std::size_t n) {
    // Identical blocking/order to the widening impl above (only the 32-block
    // dot primitive differs), so the two are bitwise identical end to end.
    // Sub-8 decode needs no LUT.
    float acc0 = 0, acc1 = 0, acc2 = 0, acc3 = 0;
    std::size_t i = 0;
    for (; i + 128 <= n; i += 128, weights += 64, act += 128) {
        __builtin_prefetch(weights + 128, 0, 3);
        __builtin_prefetch(act + 256, 0, 3);
        acc0 += static_cast<float>(int4_block_dot_sdot(weights, act)) *
                wscales[i / 32] * act_scales[i / 32];
        acc1 += static_cast<float>(int4_block_dot_sdot(weights + 16, act + 32)) *
                wscales[i / 32 + 1] * act_scales[i / 32 + 1];
        acc2 += static_cast<float>(int4_block_dot_sdot(weights + 32, act + 64)) *
                wscales[i / 32 + 2] * act_scales[i / 32 + 2];
        acc3 += static_cast<float>(int4_block_dot_sdot(weights + 48, act + 96)) *
                wscales[i / 32 + 3] * act_scales[i / 32 + 3];
    }
    float result = (acc0 + acc1) + (acc2 + acc3);
    for (; i + 32 <= n; i += 32, weights += 16, act += 32)
        result += static_cast<float>(int4_block_dot_sdot(weights, act)) *
                  wscales[i / 32] * act_scales[i / 32];
    if (i < n) {
        const float scale = wscales[i / 32] * act_scales[i / 32];
        float block_sum = 0;
        for (std::size_t j = 0; i < n; ++i, ++j) {
            const int nibble = j < 16 ? (weights[j] & 15) : ((weights[j - 16] >> 4) & 15);
            block_sum += static_cast<float>(nibble - 8) * static_cast<float>(act[j]);
        }
        result += block_sum * scale;
    }
    return result;
}
#endif

// TU-internal dispatch: SDOT when the CPU reports ASIMDDP, else widening.
// Same wiring note as dot_int8_q8_neon_internal (OUT OF SCOPE here): expose
// via a future kernels.cpp NEON Q8 selector once the runtime quantizes shared
// activations to int8 per-32; internal-only until then.
[[maybe_unused]] static float dot_int4_q8_neon_internal(const std::uint8_t* weights, const std::int8_t* act,
                                                        const float* wscales, const float* act_scales,
                                                        std::size_t n) {
    if (neon_detail::has_dotprod_runtime()) {
#ifdef __ARM_FEATURE_DOTPROD
        return dot_int4_q8_sdot_impl(weights, act, wscales, act_scales, n);
#else
        // CPU has SDOT but this TU was built baseline (no +dotprod encoding).
#endif
    }
    return dot_int4_q8_widening_impl(weights, act, wscales, act_scales, n);
}

[[maybe_unused]] static float dot_fp4_q8_widening_impl(const std::uint8_t* weights, const std::int8_t* act,
                                                       const std::uint8_t* scales, const float* act_scales,
                                                       std::size_t n) {
    // Mirrors dot_fp4_q8_avx2: 64-wide (two 32-blocks = four 16-weight groups,
    // 4 float chains) + 32-wide + scalar tail to the logical end. Per-16 E4M3
    // weight scale (halved LUT entry; the decoded element is the half value)
    // times the per-32 act scale in the block epilogue.
    const int8x16_t lut = vld1q_s8(fp4_element_lut());
    const float* scale_lut = fp4_scale_lut();
    float acc0 = 0, acc1 = 0, acc2 = 0, acc3 = 0;
    std::size_t i = 0;
    for (; i + 64 <= n; i += 64, weights += 32, act += 64) {
        __builtin_prefetch(weights + 64, 0, 3);
        __builtin_prefetch(act + 128, 0, 3);
        {
            int8x16_t s0, s1;
            nibbles_seq32(weights, lut, s0, s1);
            const float ascale = act_scales[i / 32];
            acc0 += static_cast<float>(dot16_vec_widening(s0, vld1q_s8(act))) *
                    scale_lut[scales[i / 16]] * ascale;
            acc1 += static_cast<float>(dot16_vec_widening(s1, vld1q_s8(act + 16))) *
                    scale_lut[scales[i / 16 + 1]] * ascale;
        }
        {
            int8x16_t s0, s1;
            nibbles_seq32(weights + 16, lut, s0, s1);
            const float ascale = act_scales[i / 32 + 1];
            acc2 += static_cast<float>(dot16_vec_widening(s0, vld1q_s8(act + 32))) *
                    scale_lut[scales[i / 16 + 2]] * ascale;
            acc3 += static_cast<float>(dot16_vec_widening(s1, vld1q_s8(act + 48))) *
                    scale_lut[scales[i / 16 + 3]] * ascale;
        }
    }
    float result = (acc0 + acc1) + (acc2 + acc3);
    for (; i + 32 <= n; i += 32, weights += 16, act += 32) {
        int8x16_t s0, s1;
        nibbles_seq32(weights, lut, s0, s1);
        const float ascale = act_scales[i / 32];
        result += static_cast<float>(dot16_vec_widening(s0, vld1q_s8(act))) *
                  scale_lut[scales[i / 16]] * ascale;
        result += static_cast<float>(dot16_vec_widening(s1, vld1q_s8(act + 16))) *
                  scale_lut[scales[i / 16 + 1]] * ascale;
    }
    if (i < n) {
        // Partial tail block, read only to the logical end (i % 32 == 0 here,
        // so i % 16 == 0 and scales[i/16] is the exact next per-16 scale).
        const auto* elements = fp4_element_lut();
        float block_sum = 0;
        for (std::size_t j = 0; i < n; ++i, ++j)
            block_sum += static_cast<float>(elements[(weights[j / 2] >> (4 * (i % 2))) & 15]) *
                         scale_lut[scales[i / 16]] * static_cast<float>(act[j]);
        result += block_sum * act_scales[i / 32];
    }
    return result;
}

#ifdef __ARM_FEATURE_DOTPROD
[[maybe_unused]] static float dot_fp4_q8_sdot_impl(const std::uint8_t* weights, const std::int8_t* act,
                                                   const std::uint8_t* scales, const float* act_scales,
                                                   std::size_t n) {
    // Identical blocking/order to the widening impl above (only the 16-dot
    // primitive differs), so the two are bitwise identical end to end.
    const int8x16_t lut = vld1q_s8(fp4_element_lut());
    const float* scale_lut = fp4_scale_lut();
    float acc0 = 0, acc1 = 0, acc2 = 0, acc3 = 0;
    std::size_t i = 0;
    for (; i + 64 <= n; i += 64, weights += 32, act += 64) {
        __builtin_prefetch(weights + 64, 0, 3);
        __builtin_prefetch(act + 128, 0, 3);
        {
            int8x16_t s0, s1;
            nibbles_seq32(weights, lut, s0, s1);
            const float ascale = act_scales[i / 32];
            acc0 += static_cast<float>(dot16_vec_sdot(s0, vld1q_s8(act))) *
                    scale_lut[scales[i / 16]] * ascale;
            acc1 += static_cast<float>(dot16_vec_sdot(s1, vld1q_s8(act + 16))) *
                    scale_lut[scales[i / 16 + 1]] * ascale;
        }
        {
            int8x16_t s0, s1;
            nibbles_seq32(weights + 16, lut, s0, s1);
            const float ascale = act_scales[i / 32 + 1];
            acc2 += static_cast<float>(dot16_vec_sdot(s0, vld1q_s8(act + 32))) *
                    scale_lut[scales[i / 16 + 2]] * ascale;
            acc3 += static_cast<float>(dot16_vec_sdot(s1, vld1q_s8(act + 48))) *
                    scale_lut[scales[i / 16 + 3]] * ascale;
        }
    }
    float result = (acc0 + acc1) + (acc2 + acc3);
    for (; i + 32 <= n; i += 32, weights += 16, act += 32) {
        int8x16_t s0, s1;
        nibbles_seq32(weights, lut, s0, s1);
        const float ascale = act_scales[i / 32];
        result += static_cast<float>(dot16_vec_sdot(s0, vld1q_s8(act))) *
                  scale_lut[scales[i / 16]] * ascale;
        result += static_cast<float>(dot16_vec_sdot(s1, vld1q_s8(act + 16))) *
                  scale_lut[scales[i / 16 + 1]] * ascale;
    }
    if (i < n) {
        const auto* elements = fp4_element_lut();
        float block_sum = 0;
        for (std::size_t j = 0; i < n; ++i, ++j)
            block_sum += static_cast<float>(elements[(weights[j / 2] >> (4 * (i % 2))) & 15]) *
                         scale_lut[scales[i / 16]] * static_cast<float>(act[j]);
        result += block_sum * act_scales[i / 32];
    }
    return result;
}
#endif

// TU-internal dispatch: SDOT when the CPU reports ASIMDDP, else widening.
// Same wiring note as dot_int4_q8_neon_internal - internal-only until a
// future kernels.cpp NEON Q8 selector exists.
[[maybe_unused]] static float dot_fp4_q8_neon_internal(const std::uint8_t* weights, const std::int8_t* act,
                                                       const std::uint8_t* scales, const float* act_scales,
                                                       std::size_t n) {
    if (neon_detail::has_dotprod_runtime()) {
#ifdef __ARM_FEATURE_DOTPROD
        return dot_fp4_q8_sdot_impl(weights, act, scales, act_scales, n);
#else
        // CPU has SDOT but this TU was built baseline (no +dotprod encoding).
#endif
    }
    return dot_fp4_q8_widening_impl(weights, act, scales, act_scales, n);
}

// ---- i8mm (ARMv8.6 usmmla): NOT on Snapdragon 888, future only -------------
// The 888 (X1/A78) predates ARMv8.6 maternal-multiply; there is no HWCAP2_I8MM
// on it, so any i8mm path must stay unselected there. Kept as a commented
// stub (no compiled code) so the baseline build cannot break:
//
//   #ifdef __ARM_FEATURE_MATMUL_INT8
//   // 2x2x8 int8 outer-product per instruction; sketch for a future Q8 block:
//   //   int32x4_t acc = vdupq_n_s32(0);
//   //   acc = vusmmala_s32(acc, vld1q_u8(w8mm), vld1q_s8(a8mm));  // 8x8 tile
//   //   ... one MAC per 8-element row pair, then the same per-32 float scale
//   //   epilogue as the SDOT impl above ...
//   //   runtime gate would be getauxval(AT_HWCAP2) & HWCAP2_I8MM (Linux) and
//   //   is NEVER true on the S21 FE - do not select without that HWCAP check.
//   #endif
//
// If i8mm is ever wired, it belongs beside dot_int8_q8_neon_internal with the
// same internal-dispatch shape (HWCAP2 probe -> i8mm, else SDOT, else widen).

// Decode 16 packed nibbles through the int8 LUT (vtbl1 = pshufb equivalent).
inline int8x16_t nibbles_to_s8(const std::uint8_t* packed, int low, const int8x16_t& lut) {
    uint8x16_t p = vld1q_u8(packed);
    uint8x16_t n = low ? vandq_u8(p, vdupq_n_u8(15)) : vshrq_n_u8(p, 4);
    return vqtbl1q_s8(vreinterpretq_s8_s8(lut), n);
}

// Decode 16 split-half int4 nibbles WITHOUT a LUT: int4 codes are
// offset-binary (value = code-8), which the legacy two's-complement LUT
// does not represent. Plain integer sub, exact.
inline int8x16_t nibbles_sub8(const std::uint8_t* packed, int low) {
    uint8x16_t p = vld1q_u8(packed);
    uint8x16_t n = low ? vandq_u8(p, vdupq_n_u8(15)) : vshrq_n_u8(p, 4);
    return vsubq_s8(vreinterpretq_s8_u8(n), vdupq_n_s8(8));
}

inline float32x4_t s8x16_dot_f32(const int8x16_t& w, const float* act, float scale) {
    int16x8_t w0 = vmovl_s8(vget_low_s8(w));
    int16x8_t w1 = vmovl_s8(vget_high_s8(w));
    float32x4_t acc = vmulq_f32(vcvtq_f32_s32(vmovl_s16(vget_low_s16(w0))), vld1q_f32(act));
    acc = vfmaq_f32(acc, vcvtq_f32_s32(vmovl_s16(vget_high_s16(w0))), vld1q_f32(act + 4));
    acc = vfmaq_f32(acc, vcvtq_f32_s32(vmovl_s16(vget_low_s16(w1))), vld1q_f32(act + 8));
    acc = vfmaq_f32(acc, vcvtq_f32_s32(vmovl_s16(vget_high_s16(w1))), vld1q_f32(act + 12));
    return vmulq_n_f32(acc, scale);
}

// Dot 8 decoded int8s against 8 floats with one scale (one fp4 quarter: the
// per-16 scale boundary cuts across the even/odd split, so a 32-block needs
// four of these, not two s8x16 halves - see dot_fp4_neon).
inline float32x4_t s8x8_dot_f32(int8x8_t w8, const float* act, float scale) {
    int16x8_t w16 = vmovl_s8(w8);
    float32x4_t lo = vcvtq_f32_s32(vmovl_s16(vget_low_s16(w16)));
    float32x4_t hi = vcvtq_f32_s32(vmovl_s16(vget_high_s16(w16)));
    float32x4_t acc = vmulq_f32(lo, vld1q_f32(act));
    acc = vfmaq_f32(acc, hi, vld1q_f32(act + 4));
    return vmulq_n_f32(acc, scale);
}

float dot_int4_neon(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 layout dots against linear acts; no permute.
    // Sub-8 decode (offset-binary codes; the legacy two's-complement LUT
    // does not represent them).
    float32x4_t sum0 = vdupq_n_f32(0), sum1 = vdupq_n_f32(0);
    float32x4_t sum2 = vdupq_n_f32(0), sum3 = vdupq_n_f32(0);
    std::size_t i = 0;
    // Split-half nibbles: byte j holds w[j] (low) and w[j+16] (high), so low
    // nibbles ARE w[0..15] and high nibbles ARE w[16..31] (both linear) dotted
    // against linear acts. int4 scales are per-32, so one scale covers both
    // halves of a 32-block (unlike fp4's per-16 split).
    for (; i + 64 <= n; i += 64, weights += 32, act_orig += 64) {
        __builtin_prefetch(weights + 32, 0, 3);
        __builtin_prefetch(act_orig + 64, 0, 3);
        sum0 = vaddq_f32(sum0, s8x16_dot_f32(nibbles_sub8(weights, 1), act_orig, scales[i / 32]));
        sum1 = vaddq_f32(sum1, s8x16_dot_f32(nibbles_sub8(weights, 0), act_orig + 16, scales[i / 32]));
        sum2 = vaddq_f32(sum2, s8x16_dot_f32(nibbles_sub8(weights + 16, 1), act_orig + 32, scales[i / 32 + 1]));
        sum3 = vaddq_f32(sum3, s8x16_dot_f32(nibbles_sub8(weights + 16, 0), act_orig + 48, scales[i / 32 + 1]));
    }
    for (; i + 32 <= n; i += 32, weights += 16, act_orig += 32) {
        sum0 = vaddq_f32(sum0, s8x16_dot_f32(nibbles_sub8(weights, 1), act_orig, scales[i / 32]));
        sum1 = vaddq_f32(sum1, s8x16_dot_f32(nibbles_sub8(weights, 0), act_orig + 16, scales[i / 32]));
    }
    float32x4_t sum = vaddq_f32(vaddq_f32(sum0, sum1), vaddq_f32(sum2, sum3));
    float result = vaddvq_f32(sum);
    // Partial tail block, read only to its logical end. i is a multiple of 32;
    // weights points at the current 16-byte split-half block, act_orig at
    // element i, scales+i/32 at the next per-32 scale; dot_int4_scalar handles
    // split-half indexing against linear acts.
    if (i < n) result += dot_int4_scalar(weights, act_orig, scales + i / 32, n - i);
    return result;
}

float dot_fp4_neon(const std::uint8_t* weights, const float* act_perm, const float* act_orig,
                   const std::uint8_t* scales, std::size_t n) {
    const int8x16_t lut = vld1q_s8(fp4_element_lut());
    const float* scale_lut = fp4_scale_lut();
    float32x4_t sum0 = vdupq_n_f32(0), sum1 = vdupq_n_f32(0);
    float32x4_t sum2 = vdupq_n_f32(0), sum3 = vdupq_n_f32(0);
    std::size_t i = 0;
    // E2M1 elements via vtbl LUT (half values, exact in int8), widen to FP32,
    // scale with the halved E4M3 entry. Layout matches AVX2: PERM[0..16) holds
    // evens, PERM[16..32) holds odds, but the per-16 scale boundary CUTS ACROSS
    // the split: elements 0-15 (scaleA) are evens[0..8)+odds[0..8], elements
    // 16-31 (scaleB) are evens[8..16)+odds[8..16). Each 32-block is therefore
    // four scaled quarters into the four independent accumulators (same 4-acc
    // order as dot_fp4_avx2); a two-halves/one-scale-per-half form would apply
    // scaleA to evens 16-30 that belong to scaleB. Fixed to the 4-quarter form.
    for (; i + 64 <= n; i += 64, weights += 32, act_perm += 64) {
        __builtin_prefetch(weights + 32, 0, 3);
        __builtin_prefetch(act_perm + 64, 0, 3);
        {
            int8x16_t ev = nibbles_to_s8(weights, 1, lut);
            int8x16_t od = nibbles_to_s8(weights, 0, lut);
            const float sA = scale_lut[scales[i / 16]];
            const float sB = scale_lut[scales[i / 16 + 1]];
            sum0 = vaddq_f32(sum0, s8x8_dot_f32(vget_low_s8(ev), act_perm, sA));
            sum1 = vaddq_f32(sum1, s8x8_dot_f32(vget_high_s8(ev), act_perm + 8, sB));
            sum2 = vaddq_f32(sum2, s8x8_dot_f32(vget_low_s8(od), act_perm + 16, sA));
            sum3 = vaddq_f32(sum3, s8x8_dot_f32(vget_high_s8(od), act_perm + 24, sB));
        }
        {
            int8x16_t ev = nibbles_to_s8(weights + 16, 1, lut);
            int8x16_t od = nibbles_to_s8(weights + 16, 0, lut);
            const float sA = scale_lut[scales[i / 16 + 2]];
            const float sB = scale_lut[scales[i / 16 + 3]];
            sum0 = vaddq_f32(sum0, s8x8_dot_f32(vget_low_s8(ev), act_perm + 32, sA));
            sum1 = vaddq_f32(sum1, s8x8_dot_f32(vget_high_s8(ev), act_perm + 40, sB));
            sum2 = vaddq_f32(sum2, s8x8_dot_f32(vget_low_s8(od), act_perm + 48, sA));
            sum3 = vaddq_f32(sum3, s8x8_dot_f32(vget_high_s8(od), act_perm + 56, sB));
        }
    }
    for (; i + 32 <= n; i += 32, weights += 16, act_perm += 32) {
        int8x16_t ev = nibbles_to_s8(weights, 1, lut);
        int8x16_t od = nibbles_to_s8(weights, 0, lut);
        const float sA = scale_lut[scales[i / 16]];
        const float sB = scale_lut[scales[i / 16 + 1]];
        sum0 = vaddq_f32(sum0, s8x8_dot_f32(vget_low_s8(ev), act_perm, sA));
        sum1 = vaddq_f32(sum1, s8x8_dot_f32(vget_high_s8(ev), act_perm + 8, sB));
        sum2 = vaddq_f32(sum2, s8x8_dot_f32(vget_low_s8(od), act_perm + 16, sA));
        sum3 = vaddq_f32(sum3, s8x8_dot_f32(vget_high_s8(od), act_perm + 24, sB));
    }
    float32x4_t sum = vaddq_f32(vaddq_f32(sum0, sum1), vaddq_f32(sum2, sum3));
    float result = vaddvq_f32(sum);
    // Same tail invariant as int4 (i % 32 == 0, so i % 16 == 0): the scalar
    // tail with scales+i/16 handles per-16 blocks correctly to the logical end.
    if (i < n) result += dot_fp4_scalar(weights, act_orig + i, scales + i / 16, n - i);
    return result;
}

// One canonical 32-block deinterleave: OUT[0..16)=IN evens, OUT[16..32)=IN odds.
inline void permute_one32_neon(const float* in, float* out) {
    float32x4_t r0 = vld1q_f32(in), r1 = vld1q_f32(in + 4);
    float32x4_t r2 = vld1q_f32(in + 8), r3 = vld1q_f32(in + 12);
    float32x4_t r4 = vld1q_f32(in + 16), r5 = vld1q_f32(in + 20);
    float32x4_t r6 = vld1q_f32(in + 24), r7 = vld1q_f32(in + 28);
    vst1q_f32(out, vuzp1q_f32(r0, r1));
    vst1q_f32(out + 4, vuzp1q_f32(r2, r3));
    vst1q_f32(out + 8, vuzp1q_f32(r4, r5));
    vst1q_f32(out + 12, vuzp1q_f32(r6, r7));
    vst1q_f32(out + 16, vuzp2q_f32(r0, r1));
    vst1q_f32(out + 20, vuzp2q_f32(r2, r3));
    vst1q_f32(out + 24, vuzp2q_f32(r4, r5));
    vst1q_f32(out + 28, vuzp2q_f32(r6, r7));
}

// Canonical 32-block deinterleave with NEON uzp (same layout as the AVX2
// and scalar permute; tail untouched, callers never read PERM tails).
// Verified against permute_act32_blocked (kernels.cpp): per full [c,c+32),
// OUT[c+k]=IN[c+2k] and OUT[c+16+k]=IN[c+2k+1]. With r0=IN[0..3], r1=IN[4..7],
// vuzp1q(r0,r1)=[r0[0],r0[2],r1[0],r1[2]]=IN[0,2,4,6]->OUT[0..4) (evens) and
// vuzp2q(r0,r1)=IN[1,3,5,7]->OUT[16..20) (odds); r2/r3, r4/r5, r6/r7 extend
// the same pattern to OUT[4..16)/OUT[20..32). The uzp form is therefore the
// exact even/odd 32-block split the int4/fp4 dots expect, not a reversal.
// 64-wide unroll (two 32-blocks per iteration, in order) halves loop overhead
// with identical element order, mirroring permute_act32_blocked's 64-wide.
void permute_act32_neon(const float* input, std::size_t n, float* out) {
    std::size_t c = 0;
    for (; c + 64 <= n; c += 64) {
        permute_one32_neon(input + c, out + c);
        permute_one32_neon(input + c + 32, out + c + 32);
    }
    for (; c + 32 <= n; c += 32) permute_one32_neon(input + c, out + c);
}

}  // namespace cism

#else

// Non-ARM build: scalar wrappers so the symbols always link. Dispatch
// (kernels.cpp) never selects them without NEON hardware.
namespace cism {
float dot_neon(const float* a, const float* b, std::size_t n) { return dot_scalar(a, b, n); }
float dot_int8_neon(const std::int8_t* weights, const float* input, std::size_t n) {
    return dot_int8_scalar(weights, input, n);
}
float dot_int4_neon(const std::uint8_t* weights, const float* /*act_perm*/, const float* act_orig,
                    const float* scales, std::size_t n) {
    return dot_int4_scalar(weights, act_orig, scales, n);
}
float dot_fp4_neon(const std::uint8_t* weights, const float* /*act_perm*/, const float* act_orig,
                    const std::uint8_t* scales, std::size_t n) {
    return dot_fp4_scalar(weights, act_orig, scales, n);
}
void permute_act32_neon(const float* input, std::size_t n, float* out) {
    permute_act32_blocked(input, n, out);
}
}  // namespace cism

#endif