Skip to content

feat(engine): a chained speculative proposal computes where its tensors already are - #1861

Merged
justinchuby merged 3 commits into
mainfrom
justinchuby/speculative-cuda-residency
Aug 23, 2026
Merged

justinchuby merged 3 commits into
mainfrom
justinchuby/speculative-cuda-residency

Conversation

@justinchuby

Copy link
Copy Markdown
Owner

What

A chained speculative proposal now does its tensor algebra where its tensors already are. On a native-CUDA run this removes the last host round trips in the workflow: the whole-KV download-and-re-upload on every rejection, the folded carry brought down so concat(embed(token), carry) could be built in host memory, and the full logits row downloaded per draft token to learn one token id.

Builds on merged #1723 (zero-copy rollback of contiguous cells, device_ops scaffolding, the native_workflow_parity staging ratchet), #1853 and #1858.

Why the ratchet existed, and why it is gone

#1723 landed a measured ratchet of 16 device→host materializations across speculative_decode(8, 4) on the hermetic gemma4_chained package. Measuring where they came from (backtrace on materialize_workflow_value_copy, origin/main 030a4c019) shows all 16 are rollback_speculative_state's strided host fallback: 4 rejections × 4 declared [1, 2, seq, 8] KV cells, each downloaded whole to drop a few positions and re-uploaded on the next bind. That is the single largest transfer in the workflow, and it is now a device-to-device narrowing.

Design — one seam, no special cases

crates/onnx-genai-engine/src/pipeline/device_ops.rs:

pub(crate) trait ResidentTensorOps {
    fn residency(&self) -> Residency;
    fn adopt(&self, value: &Value) -> Result<Value>;              // the ONE sanctioned crossing
    fn zeros(&self, shape: &[i64], dtype: DataType) -> Result<Value>;
    fn slice_axis(&self, v: &Value, axis, start, len) -> Result<Value>;
    fn gather_rows(&self, table: &Value, ids: &[i64]) -> Result<Value>;
    fn scatter_into_last_axis(&self, dst: &Value, offset, src: &Value) -> Result<()>;
    fn argmax_rows(&self, logits: &Value, rows: usize) -> Result<Vec<u32>>;
    // last_along_axis / truncate_axis are the two windows of slice_axis
}
  • slice_axis is the only narrowing implemented. last_along_axis and truncate_axis are windows of it, so the two paths a proposal chain uses cannot drift from the primitive (Rule 10).
  • Residency is preserved or the call fails. A device value in a build with no device implementation is an error naming both remedies — never a quiet copy. An operation whose operands straddle residencies is refused (Rule 4).
  • adopt is the only crossing, stated once per proposal for the carry seed the target left elsewhere — never per draft token.
  • Zero-copy where it is real. A contiguous window is Value::alias_with_offset, whose owner keeps the device allocation alive; a strided one is device-to-device copies. Neither crosses the bus.

Two supporting facts had to become explicit:

  • The chain's residency is where the proposer executes (WorkflowRuntime::component_execution_residency) — configured, not discovered — so the fused input is built in the right memory from the first step instead of being uploaded per step.
  • invoke_component_values takes the shape symbols a proposer's outputs declare and its inputs do not (a vocabulary, above all), proven from the workflow's own bound values for the same component in the same run. Without them the run cannot size a device buffer for those outputs and hands each one back through host memory — a full logits download per draft token. A hint contradicting an input's actual extent is a validation error, not an override; a symbol two outputs disagree about is dropped rather than guessed.

Value gains empty_cuda / fill_zero_device (an owning cudaMalloc allocation whose guard travels with the value), as_slice_f32 (borrow a [vocab, hidden] table rather than copy it) and write_raw_bytes_at (dtype-agnostic host scatter). cuda_rt grows CudaAllocation over cudaMalloc/cudaFree. Value::alias_with_offset already existed from #1723.

Rule 3: WorkflowRuntime::host_copy_of had no callers left and is deleted.

H200 evidence — gemma4_chained, native-CUDA, speculative_decode(8, 4)

