Skip to content

Stop handing activation nodes to ORT's CPU EP - #1097

Merged
justinchuby merged 9 commits into
mainfrom
deckard/no-ort-fallback
Aug 17, 2026
Merged

justinchuby merged 9 commits into
mainfrom
deckard/no-ort-fallback

Conversation

@justinchuby

@justinchuby justinchuby commented Aug 16, 2026 •

Copy link
Copy Markdown
Owner

What this changes

When our CPU EP is selected, it must not hand work to ORT's CPU EP. This PR
removes the three mechanisms by which it was doing so for the activation
and normalization families. GetCapability runs three independent fail-closed
filters and a claim must clear all of them; each of the last two was found only
after this PR had already claimed the job was done.

1. The performance-based decline policy is deleted

assignment_policy.rs (2118 lines) measured whether we beat ORT on a given
shape/dtype and, where we lost, returned ClaimPreference::defer so ORT's CPU
EP would take the node. That whole file is gone, along with the
claim_preference override in provider.rs. claim_preference_node now
returns Claim immediately.

Beyond the architectural rule, the policy could not have worked as intended:

  • A deferral splits the graph. Every declined node is a partition boundary,
    which costs fusion, prepacking and buffer reuse across it — none of which the
    per-node threshold accounted for.
  • It cannot see the thread count. Capability runs before the session's
    intra-op pool is known. Sqrt at 64 Ki wins 1.9x against a single-threaded
    ORT and loses at 0.30x against 16 threads. One number cannot be right for
    both.
  • Sometimes there is nothing to defer to. ORT has no bf16 kernel for these
    ops and no f16 kernel for most. Declining a bf16 Gelu does not get a faster
    kernel, it gets a load failure.
  • The thresholds were tuned on one host. Every number was measured on a
    single EPYC 9V74. Shipping it made every other machine's latency a guess.

Removing the override is also a small capability-time win: the default adapter
deep-clones every input Shape and collects dtypes for every node in order to
build the metadata the policy consumed.

2. The shape-inference filter was declining the same ops anyway

Deleting the policy did not, by itself, achieve the goal. GetCapability runs
a second, independent fail-closed filter (onnx-runtime-ep-plugin/src/ep.rs):
it drops any claim containing a node whose ShapeInference::for_node returns
Declined, and that match ends in _ => Declined. An op we register a
kernel for, but which is absent from that table, is silently handed to ORT no
matter what supports_op answers.
This is the same mechanism that made the
com.microsoft activations unreachable until #1082.

The trigonometric, hyperbolic and remaining activation ops were all in that
gap. Now listed: Sin, Cos, Tan, Asin, Acos, Atan, Sinh, Cosh,
Asinh, Acosh, Atanh, ThresholdedRelu, Swish, com.microsoft::Silu
and PRelu.

Silu had been deliberately excluded with the comment that ORT has no kernel
for it. That reasoning was backwards: an op ORT cannot run is precisely the one
we must never hand over.

GroupNormalization was in the same gap and is now covered too.

3. The dtype filter was declining a different set of ops

Review found a third filter, and it was still handing over one of the very ops
section 2 had just fixed.

node_passes_dtype_filter looks the node's op up in the plugin's
KernelRegistryEntry list and returns false when there is no entry. That
list is built from build_cpu_registry_with_descriptors, which recorded keys
as they were registered — but register_cnn_ops takes &mut OpRegistry and
writes past the recording wrapper. Eighteen ops were in the registry and
absent from the descriptors, so supports_op claimed each one and capability
then dropped it:

PRelu, BatchNormalization, InstanceNormalization, GroupNormalization,
Conv, ConvTranspose, MaxPool, AveragePool, GlobalMaxPool,
GlobalAveragePool, GlobalLpPool, LpPool, Resize, GridSample,
AffineGrid, Col2Im, CenterCropPad, SpaceToDepth

Four are activations or normalizations this EP owns. PRelu is the sharpest
case: section 2 gave it a shape rule, so it cleared filter two, and the dtype
filter declined it anyway.
Both pure-Rust inventory tests passed while real
ORT ran the node.

Descriptors are now derived from OpRegistry::keys() instead of a parallel
recorded list, making the two sets identical by construction rather than by
convention. They are also sorted: they get leaked into a 'static slice ORT
reads, and hash-map iteration order would make any snapshot diff flap.

descriptors_derived_from_real_registry_not_hand_maintained had been asserting
this bug as correct behaviour — it allowed a delta of up to 50 entries and
named CNN ops as the expected difference. It now asserts set equality.

The lesson, having now been caught twice: an inventory test is only as good as
its source of truth, and two review rounds passed on tests that enumerated the
wrong set. The only check that cannot be fooled this way is the end-to-end one
that asks real ORT which EP got the node.

Scope — what this does not fix

64 registered ops remain in the shape-inference gap. This PR closes the
activation/elementwise families; it does not close the gap universally. The
full list is asserted exactly by the new inventory test and summarised in
docs/performance/CPU_ACTIVATION_GAPS.md:

  • 20 data-dependent — output shape is a function of an input's values
    (NonZero, Unique, Compress, Expand, Tile, Pad, TopK, Split,
    Unsqueeze, Resize, AffineGrid, Col2Im, CenterCropPad, ...).
    Correctly declined today, though most carry a constant initializer in
    practice, so a pass that resolves initializer values at capability time could
    claim them.
  • 10 internal fusion ops created after capability, never candidates
    (FusedGemm, FusedAttention, FusedMatMulBias, the pkg.nxrt ops).
  • 34 inferrable but unwritten — this is the work. Ten are one-line
    shape-preserving rules (QuantizeLinear, DequantizeLinear,
    CastLike, ScatterND, Trilu, CumSum, ...); nine
    are pooling/CNN geometry (MaxPool, AveragePool, Global*Pool,
    ConvTranspose, GridSample, SpaceToDepth) inferrable exactly as
    build_conv already does for Conv; eight more follow from attributes
    (ArgMax, Flatten, GatherElements, Size, ...); the rest are contrib
    and model ops.

Two entries deserve singling out. com.microsoft::Attention is the attention
op in exported GenAI models, and the existing opset-23 arm is guarded to the
default domain, so we hand it over. LinearAttention (both domains) and
com.microsoft::CausalConvWithState are the Qwen3.5 / Qwen3-Next hybrid
linear-attention primitives — ORT has no kernel for them at all, so
declining them does not get a faster implementation, it gets a load failure.

Closing that remainder is the next PR. It is a different domain from the
activation kernels and each entry needs its own numeric test.

Corrections to my earlier figures

I published two wrong counts before this test existed, in opposite directions,
and the test found both.

66, then 52, now 66 again — for different reasons each time. The first
figure came from a scratch script that text-matched the table and missed its
guard arms, so it reported the whole Reduce* family (handled by op_name if is_reduction(op_name)) as a gap. Correcting that, I over-corrected and claimed
the pooling family "has rules already" — it does not; compute.rs has no
pooling arm at all. The 52 figure was also built on
build_cpu_registry_with_descriptors, which is not the registry:
register_cnn_ops writes straight to the inner OpRegistry, so 14 CNN ops and
PRelu never appear in the descriptors at all. The test now enumerates
OpRegistry::keys() — the same set supports_op consults.

MoE/QMoE/LinearAttention/CausalConvWithState were misclassified as
internal fusion ops.
They are read from exported models —
deepseek_v2_tiny_qmoe_native_e2e.rs asserts a loaded graph contains a
com.microsoft::QMoE node, and the linear-attention pair are Qwen3.5
primitives. All are real gaps.

Three gaps I had missed entirely: Unsqueeze (declines whenever axes is
input[1], i.e. every opset-13+ graph), com.microsoft::Attention, and
EyeLike.

This is the argument for the inventory test: hand-maintained prose about which
ops reach ORT was wrong three times in a row, in both directions, and each time
it read as confident.

Tests

Inventory - the test that would have caught this class of bug

every_registered_op_has_a_shape_rule_or_is_a_known_gap enumerates
OpRegistry::keys() and asserts the set of ops that decline shape inference
exactly. Registering an op without a shape rule fails; adding a shape rule
without removing the op from the list also fails. Neither direction can pass
silently, and the allowlist doubles as the gap inventory. It needs no ORT, so
it runs in every job rather than only the ORT-gated one.

The probe is a sweep, not a point: opsets {1, 13, 18, 22, 23} x arities 1..4 x
ranks 1..4, counting an op as declining only when nothing in the matrix
produces a rule. A single one-input rank-2 probe reported Conv as a gap,
because build_conv reads input_shapes[1][0] for its output channel count
and needs rank >= 3. Sweeping removes that artifact and keeps the allowlist from
encoding one arbitrary opset.

Verified both directions falsify: deleting | "Cos" fails with [("", "Cos")];
adding | "Trilu" fails with these ops now have a shape rule but are still listed as declined: [("", "Trilu")].

no_activation_or_norm_op_is_left_to_ort is a standing guard on the 39
activation and norm ops this EP owns - the families #1082, #1093 and #1097 made
reachable. It asserts in both directions: a filter of the form
registered.contains(op) && declines(op) would let a deleted registration
pass silently, which is the same hand-off by a different route. Verified by
renaming the Silu kernel key, which now fails with these ops are no longer registered by the CPU EP, so ORT will execute them: [("com.microsoft", "Silu")]. That direction immediately caught two entries I had wrong: Swish
is registered in the default domain, not com.microsoft, and HardSwish has
no kernel at all.

Assignment sweeps

Two sweeps replace the 13 deleted deferral tests. Both iterate 21 fixtures:

  • no_supported_node_is_ever_left_to_the_ort_cpu_ep - for every fixture, the
    op under test appears in our EP's node list and in no other EP's.
  • every_fixture_loads_with_cpu_fallback_disabled - loads each fixture with
    session.disable_cpu_ep_fallback, so a silent hand-off becomes a load
    failure rather than a slow success.

