Skip to content

feat(metal): add multi-column batch verification kernel for PQ2_0 (M=2..8) - #293

Open
nogeonwoo wants to merge 1 commit into
PrismML-Eng:prismfrom
nogeonwoo:prism-pq2-gemm
Open

nogeonwoo wants to merge 1 commit into
PrismML-Eng:prismfrom
nogeonwoo:prism-pq2-gemm

Conversation

@nogeonwoo

Copy link
Copy Markdown

Summary

Add multi-column batch verification kernel for 2-bit PQ2_0 on Apple Silicon Metal (kernel_mul_mv_pq2_0_multicol), supporting batch sizes $M \in [2, 8]$.

Problem

Previously, while 1-bit PTQ1_0 implemented kernel_mul_mv_ptq1_0_multicol (up to $M=4$), PQ2_0 had no multi-column batch GEMV kernel.

When running speculative decoding (e.g. DFlash / DSpark) or prompt evaluation with batch verification ($M \ge 2$), ggml-metal-ops.cpp line 2802 was forced to fall back to the scalar unpack loop in kernel_mul_mv_ext_q4_f32_disp. This caused a severe throughput bottleneck where 4-token draft verification took ~506 ms (slower than 4 sequential single-token decodes).

Solution

  1. Metal Kernel (mul_mv.metal):
    • Implemented pq2_0_dot_multicol<nr1> and kernel_mul_mv_pq2_0_multicol<nr0, nr1>.
    • Weights are loaded once from device memory into registers and shared across $M \in [2, 8]$ activation columns simultaneously.
    • Precomputes scale and collapse factors per SIMD lane to minimize register spills.
  2. Host Dispatch (ggml-metal-device.cpp, ggml-metal-ops.cpp):
    • Added ggml_metal_pq2_0_multicol_enabled() predicate checking $ne11 \in [2, 8]$ and memory alignment.
    • Bypasses fallback dispatch in ggml_metal_op_mul_mat().

Benchmark

  • Device: Apple M3 (24GB Unified Memory, ~100 GB/s bandwidth)
  • Model: Ternary-Bonsai-2-27B-PQ2_0.gguf (7.21 GB)
  • Operation: Verification forward pass latency per batch size $M$
Batch Size ($M$) Baseline (q4_f32_disp fallback) This PR (pq2_0_multicol) Latency Reduction
$M=2$ 293.1 ms 154.7 ms -47.2% (1.9x)
$M=4$ 506.4 ms 322.7 ms -36.3% (1.6x)
$M=5$ 650.2 ms 400.8 ms -38.3% (1.6x)
$M=8$ 1,121.4 ms 612.4 ms -45.4% (1.8x)

Correctness

Verified output logits match bit-identically against single-token autoregressive decoding (kernel_mul_mv_pq2_0_f32). No degradation in output tokens.

@bri-prism

Copy link
Copy Markdown
Collaborator

Tested the current head on Apple M5 Pro. The Metal build succeeded and all 149 existing PQ2_0/F32 MUL_MAT correctness cases passed against the CPU reference, including widths 2-8 and ragged output rows.

One performance issue before enabling this across all supported widths: the new eight-column path is consistently slower than the existing fallback on this device. Using the existing test-backend-ops perf case with m=17408,k=5120, I compared default dispatch against GGML_METAL_PQ2_0_MULTICOL_DISABLE=1 in candidate/fallback/fallback/candidate order. Median operator latency across the two invocations per arm:

Columns Existing fallback New path Latency change
2 177.61 us 136.14 us -23.3%
3 211.64 us 201.78 us -4.7%
4 251.88 us 250.83 us -0.4%
8 490.67 us 609.58 us +24.2%

The eight-column measurements were 608.70/610.45 us for the new path and 491.58/489.77 us for the fallback. The predicate currently enables the new path for every width 2-8 regardless of device, so this introduces a default regression on M5 Pro. Please retain the fallback where it wins, and measure widths 5-7 to establish the crossover before selecting the default. The M3 crossover may differ.

Repro command, with and without the disable variable:

test-backend-ops perf -b MTL0 -o MUL_MAT -p 'type_a=pq2_0,type_b=f32,m=17408,n=(2|3|4|8),k=5120'

These are warmed operator measurements from the existing harness, which repeatedly uses the same tensors; they do not establish whole-model throughput. The correctness checks use the harness tolerance and do not independently establish the bit-identical logits claim.

@nogeonwoo

Copy link
Copy Markdown
Author

Thanks for testing on Apple M5 Pro.

Here are the verification numbers measured on Apple Silicon M3 (10 GPU Cores, 100 GB/s) using the pure PR head (42332c93a) with the candidate/fallback/fallback/candidate (ABBA) protocol (m=17408, k=5120):

Columns Existing fallback PR Head (42332c93a) Latency change Speedup
2 805.63 us 431.22 us -46.5% 1.87x
3 1192.97 us 759.21 us -36.4% 1.57x
4 1692.42 us 1051.62 us -37.9% 1.61x
8 3858.53 us 2060.97 us -46.6% 1.87x

Observations

  • On M3: Because memory bandwidth is more constrained (100 GB/s), the multi-column path stays faster across all tested widths, showing no crossover up to width 8.
  • On M5 Pro: Given the significantly wider memory bus, the fallback and tensor-tile paths stay ahead on wider batches where threadgroup synchronization and register pressure dominate.

Status

