diff --git a/TORAD.md b/TORAD.md index 85d8f9ca5261..6b70587a6417 100644 --- a/TORAD.md +++ b/TORAD.md @@ -93,6 +93,10 @@ 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 | +| `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 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..a6a687ef184e 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; @@ -1681,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; } @@ -1758,7 +1797,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 +1816,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/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, 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); 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