From 9098196d1179bc7b0e8eafee6a1caf7c16645995 Mon Sep 17 00:00:00 2001 From: Ruslan Karymov <67752306+karusrus@users.noreply.github.com> Date: Sun, 27 Sep 2026 00:04:47 +0200 Subject: [PATCH] arm: NEON + i8mm vec_dot for PTQ1_0 PTQ1_0 had only the generic vec_dot on ARM. This adds a NEON kernel: - trits are decoded in 8-bit lanes: ((q + (q >> 1)) >> 1) >> 6 equals (3q) >> 8 for every byte value, so no widening is needed; a 128-trit block becomes eight int8x16 vectors in element order - each 32-wide Q8_0 sub-block is reduced with sdot and accumulated into a float32x4 with vmlaq_n_f32, one horizontal add per row - with __ARM_FEATURE_MATMUL_INT8, PTQ1_0 uses nrows = 2 and the nrc == 2 path computes a 2x2 tile with vmmlaq_s32, reusing decoded trits for two activation columns Snapdragon 7 Gen 4, Ternary Bonsai 2 27B, -t 5: pp64 1.22 -> 2.33 t/s, tg16 1.09 -> 1.31 t/s against the generic path. test-quantize-fns passes; greedy output matches the generic path on the same device. Co-Authored-By: Claude Opus 5.5 --- ggml/src/ggml-cpu/arch-fallback.h | 4 - ggml/src/ggml-cpu/arch/arm/quants.c | 153 ++++++++++++++++++++++++++++ ggml/src/ggml-cpu/ggml-cpu.c | 4 + 3 files changed, 157 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-cpu/arch-fallback.h b/ggml/src/ggml-cpu/arch-fallback.h index f22018a5a76f..1da64c98f697 100644 --- a/ggml/src/ggml-cpu/arch-fallback.h +++ b/ggml/src/ggml-cpu/arch-fallback.h @@ -82,10 +82,6 @@ #define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0 #define ggml_gemm_pq2_0_4x8_q8_0_generic ggml_gemm_pq2_0_4x8_q8_0 #elif defined(__aarch64__) || defined(__arm__) || defined(_M_ARM) || defined(_M_ARM64) -// PTQ1_0's NEON vec_dot is unverified on real ARM hardware and measured slower -// than the generic loop on two independent devices (Apple M5 Pro, Snapdragon 7 -// Gen 4); alias it back to generic until a validated replacement lands -#define ggml_vec_dot_ptq1_0_q8_0_generic ggml_vec_dot_ptq1_0_q8_0 // repack.cpp #define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4 #define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8 diff --git a/ggml/src/ggml-cpu/arch/arm/quants.c b/ggml/src/ggml-cpu/arch/arm/quants.c index bea782fdb57e..490a9a444ae1 100644 --- a/ggml/src/ggml-cpu/arch/arm/quants.c +++ b/ggml/src/ggml-cpu/arch/arm/quants.c @@ -1527,6 +1527,159 @@ void ggml_vec_dot_q8_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const voi *s = sumf; } +#if defined(__ARM_NEON) +// Decode one PTQ1_0 block (128 trits) into eight int8x16 vectors in element order, +// values in {-1, 0, 1}. Trit order follows ggml_vec_dot_ptq1_0_q8_0_generic: +// qs[0..15] -> values 0..79 : trit 0..4 of each byte, 16 values per trit +// qs[16..23] -> values 80..119: trit 0..4 of each byte, 8 values per trit +// qh[0..1] -> values 120..127: trit 0..3 of each byte, 2 values per trit +// A trit is ((uint8)(q * 3^n) * 3) >> 8, computed exactly as ((q + (q >> 1)) >> 1) >> 6 +// (checked for all 256 byte values). +static inline void ggml_ptq1_0_decode_neon(const block_ptq1_0 * GGML_RESTRICT b, int8x16_t e[8]) { + static const uint8_t k_qh_pow3[8] = {1, 1, 3, 3, 9, 9, 27, 27}; + const int8x16_t one = vdupq_n_s8(1); + +#define PTQ1_TRIT16(v) vsubq_s8(vreinterpretq_s8_u8(vshrq_n_u8(vhaddq_u8((v), vshrq_n_u8((v), 1)), 6)), one) +#define PTQ1_TRIT8(v) vshr_n_u8(vhadd_u8((v), vshr_n_u8((v), 1)), 6) + + const uint8x16_t a = vld1q_u8(b->qs); + e[0] = PTQ1_TRIT16(a); + e[1] = PTQ1_TRIT16(vmulq_u8(a, vdupq_n_u8(3))); + e[2] = PTQ1_TRIT16(vmulq_u8(a, vdupq_n_u8(9))); + e[3] = PTQ1_TRIT16(vmulq_u8(a, vdupq_n_u8(27))); + e[4] = PTQ1_TRIT16(vmulq_u8(a, vdupq_n_u8(81))); + + const uint8x8_t c = vld1_u8(b->qs + 16); + const uint8x8_t u0 = PTQ1_TRIT8(c); + const uint8x8_t u1 = PTQ1_TRIT8(vmul_u8(c, vdup_n_u8(3))); + const uint8x8_t u2 = PTQ1_TRIT8(vmul_u8(c, vdup_n_u8(9))); + const uint8x8_t u3 = PTQ1_TRIT8(vmul_u8(c, vdup_n_u8(27))); + const uint8x8_t u4 = PTQ1_TRIT8(vmul_u8(c, vdup_n_u8(81))); + + // qh[0..1] -> {qh0, qh1} x {1, 3, 9, 27} + const uint16_t qh16 = (uint16_t) b->qh[0] | ((uint16_t) b->qh[1] << 8); + const uint8x8_t h = vreinterpret_u8_u16(vdup_n_u16(qh16)); + const uint8x8_t uh = PTQ1_TRIT8(vmul_u8(h, vld1_u8(k_qh_pow3))); + + e[5] = vsubq_s8(vreinterpretq_s8_u8(vcombine_u8(u0, u1)), one); // values 80..95 + e[6] = vsubq_s8(vreinterpretq_s8_u8(vcombine_u8(u2, u3)), one); // values 96..111 + e[7] = vsubq_s8(vreinterpretq_s8_u8(vcombine_u8(u4, uh)), one); // values 112..127 + +#undef PTQ1_TRIT16 +#undef PTQ1_TRIT8 +} +#endif + +// PTQ1_0 x Q8_0: one 128-wide PTQ1_0 block meets four Q8_0 blocks. +// With i8mm, nrc == 2 computes a 2x2 tile (two weight rows x two activation columns) +// with SMMLA, so every decoded trit vector is used twice. +void ggml_vec_dot_ptq1_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { + assert(n % QK_PTQ1_0 == 0); +#if defined(__ARM_FEATURE_MATMUL_INT8) + assert((nrc == 2) || (nrc == 1)); +#else + assert(nrc == 1); +#endif + UNUSED(nrc); + UNUSED(bx); + UNUSED(by); + UNUSED(bs); + +#if defined(__ARM_NEON) + const int nb = n / QK_PTQ1_0; + +#if defined(__ARM_FEATURE_MATMUL_INT8) + if (nrc == 2) { + const block_ptq1_0 * GGML_RESTRICT x0 = vx; + const block_ptq1_0 * GGML_RESTRICT x1 = (const block_ptq1_0 *) ((const uint8_t *) vx + bx); + const block_q8_0 * GGML_RESTRICT y0 = vy; + const block_q8_0 * GGML_RESTRICT y1 = (const block_q8_0 *) ((const uint8_t *) vy + by); + + float32x4_t sumv = vdupq_n_f32(0.0f); + + for (int i = 0; i < nb; ++i) { + int8x16_t e0[8]; + int8x16_t e1[8]; + ggml_ptq1_0_decode_neon(&x0[i], e0); + ggml_ptq1_0_decode_neon(&x1[i], e1); + + const float dx0 = GGML_CPU_FP16_TO_FP32(x0[i].d); + const float dx1 = GGML_CPU_FP16_TO_FP32(x1[i].d); + + for (int k = 0; k < 4; ++k) { + const block_q8_0 * GGML_RESTRICT b0 = &y0[4*i + k]; + const block_q8_0 * GGML_RESTRICT b1 = &y1[4*i + k]; + + const int8x16_t y0_l = vld1q_s8(b0->qs); + const int8x16_t y0_h = vld1q_s8(b0->qs + 16); + const int8x16_t y1_l = vld1q_s8(b1->qs); + const int8x16_t y1_h = vld1q_s8(b1->qs + 16); + + const int8x16_t l0 = vreinterpretq_s8_s64(vzip1q_s64(vreinterpretq_s64_s8(e0[2*k]), vreinterpretq_s64_s8(e1[2*k]))); + const int8x16_t l1 = vreinterpretq_s8_s64(vzip2q_s64(vreinterpretq_s64_s8(e0[2*k]), vreinterpretq_s64_s8(e1[2*k]))); + const int8x16_t l2 = vreinterpretq_s8_s64(vzip1q_s64(vreinterpretq_s64_s8(e0[2*k + 1]), vreinterpretq_s64_s8(e1[2*k + 1]))); + const int8x16_t l3 = vreinterpretq_s8_s64(vzip2q_s64(vreinterpretq_s64_s8(e0[2*k + 1]), vreinterpretq_s64_s8(e1[2*k + 1]))); + + const int8x16_t r0 = vreinterpretq_s8_s64(vzip1q_s64(vreinterpretq_s64_s8(y0_l), vreinterpretq_s64_s8(y1_l))); + const int8x16_t r1 = vreinterpretq_s8_s64(vzip2q_s64(vreinterpretq_s64_s8(y0_l), vreinterpretq_s64_s8(y1_l))); + const int8x16_t r2 = vreinterpretq_s8_s64(vzip1q_s64(vreinterpretq_s64_s8(y0_h), vreinterpretq_s64_s8(y1_h))); + const int8x16_t r3 = vreinterpretq_s8_s64(vzip2q_s64(vreinterpretq_s64_s8(y0_h), vreinterpretq_s64_s8(y1_h))); + + const float dy0 = GGML_CPU_FP16_TO_FP32(b0->d); + const float dy1 = GGML_CPU_FP16_TO_FP32(b1->d); + const float32_t scale_arr[4] = { dx0*dy0, dx0*dy1, dx1*dy0, dx1*dy1 }; + + int32x4_t acc = vmmlaq_s32(vdupq_n_s32(0), l0, r0); + acc = vmmlaq_s32(acc, l1, r1); + acc = vmmlaq_s32(acc, l2, r2); + acc = vmmlaq_s32(acc, l3, r3); + + sumv = vmlaq_f32(sumv, vcvtq_f32_s32(acc), vld1q_f32(scale_arr)); + } + } + + // sumv = {x0.y0, x0.y1, x1.y0, x1.y1} -> s[0..1] = column 0, s[bs..bs+1] = column 1 + const float32x4_t sumv1 = vextq_f32(sumv, sumv, 2); + const float32x4_t sumv2 = vzip1q_f32(sumv, sumv1); + + vst1_f32(s, vget_low_f32 (sumv2)); + vst1_f32(s + bs, vget_high_f32(sumv2)); + return; + } +#endif + + const block_ptq1_0 * GGML_RESTRICT x = vx; + const block_q8_0 * GGML_RESTRICT y = vy; + + float32x4_t sumv = vdupq_n_f32(0.0f); + + for (int i = 0; i < nb; ++i) { + int8x16_t e[8]; + ggml_ptq1_0_decode_neon(&x[i], e); + + const float dx = GGML_CPU_FP16_TO_FP32(x[i].d); + + for (int k = 0; k < 4; ++k) { + const block_q8_0 * GGML_RESTRICT yb = &y[4*i + k]; +#if defined(__ARM_FEATURE_DOTPROD) + const int32x4_t p = vdotq_s32(vdotq_s32(vdupq_n_s32(0), e[2*k], vld1q_s8(yb->qs)), e[2*k + 1], vld1q_s8(yb->qs + 16)); +#else + const int8x16_t ya = vld1q_s8(yb->qs); + const int8x16_t yc = vld1q_s8(yb->qs + 16); + const int32x4_t p = vaddq_s32( + vaddq_s32(vpaddlq_s16(vmull_s8(vget_low_s8(e[2*k]), vget_low_s8(ya))), vpaddlq_s16(vmull_s8(vget_high_s8(e[2*k]), vget_high_s8(ya)))), + vaddq_s32(vpaddlq_s16(vmull_s8(vget_low_s8(e[2*k + 1]), vget_low_s8(yc))), vpaddlq_s16(vmull_s8(vget_high_s8(e[2*k + 1]), vget_high_s8(yc))))); +#endif + sumv = vmlaq_n_f32(sumv, vcvtq_f32_s32(p), dx * GGML_CPU_FP16_TO_FP32(yb->d)); + } + } + + *s = vaddvq_f32(sumv); +#else + ggml_vec_dot_ptq1_0_q8_0_generic(n, s, bs, vx, bx, vy, by, nrc); +#endif +} + void ggml_vec_dot_tq1_0_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { assert(nrc == 1); UNUSED(nrc); diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index dd40c284531a..10d4d690c862 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -247,7 +247,11 @@ static const struct ggml_type_traits_cpu type_traits_cpu[GGML_TYPE_COUNT] = { .from_float = quantize_row_ptq1_0, .vec_dot = ggml_vec_dot_ptq1_0_q8_0, .vec_dot_type = GGML_TYPE_Q8_0, +#if defined (__ARM_FEATURE_MATMUL_INT8) + .nrows = 2, +#else .nrows = 1, +#endif }, [GGML_TYPE_Q4_0] = { .from_float = quantize_row_q4_0,