before (030a4c019) after
device→host materializations 16 0
device read-back full logits row per draft token 64 B = 4 B × 16 invocations, nothing else
device input bindings (delta) 0 16 — one fused input per proposer step, bound zero-copy
device outputs kept resident (delta) 32 64 — +32 = 16 × (logits, folded carry)
embedding-table artifact reads 1 1
proposer invocations 16 16
committed tokens [30, 5, 2, 7, 30, 5, 2, 7] identical
tally proposed 16, accepted 4, rejections 4, rolled-back cells 16 identical

Free device memory is flat across 100 proposals after warm-up (repeated_proposals_do_not_retain_device_memory_native_cuda), which is what proves the aliases a step holds release the buffers they borrow.

Tests

  • native_workflow_parity::chained_proposal_stays_device_resident_native_cuda — zero materializations, exactly four bytes back per proposer invocation, ≥ 2 device-resident proposer outputs per invocation, every rolled-back cell still on device 0 at the truncated length.
  • native_workflow_parity::repeated_proposals_do_not_retain_device_memory_native_cuda — device memory flat.
  • native_workflow_parity::chained_speculative_proposal_parity{,_native_cuda} — ratchet is now assert_eq!(staged, 0); read-back held to a 4-byte-per-invocation budget; embedding table read exactly once.
  • device_ops::cuda_tests — differential against the host implementation: every slice (contiguous and strided, every axis and window of four shapes), gather, scatter, adoption, and argmax over ties / NaN / −inf rows.
  • device_ops::tests, speculative::rollback_residency_tests — host semantics, the narrowing-primitive identity, and every fail-closed rejection.

Suites run (H200, --features native-cuda)

native_workflow_parity (14), native_workflow_smoke (2), gemma4_chained_workflow (9), chained_proposer_real (2), canonical_execution_parity (7), workflow_policy_e2e (20), static_cache_workflow (7), tree_speculative (2), onnx-genai-engine --lib (641 pass), onnx-genai-ort --lib (190 pass). cargo fmt --all; cargo clippy --all-targets clean for onnx-genai-engine and onnx-genai-ort under native-cuda, native-backend, and no-default-features.

Two failures reproduce unchanged on origin/main 030a4c019 and are not from this branch: engine::runtime::tests::native_generate_rejects_over_kv_byte_budget_before_backend_run and loader::tests::selected_non_dense_candidate_fails_explicitly.

Compatibility

HostTensorOps is the semantics the driver always had, so CPU and ORT runs are byte-identical — gemma4_chained_workflow's 9 cases and the ORT half of every parity case are unchanged, and a_proposal_chain_stages_nothing_of_its_own still reports zero. Nothing here is keyed on a model, a port spelling, or a tensor shape: the sequence axis still comes from the declared serving.state_service.groups[..].sequence_axis, the fused-input split from the declared port widths, and the embedding table from the contract's {component, table}.

No changes to any continuous-batch or speculative interpreter dispatch file.

…rs already are

A proposal chain's tensors are produced by components running on a device, and
the interpreter was reading them back to operate on them. On a native-CUDA run
that meant, per proposal: the whole KV cache downloaded and re-uploaded on every
rejection to drop a few positions; the folded carry brought down so
`concat(embed(token), carry)` could be built in host memory and pushed back up;
and the full logits row downloaded per draft token to learn one token id.

`pipeline/device_ops.rs` makes "operate where the value is" expressible. One
trait — `slice_axis` with `last_along_axis`/`truncate_axis` as its two windows,
plus `gather_rows`, `scatter_into_last_axis`, `zeros`, `argmax_rows` and the
single sanctioned crossing `adopt` — with a host and a CUDA implementation that
must agree element for element. Every method returns a value in its principal
input's residency or fails naming both remedies; a value on a device this build
has no implementation for is an error, never a quiet copy to the host, and an
operation whose operands straddle residencies is refused rather than staged.

The chain now:

* truncates a rejected proposal's declared state cells on the device — a
  contiguous prefix is a pointer view and a strided one is device-to-device
  copies, so the largest transfer in the workflow disappears;
* keeps the folded carry seed and every step's carry where the component left
  them, adopting the seed once per proposal rather than per draft token;
* assembles the fused input into one device buffer allocated per proposal,
  gathering the embedding row and scattering both halves in place;
* argmaxes on the device, reading back four bytes per step;
* reads the declared embedding table out of the artifact once and mirrors it
  into the chain's residency once, for the runtime's life.

