Spaces:
Sleeping
Sleeping
Download native/kernels_neon.cpp from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 44.3 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_neon.cpp
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/native/kernels_neon.cpp
-
curl -L -o kernels_neon.cpp https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/native/kernels_neon.cpp
44.3 kB
| // 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. | |
| // 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. | |
| 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. | |
| constexpr unsigned long kHwcapDotprod = HWCAP_ASIMDDP; | |
| constexpr unsigned long kHwcapDotprod = (1UL << 20); | |
| constexpr unsigned long kHwcapDotprod = 0UL; | |
| // 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() { | |
| (void)kHwcapDotprod; | |
| return true; // All Apple Silicon ships ARMv8.2+dotprod or later. | |
| return (getauxval(AT_HWCAP) & kHwcapDotprod) != 0UL; | |
| return true; // TU built with +dotprod: encoding is always legal. | |
| return false; | |
| } | |
| } // 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; | |
| } | |
| // 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; | |
| } | |
| // 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()) { | |
| return dot_int8_q8_sdot_impl(weights, act, act_scales, n); | |
| // 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. | |
| } | |
| 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); | |
| } | |
| // 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); | |
| } | |
| // 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)); | |
| } | |
| [[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)); | |
| } | |
| [[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; | |
| } | |
| [[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; | |
| } | |
| // 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()) { | |
| return dot_int4_q8_sdot_impl(weights, act, wscales, act_scales, n); | |
| // CPU has SDOT but this TU was built baseline (no +dotprod encoding). | |
| } | |
| 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; | |
| } | |
| [[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; | |
| } | |
| // 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()) { | |
| return dot_fp4_q8_sdot_impl(weights, act, scales, act_scales, n); | |
| // CPU has SDOT but this TU was built baseline (no +dotprod encoding). | |
| } | |
| 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 | |
| // 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 | |