Commit a894dae93 for llama.cpp
commit a894dae939d426954ce54bb604824f1ae918a0c5
Author: Georgi Gerganov <ggerganov@gmail.com>
Date: Sun Sep 20 17:52:20 2026 +0300
metal : support arbitrary hc in dsv4_hc_pre (#29169)
the dsv4_hc_pre kernels hardcoded hc = 4 via a constexpr used with
simd_shuffle, so the op was rejected by supports_op for any other hc
and fell back to CPU. Kimi-K3 uses dsv4_hc_pre with hc equal to the
number of banked checkpoints in the cross-layer residual stack, which
grows with the layer index.
pass n_hc as a function constant (FC_DSV4_HC) with per-n_hc pipeline
variants, and loop over it in both pre kernels with direct loads
add test-backend-ops cases for hc = 1, 2, 3, 5, 8 and 65, gated and
not gated
Assisted-by: pi:llama.cpp/Qwen3.8-27B
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;
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index e4af4299e..c75cb3c0f 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -4263,6 +4263,7 @@ struct test_dsv4_hc_comb : public test_dsv4_hc {
struct test_dsv4_hc_pre : public test_dsv4_hc {
const int64_t n_embd;
+ const int64_t n_hc;
const int64_t n_tokens;
const bool gated;
@@ -4272,23 +4273,23 @@ struct test_dsv4_hc_pre : public test_dsv4_hc {
}
std::string vars() override {
- return VARS_TO_STR3(n_embd, n_tokens, gated);
+ return VARS_TO_STR4(n_embd, n_hc, n_tokens, gated);
}
- test_dsv4_hc_pre(int64_t n_embd = 31, int64_t n_tokens = 17, bool gated = false)
- : n_embd(n_embd), n_tokens(n_tokens), gated(gated) {}
+ test_dsv4_hc_pre(int64_t n_embd = 31, int64_t n_hc = 4, int64_t n_tokens = 17, bool gated = false)
+ : n_embd(n_embd), n_hc(n_hc), n_tokens(n_tokens), gated(gated) {}
ggml_tensor * build_graph(ggml_context * ctx) override {
- ggml_tensor * x = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens);
+ ggml_tensor * x = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, n_hc, n_tokens);
ggml_set_name(x, "x");
if (gated) {
- ggml_tensor * gate = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens);
+ ggml_tensor * gate = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, n_hc, n_tokens);
ggml_set_name(gate, "gate");
- out = ggml_dsv4_hc_pre_gated(ctx, x, gate, 1.0f/hc);
+ out = ggml_dsv4_hc_pre_gated(ctx, x, gate, 1.0f/n_hc);
} else {
- ggml_tensor * weights = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens);
+ ggml_tensor * weights = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_hc, n_tokens);
ggml_set_name(weights, "weights");
out = ggml_dsv4_hc_pre(ctx, x, weights);
@@ -9011,12 +9012,16 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_dsv4_hc_comb(n_tokens, 20));
}
- test_cases.emplace_back(new test_dsv4_hc_pre(1, 1));
- test_cases.emplace_back(new test_dsv4_hc_pre(31, 17));
- test_cases.emplace_back(new test_dsv4_hc_pre(128, 257));
- test_cases.emplace_back(new test_dsv4_hc_pre(4096, 21));
- test_cases.emplace_back(new test_dsv4_hc_pre(31, 17, true));
- test_cases.emplace_back(new test_dsv4_hc_pre(4096, 21, true));
+ test_cases.emplace_back(new test_dsv4_hc_pre(1, 4, 1));
+ test_cases.emplace_back(new test_dsv4_hc_pre(31, 4, 17));
+ test_cases.emplace_back(new test_dsv4_hc_pre(128, 4, 257));
+ test_cases.emplace_back(new test_dsv4_hc_pre(4096, 4, 21));
+ test_cases.emplace_back(new test_dsv4_hc_pre(31, 4, 17, true));
+ test_cases.emplace_back(new test_dsv4_hc_pre(4096, 4, 21, true));
+ for (int64_t n_hc : {1, 2, 3, 5, 8, 65}) {
+ test_cases.emplace_back(new test_dsv4_hc_pre(128, n_hc, 17));
+ test_cases.emplace_back(new test_dsv4_hc_pre(128, n_hc, 17, true));
+ }
test_cases.emplace_back(new test_dsv4_hc_post(1, 1));
test_cases.emplace_back(new test_dsv4_hc_post(31, 17));