ggml-cpu : AVX-512 tiled kernel for single-token flash attention (faster long-context decode) - #254
Conversation
The split-KV decode path called one_chunk once per q head, so each f16 KV row was loaded and converted once per q head sharing it. The new kernel handles all q heads of a KV head together, 16 KV rows at a time: K/V converted to f32 once per tile, scores and online softmax on the whole tile with a vector exp, and each head's VKQ accumulator updated 16 floats at a time in a register. It writes the same [M, S, VKQ] partials, so the reduction is unchanged. Used only for F16 K/V, F32 Q, no softcap, D % 16 == 0, D <= 512, G <= 16. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
bri-prism
left a comment
There was a problem hiding this comment.
Tested on an AMD EPYC Genoa (AVX-512 F/BW/VL, VNNI, BF16; 8 vCPUs), CPU only, 2B PQ2_0 model. One heads-up first: this branch is based on an older prism (bdc23b56) that segfaults while loading the model, in ggml_backend_cpu_repack_buffer_set_tensor. The merge-base crashes the same way without your change, so it isn't caused by this PR. Merged onto current prism it builds and runs cleanly, and a rebase would pick that up.
On current prism plus this PR: output matches the base (KLD 0 with batched prompts; with -ub 1, which exercises the single-token path, mean KLD 0.00015 and same top token 98.8%). Decode with -fa 1, two rounds: 16K context 22.7 to 29.6 t/s (+30%), 8K 32.5 to 33.4-34.9 t/s, and no change at empty context. Nice.
Overview
Token generation at long context on CPU is dominated by
FLASH_ATTN_EXT: withTernary-Bonsai-2-27B(16 full-attention layers, 24 q heads, 4 KV heads, head dim 256) decode drops from ~6.5 t/s at depth 0 to ~4.3 t/s at 16k, and the extra time is all in attention.The split-KV decode path (one query row) currently calls
one_chunkonce per q head, so every KV row is loaded and converted from f16 once per q head sharing it (6x with GQA 6), and each head does its dot products and online softmax row by row.This PR adds an AVX-512 kernel for that path. For each KV head and each tile of 16 KV rows it:
exp;VKQaccumulator 16 floats at a time in a register (acc = acc*ms + sum_t p[t]*V[t]).It writes the same
[M, S, VKQ]partials asone_chunk, so the existing reduction is unchanged; sinks are applied on the first chunk as before.It is used only when the build has AVX-512F/DQ + F16C and: K and V are F16, Q is F32, no logit softcap,
D % 16 == 0,D <= 512, at most 16 q heads per KV head,ne[3] == 1. Everything else goes through the existing code.Additional information
2x Xeon Gold 6262 (Cascade Lake), Linux, GCC 14,
-DGGML_NATIVE=ON.Correctness:
test-backend-ops -o FLASH_ATTN_EXTcomparing the CPU backend against the CPU reference path (test-only local change so the CPU is not skipped), with extra single-token cases (kv512 / 1000 / 4096 / 16384; D 64/128/256; GQA 1/4/6/8/12; mask, ALiBi, softcap, sinks): 5182/5182 pass. Greedy continuation of a 7511-token prompt (64 tokens) is identical to the base.Ternary-Bonsai-2-27B-PQ2_0.gguf, baseprism+ #245 (needed to load PQ2_0 on AVX-512),numactl --interleave=all llama-bench -ngl 0 --numa distribute -t 24 -p 0 -n 32 -d D:With all the other open CPU PRs applied (#245 #246 #249 #250 #251): attention time per token at 8k 31 -> 15.5 ms, generation after a 7511-token prompt 5.48 -> 6.28 t/s, tg32 at 16k 4.34 -> 4.96 t/s with default memory policy. With the PR,
numactl --interleave=allgives 5.67 t/s at 16k: without it the KV cache sits on one node and attention becomes limited by that socket's bandwidth.Requirements
🤖 Generated with Claude Code