Skip to content

perf(ep-plugin): infer output shapes into reusable storage - #2167

Merged
justinchuby merged 2 commits into
mainfrom
resch/alloc-infer-shapes
Aug 26, 2026
Merged

justinchuby merged 2 commits into
mainfrom
resch/alloc-infer-shapes

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 26, 2026 •

Copy link
Copy Markdown
Owner

Part of #1077.

infer_shapes returns a fresh Vec<Vec<usize>> for every node of every Run. Both call sites only ever read it as a slice — .iter().enumerate(), an index, a &[_] — and drop it a few hundred instructions later. The container is pure overhead.

What it costs, measured rather than assumed

Callgrind --separate-callers=3 at depth 100, differenced across 300→1300 iterations:

site Ir/node
try_allocate_in ← infer_shapes (the inner to_vec) 57.4
malloc ← box_new_uninit (the one-element vec![_]) 63.0
drop_glue::<Vec<Vec<usize>>> 78.4
free ← that drop glue 78.4
total directly attributable 277.2

On top of that sits a share of glibc's _int_free, which at 482.4 Ir/node is the single largest allocator line in the whole profile — it aggregates every Rust deallocation, so it falls in proportion to the allocations removed. It did: −108.0 Ir/node, close to the 2-of-7 allocations-per-node this removes.

The profile also settled which arm the grid actually takes, rather than leaving it to inference. There is no infer_shared_node and no broadcast_shapes in it, and box_new_uninit is present — which is what vec![single] lowers to. That is SameAsInput, doing exactly two allocations.

The change

infer_shapes_into writes into a Vec<DimVec<usize>> parked in RunScratch, which already carries the per-Run input and output storage and is already borrowed once per Run behind a Drop guard that runs on early returns and while unwinding. DimVec is inline to rank 8, so a shape of ordinary rank costs no allocation at all.

Parking it there rather than using a per-Run local is what makes depth 1 improve. A local would be allocated and dropped for its single node, leaving the fixed per-Run cost exactly where it was.

Only the arms the profile shows as hot are reimplemented. Everything else delegates to infer_shapes verbatim. Forty inference rules are the correctness-critical part of this file, and duplicating them to save an allocation on a cold path would be a bad trade.

Result

Three-point profile, two disjoint differencing intervals plus a monotonicity check, under scripts/hostlock.sh held by the outer harness across every arm:

case before after delta interval disagreement
depth 100 2,881.4 Ir/node 2,518.4 −12.6% ≤0.10%
depth 1 14,681.4 Ir/Run 14,340.8 −2.3% ≤0.18%

Depth 1 is per-Run fixed cost, where 341 Ir of saving is diluted by everything else in a Run; depth 100 shows the per-node effect directly. The two agree on the per-node saving (~341 vs ~363 Ir), which is the cross-check that matters.

Correctness

The fast/slow split is only safe while the two agree, so assert_paths_agree is a differential oracle over both the Ok shapes and the Err text — a fast arm that skipped a bounds check would otherwise still "work" right up until it indexed out of range.

Mutation-proved, because an oracle that cannot fail is decoration:

mutation result
drop the SameAsInput bounds-check message RED
omit out.clear() (stale shapes from the previous node survive) RED
return a truncated shape from the broadcast arm RED

Also covered: rank 0; rank past INLINE_RANK where DimVec spills to the heap, reused across wide → narrow → wide; and a node with fewer outputs than its predecessor, where a buffer that only overwrote would leave the previous node's extra output slot visible.

Preserved: dynamic shapes (the buffer is refilled per node, never cached), absent/optional slots, concurrency (RunScratch is thread-local and the existing re-entrancy fallback allocates fresh storage that cannot alias), and retained outputs — IntermediateBuf now takes shape.to_vec() because the buffer is overwritten by the next node, which is the same allocation the Vec clone made there before.

No ORT CPU fallback, no runtime/default MLAS.

Validation: cargo fmt --check clean; clippy -p onnx-runtime-ep-plugin --all-targets --all-features clean; 350 lib tests + the full onnx-runtime-ep-cpu-plugin release suite (including ORT e2e conformance with NXRT_REQUIRE_ORT_TESTS=1) green.


Review round

Independent review (claude-opus-4.8) found no MUST-FIX in the production change: arm equivalence verified line by line against infer_shapes, the clear-on-every-Ok-path invariant confirmed, and the IntermediateBuf aliasing claim traced through all_output_views and absent_strides_storage rather than taken on trust.

It found two real holes in the tests, both of the same shape — an assertion that could not fail — and constructed a fourth and fifth mutation that survived the entire suite:

mutation before review after
M1 SameAsInput bounds-check message dropped RED RED
M2 out.clear() deleted, SameAsInput arm RED RED
M3 broadcast arm returns a truncated shape RED RED
M4 out.clear() deleted, delegating (other =>) arm survived RED
M5 out.clear() deleted, SharedNative arm survived RED

Each arm clears the buffer separately, so proving SameAsInput clears proved nothing about the other two. assert_paths_agree always started from a fresh buffer, where clear() is a no-op, and the truncation test only ever drove the SameAsInput arm. every_arm_truncates_when_a_node_has_fewer_outputs now drives all three writing arms onto a buffer already holding a longer result.

This is load-bearing in production, not just coverage bookkeeping: the routed path reads output_shapes.iter().enumerate() as the node's output arity, so a stale tail is a phantom output slot rather than a wrong number.

The review also showed the SharedNative arm had zero test coverage, which contradicted the differential-oracle guarantee this PR states for itself — the arm is reimplemented, so it needs the oracle as much as the others, and its Resolved branch is the one production takes for Expand/Tile with concrete shape operands. Now covered on both sides of the branch plus a declining fallback, with an assertion that the declining case really declines so the test cannot pass vacuously.

Two comment corrections, both substantive:

  • clear_and_bound claimed the capacity bound existed because a spilled DimVec pins heap the other scratch vectors do not. It does not — Vec::clear() drops the elements, so any spilled DimVec's heap is already released there. The bound covers only the outer spine, same rationale as owned/slots. A comment that misstates why a bound exists is how the bound gets deleted later.
  • "Allocating nothing" was false for SharedNative-Resolved, which receives an owned Vec<Vec<usize>> and copies it in — strictly more work than returning it. Narrowed to the elementwise paths, and the doc now states what that arm is actually for: routing a declining rule's fallback through the fast arms.

The review commit is tests and comments only; the measurements above are unchanged and were re-confirmed against the final head rather than assumed.

infer_shapes returns a fresh Vec<Vec<usize>> for every node of every Run.
Callers only read it as a slice and drop it a few hundred instructions
later, so the container is pure overhead.

Measured at depth 100 with callgrind --separate-callers=3: 277.2 Ir/node
across the two allocations and their frees (57.4 inner to_vec, 63.0 for
the one-element vec![_] lowering to box_new_uninit, 78.4 drop glue, 78.4
free), on top of which sits a share of glibc _int_free -- the single
largest allocator line in the profile at 482.4 Ir/node aggregated over
every Rust deallocation.

infer_shapes_into writes into a Vec<DimVec<usize>> parked in RunScratch,
which already carries the per-Run input and output storage and is already
borrowed once per Run behind a Drop guard. Parking it there rather than
using a per-Run local is what lets depth 1 benefit: a local would be
allocated and dropped for its single node.

Only the arms the dispatch profile shows as hot are reimplemented; every
other strategy delegates to infer_shapes verbatim. Forty inference rules
are the correctness-critical part of this file and duplicating them to
save an allocation on a cold path would be a bad trade.
assert_paths_agree stands behind the split as a differential oracle over
both the Ok shapes and the Err text, so the paths cannot drift silently.

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

codecov Bot commented Aug 26, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 95.48023% with 8 lines in your changes missing coverage. Please review.
✅ Project coverage is 80.58%. Comparing base (7cb5d98) to head (59bbd21).
⚠️ Report is 9 commits behind head on main.

Files with missing lines Patch % Lines
crates/onnx-runtime-ep-plugin/src/compute.rs 95.48% 5 Missing and 3 partials ⚠️
Additional details and impacted files

Impacted file tree graph

@@           Coverage Diff            @@
##             main    #2167    +/-   ##
========================================
  Coverage   80.58%   80.58%            
========================================
  Files         431      432     +1     
  Lines      217880   218062   +182     
  Branches   217880   218062   +182     
========================================
+ Hits       175568   175718   +150     
- Misses      36481    36502    +21     
- Partials     5831     5842    +11     
Flag Coverage Δ
cli-ort-linux 72.51% <ø> (ø)
cli-ort-windows 72.10% <ø> (+0.09%) ⬆️
mlas 85.90% <ø> (ø)
offline 80.69% <95.48%> (+<0.01%) ⬆️

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

Files with missing lines Coverage Δ
crates/onnx-runtime-ep-plugin/src/compute.rs 77.50% <95.48%> (+0.77%) ⬆️

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

…d-bearing

Independent review found no production defect but two real holes in the
tests, both of the same shape: an assertion that cannot fail.

