Repository navigation
feat(engine): a chained speculative proposal computes where its tensors already are - #1861
Merged
Merged
Conversation
…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 Report✅ All modified and coverable lines are covered by tests.
Additional details and impacted files@@ 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
Flags with carried forward coverage won't be shown. Click here to find out more. 🚀 New features to boost your workflow:
|
…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>
This was referenced Aug 23, 2026
Merged
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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_opsscaffolding, thenative_workflow_paritystaging 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 hermeticgemma4_chainedpackage. Measuring where they came from (backtrace onmaterialize_workflow_value_copy, origin/main030a4c019) shows all 16 arerollback_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:slice_axisis the only narrowing implemented.last_along_axisandtruncate_axisare windows of it, so the two paths a proposal chain uses cannot drift from the primitive (Rule 10).adoptis the only crossing, stated once per proposal for the carry seed the target left elsewhere — never per draft token.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:
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_valuestakes 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.Valuegainsempty_cuda/fill_zero_device(an owningcudaMallocallocation whose guard travels with the value),as_slice_f32(borrow a[vocab, hidden]table rather than copy it) andwrite_raw_bytes_at(dtype-agnostic host scatter).cuda_rtgrowsCudaAllocationovercudaMalloc/cudaFree.Value::alias_with_offsetalready existed from #1723.Rule 3:
WorkflowRuntime::host_copy_ofhad no callers left and is deleted.H200 evidence —
gemma4_chained, native-CUDA,speculative_decode(8, 4)030a4c019)[30, 5, 2, 7, 30, 5, 2, 7]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 nowassert_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-targetsclean foronnx-genai-engineandonnx-genai-ortundernative-cuda,native-backend, and no-default-features.Two failures reproduce unchanged on
origin/main030a4c019and are not from this branch:engine::runtime::tests::native_generate_rejects_over_kv_byte_budget_before_backend_runandloader::tests::selected_non_dense_candidate_fails_explicitly.Compatibility
HostTensorOpsis 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, anda_proposal_chain_stages_nothing_of_its_ownstill reports zero. Nothing here is keyed on a model, a port spelling, or a tensor shape: the sequence axis still comes from the declaredserving.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.