diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 9657e7edb..2d3887588 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -497,25 +497,21 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexe } ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_metal_library_t lib, const ggml_tensor * op) { - const char * name = nullptr; + char name[256]; + const char * base = nullptr; switch (op->op) { case GGML_OP_DSV4_HC_COMB: - name = "kernel_dsv4_hc_comb_f32"; + base = "kernel_dsv4_hc_comb_f32"; + snprintf(name, 256, "%s", base); break; case GGML_OP_DSV4_HC_PRE: - if (ggml_get_op_params_i32(op, 1) != 0) { - name = "kernel_dsv4_hc_pre_gated_f32"; - } else { - name = "kernel_dsv4_hc_pre_f32"; - } + base = ggml_get_op_params_i32(op, 1) != 0 ? "kernel_dsv4_hc_pre_gated_f32" : "kernel_dsv4_hc_pre_f32"; + snprintf(name, 256, "%s_n_hc=%d", base, (int) op->src[0]->ne[1]); break; case GGML_OP_DSV4_HC_POST: - if (op->src[3]) { - name = "kernel_dsv4_hc_post_f32"; - } else { - name = "kernel_dsv4_hc_post_nocomb_f32"; - } + base = op->src[3] ? "kernel_dsv4_hc_post_f32" : "kernel_dsv4_hc_post_nocomb_f32"; + snprintf(name, 256, "%s", base); break; default: GGML_ABORT("fatal error"); @@ -523,7 +519,18 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_met ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); if (!res.pipeline) { - res = ggml_metal_library_compile_pipeline(lib, name, name, nullptr); + ggml_metal_cv_t cv = nullptr; + + if (op->op == GGML_OP_DSV4_HC_PRE) { + cv = ggml_metal_cv_init(); + ggml_metal_cv_set_int32(cv, (int32_t) op->src[0]->ne[1], FC_DSV4_HC + 0); + } + + res = ggml_metal_library_compile_pipeline(lib, base, name, cv); + + if (cv) { + ggml_metal_cv_free(cv); + } } return res; diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 9650de268..9c2afbd9c 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1807,7 +1807,6 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && - op->src[0]->ne[1] == 4 && ggml_is_contiguous_rows(op->src[0]) && ggml_is_contiguous_rows(op->src[1]); case GGML_OP_DSV4_HC_POST: diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index eaa4278db..490dd83a1 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -119,6 +119,7 @@ #define FC_NORM 1700 #define FC_TOPK_MOE 1800 #define FC_MOE_REDUCE 1900 +#define FC_DSV4_HC 2000 // op-specific constants #define OP_FLASH_ATTN_EXT_NQPSG 8 diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 0323dc386..29db37f87 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1466,7 +1466,6 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) { GGML_ASSERT(x->type == GGML_TYPE_F32); GGML_ASSERT(weights->type == GGML_TYPE_F32); GGML_ASSERT(op->type == GGML_TYPE_F32); - GGML_ASSERT(x->ne[1] == 4); ggml_metal_kargs_dsv4_hc_pre args = { /*.n_embd =*/ (int32_t) x->ne[0], diff --git a/ggml/src/ggml-metal/kernels/misc.metal b/ggml/src/ggml-metal/kernels/misc.metal index 877ccf2e1..279d69f8f 100644 --- a/ggml/src/ggml-metal/kernels/misc.metal +++ b/ggml/src/ggml-metal/kernels/misc.metal @@ -442,6 +442,8 @@ template [[host_name("kernel_fwht_f16_128")]] kernel kernel_fwht_f16_t kernel_fw template [[host_name("kernel_fwht_f16_256")]] kernel kernel_fwht_f16_t kernel_fwht<256, half>; template [[host_name("kernel_fwht_f16_512")]] kernel kernel_fwht_f16_t kernel_fwht<512, half>; +constant int FC_dsv4_hc_n_hc [[function_constant(FC_DSV4_HC + 0)]]; + kernel void kernel_dsv4_hc_comb_f32( constant ggml_metal_kargs_dsv4_hc_comb & args, device const char * mixes, @@ -512,29 +514,19 @@ kernel void kernel_dsv4_hc_pre_f32( ushort tiisg[[thread_index_in_simdgroup]], ushort sgitg[[simdgroup_index_in_threadgroup]], ushort3 ntg[[threads_per_threadgroup]]) { - constexpr ushort hc = 4; - const int it = tgpig.y; const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg; - float weight_lane = 0.0f; - if (tiisg < hc) { - weight_lane = *(device const float *) (weights + tiisg*args.nb_w0 + it*args.nb_w1); - } - - float w[hc]; - FOR_UNROLL (ushort ih = 0; ih < hc; ++ih) { - w[ih] = simd_shuffle(weight_lane, ih); - } - if (i0 >= args.n_embd) { return; } device const char * xb = x + i0*args.nb_x0 + it*args.nb_x2; float result = 0.0f; - FOR_UNROLL (ushort ih = 0; ih < hc; ++ih) { - result = fma(*(device const float *) (xb + ih*args.nb_x1), w[ih], result); + FOR_UNROLL (int ih = 0; ih < FC_dsv4_hc_n_hc; ++ih) { + const float xv = *(device const float *) (xb + ih*args.nb_x1); + const float wv = *(device const float *) (weights + ih*args.nb_w0 + it*args.nb_w1); + result = fma(xv, wv, result); } *(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = args.scale*result; @@ -549,8 +541,6 @@ kernel void kernel_dsv4_hc_pre_gated_f32( ushort tiisg[[thread_index_in_simdgroup]], ushort sgitg[[simdgroup_index_in_threadgroup]], ushort3 ntg[[threads_per_threadgroup]]) { - constexpr ushort hc = 4; - const int it = tgpig.y; const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg; @@ -561,9 +551,10 @@ kernel void kernel_dsv4_hc_pre_gated_f32( device const char * xb = x + i0*args.nb_x0 + it*args.nb_x2; device const char * gb = gate + i0*args.nb_w0 + it*args.nb_w2; float result = 0.0f; - FOR_UNROLL (ushort ih = 0; ih < hc; ++ih) { - const float g = 1.0f/(1.0f + exp(-*(device const float *) (gb + ih*args.nb_w1))); - result = fma(*(device const float *) (xb + ih*args.nb_x1), g, result); + FOR_UNROLL (int ih = 0; ih < FC_dsv4_hc_n_hc; ++ih) { + const float g = 1.0f/(1.0f + exp(-*(device const float *) (gb + ih*args.nb_w1))); + const float xv = *(device const float *) (xb + ih*args.nb_x1); + result = fma(xv, g, result); } *(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = args.scale*result;