From f1e45c29c882890371bdf546836cb3efe6cf7722 Mon Sep 17 00:00:00 2001 From: Marcos Damasceno Date: Sun, 27 Sep 2026 13:14:06 -0500 Subject: [PATCH 1/6] perf(server): top-n pre-sampling probabilities ride the top-k prefilter, 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. --- include/llama.h | 5 ++ src/llama-sampler.cpp | 86 ++++++++++++++++++++++ tests/test-backend-sampler.cpp | 52 +++++++++++++ tools/server/server-common.cpp | 10 ++- tools/server/server-context.cpp | 14 +++- tools/server/tests/unit/test_completion.py | 36 +++++++++ 6 files changed, 198 insertions(+), 5 deletions(-) diff --git a/include/llama.h b/include/llama.h index 7241b844cacf..2b5b7dc23558 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1418,6 +1418,11 @@ extern "C" { /// Setting k <= 0 makes this a noop LLAMA_API struct llama_sampler * llama_sampler_init_top_k (int32_t k); + /// @details The row's probabilities over every token (a softmax of its logits), changing no logit. Backend only: a + /// backend chain [row-probs, top-k] hands the host each candidate's probability under the whole row + /// (llama_get_sampled_probs_ith); on the CPU it does nothing + LLAMA_API struct llama_sampler * llama_sampler_init_row_probs(void); + /// @details Nucleus sampling described in academic paper "The Curious Case of Neural Text Degeneration" https://arxiv.org/abs/1904.09751 LLAMA_API struct llama_sampler * llama_sampler_init_top_p (float p, size_t min_keep); diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp index ac57591cb9db..e0e13f9aba2b 100644 --- a/src/llama-sampler.cpp +++ b/src/llama-sampler.cpp @@ -1541,6 +1541,13 @@ static void llama_sampler_top_k_backend_apply( data->logits = ggml_get_rows(ctx, logits_rows, top_k); ggml_set_name(data->logits, "top_k_rows"); + // probabilities a sampler before this one gave the row stay aligned with its candidates + if (data->probs) { + struct ggml_tensor * probs_rows = ggml_reshape_2d(ctx, data->probs, 1, ggml_nelements(data->probs)); + data->probs = ggml_get_rows(ctx, probs_rows, top_k); + ggml_set_name(data->probs, "top_k_probs"); + } + GGML_UNUSED(gf); } @@ -1575,6 +1582,85 @@ struct llama_sampler * llama_sampler_init_top_k(int32_t k) { ); } +// row-probs: the row's probabilities over every token (a softmax of its logits), taken on the backend and changing no +// logit. The samplers after it keep them aligned with their candidates (top-k), so a chain [row-probs, top-k] hands +// the host each of the k candidates' probability under the whole row (llama_get_sampled_probs_ith) + +struct llama_sampler_row_probs : public llama_sampler_backend { +}; + +static const char * llama_sampler_row_probs_name(const struct llama_sampler * smpl) { + auto * sctx = (llama_sampler_row_probs *) smpl->ctx; + return sctx->get_name(); +} + +static void llama_sampler_row_probs_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) { + // the probabilities are the backend's output for the caller; on the CPU no sampler after this one reads them + GGML_UNUSED(smpl); + GGML_UNUSED(cur_p); +} + +static struct llama_sampler * llama_sampler_row_probs_clone(const struct llama_sampler * smpl) { + GGML_UNUSED(smpl); + return llama_sampler_init_row_probs(); +} + +static void llama_sampler_row_probs_free(struct llama_sampler * smpl) { + delete (llama_sampler_row_probs *) smpl->ctx; +} + +static bool llama_sampler_row_probs_backend_init( + struct llama_sampler * smpl, + ggml_backend_buffer_type_t buft, + uint32_t n_outputs_max_per_seq) { + auto * sctx = (llama_sampler_row_probs *) smpl->ctx; + GGML_UNUSED(n_outputs_max_per_seq); + + const bool res = llama_sampler_backend_support(smpl, buft); + + sctx->init(res); + + return res; +} + +static void llama_sampler_row_probs_backend_apply( + struct llama_sampler * smpl, + struct ggml_context * ctx, + struct ggml_cgraph * gf, + struct llama_sampler_data * data) { + GGML_UNUSED(smpl); + GGML_UNUSED(gf); + + struct ggml_tensor * logits = ggml_reshape_1d(ctx, data->logits, ggml_nelements(data->logits)); + + data->probs = ggml_soft_max(ctx, logits); + ggml_set_name(data->probs, "row_probs"); +} + +static struct llama_sampler_i llama_sampler_row_probs_i = { + /* .name = */ llama_sampler_row_probs_name, + /* .accept = */ nullptr, + /* .apply = */ llama_sampler_row_probs_apply, + /* .reset = */ nullptr, + /* .clone = */ llama_sampler_row_probs_clone, + /* .free = */ llama_sampler_row_probs_free, + /* .backend_init = */ llama_sampler_row_probs_backend_init, + /* .backend_accept = */ nullptr, + /* .backend_apply = */ llama_sampler_row_probs_backend_apply, + /* .backend_set_input = */ nullptr, + /* .backend_reset = */ nullptr, + /* .copy_state = */ llama_sampler_backend_copy_state, +}; + +struct llama_sampler * llama_sampler_init_row_probs(void) { + return llama_sampler_init( + /* .iface = */ &llama_sampler_row_probs_i, + /* .ctx = */ new llama_sampler_row_probs { + ("row-probs"), + } + ); +} + // top-p struct llama_sampler_top_p : public llama_sampler_backend { diff --git a/tests/test-backend-sampler.cpp b/tests/test-backend-sampler.cpp index a04143c6fd22..497330596c91 100644 --- a/tests/test-backend-sampler.cpp +++ b/tests/test-backend-sampler.cpp @@ -411,6 +411,57 @@ static void test_backend_top_k_sampling(const test_params & params) { printf("backend top-k hybrid sampling test PASSED\n"); } +// [row-probs, top-k] on the backend: the k candidates are the row's k largest logits, and each one's probability is +// its softmax under the whole row (not renormalized over the k), as the CPU computes it from the full logits of the +// same prompt on a sequence without a backend sampler +static void test_backend_row_probs_top_k(const test_params & params) { + const int32_t k = 8; + llama_sampler_ptr backend_sampler_chain(llama_sampler_chain_init(llama_sampler_chain_default_params())); + llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_row_probs()); + llama_sampler_chain_add(backend_sampler_chain.get(), llama_sampler_init_top_k(k)); + std::vector backend_sampler_configs = {{ 0, backend_sampler_chain.get() }}; + + test_context test_ctx(params, backend_sampler_configs, 2); + + if (!test_ctx.decode({{0, "Hello"}, {1, "Hello"}})) { + GGML_ASSERT(false && "Failed to decode token"); + } + + const int32_t idx_backend = test_ctx.idx_for_seq(0); + const int32_t idx_cpu = test_ctx.idx_for_seq(1); + + const float * probs = llama_get_sampled_probs_ith(test_ctx.ctx.get(), idx_backend); + const llama_token * candidates = llama_get_sampled_candidates_ith(test_ctx.ctx.get(), idx_backend); + GGML_ASSERT(probs != nullptr); + GGML_ASSERT(llama_get_sampled_probs_count_ith(test_ctx.ctx.get(), idx_backend) == (uint32_t) k); + GGML_ASSERT(llama_get_sampled_candidates_count_ith(test_ctx.ctx.get(), idx_backend) == (uint32_t) k); + + const float * logits = llama_get_logits_ith(test_ctx.ctx.get(), idx_cpu); + GGML_ASSERT(logits != nullptr); + const float max_l = *std::max_element(logits, logits + test_ctx.n_vocab); + double sum = 0.0; + for (int i = 0; i < test_ctx.n_vocab; i++) { + sum += std::exp((double) logits[i] - max_l); + } + std::vector sorted(logits, logits + test_ctx.n_vocab); + std::nth_element(sorted.begin(), sorted.begin() + (k - 1), sorted.end(), std::greater()); + const float kth = sorted[k - 1]; + + double mass = 0.0; + for (int32_t i = 0; i < k; i++) { + const llama_token id = candidates[i]; + GGML_ASSERT(id >= 0 && id < test_ctx.n_vocab); + GGML_ASSERT(logits[id] >= kth); + const double p = std::exp((double) logits[id] - max_l) / sum; + printf("row-probs candidate[%d] = %d: backend %.8f, cpu %.8f\n", i, id, probs[i], p); + GGML_ASSERT(std::fabs(probs[i] - p) <= 1e-5 + 1e-4 * p); + mass += probs[i]; + } + GGML_ASSERT(mass <= 1.0 + 1e-5); + + printf("backend row-probs top-k test PASSED\n"); +} + static void test_backend_temp_sampling(const test_params & params) { { const float temp_0 = 0.8f; @@ -2134,6 +2185,7 @@ static const backend_test_case BACKEND_TESTS[] = { { "temp", test_backend_temp_sampling, true }, { "temp_ext", test_backend_temp_ext_sampling, true }, { "top_k", test_backend_top_k_sampling, true }, + { "row_probs_top_k", test_backend_row_probs_top_k, true }, { "multi_sequence", test_backend_multi_sequence_sampling, true }, { "dist", test_backend_dist_sampling, true }, { "dist_and_cpu", test_backend_dist_sampling_and_cpu, true }, diff --git a/tools/server/server-common.cpp b/tools/server/server-common.cpp index dcb258c4454b..5e6fc091b485 100644 --- a/tools/server/server-common.cpp +++ b/tools/server/server-common.cpp @@ -1503,10 +1503,14 @@ std::vector get_token_probabilities(llama_context * ctx, int i const int n_logits = llama_get_sampled_logits_count_ith(ctx, idx); + // the row's probabilities over every token, taken on the backend beside its top k (row-probs): no softmax of the k + const float * probs = llama_get_sampled_probs_ith(ctx, idx); + const bool row_probs = sampled_ids && probs && (int) llama_get_sampled_probs_count_ith(ctx, idx) == n_logits; + cur.resize(n_logits); if (sampled_ids) { for (int i = 0; i < n_logits; i++) { - cur[i] = llama_token_data{sampled_ids[i], logits[i], 0.0f}; + cur[i] = llama_token_data{sampled_ids[i], logits[i], row_probs ? probs[i] : 0.0f}; } } else { for (llama_token token_id = 0; token_id < n_logits; token_id++) { @@ -1525,6 +1529,10 @@ std::vector get_token_probabilities(llama_context * ctx, int i }); } + if (row_probs) { + return cur; + } + // apply softmax float max_l = -std::numeric_limits::infinity(); if (n_top > 0) { diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 363152467b19..0197a7786432 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -1943,14 +1943,20 @@ struct server_context_impl { // a chain drawing on the CPU from a row's k largest logits (common_sampler_backend_top_k) has the backend take // them in the graph: a decode copies k logits a row instead of n_vocab and the CPU scans none, while the same - // chain and RNG draw, so each token is the one the CPU's top k gives. Not with a reader of every logit (the - // pre-sampling probabilities, the pull detector, the lens), and taken off as the chain is when a grammar or a - // reasoning budget starts to constrain (release_backend_sampler_if_constrained) + // chain and RNG draw, so each token is the one the CPU's top k gives. The pre-sampling probabilities of at most + // k tokens ride along: row-probs gives the row's softmax over every token and the top k keeps the k's, which + // get_token_probabilities reads. Not with a reader of every logit (more pre-sampling probabilities than k, the + // pull detector, the lens), and taken off as the chain is when a grammar or a reasoning budget starts to + // constrain (release_backend_sampler_if_constrained) const int32_t top_k = common_sampler_backend_top_k(slot.smpl.get()); const bool lens_on = !params_base.lens_out.empty() && llama_n_lens_layers(ctx_tgt) > 0; - if (!slot.backend_sampler && top_k > 0 && !need_pre_sample_logits && pull_k.empty() && !lens_on && + const bool pre_probs_in_k = !need_pre_sample_logits || task.params.sampling.n_probs <= top_k; + if (!slot.backend_sampler && top_k > 0 && pre_probs_in_k && pull_k.empty() && !lens_on && common_sampler_backend_passive(slot.smpl.get())) { slot.smpl_top_k.reset(llama_sampler_chain_init(llama_sampler_chain_default_params())); + if (need_pre_sample_logits) { + llama_sampler_chain_add(slot.smpl_top_k.get(), llama_sampler_init_row_probs()); + } llama_sampler_chain_add(slot.smpl_top_k.get(), llama_sampler_init_top_k(top_k)); slot.backend_sampler = llama_set_sampler(ctx_tgt, slot.id, slot.smpl_top_k.get()); } diff --git a/tools/server/tests/unit/test_completion.py b/tools/server/tests/unit/test_completion.py index 24d9cb8fb4d7..b0daaf99d3be 100644 --- a/tools/server/tests/unit/test_completion.py +++ b/tools/server/tests/unit/test_completion.py @@ -701,6 +701,42 @@ def test_top_k_prefilter(monkeypatch, temperature, constraint, expected): assert bodies[0]["tokens"] == bodies[1]["tokens"] assert bodies[0]["completion_probabilities"] == bodies[1]["completion_probabilities"] +@pytest.mark.parametrize("temperature,n_probs", [(0.0, 5), (1.0, 5), (1.0, 40)]) +def test_top_k_prefilter_pre_sampling_probs(monkeypatch, temperature, n_probs): + """Pre-sampling probabilities of at most k tokens keep the top k on the backend (row-probs beside it): the same tokens, + and each token's top n_probs by id with their log-probabilities under the whole row, as with every logit on the CPU + (LLAMA_TOP_K_PREFILTER_LEGACY=1) up to the order of the softmax's sum.""" + global server + bodies = [] + for legacy in ("0", "1"): + monkeypatch.setenv("LLAMA_TOP_K_PREFILTER_LEGACY", legacy) + server = ServerPreset.tinyllama2() + server.start() + res = server.make_request("POST", "/completion", data={ + "prompt": "Once upon a time", + "n_predict": 40, + "temperature": temperature, + "top_k": 40, + "seed": 42, + "n_probs": n_probs, + "return_tokens": True, + }) + server.stop() + assert res.status_code == 200 + bodies.append(res.body) + + assert len(bodies[0]["tokens"]) >= 8 + assert bodies[0]["tokens"] == bodies[1]["tokens"] + probs = [b["completion_probabilities"] for b in bodies] + assert len(probs[0]) == len(probs[1]) == len(bodies[0]["tokens"]) + for on, legacy in zip(*probs): + assert on["id"] == legacy["id"] + assert on["logprob"] == pytest.approx(legacy["logprob"], abs=2e-5) + assert [t["id"] for t in on["top_logprobs"]] == [t["id"] for t in legacy["top_logprobs"]] + assert len(on["top_logprobs"]) == n_probs + for a, b in zip(on["top_logprobs"], legacy["top_logprobs"]): + assert a["logprob"] == pytest.approx(b["logprob"], abs=2e-5) + @pytest.mark.parametrize("tokenize,openai_style", [(False, False), (False, True), (True, False), (True, True)]) def test_logit_bias(tokenize, openai_style): global server From 00d8d77f3a52ee022c34bda68eaf883cb3ccd25b Mon Sep 17 00:00:00 2001 From: Marcos Damasceno Date: Sun, 27 Sep 2026 13:14:17 -0500 Subject: [PATCH 2/6] docs(torad): the row-probs row (f1e45c29c) 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. --- TORAD.md | 1 + 1 file changed, 1 insertion(+) diff --git a/TORAD.md b/TORAD.md index 85d8f9ca5261..ea967d6e1010 100644 --- a/TORAD.md +++ b/TORAD.md @@ -93,6 +93,7 @@ which pins a commit of this branch as a submodule. | `4104c47d5` | `llama-bench`'s own type-name table had no f32, so `-cts f32` exited on its arguments though f32 is the state's default and a type the engine serves, and `-ctk f32 -ctv f32` did the same for a K/V pair the CUDA flash attention runs. Measured on the RTX 5080: `-cts f32,f16` runs both rows (tg16 99.26 and 99.20), `-ctk f32 -ctv f32` runs, `-cts f64` still exits with the invalid-parameter error | (fix) | | `01f4f0fda` | the raw q4_0/q8_0 flash attention lost time to its own shared memory beside its DRAM stream: at 65,536 cells on an RTX 5080 (Ternary Bonsai 2 27B's attention: head 256, 4 KV heads at GQA 6, a bit mask) 3.43M of its 7.36M shared-memory wavefronts were bank conflicts — the V tile's dequant stored 16 bytes from 8 threads on 2 rows (4-way), and the K scales' float tile loaded 32 rows of one block a warp (4-way) — and a decode token ran 2 blocks an SM at 43,168 bytes. The dequant's store phase now spans 8 rows (store conflicts 3.2M -> 0.09M); K*Q reads each block's f16 scale from the raw rows, dropping the float tile, its pass and its barrier (load conflicts 0.23M -> 0.03M, 26,272 bytes a block, 3 blocks an SM); the stream-k fixup loads the next 8 blocks' partials before folding (6.9 -> 5.9 us); padding columns write no fixup partial; and the next raw K tile loads beside V where 2 blocks share an SM (a verify's 4-warp tile). test-backend-ops perf, the served layout, 6 rounds against the parent's library: a token -16.7 / -3.7 / -3.8 / -1.8 % and a 4-row verify -6.7 / -3.6 / -2.9 / -2.8 % at 16,384 / 65,536 / 131,072 / 245,760 cells; llama-bench as served, 4 rounds while the host swapped under other builds: tg64 +1.2 % at 65,536 (91.19 -> 92.30 tok/s) and +0.5 % at 16,384, pp4 -0.5 % and -0.1 %, inside that noise (the attention's share of the step predicts +0.3 to +0.7 %). FLASH_ATTN_EXT 3,220/3,220; greedy 64 tokens after a 7K-token prompt byte-identical to the parent | none of its own (it keeps the raw path's arithmetic): `GGML_CUDA_FATTN_Q4_0_LEGACY=1` / `GGML_CUDA_FATTN_Q8_0_LEGACY=1` restore the stock kernels | | `7537f40da` | `llama-bench` left the context's `n_rs_seq` at 0, so its pp4 measured a verify that writes one snapshot of each recurrent layer's state and conv window, where a drafting server (`n_rs_seq` = the draft's n_max) writes n_max + 1. `-rs` is a parameter axis like `-cts` and a field in every printer; the bench's context holds at least `n_rs_seq` + 2 rows, since the batch is clamped to the context and `split_equal` keeps a draft's last `n_rs_seq` + 1 rows in one larger ubatch (pp4 with `-rs 3` at depth 0 aborted there). RTX 5080, Ternary Bonsai 2 27B, q4_0 K/V, f16 state, pp4 in a 512 ubatch at depth 16,384, 32 samples each: `-rs 0` 386.15, `-rs 3` 383.76 tok/s (-0.62 %) | (new option) | +| `f1e45c29c` | a request for pre-sampling probabilities (`n_probs` without `post_sampling_probs`, and every OpenAI `logprobs` request) turned the top-k prefilter off, since they are under the whole row: every decode copied each output row's n_vocab logits to the host, the CPU chain scanned them for its top k, and `get_token_probabilities` built a 248,320-entry vector and ran 248,320 `expf` a token for a softmax whose only use of the row is its sum. On an RTX 5090 at 245,752 tokens, greedy with `n_probs` 5 decoded 93.7 tok/s against 107.4 without. `llama_sampler_init_row_probs()` gives a row its softmax 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 prefilter chain is [row-probs, top-k] and `get_token_probabilities` reads the k candidates' probabilities under the whole row. RTX 5080, Ternary Bonsai 2 27B served at -c 32768, n_probs 5, medians of 9: plain 95.65 -> 105.14 tok/s (105.92 without n_probs), drafted 166.20 -> 192.51 (192.15), the same tokens and top-5 ids; log-probabilities move by up to 1.9e-3, the CPU's float running sum over the row (the backend's softmax matches a double-precision one to 4e-7). test-backend-sampler `row_probs_top_k` (Qwen3 0.6B Q4_0: the backend's k probabilities against a CPU softmax of the full row to 1e-5; without top-k's gather it fails); `test_top_k_prefilter_pre_sampling_probs` (the same tokens and top-n ids, log-probabilities within 2e-5 of every logit on the CPU, measured gap 2.9e-6; the k's own softmax fails it) | `LLAMA_TOP_K_PREFILTER_LEGACY=1` keeps every logit on the CPU; more probabilities asked than the chain's k do too | Every switch in the last column is read once per process and parses as an integer: a `*_LEGACY` switch set to `0` is the same as unset (the change stays on), and `=0` turns off `GGML_CUDA_LORA_RANK1_FUSE` and From 818965d7205445be82efd383b26f48ba45184cf3 Mon Sep 17 00:00:00 2001 From: Marcos Damasceno Date: Sun, 27 Sep 2026 14:02:31 -0500 Subject: [PATCH 3/6] perf(cuda): one copy of the mma flash attention's tile code, the stream-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. --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 52 ++++++++++++---------------- 1 file changed, 22 insertions(+), 30 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index fc6a8bcaf3c9..f726f39ace79 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -864,7 +864,7 @@ static __device__ __forceinline__ void flash_attn_ext_raw_KQ( } template static __device__ __forceinline__ void flash_attn_ext_f16_iter( const float2 * const __restrict__ Q_f2, @@ -1496,7 +1496,11 @@ template struct mma_tile_sizes { }; #endif // defined(TURING_MMA_AVAILABLE) -template +// needs_fixup: the block ends a tile it did not start; its result goes to dst unnormalized, its KQ max and rowsum to the fixup +// buffer. is_fixup: the block ends inside the tile; its result and meta go to the fixup buffer. Neither: it did the whole tile. +// They are arguments, not template parameters: one copy of the tile's code serves every block, so the few blocks that end a +// tile run instructions the rest have already brought into the SM's instruction cache. +template static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const float2 * const __restrict__ Q_f2, const half2 * const __restrict__ K_h2, @@ -1523,7 +1527,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const int kb0_start, const int kb0_stop, const uint32_t * const __restrict__ live_bits, // if set, [kb0_start, kb0_stop) index the Q tile's live steps - const int nwords) { + const int nwords, + const bool needs_fixup, + const bool is_fixup) { #if defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE) //In this kernel Q, K, V are matrices while i, j, k are matrix indices. @@ -1673,7 +1679,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int k_VKQ_sup = nbatch_fa; const int kb0_next = live_bits ? fattn_kv_live_next(live_bits, kb0) : kb0 + 1; flash_attn_ext_f16_iter - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, mask_packed, tile_Q, tile_K, tile_V, tile_mask, tile_raw, Q_B, Q8, VKQ_C, @@ -1683,7 +1689,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr bool last_iter = true; const int k_VKQ_sup = ne11 - kb0*nbatch_fa; flash_attn_ext_f16_iter - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, mask_packed, tile_Q, tile_K, tile_V, tile_mask, tile_raw, Q_B, Q8, VKQ_C, @@ -1695,7 +1701,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int k_VKQ_sup = nbatch_fa; const int kb0_next = live_bits ? fattn_kv_live_next(live_bits, kb0) : kb0 + 1; flash_attn_ext_f16_iter - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, mask_packed, tile_Q, tile_K, tile_V, tile_mask, tile_raw, Q_B, Q8, VKQ_C, @@ -1705,7 +1711,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr bool last_iter = true; constexpr int k_VKQ_sup = nbatch_fa; flash_attn_ext_f16_iter - (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, mask_packed, tile_Q, tile_K, tile_V, tile_mask, tile_raw, Q_B, Q8, VKQ_C, @@ -2098,7 +2104,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( GGML_UNUSED_VARS(Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, mask_packed, - jt, kb0_start, kb0_stop, live_bits, nwords); + jt, kb0_start, kb0_stop, live_bits, nwords, needs_fixup, is_fixup); NO_DEVICE_CODE; #endif // defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE) } @@ -2274,28 +2280,14 @@ static __global__ void flash_attn_ext_f16( kb0_stop = min(kb0_stop, KV_max[sequence*iter_j + jt] / nbatch_fa); kb0_live = max(kb0_start, min(KV_max[iter_j*ne03 + sequence*iter_j + jt] / nbatch_fa, kb0_stop - 1)); } - if (last_tile) { - constexpr bool is_fixup = true; // Last index writes its data to fixup buffer to avoid data races with other blocks. - constexpr bool needs_fixup = false; - flash_attn_ext_f16_process_tile - (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, - ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, mask_packed != 0, jt, zt_gqa, kb0_live, kb0_stop, - live_bits, live.nwords); - } else if (kb0_start == 0) { - constexpr bool is_fixup = false; // All but (potentially) the last iterations write their data to dst rather than the fixup buffer. - constexpr bool needs_fixup = false; // CUDA block is working on an entire tile. - flash_attn_ext_f16_process_tile - (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, - ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, mask_packed != 0, jt, zt_gqa, kb0_live, kb0_stop, - live_bits, live.nwords); - } else { - constexpr bool is_fixup = false; - constexpr bool needs_fixup = true; // CUDA block is missing the beginning of a tile. - flash_attn_ext_f16_process_tile - (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, - ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, mask_packed != 0, jt, zt_gqa, kb0_live, kb0_stop, - live_bits, live.nwords); - } + // A block that ends inside the tile writes to the fixup buffer, to avoid data races with the blocks after it; one that ends + // a tile it did not start writes to dst and needs the fixup; one that did the whole tile writes its final result. + const bool is_fixup = last_tile; + const bool needs_fixup = !last_tile && kb0_start != 0; + flash_attn_ext_f16_process_tile + (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, + ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, mask_packed != 0, jt, zt_gqa, kb0_live, kb0_stop, + live_bits, live.nwords, needs_fixup, is_fixup); } #else GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, KV_live_ptr, dst_ptr, dst_meta_ptr, scale, From 6aa14679c1bc2362ee32b02022e9a4c68db44587 Mon Sep 17 00:00:00 2001 From: Marcos Damasceno Date: Sun, 27 Sep 2026 14:02:56 -0500 Subject: [PATCH 4/6] perf(cuda): the flash attention's live steps scanned once a graph, not 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. --- ggml/src/ggml-cuda/common.cuh | 68 +++++++++++++++++++++++++++++ ggml/src/ggml-cuda/fattn-common.cuh | 68 ++++++++++++++++++++--------- ggml/src/ggml-cuda/ggml-cuda.cu | 5 +++ 3 files changed, 121 insertions(+), 20 deletions(-) diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index e79d5c320b2c..9d5d04c50336 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -6,6 +6,7 @@ #include #include +#include #include #include @@ -1530,6 +1531,70 @@ struct ggml_cuda_gdn_gather_context { } }; +// The flash attention's live KV steps (flash_attn_mask_to_KV_live) for the graph being evaluated: every attention layer of a +// graph reads the same mask, so the first one scans it and the others read what it found. What the scan read and how it +// split the steps is the key; a different mask, shape or stream scans again. The memory is the context's and outlives every +// graph evaluation: a CUDA graph captured with the scan writes it again on each replay, before the layers after it read it. +// A buffer too small for a later scan is kept, not freed, since a graph captured earlier may still name it. +struct ggml_cuda_fattn_kv_live_context { + struct key_t { + const void * mask; + const void * mask_data; + cudaStream_t stream; + int64_t ne[4]; // the mask's + size_t nb[4]; // the mask's + int geometry[6]; // ncols1, nbatch_fa, iter_k, the Q tiles, the output tiles per Q tile, the sequences + + bool operator==(const key_t & o) const { + return mask == o.mask && mask_data == o.mask_data && stream == o.stream && + memcmp(ne, o.ne, sizeof(ne)) == 0 && memcmp(nb, o.nb, sizeof(nb)) == 0 && + memcmp(geometry, o.geometry, sizeof(geometry)) == 0; + } + }; + + bool valid = false; // filled in this graph evaluation + key_t key = {}; + int * ptr = nullptr; + size_t size = 0; // ints + std::vector kept; // smaller buffers graphs captured before may still name + + void reset() { + valid = false; + } + + // the steps a scan of this key found in this graph evaluation, or nullptr + const int * find(const key_t & k) const { + return valid && k == key ? ptr : nullptr; + } + + // n ints for a scan of this key to fill, found by the next find() with the same key until reset() + int * fill(const key_t & k, const size_t n) { + if (n > size) { + if (ptr != nullptr) { + kept.push_back(ptr); + } + CUDA_CHECK(cudaMalloc(&ptr, n*sizeof(int))); + size = n; + } + key = k; + valid = true; + return ptr; + } + + void release() { + for (int * p : kept) { + CUDA_CHECK(cudaFree(p)); + } + kept.clear(); + if (ptr != nullptr) { + CUDA_CHECK(cudaFree(ptr)); + ptr = nullptr; + } + size = 0; + valid = false; + } +}; + // Fused conv-state update for the GDN's causal conv (build_conv_state): GET_ROWS(conv cache, s_copy) -> RESHAPE -> // CONCAT(state, transposed new inputs) -> the rollback snapshots' CPYs into the cache, and the SSM_CONV reading the // CONCAT. The graph evaluator skips the GET_ROWS, the CONCAT and the CPYs (ggml_cuda_try_ssm_conv_state_update) and @@ -1725,6 +1790,7 @@ struct ggml_backend_cuda_context { ggml_cuda_stream_context concurrent_stream_context; ggml_cuda_gdn_gather_context gdn_gather_context; ggml_cuda_ssm_conv_update_context ssm_conv_update_context; + ggml_cuda_fattn_kv_live_context fattn_kv_live_context; ggml_cuda_pq2_prefetch pq2_next; // for the node being dispatched ~ggml_backend_cuda_context(); @@ -1745,6 +1811,8 @@ struct ggml_backend_cuda_context { ggml_cuda_ssm_conv_update_context & ssm_conv_updates() { return ssm_conv_update_context; } + ggml_cuda_fattn_kv_live_context & fattn_kv_live() { return fattn_kv_live_context; } + cublasHandle_t cublas_handle() { if (cublas_handles[device][curr_stream_no] == nullptr) { ggml_cuda_set_device(device); diff --git a/ggml/src/ggml-cuda/fattn-common.cuh b/ggml/src/ggml-cuda/fattn-common.cuh index cf4276fc02b5..2ee3da1c8af9 100644 --- a/ggml/src/ggml-cuda/fattn-common.cuh +++ b/ggml/src/ggml-cuda/fattn-common.cuh @@ -1596,6 +1596,11 @@ void launch_fattn( const bool kv_live = kv_live_ok && !kv_live_legacy && !mask_prefix && stream_k && GGML_CUDA_CC_IS_NVIDIA(cc) && mask && K->ne[1] >= 4096 && K->ne[1] % FATTN_KQ_STRIDE == 0 && (nbatch_fa == 32 || nbatch_fa == 64) && (uintptr_t) mask->data % 16 == 0 && mask->nb[1] % 16 == 0 && mask->nb[3] % 16 == 0; + // The live steps depend on the mask and the split alone, and every attention layer of a graph reads the same mask: the + // first layer scans it into the context's memory and the others read what it found (ggml_cuda_fattn_kv_live_context), + // one memset and one scan a graph instead of one a layer. GGML_CUDA_FATTN_LIVE_SCAN_EACH_LEGACY=1: every layer scans. + static const bool kv_live_scan_each_legacy = ggml_env_switch("GGML_CUDA_FATTN_LIVE_SCAN_EACH_LEGACY"); + const int * KV_live_ptr = nullptr; if (kv_live) { const size_t unit = mask_packed ? sizeof(uint16_t) : sizeof(half2); const int64_t s31 = mask->nb[1] / unit; @@ -1605,25 +1610,48 @@ void launch_fattn( const int n = ntiles_x*Q->ne[3]; const int nrows = mask->ne[1]; const int ne33 = mask->ne[3]; - - KV_live.alloc(fattn_kv_live_size(n, Q->ne[3], iter_k)); - CUDA_CHECK(cudaMemsetAsync(KV_live.ptr, 0, (n + 1)*sizeof(int), main_stream)); // the step counts and the blocks done - - const dim3 blocks_num_KV_live((iter_k + FATTN_KV_LIVE_THREADS - 1)/FATTN_KV_LIVE_THREADS, ntiles_x, Q->ne[3]); - const dim3 block_dim_KV_live(FATTN_KV_LIVE_THREADS, 1, 1); - const int ntiles_per_q_tile = ntiles_z_gqa*K->ne[2]; - - ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(blocks_num_KV_live, block_dim_KV_live, 0, main_stream); - if (mask_packed && nbatch_fa == 32) { - ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, KV_live.ptr, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); - } else if (mask_packed) { - ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, KV_live.ptr, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); - } else if (nbatch_fa == 32) { - ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, KV_live.ptr, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); - } else { - ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, KV_live.ptr, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); + const int ntiles_per_q_tile = ntiles_z_gqa*K->ne[2]; + + ggml_cuda_fattn_kv_live_context::key_t key = {}; + key.mask = mask; + key.mask_data = mask->data; + key.stream = main_stream; + for (int i = 0; i < 4; ++i) { + key.ne[i] = mask->ne[i]; + key.nb[i] = mask->nb[i]; + } + const int geometry[6] = {ncols1, nbatch_fa, iter_k, ntiles_x, ntiles_per_q_tile, (int) Q->ne[3]}; + memcpy(key.geometry, geometry, sizeof(geometry)); + + ggml_cuda_fattn_kv_live_context & scans = ctx.fattn_kv_live(); + KV_live_ptr = kv_live_scan_each_legacy ? nullptr : scans.find(key); + if (KV_live_ptr == nullptr) { + const size_t size = fattn_kv_live_size(n, Q->ne[3], iter_k); + int * live; + if (kv_live_scan_each_legacy) { + KV_live.alloc(size); + live = KV_live.ptr; + } else { + live = scans.fill(key, size); + } + CUDA_CHECK(cudaMemsetAsync(live, 0, (n + 1)*sizeof(int), main_stream)); // the step counts and the blocks done + + const dim3 blocks_num_KV_live((iter_k + FATTN_KV_LIVE_THREADS - 1)/FATTN_KV_LIVE_THREADS, ntiles_x, Q->ne[3]); + const dim3 block_dim_KV_live(FATTN_KV_LIVE_THREADS, 1, 1); + + ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(blocks_num_KV_live, block_dim_KV_live, 0, main_stream); + if (mask_packed && nbatch_fa == 32) { + ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, live, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); + } else if (mask_packed) { + ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, live, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); + } else if (nbatch_fa == 32) { + ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, live, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); + } else { + ggml_cuda_kernel_launch(flash_attn_mask_to_KV_live, launch_params, mask->data, live, iter_k, nrows, ne33, s31, s33, ntiles_per_q_tile); + } + CUDA_CHECK(cudaGetLastError()); + KV_live_ptr = live; } - CUDA_CHECK(cudaGetLastError()); } else if (mask && K->ne[1] % FATTN_KQ_STRIDE == 0 && (Q->ne[1] >= 1024 || Q->ne[3] > 1 || kv_range)) { const size_t unit = mask_packed ? sizeof(uint16_t) : sizeof(half2); const int64_t s31 = mask->nb[1] / unit; @@ -1758,7 +1786,7 @@ void launch_fattn( mask ? ((const char *) mask->data) : nullptr, sinks ? ((const char *) sinks->data) : nullptr, KV_max.ptr, - KV_live.ptr, + KV_live_ptr, !stream_k && parallel_blocks > 1 ? dst_tmp.ptr : (float *) KQV->data, dst_tmp_meta.ptr, scale, max_bias, m0, m1, n_head_log2, logit_softcap, Q->ne[0], ne01, Q->ne[2], Q->ne[3], Q->nb[1], Q->nb[2], Q->nb[3], @@ -1777,7 +1805,7 @@ void launch_fattn( const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(blocks_num_combine, block_dim_combine, 0, main_stream); ggml_cuda_kernel_launch(flash_attn_stream_k_fixup_live, launch_params, - (float *) KQV->data, dst_tmp_meta.ptr, KV_live.ptr, + (float *) KQV->data, dst_tmp_meta.ptr, KV_live_ptr, Q->ne[1], Q->ne[2], Q->ne[3], K->ne[2], (int)blocks_num.x, gqa_ratio, K->ne[1] / nbatch_fa, ntiles_x, ntiles_z_gqa, kv_live_fixup_legacy ? 1 : 0); } else if ((int)blocks_num.x % ntiles_dst == 0 && (int)blocks_num.x > ntiles_dst) { diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index ee7cf8ce925b..22e8a327f667 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -734,6 +734,10 @@ ggml_backend_cuda_context::~ggml_backend_cuda_context() { if (copy_event != nullptr) { CUDA_CHECK(cudaEventDestroy(copy_event)); } + if (fattn_kv_live_context.size != 0) { + ggml_cuda_set_device(device); + fattn_kv_live_context.release(); + } for (int i = 0; i < GGML_CUDA_MAX_DEVICES; ++i) { for (int j = 0; j < GGML_CUDA_MAX_STREAMS; ++j) { if (streams[i][j] != nullptr) { @@ -5224,6 +5228,7 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud cuda_ctx->gdn_gathers().reset(); cuda_ctx->ssm_conv_updates().reset(); + cuda_ctx->fattn_kv_live().reset(); const std::vector pq2_plan = ggml_cuda_pq2_prefetch_plan(*cuda_ctx, cgraph); From 32e695ecfb38171a71d2a24a236b2a1c683b6ebd Mon Sep 17 00:00:00 2001 From: Marcos Damasceno Date: Sun, 27 Sep 2026 14:49:11 -0500 Subject: [PATCH 5/6] perf(cuda): the live-tile split's blocks a multiple of a Q tile's output 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. 01f4f0fda 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 01f4f0fda 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 01f4f0fda'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. --- ggml/src/ggml-cuda/fattn-common.cuh | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-common.cuh b/ggml/src/ggml-cuda/fattn-common.cuh index 2ee3da1c8af9..a6a687ef184e 100644 --- a/ggml/src/ggml-cuda/fattn-common.cuh +++ b/ggml/src/ggml-cuda/fattn-common.cuh @@ -1709,9 +1709,20 @@ void launch_fattn( const int efficiency_loss_percent = nblocks_stream_k_rounded > 0 ? 100 * (nblocks_stream_k_raw - nblocks_stream_k_rounded) / nblocks_stream_k_raw : 100; - const int nblocks_stream_k = !kv_live && efficiency_loss_percent <= max_efficiency_loss_percent - ? nblocks_stream_k_rounded + // kv_live: the output tiles of a Q tile (its KV heads x GQA groups) share its live steps, so a block count that is a + // multiple of them gives every head the same blocks over the same steps. A cell's K and V rows hold the heads side by + // side (q4_0, head 256: 144 bytes each), so neighbouring heads share the 64-byte units the L2 fetches from DRAM, and + // only blocks reading a cell's rows together fetch a shared unit once. 510 blocks (3 an SM on an RTX 5090) split 4 + // heads 127.5 ways, the heads' blocks half a block apart: 12 units fetched a cell row where 9 hold it, 4/3 the K/V + // bytes from DRAM. At most ntiles_per_q_tile - 1 blocks go. GGML_CUDA_FATTN_LIVE_BLOCKS_ALIGN_LEGACY=1: every block the + // SMs hold. + static const bool kv_live_align_legacy = ggml_env_switch("GGML_CUDA_FATTN_LIVE_BLOCKS_ALIGN_LEGACY"); + const int ntiles_per_q_tile = ntiles_z_gqa*K->ne[2]; + const int nblocks_live = !kv_live_align_legacy && nblocks_stream_k_raw >= ntiles_per_q_tile + ? nblocks_stream_k_raw - nblocks_stream_k_raw % ntiles_per_q_tile : nblocks_stream_k_raw; + const int nblocks_stream_k = kv_live ? nblocks_live : + efficiency_loss_percent <= max_efficiency_loss_percent ? nblocks_stream_k_rounded : nblocks_stream_k_raw; blocks_num.x = nblocks_stream_k; } From 67291c8050aa3b6c564aebecda1045dc400e2a87 Mon Sep 17 00:00:00 2001 From: Marcos Damasceno Date: Sun, 27 Sep 2026 15:51:13 -0500 Subject: [PATCH 6/6] docs(torad): the rows for the one tile copy, the scan once a graph and the aligned live-tile split (818965d72, 6aa14679c, 32e695ecf) --- TORAD.md | 3 +++ 1 file changed, 3 insertions(+) diff --git a/TORAD.md b/TORAD.md index ea967d6e1010..6b70587a6417 100644 --- a/TORAD.md +++ b/TORAD.md @@ -94,6 +94,9 @@ which pins a commit of this branch as a submodule. | `01f4f0fda` | the raw q4_0/q8_0 flash attention lost time to its own shared memory beside its DRAM stream: at 65,536 cells on an RTX 5080 (Ternary Bonsai 2 27B's attention: head 256, 4 KV heads at GQA 6, a bit mask) 3.43M of its 7.36M shared-memory wavefronts were bank conflicts — the V tile's dequant stored 16 bytes from 8 threads on 2 rows (4-way), and the K scales' float tile loaded 32 rows of one block a warp (4-way) — and a decode token ran 2 blocks an SM at 43,168 bytes. The dequant's store phase now spans 8 rows (store conflicts 3.2M -> 0.09M); K*Q reads each block's f16 scale from the raw rows, dropping the float tile, its pass and its barrier (load conflicts 0.23M -> 0.03M, 26,272 bytes a block, 3 blocks an SM); the stream-k fixup loads the next 8 blocks' partials before folding (6.9 -> 5.9 us); padding columns write no fixup partial; and the next raw K tile loads beside V where 2 blocks share an SM (a verify's 4-warp tile). test-backend-ops perf, the served layout, 6 rounds against the parent's library: a token -16.7 / -3.7 / -3.8 / -1.8 % and a 4-row verify -6.7 / -3.6 / -2.9 / -2.8 % at 16,384 / 65,536 / 131,072 / 245,760 cells; llama-bench as served, 4 rounds while the host swapped under other builds: tg64 +1.2 % at 65,536 (91.19 -> 92.30 tok/s) and +0.5 % at 16,384, pp4 -0.5 % and -0.1 %, inside that noise (the attention's share of the step predicts +0.3 to +0.7 %). FLASH_ATTN_EXT 3,220/3,220; greedy 64 tokens after a 7K-token prompt byte-identical to the parent | none of its own (it keeps the raw path's arithmetic): `GGML_CUDA_FATTN_Q4_0_LEGACY=1` / `GGML_CUDA_FATTN_Q8_0_LEGACY=1` restore the stock kernels | | `7537f40da` | `llama-bench` left the context's `n_rs_seq` at 0, so its pp4 measured a verify that writes one snapshot of each recurrent layer's state and conv window, where a drafting server (`n_rs_seq` = the draft's n_max) writes n_max + 1. `-rs` is a parameter axis like `-cts` and a field in every printer; the bench's context holds at least `n_rs_seq` + 2 rows, since the batch is clamped to the context and `split_equal` keeps a draft's last `n_rs_seq` + 1 rows in one larger ubatch (pp4 with `-rs 3` at depth 0 aborted there). RTX 5080, Ternary Bonsai 2 27B, q4_0 K/V, f16 state, pp4 in a 512 ubatch at depth 16,384, 32 samples each: `-rs 0` 386.15, `-rs 3` 383.76 tok/s (-0.62 %) | (new option) | | `f1e45c29c` | a request for pre-sampling probabilities (`n_probs` without `post_sampling_probs`, and every OpenAI `logprobs` request) turned the top-k prefilter off, since they are under the whole row: every decode copied each output row's n_vocab logits to the host, the CPU chain scanned them for its top k, and `get_token_probabilities` built a 248,320-entry vector and ran 248,320 `expf` a token for a softmax whose only use of the row is its sum. On an RTX 5090 at 245,752 tokens, greedy with `n_probs` 5 decoded 93.7 tok/s against 107.4 without. `llama_sampler_init_row_probs()` gives a row its softmax 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 prefilter chain is [row-probs, top-k] and `get_token_probabilities` reads the k candidates' probabilities under the whole row. RTX 5080, Ternary Bonsai 2 27B served at -c 32768, n_probs 5, medians of 9: plain 95.65 -> 105.14 tok/s (105.92 without n_probs), drafted 166.20 -> 192.51 (192.15), the same tokens and top-5 ids; log-probabilities move by up to 1.9e-3, the CPU's float running sum over the row (the backend's softmax matches a double-precision one to 4e-7). test-backend-sampler `row_probs_top_k` (Qwen3 0.6B Q4_0: the backend's k probabilities against a CPU softmax of the full row to 1e-5; without top-k's gather it fails); `test_top_k_prefilter_pre_sampling_probs` (the same tokens and top-n ids, log-probabilities within 2e-5 of every logit on the CPU, measured gap 2.9e-6; the k's own softmax fails it) | `LLAMA_TOP_K_PREFILTER_LEGACY=1` keeps every logit on the CPU; more probabilities asked than the chain's k do too | +| `818965d72` | the mma flash attention compiled its tile code once per stream-k role (a block that ends inside a tile, one that ends a tile it did not start, one that does a whole tile) as template parameters, though only the epilogue's destination depends on the role: Ternary Bonsai 2 27B's decode kernel (head 256, q4_0 K/V read raw) was 184 KB of SASS in three copies. A per-block timeline on an RTX 5080 (252 blocks, 63 a KV head) had the 4 blocks that end a tile ending last, 8-9 us after the median block, their prologues 7.8-10.5 us against 3.7-4.2: only they ran the tile-ending copy, cold in every cache after 75-283 MB of K/V through the 64 MB L2. The roles are arguments of one copy (71 KB, 255 registers, 16 bytes of stack against 24; libggml-cuda 79.1 -> 73.2 MB). test-backend-ops perf, the served layout, 4 rounds against the parent's library: a token -2.8 / -4.2 / -2.2 / -0.9 % and a 4-row verify +0.4 / -3.1 / -1.0 / -1.1 % at 16,384 / 65,536 / 131,072 / 245,760 cells; llama-bench at 16,384 inside its noise. FLASH_ATTN_EXT 3,220/3,220; the served decode's texts identical | none: the same arithmetic and stores, only the code's layout changed | +| `6aa14679c` | the live-tile split (serving with more than one slot and `--kv-unified`) ran a memset and a scan of the mask before every flash attention, 3.2-4.5 us a call in the graph on an RTX 5080, though 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 layers). `ggml_cuda_fattn_kv_live_context` keeps a graph evaluation's first scan, keyed by the mask tensor and its data, shape and strides, the stream and the split; the other layers read it and any other key scans again; it resets where the graph evaluator resets its other per-graph records, so a CUDA graph captures one scan. test-backend-ops perf (its graph repeats the op: only the first copy scans), 4 rounds: a token -13.2 / -3.6 / -2.0 / -1.1 % and a 4-row verify -13.5 / -3.7 / -1.2 / -0.7 % at 16,384 / 65,536 / 131,072 / 245,760 cells (3.3-4.2 us a call). FLASH_ATTN_EXT 3,220/3,220; the served decode at 15,616 tokens, 4 slots: the same tokens in every leg, plain and drafted, its rates inside a loaded host's noise | `GGML_CUDA_FATTN_LIVE_SCAN_EACH_LEGACY=1` | +| `32e695ecf` | the live-tile split launched every block the SMs hold, where the plain stream-k split rounds down to a multiple of the output tiles. 01f4f0fda's 3 blocks an SM made a decode token's split 510 blocks on an RTX 5090 (127.5 per KV head) and 210 on an RTX 5070 Ti (52.5): each head's blocks started half a block's steps from its neighbour's. A cell's K and V rows hold the heads side by side (q4_0, head 256: 144 bytes each) and the L2 fetches 64-byte units, so neighbouring heads share one, fetched once only when they read the row together: 12 units a row where 9 hold it. ncu on the 5070 Ti at 131,072 cells: DRAM read 202-207 MB against 151.0 MB of K and V; rounded down to a multiple of the Q tile's output tiles (208 blocks), 152-153 MB. test-backend-ops perf, a token, the same library with and without the switch: 5090 -15.7 / -7.8 / -17.0 / -20.9 %, 5070 Ti -2.1 / -18.3 / -21.4 / -23.1 % at 16,384 / 65,536 / 131,072 / 245,760 cells; on the 5090 131,072 and 245,760 cells take 115.72 and 196.44 us, where 4104c47 (340 blocks) took 123.69 and 204.66 and 4b61c54 145.50 and 255.48. A 4-row verify (2 blocks an SM: 340 and 140 blocks) does not move; the 5080's 252 is a multiple already. The served decode on the 5090 at 245K tokens (the gate's pool, 4 x 294,912 with `--kv-unified`, its four questions, 3 legs each): plain 98.75 -> 107.92 tok/s (+9.3 %), drafted 205.43 -> 208.90 (+1.7 %), the texts first differing 35-132 tokens in: the split sums in another order. Held to engine-10's G1 bars at that shape with top-5 log-probabilities, 384 tokens: plain 4 first differences, each at a tie, |dlogprob| p99 0.095 and max 0.128; drafted 4, each at a tie, p99 0.070, max 0.078. The gate's one-slot legs take the live-tile path too: their bits move the same way and their rate does not (101.5 against 100.7 tok/s at 245K). FLASH_ATTN_EXT 3,220/3,220 on the 5070 Ti and the 5080 | `GGML_CUDA_FATTN_LIVE_BLOCKS_ALIGN_LEGACY=1` | Every switch in the last column is read once per process and parses as an integer: a `*_LEGACY` switch set to `0` is the same as unset (the change stays on), and `=0` turns off `GGML_CUDA_LORA_RANK1_FUSE` and