From dbaeb30d6cc0a9be3fdc4318b30113f5dc6ab7b5 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Tue, 15 Sep 2026 22:11:28 +0000 Subject: [PATCH 01/10] Initial plan From cf546932bfda635b0164eb88c38fc1b5da5b653a Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 16 Sep 2026 02:18:34 +0000 Subject: [PATCH 02/10] Implement packed sparse attention indexer Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- docs/ContribOperators.md | 158 ++- docs/OperatorKernels.md | 1 + .../packed_sparse_attention_indexer.md | 267 +++++ .../webgpu/packed_sparse_attention_indexer.md | 48 + .../packed_sparse_attention_indexer_common.h | 85 ++ .../sparse/sparse_attention_indexer_common.h | 18 +- .../contrib_ops/cuda/cuda_contrib_kernels.cc | 6 + .../sparse/packed_sparse_attention_indexer.cc | 437 +++++++ .../sparse/packed_sparse_attention_indexer.h | 37 + .../packed_sparse_attention_indexer_impl.cu | 868 ++++++++++++++ .../packed_sparse_attention_indexer_impl.h | 106 ++ .../sparse_attention_indexer_device_math.cuh | 149 +++ .../sparse/sparse_attention_indexer_impl.cu | 96 +- .../bert/packed_sparse_attention_indexer.cc | 982 ++++++++++++++++ .../bert/packed_sparse_attention_indexer.h | 151 +++ .../webgpu/webgpu_contrib_kernels.cc | 2 + .../core/graph/contrib_ops/bert_defs.cc | 383 ++++++ onnxruntime/core/graph/contrib_ops/ms_opset.h | 2 + .../python/tools/symbolic_shape_infer.py | 51 + ...packed_sparse_attention_indexer_op_test.cc | 1044 +++++++++++++++++ ...untime_test_python_symbolic_shape_infer.py | 129 ++ 21 files changed, 4931 insertions(+), 89 deletions(-) create mode 100644 docs/contrib_ops/packed_sparse_attention_indexer.md create mode 100644 docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md create mode 100644 onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h create mode 100644 onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh create mode 100644 onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc create mode 100644 onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h create mode 100644 onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index 3cc7009d0e83e..875a3483f87d4 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -79,6 +79,7 @@ Do not modify directly.* * com.microsoft.NhwcMaxPool * com.microsoft.PackedAttention * com.microsoft.PackedMultiHeadAttention + * com.microsoft.PackedSparseAttentionIndexer * com.microsoft.Pad * com.microsoft.PagedAttention * com.microsoft.QAttention @@ -4672,6 +4673,161 @@ This version of the operator has been available since version 1 of the 'com.micr +### **com.microsoft.PackedSparseAttentionIndexer** + + Packed/variable-length counterpart of SparseAttentionIndexer, for continuous-batching engines + (such as an OgaEngine-style PagedAttention model) that flatten every request's tokens into one + [total_tokens, ...] axis instead of a dense [batch_size, sequence_length, ...] axis. It selects, for + every packed query token, the sparse-attention candidates that a following SparsePagedAttention (or + similar) operator is allowed to read. + + Unlike SparseAttentionIndexer, this operator: + * takes packed query/key tensors plus cumulative_sequence_lengths (request boundaries) and + past_sequence_lengths (per-request past length) instead of a dense batch and a dense mask; + * derives ordinary causal visibility purely from that packed metadata -- there is no mask input; + * uses a single generic set of state slots (past_key_state / past_kv_buffer / past_gate_buffer / + past_state_lengths) for both policy_mode values, each with a shape that is fixed across calls + (state never grows and is never concatenated); state overflow beyond the fixed capacity is + rejected as a deterministic no-op on state rather than truncated or allowed to corrupt memory; + * additionally emits selected_counts, the exact number of active (non -1) entries per query, so + that no downstream consumer needs to scan selected_indices for its query's true count. + + Both policy_mode values keep the semantics of SparseAttentionIndexer, applied independently to each + request's own packed token range and fixed-capacity state slice: + + policy_mode = "qsa" ("query sparse attention" token indexer) + Processes each request's new tokens sequentially: appends raw indexer keys to the generic + pending buffer, and whenever it reaches compress_ratio tokens, mean-pools it, applies RMSNorm + and key_norm_weight, applies the leading/split-half rotary convention at the block's first + logical token position, and appends the prepared (already normalized and rotated) key to + key_state. Queries are scored against every causally visible complete block with + sum_h ReLU(q_h . k), the token_budget / compress_ratio highest scoring blocks are kept, and + their token indices are emitted (request-local logical positions, i.e. the same numbering as + past_sequence_lengths + local offset) followed by the causally visible tokens of the trailing + incomplete block. + + policy_mode = "csa" ("compressed sparse attention" block indexer) + Applies the same window-plan arithmetic as SparseAttentionIndexer (overlap/leftover/new window + count) independently per request, using that request's own buffer_length and new token count; + every newly closed window is compressed with the softmax-gated Ca/Cb pooling, normalized, + rotated and appended to key_state. Queries are scored against every causally visible compressed + entry with sum_h w_h * ReLU(q_h . k) and the index_topk highest scoring entry indices are + emitted. + + Common contract: + * selected_indices is int32 with a fixed capacity that only depends on attributes: + token_budget + compress_ratio - 1 for "qsa" (values are request-local token positions into the + main key/value cache, directly consumable by SparsePagedAttention configured with + attention_mode="selected_only", selected_kv_source="main") and index_topk for "csa" (values are + compressed-entry indices into key_state, directly consumable by SparsePagedAttention configured + with attention_mode="local_plus_selected", selected_kv_source="auxiliary"; key_state is + layout-compatible with a [batch_size, capacity, 1, head_size] auxiliary cache when K = V). + Unused entries are -1 and selected_counts holds the exact number of used entries. + * key_norm_weight is the effective RMSNorm multiplier, exactly as in SparseAttentionIndexer. + * Accumulation, pooling, softmax, normalization and scoring are performed in float32 and the + result is rounded once to the tensor element type. + * Ties in the top-k selection are broken by the smaller entry index, and the emitted entries are + ordered by decreasing score, so the result is deterministic. + * cos_cache / sin_cache may be shared across the batch ([max_position, rotary_width]) or + request-specific ([batch_size, max_position, rotary_width]). + * cumulative_sequence_lengths, past_sequence_lengths and past_state_lengths are read directly by + the device kernel; a zero-token request row (a repeated cumulative offset) is valid and simply + contributes no query rows for that request. + + OgaEngine integration note: this operator only defines the ORT operator; wiring + past_key_state / past_kv_buffer / past_gate_buffer / past_state_lengths as Engine-managed, + per-request fixed-size state (analogous to a paged auxiliary cache) is expected to happen in the + OgaEngine / Model Builder integration, which is out of scope for this operator definition. + +#### Version + +This version of the operator has been available since version 1 of the 'com.microsoft' operator set. + +#### Attributes + +
+
compress_ratio : int (required)
+
Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0.
+
epsilon : float
+
Epsilon of the RMS normalization applied to the compressed keys. Default is 1e-6.
+
head_weight_scale : float
+
Only for policy_mode 'csa': scale applied to head_weights. Default is 1/sqrt(num_heads). Must be omitted when policy_mode is 'qsa'.
+
index_topk : int
+
Only for policy_mode 'csa': number of compressed entries selected per query. Must be > 0. Must be omitted when policy_mode is 'qsa'.
+
policy_mode : string (required)
+
Indexer policy. Must be exactly 'qsa' (token indexer) or 'csa' (compressed block indexer).
+
scale : float
+
Scale applied to the per-head ReLU scores. Default is 1/sqrt(head_size).
+
state_capacity : int (required)
+
Fixed capacity (number of entries) of past_key_state / present_key_state. Must be > 0.
+
token_budget : int
+
Only for policy_mode 'qsa': maximum number of tokens selected from complete blocks. Must be > 0 and divisible by compress_ratio. Must be omitted when policy_mode is 'csa'.
+
+ +#### Inputs (10 - 15) + +
+
query : T
+
Packed indexer queries with shape (total_tokens, num_heads, head_size), already normalized but not yet rotated.
+
key : T
+
Packed indexer key projection of the new tokens. Shape is (total_tokens, head_size) for policy_mode 'qsa' and (total_tokens, 2 * head_size) for policy_mode 'csa', where the first head_size channels are the Ca series and the last head_size channels the Cb series.
+
key_norm_weight : T
+
Effective RMSNorm multiplier of the compressed keys, with shape (head_size).
+
cos_cache : T
+
Cosine rotary table indexed by absolute key position, shared across the batch with shape (max_rotary_sequence_length, rotary_width) or request-specific with shape (batch_size, max_rotary_sequence_length, rotary_width).
+
sin_cache : T
+
Sine rotary table with the same shape as cos_cache.
+
cumulative_sequence_lengths : M
+
Device-resident packed request boundaries with shape (batch_size + 1); cumulative_sequence_lengths[0] must be 0 and cumulative_sequence_lengths[batch_size] must equal total_tokens. Request b owns rows [cumulative_sequence_lengths[b], cumulative_sequence_lengths[b + 1]) of query/key (a repeated offset is a valid zero-token row).
+
past_sequence_lengths : M
+
Device-resident number of tokens already processed for each request before this call, with shape (batch_size). Used as the default absolute query position when position_ids is omitted (policy_mode 'qsa'), and to validate state consistency.
+
gate (optional) : T
+
Only for policy_mode 'csa': gate projection of the new tokens with shape (total_tokens, 2 * head_size).
+
position_bias (optional) : T
+
Only for policy_mode 'csa': per-slot gate bias with shape (compress_ratio, 2 * head_size).
+
head_weights (optional) : T
+
Only for policy_mode 'csa': per-head score weights with shape (total_tokens, num_heads).
+
position_ids (optional) : I
+
Optional for policy_mode 'qsa', required for policy_mode 'csa': absolute position of every packed query, with shape (total_tokens).
+
past_key_state : T
+
Generic fixed-capacity state: policy_mode 'qsa' stores prepared complete-block keys; policy_mode 'csa' stores compressed keys. Shape is (batch_size, state_capacity, head_size) and never changes across calls.
+
past_kv_buffer : T
+
Generic fixed-capacity pending-token buffer. Shape is (batch_size, 2 * compress_ratio - 1, head_size) for policy_mode 'qsa' (which only ever uses up to compress_ratio - 1 of these entries) and (batch_size, 2 * compress_ratio - 1, 2 * head_size) for policy_mode 'csa'.
+
past_gate_buffer (optional) : T
+
Only for policy_mode 'csa': buffered gate projections with the same shape as past_kv_buffer.
+
past_state_lengths : M
+
Generic per-request state length with shape (batch_size, 2). Column 0 is the key_state entry count (policy_mode 'qsa': complete-block count; 'csa': compressed-entry count); column 1 is the pending-buffer length (policy_mode 'qsa': incomplete-block length in [0, compress_ratio); 'csa': buffer length in [0, 2 * compress_ratio)).
+
+ +#### Outputs (6 - 6) + +
+
selected_indices : M
+
Selected entries with shape (total_tokens, capacity). capacity is token_budget + compress_ratio - 1 for policy_mode 'qsa' (request-local token positions into the main key/value cache) and index_topk for policy_mode 'csa' (compressed entry indices into key_state). Unused entries are -1.
+
selected_counts : M
+
Exact number of used (non -1) entries of selected_indices for every query, with shape (total_tokens).
+
present_key_state : T
+
Updated generic key state, with the same fixed shape as past_key_state.
+
present_kv_buffer : T
+
Updated generic pending-token buffer, with the same fixed shape as past_kv_buffer.
+
present_gate_buffer (optional) : T
+
Only for policy_mode 'csa': updated gate buffer with the same fixed shape as past_gate_buffer.
+
present_state_lengths : M
+
Updated generic per-request state length, with the same fixed shape as past_state_lengths.
+
+ +#### Type Constraints + +
+
T : tensor(float), tensor(float16), tensor(bfloat16)
+
Constrain floating point tensors to float, float16 and bfloat16.
+
I : tensor(int64)
+
Constrain position ids to 64-bit integer tensors.
+
M : tensor(int32)
+
Constrain packed metadata, generic state lengths and selected indices/counts to 32-bit integer tensors.
+
+ + ### **com.microsoft.Pad** Given `data` tensor, pads, mode, and value. @@ -7814,5 +7970,3 @@ No versioning maintained for experimental ops.
T : tensor(float)
Constrain input and output types to float32 tensors.
- - diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md index 59a8b35f33420..9a72c7b1dac12 100644 --- a/docs/OperatorKernels.md +++ b/docs/OperatorKernels.md @@ -1122,6 +1122,7 @@ The **OpSet Version** column uses the following notation: |NhwcConv|*in* X:**T**
*in* W:**T**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)| |PackedAttention|*in* input:**T**
*in* weights:**T**
*in* bias:**T**
*in* token_offset:**M**
*in* cumulative_sequence_length:**M**
*in* attention_bias:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| |PackedMultiHeadAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* bias:**T**
*in* token_offset:**M**
*in* cumulative_sequence_length:**M**
*in* attention_bias:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| +|PackedSparseAttentionIndexer|*in* query:**T**
*in* key:**T**
*in* key_norm_weight:**T**
*in* cos_cache:**T**
*in* sin_cache:**T**
*in* cumulative_sequence_lengths:**M**
*in* past_sequence_lengths:**M**
*in* gate:**T**
*in* position_bias:**T**
*in* head_weights:**T**
*in* position_ids:**I**
*in* past_key_state:**T**
*in* past_kv_buffer:**T**
*in* past_gate_buffer:**T**
*in* past_state_lengths:**M**
*out* selected_indices:**M**
*out* selected_counts:**M**
*out* present_key_state:**T**
*out* present_kv_buffer:**T**
*out* present_gate_buffer:**T**
*out* present_state_lengths:**M**|1+|**I** = tensor(int64)
**M** = tensor(int32)
**T** = tensor(bfloat16), tensor(float), tensor(float16)| |PagedAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* key_cache:**T_CACHE**
*in* value_cache:**T_CACHE**
*in* cumulative_sequence_length:**S**
*in* past_seqlens:**S**
*in* block_table:**S**
*in* cos_cache:**T**
*in* sin_cache:**T**
*in* slot_mapping:**S**
*in* head_sink:**T**
*in* q_norm_weight:**T**
*in* k_norm_weight:**T**
*in* k_scale:**T_KV_SCALE**
*in* v_scale:**T_KV_SCALE**
*in* attention_metadata:**S**
*out* output:**T**
*out* key_cache_out:**T_CACHE**
*out* value_cache_out:**T_CACHE**|1+|**S** = tensor(int32)
**T** = tensor(bfloat16), tensor(float16)
**T_CACHE** = tensor(bfloat16), tensor(float16), tensor(float8e4m3fn), tensor(int8), tensor(uint8)
**T_KV_SCALE** = tensor(float)| |QAttention|*in* input:**T1**
*in* weight:**T2**
*in* bias:**T3**
*in* input_scale:**T3**
*in* weight_scale:**T3**
*in* mask_index:**T4**
*in* input_zero_point:**T1**
*in* weight_zero_point:**T2**
*in* past:**T3**
*out* output:**T3**
*out* present:**T3**|1+|**T1** = tensor(int8)
**T2** = tensor(int8)
**T3** = tensor(float), tensor(float16)
**T4** = tensor(int32)| |QMoE|*in* input:**T**
*in* router_probs:**T**
*in* fc1_experts_weights:**T1**
*in* fc1_scales:**T2**
*in* fc1_experts_bias:**T**
*in* fc2_experts_weights:**T1**
*in* fc2_scales:**T2**
*in* fc2_experts_bias:**T**
*in* fc3_experts_weights:**T1**
*in* fc3_scales:**T2**
*in* fc3_experts_bias:**T**
*in* fc1_zero_points:**T1**
*in* fc2_zero_points:**T1**
*in* fc3_zero_points:**T1**
*in* router_weights:**T**
*in* fc1_global_scale:**T4**
*in* fc2_global_scale:**T4**
*in* fc1_act_scale:**T4**
*in* fc2_act_scale:**T4**
*in* fc1_act_block_scale:**T2**
*in* fc2_act_block_scale:**T2**
*out* output:**T**|1+|**T** = tensor(bfloat16), tensor(float16)
**T1** = tensor(float8e4m3fn), tensor(uint8)
**T2** = tensor(bfloat16), tensor(float16), tensor(float8e4m3fn), tensor(float8e8m0)
**T4** = tensor(float)| diff --git a/docs/contrib_ops/packed_sparse_attention_indexer.md b/docs/contrib_ops/packed_sparse_attention_indexer.md new file mode 100644 index 0000000000000..712e73d2da278 --- /dev/null +++ b/docs/contrib_ops/packed_sparse_attention_indexer.md @@ -0,0 +1,267 @@ +# PackedSparseAttentionIndexer — Operator Documentation + +This document describes the `com.microsoft::PackedSparseAttentionIndexer` contrib operator: the +packed/variable-length counterpart of `com.microsoft::SparseAttentionIndexer`, built for +continuous-batching (paged) inference engines such as an OgaEngine-style `PagedAttention` model. + +Source: +[bert_defs.cc](../../onnxruntime/core/graph/contrib_ops/bert_defs.cc) (schema), +[sparse_attention_indexer_common.h](../../onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h) +(policy enum, selected-capacity formula and CSA window-plan arithmetic, shared unmodified with +`SparseAttentionIndexer`), +[packed_sparse_attention_indexer_common.h](../../onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h) +(fixed 15-input / 6-output slot map), +[packed_sparse_attention_indexer.cc](../../onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc) / +[packed_sparse_attention_indexer_impl.cu](../../onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu) +(CUDA), +[packed_sparse_attention_indexer.cc](../../onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc) +(WebGPU, see also [the WebGPU note](webgpu/packed_sparse_attention_indexer.md)). + +--- + +## 1. Why a separate operator + +`SparseAttentionIndexer` uses dense `[batch_size, sequence_length, ...]` tensors, an explicit dense +QSA visibility mask, and state that grows by concatenation every call +(`present_key = concat(past_key, key)`, etc.). A continuous-batching / paged engine instead: + +- flattens every request's tokens into one `[total_tokens, ...]` axis (packed layout); +- schedules a different number of new tokens per request per step; +- keeps every request's KV/indexer state in a **fixed-address, fixed-capacity** slot of a state + pool (so the engine can reuse buffers and support CUDA graph capture), never a tensor that grows; +- derives causal visibility purely from packed offsets, never from a materialized `[B, S, T]` mask. + +Restructuring `SparseAttentionIndexer` in place to support both contracts was judged more invasive +and risky than the value of the shared line count (roughly 25-35% of the dense implementation is +directly reusable without change; most of the rest needs new shapes, new state semantics, or a +different launch/index mapping). `PackedSparseAttentionIndexer` is therefore a new op, version 1, +that **does not change `SparseAttentionIndexer`'s schema or behavior**. See [§8](#8-what-is-shared-vs-packed-specific). + +## 2. Operator schema + +Attributes: + +| Attribute | Constraint | Meaning | +|---|---|---| +| `policy_mode` | required, `"qsa"` or `"csa"` | selects the indexer flavour | +| `compress_ratio` | required, `> 0` | tokens folded into one block/compressed entry | +| `state_capacity` | required, `> 0` | fixed capacity (entries) of `past_key_state` | +| `token_budget` | `qsa` only, `> 0`, divisible by `compress_ratio` | selected-token budget | +| `index_topk` | `csa` only, `> 0` | selected compressed-entry count | +| `epsilon` | default `1e-6` | RMSNorm epsilon | +| `scale` | default `1/sqrt(head_size)` | per-head score scale | +| `head_weight_scale` | `csa` only, default `1/sqrt(num_heads)` | head-weight score scale | + +Inputs are **fixed at 15 indices** for both policies (unlike the dense op, which uses a different +input/output count per policy). A slot not owned by the active policy is a *positional* optional: +its `NodeProto` input name is empty rather than the slot being removed from the list, so every +later slot keeps its fixed index. + +| # | Name | Shape | Type | Policy | +|---|---|---|---|---| +| 0 | `query` | `(total_tokens, num_heads, head_size)` | T | both | +| 1 | `key` | `(total_tokens, head_size)` qsa / `(total_tokens, 2*head_size)` csa | T | both | +| 2 | `key_norm_weight` | `(head_size)` | T | both | +| 3 | `cos_cache` | `(max_position, rotary_width)` or `(batch_size, max_position, rotary_width)` | T | both | +| 4 | `sin_cache` | same shape as `cos_cache` | T | both | +| 5 | `cumulative_sequence_lengths` | `(batch_size + 1)` | int32, device-resident | both | +| 6 | `past_sequence_lengths` | `(batch_size)` | int32, device-resident | both | +| 7 | `gate` | `(total_tokens, 2*head_size)` | T | csa only | +| 8 | `position_bias` | `(compress_ratio, 2*head_size)` | T | csa only | +| 9 | `head_weights` | `(total_tokens, num_heads)` | T | csa only | +| 10 | `position_ids` | `(total_tokens)` | int64 | optional qsa / required csa | +| 11 | `past_key_state` | `(batch_size, state_capacity, head_size)` | T | both (generic) | +| 12 | `past_kv_buffer` | `(batch_size, 2*compress_ratio-1, width)` | T | both (generic) | +| 13 | `past_gate_buffer` | same shape as `past_kv_buffer` | T | csa only | +| 14 | `past_state_lengths` | `(batch_size, 2)` | int32, device-resident | both (generic) | + +Outputs are **fixed at 6 indices** for both policies (`present_gate_buffer` is declared with an +empty output name for `qsa`, the same positional-optional convention as above): + +| # | Name | Shape | Type | Policy | +|---|---|---|---|---| +| 0 | `selected_indices` | `(total_tokens, capacity)` | int32 | both | +| 1 | `selected_counts` | `(total_tokens)` | int32 | both | +| 2 | `present_key_state` | same shape as `past_key_state` | T | both | +| 3 | `present_kv_buffer` | same shape as `past_kv_buffer` | T | both | +| 4 | `present_gate_buffer` | same shape as `past_gate_buffer` | T | csa only | +| 5 | `present_state_lengths` | same shape as `past_state_lengths` | int32 | both | + +`capacity` is `token_budget + compress_ratio - 1` for `qsa` and `index_topk` for `csa`, exactly the +same formula (`SelectedCapacity`) used by `SparseAttentionIndexer`. + +**No output shape depends on tensor data.** `total_tokens` and `batch_size` come from input +*shapes* (`query.shape[0]`, `cumulative_sequence_lengths.shape[0] - 1`); every state output has +exactly the same shape as its corresponding state input. This is what makes the fixed-capacity +state design load-bearing: a growing/concatenated state (as in the dense op) would require a +data-dependent output shape, which is incompatible with CUDA graph capture and with pre-allocated +paged state pools. + +## 3. Generic state, shared by both policies + +Both policies read and write the *same four* state slots — there is no separate +`past_compressed_key` vs. `past_key` naming split as in the dense op: + +- `past_key_state` / `present_key_state`: `qsa` stores prepared (already mean-pooled, RMSNorm'd and + rotated) complete-block keys; `csa` stores compressed keys. Layout-compatible with, or cheaply + reshaped to, a `[batch_size, capacity, 1, head_size]` auxiliary paged cache when K = V. +- `past_kv_buffer` / `present_kv_buffer` (and `past_gate_buffer` / `present_gate_buffer`, `csa` + only): the generic pending-token buffer, fixed capacity `2 * compress_ratio - 1`. `qsa` only ever + uses up to `compress_ratio - 1` of these entries (a raw, not-yet-pooled block); `csa` uses the + full range for the overlap ("Ca") plus leftover ("Cb") halves of the window-plan arithmetic + reused from `SparseAttentionIndexer`. +- `past_state_lengths` / `present_state_lengths`: `(batch_size, 2)`. Column 0 is the `key_state` + entry count (`qsa`: complete-block count; `csa`: compressed-entry count); column 1 is the pending + buffer length (`qsa`: incomplete-block length in `[0, compress_ratio)`; `csa`: buffer length in + `[0, 2 * compress_ratio)`, exactly the invariant already documented for + `CsaWindowPlan`/`TryComputeCsaWindowPlan`). + +State never grows. `present_*` always has exactly the same shape as `past_*`; only the *contents* +change. Input/output aliasing is supported: every kernel reads its sources (`past_kv_buffer` / +`key`, `past_gate_buffer` / `gate`) and never re-reads `present_*`, so it is correct whether +`present_*` is a distinct allocation or the same underlying buffer as `past_*`. + +**State overflow.** If a call would close more blocks/windows than +`state_capacity - old_entry_count` allows, that request's step is rejected as a deterministic +no-op: its state and state lengths remain unchanged, and its selection outputs stay empty. This +never reads or writes outside a tensor's fixed extent and never silently truncates state. + +## 4. Packed metadata and device-side safety + +`cumulative_sequence_lengths` and `past_sequence_lengths` are **device-resident** tensors, read +directly by the kernels — never copied to the host or synchronized on. The device-visible +invariants (validated by well-behaved callers; a malformed value never causes memory corruption, +see below) are: + +- `cumulative_sequence_lengths[0] == 0`; +- `cumulative_sequence_lengths[batch_size] == total_tokens`; +- `cumulative_sequence_lengths` is nondecreasing (a repeated offset — a zero-token request row — + is valid and simply contributes no query rows for that request); +- `past_sequence_lengths[b] >= 0` and, for `qsa`, consistent with `past_state_lengths[b]` + (`key_state_length == past_sequence_length / compress_ratio`, + `buffer_length == past_sequence_length % compress_ratio`); +- `past_state_lengths[b, 0] <= state_capacity` and `past_state_lengths[b, 1]` within its policy's + valid buffer range. + +Because there is no host synchronization, the kernels cannot literally raise a C++ exception when +one of these invariants is violated by the input data (as opposed to a mismatched tensor *shape*, +which the host-side `OpKernel::Compute` and the ONNX schema still check the ordinary way). Instead, +every per-request quantity read from these tensors is **clamped into its valid range before use** +(`old_key_len = clamp(past_state_lengths[b,0], 0, state_capacity)`, etc.), and the request-token +lookup (`PackedBatchOfToken`, a binary search over `cumulative_sequence_lengths`) always returns an +index in `[0, batch_size)`. The result is that malformed metadata can make the numeric result wrong +for the affected request, but it can never read or write outside a tensor's allocated extent and +never causes overlapping writes between requests. This mirrors the "prefer deterministic safe +outputs" guidance for EPs that cannot report device-side validation errors asynchronously. + +## 5. Policy `qsa` + +Each request's packed token range is processed independently, in the same three stages as the +dense `qsa` policy but against fixed-capacity state instead of a growing cache: + +1. **Update** (one launch per request): append each raw indexer key to the generic pending buffer; + whenever it reaches `compress_ratio` tokens, mean-pool it, apply RMSNorm and `key_norm_weight`, + apply the leading/split-half rotary convention at the block's first logical token position + (`entry * compress_ratio`, where `entry` is the block's absolute index in `key_state`), and + append the **prepared** (already normalized and rotated) key to `key_state`. This differs from + the dense kernel, which stores *raw* concatenated keys and repeats the pooling/RMSNorm/rotate + work for every query; storing the prepared key once, at update time, is possible only because + packed `key_state` never needs to be re-windowed the way a dense-mask query can. +2. **Score** (one launch per query token, per candidate block): every causally visible block + (`block index < min(key_state_length, causal_threshold(position))`) is scored directly against + `key_state` with `sum_h ReLU(q_h . k)` — a single dot product, no recomputation. +3. **Select** (one launch per query token): keeps the `token_budget / compress_ratio` highest + scoring blocks (ties broken by ascending index) and appends the request-local logical token + positions `[j * compress_ratio, ..., j * compress_ratio + compress_ratio - 1]` for each selected + block `j`, followed by every causally visible position of the current incomplete block, and + writes the exact active count to `selected_counts`. + +"Request-local logical token position" means the same numbering as +`past_sequence_lengths[b] + local_offset` — i.e. the request's own absolute token position, which +is exactly what a per-request main paged KV cache is addressed by. The output is therefore directly +consumable by `SparsePagedAttention` configured with `attention_mode="selected_only"`, +`selected_kv_source="main"`. + +## 6. Policy `csa` + +Reuses `SparseAttentionIndexer`'s CSA compression, rotary, scoring, causal threshold, and +deterministic TopK semantics verbatim (the *math* is unchanged); only the layout, state, and launch +mapping are packed: + +1. **Update** (one launch per request): computes the window plan + (`overlap_length`, `leftover_length`, `new_window_count`, `present_buffer_length`, + `present_buffer_start`) for that request's own `buffer_length` and packed token count using + `TryComputeCsaWindowPlan` — the exact same function used by `SparseAttentionIndexer`'s schema + and kernel, called directly on the device (it is a small `SAI_HOST_DEVICE` inline function with + no CUDA-specific code). Every closed window is compressed with the softmax-gated Ca/Cb pooling, + normalized, rotated with the trailing convention, and appended to `key_state`; the request + rejects (caps) new windows beyond `state_capacity` as described in [§3](#3-generic-state-shared-by-both-policies). +2. **Score** (one launch per query token, per compressed entry): every entry is scored with + `sum_h w_h * ReLU(q_h . k)` and masked by the causal threshold from `position_ids` (required for + `csa`, unlike `qsa` where it is an optional override of the default `past_sequence_length + + local offset`). +3. **Select** (one launch per query token): keeps the `index_topk` highest scoring, causally + visible entries and writes the exact active count to `selected_counts`. + +`selected_indices` values are compressed-entry indices into `key_state`, consumable by +`SparsePagedAttention` configured with `attention_mode="local_plus_selected"`, +`selected_kv_source="auxiliary"`; `key_state` is layout-compatible with, or cheaply reshaped to, the +auxiliary cache contract (`[batch_size, capacity, 1, head_size]` when K = V). + +## 7. Provider support + +CUDA and WebGPU both implement version 1 of this operator, for `float32`, `float16` (CUDA also +`bfloat16`). There is intentionally no CPU kernel (only the shared constants/helpers/schema are +CPU-agnostic); a production model that uses this op targets a paged-KV engine on an accelerator. +See [the WebGPU note](webgpu/packed_sparse_attention_indexer.md) for WebGPU-specific details. + +## 8. What is shared vs. packed-specific + +Shared with `SparseAttentionIndexer`, unmodified: + +- `Policy` enum, `TryParsePolicy`, `SelectedCapacity` (selected-capacity formula) — from + `sparse_attention_indexer_common.h`; +- `CsaWindowPlan` / `TryComputeCsaWindowPlan` (CSA window-plan arithmetic) — same header, now + additionally annotated `SAI_HOST_DEVICE` so CUDA device code can call it directly; +- the CUDA device math (`sparse_attention_indexer_device_math.cuh`, newly extracted from + `sparse_attention_indexer_impl.cu` with no behavior change): FP32 block reductions + (`SaiBlockSum`), deterministic argmax/selection (`SaiBlockArgMax`, `SaiScanForNext`), the leading + and trailing RoPE conventions (`SaiLeadingRope`, `SaiTrailingRope`), and the causal-threshold + formula (`SaiCausalThreshold`, which turns out to be exactly the "number of complete blocks fully + visible to a query" formula needed by both policies here, unifying what the dense implementation + computed two different ways). + +Deliberately **not** shared (packed-specific mechanics with no dense equivalent, or dense-only +mechanics with no packed equivalent): + +- packed metadata validation and the device-side per-token/per-request lookup + (`PackedBatchOfToken`); +- fully in-place fixed-capacity state update (no growing/concatenating state, no host-visible + data-dependent output shape); +- plain causal visibility derived from packed offsets (no dense `[B, 1, S, T]` mask input, no + mask-compaction step); +- the dense op's batch-major `[B, S, ...]` launch/index mapping and host wrappers, which do not + apply to a token-major `[total_tokens, ...]` tensor. + +## 9. Testing + +`onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc` covers: + +- shape inference for `qsa` and `csa` (fixed `selected_indices`/`selected_counts` shapes, fixed + state output shapes, the strict per-slot policy validation, and the always-6-outputs contract); +- multi-request packed batches with unequal token counts, a zero-token request row, and prefill + followed by decode with independent per-request state (CUDA/WebGPU, skipped without the + respective execution provider); +- `qsa` state-capacity overflow safety; +- FP32/FP16 (and CUDA-only BF16) numeric coverage against an in-file reference that mirrors this + document's contract. + +## 10. Known limitations and follow-ups + +- The reference CUDA/WebGPU kernels prioritize correctness over throughput (see the top-of-file + comments in the `.cu`/`.cc` implementations); they are not yet tuned for large `state_capacity` + or long packed batches. +- OgaEngine / Model Builder integration (declaring `past_key_state` etc. as Engine-managed, + per-request fixed-size state, analogous to a paged auxiliary cache) is out of scope for this + operator definition and is expected in a follow-up to `microsoft/onnxruntime-genai`. +- No CPU kernel is provided. diff --git a/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md b/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md new file mode 100644 index 0000000000000..c1ce8a1549358 --- /dev/null +++ b/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md @@ -0,0 +1,48 @@ +# PackedSparseAttentionIndexer on WebGPU + +The WebGPU execution provider implements version 1 of +`com.microsoft.PackedSparseAttentionIndexer` for the `qsa` and `csa` policies. It uses the +provider-neutral schema and generic state ABI described in the +[operator documentation](../packed_sparse_attention_indexer.md). + +## Supported subset + +- packed (`total_tokens`-major) inputs, driven by device-resident + `cumulative_sequence_lengths` / `past_sequence_lengths`; +- `qsa` and `csa` policy modes; +- `float32` and `float16`; +- shared cos/sin rotary cache (`(max_position, rotary_width)`) or request-specific cache + (`(batch_size, max_position, rotary_width)`); +- generic fixed-capacity `past_key_state` / `past_kv_buffer` / `past_gate_buffer` / + `past_state_lengths` state, including input/output aliasing; +- deterministic score-descending, index-ascending top-k ties; +- `selected_counts`, the exact active-entry count per query. + +BF16 is not registered by the WebGPU kernel (CUDA only). Unknown policies and +policy-incompatible inputs or attributes are rejected. + +## Execution + +Every program is one invocation per row (one active thread per workgroup; the rest of the +workgroup is idle), exactly like the dense `SparseAttentionIndexer` WebGPU kernel: state-update +programs dispatch one row per **request**, and the score/select programs dispatch one row per +**query token**. As in the dense kernel, intermediate values (pooled/normalized/rotated keys, +per-candidate scores) are recomputed by small WGSL helper functions on demand rather than staged +into workgroup-shared arrays, both to keep every kernel correct without relying on WGSL arrays +sized by a runtime (uniform) `head_size`, and to keep the packed kernel's structure directly +comparable to the CUDA implementation's per-request/per-token update and score/select stages. All +reductions and softmax calculations accumulate in FP32, including for FP16 inputs. + +Per-request quantities (`cumulative_sequence_lengths`, `past_sequence_lengths`, +`past_state_lengths`) are read directly from device buffers inside the shaders — never on the +host — and are clamped into their valid ranges before use, so malformed packed metadata can never +cause an out-of-bounds buffer access (see the main document's device-side safety section). + +## Follow-up work + +- specialized large-candidate top-k; +- subgroup-optimized reductions; +- fused projection, pooling, and scoring; +- reduced recomputation and temporary-buffer use; +- BF16 support; +- WebGPU `SparsePagedAttention` end-to-end integration. diff --git a/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h b/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h new file mode 100644 index 0000000000000..24b7ce1584071 --- /dev/null +++ b/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h @@ -0,0 +1,85 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +// +// Shared constants for com.microsoft.PackedSparseAttentionIndexer. This op reuses the policy +// enum, selected-capacity formula and CSA window-plan arithmetic already defined for +// com.microsoft.SparseAttentionIndexer in sparse_attention_indexer_common.h; it does not modify +// that header's input/output slot layout, which stays specific to the dense operator. +// +// PackedSparseAttentionIndexer instead uses packed [total_tokens, ...] query/key tensors, +// device-resident cumulative_sequence_lengths / past_sequence_lengths, and generic fixed-capacity +// state slots that are shared by both policy_mode values (unlike the dense op's policy-specific +// state names). + +#pragma once + +#include + +#include "contrib_ops/cpu/sparse/sparse_attention_indexer_common.h" + +namespace onnxruntime { +namespace contrib { +namespace packed_sparse_attention_indexer { + +// Re-exported so callers only need to include this header. +using sparse_attention_indexer::CsaWindowPlan; +using sparse_attention_indexer::kPolicyModeCsa; +using sparse_attention_indexer::kPolicyModeQsa; +using sparse_attention_indexer::Policy; +using sparse_attention_indexer::SelectedCapacity; +using sparse_attention_indexer::TryComputeCsaWindowPlan; +using sparse_attention_indexer::TryParsePolicy; + +// Fixed input slots. Slots 7-9 belong to policy_mode="csa" only; slot 10 (position_ids) is +// optional for "qsa" and required for "csa". Every other slot is required for both policies. +enum InputIndex : int { + kQuery = 0, // [total_tokens, num_heads, head_size] + kKey = 1, // qsa: [total_tokens, head_size]; csa: [total_tokens, 2 * head_size] + kKeyNormWeight = 2, // [head_size] + kCosCache = 3, // [max_position, rotary_width] or [batch_size, max_position, rotary_width] + kSinCache = 4, // same shape as cos_cache + kCumulativeSequenceLengths = 5, // [batch_size + 1], int32 + kPastSequenceLengths = 6, // [batch_size], int32 + kGate = 7, // csa only: [total_tokens, 2 * head_size] + kPositionBias = 8, // csa only: [compress_ratio, 2 * head_size] + kHeadWeights = 9, // csa only: [total_tokens, num_heads] + kPositionIds = 10, // optional (qsa) / required (csa): [total_tokens], int64 + kPastKeyState = 11, // generic: [batch_size, state_capacity, head_size] + kPastKvBuffer = 12, // generic: [batch_size, 2 * compress_ratio - 1, width] + kPastGateBuffer = 13, // csa only: same shape as past_kv_buffer + kPastStateLengths = 14, // generic: [batch_size, 2], int32 + kInputCount = 15, +}; + +// Fixed output slots. present_gate_buffer is declared (with an empty name) but not produced for +// policy_mode="qsa". +enum OutputIndex : int { + kSelectedIndices = 0, // [total_tokens, selected_capacity], int32, unused entries -1 + kSelectedCounts = 1, // [total_tokens], int32 + kPresentKeyState = 2, // same shape as past_key_state + kPresentKvBuffer = 3, // same shape as past_kv_buffer + kPresentGateBuffer = 4, // csa only: same shape as past_gate_buffer + kPresentStateLengths = 5, // [batch_size, 2], int32 + kOutputCount = 6, +}; + +// Every PackedSparseAttentionIndexer node declares all 6 fixed outputs; present_gate_buffer is an +// empty-name optional output for policy_mode="qsa". +constexpr int kFixedOutputCount = kOutputCount; + +// Column layout of past_state_lengths / present_state_lengths. +enum StateLengthColumn : int { + kKeyStateLength = 0, // qsa: complete-block count; csa: compressed-entry count + kBufferLength = 1, // qsa: incomplete-block length in [0, compress_ratio); + // csa: pending buffer length in [0, 2 * compress_ratio) + kStateLengthColumns = 2, +}; + +// Generic pending-buffer capacity: qsa only ever uses up to compress_ratio - 1 of these slots. +SAI_HOST_DEVICE inline int64_t GenericBufferCapacity(int64_t compress_ratio) { + return 2 * compress_ratio - 1; +} + +} // namespace packed_sparse_attention_indexer +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h b/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h index d9eb3d1671cca..341315716a2ef 100644 --- a/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h +++ b/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h @@ -6,6 +6,16 @@ #include #include +// nvcc recognizes __host__/__device__ as built-in qualifiers in any translation unit it compiles +// (no CUDA header include required), but a plain host compiler does not know these tokens. This +// header is shared by CPU-only graph/schema code and by CUDA device code, so the annotation is +// only emitted when nvcc is compiling the translation unit that includes this header. +#if defined(__CUDACC__) +#define SAI_HOST_DEVICE __host__ __device__ +#else +#define SAI_HOST_DEVICE +#endif + namespace onnxruntime { namespace contrib { namespace sparse_attention_indexer { @@ -68,8 +78,8 @@ constexpr int kCsaOutputCount = 5; // Number of selected entries emitted per query. The capacity only depends on attributes, so it is // a compile-time constant of the graph rather than a function of the data. -inline int64_t SelectedCapacity(Policy policy, int64_t token_budget, int64_t index_topk, - int64_t compress_ratio) { +SAI_HOST_DEVICE inline int64_t SelectedCapacity(Policy policy, int64_t token_budget, int64_t index_topk, + int64_t compress_ratio) { return policy == Policy::kQsa ? token_budget + compress_ratio - 1 : index_topk; } @@ -87,8 +97,8 @@ struct CsaWindowPlan { int64_t present_buffer_start = 0; // offset of that buffer inside [past buffer | new tokens] }; -inline bool TryComputeCsaWindowPlan(int64_t past_buffer_length, int64_t sequence_length, - int64_t compress_ratio, CsaWindowPlan& plan) { +SAI_HOST_DEVICE inline bool TryComputeCsaWindowPlan(int64_t past_buffer_length, int64_t sequence_length, + int64_t compress_ratio, CsaWindowPlan& plan) { if (compress_ratio <= 0 || sequence_length < 0 || past_buffer_length < 0 || past_buffer_length >= 2 * compress_ratio) { return false; diff --git a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc index f9165bb5074e1..901bba840d994 100644 --- a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc @@ -256,6 +256,9 @@ class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, SparseAttention); class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, SparseAttentionIndexer); class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, SparseAttentionIndexer); class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, SparseAttentionIndexer); +class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, PackedSparseAttentionIndexer); +class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, PackedSparseAttentionIndexer); +class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, PackedSparseAttentionIndexer); class CUDA_MS_OP_THREE_TYPED_CLASS_NAME(1, uint8_t, float, int32_t, GatherBlockQuantized); class CUDA_MS_OP_THREE_TYPED_CLASS_NAME(1, uint8_t, MLFloat16, int32_t, GatherBlockQuantized); @@ -565,6 +568,9 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc new file mode 100644 index 0000000000000..eb23d13acd8b0 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc @@ -0,0 +1,437 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h" + +#include +#include +#include +#include + +#include "contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h" +#include "core/providers/cuda/cuda_common.h" +#include "core/providers/cuda/cuda_type_conversion.h" + +namespace onnxruntime { +namespace contrib { +namespace cuda { + +using namespace onnxruntime::cuda; +namespace psai = onnxruntime::contrib::packed_sparse_attention_indexer; + +#define REGISTER_KERNEL_TYPED(T) \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + PackedSparseAttentionIndexer, \ + kMSDomain, \ + 1, \ + T, \ + kCudaExecutionProvider, \ + (*KernelDefBuilder::Create()) \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("I", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("M", DataTypeImpl::GetTensorType()), \ + PackedSparseAttentionIndexer); + +REGISTER_KERNEL_TYPED(float) +REGISTER_KERNEL_TYPED(MLFloat16) +REGISTER_KERNEL_TYPED(BFloat16) + +#undef REGISTER_KERNEL_TYPED + +namespace { + +Status CheckShape(const Tensor* tensor, const char* name, std::initializer_list expected) { + ORT_RETURN_IF(tensor == nullptr, "PackedSparseAttentionIndexer: ", name, " is required"); + const TensorShape expected_shape(expected); + ORT_RETURN_IF_NOT(tensor->Shape() == expected_shape, "PackedSparseAttentionIndexer: ", name, " must have shape ", + expected_shape.ToString(), ", got ", tensor->Shape().ToString()); + return Status::OK(); +} + +Status CheckIntDimension(const char* name, int64_t value, bool allow_zero = true) { + ORT_RETURN_IF(value < (allow_zero ? 0 : 1) || value > std::numeric_limits::max(), + "PackedSparseAttentionIndexer: ", name, " must be in ", allow_zero ? "[0, INT_MAX]" : "(0, INT_MAX]", + ", got ", value); + return Status::OK(); +} + +// cos_cache / sin_cache may be shared across the batch ([max_position, rotary_width]) or +// request-specific ([batch_size, max_position, rotary_width]). +struct RotaryCacheShape { + bool batched; + int64_t max_rotary_length; + int64_t rotary_width; +}; + +Status CheckRotaryCache(const Tensor* cos_cache, const Tensor* sin_cache, int64_t batch_size, + RotaryCacheShape& out) { + ORT_RETURN_IF(cos_cache == nullptr, "PackedSparseAttentionIndexer: cos_cache is required"); + const auto& cos_shape = cos_cache->Shape(); + out.batched = cos_shape.NumDimensions() == 3; + ORT_RETURN_IF_NOT( + (out.batched && cos_shape[0] == batch_size && cos_shape[1] > 0) || + (cos_shape.NumDimensions() == 2 && cos_shape[0] > 0), + "PackedSparseAttentionIndexer: cos_cache must have shape (max_position, rotary_width) or " + "(batch_size, max_position, rotary_width), got ", + cos_shape.ToString()); + out.max_rotary_length = out.batched ? cos_shape[1] : cos_shape[0]; + out.rotary_width = out.batched ? cos_shape[2] : cos_shape[1]; + ORT_RETURN_IF_ERROR(CheckIntDimension("max_rotary_sequence_length", out.max_rotary_length, false)); + ORT_RETURN_IF_ERROR(CheckIntDimension("rotary_width", out.rotary_width, false)); + ORT_RETURN_IF_NOT(sin_cache != nullptr && sin_cache->Shape() == cos_shape, + "PackedSparseAttentionIndexer: sin_cache must have the same shape as cos_cache"); + return Status::OK(); +} + +} // namespace + +template +PackedSparseAttentionIndexer::PackedSparseAttentionIndexer(const OpKernelInfo& info) : CudaKernel(info) { + std::string policy_mode; + ORT_ENFORCE(info.GetAttr("policy_mode", &policy_mode).IsOK(), + "PackedSparseAttentionIndexer: policy_mode is required"); + ORT_ENFORCE(psai::TryParsePolicy(policy_mode, policy_), "PackedSparseAttentionIndexer: policy_mode must be '", + psai::kPolicyModeQsa, "' or '", psai::kPolicyModeCsa, "', got '", policy_mode, "'"); + + ORT_ENFORCE(info.GetAttr("compress_ratio", &compress_ratio_).IsOK(), + "PackedSparseAttentionIndexer: compress_ratio is required"); + ORT_ENFORCE(compress_ratio_ > 0 && compress_ratio_ <= std::numeric_limits::max(), + "PackedSparseAttentionIndexer: compress_ratio must be in (0, INT_MAX], got ", compress_ratio_); + + int64_t state_capacity = 0; + ORT_ENFORCE(info.GetAttr("state_capacity", &state_capacity).IsOK(), + "PackedSparseAttentionIndexer: state_capacity is required"); + ORT_ENFORCE(state_capacity > 0 && state_capacity <= std::numeric_limits::max(), + "PackedSparseAttentionIndexer: state_capacity must be in (0, INT_MAX], got ", state_capacity); + + const bool has_token_budget = info.GetAttr("token_budget", &token_budget_).IsOK(); + const bool has_index_topk = info.GetAttr("index_topk", &index_topk_).IsOK(); + float head_weight_scale = 0.0f; + has_head_weight_scale_ = info.GetAttr("head_weight_scale", &head_weight_scale).IsOK(); + + if (policy_ == psai::Policy::kQsa) { + ORT_ENFORCE(has_token_budget, + "PackedSparseAttentionIndexer: token_budget is required when policy_mode is 'qsa'"); + ORT_ENFORCE(!has_index_topk && !has_head_weight_scale_, + "PackedSparseAttentionIndexer: index_topk and head_weight_scale must not be set when policy_mode " + "is 'qsa'"); + ORT_ENFORCE(token_budget_ > 0 && token_budget_ % compress_ratio_ == 0 && + token_budget_ <= std::numeric_limits::max() - compress_ratio_ + 1, + "PackedSparseAttentionIndexer: token_budget must be > 0, divisible by compress_ratio, and produce " + "a selected capacity no greater than INT_MAX, got token_budget=", + token_budget_, " compress_ratio=", compress_ratio_); + index_topk_ = 0; + } else { + ORT_ENFORCE(has_index_topk, "PackedSparseAttentionIndexer: index_topk is required when policy_mode is 'csa'"); + ORT_ENFORCE(!has_token_budget, + "PackedSparseAttentionIndexer: token_budget must not be set when policy_mode is 'csa'"); + ORT_ENFORCE(index_topk_ > 0 && index_topk_ <= std::numeric_limits::max(), + "PackedSparseAttentionIndexer: index_topk must be in (0, INT_MAX], got ", index_topk_); + token_budget_ = 0; + } + + epsilon_ = info.GetAttrOrDefault("epsilon", 1.0e-6f); + ORT_ENFORCE(epsilon_ >= 0.0f, "PackedSparseAttentionIndexer: epsilon must be >= 0, got ", epsilon_); + has_scale_ = info.GetAttr("scale", &scale_).IsOK(); + head_weight_scale_ = head_weight_scale; +} + +template +Status PackedSparseAttentionIndexer::ComputeInternal(OpKernelContext* context) const { + const bool is_qsa = policy_ == psai::Policy::kQsa; + constexpr int kCsaOnlyInputs[] = {psai::kGate, psai::kPositionBias, psai::kHeadWeights}; + for (int index : kCsaOnlyInputs) { + const bool provided = index < context->InputCount() && context->Input(index) != nullptr; + ORT_RETURN_IF(provided != !is_qsa, "PackedSparseAttentionIndexer: input ", index, + provided ? " must be omitted for policy_mode 'qsa'" : " is required for policy_mode 'csa'"); + } + const bool position_ids_provided = + psai::kPositionIds < context->InputCount() && context->Input(psai::kPositionIds) != nullptr; + ORT_RETURN_IF(!is_qsa && !position_ids_provided, + "PackedSparseAttentionIndexer: input ", psai::kPositionIds, + " (position_ids) is required for policy_mode 'csa'"); + const bool gate_buffer_provided = + psai::kPastGateBuffer < context->InputCount() && context->Input(psai::kPastGateBuffer) != nullptr; + ORT_RETURN_IF(gate_buffer_provided != !is_qsa, "PackedSparseAttentionIndexer: input ", psai::kPastGateBuffer, + gate_buffer_provided ? " must be omitted for policy_mode 'qsa'" + : " is required for policy_mode 'csa'"); + + return is_qsa ? ComputeQsa(context) : ComputeCsa(context); +} + +template +Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) const { + using CudaT = typename OrtToCudaType::type; + + const Tensor* query = context->Input(psai::kQuery); + const Tensor* key = context->Input(psai::kKey); + const Tensor* key_norm_weight = context->Input(psai::kKeyNormWeight); + const Tensor* cos_cache = context->Input(psai::kCosCache); + const Tensor* sin_cache = context->Input(psai::kSinCache); + const Tensor* cumulative_sequence_lengths = context->Input(psai::kCumulativeSequenceLengths); + const Tensor* past_sequence_lengths = context->Input(psai::kPastSequenceLengths); + const Tensor* position_ids = context->Input(psai::kPositionIds); + const Tensor* past_key_state = context->Input(psai::kPastKeyState); + const Tensor* past_kv_buffer = context->Input(psai::kPastKvBuffer); + const Tensor* past_state_lengths = context->Input(psai::kPastStateLengths); + + ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); + const auto& query_shape = query->Shape(); + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, + "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads, head_size), " + "got ", + query_shape.ToString()); + const int64_t total_tokens = query_shape[0]; + const int64_t num_heads = query_shape[1]; + const int64_t head_size = query_shape[2]; + ORT_RETURN_IF_ERROR(CheckIntDimension("total_tokens", total_tokens)); + ORT_RETURN_IF_ERROR(CheckIntDimension("num_heads", num_heads, false)); + ORT_RETURN_IF_ERROR(CheckIntDimension("head_size", head_size, false)); + + ORT_RETURN_IF(cumulative_sequence_lengths == nullptr, + "PackedSparseAttentionIndexer: cumulative_sequence_lengths is required"); + const auto& cu_shape = cumulative_sequence_lengths->Shape(); + ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1, + "PackedSparseAttentionIndexer: cumulative_sequence_lengths must have shape (batch_size + 1), " + "got ", + cu_shape.ToString()); + const int64_t batch_size = cu_shape[0] - 1; + ORT_RETURN_IF_ERROR(CheckIntDimension("batch_size", batch_size)); + + ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size})); + ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, head_size})); + ORT_RETURN_IF_ERROR(CheckShape(key_norm_weight, "key_norm_weight", {head_size})); + if (position_ids != nullptr) { + ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens})); + } + + RotaryCacheShape rotary; + ORT_RETURN_IF_ERROR(CheckRotaryCache(cos_cache, sin_cache, batch_size, rotary)); + ORT_RETURN_IF_NOT( + rotary.rotary_width > 0 && rotary.rotary_width % 2 == 0 && rotary.rotary_width <= head_size, + "PackedSparseAttentionIndexer: policy_mode 'qsa' requires an even rotary_width in (0, head_size], got ", + rotary.rotary_width); + + ORT_RETURN_IF(past_key_state == nullptr, "PackedSparseAttentionIndexer: past_key_state is required"); + const auto& key_state_shape = past_key_state->Shape(); + ORT_RETURN_IF_NOT(key_state_shape.NumDimensions() == 3 && key_state_shape[0] == batch_size && + key_state_shape[2] == head_size, + "PackedSparseAttentionIndexer: past_key_state must have shape " + "(batch_size, state_capacity, head_size), got ", + key_state_shape.ToString()); + const int64_t state_capacity = key_state_shape[1]; + ORT_RETURN_IF_ERROR(CheckIntDimension("state_capacity", state_capacity, false)); + + const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); + ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, head_size})); + ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths", + {batch_size, psai::kStateLengthColumns})); + + const int64_t capacity = psai::SelectedCapacity(psai::Policy::kQsa, token_budget_, index_topk_, compress_ratio_); + + Tensor* selected_indices = context->Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity})); + Tensor* selected_counts = context->Output(psai::kSelectedCounts, TensorShape({total_tokens})); + Tensor* present_key_state = context->Output(psai::kPresentKeyState, key_state_shape); + Tensor* present_kv_buffer = + context->Output(psai::kPresentKvBuffer, TensorShape({batch_size, buffer_capacity, head_size})); + Tensor* present_state_lengths = + context->Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns})); + ORT_RETURN_IF(selected_indices == nullptr || selected_counts == nullptr || present_key_state == nullptr || + present_kv_buffer == nullptr || present_state_lengths == nullptr, + "PackedSparseAttentionIndexer: policy_mode 'qsa' requires selected_indices, selected_counts, " + "present_key_state, present_kv_buffer and present_state_lengths outputs"); + + PackedSparseAttentionIndexerParams params; + params.batch_size = static_cast(batch_size); + params.total_tokens = static_cast(total_tokens); + params.num_heads = static_cast(num_heads); + params.head_size = static_cast(head_size); + params.rotary_width = static_cast(rotary.rotary_width); + params.max_rotary_length = static_cast(rotary.max_rotary_length); + params.cos_cache_batched = rotary.batched; + params.compress_ratio = static_cast(compress_ratio_); + params.state_capacity = static_cast(state_capacity); + params.buffer_capacity = static_cast(buffer_capacity); + params.capacity = static_cast(capacity); + params.has_position_ids = position_ids != nullptr; + params.epsilon = epsilon_; + params.scale = has_scale_ ? scale_ : 1.0f / std::sqrt(static_cast(head_size)); + params.block_topk = static_cast(token_budget_ / compress_ratio_); + + auto float_workspace = GetScratchBuffer(GetQsaPackedWorkspaceFloatCount(params), GetComputeStream(context)); + auto overflow_flags = + GetScratchBuffer(static_cast(std::max(batch_size, 1)), GetComputeStream(context)); + + return LaunchQsaPackedSparseAttentionIndexer( + Stream(context), params, + reinterpret_cast(query->Data()), + reinterpret_cast(key->Data()), + reinterpret_cast(key_norm_weight->Data()), + reinterpret_cast(cos_cache->Data()), + reinterpret_cast(sin_cache->Data()), + cumulative_sequence_lengths->Data(), + past_sequence_lengths->Data(), + position_ids != nullptr ? position_ids->Data() : nullptr, + reinterpret_cast(past_key_state->Data()), + reinterpret_cast(past_kv_buffer->Data()), + past_state_lengths->Data(), + selected_indices->MutableData(), + selected_counts->MutableData(), + reinterpret_cast(present_key_state->MutableData()), + reinterpret_cast(present_kv_buffer->MutableData()), + present_state_lengths->MutableData(), + float_workspace.get(), + overflow_flags.get()); +} + +template +Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) const { + using CudaT = typename OrtToCudaType::type; + + const Tensor* query = context->Input(psai::kQuery); + const Tensor* key = context->Input(psai::kKey); + const Tensor* key_norm_weight = context->Input(psai::kKeyNormWeight); + const Tensor* cos_cache = context->Input(psai::kCosCache); + const Tensor* sin_cache = context->Input(psai::kSinCache); + const Tensor* cumulative_sequence_lengths = context->Input(psai::kCumulativeSequenceLengths); + const Tensor* past_sequence_lengths = context->Input(psai::kPastSequenceLengths); + const Tensor* gate = context->Input(psai::kGate); + const Tensor* position_bias = context->Input(psai::kPositionBias); + const Tensor* head_weights = context->Input(psai::kHeadWeights); + const Tensor* position_ids = context->Input(psai::kPositionIds); + const Tensor* past_key_state = context->Input(psai::kPastKeyState); + const Tensor* past_kv_buffer = context->Input(psai::kPastKvBuffer); + const Tensor* past_gate_buffer = context->Input(psai::kPastGateBuffer); + const Tensor* past_state_lengths = context->Input(psai::kPastStateLengths); + + ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); + const auto& query_shape = query->Shape(); + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, + "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads, head_size), " + "got ", + query_shape.ToString()); + const int64_t total_tokens = query_shape[0]; + const int64_t num_heads = query_shape[1]; + const int64_t head_size = query_shape[2]; + ORT_RETURN_IF_ERROR(CheckIntDimension("total_tokens", total_tokens)); + ORT_RETURN_IF_ERROR(CheckIntDimension("num_heads", num_heads, false)); + ORT_RETURN_IF_ERROR(CheckIntDimension("head_size", head_size, false)); + ORT_RETURN_IF(head_size > std::numeric_limits::max() / 2, + "PackedSparseAttentionIndexer: 2 * head_size must be no greater than INT_MAX"); + const int64_t width = 2 * head_size; + + ORT_RETURN_IF(cumulative_sequence_lengths == nullptr, + "PackedSparseAttentionIndexer: cumulative_sequence_lengths is required"); + const auto& cu_shape = cumulative_sequence_lengths->Shape(); + ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1, + "PackedSparseAttentionIndexer: cumulative_sequence_lengths must have shape (batch_size + 1), " + "got ", + cu_shape.ToString()); + const int64_t batch_size = cu_shape[0] - 1; + ORT_RETURN_IF_ERROR(CheckIntDimension("batch_size", batch_size)); + + ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size})); + ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, width})); + ORT_RETURN_IF_ERROR(CheckShape(key_norm_weight, "key_norm_weight", {head_size})); + ORT_RETURN_IF_ERROR(CheckShape(gate, "gate", {total_tokens, width})); + ORT_RETURN_IF_ERROR(CheckShape(position_bias, "position_bias", {compress_ratio_, width})); + ORT_RETURN_IF_ERROR(CheckShape(head_weights, "head_weights", {total_tokens, num_heads})); + ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens})); + + RotaryCacheShape rotary; + ORT_RETURN_IF_ERROR(CheckRotaryCache(cos_cache, sin_cache, batch_size, rotary)); + ORT_RETURN_IF_NOT(rotary.rotary_width > 0 && 2 * rotary.rotary_width <= head_size, + "PackedSparseAttentionIndexer: policy_mode 'csa' requires 0 < 2 * rotary_width <= head_size, " + "got rotary_width=", + rotary.rotary_width, " head_size=", head_size); + + ORT_RETURN_IF(past_key_state == nullptr, "PackedSparseAttentionIndexer: past_key_state is required"); + const auto& key_state_shape = past_key_state->Shape(); + ORT_RETURN_IF_NOT(key_state_shape.NumDimensions() == 3 && key_state_shape[0] == batch_size && + key_state_shape[2] == head_size, + "PackedSparseAttentionIndexer: past_key_state must have shape " + "(batch_size, state_capacity, head_size), got ", + key_state_shape.ToString()); + const int64_t state_capacity = key_state_shape[1]; + ORT_RETURN_IF_ERROR(CheckIntDimension("state_capacity", state_capacity, false)); + + const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); + ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, width})); + ORT_RETURN_IF_ERROR(CheckShape(past_gate_buffer, "past_gate_buffer", {batch_size, buffer_capacity, width})); + ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths", + {batch_size, psai::kStateLengthColumns})); + + const int64_t capacity = psai::SelectedCapacity(psai::Policy::kCsa, token_budget_, index_topk_, compress_ratio_); + + Tensor* selected_indices = context->Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity})); + Tensor* selected_counts = context->Output(psai::kSelectedCounts, TensorShape({total_tokens})); + Tensor* present_key_state = context->Output(psai::kPresentKeyState, key_state_shape); + Tensor* present_kv_buffer = + context->Output(psai::kPresentKvBuffer, TensorShape({batch_size, buffer_capacity, width})); + Tensor* present_gate_buffer = + context->Output(psai::kPresentGateBuffer, TensorShape({batch_size, buffer_capacity, width})); + Tensor* present_state_lengths = + context->Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns})); + ORT_RETURN_IF(selected_indices == nullptr || selected_counts == nullptr || present_key_state == nullptr || + present_kv_buffer == nullptr || present_gate_buffer == nullptr || + present_state_lengths == nullptr, + "PackedSparseAttentionIndexer: policy_mode 'csa' requires selected_indices, selected_counts, " + "present_key_state, present_kv_buffer, present_gate_buffer and present_state_lengths outputs"); + + PackedSparseAttentionIndexerParams params; + params.batch_size = static_cast(batch_size); + params.total_tokens = static_cast(total_tokens); + params.num_heads = static_cast(num_heads); + params.head_size = static_cast(head_size); + params.rotary_width = static_cast(rotary.rotary_width); + params.max_rotary_length = static_cast(rotary.max_rotary_length); + params.cos_cache_batched = rotary.batched; + params.compress_ratio = static_cast(compress_ratio_); + params.state_capacity = static_cast(state_capacity); + params.buffer_capacity = static_cast(buffer_capacity); + params.capacity = static_cast(capacity); + params.has_position_ids = true; + params.epsilon = epsilon_; + params.scale = has_scale_ ? scale_ : 1.0f / std::sqrt(static_cast(head_size)); + params.index_topk = static_cast(index_topk_); + params.head_weight_scale = + has_head_weight_scale_ ? head_weight_scale_ : 1.0f / std::sqrt(static_cast(num_heads)); + + auto float_workspace = GetScratchBuffer(GetCsaPackedWorkspaceFloatCount(params), GetComputeStream(context)); + auto overflow_flags = + GetScratchBuffer(static_cast(std::max(batch_size, 1)), GetComputeStream(context)); + + return LaunchCsaPackedSparseAttentionIndexer( + Stream(context), params, + reinterpret_cast(query->Data()), + reinterpret_cast(key->Data()), + reinterpret_cast(key_norm_weight->Data()), + reinterpret_cast(cos_cache->Data()), + reinterpret_cast(sin_cache->Data()), + reinterpret_cast(gate->Data()), + reinterpret_cast(position_bias->Data()), + reinterpret_cast(head_weights->Data()), + cumulative_sequence_lengths->Data(), + past_sequence_lengths->Data(), + position_ids->Data(), + reinterpret_cast(past_key_state->Data()), + reinterpret_cast(past_kv_buffer->Data()), + reinterpret_cast(past_gate_buffer->Data()), + past_state_lengths->Data(), + selected_indices->MutableData(), + selected_counts->MutableData(), + reinterpret_cast(present_key_state->MutableData()), + reinterpret_cast(present_kv_buffer->MutableData()), + reinterpret_cast(present_gate_buffer->MutableData()), + present_state_lengths->MutableData(), + float_workspace.get(), + overflow_flags.get()); +} + +template class PackedSparseAttentionIndexer; +template class PackedSparseAttentionIndexer; +template class PackedSparseAttentionIndexer; + +} // namespace cuda +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h new file mode 100644 index 0000000000000..b7ddfa9b480ce --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h @@ -0,0 +1,37 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h" +#include "core/common/common.h" +#include "core/providers/cuda/cuda_kernel.h" + +namespace onnxruntime { +namespace contrib { +namespace cuda { + +template +class PackedSparseAttentionIndexer final : public onnxruntime::cuda::CudaKernel { + public: + explicit PackedSparseAttentionIndexer(const OpKernelInfo& info); + Status ComputeInternal(OpKernelContext* context) const override; + + private: + Status ComputeQsa(OpKernelContext* context) const; + Status ComputeCsa(OpKernelContext* context) const; + + packed_sparse_attention_indexer::Policy policy_; + int64_t compress_ratio_; + int64_t token_budget_; + int64_t index_topk_; + float epsilon_; + float scale_; + float head_weight_scale_; + bool has_scale_; + bool has_head_weight_scale_; +}; + +} // namespace cuda +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu new file mode 100644 index 0000000000000..36f1a6a521b9b --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu @@ -0,0 +1,868 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +// +// Correctness-first implementation of com.microsoft.PackedSparseAttentionIndexer. It shares its +// block reductions, deterministic argmax/selection helpers, RoPE math and causal-threshold formula +// with com.microsoft.SparseAttentionIndexer via sparse_attention_indexer_device_math.cuh, and +// reuses the CSA window-plan arithmetic directly on the device from +// sparse_attention_indexer_common.h (its helpers are SAI_HOST_DEVICE). What is packed-specific: +// token-major (not batch-major) layout, cumulative_sequence_lengths / past_sequence_lengths driven +// per-request bookkeeping, fully in-place fixed-capacity state update (no growing/concatenating +// state), and plain causal visibility derived from packed metadata (no dense mask). +// +// Device-side safety: every per-request quantity (past_sequence_lengths, past_state_lengths, +// cumulative offsets) is read directly from device memory inside the kernels below -- there is no +// host readback or stream synchronization. Values are always clamped into the fixed-capacity range +// before use, so malformed metadata can make the result semantically wrong but can never cause an +// out-of-bounds access or an overlapping write. State-capacity overflow is *rejected*, not +// silently truncated: if a request's new blocks/windows would not all fit in state_capacity this +// call, the update kernel applies none of them (present_state_lengths / present_key_state / +// present_kv_buffer / present_gate_buffer for that request are left exactly as their past_* +// counterparts) and records the rejection in a small overflow_flags workspace; the select kernels +// then force that request's selected_indices/selected_counts to the deterministic safe empty +// result (-1 / 0) for this call instead of selecting against a partially updated state. See +// QsaUpdateStateKernel / CsaUpdateStateKernel / QsaSelectKernel / CsaSelectKernel. +// +// See docs/contrib_ops/cuda/packed_sparse_attention_indexer.md for the full operator contract. + +#include "contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h" + +#include +#include +#include + +#include + +#include "contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h" +#include "contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh" +#include "core/providers/cuda/cu_inc/cuda_type_helper.cuh" + +namespace onnxruntime { +namespace contrib { +namespace cuda { + +namespace psai = onnxruntime::contrib::packed_sparse_attention_indexer; + +namespace { + +// The block reductions below halve the active thread count, so this must stay a power of two. +constexpr int kThreads = 128; + +// Largest b such that cumulative_sequence_lengths[b] <= token, assuming the array is nondecreasing. +// If the data itself is malformed this may attribute a token to the wrong request, but the result +// is always an index in [0, batch_size), so it can never cause an out-of-bounds access. +__device__ __forceinline__ int PackedBatchOfToken(const int32_t* cumulative_sequence_lengths, int batch_size, + int token) { + int lo = 0; + int hi = batch_size - 1; + while (lo < hi) { + const int mid = lo + (hi - lo + 1) / 2; + if (cumulative_sequence_lengths[mid] <= token) { + lo = mid; + } else { + hi = mid - 1; + } + } + return lo; +} + +template +__global__ void ElementwiseCopyKernel(const T* src, T* dst, int64_t count) { + for (int64_t i = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; i < count; + i += static_cast(gridDim.x) * blockDim.x) { + dst[i] = src[i]; + } +} + +// --------------------------------------------------------------------------------------------- +// policy_mode = "qsa" +// --------------------------------------------------------------------------------------------- + +// One block per request: forms every newly-closed compress_ratio block (mean-pool -> RMSNorm -> +// leading RoPE -> append into fixed-capacity key_state) and publishes the raw trailing buffer. +// Reads only past_kv_buffer / key (never present_kv_buffer / present_key_state), so it is correct +// whether or not the present/past tensors are the same aliased allocation. +template +__global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, const T* cos_cache, + const T* sin_cache, const int32_t* cumulative_sequence_lengths, + const T* past_kv_buffer, const int32_t* past_state_lengths, + T* present_key_state, T* present_kv_buffer, + int32_t* present_state_lengths, int32_t* overflow_flags, + PackedSparseAttentionIndexerParams params) { + extern __shared__ float shared[]; + float* pooled = shared; + float* rotated = shared + params.head_size; + float* reduction = shared + 2 * params.head_size; + + for (int b = static_cast(blockIdx.x); b < params.batch_size; b += static_cast(gridDim.x)) { + const int req_start = cumulative_sequence_lengths[b]; + const int req_end = cumulative_sequence_lengths[b + 1]; + const int req_len = req_end > req_start ? req_end - req_start : 0; + + const int old_key_len = min(max(past_state_lengths[b * 2 + psai::kKeyStateLength], 0), params.state_capacity); + const int old_buf_len = + min(max(past_state_lengths[b * 2 + psai::kBufferLength], 0), params.compress_ratio - 1); + + const int pending = old_buf_len + req_len; + const int full_new_block_count = pending / params.compress_ratio; + const int capacity_left = params.state_capacity - old_key_len; // >= 0 by construction of old_key_len + // Reject (do not partially apply) a step that would need more than the fixed state_capacity: + // no new blocks are formed and the buffer is left exactly as it was, so a rejected step is a + // deterministic no-op on state rather than a silent partial truncation. + const bool overflowed = full_new_block_count > capacity_left; + const int new_block_count = overflowed ? 0 : full_new_block_count; + const int new_buf_len = overflowed ? old_buf_len : (pending % params.compress_ratio); + + // Barrier: every thread has now read past_state_lengths (identically) before any thread below + // writes present_state_lengths, which keeps this correct even if the two tensors alias. + __syncthreads(); + + if (threadIdx.x == 0) { + present_state_lengths[b * 2 + psai::kKeyStateLength] = old_key_len + new_block_count; + present_state_lengths[b * 2 + psai::kBufferLength] = new_buf_len; + overflow_flags[b] = overflowed ? 1 : 0; + } + + for (int k = 0; k < new_block_count; ++k) { + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + float sum = 0.0f; + for (int t = 0; t < params.compress_ratio; ++t) { + const int virtual_pos = k * params.compress_ratio + t; + sum += virtual_pos < old_buf_len + ? to_float(past_kv_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) * + params.head_size + + d]) + : to_float(key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * + params.head_size + + d]); + } + pooled[d] = sum / static_cast(params.compress_ratio); + } + __syncthreads(); + + float sum_squares = 0.0f; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + sum_squares += pooled[d] * pooled[d]; + } + sum_squares = SaiBlockSum(sum_squares, reduction); + const float inverse_rms = rsqrtf(sum_squares / static_cast(params.head_size) + params.epsilon); + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + pooled[d] = pooled[d] * inverse_rms * to_float(key_norm_weight[d]); + } + __syncthreads(); + + const int entry = old_key_len + k; + const int rope_position = + SaiClampPosition(static_cast(entry) * params.compress_ratio, params.max_rotary_length); + const int64_t cache_offset = + (static_cast(params.cos_cache_batched ? b : 0) * params.max_rotary_length + rope_position) * + params.rotary_width; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + rotated[d] = SaiLeadingRope(pooled, params.rotary_width, cos_cache + cache_offset, + sin_cache + cache_offset, d); + } + __syncthreads(); + + const int64_t out_base = (static_cast(b) * params.state_capacity + entry) * params.head_size; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + present_key_state[out_base + d] = from_float(rotated[d]); + } + __syncthreads(); + } + + // Publish the raw trailing buffer. Skipped entirely on overflow: present_kv_buffer already + // holds past_kv_buffer's contents unchanged (from the baseline copy in the Launch function + // below), which is exactly the prior valid buffer this rejected step must preserve. + if (!overflowed) { + for (int t = static_cast(threadIdx.x); t < new_buf_len; t += static_cast(blockDim.x)) { + const int virtual_pos = new_block_count * params.compress_ratio + t; + const int64_t out_base = (static_cast(b) * params.buffer_capacity + t) * params.head_size; + for (int d = 0; d < params.head_size; ++d) { + const float value = + virtual_pos < old_buf_len + ? to_float( + past_kv_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) * + params.head_size + + d]) + : to_float( + key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * params.head_size + d]); + present_kv_buffer[out_base + d] = from_float(value); + } + } + } + __syncthreads(); + } +} + +// One block per (token, head): rotates the query once so downstream scoring kernels only ever dot +// two already-rotated/prepared vectors. kUseLeadingRope selects the qsa convention (position +// defaults to past_sequence_lengths[batch] + request-local offset when position_ids is absent); +// otherwise the csa trailing convention with positions always taken from position_ids. +template +__global__ void PackedRotateQueryKernel(const T* query, const T* cos_cache, const T* sin_cache, + const int32_t* cumulative_sequence_lengths, + const int32_t* past_sequence_lengths, const int64_t* position_ids, + float* query_rotated, PackedSparseAttentionIndexerParams params) { + extern __shared__ float shared[]; + const int64_t rows = static_cast(params.total_tokens) * params.num_heads; + for (int64_t row = blockIdx.x; row < rows; row += gridDim.x) { + const int token = static_cast(row / params.num_heads); + const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token); + const int64_t base = row * params.head_size; + + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + shared[d] = to_float(query[base + d]); + } + __syncthreads(); + + const int64_t abs_position = + params.has_position_ids + ? position_ids[token] + : static_cast(past_sequence_lengths[batch]) + (token - cumulative_sequence_lengths[batch]); + const int position = SaiClampPosition(abs_position, params.max_rotary_length); + const int64_t cache_offset = + (static_cast(params.cos_cache_batched ? batch : 0) * params.max_rotary_length + position) * + params.rotary_width; + const T* cos_row = cos_cache + cache_offset; + const T* sin_row = sin_cache + cache_offset; + + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + query_rotated[base + d] = kUseLeadingRope + ? SaiLeadingRope(shared, params.rotary_width, cos_row, sin_row, d) + : SaiTrailingRope(shared, params.head_size, params.rotary_width, cos_row, + sin_row, d); + } + __syncthreads(); + } +} + +// One block per (token, key_state slot). Scores are directly dotted against the already-prepared +// present_key_state entry (unlike the dense op, no per-query pooling/normalize/rotate is repeated +// here because the packed contract stores fully-prepared blocks in key_state). +template +__global__ void QsaBlockScoreKernel(const T* present_key_state, const float* query_rotated, + const int32_t* cumulative_sequence_lengths, + const int32_t* past_sequence_lengths, const int64_t* position_ids, + const int32_t* present_state_lengths, float* block_scores, + PackedSparseAttentionIndexerParams params) { + extern __shared__ float reduction[]; + const int64_t total = static_cast(params.total_tokens) * params.state_capacity; + for (int64_t work = blockIdx.x; work < total; work += gridDim.x) { + const int token = static_cast(work / params.state_capacity); + const int block_index = static_cast(work % params.state_capacity); + const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token); + const int key_len_after = present_state_lengths[batch * 2 + psai::kKeyStateLength]; + + if (block_index >= key_len_after) { + if (threadIdx.x == 0) { + block_scores[work] = SaiNegativeInfinity(); + } + continue; + } + + const int64_t abs_position = + params.has_position_ids + ? position_ids[token] + : static_cast(past_sequence_lengths[batch]) + (token - cumulative_sequence_lengths[batch]); + const int64_t causal_count = SaiCausalThreshold(abs_position, params.compress_ratio); + const int64_t visible_block_count = causal_count < key_len_after ? causal_count : key_len_after; + if (static_cast(block_index) >= visible_block_count) { + if (threadIdx.x == 0) { + block_scores[work] = SaiNegativeInfinity(); + } + continue; + } + + const int64_t key_base = (static_cast(batch) * params.state_capacity + block_index) * params.head_size; + float score = 0.0f; + for (int head = 0; head < params.num_heads; ++head) { + const float* query_head = + query_rotated + (static_cast(token) * params.num_heads + head) * params.head_size; + float partial = 0.0f; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + partial += query_head[d] * to_float(present_key_state[key_base + d]); + } + score += fmaxf(SaiBlockSum(partial, reduction), 0.0f); + } + if (threadIdx.x == 0) { + block_scores[work] = score * params.scale; + } + __syncthreads(); + } +} + +// One block per query token. Emits the token indices of the highest scoring blocks followed by the +// causally visible tokens of the trailing incomplete block, and the exact active count. A request +// whose update step was rejected for exceeding state_capacity this call (overflow_flags[batch] set) +// always gets the safe empty result: indices stay -1 (already reset below) and count is 0. +__global__ void QsaSelectKernel(const float* block_scores, const int32_t* cumulative_sequence_lengths, + const int32_t* past_sequence_lengths, const int64_t* position_ids, + const int32_t* present_state_lengths, const int32_t* overflow_flags, + int32_t* selected_indices, int32_t* selected_counts, + PackedSparseAttentionIndexerParams params) { + extern __shared__ float shared[]; + float* shared_value = shared; + int* shared_index = reinterpret_cast(shared + blockDim.x); + + for (int token = static_cast(blockIdx.x); token < params.total_tokens; token += static_cast(gridDim.x)) { + int32_t* out_row = selected_indices + static_cast(token) * params.capacity; + for (int p = static_cast(threadIdx.x); p < params.capacity; p += static_cast(blockDim.x)) { + out_row[p] = -1; + } + __syncthreads(); + + const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token); + if (overflow_flags[batch] != 0) { + if (threadIdx.x == 0) { + selected_counts[token] = 0; + } + __syncthreads(); + continue; + } + + const int key_len_after = present_state_lengths[batch * 2 + psai::kKeyStateLength]; + const int64_t abs_position = + params.has_position_ids + ? position_ids[token] + : static_cast(past_sequence_lengths[batch]) + (token - cumulative_sequence_lengths[batch]); + const int64_t causal_count = SaiCausalThreshold(abs_position, params.compress_ratio); + const int64_t visible_64 = causal_count < key_len_after ? causal_count : static_cast(key_len_after); + const int visible_block_count = static_cast(visible_64 < 0 ? 0 : visible_64); + const int selected = params.block_topk < visible_block_count ? params.block_topk : visible_block_count; + + const float* scores_row = block_scores + static_cast(token) * params.state_capacity; + + float previous_score = 0.0f; + int previous_index = -1; + int emitted_blocks = 0; + for (int rank = 0; rank < selected; ++rank) { + float best_value = 0.0f; + int best_index = -1; + SaiScanForNext(scores_row, visible_block_count, previous_score, previous_index, &best_value, &best_index); + shared_value[threadIdx.x] = best_value; + shared_index[threadIdx.x] = best_index; + __syncthreads(); + SaiBlockArgMax(shared_value, shared_index); + previous_index = shared_index[0]; + previous_score = shared_value[0]; + __syncthreads(); + if (previous_index < 0) { + break; + } + for (int t = static_cast(threadIdx.x); t < params.compress_ratio; t += static_cast(blockDim.x)) { + out_row[rank * params.compress_ratio + t] = previous_index * params.compress_ratio + t; + } + emitted_blocks = rank + 1; + __syncthreads(); + } + + // The trailing incomplete block is always causally visible in full up to this query's own + // position; only its indices are needed (SparsePagedAttention reads the raw main cache). + const int64_t block_start = static_cast(visible_block_count) * params.compress_ratio; + const int64_t natural_tail = abs_position >= block_start ? (abs_position - block_start + 1) : 0; + const int remaining_capacity = params.capacity - emitted_blocks * params.compress_ratio; + const int64_t tail_count_64 = natural_tail < remaining_capacity ? natural_tail : remaining_capacity; + const int tail_count = static_cast(tail_count_64 < 0 ? 0 : tail_count_64); + for (int t = static_cast(threadIdx.x); t < tail_count; t += static_cast(blockDim.x)) { + out_row[emitted_blocks * params.compress_ratio + t] = static_cast(block_start + t); + } + if (threadIdx.x == 0) { + selected_counts[token] = emitted_blocks * params.compress_ratio + tail_count; + } + __syncthreads(); + } +} + +// --------------------------------------------------------------------------------------------- +// policy_mode = "csa" +// --------------------------------------------------------------------------------------------- + +// One block per request: closes every new compression window (softmax-gated pool -> RMSNorm -> +// trailing RoPE -> append into fixed-capacity key_state) using the shared CsaWindowPlan helper, +// and publishes the raw overlap+leftover buffer. Reads only past_kv_buffer / past_gate_buffer / +// key / gate (never the present_* tensors), so it is correct whether or not present/past alias. +template +__global__ void CsaUpdateStateKernel(const T* key, const T* gate, const T* key_norm_weight, const T* cos_cache, + const T* sin_cache, const T* position_bias, + const int32_t* cumulative_sequence_lengths, const T* past_kv_buffer, + const T* past_gate_buffer, const int32_t* past_state_lengths, + T* present_key_state, T* present_kv_buffer, T* present_gate_buffer, + int32_t* present_state_lengths, int32_t* overflow_flags, + PackedSparseAttentionIndexerParams params) { + extern __shared__ float shared[]; + float* pooled = shared; + float* reduction = shared + params.head_size; + const int width = 2 * params.head_size; + + for (int b = static_cast(blockIdx.x); b < params.batch_size; b += static_cast(gridDim.x)) { + const int req_start = cumulative_sequence_lengths[b]; + const int req_end = cumulative_sequence_lengths[b + 1]; + const int req_len = req_end > req_start ? req_end - req_start : 0; + + const int old_key_len = min(max(past_state_lengths[b * 2 + psai::kKeyStateLength], 0), params.state_capacity); + const int old_buf_len = + min(max(past_state_lengths[b * 2 + psai::kBufferLength], 0), params.buffer_capacity); + + sai::CsaWindowPlan plan; + const bool plan_ok = sai::TryComputeCsaWindowPlan(old_buf_len, req_len, params.compress_ratio, plan); + // plan_ok is always true here: old_buf_len is clamped into [0, buffer_capacity) == + // [0, 2 * compress_ratio) and req_len >= 0, which are exactly the documented preconditions. + const int full_new_window_count = plan_ok ? static_cast(plan.new_window_count) : 0; + const int capacity_left_raw = params.state_capacity - old_key_len; + const int capacity_left = capacity_left_raw > 0 ? capacity_left_raw : 0; + // Reject (do not partially apply) a step that would need more than the fixed state_capacity: + // no new windows are closed and the buffer is left exactly as it was, so a rejected step is a + // deterministic no-op on state rather than a silent partial truncation. + const bool overflowed = full_new_window_count > capacity_left; + const int new_window_count = overflowed ? 0 : full_new_window_count; + const int present_buffer_length = + overflowed ? old_buf_len + : (static_cast(plan.present_buffer_length) < params.buffer_capacity + ? static_cast(plan.present_buffer_length) + : params.buffer_capacity); + const int present_buffer_start = overflowed ? 0 : static_cast(plan.present_buffer_start); + const int overlap_length = static_cast(plan.overlap_length); + + // Barrier: every thread has now read past_state_lengths / computed the plan (identically) + // before any thread below writes present_state_lengths, which keeps this correct even if the + // two tensors alias. + __syncthreads(); + + if (threadIdx.x == 0) { + present_state_lengths[b * 2 + psai::kKeyStateLength] = old_key_len + new_window_count; + present_state_lengths[b * 2 + psai::kBufferLength] = present_buffer_length; + overflow_flags[b] = overflowed ? 1 : 0; + } + + for (int k = 0; k < new_window_count; ++k) { + const bool has_previous = k >= 1 || overlap_length >= params.compress_ratio; + const int previous_base = overlap_length + (k - 1) * params.compress_ratio; + const int current_base = overlap_length + k * params.compress_ratio; + + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + float max_gate = SaiNegativeInfinity(); + if (has_previous) { + for (int slot = 0; slot < params.compress_ratio; ++slot) { + const int virtual_pos = previous_base + slot; + const float value = + (virtual_pos < old_buf_len + ? to_float(past_gate_buffer[(static_cast(b) * params.buffer_capacity + + virtual_pos) * + width + + d]) + : to_float(gate[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + + d])) + + to_float(position_bias[static_cast(slot) * width + d]); + max_gate = fmaxf(max_gate, value); + } + } + for (int slot = 0; slot < params.compress_ratio; ++slot) { + const int virtual_pos = current_base + slot; + const float value = + (virtual_pos < old_buf_len + ? to_float(past_gate_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) * + width + + params.head_size + d]) + : to_float(gate[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + + params.head_size + d])) + + to_float(position_bias[static_cast(slot) * width + params.head_size + d]); + max_gate = fmaxf(max_gate, value); + } + + float denominator = 0.0f; + float accumulator = 0.0f; + if (has_previous) { + for (int slot = 0; slot < params.compress_ratio; ++slot) { + const int virtual_pos = previous_base + slot; + const float logit = + (virtual_pos < old_buf_len + ? to_float(past_gate_buffer[(static_cast(b) * params.buffer_capacity + + virtual_pos) * + width + + d]) + : to_float(gate[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + + d])) + + to_float(position_bias[static_cast(slot) * width + d]); + const float weight = __expf(logit - max_gate); + denominator += weight; + accumulator += + weight * + (virtual_pos < old_buf_len + ? to_float(past_kv_buffer[(static_cast(b) * params.buffer_capacity + + virtual_pos) * + width + + d]) + : to_float(key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + + d])); + } + } + for (int slot = 0; slot < params.compress_ratio; ++slot) { + const int virtual_pos = current_base + slot; + const float logit = + (virtual_pos < old_buf_len + ? to_float(past_gate_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) * + width + + params.head_size + d]) + : to_float(gate[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + + params.head_size + d])) + + to_float(position_bias[static_cast(slot) * width + params.head_size + d]); + const float weight = __expf(logit - max_gate); + denominator += weight; + accumulator += + weight * (virtual_pos < old_buf_len + ? to_float(past_kv_buffer[(static_cast(b) * params.buffer_capacity + + virtual_pos) * + width + + params.head_size + d]) + : to_float(key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * + width + + params.head_size + d])); + } + pooled[d] = denominator > 0.0f ? accumulator / denominator : 0.0f; + } + __syncthreads(); + + float sum_squares = 0.0f; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + sum_squares += pooled[d] * pooled[d]; + } + sum_squares = SaiBlockSum(sum_squares, reduction); + const float inverse_rms = rsqrtf(sum_squares / static_cast(params.head_size) + params.epsilon); + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + pooled[d] = pooled[d] * inverse_rms * to_float(key_norm_weight[d]); + } + __syncthreads(); + + const int entry = old_key_len + k; + const int rope_position = + SaiClampPosition(static_cast(entry) * params.compress_ratio, params.max_rotary_length); + const int64_t cache_offset = + (static_cast(params.cos_cache_batched ? b : 0) * params.max_rotary_length + rope_position) * + params.rotary_width; + const int64_t out_base = (static_cast(b) * params.state_capacity + entry) * params.head_size; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + present_key_state[out_base + d] = from_float(SaiTrailingRope( + pooled, params.head_size, params.rotary_width, cos_cache + cache_offset, sin_cache + cache_offset, d)); + } + __syncthreads(); + } + + // Publish the raw overlap+leftover buffer. Skipped entirely on overflow: present_kv_buffer / + // present_gate_buffer already hold past_kv_buffer's / past_gate_buffer's contents unchanged + // (from the baseline copy in the Launch function below), which is exactly the prior valid + // buffer this rejected step must preserve. + if (!overflowed) { + for (int t = static_cast(threadIdx.x); t < present_buffer_length; t += static_cast(blockDim.x)) { + const int virtual_pos = present_buffer_start + t; + const int64_t out_base = (static_cast(b) * params.buffer_capacity + t) * width; + for (int c = 0; c < width; ++c) { + const float key_value = + virtual_pos < old_buf_len + ? to_float( + past_kv_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) * width + c]) + : to_float(key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + c]); + const float gate_value = + virtual_pos < old_buf_len + ? to_float(past_gate_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) * + width + + c]) + : to_float(gate[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + c]); + present_kv_buffer[out_base + c] = from_float(key_value); + present_gate_buffer[out_base + c] = from_float(gate_value); + } + } + } + __syncthreads(); + } +} + +// One block per (token, key_state slot). Also applies the causal mask so the selection kernel only +// has to read scores. +template +__global__ void CsaScoreKernel(const T* present_key_state, const float* query_rotated, const T* head_weights, + const int32_t* cumulative_sequence_lengths, const int64_t* position_ids, + const int32_t* present_state_lengths, float* scores, + PackedSparseAttentionIndexerParams params) { + extern __shared__ float reduction[]; + const int64_t total = static_cast(params.total_tokens) * params.state_capacity; + for (int64_t work = blockIdx.x; work < total; work += gridDim.x) { + const int token = static_cast(work / params.state_capacity); + const int entry = static_cast(work % params.state_capacity); + const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token); + const int key_len_after = present_state_lengths[batch * 2 + psai::kKeyStateLength]; + + if (entry >= key_len_after) { + if (threadIdx.x == 0) { + scores[work] = SaiNegativeInfinity(); + } + continue; + } + const int64_t threshold = SaiCausalThreshold(position_ids[token], params.compress_ratio); + if (static_cast(entry) >= threshold) { + if (threadIdx.x == 0) { + scores[work] = SaiNegativeInfinity(); + } + continue; + } + + const int64_t key_base = (static_cast(batch) * params.state_capacity + entry) * params.head_size; + float total_score = 0.0f; + for (int head = 0; head < params.num_heads; ++head) { + const float* query_head = + query_rotated + (static_cast(token) * params.num_heads + head) * params.head_size; + float partial = 0.0f; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + partial += query_head[d] * to_float(present_key_state[key_base + d]); + } + const float dot = SaiBlockSum(partial, reduction); + if (threadIdx.x == 0) { + total_score += fmaxf(dot, 0.0f) * to_float(head_weights[static_cast(token) * params.num_heads + + head]); + } + __syncthreads(); + } + if (threadIdx.x == 0) { + scores[work] = total_score * params.scale * params.head_weight_scale; + } + __syncthreads(); + } +} + +// One block per query token. Selects the index_topk highest scoring, causally-visible compressed +// entries and the exact active count. A request whose update step was rejected for exceeding +// state_capacity this call (overflow_flags[batch] set) always gets the safe empty result: indices +// stay -1 (already reset below) and count is 0. +__global__ void CsaSelectKernel(const float* scores, const int32_t* cumulative_sequence_lengths, + const int64_t* position_ids, const int32_t* present_state_lengths, + const int32_t* overflow_flags, int32_t* selected_indices, + int32_t* selected_counts, PackedSparseAttentionIndexerParams params) { + extern __shared__ float shared[]; + float* shared_value = shared; + int* shared_index = reinterpret_cast(shared + blockDim.x); + + for (int token = static_cast(blockIdx.x); token < params.total_tokens; token += static_cast(gridDim.x)) { + int32_t* out_row = selected_indices + static_cast(token) * params.capacity; + for (int p = static_cast(threadIdx.x); p < params.capacity; p += static_cast(blockDim.x)) { + out_row[p] = -1; + } + __syncthreads(); + + const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token); + if (overflow_flags[batch] != 0) { + if (threadIdx.x == 0) { + selected_counts[token] = 0; + } + __syncthreads(); + continue; + } + + const int key_len_after = present_state_lengths[batch * 2 + psai::kKeyStateLength]; + const int64_t threshold = SaiCausalThreshold(position_ids[token], params.compress_ratio); + const int64_t visible_64 = threshold < key_len_after ? threshold : static_cast(key_len_after); + const int visible = static_cast(visible_64 < 0 ? 0 : visible_64); + const int selected = params.index_topk < visible ? params.index_topk : visible; + + const float* scores_row = scores + static_cast(token) * params.state_capacity; + float previous_score = 0.0f; + int previous_index = -1; + int emitted = 0; + for (int rank = 0; rank < selected; ++rank) { + float best_value = 0.0f; + int best_index = -1; + SaiScanForNext(scores_row, visible, previous_score, previous_index, &best_value, &best_index); + shared_value[threadIdx.x] = best_value; + shared_index[threadIdx.x] = best_index; + __syncthreads(); + SaiBlockArgMax(shared_value, shared_index); + previous_index = shared_index[0]; + previous_score = shared_value[0]; + __syncthreads(); + if (previous_index < 0) { + break; + } + if (threadIdx.x == 0) { + out_row[rank] = previous_index; + } + emitted = rank + 1; + __syncthreads(); + } + if (threadIdx.x == 0) { + selected_counts[token] = emitted; + } + __syncthreads(); + } +} + +} // namespace + +size_t GetQsaPackedWorkspaceFloatCount(const PackedSparseAttentionIndexerParams& params) { + const size_t rows = static_cast(params.total_tokens); + return rows * params.num_heads * params.head_size + rows * static_cast(std::max(params.state_capacity, 1)); +} + +size_t GetCsaPackedWorkspaceFloatCount(const PackedSparseAttentionIndexerParams& params) { + const size_t rows = static_cast(params.total_tokens); + return rows * params.num_heads * params.head_size + rows * static_cast(std::max(params.state_capacity, 1)); +} + +template +Status LaunchQsaPackedSparseAttentionIndexer( + cudaStream_t stream, const PackedSparseAttentionIndexerParams& params, const T* query, const T* key, + const T* key_norm_weight, const T* cos_cache, const T* sin_cache, const int32_t* cumulative_sequence_lengths, + const int32_t* past_sequence_lengths, const int64_t* position_ids, const T* past_key_state, + const T* past_kv_buffer, const int32_t* past_state_lengths, int32_t* selected_indices, + int32_t* selected_counts, T* present_key_state, T* present_kv_buffer, int32_t* present_state_lengths, + float* float_workspace, int32_t* overflow_flags) { + if (params.batch_size > 0) { + const int64_t key_state_elems = + static_cast(params.batch_size) * params.state_capacity * params.head_size; + if (key_state_elems > 0 && present_key_state != past_key_state) { + ElementwiseCopyKernel<<>>( + past_key_state, present_key_state, key_state_elems); + } + const int64_t buffer_elems = static_cast(params.batch_size) * params.buffer_capacity * params.head_size; + if (buffer_elems > 0 && present_kv_buffer != past_kv_buffer) { + ElementwiseCopyKernel<<>>( + past_kv_buffer, present_kv_buffer, buffer_elems); + } + if (present_state_lengths != past_state_lengths) { + const int64_t length_elems = static_cast(params.batch_size) * psai::kStateLengthColumns; + ElementwiseCopyKernel<<>>( + past_state_lengths, present_state_lengths, length_elems); + } + } + + if (params.batch_size == 0) { + return CUDA_CALL(cudaGetLastError()); + } + + const size_t value_bytes = static_cast(params.head_size) * sizeof(float); + const int state_blocks = static_cast(std::min(params.batch_size, kSaiMaxGridDimX)); + QsaUpdateStateKernel<<>>( + key, key_norm_weight, cos_cache, sin_cache, cumulative_sequence_lengths, past_kv_buffer, past_state_lengths, + present_key_state, present_kv_buffer, present_state_lengths, overflow_flags, params); + + if (params.total_tokens == 0) { + return CUDA_CALL(cudaGetLastError()); + } + + float* query_rotated = float_workspace; + float* block_scores = + float_workspace + static_cast(params.total_tokens) * params.num_heads * params.head_size; + + const int64_t rotate_rows = static_cast(params.total_tokens) * params.num_heads; + const int rotate_blocks = static_cast(std::min(rotate_rows, kSaiMaxGridDimX)); + PackedRotateQueryKernel<<>>( + query, cos_cache, sin_cache, cumulative_sequence_lengths, past_sequence_lengths, position_ids, query_rotated, + params); + + if (params.state_capacity > 0) { + const int64_t score_work = static_cast(params.total_tokens) * params.state_capacity; + const int score_blocks = static_cast(std::min(score_work, kSaiMaxGridDimX)); + QsaBlockScoreKernel<<>>( + present_key_state, query_rotated, cumulative_sequence_lengths, past_sequence_lengths, position_ids, + present_state_lengths, block_scores, params); + } + + const int token_blocks = static_cast(std::min(params.total_tokens, kSaiMaxGridDimX)); + QsaSelectKernel<<>>( + block_scores, cumulative_sequence_lengths, past_sequence_lengths, position_ids, present_state_lengths, + overflow_flags, selected_indices, selected_counts, params); + + return CUDA_CALL(cudaGetLastError()); +} + +template +Status LaunchCsaPackedSparseAttentionIndexer( + cudaStream_t stream, const PackedSparseAttentionIndexerParams& params, const T* query, const T* key, + const T* key_norm_weight, const T* cos_cache, const T* sin_cache, const T* gate, const T* position_bias, + const T* head_weights, const int32_t* cumulative_sequence_lengths, const int32_t* past_sequence_lengths, + const int64_t* position_ids, const T* past_key_state, const T* past_kv_buffer, const T* past_gate_buffer, + const int32_t* past_state_lengths, int32_t* selected_indices, int32_t* selected_counts, T* present_key_state, + T* present_kv_buffer, T* present_gate_buffer, int32_t* present_state_lengths, float* float_workspace, + int32_t* overflow_flags) { + if (params.batch_size > 0) { + const int64_t key_state_elems = + static_cast(params.batch_size) * params.state_capacity * params.head_size; + if (key_state_elems > 0 && present_key_state != past_key_state) { + ElementwiseCopyKernel<<>>( + past_key_state, present_key_state, key_state_elems); + } + const int64_t buffer_elems = + static_cast(params.batch_size) * params.buffer_capacity * 2 * params.head_size; + if (buffer_elems > 0) { + if (present_kv_buffer != past_kv_buffer) { + ElementwiseCopyKernel<<>>( + past_kv_buffer, present_kv_buffer, buffer_elems); + } + if (present_gate_buffer != past_gate_buffer) { + ElementwiseCopyKernel<<>>( + past_gate_buffer, present_gate_buffer, buffer_elems); + } + } + if (present_state_lengths != past_state_lengths) { + const int64_t length_elems = static_cast(params.batch_size) * psai::kStateLengthColumns; + ElementwiseCopyKernel<<>>( + past_state_lengths, present_state_lengths, length_elems); + } + } + + if (params.batch_size == 0) { + return CUDA_CALL(cudaGetLastError()); + } + + const size_t value_bytes = static_cast(params.head_size) * sizeof(float); + const int state_blocks = static_cast(std::min(params.batch_size, kSaiMaxGridDimX)); + CsaUpdateStateKernel<<>>( + key, gate, key_norm_weight, cos_cache, sin_cache, position_bias, cumulative_sequence_lengths, past_kv_buffer, + past_gate_buffer, past_state_lengths, present_key_state, present_kv_buffer, present_gate_buffer, + present_state_lengths, overflow_flags, params); + + if (params.total_tokens == 0) { + return CUDA_CALL(cudaGetLastError()); + } + + float* query_rotated = float_workspace; + float* scores = float_workspace + static_cast(params.total_tokens) * params.num_heads * params.head_size; + + const int64_t rotate_rows = static_cast(params.total_tokens) * params.num_heads; + const int rotate_blocks = static_cast(std::min(rotate_rows, kSaiMaxGridDimX)); + PackedRotateQueryKernel<<>>( + query, cos_cache, sin_cache, cumulative_sequence_lengths, past_sequence_lengths, position_ids, query_rotated, + params); + + if (params.state_capacity > 0) { + const int64_t score_work = static_cast(params.total_tokens) * params.state_capacity; + const int score_blocks = static_cast(std::min(score_work, kSaiMaxGridDimX)); + CsaScoreKernel<<>>( + present_key_state, query_rotated, head_weights, cumulative_sequence_lengths, position_ids, + present_state_lengths, scores, params); + } + + const int token_blocks = static_cast(std::min(params.total_tokens, kSaiMaxGridDimX)); + CsaSelectKernel<<>>( + scores, cumulative_sequence_lengths, position_ids, present_state_lengths, overflow_flags, selected_indices, + selected_counts, params); + + return CUDA_CALL(cudaGetLastError()); +} + +#define INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(T) \ + template Status LaunchQsaPackedSparseAttentionIndexer( \ + cudaStream_t, const PackedSparseAttentionIndexerParams&, const T*, const T*, const T*, const T*, \ + const T*, const int32_t*, const int32_t*, const int64_t*, const T*, const T*, const int32_t*, int32_t*, \ + int32_t*, T*, T*, int32_t*, float*, int32_t*); \ + template Status LaunchCsaPackedSparseAttentionIndexer( \ + cudaStream_t, const PackedSparseAttentionIndexerParams&, const T*, const T*, const T*, const T*, \ + const T*, const T*, const T*, const T*, const int32_t*, const int32_t*, const int64_t*, const T*, \ + const T*, const T*, const int32_t*, int32_t*, int32_t*, T*, T*, T*, int32_t*, float*); + +INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(float) +INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(half) +INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(__nv_bfloat16) + +#undef INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER + +} // namespace cuda +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h new file mode 100644 index 0000000000000..2eae2bdc4465a --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h @@ -0,0 +1,106 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include + +#include "core/providers/cuda/cuda_common.h" + +namespace onnxruntime { +namespace contrib { +namespace cuda { + +// Everything the device code needs to know about a PackedSparseAttentionIndexer call. All of it is +// derived from attributes and input *shapes* (never input values), so no device data is ever read +// on the host; per-request quantities (past_sequence_lengths, past_state_lengths, cumulative +// offsets) are read directly by the kernels below from device memory. +struct PackedSparseAttentionIndexerParams { + int batch_size = 0; + int total_tokens = 0; + int num_heads = 0; + int head_size = 0; + int rotary_width = 0; // cos_cache.shape[-1] + int max_rotary_length = 0; // cos_cache.shape[-2] + bool cos_cache_batched = false; // cos_cache rank: 3 = [batch, pos, rot], 2 = [pos, rot] + int compress_ratio = 0; + int state_capacity = 0; // past_key_state.shape[1] + int buffer_capacity = 0; // past_kv_buffer.shape[1] == 2 * compress_ratio - 1 + int capacity = 0; // selected_indices.shape[1] + bool has_position_ids = false; + float epsilon = 1e-6f; + float scale = 0.0f; + + // policy_mode = "qsa" + int block_topk = 0; // token_budget / compress_ratio + + // policy_mode = "csa" + int index_topk = 0; + float head_weight_scale = 0.0f; +}; + +// Scratch requirements, in float elements. +size_t GetQsaPackedWorkspaceFloatCount(const PackedSparseAttentionIndexerParams& params); +size_t GetCsaPackedWorkspaceFloatCount(const PackedSparseAttentionIndexerParams& params); + +// `overflow_flags` is a caller-allocated int32 scratch buffer with at least `batch_size` elements +// (unused when batch_size == 0). The update kernel writes, per request, whether this call's new +// blocks/windows would exceed state_capacity; when it does, the whole step is rejected for that +// request (present_state_lengths / present_key_state / present_kv_buffer / present_gate_buffer for +// that request are left exactly as their past_* counterparts) and the select kernels force that +// request's selected_indices/selected_counts to the safe empty result (-1 / 0) rather than +// selecting against a partially updated state. +template +Status LaunchQsaPackedSparseAttentionIndexer( + cudaStream_t stream, + const PackedSparseAttentionIndexerParams& params, + const T* query, + const T* key, + const T* key_norm_weight, + const T* cos_cache, + const T* sin_cache, + const int32_t* cumulative_sequence_lengths, + const int32_t* past_sequence_lengths, + const int64_t* position_ids, + const T* past_key_state, + const T* past_kv_buffer, + const int32_t* past_state_lengths, + int32_t* selected_indices, + int32_t* selected_counts, + T* present_key_state, + T* present_kv_buffer, + int32_t* present_state_lengths, + float* float_workspace, + int32_t* overflow_flags); + +template +Status LaunchCsaPackedSparseAttentionIndexer( + cudaStream_t stream, + const PackedSparseAttentionIndexerParams& params, + const T* query, + const T* key, + const T* key_norm_weight, + const T* cos_cache, + const T* sin_cache, + const T* gate, + const T* position_bias, + const T* head_weights, + const int32_t* cumulative_sequence_lengths, + const int32_t* past_sequence_lengths, + const int64_t* position_ids, + const T* past_key_state, + const T* past_kv_buffer, + const T* past_gate_buffer, + const int32_t* past_state_lengths, + int32_t* selected_indices, + int32_t* selected_counts, + T* present_key_state, + T* present_kv_buffer, + T* present_gate_buffer, + int32_t* present_state_lengths, + float* float_workspace, + int32_t* overflow_flags); + +} // namespace cuda +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh b/onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh new file mode 100644 index 0000000000000..de7c2d3b323ad --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh @@ -0,0 +1,149 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +// +// Device-side math shared by com.microsoft.SparseAttentionIndexer (dense, growing state) and +// com.microsoft.PackedSparseAttentionIndexer (packed, fixed-capacity state). Everything here is a +// small, self-contained helper with no dependency on either operator's tensor layout, so it is +// included (not linked) by both .cu translation units; the anonymous namespace gives each +// translation unit its own private copy, which is the normal, ODR-safe pattern for header-only +// CUDA device helpers. + +#pragma once + +#include +#include + +#include +#include + +#include "core/providers/cuda/cu_inc/cuda_type_helper.cuh" + +namespace onnxruntime { +namespace contrib { +namespace cuda { + +namespace { + +constexpr int64_t kSaiMaxGridDimX = 2147483647; + +__device__ __forceinline__ float SaiNegativeInfinity() { return -CUDART_INF_F; } + +inline int SaiGridForElements(int64_t count, int threads) { + const int64_t blocks = (count + threads - 1) / threads; + return static_cast(std::clamp(blocks, 1, 65535)); +} + +// --------------------------------------------------------------------------------------------- +// FP32 block reductions +// --------------------------------------------------------------------------------------------- + +__device__ __forceinline__ float SaiBlockSum(float value, float* shared) { + shared[threadIdx.x] = value; + __syncthreads(); + for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) { + if (threadIdx.x < stride) { + shared[threadIdx.x] += shared[threadIdx.x + stride]; + } + __syncthreads(); + } + const float total = shared[0]; + __syncthreads(); + return total; +} + +// Reduces (value, index) pairs to the largest value, breaking ties towards the smaller index. +// A negative index marks an empty slot. shared_value/shared_index must already be filled and synced. +__device__ __forceinline__ void SaiBlockArgMax(float* shared_value, int* shared_index) { + for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) { + if (threadIdx.x < stride) { + const int other_index = shared_index[threadIdx.x + stride]; + if (other_index >= 0) { + const int this_index = shared_index[threadIdx.x]; + const float other_value = shared_value[threadIdx.x + stride]; + const float this_value = shared_value[threadIdx.x]; + if (this_index < 0 || other_value > this_value || + (other_value == this_value && other_index < this_index)) { + shared_value[threadIdx.x] = other_value; + shared_index[threadIdx.x] = other_index; + } + } + } + __syncthreads(); + } +} + +// Per-thread scan for the best entry that comes strictly after (previous_score, previous_index) in +// the total order "score descending, then index ascending". Entries already emitted are therefore +// skipped without needing a visited bitmap. +__device__ __forceinline__ void SaiScanForNext(const float* scores, int count, float previous_score, + int previous_index, float* best_value, int* best_index) { + *best_index = -1; + *best_value = 0.0f; + for (int candidate = static_cast(threadIdx.x); candidate < count; + candidate += static_cast(blockDim.x)) { + const float value = scores[candidate]; + if (previous_index >= 0 && + !(value < previous_score || (value == previous_score && candidate > previous_index))) { + continue; + } + if (*best_index < 0 || value > *best_value || + (value == *best_value && candidate < *best_index)) { + *best_value = value; + *best_index = candidate; + } + } +} + +// --------------------------------------------------------------------------------------------- +// Rotary embeddings +// --------------------------------------------------------------------------------------------- + +// Split-half rotary over the leading `rotary_width` channels (the convention used by the qsa +// reference). Channels beyond `rotary_width` pass through unchanged. +template +__device__ __forceinline__ float SaiLeadingRope(const float* value, int rotary_width, const T* cos_row, + const T* sin_row, int d) { + if (d >= rotary_width) { + return value[d]; + } + const int half = rotary_width / 2; + const float paired = (d < half) ? -value[d + half] : value[d - half]; + return value[d] * to_float(cos_row[d]) + paired * to_float(sin_row[d]); +} + +// Interleaved rotary over the trailing 2 * rotary_width channels (the convention used by the csa +// reference). Each cos/sin entry covers one channel pair, matching repeat_interleave(2). +template +__device__ __forceinline__ float SaiTrailingRope(const float* value, int head_size, int rotary_width, + const T* cos_row, const T* sin_row, int d) { + const int base = head_size - 2 * rotary_width; + if (d < base) { + return value[d]; + } + const int offset = d - base; + const float paired = ((offset & 1) == 0) ? -value[d + 1] : value[d - 1]; + return value[d] * to_float(cos_row[offset >> 1]) + paired * to_float(sin_row[offset >> 1]); +} + +// --------------------------------------------------------------------------------------------- +// Causal geometry +// --------------------------------------------------------------------------------------------- + +// Highest compressed entry a query at `position` may attend to, matching (position + 1) // ratio. +__device__ __forceinline__ int64_t SaiCausalThreshold(int64_t position, int compress_ratio) { + return position < 0 ? 0 : position / compress_ratio + (position % compress_ratio == compress_ratio - 1); +} + +__device__ __forceinline__ int SaiClampPosition(int64_t position, int max_rotary_length) { + if (position < 0) { + return 0; + } + const int64_t limit = max_rotary_length - 1; + return static_cast(position < limit ? position : limit); +} + +} // namespace + +} // namespace cuda +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_impl.cu b/onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_impl.cu index ffc453754231b..40c07c1238f7a 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_impl.cu +++ b/onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_impl.cu @@ -14,6 +14,7 @@ #include +#include "contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh" #include "core/providers/cuda/cu_inc/cuda_type_helper.cuh" namespace onnxruntime { @@ -24,114 +25,43 @@ namespace { // The block reductions below halve the active thread count, so this must stay a power of two. constexpr int kThreads = 128; -constexpr int64_t kMaxGridDimX = 2147483647; +constexpr int64_t kMaxGridDimX = kSaiMaxGridDimX; -__device__ __forceinline__ float NegativeInfinity() { return -CUDART_INF_F; } +// Thin, same-signature aliases over the device math shared with the packed indexer implementation +// (sparse_attention_indexer_device_math.cuh), so the kernels below are unchanged. +__device__ __forceinline__ float NegativeInfinity() { return SaiNegativeInfinity(); } -int GridForElements(int64_t count) { - const int64_t blocks = (count + kThreads - 1) / kThreads; - return static_cast(std::clamp(blocks, 1, 65535)); -} +int GridForElements(int64_t count) { return SaiGridForElements(count, kThreads); } -// --------------------------------------------------------------------------------------------- -// Shared device helpers -// --------------------------------------------------------------------------------------------- - -__device__ __forceinline__ float BlockSum(float value, float* shared) { - shared[threadIdx.x] = value; - __syncthreads(); - for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) { - if (threadIdx.x < stride) { - shared[threadIdx.x] += shared[threadIdx.x + stride]; - } - __syncthreads(); - } - const float total = shared[0]; - __syncthreads(); - return total; -} +__device__ __forceinline__ float BlockSum(float value, float* shared) { return SaiBlockSum(value, shared); } -// Reduces (value, index) pairs to the largest value, breaking ties towards the smaller index. -// A negative index marks an empty slot. shared_value/shared_index must already be filled and synced. __device__ __forceinline__ void BlockArgMax(float* shared_value, int* shared_index) { - for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) { - if (threadIdx.x < stride) { - const int other_index = shared_index[threadIdx.x + stride]; - if (other_index >= 0) { - const int this_index = shared_index[threadIdx.x]; - const float other_value = shared_value[threadIdx.x + stride]; - const float this_value = shared_value[threadIdx.x]; - if (this_index < 0 || other_value > this_value || - (other_value == this_value && other_index < this_index)) { - shared_value[threadIdx.x] = other_value; - shared_index[threadIdx.x] = other_index; - } - } - } - __syncthreads(); - } + SaiBlockArgMax(shared_value, shared_index); } -// Per-thread scan for the best entry that comes strictly after (previous_score, previous_index) in -// the total order "score descending, then index ascending". Entries already emitted are therefore -// skipped without needing a visited bitmap. __device__ __forceinline__ void ScanForNext(const float* scores, int count, float previous_score, int previous_index, float* best_value, int* best_index) { - *best_index = -1; - *best_value = 0.0f; - for (int candidate = static_cast(threadIdx.x); candidate < count; - candidate += static_cast(blockDim.x)) { - const float value = scores[candidate]; - if (previous_index >= 0 && - !(value < previous_score || (value == previous_score && candidate > previous_index))) { - continue; - } - if (*best_index < 0 || value > *best_value || - (value == *best_value && candidate < *best_index)) { - *best_value = value; - *best_index = candidate; - } - } + SaiScanForNext(scores, count, previous_score, previous_index, best_value, best_index); } -// Split-half rotary over the leading `rotary_width` channels (the convention used by the qsa -// reference). Channels beyond `rotary_width` pass through unchanged. template __device__ __forceinline__ float LeadingRope(const float* value, int rotary_width, const T* cos_row, const T* sin_row, int d) { - if (d >= rotary_width) { - return value[d]; - } - const int half = rotary_width / 2; - const float paired = (d < half) ? -value[d + half] : value[d - half]; - return value[d] * to_float(cos_row[d]) + paired * to_float(sin_row[d]); + return SaiLeadingRope(value, rotary_width, cos_row, sin_row, d); } -// Interleaved rotary over the trailing 2 * rotary_width channels (the convention used by the csa -// reference). Each cos/sin entry covers one channel pair, matching repeat_interleave(2). template __device__ __forceinline__ float TrailingRope(const float* value, int head_size, int rotary_width, const T* cos_row, const T* sin_row, int d) { - const int base = head_size - 2 * rotary_width; - if (d < base) { - return value[d]; - } - const int offset = d - base; - const float paired = ((offset & 1) == 0) ? -value[d + 1] : value[d - 1]; - return value[d] * to_float(cos_row[offset >> 1]) + paired * to_float(sin_row[offset >> 1]); + return SaiTrailingRope(value, head_size, rotary_width, cos_row, sin_row, d); } -// Highest compressed entry a query at `position` may attend to, matching (position + 1) // ratio. __device__ __forceinline__ int64_t CausalThreshold(int64_t position, int compress_ratio) { - return position < 0 ? 0 : position / compress_ratio + (position % compress_ratio == compress_ratio - 1); + return SaiCausalThreshold(position, compress_ratio); } __device__ __forceinline__ int ClampPosition(int64_t position, int max_rotary_length) { - if (position < 0) { - return 0; - } - const int64_t limit = max_rotary_length - 1; - return static_cast(position < limit ? position : limit); + return SaiClampPosition(position, max_rotary_length); } // --------------------------------------------------------------------------------------------- diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc new file mode 100644 index 0000000000000..e12d4ddbb5c1d --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -0,0 +1,982 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. +// +// WebGPU implementation of com.microsoft.PackedSparseAttentionIndexer. Like its dense +// SparseAttentionIndexer counterpart, every program below is single-invocation-per-row +// (`if (... || local_idx != 0u) { return; }`) and recomputes reductions on demand rather than +// staging per-row intermediates in workgroup memory; this keeps every kernel correct without +// requiring WGSL arrays sized by a runtime (uniform) head_size, which WGSL does not support for +// function-local variables. See docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md and +// packed_sparse_attention_indexer_impl.cu (the CUDA implementation) for the full operator +// contract and the device-side metadata-safety argument, which apply unchanged here: every +// per-request quantity is read directly from device buffers inside the shader (never on the +// host), and values are clamped into the fixed-capacity range before use so malformed metadata +// can never cause an out-of-bounds access. + +#include "contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h" + +#include +#include +#include + +#include "contrib_ops/webgpu/webgpu_contrib_kernels.h" +#include "core/providers/webgpu/shader_helper.h" +#include "core/providers/webgpu/webgpu_supported_types.h" + +namespace onnxruntime { +namespace contrib { +namespace webgpu { + +namespace psai = onnxruntime::contrib::packed_sparse_attention_indexer; + +ONNX_OPERATOR_KERNEL_EX( + PackedSparseAttentionIndexer, + kMSDomain, + 1, + kWebGpuExecutionProvider, + (*KernelDefBuilder::Create()) + .TypeConstraint("T", WebGpuSupportedFloatTypes()) + .TypeConstraint("I", DataTypeImpl::GetTensorType()) + .TypeConstraint("M", DataTypeImpl::GetTensorType()), + PackedSparseAttentionIndexer); + +namespace { + +constexpr uint32_t kWorkgroupSize = 64; + +Status CheckShape(const Tensor* tensor, const char* name, std::initializer_list expected) { + ORT_RETURN_IF(tensor == nullptr, "PackedSparseAttentionIndexer: ", name, " is required"); + const TensorShape expected_shape(expected); + ORT_RETURN_IF_NOT(tensor->Shape() == expected_shape, "PackedSparseAttentionIndexer: ", name, " must have shape ", + expected_shape.ToString(), ", got ", tensor->Shape().ToString()); + return Status::OK(); +} + +uint32_t ToUint32(int64_t value) { return onnxruntime::narrow(value); } + +struct RotaryCacheShape { + bool batched; + int64_t max_rotary_length; + int64_t rotary_width; +}; + +Status CheckRotaryCache(const Tensor* cos_cache, const Tensor* sin_cache, int64_t batch_size, + RotaryCacheShape& out) { + ORT_RETURN_IF(cos_cache == nullptr, "PackedSparseAttentionIndexer: cos_cache is required"); + const auto& cos_shape = cos_cache->Shape(); + out.batched = cos_shape.NumDimensions() == 3; + ORT_RETURN_IF_NOT( + (out.batched && cos_shape[0] == batch_size && cos_shape[1] > 0) || + (cos_shape.NumDimensions() == 2 && cos_shape[0] > 0), + "PackedSparseAttentionIndexer: cos_cache must have shape (max_position, rotary_width) or " + "(batch_size, max_position, rotary_width), got ", + cos_shape.ToString()); + out.max_rotary_length = out.batched ? cos_shape[1] : cos_shape[0]; + out.rotary_width = out.batched ? cos_shape[2] : cos_shape[1]; + ORT_RETURN_IF_NOT(out.max_rotary_length > 0 && out.rotary_width > 0, + "PackedSparseAttentionIndexer: invalid cos_cache shape ", cos_shape.ToString()); + ORT_RETURN_IF_NOT(sin_cache != nullptr && sin_cache->Shape() == cos_shape, + "PackedSparseAttentionIndexer: sin_cache must have the same shape as cos_cache"); + return Status::OK(); +} + +} // namespace + +Status PackedSparseAttentionIndexerCopyProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& src = shader.AddInput("src", ShaderUsage::UseUniform); + const auto& dst = shader.AddOutput("dst", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + shader.MainFunctionBody() + << shader.GuardAgainstOutOfBoundsWorkgroupSizes("uniforms.total") + << " " << dst.SetByOffset("global_idx", "dst_element_t(" + src.GetByOffset("global_idx") + ")") << "\n"; + return Status::OK(); +} + +Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& key = shader.AddInput("key", ShaderUsage::UseUniform); + const auto& norm = shader.AddInput("key_norm_weight", ShaderUsage::UseUniform); + const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); + const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); + const auto& past_kv_buffer = shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); + const auto& past_state_lengths = shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); + const auto& present_key_state = + shader.AddOutput("present_key_state", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& present_kv_buffer = + shader.AddOutput("present_kv_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& present_state_lengths = shader.AddOutput("present_state_lengths", ShaderUsage::UseUniform); + const auto& overflow_flags = shader.AddOutput("overflow_flags", ShaderUsage::UseUniform); + + shader.AdditionalImplementation() + << "fn extended_key(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, d: u32) -> f32 {\n" + << " if (virtual_pos < old_buf_len) {\n" + << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * uniforms.head_size + d;\n" + << " return f32(" << past_kv_buffer.GetByOffset("idx") << ");\n" + << " }\n" + << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * uniforms.head_size + d;\n" + << " return f32(" << key.GetByOffset("idx2") << ");\n" + << "}\n" + << "fn pooled(b: u32, k: i32, old_buf_len: i32, req_start: i32, d: u32) -> f32 {\n" + << " var sum = 0.0;\n" + << " for (var t = 0u; t < uniforms.compress_ratio; t++) {\n" + << " sum += extended_key(b, k * i32(uniforms.compress_ratio) + i32(t), old_buf_len, req_start, d);\n" + << " }\n" + << " return sum / f32(uniforms.compress_ratio);\n" + << "}\n" + << "fn normalized(b: u32, k: i32, old_buf_len: i32, req_start: i32, d: u32) -> f32 {\n" + << " var square_sum = 0.0;\n" + << " for (var c = 0u; c < uniforms.head_size; c++) {\n" + << " let value = pooled(b, k, old_buf_len, req_start, c);\n" + << " square_sum += value * value;\n" + << " }\n" + << " return pooled(b, k, old_buf_len, req_start, d) * " + "inverseSqrt(square_sum / f32(uniforms.head_size) + uniforms.epsilon) * f32(" + << norm.GetByOffset("d") << ");\n" + << "}\n" + << "fn clamp_position(position: i32) -> u32 {\n" + << " if (position < 0) { return 0u; }\n" + << " return min(u32(position), uniforms.max_rotary_length - 1u);\n" + << "}\n" + << "fn rotated(b: u32, k: i32, old_buf_len: i32, req_start: i32, old_key_len: i32, d: u32) -> f32 {\n" + << " let value = normalized(b, k, old_buf_len, req_start, d);\n" + << " if (d >= uniforms.rotary_width) { return value; }\n" + << " let half = uniforms.rotary_width / 2u;\n" + << " let pair_d = select(d - half, d + half, d < half);\n" + << " let sign = select(1.0, -1.0, d < half);\n" + << " let paired = sign * normalized(b, k, old_buf_len, req_start, pair_d);\n" + << " let entry = old_key_len + k;\n" + << " let position = clamp_position(entry * i32(uniforms.compress_ratio));\n"; + if (cos_cache_batched_) { + shader.AdditionalImplementation() + << " let cache = (b * uniforms.max_rotary_length + position) * uniforms.rotary_width + d;\n"; + } else { + shader.AdditionalImplementation() << " let cache = position * uniforms.rotary_width + d;\n"; + } + shader.AdditionalImplementation() + << " return value * f32(" << cos_cache.GetByOffset("cache") << ") + paired * f32(" + << sin_cache.GetByOffset("cache") << ");\n" + << "}\n"; + + shader.MainFunctionBody() + << " let b = workgroup_idx;\n" + << " if (b >= uniforms.batch_size || local_idx != 0u) { return; }\n" + << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" + << " let req_end = " << cu_seqlens.GetByOffset("b + 1u") << ";\n" + << " let req_len = max(req_end - req_start, 0);\n" + << " let old_key_len = clamp(" << past_state_lengths.GetByOffset("b * 2u") + << ", 0, i32(uniforms.state_capacity));\n" + << " let old_buf_len = clamp(" << past_state_lengths.GetByOffset("b * 2u + 1u") + << ", 0, i32(uniforms.compress_ratio) - 1);\n" + << " let pending = old_buf_len + req_len;\n" + << " let full_new_block_count = pending / i32(uniforms.compress_ratio);\n" + << " let capacity_left = max(i32(uniforms.state_capacity) - old_key_len, 0);\n" + << " let new_block_count = min(full_new_block_count, capacity_left);\n" + << " let overflowed = new_block_count < full_new_block_count;\n" + << " let new_buf_len = select(pending % i32(uniforms.compress_ratio), 0, overflowed);\n" + << " " << present_state_lengths.SetByOffset("b * 2u", "old_key_len + new_block_count") << "\n" + << " " << present_state_lengths.SetByOffset("b * 2u + 1u", "new_buf_len") << "\n" + << " for (var k = 0; k < new_block_count; k++) {\n" + << " let entry = u32(old_key_len + k);\n" + << " for (var d = 0u; d < uniforms.head_size; d++) {\n" + << " let value = rotated(b, k, old_buf_len, req_start, old_key_len, d);\n" + << " " + << present_key_state.SetByOffset("(b * uniforms.state_capacity + entry) * uniforms.head_size + d", + "present_key_state_element_t(value)") + << "\n" + << " }\n" + << " }\n" + << " for (var t = 0; t < new_buf_len; t++) {\n" + << " let virtual_pos = new_block_count * i32(uniforms.compress_ratio) + t;\n" + << " for (var d = 0u; d < uniforms.head_size; d++) {\n" + << " let value = extended_key(b, virtual_pos, old_buf_len, req_start, d);\n" + << " " + << present_kv_buffer.SetByOffset("(b * uniforms.buffer_capacity + u32(t)) * uniforms.head_size + d", + "present_kv_buffer_element_t(value)") + << "\n" + << " }\n" + << " }\n"; + return Status::OK(); +} + +Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& query = shader.AddInput("query", ShaderUsage::UseUniform); + const auto& present_key_state = shader.AddInput("present_key_state", ShaderUsage::UseUniform); + const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); + const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); + const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); + const ShaderVariableHelper* position_ids = nullptr; + if (has_position_ids_) { + position_ids = &shader.AddInput("position_ids", ShaderUsage::UseUniform); + } + const auto& present_state_lengths = shader.AddInput("present_state_lengths", ShaderUsage::UseUniform); + const auto& selected_indices = shader.AddOutput("selected_indices", ShaderUsage::UseUniform); + const auto& selected_counts = shader.AddOutput("selected_counts", ShaderUsage::UseUniform); + + shader.AdditionalImplementation() + << "fn batch_of_token(token: u32) -> u32 {\n" + << " var b = 0u;\n" + << " for (var i = 0u; i < uniforms.batch_size; i++) {\n" + << " if (u32(" << cu_seqlens.GetByOffset("i") << ") <= token) { b = i; } else { break; }\n" + << " }\n" + << " return b;\n" + << "}\n" + << "fn clamp_position(position: i32) -> u32 {\n" + << " if (position < 0) { return 0u; }\n" + << " return min(u32(position), uniforms.max_rotary_length - 1u);\n" + << "}\n" + << "fn causal_count(position: i32) -> u32 {\n" + << " if (position < 0) { return 0u; }\n" + << " let p = u32(position);\n" + << " let cr = uniforms.compress_ratio;\n" + << " return p / cr + select(0u, 1u, p % cr == cr - 1u);\n" + << "}\n"; + if (has_position_ids_) { + shader.AdditionalImplementation() + << "fn abs_position(token: u32, b: u32) -> i32 {\n" + << " let raw = " << position_ids->GetByOffset("token", true) << ";\n" + << " if ((raw.y & 0x80000000u) != 0u) { return -2147483648; }\n" + << " if (raw.y != 0u || raw.x > 2147483647u) { return 2147483647; }\n" + << " return i32(raw.x);\n" + << "}\n"; + } else { + shader.AdditionalImplementation() + << "fn abs_position(token: u32, b: u32) -> i32 {\n" + << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" + << " return " << past_seqlens.GetByOffset("b") << " + (i32(token) - req_start);\n" + << "}\n"; + } + shader.AdditionalImplementation() + << "fn query_value(token: u32, head: u32, d: u32, position: i32, b: u32) -> f32 {\n" + << " let base = (token * uniforms.num_heads + head) * uniforms.head_size;\n" + << " var value = f32(" << query.GetByOffset("base + d") << ");\n" + << " if (d >= uniforms.rotary_width) { return value; }\n" + << " let half = uniforms.rotary_width / 2u;\n" + << " let pair_d = select(d - half, d + half, d < half);\n" + << " let sign = select(1.0, -1.0, d < half);\n" + << " let paired = sign * f32(" << query.GetByOffset("base + pair_d") << ");\n" + << " let position_clamped = clamp_position(position);\n"; + if (cos_cache_batched_) { + shader.AdditionalImplementation() + << " let cache = (b * uniforms.max_rotary_length + position_clamped) * uniforms.rotary_width + d;\n"; + } else { + shader.AdditionalImplementation() << " let cache = position_clamped * uniforms.rotary_width + d;\n"; + } + shader.AdditionalImplementation() + << " value = value * f32(" << cos_cache.GetByOffset("cache") << ") + paired * f32(" + << sin_cache.GetByOffset("cache") << ");\n" + << " return value;\n" + << "}\n" + << "fn block_score(token: u32, b: u32, block_index: u32, position: i32) -> f32 {\n" + << " let key_base = (b * uniforms.state_capacity + block_index) * uniforms.head_size;\n" + << " var score = 0.0;\n" + << " for (var head = 0u; head < uniforms.num_heads; head++) {\n" + << " var dot = 0.0;\n" + << " for (var d = 0u; d < uniforms.head_size; d++) {\n" + << " dot += query_value(token, head, d, position, b) * f32(" + << present_key_state.GetByOffset("key_base + d") << ");\n" + << " }\n" + << " score += max(dot, 0.0);\n" + << " }\n" + << " return score * uniforms.scale;\n" + << "}\n"; + + shader.MainFunctionBody() + << " let token = workgroup_idx;\n" + << " if (token >= uniforms.total_tokens || local_idx != 0u) { return; }\n" + << " let output_base = token * uniforms.capacity;\n" + << " for (var i = 0u; i < uniforms.capacity; i++) {\n" + << " " << selected_indices.SetByOffset("output_base + i", "-1") << "\n" + << " }\n" + << " let b = batch_of_token(token);\n" + << " let key_len_after = u32(" << present_state_lengths.GetByOffset("b * 2u") << ");\n" + << " let position = abs_position(token, b);\n" + << " let causal = causal_count(position);\n" + << " let visible_block_count = min(key_len_after, causal);\n" + << " let selected = min(uniforms.block_topk, visible_block_count);\n" + << " var previous_score = 0.0;\n" + << " var previous_index = -1i;\n" + << " var emitted_blocks = 0u;\n" + << " for (var rank = 0u; rank < selected; rank++) {\n" + << " var best_score = 0.0;\n" + << " var best_index = -1i;\n" + << " for (var candidate = 0u; candidate < visible_block_count; candidate++) {\n" + << " let score = block_score(token, b, candidate, position);\n" + << " if (previous_index >= 0 && !(score < previous_score || " + "(score == previous_score && i32(candidate) > previous_index))) { continue; }\n" + << " if (best_index < 0 || score > best_score || (score == best_score && i32(candidate) < " + "best_index)) {\n" + << " best_score = score;\n" + << " best_index = i32(candidate);\n" + << " }\n" + << " }\n" + << " if (best_index < 0) { break; }\n" + << " for (var t = 0u; t < uniforms.compress_ratio; t++) {\n" + << " " + << selected_indices.SetByOffset("output_base + rank * uniforms.compress_ratio + t", + "best_index * i32(uniforms.compress_ratio) + i32(t)") + << "\n" + << " }\n" + << " emitted_blocks = rank + 1u;\n" + << " previous_score = best_score;\n" + << " previous_index = best_index;\n" + << " }\n" + << " let block_start = i32(visible_block_count * uniforms.compress_ratio);\n" + << " let natural_tail = select(0, position - block_start + 1, position >= block_start);\n" + << " let remaining_capacity = i32(uniforms.capacity) - i32(emitted_blocks * uniforms.compress_ratio);\n" + << " var tail_count = min(natural_tail, remaining_capacity);\n" + << " tail_count = max(tail_count, 0);\n" + << " for (var t = 0; t < tail_count; t++) {\n" + << " " + << selected_indices.SetByOffset("output_base + emitted_blocks * uniforms.compress_ratio + u32(t)", + "block_start + t") + << "\n" + << " }\n" + << " " + << selected_counts.SetByOffset("token", "i32(emitted_blocks * uniforms.compress_ratio) + tail_count") + << "\n"; + return Status::OK(); +} + +Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& key = shader.AddInput("key", ShaderUsage::UseUniform); + const auto& gate = shader.AddInput("gate", ShaderUsage::UseUniform); + const auto& norm = shader.AddInput("key_norm_weight", ShaderUsage::UseUniform); + const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); + const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& position_bias = shader.AddInput("position_bias", ShaderUsage::UseUniform); + const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); + const auto& past_kv_buffer = shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); + const auto& past_gate_buffer = shader.AddInput("past_gate_buffer", ShaderUsage::UseUniform); + const auto& past_state_lengths = shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); + const auto& present_key_state = + shader.AddOutput("present_key_state", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& present_kv_buffer = + shader.AddOutput("present_kv_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& present_gate_buffer = + shader.AddOutput("present_gate_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); + const auto& present_state_lengths = shader.AddOutput("present_state_lengths", ShaderUsage::UseUniform); + + shader.AdditionalImplementation() + << "fn extended_key(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, channel: u32) -> f32 {\n" + << " let width = 2u * uniforms.head_size;\n" + << " if (virtual_pos < old_buf_len) {\n" + << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * width + channel;\n" + << " return f32(" << past_kv_buffer.GetByOffset("idx") << ");\n" + << " }\n" + << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\n" + << " return f32(" << key.GetByOffset("idx2") << ");\n" + << "}\n" + << "fn extended_gate(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, channel: u32) -> f32 {\n" + << " let width = 2u * uniforms.head_size;\n" + << " if (virtual_pos < old_buf_len) {\n" + << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * width + channel;\n" + << " return f32(" << past_gate_buffer.GetByOffset("idx") << ");\n" + << " }\n" + << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\n" + << " return f32(" << gate.GetByOffset("idx2") << ");\n" + << "}\n" + << "fn pooled(b: u32, k: i32, old_buf_len: i32, req_start: i32, overlap_length: i32, d: u32) -> f32 {\n" + << " let width = 2u * uniforms.head_size;\n" + << " let has_previous = k >= 1 || overlap_length >= i32(uniforms.compress_ratio);\n" + << " let previous_base = overlap_length + (k - 1) * i32(uniforms.compress_ratio);\n" + << " let current_base = overlap_length + k * i32(uniforms.compress_ratio);\n" + << " var max_gate = -3.4028234663852886e+38;\n" + << " if (has_previous) {\n" + << " for (var slot = 0u; slot < uniforms.compress_ratio; slot++) {\n" + << " let v = extended_gate(b, previous_base + i32(slot), old_buf_len, req_start, d) + f32(" + << position_bias.GetByOffset("slot * width + d") << ");\n" + << " max_gate = max(max_gate, v);\n" + << " }\n" + << " }\n" + << " for (var slot = 0u; slot < uniforms.compress_ratio; slot++) {\n" + << " let v = extended_gate(b, current_base + i32(slot), old_buf_len, req_start, " + "uniforms.head_size + d) + f32(" + << position_bias.GetByOffset("slot * width + uniforms.head_size + d") << ");\n" + << " max_gate = max(max_gate, v);\n" + << " }\n" + << " var denominator = 0.0;\n" + << " var accumulator = 0.0;\n" + << " if (has_previous) {\n" + << " for (var slot = 0u; slot < uniforms.compress_ratio; slot++) {\n" + << " let logit = extended_gate(b, previous_base + i32(slot), old_buf_len, req_start, d) + f32(" + << position_bias.GetByOffset("slot * width + d") << ");\n" + << " let weight = exp(logit - max_gate);\n" + << " denominator += weight;\n" + << " accumulator += weight * extended_key(b, previous_base + i32(slot), old_buf_len, req_start, d);\n" + << " }\n" + << " }\n" + << " for (var slot = 0u; slot < uniforms.compress_ratio; slot++) {\n" + << " let logit = extended_gate(b, current_base + i32(slot), old_buf_len, req_start, " + "uniforms.head_size + d) + f32(" + << position_bias.GetByOffset("slot * width + uniforms.head_size + d") << ");\n" + << " let weight = exp(logit - max_gate);\n" + << " denominator += weight;\n" + << " accumulator += weight * extended_key(b, current_base + i32(slot), old_buf_len, req_start, " + "uniforms.head_size + d);\n" + << " }\n" + << " return select(0.0, accumulator / denominator, denominator > 0.0);\n" + << "}\n" + << "fn clamp_position(position: i32) -> u32 {\n" + << " if (position < 0) { return 0u; }\n" + << " return min(u32(position), uniforms.max_rotary_length - 1u);\n" + << "}\n"; + + shader.MainFunctionBody() + << " let b = workgroup_idx;\n" + << " if (b >= uniforms.batch_size || local_idx != 0u) { return; }\n" + << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" + << " let req_end = " << cu_seqlens.GetByOffset("b + 1u") << ";\n" + << " let req_len = max(req_end - req_start, 0);\n" + << " let old_key_len = clamp(" << past_state_lengths.GetByOffset("b * 2u") + << ", 0, i32(uniforms.state_capacity));\n" + << " let old_buf_len = clamp(" << past_state_lengths.GetByOffset("b * 2u + 1u") + << ", 0, i32(uniforms.buffer_capacity));\n" + << " let overlap_length = select(0, i32(uniforms.compress_ratio), old_buf_len >= " + "i32(uniforms.compress_ratio));\n" + << " let leftover_length = old_buf_len - overlap_length;\n" + << " let pending = leftover_length + req_len;\n" + << " let full_new_window_count = pending / i32(uniforms.compress_ratio);\n" + << " let capacity_left = max(i32(uniforms.state_capacity) - old_key_len, 0);\n" + << " let new_window_count = min(full_new_window_count, capacity_left);\n" + << " let overflowed = new_window_count < full_new_window_count;\n" + << " var present_buffer_length = 0;\n" + << " var present_buffer_start = 0;\n" + << " if (!overflowed) {\n" + << " if (new_window_count > 0) {\n" + << " present_buffer_length = i32(uniforms.compress_ratio) + pending % i32(uniforms.compress_ratio);\n" + << " present_buffer_start = overlap_length + (new_window_count - 1) * i32(uniforms.compress_ratio);\n" + << " } else {\n" + << " present_buffer_length = old_buf_len + req_len;\n" + << " present_buffer_start = 0;\n" + << " }\n" + << " present_buffer_length = min(present_buffer_length, i32(uniforms.buffer_capacity));\n" + << " }\n" + << " " << present_state_lengths.SetByOffset("b * 2u", "old_key_len + new_window_count") << "\n" + << " " << present_state_lengths.SetByOffset("b * 2u + 1u", "present_buffer_length") << "\n" + << " for (var k = 0; k < new_window_count; k++) {\n" + << " var square_sum = 0.0;\n" + << " for (var c = 0u; c < uniforms.head_size; c++) {\n" + << " let v = pooled(b, k, old_buf_len, req_start, overlap_length, c);\n" + << " square_sum += v * v;\n" + << " }\n" + << " let inverse_rms = inverseSqrt(square_sum / f32(uniforms.head_size) + uniforms.epsilon);\n" + << " let entry = u32(old_key_len + k);\n" + << " let position = clamp_position(i32(entry) * i32(uniforms.compress_ratio));\n"; + if (cos_cache_batched_) { + shader.MainFunctionBody() + << " let cache_base = (b * uniforms.max_rotary_length + position) * uniforms.rotary_width;\n"; + } else { + shader.MainFunctionBody() << " let cache_base = position * uniforms.rotary_width;\n"; + } + shader.MainFunctionBody() + << " let rotary_base = uniforms.head_size - 2u * uniforms.rotary_width;\n" + << " for (var d = 0u; d < uniforms.head_size; d++) {\n" + << " var value = pooled(b, k, old_buf_len, req_start, overlap_length, d) * inverse_rms * f32(" + << norm.GetByOffset("d") << ");\n" + << " if (d >= rotary_base) {\n" + << " let offset = d - rotary_base;\n" + << " let pair_d = select(d - 1u, d + 1u, (offset & 1u) == 0u);\n" + << " let sign = select(1.0, -1.0, (offset & 1u) == 0u);\n" + << " let paired = sign * pooled(b, k, old_buf_len, req_start, overlap_length, pair_d) * " + "inverse_rms * f32(" + << norm.GetByOffset("pair_d") << ");\n" + << " value = value * f32(" << cos_cache.GetByOffset("cache_base + offset / 2u") << ") + paired * f32(" + << sin_cache.GetByOffset("cache_base + offset / 2u") << ");\n" + << " }\n" + << " " + << present_key_state.SetByOffset("(b * uniforms.state_capacity + entry) * uniforms.head_size + d", + "present_key_state_element_t(value)") + << "\n" + << " }\n" + << " }\n" + << " for (var t = 0; t < present_buffer_length; t++) {\n" + << " let virtual_pos = present_buffer_start + t;\n" + << " for (var c = 0u; c < 2u * uniforms.head_size; c++) {\n" + << " let key_value = extended_key(b, virtual_pos, old_buf_len, req_start, c);\n" + << " let gate_value = extended_gate(b, virtual_pos, old_buf_len, req_start, c);\n" + << " " + << present_kv_buffer.SetByOffset("(b * uniforms.buffer_capacity + u32(t)) * (2u * uniforms.head_size) + c", + "present_kv_buffer_element_t(key_value)") + << "\n" + << " " + << present_gate_buffer.SetByOffset( + "(b * uniforms.buffer_capacity + u32(t)) * (2u * uniforms.head_size) + c", + "present_gate_buffer_element_t(gate_value)") + << "\n" + << " }\n" + << " }\n"; + return Status::OK(); +} + +Status PackedSparseAttentionIndexerCsaSelectProgram::GenerateShaderCode(ShaderHelper& shader) const { + const auto& query = shader.AddInput("query", ShaderUsage::UseUniform); + const auto& present_key_state = shader.AddInput("present_key_state", ShaderUsage::UseUniform); + const auto& head_weights = shader.AddInput("head_weights", ShaderUsage::UseUniform); + const auto& position_ids = shader.AddInput("position_ids", ShaderUsage::UseUniform); + const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); + const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); + const auto& present_state_lengths = shader.AddInput("present_state_lengths", ShaderUsage::UseUniform); + const auto& selected_indices = shader.AddOutput("selected_indices", ShaderUsage::UseUniform); + const auto& selected_counts = shader.AddOutput("selected_counts", ShaderUsage::UseUniform); + + shader.AdditionalImplementation() + << "fn batch_of_token(token: u32) -> u32 {\n" + << " var b = 0u;\n" + << " for (var i = 0u; i < uniforms.batch_size; i++) {\n" + << " if (u32(" << cu_seqlens.GetByOffset("i") << ") <= token) { b = i; } else { break; }\n" + << " }\n" + << " return b;\n" + << "}\n" + << "fn clamped_position(raw: vec2) -> u32 {\n" + << " if ((raw.y & 0x80000000u) != 0u) { return 0u; }\n" + << " if (raw.y != 0u) { return 0xffffffffu; }\n" + << " return raw.x;\n" + << "}\n" + << "fn causal_count(raw: vec2) -> u32 {\n" + << " if ((raw.y & 0x80000000u) != 0u) { return 0u; }\n" + << " if (raw.y != 0u) { return 0xffffffffu; }\n" + << " let p = raw.x;\n" + << " let cr = uniforms.compress_ratio;\n" + << " return p / cr + select(0u, 1u, p % cr == cr - 1u);\n" + << "}\n" + << "fn query_value(token: u32, head: u32, d: u32, raw: vec2, b: u32) -> f32 {\n" + << " let base = (token * uniforms.num_heads + head) * uniforms.head_size;\n" + << " var value = f32(" << query.GetByOffset("base + d") << ");\n" + << " let rotary_base = uniforms.head_size - 2u * uniforms.rotary_width;\n" + << " if (d < rotary_base) { return value; }\n" + << " let offset = d - rotary_base;\n" + << " let pair_d = select(d - 1u, d + 1u, (offset & 1u) == 0u);\n" + << " let sign = select(1.0, -1.0, (offset & 1u) == 0u);\n" + << " let paired = sign * f32(" << query.GetByOffset("base + pair_d") << ");\n" + << " let position = min(clamped_position(raw), uniforms.max_rotary_length - 1u);\n"; + if (cos_cache_batched_) { + shader.AdditionalImplementation() + << " let cache = (b * uniforms.max_rotary_length + position) * uniforms.rotary_width + offset / 2u;\n"; + } else { + shader.AdditionalImplementation() + << " let cache = position * uniforms.rotary_width + offset / 2u;\n"; + } + shader.AdditionalImplementation() + << " value = value * f32(" << cos_cache.GetByOffset("cache") << ") + paired * f32(" + << sin_cache.GetByOffset("cache") << ");\n" + << " return value;\n" + << "}\n" + << "fn entry_score(token: u32, b: u32, entry: u32, raw: vec2) -> f32 {\n" + << " let key_base = (b * uniforms.state_capacity + entry) * uniforms.head_size;\n" + << " var score = 0.0;\n" + << " for (var head = 0u; head < uniforms.num_heads; head++) {\n" + << " var dot = 0.0;\n" + << " for (var d = 0u; d < uniforms.head_size; d++) {\n" + << " dot += query_value(token, head, d, raw, b) * f32(" + << present_key_state.GetByOffset("key_base + d") << ");\n" + << " }\n" + << " score += max(dot, 0.0) * f32(" << head_weights.GetByOffset("token * uniforms.num_heads + head") + << ");\n" + << " }\n" + << " return score * uniforms.scale * uniforms.head_weight_scale;\n" + << "}\n"; + + shader.MainFunctionBody() + << " let token = workgroup_idx;\n" + << " if (token >= uniforms.total_tokens || local_idx != 0u) { return; }\n" + << " let output_base = token * uniforms.capacity;\n" + << " for (var i = 0u; i < uniforms.capacity; i++) {\n" + << " " << selected_indices.SetByOffset("output_base + i", "-1") << "\n" + << " }\n" + << " let b = batch_of_token(token);\n" + << " let key_len_after = u32(" << present_state_lengths.GetByOffset("b * 2u") << ");\n" + << " let raw = " << position_ids.GetByOffset("token", true) << ";\n" + << " let threshold = min(causal_count(raw), key_len_after);\n" + << " let selected = min(uniforms.index_topk, threshold);\n" + << " var previous_score = 0.0;\n" + << " var previous_index = -1i;\n" + << " var emitted = 0u;\n" + << " for (var rank = 0u; rank < selected; rank++) {\n" + << " var best_score = 0.0;\n" + << " var best_index = -1i;\n" + << " for (var candidate = 0u; candidate < threshold; candidate++) {\n" + << " let score = entry_score(token, b, candidate, raw);\n" + << " if (previous_index >= 0 && !(score < previous_score || " + "(score == previous_score && i32(candidate) > previous_index))) { continue; }\n" + << " if (best_index < 0 || score > best_score || (score == best_score && i32(candidate) < " + "best_index)) {\n" + << " best_score = score;\n" + << " best_index = i32(candidate);\n" + << " }\n" + << " }\n" + << " if (best_index < 0) { break; }\n" + << " " << selected_indices.SetByOffset("output_base + rank", "best_index") << "\n" + << " emitted = rank + 1u;\n" + << " previous_score = best_score;\n" + << " previous_index = best_index;\n" + << " }\n" + << " " << selected_counts.SetByOffset("token", "i32(emitted)") << "\n"; + return Status::OK(); +} + +PackedSparseAttentionIndexer::PackedSparseAttentionIndexer(const OpKernelInfo& info) : WebGpuKernel(info) { + std::string policy_mode; + ORT_ENFORCE(info.GetAttr("policy_mode", &policy_mode).IsOK(), + "PackedSparseAttentionIndexer: policy_mode is required"); + ORT_ENFORCE(psai::TryParsePolicy(policy_mode, policy_), "PackedSparseAttentionIndexer: policy_mode must be '", + psai::kPolicyModeQsa, "' or '", psai::kPolicyModeCsa, "', got '", policy_mode, "'"); + ORT_ENFORCE(info.GetAttr("compress_ratio", &compress_ratio_).IsOK(), + "PackedSparseAttentionIndexer: compress_ratio is required"); + ORT_ENFORCE(compress_ratio_ > 0, "PackedSparseAttentionIndexer: compress_ratio must be > 0"); + + const bool has_token_budget = info.GetAttr("token_budget", &token_budget_).IsOK(); + const bool has_index_topk = info.GetAttr("index_topk", &index_topk_).IsOK(); + has_scale_ = info.GetAttr("scale", &scale_).IsOK(); + has_head_weight_scale_ = info.GetAttr("head_weight_scale", &head_weight_scale_).IsOK(); + if (policy_ == psai::Policy::kQsa) { + ORT_ENFORCE(has_token_budget && token_budget_ > 0 && token_budget_ % compress_ratio_ == 0, + "PackedSparseAttentionIndexer: token_budget must be > 0 and divisible by compress_ratio for qsa"); + ORT_ENFORCE(!has_index_topk && !has_head_weight_scale_, + "PackedSparseAttentionIndexer: csa attributes must be omitted for qsa"); + index_topk_ = 0; + } else { + ORT_ENFORCE(has_index_topk && index_topk_ > 0, "PackedSparseAttentionIndexer: index_topk must be > 0 for csa"); + ORT_ENFORCE(!has_token_budget, "PackedSparseAttentionIndexer: token_budget must be omitted for csa"); + token_budget_ = 0; + } + epsilon_ = info.GetAttrOrDefault("epsilon", 1.0e-6f); + ORT_ENFORCE(epsilon_ >= 0.0f, "PackedSparseAttentionIndexer: epsilon must be >= 0"); +} + +Status PackedSparseAttentionIndexer::ComputeInternal(onnxruntime::webgpu::ComputeContext& context) const { + const bool is_qsa = policy_ == psai::Policy::kQsa; + constexpr int kCsaOnlyInputs[] = {psai::kGate, psai::kPositionBias, psai::kHeadWeights}; + for (int index : kCsaOnlyInputs) { + const bool provided = index < context.InputCount() && context.Input(index) != nullptr; + ORT_RETURN_IF(provided != !is_qsa, "PackedSparseAttentionIndexer: input ", index, + provided ? " must be omitted for policy_mode 'qsa'" : " is required for policy_mode 'csa'"); + } + const bool position_ids_provided = + psai::kPositionIds < context.InputCount() && context.Input(psai::kPositionIds) != nullptr; + ORT_RETURN_IF(!is_qsa && !position_ids_provided, + "PackedSparseAttentionIndexer: position_ids is required for " + "policy_mode 'csa'"); + const bool gate_buffer_provided = + psai::kPastGateBuffer < context.InputCount() && context.Input(psai::kPastGateBuffer) != nullptr; + ORT_RETURN_IF(gate_buffer_provided != !is_qsa, "PackedSparseAttentionIndexer: past_gate_buffer ", + gate_buffer_provided ? "must be omitted for policy_mode 'qsa'" + : "is required for policy_mode 'csa'"); + return is_qsa ? ComputeQsa(context) : ComputeCsa(context); +} + +Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeContext& context) const { + const Tensor* query = context.Input(psai::kQuery); + const Tensor* key = context.Input(psai::kKey); + const Tensor* norm = context.Input(psai::kKeyNormWeight); + const Tensor* cos_cache = context.Input(psai::kCosCache); + const Tensor* sin_cache = context.Input(psai::kSinCache); + const Tensor* cu_seqlens = context.Input(psai::kCumulativeSequenceLengths); + const Tensor* past_seqlens = context.Input(psai::kPastSequenceLengths); + const Tensor* position_ids = context.Input(psai::kPositionIds); + const Tensor* past_key_state = context.Input(psai::kPastKeyState); + const Tensor* past_kv_buffer = context.Input(psai::kPastKvBuffer); + const Tensor* past_state_lengths = context.Input(psai::kPastStateLengths); + + ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); + const auto& query_shape = query->Shape(); + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, "PackedSparseAttentionIndexer: query must have rank 3"); + const int64_t total_tokens = query_shape[0]; + const int64_t num_heads = query_shape[1]; + const int64_t head_size = query_shape[2]; + ORT_RETURN_IF_NOT(num_heads > 0 && head_size > 0, "PackedSparseAttentionIndexer: invalid query dimensions"); + + ORT_RETURN_IF(cu_seqlens == nullptr, "PackedSparseAttentionIndexer: cumulative_sequence_lengths is required"); + const auto& cu_shape = cu_seqlens->Shape(); + ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1, + "PackedSparseAttentionIndexer: invalid cumulative_sequence_lengths shape"); + const int64_t batch_size = cu_shape[0] - 1; + + ORT_RETURN_IF_ERROR(CheckShape(past_seqlens, "past_sequence_lengths", {batch_size})); + ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, head_size})); + ORT_RETURN_IF_ERROR(CheckShape(norm, "key_norm_weight", {head_size})); + if (position_ids != nullptr) { + ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens})); + } + + RotaryCacheShape rotary; + ORT_RETURN_IF_ERROR(CheckRotaryCache(cos_cache, sin_cache, batch_size, rotary)); + ORT_RETURN_IF_NOT(rotary.rotary_width > 0 && rotary.rotary_width % 2 == 0 && rotary.rotary_width <= head_size, + "PackedSparseAttentionIndexer: invalid qsa rotary_width"); + + ORT_RETURN_IF(past_key_state == nullptr, "PackedSparseAttentionIndexer: past_key_state is required"); + const auto& key_state_shape = past_key_state->Shape(); + ORT_RETURN_IF_NOT(key_state_shape.NumDimensions() == 3 && key_state_shape[0] == batch_size && + key_state_shape[2] == head_size, + "PackedSparseAttentionIndexer: invalid past_key_state shape"); + const int64_t state_capacity = key_state_shape[1]; + + const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); + ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, head_size})); + ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths", + {batch_size, psai::kStateLengthColumns})); + + const int64_t capacity = psai::SelectedCapacity(psai::Policy::kQsa, token_budget_, index_topk_, compress_ratio_); + Tensor* selected_indices = context.Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity})); + Tensor* selected_counts = context.Output(psai::kSelectedCounts, TensorShape({total_tokens})); + Tensor* present_key_state = context.Output(psai::kPresentKeyState, key_state_shape); + Tensor* present_kv_buffer = + context.Output(psai::kPresentKvBuffer, TensorShape({batch_size, buffer_capacity, head_size})); + Tensor* present_state_lengths = + context.Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns})); + + if (present_key_state->DataRaw() != past_key_state->DataRaw()) { + const int64_t total = present_key_state->Shape().Size(); + if (total > 0) { + PackedSparseAttentionIndexerCopyProgram copy; + copy.SetWorkgroupSize(kWorkgroupSize) + .AddInput({past_key_state, ProgramTensorMetadataDependency::Type}) + .AddOutput({present_key_state, ProgramTensorMetadataDependency::Type}) + .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) + .AddUniformVariables({{ToUint32(total)}}); + ORT_RETURN_IF_ERROR(context.RunProgram(copy)); + } + } + if (present_kv_buffer->DataRaw() != past_kv_buffer->DataRaw()) { + const int64_t total = present_kv_buffer->Shape().Size(); + if (total > 0) { + PackedSparseAttentionIndexerCopyProgram copy; + copy.SetWorkgroupSize(kWorkgroupSize) + .AddInput({past_kv_buffer, ProgramTensorMetadataDependency::Type}) + .AddOutput({present_kv_buffer, ProgramTensorMetadataDependency::Type}) + .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) + .AddUniformVariables({{ToUint32(total)}}); + ORT_RETURN_IF_ERROR(context.RunProgram(copy)); + } + } + if (present_state_lengths->DataRaw() != past_state_lengths->DataRaw()) { + const int64_t total = present_state_lengths->Shape().Size(); + if (total > 0) { + PackedSparseAttentionIndexerCopyProgram copy; + copy.SetWorkgroupSize(kWorkgroupSize) + .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) + .AddUniformVariables({{ToUint32(total)}}); + ORT_RETURN_IF_ERROR(context.RunProgram(copy)); + } + } + + if (batch_size > 0) { + PackedSparseAttentionIndexerQsaUpdateProgram update{rotary.batched}; + update.CacheHint(rotary.batched) + .SetWorkgroupSize(kWorkgroupSize) + .AddInputs({{key, ProgramTensorMetadataDependency::Type}, + {norm, ProgramTensorMetadataDependency::Type}, + {cos_cache, ProgramTensorMetadataDependency::Type}, + {sin_cache, ProgramTensorMetadataDependency::Type}, + {cu_seqlens, ProgramTensorMetadataDependency::Type}, + {past_kv_buffer, ProgramTensorMetadataDependency::Type}}) + .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, + {present_kv_buffer, ProgramTensorMetadataDependency::Type}}) + .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .SetDispatchGroupSize(ToUint32(batch_size)) + .AddUniformVariables({{ToUint32(batch_size)}, + {ToUint32(compress_ratio_)}, + {ToUint32(state_capacity)}, + {ToUint32(buffer_capacity)}, + {ToUint32(head_size)}, + {ToUint32(rotary.rotary_width)}, + {ToUint32(rotary.max_rotary_length)}, + {epsilon_}}); + ORT_RETURN_IF_ERROR(context.RunProgram(update)); + } + + if (total_tokens == 0) { + return Status::OK(); + } + + PackedSparseAttentionIndexerQsaSelectProgram select{rotary.batched, position_ids != nullptr}; + select.CacheHint(rotary.batched, position_ids != nullptr) + .SetWorkgroupSize(kWorkgroupSize) + .AddInputs({{query, ProgramTensorMetadataDependency::Type}, + {present_key_state, ProgramTensorMetadataDependency::Type}, + {cos_cache, ProgramTensorMetadataDependency::Type}, + {sin_cache, ProgramTensorMetadataDependency::Type}, + {cu_seqlens, ProgramTensorMetadataDependency::Type}, + {past_seqlens, ProgramTensorMetadataDependency::Type}}); + if (position_ids != nullptr) { + select.AddInput({position_ids, ProgramTensorMetadataDependency::Type}); + } + select.AddInput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddOutputs({{selected_indices, ProgramTensorMetadataDependency::Type}, + {selected_counts, ProgramTensorMetadataDependency::Type}}) + .SetDispatchGroupSize(ToUint32(total_tokens)) + .AddUniformVariables({{ToUint32(total_tokens)}, + {ToUint32(batch_size)}, + {ToUint32(num_heads)}, + {ToUint32(head_size)}, + {ToUint32(rotary.rotary_width)}, + {ToUint32(rotary.max_rotary_length)}, + {ToUint32(compress_ratio_)}, + {ToUint32(state_capacity)}, + {ToUint32(capacity)}, + {ToUint32(token_budget_ / compress_ratio_)}, + {epsilon_}, + {has_scale_ ? scale_ : 1.0f / std::sqrt(static_cast(head_size))}}); + return context.RunProgram(select); +} + +Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeContext& context) const { + const Tensor* query = context.Input(psai::kQuery); + const Tensor* key = context.Input(psai::kKey); + const Tensor* norm = context.Input(psai::kKeyNormWeight); + const Tensor* cos_cache = context.Input(psai::kCosCache); + const Tensor* sin_cache = context.Input(psai::kSinCache); + const Tensor* cu_seqlens = context.Input(psai::kCumulativeSequenceLengths); + const Tensor* past_seqlens = context.Input(psai::kPastSequenceLengths); + const Tensor* gate = context.Input(psai::kGate); + const Tensor* position_bias = context.Input(psai::kPositionBias); + const Tensor* head_weights = context.Input(psai::kHeadWeights); + const Tensor* position_ids = context.Input(psai::kPositionIds); + const Tensor* past_key_state = context.Input(psai::kPastKeyState); + const Tensor* past_kv_buffer = context.Input(psai::kPastKvBuffer); + const Tensor* past_gate_buffer = context.Input(psai::kPastGateBuffer); + const Tensor* past_state_lengths = context.Input(psai::kPastStateLengths); + + ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); + const auto& query_shape = query->Shape(); + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, "PackedSparseAttentionIndexer: query must have rank 3"); + const int64_t total_tokens = query_shape[0]; + const int64_t num_heads = query_shape[1]; + const int64_t head_size = query_shape[2]; + ORT_RETURN_IF_NOT(num_heads > 0 && head_size > 0, "PackedSparseAttentionIndexer: invalid query dimensions"); + const int64_t width = 2 * head_size; + + ORT_RETURN_IF(cu_seqlens == nullptr, "PackedSparseAttentionIndexer: cumulative_sequence_lengths is required"); + const auto& cu_shape = cu_seqlens->Shape(); + ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1, + "PackedSparseAttentionIndexer: invalid cumulative_sequence_lengths shape"); + const int64_t batch_size = cu_shape[0] - 1; + + ORT_RETURN_IF_ERROR(CheckShape(past_seqlens, "past_sequence_lengths", {batch_size})); + ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, width})); + ORT_RETURN_IF_ERROR(CheckShape(norm, "key_norm_weight", {head_size})); + ORT_RETURN_IF_ERROR(CheckShape(gate, "gate", {total_tokens, width})); + ORT_RETURN_IF_ERROR(CheckShape(position_bias, "position_bias", {compress_ratio_, width})); + ORT_RETURN_IF_ERROR(CheckShape(head_weights, "head_weights", {total_tokens, num_heads})); + ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens})); + + RotaryCacheShape rotary; + ORT_RETURN_IF_ERROR(CheckRotaryCache(cos_cache, sin_cache, batch_size, rotary)); + ORT_RETURN_IF_NOT(rotary.rotary_width > 0 && 2 * rotary.rotary_width <= head_size, + "PackedSparseAttentionIndexer: invalid csa rotary_width"); + + ORT_RETURN_IF(past_key_state == nullptr, "PackedSparseAttentionIndexer: past_key_state is required"); + const auto& key_state_shape = past_key_state->Shape(); + ORT_RETURN_IF_NOT(key_state_shape.NumDimensions() == 3 && key_state_shape[0] == batch_size && + key_state_shape[2] == head_size, + "PackedSparseAttentionIndexer: invalid past_key_state shape"); + const int64_t state_capacity = key_state_shape[1]; + + const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); + ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, width})); + ORT_RETURN_IF_ERROR(CheckShape(past_gate_buffer, "past_gate_buffer", {batch_size, buffer_capacity, width})); + ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths", + {batch_size, psai::kStateLengthColumns})); + + const int64_t capacity = psai::SelectedCapacity(psai::Policy::kCsa, token_budget_, index_topk_, compress_ratio_); + Tensor* selected_indices = context.Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity})); + Tensor* selected_counts = context.Output(psai::kSelectedCounts, TensorShape({total_tokens})); + Tensor* present_key_state = context.Output(psai::kPresentKeyState, key_state_shape); + Tensor* present_kv_buffer = + context.Output(psai::kPresentKvBuffer, TensorShape({batch_size, buffer_capacity, width})); + Tensor* present_gate_buffer = + context.Output(psai::kPresentGateBuffer, TensorShape({batch_size, buffer_capacity, width})); + Tensor* present_state_lengths = + context.Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns})); + + auto copy_if_needed = [&](const Tensor* src, Tensor* dst) -> Status { + if (src->DataRaw() == dst->DataRaw()) { + return Status::OK(); + } + const int64_t total = dst->Shape().Size(); + if (total == 0) { + return Status::OK(); + } + PackedSparseAttentionIndexerCopyProgram copy; + copy.SetWorkgroupSize(kWorkgroupSize) + .AddInput({src, ProgramTensorMetadataDependency::Type}) + .AddOutput({dst, ProgramTensorMetadataDependency::Type}) + .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) + .AddUniformVariables({{ToUint32(total)}}); + return context.RunProgram(copy); + }; + ORT_RETURN_IF_ERROR(copy_if_needed(past_key_state, present_key_state)); + ORT_RETURN_IF_ERROR(copy_if_needed(past_kv_buffer, present_kv_buffer)); + ORT_RETURN_IF_ERROR(copy_if_needed(past_gate_buffer, present_gate_buffer)); + ORT_RETURN_IF_ERROR(copy_if_needed(past_state_lengths, present_state_lengths)); + + if (batch_size > 0) { + PackedSparseAttentionIndexerCsaUpdateProgram update{rotary.batched}; + update.CacheHint(rotary.batched) + .SetWorkgroupSize(kWorkgroupSize) + .AddInputs({{key, ProgramTensorMetadataDependency::Type}, + {gate, ProgramTensorMetadataDependency::Type}, + {norm, ProgramTensorMetadataDependency::Type}, + {cos_cache, ProgramTensorMetadataDependency::Type}, + {sin_cache, ProgramTensorMetadataDependency::Type}, + {position_bias, ProgramTensorMetadataDependency::Type}, + {cu_seqlens, ProgramTensorMetadataDependency::Type}, + {past_kv_buffer, ProgramTensorMetadataDependency::Type}, + {past_gate_buffer, ProgramTensorMetadataDependency::Type}}) + .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, + {present_kv_buffer, ProgramTensorMetadataDependency::Type}, + {present_gate_buffer, ProgramTensorMetadataDependency::Type}}) + .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .SetDispatchGroupSize(ToUint32(batch_size)) + .AddUniformVariables({{ToUint32(batch_size)}, + {ToUint32(compress_ratio_)}, + {ToUint32(state_capacity)}, + {ToUint32(buffer_capacity)}, + {ToUint32(head_size)}, + {ToUint32(rotary.rotary_width)}, + {ToUint32(rotary.max_rotary_length)}, + {epsilon_}}); + ORT_RETURN_IF_ERROR(context.RunProgram(update)); + } + + if (total_tokens == 0) { + return Status::OK(); + } + + PackedSparseAttentionIndexerCsaSelectProgram select{rotary.batched}; + select.CacheHint(rotary.batched) + .SetWorkgroupSize(kWorkgroupSize) + .AddInputs({{query, ProgramTensorMetadataDependency::Type}, + {present_key_state, ProgramTensorMetadataDependency::Type}, + {head_weights, ProgramTensorMetadataDependency::Type}, + {position_ids, ProgramTensorMetadataDependency::Type}, + {cos_cache, ProgramTensorMetadataDependency::Type}, + {sin_cache, ProgramTensorMetadataDependency::Type}, + {cu_seqlens, ProgramTensorMetadataDependency::Type}}) + .AddInput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddOutputs({{selected_indices, ProgramTensorMetadataDependency::Type}, + {selected_counts, ProgramTensorMetadataDependency::Type}}) + .SetDispatchGroupSize(ToUint32(total_tokens)) + .AddUniformVariables({{ToUint32(total_tokens)}, + {ToUint32(batch_size)}, + {ToUint32(num_heads)}, + {ToUint32(head_size)}, + {ToUint32(rotary.rotary_width)}, + {ToUint32(rotary.max_rotary_length)}, + {ToUint32(compress_ratio_)}, + {ToUint32(state_capacity)}, + {ToUint32(capacity)}, + {ToUint32(index_topk_)}, + {has_scale_ ? scale_ : 1.0f / std::sqrt(static_cast(head_size))}, + {has_head_weight_scale_ ? head_weight_scale_ + : 1.0f / std::sqrt(static_cast(num_heads))}}); + return context.RunProgram(select); +} + +} // namespace webgpu +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h new file mode 100644 index 0000000000000..c25b17b52e7ad --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h @@ -0,0 +1,151 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h" +#include "core/providers/webgpu/program.h" +#include "core/providers/webgpu/webgpu_kernel.h" + +namespace onnxruntime { +namespace contrib { +namespace webgpu { + +using namespace onnxruntime::webgpu; + +// Copies past_* state into present_* state unchanged (element-for-element); used as the baseline +// before the update programs below overwrite only the newly produced entries. Works for any +// tensor element type via UseElementTypeAlias, so it is reused for key_state / kv_buffer / +// gate_buffer (T) and state_lengths (int32). +class PackedSparseAttentionIndexerCopyProgram final + : public Program { + public: + PackedSparseAttentionIndexerCopyProgram() : Program{"PackedSparseAttentionIndexerCopy"} {} + Status GenerateShaderCode(ShaderHelper& shader) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( + {"total", ProgramUniformVariableDataType::Uint32}); +}; + +// One invocation per request: forms every newly-closed compress_ratio block (mean-pool -> +// RMSNorm -> leading RoPE -> append) and publishes the raw trailing buffer. +class PackedSparseAttentionIndexerQsaUpdateProgram final + : public Program { + public: + explicit PackedSparseAttentionIndexerQsaUpdateProgram(bool cos_cache_batched) + : Program{"PackedSparseAttentionIndexerQsaUpdate"}, cos_cache_batched_{cos_cache_batched} {} + Status GenerateShaderCode(ShaderHelper& shader) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( + {"batch_size", ProgramUniformVariableDataType::Uint32}, + {"compress_ratio", ProgramUniformVariableDataType::Uint32}, + {"state_capacity", ProgramUniformVariableDataType::Uint32}, + {"buffer_capacity", ProgramUniformVariableDataType::Uint32}, + {"head_size", ProgramUniformVariableDataType::Uint32}, + {"rotary_width", ProgramUniformVariableDataType::Uint32}, + {"max_rotary_length", ProgramUniformVariableDataType::Uint32}, + {"epsilon", ProgramUniformVariableDataType::Float32}); + + private: + bool cos_cache_batched_; +}; + +// One invocation per query token: rotates the query, scores it against every causally visible +// prepared key_state entry, selects the token_budget / compress_ratio highest scoring blocks, and +// appends the causally visible tokens of the trailing incomplete block. +class PackedSparseAttentionIndexerQsaSelectProgram final + : public Program { + public: + PackedSparseAttentionIndexerQsaSelectProgram(bool cos_cache_batched, bool has_position_ids) + : Program{"PackedSparseAttentionIndexerQsaSelect"}, + cos_cache_batched_{cos_cache_batched}, + has_position_ids_{has_position_ids} {} + Status GenerateShaderCode(ShaderHelper& shader) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( + {"total_tokens", ProgramUniformVariableDataType::Uint32}, + {"batch_size", ProgramUniformVariableDataType::Uint32}, + {"num_heads", ProgramUniformVariableDataType::Uint32}, + {"head_size", ProgramUniformVariableDataType::Uint32}, + {"rotary_width", ProgramUniformVariableDataType::Uint32}, + {"max_rotary_length", ProgramUniformVariableDataType::Uint32}, + {"compress_ratio", ProgramUniformVariableDataType::Uint32}, + {"state_capacity", ProgramUniformVariableDataType::Uint32}, + {"capacity", ProgramUniformVariableDataType::Uint32}, + {"block_topk", ProgramUniformVariableDataType::Uint32}, + {"epsilon", ProgramUniformVariableDataType::Float32}, + {"scale", ProgramUniformVariableDataType::Float32}); + + private: + bool cos_cache_batched_; + bool has_position_ids_; +}; + +// One invocation per request: closes every new compression window (softmax-gated pool -> RMSNorm +// -> trailing RoPE -> append) and publishes the raw overlap+leftover buffer. +class PackedSparseAttentionIndexerCsaUpdateProgram final + : public Program { + public: + explicit PackedSparseAttentionIndexerCsaUpdateProgram(bool cos_cache_batched) + : Program{"PackedSparseAttentionIndexerCsaUpdate"}, cos_cache_batched_{cos_cache_batched} {} + Status GenerateShaderCode(ShaderHelper& shader) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( + {"batch_size", ProgramUniformVariableDataType::Uint32}, + {"compress_ratio", ProgramUniformVariableDataType::Uint32}, + {"state_capacity", ProgramUniformVariableDataType::Uint32}, + {"buffer_capacity", ProgramUniformVariableDataType::Uint32}, + {"head_size", ProgramUniformVariableDataType::Uint32}, + {"rotary_width", ProgramUniformVariableDataType::Uint32}, + {"max_rotary_length", ProgramUniformVariableDataType::Uint32}, + {"epsilon", ProgramUniformVariableDataType::Float32}); + + private: + bool cos_cache_batched_; +}; + +// One invocation per query token: rotates the query, scores it against every causally visible +// compressed key_state entry, and selects the index_topk highest scoring entries. +class PackedSparseAttentionIndexerCsaSelectProgram final + : public Program { + public: + explicit PackedSparseAttentionIndexerCsaSelectProgram(bool cos_cache_batched) + : Program{"PackedSparseAttentionIndexerCsaSelect"}, cos_cache_batched_{cos_cache_batched} {} + Status GenerateShaderCode(ShaderHelper& shader) const override; + WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( + {"total_tokens", ProgramUniformVariableDataType::Uint32}, + {"batch_size", ProgramUniformVariableDataType::Uint32}, + {"num_heads", ProgramUniformVariableDataType::Uint32}, + {"head_size", ProgramUniformVariableDataType::Uint32}, + {"rotary_width", ProgramUniformVariableDataType::Uint32}, + {"max_rotary_length", ProgramUniformVariableDataType::Uint32}, + {"compress_ratio", ProgramUniformVariableDataType::Uint32}, + {"state_capacity", ProgramUniformVariableDataType::Uint32}, + {"capacity", ProgramUniformVariableDataType::Uint32}, + {"index_topk", ProgramUniformVariableDataType::Uint32}, + {"scale", ProgramUniformVariableDataType::Float32}, + {"head_weight_scale", ProgramUniformVariableDataType::Float32}); + + private: + bool cos_cache_batched_; +}; + +class PackedSparseAttentionIndexer final : public WebGpuKernel { + public: + explicit PackedSparseAttentionIndexer(const OpKernelInfo& info); + Status ComputeInternal(onnxruntime::webgpu::ComputeContext& context) const override; + + private: + Status ComputeQsa(onnxruntime::webgpu::ComputeContext& context) const; + Status ComputeCsa(onnxruntime::webgpu::ComputeContext& context) const; + + packed_sparse_attention_indexer::Policy policy_; + int64_t compress_ratio_; + int64_t token_budget_; + int64_t index_topk_; + float epsilon_; + float scale_; + float head_weight_scale_; + bool has_scale_; + bool has_head_weight_scale_; +}; + +} // namespace webgpu +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc b/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc index 8f6e2f458c526..1d932e849e451 100644 --- a/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc @@ -12,6 +12,7 @@ #include "contrib_ops/webgpu/bert/linear_attention_gates.h" #include "contrib_ops/webgpu/bert/paged_attention.h" #include "contrib_ops/webgpu/bert/sparse_attention_indexer.h" +#include "contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h" #include "core/framework/op_kernel.h" @@ -54,6 +55,7 @@ static const BuildKernelCreateInfoFn build_kernel_create_info_function_table[] = BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, // LayerNormalization used to be a contrib op that (incorrectly) used kOnnxDomain so we need to version it diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index 7120200ba4751..93435b40f8bc0 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -13,6 +13,7 @@ #include "core/graph/contrib_ops/shape_inference_functions.h" #include "contrib_ops/cpu/bert/attention_common.h" #include "contrib_ops/cpu/sparse/sparse_attention_indexer_common.h" +#include "contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h" // Suppress a warning: global initializer calls a non-constexpr function 'symbol' which is from // ONNX_OPERATOR_SET_SCHEMA_EX macro and only happens in debug build #if defined(_WIN32) && !defined(NDEBUG) @@ -2297,6 +2298,388 @@ ONNX_MS_OPERATOR_SET_SCHEMA( SparseAttentionIndexerTypeAndShapeInference(ctx); })); +namespace psai = ::onnxruntime::contrib::packed_sparse_attention_indexer; + +namespace { + +bool PackedSparseAttentionIndexerHasInput(ONNX_NAMESPACE::InferenceContext& ctx, int index) { + return static_cast(index) < ctx.getNumInputs() && ctx.getInputType(index) != nullptr; +} + +const ONNX_NAMESPACE::TensorShapeProto* PackedSparseAttentionIndexerShape(ONNX_NAMESPACE::InferenceContext& ctx, + int index, int expected_rank) { + if (!PackedSparseAttentionIndexerHasInput(ctx, index) || !hasInputShape(ctx, index)) { + return nullptr; + } + const auto& shape = getInputShape(ctx, index); + if (shape.dim_size() != expected_rank) { + fail_shape_inference("PackedSparseAttentionIndexer: input ", index, " must have rank ", expected_rank, + ", got rank ", shape.dim_size()); + } + return &shape; +} + +} // namespace + +void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& ctx) { + const std::string policy_mode = getAttribute(ctx, "policy_mode", std::string()); + psai::Policy policy = psai::Policy::kQsa; + if (!psai::TryParsePolicy(policy_mode, policy)) { + fail_shape_inference("PackedSparseAttentionIndexer: policy_mode must be 'qsa' or 'csa', got '", policy_mode, + "'"); + } + const bool is_qsa = policy == psai::Policy::kQsa; + + const int64_t compress_ratio = getAttribute(ctx, "compress_ratio", static_cast(0)); + if (compress_ratio <= 0 || compress_ratio > std::numeric_limits::max()) { + fail_shape_inference("PackedSparseAttentionIndexer: compress_ratio must be in (0, INT_MAX], got ", + compress_ratio); + } + const int64_t state_capacity = getAttribute(ctx, "state_capacity", static_cast(0)); + if (state_capacity <= 0 || state_capacity > std::numeric_limits::max()) { + fail_shape_inference("PackedSparseAttentionIndexer: state_capacity must be in (0, INT_MAX], got ", + state_capacity); + } + + const int64_t token_budget = getAttribute(ctx, "token_budget", static_cast(0)); + const int64_t index_topk = getAttribute(ctx, "index_topk", static_cast(0)); + if (is_qsa) { + if (ctx.getAttribute("index_topk") != nullptr || ctx.getAttribute("head_weight_scale") != nullptr) { + fail_shape_inference( + "PackedSparseAttentionIndexer: index_topk and head_weight_scale must not be set when policy_mode is " + "'qsa'"); + } + if (token_budget <= 0 || token_budget % compress_ratio != 0 || + token_budget > std::numeric_limits::max() - compress_ratio + 1) { + fail_shape_inference( + "PackedSparseAttentionIndexer: policy_mode 'qsa' requires token_budget > 0, divisible by " + "compress_ratio, and a selected capacity no greater than INT_MAX, got token_budget=", + token_budget, " compress_ratio=", compress_ratio); + } + } else { + if (ctx.getAttribute("token_budget") != nullptr) { + fail_shape_inference("PackedSparseAttentionIndexer: token_budget must not be set when policy_mode is 'csa'"); + } + if (index_topk <= 0 || index_topk > std::numeric_limits::max()) { + fail_shape_inference( + "PackedSparseAttentionIndexer: policy_mode 'csa' requires index_topk in (0, INT_MAX], got ", index_topk); + } + } + + // Strict policy input validation: every fixed slot required by every policy must be provided; + // csa-only slots must be provided iff policy_mode is 'csa'; position_ids is optional for 'qsa' + // and required for 'csa'. + for (int index : {psai::kQuery, psai::kKey, psai::kKeyNormWeight, psai::kCosCache, psai::kSinCache, + psai::kCumulativeSequenceLengths, psai::kPastSequenceLengths, psai::kPastKeyState, + psai::kPastKvBuffer, psai::kPastStateLengths}) { + if (!PackedSparseAttentionIndexerHasInput(ctx, index)) { + fail_shape_inference("PackedSparseAttentionIndexer: input ", index, " is required for every policy_mode"); + } + } + for (int index : {psai::kGate, psai::kPositionBias, psai::kHeadWeights, psai::kPastGateBuffer}) { + if (PackedSparseAttentionIndexerHasInput(ctx, index) == is_qsa) { + fail_shape_inference("PackedSparseAttentionIndexer: input ", index, + is_qsa ? " must be omitted when policy_mode is 'qsa'" + : " is required when policy_mode is 'csa'"); + } + } + if (!is_qsa && !PackedSparseAttentionIndexerHasInput(ctx, psai::kPositionIds)) { + fail_shape_inference("PackedSparseAttentionIndexer: input ", psai::kPositionIds, + " (position_ids) is required when policy_mode is 'csa'"); + } + + if (ctx.getNumOutputs() != static_cast(psai::kFixedOutputCount)) { + fail_shape_inference("PackedSparseAttentionIndexer: exactly ", psai::kFixedOutputCount, + " declared outputs are required, got ", ctx.getNumOutputs()); + } + + updateOutputElemType(ctx, psai::kSelectedIndices, ONNX_NAMESPACE::TensorProto_DataType_INT32); + updateOutputElemType(ctx, psai::kSelectedCounts, ONNX_NAMESPACE::TensorProto_DataType_INT32); + updateOutputElemType(ctx, psai::kPresentStateLengths, ONNX_NAMESPACE::TensorProto_DataType_INT32); + propagateElemTypeFromInputToOutput(ctx, psai::kPastKeyState, psai::kPresentKeyState); + propagateElemTypeFromInputToOutput(ctx, psai::kPastKvBuffer, psai::kPresentKvBuffer); + if (!is_qsa) { + // present_gate_buffer keeps its fixed positional slot (with an empty name) for policy_mode + // 'qsa'; only propagate its type/shape when it is actually produced. + propagateElemTypeFromInputToOutput(ctx, psai::kPastGateBuffer, psai::kPresentGateBuffer); + } + + (void)PackedSparseAttentionIndexerShape(ctx, psai::kKeyNormWeight, 1); + (void)PackedSparseAttentionIndexerShape(ctx, psai::kKey, 2); + (void)PackedSparseAttentionIndexerShape(ctx, psai::kCumulativeSequenceLengths, 1); + (void)PackedSparseAttentionIndexerShape(ctx, psai::kPastSequenceLengths, 1); + if (!is_qsa) { + (void)PackedSparseAttentionIndexerShape(ctx, psai::kGate, 2); + (void)PackedSparseAttentionIndexerShape(ctx, psai::kPositionBias, 2); + (void)PackedSparseAttentionIndexerShape(ctx, psai::kHeadWeights, 2); + } + if (PackedSparseAttentionIndexerHasInput(ctx, psai::kPositionIds)) { + (void)PackedSparseAttentionIndexerShape(ctx, psai::kPositionIds, 1); + } + + const auto* query_shape = PackedSparseAttentionIndexerShape(ctx, psai::kQuery, 3); + if (query_shape != nullptr) { + const auto& total_tokens_dim = query_shape->dim(0); + const auto& num_heads_dim = query_shape->dim(1); + if (num_heads_dim.has_dim_value() && num_heads_dim.dim_value() <= 0) { + fail_shape_inference("PackedSparseAttentionIndexer: num_heads must be > 0, got ", num_heads_dim.dim_value()); + } + + const int64_t capacity = psai::SelectedCapacity(policy, token_budget, index_topk, compress_ratio); + ONNX_NAMESPACE::TensorShapeProto selected_shape; + SparseAttentionIndexerAppendDim(selected_shape, total_tokens_dim); + selected_shape.add_dim()->set_dim_value(capacity); + updateOutputShape(ctx, psai::kSelectedIndices, selected_shape); + + ONNX_NAMESPACE::TensorShapeProto counts_shape; + SparseAttentionIndexerAppendDim(counts_shape, total_tokens_dim); + updateOutputShape(ctx, psai::kSelectedCounts, counts_shape); + } + + // State never grows: present_* always has exactly the same fixed shape as past_*. + const auto* key_state_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKeyState, 3); + if (key_state_shape != nullptr) { + updateOutputShape(ctx, psai::kPresentKeyState, *key_state_shape); + } + const auto* kv_buffer_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKvBuffer, 3); + if (kv_buffer_shape != nullptr) { + updateOutputShape(ctx, psai::kPresentKvBuffer, *kv_buffer_shape); + if (!is_qsa) { + updateOutputShape(ctx, psai::kPresentGateBuffer, *kv_buffer_shape); + } + } + const auto* state_lengths_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastStateLengths, 2); + if (state_lengths_shape != nullptr) { + updateOutputShape(ctx, psai::kPresentStateLengths, *state_lengths_shape); + } +} + +constexpr const char* PackedSparseAttentionIndexer_ver1_doc = R"DOC( +Packed/variable-length counterpart of SparseAttentionIndexer, for continuous-batching engines +(such as an OgaEngine-style PagedAttention model) that flatten every request's tokens into one +[total_tokens, ...] axis instead of a dense [batch_size, sequence_length, ...] axis. It selects, for +every packed query token, the sparse-attention candidates that a following SparsePagedAttention (or +similar) operator is allowed to read. + +Unlike SparseAttentionIndexer, this operator: + * takes packed query/key tensors plus cumulative_sequence_lengths (request boundaries) and + past_sequence_lengths (per-request past length) instead of a dense batch and a dense mask; + * derives ordinary causal visibility purely from that packed metadata -- there is no mask input; + * uses a single generic set of state slots (past_key_state / past_kv_buffer / past_gate_buffer / + past_state_lengths) for both policy_mode values, each with a shape that is fixed across calls + (state never grows and is never concatenated); a step that would overflow the fixed capacity is + rejected as a deterministic no-op on state rather than truncated or allowed to corrupt memory; + * additionally emits selected_counts, the exact number of active (non -1) entries per query, so + that no downstream consumer needs to scan selected_indices for its query's true count. + +Both policy_mode values keep the semantics of SparseAttentionIndexer, applied independently to each +request's own packed token range and fixed-capacity state slice: + + policy_mode = "qsa" ("query sparse attention" token indexer) + Processes each request's new tokens sequentially: appends raw indexer keys to the generic + pending buffer, and whenever it reaches compress_ratio tokens, mean-pools it, applies RMSNorm + and key_norm_weight, applies the leading/split-half rotary convention at the block's first + logical token position, and appends the prepared (already normalized and rotated) key to + key_state. Queries are scored against every causally visible complete block with + sum_h ReLU(q_h . k), the token_budget / compress_ratio highest scoring blocks are kept, and + their token indices are emitted (request-local logical positions, i.e. the same numbering as + past_sequence_lengths + local offset) followed by the causally visible tokens of the trailing + incomplete block. + + policy_mode = "csa" ("compressed sparse attention" block indexer) + Applies the same window-plan arithmetic as SparseAttentionIndexer (overlap/leftover/new window + count) independently per request, using that request's own buffer_length and new token count; + every newly closed window is compressed with the softmax-gated Ca/Cb pooling, normalized, + rotated and appended to key_state. Queries are scored against every causally visible compressed + entry with sum_h w_h * ReLU(q_h . k) and the index_topk highest scoring entry indices are + emitted. + +Common contract: + * selected_indices is int32 with a fixed capacity that only depends on attributes: + token_budget + compress_ratio - 1 for "qsa" (values are request-local token positions into the + main key/value cache, directly consumable by SparsePagedAttention configured with + attention_mode="selected_only", selected_kv_source="main") and index_topk for "csa" (values are + compressed-entry indices into key_state, directly consumable by SparsePagedAttention configured + with attention_mode="local_plus_selected", selected_kv_source="auxiliary"; key_state is + layout-compatible with a [batch_size, capacity, 1, head_size] auxiliary cache when K = V). + Unused entries are -1 and selected_counts holds the exact number of used entries. + * key_norm_weight is the effective RMSNorm multiplier, exactly as in SparseAttentionIndexer. + * Accumulation, pooling, softmax, normalization and scoring are performed in float32 and the + result is rounded once to the tensor element type. + * Ties in the top-k selection are broken by the smaller entry index, and the emitted entries are + ordered by decreasing score, so the result is deterministic. + * cos_cache / sin_cache may be shared across the batch ([max_position, rotary_width]) or + request-specific ([batch_size, max_position, rotary_width]). + * cumulative_sequence_lengths, past_sequence_lengths and past_state_lengths are read directly by + the device kernel; a zero-token request row (a repeated cumulative offset) is valid and simply + contributes no query rows for that request. + +OgaEngine integration note: this operator only defines the ORT contrib op; wiring +past_key_state / past_kv_buffer / past_gate_buffer / past_state_lengths as Engine-managed, +per-request fixed-size state (analogous to a paged auxiliary cache) is expected to happen in the +OgaEngine / Model Builder integration, which is out of scope for this operator definition. +)DOC"; + +ONNX_MS_OPERATOR_SET_SCHEMA( + PackedSparseAttentionIndexer, 1, + OpSchema() + .SetDoc(PackedSparseAttentionIndexer_ver1_doc) + .Attr("policy_mode", + "Indexer policy. Must be exactly 'qsa' (token indexer) or 'csa' (compressed block indexer).", + AttributeProto::STRING) + .Attr("compress_ratio", + "Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0.", + AttributeProto::INT) + .Attr("state_capacity", + "Fixed capacity (number of entries) of past_key_state / present_key_state. Must be > 0.", + AttributeProto::INT) + .Attr("token_budget", + "Only for policy_mode 'qsa': maximum number of tokens selected from complete blocks. " + "Must be > 0 and divisible by compress_ratio. Must be omitted when policy_mode is 'csa'.", + AttributeProto::INT, + OPTIONAL_VALUE) + .Attr("index_topk", + "Only for policy_mode 'csa': number of compressed entries selected per query. Must be > 0. " + "Must be omitted when policy_mode is 'qsa'.", + AttributeProto::INT, + OPTIONAL_VALUE) + .Attr("epsilon", + "Epsilon of the RMS normalization applied to the compressed keys. Default is 1e-6.", + AttributeProto::FLOAT, + 1.0e-6f) + .Attr("scale", + "Scale applied to the per-head ReLU scores. Default is 1/sqrt(head_size).", + AttributeProto::FLOAT, + OPTIONAL_VALUE) + .Attr("head_weight_scale", + "Only for policy_mode 'csa': scale applied to head_weights. Default is 1/sqrt(num_heads). " + "Must be omitted when policy_mode is 'qsa'.", + AttributeProto::FLOAT, + OPTIONAL_VALUE) + .Input(0, + "query", + "Packed indexer queries with shape (total_tokens, num_heads, head_size), already normalized but " + "not yet rotated.", + "T") + .Input(1, + "key", + "Packed indexer key projection of the new tokens. Shape is (total_tokens, head_size) for " + "policy_mode 'qsa' and (total_tokens, 2 * head_size) for policy_mode 'csa', where the first " + "head_size channels are the Ca series and the last head_size channels the Cb series.", + "T") + .Input(2, + "key_norm_weight", + "Effective RMSNorm multiplier of the compressed keys, with shape (head_size).", + "T") + .Input(3, + "cos_cache", + "Cosine rotary table indexed by absolute key position, shared across the batch with shape " + "(max_rotary_sequence_length, rotary_width) or request-specific with shape " + "(batch_size, max_rotary_sequence_length, rotary_width).", + "T") + .Input(4, + "sin_cache", + "Sine rotary table with the same shape as cos_cache.", + "T") + .Input(5, + "cumulative_sequence_lengths", + "Device-resident packed request boundaries with shape (batch_size + 1); " + "cumulative_sequence_lengths[0] must be 0 and cumulative_sequence_lengths[batch_size] must equal " + "total_tokens. Request b owns rows [cumulative_sequence_lengths[b], " + "cumulative_sequence_lengths[b + 1]) of query/key (a repeated offset is a valid zero-token row).", + "M") + .Input(6, + "past_sequence_lengths", + "Device-resident number of tokens already processed for each request before this call, with " + "shape (batch_size). Used as the default absolute query position when position_ids is omitted " + "(policy_mode 'qsa'), and to validate state consistency.", + "M") + .Input(7, + "gate", + "Only for policy_mode 'csa': gate projection of the new tokens with shape " + "(total_tokens, 2 * head_size).", + "T", + OpSchema::Optional) + .Input(8, + "position_bias", + "Only for policy_mode 'csa': per-slot gate bias with shape (compress_ratio, 2 * head_size).", + "T", + OpSchema::Optional) + .Input(9, + "head_weights", + "Only for policy_mode 'csa': per-head score weights with shape (total_tokens, num_heads).", + "T", + OpSchema::Optional) + .Input(10, + "position_ids", + "Optional for policy_mode 'qsa', required for policy_mode 'csa': absolute position of every " + "packed query, with shape (total_tokens).", + "I", + OpSchema::Optional) + .Input(11, + "past_key_state", + "Generic fixed-capacity state: policy_mode 'qsa' stores prepared complete-block keys; " + "policy_mode 'csa' stores compressed keys. Shape is (batch_size, state_capacity, head_size) and " + "never changes across calls.", + "T") + .Input(12, + "past_kv_buffer", + "Generic fixed-capacity pending-token buffer. Shape is (batch_size, 2 * compress_ratio - 1, " + "head_size) for policy_mode 'qsa' (which only ever uses up to compress_ratio - 1 of these " + "entries) and (batch_size, 2 * compress_ratio - 1, 2 * head_size) for policy_mode 'csa'.", + "T") + .Input(13, + "past_gate_buffer", + "Only for policy_mode 'csa': buffered gate projections with the same shape as past_kv_buffer.", + "T", + OpSchema::Optional) + .Input(14, + "past_state_lengths", + "Generic per-request state length with shape (batch_size, 2). Column 0 is the key_state entry " + "count (policy_mode 'qsa': complete-block count; 'csa': compressed-entry count); column 1 is the " + "pending-buffer length (policy_mode 'qsa': incomplete-block length in [0, compress_ratio); 'csa': " + "buffer length in [0, 2 * compress_ratio)).", + "M") + .Output(0, + "selected_indices", + "Selected entries with shape (total_tokens, capacity). capacity is " + "token_budget + compress_ratio - 1 for policy_mode 'qsa' (request-local token positions into the " + "main key/value cache) and index_topk for policy_mode 'csa' (compressed entry indices into " + "key_state). Unused entries are -1.", + "M") + .Output(1, + "selected_counts", + "Exact number of used (non -1) entries of selected_indices for every query, with shape " + "(total_tokens).", + "M") + .Output(2, + "present_key_state", + "Updated generic key state, with the same fixed shape as past_key_state.", + "T") + .Output(3, + "present_kv_buffer", + "Updated generic pending-token buffer, with the same fixed shape as past_kv_buffer.", + "T") + .Output(4, + "present_gate_buffer", + "Only for policy_mode 'csa': updated gate buffer with the same fixed shape as past_gate_buffer.", + "T", + OpSchema::Optional) + .Output(5, + "present_state_lengths", + "Updated generic per-request state length, with the same fixed shape as past_state_lengths.", + "M") + .TypeConstraint("T", + {"tensor(float)", "tensor(float16)", "tensor(bfloat16)"}, + "Constrain floating point tensors to float, float16 and bfloat16.") + .TypeConstraint("I", {"tensor(int64)"}, "Constrain position ids to 64-bit integer tensors.") + .TypeConstraint("M", {"tensor(int32)"}, + "Constrain packed metadata, generic state lengths and selected indices/counts to 32-bit " + "integer tensors.") + .TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) { + PackedSparseAttentionIndexerTypeAndShapeInference(ctx); + })); + constexpr const char* Longformer_Attention_doc = R"DOC( Longformer Self Attention with a local context and a global context. Tokens attend locally: Each token attends to its W previous tokens and W succeeding tokens with W being the window length. A selected few tokens diff --git a/onnxruntime/core/graph/contrib_ops/ms_opset.h b/onnxruntime/core/graph/contrib_ops/ms_opset.h index c607b7f81cb7a..fdef7c98786f7 100644 --- a/onnxruntime/core/graph/contrib_ops/ms_opset.h +++ b/onnxruntime/core/graph/contrib_ops/ms_opset.h @@ -117,6 +117,7 @@ class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, SkipLayerNormalization); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, SkipSimplifiedLayerNormalization); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, SparseAttention); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, SparseAttentionIndexer); +class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, PackedSparseAttentionIndexer); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, SparseToDenseMatMul); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, Tokenizer); class ONNX_OPERATOR_SET_SCHEMA_CLASS_NAME(Microsoft, 1, TorchEmbedding); @@ -244,6 +245,7 @@ class OpSet_Microsoft_ver1 { fn(GetOpSchema()); fn(GetOpSchema()); fn(GetOpSchema()); + fn(GetOpSchema()); fn(GetOpSchema()); fn(GetOpSchema()); fn(GetOpSchema()); diff --git a/onnxruntime/python/tools/symbolic_shape_infer.py b/onnxruntime/python/tools/symbolic_shape_infer.py index ed2a384f30401..1bbc0e2f64c0f 100755 --- a/onnxruntime/python/tools/symbolic_shape_infer.py +++ b/onnxruntime/python/tools/symbolic_shape_infer.py @@ -238,6 +238,7 @@ def __init__(self, int_max, auto_merge, guess_output_rank, verbose, prefix=""): "SkipSimplifiedLayerNormalization": self._infer_SkipLayerNormalization, "SparseAttention": self._infer_SparseAttention, "SparseAttentionIndexer": self._infer_SparseAttentionIndexer, + "PackedSparseAttentionIndexer": self._infer_PackedSparseAttentionIndexer, "UnfoldTensor": self._infer_UnfoldTensor, } self.aten_op_dispatcher_ = { @@ -498,6 +499,7 @@ def _onnx_infer_single_node(self, node): "SkipSimplifiedLayerNormalization", "SparseAttention", "SparseAttentionIndexer", + "PackedSparseAttentionIndexer", "SkipGroupNorm", "QLinearAdd", "QLinearMul", @@ -2697,6 +2699,55 @@ def past_shape(index): set_output(3, [query_shape[0], present_buffer_length, past_buffer_shape[2]]) set_output(4, [query_shape[0], present_buffer_length, past_buffer_shape[2]]) + def _infer_PackedSparseAttentionIndexer(self, node): # noqa: N802 + policy_mode = get_attribute(node, "policy_mode", b"") + if isinstance(policy_mode, bytes): + policy_mode = policy_mode.decode("utf-8") + compress_ratio = get_attribute(node, "compress_ratio", 0) + query_shape = self._get_sympy_shape(node, 0) + total_tokens = query_shape[0] + output_dtype = self.known_vi_[node.input[0]].type.tensor_type.elem_type + + if policy_mode == "qsa": + capacity = get_attribute(node, "token_budget", 0) + compress_ratio - 1 + else: + capacity = get_attribute(node, "index_topk", 0) + + vi = self.known_vi_[node.output[0]] + vi.CopyFrom( + helper.make_tensor_value_info( + node.output[0], + onnx.TensorProto.INT32, + get_shape_from_sympy_shape([total_tokens, capacity]), + ) + ) + if len(node.output) > 1 and node.output[1]: + vi = self.known_vi_[node.output[1]] + vi.CopyFrom( + helper.make_tensor_value_info( + node.output[1], onnx.TensorProto.INT32, get_shape_from_sympy_shape([total_tokens]) + ) + ) + + def copy_state_output(output_index, input_index, dtype): + # State never grows: present_* always has exactly the same fixed shape as past_*, so + # shape inference can simply copy it, unlike the dense op's growing/concatenated state. + if output_index >= len(node.output) or not node.output[output_index]: + return + if input_index >= len(node.input) or not node.input[input_index]: + return + shape = self._get_sympy_shape(node, input_index) + out_vi = self.known_vi_[node.output[output_index]] + out_vi.CopyFrom( + helper.make_tensor_value_info(node.output[output_index], dtype, get_shape_from_sympy_shape(shape)) + ) + + copy_state_output(2, 11, output_dtype) # present_key_state <- past_key_state + copy_state_output(3, 12, output_dtype) # present_kv_buffer <- past_kv_buffer + if policy_mode != "qsa": + copy_state_output(4, 13, output_dtype) # present_gate_buffer <- past_gate_buffer + copy_state_output(5, 14, onnx.TensorProto.INT32) # present_state_lengths <- past_state_lengths + def _infer_SkipGroupNorm(self, node): # noqa: N802 self._propagate_shape_and_type(node, 0, 0) if len(node.output) > 1: diff --git a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc new file mode 100644 index 0000000000000..e6197057f623d --- /dev/null +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -0,0 +1,1044 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +// Coverage for com.microsoft.PackedSparseAttentionIndexer. +// +// The "ShapeInference" suite drives Graph::Resolve() directly, so it runs in every build: it pins +// the fixed selected_indices/selected_counts shapes, the fixed (never data-dependent) state output +// shapes, and the strict policy validation. Cases that are expected to fail shape inference call +// fail_shape_inference, which aborts in ORT_NO_EXCEPTIONS builds, so they are compiled out there. +// +// The numeric suite needs the CUDA or WebGPU execution provider (the operator has no CPU kernel) +// and is skipped when neither is available. Expectations come from a float reference in this file +// that mirrors the operator contract (mean-pool/RMSNorm/RoPE for qsa, softmax-gated pooling for +// csa, deterministic top-k selection); the inputs are rounded to the tested element type first so +// the reference sees exactly what the kernel reads. + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "gtest/gtest.h" + +#include "contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h" +#include "core/graph/constants.h" +#include "core/graph/model.h" +#include "test/common/tensor_op_test_utils.h" +#include "test/providers/provider_test_utils.h" +#include "test/test_environment.h" +#include "test/unittest_util/graph_transform_test_builder.h" +#include "test/util/include/asserts.h" +#include "test/util/include/default_providers.h" + +namespace onnxruntime { +namespace test { + +namespace psai = ::onnxruntime::contrib::packed_sparse_attention_indexer; + +namespace { + +constexpr int kOnnxOpsetVersion = 17; + +// --------------------------------------------------------------------------------------------- +// Shape inference helpers +// --------------------------------------------------------------------------------------------- + +Status BuildAndResolve(const std::function& add_node, + std::unique_ptr& model) { + std::unordered_map domain_to_version; + domain_to_version[kOnnxDomain] = kOnnxOpsetVersion; + domain_to_version[kMSDomain] = 1; + + model = std::unique_ptr(new Model("packed_sparse_attention_indexer", /*is_onnx_domain_only=*/false, + ModelMetaData(), PathString(), IOnnxRuntimeOpSchemaRegistryList(), + domain_to_version, {}, DefaultLoggingManager().DefaultLogger())); + + ModelTestBuilder builder(model->MainGraph()); + add_node(builder); + builder.SetGraphOutputs(); + return model->MainGraph().Resolve(); +} + +void ExpectShape(const Graph& graph, const std::string& name, ONNX_NAMESPACE::TensorProto_DataType elem_type, + const std::vector& expected) { + const NodeArg* arg = graph.GetNodeArg(name); + ASSERT_NE(arg, nullptr); + const ONNX_NAMESPACE::TypeProto* type = arg->TypeAsProto(); + ASSERT_NE(type, nullptr); + ASSERT_TRUE(type->has_tensor_type()); + EXPECT_EQ(type->tensor_type().elem_type(), static_cast(elem_type)); + const ONNX_NAMESPACE::TensorShapeProto& shape = type->tensor_type().shape(); + ASSERT_EQ(shape.dim_size(), static_cast(expected.size())); + for (int i = 0; i < shape.dim_size(); ++i) { + ASSERT_TRUE(shape.dim(i).has_dim_value()) << "dimension " << i << " of " << name << " is not static"; + EXPECT_EQ(shape.dim(i).dim_value(), expected[static_cast(i)]) << "dimension " << i << " of " << name; + } +} + +struct GraphOptions { + int64_t batch_size = 2; + int64_t total_tokens = 5; + int64_t num_heads = 2; + int64_t head_size = 8; + int64_t rotary_width = 8; + int64_t compress_ratio = 2; + int64_t state_capacity = 6; + int64_t token_budget = 4; + bool add_index_topk = false; + bool add_csa_inputs = false; + bool add_position_ids = false; + int output_count = psai::kFixedOutputCount; + std::string policy_mode = psai::kPolicyModeQsa; +}; + +int64_t BufferCapacity(int64_t compress_ratio) { return 2 * compress_ratio - 1; } + +// Builds a fixed 15-input node; csa-only slots are left empty for policy_mode "qsa", as the schema +// requires. position_ids (slot 10) is optional for "qsa" and forced on for "csa". +void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { + const bool is_csa = options.policy_mode == psai::kPolicyModeCsa; + const int64_t width = is_csa ? 2 * options.head_size : options.head_size; + const int64_t buffer_capacity = BufferCapacity(options.compress_ratio); + NodeArg& empty = builder.graph_.GetOrCreateNodeArg("", nullptr); + + std::vector inputs{ + builder.MakeInput( + std::vector{options.total_tokens, options.num_heads, options.head_size}), + builder.MakeInput(std::vector{options.total_tokens, width}), + builder.MakeInput(std::vector{options.head_size}), + builder.MakeInput(std::vector{64, options.rotary_width}), + builder.MakeInput(std::vector{64, options.rotary_width}), + builder.MakeInput(std::vector{options.batch_size + 1}), + builder.MakeInput(std::vector{options.batch_size}), + }; + if (is_csa || options.add_csa_inputs) { + inputs.push_back(builder.MakeInput(std::vector{options.total_tokens, width})); + inputs.push_back(builder.MakeInput(std::vector{options.compress_ratio, width})); + inputs.push_back(builder.MakeInput(std::vector{options.total_tokens, options.num_heads})); + } else { + inputs.push_back(&empty); + inputs.push_back(&empty); + inputs.push_back(&empty); + } + if (is_csa || options.add_position_ids) { + inputs.push_back(builder.MakeInput(std::vector{options.total_tokens})); + } else { + inputs.push_back(&empty); + } + inputs.push_back( + builder.MakeInput(std::vector{options.batch_size, options.state_capacity, options.head_size})); + inputs.push_back( + builder.MakeInput(std::vector{options.batch_size, buffer_capacity, width})); + if (is_csa) { + inputs.push_back(builder.MakeInput(std::vector{options.batch_size, buffer_capacity, width})); + } else { + inputs.push_back(&empty); + } + inputs.push_back(builder.MakeInput(std::vector{options.batch_size, 2})); + + std::vector outputs; + for (int i = 0; i < options.output_count; ++i) { + outputs.push_back(i == psai::kPresentGateBuffer && !is_csa ? &empty : builder.MakeOutput()); + } + Node& node = builder.AddNode("PackedSparseAttentionIndexer", inputs, outputs, kMSDomain); + node.AddAttribute("policy_mode", options.policy_mode); + node.AddAttribute("compress_ratio", options.compress_ratio); + node.AddAttribute("state_capacity", options.state_capacity); + if (is_csa || options.add_index_topk) { + node.AddAttribute("index_topk", static_cast(3)); + } else { + node.AddAttribute("token_budget", options.token_budget); + } +} + +} // namespace + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, QsaInfersFixedCapacityAndState) { + GraphOptions options; + std::unique_ptr model; + ASSERT_STATUS_OK(BuildAndResolve([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, model)); + + const Graph& graph = model->MainGraph(); + const Node& node = *graph.Nodes().begin(); + const int64_t capacity = options.token_budget + options.compress_ratio - 1; + ExpectShape(graph, node.OutputDefs()[psai::kSelectedIndices]->Name(), ONNX_NAMESPACE::TensorProto_DataType_INT32, + {options.total_tokens, capacity}); + ExpectShape(graph, node.OutputDefs()[psai::kSelectedCounts]->Name(), ONNX_NAMESPACE::TensorProto_DataType_INT32, + {options.total_tokens}); + ExpectShape(graph, node.OutputDefs()[psai::kPresentKeyState]->Name(), ONNX_NAMESPACE::TensorProto_DataType_FLOAT, + {options.batch_size, options.state_capacity, options.head_size}); + ExpectShape(graph, node.OutputDefs()[psai::kPresentKvBuffer]->Name(), ONNX_NAMESPACE::TensorProto_DataType_FLOAT, + {options.batch_size, BufferCapacity(options.compress_ratio), options.head_size}); + ExpectShape(graph, node.OutputDefs()[psai::kPresentStateLengths]->Name(), + ONNX_NAMESPACE::TensorProto_DataType_INT32, {options.batch_size, 2}); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, CsaInfersFixedCapacityAndState) { + GraphOptions options; + options.policy_mode = psai::kPolicyModeCsa; + std::unique_ptr model; + ASSERT_STATUS_OK(BuildAndResolve([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, model)); + + const Graph& graph = model->MainGraph(); + const Node& node = *graph.Nodes().begin(); + ExpectShape(graph, node.OutputDefs()[psai::kSelectedIndices]->Name(), ONNX_NAMESPACE::TensorProto_DataType_INT32, + {options.total_tokens, 3}); + ExpectShape(graph, node.OutputDefs()[psai::kSelectedCounts]->Name(), ONNX_NAMESPACE::TensorProto_DataType_INT32, + {options.total_tokens}); + ExpectShape(graph, node.OutputDefs()[psai::kPresentKeyState]->Name(), ONNX_NAMESPACE::TensorProto_DataType_FLOAT, + {options.batch_size, options.state_capacity, options.head_size}); + ExpectShape(graph, node.OutputDefs()[psai::kPresentKvBuffer]->Name(), ONNX_NAMESPACE::TensorProto_DataType_FLOAT, + {options.batch_size, BufferCapacity(options.compress_ratio), 2 * options.head_size}); + ExpectShape(graph, node.OutputDefs()[psai::kPresentGateBuffer]->Name(), ONNX_NAMESPACE::TensorProto_DataType_FLOAT, + {options.batch_size, BufferCapacity(options.compress_ratio), 2 * options.head_size}); + ExpectShape(graph, node.OutputDefs()[psai::kPresentStateLengths]->Name(), + ONNX_NAMESPACE::TensorProto_DataType_INT32, {options.batch_size, 2}); +} + +#ifndef ORT_NO_EXCEPTIONS + +namespace { + +void ExpectResolveFailure(const std::function& add_node, + const std::string& expected_message) { + std::unique_ptr model; + const Status status = BuildAndResolve(add_node, model); + ASSERT_FALSE(status.IsOK()) << "expected shape inference to reject the node"; + EXPECT_NE(status.ErrorMessage().find(expected_message), std::string::npos) + << "actual message: " << status.ErrorMessage(); +} + +} // namespace + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsUnknownPolicyMode) { + GraphOptions options; + options.policy_mode = "qsa_v2"; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "policy_mode must be 'qsa' or 'csa'"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsQsaWithCsaAttribute) { + GraphOptions options; + options.add_index_topk = true; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "index_topk and head_weight_scale must not be set"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsZeroNumHeads) { + GraphOptions options; + options.num_heads = 0; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "num_heads must be > 0"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsQsaTokenBudgetNotDivisibleByCompressRatio) { + GraphOptions options; + options.token_budget = 5; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "requires token_budget > 0, divisible by compress_ratio"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsQsaWithCsaInput) { + GraphOptions options; + options.add_csa_inputs = true; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "must be omitted when policy_mode is 'qsa'"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsQsaWithPositionIdsAndCsaInputsMismatch) { + // position_ids alone is allowed for qsa (optional); only the csa-only slots must be omitted. + GraphOptions options; + options.add_position_ids = true; + std::unique_ptr model; + ASSERT_STATUS_OK(BuildAndResolve([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, model)); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsCsaMissingPositionIds) { + GraphOptions options; + options.policy_mode = psai::kPolicyModeCsa; + options.add_position_ids = false; + // Force the csa node to omit position_ids by rebuilding without it. + std::unique_ptr model; + const Status status = BuildAndResolve( + [&options](ModelTestBuilder& builder) { + GraphOptions local = options; + NodeArg& empty = builder.graph_.GetOrCreateNodeArg("", nullptr); + const int64_t width = 2 * local.head_size; + const int64_t buffer_capacity = BufferCapacity(local.compress_ratio); + std::vector inputs{ + builder.MakeInput(std::vector{local.total_tokens, local.num_heads, local.head_size}), + builder.MakeInput(std::vector{local.total_tokens, width}), + builder.MakeInput(std::vector{local.head_size}), + builder.MakeInput(std::vector{64, local.rotary_width}), + builder.MakeInput(std::vector{64, local.rotary_width}), + builder.MakeInput(std::vector{local.batch_size + 1}), + builder.MakeInput(std::vector{local.batch_size}), + builder.MakeInput(std::vector{local.total_tokens, width}), + builder.MakeInput(std::vector{local.compress_ratio, width}), + builder.MakeInput(std::vector{local.total_tokens, local.num_heads}), + &empty, // position_ids omitted: invalid for csa + builder.MakeInput( + std::vector{local.batch_size, local.state_capacity, local.head_size}), + builder.MakeInput(std::vector{local.batch_size, buffer_capacity, width}), + builder.MakeInput(std::vector{local.batch_size, buffer_capacity, width}), + builder.MakeInput(std::vector{local.batch_size, 2}), + }; + std::vector outputs; + for (int i = 0; i < psai::kFixedOutputCount; ++i) { + outputs.push_back(builder.MakeOutput()); + } + Node& node = builder.AddNode("PackedSparseAttentionIndexer", inputs, outputs, kMSDomain); + node.AddAttribute("policy_mode", local.policy_mode); + node.AddAttribute("compress_ratio", local.compress_ratio); + node.AddAttribute("state_capacity", local.state_capacity); + node.AddAttribute("index_topk", static_cast(3)); + }, + model); + ASSERT_FALSE(status.IsOK()); + EXPECT_NE(status.ErrorMessage().find("position_ids"), std::string::npos) << status.ErrorMessage(); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsWrongOutputCount) { + GraphOptions options; + options.output_count = 4; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "exactly 6 declared outputs"); +} + +#endif // ORT_NO_EXCEPTIONS + +// --------------------------------------------------------------------------------------------- +// Numeric behaviour (CUDA / WebGPU only) +// --------------------------------------------------------------------------------------------- + +namespace { + +enum class ProviderKind { + Cuda, + WebGpu, +}; + +std::unique_ptr CreateProvider(ProviderKind provider_kind) { + if (provider_kind == ProviderKind::Cuda) { + return DefaultCudaExecutionProvider(); + } +#ifdef USE_WEBGPU + return DefaultWebGpuExecutionProvider(); +#else + return nullptr; +#endif +} + +void RunOnProvider(OpTester& test, std::unique_ptr provider) { + std::vector> providers; + providers.push_back(std::move(provider)); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &providers); +} + +std::vector MakeWave(size_t count, float phase, float step) { + std::vector values(count); + for (size_t i = 0; i < count; ++i) { + values[i] = std::sin(phase + step * static_cast(i)); + } + return values; +} + +template +std::vector ToElementType(const std::vector& data) { + if constexpr (std::is_same_v) { + return ToFloat16(data); + } else if constexpr (std::is_same_v) { + return ToBFloat16(data); + } else { + return data; + } +} + +// Rounds through the tested element type so the reference consumes exactly the kernel's inputs. +template +std::vector RoundTrip(const std::vector& data) { + if constexpr (std::is_same_v) { + return data; + } else { + std::vector converted = ToElementType(data); + std::vector result(data.size()); + for (size_t i = 0; i < data.size(); ++i) { + result[i] = converted[i].ToFloat(); + } + return result; + } +} + +// Split-half rotary over the leading rotary_width channels. +std::vector LeadingRope(const std::vector& value, int rotary_width, const float* cos_row, + const float* sin_row) { + const int head_size = static_cast(value.size()); + const int half = rotary_width / 2; + std::vector result(value); + for (int d = 0; d < rotary_width && d < head_size; ++d) { + const float paired = (d < half) ? -value[static_cast(d + half)] : value[static_cast(d - half)]; + result[static_cast(d)] = value[static_cast(d)] * cos_row[d] + paired * sin_row[d]; + } + return result; +} + +// Interleaved rotary over the trailing 2 * rotary_width channels. +std::vector TrailingRope(const std::vector& value, int rotary_width, const float* cos_row, + const float* sin_row) { + const int head_size = static_cast(value.size()); + const int base = head_size - 2 * rotary_width; + std::vector result(value); + for (int d = base; d < head_size; ++d) { + const int offset = d - base; + const float paired = ((offset & 1) == 0) ? -value[static_cast(d + 1)] : value[static_cast(d - 1)]; + result[static_cast(d)] = + value[static_cast(d)] * cos_row[offset >> 1] + paired * sin_row[offset >> 1]; + } + return result; +} + +std::vector RmsNormalize(const std::vector& value, const std::vector& weight, float epsilon) { + float sum_squares = 0.0f; + for (float element : value) { + sum_squares += element * element; + } + const float inverse_rms = 1.0f / std::sqrt(sum_squares / static_cast(value.size()) + epsilon); + std::vector result(value.size()); + for (size_t d = 0; d < value.size(); ++d) { + result[d] = value[d] * inverse_rms * weight[d]; + } + return result; +} + +// Order used by both selection kernels: score descending, then entry index ascending. +std::vector RankByScore(const std::vector& scores, int count) { + std::vector order(static_cast(count)); + std::iota(order.begin(), order.end(), 0); + std::stable_sort(order.begin(), order.end(), [&scores](int left, int right) { + if (scores[static_cast(left)] != scores[static_cast(right)]) { + return scores[static_cast(left)] > scores[static_cast(right)]; + } + return left < right; + }); + return order; +} + +int CausalThreshold(int position, int compress_ratio) { + if (position < 0) return 0; + return position / compress_ratio + (position % compress_ratio == compress_ratio - 1 ? 1 : 0); +} + +struct QsaPackedProblem { + int batch_size = 2; + std::vector cumulative_sequence_lengths; // [batch + 1] + std::vector past_sequence_lengths; // [batch] + std::vector past_state_lengths; // [batch, 2] + int head_size = 4; + int num_heads = 2; + int rotary_width = 4; + int compress_ratio = 2; + int token_budget = 4; + int state_capacity = 4; + int max_position = 64; + float epsilon = 1.0e-6f; + std::optional scale; + + std::vector query; + std::vector key; + std::vector key_norm_weight; + std::vector cos_cache; // shared: [max_position, rotary_width] + std::vector sin_cache; + std::vector past_key_state; // [batch, state_capacity, head_size] + std::vector past_kv_buffer; // [batch, buffer_capacity, head_size] + + int TotalTokens() const { return cumulative_sequence_lengths.back(); } + int BufferCapacity() const { return 2 * compress_ratio - 1; } + int Capacity() const { return token_budget + compress_ratio - 1; } +}; + +struct QsaPackedResult { + std::vector selected_indices; + std::vector selected_counts; + std::vector present_key_state; + std::vector present_kv_buffer; + std::vector present_state_lengths; +}; + +void QsaPackedReference(const QsaPackedProblem& p, QsaPackedResult& out) { + const int total_tokens = p.TotalTokens(); + const int head_size = p.head_size; + const int block_topk = p.token_budget / p.compress_ratio; + const float scale = p.scale.value_or(1.0f / std::sqrt(static_cast(head_size))); + const int buffer_capacity = p.BufferCapacity(); + const int capacity = p.Capacity(); + + out.present_key_state = p.past_key_state; + out.present_kv_buffer.assign(static_cast(p.batch_size) * buffer_capacity * head_size, 0.0f); + out.present_state_lengths.assign(static_cast(p.batch_size) * 2, 0); + + std::vector key_len_after(static_cast(p.batch_size)); + for (int b = 0; b < p.batch_size; ++b) { + const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; + const int req_end = p.cumulative_sequence_lengths[static_cast(b) + 1]; + const int req_len = std::max(req_end - req_start, 0); + const int old_key_len = + std::clamp(p.past_state_lengths[static_cast(b) * 2 + 0], 0, p.state_capacity); + const int old_buf_len = + std::clamp(p.past_state_lengths[static_cast(b) * 2 + 1], 0, p.compress_ratio - 1); + const int pending = old_buf_len + req_len; + const int new_block_count = std::min(pending / p.compress_ratio, std::max(p.state_capacity - old_key_len, 0)); + const int new_buf_len = new_block_count < pending / p.compress_ratio ? 0 : pending % p.compress_ratio; + + auto raw_value = [&](int virtual_pos, int d) -> float { + return virtual_pos < old_buf_len + ? p.past_kv_buffer[(static_cast(b) * buffer_capacity + virtual_pos) * head_size + d] + : p.key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * head_size + d]; + }; + + for (int k = 0; k < new_block_count; ++k) { + std::vector pooled(static_cast(head_size), 0.0f); + for (int t = 0; t < p.compress_ratio; ++t) { + for (int d = 0; d < head_size; ++d) { + pooled[static_cast(d)] += raw_value(k * p.compress_ratio + t, d); + } + } + for (float& v : pooled) v /= static_cast(p.compress_ratio); + pooled = RmsNormalize(pooled, p.key_norm_weight, p.epsilon); + const int entry = old_key_len + k; + const int position = std::min(entry * p.compress_ratio, p.max_position - 1); + pooled = LeadingRope(pooled, p.rotary_width, p.cos_cache.data() + position * p.rotary_width, + p.sin_cache.data() + position * p.rotary_width); + for (int d = 0; d < head_size; ++d) { + out.present_key_state[(static_cast(b) * p.state_capacity + entry) * head_size + d] = + pooled[static_cast(d)]; + } + } + for (int t = 0; t < new_buf_len; ++t) { + for (int d = 0; d < head_size; ++d) { + out.present_kv_buffer[(static_cast(b) * buffer_capacity + t) * head_size + d] = + raw_value(new_block_count * p.compress_ratio + t, d); + } + } + out.present_state_lengths[static_cast(b) * 2 + 0] = old_key_len + new_block_count; + out.present_state_lengths[static_cast(b) * 2 + 1] = new_buf_len; + key_len_after[static_cast(b)] = old_key_len + new_block_count; + } + + out.selected_indices.assign(static_cast(total_tokens) * capacity, -1); + out.selected_counts.assign(static_cast(total_tokens), 0); + for (int b = 0; b < p.batch_size; ++b) { + const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; + const int req_end = p.cumulative_sequence_lengths[static_cast(b) + 1]; + for (int token = req_start; token < req_end; ++token) { + const int position = p.past_sequence_lengths[static_cast(b)] + (token - req_start); + const int causal_count = CausalThreshold(position, p.compress_ratio); + const int visible_block_count = std::min(key_len_after[static_cast(b)], causal_count); + const int selected = std::min(block_topk, visible_block_count); + + std::vector> rotated_query(static_cast(p.num_heads)); + for (int h = 0; h < p.num_heads; ++h) { + const size_t base = (static_cast(token) * p.num_heads + h) * head_size; + std::vector head(p.query.begin() + base, p.query.begin() + base + head_size); + const int clamped_position = std::min(std::max(position, 0), p.max_position - 1); + rotated_query[static_cast(h)] = + LeadingRope(head, p.rotary_width, p.cos_cache.data() + clamped_position * p.rotary_width, + p.sin_cache.data() + clamped_position * p.rotary_width); + } + + std::vector scores(static_cast(visible_block_count), 0.0f); + for (int j = 0; j < visible_block_count; ++j) { + float score = 0.0f; + for (int h = 0; h < p.num_heads; ++h) { + float dot = 0.0f; + for (int d = 0; d < head_size; ++d) { + dot += rotated_query[static_cast(h)][static_cast(d)] * + out.present_key_state[(static_cast(b) * p.state_capacity + j) * head_size + d]; + } + score += std::max(dot, 0.0f); + } + scores[static_cast(j)] = score * scale; + } + const std::vector order = RankByScore(scores, visible_block_count); + const int emitted_blocks = std::min(selected, visible_block_count); + int32_t* out_row = out.selected_indices.data() + static_cast(token) * capacity; + for (int rank = 0; rank < emitted_blocks; ++rank) { + for (int t = 0; t < p.compress_ratio; ++t) { + out_row[rank * p.compress_ratio + t] = order[static_cast(rank)] * p.compress_ratio + t; + } + } + const int block_start = visible_block_count * p.compress_ratio; + const int natural_tail = position >= block_start ? position - block_start + 1 : 0; + const int remaining_capacity = capacity - emitted_blocks * p.compress_ratio; + const int tail_count = std::clamp(natural_tail, 0, remaining_capacity); + for (int t = 0; t < tail_count; ++t) { + out_row[emitted_blocks * p.compress_ratio + t] = block_start + t; + } + out.selected_counts[static_cast(token)] = emitted_blocks * p.compress_ratio + tail_count; + } + } +} + +QsaPackedProblem MakeQsaPackedProblem(QsaPackedProblem problem = {}) { + if (problem.cumulative_sequence_lengths.empty()) { + // Default: 2 requests, 2 and 3 tokens respectively (unequal lengths). + problem.cumulative_sequence_lengths = {0, 2, 5}; + problem.past_sequence_lengths = {3, 0}; + } + const int total_tokens = problem.TotalTokens(); + const int buffer_capacity = problem.BufferCapacity(); + if (problem.past_state_lengths.empty()) { + problem.past_state_lengths.assign(static_cast(problem.batch_size) * 2, 0); + for (int b = 0; b < problem.batch_size; ++b) { + problem.past_state_lengths[static_cast(b) * 2 + 0] = problem.past_sequence_lengths[static_cast(b)] / problem.compress_ratio; + problem.past_state_lengths[static_cast(b) * 2 + 1] = problem.past_sequence_lengths[static_cast(b)] % problem.compress_ratio; + } + } + + problem.query = + MakeWave(static_cast(total_tokens) * problem.num_heads * problem.head_size, 0.35f, 0.41f); + problem.key = MakeWave(static_cast(total_tokens) * problem.head_size, 1.10f, 0.29f); + problem.key_norm_weight = MakeWave(static_cast(problem.head_size), 0.70f, 0.17f); + problem.cos_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 0.20f, 0.13f); + problem.sin_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 0.90f, 0.19f); + problem.past_key_state = + MakeWave(static_cast(problem.batch_size) * problem.state_capacity * problem.head_size, 0.05f, 0.23f); + problem.past_kv_buffer = + MakeWave(static_cast(problem.batch_size) * buffer_capacity * problem.head_size, 0.60f, 0.09f); + return problem; +} + +template +void RunQsaPackedTest(float tolerance, QsaPackedProblem problem = MakeQsaPackedProblem(), + ProviderKind provider_kind = ProviderKind::Cuda) { + auto provider = CreateProvider(provider_kind); + if (provider == nullptr) { + GTEST_SKIP() << (provider_kind == ProviderKind::Cuda ? "CUDA" : "WebGPU") + << " execution provider is not available"; + } + + problem.query = RoundTrip(problem.query); + problem.key = RoundTrip(problem.key); + problem.key_norm_weight = RoundTrip(problem.key_norm_weight); + problem.cos_cache = RoundTrip(problem.cos_cache); + problem.sin_cache = RoundTrip(problem.sin_cache); + problem.past_key_state = RoundTrip(problem.past_key_state); + problem.past_kv_buffer = RoundTrip(problem.past_kv_buffer); + + QsaPackedResult expected; + QsaPackedReference(problem, expected); + + const int64_t total_tokens = problem.TotalTokens(); + const int64_t batch_size = problem.batch_size; + const int64_t head_size = problem.head_size; + const int64_t buffer_capacity = problem.BufferCapacity(); + + OpTester test("PackedSparseAttentionIndexer", 1, onnxruntime::kMSDomain); + test.AddAttribute("policy_mode", std::string(psai::kPolicyModeQsa)); + test.AddAttribute("compress_ratio", static_cast(problem.compress_ratio)); + test.AddAttribute("state_capacity", static_cast(problem.state_capacity)); + test.AddAttribute("token_budget", static_cast(problem.token_budget)); + if (problem.scale.has_value()) { + test.AddAttribute("scale", *problem.scale); + } + test.AddInput("query", {total_tokens, problem.num_heads, head_size}, ToElementType(problem.query)); + test.AddInput("key", {total_tokens, head_size}, ToElementType(problem.key)); + test.AddInput("key_norm_weight", {head_size}, ToElementType(problem.key_norm_weight)); + test.AddInput("cos_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.cos_cache)); + test.AddInput("sin_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.sin_cache)); + test.AddInput("cumulative_sequence_lengths", {batch_size + 1}, problem.cumulative_sequence_lengths); + test.AddInput("past_sequence_lengths", {batch_size}, problem.past_sequence_lengths); + test.AddOptionalInputEdge(); // gate + test.AddOptionalInputEdge(); // position_bias + test.AddOptionalInputEdge(); // head_weights + test.AddOptionalInputEdge(); // position_ids + test.AddInput("past_key_state", {batch_size, problem.state_capacity, head_size}, + ToElementType(problem.past_key_state)); + test.AddInput("past_kv_buffer", {batch_size, buffer_capacity, head_size}, + ToElementType(problem.past_kv_buffer)); + test.AddOptionalInputEdge(); // past_gate_buffer + test.AddInput("past_state_lengths", {batch_size, 2}, problem.past_state_lengths); + + test.AddOutput("selected_indices", {total_tokens, problem.Capacity()}, expected.selected_indices); + test.AddOutput("selected_counts", {total_tokens}, expected.selected_counts); + test.AddOutput("present_key_state", {batch_size, problem.state_capacity, head_size}, + ToElementType(expected.present_key_state), false, 0.0f, tolerance); + test.AddOutput("present_kv_buffer", {batch_size, buffer_capacity, head_size}, + ToElementType(expected.present_kv_buffer), false, 0.0f, tolerance); + test.AddOptionalOutputEdge(); // present_gate_buffer + test.AddOutput("present_state_lengths", {batch_size, 2}, expected.present_state_lengths); + RunOnProvider(test, std::move(provider)); +} + +} // namespace + +TEST(PackedSparseAttentionIndexerTest, QsaFloat) { RunQsaPackedTest(1.0e-5f); } + +TEST(PackedSparseAttentionIndexerTest, QsaFloat16) { RunQsaPackedTest(2.0e-3f); } + +TEST(PackedSparseAttentionIndexerTest, QsaBFloat16) { RunQsaPackedTest(2.0e-2f); } + +// Prefill followed by decode: request 0 continues an existing 3-token history, request 1 starts +// fresh; the two requests keep independent block/tail state. +TEST(PackedSparseAttentionIndexerTest, QsaPrefillThenDecodeIndependentState) { + QsaPackedProblem problem; + problem.cumulative_sequence_lengths = {0, 1, 4}; + problem.past_sequence_lengths = {5, 0}; + RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(std::move(problem))); +} + +// A zero-token request row (repeated cumulative offset) must not affect the other request and +// must not read out of bounds. +TEST(PackedSparseAttentionIndexerTest, QsaZeroTokenRequestRow) { + QsaPackedProblem problem; + problem.batch_size = 3; + problem.cumulative_sequence_lengths = {0, 2, 2, 5}; + problem.past_sequence_lengths = {3, 7, 0}; + RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(std::move(problem))); +} + +// state_capacity smaller than what the incoming tokens would naturally produce: the update must +// deterministically drop the overflowing blocks rather than corrupt memory. +TEST(PackedSparseAttentionIndexerTest, QsaStateCapacityOverflowIsSafe) { + QsaPackedProblem problem; + problem.cumulative_sequence_lengths = {0, 6}; + problem.past_sequence_lengths = {4}; + problem.batch_size = 1; + problem.state_capacity = 2; // only room for 2 - 2(already used) = 0 new complete blocks + problem.past_state_lengths = {2, 0}; + RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(std::move(problem))); +} + +namespace { + +struct CsaPackedProblem { + int batch_size = 2; + std::vector cumulative_sequence_lengths; + std::vector past_sequence_lengths; // unused by csa math, kept for symmetry with qsa + std::vector past_state_lengths; // [batch, 2]: compressed count, buffer length + int head_size = 4; + int num_heads = 2; + int rotary_width = 2; + int compress_ratio = 2; + int index_topk = 3; + int state_capacity = 4; + int max_position = 64; + float epsilon = 1.0e-6f; + std::optional scale; + std::optional head_weight_scale; + + std::vector query; + std::vector key; + std::vector key_norm_weight; + std::vector cos_cache; + std::vector sin_cache; + std::vector gate; + std::vector position_bias; + std::vector head_weights; + std::vector position_ids; + std::vector past_key_state; + std::vector past_kv_buffer; + std::vector past_gate_buffer; + + int Width() const { return 2 * head_size; } + int TotalTokens() const { return cumulative_sequence_lengths.back(); } + int BufferCapacity() const { return 2 * compress_ratio - 1; } +}; + +struct CsaPackedResult { + std::vector selected_indices; + std::vector selected_counts; + std::vector present_key_state; + std::vector present_kv_buffer; + std::vector present_gate_buffer; + std::vector present_state_lengths; +}; + +void CsaPackedReference(const CsaPackedProblem& p, CsaPackedResult& out) { + const int total_tokens = p.TotalTokens(); + const int head_size = p.head_size; + const int width = p.Width(); + const int buffer_capacity = p.BufferCapacity(); + const int capacity = p.index_topk; + const float scale = p.scale.value_or(1.0f / std::sqrt(static_cast(head_size))); + const float head_weight_scale = p.head_weight_scale.value_or(1.0f / std::sqrt(static_cast(p.num_heads))); + + out.present_key_state = p.past_key_state; + out.present_kv_buffer.assign(static_cast(p.batch_size) * buffer_capacity * width, 0.0f); + out.present_gate_buffer.assign(static_cast(p.batch_size) * buffer_capacity * width, 0.0f); + out.present_state_lengths.assign(static_cast(p.batch_size) * 2, 0); + std::vector key_len_after(static_cast(p.batch_size)); + + for (int b = 0; b < p.batch_size; ++b) { + const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; + const int req_end = p.cumulative_sequence_lengths[static_cast(b) + 1]; + const int req_len = std::max(req_end - req_start, 0); + const int old_key_len = + std::clamp(p.past_state_lengths[static_cast(b) * 2 + 0], 0, p.state_capacity); + const int old_buf_len = + std::clamp(p.past_state_lengths[static_cast(b) * 2 + 1], 0, buffer_capacity); + const int overlap_length = old_buf_len >= p.compress_ratio ? p.compress_ratio : 0; + const int leftover_length = old_buf_len - overlap_length; + const int pending = leftover_length + req_len; + const int full_new_window_count = pending / p.compress_ratio; + const int new_window_count = + std::min(full_new_window_count, std::max(p.state_capacity - old_key_len, 0)); + const bool overflowed = new_window_count < full_new_window_count; + int present_buffer_length = 0; + int present_buffer_start = 0; + if (!overflowed) { + if (new_window_count > 0) { + present_buffer_length = p.compress_ratio + pending % p.compress_ratio; + present_buffer_start = overlap_length + (new_window_count - 1) * p.compress_ratio; + } else { + present_buffer_length = old_buf_len + req_len; + present_buffer_start = 0; + } + present_buffer_length = std::min(present_buffer_length, buffer_capacity); + } + + auto extended_key = [&](int virtual_pos, int channel) -> float { + return virtual_pos < old_buf_len + ? p.past_kv_buffer[(static_cast(b) * buffer_capacity + virtual_pos) * width + channel] + : p.key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + channel]; + }; + auto extended_gate = [&](int virtual_pos, int channel) -> float { + return virtual_pos < old_buf_len + ? p.past_gate_buffer[(static_cast(b) * buffer_capacity + virtual_pos) * width + channel] + : p.gate[(static_cast(req_start) + (virtual_pos - old_buf_len)) * width + channel]; + }; + + for (int k = 0; k < new_window_count; ++k) { + const bool has_previous = k >= 1 || overlap_length >= p.compress_ratio; + const int previous_base = overlap_length + (k - 1) * p.compress_ratio; + const int current_base = overlap_length + k * p.compress_ratio; + std::vector pooled(static_cast(head_size), 0.0f); + for (int d = 0; d < head_size; ++d) { + float max_gate = -std::numeric_limits::infinity(); + if (has_previous) { + for (int slot = 0; slot < p.compress_ratio; ++slot) { + max_gate = std::max(max_gate, extended_gate(previous_base + slot, d) + p.position_bias[static_cast(slot) * width + d]); + } + } + for (int slot = 0; slot < p.compress_ratio; ++slot) { + max_gate = std::max(max_gate, extended_gate(current_base + slot, head_size + d) + + p.position_bias[static_cast(slot) * width + head_size + d]); + } + float denom = 0.0f, acc = 0.0f; + if (has_previous) { + for (int slot = 0; slot < p.compress_ratio; ++slot) { + const float logit = extended_gate(previous_base + slot, d) + p.position_bias[static_cast(slot) * width + d]; + const float w = std::exp(logit - max_gate); + denom += w; + acc += w * extended_key(previous_base + slot, d); + } + } + for (int slot = 0; slot < p.compress_ratio; ++slot) { + const float logit = extended_gate(current_base + slot, head_size + d) + + p.position_bias[static_cast(slot) * width + head_size + d]; + const float w = std::exp(logit - max_gate); + denom += w; + acc += w * extended_key(current_base + slot, head_size + d); + } + pooled[static_cast(d)] = denom > 0.0f ? acc / denom : 0.0f; + } + pooled = RmsNormalize(pooled, p.key_norm_weight, p.epsilon); + const int entry = old_key_len + k; + const int position = std::min(entry * p.compress_ratio, p.max_position - 1); + pooled = TrailingRope(pooled, p.rotary_width, p.cos_cache.data() + position * p.rotary_width, + p.sin_cache.data() + position * p.rotary_width); + for (int d = 0; d < head_size; ++d) { + out.present_key_state[(static_cast(b) * p.state_capacity + entry) * head_size + d] = + pooled[static_cast(d)]; + } + } + for (int t = 0; t < present_buffer_length; ++t) { + const int virtual_pos = present_buffer_start + t; + for (int c = 0; c < width; ++c) { + out.present_kv_buffer[(static_cast(b) * buffer_capacity + t) * width + c] = extended_key(virtual_pos, c); + out.present_gate_buffer[(static_cast(b) * buffer_capacity + t) * width + c] = extended_gate(virtual_pos, c); + } + } + out.present_state_lengths[static_cast(b) * 2 + 0] = old_key_len + new_window_count; + out.present_state_lengths[static_cast(b) * 2 + 1] = present_buffer_length; + key_len_after[static_cast(b)] = old_key_len + new_window_count; + } + + out.selected_indices.assign(static_cast(total_tokens) * capacity, -1); + out.selected_counts.assign(static_cast(total_tokens), 0); + for (int b = 0; b < p.batch_size; ++b) { + const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; + const int req_end = p.cumulative_sequence_lengths[static_cast(b) + 1]; + for (int token = req_start; token < req_end; ++token) { + const int position = static_cast(p.position_ids[static_cast(token)]); + const int threshold = std::min(CausalThreshold(position, p.compress_ratio), key_len_after[static_cast(b)]); + const int selected = std::min(p.index_topk, threshold); + + std::vector> rotated_query(static_cast(p.num_heads)); + for (int h = 0; h < p.num_heads; ++h) { + const size_t base = (static_cast(token) * p.num_heads + h) * head_size; + std::vector head(p.query.begin() + base, p.query.begin() + base + head_size); + const int clamped_position = std::min(std::max(position, 0), p.max_position - 1); + rotated_query[static_cast(h)] = + TrailingRope(head, p.rotary_width, p.cos_cache.data() + clamped_position * p.rotary_width, + p.sin_cache.data() + clamped_position * p.rotary_width); + } + + std::vector scores(static_cast(threshold), 0.0f); + for (int e = 0; e < threshold; ++e) { + float score = 0.0f; + for (int h = 0; h < p.num_heads; ++h) { + float dot = 0.0f; + for (int d = 0; d < head_size; ++d) { + dot += rotated_query[static_cast(h)][static_cast(d)] * + out.present_key_state[(static_cast(b) * p.state_capacity + e) * head_size + d]; + } + score += std::max(dot, 0.0f) * p.head_weights[static_cast(token) * p.num_heads + h]; + } + scores[static_cast(e)] = score * scale * head_weight_scale; + } + const std::vector order = RankByScore(scores, threshold); + int32_t* out_row = out.selected_indices.data() + static_cast(token) * capacity; + for (int rank = 0; rank < selected; ++rank) { + out_row[rank] = order[static_cast(rank)]; + } + out.selected_counts[static_cast(token)] = selected; + } + } +} + +CsaPackedProblem MakeCsaPackedProblem(CsaPackedProblem problem = {}) { + if (problem.cumulative_sequence_lengths.empty()) { + problem.cumulative_sequence_lengths = {0, 2, 5}; + } + if (problem.past_state_lengths.empty()) { + problem.past_state_lengths.assign(static_cast(problem.batch_size) * 2, 0); + } + if (problem.past_sequence_lengths.empty()) { + problem.past_sequence_lengths.assign(static_cast(problem.batch_size), 0); + } + const int total_tokens = problem.TotalTokens(); + const int width = problem.Width(); + const int buffer_capacity = problem.BufferCapacity(); + + problem.query = MakeWave(static_cast(total_tokens) * problem.num_heads * problem.head_size, 0.25f, 0.37f); + problem.key = MakeWave(static_cast(total_tokens) * width, 0.60f, 0.21f); + problem.key_norm_weight = MakeWave(static_cast(problem.head_size), 0.45f, 0.31f); + problem.cos_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 0.15f, 0.27f); + problem.sin_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 1.05f, 0.33f); + problem.gate = MakeWave(static_cast(total_tokens) * width, 0.80f, 0.24f); + problem.position_bias = MakeWave(static_cast(problem.compress_ratio) * width, 0.33f, 0.11f); + problem.head_weights = MakeWave(static_cast(total_tokens) * problem.num_heads, 1.30f, 0.47f); + problem.past_key_state = + MakeWave(static_cast(problem.batch_size) * problem.state_capacity * problem.head_size, 0.50f, 0.39f); + problem.past_kv_buffer = MakeWave(static_cast(problem.batch_size) * buffer_capacity * width, 0.95f, 0.18f); + problem.past_gate_buffer = MakeWave(static_cast(problem.batch_size) * buffer_capacity * width, 1.45f, 0.22f); + + if (problem.position_ids.empty()) { + problem.position_ids.assign(static_cast(total_tokens), 0); + for (int b = 0; b < problem.batch_size; ++b) { + const int req_start = problem.cumulative_sequence_lengths[static_cast(b)]; + const int req_end = problem.cumulative_sequence_lengths[static_cast(b) + 1]; + for (int token = req_start; token < req_end; ++token) { + problem.position_ids[static_cast(token)] = 2 + (token - req_start); + } + } + } + return problem; +} + +template +void RunCsaPackedTest(const CsaPackedProblem& base, float tolerance, + ProviderKind provider_kind = ProviderKind::Cuda) { + auto provider = CreateProvider(provider_kind); + if (provider == nullptr) { + GTEST_SKIP() << (provider_kind == ProviderKind::Cuda ? "CUDA" : "WebGPU") + << " execution provider is not available"; + } + + CsaPackedProblem problem = base; + problem.query = RoundTrip(problem.query); + problem.key = RoundTrip(problem.key); + problem.key_norm_weight = RoundTrip(problem.key_norm_weight); + problem.cos_cache = RoundTrip(problem.cos_cache); + problem.sin_cache = RoundTrip(problem.sin_cache); + problem.gate = RoundTrip(problem.gate); + problem.position_bias = RoundTrip(problem.position_bias); + problem.head_weights = RoundTrip(problem.head_weights); + problem.past_key_state = RoundTrip(problem.past_key_state); + problem.past_kv_buffer = RoundTrip(problem.past_kv_buffer); + problem.past_gate_buffer = RoundTrip(problem.past_gate_buffer); + + CsaPackedResult expected; + CsaPackedReference(problem, expected); + + const int64_t total_tokens = problem.TotalTokens(); + const int64_t batch_size = problem.batch_size; + const int64_t head_size = problem.head_size; + const int64_t width = problem.Width(); + const int64_t buffer_capacity = problem.BufferCapacity(); + + OpTester test("PackedSparseAttentionIndexer", 1, onnxruntime::kMSDomain); + test.AddAttribute("policy_mode", std::string(psai::kPolicyModeCsa)); + test.AddAttribute("compress_ratio", static_cast(problem.compress_ratio)); + test.AddAttribute("state_capacity", static_cast(problem.state_capacity)); + test.AddAttribute("index_topk", static_cast(problem.index_topk)); + if (problem.scale.has_value()) test.AddAttribute("scale", *problem.scale); + if (problem.head_weight_scale.has_value()) test.AddAttribute("head_weight_scale", *problem.head_weight_scale); + + test.AddInput("query", {total_tokens, problem.num_heads, head_size}, ToElementType(problem.query)); + test.AddInput("key", {total_tokens, width}, ToElementType(problem.key)); + test.AddInput("key_norm_weight", {head_size}, ToElementType(problem.key_norm_weight)); + test.AddInput("cos_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.cos_cache)); + test.AddInput("sin_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.sin_cache)); + test.AddInput("cumulative_sequence_lengths", {batch_size + 1}, problem.cumulative_sequence_lengths); + test.AddInput("past_sequence_lengths", {batch_size}, problem.past_sequence_lengths); + test.AddInput("gate", {total_tokens, width}, ToElementType(problem.gate)); + test.AddInput("position_bias", {problem.compress_ratio, width}, ToElementType(problem.position_bias)); + test.AddInput("head_weights", {total_tokens, problem.num_heads}, ToElementType(problem.head_weights)); + test.AddInput("position_ids", {total_tokens}, problem.position_ids); + test.AddInput("past_key_state", {batch_size, problem.state_capacity, head_size}, + ToElementType(problem.past_key_state)); + test.AddInput("past_kv_buffer", {batch_size, buffer_capacity, width}, ToElementType(problem.past_kv_buffer)); + test.AddInput("past_gate_buffer", {batch_size, buffer_capacity, width}, + ToElementType(problem.past_gate_buffer)); + test.AddInput("past_state_lengths", {batch_size, 2}, problem.past_state_lengths); + + test.AddOutput("selected_indices", {total_tokens, problem.index_topk}, expected.selected_indices); + test.AddOutput("selected_counts", {total_tokens}, expected.selected_counts); + test.AddOutput("present_key_state", {batch_size, problem.state_capacity, head_size}, + ToElementType(expected.present_key_state), false, 0.0f, tolerance); + test.AddOutput("present_kv_buffer", {batch_size, buffer_capacity, width}, + ToElementType(expected.present_kv_buffer), false, 0.0f, tolerance); + test.AddOutput("present_gate_buffer", {batch_size, buffer_capacity, width}, + ToElementType(expected.present_gate_buffer), false, 0.0f, tolerance); + test.AddOutput("present_state_lengths", {batch_size, 2}, expected.present_state_lengths); + RunOnProvider(test, std::move(provider)); +} + +} // namespace + +TEST(PackedSparseAttentionIndexerTest, CsaFloat) { RunCsaPackedTest(MakeCsaPackedProblem(), 1.0e-5f); } + +TEST(PackedSparseAttentionIndexerTest, CsaFloat16) { RunCsaPackedTest(MakeCsaPackedProblem(), 4.0e-3f); } + +TEST(PackedSparseAttentionIndexerTest, CsaBFloat16) { RunCsaPackedTest(MakeCsaPackedProblem(), 3.0e-2f); } + +// Prefill followed by decode: continues a previously compressed entry and a partially filled +// buffer for one request, while another request starts fresh. +TEST(PackedSparseAttentionIndexerTest, CsaPrefillThenDecodeIndependentState) { + CsaPackedProblem problem; + problem.cumulative_sequence_lengths = {0, 1, 4}; + problem.past_state_lengths = {1, 1, 0, 0}; + RunCsaPackedTest(MakeCsaPackedProblem(std::move(problem)), 1.0e-5f); +} + +} // namespace test +} // namespace onnxruntime diff --git a/onnxruntime/test/python/onnxruntime_test_python_symbolic_shape_infer.py b/onnxruntime/test/python/onnxruntime_test_python_symbolic_shape_infer.py index 9986fe62895ef..fc37603b42942 100644 --- a/onnxruntime/test/python/onnxruntime_test_python_symbolic_shape_infer.py +++ b/onnxruntime/test/python/onnxruntime_test_python_symbolic_shape_infer.py @@ -265,6 +265,135 @@ def test_sparse_attention_indexer_csa_symbolic_fallback(self): self.assertTrue(buffer_shape[1].startswith("SparseAttentionIndexer_")) self.assertEqual(self._tensor_shape(outputs["present_gate_buffer"]), buffer_shape) + def _infer_packed_sparse_attention_indexer(self, node, inputs): + outputs = [helper.make_tensor_value_info(name, TensorProto.UNDEFINED, None) for name in node.output if name] + graph = helper.make_graph([node], "PackedSparseAttentionIndexer_Test", inputs, outputs) + model = helper.make_model( + graph, + opset_imports=[helper.make_opsetid("", 17), helper.make_opsetid("com.microsoft", 1)], + ) + return SymbolicShapeInference.infer_shapes(model, auto_merge=True) + + def test_packed_sparse_attention_indexer_qsa(self): + node = helper.make_node( + "PackedSparseAttentionIndexer", + [ + "query", + "key", + "key_norm_weight", + "cos_cache", + "sin_cache", + "cumulative_sequence_lengths", + "past_sequence_lengths", + "", + "", + "", + "", + "past_key_state", + "past_kv_buffer", + "", + "past_state_lengths", + ], + [ + "selected_indices", + "selected_counts", + "present_key_state", + "present_kv_buffer", + "", + "present_state_lengths", + ], + domain="com.microsoft", + policy_mode="qsa", + compress_ratio=4, + state_capacity=5, + token_budget=8, + ) + inputs = [ + helper.make_tensor_value_info("query", TensorProto.FLOAT16, ["total_tokens", 2, 8]), + helper.make_tensor_value_info("key", TensorProto.FLOAT16, ["total_tokens", 8]), + helper.make_tensor_value_info("key_norm_weight", TensorProto.FLOAT16, [8]), + helper.make_tensor_value_info("cos_cache", TensorProto.FLOAT16, [64, 8]), + helper.make_tensor_value_info("sin_cache", TensorProto.FLOAT16, [64, 8]), + helper.make_tensor_value_info("cumulative_sequence_lengths", TensorProto.INT32, [3]), + helper.make_tensor_value_info("past_sequence_lengths", TensorProto.INT32, [2]), + helper.make_tensor_value_info("past_key_state", TensorProto.FLOAT16, [2, 5, 8]), + helper.make_tensor_value_info("past_kv_buffer", TensorProto.FLOAT16, [2, 7, 8]), + helper.make_tensor_value_info("past_state_lengths", TensorProto.INT32, [2, 2]), + ] + + inferred = self._infer_packed_sparse_attention_indexer(node, inputs) + outputs = {output.name: output for output in inferred.graph.output} + self.assertEqual(self._tensor_shape(outputs["selected_indices"]), ["total_tokens", 11]) + self.assertEqual(outputs["selected_indices"].type.tensor_type.elem_type, TensorProto.INT32) + self.assertEqual(self._tensor_shape(outputs["selected_counts"]), ["total_tokens"]) + self.assertEqual(outputs["selected_counts"].type.tensor_type.elem_type, TensorProto.INT32) + self.assertEqual(self._tensor_shape(outputs["present_key_state"]), [2, 5, 8]) + self.assertEqual(outputs["present_key_state"].type.tensor_type.elem_type, TensorProto.FLOAT16) + self.assertEqual(self._tensor_shape(outputs["present_kv_buffer"]), [2, 7, 8]) + self.assertEqual(self._tensor_shape(outputs["present_state_lengths"]), [2, 2]) + self.assertEqual(outputs["present_state_lengths"].type.tensor_type.elem_type, TensorProto.INT32) + + def test_packed_sparse_attention_indexer_csa(self): + node = helper.make_node( + "PackedSparseAttentionIndexer", + [ + "query", + "key", + "key_norm_weight", + "cos_cache", + "sin_cache", + "cumulative_sequence_lengths", + "past_sequence_lengths", + "gate", + "position_bias", + "head_weights", + "position_ids", + "past_key_state", + "past_kv_buffer", + "past_gate_buffer", + "past_state_lengths", + ], + [ + "selected_indices", + "selected_counts", + "present_key_state", + "present_kv_buffer", + "present_gate_buffer", + "present_state_lengths", + ], + domain="com.microsoft", + policy_mode="csa", + compress_ratio=4, + state_capacity=6, + index_topk=3, + ) + inputs = [ + helper.make_tensor_value_info("query", TensorProto.FLOAT, ["total_tokens", 2, 8]), + helper.make_tensor_value_info("key", TensorProto.FLOAT, ["total_tokens", 16]), + helper.make_tensor_value_info("key_norm_weight", TensorProto.FLOAT, [8]), + helper.make_tensor_value_info("cos_cache", TensorProto.FLOAT, [2, 64, 4]), + helper.make_tensor_value_info("sin_cache", TensorProto.FLOAT, [2, 64, 4]), + helper.make_tensor_value_info("cumulative_sequence_lengths", TensorProto.INT32, [3]), + helper.make_tensor_value_info("past_sequence_lengths", TensorProto.INT32, [2]), + helper.make_tensor_value_info("gate", TensorProto.FLOAT, ["total_tokens", 16]), + helper.make_tensor_value_info("position_bias", TensorProto.FLOAT, [4, 16]), + helper.make_tensor_value_info("head_weights", TensorProto.FLOAT, ["total_tokens", 2]), + helper.make_tensor_value_info("position_ids", TensorProto.INT64, ["total_tokens"]), + helper.make_tensor_value_info("past_key_state", TensorProto.FLOAT, [2, 6, 8]), + helper.make_tensor_value_info("past_kv_buffer", TensorProto.FLOAT, [2, 7, 16]), + helper.make_tensor_value_info("past_gate_buffer", TensorProto.FLOAT, [2, 7, 16]), + helper.make_tensor_value_info("past_state_lengths", TensorProto.INT32, [2, 2]), + ] + + inferred = self._infer_packed_sparse_attention_indexer(node, inputs) + outputs = {output.name: output for output in inferred.graph.output} + self.assertEqual(self._tensor_shape(outputs["selected_indices"]), ["total_tokens", 3]) + self.assertEqual(self._tensor_shape(outputs["selected_counts"]), ["total_tokens"]) + self.assertEqual(self._tensor_shape(outputs["present_key_state"]), [2, 6, 8]) + self.assertEqual(self._tensor_shape(outputs["present_kv_buffer"]), [2, 7, 16]) + self.assertEqual(self._tensor_shape(outputs["present_gate_buffer"]), [2, 7, 16]) + self.assertEqual(self._tensor_shape(outputs["present_state_lengths"]), [2, 2]) + def test_unsqueeze_opset_11(self): graph = helper.make_graph( [ From 6f8041c0358ccae3c04abad29bdb0afe3e9c02ed Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 16 Sep 2026 02:25:01 +0000 Subject: [PATCH 03/10] Harden packed indexer state validation Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- .../sparse/packed_sparse_attention_indexer.cc | 39 ++++--- .../sparse/packed_sparse_attention_indexer.h | 1 + .../packed_sparse_attention_indexer_impl.cu | 77 ++++++++----- .../bert/packed_sparse_attention_indexer.cc | 105 ++++++++++++----- .../bert/packed_sparse_attention_indexer.h | 3 + .../core/graph/contrib_ops/bert_defs.cc | 4 + ...packed_sparse_attention_indexer_op_test.cc | 107 ++++++++++++++---- 7 files changed, 244 insertions(+), 92 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc index eb23d13acd8b0..2b289bb80443e 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc @@ -19,17 +19,19 @@ namespace cuda { using namespace onnxruntime::cuda; namespace psai = onnxruntime::contrib::packed_sparse_attention_indexer; -#define REGISTER_KERNEL_TYPED(T) \ - ONNX_OPERATOR_TYPED_KERNEL_EX( \ - PackedSparseAttentionIndexer, \ - kMSDomain, \ - 1, \ - T, \ - kCudaExecutionProvider, \ - (*KernelDefBuilder::Create()) \ - .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ - .TypeConstraint("I", DataTypeImpl::GetTensorType()) \ - .TypeConstraint("M", DataTypeImpl::GetTensorType()), \ +#define REGISTER_KERNEL_TYPED(T) \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + PackedSparseAttentionIndexer, \ + kMSDomain, \ + 1, \ + T, \ + kCudaExecutionProvider, \ + (*KernelDefBuilder::Create()) \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("I", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("M", DataTypeImpl::GetTensorType()) \ + .MayInplace(11, 2) \ + .MayInplace(14, 5), \ PackedSparseAttentionIndexer); REGISTER_KERNEL_TYPED(float) @@ -98,11 +100,10 @@ PackedSparseAttentionIndexer::PackedSparseAttentionIndexer(const OpKernelInfo ORT_ENFORCE(compress_ratio_ > 0 && compress_ratio_ <= std::numeric_limits::max(), "PackedSparseAttentionIndexer: compress_ratio must be in (0, INT_MAX], got ", compress_ratio_); - int64_t state_capacity = 0; - ORT_ENFORCE(info.GetAttr("state_capacity", &state_capacity).IsOK(), + ORT_ENFORCE(info.GetAttr("state_capacity", &state_capacity_).IsOK(), "PackedSparseAttentionIndexer: state_capacity is required"); - ORT_ENFORCE(state_capacity > 0 && state_capacity <= std::numeric_limits::max(), - "PackedSparseAttentionIndexer: state_capacity must be in (0, INT_MAX], got ", state_capacity); + ORT_ENFORCE(state_capacity_ > 0 && state_capacity_ <= std::numeric_limits::max(), + "PackedSparseAttentionIndexer: state_capacity must be in (0, INT_MAX], got ", state_capacity_); const bool has_token_budget = info.GetAttr("token_budget", &token_budget_).IsOK(); const bool has_index_topk = info.GetAttr("index_topk", &index_topk_).IsOK(); @@ -197,6 +198,8 @@ Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) con cu_shape.ToString()); const int64_t batch_size = cu_shape[0] - 1; ORT_RETURN_IF_ERROR(CheckIntDimension("batch_size", batch_size)); + ORT_RETURN_IF(batch_size == 0 && total_tokens != 0, + "PackedSparseAttentionIndexer: total_tokens must be 0 when batch_size is 0"); ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, head_size})); @@ -221,6 +224,8 @@ Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) con key_state_shape.ToString()); const int64_t state_capacity = key_state_shape[1]; ORT_RETURN_IF_ERROR(CheckIntDimension("state_capacity", state_capacity, false)); + ORT_RETURN_IF_NOT(state_capacity == state_capacity_, + "PackedSparseAttentionIndexer: past_key_state capacity must match the state_capacity attribute"); const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, head_size})); @@ -329,6 +334,8 @@ Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) con cu_shape.ToString()); const int64_t batch_size = cu_shape[0] - 1; ORT_RETURN_IF_ERROR(CheckIntDimension("batch_size", batch_size)); + ORT_RETURN_IF(batch_size == 0 && total_tokens != 0, + "PackedSparseAttentionIndexer: total_tokens must be 0 when batch_size is 0"); ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, width})); @@ -354,6 +361,8 @@ Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) con key_state_shape.ToString()); const int64_t state_capacity = key_state_shape[1]; ORT_RETURN_IF_ERROR(CheckIntDimension("state_capacity", state_capacity, false)); + ORT_RETURN_IF_NOT(state_capacity == state_capacity_, + "PackedSparseAttentionIndexer: past_key_state capacity must match the state_capacity attribute"); const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, width})); diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h index b7ddfa9b480ce..c8f3ec19f2082 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h @@ -23,6 +23,7 @@ class PackedSparseAttentionIndexer final : public onnxruntime::cuda::CudaKernel packed_sparse_attention_indexer::Policy policy_; int64_t compress_ratio_; + int64_t state_capacity_; int64_t token_budget_; int64_t index_topk_; float epsilon_; diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu index 36f1a6a521b9b..3e005ce00ba65 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu @@ -85,7 +85,8 @@ __global__ void ElementwiseCopyKernel(const T* src, T* dst, int64_t count) { template __global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, const T* cos_cache, const T* sin_cache, const int32_t* cumulative_sequence_lengths, - const T* past_kv_buffer, const int32_t* past_state_lengths, + const int32_t* past_sequence_lengths, const T* past_kv_buffer, + const int32_t* past_state_lengths, T* present_key_state, T* present_kv_buffer, int32_t* present_state_lengths, int32_t* overflow_flags, PackedSparseAttentionIndexerParams params) { @@ -97,11 +98,20 @@ __global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, con for (int b = static_cast(blockIdx.x); b < params.batch_size; b += static_cast(gridDim.x)) { const int req_start = cumulative_sequence_lengths[b]; const int req_end = cumulative_sequence_lengths[b + 1]; - const int req_len = req_end > req_start ? req_end - req_start : 0; - - const int old_key_len = min(max(past_state_lengths[b * 2 + psai::kKeyStateLength], 0), params.state_capacity); - const int old_buf_len = - min(max(past_state_lengths[b * 2 + psai::kBufferLength], 0), params.compress_ratio - 1); + const int raw_key_len = past_state_lengths[b * 2 + psai::kKeyStateLength]; + const int raw_buf_len = past_state_lengths[b * 2 + psai::kBufferLength]; + const int past_sequence_length = past_sequence_lengths[b]; + const bool invalid_metadata = + cumulative_sequence_lengths[0] != 0 || + cumulative_sequence_lengths[params.batch_size] != params.total_tokens || + req_start < 0 || req_end < req_start || req_end > params.total_tokens || + past_sequence_length < 0 || raw_key_len < 0 || raw_key_len > params.state_capacity || + raw_buf_len < 0 || raw_buf_len >= params.compress_ratio || + raw_key_len != past_sequence_length / params.compress_ratio || + raw_buf_len != past_sequence_length % params.compress_ratio; + const int req_len = invalid_metadata ? 0 : req_end - req_start; + const int old_key_len = min(max(raw_key_len, 0), params.state_capacity); + const int old_buf_len = min(max(raw_buf_len, 0), params.compress_ratio - 1); const int pending = old_buf_len + req_len; const int full_new_block_count = pending / params.compress_ratio; @@ -109,9 +119,9 @@ __global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, con // Reject (do not partially apply) a step that would need more than the fixed state_capacity: // no new blocks are formed and the buffer is left exactly as it was, so a rejected step is a // deterministic no-op on state rather than a silent partial truncation. - const bool overflowed = full_new_block_count > capacity_left; - const int new_block_count = overflowed ? 0 : full_new_block_count; - const int new_buf_len = overflowed ? old_buf_len : (pending % params.compress_ratio); + const bool rejected = invalid_metadata || full_new_block_count > capacity_left; + const int new_block_count = rejected ? 0 : full_new_block_count; + const int new_buf_len = rejected ? old_buf_len : (pending % params.compress_ratio); // Barrier: every thread has now read past_state_lengths (identically) before any thread below // writes present_state_lengths, which keeps this correct even if the two tensors alias. @@ -120,7 +130,7 @@ __global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, con if (threadIdx.x == 0) { present_state_lengths[b * 2 + psai::kKeyStateLength] = old_key_len + new_block_count; present_state_lengths[b * 2 + psai::kBufferLength] = new_buf_len; - overflow_flags[b] = overflowed ? 1 : 0; + overflow_flags[b] = rejected ? 1 : 0; } for (int k = 0; k < new_block_count; ++k) { @@ -384,7 +394,8 @@ __global__ void QsaSelectKernel(const float* block_scores, const int32_t* cumula template __global__ void CsaUpdateStateKernel(const T* key, const T* gate, const T* key_norm_weight, const T* cos_cache, const T* sin_cache, const T* position_bias, - const int32_t* cumulative_sequence_lengths, const T* past_kv_buffer, + const int32_t* cumulative_sequence_lengths, + const int32_t* past_sequence_lengths, const T* past_kv_buffer, const T* past_gate_buffer, const int32_t* past_state_lengths, T* present_key_state, T* present_kv_buffer, T* present_gate_buffer, int32_t* present_state_lengths, int32_t* overflow_flags, @@ -397,11 +408,17 @@ __global__ void CsaUpdateStateKernel(const T* key, const T* gate, const T* key_n for (int b = static_cast(blockIdx.x); b < params.batch_size; b += static_cast(gridDim.x)) { const int req_start = cumulative_sequence_lengths[b]; const int req_end = cumulative_sequence_lengths[b + 1]; - const int req_len = req_end > req_start ? req_end - req_start : 0; - - const int old_key_len = min(max(past_state_lengths[b * 2 + psai::kKeyStateLength], 0), params.state_capacity); - const int old_buf_len = - min(max(past_state_lengths[b * 2 + psai::kBufferLength], 0), params.buffer_capacity); + const int raw_key_len = past_state_lengths[b * 2 + psai::kKeyStateLength]; + const int raw_buf_len = past_state_lengths[b * 2 + psai::kBufferLength]; + const bool invalid_metadata = + cumulative_sequence_lengths[0] != 0 || + cumulative_sequence_lengths[params.batch_size] != params.total_tokens || + req_start < 0 || req_end < req_start || req_end > params.total_tokens || + past_sequence_lengths[b] < 0 || raw_key_len < 0 || raw_key_len > params.state_capacity || + raw_buf_len < 0 || raw_buf_len > params.buffer_capacity; + const int req_len = invalid_metadata ? 0 : req_end - req_start; + const int old_key_len = min(max(raw_key_len, 0), params.state_capacity); + const int old_buf_len = min(max(raw_buf_len, 0), params.buffer_capacity); sai::CsaWindowPlan plan; const bool plan_ok = sai::TryComputeCsaWindowPlan(old_buf_len, req_len, params.compress_ratio, plan); @@ -413,14 +430,14 @@ __global__ void CsaUpdateStateKernel(const T* key, const T* gate, const T* key_n // Reject (do not partially apply) a step that would need more than the fixed state_capacity: // no new windows are closed and the buffer is left exactly as it was, so a rejected step is a // deterministic no-op on state rather than a silent partial truncation. - const bool overflowed = full_new_window_count > capacity_left; - const int new_window_count = overflowed ? 0 : full_new_window_count; + const bool rejected = invalid_metadata || full_new_window_count > capacity_left; + const int new_window_count = rejected ? 0 : full_new_window_count; const int present_buffer_length = - overflowed ? old_buf_len - : (static_cast(plan.present_buffer_length) < params.buffer_capacity - ? static_cast(plan.present_buffer_length) - : params.buffer_capacity); - const int present_buffer_start = overflowed ? 0 : static_cast(plan.present_buffer_start); + rejected ? old_buf_len + : (static_cast(plan.present_buffer_length) < params.buffer_capacity + ? static_cast(plan.present_buffer_length) + : params.buffer_capacity); + const int present_buffer_start = rejected ? 0 : static_cast(plan.present_buffer_start); const int overlap_length = static_cast(plan.overlap_length); // Barrier: every thread has now read past_state_lengths / computed the plan (identically) @@ -431,7 +448,7 @@ __global__ void CsaUpdateStateKernel(const T* key, const T* gate, const T* key_n if (threadIdx.x == 0) { present_state_lengths[b * 2 + psai::kKeyStateLength] = old_key_len + new_window_count; present_state_lengths[b * 2 + psai::kBufferLength] = present_buffer_length; - overflow_flags[b] = overflowed ? 1 : 0; + overflow_flags[b] = rejected ? 1 : 0; } for (int k = 0; k < new_window_count; ++k) { @@ -739,8 +756,8 @@ Status LaunchQsaPackedSparseAttentionIndexer( const size_t value_bytes = static_cast(params.head_size) * sizeof(float); const int state_blocks = static_cast(std::min(params.batch_size, kSaiMaxGridDimX)); QsaUpdateStateKernel<<>>( - key, key_norm_weight, cos_cache, sin_cache, cumulative_sequence_lengths, past_kv_buffer, past_state_lengths, - present_key_state, present_kv_buffer, present_state_lengths, overflow_flags, params); + key, key_norm_weight, cos_cache, sin_cache, cumulative_sequence_lengths, past_sequence_lengths, past_kv_buffer, + past_state_lengths, present_key_state, present_kv_buffer, present_state_lengths, overflow_flags, params); if (params.total_tokens == 0) { return CUDA_CALL(cudaGetLastError()); @@ -814,9 +831,9 @@ Status LaunchCsaPackedSparseAttentionIndexer( const size_t value_bytes = static_cast(params.head_size) * sizeof(float); const int state_blocks = static_cast(std::min(params.batch_size, kSaiMaxGridDimX)); CsaUpdateStateKernel<<>>( - key, gate, key_norm_weight, cos_cache, sin_cache, position_bias, cumulative_sequence_lengths, past_kv_buffer, - past_gate_buffer, past_state_lengths, present_key_state, present_kv_buffer, present_gate_buffer, - present_state_lengths, overflow_flags, params); + key, gate, key_norm_weight, cos_cache, sin_cache, position_bias, cumulative_sequence_lengths, + past_sequence_lengths, past_kv_buffer, past_gate_buffer, past_state_lengths, present_key_state, + present_kv_buffer, present_gate_buffer, present_state_lengths, overflow_flags, params); if (params.total_tokens == 0) { return CUDA_CALL(cudaGetLastError()); @@ -855,7 +872,7 @@ Status LaunchCsaPackedSparseAttentionIndexer( template Status LaunchCsaPackedSparseAttentionIndexer( \ cudaStream_t, const PackedSparseAttentionIndexerParams&, const T*, const T*, const T*, const T*, \ const T*, const T*, const T*, const T*, const int32_t*, const int32_t*, const int64_t*, const T*, \ - const T*, const T*, const int32_t*, int32_t*, int32_t*, T*, T*, T*, int32_t*, float*); + const T*, const T*, const int32_t*, int32_t*, int32_t*, T*, T*, T*, int32_t*, float*, int32_t*); INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(float) INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(half) diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc index e12d4ddbb5c1d..a277d7dc12264 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -97,6 +97,7 @@ Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHe const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); + const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); const auto& past_kv_buffer = shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); const auto& past_state_lengths = shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); const auto& present_key_state = @@ -161,17 +162,29 @@ Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHe << " if (b >= uniforms.batch_size || local_idx != 0u) { return; }\n" << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" << " let req_end = " << cu_seqlens.GetByOffset("b + 1u") << ";\n" - << " let req_len = max(req_end - req_start, 0);\n" - << " let old_key_len = clamp(" << past_state_lengths.GetByOffset("b * 2u") - << ", 0, i32(uniforms.state_capacity));\n" - << " let old_buf_len = clamp(" << past_state_lengths.GetByOffset("b * 2u + 1u") - << ", 0, i32(uniforms.compress_ratio) - 1);\n" + << " let old_key_len = " << past_state_lengths.GetByOffset("b * 2u") << ";\n" + << " let old_buf_len = " << past_state_lengths.GetByOffset("b * 2u + 1u") << ";\n" + << " let past_len = " << past_seqlens.GetByOffset("b") << ";\n" + << " let invalid = " << cu_seqlens.GetByOffset("0") << " != 0 || " + << cu_seqlens.GetByOffset("uniforms.batch_size") << " != i32(uniforms.total_tokens) || " + << "req_start < 0 || req_end < req_start || req_end > i32(uniforms.total_tokens) || past_len < 0 || " + << "old_key_len < 0 || old_key_len > i32(uniforms.state_capacity) || old_buf_len < 0 || " + << "old_buf_len >= i32(uniforms.compress_ratio) || " + << "old_key_len != past_len / i32(uniforms.compress_ratio) || " + << "old_buf_len != past_len % i32(uniforms.compress_ratio);\n" + << " if (invalid) {\n" + << " " << overflow_flags.SetByOffset("b", "1") << "\n" + << " return;\n" + << " }\n" + << " let req_len = req_end - req_start;\n" << " let pending = old_buf_len + req_len;\n" << " let full_new_block_count = pending / i32(uniforms.compress_ratio);\n" << " let capacity_left = max(i32(uniforms.state_capacity) - old_key_len, 0);\n" - << " let new_block_count = min(full_new_block_count, capacity_left);\n" - << " let overflowed = new_block_count < full_new_block_count;\n" - << " let new_buf_len = select(pending % i32(uniforms.compress_ratio), 0, overflowed);\n" + << " let overflowed = full_new_block_count > capacity_left;\n" + << " " << overflow_flags.SetByOffset("b", "select(0, 1, overflowed)") << "\n" + << " if (overflowed) { return; }\n" + << " let new_block_count = full_new_block_count;\n" + << " let new_buf_len = pending % i32(uniforms.compress_ratio);\n" << " " << present_state_lengths.SetByOffset("b * 2u", "old_key_len + new_block_count") << "\n" << " " << present_state_lengths.SetByOffset("b * 2u + 1u", "new_buf_len") << "\n" << " for (var k = 0; k < new_block_count; k++) {\n" @@ -209,6 +222,7 @@ Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHe position_ids = &shader.AddInput("position_ids", ShaderUsage::UseUniform); } const auto& present_state_lengths = shader.AddInput("present_state_lengths", ShaderUsage::UseUniform); + const auto& overflow_flags = shader.AddInput("overflow_flags", ShaderUsage::UseUniform); const auto& selected_indices = shader.AddOutput("selected_indices", ShaderUsage::UseUniform); const auto& selected_counts = shader.AddOutput("selected_counts", ShaderUsage::UseUniform); @@ -288,6 +302,10 @@ Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHe << " " << selected_indices.SetByOffset("output_base + i", "-1") << "\n" << " }\n" << " let b = batch_of_token(token);\n" + << " if (" << overflow_flags.GetByOffset("b") << " != 0) {\n" + << " " << selected_counts.SetByOffset("token", "0") << "\n" + << " return;\n" + << " }\n" << " let key_len_after = u32(" << present_state_lengths.GetByOffset("b * 2u") << ");\n" << " let position = abs_position(token, b);\n" << " let causal = causal_count(position);\n" @@ -345,6 +363,7 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); const auto& position_bias = shader.AddInput("position_bias", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); + const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); const auto& past_kv_buffer = shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); const auto& past_gate_buffer = shader.AddInput("past_gate_buffer", ShaderUsage::UseUniform); const auto& past_state_lengths = shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); @@ -355,6 +374,7 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe const auto& present_gate_buffer = shader.AddOutput("present_gate_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_state_lengths = shader.AddOutput("present_state_lengths", ShaderUsage::UseUniform); + const auto& overflow_flags = shader.AddOutput("overflow_flags", ShaderUsage::UseUniform); shader.AdditionalImplementation() << "fn extended_key(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, channel: u32) -> f32 {\n" @@ -426,31 +446,39 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe << " if (b >= uniforms.batch_size || local_idx != 0u) { return; }\n" << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" << " let req_end = " << cu_seqlens.GetByOffset("b + 1u") << ";\n" - << " let req_len = max(req_end - req_start, 0);\n" - << " let old_key_len = clamp(" << past_state_lengths.GetByOffset("b * 2u") - << ", 0, i32(uniforms.state_capacity));\n" - << " let old_buf_len = clamp(" << past_state_lengths.GetByOffset("b * 2u + 1u") - << ", 0, i32(uniforms.buffer_capacity));\n" + << " let old_key_len = " << past_state_lengths.GetByOffset("b * 2u") << ";\n" + << " let old_buf_len = " << past_state_lengths.GetByOffset("b * 2u + 1u") << ";\n" + << " let invalid = " << cu_seqlens.GetByOffset("0") << " != 0 || " + << cu_seqlens.GetByOffset("uniforms.batch_size") << " != i32(uniforms.total_tokens) || " + << "req_start < 0 || req_end < req_start || req_end > i32(uniforms.total_tokens) || " + << past_seqlens.GetByOffset("b") << " < 0 || old_key_len < 0 || " + << "old_key_len > i32(uniforms.state_capacity) || old_buf_len < 0 || " + << "old_buf_len > i32(uniforms.buffer_capacity);\n" + << " if (invalid) {\n" + << " " << overflow_flags.SetByOffset("b", "1") << "\n" + << " return;\n" + << " }\n" + << " let req_len = req_end - req_start;\n" << " let overlap_length = select(0, i32(uniforms.compress_ratio), old_buf_len >= " "i32(uniforms.compress_ratio));\n" << " let leftover_length = old_buf_len - overlap_length;\n" << " let pending = leftover_length + req_len;\n" << " let full_new_window_count = pending / i32(uniforms.compress_ratio);\n" << " let capacity_left = max(i32(uniforms.state_capacity) - old_key_len, 0);\n" - << " let new_window_count = min(full_new_window_count, capacity_left);\n" - << " let overflowed = new_window_count < full_new_window_count;\n" + << " let overflowed = full_new_window_count > capacity_left;\n" + << " " << overflow_flags.SetByOffset("b", "select(0, 1, overflowed)") << "\n" + << " if (overflowed) { return; }\n" + << " let new_window_count = full_new_window_count;\n" << " var present_buffer_length = 0;\n" << " var present_buffer_start = 0;\n" - << " if (!overflowed) {\n" - << " if (new_window_count > 0) {\n" - << " present_buffer_length = i32(uniforms.compress_ratio) + pending % i32(uniforms.compress_ratio);\n" - << " present_buffer_start = overlap_length + (new_window_count - 1) * i32(uniforms.compress_ratio);\n" - << " } else {\n" - << " present_buffer_length = old_buf_len + req_len;\n" - << " present_buffer_start = 0;\n" - << " }\n" - << " present_buffer_length = min(present_buffer_length, i32(uniforms.buffer_capacity));\n" + << " if (new_window_count > 0) {\n" + << " present_buffer_length = i32(uniforms.compress_ratio) + pending % i32(uniforms.compress_ratio);\n" + << " present_buffer_start = overlap_length + (new_window_count - 1) * i32(uniforms.compress_ratio);\n" + << " } else {\n" + << " present_buffer_length = old_buf_len + req_len;\n" + << " present_buffer_start = 0;\n" << " }\n" + << " present_buffer_length = min(present_buffer_length, i32(uniforms.buffer_capacity));\n" << " " << present_state_lengths.SetByOffset("b * 2u", "old_key_len + new_window_count") << "\n" << " " << present_state_lengths.SetByOffset("b * 2u + 1u", "present_buffer_length") << "\n" << " for (var k = 0; k < new_window_count; k++) {\n" @@ -517,6 +545,7 @@ Status PackedSparseAttentionIndexerCsaSelectProgram::GenerateShaderCode(ShaderHe const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); const auto& present_state_lengths = shader.AddInput("present_state_lengths", ShaderUsage::UseUniform); + const auto& overflow_flags = shader.AddInput("overflow_flags", ShaderUsage::UseUniform); const auto& selected_indices = shader.AddOutput("selected_indices", ShaderUsage::UseUniform); const auto& selected_counts = shader.AddOutput("selected_counts", ShaderUsage::UseUniform); @@ -585,6 +614,10 @@ Status PackedSparseAttentionIndexerCsaSelectProgram::GenerateShaderCode(ShaderHe << " " << selected_indices.SetByOffset("output_base + i", "-1") << "\n" << " }\n" << " let b = batch_of_token(token);\n" + << " if (" << overflow_flags.GetByOffset("b") << " != 0) {\n" + << " " << selected_counts.SetByOffset("token", "0") << "\n" + << " return;\n" + << " }\n" << " let key_len_after = u32(" << present_state_lengths.GetByOffset("b * 2u") << ");\n" << " let raw = " << position_ids.GetByOffset("token", true) << ";\n" << " let threshold = min(causal_count(raw), key_len_after);\n" @@ -624,6 +657,8 @@ PackedSparseAttentionIndexer::PackedSparseAttentionIndexer(const OpKernelInfo& i ORT_ENFORCE(info.GetAttr("compress_ratio", &compress_ratio_).IsOK(), "PackedSparseAttentionIndexer: compress_ratio is required"); ORT_ENFORCE(compress_ratio_ > 0, "PackedSparseAttentionIndexer: compress_ratio must be > 0"); + ORT_ENFORCE(info.GetAttr("state_capacity", &state_capacity_).IsOK() && state_capacity_ > 0, + "PackedSparseAttentionIndexer: state_capacity must be a positive integer"); const bool has_token_budget = info.GetAttr("token_budget", &token_budget_).IsOK(); const bool has_index_topk = info.GetAttr("index_topk", &index_topk_).IsOK(); @@ -691,6 +726,8 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1, "PackedSparseAttentionIndexer: invalid cumulative_sequence_lengths shape"); const int64_t batch_size = cu_shape[0] - 1; + ORT_RETURN_IF(batch_size == 0 && total_tokens != 0, + "PackedSparseAttentionIndexer: total_tokens must be 0 when batch_size is 0"); ORT_RETURN_IF_ERROR(CheckShape(past_seqlens, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, head_size})); @@ -710,6 +747,8 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont key_state_shape[2] == head_size, "PackedSparseAttentionIndexer: invalid past_key_state shape"); const int64_t state_capacity = key_state_shape[1]; + ORT_RETURN_IF_NOT(state_capacity == state_capacity_, + "PackedSparseAttentionIndexer: past_key_state capacity must match the state_capacity attribute"); const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, head_size})); @@ -724,6 +763,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont context.Output(psai::kPresentKvBuffer, TensorShape({batch_size, buffer_capacity, head_size})); Tensor* present_state_lengths = context.Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns})); + Tensor overflow_flags = context.CreateGPUTensor(DataTypeImpl::GetType(), TensorShape({batch_size})); if (present_key_state->DataRaw() != past_key_state->DataRaw()) { const int64_t total = present_key_state->Shape().Size(); @@ -771,13 +811,16 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont {cos_cache, ProgramTensorMetadataDependency::Type}, {sin_cache, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}, + {past_seqlens, ProgramTensorMetadataDependency::Type}, {past_kv_buffer, ProgramTensorMetadataDependency::Type}}) .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) .AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, {present_kv_buffer, ProgramTensorMetadataDependency::Type}}) .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddOutput({&overflow_flags, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32(batch_size)) .AddUniformVariables({{ToUint32(batch_size)}, + {ToUint32(total_tokens)}, {ToUint32(compress_ratio_)}, {ToUint32(state_capacity)}, {ToUint32(buffer_capacity)}, @@ -804,7 +847,8 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont if (position_ids != nullptr) { select.AddInput({position_ids, ProgramTensorMetadataDependency::Type}); } - select.AddInput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + select.AddInputs({{present_state_lengths, ProgramTensorMetadataDependency::Type}, + {&overflow_flags, ProgramTensorMetadataDependency::Type}}) .AddOutputs({{selected_indices, ProgramTensorMetadataDependency::Type}, {selected_counts, ProgramTensorMetadataDependency::Type}}) .SetDispatchGroupSize(ToUint32(total_tokens)) @@ -854,6 +898,8 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1, "PackedSparseAttentionIndexer: invalid cumulative_sequence_lengths shape"); const int64_t batch_size = cu_shape[0] - 1; + ORT_RETURN_IF(batch_size == 0 && total_tokens != 0, + "PackedSparseAttentionIndexer: total_tokens must be 0 when batch_size is 0"); ORT_RETURN_IF_ERROR(CheckShape(past_seqlens, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, width})); @@ -874,6 +920,8 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont key_state_shape[2] == head_size, "PackedSparseAttentionIndexer: invalid past_key_state shape"); const int64_t state_capacity = key_state_shape[1]; + ORT_RETURN_IF_NOT(state_capacity == state_capacity_, + "PackedSparseAttentionIndexer: past_key_state capacity must match the state_capacity attribute"); const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_); ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, width})); @@ -891,6 +939,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont context.Output(psai::kPresentGateBuffer, TensorShape({batch_size, buffer_capacity, width})); Tensor* present_state_lengths = context.Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns})); + Tensor overflow_flags = context.CreateGPUTensor(DataTypeImpl::GetType(), TensorShape({batch_size})); auto copy_if_needed = [&](const Tensor* src, Tensor* dst) -> Status { if (src->DataRaw() == dst->DataRaw()) { @@ -924,6 +973,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont {sin_cache, ProgramTensorMetadataDependency::Type}, {position_bias, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}, + {past_seqlens, ProgramTensorMetadataDependency::Type}, {past_kv_buffer, ProgramTensorMetadataDependency::Type}, {past_gate_buffer, ProgramTensorMetadataDependency::Type}}) .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) @@ -931,8 +981,10 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont {present_kv_buffer, ProgramTensorMetadataDependency::Type}, {present_gate_buffer, ProgramTensorMetadataDependency::Type}}) .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddOutput({&overflow_flags, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32(batch_size)) .AddUniformVariables({{ToUint32(batch_size)}, + {ToUint32(total_tokens)}, {ToUint32(compress_ratio_)}, {ToUint32(state_capacity)}, {ToUint32(buffer_capacity)}, @@ -957,7 +1009,8 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont {cos_cache, ProgramTensorMetadataDependency::Type}, {sin_cache, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}}) - .AddInput({present_state_lengths, ProgramTensorMetadataDependency::Type}) + .AddInputs({{present_state_lengths, ProgramTensorMetadataDependency::Type}, + {&overflow_flags, ProgramTensorMetadataDependency::Type}}) .AddOutputs({{selected_indices, ProgramTensorMetadataDependency::Type}, {selected_counts, ProgramTensorMetadataDependency::Type}}) .SetDispatchGroupSize(ToUint32(total_tokens)) diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h index c25b17b52e7ad..7be03c2a75310 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h @@ -36,6 +36,7 @@ class PackedSparseAttentionIndexerQsaUpdateProgram final Status GenerateShaderCode(ShaderHelper& shader) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( {"batch_size", ProgramUniformVariableDataType::Uint32}, + {"total_tokens", ProgramUniformVariableDataType::Uint32}, {"compress_ratio", ProgramUniformVariableDataType::Uint32}, {"state_capacity", ProgramUniformVariableDataType::Uint32}, {"buffer_capacity", ProgramUniformVariableDataType::Uint32}, @@ -88,6 +89,7 @@ class PackedSparseAttentionIndexerCsaUpdateProgram final Status GenerateShaderCode(ShaderHelper& shader) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( {"batch_size", ProgramUniformVariableDataType::Uint32}, + {"total_tokens", ProgramUniformVariableDataType::Uint32}, {"compress_ratio", ProgramUniformVariableDataType::Uint32}, {"state_capacity", ProgramUniformVariableDataType::Uint32}, {"buffer_capacity", ProgramUniformVariableDataType::Uint32}, @@ -137,6 +139,7 @@ class PackedSparseAttentionIndexer final : public WebGpuKernel { packed_sparse_attention_indexer::Policy policy_; int64_t compress_ratio_; + int64_t state_capacity_; int64_t token_budget_; int64_t index_topk_; float epsilon_; diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index 93435b40f8bc0..cdc93a5cec6b2 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -2439,6 +2439,10 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference // State never grows: present_* always has exactly the same fixed shape as past_*. const auto* key_state_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKeyState, 3); if (key_state_shape != nullptr) { + if (key_state_shape->dim(1).has_dim_value() && key_state_shape->dim(1).dim_value() != state_capacity) { + fail_shape_inference("PackedSparseAttentionIndexer: past_key_state dimension 1 must equal state_capacity (", + state_capacity, "), got ", key_state_shape->dim(1).dim_value()); + } updateOutputShape(ctx, psai::kPresentKeyState, *key_state_shape); } const auto* kv_buffer_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKvBuffer, 3); diff --git a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc index e6197057f623d..331247e8ded5b 100644 --- a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -91,6 +91,7 @@ struct GraphOptions { int64_t rotary_width = 8; int64_t compress_ratio = 2; int64_t state_capacity = 6; + int64_t input_state_capacity = -1; int64_t token_budget = 4; bool add_index_topk = false; bool add_csa_inputs = false; @@ -133,8 +134,10 @@ void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { } else { inputs.push_back(&empty); } + const int64_t input_state_capacity = + options.input_state_capacity >= 0 ? options.input_state_capacity : options.state_capacity; inputs.push_back( - builder.MakeInput(std::vector{options.batch_size, options.state_capacity, options.head_size})); + builder.MakeInput(std::vector{options.batch_size, input_state_capacity, options.head_size})); inputs.push_back( builder.MakeInput(std::vector{options.batch_size, buffer_capacity, width})); if (is_csa) { @@ -181,6 +184,13 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, QsaInfersFixedCapacityAndSt ONNX_NAMESPACE::TensorProto_DataType_INT32, {options.batch_size, 2}); } +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsStateCapacityShapeMismatch) { + GraphOptions options; + options.input_state_capacity = options.state_capacity + 1; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "past_key_state dimension 1 must equal state_capacity"); +} + TEST(PackedSparseAttentionIndexerShapeInferenceTest, CsaInfersFixedCapacityAndState) { GraphOptions options; options.policy_mode = psai::kPolicyModeCsa; @@ -481,10 +491,11 @@ void QsaPackedReference(const QsaPackedProblem& p, QsaPackedResult& out) { const int capacity = p.Capacity(); out.present_key_state = p.past_key_state; - out.present_kv_buffer.assign(static_cast(p.batch_size) * buffer_capacity * head_size, 0.0f); - out.present_state_lengths.assign(static_cast(p.batch_size) * 2, 0); + out.present_kv_buffer = p.past_kv_buffer; + out.present_state_lengths = p.past_state_lengths; std::vector key_len_after(static_cast(p.batch_size)); + std::vector overflowed(static_cast(p.batch_size), false); for (int b = 0; b < p.batch_size; ++b) { const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; const int req_end = p.cumulative_sequence_lengths[static_cast(b) + 1]; @@ -494,8 +505,14 @@ void QsaPackedReference(const QsaPackedProblem& p, QsaPackedResult& out) { const int old_buf_len = std::clamp(p.past_state_lengths[static_cast(b) * 2 + 1], 0, p.compress_ratio - 1); const int pending = old_buf_len + req_len; - const int new_block_count = std::min(pending / p.compress_ratio, std::max(p.state_capacity - old_key_len, 0)); - const int new_buf_len = new_block_count < pending / p.compress_ratio ? 0 : pending % p.compress_ratio; + const int full_new_block_count = pending / p.compress_ratio; + overflowed[static_cast(b)] = full_new_block_count > std::max(p.state_capacity - old_key_len, 0); + if (overflowed[static_cast(b)]) { + key_len_after[static_cast(b)] = old_key_len; + continue; + } + const int new_block_count = full_new_block_count; + const int new_buf_len = pending % p.compress_ratio; auto raw_value = [&](int virtual_pos, int d) -> float { return virtual_pos < old_buf_len @@ -535,6 +552,9 @@ void QsaPackedReference(const QsaPackedProblem& p, QsaPackedResult& out) { out.selected_indices.assign(static_cast(total_tokens) * capacity, -1); out.selected_counts.assign(static_cast(total_tokens), 0); for (int b = 0; b < p.batch_size; ++b) { + if (overflowed[static_cast(b)]) { + continue; + } const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; const int req_end = p.cumulative_sequence_lengths[static_cast(b) + 1]; for (int token = req_start; token < req_end; ++token) { @@ -705,7 +725,7 @@ TEST(PackedSparseAttentionIndexerTest, QsaZeroTokenRequestRow) { } // state_capacity smaller than what the incoming tokens would naturally produce: the update must -// deterministically drop the overflowing blocks rather than corrupt memory. +// reject the request's step without partially changing its state. TEST(PackedSparseAttentionIndexerTest, QsaStateCapacityOverflowIsSafe) { QsaPackedProblem problem; problem.cumulative_sequence_lengths = {0, 6}; @@ -716,6 +736,22 @@ TEST(PackedSparseAttentionIndexerTest, QsaStateCapacityOverflowIsSafe) { RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(std::move(problem))); } +#ifdef USE_WEBGPU +TEST(PackedSparseAttentionIndexerWebGpuTest, QsaFloat) { + RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(), ProviderKind::WebGpu); +} + +TEST(PackedSparseAttentionIndexerWebGpuTest, QsaStateCapacityOverflowIsRejected) { + QsaPackedProblem problem; + problem.batch_size = 1; + problem.cumulative_sequence_lengths = {0, 6}; + problem.past_sequence_lengths = {4}; + problem.state_capacity = 2; + problem.past_state_lengths = {2, 0}; + RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(std::move(problem)), ProviderKind::WebGpu); +} +#endif + namespace { struct CsaPackedProblem { @@ -771,10 +807,11 @@ void CsaPackedReference(const CsaPackedProblem& p, CsaPackedResult& out) { const float head_weight_scale = p.head_weight_scale.value_or(1.0f / std::sqrt(static_cast(p.num_heads))); out.present_key_state = p.past_key_state; - out.present_kv_buffer.assign(static_cast(p.batch_size) * buffer_capacity * width, 0.0f); - out.present_gate_buffer.assign(static_cast(p.batch_size) * buffer_capacity * width, 0.0f); - out.present_state_lengths.assign(static_cast(p.batch_size) * 2, 0); + out.present_kv_buffer = p.past_kv_buffer; + out.present_gate_buffer = p.past_gate_buffer; + out.present_state_lengths = p.past_state_lengths; std::vector key_len_after(static_cast(p.batch_size)); + std::vector overflowed(static_cast(p.batch_size), false); for (int b = 0; b < p.batch_size; ++b) { const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; @@ -788,21 +825,22 @@ void CsaPackedReference(const CsaPackedProblem& p, CsaPackedResult& out) { const int leftover_length = old_buf_len - overlap_length; const int pending = leftover_length + req_len; const int full_new_window_count = pending / p.compress_ratio; - const int new_window_count = - std::min(full_new_window_count, std::max(p.state_capacity - old_key_len, 0)); - const bool overflowed = new_window_count < full_new_window_count; + overflowed[static_cast(b)] = full_new_window_count > std::max(p.state_capacity - old_key_len, 0); + if (overflowed[static_cast(b)]) { + key_len_after[static_cast(b)] = old_key_len; + continue; + } + const int new_window_count = full_new_window_count; int present_buffer_length = 0; int present_buffer_start = 0; - if (!overflowed) { - if (new_window_count > 0) { - present_buffer_length = p.compress_ratio + pending % p.compress_ratio; - present_buffer_start = overlap_length + (new_window_count - 1) * p.compress_ratio; - } else { - present_buffer_length = old_buf_len + req_len; - present_buffer_start = 0; - } - present_buffer_length = std::min(present_buffer_length, buffer_capacity); + if (new_window_count > 0) { + present_buffer_length = p.compress_ratio + pending % p.compress_ratio; + present_buffer_start = overlap_length + (new_window_count - 1) * p.compress_ratio; + } else { + present_buffer_length = old_buf_len + req_len; + present_buffer_start = 0; } + present_buffer_length = std::min(present_buffer_length, buffer_capacity); auto extended_key = [&](int virtual_pos, int channel) -> float { return virtual_pos < old_buf_len @@ -874,6 +912,9 @@ void CsaPackedReference(const CsaPackedProblem& p, CsaPackedResult& out) { out.selected_indices.assign(static_cast(total_tokens) * capacity, -1); out.selected_counts.assign(static_cast(total_tokens), 0); for (int b = 0; b < p.batch_size; ++b) { + if (overflowed[static_cast(b)]) { + continue; + } const int req_start = p.cumulative_sequence_lengths[static_cast(b)]; const int req_end = p.cumulative_sequence_lengths[static_cast(b) + 1]; for (int token = req_start; token < req_end; ++token) { @@ -1040,5 +1081,29 @@ TEST(PackedSparseAttentionIndexerTest, CsaPrefillThenDecodeIndependentState) { RunCsaPackedTest(MakeCsaPackedProblem(std::move(problem)), 1.0e-5f); } +TEST(PackedSparseAttentionIndexerTest, CsaStateCapacityOverflowIsRejected) { + CsaPackedProblem problem; + problem.batch_size = 1; + problem.cumulative_sequence_lengths = {0, 4}; + problem.state_capacity = 1; + problem.past_state_lengths = {1, 0}; + RunCsaPackedTest(MakeCsaPackedProblem(std::move(problem)), 1.0e-5f); +} + +#ifdef USE_WEBGPU +TEST(PackedSparseAttentionIndexerWebGpuTest, CsaFloat) { + RunCsaPackedTest(MakeCsaPackedProblem(), 1.0e-5f, ProviderKind::WebGpu); +} + +TEST(PackedSparseAttentionIndexerWebGpuTest, CsaStateCapacityOverflowIsRejected) { + CsaPackedProblem problem; + problem.batch_size = 1; + problem.cumulative_sequence_lengths = {0, 4}; + problem.state_capacity = 1; + problem.past_state_lengths = {1, 0}; + RunCsaPackedTest(MakeCsaPackedProblem(std::move(problem)), 1.0e-5f, ProviderKind::WebGpu); +} +#endif + } // namespace test } // namespace onnxruntime From 5b9d10ebd776c303c2d34c4980d1831f1aa8c4e2 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 16 Sep 2026 02:35:03 +0000 Subject: [PATCH 04/10] Validate packed indexer capacity bounds Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- docs/ContribOperators.md | 2 +- .../sparse/packed_sparse_attention_indexer.cc | 7 ++++-- .../bert/packed_sparse_attention_indexer.cc | 5 ++++- .../core/graph/contrib_ops/bert_defs.cc | 12 ++++++---- ...packed_sparse_attention_indexer_op_test.cc | 22 +++++++++++++------ 5 files changed, 33 insertions(+), 15 deletions(-) diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index 875a3483f87d4..f94548e123ae0 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -4747,7 +4747,7 @@ This version of the operator has been available since version 1 of the 'com.micr
compress_ratio : int (required)
-
Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0.
+
Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0 and 2 * compress_ratio - 1 must not exceed INT_MAX.
epsilon : float
Epsilon of the RMS normalization applied to the compressed keys. Default is 1e-6.
head_weight_scale : float
diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc index 2b289bb80443e..7c4d32887f15c 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc @@ -97,8 +97,11 @@ PackedSparseAttentionIndexer::PackedSparseAttentionIndexer(const OpKernelInfo ORT_ENFORCE(info.GetAttr("compress_ratio", &compress_ratio_).IsOK(), "PackedSparseAttentionIndexer: compress_ratio is required"); - ORT_ENFORCE(compress_ratio_ > 0 && compress_ratio_ <= std::numeric_limits::max(), - "PackedSparseAttentionIndexer: compress_ratio must be in (0, INT_MAX], got ", compress_ratio_); + ORT_ENFORCE(compress_ratio_ > 0 && + compress_ratio_ <= (static_cast(std::numeric_limits::max()) + 1) / 2, + "PackedSparseAttentionIndexer: compress_ratio must be positive and produce a generic buffer capacity " + "no greater than INT_MAX, got ", + compress_ratio_); ORT_ENFORCE(info.GetAttr("state_capacity", &state_capacity_).IsOK(), "PackedSparseAttentionIndexer: state_capacity is required"); diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc index a277d7dc12264..7c13924c84789 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -656,7 +656,10 @@ PackedSparseAttentionIndexer::PackedSparseAttentionIndexer(const OpKernelInfo& i psai::kPolicyModeQsa, "' or '", psai::kPolicyModeCsa, "', got '", policy_mode, "'"); ORT_ENFORCE(info.GetAttr("compress_ratio", &compress_ratio_).IsOK(), "PackedSparseAttentionIndexer: compress_ratio is required"); - ORT_ENFORCE(compress_ratio_ > 0, "PackedSparseAttentionIndexer: compress_ratio must be > 0"); + ORT_ENFORCE(compress_ratio_ > 0 && + compress_ratio_ <= (static_cast(std::numeric_limits::max()) + 1) / 2, + "PackedSparseAttentionIndexer: compress_ratio must be positive and produce a generic buffer capacity " + "no greater than INT_MAX"); ORT_ENFORCE(info.GetAttr("state_capacity", &state_capacity_).IsOK() && state_capacity_ > 0, "PackedSparseAttentionIndexer: state_capacity must be a positive integer"); diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index cdc93a5cec6b2..6c9058e4a39f2 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -2331,9 +2331,12 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference const bool is_qsa = policy == psai::Policy::kQsa; const int64_t compress_ratio = getAttribute(ctx, "compress_ratio", static_cast(0)); - if (compress_ratio <= 0 || compress_ratio > std::numeric_limits::max()) { - fail_shape_inference("PackedSparseAttentionIndexer: compress_ratio must be in (0, INT_MAX], got ", - compress_ratio); + if (compress_ratio <= 0 || + compress_ratio > (static_cast(std::numeric_limits::max()) + 1) / 2) { + fail_shape_inference( + "PackedSparseAttentionIndexer: compress_ratio must be positive and produce a generic " + "buffer capacity no greater than INT_MAX, got ", + compress_ratio); } const int64_t state_capacity = getAttribute(ctx, "state_capacity", static_cast(0)); if (state_capacity <= 0 || state_capacity > std::numeric_limits::max()) { @@ -2532,7 +2535,8 @@ ONNX_MS_OPERATOR_SET_SCHEMA( "Indexer policy. Must be exactly 'qsa' (token indexer) or 'csa' (compressed block indexer).", AttributeProto::STRING) .Attr("compress_ratio", - "Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0.", + "Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0 and " + "2 * compress_ratio - 1 must not exceed INT_MAX.", AttributeProto::INT) .Attr("state_capacity", "Fixed capacity (number of entries) of past_key_state / present_key_state. Must be > 0.", diff --git a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc index 331247e8ded5b..9836440c9cfa6 100644 --- a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -184,13 +184,6 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, QsaInfersFixedCapacityAndSt ONNX_NAMESPACE::TensorProto_DataType_INT32, {options.batch_size, 2}); } -TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsStateCapacityShapeMismatch) { - GraphOptions options; - options.input_state_capacity = options.state_capacity + 1; - ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, - "past_key_state dimension 1 must equal state_capacity"); -} - TEST(PackedSparseAttentionIndexerShapeInferenceTest, CsaInfersFixedCapacityAndState) { GraphOptions options; options.policy_mode = psai::kPolicyModeCsa; @@ -228,6 +221,13 @@ void ExpectResolveFailure(const std::function& } // namespace +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsStateCapacityShapeMismatch) { + GraphOptions options; + options.input_state_capacity = options.state_capacity + 1; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "past_key_state dimension 1 must equal state_capacity"); +} + TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsUnknownPolicyMode) { GraphOptions options; options.policy_mode = "qsa_v2"; @@ -249,6 +249,14 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsZeroNumHeads) { "num_heads must be > 0"); } +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsGenericBufferCapacityOverflow) { + GraphOptions options; + options.policy_mode = psai::kPolicyModeCsa; + options.compress_ratio = static_cast(std::numeric_limits::max()) / 2 + 2; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "generic buffer capacity no greater than INT_MAX"); +} + TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsQsaTokenBudgetNotDivisibleByCompressRatio) { GraphOptions options; options.token_budget = 5; From de0e5874953f45033b1d7f170c76adb575856853 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 16 Sep 2026 02:48:00 +0000 Subject: [PATCH 05/10] Fix packed indexer output-count test Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- .../test/contrib_ops/packed_sparse_attention_indexer_op_test.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc index 9836440c9cfa6..53bccaa1f83e3 100644 --- a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -328,7 +328,7 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsWrongOutputCount) { GraphOptions options; options.output_count = 4; ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, - "exactly 6 declared outputs"); + "output size 4 not in range [min=6, max=6]"); } #endif // ORT_NO_EXCEPTIONS From df26284422c54c35f62761747c7a75323ec8868b Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 16 Sep 2026 19:08:55 +0000 Subject: [PATCH 06/10] Address packed indexer review feedback Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- .../packed_sparse_attention_indexer.md | 7 +- .../packed_sparse_attention_indexer_impl.cu | 8 +- .../bert/packed_sparse_attention_indexer.cc | 103 ++++++++--- .../bert/packed_sparse_attention_indexer.h | 22 ++- .../core/graph/contrib_ops/bert_defs.cc | 173 ++++++++++++++++-- ...packed_sparse_attention_indexer_op_test.cc | 147 +++++++++++++-- 6 files changed, 383 insertions(+), 77 deletions(-) diff --git a/docs/contrib_ops/packed_sparse_attention_indexer.md b/docs/contrib_ops/packed_sparse_attention_indexer.md index 712e73d2da278..883a074a657aa 100644 --- a/docs/contrib_ops/packed_sparse_attention_indexer.md +++ b/docs/contrib_ops/packed_sparse_attention_indexer.md @@ -117,9 +117,10 @@ Both policies read and write the *same four* state slots — there is no separat `CsaWindowPlan`/`TryComputeCsaWindowPlan`). State never grows. `present_*` always has exactly the same shape as `past_*`; only the *contents* -change. Input/output aliasing is supported: every kernel reads its sources (`past_kv_buffer` / -`key`, `past_gate_buffer` / `gate`) and never re-reads `present_*`, so it is correct whether -`present_*` is a distinct allocation or the same underlying buffer as `past_*`. +change. Input/output aliasing is supported. CUDA avoids unsafe buffer aliases; WebGPU omits an +aliased `past_*` read-only binding and reads the prior contents through the matching read-write +`present_*` binding. Each request is handled by one invocation, and buffer compaction reads entries +at or above the destination index before overwriting them. **State overflow.** If a call would close more blocks/windows than `state_capacity - old_entry_count` allows, that request's step is rejected as a deterministic diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu index 3e005ce00ba65..f2bcc3f29c2bf 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu @@ -183,7 +183,7 @@ __global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, con // Publish the raw trailing buffer. Skipped entirely on overflow: present_kv_buffer already // holds past_kv_buffer's contents unchanged (from the baseline copy in the Launch function // below), which is exactly the prior valid buffer this rejected step must preserve. - if (!overflowed) { + if (!rejected) { for (int t = static_cast(threadIdx.x); t < new_buf_len; t += static_cast(blockDim.x)) { const int virtual_pos = new_block_count * params.compress_ratio + t; const int64_t out_base = (static_cast(b) * params.buffer_capacity + t) * params.head_size; @@ -420,8 +420,8 @@ __global__ void CsaUpdateStateKernel(const T* key, const T* gate, const T* key_n const int old_key_len = min(max(raw_key_len, 0), params.state_capacity); const int old_buf_len = min(max(raw_buf_len, 0), params.buffer_capacity); - sai::CsaWindowPlan plan; - const bool plan_ok = sai::TryComputeCsaWindowPlan(old_buf_len, req_len, params.compress_ratio, plan); + psai::CsaWindowPlan plan; + const bool plan_ok = psai::TryComputeCsaWindowPlan(old_buf_len, req_len, params.compress_ratio, plan); // plan_ok is always true here: old_buf_len is clamped into [0, buffer_capacity) == // [0, 2 * compress_ratio) and req_len >= 0, which are exactly the documented preconditions. const int full_new_window_count = plan_ok ? static_cast(plan.new_window_count) : 0; @@ -568,7 +568,7 @@ __global__ void CsaUpdateStateKernel(const T* key, const T* gate, const T* key_n // present_gate_buffer already hold past_kv_buffer's / past_gate_buffer's contents unchanged // (from the baseline copy in the Launch function below), which is exactly the prior valid // buffer this rejected step must preserve. - if (!overflowed) { + if (!rejected) { for (int t = static_cast(threadIdx.x); t < present_buffer_length; t += static_cast(blockDim.x)) { const int virtual_pos = present_buffer_start + t; const int64_t out_base = (static_cast(b) * params.buffer_capacity + t) * width; diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc index 7c13924c84789..725e8a1c09d0f 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -37,7 +37,11 @@ ONNX_OPERATOR_KERNEL_EX( (*KernelDefBuilder::Create()) .TypeConstraint("T", WebGpuSupportedFloatTypes()) .TypeConstraint("I", DataTypeImpl::GetTensorType()) - .TypeConstraint("M", DataTypeImpl::GetTensorType()), + .TypeConstraint("M", DataTypeImpl::GetTensorType()) + .MayInplace(psai::kPastKeyState, psai::kPresentKeyState) + .MayInplace(psai::kPastKvBuffer, psai::kPresentKvBuffer) + .MayInplace(psai::kPastGateBuffer, psai::kPresentGateBuffer) + .MayInplace(psai::kPastStateLengths, psai::kPresentStateLengths), PackedSparseAttentionIndexer); namespace { @@ -98,20 +102,30 @@ Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHe const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); - const auto& past_kv_buffer = shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); - const auto& past_state_lengths = shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); + const ShaderVariableHelper* past_kv_buffer = nullptr; + if (!kv_buffer_aliases_) { + past_kv_buffer = &shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); + } + const ShaderVariableHelper* past_state_lengths = nullptr; + if (!state_lengths_aliases_) { + past_state_lengths = &shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); + } const auto& present_key_state = shader.AddOutput("present_key_state", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_kv_buffer = shader.AddOutput("present_kv_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_state_lengths = shader.AddOutput("present_state_lengths", ShaderUsage::UseUniform); const auto& overflow_flags = shader.AddOutput("overflow_flags", ShaderUsage::UseUniform); + const ShaderVariableHelper& kv_buffer_history = + kv_buffer_aliases_ ? present_kv_buffer : *past_kv_buffer; + const ShaderVariableHelper& state_lengths_history = + state_lengths_aliases_ ? present_state_lengths : *past_state_lengths; shader.AdditionalImplementation() << "fn extended_key(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, d: u32) -> f32 {\n" << " if (virtual_pos < old_buf_len) {\n" << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * uniforms.head_size + d;\n" - << " return f32(" << past_kv_buffer.GetByOffset("idx") << ");\n" + << " return f32(" << kv_buffer_history.GetByOffset("idx") << ");\n" << " }\n" << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * uniforms.head_size + d;\n" << " return f32(" << key.GetByOffset("idx2") << ");\n" @@ -162,8 +176,8 @@ Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHe << " if (b >= uniforms.batch_size || local_idx != 0u) { return; }\n" << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" << " let req_end = " << cu_seqlens.GetByOffset("b + 1u") << ";\n" - << " let old_key_len = " << past_state_lengths.GetByOffset("b * 2u") << ";\n" - << " let old_buf_len = " << past_state_lengths.GetByOffset("b * 2u + 1u") << ";\n" + << " let old_key_len = " << state_lengths_history.GetByOffset("b * 2u") << ";\n" + << " let old_buf_len = " << state_lengths_history.GetByOffset("b * 2u + 1u") << ";\n" << " let past_len = " << past_seqlens.GetByOffset("b") << ";\n" << " let invalid = " << cu_seqlens.GetByOffset("0") << " != 0 || " << cu_seqlens.GetByOffset("uniforms.batch_size") << " != i32(uniforms.total_tokens) || " @@ -364,9 +378,18 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe const auto& position_bias = shader.AddInput("position_bias", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); - const auto& past_kv_buffer = shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); - const auto& past_gate_buffer = shader.AddInput("past_gate_buffer", ShaderUsage::UseUniform); - const auto& past_state_lengths = shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); + const ShaderVariableHelper* past_kv_buffer = nullptr; + if (!kv_buffer_aliases_) { + past_kv_buffer = &shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); + } + const ShaderVariableHelper* past_gate_buffer = nullptr; + if (!gate_buffer_aliases_) { + past_gate_buffer = &shader.AddInput("past_gate_buffer", ShaderUsage::UseUniform); + } + const ShaderVariableHelper* past_state_lengths = nullptr; + if (!state_lengths_aliases_) { + past_state_lengths = &shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); + } const auto& present_key_state = shader.AddOutput("present_key_state", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_kv_buffer = @@ -375,13 +398,19 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe shader.AddOutput("present_gate_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_state_lengths = shader.AddOutput("present_state_lengths", ShaderUsage::UseUniform); const auto& overflow_flags = shader.AddOutput("overflow_flags", ShaderUsage::UseUniform); + const ShaderVariableHelper& kv_buffer_history = + kv_buffer_aliases_ ? present_kv_buffer : *past_kv_buffer; + const ShaderVariableHelper& gate_buffer_history = + gate_buffer_aliases_ ? present_gate_buffer : *past_gate_buffer; + const ShaderVariableHelper& state_lengths_history = + state_lengths_aliases_ ? present_state_lengths : *past_state_lengths; shader.AdditionalImplementation() << "fn extended_key(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, channel: u32) -> f32 {\n" << " let width = 2u * uniforms.head_size;\n" << " if (virtual_pos < old_buf_len) {\n" << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * width + channel;\n" - << " return f32(" << past_kv_buffer.GetByOffset("idx") << ");\n" + << " return f32(" << kv_buffer_history.GetByOffset("idx") << ");\n" << " }\n" << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\n" << " return f32(" << key.GetByOffset("idx2") << ");\n" @@ -390,7 +419,7 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe << " let width = 2u * uniforms.head_size;\n" << " if (virtual_pos < old_buf_len) {\n" << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * width + channel;\n" - << " return f32(" << past_gate_buffer.GetByOffset("idx") << ");\n" + << " return f32(" << gate_buffer_history.GetByOffset("idx") << ");\n" << " }\n" << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\n" << " return f32(" << gate.GetByOffset("idx2") << ");\n" @@ -446,8 +475,8 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe << " if (b >= uniforms.batch_size || local_idx != 0u) { return; }\n" << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" << " let req_end = " << cu_seqlens.GetByOffset("b + 1u") << ";\n" - << " let old_key_len = " << past_state_lengths.GetByOffset("b * 2u") << ";\n" - << " let old_buf_len = " << past_state_lengths.GetByOffset("b * 2u + 1u") << ";\n" + << " let old_key_len = " << state_lengths_history.GetByOffset("b * 2u") << ";\n" + << " let old_buf_len = " << state_lengths_history.GetByOffset("b * 2u + 1u") << ";\n" << " let invalid = " << cu_seqlens.GetByOffset("0") << " != 0 || " << cu_seqlens.GetByOffset("uniforms.batch_size") << " != i32(uniforms.total_tokens) || " << "req_start < 0 || req_end < req_start || req_end > i32(uniforms.total_tokens) || " @@ -806,19 +835,25 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont } if (batch_size > 0) { - PackedSparseAttentionIndexerQsaUpdateProgram update{rotary.batched}; - update.CacheHint(rotary.batched) + const bool kv_buffer_aliases = present_kv_buffer->DataRaw() == past_kv_buffer->DataRaw(); + const bool state_lengths_aliases = present_state_lengths->DataRaw() == past_state_lengths->DataRaw(); + PackedSparseAttentionIndexerQsaUpdateProgram update{rotary.batched, kv_buffer_aliases, state_lengths_aliases}; + update.CacheHint(rotary.batched, kv_buffer_aliases, state_lengths_aliases) .SetWorkgroupSize(kWorkgroupSize) .AddInputs({{key, ProgramTensorMetadataDependency::Type}, {norm, ProgramTensorMetadataDependency::Type}, {cos_cache, ProgramTensorMetadataDependency::Type}, {sin_cache, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}, - {past_seqlens, ProgramTensorMetadataDependency::Type}, - {past_kv_buffer, ProgramTensorMetadataDependency::Type}}) - .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) - .AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, - {present_kv_buffer, ProgramTensorMetadataDependency::Type}}) + {past_seqlens, ProgramTensorMetadataDependency::Type}}); + if (!kv_buffer_aliases) { + update.AddInput({past_kv_buffer, ProgramTensorMetadataDependency::Type}); + } + if (!state_lengths_aliases) { + update.AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}); + } + update.AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, + {present_kv_buffer, ProgramTensorMetadataDependency::Type}}) .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) .AddOutput({&overflow_flags, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32(batch_size)) @@ -966,8 +1001,12 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_ERROR(copy_if_needed(past_state_lengths, present_state_lengths)); if (batch_size > 0) { - PackedSparseAttentionIndexerCsaUpdateProgram update{rotary.batched}; - update.CacheHint(rotary.batched) + const bool kv_buffer_aliases = present_kv_buffer->DataRaw() == past_kv_buffer->DataRaw(); + const bool gate_buffer_aliases = present_gate_buffer->DataRaw() == past_gate_buffer->DataRaw(); + const bool state_lengths_aliases = present_state_lengths->DataRaw() == past_state_lengths->DataRaw(); + PackedSparseAttentionIndexerCsaUpdateProgram update{ + rotary.batched, kv_buffer_aliases, gate_buffer_aliases, state_lengths_aliases}; + update.CacheHint(rotary.batched, kv_buffer_aliases, gate_buffer_aliases, state_lengths_aliases) .SetWorkgroupSize(kWorkgroupSize) .AddInputs({{key, ProgramTensorMetadataDependency::Type}, {gate, ProgramTensorMetadataDependency::Type}, @@ -976,13 +1015,19 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont {sin_cache, ProgramTensorMetadataDependency::Type}, {position_bias, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}, - {past_seqlens, ProgramTensorMetadataDependency::Type}, - {past_kv_buffer, ProgramTensorMetadataDependency::Type}, - {past_gate_buffer, ProgramTensorMetadataDependency::Type}}) - .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) - .AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, - {present_kv_buffer, ProgramTensorMetadataDependency::Type}, - {present_gate_buffer, ProgramTensorMetadataDependency::Type}}) + {past_seqlens, ProgramTensorMetadataDependency::Type}}); + if (!kv_buffer_aliases) { + update.AddInput({past_kv_buffer, ProgramTensorMetadataDependency::Type}); + } + if (!gate_buffer_aliases) { + update.AddInput({past_gate_buffer, ProgramTensorMetadataDependency::Type}); + } + if (!state_lengths_aliases) { + update.AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}); + } + update.AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, + {present_kv_buffer, ProgramTensorMetadataDependency::Type}, + {present_gate_buffer, ProgramTensorMetadataDependency::Type}}) .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) .AddOutput({&overflow_flags, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32(batch_size)) diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h index 7be03c2a75310..1fafc10110861 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h @@ -31,8 +31,12 @@ class PackedSparseAttentionIndexerCopyProgram final class PackedSparseAttentionIndexerQsaUpdateProgram final : public Program { public: - explicit PackedSparseAttentionIndexerQsaUpdateProgram(bool cos_cache_batched) - : Program{"PackedSparseAttentionIndexerQsaUpdate"}, cos_cache_batched_{cos_cache_batched} {} + PackedSparseAttentionIndexerQsaUpdateProgram(bool cos_cache_batched, bool kv_buffer_aliases, + bool state_lengths_aliases) + : Program{"PackedSparseAttentionIndexerQsaUpdate"}, + cos_cache_batched_{cos_cache_batched}, + kv_buffer_aliases_{kv_buffer_aliases}, + state_lengths_aliases_{state_lengths_aliases} {} Status GenerateShaderCode(ShaderHelper& shader) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( {"batch_size", ProgramUniformVariableDataType::Uint32}, @@ -47,6 +51,8 @@ class PackedSparseAttentionIndexerQsaUpdateProgram final private: bool cos_cache_batched_; + bool kv_buffer_aliases_; + bool state_lengths_aliases_; }; // One invocation per query token: rotates the query, scores it against every causally visible @@ -84,8 +90,13 @@ class PackedSparseAttentionIndexerQsaSelectProgram final class PackedSparseAttentionIndexerCsaUpdateProgram final : public Program { public: - explicit PackedSparseAttentionIndexerCsaUpdateProgram(bool cos_cache_batched) - : Program{"PackedSparseAttentionIndexerCsaUpdate"}, cos_cache_batched_{cos_cache_batched} {} + PackedSparseAttentionIndexerCsaUpdateProgram(bool cos_cache_batched, bool kv_buffer_aliases, + bool gate_buffer_aliases, bool state_lengths_aliases) + : Program{"PackedSparseAttentionIndexerCsaUpdate"}, + cos_cache_batched_{cos_cache_batched}, + kv_buffer_aliases_{kv_buffer_aliases}, + gate_buffer_aliases_{gate_buffer_aliases}, + state_lengths_aliases_{state_lengths_aliases} {} Status GenerateShaderCode(ShaderHelper& shader) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( {"batch_size", ProgramUniformVariableDataType::Uint32}, @@ -100,6 +111,9 @@ class PackedSparseAttentionIndexerCsaUpdateProgram final private: bool cos_cache_batched_; + bool kv_buffer_aliases_; + bool gate_buffer_aliases_; + bool state_lengths_aliases_; }; // One invocation per query token: rotates the query, scores it against every causally visible diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index ae35cd73f11fe..dce092fd0a7b8 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -4,6 +4,7 @@ #include #include #include +#include #include #include "core/graph/constants.h" @@ -2395,6 +2396,11 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference fail_shape_inference("PackedSparseAttentionIndexer: exactly ", psai::kFixedOutputCount, " declared outputs are required, got ", ctx.getNumOutputs()); } + if (ctx.hasOutput(psai::kPresentGateBuffer) == is_qsa) { + fail_shape_inference("PackedSparseAttentionIndexer: output ", psai::kPresentGateBuffer, + is_qsa ? " must be omitted when policy_mode is 'qsa'" + : " is required when policy_mode is 'csa'"); + } updateOutputElemType(ctx, psai::kSelectedIndices, ONNX_NAMESPACE::TensorProto_DataType_INT32); updateOutputElemType(ctx, psai::kSelectedCounts, ONNX_NAMESPACE::TensorProto_DataType_INT32); @@ -2407,20 +2413,156 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference propagateElemTypeFromInputToOutput(ctx, psai::kPastGateBuffer, psai::kPresentGateBuffer); } - (void)PackedSparseAttentionIndexerShape(ctx, psai::kKeyNormWeight, 1); - (void)PackedSparseAttentionIndexerShape(ctx, psai::kKey, 2); - (void)PackedSparseAttentionIndexerShape(ctx, psai::kCumulativeSequenceLengths, 1); - (void)PackedSparseAttentionIndexerShape(ctx, psai::kPastSequenceLengths, 1); + const auto* query_shape = PackedSparseAttentionIndexerShape(ctx, psai::kQuery, 3); + const auto* key_shape = PackedSparseAttentionIndexerShape(ctx, psai::kKey, 2); + const auto* norm_shape = PackedSparseAttentionIndexerShape(ctx, psai::kKeyNormWeight, 1); + const auto* cumulative_shape = + PackedSparseAttentionIndexerShape(ctx, psai::kCumulativeSequenceLengths, 1); + const auto* past_sequence_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastSequenceLengths, 1); + const ONNX_NAMESPACE::TensorShapeProto* cos_shape = nullptr; + const ONNX_NAMESPACE::TensorShapeProto* sin_shape = nullptr; + for (const auto [index, name, shape_out] : + {std::tuple{ + psai::kCosCache, "cos_cache", &cos_shape}, + {psai::kSinCache, "sin_cache", &sin_shape}}) { + if (hasInputShape(ctx, index)) { + const auto& shape = getInputShape(ctx, index); + if (shape.dim_size() != 2 && shape.dim_size() != 3) { + fail_shape_inference("PackedSparseAttentionIndexer: ", name, " must have rank 2 or 3, got rank ", + shape.dim_size()); + } + *shape_out = &shape; + } + } + const auto* key_state_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKeyState, 3); + const auto* kv_buffer_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKvBuffer, 3); + const auto* state_lengths_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastStateLengths, 2); + const ONNX_NAMESPACE::TensorShapeProto* gate_shape = nullptr; + const ONNX_NAMESPACE::TensorShapeProto* position_bias_shape = nullptr; + const ONNX_NAMESPACE::TensorShapeProto* head_weights_shape = nullptr; + const ONNX_NAMESPACE::TensorShapeProto* gate_buffer_shape = nullptr; if (!is_qsa) { - (void)PackedSparseAttentionIndexerShape(ctx, psai::kGate, 2); - (void)PackedSparseAttentionIndexerShape(ctx, psai::kPositionBias, 2); - (void)PackedSparseAttentionIndexerShape(ctx, psai::kHeadWeights, 2); + gate_shape = PackedSparseAttentionIndexerShape(ctx, psai::kGate, 2); + position_bias_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPositionBias, 2); + head_weights_shape = PackedSparseAttentionIndexerShape(ctx, psai::kHeadWeights, 2); + gate_buffer_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastGateBuffer, 3); } + const ONNX_NAMESPACE::TensorShapeProto* position_ids_shape = nullptr; if (PackedSparseAttentionIndexerHasInput(ctx, psai::kPositionIds)) { - (void)PackedSparseAttentionIndexerShape(ctx, psai::kPositionIds, 1); + position_ids_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPositionIds, 1); + } + + auto require_equal_dims = [](const ONNX_NAMESPACE::TensorShapeProto* lhs, int lhs_index, + const ONNX_NAMESPACE::TensorShapeProto* rhs, int rhs_index, + const char* description) { + if (lhs != nullptr && rhs != nullptr && lhs->dim(lhs_index).has_dim_value() && + rhs->dim(rhs_index).has_dim_value() && + lhs->dim(lhs_index).dim_value() != rhs->dim(rhs_index).dim_value()) { + fail_shape_inference("PackedSparseAttentionIndexer: ", description); + } + }; + auto require_dim_value = [](const ONNX_NAMESPACE::TensorShapeProto* shape, int index, int64_t expected, + const char* description) { + if (shape != nullptr && shape->dim(index).has_dim_value() && shape->dim(index).dim_value() != expected) { + fail_shape_inference("PackedSparseAttentionIndexer: ", description, " (", expected, "), got ", + shape->dim(index).dim_value()); + } + }; + + const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio); + require_equal_dims(key_shape, 0, query_shape, 0, "key dimension 0 must equal query dimension 0"); + require_equal_dims(norm_shape, 0, query_shape, 2, "key_norm_weight dimension 0 must equal head_size"); + require_equal_dims(past_sequence_shape, 0, key_state_shape, 0, + "past_sequence_lengths dimension 0 must equal the state batch dimension"); + require_equal_dims(key_state_shape, 0, kv_buffer_shape, 0, + "past_key_state and past_kv_buffer batch dimensions must match"); + require_equal_dims(state_lengths_shape, 0, key_state_shape, 0, + "past_state_lengths dimension 0 must equal the state batch dimension"); + require_equal_dims(key_state_shape, 2, query_shape, 2, + "past_key_state dimension 2 must equal head_size"); + require_dim_value(key_state_shape, 1, state_capacity, "past_key_state dimension 1 must equal state_capacity"); + require_dim_value(kv_buffer_shape, 1, buffer_capacity, + "past_kv_buffer dimension 1 must equal 2 * compress_ratio - 1"); + require_dim_value(state_lengths_shape, 1, psai::kStateLengthColumns, + "past_state_lengths dimension 1 must equal 2"); + if (cumulative_shape != nullptr && past_sequence_shape != nullptr && + cumulative_shape->dim(0).has_dim_value() && past_sequence_shape->dim(0).has_dim_value() && + cumulative_shape->dim(0).dim_value() != past_sequence_shape->dim(0).dim_value() + 1) { + fail_shape_inference( + "PackedSparseAttentionIndexer: cumulative_sequence_lengths dimension 0 must equal batch_size + 1"); + } + + if (query_shape != nullptr && query_shape->dim(2).has_dim_value()) { + const int64_t head_size = query_shape->dim(2).dim_value(); + if (head_size <= 0) { + fail_shape_inference("PackedSparseAttentionIndexer: head_size must be > 0, got ", head_size); + } + if (!is_qsa && head_size > std::numeric_limits::max() / 2) { + fail_shape_inference("PackedSparseAttentionIndexer: 2 * head_size exceeds INT64_MAX"); + } + const int64_t width = is_qsa ? head_size : 2 * head_size; + require_dim_value(key_shape, 1, width, "key dimension 1 must equal the policy-specific width"); + require_dim_value(kv_buffer_shape, 2, width, + "past_kv_buffer dimension 2 must equal the policy-specific width"); + if (!is_qsa) { + require_dim_value(gate_shape, 1, width, "gate dimension 1 must equal 2 * head_size"); + require_dim_value(position_bias_shape, 1, width, "position_bias dimension 1 must equal 2 * head_size"); + require_dim_value(gate_buffer_shape, 2, width, "past_gate_buffer dimension 2 must equal 2 * head_size"); + } + } + if (!is_qsa) { + require_equal_dims(gate_shape, 0, query_shape, 0, "gate dimension 0 must equal query dimension 0"); + require_equal_dims(head_weights_shape, 0, query_shape, 0, + "head_weights dimension 0 must equal query dimension 0"); + require_equal_dims(head_weights_shape, 1, query_shape, 1, + "head_weights dimension 1 must equal num_heads"); + require_dim_value(position_bias_shape, 0, compress_ratio, + "position_bias dimension 0 must equal compress_ratio"); + require_equal_dims(gate_buffer_shape, 0, kv_buffer_shape, 0, + "past_gate_buffer and past_kv_buffer batch dimensions must match"); + require_equal_dims(gate_buffer_shape, 1, kv_buffer_shape, 1, + "past_gate_buffer and past_kv_buffer capacities must match"); + require_equal_dims(gate_buffer_shape, 2, kv_buffer_shape, 2, + "past_gate_buffer and past_kv_buffer widths must match"); + } + require_equal_dims(position_ids_shape, 0, query_shape, 0, + "position_ids dimension 0 must equal query dimension 0"); + if (cos_shape != nullptr && sin_shape != nullptr) { + if (cos_shape->dim_size() != sin_shape->dim_size()) { + fail_shape_inference("PackedSparseAttentionIndexer: cos_cache and sin_cache ranks must match"); + } + for (int i = 0; i < cos_shape->dim_size(); ++i) { + require_equal_dims(cos_shape, i, sin_shape, i, "cos_cache and sin_cache dimensions must match"); + } + } + for (const auto* cache_shape : {cos_shape, sin_shape}) { + if (cache_shape == nullptr) { + continue; + } + if (cache_shape->dim_size() == 3) { + require_equal_dims(cache_shape, 0, past_sequence_shape, 0, + "batched rotary cache dimension 0 must equal batch_size"); + } + const int position_dim = cache_shape->dim_size() - 2; + const int rotary_dim = cache_shape->dim_size() - 1; + if (cache_shape->dim(position_dim).has_dim_value() && cache_shape->dim(position_dim).dim_value() <= 0) { + fail_shape_inference("PackedSparseAttentionIndexer: rotary cache max_position must be > 0"); + } + if (cache_shape->dim(rotary_dim).has_dim_value()) { + const int64_t rotary_width = cache_shape->dim(rotary_dim).dim_value(); + if (rotary_width <= 0 || (is_qsa && rotary_width % 2 != 0)) { + fail_shape_inference("PackedSparseAttentionIndexer: invalid rotary cache width ", rotary_width); + } + if (query_shape != nullptr && query_shape->dim(2).has_dim_value()) { + const int64_t head_size = query_shape->dim(2).dim_value(); + if ((is_qsa && rotary_width > head_size) || + (!is_qsa && rotary_width > head_size / 2)) { + fail_shape_inference("PackedSparseAttentionIndexer: rotary cache width is incompatible with head_size"); + } + } + } } - const auto* query_shape = PackedSparseAttentionIndexerShape(ctx, psai::kQuery, 3); if (query_shape != nullptr) { const auto& total_tokens_dim = query_shape->dim(0); const auto& num_heads_dim = query_shape->dim(1); @@ -2440,22 +2582,15 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference } // State never grows: present_* always has exactly the same fixed shape as past_*. - const auto* key_state_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKeyState, 3); if (key_state_shape != nullptr) { - if (key_state_shape->dim(1).has_dim_value() && key_state_shape->dim(1).dim_value() != state_capacity) { - fail_shape_inference("PackedSparseAttentionIndexer: past_key_state dimension 1 must equal state_capacity (", - state_capacity, "), got ", key_state_shape->dim(1).dim_value()); - } updateOutputShape(ctx, psai::kPresentKeyState, *key_state_shape); } - const auto* kv_buffer_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastKvBuffer, 3); if (kv_buffer_shape != nullptr) { updateOutputShape(ctx, psai::kPresentKvBuffer, *kv_buffer_shape); - if (!is_qsa) { - updateOutputShape(ctx, psai::kPresentGateBuffer, *kv_buffer_shape); - } } - const auto* state_lengths_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastStateLengths, 2); + if (gate_buffer_shape != nullptr) { + updateOutputShape(ctx, psai::kPresentGateBuffer, *gate_buffer_shape); + } if (state_lengths_shape != nullptr) { updateOutputShape(ctx, psai::kPresentStateLengths, *state_lengths_shape); } diff --git a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc index 53bccaa1f83e3..ba7c3baf2fca3 100644 --- a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -93,9 +93,13 @@ struct GraphOptions { int64_t state_capacity = 6; int64_t input_state_capacity = -1; int64_t token_budget = 4; + int64_t key_total_tokens = -1; + int64_t kv_buffer_capacity = -1; + int64_t gate_buffer_width = -1; bool add_index_topk = false; bool add_csa_inputs = false; bool add_position_ids = false; + bool invert_gate_output_presence = false; int output_count = psai::kFixedOutputCount; std::string policy_mode = psai::kPolicyModeQsa; }; @@ -113,7 +117,8 @@ void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { std::vector inputs{ builder.MakeInput( std::vector{options.total_tokens, options.num_heads, options.head_size}), - builder.MakeInput(std::vector{options.total_tokens, width}), + builder.MakeInput( + std::vector{options.key_total_tokens >= 0 ? options.key_total_tokens : options.total_tokens, width}), builder.MakeInput(std::vector{options.head_size}), builder.MakeInput(std::vector{64, options.rotary_width}), builder.MakeInput(std::vector{64, options.rotary_width}), @@ -139,9 +144,12 @@ void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { inputs.push_back( builder.MakeInput(std::vector{options.batch_size, input_state_capacity, options.head_size})); inputs.push_back( - builder.MakeInput(std::vector{options.batch_size, buffer_capacity, width})); + builder.MakeInput(std::vector{ + options.batch_size, options.kv_buffer_capacity >= 0 ? options.kv_buffer_capacity : buffer_capacity, width})); if (is_csa) { - inputs.push_back(builder.MakeInput(std::vector{options.batch_size, buffer_capacity, width})); + inputs.push_back(builder.MakeInput( + std::vector{options.batch_size, buffer_capacity, + options.gate_buffer_width >= 0 ? options.gate_buffer_width : width})); } else { inputs.push_back(&empty); } @@ -149,7 +157,8 @@ void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { std::vector outputs; for (int i = 0; i < options.output_count; ++i) { - outputs.push_back(i == psai::kPresentGateBuffer && !is_csa ? &empty : builder.MakeOutput()); + const bool gate_output_present = is_csa != options.invert_gate_output_presence; + outputs.push_back(i == psai::kPresentGateBuffer && !gate_output_present ? &empty : builder.MakeOutput()); } Node& node = builder.AddNode("PackedSparseAttentionIndexer", inputs, outputs, kMSDomain); node.AddAttribute("policy_mode", options.policy_mode); @@ -228,6 +237,43 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsStateCapacityShapeMi "past_key_state dimension 1 must equal state_capacity"); } +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsQsaGateOutput) { + GraphOptions options; + options.invert_gate_output_presence = true; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "must be omitted when policy_mode is 'qsa'"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsCsaMissingGateOutput) { + GraphOptions options; + options.policy_mode = psai::kPolicyModeCsa; + options.invert_gate_output_presence = true; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "is required when policy_mode is 'csa'"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsMismatchedKeyTokenCount) { + GraphOptions options; + options.key_total_tokens = options.total_tokens + 1; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "key dimension 0 must equal query dimension 0"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsWrongKvBufferCapacity) { + GraphOptions options; + options.kv_buffer_capacity = BufferCapacity(options.compress_ratio) + 1; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "past_kv_buffer dimension 1 must equal 2 * compress_ratio - 1"); +} + +TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsMismatchedCsaGateBufferWidth) { + GraphOptions options; + options.policy_mode = psai::kPolicyModeCsa; + options.gate_buffer_width = 2 * options.head_size + 1; + ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, + "past_gate_buffer dimension 2 must equal 2 * head_size"); +} + TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsUnknownPolicyMode) { GraphOptions options; options.policy_mode = "qsa_v2"; @@ -395,6 +441,20 @@ std::vector RoundTrip(const std::vector& data) { } } +template +std::vector CopyTensorToFloat(const Tensor& tensor) { + const auto values = tensor.DataAsSpan(); + std::vector result(values.size()); + for (size_t i = 0; i < values.size(); ++i) { + if constexpr (std::is_same_v) { + result[i] = values[i]; + } else { + result[i] = values[i].ToFloat(); + } + } + return result; +} + // Split-half rotary over the leading rotary_width channels. std::vector LeadingRope(const std::vector& value, int rotary_width, const float* cos_row, const float* sin_row) { @@ -645,7 +705,7 @@ QsaPackedProblem MakeQsaPackedProblem(QsaPackedProblem problem = {}) { template void RunQsaPackedTest(float tolerance, QsaPackedProblem problem = MakeQsaPackedProblem(), - ProviderKind provider_kind = ProviderKind::Cuda) { + ProviderKind provider_kind = ProviderKind::Cuda, QsaPackedResult* actual = nullptr) { auto provider = CreateProvider(provider_kind); if (provider == nullptr) { GTEST_SKIP() << (provider_kind == ProviderKind::Cuda ? "CUDA" : "WebGPU") @@ -703,6 +763,13 @@ void RunQsaPackedTest(float tolerance, QsaPackedProblem problem = MakeQsaPackedP test.AddOptionalOutputEdge(); // present_gate_buffer test.AddOutput("present_state_lengths", {batch_size, 2}, expected.present_state_lengths); RunOnProvider(test, std::move(provider)); + if (actual != nullptr) { + const auto& fetches = test.GetFetches(); + actual->present_key_state = CopyTensorToFloat(fetches[2].Get()); + actual->present_kv_buffer = CopyTensorToFloat(fetches[3].Get()); + actual->present_state_lengths.assign(fetches[4].Get().DataAsSpan().begin(), + fetches[4].Get().DataAsSpan().end()); + } } } // namespace @@ -713,13 +780,26 @@ TEST(PackedSparseAttentionIndexerTest, QsaFloat16) { RunQsaPackedTest TEST(PackedSparseAttentionIndexerTest, QsaBFloat16) { RunQsaPackedTest(2.0e-2f); } -// Prefill followed by decode: request 0 continues an existing 3-token history, request 1 starts -// fresh; the two requests keep independent block/tail state. TEST(PackedSparseAttentionIndexerTest, QsaPrefillThenDecodeIndependentState) { - QsaPackedProblem problem; - problem.cumulative_sequence_lengths = {0, 1, 4}; - problem.past_sequence_lengths = {5, 0}; - RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(std::move(problem))); + if (DefaultCudaExecutionProvider() == nullptr) { + GTEST_SKIP() << "CUDA execution provider is not available"; + } + + QsaPackedProblem prefill; + prefill.cumulative_sequence_lengths = {0, 3, 5}; + prefill.past_sequence_lengths = {0, 0}; + QsaPackedResult prefill_outputs; + RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(std::move(prefill)), ProviderKind::Cuda, + &prefill_outputs); + + QsaPackedProblem decode; + decode.cumulative_sequence_lengths = {0, 1, 3}; + decode.past_sequence_lengths = {3, 2}; + decode = MakeQsaPackedProblem(std::move(decode)); + decode.past_key_state = std::move(prefill_outputs.present_key_state); + decode.past_kv_buffer = std::move(prefill_outputs.present_kv_buffer); + decode.past_state_lengths = std::move(prefill_outputs.present_state_lengths); + RunQsaPackedTest(1.0e-5f, std::move(decode)); } // A zero-token request row (repeated cumulative offset) must not affect the other request and @@ -749,6 +829,10 @@ TEST(PackedSparseAttentionIndexerWebGpuTest, QsaFloat) { RunQsaPackedTest(1.0e-5f, MakeQsaPackedProblem(), ProviderKind::WebGpu); } +TEST(PackedSparseAttentionIndexerWebGpuTest, QsaFloat16) { + RunQsaPackedTest(2.0e-3f, MakeQsaPackedProblem(), ProviderKind::WebGpu); +} + TEST(PackedSparseAttentionIndexerWebGpuTest, QsaStateCapacityOverflowIsRejected) { QsaPackedProblem problem; problem.batch_size = 1; @@ -1005,7 +1089,7 @@ CsaPackedProblem MakeCsaPackedProblem(CsaPackedProblem problem = {}) { template void RunCsaPackedTest(const CsaPackedProblem& base, float tolerance, - ProviderKind provider_kind = ProviderKind::Cuda) { + ProviderKind provider_kind = ProviderKind::Cuda, CsaPackedResult* actual = nullptr) { auto provider = CreateProvider(provider_kind); if (provider == nullptr) { GTEST_SKIP() << (provider_kind == ProviderKind::Cuda ? "CUDA" : "WebGPU") @@ -1070,6 +1154,14 @@ void RunCsaPackedTest(const CsaPackedProblem& base, float tolerance, ToElementType(expected.present_gate_buffer), false, 0.0f, tolerance); test.AddOutput("present_state_lengths", {batch_size, 2}, expected.present_state_lengths); RunOnProvider(test, std::move(provider)); + if (actual != nullptr) { + const auto& fetches = test.GetFetches(); + actual->present_key_state = CopyTensorToFloat(fetches[2].Get()); + actual->present_kv_buffer = CopyTensorToFloat(fetches[3].Get()); + actual->present_gate_buffer = CopyTensorToFloat(fetches[4].Get()); + actual->present_state_lengths.assign(fetches[5].Get().DataAsSpan().begin(), + fetches[5].Get().DataAsSpan().end()); + } } } // namespace @@ -1080,13 +1172,28 @@ TEST(PackedSparseAttentionIndexerTest, CsaFloat16) { RunCsaPackedTest TEST(PackedSparseAttentionIndexerTest, CsaBFloat16) { RunCsaPackedTest(MakeCsaPackedProblem(), 3.0e-2f); } -// Prefill followed by decode: continues a previously compressed entry and a partially filled -// buffer for one request, while another request starts fresh. TEST(PackedSparseAttentionIndexerTest, CsaPrefillThenDecodeIndependentState) { - CsaPackedProblem problem; - problem.cumulative_sequence_lengths = {0, 1, 4}; - problem.past_state_lengths = {1, 1, 0, 0}; - RunCsaPackedTest(MakeCsaPackedProblem(std::move(problem)), 1.0e-5f); + if (DefaultCudaExecutionProvider() == nullptr) { + GTEST_SKIP() << "CUDA execution provider is not available"; + } + + CsaPackedProblem prefill; + prefill.cumulative_sequence_lengths = {0, 3, 5}; + prefill.past_sequence_lengths = {0, 0}; + CsaPackedResult prefill_outputs; + RunCsaPackedTest(MakeCsaPackedProblem(std::move(prefill)), 1.0e-5f, ProviderKind::Cuda, + &prefill_outputs); + + CsaPackedProblem decode; + decode.cumulative_sequence_lengths = {0, 1, 3}; + decode.past_sequence_lengths = {3, 2}; + decode.position_ids = {3, 2, 3}; + decode = MakeCsaPackedProblem(std::move(decode)); + decode.past_key_state = std::move(prefill_outputs.present_key_state); + decode.past_kv_buffer = std::move(prefill_outputs.present_kv_buffer); + decode.past_gate_buffer = std::move(prefill_outputs.present_gate_buffer); + decode.past_state_lengths = std::move(prefill_outputs.present_state_lengths); + RunCsaPackedTest(decode, 1.0e-5f); } TEST(PackedSparseAttentionIndexerTest, CsaStateCapacityOverflowIsRejected) { @@ -1103,6 +1210,10 @@ TEST(PackedSparseAttentionIndexerWebGpuTest, CsaFloat) { RunCsaPackedTest(MakeCsaPackedProblem(), 1.0e-5f, ProviderKind::WebGpu); } +TEST(PackedSparseAttentionIndexerWebGpuTest, CsaFloat16) { + RunCsaPackedTest(MakeCsaPackedProblem(), 4.0e-3f, ProviderKind::WebGpu); +} + TEST(PackedSparseAttentionIndexerWebGpuTest, CsaStateCapacityOverflowIsRejected) { CsaPackedProblem problem; problem.batch_size = 1; From 8ece210cf3df0bc1c6b9ac883b4e2e99b65a2cf8 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 16 Sep 2026 21:16:59 +0000 Subject: [PATCH 07/10] Fix packed indexer CI failures Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- onnxruntime/core/graph/contrib_ops/bert_defs.cc | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index dce092fd0a7b8..37ce43b089b24 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -2396,7 +2396,8 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference fail_shape_inference("PackedSparseAttentionIndexer: exactly ", psai::kFixedOutputCount, " declared outputs are required, got ", ctx.getNumOutputs()); } - if (ctx.hasOutput(psai::kPresentGateBuffer) == is_qsa) { + const bool has_gate_output = ctx.getOutputType(psai::kPresentGateBuffer) != nullptr; + if (has_gate_output == is_qsa) { fail_shape_inference("PackedSparseAttentionIndexer: output ", psai::kPresentGateBuffer, is_qsa ? " must be omitted when policy_mode is 'qsa'" : " is required when policy_mode is 'csa'"); @@ -2421,7 +2422,7 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference const auto* past_sequence_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastSequenceLengths, 1); const ONNX_NAMESPACE::TensorShapeProto* cos_shape = nullptr; const ONNX_NAMESPACE::TensorShapeProto* sin_shape = nullptr; - for (const auto [index, name, shape_out] : + for (const auto& [index, name, shape_out] : {std::tuple{ psai::kCosCache, "cos_cache", &cos_shape}, {psai::kSinCache, "sin_cache", &sin_shape}}) { From d0b6fafab8ca67b9fcfaa7ffe1d44e06b7414ae9 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 17 Sep 2026 01:30:45 +0000 Subject: [PATCH 08/10] Fix packed indexer schema validation Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- docs/ContribOperators.md | 138 +++++++++++++----- docs/OperatorKernels.md | 12 +- .../sparse/packed_sparse_attention_indexer.cc | 4 + .../bert/packed_sparse_attention_indexer.cc | 4 + .../core/graph/contrib_ops/bert_defs.cc | 7 - ...packed_sparse_attention_indexer_op_test.cc | 21 +-- 6 files changed, 114 insertions(+), 72 deletions(-) diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index 64988ff7f73c6..f3d97c1cc4599 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -1813,8 +1813,10 @@ This version of the operator has been available since version 1 of the 'com.micr gate = sigmoid(sign(dot) * sqrt(max(abs(dot), 1e-6))) where dot = sum(RMSNorm(key) * RMSNorm(query)) / sqrt(hidden_size). - The output is gate * value, broadcast across the hyper-connections. The final Engram residual - value + short_conv(value) is then expressed with RMSNorm, CausalConvWithState and Add. + The output is gate * value, broadcast across the hyper-connections. The optional gated_value_normed + output applies RMSNorm to gate * value with conv_norm_scale, which can feed a following + CausalConvWithState. The final Engram residual value + short_conv(value) is then expressed with + RMSNorm, CausalConvWithState and Add. #### Version @@ -1827,7 +1829,7 @@ This version of the operator has been available since version 1 of the 'com.micr
Epsilon used by both RMS normalization steps. Default is 1e-5.
-#### Inputs +#### Inputs (5 - 6)
key : T
@@ -1840,13 +1842,17 @@ This version of the operator has been available since version 1 of the 'com.micr
RMSNorm scale for keys with shape (hc_mult, hidden_size).
query_norm_scale : T
RMSNorm scale for queries with shape (hc_mult, hidden_size).
+
conv_norm_scale (optional) : T
+
Optional RMSNorm scale for the gated value, with shape (hc_mult, hidden_size). Required when gated_value_normed is requested.
-#### Outputs +#### Outputs (1 - 2)
output : T
Gated value tensor with shape (batch_size, sequence_length, hc_mult, hidden_size).
+
gated_value_normed (optional) : T
+
Optional RMS-normalized gated value tensor with shape (batch_size, sequence_length, hc_mult, hidden_size).
#### Type Constraints @@ -2435,13 +2441,22 @@ This version of the operator has been available since version 1 of the 'com.micr GatherBlockQuantized is a Gather with data quantized. It is similar to Gather (https://github.com/onnx/onnx/blob/main/docs/Operators.md#gather) with differences: 1. Input `data` is a constant. It is quantized block-wise along attribute `quantize_axis` with block size specified by attribute `block_size`. `block_size` must be a power of 2 and not smaller than 16, like 16, 32, 64, 128, ... + For an FP8 or FP4 `data` type (see point 6 below), `block_size` may also be 0, meaning the entire `quantize_axis` + dimension forms a single block (i.e. one scale per row). 2. Input `data`'s scale and zero point are specified by input `scales` and `zero_points`. `scales` and `zero_points` are also constants. If `zero_points` is not provided, the default value is 0 for int4/uint4, or 2^(bits-1) for uint8. + `zero_points` must not be provided when `data` is an FP8 or FP4 type: FP8/FP4 quantization is symmetric. 3. During the op execution, `data` and `indices` are first used to generate the quantized output. Then, `scales` and `zero_points` are used to dequantize the output. 4. The `output` and `scales` have the same type. The `data` and `zero_points` have the same type. 5. For uint8 data, the `gather_axis` must be 0. The supported `bits` values for uint8 data are 2, 4, and 8; for `bits` < 8 the values are packed along the last dimension (low-order bits first). + 6. `data` may also be an FP8 type (float8e4m3fn, float8e4m3fnuz, float8e5m2 or float8e5m2fnuz) or an FP4 type + (float4e2m1), rather than an integer block-quantized type. In that case `bits` is ignored, there is + no `zero_points` input, and dequantization is simply `output[...] = float(data[...]) * scales[block_index(...)]`. + On any axis other than `quantize_axis`, the corresponding `scales` dimension must either equal `data`'s + dimension, or be 1, in which case the scale is broadcast along that axis (e.g. a single scale shared by + every row, as with a per-tensor scale applied to an entire embedding table). #### Version @@ -2451,9 +2466,9 @@ This version of the operator has been available since version 1 of the 'com.micr
bits : int
-
Number of bits used for weight quantization. Must be 2, 4 or 8.
+
Number of bits used for weight quantization. Must be 2, 4 or 8. Ignored when `data` is an FP8 or FP4 type.
block_size : int
-
(Optional) block size used for weight quantization. It needs to be a power of 2 and not smaller than 16.
+
(Optional) block size used for weight quantization. It needs to be a power of 2 and not smaller than 16, or 0. A value of 0 is only valid for an FP8 or FP4 `data` type and means the entire `quantize_axis` dimension forms a single block.
gather_axis : int
(Optional) Which axis to gather on. Negative value means counting dimensions from the back. Accepted range is [-r, r-1] where r = rank(data).
quantize_axis : int
@@ -2466,11 +2481,11 @@ This version of the operator has been available since version 1 of the 'com.micr
data : T1
Tensor of rank r >= 1. Block-wise quantized.
indices : Tind
-
Tensor of int32/int64 indices, of any rank q. All index values are expected to be within bounds [-s, s-1] along axis of size s. It is an error if any of the index values are out of bounds.
+
Tensor of int32/int64 indices, of any rank q. Values in [-s, s-1] select elements along an axis of size s. Unlike ONNX Gather, an out-of-range index produces zeros for the corresponding output slice.
scales : T2
-
quantization scale
+
quantization scale. Same rank as data. On axes other than quantize_axis, a dimension of 1 broadcasts the scale along that axis (e.g. a single per-tensor scale for the whole table); only applicable when `data` is an FP8 or FP4 type.
zero_points (optional) : T1
-
quantization zero points
+
quantization zero points. Must not be provided when `data` is an FP8 or FP4 type.
#### Outputs @@ -2483,7 +2498,7 @@ This version of the operator has been available since version 1 of the 'com.micr #### Type Constraints
-
T1 : tensor(int4), tensor(uint4), tensor(uint8)
+
T1 : tensor(int4), tensor(uint4), tensor(uint8), tensor(float8e4m3fn), tensor(float8e4m3fnuz), tensor(float8e5m2), tensor(float8e5m2fnuz), tensor(float4e2m1)
Constrain quantized types.
T2 : tensor(float), tensor(float16), tensor(bfloat16)
Constrain dequantized types.
@@ -2925,6 +2940,25 @@ This version of the operator has been available since version 1 of the 'com.micr **Cache Format:** The past and present KV cache tensors are expected in a BNSH format: `(batch_size, num_heads, cache_sequence_length, head_size)`, where `cache_sequence_length` is the length of the cached key/value sequences, or the maximum sequence length when past and present buffer sharing is used. + **Windowed KV Cache (`sliding_window_cache` attribute):** + When `sliding_window_cache` is 1, the past/present buffers are window-sized instead of full-length and the operator evicts internally. Let `C` be the cache capacity (dimension 2 of `past_key`, which is also the sequence dimension of `present_key`), `W` be `local_window_size`, and `T` be the absolute number of tokens processed so far by this batch entry, i.e. `seqlens_k[b] + 1`. The scalar `total_sequence_length` input is only the batch maximum of `T`; the layout below is per batch entry, so a ragged batch gets a different resident range per entry. `C` must be at least `W`. + + After a step, rows `[0, L)` of `present_key` and `present_value` hold the `L` most recent positions in increasing position order, so row `i` holds absolute position `T - L + i`. The retained positions are always physically contiguous and start at row 0; the layout never wraps around, so a ring-buffer layout cannot be exposed through these outputs. Rows `[L, C)` are unspecified. The resident count `L` is a function of `T` alone: + + ``` + G = C - W + 1 + L(T) = T if T <= C + L(T) = T - G * ceil((T - C) / G) otherwise + ``` + + Hence `min(T, W) <= L(T) <= min(T, C)`: the whole window stays resident, and eviction reclaims `G` positions at once rather than one position per step, so consumers must not assume that the cache is kept full at `min(T, C)`. + + Because `L` depends only on `T`, the resulting layout is independent of how the tokens were split into steps: a multi-token step of `S` tokens (speculative decoding, chunked prefill) leaves exactly the layout that the same tokens would produce one at a time. Any `S >= 1` is accepted, including `S > C`; a step that would evict positions it still has to read is staged internally, so the capacity does not have to cover the step. When past context is present, the existing operator restriction still applies: `sequence_length > 1` requires `batch_size == 1`. + + An execution provider may accept only part of the `C >= W` range. A configuration with `C < W` (equivalently, `W > C`) is invalid and is rejected with `INVALID_ARGUMENT`. The CUDA implementation requires `C == W`, so there `G` is 1 and `L(T)` is `min(T, C)`; a larger capacity is rejected. The CPU implementation accepts any `C >= W`, and slack above the window amortizes compaction over `G` steps. + + To drop the last `k` tokens, for example after rejecting speculative draft tokens, re-run with the smaller `total_sequence_length` and `seqlens_k` and leave the buffer untouched. That is exact when `L(T - k) == L(T) - k`, which callers can evaluate with the formula above. Otherwise the shorter layout needs positions that have already been evicted, and the window has to be re-materialized. + **Quantization:** When quantization is enabled, `past_key` and `past_value` inputs can be of type `float8e4m3fn`, `uint8` or `int8`. The corresponding `k_scale` and `v_scale` tensors must be provided. The operator will output `present_key` and `present_value` in same format as the `past_key` and `past_value`. @@ -2968,7 +3002,7 @@ This version of the operator has been available since version 1 of the 'com.micr
scale : float
Custom scale will be used if specified. Default value is 1/sqrt(head_size)
sliding_window_cache : int
-
Set to 1 when the past/present KV buffers are window-sized instead of holding the whole sequence. The op then keeps only the min(total_sequence_length, cache_capacity) most recent tokens, contiguously, using cache-relative indexing and evicting from the front as needed. Requires local_window_size > 0 and a cache capacity of at least local_window_size. Multi-token steps may use a temporary staging buffer, so the capacity need not cover the entire step. Default value is 0 (full-length cache).
+
Set to 1 when the past/present KV buffers are window-sized instead of holding the whole sequence. The op then evicts internally and indexes the buffers in cache-relative coordinates, keeping the most recent positions contiguously at rows [0, L) with min(T, local_window_size) <= L <= min(T, capacity), where T is seqlens_k[b] + 1 for that batch entry. Requires local_window_size > 0 and a cache capacity of at least local_window_size; a smaller capacity (W > C) is rejected with INVALID_ARGUMENT. The CUDA implementation additionally requires the capacity to equal local_window_size. Multi-token steps of any length are supported and produce the same layout as single-token steps, so the capacity need not cover the entire step. When past context is present, sequence_length > 1 requires batch_size == 1. See the Windowed KV Cache section of the operator description for the exact resident-range, eviction and rollback contract. Default value is 0 (full-length cache).
smooth_softmax : int
Use a smooth factor in softmax.
softcap : float
@@ -2987,9 +3021,9 @@ This version of the operator has been available since version 1 of the 'com.micr
value (optional) : T
Value with shape (batch_size, kv_sequence_length, kv_hidden_size)
past_key (optional) : T_CACHE
-
past state key with support for format BNSH. When past_key uses same tensor as present_key(k-v cache), it is of length max_sequence_length... otherwise of length past_sequence_length.
+
past state key with support for format BNSH. When past_key uses same tensor as present_key(k-v cache), it is of length max_sequence_length... otherwise of length past_sequence_length. When sliding_window_cache is 1 this length is the window cache capacity C, which is chosen by the caller independently of the sequence length and must be at least local_window_size.
past_value (optional) : T_CACHE
-
past state value with support for format BNSH. When past_value uses same tensor as present_value(k-v cache), it is of length max_sequence_length... otherwise of length past_sequence_length.
+
past state value with support for format BNSH. When past_value uses same tensor as present_value(k-v cache), it is of length max_sequence_length... otherwise of length past_sequence_length. When sliding_window_cache is 1 this length is the window cache capacity C, which is chosen by the caller independently of the sequence length and must be at least local_window_size.
seqlens_k : M
1D Tensor of shape (batch_size). Equivalent to (total_sequence_lengths - 1).
total_sequence_length : M
@@ -3001,7 +3035,7 @@ This version of the operator has been available since version 1 of the 'com.micr
position_ids (optional) : tensor(int64)
2D tensor with shape (batch_size, sequence_length). When processing the first prompt the kernel uses only the first element
attention_bias (optional) : T
-
additional add to QxK' with shape (batch_size or 1, num_heads or 1, sequence_length, total_sequence_length)
+
additional add to QxK' with shape (batch_size or 1, num_heads or 1, sequence_length, total_sequence_length). The last dimension is indexed by absolute key position and stays total_sequence_length when sliding_window_cache is 1: it is not reduced to the cache capacity or to local_window_size. The operator reads the columns of the positions that are resident in the cache and ignores the rest. CPU supports this windowed absolute-column indexing; CUDA rejects attention_bias when sliding_window_cache is 1.
head_sink (optional) : T
1D tensor with shape (num_heads). Each head has a smooth factor adding to the denominator of softmax.
k_scale (optional) : T_KV_SCALE
@@ -3441,18 +3475,18 @@ This version of the operator has been available since version 1 of the 'com.micr ### **com.microsoft.MatMulBlockQuantizedFp8Weight** - Weight-only block-scaled FP8 (E4M3) matrix multiplication. + Block-scaled FP8 (E4M3) matrix multiplication with optional FP8 activation quantization. - The weight tensor B is FP8 E4M3 of shape [N, K] with one FP32 scale per `block_size` consecutive - K values (`b_scale` of shape [N, ceil(K / block_size)]). The dequantized weight value is - `fp8_e4m3(B[n, k]) * b_scale[n, k / block_size]`. The weight is dequantized to the activation - type (FP16/BF16) and multiplied with the FP16/BF16 activation A. This path is architecture - independent and runs on any CUDA architecture (SM80+). + The weight tensor B has shape [N, K] with one FP32 scale per `block_size` consecutive K values + (`b_scale` of shape [N, ceil(K / block_size)]). The scaled weight value is + `B_scaled[n, k] = fp8_e4m3(B[n, k]) * b_scale[n, k / block_size]`. - When the optional `a_scale` (a single fp32 scalar) is provided, the activation A is statically - quantized to FP8 E4M3 and dequantized back (`a_deq = fp8_e4m3(A / a_scale) * a_scale`) before the - matmul, realizing W8A8 activation numerics. When `a_scale` is omitted the activation is kept at - full FP16/BF16 precision (weight-only W8A16). + When the optional scalar `a_scale` is provided, the activation values used in the multiplication + are `A_scaled = fp8_e4m3(A / a_scale) * a_scale` (W8A8). Otherwise, A retains its FP16/BF16 + precision (weight-only W8A16). + + The operator multiplies the activation by the transpose of B_scaled and adds the optional bias. + The output has shape [..., N] and the same element type as A. #### Version @@ -3475,7 +3509,7 @@ This version of the operator has been available since version 1 of the 'com.micr
b_scale : T2
Per-block FP32 weight scales of shape [N, ceil(K / block_size)].
a_scale (optional) : T2
-
Optional global fp32 activation scale (scalar). When present, A is statically quantized to FP8 E4M3 with this scale and dequantized back before the matmul (W8A8 numerics); when absent, A stays in full FP16/BF16 precision.
+
Optional global fp32 activation scale (scalar). When present, A is statically quantized to FP8 E4M3 with this scale (W8A8 numerics); when absent, A retains its FP16/BF16 precision.
bias (optional) : T
Optional bias of shape [N].
@@ -4282,12 +4316,26 @@ This version of the operator has been available since version 1 of the 'com.micr across invocations (chunked prefill or autoregressive decode), the optional past_ids input carries those preceding ids and present_ids returns the ids to pass to the next call. Both have shape (batch_size, max_ngram_size - 1) and are right-aligned, so the last slot is the most recent id. - Positions before the start of the whole sequence use pad_id. Running the op once over a full sequence - and running it over consecutive chunks while threading present_ids into past_ids produce identical - hash ids. When past_ids is omitted the missing history is pad_id, which matches a fresh sequence. + Positions before the start of the whole sequence use pad_id, or eos_token_id when it is provided. + Running the op once over a full sequence and running it over consecutive chunks while threading + present_ids into past_ids produce identical hash ids, including when reset_on_eos is enabled. When + segment_ids is used, segment boundaries are applied only within the current input_ids chunk and are + not inferred from past_ids. When past_ids is omitted the missing history is pad_id, or eos_token_id + when it is provided. past_ids and present_ids may use the same allocation. Such in-place execution is transaction-safe only when the whole operator call is unconditionally committed; a caller that may select a prefix or roll back must preserve past_ids. + + Optional inputs add packed-sequence and Qwen4-Exp-style n-gram embedding support: + + - eos_token_id, when provided together with reset_on_eos != 0, causes causal history to reset at EOS + boundaries: any shifted position at or before the most recent EOS strictly before the current + position is replaced with eos_token_id instead of the real token. + - segment_ids, when provided, additionally resets causal history at any position whose segment id + differs from the immediately preceding position's segment id within input_ids. Segment boundaries + are not checked against past_ids history. + - head_offsets, when provided, adds a fixed per-output-head offset after the modulo by the head's + vocabulary size, letting all heads across all n-gram orders share one flat embedding table. #### Version @@ -4302,19 +4350,27 @@ This version of the operator has been available since version 1 of the 'com.micr
Number of hash heads emitted for each n-gram order.
pad_id : int (required)
Compressed tokenizer id used to pad causal shifts before the beginning of a sequence.
+
reset_on_eos : int
+
When non-zero and the eos_token_id input is provided, reset causal n-gram history at EOS boundaries as described in the op doc. Default is 0 (disabled), which preserves the original pad_id-only behavior.
-#### Inputs (3 - 4) +#### Inputs (3 - 7)
input_ids : M
Compressed tokenizer ids with shape (batch_size, sequence_length).
multipliers : M
-
Per-shift hash multipliers with shape (max_ngram_size). Conventionally odd, but any value is accepted.
+
Per-shift hash multipliers with shape at least (max_ngram_size). Conventionally odd, but any value is accepted.
vocab_sizes : M
Per-output-head vocabulary sizes, conventionally prime, with shape ((max_ngram_size - 1) * n_head_per_ngram). Every entry must be strictly positive. The CPU implementation rejects a non-positive entry; GPU implementations guard the modulo to avoid a device-side division by zero and emit a hash id of 0 for that head.
past_ids (optional) : M
-
Optional compressed tokenizer ids for the max_ngram_size - 1 positions that precede this call, with shape (batch_size, max_ngram_size - 1). Right-aligned, so the last slot is the most recent id. If omitted the history is pad_id.
+
Optional compressed tokenizer ids for the max_ngram_size - 1 positions that precede this call, with shape (batch_size, max_ngram_size - 1). Right-aligned, so the last slot is the most recent id. If omitted the history is pad_id, or eos_token_id when provided.
+
head_offsets (optional) : M
+
Optional per-output-head additive offset with shape ((max_ngram_size - 1) * n_head_per_ngram), added after the modulo.
+
eos_token_id (optional) : M
+
Optional scalar end-of-sequence token id, same type as input_ids. Required for reset_on_eos to take effect and for EOS-based substitution of unavailable prior context; see the op doc.
+
segment_ids (optional) : tensor(int32)
+
Optional per-token segment id with shape (batch_size, sequence_length), used to reset causal history at packed-sequence boundaries within input_ids.
#### Outputs (1 - 2) @@ -4687,14 +4743,14 @@ This version of the operator has been available since version 1 of the 'com.micr * derives ordinary causal visibility purely from that packed metadata -- there is no mask input; * uses a single generic set of state slots (past_key_state / past_kv_buffer / past_gate_buffer / past_state_lengths) for both policy_mode values, each with a shape that is fixed across calls - (state never grows and is never concatenated); state overflow beyond the fixed capacity is + (state never grows and is never concatenated); a step that would overflow the fixed capacity is rejected as a deterministic no-op on state rather than truncated or allowed to corrupt memory; * additionally emits selected_counts, the exact number of active (non -1) entries per query, so that no downstream consumer needs to scan selected_indices for its query's true count. - + Both policy_mode values keep the semantics of SparseAttentionIndexer, applied independently to each request's own packed token range and fixed-capacity state slice: - + policy_mode = "qsa" ("query sparse attention" token indexer) Processes each request's new tokens sequentially: appends raw indexer keys to the generic pending buffer, and whenever it reaches compress_ratio tokens, mean-pools it, applies RMSNorm @@ -4705,7 +4761,7 @@ This version of the operator has been available since version 1 of the 'com.micr their token indices are emitted (request-local logical positions, i.e. the same numbering as past_sequence_lengths + local offset) followed by the causally visible tokens of the trailing incomplete block. - + policy_mode = "csa" ("compressed sparse attention" block indexer) Applies the same window-plan arithmetic as SparseAttentionIndexer (overlap/leftover/new window count) independently per request, using that request's own buffer_length and new token count; @@ -4713,7 +4769,7 @@ This version of the operator has been available since version 1 of the 'com.micr rotated and appended to key_state. Queries are scored against every causally visible compressed entry with sum_h w_h * ReLU(q_h . k) and the index_topk highest scoring entry indices are emitted. - + Common contract: * selected_indices is int32 with a fixed capacity that only depends on attributes: token_budget + compress_ratio - 1 for "qsa" (values are request-local token positions into the @@ -4733,8 +4789,8 @@ This version of the operator has been available since version 1 of the 'com.micr * cumulative_sequence_lengths, past_sequence_lengths and past_state_lengths are read directly by the device kernel; a zero-token request row (a repeated cumulative offset) is valid and simply contributes no query rows for that request. - - OgaEngine integration note: this operator only defines the ORT operator; wiring + + OgaEngine integration note: this operator only defines the ORT contrib op; wiring past_key_state / past_kv_buffer / past_gate_buffer / past_state_lengths as Engine-managed, per-request fixed-size state (analogous to a paged auxiliary cache) is expected to happen in the OgaEngine / Model Builder integration, which is out of scope for this operator definition. @@ -4764,7 +4820,7 @@ This version of the operator has been available since version 1 of the 'com.micr
Only for policy_mode 'qsa': maximum number of tokens selected from complete blocks. Must be > 0 and divisible by compress_ratio. Must be omitted when policy_mode is 'csa'.
-#### Inputs (10 - 15) +#### Inputs
query : T
@@ -4799,7 +4855,7 @@ This version of the operator has been available since version 1 of the 'com.micr
Generic per-request state length with shape (batch_size, 2). Column 0 is the key_state entry count (policy_mode 'qsa': complete-block count; 'csa': compressed-entry count); column 1 is the pending-buffer length (policy_mode 'qsa': incomplete-block length in [0, compress_ratio); 'csa': buffer length in [0, 2 * compress_ratio)).
-#### Outputs (6 - 6) +#### Outputs
selected_indices : M
@@ -7964,3 +8020,5 @@ No versioning maintained for experimental ops.
T : tensor(float)
Constrain input and output types to float32 tensors.
+ + diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md index 1efbe5c280c2b..dcd3a06cc3314 100644 --- a/docs/OperatorKernels.md +++ b/docs/OperatorKernels.md @@ -582,7 +582,7 @@ The **OpSet Version** column uses the following notation: |DynamicQuantizeMatMul|*in* A:**T1**
*in* B:**T2**
*in* b_scale:**T1**
*in* b_zero_point:**T2**
*in* bias:**T1**
*out* Y:**T1**|1+|**T1** = tensor(float)
**T2** = tensor(int8), tensor(uint8)| |DynamicTimeWarping|*in* input:**F**
*out* output:**I**|1+|**F** = tensor(float)
**I** = tensor(int32)| |EmbedLayerNormalization|*in* input_ids:**T1**
*in* segment_ids:**T1**
*in* word_embedding:**T**
*in* position_embedding:**T**
*in* segment_embedding:**T**
*in* gamma:**T**
*in* beta:**T**
*in* mask:**T1**
*in* position_ids:**T1**
*out* output:**T**
*out* mask_index:**T1**
*out* embedding_sum:**T**|1+|**T** = tensor(float)| -|EngramGate|*in* key:**T**
*in* query:**T**
*in* value:**T**
*in* key_norm_scale:**T**
*in* query_norm_scale:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| +|EngramGate|*in* key:**T**
*in* query:**T**
*in* value:**T**
*in* key_norm_scale:**T**
*in* query_norm_scale:**T**
*in* conv_norm_scale:**T**
*out* output:**T**
*out* gated_value_normed:**T**|1+|**T** = tensor(float), tensor(float16)| |ExpandDims|*in* X:**T**
*in* axis:**tensor(int32)**
*out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)
**axis** = tensor(int32)| |FastGelu|*in* X:**T**
*in* bias:**T**
*out* Y:**T**|1+|**T** = tensor(float)| |FusedConv|*in* X:**T**
*in* W:**T**
*in* B:**T**
*in* Z:**T**
*out* Y:**T**|1+|**T** = tensor(float)| @@ -590,7 +590,7 @@ The **OpSet Version** column uses the following notation: |FusedMatMul|*in* A:**T**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float)| |GatedAdd|*in* X:**T**
*in* Y:**T**
*in* gate:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| |GatedRMSNorm|*in* X:**T**
*in* scale:**T**
*in* gate:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)| -|GatherBlockQuantized|*in* data:**T1**
*in* indices:**Tind**
*in* scales:**T2**
*in* zero_points:**T1**
*out* output:**T2**|1+|**T1** = tensor(int4), tensor(uint4), tensor(uint8)
**T2** = tensor(float), tensor(float16)
**Tind** = tensor(int32), tensor(int64)| +|GatherBlockQuantized|*in* data:**T1**
*in* indices:**Tind**
*in* scales:**T2**
*in* zero_points:**T1**
*out* output:**T2**|1+|**T1** = tensor(float4e2m1), tensor(float8e4m3fn), tensor(float8e4m3fnuz), tensor(float8e5m2), tensor(float8e5m2fnuz), tensor(int4), tensor(uint4), tensor(uint8)
**T2** = tensor(float), tensor(float16)
**Tind** = tensor(int32), tensor(int64)| |GatherND|*in* data:**T**
*in* indices:**Tind**
*out* output:**T**|1+|**T** = tensor(bfloat16), tensor(bool), tensor(double), tensor(float), tensor(float16), tensor(int16), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint16), tensor(uint32), tensor(uint64), tensor(uint8)
**Tind** = tensor(int32), tensor(int64)| |Gelu|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(float)| |GreedySearch|*in* input_ids:**I**
*in* max_length:**I**
*in* min_length:**I**
*in* repetition_penalty:**T**
*in* vocab_mask:**I**
*in* prefix_vocab_mask:**I**
*in* attention_mask:**I**
*out* sequences:**I**|1+|**T** = tensor(float)| @@ -609,7 +609,7 @@ The **OpSet Version** column uses the following notation: |MoE|*in* input:**T**
*in* router_probs:**T**
*in* fc1_experts_weights:**T**
*in* fc1_experts_bias:**T**
*in* fc2_experts_weights:**T**
*in* fc2_experts_bias:**T**
*in* fc3_experts_weights:**T**
*in* fc3_experts_bias:**T**
*out* output:**T**|1+|**T** = tensor(float)| |MultiHeadAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* bias:**T**
*in* key_padding_mask:**M**
*in* attention_bias:**T**
*in* past_key:**T**
*in* past_value:**T**
*in* past_sequence_length:**M**
*in* cache_indirection:**M**
*out* output:**T**
*out* present_key:**T**
*out* present_value:**T**
*out* qk:**QK**|1+|**T** = tensor(float)| |MurmurHash3|*in* X:**T1**
*out* Y:**T2**|1+|**T1** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint32), tensor(uint64)
**T2** = tensor(int32), tensor(uint32)| -|NGramHashMapping|*in* input_ids:**M**
*in* multipliers:**M**
*in* vocab_sizes:**M**
*in* past_ids:**M**
*out* hash_ids:**M**
*out* present_ids:**M**|1+|**M** = tensor(int32), tensor(int64)| +|NGramHashMapping|*in* input_ids:**M**
*in* multipliers:**M**
*in* vocab_sizes:**M**
*in* past_ids:**M**
*in* head_offsets:**M**
*in* eos_token_id:**M**
*in* segment_ids:**tensor(int32)**
*out* hash_ids:**M**
*out* present_ids:**M**|1+|**M** = tensor(int32), tensor(int64)| |NGramRepeatBlock|*in* input_ids:**Tid**
*in* scores:**T**
*out* scores_out:**T**|1+|**T** = tensor(float)
**Tid** = tensor(int64)| |NhwcMaxPool|*in* x:**T**
*out* y:**T**|1+|**T** = tensor(int8), tensor(uint8)| |Pad|*in* data:**T**
*in* pads:**tensor(int64)**
*in* value:**T**
*out* output:**T**|1+|**T** = tensor(float)| @@ -1089,7 +1089,7 @@ The **OpSet Version** column uses the following notation: |DequantizeWithOrder|*in* input:**Q**
*in* scale_input:**S**
*out* output:**F**|1+|**F** = tensor(float), tensor(float16)
**Q** = tensor(int8)
**S** = tensor(float)| |DynamicTimeWarping|*in* input:**F**
*out* output:**I**|1+|**F** = tensor(float)
**I** = tensor(int32)| |EmbedLayerNormalization|*in* input_ids:**T1**
*in* segment_ids:**T1**
*in* word_embedding:**T**
*in* position_embedding:**T**
*in* segment_embedding:**T**
*in* gamma:**T**
*in* beta:**T**
*in* mask:**T1**
*in* position_ids:**T1**
*out* output:**T**
*out* mask_index:**T1**
*out* embedding_sum:**T**|1+|**T** = tensor(float), tensor(float16)| -|EngramGate|*in* key:**T**
*in* query:**T**
*in* value:**T**
*in* key_norm_scale:**T**
*in* query_norm_scale:**T**
*out* output:**T**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| +|EngramGate|*in* key:**T**
*in* query:**T**
*in* value:**T**
*in* key_norm_scale:**T**
*in* query_norm_scale:**T**
*in* conv_norm_scale:**T**
*out* output:**T**
*out* gated_value_normed:**T**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| |FastGelu|*in* X:**T**
*in* bias:**T**
*out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)| |FusedConv|*in* X:**T**
*in* W:**T**
*in* B:**T**
*in* Z:**T**
*out* Y:**T**|1+|**T** = tensor(float)| |FusedMatMul|*in* A:**T**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)| @@ -1097,7 +1097,7 @@ The **OpSet Version** column uses the following notation: |GatedDeltaNet|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* cu_seqlens:**TI**
*in* decay:**TS**
*in* beta:**TS**
*in* initial_state:**TS**
*in* a_log:**TS**
*in* dt_bias:**TS**
*in* capture_count:**TI**
*in* state_update_active:**TI**
*out* output:**T**
*out* final_state:**TS**
*out* state_update:**TS**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)
**TI** = tensor(int32)
**TS** = tensor(float)| |GatedRMSNorm|*in* X:**T**
*in* scale:**T**
*in* gate:**T**
*out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| |GatedRelativePositionBias|*in* query_layer:**T**
*in* query_bias:**T**
*in* rel_pos:**T**
*in* weight:**T**
*in* bias:**T**
*in* eco_a:**T**
*in* token_offset:**M**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| -|GatherBlockQuantized|*in* data:**T1**
*in* indices:**Tind**
*in* scales:**T2**
*in* zero_points:**T1**
*out* output:**T2**|1+|**T1** = tensor(int4), tensor(uint4), tensor(uint8)
**T2** = tensor(bfloat16), tensor(float), tensor(float16)
**Tind** = tensor(int32), tensor(int64)| +|GatherBlockQuantized|*in* data:**T1**
*in* indices:**Tind**
*in* scales:**T2**
*in* zero_points:**T1**
*out* output:**T2**|1+|**T1** = tensor(float4e2m1), tensor(float8e4m3fn), tensor(float8e4m3fnuz), tensor(float8e5m2), tensor(float8e5m2fnuz), tensor(int4), tensor(uint4), tensor(uint8)
**T2** = tensor(bfloat16), tensor(float), tensor(float16)
**Tind** = tensor(int32), tensor(int64)| |Gelu|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)| |GemmFloat8|*in* A:**TA**
*in* B:**TB**
*in* C:**TC**
*in* scaleA:**TS**
*in* scaleB:**TS**
*in* scaleY:**TS**
*out* Y:**TR**|1+|**TA** = tensor(bfloat16), tensor(float), tensor(float16), tensor(float8e4m3fn), tensor(float8e5m2)
**TB** = tensor(bfloat16), tensor(float), tensor(float16), tensor(float8e4m3fn), tensor(float8e5m2)
**TR** = tensor(bfloat16), tensor(float), tensor(float16), tensor(float8e4m3fn), tensor(float8e5m2)
**TS** = tensor(float)| |GemmaRotaryEmbedding|*in* emb:**U**
*in* q:**T**
*in* q_rot:**T**
*in* k:**T**
*in* k_rot:**T**
*out* output1:**T**
*out* output2:**T**|1+|**T** = tensor(float16)
**U** = tensor(float)| @@ -1117,7 +1117,7 @@ The **OpSet Version** column uses the following notation: |MatMulNBits|*in* A:**T1**
*in* B:**T2**
*in* scales:**T1**
*in* zero_points:**T3**
*in* g_idx:**T4**
*in* bias:**T1**
*out* Y:**T1**|1+|**T1** = tensor(bfloat16), tensor(float), tensor(float16)
**T2** = tensor(uint8)
**T3** = tensor(bfloat16), tensor(float), tensor(float16), tensor(uint8)| |MoE|*in* input:**T**
*in* router_probs:**T**
*in* fc1_experts_weights:**T**
*in* fc1_experts_bias:**T**
*in* fc2_experts_weights:**T**
*in* fc2_experts_bias:**T**
*in* fc3_experts_weights:**T**
*in* fc3_experts_bias:**T**
*out* output:**T**|1+|**T** = tensor(bfloat16), tensor(float), tensor(float16)| |MultiHeadAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* bias:**T**
*in* key_padding_mask:**M**
*in* attention_bias:**T**
*in* past_key:**T**
*in* past_value:**T**
*in* past_sequence_length:**M**
*in* cache_indirection:**M**
*out* output:**T**
*out* present_key:**T**
*out* present_value:**T**
*out* qk:**QK**|1+|**QK** = tensor(bfloat16), tensor(float), tensor(float16)
**T** = tensor(bfloat16), tensor(float), tensor(float16)| -|NGramHashMapping|*in* input_ids:**M**
*in* multipliers:**M**
*in* vocab_sizes:**M**
*in* past_ids:**M**
*out* hash_ids:**M**
*out* present_ids:**M**|1+|**M** = tensor(int32), tensor(int64)| +|NGramHashMapping|*in* input_ids:**M**
*in* multipliers:**M**
*in* vocab_sizes:**M**
*in* past_ids:**M**
*in* head_offsets:**M**
*in* eos_token_id:**M**
*in* segment_ids:**tensor(int32)**
*out* hash_ids:**M**
*out* present_ids:**M**|1+|**M** = tensor(int32), tensor(int64)| |NGramRepeatBlock|*in* input_ids:**Tid**
*in* scores:**T**
*out* scores_out:**T**|1+|**T** = tensor(float)
**Tid** = tensor(int64)| |NhwcConv|*in* X:**T**
*in* W:**T**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)| |PackedAttention|*in* input:**T**
*in* weights:**T**
*in* bias:**T**
*in* token_offset:**M**
*in* cumulative_sequence_length:**M**
*in* attention_bias:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)| diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc index 7c4d32887f15c..19d69eb1cfc67 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc @@ -237,6 +237,10 @@ Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) con const int64_t capacity = psai::SelectedCapacity(psai::Policy::kQsa, token_budget_, index_topk_, compress_ratio_); + Tensor* present_gate_buffer = context->Output(psai::kPresentGateBuffer, TensorShape({0})); + ORT_RETURN_IF(present_gate_buffer != nullptr, + "PackedSparseAttentionIndexer: output ", psai::kPresentGateBuffer, + " must be omitted for policy_mode 'qsa'"); Tensor* selected_indices = context->Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity})); Tensor* selected_counts = context->Output(psai::kSelectedCounts, TensorShape({total_tokens})); Tensor* present_key_state = context->Output(psai::kPresentKeyState, key_state_shape); diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc index 725e8a1c09d0f..3cdf07349d452 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -788,6 +788,10 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont {batch_size, psai::kStateLengthColumns})); const int64_t capacity = psai::SelectedCapacity(psai::Policy::kQsa, token_budget_, index_topk_, compress_ratio_); + Tensor* present_gate_buffer = context.Output(psai::kPresentGateBuffer, TensorShape({0})); + ORT_RETURN_IF(present_gate_buffer != nullptr, + "PackedSparseAttentionIndexer: output ", psai::kPresentGateBuffer, + " must be omitted for policy_mode 'qsa'"); Tensor* selected_indices = context.Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity})); Tensor* selected_counts = context.Output(psai::kSelectedCounts, TensorShape({total_tokens})); Tensor* present_key_state = context.Output(psai::kPresentKeyState, key_state_shape); diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index 37ce43b089b24..e50fbf90bbbd5 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -2396,13 +2396,6 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference fail_shape_inference("PackedSparseAttentionIndexer: exactly ", psai::kFixedOutputCount, " declared outputs are required, got ", ctx.getNumOutputs()); } - const bool has_gate_output = ctx.getOutputType(psai::kPresentGateBuffer) != nullptr; - if (has_gate_output == is_qsa) { - fail_shape_inference("PackedSparseAttentionIndexer: output ", psai::kPresentGateBuffer, - is_qsa ? " must be omitted when policy_mode is 'qsa'" - : " is required when policy_mode is 'csa'"); - } - updateOutputElemType(ctx, psai::kSelectedIndices, ONNX_NAMESPACE::TensorProto_DataType_INT32); updateOutputElemType(ctx, psai::kSelectedCounts, ONNX_NAMESPACE::TensorProto_DataType_INT32); updateOutputElemType(ctx, psai::kPresentStateLengths, ONNX_NAMESPACE::TensorProto_DataType_INT32); diff --git a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc index ba7c3baf2fca3..f97f2bfaeb513 100644 --- a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -88,7 +88,7 @@ struct GraphOptions { int64_t total_tokens = 5; int64_t num_heads = 2; int64_t head_size = 8; - int64_t rotary_width = 8; + int64_t rotary_width = 4; int64_t compress_ratio = 2; int64_t state_capacity = 6; int64_t input_state_capacity = -1; @@ -99,7 +99,6 @@ struct GraphOptions { bool add_index_topk = false; bool add_csa_inputs = false; bool add_position_ids = false; - bool invert_gate_output_presence = false; int output_count = psai::kFixedOutputCount; std::string policy_mode = psai::kPolicyModeQsa; }; @@ -157,8 +156,7 @@ void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { std::vector outputs; for (int i = 0; i < options.output_count; ++i) { - const bool gate_output_present = is_csa != options.invert_gate_output_presence; - outputs.push_back(i == psai::kPresentGateBuffer && !gate_output_present ? &empty : builder.MakeOutput()); + outputs.push_back(i == psai::kPresentGateBuffer && !is_csa ? &empty : builder.MakeOutput()); } Node& node = builder.AddNode("PackedSparseAttentionIndexer", inputs, outputs, kMSDomain); node.AddAttribute("policy_mode", options.policy_mode); @@ -237,21 +235,6 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsStateCapacityShapeMi "past_key_state dimension 1 must equal state_capacity"); } -TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsQsaGateOutput) { - GraphOptions options; - options.invert_gate_output_presence = true; - ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, - "must be omitted when policy_mode is 'qsa'"); -} - -TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsCsaMissingGateOutput) { - GraphOptions options; - options.policy_mode = psai::kPolicyModeCsa; - options.invert_gate_output_presence = true; - ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, - "is required when policy_mode is 'csa'"); -} - TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsMismatchedKeyTokenCount) { GraphOptions options; options.key_total_tokens = options.total_tokens + 1; From e81906b1db7df7d9bdd7d387ab26a78fabf19466 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 17 Sep 2026 03:59:00 +0000 Subject: [PATCH 09/10] Reduce packed indexer WebGPU bindings Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com> --- .../bert/packed_sparse_attention_indexer.cc | 189 +++++++++--------- .../bert/packed_sparse_attention_indexer.h | 23 +-- 2 files changed, 96 insertions(+), 116 deletions(-) diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc index 3cdf07349d452..b4cc3c0dd7a8a 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -91,35 +91,48 @@ Status PackedSparseAttentionIndexerCopyProgram::GenerateShaderCode(ShaderHelper& const auto& dst = shader.AddOutput("dst", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); shader.MainFunctionBody() << shader.GuardAgainstOutOfBoundsWorkgroupSizes("uniforms.total") - << " " << dst.SetByOffset("global_idx", "dst_element_t(" + src.GetByOffset("global_idx") + ")") << "\n"; + << " " + << dst.SetByOffset("uniforms.dst_offset + global_idx", + "dst_element_t(" + src.GetByOffset("global_idx") + ")") + << "\n"; + return Status::OK(); +} + +static Status PackTwoTensors(onnxruntime::webgpu::ComputeContext& context, + const Tensor& first, + const Tensor& second, + Tensor& packed) { + uint32_t dst_offset = 0; + for (const Tensor* source : {&first, &second}) { + const uint32_t total = ToUint32(source->Shape().Size()); + if (total > 0) { + PackedSparseAttentionIndexerCopyProgram copy; + copy.SetWorkgroupSize(kWorkgroupSize) + .AddInput({source, ProgramTensorMetadataDependency::Type}) + .AddOutput({&packed, ProgramTensorMetadataDependency::Type}) + .SetDispatchGroupSize((total + kWorkgroupSize - 1) / kWorkgroupSize) + .AddUniformVariables({{total}, {dst_offset}}); + ORT_RETURN_IF_ERROR(context.RunProgram(copy)); + } + dst_offset += total; + } return Status::OK(); } Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHelper& shader) const { const auto& key = shader.AddInput("key", ShaderUsage::UseUniform); const auto& norm = shader.AddInput("key_norm_weight", ShaderUsage::UseUniform); - const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); - const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& rotary_cache = shader.AddInput("rotary_cache", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); - const ShaderVariableHelper* past_kv_buffer = nullptr; - if (!kv_buffer_aliases_) { - past_kv_buffer = &shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); - } - const ShaderVariableHelper* past_state_lengths = nullptr; - if (!state_lengths_aliases_) { - past_state_lengths = &shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); - } const auto& present_key_state = shader.AddOutput("present_key_state", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_kv_buffer = shader.AddOutput("present_kv_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_state_lengths = shader.AddOutput("present_state_lengths", ShaderUsage::UseUniform); const auto& overflow_flags = shader.AddOutput("overflow_flags", ShaderUsage::UseUniform); - const ShaderVariableHelper& kv_buffer_history = - kv_buffer_aliases_ ? present_kv_buffer : *past_kv_buffer; - const ShaderVariableHelper& state_lengths_history = - state_lengths_aliases_ ? present_state_lengths : *past_state_lengths; + const ShaderVariableHelper& kv_buffer_history = present_kv_buffer; + const ShaderVariableHelper& state_lengths_history = present_state_lengths; shader.AdditionalImplementation() << "fn extended_key(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, d: u32) -> f32 {\n" @@ -167,8 +180,10 @@ Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHe shader.AdditionalImplementation() << " let cache = position * uniforms.rotary_width + d;\n"; } shader.AdditionalImplementation() - << " return value * f32(" << cos_cache.GetByOffset("cache") << ") + paired * f32(" - << sin_cache.GetByOffset("cache") << ");\n" + << " let cache_size = select(1u, uniforms.batch_size, " << (cos_cache_batched_ ? "true" : "false") + << ") * uniforms.max_rotary_length * uniforms.rotary_width;\n" + << " return value * f32(" << rotary_cache.GetByOffset("cache") << ") + paired * f32(" + << rotary_cache.GetByOffset("cache_size + cache") << ");\n" << "}\n"; shader.MainFunctionBody() @@ -227,8 +242,7 @@ Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHe Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHelper& shader) const { const auto& query = shader.AddInput("query", ShaderUsage::UseUniform); const auto& present_key_state = shader.AddInput("present_key_state", ShaderUsage::UseUniform); - const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); - const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& rotary_cache = shader.AddInput("rotary_cache", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); const ShaderVariableHelper* position_ids = nullptr; @@ -290,8 +304,10 @@ Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHe shader.AdditionalImplementation() << " let cache = position_clamped * uniforms.rotary_width + d;\n"; } shader.AdditionalImplementation() - << " value = value * f32(" << cos_cache.GetByOffset("cache") << ") + paired * f32(" - << sin_cache.GetByOffset("cache") << ");\n" + << " let cache_size = select(1u, uniforms.batch_size, " << (cos_cache_batched_ ? "true" : "false") + << ") * uniforms.max_rotary_length * uniforms.rotary_width;\n" + << " value = value * f32(" << rotary_cache.GetByOffset("cache") << ") + paired * f32(" + << rotary_cache.GetByOffset("cache_size + cache") << ");\n" << " return value;\n" << "}\n" << "fn block_score(token: u32, b: u32, block_index: u32, position: i32) -> f32 {\n" @@ -370,26 +386,11 @@ Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHe } Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHelper& shader) const { - const auto& key = shader.AddInput("key", ShaderUsage::UseUniform); - const auto& gate = shader.AddInput("gate", ShaderUsage::UseUniform); + const auto& key_gate = shader.AddInput("key_gate", ShaderUsage::UseUniform); const auto& norm = shader.AddInput("key_norm_weight", ShaderUsage::UseUniform); - const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); - const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& rotary_cache = shader.AddInput("rotary_cache", ShaderUsage::UseUniform); const auto& position_bias = shader.AddInput("position_bias", ShaderUsage::UseUniform); - const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); - const auto& past_seqlens = shader.AddInput("past_sequence_lengths", ShaderUsage::UseUniform); - const ShaderVariableHelper* past_kv_buffer = nullptr; - if (!kv_buffer_aliases_) { - past_kv_buffer = &shader.AddInput("past_kv_buffer", ShaderUsage::UseUniform); - } - const ShaderVariableHelper* past_gate_buffer = nullptr; - if (!gate_buffer_aliases_) { - past_gate_buffer = &shader.AddInput("past_gate_buffer", ShaderUsage::UseUniform); - } - const ShaderVariableHelper* past_state_lengths = nullptr; - if (!state_lengths_aliases_) { - past_state_lengths = &shader.AddInput("past_state_lengths", ShaderUsage::UseUniform); - } + const auto& sequence_metadata = shader.AddInput("sequence_metadata", ShaderUsage::UseUniform); const auto& present_key_state = shader.AddOutput("present_key_state", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_kv_buffer = @@ -398,12 +399,9 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe shader.AddOutput("present_gate_buffer", ShaderUsage::UseUniform | ShaderUsage::UseElementTypeAlias); const auto& present_state_lengths = shader.AddOutput("present_state_lengths", ShaderUsage::UseUniform); const auto& overflow_flags = shader.AddOutput("overflow_flags", ShaderUsage::UseUniform); - const ShaderVariableHelper& kv_buffer_history = - kv_buffer_aliases_ ? present_kv_buffer : *past_kv_buffer; - const ShaderVariableHelper& gate_buffer_history = - gate_buffer_aliases_ ? present_gate_buffer : *past_gate_buffer; - const ShaderVariableHelper& state_lengths_history = - state_lengths_aliases_ ? present_state_lengths : *past_state_lengths; + const ShaderVariableHelper& kv_buffer_history = present_kv_buffer; + const ShaderVariableHelper& gate_buffer_history = present_gate_buffer; + const ShaderVariableHelper& state_lengths_history = present_state_lengths; shader.AdditionalImplementation() << "fn extended_key(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, channel: u32) -> f32 {\n" @@ -413,7 +411,7 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe << " return f32(" << kv_buffer_history.GetByOffset("idx") << ");\n" << " }\n" << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\n" - << " return f32(" << key.GetByOffset("idx2") << ");\n" + << " return f32(" << key_gate.GetByOffset("idx2") << ");\n" << "}\n" << "fn extended_gate(b: u32, virtual_pos: i32, old_buf_len: i32, req_start: i32, channel: u32) -> f32 {\n" << " let width = 2u * uniforms.head_size;\n" @@ -422,7 +420,8 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe << " return f32(" << gate_buffer_history.GetByOffset("idx") << ");\n" << " }\n" << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\n" - << " return f32(" << gate.GetByOffset("idx2") << ");\n" + << " let gate_base = uniforms.total_tokens * 2u * uniforms.head_size;\n" + << " return f32(" << key_gate.GetByOffset("gate_base + idx2") << ");\n" << "}\n" << "fn pooled(b: u32, k: i32, old_buf_len: i32, req_start: i32, overlap_length: i32, d: u32) -> f32 {\n" << " let width = 2u * uniforms.head_size;\n" @@ -473,14 +472,14 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe shader.MainFunctionBody() << " let b = workgroup_idx;\n" << " if (b >= uniforms.batch_size || local_idx != 0u) { return; }\n" - << " let req_start = " << cu_seqlens.GetByOffset("b") << ";\n" - << " let req_end = " << cu_seqlens.GetByOffset("b + 1u") << ";\n" + << " let req_start = " << sequence_metadata.GetByOffset("b") << ";\n" + << " let req_end = " << sequence_metadata.GetByOffset("b + 1u") << ";\n" << " let old_key_len = " << state_lengths_history.GetByOffset("b * 2u") << ";\n" << " let old_buf_len = " << state_lengths_history.GetByOffset("b * 2u + 1u") << ";\n" - << " let invalid = " << cu_seqlens.GetByOffset("0") << " != 0 || " - << cu_seqlens.GetByOffset("uniforms.batch_size") << " != i32(uniforms.total_tokens) || " + << " let invalid = " << sequence_metadata.GetByOffset("0") << " != 0 || " + << sequence_metadata.GetByOffset("uniforms.batch_size") << " != i32(uniforms.total_tokens) || " << "req_start < 0 || req_end < req_start || req_end > i32(uniforms.total_tokens) || " - << past_seqlens.GetByOffset("b") << " < 0 || old_key_len < 0 || " + << sequence_metadata.GetByOffset("uniforms.batch_size + 1u + b") << " < 0 || old_key_len < 0 || " << "old_key_len > i32(uniforms.state_capacity) || old_buf_len < 0 || " << "old_buf_len > i32(uniforms.buffer_capacity);\n" << " if (invalid) {\n" @@ -537,8 +536,12 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe << " let paired = sign * pooled(b, k, old_buf_len, req_start, overlap_length, pair_d) * " "inverse_rms * f32(" << norm.GetByOffset("pair_d") << ");\n" - << " value = value * f32(" << cos_cache.GetByOffset("cache_base + offset / 2u") << ") + paired * f32(" - << sin_cache.GetByOffset("cache_base + offset / 2u") << ");\n" + << " let cache_size = select(1u, uniforms.batch_size, " + << (cos_cache_batched_ ? "true" : "false") + << ") * uniforms.max_rotary_length * uniforms.rotary_width;\n" + << " value = value * f32(" << rotary_cache.GetByOffset("cache_base + offset / 2u") + << ") + paired * f32(" + << rotary_cache.GetByOffset("cache_size + cache_base + offset / 2u") << ");\n" << " }\n" << " " << present_key_state.SetByOffset("(b * uniforms.state_capacity + entry) * uniforms.head_size + d", @@ -570,8 +573,7 @@ Status PackedSparseAttentionIndexerCsaSelectProgram::GenerateShaderCode(ShaderHe const auto& present_key_state = shader.AddInput("present_key_state", ShaderUsage::UseUniform); const auto& head_weights = shader.AddInput("head_weights", ShaderUsage::UseUniform); const auto& position_ids = shader.AddInput("position_ids", ShaderUsage::UseUniform); - const auto& cos_cache = shader.AddInput("cos_cache", ShaderUsage::UseUniform); - const auto& sin_cache = shader.AddInput("sin_cache", ShaderUsage::UseUniform); + const auto& rotary_cache = shader.AddInput("rotary_cache", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); const auto& present_state_lengths = shader.AddInput("present_state_lengths", ShaderUsage::UseUniform); const auto& overflow_flags = shader.AddInput("overflow_flags", ShaderUsage::UseUniform); @@ -616,8 +618,10 @@ Status PackedSparseAttentionIndexerCsaSelectProgram::GenerateShaderCode(ShaderHe << " let cache = position * uniforms.rotary_width + offset / 2u;\n"; } shader.AdditionalImplementation() - << " value = value * f32(" << cos_cache.GetByOffset("cache") << ") + paired * f32(" - << sin_cache.GetByOffset("cache") << ");\n" + << " let cache_size = select(1u, uniforms.batch_size, " << (cos_cache_batched_ ? "true" : "false") + << ") * uniforms.max_rotary_length * uniforms.rotary_width;\n" + << " value = value * f32(" << rotary_cache.GetByOffset("cache") << ") + paired * f32(" + << rotary_cache.GetByOffset("cache_size + cache") << ");\n" << " return value;\n" << "}\n" << "fn entry_score(token: u32, b: u32, entry: u32, raw: vec2) -> f32 {\n" @@ -787,6 +791,10 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths", {batch_size, psai::kStateLengthColumns})); + Tensor rotary_cache = + context.CreateGPUTensor(cos_cache->DataType(), TensorShape({2 * cos_cache->Shape().Size()})); + ORT_RETURN_IF_ERROR(PackTwoTensors(context, *cos_cache, *sin_cache, rotary_cache)); + const int64_t capacity = psai::SelectedCapacity(psai::Policy::kQsa, token_budget_, index_topk_, compress_ratio_); Tensor* present_gate_buffer = context.Output(psai::kPresentGateBuffer, TensorShape({0})); ORT_RETURN_IF(present_gate_buffer != nullptr, @@ -809,7 +817,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont .AddInput({past_key_state, ProgramTensorMetadataDependency::Type}) .AddOutput({present_key_state, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) - .AddUniformVariables({{ToUint32(total)}}); + .AddUniformVariables({{ToUint32(total)}, {0u}}); ORT_RETURN_IF_ERROR(context.RunProgram(copy)); } } @@ -821,7 +829,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont .AddInput({past_kv_buffer, ProgramTensorMetadataDependency::Type}) .AddOutput({present_kv_buffer, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) - .AddUniformVariables({{ToUint32(total)}}); + .AddUniformVariables({{ToUint32(total)}, {0u}}); ORT_RETURN_IF_ERROR(context.RunProgram(copy)); } } @@ -833,29 +841,20 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont .AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}) .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) - .AddUniformVariables({{ToUint32(total)}}); + .AddUniformVariables({{ToUint32(total)}, {0u}}); ORT_RETURN_IF_ERROR(context.RunProgram(copy)); } } if (batch_size > 0) { - const bool kv_buffer_aliases = present_kv_buffer->DataRaw() == past_kv_buffer->DataRaw(); - const bool state_lengths_aliases = present_state_lengths->DataRaw() == past_state_lengths->DataRaw(); - PackedSparseAttentionIndexerQsaUpdateProgram update{rotary.batched, kv_buffer_aliases, state_lengths_aliases}; - update.CacheHint(rotary.batched, kv_buffer_aliases, state_lengths_aliases) + PackedSparseAttentionIndexerQsaUpdateProgram update{rotary.batched}; + update.CacheHint(rotary.batched) .SetWorkgroupSize(kWorkgroupSize) .AddInputs({{key, ProgramTensorMetadataDependency::Type}, {norm, ProgramTensorMetadataDependency::Type}, - {cos_cache, ProgramTensorMetadataDependency::Type}, - {sin_cache, ProgramTensorMetadataDependency::Type}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}, {past_seqlens, ProgramTensorMetadataDependency::Type}}); - if (!kv_buffer_aliases) { - update.AddInput({past_kv_buffer, ProgramTensorMetadataDependency::Type}); - } - if (!state_lengths_aliases) { - update.AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}); - } update.AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, {present_kv_buffer, ProgramTensorMetadataDependency::Type}}) .AddOutput({present_state_lengths, ProgramTensorMetadataDependency::Type}) @@ -882,8 +881,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont .SetWorkgroupSize(kWorkgroupSize) .AddInputs({{query, ProgramTensorMetadataDependency::Type}, {present_key_state, ProgramTensorMetadataDependency::Type}, - {cos_cache, ProgramTensorMetadataDependency::Type}, - {sin_cache, ProgramTensorMetadataDependency::Type}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}, {past_seqlens, ProgramTensorMetadataDependency::Type}}); if (position_ids != nullptr) { @@ -971,6 +969,16 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths", {batch_size, psai::kStateLengthColumns})); + Tensor key_gate = + context.CreateGPUTensor(key->DataType(), TensorShape({key->Shape().Size() + gate->Shape().Size()})); + ORT_RETURN_IF_ERROR(PackTwoTensors(context, *key, *gate, key_gate)); + Tensor rotary_cache = + context.CreateGPUTensor(cos_cache->DataType(), TensorShape({2 * cos_cache->Shape().Size()})); + ORT_RETURN_IF_ERROR(PackTwoTensors(context, *cos_cache, *sin_cache, rotary_cache)); + Tensor sequence_metadata = context.CreateGPUTensor( + cu_seqlens->DataType(), TensorShape({cu_seqlens->Shape().Size() + past_seqlens->Shape().Size()})); + ORT_RETURN_IF_ERROR(PackTwoTensors(context, *cu_seqlens, *past_seqlens, sequence_metadata)); + const int64_t capacity = psai::SelectedCapacity(psai::Policy::kCsa, token_budget_, index_topk_, compress_ratio_); Tensor* selected_indices = context.Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity})); Tensor* selected_counts = context.Output(psai::kSelectedCounts, TensorShape({total_tokens})); @@ -996,7 +1004,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont .AddInput({src, ProgramTensorMetadataDependency::Type}) .AddOutput({dst, ProgramTensorMetadataDependency::Type}) .SetDispatchGroupSize(ToUint32((total + kWorkgroupSize - 1) / kWorkgroupSize)) - .AddUniformVariables({{ToUint32(total)}}); + .AddUniformVariables({{ToUint32(total)}, {0u}}); return context.RunProgram(copy); }; ORT_RETURN_IF_ERROR(copy_if_needed(past_key_state, present_key_state)); @@ -1005,30 +1013,14 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_ERROR(copy_if_needed(past_state_lengths, present_state_lengths)); if (batch_size > 0) { - const bool kv_buffer_aliases = present_kv_buffer->DataRaw() == past_kv_buffer->DataRaw(); - const bool gate_buffer_aliases = present_gate_buffer->DataRaw() == past_gate_buffer->DataRaw(); - const bool state_lengths_aliases = present_state_lengths->DataRaw() == past_state_lengths->DataRaw(); - PackedSparseAttentionIndexerCsaUpdateProgram update{ - rotary.batched, kv_buffer_aliases, gate_buffer_aliases, state_lengths_aliases}; - update.CacheHint(rotary.batched, kv_buffer_aliases, gate_buffer_aliases, state_lengths_aliases) + PackedSparseAttentionIndexerCsaUpdateProgram update{rotary.batched}; + update.CacheHint(rotary.batched) .SetWorkgroupSize(kWorkgroupSize) - .AddInputs({{key, ProgramTensorMetadataDependency::Type}, - {gate, ProgramTensorMetadataDependency::Type}, + .AddInputs({{&key_gate, ProgramTensorMetadataDependency::Type}, {norm, ProgramTensorMetadataDependency::Type}, - {cos_cache, ProgramTensorMetadataDependency::Type}, - {sin_cache, ProgramTensorMetadataDependency::Type}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, {position_bias, ProgramTensorMetadataDependency::Type}, - {cu_seqlens, ProgramTensorMetadataDependency::Type}, - {past_seqlens, ProgramTensorMetadataDependency::Type}}); - if (!kv_buffer_aliases) { - update.AddInput({past_kv_buffer, ProgramTensorMetadataDependency::Type}); - } - if (!gate_buffer_aliases) { - update.AddInput({past_gate_buffer, ProgramTensorMetadataDependency::Type}); - } - if (!state_lengths_aliases) { - update.AddInput({past_state_lengths, ProgramTensorMetadataDependency::Type}); - } + {&sequence_metadata, ProgramTensorMetadataDependency::Type}}); update.AddOutputs({{present_key_state, ProgramTensorMetadataDependency::Type}, {present_kv_buffer, ProgramTensorMetadataDependency::Type}, {present_gate_buffer, ProgramTensorMetadataDependency::Type}}) @@ -1058,8 +1050,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont {present_key_state, ProgramTensorMetadataDependency::Type}, {head_weights, ProgramTensorMetadataDependency::Type}, {position_ids, ProgramTensorMetadataDependency::Type}, - {cos_cache, ProgramTensorMetadataDependency::Type}, - {sin_cache, ProgramTensorMetadataDependency::Type}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}}) .AddInputs({{present_state_lengths, ProgramTensorMetadataDependency::Type}, {&overflow_flags, ProgramTensorMetadataDependency::Type}}) diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h index 1fafc10110861..2607a1492e48f 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h @@ -23,7 +23,8 @@ class PackedSparseAttentionIndexerCopyProgram final PackedSparseAttentionIndexerCopyProgram() : Program{"PackedSparseAttentionIndexerCopy"} {} Status GenerateShaderCode(ShaderHelper& shader) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( - {"total", ProgramUniformVariableDataType::Uint32}); + {"total", ProgramUniformVariableDataType::Uint32}, + {"dst_offset", ProgramUniformVariableDataType::Uint32}); }; // One invocation per request: forms every newly-closed compress_ratio block (mean-pool -> @@ -31,12 +32,9 @@ class PackedSparseAttentionIndexerCopyProgram final class PackedSparseAttentionIndexerQsaUpdateProgram final : public Program { public: - PackedSparseAttentionIndexerQsaUpdateProgram(bool cos_cache_batched, bool kv_buffer_aliases, - bool state_lengths_aliases) + explicit PackedSparseAttentionIndexerQsaUpdateProgram(bool cos_cache_batched) : Program{"PackedSparseAttentionIndexerQsaUpdate"}, - cos_cache_batched_{cos_cache_batched}, - kv_buffer_aliases_{kv_buffer_aliases}, - state_lengths_aliases_{state_lengths_aliases} {} + cos_cache_batched_{cos_cache_batched} {} Status GenerateShaderCode(ShaderHelper& shader) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( {"batch_size", ProgramUniformVariableDataType::Uint32}, @@ -51,8 +49,6 @@ class PackedSparseAttentionIndexerQsaUpdateProgram final private: bool cos_cache_batched_; - bool kv_buffer_aliases_; - bool state_lengths_aliases_; }; // One invocation per query token: rotates the query, scores it against every causally visible @@ -90,13 +86,9 @@ class PackedSparseAttentionIndexerQsaSelectProgram final class PackedSparseAttentionIndexerCsaUpdateProgram final : public Program { public: - PackedSparseAttentionIndexerCsaUpdateProgram(bool cos_cache_batched, bool kv_buffer_aliases, - bool gate_buffer_aliases, bool state_lengths_aliases) + explicit PackedSparseAttentionIndexerCsaUpdateProgram(bool cos_cache_batched) : Program{"PackedSparseAttentionIndexerCsaUpdate"}, - cos_cache_batched_{cos_cache_batched}, - kv_buffer_aliases_{kv_buffer_aliases}, - gate_buffer_aliases_{gate_buffer_aliases}, - state_lengths_aliases_{state_lengths_aliases} {} + cos_cache_batched_{cos_cache_batched} {} Status GenerateShaderCode(ShaderHelper& shader) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES( {"batch_size", ProgramUniformVariableDataType::Uint32}, @@ -111,9 +103,6 @@ class PackedSparseAttentionIndexerCsaUpdateProgram final private: bool cos_cache_batched_; - bool kv_buffer_aliases_; - bool gate_buffer_aliases_; - bool state_lengths_aliases_; }; // One invocation per query token: rotates the query, scores it against every causally visible From f32b9bff11fb914ff1a651e34033c1a8f69aa742 Mon Sep 17 00:00:00 2001 From: Kunal Vaishnavi Date: Fri, 18 Sep 2026 21:35:11 +0000 Subject: [PATCH 10/10] Fuse query normalization into PackedSparseAttentionIndexer --- docs/ContribOperators.md | 6 +- .../packed_sparse_attention_indexer.md | 33 ++++---- .../webgpu/packed_sparse_attention_indexer.md | 2 + .../packed_sparse_attention_indexer_common.h | 33 ++++---- .../sparse/packed_sparse_attention_indexer.cc | 36 +++++++-- .../packed_sparse_attention_indexer_impl.cu | 50 ++++++++---- .../packed_sparse_attention_indexer_impl.h | 2 + .../bert/packed_sparse_attention_indexer.cc | 61 ++++++++++++--- .../core/graph/contrib_ops/bert_defs.cc | 77 ++++++++++++------- ...packed_sparse_attention_indexer_op_test.cc | 26 +++++-- 10 files changed, 221 insertions(+), 105 deletions(-) diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index f3d97c1cc4599..2ac4253ac4cfe 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -4805,7 +4805,7 @@ This version of the operator has been available since version 1 of the 'com.micr
compress_ratio : int (required)
Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0 and 2 * compress_ratio - 1 must not exceed INT_MAX.
epsilon : float
-
Epsilon of the RMS normalization applied to the compressed keys. Default is 1e-6.
+
Epsilon of the RMS normalization applied to queries and compressed keys. Default is 1e-6.
head_weight_scale : float
Only for policy_mode 'csa': scale applied to head_weights. Default is 1/sqrt(num_heads). Must be omitted when policy_mode is 'qsa'.
index_topk : int
@@ -4824,9 +4824,11 @@ This version of the operator has been available since version 1 of the 'com.micr
query : T
-
Packed indexer queries with shape (total_tokens, num_heads, head_size), already normalized but not yet rotated.
+
Packed indexer queries with shape (total_tokens, num_heads * head_size), before normalization, logical reshape, and rotary embedding.
key : T
Packed indexer key projection of the new tokens. Shape is (total_tokens, head_size) for policy_mode 'qsa' and (total_tokens, 2 * head_size) for policy_mode 'csa', where the first head_size channels are the Ca series and the last head_size channels the Cb series.
+
query_norm_weight : T
+
Effective RMSNorm multiplier of the queries, with shape (head_size).
key_norm_weight : T
Effective RMSNorm multiplier of the compressed keys, with shape (head_size).
cos_cache : T
diff --git a/docs/contrib_ops/packed_sparse_attention_indexer.md b/docs/contrib_ops/packed_sparse_attention_indexer.md index 883a074a657aa..6f7f7e8ba0bf5 100644 --- a/docs/contrib_ops/packed_sparse_attention_indexer.md +++ b/docs/contrib_ops/packed_sparse_attention_indexer.md @@ -10,7 +10,7 @@ Source: (policy enum, selected-capacity formula and CSA window-plan arithmetic, shared unmodified with `SparseAttentionIndexer`), [packed_sparse_attention_indexer_common.h](../../onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h) -(fixed 15-input / 6-output slot map), +(fixed 16-input / 6-output slot map), [packed_sparse_attention_indexer.cc](../../onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc) / [packed_sparse_attention_indexer_impl.cu](../../onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu) (CUDA), @@ -52,28 +52,29 @@ Attributes: | `scale` | default `1/sqrt(head_size)` | per-head score scale | | `head_weight_scale` | `csa` only, default `1/sqrt(num_heads)` | head-weight score scale | -Inputs are **fixed at 15 indices** for both policies (unlike the dense op, which uses a different +Inputs are **fixed at 16 indices** for both policies (unlike the dense op, which uses a different input/output count per policy). A slot not owned by the active policy is a *positional* optional: its `NodeProto` input name is empty rather than the slot being removed from the list, so every later slot keeps its fixed index. | # | Name | Shape | Type | Policy | |---|---|---|---|---| -| 0 | `query` | `(total_tokens, num_heads, head_size)` | T | both | +| 0 | `query` | `(total_tokens, num_heads*head_size)` | T | both | | 1 | `key` | `(total_tokens, head_size)` qsa / `(total_tokens, 2*head_size)` csa | T | both | -| 2 | `key_norm_weight` | `(head_size)` | T | both | -| 3 | `cos_cache` | `(max_position, rotary_width)` or `(batch_size, max_position, rotary_width)` | T | both | -| 4 | `sin_cache` | same shape as `cos_cache` | T | both | -| 5 | `cumulative_sequence_lengths` | `(batch_size + 1)` | int32, device-resident | both | -| 6 | `past_sequence_lengths` | `(batch_size)` | int32, device-resident | both | -| 7 | `gate` | `(total_tokens, 2*head_size)` | T | csa only | -| 8 | `position_bias` | `(compress_ratio, 2*head_size)` | T | csa only | -| 9 | `head_weights` | `(total_tokens, num_heads)` | T | csa only | -| 10 | `position_ids` | `(total_tokens)` | int64 | optional qsa / required csa | -| 11 | `past_key_state` | `(batch_size, state_capacity, head_size)` | T | both (generic) | -| 12 | `past_kv_buffer` | `(batch_size, 2*compress_ratio-1, width)` | T | both (generic) | -| 13 | `past_gate_buffer` | same shape as `past_kv_buffer` | T | csa only | -| 14 | `past_state_lengths` | `(batch_size, 2)` | int32, device-resident | both (generic) | +| 2 | `query_norm_weight` | `(head_size)` | T | both | +| 3 | `key_norm_weight` | `(head_size)` | T | both | +| 4 | `cos_cache` | `(max_position, rotary_width)` or `(batch_size, max_position, rotary_width)` | T | both | +| 5 | `sin_cache` | same shape as `cos_cache` | T | both | +| 6 | `cumulative_sequence_lengths` | `(batch_size + 1)` | int32, device-resident | both | +| 7 | `past_sequence_lengths` | `(batch_size)` | int32, device-resident | both | +| 8 | `gate` | `(total_tokens, 2*head_size)` | T | csa only | +| 9 | `position_bias` | `(compress_ratio, 2*head_size)` | T | csa only | +| 10 | `head_weights` | `(total_tokens, num_heads)` | T | csa only | +| 11 | `position_ids` | `(total_tokens)` | int64 | optional qsa / required csa | +| 12 | `past_key_state` | `(batch_size, state_capacity, head_size)` | T | both (generic) | +| 13 | `past_kv_buffer` | `(batch_size, 2*compress_ratio-1, width)` | T | both (generic) | +| 14 | `past_gate_buffer` | same shape as `past_kv_buffer` | T | csa only | +| 15 | `past_state_lengths` | `(batch_size, 2)` | int32, device-resident | both (generic) | Outputs are **fixed at 6 indices** for both policies (`present_gate_buffer` is declared with an empty output name for `qsa`, the same positional-optional convention as above): diff --git a/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md b/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md index c1ce8a1549358..33ade715ef995 100644 --- a/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md +++ b/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md @@ -32,6 +32,8 @@ into workgroup-shared arrays, both to keep every kernel correct without relying sized by a runtime (uniform) `head_size`, and to keep the packed kernel's structure directly comparable to the CUDA implementation's per-request/per-token update and score/select stages. All reductions and softmax calculations accumulate in FP32, including for FP16 inputs. +Raw flattened queries are RMS-normalized per logical head with `query_norm_weight` before the +policy-specific rotary embedding, matching the CUDA implementation. Per-request quantities (`cumulative_sequence_lengths`, `past_sequence_lengths`, `past_state_lengths`) are read directly from device buffers inside the shaders — never on the diff --git a/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h b/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h index 24b7ce1584071..0d55ae4c5e5e5 100644 --- a/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h +++ b/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h @@ -30,25 +30,26 @@ using sparse_attention_indexer::SelectedCapacity; using sparse_attention_indexer::TryComputeCsaWindowPlan; using sparse_attention_indexer::TryParsePolicy; -// Fixed input slots. Slots 7-9 belong to policy_mode="csa" only; slot 10 (position_ids) is +// Fixed input slots. Slots 8-10 belong to policy_mode="csa" only; slot 11 (position_ids) is // optional for "qsa" and required for "csa". Every other slot is required for both policies. enum InputIndex : int { - kQuery = 0, // [total_tokens, num_heads, head_size] + kQuery = 0, // [total_tokens, num_heads * head_size] kKey = 1, // qsa: [total_tokens, head_size]; csa: [total_tokens, 2 * head_size] - kKeyNormWeight = 2, // [head_size] - kCosCache = 3, // [max_position, rotary_width] or [batch_size, max_position, rotary_width] - kSinCache = 4, // same shape as cos_cache - kCumulativeSequenceLengths = 5, // [batch_size + 1], int32 - kPastSequenceLengths = 6, // [batch_size], int32 - kGate = 7, // csa only: [total_tokens, 2 * head_size] - kPositionBias = 8, // csa only: [compress_ratio, 2 * head_size] - kHeadWeights = 9, // csa only: [total_tokens, num_heads] - kPositionIds = 10, // optional (qsa) / required (csa): [total_tokens], int64 - kPastKeyState = 11, // generic: [batch_size, state_capacity, head_size] - kPastKvBuffer = 12, // generic: [batch_size, 2 * compress_ratio - 1, width] - kPastGateBuffer = 13, // csa only: same shape as past_kv_buffer - kPastStateLengths = 14, // generic: [batch_size, 2], int32 - kInputCount = 15, + kQueryNormWeight = 2, // [head_size] + kKeyNormWeight = 3, // [head_size] + kCosCache = 4, // [max_position, rotary_width] or [batch_size, max_position, rotary_width] + kSinCache = 5, // same shape as cos_cache + kCumulativeSequenceLengths = 6, // [batch_size + 1], int32 + kPastSequenceLengths = 7, // [batch_size], int32 + kGate = 8, // csa only: [total_tokens, 2 * head_size] + kPositionBias = 9, // csa only: [compress_ratio, 2 * head_size] + kHeadWeights = 10, // csa only: [total_tokens, num_heads] + kPositionIds = 11, // optional (qsa) / required (csa): [total_tokens], int64 + kPastKeyState = 12, // generic: [batch_size, state_capacity, head_size] + kPastKvBuffer = 13, // generic: [batch_size, 2 * compress_ratio - 1, width] + kPastGateBuffer = 14, // csa only: same shape as past_kv_buffer + kPastStateLengths = 15, // generic: [batch_size, 2], int32 + kInputCount = 16, }; // Fixed output slots. present_gate_buffer is declared (with an empty name) but not produced for diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc index 19d69eb1cfc67..6e01ecf2c6d6a 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc @@ -169,6 +169,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) con const Tensor* query = context->Input(psai::kQuery); const Tensor* key = context->Input(psai::kKey); + const Tensor* query_norm_weight = context->Input(psai::kQueryNormWeight); const Tensor* key_norm_weight = context->Input(psai::kKeyNormWeight); const Tensor* cos_cache = context->Input(psai::kCosCache); const Tensor* sin_cache = context->Input(psai::kSinCache); @@ -181,13 +182,20 @@ Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) con ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); const auto& query_shape = query->Shape(); - ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, - "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads, head_size), " + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 2, + "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads * head_size), " "got ", query_shape.ToString()); const int64_t total_tokens = query_shape[0]; - const int64_t num_heads = query_shape[1]; - const int64_t head_size = query_shape[2]; + ORT_RETURN_IF(query_norm_weight == nullptr, "PackedSparseAttentionIndexer: query_norm_weight is required"); + const auto& query_norm_shape = query_norm_weight->Shape(); + ORT_RETURN_IF_NOT(query_norm_shape.NumDimensions() == 1, + "PackedSparseAttentionIndexer: query_norm_weight must have shape (head_size), got ", + query_norm_shape.ToString()); + const int64_t head_size = query_norm_shape[0]; + ORT_RETURN_IF(head_size <= 0 || query_shape[1] % head_size != 0, + "PackedSparseAttentionIndexer: query width must be divisible by head_size"); + const int64_t num_heads = query_shape[1] / head_size; ORT_RETURN_IF_ERROR(CheckIntDimension("total_tokens", total_tokens)); ORT_RETURN_IF_ERROR(CheckIntDimension("num_heads", num_heads, false)); ORT_RETURN_IF_ERROR(CheckIntDimension("head_size", head_size, false)); @@ -206,6 +214,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) con ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, head_size})); + ORT_RETURN_IF_ERROR(CheckShape(query_norm_weight, "query_norm_weight", {head_size})); ORT_RETURN_IF_ERROR(CheckShape(key_norm_weight, "key_norm_weight", {head_size})); if (position_ids != nullptr) { ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens})); @@ -278,6 +287,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) con Stream(context), params, reinterpret_cast(query->Data()), reinterpret_cast(key->Data()), + reinterpret_cast(query_norm_weight->Data()), reinterpret_cast(key_norm_weight->Data()), reinterpret_cast(cos_cache->Data()), reinterpret_cast(sin_cache->Data()), @@ -302,6 +312,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) con const Tensor* query = context->Input(psai::kQuery); const Tensor* key = context->Input(psai::kKey); + const Tensor* query_norm_weight = context->Input(psai::kQueryNormWeight); const Tensor* key_norm_weight = context->Input(psai::kKeyNormWeight); const Tensor* cos_cache = context->Input(psai::kCosCache); const Tensor* sin_cache = context->Input(psai::kSinCache); @@ -318,13 +329,20 @@ Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) con ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); const auto& query_shape = query->Shape(); - ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, - "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads, head_size), " + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 2, + "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads * head_size), " "got ", query_shape.ToString()); const int64_t total_tokens = query_shape[0]; - const int64_t num_heads = query_shape[1]; - const int64_t head_size = query_shape[2]; + ORT_RETURN_IF(query_norm_weight == nullptr, "PackedSparseAttentionIndexer: query_norm_weight is required"); + const auto& query_norm_shape = query_norm_weight->Shape(); + ORT_RETURN_IF_NOT(query_norm_shape.NumDimensions() == 1, + "PackedSparseAttentionIndexer: query_norm_weight must have shape (head_size), got ", + query_norm_shape.ToString()); + const int64_t head_size = query_norm_shape[0]; + ORT_RETURN_IF(head_size <= 0 || query_shape[1] % head_size != 0, + "PackedSparseAttentionIndexer: query width must be divisible by head_size"); + const int64_t num_heads = query_shape[1] / head_size; ORT_RETURN_IF_ERROR(CheckIntDimension("total_tokens", total_tokens)); ORT_RETURN_IF_ERROR(CheckIntDimension("num_heads", num_heads, false)); ORT_RETURN_IF_ERROR(CheckIntDimension("head_size", head_size, false)); @@ -346,6 +364,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) con ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, width})); + ORT_RETURN_IF_ERROR(CheckShape(query_norm_weight, "query_norm_weight", {head_size})); ORT_RETURN_IF_ERROR(CheckShape(key_norm_weight, "key_norm_weight", {head_size})); ORT_RETURN_IF_ERROR(CheckShape(gate, "gate", {total_tokens, width})); ORT_RETURN_IF_ERROR(CheckShape(position_bias, "position_bias", {compress_ratio_, width})); @@ -421,6 +440,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) con Stream(context), params, reinterpret_cast(query->Data()), reinterpret_cast(key->Data()), + reinterpret_cast(query_norm_weight->Data()), reinterpret_cast(key_norm_weight->Data()), reinterpret_cast(cos_cache->Data()), reinterpret_cast(sin_cache->Data()), diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu index f2bcc3f29c2bf..620247fc41d48 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu @@ -209,11 +209,13 @@ __global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, con // defaults to past_sequence_lengths[batch] + request-local offset when position_ids is absent); // otherwise the csa trailing convention with positions always taken from position_ids. template -__global__ void PackedRotateQueryKernel(const T* query, const T* cos_cache, const T* sin_cache, +__global__ void PackedRotateQueryKernel(const T* query, const T* query_norm_weight, + const T* cos_cache, const T* sin_cache, const int32_t* cumulative_sequence_lengths, const int32_t* past_sequence_lengths, const int64_t* position_ids, float* query_rotated, PackedSparseAttentionIndexerParams params) { extern __shared__ float shared[]; + float* reduction = shared + params.head_size; const int64_t rows = static_cast(params.total_tokens) * params.num_heads; for (int64_t row = blockIdx.x; row < rows; row += gridDim.x) { const int token = static_cast(row / params.num_heads); @@ -225,6 +227,17 @@ __global__ void PackedRotateQueryKernel(const T* query, const T* cos_cache, cons } __syncthreads(); + float sum_squares = 0.0f; + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + sum_squares += shared[d] * shared[d]; + } + sum_squares = SaiBlockSum(sum_squares, reduction); + const float inverse_rms = rsqrtf(sum_squares / static_cast(params.head_size) + params.epsilon); + for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) { + shared[d] = shared[d] * inverse_rms * to_float(query_norm_weight[d]); + } + __syncthreads(); + const int64_t abs_position = params.has_position_ids ? position_ids[token] @@ -725,7 +738,8 @@ size_t GetCsaPackedWorkspaceFloatCount(const PackedSparseAttentionIndexerParams& template Status LaunchQsaPackedSparseAttentionIndexer( cudaStream_t stream, const PackedSparseAttentionIndexerParams& params, const T* query, const T* key, - const T* key_norm_weight, const T* cos_cache, const T* sin_cache, const int32_t* cumulative_sequence_lengths, + const T* query_norm_weight, const T* key_norm_weight, const T* cos_cache, const T* sin_cache, + const int32_t* cumulative_sequence_lengths, const int32_t* past_sequence_lengths, const int64_t* position_ids, const T* past_key_state, const T* past_kv_buffer, const int32_t* past_state_lengths, int32_t* selected_indices, int32_t* selected_counts, T* present_key_state, T* present_kv_buffer, int32_t* present_state_lengths, @@ -769,9 +783,9 @@ Status LaunchQsaPackedSparseAttentionIndexer( const int64_t rotate_rows = static_cast(params.total_tokens) * params.num_heads; const int rotate_blocks = static_cast(std::min(rotate_rows, kSaiMaxGridDimX)); - PackedRotateQueryKernel<<>>( - query, cos_cache, sin_cache, cumulative_sequence_lengths, past_sequence_lengths, position_ids, query_rotated, - params); + PackedRotateQueryKernel<<>>( + query, query_norm_weight, cos_cache, sin_cache, cumulative_sequence_lengths, past_sequence_lengths, + position_ids, query_rotated, params); if (params.state_capacity > 0) { const int64_t score_work = static_cast(params.total_tokens) * params.state_capacity; @@ -792,7 +806,8 @@ Status LaunchQsaPackedSparseAttentionIndexer( template Status LaunchCsaPackedSparseAttentionIndexer( cudaStream_t stream, const PackedSparseAttentionIndexerParams& params, const T* query, const T* key, - const T* key_norm_weight, const T* cos_cache, const T* sin_cache, const T* gate, const T* position_bias, + const T* query_norm_weight, const T* key_norm_weight, const T* cos_cache, const T* sin_cache, const T* gate, + const T* position_bias, const T* head_weights, const int32_t* cumulative_sequence_lengths, const int32_t* past_sequence_lengths, const int64_t* position_ids, const T* past_key_state, const T* past_kv_buffer, const T* past_gate_buffer, const int32_t* past_state_lengths, int32_t* selected_indices, int32_t* selected_counts, T* present_key_state, @@ -844,9 +859,9 @@ Status LaunchCsaPackedSparseAttentionIndexer( const int64_t rotate_rows = static_cast(params.total_tokens) * params.num_heads; const int rotate_blocks = static_cast(std::min(rotate_rows, kSaiMaxGridDimX)); - PackedRotateQueryKernel<<>>( - query, cos_cache, sin_cache, cumulative_sequence_lengths, past_sequence_lengths, position_ids, query_rotated, - params); + PackedRotateQueryKernel<<>>( + query, query_norm_weight, cos_cache, sin_cache, cumulative_sequence_lengths, past_sequence_lengths, + position_ids, query_rotated, params); if (params.state_capacity > 0) { const int64_t score_work = static_cast(params.total_tokens) * params.state_capacity; @@ -864,14 +879,15 @@ Status LaunchCsaPackedSparseAttentionIndexer( return CUDA_CALL(cudaGetLastError()); } -#define INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(T) \ - template Status LaunchQsaPackedSparseAttentionIndexer( \ - cudaStream_t, const PackedSparseAttentionIndexerParams&, const T*, const T*, const T*, const T*, \ - const T*, const int32_t*, const int32_t*, const int64_t*, const T*, const T*, const int32_t*, int32_t*, \ - int32_t*, T*, T*, int32_t*, float*, int32_t*); \ - template Status LaunchCsaPackedSparseAttentionIndexer( \ - cudaStream_t, const PackedSparseAttentionIndexerParams&, const T*, const T*, const T*, const T*, \ - const T*, const T*, const T*, const T*, const int32_t*, const int32_t*, const int64_t*, const T*, \ +#define INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(T) \ + template Status LaunchQsaPackedSparseAttentionIndexer( \ + cudaStream_t, const PackedSparseAttentionIndexerParams&, const T*, const T*, const T*, const T*, \ + const T*, const T*, const int32_t*, const int32_t*, const int64_t*, const T*, const T*, const int32_t*, \ + int32_t*, \ + int32_t*, T*, T*, int32_t*, float*, int32_t*); \ + template Status LaunchCsaPackedSparseAttentionIndexer( \ + cudaStream_t, const PackedSparseAttentionIndexerParams&, const T*, const T*, const T*, const T*, \ + const T*, const T*, const T*, const T*, const T*, const int32_t*, const int32_t*, const int64_t*, const T*, \ const T*, const T*, const int32_t*, int32_t*, int32_t*, T*, T*, T*, int32_t*, float*, int32_t*); INSTANTIATE_PACKED_SPARSE_ATTENTION_INDEXER(float) diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h index 2eae2bdc4465a..9d92076c2ff51 100644 --- a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h @@ -56,6 +56,7 @@ Status LaunchQsaPackedSparseAttentionIndexer( const PackedSparseAttentionIndexerParams& params, const T* query, const T* key, + const T* query_norm_weight, const T* key_norm_weight, const T* cos_cache, const T* sin_cache, @@ -79,6 +80,7 @@ Status LaunchCsaPackedSparseAttentionIndexer( const PackedSparseAttentionIndexerParams& params, const T* query, const T* key, + const T* query_norm_weight, const T* key_norm_weight, const T* cos_cache, const T* sin_cache, diff --git a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc index b4cc3c0dd7a8a..461cc03666525 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -241,6 +241,7 @@ Status PackedSparseAttentionIndexerQsaUpdateProgram::GenerateShaderCode(ShaderHe Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHelper& shader) const { const auto& query = shader.AddInput("query", ShaderUsage::UseUniform); + const auto& query_norm = shader.AddInput("query_norm_weight", ShaderUsage::UseUniform); const auto& present_key_state = shader.AddInput("present_key_state", ShaderUsage::UseUniform); const auto& rotary_cache = shader.AddInput("rotary_cache", ShaderUsage::UseUniform); const auto& cu_seqlens = shader.AddInput("cumulative_sequence_lengths", ShaderUsage::UseUniform); @@ -288,14 +289,25 @@ Status PackedSparseAttentionIndexerQsaSelectProgram::GenerateShaderCode(ShaderHe << "}\n"; } shader.AdditionalImplementation() + << "fn normalized_query_value(token: u32, head: u32, d: u32) -> f32 {\n" + << " let base = (token * uniforms.num_heads + head) * uniforms.head_size;\n" + << " var square_sum = 0.0;\n" + << " for (var k = 0u; k < uniforms.head_size; k++) {\n" + << " let value = f32(" << query.GetByOffset("base + k") << ");\n" + << " square_sum += value * value;\n" + << " }\n" + << " return f32(" << query.GetByOffset("base + d") + << ") * inverseSqrt(square_sum / f32(uniforms.head_size) + uniforms.epsilon) * f32(" + << query_norm.GetByOffset("d") << ");\n" + << "}\n" << "fn query_value(token: u32, head: u32, d: u32, position: i32, b: u32) -> f32 {\n" << " let base = (token * uniforms.num_heads + head) * uniforms.head_size;\n" - << " var value = f32(" << query.GetByOffset("base + d") << ");\n" + << " var value = normalized_query_value(token, head, d);\n" << " if (d >= uniforms.rotary_width) { return value; }\n" << " let half = uniforms.rotary_width / 2u;\n" << " let pair_d = select(d - half, d + half, d < half);\n" << " let sign = select(1.0, -1.0, d < half);\n" - << " let paired = sign * f32(" << query.GetByOffset("base + pair_d") << ");\n" + << " let paired = sign * normalized_query_value(token, head, pair_d);\n" << " let position_clamped = clamp_position(position);\n"; if (cos_cache_batched_) { shader.AdditionalImplementation() @@ -570,6 +582,7 @@ Status PackedSparseAttentionIndexerCsaUpdateProgram::GenerateShaderCode(ShaderHe Status PackedSparseAttentionIndexerCsaSelectProgram::GenerateShaderCode(ShaderHelper& shader) const { const auto& query = shader.AddInput("query", ShaderUsage::UseUniform); + const auto& query_norm = shader.AddInput("query_norm_weight", ShaderUsage::UseUniform); const auto& present_key_state = shader.AddInput("present_key_state", ShaderUsage::UseUniform); const auto& head_weights = shader.AddInput("head_weights", ShaderUsage::UseUniform); const auto& position_ids = shader.AddInput("position_ids", ShaderUsage::UseUniform); @@ -600,15 +613,26 @@ Status PackedSparseAttentionIndexerCsaSelectProgram::GenerateShaderCode(ShaderHe << " let cr = uniforms.compress_ratio;\n" << " return p / cr + select(0u, 1u, p % cr == cr - 1u);\n" << "}\n" + << "fn normalized_query_value(token: u32, head: u32, d: u32) -> f32 {\n" + << " let base = (token * uniforms.num_heads + head) * uniforms.head_size;\n" + << " var square_sum = 0.0;\n" + << " for (var k = 0u; k < uniforms.head_size; k++) {\n" + << " let value = f32(" << query.GetByOffset("base + k") << ");\n" + << " square_sum += value * value;\n" + << " }\n" + << " return f32(" << query.GetByOffset("base + d") + << ") * inverseSqrt(square_sum / f32(uniforms.head_size) + uniforms.epsilon) * f32(" + << query_norm.GetByOffset("d") << ");\n" + << "}\n" << "fn query_value(token: u32, head: u32, d: u32, raw: vec2, b: u32) -> f32 {\n" << " let base = (token * uniforms.num_heads + head) * uniforms.head_size;\n" - << " var value = f32(" << query.GetByOffset("base + d") << ");\n" + << " var value = normalized_query_value(token, head, d);\n" << " let rotary_base = uniforms.head_size - 2u * uniforms.rotary_width;\n" << " if (d < rotary_base) { return value; }\n" << " let offset = d - rotary_base;\n" << " let pair_d = select(d - 1u, d + 1u, (offset & 1u) == 0u);\n" << " let sign = select(1.0, -1.0, (offset & 1u) == 0u);\n" - << " let paired = sign * f32(" << query.GetByOffset("base + pair_d") << ");\n" + << " let paired = sign * normalized_query_value(token, head, pair_d);\n" << " let position = min(clamped_position(raw), uniforms.max_rotary_length - 1u);\n"; if (cos_cache_batched_) { shader.AdditionalImplementation() @@ -739,6 +763,7 @@ Status PackedSparseAttentionIndexer::ComputeInternal(onnxruntime::webgpu::Comput Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeContext& context) const { const Tensor* query = context.Input(psai::kQuery); const Tensor* key = context.Input(psai::kKey); + const Tensor* query_norm = context.Input(psai::kQueryNormWeight); const Tensor* norm = context.Input(psai::kKeyNormWeight); const Tensor* cos_cache = context.Input(psai::kCosCache); const Tensor* sin_cache = context.Input(psai::kSinCache); @@ -751,10 +776,15 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); const auto& query_shape = query->Shape(); - ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, "PackedSparseAttentionIndexer: query must have rank 3"); + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 2, "PackedSparseAttentionIndexer: query must have rank 2"); const int64_t total_tokens = query_shape[0]; - const int64_t num_heads = query_shape[1]; - const int64_t head_size = query_shape[2]; + ORT_RETURN_IF(query_norm == nullptr, "PackedSparseAttentionIndexer: query_norm_weight is required"); + const auto& query_norm_shape = query_norm->Shape(); + ORT_RETURN_IF_NOT(query_norm_shape.NumDimensions() == 1 && query_norm_shape[0] > 0 && + query_shape[1] % query_norm_shape[0] == 0, + "PackedSparseAttentionIndexer: invalid flattened query dimensions"); + const int64_t head_size = query_norm_shape[0]; + const int64_t num_heads = query_shape[1] / head_size; ORT_RETURN_IF_NOT(num_heads > 0 && head_size > 0, "PackedSparseAttentionIndexer: invalid query dimensions"); ORT_RETURN_IF(cu_seqlens == nullptr, "PackedSparseAttentionIndexer: cumulative_sequence_lengths is required"); @@ -767,6 +797,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_ERROR(CheckShape(past_seqlens, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, head_size})); + ORT_RETURN_IF_ERROR(CheckShape(query_norm, "query_norm_weight", {head_size})); ORT_RETURN_IF_ERROR(CheckShape(norm, "key_norm_weight", {head_size})); if (position_ids != nullptr) { ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens})); @@ -880,6 +911,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont select.CacheHint(rotary.batched, position_ids != nullptr) .SetWorkgroupSize(kWorkgroupSize) .AddInputs({{query, ProgramTensorMetadataDependency::Type}, + {query_norm, ProgramTensorMetadataDependency::Type}, {present_key_state, ProgramTensorMetadataDependency::Type}, {&rotary_cache, ProgramTensorMetadataDependency::Type}, {cu_seqlens, ProgramTensorMetadataDependency::Type}, @@ -910,6 +942,7 @@ Status PackedSparseAttentionIndexer::ComputeQsa(onnxruntime::webgpu::ComputeCont Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeContext& context) const { const Tensor* query = context.Input(psai::kQuery); const Tensor* key = context.Input(psai::kKey); + const Tensor* query_norm = context.Input(psai::kQueryNormWeight); const Tensor* norm = context.Input(psai::kKeyNormWeight); const Tensor* cos_cache = context.Input(psai::kCosCache); const Tensor* sin_cache = context.Input(psai::kSinCache); @@ -926,10 +959,15 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required"); const auto& query_shape = query->Shape(); - ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3, "PackedSparseAttentionIndexer: query must have rank 3"); + ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 2, "PackedSparseAttentionIndexer: query must have rank 2"); const int64_t total_tokens = query_shape[0]; - const int64_t num_heads = query_shape[1]; - const int64_t head_size = query_shape[2]; + ORT_RETURN_IF(query_norm == nullptr, "PackedSparseAttentionIndexer: query_norm_weight is required"); + const auto& query_norm_shape = query_norm->Shape(); + ORT_RETURN_IF_NOT(query_norm_shape.NumDimensions() == 1 && query_norm_shape[0] > 0 && + query_shape[1] % query_norm_shape[0] == 0, + "PackedSparseAttentionIndexer: invalid flattened query dimensions"); + const int64_t head_size = query_norm_shape[0]; + const int64_t num_heads = query_shape[1] / head_size; ORT_RETURN_IF_NOT(num_heads > 0 && head_size > 0, "PackedSparseAttentionIndexer: invalid query dimensions"); const int64_t width = 2 * head_size; @@ -943,6 +981,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont ORT_RETURN_IF_ERROR(CheckShape(past_seqlens, "past_sequence_lengths", {batch_size})); ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, width})); + ORT_RETURN_IF_ERROR(CheckShape(query_norm, "query_norm_weight", {head_size})); ORT_RETURN_IF_ERROR(CheckShape(norm, "key_norm_weight", {head_size})); ORT_RETURN_IF_ERROR(CheckShape(gate, "gate", {total_tokens, width})); ORT_RETURN_IF_ERROR(CheckShape(position_bias, "position_bias", {compress_ratio_, width})); @@ -1047,6 +1086,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont select.CacheHint(rotary.batched) .SetWorkgroupSize(kWorkgroupSize) .AddInputs({{query, ProgramTensorMetadataDependency::Type}, + {query_norm, ProgramTensorMetadataDependency::Type}, {present_key_state, ProgramTensorMetadataDependency::Type}, {head_weights, ProgramTensorMetadataDependency::Type}, {position_ids, ProgramTensorMetadataDependency::Type}, @@ -1067,6 +1107,7 @@ Status PackedSparseAttentionIndexer::ComputeCsa(onnxruntime::webgpu::ComputeCont {ToUint32(state_capacity)}, {ToUint32(capacity)}, {ToUint32(index_topk_)}, + {epsilon_}, {has_scale_ ? scale_ : 1.0f / std::sqrt(static_cast(head_size))}, {has_head_weight_scale_ ? head_weight_scale_ : 1.0f / std::sqrt(static_cast(num_heads))}}); diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index 386432c01e52f..870110b3f2dac 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -2421,9 +2421,10 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference propagateElemTypeFromInputToOutput(ctx, psai::kPastGateBuffer, psai::kPresentGateBuffer); } - const auto* query_shape = PackedSparseAttentionIndexerShape(ctx, psai::kQuery, 3); + const auto* query_shape = PackedSparseAttentionIndexerShape(ctx, psai::kQuery, 2); const auto* key_shape = PackedSparseAttentionIndexerShape(ctx, psai::kKey, 2); - const auto* norm_shape = PackedSparseAttentionIndexerShape(ctx, psai::kKeyNormWeight, 1); + const auto* query_norm_shape = PackedSparseAttentionIndexerShape(ctx, psai::kQueryNormWeight, 1); + const auto* key_norm_shape = PackedSparseAttentionIndexerShape(ctx, psai::kKeyNormWeight, 1); const auto* cumulative_shape = PackedSparseAttentionIndexerShape(ctx, psai::kCumulativeSequenceLengths, 1); const auto* past_sequence_shape = PackedSparseAttentionIndexerShape(ctx, psai::kPastSequenceLengths, 1); @@ -2479,14 +2480,15 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio); require_equal_dims(key_shape, 0, query_shape, 0, "key dimension 0 must equal query dimension 0"); - require_equal_dims(norm_shape, 0, query_shape, 2, "key_norm_weight dimension 0 must equal head_size"); + require_equal_dims(query_norm_shape, 0, key_norm_shape, 0, + "query_norm_weight and key_norm_weight dimensions must match"); require_equal_dims(past_sequence_shape, 0, key_state_shape, 0, "past_sequence_lengths dimension 0 must equal the state batch dimension"); require_equal_dims(key_state_shape, 0, kv_buffer_shape, 0, "past_key_state and past_kv_buffer batch dimensions must match"); require_equal_dims(state_lengths_shape, 0, key_state_shape, 0, "past_state_lengths dimension 0 must equal the state batch dimension"); - require_equal_dims(key_state_shape, 2, query_shape, 2, + require_equal_dims(key_state_shape, 2, query_norm_shape, 0, "past_key_state dimension 2 must equal head_size"); require_dim_value(key_state_shape, 1, state_capacity, "past_key_state dimension 1 must equal state_capacity"); require_dim_value(kv_buffer_shape, 1, buffer_capacity, @@ -2500,8 +2502,8 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference "PackedSparseAttentionIndexer: cumulative_sequence_lengths dimension 0 must equal batch_size + 1"); } - if (query_shape != nullptr && query_shape->dim(2).has_dim_value()) { - const int64_t head_size = query_shape->dim(2).dim_value(); + if (query_norm_shape != nullptr && query_norm_shape->dim(0).has_dim_value()) { + const int64_t head_size = query_norm_shape->dim(0).dim_value(); if (head_size <= 0) { fail_shape_inference("PackedSparseAttentionIndexer: head_size must be > 0, got ", head_size); } @@ -2522,8 +2524,14 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference require_equal_dims(gate_shape, 0, query_shape, 0, "gate dimension 0 must equal query dimension 0"); require_equal_dims(head_weights_shape, 0, query_shape, 0, "head_weights dimension 0 must equal query dimension 0"); - require_equal_dims(head_weights_shape, 1, query_shape, 1, - "head_weights dimension 1 must equal num_heads"); + if (head_weights_shape != nullptr && query_shape != nullptr && + head_weights_shape->dim(1).has_dim_value() && query_shape->dim(1).has_dim_value() && + query_norm_shape != nullptr && query_norm_shape->dim(0).has_dim_value() && + query_norm_shape->dim(0).dim_value() > 0 && + query_shape->dim(1).dim_value() / query_norm_shape->dim(0).dim_value() != + head_weights_shape->dim(1).dim_value()) { + fail_shape_inference("PackedSparseAttentionIndexer: head_weights dimension 1 must equal num_heads"); + } require_dim_value(position_bias_shape, 0, compress_ratio, "position_bias dimension 0 must equal compress_ratio"); require_equal_dims(gate_buffer_shape, 0, kv_buffer_shape, 0, @@ -2561,8 +2569,8 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference if (rotary_width <= 0 || (is_qsa && rotary_width % 2 != 0)) { fail_shape_inference("PackedSparseAttentionIndexer: invalid rotary cache width ", rotary_width); } - if (query_shape != nullptr && query_shape->dim(2).has_dim_value()) { - const int64_t head_size = query_shape->dim(2).dim_value(); + if (query_norm_shape != nullptr && query_norm_shape->dim(0).has_dim_value()) { + const int64_t head_size = query_norm_shape->dim(0).dim_value(); if ((is_qsa && rotary_width > head_size) || (!is_qsa && rotary_width > head_size / 2)) { fail_shape_inference("PackedSparseAttentionIndexer: rotary cache width is incompatible with head_size"); @@ -2573,9 +2581,15 @@ void PackedSparseAttentionIndexerTypeAndShapeInference(ONNX_NAMESPACE::Inference if (query_shape != nullptr) { const auto& total_tokens_dim = query_shape->dim(0); - const auto& num_heads_dim = query_shape->dim(1); - if (num_heads_dim.has_dim_value() && num_heads_dim.dim_value() <= 0) { - fail_shape_inference("PackedSparseAttentionIndexer: num_heads must be > 0, got ", num_heads_dim.dim_value()); + const auto& query_width_dim = query_shape->dim(1); + if (query_width_dim.has_dim_value() && query_width_dim.dim_value() <= 0) { + fail_shape_inference("PackedSparseAttentionIndexer: query width must be > 0, got ", + query_width_dim.dim_value()); + } + if (query_width_dim.has_dim_value() && query_norm_shape != nullptr && + query_norm_shape->dim(0).has_dim_value() && query_norm_shape->dim(0).dim_value() > 0 && + query_width_dim.dim_value() % query_norm_shape->dim(0).dim_value() != 0) { + fail_shape_inference("PackedSparseAttentionIndexer: query width must be divisible by head_size"); } const int64_t capacity = psai::SelectedCapacity(policy, token_budget, index_topk, compress_ratio); @@ -2653,7 +2667,8 @@ Common contract: with attention_mode="local_plus_selected", selected_kv_source="auxiliary"; key_state is layout-compatible with a [batch_size, capacity, 1, head_size] auxiliary cache when K = V). Unused entries are -1 and selected_counts holds the exact number of used entries. - * key_norm_weight is the effective RMSNorm multiplier, exactly as in SparseAttentionIndexer. + * query_norm_weight and key_norm_weight are the effective RMSNorm multipliers, exactly as in + SparseAttentionIndexer. * Accumulation, pooling, softmax, normalization and scoring are performed in float32 and the result is rounded once to the tensor element type. * Ties in the top-k selection are broken by the smaller entry index, and the emitted entries are @@ -2695,7 +2710,7 @@ ONNX_MS_OPERATOR_SET_SCHEMA( AttributeProto::INT, OPTIONAL_VALUE) .Attr("epsilon", - "Epsilon of the RMS normalization applied to the compressed keys. Default is 1e-6.", + "Epsilon of the RMS normalization applied to queries and compressed keys. Default is 1e-6.", AttributeProto::FLOAT, 1.0e-6f) .Attr("scale", @@ -2709,8 +2724,8 @@ ONNX_MS_OPERATOR_SET_SCHEMA( OPTIONAL_VALUE) .Input(0, "query", - "Packed indexer queries with shape (total_tokens, num_heads, head_size), already normalized but " - "not yet rotated.", + "Packed indexer queries with shape (total_tokens, num_heads * head_size), before normalization, " + "logical reshape, and rotary embedding.", "T") .Input(1, "key", @@ -2719,72 +2734,76 @@ ONNX_MS_OPERATOR_SET_SCHEMA( "head_size channels are the Ca series and the last head_size channels the Cb series.", "T") .Input(2, + "query_norm_weight", + "Effective RMSNorm multiplier of the queries, with shape (head_size).", + "T") + .Input(3, "key_norm_weight", "Effective RMSNorm multiplier of the compressed keys, with shape (head_size).", "T") - .Input(3, + .Input(4, "cos_cache", "Cosine rotary table indexed by absolute key position, shared across the batch with shape " "(max_rotary_sequence_length, rotary_width) or request-specific with shape " "(batch_size, max_rotary_sequence_length, rotary_width).", "T") - .Input(4, + .Input(5, "sin_cache", "Sine rotary table with the same shape as cos_cache.", "T") - .Input(5, + .Input(6, "cumulative_sequence_lengths", "Device-resident packed request boundaries with shape (batch_size + 1); " "cumulative_sequence_lengths[0] must be 0 and cumulative_sequence_lengths[batch_size] must equal " "total_tokens. Request b owns rows [cumulative_sequence_lengths[b], " "cumulative_sequence_lengths[b + 1]) of query/key (a repeated offset is a valid zero-token row).", "M") - .Input(6, + .Input(7, "past_sequence_lengths", "Device-resident number of tokens already processed for each request before this call, with " "shape (batch_size). Used as the default absolute query position when position_ids is omitted " "(policy_mode 'qsa'), and to validate state consistency.", "M") - .Input(7, + .Input(8, "gate", "Only for policy_mode 'csa': gate projection of the new tokens with shape " "(total_tokens, 2 * head_size).", "T", OpSchema::Optional) - .Input(8, + .Input(9, "position_bias", "Only for policy_mode 'csa': per-slot gate bias with shape (compress_ratio, 2 * head_size).", "T", OpSchema::Optional) - .Input(9, + .Input(10, "head_weights", "Only for policy_mode 'csa': per-head score weights with shape (total_tokens, num_heads).", "T", OpSchema::Optional) - .Input(10, + .Input(11, "position_ids", "Optional for policy_mode 'qsa', required for policy_mode 'csa': absolute position of every " "packed query, with shape (total_tokens).", "I", OpSchema::Optional) - .Input(11, + .Input(12, "past_key_state", "Generic fixed-capacity state: policy_mode 'qsa' stores prepared complete-block keys; " "policy_mode 'csa' stores compressed keys. Shape is (batch_size, state_capacity, head_size) and " "never changes across calls.", "T") - .Input(12, + .Input(13, "past_kv_buffer", "Generic fixed-capacity pending-token buffer. Shape is (batch_size, 2 * compress_ratio - 1, " "head_size) for policy_mode 'qsa' (which only ever uses up to compress_ratio - 1 of these " "entries) and (batch_size, 2 * compress_ratio - 1, 2 * head_size) for policy_mode 'csa'.", "T") - .Input(13, + .Input(14, "past_gate_buffer", "Only for policy_mode 'csa': buffered gate projections with the same shape as past_kv_buffer.", "T", OpSchema::Optional) - .Input(14, + .Input(15, "past_state_lengths", "Generic per-request state length with shape (batch_size, 2). Column 0 is the key_state entry " "count (policy_mode 'qsa': complete-block count; 'csa': compressed-entry count); column 1 is the " diff --git a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc index f97f2bfaeb513..c455ee1b64a44 100644 --- a/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -105,8 +105,8 @@ struct GraphOptions { int64_t BufferCapacity(int64_t compress_ratio) { return 2 * compress_ratio - 1; } -// Builds a fixed 15-input node; csa-only slots are left empty for policy_mode "qsa", as the schema -// requires. position_ids (slot 10) is optional for "qsa" and forced on for "csa". +// Builds a fixed 16-input node; csa-only slots are left empty for policy_mode "qsa", as the schema +// requires. position_ids (slot 11) is optional for "qsa" and forced on for "csa". void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { const bool is_csa = options.policy_mode == psai::kPolicyModeCsa; const int64_t width = is_csa ? 2 * options.head_size : options.head_size; @@ -115,10 +115,11 @@ void AddNode(ModelTestBuilder& builder, const GraphOptions& options) { std::vector inputs{ builder.MakeInput( - std::vector{options.total_tokens, options.num_heads, options.head_size}), + std::vector{options.total_tokens, options.num_heads * options.head_size}), builder.MakeInput( std::vector{options.key_total_tokens >= 0 ? options.key_total_tokens : options.total_tokens, width}), builder.MakeInput(std::vector{options.head_size}), + builder.MakeInput(std::vector{options.head_size}), builder.MakeInput(std::vector{64, options.rotary_width}), builder.MakeInput(std::vector{64, options.rotary_width}), builder.MakeInput(std::vector{options.batch_size + 1}), @@ -275,7 +276,7 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsZeroNumHeads) { GraphOptions options; options.num_heads = 0; ExpectResolveFailure([&options](ModelTestBuilder& builder) { AddNode(builder, options); }, - "num_heads must be > 0"); + "query width must be > 0"); } TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsGenericBufferCapacityOverflow) { @@ -321,9 +322,10 @@ TEST(PackedSparseAttentionIndexerShapeInferenceTest, RejectsCsaMissingPositionId const int64_t width = 2 * local.head_size; const int64_t buffer_capacity = BufferCapacity(local.compress_ratio); std::vector inputs{ - builder.MakeInput(std::vector{local.total_tokens, local.num_heads, local.head_size}), + builder.MakeInput(std::vector{local.total_tokens, local.num_heads * local.head_size}), builder.MakeInput(std::vector{local.total_tokens, width}), builder.MakeInput(std::vector{local.head_size}), + builder.MakeInput(std::vector{local.head_size}), builder.MakeInput(std::vector{64, local.rotary_width}), builder.MakeInput(std::vector{64, local.rotary_width}), builder.MakeInput(std::vector{local.batch_size + 1}), @@ -514,6 +516,7 @@ struct QsaPackedProblem { std::vector query; std::vector key; + std::vector query_norm_weight; std::vector key_norm_weight; std::vector cos_cache; // shared: [max_position, rotary_width] std::vector sin_cache; @@ -618,6 +621,7 @@ void QsaPackedReference(const QsaPackedProblem& p, QsaPackedResult& out) { for (int h = 0; h < p.num_heads; ++h) { const size_t base = (static_cast(token) * p.num_heads + h) * head_size; std::vector head(p.query.begin() + base, p.query.begin() + base + head_size); + head = RmsNormalize(head, p.query_norm_weight, p.epsilon); const int clamped_position = std::min(std::max(position, 0), p.max_position - 1); rotated_query[static_cast(h)] = LeadingRope(head, p.rotary_width, p.cos_cache.data() + clamped_position * p.rotary_width, @@ -676,6 +680,7 @@ QsaPackedProblem MakeQsaPackedProblem(QsaPackedProblem problem = {}) { problem.query = MakeWave(static_cast(total_tokens) * problem.num_heads * problem.head_size, 0.35f, 0.41f); problem.key = MakeWave(static_cast(total_tokens) * problem.head_size, 1.10f, 0.29f); + problem.query_norm_weight = MakeWave(static_cast(problem.head_size), 1.30f, 0.23f); problem.key_norm_weight = MakeWave(static_cast(problem.head_size), 0.70f, 0.17f); problem.cos_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 0.20f, 0.13f); problem.sin_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 0.90f, 0.19f); @@ -697,6 +702,7 @@ void RunQsaPackedTest(float tolerance, QsaPackedProblem problem = MakeQsaPackedP problem.query = RoundTrip(problem.query); problem.key = RoundTrip(problem.key); + problem.query_norm_weight = RoundTrip(problem.query_norm_weight); problem.key_norm_weight = RoundTrip(problem.key_norm_weight); problem.cos_cache = RoundTrip(problem.cos_cache); problem.sin_cache = RoundTrip(problem.sin_cache); @@ -719,8 +725,9 @@ void RunQsaPackedTest(float tolerance, QsaPackedProblem problem = MakeQsaPackedP if (problem.scale.has_value()) { test.AddAttribute("scale", *problem.scale); } - test.AddInput("query", {total_tokens, problem.num_heads, head_size}, ToElementType(problem.query)); + test.AddInput("query", {total_tokens, problem.num_heads * head_size}, ToElementType(problem.query)); test.AddInput("key", {total_tokens, head_size}, ToElementType(problem.key)); + test.AddInput("query_norm_weight", {head_size}, ToElementType(problem.query_norm_weight)); test.AddInput("key_norm_weight", {head_size}, ToElementType(problem.key_norm_weight)); test.AddInput("cos_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.cos_cache)); test.AddInput("sin_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.sin_cache)); @@ -847,6 +854,7 @@ struct CsaPackedProblem { std::vector query; std::vector key; + std::vector query_norm_weight; std::vector key_norm_weight; std::vector cos_cache; std::vector sin_cache; @@ -1001,6 +1009,7 @@ void CsaPackedReference(const CsaPackedProblem& p, CsaPackedResult& out) { for (int h = 0; h < p.num_heads; ++h) { const size_t base = (static_cast(token) * p.num_heads + h) * head_size; std::vector head(p.query.begin() + base, p.query.begin() + base + head_size); + head = RmsNormalize(head, p.query_norm_weight, p.epsilon); const int clamped_position = std::min(std::max(position, 0), p.max_position - 1); rotated_query[static_cast(h)] = TrailingRope(head, p.rotary_width, p.cos_cache.data() + clamped_position * p.rotary_width, @@ -1046,6 +1055,7 @@ CsaPackedProblem MakeCsaPackedProblem(CsaPackedProblem problem = {}) { problem.query = MakeWave(static_cast(total_tokens) * problem.num_heads * problem.head_size, 0.25f, 0.37f); problem.key = MakeWave(static_cast(total_tokens) * width, 0.60f, 0.21f); + problem.query_norm_weight = MakeWave(static_cast(problem.head_size), 1.25f, 0.19f); problem.key_norm_weight = MakeWave(static_cast(problem.head_size), 0.45f, 0.31f); problem.cos_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 0.15f, 0.27f); problem.sin_cache = MakeWave(static_cast(problem.max_position) * problem.rotary_width, 1.05f, 0.33f); @@ -1082,6 +1092,7 @@ void RunCsaPackedTest(const CsaPackedProblem& base, float tolerance, CsaPackedProblem problem = base; problem.query = RoundTrip(problem.query); problem.key = RoundTrip(problem.key); + problem.query_norm_weight = RoundTrip(problem.query_norm_weight); problem.key_norm_weight = RoundTrip(problem.key_norm_weight); problem.cos_cache = RoundTrip(problem.cos_cache); problem.sin_cache = RoundTrip(problem.sin_cache); @@ -1109,8 +1120,9 @@ void RunCsaPackedTest(const CsaPackedProblem& base, float tolerance, if (problem.scale.has_value()) test.AddAttribute("scale", *problem.scale); if (problem.head_weight_scale.has_value()) test.AddAttribute("head_weight_scale", *problem.head_weight_scale); - test.AddInput("query", {total_tokens, problem.num_heads, head_size}, ToElementType(problem.query)); + test.AddInput("query", {total_tokens, problem.num_heads * head_size}, ToElementType(problem.query)); test.AddInput("key", {total_tokens, width}, ToElementType(problem.key)); + test.AddInput("query_norm_weight", {head_size}, ToElementType(problem.query_norm_weight)); test.AddInput("key_norm_weight", {head_size}, ToElementType(problem.key_norm_weight)); test.AddInput("cos_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.cos_cache)); test.AddInput("sin_cache", {problem.max_position, problem.rotary_width}, ToElementType(problem.sin_cache));