Each arm of infer_shapes_into clears the buffer separately, so proving
that SameAsInput clears proves nothing about the others. Deleting the
out.clear() from the delegating arm, or from the shared-native arm,
survived the whole suite -- the truncation test only ever drove the
SameAsInput arm, and assert_paths_agree always started from a fresh
buffer, where clear() is a no-op. Both are load-bearing in production:
the routed path reads output_shapes.iter().enumerate() as the node's
output arity, so a stale tail is a phantom output slot.

every_arm_truncates_when_a_node_has_fewer_outputs now drives all three
writing arms onto a buffer already holding a longer result. Both
deletions now fail.

The shared-native arm was also entirely untested, which contradicted the
oracle guarantee this PR claims for it -- it is reimplemented, so it
needs the differential oracle as much as the others, and its Resolved
branch is the one production takes for Expand/Tile with concrete shape
operands. Covered on both sides of the branch plus a failing fallback,
with an assert that the declining case really declines so the test
cannot pass vacuously.

Also corrected two comments the review showed were overstated: the
capacity bound covers only the outer Vec's spine, since clear() already
released any spilled DimVec's heap; and 'allocating nothing' is not true
of the shared-native arm, which still copies an owned Vec in.

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

Copy link
Copy Markdown
Owner Author

Review addressed. No MUST-FIX was found, but the two SHOULD-FIX items were both real and both of the same shape — an assertion that could not fail — so they are worth stating plainly.

The fourth and fifth mutations you constructed were correct. Deleting out.clear() from the delegating arm, and from the shared-native arm, each survived the entire suite. Your diagnosis of why is exactly right: assert_paths_agree always starts from a fresh buffer, where clear() is a no-op, and the truncation test only ever drove the SameAsInput arm — so it proved that one arm clears and silently claimed the property for all three. Each arm clears separately; the test has to drive each one separately.

every_arm_truncates_when_a_node_has_fewer_outputs now runs all three writing arms onto a buffer already holding a longer result. Re-mutated:

mutation before now
out.clear() deleted, delegating arm survived RED
out.clear() deleted, shared-native arm survived RED

That matters more than the coverage gap it closes, because this is production-load-bearing: the routed path reads output_shapes.iter().enumerate() as the node's output arity, so a stale tail is not a wrong number — it is a phantom output slot.

On the untested shared-native arm — you were right that it contradicted the guarantee the PR states for itself, and right that Resolved is the branch production takes for Expand/Tile with concrete shape operands. Now covered on both sides of the branch, plus a fallback that itself fails so the error text is compared too. I added assert_eq!(infer_shared_node(&unknown, ...), SymbolicOrUnknown) before the declining case, because that test passes vacuously if the rule ever starts resolving.

Both NITs taken, and both were substantive rather than cosmetic.

  • The clear_and_bound comment claimed the bound existed because a spilled DimVec pins heap the other vectors do not. You checked and that heap is already released by clear() dropping the elements. The bound covers only the outer spine, same rationale as owned/slots. Corrected — a comment that misstates why a bound exists is how the bound gets removed later.
  • "Allocating nothing on the paths the dispatch grid actually takes" was false for shared-native Resolved, which still receives an owned Vec<Vec<usize>> and copies it in — strictly more work than Ok(shapes). Narrowed to the elementwise paths, and the doc now says what that arm is actually for: routing a declining rule's fallback through the fast arms, not saving an allocation on resolve.

Your §3 verification of the IntermediateBuf lifetime claim is the check I most wanted an independent pass on, since that is the one place a retained borrow into reused storage would be use-after-overwrite rather than a wrong number. Thank you for tracing all_output_views and absent_strides_storage rather than taking the claim.

351 lib tests green; fmt and clippy clean. Measurements unchanged — this commit is tests and comments only.

@justinchuby
justinchuby enabled auto-merge (squash) August 26, 2026 01:34
@justinchuby
justinchuby merged commit bc79b3f into main Aug 26, 2026
16 of 20 checks passed
@justinchuby
justinchuby deleted the resch/alloc-infer-shapes branch August 26, 2026 01:43
justinchuby added a commit that referenced this pull request Aug 26, 2026
Part of #1077. Follows #2167, same method: attribute, implement, measure
across two disjoint intervals, mutation-prove.

## What

Every buffer-sink output built **two heap vectors per node** — a
`Vec<usize>` for the shape and a `Vec<i64>` for the strides — and freed
them one node later. `DimVec` already exists in this crate for exactly
this shape of problem: inline up to `INLINE_RANK` (8), spilling to the
heap only past it. `output_shapes` is already a `DimVec` since #2167, so
the shape copy becomes an inline copy rather than an allocation.

