From e13512ab64cd6560b77133148bd26a8d328b02d9 Mon Sep 17 00:00:00 2001 From: ynankani Date: Mon, 13 Jul 2026 18:42:31 +0000 Subject: [PATCH 01/11] CUDA: XOR swizzle flash attn K,V smem fp16 tiles Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 41 ++++++---- ggml/src/ggml-cuda/fattn-swizzle.cuh | 110 +++++++++++++++++++++++++++ tests/test-backend-ops.cpp | 14 +++- 3 files changed, 146 insertions(+), 19 deletions(-) create mode 100644 ggml/src/ggml-cuda/fattn-swizzle.cuh diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 7f4cfd5511ff..143328e9d4e0 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -2,6 +2,7 @@ #include "cp-async.cuh" #include "mma.cuh" #include "fattn-common.cuh" +#include "fattn-swizzle.cuh" using namespace ggml_cuda_mma; @@ -66,7 +67,7 @@ static constexpr __host__ __device__ fattn_mma_config ggml_cuda_fattn_mma_get_co GGML_CUDA_FATTN_MMA_CONFIG_CASE(192, 128, 32, 128, 2, 32, 96, 64, 64, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(192, 128, 64, 128, 2, 32, 96, 64, 64, 2, true); - GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 8, 64, 4, 64, 128, 128, 128, 2, true); + GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 8, 128, 2, 64, 128, 128, 128, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 16, 64, 4, 32, 128, 128, 128, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 32, 128, 2, 32, 128, 128, 128, 2, true); GGML_CUDA_FATTN_MMA_CONFIG_CASE(256, 256, 64, 128, 2, 32, 128, 128, 128, 2, true); @@ -397,7 +398,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - cp_async_cg_16(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i*stride_KV + k*h2_per_chunk); + const int smem_offs_b = ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk); + cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i*stride_KV + k*h2_per_chunk); } } }; @@ -432,7 +434,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, + ggml_cuda_memcpy_1<16>((char*)tile_KV + ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk), !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); } } @@ -568,9 +570,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols); constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); - constexpr int stride_tile_K = nbatch_K2 + 4; - - constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : nbatch_V2 + 4; + // swizzle the tile stride for K and V based on the batch size. + constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); + constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); const int k_VKQ_0 = kb0 * nbatch_fa; #if defined(TURING_MMA_AVAILABLE) @@ -623,7 +625,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( #pragma unroll for (int k_KQ_0 = k0_start; k_KQ_0 < k0_stop; k_KQ_0 += T_A_KQ::J) { T_A_KQ K_A; - load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix( + K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[k_KQ_0/T_A_KQ::J]); } else { @@ -649,7 +652,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i_KQ_0 = i_KQ_00 + (threadIdx.y % np)*T_A_KQ::I; T_A_KQ K_A; - load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix( + K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[0]); @@ -978,7 +982,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::J; T_A_VKQ A; // Transposed in SRAM but not in registers, gets transposed on load. - load_ldmatrix_trans(A, tile_V_i + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); + const int v_lin = (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; + ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans( + A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); if constexpr (T_B_KQ::I == 8) { mma(VKQ_C[i_VKQ_0/T_A_VKQ::I], A, B[k00/(np*T_A_VKQ::J)]); } else { @@ -1004,7 +1010,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::I; T_A_VKQ A; // Transposed in both SRAM and registers, load normally. - load_ldmatrix(A, tile_V_i + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); + const int v_lin = (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; + ggml_cuda_fattn_smem_swizzle::load_ldmatrix( + A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); mma(VKQ_C[i_VKQ_0/i0_stride], B[k00/(np*T_A_VKQ::I)], A); } } @@ -1168,9 +1176,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( static_assert(nwarps * (cols_per_warp/ncols2) % ncols1 == 0, "bad nwarps"); constexpr int stride_tile_Q = DKQ/2 + 4; - constexpr int stride_tile_K = nbatch_K2 + 4; - - constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : nbatch_V2 + 4; + // swizzle the tile stride for K and V based on the batch size. + constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); + constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); constexpr int stride_tile_KV_max = stride_tile_K > stride_tile_V ? stride_tile_K : stride_tile_V; extern __shared__ half2 tile_Q[]; @@ -1914,8 +1922,11 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml constexpr bool V_is_K_view = DKQ == 576; // Guaranteed by the kernel selection logic in fattn.cu - const size_t nbytes_shared_KV_1stage = nbatch_fa * std::max(nbatch_K2 + 4, nbatch_V2 + 4) * sizeof(half2); - const size_t nbytes_shared_KV_2stage = nbatch_fa * (nbatch_K2 + 4 + nbatch_V2 + 4) * sizeof(half2); + // KV tile strides must match flash_attn_ext_f16_iter / _process_tile. + const int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2, cc); + const int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2, cc); + const size_t nbytes_shared_KV_1stage = nbatch_fa * std::max(stride_tile_K, stride_tile_V) * sizeof(half2); + const size_t nbytes_shared_KV_2stage = nbatch_fa * (stride_tile_K + stride_tile_V) * sizeof(half2); const size_t nbytes_shared_Q = ncols * (DKQ/2 + 4) * sizeof(half2); const size_t nbytes_shared_mask = ncols1 * (nbatch_fa/2 + 4) * sizeof(half2); const size_t nbytes_shared_combine = nwarps*cols_per_warp * (nbatch_combine + 4) * sizeof(half2); diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh new file mode 100644 index 000000000000..85da0b6ef7ee --- /dev/null +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -0,0 +1,110 @@ +#pragma once + +#include "common.cuh" +#include "mma.cuh" + +// XOR swizzle for K/V SMEM tiles to avoid bank conflicts without row padding (Turing+ only). +// Stride must be a power-of-two >= 32 half2 columns,otherwise we keep +4 row padding. + +static __host__ __device__ constexpr bool ggml_cuda_fattn_swz_pow2_stride(const int nbatch_2) { + return nbatch_2 >= 32 && (nbatch_2 & (nbatch_2 - 1)) == 0; +} + +static __device__ constexpr bool ggml_cuda_fattn_swz_enabled(const int nbatch_2) { +#if defined(TURING_MMA_AVAILABLE) + return ggml_cuda_fattn_swz_pow2_stride(nbatch_2); +#else + GGML_UNUSED(nbatch_2); + return false; +#endif +} + +static __host__ bool ggml_cuda_fattn_swz_enabled(const int nbatch_2, const int cc) { +#ifdef GGML_USE_HIP + GGML_UNUSED(nbatch_2); + GGML_UNUSED(cc); + return false; +#else + return turing_mma_available(cc) && ggml_cuda_fattn_swz_pow2_stride(nbatch_2); +#endif +} + +static __device__ constexpr int ggml_cuda_fattn_swz_tile_stride(const int nbatch_2) { + return ggml_cuda_fattn_swz_enabled(nbatch_2) ? nbatch_2 : nbatch_2 + 4; +} + +static __host__ int ggml_cuda_fattn_swz_tile_stride(const int nbatch_2, const int cc) { + return ggml_cuda_fattn_swz_enabled(nbatch_2, cc) ? nbatch_2 : nbatch_2 + 4; +} + +// Swizzled byte offset for tile element (row, col_h2); same map used for writes and reads. +template +static __device__ __forceinline__ int ggml_cuda_fattn_swz_bytes_rc(const int row, const int col_h2) { + int off_bytes = (row * stride_h2 + col_h2) * (int) sizeof(half2); + if constexpr (ggml_cuda_fattn_swz_enabled(stride_h2)) { + off_bytes ^= (row & 7) << 4; + } + return off_bytes; +} + +namespace ggml_cuda_fattn_smem_swizzle { + +// ldmatrix.x4 from a 32-bit .shared address (lower register pressure than 64-bit generic pointers). +static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4(int * xi, const uint32_t saddr) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];" + : "=r"(xi[0]), "=r"(xi[1]), "=r"(xi[2]), "=r"(xi[3]) + : "r"(saddr)); +} +static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4_trans(int * xi, const uint32_t saddr) { + asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.b16 {%0, %1, %2, %3}, [%4];" + : "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3]) + : "r"(saddr)); +} + +// Per-lane swizzled .shared address for tile<16,8> ldmatrix. +template +static __device__ __forceinline__ uint32_t ggml_cuda_fattn_swz_saddr( + const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) { + const int lane_row = threadIdx.x % I; + const int lane_col = (threadIdx.x / I) * (J / 2); + const uint32_t base = __cvta_generic_to_shared(tile_base); + uint32_t byte_off = (uint32_t)((base_row + lane_row) * stride_h2 + base_col_h2 + lane_col) * (uint32_t)sizeof(half2); + if constexpr (ggml_cuda_fattn_swz_enabled(stride_h2)) { + byte_off ^= (uint32_t)(((base_row + lane_row) & 7) << 4); + } + return base + byte_off; +} + +template +static __device__ __forceinline__ void load_ldmatrix( + TileT & t, half2 * tile_base, const int base_row, const int base_col_h2) { + using Tile = typename std::remove_reference::type; + constexpr int I = Tile::I; + constexpr int J = Tile::J; +#if defined(TURING_MMA_AVAILABLE) + if constexpr (I == 16 && J == 8 && ggml_cuda_fattn_swz_enabled(stride_h2)) { + const uint32_t saddr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); + ggml_cuda_fattn_ldmatrix_x4((int *) t.x, saddr); + return; + } +#endif // TURING_MMA_AVAILABLE + ggml_cuda_mma::load_ldmatrix(t, tile_base + base_row * stride_h2 + base_col_h2, stride_h2); +} + +template +static __device__ __forceinline__ void load_ldmatrix_trans( + TileT & t, half2 * tile_base, const int base_row, const int base_col_h2) { + using Tile = typename std::remove_reference::type; + constexpr int I = Tile::I; + constexpr int J = Tile::J; +#if defined(TURING_MMA_AVAILABLE) + if constexpr (I == 16 && J == 8 && ggml_cuda_fattn_swz_enabled(stride_h2)) { + const uint32_t saddr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); + ggml_cuda_fattn_ldmatrix_x4_trans((int *) t.x, saddr); + return; + } +#endif // TURING_MMA_AVAILABLE + ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + base_row * stride_h2 + base_col_h2, stride_h2); +} + +} // namespace ggml_cuda_fattn_smem_swizzle diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 996b88db296e..ba1c7085464e 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10074,6 +10074,10 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 1024, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}, false)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, false)); + // FLASH_ATTN_EXT MMA: non-pow2 head size and MLA K/V view. + test_cases.emplace_back(new test_flash_attn_ext(192, 128, 8, {8, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {20, 1}, 512, 1, true, false, 0, 0, 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})); @@ -10470,10 +10474,12 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1}, 10000, 512, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1}, 20000, 512, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); - for (int kv : { 4096, 8192, 16384, }) { - for (int hs : { 64, 128, }) { - for (int nr : { 1, 4, }) { - test_cases.emplace_back(new test_flash_attn_ext(hs, hs, 8, {nr, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + for (int kv : { 4096, 8192, 16384,32768, 65536, }) { + for (int hs : { 64, 128, 256, }) { + for (int nr : { 1, 4, 8, }) { + for (int nb : { 1, 4096, }) { + test_cases.emplace_back(new test_flash_attn_ext(hs, hs, 8, {nr, 1}, kv, nb, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + } } } } From 8cb5912bfee880b03e2ef96991eb387797f3cad0 Mon Sep 17 00:00:00 2001 From: ynankani Date: Thu, 16 Jul 2026 11:48:27 +0000 Subject: [PATCH 02/11] Fix use 64bit generic pointer instead of 32bit shared pointer Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-swizzle.cuh | 25 ++++++++++++------------- 1 file changed, 12 insertions(+), 13 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh index 85da0b6ef7ee..763e9db91f4f 100644 --- a/ggml/src/ggml-cuda/fattn-swizzle.cuh +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -49,30 +49,29 @@ static __device__ __forceinline__ int ggml_cuda_fattn_swz_bytes_rc(const int row namespace ggml_cuda_fattn_smem_swizzle { -// ldmatrix.x4 from a 32-bit .shared address (lower register pressure than 64-bit generic pointers). -static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4(int * xi, const uint32_t saddr) { +// ldmatrix.x4 via 64-bit generic pointer. +static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4(int * xi, const half2 * addr) { asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];" : "=r"(xi[0]), "=r"(xi[1]), "=r"(xi[2]), "=r"(xi[3]) - : "r"(saddr)); + : "l"(addr)); } -static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4_trans(int * xi, const uint32_t saddr) { +static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4_trans(int * xi, const half2 * addr) { asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.b16 {%0, %1, %2, %3}, [%4];" : "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3]) - : "r"(saddr)); + : "l"(addr)); } -// Per-lane swizzled .shared address for tile<16,8> ldmatrix. +// Per-lane swizzled generic pointer for tile<16,8> ldmatrix. template -static __device__ __forceinline__ uint32_t ggml_cuda_fattn_swz_saddr( +static __device__ __forceinline__ const half2 * ggml_cuda_fattn_swz_saddr( const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) { const int lane_row = threadIdx.x % I; const int lane_col = (threadIdx.x / I) * (J / 2); - const uint32_t base = __cvta_generic_to_shared(tile_base); uint32_t byte_off = (uint32_t)((base_row + lane_row) * stride_h2 + base_col_h2 + lane_col) * (uint32_t)sizeof(half2); if constexpr (ggml_cuda_fattn_swz_enabled(stride_h2)) { byte_off ^= (uint32_t)(((base_row + lane_row) & 7) << 4); } - return base + byte_off; + return (const half2 *) ((const char *) tile_base + byte_off); } template @@ -83,8 +82,8 @@ static __device__ __forceinline__ void load_ldmatrix( constexpr int J = Tile::J; #if defined(TURING_MMA_AVAILABLE) if constexpr (I == 16 && J == 8 && ggml_cuda_fattn_swz_enabled(stride_h2)) { - const uint32_t saddr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); - ggml_cuda_fattn_ldmatrix_x4((int *) t.x, saddr); + const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); + ggml_cuda_fattn_ldmatrix_x4((int *) t.x, addr); return; } #endif // TURING_MMA_AVAILABLE @@ -99,8 +98,8 @@ static __device__ __forceinline__ void load_ldmatrix_trans( constexpr int J = Tile::J; #if defined(TURING_MMA_AVAILABLE) if constexpr (I == 16 && J == 8 && ggml_cuda_fattn_swz_enabled(stride_h2)) { - const uint32_t saddr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); - ggml_cuda_fattn_ldmatrix_x4_trans((int *) t.x, saddr); + const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); + ggml_cuda_fattn_ldmatrix_x4_trans((int *) t.x, addr); return; } #endif // TURING_MMA_AVAILABLE From 73e5d1e47d3c1f8bfa2c09d6cd6455e175272cd2 Mon Sep 17 00:00:00 2001 From: ynankani Date: Wed, 29 Jul 2026 15:04:02 +0000 Subject: [PATCH 03/11] fix shared memory race in FA on DGX Spark --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 143328e9d4e0..7c12fb7805ff 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -1443,6 +1443,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const int jc_cwm = threadIdx.y*(2*T_C_VKQ::J) + 2*T_C_VKQ::get_j(-1) + jc_cwmo; // jc combine write meta const float2 KQ_cmr = make_float2(KQ_max[jc_cwmo], KQ_rowsum[jc_cwmo]); // KQ combine max rowsum + __syncthreads(); + if (((!needs_fixup && !is_fixup) || np > 1) && threadIdx.x < 2*T_C_VKQ::J) { // Use the 16 bytes of padding in each row to store the meta data: KQ max, KQ rowsum, KQ max scale. ((float2 *) tile_Q)[jc_cwm*(tile_stride/2) + nbatch_combine/2] = KQ_cmr; @@ -1479,6 +1481,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const bool thread_should_write = T_C_KQ::J == 8 || T_C_KQ::get_j(threadIdx.x & 2) < 8; #endif // defined(TURING_MMA_AVAILABLE) + __syncthreads(); + if (((!needs_fixup && !is_fixup) || np > 1) && thread_should_write) { ((float2 *) tile_Q)[jc_cwm*(tile_stride/2) + nbatch_combine/2] = KQ_cmr; } From 50b0d08f5cdbd057da7b487b2371af6b0e4bd99a Mon Sep 17 00:00:00 2001 From: ynankani Date: Fri, 31 Jul 2026 10:44:30 +0000 Subject: [PATCH 04/11] Handle corener case Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 27 +++++++++++++++------------ ggml/src/ggml-cuda/fattn-swizzle.cuh | 22 ++++++++++++---------- 2 files changed, 27 insertions(+), 22 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 7c12fb7805ff..3e75bc06b775 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -361,7 +361,7 @@ static constexpr __device__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, c // ------------------------------------------------------------------------------------------------------------------ -template +template static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( const half2 * const __restrict__ KV, half2 * const __restrict__ tile_KV, const int D2, const int stride_KV, const int i_sup) { constexpr int warp_size = ggml_cuda_get_physical_warp_size(); @@ -398,7 +398,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - const int smem_offs_b = ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk); + const int smem_offs_b = ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk); cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i*stride_KV + k*h2_per_chunk); } } @@ -434,7 +434,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - ggml_cuda_memcpy_1<16>((char*)tile_KV + ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk), + ggml_cuda_memcpy_1<16>((char*)tile_KV + ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk), !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); } } @@ -571,6 +571,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); // swizzle the tile stride for K and V based on the batch size. + constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); + constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_swz_enabled(nbatch_V2); constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); @@ -590,7 +592,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr bool use_cp_async = true; cp_async_wait_all(); __syncthreads(); - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (V_h2 + int64_t(k_VKQ_0)*stride_V, tile_V, nbatch_V2, stride_V, k_VKQ_sup); } else { constexpr bool use_cp_async = nstages == 1; @@ -609,7 +611,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( if constexpr (nstages <= 1) { const int k0_diff = k0_stop - k0_start; constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (K_h2 + int64_t(k_VKQ_0)*stride_K + k0_start, tile_K, k0_diff, stride_K, k_VKQ_sup); if (use_cp_async) { cp_async_wait_all(); @@ -625,7 +627,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( #pragma unroll for (int k_KQ_0 = k0_start; k_KQ_0 < k0_stop; k_KQ_0 += T_A_KQ::J) { T_A_KQ K_A; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix( + ggml_cuda_fattn_smem_swizzle::load_ldmatrix( K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[k_KQ_0/T_A_KQ::J]); @@ -652,7 +654,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i_KQ_0 = i_KQ_00 + (threadIdx.y % np)*T_A_KQ::I; T_A_KQ K_A; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix( + ggml_cuda_fattn_smem_swizzle::load_ldmatrix( K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { @@ -947,7 +949,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( flash_attn_ext_f16_load_mask (mask_h + k_VKQ_0 + nbatch_fa, tile_mask, stride_mask, k_VKQ_sup, jt*ncols1, ne01); } - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (K_h2 + int64_t(k_VKQ_0 + nbatch_fa)*stride_K, tile_K, nbatch_K2, stride_K, k_VKQ_sup); } } @@ -963,7 +965,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i0_diff = i0_stop - i0_start; if (!V_is_K_view || i0_stop > 2*nbatch_K2) { constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (V_h2 + int64_t(k_VKQ_0)*stride_V + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_sup); if (use_cp_async) { cp_async_wait_all(); @@ -983,7 +985,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( T_A_VKQ A; // Transposed in SRAM but not in registers, gets transposed on load. const int v_lin = (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans( + ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans( A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); if constexpr (T_B_KQ::I == 8) { mma(VKQ_C[i_VKQ_0/T_A_VKQ::I], A, B[k00/(np*T_A_VKQ::J)]); @@ -1011,7 +1013,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( T_A_VKQ A; // Transposed in both SRAM and registers, load normally. const int v_lin = (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix( + ggml_cuda_fattn_smem_swizzle::load_ldmatrix( A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); mma(VKQ_C[i_VKQ_0/i0_stride], B[k00/(np*T_A_VKQ::I)], A); } @@ -1177,6 +1179,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int stride_tile_Q = DKQ/2 + 4; // swizzle the tile stride for K and V based on the batch size. + constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); constexpr int stride_tile_KV_max = stride_tile_K > stride_tile_V ? stride_tile_K : stride_tile_V; @@ -1273,7 +1276,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( flash_attn_ext_f16_load_mask (mask_h + kb0*nbatch_fa, tile_mask, stride_mask, k_VKQ_sup, jt*ncols1, ne01); } - flash_attn_ext_f16_load_tile + flash_attn_ext_f16_load_tile (K_h2 + int64_t(kb0)*nbatch_fa*stride_K, tile_K, nbatch_K2, stride_K, k_VKQ_sup); } diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh index 763e9db91f4f..c33b5e332b06 100644 --- a/ggml/src/ggml-cuda/fattn-swizzle.cuh +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -38,10 +38,11 @@ static __host__ int ggml_cuda_fattn_swz_tile_stride(const int nbatch_2, const in } // Swizzled byte offset for tile element (row, col_h2); same map used for writes and reads. -template +template static __device__ __forceinline__ int ggml_cuda_fattn_swz_bytes_rc(const int row, const int col_h2) { + static_assert(!swz || ggml_cuda_fattn_swz_pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); int off_bytes = (row * stride_h2 + col_h2) * (int) sizeof(half2); - if constexpr (ggml_cuda_fattn_swz_enabled(stride_h2)) { + if constexpr (swz) { off_bytes ^= (row & 7) << 4; } return off_bytes; @@ -62,27 +63,28 @@ static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4_trans(int * x } // Per-lane swizzled generic pointer for tile<16,8> ldmatrix. -template +template static __device__ __forceinline__ const half2 * ggml_cuda_fattn_swz_saddr( const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) { + static_assert(!swz || ggml_cuda_fattn_swz_pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); const int lane_row = threadIdx.x % I; const int lane_col = (threadIdx.x / I) * (J / 2); uint32_t byte_off = (uint32_t)((base_row + lane_row) * stride_h2 + base_col_h2 + lane_col) * (uint32_t)sizeof(half2); - if constexpr (ggml_cuda_fattn_swz_enabled(stride_h2)) { + if constexpr (swz) { byte_off ^= (uint32_t)(((base_row + lane_row) & 7) << 4); } return (const half2 *) ((const char *) tile_base + byte_off); } -template +template static __device__ __forceinline__ void load_ldmatrix( TileT & t, half2 * tile_base, const int base_row, const int base_col_h2) { using Tile = typename std::remove_reference::type; constexpr int I = Tile::I; constexpr int J = Tile::J; #if defined(TURING_MMA_AVAILABLE) - if constexpr (I == 16 && J == 8 && ggml_cuda_fattn_swz_enabled(stride_h2)) { - const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); + if constexpr (I == 16 && J == 8 && swz) { + const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); ggml_cuda_fattn_ldmatrix_x4((int *) t.x, addr); return; } @@ -90,15 +92,15 @@ static __device__ __forceinline__ void load_ldmatrix( ggml_cuda_mma::load_ldmatrix(t, tile_base + base_row * stride_h2 + base_col_h2, stride_h2); } -template +template static __device__ __forceinline__ void load_ldmatrix_trans( TileT & t, half2 * tile_base, const int base_row, const int base_col_h2) { using Tile = typename std::remove_reference::type; constexpr int I = Tile::I; constexpr int J = Tile::J; #if defined(TURING_MMA_AVAILABLE) - if constexpr (I == 16 && J == 8 && ggml_cuda_fattn_swz_enabled(stride_h2)) { - const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); + if constexpr (I == 16 && J == 8 && swz) { + const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); ggml_cuda_fattn_ldmatrix_x4_trans((int *) t.x, addr); return; } From 24da80813a4f6b1ecf168b08bad4db21fe19df99 Mon Sep 17 00:00:00 2001 From: ynankani Date: Tue, 4 Aug 2026 16:29:13 +0000 Subject: [PATCH 05/11] Add swizzle test cases and gate sync for swizzled path only Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 17 ++++++++++++----- tests/test-backend-ops.cpp | 7 +++++++ 2 files changed, 19 insertions(+), 5 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 3e75bc06b775..4e52f8fdcf17 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -571,10 +571,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); // swizzle the tile stride for K and V based on the batch size. - constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); - constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_swz_enabled(nbatch_V2); constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); + constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); + constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_swz_enabled(nbatch_V2); const int k_VKQ_0 = kb0 * nbatch_fa; #if defined(TURING_MMA_AVAILABLE) @@ -1179,10 +1179,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int stride_tile_Q = DKQ/2 + 4; // swizzle the tile stride for K and V based on the batch size. - constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); constexpr int stride_tile_KV_max = stride_tile_K > stride_tile_V ? stride_tile_K : stride_tile_V; + constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); + constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_swz_enabled(nbatch_V2); extern __shared__ half2 tile_Q[]; half2 * tile_K = Q_in_reg ? tile_Q : tile_Q + ncols * stride_tile_Q; @@ -1441,12 +1442,16 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int tile_stride = nbatch_combine + 4; static_assert((DV/2) % nbatch_combine == 0, "bad nbatch_combine"); + constexpr bool combine_needs_sync = swz_K || swz_V; + if constexpr (cols_per_warp == 8) { const int jc_cwmo = (threadIdx.x % (2*T_C_VKQ::J)) / T_C_VKQ::J; // jc combine write meta offset const int jc_cwm = threadIdx.y*(2*T_C_VKQ::J) + 2*T_C_VKQ::get_j(-1) + jc_cwmo; // jc combine write meta const float2 KQ_cmr = make_float2(KQ_max[jc_cwmo], KQ_rowsum[jc_cwmo]); // KQ combine max rowsum - __syncthreads(); + if constexpr (combine_needs_sync) { + __syncthreads(); + } if (((!needs_fixup && !is_fixup) || np > 1) && threadIdx.x < 2*T_C_VKQ::J) { // Use the 16 bytes of padding in each row to store the meta data: KQ max, KQ rowsum, KQ max scale. @@ -1484,7 +1489,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const bool thread_should_write = T_C_KQ::J == 8 || T_C_KQ::get_j(threadIdx.x & 2) < 8; #endif // defined(TURING_MMA_AVAILABLE) - __syncthreads(); + if constexpr (combine_needs_sync) { + __syncthreads(); + } if (((!needs_fixup && !is_fixup) || np > 1) && thread_should_write) { ((float2 *) tile_Q)[jc_cwm*(tile_stride/2) + nbatch_combine/2] = KQ_cmr; diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index ba1c7085464e..0aa4b3c99396 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10078,6 +10078,13 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext(192, 128, 8, {8, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {20, 1}, 512, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // FLASH_ATTN_EXT MMA, swizzled K/V tiles: nbatch_K2 = 32, 64, 128, 256. + test_cases.emplace_back(new test_flash_attn_ext( 64, 64, 8, {8, 1}, 4096, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {4, 1}, 4096, 8, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + 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)); + + 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})); From bc56f686ad6e4b05b33cf6617855b5c978140586 Mon Sep 17 00:00:00 2001 From: ynankani Date: Tue, 11 Aug 2026 05:58:23 +0000 Subject: [PATCH 06/11] gate CUDA PTX Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-swizzle.cuh | 2 ++ 1 file changed, 2 insertions(+) diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh index c33b5e332b06..abb58ab29f5e 100644 --- a/ggml/src/ggml-cuda/fattn-swizzle.cuh +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -50,6 +50,7 @@ static __device__ __forceinline__ int ggml_cuda_fattn_swz_bytes_rc(const int row namespace ggml_cuda_fattn_smem_swizzle { +#if defined(TURING_MMA_AVAILABLE) // ldmatrix.x4 via 64-bit generic pointer. static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4(int * xi, const half2 * addr) { asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];" @@ -61,6 +62,7 @@ static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4_trans(int * x : "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3]) : "l"(addr)); } +#endif // defined(TURING_MMA_AVAILABLE) // Per-lane swizzled generic pointer for tile<16,8> ldmatrix. template From c9a90f7c3057643651d6d40f6d8bee39ab3dc86b Mon Sep 17 00:00:00 2001 From: ynankani Date: Tue, 11 Aug 2026 08:42:28 +0000 Subject: [PATCH 07/11] offset calculation specific for swizzle branch Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 41 +++++++++++++++++++++------- 1 file changed, 31 insertions(+), 10 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 4e52f8fdcf17..af5f18a8799b 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -398,8 +398,13 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - const int smem_offs_b = ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk); - cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i*stride_KV + k*h2_per_chunk); + if constexpr (swz) { + const int smem_offs_b = ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk); + cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i*stride_KV + k*h2_per_chunk); + } else { + cp_async_cg_16( + tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i*stride_KV + k*h2_per_chunk); + } } } }; @@ -434,8 +439,15 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) { const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); - ggml_cuda_memcpy_1<16>((char*)tile_KV + ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk), - !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + if constexpr (swz) { + ggml_cuda_memcpy_1<16>( + (char *) tile_KV + ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk), + !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + } else { + ggml_cuda_memcpy_1<16>( + tile_KV + i*stride_tile + k*4, + !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + } } } }; @@ -984,9 +996,14 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::J; T_A_VKQ A; // Transposed in SRAM but not in registers, gets transposed on load. - const int v_lin = (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans( - A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); + if constexpr (swz_V) { + const int v_lin = (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; + ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans( + A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); + } else { + load_ldmatrix_trans( + A, tile_V_i + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); + } if constexpr (T_B_KQ::I == 8) { mma(VKQ_C[i_VKQ_0/T_A_VKQ::I], A, B[k00/(np*T_A_VKQ::J)]); } else { @@ -1012,9 +1029,13 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::I; T_A_VKQ A; // Transposed in both SRAM and registers, load normally. - const int v_lin = (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix( - A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); + if constexpr (swz_V) { + const int v_lin = (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; + ggml_cuda_fattn_smem_swizzle::load_ldmatrix( + A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); + } else { + load_ldmatrix(A, tile_V_i + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); + } mma(VKQ_C[i_VKQ_0/i0_stride], B[k00/(np*T_A_VKQ::I)], A); } } From 2a57183070ad1e776883d827c440c5349a32d327 Mon Sep 17 00:00:00 2001 From: ynankani Date: Fri, 14 Aug 2026 08:12:25 +0000 Subject: [PATCH 08/11] Reafctor code Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 54 +++++++-------- ggml/src/ggml-cuda/fattn-swizzle.cuh | 100 +++++++++++---------------- 2 files changed, 69 insertions(+), 85 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index af5f18a8799b..3b886721510a 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -399,11 +399,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); if constexpr (swz) { - const int smem_offs_b = ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk); + const int smem_offs_b = ggml_cuda_fattn_smem_swizzle::bytes_rc(i, k*h2_per_chunk); cp_async_cg_16(tile_KV_32 + smem_offs_b, KV + i*stride_KV + k*h2_per_chunk); } else { - cp_async_cg_16( - tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i*stride_KV + k*h2_per_chunk); + cp_async_cg_16(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i*stride_KV + k*h2_per_chunk); } } } @@ -440,12 +439,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k); if constexpr (swz) { - ggml_cuda_memcpy_1<16>( - (char *) tile_KV + ggml_cuda_fattn_swz_bytes_rc(i, k*h2_per_chunk), + ggml_cuda_memcpy_1<16>((char *) tile_KV + ggml_cuda_fattn_smem_swizzle::bytes_rc(i, k*h2_per_chunk), !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); } else { - ggml_cuda_memcpy_1<16>( - tile_KV + i*stride_tile + k*4, + ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); } } @@ -583,10 +580,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); // swizzle the tile stride for K and V based on the batch size. - constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); - constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); - constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); - constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_swz_enabled(nbatch_V2); + constexpr int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2); + constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2); + constexpr bool swz_K = ggml_cuda_fattn_smem_swizzle::enabled(nbatch_K2); + constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_smem_swizzle::enabled(nbatch_V2); const int k_VKQ_0 = kb0 * nbatch_fa; #if defined(TURING_MMA_AVAILABLE) @@ -639,8 +636,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( #pragma unroll for (int k_KQ_0 = k0_start; k_KQ_0 < k0_stop; k_KQ_0 += T_A_KQ::J) { T_A_KQ K_A; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix( - K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); + if constexpr (swz_K) { + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); + } else { + load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); + } if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[k_KQ_0/T_A_KQ::J]); } else { @@ -666,8 +666,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i_KQ_0 = i_KQ_00 + (threadIdx.y % np)*T_A_KQ::I; T_A_KQ K_A; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix( - K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); + if constexpr (swz_K) { + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); + } else { + load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); + } if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[0]); @@ -998,11 +1001,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( T_A_VKQ A; // Transposed in SRAM but not in registers, gets transposed on load. if constexpr (swz_V) { const int v_lin = (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans( - A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans(A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); } else { - load_ldmatrix_trans( - A, tile_V_i + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); + load_ldmatrix_trans(A, tile_V_i + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); } if constexpr (T_B_KQ::I == 8) { mma(VKQ_C[i_VKQ_0/T_A_VKQ::I], A, B[k00/(np*T_A_VKQ::J)]); @@ -1031,8 +1032,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( T_A_VKQ A; // Transposed in both SRAM and registers, load normally. if constexpr (swz_V) { const int v_lin = (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix( - A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); } else { load_ldmatrix(A, tile_V_i + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); } @@ -1200,11 +1200,11 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int stride_tile_Q = DKQ/2 + 4; // swizzle the tile stride for K and V based on the batch size. - constexpr int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2); - constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2); + constexpr int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2); + constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2); constexpr int stride_tile_KV_max = stride_tile_K > stride_tile_V ? stride_tile_K : stride_tile_V; - constexpr bool swz_K = ggml_cuda_fattn_swz_enabled(nbatch_K2); - constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_swz_enabled(nbatch_V2); + constexpr bool swz_K = ggml_cuda_fattn_smem_swizzle::enabled(nbatch_K2); + constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_smem_swizzle::enabled(nbatch_V2); extern __shared__ half2 tile_Q[]; half2 * tile_K = Q_in_reg ? tile_Q : tile_Q + ncols * stride_tile_Q; @@ -1958,8 +1958,8 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml constexpr bool V_is_K_view = DKQ == 576; // Guaranteed by the kernel selection logic in fattn.cu // KV tile strides must match flash_attn_ext_f16_iter / _process_tile. - const int stride_tile_K = ggml_cuda_fattn_swz_tile_stride(nbatch_K2, cc); - const int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_swz_tile_stride(nbatch_V2, cc); + const int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2, cc); + const int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2, cc); const size_t nbytes_shared_KV_1stage = nbatch_fa * std::max(stride_tile_K, stride_tile_V) * sizeof(half2); const size_t nbytes_shared_KV_2stage = nbatch_fa * (stride_tile_K + stride_tile_V) * sizeof(half2); const size_t nbytes_shared_Q = ncols * (DKQ/2 + 4) * sizeof(half2); diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh index abb58ab29f5e..00e72939e229 100644 --- a/ggml/src/ggml-cuda/fattn-swizzle.cuh +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -6,108 +6,92 @@ // XOR swizzle for K/V SMEM tiles to avoid bank conflicts without row padding (Turing+ only). // Stride must be a power-of-two >= 32 half2 columns,otherwise we keep +4 row padding. -static __host__ __device__ constexpr bool ggml_cuda_fattn_swz_pow2_stride(const int nbatch_2) { +namespace ggml_cuda_fattn_smem_swizzle { + +static __host__ __device__ constexpr bool pow2_stride(const int nbatch_2) { return nbatch_2 >= 32 && (nbatch_2 & (nbatch_2 - 1)) == 0; } -static __device__ constexpr bool ggml_cuda_fattn_swz_enabled(const int nbatch_2) { +static __device__ constexpr bool enabled(const int nbatch_2) { #if defined(TURING_MMA_AVAILABLE) - return ggml_cuda_fattn_swz_pow2_stride(nbatch_2); + return pow2_stride(nbatch_2); #else GGML_UNUSED(nbatch_2); return false; -#endif +#endif // defined(TURING_MMA_AVAILABLE) } -static __host__ bool ggml_cuda_fattn_swz_enabled(const int nbatch_2, const int cc) { +static __host__ bool enabled(const int nbatch_2, const int cc) { #ifdef GGML_USE_HIP GGML_UNUSED(nbatch_2); GGML_UNUSED(cc); return false; #else - return turing_mma_available(cc) && ggml_cuda_fattn_swz_pow2_stride(nbatch_2); -#endif + return turing_mma_available(cc) && pow2_stride(nbatch_2); +#endif // GGML_USE_HIP } -static __device__ constexpr int ggml_cuda_fattn_swz_tile_stride(const int nbatch_2) { - return ggml_cuda_fattn_swz_enabled(nbatch_2) ? nbatch_2 : nbatch_2 + 4; +static __device__ constexpr int tile_stride(const int nbatch_2) { + return enabled(nbatch_2) ? nbatch_2 : nbatch_2 + 4; } -static __host__ int ggml_cuda_fattn_swz_tile_stride(const int nbatch_2, const int cc) { - return ggml_cuda_fattn_swz_enabled(nbatch_2, cc) ? nbatch_2 : nbatch_2 + 4; +static __host__ int tile_stride(const int nbatch_2, const int cc) { + return enabled(nbatch_2, cc) ? nbatch_2 : nbatch_2 + 4; } // Swizzled byte offset for tile element (row, col_h2); same map used for writes and reads. -template -static __device__ __forceinline__ int ggml_cuda_fattn_swz_bytes_rc(const int row, const int col_h2) { - static_assert(!swz || ggml_cuda_fattn_swz_pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); - int off_bytes = (row * stride_h2 + col_h2) * (int) sizeof(half2); - if constexpr (swz) { - off_bytes ^= (row & 7) << 4; - } - return off_bytes; +template +static __device__ __forceinline__ int bytes_rc(const int row, const int col_h2) { + static_assert(pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); + return ((row * stride_h2 + col_h2) * (int) sizeof(half2)) ^ ((row & 7) << 4); } -namespace ggml_cuda_fattn_smem_swizzle { - #if defined(TURING_MMA_AVAILABLE) // ldmatrix.x4 via 64-bit generic pointer. -static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4(int * xi, const half2 * addr) { +static __device__ __forceinline__ void ldmatrix_x4(int * xi, const half2 * addr) { asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];" : "=r"(xi[0]), "=r"(xi[1]), "=r"(xi[2]), "=r"(xi[3]) : "l"(addr)); } -static __device__ __forceinline__ void ggml_cuda_fattn_ldmatrix_x4_trans(int * xi, const half2 * addr) { + +static __device__ __forceinline__ void ldmatrix_x4_trans(int * xi, const half2 * addr) { asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.b16 {%0, %1, %2, %3}, [%4];" : "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3]) : "l"(addr)); } #endif // defined(TURING_MMA_AVAILABLE) -// Per-lane swizzled generic pointer for tile<16,8> ldmatrix. -template -static __device__ __forceinline__ const half2 * ggml_cuda_fattn_swz_saddr( - const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) { - static_assert(!swz || ggml_cuda_fattn_swz_pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); - const int lane_row = threadIdx.x % I; - const int lane_col = (threadIdx.x / I) * (J / 2); - uint32_t byte_off = (uint32_t)((base_row + lane_row) * stride_h2 + base_col_h2 + lane_col) * (uint32_t)sizeof(half2); - if constexpr (swz) { - byte_off ^= (uint32_t)(((base_row + lane_row) & 7) << 4); - } +// Per-lane swizzled address for one tile<16, 8, half2> ldmatrix: 16 rows, 4 half2 columns per lane. +template +static __device__ __forceinline__ const half2 * lane_addr( + const half2 * tile_base, const int base_row, const int base_col_h2) { + static_assert(pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); + const int row = base_row + threadIdx.x % 16; + const int col = base_col_h2 + (threadIdx.x / 16) * 4; + const uint32_t byte_off = (uint32_t) ((row * stride_h2 + col) * (int) sizeof(half2)) ^ (uint32_t) ((row & 7) << 4); return (const half2 *) ((const char *) tile_base + byte_off); } -template +template static __device__ __forceinline__ void load_ldmatrix( - TileT & t, half2 * tile_base, const int base_row, const int base_col_h2) { - using Tile = typename std::remove_reference::type; - constexpr int I = Tile::I; - constexpr int J = Tile::J; + ggml_cuda_mma::tile<16, 8, half2> & t, const half2 * tile_base, const int base_row, const int base_col_h2) { #if defined(TURING_MMA_AVAILABLE) - if constexpr (I == 16 && J == 8 && swz) { - const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); - ggml_cuda_fattn_ldmatrix_x4((int *) t.x, addr); - return; - } -#endif // TURING_MMA_AVAILABLE - ggml_cuda_mma::load_ldmatrix(t, tile_base + base_row * stride_h2 + base_col_h2, stride_h2); + ldmatrix_x4((int *) t.x, lane_addr(tile_base, base_row, base_col_h2)); +#else + GGML_UNUSED_VARS(t, tile_base, base_row, base_col_h2); + NO_DEVICE_CODE; +#endif // defined(TURING_MMA_AVAILABLE) } -template +template static __device__ __forceinline__ void load_ldmatrix_trans( - TileT & t, half2 * tile_base, const int base_row, const int base_col_h2) { - using Tile = typename std::remove_reference::type; - constexpr int I = Tile::I; - constexpr int J = Tile::J; + ggml_cuda_mma::tile<16, 8, half2> & t, const half2 * tile_base, const int base_row, const int base_col_h2) { #if defined(TURING_MMA_AVAILABLE) - if constexpr (I == 16 && J == 8 && swz) { - const half2 * addr = ggml_cuda_fattn_swz_saddr(tile_base, base_row, base_col_h2, I, J); - ggml_cuda_fattn_ldmatrix_x4_trans((int *) t.x, addr); - return; - } -#endif // TURING_MMA_AVAILABLE - ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + base_row * stride_h2 + base_col_h2, stride_h2); + ldmatrix_x4_trans((int *) t.x, lane_addr(tile_base, base_row, base_col_h2)); +#else + GGML_UNUSED_VARS(t, tile_base, base_row, base_col_h2); + NO_DEVICE_CODE; +#endif // defined(TURING_MMA_AVAILABLE) } } // namespace ggml_cuda_fattn_smem_swizzle From de01a0cc3f393c206878aa746847194b5a167ffb Mon Sep 17 00:00:00 2001 From: ynankani Date: Tue, 18 Aug 2026 14:46:29 +0000 Subject: [PATCH 09/11] Refactor FA swizzle ldmatrix if/else into helpers (K row/col, V offset) Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 26 ++-------- ggml/src/ggml-cuda/fattn-swizzle.cuh | 75 +++++++++++++++++++--------- 2 files changed, 56 insertions(+), 45 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 3b886721510a..387e70fa1497 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -636,11 +636,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( #pragma unroll for (int k_KQ_0 = k0_start; k_KQ_0 < k0_stop; k_KQ_0 += T_A_KQ::J) { T_A_KQ K_A; - if constexpr (swz_K) { - ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); - } else { - load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); - } + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[k_KQ_0/T_A_KQ::J]); } else { @@ -666,11 +662,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int i_KQ_0 = i_KQ_00 + (threadIdx.y % np)*T_A_KQ::I; T_A_KQ K_A; - if constexpr (swz_K) { - ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); - } else { - load_ldmatrix(K_A, tile_K + i_KQ_0*stride_tile_K + (k_KQ_0 - k0_start), stride_tile_K); - } + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start); if constexpr (cols_per_warp == 8) { mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[0]); @@ -999,12 +991,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::J; T_A_VKQ A; // Transposed in SRAM but not in registers, gets transposed on load. - if constexpr (swz_V) { - const int v_lin = (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans(A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); - } else { - load_ldmatrix_trans(A, tile_V_i + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); - } + ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans(A, tile_V, (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2); if constexpr (T_B_KQ::I == 8) { mma(VKQ_C[i_VKQ_0/T_A_VKQ::I], A, B[k00/(np*T_A_VKQ::J)]); } else { @@ -1030,12 +1017,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::I; T_A_VKQ A; // Transposed in both SRAM and registers, load normally. - if constexpr (swz_V) { - const int v_lin = (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2; - ggml_cuda_fattn_smem_swizzle::load_ldmatrix(A, tile_V, v_lin / stride_tile_V, v_lin % stride_tile_V); - } else { - load_ldmatrix(A, tile_V_i + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V); - } + ggml_cuda_fattn_smem_swizzle::load_ldmatrix(A, tile_V, (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2); mma(VKQ_C[i_VKQ_0/i0_stride], B[k00/(np*T_A_VKQ::I)], A); } } diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh index 00e72939e229..4746dd6a300c 100644 --- a/ggml/src/ggml-cuda/fattn-swizzle.cuh +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -39,59 +39,88 @@ static __host__ int tile_stride(const int nbatch_2, const int cc) { return enabled(nbatch_2, cc) ? nbatch_2 : nbatch_2 + 4; } -// Swizzled byte offset for tile element (row, col_h2); same map used for writes and reads. +// Swizzled byte offset for tile element (row, col_h2), same map used for writes and reads. template static __device__ __forceinline__ int bytes_rc(const int row, const int col_h2) { static_assert(pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); return ((row * stride_h2 + col_h2) * (int) sizeof(half2)) ^ ((row & 7) << 4); } -#if defined(TURING_MMA_AVAILABLE) // ldmatrix.x4 via 64-bit generic pointer. static __device__ __forceinline__ void ldmatrix_x4(int * xi, const half2 * addr) { +#if defined(TURING_MMA_AVAILABLE) asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];" : "=r"(xi[0]), "=r"(xi[1]), "=r"(xi[2]), "=r"(xi[3]) : "l"(addr)); +#else + GGML_UNUSED_VARS(xi, addr); + NO_DEVICE_CODE; +#endif // defined(TURING_MMA_AVAILABLE) } static __device__ __forceinline__ void ldmatrix_x4_trans(int * xi, const half2 * addr) { +#if defined(TURING_MMA_AVAILABLE) asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.b16 {%0, %1, %2, %3}, [%4];" : "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3]) : "l"(addr)); -} +#else + GGML_UNUSED_VARS(xi, addr); + NO_DEVICE_CODE; #endif // defined(TURING_MMA_AVAILABLE) +} // Per-lane swizzled address for one tile<16, 8, half2> ldmatrix: 16 rows, 4 half2 columns per lane. template static __device__ __forceinline__ const half2 * lane_addr( - const half2 * tile_base, const int base_row, const int base_col_h2) { + const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) { static_assert(pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); - const int row = base_row + threadIdx.x % 16; - const int col = base_col_h2 + (threadIdx.x / 16) * 4; - const uint32_t byte_off = (uint32_t) ((row * stride_h2 + col) * (int) sizeof(half2)) ^ (uint32_t) ((row & 7) << 4); + const int lane_row = threadIdx.x % I; + const int lane_col = (threadIdx.x / I) * (J / 2); + uint32_t byte_off = (uint32_t) ((base_row + lane_row)*stride_h2 + base_col_h2 + lane_col) * (uint32_t) sizeof(half2); + byte_off ^= (uint32_t) (((base_row + lane_row) & 7) << 4); return (const half2 *) ((const char *) tile_base + byte_off); } -template +template static __device__ __forceinline__ void load_ldmatrix( - ggml_cuda_mma::tile<16, 8, half2> & t, const half2 * tile_base, const int base_row, const int base_col_h2) { -#if defined(TURING_MMA_AVAILABLE) - ldmatrix_x4((int *) t.x, lane_addr(tile_base, base_row, base_col_h2)); -#else - GGML_UNUSED_VARS(t, tile_base, base_row, base_col_h2); - NO_DEVICE_CODE; -#endif // defined(TURING_MMA_AVAILABLE) + TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) { + if constexpr (swz) { + static_assert(std::is_same_v>, + "the swizzled layout is only supported for tile<16, 8, half2>"); + ldmatrix_x4((int *) t.x, lane_addr(tile_base, base_row, base_col_h2, TileT::I, TileT::J)); + } else { + ggml_cuda_mma::load_ldmatrix(t, tile_base + base_row*stride_h2 + base_col_h2, stride_h2); + } } -template +template +static __device__ __forceinline__ void load_ldmatrix(TileT & t, const half2 * tile_base, const int off_h2) { + if constexpr (swz) { + load_ldmatrix(t, tile_base, off_h2 / stride_h2, off_h2 % stride_h2); + } else { + ggml_cuda_mma::load_ldmatrix(t, tile_base + off_h2, stride_h2); + } +} + +template static __device__ __forceinline__ void load_ldmatrix_trans( - ggml_cuda_mma::tile<16, 8, half2> & t, const half2 * tile_base, const int base_row, const int base_col_h2) { -#if defined(TURING_MMA_AVAILABLE) - ldmatrix_x4_trans((int *) t.x, lane_addr(tile_base, base_row, base_col_h2)); -#else - GGML_UNUSED_VARS(t, tile_base, base_row, base_col_h2); - NO_DEVICE_CODE; -#endif // defined(TURING_MMA_AVAILABLE) + TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) { + if constexpr (swz) { + static_assert(std::is_same_v>, + "the swizzled layout is only supported for tile<16, 8, half2>"); + ldmatrix_x4_trans((int *) t.x, lane_addr(tile_base, base_row, base_col_h2, TileT::I, TileT::J)); + } else { + ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + base_row*stride_h2 + base_col_h2, stride_h2); + } +} + +template +static __device__ __forceinline__ void load_ldmatrix_trans(TileT & t, const half2 * tile_base, const int off_h2) { + if constexpr (swz) { + load_ldmatrix_trans(t, tile_base, off_h2 / stride_h2, off_h2 % stride_h2); + } else { + ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + off_h2, stride_h2); + } } } // namespace ggml_cuda_fattn_smem_swizzle From 8ccc1daaf6395a71d0ec33edbb5314de81bdec7b Mon Sep 17 00:00:00 2001 From: ynankani Date: Mon, 24 Aug 2026 18:45:00 +0000 Subject: [PATCH 10/11] rebase and update test case args Signed-off-by: ynankani --- tests/test-backend-ops.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 0aa4b3c99396..babad87bede7 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10076,7 +10076,7 @@ static std::vector> make_test_cases_eval() { // FLASH_ATTN_EXT MMA: non-pow2 head size and MLA K/V view. test_cases.emplace_back(new test_flash_attn_ext(192, 128, 8, {8, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); - test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {20, 1}, 512, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {20, 1}, 512, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true)); // FLASH_ATTN_EXT MMA, swizzled K/V tiles: nbatch_K2 = 32, 64, 128, 256. test_cases.emplace_back(new test_flash_attn_ext( 64, 64, 8, {8, 1}, 4096, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); From 088d425c59ea5932aeb8f5da237efa775c7ca486 Mon Sep 17 00:00:00 2001 From: ynankani Date: Mon, 31 Aug 2026 08:25:59 +0000 Subject: [PATCH 11/11] Allow swizzle for non-pow2 shapes, for which nbatch_2%32==0 Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-swizzle.cuh | 14 +++++++------- tests/test-backend-ops.cpp | 9 +++++---- 2 files changed, 12 insertions(+), 11 deletions(-) diff --git a/ggml/src/ggml-cuda/fattn-swizzle.cuh b/ggml/src/ggml-cuda/fattn-swizzle.cuh index 4746dd6a300c..44338c8db08d 100644 --- a/ggml/src/ggml-cuda/fattn-swizzle.cuh +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -4,17 +4,17 @@ #include "mma.cuh" // XOR swizzle for K/V SMEM tiles to avoid bank conflicts without row padding (Turing+ only). -// Stride must be a power-of-two >= 32 half2 columns,otherwise we keep +4 row padding. +// Stride must be a multiple of 32 half2 columns, otherwise we keep +4 row padding. namespace ggml_cuda_fattn_smem_swizzle { -static __host__ __device__ constexpr bool pow2_stride(const int nbatch_2) { - return nbatch_2 >= 32 && (nbatch_2 & (nbatch_2 - 1)) == 0; +static __host__ __device__ constexpr bool bank_aligned(const int nbatch_2) { + return nbatch_2 >= 32 && nbatch_2 % 32 == 0; } static __device__ constexpr bool enabled(const int nbatch_2) { #if defined(TURING_MMA_AVAILABLE) - return pow2_stride(nbatch_2); + return bank_aligned(nbatch_2); #else GGML_UNUSED(nbatch_2); return false; @@ -27,7 +27,7 @@ static __host__ bool enabled(const int nbatch_2, const int cc) { GGML_UNUSED(cc); return false; #else - return turing_mma_available(cc) && pow2_stride(nbatch_2); + return turing_mma_available(cc) && bank_aligned(nbatch_2); #endif // GGML_USE_HIP } @@ -42,7 +42,7 @@ static __host__ int tile_stride(const int nbatch_2, const int cc) { // Swizzled byte offset for tile element (row, col_h2), same map used for writes and reads. template static __device__ __forceinline__ int bytes_rc(const int row, const int col_h2) { - static_assert(pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); + static_assert(bank_aligned(stride_h2), "swizzled tile needs a stride that is a multiple of 32"); return ((row * stride_h2 + col_h2) * (int) sizeof(half2)) ^ ((row & 7) << 4); } @@ -73,7 +73,7 @@ static __device__ __forceinline__ void ldmatrix_x4_trans(int * xi, const half2 * template static __device__ __forceinline__ const half2 * lane_addr( const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) { - static_assert(pow2_stride(stride_h2), "swizzled tile needs a pow2 stride"); + static_assert(bank_aligned(stride_h2), "swizzled tile needs a stride that is a multiple of 32"); const int lane_row = threadIdx.x % I; const int lane_col = (threadIdx.x / I) * (J / 2); uint32_t byte_off = (uint32_t) ((base_row + lane_row)*stride_h2 + base_col_h2 + lane_col) * (uint32_t) sizeof(half2); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index babad87bede7..c0b0e85c0c79 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10078,13 +10078,12 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext(192, 128, 8, {8, 1}, 4096, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {20, 1}, 512, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true)); - // FLASH_ATTN_EXT MMA, swizzled K/V tiles: nbatch_K2 = 32, 64, 128, 256. + // FLASH_ATTN_EXT MMA, swizzled K/V tiles, power-of-two stride: nbatch_K2 = 32, 64, 128, 256. test_cases.emplace_back(new test_flash_attn_ext( 64, 64, 8, {8, 1}, 4096, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {4, 1}, 4096, 8, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); 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)); - 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})); @@ -10482,10 +10481,12 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1}, 20000, 512, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); for (int kv : { 4096, 8192, 16384,32768, 65536, }) { - for (int hs : { 64, 128, 256, }) { + for (int hs : { 64, 128, 256, 576, }) { + const int hsv = hs == 576 ? 512 : hs; + const bool v_view = hs == 576; for (int nr : { 1, 4, 8, }) { for (int nb : { 1, 4096, }) { - test_cases.emplace_back(new test_flash_attn_ext(hs, hs, 8, {nr, 1}, kv, nb, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(hs, hsv, 8, {nr, 1}, kv, nb, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, v_view)); } } }