…MLA (slice 3A.2) (#1978)
## Slice 3A.2 — native onnx-genai CUDA support for
`com.microsoft::PagedAttention` LATENT (GLM-5.2 dense MLA)
Builds on the merged audit/validator/oracle (#1940) and the
token-major/LATENT KV index emission (#1955). This slice adds the
**native CUDA kernel** for the exact ORT v1
`com.microsoft::PagedAttention` **LATENT** (absorbed-MLA) subset that
GLM-5.2 dense MLA needs, with `onnx-genai-kv` remaining the **sole**
page/cache authority.
**Depends on / stacks after #1955** (already merged into `main`).
### What this adds
- `crates/onnx-runtime-ep-cuda/src/kernels/paged_attention.rs` — NVRTC
f16/bf16 write + attention kernels implemented **from the oracle
equations** (not copied upstream source): partial-RoPE suffix write into
the paged latent cache, online softmax over the latent cache honoring
`local_window_size`/softcap, V taken from the leading `v_head_size`
channels of the same latent row. Plus `PagedAttentionFactory`,
`PagedAttentionLatentKernel`, and `unsupported_reason()`.
- Five-place op registration in `kernels/mod.rs` + the
`unsupported_reason` arm in `provider.rs::supports_op`.
### Invariants held
- **Default-off / typed subset only.** The op is claimed **only** when
the typed geometry+dtype validator proves the exact supported subset
(fp16/bf16 LATENT, single latent KV head, GLM qk=192/v=128/partial RoPE,
block pow2 ≥ 16). Every unsupported optional mode returns a **typed
NotImplemented** reason rather than silently miscomputing: non-LATENT
layout, quantized cache (k/v quant type + int4/float4 cache dtype),
`head_sink`, q/k-norm, k/v scales, present `value`/`value_cache`,
non-f16/bf16.
- **One-authority.** No op-side allocation and no second KV manager —
the kernel consumes/mutates the caller's page buffers in place.
**In-place `key_cache_out` alias contract enforced** (non-aliased cache
output is rejected).
- **Capture/replay safe.** Warmed kernel signature; no host
sync/allocation during capture. Proven by a capture + 3×replay test that
is bit-equal to eager.
- **No Mobius changes / no export claims** in this slice. No full-size
performance claim.
### Tests (native CUDA vs `onnx-genai-paged-attention` oracle, verified
on an idle A100)
- Parity: tiny prefill/decode; GLM dims (qk=192/v=128, partial RoPE)
first-token/prefill/decode; block pow2≥16 boundary + multi-request; slot
`-1` skip; no-rotary; eager equivalence (fp16 err ~2e-4, bf16 ~2e-3 —
well under tol).
- Capture + 3 replays bit-equal to eager (`check_capture_error()==0`).
- Rejections: non-aliased `key_cache_out` and missing required input
(`block_table`).
- Measurement (tiny-shape correctness gate, **not** a perf claim):
CUDA-event timing, n=5, prefill + decode, page/VRAM accounting, op-side
alloc = 0. Prefill med ~0.050 ms, decode med ~0.073 ms on tiny GLM
shapes.
- 5 pure unit tests cover the typed rejections without a GPU.
- Existing `onnx-runtime-ep-cuda` lib suite green (559 passed) — no
regressions; KV geometry validation (#1955 + Gaff's sibling constraints)
unchanged.
### Reviewer / gates
- **Draft** — for independent review by **Gaff or Roy** (final approval
required; reviewer excludes Leon/Sapper). Do **not** merge without
explicit approval.
- Remaining gates: full-size GLM checkpoint runs (correctness +
measurement) coordinated with the ongoing GLM GGUF/safetensors work; 3B
Mobius opt-in `--paged-attention` export only starts **after** 3A
approval.
_GPU tests are gated behind the `gpu-tests` cargo feature and the CUDA
env; run on A100 with `--features gpu-tests`._
---------
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Slice 3A.1 — KV-authority half of native PagedAttention integration
First independently-reviewable increment of slice 3A. Follows #1940 (merged
corrected audit + typed validator/oracle crate
onnx-genai-paged-attention).Slice 3A as briefed (native CUDA
com.microsoft::PagedAttentionLATENT kernel +KV APIs + full parity/capture matrix + A100 measurement) is larger than one
safe, verifiable increment. Per the "land in order / no dead code / each slice
independently reviewed" directive, this PR lands Requirement 1 only: additive
token-major / LATENT paged-index emission in
onnx-genai-kv. It is a strictprerequisite for the CUDA kernel and is fully validated on host (no GPU).
What it adds (strictly additive, read-only)
crates/onnx-genai-kv/src/paged_index.rs:PagedIndexPlan::build(&PageTable, &[PagedRequest])emits the exact ORT v1int32 index tensors from the existing page authority:
block_table[num_seqs, max_num_blocks_per_seq]— physicalPageIdperlogical block, padded with
PAGED_BLOCK_TABLE_PAD;slot_mapping[token_count]—page_id*block_size + offsetper query token;PAGED_SLOT_EMPTY == -1documented skip sentinel;cumulative_sequence_length[num_seqs+1](cu_seqlens_q),past_seqlens,derived
context_lens.LatentCacheGeometry+validate(), and canonicaltoken_major_element_offset/
latent_element_offsetso a CUDA kernel and the CPU oracle index the cachethrough one formula.
block_size= power-of-two ≥ 16 (matchescheck_kv_cache).PagedKvCache::emit_paged_index_planconvenience.One-authority invariant (enforced in code)
onnx-genai-kvstays the sole owner of page allocation/lifetime.paged_indexis a read-only view — it allocates/frees/mutates nothing, no second manager, no
op-side allocation. The physical block id is the
PageId(
KvViewKind::VirtuallyContiguous); a kernel binds these host-emitted indices asstable device inputs and updates caller-owned cache tensors in place.
Byte-identity
No change to
Page/PageTablestorage; existing head-major layout untouched. Aread-only/leak-free test asserts
materialize_sequenceand poolusage()/stats()are identical before and after emission.Typed rejections (never silent miscompute) — all tested
non-pow2/
<16block size, windowed/attention-sink (non-contiguous) sequences,query>context, missing backing pages, i32 block-id/slot overflow.
Tests
3 unit + 12 integration (prefill/decode slot math, exact pow2 boundaries
16/32/64, multi-request row-major+padding, page reuse after free, read-only/
leak-free, every typed rejection). Full crate suite (147+…) and
cargo clippy -p onnx-genai-kv --testsclean."No dead code" note
First production consumer is the native CUDA kernel (slice 3A.2). This mirrors
the crate's existing
kv_capacity_bucket/ensure_kv_capacity/KvCapacityGrowthBackendauthority seam (exercised by tests, consumed bybackends). Its consumer lands in the next, independently-reviewed slice.
Remaining gates
quantized cache → typed NotImplemented; then A100 measurement (n≥3).
--paged-attentionexport, only after 3A approval.Reviewer: independent, excluding Leon (author). Final approval: Gaff or Roy.
Do not merge without explicit approval. Draft.