Subject: [PATCH] vulkan : fa cm1: shift the last partial KV block instead of staging every block

Assisted-by: Claude Code (Claude Opus 5.5)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-common.h b/ggml/src/ggml-vulkan/ggml-vulkan-common.h
index 4ae5fea..e087f82 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-common.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-common.h
@@ -31,7 +31,7 @@ bool ggml_vk_intel_windows_driver_in_range(uint32_t driver_version, uint32_t low
 // shaders
 void ggml_vk_destroy_pipeline(vk::Device& device, vk_pipeline& pipeline);
 vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_t hsk, uint32_t hsv, uint32_t n_rows, uint32_t n_kv, ggml_type k_type, ggml_type v_type, bool f32acc);
-vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc, bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type);
+vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc, bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type, bool use_kv_tail_shift = false);
 uint32_t get_subgroup_size(const std::string &pipeline_name, const vk_device_architecture &arch);
 void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested = nullptr);
 bool ggml_vk_flash_attn_scalar_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc, ggml_type k_type, ggml_type v_type);
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index fe11023..e90ed0c 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -1365,7 +1365,8 @@ vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_t hsk,
 }
 
 vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc,
-                                                  bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type) {
+                                                  bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type,
+                                                  bool use_kv_tail_shift) {
     const bool old_amd_windows = device->vendor_id == VK_VENDOR_ID_AMD && device->driver_id == vk::DriverId::eAmdProprietary &&
                                  (device->architecture == AMD_GCN || device->architecture == AMD_RDNA1 || device->architecture == AMD_RDNA2);
 
@@ -1373,7 +1374,8 @@ vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_
                      (use_mask          ? 2 : 0) |
                      (use_logit_softcap ? 4 : 0) |
                      (old_amd_windows   ? 8 : 0) |
-                     (use_sparse        ? 16 : 0);
+                     (use_sparse        ? 16 : 0) |
+                     (use_kv_tail_shift ? 32 : 0);
 
     const uint32_t subgroup_size = params.disable_subgroups ? 0 : params.subgroup_size;
 
@@ -8144,9 +8146,9 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
         nbv3_eff = (uint32_t)((uint64_t)HSV * KV * nev2 * sizeof(ggml_fp16_t));
     }
     const uint32_t alignment = tuning_params.block_cols;
-    bool aligned = (KV % alignment) == 0 &&
-                   // the "aligned" shader variant will forcibly align strides, for performance
-                   (q_stride & 7) == 0 && (k_stride & 7) == 0 && (v_stride & 7) == 0;
+    const bool strides_aligned = (q_stride & 7) == 0 && (k_stride & 7) == 0 && (v_stride & 7) == 0;
+    // the "aligned" shader variant will forcibly align strides, for performance
+    bool aligned = (KV % alignment) == 0 && strides_aligned;
 
     // Need to use the coopmat2 variant that clamps loads when HSK/HSV aren't sufficiently aligned.
     if (((HSK | HSV) % 16) != 0 && tuning_params.path == FA_COOPMAT2) {
@@ -8156,8 +8158,15 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
     // Only use mask opt when the mask is fairly large. This hasn't been tuned extensively.
     bool use_mask_opt = mask && !use_sparse && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16
                         && (ctx->device->architecture != vk_device_architecture::AMD_GCN || HSK > 256 || HSV > 256);
+    // KV not a multiple of Bc: coopmat1 loads the last block as chunks that end at KV instead of staging every block
+    const bool use_kv_tail_shift = !aligned && strides_aligned && KV >= 16 &&
+                                   tuning_params.path == FA_COOPMAT1 && !tuning_params.shmem_staging &&
+                                   !use_sparse &&
+                                   k_type_eff == GGML_TYPE_F16 && v_type_eff == GGML_TYPE_F16 &&
+                                   HSK % 16 == 0 && HSV % 16 == 0;
     vk_fa_pipeline_state fa_pipeline_state = get_fa_pipeline_state(ctx->device, tuning_params, HSK, HSV, aligned, f32acc,
-                                                                   mask != nullptr, use_mask_opt, logit_softcap != 0, use_sparse, k_type_eff, v_type_eff);
+                                                                   mask != nullptr, use_mask_opt, logit_softcap != 0, use_sparse, k_type_eff, v_type_eff,
+                                                                   use_kv_tail_shift);
 
     vk_pipeline pipeline = nullptr;
 
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
index 2e0e23b..1f16e18 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
@@ -26,6 +26,8 @@ const bool LOGIT_SOFTCAP   = (Flags & 4) != 0;
 const bool OLD_AMD_WINDOWS = (Flags & 8) != 0;
 // Sparse: gather binding-7 indices instead of scanning [0,KV); p.split_kv = n_kv_max.
 const bool USE_SPARSE      = (Flags & 16) != 0;