`contiguous_strides` returns a `DimVec<i64>` too, which empties its
callers as well: `absent_slot_strides` was building a `Vec<Vec<i64>>`,
one inner allocation per absent slot.

## Measured

Depth 100, `grid_relu_100_tiny`, single-threaded, callgrind Ir, three
points (300/800/1300) giving **two disjoint differencing intervals plus
a monotonicity check**:

| | before | after | delta |
|---|---:|---:|---:|
| depth 100 | 2,518.4 Ir/node | **2,306.5 .. 2,310.5** | **−8.3%** |
| depth 1 | 14,340.8 Ir/`Run` | 14,313.8 .. 14,347.5 | flat |

Interval disagreement 0.17% (depth 100) and 0.23% (depth 1).

**Depth 1 being flat is the expected result, not a disappointing one.**
A single-node graph has no intermediate buffers at all — its output goes
straight to an ORT sink — so there is nothing here for it to save. A
depth-1 improvement would have been evidence I had measured something
other than what I changed.

## Mechanism

**Correction to an earlier version of this description.** It claimed
`contiguous_strides` "leaves the profile entirely". That was wrong, and
it was wrong because I read a *truncated top-30 list of source lines*
and concluded from an absence, instead of differencing the before and
after profiles by symbol. `contiguous_strides` does not leave the
profile — it gets **more expensive**, 51.5 → 75.2 Ir/node, because
indexing through a `DimVec` enum costs more than indexing a slice. The
−8.3% is unaffected and independently measured, but the story underneath
it is a good deal more interesting than the one I first told.

Before/after self cost by symbol, both from
`-Cdebuginfo=line-tables-only` builds differenced 300→1300 at depth 100:

| symbol | before | after | delta |
|---|---:|---:|---:|
| `_int_free` | 230.0 | 122.8 | **−107.3** |
| `malloc` | 189.9 | 100.5 | **−89.4** |
| `free` | 165.8 | 88.5 | **−77.2** |
| `Vec<Vec<i64>>` `SpecFromIter` | 37.0 | — | −37.0 |
| `compute_execute::{closure#0}::{closure#0}` | 705.9 | 681.1 | −24.7 |
| `drop_glue::<Vec<Vec<usize>>>` | 19.0 | — | −19.0 |
| `__rustc::__rdl_alloc` | 36.5 | 18.7 | −17.8 |
| **`__memcpy_avx_unaligned_erms`** | 25.3 | 103.6 | **+78.4** |
| `Vec<DimVec<i64>>` `SpecFromIter` | — | 44.0 | +44.0 |
| **`contiguous_strides`** | 51.5 | 75.2 | **+23.8** |
| `drop_glue::<Vec<DimVec<i64>>>` | — | 19.0 | +19.0 |
| `Map<Enumerate<Iter<RoutedSlot>>>` | 95.1 | 102.9 | +7.9 |

**−372 Ir/node of allocator and allocation-shaped work removed, +173
given back as `memcpy` and inline-enum overhead, net −209.** That
matches the independent `ir3` interval measurement (−208 to −212) and
the whole-profile totals (2,518.0 → 2,309.2).

So the trade this PR makes is bigger in both directions than the
headline suggests, and the `memcpy` line is the receipt for the
struct-widening cost predicted below — 25.3 → 103.6 Ir/node, from moving
a ~112-byte-wider struct twice per node. **That is the next lever**, and
it is now quantified rather than suspected: eliminating one of the two
moves is worth up to ~39 Ir/node on its own.

## The trade, stated plainly

`IntermediateBuf` gets **wider**: two `DimVec`s are larger than two
`Vec`s, and the routed path moves the struct twice (staged into
`new_bufs`, then installed into `intermediates` after the kernel runs —
the staging is what keeps a new buffer from clobbering an input that
shares its index). That shows up in the profile as more `memcpy`, and it
is why the win is 8.3% and not more.

The two `malloc`/`free` pairs it removes are worth more than the wider
move, and — as in #2167 — the shared free path shrinks with total
allocator traffic on top of the directly attributable saving.

`contiguous_strides` now zeroes and sets the innermost stride rather
than filling with ones, because every other element is overwritten by
the loop immediately after.

## Visibility

`IntermediateBuf` and its fields drop from `pub` to `pub(crate)`.
`DimVec` is `pub(crate)`, so a `pub` field of that type is a
`private_interfaces` warning. The struct is an internal execution detail
that is never named outside `compute.rs`; the alternative — widening
`DimVec` to `pub` — would export an internal representation to make a
warning go away.

## Tests

Three new, all differential or boundary-focused rather than restatements
of the implementation:

