Skip to content

fix(session): expose bound inputs to control flow - #443

Merged
justinchuby merged 2 commits into
mainfrom
squad/384-value1414
Jul 30, 2026
Merged

justinchuby merged 2 commits into
mainfrom
squad/384-value1414

Conversation

@justinchuby

@justinchuby justinchuby commented Jul 30, 2026 •

Copy link
Copy Markdown
Owner

Summary

  • make persistent device input/output bindings visible to control-flow materialization
  • route If, Loop, and Scan output storage through the general external-output-aware path
  • add a two-step persistent-state regression proving capture reads and bound-output writeback

Root cause and correctness fix

value#1414 is past_key_values.0.recurrent_state, captured by the Scan inside Qwen3.6's inlined com.microsoft::LinearAttention. Control flow first failed because capture materialization ignored external bindings. Once reads worked, control-flow outputs still went only to executor-owned buffers, bypassing the persistent present.*.recurrent_state binding. Every decode step reread stale initial state and repeated the prompt.

The fix treats external bindings as first-class at both boundaries: captures/formal inputs can be materialized from bound storage, and control-flow output writes reuse store_raw_tensor_output. It is general across externally bound If, Loop, and Scan tensors and does not depend on state names, ranks, or Qwen geometry.

Decisive Qwen3.6-27B parity evidence

Prompt: The capital of France is; greedy, first 16 generated tokens.

  • native CUDA: [11751, 13, 271, 248068, 271, 248069, 271, 4639, 369, 4252, 13, 11751, 369, 279, 6511, 321]
  • trusted ORT-CPU: [11751, 13, 271, 248068, 271, 248069, 271, 4639, 369, 4252, 13, 11751, 369, 279, 6511, 321]
  • text: ` Paris.

That is correct. Paris is the capital and`

  • native CUDA: 5.26 tok/s; ORT-CPU: 1.67 tok/s

ORT-CUDA 1.27 and 1.28 abort during graph optimization for this model, so ORT-CPU with basic optimization is the trusted reference. Verdict: FIXED; exact token parity proven.

Regression

if_materializes_outer_capture_from_persistent_device_input binds X/Y to one device allocation, runs twice, and asserts [5, 9] -> [6, 10] -> [7, 11].

Validation

  • cargo test -p onnx-runtime-session --test control_flow: 23 passed
  • formatting check passed
  • session clippy passed
  • engine native-backend clippy passed with and without CUDA
  • real Qwen3.6-27B native CUDA decode on GPU4

References #384

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby
justinchuby force-pushed the squad/384-value1414 branch from c7aa193 to a3143e0 Compare July 30, 2026 09:48
@codecov

codecov Bot commented Jul 30, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 60.56338% with 28 lines in your changes missing coverage. Please review.
✅ Project coverage is 81.31%. Comparing base (f3b9bc8) to head (815164b).

Files with missing lines Patch % Lines
.../onnx-runtime-session/src/executor/control_flow.rs 58.33% 8 Missing and 17 partials ⚠️
crates/onnx-runtime-session/src/executor/state.rs 70.00% 3 Missing ⚠️
Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main     #443      +/-   ##
==========================================
+ Coverage   80.58%   81.31%   +0.73%     
==========================================
  Files         314      314              
  Lines      122493   122520      +27     
  Branches   122493   122520      +27     
==========================================
+ Hits        98711    99633     +922     
+ Misses      19763    18863     -900     
- Partials     4019     4024       +5     
Flag Coverage Δ
cli-ort-linux 83.27% <ø> (ø)
cli-ort-windows 82.78% <ø> (ø)
mlas 78.72% <ø> (+0.81%) ⬆️
offline 81.27% <60.56%> (+0.76%) ⬆️

Flags with carried forward coverage won't be shown. Click here to find out more.

Files with missing lines Coverage Δ
...ates/onnx-runtime-session/src/executor/dispatch.rs 73.23% <100.00%> (ø)
crates/onnx-runtime-session/src/executor/state.rs 83.63% <70.00%> (-1.37%) ⬇️
.../onnx-runtime-session/src/executor/control_flow.rs 64.18% <58.33%> (+0.32%) ⬆️

... and 8 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

github-actions Bot commented Jul 30, 2026 •

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_f32_threads=8/1x256x256 34.82 µs 51.55 µs +48.1%
⚠️ matmul/small_generic_f16_threads=8/1x256x256 30.87 µs 37.60 µs +21.8%
⚠️ matmul/medium_generic_f16_threads=8/32x512x512 36.06 µs 42.67 µs +18.3%
⚠️ gather/large_bf16_threads=1-internal/131072 14.55 µs 16.83 µs +15.7%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 1.92 ms 2.20 ms +14.8%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.06 ms 10.35 ms +14.3%
✅ matmul/large_generic_f32_threads=8/32x1024x1024 4.57 ms 5.14 ms +12.5%
✅ matmul/small_generic_f32_threads=1/1x256x256 38.46 µs 42.56 µs +10.6%
✅ add/medium_f16_threads=1-internal/262144 2.59 ms 2.87 ms +10.6%
✅ matmul/small_generic_f16_threads=1/1x256x256 32.65 µs 36.10 µs +10.6%
✅ matmul/medium_generic_bf16_threads=8/32x512x512 513.16 µs 565.47 µs +10.2%
✅ matmul/large_generic_bf16_threads=8/32x1024x1024 2.06 ms 2.23 ms +7.9%
✅ add/small_bf16_threads=1-internal/1024 13.35 µs 14.24 µs +6.7%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 82.82 µs 88.30 µs +6.6%
✅ add/medium_bf16_threads=1-internal/262144 2.63 ms 2.79 ms +6.1%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 93.86 µs 98.55 µs +5.0%
✅ add/large_f16_threads=1-internal/4194304 42.86 ms 44.92 ms +4.8%
✅ gather/small_f32_threads=1-internal/4096 675.1 ns 706.9 ns +4.7%
✅ gather/medium_f32_threads=1-internal/32768 4.17 µs 4.33 µs +4.0%
✅ matmul/medium_generic_f32_threads=8/32x512x512 1.70 ms 1.75 ms +3.2%
✅ add/large_f32_threads=1-internal/4194304 40.61 ms 41.78 ms +2.9%
✅ gather/small_f16_threads=1-internal/4096 482.7 ns 495.6 ns +2.7%
✅ reduce_mean/small_f32_threads=1-internal/4096 15.39 µs 15.72 µs +2.2%
✅ reduce_mean/large_f32_threads=1-internal/262144 984.40 µs 1.00 ms +2.0%
✅ tokenization/encode_tokens_per_second 402.72 µs 408.40 µs +1.4%
✅ reduce_mean/medium_f32_threads=1-internal/65536 247.80 µs 250.70 µs +1.2%
✅ matmul/medium_generic_f16_threads=1/32x512x512 35.10 µs 35.39 µs +0.8%
✅ gather/medium_f16_threads=1-internal/32768 2.53 µs 2.54 µs +0.7%
✅ add/medium_f32_threads=1-internal/262144 2.61 ms 2.63 ms +0.6%
✅ sampling_latency/greedy_per_token 3.31 µs 3.32 µs +0.4%
✅ gather/large_f32_threads=1-internal/131072 42.30 µs 42.15 µs -0.4%
✅ gather/medium_bf16_threads=1-internal/32768 2.60 µs 2.58 µs -0.8%
✅ add/small_f16_threads=1-internal/1024 14.22 µs 13.91 µs -2.2%
✅ matmul/small_generic_bf16_threads=1/1x256x256 38.86 µs 37.39 µs -3.8%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 564.17 µs 541.87 µs -4.0%
✅ tokenization/decode_tokens_per_second 7.14 ms 6.77 ms -5.2%
✅ gather/large_f16_threads=1-internal/131072 17.57 µs 16.57 µs -5.7%
✅ matmul/small_generic_bf16_threads=8/1x256x256 42.48 µs 38.88 µs -8.5%
✅ add/small_f32_threads=1-internal/1024 224.4 ns 204.9 ns -8.7%
✅ kv_cache/alloc_dealloc_pages 40.62 µs 36.73 µs -9.6%
✅ sampling_latency/top_p_per_token 1.06 ms 955.68 µs -10.0%
✅ sampling_latency/top_k_per_token 524.85 µs 455.09 µs -13.3%
✅ gather/small_bf16_threads=1-internal/4096 580.5 ns 502.8 ns -13.4%
✅ logit_processing/seven_processor_chain_per_step 1.37 ms 1.17 ms -14.5%
🟢 add/large_bf16_threads=1-internal/4194304 51.94 ms 43.71 ms -15.8%
🟢 matmul/medium_generic_f32_threads=1/32x512x512 3.10 ms 2.59 ms -16.4%
🟢 sampling_latency/min_p_per_token 395.00 µs 328.43 µs -16.9%
🟢 grammar_masking/llguidance_compute_mask/32 91.52 µs 73.15 µs -20.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: { 4.01 3.85 6.05 }
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)

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

Copy link
Copy Markdown
Owner Author

VERDICT: APPROVE

Independent review of PR #443 (branch squad/384-value1414, commit 815164b). Author Mary locked out; reviewed on merits in scratch worktree melina-443. CPU-testable executor logic; no GPU needed.

Summary

Both bugs are fixed correctly and, critically, the write-back is GENERAL — not a benchmark-specific hack.

1. Write-back fix (CRITICAL) — CORRECT & GENERAL

Control-flow output storage now routes every store through the pre-existing, external-output-aware path store_raw_tensor_output (sequence_ops.rs:467-487). When a control-flow output vid is a persistent binding, it copies into the device binding via writable_buffer() (sequence_ops.rs:477-479); otherwise it falls back to the executor-owned buffer (:480-484). This is applied uniformly to ALL control-flow output kinds:

  • If outputs: control_flow.rs:683
  • Loop carried finals: control_flow.rs:984; Loop scan outputs: :986-993
  • Scan state finals: control_flow.rs:1352; Scan scan-outputs: :1361-1367
    The old code (removed at control_flow.rs:454-459 region) always allocated a fresh executor buffer and never touched external.outputs, so present..recurrent_state never replaced the paired past. binding -> stale zeros re-read each step -> repeated token. Keyed purely on external.outputs.contains_key(&vid); no tensor-name/rank/model special-casing. DRY (reuses the sequence-op path). Meets Justin's GENERAL rule.

2. Capture fix — scoping sound, no leak

value_tensor now reads external.inputs (falling back to external.outputs) when the vid is externally bound (control_flow.rs:251-272), and prepare_subgraph counts external inputs/outputs as "materialized" (control_flow.rs:486-495). Scope resolution is unchanged: a subgraph free variable is resolved by name against the outer name_index, gated by resolved.contains_key(&vid) && materialized. External bindings are top-level graph I/O values that are legitimately in enclosing scope for inlined control flow; they were previously wrongly EXCLUDED (cause of "value#1414 not produced"). No new value becomes visible that ONNX name-capture semantics didn't already permit. dtype/byte-capacity guard at :262-267 rejects mismatched bindings. readable_buffer() null-checks (state.rs:584-596). No scope leak found.

3. No default-path regression

cargo test -p onnx-runtime-session: 228 passed; 0 failed (includes GQA/plain-KV decode tests executor.rs:370,378). The new path only diverges when external.outputs.contains_key(&vid); plain KV-attention decode is unaffected. PASS.

4. Regression test — non-tautological, named

if_materializes_outer_capture_from_persistent_device_input (tests/control_flow.rs:328-359). Binds one device allocation as both input X and output Y (allocate_device_binding("X", Some("Y"))); the If subgraph CAPTURES outer X (if_branch -> capture(b,"X"), control_flow.rs test :121) and Adds ones. Runs twice asserting [5,9]->[6,10]->[7,11]. Without the write-back, step 2 re-reads [5,9] and yields [6,10], so the assert_eq at step 2 FAILS. It exercises both the capture read and the output write-back and would catch a reintroduced stale-state bug. Not a tautology.

5. Control-flow suite

cargo test -p onnx-runtime-session --test control_flow: 23 passed; 0 failed (incl. the new regression test).

6. Feature sets / lint

7. Aliasing / ordering (advisory) — no hazard

Write-back copies from a host Tensor (tensor.as_bytes()), i.e. the subgraph output was already device->host materialized (synchronizing), then host->device into the binding. Within a step, captured input X is read during prepare_subgraph BEFORE the subgraph runs; output Y is written AFTER — correct read-before-write ordering even when X and Y alias one allocation. copy_from_host is stream-ordered/synchronous, so the write-back cannot race the just-finished subgraph kernels. Note (non-blocking): recurrent state incurs a device->host->device round-trip per step; correctness-safe and consistent with existing control-flow output handling.

Verdict: APPROVE. Write-back is general; scoping sound; 228 session + 23 control_flow tests green; regression test is real. No blocking concerns.

— Melina (independent review)

@justinchuby
justinchuby marked this pull request as ready for review July 30, 2026 10:44
@justinchuby
justinchuby merged commit 1ba215e into main Jul 30, 2026
15 checks passed
@justinchuby
justinchuby deleted the squad/384-value1414 branch July 30, 2026 10:45
justinchuby added a commit that referenced this pull request Jul 31, 2026
…ture scoping (#556)

State consolidation only (no code). Merges 3 decision notes into
decisions.md (20332→20465 B, under gate).

**Key durable record:** the ⚑ PENDING-JUSTIN 27B Scan→CUDA-capture
decision — the ~15-30× decode lever is NOT a contained increment
(prefill/decode share one plan/session; static single-trip inline
corrupts prefill). Correct fix = runtime shape-conditional dual-path
capture infra, which reshapes the delicate #443/#543 control-flow
capture core. Awaiting @justinchuby go-ahead before touching that core.

Also records: 27B decode profile (Scan=56.5%, ~35× off HBM roofline,
structural not kernel), and the #554 session-reuse fix. Round-8 wave
record archived; histories compacted.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 1, 2026
Transpose used launch_metadata, which per call did alloc_raw + htod
(perm/stride upload) + synchronize() + free_raw. The synchronize() only
guards the immediate free but serializes the stream on every one of the
~624 Transpose calls per qwen3.6-27b decode step, in both the eager
(Scan-child) and captured paths, and left Transpose reporting
CaptureSupport::Unsupported so it could never fold into the parent CUDA
graph.

Migrate Transpose to the PersistentMetadata + launch_persistent_metadata
pattern already blessed for Expand and Slice (#443/#543): cache the
perm/stride buffer, upload/sync once at warm-up, then run with no
per-call alloc/htod/sync. Once warmed on its exact shape/dtype/perm
signature the kernel reports Supported and errors on mid-capture
signature change (identical guard to Expand). Same transpose_bytes
kernel and identical metadata bytes -> byte-exact.

Measured on stock qwen3.6-27b-int4 (auto-derived io, #573): decode
165.27 -> 156.74 ms/tok (~5.4%, ~1.05x), token IDs byte-identical to the
CPU fp32 oracle (ORT-CUDA crashes on this model). Does not beat ORT
alone; this is one op of the Lever-2 non-Scan swarm.

Tile deferred (separate PR): it shares launch_metadata but also reads
repeats device->host via host_ints, needing repeats-warming.

Tests: new transpose_warmed_metadata_captures_and_matches_eager EP unit
test (unwarmed declines -> warms -> capture+replay byte-identical to
eager); full construction_gpu suite 19/19; 27b autoderive oracle e2e and
qwen3-0.6b e2e both byte-identical.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 1, 2026
…a (Lever 2 PR-1) (#574)

## Lever 2 PR-1 — make CUDA `Transpose` capture-eligible via persistent
metadata

Part of the Lever-2 non-Scan swarm from `cohaagen-27b-decode-perf.md`
toward closing the native-vs-ORT gap on the qwen3.6-27b (dense-GQA +
LinearAttention hybrid). **Transpose only** — bounded, byte-exact,
single-op.

### Problem
`TransposeKernel` used `launch_metadata`, which on **every call** did
`alloc_raw` + `htod` (perm/stride upload) + `synchronize()` +
`free_raw`. The `synchronize()` only guards the immediate `free_raw`,
but it serializes the stream on every one of the ~624 Transpose calls
per 27b decode step — in **both** the eager (Scan-child) and captured
paths — and left Transpose reporting `CaptureSupport::Unsupported`, so
it could never fold into the parent CUDA graph.

### Fix
Migrate `Transpose` to the in-repo `PersistentMetadata` +
`launch_persistent_metadata` pattern already blessed for **Expand** and
**Slice** (#443/#543): cache the perm/stride buffer, upload/sync
**once** at warm-up, then run with **no per-call alloc/htod/sync**. Once
warmed on its exact shape/dtype/perm signature the kernel reports
`Supported` and errors on a mid-capture signature change (identical
guard to Expand). Same `transpose_bytes` kernel, identical metadata
bytes → **byte-exact**.

### Blast radius
- ONLY `crates/onnx-runtime-ep-cuda/src/kernels/movement.rs`
(TransposeKernel + factory) + one new EP test.
- Shared capture machinery (segmenter, `subgraph_graph_capturable`,
`capture_quarantine_ops`, run.rs) untouched — Transpose just starts
reporting `Supported` once warmed.
- **Tile deferred** (separate PR): shares `launch_metadata` but also
reads `repeats` device→host via `host_ints`, needing repeats-warming —
not the same trivial mechanism.

### Measured yield (stock qwen3.6-27b-int4, CUDA; --steady --decode-skip
8 --warmups 2 --runs 3 --tokens 64)

| config | decode ms/tok | tok/s | first 16 token IDs |
|---|---:|---:|---|
| baseline (origin/main) | 165.27 | 6.05 | [11751, 13, 271, 248068, 271,
248069, 271, 4639, 369, 4252, 13, 11751, 369, 279, 6511, 321] |
| this PR | 156.74 | 6.38 | [11751, 13, 271, 248068, 271, 248069, 271,
4639, 369, 4252, 13, 11751, 369, 279, 6511, 321] |

**-8.5 ms/tok (~5.4%, ~1.05×), token IDs byte-identical.**

Honest caveat: this is one op of the Lever-2 swarm and does **NOT** beat
ORT (17.38 tok/s / 57 ms/tok) alone — that is the two-lane result (full
Lever-2 swarm + Inc-1b Scan-body capture). Body-internal Transposes get
the desync benefit now; their capture-fold benefit needs Inc-1b (the
Scan child never captures today).

### Correctness (all pass)
- **New EP unit test**
`transpose_warmed_metadata_captures_and_matches_eager`: unwarmed →
declines capture; first eager run warms → `Supported`; cached-metadata
re-run byte-identical; capture+replay byte-identical to eager.
Non-vacuous both directions.
- **Full `construction_gpu` EP suite:** 19/19 pass — capture surface
unregressed.
- **27b oracle gate**
`native_autoderive_io_cuda_e2e::stock_export_auto_derives_io_and_matches_cpu_oracle`:
PASS — stock 27b native CUDA == CPU fp32 oracle, byte-identical.
(ORT-CUDA crashes on this model → CPU fp32 is the oracle, per #384.)
- **Small-model gate** `qwen3_0_6b_native_cuda_e2e`: PASS; decode token
IDs == pinned golden.
- `cargo fmt --all --check` clean.

Decision note: `.squad/decisions/inbox/cohaagen-lever2-transpose-pr.md`.

Do **not** auto-merge — Harry reviews first (author-lockout).

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 1, 2026
…ver 2 PR-2) (#577)

## What this unlocks

Lever 2 PR-2. Migrates the CUDA **Tile** kernel to the
`PersistentMetadata` / `launch_persistent_metadata` pattern already used
by Transpose (#574) and Expand, so a warmed fixed-decode-signature Tile
becomes **CUDA-graph capture-eligible** and drops its per-call host
sync.

Post-#574, Tile was the last `launch_metadata` caller in `movement.rs`.
Per call it did: device `alloc_raw` + `htod` of the metadata + kernel
launch + **`synchronize()`** (to guard the `free_raw`) + free — plus a
blocking `dtoh` read of the `repeats` input. On the qwen3.6-27b
linear-attention hybrid, Tile runs ~576×/step inside eager Scan bodies,
so those syncs are on the decode hot path.

## Correctness — byte-exact vs CPU fp32 oracle

ORT-CUDA crashes on the 27b, so the CPU fp32 backend is the oracle.

- **27b stock export** native CUDA == CPU fp32 oracle,
**byte-identical** (`native_autoderive_io_cuda_e2e`, `#[ignore]` GPU e2e
— passed).
- **qwen3-0.6b** native CUDA == ORT golden, **byte-identical**
(`qwen3_0_6b_native_cuda_e2e` — passed); capture engagement unchanged
(captures=2, replays=34, fallbacks=0).
- New **positive** test
`tile_warmed_metadata_captures_and_matches_eager`: unwarmed declines
capture; warmed reports `Supported`; eager re-run and capture+replay
both byte-identical.
- New **negative** test `tile_rejects_signature_change_during_capture`:
feeding a changed signature mid-capture returns `Err` "changed during
CUDA graph capture" — locks the guard's teeth.

The `repeats` input (a blocking `dtoh`, illegal mid-capture) is now read
**only on the unwarmed eager path** to validate `output[i] == input[i] *
repeats[i]`. Under a stable `(dtype, input_shape, output_shape)`
signature, `repeats` is mathematically fixed by shape inference, so
gating that read behind the warmed signature is byte-exact. Kernel
`tile_bytes` metadata (`[output.shape, input.shape, input_strides]`) is
unchanged; both launch paths pass identical args.

## Honest perf

Stock 27b, `profile_native --steady --decode-skip 8 --warmups 2 --runs 3
--tokens 64`:

| | ms/tok | tok/s |
|---|---|---|
| Baseline (origin/main @15288d7a) | 156.3 | 6.40 |
| This PR | 150.0 | 6.67 |

**~4.0% faster**, tokens byte-identical. This does **not** beat ORT
(17.38 tok/s) on its own — that requires the Inc-1b Scan-body capture
lane (awaiting Justin's greenlight) plus removal of the shared eager
`synchronize()` on the ~50 always-sync kernels. This is one more fat op
off the eager-sync hot path, in the same reviewable single-op shape as
#574.

## Blast radius

Tile kernel path + its two tests only. Dead `launch_metadata` removed
(Tile was its sole remaining caller). **No** change to shared capture
machinery, the #443/#543 capture-correctness invariants, weight offload,
session-reuse, or GAP-3.

Refs #574 (Lever 2 lineage), #443/#543 (capture invariants).

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 1, 2026
…, unwired) (#580)

## Inc-1b PR-1 — `inline_single_trip_scan_bodies` graph transform

First bounded step of the Inc-1b lane (`cohaagen-27b-inc1b-design.md`
§1). Adds a new IR graph transform in `onnx-runtime-ir` and nothing
else.

### What the transform does
`inline_single_trip_scan_bodies(&Graph) -> Graph` lowers each
**single-trip-eligible** `Scan` node to one straight-line iteration of
its body spliced into the parent graph:

- **Eligibility (structural, safe-by-construction):** default-domain
`Scan` with a registered `body` subgraph, carrying **recurrent state**
(`num_state >= 1`, i.e. a LinearAttention body threading `state_pairs`),
whose every scan input's scan axis is **not a static extent > 1**
(`Static(1)` or dynamic — the decode case), with **no nested control
flow**, and whose lexical captures all resolve by name in parent scope.
Anything else is left untouched.
- **Scan input → body:** a `Squeeze` drops the size-1 scan axis and
feeds the body scan-input formal.
- **Body splice:** every body node is cloned into the parent with
**fresh, remapped `ValueId`s** (`remap: body ValueId → new parent
ValueId`). Body state formals bind to the Scan's parent state-input
values; the body scan-input formal binds to the squeezed slice;
**lexical captures resolve by name** to values already in parent scope
(same resolution `ChildExecutor::new` does via `capture_names`);
**body-local initializers are promoted** to parent initializers.
- **Body outputs → Scan outputs:** the first `num_state` body outputs
become the parent present-state values (written directly by their
producer; an `Identity` is used only for a rare pass-through); scan
outputs get an `Unsqueeze` re-adding the size-1 axis to match the Scan's
declared scan-output shape.
- **Delete** the `Scan` node and drop its body subgraph.

The transform is a **pure `Graph -> Graph`** function — the input graph
is never mutated.

### Emphatically UNWIRED / zero behavior change / dead until PR-2
This pass is **not called from any executor, decoder, loader, or run
path**. It is dead code exercised **only** by the unit tests. The
prefill / multi-trip / dense paths are completely unaffected. No changes
to `control_flow.rs`, `run.rs`, `state.rs`, or decode routing.

- **PR-2** (later) will build a `decode_inline_exec` from the
transformed graph and route decode to it behind a **default-off** flag
(eager, transform-correctness proof).
- **PR-3** (later) lets the inlined body **capture** — that is the
highest-blast-radius step and **still needs Justin's greenlight +
capture-team (Mary / #443–#543) sign-off.** **This PR does not** and is
safe to review/merge on its own.

### Crate-boundary note (why inference isn't called inside the pass)
`onnx-runtime-ir` is the dependency-free base contract and cannot depend
on `onnx-runtime-shape-inference` (which depends on it) — not even as a
dev-dependency, since that pulls in a second, incompatible copy of the
IR types. So the transform performs the **structural merge only**
(mirroring the graph-building half of `ChildExecutor::compile`); a
caller (PR-2) re-runs `InferenceRegistry::infer_graph` (Permissive) to
re-resolve interior shapes, exactly as `ChildExecutor::compile` does
today. The "inference re-converges" check is therefore placed in the
shape-inference crate's integration tests, on the side of the edge that
can see both.

### Tests (CPU, fast, non-ignored)
`onnx-runtime-ir` — `src/scan_inline/tests.rs` (5 tests, shared
tiny-graph builder):
- **`lowers_single_trip_scan_to_straight_line_body`** — proves the Scan
is removed; body ops (MatMul + 3 Adds) are present; one `Squeeze` + one
`Unsqueeze` appear; the parent present-state is produced by an inlined
`Add`; the scan output is produced by `Unsqueeze`; the lexical capture
`w` resolves by name to the pre-existing parent weight the inlined
MatMul reads; the body-local `bias` initializer is promoted; the graph
is structurally valid; and the rewired boundary/interior values carry
the expected static shapes.
- **`leaves_multi_trip_scan_untouched`** — recurrent Scan with a static
scan-axis extent of 3 is a no-op (must not statically collapse a genuine
multi-trip Scan).
- **`leaves_non_recurrent_scan_untouched`** — a pure element-wise map
Scan (`num_state == 0`) is a no-op.
- **`leaves_dense_scanless_graph_untouched`** — a dense Scan-free graph
is returned structurally unchanged (non-vacuous both directions).
- **`remap_produces_fresh_non_colliding_interior_ids`** — with
body/parent id ranges deliberately overlapping, asserts every
node-referenced id is live, every interior value id is freshly allocated
above the pre-transform high-water mark (no leaked body id, no alias of
a pre-existing parent value), and `validate()` passes (no id collisions
/ consistent edges).

`onnx-runtime-shape-inference` — `tests/scan_inline_inference.rs` (1
test):
- **`permissive_inference_reconverges_over_inlined_scan_body`** —
inlines the eligible hybrid, blanks every produced value's shape, then
asserts Permissive whole-graph inference re-derives the interior +
boundary shapes (present state `[2,4]`, scan output `[1,2,4]`) from the
graph inputs/initializers through the inlined nodes.

### Verification
- `cargo test -p onnx-runtime-ir` — 67 passed (incl. the 5 new).
- `cargo test -p onnx-runtime-shape-inference --test
scan_inline_inference` — 1 passed.
- `cargo fmt --all --check` — clean.
- `cargo clippy -p onnx-runtime-ir --all-targets -- -D warnings` — exit
0; `cargo clippy -p onnx-runtime-shape-inference --tests -- -D warnings`
— exit 0.
- `cargo check` for `onnx-runtime-ir`, `-shape-inference`, `-optimizer`,
`-loader` (its dependents) — clean. (A workspace-wide `cargo check`
fails only on the unrelated `onnx-runtime-cpuinfo` cmake vendor
submodule, which is not populated in a fresh git worktree —
environmental, not from this change; PR-1 does not touch ep-cuda or any
GPU code.)

### Scope held
Size **M**, **LOW** blast radius (new pass, dead until PR-2), exactly as
the design's staging table predicts. No execution/capture surfaces were
touched to compile or test it.

Per Inc-1b design: `cohaagen-27b-inc1b-design.md`. Please review before
merge — **do not enable auto-merge** (independent reviewer: Harry).

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 2, 2026
…R, flag default-OFF) (#588)

## Inc-1b PR-2 — wire a decode-specialized inlined-body Executor
(dual-plan EAGER, flag default-OFF)

Wires PR-1's `inline_single_trip_scan_bodies` transform (#580) into a
**second, decode-specialized `Executor`** and routes single-token decode
to it. Eager (capture stays OFF). **Flag default-OFF ⇒ zero behavior
change unless explicitly enabled.**

### What this does
- **Sibling Executor** (`onnx-runtime-session`):
`Executor::build_decode_inline_sibling` runs the transform, re-resolves
interior shapes with Permissive inference (mirrors
`ChildExecutor::compile`), and builds a second executor that **shares
the main exec's `Arc<WeightStore>` and `Arc<dyn ExecutionProvider>`**.
Returns `None` for a dense (non-hybrid) decoder. The prefill/main exec
is left **byte-identical**.
- **Session API**: `InferenceSession::{enable_decode_inline,
decode_inline_ready, run_decode_inline_with_device_bindings}`. The
sibling binds the **identical persistent device state buffers** the main
exec used at the prefill→decode hand-off (bindings resolve by name; the
transform leaves graph input/output names+order unchanged ⇒
recurrent-state continuity is automatic — design §3, the integration
invariant).
- **Engine flag + routing** (`onnx-genai-engine`):
`ONNX_GENAI_DECODE_INLINE_SCAN` (default OFF; truthy
`1`/`true`/`yes`/`on`). Lazy build at the first single-token decode
step. Single-token decode routes to the sibling on all three native
single-token paths — CUDA greedy device-argmax fast path, CUDA logits
path, and CPU in-place path. Greedy reuses the existing device-argmax
kernel, so tie-breaking is byte-identical and full logits never
round-trip to host.

### Flag default-OFF — zero behavior change when off
When `ONNX_GENAI_DECODE_INLINE_SCAN` is unset/falsy, the sibling is
never built, `decode_inline` latches `Disabled`, and every decode step
uses today's Scan child-session path unchanged. An ordinary session is
byte-identical to `main` and pays nothing.

### Harry's 4 mandatory guards → tests
1. **Byte-identical parity + final recurrent state** —
`decode_inline_sibling_is_byte_exact_with_scan_and_preserves_state`
(session): N decode steps, Scan plan vs inline plan, per-token outputs
byte-identical AND final recurrent state identical.
2. **Runtime scan-axis extent==1 assertion + fallback** —
`route_decode_inline_decision` (pure) +
`decode_inline_routes_only_single_token_when_enabled` /
`decode_inline_never_routes_when_disabled_or_unbuilt` (engine): only
single-token (extent-1) steps route to the sibling; multi-token steps
fall back to the main Scan exec so a wrongly-collapsed graph is never
run.
3. **Persistent state-buffer continuity** —
`decode_inline_sibling_preserves_persistent_state_across_prefill_handoff`
(session).
4. **state_pairs ordering + shape check** —
`decode_inline_sibling_preserves_state_output_order_and_resolves_shapes`
(session): first `num_state` present outputs map to present-state in
`io.state_pairs` order; inlined-interior shapes resolve (Permissive)
before use.

Plus `decode_inline_sibling_none_for_dense_graph` and
`decode_inline_flag_defaults_off_and_parses_truthy`.

### Measured OFF vs ON — Qwen3.6-27B int4 hybrid (H200)
`profile_native --ep cuda --backend native --steady --decode-skip 8
--warmups 2 --runs 3 --tokens 64`

| flag | decode ms/tok (medians, 4 runs) | tok/s |
|------|----------------------------------|-------|
| OFF  | ~153 (149.2 / 155.9 / 157.7 / 150.6) | ~6.5 |
| ON   | ~122 (117.7 / 124.0 / 121.9 / 126.7) | ~8.2 |

**~1.26× decode speedup** (design §5 predicted ~1.28×; the 167→130
ms/tok absolutes were on a slower baseline — same ratio here).
**Generated token ids byte-identical OFF vs ON on every run.** The eager
inline plan beats even the CUDA-graph-captured baseline because the
captured Scan operator still pays real per-step child-dispatch +
loop-state-collect work inside each replay; inlining removes that
boundary entirely (design §5).

### Byte-exact GPU e2e
`native_autoderive_io_cuda_e2e.rs` (`#[ignore]`, stock 27b == CPU fp32
oracle, expected ids
`[11751,13,271,248068,271,248069,271,4639,369,4252,13,11751,369,279,6511,321]`)
run with the flag ON — see PR comment for the run result.

### Scope
- Only native single-token decode paths route to the sibling;
routed-port/`inputs_embeds` and multi-token steps keep the main exec.
- **Capture is OUT OF SCOPE** — device-graph capture of the inlined body
is **PR-3 (greenlight-gated, capture-team #443/#543 sign-off
required)**. This PR touches none of the capture surface.
- fmt clean; `clippy -D warnings` exit 0 on `onnx-runtime-session` and
`onnx-genai-engine` (incl. `--features cuda,native-backend`).

Design: `cohaagen-27b-inc1b-design.md` §1–3; transform PR #580.

**Do not auto-merge — Harry reviews first (author-lockout).**

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

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 2, 2026
…et-A) (#589)

# Inc-1b PR-3 — capture-fold the decode-inline sibling (flag-gated,
bucket-A)

Part of the Inc-1b 27B decode-perf lane. PR-1 (#580,
inline_single_trip_scan_bodies) and PR-2 (#588, the eager decode-inline
sibling behind ONNX_GENAI_DECODE_INLINE_SCAN, default OFF) are merged.
This is **PR-3, the capture step**: let the decode-inline sibling's
inlined body ops fold into the CUDA-graph capture.

Bucket-(A)-only per my accepted scope note
(.squad/decisions/inbox/cohaagen-inc1b-pr3-scope.md): **no change to the
shared #443/#543 capture surface**. The sibling is an ordinary Executor
whose plan has no Scan after inlining, so it reuses the segmenter /
warm-seed / quarantine machinery verbatim; PR-3 only drives it through
the existing capture state machine.

## Flag-gated, default-OFF (structural no-op)
Capture engages only when ONNX_GENAI_DECODE_INLINE_SCAN gates
route_inline. Flag off: the sibling is never built and the inline branch
is never taken — byte-identical to current main.

## Harry's 4 PR-3 invariants (each with a non-vacuous test)
1. **Re-introduce check_device_capture_error()** on the sibling capture
path, piggybacked on the single logits/greedy device-to-host sync
(detection-before-consumption). The latch lives on the shared EP, so the
poll observes the sibling's captured-replay result; a latched violation
rejects the token and invalidates the graph.
2. **Capture-engagement test**
decode_inline_sibling_folds_body_into_captured_graph_byte_exact: the
inlined body folds into >= 1 captured segment while staying byte-exact
with the eager sibling run.
3. **Inlined-interior shapes join the warm-seeded snapshots** — proven
by the engagement test capturing after an eager warmup (warm-seed is the
precondition for capture engaging).
4. **Scope-lock** route_decode_inline_decision refuses inputs_embeds /
Routed step-input decoders (new has_eager_step_inputs arg); covered by
decode_inline_never_routes_when_decoder_has_eager_step_inputs.

## Single-slot / single-latch EP safety
Prefill runs the main exec (multi-token, non-capturable); ALL
single-token decode routes to the sibling; the main capture machine
stays dormant. So the shared EP's one graph slot + one capture-error
latch are owned solely by the sibling — no double-capture, no
cross-latch bleed. invalidate_graph now also resets the sibling's
host-side capture schedule so KV-growth / shape-change re-warms instead
of replaying a dropped graph.

## Correctness gate (real 27B, H200, qwen3.6-27b-int4-cuda)
Byte-exact vs the CPU fp32 oracle with capture engaged (flag ON), greedy
ids:

[11751, 13, 271, 248068, 271, 248069, 271, 4639, 369, 4252, 13, 11751,
369, 279, 6511, 321]

native_autoderive_io_cuda_e2e passed with
ONNX_GENAI_DECODE_INLINE_SCAN=1 (CUDA tokens == CPU oracle tokens).

## Measured decode perf (27B, H200, ms/tok, prefill+load cancelled by
token-count delta)
- flag OFF (main eager path): 143.8 ms/tok
- flag ON + capture: 70.1 ms/tok  => **2.05x**

ORT-CUDA crashes on this hybrid linear-attention export (documented
stl_vector assertion), so there is no live ORT baseline for this model;
the trusted reference is our native CPU fp32 oracle.

## Mutation map
- crates/onnx-runtime-session/src/lib.rs — 5 additive sibling capture
wrappers on InferenceSession + the CUDA capture-engagement test.
- crates/onnx-genai-engine/src/native_decode/cuda.rs —
inline_graph_phase field; new run_one_token_inline capture state
machine; invalidate_graph resets the sibling too; both inline branches
drive capture + re-add the capture-error poll; has_eager_step_inputs
widened to pub(super).
- crates/onnx-genai-engine/src/native_decode/mod.rs —
route_decode_inline_decision gains has_eager_step_inputs.
- crates/onnx-genai-engine/src/native_decode/tests.rs — updated calls +
new scope-lock test.

## Not in scope
Does NOT flip the default to ON — that stays Justin's decision. Build
evidence: .squad/decisions/inbox/cohaagen-inc1b-pr3-build.md.

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