mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-17 20:31:47 +02:00
metal : add MUL_MAT + UNARY and MUL_MAT + ADD + UNARY fusion
Adds dense mat-vec activation fusion for sigmoid/silu and bias+softplus. The mat-vec kernels apply the activation/bias epilogue via function constants, avoiding the separate unary/add passes. Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp
This commit is contained in:
@@ -726,7 +726,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_solve_tri(ggml_m
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext(ggml_metal_library_t lib, const ggml_tensor * op, int nsg, int nxpsg, int r1ptg) {
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext(ggml_metal_library_t lib, const ggml_tensor * op, int nsg, int nxpsg, int r1ptg, int act_op, bool has_bias) {
|
||||
char base[256];
|
||||
char name[256];
|
||||
|
||||
@@ -739,7 +739,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext(ggml_
|
||||
GGML_ASSERT(ne12 <= INT16_MAX && r2 <= INT16_MAX && r3 <= INT16_MAX);
|
||||
|
||||
snprintf(base, 256, "kernel_mul_mv_ext_%s_%s_r1_%d", ggml_type_name(tsrc0), ggml_type_name(tsrc1), r1ptg);
|
||||
snprintf(name, 256, "%s_nsg=%d_nxpsg=%d_ne12=%d_r2=%d_r3=%d", base, nsg, nxpsg, ne12, r2, r3);
|
||||
snprintf(name, 256, "%s_nsg=%d_nxpsg=%d_ne12=%d_r2=%d_r3=%d_act=%d_bias=%d", base, nsg, nxpsg, ne12, r2, r3, act_op, has_bias ? 1 : 0);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
@@ -750,6 +750,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext(ggml_
|
||||
ggml_metal_cv_set_int16(cv, (int16_t) ne12, FC_MUL_MV + 2);
|
||||
ggml_metal_cv_set_int16(cv, (int16_t) r2, FC_MUL_MV + 3);
|
||||
ggml_metal_cv_set_int16(cv, (int16_t) r3, FC_MUL_MV + 4);
|
||||
ggml_metal_cv_set_int16(cv, (int16_t) act_op, FC_MUL_MV + 6);
|
||||
ggml_metal_cv_set_bool (cv, has_bias, FC_MUL_MV + 7);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
@@ -821,7 +823,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_meta
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_metal_library_t lib, const ggml_tensor * op) {
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_metal_library_t lib, const ggml_tensor * op, int act_op, bool has_bias) {
|
||||
GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
|
||||
GGML_TENSOR_LOCALS( int32_t, ne1, op->src[1], ne);
|
||||
|
||||
@@ -1038,7 +1040,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
|
||||
const int16_t r3 = (int16_t) (ne13 / ne03);
|
||||
|
||||
snprintf(base, 256, "kernel_mul_mv_%s_%s%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1), suffix);
|
||||
snprintf(name, 256, "%s_nsg=%d_ne12=%d_r2=%d_r3=%d_split=%d", base, nsg, ne12, r2, r3, split);
|
||||
snprintf(name, 256, "%s_nsg=%d_ne12=%d_r2=%d_r3=%d_split=%d_act=%d_bias=%d", base, nsg, ne12, r2, r3, split, act_op, has_bias ? 1 : 0);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
@@ -1049,6 +1051,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv(ggml_meta
|
||||
ggml_metal_cv_set_int16(cv, r2, FC_MUL_MV + 3);
|
||||
ggml_metal_cv_set_int16(cv, r3, FC_MUL_MV + 4);
|
||||
ggml_metal_cv_set_bool (cv, split, FC_MUL_MV + 5);
|
||||
ggml_metal_cv_set_int16(cv, (int16_t) act_op, FC_MUL_MV + 6);
|
||||
ggml_metal_cv_set_bool (cv, has_bias, FC_MUL_MV + 7);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
|
||||
@@ -134,9 +134,9 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan_
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_gated_delta_net (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_solve_tri (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext (ggml_metal_library_t lib, const struct ggml_tensor * op, int nsg, int nxpsg, int r1ptg);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext (ggml_metal_library_t lib, const struct ggml_tensor * op, int nsg, int nxpsg, int r1ptg, int act_op, bool has_bias);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv (ggml_metal_library_t lib, const struct ggml_tensor * op, int act_op, bool has_bias);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id_map0 (ggml_metal_library_t lib, int ne02, int ne20);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_id (ggml_metal_library_t lib, const struct ggml_tensor * op);
|
||||
|
||||
@@ -10,6 +10,14 @@
|
||||
|
||||
// ---- helpers -------------------------------------------------------------
|
||||
|
||||
// follow the view/reshape chain to the underlying tensor
|
||||
static const ggml_tensor * ggml_metal_fusion_view_src(const ggml_tensor * t) {
|
||||
while (t && t->view_src) {
|
||||
t = t->view_src;
|
||||
}
|
||||
return t;
|
||||
}
|
||||
|
||||
// true if two tensors live in the same Metal buffer
|
||||
static bool ggml_metal_fusion_same_buffer(const ggml_tensor * a, const ggml_tensor * b) {
|
||||
if (!a || !b) {
|
||||
@@ -78,6 +86,109 @@ static bool ggml_metal_fusion_check_norm(
|
||||
return true;
|
||||
}
|
||||
|
||||
// MUL_MAT + UNARY (sigmoid/silu): dense mat-vec output followed directly by a supported unary op
|
||||
static bool ggml_metal_fusion_check_mul_mat_unary(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_UNUSED(fusion);
|
||||
GGML_UNUSED(gf);
|
||||
GGML_UNUSED(node_idxs);
|
||||
GGML_UNUSED(idx);
|
||||
GGML_UNUSED(mode);
|
||||
|
||||
const ggml_tensor * mm = nodes[0];
|
||||
const ggml_tensor * un = nodes[1];
|
||||
|
||||
if (mm->op != GGML_OP_MUL_MAT || un->op != GGML_OP_UNARY || un->src[1]) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (ggml_metal_fusion_view_src(un->src[0]) != mm) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_unary_op un_op = ggml_get_unary_op(un);
|
||||
if (un_op != GGML_UNARY_OP_SIGMOID && un_op != GGML_UNARY_OP_SILU) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (mm->type != GGML_TYPE_F32 || un->type != GGML_TYPE_F32 || !ggml_is_contiguous_rows(un)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// only mat-vec paths for now
|
||||
if (mm->src[1]->type != GGML_TYPE_F32 || mm->src[1]->ne[1] > 8) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_type wt = mm->src[0]->type;
|
||||
if (wt != GGML_TYPE_F32 && wt != GGML_TYPE_F16 && wt != GGML_TYPE_BF16 && wt != GGML_TYPE_Q8_0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
// MUL_MAT + ADD (bias) + UNARY (softplus): dense mat-vec output with a 1D bias and softplus
|
||||
static bool ggml_metal_fusion_check_mul_mat_add_unary(
|
||||
const ggml_metal_fusion * fusion,
|
||||
const ggml_tensor * const * nodes,
|
||||
const ggml_cgraph * gf,
|
||||
const int * node_idxs,
|
||||
int idx,
|
||||
ggml_metal_fusion_mode mode) {
|
||||
GGML_UNUSED(fusion);
|
||||
GGML_UNUSED(gf);
|
||||
GGML_UNUSED(node_idxs);
|
||||
GGML_UNUSED(idx);
|
||||
GGML_UNUSED(mode);
|
||||
|
||||
const ggml_tensor * mm = nodes[0];
|
||||
const ggml_tensor * add = nodes[1];
|
||||
const ggml_tensor * un = nodes[2];
|
||||
const ggml_tensor * bias = add->src[1];
|
||||
|
||||
if (mm->op != GGML_OP_MUL_MAT || add->op != GGML_OP_ADD ||
|
||||
un->op != GGML_OP_UNARY || un->src[0] != add || !bias) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (ggml_metal_fusion_view_src(add->src[0]) != mm) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (ggml_get_unary_op(un) != GGML_UNARY_OP_SOFTPLUS) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (mm->type != GGML_TYPE_F32 || add->type != GGML_TYPE_F32 || un->type != GGML_TYPE_F32 ||
|
||||
!ggml_is_contiguous_rows(add) || !ggml_is_contiguous_rows(un)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (bias->type != GGML_TYPE_F32 || !ggml_is_contiguous(bias) ||
|
||||
bias->ne[1] != 1 || bias->ne[2] != 1 || bias->ne[3] != 1 ||
|
||||
bias->ne[0] != mm->ne[0]) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// only mat-vec paths for now
|
||||
if (mm->src[1]->type != GGML_TYPE_F32 || mm->src[1]->ne[1] > 8) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const ggml_type wt = mm->src[0]->type;
|
||||
if (wt != GGML_TYPE_F32 && wt != GGML_TYPE_F16 && wt != GGML_TYPE_BF16 && wt != GGML_TYPE_Q8_0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
// ADD x N: each ADD reads the previous ADD as src0, and all addends must share layout
|
||||
// (and, in FULL mode, live in the same Metal buffer)
|
||||
static bool ggml_metal_fusion_check_add_chain(
|
||||
@@ -568,6 +679,9 @@ static const ggml_op ops_topk_moe_norm_scale[] = {
|
||||
GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_SCALE
|
||||
};
|
||||
|
||||
static const ggml_op ops_mul_mat_unary[] = { GGML_OP_MUL_MAT, GGML_OP_UNARY };
|
||||
static const ggml_op ops_mul_mat_add_unary[] = { GGML_OP_MUL_MAT, GGML_OP_ADD, GGML_OP_UNARY };
|
||||
|
||||
static const ggml_op ops_moe_reduce_2[] = { GGML_OP_MUL, GGML_OP_ADD };
|
||||
static const ggml_op ops_moe_reduce_3[] = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD };
|
||||
static const ggml_op ops_moe_reduce_4[] = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
|
||||
@@ -602,6 +716,8 @@ static const ggml_metal_fusion ggml_metal_fusions[] = {
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_6, 6, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, 7, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, 8, true, ggml_metal_fusion_check_moe_reduce },
|
||||
{ GGML_METAL_FUSION_MUL_MAT_UNARY, ops_mul_mat_unary, 2, true, ggml_metal_fusion_check_mul_mat_unary },
|
||||
{ GGML_METAL_FUSION_MUL_MAT_ADD_UNARY, ops_mul_mat_add_unary, 3, true, ggml_metal_fusion_check_mul_mat_add_unary },
|
||||
};
|
||||
|
||||
const ggml_metal_fusion * ggml_metal_fusion_all(int * n) {
|
||||
|
||||
@@ -38,6 +38,8 @@ typedef enum ggml_metal_fusion_id {
|
||||
GGML_METAL_FUSION_GDN_CACHE, // GATED_DELTA_NET + CPY (write snapshots into the recurrent cache)
|
||||
GGML_METAL_FUSION_TOPK_MOE, // SOFT_MAX + ARGSORT + GET_ROWS + norm/scale (MoE routing)
|
||||
GGML_METAL_FUSION_MOE_REDUCE, // MUL + expert VIEWs + ADD chain (MoE output reduction)
|
||||
GGML_METAL_FUSION_MUL_MAT_UNARY, // MUL_MAT + UNARY (sigmoid/silu)
|
||||
GGML_METAL_FUSION_MUL_MAT_ADD_UNARY, // MUL_MAT + ADD (bias) + UNARY (softplus)
|
||||
} ggml_metal_fusion_id;
|
||||
|
||||
struct ggml_metal_fusion {
|
||||
|
||||
@@ -2405,6 +2405,47 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
|
||||
return ggml_metal_op_fwht(ctx, idx);
|
||||
}
|
||||
}
|
||||
|
||||
int n_fuse = 1;
|
||||
bool fused_mul_mat = false;
|
||||
bool fused_mul_mat_bias = false;
|
||||
int act_op = 0;
|
||||
ggml_tensor * bias = nullptr;
|
||||
|
||||
{
|
||||
int n = 1;
|
||||
const ggml_metal_fusion * fusion = ctx->can_fuse(idx, GGML_METAL_FUSION_FULL, &n);
|
||||
|
||||
if (fusion && (fusion->id == GGML_METAL_FUSION_MUL_MAT_UNARY ||
|
||||
fusion->id == GGML_METAL_FUSION_MUL_MAT_ADD_UNARY)) {
|
||||
n_fuse = n;
|
||||
fused_mul_mat = true;
|
||||
|
||||
ctx->count_fusions(fusion);
|
||||
|
||||
if (ggml_metal_fusion_info_debug(ctx->finfo) > 1) {
|
||||
if (fusion->id == GGML_METAL_FUSION_MUL_MAT_ADD_UNARY) {
|
||||
GGML_LOG_DEBUG("%s: fuse: MUL_MAT + ADD + UNARY\n", __func__);
|
||||
} else {
|
||||
GGML_LOG_DEBUG("%s: fuse: MUL_MAT + UNARY\n", __func__);
|
||||
}
|
||||
}
|
||||
|
||||
if (fusion->id == GGML_METAL_FUSION_MUL_MAT_ADD_UNARY) {
|
||||
fused_mul_mat_bias = true;
|
||||
bias = ctx->node(idx + 1)->src[1];
|
||||
}
|
||||
|
||||
const ggml_unary_op un_op = ggml_get_unary_op(ctx->node(idx + n_fuse - 1));
|
||||
switch (un_op) {
|
||||
case GGML_UNARY_OP_SIGMOID: act_op = 1; break;
|
||||
case GGML_UNARY_OP_SILU: act_op = 2; break;
|
||||
case GGML_UNARY_OP_SOFTPLUS: act_op = 3; break;
|
||||
default: GGML_ABORT("fatal error");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const ggml_metal_device_props * props_dev = ggml_metal_device_get_props(ctx->dev);
|
||||
|
||||
GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne);
|
||||
@@ -2494,7 +2535,7 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
|
||||
GGML_ABORT("unsupported ne11");
|
||||
};
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mv_ext(lib, op, nsg, nxpsg, r1ptg);
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mv_ext(lib, op, nsg, nxpsg, r1ptg, act_op, fused_mul_mat_bias);
|
||||
|
||||
ggml_metal_kargs_mul_mv_ext args = {
|
||||
/*.ne00 =*/ ne00,
|
||||
@@ -2521,7 +2562,8 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), 1);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 2);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 3);
|
||||
ggml_metal_encoder_set_buffer (enc, fused_mul_mat ? ggml_metal_get_buffer_id(ctx->node(idx + n_fuse - 1)) : ggml_metal_get_buffer_id(op), 3);
|
||||
ggml_metal_encoder_set_buffer (enc, fused_mul_mat_bias ? ggml_metal_get_buffer_id(bias) : ggml_metal_get_buffer_id(op), 4);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, ((ne01 + r0ptg - 1)/r0ptg), ((ne11 + r1ptg - 1)/r1ptg), ne12*ne13, 32, nsg, 1);
|
||||
} else if (ggml_metal_op_mul_mat_use_mm(op, props_dev->has_simdgroup_mm)) {
|
||||
@@ -2571,7 +2613,7 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, ((ne11 + nr1 - 1) / nr1), ((ne01 + nr0 - 1) / nr0), ne12 * ne13, 32, nsg, 1);
|
||||
} else {
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mv(lib, op);
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mv(lib, op, act_op, fused_mul_mat_bias);
|
||||
|
||||
const int nr0 = pipeline.nr0;
|
||||
const int nr1 = pipeline.nr1;
|
||||
@@ -2605,7 +2647,8 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), 1);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 2);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 3);
|
||||
ggml_metal_encoder_set_buffer (enc, fused_mul_mat ? ggml_metal_get_buffer_id(ctx->node(idx + n_fuse - 1)) : ggml_metal_get_buffer_id(op), 3);
|
||||
ggml_metal_encoder_set_buffer (enc, fused_mul_mat_bias ? ggml_metal_get_buffer_id(bias) : ggml_metal_get_buffer_id(op), 4);
|
||||
|
||||
ggml_metal_encoder_set_threadgroup_memory_size(enc, smem, 0);
|
||||
|
||||
@@ -2619,7 +2662,7 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
|
||||
}
|
||||
}
|
||||
|
||||
return 1;
|
||||
return n_fuse;
|
||||
}
|
||||
|
||||
size_t ggml_metal_op_mul_mat_id_extra_tpe(const ggml_tensor * op) {
|
||||
|
||||
@@ -1,5 +1,28 @@
|
||||
#include "common.h"
|
||||
#include "dequantize.h"
|
||||
|
||||
constant short FC_mul_mv_nsg [[function_constant(FC_MUL_MV + 0)]];
|
||||
constant short FC_mul_mv_nxpsg [[function_constant(FC_MUL_MV + 1)]];
|
||||
constant short FC_mul_mv_ne12 [[function_constant(FC_MUL_MV + 2)]];
|
||||
constant short FC_mul_mv_r2 [[function_constant(FC_MUL_MV + 3)]];
|
||||
constant short FC_mul_mv_r3 [[function_constant(FC_MUL_MV + 4)]];
|
||||
constant bool FC_mul_mv_split [[function_constant(FC_MUL_MV + 5)]];
|
||||
constant short FC_mul_mv_act [[function_constant(FC_MUL_MV + 6)]];
|
||||
constant bool FC_mul_mv_has_bias [[function_constant(FC_MUL_MV + 7)]];
|
||||
|
||||
static inline float mul_mv_epilogue(float tot, device const float * bias, int row) {
|
||||
if (FC_mul_mv_has_bias) {
|
||||
tot += bias[row];
|
||||
}
|
||||
|
||||
switch (FC_mul_mv_act) {
|
||||
case 1: return 1.0f / (1.0f + exp(-tot)); // SIGMOID
|
||||
case 2: return tot / (1.0f + exp(-tot)); // SILU
|
||||
case 3: return max(tot, 0.0f) + log(1.0f + exp(-fabs(tot))); // SOFTPLUS
|
||||
default: return tot;
|
||||
}
|
||||
}
|
||||
|
||||
// Q1_0 dot product: dot = d * (2 * Σ(yl[i] where bit=1) - sumy)
|
||||
inline float block_q_n_dot_y(device const block_q1_0 * qb_curr, float sumy, thread float * yl, int il) {
|
||||
device const uint8_t * qs = qb_curr->qs + il / 8;
|
||||
@@ -174,7 +197,8 @@ static inline void helper_mv_reduce_and_write(
|
||||
const int ne01,
|
||||
ushort tiisg,
|
||||
ushort sgitg,
|
||||
threadgroup char * shmem) {
|
||||
threadgroup char * shmem,
|
||||
device const float * bias) {
|
||||
constexpr short NW = N_SIMDWIDTH;
|
||||
|
||||
threadgroup float * shmem_f32[NR0];
|
||||
@@ -203,18 +227,11 @@ static inline void helper_mv_reduce_and_write(
|
||||
float tot = simd_sum(shmem_f32[row][tiisg]);
|
||||
|
||||
if (tiisg == 0 && sgitg == 0) {
|
||||
dst_f32[r0 + row] = tot;
|
||||
dst_f32[r0 + row] = mul_mv_epilogue(tot, bias, r0 + row);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
constant short FC_mul_mv_nsg [[function_constant(FC_MUL_MV + 0)]];
|
||||
constant short FC_mul_mv_nxpsg [[function_constant(FC_MUL_MV + 1)]];
|
||||
constant short FC_mul_mv_ne12 [[function_constant(FC_MUL_MV + 2)]];
|
||||
constant short FC_mul_mv_r2 [[function_constant(FC_MUL_MV + 3)]];
|
||||
constant short FC_mul_mv_r3 [[function_constant(FC_MUL_MV + 4)]];
|
||||
constant bool FC_mul_mv_split [[function_constant(FC_MUL_MV + 5)]];
|
||||
|
||||
template<typename block_q_type, short NR0, typename args_t>
|
||||
void mul_vec_q_n_f32_impl(
|
||||
args_t args,
|
||||
@@ -519,7 +536,8 @@ void kernel_mul_mv_q8_0_f32_impl(
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiisg,
|
||||
ushort sgitg) {
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
const short NSG = FC_mul_mv_nsg;
|
||||
|
||||
constexpr short NW = N_SIMDWIDTH;
|
||||
@@ -581,7 +599,7 @@ void kernel_mul_mv_q8_0_f32_impl(
|
||||
|
||||
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
|
||||
|
||||
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01, tiisg, sgitg, shmem);
|
||||
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01, tiisg, sgitg, shmem, bias);
|
||||
}
|
||||
|
||||
[[host_name("kernel_mul_mv_q8_0_f32")]]
|
||||
@@ -590,11 +608,12 @@ kernel void kernel_mul_mv_q8_0_f32(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
threadgroup char * shmem [[threadgroup(0)]],
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||
kernel_mul_mv_q8_0_f32_impl<N_R0_Q8_0, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||
kernel_mul_mv_q8_0_f32_impl<N_R0_Q8_0, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias);
|
||||
}
|
||||
|
||||
// mat-vec kernel processing in chunks of float4
|
||||
@@ -605,6 +624,7 @@ void kernel_mul_mv_ext_q4_f32_impl(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||
@@ -695,7 +715,7 @@ void kernel_mul_mv_ext_q4_f32_impl(
|
||||
device float * dst_f32 = (device float *) dst + (uint64_t)i1m*args.ne0*args.ne1 + (uint64_t)(i11 + ir1)*args.ne0;
|
||||
|
||||
if (i01 < args.ne01) {
|
||||
dst_f32[i01] = sumf[ir1];
|
||||
dst_f32[i01] = mul_mv_epilogue(sumf[ir1], bias, i01);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -708,6 +728,7 @@ void kernel_mul_mv_ext_q4x4_f32_impl(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||
@@ -802,7 +823,7 @@ void kernel_mul_mv_ext_q4x4_f32_impl(
|
||||
device float * dst_f32 = (device float *) dst + (uint64_t)i1m*args.ne0*args.ne1 + (uint64_t)(i11 + ir1)*args.ne0;
|
||||
|
||||
if (i01 < args.ne01) {
|
||||
dst_f32[i01] = sumf[ir1];
|
||||
dst_f32[i01] = mul_mv_epilogue(sumf[ir1], bias, i01);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -816,10 +837,11 @@ kernel void kernel_mul_mv_ext_q4_f32_disp(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||
kernel_mul_mv_ext_q4_f32_impl<r1ptg, q_t, epb/4, deq_t4>(args, src0, src1, dst, tgpig, tiisg, sgitg);
|
||||
kernel_mul_mv_ext_q4_f32_impl<r1ptg, q_t, epb/4, deq_t4>(args, src0, src1, dst, bias, tgpig, tiisg, sgitg);
|
||||
}
|
||||
|
||||
template<short r1ptg, typename q_t, short epb, void (*deq_t4x4)(device const q_t *, short, thread float4x4 &)>
|
||||
@@ -828,10 +850,11 @@ kernel void kernel_mul_mv_ext_q4x4_f32_disp(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||
kernel_mul_mv_ext_q4x4_f32_impl<r1ptg, q_t, epb/16, deq_t4x4>(args, src0, src1, dst, tgpig, tiisg, sgitg);
|
||||
kernel_mul_mv_ext_q4x4_f32_impl<r1ptg, q_t, epb/16, deq_t4x4>(args, src0, src1, dst, bias, tgpig, tiisg, sgitg);
|
||||
}
|
||||
|
||||
typedef decltype(kernel_mul_mv_ext_q4_f32_disp <2, block_q8_0, 32, dequantize_q8_0_t4>) mul_mv_ext_q4_f32_t;
|
||||
@@ -933,7 +956,8 @@ void kernel_mul_mv_t_t_impl(
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiisg,
|
||||
ushort sgitg) {
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
const short NSG = FC_mul_mv_nsg;
|
||||
|
||||
constexpr short NW = N_SIMDWIDTH;
|
||||
@@ -1001,7 +1025,7 @@ void kernel_mul_mv_t_t_impl(
|
||||
|
||||
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
|
||||
|
||||
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01, tiisg, sgitg, shmem);
|
||||
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01, tiisg, sgitg, shmem, bias);
|
||||
}
|
||||
|
||||
template<typename T0, typename T1, typename args_t>
|
||||
@@ -1013,12 +1037,13 @@ void kernel_mul_mv_t_t_disp(
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiisg,
|
||||
ushort sgitg) {
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
switch (args.nr0) {
|
||||
//case 1: kernel_mul_mv_t_t_impl<T0, T1, 1, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
case 2: kernel_mul_mv_t_t_impl<T0, T1, 2, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
//case 3: kernel_mul_mv_t_t_impl<T0, T1, 3, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
//case 4: kernel_mul_mv_t_t_impl<T0, T1, 4, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
//case 1: kernel_mul_mv_t_t_impl<T0, T1, 1, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
case 2: kernel_mul_mv_t_t_impl<T0, T1, 2, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
//case 3: kernel_mul_mv_t_t_impl<T0, T1, 3, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
//case 4: kernel_mul_mv_t_t_impl<T0, T1, 4, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1028,11 +1053,12 @@ kernel void kernel_mul_mv_t_t(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
threadgroup char * shmem [[threadgroup(0)]],
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||
kernel_mul_mv_t_t_disp<T0, T1, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||
kernel_mul_mv_t_t_disp<T0, T1, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias);
|
||||
}
|
||||
|
||||
typedef decltype(kernel_mul_mv_t_t<half, half>) mul_mv_t_t;
|
||||
@@ -1054,7 +1080,8 @@ void kernel_mul_mv_t_t_4_impl(
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiisg,
|
||||
ushort sgitg) {
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
const short NSG = FC_mul_mv_nsg;
|
||||
|
||||
constexpr short NW = N_SIMDWIDTH;
|
||||
@@ -1125,7 +1152,7 @@ void kernel_mul_mv_t_t_4_impl(
|
||||
|
||||
device float * dst_f32 = (device float *) dst + (uint64_t)im*args.ne0*args.ne1 + (uint64_t)r1*args.ne0;
|
||||
|
||||
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01, tiisg, sgitg, shmem);
|
||||
helper_mv_reduce_and_write<NR0>(dst_f32, sumf, r0, args.ne01, tiisg, sgitg, shmem, bias);
|
||||
}
|
||||
|
||||
template<typename T0, typename T04, typename T1, typename T14, typename args_t>
|
||||
@@ -1137,12 +1164,13 @@ void kernel_mul_mv_t_t_4_disp(
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiisg,
|
||||
ushort sgitg) {
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
switch (args.nr0) {
|
||||
//case 1: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 1, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
case 2: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 2, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
//case 3: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 3, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
//case 4: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 4, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg); break;
|
||||
//case 1: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 1, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
case 2: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 2, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
//case 3: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 3, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
//case 4: kernel_mul_mv_t_t_4_impl<T0, T04, T1, T14, 4, args_t>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias); break;
|
||||
};
|
||||
}
|
||||
|
||||
@@ -1152,11 +1180,12 @@ kernel void kernel_mul_mv_t_t_4(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
threadgroup char * shmem [[threadgroup(0)]],
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]],
|
||||
ushort sgitg[[simdgroup_index_in_threadgroup]]) {
|
||||
kernel_mul_mv_t_t_4_disp<T0, T04, T1, T14, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||
kernel_mul_mv_t_t_4_disp<T0, T04, T1, T14, constant ggml_metal_kargs_mul_mv &>(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias);
|
||||
}
|
||||
|
||||
typedef decltype(kernel_mul_mv_t_t_4<half, half4, half, half4>) mul_mv_t_t_4;
|
||||
@@ -1176,7 +1205,8 @@ void kernel_mul_mv_t_t_short_impl(
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
uint3 tgpig,
|
||||
ushort tiisg) {
|
||||
ushort tiisg,
|
||||
device const float * bias) {
|
||||
const int r0 = tgpig.x*32 + tiisg;
|
||||
const int r1 = tgpig.y;
|
||||
const int im = tgpig.z;
|
||||
@@ -1204,7 +1234,7 @@ void kernel_mul_mv_t_t_short_impl(
|
||||
res += (float) x[i] * (float) y[i];
|
||||
}
|
||||
|
||||
dst_f32[(uint64_t)r1*args.ne0 + r0] = res;
|
||||
dst_f32[(uint64_t)r1*args.ne0 + r0] = mul_mv_epilogue(res, bias, r0);
|
||||
}
|
||||
|
||||
template<typename T0, typename T1>
|
||||
@@ -1213,6 +1243,7 @@ kernel void kernel_mul_mv_t_t_short(
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
device const float * bias,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
ushort tiisg[[thread_index_in_simdgroup]]) {
|
||||
kernel_mul_mv_t_t_short_impl<T0, T1, constant ggml_metal_kargs_mul_mv &>(
|
||||
@@ -1221,7 +1252,8 @@ kernel void kernel_mul_mv_t_t_short(
|
||||
src1,
|
||||
dst,
|
||||
tgpig,
|
||||
tiisg);
|
||||
tiisg,
|
||||
bias);
|
||||
}
|
||||
|
||||
typedef decltype(kernel_mul_mv_t_t_short<half, half>) mul_mv_t_t_short_t;
|
||||
@@ -3246,7 +3278,8 @@ typedef void (kernel_mul_mv_disp_t)(
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
uint3 tgpig,
|
||||
ushort tiisg);
|
||||
ushort tiisg,
|
||||
device const float * bias);
|
||||
|
||||
typedef void (kernel_mul_mv2_disp_t)(
|
||||
ggml_metal_kargs_mul_mv args,
|
||||
@@ -3256,7 +3289,8 @@ typedef void (kernel_mul_mv2_disp_t)(
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiisg,
|
||||
ushort sgitg);
|
||||
ushort sgitg,
|
||||
device const float * bias);
|
||||
|
||||
template<kernel_mul_mv_disp_t disp_fn>
|
||||
void mmv_fn(
|
||||
@@ -3268,8 +3302,9 @@ void mmv_fn(
|
||||
uint3 tgpig,
|
||||
ushort tiitg,
|
||||
ushort tiisg,
|
||||
ushort sgitg) {
|
||||
disp_fn(args, src0, src1, dst, tgpig, tiisg);
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
disp_fn(args, src0, src1, dst, tgpig, tiisg, bias);
|
||||
}
|
||||
|
||||
template<kernel_mul_mv2_disp_t disp_fn>
|
||||
@@ -3282,7 +3317,33 @@ void mmv_fn(
|
||||
uint3 tgpig,
|
||||
ushort tiitg,
|
||||
ushort tiisg,
|
||||
ushort sgitg) {
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
disp_fn(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg, bias);
|
||||
}
|
||||
|
||||
typedef void (kernel_mul_mv2_disp_no_bias_t)(
|
||||
ggml_metal_kargs_mul_mv args,
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiisg,
|
||||
ushort sgitg);
|
||||
|
||||
template<kernel_mul_mv2_disp_no_bias_t disp_fn>
|
||||
void mmv_fn(
|
||||
ggml_metal_kargs_mul_mv args,
|
||||
device const char * src0,
|
||||
device const char * src1,
|
||||
device char * dst,
|
||||
threadgroup char * shmem,
|
||||
uint3 tgpig,
|
||||
ushort tiitg,
|
||||
ushort tiisg,
|
||||
ushort sgitg,
|
||||
device const float * bias) {
|
||||
disp_fn(args, src0, src1, dst, shmem, tgpig, tiisg, sgitg);
|
||||
}
|
||||
|
||||
@@ -3349,7 +3410,8 @@ kernel void kernel_mul_mv_id(
|
||||
tgpig,
|
||||
tiitg,
|
||||
tiisg,
|
||||
sgitg);
|
||||
sgitg,
|
||||
/* bias */ nullptr);
|
||||
}
|
||||
|
||||
typedef decltype(kernel_mul_mv_id<mmv_fn<kernel_mul_mv_t_t_disp<float, float>>>) kernel_mul_mv_id_t;
|
||||
|
||||
@@ -11,10 +11,12 @@ bailingmoe ,1 ,any ,RMS_NORM+MUL , 5
|
||||
bailingmoe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 2
|
||||
bailingmoe2 ,1 ,any ,ADD+ADD , 1
|
||||
bailingmoe2 ,1 ,any ,MUL+ADD , 2
|
||||
bailingmoe2 ,1 ,decode ,MUL_MAT+UNARY , 1
|
||||
bailingmoe2 ,1 ,any ,RMS_NORM+MUL , 9
|
||||
bailingmoe3 ,1 ,any ,ADD+ADD , 1
|
||||
bailingmoe3 ,1 ,any ,GATED_DELTA_NET+CPY , 1
|
||||
bailingmoe3 ,1 ,any ,MUL+ADD , 2
|
||||
bailingmoe3 ,1 ,decode ,MUL_MAT+UNARY , 5
|
||||
bailingmoe3 ,1 ,any ,RMS_NORM+MUL , 8
|
||||
bailingmoe3 ,1 ,any ,RMS_NORM+SCALE , 2
|
||||
bloom ,0 ,any ,NORM+MUL+ADD , 6
|
||||
@@ -32,15 +34,18 @@ deepseek ,0 ,any ,RMS_NORM+MUL , 5
|
||||
deepseek ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 1
|
||||
deepseek2 ,0 ,any ,ADD+ADD , 1
|
||||
deepseek2 ,0 ,any ,MUL+ADD , 2
|
||||
deepseek2 ,0 ,decode ,MUL_MAT+UNARY , 1
|
||||
deepseek2 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
deepseek32 ,0 ,any ,ADD+ADD , 1
|
||||
deepseek32 ,0 ,any ,MUL+ADD , 2
|
||||
deepseek32 ,0 ,decode ,MUL_MAT+UNARY , 1
|
||||
deepseek32 ,0 ,any ,NORM+MUL+ADD , 2
|
||||
deepseek32 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
deepseek4 ,0 ,any ,MUL+ADD , 8
|
||||
deepseek4 ,0 ,any ,RMS_NORM+MUL , 20
|
||||
dots1 ,0 ,any ,ADD+ADD , 1
|
||||
dots1 ,0 ,any ,MUL+ADD , 2
|
||||
dots1 ,0 ,decode ,MUL_MAT+UNARY , 1
|
||||
dots1 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
dream ,0 ,any ,RMS_NORM+MUL , 5
|
||||
ernie4_5-moe ,1 ,any ,ADD+ADD , 1
|
||||
@@ -59,12 +64,14 @@ gemma2 ,0 ,any ,RMS_NORM+MUL , 5
|
||||
gemma2 ,0 ,any ,RMS_NORM+MUL+ADD , 4
|
||||
glm-dsa ,0 ,any ,ADD+ADD , 1
|
||||
glm-dsa ,0 ,any ,MUL+ADD , 2
|
||||
glm-dsa ,0 ,decode ,MUL_MAT+UNARY , 1
|
||||
glm-dsa ,0 ,any ,NORM+MUL+ADD , 2
|
||||
glm-dsa ,0 ,any ,RMS_NORM+MUL , 9
|
||||
glm4 ,0 ,any ,RMS_NORM+MUL , 5
|
||||
glm4 ,0 ,any ,RMS_NORM+MUL+ADD , 4
|
||||
glm4moe ,1 ,any ,ADD+ADD , 1
|
||||
glm4moe ,1 ,any ,MUL+ADD , 2
|
||||
glm4moe ,1 ,decode ,MUL_MAT+UNARY , 1
|
||||
glm4moe ,1 ,any ,RMS_NORM+MUL , 9
|
||||
gpt-oss ,0 ,any ,MUL+ADD , 4
|
||||
gpt-oss ,0 ,any ,RMS_NORM+MUL , 5
|
||||
@@ -94,8 +101,10 @@ hunyuan-moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV,
|
||||
hunyuan_vl ,0 ,any ,RMS_NORM+MUL , 9
|
||||
hy_v3 ,0 ,any ,ADD+ADD , 2
|
||||
hy_v3 ,0 ,any ,MUL+ADD , 4
|
||||
hy_v3 ,0 ,decode ,MUL_MAT+UNARY , 2
|
||||
hy_v3 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
hy_v4 ,0 ,any ,MUL+ADD , 2
|
||||
hy_v4 ,0 ,decode ,MUL_MAT+UNARY , 3
|
||||
hy_v4 ,0 ,any ,NORM+MUL+ADD , 1
|
||||
hy_v4 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
internlm2 ,0 ,any ,RMS_NORM+MUL , 5
|
||||
@@ -104,15 +113,19 @@ jais2 ,0 ,any ,NORM+MUL+ADD , 5
|
||||
jamba ,0 ,any ,RMS_NORM+MUL , 8
|
||||
kimi-k3 ,0 ,any ,GATED_DELTA_NET+CPY , 1
|
||||
kimi-k3 ,0 ,any ,MUL+ADD , 2
|
||||
kimi-k3 ,0 ,decode ,MUL_MAT+UNARY , 4
|
||||
kimi-k3 ,0 ,any ,RMS_NORM+MUL , 17
|
||||
kimi-k3 ,0 ,any ,RMS_NORM+SCALE , 2
|
||||
kimi-linear ,0 ,any ,ADD+ADD , 1
|
||||
kimi-linear ,0 ,any ,GATED_DELTA_NET+CPY , 1
|
||||
kimi-linear ,0 ,any ,MUL+ADD , 2
|
||||
kimi-linear ,0 ,decode ,MUL_MAT+ADD+UNARY , 1
|
||||
kimi-linear ,0 ,decode ,MUL_MAT+UNARY , 3
|
||||
kimi-linear ,0 ,any ,RMS_NORM+MUL , 7
|
||||
kimi-linear ,0 ,any ,RMS_NORM+SCALE , 2
|
||||
lfm2 ,0 ,any ,RMS_NORM+MUL , 7
|
||||
lfm2moe ,1 ,any ,MUL+ADD , 2
|
||||
lfm2moe ,1 ,decode ,MUL_MAT+UNARY , 1
|
||||
lfm2moe ,1 ,any ,RMS_NORM+MUL , 7
|
||||
llada ,0 ,any ,RMS_NORM+MUL , 5
|
||||
llada-moe ,1 ,any ,MUL+ADD , 4
|
||||
@@ -123,6 +136,7 @@ llama ,0 ,any ,MUL+ADD , 4
|
||||
llama ,0 ,any ,RMS_NORM+MUL , 5
|
||||
llama ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
llama4 ,0 ,any ,ADD+ADD , 2
|
||||
llama4 ,0 ,decode ,MUL_MAT+UNARY , 2
|
||||
llama4 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
maincoder ,0 ,any ,RMS_NORM+MUL , 9
|
||||
mamba ,0 ,any ,RMS_NORM+MUL , 3
|
||||
@@ -133,12 +147,15 @@ minicpm ,0 ,any ,RMS_NORM+MUL , 5
|
||||
minicpm ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
minicpm3 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
minimax-01 ,0 ,any ,MUL+ADD , 4
|
||||
minimax-01 ,0 ,decode ,MUL_MAT+UNARY , 2
|
||||
minimax-01 ,0 ,any ,RMS_NORM+MUL , 6
|
||||
minimax-01 ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
minimax-m2 ,0 ,any ,MUL+ADD , 4
|
||||
minimax-m2 ,0 ,decode ,MUL_MAT+UNARY , 2
|
||||
minimax-m2 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
minimax-m3 ,0 ,any ,ADD+ADD , 1
|
||||
minimax-m3 ,0 ,any ,MUL+ADD , 2
|
||||
minimax-m3 ,0 ,decode ,MUL_MAT+UNARY , 1
|
||||
minimax-m3 ,0 ,any ,RMS_NORM+MUL , 11
|
||||
mistral3 ,0 ,any ,RMS_NORM+MUL , 5
|
||||
mistral3 ,0 ,any ,MUL+ADD , 4
|
||||
@@ -146,6 +163,7 @@ mistral3 ,0 ,any ,RMS_NORM+MUL , 5
|
||||
mistral3 ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
mistral4 ,0 ,any ,ADD+ADD , 1
|
||||
mistral4 ,0 ,any ,MUL+ADD , 2
|
||||
mistral4 ,0 ,decode ,MUL_MAT+UNARY , 1
|
||||
mistral4 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
mpt ,0 ,any ,NORM+MUL+ADD , 5
|
||||
nanbeige ,0 ,any ,RMS_NORM+MUL , 5
|
||||
@@ -174,6 +192,7 @@ qwen ,0 ,any ,RMS_NORM+MUL , 5
|
||||
qwen2 ,0 ,any ,RMS_NORM+MUL , 5
|
||||
qwen2moe ,1 ,any ,ADD+ADD , 2
|
||||
qwen2moe ,1 ,any ,MUL+ADD , 4
|
||||
qwen2moe ,1 ,decode ,MUL_MAT+UNARY , 2
|
||||
qwen2moe ,1 ,any ,RMS_NORM+MUL , 5
|
||||
qwen2moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS , 2
|
||||
qwen2vl ,0 ,any ,RMS_NORM+MUL , 5
|
||||
@@ -184,6 +203,7 @@ qwen35 ,0 ,any ,RMS_NORM+SCALE , 2
|
||||
qwen35moe ,1 ,any ,ADD+ADD , 2
|
||||
qwen35moe ,1 ,any ,GATED_DELTA_NET+CPY , 1
|
||||
qwen35moe ,1 ,any ,MUL+ADD , 4
|
||||
qwen35moe ,1 ,decode ,MUL_MAT+UNARY , 2
|
||||
qwen35moe ,1 ,any ,RMS_NORM+MUL , 8
|
||||
qwen35moe ,1 ,any ,RMS_NORM+SCALE , 2
|
||||
qwen35moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
@@ -193,6 +213,7 @@ qwen3moe ,1 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV,
|
||||
qwen3next ,0 ,any ,ADD+ADD , 2
|
||||
qwen3next ,0 ,any ,GATED_DELTA_NET+CPY , 1
|
||||
qwen3next ,0 ,any ,MUL+ADD , 4
|
||||
qwen3next ,0 ,decode ,MUL_MAT+UNARY , 3
|
||||
qwen3next ,0 ,any ,RMS_NORM+MUL , 8
|
||||
qwen3next ,0 ,any ,RMS_NORM+SCALE , 2
|
||||
qwen3next ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
@@ -205,6 +226,7 @@ qwen4exp ,0 ,any ,ADD+ADD+ADD , 5
|
||||
qwen4exp ,0 ,any ,ADD+ADD+ADD+ADD+ADD+ADD+ADD , 9
|
||||
qwen4exp ,0 ,any ,GATED_DELTA_NET+CPY , 1
|
||||
qwen4exp ,0 ,any ,MUL+ADD , 4
|
||||
qwen4exp ,0 ,decode ,MUL_MAT+UNARY , 7
|
||||
qwen4exp ,0 ,any ,RMS_NORM+MUL , 13
|
||||
qwen4exp ,0 ,any ,RMS_NORM+SCALE , 2
|
||||
qwen4exp ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
@@ -215,6 +237,7 @@ rnd1 ,0 ,any ,RMS_NORM+MUL , 9
|
||||
rnd1 ,0 ,any ,SOFT_MAX+ARGSORT+GET_ROWS+SUM_ROWS+CLAMP+DIV, 2
|
||||
seed_oss ,0 ,any ,RMS_NORM+MUL , 5
|
||||
smallthinker ,0 ,any ,MUL+ADD , 4
|
||||
smallthinker ,0 ,decode ,MUL_MAT+UNARY , 1
|
||||
smallthinker ,0 ,any ,RMS_NORM+MUL , 5
|
||||
smollm3 ,0 ,any ,RMS_NORM+MUL , 5
|
||||
stablelm ,0 ,any ,NORM+MUL , 4
|
||||
|
||||
|
@@ -6890,6 +6890,96 @@ struct test_moe_reduce : public test_case {
|
||||
}
|
||||
};
|
||||
|
||||
// GGML_OP_MUL_MAT + GGML_OP_UNARY (sigmoid/silu)
|
||||
struct test_mul_mat_unary : public test_case {
|
||||
const ggml_type type;
|
||||
const ggml_unary_op un_op;
|
||||
const int64_t m;
|
||||
const int64_t n;
|
||||
const int64_t k;
|
||||
|
||||
std::string vars() override {
|
||||
return VARS_TO_STR5(type, un_op, m, n, k);
|
||||
}
|
||||
|
||||
test_mul_mat_unary(ggml_type type, ggml_unary_op un_op, int64_t m = 64, int64_t n = 4, int64_t k = 128)
|
||||
: type(type), un_op(un_op), m(m), n(n), k(k) {}
|
||||
|
||||
double max_nmse_err() override { return 5e-4; }
|
||||
|
||||
std::string op_desc(ggml_tensor * t) override {
|
||||
GGML_UNUSED(t);
|
||||
return "MUL_MAT_UNARY";
|
||||
}
|
||||
|
||||
bool run_whole_graph() override { return true; }
|
||||
|
||||
ggml_tensor * build_graph(ggml_context * ctx) override {
|
||||
ggml_tensor * w = ggml_new_tensor_2d(ctx, type, k, m);
|
||||
ggml_tensor * a = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, k, n);
|
||||
ggml_set_name(w, "w");
|
||||
ggml_set_name(a, "a");
|
||||
|
||||
ggml_tensor * out = ggml_mul_mat(ctx, w, a);
|
||||
out = un_op == GGML_UNARY_OP_SIGMOID ? ggml_sigmoid(ctx, out) : ggml_silu(ctx, out);
|
||||
ggml_set_name(out, "out");
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
void initialize_tensors(ggml_context * ctx) override {
|
||||
for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) {
|
||||
init_tensor_uniform(t);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// GGML_OP_MUL_MAT + GGML_OP_ADD (bias) + GGML_OP_UNARY (softplus)
|
||||
struct test_mul_mat_add_unary : public test_case {
|
||||
const ggml_type type;
|
||||
const int64_t m;
|
||||
const int64_t n;
|
||||
const int64_t k;
|
||||
|
||||
std::string vars() override {
|
||||
return VARS_TO_STR4(type, m, n, k);
|
||||
}
|
||||
|
||||
test_mul_mat_add_unary(ggml_type type, int64_t m = 64, int64_t n = 4, int64_t k = 128)
|
||||
: type(type), m(m), n(n), k(k) {}
|
||||
|
||||
double max_nmse_err() override { return 5e-4; }
|
||||
|
||||
std::string op_desc(ggml_tensor * t) override {
|
||||
GGML_UNUSED(t);
|
||||
return "MUL_MAT_ADD_UNARY";
|
||||
}
|
||||
|
||||
bool run_whole_graph() override { return true; }
|
||||
|
||||
ggml_tensor * build_graph(ggml_context * ctx) override {
|
||||
ggml_tensor * w = ggml_new_tensor_2d(ctx, type, k, m);
|
||||
ggml_tensor * a = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, k, n);
|
||||
ggml_tensor * b = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, m);
|
||||
ggml_set_name(w, "w");
|
||||
ggml_set_name(a, "a");
|
||||
ggml_set_name(b, "b");
|
||||
|
||||
ggml_tensor * out = ggml_mul_mat(ctx, w, a);
|
||||
out = ggml_add(ctx, out, b);
|
||||
out = ggml_softplus(ctx, out);
|
||||
ggml_set_name(out, "out");
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
void initialize_tensors(ggml_context * ctx) override {
|
||||
for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) {
|
||||
init_tensor_uniform(t);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
struct test_mul_mat_vec_fusion : public test_case {
|
||||
const ggml_type type;
|
||||
const ggml_glu_op glu_op;
|
||||
@@ -10872,6 +10962,14 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
|
||||
false, 16, 8, false, false, true, false, { 1, 1 }));
|
||||
}
|
||||
|
||||
for (ggml_type type : { GGML_TYPE_F32, GGML_TYPE_Q8_0 }) {
|
||||
for (int64_t n : { 1, 4 }) {
|
||||
test_cases.emplace_back(new test_mul_mat_unary(type, GGML_UNARY_OP_SIGMOID, 64, n, 128));
|
||||
test_cases.emplace_back(new test_mul_mat_unary(type, GGML_UNARY_OP_SILU, 64, n, 128));
|
||||
test_cases.emplace_back(new test_mul_mat_add_unary(type, 64, n, 128));
|
||||
}
|
||||
}
|
||||
|
||||
for (auto gate : {GATING_FUNC_SOFTMAX, GATING_FUNC_SIGMOID, GATING_FUNC_SOFTMAX_WEIGHT, GATING_FUNC_SQRT_SOFTPLUS}) {
|
||||
for (bool with_norm : {false, true}) {
|
||||
for (bool bias_probs : {false, true}) {
|
||||
|
||||
Reference in New Issue
Block a user