From d4ca15de27a8502218832ee85caf0fdf4222db3b Mon Sep 17 00:00:00 2001 From: Jason Titus Date: Wed, 23 Sep 2026 20:08:36 -0700 Subject: [PATCH] metal : add PTQ1_0 mat-vec for two to four columns Assisted-by: OpenAI Codex --- ggml/src/ggml-metal/ggml-metal-device.cpp | 15 +++ ggml/src/ggml-metal/ggml-metal-device.h | 2 + ggml/src/ggml-metal/ggml-metal-ops.cpp | 2 +- ggml/src/ggml-metal/kernels/mul_mv.metal | 134 ++++++++++++++++++++++ tests/test-backend-ops.cpp | 9 +- 5 files changed, 160 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index e0856c91b3e1..3420ba85116c 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -882,6 +882,13 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_meta return res; } +bool ggml_metal_ptq1_multicol_enabled(const ggml_tensor * op) { + static const bool enabled = getenv("GGML_METAL_PTQ1_MULTICOL") && atoi(getenv("GGML_METAL_PTQ1_MULTICOL")) == 1; + return enabled && op->src[0]->type == GGML_TYPE_PTQ1_0 && op->src[1]->type == GGML_TYPE_F32 && + op->src[0]->ne[0] % ggml_blck_size(GGML_TYPE_PTQ1_0) == 0 && op->src[1]->nb[0] == sizeof(float) && + op->src[1]->ne[1] >= 2 && op->src[1]->ne[1] <= 4; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_metal_library_t lib, const ggml_tensor * op) { GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); GGML_TENSOR_LOCALS( int32_t, ne1, op->src[1], ne); @@ -899,6 +906,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta const ggml_type tsrc1 = op->src[1]->type; const char * suffix = ""; + char ptq1_suffix[16]; // use custom matrix x vector kernel switch (tsrc0) { @@ -950,6 +958,13 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta { nsg = N_SG_PTQ1_0; nr0 = N_R0_PTQ1_0; + if (ggml_metal_ptq1_multicol_enabled(op)) { + nr0 = 4; + nsg = 1; + nr1 = ne11; + snprintf(ptq1_suffix, sizeof(ptq1_suffix), "_mc_c%d", nr1); + suffix = ptq1_suffix; + } } break; case GGML_TYPE_Q4_0: { diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 76cb4b2646ff..0d305888f925 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -6,6 +6,8 @@ extern "C" { #endif +bool ggml_metal_ptq1_multicol_enabled(const struct ggml_tensor * op); + struct ggml_metal_buffer_id { void * metal; // id size_t offs; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index de5be248470a..c41bae469ac5 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -2800,7 +2800,7 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) { (op->src[0]->type == GGML_TYPE_Q1_0 && q1_0_ext_enable) || op->src[0]->type == GGML_TYPE_Q2_0 || op->src[0]->type == GGML_TYPE_PQ2_0 || - op->src[0]->type == GGML_TYPE_PTQ1_0 || + (op->src[0]->type == GGML_TYPE_PTQ1_0 && !ggml_metal_ptq1_multicol_enabled(op)) || op->src[0]->type == GGML_TYPE_Q4_0 || op->src[0]->type == GGML_TYPE_Q4_1 || op->src[0]->type == GGML_TYPE_Q5_0 || diff --git a/ggml/src/ggml-metal/kernels/mul_mv.metal b/ggml/src/ggml-metal/kernels/mul_mv.metal index de7dbaffea02..7d3bbc0bee89 100644 --- a/ggml/src/ggml-metal/kernels/mul_mv.metal +++ b/ggml/src/ggml-metal/kernels/mul_mv.metal @@ -1019,6 +1019,140 @@ void kernel_mul_mv_ptq1_0_f32_impl( } } +// Keep the single-vector arithmetic order while sharing decoded weights across columns. +template +inline void ptq1_0_dot_multicol(device const block_ptq1_0 * qb, + thread const float (&yl)[nr1][17], thread const float (&sumy)[nr1], short it, + thread float (&sumf)[nr1]) { + float acc[nr1] = {}; + FOR_UNROLL (short byte = 0; byte < 3; ++byte) { + const float u = (float) qb->qs[byte < 2 ? 2*it + byte : 16 + it] * (1.0f/256.0f); + const float g[5] = {floor(3.0f*u), floor(9.0f*u), floor(27.0f*u), floor(81.0f*u), floor(243.0f*u)}; + FOR_UNROLL (short col = 0; col < nr1; ++col) { + FOR_UNROLL (short n = 0; n < 5; ++n) { + acc[col] += g[n] * yl[col][5*byte + n]; + } + } + } + const float u = (float) qb->qh[it & 1] * (1.0f/256.0f); + const float p0 = yl[0][16]; + const float trit = floor(3.0f*p0*u) - 3.0f*floor(p0*u); + const float d = (float) qb->d; + FOR_UNROLL (short col = 0; col < nr1; ++col) { + acc[col] += trit * yl[col][15]; + sumf[col] += (acc[col] - sumy[col]) * d; + } +} + +template +kernel void kernel_mul_mv_ptq1_0_multicol( + constant ggml_metal_kargs_mul_mv & args, + device const char * src0, + device const char * src1, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + const short NSG = FC_mul_mv_nsg; + + const int nb = args.ne00/QK_PTQ1_0; + + const int r0 = tgpig.x; + const int r1 = tgpig.y * nr1; + const int im = tgpig.z; + + const int first_row = (r0 * NSG + sgitg) * nr0; + + const uint i12 = im%FC_mul_mv_ne12; + const uint i13 = im/FC_mul_mv_ne12; + + const uint64_t offset1 = r1*args.nb11 + (i12)*args.nb12 + (i13)*args.nb13; + + device const float * y = (device const float *) (src1 + offset1); + + device const block_ptq1_0 * ax[nr0]; + for (int row = 0; row < nr0; ++row) { + const uint64_t offset0 = min(first_row + row, args.ne01 - 1)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03; + ax[row] = (device const block_ptq1_0 *) ((device char *) src0 + offset0); + } + + // 15 collapse coefficients, the qh activation, and the qh trit's 3^n + float yl[nr1][17]; + float sumf[nr0][nr1] = {}; + + // Eight threads cover one block, with each thread reading whole bytes. + const short ix = (tiisg/8); + const short it = (tiisg%8); + + device const float * yb = y + ix*QK_PTQ1_0; + + { + const float pow3f[4] = {1.0f, 3.0f, 9.0f, 27.0f}; + FOR_UNROLL (short col = 0; col < nr1; ++col) { + yl[col][16] = pow3f[it >> 1]; + } + } + + for (int ib = ix; ib < nb; ib += N_SIMDWIDTH/8) { + // Reuse collapse coefficients across rows: c[k-1] = y_{k-1} - 3*y_k, c[4] = y_4. + float sumy[nr1] = {}; + FOR_UNROLL (short col = 0; col < nr1; ++col) { + device const float * yc = (device const float *) ((device const char *) yb + col*args.nb11); + + FOR_UNROLL (short k = 0; k < 2; ++k) { + const short m = 2*it + k; + float y[5]; + FOR_UNROLL (short n = 0; n < 5; ++n) { + y[n] = yc[n*16 + m]; + sumy[col] += y[n]; + } + FOR_UNROLL (short n = 0; n < 4; ++n) { + yl[col][5*k + n] = y[n] - 3.0f*y[n+1]; + } + yl[col][5*k + 4] = y[4]; + } + { + float y[5]; + FOR_UNROLL (short n = 0; n < 5; ++n) { + y[n] = yc[80 + n*8 + it]; + sumy[col] += y[n]; + } + FOR_UNROLL (short n = 0; n < 4; ++n) { + yl[col][10 + n] = y[n] - 3.0f*y[n+1]; + } + yl[col][14] = y[4]; + } + { + const float v = yc[120 + it]; + yl[col][15] = v; + sumy[col] += v; + } + } + + FOR_UNROLL (short row = 0; row < nr0; row++) { + ptq1_0_dot_multicol(ax[row] + ib, yl, sumy, it, sumf[row]); + } + + yb += QK_PTQ1_0 * (N_SIMDWIDTH/8); + } + + device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0; + + for (int row = 0; row < nr0; ++row) { + FOR_UNROLL (short col = 0; col < nr1; ++col) { + const float tot = simd_sum(sumf[row][col]); + if (tiisg == 0 && first_row + row < args.ne01) { + dst_f32[(uint64_t) col*args.ne0 + first_row + row] = tot; + } + } + } +} + +typedef decltype(kernel_mul_mv_ptq1_0_multicol<4, 2>) mul_mv_ptq1_multicol_t; +template [[host_name("kernel_mul_mv_ptq1_0_f32_mc_c2")]] kernel mul_mv_ptq1_multicol_t kernel_mul_mv_ptq1_0_multicol<4, 2>; +template [[host_name("kernel_mul_mv_ptq1_0_f32_mc_c3")]] kernel mul_mv_ptq1_multicol_t kernel_mul_mv_ptq1_0_multicol<4, 3>; +template [[host_name("kernel_mul_mv_ptq1_0_f32_mc_c4")]] kernel mul_mv_ptq1_multicol_t kernel_mul_mv_ptq1_0_multicol<4, 4>; + [[host_name("kernel_mul_mv_ptq1_0_f32")]] kernel void kernel_mul_mv_ptq1_0_f32( constant ggml_metal_kargs_mul_mv & args, diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 6510e15f5446..1755c5b0a637 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9400,6 +9400,13 @@ static std::vector> make_test_cases_eval() { } } + // PTQ1_0 small batches, row tails and broadcast dimensions. + for (int n : {1, 2, 3, 4, 8}) { + for (int k : {128, 384, 5120}) { + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PTQ1_0, GGML_TYPE_F32, 7, n, k, {2, 2}, {2, 1})); + } + } + // BF16 is absent from base_types: add the 3 standard non-contig permutations explicitly test_cases.emplace_back(new test_mul_mat(GGML_TYPE_BF16, GGML_TYPE_F32, 16, 1, 256, {2, 3}, {1, 1}, {0, 2, 1, 3})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_BF16, GGML_TYPE_F32, 16, 1, 256, {2, 3}, {1, 1}, {0, 1, 3, 2})); @@ -10310,7 +10317,7 @@ static std::vector> make_test_cases_perf() { } // batched decode (several sequences per step) through the mat-vec path for (ggml_type t : {GGML_TYPE_PTQ1_0, GGML_TYPE_PQ2_0, GGML_TYPE_Q4_0}) { - for (int n : {2, 4, 8}) { + for (int n : {2, 3, 4, 8}) { test_cases.emplace_back(new test_mul_mat(t, GGML_TYPE_F32, 17408, n, 5120, {1, 1}, {1, 1})); } }