diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index ad42b09ca730..4629dbecd9c9 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -909,6 +909,14 @@ bool ggml_metal_ptq1_multicol_enabled(const ggml_tensor * op) { op->src[1]->ne[1] >= 2 && op->src[1]->ne[1] <= 4; } +bool ggml_metal_pq2_0_multicol_enabled(const ggml_tensor * op, bool has_tensor) { + static const bool disabled = getenv("GGML_METAL_PQ2_0_MULTICOL_DISABLE") && atoi(getenv("GGML_METAL_PQ2_0_MULTICOL_DISABLE")) == 1; + return !disabled && !has_tensor && op->src[0]->type == GGML_TYPE_PQ2_0 && op->src[1]->type == GGML_TYPE_F32 && + op->src[0]->ne[0] % ggml_blck_size(GGML_TYPE_PQ2_0) == 0 && op->src[1]->nb[0] == sizeof(float) && + op->src[1]->ne[1] >= 2 && op->src[1]->ne[1] <= 8; +} + + 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); @@ -927,6 +935,9 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta const char * suffix = ""; char ptq1_suffix[16]; + char pq2_suffix[16]; + + const bool has_tensor = ggml_metal_device_get_props(ggml_metal_library_get_device(lib))->has_tensor; // use custom matrix x vector kernel switch (tsrc0) { @@ -971,28 +982,35 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta } break; case GGML_TYPE_PQ2_0: { - nsg = N_SG_PQ2_0; - nr0 = N_R0_PQ2_0; - // multi-column variants decode each weight byte once per nr1 src1 columns - // (spec-decode verify batches). GGML_METAL_PQ2_0_NR1=1 restores mul_mv_ext. - static const int nr1_env = getenv("GGML_METAL_PQ2_0_NR1") ? atoi(getenv("GGML_METAL_PQ2_0_NR1")) : 0; - if (nr1_env >= 2 && nr1_env <= 4 && ne11 >= 2) { - nr1 = std::min(nr1_env, ne11); - } else if (nr1_env != 1 && (ne11 == 2 || ne11 == 3)) { + nsg = N_SG_PQ2_0; + nr0 = N_R0_PQ2_0; + if (ggml_metal_pq2_0_multicol_enabled(op, has_tensor)) { + nr0 = 4; + nsg = 1; nr1 = ne11; + snprintf(pq2_suffix, sizeof(pq2_suffix), "_mc_c%d", nr1); + suffix = pq2_suffix; + } else { + // multi-column variants decode each weight byte once per nr1 src1 columns + // (spec-decode verify batches). GGML_METAL_PQ2_0_NR1=1 restores mul_mv_ext. + static const int nr1_env = getenv("GGML_METAL_PQ2_0_NR1") ? atoi(getenv("GGML_METAL_PQ2_0_NR1")) : 0; + if (nr1_env >= 2 && nr1_env <= 4 && ne11 >= 2) { + nr1 = std::min(nr1_env, ne11); + } else if (nr1_env != 1 && (ne11 == 2 || ne11 == 3)) { + nr1 = ne11; + } + if (nr1 > 1) { + static const int nr0_env = + getenv("GGML_METAL_PQ2_0_NC_NR0") ? atoi(getenv("GGML_METAL_PQ2_0_NC_NR0")) : 4; + static const int nsg_env = + getenv("GGML_METAL_PQ2_0_NC_NSG") ? atoi(getenv("GGML_METAL_PQ2_0_NC_NSG")) : N_SG_PQ2_0; + nr0 = (nr0_env == 2 || nr0_env == 8) ? nr0_env : 4; + nsg = nsg_env >= 1 && nsg_env <= 8 ? nsg_env : N_SG_PQ2_0; + snprintf(ptq1_suffix, sizeof(ptq1_suffix), "_nr1_%d_r%d", nr1, nr0); + suffix = ptq1_suffix; + } } - if (nr1 > 1) { - static const int nr0_env = - getenv("GGML_METAL_PQ2_0_NC_NR0") ? atoi(getenv("GGML_METAL_PQ2_0_NC_NR0")) : 4; - static const int nsg_env = - getenv("GGML_METAL_PQ2_0_NC_NSG") ? atoi(getenv("GGML_METAL_PQ2_0_NC_NSG")) : N_SG_PQ2_0; - nr0 = (nr0_env == 2 || nr0_env == 8) ? nr0_env : 4; - nsg = nsg_env >= 1 && nsg_env <= 8 ? nsg_env : N_SG_PQ2_0; - snprintf(ptq1_suffix, sizeof(ptq1_suffix), "_nr1_%d_r%d", nr1, nr0); - suffix = ptq1_suffix; - } - } - break; + } break; case GGML_TYPE_PTQ1_0: { nsg = N_SG_PTQ1_0; diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 5efc4c82f357..60215687b590 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -7,6 +7,7 @@ extern "C" { #endif bool ggml_metal_ptq1_multicol_enabled(const struct ggml_tensor * op); +bool ggml_metal_pq2_0_multicol_enabled(const struct ggml_tensor * op, bool has_tensor); struct ggml_metal_buffer_id { void * metal; // id diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index dc2a9072ccae..9ea88dc61fa5 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -2943,7 +2943,7 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) { op->src[0]->type == GGML_TYPE_BF16 || (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 && (pq2_0_ext_enable || ne11 >= 4)) || + (op->src[0]->type == GGML_TYPE_PQ2_0 && (pq2_0_ext_enable || (!ggml_metal_pq2_0_multicol_enabled(op, props_dev->has_tensor) && ne11 >= 4))) || (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 || diff --git a/ggml/src/ggml-metal/kernels/mul_mv.metal b/ggml/src/ggml-metal/kernels/mul_mv.metal index d0a39973dd37..01f8124471ef 100644 --- a/ggml/src/ggml-metal/kernels/mul_mv.metal +++ b/ggml/src/ggml-metal/kernels/mul_mv.metal @@ -1193,6 +1193,39 @@ inline float q2_dot_coeffs(device const block_t * qb, float sumy, thread const f return qb->d * (acc - sumy); } +template +inline void pq2_0_dot_multicol( + device const block_t * qb, + thread const float yl[nr1][16], + thread const float sumy[nr1], + int il, + thread float sumf[nr1]) { + device const uint8_t * qs = qb->qs + (il / 4); + + float acc[nr1] = {}; + + FOR_UNROLL (short j = 0; j < 4; j++) { + const float b = (float) qs[j]; + const float u = b * (1.0f/256.0f); + const float f4 = floor( 4.0f*u); + const float f16 = floor(16.0f*u); + const float f64 = floor(64.0f*u); + + FOR_UNROLL (short col = 0; col < nr1; ++col) { + acc[col] += f4 * yl[col][4*j + 0]; + acc[col] += f16 * yl[col][4*j + 1]; + acc[col] += f64 * yl[col][4*j + 2]; + acc[col] += b * yl[col][4*j + 3]; + } + } + + const float d = (float) qb->d; + FOR_UNROLL (short col = 0; col < nr1; ++col) { + sumf[col] += d * (acc[col] - sumy[col]); + } +} + + template void kernel_mul_mv_pq2_0_f32_impl( args_t args, @@ -1404,6 +1437,83 @@ void kernel_mul_mv_pq2_0_f32_nc_impl(args_t args, } } +template +kernel void kernel_mul_mv_pq2_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_PQ2_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_pq2_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_pq2_0 *) ((device char *) src0 + offset0); + } + + float yl[nr1][16]; + float sumf[nr0][nr1] = {}; + + const short ix = (tiisg/8); + const short il = (tiisg%8)*16; + + device const float * yb = y + ix*QK_PQ2_0 + il; + + for (int ib = ix; ib < nb; ib += N_SIMDWIDTH/8) { + 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 j = 0; j < 4; j++) { + const float y0 = yc[4*j + 0]; + const float y1 = yc[4*j + 1]; + const float y2 = yc[4*j + 2]; + const float y3 = yc[4*j + 3]; + sumy[col] += (y0 + y1) + (y2 + y3); + yl[col][4*j + 0] = y3 - 4.0f*y2; + yl[col][4*j + 1] = y2 - 4.0f*y1; + yl[col][4*j + 2] = y1 - 4.0f*y0; + yl[col][4*j + 3] = y0; + } + } + + FOR_UNROLL (short row = 0; row < nr0; row++) { + pq2_0_dot_multicol(ax[row] + ib, yl, sumy, il, sumf[row]); + } + + yb += QK_PQ2_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; + } + } + } +} + template kernel void kernel_mul_mv_pq2_0_f32_nc(constant ggml_metal_kargs_mul_mv & args, const device char * src0, @@ -1426,6 +1536,17 @@ template [[host_name("kernel_mul_mv_pq2_0_f32_nr1_4_r4")]] kernel mul_mv_pq2_0_n template [[host_name("kernel_mul_mv_pq2_0_f32_nr1_2_r8")]] kernel mul_mv_pq2_0_nc_t kernel_mul_mv_pq2_0_f32_nc<8, 2>; template [[host_name("kernel_mul_mv_pq2_0_f32_nr1_3_r8")]] kernel mul_mv_pq2_0_nc_t kernel_mul_mv_pq2_0_f32_nc<8, 3>; template [[host_name("kernel_mul_mv_pq2_0_f32_nr1_4_r8")]] kernel mul_mv_pq2_0_nc_t kernel_mul_mv_pq2_0_f32_nc<8, 4>; + +typedef decltype(kernel_mul_mv_pq2_0_multicol<4, 2>) mul_mv_pq2_multicol_t; +template [[host_name("kernel_mul_mv_pq2_0_f32_mc_c2")]] kernel mul_mv_pq2_multicol_t kernel_mul_mv_pq2_0_multicol<4, 2>; +template [[host_name("kernel_mul_mv_pq2_0_f32_mc_c3")]] kernel mul_mv_pq2_multicol_t kernel_mul_mv_pq2_0_multicol<4, 3>; +template [[host_name("kernel_mul_mv_pq2_0_f32_mc_c4")]] kernel mul_mv_pq2_multicol_t kernel_mul_mv_pq2_0_multicol<4, 4>; +template [[host_name("kernel_mul_mv_pq2_0_f32_mc_c5")]] kernel mul_mv_pq2_multicol_t kernel_mul_mv_pq2_0_multicol<4, 5>; +template [[host_name("kernel_mul_mv_pq2_0_f32_mc_c6")]] kernel mul_mv_pq2_multicol_t kernel_mul_mv_pq2_0_multicol<4, 6>; +template [[host_name("kernel_mul_mv_pq2_0_f32_mc_c7")]] kernel mul_mv_pq2_multicol_t kernel_mul_mv_pq2_0_multicol<4, 7>; +template [[host_name("kernel_mul_mv_pq2_0_f32_mc_c8")]] kernel mul_mv_pq2_multicol_t kernel_mul_mv_pq2_0_multicol<4, 8>; + + kernel void kernel_mul_mv_q4_0_f32( constant ggml_metal_kargs_mul_mv & args, device const char * src0, diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 3e46715a2096..8c929eed4e18 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9338,9 +9338,13 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PQ2_0, GGML_TYPE_F32, 97, n, k, {1, 1}, {1, 1})); } } - for (int64_t n : {9, 16, 17, 32, 64, 512}) { + // PQ2_0 multi-column verify widths (2 to 8 columns) and few-row tensor tile sizes + for (int64_t n : {2, 3, 4, 5, 6, 7, 8, 9, 16, 17, 32, 64, 512}) { test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PQ2_0, GGML_TYPE_F32, 4096, n, 17408, {1, 1}, {1, 1})); } + for (int64_t n = 2; n <= 8; ++n) { + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PQ2_0, GGML_TYPE_F32, 17408, n, 5120, {1, 1}, {1, 1})); + } test_cases.emplace_back(new test_mul_mat(GGML_TYPE_Q4_0, GGML_TYPE_F32, 2880, 32, 2880, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_Q8_0, GGML_TYPE_F32, 2880, 32, 2880, {1, 1}, {1, 1}));