Commit b74f590ea for llama.cpp
commit b74f590eafec2fafc6e0e98ee93b2e5d3efa9042
Author: Siavash Norouzi <35790025+siavashnorouzi@users.noreply.github.com>
Date: Sun Sep 6 23:23:21 2026 -0700
ggml-cuda: fix divergent barrier in f16 flash attention (#27870)
* ggml-cuda: fix divergent barrier in f16 flash attention
* ggml-cuda: avoid duplicate metadata pointer setup
diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh
index 126a4c452..bc5060e81 100644
--- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh
+++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh
@@ -1545,77 +1545,77 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
}
}
- if (np > 1 && threadIdx.y % np == 0) {
- // Combine the meta data for parallel warps via shared memory.
- // Warps with threadIdx.y % np != 0 must NOT return early.
- // All threads must return simultaneously to avoid race conditions with work on the next tile.
-
+ if (np > 1) {
constexpr int nmeta = np*cols_per_warp >= warp_size ? np*cols_per_warp/warp_size : 1;
+ float KQ_cmn;
+ float KQ_cms[nmeta];
+ float KQ_crs;
+
const int jc_meta = threadIdx.y*cols_per_warp + (np*cols_per_warp < warp_size ? threadIdx.x % (np*cols_per_warp) : threadIdx.x);
float2 * const meta_ptr = ((float2 *) tile_Q) + jc_meta*(tile_stride/2) + nbatch_combine/2;
- float2 meta[nmeta];
+
+ if (threadIdx.y % np == 0) {
+ // Combine the meta data for parallel warps via shared memory.
+ float2 meta[nmeta];
#pragma unroll
- for (int imeta = 0; imeta < nmeta; ++imeta) {
- meta[imeta] = meta_ptr[imeta * warp_size * tile_stride/2];
- }
+ for (int imeta = 0; imeta < nmeta; ++imeta) {
+ meta[imeta] = meta_ptr[imeta * warp_size * tile_stride/2];
+ }
- float KQ_cmn = meta[0].x; // KQ combine max new, max between all parallel warps.
+ KQ_cmn = meta[0].x; // KQ combine max new, max between all parallel warps.
#pragma unroll
- for (int imeta = 1; imeta < nmeta; ++imeta) {
- KQ_cmn = fmaxf(KQ_cmn, meta[imeta].x);
- }
+ for (int imeta = 1; imeta < nmeta; ++imeta) {
+ KQ_cmn = fmaxf(KQ_cmn, meta[imeta].x);
+ }
#pragma unroll
- for (int offset = np*cols_per_warp/2; offset >= cols_per_warp; offset >>= 1) {
- if (offset < warp_size) {
- KQ_cmn = fmaxf(KQ_cmn, __shfl_xor_sync(0xFFFFFFFF, KQ_cmn, offset, warp_size));
+ for (int offset = np*cols_per_warp/2; offset >= cols_per_warp; offset >>= 1) {
+ if (offset < warp_size) {
+ KQ_cmn = fmaxf(KQ_cmn, __shfl_xor_sync(0xFFFFFFFF, KQ_cmn, offset, warp_size));
+ }
}
- }
- float KQ_cms[nmeta]; // KQ combine max scale per warp.
#pragma unroll
- for (int imeta = 0; imeta < nmeta; ++imeta) {
- KQ_cms[imeta] = expf(meta[imeta].x - KQ_cmn);
- }
+ for (int imeta = 0; imeta < nmeta; ++imeta) {
+ KQ_cms[imeta] = expf(meta[imeta].x - KQ_cmn);
+ }
- float KQ_crs = KQ_cms[0]*meta[0].y; // KQ combine rowsum, scaled sum of all parallel warps.
+ KQ_crs = KQ_cms[0]*meta[0].y; // KQ combine rowsum, scaled sum of all parallel warps.
#pragma unroll
- for (int imeta = 1; imeta < nmeta; ++imeta) {
- KQ_crs += KQ_cms[imeta]*meta[imeta].y;
- }
+ for (int imeta = 1; imeta < nmeta; ++imeta) {
+ KQ_crs += KQ_cms[imeta]*meta[imeta].y;
+ }
#pragma unroll
- for (int offset = np*cols_per_warp/2; offset >= cols_per_warp; offset >>= 1) {
- if (offset < warp_size) {
- KQ_crs += __shfl_xor_sync(0xFFFFFFFF, KQ_crs, offset, warp_size);
+ for (int offset = np*cols_per_warp/2; offset >= cols_per_warp; offset >>= 1) {
+ if (offset < warp_size) {
+ KQ_crs += __shfl_xor_sync(0xFFFFFFFF, KQ_crs, offset, warp_size);
+ }
}
}
__syncthreads();
- // Write back combined meta data:
+ if (threadIdx.y % np == 0) {
+ // Write back combined meta data:
#pragma unroll
- for (int imeta = 0; imeta < nmeta; ++imeta) {
- if (np*cols_per_warp >= warp_size || threadIdx.x < np*cols_per_warp) {
- // Combined KQ max scale + rowsum.
- meta_ptr[imeta * warp_size * tile_stride/2] = make_float2(KQ_cms[imeta], KQ_crs);
+ for (int imeta = 0; imeta < nmeta; ++imeta) {
+ if (np*cols_per_warp >= warp_size || threadIdx.x < np*cols_per_warp) {
+ // Combined KQ max scale + rowsum.
+ meta_ptr[imeta * warp_size * tile_stride/2] = make_float2(KQ_cms[imeta], KQ_crs);
+ }
}
- }
- // Combined KQ max + rowsum.
- static_assert(cols_per_warp <= warp_size);
- if (needs_fixup && (cols_per_warp == warp_size || threadIdx.x < cols_per_warp)) {
- float2 * dstk_fixup_meta = dstk_fixup + blockIdx.x*ncols;
- dstk_fixup_meta[(threadIdx.y/np)*cols_per_warp + threadIdx.x] = make_float2(KQ_cmn, KQ_crs);
- }
- if (is_fixup && (cols_per_warp == warp_size || threadIdx.x < cols_per_warp)) {
- float2 * dstk_fixup_meta = dstk_fixup + (gridDim.x + blockIdx.x)*ncols;
- dstk_fixup_meta[(threadIdx.y/np)*cols_per_warp + threadIdx.x] = make_float2(KQ_cmn, KQ_crs);
+ // Combined KQ max + rowsum.
+ static_assert(cols_per_warp <= warp_size);
+ if (needs_fixup && (cols_per_warp == warp_size || threadIdx.x < cols_per_warp)) {
+ float2 * dstk_fixup_meta = dstk_fixup + blockIdx.x*ncols;
+ dstk_fixup_meta[(threadIdx.y/np)*cols_per_warp + threadIdx.x] = make_float2(KQ_cmn, KQ_crs);
+ }
+ if (is_fixup && (cols_per_warp == warp_size || threadIdx.x < cols_per_warp)) {
+ float2 * dstk_fixup_meta = dstk_fixup + (gridDim.x + blockIdx.x)*ncols;
+ dstk_fixup_meta[(threadIdx.y/np)*cols_per_warp + threadIdx.x] = make_float2(KQ_cmn, KQ_crs);
+ }
}
- } else if (np > 1) {
- // Warps with threadIdx.y % np == 0 execute a __syncthreads() in the if branch.
- // Therefore, all other warps also need to execute a __syncthreads().
- // Otherwise the points at which warps synchronize with each other would become misaligned.
- __syncthreads();
}
#pragma unroll