diff --git a/TORAD.md b/TORAD.md index 9b643bb79dc6..85d8f9ca5261 100644 --- a/TORAD.md +++ b/TORAD.md @@ -91,6 +91,8 @@ which pins a commit of this branch as a submodule. | `b2eb4336a` | `llama-bench` had no parameter axis for the recurrent state's cache type, so a bench of a Gated DeltaNet model always ran the context's default `type_s` (f32) instead of what the server serves — Ternary Bonsai 2 27B serves an f16 state, and a bench at f32 moved twice its state bytes (48 layers x 48 x 128 x 128). `-cts`, matching `-ctk`/`-ctv`, is now a field in every printer | (new option) | | `16ad036c4` | between two PQ2_0 matmuls a decode token runs a chain of small kernels (norm, FWHT, q8_1 quantize, the Gated DeltaNet layer's conv/recurrence/gated norm), 1-3 us each, and PDL lets the next matmul request its own weights only under the one kernel before it — on an RTX 5080 at depth 16,384 that chain is ~1.3 ms of a 10 ms token against ~50 us floors. A PQ2_0 launch's blocks now take every gridDim.x-th tile from their index (so what a launch reads first is its matrices' heads), and each block of the launch before it, once its boxes have landed, prefetches its share of those heads into L2 (`cp.async.bulk.prefetch.L2`), sized from `GGML_CUDA_PQ2_PREFETCH_US` (default 2, at most a quarter of the card's L2; 0 disables it) times the device's DRAM rate. Ternary Bonsai 2 27B, llama-bench tg64, served state, graphs on, medians of 20 reps: RTX 5080 +2.4 % at depth 16,384 (100.54 -> 102.97 tok/s), +2.7 % at 0 (103.75 -> 106.55); RTX 5070 Ti +1.6 % (78.98 -> 80.21), +1.1 % (79.83 -> 80.71); RTX 5090, against engine-02512a3's `libggml-cuda` under this build's binaries, +4.84 % at 16,384 (152.93 -> 160.34) and +4.51 % at 0 (157.57 -> 164.68); greedy 128 tokens byte-identical to the parent at 2 and 8 us on the 5080 and at the default on the 5070 Ti. `test-backend-ops`: the 296 PQ2_0 cases and 24 MUL_MAT_GROUP cases pass with the prefetch off, at 8 us and at 64 us | `GGML_CUDA_PQ2_PREFETCH_US=0` | | `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) | 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-backend-meta.cpp b/ggml/src/ggml-backend-meta.cpp index fe58ea3bb7a6..eafb6adb1e0f 100644 --- a/ggml/src/ggml-backend-meta.cpp +++ b/ggml/src/ggml-backend-meta.cpp @@ -1083,7 +1083,35 @@ static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state( split_state.ne[j] *= tensor->ne[split_state.axis]; if (split_state.ne[j] != 0 || tensor->src[i]->ne[src_ss[i].axis] != 0) { const int64_t div = tensor->src[i]->ne[src_ss[i].axis] * split_state.nr[0]; - GGML_ASSERT(split_state.ne[j] % div == 0); + if (split_state.ne[j] % div != 0) { + char segbuf[256]; + int segpos = 0; + for (size_t s = 0; s < src_ss[i].n_segments && segpos < 200; s++) { + segpos += snprintf(segbuf + segpos, sizeof(segbuf) - segpos, "%s[%lld]*%u", + s > 0 ? "+" : "", (long long) src_ss[i].ne[s*n_bufs + j], + src_ss[i].nr[s]); + } + char chainbuf[512]; + int chainpos = 0; + // names/ops/dims only: calling back into the split computer here + // recursed without bound (its cache is being filled by this walk) + const ggml_tensor * anc = tensor; + for (int depth = 0; depth < 6 && anc != nullptr && chainpos < 430; depth++) { + chainpos += snprintf(chainbuf + chainpos, sizeof(chainbuf) - chainpos, + "%s%s:%s[%lld]", depth > 0 ? " <- " : "", + anc->name, ggml_op_name(anc->op), (long long) anc->ne[0]); + anc = anc->src[0] == anc ? nullptr : anc->src[0]; + } + GGML_ABORT("tensor split ratio is not integral: node=%s op=%s dst=[%lld,%lld,%lld,%lld] dstaxis=%d src%zu=[%lld,%lld,%lld,%lld] axis=%d segs=%s num=%lld den=%lld chain=%s", + tensor->name, ggml_op_name(tensor->op), + (long long) tensor->ne[0], (long long) tensor->ne[1], + (long long) tensor->ne[2], (long long) tensor->ne[3], + split_state.axis, i, + (long long) tensor->src[i]->ne[0], (long long) tensor->src[i]->ne[1], + (long long) tensor->src[i]->ne[2], (long long) tensor->src[i]->ne[3], + src_ss[i].axis, segbuf, (long long) split_state.ne[j], (long long) div, + chainbuf); + } split_state.ne[j] /= div; } } diff --git a/ggml/src/ggml-cuda/cp-async.cuh b/ggml/src/ggml-cuda/cp-async.cuh index 63d0c482ff72..eb668caa377a 100644 --- a/ggml/src/ggml-cuda/cp-async.cuh +++ b/ggml/src/ggml-cuda/cp-async.cuh @@ -55,3 +55,23 @@ static __device__ __forceinline__ void cp_async_wait_all() { NO_DEVICE_CODE; #endif // CP_ASYNC_AVAILABLE } + +// Closes a group of this thread's asynchronous copies: the ones issued since the last group. +static __device__ __forceinline__ void cp_async_commit_group() { +#ifdef CP_ASYNC_AVAILABLE + asm volatile("cp.async.commit_group;"); +#else + NO_DEVICE_CODE; +#endif // CP_ASYNC_AVAILABLE +} + +// Makes each thread wait until at most its n newest groups are still in flight, every older one done. +// As cp_async_wait_all, no synchronization beyond the thread. +template +static __device__ __forceinline__ void cp_async_wait_group() { +#ifdef CP_ASYNC_AVAILABLE + asm volatile("cp.async.wait_group %0;" : : "n"(n)); +#else + NO_DEVICE_CODE; +#endif // CP_ASYNC_AVAILABLE +} diff --git a/ggml/src/ggml-cuda/fattn-common.cuh b/ggml/src/ggml-cuda/fattn-common.cuh index 72c5f64a011f..cf4276fc02b5 100644 --- a/ggml/src/ggml-cuda/fattn-common.cuh +++ b/ggml/src/ggml-cuda/fattn-common.cuh @@ -1071,6 +1071,75 @@ static __global__ void flash_attn_mask_to_KV_live( } } +// A stream-k tile's fold: the partial results of blocks b_last - 1 down to b_first, onto the one of block b_last (dst_val, +// max_val, rowsum), the order and the formulas every fixup kernel used, so the bits are theirs. data: the blocks' partial +// results (block b's column jc at b*ncols*D + jc*D); meta: their KQ max and rowsum (at b*ncols + jc); skip(b): the block did +// no work on this tile. Each of the next FATTN_FIXUP_CHUNK blocks' loads goes out before the current ones fold, so a fold +// waits on its fmaxf and its FMAs, not on a load per block: with one load per block in turn, a decode token's fixup at +// 65,536 cells took 6.9 us (ncu, RTX 5080, 42 blocks a tile), folded 8 ahead 5.9 us at 63. +#define FATTN_FIXUP_CHUNK 8 + +template +static __device__ __forceinline__ void flash_attn_fixup_fold( + float & dst_val, float & max_val, float & rowsum, const float * __restrict__ data, const float2 * __restrict__ meta, + const int jc, const int tid, const int b_first, const int b_last, const skip_t & skip) { + constexpr int C = FATTN_FIXUP_CHUNK; + + float add[2][C]; + float2 mk[2][C]; + const auto load = [&](const int buf, const int b0) { +#pragma unroll + for (int k = 0; k < C; ++k) { + const int b = b0 - k; + if (b >= b_first) { + add[buf][k] = data[b*ncols*D + jc*D + tid]; + mk[buf][k] = meta[b*ncols + jc]; + } + } + }; + const auto fold = [&](const int buf, const int b0) { +#pragma unroll + for (int k = 0; k < C; ++k) { + const int b = b0 - k; + if (b < b_first || skip(b)) { + continue; + } + const float2 tmp = mk[buf][k]; + + const float max_val_new = fmaxf(max_val, tmp.x); + + const float diff_val = max_val - max_val_new; + const float diff_add = tmp.x - max_val_new; + + const float scale_val = diff_val >= SOFTMAX_FTZ_THRESHOLD ? expf(diff_val) : 0.0f; + const float scale_add = diff_add >= SOFTMAX_FTZ_THRESHOLD ? expf(diff_add) : 0.0f; + + dst_val = scale_val*dst_val + scale_add*add[buf][k]; + rowsum = scale_val*rowsum + scale_add*tmp.y; + + max_val = max_val_new; + } + }; + + int b0 = b_last - 1; + load(0, b0); + while (b0 >= b_first) { + if (b0 - C >= b_first) { + load(1, b0 - C); + } + fold(0, b0); + b0 -= C; + if (b0 < b_first) { + break; + } + if (b0 - C >= b_first) { + load(0, b0 - C); + } + fold(1, b0); + b0 -= C; + } +} + template // D == head size __launch_bounds__(D, 1) static __global__ void flash_attn_stream_k_fixup_uniform( @@ -1130,24 +1199,8 @@ static __global__ void flash_attn_stream_k_fixup_uniform( } // Combine with all previous blocks in this tile. - for (int bidx = b_last - 1; bidx >= b_first; --bidx) { - const float dst_add = dst_fixup_data[bidx*ncols*D + jc*D + tid]; - - const float2 tmp = dst_fixup[(nblocks_stream_k + bidx)*ncols + jc]; - - const float max_val_new = fmaxf(max_val, tmp.x); - - const float diff_val = max_val - max_val_new; - const float diff_add = tmp.x - max_val_new; - - const float scale_val = diff_val >= SOFTMAX_FTZ_THRESHOLD ? expf(diff_val) : 0.0f; - const float scale_add = diff_add >= SOFTMAX_FTZ_THRESHOLD ? expf(diff_add) : 0.0f; - - dst_val = scale_val*dst_val + scale_add*dst_add; - rowsum = scale_val*rowsum + scale_add*tmp.y; - - max_val = max_val_new; - } + flash_attn_fixup_fold(dst_val, max_val, rowsum, dst_fixup_data, dst_fixup + nblocks_stream_k*ncols, jc, tid, + b_first, b_last, [](const int) { return false; }); // Write back final result: *dst = dst_val / rowsum; @@ -1329,28 +1382,10 @@ static __global__ void flash_attn_stream_k_fixup_live( // A block can have no unit of work only when there are fewer units than blocks; the test is two 64-bit divisions ahead of // every load, so it runs only then. const bool maybe_empty = fixup_legacy || total_work < nblocks; - for (int bidx = b_last - 1; bidx >= b_first; --bidx) { - if (maybe_empty && int64_t(bidx)*total_work / nblocks == int64_t(bidx + 1)*total_work / nblocks) { - continue; // Did not have any data. - } - - const float dst_add = dst_fixup_data[bidx*ncols*D + jc*D + tid]; - - const float2 tmp = dst_fixup[(nblocks + bidx)*ncols + jc]; - - const float max_val_new = fmaxf(max_val, tmp.x); - - const float diff_val = max_val - max_val_new; - const float diff_add = tmp.x - max_val_new; - - const float scale_val = diff_val >= SOFTMAX_FTZ_THRESHOLD ? expf(diff_val) : 0.0f; - const float scale_add = diff_add >= SOFTMAX_FTZ_THRESHOLD ? expf(diff_add) : 0.0f; - - dst_val = scale_val*dst_val + scale_add*dst_add; - rowsum = scale_val*rowsum + scale_add*tmp.y; - - max_val = max_val_new; - } + flash_attn_fixup_fold(dst_val, max_val, rowsum, dst_fixup_data, dst_fixup + nblocks*ncols, jc, tid, + b_first, b_last, [&](const int b) { // did not have any data + return maybe_empty && int64_t(b)*total_work / nblocks == int64_t(b + 1)*total_work / nblocks; + }); *dst = dst_val / rowsum; } diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 9dd8bff5dbb0..fc6a8bcaf3c9 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -580,6 +580,14 @@ static constexpr __host__ __device__ int fattn_raw_row_bytes() { return (D/QK8_0)*fattn_raw_block_bytes(); } +// The K tile's shared memory in half2, the multi-stage pipeline's: f16 K, nbatch_fa rows of stride_tile_K; raw K, none, K*Q +// reads the raw tile (flash_attn_ext_raw_KQ). An f16-sized region for raw K was 16 KB of the 43 KB that held the head-256 +// decode tile at 2 blocks per SM; at 26 KB it runs 3. +static constexpr __host__ __device__ int fattn_mma_tile_K_h2( + const ggml_type type_K, const int nbatch_fa, const int stride_tile_K) { + return type_K == GGML_TYPE_F16 ? nbatch_fa*stride_tile_K : 0; +} + template static __device__ __forceinline__ void flash_attn_ext_raw_load( const char * const __restrict__ KV, char * const __restrict__ tile_raw, const int stride_KV) { @@ -607,6 +615,10 @@ static __device__ __forceinline__ void flash_attn_ext_raw_load( // Each thread converts 2 blocks (36 bytes for q4_0, 68 for q8_0, 4-byte aligned). The result is bit-identical to // dequantize_block_q4_0 and dequantize_block_q8_0: the integer values are exact in f16 and one __hmul2 rounds the // product with the block scale once, as the f32 product converted to f16 does. +// 8 consecutive threads take the same block pair of 8 consecutive rows. A 16-byte store is served 8 threads at a time, +// and a row stride of an odd number of 16 bytes puts 8 rows in 8 distinct bank quads; 8 threads on 2 rows hit 2, a +// 4-way conflict that was 3.2M of the 4.4M shared store wavefronts (ncu, RTX 5080, head 256 at 65,536 cells). A warp +// still loads 8 rows x 4 block pairs, the same conflict-free set of raw words as before. template static __device__ __forceinline__ void flash_attn_ext_raw_dequant_tile( const char * const __restrict__ tile_raw, half2 * const __restrict__ tile) { @@ -616,6 +628,8 @@ static __device__ __forceinline__ void flash_attn_ext_raw_dequant_tile( constexpr int npairs = nbatch_fa*pairs_per_row; constexpr int pair_words = 2*fattn_raw_block_bytes()/sizeof(int); static_assert(D % (2*QK8_0) == 0, "bad D"); + static_assert(nbatch_fa % 8 == 0, "bad nbatch_fa"); + static_assert(stride_tile*sizeof(half2) % 32 == 16, "the row stride must be an odd number of 16 bytes"); #pragma unroll for (int p0 = 0; p0 < npairs; p0 += nwarps*warp_size) { @@ -625,8 +639,8 @@ static __device__ __forceinline__ void flash_attn_ext_raw_dequant_tile( break; } - const int i = p / pairs_per_row; - const int bp = p % pairs_per_row; + const int i = p % 8 + 8*(p / (8*pairs_per_row)); + const int bp = (p / 8) % pairs_per_row; const int * src = (const int *) (tile_raw + i*fattn_raw_row_bytes() + bp*(2*fattn_raw_block_bytes())); int w[pair_words]; @@ -771,27 +785,6 @@ static __device__ __forceinline__ void flash_attn_ext_raw_quantize_Q( } } -// The scale of every block of a raw tile as float: tile_d[b*nbatch_fa + i] is block b of row i. -template -static __device__ __forceinline__ void flash_attn_ext_raw_load_d( - const char * const __restrict__ tile_raw, float * const __restrict__ tile_d) { - constexpr int warp_size = ggml_cuda_get_physical_warp_size(); - constexpr int nd = nbatch_fa*(D/QK4_0); - -#pragma unroll - for (int i0 = 0; i0 < nd; i0 += nwarps*warp_size) { - const int i = i0 + threadIdx.y*warp_size + threadIdx.x; - - if (i0 + nwarps*warp_size > nd && i >= nd) { - break; - } - - const int row = i % nbatch_fa; - const int b = i / nbatch_fa; - tile_d[i] = __half2float(*((const half *) (tile_raw + row*fattn_raw_row_bytes() + b*fattn_raw_block_bytes()))); - } -} - // KQ for the warp's 16 Q columns and its K rows of a raw tile: one int8 m16n8k32 MMA per 32 value block. The result is // the float 12582912 + dot (C input 0x4B400000), exact for |dot| < 2^22 (here <= 32*127*128), so one FADD recovers the // integer dot product of the block and one FFMA applies the block's scale in f32: the arithmetic of @@ -800,10 +793,12 @@ static __device__ __forceinline__ void flash_attn_ext_raw_load_d( // exactly the B fragment of thread (g, t) for K row g. The nibbles enter unsigned (s8 x u8); the C input // 0x4B400000 - 8*sum(q) removes the q4_0 offset of 8. // q8_0: qs words t and 4+t of a block hold values 4t..4t+3 and 16+4t..16+4t+3, the B fragment as stored (s8 x s8). +// The block scales of the C fragment's K rows 2t and 2t+1 are read as f16 from the raw rows: 4 addresses a warp, in 4 +// distinct banks. Converted to a float tile first, they took a pass and a barrier of their own each step, its loads 4-way +// bank conflicts. template static __device__ __forceinline__ void flash_attn_ext_raw_KQ( - const char * const __restrict__ tile_raw, const float * const __restrict__ tile_d, - const fattn_raw_Q8 & Q8, T_C_KQ * const __restrict__ KQ_C) { + const char * const __restrict__ tile_raw, const fattn_raw_Q8 & Q8, T_C_KQ * const __restrict__ KQ_C) { #ifdef AMPERE_MMA_AVAILABLE static_assert(type_K == GGML_TYPE_Q4_0 || type_K == GGML_TYPE_Q8_0, "no raw K*Q of this type"); static_assert(T_C_KQ::I == 16 && T_C_KQ::J == 16, "bad KQ tile"); @@ -818,8 +813,8 @@ static __device__ __forceinline__ void flash_attn_ext_raw_KQ( const int i0 = i00 + (threadIdx.y % np)*T_C_KQ::J; #pragma unroll for (int h = 0; h < 2; ++h) { - const int * src = (const int *) (tile_raw + (i0 + 8*h + g)*fattn_raw_row_bytes()); - const float * d = tile_d + i0 + 8*h + 2*t; + const int * src = (const int *) (tile_raw + (i0 + 8*h + g)*fattn_raw_row_bytes()); + const char * d = tile_raw + (i0 + 8*h + 2*t)*fattn_raw_row_bytes(); float acc[4] = {0.0f, 0.0f, 0.0f, 0.0f}; #pragma unroll for (int bp = 0; bp < D/(2*QK8_0); ++bp) { @@ -846,7 +841,9 @@ static __device__ __forceinline__ void flash_attn_ext_raw_KQ( : "r"(Q8.a[b][0]), "r"(Q8.a[b][1]), "r"(Q8.a[b][2]), "r"(Q8.a[b][3]), "r"(lo), "r"(hi), "r"(0x4B400000)); } - const float2 dk = *((const float2 *) (d + b*nbatch_fa)); + const float2 dk = make_float2( + __half2float(*((const half *) (d + b*fattn_raw_block_bytes()))), + __half2float(*((const half *) (d + fattn_raw_row_bytes() + b*fattn_raw_block_bytes())))); acc[0] = fmaf(__int_as_float(dot[0]) - bias, dk.x, acc[0]); acc[1] = fmaf(__int_as_float(dot[1]) - bias, dk.y, acc[1]); acc[2] = fmaf(__int_as_float(dot[2]) - bias, dk.x, acc[2]); @@ -861,7 +858,7 @@ static __device__ __forceinline__ void flash_attn_ext_raw_KQ( } } #else - GGML_UNUSED_VARS(tile_raw, tile_d, Q8, KQ_C); + GGML_UNUSED_VARS(tile_raw, Q8, KQ_C); NO_DEVICE_CODE; #endif // AMPERE_MMA_AVAILABLE } @@ -913,6 +910,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr bool raw_KV = type_K != GGML_TYPE_F16; static_assert(raw_KV == (type_V != GGML_TYPE_F16), "K and V are both raw or both f16"); static_assert(!raw_KV || cols_per_warp == 16, "raw K*Q needs 16 Q columns per warp"); + // The next raw K tile loads as soon as K*Q has read this one, in flight beside this step's V through the softmax, + // where 2 blocks share an SM: the 4-warp tiles of a verify, whose K then loaded only after V arrived, 1-3 % slower at + // 65,536 cells. With 3 blocks an SM (a decode token's 2-warp tile) K loads after V arrives: loaded early it ran 1-2 % + // slower at 16,384 and 245,760 cells (RTX 5080, head 256). + constexpr bool raw_K_early = raw_KV && ggml_cuda_fattn_mma_get_occupancy(DKQ, DV, ncols) <= 2; constexpr int stride_tile_K = nbatch_K2 + 4; @@ -935,12 +937,12 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( cp_async_wait_all(); __syncthreads(); if constexpr (raw_KV) { - // The raw K tile has arrived: start loading the raw V tile. K*Q reads the K values raw - // (flash_attn_ext_raw_KQ) and needs only the block scales as float, kept where tile_K would be. + // The raw K tile has arrived: start loading the raw V tile. K*Q reads K raw (flash_attn_ext_raw_KQ). flash_attn_ext_raw_load ((const char *) V_h2 + int64_t(k_VKQ_0)*stride_V, tile_raw + nbatch_fa*fattn_raw_row_bytes(), stride_V); - flash_attn_ext_raw_load_d(tile_raw, (float *) tile_K); - __syncthreads(); + if constexpr (raw_K_early) { + cp_async_commit_group(); // the raw V tile, which the wait before the VKQ tile waits for alone + } } else { flash_attn_ext_f16_load_tile (V_h2 + int64_t(k_VKQ_0)*stride_V, tile_V, nbatch_V2, stride_V, k_VKQ_sup); @@ -973,7 +975,13 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( // Calculate tile of KQ: if constexpr (raw_KV) { - flash_attn_ext_raw_KQ(tile_raw, (const float *) tile_K, Q8, KQ_C); + flash_attn_ext_raw_KQ(tile_raw, Q8, KQ_C); + if constexpr (raw_K_early && !last_iter) { + __syncthreads(); // every warp has read the raw K tile + flash_attn_ext_raw_load + ((const char *) K_h2 + int64_t(kb0_next)*nbatch_fa*stride_K, tile_raw, stride_K); + cp_async_commit_group(); + } } else if constexpr (Q_in_reg) { #pragma unroll for (int i_KQ_00 = 0; i_KQ_00 < nbatch_fa; i_KQ_00 += np*T_A_KQ::I) { @@ -1292,9 +1300,13 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( if constexpr (nstages > 1) { static_assert(!V_is_K_view, "K data reuse not implemented multi-stage loading"); - // Preload K tile for next iteration: + // Preload K tile for next iteration (raw_K_early: loading since K*Q, the wait is for this step's V alone): constexpr bool use_cp_async = true; - cp_async_wait_all(); + if constexpr (raw_K_early) { + cp_async_wait_group(); + } else { + cp_async_wait_all(); + } __syncthreads(); if (!last_iter) { if (ncols2 > 1 || mask_h) { @@ -1302,8 +1314,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( (mask_h, tile_mask, stride_mask, kb0_next*nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, mask_packed); } if constexpr (raw_KV) { - flash_attn_ext_raw_load - ((const char *) K_h2 + int64_t(kb0_next)*nbatch_fa*stride_K, tile_raw, stride_K); + if constexpr (!raw_K_early) { + flash_attn_ext_raw_load + ((const char *) K_h2 + int64_t(kb0_next)*nbatch_fa*stride_K, tile_raw, stride_K); + } } else { flash_attn_ext_f16_load_tile (K_h2 + int64_t(kb0_next)*nbatch_fa*stride_K, tile_K, nbatch_K2, stride_K, k_VKQ_sup); @@ -1547,7 +1561,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( extern __shared__ half2 tile_Q[]; half2 * tile_K = Q_in_reg ? tile_Q : tile_Q + ncols * stride_tile_Q; - half2 * tile_V = nstages > 1 ? tile_K + nbatch_fa * stride_tile_K : tile_K; + half2 * tile_V = nstages > 1 ? tile_K + fattn_mma_tile_K_h2(type_K, nbatch_fa, stride_tile_K) : tile_K; half * tile_mask = (half *) (nstages > 1 ? tile_V + nbatch_fa * stride_tile_V : tile_V + nbatch_fa * stride_tile_KV_max); char * tile_raw = (char *) (tile_mask + ncols1*(nbatch_fa + 8)); // raw K/V only: raw K tile, then raw V tile @@ -2040,7 +2054,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const int j_dst = jc_dst / ncols2; const int c_dst = jc_dst % ncols2; - if (!is_fixup && ((ncols1 > 1 && jt*ncols1 + j_dst >= int(ne01.z)) || (ncols2 > 1 && zt_gqa*ncols2 + c_dst >= gqa_ratio))) { + // A column past the Q rows or the GQA group is padding. Its partial is not written for the stream-k fixup + // either, which skips the same columns: 10 of the 16 of a decode token (2 rows x 8 heads for 1 x 6), + // whose blocks write 6 KB each instead of 16 as they finish. + if ((ncols1 > 1 && jt*ncols1 + j_dst >= int(ne01.z)) || (ncols2 > 1 && zt_gqa*ncols2 + c_dst >= gqa_ratio)) { continue; } @@ -2318,7 +2335,8 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml constexpr bool raw_KV = type_K != GGML_TYPE_F16; const size_t nbytes_shared_KV_1stage = nbatch_fa * std::max(nbatch_K2 + 4, nbatch_V2 + 4) * sizeof(half2); - const size_t nbytes_shared_KV_2stage = nbatch_fa * (nbatch_K2 + 4 + nbatch_V2 + 4) * sizeof(half2); + const size_t nbytes_shared_KV_2stage = (fattn_mma_tile_K_h2(type_K, nbatch_fa, nbatch_K2 + 4) + nbatch_fa*(nbatch_V2 + 4)) + * sizeof(half2); const size_t nbytes_shared_Q = ncols * (DKQ/2 + 4) * sizeof(half2); const size_t nbytes_shared_mask = ncols1 * (nbatch_fa/2 + 4) * sizeof(half2); const size_t nbytes_shared_combine = nwarps*cols_per_warp * (nbatch_combine + 4) * sizeof(half2); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index bca4ffc01a90..d9406f6abe64 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -11687,6 +11687,17 @@ static std::vector> make_test_cases_perf() { } } + // Ternary Bonsai 2 27B's attention as served (q4_0 K/V, a bit-packed mask, 4 KV heads at GQA 6, head size 256) at + // the depths it serves, for a decode token and an MTP verify at n_max 3. K and V are laid out as the KV cache's + // views are (llama_kv_cache::get_k, permuted in build_attn_mha): the 4 heads of a cell together, a head's cells + // 4*144 bytes apart + for (int64_t kv : { 16384, 65536, 131072, 245760 }) { + for (int64_t nb : { 1, 4 }) { + test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, kv, nb, true, false, 0, 0, GGML_PREC_F32, + GGML_TYPE_Q4_0, GGML_TYPE_Q4_0, {0, 2, 1, 3}, false, false, /*mask_bits=*/true)); + } + } + // SWIGLU at a 27B-class FFN width, fused [gate|up] vs split operands // note: same bytes either way, so a backend that indexes them differently shows it here for (ggml_type type : {GGML_TYPE_F16, GGML_TYPE_F32}) { diff --git a/tools/llama-bench/llama-bench.cpp b/tools/llama-bench/llama-bench.cpp index 7dc5ccdeda92..0a9453ae3a61 100644 --- a/tools/llama-bench/llama-bench.cpp +++ b/tools/llama-bench/llama-bench.cpp @@ -334,6 +334,7 @@ struct cmd_params { std::vector type_k; std::vector type_v; std::vector type_s; + std::vector n_rs_seq; std::vector n_threads; std::vector cpu_mask; std::vector cpu_strict; @@ -379,6 +380,7 @@ static const cmd_params cmd_params_defaults = { /* type_k */ { GGML_TYPE_F16 }, /* type_v */ { GGML_TYPE_F16 }, /* type_s */ { GGML_TYPE_F32 }, + /* n_rs_seq */ { 0 }, /* n_threads */ { common_cpu_get_num_math() }, /* cpu_mask */ { "0x0" }, /* cpu_strict */ { false }, @@ -451,6 +453,8 @@ static void print_usage(int /* argc */, char ** argv) { printf(" -ctk, --cache-type-k (default: %s)\n", join(transform_to_str(cmd_params_defaults.type_k, ggml_type_name), ",").c_str()); printf(" -ctv, --cache-type-v (default: %s)\n", join(transform_to_str(cmd_params_defaults.type_v, ggml_type_name), ",").c_str()); printf(" -cts, --cache-type-s recurrent state (default: %s)\n", join(transform_to_str(cmd_params_defaults.type_s, ggml_type_name), ",").c_str()); + printf(" -rs, --n-rs-seq recurrent-state snapshots a sequence keeps for a draft's rollback,\n" + " a drafting server's n_max (default: %s)\n", join(cmd_params_defaults.n_rs_seq, ",").c_str()); printf(" -t, --threads (default: %s)\n", join(cmd_params_defaults.n_threads, ",").c_str()); printf(" -C, --cpu-mask (default: %s)\n", join(cmd_params_defaults.cpu_mask, ",").c_str()); printf(" --cpu-strict <0|1> (default: %s)\n", join(cmd_params_defaults.cpu_strict, ",").c_str()); @@ -677,6 +681,13 @@ static cmd_params parse_cmd_params(int argc, char ** argv) { break; } params.type_s.insert(params.type_s.end(), types.begin(), types.end()); + } else if (arg == "-rs" || arg == "--n-rs-seq") { + if (++i >= argc) { + invalid_param = true; + break; + } + auto p = parse_int_range(argv[i]); + params.n_rs_seq.insert(params.n_rs_seq.end(), p.begin(), p.end()); } else if (arg == "-dev" || arg == "--device") { if (++i >= argc) { invalid_param = true; @@ -1154,6 +1165,9 @@ static cmd_params parse_cmd_params(int argc, char ** argv) { if (params.type_s.empty()) { params.type_s = cmd_params_defaults.type_s; } + if (params.n_rs_seq.empty()) { + params.n_rs_seq = cmd_params_defaults.n_rs_seq; + } if (params.n_gpu_layers.empty()) { params.n_gpu_layers = cmd_params_defaults.n_gpu_layers; } @@ -1225,6 +1239,7 @@ struct cmd_params_instance { ggml_type type_k; ggml_type type_v; ggml_type type_s; + int n_rs_seq; int n_threads; std::string cpu_mask; bool cpu_strict; @@ -1308,12 +1323,15 @@ struct cmd_params_instance { llama_context_params to_llama_cparams() const { llama_context_params cparams = llama_context_default_params(); - cparams.n_ctx = n_prompt + n_gen + n_depth; + // the batch is clamped to the context, and a ubatch must hold more rows than a draft's rollback keeps whole + // (split_equal's n_keep_tail, n_rs_seq + 1): pp4 at depth 0 with -rs 3 would otherwise get a 4-row ubatch + cparams.n_ctx = std::max(n_prompt + n_gen + n_depth, n_rs_seq > 0 ? n_rs_seq + 2 : 0); cparams.n_batch = n_batch; cparams.n_ubatch = n_ubatch; cparams.type_k = type_k; cparams.type_v = type_v; cparams.type_s = type_s; + cparams.n_rs_seq = n_rs_seq; cparams.offload_kqv = !no_kv_offload; cparams.flash_attn_type = flash_attn; cparams.embeddings = embeddings; @@ -1348,6 +1366,7 @@ static std::vector get_cmd_params_instances(const cmd_param for (const auto & tk : params.type_k) for (const auto & tv : params.type_v) for (const auto & tst : params.type_s) + for (const auto & nrs : params.n_rs_seq) for (const auto & nkvo : params.no_kv_offload) for (const auto & fa : params.flash_attn) for (const auto & nt : params.n_threads) @@ -1369,6 +1388,7 @@ static std::vector get_cmd_params_instances(const cmd_param /* .type_k = */ tk, /* .type_v = */ tv, /* .type_s = */ tst, + /* .n_rs_seq = */ nrs, /* .n_threads = */ nt, /* .cpu_mask = */ cm, /* .cpu_strict = */ cs, @@ -1406,6 +1426,7 @@ static std::vector get_cmd_params_instances(const cmd_param /* .type_k = */ tk, /* .type_v = */ tv, /* .type_s = */ tst, + /* .n_rs_seq = */ nrs, /* .n_threads = */ nt, /* .cpu_mask = */ cm, /* .cpu_strict = */ cs, @@ -1443,6 +1464,7 @@ static std::vector get_cmd_params_instances(const cmd_param /* .type_k = */ tk, /* .type_v = */ tv, /* .type_s = */ tst, + /* .n_rs_seq = */ nrs, /* .n_threads = */ nt, /* .cpu_mask = */ cm, /* .cpu_strict = */ cs, @@ -1489,6 +1511,7 @@ struct test { ggml_type type_k; ggml_type type_v; ggml_type type_s; + int n_rs_seq; int n_gpu_layers; int n_cpu_moe; llama_split_mode split_mode; @@ -1529,6 +1552,7 @@ struct test { type_k = inst.type_k; type_v = inst.type_v; type_s = inst.type_s; + n_rs_seq = inst.n_rs_seq; n_gpu_layers = inst.n_gpu_layers; n_cpu_moe = inst.n_cpu_moe; split_mode = inst.split_mode; @@ -1598,8 +1622,8 @@ struct test { "build_commit", "build_number", "cpu_info", "gpu_info", "backends", "model_filename", "model_type", "model_size", "model_n_params", "n_batch", "n_ubatch", "n_threads", "cpu_mask", "cpu_strict", "poll", - "type_k", "type_v", "type_s", "n_gpu_layers", "n_cpu_moe", - "split_mode", + "type_k", "type_v", "type_s", "n_rs_seq", "n_gpu_layers", + "n_cpu_moe", "split_mode", "main_gpu", "no_kv_offload", "flash_attn", "devices", "tensor_split", "tensor_buft_overrides", "load_mode", "embeddings", "no_op_offload", "no_host", "fit_target", "fit_min_ctx", @@ -1616,7 +1640,7 @@ struct test { field == "poll" || field == "model_size" || field == "model_n_params" || field == "n_gpu_layers" || field == "main_gpu" || field == "n_prompt" || field == "n_gen" || field == "n_depth" || field == "avg_ns" || field == "stddev_ns" || field == "no_op_offload" || field == "n_cpu_moe" || - field == "fit_target" || field == "fit_min_ctx" || field == "flash_attn") { + field == "fit_target" || field == "fit_min_ctx" || field == "flash_attn" || field == "n_rs_seq") { return INT; } if (field == "f16_kv" || field == "no_kv_offload" || field == "cpu_strict" || @@ -1687,6 +1711,7 @@ struct test { ggml_type_name(type_k), ggml_type_name(type_v), ggml_type_name(type_s), + std::to_string(n_rs_seq), std::to_string(n_gpu_layers), std::to_string(n_cpu_moe), split_mode_str(split_mode), @@ -1990,6 +2015,9 @@ struct markdown_printer : public printer { if (params.type_s.size() > 1 || params.type_s != cmd_params_defaults.type_s) { fields.emplace_back("type_s"); } + if (params.n_rs_seq.size() > 1 || params.n_rs_seq != cmd_params_defaults.n_rs_seq) { + fields.emplace_back("n_rs_seq"); + } if (params.main_gpu.size() > 1 || params.main_gpu != cmd_params_defaults.main_gpu) { fields.emplace_back("main_gpu"); }