llama + CUDA: flash_attn_ext_rows (fix unified kv for multi-seq)

This commit is contained in:
Aman Gupta
2026-09-27 15:32:38 +08:00
parent 85ca3b52c3
commit 7908e9e8ce
27 changed files with 626 additions and 68 deletions
+18
View File
@@ -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),
+4
View File
@@ -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;
+28 -8
View File
@@ -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<int32_t>(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);
+4
View File
@@ -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:
+18 -5
View File
@@ -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],
+16 -7
View File
@@ -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<swz, half2>(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<stride_tile_V, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(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<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(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<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(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,
+3 -1
View File
@@ -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,
+3 -1
View File
@@ -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,
+16 -3
View File
@@ -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<int64_t>(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<int64_t>(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;
+1 -1
View File
@@ -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;
+5
View File
@@ -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];
+4
View File
@@ -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 &&
+4
View File
@@ -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;
+5
View File
@@ -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"};
}
+2 -1
View File
@@ -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;
}
+4
View File
@@ -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];
+2 -1
View File
@@ -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 ||
+39 -3
View File
@@ -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,
+17
View File
@@ -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);
+2
View File
@@ -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;
+116 -21
View File
@@ -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<llm_graph_input_attn_kv> 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<llm_graph_input_attn_kv>(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<const llama_memory_hybrid_context *>(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<llm_graph_input_mem_hybrid>(cparams, std::move(inp_attn), std::move(inp_rs), mctx_cur);
+20 -2
View File
@@ -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;
+116 -11
View File
@@ -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<typename T, bool causal, bool swa, bool is_2d, bool alibi>
@@ -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<T>(0.0f);
const T mask_drop = llama_cast<T>(-INFINITY);
@@ -1586,6 +1644,8 @@ static void set_input_kq_mask_impl(const args_set_input_kq_mask & args, T * data
std::unordered_map<llama_seq_id, uint32_t> seq_srct;
std::unordered_map<llama_seq_id, std::vector<uint32_t>> 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<T>(static_cast<float>(-std::abs(p0 - p1)));
data[idst + c] = llama_cast<T>(static_cast<float>(-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<ggml_fp16_t>(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 {
+15 -2
View File
@@ -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;
};
+54
View File
@@ -3,6 +3,7 @@
#include "llama.h"
#include "llama-cparams.h"
#include <algorithm>
#include <bitset>
#include <cassert>
#include <cstring>
@@ -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<int32_t> & 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<std::pair<llama_pos, uint32_t>> 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<int32_t> seq_cells_list[LLAMA_MAX_SEQ];
mutable std::bitset<LLAMA_MAX_SEQ> 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<size_t>(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
+2 -1
View File
@@ -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
+108
View File
@@ -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<int32_t> 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<int32_t> 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<int64_t>(n_cells, kv - (s*37) % std::max<int64_t>(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<float> f(ggml_nelements(t));
std::uniform_real_distribution<float> 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<ggml_fp16_t> 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<std::unique_ptr<test_case>> 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));