Two supporting facts had to become explicit. The chain's residency is where the
proposer *executes* (`component_execution_residency`), which is configured
rather than discovered, so the fused input is built in the right memory from the
first step. And `invoke_component_values` accepts the shape symbols a
proposer's outputs declare but its inputs do not — a vocabulary, above all —
proven from the workflow's own bound values: an output whose shape cannot be
resolved has no device buffer sized for it and comes back through host memory.
A hint contradicting an input's actual extent is a validation error, not an
override.

`Value` gains the primitives this needs: `empty_cuda`/`fill_zero_device` (an
owning `cudaMalloc` allocation whose guard travels with the value),
`as_slice_f32` (borrow a `[vocab, hidden]` table instead of copying it), and
`write_raw_bytes_at` (the dtype-agnostic host scatter). `cuda_rt` grows
`CudaAllocation` over `cudaMalloc`/`cudaFree`.

Measured on the hermetic `gemma4_chained` package, native-CUDA on an H200,
across `speculative_decode(8, 4)`:

* device→host materializations: 16 -> 0;
* device read-back: exactly 4 bytes per proposer invocation (64 bytes for 16),
  and nothing else;
* every rolled-back state cell still resident on device 0 at the truncated
  length;
* free device memory flat across 100 proposals;
* tokens, acceptance tally and greedy reference unchanged on ORT, native-CPU
  and native-CUDA.

`device_ops::cuda_tests` is the differential layer beneath that: every device
slice (contiguous and strided, every axis and window of four shapes), gather,
scatter, adoption and argmax — ties and NaN included — must equal what the host
implementation produces.

Rule 3: `host_copy_of` had no callers left and is gone.

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

codecov Bot commented Aug 23, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 72.56%. Comparing base (030a4c0) to head (10cb418).
⚠️ Report is 9 commits behind head on main.

❗ There is a different number of reports uploaded between BASE (030a4c0) and HEAD (10cb418). Click for more details.

HEAD has 4 uploads less than BASE
Flag BASE (030a4c0) HEAD (10cb418)
offline 3 0
mlas 1 0
Additional details and impacted files

Impacted file tree graph

@@             Coverage Diff             @@
##             main    #1861       +/-   ##
===========================================
- Coverage   80.29%   72.56%    -7.73%     
===========================================
  Files         414       12      -402     
  Lines      203362     5231   -198131     
  Branches   203362     5231   -198131     
===========================================
- Hits       163290     3796   -159494     
+ Misses      34518     1307    -33211     
+ Partials     5554      128     -5426     
Flag Coverage Δ
cli-ort-linux 72.51% <ø> (ø)
cli-ort-windows 72.10% <ø> (+0.09%) ⬆️
mlas ?
offline ?

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

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.
  • 📦 JS Bundle Analysis: Save yourself from yourself by tracking and limiting bundle sizes in JS merges.

Copilot AI added 2 commits August 23, 2026 17:52
…vider's stream

Independent review found the seam's device-to-device copies unordered with
respect to the kernels that read what they write, plus five smaller defects.

**The copies were never fenced (correctness).** `CudaTensorOps` issues
`cudaMemcpy` with `cudaMemcpyDeviceToDevice`, which does no host-side
synchronization and is enqueued on cudart's legacy default stream. Both CUDA
execution providers run kernels on streams created `cudaStreamNonBlocking`,
which are exempt from that stream's implicit ordering — so a proposer step could
launch against a fused input this seam was still assembling, and, worse, a
target run could read a KV cell a rejection was still truncating. That is a
silent wrong answer proportional to how much of the copy is outstanding: never
on a 128-byte fixture, megabytes on a real cache. Every write batch now ends in
one `fence()`, the same barrier and the same reasoning that already bracket the
shared-KV grow copy. Destinations are allocated through `zeroed_unfenced` so a
row-wise gather still costs one barrier, not one per row. The other direction
needed nothing: every device value entering the seam was published by the
component boundary, which drains the producing stream first.

**The per-step carry used the chain's ops instead of its own.** Every other site
derives operations from the value; this one asked the chain-wide residency, so a
backend that publishes an output somewhere other than where the chain assembles
— an ORT session on a CUDA device — would have hit a host implementation with a
device value and failed with a misleading remedy. It now narrows where the value
is and adopts once, exactly as the carry seed does.

