// 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 #include #if defined(CISM_HAVE_NEON) && defined(__aarch64__) #include #if defined(__linux__) #include // HWCAP_ASIMDDP lives in on glibc but 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() #include #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(*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(dot16_s8_widening(weights, act) + dot16_s8_widening(weights + 16, act + 16)) * act_scales[i / 32]; acc1 += static_cast(dot16_s8_widening(weights + 32, act + 32) + dot16_s8_widening(weights + 48, act + 48)) * act_scales[i / 32 + 1]; acc2 += static_cast(dot16_s8_widening(weights + 64, act + 64) + dot16_s8_widening(weights + 80, act + 80)) * act_scales[i / 32 + 2]; acc3 += static_cast(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(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(weights[j]) * static_cast(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(dot16_s8_sdot(weights, act) + dot16_s8_sdot(weights + 16, act + 16)) * act_scales[i / 32]; acc1 += static_cast(dot16_s8_sdot(weights + 32, act + 32) + dot16_s8_sdot(weights + 48, act + 48)) * act_scales[i / 32 + 1]; acc2 += static_cast(dot16_s8_sdot(weights + 64, act + 64) + dot16_s8_sdot(weights + 80, act + 80)) * act_scales[i / 32 + 2]; acc3 += static_cast(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(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(weights[j]) * static_cast(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(int4_block_dot_wide(weights, act)) * wscales[i / 32] * act_scales[i / 32]; acc1 += static_cast(int4_block_dot_wide(weights + 16, act + 32)) * wscales[i / 32 + 1] * act_scales[i / 32 + 1]; acc2 += static_cast(int4_block_dot_wide(weights + 32, act + 64)) * wscales[i / 32 + 2] * act_scales[i / 32 + 2]; acc3 += static_cast(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(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(nibble - 8) * static_cast(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(int4_block_dot_sdot(weights, act)) * wscales[i / 32] * act_scales[i / 32]; acc1 += static_cast(int4_block_dot_sdot(weights + 16, act + 32)) * wscales[i / 32 + 1] * act_scales[i / 32 + 1]; acc2 += static_cast(int4_block_dot_sdot(weights + 32, act + 64)) * wscales[i / 32 + 2] * act_scales[i / 32 + 2]; acc3 += static_cast(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(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(nibble - 8) * static_cast(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(dot16_vec_widening(s0, vld1q_s8(act))) * scale_lut[scales[i / 16]] * ascale; acc1 += static_cast(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(dot16_vec_widening(s0, vld1q_s8(act + 32))) * scale_lut[scales[i / 16 + 2]] * ascale; acc3 += static_cast(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(dot16_vec_widening(s0, vld1q_s8(act))) * scale_lut[scales[i / 16]] * ascale; result += static_cast(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(elements[(weights[j / 2] >> (4 * (i % 2))) & 15]) * scale_lut[scales[i / 16]] * static_cast(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(dot16_vec_sdot(s0, vld1q_s8(act))) * scale_lut[scales[i / 16]] * ascale; acc1 += static_cast(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(dot16_vec_sdot(s0, vld1q_s8(act + 32))) * scale_lut[scales[i / 16 + 2]] * ascale; acc3 += static_cast(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(dot16_vec_sdot(s0, vld1q_s8(act))) * scale_lut[scales[i / 16]] * ascale; result += static_cast(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(elements[(weights[j / 2] >> (4 * (i % 2))) & 15]) * scale_lut[scales[i / 16]] * static_cast(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