Xenova's picture
Xenova HF Staff
sync c928d21e6cc1
2fad9e5 verified
|
Raw
History Blame
9.74 kB
---
library_name: kernels
license: apache-2.0
tags:
- kernel
- webgpu
- wgsl
---
# com.microsoft.GroupQueryAttention
`com.microsoft` · ONNX Runtime contrib operator · contrib since_version 1
## Description
Grouped-query attention for explicit BSH Q/K/V and BNSH caches. Direct Q/K/V supports bidirectional attention or causal local windows and may store its generated float cache independently as float16 or float32; existing unquantized cache inputs match the Q/K/V dtype. Causal cache paths support rotary embeddings, sliding windows, bias, head sinks, softcap, smooth softmax, and paired Q/K RMS normalization. Int8/int4 caches require float32 Q/K/V and output; int4 is prompt-only. Packed QKV, position IDs, interleaved rotary, bfloat16/float8, and diagnostic QK output are not implemented.
See the [ONNX Runtime `GroupQueryAttention` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.GroupQueryAttention) for the reference semantics.
## Inputs
| Name | Bind key | Logical dtype | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- | --- |
| `query` | `queryT` | `T` | `3` | — | Query tensor of shape `(batch_size, sequence_length, num_heads * head_size)`. | required |
| `key` | `keyT` | `T` | `3` | — | Key tensor of shape `(batch_size, kv_sequence_length, kv_num_heads * head_size)`. | required |
| `value` | `valueT` | `T` | `3` | — | Value tensor of shape `(batch_size, kv_sequence_length, kv_num_heads * head_size)`. | required |
| `past_key` | `pastKeyT` | `T_CACHE` | `4` | — | Optional cached key state in BNSH format. Its sequence axis is `max_sequence_length` when the past and present buffers are shared, otherwise `past_sequence_length`; int4 stores each signed value as a +8-biased nibble, with the even head coordinate low, packing two values per logical uint8 element and widening each byte to one u32 WebGPU buffer word. | optional |
| `past_value` | `pastValueT` | `T_CACHE` | `4` | — | Optional cached value state in BNSH format with the same length and packing semantics as `past_key`. | optional |
| `seqlens_k` | `seqlensKT` | `M` | `1` | — | Int32 tensor of shape `(batch_size)` containing each sample's total sequence length minus one. | required |
| `total_sequence_length` | `totalSequenceLengthT` | `M` | `1` | — | Length-one int32 tensor containing the maximum total sequence length (past plus new) in the batch. | required |
| `cos_cache` | `cosCacheT` | `T` | `2` | — | Optional cosine cache for rotary embeddings with shape `(max_sequence_length, head_size / 2)`. | optional |
| `sin_cache` | `sinCacheT` | `T` | `2` | — | Optional sine cache for rotary embeddings with shape `(max_sequence_length, head_size / 2)`. | optional |
| `attention_bias` | `attentionBiasT` | `T` | `4` | — | Optional additive term for QK scores with shape `(batch_size or 1, num_heads or 1, sequence_length, total_sequence_length)`; the first two dimensions broadcast. | optional |
| `head_sink` | `headSinkT` | `T` | `1` | — | Optional per-head smooth factor of shape `(num_heads)` added to the softmax denominator. | optional |
| `k_scale` | `kScaleT` | `T_KV_SCALE` | `1` | — | Optional float32 key-cache scale: one value for `PER_TENSOR`, or `kv_num_heads * head_size` values for `PER_CHANNEL`. | optional |
| `v_scale` | `vScaleT` | `T_KV_SCALE` | `1` | — | Optional float32 value-cache scale with the same shape convention as `k_scale`. | optional |
| `q_norm_weight` | `qNormWeightT` | `T` | `1` | — | Optional per-head RMS-normalization weight of shape `(head_size)` applied to queries before rotary embedding. It must be provided together with `k_norm_weight`. | optional |
| `k_norm_weight` | `kNormWeightT` | `T` | `1` | — | Optional per-head RMS-normalization weight of shape `(head_size)` applied to keys before rotary embedding. It must be provided together with `q_norm_weight`. | optional |
## Outputs
| Name | Bind key | Logical dtype | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- | --- |
| `output` | `outputT` | `T` | `3` | same as `query` | Attention output of shape `(batch_size, sequence_length, hidden_size)`. | required |
| `present_key` | `presentKeyT` | `T_CACHE` | `4` | — | Updated key cache in BNSH format. Its sequence axis is `max_sequence_length` for a shared buffer, otherwise `past_sequence_length + kv_sequence_length`; int4 stores each signed value as a +8-biased nibble, with the even head coordinate low, packing two values per logical uint8 element and widening each byte to one u32 WebGPU buffer word. | required |
| `present_value` | `presentValueT` | `T_CACHE` | `4` | — | Updated value cache in BNSH format with the same length and packing semantics as `present_key`. | required |
## Attributes
Attributes and default values (overridable per request):
| Attribute | Default | Description |
| --- | --- | --- |
| `causal` | `1` | Whether to apply a causal mask. Set to 0 for bidirectional attention; `local_window_size` must then be -1. |
| `do_rotary` | `0` | Set to 1 to apply rotary position embeddings. The default 0 disables them. |
| `k_quant_type` | `"NONE"` | Key-cache quantization mode: `NONE`, `PER_TENSOR`, or `PER_CHANNEL`. |
| `local_window_size` | `-1` | Left window size for causal local attention. The default -1 disables local attention, and the value must be -1 when `causal` is 0. |
| `qk_norm_epsilon` | `0.000001` | Epsilon for the per-head Q/K RMS normalization applied when both normalization weights are provided. |
| `sliding_window_cache` | `0` | Set to 1 when past/present caches are fixed-size window buffers that evict old tokens from the front. Requires `local_window_size > 0` and enough cache capacity. |
| `smooth_softmax` | `-1` | Set to 1 to enable the smooth-softmax denominator term. |
| `softcap` | `0` | Positive softcap applied to attention scores. The default 0 disables soft-capping. |
| `v_quant_type` | `"NONE"` | Value-cache quantization mode: `NONE`, `PER_TENSOR`, or `PER_CHANNEL`. |
| `kv_cache_bit_width` | — | Quantized cache bit width, either 8 or 4. Four-bit values are packed two per uint8 element. |
| `kv_num_heads` | — | Number of key/value attention heads. |
| `num_heads` | — | Number of query attention heads. |
| `scale` | — | Optional QK score scale; zero or omission selects `1 / sqrt(head_size)`. |
## Type constraints
| Variable | Allowed dtypes |
| --- | --- |
| `T` | `float32`, `float16` |
| `T_CACHE` | `float32`, `float16`, `uint8`, `int8` |
| `T_KV_SCALE` | `float32` |
| `M` | `int32` |
## Device requirements
Some implementation variants require `subgroup-matrix`, `shader-f16`, and `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.
## Files
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, provenance)
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
- [`test.json`](build/webgpu/test.json) — correctness cases
- [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
- [`attention-rank4-tiled.wgsl.jinja`](build/webgpu/attention-rank4-tiled.wgsl.jinja)
- [`attn-flash-decode-splitk-merge.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk-merge.wgsl.jinja)
- [`attn-flash-decode-splitk.wgsl.jinja`](build/webgpu/attn-flash-decode-splitk.wgsl.jinja)
- [`attn-flash-online.wgsl.jinja`](build/webgpu/attn-flash-online.wgsl.jinja)
- [`attn-flash-prefill-cluster.wgsl.jinja`](build/webgpu/attn-flash-prefill-cluster.wgsl.jinja)
- [`attn-flash-q32-broadcast.wgsl.jinja`](build/webgpu/attn-flash-q32-broadcast.wgsl.jinja)
- [`attn-materialized-rowstats-combine-f32.wgsl.jinja`](build/webgpu/attn-materialized-rowstats-combine-f32.wgsl.jinja)
- [`attn-materialized-sgmat-f32.wgsl.jinja`](build/webgpu/attn-materialized-sgmat-f32.wgsl.jinja)
- [`attn-online-scalar.wgsl.jinja`](build/webgpu/attn-online-scalar.wgsl.jinja)
- [`gqa-attention.wgsl.jinja`](build/webgpu/gqa-attention.wgsl.jinja)
- [`gqa-present.wgsl.jinja`](build/webgpu/gqa-present.wgsl.jinja)
- [`gqa-qprep.wgsl.jinja`](build/webgpu/gqa-qprep.wgsl.jinja)
## Use with `@huggingface/kernels`
The loader automatically allocates outputs whose metadata it can derive from the manifest contract and this call.
The explicit `outputs` entries provide shape and logical dtype metadata for the results listed below:
- `presentKeyT`
- `presentValueT`
Each entry either requests an optional result or supplies metadata that cannot be inferred from the inputs.
The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
Replace each `*Data` placeholder with a typed array containing the corresponding input data.
```js
import { getKernel } from "@huggingface/kernels";
const kernel = await getKernel("webgpu-kernels/com.microsoft.GroupQueryAttention", { version: 1 });
// Explicit destinations request optional results or supply metadata that cannot be inferred.
const { outputT, presentKeyT, presentValueT } = await kernel({
queryT: { data: queryTData, shape: [2, 1, 16] },
keyT: { data: keyTData, shape: [2, 1, 8] },
valueT: { data: valueTData, shape: [2, 1, 8] },
pastKeyT: { data: pastKeyTData, shape: [2, 1, 8, 8] },
pastValueT: { data: pastValueTData, shape: [2, 1, 8, 8] },
seqlensKT: { data: seqlensKTData, shape: [2] },
totalSequenceLengthT: { data: totalSequenceLengthTData, shape: [1] },
}, {
attrs: { num_heads: 2, kv_num_heads: 1 },
outputs: {
presentKeyT: { shape: [2, 1, 8, 8], dtype: "float32" },
presentValueT: { shape: [2, 1, 8, 8], dtype: "float32" },
},
});
```