Skip to content
Open
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
58 changes: 38 additions & 20 deletions ggml/src/ggml-metal/ggml-metal-device.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand All @@ -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) {
Expand Down Expand Up @@ -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<int>(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<int>(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;
Expand Down
1 change: 1 addition & 0 deletions ggml/src/ggml-metal/ggml-metal-device.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<MTLBuffer>
Expand Down
2 changes: 1 addition & 1 deletion ggml/src/ggml-metal/ggml-metal-ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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 ||
Expand Down
121 changes: 121 additions & 0 deletions ggml/src/ggml-metal/kernels/mul_mv.metal
Original file line number Diff line number Diff line change
Expand Up @@ -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 <int nr1, typename block_t>
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<int nr0, typename args_t>
void kernel_mul_mv_pq2_0_f32_impl(
args_t args,
Expand Down Expand Up @@ -1404,6 +1437,83 @@ void kernel_mul_mv_pq2_0_f32_nc_impl(args_t args,
}
}

template<int nr0, int nr1>
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<nr1>(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 <int nr0, int nr1>
kernel void kernel_mul_mv_pq2_0_f32_nc(constant ggml_metal_kargs_mul_mv & args,
const device char * src0,
Expand All @@ -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,
Expand Down
6 changes: 5 additions & 1 deletion tests/test-backend-ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9338,9 +9338,13 @@ static std::vector<std::unique_ptr<test_case>> 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}));
Expand Down