Commit 2ca15f540 for llama.cpp
commit 2ca15f5404760548c39e7b92bd43116a09414a1a
Author: Johannes Gäßler <johannesg@5d6.de>
Date: Sun Oct 4 22:48:30 2026 +0200
CUDA: refactor swizzling code (#29612)
* CUDA: refactor swizzling code
* fix templates/loop bounds
diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh
index 26ce46e97..9d221d23e 100644
--- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh
+++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh
@@ -329,32 +329,6 @@ 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.
@@ -372,6 +346,20 @@ static __host__ int get_cols_per_warp(const int cc) {
}
}
+static __host__ bool ggml_cuda_fattn_mma_get_swizzled(const int DKQ, const int DV, const int ncols, const int cc) {
+ return turing_mma_available(cc) &&
+ ggml_cuda_fattn_mma_get_nbatch_K2(DKQ, DV, ncols, cc) % 32 == 0 && ggml_cuda_fattn_mma_get_nbatch_V2(DKQ, DV, ncols, cc) % 32 == 0;
+}
+
+static constexpr __device__ bool ggml_cuda_fattn_mma_get_swizzled(const int DKQ, const int DV, const int ncols) {
+#ifdef TURING_MMA_AVAILABLE
+ return ggml_cuda_fattn_mma_get_nbatch_K2(DKQ, DV, ncols) % 32 == 0 && ggml_cuda_fattn_mma_get_nbatch_V2(DKQ, DV, ncols) % 32 == 0;
+#else
+ GGML_UNUSED_VARS(DKQ, DV, ncols);
+ return false;
+#endif // TURING_MMA_AVAILABLE
+}
+
// ------------------------------------------------------------------------------------------------------------------
static __host__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, const int DV, const int ncols1, const int ncols2, const int cc) {
@@ -392,14 +380,15 @@ static constexpr __device__ int ggml_cuda_fattn_mma_get_nstages(
// ------------------------------------------------------------------------------------------------------------------
-template<int stride_tile, bool swz, int nwarps, int nbatch_fa, bool use_cp_async, bool oob_check, bool use_sparse>
+template<int stride_tile, int nwarps, int nbatch_fa, bool use_cp_async, bool oob_check, bool use_sparse>
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 k_VKQ_0, const int i_sup, const int32_t * const __restrict__ indices) {
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
// K/V data is loaded with decreasing granularity for D for better memory bandwidth.
// The minimum granularity is 16 bytes.
- constexpr int h2_per_chunk = 16/sizeof(half2);
+ constexpr int chunk_size = 16;
+ constexpr int h2_per_chunk = chunk_size / sizeof(half2);
const int chunks_per_row = D2 / h2_per_chunk;
if constexpr (use_cp_async) {
static_assert(warp_size == 32, "bad warp_size");
@@ -439,7 +428,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);
- 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);
+ cp_async_cg_16<preload>(tile_KV_32 + swizzle<stride_tile*sizeof(half2), char>(i*stride_tile*sizeof(half2) + k*chunk_size, i), KV + i_KV*stride_KV + k*h2_per_chunk);
}
}
};
@@ -481,7 +470,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;
}
- ggml_cuda_memcpy_1<16>((char *) tile_KV + swizzle_bytes<swz, half2>(i, k*h2_per_chunk, stride_tile), src);
+ ggml_cuda_memcpy_1<16>(swizzle<stride_tile>(tile_KV, i*stride_tile + k*h2_per_chunk, i), src);
}
}
};
@@ -624,9 +613,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);
- 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 bool swz = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols);
+ constexpr int stride_tile_K = swz ? nbatch_K2 : nbatch_K2 + 4;
+ constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : (swz ? nbatch_V2 : nbatch_V2 + 4);
const int k_VKQ_0 = kb0 * nbatch_fa;
#if defined(TURING_MMA_AVAILABLE)
@@ -644,7 +633,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, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_V, 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
@@ -664,7 +653,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, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_K, 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();
@@ -680,7 +669,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;
- load_ldmatrix<swz>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start, stride_tile_K);
+ load_ldmatrix_swizzled<stride_tile_K>(K_A, tile_K, i_KQ_0*stride_tile_K + 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 {
@@ -706,7 +695,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;
- load_ldmatrix<swz>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start, stride_tile_K);
+ load_ldmatrix_swizzled<stride_tile_K>(K_A, tile_K, i_KQ_0*stride_tile_K + 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]);
@@ -1001,7 +990,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, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_K, 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);
}
}
@@ -1017,7 +1006,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, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_V, 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();
@@ -1025,7 +1014,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
__syncthreads();
}
}
- const half2 * tile_V_i = !V_is_K_view || i0_stop > 2*nbatch_K2 ? tile_V : tile_V + i0_start/2;
+ const int tile_V_offset_i = !V_is_K_view || i0_stop > 2*nbatch_K2 ? 0 : i0_start/2;
#if defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)
#pragma unroll
@@ -1036,7 +1025,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.
- load_ldmatrix_trans<swz>(A, tile_V, 2*k0, (int)(tile_V_i - tile_V) + (i_VKQ_0 - i0_start)/2, stride_tile_V);
+ load_ldmatrix_trans_swizzled<stride_tile_V>(A, tile_V, tile_V_offset_i + 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 {
@@ -1062,8 +1051,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.
- 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);
+ load_ldmatrix_swizzled<stride_tile_V>(A, tile_V, tile_V_offset_i + 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);
}
}
@@ -1253,10 +1241,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;
- 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 bool swz = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols);
+ constexpr int stride_tile_Q = DKQ/2 + 4;
+ constexpr int stride_tile_K = swz ? nbatch_K2 : nbatch_K2 + 4;
+ constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : (swz ? nbatch_V2 : nbatch_V2 + 4);
constexpr int stride_tile_KV_max = stride_tile_K > stride_tile_V ? stride_tile_K : stride_tile_V;
extern __shared__ half2 tile_Q[];
@@ -1354,7 +1342,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, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
+ flash_attn_ext_f16_load_tile<stride_tile_K, 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);
}
@@ -2039,9 +2027,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 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 bool swz = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols, cc);
+ const int stride_tile_K = swz ? nbatch_K2 : nbatch_K2 + 4;
+ const int stride_tile_V = V_is_K_view ? stride_tile_K : (swz ? nbatch_V2 : nbatch_V2 + 4);
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/mma.cuh b/ggml/src/ggml-cuda/mma.cuh
index 3583ba5e1..93a51aa0b 100644
--- a/ggml/src/ggml-cuda/mma.cuh
+++ b/ggml/src/ggml-cuda/mma.cuh
@@ -782,18 +782,27 @@ 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 <int stride, typename T>
+ static __device__ __forceinline__ uint32_t swizzle(const uint32_t offset, const uint32_t i) {
+ static_assert(sizeof(T) <= 4, "unsupported type size");
+ constexpr int stride_bytes = stride*sizeof(T);
+ static_assert(stride_bytes % 16 == 0, "bad stride");
+ constexpr uint32_t shift = sizeof(T) == 1 ? 4 : (sizeof(T) == 2 ? 3 : 2);
+ if (stride_bytes % 32 != 0) {
+ return offset; // Equivalent to padding with 16 bytes.
+ }
+ if (stride_bytes % 64 != 0) {
+ return offset ^ (((i / 4) % 2) << shift);
+ }
+ if (stride_bytes % 128 != 0) {
+ return offset ^ (((i / 2) % 4) << shift);
+ }
+ return offset ^ ((i % 8) << shift);
}
- 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 <int stride, typename T>
+ static __device__ __forceinline__ T * swizzle(T * ptr, const uint32_t offset, const uint32_t i) {
+ return ptr + swizzle<stride, T>(offset, i);
}
template <typename T>
@@ -872,29 +881,6 @@ 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, int J, typename T, data_layout dl>
- static __device__ __forceinline__ void load_ldmatrix(
- tile<I, J, 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)
- static_assert(I == 16, "bad tile width");
- static_assert(J == 8, "bad tile height");
- 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);
@@ -902,10 +888,15 @@ namespace ggml_cuda_mma {
static __device__ __forceinline__ void load_ldmatrix(
tile<8, 4, half2, DATA_LAYOUT_J_MAJOR_MIRRORED> & t, const half2 * __restrict__ xs0, const int stride) {
+#ifdef VOLTA_MMA_AVAILABLE
#pragma unroll
for (int l0 = 0; l0 < t.ne; l0 += 2) {
ggml_cuda_memcpy_1<2*sizeof(half2)>(t.x + l0, xs0 + t.get_i(l0)*stride + t.get_j(l0));
}
+#else
+ GGML_UNUSED_VARS(t, xs0, stride);
+ NO_DEVICE_CODE;
+#endif // VOLTA_MMA_AVAILABLE
}
static __device__ __forceinline__ void load_ldmatrix(
@@ -954,25 +945,112 @@ 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;
+ template <int stride, int I, int J, typename T, data_layout dl>
+ static __device__ __forceinline__ void load_ldmatrix_swizzled(
+ tile<I, J, T, dl> & t, const T * __restrict__ xs0, const int offset) {
+#if defined(TURING_MMA_AVAILABLE)
+ static_assert(I == 16, "bad tile width");
+ static_assert(J == 8, "bad tile height");
+ const int i = threadIdx.x % t.I;
+ const int j = (threadIdx.x / t.I) * (t.J / 2);
+ int offset_ij = offset + i * stride + j;
+ offset_ij = swizzle<stride, T>(offset_ij, i);
+ 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"(xs0 + offset_ij));
+#elif defined(VOLTA_MMA_AVAILABLE)
+#pragma unroll
+ for (int o = 0; o < t.ne; o += 4) {
+ const int offset_ij = offset + t.get_i(o) * stride + o;
+ ggml_cuda_memcpy_1<4*sizeof(T)>(t.x + o, swizzle<stride>(xs0, offset_ij, t.get_i(o)));
+ }
+#elif defined(AMD_WMMA_AVAILABLE)
+#ifdef RDNA3
+ static_assert(dl == DATA_LAYOUT_I_MAJOR_MIRRORED, "bad data layout");
+ static_assert(sizeof(t.x) == 32, "bad ne");
+ static_assert(I == 16, "bad tile width");
+ static_assert(J == 8, "bad tile height");
+#pragma unroll
+ for (int o = 0; o < 8; o += 4) {
+ const int offset_ij = offset + t.get_i(0) * stride + o;
+ ggml_cuda_memcpy_1<16>(t.x + o, swizzle<stride>(xs0, offset_ij, t.get_i(0)));
+ }
+#else
+ static_assert(dl == DATA_LAYOUT_I_MAJOR, "bad data layout");
+ static_assert(sizeof(t.x) == 16, "bad ne");
+ const int offset_ij = offset + t.get_i(0)*stride + t.get_j(0);
+ ggml_cuda_memcpy_1<16>(t.x, swizzle<stride>(xs0, offset_ij, t.get_i(0)));
+#endif // RDNA3
+#elif defined(AMD_MFMA_AVAILABLE)
+ static_assert(sizeof(t.x) == 8, "bad ne");
+ const int offset_ij = offset + t.get_i(0)*stride + t.get_j(0);
+ ggml_cuda_memcpy_1<8>(t.x, swizzle<stride>(xs0, offset_ij, t.get_i(0)));
+#else
+ GGML_UNUSED_VARS(t, xs0, offset);
+ NO_DEVICE_CODE;
+#endif // defined(TURING_MMA_AVAILABLE)
+ }
+
+ template <int stride>
+ static __device__ __forceinline__ void load_ldmatrix_swizzled(
+ tile<8, 4, half2, DATA_LAYOUT_J_MAJOR_MIRRORED> & t, const half2 * __restrict__ xs0, const int offset) {
+#ifdef VOLTA_MMA_AVAILABLE
+#pragma unroll
+ for (int l0 = 0; l0 < t.ne; l0 += 2) {
+ const int offset_ij = offset + t.get_i(l0)*stride + t.get_j(l0);
+ ggml_cuda_memcpy_1<2*sizeof(half2)>(t.x + l0, swizzle<stride>(xs0, offset_ij, t.get_i(l0)));
}
+#else
+ GGML_UNUSED_VARS(t, xs0, offset);
+ NO_DEVICE_CODE;
+#endif // VOLTA_MMA_AVAILABLE
+ }
+
+ template <int stride, int I, typename T, data_layout dl>
+ static __device__ __forceinline__ void load_ldmatrix_trans_swizzled(
+ tile<I, 8, T, dl> & t, const T * __restrict__ xs0, const int offset) {
#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);
+ const int i = threadIdx.x % t.I;
+ const int j = (threadIdx.x / t.I) * (t.J / 2);
+ int offset_ij = offset + i * stride + j;
+ offset_ij = swizzle<stride, T>(offset_ij, i);
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)));
+ : "l"(xs0 + offset_ij));
+#elif defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
+ static_assert(dl == DATA_LAYOUT_I_MAJOR || dl == DATA_LAYOUT_I_MAJOR_MIRRORED, "bad data layout");
+ if constexpr (I == 32) {
+#pragma unroll
+ for (int l0 = 0; l0 < t.ne/2; ++l0) {
+ half2 tmp[2];
+#pragma unroll
+ for (int o = 0; o < 2; ++o) {
+ const int j = 2*t.get_j(l0) + o;
+ int offset_ij = offset + j*stride + t.get_i(l0)/2;
+ offset_ij = swizzle<stride, T>(offset_ij, j);
+ tmp[o] = xs0[offset_ij];
+ }
+
+ t.x[l0] = __lows2half2(tmp[0], tmp[1]);
+ t.x[l0 + t.ne/2] = __highs2half2(tmp[0], tmp[1]);
+ }
+ } else {
+ half * xh = (half *) t.x;
+#pragma unroll
+ for (int l = 0; l < t.ne; ++l) {
+#pragma unroll
+ for (int o = 0; o < 2; ++o) {
+ const int j = 2*t.get_j(l) + o;
+ xh[2*l + o] = ((const half *) xs0)[swizzle<2*stride, half>(2*offset + j*(2*stride) + t.get_i(l), j)];
+ }
+ }
+ }
#else
- GGML_UNUSED_VARS(t, tile_base, i0, j0, stride);
+ GGML_UNUSED_VARS(t, xs0, offset);
NO_DEVICE_CODE;
#endif // defined(TURING_MMA_AVAILABLE)
}