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
15 changes: 15 additions & 0 deletions ggml/src/ggml-metal/ggml-metal-device.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand All @@ -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) {
Expand Down Expand Up @@ -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:
{
Expand Down
2 changes: 2 additions & 0 deletions ggml/src/ggml-metal/ggml-metal-device.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<MTLBuffer>
size_t offs;
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 @@ -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 ||
Expand Down
134 changes: 134 additions & 0 deletions ggml/src/ggml-metal/kernels/mul_mv.metal
Original file line number Diff line number Diff line change
Expand Up @@ -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<int nr1>
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<int nr0, int nr1>
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<nr1>(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,
Expand Down
9 changes: 8 additions & 1 deletion tests/test-backend-ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9400,6 +9400,13 @@ static std::vector<std::unique_ptr<test_case>> 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}));
Expand Down Expand Up @@ -10310,7 +10317,7 @@ static std::vector<std::unique_ptr<test_case>> 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}));
}
}
Expand Down