Created
August 13, 2026 04:58
-
-
Save kishida/4729c4c3d3d82c7b89e98902260ad1b5 to your computer and use it in GitHub Desktop.
patch for Metal of unsloth's q1 extension
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp | |
| index c153bd821..64c3d3d70 100644 | |
| --- a/ggml/src/ggml-metal/ggml-metal-device.cpp | |
| +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp | |
| @@ -941,6 +941,21 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta | |
| nsg = N_SG_IQ1_M; | |
| nr0 = N_R0_IQ1_M; | |
| } break; | |
| + case GGML_TYPE_IQ1_XS: | |
| + { | |
| + nsg = N_SG_IQ1_XS; | |
| + nr0 = N_R0_IQ1_XS; | |
| + } break; | |
| + case GGML_TYPE_IQ1_XXS: | |
| + { | |
| + nsg = N_SG_IQ1_XXS; | |
| + nr0 = N_R0_IQ1_XXS; | |
| + } break; | |
| + case GGML_TYPE_IQ1_XXXS: | |
| + { | |
| + nsg = N_SG_IQ1_XXXS; | |
| + nr0 = N_R0_IQ1_XXXS; | |
| + } break; | |
| case GGML_TYPE_IQ4_NL: | |
| { | |
| nsg = N_SG_IQ4_NL; | |
| @@ -1170,6 +1185,21 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id(ggml_m | |
| nsg = N_SG_IQ1_M; | |
| nr0 = N_R0_IQ1_M; | |
| } break; | |
| + case GGML_TYPE_IQ1_XS: | |
| + { | |
| + nsg = N_SG_IQ1_XS; | |
| + nr0 = N_R0_IQ1_XS; | |
| + } break; | |
| + case GGML_TYPE_IQ1_XXS: | |
| + { | |
| + nsg = N_SG_IQ1_XXS; | |
| + nr0 = N_R0_IQ1_XXS; | |
| + } break; | |
| + case GGML_TYPE_IQ1_XXXS: | |
| + { | |
| + nsg = N_SG_IQ1_XXXS; | |
| + nr0 = N_R0_IQ1_XXXS; | |
| + } break; | |
| case GGML_TYPE_IQ4_NL: | |
| { | |
| nsg = N_SG_IQ4_NL; | |
| diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h | |
| index e173b91c0..ea726ca51 100644 | |
| --- a/ggml/src/ggml-metal/ggml-metal-impl.h | |
| +++ b/ggml/src/ggml-metal/ggml-metal-impl.h | |
| @@ -66,6 +66,15 @@ | |
| #define N_R0_IQ1_M 4 | |
| #define N_SG_IQ1_M 2 | |
| +#define N_R0_IQ1_XS 4 | |
| +#define N_SG_IQ1_XS 2 | |
| + | |
| +#define N_R0_IQ1_XXS 4 | |
| +#define N_SG_IQ1_XXS 2 | |
| + | |
| +#define N_R0_IQ1_XXXS 4 | |
| +#define N_SG_IQ1_XXXS 2 | |
| + | |
| #define N_R0_IQ2_XXS 4 | |
| #define N_SG_IQ2_XXS 2 | |
| diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal | |
| index 92258b737..c6135f277 100644 | |
| --- a/ggml/src/ggml-metal/ggml-metal.metal | |
| +++ b/ggml/src/ggml-metal/ggml-metal.metal | |
| @@ -946,6 +946,68 @@ void dequantize_iq1_s(device const block_iq1_s * xb, short il, thread type4x4 & | |
| } | |
| } | |
| +template <typename type4x4> | |
| +void dequantize_iq1_xs(device const block_iq1_xs * xb, short il, thread type4x4 & reg) { | |
| + // il is 0...15 for QK_K = 256 => index of block of 32 is il/2 | |
| + const int ib32 = il/2; | |
| + il = il%2; | |
| + const float d = xb->d; | |
| + device const uint8_t * qs = xb->qs + 4*ib32 + 2*il; | |
| + const uint8_t nib = (xb->sc[ib32/2] >> 4*(ib32%2)) & 0xf; | |
| + const float dl = d * (2*(nib & 7) + 1); | |
| + const float ml = dl * (nib & 8 ? -1 - IQ1S_DELTA : -1 + IQ1S_DELTA); | |
| + const uint8_t h = xb->qh[ib32] >> 4*il; | |
| + constant uint8_t * grid1 = (constant uint8_t *)(iq1_xs_grid_gpu + (qs[0] | ((h << 8) & 0x300))); | |
| + constant uint8_t * grid2 = (constant uint8_t *)(iq1_xs_grid_gpu + (qs[1] | ((h << 6) & 0x300))); | |
| + for (int i = 0; i < 4; ++i) { | |
| + reg[0][i] = dl * (grid1[i] & 0xf) + ml; | |
| + reg[1][i] = dl * (grid1[i] >> 4) + ml; | |
| + reg[2][i] = dl * (grid2[i] & 0xf) + ml; | |
| + reg[3][i] = dl * (grid2[i] >> 4) + ml; | |
| + } | |
| +} | |
| + | |
| +template <typename type4x4> | |
| +void dequantize_iq1_xxs(device const block_iq1_xxs * xb, short il, thread type4x4 & reg) { | |
| + // il is 0...15 for QK_K = 256 => index of block of 32 is il/2 | |
| + const int ib32 = il/2; | |
| + il = il%2; | |
| + const float d = xb->d; | |
| + device const uint8_t * qs = xb->qs + 4*ib32 + 2*il; | |
| + const uint8_t qh = xb->qh[ib32]; | |
| + const float dl = d * (2*((qh >> 4) & 7) + 1); | |
| + const float ml = dl * (qh & 0x80 ? -1 - IQ1S_DELTA : -1 + IQ1S_DELTA); | |
| + const uint8_t h = qh >> 2*il; | |
| + constant uint8_t * grid1 = (constant uint8_t *)(iq1_xxs_grid_gpu + (qs[0] | ((h << 8) & 0x100))); | |
| + constant uint8_t * grid2 = (constant uint8_t *)(iq1_xxs_grid_gpu + (qs[1] | ((h << 7) & 0x100))); | |
| + for (int i = 0; i < 4; ++i) { | |
| + reg[0][i] = dl * (grid1[i] & 0xf) + ml; | |
| + reg[1][i] = dl * (grid1[i] >> 4) + ml; | |
| + reg[2][i] = dl * (grid2[i] & 0xf) + ml; | |
| + reg[3][i] = dl * (grid2[i] >> 4) + ml; | |
| + } | |
| +} | |
| + | |
| +template <typename type4x4> | |
| +void dequantize_iq1_xxxs(device const block_iq1_xxxs * xb, short il, thread type4x4 & reg) { | |
| + // il is 0...15 for QK_K = 256 => index of block of 32 is il/2 | |
| + const int ib32 = il/2; | |
| + il = il%2; | |
| + const float d = xb->d; | |
| + device const uint8_t * qs = xb->qs + 4*ib32 + 2*il; | |
| + const uint8_t nib = (xb->sc[ib32/2] >> 4*(ib32%2)) & 0xf; | |
| + const float dl = d * (2*(nib & 7) + 1); | |
| + const float ml = dl * (nib & 8 ? -1 - IQ1S_DELTA : -1 + IQ1S_DELTA); | |
| + constant uint8_t * grid1 = (constant uint8_t *)(iq1_xxxs_grid_gpu + qs[0]); | |
| + constant uint8_t * grid2 = (constant uint8_t *)(iq1_xxxs_grid_gpu + qs[1]); | |
| + for (int i = 0; i < 4; ++i) { | |
| + reg[0][i] = dl * (grid1[i] & 0xf) + ml; | |
| + reg[1][i] = dl * (grid1[i] >> 4) + ml; | |
| + reg[2][i] = dl * (grid2[i] & 0xf) + ml; | |
| + reg[3][i] = dl * (grid2[i] >> 4) + ml; | |
| + } | |
| +} | |
| + | |
| template <typename type4x4> | |
| void dequantize_iq1_m(device const block_iq1_m * xb, short il, thread type4x4 & reg) { | |
| // il is 0...15 for QK_K = 256 => index of block of 32 is il/2 | |
| @@ -9399,6 +9461,309 @@ kernel void kernel_mul_mv_iq1_s_f32( | |
| kernel_mul_mv_iq1_s_f32_impl<N_R0_IQ1_S, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg); | |
| } | |
| +template<int nr0, typename args_t> | |
| +void kernel_mul_mv_iq1_xs_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 offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03; | |
| + const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13; | |
| + | |
| + device const block_iq1_xs * x = (device const block_iq1_xs *) (src0 + offset0); | |
| + device const float * y = (device const float *) (src1 + offset1); | |
| + | |
| + float yl[32]; | |
| + float sumf[nr0]={0.f}; | |
| + | |
| + const int nb32 = nb * (QK_K / 32); | |
| + | |
| + const short ix = tiisg; | |
| + | |
| + device const float * y4 = y + 32 * ix; | |
| + | |
| + for (int ib32 = ix; ib32 < nb32; ib32 += 32) { | |
| + float sumy = 0; | |
| + for (short i = 0; i < 32; ++i) { | |
| + yl[i] = y4[i]; | |
| + sumy += yl[i]; | |
| + } | |
| + | |
| + const int ibl = ib32 / (QK_K / 32); | |
| + const int ib = ib32 % (QK_K / 32); | |
| + | |
| + device const block_iq1_xs * xr = x + ibl; | |
| + device const uint8_t * qs = xr->qs + 4 * ib; | |
| + device const uint8_t * qh = xr->qh + ib; | |
| + device const uint8_t * sc = xr->sc + ib/2; | |
| + device const half * dh = &xr->d; | |
| + | |
| + for (short row = 0; row < nr0; row++) { | |
| + constant uint8_t * grid1 = (constant uint8_t *)(iq1_xs_grid_gpu + (qs[0] | ((qh[0] << 8) & 0x300))); | |
| + constant uint8_t * grid2 = (constant uint8_t *)(iq1_xs_grid_gpu + (qs[1] | ((qh[0] << 6) & 0x300))); | |
| + constant uint8_t * grid3 = (constant uint8_t *)(iq1_xs_grid_gpu + (qs[2] | ((qh[0] << 4) & 0x300))); | |
| + constant uint8_t * grid4 = (constant uint8_t *)(iq1_xs_grid_gpu + (qs[3] | ((qh[0] << 2) & 0x300))); | |
| + | |
| + const uint8_t nib = (sc[0] >> 4*(ib%2)) & 0xf; | |
| + | |
| + float sum = 0; | |
| + for (short j = 0; j < 4; ++j) { | |
| + sum += yl[j+ 0] * (grid1[j] & 0xf) + yl[j+ 4] * (grid1[j] >> 4) | |
| + + yl[j+ 8] * (grid2[j] & 0xf) + yl[j+12] * (grid2[j] >> 4) | |
| + + yl[j+16] * (grid3[j] & 0xf) + yl[j+20] * (grid3[j] >> 4) | |
| + + yl[j+24] * (grid4[j] & 0xf) + yl[j+28] * (grid4[j] >> 4); | |
| + } | |
| + sumf[row] += (float)dh[0] * (sum + sumy * (nib & 8 ? -1 - IQ1S_DELTA : -1 + IQ1S_DELTA)) * (2*(nib & 7) + 1); | |
| + | |
| + dh += args.nb01/2; | |
| + qs += args.nb01; | |
| + qh += args.nb01; | |
| + sc += args.nb01; | |
| + } | |
| + | |
| + y4 += 32 * 32; | |
| + } | |
| + | |
| + 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 && first_row + row < args.ne0; ++row) { | |
| + float sum_all = simd_sum(sumf[row]); | |
| + if (tiisg == 0) { | |
| + dst_f32[first_row + row] = sum_all; | |
| + } | |
| + } | |
| +} | |
| + | |
| +[[host_name("kernel_mul_mv_iq1_xs_f32")]] | |
| +kernel void kernel_mul_mv_iq1_xs_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_iq1_xs_f32_impl<N_R0_IQ1_XS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg); | |
| +} | |
| + | |
| +template<int nr0, typename args_t> | |
| +void kernel_mul_mv_iq1_xxs_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 offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03; | |
| + const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13; | |
| + | |
| + device const block_iq1_xxs * x = (device const block_iq1_xxs *) (src0 + offset0); | |
| + device const float * y = (device const float *) (src1 + offset1); | |
| + | |
| + float yl[32]; | |
| + float sumf[nr0]={0.f}; | |
| + | |
| + const int nb32 = nb * (QK_K / 32); | |
| + | |
| + const short ix = tiisg; | |
| + | |
| + device const float * y4 = y + 32 * ix; | |
| + | |
| + for (int ib32 = ix; ib32 < nb32; ib32 += 32) { | |
| + float sumy = 0; | |
| + for (short i = 0; i < 32; ++i) { | |
| + yl[i] = y4[i]; | |
| + sumy += yl[i]; | |
| + } | |
| + | |
| + const int ibl = ib32 / (QK_K / 32); | |
| + const int ib = ib32 % (QK_K / 32); | |
| + | |
| + device const block_iq1_xxs * xr = x + ibl; | |
| + device const uint8_t * qs = xr->qs + 4 * ib; | |
| + device const uint8_t * qh = xr->qh + ib; | |
| + device const half * dh = &xr->d; | |
| + | |
| + for (short row = 0; row < nr0; row++) { | |
| + constant uint8_t * grid1 = (constant uint8_t *)(iq1_xxs_grid_gpu + (qs[0] | ((qh[0] << 8) & 0x100))); | |
| + constant uint8_t * grid2 = (constant uint8_t *)(iq1_xxs_grid_gpu + (qs[1] | ((qh[0] << 7) & 0x100))); | |
| + constant uint8_t * grid3 = (constant uint8_t *)(iq1_xxs_grid_gpu + (qs[2] | ((qh[0] << 6) & 0x100))); | |
| + constant uint8_t * grid4 = (constant uint8_t *)(iq1_xxs_grid_gpu + (qs[3] | ((qh[0] << 5) & 0x100))); | |
| + | |
| + float sum = 0; | |
| + for (short j = 0; j < 4; ++j) { | |
| + sum += yl[j+ 0] * (grid1[j] & 0xf) + yl[j+ 4] * (grid1[j] >> 4) | |
| + + yl[j+ 8] * (grid2[j] & 0xf) + yl[j+12] * (grid2[j] >> 4) | |
| + + yl[j+16] * (grid3[j] & 0xf) + yl[j+20] * (grid3[j] >> 4) | |
| + + yl[j+24] * (grid4[j] & 0xf) + yl[j+28] * (grid4[j] >> 4); | |
| + } | |
| + sumf[row] += (float)dh[0] * (sum + sumy * (qh[0] & 0x80 ? -1 - IQ1S_DELTA : -1 + IQ1S_DELTA)) * (2*((qh[0] >> 4) & 7) + 1); | |
| + | |
| + dh += args.nb01/2; | |
| + qs += args.nb01; | |
| + qh += args.nb01; | |
| + } | |
| + | |
| + y4 += 32 * 32; | |
| + } | |
| + | |
| + 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 && first_row + row < args.ne0; ++row) { | |
| + float sum_all = simd_sum(sumf[row]); | |
| + if (tiisg == 0) { | |
| + dst_f32[first_row + row] = sum_all; | |
| + } | |
| + } | |
| +} | |
| + | |
| +[[host_name("kernel_mul_mv_iq1_xxs_f32")]] | |
| +kernel void kernel_mul_mv_iq1_xxs_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_iq1_xxs_f32_impl<N_R0_IQ1_XXS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg); | |
| +} | |
| + | |
| +template<int nr0, typename args_t> | |
| +void kernel_mul_mv_iq1_xxxs_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 offset0 = first_row*args.nb01 + (i12/FC_mul_mv_r2)*args.nb02 + (i13/FC_mul_mv_r3)*args.nb03; | |
| + const uint64_t offset1 = r1*args.nb11 + (i12 )*args.nb12 + (i13 )*args.nb13; | |
| + | |
| + device const block_iq1_xxxs * x = (device const block_iq1_xxxs *) (src0 + offset0); | |
| + device const float * y = (device const float *) (src1 + offset1); | |
| + | |
| + float yl[32]; | |
| + float sumf[nr0]={0.f}; | |
| + | |
| + const int nb32 = nb * (QK_K / 32); | |
| + | |
| + const short ix = tiisg; | |
| + | |
| + device const float * y4 = y + 32 * ix; | |
| + | |
| + for (int ib32 = ix; ib32 < nb32; ib32 += 32) { | |
| + float sumy = 0; | |
| + for (short i = 0; i < 32; ++i) { | |
| + yl[i] = y4[i]; | |
| + sumy += yl[i]; | |
| + } | |
| + | |
| + const int ibl = ib32 / (QK_K / 32); | |
| + const int ib = ib32 % (QK_K / 32); | |
| + | |
| + device const block_iq1_xxxs * xr = x + ibl; | |
| + device const uint8_t * qs = xr->qs + 4 * ib; | |
| + device const uint8_t * sc = xr->sc + ib/2; | |
| + device const half * dh = &xr->d; | |
| + | |
| + for (short row = 0; row < nr0; row++) { | |
| + constant uint8_t * grid1 = (constant uint8_t *)(iq1_xxxs_grid_gpu + qs[0]); | |
| + constant uint8_t * grid2 = (constant uint8_t *)(iq1_xxxs_grid_gpu + qs[1]); | |
| + constant uint8_t * grid3 = (constant uint8_t *)(iq1_xxxs_grid_gpu + qs[2]); | |
| + constant uint8_t * grid4 = (constant uint8_t *)(iq1_xxxs_grid_gpu + qs[3]); | |
| + | |
| + const uint8_t nib = (sc[0] >> 4*(ib%2)) & 0xf; | |
| + | |
| + float sum = 0; | |
| + for (short j = 0; j < 4; ++j) { | |
| + sum += yl[j+ 0] * (grid1[j] & 0xf) + yl[j+ 4] * (grid1[j] >> 4) | |
| + + yl[j+ 8] * (grid2[j] & 0xf) + yl[j+12] * (grid2[j] >> 4) | |
| + + yl[j+16] * (grid3[j] & 0xf) + yl[j+20] * (grid3[j] >> 4) | |
| + + yl[j+24] * (grid4[j] & 0xf) + yl[j+28] * (grid4[j] >> 4); | |
| + } | |
| + sumf[row] += (float)dh[0] * (sum + sumy * (nib & 8 ? -1 - IQ1S_DELTA : -1 + IQ1S_DELTA)) * (2*(nib & 7) + 1); | |
| + | |
| + dh += args.nb01/2; | |
| + qs += args.nb01; | |
| + sc += args.nb01; | |
| + } | |
| + | |
| + y4 += 32 * 32; | |
| + } | |
| + | |
| + 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 && first_row + row < args.ne0; ++row) { | |
| + float sum_all = simd_sum(sumf[row]); | |
| + if (tiisg == 0) { | |
| + dst_f32[first_row + row] = sum_all; | |
| + } | |
| + } | |
| +} | |
| + | |
| +[[host_name("kernel_mul_mv_iq1_xxxs_f32")]] | |
| +kernel void kernel_mul_mv_iq1_xxxs_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_iq1_xxxs_f32_impl<N_R0_IQ1_XXXS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, nullptr, tgpig, tiisg, sgitg); | |
| +} | |
| + | |
| template<int nr0, typename args_t> | |
| void kernel_mul_mv_iq1_m_f32_impl( | |
| args_t args, | |
| @@ -9913,6 +10278,9 @@ template [[host_name("kernel_get_rows_iq3_s")]] kernel get_rows_q_t kernel_get | |
| template [[host_name("kernel_get_rows_iq2_s")]] kernel get_rows_q_t kernel_get_rows_q<block_iq2_s, QK_NL, dequantize_iq2_s>; | |
| template [[host_name("kernel_get_rows_iq1_s")]] kernel get_rows_q_t kernel_get_rows_q<block_iq1_s, QK_NL, dequantize_iq1_s>; | |
| 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_iq1_xs")]] kernel get_rows_q_t kernel_get_rows_q<block_iq1_xs, QK_NL, dequantize_iq1_xs>; | |
| +template [[host_name("kernel_get_rows_iq1_xxs")]] kernel get_rows_q_t kernel_get_rows_q<block_iq1_xxs, QK_NL, dequantize_iq1_xxs>; | |
| +template [[host_name("kernel_get_rows_iq1_xxxs")]] kernel get_rows_q_t kernel_get_rows_q<block_iq1_xxxs, QK_NL, dequantize_iq1_xxxs>; | |
| 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>; | |
| @@ -10784,6 +11152,9 @@ template [[host_name("kernel_mul_mm_iq3_s_f32")]] kernel mul_mm_t kernel_mul_m | |
| template [[host_name("kernel_mul_mm_iq2_s_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq2_s, QK_NL, dequantize_iq2_s, float, float4x4, float, float2x4>; | |
| template [[host_name("kernel_mul_mm_iq1_s_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_s, QK_NL, dequantize_iq1_s, float, float4x4, float, float2x4>; | |
| 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_iq1_xs_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xs, QK_NL, dequantize_iq1_xs, float, float4x4, float, float2x4>; | |
| +template [[host_name("kernel_mul_mm_iq1_xxs_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxs, QK_NL, dequantize_iq1_xxs, float, float4x4, float, float2x4>; | |
| +template [[host_name("kernel_mul_mm_iq1_xxxs_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxxs, QK_NL, dequantize_iq1_xxxs, 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>; | |
| @@ -10809,6 +11180,9 @@ template [[host_name("kernel_mul_mm_iq3_s_f16")]] kernel mul_mm_t kernel_mul_m | |
| template [[host_name("kernel_mul_mm_iq2_s_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq2_s, QK_NL, dequantize_iq2_s, float, float4x4, half, half2x4>; | |
| template [[host_name("kernel_mul_mm_iq1_s_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_s, QK_NL, dequantize_iq1_s, float, float4x4, half, half2x4>; | |
| 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_iq1_xs_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xs, QK_NL, dequantize_iq1_xs, float, float4x4, half, half2x4>; | |
| +template [[host_name("kernel_mul_mm_iq1_xxs_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxs, QK_NL, dequantize_iq1_xxs, float, float4x4, half, half2x4>; | |
| +template [[host_name("kernel_mul_mm_iq1_xxxs_f16")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxxs, QK_NL, dequantize_iq1_xxxs, 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>; | |
| @@ -10843,6 +11217,9 @@ template [[host_name("kernel_mul_mm_id_iq3_s_f32")]] kernel mul_mm_id kernel_m | |
| template [[host_name("kernel_mul_mm_id_iq2_s_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq2_s, QK_NL, dequantize_iq2_s, float, float4x4, float, float2x4>; | |
| template [[host_name("kernel_mul_mm_id_iq1_s_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_s, QK_NL, dequantize_iq1_s, float, float4x4, float, float2x4>; | |
| 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_iq1_xs_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xs, QK_NL, dequantize_iq1_xs, float, float4x4, float, float2x4>; | |
| +template [[host_name("kernel_mul_mm_id_iq1_xxs_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxs, QK_NL, dequantize_iq1_xxs, float, float4x4, float, float2x4>; | |
| +template [[host_name("kernel_mul_mm_id_iq1_xxxs_f32")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxxs, QK_NL, dequantize_iq1_xxxs, 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>; | |
| @@ -10868,6 +11245,9 @@ template [[host_name("kernel_mul_mm_id_iq3_s_f16")]] kernel mul_mm_id kernel_m | |
| template [[host_name("kernel_mul_mm_id_iq2_s_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq2_s, QK_NL, dequantize_iq2_s, float, float4x4, half, half2x4>; | |
| template [[host_name("kernel_mul_mm_id_iq1_s_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_s, QK_NL, dequantize_iq1_s, float, float4x4, half, half2x4>; | |
| 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_iq1_xs_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xs, QK_NL, dequantize_iq1_xs, float, float4x4, half, half2x4>; | |
| +template [[host_name("kernel_mul_mm_id_iq1_xxs_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxs, QK_NL, dequantize_iq1_xxs, float, float4x4, half, half2x4>; | |
| +template [[host_name("kernel_mul_mm_id_iq1_xxxs_f16")]] kernel mul_mm_id kernel_mul_mm_id<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, block_iq1_xxxs, QK_NL, dequantize_iq1_xxxs, 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>; | |
| @@ -11020,6 +11400,9 @@ template [[host_name("kernel_mul_mv_id_q5_K_f32")]] kernel kernel_mul_mv_id_t | |
| template [[host_name("kernel_mul_mv_id_q6_K_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_q6_K_f32_impl <N_R0_Q6_K>>>; | |
| template [[host_name("kernel_mul_mv_id_iq1_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_s_f32_impl <N_R0_IQ1_S>>>; | |
| template [[host_name("kernel_mul_mv_id_iq1_m_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_m_f32_impl <N_R0_IQ1_M>>>; | |
| +template [[host_name("kernel_mul_mv_id_iq1_xs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_xs_f32_impl <N_R0_IQ1_XS>>>; | |
| +template [[host_name("kernel_mul_mv_id_iq1_xxs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_xxs_f32_impl<N_R0_IQ1_XXS>>>; | |
| +template [[host_name("kernel_mul_mv_id_iq1_xxxs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq1_xxxs_f32_impl<N_R0_IQ1_XXXS>>>; | |
| template [[host_name("kernel_mul_mv_id_iq2_xxs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_xxs_f32_impl<N_R0_IQ2_XXS>>>; | |
| template [[host_name("kernel_mul_mv_id_iq2_xs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq2_xs_f32_impl <N_R0_IQ2_XS>>>; | |
| template [[host_name("kernel_mul_mv_id_iq3_xxs_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq3_xxs_f32_impl<N_R0_IQ3_XXS>>>; |
Author
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
patch for this PR
unslothai/llama.cpp#61