diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.cpp index 7f8da06d173a..9d97e3161f84 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.cpp +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat.cpp @@ -1107,7 +1107,7 @@ static bool match_q6_k_token1_final_projection_q8_dispatch(const DispatchMatchCo static bool match_mul_mat_postops_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { const MulMatPostOpsMatch match = match_mul_mat_postops(context); - if (!match.matched()) { + if (!match.matched() || match.input_size % 256 != 0) { return false; } diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp index 0c981eb81657..cfda7ae003a0 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp @@ -15,9 +15,13 @@ #include "dispatch-small-rows.h" #include "ggml.h" +#include "graph/graph-matcher.h" #include "kernel-corpus/kernel-corpus-catalog-verify.h" #include +#include +#include +#include #include namespace ggml::hrx { @@ -28,6 +32,19 @@ static constexpr KernelCatalogRef kSumRowsKernel = GGML_HRX_KERNEL_REF("loom static constexpr KernelCatalogRef kArgsortRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_argsort_rows_f32"); static constexpr KernelCatalogRef kGetRowsSmallKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_get_rows_small_f32"); static constexpr KernelCatalogRef kCopyStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_strided_f32"); +static constexpr KernelCatalogRef kNormRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_norm_rows_f32"); +static constexpr KernelCatalogRef kBinaryStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_binary_strided_f32"); +static constexpr KernelCatalogRef kClampKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_clamp_f32"); +static constexpr KernelCatalogRef kClampInplaceKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_clamp_inplace_f32"); +static constexpr KernelCatalogRef kCopyF32F16Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_copy_strided_f32_f16"); +static constexpr KernelCatalogRef kAttentionStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_attention_strided_f32_f16"); +static constexpr KernelCatalogRef kAttentionRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_attention_rows_f32_f16"); +static constexpr KernelCatalogRef kRopeRotateHalfKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_rope_rotate_half_f32"); +static constexpr KernelCatalogRef kGegluStridedKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_geglu_strided_f32"); +static constexpr KernelCatalogRef kMulMatSmallF16Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_small_f16_f32"); +static constexpr KernelCatalogRef kMulMatSmallF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_small_f32_f32"); +static constexpr KernelCatalogRef kMulMatSmallQ8Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_small_q8_0_f32"); +static constexpr KernelCatalogRef kMulMatRowsQ8Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_rows_q8_0_f32"); static bool packed(const Value & value, size_t element_size) { size_t stride = element_size; @@ -133,6 +150,32 @@ static bool match_argsort_rows(const DispatchMatchContext & context, DispatchMat return true; } +// NORM (LayerNorm without affine; its weight and bias are separate MUL/ADD nodes) on packed F32 rows +static bool match_norm_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_NORM, GGML_TYPE_F32, input, output) || !same_shape(*input, *output) || + input->ne[0] > 65536) { + return false; + } + const RmsNormParams * params = op_params_as(context.root_node->params); + if (params == nullptr) { + return false; + } + std::ostringstream eps; + eps.precision(9); + eps << params->eps; + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kNormRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + dispatch.kernel.compile_parameters.emplace("ggml.norm_rows_f32.epsilon", eps.str()); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + // GET_ROWS within batches: source [W, S, B], ids [R, B] (I32, rows may be strided, as the first k // columns of an ARGSORT are), output [W, R, B]. Only rows the // regular get_rows kernel does not take (narrower than 4 floats, or not a multiple of 4). @@ -255,11 +298,425 @@ static bool match_repeat_broadcast(const DispatchMatchContext & context, Dispatc return true; } +// ADD / SUB / MUL / DIV of F32 values with any element strides, broadcasting either input (ggml's +// rule: an input dim of 1 against a larger output dim), into a packed output. Registered below the +// packed binary kernels (priority -10), so it only takes what they refuse: strided views mostly. +static bool match_binary_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->inputs.size() != 2) { + return false; + } + int64_t op = -1; + switch (node->op) { + case GGML_OP_ADD: op = 0; break; + case GGML_OP_SUB: op = 1; break; + case GGML_OP_MUL: op = 2; break; + case GGML_OP_DIV: op = 3; break; + default: return false; + } + const Value * lhs = context.graph.values().find(node->inputs[0]); + const Value * rhs = context.graph.values().find(node->inputs[1]); + const Value * output = context.graph.values().find(node->output); + if (lhs == nullptr || rhs == nullptr || output == nullptr || lhs->type != GGML_TYPE_F32 || rhs->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || + lhs->storage == output->storage || rhs->storage == output->storage || output->element_count > 268435456) { + return false; + } + int64_t a[GGML_MAX_DIMS], b[GGML_MAX_DIMS]; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + for (const Value * v : { lhs, rhs }) { + if (v->nb[i] % sizeof(float) != 0 || (v->ne[i] != 1 && v->ne[i] != output->ne[i])) { + return false; + } + } + a[i] = lhs->ne[i] == 1 ? 0 : static_cast(lhs->nb[i] / sizeof(float)); + b[i] = rhs->ne[i] == 1 ? 0 : static_cast(rhs->nb[i] / sizeof(float)); + } + const int64_t a_extent = static_cast(lhs->byte_count / sizeof(float)); + const int64_t b_extent = static_cast(rhs->byte_count / sizeof(float)); + if (a_extent < 1 || b_extent < 1 || a_extent > 268435456 || b_extent > 268435456) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kBinaryStridedKernel); + static const char * const ne_names[] = { "ne0", "ne1", "ne2", "ne3" }; + static const char * const a_names[] = { "a0", "a1", "a2", "a3" }; + static const char * const b_names[] = { "b0", "b1", "b2", "b3" }; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.integer_parameters.emplace(ne_names[i], output->ne[i]); + dispatch.kernel.integer_parameters.emplace(a_names[i], a[i]); + dispatch.kernel.integer_parameters.emplace(b_names[i], b[i]); + } + dispatch.kernel.integer_parameters.emplace("a_extent", a_extent); + dispatch.kernel.integer_parameters.emplace("b_extent", b_extent); + dispatch.kernel.integer_parameters.emplace("op", op); + dispatch.bindings.push_back({ lhs->id, 0, lhs->byte_count }); + dispatch.bindings.push_back({ rhs->id, 0, rhs->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +static std::string f32_config(float value) { + std::ostringstream out; + out.precision(9); + out << value; + return out.str(); +} + +// CLAMP of a packed F32 tensor on its own (fused router clamps match first, at their own priority) +static bool match_clamp(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CLAMP || node->inputs.size() != 1) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * output = context.graph.values().find(node->output); + const ClampParams * params = op_params_as(node->params); + if (input == nullptr || output == nullptr || params == nullptr || input->type != GGML_TYPE_F32 || + output->type != GGML_TYPE_F32 || !packed(*input, sizeof(float)) || !packed(*output, sizeof(float)) || + !same_shape(*input, *output) || output->element_count > 268435456) { + return false; + } + // ggml_clamp is in place: the output is a view of the input, same layout + const bool in_place = output->alias_source.value >= 0; + if (in_place ? (output->storage != input->storage || output->storage_offset != input->storage_offset) + : input->storage == output->storage) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(in_place ? kClampInplaceKernel : kClampKernel); + dispatch.kernel.integer_parameters.emplace("element_count", output->element_count); + dispatch.kernel.compile_parameters.emplace("ggml.clamp_f32.min", f32_config(params->min)); + dispatch.kernel.compile_parameters.emplace("ggml.clamp_f32.max", f32_config(params->max)); + if (!in_place) { + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + } + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +// CPY F32 (any element strides) -> packed F16: ggml_cpy(src, dst) returns a view of dst, so the +// output aliases the second input, which only gives the destination's layout +static bool match_copy_f32_f16(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CPY || node->inputs.empty() || node->inputs.size() > 2) { + return false; + } + const Value * input = context.graph.values().find(node->inputs[0]); + const Value * output = context.graph.values().find(node->output); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F16 || + !packed(*output, ggml_type_size(GGML_TYPE_F16)) || input->storage == output->storage || + input->element_count != output->element_count || output->element_count > 268435456) { + return false; + } + int64_t strides[GGML_MAX_DIMS]; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (input->nb[i] % sizeof(float) != 0 || input->ne[i] != output->ne[i]) { + return false; + } + strides[i] = static_cast(input->nb[i] / sizeof(float)); + } + const int64_t extent = static_cast(input->byte_count / sizeof(float)); + if (extent < 1 || extent > 268435456) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kCopyF32F16Kernel); + static const char * const ne_names[] = { "ne0", "ne1", "ne2", "ne3" }; + static const char * const s_names[] = { "s0", "s1", "s2", "s3" }; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + dispatch.kernel.integer_parameters.emplace(ne_names[i], output->ne[i]); + dispatch.kernel.integer_parameters.emplace(s_names[i], strides[i]); + } + dispatch.kernel.integer_parameters.emplace("source_extent", extent); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + +// FLASH_ATTN_EXT with the layouts the flash-attention kernels refuse (one contiguous block per head, +// as encoders lay them out): F32 query [d, n_q, h], F16 key/value [d, n_kv, h_kv], F16 mask +// [n_kv, >= n_q], output [dv, h, n_q] packed; no ALiBi, no softcap. Registered below them. +static bool match_attention_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_FLASH_ATTN_EXT || node->inputs.size() != 4) { + return false; + } + const Value * q = context.graph.values().find(node->inputs[0]); + const Value * k = context.graph.values().find(node->inputs[1]); + const Value * v = context.graph.values().find(node->inputs[2]); + const Value * mask = context.graph.values().find(node->inputs[3]); + const Value * out = context.graph.values().find(node->output); + const FlashAttnExtParams * params = op_params_as(node->params); + if (q == nullptr || k == nullptr || v == nullptr || mask == nullptr || out == nullptr || params == nullptr || + q->type != GGML_TYPE_F32 || k->type != GGML_TYPE_F16 || v->type != GGML_TYPE_F16 || mask->type != GGML_TYPE_F16 || + out->type != GGML_TYPE_F32 || params->max_bias != 0.0f || params->logit_softcap != 0.0f) { + return false; + } + const int64_t d = q->ne[0], dv = v->ne[0], nq = q->ne[1], nkv = k->ne[1], nh = q->ne[2], nhkv = k->ne[2]; + if (k->ne[0] != d || v->ne[1] != nkv || v->ne[2] != nhkv || nhkv < 1 || nh % nhkv != 0 || d > 1024 || dv > 1024 || + dv > 1024 || q->ne[3] != 1 || k->ne[3] != 1 || v->ne[3] != 1 || mask->ne[0] < nkv || mask->ne[1] < nq || + mask->ne[2] != 1 || mask->ne[3] != 1 || out->ne[0] != dv || out->ne[1] != nh || out->ne[2] != nq || + out->ne[3] != 1 || !packed(*out, sizeof(float)) || q->nb[0] != sizeof(float) || k->nb[0] != 2 || v->nb[0] != 2 || + mask->nb[0] != 2 || q->nb[1] % 4 || q->nb[2] % 4 || k->nb[1] % 2 || k->nb[2] % 2 || v->nb[1] % 2 || v->nb[2] % 2 || + mask->nb[1] % 2 || nq > 65536 || nkv > 65536 || nh > 1024 || dv < 1 || dv > 1024 || (dv % 32) != 0) { + return false; + } + Dispatch dispatch; + // scores once per (query, head) in workgroup memory when they fit (ONEBIT_HRX_ATTN_PER_LANE=1: the old way) + const bool rows = nkv <= 2048 && d <= 256 && dv <= 256 && std::getenv("ONEBIT_HRX_ATTN_PER_LANE") == nullptr; + dispatch.kernel = make_kernel_specialization(rows ? kAttentionRowsKernel : kAttentionStridedKernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("qk_size", d); + ip.emplace("v_size", dv); + ip.emplace("q_count", nq); + ip.emplace("kv_count", nkv); + ip.emplace("head_count", nh); + ip.emplace("kv_head_count", nhkv); + ip.emplace("q_s1", static_cast(q->nb[1] / 4)); + ip.emplace("q_s2", static_cast(q->nb[2] / 4)); + ip.emplace("k_s1", static_cast(k->nb[1] / 2)); + ip.emplace("k_s2", static_cast(k->nb[2] / 2)); + ip.emplace("v_s1", static_cast(v->nb[1] / 2)); + ip.emplace("v_s2", static_cast(v->nb[2] / 2)); + ip.emplace("m_s1", static_cast(mask->nb[1] / 2)); + ip.emplace("q_extent", static_cast(q->byte_count / 4)); + ip.emplace("k_extent", static_cast(k->byte_count / 2)); + ip.emplace("v_extent", static_cast(v->byte_count / 2)); + ip.emplace("m_extent", static_cast(mask->byte_count / 2)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_strided.scale", f32_config(params->scale)); + for (const Value * b : { q, k, v, mask, out }) { + dispatch.bindings.push_back({ b->id, 0, b->byte_count }); + } + finish(context, match, std::move(dispatch)); + return true; +} + +// MUL_MAT of a small F16/F32 weight [K, N] (rows packed, any K) or Q8_0 weight (K a multiple of 32) with F32 columns [K, T] into a +// packed [N, T]: heads and projections the tiled matmul kernels refuse (K not a multiple of 256). +// Registered below them; one workitem per output, so it is for small N x T only. +static bool match_mul_mat_small(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_MUL_MAT || node->inputs.size() != 2) { + return false; + } + const Value * w = context.graph.values().find(node->inputs[0]); + const Value * x = context.graph.values().find(node->inputs[1]); + const Value * out = context.graph.values().find(node->output); + const bool q8 = w != nullptr && w->type == GGML_TYPE_Q8_0; + if (w == nullptr || x == nullptr || out == nullptr || + (w->type != GGML_TYPE_F16 && w->type != GGML_TYPE_F32 && !q8) || x->type != GGML_TYPE_F32 || + out->type != GGML_TYPE_F32 || !packed(*out, sizeof(float))) { + return false; + } + // element size, or for Q8_0 the block size (w_s1 then counts blocks) + const size_t wsz = q8 ? ggml_type_size(GGML_TYPE_Q8_0) : ggml_type_size(w->type); + const int64_t k = w->ne[0], n = w->ne[1], t = x->ne[1]; + if (x->ne[0] != k || w->ne[2] != 1 || w->ne[3] != 1 || x->ne[2] != 1 || x->ne[3] != 1 || out->ne[0] != n || + out->ne[1] != t || out->ne[2] != 1 || out->ne[3] != 1 || (!q8 && w->nb[0] != wsz) || w->nb[1] % wsz != 0 || + (q8 && k % 32 != 0) || x->nb[0] != sizeof(float) || x->nb[1] % sizeof(float) != 0 || n * t > (1 << 22) || + k > 1048576) { + return false; + } + Dispatch dispatch; + // Q8_0 with enough K: a workgroup per output, lanes over the blocks (ONEBIT_HRX_Q8_PER_OUTPUT=1: the old way) + const bool rows = q8 && k >= 32 * 64 && std::getenv("ONEBIT_HRX_Q8_PER_OUTPUT") == nullptr; + dispatch.kernel = make_kernel_specialization(rows ? kMulMatRowsQ8Kernel + : q8 ? kMulMatSmallQ8Kernel + : w->type == GGML_TYPE_F16 ? kMulMatSmallF16Kernel : kMulMatSmallF32Kernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("k_size", k); + ip.emplace("n_size", n); + ip.emplace("t_count", t); + ip.emplace("w_s1", static_cast(w->nb[1] / wsz)); + ip.emplace("x_s1", static_cast(x->nb[1] / sizeof(float))); + ip.emplace("w_extent", static_cast(w->byte_count / wsz)); + ip.emplace("x_extent", static_cast(x->byte_count / sizeof(float))); + for (const Value * b : { w, x, out }) { + dispatch.bindings.push_back({ b->id, 0, b->byte_count }); + } + finish(context, match, std::move(dispatch)); + return true; +} + +// the node producing `value` when it is `op` and `value` has no other consumer +static const GraphNode * sole_producer(const Graph & graph, ValueId value, ggml_op op) { + const GraphNode * node = graph.index().producer(value); + return node != nullptr && node->op == op && graph.index().has_single_consumer(value) ? node : nullptr; +} + +// Rotate-half RoPE lowered to eight nodes (ModernBERT through ggmlc): +// out = ADD(MUL(x, cos), MUL(CONT(CONCAT(NEG(CONT(x[half:])), CONT(x[:half]))), sin)) +// with x packed F32 [d, T, H] and cos/sin [d, T] broadcast over heads: one kernel for all eight. +static bool match_rope_rotate_half(const DispatchMatchContext & context, DispatchMatch & match) { + // rooted at x * cos, the chain's first node in graph order (a fused match covers later nodes only) + const Graph & graph = context.graph; + const GraphNode * mul_cos = context.root_node; + if (mul_cos == nullptr || mul_cos->op != GGML_OP_MUL || mul_cos->inputs.size() != 2 || !graph.has_index() || + !graph.index().has_single_consumer(mul_cos->output)) { + return false; + } + const GraphNode * add = graph.index().consumers(mul_cos->output).front(); + if (add == nullptr || add->op != GGML_OP_ADD || add->inputs.size() != 2) { + return false; + } + for (int order = 0; order < 2; ++order) { + if (add->inputs[order] != mul_cos->output) { + continue; + } + const GraphNode * mul_sin = sole_producer(graph, add->inputs[1 - order], GGML_OP_MUL); + if (mul_sin == nullptr || mul_sin->inputs.size() != 2) { + continue; + } + const GraphNode * cat_cont = sole_producer(graph, mul_sin->inputs[0], GGML_OP_CONT); + if (cat_cont == nullptr) { + continue; + } + const GraphNode * concat = sole_producer(graph, cat_cont->inputs[0], GGML_OP_CONCAT); + if (concat == nullptr || concat->inputs.size() != 2) { + continue; + } + const GraphNode * neg = sole_producer(graph, concat->inputs[0], GGML_OP_UNARY); + const GraphNode * lo_cont = sole_producer(graph, concat->inputs[1], GGML_OP_CONT); + if (neg == nullptr || lo_cont == nullptr || neg->inputs.size() != 1) { + continue; + } + const UnaryParams * neg_params = op_params_as(neg->params); + const GraphNode * hi_cont = sole_producer(graph, neg->inputs[0], GGML_OP_CONT); + if (neg_params == nullptr || neg_params->op != UnaryKind::Neg || hi_cont == nullptr) { + continue; + } + const Value * x = graph.values().find(mul_cos->inputs[0]); + const Value * cs = graph.values().find(mul_cos->inputs[1]); + const Value * sn = graph.values().find(mul_sin->inputs[1]); + const Value * hi = graph.values().find(hi_cont->inputs[0]); + const Value * lo = graph.values().find(lo_cont->inputs[0]); + const Value * out = graph.values().find(add->output); + if (x == nullptr || cs == nullptr || sn == nullptr || hi == nullptr || lo == nullptr || out == nullptr) { + continue; + } + const int64_t d = x->ne[0], t = x->ne[1], h = x->ne[2], half = d / 2; + bool ok = x->type == GGML_TYPE_F32 && cs->type == GGML_TYPE_F32 && sn->type == GGML_TYPE_F32 && + out->type == GGML_TYPE_F32 && packed(*x, sizeof(float)) && packed(*out, sizeof(float)) && + same_shape(*x, *out) && x->ne[3] == 1 && d % 2 == 0 && d <= 4096 && h <= 4096 && + out->storage != x->storage && out->alias_source.value < 0; + // the two halves are views of x at its offset and half a row further, with x's strides + ok = ok && hi->storage == x->storage && lo->storage == x->storage && hi->ne[0] == half && lo->ne[0] == half && + lo->storage_offset == x->storage_offset && hi->storage_offset == x->storage_offset + half * sizeof(float) && + hi->nb == x->nb && lo->nb == x->nb && hi->ne[1] == t && hi->ne[2] == h && lo->ne[1] == t && lo->ne[2] == h; + // cos and sin: [d, T], rows possibly strided, broadcast over heads + for (const Value * v : { cs, sn }) { + ok = ok && v->ne[0] == d && v->ne[1] == t && v->ne[2] == 1 && v->ne[3] == 1 && v->nb[0] == sizeof(float) && + v->nb[1] % sizeof(float) == 0 && v->storage != out->storage; + } + if (!ok) { + continue; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kRopeRotateHalfKernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("ne0", d); + ip.emplace("ne1", t); + ip.emplace("ne2", h); + ip.emplace("cos_s1", static_cast(cs->nb[1] / sizeof(float))); + ip.emplace("sin_s1", static_cast(sn->nb[1] / sizeof(float))); + ip.emplace("cos_extent", static_cast(cs->byte_count / sizeof(float))); + ip.emplace("sin_extent", static_cast(sn->byte_count / sizeof(float))); + for (const Value * b : { x, cs, sn, out }) { + dispatch.bindings.push_back({ b->id, 0, b->byte_count }); + } + match.covered_nodes.push_back(context.root_index); + for (const GraphNode * n : { hi_cont, neg, lo_cont, concat, cat_cont, mul_sin, add }) { + if (!append_covered_node_index_once(graph, context.covered_nodes, n, match.covered_nodes)) { + return false; + } + } + match.dispatches.push_back(std::move(dispatch)); + return true; + } + return false; +} + +// GEGLU lowered to CONT(gate view) -> GELU -> MUL(., up view), rooted at the CONT: one kernel +// reads both strided halves and writes gelu(gate) * up. +static bool match_geglu_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const Graph & graph = context.graph; + const GraphNode * cont = context.root_node; + if (cont == nullptr || cont->op != GGML_OP_CONT || cont->inputs.size() != 1 || !graph.has_index() || + !graph.index().has_single_consumer(cont->output)) { + return false; + } + const GraphNode * gelu = graph.index().consumers(cont->output).front(); + const UnaryParams * gp = gelu != nullptr && gelu->op == GGML_OP_UNARY ? op_params_as(gelu->params) : nullptr; + if (gp == nullptr || gp->op != UnaryKind::Gelu || !graph.index().has_single_consumer(gelu->output)) { + return false; + } + const GraphNode * mul = graph.index().consumers(gelu->output).front(); + if (mul == nullptr || mul->op != GGML_OP_MUL || mul->inputs.size() != 2 || mul->inputs[0] != gelu->output) { + return false; + } + const Value * a = graph.values().find(cont->inputs[0]); + const Value * b = graph.values().find(mul->inputs[1]); + const Value * out = graph.values().find(mul->output); + if (a == nullptr || b == nullptr || out == nullptr || a->type != GGML_TYPE_F32 || b->type != GGML_TYPE_F32 || + out->type != GGML_TYPE_F32 || !packed(*out, sizeof(float)) || !same_shape(*a, *out) || !same_shape(*b, *out) || + out->ne[2] != 1 || out->ne[3] != 1 || a->nb[0] != sizeof(float) || b->nb[0] != sizeof(float) || + a->nb[1] % sizeof(float) != 0 || b->nb[1] % sizeof(float) != 0 || out->storage == a->storage || + out->storage == b->storage || out->alias_source.value >= 0) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGegluStridedKernel); + auto & ip = dispatch.kernel.integer_parameters; + ip.emplace("n_size", out->ne[0]); + ip.emplace("t_count", out->ne[1]); + ip.emplace("a_s1", static_cast(a->nb[1] / sizeof(float))); + ip.emplace("b_s1", static_cast(b->nb[1] / sizeof(float))); + ip.emplace("a_extent", static_cast(a->byte_count / sizeof(float))); + ip.emplace("b_extent", static_cast(b->byte_count / sizeof(float))); + for (const Value * v : { a, b, out }) { + dispatch.bindings.push_back({ v->id, 0, v->byte_count }); + } + match.covered_nodes.push_back(context.root_index); + for (const GraphNode * n : { gelu, mul }) { + if (!append_covered_node_index_once(graph, context.covered_nodes, n, match.covered_nodes)) { + return false; + } + } + match.dispatches.push_back(std::move(dispatch)); + return true; +} + } // namespace void register_small_rows_dispatches(DispatchRegistryBuilder & registry) { registry.add({ "common.softmax_rows_f32", GGML_OP_SOFT_MAX, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, match_softmax_rows }); + registry.add({ "common.binary_strided_f32.add", GGML_OP_ADD, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.binary_strided_f32.sub", GGML_OP_SUB, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.binary_strided_f32.mul", GGML_OP_MUL, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.binary_strided_f32.div", GGML_OP_DIV, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_binary_strided }); + registry.add({ "common.geglu_strided_f32", GGML_OP_CONT, DispatchMatchKind::Fused, 300, DispatchSource::Common, + match_geglu_strided }); + registry.add({ "common.rope_rotate_half_f32", GGML_OP_MUL, DispatchMatchKind::Fused, 300, DispatchSource::Common, + match_rope_rotate_half }); + registry.add({ "common.mul_mat_small_f32", GGML_OP_MUL_MAT, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_mul_mat_small }); + registry.add({ "common.attention_strided_f32_f16", GGML_OP_FLASH_ATTN_EXT, DispatchMatchKind::SingleOp, -10, + DispatchSource::Common, match_attention_strided }); + registry.add({ "common.copy_strided_f32_f16", GGML_OP_CPY, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_copy_f32_f16 }); + registry.add({ "common.clamp_f32", GGML_OP_CLAMP, DispatchMatchKind::SingleOp, -10, DispatchSource::Common, + match_clamp }); + registry.add({ "common.norm_rows_f32", GGML_OP_NORM, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_norm_rows }); registry.add({ "common.sum_rows_f32", GGML_OP_SUM_ROWS, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, match_sum_rows }); registry.add({ "common.argsort_rows_f32", GGML_OP_ARGSORT, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h index 6bb09ffca247..76b6464e177e 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h @@ -18,9 +18,9 @@ namespace ggml::hrx { -// SOFT_MAX (no mask, scale 1), SUM_ROWS, ARGSORT and narrow GET_ROWS on short F32 rows, such +// SOFT_MAX (no mask, scale 1), SUM_ROWS, NORM, ARGSORT and narrow GET_ROWS on short F32 rows, such // as a MoE router's, CONT of strided F32 views and broadcast REPEAT: ggml_softmax_rows_f32, ggml_sum_rows_f32, -// ggml_argsort_rows_f32, ggml_get_rows_small_f32 and ggml_copy_strided_f32. +// ggml_argsort_rows_f32, ggml_norm_rows_f32, ggml_get_rows_small_f32 and ggml_copy_strided_f32. void register_small_rows_dispatches(DispatchRegistryBuilder & registry); } // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/ggml-hrx.cpp b/ggml/src/ggml-hrx/ggml-hrx.cpp index e4486ecc4ef8..5826900a17e0 100644 --- a/ggml/src/ggml-hrx/ggml-hrx.cpp +++ b/ggml/src/ggml-hrx/ggml-hrx.cpp @@ -620,6 +620,7 @@ static bool eager_capability_declared(enum ggml_op op) { case GGML_OP_MUL: case GGML_OP_MUL_MAT: case GGML_OP_MUL_MAT_ID: + case GGML_OP_NORM: case GGML_OP_PERMUTE: case GGML_OP_REPEAT: // broadcast-only, through the strided copy (dispatch-small-rows.cpp) case GGML_OP_RESHAPE: diff --git a/ggml/src/ggml-hrx/graph/op-params.cpp b/ggml/src/ggml-hrx/graph/op-params.cpp index 015b6b7b77e3..9820eb1504c0 100644 --- a/ggml/src/ggml-hrx/graph/op-params.cpp +++ b/ggml/src/ggml-hrx/graph/op-params.cpp @@ -258,6 +258,7 @@ OpParams import_op_params(const ggml_tensor & tensor) { switch (tensor.op) { case GGML_OP_RMS_NORM: case GGML_OP_L2_NORM: + case GGML_OP_NORM: return RmsNormParams{ ggml_get_op_params_f32(&tensor, 0) }; case GGML_OP_SOFT_MAX: return SoftMaxParams{ @@ -312,6 +313,7 @@ bool op_params_equivalent(ggml_op op, const OpParams & lhs, const OpParams & rhs switch (op) { case GGML_OP_RMS_NORM: case GGML_OP_L2_NORM: + case GGML_OP_NORM: return rms_norm_params_equivalent(lhs, rhs); case GGML_OP_SOFT_MAX: return soft_max_params_equivalent(lhs, rhs); diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json index f9ba0b1336f5..f05c855020d9 100644 --- a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json @@ -6135,6 +6135,1207 @@ "library_sources": [] }, "compile_dependencies": [] + }, + { + "name": "ggml_norm_rows_f32", + "family": "loom_libs", + "symbol": "ggml_norm_rows_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_binary_strided_f32", + "family": "loom_libs", + "symbol": "ggml_binary_strided_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "a0", + "type": "index" + }, + { + "name": "a1", + "type": "index" + }, + { + "name": "a2", + "type": "index" + }, + { + "name": "a3", + "type": "index" + }, + { + "name": "b0", + "type": "index" + }, + { + "name": "b1", + "type": "index" + }, + { + "name": "b2", + "type": "index" + }, + { + "name": "b3", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + }, + { + "name": "op", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "a0", + "type": "index" + }, + { + "name": "a1", + "type": "index" + }, + { + "name": "a2", + "type": "index" + }, + { + "name": "a3", + "type": "index" + }, + { + "name": "b0", + "type": "index" + }, + { + "name": "b1", + "type": "index" + }, + { + "name": "b2", + "type": "index" + }, + { + "name": "b3", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + }, + { + "name": "op", + "type": "index" + } + ], + "bindings": [ + "lhs", + "rhs", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_clamp_f32", + "family": "loom_libs", + "symbol": "ggml_clamp_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_clamp_inplace_f32", + "family": "loom_libs", + "symbol": "ggml_clamp_inplace_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "element_count", + "type": "index" + } + ], + "bindings": [ + "data" + ], + "binding_access": [ + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_copy_strided_f32_f16", + "family": "loom_libs", + "symbol": "ggml_copy_strided_f32_f16", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "s0", + "type": "index" + }, + { + "name": "s1", + "type": "index" + }, + { + "name": "s2", + "type": "index" + }, + { + "name": "s3", + "type": "index" + }, + { + "name": "source_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "ne3", + "type": "index" + }, + { + "name": "s0", + "type": "index" + }, + { + "name": "s1", + "type": "index" + }, + { + "name": "s2", + "type": "index" + }, + { + "name": "s3", + "type": "index" + }, + { + "name": "source_extent", + "type": "index" + } + ], + "bindings": [ + "input", + "output" + ], + "binding_access": [ + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_attention_strided_f32_f16", + "family": "loom_libs", + "symbol": "ggml_attention_strided_f32_f16", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_small_f16_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_small_f16_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_small_f32_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_small_f32_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_small_q8_0_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_small_q8_0_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_mul_mat_rows_q8_0_f32", + "family": "loom_libs", + "symbol": "ggml_mul_mat_rows_q8_0_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "k_size", + "type": "index" + }, + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "w_s1", + "type": "index" + }, + { + "name": "x_s1", + "type": "index" + }, + { + "name": "w_extent", + "type": "index" + }, + { + "name": "x_extent", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_attention_rows_f32_f16", + "family": "loom_libs", + "symbol": "ggml_attention_rows_f32_f16", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "qk_size", + "type": "index" + }, + { + "name": "v_size", + "type": "index" + }, + { + "name": "q_count", + "type": "index" + }, + { + "name": "kv_count", + "type": "index" + }, + { + "name": "head_count", + "type": "index" + }, + { + "name": "kv_head_count", + "type": "index" + }, + { + "name": "q_s1", + "type": "index" + }, + { + "name": "q_s2", + "type": "index" + }, + { + "name": "k_s1", + "type": "index" + }, + { + "name": "k_s2", + "type": "index" + }, + { + "name": "v_s1", + "type": "index" + }, + { + "name": "v_s2", + "type": "index" + }, + { + "name": "m_s1", + "type": "index" + }, + { + "name": "q_extent", + "type": "index" + }, + { + "name": "k_extent", + "type": "index" + }, + { + "name": "v_extent", + "type": "index" + }, + { + "name": "m_extent", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "value", + "mask", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_rope_rotate_half_f32", + "family": "loom_libs", + "symbol": "ggml_rope_rotate_half_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "cos_s1", + "type": "index" + }, + { + "name": "sin_s1", + "type": "index" + }, + { + "name": "cos_extent", + "type": "index" + }, + { + "name": "sin_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "ne0", + "type": "index" + }, + { + "name": "ne1", + "type": "index" + }, + { + "name": "ne2", + "type": "index" + }, + { + "name": "cos_s1", + "type": "index" + }, + { + "name": "sin_s1", + "type": "index" + }, + { + "name": "cos_extent", + "type": "index" + }, + { + "name": "sin_extent", + "type": "index" + } + ], + "bindings": [ + "input", + "cos", + "sin", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_geglu_strided_f32", + "family": "loom_libs", + "symbol": "ggml_geglu_strided_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "a_s1", + "type": "index" + }, + { + "name": "b_s1", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "n_size", + "type": "index" + }, + { + "name": "t_count", + "type": "index" + }, + { + "name": "a_s1", + "type": "index" + }, + { + "name": "b_s1", + "type": "index" + }, + { + "name": "a_extent", + "type": "index" + }, + { + "name": "b_extent", + "type": "index" + } + ], + "bindings": [ + "gate", + "up", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/small_rows_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] } ], "link_modules": [], diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom index 9e055426900c..881c3b9095d4 100644 --- a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom @@ -14,7 +14,8 @@ // limitations under the License. // // Row ops for short rows (MoE routers, head groups): SOFT_MAX (no mask, scale 1), SUM_ROWS, -// ARGSORT and GET_ROWS. They run where a model's router or head-group reduction would +// ARGSORT and GET_ROWS; and NORM (LayerNorm without affine, as encoders such as ModernBERT use), +// one workgroup per row. They run where a model's router or head-group reduction would // otherwise leave the GPU for a few dozen floats per token. One workitem per row (softmax, // sum) or per element (argsort rank, gather); rows are short, so no cross-lane reduction. @@ -265,3 +266,840 @@ kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_copy_strided_f } kernel.return } + +// NORM: y = (x - mean) / sqrt(var + eps) per row, var of the centered values (ggml's order). +// One 256-lane workgroup per row: lanes stride the columns, two workgroup reductions. +config.decl @ggml.norm_rows_f32.epsilon : f32 + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_norm_rows_f32") @ggml_norm_rows_f32(%column_count: index, %row_count: index) { + %one = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%row_count, %one, %one) workgroup_size(%c256, %one, %one) : index +} launch(%column_count: index, %row_count: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 65536)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %eps = config.get @ggml.norm_rows_f32.epsilon : f32 + %c256 = index.constant 256 : index + %c0_f32 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %row0 = kernel.workgroup.id : index + %row = index.assume %row0 [range(%row0, 0, 16777215)] : index + %lane = kernel.workitem.id : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%rows]x[%cols]xf32> + %sum = scf.for %column = [%lane to %cols step %c256](%acc = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %next = scalar.addf %acc, %value : f32 + scf.yield %next : f32 + } + %row_sum = kernel.workgroup.reduce %sum : f32 + %cols_i32 = index.cast %cols : index to i32 + %cols_f32 = scalar.sitofp %cols_i32 : i32 to f32 + %mean = scalar.divf %row_sum, %cols_f32 : f32 + %squares = scf.for %column = [%lane to %cols step %c256](%acc = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %centered = scalar.subf %value, %mean : f32 + %square = scalar.mulf %centered, %centered : f32 + %next = scalar.addf %acc, %square : f32 + scf.yield %next : f32 + } + %row_squares = kernel.workgroup.reduce %squares : f32 + %variance = scalar.divf %row_squares, %cols_f32 : f32 + %biased = scalar.addf %variance, %eps : f32 + %scale = scalar.rsqrtf %biased : f32 + %unused = scf.for %column = [%lane to %cols step %c256](%carry = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %centered = scalar.subf %value, %mean : f32 + %normalized = scalar.mulf %centered, %scale : f32 + view.store %normalized, %output_view[%row, %column] : f32, view<[%rows]x[%cols]xf32> + scf.yield %carry : f32 + } + kernel.return +} + +// ADD / SUB / MUL / DIV with either input strided or broadcast (stride 0 on a broadcast dim) into a +// packed output: the cases the packed binary kernels do not take. One workitem per output element. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_binary_strided_f32") @ggml_binary_strided_f32(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %a0: index, %a1: index, %a2: index, %a3: index, %b0: index, %b1: index, %b2: index, %b3: index, %a_extent: index, %b_extent: index, %op: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %e01 = index.mul %ne0, %ne1 : index + %e012 = index.mul %e01, %ne2 : index + %elements = index.mul %e012, %ne3 : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %a0: index, %a1: index, %a2: index, %a3: index, %b0: index, %b1: index, %b2: index, %b3: index, %a_extent: index, %b_extent: index, %op: index, %lhs: buffer, %rhs: buffer, %output: buffer) { + %n0 = index.assume %ne0 [range(%ne0, 1, 16777216)] : index + %n1 = index.assume %ne1 [range(%ne1, 1, 16777216)] : index + %n2 = index.assume %ne2 [range(%ne2, 1, 16777216)] : index + %n3 = index.assume %ne3 [range(%ne3, 1, 16777216)] : index + %ae = index.assume %a_extent [range(%a_extent, 1, 268435456)] : index + %be = index.assume %b_extent [range(%b_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %linear0 = index.add %base, %lane : index + %e01 = index.mul %n0, %n1 : index + %e012 = index.mul %e01, %n2 : index + %elements = index.mul %e012, %n3 : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %i0 = index.rem %linear, %n0 : index + %q0 = index.div %linear, %n0 : index + %i1 = index.rem %q0, %n1 : index + %q1 = index.div %q0, %n1 : index + %i2 = index.rem %q1, %n2 : index + %i3 = index.div %q1, %n2 : index + %x0 = index.mul %i0, %a0 : index + %x1 = index.mul %i1, %a1 : index + %x2 = index.mul %i2, %a2 : index + %x3 = index.mul %i3, %a3 : index + %x01 = index.add %x0, %x1 : index + %x23 = index.add %x2, %x3 : index + %xa0 = index.add %x01, %x23 : index + %xin = index.cmp ult, %xa0, %ae : index + %xa = scf.select %xin, %xa0, %c0 : index + %y0 = index.mul %i0, %b0 : index + %y1 = index.mul %i1, %b1 : index + %y2 = index.mul %i2, %b2 : index + %y3 = index.mul %i3, %b3 : index + %y01 = index.add %y0, %y1 : index + %y23 = index.add %y2, %y3 : index + %yb0 = index.add %y01, %y23 : index + %yin = index.cmp ult, %yb0, %be : index + %yb = scf.select %yin, %yb0, %c0 : index + %lhs_na, %rhs_na, %out_na = buffer.assume.noalias %lhs, %rhs, %output : buffer, buffer, buffer + %lhs_view = buffer.view %lhs_na[%zero_offset] : buffer -> view<[%ae]xf32> + %rhs_view = buffer.view %rhs_na[%zero_offset] : buffer -> view<[%be]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%elements]xf32> + %a = view.load %lhs_view[%xa] : view<[%ae]xf32> -> f32 + %b = view.load %rhs_view[%yb] : view<[%be]xf32> -> f32 + %sum = scalar.addf %a, %b : f32 + %difference = scalar.subf %a, %b : f32 + %product = scalar.mulf %a, %b : f32 + %quotient = scalar.divf %a, %b : f32 + %is_add = index.cmp eq, %op, %c0 : index + %is_sub = index.cmp eq, %op, %c1 : index + %is_mul = index.cmp eq, %op, %c2 : index + %r0 = scf.select %is_mul, %product, %quotient : f32 + %r1 = scf.select %is_sub, %difference, %r0 : f32 + %result = scf.select %is_add, %sum, %r1 : f32 + scf.if %valid { + view.store %result, %out_view[%linear] : f32, view<[%elements]xf32> + } + kernel.return +} + +// CLAMP of a packed F32 tensor, standalone (the MoE router's clamp is fused elsewhere). +config.decl @ggml.clamp_f32.min : f32 + +config.decl @ggml.clamp_f32.max : f32 + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_clamp_f32") @ggml_clamp_f32(%element_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %element_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%element_count: index, %input: buffer, %output: buffer) { + %n = index.assume %element_count [range(%element_count, 1, 268435456)] : index + %lo = config.get @ggml.clamp_f32.min : f32 + %hi = config.get @ggml.clamp_f32.max : f32 + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %i0 = index.add %base, %lane : index + %valid = index.cmp ult, %i0, %n : index + %i = scf.select %valid, %i0, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%n]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%n]xf32> + %value = view.load %input_view[%i] : view<[%n]xf32> -> f32 + %above = scalar.maxnumf %value, %lo : f32 + %clamped = scalar.minnumf %above, %hi : f32 + scf.if %valid { + view.store %clamped, %output_view[%i] : f32, view<[%n]xf32> + } + kernel.return +} + +// CLAMP in place (ggml_clamp's output is a view of its input): one buffer, read then write. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_clamp_inplace_f32") @ggml_clamp_inplace_f32(%element_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %element_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%element_count: index, %data: buffer) { + %n = index.assume %element_count [range(%element_count, 1, 268435456)] : index + %lo = config.get @ggml.clamp_f32.min : f32 + %hi = config.get @ggml.clamp_f32.max : f32 + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %i0 = index.add %base, %lane : index + %valid = index.cmp ult, %i0, %n : index + %i = scf.select %valid, %i0, %c0 : index + %view = buffer.view %data[%zero_offset] : buffer -> view<[%n]xf32> + %value = view.load %view[%i] : view<[%n]xf32> -> f32 + %above = scalar.maxnumf %value, %lo : f32 + %clamped = scalar.minnumf %above, %hi : f32 + scf.if %valid { + view.store %clamped, %view[%i] : f32, view<[%n]xf32> + } + kernel.return +} + +// CPY of an F32 tensor (any element strides) into a packed F16 destination, such as an attention +// mask converted for flash attention. One workitem per element. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_copy_strided_f32_f16") @ggml_copy_strided_f32_f16(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %s0: index, %s1: index, %s2: index, %s3: index, %source_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %e01 = index.mul %ne0, %ne1 : index + %e012 = index.mul %e01, %ne2 : index + %elements = index.mul %e012, %ne3 : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%ne0: index, %ne1: index, %ne2: index, %ne3: index, %s0: index, %s1: index, %s2: index, %s3: index, %source_extent: index, %input: buffer, %output: buffer) { + %n0 = index.assume %ne0 [range(%ne0, 1, 16777216)] : index + %n1 = index.assume %ne1 [range(%ne1, 1, 16777216)] : index + %n2 = index.assume %ne2 [range(%ne2, 1, 16777216)] : index + %n3 = index.assume %ne3 [range(%ne3, 1, 16777216)] : index + %extent = index.assume %source_extent [range(%source_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %linear0 = index.add %base, %lane : index + %e01 = index.mul %n0, %n1 : index + %e012 = index.mul %e01, %n2 : index + %elements = index.mul %e012, %n3 : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %i0 = index.rem %linear, %n0 : index + %q0 = index.div %linear, %n0 : index + %i1 = index.rem %q0, %n1 : index + %q1 = index.div %q0, %n1 : index + %i2 = index.rem %q1, %n2 : index + %i3 = index.div %q1, %n2 : index + %o0 = index.mul %i0, %s0 : index + %o1 = index.mul %i1, %s1 : index + %o2 = index.mul %i2, %s2 : index + %o3 = index.mul %i3, %s3 : index + %o01 = index.add %o0, %o1 : index + %o23 = index.add %o2, %o3 : index + %source0 = index.add %o01, %o23 : index + %in_range = index.cmp ult, %source0, %extent : index + %source = scf.select %in_range, %source0, %c0 : index + %input_noalias, %output_noalias = buffer.assume.noalias %input, %output : buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%extent]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%elements]xf16> + %value = view.load %input_view[%source] : view<[%extent]xf32> -> f32 + %half = scalar.fptrunc %value : f32 to f16 + scf.if %valid { + view.store %half, %output_view[%linear] : f16, view<[%elements]xf16> + } + kernel.return +} + +// FLASH_ATTN_EXT for encoder-sized inputs with any Q/K/V/mask strides (F32 query, F16 key, value +// and mask, no ALiBi or softcap): the layouts the llama-shaped flash-attention kernels do not +// take, such as one contiguous block per head. One workgroup per (query token, head), one lane +// per value dimension, online softmax over the keys (running max and sum), so no lane has to +// wait for another; each lane computes the query-key dot products itself. +config.decl @ggml.attention_strided.scale : f32 + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_attention_strided_f32_f16") @ggml_attention_strided_f32_f16(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index) { + %one = index.constant 1 : index + %groups = index.mul %q_count, %head_count : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%v_size, %one, %one) : index +} launch(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) { + %d = index.assume %qk_size [range(%qk_size, 1, 1024)] : index + %dv = index.assume %v_size [range(%v_size, 1, 1024)] : index + %nq = index.assume %q_count [range(%q_count, 1, 65536)] : index + %nkv = index.assume %kv_count [range(%kv_count, 1, 65536)] : index + %nh = index.assume %head_count [range(%head_count, 1, 1024)] : index + %nhkv = index.assume %kv_head_count [range(%kv_head_count, 1, 1024)] : index + %qe = index.assume %q_extent [range(%q_extent, 1, 268435456)] : index + %ke = index.assume %k_extent [range(%k_extent, 1, 268435456)] : index + %ve = index.assume %v_extent [range(%v_extent, 1, 268435456)] : index + %me = index.assume %m_extent [range(%m_extent, 1, 268435456)] : index + %scale = config.get @ggml.attention_strided.scale : f32 + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %zero_offset = index.constant 0 : offset + %f0 = scalar.constant 0.0 : f32 + %lowest = scalar.constant -3.40282347e+38 : f32 + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %qi = index.div %group, %nh : index + %h = index.rem %group, %nh : index + %ratio = index.div %nh, %nhkv : index + %hkv = index.div %h, %ratio : index + %q_na, %k_na, %v_na, %m_na, %o_na = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %q_view = buffer.view %q_na[%zero_offset] : buffer -> view<[%qe]xf32> + %k_view = buffer.view %k_na[%zero_offset] : buffer -> view<[%ke]xf16> + %v_view = buffer.view %v_na[%zero_offset] : buffer -> view<[%ve]xf16> + %m_view = buffer.view %m_na[%zero_offset] : buffer -> view<[%me]xf16> + %out_count0 = index.mul %dv, %nh : index + %out_count = index.mul %out_count0, %nq : index + %o_view = buffer.view %o_na[%zero_offset] : buffer -> view<[%out_count]xf32> + %qa = index.mul %qi, %q_s1 : index + %qb = index.mul %h, %q_s2 : index + %q_base = index.add %qa, %qb : index + %kb = index.mul %hkv, %k_s2 : index + %vb = index.mul %hkv, %v_s2 : index + %m_base = index.mul %qi, %m_s1 : index + %m_final, %l_final, %acc_final = scf.for %j = [%c0 to %nkv step %c1](%m = %lowest : f32, %l = %f0 : f32, %acc = %f0 : f32) -> (f32, f32, f32) { + %kj = index.mul %j, %k_s1 : index + %k_base = index.add %kj, %kb : index + %dot = scf.for %t = [%c0 to %d step %c1](%sum = %f0 : f32) -> (f32) { + %qt0 = index.add %q_base, %t : index + %qt = index.assume %qt0 [range(%qt0, 0, 268435455)] : index + %kt0 = index.add %k_base, %t : index + %kt = index.assume %kt0 [range(%kt0, 0, 268435455)] : index + %qv = view.load %q_view[%qt] : view<[%qe]xf32> -> f32 + %kv16 = view.load %k_view[%kt] : view<[%ke]xf16> -> f16 + %kv = scalar.extf %kv16 : f16 to f32 + %next = scalar.fmaf %qv, %kv, %sum : f32 + scf.yield %next : f32 + } + %mi0 = index.add %m_base, %j : index + %mi = index.assume %mi0 [range(%mi0, 0, 268435455)] : index + %mask16 = view.load %m_view[%mi] : view<[%me]xf16> -> f16 + %maskv = scalar.extf %mask16 : f16 to f32 + %scaled = scalar.mulf %dot, %scale : f32 + %score = scalar.addf %scaled, %maskv : f32 + %m_new = scalar.maxnumf %m, %score : f32 + %d_old = scalar.subf %m, %m_new : f32 + %corr = scalar.expf %d_old : f32 + %d_new = scalar.subf %score, %m_new : f32 + %p = scalar.expf %d_new : f32 + %l_scaled = scalar.mulf %l, %corr : f32 + %l_new = scalar.addf %l_scaled, %p : f32 + %vj = index.mul %j, %v_s1 : index + %v_row = index.add %vj, %vb : index + %vi0 = index.add %v_row, %lane : index + %vi = index.assume %vi0 [range(%vi0, 0, 268435455)] : index + %v16 = view.load %v_view[%vi] : view<[%ve]xf16> -> f16 + %vv = scalar.extf %v16 : f16 to f32 + %acc_scaled = scalar.mulf %acc, %corr : f32 + %acc_new = scalar.fmaf %p, %vv, %acc_scaled : f32 + scf.yield %m_new, %l_new, %acc_new : f32, f32, f32 + } + %result = scalar.divf %acc_final, %l_final : f32 + %oh = index.mul %h, %dv : index + %oq0 = index.mul %qi, %nh : index + %oq = index.mul %oq0, %dv : index + %o0 = index.add %oq, %oh : index + %oi0 = index.add %o0, %lane : index + %oi = index.assume %oi0 [range(%oi0, 0, 268435455)] : index + view.store %result, %o_view[%oi] : f32, view<[%out_count]xf32> + kernel.return +} + +// MUL_MAT of a small F16 weight [K, N] (any K, such as a classification head's) with F32 columns +// [K, T] (any element strides), into a packed [N, T]: one workitem per output, a loop over K. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_small_f16_f32") @ggml_mul_mat_small_f16_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %outputs = index.mul %n_size, %t_count : index + %rounded = index.add %outputs, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 1, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %we = index.assume %w_extent [range(%w_extent, 1, 268435456)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %o0 = index.add %base, %lane : index + %outputs = index.mul %n, %t : index + %valid = index.cmp ult, %o0, %outputs : index + %o = scf.select %valid, %o0, %c0 : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %w_na, %x_na, %out_na = buffer.assume.noalias %weight, %input, %output : buffer, buffer, buffer + %w_view = buffer.view %w_na[%zero_offset] : buffer -> view<[%we]xf16> + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %w_base = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %dot = scf.for %i = [%c0 to %k step %c1](%sum = %f0 : f32) -> (f32) { + %wi0 = index.add %w_base, %i : index + %wi = index.assume %wi0 [range(%wi0, 0, 268435455)] : index + %xi0 = index.add %x_base, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %w16 = view.load %w_view[%wi] : view<[%we]xf16> -> f16 + %w = scalar.extf %w16 : f16 to f32 + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %w, %x, %sum : f32 + scf.yield %next : f32 + } + scf.if %valid { + view.store %dot, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// MUL_MAT of a small F32 weight [K, N] (any K, such as a classification head's) with F32 columns +// [K, T] (any element strides), into a packed [N, T]: one workitem per output, a loop over K. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_small_f32_f32") @ggml_mul_mat_small_f32_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %outputs = index.mul %n_size, %t_count : index + %rounded = index.add %outputs, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 1, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %we = index.assume %w_extent [range(%w_extent, 1, 268435456)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %o0 = index.add %base, %lane : index + %outputs = index.mul %n, %t : index + %valid = index.cmp ult, %o0, %outputs : index + %o = scf.select %valid, %o0, %c0 : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %w_na, %x_na, %out_na = buffer.assume.noalias %weight, %input, %output : buffer, buffer, buffer + %w_view = buffer.view %w_na[%zero_offset] : buffer -> view<[%we]xf32> + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %w_base = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %dot = scf.for %i = [%c0 to %k step %c1](%sum = %f0 : f32) -> (f32) { + %wi0 = index.add %w_base, %i : index + %wi = index.assume %wi0 [range(%wi0, 0, 268435455)] : index + %xi0 = index.add %x_base, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %w = view.load %w_view[%wi] : view<[%we]xf32> -> f32 + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %w, %x, %sum : f32 + scf.yield %next : f32 + } + scf.if %valid { + view.store %dot, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// MUL_MAT of a Q8_0 weight [K, N] with K any multiple of 32 (the tiled kernels need 256), with F32 +// columns [K, T] (any element strides), into a packed [N, T]: one workitem per output, a loop over +// the 34-byte blocks (f16 scale, 32 int8 codes). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_small_q8_0_f32") @ggml_mul_mat_small_q8_0_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %outputs = index.mul %n_size, %t_count : index + %rounded = index.add %outputs, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 32, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %block_bytes = index.constant 34 : offset + %code_offset = index.constant 2 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %o0 = index.add %base, %lane : index + %outputs = index.mul %n, %t : index + %valid = index.cmp ult, %o0, %outputs : index + %o = scf.select %valid, %o0, %c0 : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %x_na, %out_na = buffer.assume.noalias %input, %output : buffer, buffer + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %blocks = index.div %k, %c32 : index + %row_block = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %dot = scf.for %b = [%c0 to %blocks step %c1](%sum = %f0 : f32) -> (f32) { + %gb = index.add %row_block, %b : index + %block_base = index.scale %gb, %block_bytes : index, offset -> offset + %code_base = index.add %block_base, %code_offset : offset + %d_view = buffer.view %weight[%block_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_base] : buffer -> view<32xi8> + %d16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d16 : f16 to f32 + %xb = index.mul %b, %c32 : index + %xrow = index.add %x_base, %xb : index + %partial = scf.for %i = [%c0 to %c32 step %c1](%acc = %f0 : f32) -> (f32) { + %code = view.load %code_view[%i] : view<32xi8> -> i8 + %q = scalar.sitofp %code : i8 to f32 + %xi0 = index.add %xrow, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %q, %x, %acc : f32 + scf.yield %next : f32 + } + %next_sum = scalar.fmaf %d, %partial, %sum : f32 + scf.yield %next_sum : f32 + } + scf.if %valid { + view.store %dot, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// Q8_0 MUL_MAT, one 64-lane workgroup per output: the lanes split the K blocks (reads along each +// weight row stay together), then one workgroup reduction. For K = 2624 (ModernBERT's FFN down +// projection) the per-output kernel above reads 64 different rows per wave. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_mul_mat_rows_q8_0_f32") @ggml_mul_mat_rows_q8_0_f32(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %outputs = index.mul %n_size, %t_count : index + kernel.launch.config workgroups(%outputs, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%k_size: index, %n_size: index, %t_count: index, %w_s1: index, %x_s1: index, %w_extent: index, %x_extent: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %k_size [range(%k_size, 32, 1048576)] : index + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %xe = index.assume %x_extent [range(%x_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c32 = index.constant 32 : index + %c64 = index.constant 64 : index + %f0 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %block_bytes = index.constant 34 : offset + %code_offset = index.constant 2 : offset + %o = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %col = index.rem %o, %n : index + %tok = index.div %o, %n : index + %x_na, %out_na = buffer.assume.noalias %input, %output : buffer, buffer + %x_view = buffer.view %x_na[%zero_offset] : buffer -> view<[%xe]xf32> + %outputs = index.mul %n, %t : index + %out_view = buffer.view %out_na[%zero_offset] : buffer -> view<[%outputs]xf32> + %blocks = index.div %k, %c32 : index + %row_block = index.mul %col, %w_s1 : index + %x_base = index.mul %tok, %x_s1 : index + %mine = scf.for %b = [%lane to %blocks step %c64](%sum = %f0 : f32) -> (f32) { + %gb = index.add %row_block, %b : index + %block_base = index.scale %gb, %block_bytes : index, offset -> offset + %code_base = index.add %block_base, %code_offset : offset + %d_view = buffer.view %weight[%block_base] : buffer -> view<1xf16> + %code_view = buffer.view %weight[%code_base] : buffer -> view<32xi8> + %d16 = view.load %d_view[%c0] : view<1xf16> -> f16 + %d = scalar.extf %d16 : f16 to f32 + %xb = index.mul %b, %c32 : index + %xrow = index.add %x_base, %xb : index + %partial = scf.for %i = [%c0 to %c32 step %c1](%acc = %f0 : f32) -> (f32) { + %code = view.load %code_view[%i] : view<32xi8> -> i8 + %q = scalar.sitofp %code : i8 to f32 + %xi0 = index.add %xrow, %i : index + %xi = index.assume %xi0 [range(%xi0, 0, 268435455)] : index + %x = view.load %x_view[%xi] : view<[%xe]xf32> -> f32 + %next = scalar.fmaf %q, %x, %acc : f32 + scf.yield %next : f32 + } + %next_sum = scalar.fmaf %d, %partial, %sum : f32 + scf.yield %next_sum : f32 + } + %total = kernel.workgroup.reduce %mine : f32 + %first = index.cmp eq, %lane, %c0 : index + scf.if %first { + view.store %total, %out_view[%o] : f32, view<[%outputs]xf32> + } + kernel.return +} + +// The same attention with each score computed once: one 64-lane workgroup per (query token, head) +// stages the query row in workgroup memory, each lane scores its keys (j = lane, lane + 64, ...) +// into a workgroup score row, workgroup max and sum reductions give the softmax, and the lanes +// then take the value dimensions (reads along each value row). Up to 2048 keys, head sizes <= 256. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_attention_rows_f32_f16") @ggml_attention_rows_f32_f16(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %groups = index.mul %q_count, %head_count : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%qk_size: index, %v_size: index, %q_count: index, %kv_count: index, %head_count: index, %kv_head_count: index, %q_s1: index, %q_s2: index, %k_s1: index, %k_s2: index, %v_s1: index, %v_s2: index, %m_s1: index, %q_extent: index, %k_extent: index, %v_extent: index, %m_extent: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %output: buffer) { + %d = index.assume %qk_size [range(%qk_size, 1, 256)] : index + %dv = index.assume %v_size [range(%v_size, 1, 256)] : index + %nq = index.assume %q_count [range(%q_count, 1, 65536)] : index + %nkv = index.assume %kv_count [range(%kv_count, 1, 2048)] : index + %nh = index.assume %head_count [range(%head_count, 1, 1024)] : index + %nhkv = index.assume %kv_head_count [range(%kv_head_count, 1, 1024)] : index + %qe = index.assume %q_extent [range(%q_extent, 1, 268435456)] : index + %ke = index.assume %k_extent [range(%k_extent, 1, 268435456)] : index + %ve = index.assume %v_extent [range(%v_extent, 1, 268435456)] : index + %me = index.assume %m_extent [range(%m_extent, 1, 268435456)] : index + %scale = config.get @ggml.attention_strided.scale : f32 + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %zero = index.constant 0 : offset + %f0 = scalar.constant 0.0 : f32 + %lowest = scalar.constant -3.40282347e+38 : f32 + %q_bytes = index.constant 1024 : offset + %s_bytes = index.constant 8192 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %qi = index.div %group, %nh : index + %h = index.rem %group, %nh : index + %ratio = index.div %nh, %nhkv : index + %hkv = index.div %h, %ratio : index + %q_na, %k_na, %v_na, %m_na, %o_na = buffer.assume.noalias %query, %key, %value, %mask, %output : buffer, buffer, buffer, buffer, buffer + %q_view = buffer.view %q_na[%zero] : buffer -> view<[%qe]xf32> + %k_view = buffer.view %k_na[%zero] : buffer -> view<[%ke]xf16> + %v_view = buffer.view %v_na[%zero] : buffer -> view<[%ve]xf16> + %m_view = buffer.view %m_na[%zero] : buffer -> view<[%me]xf16> + %out_count0 = index.mul %dv, %nh : index + %out_count = index.mul %out_count0, %nq : index + %o_view = buffer.view %o_na[%zero] : buffer -> view<[%out_count]xf32> + %q_shared = buffer.alloca align(16) %q_bytes : buffer + %q_row = buffer.view %q_shared[%zero] : buffer -> view<256xf32> + %s_shared = buffer.alloca align(16) %s_bytes : buffer + %s_row = buffer.view %s_shared[%zero] : buffer -> view<2048xf32> + %qa = index.mul %qi, %q_s1 : index + %qb = index.mul %h, %q_s2 : index + %q_base = index.add %qa, %qb : index + %u0 = scf.for %t = [%lane to %d step %c64](%carry0 = %f0 : f32) -> (f32) { + %qt0 = index.add %q_base, %t : index + %qt = index.assume %qt0 [range(%qt0, 0, 268435455)] : index + %qv = view.load %q_view[%qt] : view<[%qe]xf32> -> f32 + view.store %qv, %q_row[%t] : f32, view<256xf32> + scf.yield %carry0 : f32 + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %kb = index.mul %hkv, %k_s2 : index + %m_base = index.mul %qi, %m_s1 : index + %lane_max = scf.for %j = [%lane to %nkv step %c64](%mx = %lowest : f32) -> (f32) { + %kj = index.mul %j, %k_s1 : index + %k_base = index.add %kj, %kb : index + %dot = scf.for %t = [%c0 to %d step %c1](%sum = %f0 : f32) -> (f32) { + %kt0 = index.add %k_base, %t : index + %kt = index.assume %kt0 [range(%kt0, 0, 268435455)] : index + %qv = view.load %q_row[%t] : view<256xf32> -> f32 + %kv16 = view.load %k_view[%kt] : view<[%ke]xf16> -> f16 + %kv = scalar.extf %kv16 : f16 to f32 + %next = scalar.fmaf %qv, %kv, %sum : f32 + scf.yield %next : f32 + } + %mi0 = index.add %m_base, %j : index + %mi = index.assume %mi0 [range(%mi0, 0, 268435455)] : index + %mask16 = view.load %m_view[%mi] : view<[%me]xf16> -> f16 + %maskv = scalar.extf %mask16 : f16 to f32 + %scaled = scalar.mulf %dot, %scale : f32 + %score = scalar.addf %scaled, %maskv : f32 + view.store %score, %s_row[%j] : f32, view<2048xf32> + %next_max = scalar.maxnumf %mx, %score : f32 + scf.yield %next_max : f32 + } + %row_max = kernel.workgroup.reduce %lane_max : f32 + %lane_sum = scf.for %j = [%lane to %nkv step %c64](%acc = %f0 : f32) -> (f32) { + %score = view.load %s_row[%j] : view<2048xf32> -> f32 + %shifted = scalar.subf %score, %row_max : f32 + %p = scalar.expf %shifted : f32 + view.store %p, %s_row[%j] : f32, view<2048xf32> + %next = scalar.addf %acc, %p : f32 + scf.yield %next : f32 + } + %row_sum = kernel.workgroup.reduce %lane_sum : f32 + kernel.barrier scope(workgroup) ordering(acq_rel) + %vb = index.mul %hkv, %v_s2 : index + %oh = index.mul %h, %dv : index + %oq0 = index.mul %qi, %nh : index + %oq = index.mul %oq0, %dv : index + %o_base = index.add %oq, %oh : index + %u1 = scf.for %c = [%lane to %dv step %c64](%carry1 = %f0 : f32) -> (f32) { + %acc_v = scf.for %j = [%c0 to %nkv step %c1](%acc = %f0 : f32) -> (f32) { + %p = view.load %s_row[%j] : view<2048xf32> -> f32 + %vj = index.mul %j, %v_s1 : index + %v_row = index.add %vj, %vb : index + %vi0 = index.add %v_row, %c : index + %vi = index.assume %vi0 [range(%vi0, 0, 268435455)] : index + %v16 = view.load %v_view[%vi] : view<[%ve]xf16> -> f16 + %vv = scalar.extf %v16 : f16 to f32 + %next = scalar.fmaf %p, %vv, %acc : f32 + scf.yield %next : f32 + } + %result = scalar.divf %acc_v, %row_sum : f32 + %oi0 = index.add %o_base, %c : index + %oi = index.assume %oi0 [range(%oi0, 0, 268435455)] : index + view.store %result, %o_view[%oi] : f32, view<[%out_count]xf32> + scf.yield %carry1 : f32 + } + kernel.return +} + +// Rotate-half RoPE as graph compilers lower it (x * cos + concat(-x[half:], x[:half]) * sin, eight +// nodes), in one pass: x packed [d, T, H], cos and sin [d, T] (row strides given), broadcast over +// heads. One workitem per element. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_rope_rotate_half_f32") @ggml_rope_rotate_half_f32(%ne0: index, %ne1: index, %ne2: index, %cos_s1: index, %sin_s1: index, %cos_extent: index, %sin_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %e01 = index.mul %ne0, %ne1 : index + %elements = index.mul %e01, %ne2 : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%ne0: index, %ne1: index, %ne2: index, %cos_s1: index, %sin_s1: index, %cos_extent: index, %sin_extent: index, %input: buffer, %cos: buffer, %sin: buffer, %output: buffer) { + %d = index.assume %ne0 [range(%ne0, 2, 4096)] : index + %t = index.assume %ne1 [range(%ne1, 1, 1048576)] : index + %h = index.assume %ne2 [range(%ne2, 1, 4096)] : index + %ce = index.assume %cos_extent [range(%cos_extent, 1, 268435456)] : index + %se = index.assume %sin_extent [range(%sin_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c2 = index.constant 2 : index + %c64 = index.constant 64 : index + %zero = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %lin0 = index.add %base, %lane : index + %e01 = index.mul %d, %t : index + %elements = index.mul %e01, %h : index + %valid = index.cmp ult, %lin0, %elements : index + %lin = scf.select %valid, %lin0, %c0 : index + %i0 = index.rem %lin, %d : index + %q0 = index.div %lin, %d : index + %ti = index.rem %q0, %t : index + %half = index.div %d, %c2 : index + %low = index.cmp ult, %i0, %half : index + %up = index.add %lin, %half : index + %down0 = index.sub %lin, %half : index + %down = scf.select %low, %lin, %down0 : index + %partner0 = scf.select %low, %up, %down : index + %partner_in = index.cmp ult, %partner0, %elements : index + %partner = scf.select %partner_in, %partner0, %c0 : index + %x_na, %c_na, %s_na, %o_na = buffer.assume.noalias %input, %cos, %sin, %output : buffer, buffer, buffer, buffer + %x_view = buffer.view %x_na[%zero] : buffer -> view<[%elements]xf32> + %c_view = buffer.view %c_na[%zero] : buffer -> view<[%ce]xf32> + %s_view = buffer.view %s_na[%zero] : buffer -> view<[%se]xf32> + %o_view = buffer.view %o_na[%zero] : buffer -> view<[%elements]xf32> + %x = view.load %x_view[%lin] : view<[%elements]xf32> -> f32 + %other = view.load %x_view[%partner] : view<[%elements]xf32> -> f32 + %neg_other = scalar.negf %other : f32 + %rot = scf.select %low, %neg_other, %other : f32 + %cr = index.mul %ti, %cos_s1 : index + %ci0 = index.add %cr, %i0 : index + %ci_in = index.cmp ult, %ci0, %ce : index + %ci = scf.select %ci_in, %ci0, %c0 : index + %sr = index.mul %ti, %sin_s1 : index + %si0 = index.add %sr, %i0 : index + %si_in = index.cmp ult, %si0, %se : index + %si = scf.select %si_in, %si0, %c0 : index + %cv = view.load %c_view[%ci] : view<[%ce]xf32> -> f32 + %sv = view.load %s_view[%si] : view<[%se]xf32> -> f32 + %xc = scalar.mulf %x, %cv : f32 + %result = scalar.fmaf %rot, %sv, %xc : f32 + scf.if %valid { + view.store %result, %o_view[%lin] : f32, view<[%elements]xf32> + } + kernel.return +} + +// GEGLU lowered to three nodes (CONT of the gate half, GELU, MUL by the up half): gelu(a) * b with +// a and b strided F32 views [n, T] (rows a_s1 / b_s1 apart), into a packed [n, T]. GELU in ggml's +// tanh form, 0.5 x (1 + tanh(sqrt(2/pi) x (1 + 0.044715 x^2))), tanh(z) = 1 - 2 / (exp(2z) + 1). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_geglu_strided_f32") @ggml_geglu_strided_f32(%n_size: index, %t_count: index, %a_s1: index, %b_s1: index, %a_extent: index, %b_extent: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %elements = index.mul %n_size, %t_count : index + %rounded = index.add %elements, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%n_size: index, %t_count: index, %a_s1: index, %b_s1: index, %a_extent: index, %b_extent: index, %gate: buffer, %up: buffer, %output: buffer) { + %n = index.assume %n_size [range(%n_size, 1, 1048576)] : index + %t = index.assume %t_count [range(%t_count, 1, 1048576)] : index + %ae = index.assume %a_extent [range(%a_extent, 1, 268435456)] : index + %be = index.assume %b_extent [range(%b_extent, 1, 268435456)] : index + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %zero = index.constant 0 : offset + %half = scalar.constant 0.5 : f32 + %one_f = scalar.constant 1.0 : f32 + %two = scalar.constant 2.0 : f32 + %k0 = scalar.constant 0.7978845608 : f32 + %k1 = scalar.constant 0.044715 : f32 + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %lin0 = index.add %base, %lane : index + %elements = index.mul %n, %t : index + %valid = index.cmp ult, %lin0, %elements : index + %lin = scf.select %valid, %lin0, %c0 : index + %i = index.rem %lin, %n : index + %row = index.div %lin, %n : index + %ar = index.mul %row, %a_s1 : index + %ai0 = index.add %ar, %i : index + %ai_in = index.cmp ult, %ai0, %ae : index + %ai = scf.select %ai_in, %ai0, %c0 : index + %br = index.mul %row, %b_s1 : index + %bi0 = index.add %br, %i : index + %bi_in = index.cmp ult, %bi0, %be : index + %bi = scf.select %bi_in, %bi0, %c0 : index + %a_na, %b_na, %o_na = buffer.assume.noalias %gate, %up, %output : buffer, buffer, buffer + %a_view = buffer.view %a_na[%zero] : buffer -> view<[%ae]xf32> + %b_view = buffer.view %b_na[%zero] : buffer -> view<[%be]xf32> + %o_view = buffer.view %o_na[%zero] : buffer -> view<[%elements]xf32> + %x = view.load %a_view[%ai] : view<[%ae]xf32> -> f32 + %u = view.load %b_view[%bi] : view<[%be]xf32> -> f32 + %x2 = scalar.mulf %x, %x : f32 + %poly = scalar.fmaf %k1, %x2, %one_f : f32 + %inner0 = scalar.mulf %x, %poly : f32 + %inner = scalar.mulf %k0, %inner0 : f32 + %twice = scalar.mulf %inner, %two : f32 + %e = scalar.expf %twice : f32 + %e1 = scalar.addf %e, %one_f : f32 + %frac = scalar.divf %two, %e1 : f32 + %tanh = scalar.subf %one_f, %frac : f32 + %onep = scalar.addf %one_f, %tanh : f32 + %hx = scalar.mulf %half, %x : f32 + %g = scalar.mulf %hx, %onep : f32 + %result = scalar.mulf %g, %u : f32 + scf.if %valid { + view.store %result, %o_view[%lin] : f32, view<[%elements]xf32> + } + kernel.return +}