Verified the sweep falsifies too: with | "Sin" removed from the table it
fails with [sin_assignment_f32] ours=[], others=["Sin"].

Both run fail-closed in CI: conformance_setup panics rather than skipping when
NXRT_REQUIRE_ORT_TESTS=1, which the CLI ORT job sets. (That job's plugin
step was itself being skipped whenever an earlier step failed - fixed
separately in #1096.)

erf_reference would have become dead code when the deferral tests were
deleted. Rather than remove it, it is now used by
float16_biasgelu_runs_on_our_ep_with_correct_numerics, which checks our
kernel's numerics - more important now that we always execute BiasGelu
instead of sometimes declining it.

Docs

CPU_MATMUL_ASSIGNMENT.md is reframed from claim/defer to win/gap. Every
measurement is unchanged — the numbers still say exactly where we are slower
than ORT; they are now a work list rather than a decline table.

CPU_ACTIVATION_GAPS.md is new: every range where we still lose, and the two
root causes. Neither is polynomial accuracy:

  1. Our elementwise kernels are single-threaded while ORT splits across its
    intra-op pool. This is the flat ~0.7-0.8x plateau across the f32 activation
    family, and it is worth more than any further approximation work.
    KernelContext_ParallelFor is the untried lever.
  2. ~1.2 us of fixed per-node plugin dispatch overhead, which is what the
    ~0.75-0.8x at n=1 measures.

Validation

cargo fmt, clippy clean on the three affected crates, 1266 ep-cpu lib tests, 224 ep-plugin unit tests, 2 inventory tests, and 37 plugin E2E tests under NXRT_REQUIRE_ORT_TESTS=1 with real ORT.

Update — third decline path (commit 67a611353)

Opus review round 4 returned NO-GO on the grounds that a fourth decline
path existed and the PR's central claim was therefore still false. It was
right. See section 3 above.

Independently verified before fixing: 18 of the 177 unique (domain, op_type)
pairs in the registry had no descriptor, including PRelu,
BatchNormalization, InstanceNormalization and GroupNormalization.
Counting individual registry entries (op_type + domain + since_version,
the unit OpRegistry::len() reports) the registry holds 208, and descriptors
now match it exactly — descriptors_derived_from_real_registry_not_hand_maintained
asserts that equality.

New guards, each verified to fail when its invariant is broken:

guard falsified by observed failure
every_registered_op_has_a_kernel_registry_entry filtering PRelu out of descriptors 1 registered ops have no kernel-registry entry ... ["::PRelu"]
activation_and_norm_ops_clear_every_capability_filter same ::PRelu: no kernel-registry entry (dtype filter declines it)
prelu_assignment_f32 (real ORT) same ours=[], others=["PRelu"] → 'PRelu' must run on this EP
descriptors_derived_from_real_registry_not_hand_maintained same names the missing ops instead of tolerating a delta of 50

After the fix, real ORT reports ours=["PRelu"], others=[] and
ours=["GroupNormalization"], others=[].

Writing activation_and_norm_ops_clear_every_capability_filter immediately
found two more real gaps: Celu and Mish have no kernel at all. That is a
missing feature rather than a decline, so they are excluded from that test with
a comment naming them, and recorded in CPU_ACTIVATION_GAPS.md under a new
"Activations with no kernel at all" section rather than quietly dropped.

Re-validated: cargo fmt --all --check, scoped clippy clean, 1267 ep-cpu lib
tests, 224 ep-plugin unit tests, and 57 plugin tests under
NXRT_REQUIRE_ORT_TESTS=1 (37 E2E + 4 coverage + 9 + 6 + 1).

Update — Opus round 5: GO WITH FINDINGS

Round 5 confirmed the central claim now holds, verified against real ORT
(ours=["PRelu"], others=[] and ours=["GroupNormalization"], others=[]), and
traced every node-removing gate in ep_get_capability_inner to confirm no
fifth decline path exists for the activation/norm families. Two findings, both
addressed:

Finding 1 (minor, real) — Conv advertised a dtype its kernel rejects.
Giving Conv a KernelRegistryEntry was itself an over-claim:
supported_dtypes_for_op("Conv", "") returned FLOAT_DTYPES, which includes
f64, but ConvKernel::execute rejects anything outside f32/f16/bf16. Before
this PR Conv had no descriptor so an f64 Conv was declined; after it, the
node would clear both filters, compile, and then fail at Run. Fixed by adding
MLAS_FLOAT_DTYPES and giving Conv its own arm.

This is not a re-introduced fallback. Advertising a dtype we cannot execute is
a different thing from declining one we can: f64 Conv is genuinely
unsupported, and reporting that honestly at capability time is correct. The
rest of the CNN family really does dispatch f64 through dispatch_float! and
keeps FLOAT_DTYPES. Pinned by conv_does_not_advertise_a_dtype_its_kernel_rejects.

Finding 2 (minor, docs) — stale counts in this description. The breakdown
said 20 + 10 + 36 = 66 while the total said 65, and still listed
GroupNormalization among the unwritten shape-preserving rules even though
this PR wrote one. Corrected to 35 above. The "177" figure was unique
(domain, op_type) pairs, not OpRegistry::len() (208) — both are now stated
explicitly. The committed CPU_ACTIVATION_GAPS.md was already correct.


Merge with main

main moved under this branch and #1101 independently hit the same dtype-filter
bug class, for MatMulNBits and QLinearMatMul, adding FLOAT_COMPUTE_DTYPES —
byte-for-byte the same f32/f16/bf16 set as this branch's MLAS_FLOAT_DTYPES.
Two people finding the same trap independently is the strongest evidence that
the fail-closed dtype filter needed the systematic fix in this PR rather than
another per-op patch. kernels/mod.rs was resolved by taking main's block
verbatim, deleting the duplicate constant, pointing Conv at
FLOAT_COMPUTE_DTYPES and folding this branch's rationale into main's doc
comment.

The merge then made the inventory test earn its keep twice:

The probe was lying about com.microsoft::MatMulNBits. declines() supplied
no attributes, deliberately: an attribute-dependent rule should fall back to the
ONNX default when the attribute is absent, and the empty bundle keeps that
honest. But MatMulNBits derives its output width from N, and a node without
N is malformed, not defaulted — for_node declines it, correctly. No real
graph reaches that path, so the test was reporting a gap that does not exist.
The probe now sweeps a plausible attribute bundle alongside the empty one,
so default-fallback rules are still checked against absence while rules that
cannot default are modelled the way production sees them.

With that fixed, the first assert stopped masking the second.
QLinearMatMul gained a real rule in main and was stale in this branch's
DECLINED list. Removed. That is the drift check doing exactly the job it was
written for — the list cannot quietly stop describing reality in either
direction.

Gap count 65 → 64; group 3, 35 → 34. CPU_ACTIVATION_GAPS.md updated to match.

Post-merge validation, all green: cargo fmt --all -- --check; scoped clippy
(onnx-runtime-ep-cpu, -ep-plugin, -ep-cpu-plugin, --all-targets) clean;
1280 ep-cpu lib tests; and the full plugin suite under
NXRT_REQUIRE_ORT_TESTS=1 — 40 real-ORT E2E assignment tests plus the 4
inventory tests.

The Fast lane on this branch will stay red on cargo fmt --all -- --check
until #1107 lands. That failure is main's, from #1101, in
matmul_nbits.rs — a file this PR does not touch. The fix is deliberately
not duplicated here; this branch will pick it up by merging main once
#1107 is in.

When this EP is selected it now claims every node it supports. The
performance-based decline is removed entirely.

`assignment_policy` measured each op against ORT's CPU kernel and returned
`DeferToHost` for the shape/dtype ranges it lost, so ORT's CPU EP ran
those nodes instead. Treated purely as a per-node latency question that
was defensible, and every threshold in it was measured rather than
guessed. As an architecture it is wrong: selecting this EP is a request
for this EP, and a range where we are slower is a kernel to optimize, not
a node to give away.

The cost of the deferral was never just the node either. It splits the
graph, so the session pays a partition boundary and gives up the fusion,
prepack and buffer reuse that only hold inside one partition — none of
which the per-node A/B that justified the threshold could see. It also
made latency depend on thresholds tuned on one machine, which is why
`Sqrt` had to be declined despite winning 1.9x single-threaded: it
inverts to 0.30x at sixteen threads and the thread count is not visible
at capability time.

So `claim_preference` is gone rather than reduced. `claim_preference_node`
is overridden to answer before the default adapter deep-clones every input
shape and collects input dtypes for every node in the graph, which nothing
reads any more.

The measurements are kept, reframed as a work list: docs/performance/
CPU_ACTIVATION_GAPS.md records every range where this EP is still slower
than ORT, and CPU_MATMUL_ASSIGNMENT.md becomes evidence for the same. Each
row under 1.00 is now an open gap rather than a threshold.

Two falsifiers replace the thirteen tests that asserted deferral. They
sweep all twenty assignment fixtures — built to straddle the old
boundaries, so exactly the right corpus for the opposite assertion — and
check ORT's own node-to-EP assignment, which fails if a decline is
reintroduced anywhere in the path rather than only in the policy
function. The second runs them with `disable_cpu_ep_fallback=1`, where an
unclaimed node makes session creation fail outright, so loading is itself
proof that nothing was given away.

float16 `BiasGelu` keeps its numeric check, now against our kernel rather
than ORT's: we execute it, so verifying the arithmetic matters more than
it did when ORT computed it.

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

codecov Bot commented Aug 16, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 95.29412% with 4 lines in your changes missing coverage. Please review.
✅ Project coverage is 79.90%. Comparing base (b7fa5e1) to head (37a3b7a).

Files with missing lines Patch % Lines
crates/onnx-runtime-ep-cpu/src/kernels/mod.rs 91.83% 3 Missing and 1 partial ⚠️
Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main    #1097      +/-   ##
==========================================
+ Coverage   79.86%   79.90%   +0.04%     
==========================================
  Files         367      368       +1     
  Lines      157553   159485    +1932     
  Branches   157553   159485    +1932     
==========================================
+ Hits       125823   127438    +1615     
- Misses      27007    27332     +325     
+ Partials     4723     4715       -8     
Flag Coverage Δ
cli-ort-linux 83.79% <ø> (ø)
cli-ort-windows 83.40% <ø> (ø)
mlas 84.54% <ø> (?)
offline 79.67% <95.29%> (-0.04%) ⬇️

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

Files with missing lines Coverage Δ
crates/onnx-runtime-ep-cpu/src/provider.rs 88.96% <100.00%> (+0.96%) ⬆️
crates/onnx-runtime-ep-plugin/src/compute.rs 79.93% <100.00%> (+2.54%) ⬆️
crates/onnx-runtime-ep-cpu/src/kernels/mod.rs 94.98% <91.83%> (+2.03%) ⬆️

... and 8 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@github-actions

github-actions Bot commented Aug 17, 2026 •

Copy link
Copy Markdown

🔴 Benchmark Regression Detected

Comparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).