- `contiguous_strides_matches_the_ir_oracle_across_the_inline_boundary`
— walks rank 0..=`INLINE_RANK`+3 against
`onnx_runtime_ir::compute_contiguous_strides`, which is the same
algorithm **in a crate this change does not touch**, so it is a real
oracle. Non-uniform extents, so a transposed or off-by-one stride cannot
coincide with the right answer. Asserts it saw both representations, so
it cannot pass vacuously if the boundary moves.
- `contiguous_strides_spills_rather_than_truncating` — truncation at
`INLINE_RANK` would still produce plausible-looking leading strides, so
length and innermost/outermost values are pinned separately.
- `an_intermediate_buf_owns_its_shape_past_the_inline_rank` — the buf's
shape is copied from a slot in reusable scratch that the next node
overwrites. For a spilled rank, owning means a **deep** copy; the test
mutates the source after construction and demands the buf is unaffected.

**Mutation-proved — all six go red:**

| mutation | result |
|---|---|
| M1 innermost stride never set to 1 | RED |
| M2 `zeroed(len.min(INLINE_RANK))` — silent truncation at the spill
boundary | RED |
| M3 off-by-one loop bound, second-innermost stride left unset | RED |
| M4 `absent_slot_strides` returns a bogus 1-element vector for an
unknown shape | RED |
| **M5 `(shape[i + 1] as i64).max(2)` in the recurrence** | **RED
(survived until review)** |
| **M6 `view()` truncates a spilled shape at `INLINE_RANK`** | **RED
(new test)** |

## Review round

Independent review found **no MUST-FIX** — it verified the rewrite is
exactly equivalent for every rank including 0 and 1, that the clone is a
genuine deep copy in both representations, that dropping `output.shape`
at the kernel-sized site is correct after the partial move of
`output.bytes`, and that nothing outside `compute.rs` names
`IntermediateBuf`.

It found two real problems, both in the tests.

**M5 survived the whole suite.** Replacing the recurrence's multiplicand
with `(shape[i + 1] as i64).max(2)` is the *identity* for every shape
the tests used, because every one of them was built from extents of 2 or
more.

The reasoning that opened the hole was in my own comment: *distinct,
non-uniform extents so a transposition cannot coincide with the right
answer*. That is a good argument for including 2, 3, 4 and a bad one for
excluding 1. A stride **on** a size-1 axis is inert — its index is
always zero — which is what makes it tempting to leave out. But that
axis is still a **multiplicand** for every axis outside it, so an error
there propagates into strides that are live. `[2,1,3]` should be
`[3,3,1]`; the mutant gives `[6,3,1]`, and element `(1,·,·)` reads three
floats past where it should. Size-1 axes are ubiquitous — broadcasting,
unsqueezed axes, NCHW with C=1.

The sweep now carries interior and trailing unit axes on both sides of
the spill boundary, an all-ones spilled shape, and a zero extent, and
asserts it *exercised* a size-1 interior axis so it cannot quietly
regress to all-large extents again.

**`an_intermediate_buf_owns_its_shape_past_the_inline_rank` was a
tautology, and is gone.** It asserted that `DimVec::clone` deep-copies —
a compiler guarantee, since there is no safe `Clone` that shares a
`Vec`'s buffer, and the aliasing alternative (a move) does not compile.
No mutation could make it fail, and it never called the routed
construction site it claimed to be about; it built its own struct
literal. That is precisely the failure the previous review round caught
me on, reproduced one PR later.

Replaced with `a_spilled_intermediate_buf_view_reports_every_dimension`,
which pins something that *can* go wrong — `view()` handing the kernel a
truncated slice for a spilled rank — and is mutation-proved by M6.

The reviewer also noted
`contiguous_strides_spills_rather_than_truncating` is redundant with the
oracle sweep for *detection*. Kept deliberately, with the reason now in
the comment: it states the answer in closed form, and the oracle is the
same algorithm by construction — which makes it a good check on
representation and initialisation and a poor one on the algorithm
itself.

The review commit is tests only; both hunks fall inside `mod tests`, so
the measurements above stand against the final head rather than being
re-asserted.

## Validation

`cargo fmt` clean; `clippy -p onnx-runtime-ep-plugin --all-targets
--all-features` **0 warnings**; 352 lib tests; full
`onnx-runtime-ep-cpu-plugin` release suite with
`NXRT_REQUIRE_ORT_TESTS=1` green, including the ORT e2e conformance
tests that drive this path for real.

All builds and measurements ran under `scripts/hostlock.sh` (#1806) held
by the outer harness across every arm.

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant