diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index a9038f1f4d..c8b23b2d5c 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -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)) { diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 2d38875884..dc6b695eb0 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -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); diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 29db37f873..527892e18a 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -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); diff --git a/ggml/src/ggml-metal/kernels/mul_mm.metal b/ggml/src/ggml-metal/kernels/mul_mm.metal index 71d9911491..a25838f926 100644 --- a/ggml/src/ggml-metal/kernels/mul_mm.metal +++ b/ggml/src/ggml-metal/kernels/mul_mm.metal @@ -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]; diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index da0e24fcd9..f7e27703f2 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -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: diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 86f0e958a5..9c5dc768ef 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -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: { diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 02ae8bd92a..07ca49ad09 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2302,6 +2302,10 @@ ggml_tensor * llm_graph_context::build_moe_ffn( } experts = build_lora_mm_id(down_exps, cur, selected_experts, down_exps_s); // [n_embd, n_expert_used, n_tokens] + if (arch == LLM_ARCH_MISTRAL4) { + // src1 can exceed F16 range + ggml_prec_set_src(experts, GGML_PREC_F32, 1); + } cb(experts, "ffn_moe_down", il); if (down_exps_s) { diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 5f030c4068..505112f6be 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -5186,6 +5186,11 @@ struct test_mul_mat_id : public test_case { ggml_tensor * out = ggml_mul_mat_id(ctx, as, b, ids); ggml_set_name(out, "out"); + if (amax > 65504.0f) { + // src1 exceeds F16 range + ggml_prec_set_src(out, GGML_PREC_F32, 1); + } + return out; } @@ -10185,11 +10190,10 @@ static std::vector> make_test_cases_eval() { } // test src1 f16 overflow - // TODO: https://github.com/ggml-org/llama.cpp/pull/26223#issuecomment-5585815365 - //for (int n : {16, 32, 64}) { - // test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q4_K, GGML_TYPE_F32, 128, 4, false, 4096, n, 2048, 1e5f)); - // test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q8_0, GGML_TYPE_F32, 8, 2, false, 512, n, 256, 1e5f)); - //} + for (int n : {16, 32, 64}) { + test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q4_K, GGML_TYPE_F32, 128, 4, false, 4096, n, 2048, 1e5f)); + test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q8_0, GGML_TYPE_F32, 8, 2, false, 512, n, 256, 1e5f)); + } for (ggml_type type_a : base_types) { for (ggml_type type_b : {GGML_TYPE_F32 /*, GGML_TYPE_F16 */}) {