**"Only token ids cross the bus" was not quite true.** Adoption is the seam's
other crossing, and `HostTensorOps::adopt` counted nothing. `adopt_into` now
records what it brings back, so `device_readback_bytes` plus
`host_staging_count` genuinely account for every returning byte.

**An empty result's residency was path-dependent.** A zero-length slice took the
alias branch (device) or the fallback (host) depending only on whether the
leading extents were 1. Emptiness is now decided first, and the module documents
the exception: a result with no elements has no device address to publish, which
is how the component boundary already publishes an empty output.

**The CUDA tests gated themselves on the API under test.** `empty_cuda`
regressing would have made all seven pass silently; the probe is now
`device_memory_info`.

**The leak test asserted on a device-global reading.** `cudaMemGetInfo` moves
under you on a shared GPU. `cuda_rt::live_allocations()` counts this process's
own `CudaAllocation`s, which is exactly attributable; free memory is still
checked, but as a bound with a stated noise tolerance.

Also: `component_execution_residency` returns `Result` and names an unusable
ordinal rather than degrading a device chain to a host one in silence;
`write_raw_bytes_at` uses `ptr::copy`, because nothing in its signature stops a
caller passing bytes borrowed from a value that aliases the destination;
`CudaAllocation`'s `unsafe impl Send/Sync` is deleted (all fields are integers,
so the compiler already grants both, and the assertion would only have kept them
true by fiat); and the adopt/hint error messages name the shape, dtype and
device they are about.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Signed-off-by: Copilot <223556219+Copilot@users.noreply.github.com>
A second review pass confirmed the fence closes the stream-ordering hazard and
found six smaller things, all fixed here.

**The leak test could be moved by its own siblings.** `live_allocations()` is
process-global, and three other `native-cuda` tests in `native_workflow_parity`
build and drop CUDA engines — each holding an embedding mirror — while it loops.
It passed by a fraction of a second, and one more CUDA test in that binary would
have broken it deterministically. Cargo gives each integration test file its own
process, so the measurement moved to `chained_proposal_device_memory`, a file of
one test, with the reason stated in its module doc. The free-bytes backstop is
now a coarse check on retention the count cannot see, with the tolerance
tightened to 64 MiB and its role stated instead of implied.

**The embedding-table upload was ordered by an unstated invariant.** It happened
to be safe because its only consumer is a copy on the same legacy stream that
*is* fenced. It now goes through the seam's own `adopt`, so it is ordered by the
same fence as every other device write here — and the hand-rolled allocate-plus-
copy it replaces was a second implementation of adoption (Rule 10).

**The public counter's documentation was stale.** `Engine::device_readback_bytes`
and its runtime counterpart still said "the token id a device argmax produces";
adoption bytes have counted since the previous commit. A device-resident chain
still spends only four bytes per step, and the docs now say both things.

**`zeros` did not obey the module's own empty-result exception.** A zero-element
shape reached `Value::empty_cuda`, which refuses it. It now returns the same
host-empty tensor every other branch does, so the rule holds everywhere rather
than in most places.

**The fence's cost is now stated where the fence is.** It is per operation, two
to three more per step than correctness needs. Hoisting it to the caller was
considered and rejected: a fence a caller can forget buys microseconds and pays
for them in a bug that presents as wrong output on large models and as nothing
at all on the fixtures. Making it free means issuing the copies on the
provider's own stream, which needs a stream handle this seam does not have; that
is written down rather than left for the next reader to rediscover.

Also: the architecture doc describes what the memory test now asserts.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Signed-off-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby
justinchuby merged commit 763ae30 into main Aug 23, 2026
7 of 15 checks passed
@justinchuby
justinchuby deleted the justinchuby/speculative-cuda-residency branch August 23, 2026 18:40
justinchuby added a commit that referenced this pull request Aug 23, 2026
The hermetic `gemma4_chained` fixture proves the interpreter's chained-proposal
semantics on a 32-token vocabulary whose target is a lookup table. It cannot
prove that a *published* package decodes, and that gap is where the real
failures were: every one of them passes the fixture.

Driving the published Gemma4-E2B speculative pair through the interpreter on an
H200 found five, four of them here and one in the packages themselves. Each is
fixed at the layer that owns it.

