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}));