Repository navigation
feat(xla): define operator-level MLX/IREE numeric contracts #932
Description
Activity
- addedstatus:readyReady to be worked onReady to be worked ontype:enhancementNew features, capabilities, or significant additionsNew features, capabilities, or significant additionspriority:highHigh priorityHigh priorityarea:architectureArchitecture and code structure changesArchitecture and code structure changesarea:coremlxcel-core: MLX FFI, primitives, KV cache, layersmlxcel-core: MLX FFI, primitives, KV cache, layersarea:inferenceGeneration, sampling, decoding (incl. speculative, DRY)Generation, sampling, decoding (incl. speculative, DRY)
on Jul 26, 2026 5 remaining items
- added a commit that references this issue
on Jul 27, 2026 Progress update:
- feat(xla): version auxiliary dtype contracts #934 merged: auxiliary dtype contracts are now versioned.
- feat(xla): bind operator numeric contracts #937 merged: operator-level numeric contracts are bound to backend artifacts.
- feat(xla): add bounded numeric oracle harness #938 merged: bounded numeric oracle reports exact first divergence and keeps production qualification conservative.
- feat(xla): run dense matmul numeric probe #940 merged: a real in-process local-task IREE dense-matmul probe now compiles, loads, invokes, and compares exact f32 bits against the canonical Rust reference. The observed fixture passed 4/4 elements with zero divergence, while the backend algorithm remains explicitly unobserved and therefore not production-qualified.
The seven family PRs remain drafts until their family-specific intermediate/output/token gates pass. Next work is to expand the CPU probe surface to normalization, activation/projection, residual, and softmax/attention primitives before attempting family-specific qualification. CUDA-dependent Q4/SSCP parity remains blocked on a GPU-capable runner.
Shared numeric-contract foundation progress:
- feat(xla): expand core numeric probes #941 merged at 8795726, extending the real in-process CPU IREE suite from dense matmul to residual add, RMSNorm, and attention softmax.
- feat(xla): share contract-sensitive numeric emitter helpers #942 merged at 9465bba, routing production Qwen2-VL/shared vision LayerNorm, exact/tanh GELU, and stable softmax through the same contract-sensitive helpers used by probes. It also added LayerNorm, SiLU projection, and GELU projection probes.
Current CPU evidence: 7 operations, 48 output elements, 0 failures. Dense matmul, residual add, RMSNorm, attention softmax, LayerNorm, and SiLU projection are bit-identical to the canonical references; GELU projection max absolute error is 4.7683716e-7 within the unchanged contract. All reports remain canonical-decomposition and production_qualified=false, so no family capability changed.
Next CPU-supported boundary is the explicit prefix reduction/scan contract. Post-#923 Q4 qmv/qmm and production SSCP convolution parity still require the CUDA environment.
Prefix-reduction foundation is now merged in #943 at 15a3711. The shared helper pins the MLX CUDA one-block contiguous scan hierarchy (4 values/thread, 32-thread staged warp scans, first-warp total scan, width <= 4096) and records association=cuda-one-block-scan in operator artifact identity.
A cancellation-sensitive 2x130 probe distinguishes this schedule from naive sequential accumulation and passed all 260 outputs bit-exactly through local-task CPU IREE. The full suite is now 8 operations / 308 outputs with 0 failures; all reports remain canonical-decomposition and production_qualified=false. Next I will rebase draft #922 and replace its family-local cumulative scan body with this shared helper; the unrelated conv1 cuDNN/IREE reduction-plan blocker remains unchanged.
Merged #944 with an independently executable CPU IREE affine Q4 dequant probe. The probe uploads actual u32 packed lanes plus f16 scale/bias resident weights, pins least-significant-first lane order, group-8 metadata, and separate multiply-then-add evaluation, and passes 32/32 outputs bit-exact. The complete local-task suite now passes 9 operators / 340 outputs with production_qualified=false throughout. Native CUDA qmv/qmm parity remains CUDA-runner-gated.
Shared diagnostic runtime foundation merged via #945. The diagnostics-only helper bounds IREE local-task topology to one group, uses the host-default pthread stack, caches the process-global parse result, and leaves production startup unchanged. Local validation covered formatting, two structural tests, diagnostics-feature compilation, and Clippy with pre-existing warnings only. Final native linkage and runtime execution are now being exercised through the Youtu family oracle in #919.
Youtu-VL bounded numeric-contract evidence (
#919 @ 20c2bffd)The real pinned
tencent/Youtu-VL-4B-Instruct@8d30a0e...HF-eager capture and
local-task IREE diagnostic graph now execute with a bounded fail-fast runner.Resolved semantic boundaries:
- all 18 checkpoint artifact hashes and fixture hash match;
- placeholder cardinality now expands exactly per image;
- stale shard-index hints no longer select the wrong tensor file;
- resized pixels, flattened/channel-fast patch layout, spatial metadata,
window ordering, and vision RoPE match exactly; - patch projection passes the unchanged gate (
max_absolute=0.015625); - row-wise LayerNorm decomposition matches HF exactly.
First unresolved production-reference boundary:
- stage:
layer.7.full - first failure: flat index 310
- actual/reference:
0.3125/0.28515625 - absolute/max absolute:
0.02734375/0.203125 - unchanged-threshold mismatch count: 2,999
Feeding the IREE patch projection into HF layer 0 remains within
0.015625of
the HF reference, but the IREE layer-0 output reaches0.03125; drift begins
inside BF16 projection/attention contractions and accumulates. Explicitly
typing contraction inputs as BF16 produced byte-identical local-task IREE
outputs, so the family-local workaround was removed. Shared LayerNorm,
RMSNorm, GELU, and softmax emitter helpers are used. No tolerance changed.Classification under #932: canonical decompositions execute, but this target
does not prove the pinned production-reference BF16 contraction schedule.
Youtu-VL remains draft/fail-closed; language logits/KV, exact greedy tokens,
reset/reuse, and mixed batching remain unqualified.Current-main family qualification checkpoint
Shared foundation is merged on
main@ec6fbb6dthrough #934, #937, #938, and #940–#945. The local-task CPU suite covers 9 operator contracts / 340 outputs with zero contract failures, while every report remains correctly classified as canonical decomposition andproduction_qualified=false.All seven family branches are now rebased onto this foundation, draft, mergeable/clean, and green on their current CI heads:
Family PR head Current qualification result Molmo2 #916 684c6ee1Actual gate fails first at projector.output_all, flat 358403,max_abs=0.1640625,rms=0.0026185848; preceding product passes.Gemma3 VLM #917 9ab13992Current actual was bounded before model start by the one-time pinned-MLX rebuild. Last valid actual remains block 1, flat 658517, 8 failures, max_abs=0.6672516; no current-head parity claim.Molmo #918 a4047b22Assets and hashes are pinned, but no release runner was produced inside the bounded link window. A goldaudit also invalidated Cargo/MLX fingerprints before link and was stopped immediately; no actual comparison claim.Youtu-VL #919 20c2bffdActual gate reaches layer 7; first fail flat 310, actual/reference 0.3125/0.28515625,max_abs=0.203125, 2,999 mismatches.Qwen2.5-VL #920 02d41a18Current structural tests pass; this environment exposes no NVIDIA device, so post-#923 production Q4 qmv/qmm evidence remains unavailable. Qwen3-VL #921 a4b641e6Current Qwen3/DeepStack tests pass; the stale pre-#923 Q4 evidence is not promoted, and no NVIDIA device is exposed for reproduction. Gemma3n audio #922 0a5aa90dShared prefix scan is adopted; diagnostics and 17 audio tests pass. The production first SSCP convolution still has 6 BF16-ULP residual failures with max_abs=0.25, without a represented cuDNN-equivalent schedule.Qualification decision: none of #916–#922 satisfies its unchanged original intermediate plus token-exact gates, so none is merged or advertised. No tolerance was relaxed, and no canonical/mathematical proxy is reported as production parity.
Remaining #932 blockers are production CUDA Q4 qmv/qmm evidence, a reproducible production SSCP convolution contract, and current-head actual completion for the environment-bounded families. Until a GPU-capable/cached runner or a reproducible backend schedule is available, #932 and the affected family capabilities must remain open and fail-closed.
Refs #566
- addedstatus:blockedBlocked by dependencies or other issuesBlocked by dependencies or other issuesand removedstatus:in-progressCurrently being worked onCurrently being worked on
on Jul 27, 2026 Three host and methodology defects found while running the family gates on GB10
Running the #566 family gates on the GB10 qualification host surfaced three problems that
sit underneath the numeric work. The first two are build-environment defects; the third
invalidates the reference policy of past MLX-vs-IREE evidence and is the one that matters
for this issue.1. The MLX-vs-IREE reference has been running TF32, not F32
mlx/utils.h:196defaultsMLX_ENABLE_TF32to1:inline bool enable_tf32() { static bool enable_tf32_ = get_var("MLX_ENABLE_TF32", 1); return enable_tf32_; }
With it on,
mlx/backend/cuda/gemms/cublas_gemm.cppselects
CUBLAS_COMPUTE_32F_FAST_TF32for f32 contractions, which carries a 10-bit mantissa.
src/lib/mlxcel-xla/README.mdalready states the policy explicitly:The companion command sets
MLX_ENABLE_TF32=0: the declared reference policy is true F32,
while MLX CUDA's default FAST_TF32 mode is an intentional, lower-precision throughput
policy and is not labeled as F32 evidence.None of the family runners set it, and
CUDA_TEST.mddoes not list it. So MLX CUDA has been
acting as a TF32 reference against a true-F32 IREE candidate, and the resulting gap has been
attributed to the IREE graphs.Measured on #918 (Molmo,
torch_dtype: float32, unquantized vision tower, and a Molmo
emitter that applies no precision demotion, so both sides should be plain f32):stage MLX_ENABLE_TF32default (1)MLX_ENABLE_TF32=0vision.patch_embedding1.336575e-3, 0 fail 1.336575e-3, 0 fail vision.selected_layer_142.258102, 0 fail 1.991226, 0 fail vision.selected_layer_212.317841, 4 fail 2.046043, 4 fail vision.projector2.251205e-1, 0 fail 1.643753e-1, 0 fail prepared.sparse_add_visual_rows2.500000e-1, 0 fail 1.875000e-1, 0 fail decoder.prefill_logits1.176004e-1, 0 fail 1.173029e-1, 0 fail Turning TF32 off removes 12 percent of the drift at the selected ViT layers and 25 to 27
percent at the projector and the merge. The patch embedding is unchanged, consistent with
mlx/backend/cuda/conv.cpp:139gating the convolution path onenable_tf32()separately.This does not by itself close #918: the same four elements at flat index 514613 still fail.
The point is narrower and important. Any past or future statement of the form "IREE diverges
from the MLX F32 reference by X" is not valid F32 evidence unlessMLX_ENABLE_TF32=0was
set for that run, and the residual after setting it is the only number that should be
attributed to the graph.Suggested follow-ups:
- have the family runners assert
MLX_ENABLE_TF32=0rather than rely on the operator, in
the same spirit as the diagnostics-only local-task thread configuration - record the value in every report alongside the other backend identity fields required by
feat(xla): add bounded numeric oracle harness #938 - add it to the
CUDA_TEST.mdenvironment section
2.
build.rssilently targets the wrong GPU whennvidia-smiis not visiblesrc/lib/mlxcel-core/build.rs:274:let cuda_arch = env::var("MLX_CUDA_ARCHITECTURES") .unwrap_or_else(|_| detect_cuda_arch().unwrap_or_else(|| "90a".to_string()));
detect_cuda_arch()shells out tonvidia-smi. In a sandbox without GPU passthrough, which
CUDA_TEST.mddocuments as the situation for the earlier sessions, detection fails and the
build silently produces90abinaries that cannot launch on this host's sm_121. The failure
surfaces much later and looks like a runtime or numeric fault:cudaGraphAddKernelNode(...) failed: invalid resource handlewith CUDA graphs enabledcudaLaunchKernelExC(...) failed: no kernel image is available for execution on the device
withMLX_USE_CUDA_GRAPHS=0
Wrong-architecture MLX builds were present in the #866, #869 and #878 worktrees. A loud
failure would be better than a silent fallback, since the fallback cannot run on the host
that just failed to be detected.3.
build.rsauto-detection does not match the shipped release architecturesm_arch_with_suffix()appends the architecture-specificasuffix for any SM at or above
90, so auto-detection on GB10 yields121a. The release workflow ships
MLX_CUDA_ARCHITECTURES: "90a;100;121"anddocs/installation.mddocuments plain121for
GB10. The suffix is only documented as required for Hopper, whereqmm_sm90is gated on
90a; there is no equivalent Blackwell gate.An auto-detected local build therefore has a different kernel set from the released binary.
Verified on #918 that121aand121produce byte-identical numeric results
(vision.selected_layer_21max_abs2.317841, 4 failures, first failure 514613 under both),
so this is a reproducibility and provenance problem rather than a numeric one, but
qualification evidence should be produced with the shipped value.Environment used
GB10, driver 580.159.03, CUDA 13.0, IREE 3.12.0rc20260721 (compiler and source-built
runtime,dc9601f8), Rust 1.93.1, MLX pinb7c3dd6d27f4,MLX_CUDA_ARCHITECTURES=121.- have the family runners assert
Scope of the TF32 reference-policy defect, measured across all five family gates
Follow-up to the report above. I re-ran every family gate I could reach at current head with
MLX_ENABLE_TF32=0against the MLX default, changing nothing else, to bound which results the
policy actually affects. The answer is narrower than the initial report implied and is worth
recording precisely.PR family effect of MLX_ENABLE_TF32=0verdict #916 Molmo2 vit.block.00.0028572083 to 0.000009536743 (300x)decides the gate: FAILED at first block to full PASS #918 Molmo v1 selected layers 12 percent, projector and merge 25 to 27 percent improves, 4 failures remain #920 Qwen2.5-VL none, identical to six decimals no effect #921 Qwen3-VL none, all 39 failing comparisons identical no effect #917 Gemma3 VLM none, identical at every stage no effect The discriminator is whether the compared path performs genuine f32 contractions.
CUBLAS_COMPUTE_32F_FAST_TF32is only selected for f32 operands, so towers resident at a
narrow checkpoint dtype never take that route. The Qwen2.5-VL report makes this visible
directly: every emitted-graph stage is bit-identical between the two runs, while the two
host-side dense-f32 controls move
(mlp_gate_projection_iree_norm2_dense_f32_control0.000122 to 0.000076,
mlp_up_projection_iree_norm2_dense_f32_control0.000290 to 0.000122). The flag reaches the
f32 host controls and not the f16 tower.Practical consequences:
- The Molmo2 result previously recorded as a numeric blocker (first strict CUDA failure at
selected layer 18,projector.output_allatmax_abs=0.1640625) was a reference-policy
artifact. That gate passes at329fbd99. - The Qwen2.5-VL, Qwen3-VL and Gemma3 divergences are properties of the emitted graphs and
are unchanged. Their blockers stand as recorded. - Asserting the policy in the runners is still worth doing, because a reader cannot tell
from a report which category a given family falls into, and feat(xla): add bounded numeric oracle harness #938 requires backend identity
to be recorded rather than inferred.
Every result above was produced with
MLX_CUDA_ARCHITECTURES=121on GB10, IREE
3.12.0rc20260721, MLX pinb7c3dd6d27f4, Rust 1.93.1, tolerances unchanged.- The Molmo2 result previously recorded as a numeric blocker (first strict CUDA failure at
Reduction-order contract: decided
The maintainer resolved this on 2026-07-30. A family PR is gated on exactly two
things:- Every operator's
same_inputcontrol passes at its existing threshold. No
absolute, relative, RMS or ULP threshold is relaxed for any model. - Deterministic greedy output is token-exact against MLX on the pinned
fixture.
Accumulated intermediate-value drift is still measured and reported as evidence,
but it is not a gate.Why
Per-operator parity is already established. #921's
.same_inputcontrols all pass,
so each operator agrees with MLX when handed identical inputs; what diverges is the
composition over many blocks. #922's convolution oracle reports
schedule_outcomes=[0.0, 1.0], an exact-tie cancellation, which shows the residual
is reduction order rather than precision. StableHLO offers no portable way to
pin a reduction tree, and a previous audit already showed
precision_config=HIGHESTselects the same IREE schedule, so matching MLX's
reduction order would require custom kernels or IREE codegen changes of unbounded
scope.Gating on per-operator parity plus token exactness keeps every existing numeric
threshold intact while testing the property a production consumer actually
experiences.Consequences
PR previously blocking observation now judged by #918 Molmo 4 of 1,181,696 values at selected_layer_21, same index on local-task / local-sync / cudatoken gate #921 Qwen3-VL 39 comparisons, block_2(2 values) growing toblock_23(89,017)token gate #920 Qwen2.5-VL 22 failing stages, post_full_layer_23max_abs=105.27258token gate #917 Gemma3 VLM block1.output8 failuresmax_abs=0.6672516growing to 1,408,110token gate, but the block-1 onset still needs an explanation: a divergence that appears abruptly at one block is not the same signature as gradual drift and may be a real defect This makes the missing token-exactness gates the critical path for these PRs rather
than further intermediate-divergence investigation.Precondition found while implementing this
#963: every OpenXLA path returned the terminating EOS id as generated output, while
the eager MLX paths drop it before detokenization. Measured on GB10 with
molmo2-4b: eager CLI returnsWhite(1 token),MLXCEL_BACKEND=xlaCLI returns
White<|im_end|>(2 tokens). Until that is fixed, every family fails a
token-exact gate on one trailing token, so #963 lands before any family PR is
judged by criterion 2.Refs #566.
- Every operator's
Parent and affected work
ec6fbb6d(shared foundation through fix(xla): bound diagnostic local-task threads #945)Problem
The remaining OpenXLA multimodal family ports now fail at a shared boundary rather than at missing model topology: MLX CUDA and StableHLO/IREE can execute mathematically equivalent operators with different operand materialization, accumulator precision, reassociation, reduction trees, fused-kernel ordering, or backend-selected convolution plans. The current family implementations encode some of these choices locally, so each full-model oracle run is rediscovering the same backend contract one architecture at a time.
This is not a request to relax tolerances. The existing family gates must remain unchanged, and final greedy output remains token-exact. If an MLX kernel has no deterministic, representable execution contract, the affected XLA capability must remain fail-closed until a reproducible contract is implemented or the reference path itself is made deterministic.
The reference baseline also changed after the current seven branches were cut: #923 fixed a shared-memory race in MLX CUDA
qmm_sm80and changed the kernel artifact identity. Any Q4 evidence collected beforef7a877e7must therefore be treated as historical until it is reproduced on the deterministic kernel.Current evidence
684c6ee1projector.output_all; first failure is flat index 358403 withmax_abs=0.1640625,rms=0.0026185848. The immediately precedingprojector.productpasses (max_abs=0.037597656,rms=0.0002196383), isolating final dense-projection accumulation/association.9ab13992siglip.hidden.block1.outputat002bb77e: flat index 658517, 8 failures,max_abs=0.6672516; current head passes the exact example check and focused Gemma3 tests but makes no new parity claim.a4047b2220c2bffdlayer.7.full, flat index 310: actual/reference0.3125/0.28515625,max_abs=0.203125, 2,999 unchanged-threshold mismatches.02d41a18a4b641e60a5aa90dmax_abs=0.25; no reproducible cuDNN-plan-equivalent IREE schedule is represented, so the capability remains disabled.Goal
Define and enforce one versioned, operator-level numeric contract for MLX-to-StableHLO/IREE qualification, then use bounded micro-oracles to resolve shared drift before repeating full-model gates.
Required implementation
f7a877e7baseline before using their evidence. Re-run only the bounded operator probes first; repeat a heavyweight full-model oracle only after the first divergent operator boundary passes.Merge strategy
Implement the shared descriptor, fingerprint validation, explicit helpers, and micro-oracle harness in a narrow foundation PR first. After that PR lands, rebase the seven draft family PRs sequentially, resolve their shared-file overlap against the foundation, and run the affected bounded probes. A family PR becomes ready only after its original intermediate and token-exact gates pass; diagnostic or fail-closed infrastructure may be split into a separate PR when it is independently complete and does not expose the unfinished family.
Non-goals
Validation
mlxcel-xlastructural/golden tests, native IREE compile/load/invoke checks, text-only XLA architecture oracle, and continuous-batch regression suites.Acceptance criteria