From dbaeb30d6cc0a9be3fdc4318b30113f5dc6ab7b5 Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Tue, 15 Sep 2026 22:11:28 +0000
Subject: [PATCH 01/10] Initial plan
From cf546932bfda635b0164eb88c38fc1b5da5b653a Mon Sep 17 00:00:00 2001
From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com>
Date: Wed, 16 Sep 2026 02:18:34 +0000
Subject: [PATCH 02/10] Implement packed sparse attention indexer
Co-authored-by: kunal-vaishnavi <115581922+kunal-vaishnavi@users.noreply.github.com>
---
docs/ContribOperators.md | 158 ++-
docs/OperatorKernels.md | 1 +
.../packed_sparse_attention_indexer.md | 267 +++++
.../webgpu/packed_sparse_attention_indexer.md | 48 +
.../packed_sparse_attention_indexer_common.h | 85 ++
.../sparse/sparse_attention_indexer_common.h | 18 +-
.../contrib_ops/cuda/cuda_contrib_kernels.cc | 6 +
.../sparse/packed_sparse_attention_indexer.cc | 437 +++++++
.../sparse/packed_sparse_attention_indexer.h | 37 +
.../packed_sparse_attention_indexer_impl.cu | 868 ++++++++++++++
.../packed_sparse_attention_indexer_impl.h | 106 ++
.../sparse_attention_indexer_device_math.cuh | 149 +++
.../sparse/sparse_attention_indexer_impl.cu | 96 +-
.../bert/packed_sparse_attention_indexer.cc | 982 ++++++++++++++++
.../bert/packed_sparse_attention_indexer.h | 151 +++
.../webgpu/webgpu_contrib_kernels.cc | 2 +
.../core/graph/contrib_ops/bert_defs.cc | 383 ++++++
onnxruntime/core/graph/contrib_ops/ms_opset.h | 2 +
.../python/tools/symbolic_shape_infer.py | 51 +
...packed_sparse_attention_indexer_op_test.cc | 1044 +++++++++++++++++
...untime_test_python_symbolic_shape_infer.py | 129 ++
21 files changed, 4931 insertions(+), 89 deletions(-)
create mode 100644 docs/contrib_ops/packed_sparse_attention_indexer.md
create mode 100644 docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md
create mode 100644 onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h
create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc
create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h
create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu
create mode 100644 onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h
create mode 100644 onnxruntime/contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh
create mode 100644 onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc
create mode 100644 onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.h
create mode 100644 onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc
diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md
index 3cc7009d0e83e..875a3483f87d4 100644
--- a/docs/ContribOperators.md
+++ b/docs/ContribOperators.md
@@ -79,6 +79,7 @@ Do not modify directly.*
* com.microsoft.NhwcMaxPool
* com.microsoft.PackedAttention
* com.microsoft.PackedMultiHeadAttention
+ * com.microsoft.PackedSparseAttentionIndexer
* com.microsoft.Pad
* com.microsoft.PagedAttention
* com.microsoft.QAttention
@@ -4672,6 +4673,161 @@ This version of the operator has been available since version 1 of the 'com.micr
+### **com.microsoft.PackedSparseAttentionIndexer**
+
+ Packed/variable-length counterpart of SparseAttentionIndexer, for continuous-batching engines
+ (such as an OgaEngine-style PagedAttention model) that flatten every request's tokens into one
+ [total_tokens, ...] axis instead of a dense [batch_size, sequence_length, ...] axis. It selects, for
+ every packed query token, the sparse-attention candidates that a following SparsePagedAttention (or
+ similar) operator is allowed to read.
+
+ Unlike SparseAttentionIndexer, this operator:
+ * takes packed query/key tensors plus cumulative_sequence_lengths (request boundaries) and
+ past_sequence_lengths (per-request past length) instead of a dense batch and a dense mask;
+ * derives ordinary causal visibility purely from that packed metadata -- there is no mask input;
+ * uses a single generic set of state slots (past_key_state / past_kv_buffer / past_gate_buffer /
+ past_state_lengths) for both policy_mode values, each with a shape that is fixed across calls
+ (state never grows and is never concatenated); state overflow beyond the fixed capacity is
+ rejected as a deterministic no-op on state rather than truncated or allowed to corrupt memory;
+ * additionally emits selected_counts, the exact number of active (non -1) entries per query, so
+ that no downstream consumer needs to scan selected_indices for its query's true count.
+
+ Both policy_mode values keep the semantics of SparseAttentionIndexer, applied independently to each
+ request's own packed token range and fixed-capacity state slice:
+
+ policy_mode = "qsa" ("query sparse attention" token indexer)
+ Processes each request's new tokens sequentially: appends raw indexer keys to the generic
+ pending buffer, and whenever it reaches compress_ratio tokens, mean-pools it, applies RMSNorm
+ and key_norm_weight, applies the leading/split-half rotary convention at the block's first
+ logical token position, and appends the prepared (already normalized and rotated) key to
+ key_state. Queries are scored against every causally visible complete block with
+ sum_h ReLU(q_h . k), the token_budget / compress_ratio highest scoring blocks are kept, and
+ their token indices are emitted (request-local logical positions, i.e. the same numbering as
+ past_sequence_lengths + local offset) followed by the causally visible tokens of the trailing
+ incomplete block.
+
+ policy_mode = "csa" ("compressed sparse attention" block indexer)
+ Applies the same window-plan arithmetic as SparseAttentionIndexer (overlap/leftover/new window
+ count) independently per request, using that request's own buffer_length and new token count;
+ every newly closed window is compressed with the softmax-gated Ca/Cb pooling, normalized,
+ rotated and appended to key_state. Queries are scored against every causally visible compressed
+ entry with sum_h w_h * ReLU(q_h . k) and the index_topk highest scoring entry indices are
+ emitted.
+
+ Common contract:
+ * selected_indices is int32 with a fixed capacity that only depends on attributes:
+ token_budget + compress_ratio - 1 for "qsa" (values are request-local token positions into the
+ main key/value cache, directly consumable by SparsePagedAttention configured with
+ attention_mode="selected_only", selected_kv_source="main") and index_topk for "csa" (values are
+ compressed-entry indices into key_state, directly consumable by SparsePagedAttention configured
+ with attention_mode="local_plus_selected", selected_kv_source="auxiliary"; key_state is
+ layout-compatible with a [batch_size, capacity, 1, head_size] auxiliary cache when K = V).
+ Unused entries are -1 and selected_counts holds the exact number of used entries.
+ * key_norm_weight is the effective RMSNorm multiplier, exactly as in SparseAttentionIndexer.
+ * Accumulation, pooling, softmax, normalization and scoring are performed in float32 and the
+ result is rounded once to the tensor element type.
+ * Ties in the top-k selection are broken by the smaller entry index, and the emitted entries are
+ ordered by decreasing score, so the result is deterministic.
+ * cos_cache / sin_cache may be shared across the batch ([max_position, rotary_width]) or
+ request-specific ([batch_size, max_position, rotary_width]).
+ * cumulative_sequence_lengths, past_sequence_lengths and past_state_lengths are read directly by
+ the device kernel; a zero-token request row (a repeated cumulative offset) is valid and simply
+ contributes no query rows for that request.
+
+ OgaEngine integration note: this operator only defines the ORT operator; wiring
+ past_key_state / past_kv_buffer / past_gate_buffer / past_state_lengths as Engine-managed,
+ per-request fixed-size state (analogous to a paged auxiliary cache) is expected to happen in the
+ OgaEngine / Model Builder integration, which is out of scope for this operator definition.
+
+#### Version
+
+This version of the operator has been available since version 1 of the 'com.microsoft' operator set.
+
+#### Attributes
+
+
+- compress_ratio : int (required)
+- Number of consecutive tokens folded into one compressed/pooled entry. Must be > 0.
+- epsilon : float
+- Epsilon of the RMS normalization applied to the compressed keys. Default is 1e-6.
+- head_weight_scale : float
+- Only for policy_mode 'csa': scale applied to head_weights. Default is 1/sqrt(num_heads). Must be omitted when policy_mode is 'qsa'.
+- index_topk : int
+- Only for policy_mode 'csa': number of compressed entries selected per query. Must be > 0. Must be omitted when policy_mode is 'qsa'.
+- policy_mode : string (required)
+- Indexer policy. Must be exactly 'qsa' (token indexer) or 'csa' (compressed block indexer).
+- scale : float
+- Scale applied to the per-head ReLU scores. Default is 1/sqrt(head_size).
+- state_capacity : int (required)
+- Fixed capacity (number of entries) of past_key_state / present_key_state. Must be > 0.
+- token_budget : int
+- Only for policy_mode 'qsa': maximum number of tokens selected from complete blocks. Must be > 0 and divisible by compress_ratio. Must be omitted when policy_mode is 'csa'.
+
+
+#### Inputs (10 - 15)
+
+
+- query : T
+- Packed indexer queries with shape (total_tokens, num_heads, head_size), already normalized but not yet rotated.
+- key : T
+- Packed indexer key projection of the new tokens. Shape is (total_tokens, head_size) for policy_mode 'qsa' and (total_tokens, 2 * head_size) for policy_mode 'csa', where the first head_size channels are the Ca series and the last head_size channels the Cb series.
+- key_norm_weight : T
+- Effective RMSNorm multiplier of the compressed keys, with shape (head_size).
+- cos_cache : T
+- Cosine rotary table indexed by absolute key position, shared across the batch with shape (max_rotary_sequence_length, rotary_width) or request-specific with shape (batch_size, max_rotary_sequence_length, rotary_width).
+- sin_cache : T
+- Sine rotary table with the same shape as cos_cache.
+- cumulative_sequence_lengths : M
+- Device-resident packed request boundaries with shape (batch_size + 1); cumulative_sequence_lengths[0] must be 0 and cumulative_sequence_lengths[batch_size] must equal total_tokens. Request b owns rows [cumulative_sequence_lengths[b], cumulative_sequence_lengths[b + 1]) of query/key (a repeated offset is a valid zero-token row).
+- past_sequence_lengths : M
+- Device-resident number of tokens already processed for each request before this call, with shape (batch_size). Used as the default absolute query position when position_ids is omitted (policy_mode 'qsa'), and to validate state consistency.
+- gate (optional) : T
+- Only for policy_mode 'csa': gate projection of the new tokens with shape (total_tokens, 2 * head_size).
+- position_bias (optional) : T
+- Only for policy_mode 'csa': per-slot gate bias with shape (compress_ratio, 2 * head_size).
+- head_weights (optional) : T
+- Only for policy_mode 'csa': per-head score weights with shape (total_tokens, num_heads).
+- position_ids (optional) : I
+- Optional for policy_mode 'qsa', required for policy_mode 'csa': absolute position of every packed query, with shape (total_tokens).
+- past_key_state : T
+- Generic fixed-capacity state: policy_mode 'qsa' stores prepared complete-block keys; policy_mode 'csa' stores compressed keys. Shape is (batch_size, state_capacity, head_size) and never changes across calls.
+- past_kv_buffer : T
+- Generic fixed-capacity pending-token buffer. Shape is (batch_size, 2 * compress_ratio - 1, head_size) for policy_mode 'qsa' (which only ever uses up to compress_ratio - 1 of these entries) and (batch_size, 2 * compress_ratio - 1, 2 * head_size) for policy_mode 'csa'.
+- past_gate_buffer (optional) : T
+- Only for policy_mode 'csa': buffered gate projections with the same shape as past_kv_buffer.
+- past_state_lengths : M
+- Generic per-request state length with shape (batch_size, 2). Column 0 is the key_state entry count (policy_mode 'qsa': complete-block count; 'csa': compressed-entry count); column 1 is the pending-buffer length (policy_mode 'qsa': incomplete-block length in [0, compress_ratio); 'csa': buffer length in [0, 2 * compress_ratio)).
+
+
+#### Outputs (6 - 6)
+
+
+- selected_indices : M
+- Selected entries with shape (total_tokens, capacity). capacity is token_budget + compress_ratio - 1 for policy_mode 'qsa' (request-local token positions into the main key/value cache) and index_topk for policy_mode 'csa' (compressed entry indices into key_state). Unused entries are -1.
+- selected_counts : M
+- Exact number of used (non -1) entries of selected_indices for every query, with shape (total_tokens).
+- present_key_state : T
+- Updated generic key state, with the same fixed shape as past_key_state.
+- present_kv_buffer : T
+- Updated generic pending-token buffer, with the same fixed shape as past_kv_buffer.
+- present_gate_buffer (optional) : T
+- Only for policy_mode 'csa': updated gate buffer with the same fixed shape as past_gate_buffer.
+- present_state_lengths : M
+- Updated generic per-request state length, with the same fixed shape as past_state_lengths.
+
+
+#### Type Constraints
+
+
+- T : tensor(float), tensor(float16), tensor(bfloat16)
+- Constrain floating point tensors to float, float16 and bfloat16.
+- I : tensor(int64)
+- Constrain position ids to 64-bit integer tensors.
+- M : tensor(int32)
+- Constrain packed metadata, generic state lengths and selected indices/counts to 32-bit integer tensors.
+
+
+
### **com.microsoft.Pad**
Given `data` tensor, pads, mode, and value.
@@ -7814,5 +7970,3 @@ No versioning maintained for experimental ops.
T : tensor(float)
Constrain input and output types to float32 tensors.
-
-
diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md
index 59a8b35f33420..9a72c7b1dac12 100644
--- a/docs/OperatorKernels.md
+++ b/docs/OperatorKernels.md
@@ -1122,6 +1122,7 @@ The **OpSet Version** column uses the following notation:
|NhwcConv|*in* X:**T**
*in* W:**T**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|PackedAttention|*in* input:**T**
*in* weights:**T**
*in* bias:**T**
*in* token_offset:**M**
*in* cumulative_sequence_length:**M**
*in* attention_bias:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
|PackedMultiHeadAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* bias:**T**
*in* token_offset:**M**
*in* cumulative_sequence_length:**M**
*in* attention_bias:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
+|PackedSparseAttentionIndexer|*in* query:**T**
*in* key:**T**
*in* key_norm_weight:**T**
*in* cos_cache:**T**
*in* sin_cache:**T**
*in* cumulative_sequence_lengths:**M**
*in* past_sequence_lengths:**M**
*in* gate:**T**
*in* position_bias:**T**
*in* head_weights:**T**
*in* position_ids:**I**
*in* past_key_state:**T**
*in* past_kv_buffer:**T**
*in* past_gate_buffer:**T**
*in* past_state_lengths:**M**
*out* selected_indices:**M**
*out* selected_counts:**M**
*out* present_key_state:**T**
*out* present_kv_buffer:**T**
*out* present_gate_buffer:**T**
*out* present_state_lengths:**M**|1+|**I** = tensor(int64)
**M** = tensor(int32)
**T** = tensor(bfloat16), tensor(float), tensor(float16)|
|PagedAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* key_cache:**T_CACHE**
*in* value_cache:**T_CACHE**
*in* cumulative_sequence_length:**S**
*in* past_seqlens:**S**
*in* block_table:**S**
*in* cos_cache:**T**
*in* sin_cache:**T**
*in* slot_mapping:**S**
*in* head_sink:**T**
*in* q_norm_weight:**T**
*in* k_norm_weight:**T**
*in* k_scale:**T_KV_SCALE**
*in* v_scale:**T_KV_SCALE**
*in* attention_metadata:**S**
*out* output:**T**
*out* key_cache_out:**T_CACHE**
*out* value_cache_out:**T_CACHE**|1+|**S** = tensor(int32)
**T** = tensor(bfloat16), tensor(float16)
**T_CACHE** = tensor(bfloat16), tensor(float16), tensor(float8e4m3fn), tensor(int8), tensor(uint8)
**T_KV_SCALE** = tensor(float)|
|QAttention|*in* input:**T1**
*in* weight:**T2**
*in* bias:**T3**
*in* input_scale:**T3**
*in* weight_scale:**T3**
*in* mask_index:**T4**
*in* input_zero_point:**T1**
*in* weight_zero_point:**T2**
*in* past:**T3**
*out* output:**T3**
*out* present:**T3**|1+|**T1** = tensor(int8)
**T2** = tensor(int8)
**T3** = tensor(float), tensor(float16)
**T4** = tensor(int32)|
|QMoE|*in* input:**T**
*in* router_probs:**T**
*in* fc1_experts_weights:**T1**
*in* fc1_scales:**T2**
*in* fc1_experts_bias:**T**
*in* fc2_experts_weights:**T1**
*in* fc2_scales:**T2**
*in* fc2_experts_bias:**T**
*in* fc3_experts_weights:**T1**
*in* fc3_scales:**T2**
*in* fc3_experts_bias:**T**
*in* fc1_zero_points:**T1**
*in* fc2_zero_points:**T1**
*in* fc3_zero_points:**T1**
*in* router_weights:**T**
*in* fc1_global_scale:**T4**
*in* fc2_global_scale:**T4**
*in* fc1_act_scale:**T4**
*in* fc2_act_scale:**T4**
*in* fc1_act_block_scale:**T2**
*in* fc2_act_block_scale:**T2**
*out* output:**T**|1+|**T** = tensor(bfloat16), tensor(float16)
**T1** = tensor(float8e4m3fn), tensor(uint8)
**T2** = tensor(bfloat16), tensor(float16), tensor(float8e4m3fn), tensor(float8e8m0)
**T4** = tensor(float)|
diff --git a/docs/contrib_ops/packed_sparse_attention_indexer.md b/docs/contrib_ops/packed_sparse_attention_indexer.md
new file mode 100644
index 0000000000000..712e73d2da278
--- /dev/null
+++ b/docs/contrib_ops/packed_sparse_attention_indexer.md
@@ -0,0 +1,267 @@
+# PackedSparseAttentionIndexer — Operator Documentation
+
+This document describes the `com.microsoft::PackedSparseAttentionIndexer` contrib operator: the
+packed/variable-length counterpart of `com.microsoft::SparseAttentionIndexer`, built for
+continuous-batching (paged) inference engines such as an OgaEngine-style `PagedAttention` model.
+
+Source:
+[bert_defs.cc](../../onnxruntime/core/graph/contrib_ops/bert_defs.cc) (schema),
+[sparse_attention_indexer_common.h](../../onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h)
+(policy enum, selected-capacity formula and CSA window-plan arithmetic, shared unmodified with
+`SparseAttentionIndexer`),
+[packed_sparse_attention_indexer_common.h](../../onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h)
+(fixed 15-input / 6-output slot map),
+[packed_sparse_attention_indexer.cc](../../onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc) /
+[packed_sparse_attention_indexer_impl.cu](../../onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu)
+(CUDA),
+[packed_sparse_attention_indexer.cc](../../onnxruntime/contrib_ops/webgpu/bert/packed_sparse_attention_indexer.cc)
+(WebGPU, see also [the WebGPU note](webgpu/packed_sparse_attention_indexer.md)).
+
+---
+
+## 1. Why a separate operator
+
+`SparseAttentionIndexer` uses dense `[batch_size, sequence_length, ...]` tensors, an explicit dense
+QSA visibility mask, and state that grows by concatenation every call
+(`present_key = concat(past_key, key)`, etc.). A continuous-batching / paged engine instead:
+
+- flattens every request's tokens into one `[total_tokens, ...]` axis (packed layout);
+- schedules a different number of new tokens per request per step;
+- keeps every request's KV/indexer state in a **fixed-address, fixed-capacity** slot of a state
+ pool (so the engine can reuse buffers and support CUDA graph capture), never a tensor that grows;
+- derives causal visibility purely from packed offsets, never from a materialized `[B, S, T]` mask.
+
+Restructuring `SparseAttentionIndexer` in place to support both contracts was judged more invasive
+and risky than the value of the shared line count (roughly 25-35% of the dense implementation is
+directly reusable without change; most of the rest needs new shapes, new state semantics, or a
+different launch/index mapping). `PackedSparseAttentionIndexer` is therefore a new op, version 1,
+that **does not change `SparseAttentionIndexer`'s schema or behavior**. See [§8](#8-what-is-shared-vs-packed-specific).
+
+## 2. Operator schema
+
+Attributes:
+
+| Attribute | Constraint | Meaning |
+|---|---|---|
+| `policy_mode` | required, `"qsa"` or `"csa"` | selects the indexer flavour |
+| `compress_ratio` | required, `> 0` | tokens folded into one block/compressed entry |
+| `state_capacity` | required, `> 0` | fixed capacity (entries) of `past_key_state` |
+| `token_budget` | `qsa` only, `> 0`, divisible by `compress_ratio` | selected-token budget |
+| `index_topk` | `csa` only, `> 0` | selected compressed-entry count |
+| `epsilon` | default `1e-6` | RMSNorm epsilon |
+| `scale` | default `1/sqrt(head_size)` | per-head score scale |
+| `head_weight_scale` | `csa` only, default `1/sqrt(num_heads)` | head-weight score scale |
+
+Inputs are **fixed at 15 indices** for both policies (unlike the dense op, which uses a different
+input/output count per policy). A slot not owned by the active policy is a *positional* optional:
+its `NodeProto` input name is empty rather than the slot being removed from the list, so every
+later slot keeps its fixed index.
+
+| # | Name | Shape | Type | Policy |
+|---|---|---|---|---|
+| 0 | `query` | `(total_tokens, num_heads, head_size)` | T | both |
+| 1 | `key` | `(total_tokens, head_size)` qsa / `(total_tokens, 2*head_size)` csa | T | both |
+| 2 | `key_norm_weight` | `(head_size)` | T | both |
+| 3 | `cos_cache` | `(max_position, rotary_width)` or `(batch_size, max_position, rotary_width)` | T | both |
+| 4 | `sin_cache` | same shape as `cos_cache` | T | both |
+| 5 | `cumulative_sequence_lengths` | `(batch_size + 1)` | int32, device-resident | both |
+| 6 | `past_sequence_lengths` | `(batch_size)` | int32, device-resident | both |
+| 7 | `gate` | `(total_tokens, 2*head_size)` | T | csa only |
+| 8 | `position_bias` | `(compress_ratio, 2*head_size)` | T | csa only |
+| 9 | `head_weights` | `(total_tokens, num_heads)` | T | csa only |
+| 10 | `position_ids` | `(total_tokens)` | int64 | optional qsa / required csa |
+| 11 | `past_key_state` | `(batch_size, state_capacity, head_size)` | T | both (generic) |
+| 12 | `past_kv_buffer` | `(batch_size, 2*compress_ratio-1, width)` | T | both (generic) |
+| 13 | `past_gate_buffer` | same shape as `past_kv_buffer` | T | csa only |
+| 14 | `past_state_lengths` | `(batch_size, 2)` | int32, device-resident | both (generic) |
+
+Outputs are **fixed at 6 indices** for both policies (`present_gate_buffer` is declared with an
+empty output name for `qsa`, the same positional-optional convention as above):
+
+| # | Name | Shape | Type | Policy |
+|---|---|---|---|---|
+| 0 | `selected_indices` | `(total_tokens, capacity)` | int32 | both |
+| 1 | `selected_counts` | `(total_tokens)` | int32 | both |
+| 2 | `present_key_state` | same shape as `past_key_state` | T | both |
+| 3 | `present_kv_buffer` | same shape as `past_kv_buffer` | T | both |
+| 4 | `present_gate_buffer` | same shape as `past_gate_buffer` | T | csa only |
+| 5 | `present_state_lengths` | same shape as `past_state_lengths` | int32 | both |
+
+`capacity` is `token_budget + compress_ratio - 1` for `qsa` and `index_topk` for `csa`, exactly the
+same formula (`SelectedCapacity`) used by `SparseAttentionIndexer`.
+
+**No output shape depends on tensor data.** `total_tokens` and `batch_size` come from input
+*shapes* (`query.shape[0]`, `cumulative_sequence_lengths.shape[0] - 1`); every state output has
+exactly the same shape as its corresponding state input. This is what makes the fixed-capacity
+state design load-bearing: a growing/concatenated state (as in the dense op) would require a
+data-dependent output shape, which is incompatible with CUDA graph capture and with pre-allocated
+paged state pools.
+
+## 3. Generic state, shared by both policies
+
+Both policies read and write the *same four* state slots — there is no separate
+`past_compressed_key` vs. `past_key` naming split as in the dense op:
+
+- `past_key_state` / `present_key_state`: `qsa` stores prepared (already mean-pooled, RMSNorm'd and
+ rotated) complete-block keys; `csa` stores compressed keys. Layout-compatible with, or cheaply
+ reshaped to, a `[batch_size, capacity, 1, head_size]` auxiliary paged cache when K = V.
+- `past_kv_buffer` / `present_kv_buffer` (and `past_gate_buffer` / `present_gate_buffer`, `csa`
+ only): the generic pending-token buffer, fixed capacity `2 * compress_ratio - 1`. `qsa` only ever
+ uses up to `compress_ratio - 1` of these entries (a raw, not-yet-pooled block); `csa` uses the
+ full range for the overlap ("Ca") plus leftover ("Cb") halves of the window-plan arithmetic
+ reused from `SparseAttentionIndexer`.
+- `past_state_lengths` / `present_state_lengths`: `(batch_size, 2)`. Column 0 is the `key_state`
+ entry count (`qsa`: complete-block count; `csa`: compressed-entry count); column 1 is the pending
+ buffer length (`qsa`: incomplete-block length in `[0, compress_ratio)`; `csa`: buffer length in
+ `[0, 2 * compress_ratio)`, exactly the invariant already documented for
+ `CsaWindowPlan`/`TryComputeCsaWindowPlan`).
+
+State never grows. `present_*` always has exactly the same shape as `past_*`; only the *contents*
+change. Input/output aliasing is supported: every kernel reads its sources (`past_kv_buffer` /
+`key`, `past_gate_buffer` / `gate`) and never re-reads `present_*`, so it is correct whether
+`present_*` is a distinct allocation or the same underlying buffer as `past_*`.
+
+**State overflow.** If a call would close more blocks/windows than
+`state_capacity - old_entry_count` allows, that request's step is rejected as a deterministic
+no-op: its state and state lengths remain unchanged, and its selection outputs stay empty. This
+never reads or writes outside a tensor's fixed extent and never silently truncates state.
+
+## 4. Packed metadata and device-side safety
+
+`cumulative_sequence_lengths` and `past_sequence_lengths` are **device-resident** tensors, read
+directly by the kernels — never copied to the host or synchronized on. The device-visible
+invariants (validated by well-behaved callers; a malformed value never causes memory corruption,
+see below) are:
+
+- `cumulative_sequence_lengths[0] == 0`;
+- `cumulative_sequence_lengths[batch_size] == total_tokens`;
+- `cumulative_sequence_lengths` is nondecreasing (a repeated offset — a zero-token request row —
+ is valid and simply contributes no query rows for that request);
+- `past_sequence_lengths[b] >= 0` and, for `qsa`, consistent with `past_state_lengths[b]`
+ (`key_state_length == past_sequence_length / compress_ratio`,
+ `buffer_length == past_sequence_length % compress_ratio`);
+- `past_state_lengths[b, 0] <= state_capacity` and `past_state_lengths[b, 1]` within its policy's
+ valid buffer range.
+
+Because there is no host synchronization, the kernels cannot literally raise a C++ exception when
+one of these invariants is violated by the input data (as opposed to a mismatched tensor *shape*,
+which the host-side `OpKernel::Compute` and the ONNX schema still check the ordinary way). Instead,
+every per-request quantity read from these tensors is **clamped into its valid range before use**
+(`old_key_len = clamp(past_state_lengths[b,0], 0, state_capacity)`, etc.), and the request-token
+lookup (`PackedBatchOfToken`, a binary search over `cumulative_sequence_lengths`) always returns an
+index in `[0, batch_size)`. The result is that malformed metadata can make the numeric result wrong
+for the affected request, but it can never read or write outside a tensor's allocated extent and
+never causes overlapping writes between requests. This mirrors the "prefer deterministic safe
+outputs" guidance for EPs that cannot report device-side validation errors asynchronously.
+
+## 5. Policy `qsa`
+
+Each request's packed token range is processed independently, in the same three stages as the
+dense `qsa` policy but against fixed-capacity state instead of a growing cache:
+
+1. **Update** (one launch per request): append each raw indexer key to the generic pending buffer;
+ whenever it reaches `compress_ratio` tokens, mean-pool it, apply RMSNorm and `key_norm_weight`,
+ apply the leading/split-half rotary convention at the block's first logical token position
+ (`entry * compress_ratio`, where `entry` is the block's absolute index in `key_state`), and
+ append the **prepared** (already normalized and rotated) key to `key_state`. This differs from
+ the dense kernel, which stores *raw* concatenated keys and repeats the pooling/RMSNorm/rotate
+ work for every query; storing the prepared key once, at update time, is possible only because
+ packed `key_state` never needs to be re-windowed the way a dense-mask query can.
+2. **Score** (one launch per query token, per candidate block): every causally visible block
+ (`block index < min(key_state_length, causal_threshold(position))`) is scored directly against
+ `key_state` with `sum_h ReLU(q_h . k)` — a single dot product, no recomputation.
+3. **Select** (one launch per query token): keeps the `token_budget / compress_ratio` highest
+ scoring blocks (ties broken by ascending index) and appends the request-local logical token
+ positions `[j * compress_ratio, ..., j * compress_ratio + compress_ratio - 1]` for each selected
+ block `j`, followed by every causally visible position of the current incomplete block, and
+ writes the exact active count to `selected_counts`.
+
+"Request-local logical token position" means the same numbering as
+`past_sequence_lengths[b] + local_offset` — i.e. the request's own absolute token position, which
+is exactly what a per-request main paged KV cache is addressed by. The output is therefore directly
+consumable by `SparsePagedAttention` configured with `attention_mode="selected_only"`,
+`selected_kv_source="main"`.
+
+## 6. Policy `csa`
+
+Reuses `SparseAttentionIndexer`'s CSA compression, rotary, scoring, causal threshold, and
+deterministic TopK semantics verbatim (the *math* is unchanged); only the layout, state, and launch
+mapping are packed:
+
+1. **Update** (one launch per request): computes the window plan
+ (`overlap_length`, `leftover_length`, `new_window_count`, `present_buffer_length`,
+ `present_buffer_start`) for that request's own `buffer_length` and packed token count using
+ `TryComputeCsaWindowPlan` — the exact same function used by `SparseAttentionIndexer`'s schema
+ and kernel, called directly on the device (it is a small `SAI_HOST_DEVICE` inline function with
+ no CUDA-specific code). Every closed window is compressed with the softmax-gated Ca/Cb pooling,
+ normalized, rotated with the trailing convention, and appended to `key_state`; the request
+ rejects (caps) new windows beyond `state_capacity` as described in [§3](#3-generic-state-shared-by-both-policies).
+2. **Score** (one launch per query token, per compressed entry): every entry is scored with
+ `sum_h w_h * ReLU(q_h . k)` and masked by the causal threshold from `position_ids` (required for
+ `csa`, unlike `qsa` where it is an optional override of the default `past_sequence_length +
+ local offset`).
+3. **Select** (one launch per query token): keeps the `index_topk` highest scoring, causally
+ visible entries and writes the exact active count to `selected_counts`.
+
+`selected_indices` values are compressed-entry indices into `key_state`, consumable by
+`SparsePagedAttention` configured with `attention_mode="local_plus_selected"`,
+`selected_kv_source="auxiliary"`; `key_state` is layout-compatible with, or cheaply reshaped to, the
+auxiliary cache contract (`[batch_size, capacity, 1, head_size]` when K = V).
+
+## 7. Provider support
+
+CUDA and WebGPU both implement version 1 of this operator, for `float32`, `float16` (CUDA also
+`bfloat16`). There is intentionally no CPU kernel (only the shared constants/helpers/schema are
+CPU-agnostic); a production model that uses this op targets a paged-KV engine on an accelerator.
+See [the WebGPU note](webgpu/packed_sparse_attention_indexer.md) for WebGPU-specific details.
+
+## 8. What is shared vs. packed-specific
+
+Shared with `SparseAttentionIndexer`, unmodified:
+
+- `Policy` enum, `TryParsePolicy`, `SelectedCapacity` (selected-capacity formula) — from
+ `sparse_attention_indexer_common.h`;
+- `CsaWindowPlan` / `TryComputeCsaWindowPlan` (CSA window-plan arithmetic) — same header, now
+ additionally annotated `SAI_HOST_DEVICE` so CUDA device code can call it directly;
+- the CUDA device math (`sparse_attention_indexer_device_math.cuh`, newly extracted from
+ `sparse_attention_indexer_impl.cu` with no behavior change): FP32 block reductions
+ (`SaiBlockSum`), deterministic argmax/selection (`SaiBlockArgMax`, `SaiScanForNext`), the leading
+ and trailing RoPE conventions (`SaiLeadingRope`, `SaiTrailingRope`), and the causal-threshold
+ formula (`SaiCausalThreshold`, which turns out to be exactly the "number of complete blocks fully
+ visible to a query" formula needed by both policies here, unifying what the dense implementation
+ computed two different ways).
+
+Deliberately **not** shared (packed-specific mechanics with no dense equivalent, or dense-only
+mechanics with no packed equivalent):
+
+- packed metadata validation and the device-side per-token/per-request lookup
+ (`PackedBatchOfToken`);
+- fully in-place fixed-capacity state update (no growing/concatenating state, no host-visible
+ data-dependent output shape);
+- plain causal visibility derived from packed offsets (no dense `[B, 1, S, T]` mask input, no
+ mask-compaction step);
+- the dense op's batch-major `[B, S, ...]` launch/index mapping and host wrappers, which do not
+ apply to a token-major `[total_tokens, ...]` tensor.
+
+## 9. Testing
+
+`onnxruntime/test/contrib_ops/packed_sparse_attention_indexer_op_test.cc` covers:
+
+- shape inference for `qsa` and `csa` (fixed `selected_indices`/`selected_counts` shapes, fixed
+ state output shapes, the strict per-slot policy validation, and the always-6-outputs contract);
+- multi-request packed batches with unequal token counts, a zero-token request row, and prefill
+ followed by decode with independent per-request state (CUDA/WebGPU, skipped without the
+ respective execution provider);
+- `qsa` state-capacity overflow safety;
+- FP32/FP16 (and CUDA-only BF16) numeric coverage against an in-file reference that mirrors this
+ document's contract.
+
+## 10. Known limitations and follow-ups
+
+- The reference CUDA/WebGPU kernels prioritize correctness over throughput (see the top-of-file
+ comments in the `.cu`/`.cc` implementations); they are not yet tuned for large `state_capacity`
+ or long packed batches.
+- OgaEngine / Model Builder integration (declaring `past_key_state` etc. as Engine-managed,
+ per-request fixed-size state, analogous to a paged auxiliary cache) is out of scope for this
+ operator definition and is expected in a follow-up to `microsoft/onnxruntime-genai`.
+- No CPU kernel is provided.
diff --git a/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md b/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md
new file mode 100644
index 0000000000000..c1ce8a1549358
--- /dev/null
+++ b/docs/contrib_ops/webgpu/packed_sparse_attention_indexer.md
@@ -0,0 +1,48 @@
+# PackedSparseAttentionIndexer on WebGPU
+
+The WebGPU execution provider implements version 1 of
+`com.microsoft.PackedSparseAttentionIndexer` for the `qsa` and `csa` policies. It uses the
+provider-neutral schema and generic state ABI described in the
+[operator documentation](../packed_sparse_attention_indexer.md).
+
+## Supported subset
+
+- packed (`total_tokens`-major) inputs, driven by device-resident
+ `cumulative_sequence_lengths` / `past_sequence_lengths`;
+- `qsa` and `csa` policy modes;
+- `float32` and `float16`;
+- shared cos/sin rotary cache (`(max_position, rotary_width)`) or request-specific cache
+ (`(batch_size, max_position, rotary_width)`);
+- generic fixed-capacity `past_key_state` / `past_kv_buffer` / `past_gate_buffer` /
+ `past_state_lengths` state, including input/output aliasing;
+- deterministic score-descending, index-ascending top-k ties;
+- `selected_counts`, the exact active-entry count per query.
+
+BF16 is not registered by the WebGPU kernel (CUDA only). Unknown policies and
+policy-incompatible inputs or attributes are rejected.
+
+## Execution
+
+Every program is one invocation per row (one active thread per workgroup; the rest of the
+workgroup is idle), exactly like the dense `SparseAttentionIndexer` WebGPU kernel: state-update
+programs dispatch one row per **request**, and the score/select programs dispatch one row per
+**query token**. As in the dense kernel, intermediate values (pooled/normalized/rotated keys,
+per-candidate scores) are recomputed by small WGSL helper functions on demand rather than staged
+into workgroup-shared arrays, both to keep every kernel correct without relying on WGSL arrays
+sized by a runtime (uniform) `head_size`, and to keep the packed kernel's structure directly
+comparable to the CUDA implementation's per-request/per-token update and score/select stages. All
+reductions and softmax calculations accumulate in FP32, including for FP16 inputs.
+
+Per-request quantities (`cumulative_sequence_lengths`, `past_sequence_lengths`,
+`past_state_lengths`) are read directly from device buffers inside the shaders — never on the
+host — and are clamped into their valid ranges before use, so malformed packed metadata can never
+cause an out-of-bounds buffer access (see the main document's device-side safety section).
+
+## Follow-up work
+
+- specialized large-candidate top-k;
+- subgroup-optimized reductions;
+- fused projection, pooling, and scoring;
+- reduced recomputation and temporary-buffer use;
+- BF16 support;
+- WebGPU `SparsePagedAttention` end-to-end integration.
diff --git a/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h b/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h
new file mode 100644
index 0000000000000..24b7ce1584071
--- /dev/null
+++ b/onnxruntime/contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h
@@ -0,0 +1,85 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+//
+// Shared constants for com.microsoft.PackedSparseAttentionIndexer. This op reuses the policy
+// enum, selected-capacity formula and CSA window-plan arithmetic already defined for
+// com.microsoft.SparseAttentionIndexer in sparse_attention_indexer_common.h; it does not modify
+// that header's input/output slot layout, which stays specific to the dense operator.
+//
+// PackedSparseAttentionIndexer instead uses packed [total_tokens, ...] query/key tensors,
+// device-resident cumulative_sequence_lengths / past_sequence_lengths, and generic fixed-capacity
+// state slots that are shared by both policy_mode values (unlike the dense op's policy-specific
+// state names).
+
+#pragma once
+
+#include
+
+#include "contrib_ops/cpu/sparse/sparse_attention_indexer_common.h"
+
+namespace onnxruntime {
+namespace contrib {
+namespace packed_sparse_attention_indexer {
+
+// Re-exported so callers only need to include this header.
+using sparse_attention_indexer::CsaWindowPlan;
+using sparse_attention_indexer::kPolicyModeCsa;
+using sparse_attention_indexer::kPolicyModeQsa;
+using sparse_attention_indexer::Policy;
+using sparse_attention_indexer::SelectedCapacity;
+using sparse_attention_indexer::TryComputeCsaWindowPlan;
+using sparse_attention_indexer::TryParsePolicy;
+
+// Fixed input slots. Slots 7-9 belong to policy_mode="csa" only; slot 10 (position_ids) is
+// optional for "qsa" and required for "csa". Every other slot is required for both policies.
+enum InputIndex : int {
+ kQuery = 0, // [total_tokens, num_heads, head_size]
+ kKey = 1, // qsa: [total_tokens, head_size]; csa: [total_tokens, 2 * head_size]
+ kKeyNormWeight = 2, // [head_size]
+ kCosCache = 3, // [max_position, rotary_width] or [batch_size, max_position, rotary_width]
+ kSinCache = 4, // same shape as cos_cache
+ kCumulativeSequenceLengths = 5, // [batch_size + 1], int32
+ kPastSequenceLengths = 6, // [batch_size], int32
+ kGate = 7, // csa only: [total_tokens, 2 * head_size]
+ kPositionBias = 8, // csa only: [compress_ratio, 2 * head_size]
+ kHeadWeights = 9, // csa only: [total_tokens, num_heads]
+ kPositionIds = 10, // optional (qsa) / required (csa): [total_tokens], int64
+ kPastKeyState = 11, // generic: [batch_size, state_capacity, head_size]
+ kPastKvBuffer = 12, // generic: [batch_size, 2 * compress_ratio - 1, width]
+ kPastGateBuffer = 13, // csa only: same shape as past_kv_buffer
+ kPastStateLengths = 14, // generic: [batch_size, 2], int32
+ kInputCount = 15,
+};
+
+// Fixed output slots. present_gate_buffer is declared (with an empty name) but not produced for
+// policy_mode="qsa".
+enum OutputIndex : int {
+ kSelectedIndices = 0, // [total_tokens, selected_capacity], int32, unused entries -1
+ kSelectedCounts = 1, // [total_tokens], int32
+ kPresentKeyState = 2, // same shape as past_key_state
+ kPresentKvBuffer = 3, // same shape as past_kv_buffer
+ kPresentGateBuffer = 4, // csa only: same shape as past_gate_buffer
+ kPresentStateLengths = 5, // [batch_size, 2], int32
+ kOutputCount = 6,
+};
+
+// Every PackedSparseAttentionIndexer node declares all 6 fixed outputs; present_gate_buffer is an
+// empty-name optional output for policy_mode="qsa".
+constexpr int kFixedOutputCount = kOutputCount;
+
+// Column layout of past_state_lengths / present_state_lengths.
+enum StateLengthColumn : int {
+ kKeyStateLength = 0, // qsa: complete-block count; csa: compressed-entry count
+ kBufferLength = 1, // qsa: incomplete-block length in [0, compress_ratio);
+ // csa: pending buffer length in [0, 2 * compress_ratio)
+ kStateLengthColumns = 2,
+};
+
+// Generic pending-buffer capacity: qsa only ever uses up to compress_ratio - 1 of these slots.
+SAI_HOST_DEVICE inline int64_t GenericBufferCapacity(int64_t compress_ratio) {
+ return 2 * compress_ratio - 1;
+}
+
+} // namespace packed_sparse_attention_indexer
+} // namespace contrib
+} // namespace onnxruntime
diff --git a/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h b/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h
index d9eb3d1671cca..341315716a2ef 100644
--- a/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h
+++ b/onnxruntime/contrib_ops/cpu/sparse/sparse_attention_indexer_common.h
@@ -6,6 +6,16 @@
#include
#include
+// nvcc recognizes __host__/__device__ as built-in qualifiers in any translation unit it compiles
+// (no CUDA header include required), but a plain host compiler does not know these tokens. This
+// header is shared by CPU-only graph/schema code and by CUDA device code, so the annotation is
+// only emitted when nvcc is compiling the translation unit that includes this header.
+#if defined(__CUDACC__)
+#define SAI_HOST_DEVICE __host__ __device__
+#else
+#define SAI_HOST_DEVICE
+#endif
+
namespace onnxruntime {
namespace contrib {
namespace sparse_attention_indexer {
@@ -68,8 +78,8 @@ constexpr int kCsaOutputCount = 5;
// Number of selected entries emitted per query. The capacity only depends on attributes, so it is
// a compile-time constant of the graph rather than a function of the data.
-inline int64_t SelectedCapacity(Policy policy, int64_t token_budget, int64_t index_topk,
- int64_t compress_ratio) {
+SAI_HOST_DEVICE inline int64_t SelectedCapacity(Policy policy, int64_t token_budget, int64_t index_topk,
+ int64_t compress_ratio) {
return policy == Policy::kQsa ? token_budget + compress_ratio - 1 : index_topk;
}
@@ -87,8 +97,8 @@ struct CsaWindowPlan {
int64_t present_buffer_start = 0; // offset of that buffer inside [past buffer | new tokens]
};
-inline bool TryComputeCsaWindowPlan(int64_t past_buffer_length, int64_t sequence_length,
- int64_t compress_ratio, CsaWindowPlan& plan) {
+SAI_HOST_DEVICE inline bool TryComputeCsaWindowPlan(int64_t past_buffer_length, int64_t sequence_length,
+ int64_t compress_ratio, CsaWindowPlan& plan) {
if (compress_ratio <= 0 || sequence_length < 0 || past_buffer_length < 0 ||
past_buffer_length >= 2 * compress_ratio) {
return false;
diff --git a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc
index f9165bb5074e1..901bba840d994 100644
--- a/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc
+++ b/onnxruntime/contrib_ops/cuda/cuda_contrib_kernels.cc
@@ -256,6 +256,9 @@ class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, SparseAttention);
class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, SparseAttentionIndexer);
class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, SparseAttentionIndexer);
class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, SparseAttentionIndexer);
+class CUDA_MS_OP_TYPED_CLASS_NAME(1, float, PackedSparseAttentionIndexer);
+class CUDA_MS_OP_TYPED_CLASS_NAME(1, MLFloat16, PackedSparseAttentionIndexer);
+class CUDA_MS_OP_TYPED_CLASS_NAME(1, BFloat16, PackedSparseAttentionIndexer);
class CUDA_MS_OP_THREE_TYPED_CLASS_NAME(1, uint8_t, float, int32_t, GatherBlockQuantized);
class CUDA_MS_OP_THREE_TYPED_CLASS_NAME(1, uint8_t, MLFloat16, int32_t, GatherBlockQuantized);
@@ -565,6 +568,9 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc
new file mode 100644
index 0000000000000..eb23d13acd8b0
--- /dev/null
+++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.cc
@@ -0,0 +1,437 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+#include "contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h"
+
+#include
+#include
+#include
+#include
+
+#include "contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h"
+#include "core/providers/cuda/cuda_common.h"
+#include "core/providers/cuda/cuda_type_conversion.h"
+
+namespace onnxruntime {
+namespace contrib {
+namespace cuda {
+
+using namespace onnxruntime::cuda;
+namespace psai = onnxruntime::contrib::packed_sparse_attention_indexer;
+
+#define REGISTER_KERNEL_TYPED(T) \
+ ONNX_OPERATOR_TYPED_KERNEL_EX( \
+ PackedSparseAttentionIndexer, \
+ kMSDomain, \
+ 1, \
+ T, \
+ kCudaExecutionProvider, \
+ (*KernelDefBuilder::Create()) \
+ .TypeConstraint("T", DataTypeImpl::GetTensorType()) \
+ .TypeConstraint("I", DataTypeImpl::GetTensorType()) \
+ .TypeConstraint("M", DataTypeImpl::GetTensorType()), \
+ PackedSparseAttentionIndexer);
+
+REGISTER_KERNEL_TYPED(float)
+REGISTER_KERNEL_TYPED(MLFloat16)
+REGISTER_KERNEL_TYPED(BFloat16)
+
+#undef REGISTER_KERNEL_TYPED
+
+namespace {
+
+Status CheckShape(const Tensor* tensor, const char* name, std::initializer_list expected) {
+ ORT_RETURN_IF(tensor == nullptr, "PackedSparseAttentionIndexer: ", name, " is required");
+ const TensorShape expected_shape(expected);
+ ORT_RETURN_IF_NOT(tensor->Shape() == expected_shape, "PackedSparseAttentionIndexer: ", name, " must have shape ",
+ expected_shape.ToString(), ", got ", tensor->Shape().ToString());
+ return Status::OK();
+}
+
+Status CheckIntDimension(const char* name, int64_t value, bool allow_zero = true) {
+ ORT_RETURN_IF(value < (allow_zero ? 0 : 1) || value > std::numeric_limits::max(),
+ "PackedSparseAttentionIndexer: ", name, " must be in ", allow_zero ? "[0, INT_MAX]" : "(0, INT_MAX]",
+ ", got ", value);
+ return Status::OK();
+}
+
+// cos_cache / sin_cache may be shared across the batch ([max_position, rotary_width]) or
+// request-specific ([batch_size, max_position, rotary_width]).
+struct RotaryCacheShape {
+ bool batched;
+ int64_t max_rotary_length;
+ int64_t rotary_width;
+};
+
+Status CheckRotaryCache(const Tensor* cos_cache, const Tensor* sin_cache, int64_t batch_size,
+ RotaryCacheShape& out) {
+ ORT_RETURN_IF(cos_cache == nullptr, "PackedSparseAttentionIndexer: cos_cache is required");
+ const auto& cos_shape = cos_cache->Shape();
+ out.batched = cos_shape.NumDimensions() == 3;
+ ORT_RETURN_IF_NOT(
+ (out.batched && cos_shape[0] == batch_size && cos_shape[1] > 0) ||
+ (cos_shape.NumDimensions() == 2 && cos_shape[0] > 0),
+ "PackedSparseAttentionIndexer: cos_cache must have shape (max_position, rotary_width) or "
+ "(batch_size, max_position, rotary_width), got ",
+ cos_shape.ToString());
+ out.max_rotary_length = out.batched ? cos_shape[1] : cos_shape[0];
+ out.rotary_width = out.batched ? cos_shape[2] : cos_shape[1];
+ ORT_RETURN_IF_ERROR(CheckIntDimension("max_rotary_sequence_length", out.max_rotary_length, false));
+ ORT_RETURN_IF_ERROR(CheckIntDimension("rotary_width", out.rotary_width, false));
+ ORT_RETURN_IF_NOT(sin_cache != nullptr && sin_cache->Shape() == cos_shape,
+ "PackedSparseAttentionIndexer: sin_cache must have the same shape as cos_cache");
+ return Status::OK();
+}
+
+} // namespace
+
+template
+PackedSparseAttentionIndexer::PackedSparseAttentionIndexer(const OpKernelInfo& info) : CudaKernel(info) {
+ std::string policy_mode;
+ ORT_ENFORCE(info.GetAttr("policy_mode", &policy_mode).IsOK(),
+ "PackedSparseAttentionIndexer: policy_mode is required");
+ ORT_ENFORCE(psai::TryParsePolicy(policy_mode, policy_), "PackedSparseAttentionIndexer: policy_mode must be '",
+ psai::kPolicyModeQsa, "' or '", psai::kPolicyModeCsa, "', got '", policy_mode, "'");
+
+ ORT_ENFORCE(info.GetAttr("compress_ratio", &compress_ratio_).IsOK(),
+ "PackedSparseAttentionIndexer: compress_ratio is required");
+ ORT_ENFORCE(compress_ratio_ > 0 && compress_ratio_ <= std::numeric_limits::max(),
+ "PackedSparseAttentionIndexer: compress_ratio must be in (0, INT_MAX], got ", compress_ratio_);
+
+ int64_t state_capacity = 0;
+ ORT_ENFORCE(info.GetAttr("state_capacity", &state_capacity).IsOK(),
+ "PackedSparseAttentionIndexer: state_capacity is required");
+ ORT_ENFORCE(state_capacity > 0 && state_capacity <= std::numeric_limits::max(),
+ "PackedSparseAttentionIndexer: state_capacity must be in (0, INT_MAX], got ", state_capacity);
+
+ const bool has_token_budget = info.GetAttr("token_budget", &token_budget_).IsOK();
+ const bool has_index_topk = info.GetAttr("index_topk", &index_topk_).IsOK();
+ float head_weight_scale = 0.0f;
+ has_head_weight_scale_ = info.GetAttr("head_weight_scale", &head_weight_scale).IsOK();
+
+ if (policy_ == psai::Policy::kQsa) {
+ ORT_ENFORCE(has_token_budget,
+ "PackedSparseAttentionIndexer: token_budget is required when policy_mode is 'qsa'");
+ ORT_ENFORCE(!has_index_topk && !has_head_weight_scale_,
+ "PackedSparseAttentionIndexer: index_topk and head_weight_scale must not be set when policy_mode "
+ "is 'qsa'");
+ ORT_ENFORCE(token_budget_ > 0 && token_budget_ % compress_ratio_ == 0 &&
+ token_budget_ <= std::numeric_limits::max() - compress_ratio_ + 1,
+ "PackedSparseAttentionIndexer: token_budget must be > 0, divisible by compress_ratio, and produce "
+ "a selected capacity no greater than INT_MAX, got token_budget=",
+ token_budget_, " compress_ratio=", compress_ratio_);
+ index_topk_ = 0;
+ } else {
+ ORT_ENFORCE(has_index_topk, "PackedSparseAttentionIndexer: index_topk is required when policy_mode is 'csa'");
+ ORT_ENFORCE(!has_token_budget,
+ "PackedSparseAttentionIndexer: token_budget must not be set when policy_mode is 'csa'");
+ ORT_ENFORCE(index_topk_ > 0 && index_topk_ <= std::numeric_limits::max(),
+ "PackedSparseAttentionIndexer: index_topk must be in (0, INT_MAX], got ", index_topk_);
+ token_budget_ = 0;
+ }
+
+ epsilon_ = info.GetAttrOrDefault("epsilon", 1.0e-6f);
+ ORT_ENFORCE(epsilon_ >= 0.0f, "PackedSparseAttentionIndexer: epsilon must be >= 0, got ", epsilon_);
+ has_scale_ = info.GetAttr("scale", &scale_).IsOK();
+ head_weight_scale_ = head_weight_scale;
+}
+
+template
+Status PackedSparseAttentionIndexer::ComputeInternal(OpKernelContext* context) const {
+ const bool is_qsa = policy_ == psai::Policy::kQsa;
+ constexpr int kCsaOnlyInputs[] = {psai::kGate, psai::kPositionBias, psai::kHeadWeights};
+ for (int index : kCsaOnlyInputs) {
+ const bool provided = index < context->InputCount() && context->Input(index) != nullptr;
+ ORT_RETURN_IF(provided != !is_qsa, "PackedSparseAttentionIndexer: input ", index,
+ provided ? " must be omitted for policy_mode 'qsa'" : " is required for policy_mode 'csa'");
+ }
+ const bool position_ids_provided =
+ psai::kPositionIds < context->InputCount() && context->Input(psai::kPositionIds) != nullptr;
+ ORT_RETURN_IF(!is_qsa && !position_ids_provided,
+ "PackedSparseAttentionIndexer: input ", psai::kPositionIds,
+ " (position_ids) is required for policy_mode 'csa'");
+ const bool gate_buffer_provided =
+ psai::kPastGateBuffer < context->InputCount() && context->Input(psai::kPastGateBuffer) != nullptr;
+ ORT_RETURN_IF(gate_buffer_provided != !is_qsa, "PackedSparseAttentionIndexer: input ", psai::kPastGateBuffer,
+ gate_buffer_provided ? " must be omitted for policy_mode 'qsa'"
+ : " is required for policy_mode 'csa'");
+
+ return is_qsa ? ComputeQsa(context) : ComputeCsa(context);
+}
+
+template
+Status PackedSparseAttentionIndexer::ComputeQsa(OpKernelContext* context) const {
+ using CudaT = typename OrtToCudaType::type;
+
+ const Tensor* query = context->Input(psai::kQuery);
+ const Tensor* key = context->Input(psai::kKey);
+ const Tensor* key_norm_weight = context->Input(psai::kKeyNormWeight);
+ const Tensor* cos_cache = context->Input(psai::kCosCache);
+ const Tensor* sin_cache = context->Input(psai::kSinCache);
+ const Tensor* cumulative_sequence_lengths = context->Input(psai::kCumulativeSequenceLengths);
+ const Tensor* past_sequence_lengths = context->Input(psai::kPastSequenceLengths);
+ const Tensor* position_ids = context->Input(psai::kPositionIds);
+ const Tensor* past_key_state = context->Input(psai::kPastKeyState);
+ const Tensor* past_kv_buffer = context->Input(psai::kPastKvBuffer);
+ const Tensor* past_state_lengths = context->Input(psai::kPastStateLengths);
+
+ ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required");
+ const auto& query_shape = query->Shape();
+ ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3,
+ "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads, head_size), "
+ "got ",
+ query_shape.ToString());
+ const int64_t total_tokens = query_shape[0];
+ const int64_t num_heads = query_shape[1];
+ const int64_t head_size = query_shape[2];
+ ORT_RETURN_IF_ERROR(CheckIntDimension("total_tokens", total_tokens));
+ ORT_RETURN_IF_ERROR(CheckIntDimension("num_heads", num_heads, false));
+ ORT_RETURN_IF_ERROR(CheckIntDimension("head_size", head_size, false));
+
+ ORT_RETURN_IF(cumulative_sequence_lengths == nullptr,
+ "PackedSparseAttentionIndexer: cumulative_sequence_lengths is required");
+ const auto& cu_shape = cumulative_sequence_lengths->Shape();
+ ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1,
+ "PackedSparseAttentionIndexer: cumulative_sequence_lengths must have shape (batch_size + 1), "
+ "got ",
+ cu_shape.ToString());
+ const int64_t batch_size = cu_shape[0] - 1;
+ ORT_RETURN_IF_ERROR(CheckIntDimension("batch_size", batch_size));
+
+ ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size}));
+ ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, head_size}));
+ ORT_RETURN_IF_ERROR(CheckShape(key_norm_weight, "key_norm_weight", {head_size}));
+ if (position_ids != nullptr) {
+ ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens}));
+ }
+
+ RotaryCacheShape rotary;
+ ORT_RETURN_IF_ERROR(CheckRotaryCache(cos_cache, sin_cache, batch_size, rotary));
+ ORT_RETURN_IF_NOT(
+ rotary.rotary_width > 0 && rotary.rotary_width % 2 == 0 && rotary.rotary_width <= head_size,
+ "PackedSparseAttentionIndexer: policy_mode 'qsa' requires an even rotary_width in (0, head_size], got ",
+ rotary.rotary_width);
+
+ ORT_RETURN_IF(past_key_state == nullptr, "PackedSparseAttentionIndexer: past_key_state is required");
+ const auto& key_state_shape = past_key_state->Shape();
+ ORT_RETURN_IF_NOT(key_state_shape.NumDimensions() == 3 && key_state_shape[0] == batch_size &&
+ key_state_shape[2] == head_size,
+ "PackedSparseAttentionIndexer: past_key_state must have shape "
+ "(batch_size, state_capacity, head_size), got ",
+ key_state_shape.ToString());
+ const int64_t state_capacity = key_state_shape[1];
+ ORT_RETURN_IF_ERROR(CheckIntDimension("state_capacity", state_capacity, false));
+
+ const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_);
+ ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, head_size}));
+ ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths",
+ {batch_size, psai::kStateLengthColumns}));
+
+ const int64_t capacity = psai::SelectedCapacity(psai::Policy::kQsa, token_budget_, index_topk_, compress_ratio_);
+
+ Tensor* selected_indices = context->Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity}));
+ Tensor* selected_counts = context->Output(psai::kSelectedCounts, TensorShape({total_tokens}));
+ Tensor* present_key_state = context->Output(psai::kPresentKeyState, key_state_shape);
+ Tensor* present_kv_buffer =
+ context->Output(psai::kPresentKvBuffer, TensorShape({batch_size, buffer_capacity, head_size}));
+ Tensor* present_state_lengths =
+ context->Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns}));
+ ORT_RETURN_IF(selected_indices == nullptr || selected_counts == nullptr || present_key_state == nullptr ||
+ present_kv_buffer == nullptr || present_state_lengths == nullptr,
+ "PackedSparseAttentionIndexer: policy_mode 'qsa' requires selected_indices, selected_counts, "
+ "present_key_state, present_kv_buffer and present_state_lengths outputs");
+
+ PackedSparseAttentionIndexerParams params;
+ params.batch_size = static_cast(batch_size);
+ params.total_tokens = static_cast(total_tokens);
+ params.num_heads = static_cast(num_heads);
+ params.head_size = static_cast(head_size);
+ params.rotary_width = static_cast(rotary.rotary_width);
+ params.max_rotary_length = static_cast(rotary.max_rotary_length);
+ params.cos_cache_batched = rotary.batched;
+ params.compress_ratio = static_cast(compress_ratio_);
+ params.state_capacity = static_cast(state_capacity);
+ params.buffer_capacity = static_cast(buffer_capacity);
+ params.capacity = static_cast(capacity);
+ params.has_position_ids = position_ids != nullptr;
+ params.epsilon = epsilon_;
+ params.scale = has_scale_ ? scale_ : 1.0f / std::sqrt(static_cast(head_size));
+ params.block_topk = static_cast(token_budget_ / compress_ratio_);
+
+ auto float_workspace = GetScratchBuffer(GetQsaPackedWorkspaceFloatCount(params), GetComputeStream(context));
+ auto overflow_flags =
+ GetScratchBuffer(static_cast(std::max(batch_size, 1)), GetComputeStream(context));
+
+ return LaunchQsaPackedSparseAttentionIndexer(
+ Stream(context), params,
+ reinterpret_cast(query->Data()),
+ reinterpret_cast(key->Data()),
+ reinterpret_cast(key_norm_weight->Data()),
+ reinterpret_cast(cos_cache->Data()),
+ reinterpret_cast(sin_cache->Data()),
+ cumulative_sequence_lengths->Data(),
+ past_sequence_lengths->Data(),
+ position_ids != nullptr ? position_ids->Data() : nullptr,
+ reinterpret_cast(past_key_state->Data()),
+ reinterpret_cast(past_kv_buffer->Data()),
+ past_state_lengths->Data(),
+ selected_indices->MutableData(),
+ selected_counts->MutableData(),
+ reinterpret_cast(present_key_state->MutableData()),
+ reinterpret_cast(present_kv_buffer->MutableData()),
+ present_state_lengths->MutableData(),
+ float_workspace.get(),
+ overflow_flags.get());
+}
+
+template
+Status PackedSparseAttentionIndexer::ComputeCsa(OpKernelContext* context) const {
+ using CudaT = typename OrtToCudaType::type;
+
+ const Tensor* query = context->Input(psai::kQuery);
+ const Tensor* key = context->Input(psai::kKey);
+ const Tensor* key_norm_weight = context->Input(psai::kKeyNormWeight);
+ const Tensor* cos_cache = context->Input(psai::kCosCache);
+ const Tensor* sin_cache = context->Input(psai::kSinCache);
+ const Tensor* cumulative_sequence_lengths = context->Input(psai::kCumulativeSequenceLengths);
+ const Tensor* past_sequence_lengths = context->Input(psai::kPastSequenceLengths);
+ const Tensor* gate = context->Input(psai::kGate);
+ const Tensor* position_bias = context->Input(psai::kPositionBias);
+ const Tensor* head_weights = context->Input(psai::kHeadWeights);
+ const Tensor* position_ids = context->Input(psai::kPositionIds);
+ const Tensor* past_key_state = context->Input(psai::kPastKeyState);
+ const Tensor* past_kv_buffer = context->Input(psai::kPastKvBuffer);
+ const Tensor* past_gate_buffer = context->Input(psai::kPastGateBuffer);
+ const Tensor* past_state_lengths = context->Input(psai::kPastStateLengths);
+
+ ORT_RETURN_IF(query == nullptr, "PackedSparseAttentionIndexer: query is required");
+ const auto& query_shape = query->Shape();
+ ORT_RETURN_IF_NOT(query_shape.NumDimensions() == 3,
+ "PackedSparseAttentionIndexer: query must have shape (total_tokens, num_heads, head_size), "
+ "got ",
+ query_shape.ToString());
+ const int64_t total_tokens = query_shape[0];
+ const int64_t num_heads = query_shape[1];
+ const int64_t head_size = query_shape[2];
+ ORT_RETURN_IF_ERROR(CheckIntDimension("total_tokens", total_tokens));
+ ORT_RETURN_IF_ERROR(CheckIntDimension("num_heads", num_heads, false));
+ ORT_RETURN_IF_ERROR(CheckIntDimension("head_size", head_size, false));
+ ORT_RETURN_IF(head_size > std::numeric_limits::max() / 2,
+ "PackedSparseAttentionIndexer: 2 * head_size must be no greater than INT_MAX");
+ const int64_t width = 2 * head_size;
+
+ ORT_RETURN_IF(cumulative_sequence_lengths == nullptr,
+ "PackedSparseAttentionIndexer: cumulative_sequence_lengths is required");
+ const auto& cu_shape = cumulative_sequence_lengths->Shape();
+ ORT_RETURN_IF_NOT(cu_shape.NumDimensions() == 1 && cu_shape[0] >= 1,
+ "PackedSparseAttentionIndexer: cumulative_sequence_lengths must have shape (batch_size + 1), "
+ "got ",
+ cu_shape.ToString());
+ const int64_t batch_size = cu_shape[0] - 1;
+ ORT_RETURN_IF_ERROR(CheckIntDimension("batch_size", batch_size));
+
+ ORT_RETURN_IF_ERROR(CheckShape(past_sequence_lengths, "past_sequence_lengths", {batch_size}));
+ ORT_RETURN_IF_ERROR(CheckShape(key, "key", {total_tokens, width}));
+ ORT_RETURN_IF_ERROR(CheckShape(key_norm_weight, "key_norm_weight", {head_size}));
+ ORT_RETURN_IF_ERROR(CheckShape(gate, "gate", {total_tokens, width}));
+ ORT_RETURN_IF_ERROR(CheckShape(position_bias, "position_bias", {compress_ratio_, width}));
+ ORT_RETURN_IF_ERROR(CheckShape(head_weights, "head_weights", {total_tokens, num_heads}));
+ ORT_RETURN_IF_ERROR(CheckShape(position_ids, "position_ids", {total_tokens}));
+
+ RotaryCacheShape rotary;
+ ORT_RETURN_IF_ERROR(CheckRotaryCache(cos_cache, sin_cache, batch_size, rotary));
+ ORT_RETURN_IF_NOT(rotary.rotary_width > 0 && 2 * rotary.rotary_width <= head_size,
+ "PackedSparseAttentionIndexer: policy_mode 'csa' requires 0 < 2 * rotary_width <= head_size, "
+ "got rotary_width=",
+ rotary.rotary_width, " head_size=", head_size);
+
+ ORT_RETURN_IF(past_key_state == nullptr, "PackedSparseAttentionIndexer: past_key_state is required");
+ const auto& key_state_shape = past_key_state->Shape();
+ ORT_RETURN_IF_NOT(key_state_shape.NumDimensions() == 3 && key_state_shape[0] == batch_size &&
+ key_state_shape[2] == head_size,
+ "PackedSparseAttentionIndexer: past_key_state must have shape "
+ "(batch_size, state_capacity, head_size), got ",
+ key_state_shape.ToString());
+ const int64_t state_capacity = key_state_shape[1];
+ ORT_RETURN_IF_ERROR(CheckIntDimension("state_capacity", state_capacity, false));
+
+ const int64_t buffer_capacity = psai::GenericBufferCapacity(compress_ratio_);
+ ORT_RETURN_IF_ERROR(CheckShape(past_kv_buffer, "past_kv_buffer", {batch_size, buffer_capacity, width}));
+ ORT_RETURN_IF_ERROR(CheckShape(past_gate_buffer, "past_gate_buffer", {batch_size, buffer_capacity, width}));
+ ORT_RETURN_IF_ERROR(CheckShape(past_state_lengths, "past_state_lengths",
+ {batch_size, psai::kStateLengthColumns}));
+
+ const int64_t capacity = psai::SelectedCapacity(psai::Policy::kCsa, token_budget_, index_topk_, compress_ratio_);
+
+ Tensor* selected_indices = context->Output(psai::kSelectedIndices, TensorShape({total_tokens, capacity}));
+ Tensor* selected_counts = context->Output(psai::kSelectedCounts, TensorShape({total_tokens}));
+ Tensor* present_key_state = context->Output(psai::kPresentKeyState, key_state_shape);
+ Tensor* present_kv_buffer =
+ context->Output(psai::kPresentKvBuffer, TensorShape({batch_size, buffer_capacity, width}));
+ Tensor* present_gate_buffer =
+ context->Output(psai::kPresentGateBuffer, TensorShape({batch_size, buffer_capacity, width}));
+ Tensor* present_state_lengths =
+ context->Output(psai::kPresentStateLengths, TensorShape({batch_size, psai::kStateLengthColumns}));
+ ORT_RETURN_IF(selected_indices == nullptr || selected_counts == nullptr || present_key_state == nullptr ||
+ present_kv_buffer == nullptr || present_gate_buffer == nullptr ||
+ present_state_lengths == nullptr,
+ "PackedSparseAttentionIndexer: policy_mode 'csa' requires selected_indices, selected_counts, "
+ "present_key_state, present_kv_buffer, present_gate_buffer and present_state_lengths outputs");
+
+ PackedSparseAttentionIndexerParams params;
+ params.batch_size = static_cast(batch_size);
+ params.total_tokens = static_cast(total_tokens);
+ params.num_heads = static_cast(num_heads);
+ params.head_size = static_cast(head_size);
+ params.rotary_width = static_cast(rotary.rotary_width);
+ params.max_rotary_length = static_cast(rotary.max_rotary_length);
+ params.cos_cache_batched = rotary.batched;
+ params.compress_ratio = static_cast(compress_ratio_);
+ params.state_capacity = static_cast(state_capacity);
+ params.buffer_capacity = static_cast(buffer_capacity);
+ params.capacity = static_cast(capacity);
+ params.has_position_ids = true;
+ params.epsilon = epsilon_;
+ params.scale = has_scale_ ? scale_ : 1.0f / std::sqrt(static_cast(head_size));
+ params.index_topk = static_cast(index_topk_);
+ params.head_weight_scale =
+ has_head_weight_scale_ ? head_weight_scale_ : 1.0f / std::sqrt(static_cast(num_heads));
+
+ auto float_workspace = GetScratchBuffer(GetCsaPackedWorkspaceFloatCount(params), GetComputeStream(context));
+ auto overflow_flags =
+ GetScratchBuffer(static_cast(std::max(batch_size, 1)), GetComputeStream(context));
+
+ return LaunchCsaPackedSparseAttentionIndexer(
+ Stream(context), params,
+ reinterpret_cast(query->Data()),
+ reinterpret_cast(key->Data()),
+ reinterpret_cast(key_norm_weight->Data()),
+ reinterpret_cast(cos_cache->Data()),
+ reinterpret_cast(sin_cache->Data()),
+ reinterpret_cast(gate->Data()),
+ reinterpret_cast(position_bias->Data()),
+ reinterpret_cast(head_weights->Data()),
+ cumulative_sequence_lengths->Data(),
+ past_sequence_lengths->Data(),
+ position_ids->Data(),
+ reinterpret_cast(past_key_state->Data()),
+ reinterpret_cast(past_kv_buffer->Data()),
+ reinterpret_cast(past_gate_buffer->Data()),
+ past_state_lengths->Data(),
+ selected_indices->MutableData(),
+ selected_counts->MutableData(),
+ reinterpret_cast(present_key_state->MutableData()),
+ reinterpret_cast(present_kv_buffer->MutableData()),
+ reinterpret_cast(present_gate_buffer->MutableData()),
+ present_state_lengths->MutableData(),
+ float_workspace.get(),
+ overflow_flags.get());
+}
+
+template class PackedSparseAttentionIndexer;
+template class PackedSparseAttentionIndexer;
+template class PackedSparseAttentionIndexer;
+
+} // namespace cuda
+} // namespace contrib
+} // namespace onnxruntime
diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h
new file mode 100644
index 0000000000000..b7ddfa9b480ce
--- /dev/null
+++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer.h
@@ -0,0 +1,37 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+#pragma once
+
+#include "contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h"
+#include "core/common/common.h"
+#include "core/providers/cuda/cuda_kernel.h"
+
+namespace onnxruntime {
+namespace contrib {
+namespace cuda {
+
+template
+class PackedSparseAttentionIndexer final : public onnxruntime::cuda::CudaKernel {
+ public:
+ explicit PackedSparseAttentionIndexer(const OpKernelInfo& info);
+ Status ComputeInternal(OpKernelContext* context) const override;
+
+ private:
+ Status ComputeQsa(OpKernelContext* context) const;
+ Status ComputeCsa(OpKernelContext* context) const;
+
+ packed_sparse_attention_indexer::Policy policy_;
+ int64_t compress_ratio_;
+ int64_t token_budget_;
+ int64_t index_topk_;
+ float epsilon_;
+ float scale_;
+ float head_weight_scale_;
+ bool has_scale_;
+ bool has_head_weight_scale_;
+};
+
+} // namespace cuda
+} // namespace contrib
+} // namespace onnxruntime
diff --git a/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu
new file mode 100644
index 0000000000000..36f1a6a521b9b
--- /dev/null
+++ b/onnxruntime/contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.cu
@@ -0,0 +1,868 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+//
+// Correctness-first implementation of com.microsoft.PackedSparseAttentionIndexer. It shares its
+// block reductions, deterministic argmax/selection helpers, RoPE math and causal-threshold formula
+// with com.microsoft.SparseAttentionIndexer via sparse_attention_indexer_device_math.cuh, and
+// reuses the CSA window-plan arithmetic directly on the device from
+// sparse_attention_indexer_common.h (its helpers are SAI_HOST_DEVICE). What is packed-specific:
+// token-major (not batch-major) layout, cumulative_sequence_lengths / past_sequence_lengths driven
+// per-request bookkeeping, fully in-place fixed-capacity state update (no growing/concatenating
+// state), and plain causal visibility derived from packed metadata (no dense mask).
+//
+// Device-side safety: every per-request quantity (past_sequence_lengths, past_state_lengths,
+// cumulative offsets) is read directly from device memory inside the kernels below -- there is no
+// host readback or stream synchronization. Values are always clamped into the fixed-capacity range
+// before use, so malformed metadata can make the result semantically wrong but can never cause an
+// out-of-bounds access or an overlapping write. State-capacity overflow is *rejected*, not
+// silently truncated: if a request's new blocks/windows would not all fit in state_capacity this
+// call, the update kernel applies none of them (present_state_lengths / present_key_state /
+// present_kv_buffer / present_gate_buffer for that request are left exactly as their past_*
+// counterparts) and records the rejection in a small overflow_flags workspace; the select kernels
+// then force that request's selected_indices/selected_counts to the deterministic safe empty
+// result (-1 / 0) for this call instead of selecting against a partially updated state. See
+// QsaUpdateStateKernel / CsaUpdateStateKernel / QsaSelectKernel / CsaSelectKernel.
+//
+// See docs/contrib_ops/cuda/packed_sparse_attention_indexer.md for the full operator contract.
+
+#include "contrib_ops/cuda/sparse/packed_sparse_attention_indexer_impl.h"
+
+#include
+#include
+#include
+
+#include
+
+#include "contrib_ops/cpu/sparse/packed_sparse_attention_indexer_common.h"
+#include "contrib_ops/cuda/sparse/sparse_attention_indexer_device_math.cuh"
+#include "core/providers/cuda/cu_inc/cuda_type_helper.cuh"
+
+namespace onnxruntime {
+namespace contrib {
+namespace cuda {
+
+namespace psai = onnxruntime::contrib::packed_sparse_attention_indexer;
+
+namespace {
+
+// The block reductions below halve the active thread count, so this must stay a power of two.
+constexpr int kThreads = 128;
+
+// Largest b such that cumulative_sequence_lengths[b] <= token, assuming the array is nondecreasing.
+// If the data itself is malformed this may attribute a token to the wrong request, but the result
+// is always an index in [0, batch_size), so it can never cause an out-of-bounds access.
+__device__ __forceinline__ int PackedBatchOfToken(const int32_t* cumulative_sequence_lengths, int batch_size,
+ int token) {
+ int lo = 0;
+ int hi = batch_size - 1;
+ while (lo < hi) {
+ const int mid = lo + (hi - lo + 1) / 2;
+ if (cumulative_sequence_lengths[mid] <= token) {
+ lo = mid;
+ } else {
+ hi = mid - 1;
+ }
+ }
+ return lo;
+}
+
+template
+__global__ void ElementwiseCopyKernel(const T* src, T* dst, int64_t count) {
+ for (int64_t i = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; i < count;
+ i += static_cast(gridDim.x) * blockDim.x) {
+ dst[i] = src[i];
+ }
+}
+
+// ---------------------------------------------------------------------------------------------
+// policy_mode = "qsa"
+// ---------------------------------------------------------------------------------------------
+
+// One block per request: forms every newly-closed compress_ratio block (mean-pool -> RMSNorm ->
+// leading RoPE -> append into fixed-capacity key_state) and publishes the raw trailing buffer.
+// Reads only past_kv_buffer / key (never present_kv_buffer / present_key_state), so it is correct
+// whether or not the present/past tensors are the same aliased allocation.
+template
+__global__ void QsaUpdateStateKernel(const T* key, const T* key_norm_weight, const T* cos_cache,
+ const T* sin_cache, const int32_t* cumulative_sequence_lengths,
+ const T* past_kv_buffer, const int32_t* past_state_lengths,
+ T* present_key_state, T* present_kv_buffer,
+ int32_t* present_state_lengths, int32_t* overflow_flags,
+ PackedSparseAttentionIndexerParams params) {
+ extern __shared__ float shared[];
+ float* pooled = shared;
+ float* rotated = shared + params.head_size;
+ float* reduction = shared + 2 * params.head_size;
+
+ for (int b = static_cast(blockIdx.x); b < params.batch_size; b += static_cast(gridDim.x)) {
+ const int req_start = cumulative_sequence_lengths[b];
+ const int req_end = cumulative_sequence_lengths[b + 1];
+ const int req_len = req_end > req_start ? req_end - req_start : 0;
+
+ const int old_key_len = min(max(past_state_lengths[b * 2 + psai::kKeyStateLength], 0), params.state_capacity);
+ const int old_buf_len =
+ min(max(past_state_lengths[b * 2 + psai::kBufferLength], 0), params.compress_ratio - 1);
+
+ const int pending = old_buf_len + req_len;
+ const int full_new_block_count = pending / params.compress_ratio;
+ const int capacity_left = params.state_capacity - old_key_len; // >= 0 by construction of old_key_len
+ // Reject (do not partially apply) a step that would need more than the fixed state_capacity:
+ // no new blocks are formed and the buffer is left exactly as it was, so a rejected step is a
+ // deterministic no-op on state rather than a silent partial truncation.
+ const bool overflowed = full_new_block_count > capacity_left;
+ const int new_block_count = overflowed ? 0 : full_new_block_count;
+ const int new_buf_len = overflowed ? old_buf_len : (pending % params.compress_ratio);
+
+ // Barrier: every thread has now read past_state_lengths (identically) before any thread below
+ // writes present_state_lengths, which keeps this correct even if the two tensors alias.
+ __syncthreads();
+
+ if (threadIdx.x == 0) {
+ present_state_lengths[b * 2 + psai::kKeyStateLength] = old_key_len + new_block_count;
+ present_state_lengths[b * 2 + psai::kBufferLength] = new_buf_len;
+ overflow_flags[b] = overflowed ? 1 : 0;
+ }
+
+ for (int k = 0; k < new_block_count; ++k) {
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ float sum = 0.0f;
+ for (int t = 0; t < params.compress_ratio; ++t) {
+ const int virtual_pos = k * params.compress_ratio + t;
+ sum += virtual_pos < old_buf_len
+ ? to_float(past_kv_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) *
+ params.head_size +
+ d])
+ : to_float(key[(static_cast(req_start) + (virtual_pos - old_buf_len)) *
+ params.head_size +
+ d]);
+ }
+ pooled[d] = sum / static_cast(params.compress_ratio);
+ }
+ __syncthreads();
+
+ float sum_squares = 0.0f;
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ sum_squares += pooled[d] * pooled[d];
+ }
+ sum_squares = SaiBlockSum(sum_squares, reduction);
+ const float inverse_rms = rsqrtf(sum_squares / static_cast(params.head_size) + params.epsilon);
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ pooled[d] = pooled[d] * inverse_rms * to_float(key_norm_weight[d]);
+ }
+ __syncthreads();
+
+ const int entry = old_key_len + k;
+ const int rope_position =
+ SaiClampPosition(static_cast(entry) * params.compress_ratio, params.max_rotary_length);
+ const int64_t cache_offset =
+ (static_cast(params.cos_cache_batched ? b : 0) * params.max_rotary_length + rope_position) *
+ params.rotary_width;
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ rotated[d] = SaiLeadingRope(pooled, params.rotary_width, cos_cache + cache_offset,
+ sin_cache + cache_offset, d);
+ }
+ __syncthreads();
+
+ const int64_t out_base = (static_cast(b) * params.state_capacity + entry) * params.head_size;
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ present_key_state[out_base + d] = from_float(rotated[d]);
+ }
+ __syncthreads();
+ }
+
+ // Publish the raw trailing buffer. Skipped entirely on overflow: present_kv_buffer already
+ // holds past_kv_buffer's contents unchanged (from the baseline copy in the Launch function
+ // below), which is exactly the prior valid buffer this rejected step must preserve.
+ if (!overflowed) {
+ for (int t = static_cast(threadIdx.x); t < new_buf_len; t += static_cast(blockDim.x)) {
+ const int virtual_pos = new_block_count * params.compress_ratio + t;
+ const int64_t out_base = (static_cast(b) * params.buffer_capacity + t) * params.head_size;
+ for (int d = 0; d < params.head_size; ++d) {
+ const float value =
+ virtual_pos < old_buf_len
+ ? to_float(
+ past_kv_buffer[(static_cast(b) * params.buffer_capacity + virtual_pos) *
+ params.head_size +
+ d])
+ : to_float(
+ key[(static_cast(req_start) + (virtual_pos - old_buf_len)) * params.head_size + d]);
+ present_kv_buffer[out_base + d] = from_float(value);
+ }
+ }
+ }
+ __syncthreads();
+ }
+}
+
+// One block per (token, head): rotates the query once so downstream scoring kernels only ever dot
+// two already-rotated/prepared vectors. kUseLeadingRope selects the qsa convention (position
+// defaults to past_sequence_lengths[batch] + request-local offset when position_ids is absent);
+// otherwise the csa trailing convention with positions always taken from position_ids.
+template
+__global__ void PackedRotateQueryKernel(const T* query, const T* cos_cache, const T* sin_cache,
+ const int32_t* cumulative_sequence_lengths,
+ const int32_t* past_sequence_lengths, const int64_t* position_ids,
+ float* query_rotated, PackedSparseAttentionIndexerParams params) {
+ extern __shared__ float shared[];
+ const int64_t rows = static_cast(params.total_tokens) * params.num_heads;
+ for (int64_t row = blockIdx.x; row < rows; row += gridDim.x) {
+ const int token = static_cast(row / params.num_heads);
+ const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token);
+ const int64_t base = row * params.head_size;
+
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ shared[d] = to_float(query[base + d]);
+ }
+ __syncthreads();
+
+ const int64_t abs_position =
+ params.has_position_ids
+ ? position_ids[token]
+ : static_cast(past_sequence_lengths[batch]) + (token - cumulative_sequence_lengths[batch]);
+ const int position = SaiClampPosition(abs_position, params.max_rotary_length);
+ const int64_t cache_offset =
+ (static_cast(params.cos_cache_batched ? batch : 0) * params.max_rotary_length + position) *
+ params.rotary_width;
+ const T* cos_row = cos_cache + cache_offset;
+ const T* sin_row = sin_cache + cache_offset;
+
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ query_rotated[base + d] = kUseLeadingRope
+ ? SaiLeadingRope(shared, params.rotary_width, cos_row, sin_row, d)
+ : SaiTrailingRope(shared, params.head_size, params.rotary_width, cos_row,
+ sin_row, d);
+ }
+ __syncthreads();
+ }
+}
+
+// One block per (token, key_state slot). Scores are directly dotted against the already-prepared
+// present_key_state entry (unlike the dense op, no per-query pooling/normalize/rotate is repeated
+// here because the packed contract stores fully-prepared blocks in key_state).
+template
+__global__ void QsaBlockScoreKernel(const T* present_key_state, const float* query_rotated,
+ const int32_t* cumulative_sequence_lengths,
+ const int32_t* past_sequence_lengths, const int64_t* position_ids,
+ const int32_t* present_state_lengths, float* block_scores,
+ PackedSparseAttentionIndexerParams params) {
+ extern __shared__ float reduction[];
+ const int64_t total = static_cast(params.total_tokens) * params.state_capacity;
+ for (int64_t work = blockIdx.x; work < total; work += gridDim.x) {
+ const int token = static_cast(work / params.state_capacity);
+ const int block_index = static_cast(work % params.state_capacity);
+ const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token);
+ const int key_len_after = present_state_lengths[batch * 2 + psai::kKeyStateLength];
+
+ if (block_index >= key_len_after) {
+ if (threadIdx.x == 0) {
+ block_scores[work] = SaiNegativeInfinity();
+ }
+ continue;
+ }
+
+ const int64_t abs_position =
+ params.has_position_ids
+ ? position_ids[token]
+ : static_cast(past_sequence_lengths[batch]) + (token - cumulative_sequence_lengths[batch]);
+ const int64_t causal_count = SaiCausalThreshold(abs_position, params.compress_ratio);
+ const int64_t visible_block_count = causal_count < key_len_after ? causal_count : key_len_after;
+ if (static_cast(block_index) >= visible_block_count) {
+ if (threadIdx.x == 0) {
+ block_scores[work] = SaiNegativeInfinity();
+ }
+ continue;
+ }
+
+ const int64_t key_base = (static_cast(batch) * params.state_capacity + block_index) * params.head_size;
+ float score = 0.0f;
+ for (int head = 0; head < params.num_heads; ++head) {
+ const float* query_head =
+ query_rotated + (static_cast(token) * params.num_heads + head) * params.head_size;
+ float partial = 0.0f;
+ for (int d = static_cast(threadIdx.x); d < params.head_size; d += static_cast(blockDim.x)) {
+ partial += query_head[d] * to_float(present_key_state[key_base + d]);
+ }
+ score += fmaxf(SaiBlockSum(partial, reduction), 0.0f);
+ }
+ if (threadIdx.x == 0) {
+ block_scores[work] = score * params.scale;
+ }
+ __syncthreads();
+ }
+}
+
+// One block per query token. Emits the token indices of the highest scoring blocks followed by the
+// causally visible tokens of the trailing incomplete block, and the exact active count. A request
+// whose update step was rejected for exceeding state_capacity this call (overflow_flags[batch] set)
+// always gets the safe empty result: indices stay -1 (already reset below) and count is 0.
+__global__ void QsaSelectKernel(const float* block_scores, const int32_t* cumulative_sequence_lengths,
+ const int32_t* past_sequence_lengths, const int64_t* position_ids,
+ const int32_t* present_state_lengths, const int32_t* overflow_flags,
+ int32_t* selected_indices, int32_t* selected_counts,
+ PackedSparseAttentionIndexerParams params) {
+ extern __shared__ float shared[];
+ float* shared_value = shared;
+ int* shared_index = reinterpret_cast(shared + blockDim.x);
+
+ for (int token = static_cast(blockIdx.x); token < params.total_tokens; token += static_cast(gridDim.x)) {
+ int32_t* out_row = selected_indices + static_cast(token) * params.capacity;
+ for (int p = static_cast(threadIdx.x); p < params.capacity; p += static_cast(blockDim.x)) {
+ out_row[p] = -1;
+ }
+ __syncthreads();
+
+ const int batch = PackedBatchOfToken(cumulative_sequence_lengths, params.batch_size, token);
+ if (overflow_flags[batch] != 0) {
+ if (threadIdx.x == 0) {
+ selected_counts[token] = 0;
+ }
+ __syncthreads();
+ continue;
+ }
+
+ const int key_len_after = present_state_lengths[batch * 2 + psai::kKeyStateLength];
+ const int64_t abs_position =
+ params.has_position_ids
+ ? position_ids[token]
+ : static_cast(past_sequence_lengths[batch]) + (token - cumulative_sequence_lengths[batch]);
+ const int64_t causal_count = SaiCausalThreshold(abs_position, params.compress_ratio);
+ const int64_t visible_64 = causal_count < key_len_after ? causal_count : static_cast(key_len_after);
+ const int visible_block_count = static_cast(visible_64 < 0 ? 0 : visible_64);
+ const int selected = params.block_topk < visible_block_count ? params.block_topk : visible_block_count;
+
+ const float* scores_row = block_scores + static_cast(token) * params.state_capacity;
+
+ float previous_score = 0.0f;
+ int previous_index = -1;
+ int emitted_blocks = 0;
+ for (int rank = 0; rank < selected; ++rank) {
+ float best_value = 0.0f;
+ int best_index = -1;
+ SaiScanForNext(scores_row, visible_block_count, previous_score, previous_index, &best_value, &best_index);
+ shared_value[threadIdx.x] = best_value;
+ shared_index[threadIdx.x] = best_index;
+ __syncthreads();
+ SaiBlockArgMax(shared_value, shared_index);
+ previous_index = shared_index[0];
+ previous_score = shared_value[0];
+ __syncthreads();
+ if (previous_index < 0) {
+ break;
+ }
+ for (int t = static_cast(threadIdx.x); t < params.compress_ratio; t += static_cast(blockDim.x)) {
+ out_row[rank * params.compress_ratio + t] = previous_index * params.compress_ratio + t;
+ }
+ emitted_blocks = rank + 1;
+ __syncthreads();
+ }
+
+ // The trailing incomplete block is always causally visible in full up to this query's own
+ // position; only its indices are needed (SparsePagedAttention reads the raw main cache).
+ const int64_t block_start = static_cast(visible_block_count) * params.compress_ratio;
+ const int64_t natural_tail = abs_position >= block_start ? (abs_position - block_start + 1) : 0;
+ const int remaining_capacity = params.capacity - emitted_blocks * params.compress_ratio;
+ const int64_t tail_count_64 = natural_tail < remaining_capacity ? natural_tail : remaining_capacity;
+ const int tail_count = static_cast(tail_count_64 < 0 ? 0 : tail_count_64);
+ for (int t = static_cast(threadIdx.x); t < tail_count; t += static_cast(blockDim.x)) {
+ out_row[emitted_blocks * params.compress_ratio + t] = static_cast(block_start + t);
+ }
+ if (threadIdx.x == 0) {
+ selected_counts[token] = emitted_blocks * params.compress_ratio + tail_count;
+ }
+ __syncthreads();
+ }
+}
+
+// ---------------------------------------------------------------------------------------------
+// policy_mode = "csa"
+// ---------------------------------------------------------------------------------------------
+
+// One block per request: closes every new compression window (softmax-gated pool -> RMSNorm ->
+// trailing RoPE -> append into fixed-capacity key_state) using the shared CsaWindowPlan helper,
+// and publishes the raw overlap+leftover buffer. Reads only past_kv_buffer / past_gate_buffer /
+// key / gate (never the present_* tensors), so it is correct whether or not present/past alias.
+template