chore(train): engine-10 continued: pre-sampling probabilities on the prefilter, the live-tile attention's one copy, one scan and aligned split - #72
Merged
Conversation
…er, the row's softmax taken on the backend A request for pre-sampling probabilities (n_probs without post_sampling_probs) turned the top-k prefilter off, since the probabilities are under the whole row: every decode copied each output row's n_vocab logits to the host, the CPU chain scanned them all for its top k, and get_token_probabilities built a 248,320-entry vector, sorted its head and ran 248,320 expf a token for a softmax whose only use of the whole row is its sum. On an RTX 5090 at 245,752 tokens, Ternary Bonsai 2 27B greedy with n_probs 5 decoded 93.7 tok/s against 107.4 without, and drafted 146-167 against 180-214 (the engine gate's P11/P11n and D11/D11n legs). llama_sampler_init_row_probs() gives a row the probabilities of its logits on the backend and changes no logit, and a backend top-k keeps probabilities a sampler before it produced aligned with its candidates. With at most k probabilities asked, the server's prefilter chain is [row-probs, top-k]: a decode copies k logits, ids and probabilities a row, the CPU chain draws from the k as before, and get_token_probabilities reads each candidate's probability under the whole row instead of renormalizing the k. More than k asked keeps every logit on the CPU; LLAMA_TOP_K_PREFILTER_LEGACY=1 still does for every request. RTX 5080, Ternary Bonsai 2 27B served (q4_0 K/V, f16 state, -c 32768), three prompts greedy at 256 tokens with n_probs 5, 3 rounds rotated, medians of 9 against the same requests without n_probs: plain 95.65 tok/s with every logit on the CPU (-9.7 %), 105.14 with the prefilter (-0.7 %), 105.92 without; with the MTP draft 166.20 (-13.5 %), 192.51 (+0.2 %), 192.15. The same tokens on every leg and the same top-5 ids; log-probabilities differ by up to 1.9e-3, the CPU's float running sum over 248,320 logits (the backend's softmax matches a double-precision one to 4e-7 below). test-backend-sampler row_probs_top_k: the k candidates' probabilities against a CPU softmax of the same prompt's full logits on a sequence without a backend sampler, to 1e-5 (Qwen3 0.6B Q4_0, the suite's CI model: 22/22); without top-k's gather it fails on the count. test_top_k_prefilter_pre_sampling_probs: the same tokens, top-n ids and log-probabilities within 2e-5 as with every logit on the CPU (measured gap 2.9e-6), greedy and seeded at n_probs 5 and 40; reading the k's softmax instead fails all three by 2.3e-4 and more. The server's unit suite without the slow tests: 362 pass; test_compat_anthropic's 27 fail only where port 8082 is taken by another local service, and all 28 pass with its server on a free port.
Pre-sampling probabilities on the top-k prefilter: what n_probs cost a served decode, the row-probs sampler and the top-k's gather, the A/B on a 5080 and the tests shown able to fail.
…am-k roles as arguments The mma kernel instantiated its tile code (flash_attn_ext_f16_process_tile, and flash_attn_ext_f16_iter under it) once per stream-k role, as template parameters: a block that ends inside a tile (it writes its partial to the fixup buffer), one that ends a tile it did not start (it writes to dst and needs the fixup), one that did a whole tile. The loop never reads the role; only the epilogue's destination depends on it. Ternary Bonsai 2 27B's decode kernel (head 256, q4_0 K/V read raw) was 184 KB of SASS for three copies. A per-block timeline (globaltimer at entry, before the main loop, after it and at the end; a probe build, not landed) of one decode token at 65,536 and 245,760 cells on an RTX 5080 (252 blocks, 4 KV heads, 63 blocks a head) showed the blocks that end a tile, 62, 125, 188 and 251, ending last at both depths, 7.9 and 8.9 us after the median block; their main loops ran as fast as every other block's, their prologues took 7.8-10.5 us against a median of 3.7-4.2 and their epilogues 1.2-2.2 us against 1.0. Only those 4 blocks of 252 ran the needs-fixup copy, so its instructions were in no SM's instruction cache and, after 75-283 MB of K/V streamed through the 64 MB L2, in no cache at all. The roles are now arguments of one copy: the kernel is 71 KB of SASS (255 registers, 16 bytes of stack against 24), libggml-cuda 79.1 -> 73.2 MB. With the probe, the last block ends 3.3-4.2 us after the median, and the tile-ending blocks among the rest. Against the parent's libggml-cuda under one test-backend-ops (the served attention, K and V laid out as the cache's views are), 4 rounds, the order rotated, medians (a first run of 4 agreed within 0.6 points): a token: 16,384 -2.8 % 65,536 -4.2 % 131,072 -2.2 % 245,760 -0.9 % (-0.7, -4.5, -4.2, -2.9 us) a 4-row verify: 16,384 +0.4 % 65,536 -3.1 % 131,072 -1.0 % 245,760 -1.1 % (+0.1, -3.4, -2.0, -3.7 us) (at 16,384 cells the test's K and V stay in the L2 between runs, as they do not in a model's graph). llama-bench at 16,384, one sequence (the mask-prefix path), 4 rounds of 5 samples after 2 warm-ups: tg64 103.64 -> 103.47 tok/s, pp4 (-rs 3) 381.97 -> 376.14; the next commit's library, the same code on that path, read 103.75 and 381.18: inside the noise. The served drafted decode at bench-head's short depths, 4 slots: 211.82 -> 212.56 tok/s. Checks: FLASH_ATTN_EXT 3,220/3,220. The served drafted decode (the served pack, 4 slots with --kv-unified: the live-tile path), bench-head's six greedy requests at 512 tokens: the texts identical across all six legs of both arms; every value is computed as before and written where it was. No switch: the same arithmetic and the same stores, only the code's layout changed.
…t once a layer The mma kernel's live-tile split (KV_live) runs a memset and a scan of the mask (flash_attn_mask_to_KV_live) before every flash attention: 3.2-4.5 us a call in the graph on an RTX 5080 at a decode token and a 4-row verify (nsys, each node timed: the memset 0.26 us, 0.6 us to the scan, the scan 1.6-3.0, 0.8 us to the attention). The live steps depend on the mask and the split alone, and every attention layer of a graph reads the same mask (Ternary Bonsai 2 27B: 16 of them), so 15 of 16 scans found what the first had. ggml_cuda_fattn_kv_live_context keeps the first scan of a graph evaluation, keyed by what it read and how it split (the mask tensor and its data, shape and strides, the stream, ncols1, nbatch_fa, the steps, the Q and output tiles, the sequences); the others read it, and any other key scans again. It is reset where the graph evaluator resets its other per-graph records (the Gated DeltaNet gathers, the conv-state updates), so a CUDA graph captures the scan once, in its first attention, and every replay writes it there before the layers after read it. The memory is the context's and is freed only with it: a buffer too small for a later key is kept, since a graph captured before may still name it. Serving takes this path with more than one slot (--kv-unified: the mask-prefix hint is off). test-backend-ops perf against the previous commit's library (the served attention; its graph repeats the op, so as in a model only the first copy scans), 4 rounds, the order rotated, medians: a token: 16,384 -13.2 % 65,536 -3.6 % 131,072 -2.0 % 245,760 -1.1 % (-3.3, -3.6, -3.8, -3.8 us) a 4-row verify: 16,384 -13.5 % 65,536 -3.7 % 131,072 -1.2 % 245,760 -0.7 % (-4.2, -3.9, -2.4, -2.4 us) the scan's own 3.2-4.5 us. The served decode (the public pack, 4 slots with --kv-unified, a 15,616-token chat prompt, 256 greedy tokens, a fresh server a leg, 3 legs an arm on a host loaded by a render and a Rust build): the best legs plain 101.50 -> 101.16 tok/s and drafted 205.37 -> 211.80 against the parent of the previous commit, inside that host's noise (15 scans of about 4 us in a 10 ms token). Checks: FLASH_ATTN_EXT 3,220/3,220 on an RTX 5080; the served decode's tokens identical in every leg of the three libraries, plain and drafted (every attention reads the live steps the scan it replaces would have written). GGML_CUDA_FATTN_LIVE_SCAN_EACH_LEGACY=1: every layer scans into the pool, as before.
…put tiles, the heads' blocks aligned The live-tile split (KV_live) launched every block the SMs hold, where the plain stream-k split rounds down to a multiple of the output tiles. 01f4f0f put a decode token's tile at 3 blocks an SM, so an RTX 5090 (170 SMs) splits Ternary Bonsai 2 27B's 4 KV heads 510 ways, 127.5 a head, and an RTX 5070 Ti (70 SMs) 210 ways, 52.5 a head: head h's blocks start half a block's steps from head h-1's. A cell's K and V rows hold the 4 heads side by side, 144 bytes each, so neighbouring heads share the 64-byte units the L2 fetches from DRAM: blocks reading a cell's rows together fetch a shared unit once, blocks half a range apart fetch it twice, 12 units a cell row where 9 hold it. The block count is now rounded down to a multiple of the Q tile's output tiles (its KV heads x GQA groups), at most ntiles_per_q_tile - 1 blocks fewer (510 -> 508, 210 -> 208; an RTX 5080's 252 is a multiple already). ncu on an RTX 5070 Ti, a decode token's attention at 131,072 cells (K and V 151.0 MB): DRAM bytes read 202-207 MB with 210 blocks, 152-153 MB with 208 (3 launches each); L2 read hits 33.9 -> 46.1 %; at 245,760 cells (283.1 MB) the L2's misses 377.3 -> 282.9 MB. The split read the cache 4/3 times; aligned, once. test-backend-ops perf, the served attention (q4_0 K/V read raw, a bit mask, the cache's layout), the same library with and without the switch, medians, a token: RTX 5090, 6 rounds: 16,384 -15.7 % 65,536 -7.8 % 131,072 -17.0 % 245,760 -20.9 % (19.60 / 37.84 / 115.72 / 196.44 us) RTX 5070 Ti, 4 rounds: 16,384 -2.1 % 65,536 -18.3 % 131,072 -21.4 % 245,760 -23.1 % a 4-row verify runs 2 blocks an SM (340 and 140 blocks, multiples of 4 already): within 0.2 % on the 5090, 1.7 % on the 5070 Ti. On the 5090 against 4104c47's published library, before 01f4f0f put the token at 3 blocks an SM (340 blocks): 131,072 123.69 -> 115.72 us and 245,760 204.66 -> 196.44, where 01f4f0f's split had them at 145.50 and 255.48. Checks: FLASH_ATTN_EXT 3,220/3,220 on the RTX 5070 Ti (its token split moved 210 -> 208) and on the RTX 5080; on the RTX 5090 its first 1,015 cases passed before the run was stopped (one host thread bounds it there at 45 minutes). GGML_CUDA_FATTN_LIVE_BLOCKS_ALIGN_LEGACY=1: every block the SMs hold, as before.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
The train engine-10 past #71: 6 commits, main (95d2326) to 67291c8. rig pins 32e695e, engine-32e695e.
n_probs(or any OpenAIlogprobs) turned the prefilter off, so every token copied the row's 248,320 logits to the host for a CPU softmax.llama_sampler_init_row_probs()takes the row's softmax on the backend. RTX 5080 at -c 32768, n_probs 5: plain 95.65 → 105.14 tok/s, drafted 166.20 → 192.51, the same tokens and top-5 ids.LLAMA_TOP_K_PREFILTER_LEGACY=1keeps every logit on the CPU.GGML_CUDA_FATTN_LIVE_SCAN_EACH_LEGACY=1scans every layer.GGML_CUDA_FATTN_LIVE_BLOCKS_ALIGN_LEGACY=1keeps every block.FLASH_ATTN_EXT 3,220/3,220 on an RTX 5080 and an RTX 5070 Ti.
Gate on 00d8d77's prebuilt (RTX 5090, against engine-4b61c54, the 245K prompt's four questions greedy with top-5 log-probabilities, GATE PASS):
n_probs: the same tokens and drafting;Gate on 32e695e's prebuilt (RTX 5090, against engine-4b61c54, the same four questions, GATE PASS). The aligned split sums in another order, so the pairs are held to engine-10's G1 bars:
n_probs: the same tokens and drafting;One-copy and scan-once are bit-identical to their parents on an RTX 5080 (15,616-token prompt, greedy with top-5 log-probabilities, one slot and four).
e2e-driver-only.shwith the prebuilt on a 5070 Ti: PASS.