Skip to content

quant: NEON arm for batched matmul_fp8 (bit-exact, 6–10× at prefill shapes) - #1761

Merged
JustVugg merged 2 commits into
JustVugg:devfrom
mfethe1:neon-fp8-batch
Oct 1, 2026
Merged

JustVugg merged 2 commits into
JustVugg:devfrom
mfethe1:neon-fp8-batch

Conversation

@mfethe1

@mfethe1 mfethe1 commented Sep 26, 2026

Copy link
Copy Markdown
Contributor

NEON arm for the batched FP8 matmul (matmul_fp8)

Follow-up to #1696 / #1730. @Freizeitminister's prefill breakdown on v1.12.1 (M5 Pro, REAP 150B, comment) showed that on Apple Silicon the next-largest cost after the FP4 experts is the FP8 attention projections: 59 s of the 577-token TTFT (27%), 39 s of it in wo, scaling linearly with prompt length. The reason is structural: fp8_batch_compute's tiled fast path is #ifdef __AVX2__ only (and v4_ds_fp8_view only sets block_rows = 8 under AVX2), so arm64 falls through to the scalar matmul_fp8. Thanks for publishing those numbers and the breakdown. They are the end-to-end TTFT figures the #1730 thread was asking for, and they pointed straight at this gap.

Change

  • The two existing matmul_fp8 kernels (clang four-row and GCC upstream) are renamed to matmul_fp8_scalar. Their bodies are untouched.
  • New matmul_fp8_neon (#ifdef __ARM_NEON):
    • Covers 16 output rows × 4 tokens per pass.
    • Decodes each 128-column block's rows once into a column-major tile, then runs one vfmaq_f32 chain per (row, token) from zero.
    • Uses the same fused multiply-adds, in the same order, that clang (-ffp-contract=on) and GCC (arm64 default fast) emit for the scalar form. The double-precision block combine is unchanged.
  • matmul_fp8 dispatches to the NEON arm when S ≥ 4. Decode (S < 4) stays scalar, since one token can't pay back the tile decode.
  • x86, CUDA, fp8_format.h and deepseek_v4.c are untouched.

Exactness

The output is bit-identical to the scalar kernel, not just within tolerance.

  • tests/test_fp8_passthrough adds run_batch_exact, which compares matmul_fp8 against matmul_fp8_scalar with memcmp. It covers:
    • O not a multiple of 16 (zero-padded row tile)
    • S not a multiple of 4 (clamped token lanes)
    • S > 64 (chunk edge)
    • I not a multiple of 128
    • one NaN byte
  • Mutation check: perturbing one FMA lane, or dropping the 4th token's store, fails 6/6 cases.
  • make check: exit 0 on this M4 (clang, OpenMP). All C suites ok, Ran 1515 tests … OK (skipped=130).

Kernel numbers (Apple M4, 10 cores, OpenMP, best of 5, bit-exact on every shape)

S × I × O clang: scalar → NEON gcc-15 -O3: scalar → NEON
577 × 8192 × 4096 1259 → 145 ms (8.7×) 2843 → 278 ms (10.2×)
577 × 4096 × 1024 114 → 18.8 ms (6.1×) 295 → 28.0 ms (10.5×)
189 × 4096 × 1024 35.6 → 5.1 ms (7.0×) 109 → 11.7 ms (9.4×)
189 × 1024 × 4096 33.3 → 4.4 ms (7.6×) 135 → 12.8 ms (10.5×)
128 × 4096 × 1024 23.6 → 3.3 ms (7.1×) —
16 × 4096 × 1024 2.96 → 0.59 ms (5.1×) 9.18 → 1.69 ms (5.5×)
4 × 4096 × 1024 0.86 → 0.37 ms (2.3×) 2.87 → 0.83 ms (3.5×)

Test boundary: these are synthetic-weight kernel timings on one M4. I haven't measured TTFT on real weights: the model doesn't fit in the free disk on the machine I have. If the kernel gain carries over, the 59 s FP8 share of the 577-token run should drop to single-digit seconds, but that is an estimate. @Freizeitminister, if your offer to re-run the prefill breakdown on a branch still stands, mfethe1:neon-fp8-batch is this PR. The FP8 and attention rows would confirm or refute it.

fp8_batch_compute's tiled path is AVX2-only, so arm64 prefill ran the scalar
matmul_fp8 -- 27% of a 577-token TTFT on M5 Pro (JustVugg#1696 breakdown).  Add a
16-row x 4-token vfmaq kernel reproducing the scalar FMA chain and double
block combine exactly; dispatch at S>=4.  6-8.7x (clang) / 9-10.5x (gcc) on
M4 at prefill shapes; test_fp8_passthrough asserts bit-identity.
@mfethe1

mfethe1 commented Sep 27, 2026

Copy link
Copy Markdown
Contributor Author

Independent reproduction on a second M4 (Mac mini, 10 cores, 24 GB, AC power) — bit-exact on every shape, both compilers, and the speedups hold:

Built from this PR branch (3a3ba2b), no source edits. tests/test_fp8_passthrough passes under both gcc-15 15.2.0 (-O3 -march=armv8.2-a+dotprod -fopenmp) and Apple clang 17 (-O3 + libomp). The bench below interleaves scalar vs dispatched matmul_fp8 best-of-7 with a memcmp over the full output each round:

S × I × O gcc-15 scalar → NEON clang scalar → NEON bit-exact
577×8192×4096 (wo-shaped) 1966 → 211 ms (9.3×) 1046 → 115 ms (9.1×) yes
577×4096×1024 233 → 26 ms (9.1×) 124 → 16 ms (7.9×) yes
189×4096×1024 76 → 8.4 ms (9.0×) 40 → 5.2 ms (7.8×) yes
189×1024×4096 77 → 8.5 ms (9.0×) 41 → 5.2 ms (8.0×) yes
16×4096×1024 6.5 → 1.1 ms (6.2×) 3.2 → 0.6 ms (5.2×) yes
4×4096×1024 (dispatch boundary) 1.7 → 0.5 ms (3.4×) 0.8 → 0.4 ms (2.3×) yes
1×4096×1024 (decode, scalar path) 0.96× — correctly not dispatched 0.99× yes

Numbers land within a few percent of the PR table on gcc and slightly under on clang, same shape-to-shape profile. The S≥4 dispatch boundary behaves as documented.

Also took a first-pass census of the loose thread from #1696 (other __AVX2__-only fast paths falling to scalar on arm64): beyond the two already fixed in #1730/#1761, there are AVX2-gated sites in inkling.c, qwen36.c, qwen38_core.h, olmoe.c, kimi_k3.c, expert_ffn.h, sparse_attn.h, fused_simd.h, gsgemv.h/qgemv.h (these two have SSE4.1 arms that also miss arm64), and more in deepseek_v4.c/deepseek_v41.c. Happy to file the triaged follow-up issue if wanted — several look like the same one-day NEON-port shape as this PR, but which ones are hot needs a profile per engine.

@JustVugg

Copy link
Copy Markdown
Owner

Thanks for the kernel. The "bit-exact" claim does not hold on SVE CPUs (Graviton3/4, Grace). We compared the new NEON path with the scalar one using gcc for aarch64 under qemu and full-mantissa inputs:

build outputs that differ
-mcpu=neoverse-v1, neoverse-v2, armv8.2-a+sve 3054 of 4096
armv8.2-a (no SVE) 0 of 4096

On SVE, gcc vectorises the scalar loop with fadda and does not fuse the multiply-add, while your NEON path does. The test cannot see it because rndf() in test_fp8_passthrough.c makes every product exact, so it also gives 0 differences on SVE.

Suggested fix: use fmaf in the scalar loop on aarch64 (or skip the NEON arm when __ARM_FEATURE_SVE is set), and give the test full-mantissa inputs so it would catch this.

…a test inputs

SVE review (Graviton3/4, Grace): gcc auto-vectorises the scalar loop with
unfused fadda while the NEON arm uses vfmaq, so bit-identity broke on
neoverse-v1/v2 and armv8.2-a+sve (3054/4096 outputs differed).  The test
could not see it because rndf()'s 20-bit-mantissa inputs make every
e4m3*x product exactly representable.

Fix both sides of the gap: __builtin_fmaf in the scalar accumulation on
__aarch64__ (both the clang four-row and gcc one-row arms) pins the chain
fused regardless of -ffp-contract or the vectoriser, and run_batch_exact
now draws x from full 23-bit-mantissa [1,2) values so fused-vs-unfused
divergence is observable.

Evidence (M4): pre-fix + new inputs + -ffp-contract=off fails 6/6 shapes
(non-fused scalar is the SVE-fadda mechanism class); post-fix passes
gcc-15 and clang-17 across default and -ffp-contract=off.  Default-build
scalar output on non-SVE aarch64 is unchanged (the chain was already
contracted there).
@mfethe1

mfethe1 commented Sep 28, 2026

Copy link
Copy Markdown
Contributor Author

Good catch, and the diagnosis is exactly right — thank you for running it under qemu.

Fixed in e2e3e0c, both sides of the gap:

  1. Kernel (c/quant.h): the scalar accumulation now uses __builtin_fmaf on __aarch64__ (both the clang four-row and gcc one-row arms), pinning the chain fused regardless of -ffp-contract or the auto-vectoriser, so the SVE fadda lowering can't diverge from vfmaq anymore. I went with the fmaf route rather than #ifdef __ARM_FEATURE_SVE-skipping the NEON arm: it also keeps the reference correct for anyone building with -ffp-contract=off, and the default-build scalar output on non-SVE aarch64 is unchanged (the chain was already contracted there, so this is a no-op on M-series).

  2. Test (test_fp8_passthrough.c): run_batch_exact now draws x from full 23-bit-mantissa values in [1,2) — e4m3 * x needs 27 bits, so fused vs unfused chains round differently and the test can actually catch this class. You're right that the old rndf() (20-bit mantissa) made every product exactly representable; confirmed on M4 that even -ffp-contract=off passed the old inputs — the blindness was total.

Evidence on M4 (gcc-15 15.2.0 and clang-17, -O3 -march=armv8.2-a+dotprod):

build old inputs full-mantissa inputs
pre-fix, default pass pass
pre-fix, -ffp-contract=off pass (blind) fail 6/6 shapes
post-fix, default pass pass
post-fix, -ffp-contract=off — pass

The unfused non-SVE build is the same mechanism class as your SVE fadda row, so I can't run your exact matrix locally — but if your qemu harness is still set up, the full-mantissa inputs + fmaf build should now show 0/4096 on neoverse-v1/v2/armv8.2-a+sve too. If any SVE cell still differs, that would point at a second divergence source and I'd want to see it.

@Freizeitminister

Copy link
Copy Markdown

@mfethe1 ran it. Same weights (REAP 150B), same M5 Pro (48 GB), same three prompts as the #1696 breakdown, RAM_GB=24, DSV4_ATTN_PROF=1, AC power. To isolate the PR I built its merge base on dev (4e28e398) and the PR head (e2e3e0c) with the same toolchain (Apple clang 14, OpenMP) and ran them back to back. test_fp8_passthrough with the new full-mantissa inputs passes here too.

TTFT

prompt base 4e28e398 #1761 e2e3e0c
189 tok 60.9 s 44.0 s −28 %
577 tok 214.7 s 162.1 s −25 %
189 tok warm 66.0 s 51.0 s −23 %

Outputs are identical between base and PR for all three prompts (the 577-token one also matches the 1.12.1 run from #1696 word for word), so the bit-exact claim holds on real weights here. Decode is unchanged within noise, as expected with S < 4 staying scalar.

Breakdown, seconds summed over all 43 layers (blockprof sums match TTFT within 0.2 s):

189 base → PR 577 base → PR 189 warm base → PR
attention block 33.1 → 19.8 121.0 → 79.0 34.1 → 20.8
– FP8 projections (qa+qb+kv+wo) 19.2 → 5.3 58.7 → 15.6 19.5 → 5.8
– of which wo 12.7 → 3.9 39.1 → 11.4 13.0 → 4.4
– core attention (attn) 10.2 → 10.5 50.0 → 50.5 10.9 → 11.0
– compressor + indexer 3.6 → 3.9 12.1 → 12.8 3.6 → 4.0
MoE 24.7 → 21.2 84.5 → 73.7 28.9 → 27.1
hc1+hc2+hc3 3.0 → 3.0 9.1 → 9.2 3.0 → 3.1

What I read from it:

  1. FP8 projections: 3.4–3.8× end to end (58.7 → 15.6 s at 577 tokens). That's short of your 6–10× kernel numbers and of the single-digit estimate. I haven't split the remaining 15.6 s further. It may include work around the kernel inside the same attnprof windows (e.g. the activation qdq), but that is a guess, not measured.
  2. MoE also got ~11 s faster at 577 tokens, although the PR doesn't touch it. I think that's the shared expert: v4_shared_expert_forward_batch_ref runs gate/up/down through coli_fp8_matmul_batch_ref → matmul_fp8, so it picks up the NEON arm too. That fits the numbers, but I didn't profile the shared expert separately.
  3. New order at 577 tokens: MoE 73.7 s (45 %), core attention 50.5 s (31 %), FP8 projections 15.6 s (10 %). Core attention is untouched and still grows faster than the prompt.

Test conditions: the base run landed within 1 % of the 1.12.1 TTFT from #1696 (214.7 vs 217.0 s), so the 55 dev commits in between don't move this. Both runs had the full 24 GiB budget with no swap. macOS compressed more of the expert cache than on Sep 25 (browsers open), but by a similar amount in both runs.

Taken together with #1730, the 577-token TTFT on this machine went from 371 s (pre-1.12.1) to 162 s. Thanks!

@JustVugg
JustVugg merged commit 47f3569 into JustVugg:dev Oct 1, 2026
30 checks passed
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.

3 participants