diff --git a/ggml/src/ggml-hrx/CMakeLists.txt b/ggml/src/ggml-hrx/CMakeLists.txt index c16c10c0136f..d10f2375afa0 100644 --- a/ggml/src/ggml-hrx/CMakeLists.txt +++ b/ggml/src/ggml-hrx/CMakeLists.txt @@ -243,6 +243,8 @@ ggml_add_backend_library(ggml-hrx dispatch_registration/common/dispatch-swiglu-oai.h dispatch_registration/common/moe-placement-guard.cpp dispatch_registration/common/moe-placement-guard.h + dispatch_registration/common/dispatch-attention-sink.cpp + dispatch_registration/common/dispatch-attention-sink.h dispatch_registration/common/dispatch-softplus.cpp dispatch_registration/common/dispatch-softplus.h dispatch_registration/common/dispatch-scale.h diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.cpp new file mode 100644 index 000000000000..da65f4bf4c8f --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.cpp @@ -0,0 +1,125 @@ +// 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. + +// Attention sinks for FLASH_ATTN_EXT (gpt-oss): the FlashAttention dispatch runs without the sink and +// ops/attention_sink_f32.loom then rescales its output in place, row by row, by S / (S + exp(sink - M)), which is +// exact (see the kernel). dispatch-flash-attention.cpp accepts the 5-input node when attention_sinks_supported +// holds and appends this dispatch after its own. + +#include "dispatch-attention-sink.h" + +#include "dispatch-mul-mat-common.h" +#include "graph/op-params.h" +#include "kernel-corpus/kernel-corpus-catalog-verify.h" + +#include +#include +#include + +namespace ggml::hrx { + +namespace { + +static constexpr KernelCatalogRef kAttentionSinkF32Kernel = GGML_HRX_KERNEL_REF("loom_libs", "ggml_attention_sink_f32"); + +std::string index_config(int64_t value) { + return std::to_string(value); +} + +// True when the two values share storage and their byte ranges intersect. +bool overlaps(const Graph & graph, const Value & lhs, const Value & rhs) { + if (lhs.id == rhs.id) { + return true; + } + if (!graph.values().same_storage(lhs.id, rhs.id)) { + return false; + } + const size_t lhs_end = lhs.storage_offset + lhs.byte_count; + const size_t rhs_end = rhs.storage_offset + rhs.byte_count; + return lhs.storage_offset < rhs_end && rhs.storage_offset < lhs_end; +} + +} // namespace + +bool attention_sinks_supported(const Graph & graph, const GraphNode & node) { + if (node.op != GGML_OP_FLASH_ATTN_EXT || node.inputs.size() != 5) { + return false; + } + const Value * query = graph.values().find(node.inputs[0]); + const Value * sinks = graph.values().find(node.inputs[4]); + if (query == nullptr || sinks == nullptr || sinks->type != GGML_TYPE_F32 || !sinks->contiguous) { + return false; + } + const int64_t query_head_count = query->ne[2]; + return sinks->ne[0] == query_head_count && sinks->ne[1] == 1 && sinks->ne[2] == 1 && sinks->ne[3] == 1; +} + +bool append_attention_sink_dispatch(const Graph & graph, const GraphNode & node, DispatchMatch & dispatch_match) { + if (!attention_sinks_supported(graph, node)) { + return false; + } + const Value * query = graph.values().find(node.inputs[0]); + const Value * key = graph.values().find(node.inputs[1]); + const Value * mask = graph.values().find(node.inputs[3]); + const Value * sinks = graph.values().find(node.inputs[4]); + const Value * output = graph.values().find(node.output); + const FlashAttnExtParams * params = op_params_as(node.params); + if (key == nullptr || mask == nullptr || output == nullptr || params == nullptr) { + return false; + } + // The FlashAttention matchers have checked these layouts (query [tokens][heads][d] f32, key + // [capacity][kv_heads][d] f16, output [tokens][heads][dv] f32); the mask rows must be key_count apart. + const int64_t tokens = query->ne[1]; + const int64_t query_head_count = query->ne[2]; + const int64_t qk_head_size = query->ne[0]; + const int64_t key_capacity = key->ne[1]; + const int64_t key_value_head_count = key->ne[2]; + const int64_t value_head_size = output->ne[0]; + const int64_t key_count = mask->ne[0]; + if (mask->type != GGML_TYPE_F16 || mask->nb[0] != sizeof(ggml_fp16_t) || + mask->nb[1] != static_cast(key_count) * sizeof(ggml_fp16_t) || key_count < 1 || + key_count > key_capacity || query_head_count > 256 || key_value_head_count > 256 || qk_head_size % 16 != 0 || + value_head_size % 16 != 0 || qk_head_size > 576 || value_head_size > 576) { + return false; + } + // The output is rewritten in place after FlashAttention wrote it; no input may overlap its bytes (views of one + // allocation share a storage root but are distinct values, so compare storage, not value ids). + for (const Value * input : { query, key, mask, sinks }) { + if (overlaps(graph, *input, *output)) { + return false; + } + } + + Dispatch dispatch; + dispatch.kernel = make_kernel_specialization(kAttentionSinkF32Kernel); + dispatch.kernel.integer_parameters.emplace("tokens", tokens); + dispatch.kernel.integer_parameters.emplace("key_count", key_count); + dispatch.kernel.integer_parameters.emplace("key_capacity", key_capacity); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.query_head_count", index_config(query_head_count)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.key_value_head_count", + index_config(key_value_head_count)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.qk_head_size", index_config(qk_head_size)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.value_head_size", index_config(value_head_size)); + dispatch.kernel.compile_parameters.emplace("ggml.attention_sink.scale", common_to_config_value(params->scale)); + dispatch.bindings.push_back({ query->id, 0, query->byte_count }); + dispatch.bindings.push_back({ key->id, 0, key->byte_count }); + dispatch.bindings.push_back({ mask->id, 0, mask->byte_count }); + dispatch.bindings.push_back({ sinks->id, 0, sinks->byte_count }); + dispatch.bindings.push_back({ output->id, 0, output->byte_count }); + dispatch_match.dispatches.push_back(std::move(dispatch)); + return true; +} + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.h new file mode 100644 index 000000000000..75c91aa2fdb1 --- /dev/null +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-attention-sink.h @@ -0,0 +1,29 @@ +// 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" +#include "graph/graph.h" + +namespace ggml::hrx { + +// A 5-input FLASH_ATTN_EXT whose fifth input is F32 sinks, one per query head. +bool attention_sinks_supported(const Graph & graph, const GraphNode & node); + +// Appends the in-place sink rescale of the node's output (run after the FlashAttention dispatch). +bool append_attention_sink_dispatch(const Graph & graph, const GraphNode & node, DispatchMatch & dispatch_match); + +} // namespace ggml::hrx diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp index cfcfb942d750..221526e1de25 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-flash-attention.cpp @@ -1,3 +1,4 @@ +#include "dispatch-attention-sink.h" #include "dispatch-flash-attention.h" #include "ggml.h" @@ -210,7 +211,8 @@ static FlashAttentionMatch match_flash_attention_f32_f16(const Graph & gra const CommandPlan & plan, const GraphNode * node) { FlashAttentionMatch match; - if (node == nullptr || node->op != GGML_OP_FLASH_ATTN_EXT || node->inputs.size() != 4) { + if (node == nullptr || node->op != GGML_OP_FLASH_ATTN_EXT || + (node->inputs.size() != 4 && !attention_sinks_supported(graph, *node))) { return match; } @@ -463,7 +465,8 @@ static DispatchBinding prepare_flash_attention_value(const DispatchMatchContext static bool match_flash_attention_gate_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) { const FlashAttentionMatch match = match_flash_attention_f32_f16(context.graph, context.plan, context.root_node); - if (!match.matched() || match.output_layout == nullptr || match.output_layout->op != GGML_OP_RESHAPE) { + if (!match.matched() || context.root_node->inputs.size() != 4 || match.output_layout == nullptr || + match.output_layout->op != GGML_OP_RESHAPE) { return false; } @@ -574,6 +577,10 @@ static bool match_flash_attention_f32_f16_dispatch(const DispatchMatchContext & } } dispatch_match.dispatches.push_back(std::move(dispatch)); + if (context.root_node->inputs.size() == 5 && + !append_attention_sink_dispatch(context.graph, *context.root_node, dispatch_match)) { + return false; + } return true; } 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 7b732475e1e8..15030b241b02 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 @@ -259,6 +259,9 @@ }, { "path": "ops/swiglu_oai_f32.loom" + }, + { + "path": "ops/attention_sink_f32.loom" } ], "exports": [ @@ -8071,6 +8074,63 @@ "library_sources": [] }, "compile_dependencies": [] + }, + { + "name": "ggml_attention_sink_f32", + "family": "loom_libs", + "symbol": "ggml_attention_sink_f32", + "source": "ops/attention_sink_f32.loom", + "workload_parameters": [ + { + "name": "tokens", + "type": "index" + }, + { + "name": "key_count", + "type": "index" + }, + { + "name": "key_capacity", + "type": "index" + } + ], + "launch_parameters": [ + { + "name": "tokens", + "type": "index" + }, + { + "name": "key_count", + "type": "index" + }, + { + "name": "key_capacity", + "type": "index" + } + ], + "bindings": [ + "query", + "key", + "mask", + "sinks", + "output" + ], + "binding_access": [ + "read", + "read", + "read", + "read", + "read_write" + ], + "target_selector": "", + "compile_recipe": { + "mode": "direct", + "primary_sources": [ + "ops/attention_sink_f32.loom" + ], + "library_sources": [] + }, + "compile_dependencies": [] } ], "link_modules": [], diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/attention_sink_f32.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/attention_sink_f32.loom new file mode 100644 index 000000000000..8c477cef1610 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/attention_sink_f32.loom @@ -0,0 +1,329 @@ +// 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. + +// Attention sinks (gpt-oss) applied after FLASH_ATTN_EXT, in place on its output. A sink adds one logit per head to +// the softmax denominator and nothing to the numerator, so with M the row maximum and S the row sum of +// exp(scale q.k + mask - M): +// output_with_sink = output_without_sink * S / (S + exp(sink - M)) +// This kernel recomputes M and S for each (query token, query head) row with FlashAttention's conventions (scores +// scale * q.k + mask, mask entries below -1e30, i.e. -inf, skipped) and rescales the row. A row with every key +// masked has S = 0 and is written as zeros. +// Layouts (as the HRX FlashAttention matchers require): query [tokens][query_heads][qk_head_size] f32, key +// [key_capacity][key_value_heads][qk_head_size] f16, mask [tokens][key_count] f16, sinks [query_heads] f32, output +// [tokens][query_heads][value_head_size] f32. Query head h reads key-value head h / (query_heads / key_value_heads). + +amdgpu.target @ggml_attention_sink_gfx11_wave32 {subgroup_size = 32} + +config.decl @ggml.attention_sink.query_head_count : %value: index where [range(%value, 1, 256)] + +config.decl @ggml.attention_sink.key_value_head_count : %value: index where [range(%value, 1, 256)] + +config.decl @ggml.attention_sink.qk_head_size : %value: index where [range(%value, 16, 576), mul(%value, 16)] + +config.decl @ggml.attention_sink.value_head_size : %value: index where [range(%value, 16, 576), mul(%value, 16)] + +config.decl @ggml.attention_sink.scale : f32 + +kernel.def target(@ggml_attention_sink_gfx11_wave32) export("ggml_attention_sink_f32") @ggml_attention_sink_f32(%tokens: index, %key_count: index, %key_capacity: index) { + %query_head_count = config.get @ggml.attention_sink.query_head_count : index + %c1 = index.constant 1 : index + %c256 = index.constant 256 : index + kernel.launch.config workgroups(%query_head_count, %tokens, %c1) workgroup_size(%c256, %c1, %c1) : index +} launch(%tokens: index, %key_count: index, %key_capacity: index, %query: buffer, %key: buffer, %mask: buffer, %sinks: buffer, %output: buffer) where [range(%tokens, 1, 65536), range(%key_count, 1, 1048576), range(%key_capacity, 1, 1048576)] { + %query_head_count0 = config.get @ggml.attention_sink.query_head_count : index + %key_value_head_count0 = config.get @ggml.attention_sink.key_value_head_count : index + %qk_head_size0 = config.get @ggml.attention_sink.qk_head_size : index + %value_head_size0 = config.get @ggml.attention_sink.value_head_size : index + %scale = config.get @ggml.attention_sink.scale : f32 + %query_head_count = index.assume %query_head_count0 [range(%query_head_count0, 1, 256)] : index + %key_value_head_count = index.assume %key_value_head_count0 [range(%key_value_head_count0, 1, 256)] : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 16, 576), mul(%value_head_size0, 16)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c4 = index.constant 4 : index + %c8 = index.constant 8 : index + %c256 = index.constant 256 : index + %c0_offset = index.constant 0 : offset + %stage_bytes = index.constant 64 : offset + %zero = scalar.constant 0.0 : f32 + %one = scalar.constant 1.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %head0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %workitem = kernel.workitem.id : index + %subgroup = kernel.subgroup.id : index + %lane = kernel.subgroup.lane.id : index + %head, %launch_heads = index.assume %head0, %query_head_count [lt(%head0, %query_head_count)] : index, index + %token, %launch_tokens = index.assume %token0, %tokens [lt(%token0, %tokens)] : index, index + %heads_per_key_value_head0 = index.div %query_head_count, %key_value_head_count : index + %heads_per_key_value_head = index.assume %heads_per_key_value_head0 [range(%heads_per_key_value_head0, 1, 256)] : index + %key_value_head0 = index.div %head, %heads_per_key_value_head : index + %key_value_head, %launch_key_value_heads = index.assume %key_value_head0, %key_value_head_count [lt(%key_value_head0, %key_value_head_count)] : index, index + %bounded_key_count, %launch_key_capacity = index.assume %key_count, %key_capacity [le(%key_count, %key_capacity)] : index, index + %query_view = buffer.view %query[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> + %key_view = buffer.view %key[%c0_offset] : buffer -> view<[%launch_key_capacity]x[%launch_key_value_heads]x[%qk_head_size]xf16> + %mask_view = buffer.view %mask[%c0_offset] : buffer -> view<[%launch_tokens]x[%bounded_key_count]xf16> + %sink_view = buffer.view %sinks[%c0_offset] : buffer -> view<[%launch_heads]xf32> + %output_view = buffer.view %output[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + %stage = buffer.alloca align(16) %stage_bytes : buffer + %stage_view = buffer.view %stage[%c0_offset] : buffer -> view<16xf32> + // Per workitem: online maximum and sum over keys workitem, workitem + 256, ... + %local_max, %local_sum = scf.for %key_index = [%workitem to %bounded_key_count step %c256](%running_max = %negative_large : f32, %running_sum = %zero : f32) -> (f32, f32) { + %mask_f16 = view.load %mask_view[%token, %key_index] : view<[%launch_tokens]x[%bounded_key_count]xf16> -> f16 + %mask_value = scalar.extf %mask_f16 : f16 to f32 + %active = scalar.cmpf ogt, %mask_value, %negative_large : f32 + %next_max, %next_sum = scf.if %active -> (f32, f32) { + %key_row, %row_bound = index.assume %key_index, %key_capacity [lt(%key_index, %key_capacity)] : index, index + %dot = scf.for %channel = [%c0 to %qk_head_size step %c4](%acc = %zero : f32) -> (f32) { + %q = vector.load %query_view[%token, %head, %channel] : view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> -> vector<4xf32> + %k16 = vector.load %key_view[%key_row, %key_value_head, %channel] : view<[%launch_key_capacity]x[%launch_key_value_heads]x[%qk_head_size]xf16> -> vector<4xf16> + %k = vector.extf %k16 : vector<4xf16> to vector<4xf32> + %qk = vector.mulf %q, %k : vector<4xf32> + %sum4 = vector.reduce %qk, %acc : vector<4xf32>, f32 + scf.yield %sum4 : f32 + } + %scaled = scalar.mulf %dot, %scale : f32 + %score = scalar.addf %scaled, %mask_value : f32 + %grows = scalar.cmpf ogt, %score, %running_max : f32 + %new_max = scf.select %grows, %score, %running_max : f32 + %old_shift = scalar.subf %running_max, %new_max : f32 + %old_scale = scalar.expf %old_shift : f32 + %score_shift = scalar.subf %score, %new_max : f32 + %score_exp = scalar.expf %score_shift : f32 + %kept = scalar.mulf %running_sum, %old_scale : f32 + %new_sum = scalar.addf %kept, %score_exp : f32 + scf.yield %new_max, %new_sum : f32, f32 + } else { + scf.yield %running_max, %running_sum : f32, f32 + } + scf.yield %next_max, %next_sum : f32, f32 + } + // Row maximum: subgroup reduce, then across the eight subgroups through workgroup memory. + %subgroup_max = kernel.subgroup.reduce %local_max : f32 + %lane_is_zero = index.cmp eq, %lane, %c0 : index + %subgroup_slot0 = index.assume %subgroup [range(%subgroup, 0, 7)] : index + scf.if %lane_is_zero { + view.store %subgroup_max, %stage_view[%subgroup_slot0] : f32, view<16xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_max = scf.for %slot = [%c0 to %c8 step %c1](%acc_max = %negative_large : f32) -> (f32) { + %slot_max = view.load %stage_view[%slot] : view<16xf32> -> f32 + %m = scalar.maxnumf %acc_max, %slot_max : f32 + scf.yield %m : f32 + } + // Row sum relative to the row maximum. + %local_shift = scalar.subf %local_max, %row_max : f32 + %local_rescale = scalar.expf %local_shift : f32 + %local_rescaled = scalar.mulf %local_sum, %local_rescale : f32 + %subgroup_sum = kernel.subgroup.reduce %local_rescaled : f32 + %subgroup_slot1 = index.add %subgroup_slot0, %c8 : index + scf.if %lane_is_zero { + view.store %subgroup_sum, %stage_view[%subgroup_slot1] : f32, view<16xf32> + } + kernel.barrier scope(workgroup) ordering(acq_rel) + %row_sum_total = scf.for %slot0 = [%c0 to %c8 step %c1](%acc_s = %zero : f32) -> (f32) { + %slot = index.add %slot0, %c8 : index + %slot_sum = view.load %stage_view[%slot] : view<16xf32> -> f32 + %s = scalar.addf %acc_s, %slot_sum : f32 + scf.yield %s : f32 + } + %sink = view.load %sink_view[%head] : view<[%launch_heads]xf32> -> f32 + %sink_shift = scalar.subf %sink, %row_max : f32 + %sink_exp = scalar.expf %sink_shift : f32 + %denominator = scalar.addf %row_sum_total, %sink_exp : f32 + %factor0 = scalar.divf %row_sum_total, %denominator : f32 + %any_key = scalar.cmpf ogt, %row_sum_total, %zero : f32 + %factor = scf.select %any_key, %factor0, %zero : f32 + scf.for %channel = [%workitem to %value_head_size step %c256] { + %value = view.load %output_view[%token, %head, %channel] : view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> -> f32 + %scaled_value = scalar.mulf %value, %factor : f32 + %result = scf.select %any_key, %scaled_value, %zero : f32 + view.store %result, %output_view[%token, %head, %channel] : f32, view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + } + kernel.return +} + +// Reference for the cases: full attention for one (query head, token) row per workitem, with (use_sink = 1) or +// without (use_sink = 0) the sink in the denominator. value is [key_capacity][key_value_heads][value_head_size] f16. +kernel.def target(@ggml_attention_sink_gfx11_wave32) export("ggml_attention_sink_reference_f32") @ggml_attention_sink_reference_f32(%tokens: index, %key_count: index, %key_capacity: index, %use_sink: index) { + %query_head_count = config.get @ggml.attention_sink.query_head_count : index + %c1 = index.constant 1 : index + kernel.launch.config workgroups(%query_head_count, %tokens, %c1) workgroup_size(%c1, %c1, %c1) : index +} launch(%tokens: index, %key_count: index, %key_capacity: index, %use_sink: index, %query: buffer, %key: buffer, %value: buffer, %mask: buffer, %sinks: buffer, %output: buffer) where [range(%tokens, 1, 65536), range(%key_count, 1, 1048576), range(%key_capacity, 1, 1048576), range(%use_sink, 0, 1)] { + %query_head_count0 = config.get @ggml.attention_sink.query_head_count : index + %key_value_head_count0 = config.get @ggml.attention_sink.key_value_head_count : index + %qk_head_size0 = config.get @ggml.attention_sink.qk_head_size : index + %value_head_size0 = config.get @ggml.attention_sink.value_head_size : index + %scale = config.get @ggml.attention_sink.scale : f32 + %query_head_count = index.assume %query_head_count0 [range(%query_head_count0, 1, 256)] : index + %key_value_head_count = index.assume %key_value_head_count0 [range(%key_value_head_count0, 1, 256)] : index + %qk_head_size = index.assume %qk_head_size0 [range(%qk_head_size0, 16, 576), mul(%qk_head_size0, 16)] : index + %value_head_size = index.assume %value_head_size0 [range(%value_head_size0, 16, 576), mul(%value_head_size0, 16)] : index + %c0 = index.constant 0 : index + %c1 = index.constant 1 : index + %c0_offset = index.constant 0 : offset + %zero = scalar.constant 0.0 : f32 + %negative_large = scalar.constant -1e+30 : f32 + %head0 = kernel.workgroup.id : index + %token0 = kernel.workgroup.id : index + %head, %launch_heads = index.assume %head0, %query_head_count [lt(%head0, %query_head_count)] : index, index + %token, %launch_tokens = index.assume %token0, %tokens [lt(%token0, %tokens)] : index, index + %ratio0 = index.div %query_head_count, %key_value_head_count : index + %ratio = index.assume %ratio0 [range(%ratio0, 1, 256)] : index + %kv_head0 = index.div %head, %ratio : index + %kv_head, %launch_kv_heads = index.assume %kv_head0, %key_value_head_count [lt(%kv_head0, %key_value_head_count)] : index, index + %count, %capacity = index.assume %key_count, %key_capacity [le(%key_count, %key_capacity)] : index, index + %qv = buffer.view %query[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> + %kv = buffer.view %key[%c0_offset] : buffer -> view<[%capacity]x[%launch_kv_heads]x[%qk_head_size]xf16> + %vv = buffer.view %value[%c0_offset] : buffer -> view<[%capacity]x[%launch_kv_heads]x[%value_head_size]xf16> + %mv = buffer.view %mask[%c0_offset] : buffer -> view<[%launch_tokens]x[%count]xf16> + %sv = buffer.view %sinks[%c0_offset] : buffer -> view<[%launch_heads]xf32> + %ov = buffer.view %output[%c0_offset] : buffer -> view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + %row_max = scf.for %j = [%c0 to %count step %c1](%m = %negative_large : f32) -> (f32) { + %jk, %jcap = index.assume %j, %key_capacity [lt(%j, %key_capacity)] : index, index + %mk16 = view.load %mv[%token, %j] : view<[%launch_tokens]x[%count]xf16> -> f16 + %mk = scalar.extf %mk16 : f16 to f32 + %dot = scf.for %c = [%c0 to %qk_head_size step %c1](%a = %zero : f32) -> (f32) { + %q = view.load %qv[%token, %head, %c] : view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> -> f32 + %k16 = view.load %kv[%jk, %kv_head, %c] : view<[%capacity]x[%launch_kv_heads]x[%qk_head_size]xf16> -> f16 + %k = scalar.extf %k16 : f16 to f32 + %p = scalar.mulf %q, %k : f32 + %a1 = scalar.addf %a, %p : f32 + scf.yield %a1 : f32 + } + %sd = scalar.mulf %dot, %scale : f32 + %s = scalar.addf %sd, %mk : f32 + %live = scalar.cmpf ogt, %mk, %negative_large : f32 + %bigger = scalar.cmpf ogt, %s, %m : f32 + %take = scalar.andi %live, %bigger : i1 + %m1 = scf.select %take, %s, %m : f32 + scf.yield %m1 : f32 + } + %sink_raw = view.load %sv[%head] : view<[%launch_heads]xf32> -> f32 + %has_sink = index.cmp eq, %use_sink, %c1 : index + scf.for %c = [%c0 to %value_head_size step %c1] { + %num, %den = scf.for %j = [%c0 to %count step %c1](%n = %zero : f32, %d = %zero : f32) -> (f32, f32) { + %jk, %jcap = index.assume %j, %key_capacity [lt(%j, %key_capacity)] : index, index + %mk16 = view.load %mv[%token, %j] : view<[%launch_tokens]x[%count]xf16> -> f16 + %mk = scalar.extf %mk16 : f16 to f32 + %live = scalar.cmpf ogt, %mk, %negative_large : f32 + %dot = scf.for %i = [%c0 to %qk_head_size step %c1](%a = %zero : f32) -> (f32) { + %q = view.load %qv[%token, %head, %i] : view<[%launch_tokens]x[%launch_heads]x[%qk_head_size]xf32> -> f32 + %k16 = view.load %kv[%jk, %kv_head, %i] : view<[%capacity]x[%launch_kv_heads]x[%qk_head_size]xf16> -> f16 + %k = scalar.extf %k16 : f16 to f32 + %p = scalar.mulf %q, %k : f32 + %a1 = scalar.addf %a, %p : f32 + scf.yield %a1 : f32 + } + %sd = scalar.mulf %dot, %scale : f32 + %s = scalar.addf %sd, %mk : f32 + %sh = scalar.subf %s, %row_max : f32 + %e0 = scalar.expf %sh : f32 + %e = scf.select %live, %e0, %zero : f32 + %v16 = view.load %vv[%jk, %kv_head, %c] : view<[%capacity]x[%launch_kv_heads]x[%value_head_size]xf16> -> f16 + %v = scalar.extf %v16 : f16 to f32 + %ev = scalar.mulf %e, %v : f32 + %n1 = scalar.addf %n, %ev : f32 + %d1 = scalar.addf %d, %e : f32 + scf.yield %n1, %d1 : f32, f32 + } + %sink_shift = scalar.subf %sink_raw, %row_max : f32 + %sink_e = scalar.expf %sink_shift : f32 + %sink_term = scf.select %has_sink, %sink_e, %zero : f32 + %den_total = scalar.addf %den, %sink_term : f32 + %ratio_out = scalar.divf %num, %den_total : f32 + %any = scalar.cmpf ogt, %den, %zero : f32 + %out = scf.select %any, %ratio_out, %zero : f32 + view.store %out, %ov[%token, %head, %c] : f32, view<[%launch_tokens]x[%launch_heads]x[%value_head_size]xf32> + } + kernel.return +} + +// Cases: 8 query heads over 2 key-value heads (gpt-oss's 64:8 grouping in small), head sizes 64, scale 0.125. +// Run with --config=ggml.attention_sink.query_head_count=8 --config=ggml.attention_sink.key_value_head_count=2 +// --config=ggml.attention_sink.qk_head_size=64 --config=ggml.attention_sink.value_head_size=64 --config=ggml.attention_sink.scale=0.125 + +check.case public @ggml_attention_sink_prefill_case { + %q_seed = check.param.seed base(7300000000000068001) count(1) : i64 + %k_seed = check.param.seed base(7300000000000068002) count(1) : i64 + %v_seed = check.param.seed base(7300000000000068003) count(1) : i64 + %m_seed = check.param.seed base(7300000000000068004) count(1) : i64 + %s_seed = check.param.seed base(7300000000000068005) count(1) : i64 + %tokens = check.literal value(3) : index + %keys = check.literal value(300) : index + %no_sink = check.literal value(0) : index + %sink = check.literal value(1) : index + %query = check.generate.random.uniform seed(%q_seed) range(-1.0 to 1.0) : tensor<1536xf32> + %key = check.generate.random.uniform seed(%k_seed) range(-1.0 to 1.0) : tensor<38400xf16> + %value = check.generate.random.uniform seed(%v_seed) range(-1.0 to 1.0) : tensor<38400xf16> + %mask = check.generate.random.uniform seed(%m_seed) range(-3.0 to 0.0) : tensor<900xf16> + %sinks = check.generate.random.uniform seed(%s_seed) range(-2.0 to 2.0) : tensor<8xf32> + %actual = check.generate.fill value(-7.0) : tensor<1536xf32> + %expected = check.generate.fill value(7.0) : tensor<1536xf32> + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %no_sink](%tokens, %keys, %keys, %no_sink, %query, %key, %value, %mask, %sinks, %actual) : [index, index, index, index](index, index, index, index, tensor<1536xf32>, tensor<38400xf16>, tensor<38400xf16>, tensor<900xf16>, tensor<8xf32>, tensor<1536xf32>) + kernel.launch @ggml_attention_sink_f32[%tokens, %keys, %keys](%tokens, %keys, %keys, %query, %key, %mask, %sinks, %actual) : [index, index, index](index, index, index, tensor<1536xf32>, tensor<38400xf16>, tensor<900xf16>, tensor<8xf32>, tensor<1536xf32>) + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %sink](%tokens, %keys, %keys, %sink, %query, %key, %value, %mask, %sinks, %expected) : [index, index, index, index](index, index, index, index, tensor<1536xf32>, tensor<38400xf16>, tensor<38400xf16>, tensor<900xf16>, tensor<8xf32>, tensor<1536xf32>) + check.expect.close actual(%actual) expected(%expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<1536xf32> + check.return +} + +check.case public @ggml_attention_sink_decode_case { + %q_seed = check.param.seed base(7300000000000068011) count(1) : i64 + %k_seed = check.param.seed base(7300000000000068012) count(1) : i64 + %v_seed = check.param.seed base(7300000000000068013) count(1) : i64 + %m_seed = check.param.seed base(7300000000000068014) count(1) : i64 + %s_seed = check.param.seed base(7300000000000068015) count(1) : i64 + %tokens = check.literal value(1) : index + %keys = check.literal value(1000) : index + %no_sink = check.literal value(0) : index + %sink = check.literal value(1) : index + %query = check.generate.random.uniform seed(%q_seed) range(-1.0 to 1.0) : tensor<512xf32> + %key = check.generate.random.uniform seed(%k_seed) range(-1.0 to 1.0) : tensor<128000xf16> + %value = check.generate.random.uniform seed(%v_seed) range(-1.0 to 1.0) : tensor<128000xf16> + %mask = check.generate.random.uniform seed(%m_seed) range(-3.0 to 0.0) : tensor<1000xf16> + %sinks = check.generate.random.uniform seed(%s_seed) range(-2.0 to 2.0) : tensor<8xf32> + %actual = check.generate.fill value(-7.0) : tensor<512xf32> + %expected = check.generate.fill value(7.0) : tensor<512xf32> + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %no_sink](%tokens, %keys, %keys, %no_sink, %query, %key, %value, %mask, %sinks, %actual) : [index, index, index, index](index, index, index, index, tensor<512xf32>, tensor<128000xf16>, tensor<128000xf16>, tensor<1000xf16>, tensor<8xf32>, tensor<512xf32>) + kernel.launch @ggml_attention_sink_f32[%tokens, %keys, %keys](%tokens, %keys, %keys, %query, %key, %mask, %sinks, %actual) : [index, index, index](index, index, index, tensor<512xf32>, tensor<128000xf16>, tensor<1000xf16>, tensor<8xf32>, tensor<512xf32>) + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %sink](%tokens, %keys, %keys, %sink, %query, %key, %value, %mask, %sinks, %expected) : [index, index, index, index](index, index, index, index, tensor<512xf32>, tensor<128000xf16>, tensor<128000xf16>, tensor<1000xf16>, tensor<8xf32>, tensor<512xf32>) + check.expect.close actual(%actual) expected(%expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<512xf32> + check.return +} + +check.case public @ggml_attention_sink_masked_case { + %q_seed = check.param.seed base(7300000000000068021) count(1) : i64 + %k_seed = check.param.seed base(7300000000000068022) count(1) : i64 + %v_seed = check.param.seed base(7300000000000068023) count(1) : i64 + %m_seed = check.param.seed base(7300000000000068024) count(1) : i64 + %s_seed = check.param.seed base(7300000000000068025) count(1) : i64 + %tokens = check.literal value(2) : index + %keys = check.literal value(40) : index + %no_sink = check.literal value(0) : index + %sink = check.literal value(1) : index + %query = check.generate.random.uniform seed(%q_seed) range(-1.0 to 1.0) : tensor<1024xf32> + %key = check.generate.random.uniform seed(%k_seed) range(-1.0 to 1.0) : tensor<5120xf16> + %value = check.generate.random.uniform seed(%v_seed) range(-1.0 to 1.0) : tensor<5120xf16> + %mask = check.generate.fill value(-1.0e30) : tensor<80xf16> + %sinks = check.generate.random.uniform seed(%s_seed) range(-2.0 to 2.0) : tensor<8xf32> + %actual = check.generate.fill value(-7.0) : tensor<1024xf32> + %expected = check.generate.fill value(7.0) : tensor<1024xf32> + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %no_sink](%tokens, %keys, %keys, %no_sink, %query, %key, %value, %mask, %sinks, %actual) : [index, index, index, index](index, index, index, index, tensor<1024xf32>, tensor<5120xf16>, tensor<5120xf16>, tensor<80xf16>, tensor<8xf32>, tensor<1024xf32>) + kernel.launch @ggml_attention_sink_f32[%tokens, %keys, %keys](%tokens, %keys, %keys, %query, %key, %mask, %sinks, %actual) : [index, index, index](index, index, index, tensor<1024xf32>, tensor<5120xf16>, tensor<80xf16>, tensor<8xf32>, tensor<1024xf32>) + kernel.launch @ggml_attention_sink_reference_f32[%tokens, %keys, %keys, %sink](%tokens, %keys, %keys, %sink, %query, %key, %value, %mask, %sinks, %expected) : [index, index, index, index](index, index, index, index, tensor<1024xf32>, tensor<5120xf16>, tensor<5120xf16>, tensor<80xf16>, tensor<8xf32>, tensor<1024xf32>) + check.expect.close actual(%actual) expected(%expected) atol(1.0e-5) rtol(1.0e-4) nan(same) : tensor<1024xf32> + check.return +} diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 77f7e6afc66e..219bbbaacaa2 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -385,4 +385,9 @@ if (TARGET ggml-hrx) target_link_libraries(test-hrx-fa-masked-v PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) target_include_directories(test-hrx-fa-masked-v PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) add_test(NAME test-hrx-fa-masked-v COMMAND test-hrx-fa-masked-v) + + add_executable(test-hrx-attention-sink test-hrx-attention-sink.cpp) + target_link_libraries(test-hrx-attention-sink PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-attention-sink PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-attention-sink COMMAND test-hrx-attention-sink) endif() diff --git a/tests/test-hrx-attention-sink.cpp b/tests/test-hrx-attention-sink.cpp new file mode 100644 index 000000000000..702dea88672c --- /dev/null +++ b/tests/test-hrx-attention-sink.cpp @@ -0,0 +1,75 @@ +// 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. + +// The attention-sink dispatch rewrites FlashAttention's output in place, so it must refuse a node whose output +// shares storage with one of its inputs, including a different value (view) of the same allocation. + +#include "dispatch_registration/common/dispatch-attention-sink.h" +#include "ggml.h" +#include "graph/graph.h" +#include "graph/op-params.h" + +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static bool append_for(bool alias_output_onto_query) { + ggml_init_params params = { 16 * ggml_tensor_overhead(), nullptr, true }; + ggml_context * ctx = ggml_init(params); + const int64_t d = 64, tokens = 8, heads = 8, kv_heads = 2, keys = 32; // tokens == heads: output and query share a layout + ggml_tensor * q = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, d, tokens, heads); + ggml_tensor * k = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, d, keys, kv_heads); + ggml_tensor * v = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, d, keys, kv_heads); + ggml_tensor * mask = ggml_new_tensor_2d(ctx, GGML_TYPE_F16, keys, tokens); + ggml_tensor * sinks = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, heads); + ggml_tensor * fa = ggml_flash_attn_ext(ctx, q, k, v, mask, 0.125f, 0.0f, 0.0f); + ggml_flash_attn_ext_add_sinks(fa, sinks); + + ggml::hrx::Graph graph; + std::vector inputs; + for (ggml_tensor * source : fa->src) { + if (source != nullptr) { + inputs.push_back(graph.values().get_or_add_tensor_value(source, ggml::hrx::ValueKind::External)); + } + } + const ggml::hrx::ValueId output = graph.values().get_or_add_tensor_value(fa, ggml::hrx::ValueKind::Transient); + if (alias_output_onto_query) { + REQUIRE(graph.values().alias_storage(output, inputs[0]).success()); + REQUIRE(graph.values().same_storage(output, inputs[0])); + REQUIRE(output.value != inputs[0].value); + } + ggml::hrx::GraphNode & node = graph.add_node(fa->op, output, inputs); + node.params = ggml::hrx::import_op_params(*fa); + REQUIRE(ggml::hrx::attention_sinks_supported(graph, node)); + ggml::hrx::DispatchMatch match; + const bool appended = ggml::hrx::append_attention_sink_dispatch(graph, node, match); + ggml_free(ctx); + return appended; +} + +int main() { + REQUIRE(append_for(false)); + REQUIRE(!append_for(true)); + std::printf("test-hrx-attention-sink: disjoint output accepted, output aliasing the query refused\n"); + return 0; +}