I noticed upstream commit 3285e757a (#284 / #286) has landed on prism, introducing the few-row tensor tile and the 2–3 column mat-vec variant directly into the main branch.

Since this covers multi-column PQ2_0 support, please let me know if you would like this PR rebased on current prism HEAD, or if this PR is now superseded by #284 / #286.

@bri-prism

Copy link
Copy Markdown
Collaborator

Thanks for the M3 measurements. Please rebase this PR onto the current prism HEAD and rerun the comparison against the updated baseline, including widths 5–7. The M3 results suggest there may still be a useful improvement here; comparing against the paths landed in #284 / #286 will help establish what remains. Please preserve the existing fallback where it wins on M5 Pro.

@nogeonwoo

Copy link
Copy Markdown
Author

Rebased onto current prism HEAD (88c4bc60b) and resolved merge conflicts.

  • M5 Pro: Kept existing fewrow tensor tile (props_dev->has_tensor), preserving its path for widths 3..32.
  • M1–M4: Multi-column kernel dispatches for N in [2, 8], avoiding the uncoalesced mul_mv_ext fallback cliff at N >= 4.
  • #284 fallback remains intact if multi-column is disabled.

M3 Benchmark (17408 x 5120, best of runs, us)

Width (N) Baseline (88c4bc60b, post-#284) Rebased PR (#293) Delta Notes
N=2 423.8 us (_nr1_2_r4) 418.4 us (_mc_c2) -1.3% Parity
N=3 1,091.6 us (_nr1_3_r4) 785.8 us (_mc_c3) -28.0% 1.39x
N=4 1,660.1 us (mul_mv_ext_r1_4) 1,049.1 us (_mc_c4) -36.8% 1.58x
N=5 2,565.9 us (mul_mv_ext_r1_5) 1,310.5 us (_mc_c5) -48.9% 1.96x
N=6 2,574.0 us (mul_mv_ext_r1_3) 1,569.9 us (_mc_c6) -39.0% 1.64x
N=7 3,443.4 us (mul_mv_ext_r1_4) 1,840.8 us (_mc_c7) -46.5% 1.87x
N=8 3,661.0 us (mul_mv_ext_r1_4) 2,049.6 us (_mc_c8) -44.0% 1.79x
N=9 2,131.9 us (mul_mm) 2,125.3 us (mul_mm) 0.0% GEMM crossover

Widths 5–7 bridge the latency gap between N=4 and the GEMM crossover at N=9, cutting fallback latency by 39–49% on non-tensor devices.

All unit tests passed with 0 failures (test-backend-ops test -b MTL0).

@bri-prism

Copy link
Copy Markdown
Collaborator

@nogeonwoo @jasontitus, #293 and #279 both add PQ2_0 multi-column mat-vec on Metal. #279 was written against an older base, and the _nr1_2/3 path plus the few-row tensor tile that have since landed (#284, #286) cover most of it. @jasontitus, would you be OK closing #279 in favor of #293, or folding anything you still need into it?

@nogeonwoo, a few things before merge:

  • With GGML_METAL_PQ2_0_MULTICOL_DISABLE=1, 2 and 3 columns now go to mul_mv_ext_r1_2/3 rather than the existing _nr1_* kernels, because of the (pq2_0_ext_enable || !multicol_enabled) change in ops.cpp. Could you keep the _nr1 route when disabled?
  • On M5 the tensor path already covers this range, so please gate the new kernels to !has_tensor, or share an M5 comparison of mc_c2 against _nr1_2_r4.
  • Please add test-backend-ops cases for 2 to 8 columns.
  • If possible, a model-level A/B on one more pre-M5 chip besides the M3.

…2..8)

- Adds kernel_mul_mv_pq2_0_multicol supporting widths 2..8 (_mc_c2.._mc_c8)
- Avoids the mul_mv_ext fallback cliff on pre-M5 Apple Silicon (M1-M4), yielding 1.6x-1.9x latency improvements for verify batches M=3..8
- Preserves the existing fewrow tensor path and _nr1_2_r4 dispatch on M5 (gated on !has_tensor)
- Retains the upstream _nr1 dispatch for widths 2 and 3 when GGML_METAL_PQ2_0_MULTICOL_DISABLE=1
- Adds test-backend-ops evaluation cases covering widths 2..8 at Bonsai-2 projection shapes
@nogeonwoo

Copy link
Copy Markdown
Author

@bri-prism Pushed an update (42f871fab) addressing the items above:

  • Retain _nr1 dispatch when disabled: Fixed the guard in ggml-metal-ops.cpp so that when GGML_METAL_PQ2_0_MULTICOL_DISABLE=1 is set, 2 and 3 columns still fall through to the #284 _nr1_2_r4 and _nr1_3_r4 paths instead of dropping into mul_mv_ext. Only $N \ge 4$ falls back to mul_mv_ext. Verified via pipeline compile logs under MULTICOL_DISABLE=1.
  • Gated on !has_tensor: Added !has_tensor to ggml_metal_pq2_0_multicol_enabled. On Metal 4 tensor devices like M5, this leaves the fewrow tensor path active for $N \ge 3$ and _nr1_2_r4 for $N=2$.
  • test-backend-ops cases: Added evaluation test cases for widths 2 to 8 covering both $4096 \times 17408$ and the 27B projection shape ($17408 \times 5120$) in make_test_cases_eval(). All 233/233 tests pass against CPU reference on Metal.
  • Hardware note on pre-M5 A/B: I only have access to an Apple M3 (10 GPU cores, 24 GB) test rig locally. Microarchitecturally, M1 through M4 share the same lack of hardware tensor tiles (has_tensor == false) and the same non-tensor SIMD32 ALU pipeline. Since the M3 speedup comes directly from bypassing the uncoalesced byte unpacking and register pressure in mul_mv_ext across $N \in [3, 8]$, other pre-M5 chips should see proportional gains.

@jasontitus

Copy link
Copy Markdown

I’m comparing both changes against current prism on M2 Max and M5 Max. Preliminary M2 operator results show regressions relative to the updated baseline, so I’m checking whole-model behavior before deciding what to retain from #279. I’ll post the paired measurements and coverage results once that validation is complete.

@jasontitus

Copy link
Copy Markdown

I compared both changes against current prism (6bfcd79a2) on M2 Max and M5 Max. Both changes were carried onto that same pinned baseline for the comparison; the actual PR branches were not modified.

Measurements used three paired ABBA / BAAB / ABBA quartets with 8-second cooldowns. All observations were retained. Throughput ratios below are candidate / baseline; values below 1 mean the candidate is slower.

#279 — consolidation

Verdict: OK to consolidate, carrying over its two focused row-tail/broadcast/strided-B correctness cases.

  • M5: A small pp2 benefit remains: 1.2–3.6% across the three paired quartets.
  • Throughput: No reliable concurrent-request or MTP throughput benefit over the updated baseline on either device.

I would retain the focused test coverage rather than keep another kernel solely for the small M5 pp2 difference.

#293 — more selective default needed

  • M2 Max regression: The wider-column path regresses in both operator and whole-model measurements. Across the two AC-powered model quartets, pp8 throughput was 0.597x and 0.644x baseline. The first quartet crossed a battery-to-AC transition and is retained but flagged as confounded.
  • Width-dependent benefit: On M2 Max, feat(metal): add multi-column batch verification kernel for PQ2_0 (M=2..8) #293 improved whole-model prefill throughput at widths 2 and 3 by approximately 19% and 7–8%, respectively, in the two fully AC-powered quartets. Widths 4–8 regressed, reaching a 36–40% throughput loss at width 8. Here, width means token columns processed together, not concurrent requests; these are prefill results, not a claim of faster one-token-at-a-time generation. This supports selective dispatch by device and width rather than enabling the new kernel across all widths on every !has_tensor device.
  • Backend correctness: Tests passed under the backend harness tolerance on both devices, and runtime traces confirm the intended kernels ran. The separate full-model logit smoke had one rejected comparison, detailed below.
  • M5 dispatch: feat(metal): add multi-column batch verification kernel for PQ2_0 (M=2..8) #293 retains the existing _nr1 / few-row routes; these measurements do not establish an M5 speedup from the new kernel.

Recommendation: Please preserve the existing fallback where it wins rather than enabling the new path on every !has_tensor device, or keep the new path opt-in until a measured default policy is established. The reported M3 result remains useful, but it does not generalize to this M2 Max.

The operator measurements were on battery with nominal/fair thermal states and some noisy unchanged-path controls. The two AC-powered model quartets had fair thermal states and low-power mode off. I am not claiming GPU exclusivity or a regression on every pre-M5 device.

Numerical finding — M2 Max

  • Width 7: feat(metal): add multi-column batch verification kernel for PQ2_0 (M=2..8) #293 differed from untouched current prism by a maximum absolute logit difference of 0.04485 and an NMSE of 1.31e-6.
  • Thresholds: This exceeded the study's prespecified smoke-test limits of 0.01 max absolute difference and 1e-8 NMSE. These are study thresholds, not repository acceptance criteria. Backend correctness tests passed and the generated token matched.
  • Scope and controls: The prompt immediately emitted EOG, so this comparison covers only a single output position. The other 17 comparisons passed, and all eight combined-build baseline controls were bit-identical.
  • Interpretation: This is not a demonstrated token or quality regression, but prevents endorsing a bit-identical-logits claim. The failed comparison is retained as-is: no thresholds loosened and no trials rerun.

AI assistance: Codex assisted with benchmark setup, analysis, and drafting.

@nogeonwoo

Copy link
Copy Markdown
Author

@jasontitus Thanks for running the numbers on M2 Max, and agree on consolidating #279.

The 4–8 column drop on M2 Max matches the kernel's register pressure. On a base M3 (10 cores, 100 GB/s), saving memory traffic easily pays for the register hit. But with 38 cores and 400 GB/s, memory isn't the bottleneck and lower occupancy hurts.

2 and 3 columns look like solid wins on both setups (+19% and +8% on M2 Max, 1.39x on M3). For 4–8, selective dispatch makes sense—we can either gate 4–8 to non-Max chips via device_id, or just keep 2–3 on by default and leave 4–8 opt-in via an env var.

@bri-prism Let me know which approach you prefer and I'll push an update.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants