From cbf240af30a44e2108d7e40a5f4333088241cced Mon Sep 17 00:00:00 2001 From: Justin Chu Date: Fri, 31 Jul 2026 15:04:14 +0000 Subject: [PATCH] feat(engine): flag-gated single-trip Scan inline dual-path (slice 1a) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A Scan whose RUNTIME trip_count == 1 (a single decode step) runs its body once straight-line instead of the generic exec_scan loop, while prefill (trip_count = prompt_len > 1) keeps the unchanged loop. Selection is keyed on the observed trip_count at execution time — NOT a graph rewrite — because prefill and decode share one executor/plan, so a static single-trip bake would corrupt prefill. Gated by ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP (default OFF): flag OFF is zero behavior change. Both the loop and the inline path drive the body through a shared run_scan_body_step helper and share the finishing code, so the inline path is byte-exact with a one-iteration loop by construction (DRY, general — no op/model special-casing). Correctness-only foundation; no capture changes (slice 1b will let the inlined body enter CUDA-graph capture). No changes to plan_capture_region / node_capture_reason / StructuralCaptureDecline. Evidence: - CPU test scan_single_trip_inline_is_byte_exact_and_runtime_keyed: byte-exact vs loop over both outputs, engages only at trip_count==1, count==0 on prefill (runtime-keyed). Mutation-checked non-vacuous. - CUDA-gated regression cuda_scan_single_trip_inline_is_byte_exact_and_runtime_keyed (device 4): same on real ORT-CUDA, via scan_inline_single_trip_count. - On-model 27B (qwen3.6-27b int4, device 4, greedy 48 tok): token ids IDENTICAL flag OFF vs ON across prefill + 48 decode steps. - Re-ran #554 (recurrent-state reset) and #544 (weight page-in WAR) green. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .squad/decisions/inbox/mary-scan-1a.md | 99 ++++++++++ .../src/executor/build.rs | 11 ++ .../src/executor/control_flow.rs | 166 +++++++++++------ .../src/executor/state.rs | 35 ++++ .../src/executor/tests.rs | 146 +++++++++++++++ crates/onnx-runtime-session/src/lib.rs | 10 + .../tests/cuda_scan_inline_single_trip.rs | 175 ++++++++++++++++++ 7 files changed, 586 insertions(+), 56 deletions(-) create mode 100644 .squad/decisions/inbox/mary-scan-1a.md create mode 100644 crates/onnx-runtime-session/tests/cuda_scan_inline_single_trip.rs diff --git a/.squad/decisions/inbox/mary-scan-1a.md b/.squad/decisions/inbox/mary-scan-1a.md new file mode 100644 index 0000000000..f3e8ad7ef6 --- /dev/null +++ b/.squad/decisions/inbox/mary-scan-1a.md @@ -0,0 +1,99 @@ +# Decision — Scan single-trip inline dual-path, SLICE 1a (Mary) + +**Date:** 2026-07-31 · **Branch:** `feat/27b-scan-capture-1a` (off origin/main) · **Author:** mary +**Status:** committed, NOT PR'd — awaiting Justin's independent review + open/merge. +**Scope:** correctness-only host-execution dual-path. NO capture changes (that is slice 1b). + +## What this is +The GREEN-LIT Approach-1 **1a** from the PENDING-JUSTIN root-cause: make a `Scan` +whose **runtime** scan-axis length (`trip_count`) is exactly 1 (a single decode +step) execute its body **once, straight-line**, instead of the generic +`exec_scan` loop — while prefill (`trip_count = prompt_len > 1`) keeps the +unchanged loop. Foundation for 1b (letting that inlined body enter CUDA-graph +capture). + +## Mechanism (where the selection happens) +- File: `crates/onnx-runtime-session/src/executor/control_flow.rs`, in `exec_scan` + (after `trip_count`/axes/slices are resolved, right before the iteration loop). +- Branch: `if self.scan_inline_single_trip_enabled && trip_count == 1 { inline } + else { existing loop }`. The condition is evaluated at **execution time** on the + observed `trip_count`, NOT a graph rewrite — this is the whole point: prefill + and decode **share one InferenceSession/executor/plan**, so a static single-trip + bake would corrupt prefill. Runtime keying is the correctness guarantee. +- **DRY:** both the loop and the inline path drive the body through one shared + helper `run_scan_body_step` (run subgraph once → validate output count → + validate carried-state dtype/shape → split next-state / scan-outputs), and both + share the identical finishing code (state store + `TensorStackAccumulator:: + finish_scan`). The inline path is therefore **byte-exact with a one-iteration + loop by construction** — they cannot diverge. No op- or model-name special-casing; + works for ANY single-trip Scan (num_scan_inputs, axes, directions all honored). + +## Flag (default OFF) +- Env: `ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP` — ON only on `1`/`true`/`on` + (case-insensitive, trimmed). Unset/empty/`0`/unrecognized ⇒ OFF. +- Read once at session build (`scan_inline_single_trip_env_enabled()` in + `state.rs`), stored as `Executor::scan_inline_single_trip_enabled`. +- **Flag OFF ⇒ zero behavior change**: every trip_count uses the loop; the only + code delta on that path is that the loop body was factored into + `run_scan_body_step` (behavior-identical, proven by the tests below). + +## Observability (non-vacuity) +- `Executor::scan_inline_single_trip_count` counts every engagement; surfaced as + `InferenceSession::scan_inline_single_trip_count()` (mirrors `decode_memo_counts`). + +## Byte-exact evidence +1. **CPU unit test** (always-on, deterministic) — + `executor::tests::scan_single_trip_inline_is_byte_exact_and_runtime_keyed`: + synthetic multi-node Scan body (Add→Mul→Sub, 2 scan inputs, 1 state + 1 scan + output). Asserts: flag-OFF count==0; flag-ON at trip_count==1 count==1 and + output **byte-identical** over BOTH outputs vs the loop; and at trip_count==3 + (prefill) flag-ON count stays **0** (runtime-keyed, not static) with output == + loop. Mutation-checked non-vacuous: forcing the branch to never engage flips + the count assertion to FAIL (verified: `left: 0, right: 1`). +2. **CUDA-gated regression test** (own binary so no sibling races the env flag) — + `tests/cuda_scan_inline_single_trip.rs:: + cuda_scan_single_trip_inline_is_byte_exact_and_runtime_keyed`: same assertions + on real ORT-CUDA (device 4). PASSED. +3. **On-model 27B** (qwen3.6-27b-int4-cuda, qwen36-conv1d io-overlay, device 4, + greedy, prompt "The history of computing began", 48 tokens, --steady): + token id sequences **IDENTICAL** flag-OFF vs flag-ON, covering prefill + (~790 ms, prompt_len>1) AND 48 single-trip decode steps (48 LinearAttention + Scans/step): + `[303,279,220,16,24,19,15,82,440,279,4257,314,279,1118,13934,17943,11,1680, + 430,279,5025,40,1646,11,864,557,5617,303,220,16,24,19,20,13,4081,3988,17943, + 998,3349,11,11064,11,321,2483,4927,13017,13,4213]`. Throughput ~6.1→5.8 tok/s + (within noise; 1a is host-execution-identical, no capture yet — as expected). + On-model engagement is proven by the counter in test (2); token-identity here + is the end-to-end correctness lock. + +## Regressions re-run (all PASS, device 4) +- #554 session-reuse recurrent-state reset: + `native_cuda_reused_session_rezeros_recurrent_state` ✅ +- #544 async fence-ordered weight page-in: `cuda_prefetch_war:: + drive_double_buffer_war_safe_across_waves` ✅ +- CUDA Scan/Sequence oracle: `cuda_control_flow_safety` ✅ +- Full CPU suites: session lib (105) + control_flow (23) + executor (32) ✅ + +## Files changed +- `executor/control_flow.rs` — runtime dual-path branch + shared + `run_scan_body_step` helper. +- `executor/state.rs` — flag field + counter field + env parser. +- `executor/build.rs` — field init + `scan_inline_single_trip_count()` accessor. +- `lib.rs` — public `scan_inline_single_trip_count()`. +- `executor/tests.rs` — CPU byte-exact + runtime-keyed test. +- `tests/cuda_scan_inline_single_trip.rs` — CUDA-gated regression (new). + +## Contained-slice check +1a stayed contained: NO changes to `provider.rs:plan_capture_region`, +`executor/capture.rs:node_capture_reason`, or any StructuralCaptureDecline logic. +Scan remains structurally declined at the capture seam and runs eager in both +paths — no capture interaction. + +## Slice 1b will add (handoff) +- Let the single-trip inlined body **enter CUDA-graph capture** (fold body nodes + into the parent capture region / grant the trip_count==1 Scan a capture + exemption). Blast radius: `provider.rs:458` + `executor/capture.rs`. +- Validate captures/replays counters RISE and assert 27B tokens byte-identical to + the locked reference (the sequence above is the 1a reference). +- 1a already gives 1b a clean, distinct straight-line code path to recognize; the + `scan_inline_single_trip_count` counter is the engagement tripwire to reuse. diff --git a/crates/onnx-runtime-session/src/executor/build.rs b/crates/onnx-runtime-session/src/executor/build.rs index 7e19730a16..edd3ae5b0a 100644 --- a/crates/onnx-runtime-session/src/executor/build.rs +++ b/crates/onnx-runtime-session/src/executor/build.rs @@ -516,6 +516,8 @@ impl Executor { decode_view_plan_disabled: false, compute_in_place_enabled: compute_in_place_env_enabled(), compute_in_place_alias_count: 0, + scan_inline_single_trip_enabled: scan_inline_single_trip_env_enabled(), + scan_inline_single_trip_count: 0, kernel_bindings: vec![None; plan_len], }; @@ -1353,6 +1355,15 @@ impl Executor { ) } + /// How many times the single-trip `Scan` inline path engaged over this + /// executor's lifetime. `> 0` after a decode run proves the dual-path is + /// non-vacuously firing (an on-model A/B reads this to reject a silently + /// gated-out pass); stays 0 whenever the flag is OFF or every `Scan` runs at + /// `trip_count != 1`. + pub(crate) fn scan_inline_single_trip_count(&self) -> u64 { + self.scan_inline_single_trip_count + } + /// F5 Stage 2 replay guard: every retained view's source buffer must still be /// the identical allocation (same base pointer *and* capacity) it was under /// when the plan was built. A realloc or move — even one that preserves the diff --git a/crates/onnx-runtime-session/src/executor/control_flow.rs b/crates/onnx-runtime-session/src/executor/control_flow.rs index 0f800a9b0a..5ecbba3c46 100644 --- a/crates/onnx-runtime-session/src/executor/control_flow.rs +++ b/crates/onnx-runtime-session/src/executor/control_flow.rs @@ -1287,66 +1287,62 @@ impl Executor { scan_slices.push(Tensor::from_raw(input.dtype, shape, &bytes)?); } } - for step in 0..trip_count { - if step != 0 { - for (index, (((input, &axis), &direction), slice)) in scan_inputs - .iter() - .zip(&input_axes) - .zip(&input_directions) - .zip(scan_slices.iter_mut()) - .enumerate() - { - let source_index = if direction == 0 { - step - } else { - trip_count - 1 - step - }; - let (_, bytes) = scan_slice(input, axis, source_index, index)?; - slice.overwrite_bytes(&bytes)?; - } - } - let mut formal: Vec<&Tensor> = Vec::with_capacity(num_state + num_scan_inputs); - formal.extend(state.iter()); - formal.extend(scan_slices.iter()); - - let outs = self.run_subgraph(&prepared, &formal)?; - drop(formal); - let expected = num_state + num_scan_outputs; - if outs.len() != expected { - return Err(SessionError::OutputShapeCountMismatch { - op: "Scan/body".to_string(), - expected, - got: outs.len(), - }); + // Runtime dual-path (slice 1a). A single-trip Scan — trip_count == 1, the + // per-token decode regime — runs its body ONCE straight-line under the + // opt-in `scan_inline_single_trip_enabled` flag, instead of the generic + // loop. The selection is keyed on the RUNTIME trip_count, never the graph: + // prefill (trip_count = prompt_len > 1) and decode (trip_count == 1) share + // this same executor/plan, so a static single-trip rewrite would corrupt + // prefill. Flag OFF, or any trip_count other than 1, takes the unchanged + // loop path below. Both regimes drive the body through the identical + // `run_scan_body_step` and share the finishing code that follows, so the + // inline path is byte-exact with a one-iteration loop by construction. + if self.scan_inline_single_trip_enabled && trip_count == 1 { + self.scan_inline_single_trip_count += 1; + let (next_state, scan_outs) = self.run_scan_body_step( + &prepared, + &state, + &scan_slices, + num_state, + num_scan_outputs, + &state_specs, + )?; + state = next_state; + for (acc, tensor) in scan_acc.iter_mut().zip(scan_outs) { + acc.push(tensor)?; } - let mut it = outs.into_iter(); - let next_state: Vec = (&mut it).take(num_state).collect(); - for (index, (tensor, (expected_dtype, expected_shape))) in - next_state.iter().zip(&state_specs).enumerate() - { - if tensor.dtype != *expected_dtype { - return Err(SessionError::ControlFlow { - op: "Scan".to_string(), - reason: format!( - "state output {index} dtype mismatch: expected {expected_dtype:?}, got {:?}", - tensor.dtype - ), - }); + } else { + for step in 0..trip_count { + if step != 0 { + for (index, (((input, &axis), &direction), slice)) in scan_inputs + .iter() + .zip(&input_axes) + .zip(&input_directions) + .zip(scan_slices.iter_mut()) + .enumerate() + { + let source_index = if direction == 0 { + step + } else { + trip_count - 1 - step + }; + let (_, bytes) = scan_slice(input, axis, source_index, index)?; + slice.overwrite_bytes(&bytes)?; + } } - if tensor.shape != *expected_shape { - return Err(SessionError::ControlFlow { - op: "Scan".to_string(), - reason: format!( - "state output {index} shape mismatch: expected {expected_shape:?}, got {:?}", - tensor.shape - ), - }); + let (next_state, scan_outs) = self.run_scan_body_step( + &prepared, + &state, + &scan_slices, + num_state, + num_scan_outputs, + &state_specs, + )?; + state = next_state; + for (acc, tensor) in scan_acc.iter_mut().zip(scan_outs) { + acc.push(tensor)?; } } - state = next_state; - for acc in scan_acc.iter_mut() { - acc.push(it.next().expect("scan output present"))?; - } } for (i, t) in state.iter().enumerate() { @@ -1370,6 +1366,64 @@ impl Executor { } Ok(()) } + + /// Run the `Scan` body once for the current formal inputs (carried state + /// followed by this step's scan-input slices), validate the body's output + /// count and the carried-state dtype/shape invariants, and split the result + /// into the next carried state and this step's scan outputs. Shared verbatim + /// by the generic multi-trip loop and the single-trip inline dual-path so the + /// two body-execution regimes can never diverge — the sole guarantee behind + /// slice 1a's byte-exactness. + fn run_scan_body_step( + &mut self, + prepared: &PreparedSubgraph, + state: &[Tensor], + scan_slices: &[Tensor], + num_state: usize, + num_scan_outputs: usize, + state_specs: &[(DataType, Vec)], + ) -> Result<(Vec, Vec)> { + let mut formal: Vec<&Tensor> = Vec::with_capacity(state.len() + scan_slices.len()); + formal.extend(state.iter()); + formal.extend(scan_slices.iter()); + + let outs = self.run_subgraph(prepared, &formal)?; + drop(formal); + let expected = num_state + num_scan_outputs; + if outs.len() != expected { + return Err(SessionError::OutputShapeCountMismatch { + op: "Scan/body".to_string(), + expected, + got: outs.len(), + }); + } + let mut it = outs.into_iter(); + let next_state: Vec = (&mut it).take(num_state).collect(); + for (index, (tensor, (expected_dtype, expected_shape))) in + next_state.iter().zip(state_specs).enumerate() + { + if tensor.dtype != *expected_dtype { + return Err(SessionError::ControlFlow { + op: "Scan".to_string(), + reason: format!( + "state output {index} dtype mismatch: expected {expected_dtype:?}, got {:?}", + tensor.dtype + ), + }); + } + if tensor.shape != *expected_shape { + return Err(SessionError::ControlFlow { + op: "Scan".to_string(), + reason: format!( + "state output {index} shape mismatch: expected {expected_shape:?}, got {:?}", + tensor.shape + ), + }); + } + } + let scan_outs: Vec = it.collect(); + Ok((next_state, scan_outs)) + } } fn scan_slice( diff --git a/crates/onnx-runtime-session/src/executor/state.rs b/crates/onnx-runtime-session/src/executor/state.rs index 0f1ad15f59..16631562e7 100644 --- a/crates/onnx-runtime-session/src/executor/state.rs +++ b/crates/onnx-runtime-session/src/executor/state.rs @@ -232,6 +232,23 @@ pub(crate) struct Executor { pub(super) compute_in_place_enabled: bool, /// Successful dead-input buffer aliases, retained for parity/safety tests. pub(super) compute_in_place_alias_count: u64, + /// Opt-in (default OFF) master switch for the single-trip `Scan` inline + /// dual-path (`ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP`). When ON, a `Scan` whose + /// runtime scan-axis length is exactly 1 (a single decode step) runs its body + /// once straight-line instead of the generic `exec_scan` loop; any other + /// trip count — including prefill at `prompt_len > 1` — keeps the unchanged + /// loop. The selection is at RUNTIME, keyed on the observed trip count, never + /// baked into the graph: prefill and decode share one executor/plan, so a + /// static single-trip rewrite would corrupt prefill. Flag OFF ⇒ every trip + /// count uses the loop (zero behavior change). Slice 1a is host-execution + /// only; it does not interact with device-graph capture. + pub(super) scan_inline_single_trip_enabled: bool, + /// Diagnostic: how many times the single-trip `Scan` inline path actually + /// engaged over this executor's lifetime. `> 0` after a decode run proves the + /// dual-path is non-vacuously firing (an on-model A/B and the CUDA-gated + /// regression test read it to reject a silently-gated-out pass); it stays 0 + /// whenever the flag is OFF or every `Scan` runs at `trip_count != 1`. + pub(super) scan_inline_single_trip_count: u64, /// Per-plan-node kernel pre-binding (Stage 3). Each slot stores the /// [`KernelKey`] from the most recent successful kernel lookup for that plan /// node. On subsequent dispatch, if the current input shapes match the stored @@ -525,6 +542,24 @@ pub(super) fn decode_memo_verify_env_enabled() -> bool { ) } +/// Whether the single-trip `Scan` inline dual-path is enabled +/// (`ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP`). Default OFF: this is an opt-in +/// correctness-foundation path, so it engages only on an explicit ON value +/// (`1`/`true`/`on`, case-insensitive, whitespace-trimmed). Every other state — +/// unset, empty, `0`, or unrecognized — leaves it OFF, so the executor keeps +/// running the unchanged `exec_scan` loop for every trip count. +pub(super) fn scan_inline_single_trip_env_enabled() -> bool { + matches!( + std::env::var("ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP") + .ok() + .as_deref() + .map(str::trim) + .map(str::to_ascii_lowercase) + .as_deref(), + Some("1") | Some("true") | Some("on") + ) +} + /// Per-input geometry the run loop resolves once per node: the raw base pointer /// of the backing (root) buffer plus the real view (shape, element strides — /// possibly non-contiguous or negative — and byte offset) to read it through. diff --git a/crates/onnx-runtime-session/src/executor/tests.rs b/crates/onnx-runtime-session/src/executor/tests.rs index 61445463f4..6a7d2d9520 100644 --- a/crates/onnx-runtime-session/src/executor/tests.rs +++ b/crates/onnx-runtime-session/src/executor/tests.rs @@ -3312,3 +3312,149 @@ fn kernel_prebinding_fallback_fires_on_shape_change() { "after shape change, the updated binding must serve the fast path" ); } + +/// Build a parent graph with a single `Scan` over a **multi-node** body so the +/// single-trip inline dual-path and the generic loop are exercised on identical +/// non-trivial work. `steps` is the scan-axis length (`1` = a decode step; `>1` +/// = a prefill-shaped run). The body threads carried state through two scan +/// inputs across three ops and emits one carried-state output plus one +/// per-iteration scan output: +/// +/// `state_x = Add(state, x)` +/// `state_out = Mul(state_x, y)` (next carried state) +/// `scan_out = Sub(state_out, x)` (stacked on the scan axis) +fn scan_inline_test_graph(steps: usize) -> Graph { + const W: usize = 3; + + let mut body = Graph::new(); + body.opset_imports.insert(String::new(), 17); + let state = body.create_named_value("state", DataType::Float32, static_shape([W])); + let x = body.create_named_value("x", DataType::Float32, static_shape([W])); + let y = body.create_named_value("y", DataType::Float32, static_shape([W])); + body.add_input(state); + body.add_input(x); + body.add_input(y); + let state_x = body.create_named_value("state_x", DataType::Float32, static_shape([W])); + body.insert_node(Node::new( + NodeId(0), + "Add", + vec![Some(state), Some(x)], + vec![state_x], + )); + let state_out = body.create_named_value("state_out", DataType::Float32, static_shape([W])); + body.insert_node(Node::new( + NodeId(0), + "Mul", + vec![Some(state_x), Some(y)], + vec![state_out], + )); + let scan_out = body.create_named_value("scan_out", DataType::Float32, static_shape([W])); + body.insert_node(Node::new( + NodeId(0), + "Sub", + vec![Some(state_out), Some(x)], + vec![scan_out], + )); + body.add_output(state_out); + body.add_output(scan_out); + + let mut graph = Graph::new(); + graph.opset_imports.insert(String::new(), 17); + let initial = init_inline(&mut graph, "initial", &[W], vec![0.0; W]); + let x_in = graph.create_named_value("X", DataType::Float32, static_shape([steps, W])); + let y_in = graph.create_named_value("Y", DataType::Float32, static_shape([steps, W])); + graph.add_input(x_in); + graph.add_input(y_in); + let final_state = graph.create_named_value("final_state", DataType::Float32, static_shape([W])); + let scan_output = + graph.create_named_value("scan_output", DataType::Float32, static_shape([steps, W])); + let mut scan = Node::new( + NodeId(0), + "Scan", + vec![Some(initial), Some(x_in), Some(y_in)], + vec![final_state, scan_output], + ); + scan.attributes + .insert("num_scan_inputs".to_string(), Attribute::Int(2)); + let scan_id = graph.insert_node(scan); + graph.subgraphs.insert((scan_id, "body".to_string()), body); + graph.add_output(final_state); + graph.add_output(scan_output); + graph +} + +fn init_inline(graph: &mut Graph, name: &str, dims: &[usize], data: Vec) -> ValueId { + use onnx_runtime_ir::{TensorData, WeightRef}; + let bytes: Vec = data.iter().flat_map(|v| v.to_le_bytes()).collect(); + let value = + graph.create_named_value(name, DataType::Float32, static_shape(dims.iter().copied())); + graph.set_initializer( + value, + WeightRef::Inline(TensorData::from_raw( + DataType::Float32, + dims.to_vec(), + bytes, + )), + ); + value +} + +fn run_scan_inline_graph(steps: usize, inline: bool) -> (Vec>, u64) { + let mut exec = Executor::build( + scan_inline_test_graph(steps), + Arc::new(WeightStore::new()), + auto_detect_cpu_ep().unwrap(), + ) + .unwrap(); + exec.scan_inline_single_trip_enabled = inline; + + let n = steps * 3; + let x: Vec = (0..n).map(|i| (i as f32) + 1.0).collect(); + let y: Vec = (0..n).map(|i| (i as f32) * 0.5 + 2.0).collect(); + let x_t = Tensor::from_f32(&[steps, 3], &x).unwrap(); + let y_t = Tensor::from_f32(&[steps, 3], &y).unwrap(); + let outputs = exec.run(&[("X", &x_t), ("Y", &y_t)]).unwrap(); + let bytes = outputs.iter().map(|t| t.as_bytes().to_vec()).collect(); + (bytes, exec.scan_inline_single_trip_count()) +} + +/// Slice-1a correctness gate. Proves the flag-gated single-trip `Scan` inline +/// dual-path is (1) **byte-exact** with the generic `exec_scan` loop, and (2) +/// **non-vacuously engaged** and **runtime-keyed** — engaging only at +/// `trip_count == 1` and never on a prefill-shaped (`trip_count > 1`) run, even +/// with the flag ON. The shared-plan tripwire: a static single-trip rewrite +/// would fire on prefill too; this asserts it does not. Byte-equality is checked +/// over BOTH the carried `final_state` and the stacked `scan_output`, so a wrong +/// inline path (dropped/duplicated body run, mis-stacked scan axis, or skipped +/// state thread) makes the test FAIL. +#[test] +fn scan_single_trip_inline_is_byte_exact_and_runtime_keyed() { + // Decode regime (trip_count == 1): inline path engages exactly once and is + // byte-identical to the loop over every output. + let (loop_out, loop_count) = run_scan_inline_graph(1, false); + let (inline_out, inline_count) = run_scan_inline_graph(1, true); + assert_eq!(loop_count, 0, "flag OFF must never engage the inline path"); + assert_eq!( + inline_count, 1, + "flag ON at trip_count==1 must engage the inline path exactly once" + ); + assert_eq!( + inline_out, loop_out, + "single-trip inline output must be byte-exact with the loop path" + ); + + // Prefill regime (trip_count == 3): even with the flag ON the inline path + // must NOT engage (runtime-keyed, not a static rewrite), and the output must + // still match the loop. + let (prefill_loop, prefill_loop_count) = run_scan_inline_graph(3, false); + let (prefill_inline, prefill_inline_count) = run_scan_inline_graph(3, true); + assert_eq!(prefill_loop_count, 0, "loop path never counts"); + assert_eq!( + prefill_inline_count, 0, + "flag ON must NOT inline a prefill (trip_count>1) Scan — the shared-plan tripwire" + ); + assert_eq!( + prefill_inline, prefill_loop, + "prefill output must be identical flag-on vs flag-off" + ); +} diff --git a/crates/onnx-runtime-session/src/lib.rs b/crates/onnx-runtime-session/src/lib.rs index 22588f5062..5e7944a20b 100644 --- a/crates/onnx-runtime-session/src/lib.rs +++ b/crates/onnx-runtime-session/src/lib.rs @@ -1055,6 +1055,16 @@ impl InferenceSession { self.exec.decode_view_plan_counts() } + /// How many times the single-trip `Scan` inline dual-path + /// (`ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP`) engaged over this session's + /// lifetime. `> 0` after a decode run proves the runtime `trip_count == 1` + /// inline path actually fired (not a silently gated-out pass); an on-model + /// flag-on/flag-off A/B reads this alongside the token stream to prove the + /// dual-path is both engaged and byte-exact. + pub fn scan_inline_single_trip_count(&self) -> u64 { + self.exec.scan_inline_single_trip_count() + } + /// Run with persistent device allocations supplying graph inputs and, /// optionally, aliasing graph outputs. Bound outputs are returned as `None` /// because their bytes remain resident in the caller-owned allocation. diff --git a/crates/onnx-runtime-session/tests/cuda_scan_inline_single_trip.rs b/crates/onnx-runtime-session/tests/cuda_scan_inline_single_trip.rs new file mode 100644 index 0000000000..1363a79376 --- /dev/null +++ b/crates/onnx-runtime-session/tests/cuda_scan_inline_single_trip.rs @@ -0,0 +1,175 @@ +//! CUDA-gated regression for slice 1a: the flag-gated single-trip `Scan` inline +//! dual-path must engage on device at `trip_count == 1`, stay byte-exact with +//! the generic `exec_scan` loop, and never inline a prefill-shaped +//! (`trip_count > 1`) run — the shared-plan tripwire. This lives in its own test +//! binary so no sibling test races the process-global +//! `ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP` env flag it toggles at session build. +#![cfg(feature = "cuda")] + +use std::sync::{Mutex, OnceLock}; + +use onnx_runtime_ir::{ + Attribute, DataType, Graph, Node, NodeId, TensorData, ValueId, WeightRef, static_shape, +}; +use onnx_runtime_loader::{Model, encode_model}; +use onnx_runtime_session::{DevicePreference, InferenceSession, Tensor}; + +const W: usize = 3; + +fn f32_bytes(data: &[f32]) -> Vec { + data.iter().flat_map(|v| v.to_le_bytes()).collect() +} + +fn init(graph: &mut Graph, name: &str, dims: &[usize], data: &[f32]) -> ValueId { + let value = + graph.create_named_value(name, DataType::Float32, static_shape(dims.iter().copied())); + graph.set_initializer( + value, + WeightRef::Inline(TensorData::from_raw( + DataType::Float32, + dims.to_vec(), + f32_bytes(data), + )), + ); + value +} + +fn body_op(body: &mut Graph, op_type: &str, inputs: &[ValueId], name: &str) -> ValueId { + let out = body.create_named_value(name, DataType::Float32, static_shape([W])); + body.insert_node(Node::new( + NodeId(0), + op_type, + inputs.iter().copied().map(Some).collect(), + vec![out], + )); + out +} + +/// A multi-node `Scan` body threading carried state through two scan inputs: +/// `state_x = Add(state, x)` +/// `state_out = Mul(state_x, y)` (next carried state) +/// `scan_out = Sub(state_out, x)` (stacked on the scan axis) +fn scan_body() -> Graph { + let mut body = Graph::new(); + body.opset_imports.insert(String::new(), 17); + let state = body.create_named_value("state", DataType::Float32, static_shape([W])); + let x = body.create_named_value("x", DataType::Float32, static_shape([W])); + let y = body.create_named_value("y", DataType::Float32, static_shape([W])); + body.add_input(state); + body.add_input(x); + body.add_input(y); + let state_x = body_op(&mut body, "Add", &[state, x], "state_x"); + let state_out = body_op(&mut body, "Mul", &[state_x, y], "state_out"); + let scan_out = body_op(&mut body, "Sub", &[state_out, x], "scan_out"); + body.add_output(state_out); + body.add_output(scan_out); + body +} + +fn scan_model(steps: usize) -> Vec { + let mut graph = Graph::new(); + graph.opset_imports.insert(String::new(), 17); + let initial = init(&mut graph, "initial", &[W], &[0.0; W]); + let x_in = graph.create_named_value("X", DataType::Float32, static_shape([steps, W])); + let y_in = graph.create_named_value("Y", DataType::Float32, static_shape([steps, W])); + graph.add_input(x_in); + graph.add_input(y_in); + let final_state = graph.create_named_value("final_state", DataType::Float32, static_shape([W])); + let scan_output = + graph.create_named_value("scan_output", DataType::Float32, static_shape([steps, W])); + let body = scan_body(); + let mut scan = Node::new( + NodeId(0), + "Scan", + vec![Some(initial), Some(x_in), Some(y_in)], + vec![final_state, scan_output], + ); + scan.attributes + .insert("num_scan_inputs".into(), Attribute::Int(2)); + scan.attributes + .insert("body".into(), Attribute::Graph(Box::new(body.clone()))); + let scan_id = graph.insert_node(scan); + graph.subgraphs.insert((scan_id, "body".into()), body); + graph.add_output(final_state); + graph.add_output(scan_output); + encode_model(&Model::new(&graph)).expect("encode scan model") +} + +/// Serializes the env-flag mutation so the CUDA runtime's own session builds +/// (this binary only) never observe a torn flag value. +fn env_lock() -> &'static Mutex<()> { + static ENV_LOCK: OnceLock> = OnceLock::new(); + ENV_LOCK.get_or_init(|| Mutex::new(())) +} + +/// Build a CUDA session with the inline flag forced to `inline` for the duration +/// of the build (it is read once at build), run the feeds, and return the raw +/// output bytes plus the inline-engagement counter. +fn run_cuda(bytes: &[u8], feeds: &[(&str, &Tensor)], inline: bool) -> (Vec>, u64) { + let _guard = env_lock().lock().expect("env lock"); + let prev = std::env::var_os("ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP"); + // SAFETY: all mutations of this var in this binary are serialized by env_lock. + unsafe { + std::env::set_var( + "ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP", + if inline { "1" } else { "0" }, + ); + } + let mut session = InferenceSession::builder() + .model_bytes(bytes) + .device(DevicePreference::Gpu { index: Some(0) }) + .build() + .expect("build CUDA session"); + // SAFETY: serialized by env_lock; restore the prior value now that the flag + // has been consumed at build. + unsafe { + match prev { + Some(v) => std::env::set_var("ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP", v), + None => std::env::remove_var("ONNX_GENAI_SCAN_INLINE_SINGLE_TRIP"), + } + } + let outputs = session.run(feeds).expect("run CUDA session"); + let out_bytes = outputs.iter().map(|t| t.as_bytes().to_vec()).collect(); + (out_bytes, session.scan_inline_single_trip_count()) +} + +#[test] +fn cuda_scan_single_trip_inline_is_byte_exact_and_runtime_keyed() { + // Decode regime: trip_count == 1. Flag ON engages the inline path exactly + // once and is byte-identical to the flag-OFF loop over both outputs. + let decode = scan_model(1); + let x1 = Tensor::from_f32(&[1, W], &[1.0, 2.0, 3.0]).unwrap(); + let y1 = Tensor::from_f32(&[1, W], &[2.0, 2.5, 3.0]).unwrap(); + let feeds1 = [("X", &x1), ("Y", &y1)]; + + let (loop_out, loop_count) = run_cuda(&decode, &feeds1, false); + let (inline_out, inline_count) = run_cuda(&decode, &feeds1, true); + assert_eq!(loop_count, 0, "flag OFF must never engage the inline path"); + assert_eq!( + inline_count, 1, + "flag ON at trip_count==1 must engage the inline path exactly once on CUDA" + ); + assert_eq!( + inline_out, loop_out, + "single-trip inline output must be byte-exact with the loop path on CUDA" + ); + + // Prefill regime: trip_count == 3. Even with the flag ON the inline path must + // NOT engage (runtime-keyed, not a static rewrite), and output must match. + let prefill = scan_model(3); + let x3 = Tensor::from_f32(&[3, W], &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]).unwrap(); + let y3 = Tensor::from_f32(&[3, W], &[2.0, 2.5, 3.0, 3.5, 4.0, 4.5, 5.0, 5.5, 6.0]).unwrap(); + let feeds3 = [("X", &x3), ("Y", &y3)]; + + let (prefill_loop, prefill_loop_count) = run_cuda(&prefill, &feeds3, false); + let (prefill_inline, prefill_inline_count) = run_cuda(&prefill, &feeds3, true); + assert_eq!(prefill_loop_count, 0, "loop path never counts"); + assert_eq!( + prefill_inline_count, 0, + "flag ON must NOT inline a prefill (trip_count>1) Scan on CUDA — shared-plan tripwire" + ); + assert_eq!( + prefill_inline, prefill_loop, + "prefill output must be identical flag-on vs flag-off on CUDA" + ); +}