mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-10-03 11:27:31 -05:00
metal : gate mul_mm_id src1 rescale behind ggml_prec (llama/29029)
* metal : gate mul_mm_id src1 rescale behind ggml_prec Assisted-by: Claude Fable 5.1 * ggml-webgpu: reject MUL_MAT_ID when src1 precision is F32 * cuda/vulkan: reject MUL_MAT_ID in supports_op when src1 prec is F32 fix `supports_op` to return false for failing backends when the specified src1 precision is f32 Assisted-by: Claude Fable 5.1 --------- Co-authored-by: yomaytk <yoshimura.masashi.frbs@gmail.com>
This commit is contained in:
committed by
Georgi Gerganov
co-authored by
yomaytk
parent
ff565cf322
commit
ad0058be6e
@@ -5131,6 +5131,9 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
if (b->type == GGML_TYPE_F16 && a->type != GGML_TYPE_F16) {
|
||||
return false;
|
||||
}
|
||||
if (op->op == GGML_OP_MUL_MAT_ID && ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
|
||||
return false;
|
||||
}
|
||||
#ifdef GGML_USE_MUSA
|
||||
const int cc = ggml_cuda_info().devices[dev_ctx->device].cc;
|
||||
if (b->ne[2]*b->ne[3] > 1 && !ggml_is_transposed(a) && !ggml_is_transposed(b)) {
|
||||
|
||||
@@ -1156,14 +1156,18 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id(ggml_m
|
||||
|
||||
const bool bc_inp = op->src[0]->ne[0] % 32 != 0;
|
||||
|
||||
// src1 prec [TAG_GGML_PREC]
|
||||
const bool amax = ggml_get_op_params_i32(op, 3) == GGML_PREC_F32;
|
||||
|
||||
snprintf(base, 256, "kernel_mul_mm_id_%s_%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1));
|
||||
snprintf(name, 256, "%s_bci=%d", base, bc_inp);
|
||||
snprintf(name, 256, "%s_bci=%d_amax=%d", base, bc_inp, amax);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
|
||||
ggml_metal_cv_set_bool(cv, bc_inp, FC_MUL_MM + 0);
|
||||
ggml_metal_cv_set_bool(cv, amax, FC_MUL_MM + 6);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
|
||||
@@ -2719,9 +2719,12 @@ int ggml_metal_op_mul_mat_id(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_buffer_id bid_amax = bid_ids;
|
||||
bid_amax.offs += ggml_metal_op_mul_mat_id_extra_ids(op);
|
||||
|
||||
// src1 prec [TAG_GGML_PREC]
|
||||
const bool use_amax = ggml_get_op_params_i32(op, 3) == GGML_PREC_F32;
|
||||
|
||||
// src1 rescale factors, computed before the matmul
|
||||
// ref: https://github.com/ggml-org/llama.cpp/pull/26223
|
||||
{
|
||||
if (use_amax) {
|
||||
ggml_metal_kargs_mul_mm_id_amax args = {
|
||||
/*.ne00 =*/ ne10,
|
||||
/*.ne01 =*/ ne11,
|
||||
@@ -2779,17 +2782,17 @@ int ggml_metal_op_mul_mat_id(ggml_metal_op_t ctx, int idx) {
|
||||
|
||||
ggml_metal_op_concurrency_reset(ctx);
|
||||
|
||||
{
|
||||
if (use_amax) {
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mm_id_amax(lib);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_buffer (enc, bid_amax, 0);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, 1, 1, 1, 32, 1, 1);
|
||||
}
|
||||
|
||||
// the next kernel has to wait for the amax data
|
||||
ggml_metal_op_concurrency_reset(ctx);
|
||||
// the next kernel has to wait for the amax data
|
||||
ggml_metal_op_concurrency_reset(ctx);
|
||||
}
|
||||
|
||||
{
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mm_id(lib, op);
|
||||
|
||||
@@ -7,6 +7,7 @@ constant short FC_mul_mm_ne12 [[function_constant(FC_MUL_MM + 2)]];
|
||||
constant short FC_mul_mm_ne13 [[function_constant(FC_MUL_MM + 3)]];
|
||||
constant short FC_mul_mm_r2 [[function_constant(FC_MUL_MM + 4)]];
|
||||
constant short FC_mul_mm_r3 [[function_constant(FC_MUL_MM + 5)]];
|
||||
constant bool FC_mul_mm_id_amax [[function_constant(FC_MUL_MM + 6)]];
|
||||
|
||||
// each block_q contains 16*nl weights
|
||||
#ifdef GGML_METAL_HAS_TENSOR
|
||||
@@ -584,8 +585,8 @@ kernel void kernel_mul_mm_id(
|
||||
const short lb1 = (short) tiitg/NL1; // 0 .. NR1-1, this thread's row of the B tile
|
||||
|
||||
// power-of-two rescaling
|
||||
const float s1_inv = ((device const float *) amax)[0];
|
||||
const float s1_scale = ((device const float *) amax)[1];
|
||||
const float s1_inv = FC_mul_mm_id_amax ? ((device const float *) amax)[0] : 1.0f;
|
||||
const float s1_scale = FC_mul_mm_id_amax ? ((device const float *) amax)[1] : 1.0f;
|
||||
|
||||
#ifndef GGML_METAL_HAS_TENSOR
|
||||
S0_8x8 ma[4];
|
||||
|
||||
@@ -14953,6 +14953,9 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
|
||||
// If there's not enough shared memory for row_ids and the result tile, fallback to CPU
|
||||
return false;
|
||||
}
|
||||
if (ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
switch (src0_type) {
|
||||
case GGML_TYPE_F32:
|
||||
|
||||
@@ -4506,6 +4506,9 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
||||
default:
|
||||
break;
|
||||
}
|
||||
if (ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
|
||||
supports_op = false;
|
||||
}
|
||||
break;
|
||||
case GGML_OP_FLASH_ATTN_EXT:
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user