Commit 1884824fd for llama.cpp
commit 1884824fdaeb19af27cd8ba0ca50fcbed28eca44
Author: ynankani <ynankani@nvidia.com>
Date: Mon Sep 21 10:58:28 2026 +0000
CUDA: Follow up of #25635, refactoring FA shared smem swizzle (#28536)
* remove explicit swz value in config and rebase
Signed-off-by: ynankani <ynankani@nvidia.com>
* address review comments
Signed-off-by: ynankani <ynankani@nvidia.com>
---------
Signed-off-by: ynankani <ynankani@nvidia.com>
diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh
index df7dd8615..84219ca92 100644
--- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh
+++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh
@@ -2,7 +2,6 @@
#include "cp-async.cuh"
#include "mma.cuh"
#include "fattn-common.cuh"
-#include "fattn-swizzle.cuh"
using namespace ggml_cuda_mma;
@@ -327,6 +326,32 @@ static constexpr __device__ bool ggml_cuda_fattn_mma_get_Q_in_reg(const int DKQ,
return ggml_cuda_fattn_mma_get_config(DKQ, DV, ncols).Q_in_reg;
}
+// Swizzling needs a tile stride that is a multiple of 32 half2 columns.
+static constexpr __host__ __device__ bool ggml_cuda_fattn_mma_bank_aligned(const int nbatch_2) {
+ return nbatch_2 >= 32 && nbatch_2 % 32 == 0;
+}
+
+// Swizzling needs ldmatrix, on other hardware the tiles keep the row padding.
+static __host__ bool ggml_cuda_fattn_mma_get_swizzled(const int DKQ, const int DV, const int ncols1, const int ncols2, const int cc) {
+ const fattn_mma_config cfg = ggml_cuda_fattn_mma_get_config(DKQ, DV, ncols1*ncols2, cc);
+ return turing_mma_available(cc) && ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_K2) && ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_V2);
+}
+
+static constexpr __device__ bool ggml_cuda_fattn_mma_get_swizzled(const int DKQ, const int DV, const int ncols1, const int ncols2) {
+#if defined(TURING_MMA_AVAILABLE)
+ const fattn_mma_config cfg = ggml_cuda_fattn_mma_get_config(DKQ, DV, ncols1*ncols2);
+ return ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_K2) && ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_V2);
+#else
+ GGML_UNUSED_VARS(DKQ, DV, ncols1, ncols2);
+ return false;
+#endif // defined(TURING_MMA_AVAILABLE)
+}
+
+// Row padding is only needed if the tile is not swizzled.
+static constexpr __host__ __device__ int ggml_cuda_fattn_mma_get_stride_tile(const int nbatch_2, const bool swizzled) {
+ return swizzled ? nbatch_2 : nbatch_2 + 4;
+}
+
static constexpr __device__ int get_cols_per_thread() {
#if defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)
return 1; // AMD has a single column per thread.
@@ -411,12 +436,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);
- if constexpr (swz) {
- const int smem_offs_b = ggml_cuda_fattn_smem_swizzle::bytes_rc<stride_tile>(i, k*h2_per_chunk);
- cp_async_cg_16<preload>(tile_KV_32 + smem_offs_b, KV + i_KV*stride_KV + k*h2_per_chunk);
- } else {
- cp_async_cg_16<preload>(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i_KV*stride_KV + k*h2_per_chunk);
- }
+ cp_async_cg_16<preload>(tile_KV_32 + swizzle_bytes<swz, half2>(i, k*h2_per_chunk, stride_tile), KV + i_KV*stride_KV + k*h2_per_chunk);
}
}
};
@@ -458,11 +478,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile(
} else {
src = !oob_check || i < i_sup ? KV + int64_t(k_VKQ_0 + i)*stride_KV + k*h2_per_chunk : zero;
}
- if constexpr (swz) {
- ggml_cuda_memcpy_1<16>((char *) tile_KV + ggml_cuda_fattn_smem_swizzle::bytes_rc<stride_tile>(i, k*h2_per_chunk), src);
- } else {
- ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, src);
- }
+ ggml_cuda_memcpy_1<16>((char *) tile_KV + swizzle_bytes<swz, half2>(i, k*h2_per_chunk, stride_tile), src);
}
}
};
@@ -605,11 +621,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, use_sparse);
- // swizzle the tile stride for K and V based on the batch size.
- 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);
+ constexpr bool swz = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols1, ncols2);
+ constexpr int stride_tile_K = ggml_cuda_fattn_mma_get_stride_tile(nbatch_K2, swz);
+ constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_mma_get_stride_tile(nbatch_V2, swz);
const int k_VKQ_0 = kb0 * nbatch_fa;
#if defined(TURING_MMA_AVAILABLE)
@@ -627,7 +641,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<stride_tile_V, swz_V, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_V, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(V_h2, tile_V, nbatch_V2, stride_V, k_VKQ_0, k_VKQ_sup, nullptr);
} else {
// the sparse mask values are gathered per element, always load them synchronously
@@ -647,7 +661,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<stride_tile_K, swz_K, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(K_h2 + k0_start, tile_K, k0_diff, stride_K, k_VKQ_0, k_VKQ_sup, indices);
if (use_cp_async) {
cp_async_wait_all();
@@ -663,7 +677,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<stride_tile_K, swz_K>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start);
+ load_ldmatrix<swz>(K_A, tile_K, i_KQ_0, 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 {
@@ -689,7 +703,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<stride_tile_K, swz_K>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start);
+ load_ldmatrix<swz>(K_A, tile_K, i_KQ_0, 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]);
@@ -984,7 +998,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
flash_attn_ext_f16_load_mask<ncols1, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(mask_h, tile_mask, stride_mask, k_VKQ_0 + nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr);
}
- flash_attn_ext_f16_load_tile<stride_tile_K, swz_K, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(K_h2, tile_K, nbatch_K2, stride_K, k_VKQ_0 + nbatch_fa, k_VKQ_sup, nullptr);
}
}
@@ -1000,7 +1014,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<stride_tile_V, swz_V, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_V, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(V_h2 + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_0, k_VKQ_sup, indices);
if (use_cp_async) {
cp_async_wait_all();
@@ -1019,7 +1033,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.
- ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans<stride_tile_V, swz_V>(A, tile_V, (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2);
+ load_ldmatrix_trans<swz>(A, tile_V, 2*k0, (int)(tile_V_i - 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 {
@@ -1045,7 +1059,8 @@ 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.
- ggml_cuda_fattn_smem_swizzle::load_ldmatrix<stride_tile_V, swz_V>(A, tile_V, (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2);
+ static_assert(!swz, "Volta has no ldmatrix");
+ 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);
}
}
@@ -1236,12 +1251,10 @@ 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;
- // swizzle the tile stride for K and V based on the batch size.
- 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 = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols1, ncols2);
+ constexpr int stride_tile_K = ggml_cuda_fattn_mma_get_stride_tile(nbatch_K2, swz);
+ constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_mma_get_stride_tile(nbatch_V2, swz);
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_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;
@@ -1338,7 +1351,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
flash_attn_ext_f16_load_mask<ncols1, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(mask_h, tile_mask, stride_mask, kb0*nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr);
}
- flash_attn_ext_f16_load_tile<stride_tile_K, swz_K, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(K_h2, tile_K, nbatch_K2, stride_K, kb0*nbatch_fa, k_VKQ_sup, nullptr);
}
@@ -1503,14 +1516,12 @@ 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
- if constexpr (combine_needs_sync) {
+ if constexpr (swz) {
__syncthreads();
}
@@ -1550,7 +1561,7 @@ 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)
- if constexpr (combine_needs_sync) {
+ if constexpr (swz) {
__syncthreads();
}
@@ -2025,8 +2036,9 @@ 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_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 bool swizzled = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols1, ncols2, cc);
+ const int stride_tile_K = ggml_cuda_fattn_mma_get_stride_tile(nbatch_K2, swizzled);
+ const int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_mma_get_stride_tile(nbatch_V2, swizzled);
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
deleted file mode 100644
index 44338c8db..000000000
--- a/ggml/src/ggml-cuda/fattn-swizzle.cuh
+++ /dev/null
@@ -1,126 +0,0 @@
-#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 multiple of 32 half2 columns, otherwise we keep +4 row padding.
-
-namespace ggml_cuda_fattn_smem_swizzle {
-
-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 bank_aligned(nbatch_2);
-#else
- GGML_UNUSED(nbatch_2);
- return false;
-#endif // defined(TURING_MMA_AVAILABLE)
-}
-
-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) && bank_aligned(nbatch_2);
-#endif // GGML_USE_HIP
-}
-
-static __device__ constexpr int tile_stride(const int nbatch_2) {
- return enabled(nbatch_2) ? 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<int stride_h2>
-static __device__ __forceinline__ int bytes_rc(const int row, const int col_h2) {
- 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);
-}
-
-// 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<int stride_h2>
-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(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);
- byte_off ^= (uint32_t) (((base_row + lane_row) & 7) << 4);
- return (const half2 *) ((const char *) tile_base + byte_off);
-}
-
-template<int stride_h2, bool swz, typename TileT>
-static __device__ __forceinline__ void load_ldmatrix(
- TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) {
- if constexpr (swz) {
- static_assert(std::is_same_v<TileT, ggml_cuda_mma::tile<16, 8, half2>>,
- "the swizzled layout is only supported for tile<16, 8, half2>");
- ldmatrix_x4((int *) t.x, lane_addr<stride_h2>(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<int stride_h2, bool swz, typename TileT>
-static __device__ __forceinline__ void load_ldmatrix(TileT & t, const half2 * tile_base, const int off_h2) {
- if constexpr (swz) {
- load_ldmatrix<stride_h2, swz>(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<int stride_h2, bool swz, typename TileT>
-static __device__ __forceinline__ void load_ldmatrix_trans(
- TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) {
- if constexpr (swz) {
- static_assert(std::is_same_v<TileT, ggml_cuda_mma::tile<16, 8, half2>>,
- "the swizzled layout is only supported for tile<16, 8, half2>");
- ldmatrix_x4_trans((int *) t.x, lane_addr<stride_h2>(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<int stride_h2, bool swz, typename TileT>
-static __device__ __forceinline__ void load_ldmatrix_trans(TileT & t, const half2 * tile_base, const int off_h2) {
- if constexpr (swz) {
- load_ldmatrix_trans<stride_h2, swz>(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
diff --git a/ggml/src/ggml-cuda/mma.cuh b/ggml/src/ggml-cuda/mma.cuh
index 8d7c69dc3..6af2b6a14 100644
--- a/ggml/src/ggml-cuda/mma.cuh
+++ b/ggml/src/ggml-cuda/mma.cuh
@@ -782,6 +782,20 @@ namespace ggml_cuda_mma {
}
}
+ // Byte offset of tile element (i, j). If swz, XOR swizzle it to avoid bank conflicts without row padding.
+ template <bool swz, typename T>
+ static __device__ __forceinline__ int swizzle_bytes(const int i, const int j, const int stride) {
+ static_assert(!swz || sizeof(T) == 4, "swizzled tiles need 32 bit elements");
+ const int off = (i*stride + j) * (int) sizeof(T);
+ return swz ? off ^ ((i & 7) << 4) : off;
+ }
+
+ template <bool swz, typename T>
+ static __device__ __forceinline__ const T * swizzle(
+ const T * __restrict__ tile_base, const int i, const int j, const int stride) {
+ return (const T *) ((const char *) tile_base + swizzle_bytes<swz, T>(i, j, stride));
+ }
+
template <typename T>
static __device__ __forceinline__ void load_ldmatrix(
tile<8, 8, T> & t, const T * __restrict__ xs0, const int stride) {
@@ -858,6 +872,27 @@ namespace ggml_cuda_mma {
#endif // TURING_MMA_AVAILABLE
}
+ // Load from tile element (i0, j0), swz tells if the tile is stored swizzled.
+ template <bool swz, typename T, data_layout dl>
+ static __device__ __forceinline__ void load_ldmatrix(
+ tile<16, 8, T, dl> & t, const T * __restrict__ tile_base, const int i0, const int j0, const int stride) {
+ if constexpr (!swz) {
+ load_ldmatrix(t, tile_base + i0*stride + j0, stride);
+ return;
+ }
+#if defined(TURING_MMA_AVAILABLE)
+ const int i = i0 + threadIdx.x % t.I;
+ const int j = j0 + (threadIdx.x / t.I) * (t.J / 2);
+ int * xi = (int *) t.x;
+ 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"(swizzle<true>(tile_base, i, j, stride)));
+#else
+ GGML_UNUSED_VARS(t, tile_base, i0, j0, stride);
+ NO_DEVICE_CODE;
+#endif // defined(TURING_MMA_AVAILABLE)
+ }
+
static __device__ __forceinline__ void load_ldmatrix(
tile<8, 4, half2, DATA_LAYOUT_I_MAJOR_MIRRORED> & t, const half2 * __restrict__ xs0, const int stride) {
ggml_cuda_memcpy_1<4*sizeof(half2)>(t.x, xs0 + t.get_i(0)*stride);
@@ -917,6 +952,29 @@ namespace ggml_cuda_mma {
#endif // TURING_MMA_AVAILABLE
}
+ // Load from tile element (i0, j0), swz tells if the tile is stored swizzled.
+ template <bool swz, int I, typename T, data_layout dl>
+ static __device__ __forceinline__ void load_ldmatrix_trans(
+ tile<I, 8, T, dl> & t, const T * __restrict__ tile_base, const int i0, const int j0, const int stride) {
+ if constexpr (!swz) {
+ load_ldmatrix_trans(t, tile_base + i0*stride + j0, stride);
+ return;
+ }
+#if defined(TURING_MMA_AVAILABLE)
+ static_assert(I == 16, "bad tile width");
+ static_assert(dl == DATA_LAYOUT_I_MAJOR, "bad data layout");
+ const int i = i0 + threadIdx.x % t.I;
+ const int j = j0 + (threadIdx.x / t.I) * (t.J / 2);
+ int * xi = (int *) t.x;
+ 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"(swizzle<true>(tile_base, i, j, stride)));
+#else
+ GGML_UNUSED_VARS(t, tile_base, i0, j0, stride);
+ NO_DEVICE_CODE;
+#endif // defined(TURING_MMA_AVAILABLE)
+ }
+
static __device__ __forceinline__ void mma(
tile<16, 8, int> & D, const tile<16, 4, int> & A, const tile<8, 4, int> & B) {
#ifdef TURING_MMA_AVAILABLE