mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-09-29 09:27:44 -05:00
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:
@@ -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;
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user