+// KV not a multiple of Bc: cm1 loads the last block as chunks that end at KV instead of staging it.
+const bool KV_TAIL_SHIFT   = (Flags & 32) != 0;
 
 // Round up head sizes to a multiple of 16, for coopmat1/coopmat2 paths
 const uint32_t HSK_pad = (HSK + 15) & ~15;
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
index aa9dd62..982b63b 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
@@ -62,6 +62,37 @@ shared O_TYPEV4 pvsh[MatBc * osh_stride];
 
 shared ACC_TYPE slope[Br];
 
+// First KV row of 16-row chunk ch in block j. With shift, chunks of the last block end at KV at the latest.
+uint kv_chunk_row(uint j, uint ch, bool shift) {
+    const uint row = j * Bc + ch * MatBc;
+    return shift ? min(row, KV - MatBc) : row;
+}
+
+// KV row held by column col of block j.
+uint kv_col_row(uint j, uint col, bool shift) {
+    return kv_chunk_row(j, col / MatBc, shift) + col % MatBc;
+}
+
+// Whether column col of block j holds a key that is in range and not repeated from an earlier chunk.
+bool kv_col_valid(uint j, uint col, bool shift) {
+    if (shift) {
+        const uint ch = col / MatBc;
+        return col % MatBc >= j * Bc + ch * MatBc - kv_chunk_row(j, ch, true);
+    }
+    return j * Bc + col < KV;
+}
+
+// fa_kv_index with the tail-block shift applied; identical to fa_kv_index when KV_TAIL_SHIFT is off.
+bool fa_kv_index_ts(uint lin, out uint kv_col) {
+    if (KV_TAIL_SHIFT) {
+        const uint j = lin / Bc;
+        const uint col = lin % Bc;
+        kv_col = kv_col_row(j, col, true);
+        return kv_col_valid(j, col, true);
+    }
+    return fa_kv_index(lin, kv_col);
+}
+
 void main() {
 #ifdef NEEDS_INIT_IQ_SHMEM
     if (fa_type_needs_shmem(FaTypeK) || fa_type_needs_shmem(FaTypeV)) {
@@ -152,6 +183,10 @@ void main() {
 
     [[dont_unroll]]
     for (uint32_t j = start_j; j < end_j; ++j) {
+        // The tail block (KV not a multiple of Bc) loads shifted chunks ending at KV;
+        // all other blocks are fully in range and need no bounds check.
+        const bool kv_bounds = KV_bounds_check && (!KV_TAIL_SHIFT || (j + 1) * Bc > KV);
+        const bool kv_shift = KV_TAIL_SHIFT && kv_bounds;
 
         [[unroll]] for (uint32_t idx = 0; idx < mask_cache.length(); ++idx) {
             mask_cache[idx] = f16vec4(0);
@@ -177,7 +212,7 @@ void main() {
                     uint32_t r = (idx + tid) % (Br / 4);
                     if (idx + tid < Bc * Br / 4 || idx + gl_WorkGroupSize.x <= Bc * Br / 4) {
                         uint32_t kcol;
-                        bool kv_active = fa_kv_index(j * Bc + c, kcol);
+                        bool kv_active = fa_kv_index_ts(j * Bc + c, kcol);
                         if (kv_active) {
                             f16vec4 m;
                             if (USE_SPARSE) {
@@ -186,25 +221,25 @@ void main() {
                                 m = f16vec4(mv);
                                 max_mask = max(max_mask, float(mv));
                             } else if (!nem1_bounds_check || i * Br + r * 4 + 3 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 2) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 3) * m_stride + (j * Bc + c)]);
+                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + kcol],
+                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + kcol],
+                                            data_m[m_offset + (i * Br + r * 4 + 2) * m_stride + kcol],
+                                            data_m[m_offset + (i * Br + r * 4 + 3) * m_stride + kcol]);
                                 max_mask = max(max(max(max(max_mask, float(m[0])), float(m[1])), float(m[2])), float(m[3]));
                             } else if (i * Br + r * 4 + 2 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 2) * m_stride + (j * Bc + c)],
+                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + kcol],
+                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + kcol],
+                                            data_m[m_offset + (i * Br + r * 4 + 2) * m_stride + kcol],
                                             0.0);
                                 max_mask = max(max(max(max_mask, float(m[0])), float(m[1])), float(m[2]));
                             } else if (i * Br + r * 4 + 1 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + (j * Bc + c)],
