Repository navigation
quant: NEON arm for batched matmul_fp8 (bit-exact, 6–10× at prefill shapes) - #1761
Conversation
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.
|
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 (
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 |
|
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:
On SVE, gcc vectorises the scalar loop with Suggested fix: use |
…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).
|
Good catch, and the diagnosis is exactly right — thank you for running it under qemu. Fixed in e2e3e0c, both sides of the gap:
Evidence on M4 (gcc-15 15.2.0 and clang-17,
The unfused non-SVE build is the same mechanism class as your SVE |
|
@mfethe1 ran it. Same weights (REAP 150B), same M5 Pro (48 GB), same three prompts as the #1696 breakdown, TTFT
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):
What I read from it:
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 Taken together with #1730, the 577-token TTFT on this machine went from 371 s (pre-1.12.1) to 162 s. Thanks! |
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 (andv4_ds_fp8_viewonly setsblock_rows = 8under AVX2), so arm64 falls through to the scalarmatmul_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
matmul_fp8kernels (clang four-row and GCC upstream) are renamed tomatmul_fp8_scalar. Their bodies are untouched.matmul_fp8_neon(#ifdef __ARM_NEON):vfmaq_f32chain per (row, token) from zero.-ffp-contract=on) and GCC (arm64 defaultfast) emit for the scalar form. The double-precision block combine is unchanged.matmul_fp8dispatches to the NEON arm when S ≥ 4. Decode (S < 4) stays scalar, since one token can't pay back the tile decode.fp8_format.handdeepseek_v4.care untouched.Exactness
The output is bit-identical to the scalar kernel, not just within tolerance.
tests/test_fp8_passthroughaddsrun_batch_exact, which comparesmatmul_fp8againstmatmul_fp8_scalarwithmemcmp. It covers: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)
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-batchis this PR. The FP8 and attention rows would confirm or refute it.