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 +} 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);