Skip to content

ggml-cpu : AVX-512 tiled kernel for single-token flash attention (faster long-context decode) - #254

Merged
bri-prism merged 1 commit into
PrismML-Eng:prismfrom
lenny76:cpu-fa-decode-avx512
Sep 25, 2026
Merged

bri-prism merged 1 commit into
PrismML-Eng:prismfrom
lenny76:cpu-fa-decode-avx512

Conversation

@lenny76

@lenny76 lenny76 commented Sep 23, 2026

Copy link
Copy Markdown

Overview

Token generation at long context on CPU is dominated by FLASH_ATTN_EXT: with Ternary-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_chunk once 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:

  • converts the K and V rows of the tile from f16 to f32 once, for all the q heads of the group;
  • computes the 16 scores of each head (q pre-scaled), applies mask/ALiBi slope, and does the online softmax on the whole tile with a vector exp;
  • updates each head's f32 VKQ accumulator 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 as one_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_EXT comparing the CPU backend against the CPU reference path (test-only local change so the CPU is not skipped), with extra single-token cases (kv 512 / 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, base prism + #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:

depth base tg32 PR tg32
0 5.43 5.39
8192 4.71 4.99
16384 4.02 4.61 / 4.60

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=all gives 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

  • I have read and agree with the contributing guidelines
  • AI usage disclosure: YES. Claude Code (Anthropic) wrote the patch and ran the measurements above on my machine, under my direction.

🤖 Generated with Claude Code

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 bri-prism left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@bri-prism
bri-prism merged commit 78e1640 into PrismML-Eng:prism Sep 25, 2026
3 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants