Skip to content

fix(engine): bind fixed CUDA decoder state by shape - #436

Merged
justinchuby merged 1 commit into
mainfrom
squad/384-conv-state-rank
Jul 30, 2026
Merged

justinchuby merged 1 commit into
mainfrom
squad/384-conv-state-rank

Conversation

@justinchuby

Copy link
Copy Markdown
Owner

Summary

  • distinguish metadata-declared fixed state_pairs from growable KV pairs in native CUDA decode
  • allocate fixed recurrent state at its declared rank and static geometry, zero-initialize it, and keep it outside KV capacity growth
  • preserve rank-4 growable KV physical/logical shapes and bytes-per-token accounting
  • add pure shape-contract tests and a GPU smoke that binds and executes a rank-3 FP16 fixed state alongside rank-4 KV

Root cause

DecodeCudaState treated every past/present pair as a rank-4 KV cache. It forced axis 2 to the KV capacity bucket and later changed axis 2 on every decode step and capacity growth. Metadata state_pairs already identify fixed replace-semantics recurrent tensors, but that distinction was discarded before CUDA binding allocation.

Real-model evidence

Qwen3.6-27B INT4 previously failed while allocating past_key_values.12.conv_state as rank-3 FP16. With this change it clears CUDA state allocation and begins the forward pass. The next blocker is the CUDA Conv kernel declining rank-3 1-D convolution at node __fn0_Conv_node_12; the old line-480 state allocation error is gone.

Validation

  • native_decode::tests::cuda_persistent_state_shapes_preserve_growing_kv_and_fixed_recurrent_geometry
  • native_decode::tests::cuda_fixed_state_shapes_reject_unbounded_non_batch_dimensions
  • native_decode::tests::native_cuda_binds_rank3_fixed_state_without_changing_growing_kv on GPU4
  • cargo fmt --all -- --check
  • cargo clippy -p onnx-genai-engine --all-targets --features native-backend,cuda -- -D warnings

References #384

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@codecov

codecov Bot commented Jul 30, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 80.58%. Comparing base (632a016) to head (d291d91).
⚠️ Report is 3 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main     #436      +/-   ##
==========================================
- Coverage   81.31%   80.58%   -0.74%     
==========================================
  Files         314      314              
  Lines      122493   122493              
  Branches   122493   122493              
==========================================
- Hits        99606    98711     -895     
- Misses      18862    19763     +901     
+ Partials     4025     4019       -6     
Flag Coverage Δ
cli-ort-linux 83.27% <ø> (ø)
cli-ort-windows 82.67% <ø> (-0.11%) ⬇️
mlas 77.91% <ø> (-0.82%) ⬇️
offline 80.50% <ø> (-0.76%) ⬇️

Flags with carried forward coverage won't be shown. Click here to find out more.
see 7 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@github-actions

Copy link
Copy Markdown

🔴 Benchmark Regression Detected

Comparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).

ℹ️ Absolute times are informational only — they vary with runner load. The % change column is the reliable signal because both sides ran under identical conditions.

