Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 0 additions & 4 deletions ggml/src/ggml-cpu/arch-fallback.h
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
153 changes: 153 additions & 0 deletions ggml/src/ggml-cpu/arch/arm/quants.c
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down
4 changes: 4 additions & 0 deletions ggml/src/ggml-cpu/ggml-cpu.c
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down