**The serving contract's own controls were demanded of the caller.**
`serving.{active, done, accepted_len}` names values the runtime's token policy
writes; the packages declare them as required application inputs with no
default, so admission refused every request -- asking an application whether the
row it just submitted is already `done`. Package admission now seeds a workflow
input that a serving control resolves to from that control's declared role
(`active = true`, `done = false`, `accepted_len = 0`), exactly as it already
does for a runtime-managed state seed. A caller-supplied value still wins, a
package-declared default still wins over the derived one, and an input the
serving contract does not name stays as required as it was. Required-input
validation is not weakened and nothing is model-specific.

**The chained proposal path was float32-only.**
`token_embedding.table` had to be float32 and the fused proposer input was
allocated float32, so a real fp16 export could not start a chain. The chain now
takes its arithmetic currency from the fused input's *declared* contract and
requires the declared table to agree, naming both sides when it does not. The
currency is the element type of the value the workflow bound to that port.
`EmbeddingTable` holds the initializer in its own element type; `row()` widens
one row rather than the table, which at a real vocabulary is gigabytes.

**A borrowed cache was read from the seed, not from its owner.**
`serving.state_service.groups[*].ports.<proposer>` declares `access: read_only`
aliases: the port *is* the cell another component owns. The chain bound those
ports from the value the pass started with -- for a pass over a fresh context,
an empty cache -- so the drafter conditioned on nothing and the target
contradicted every draft. The fixture's target writes zeros into its cache, so
the fixture cannot tell the difference; acceptance was 0 of 12 on the real pair
and 12 of 70 after. The chain now binds a read-only alias from the value the
owning component's declared `output` port produced.

**The declared embedding was the unscaled table.**
A target whose graph multiplies a gathered row by a normalizer before its
backbone reads it does not hand a proposer the initializer's raw rows.
`token_embedding` gains an optional `scale` for exactly that: nothing in a
`[vocab, hidden]` initializer says whether one was applied, and guessing would
put a per-architecture heuristic in the runtime. It is applied once to the
table, never per gathered row.

The new gate is `chained_speculative_real_evidence`, driven by the
package-agnostic `tests/common/real_workflow.rs`: geometry is read from the
declared contracts and the graphs' own static dimensions, never written down.
A verified block commits the target's tokens whatever the proposer said, so
token-for-token equality with the standalone target package -- which the gate
requires at every width -- proves the two packages hold the same model and not
that the chain works. That claim rests on a statistic a broken chain cannot
fake: the target must confirm the proposer's first draft in at least one round
in four. It was 0 of 12 rounds before the borrowed-cache and normalizer fixes
and 4 of 10 after, at every width. The gate also requires real proposals, an
exercised rejection and rollback, the declared embedding table resolving to
non-zero weights of the declared width, the declared folded-carry seed existing
in the verification pass, and the device-residency property from #1861. Under
`ONNX_GENAI_REQUIRE_SPECULATIVE_EVIDENCE=1` a missing package, an unreadable
path, or partial width coverage is a failure rather than a skip.

Recorded run (H200, ORT 1.29.0 gpu_cuda13, CUDAExecutionProvider), 24 greedy
tokens, identical at every width and identical to the target package's own
greedy stream:

  drafts/proposal 1: 10 rounds, 4 with an accepted draft, 10 proposed,
                     4 accepted, 6 rejected, 6 rollbacks
  drafts/proposal 2: 10 rounds, 4 with an accepted draft, 20 proposed,
                     4 accepted, 16 rejected, 10 rollbacks
  drafts/proposal 4: 10 rounds, 4 with an accepted draft, 40 proposed,
                     4 accepted, 36 rejected, 10 rollbacks
  total: 30 rounds, 70 proposed, 12 accepted, 58 rejected, 26 rollbacks,
         780 cells rolled back, 100 proposer invocations, host_staging 0

The fifth defect is package-only and is fixed by republishing: every
sequence-like axis was called `sequence`, though the graphs distinguish the
query length, the cache a pass starts from, the cache it ends with, and the
borrowed caches a drafter reads. `scripts/name_sequence_axes.py` performs that
rename on a published `inference_metadata.yaml` without disturbing its authored
comments. See `docs/genai/CHAINED_SPECULATIVE_EVIDENCE.md` for the revisions
and how to reproduce.

Signed-off-by: justinchuby <justinchu@microsoft.com>
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.

2 participants