test1111111 / native /kernels_neon.cpp
spitfire4794's picture
Space: AVX512 kernels (zmm fp32/int8, parallel VNNI-Q8) + 2T cap + NEON int4 fix
40d0cd2
Raw History Blame Contribute Delete
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.
#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