metal : fix idle threads in mul_mv_iq3_xxs for ne00 < 1024 (#28086)
* metal : fix half-idle simdgroup in kernel_mul_mv_iq3_xxs_f32 for ne00 < 1024 * metal : keep N_R0_IQ3_XXS = 4, dispatch a separate 8-row split kernel for ne00/32 < 32 The plain kernel is unchanged from master (4 rows per simdgroup, one thread per chunk). The row-split mapping now lives in a separate kernel_mul_mv_iq3_xxs_f32_split instantiation with N_R0_IQ3_XXS_SPLIT = 8, and the host selects it only when ne00/32 < 32 and divides 32, so wide matrices keep the master kernel bit for bit. * metal : select the iq3_xxs row split with a function constant instead of a separate kernel
This commit is contained in:
@@ -839,6 +839,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
|
|||||||
|
|
||||||
const char * suffix = "";
|
const char * suffix = "";
|
||||||
|
|
||||||
|
bool split = false;
|
||||||
|
|
||||||
// use custom matrix x vector kernel
|
// use custom matrix x vector kernel
|
||||||
switch (tsrc0) {
|
switch (tsrc0) {
|
||||||
case GGML_TYPE_F32:
|
case GGML_TYPE_F32:
|
||||||
@@ -942,6 +944,13 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
|
|||||||
nsg = N_SG_IQ3_XXS;
|
nsg = N_SG_IQ3_XXS;
|
||||||
nr0 = N_R0_IQ3_XXS;
|
nr0 = N_R0_IQ3_XXS;
|
||||||
smem = 256*4+128;
|
smem = 256*4+128;
|
||||||
|
|
||||||
|
// split the rows across threads when there are fewer than 32 chunks per row
|
||||||
|
const int nb32 = ne00/32;
|
||||||
|
if (nb32 < 32 && (32 % nb32) == 0) {
|
||||||
|
nr0 = N_R0_IQ3_XXS_SPLIT;
|
||||||
|
split = true;
|
||||||
|
}
|
||||||
} break;
|
} break;
|
||||||
case GGML_TYPE_IQ3_S:
|
case GGML_TYPE_IQ3_S:
|
||||||
{
|
{
|
||||||
@@ -993,7 +1002,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
|
|||||||
const int16_t r3 = (int16_t) (ne13 / ne03);
|
const int16_t r3 = (int16_t) (ne13 / ne03);
|
||||||
|
|
||||||
snprintf(base, 256, "kernel_mul_mv_%s_%s%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1), suffix);
|
snprintf(base, 256, "kernel_mul_mv_%s_%s%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1), suffix);
|
||||||
snprintf(name, 256, "%s_nsg=%d_ne12=%d_r2=%d_r3=%d", base, nsg, ne12, r2, r3);
|
snprintf(name, 256, "%s_nsg=%d_ne12=%d_r2=%d_r3=%d_split=%d", base, nsg, ne12, r2, r3, split);
|
||||||
|
|
||||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||||
if (!res.pipeline) {
|
if (!res.pipeline) {
|
||||||
@@ -1003,6 +1012,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
|
|||||||
ggml_metal_cv_set_int16(cv, (int16_t) ne12, FC_MUL_MV + 2);
|
ggml_metal_cv_set_int16(cv, (int16_t) ne12, FC_MUL_MV + 2);
|
||||||
ggml_metal_cv_set_int16(cv, r2, FC_MUL_MV + 3);
|
ggml_metal_cv_set_int16(cv, r2, FC_MUL_MV + 3);
|
||||||
ggml_metal_cv_set_int16(cv, r3, FC_MUL_MV + 4);
|
ggml_metal_cv_set_int16(cv, r3, FC_MUL_MV + 4);
|
||||||
|
ggml_metal_cv_set_bool (cv, split, FC_MUL_MV + 5);
|
||||||
|
|
||||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||||
|
|
||||||
@@ -1081,6 +1091,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id(ggml_m
|
|||||||
|
|
||||||
const char * suffix = "";
|
const char * suffix = "";
|
||||||
|
|
||||||
|
bool split = false;
|
||||||
|
|
||||||
// use custom matrix x vector kernel
|
// use custom matrix x vector kernel
|
||||||
switch (tsrc0) {
|
switch (tsrc0) {
|
||||||
case GGML_TYPE_F32:
|
case GGML_TYPE_F32:
|
||||||
@@ -1177,6 +1189,13 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id(ggml_m
|
|||||||
nsg = N_SG_IQ3_XXS;
|
nsg = N_SG_IQ3_XXS;
|
||||||
nr0 = N_R0_IQ3_XXS;
|
nr0 = N_R0_IQ3_XXS;
|
||||||
smem = 256*4+128;
|
smem = 256*4+128;
|
||||||
|
|
||||||
|
// split the rows across threads when there are fewer than 32 chunks per row
|
||||||
|
const int nb32 = ne00/32;
|
||||||
|
if (nb32 < 32 && (32 % nb32) == 0) {
|
||||||
|
nr0 = N_R0_IQ3_XXS_SPLIT;
|
||||||
|
split = true;
|
||||||
|
}
|
||||||
} break;
|
} break;
|
||||||
case GGML_TYPE_IQ3_S:
|
case GGML_TYPE_IQ3_S:
|
||||||
{
|
{
|
||||||
@@ -1224,7 +1243,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id(ggml_m
|
|||||||
};
|
};
|
||||||
|
|
||||||
snprintf(base, 256, "kernel_mul_mv_id_%s_%s%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1), suffix);
|
snprintf(base, 256, "kernel_mul_mv_id_%s_%s%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1), suffix);
|
||||||
snprintf(name, 256, "%s_nsg=%d", base, nsg);
|
snprintf(name, 256, "%s_nsg=%d_split=%d", base, nsg, split);
|
||||||
|
|
||||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||||
if (!res.pipeline) {
|
if (!res.pipeline) {
|
||||||
@@ -1234,6 +1253,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id(ggml_m
|
|||||||
ggml_metal_cv_set_int16(cv, 1, FC_MUL_MV + 2);
|
ggml_metal_cv_set_int16(cv, 1, FC_MUL_MV + 2);
|
||||||
ggml_metal_cv_set_int16(cv, 1, FC_MUL_MV + 3);
|
ggml_metal_cv_set_int16(cv, 1, FC_MUL_MV + 3);
|
||||||
ggml_metal_cv_set_int16(cv, 1, FC_MUL_MV + 4);
|
ggml_metal_cv_set_int16(cv, 1, FC_MUL_MV + 4);
|
||||||
|
ggml_metal_cv_set_bool (cv, split, FC_MUL_MV + 5);
|
||||||
|
|
||||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||||
|
|
||||||
|
|||||||
@@ -77,6 +77,7 @@
|
|||||||
|
|
||||||
#define N_R0_IQ3_XXS 4
|
#define N_R0_IQ3_XXS 4
|
||||||
#define N_SG_IQ3_XXS 2
|
#define N_SG_IQ3_XXS 2
|
||||||
|
#define N_R0_IQ3_XXS_SPLIT 8
|
||||||
|
|
||||||
#define N_R0_IQ3_S 4
|
#define N_R0_IQ3_S 4
|
||||||
#define N_SG_IQ3_S 2
|
#define N_SG_IQ3_S 2
|
||||||
|
|||||||
@@ -213,6 +213,7 @@ constant short FC_mul_mv_nxpsg [[function_constant(FC_MUL_MV + 1)]];
|
|||||||
constant short FC_mul_mv_ne12 [[function_constant(FC_MUL_MV + 2)]];
|
constant short FC_mul_mv_ne12 [[function_constant(FC_MUL_MV + 2)]];
|
||||||
constant short FC_mul_mv_r2 [[function_constant(FC_MUL_MV + 3)]];
|
constant short FC_mul_mv_r2 [[function_constant(FC_MUL_MV + 3)]];
|
||||||
constant short FC_mul_mv_r3 [[function_constant(FC_MUL_MV + 4)]];
|
constant short FC_mul_mv_r3 [[function_constant(FC_MUL_MV + 4)]];
|
||||||
|
constant bool FC_mul_mv_split [[function_constant(FC_MUL_MV + 5)]];
|
||||||
|
|
||||||
template<typename block_q_type, short NR0, typename args_t>
|
template<typename block_q_type, short NR0, typename args_t>
|
||||||
void mul_vec_q_n_f32_impl(
|
void mul_vec_q_n_f32_impl(
|
||||||
@@ -2092,6 +2093,7 @@ kernel void kernel_mul_mv_iq2_xs_f32(
|
|||||||
kernel_mul_mv_iq2_xs_f32_impl<N_R0_IQ2_XS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
kernel_mul_mv_iq2_xs_f32_impl<N_R0_IQ2_XS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// FC_mul_mv_split: for nb32 < 32 (nb32 divides 32), 32/nb32 threads share each chunk and each takes a slice of the rows
|
||||||
template<int nr0, typename args_t>
|
template<int nr0, typename args_t>
|
||||||
void kernel_mul_mv_iq3_xxs_f32_impl(
|
void kernel_mul_mv_iq3_xxs_f32_impl(
|
||||||
args_t args,
|
args_t args,
|
||||||
@@ -2138,11 +2140,18 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
|
|||||||
threadgroup_barrier(mem_flags::mem_threadgroup);
|
threadgroup_barrier(mem_flags::mem_threadgroup);
|
||||||
}
|
}
|
||||||
|
|
||||||
const int ix = tiisg;
|
const short ntx = FC_mul_mv_split ? nb32 : 32;
|
||||||
|
const short nrep = 32 / ntx;
|
||||||
|
|
||||||
|
const short ix = tiisg % ntx;
|
||||||
|
const short irep = tiisg / ntx;
|
||||||
|
|
||||||
|
const short row0 = (nr0 * irep ) / nrep;
|
||||||
|
const short row1 = (nr0 * (irep + 1)) / nrep;
|
||||||
|
|
||||||
device const float * y4 = y + 32 * ix;
|
device const float * y4 = y + 32 * ix;
|
||||||
|
|
||||||
for (int ib32 = ix; ib32 < nb32; ib32 += 32) {
|
for (int ib32 = ix; ib32 < nb32; ib32 += ntx) {
|
||||||
for (short i = 0; i < 32; ++i) {
|
for (short i = 0; i < 32; ++i) {
|
||||||
yl[i] = y4[i];
|
yl[i] = y4[i];
|
||||||
}
|
}
|
||||||
@@ -2151,11 +2160,11 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
|
|||||||
const int ib = ib32 % (QK_K / 32);
|
const int ib = ib32 % (QK_K / 32);
|
||||||
|
|
||||||
device const block_iq3_xxs * xr = x + ibl;
|
device const block_iq3_xxs * xr = x + ibl;
|
||||||
device const uint8_t * q3 = xr->qs + 8 * ib;
|
device const uint8_t * q3 = xr->qs + 8 * ib + (uint64_t) row0*args.nb01;
|
||||||
device const uint16_t * gas = (device const uint16_t *)(xr->qs + QK_K/4) + 2 * ib;
|
device const uint16_t * gas = (device const uint16_t *)(xr->qs + QK_K/4) + 2 * ib + (uint64_t) row0*args.nb01/2;
|
||||||
device const half * dh = &xr->d;
|
device const half * dh = &xr->d + (uint64_t) row0*args.nb01/2;
|
||||||
|
|
||||||
for (short row = 0; row < nr0; row++) {
|
for (short row = row0; row < row1; row++) {
|
||||||
const float db = dh[0];
|
const float db = dh[0];
|
||||||
const uint32_t aux32 = gas[0] | (gas[1] << 16);
|
const uint32_t aux32 = gas[0] | (gas[1] << 16);
|
||||||
const float d = db * (0.5f + (aux32 >> 28));
|
const float d = db * (0.5f + (aux32 >> 28));
|
||||||
@@ -2177,7 +2186,7 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
|
|||||||
gas += args.nb01/2;
|
gas += args.nb01/2;
|
||||||
}
|
}
|
||||||
|
|
||||||
y4 += 32 * 32;
|
y4 += 32 * ntx;
|
||||||
}
|
}
|
||||||
|
|
||||||
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
|
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
|
||||||
@@ -2190,6 +2199,23 @@ void kernel_mul_mv_iq3_xxs_f32_impl(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
template<typename args_t>
|
||||||
|
void kernel_mul_mv_iq3_xxs_f32_disp(
|
||||||
|
args_t args,
|
||||||
|
device const char * src0,
|
||||||
|
device const char * src1,
|
||||||
|
device char * dst,
|
||||||
|
threadgroup char * shmem,
|
||||||
|
uint3 tgpig,
|
||||||
|
ushort tiisg,
|
||||||
|
ushort sgitg) {
|
||||||
|
if (FC_mul_mv_split) {
|
||||||
|
kernel_mul_mv_iq3_xxs_f32_impl<N_R0_IQ3_XXS_SPLIT, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||||
|
} else {
|
||||||
|
kernel_mul_mv_iq3_xxs_f32_impl<N_R0_IQ3_XXS, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
[[host_name("kernel_mul_mv_iq3_xxs_f32")]]
|
[[host_name("kernel_mul_mv_iq3_xxs_f32")]]
|
||||||
kernel void kernel_mul_mv_iq3_xxs_f32(
|
kernel void kernel_mul_mv_iq3_xxs_f32(
|
||||||
constant ggml_metal_kargs_mul_mv & args,
|
constant ggml_metal_kargs_mul_mv & args,
|
||||||
@@ -2201,7 +2227,7 @@ kernel void kernel_mul_mv_iq3_xxs_f32(
|
|||||||
ushort tiisg[[thread_index_in_simdgroup]],
|
ushort tiisg[[thread_index_in_simdgroup]],
|
||||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||||
|
|
||||||
kernel_mul_mv_iq3_xxs_f32_impl<N_R0_IQ3_XXS, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
kernel_mul_mv_iq3_xxs_f32_disp<constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||||
}
|
}
|
||||||
|
|
||||||
template<int nr0, typename args_t>
|
template<int nr0, typename args_t>
|
||||||
@@ -3217,7 +3243,7 @@ template [[host_name("kernel_mul_mv_id_iq1_s_f32")]] kernel kernel_mul_mv_id_t
|
|||||||
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_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_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_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_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>>>;
|
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_disp<ggml_metal_kargs_mul_mv>>>;
|
||||||
template [[host_name("kernel_mul_mv_id_iq3_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq3_s_f32_impl <N_R0_IQ3_S>>>;
|
template [[host_name("kernel_mul_mv_id_iq3_s_f32")]] kernel kernel_mul_mv_id_t kernel_mul_mv_id<mmv_fn<kernel_mul_mv_iq3_s_f32_impl <N_R0_IQ3_S>>>;
|
||||||
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_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_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>>>;
|
||||||
|
|||||||
Reference in New Issue
Block a user