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)
     }