Commit 926862e57 for llama.cpp
commit 926862e574617d5e5ab9e9c9bae317f98237f583
Author: Ethan Guo <ethanguo.dev@gmail.com>
Date: Fri Oct 2 18:19:21 2026 +0800
metal : add tensor API flash attention kernel for F16 KV (#29570)
* metal : add tensor API flash attention kernel for F16 KV
* cont : add tensor FA kernels for DK=DV=512 and DK=576, DV=512
* cont : support attention sinks, ALiBi and logit softcap in the tensor FA kernel
* cont : add tensor FA kernel for DK=192, DV=128
diff --git a/ggml/src/ggml-metal/CMakeLists.txt b/ggml/src/ggml-metal/CMakeLists.txt
index 68532a984..08408a2d4 100644
--- a/ggml/src/ggml-metal/CMakeLists.txt
+++ b/ggml/src/ggml-metal/CMakeLists.txt
@@ -232,10 +232,19 @@ else()
VERBATIM
)
+ set(AIR_FA_TENSOR "${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/fa_f16_tensor.air")
+ add_custom_command(
+ OUTPUT ${AIR_FA_TENSOR}
+ COMMAND xcrun -sdk ${METAL_SDK} metal ${XC_FLAGS_TENSOR} -DGGML_METAL_HAS_TENSOR -I ${CMAKE_RUNTIME_OUTPUT_DIRECTORY} -c ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/kernels/fa_f16.metal -o ${AIR_FA_TENSOR}
+ DEPENDS kernels/fa_f16.metal ${METALLIB_KERNELS_FA_SHARED} kernels/common.h kernels/dequantize.h ${METALLIB_COMMON} ggml-metal-impl.h
+ COMMENT "Compiling kernels/fa_f16.metal (tensor API)"
+ VERBATIM
+ )
+
add_custom_command(
OUTPUT ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib
- COMMAND xcrun -sdk ${METAL_SDK} metallib ${AIR_MM_TENSOR} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib
- DEPENDS ${AIR_MM_TENSOR}
+ COMMAND xcrun -sdk ${METAL_SDK} metallib ${AIR_MM_TENSOR} ${AIR_FA_TENSOR} -o ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-tensor.metallib
+ DEPENDS ${AIR_MM_TENSOR} ${AIR_FA_TENSOR}
COMMENT "Linking tensor API Metal kernels into ggml-tensor.metallib"
)
@@ -248,7 +257,7 @@ else()
COMMAND rm -f ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-common.h
COMMAND rm -f ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/ggml-metal-impl.h
COMMAND rm -rf ${CMAKE_RUNTIME_OUTPUT_DIRECTORY}/kernels
- DEPENDS ${AIR_FILES} ${AIR_MM_TENSOR}
+ DEPENDS ${AIR_FILES} ${AIR_MM_TENSOR} ${AIR_FA_TENSOR}
COMMENT "Linking Metal kernels into default.metallib"
)
diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp
index 95b6c513f..8cf2c8212 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-device.cpp
@@ -1715,6 +1715,41 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_b
return res;
}
+ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_tensor(
+ ggml_metal_library_t lib,
+ const ggml_tensor * op,
+ bool has_mask,
+ bool has_sinks,
+ bool has_bias,
+ bool has_scap) {
+ assert(op->op == GGML_OP_FLASH_ATTN_EXT);
+
+ char base[256];
+ char name[256];
+
+ const int32_t dk = (int32_t) op->src[1]->ne[0];
+ const int32_t dv = (int32_t) op->src[2]->ne[0];
+
+ snprintf(base, 256, "kernel_flash_attn_ext_tensor_f16_dk%d_dv%d", dk, dv);
+ snprintf(name, 256, "%s_mask=%d_sinks=%d_bias=%d_scap=%d", base, has_mask, has_sinks, has_bias, has_scap);
+
+ ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
+ if (!res.pipeline) {
+ ggml_metal_cv_t cv = ggml_metal_cv_init();
+
+ ggml_metal_cv_set_bool(cv, has_mask, FC_FLASH_ATTN_EXT_TENSOR + 0);
+ ggml_metal_cv_set_bool(cv, has_sinks, FC_FLASH_ATTN_EXT_TENSOR + 1);
+ ggml_metal_cv_set_bool(cv, has_bias, FC_FLASH_ATTN_EXT_TENSOR + 2);
+ ggml_metal_cv_set_bool(cv, has_scap, FC_FLASH_ATTN_EXT_TENSOR + 3);
+
+ res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+ ggml_metal_cv_free(cv);
+ }
+
+ return res;
+}
+
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext(
ggml_metal_library_t lib,
const ggml_tensor * op,
diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h
index 1bdaecc73..794fe979e 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.h
+++ b/ggml/src/ggml-metal/ggml-metal-device.h
@@ -194,6 +194,14 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_att
int32_t nqptg,
int32_t ncpsg);
+struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext_tensor(
+ ggml_metal_library_t lib,
+ const struct ggml_tensor * op,
+ bool has_mask,
+ bool has_sinks,
+ bool has_bias,
+ bool has_scap);
+
struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_flash_attn_ext(
ggml_metal_library_t lib,
const struct ggml_tensor * op,
diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h
index eb1868f62..3a34c81a4 100644
--- a/ggml/src/ggml-metal/ggml-metal-impl.h
+++ b/ggml/src/ggml-metal/ggml-metal-impl.h
@@ -121,11 +121,17 @@
#define FC_MOE_REDUCE 1900
#define FC_DSV4_HC 2000
#define FC_PAD 2100
+#define FC_FLASH_ATTN_EXT_TENSOR 2200
// op-specific constants
#define OP_FLASH_ATTN_EXT_NQPSG 8
#define OP_FLASH_ATTN_EXT_NCPSG 64
+#define OP_FLASH_ATTN_EXT_TENSOR_NQPSG 32
+#define OP_FLASH_ATTN_EXT_TENSOR_NQPSG_LARGE 16
+#define OP_FLASH_ATTN_EXT_TENSOR_NCPSG 64
+#define OP_FLASH_ATTN_EXT_TENSOR_NSG 8
+
#define OP_FLASH_ATTN_EXT_VEC_NQPSG 1
#define OP_FLASH_ATTN_EXT_VEC_NCPSG 32
diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp
index ab6d4f065..ed4fe47dd 100644
--- a/ggml/src/ggml-metal/ggml-metal-ops.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp
@@ -2982,6 +2982,47 @@ static bool ggml_metal_op_flash_attn_ext_use_kv_f16(const ggml_tensor * op) {
}
}
+static bool ggml_metal_op_flash_attn_ext_use_tensor(const ggml_tensor * op, bool has_tensor) {
+ assert(op->op == GGML_OP_FLASH_ATTN_EXT);
+
+ if (!has_tensor || ggml_metal_op_flash_attn_ext_use_vec(op)) {
+ return false;
+ }
+
+ const int64_t ne01 = op->src[0]->ne[1];
+ const int64_t ne02 = op->src[0]->ne[2];
+ const int64_t ne03 = op->src[0]->ne[3];
+
+ const int64_t dk = op->src[1]->ne[0];
+ const int64_t dv = op->src[2]->ne[0];
+
+ const bool dk_dv_ok = (dk == 64 && dv == 64) ||
+ (dk == 128 && dv == 128) ||
+ (dk == 192 && dv == 128) ||
+ (dk == 256 && dv == 256) ||
+ (dk == 512 && dv == 512) ||
+ (dk == 576 && dv == 512);
+
+ if (!dk_dv_ok) {
+ return false;
+ }
+
+ // large heads use fewer queries per threadgroup, so that the queries fit in threadgroup memory
+ const int64_t nqptg = dk >= 512 ? OP_FLASH_ATTN_EXT_TENSOR_NQPSG_LARGE : OP_FLASH_ATTN_EXT_TENSOR_NQPSG;
+
+ // few heads and small batches do not fill the GPU - the half8x8 kernel is faster there
+ // TODO: tune per device
+ if (((ne01 + nqptg - 1)/nqptg)*ne02*ne03*dk < 8192) {
+ return false;
+ }
+
+ if (op->src[1]->type != GGML_TYPE_F16 && !ggml_metal_op_flash_attn_ext_use_kv_f16(op)) {
+ return false;
+ }
+
+ return op->src[1]->ne[1] % OP_FLASH_ATTN_EXT_TENSOR_NCPSG == 0;
+}
+
// returns the n_kv_max hint if the sparse path is available for this op, or 0 otherwise
// the mask (src[3]) remains the single source of truth: finite entries are the valid KV positions,
// n_kv_max is only an upper bound on their number per mask row, used to size the index lists
@@ -3418,7 +3459,101 @@ int ggml_metal_op_flash_attn_ext(ggml_metal_op_t ctx, int idx) {
}
}
- if (!use_sparse && !ggml_metal_op_flash_attn_ext_use_vec(op)) {
+ if (!use_sparse && ggml_metal_op_flash_attn_ext_use_tensor(op, props_dev->has_tensor)) {
+ // tensor API kernel
+ const int nqptg = ne00 >= 512 ? OP_FLASH_ATTN_EXT_TENSOR_NQPSG_LARGE : OP_FLASH_ATTN_EXT_TENSOR_NQPSG; // queries per threadgroup
+ const int ncpsg = OP_FLASH_ATTN_EXT_TENSOR_NCPSG; // cache values per threadgroup
+ const int nsg = OP_FLASH_ATTN_EXT_TENSOR_NSG;
+
+ if (has_mask) {
+ assert(ggml_metal_op_flash_attn_ext_extra_blk(op) != 0);
+
+ ggml_metal_kargs_flash_attn_ext_blk args0 = {
+ /*.ne01 =*/ ne01,
+ /*.ne30 =*/ ne30,
+ /*.ne31 =*/ ne31,
+ /*.ne32 =*/ ne32,
+ /*.ne33 =*/ ne33,
+ /*.nb31 =*/ nb31,
+ /*.nb32 =*/ nb32,
+ /*.nb33 =*/ nb33,
+ };
+
+ auto pipeline0 = ggml_metal_library_get_pipeline_flash_attn_ext_blk(lib, op, nqptg, ncpsg);
+
+ ggml_metal_encoder_set_pipeline(enc, pipeline0);
+ ggml_metal_encoder_set_bytes (enc, &args0, sizeof(args0), 0);
+ ggml_metal_encoder_set_buffer (enc, bid_src3, 1);
+ ggml_metal_encoder_set_buffer (enc, bid_blk, 2);
+
+ const int32_t nblk1 = ((ne01 + nqptg - 1)/nqptg);
+ const int32_t nblk0 = ((ne30 + ncpsg - 1)/ncpsg);
+
+ ggml_metal_encoder_dispatch_threadgroups(enc, nblk0, nblk1, ne32*ne33, 32, 1, 1);
+
+ ggml_metal_op_concurrency_reset(ctx);
+ }
+
+ const int32_t ns10 = nb11_attn/nb10_attn;
+ const int32_t ns20 = nb21_attn/nb20_attn;
+
+ ggml_metal_kargs_flash_attn_ext args = {
+ /*.ne01 =*/ ne01,
+ /*.ne02 =*/ ne02,
+ /*.ne03 =*/ ne03,
+ /*.nb01 =*/ nb01,
+ /*.nb02 =*/ nb02,
+ /*.nb03 =*/ nb03,
+ /*.ne11 =*/ ne11,
+ /*.ne_12_2 =*/ ne12,
+ /*.ne_12_3 =*/ ne13,
+ /*.ns10 =*/ ns10,
+ /*.nb11 =*/ nb11_attn,
+ /*.nb12 =*/ nb12_attn,
+ /*.nb13 =*/ nb13_attn,
+ /*.ns20 =*/ ns20,
+ /*.nb21 =*/ nb21_attn,
+ /*.nb22 =*/ nb22_attn,
+ /*.nb23 =*/ nb23_attn,
+ /*.ne31 =*/ ne31,
+ /*.ne32 =*/ ne32,
+ /*.ne33 =*/ ne33,
+ /*.nb31 =*/ nb31,
+ /*.nb32 =*/ nb32,
+ /*.nb33 =*/ nb33,
+ /*.ne1 =*/ ne1,
+ /*.ne2 =*/ ne2,
+ /*.ne3 =*/ ne3,
+ /*.scale =*/ scale,
+ /*.max_bias =*/ max_bias,
+ /*.m0 =*/ m0,
+ /*.m1 =*/ m1,
+ /*.n_head_log2 =*/ n_head_log2,
+ /*.logit_softcap =*/ logit_softcap,
+ };
+
+ // shared memory layout: queries (half), scores (float), probabilities (half), row scale (float), rescale flag (int)
+ const size_t smem = GGML_PAD(nqptg*ne00*sizeof(ggml_fp16_t) + nqptg*ncpsg*(sizeof(float) + sizeof(ggml_fp16_t)) + nqptg*sizeof(float) + sizeof(int32_t), 16);
+
+ auto pipeline = ggml_metal_library_get_pipeline_flash_attn_ext_tensor(lib, op, has_mask, has_sinks, has_bias, has_scap);
+
+ GGML_ASSERT(nsg*32 <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
+ GGML_ASSERT(smem <= props_dev->max_theadgroup_memory_size);
+
+ ggml_metal_encoder_set_pipeline(enc, pipeline);
+ ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
+ ggml_metal_encoder_set_buffer (enc, bid_src0, 1);
+ ggml_metal_encoder_set_buffer (enc, bid_k, 2);
+ ggml_metal_encoder_set_buffer (enc, bid_v, 3);
+ ggml_metal_encoder_set_buffer (enc, bid_src3, 4);
+ ggml_metal_encoder_set_buffer (enc, bid_src4, 5);
+ ggml_metal_encoder_set_buffer (enc, bid_blk, 6);
+ ggml_metal_encoder_set_buffer (enc, bid_dst, 7);
+
+ ggml_metal_encoder_set_threadgroup_memory_size(enc, smem, 0);
+
+ ggml_metal_encoder_dispatch_threadgroups(enc, (ne01 + nqptg - 1)/nqptg, ne02, ne03, 32, nsg, 1);
+ } else if (!use_sparse && !ggml_metal_op_flash_attn_ext_use_vec(op)) {
// half8x8 kernel
const int nqptg = OP_FLASH_ATTN_EXT_NQPSG; // queries per threadgroup
const int ncpsg = OP_FLASH_ATTN_EXT_NCPSG; // cache values per simdgroup
diff --git a/ggml/src/ggml-metal/kernels/fa_f16.metal b/ggml/src/ggml-metal/kernels/fa_f16.metal
index f46eb2cd1..2dfd3799d 100644
--- a/ggml/src/ggml-metal/kernels/fa_f16.metal
+++ b/ggml/src/ggml-metal/kernels/fa_f16.metal
@@ -73,3 +73,291 @@ template [[host_name("kernel_flash_attn_ext_bf16_dk576_dv512")]] kernel flash_at
#undef FA_TYPES
#undef FA_TYPES_BF
#undef FA_TYPES_F32
+
+#ifdef GGML_METAL_HAS_TENSOR
+
+constant bool FC_flash_attn_ext_tensor_has_mask [[function_constant(FC_FLASH_ATTN_EXT_TENSOR + 0)]];
+constant bool FC_flash_attn_ext_tensor_has_sinks [[function_constant(FC_FLASH_ATTN_EXT_TENSOR + 1)]];
+constant bool FC_flash_attn_ext_tensor_has_bias [[function_constant(FC_FLASH_ATTN_EXT_TENSOR + 2)]];
+constant bool FC_flash_attn_ext_tensor_has_scap [[function_constant(FC_FLASH_ATTN_EXT_TENSOR + 3)]];
+
+// ref: https://arxiv.org/pdf/2307.08691.pdf
+template<
+ short DK, // K head size
+ short DV, // V head size
+ short Q = OP_FLASH_ATTN_EXT_TENSOR_NQPSG, // queries per threadgroup
+ short C = OP_FLASH_ATTN_EXT_TENSOR_NCPSG, // cache items per threadgroup
+ short NSG = OP_FLASH_ATTN_EXT_TENSOR_NSG> // number of simd groups
+kernel void kernel_flash_attn_ext_tensor(
+ constant ggml_metal_kargs_flash_attn_ext & args,
+ device const char * q,
+ device const char * k,
+ device const char * v,
+ device const char * mask,
+ device const char * sinks,
+ device const char * blk,
+ device char * dst,
+ threadgroup char * shmem [[threadgroup(0)]],
+ uint3 tgpig [[threadgroup_position_in_grid]],
+ ushort tiisg [[thread_index_in_simdgroup]],
+ ushort sgitg [[simdgroup_index_in_threadgroup]]) {
+ constexpr short NW = N_SIMDWIDTH;
+ constexpr short NT = NW*NSG;
+ constexpr short NQ = Q/NSG;
+ constexpr short NC = C/NW; // columns per thread
+
+ static_assert(DK % 4 == 0, "DK must be divisible by 4");
+ static_assert(Q % NSG == 0, "Q must be divisible by NSG");
+ static_assert(C % NW == 0, "C must be divisible by NW");
+
+ const int iq3 = tgpig[2];
+ const int iq2 = tgpig[1];
+ const int iq1 = tgpig[0]*Q;
+
+ const short tiitg = sgitg*NW + tiisg;
+
+ threadgroup half * sq = (threadgroup half *) shmem; // [Q, DK] queries
+ threadgroup float * ss = (threadgroup float *) (sq + Q*DK); // [Q, C] scores
+ threadgroup half * sp = (threadgroup half *) (ss + Q*C); // [Q, C] probabilities
+ threadgroup float * sr = (threadgroup float *) (sp + Q*C); // [Q] per-row scale of O
+ threadgroup int * sf = (threadgroup int *) (sr + Q); // [1] last iteration (ic0 + 1) that rescaled O
+
+ q += iq1*args.nb01 + iq2*args.nb02 + iq3*args.nb03;
+
+ {
+ const int ikv2 = iq2/(args.ne02/args.ne_12_2);
+ const int ikv3 = iq3/(args.ne03/args.ne_12_3);
+
+ k += ikv2*args.nb12 + ikv3*args.nb13;
+ v += ikv2*args.nb22 + ikv3*args.nb23;
+ }
+
+ // with softcap the scale is small (scale/softcap), so it is applied to the scores to keep the precision of Q
+ const float qscale = FC_flash_attn_ext_tensor_has_scap ? 1.0f : args.scale;
+
+ // load the queries, with the scale folded in
+ for (int i = tiitg; i < Q*DK/4; i += NT) {
+ const int j = i/(DK/4);
+
+ float4 q4 = 0.0f;
+ if (iq1 + j < args.ne01) {
+ q4 = ((device const float4 *) (q + j*args.nb01))[i%(DK/4)];
+ }
+
+ ((threadgroup half4 *) sq)[i] = (half4) (q4*qscale);
+ }
+
+ device const half * pm[NQ];
+
+ FOR_UNROLL (short jj = 0; jj < NQ; ++jj) {
+ const short j = jj*NSG + sgitg;
+
+ pm[jj] = (device const half *) (mask + (iq1 + j)*args.nb31 + (iq2%args.ne32)*args.nb32 + (iq3%args.ne33)*args.nb33);
+ }
+
+ {
+ const int nblk1 = (args.ne01 + Q - 1)/Q;
+ const int nblk0 = (args.ne11 + C - 1)/C;
+
+ blk += (((iq3%args.ne33)*args.ne32 + (iq2%args.ne32))*nblk1 + iq1/Q)*nblk0;
+ }
+
+ float M[NQ];
+ float S[NQ];
+
+ FOR_UNROLL (short jj = 0; jj < NQ; ++jj) {
+ M[jj] = -FLT_MAX/2;
+ S[jj] = 0.0f;
+ }
+
+ float slope = 1.0f;
+
+ // ALiBi
+ if (FC_flash_attn_ext_tensor_has_bias) {
+ const short h = iq2;
+
+ const float base = h < args.n_head_log2 ? args.m0 : args.m1;
+ const short exph = h < args.n_head_log2 ? h + 1 : 2*(h - args.n_head_log2) + 1;
+
+ slope = pow(base, exph);
+ }
+
+ const int sk = args.ns10;
+ const int sv = args.ns20;
+
+ auto tq = tensor(sq, dextents<int32_t, 2>(DK, Q));
+ auto ts = tensor(ss, dextents<int32_t, 2>(C, Q));
+ auto tp = tensor(sp, dextents<int32_t, 2>(C, Q));
+
+ mpp::tensor_ops::matmul2d<
+ mpp::tensor_ops::matmul2d_descriptor(Q, C, DK, false, true, false, mpp::tensor_ops::matmul2d_descriptor::mode::multiply),
+ execution_simdgroups<NSG>> mm_qk;
+
+ mpp::tensor_ops::matmul2d<
+ mpp::tensor_ops::matmul2d_descriptor(Q, DV, C, false, false, false, mpp::tensor_ops::matmul2d_descriptor::mode::multiply_accumulate),
+ execution_simdgroups<NSG>> mm_pv;
+
+ auto tv0 = tensor((device half *) v, dextents<int32_t, 2>(DV, C), array<int, 2>({1, sv}));
+
+ // the O matrix from the paper
+ auto co = mm_pv.template get_destination_cooperative_tensor<decltype(tp), decltype(tv0), float>();
+
+ FOR_UNROLL (short i = 0; i < co.get_capacity(); ++i) {
+ if (co.is_valid_element(i)) {
+ co[i] = 0.0f;
+ }
+ }
+
+ if (tiitg == 0) {
+ sf[0] = 0;
+ }
+
+ threadgroup_barrier(mem_flags::mem_threadgroup);
+
+ // the host guarantees ne11 % C == 0
+ for (int ic0 = 0, ic = 0; ic < args.ne11; ++ic0, ic += C) {
+ char blk_cur = 1;
+
+ if (FC_flash_attn_ext_tensor_has_mask) {
+ blk_cur = blk[ic0];
+
+ if (blk_cur == 0) {
+ continue;
+ }
+ }
+
+ // Q*K^T
+ {
+ auto tk = tensor((device half *) (k + (uint64_t) ic*args.nb11), dextents<int32_t, 2>(DK, C), array<int, 2>({1, sk}));
+
+ mm_qk.run(tq, tk, ts);
+ }
+
+ threadgroup_barrier(mem_flags::mem_threadgroup);
+
+ // online softmax
+ FOR_UNROLL (short jj = 0; jj < NQ; ++jj) {
+ const short j = jj*NSG + sgitg;
+
+ float s[NC];
+
+ FOR_UNROLL (short ii = 0; ii < NC; ++ii) {
+ s[ii] = ss[j*C + ii*NW + tiisg];
+ }
+
+ if (FC_flash_attn_ext_tensor_has_scap) {
+ FOR_UNROLL (short ii = 0; ii < NC; ++ii) {
+ s[ii] = args.logit_softcap*precise::tanh(s[ii]*args.scale);
+ }
+ }
+
+ if (FC_flash_attn_ext_tensor_has_mask && blk_cur != 2 && iq1 + j < args.ne31) {
+ FOR_UNROLL (short ii = 0; ii < NC; ++ii) {
+ s[ii] += slope*(float) pm[jj][ic + ii*NW + tiisg];
+ }
+ }
+
+ float m = M[jj];
+
+ FOR_UNROLL (short ii = 0; ii < NC; ++ii) {
+ m = max(m, s[ii]);
+ }
+
+ m = simd_max(m);
+
+ // lazy rescaling: move the running max only when it grows by more than 8 (e^8 fits in half)
+ float ms = 1.0f;
+
+ if (m > M[jj] + 8.0f) {
+ ms = exp(M[jj] - m);
+ M[jj] = m;
+
+ if (tiisg == 0) {
+ sf[0] = ic0 + 1;
+ }
+ }
+
+ float sum = 0.0f;
+
+ FOR_UNROLL (short ii = 0; ii < NC; ++ii) {
+ // the sum uses the same rounded values as P*V
+ const half p = (half) exp(s[ii] - M[jj]);
+
+ sp[j*C + ii*NW + tiisg] = p;
+
+ sum += (float) p;
+ }
+
+ S[jj] = S[jj]*ms + simd_sum(sum);
+
+ if (tiisg == 0) {
+ sr[j] = ms;
+ }
+ }
+
+ threadgroup_barrier(mem_flags::mem_threadgroup);
+
+ // O = diag(ms)*O + P*V
+ if (sf[0] == ic0 + 1) {
+ FOR_UNROLL (short i = 0; i < co.get_capacity(); ++i) {
+ if (co.is_valid_element(i)) {
+ co[i] *= sr[co.get_multidimensional_index(i)[1]];
+ }
+ }
+ }
+
+ {
+ auto tv = tensor((device half *) (v + (uint64_t) ic*args.nb21), dextents<int32_t, 2>(DV, C), array<int, 2>({1, sv}));
+
+ mm_pv.run(tp, tv, co);
+ }
+
+ threadgroup_barrier(mem_flags::mem_threadgroup);
+ }
+
+ FOR_UNROLL (short jj = 0; jj < NQ; ++jj) {
+ const short j = jj*NSG + sgitg;
+
+ // the sink only adds to the denominator - its rescale of O is folded into the final scale
+ float ms = 1.0f;
+
+ if (FC_flash_attn_ext_tensor_has_sinks) {
+ const float s = ((device const float *) sinks)[iq2];
+ const float m = max(M[jj], s);
+
+ ms = exp(M[jj] - m);
+
+ S[jj] = S[jj]*ms + exp(s - m);
+ }
+
+ if (tiisg == 0) {
+ sr[j] = S[jj] == 0.0f ? 0.0f : ms/S[jj];
+ }
+ }
+
+ threadgroup_barrier(mem_flags::mem_threadgroup);
+
+ FOR_UNROLL (short i = 0; i < co.get_capacity(); ++i) {
+ if (co.is_valid_element(i)) {
+ co[i] *= sr[co.get_multidimensional_index(i)[1]];
+ }
+ }
+
+ // store to global memory - rows past ne01 are clipped by the tensor extents
+ device float * pdst = (device float *) dst + ((uint64_t) iq3*args.ne2*args.ne1 + iq2 + (uint64_t) iq1*args.ne1)*DV;
+
+ auto td = tensor(pdst, dextents<int32_t, 2>(DV, args.ne01 - iq1), array<int, 2>({1, args.ne1*DV}));
+
+ co.store(td);
+}
+
+typedef decltype(kernel_flash_attn_ext_tensor<64, 64>) flash_attn_ext_tensor_t;
+
+template [[host_name("kernel_flash_attn_ext_tensor_f16_dk64_dv64" )]] kernel flash_attn_ext_tensor_t kernel_flash_attn_ext_tensor<64, 64>;
+template [[host_name("kernel_flash_attn_ext_tensor_f16_dk128_dv128")]] kernel flash_attn_ext_tensor_t kernel_flash_attn_ext_tensor<128, 128>;
+template [[host_name("kernel_flash_attn_ext_tensor_f16_dk192_dv128")]] kernel flash_attn_ext_tensor_t kernel_flash_attn_ext_tensor<192, 128>;
+template [[host_name("kernel_flash_attn_ext_tensor_f16_dk256_dv256")]] kernel flash_attn_ext_tensor_t kernel_flash_attn_ext_tensor<256, 256>;
+template [[host_name("kernel_flash_attn_ext_tensor_f16_dk512_dv512")]] kernel flash_attn_ext_tensor_t kernel_flash_attn_ext_tensor<512, 512, OP_FLASH_ATTN_EXT_TENSOR_NQPSG_LARGE>;
+template [[host_name("kernel_flash_attn_ext_tensor_f16_dk576_dv512")]] kernel flash_attn_ext_tensor_t kernel_flash_attn_ext_tensor<576, 512, OP_FLASH_ATTN_EXT_TENSOR_NQPSG_LARGE>;
+
+#endif // GGML_METAL_HAS_TENSOR
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 8bd4e9422..2682d9c06 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -8075,6 +8075,22 @@ struct test_flash_attn_ext : public test_case {
}
};
+// large Q values, so the online softmax has to rescale the partial results
+struct test_flash_attn_ext_large_logits : public test_flash_attn_ext {
+ static constexpr int q_range = 20;
+
+ using test_flash_attn_ext::test_flash_attn_ext;
+
+ std::string vars() override {
+ return test_flash_attn_ext::vars() + ",q_range=" + std::to_string(q_range);
+ }
+
+ void initialize_tensors(ggml_context * ctx) override {
+ test_flash_attn_ext::initialize_tensors(ctx);
+ init_tensor_uniform(ggml_get_tensor(ctx, "q"), -(float) q_range, (float) q_range);
+ }
+};
+
// GGML_OP_CROSS_ENTROPY_LOSS
struct test_cross_entropy_loss : public test_case {
const ggml_type type;
@@ -11192,6 +11208,12 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {2, 1}, 1024, 32, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 4, {2, 1}, 1024, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+ // FLASH_ATTN_EXT: large logits
+ test_cases.emplace_back(new test_flash_attn_ext_large_logits( 64, 64, 16, {4, 1}, 1024, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+ test_cases.emplace_back(new test_flash_attn_ext_large_logits(128, 128, 8, {4, 1}, 1024, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+ test_cases.emplace_back(new test_flash_attn_ext_large_logits(256, 256, 4, {4, 1}, 1024, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+ test_cases.emplace_back(new test_flash_attn_ext_large_logits(256, 256, 4, {4, 1}, 1024, 75, true, false, 0, 10.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
+
test_cases.emplace_back(new test_cross_entropy_loss (GGML_TYPE_F32, { 10, 5, 4, 3}));
test_cases.emplace_back(new test_cross_entropy_loss (GGML_TYPE_F32, {30000, 1, 1, 1}));
test_cases.emplace_back(new test_cross_entropy_loss_back(GGML_TYPE_F32, { 10, 5, 4, 3}));