Spaces:
Sleeping
Sleeping
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
|