|
|
|
@@ -468,6 +468,34 @@ void quantize_iq4_nl(device const float * src, device block_iq4_nl & dst) {
|
|
|
|
|
dst.d = sumq2 > 0 ? sumqx/sumq2 : d;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
void quantize_tq2_0(device const float * src, device block_tq2_0 & dst) {
|
|
|
|
|
#pragma METAL fp math_mode(safe)
|
|
|
|
|
float amax = 0.0f; // absolute max
|
|
|
|
|
|
|
|
|
|
for (int j = 0; j < QK_K; j++) {
|
|
|
|
|
const float v = src[j];
|
|
|
|
|
amax = MAX(amax, fabs(v));
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
const float d = amax;
|
|
|
|
|
const float id = d ? 1.0f/d : 0.0f;
|
|
|
|
|
|
|
|
|
|
dst.d = (half) d;
|
|
|
|
|
|
|
|
|
|
for (int j = 0; j < QK_K/4; j += 32) {
|
|
|
|
|
for (int m = 0; m < 32; ++m) {
|
|
|
|
|
uint8_t q = 0;
|
|
|
|
|
for (int n = 0; n < 4; ++n) {
|
|
|
|
|
// -1, 0, 1 -> 0, 1, 2
|
|
|
|
|
int xi = (int)round(src[m + n*32] * id) + 1;
|
|
|
|
|
q += (uint8_t)((xi & 3) << (2*n));
|
|
|
|
|
}
|
|
|
|
|
dst.qs[j + m] = q;
|
|
|
|
|
}
|
|
|
|
|
src += 4*32;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <typename type4x4>
|
|
|
|
|
void dequantize_q4_1(device const block_q4_1 * xb, short il, thread type4x4 & reg) {
|
|
|
|
|
device const uint16_t * qs = ((device const uint16_t *)xb + 2);
|
|
|
|
@@ -1021,6 +1049,25 @@ void dequantize_iq4_xs(device const block_iq4_xs * xb, short il, thread type4x4
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template <typename type4x4>
|
|
|
|
|
void dequantize_tq2_0(device const block_tq2_0 * xb, short il, thread type4x4 & reg) {
|
|
|
|
|
device const uint8_t * qs = xb->qs;
|
|
|
|
|
const float d = xb->d;
|
|
|
|
|
|
|
|
|
|
float4x4 reg_f;
|
|
|
|
|
|
|
|
|
|
// 2 bits per element, 4 elements per byte, 128 elements per 32-byte group
|
|
|
|
|
const short base = il * 16;
|
|
|
|
|
for (int k = 0; k < 16; k++) {
|
|
|
|
|
const int i = base + k;
|
|
|
|
|
const int byte = ((i >> 7) & 1) * 32 + (i & 31);
|
|
|
|
|
const int l = (i >> 5) & 3;
|
|
|
|
|
reg_f[k/4][k%4] = d * (float)(((qs[byte] >> (2*l)) & 3) - 1);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
reg = (type4x4) reg_f;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
enum ggml_sort_order {
|
|
|
|
|
GGML_SORT_ORDER_ASC,
|
|
|
|
|
GGML_SORT_ORDER_DESC,
|
|
|
|
@@ -8001,6 +8048,7 @@ template [[host_name("kernel_cpy_f32_q4_1")]] kernel cpy_f_q_t kernel_cpy_f32_
|
|
|
|
|
template [[host_name("kernel_cpy_f32_q5_0")]] kernel cpy_f_q_t kernel_cpy_f32_q<QK5_0, block_q5_0, quantize_q5_0>;
|
|
|
|
|
template [[host_name("kernel_cpy_f32_q5_1")]] kernel cpy_f_q_t kernel_cpy_f32_q<QK5_1, block_q5_1, quantize_q5_1>;
|
|
|
|
|
template [[host_name("kernel_cpy_f32_iq4_nl")]] kernel cpy_f_q_t kernel_cpy_f32_q<QK4_NL, block_iq4_nl, quantize_iq4_nl>;
|
|
|
|
|
template [[host_name("kernel_cpy_f32_tq2_0")]] kernel cpy_f_q_t kernel_cpy_f32_q<QK_K, block_tq2_0, quantize_tq2_0>;
|
|
|
|
|
|
|
|
|
|
template<typename T4x4, typename block_q, short nl, void (*dequantize_func)(device const block_q *, short, thread T4x4 &)>
|
|
|
|
|
kernel void kernel_cpy_q_f32(
|
|
|
|
@@ -8048,6 +8096,8 @@ template [[host_name("kernel_cpy_q5_0_f32")]] kernel cpy_q_f_t kernel_cpy_q_f32<
|
|
|
|
|
template [[host_name("kernel_cpy_q5_1_f32")]] kernel cpy_q_f_t kernel_cpy_q_f32<float4x4, block_q5_1, 2, dequantize_q5_1>;
|
|
|
|
|
template [[host_name("kernel_cpy_q8_0_f32")]] kernel cpy_q_f_t kernel_cpy_q_f32<float4x4, block_q8_0, 2, dequantize_q8_0>;
|
|
|
|
|
|
|
|
|
|
template [[host_name("kernel_cpy_tq2_0_f32")]] kernel cpy_q_f_t kernel_cpy_q_f32<float4x4, block_tq2_0, QK_NL, dequantize_tq2_0>;
|
|
|
|
|
|
|
|
|
|
template [[host_name("kernel_cpy_q1_0_f16")]] kernel cpy_q_f_t kernel_cpy_q_f32<half4x4, block_q1_0, 8, dequantize_q1_0>;
|
|
|
|
|
template [[host_name("kernel_cpy_q2_0_f16")]] kernel cpy_q_f_t kernel_cpy_q_f32<half4x4, block_q2_0, 4, dequantize_q2_0>;
|
|
|
|
|
template [[host_name("kernel_cpy_q4_0_f16")]] kernel cpy_q_f_t kernel_cpy_q_f32<half4x4, block_q4_0, 2, dequantize_q4_0>;
|
|
|
|
@@ -8056,6 +8106,8 @@ template [[host_name("kernel_cpy_q5_0_f16")]] kernel cpy_q_f_t kernel_cpy_q_f32<
|
|
|
|
|
template [[host_name("kernel_cpy_q5_1_f16")]] kernel cpy_q_f_t kernel_cpy_q_f32<half4x4, block_q5_1, 2, dequantize_q5_1>;
|
|
|
|
|
template [[host_name("kernel_cpy_q8_0_f16")]] kernel cpy_q_f_t kernel_cpy_q_f32<half4x4, block_q8_0, 2, dequantize_q8_0>;
|
|
|
|
|
|
|
|
|
|
template [[host_name("kernel_cpy_tq2_0_f16")]] kernel cpy_q_f_t kernel_cpy_q_f32<half4x4, block_tq2_0, QK_NL, dequantize_tq2_0>;
|
|
|
|
|
|
|
|
|
|
template<typename T>
|
|
|
|
|
kernel void kernel_concat(
|
|
|
|
|
constant ggml_metal_kargs_concat & args,
|
|
|
|
@@ -9822,6 +9874,121 @@ kernel void kernel_mul_mv_mxfp4_f32(
|
|
|
|
|
kernel_mul_mv_mxfp4_f32_impl<N_R0_MXFP4, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template<int nr0, typename args_t>
|
|
|
|
|
void kernel_mul_mv_tq2_0_f32_impl(
|
|
|
|
|
args_t args,
|
|
|
|
|
device const char * src0,
|
|
|
|
|
device const char * src1,
|
|
|
|
|
device char * dst,
|
|
|
|
|
threadgroup char * shmem,
|
|
|
|
|
uint3 tgpig,
|
|
|
|
|
ushort tiisg,
|
|
|
|
|
ushort sgitg) {
|
|
|
|
|
const short NSG = FC_mul_mv_nsg;
|
|
|
|
|
|
|
|
|
|
const int nb = args.ne00/QK_K;
|
|
|
|
|
|
|
|
|
|
const int r0 = tgpig.x;
|
|
|
|
|
const int r1 = tgpig.y;
|
|
|
|
|
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_tq2_0 * ax[nr0];
|
|
|
|
|
for (int row = 0; row < nr0; ++row) {
|
|
|
|
|
const uint64_t offset0 = (first_row + row)*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03;
|
|
|
|
|
ax[row] = (device const block_tq2_0 *) ((device char *) src0 + offset0);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
float sumf[nr0] = {0.f};
|
|
|
|
|
|
|
|
|
|
// 8 threads per block, NBLOCK blocks per pass, 2 halves per block per pass
|
|
|
|
|
constexpr short NBLOCK = 4;
|
|
|
|
|
|
|
|
|
|
constexpr short NB = N_SIMDWIDTH/NBLOCK; // threads per block
|
|
|
|
|
|
|
|
|
|
const short blk = tiisg / NB; // 0..NBLOCK-1, block handled by this thread
|
|
|
|
|
const short htg = tiisg % NB; // 0..NB-1, thread within block (0..7)
|
|
|
|
|
|
|
|
|
|
// byte and y base offsets within the block (32 elements per thread, 4 per byte)
|
|
|
|
|
device const float4 * yb4 = (device const float4 *)(y + 4*htg + blk*QK_K);
|
|
|
|
|
|
|
|
|
|
// hoisted per-byte coefficients (from y) and total y-sum, shared across rows
|
|
|
|
|
// ref: https://github.com/ggml-org/llama.cpp/pull/26980
|
|
|
|
|
float4 coef[4];
|
|
|
|
|
|
|
|
|
|
for (int ib = blk; ib < nb; ib += NBLOCK) {
|
|
|
|
|
FOR_UNROLL (short h0 = 0; h0 < 2; ++h0) {
|
|
|
|
|
const float4 y0 = yb4[ 0 + 32*h0];
|
|
|
|
|
const float4 y1 = yb4[ 8 + 32*h0];
|
|
|
|
|
const float4 y2 = yb4[16 + 32*h0];
|
|
|
|
|
const float4 y3 = yb4[24 + 32*h0];
|
|
|
|
|
|
|
|
|
|
float sumy = 0.f;
|
|
|
|
|
FOR_UNROLL (short j = 0; j < 4; ++j) {
|
|
|
|
|
coef[j] = float4(
|
|
|
|
|
y0[j],
|
|
|
|
|
y1[j] - 4.0f*y0[j],
|
|
|
|
|
y2[j] - 4.0f*y1[j],
|
|
|
|
|
y3[j] - 4.0f*y2[j]);
|
|
|
|
|
|
|
|
|
|
sumy += (y0[j] + y1[j]) + (y2[j] + y3[j]);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
FOR_UNROLL (short row = 0; row < nr0; ++row) {
|
|
|
|
|
device const block_tq2_0 & xb = ax[row][ib];
|
|
|
|
|
device const uchar * qs = xb.qs + 4*htg + 32*h0;
|
|
|
|
|
|
|
|
|
|
float sum = -sumy;
|
|
|
|
|
FOR_UNROLL (short j = 0; j < 4; ++j) {
|
|
|
|
|
// express the 2-bit field shifts (v>>2, v>>4, v>>6) as float floor ops
|
|
|
|
|
const float v = (float)qs[j];
|
|
|
|
|
|
|
|
|
|
const float f0 = v;
|
|
|
|
|
const float f1 = floor(v*0.25f); // v>>2
|
|
|
|
|
const float f2 = floor(v*0.0625); // v>>4
|
|
|
|
|
const float f3 = floor(v*0.015625); // v>>6
|
|
|
|
|
|
|
|
|
|
sum += coef[j][0]*f0 + coef[j][1]*f1 + coef[j][2]*f2 + coef[j][3]*f3;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
sumf[row] += xb.d * sum;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
yb4 += QK_K * NBLOCK / 4;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
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) {
|
|
|
|
|
const float tot = simd_sum(sumf[row]);
|
|
|
|
|
if (tiisg == 0 && first_row + row < args.ne01) {
|
|
|
|
|
dst_f32[first_row + row] = tot;
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
[[host_name("kernel_mul_mv_tq2_0_f32")]]
|
|
|
|
|
kernel void kernel_mul_mv_tq2_0_f32(
|
|
|
|
|
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]]) {
|
|
|
|
|
|
|
|
|
|
kernel_mul_mv_tq2_0_f32_impl<N_R0_TQ2_0, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg);
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template<typename block_q, short nl, void (*dequantize_func)(device const block_q *, short, thread float4x4 &)>
|
|
|
|
|
kernel void kernel_get_rows_q(
|
|
|
|
|
constant ggml_metal_kargs_get_rows & args,
|
|
|
|
@@ -9915,6 +10082,38 @@ template [[host_name("kernel_get_rows_iq1_s")]] kernel get_rows_q_t kernel_get
|
|
|
|
|
template [[host_name("kernel_get_rows_iq1_m")]] kernel get_rows_q_t kernel_get_rows_q<block_iq1_m, QK_NL, dequantize_iq1_m>;
|
|
|
|
|
template [[host_name("kernel_get_rows_iq4_nl")]] kernel get_rows_q_t kernel_get_rows_q<block_iq4_nl, 2, dequantize_iq4_nl>;
|
|
|
|
|
template [[host_name("kernel_get_rows_iq4_xs")]] kernel get_rows_q_t kernel_get_rows_q<block_iq4_xs, QK_NL, dequantize_iq4_xs>;
|
|
|
|
|
template [[host_name("kernel_get_rows_tq2_0")]] kernel get_rows_q_t kernel_get_rows_q<block_tq2_0, QK_NL, dequantize_tq2_0>;
|
|
|
|
|
|
|
|
|
|
template<typename TS, typename TI, short QK, typename block_q, void (*quantize_func)(device const float *, device block_q &)>
|
|
|
|
|
kernel void kernel_set_rows_q(
|
|
|
|
|
constant ggml_metal_kargs_set_rows & args,
|
|
|
|
|
device const void * src0,
|
|
|
|
|
device const void * src1,
|
|
|
|
|
device float * dst,
|
|
|
|
|
uint3 tgpig[[threadgroup_position_in_grid]],
|
|
|
|
|
uint tiitg[[thread_index_in_threadgroup]],
|
|
|
|
|
uint3 tptg [[threads_per_threadgroup]]) {
|
|
|
|
|
const int32_t i03 = tgpig.z;
|
|
|
|
|
const int32_t i02 = tgpig.y;
|
|
|
|
|
|
|
|
|
|
const int32_t i12 = i03%args.ne12;
|
|
|
|
|
const int32_t i11 = i02%args.ne11;
|
|
|
|
|
|
|
|
|
|
const int32_t i01 = tgpig.x*tptg.y + tiitg/tptg.x;
|
|
|
|
|
if (i01 >= args.ne01) {
|
|
|
|
|
return;
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
const int32_t i10 = i01;
|
|
|
|
|
const TI i1 = ((const device TI *) ((const device char *) src1 + i10*args.nb10 + i11*args.nb11 + i12*args.nb12))[0];
|
|
|
|
|
|
|
|
|
|
device block_q * dst_row = ( device block_q *) (( device char *) dst + i1*args.nb1 + i02*args.nb2 + i03*args.nb3);
|
|
|
|
|
const device TS * src_row = (const device TS *) ((const device char *) src0 + i01*args.nb01 + i02*args.nb02 + i03*args.nb03);
|
|
|
|
|
|
|
|
|
|
for (int ind = tiitg%tptg.x; ind < args.nk0; ind += tptg.x) {
|
|
|
|
|
quantize_func(src_row + QK*ind, dst_row[ind]);
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
template<typename TS, typename TI, typename block_q, void (*quantize_func)(device const float *, device block_q &)>
|
|
|
|
|
kernel void kernel_set_rows_q32(
|
|
|
|
@@ -10011,6 +10210,11 @@ template [[host_name("kernel_set_rows_f32_i32_q5_1")]] kernel set_rows_q32_t k
|
|
|
|
|
template [[host_name("kernel_set_rows_f32_i64_iq4_nl")]] kernel set_rows_q32_t kernel_set_rows_q32<float, int64_t, block_iq4_nl, quantize_iq4_nl>;
|
|
|
|
|
template [[host_name("kernel_set_rows_f32_i32_iq4_nl")]] kernel set_rows_q32_t kernel_set_rows_q32<float, int32_t, block_iq4_nl, quantize_iq4_nl>;
|
|
|
|
|
|
|
|
|
|
typedef decltype(kernel_set_rows_q<float, int64_t, QK_K, block_tq2_0, quantize_tq2_0>) set_rows_qK_t;
|
|
|
|
|
|
|
|
|
|
template [[host_name("kernel_set_rows_f32_i64_tq2_0")]] kernel set_rows_qK_t kernel_set_rows_q<float, int64_t, QK_K, block_tq2_0, quantize_tq2_0>;
|
|
|
|
|
template [[host_name("kernel_set_rows_f32_i32_tq2_0")]] kernel set_rows_qK_t kernel_set_rows_q<float, int32_t, QK_K, block_tq2_0, quantize_tq2_0>;
|
|
|
|
|
|
|
|
|
|
kernel void kernel_diag_f32(
|
|
|
|
|
constant ggml_metal_kargs_diag & args,
|
|
|
|
|
device const char * src0,
|
|
|
|
@@ -10786,6 +10990,7 @@ template [[host_name("kernel_mul_mm_iq1_s_f32")]] kernel mul_mm_t kernel_mul_m
|
|
|
|
|
template [[host_name("kernel_mul_mm_iq1_m_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_m, QK_NL, dequantize_iq1_m, float, float4x4, float, float2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_iq4_nl_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_nl, 2, dequantize_iq4_nl, float, float4x4, float, float2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_iq4_xs_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_xs, QK_NL, dequantize_iq4_xs, float, float4x4, float, float2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_tq2_0_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_tq2_0, QK_NL, dequantize_tq2_0, float, float4x4, float, float2x4>;
|
|
|
|
|
|
|
|
|
|
template [[host_name("kernel_mul_mm_f32_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, float4x4, 1, dequantize_f32, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_f16_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, half4x4, 1, dequantize_f16, half, half4x4, half, half2x4>;
|
|
|
|
@@ -10811,6 +11016,7 @@ template [[host_name("kernel_mul_mm_iq1_s_f16")]] kernel mul_mm_t kernel_mul_m
|
|
|
|
|
template [[host_name("kernel_mul_mm_iq1_m_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_m, QK_NL, dequantize_iq1_m, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_iq4_nl_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_nl, 2, dequantize_iq4_nl, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_iq4_xs_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_xs, QK_NL, dequantize_iq4_xs, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_tq2_0_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_tq2_0, QK_NL, dequantize_tq2_0, float, float4x4, half, half2x4>;
|
|
|
|
|
|
|
|
|
|
//
|
|
|
|
|
// indirect matrix-matrix multiplication
|
|
|
|
@@ -10845,6 +11051,7 @@ template [[host_name("kernel_mul_mm_id_iq1_s_f32")]] kernel mul_mm_id kernel_m
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_iq1_m_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_m, QK_NL, dequantize_iq1_m, float, float4x4, float, float2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_iq4_nl_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_nl, 2, dequantize_iq4_nl, float, float4x4, float, float2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_iq4_xs_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_xs, QK_NL, dequantize_iq4_xs, float, float4x4, float, float2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_tq2_0_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_tq2_0, QK_NL, dequantize_tq2_0, float, float4x4, float, float2x4>;
|
|
|
|
|
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_f32_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, float4x4, 1, dequantize_f32, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_f16_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, half4x4, 1, dequantize_f16, half, half4x4, half, half2x4>;
|
|
|
|
@@ -10870,6 +11077,7 @@ template [[host_name("kernel_mul_mm_id_iq1_s_f16")]] kernel mul_mm_id kernel_m
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_iq1_m_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_m, QK_NL, dequantize_iq1_m, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_iq4_nl_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_nl, 2, dequantize_iq4_nl, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_iq4_xs_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq4_xs, QK_NL, dequantize_iq4_xs, float, float4x4, half, half2x4>;
|
|
|
|
|
template [[host_name("kernel_mul_mm_id_tq2_0_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_tq2_0, QK_NL, dequantize_tq2_0, float, float4x4, half, half2x4>;
|
|
|
|
|
|
|
|
|
|
//
|
|
|
|
|
// matrix-vector multiplication
|
|
|
|
@@ -11027,6 +11235,7 @@ template [[host_name("kernel_mul_mv_id_iq3_s_f32")]] kernel kernel_mul_mv_id_t
|
|
|
|
|
template [[host_name("kernel_mul_mv_id_iq2_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_s_f32_impl <N_R0_IQ2_S>>>;
|
|
|
|
|
template [[host_name("kernel_mul_mv_id_iq4_nl_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq4_nl_f32_impl <N_R0_IQ4_NL>>>;
|
|
|
|
|
template [[host_name("kernel_mul_mv_id_iq4_xs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq4_xs_f32_impl <N_R0_IQ4_XS>>>;
|
|
|
|
|
template [[host_name("kernel_mul_mv_id_tq2_0_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_tq2_0_f32_impl <N_R0_TQ2_0>>>;
|
|
|
|
|
|
|
|
|
|
kernel void kernel_pool_2d_max_f32(
|
|
|
|
|
constant ggml_metal_kargs_pool_2d & args,
|
|
|
|
|