From 1f09354cd4df19052d6894a12ef67ecb3e471a77 Mon Sep 17 00:00:00 2001 From: bong-water-water-bong Date: Sat, 26 Sep 2026 19:41:47 -0300 Subject: [PATCH 1/2] ggml-hrx: grouped F16 matmul and short-row kernels (softmax, sum_rows, argsort, narrow get_rows, strided copy, broadcast repeat) New loom kernels with their dispatch registrations, for ops HRX sent to the CPU: - ggml_grouped_mul_mat_f16_f32: MUL_MAT with a batched F16 weight (one matrix per group, no broadcast), such as ZAYA's grouped convolution. This also lets the loader place such weights on HRX. - ggml_softmax_rows_f32 (no mask, scale 1), ggml_sum_rows_f32, ggml_argsort_rows_f32 (rank per element, ties by index), ggml_get_rows_small_f32 (rows narrower than 4 or not a multiple of 4; strided ids, as top-k views are), ggml_copy_strided_f32 (CONT of strided views, and REPEAT that only broadcasts, with zero strides). Each matcher claims only cases the existing kernels do not take. AMD files get one-line hookups: CMakeLists, dispatch-common, REPEAT in the declared ops. test-backend-ops -b HRX0, all against the CPU: MUL_MAT (batched F16 cases), SOFT_MAX 10/10, SUM_ROWS 6/6, ARGSORT 48/48, GET_ROWS 17/17, CONT 2/2, CPY 53/53, REPEAT 5/5. --- ggml/src/ggml-hrx/CMakeLists.txt | 4 + .../common/dispatch-common.cpp | 4 + .../common/dispatch-grouped-mul-mat.cpp | 94 +++++ .../common/dispatch-grouped-mul-mat.h | 25 ++ .../common/dispatch-small-rows.cpp | 275 +++++++++++++ .../common/dispatch-small-rows.h | 26 ++ ggml/src/ggml-hrx/ggml-hrx.cpp | 1 + .../kernels/loom-libs/manifest.json | 372 ++++++++++++++++++ .../ops/grouped_mul_mat_f16_f32.loom | 73 ++++ .../kernels/loom-libs/ops/small_rows_f32.loom | 267 +++++++++++++ 10 files changed, 1141 insertions(+) create mode 100644 ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.cpp create mode 100644 ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.h create mode 100644 ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp create mode 100644 ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h create mode 100644 ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/grouped_mul_mat_f16_f32.loom create mode 100644 ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom diff --git a/ggml/src/ggml-hrx/CMakeLists.txt b/ggml/src/ggml-hrx/CMakeLists.txt index 3e1b7533c352..54fd1e2b95a0 100644 --- a/ggml/src/ggml-hrx/CMakeLists.txt +++ b/ggml/src/ggml-hrx/CMakeLists.txt @@ -202,6 +202,8 @@ ggml_add_backend_library(ggml-hrx dispatch_registration/common/dispatch-get-rows.cpp dispatch_registration/common/dispatch-get-rows.h dispatch_registration/common/dispatch-glu.cpp + dispatch_registration/common/dispatch-grouped-mul-mat.cpp + dispatch_registration/common/dispatch-grouped-mul-mat.h dispatch_registration/common/dispatch-glu.h dispatch_registration/common/dispatch-mul-mat-id.cpp dispatch_registration/common/dispatch-mul-mat-id-common.h @@ -219,6 +221,8 @@ ggml_add_backend_library(ggml-hrx dispatch_registration/common/dispatch-rmsnorm.cpp dispatch_registration/common/dispatch-rmsnorm.h dispatch_registration/common/dispatch-scale.cpp + dispatch_registration/common/dispatch-small-rows.cpp + dispatch_registration/common/dispatch-small-rows.h dispatch_registration/common/dispatch-scale.h dispatch_registration/common/dispatch-unary.cpp dispatch_registration/common/dispatch-unary.h diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.cpp index 9ccaebe22270..b53bacbd1bf0 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.cpp +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-common.cpp @@ -13,6 +13,8 @@ #include "dispatch-rmsnorm.h" #include "dispatch-rope-set-rows.h" #include "dispatch-scale.h" +#include "dispatch-grouped-mul-mat.h" +#include "dispatch-small-rows.h" #include "dispatch-unary.h" namespace ggml::hrx { @@ -24,12 +26,14 @@ void register_common_dispatches(DispatchRegistryBuilder & registry) { register_gated_mul_mat_id_dispatches(registry); register_gated_mul_mat_dispatches(registry); register_gather_add_dispatch(registry); + register_grouped_mul_mat_dispatch(registry); register_get_rows_dispatches(registry); register_glu_dispatches(registry); register_mul_mat_id_dispatches(registry); register_mul_mat_dispatches(registry); register_rope_set_rows_dispatches(registry); register_scale_dispatch(registry); + register_small_rows_dispatches(registry); register_unary_dispatch(registry); register_rmsnorm_dispatches(registry); } diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.cpp new file mode 100644 index 000000000000..f3955e098b26 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.cpp @@ -0,0 +1,94 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#include "dispatch-grouped-mul-mat.h" + +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kGroupedMulMatF16F32Kernel = + GGML_HRX_KERNEL_REF("loom_libs", "ggml_grouped_mul_mat_f16_f32"); + +// ne[0..2] as given, ne[3] == 1, and densely packed with the given element size +static bool packed_3d(const Value & value, int64_t ne0, int64_t ne1, int64_t ne2, size_t element_size) { + if (value.ne[0] != ne0 || value.ne[1] != ne1 || value.ne[2] != ne2 || value.ne[3] != 1) { + return false; + } + return value.nb[0] == element_size && value.nb[1] == element_size * static_cast(ne0) && + value.nb[2] == value.nb[1] * static_cast(ne1); +} + +static bool match_grouped_mul_mat_dispatch(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 * weight = context.graph.values().find(node->inputs[0]); + const Value * input = context.graph.values().find(node->inputs[1]); + const Value * output = context.graph.values().find(node->output); + if (weight == nullptr || input == nullptr || output == nullptr) { + return false; + } + if (weight->type != GGML_TYPE_F16 || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32) { + return false; + } + const int64_t k = weight->ne[0]; + const int64_t n = weight->ne[1]; + const int64_t g = weight->ne[2]; + const int64_t m = input->ne[1]; + // only the batched case: 2-D weights belong to the regular MUL_MAT kernels + if (g < 2 || g > 4096 || k < 1 || k > 65536 || n < 1 || n > 65536 || m < 1 || m > 65535) { + return false; + } + if (!packed_3d(*weight, k, n, g, sizeof(uint16_t)) || !packed_3d(*input, k, m, g, sizeof(float)) || + !packed_3d(*output, n, m, g, sizeof(float)) || output->alias_source.value >= 0 || + input->storage == output->storage || weight->storage == output->storage) { + return false; + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGroupedMulMatF16F32Kernel); + dispatch.kernel.integer_parameters.emplace("input_size", k); + dispatch.kernel.integer_parameters.emplace("output_size", n); + dispatch.kernel.integer_parameters.emplace("token_count", m); + dispatch.kernel.integer_parameters.emplace("group_count", g); + dispatch.bindings.push_back({ weight->id, 0, weight->byte_count }); + dispatch.bindings.push_back({ input->id, 0, input->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace + +void register_grouped_mul_mat_dispatch(DispatchRegistryBuilder & registry) { + registry.add({ + "common.grouped_mul_mat_f16_f32", + GGML_OP_MUL_MAT, + DispatchMatchKind::SingleOp, + 0, + DispatchSource::Common, + match_grouped_mul_mat_dispatch, + }); +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.h new file mode 100644 index 000000000000..39c65942472f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-grouped-mul-mat.h @@ -0,0 +1,25 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +// MUL_MAT with a batched F16 weight (one matrix per group, no broadcast), such as ZAYA's +// grouped convolution: ggml_grouped_mul_mat_f16_f32. +void register_grouped_mul_mat_dispatch(DispatchRegistryBuilder & registry); + +} // namespace ggml::hrx 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 new file mode 100644 index 000000000000..0c981eb81657 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp @@ -0,0 +1,275 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#include "dispatch-small-rows.h" + +#include "ggml.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include + +namespace ggml::hrx { +namespace { + +static constexpr KernelCatalogRef kSoftmaxRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_softmax_rows_f32"); +static constexpr KernelCatalogRef kSumRowsKernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_sum_rows_f32"); +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 bool packed(const Value & value, size_t element_size) { + size_t stride = element_size; + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (value.ne[i] <= 0 || value.nb[i] != stride) { + return false; + } + stride *= static_cast(value.ne[i]); + } + return true; +} + +static int64_t rows_of(const Value & value) { + return value.ne[1] * value.ne[2] * value.ne[3]; +} + +static bool same_shape(const Value & a, const Value & b) { + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + if (a.ne[i] != b.ne[i]) { + return false; + } + } + return true; +} + +// the single input and the output of a row op, both packed F32 (the output may be I32) +static bool row_op_values(const DispatchMatchContext & context, ggml_op op, ggml_type output_type, + const Value *& input, const Value *& output) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != op || node->inputs.size() != 1) { + return false; + } + input = context.graph.values().find(node->inputs[0]); + output = context.graph.values().find(node->output); + return input != nullptr && output != nullptr && input->type == GGML_TYPE_F32 && output->type == output_type && + packed(*input, sizeof(float)) && packed(*output, ggml_type_size(output_type)) && + output->alias_source.value < 0 && input->storage != output->storage && rows_of(*input) <= 16777216; +} + +static void finish(const DispatchMatchContext & context, DispatchMatch & match, Dispatch && dispatch) { + match.covered_nodes.push_back(context.root_index); + match.dispatches.push_back(std::move(dispatch)); +} + +static bool match_softmax_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_SOFT_MAX, GGML_TYPE_F32, input, output) || !same_shape(*input, *output) || + input->ne[0] > 4096) { + return false; + } + const SoftMaxParams * params = op_params_as(context.root_node->params); + if (params == nullptr || params->scale != 1.0f || params->max_bias != 0.0f) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSoftmaxRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + 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; +} + +static bool match_sum_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_SUM_ROWS, GGML_TYPE_F32, input, output) || input->ne[0] > 4096 || + output->ne[0] != 1 || output->ne[1] != input->ne[1] || output->ne[2] != input->ne[2] || + output->ne[3] != input->ne[3]) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kSumRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + 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; +} + +static bool match_argsort_rows(const DispatchMatchContext & context, DispatchMatch & match) { + const Value * input = nullptr; + const Value * output = nullptr; + if (!row_op_values(context, GGML_OP_ARGSORT, GGML_TYPE_I32, input, output) || !same_shape(*input, *output) || + input->ne[0] > 1024) { + return false; + } + const ArgsortParams * params = op_params_as(context.root_node->params); + if (params == nullptr) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kArgsortRowsKernel); + dispatch.kernel.integer_parameters.emplace("column_count", input->ne[0]); + dispatch.kernel.integer_parameters.emplace("row_count", rows_of(*input)); + dispatch.kernel.integer_parameters.emplace("descending", params->order == GGML_SORT_ORDER_DESC ? 1 : 0); + 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). +static bool match_get_rows_small(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_GET_ROWS || node->inputs.size() != 2) { + return false; + } + const Value * source = context.graph.values().find(node->inputs[0]); + const Value * ids = context.graph.values().find(node->inputs[1]); + const Value * output = context.graph.values().find(node->output); + if (source == nullptr || ids == nullptr || output == nullptr || source->type != GGML_TYPE_F32 || + ids->type != GGML_TYPE_I32 || output->type != GGML_TYPE_F32 || !packed(*source, sizeof(float)) || + !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || ids->nb[0] != sizeof(int32_t) || + ids->nb[1] % sizeof(int32_t) != 0 || static_cast(ids->nb[1] / sizeof(int32_t)) < ids->ne[0]) { + return false; + } + const int64_t w = source->ne[0]; + const int64_t s = source->ne[1]; + const int64_t b = source->ne[2]; + const int64_t r = ids->ne[0]; + if ((w >= 4 && w % 4 == 0) || w > 65536 || s > 65536 || b > 65536 || r > 65536 || source->ne[3] != 1 || + ids->ne[1] != b || ids->ne[2] != 1 || ids->ne[3] != 1 || output->ne[0] != w || output->ne[1] != r || + output->ne[2] != b || output->ne[3] != 1) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kGetRowsSmallKernel); + dispatch.kernel.integer_parameters.emplace("width", w); + dispatch.kernel.integer_parameters.emplace("id_count", r); + dispatch.kernel.integer_parameters.emplace("batch_count", b); + dispatch.kernel.integer_parameters.emplace("source_rows", s); + dispatch.kernel.integer_parameters.emplace("id_stride", static_cast(ids->nb[1] / sizeof(int32_t))); + dispatch.bindings.push_back({ source->id, 0, source->byte_count }); + dispatch.bindings.push_back({ ids->id, 0, ids->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + finish(context, match, std::move(dispatch)); + return true; +} + + +// CONT of a strided F32 view (permuted, transposed or sliced) into a packed output. Only sources +// that are not packed: packed copies belong to the regular copy kernel. +static bool match_copy_strided(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_CONT || 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); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + packed(*input, sizeof(float)) || !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || + input->storage == output->storage || input->element_count != output->element_count) { + 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 || output->element_count > 268435456) { + return false; + } + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kCopyStridedKernel); + 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; +} + + +// REPEAT that only broadcasts (every input dim is 1 or the output's): the strided copy with a zero +// stride on the broadcast dims. Tiling repeats (output a multiple of a larger input) are not claimed. +static bool match_repeat_broadcast(const DispatchMatchContext & context, DispatchMatch & match) { + const GraphNode * node = context.root_node; + if (node == nullptr || node->op != GGML_OP_REPEAT || 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); + if (input == nullptr || output == nullptr || input->type != GGML_TYPE_F32 || output->type != GGML_TYPE_F32 || + !packed(*output, sizeof(float)) || output->alias_source.value >= 0 || input->storage == output->storage || + 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] != 1 && input->ne[i] != output->ne[i])) { + return false; + } + strides[i] = input->ne[i] == 1 ? 0 : 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(kCopyStridedKernel); + 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; +} + +} // 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.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, + match_argsort_rows }); + registry.add({ "common.get_rows_small_f32", GGML_OP_GET_ROWS, DispatchMatchKind::SingleOp, 0, + DispatchSource::Common, match_get_rows_small }); + registry.add({ "common.copy_strided_f32", GGML_OP_CONT, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_copy_strided }); + registry.add({ "common.repeat_broadcast_f32", GGML_OP_REPEAT, DispatchMatchKind::SingleOp, 0, DispatchSource::Common, + match_repeat_broadcast }); +} + +} // namespace ggml::hrx 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 new file mode 100644 index 000000000000..6bb09ffca247 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.h @@ -0,0 +1,26 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +#pragma once + +#include "dispatch_registration/dispatch-registry.h" + +namespace ggml::hrx { + +// SOFT_MAX (no mask, scale 1), SUM_ROWS, 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. +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 158def6bcc23..e4486ecc4ef8 100644 --- a/ggml/src/ggml-hrx/ggml-hrx.cpp +++ b/ggml/src/ggml-hrx/ggml-hrx.cpp @@ -621,6 +621,7 @@ static bool eager_capability_declared(enum ggml_op op) { case GGML_OP_MUL_MAT: case GGML_OP_MUL_MAT_ID: case GGML_OP_PERMUTE: + case GGML_OP_REPEAT: // broadcast-only, through the strided copy (dispatch-small-rows.cpp) case GGML_OP_RESHAPE: case GGML_OP_RMS_NORM: case GGML_OP_ROPE: 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 8653c2138d2c..f9ba0b1336f5 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 @@ -220,6 +220,12 @@ }, { "path": "motifs/f32_f16.loom" + }, + { + "path": "ops/grouped_mul_mat_f16_f32.loom" + }, + { + "path": "ops/small_rows_f32.loom" } ], "exports": [ @@ -5763,6 +5769,372 @@ "motifs/unary_f32_apply.loom", "motifs/quantize_q8_1_x4.loom" ] + }, + { + "name": "ggml_grouped_mul_mat_f16_f32", + "family": "loom_libs", + "symbol": "ggml_grouped_mul_mat_f16_f32", + "source": "ops/grouped_mul_mat_f16_f32.loom", + "workload_parameters": [ + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + }, + { + "name": "group_count", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "input_size", + "type": "index" + }, + { + "name": "output_size", + "type": "index" + }, + { + "name": "token_count", + "type": "index" + }, + { + "name": "group_count", + "type": "index" + } + ], + "bindings": [ + "weight", + "input", + "output" + ], + "binding_access": [ + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/grouped_mul_mat_f16_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] + }, + { + "name": "ggml_softmax_rows_f32", + "family": "loom_libs", + "symbol": "ggml_softmax_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_sum_rows_f32", + "family": "loom_libs", + "symbol": "ggml_sum_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_argsort_rows_f32", + "family": "loom_libs", + "symbol": "ggml_argsort_rows_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "descending", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "column_count", + "type": "index" + }, + { + "name": "row_count", + "type": "index" + }, + { + "name": "descending", + "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_get_rows_small_f32", + "family": "loom_libs", + "symbol": "ggml_get_rows_small_f32", + "source": "ops/small_rows_f32.loom", + "workload_parameters": [ + { + "name": "width", + "type": "index" + }, + { + "name": "id_count", + "type": "index" + }, + { + "name": "batch_count", + "type": "index" + }, + { + "name": "source_rows", + "type": "index" + }, + { + "name": "id_stride", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "width", + "type": "index" + }, + { + "name": "id_count", + "type": "index" + }, + { + "name": "batch_count", + "type": "index" + }, + { + "name": "source_rows", + "type": "index" + }, + { + "name": "id_stride", + "type": "index" + } + ], + "bindings": [ + "input", + "ids", + "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_copy_strided_f32", + "family": "loom_libs", + "symbol": "ggml_copy_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": "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": [] } ], "link_modules": [], diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/grouped_mul_mat_f16_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/grouped_mul_mat_f16_f32.loom new file mode 100644 index 000000000000..3042d9e87b05 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/grouped_mul_mat_f16_f32.loom @@ -0,0 +1,73 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Batched ("grouped") MUL_MAT with F16 weights: for every group g and token m, +// output[g][m][n] = sum_k weight[g][n][k] * input[g][m][k] (F32 accumulation) +// The weight has one [N x K] matrix per group and no broadcast. ZAYA's CCA grouped +// convolution is one of these per tap (10 groups of 128 x 128). +// One workitem computes one output element; workgroups tile (N / 64, M, G). + +amdgpu.target @ggml_grouped_mul_mat_f16_f32_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_grouped_mul_mat_f16_f32_gfx11_wave64) export("ggml_grouped_mul_mat_f16_f32") @ggml_grouped_mul_mat_f16_f32(%input_size: index, %output_size: index, %token_count: index, %group_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %output_size, %rounding : index + %tiles = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%tiles, %token_count, %group_count) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%input_size: index, %output_size: index, %token_count: index, %group_count: index, %weight: buffer, %input: buffer, %output: buffer) { + %k = index.assume %input_size [range(%input_size, 1, 65536)] : index + %n = index.assume %output_size [range(%output_size, 1, 65536)] : index + %m = index.assume %token_count [range(%token_count, 1, 65536)] : index + %g = index.assume %group_count [range(%group_count, 1, 4096)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c0_f32 = scalar.constant 0.0 : f32 + %zero_offset = index.constant 0 : offset + %tile = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %group0 = kernel.workgroup.id : index + %lane0 = kernel.workitem.id : index + %lane = index.assume %lane0 [range(%lane0, 0, 63)] : index + %token = index.assume %token0 [range(%token0, 0, 65535)] : index + %group = index.assume %group0 [range(%group0, 0, 4095)] : index + %tile_base = index.mul %tile, %c64 : index + %column0 = index.add %tile_base, %lane : index + %valid = index.cmp ult, %column0, %n : index + %column = scf.select %valid, %column0, %c0 : index + %weight_rows = index.mul %g, %n : index + %input_rows = index.mul %g, %m : index + %weight_row_base = index.mul %group, %n : index + %weight_row = index.add %weight_row_base, %column : index + %input_row_base = index.mul %group, %m : index + %input_row = index.add %input_row_base, %token : index + %weight_noalias, %input_noalias, %output_noalias = buffer.assume.noalias %weight, %input, %output : buffer, buffer, buffer + %weight_view = buffer.view %weight_noalias[%zero_offset] : buffer -> view<[%weight_rows]x[%k]xf16> + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%input_rows]x[%k]xf32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%input_rows]x[%n]xf32> + %sum = scf.for %channel = [%c0 to %k step %c1](%accumulator = %c0_f32 : f32) -> (f32) { + %weight_f16 = view.load %weight_view[%weight_row, %channel] : view<[%weight_rows]x[%k]xf16> -> f16 + %weight_value = scalar.extf %weight_f16 : f16 to f32 + %input_value = view.load %input_view[%input_row, %channel] : view<[%input_rows]x[%k]xf32> -> f32 + %next_accumulator = scalar.fmaf %weight_value, %input_value, %accumulator : f32 + scf.yield %next_accumulator : f32 + } + scf.if %valid { + view.store %sum, %output_view[%input_row, %column] : f32, view<[%input_rows]x[%n]xf32> + } + kernel.return +} 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 new file mode 100644 index 000000000000..9e055426900c --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/small_rows_f32.loom @@ -0,0 +1,267 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// 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 +// 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. + +amdgpu.target @ggml_small_rows_f32_gfx11_wave64 {subgroup_size = 64} + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_softmax_rows_f32") @ggml_softmax_rows_f32(%column_count: index, %row_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %row_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%column_count: index, %row_count: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 4096)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c0_f32 = scalar.constant 0.0 : f32 + %lowest = scalar.constant -3.40282347e+38 : f32 + %zero_offset = index.constant 0 : offset + %group = kernel.workgroup.id : index + %lane = kernel.workitem.id : index + %base = index.mul %group, %c64 : index + %row0 = index.add %base, %lane : index + %valid = index.cmp ult, %row0, %rows : index + %row = scf.select %valid, %row0, %c0 : 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> + %maximum = scf.for %column = [%c0 to %cols step %c1](%running = %lowest : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %next = scalar.maxnumf %running, %value : f32 + scf.yield %next : f32 + } + %total = scf.for %column = [%c0 to %cols step %c1](%running = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %shifted = scalar.subf %value, %maximum : f32 + %exponential = scalar.expf %shifted : f32 + %next = scalar.addf %running, %exponential : f32 + scf.yield %next : f32 + } + scf.if %valid { + %unused = scf.for %column = [%c0 to %cols step %c1](%carry = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %shifted = scalar.subf %value, %maximum : f32 + %exponential = scalar.expf %shifted : f32 + %probability = scalar.divf %exponential, %total : f32 + view.store %probability, %output_view[%row, %column] : f32, view<[%rows]x[%cols]xf32> + scf.yield %carry : f32 + } + } + kernel.return +} + +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_sum_rows_f32") @ggml_sum_rows_f32(%column_count: index, %row_count: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %rounded = index.add %row_count, %rounding : index + %groups = index.div %rounded, %sixtyfour : index + kernel.launch.config workgroups(%groups, %one, %one) workgroup_size(%sixtyfour, %one, %one) : index +} launch(%column_count: index, %row_count: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 4096)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c64 = index.constant 64 : index + %c0_f32 = 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 + %row0 = index.add %base, %lane : index + %valid = index.cmp ult, %row0, %rows : index + %row = scf.select %valid, %row0, %c0 : 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]xf32> + %total = scf.for %column = [%c0 to %cols step %c1](%running = %c0_f32 : f32) -> (f32) { + %value = view.load %input_view[%row, %column] : view<[%rows]x[%cols]xf32> -> f32 + %next = scalar.addf %running, %value : f32 + scf.yield %next : f32 + } + scf.if %valid { + view.store %total, %output_view[%row] : f32, view<[%rows]xf32> + } + kernel.return +} + +// ARGSORT: element i of a row goes to position rank(i), where rank counts the elements that +// sort before it (larger ones for descending order, smaller for ascending; ties by index). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_argsort_rows_f32") @ggml_argsort_rows_f32(%column_count: index, %row_count: index, %descending: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %elements = index.mul %column_count, %row_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(%column_count: index, %row_count: index, %descending: index, %input: buffer, %output: buffer) { + %cols = index.assume %column_count [range(%column_count, 1, 1024)] : index + %rows = index.assume %row_count [range(%row_count, 1, 16777216)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : 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 + %elements = index.mul %cols, %rows : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %row = index.div %linear, %cols : index + %element = index.rem %linear, %cols : index + %want_descending = index.cmp ne, %descending, %c0 : 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]xi32> + %mine = view.load %input_view[%row, %element] : view<[%rows]x[%cols]xf32> -> f32 + %rank = scf.for %other = [%c0 to %cols step %c1](%count = %c0 : index) -> (index) { + %value = view.load %input_view[%row, %other] : view<[%rows]x[%cols]xf32> -> f32 + %greater = scalar.cmpf ogt, %value, %mine : f32 + %less = scalar.cmpf olt, %value, %mine : f32 + %equal = scalar.cmpf oeq, %value, %mine : f32 + %earlier = index.cmp ult, %other, %element : index + %tie_before = scalar.andi %equal, %earlier : i1 + %strict = scf.select %want_descending, %greater, %less : i1 + %before = scalar.ori %strict, %tie_before : i1 + %step = scf.select %before, %c1, %c0 : index + %next = index.add %count, %step : index + scf.yield %next : index + } + scf.if %valid { + %element_i32 = index.cast %element : index to i32 + view.store %element_i32, %output_view[%row, %rank] : i32, view<[%rows]x[%cols]xi32> + } + kernel.return +} + +// GET_ROWS within batches, for narrow rows: output[b][r][c] = input[b][ids[b][r]][c]. The ids may be a +// view with a row stride (the first k columns of an ARGSORT, as top-k selection makes). +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_get_rows_small_f32") @ggml_get_rows_small_f32(%width: index, %id_count: index, %batch_count: index, %source_rows: index, %id_stride: index) { + %one = index.constant 1 : index + %sixtyfour = index.constant 64 : index + %rounding = index.constant 63 : index + %per_batch = index.mul %width, %id_count : index + %elements = index.mul %per_batch, %batch_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(%width: index, %id_count: index, %batch_count: index, %source_rows: index, %id_stride: index, %input: buffer, %ids: buffer, %output: buffer) { + %w = index.assume %width [range(%width, 1, 65536)] : index + %r = index.assume %id_count [range(%id_count, 1, 65536)] : index + %b = index.assume %batch_count [range(%batch_count, 1, 65536)] : index + %s = index.assume %source_rows [range(%source_rows, 1, 65536)] : index + %t = index.assume %id_stride [range(%id_stride, 1, 1048576)] : index + %c0 = index.constant 0 : index + %c64 = index.constant 64 : index + %c0_i32 = scalar.constant 0 : i32 + %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 + %per_batch = index.mul %w, %r : index + %elements = index.mul %per_batch, %b : index + %valid = index.cmp ult, %linear0, %elements : index + %linear = scf.select %valid, %linear0, %c0 : index + %batch = index.div %linear, %per_batch : index + %within = index.rem %linear, %per_batch : index + %slot = index.div %within, %w : index + %column = index.rem %within, %w : index + %source_total = index.mul %s, %b : index + %id_total = index.mul %r, %b : index + %output_total = index.mul %id_total, %w : index + %input_noalias, %ids_noalias, %output_noalias = buffer.assume.noalias %input, %ids, %output : buffer, buffer, buffer + %input_view = buffer.view %input_noalias[%zero_offset] : buffer -> view<[%source_total]x[%w]xf32> + %ids_view = buffer.view %ids_noalias[%zero_offset] : buffer -> view<[%b]x[%t]xi32> + %output_view = buffer.view %output_noalias[%zero_offset] : buffer -> view<[%output_total]xf32> + %id_raw = view.load %ids_view[%batch, %slot] : view<[%b]x[%t]xi32> -> i32 + %id_nonnegative = scalar.cmpi sge, %id_raw, %c0_i32 : i32 + %id_safe_i32 = scf.select %id_nonnegative, %id_raw, %c0_i32 : i32 + %id0 = index.cast %id_safe_i32 : i32 to index + %in_range = index.cmp ult, %id0, %s : index + %id = scf.select %in_range, %id0, %c0 : index + %source_base = index.mul %batch, %s : index + %source_row = index.add %source_base, %id : index + %value = view.load %input_view[%source_row, %column] : view<[%source_total]x[%w]xf32> -> f32 + scf.if %valid { + view.store %value, %output_view[%linear] : f32, view<[%output_total]xf32> + } + kernel.return +} + +// CONT of a strided (permuted, transposed or sliced) F32 view into a packed output: one workitem +// per output element, reading the source through its four element strides. +kernel.def target(@ggml_small_rows_f32_gfx11_wave64) export("ggml_copy_strided_f32") @ggml_copy_strided_f32(%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]xf32> + %value = view.load %input_view[%source] : view<[%extent]xf32> -> f32 + scf.if %valid { + view.store %value, %output_view[%linear] : f32, view<[%elements]xf32> + } + kernel.return +} From f56203e110788156f9e62bdd70457b8f42cfab6d Mon Sep 17 00:00:00 2001 From: bong-water-water-bong Date: Sat, 26 Sep 2026 19:41:47 -0300 Subject: [PATCH 2/2] zaya: decode graph that runs entirely on HRX - One token per sequence: the conv input transpose is a reshape, not a copy. - MEAN over the query-head group as SUM_ROWS + SCALE. - Qpre/Kpre reshape the contiguous matmul outputs instead of copying them. With the block entirely on HRX the copy made prefill wrong (perplexity 71.8). With the HRX kernels of the previous commit, ZAYA's decode graph on HRX0 goes from 641 graph splits to 1. --- src/models/zaya.cpp | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/src/models/zaya.cpp b/src/models/zaya.cpp index d952d646ef29..9738ff86116f 100644 --- a/src/models/zaya.cpp +++ b/src/models/zaya.cpp @@ -314,8 +314,12 @@ llama_model_zaya::graph::graph(const llama_model & model, const llm_graph_ ggml_tensor * QKraw = ggml_concat(ctx0, Qraw, Kraw, 0); cb(QKraw, "QKraw", il); - ggml_tensor * Qpre = ggml_reshape_3d(ctx0, ggml_cont(ctx0, Qraw), n_embd_head, n_head, n_tokens); - ggml_tensor * Kpre = ggml_reshape_3d(ctx0, ggml_cont(ctx0, Kraw), n_embd_head, n_head_kv, n_tokens); + // Qraw and Kraw are fresh matmul outputs, already contiguous: reshape them in place. A CONT + // here copies for nothing, and with this block entirely on HRX the copied K read back wrong + // in 512-token batches (wikitext perplexity 71.8 instead of 21.6; right with a CPU split + // after the copy, or without the copy). Decode was unaffected. + ggml_tensor * Qpre = ggml_reshape_3d(ctx0, ggml_is_contiguous(Qraw) ? Qraw : ggml_cont(ctx0, Qraw), n_embd_head, n_head, n_tokens); + ggml_tensor * Kpre = ggml_reshape_3d(ctx0, ggml_is_contiguous(Kraw) ? Kraw : ggml_cont(ctx0, Kraw), n_embd_head, n_head_kv, n_tokens); ggml_tensor * Kpre_grouped = ggml_reshape_4d(ctx0, Kpre, n_embd_head, 1, n_head_kv, n_tokens); Kpre_grouped = ggml_repeat_4d(ctx0, Kpre_grouped, n_embd_head, n_gqa, n_head_kv, n_tokens); @@ -326,7 +330,7 @@ llama_model_zaya::graph::graph(const llama_model & model, const llm_graph_ ggml_tensor * Qgroup = ggml_reshape_4d(ctx0, Qpre, n_embd_head, n_gqa, n_head_kv, n_tokens); Qgroup = ggml_permute(ctx0, Qgroup, 1, 0, 2, 3); Qgroup = ggml_cont(ctx0, Qgroup); - ggml_tensor * Qmean = ggml_mean(ctx0, Qgroup); + ggml_tensor * Qmean = ggml_scale(ctx0, ggml_sum_rows(ctx0, Qgroup), 1.0f/n_gqa); // MEAN, as SUM_ROWS + SCALE Qmean = ggml_reshape_3d(ctx0, Qmean, n_embd_head, n_head_kv, n_tokens); ggml_tensor * qk_mean_k = ggml_scale(ctx0, ggml_add(ctx0, Qmean, Kpre), 0.5f); cb(qk_mean_k, "qk_mean_k", il); @@ -334,7 +338,9 @@ llama_model_zaya::graph::graph(const llama_model & model, const llm_graph_ // [n_qk, T, S] -> [T, n_qk, S]: split the sequences before transposing, or with more than // one sequence in the ubatch a channel's row would run across all of them ggml_tensor * QKraw_t = ggml_reshape_3d(ctx0, QKraw, n_qk, n_seq_tokens, n_seqs); - QKraw_t = ggml_cont(ctx0, ggml_transpose(ctx0, QKraw_t)); + // with one token per sequence the transpose moves no data: a reshape does it without a copy + QKraw_t = n_seq_tokens == 1 ? ggml_reshape_3d(ctx0, QKraw_t, 1, n_qk, n_seqs) + : ggml_cont(ctx0, ggml_transpose(ctx0, QKraw_t)); ggml_tensor * conv_input = ggml_concat(ctx0, conv_state, QKraw_t, 0); cb(conv_input, "cca_conv_input", il);