metal : support arbitrary hc in dsv4_hc_pre (llama/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
This commit is contained in:
Georgi Gerganov
2026-09-23 20:46:47 +03:00
parent 3d949a3610
commit 984e400cc0
5 changed files with 31 additions and 34 deletions
+20 -13
View File
@@ -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;
-1
View File
@@ -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:
+1
View File
@@ -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
-1
View File
@@ -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],
+10 -19
View File
@@ -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;