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]; -#pragma unroll - 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. + if (threadIdx.y % np == 0) { + // Combine the meta data for parallel warps via shared memory. + float2 meta[nmeta]; #pragma unroll - 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 imeta = 0; imeta < nmeta; ++imeta) { + meta[imeta] = meta_ptr[imeta * warp_size * tile_stride/2]; } - } - float KQ_cms[nmeta]; // KQ combine max scale per warp. + KQ_cmn = meta[0].x; // KQ combine max new, max between all parallel warps. #pragma unroll - for (int imeta = 0; imeta < nmeta; ++imeta) { - KQ_cms[imeta] = expf(meta[imeta].x - KQ_cmn); - } + 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)); + } + } - float 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 = 0; imeta < nmeta; ++imeta) { + KQ_cms[imeta] = expf(meta[imeta].x - KQ_cmn); + } + + KQ_crs = KQ_cms[0]*meta[0].y; // KQ combine rowsum, scaled sum of all parallel warps. #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 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); + } } } __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