From e4b9af007beae34abba094cefd187658e89bbac8 Mon Sep 17 00:00:00 2001 From: ynankani Date: Mon, 31 Aug 2026 20:18:01 +0000 Subject: [PATCH] CUDA: XOR swizzle flash attn K,V smem fp16 tiles (#25635) * CUDA: XOR swizzle flash attn K,V smem fp16 tiles Signed-off-by: ynankani * Fix use 64bit generic pointer instead of 32bit shared pointer Signed-off-by: ynankani * fix shared memory race in FA on DGX Spark * Handle corener case Signed-off-by: ynankani * Add swizzle test cases and gate sync for swizzled path only Signed-off-by: ynankani * gate CUDA PTX Signed-off-by: ynankani * offset calculation specific for swizzle branch Signed-off-by: ynankani * Reafctor code Signed-off-by: ynankani * Refactor FA swizzle ldmatrix if/else into helpers (K row/col, V offset) Signed-off-by: ynankani * rebase and update test case args Signed-off-by: ynankani * Allow swizzle for non-pow2 shapes, for which nbatch_2%32==0 Signed-off-by: ynankani --------- Signed-off-by: ynankani --- ggml/src/ggml-cuda/fattn-mma-f16.cuh | 72 ++++++++++----- ggml/src/ggml-cuda/fattn-swizzle.cuh | 126 +++++++++++++++++++++++++++ tests/test-backend-ops.cpp | 22 ++++- 3 files changed, 194 insertions(+), 26 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 7f4cfd5511..387e70fa14 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); @@ -360,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(); @@ -397,7 +398,12 @@ 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); + if constexpr (swz) { + 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); + } } } }; @@ -432,8 +438,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); - 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); + if constexpr (swz) { + 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, + !oob_check || i < i_sup ? KV + i*stride_KV + k*h2_per_chunk : zero); + } } } }; @@ -568,9 +579,11 @@ 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_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) @@ -588,7 +601,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; @@ -607,7 +620,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(); @@ -623,7 +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; - 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 +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; - 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]); @@ -943,7 +956,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); } } @@ -959,7 +972,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(); @@ -978,7 +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. - 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 { @@ -1004,7 +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. - 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); } } @@ -1168,10 +1181,12 @@ 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_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_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; @@ -1265,7 +1280,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); } @@ -1430,11 +1445,17 @@ 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) { + __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; @@ -1471,6 +1492,10 @@ 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) { + __syncthreads(); + } + if (((!needs_fixup && !is_fixup) || np > 1) && thread_should_write) { ((float2 *) tile_Q)[jc_cwm*(tile_stride/2) + nbatch_combine/2] = KQ_cmr; } @@ -1914,8 +1939,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_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); 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 0000000000..44338c8db0 --- /dev/null +++ b/ggml/src/ggml-cuda/fattn-swizzle.cuh @@ -0,0 +1,126 @@ +#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 +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 +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 +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>, + "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 +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( + 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 diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 2c865784a9..d61b379282 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10183,6 +10183,16 @@ 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, {0, 1, 2, 3}, true, true)); + + // 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})); @@ -10581,10 +10591,14 @@ 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, 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, 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)); + } } } }