ℹ️ Absolute times are informational only — they vary with runner load. The % change column is the reliable signal because both sides ran under identical conditions.

Status Scenario Base PR Change
🔴 matmul/medium_generic_f32_threads=8/32x512x512 1.14 ms 1.89 ms +65.1%
🔴 block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 66.15 µs 100.25 µs +51.5%
🔴 block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 58.59 µs 81.98 µs +39.9%
⚠️ matmul/medium_generic_bf16_threads=8/32x512x512 452.68 µs 558.60 µs +23.4%
⚠️ matmul/large_generic_bf16_threads=8/32x1024x1024 1.54 ms 1.83 ms +19.0%
⚠️ qwen3_sampling_processors/top_p_fast_after_top_k 536.45 µs 632.70 µs +17.9%
⚠️ matmul/small_generic_f16_threads=1/1x256x256 35.77 µs 41.74 µs +16.7%
✅ tokenization/encode_tokens_per_second 435.70 µs 499.17 µs +14.6%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 3.99 ms 4.54 ms +13.8%
✅ matmul/small_generic_f32_threads=1/1x256x256 40.27 µs 45.80 µs +13.7%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 5.85 ms 6.59 ms +12.7%
✅ matmul/small_generic_f32_threads=8/1x256x256 52.61 µs 57.47 µs +9.2%
✅ sampling_latency/min_p_per_token 227.02 µs 246.14 µs +8.4%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 98.87 µs 105.17 µs +6.4%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 627.70 µs 647.54 µs +3.2%
✅ qwen3_sampling_processors/top_k_partial_selection 167.53 µs 171.69 µs +2.5%
✅ gather/large_bf16_threads=1-internal/131072 14.35 µs 14.68 µs +2.3%
✅ logit_processing/seven_processor_chain_per_step 345.56 µs 349.35 µs +1.1%
✅ kv_cache/alloc_dealloc_pages 39.93 µs 40.24 µs +0.8%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 9.76 ms 9.71 ms -0.6%
✅ matmul/medium_generic_f32_threads=1/32x512x512 3.22 ms 3.17 ms -1.3%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 2.18 ms 2.13 ms -2.3%
✅ qwen3_sampling_processors/top_k_top_p_fast 741.23 µs 722.20 µs -2.6%
✅ tokenization/decode_tokens_per_second 7.66 ms 7.46 ms -2.6%
✅ matmul/medium_generic_f16_threads=8/32x512x512 41.86 µs 40.56 µs -3.1%
✅ block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 604.47 µs 583.35 µs -3.5%
✅ gather/medium_f16_threads=1-internal/32768 3.05 µs 2.93 µs -4.0%
✅ gather/small_f16_threads=1-internal/4096 522.1 ns 497.9 ns -4.6%
✅ gather/small_bf16_threads=1-internal/4096 606.7 ns 570.2 ns -6.0%
✅ block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 851.98 µs 799.06 µs -6.2%
✅ gather/medium_bf16_threads=1-internal/32768 2.86 µs 2.67 µs -6.6%
✅ gather/large_f32_threads=1-internal/131072 38.10 µs 35.21 µs -7.6%
✅ matmul/medium_generic_f16_threads=1/32x512x512 39.91 µs 36.87 µs -7.6%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 2.35 ms 2.16 ms -8.1%
✅ sampling_latency/top_p_per_token 488.28 µs 448.83 µs -8.1%
✅ gather/large_f16_threads=1-internal/131072 14.58 µs 13.33 µs -8.6%
✅ block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 218.84 µs 199.00 µs -9.1%
✅ add/large_bf16_threads=1-internal/4194304 2.47 ms 2.19 ms -11.5%
✅ reduce_mean/medium_f32_threads=1-internal/65536 317.62 µs 279.87 µs -11.9%
✅ add/small_bf16_threads=1-internal/1024 549.3 ns 482.7 ns -12.1%
✅ sampling_latency/greedy_per_token 4.09 µs 3.59 µs -12.3%
✅ reduce_mean/small_f32_threads=1-internal/4096 19.49 µs 16.73 µs -14.2%
🟢 add/large_f32_threads=1-internal/4194304 1.15 ms 959.98 µs -16.6%
🟢 add/medium_bf16_threads=1-internal/262144 164.18 µs 136.16 µs -17.1%
🟢 grammar_masking/llguidance_compute_mask/32 98.40 µs 81.29 µs -17.4%
🟢 add/small_f32_threads=1-internal/1024 273.7 ns 225.3 ns -17.7%
🟢 matmul/large_generic_f16_threads=1/32x1024x1024 114.58 µs 93.97 µs -18.0%
🟢 sampling_latency/top_k_per_token 80.93 µs 64.19 µs -20.7%
🟢 matmul/large_generic_f32_threads=8/32x1024x1024 5.78 ms 4.44 ms -23.2%
🟢 add/medium_f16_threads=1-internal/262144 142.94 µs 108.19 µs -24.3%
🟢 add/large_f16_threads=1-internal/4194304 3.39 ms 2.56 ms -24.6%
🟢 matmul/small_generic_f16_threads=8/1x256x256 46.45 µs 34.72 µs -25.3%
🟢 add/small_f16_threads=1-internal/1024 677.8 ns 499.1 ns -26.4%
🟢 matmul/small_generic_bf16_threads=1/1x256x256 44.49 µs 32.68 µs -26.5%
🟢 reduce_mean/large_f32_threads=1-internal/262144 1.53 ms 1.10 ms -27.7%
🟢 matmul/small_generic_bf16_threads=8/1x256x256 58.74 µs 42.38 µs -27.9%
🟢 add/medium_f32_threads=1-internal/262144 36.76 µs 25.99 µs -29.3%
🟢 gather/small_f32_threads=1-internal/4096 1.06 µs 734.3 ns -30.8%
🟢 gather/medium_f32_threads=1-internal/32768 6.49 µs 4.42 µs -31.8%

Visual flags: ⚠️ ≥ 15% slower, 🔴 ≥ 30% slower — calibrated against measured runner noise (~27% worst-case on multi-threaded matmul)

Host info
CPU: Apple M1 (Virtual)
Cores: 3
OS: Darwin 25.5.0 arm64
Rust: rustc 1.97.1 (8bab26f4f 2026-07-14)
Load avg: { 4.46 3.68 5.10 }
What this cannot catch
  • Regressions in code paths not covered by these benchmarks (e.g., end-to-end decode with a real model)
  • Sub-threshold regressions that compound over multiple PRs
  • Performance changes that only manifest under GPU execution
  • Latency changes in the ORT integration path (these benchmarks exercise the native Rust kernels)

Deleting the performance-based decline policy was not enough. `GetCapability`
runs a second, independent fail-closed filter: it drops any claim containing a
node whose `ShapeInference::for_node` returns `Declined`, and that table ends in
`_ => Declined`. An op we register a kernel for but forgot to list there is
handed to ORT regardless of what `supports_op` says — the same mechanism that
made the `com.microsoft` activations unreachable until #1082.

The trigonometric, hyperbolic and remaining activation ops were all in that gap.
Add them to the `SameAsInput(0)` arm: Sin, Cos, Tan, Asin, Acos, Atan, Sinh,
Cosh, Asinh, Acosh, Atanh, ThresholdedRelu, Swish and `com.microsoft` Silu.
`Silu` was previously excluded on the grounds that ORT has no kernel for it,
which had it exactly backwards: an op ORT cannot run is precisely the one we
must never hand over. Add PRelu to the broadcast arm, whose slope is
unidirectionally broadcastable to the input.

Add a `Sin` assignment fixture and extend the sweep to cover it. Verified the
falsifier actually falsifies: with `| "Sin"` removed from the table the sweep
fails with `[sin_assignment_f32] ours=[], others=["Sin"]`.

Document the 66 registered ops still in the gap, split into the ones correctly
declined as data-dependent, the internal fusion ops that never appear in an
input graph, and the ~55 that are inferrable but unwritten.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby justinchuby changed the title Never hand a node to ORT's CPU EP Stop handing activation nodes to ORT's CPU EP Aug 17, 2026
The per-op assignment fixtures cannot catch an op that is registered but
missing from `ShapeInference::for_node`: they only prove the ops they name are
assigned, and a newly-registered op has no fixture by construction. That is
exactly the bug #1082 and this PR both had to fix by hand.

Add `every_registered_op_has_a_shape_rule_or_is_a_known_gap`, which enumerates
the CPU registry and asserts the set of declining ops *exactly*. Registering an
op without a shape rule fails; adding a shape rule without removing the op from
the list also fails. Neither direction can pass silently. It runs without ORT,
so it is covered by every job, not only the ORT-gated one.

Verified both directions falsify: deleting `| "Cos"` fails with
`[("", "Cos")]`; adding `| "Trilu"` fails with `these ops now have a shape rule
but are still listed as declined: [("", "Trilu")]`.

