Skip to content

perf(speculative): the qmv_wide exactness buy-back pins the whole process narrow, taxing batched decode that never asked for it #1261

Description

@inureyes

Problem

When the MTP exactness probe fails under qmv_wide and passes without it, the gate disables the kernel for the whole process and deliberately never re-enables it (retry_without_qmv_wide in src/models/speculative_exactness.rs: "it is a per-process kernel selection"). That is correct for the MTP verify forward, which is the block the contract is about, but the switch sits on the dispatch path of every quantized matmul in the process (mlxcel_qmv_wide_flag() in the quantized.cpp overlay), so everything else the same server process runs also drops to the narrow kernel:

  • batched decode at B >= 2, whose per-step projections are exactly the M >= 2 shape qmv_wide exists for;
  • prefill for other requests admitted while MTP is engaged;
  • any non-MTP request the scheduler serves alongside the speculative one.

On Apple GPU generation 15+ this fires by default now: the probe fails, the retry passes, and the process is pinned narrow (#1199 for Qwen, #1258 for Gemma 4). The measured verify-side cost is 17 to 20% (Qwen) and ~23% (Gemma 4); what the rest of the process pays has never been measured.

Step 1: measure the tax before designing the scope

Nothing below is worth building until the collateral cost is a number. On one generation 15+ host, mlxcel-server with a batch-capable target:

  • batched decode throughput at B = 2, 4, 8 with MLXCEL_QMV_WIDE=1 pinned vs =0 pinned (no drafter involved, so the comparison isolates the kernel);
  • a mixed workload: one MTP stream plus N classic streams, default env (gate flips the process narrow) vs MLXCEL_MTP_ALLOW_INEXACT=1 with the switch left wide, reading the classic streams' throughput only.

If the B-sweep delta is small on production shapes, the right fix may be documenting the tax and stopping there.

Step 2 (if the tax is material): scope the exact kernel to the verify forward

The hard part is recorded in the overlay's own comment: MLX evaluates lazily and dispatch happens at eval time on MLX's scheduler thread, so "set off, run verify, set on" from the caller thread does not bracket the kernels it means to bracket, and a caller-side thread-local never reaches the dispatch site. Candidate shapes, in rough order of invasiveness:

  1. Bracket with explicit synchronization: flip narrow, force the verify block's eval, flip wide. Costs pipeline stalls; perf(speculative): attribute the MTP drafter step and count the accept hook #1194 measured syncs at ~13% on the drafter step, so this can eat the savings.
  2. A dispatch-side predicate keyed on something the verify ops carry (a dedicated stream, or an op annotation plumbed through the overlay). More surgery in the overlay, but no stalls.
  3. Per-call kernel selection for just the verify projections via a dedicated entry point in the overlay, leaving the global flag untouched.

Whichever shape wins must keep the probe's contract: the verify block's kernel selection at probe time has to match its selection at serve time, or the probe certifies the wrong thing.

Acceptance criteria

  • The B = 2/4/8 batched-decode tax of the narrow pin is measured and recorded on a generation 15+ host.
  • Either the tax is documented as accepted (with the number), or the exact kernel is scoped so non-verify work keeps qmv_wide, with byte-identity of the MTP stream re-verified after the change.
  • The probe measures the same kernel configuration the verify forward serves with.

References

Activity

  1. added and removed on Aug 21, 2026
  2. inureyes commented on Aug 21, 2026

    @inureyes
    MemberAuthor

    Deferred to a generation 15+ host, with two premise corrections

    Picked this up for implementation from an M1 Ultra session and stopped before writing code, because Step 1 is not measurable here and the issue gates everything else on it ("Nothing below is worth building until the collateral cost is a number"). Recording what the audit found so the next attempt does not re-derive it, and does not walk into the trap below.

    The generation 13 trap: running Step 1 here would produce an artifact, not a measurement

    The patched predicate in src/lib/mlx-cpp/patches/mlx/backend/metal/quantized.cpp:584-586 is an AND, so the mlxcel flag is an off-switch only and can never turn qmv_wide on:

    return mlxcel_qmv_wide_flag().load(std::memory_order_relaxed) &&
        (mode != "affine" || d.get_architecture_gen() >= 15);

    For an affine 4-bit checkpoint on generation 13, use_qmv_wide is therefore false in both arms of the proposed A/B, and MLXCEL_QMV_WIDE=1 versus =0 measures nothing but scheduler noise. That matters more than it first looks: a near-zero delta obtained that way would falsely satisfy this issue's own off-ramp ("If the B-sweep delta is small on production shapes, the right fix may be documenting the tax and stopping there"). The B = 2/4/8 sweep in the acceptance criteria only means anything on generation 15 and later, or on a non-affine checkpoint (see below).

    The problem statement is narrower than the defect

    The body says the narrow pin "fires by default" on Apple GPU generation 15 and later. The module documentation at src/models/speculative_exactness.rs:36-38 records something wider: an mxfp4 projection diverges from M = 2 on an M1 Ultra and an M5 Max, unlike the affine row, which only breaks from generation 15. Since use_qmv_wide is mode != "affine" || arch_gen >= 15, the flag is live on every generation for block-float modes.

    Gemma 4 is one of the two families that reach mtp_exactness_gate (src/models/qwen3_5.rs:1523 and src/models/gemma4.rs:5147), and ModelOpt NVFP4 Gemma 4 checkpoints are supported (is_modelopt_nvfp4 in src/models/gemma4.rs:4204, repacked by sanitize_gemma4_nvfp4_weights). So on a generation 13 host serving a non-affine Gemma 4 checkpoint, the probe should fail, retry_without_qmv_wide should fire, and the process should be pinned narrow, which is exactly the collateral damage this issue is about.

    To be clear about the evidence: this is read from the source, not yet observed. It was not confirmed empirically because the local release binary predates #1258, which is the change that gave Gemma 4 a real probe instead of an assumption. Whoever picks this up should either confirm or refute it, since it changes the blast radius in the problem statement and it also means the bug has a reproduction that does not require generation 15+ hardware.

    What is and is not measurable per host

    Measurement Generation 13 Generation 15+
    affine 4-bit B = 2/4/8 sweep, the shape #1199 and #1258 actually hit not measurable, both arms narrow required, and the real acceptance criterion
    mixed workload, one MTP stream plus N classic streams not measurable, same reason required
    non-affine (mxfp4 / nvfp4) B-sweep measurable, flag is live on every generation measurable
    reproduction of the narrow pin itself plausible via a non-affine MTP checkpoint, unconfirmed fires by default

    A non-affine sweep on generation 13 is a real kernel-level comparison of the same two kernels, but it is not the production shape, so it cannot stand in for the acceptance criterion. It is worth running only as supporting evidence, and it should be labelled as such if it is.

    Suggested sequencing for the generation 15+ attempt

    Step 1 first, as the body already says. If the tax turns out to be material, the overlay comment at src/lib/mlx-cpp/patches/mlx/backend/metal/quantized.cpp:560-569 already states the constraint that decides between the three candidate shapes: MLX evaluates lazily and dispatch happens at eval time on MLX's scheduler thread, so a caller-side "set off, run verify, set on" does not bracket the kernels it means to, and a caller-side thread-local never reaches the dispatch site. Candidate 1 (sync-bracketing) should be costed against #1194, which measured syncs at about 13% on the drafter step.

  3. inureyes commented on Aug 21, 2026

    @inureyes
    MemberAuthor

    Addendum: an M3 Ultra qualifies for Step 1, but its batch limits are M1-Ultra-class

    Follow-up to the comment above, now that this is being picked up on an M3 Ultra.

    Generation mapping, confirmed. arch_gen_ is just the two digits before the last character of the Metal architecture name (mlx/backend/metal/device.cpp:593-601), so applegpu_g<NN><size> gives generation NN. Verified on the M1 Ultra here, which reports applegpu_g13d, matching the "generation 13" in src/models/speculative_exactness.rs:33. An M3 Ultra is therefore applegpu_g15d, generation 15.

    That clears the bar in use_qmv_wide, so affine 4-bit takes qmv_wide at M >= 2 there, the narrow pin fires by default, and the B = 2/4/8 sweep is a real measurement rather than the identity it collapses to on generation 13. The trap described above does not apply to an M3 Ultra run.

    But the size class still routes to the 'd' branch. Every generation-15 and generation-17 branch of get_qmv_batch_limit (src/lib/mlx-cpp/patches/mlx/backend/metal/quantized.cpp:108-145) is guarded by arch_size != 'd'. An Ultra is 'd', so an M3 Ultra falls through to the arch_gen >= 13 arm and gets case 'd': 32 / 18 / 12 by operand size, identical to an M1 Ultra, not the 13 / 15 / 13 the generation-15 arm would give.

    So an M3 Ultra is a genuinely mixed configuration: qmv_wide behaves as generation 15, while the matrix-matrix takeover threshold behaves as generation 13. Two consequences for this issue:

    • The --draft-block-size value at which byte-identity is forfeited for an unrelated reason sits at a different width than it does on an M5 Max, so a block width that is safe on one is not automatically safe on the other.
    • Numbers measured on an M3 Ultra should be recorded as "generation 15, size class d" rather than generalised to "generation 15+". The 'd' guard means the Ultra parts are their own row in this table, on every generation.

    Possible doc drift, worth a check by anyone with an M5. The module note at src/models/speculative_exactness.rs:39-42 says the limit "is 12 on an M1 Ultra (arch_size == 'd') and 10 on an M5 Max (the default branch)". The M1 Ultra half still matches the code. The M5 Max half does not appear to: an M5 Max is generation 17 and not 'd', so it takes the first branch (33 / 25 / 13), while default: exists only inside the arch_gen >= 13 arm. That reads like the note was written before the generation-15 and generation-17 branches arrived and was not revisited on the pin bump. Not verified against hardware here, since this session has no M5, so treat it as a flag rather than a finding.

  4. self-assigned this
    on Aug 31, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions