mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-02 10:57:33 -05:00
llama + CUDA: flash_attn_ext_rows (fix unified kv for multi-seq)
This commit is contained in:
@@ -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),
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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];
|
||||
|
||||
@@ -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 &&
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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"};
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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];
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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;
|
||||
};
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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));
|
||||
|
||||
Reference in New Issue
Block a user