Status Scenario Base PR Change
🔴 matmul/small_generic_f16_threads=8/1x256x256 31.49 µs 56.71 µs +80.1%
🔴 matmul/small_generic_bf16_threads=1/1x256x256 32.67 µs 47.79 µs +46.3%
🔴 matmul/small_generic_f32_threads=8/1x256x256 37.96 µs 53.23 µs +40.2%
🔴 matmul/large_generic_f32_threads=8/32x1024x1024 3.75 ms 5.20 ms +38.5%
🔴 matmul/small_generic_f16_threads=1/1x256x256 30.38 µs 40.74 µs +34.1%
⚠️ matmul/small_generic_bf16_threads=8/1x256x256 35.17 µs 44.41 µs +26.3%
⚠️ matmul/large_generic_f16_threads=1/32x1024x1024 76.91 µs 95.99 µs +24.8%
⚠️ matmul/medium_generic_f16_threads=1/32x512x512 29.41 µs 34.03 µs +15.7%
⚠️ matmul/medium_generic_bf16_threads=8/32x512x512 411.60 µs 476.14 µs +15.7%
✅ gather/large_f16_threads=1-internal/131072 12.69 µs 14.29 µs +12.6%
✅ gather/large_bf16_threads=1-internal/131072 13.25 µs 14.81 µs +11.8%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.32 ms 2.57 ms +10.7%
✅ sampling_latency/top_k_per_token 438.64 µs 481.07 µs +9.7%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.34 ms 10.08 ms +7.9%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 1.85 ms 1.99 ms +7.7%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 85.35 µs 90.97 µs +6.6%
✅ sampling_latency/greedy_per_token 3.02 µs 3.20 µs +6.0%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 517.63 µs 548.12 µs +5.9%
✅ sampling_latency/top_p_per_token 904.97 µs 955.19 µs +5.5%
✅ sampling_latency/min_p_per_token 315.73 µs 327.83 µs +3.8%
✅ tokenization/encode_tokens_per_second 355.18 µs 368.42 µs +3.7%
✅ tokenization/decode_tokens_per_second 5.76 ms 5.97 ms +3.6%
✅ add/medium_f32_threads=1-internal/262144 2.48 ms 2.52 ms +1.9%
✅ kv_cache/alloc_dealloc_pages 35.90 µs 36.13 µs +0.6%
✅ grammar_masking/llguidance_compute_mask/32 68.66 µs 68.91 µs +0.4%
✅ logit_processing/seven_processor_chain_per_step 1.09 ms 1.09 ms -0.4%
✅ gather/large_f32_threads=1-internal/131072 29.77 µs 28.30 µs -4.9%
✅ matmul/small_generic_f32_threads=1/1x256x256 38.24 µs 35.76 µs -6.5%
✅ gather/small_f16_threads=1-internal/4096 476.1 ns 444.8 ns -6.6%
✅ add/large_f32_threads=1-internal/4194304 40.54 ms 37.04 ms -8.6%
✅ add/small_bf16_threads=1-internal/1024 12.97 µs 11.83 µs -8.8%
✅ add/medium_f16_threads=1-internal/262144 2.62 ms 2.37 ms -9.7%
✅ matmul/large_generic_bf16_threads=8/32x1024x1024 1.52 ms 1.37 ms -10.2%
✅ gather/small_bf16_threads=1-internal/4096 488.8 ns 437.4 ns -10.5%
✅ add/medium_bf16_threads=1-internal/262144 2.70 ms 2.41 ms -10.6%
✅ add/small_f16_threads=1-internal/1024 13.39 µs 11.96 µs -10.7%
✅ gather/medium_f16_threads=1-internal/32768 2.54 µs 2.25 µs -11.4%
✅ gather/medium_bf16_threads=1-internal/32768 2.54 µs 2.25 µs -11.6%
✅ gather/small_f32_threads=1-internal/4096 729.3 ns 640.6 ns -12.2%
✅ matmul/medium_generic_f32_threads=8/32x512x512 1.19 ms 1.04 ms -12.3%
✅ reduce_mean/large_f32_threads=1-internal/262144 1.05 ms 910.52 µs -13.1%
✅ gather/medium_f32_threads=1-internal/32768 3.94 µs 3.36 µs -14.9%
🟢 add/small_f32_threads=1-internal/1024 222.4 ns 187.5 ns -15.7%
🟢 matmul/medium_generic_f16_threads=8/32x512x512 35.42 µs 29.82 µs -15.8%
🟢 add/large_f16_threads=1-internal/4194304 46.98 ms 38.90 ms -17.2%
🟢 reduce_mean/medium_f32_threads=1-internal/65536 277.68 µs 226.75 µs -18.3%
🟢 add/large_bf16_threads=1-internal/4194304 49.67 ms 39.29 ms -20.9%
🟢 reduce_mean/small_f32_threads=1-internal/4096 19.68 µs 13.94 µs -29.1%

Visual flags: ⚠️ ≥ 15% slower, 🔴 ≥ 30% slower — calibrated against measured runner noise (~27% worst-case on multi-threaded matmul)

Host info
CPU: Apple M1 (Virtual)
Cores: 3
OS: Darwin 25.4.0 arm64
Rust: rustc 1.97.1 (8bab26f4f 2026-07-14)
Load avg: { 3.26 3.70 5.28 }
What this cannot catch
  • Regressions in code paths not covered by these benchmarks (e.g., end-to-end decode with a real model)
  • Sub-threshold regressions that compound over multiple PRs
  • Performance changes that only manifest under GPU execution
  • Latency changes in the ORT integration path (these benchmarks exercise the native Rust kernels)

@justinchuby

Copy link
Copy Markdown
Owner Author

VERDICT: APPROVE

Independent review by Lori (opus) — I did not author this code. Reviewed against origin/main, built and ran on GPU 5.

What I verified

1. KV (rank-4 seq-growing) path is preserved — no regression

  • The old inline rank-4 allocation was extracted verbatim into persistent_state_shapes(..., fixed=false): same rank-4 requirement, same axis0→1, axis2→max_len, other axes Dim::Static, symbolic→bail, and logical[2]=0. Behavior is byte-for-byte for the growing KV case.
  • fixed_state_inputs is built only from io.state_pairs (load.rs). For a pure KV model (no state_pairs) the set is empty, so: kv_bytes_per_token skips nothing, the added sort_by_key(contains) is a stable no-op, the KV filter includes all pairs, and the fixed loop does nothing → identical to main.
  • kv_binding_range = kv_start..kv_end still covers exactly the non-fixed bindings in the same name-sorted order (stable sort_by_key keeps false group first). set_logical_len/eviction only touches that range, so fixed states are never sequence-grown, and KV indices are unchanged. GPU test confirmed kv_binding_range.len()==2 (conv_state excluded) and KV still grows to logical[2]==1 after a decode step.

2. Generalizes by metadata, not by name

  • "Fixed vs seq-growing" is decided by membership in io.state_pairs (metadata), not a "conv_state" string check and not a rank-3 special case. persistent_state_shapes(fixed=true) handles arbitrary rank generically. Unit test covers rank-3 fixed conv_state and rank-4 fixed recurrent_state.

3. Allocation correctness (rank-3 FP16)

  • Declared shape → axis0→1, other axes static, symbolic→bail with a clear message. checked_shape_bytes (checked mul + dtype checked_storage_bytes, FP16=2B) guards overflow; buffer is memset-zeroed to seed the initial recurrent state. This matches the existing eager path (make_empty_input_tensor, tensor.rs:223), which already seeds fixed states at full static extent while growable KV starts at 0 — so this change aligns CUDA with eager semantics rather than changing them.

4. Tests are meaningful and pass

  • cargo test -p onnx-genai-engine --features native-backend,cuda --lib native_decode: 50 passed, 0 failed (288 filtered).
  • New cuda_persistent_state_shapes_preserve_growing_kv_and_fixed_recurrent_geometry and cuda_fixed_state_shapes_reject_unbounded_non_batch_dimensions: pass.
  • GPU binding test, actually exercised on GPU 5 (ONNX_GENAI_RUN_CUDA_SMOKE=1 CUDA_VISIBLE_DEVICES=5 taskset -c 5), 0.82s, 1 passed — verifies the rank-3 FP16 conv_state binds at [1,4,3], is zero-filled before and after a decode step, sits outside kv_binding_range, and the real KV grows. Not a no-op.

5. fmt / clippy

  • cargo fmt --all --check: clean.
  • cargo clippy -p onnx-genai-engine --features native-backend,cuda --all-targets -- -D warnings: clean (0 warnings).

Minor, non-blocking observations (no fix required to merge)

  • Fixed-state bytes are excluded from kv_bytes_per_token, so they are not subtracted from the KV max_len capacity budget. For large rank-4 recurrent states this slightly over-estimates available max_len; worst case is a loud OOM at allocation, never silent mis-sizing. Consider accounting for fixed-state footprint in a follow-up.
  • persistent_state_shapes(fixed=true) forces axis0→1 (assumes axis 0 is batch), consistent with the KV path and the eager seeder. Fine for decode batch=1.
  • tests.rs: an unrelated .iter()→.iter_mut() in native_target_step_preserves_token_driven_binding is spurious but harmless.

Note for future runs

The suggested cargo test -p onnx-genai-engine native_decode (no features) runs 0 native_decode tests — the module is gated behind native-backend. Use --features native-backend,cuda.

I did not run the full 27B E2E (optional); the rank-3 FP16 binding test is the core deliverable and it passes on GPU.

@justinchuby
justinchuby marked this pull request as ready for review July 30, 2026 08:19
@justinchuby
justinchuby enabled auto-merge (squash) July 30, 2026 08:19
@justinchuby
justinchuby merged commit f0b7193 into main Jul 30, 2026
13 of 14 checks passed
@justinchuby
justinchuby deleted the squad/384-conv-state-rank branch July 30, 2026 08:19
justinchuby added a commit that referenced this pull request Jul 30, 2026
## Summary
- add rank-3 NCL Conv support with an output-owned NVRTC kernel
- support groups/depthwise convolution, stride, dilation, optional bias,
and asymmetric causal padding
- preserve the existing rank-4 NCHW cuDNN path unchanged
- add GPU-vs-CPU EP parity for basic, depthwise causal FP16, and grouped
strided/dilated Conv1D

## Implementation choice
Rank-3 tensors cannot always be lifted directly into the existing cuDNN
path because the cuDNN legacy forward API used here requires symmetric
padding, while hybrid LLM blocks require asymmetric causal padding such
as [3, 0]. A native output-owned kernel handles the complete 1-D ONNX
geometry without staging buffers or special-casing model names.

## Qwen3.6-27B evidence
The native CUDA probe now clears the former rank-3 convolution failure
at __fn0_Conv_node_12. Execution proceeds into the layer-0
linear-attention block and stops at the next independent blocker: no
inferred shape for the Silu output
v_model.layers.0.linear_attn.conv1d.CausalConvWithState_56_0.