+                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + kcol],
+                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + kcol],
                                             0.0,
                                             0.0);
                                 max_mask = max(max(max_mask, float(m[0])), float(m[1]));
                             } else if (i * Br + r * 4 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
+                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + kcol],
                                             0.0,
                                             0.0,
                                             0.0);
@@ -277,7 +312,7 @@ void main() {
             if (SHMEM_STAGING == 0) {
             // For quants we always need to dequant into kvsh; for f16/bf16 we can load
             // directly from global memory when alignment / bounds allow it.
-            const bool stage_k = USE_DECODE_K || KV_bounds_check || USE_SPARSE || d * 16 + 16 > HSK;
+            const bool stage_k = USE_DECODE_K || (kv_bounds && !KV_TAIL_SHIFT) || USE_SPARSE || d * 16 + 16 > HSK;
             if (stage_k) {
                 barrier();
                 [[unroll]] for (uint32_t idx = 0; idx < Bc * MatBr / 4; idx += gl_WorkGroupSize.x) {
@@ -286,7 +321,7 @@ void main() {
                     if (idx + tid < Bc * MatBr / 4) {
                         FLOAT_TYPEV4 K_Tf = FLOAT_TYPEV4(0);
                         uint32_t kcol;
-                        bool kv_active = fa_kv_index(j * Bc + row, kcol);
+                        bool kv_active = fa_kv_index_ts(j * Bc + row, kcol);
                         if (kv_active && (HSK == HSK_pad || d * 16 + col_vec * 4 < HSK)) {
 #if !defined(BFLOAT16)
                             if (USE_DECODE_K) {
@@ -311,7 +346,7 @@ void main() {
                 uint coord = (gl_SubgroupID * MatBc) * kvsh_stride;
                 coopMatLoad(KMat, kvsh, coord, kvsh_stride, gl_CooperativeMatrixLayoutRowMajor);
             } else {
-                const uint coord = k_offset / 4 + (j * Bc + gl_SubgroupID * MatBc) * k_stride / 4 + d * 16 / 4;
+                const uint coord = k_offset / 4 + kv_chunk_row(j, gl_SubgroupID, kv_shift) * k_stride / 4 + d * 16 / 4;
                 coopMatLoad(KMat, data_kv4, coord, k_stride / 4, gl_CooperativeMatrixLayoutRowMajor);
             }
             } else {
@@ -344,7 +379,7 @@ void main() {
                 uint32_t c = (idx + tid) / (Br / 4);
                 uint32_t r = (idx + tid) % (Br / 4);
                 if (idx + tid < Bc * Br / 4 || idx + gl_WorkGroupSize.x <= Bc * Br / 4) {
-                    if (!KV_bounds_check || j * Bc + c < KV) {
+                    if (!kv_bounds || kv_col_valid(j, c, kv_shift)) {
                         // Mask nem1 bounds check is handled when loading masks
                         ACC_TYPEV4 masks = ACC_TYPEV4(mask_cache[idx / WorkGroupSize]);
                         ACC_TYPEV4 slopes = ACC_TYPEV4(slope[r * 4], slope[r * 4 + 1], slope[r * 4 + 2], slope[r * 4 + 3]);
@@ -362,7 +397,7 @@ void main() {
 
             float rowmaxf = NEG_FLT_MAX_OVER_2;
             [[unroll]] for (uint32_t c = 0; c < cols_per_thread; ++c) {
-                if (KV_bounds_check && j * Bc + c * cols_per_iter + col_tid >= KV) {
+                if (kv_bounds && !kv_col_valid(j, c * cols_per_iter + col_tid, kv_shift)) {
                     continue;
                 }
                 rowmaxf = max(rowmaxf, float(sfsh[r_vec + (c * cols_per_iter + col_tid) * sfshstride][r_comp]));
@@ -395,7 +430,7 @@ void main() {
 
             [[unroll]] for (uint32_t r = 0; r < rows_per_thread; r += 4) {
                 const uint row = tile_row(r);
-                if (KV_bounds_check && j * Bc + col >= KV) {
+                if (kv_bounds && !kv_col_valid(j, col, kv_shift)) {
                     Psh[col * psh_stride + row / 4] = FLOAT_TYPEV4(0.0f);
                 } else {
                     const vec4 mfvec = vec4(Mf[r], Mf[r + 1], Mf[r + 2], Mf[r + 3]);
@@ -456,7 +491,7 @@ void main() {
             if (SHMEM_STAGING == 0) {
             // For quants we always preload via kvsh. For f16/bf16 we only preload when
             // alignment / bounds force it (otherwise we coopMatLoad direct from data_vv4).
-            const bool stage_v = USE_DECODE_V || KV_bounds_check || USE_SPARSE;
+            const bool stage_v = USE_DECODE_V || (kv_bounds && !KV_TAIL_SHIFT) || USE_SPARSE;
             if (stage_v) {
                 [[unroll]] for (uint32_t i = 0; i < v_loads_per_thread; ++i) {
                     const uint idx = i * gl_WorkGroupSize.x + tid;
@@ -464,14 +499,14 @@ void main() {
                     const uint col = idx % v_cols;
 
                     uint32_t v_row;
-                    bool kv_active = fa_kv_index(j * Bc + row, v_row);
+                    bool kv_active = fa_kv_index_ts(j * Bc + row, v_row);
                     const uint v_col = hsv_tile * MatBc * row_split + col * 4;
 
                     const uint coord = v_row * v_stride * BLOCK_SIZE_V + v_col;
                     const uint ib = coord / BLOCK_SIZE_V;
                     const uint iqs = coord % BLOCK_SIZE_V;
 
-                    if (USE_SPARSE ? (kv_active && v_col < HSV) : (!KV_bounds_check || (v_row < KV && v_col < HSV))) {
+                    if (USE_SPARSE ? (kv_active && v_col < HSV) : (!kv_bounds || (v_row < KV && v_col < HSV))) {
 #if !defined(BFLOAT16)
                         if (USE_DECODE_V) {
                             kvsh[row * vsh_stride + col] = dequantize4(ib, iqs, v_offset, BINDING_IDX_V);
@@ -495,9 +530,9 @@ void main() {
                     coopMatLoad(KMat, Psh, bc_chunk * MatBc * psh_stride, psh_stride, gl_CooperativeMatrixLayoutColumnMajor);
 
                     if (SHMEM_STAGING == 0) {
-                    if (!USE_DECODE_V && !KV_bounds_check && !USE_SPARSE) {
+                    if (!USE_DECODE_V && !(kv_bounds && !KV_TAIL_SHIFT) && !USE_SPARSE) {
                         // F16/BF16 values can be loaded directly from global memory
-                        const uint v_tile_row = j * Bc + bc_chunk * MatBc;
+                        const uint v_tile_row = kv_chunk_row(j, bc_chunk, kv_shift);
                         const uint v_tile_offset = v_offset / 4 + v_tile_row * v_stride / 4 + hsv_offset / 4;
                         coopMatLoad(QMat, data_vv4, v_tile_offset, v_stride / 4, gl_CooperativeMatrixLayoutRowMajor);
                     } else {
