From 7908e9e8ceab505c4dd3a21dfd2fd98bcffa3572 Mon Sep 17 00:00:00 2001 From: Aman Gupta Date: Sun, 27 Sep 2026 11:27:18 +0800 Subject: [PATCH] llama + CUDA: flash_attn_ext_rows (fix unified kv for multi-seq) --- ggml/include/ggml.h | 18 +++ ggml/src/ggml-cann/ggml-cann.cpp | 4 + ggml/src/ggml-cpu/ops.cpp | 36 ++++-- ggml/src/ggml-cpu/spacemit/ime.cpp | 4 + ggml/src/ggml-cuda/fattn-common.cuh | 23 +++- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 23 ++-- ggml/src/ggml-cuda/fattn-tile.cuh | 4 +- ggml/src/ggml-cuda/fattn-vec.cuh | 4 +- ggml/src/ggml-cuda/fattn.cu | 19 +++- ggml/src/ggml-et/ggml-et.cpp | 2 +- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 5 + ggml/src/ggml-metal/ggml-metal-device.m | 4 + ggml/src/ggml-opencl/ggml-opencl.cpp | 4 + ggml/src/ggml-openvino/ggml-openvino.cpp | 5 + ggml/src/ggml-sycl/ggml-sycl.cpp | 3 +- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 4 + ggml/src/ggml-webgpu/ggml-webgpu.cpp | 3 +- ggml/src/ggml.c | 42 ++++++- src/llama-context.cpp | 17 +++ src/llama-cparams.h | 2 + src/llama-graph.cpp | 137 +++++++++++++++++++---- src/llama-graph.h | 22 +++- src/llama-kv-cache.cpp | 127 +++++++++++++++++++-- src/llama-kv-cache.h | 17 ++- src/llama-kv-cells.h | 54 +++++++++ src/models/qwen4exp.cpp | 3 +- tests/test-backend-ops.cpp | 108 ++++++++++++++++++ 27 files changed, 626 insertions(+), 68 deletions(-) diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 224bdef927..6e4469a551 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -2494,6 +2494,24 @@ extern "C" { float max_bias, float logit_softcap); + // same as ggml_flash_attn_ext, but each q slice attends to its own list of K/V rows + // q: [n_embd_k, n_batch, n_head, ne3] + // k: [n_embd_k, n_rows, n_head_kv, 1 ] shared by all q slices + // v: [n_embd_v, n_rows, n_head_kv, 1 ] + // mask: [n_kv, n_batch, ne32, ne3] + // kv_rows: [n_kv, ne3] I32, mask column i of slice i3 uses row kv_rows[i, i3] of k and v + // negative rows are padding, their mask entries must be -inf + GGML_API struct ggml_tensor * ggml_flash_attn_ext_rows( + struct ggml_context * ctx, + struct ggml_tensor * q, + struct ggml_tensor * k, + struct ggml_tensor * v, + struct ggml_tensor * mask, + struct ggml_tensor * kv_rows, + float scale, + float max_bias, + float logit_softcap); + GGML_DEPRECATED(GGML_API void ggml_flash_attn_ext_set_prec( struct ggml_tensor * a, enum ggml_prec prec), diff --git a/ggml/src/ggml-cann/ggml-cann.cpp b/ggml/src/ggml-cann/ggml-cann.cpp index c2745014a1..391beca6cf 100644 --- a/ggml/src/ggml-cann/ggml-cann.cpp +++ b/ggml/src/ggml-cann/ggml-cann.cpp @@ -2660,6 +2660,10 @@ static bool ggml_backend_cann_supports_op(ggml_backend_dev_t dev, const ggml_ten return true; case GGML_OP_FLASH_ATTN_EXT: { + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + if (op->src[5] != nullptr) { + return false; + } #ifdef ASCEND_310P // FA not support on 310p device return false; diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index ba00a0a73e..1e37e9b14f 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -8624,6 +8624,7 @@ static void ggml_compute_forward_flash_attn_ext_f16_one_chunk( const ggml_tensor * v = dst->src[2]; const ggml_tensor * mask = dst->src[3]; const ggml_tensor * sinks = dst->src[4]; + const ggml_tensor * kv_rows = dst->src[5]; GGML_TENSOR_LOCALS(int64_t, neq, q, ne) GGML_TENSOR_LOCALS(size_t, nbq, q, nb) @@ -8731,6 +8732,8 @@ static void ggml_compute_forward_flash_attn_ext_f16_one_chunk( const float * pq = (const float *) ((char *) q->data + (iq1*nbq1 + iq2*nbq2 + iq3*nbq3)); q_to_vec_dot(pq, Q_q, DK); + const int32_t * rows = kv_rows ? (const int32_t *) ((const char *) kv_rows->data + iq3*kv_rows->nb[1]) : nullptr; + // online softmax / attention // loop over n_kv and n_head_kv // ref: https://arxiv.org/pdf/2112.05682.pdf @@ -8741,9 +8744,14 @@ static void ggml_compute_forward_flash_attn_ext_f16_one_chunk( continue; } + const int64_t ik1 = rows ? rows[ic] : ic; + if (ik1 < 0) { + continue; + } + float s; // KQ value - const char * k_data = (const char *) k->data + ( ic*nbk1 + ik2*nbk2 + ik3*nbk3); + const char * k_data = (const char *) k->data + (ik1*nbk1 + ik2*nbk2 + ik3*nbk3); kq_vec_dot(DK, &s, 0, k_data, 0, Q_q, 0, 1); s = s*scale; // scale KQ value @@ -8759,7 +8767,7 @@ static void ggml_compute_forward_flash_attn_ext_f16_one_chunk( float ms = 1.0f; // upon new higher max val, scale VKQ and KQ sum with this value float vs = 1.0f; // post-softmax KQ value, expf(s - M) - const char * v_data = ((const char *) v->data + (ic*nbv1 + iv2*nbv2 + iv3*nbv3)); + const char * v_data = ((const char *) v->data + (ik1*nbv1 + iv2*nbv2 + iv3*nbv3)); if (v->type == GGML_TYPE_F16) { if (s > M) { @@ -8926,6 +8934,10 @@ static void ggml_compute_forward_flash_attn_ext_tiled( static constexpr int Q_TILE_SZ = ggml_fa_tile_config::Q; static constexpr int KV_TILE_SZ = ggml_fa_tile_config::KV; + // with kv_rows the loop runs over the mask columns and reads K/V rows through kv_rows + const ggml_tensor * kv_rows = dst->src[5]; + const int64_t n_kv_cols = kv_rows ? kv_rows->ne[0] : nek1; + int ir = ir0; while (ir < ir1) { // q indices for the start of this tile @@ -8933,6 +8945,10 @@ static void ggml_compute_forward_flash_attn_ext_tiled( const int iq2 = (ir - iq3*neq2*neq1)/neq1; const int iq1 = (ir - iq3*neq2*neq1 - iq2*neq1); + // padding rows are masked, read row 0 for them + const int32_t * rows = kv_rows ? (const int32_t *) ((const char *) kv_rows->data + iq3*kv_rows->nb[1]) : nullptr; + auto kv_row = [rows](int64_t i) -> int64_t { return rows ? std::max(rows[i], 0) : i; }; + // Number of valid rows in this tile: // - limited by tile size (Q_TILE_SZ) // - limited by chunk boundary (ir1 - ir) @@ -8992,8 +9008,8 @@ static void ggml_compute_forward_flash_attn_ext_tiled( memset(K_f32, 0, DK * KV_TILE_SZ * sizeof(float)); memset(V32, 0, KV_TILE_SZ * DV * sizeof(float)); - for (int64_t ic = 0; ic < nek1; ic += KV_TILE_SZ) { - const int kv_tile = (int)std::min((int64_t)KV_TILE_SZ, nek1 - ic); + for (int64_t ic = 0; ic < n_kv_cols; ic += KV_TILE_SZ) { + const int kv_tile = (int)std::min((int64_t)KV_TILE_SZ, n_kv_cols - ic); // skip the tile entirely if all the masks are -inf if (mask) { @@ -9020,7 +9036,7 @@ static void ggml_compute_forward_flash_attn_ext_tiled( // Pack K tile transposed: K_f32[dk][kv] so KV_TILE is contiguous (SIMD dim) // Zero-pad the last tile so the GEMM always operates on KV_TILE_SZ columns for (int tk = 0; tk < kv_tile; tk++) { - const char * k_data = (const char *)k->data + (ic + tk)*nbk1 + ik2*nbk2 + ik3*nbk3; + const char * k_data = (const char *)k->data + kv_row(ic + tk)*nbk1 + ik2*nbk2 + ik3*nbk3; if (kv_type == GGML_TYPE_F16) { const ggml_fp16_t * k_f16 = (const ggml_fp16_t *)k_data; for (int64_t dk = 0; dk < DK; dk++) { @@ -9085,7 +9101,7 @@ static void ggml_compute_forward_flash_attn_ext_tiled( // V accumulation: VKQ32 += softmax(KQ) * V // Pack V tile to contiguous F32, zero-padded for (int tk = 0; tk < kv_tile; tk++) { - const char * v_data = (const char *)v->data + (ic + tk)*nbv1 + iv2*nbv2 + iv3*nbv3; + const char * v_data = (const char *)v->data + kv_row(ic + tk)*nbv1 + iv2*nbv2 + iv3*nbv3; if (kv_type == GGML_TYPE_F16) { ggml_cpu_fp16_to_fp32((const ggml_fp16_t *)v_data, V32 + tk * DV, DV); } else { @@ -9257,8 +9273,12 @@ static void ggml_compute_forward_flash_attn_ext_f16( // When use_ref is set, force the vec-only reference implementation (no tiling, no KV-chunking) const bool use_ref = params->use_ref; + // with kv_rows the loops run over the mask columns, the split KV path does not support it + const ggml_tensor * kv_rows = dst->src[5]; + const int64_t n_kv_cols = kv_rows ? kv_rows->ne[0] : nek1; + const bool kv_is_f32_or_f16 = (k->type == GGML_TYPE_F32 || k->type == GGML_TYPE_F16); - const bool use_split_kv_path = !use_ref && (neq1 == 1 && neq3 == 1) && kv_is_f32_or_f16 && (k->type == v->type) && q->type == GGML_TYPE_F32 && nek1 >= 512; + const bool use_split_kv_path = !use_ref && !kv_rows && (neq1 == 1 && neq3 == 1) && kv_is_f32_or_f16 && (k->type == v->type) && q->type == GGML_TYPE_F32 && nek1 >= 512; if (use_split_kv_path) { const int64_t chunk_size = (nek1 + nth - 1) / nth; @@ -9337,7 +9357,7 @@ static void ggml_compute_forward_flash_attn_ext_f16( if (use_tiled) { ggml_compute_forward_flash_attn_ext_tiled(params, dst, ir0, ir1); } else { - ggml_compute_forward_flash_attn_ext_f16_one_chunk(params, dst, ir0, ir1, 0, nek1, nullptr, 0); + ggml_compute_forward_flash_attn_ext_f16_one_chunk(params, dst, ir0, ir1, 0, n_kv_cols, nullptr, 0); } current_chunk = ggml_threadpool_chunk_add(params->threadpool, 1); diff --git a/ggml/src/ggml-cpu/spacemit/ime.cpp b/ggml/src/ggml-cpu/spacemit/ime.cpp index 29d683270e..3fd2c6da28 100644 --- a/ggml/src/ggml-cpu/spacemit/ime.cpp +++ b/ggml/src/ggml-cpu/spacemit/ime.cpp @@ -1035,6 +1035,10 @@ class tensor_traits_common : public tensor_traits_base { return true; } case GGML_OP_FLASH_ATTN_EXT: + // K/V row lists are handled by the generic CPU path + if (op->src[5] != nullptr) { + return false; + } forward_flash_attn_ext_f16(params, op); return true; case GGML_OP_CONT: diff --git a/ggml/src/ggml-cuda/fattn-common.cuh b/ggml/src/ggml-cuda/fattn-common.cuh index 6d1ce52dbc..e99ab6a57a 100644 --- a/ggml/src/ggml-cuda/fattn-common.cuh +++ b/ggml/src/ggml-cuda/fattn-common.cuh @@ -25,6 +25,7 @@ typedef void (* fattn_kernel_t)( const char * __restrict__ mask, const char * __restrict__ sinks, const int * __restrict__ KV_max, + const int * __restrict__ KV_rows, float * __restrict__ dst, float2 * __restrict__ dst_meta, const float scale, @@ -986,8 +987,13 @@ void launch_fattn( const bool V_is_K_view = V->view_src && (V->view_src == K || (V->view_src == K->view_src && V->view_offs == K->view_offs)); - const ggml_tensor * mask = dst->src[3]; - const ggml_tensor * sinks = dst->src[4]; + const ggml_tensor * mask = dst->src[3]; + const ggml_tensor * sinks = dst->src[4]; + const ggml_tensor * kv_rows = dst->src[5]; + + // with kv_rows the kernel iterates over the mask columns and reads K/V rows through kv_rows + const int64_t n_kv_cols = kv_rows ? kv_rows->ne[0] : K->ne[1]; + GGML_ASSERT(!kv_rows || !use_sparse); ggml_tensor * KQV = dst; @@ -1109,7 +1115,7 @@ void launch_fattn( // Optional optimization where the mask is scanned to determine whether part of the calculation can be skipped. // Only worth the overhead if there is at lease one FATTN_KQ_STRIDE x FATTN_KQ_STRIDE square to be skipped or // multiple sequences of possibly different lengths. - if (!use_sparse && mask && K->ne[1] % FATTN_KQ_STRIDE == 0 && (Q->ne[1] >= 1024 || Q->ne[3] > 1)) { + if (!use_sparse && mask && n_kv_cols % FATTN_KQ_STRIDE == 0 && (Q->ne[1] >= 1024 || Q->ne[3] > 1)) { const int64_t s31 = mask->nb[1] / sizeof(half2); const int64_t s33 = mask->nb[3] / sizeof(half2); @@ -1117,7 +1123,7 @@ void launch_fattn( const dim3 block_dim_KV_max(FATTN_KQ_STRIDE/2, 1, 1); const int ne_KV_max = blocks_num_KV_max.x*blocks_num_KV_max.y; - const int iter_k = K->ne[1] / FATTN_KQ_STRIDE; + const int iter_k = n_kv_cols / FATTN_KQ_STRIDE; KV_max.alloc(ne_KV_max); ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(blocks_num_KV_max, block_dim_KV_max, 0, main_stream); @@ -1126,13 +1132,19 @@ void launch_fattn( CUDA_CHECK(cudaGetLastError()); } + // K/V are shared by all slices when read through kv_rows + if (kv_rows) { + nb13 = 0; + nb23 = 0; + } + const dim3 block_dim(warp_size, nwarps, 1); int max_blocks_per_sm = 1; // Max. number of active blocks limited by occupancy. CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&max_blocks_per_sm, fattn_kernel, block_dim.x * block_dim.y * block_dim.z, nbytes_shared)); GGML_ASSERT(max_blocks_per_sm > 0); int parallel_blocks = max_blocks_per_sm; - const int64_t n_kv = use_sparse ? n_kv_max : K->ne[1]; + const int64_t n_kv = use_sparse ? n_kv_max : n_kv_cols; const int ntiles_KV = (n_kv + nbatch_fa - 1) / nbatch_fa; // Max. number of parallel blocks limited by KV cache length. dim3 blocks_num; @@ -1243,6 +1255,7 @@ void launch_fattn( mask ? ((const char *) mask->data) : nullptr, sinks ? ((const char *) sinks->data) : nullptr, KV_max.ptr, + kv_rows ? (const int *) kv_rows->data : nullptr, !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], diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index abb99a354a..4125c90585 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -428,6 +428,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( // padded slots gather row 0, the -inf mask removes their contribution const int32_t index = i < i_sup ? indices[k_VKQ_0 + i] : 0; i_KV = index >= 0 ? index : 0; + } else if (indices) { + // K/V rows given per mask column, negative rows are padding + i_KV = max(indices[k_VKQ_0 + i], 0); } else { i_KV = k_VKQ_0 + i; } @@ -475,8 +478,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( if constexpr (use_sparse) { const int32_t index = i < i_sup ? indices[k_VKQ_0 + i] : -1; src = index >= 0 ? KV + int64_t(index)*stride_KV + k*h2_per_chunk : zero; + } else if (!oob_check || i < i_sup) { + const int64_t i_KV = indices ? max(indices[k_VKQ_0 + i], 0) : k_VKQ_0 + i; + src = KV + i_KV*stride_KV + k*h2_per_chunk; } else { - src = !oob_check || i < i_sup ? KV + int64_t(k_VKQ_0 + i)*stride_KV + k*h2_per_chunk : zero; + src = zero; } ggml_cuda_memcpy_1<16>((char *) tile_KV + swizzle_bytes(i, k*h2_per_chunk, stride_tile), src); } @@ -642,7 +648,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( cp_async_wait_all(); __syncthreads(); flash_attn_ext_f16_load_tile - (V_h2, tile_V, nbatch_V2, stride_V, k_VKQ_0, k_VKQ_sup, nullptr); + (V_h2, tile_V, nbatch_V2, stride_V, k_VKQ_0, k_VKQ_sup, indices); } else { // the sparse mask values are gathered per element, always load them synchronously constexpr bool use_cp_async = nstages == 1 && !use_sparse; @@ -999,7 +1005,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( (mask_h, tile_mask, stride_mask, k_VKQ_0 + nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr); } flash_attn_ext_f16_load_tile - (K_h2, tile_K, nbatch_K2, stride_K, k_VKQ_0 + nbatch_fa, k_VKQ_sup, nullptr); + (K_h2, tile_K, nbatch_K2, stride_K, k_VKQ_0 + nbatch_fa, k_VKQ_sup, indices); } } @@ -1352,7 +1358,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( (mask_h, tile_mask, stride_mask, kb0*nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr); } flash_attn_ext_f16_load_tile - (K_h2, tile_K, nbatch_K2, stride_K, kb0*nbatch_fa, k_VKQ_sup, nullptr); + (K_h2, tile_K, nbatch_K2, stride_K, kb0*nbatch_fa, k_VKQ_sup, indices); } // kb0_start is always < kb0_stop so the last iter can be executed unconditionally. @@ -1811,6 +1817,7 @@ static __global__ void flash_attn_ext_f16( const char * mask_ptr, const char * sinks_ptr, const int * KV_max_ptr, + const int * KV_rows_ptr, float * dst_ptr, float2 * dst_meta_ptr, const float scale, @@ -1933,7 +1940,8 @@ static __global__ void flash_attn_ext_f16( const half2 * V_h2 = V_is_K_view ? K_h2 : (const half2 *) (V + nb23*sequence + nb22*z_KV); const float * sinks_f = sinks ? (const float *) sinks + zt_Q : nullptr; - const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*iter_j + jt)*ne11 : nullptr; + const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*iter_j + jt)*ne11 : + KV_rows_ptr ? KV_rows_ptr + int64_t(sequence)*ne11 : nullptr; const float slope = ncols2 == 1 ? get_alibi_slope(max_bias, zt_Q, n_head_log2, m0, m1) : 1.0f; @@ -1982,7 +1990,8 @@ static __global__ void flash_attn_ext_f16( const half2 * V_h2 = V_is_K_view ? K_h2 : (const half2 *) (V + nb23*sequence + nb22*z_KV); const float * sinks_f = sinks ? (const float *) sinks + zt_Q : nullptr; - const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*iter_j + jt)*ne11 : nullptr; + const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*iter_j + jt)*ne11 : + KV_rows_ptr ? KV_rows_ptr + int64_t(sequence)*ne11 : nullptr; const float slope = ncols2 == 1 ? get_alibi_slope(max_bias, zt_Q, n_head_log2, m0, m1) : 1.0f; @@ -1998,7 +2007,7 @@ static __global__ void flash_attn_ext_f16( (Q_f2, K_h2, V_h2, mask_h, indices, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, zt_gqa, kb0_start, kb0_stop); #else - GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, dst_ptr, dst_meta_ptr, scale, + GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, KV_rows_ptr, dst_ptr, dst_meta_ptr, scale, max_bias, m0, m1, n_head_log2, logit_softcap, ne00, ne01, ne02, ne03, nb01, nb02, nb03, diff --git a/ggml/src/ggml-cuda/fattn-tile.cuh b/ggml/src/ggml-cuda/fattn-tile.cuh index 8981ab804c..5994ff8568 100644 --- a/ggml/src/ggml-cuda/fattn-tile.cuh +++ b/ggml/src/ggml-cuda/fattn-tile.cuh @@ -797,6 +797,7 @@ static __global__ void flash_attn_tile( const char * mask_ptr, const char * sinks_ptr, const int * KV_max_ptr, + const int * KV_rows_ptr, float * dst_ptr, float2 * dst_meta_ptr, const float scale, @@ -821,6 +822,7 @@ static __global__ void flash_attn_tile( const int * GGML_CUDA_RESTRICT KV_max = KV_max_ptr; float * GGML_CUDA_RESTRICT dst = dst_ptr; float2 * GGML_CUDA_RESTRICT dst_meta = dst_meta_ptr; + GGML_UNUSED(KV_rows_ptr); // Skip unused kernel variants for faster compilation: @@ -1132,7 +1134,7 @@ static __global__ void flash_attn_tile( } } #else - GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, dst_ptr, dst_meta_ptr, scale, + GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, KV_rows_ptr, dst_ptr, dst_meta_ptr, scale, max_bias, m0, m1, n_head_log2, logit_softcap, ne00, ne01, ne02, ne03, nb01, nb02, nb03, diff --git a/ggml/src/ggml-cuda/fattn-vec.cuh b/ggml/src/ggml-cuda/fattn-vec.cuh index 57a2855659..8eb29ef802 100644 --- a/ggml/src/ggml-cuda/fattn-vec.cuh +++ b/ggml/src/ggml-cuda/fattn-vec.cuh @@ -25,6 +25,7 @@ static __global__ void flash_attn_ext_vec( const char * mask_ptr, const char * sinks_ptr, const int * KV_max_ptr, + const int * KV_rows_ptr, float * dst_ptr, float2 * dst_meta_ptr, const float scale, @@ -50,6 +51,7 @@ static __global__ void flash_attn_ext_vec( const int * GGML_CUDA_RESTRICT KV_max = KV_max_ptr; float * GGML_CUDA_RESTRICT dst = dst_ptr; float2 * GGML_CUDA_RESTRICT dst_meta = dst_meta_ptr; + GGML_UNUSED(KV_rows_ptr); // Skip unused kernel variants for faster compilation: if (use_logit_softcap && !(D == 128 || D == 256)) { @@ -512,7 +514,7 @@ static __global__ void flash_attn_ext_vec( dst_meta[((sequence*int(ne01.z) + ic0 + tid)*ne02 + head)*gridDim.y + blockIdx.y] = make_float2(KQ_max[tid], KQ_sum[tid]); } #else - GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, dst_ptr, dst_meta_ptr, scale, + GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, KV_rows_ptr, dst_ptr, dst_meta_ptr, scale, max_bias, m0, m1, n_head_log2, logit_softcap, ne00, ne01, ne02, ne03, nb01, nb02, nb03, diff --git a/ggml/src/ggml-cuda/fattn.cu b/ggml/src/ggml-cuda/fattn.cu index d1fcf58cb1..3afc6397d0 100644 --- a/ggml/src/ggml-cuda/fattn.cu +++ b/ggml/src/ggml-cuda/fattn.cu @@ -144,7 +144,7 @@ bool ggml_cuda_flash_attn_ext_mma_f16_shall_use_sparse(const int cc, const ggml_ // the dense kernel handles up to 64/ncols2 queries per K/V pass, the single-query gather has to beat that const int64_t n_gather = (ncols1 == 1 ? std::min(Q->ne[1], 64/ncols2) : ncols1) * (int64_t) n_kv_max; - return GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && + return GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && dst->src[5] == nullptr && mask != nullptr && n_kv_max > 0 && max_bias == 0.0f && logit_softcap == 0.0f && mask->ne[0] == K->ne[1] && mask->ne[1] >= Q->ne[1] && mask->ne[2] == 1 && K->ne[1] >= std::max(4096, 2*n_gather); @@ -199,12 +199,15 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_con const ggml_tensor * V = dst->src[2]; const ggml_tensor * mask = dst->src[3]; + // with kv_rows the kernel iterates over the mask columns + const int64_t n_kv = dst->src[5] ? dst->src[5]->ne[0] : K->ne[1]; + float max_bias = 0.0f; memcpy(&max_bias, (const float *) KQV->op_params + 1, sizeof(float)); // Edge cases like no mask, ALiBi, unpadded K/V, or misaligned addresses for large data transfers // are put into the template specialization without GQA optimizations. - bool use_gqa_opt = mask && max_bias == 0.0f && K->ne[1] % FATTN_KQ_STRIDE == 0; + bool use_gqa_opt = mask && max_bias == 0.0f && n_kv % FATTN_KQ_STRIDE == 0; for (const ggml_tensor * t : {Q, K, V, mask}) { if (t == nullptr || ggml_is_quantized(t->type)) { continue; @@ -549,6 +552,10 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const const ggml_tensor * K = dst->src[1]; const ggml_tensor * V = dst->src[2]; const ggml_tensor * mask = dst->src[3]; + const ggml_tensor * kv_rows = dst->src[5]; + + // with kv_rows the kernel iterates over the mask columns + const int64_t n_kv = kv_rows ? kv_rows->ne[0] : K->ne[1]; const int gqa_ratio = Q->ne[2] / K->ne[2]; GGML_ASSERT(Q->ne[2] % K->ne[2] == 0); @@ -558,7 +565,7 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const // The effective batch size for the kernel can be increased by gqa_ratio. // The kernel versions without this optimization are also used for ALiBi, if there is no mask, or if the KV cache is not padded, - bool gqa_opt_applies = gqa_ratio >= 2 && mask && max_bias == 0.0f && K->ne[1] % FATTN_KQ_STRIDE == 0; + bool gqa_opt_applies = gqa_ratio >= 2 && mask && max_bias == 0.0f && n_kv % FATTN_KQ_STRIDE == 0; for (const ggml_tensor * t : {Q, K, V, mask}) { if (t == nullptr || ggml_is_quantized(t->type)) { continue; @@ -630,6 +637,12 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const return BEST_FATTN_KERNEL_NONE; } + // only the MMA kernel reads K/V through kv_rows + if (kv_rows) { + return GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && mask && Q->ne[0] != 40 && Q->ne[0] != 72 ? + BEST_FATTN_KERNEL_MMA_F16 : BEST_FATTN_KERNEL_NONE; + } + // For small batch sizes the vector kernel may be preferable over the kernels optimized for large batch sizes: // 192 satisfies % 64 == 0 but has no vec instance (DKQ != DV); force it onto the MMA path. const bool can_use_vector_kernel = Q->ne[0] <= 256 && Q->ne[0] % 64 == 0 && Q->ne[0] != 192 && K->ne[1] % FATTN_KQ_STRIDE == 0; diff --git a/ggml/src/ggml-et/ggml-et.cpp b/ggml/src/ggml-et/ggml-et.cpp index 61c31d6f29..021a5fa7b9 100644 --- a/ggml/src/ggml-et/ggml-et.cpp +++ b/ggml/src/ggml-et/ggml-et.cpp @@ -1269,7 +1269,7 @@ static bool ggml_backend_et_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_OP_FLASH_ATTN_EXT: if (op->type == GGML_TYPE_F32 && op->src[0] && op->src[0]->type == GGML_TYPE_F32 && op->src[1] && (op->src[1]->type == GGML_TYPE_F32 || op->src[1]->type == GGML_TYPE_F16) && op->src[2] && - (op->src[2]->type == GGML_TYPE_F32 || op->src[2]->type == GGML_TYPE_F16) && op->src[4] == nullptr && + (op->src[2]->type == GGML_TYPE_F32 || op->src[2]->type == GGML_TYPE_F16) && op->src[4] == nullptr && op->src[5] == nullptr && ggml_is_contiguous_rows(op) && ggml_is_contiguous_rows(op->src[0])) { float max_bias = 0.0f; float logit_softcap = 0.0f; diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index 521a97a730..fa0c200aa3 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -4503,6 +4503,11 @@ static bool ggml_hexagon_precompute_flash_attn_params( } static bool ggml_hexagon_supported_flash_attn_ext(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) { + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + if (op->src[5] != nullptr) { + return false; + } + const struct ggml_tensor * src0 = op->src[0]; const struct ggml_tensor * src1 = op->src[1]; const struct ggml_tensor * src2 = op->src[2]; diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index fa58b89653..70133d3912 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1732,6 +1732,10 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te case GGML_OP_ROLL: return ggml_is_contiguous(op->src[0]); case GGML_OP_FLASH_ATTN_EXT: + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + if (op->src[5] != NULL) { + return false; + } // for new head sizes, add checks here if (op->src[0]->ne[0] != 32 && op->src[0]->ne[0] != 40 && diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 19cd02583e..cd281a3ce6 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -9182,6 +9182,10 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te case GGML_OP_MEAN: return op->src[0]->type == GGML_TYPE_F32; case GGML_OP_FLASH_ATTN_EXT: { + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + if (op->src[5] != nullptr) { + return false; + } #ifdef GGML_OPENCL_USE_ADRENO_KERNELS if (use_fa_bin_kernels_prefill(backend_ctx, op->src[0], op->src[1], op->src[2])) { return true; diff --git a/ggml/src/ggml-openvino/ggml-openvino.cpp b/ggml/src/ggml-openvino/ggml-openvino.cpp index 02c5962238..4373fb549b 100644 --- a/ggml/src/ggml-openvino/ggml-openvino.cpp +++ b/ggml/src/ggml-openvino/ggml-openvino.cpp @@ -1259,6 +1259,11 @@ static ggml_openvino_op_support is_op_supported_case(const ggml_tensor * op) { return {false, "FLASH_ATTN_EXT gemma3n pattern on GPU is not supported"}; } + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + if (op->src[5] != nullptr) { + return {false, "FLASH_ATTN_EXT with K/V row lists is not supported"}; + } + if (op->src[4] != nullptr) { return {false, "FLASH_ATTN_EXT with sinks is not supported"}; } diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 99ebdb3e15..c370325037 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -6822,7 +6822,8 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons case GGML_OP_SOLVE_TRI: return op->src[0]->ne[0] <= SYCL_SOLVE_TRI_MAX_N && op->src[1]->ne[0] <= SYCL_SOLVE_TRI_MAX_K; case GGML_OP_FLASH_ATTN_EXT: - return ggml_sycl_flash_attn_ext_supported(device, op); + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + return op->src[5] == nullptr && ggml_sycl_flash_attn_ext_supported(device, op); default: return false; } diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index ce40baec2d..09b96edc9b 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -15377,6 +15377,10 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm } case GGML_OP_FLASH_ATTN_EXT: { + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + if (op->src[5] != nullptr) { + return false; + } bool coopmat2 = device->coopmat2; uint32_t HSK = op->src[1]->ne[0]; uint32_t HSV = op->src[2]->ne[0]; diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 9c5dc768ef..6cc70297f2 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -4514,7 +4514,8 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const { // conservative support checks for whether the more resource-intensive shader paths // can be used, to avoid cases where flash_attn is assigned to the CPU later on - supports_op = src0->type == GGML_TYPE_F32 && + // K/V row lists (ggml_flash_attn_ext_rows) are not supported + supports_op = op->src[5] == nullptr && src0->type == GGML_TYPE_F32 && (src1->type == GGML_TYPE_F32 || src1->type == GGML_TYPE_F16 || src1->type == GGML_TYPE_Q4_0 || src1->type == GGML_TYPE_Q8_0) && (src2->type == GGML_TYPE_F32 || src2->type == GGML_TYPE_F16 || diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index cdc10a7597..ef5202e4d5 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -5494,20 +5494,30 @@ struct ggml_tensor * ggml_arange( // ggml_flash_attn_ext -struct ggml_tensor * ggml_flash_attn_ext( +static struct ggml_tensor * ggml_flash_attn_ext_impl( struct ggml_context * ctx, struct ggml_tensor * q, struct ggml_tensor * k, struct ggml_tensor * v, struct ggml_tensor * mask, + struct ggml_tensor * kv_rows, float scale, float max_bias, float logit_softcap) { GGML_ASSERT(ggml_can_mul_mat(k, q)); // TODO: check if vT can be multiplied by (k*qT) - GGML_ASSERT(q->ne[3] == k->ne[3]); - GGML_ASSERT(q->ne[3] == v->ne[3]); + if (kv_rows) { + // k and v are shared by all q slices, kv_rows picks the rows of each slice + GGML_ASSERT(mask && mask->ne[0] == kv_rows->ne[0]); + GGML_ASSERT(kv_rows->type == GGML_TYPE_I32 && ggml_is_contiguous(kv_rows)); + GGML_ASSERT(kv_rows->ne[1] == q->ne[3] && kv_rows->ne[2] == 1 && kv_rows->ne[3] == 1); + GGML_ASSERT(k->ne[3] == 1); + GGML_ASSERT(v->ne[3] == 1); + } else { + GGML_ASSERT(q->ne[3] == k->ne[3]); + GGML_ASSERT(q->ne[3] == v->ne[3]); + } if (mask) { GGML_ASSERT(mask->type == GGML_TYPE_F16); @@ -5534,10 +5544,36 @@ struct ggml_tensor * ggml_flash_attn_ext( result->src[1] = k; result->src[2] = v; result->src[3] = mask; + result->src[5] = kv_rows; return result; } +struct ggml_tensor * ggml_flash_attn_ext( + struct ggml_context * ctx, + struct ggml_tensor * q, + struct ggml_tensor * k, + struct ggml_tensor * v, + struct ggml_tensor * mask, + float scale, + float max_bias, + float logit_softcap) { + return ggml_flash_attn_ext_impl(ctx, q, k, v, mask, NULL, scale, max_bias, logit_softcap); +} + +struct ggml_tensor * ggml_flash_attn_ext_rows( + struct ggml_context * ctx, + struct ggml_tensor * q, + struct ggml_tensor * k, + struct ggml_tensor * v, + struct ggml_tensor * mask, + struct ggml_tensor * kv_rows, + float scale, + float max_bias, + float logit_softcap) { + return ggml_flash_attn_ext_impl(ctx, q, k, v, mask, kv_rows, scale, max_bias, logit_softcap); +} + void ggml_flash_attn_ext_set_prec( struct ggml_tensor * a, diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 99e55da687..3118e443c5 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -45,6 +45,12 @@ static const llm_fused_op_probe llm_fused_op_flash_attn_probe = { /*.n_tokens_per_seq =*/ 1, }; +static const llm_fused_op_probe llm_fused_op_flash_attn_kv_rows_probe = { + /*.op =*/ LLM_FUSED_OP_FLASH_ATTN_KV_ROWS, + /*.name =*/ "Flash Attention over K/V rows", + /*.n_tokens_per_seq =*/ 1, +}; + static const llm_fused_op_probe llm_fused_op_gdn_ar_probe = { /*.op =*/ LLM_FUSED_OP_GDN_AR, /*.name =*/ "fused Gated Delta Net (autoregressive)", @@ -273,6 +279,10 @@ llama_context::llama_context( cparams.op_offload = params.op_offload; cparams.kv_unified = params.kv_unified; + // resolved after flash attention, see resolve_fused_ops() + cparams.fused_kv_rows = false; + cparams.auto_fkvr = cparams.kv_unified; + // initialized later cparams.pipeline_parallel = false; @@ -558,6 +568,13 @@ void llama_context::resolve_fused_ops(const llama_memory_context_i * mctx, uint3 cparams.auto_fa = false; } + // the fallback attends to the whole unified cache with a dense mask + if (cparams.auto_fkvr) { + cparams.fused_kv_rows = cparams.flash_attn; + resolve(llm_fused_op_flash_attn_kv_rows_probe, cparams.fused_kv_rows); + cparams.auto_fkvr = false; + } + if (cparams.auto_fgdn) { LLAMA_LOG_INFO("%s: resolving fused Gated Delta Net support:\n", func); resolve(llm_fused_op_gdn_ar_probe, cparams.fused_gdn_ar); diff --git a/src/llama-cparams.h b/src/llama-cparams.h index b592de18c7..9570dcf898 100644 --- a/src/llama-cparams.h +++ b/src/llama-cparams.h @@ -48,6 +48,8 @@ struct llama_cparams { bool fused_dsv4_hc_comb; bool fused_dsv4_hc_post; bool auto_fhc; + bool fused_kv_rows; // unified cache: attend to the cells of each sequence only (ggml_flash_attn_ext_rows) + bool auto_fkvr; bool no_perf; bool warmup; // TODO: remove [TAG_LLAMA_GRAPH_NO_WARMUP] bool op_offload; diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 0b3bab6123..42e0564e4c 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -64,6 +64,68 @@ static bool can_reuse_kq_mask( return res; } +// attention over the cells of each sequence of a unified cache, see llama_kv_cache::get_n_kv_rows +static uint32_t get_n_kv_rows(const llama_kv_cache_context * mctx, const llama_cparams & cparams) { + return cparams.flash_attn && cparams.fused_kv_rows ? mctx->get_n_kv_rows() : 0; +} + +static bool can_reuse_kq_mask_rows( + ggml_tensor * kq_mask, + ggml_tensor * kv_rows, + uint32_t n_kv, + bool allow_kv_rows, + const llama_kv_cache_context * mctx, + const llama_ubatch & ubatch, + const llama_cparams & cparams) { + const uint32_t n_kv_rows = allow_kv_rows ? get_n_kv_rows(mctx, cparams) : 0; + + if (!kv_rows) { + return n_kv_rows == 0 && can_reuse_kq_mask(kq_mask, mctx, ubatch, cparams); + } + + const uint32_t n_slices = llama_kv_cache::get_n_kv_slices(ubatch); + + bool res = n_kv_rows > 0; + + res &= n_kv == mctx->get_n_kv(); + + res &= kv_rows->ne[0] == n_kv_rows; + res &= kv_rows->ne[1] == n_slices; + + res &= kq_mask->ne[0] == n_kv_rows; + res &= kq_mask->ne[1] == ubatch.n_tokens/n_slices; + res &= kq_mask->ne[3] == n_slices; + + return res; +} + +// KQ mask over the cell lists of kv_rows, or over the whole cache if kv_rows is not used +static ggml_tensor * build_attn_inp_kq_mask_rows( + ggml_context * ctx, + const llama_kv_cache_context * mctx, + const llama_ubatch & ubatch, + const llama_cparams & cparams, + bool allow_kv_rows, + ggml_tensor ** kv_rows) { + const uint32_t n_kv_rows = allow_kv_rows ? get_n_kv_rows(mctx, cparams) : 0; + if (n_kv_rows == 0) { + *kv_rows = nullptr; + return build_attn_inp_kq_mask(ctx, mctx, ubatch, cparams); + } + + const uint32_t n_slices = llama_kv_cache::get_n_kv_slices(ubatch); + + *kv_rows = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, n_kv_rows, n_slices); + ggml_set_input(*kv_rows); + ggml_set_name(*kv_rows, "attn_inp_kv_rows"); + + ggml_tensor * res = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, n_kv_rows, ubatch.n_tokens/n_slices, 1, n_slices); + ggml_set_input(res); + ggml_set_name(res, "attn_inp_kq_mask"); + + return res; +} + // impl void llm_graph_input_embd::set_input(const llama_ubatch * ubatch) { @@ -471,10 +533,14 @@ void llm_graph_input_attn_kv::set_input(const llama_ubatch * ubatch) { mctx->set_input_k_idxs(self_k_idxs, ubatch); mctx->set_input_v_idxs(self_v_idxs, ubatch); + if (self_kv_rows && self_kv_rows->buffer) { + mctx->set_input_kv_rows(self_kv_rows, ubatch); + } + // the mask is left unallocated when the graph only stores K/V without attending // (e.g. DFlash's KV-injection pass) if (self_kq_mask && self_kq_mask->buffer) { - mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn); + mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn, self_kv_rows); } if (self_k_rot && self_k_rot->buffer) { @@ -496,7 +562,7 @@ bool llm_graph_input_attn_kv::can_reuse(const llm_graph_params & params) { res &= self_k_idxs->ne[0] == params.ubatch.n_tokens; //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there - res &= can_reuse_kq_mask(self_kq_mask, mctx, params.ubatch, params.cparams); + res &= can_reuse_kq_mask_rows(self_kq_mask, self_kv_rows, n_kv, allow_kv_rows, mctx, params.ubatch, params.cparams); return res; } @@ -618,9 +684,13 @@ void llm_graph_input_attn_kv_iswa::set_input(const llama_ubatch * ubatch) { } } + if (self_kv_rows && self_kv_rows->buffer) { + mctx->get_base()->set_input_kv_rows(self_kv_rows, ubatch); + } + // the kq mask guards on its own buffer: shared cells leave idxs unbacked while the mask stays live if (self_kq_mask && self_kq_mask->buffer) { - mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn); + mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn, self_kv_rows); } // swa tensors may not be allocated if there are no SWA attention layers @@ -631,8 +701,12 @@ void llm_graph_input_attn_kv_iswa::set_input(const llama_ubatch * ubatch) { } } + if (self_kv_rows_swa && self_kv_rows_swa->buffer) { + mctx->get_swa()->set_input_kv_rows(self_kv_rows_swa, ubatch); + } + if (self_kq_mask_swa && self_kq_mask_swa->buffer) { - mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn); + mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn, self_kv_rows_swa); } if (self_k_rot && self_k_rot->buffer) { @@ -666,7 +740,7 @@ bool llm_graph_input_attn_kv_iswa::can_reuse(const llm_graph_params & params) { } if (self_kq_mask && self_kq_mask->buffer) { - res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams); + res &= can_reuse_kq_mask_rows(self_kq_mask, self_kv_rows, n_kv, true, mctx->get_base(), params.ubatch, params.cparams); } // swa tensors may not be allocated if there are no SWA attention layers @@ -676,7 +750,7 @@ bool llm_graph_input_attn_kv_iswa::can_reuse(const llm_graph_params & params) { } if (self_kq_mask_swa && self_kq_mask_swa->buffer) { - res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams); + res &= can_reuse_kq_mask_rows(self_kq_mask_swa, self_kv_rows_swa, n_kv_swa, true, mctx->get_swa(), params.ubatch, params.cparams); } return res; @@ -1090,7 +1164,11 @@ void llm_graph_input_mem_hybrid::set_input(const llama_ubatch * ubatch) { mctx->get_attn()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch); mctx->get_attn()->set_input_v_idxs(inp_attn->self_v_idxs, ubatch); - mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn); + if (inp_attn->self_kv_rows) { + mctx->get_attn()->set_input_kv_rows(inp_attn->self_kv_rows, ubatch); + } + + mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn, inp_attn->self_kv_rows); if (inp_attn->self_k_rot) { mctx->get_attn()->set_input_k_rot(inp_attn->self_k_rot); @@ -1123,7 +1201,7 @@ bool llm_graph_input_mem_hybrid::can_reuse(const llm_graph_params & params) { res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens; //res &= inp_attn->self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there - res &= can_reuse_kq_mask(inp_attn->self_kq_mask, mctx->get_attn(), params.ubatch, params.cparams); + res &= can_reuse_kq_mask_rows(inp_attn->self_kq_mask, inp_attn->self_kv_rows, inp_attn->n_kv, inp_attn->allow_kv_rows, mctx->get_attn(), params.ubatch, params.cparams); res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs(); @@ -2618,11 +2696,12 @@ ggml_tensor * llm_graph_context::build_attn_mha( ggml_tensor * v_mla, int64_t n_kv_max, float kq_scale, - int il) const { + int il, + ggml_tensor * kv_rows) const { const bool v_trans = v->nb[1] > v->nb[2]; // split the batch into streams if needed - const auto n_stream = k->ne[3]; + const auto n_stream = kv_rows ? kv_rows->ne[1] : k->ne[3]; q = ggml_view_4d(ctx0, q, q->ne[0], q->ne[1], q->ne[2]/n_stream, n_stream, q->nb[1], q->nb[2], q->nb[3]/n_stream, 0); @@ -2633,6 +2712,7 @@ ggml_tensor * llm_graph_context::build_attn_mha( ggml_tensor * cur; const bool use_flash_attn = cparams.flash_attn && kq_b == nullptr; + GGML_ASSERT(use_flash_attn || kv_rows == nullptr); if (use_flash_attn) { GGML_ASSERT(kq_b == nullptr && "Flash attention does not support KQ bias yet"); @@ -2649,9 +2729,14 @@ ggml_tensor * llm_graph_context::build_attn_mha( v = ggml_cast(ctx0, v, GGML_TYPE_F16); } - cur = ggml_flash_attn_ext(ctx0, q, k, v, kq_mask, kq_scale, hparams.f_max_alibi_bias, - hparams.attn_soft_cap ? hparams.f_attn_logit_softcapping : 0.0f); - res->add_fused_node({LLM_FUSED_OP_FLASH_ATTN, cur, il}); + if (kv_rows) { + cur = ggml_flash_attn_ext_rows(ctx0, q, k, v, kq_mask, kv_rows, kq_scale, hparams.f_max_alibi_bias, + hparams.attn_soft_cap ? hparams.f_attn_logit_softcapping : 0.0f); + } else { + cur = ggml_flash_attn_ext(ctx0, q, k, v, kq_mask, kq_scale, hparams.f_max_alibi_bias, + hparams.attn_soft_cap ? hparams.f_attn_logit_softcapping : 0.0f); + } + res->add_fused_node({kv_rows ? LLM_FUSED_OP_FLASH_ATTN_KV_ROWS : LLM_FUSED_OP_FLASH_ATTN, cur, il}); ggml_flash_attn_ext_add_sinks(cur, sinks); GGML_ASSERT(n_kv_max >= 0 && n_kv_max <= INT32_MAX); @@ -2829,18 +2914,23 @@ static std::unique_ptr build_attn_inp_kv_impl( const llama_ubatch & ubatch, const llama_hparams & hparams, const llama_cparams & cparams, - const llama_kv_cache_context * mctx_cur) { + const llama_kv_cache_context * mctx_cur, + bool allow_kv_rows = true) { auto inp = std::make_unique(hparams, cparams, mctx_cur); + inp->allow_kv_rows = allow_kv_rows; + { GGML_ASSERT(hparams.swa_type == LLAMA_SWA_TYPE_NONE && "Use llama_kv_cache_iswa for SWA"); inp->self_k_idxs = mctx_cur->build_input_k_idxs(ctx0, ubatch); inp->self_v_idxs = mctx_cur->build_input_v_idxs(ctx0, ubatch); - inp->self_kq_mask = build_attn_inp_kq_mask(ctx0, mctx_cur, ubatch, cparams); + inp->self_kq_mask = build_attn_inp_kq_mask_rows(ctx0, mctx_cur, ubatch, cparams, allow_kv_rows, &inp->self_kv_rows); inp->self_kq_mask_cnv = inp->self_kq_mask; + + inp->n_kv = mctx_cur->get_n_kv(); } inp->self_k_rot = mctx_cur->build_input_k_rot(ctx0); @@ -2905,7 +2995,7 @@ ggml_tensor * llm_graph_context::build_attn( ggml_tensor * k = mctx_cur->get_k(ctx0, il); ggml_tensor * v = mctx_cur->get_v(ctx0, il); - ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask, sinks, v_mla, 0, kq_scale, il); + ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask, sinks, v_mla, 0, kq_scale, il, inp->self_kv_rows); cb(cur, "kqv_out", il); if (inp->self_v_rot) { @@ -3155,12 +3245,13 @@ ggml_tensor * llm_graph_context::build_attn( } const auto & kq_mask = is_swa ? inp->get_kq_mask_swa() : inp->get_kq_mask(); + const auto & kv_rows = is_swa ? inp->self_kv_rows_swa : inp->self_kv_rows; ggml_tensor * q = q_cur; ggml_tensor * k = mctx_cur->get_k(ctx0, il); ggml_tensor * v = mctx_cur->get_v(ctx0, il); - ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask, sinks, v_mla, 0, kq_scale, il); + ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask, sinks, v_mla, 0, kq_scale, il, kv_rows); cb(cur, "kqv_out", il); if (v_rot) { @@ -3406,8 +3497,10 @@ llm_graph_input_attn_kv_iswa * llm_graph_context::build_attn_inp_kv_iswa() const inp->self_k_idxs = mctx_cur->get_base()->build_input_k_idxs(ctx0, ubatch); inp->self_v_idxs = mctx_cur->get_base()->build_input_v_idxs(ctx0, ubatch); - inp->self_kq_mask = build_attn_inp_kq_mask(ctx0, mctx_cur->get_base(), ubatch, cparams); + inp->self_kq_mask = build_attn_inp_kq_mask_rows(ctx0, mctx_cur->get_base(), ubatch, cparams, true, &inp->self_kv_rows); inp->self_kq_mask_cnv = inp->self_kq_mask; + + inp->n_kv = mctx_cur->get_base()->get_n_kv(); } { @@ -3416,8 +3509,10 @@ llm_graph_input_attn_kv_iswa * llm_graph_context::build_attn_inp_kv_iswa() const inp->self_k_idxs_swa = mctx_cur->get_swa()->build_input_k_idxs(ctx0, ubatch); inp->self_v_idxs_swa = mctx_cur->get_swa()->build_input_v_idxs(ctx0, ubatch); - inp->self_kq_mask_swa = build_attn_inp_kq_mask(ctx0, mctx_cur->get_swa(), ubatch, cparams); + inp->self_kq_mask_swa = build_attn_inp_kq_mask_rows(ctx0, mctx_cur->get_swa(), ubatch, cparams, true, &inp->self_kv_rows_swa); inp->self_kq_mask_swa_cnv = inp->self_kq_mask_swa; + + inp->n_kv_swa = mctx_cur->get_swa()->get_n_kv(); } inp->self_k_rot = mctx_cur->get_base()->build_input_k_rot(ctx0); @@ -3605,11 +3700,11 @@ ggml_tensor * llm_graph_context::build_rwkv_token_shift_store( ); } -llm_graph_input_mem_hybrid * llm_graph_context::build_inp_mem_hybrid() const { +llm_graph_input_mem_hybrid * llm_graph_context::build_inp_mem_hybrid(bool allow_kv_rows) const { const auto * mctx_cur = static_cast(mctx); auto inp_rs = build_rs_inp_impl (ctx0, ubatch, mctx_cur->get_recr()); - auto inp_attn = build_attn_inp_kv_impl(ctx0, ubatch, hparams, cparams, mctx_cur->get_attn()); + auto inp_attn = build_attn_inp_kv_impl(ctx0, ubatch, hparams, cparams, mctx_cur->get_attn(), allow_kv_rows); auto inp = std::make_unique(cparams, std::move(inp_attn), std::move(inp_rs), mctx_cur); diff --git a/src/llama-graph.h b/src/llama-graph.h index 3daa425bc0..eb0d44622d 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -44,6 +44,7 @@ enum llm_graph_type { enum llm_fused_op { LLM_FUSED_OP_FLASH_ATTN, + LLM_FUSED_OP_FLASH_ATTN_KV_ROWS, LLM_FUSED_OP_GDN_AR, LLM_FUSED_OP_GDN_CH, LLM_FUSED_OP_LIGHTNING_INDEXER, @@ -346,6 +347,15 @@ public: ggml_tensor * self_kq_mask = nullptr; // F32/F16 [n_kv, n_batch/n_stream, 1, n_stream] ggml_tensor * self_kq_mask_cnv = nullptr; // [n_kv, n_batch/n_stream, 1, n_stream] + // attention over the cells of each sequence, the mask columns follow these rows + ggml_tensor * self_kv_rows = nullptr; // I32 [n_kv_rows, n_slices] + + // cells of the K/V views, the rows must stay inside them + uint32_t n_kv = 0; + + // false if the model reads the KQ mask of the whole cache itself + bool allow_kv_rows = true; + // note: assumes v_rot^2 == I ggml_tensor * self_k_rot = nullptr; ggml_tensor * self_v_rot = nullptr; @@ -516,6 +526,13 @@ public: ggml_tensor * self_kq_mask_swa = nullptr; // F32/F16 [n_kv, n_batch/n_stream, 1, n_stream] ggml_tensor * self_kq_mask_swa_cnv = nullptr; // [n_kv, n_batch/n_stream, 1, n_stream] + // attention over the cells of each sequence, see llm_graph_input_attn_kv + ggml_tensor * self_kv_rows = nullptr; // I32 [n_kv_rows, n_slices] + ggml_tensor * self_kv_rows_swa = nullptr; // I32 [n_kv_rows, n_slices] + + uint32_t n_kv = 0; + uint32_t n_kv_swa = 0; + ggml_tensor * self_k_rot = nullptr; ggml_tensor * self_v_rot = nullptr; @@ -1192,7 +1209,8 @@ struct llm_graph_context { ggml_tensor * v_mla, // [n_embd_head_v_mla, n_embd_head_v, n_head_v] int64_t n_kv_max, float kq_scale, - int il) const; + int il, + ggml_tensor * kv_rows = nullptr) const; // [n_kv, n_stream], k and v are shared by the streams llm_graph_input_attn_no_cache * build_attn_inp_no_cache() const; @@ -1359,7 +1377,7 @@ struct llm_graph_context { // hybrid // - llm_graph_input_mem_hybrid * build_inp_mem_hybrid() const; + llm_graph_input_mem_hybrid * build_inp_mem_hybrid(bool allow_kv_rows = true) const; llm_graph_input_mem_hybrid_k * build_inp_mem_hybrid_k() const; llm_graph_input_mem_hybrid_iswa * build_inp_mem_hybrid_iswa() const; diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 1c93f908b0..e15624b5da 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -1263,6 +1263,59 @@ uint32_t llama_kv_cache::get_n_kv(const slot_info & sinfo) const { return result; } +uint32_t llama_kv_cache::get_n_kv_slices(const llama_ubatch & ubatch) { + if (ubatch.n_seqs_unq == 1) { + return 1; + } + + // each sequence set of an equal split is a single sequence + if (ubatch.equal_seqs() && ubatch.n_seqs == ubatch.n_seqs_unq) { + return ubatch.n_seqs; + } + + // one token per sequence + if (ubatch.n_tokens == ubatch.n_seqs_unq) { + for (uint32_t i = 0; i < ubatch.n_tokens; ++i) { + if (ubatch.n_seq_id[i] != 1) { + return 0; + } + } + + return ubatch.n_tokens; + } + + return 0; +} + +uint32_t llama_kv_cache::get_n_kv_rows(const llama_ubatch & ubatch, uint32_t n_kv) const { + if (n_stream != 1 || get_n_kv_slices(ubatch) == 0) { + return 0; + } + + const auto & cells = v_cells[0]; + + uint32_t n_max = 0; + uint32_t n_sum = 0; + + for (uint32_t s = 0; s < ubatch.n_seqs_unq; ++s) { + const uint32_t n_cells = cells.seq_n_cells(ubatch.seq_id_unq[s]); + + n_max = std::max(n_max, n_cells); + n_sum += n_cells; + } + + // pad like n_kv so that the graph can be reused + const uint32_t n_pad_cur = std::max(n_pad, 256u); + const uint32_t n_rows = GGML_PAD(n_max, n_pad_cur); + + // cells shared by the sequences (f.ex. a common prompt) are cheaper to attend once over the whole cache + if (n_rows >= n_kv || n_sum > n_kv) { + return 0; + } + + return n_rows; +} + ggml_tensor * llama_kv_cache::get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const { const int32_t ikv = map_layer_ids.at(il); @@ -1551,6 +1604,9 @@ struct args_set_input_kq_mask { int64_t n_kv; int64_t n_stream; int64_t n_tps; + + // cell of each mask column per stream, nullptr if the columns are the cells + const int32_t * rows; }; template @@ -1568,6 +1624,8 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data const int64_t n_stream = args.n_stream; const int64_t n_tps = args.n_tps; + const int32_t * rows_all = args.rows; + const T mask_keep = llama_cast(0.0f); const T mask_drop = llama_cast(-INFINITY); @@ -1586,6 +1644,8 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data std::unordered_map seq_srct; std::unordered_map> seq_idxs; + const int32_t * rows = rows_all ? rows_all + s*n_kv : nullptr; + for (uint32_t ii = 0; ii < n_tps; ++ii) { const uint32_t i = s*n_tps + ii; @@ -1630,7 +1690,8 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data } for (uint32_t jj = 0; jj < n_kv; ++jj) { - uint32_t j = jj; + // column of the mask + uint32_t c = jj; // we have an exiting mask for this sequence -> update just seq_idxs if (!alibi) { @@ -1639,11 +1700,14 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data break; } - j = idxs[jj]; + c = idxs[jj]; } } - if (cells.is_empty(j)) { + // cell of the column, negative for padding + const int32_t j = rows ? rows[c] : (int32_t) c; + + if (j < 0 || cells.is_empty(j)) { goto skip; } @@ -1658,7 +1722,7 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data if (!prev) { // record all cells for which: p0 >= seq_pos_min[seq_id] - n_swa - 32 if (p0 + (int32_t) (n_swa + 32) >= seq_pos_min[seq_id]) { - idxs.push_back(j); + idxs.push_back(c); } } } @@ -1691,14 +1755,14 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data } if (alibi) { - data[idst + j] = llama_cast(static_cast(-std::abs(p0 - p1))); + data[idst + c] = llama_cast(static_cast(-std::abs(p0 - p1))); } else { - data[idst + j] = mask_keep; + data[idst + c] = mask_keep; } continue; skip: - data[idst + j] = mask_drop; + data[idst + c] = mask_drop; } } } @@ -1743,7 +1807,7 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data } } -void llama_kv_cache::set_input_kq_mask(ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn) const { +void llama_kv_cache::set_input_kq_mask(ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn, const ggml_tensor * kv_rows) const { const uint32_t n_tokens = ubatch->n_tokens; GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); @@ -1774,8 +1838,14 @@ void llama_kv_cache::set_input_kq_mask(ggml_tensor * dst, const llama_ubatch * u /*.n_kv =*/ n_kv, /*.n_stream =*/ n_stream, /*.n_tps =*/ n_tps, + /*.rows =*/ kv_rows ? (const int32_t *) kv_rows->data : nullptr, }; + if (kv_rows) { + GGML_ASSERT(ggml_backend_buffer_is_host(kv_rows->buffer)); + GGML_ASSERT(kv_rows->ne[0] == n_kv && kv_rows->ne[1] == n_stream); + } + if (dst->type == GGML_TYPE_F16) { set_input_kq_mask_impl(args, (ggml_fp16_t *) dst->data, causal_attn); } else { @@ -1787,6 +1857,29 @@ void llama_kv_cache::set_input_kq_mask(ggml_tensor * dst, const llama_ubatch * u //LLAMA_LOG_ERROR("%s: kq mask time: %0.3f ms\n", __func__, (t_end - t_start)/1000.0); } +void llama_kv_cache::set_input_kv_rows(ggml_tensor * dst, const llama_ubatch * ubatch) const { + GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); + GGML_ASSERT(n_stream == 1); + + const int64_t n_rows = dst->ne[0]; + const int64_t n_slices = dst->ne[1]; + const int64_t n_tps = ubatch->n_tokens/n_slices; + + const auto & cells = v_cells[0]; + + int32_t * data = (int32_t *) dst->data; + + for (int64_t s = 0; s < n_slices; ++s) { + int32_t * rows = data + s*n_rows; + + const auto & cells_seq = cells.seq_cells(ubatch->seq_id[s*n_tps][0]); + GGML_ASSERT((int64_t) cells_seq.size() <= n_rows); + + std::copy(cells_seq.begin(), cells_seq.end(), rows); + std::fill(rows + cells_seq.size(), rows + n_rows, -1); + } +} + void llama_kv_cache::set_input_pos_bucket(ggml_tensor * dst, const llama_ubatch * ubatch) const { const int64_t n_tokens = ubatch->n_tokens; @@ -2780,6 +2873,9 @@ llama_kv_cache_context::llama_kv_cache_context( llama_kv_cache * kv) : status(LLAMA_MEMORY_STATUS_SUCCESS), kv(kv) { n_kv = kv->get_size(); + // worst case for graph reservation, only a unified cache splits into per-sequence lists + n_kv_rows = kv->get_n_stream() == 1 ? n_kv : 0; + const uint32_t n_stream = kv->get_n_stream(); // create a dummy slot info - the actual data is irrelevant. we just need to build the graph @@ -2832,7 +2928,8 @@ bool llama_kv_cache_context::apply() { } kv->apply_ubatch(sinfos[i_cur], ubatches[i_cur]); - n_kv = kv->get_n_kv(sinfos[i_cur]); + n_kv = kv->get_n_kv(sinfos[i_cur]); + n_kv_rows = kv->get_n_kv_rows(ubatches[i_cur], n_kv); return true; } @@ -2851,6 +2948,10 @@ uint32_t llama_kv_cache_context::get_n_kv() const { return n_kv; } +uint32_t llama_kv_cache_context::get_n_kv_rows() const { + return n_kv_rows; +} + ggml_type llama_kv_cache_context::type_k() const { return kv->type_k(); } @@ -2903,8 +3004,12 @@ void llama_kv_cache_context::set_input_v_idxs(ggml_tensor * dst, const llama_uba kv->set_input_v_idxs(dst, ubatch, sinfos[i_cur]); } -void llama_kv_cache_context::set_input_kq_mask(ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn) const { - kv->set_input_kq_mask(dst, ubatch, causal_attn); +void llama_kv_cache_context::set_input_kq_mask(ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn, const ggml_tensor * kv_rows) const { + kv->set_input_kq_mask(dst, ubatch, causal_attn, kv_rows); +} + +void llama_kv_cache_context::set_input_kv_rows(ggml_tensor * dst, const llama_ubatch * ubatch) const { + kv->set_input_kv_rows(dst, ubatch); } void llama_kv_cache_context::set_input_pos_bucket(ggml_tensor * dst, const llama_ubatch * ubatch) const { diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index 5051f43433..fe843a99a5 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -188,6 +188,13 @@ public: uint32_t get_n_kv(const slot_info & sinfo) const; + // attention over the cells of each sequence of the ubatch instead of the whole cache (see ggml_flash_attn_ext_rows) + // the ubatch splits into n_slices groups of tokens with one sequence each, 0 if it does not + static uint32_t get_n_kv_slices(const llama_ubatch & ubatch); + + // padded length of the cell list of each slice, 0 if the ubatch should attend to the whole cache + uint32_t get_n_kv_rows(const llama_ubatch & ubatch, uint32_t n_kv) const; + // get views of the current state of the cache ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; @@ -229,7 +236,8 @@ public: void set_input_k_shift(ggml_tensor * dst) const; - void set_input_kq_mask (ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn) const; + void set_input_kq_mask (ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn, const ggml_tensor * kv_rows = nullptr) const; + void set_input_kv_rows (ggml_tensor * dst, const llama_ubatch * ubatch) const; void set_input_pos_bucket(ggml_tensor * dst, const llama_ubatch * ubatch) const; void set_input_k_rot(ggml_tensor * dst) const; @@ -395,6 +403,7 @@ public: // uint32_t get_n_kv() const; + uint32_t get_n_kv_rows() const; ggml_type type_k() const; ggml_type type_v() const; @@ -425,7 +434,8 @@ public: void set_input_v_idxs(ggml_tensor * dst, const llama_ubatch * ubatch) const; void set_input_k_shift (ggml_tensor * dst) const; - void set_input_kq_mask (ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn) const; + void set_input_kq_mask (ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn, const ggml_tensor * kv_rows = nullptr) const; + void set_input_kv_rows (ggml_tensor * dst, const llama_ubatch * ubatch) const; void set_input_pos_bucket(ggml_tensor * dst, const llama_ubatch * ubatch) const; void set_input_k_rot(ggml_tensor * dst) const; @@ -466,4 +476,7 @@ private: // a heuristic, to avoid attending the full cache if it is not yet utilized // as the cache gets filled, the benefit from this heuristic disappears int32_t n_kv; + + // see llama_kv_cache::get_n_kv_rows + uint32_t n_kv_rows = 0; }; diff --git a/src/llama-kv-cells.h b/src/llama-kv-cells.h index 5d567a6ed0..533130dd60 100644 --- a/src/llama-kv-cells.h +++ b/src/llama-kv-cells.h @@ -3,6 +3,7 @@ #include "llama.h" #include "llama-cparams.h" +#include #include #include #include @@ -52,6 +53,8 @@ public: for (uint32_t s = 0; s < LLAMA_MAX_SEQ; ++s) { seq_pos[s].clear(); } + + seq_cells_ok.reset(); } void reset_shift() { @@ -360,6 +363,34 @@ public: return -1; } + // number of cells that contain seq_id + uint32_t seq_n_cells(llama_seq_id seq_id) const { + assert(seq_id >= 0); + assert(seq_id < LLAMA_MAX_SEQ); + + return seq_pos[seq_id].size(); + } + + // the cells that contain seq_id, in no particular order + const std::vector & seq_cells(llama_seq_id seq_id) const { + assert(seq_id >= 0); + assert(seq_id < LLAMA_MAX_SEQ); + + auto & res = seq_cells_list[seq_id]; + + if (!seq_cells_ok.test(seq_id)) { + res.clear(); + res.reserve(seq_pos[seq_id].size()); + for (const auto & [p, i] : seq_pos[seq_id]) { + res.push_back(i); + } + + seq_cells_ok.set(seq_id); + } + + return res; + } + // the minimum position of sequence seq_id currently present in any of the cells // return -1 if the sequence is not present llama_pos seq_pos_min(llama_seq_id seq_id) const { @@ -523,16 +554,39 @@ private: // std::set> seq_pos[LLAMA_MAX_SEQ]; + // flat copy of the cells of each seq for fast iteration: new cells are appended, a removal rebuilds it on the next use + mutable std::vector seq_cells_list[LLAMA_MAX_SEQ]; + mutable std::bitset seq_cells_ok; + // helper functions for updating `seq_pos`, once cell at a time: void seq_pos_dec(llama_seq_id s, uint32_t i) { const auto n = seq_pos[s].erase({ pos[i], i }); assert(n == 1); GGML_UNUSED(n); + + // recently added cells are removed often (f.ex. prepare() reverts a trial apply), look for them near the end + if (seq_cells_ok.test(s)) { + auto & list = seq_cells_list[s]; + + const size_t n_search = std::min(list.size(), 4096); + for (size_t k = 0; k < n_search; ++k) { + if (list[list.size() - 1 - k] == (int32_t) i) { + list.erase(list.end() - 1 - k); + return; + } + } + + seq_cells_ok.reset(s); + } } void seq_pos_inc(llama_seq_id s, uint32_t i) { seq_pos[s].insert({ pos[i], i }); + + if (seq_cells_ok.test(s)) { + seq_cells_list[s].push_back(i); + } } // remove cell i diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp index f33989de02..bc48544852 100644 --- a/src/models/qwen4exp.cpp +++ b/src/models/qwen4exp.cpp @@ -363,7 +363,8 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa cb(inpL, "model.input_embed", -1); ggml_build_forward_expand(gf, inpL); - auto * inp = build_inp_mem_hybrid(); + // QSA reads the KQ mask of the whole cache + auto * inp = build_inp_mem_hybrid(false); // qwen4exp always builds llama_memory_hybrid_idx, so this downcast is safe // the indexer cache inside it is absent when the GGUF has no indexer tensors diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 32c3f3d2a3..1ae22e30bb 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -7994,6 +7994,102 @@ struct test_flash_attn_ext : public test_case { } }; +// GGML_OP_FLASH_ATTN_EXT with kv_rows: each q slice attends to its own list of K/V cache rows +struct test_flash_attn_ext_kv_rows : public test_case { + const int64_t hs; + const int64_t nh; // K/V heads + const int64_t gqa; + const int64_t n_cells; // rows of the K/V cache + const int64_t kv; // list length per slice + const int64_t nb; // queries per slice + const int64_t n_slices; + const ggml_type type_KV; + + std::vector rows; + + std::string vars() override { + return VARS_TO_STR8(hs, nh, gqa, n_cells, kv, nb, n_slices, type_KV); + } + + double max_nmse_err() override { + return 5e-4; + } + + test_flash_attn_ext_kv_rows(int64_t hs = 256, int64_t nh = 2, int64_t gqa = 8, int64_t n_cells = 2048, int64_t kv = 512, + int64_t nb = 8, int64_t n_slices = 4, ggml_type type_KV = GGML_TYPE_F16) + : hs(hs), nh(nh), gqa(gqa), n_cells(n_cells), kv(kv), nb(nb), n_slices(n_slices), type_KV(type_KV) {} + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh*gqa, n_slices); + ggml_set_name(q, "q"); + + // cache layout [hs, nh, n_cells], shared by all slices + ggml_tensor * k0 = ggml_new_tensor_3d(ctx, type_KV, hs, nh, n_cells); + ggml_tensor * v0 = ggml_new_tensor_3d(ctx, type_KV, hs, nh, n_cells); + ggml_set_name(k0, "k0"); + ggml_set_name(v0, "v0"); + + ggml_tensor * k = ggml_permute(ctx, k0, 0, 2, 1, 3); + ggml_tensor * v = ggml_permute(ctx, v0, 0, 2, 1, 3); + + ggml_tensor * m = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, kv, nb, 1, n_slices); + ggml_set_name(m, "m"); + + ggml_tensor * r = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, kv, n_slices); + ggml_set_name(r, "r"); + + ggml_tensor * out = ggml_flash_attn_ext_rows(ctx, q, k, v, m, r, 1.0f/sqrtf(hs), 0.0f, 0.0f); + ggml_prec_set_acc(out, GGML_PREC_F32); + ggml_set_name(out, "out"); + + return out; + } + + void initialize_tensors(ggml_context * ctx) override { + std::mt19937 gen(0x6B76); + + // per slice: distinct random cells, the tail of the list is padding + rows.assign(kv*n_slices, -1); + std::vector order(n_cells); + for (int64_t s = 0; s < n_slices; ++s) { + for (int64_t i = 0; i < n_cells; ++i) { + order[i] = i; + } + std::shuffle(order.begin(), order.end(), gen); + const int64_t n_used = std::min(n_cells, kv - (s*37) % std::max(1, kv/2)); + for (int64_t i = 0; i < n_used; ++i) { + rows[s*kv + i] = order[i]; + } + } + + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) { + if (t->view_src) { + continue; + } + if (strcmp(t->name, "r") == 0) { + ggml_backend_tensor_set(t, rows.data(), 0, rows.size()*sizeof(int32_t)); + } else if (strcmp(t->name, "m") == 0) { + std::vector f(ggml_nelements(t)); + std::uniform_real_distribution dis(-1.0f, 1.0f); + for (int64_t s = 0; s < n_slices; ++s) { + for (int64_t j = 0; j < nb; ++j) { + for (int64_t i = 0; i < kv; ++i) { + const bool pad = rows[s*kv + i] < 0; + const bool drop = (i + 3*j + 5*s) % 11 == 0; + f[(s*nb + j)*kv + i] = pad || drop ? -INFINITY : dis(gen); + } + } + } + std::vector h(f.size()); + ggml_fp32_to_fp16_row(f.data(), h.data(), f.size()); + ggml_backend_tensor_set(t, h.data(), 0, h.size()*sizeof(ggml_fp16_t)); + } else { + init_tensor_uniform(t); + } + } + } +}; + // GGML_OP_CROSS_ENTROPY_LOSS struct test_cross_entropy_loss : public test_case { const ggml_type type; @@ -11013,6 +11109,18 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 8, 1}, 4096, 4, true, false, 8.0f, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512)); + // flash attention over a list of K/V rows per slice (unified KV cache) + for (int64_t nb : {1, 3, 8, 64, 512}) { + for (int64_t n_slices : {1, 4}) { + test_cases.emplace_back(new test_flash_attn_ext_kv_rows(256, 2, 8, 4096, 1024, nb, n_slices)); + test_cases.emplace_back(new test_flash_attn_ext_kv_rows(128, 8, 4, 4096, 512, nb, n_slices)); + test_cases.emplace_back(new test_flash_attn_ext_kv_rows(128, 4, 1, 1000, 256, nb, n_slices)); + test_cases.emplace_back(new test_flash_attn_ext_kv_rows( 64, 4, 2, 3000, 768, nb, n_slices, GGML_TYPE_Q8_0)); + test_cases.emplace_back(new test_flash_attn_ext_kv_rows(512, 2, 8, 4096, 1024, nb, n_slices)); + test_cases.emplace_back(new test_flash_attn_ext_kv_rows(256, 8, 2, 4096, 1280, nb, n_slices)); + } + } + // sparse mask with large batch size test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 4096, 64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 512)); test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 4096, 64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));