Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 99 additions & 0 deletions .squad/decisions/inbox/mary-scan-1a.md
Original file line number Diff line number Diff line change
@@ -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.
11 changes: 11 additions & 0 deletions crates/onnx-runtime-session/src/executor/build.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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],
};

Expand Down Expand Up @@ -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
Expand Down
166 changes: 110 additions & 56 deletions crates/onnx-runtime-session/src/executor/control_flow.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<Tensor> = (&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() {
Expand All @@ -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<usize>)],
) -> Result<(Vec<Tensor>, Vec<Tensor>)> {
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<Tensor> = (&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<Tensor> = it.collect();
Ok((next_state, scan_outs))
}
}

fn scan_slice(
Expand Down
35 changes: 35 additions & 0 deletions crates/onnx-runtime-session/src/executor/state.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
Loading
Loading