Skip to content

metal: 1-bit qmv_fast use 1 pack/thread for occupancy (~9% on M5) - #3

Merged
bri-prism merged 1 commit into
prismfrom
perf/1bit-qmv-occupancy-prism
Jun 5, 2026
Merged

bri-prism merged 1 commit into
prismfrom
perf/1bit-qmv-occupancy-prism

Conversation

@bri-prism

Copy link
Copy Markdown

What

affine_qmv_fast (the decode/token-gen quantized matvec) sets
packs_per_thread = bits == 2 ? 1 : 2, so 1-bit falls into the : 2 case and gets
values_per_thread = pack_factor * 2 = 64 → x_thread[64] (~256 B/thread). That register
footprint drops occupancy, so 1-bit decode can't hide DRAM latency and under-saturates
bandwidth (~75% on M5 Pro)
while 2-bit/4-bit saturate (~90/96%).

Use 1 pack/thread for bits <= 2 so 1-bit matches 2-bit's values_per_thread = 32
footprint. (2-bit already used 1 pack, so it is unchanged — it serves as a control below.)

Why it helps

  • Occupancy: halves per-thread register/x_thread pressure for 1-bit → better latency hiding.
  • Bonus correctness/robustness: scale_step_per_thread = group_size / values_per_thread was
    0 for group_size = 32 at 1-bit (vpt=64); with vpt=32 it is well-defined for g32.

Measured (Apple M5 Pro, distinct-weight DRAM-bound matvec, best-of-10)

2-bit is the drift control (untouched by this change):

before (packs=2) after (packs=1) Δ
1-bit 24.0 µs/matvec 21.9 µs/matvec ~9% faster (75%→82% BW)
2-bit (control) 33.2 µs 32.9 µs 0.8% (drift)

Correct: quantized_matmul vs dequantize+matmul rel_err 2.7e-4. The control moving only 0.8%
confirms the 1-bit gain is real, not machine noise.

Notes

  • Measured on a contended M5 (shared with another profiling session) — the ~9% is robust vs the
    control, but a quiet machine would tighten the headline number.
  • Remaining 1-bit under-saturation (82%→~95%) is dominated by scale+bias traffic (~33% of 1-bit
    bytes at group 64) — a follow-up (int8/shared scales), not this PR.

affine qmv_fast set packs_per_thread=2 for all bits except 2-bit, so 1-bit
got values_per_thread=64 (x_thread[64], ~256B/thread) -> low occupancy ->
1-bit decode saturates only ~75% of M5 DRAM BW vs ~90/96% for 2/4-bit.
Use 1 pack/thread for bits<=2 (values_per_thread=32), matching 2-bit's
register footprint.

Measured on M5 Pro (distinct-weight DRAM-bound, 2-bit as drift control):
1-bit 24.0 -> 21.9 us/matvec (~9%, 75->82% BW), 2-bit control 33.2->32.9
(0.8% drift); correct, rel_err 2.7e-4. Also makes scale_step_per_thread
(=group_size/values_per_thread) well-defined for group_size=32 at 1-bit.
@bri-prism
bri-prism requested a review from khosravipasha June 5, 2026 01:13
@bri-prism

bri-prism commented Jun 5, 2026 •

Copy link
Copy Markdown
Author

End-to-end validation on the test 1-bit model (MLX, M5 Pro; 128 tok, temp 0, 3 reps each):

tokens/sec
before (packs=2) 31.0 (30.97/30.99/31.03)
after (packs=1) 34.2 (34.13/34.25/34.32)
decode speedup +10.3%

Both stable across reps. The occupancy fix lifts every projection matvec in the decode path, so the end-to-end gain slightly exceeds the isolated qmv micro-bench (~9%).

@bri-prism

bri-prism commented Jun 5, 2026 •

Copy link
Copy Markdown
Author

Context-window scaling (test 1-bit model, M5 Pro, decode t/s after prompt prefill):

context (prompt tok) before after win
426 30.8 34.1 +10.7%
1,650 30.0 33.1 +10.6%
6,570 27.5 30.2 +9.8%
13,122 24.8 27.0 +8.7%

The win persists across context, diluting only slightly (10.7%→8.7%) as the attention KV cache grows to a larger share of decode (the hybrid-attention design keeps that share small — decode falls only ~19% over 30× context).

@bri-prism

bri-prism commented Jun 5, 2026 •

Copy link
Copy Markdown
Author

Losslessness confirmed — bit-exact on the test 1-bit model. Teacher-forced prefill over a fixed passage (33 positions × full vocab), before vs after:

mean KLD: 0.000e+00   max KLD: 0.000e+00
top-1 agreement: 100.00% (33/33)
logit |Δ|: mean 0.0, max 0.0   ->  arrays identical

The occupancy change (vpt 64→32) reorders the K-loop tiling but yields byte-identical logits, so output is unchanged while decode is +10%. (Note: the 1-bit checkpoint itself currently generates incoherent text — a separate, pre-existing conversion-coherence issue unrelated to this kernel; identical before/after.)

@bri-prism
bri-prism merged commit 88c9c20 into prism Jun 5, 2026
4 checks passed
@bri-prism
bri-prism deleted the perf/1bit-qmv-occupancy-prism branch June 5, 2026 14:49
@khosravipasha

Copy link
Copy Markdown
Collaborator

On m4 got 20% speed up so seems great :)

khosravipasha pushed a commit that referenced this pull request Jun 9, 2026
affine qmv_fast set packs_per_thread=2 for all bits except 2-bit, so 1-bit
got values_per_thread=64 (x_thread[64], ~256B/thread) -> low occupancy ->
1-bit decode saturates only ~75% of M5 DRAM BW vs ~90/96% for 2/4-bit.
Use 1 pack/thread for bits<=2 (values_per_thread=32), matching 2-bit's
register footprint.

Measured on M5 Pro (distinct-weight DRAM-bound, 2-bit as drift control):
1-bit 24.0 -> 21.9 us/matvec (~9%, 75->82% BW), 2-bit control 33.2->32.9
(0.8% drift); correct, rel_err 2.7e-4. Also makes scale_step_per_thread
(=group_size/values_per_thread) well-defined for group_size=32 at 1-bit.
khosravipasha pushed a commit that referenced this pull request Jun 9, 2026
affine qmv_fast set packs_per_thread=2 for all bits except 2-bit, so 1-bit
got values_per_thread=64 (x_thread[64], ~256B/thread) -> low occupancy ->
1-bit decode saturates only ~75% of M5 DRAM BW vs ~90/96% for 2/4-bit.
Use 1 pack/thread for bits<=2 (values_per_thread=32), matching 2-bit's
register footprint.

Measured on M5 Pro (distinct-weight DRAM-bound, 2-bit as drift control):
1-bit 24.0 -> 21.9 us/matvec (~9%, 75->82% BW), 2-bit control 33.2->32.9
(0.8% drift); correct, rel_err 2.7e-4. Also makes scale_step_per_thread
(=group_size/values_per_thread) well-defined for group_size=32 at 1-bit.
bri-prism added a commit to PrismML-Eng/mlx-swift that referenced this pull request Jun 10, 2026
…eader

Ports PrismML-Eng/mlx#3 into the pre-generated metal kernel header that mlx-swift compiles
into its metallib: qmv_fast_impl uses packs_per_thread = bits <= 2 ? 1 : 2, so 1-bit (Q1_0)
matches 2-bit (vpt=32) for better occupancy. ~+11% 1-bit decode on M5 (bit-exact, validated).
Follow-up: bump the Cmlx/mlx submodule to an mlx rev including #3 to make this permanent.
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.

2 participants