The inventory corrects the counts I published. My scratch script text-matched
the table and so missed its guard arms — `op_name if is_reduction(op_name)` and
the pooling arms — which made it report ops as gaps that are in fact handled.
The real figure is 52, not 66: the whole `Reduce*` family, `MaxPool`,
`AveragePool`, the `Global*Pool` family, `Resize`, `GroupNormalization`,
`ConvTranspose`, `Col2Im`, `GridSample`, `AffineGrid`, `SpaceToDepth`,
`CenterCropPad` and `LpPool` all have rules already.

It also found three real gaps I had missed and reclassified two I had wrong:

- `Unsqueeze` declines whenever `axes` arrives as input[1], which is every
  opset-13-or-later graph. A frequent op, handed to ORT.
- `com.microsoft::Attention` declines: the opset-23 arm is guarded to the
  default domain, and the contrib op's packed-QKV signature needs its own rule.
  This is *the* attention op in exported GenAI models.
- `EyeLike` was absent from my list entirely.
- `MoE` and `QMoE` are not internal fusion ops. They are read from exported
  models — `deepseek_v2_tiny_qmoe_native_e2e.rs` asserts a loaded graph
  contains a `com.microsoft::QMoE` node — so they belong in the unwritten
  group, not the never-a-candidate group.
- `PackedMultiHeadAttention` is likewise a contrib op from real models.

Rewrite the doc's gap section against the measured list, and add
`no_activation_or_norm_op_is_left_to_ort` as a standing guard on the 41
activation and norm ops this EP owns.

Also bring `for_op_domain`'s name-only table into line with `for_node` and
delete its `Silu` comment, which still argued that an op ORT cannot run is one
we should decline — the reasoning this PR reverses.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby and others added 2 commits August 17, 2026 02:04
The inventory I added had the blind spot it was written to prevent.
`build_cpu_registry_with_descriptors` does not describe the registry:
`register_cnn_ops` writes straight to the inner `OpRegistry`, so `MaxPool`,
`AveragePool`, the `Global*Pool` family, `LpPool`, `ConvTranspose`, `Resize`,
`GridSample`, `SpaceToDepth`, `GroupNormalization`, `AffineGrid`, `Col2Im`,
`CenterCropPad` and `PRelu` never appear in the descriptors. `supports_op` keys
off the registry, so every one of them was claimed, dropped by the shape
filter, and executed on ORT — invisible to the test.

Enumerate `OpRegistry::keys()` instead. That is the same set `supports_op`
consults, which is what makes it the right source of truth.

Sweep the probe rather than fixing it at one point. A bare one-input rank-2
node reported `Conv` as declining, but `build_conv` reads `input_shapes[1][0]`
for its output channel count and needs rank >= 3 — a probe artifact, not a gap.
Sweep opsets {1,13,18,22,23} x arities 1..4 x ranks 1..4 and count an op as
declining only when nothing in the matrix produces a rule. `Conv` drops out;
the 14 genuine CNN gaps do not.

The real figure is 66, and my previous "correction" was itself wrong in the
other direction: I claimed the pooling family "has rules already", but
`compute.rs` has no pooling arm at all. Only `Reduce*` was mis-listed
originally, via the `is_reduction` guard arm. Rewrite the doc against the
measured list: 20 data-dependent, 10 internal, 36 inferrable-but-unwritten.

Reclassify `LinearAttention` (both domains) and `com.microsoft::CausalConvWithState`
out of the internal-fusion group. They are the Qwen3.5 / Qwen3-Next hybrid
linear-attention primitives, read from exported models — the same mistake I
made with MoE/QMoE. ORT has no kernel for them at all, so declining them does
not get a faster implementation, it gets a load failure.

Assert `no_activation_or_norm_op_is_left_to_ort` in both directions. Filtering
on `registered.contains(op) && declines(op)` let a deleted registration pass
silently, which is the same hand-off by a different route: renaming the `Silu`
kernel left the test green. Now it fails with `these ops are no longer
registered by the CPU EP, so ORT will execute them: [("com.microsoft",
"Silu")]`. That direction also caught two entries I had wrong — `Swish` is
registered in the default domain, not `com.microsoft`, and `HardSwish` has no
kernel at all.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
# Conflicts:
#	docs/performance/CPU_MATMUL_ASSIGNMENT.md
@justinchuby

Copy link
Copy Markdown
Owner Author

Matmul owner, ack — and a finding that changes what "never defer" costs.

While writing the matmul-family half of the conformance sweep this PR asks for, the four matmul ops turned out not to have been reachable in the first place. com.microsoft::MatMulNBits and QLinearMatMul have never been claimed by this EP through the plugin path, whatever assignment_policy said:

  • supported_dtypes_for_op("MatMulNBits", "com.microsoft") was F32_ONLY, and node_passes_dtype_filter requires every edge dtype to be listed — the uint8 packed weight fails on input 1.
  • ShapeInference::for_node aliased MatMulNBits to Self::MatMul, which matmuls the activation against the packed [N, blocks, blob] weight: [1,256] × [4096,8,16] → [4096,1,16].
  • QLinearMatMul was in neither shape table, so it resolved to Declined and the fail-closed shape filter dropped the claim — exactly the failure mode fix(ep-plugin): let com.microsoft activations reach the EP, then assign them honestly #1082 fixed for the com.microsoft activations.
  • input_slots numbered ORT inputs positionally, but ORT's fused-node metadata is a set, so any node naming one value twice ran off the end of the bound array. Mul(X, X) on a plain graph input reproduces it.

Fixes plus a 10-case generated-model sweep are in #1101, based on main so it can land independently of this PR.

Two implications for this one:

  1. The ASSIGNMENT_FIXTURES list here is all unary/activation — 20 fixtures, no matmul family. fix(ep): claim and correctly execute MatMulNBits and QLinearMatMul #1101 adds the matmul half separately, so between the two the family is covered.
  2. Removing the defer gate is necessary but not sufficient. claim_preference_node is only consulted when host_fallback_available, so with session.disable_cpu_ep_fallback=1 the gate was already off — and MatMulNBits still escaped, through the dtype and shape filters instead. A sweep that only runs with fallback disabled cannot see a defer regression; one that only checks the assignment record cannot see a filter regression. Worth running each fixture in both modes here.

Same class of bug, not mine to fix — two more supported_dtypes entries that look like they exclude their own op:

  • ("MoE" | "QMoE", "com.microsoft") => F32_ONLY — both carry uint8 weight edges.
  • ("GatherBlockQuantized", "com.microsoft") => FLOAT_DTYPES — same.

If either is meant to be claimed today, it almost certainly is not being. #1101 adds KernelRegistryEntry::input_dtype_constraints for the per-slot case, which is what these need if their kernels are narrower than the union of their edge dtypes (e.g. MatMulNBits accepts only uint8 zero points while the op spec permits float16, so the union alone claims nodes the kernel then fails on).

Review found a filter neither of the previous two rounds looked at.
`GetCapability` runs three independent fail-closed checks, not two, and
a claim has to clear all of them. This PR had fixed the assignment
policy and the shape table; the dtype filter was still declining ops.

`node_passes_dtype_filter` looks the node's op up in the plugin's
`KernelRegistryEntry` list and returns false when there is no entry.
That list is built from `build_cpu_registry_with_descriptors`, which
recorded keys as they were registered -- but `register_cnn_ops` takes
`&mut OpRegistry` and writes past the recording wrapper. Eighteen ops
were in the registry and absent from the descriptors, so `supports_op`
claimed each one and capability then dropped it:

  PRelu, BatchNormalization, InstanceNormalization, GroupNormalization,
  Conv, ConvTranspose, MaxPool, AveragePool, GlobalMaxPool,
  GlobalAveragePool, GlobalLpPool, LpPool, Resize, GridSample,
  AffineGrid, Col2Im, CenterCropPad, SpaceToDepth

Four of those are activations or normalisations this EP is supposed to
own. `PRelu` is the sharpest case: the previous commit gave it a shape
rule, so it cleared filter two, and the dtype filter declined it anyway.
Both pure-Rust inventory tests passed while real ORT ran the node.

Derive descriptors from `OpRegistry::keys()` instead of a parallel
recorded list, so the two sets are identical by construction rather than
by convention, and sort them -- they are leaked into a `'static` slice
ORT reads, and hash-map order would make any snapshot diff flap.

`GroupNormalization` needed the other half too: it had no shape rule.
Added to `SameAsInput(0)` in both tables and removed from the allowlist.

The existing `descriptors_derived_from_real_registry_not_hand_maintained`
had been asserting this bug as correct behaviour -- it allowed a delta of
up to 50 entries and named CNN ops as the expected difference. It now
asserts set equality and names the ops it finds missing.

Three new guards, each verified to fail when its invariant is broken:

  every_registered_op_has_a_kernel_registry_entry -- filtering PRelu out
  of the descriptors gives `1 registered ops have no kernel-registry
  entry ... ["::PRelu"]`.

  activation_and_norm_ops_clear_every_capability_filter -- checks the
  ops this EP owns against registration, descriptors and shape rules
  together, because clearing one filter and failing another is exactly
  how PRelu stayed broken. Writing it immediately found two more real
  gaps: Celu and Mish have no kernel at all. They are excluded with a
  comment and recorded in the gaps doc rather than quietly dropped.

  prelu_assignment_f32 and groupnorm_assignment_f32 -- real ORT
  end-to-end, reading back the node's assigned EP. Before the fix ORT
  reports `ours=[], others=["PRelu"]`; after, `ours=["PRelu"], others=[]`.
  This is the only check that cannot be fooled by enumerating the wrong
  set in Rust, which is how the last two rounds passed.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 17, 2026
…1101)

## The finding

Our CPU EP has **never executed `MatMulNBits` or `QLinearMatMul`**
through the ORT plugin path. Not "deferred them under policy" — never
claimed them at all. Every int4/int8 quantized model that has run
through this EP had its matmuls executed by ORT's own kernels, silently,
regardless of what `assignment_policy` said and regardless of
`session.disable_cpu_ep_fallback`.

That invalidates the reachability half of every plugin-path claim in
`CPU_MATMUL_ASSIGNMENT.md`. It also means the recent architectural
direction — *this EP must never hand a node to ORT's CPU EP* — could not
have been satisfied by removing the defer matrix alone.

There are four independent defects, each of which is sufficient on its
own to lose the op.

### 1. The dtype filter excluded `MatMulNBits` outright

`node_passes_dtype_filter` (`ep.rs`) returns `false` unless **every**
input and output dtype of the node appears in `entry.supported_dtypes`.
`supported_dtypes_for_op("MatMulNBits", "com.microsoft")` returned
`F32_ONLY`.

`MatMulNBits` inputs are `A: float`, `B: uint8`, `scales: float`,
`zero_points: uint8`, `g_idx: int32`. The `uint8` weight fails the
filter on input 1, so the claim is dropped before anything else is
consulted. The correct set is the one the kernel itself enforces at
`matmul_nbits.rs:780-815`: float32/float16/bfloat16 on `A`/`scales`/`Y`,
`uint8` on `B` and `zero_points`, `int32` on `g_idx`.

### 2. `MatMulNBits` used the plain-matmul shape rule

`ShapeInference::for_node` mapped `"MatMul" | "MatMulNBits"` to
`Self::MatMul`. `B` is not a matmul operand — it is the packed `[N,
blocks_per_col, blob_bytes]` weight — so the rule broadcast the
activation against the packing:

| K | N | inferred | correct |
|---|---|---|---|
| 256 | 4096 | `[4096, 1, 16]` | `[1, 4096]` |
| 64 | 8 | `[8, 1, 16]` | `[1, 8]` |

The correct rule is the activation shape with its last dim replaced by
the `N` **attribute**, which is what `onnx-runtime-shape-inference`'s
`quantized_matmul` (`handlers/linalg.rs:36`) already does. Because `N`
is only visible on the node, the attribute-free `for_op_domain` table
must decline rather than fall back to the matmul rule.

Defect 1 masked defect 2: fixing only the dtypes turns a silent escape
into a hard `Y must have shape [1, 4096], got [4096, 1, 16]`.

### 3. `QLinearMatMul` was in neither shape table

Missing from `for_op_domain` and `for_node`, so it resolved to
`ShapeInference::Declined` and `GetCapability`'s fail-closed shape
filter dropped the whole claim. This is the same failure mode as the
`com.microsoft` activations fixed in #1082.

Note that the `MatMul` rule would have been wrong here too: the ONNX
operand order is `a, a_scale, a_zero_point, b, b_scale, b_zero_point,
y_scale, y_zero_point`, so the operands are inputs **0 and 3**.

### 4. `input_slots` numbered ORT inputs by position, not by value

Once claimed, `qlinear_u8` still failed:

```
Compute: input slot 7 maps to ORT input 7, but ORT bound only 7 input(s) to this fused node
```

ORT's fused-node metadata carries a **set** of input names. A node that
names the same value twice is bound once, and the sequential numbering
shifted every later slot — the last one past the end of the array ORT
actually passes. This is not exotic: a real quantized `QLinearMatMul`
routinely shares one zero-point or scale initializer across its
`a`/`b`/`y` triples, and any `Mul(x, x)`-shaped node does the same. On
`main` this surfaced as an opaque `Compute: internal panic` from an
out-of-bounds index across the C ABI.

Slots are now assigned per distinct value in first-appearance order, and
the fast path fails closed with the message above instead of panicking.

### 5. The union dtype list over-claimed what the kernels accept

Found by review of the first commit. `supported_dtypes` is a flat set
and the node filter tests membership in it, so on a mixed-dtype op it is
strictly weaker than the kernel's own rule. `com.microsoft::MatMulNBits`
spec-allows `zero_points` in `uint8`, `int32`, `float16` or `float`
(`contrib_defs.cc` T3), and **ORT builds a session for every one of
them** — I probed all three against a plain-ORT session and all three
loaded. Our kernel implements packed `uint8` only. Fixing defect 1 by
advertising the union therefore claimed a float16-zero-point node and
then failed inside `execute` with `zero_points must have dtype Uint8` —
turning a model that previously produced an answer into a hard error.

`KernelRegistryEntry` gains `input_dtype_constraints: &'static [(usize,
&'static [DataType])]`. A listed slot is checked against its own set
instead of the union; unlisted and absent slots keep the union, which is
exact for every uniform op. Populated straight from the kernels'
`require_dtype` calls:

| op | slot | allowed |
|---|---|---|
| `MatMulNBits` | 0 `A`, 2 `scales`, 5 `bias` | f32 / f16 / bf16 |
| | 1 `B`, 3 `zero_points` | uint8 |
| | 4 `g_idx` | int32 |
| `QLinearMatMul` | 1 `a_scale`, 4 `b_scale`, 6 `y_scale` | f32 |
| | 0, 2, 3, 5, 7 | uint8 / int8 |

The `QLinearMatMul` half matters for the same reason: its scales flow
through `ARITH_DTYPES`, which admits f16/bf16, and opset 21 spec-allows
those — so defect 3's fix would newly have claimed a config
`qlinear_matmul.rs:316` rejects.

### Remaining gap, stated plainly

Float / float16 / bfloat16 `zero_points` on `MatMulNBits` is a
**quantization semantic this EP does not implement**. It is declined at
`GetCapability`, and
`a_float_zero_point_matmul_nbits_is_declined_not_claimed_and_failed`
pins that it is declined *and still computed* rather than claimed and
failed. That is a genuine coverage gap in the matmul family, recorded
here rather than hidden: closing it means implementing the unpacked
float zero-point dequant path, which is a new numeric semantic and does
not belong in an urgent claim fix.

## Falsifiers

`no_matmul_family_node_escapes_to_the_ort_cpu_ep` builds ten models from
scratch, runs each under a session with
`session.disable_cpu_ep_fallback=1`, asserts via ORT's own
`record_ep_graph_assignment_info` that the op is on `cpu_ep` and that
`ops_not_on_our_ep()` is empty, then compares elementwise against a
second session with **no** EP appended — i.e. pure ORT CPU. A
non-degeneracy guard requires at least half the outputs to be non-zero
and finite first, so a both-sides-zero comparison cannot pass vacuously.

| case | shape | worst \|ours − ORT\| | tolerance |
|---|---|---|---|
| `matmul_f32_decode` | M1 K256 N4096 | 4e-6 | 1e-3 |
| `matmul_f32_prefill` | M256 K256 N4096 | 5e-6 over 1,048,576 elems |
1e-3 |
| `matmul_f16_decode` | M1 K256 N4096 | 4.88e-4 | 5e-2 |
| `gemm_f16_decode` | M1 K256 N4096 | 4.88e-4 | 5e-2 |
| `nbits4_decode` | bits4 blk32 M1 K256 N4096 | 1.83e-4 | 0.5 |
| `nbits8_prefill` | bits8 blk32 M256 K256 N4096 | 2.93e-3 over
1,048,576 elems | 0.5 |
| `nbits4_block512` | bits4 blk512 M1 K512 N2048 | ours runs; ORT
refuses the model | — |
| `nbits4_f16_decode` | bits4 blk32 M1 K256 N4096, f16 activation | 0.0
| 4.0 |
| `qlinear_u8` | M1 K256 N4096 | 0.0 (bit-exact) | 1 |
| `qlinear_i8` | M1 K256 N4096 | 0.0 (bit-exact) | 1 |

Two further cases outside the sweep: `mul_self` (defect 4) and
`nbits_float_zp` (the declined-not-crashed guard).

`nbits4_block512` is the negative control: ORT rejects block_size 512 at
kernel construction (`matmul_nbits.cc:131`), we run it, and the test
asserts ORT refuses so the claim in the assignment doc cannot go stale
silently.

`nbits4_f16_decode` exists because fixing defect 1 *widens* what we
claim — f16 activations were unreachable before, and this pins that they
are now correct rather than merely accepted.

`a_node_that_names_one_value_twice_is_bound_once` is the direct
falsifier for defect 4: `Mul(X, X)` on a **graph input**, no
initializers, no constant-sharing pass involved. It asserts the node
lands on our EP and that every element equals `x*x` exactly.

### Each fix falsified independently

Reverting one source file at a time on this branch, with the test
unchanged:

| reverted | first failure |
|---|---|
| all three | `nbits4_decode` — session creation fails, `ours=[]`,
`others=["MatMulNBits"]` |
| `compute.rs` only (shape rules) | `nbits4_decode` — `MatMulNBits: Y
must have shape [1, 4096], got [4096, 1, 16]` |
| `ep.rs` only (slot dedup) | `qlinear_u8` — `input slot 7 maps to ORT
input 7, but ORT bound only 7 input(s)` |
| slot mapping only, back to positional | `mul_self` — `input slot 1
maps to ORT input 1, but ORT bound only 1 input(s)` |

Plus seven hardware-independent unit tests: four in `compute.rs` pinning
the two new shape rules — including an explicit `assert_eq!(wrong,
vec![vec![4096, 1, 16]])` so the plain-`MatMul` alias cannot return
unnoticed — two in `kernels/mod.rs` asserting the quantized edge dtypes
are advertised, and one in `ep.rs` asserting the per-slot filter
declines a float16-zero-point node and a uint8-activation node while
accepting the packed-uint8 one.

## Performance

**No performance claim.** This PR changes nothing about the kernels; it
changes which kernels run at all. It does not make anything faster, and
the ratios in `CPU_MATMUL_ASSIGNMENT.md` are not restated here — see
below.

## What this does to the published matrix

The `MatMulNBits` and `QLinearMatMul` rows in `CPU_MATMUL_ASSIGNMENT.md`
were measured with `bench_generic`, which drives the kernels
**natively**, so those numbers remain valid as kernel-level
measurements. What is now known to be wrong is the implication that the
plugin path ever reached those kernels. Re-stating the matrix as
plugin-path numbers is deliberately **not** in this PR — it is a doc
change that depends on #1097's rewrite of the same file, and folding it
in here would make an urgent correctness fix conflict with an open
architectural PR.

## Related, not fixed here

Same class of bug, outside this PR's remit, reported to the owners
rather than fixed silently:

- `("MoE" | "QMoE", "com.microsoft") => F32_ONLY` — both carry `uint8`
weight edges.
- `("GatherBlockQuantized", "com.microsoft") => FLOAT_DTYPES` — same.

Both are very likely dropped by the dtype filter for the same reason as
defect 1.

## Verification

- `cargo test -p onnx-runtime-ep-cpu-plugin --test plugin_ort_e2e` — 50
passed (was 47; +3)
- `cargo test -p onnx-runtime-ep-plugin --lib` — 229 passed (was 224;
+5)
- `cargo test -p onnx-runtime-ep-cpu --lib` — 1317 passed; with
`--features mlas` 1334 passed (+2 dtype tests)
- `cargo test -p onnx-runtime-ep-shared-mock-plugin` — passes
(registry-entry field addition)
- `cargo fmt --all -- --check` clean
- `cargo clippy --all-targets` clean **both** with and without
`--features mlas`

Host: AMD EPYC 9V74, 32 vCPU, AVX2+FMA+F16C, ORT 1.27.0.


## Independent review

Reviewed by an independent Opus reviewer against a separate worktree,
with the ORT 1.27.0 source for `partitioning_utils.cc` pulled at the
prebuilt's own commit. Verdict was REQUEST CHANGES on the over-claim
above; all findings are addressed in `905f7c47c`:

| finding | resolution |
|---|---|
| major: float zero-point `MatMulNBits` claimed then fails at Run |
per-slot constraints; declined; regression test |
| major: f16/bf16-scale `QLinearMatMul` claimed then fails at Run | same
mechanism, `QLinearMatMul` slot table |
| minor: sweep does not exercise the slot dedup | it does — reverting
only the slot mapping fails `qlinear_u8` — but a direct `Mul(X, X)` case
was added anyway, since that one covers a **graph input** rather than an
initializer |
| nit: line-wrapped error string | fixed |
| nit: "nine models" / "the two noted" | fixed |

Round two returned APPROVE. The reviewer disabled
`input_dtype_constraints_for_op` for `MatMulNBits` and confirmed the
float-zero-point test then fails, i.e. the per-slot list is load-bearing
and nothing else catches it; wrote an independent opset-21 f16-scale
`QLinearMatMul` probe and confirmed it is declined and still computed;
reverted the slot mapping to positional and confirmed both `mul_self`
and `qlinear_u8` fail; and checked all fourteen constrained slots
against the two kernels' `require_dtype` calls. One useful correction:
for `QLinearMatMul` the per-slot list is redundant with a pre-existing
native `unsupported_reason` gate, so it is defence in depth there and
load-bearing only for `MatMulNBits`.

The reviewer independently confirmed from ORT's `MakeComputeCapability`
that the fused-node input set is deduplicated by `NodeArg` identity in
first-appearance order, over graph inputs and initializers alike, and
that the multi-node routing path already keys on value position so it
stays consistent with the fast path.


## Post-review fix

`f8f7a486a` — the generated-model element-type constants were annotated
`u32`. `ONNXTensorElementDataType` is a `c_uint`, which bindgen renders
as `u32` on Linux but `i32` on Windows, so `Rust (Windows ARM64)` and
`Rust coverage (Windows x86_64)` failed to compile the test file while
every Linux lane was green. They now use ORT's own alias. Test-only, no
behaviour change.

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby and others added 3 commits August 17, 2026 04:43
Giving `Conv` a `KernelRegistryEntry` in the previous commit was itself
an over-claim. `supported_dtypes_for_op("Conv", "")` returned
`FLOAT_DTYPES`, which includes f64, but `ConvKernel::execute` rejects
anything outside f32/f16/bf16 -- MLAS computes in f32 and widens or
narrows around it. Before that commit `Conv` had no descriptor at all,
so an f64 `Conv` never reached us; after it, the node cleared the shape
filter via `build_conv` and the dtype filter via f64's presence in the
list, got claimed and compiled, and then failed at `Run`.

Add `MLAS_FLOAT_DTYPES` and give `Conv` its own arm. The rest of the CNN
family genuinely dispatches f64 through `dispatch_float!` and keeps
`FLOAT_DTYPES`.

This is not the fallback policy coming back. Declining a dtype we can
execute and declining one we cannot are different things: f64 `Conv` is
genuinely unsupported, and saying so at capability time is the honest
answer -- the alternative is a hard failure mid-session.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
`main` (#1101) independently hit the same dtype-filter bug class for
`MatMulNBits`/`QLinearMatMul` and added `FLOAT_COMPUTE_DTYPES` — the
identical f32/f16/bf16 set as this branch's `MLAS_FLOAT_DTYPES`. Resolved
`kernels/mod.rs` by taking main's block verbatim, dropping the duplicate
constant, pointing `Conv` at `FLOAT_COMPUTE_DTYPES` and folding this
branch's rationale into main's doc comment.

The merge also surfaced two shape-table drifts, both caught by
`every_registered_op_has_a_shape_rule_or_is_a_known_gap`:

* The probe supplied no attributes, so `com.microsoft::MatMulNBits` —
  whose rule derives the output width from `N` — looked like a gap it is
  not. A node without `N` is malformed rather than defaulted, so no real
  graph hits that path. The probe now sweeps an attribute bundle
  alongside the empty one, keeping default-fallback rules honest while
  modelling production for rules that cannot default.

* With that fixed, the first assert stopped masking the second:
  `QLinearMatMul` gained a rule in main and was stale in `DECLINED`.
  Removed. Gap count 65 -> 64, group 3 35 -> 34.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@justinchuby
justinchuby marked this pull request as ready for review August 17, 2026 10:34
@justinchuby
justinchuby merged commit a2c97e7 into main Aug 17, 2026
12 of 18 checks passed
@justinchuby
justinchuby deleted the deckard/no-ort-fallback branch August 17, 2026 10:34
justinchuby added a commit that referenced this pull request Aug 17, 2026
…1121)

## What this is

`tanh_ps` and `sigmoid_ps` are the two AVX2 primitives most of the CPU
activation family is built on: `Tanh`, `Sigmoid`, SiLU / Swish,
`QuickGelu` and
`FastGelu` all end up in one of them. Both already carry the Eigen
rational
that MLAS uses — the polynomial was never the problem. What they also
carried
was a redundant saturation step, and it is not free.

The shape of both kernels is: clamp the *input* to the rational's valid
range
(`±9` for `tanh`, `±18` for `logistic`), evaluate `P(v²)·v / Q(v²)`,
clamp the
*output* to the function's mathematical range. On top of that, both then
ran a
second saturation:

```rust
let above = _mm256_cmp_ps(x, _mm256_set1_ps(tanh_c::UPPER), _CMP_GT_OQ);
let below = _mm256_cmp_ps(x, _mm256_set1_ps(tanh_c::LOWER), _CMP_LT_OQ);
let r = _mm256_blendv_ps(poly, _mm256_set1_ps(1.0), above);
_mm256_blendv_ps(r, _mm256_set1_ps(-1.0), below)
```

Four vector ops, two of them `vblendvps`, which is two uops on Zen.

**It never changes a bit.** The input clamp happens *first*, so the
largest
argument the rational ever sees is the clamp point itself — and there it
has
already reached the output bound:

| | at the clamp point | output clamp gives |
|---|---|---|
| `tanh`, `v = 9` | `p` and `q` are **bit-equal** (`0x3fcd33e9`), so
`p/q` is exactly `1.0` | `1.0` |
| `tanh`, `v = -9` | exactly `-1.0` | `-1.0` |
| `logistic`, `v = 18` | `p/q + 0.5` is exactly `1.0` | `1.0` |
| `logistic`, `v = -18` | `p/q + 0.5` is `-5.96e-8` | `0.0` |

Three of the four are *equality*, not slack, and the inclusive
`minps`/`maxps`
clamp passes equality through unchanged. (The rational's overshoot to
`1.0000001` happens strictly *inside* the range, near `|v| = 8.9999971`
— that
is what the output clamp is there for, and it is unrelated to
saturation.)

Every input past the range, `±Inf` included, therefore already exits
with the
saturated value. The blends were re-deriving a result the clamp had
produced two
instructions earlier.

## Proof, not argument

`saturation_blend_is_redundant_exhaustively` keeps the deleted sequence
as a
reference implementation and checks the shipped kernels against it over
**every
finite `f32` beyond the clamp, both signs** — 1 047 527 424 values for
`tanh`
and 1 039 138 816 for `sigmoid`, 2 086 666 240 in total — plus `±Inf`,
both
signed `NaN`s, `f32::MAX`/`MIN`, and both clamp boundaries at ±2 ULP.

Bit-identical everywhere. Not within tolerance — identical.

The sweep runs under the default rounding mode. The property was also
checked
by hand under round-down, round-up and round-toward-zero, and holds in
all
three; only round-to-nearest is exercised in CI.

It runs in 19 s in release and is `#[ignore]`d because an unoptimised
test build
takes minutes. Two non-ignored tests cover the decade above each
boundary
(32.5 M and 21.7 M values) and the special values, and run in 12 s in
the
default `cargo test` profile, so the property stays guarded on every CI
run.

## Also

The three `vblendvps`-against-zero that pin `-Inf → 0` in
`tanh_gelu_ps`,
`quick_gelu_ps` and `erf_gelu_ps` become `vandnps`. Blending against
zero *is*
an `andnot` when the mask is all-ones or all-zeros, which is exactly
what
`vcmpps` produces. One uop instead of two, same bits.

## Measured

Same-machine, alternating build A/B through a **real ORT session** — one
node,
one input, `intra_op_num_threads=1`, `RAYON_NUM_THREADS=1`, p50 of 30
runs,
median of 3 alternating rounds per build, with our EP's assignment
asserted from
ORT's own profiler on every row. AMD EPYC 9V74, AVX2+FMA, ORT 1.28.0.

Measured on top of #1097 and #1105, because on `main` today the
assignment
policy declines these ops to ORT's CPU EP, so our kernel never runs and
the
measurement is vacuous. (The first run of this A/B *was* vacuous for
exactly
that reason — both columns were ORT to within 0.1%. The harness's
assignment
check is what caught it.)

| op | elements | before, us | after, us | **speedup** |
|---|---|---|---|---|
| `Tanh` | 16384 | 13.49 | 11.71 | **1.15x** |
| `Tanh` | 65536 | 38.33 | 31.65 | **1.21x** |
| `Tanh` | 262144 | 131.87 | 105.48 | **1.25x** |
| `Tanh` | 1048576 | 533.18 | 438.21 | **1.22x** |
| `Tanh` | 4194304 | 2272.02 | 1889.22 | **1.20x** |
| `Sigmoid` | 16384 | 13.91 | 11.94 | **1.17x** |
| `Sigmoid` | 65536 | 39.96 | 35.27 | **1.13x** |
| `Sigmoid` | 262144 | 143.08 | 123.31 | **1.16x** |
| `Sigmoid` | 1048576 | 542.67 | 408.28 | **1.33x** |
| `Sigmoid` | 4194304 | 2052.81 | 1883.35 | **1.09x** |
| `FastGelu` | 65536 | 66.13 | 55.53 | **1.19x** |
| `FastGelu` | 262144 | 244.24 | 201.70 | **1.21x** |
| `FastGelu` | 1048576 | 969.63 | 782.34 | **1.24x** |
| `FastGelu` | 4194304 | 3994.51 | 3419.07 | **1.17x** |
| `QuickGelu` | 65536 | 49.36 | 42.27 | **1.17x** |
| `QuickGelu` | 262144 | 176.89 | 148.68 | **1.19x** |
| `QuickGelu` | 4194304 | 3016.00 | 2438.27 | **1.24x** |
| exact `Gelu` | 65536 | 111.12 | 108.26 | 1.03x |
| exact `Gelu` | 4194304 | 6898.34 | 6753.38 | 1.02x |
| `Erf` *(control)* | 65536 | 86.78 | 86.88 | 1.00x |
| `Erf` *(control)* | 1048576 | 1286.59 | 1286.52 | 1.00x |
| `Sqrt` *(control)* | 65536 | 21.63 | 21.38 | 1.01x |
| `Sqrt` *(control)* | 1048576 | 244.85 | 244.64 | 1.00x |

`Erf` and `Sqrt` touch neither primitive and are flat, which is the
check that
the rest is real and not drift. Exact `Gelu` goes through `erf_ps`, so
it only
collects the `andnot`, and moves ~2-3%.

## Against ORT

Same runs, ORT's CPU EP as the control, ORT time over ours — above
`1.00` we
win:

| op | 4096 | 16384 | 65536 | 262144 | 1048576 | 4194304 |
|---|---|---|---|---|---|---|
| `Tanh` | 0.73 -> 0.77 | 0.72 -> 0.83 | 0.72 -> 0.87 | 0.76 -> 0.95 |
0.88 -> **1.08** | 0.96 -> **1.16** |
| `Sigmoid` | 0.75 -> 0.81 | 0.73 -> 0.85 | 0.74 -> 0.84 | 0.75 -> 0.87
| 0.70 -> 0.92 | 0.73 -> 0.79 |
| `QuickGelu` | 0.83 -> 0.88 | 0.89 -> **1.03** | 0.97 -> **1.14** |
1.03 -> **1.23** | 1.05 -> **1.18** | 1.02 -> **1.26** |
| `FastGelu` | 0.73 -> 0.79 | 0.68 -> 0.76 | 0.67 -> 0.80 | 0.67 -> 0.81
| 0.66 -> 0.81 | 0.72 -> 0.84 |
| exact `Gelu` | 0.71 -> 0.72 | 0.67 -> 0.70 | 0.67 -> 0.69 | 0.69 ->
0.71 | 0.69 -> 0.71 | 0.71 -> 0.73 |

**This does not claim we now beat ORT.** It moves every one of these
families
toward it, takes `Tanh` past it at >=1 Mi and `QuickGelu` past it from
16 Ki up,
and leaves `FastGelu`, exact `Gelu` and `Erf` still behind. The
remaining gap is
not in these two primitives, and the next steps are elsewhere:
`erf_ps`'s
`exp_ps` tail, and the ~1.7-2.7 us fixed per-node plugin overhead that
dominates
at 4096 elements.

## Direction

This is an absorption in the sense of
`docs/performance/ABSORBING_MLAS.md`:
the win lands in the native kernel, in the **default build**, with no
`mlas`
feature and no dependency. Reading MLAS's `tanh.cpp` and `logistic.cpp`
is what
made the redundancy visible — MLAS does not have this step, because it
does not
promise the output range we promise, and comparing the two instruction
sequences
is what showed that our extra promise costs nothing to keep and four ops
to
re-state.

## Correctness

- `saturation_blend_is_redundant_exhaustively` — 2 086 666 240 values,
bit-identical.
- Two boundary sweeps + special values, in the default test profile.
- Full `onnx-runtime-ep-cpu` lib suite: 1326 passed.
- No tolerance was relaxed and no reference output changed; every
pre-existing
  `tanh`/`sigmoid`/`gelu` accuracy test passes unmodified.

## Limitations

- x86-64 AVX2+FMA only. The scalar and NEON paths are untouched.
- The A/B was run on one host. The instruction-count argument is
  machine-independent; the exact percentages are not.
- The `1048576` and `4194304` rows sit where ORT's own timing is least
stable
(its control column moved up to 1.6x between rounds on `Tanh`); medians
of
three alternating rounds are reported, and the `Erf`/`Sqrt` controls are
the
  evidence that the reported deltas are larger than that drift.

## Independent review

Reviewed by an independent Opus reviewer, verdict **GO WITH FINDINGS**.
The
reviewer re-derived the exhaustive proof independently (2 095 054 848
`tanh` and
2 078 277 632 `sigmoid` values, zero disagreements), confirmed `andnot ≡
blendv`
across quiet, signalling and non-canonical `NaN` payloads, `±0`, `±Inf`
and
subnormals, and confirmed the redundancy holds under all four MXCSR
rounding
modes.

All findings are applied:

1. **The margin claim was wrong, and it was load-bearing.** The comment
and this
body said `p/q` at `v = 9` is `1.0000001`. It is exactly `1.0` — `p` and
`q`
are the same bit pattern — so the safety argument rests on *equality
with an
inclusive clamp*, not on slack, and the old wording also contradicted a
correct comment fifteen lines above it. Both the comment and the body
now
state the exact values, and the sigmoid comment no longer claims the
result
lands "outside `[0, 1]` on both ends" when at `+18` it is exactly on the
   boundary. Corrected in `d32e52b`.
2. Signalling and non-canonical-payload `NaN`s added to the
special-value test.
3. Rounding-mode assumption noted in the kernel comment and above.
4. The reviewer's remaining point is that this PR is necessary but not
sufficient: on `main` the assignment policy declines these ops, so the
kernels are not reached until #1097 lands. That is why the benchmark is
   stacked, and it is stated in the section above.

---------

Co-authored-by: Deckard <deckard@users.noreply.github.com>
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby added a commit that referenced this pull request Aug 17, 2026
## What

Every activation in `simd_activations.rs` ran on a single core no matter
how large the tensor was. At prefill widths that leaves most of the
machine idle: a 512x4096 activation is 2M independent elements at a few
cycles each — pure throughput work with nothing forcing it to be serial.

This routes the `dispatch!` macro, `quick_gelu_f32_slice`, and the two
fused bias kernels (`tanh_gelu_bias_f32_slice`,
`erf_gelu_bias_f32_slice`) through a chunked runner backed by the rayon
pool that the GEMM kernels already use.

## Measured

Same binary, `RAYON_NUM_THREADS=1` vs `=16`, median of three interleaved
repeats, AMD EPYC 9V74 (32 vCPU, avx2/f16c/fma). ns/element:

| case (2,097,152 elems) | 1 thread | 16 threads | speedup |
| --- | ---: | ---: | ---: |
| Erf/f32/prefill512x4096 | 1.5865 | 0.2552 | **6.22x** |
| Sqrt/f32/prefill512x4096 | 0.5768 | 0.0988 | **5.84x** |
| FastGelu/f32/prefill512x4096 | 1.1671 | 0.2149 | **5.43x** |
| Gelu/f32/prefill512x4096 | 1.1655 | 0.2192 | **5.32x** |
| QuickGelu/f32/prefill512x4096 | 0.9693 | 0.1870 | **5.18x** |
| Sigmoid/f32/prefill512x4096 | 0.7823 | 0.1673 | **4.68x** |
| Tanh/f32/prefill512x4096 | 0.7874 | 0.1746 | **4.51x** |

No regressions anywhere else: decode4096 0.98–1.02x, decode3072
0.99–1.03x, small64 0.93–1.14x, tiny16 0.95–1.01x — all inside this
bench's noise band. Short calls test their length before rayon is
touched at all, so they never reach the pool.

## Methodology — why this is a same-binary comparison

`activation_bench.rs`'s own header documents a uniform 0.70–0.82x offset
on *byte-identical* kernels between separately-built binaries.
Interleaving does not remove it. An earlier revision of this change
measured a "uniform regression" that was entirely this artifact — proven
because untouched `QuickGelu` regressed by the same factor.

So the A/B here varies only the thread count within one binary. In the
first run `QuickGelu` was left serial deliberately and came out at
exactly 1.00x, confirming the harness attributes nothing to an unchanged
kernel; it is parallelised in this PR and now scales at 5.18x.

These are self-relative figures. **They are not a claim against ORT** —
ORT parallelises too, and a direct matched-thread A/B against ORT is
separate follow-up work.

## The thresholds were wrong, and only a real session showed it

The table above is a **self-relative** measurement, and self-relative is
exactly the measurement that cannot see this bug. A benchmark loop calls
the
kernel back-to-back, so rayon's workers never park and the fork/join
looks
nearly free. A real ORT session does the opposite: one short burst per
node,
with the rest of the graph in between, so almost every call wakes the
pool from
scratch — measured at roughly **50 us**.

Run through a real ORT session (same single-node graph, same input, our
EP
against ORT's CPU EP, matched intra-op threads, p50 of 200 interleaved
runs,
EP assignment asserted from ORT's own profiler), the originally proposed
`PAR_MIN_LEN` of 16 Ki was not a small loss but a **5x** one:

| elements | `PAR_MIN_LEN` = 16 Ki | `PAR_MIN_LEN` = 1 Mi |
|---|---|---|
| 16384 | 0.19x | 1.20x |
| 32768 | 0.17x | 1.43x |
| 65536 | 0.16x | 1.60x |
| 131072 | 0.21x | 1.68x |
| 262144 | 0.98x | 1.80x |

(`Sqrt`/f32 — the cleanest case, because it has no fix-up pass. Higher
is
better; these are ORT time over our time.) Every one of those sizes was
*already* faster than ORT on the serial path. The split took a 1.2-1.8x
win and
turned it into a 5x loss, on exactly the sizes a decode step uses.

So the thresholds are now `PAR_MIN_LEN = 1 Mi` and `PAR_MIN_CHUNK = 256
Ki`:
at ~50 us of wake-up and ~0.3 ns/element, break-even is around 160 Ki
elements
per chunk. Below 1 Mi the kernel stays on the serial path, which the
table
above shows is where it belongs.

## What this PR does and does not claim against ORT

With the corrected thresholds this change is a strict improvement on our
own
serial path wherever it engages, and it no longer makes any size slower
than it
was. It does **not** yet make us faster than ORT at high thread counts:
at
sixteen matched threads ORT still scales better than a rayon pool can
from
inside a plugin EP, because ORT's intra-op pool is already hot when the
node
starts and ours is not. Measured at 4 Mi, `Tanh`/f32, sixteen threads:
ORT
202 us, this PR 396 us, our serial path 1977 us. So the split is worth
having —
it is 5x better than not splitting — and it is still not enough.

Closing that gap needs the *host's* pool rather than one of our own, via
`OrtApi::KernelContext_ParallelFor`. A prototype of that is measured and
is a
large improvement over rayon at prefill sizes (`Tanh` 4 Mi: 816 us rayon
->
307 us host pool), but it carries a ~40-105 us fixed cost per call of
its own
that has to be understood before it can ship. That is separate follow-up
work
and is not in this PR.

## Three things this had to get right

**Do not touch rayon before checking the length.**
`rayon::current_num_threads()` reaches the global registry, and
initialising it spawns the pool — about 1.4 µs, more than an entire
4096-element activation. An earlier revision measured a uniform 0.6–0.8x
regression on every short case until that check moved above it.

**Chunking must not change numerics.** The vector and scalar paths round
differently, so the path is chosen once for the whole slice and every
chunk inherits it. Chunks are floored at `PAR_MIN_CHUNK` and rounded to
whole vectors so none can fall out of the vector path. The bias kernels
index `bias[i % width]`, so their chunks are whole multiples of `width`
— a mid-row cut would rotate the bias for everything after it.

**f16/bf16 are deliberately left serial.** They widen into an f32
scratch, compute, then narrow back, and parallelising only the middle
layer measured *slower*: Sqrt/f16 0.59x, Tanh/f16 0.79x, with bf16
prefill swinging 1.6–3.4 ns/element across repeats where f32 held to
±5%. Spreading 8 MB of scratch across sixteen private caches for a
serial narrow to pull back costs more locality than the arithmetic
saves. Parallelising the bulk conversions too was tried and was worse
still. The narrow-output arm of `write_mapped_reading` now runs under
`serial_scope`, pinning those paths to their previous behaviour —
measured 0.98–1.01x.

Fusing widen/compute/narrow into one pass per chunk is the real fix for
f16/bf16 and belongs in its own change, measured on its own.

## Correctness

Parallel output must be **bit-identical** to serial, not merely close —
these kernels are exactly chunk-independent, so anything else is a bug.

- `unary_kernels_are_thread_count_invariant`,
`quick_gelu_is_thread_count_invariant`,
`bias_kernels_are_thread_count_invariant` compare a one-thread pool
against the multi-threaded global pool over row widths (1, 3, 7, 11, 64,
4096, 4099) chosen to be coprime with the lane count and not to divide
the chunk size.
- The chunk policy is a pure function, so
`chunk_policy_holds_across_thread_counts_and_lengths` sweeps it over
thread counts and lengths this host cannot produce (up to 4096 threads,
4M elements).
- `serial_scope_suppresses_the_split` and
`serial_scope_is_restored_after_a_panic` pin the f16/bf16 guard,
including documenting that it is not unwind-safe.

A note on how the parallel side is obtained: `ThreadPool::install` runs
its closure *on a pool worker*, which trips the nesting guard and
silently serialises. An earlier version of these tests used `install`
for both sides and **passed with the row alignment deliberately
broken**. The parallel run is now a plain direct call.

All three guards were verified to fail when their invariant is broken:
- `+1` on the row chunk → `rows: chunk 8194 cuts row width 3 in half`,
and a real numeric divergence at element 8194.
- dropping the vector floor → `chunk 4096 could drop below the vector
threshold`.

`cargo fmt --check`, scoped clippy, and 1316 `onnx-runtime-ep-cpu` lib
tests are green.

## Limitations

- f16/bf16 unchanged by design (see above).
- Does not reach ORT's scaling at high thread counts (see above). It is
a
  large improvement on our own serial path, not a win against ORT.
- Tuned on one host. `PAR_MIN_LEN` (1 Mi) and `PAR_MIN_CHUNK` (256 Ki)
are
deliberately conservative, and the sub-threshold path is provably
identical
  to before — which is now most of the range.
- The wake-up cost is a property of a parked pool, not of rayon
specifically.
  Any pool we own rather than borrow will have it.

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

## Independent re-review

The thresholds changed after the first review's GO, so it was
re-reviewed.
Verdict **GO WITH FINDINGS**. The reviewer independently confirmed the
parallel
path is sound (rayon's safe `par_chunks_mut`/`par_chunks` zip, no manual
`Send`/`Sync`, no raw pointers, path chosen once for the whole slice so
no chunk
can drop to scalar), confirmed the bias kernels' chunks always start on
a row
boundary, and confirmed the invariance test's chunk boundaries really
are
adversarial (262 146 and 262 336 — not multiples of 8 — so a mid-row cut
would
be caught). 1331 tests pass. No interaction with #1097's rule: these
constants
are referenced only inside `simd_activations.rs` and change *how* our
kernel
runs, never *whether* the node is ours.

Findings, and what was done:

1. **One global threshold, derived from the cheapest kernel, applied to
all of
them.** The reviewer's point was that break-even scales with per-element
cost, so `Erf` and exact `Gelu` — 2-3x more work per element than `Sqrt`
—
are likely under-threaded between 256 Ki and 1 Mi. **Measured, and
correct.**
   At 16 intra-op threads, dropping `PAR_MIN_LEN` to 256 Ki gives:

   | op | 262144 | 524288 | 1048576 | 2097152 | 4194304 |
   |---|---|---|---|---|---|
| exact `Gelu` | **1.94x** | **2.35x** | **1.58x** | **1.26x** | 0.96x |
| `FastGelu` | **1.32x** | **2.03x** | **1.33x** | **1.28x** | 0.92x |
   | `Erf` | **1.28x** | **1.78x** | **1.52x** | **1.23x** | 0.79x |
   | `Sqrt` | **0.43x** | 0.72x | 1.00x | 1.16x | 0.97x |

So the reviewer is right for the transcendentals and the current value
is
right for `Sqrt`, which is the kernel that would be destroyed by a lower
one.
A single compromise value cannot serve both; the fix is a per-kernel
cost
class, which needs plumbing through four generic entry points and is
left to
a follow-up. This PR records the measurement and the reasoning beside
the
constant (`967191438`) instead of guessing, and keeps the conservative
value:
   too high costs throughput, too low cost 5x.

2. **The 256 Ki chunk floor caps a 1 Mi tensor at four workers.** True,
and it
is the same trade-off as finding 1 — it follows the cost class, so it is
   tracked with it.

3. **Thread configuration of the headline table was ambiguous.** Fixed:
the
doc-comment now states that both sides ran at `intra_op_num_threads = 1`
with
our pool at `RAYON_NUM_THREADS = 32`, so the "1.8x serial win" is
against
   *single-threaded* ORT, not multi-threaded ORT.

4. **`N` in `thread_invariance` is divisible by 3**, contradicting a
comment
claiming coprimality with every row width. Comment corrected to say
where the
   adversarial coverage actually comes from.

## What this PR does not fix

At 16 intra-op threads our elementwise kernels remain far behind ORT —
0.06-0.50x across these families, on both threshold settings. ORT scales
these
ops ~14x from 1 to 16 threads; we manage ~6x. That gap is not a
threshold
problem and is not addressed here; it is a property of owning a pool
that has
to be woken per node, and the next step is to use the host's pool
instead of
our own. Recorded so the merged state is not mistaken for a solved one.

---------

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Co-authored-by: Deckard <deckard@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