diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index ae4a2c5ad8ad4..97e14ecdc919f 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 @@ -4326,6 +4327,17 @@ This version of the operator has been available since version 1 of the 'com.micr 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. Optional inputs add packed-sequence and Qwen4-Exp-style n-gram embedding support: @@ -4730,6 +4742,163 @@ 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); 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. + +#### 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 and 2 * compress_ratio - 1 must not exceed INT_MAX.
+
epsilon : float
+
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
+
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 + +
+
query : T
+
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
+
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 + +
+
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. @@ -7868,3 +8037,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 d253fecc316b3..dcd3a06cc3314 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..6f7f7e8ba0bf5 --- /dev/null +++ b/docs/contrib_ops/packed_sparse_attention_indexer.md @@ -0,0 +1,269 @@ +# 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 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), +[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 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 | +| 1 | `key` | `(total_tokens, head_size)` qsa / `(total_tokens, 2*head_size)` csa | T | both | +| 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): + +| # | 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. 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 +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..33ade715ef995 --- /dev/null +++ b/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md @@ -0,0 +1,50 @@ +# 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. +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 +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..0d55ae4c5e5e5 --- /dev/null +++ b/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h @@ -0,0 +1,86 @@ +// 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 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] + kKey = 1, // qsa: [total_tokens, head_size]; csa: [total_tokens, 2 * head_size] + 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 +// 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 64a24ed3ec74c..02f652d7a858f 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 { @@ -67,8 +77,8 @@ constexpr int kCsaOutputCount = 3; // 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; } @@ -86,8 +96,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 124a8f1bcc9a3..167faab950a78 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); @@ -601,6 +604,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..6e01ecf2c6d6a --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc @@ -0,0 +1,473 @@ +// 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()) \ + .MayInplace(11, 2) \ + .MayInplace(14, 5), \ + 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_ <= (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"); + 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* 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); + 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() == 2, + "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads * head_size), " + "got ", + query_shape.ToString()); + const int64_t total_tokens = query_shape[0]; + 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)); + + 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(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})); + 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})); + } + + 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)); + 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})); + 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* 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); + 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(query_norm_weight->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* 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); + 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() == 2, + "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads * head_size), " + "got ", + query_shape.ToString()); + const int64_t total_tokens = query_shape[0]; + 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)); + 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(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})); + 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})); + 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)); + 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})); + 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(query_norm_weight->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..c8f3ec19f2082 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h @@ -0,0 +1,38 @@ +// 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 state_capacity_; + 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..620247fc41d48 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu @@ -0,0 +1,901 @@ +// 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 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) { + 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 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; + 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 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. + __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] = rejected ? 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 (!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; + 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* 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); + 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(); + + 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] + : 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 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, + 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 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); + + 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; + 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 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 = + 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) + // 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] = rejected ? 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 (!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; + 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* 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, + 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_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, 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; + 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* 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, + 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_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, 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; + 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 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) +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..9d92076c2ff51 --- /dev/null +++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h @@ -0,0 +1,108 @@ +// 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* 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, + float* float_workspace, + int32_t* overflow_flags); + +template +Status LaunchCsaPackedSparseAttentionIndexer( + cudaStream_t stream, + 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, + 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 a897fd04a2a63..ecccf91b7a038 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..461cc03666525 --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc @@ -0,0 +1,1119 @@ +// 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()) + .MayInplace(psai::kPastKeyState, psai::kPresentKeyState) + .MayInplace(psai::kPastKvBuffer, psai::kPresentKvBuffer) + .MayInplace(psai::kPastGateBuffer, psai::kPresentGateBuffer) + .MayInplace(psai::kPastStateLengths, psai::kPresentStateLengths), + 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("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& 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 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 = 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" + << " if (virtual_pos < old_buf_len) {\n" + << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * uniforms.head_size + d;\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" + << "}\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() + << " 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() + << " 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 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) || " + << "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 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" + << " 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& 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); + 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& 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); + + 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 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 = 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 * normalized_query_value(token, head, 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() + << " 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" + << " 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" + << " 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" + << " 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_gate = shader.AddInput("key_gate", ShaderUsage::UseUniform); + const auto& norm = shader.AddInput("key_norm_weight", ShaderUsage::UseUniform); + const auto& rotary_cache = shader.AddInput("rotary_cache", ShaderUsage::UseUniform); + const auto& position_bias = shader.AddInput("position_bias", 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 = + 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); + const auto& overflow_flags = shader.AddOutput("overflow_flags", ShaderUsage::UseUniform); + 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" + << " 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(" << kv_buffer_history.GetByOffset("idx") << ");\n" + << " }\n" + << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\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" + << " if (virtual_pos < old_buf_len) {\n" + << " let idx = (b * uniforms.buffer_capacity + u32(virtual_pos)) * width + channel;\n" + << " return f32(" << gate_buffer_history.GetByOffset("idx") << ");\n" + << " }\n" + << " let idx2 = u32(req_start + virtual_pos - old_buf_len) * width + channel;\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" + << " 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 = " << 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 = " << 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) || " + << 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" + << " " << 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 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 (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" + << " 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" + << " 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", + "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& 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); + 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); + 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 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 = 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 * 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() + << " 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() + << " 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" + << " 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" + << " 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" + << " 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 && + 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"); + + 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* 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); + 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() == 2, "PackedSparseAttentionIndexer: query must have rank 2"); + const int64_t total_tokens = query_shape[0]; + 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"); + 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(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})); + 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})); + } + + 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]; + 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})); + 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, + "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); + 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})); + 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(); + 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)}, {0u}}); + 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)}, {0u}}); + 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)}, {0u}}); + 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}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, + {cu_seqlens, ProgramTensorMetadataDependency::Type}, + {past_seqlens, 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)) + .AddUniformVariables({{ToUint32(batch_size)}, + {ToUint32(total_tokens)}, + {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}, + {query_norm, ProgramTensorMetadataDependency::Type}, + {present_key_state, ProgramTensorMetadataDependency::Type}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, + {cu_seqlens, ProgramTensorMetadataDependency::Type}, + {past_seqlens, ProgramTensorMetadataDependency::Type}}); + if (position_ids != nullptr) { + select.AddInput({position_ids, 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)) + .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* 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); + 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() == 2, "PackedSparseAttentionIndexer: query must have rank 2"); + const int64_t total_tokens = query_shape[0]; + 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; + + 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(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})); + 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})); + 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]; + 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})); + 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})); + + 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})); + 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})); + 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()) { + 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)}, {0u}}); + 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_gate, ProgramTensorMetadataDependency::Type}, + {norm, ProgramTensorMetadataDependency::Type}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, + {position_bias, ProgramTensorMetadataDependency::Type}, + {&sequence_metadata, 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)) + .AddUniformVariables({{ToUint32(batch_size)}, + {ToUint32(total_tokens)}, + {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}, + {query_norm, ProgramTensorMetadataDependency::Type}, + {present_key_state, ProgramTensorMetadataDependency::Type}, + {head_weights, ProgramTensorMetadataDependency::Type}, + {position_ids, ProgramTensorMetadataDependency::Type}, + {&rotary_cache, ProgramTensorMetadataDependency::Type}, + {cu_seqlens, 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)) + .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_)}, + {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))}}); + 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..2607a1492e48f --- /dev/null +++ b/onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h @@ -0,0 +1,157 @@ +// 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}, + {"dst_offset", 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}, + {"total_tokens", 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}, + {"total_tokens", 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 state_capacity_; + 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 c90bc09680fa7..0ef6cb7e09651 100644 --- a/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc +++ b/onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc @@ -13,6 +13,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" @@ -56,6 +57,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 deba6cfe244dd..7b1d93f793706 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" @@ -13,6 +14,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) @@ -2359,6 +2361,543 @@ 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 > (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()) { + 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); + } + + const auto* query_shape = PackedSparseAttentionIndexerShape(ctx, psai::kQuery, 2); + const auto* key_shape = PackedSparseAttentionIndexerShape(ctx, psai::kKey, 2); + 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); + 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) { + 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)) { + 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(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_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, + "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_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); + } + 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"); + 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, + "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_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"); + } + } + } + } + + if (query_shape != nullptr) { + const auto& total_tokens_dim = query_shape->dim(0); + 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); + 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_*. + if (key_state_shape != nullptr) { + updateOutputShape(ctx, psai::kPresentKeyState, *key_state_shape); + } + if (kv_buffer_shape != nullptr) { + updateOutputShape(ctx, psai::kPresentKvBuffer, *kv_buffer_shape); + } + if (gate_buffer_shape != nullptr) { + updateOutputShape(ctx, psai::kPresentGateBuffer, *gate_buffer_shape); + } + 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. + * 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 + 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 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.", + 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 queries and 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), before normalization, " + "logical reshape, and rotary embedding.", + "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, + "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(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(5, + "sin_cache", + "Sine rotary table with the same shape as cos_cache.", + "T") + .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(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(8, + "gate", + "Only for policy_mode 'csa': gate projection of the new tokens with shape " + "(total_tokens, 2 * head_size).", + "T", + OpSchema::Optional) + .Input(9, + "position_bias", + "Only for policy_mode 'csa': per-slot gate bias with shape (compress_ratio, 2 * head_size).", + "T", + OpSchema::Optional) + .Input(10, + "head_weights", + "Only for policy_mode 'csa': per-head score weights with shape (total_tokens, num_heads).", + "T", + OpSchema::Optional) + .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(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(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(14, + "past_gate_buffer", + "Only for policy_mode 'csa': buffered gate projections with the same shape as past_kv_buffer.", + "T", + OpSchema::Optional) + .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 " + "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 8365f039e7a8c..e3aec0594ed39 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", @@ -2708,6 +2710,55 @@ def past_shape(index): set_output(1, [query_shape[0], present_compressed_length, query_shape[3]]) set_output(2, [2, query_shape[0], present_buffer_length, past_buffer_shape[3]]) + 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..c455ee1b64a44 --- /dev/null +++ b/onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc @@ -0,0 +1,1223 @@ +// 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 = 4; + int64_t compress_ratio = 2; + 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; + 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 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; + 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.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}), + 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); + } + 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, input_state_capacity, options.head_size})); + inputs.push_back( + 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, + options.gate_buffer_width >= 0 ? options.gate_buffer_width : 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, 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, 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"; + 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); }, + "query width 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; + 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{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); }, + "output size 4 not in range [min=6, max=6]"); +} + +#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; + } +} + +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) { + 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 query_norm_weight; + 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 = 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]; + 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 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 + ? 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) { + 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) { + 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); + 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, + 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.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); + 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, QsaPackedResult* actual = nullptr) { + 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.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); + 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("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)); + 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)); + 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 + +TEST(PackedSparseAttentionIndexerTest, QsaFloat) { RunQsaPackedTest(1.0e-5f); } + +TEST(PackedSparseAttentionIndexerTest, QsaFloat16) { RunQsaPackedTest(2.0e-3f); } + +TEST(PackedSparseAttentionIndexerTest, QsaBFloat16) { RunQsaPackedTest(2.0e-2f); } + +TEST(PackedSparseAttentionIndexerTest, QsaPrefillThenDecodeIndependentState) { + 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 +// 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 +// reject the request's step without partially changing its state. +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))); +} + +#ifdef USE_WEBGPU +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; + 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 { + 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 query_norm_weight; + 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 = 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)]; + 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; + 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 (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) { + 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) { + 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); + 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, + 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.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); + 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, CsaPackedResult* actual = nullptr) { + 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.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); + 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("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)); + 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)); + 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 + +TEST(PackedSparseAttentionIndexerTest, CsaFloat) { RunCsaPackedTest(MakeCsaPackedProblem(), 1.0e-5f); } + +TEST(PackedSparseAttentionIndexerTest, CsaFloat16) { RunCsaPackedTest(MakeCsaPackedProblem(), 4.0e-3f); } + +TEST(PackedSparseAttentionIndexerTest, CsaBFloat16) { RunCsaPackedTest(MakeCsaPackedProblem(), 3.0e-2f); } + +TEST(PackedSparseAttentionIndexerTest, CsaPrefillThenDecodeIndependentState) { + 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) { + 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, CsaFloat16) { + RunCsaPackedTest(MakeCsaPackedProblem(), 4.0e-3f, 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 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 7aabc05004f03..d1e936529a0d9 100644 --- a/onnxruntime/test/python/onnxruntime_test_python_symbolic_shape_infer.py +++ b/onnxruntime/test/python/onnxruntime_test_python_symbolic_shape_infer.py @@ -341,6 +341,135 @@ def test_sparse_attention_indexer_csa_symbolic_fallback(self): self.assertTrue(compressed_shape[1].startswith("SparseAttentionIndexer_")) self.assertTrue(buffer_shape[2].startswith("SparseAttentionIndexer_")) + 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( [