## Validation
- cargo test -p onnx-runtime-ep-cuda --test conv_gpu
- cargo test -p onnx-runtime-ep-cuda --test cuda_conformance_gpu
every_covered_op_has_a_conformance_entry
- cargo test -p onnx-runtime-ep-cuda --test cuda_conformance_gpu
conformance_sweep_matches_cpu
- cargo fmt --all -- --check
- cargo clippy -p onnx-runtime-ep-cuda --lib -- -D warnings
- all-target clippy remains at the same 48 pre-existing errors, with no
new Conv1D warning

## Stack
This PR is based on and requires #436. Merge order: #436, then this PR.

References #384

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Jul 30, 2026
This is a follow-up build hotfix for the native-backend clippy check
failure.

**Root cause:** LoopStatePair import at the module level is only used
inside the #[cfg(feature = "cuda")] build_cuda_decoder_with_fixed_state
function. In non-cuda builds (native-backend), the import is unused,
causing `clippy -D warnings` to fail with:
```
error: unused import: LoopStatePair
```

**Fix:** Move the LoopStatePair import behind #[cfg(feature = "cuda")]
to match its actual usage, following the same cfg-gate fix pattern from
#436.

**CI status:** Verified that `cargo clippy --locked --all-targets -p
onnx-genai-engine --features native-backend -- -D warnings` passes with
this fix.

Fixes #438 native-backend CI build failure.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Jul 30, 2026
…y exclusive (#446)

## Summary

Fast-follow to the merged live GPU weight offload (#444, closing #63):
make **weight offload** and **CUDA graph capture** explicitly **mutually
exclusive**.

### Why
Live weight offload pages weights host↔device with `cuMemAlloc` /
`cuMemcpyHtoD` / `cuMemFree` — all **illegal during CUDA graph
capture**. The previous code only skipped the residency stream-sync
while capturing (`if !is_capturing()`), which *treated capture as a
supported-but-degraded mode* even though the alloc/copy/free ops it
guards can never legally run under capture. So enabling
`ONNX_GENAI_WEIGHT_OFFLOAD=1` together with graph capture (which
auto-enables for owned CUDA KV) was a latent foot-gun — flagged in
Lori's review of #444.

### What changed
- **Mutual exclusion at the decision point.**
`resolve_graph_capture_enabled` (native decode session load) now takes a
`weight_offload_enabled` input with **highest precedence**: when offload
is on, capture resolves to **OFF** — beating structural auto-enable, an
explicit `ONNX_GENAI_CUDA_GRAPH=1`, and a programmatic `Some(true)`.
Graceful: **offload wins, capture is skipped, logged once**
(`std::sync::Once`): `"weight offload is incompatible with CUDA graph
capture; capture disabled"`.
- **Dropped the dead capture branch.** With capture now impossible while
offload runs, the `is_capturing()` guard in
`CudaWeightResidency::admit()` is unreachable — removed; `admit()` now
always synchronizes before eviction (the sync, and the paging ops it
protects, are never capture-illegal anymore).
- **cfg-correct across feature sets.** The offload query is a CUDA-EP
feature; behind `#[cfg(feature = "cuda")]` it reads
`DeviceOffloadPolicy::from_env().enabled`, and is a plain `false` on the
non-cuda (`native-backend`-only) build — no
unused-import-across-feature-sets breakage (the class that bit
#436/#441).

### Before / After
| | offload off | offload on |
|---|---|---|
| **Before** | capture per structural/env/programmatic | capture may
still auto-enable → paging ops run under capture = UB/foot-gun |
| **After** | capture per structural/env/programmatic (unchanged) |
capture forced OFF, one-time log; `admit()` always syncs safely |

## Validation
- **New test** `weight_offload_forces_graph_capture_off`: offload beats
safe structure, explicit env=1, and programmatic `Some(true)`; with
offload off the same safe structure still enables capture (proves the
exclusion is genuinely offload-caused). ✅
- Existing resolver tests updated for the new parameter; all 5 pass. ✅
- **GPU** (`CUDA_VISIBLE_DEVICES=0 taskset -c 0`): `weight_offload_gpu`
— 5/5 pass with the simplified `admit()` sync (residency page-in / reuse
/ eviction / referenced-page-pin all still correct). ✅
- `cargo fmt --all --check` clean. ✅
- `cargo clippy` clean for changed files on **both**
`cuda,native-backend` **and** `native-backend`-only (non-cuda) feature
sets. ✅

## Notes
- Draft — **do not merge**. Stacked as a small safety fast-follow on the
merged #444.
- No change to the non-offload default fast path; capture behavior is
byte-identical when offload is disabled.

Refs #444, #63

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant