diff --git a/csrc/flash_api/flash_api.cpp b/csrc/flash_api/flash_api.cpp index fe46743..b19304f 100644 --- a/csrc/flash_api/flash_api.cpp +++ b/csrc/flash_api/flash_api.cpp @@ -145,7 +145,7 @@ mha_fwd_kvcache_mla( at::Tensor vcache = vcache_.has_value() ? vcache_.value() : kcache; auto q_dtype = q.dtype(); - TORCH_CHECK(q_dtype == torch::kBFloat16); + TORCH_CHECK(q_dtype == torch::kBFloat16 || q_dtype == torch::kFloat16); TORCH_CHECK(kcache.dtype() == q_dtype, "query and key must have the same dtype"); CHECK_DEVICE(q); CHECK_DEVICE(kcache); CHECK_DEVICE(vcache); @@ -259,7 +259,7 @@ mha_fwd_kvcache_mla( auto stream = at::cuda::getCurrentCUDAStream().stream(); TORCH_CHECK(head_size == 576); - params.is_bf16 = true; + params.is_bf16 = q_dtype == torch::kBFloat16; run_mha_fwd(params,stream, /*force_split_kernel*/true); out = out.view({batch_size, seqlen_q_ori, ngroups, num_heads_k, head_size_v}).transpose(2, 3) .reshape({batch_size, seqlen_q_ori, num_heads_ori, head_size_v}); diff --git a/csrc/flash_dispatch/flash_fwd_dispatch_template.h b/csrc/flash_dispatch/flash_fwd_dispatch_template.h index b348974..2fffbdc 100644 --- a/csrc/flash_dispatch/flash_fwd_dispatch_template.h +++ b/csrc/flash_dispatch/flash_fwd_dispatch_template.h @@ -37,17 +37,19 @@ namespace mcFlashAttn { constexpr static int Num_Stages = 2; FP16_SWITCH(!params.is_bf16, [&] { - if (params.seqlen_q >= 64) { - constexpr static int kBlockM = 64; - constexpr static int kBlockN = 16; - constexpr static int kNWarps = 8; - run_flash_splitkv_fwd_template(params, stream); - } else { - constexpr static int kBlockM = 32; - constexpr static int kBlockN = 16; - constexpr static int kNWarps = 4; - run_flash_splitkv_fwd_template(params, stream); - } + BOOL_SWITCH(params.num_splits > 1, Is_splits, [&] { + if (params.seqlen_q >= 64) { + constexpr static int kBlockM = 64; + constexpr static int kBlockN = 16; + constexpr static int kNWarps = 8; + run_flash_splitkv_fwd_template(params, stream); + } else { + constexpr static int kBlockM = 32; + constexpr static int kBlockN = 16; + constexpr static int kNWarps = 4; + run_flash_splitkv_fwd_template(params, stream); + } + }); }); } diff --git a/csrc/flash_kernel/flash_fwd_split_kernel_k64_32x16_4waves.h b/csrc/flash_kernel/flash_fwd_split_kernel_k64_32x16_4waves.h index 71abe20..b5f5e57 100644 --- a/csrc/flash_kernel/flash_fwd_split_kernel_k64_32x16_4waves.h +++ b/csrc/flash_kernel/flash_fwd_split_kernel_k64_32x16_4waves.h @@ -363,14 +363,13 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa } // if (cute::thread0()) { print(lse); } if constexpr (!Split) { - // use smem for O (mtreg->smem->mtreg->global) Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N) // Partition sO to match the accumulator partitioning using SmemTiledCopyO = typename Kernel_traits::SmemCopyAtomO; CONVERT_TENSOR_TYPE(ElementAccum, ElementO, acc_o_copy, rO) int warp_offset = warp_idx * 16 * 64; int thread_offset = lane_idx % 16 * 64 + lane_idx / 16 * 16; - Element *Osmem_ptr_sts = reinterpret_cast(smem_) + warp_offset + thread_offset; + ElementO *Osmem_ptr_sts = reinterpret_cast(smem_) + warp_offset + thread_offset; Tensor tOsO = make_tensor(make_smem_ptr(Osmem_ptr_sts), make_layout(Shape<_16, _4>{}, Stride<_1, Int<16*64*kNWarps>>{})); Tensor tOrO = make_tensor(rO.data(), make_layout(Shape<_16, _4>{}, @@ -427,27 +426,47 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM ); } else { - // don't use smem for O (mtreg->global) + Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N) const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM) * params.d_v; const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM; - // Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast(params.oaccum_ptr) + row_offset_oaccum), - // Shape, Int>{}, - // make_stride(kHeadDimV, _1{})); + Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast(params.oaccum_ptr) + row_offset_oaccum), + Shape, Int>{}, + make_stride(kHeadDimV, _1{})); Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast(params.softmax_lseaccum_ptr) + row_offset_lseaccum), Shape>{}, Stride<_1>{}); // if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); } - Tensor taccOrOaccum = make_tensor(acc_o_copy.data(), make_layout(Shape<_16, _4>{}, - Stride<_1, _16>{})); - int warp_offset = warp_idx / kAtomLayoutMO * 64 + warp_idx % kAtomLayoutMO * 16 * kHeadDimV; - int thread_offset = lane_idx % 16 * kHeadDimV + lane_idx / 16 * 16; - ElementO *Osmem_ptr_stg = reinterpret_cast(params.oaccum_ptr) + row_offset_oaccum + warp_offset + thread_offset; - Tensor taccOgOaccum = make_tensor(make_gmem_ptr(Osmem_ptr_stg), make_layout(Shape<_16, _4>{}, - Stride<_1, Int<128>>{})); - // Tensor taccOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum); + Tensor taccOrOaccum = make_tensor(acc_o_copy.data(), acc_o_copy.layout()); + int warp_offset = warp_idx * 16 * 64; + int thread_offset = lane_idx % 16 * 64 + lane_idx / 16 * 16; + ElementO *accOsmem_ptr_sts = reinterpret_cast(smem_) + warp_offset + thread_offset; + Tensor taccOsOaccum = make_tensor(make_smem_ptr(accOsmem_ptr_sts), make_layout(Shape<_4, _4, _4>{}, + Stride<_1, _4, Int<16*64*kNWarps>>{})); + if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); } + int O_swizzle_row_sts = tidx % 4; + #pragma unroll + for (int i = 0; i < 4; i++) { + cute::copy(taccOrOaccum(_, make_coord(i, _)), taccOsOaccum(_, O_swizzle_row_sts ^ i, _)); + } + + GmemTiledCopyO gmem_tiled_copy_Oaccum; + auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx); + // Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N) + int O_swizzle_row_lds = tidx / 16 % 4; + int O_swizzle_col_lds = tidx % 16 % 4; + int O_swizzle_col_lds_new = O_swizzle_col_lds ^ O_swizzle_row_lds; + ElementO *accOsmem_ptr_lds = reinterpret_cast(smem_) + (tidx + O_swizzle_col_lds_new - O_swizzle_col_lds) * 4; + Tensor tOsOaccum = make_tensor(make_smem_ptr(accOsmem_ptr_lds), make_layout(Shape<_4, _2, Int>{}, + Stride<_1, Int<16*64>, Int<16*64*2>>{})); + Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum); + + flash::sync_threads(); + + Tensor tOrOaccum = make_tensor(shape(tOgOaccum)); + cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum); Tensor caccO = make_identity_tensor(Shape, Int>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k) Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K) @@ -463,9 +482,12 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_32x16_4wa } } + Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + // Repeat the partitioning with identity layouts + Tensor tOcaccO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k) // Clear_OOB_K must be false since we don't want to write zeros to gmem - flash::copy_reg_to_global4x4fp32( - taccOrOaccum, taccOgOaccum, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM + flash::copy_reg_to_global( + tOrOaccum, tOgOaccum, tOcaccO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM ); } } diff --git a/csrc/flash_kernel/flash_fwd_split_kernel_k64_64x16_8waves.h b/csrc/flash_kernel/flash_fwd_split_kernel_k64_64x16_8waves.h index 24b96f3..702569c 100644 --- a/csrc/flash_kernel/flash_fwd_split_kernel_k64_64x16_8waves.h +++ b/csrc/flash_kernel/flash_fwd_split_kernel_k64_64x16_8waves.h @@ -319,7 +319,6 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8wa } // if (cute::thread0()) { print(lse); } if constexpr (!Split) { - // use smem for O (mtreg->smem->mtreg->global) Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N) // Partition sO to match the accumulator partitioning using SmemTiledCopyO = typename Kernel_traits::SmemCopyAtomO; @@ -383,28 +382,53 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8wa tOrOaccum, tOgOaccum, tOcO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM ); } else { - // don't use smem for O (mtreg->global) + Tensor sOaccum = make_tensor(make_smem_ptr(reinterpret_cast(smem_)), typename Kernel_traits::SmemLayoutO{}); // (SMEM_M,SMEM_N) const index_t row_offset_oaccum = (((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM) * params.d_v; const index_t row_offset_lseaccum = ((n_split_idx * params.b + bidb) * params.h + bidh) * params.seqlen_q + m_block * kBlockM; - // Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast(params.oaccum_ptr) + row_offset_oaccum), - // Shape, Int>{}, - // make_stride(kHeadDimV, _1{})); + Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast(params.oaccum_ptr) + row_offset_oaccum), + Shape, Int>{}, + make_stride(kHeadDimV, _1{})); Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast(params.softmax_lseaccum_ptr) + row_offset_lseaccum), Shape>{}, Stride<_1>{}); // if (tidx == 0) { printf("row_offset_o = %d, bidh = %d, gOaccum = %p\n", row_offset_o, bidh, gOaccum.data()); } - Tensor taccOrOaccum = make_tensor(acc_o_copy.data(), make_layout(Shape<_16, _4>{}, - Stride<_1, _16>{})); - int warp_offset = warp_idx / kAtomLayoutMO * 64 + warp_idx % kAtomLayoutMO * 16 * kHeadDimV; - int thread_offset = lane_idx % 16 * kHeadDimV + lane_idx / 16 * 16; - ElementO *Osmem_ptr_stg = reinterpret_cast(params.oaccum_ptr) + row_offset_oaccum + warp_offset + thread_offset; - Tensor taccOgOaccum = make_tensor(make_gmem_ptr(Osmem_ptr_stg), make_layout(Shape<_16, _4>{}, - Stride<_1, Int<128>>{})); - // Tensor taccOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum); + Tensor taccOrOaccum = make_tensor(acc_o_copy.data(), acc_o_copy.layout()); + + int warp_offset = warp_idx * 16 * 64; + int thread_offset = lane_idx % 16 * 64 + lane_idx / 16 * 16; + ElementO *accOsmem_ptr_sts = reinterpret_cast(smem_) + warp_offset + thread_offset; + Tensor taccOsOaccum = make_tensor(make_smem_ptr(accOsmem_ptr_sts), make_layout(Shape<_4, _4, _2>{}, + Stride<_1, _4, Int<16*64*kNWarps>>{})); + if constexpr (Kernel_traits::Share_Q_K_smem) { flash::sync_threads(); } + int O_swizzle_row_sts = tidx % 4; + #pragma unroll + for (int k = 0; k < 2; k++) { + #pragma unroll + for (int i = 0; i < 4; i++) { + cute::copy(taccOrOaccum(_, make_coord(i, k)), taccOsOaccum(_, O_swizzle_row_sts ^ i, k)); + } + } + + GmemTiledCopyO gmem_tiled_copy_Oaccum; + auto gmem_thr_copy_Oaccum = gmem_tiled_copy_Oaccum.get_thread_slice(tidx); + // Tensor tOsOaccum = gmem_thr_copy_Oaccum.partition_S(sOaccum); // ((Atom,AtomNum),ATOM_M,ATOM_N) + int O_swizzle_row_lds = tidx / 16 % 4; + int O_swizzle_col_lds = tidx % 16 % 4; + int O_swizzle_col_lds_new = O_swizzle_col_lds ^ O_swizzle_row_lds; + ElementO *accOsmem_ptr_lds = reinterpret_cast(smem_) + (tidx + O_swizzle_col_lds_new - O_swizzle_col_lds) * 4; + Tensor tOsOaccum = make_tensor(make_smem_ptr(accOsmem_ptr_lds), make_layout(Shape<_4, _2, Int>{}, + Stride<_1, Int<32*64>, Int<32*64*2>>{})); + Tensor tOgOaccum = gmem_thr_copy_Oaccum.partition_D(gOaccum); + + flash::sync_threads(); + + Tensor tOrOaccum = make_tensor(shape(tOgOaccum)); + cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum); + flash::sync_threads(); Tensor caccO = make_identity_tensor(Shape, Int>{}); // (BLK_M,BLK_K) -> (blk_m,blk_k) Tensor taccOcO = thr_mma_o.partition_C(caccO); // (MMA,MMA_M,MMA_K) static_assert(decltype(size<0>(taccOcO))::value == 4); @@ -419,10 +443,31 @@ __forceinline__ __device__ void compute_attn_1rowblock_splitkv_k64_mla_64x16_8wa } } + Tensor cO = make_identity_tensor(make_shape(size<0>(sOaccum), size<1>(sOaccum))); // (BLK_M,BLK_K) -> (blk_m,blk_k) + // Repeat the partitioning with identity layouts + Tensor tOcaccO = gmem_thr_copy_Oaccum.partition_D(cO); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k) // Clear_OOB_K must be false since we don't want to write zeros to gmem - flash::copy_reg_to_global4x4fp32( - taccOrOaccum, taccOgOaccum, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM + flash::copy_reg_to_global( + tOrOaccum, tOgOaccum, tOcaccO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM ); + + // left global O data + #pragma unroll + for (int k = 2; k < 4; k++) { + #pragma unroll + for (int i = 0; i < 4; i++) { + cute::copy(taccOrOaccum(_, make_coord(i, k)), taccOsOaccum(_, O_swizzle_row_sts ^ i, k - 2)); + } + } + flash::sync_threads(); + cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum); + tOgOaccum.data() = tOgOaccum.data() + (kHeadDimV/2); + + + flash::copy_reg_to_global( + tOrOaccum, tOgOaccum, tOcaccO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM + ); + } } diff --git a/csrc/flash_kernel/kernel_traits.h b/csrc/flash_kernel/kernel_traits.h index e9b61bd..a52bf0d 100644 --- a/csrc/flash_kernel/kernel_traits.h +++ b/csrc/flash_kernel/kernel_traits.h @@ -160,7 +160,7 @@ struct Flash_fwd_kernel_traits : public Base { static constexpr int kSmemKSize = size(SmemLayoutK{}) * sizeof(Element); static constexpr int kSmemVSize = size(SmemLayoutV{}) * sizeof(Element); static constexpr int kSmemKVSize = kSmemKSize + kSmemVSize; - static constexpr int kSmemSize = Share_Q_K_smem ? std::max(std::min(kSmemQSize, 64 * 1024), kSmemKSize) : kSmemQSize + kSmemKSize; + static constexpr int kSmemSize = Share_Q_K_smem ? std::max(std::min((Is_Splits_ ? 2 : 1) * kSmemQSize, 64 * 1024), kSmemKSize) : kSmemQSize + kSmemKSize; static constexpr int kRegSize = kSmemSize / sizeof(uint32_t) / kNThreads; static constexpr int kGmemElemsPerLoadB128 = sizeof(cute::uint128_t) / sizeof(Element); @@ -222,7 +222,7 @@ struct Flash_fwd_kernel_traits : public Base { Stride< _16, _1>> >; using GmemTiledCopyOaccum = decltype( - make_tiled_copy(Copy_Atom{}, + make_tiled_copy(Copy_Atom, ElementAccum>{}, GmemLayoutAtomOaccum{}, Layout>{})); // Val layout, 4 vals per store }; diff --git a/csrc/utils/utils.h b/csrc/utils/utils.h index 89a6870..1685188 100644 --- a/csrc/utils/utils.h +++ b/csrc/utils/utils.h @@ -236,10 +236,7 @@ __forceinline__ __device__ void gemm_opt(Tensor0 &acc, Tensor1 &tCrA, Tensor2 &t if (!A_in_regs) { cute::copy(smem_tiled_copy_A, tCsA(_, _, i + 1), tCrA_copy_view(_, _, i + 1)); } if (!B_in_regs) { cute::copy(smem_tiled_copy_B, tCsB(_, _, i + 1), tCrB_copy_view(_, _, i + 1)); } } - //ToDo: remove this after compiler has been updated - __builtin_mxc_schedbound_begin(); cute::gemm(tiled_mma, tCrA(_, _, i), tCrB(_, _, i), acc); - __builtin_mxc_schedbound_end(); } } @@ -631,12 +628,12 @@ __forceinline__ __device__ void apply_softcap(Tensor &tensor, co // resolves offset of a slice of a paged kv copy from gmem. // assumes that the tensor has already been positioned at the correct head. __forceinline__ __device__ -int resolve_thread_kv_page_slice_offset(const int page_block_size, const int* block_table, const int page_stride, const int row_stride, const int row_offset, const int col_offset) { +int64_t resolve_thread_kv_page_slice_offset(const int page_block_size, const int* block_table, const int page_stride, const int row_stride, const int row_offset, const int col_offset) { const int virtual_page_idx = row_offset / page_block_size; const int page_offset = row_offset - virtual_page_idx * page_block_size; - return block_table[virtual_page_idx] * page_stride - + page_offset * row_stride + return ((int64_t) block_table[virtual_page_idx]) * ((int64_t) page_stride) + + page_offset * ((int64_t) row_stride) + col_offset; } @@ -674,7 +671,7 @@ __forceinline__ __device__ void copy_b128_page_one(Tensor cons bool row_mask = Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN; const int row_offset = tidx / kGmemThreadsPerRow * kGmemRowsPerThread + kNThreads / kGmemThreadsPerRow * m + n_block * kBlockN; const int col_offset = tidx % kGmemThreadsPerRow * kElementPerThread; - const int global_kv_page_offset = flash::resolve_thread_kv_page_slice_offset(page_block_size, block_table, page_stride, row_stride, row_offset, col_offset); + const int64_t global_kv_page_offset = flash::resolve_thread_kv_page_slice_offset(page_block_size, block_table, page_stride, row_stride, row_offset, col_offset); #pragma unroll for (int k = 0; k < size<2>(S); ++k) { auto src_ptr = (VecType *)(S_base.data().get() + global_kv_page_offset + get<2>(S.stride()) * k); @@ -724,7 +721,7 @@ __forceinline__ __device__ void copy_b64_page_one(Tensor const bool row_mask = Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN; const int row_offset = tidx / kGmemThreadsPerRow * kGmemRowsPerThread + kNThreads / kGmemThreadsPerRow * m + n_block * kBlockN; const int col_offset = tidx % kGmemThreadsPerRow * kElementPerThread; - const int global_kv_page_offset = flash::resolve_thread_kv_page_slice_offset(page_block_size, block_table, page_stride, row_stride, row_offset, col_offset); + const int64_t global_kv_page_offset = flash::resolve_thread_kv_page_slice_offset(page_block_size, block_table, page_stride, row_stride, row_offset, col_offset); #pragma unroll for (int k = 0; k < size<2>(S); ++k) { auto src_ptr = (VecType *)(S_base.data().get() + global_kv_page_offset + get<2>(S.stride()) * k); @@ -772,7 +769,7 @@ __forceinline__ __device__ void copy_b32_page_one(Tensor const bool row_mask = Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN; const int row_offset = tidx / kGmemThreadsPerRow * kGmemRowsPerThread + kNThreads / kGmemThreadsPerRow * m + n_block * kBlockN; const int col_offset = tidx % kGmemThreadsPerRow * kElementPerThread; - const int global_kv_page_offset = flash::resolve_thread_kv_page_slice_offset(page_block_size, block_table, page_stride, row_stride, row_offset, col_offset); + const int64_t global_kv_page_offset = flash::resolve_thread_kv_page_slice_offset(page_block_size, block_table, page_stride, row_stride, row_offset, col_offset); #pragma unroll for (int k = 0; k < size<2>(S); ++k) { auto src_ptr = (VecType *)(S_base.data().get() + global_kv_page_offset + get<2>(S.stride()) * k); @@ -843,29 +840,6 @@ __forceinline__ __device__ void lds4x4_with_swizzle424(Tensor0 const& tCsA, Tens } } - -// resolves offset of a slice of a paged kv copy from gmem. -// assumes that the tensor has already been positioned at the correct head. -template -__forceinline__ __device__ -int resolve_thread_kv_page_slice_offset(const int tidx, const int n_block_max, const int page_block_size, - const int* block_table, const int page_stride, const int row_stride, const int row_idx = 0) { - constexpr int kGmemThreadsPerRow = Kernel_traits::kGmemThreadsPerRow; - constexpr int kGmemRowsPerThread = Kernel_traits::kGmemRowsPerThread; - constexpr int kGmemElemsPerLoad = Kernel_traits::kGmemElemsPerLoad; - constexpr int kBlockN = Kernel_traits::kBlockN; - - const int col_offset = tidx % kGmemThreadsPerRow * kGmemElemsPerLoad; - const int block_row_offset = tidx / kGmemThreadsPerRow * kGmemRowsPerThread; - const int global_row_offset = block_row_offset + (n_block_max - 1) * kBlockN; - const int page_offset = global_row_offset % page_block_size; - const int virtual_page_idx = global_row_offset / page_block_size + row_idx; - - return block_table[virtual_page_idx] * page_stride - + page_offset * row_stride - + col_offset; -} - template __forceinline__ __device__ decltype(auto) permute_4x4_b16(Tensor &t) { using data_type = typename Engine::value_type; diff --git a/tests/test_flash_mla.py b/tests/test_flash_mla.py index 9c68e08..2091dcb 100644 --- a/tests/test_flash_mla.py +++ b/tests/test_flash_mla.py @@ -6,7 +6,6 @@ import torch import triton import pytest -# from flash_mla import get_mla_metadata, flash_mla_with_kvcache from flash_mla import ( get_mla_metadata, flash_mla_with_kvcache @@ -93,15 +92,11 @@ def test_flash_mla(b, s_q, mean_sk, h_q, h_kv, d, dv, causal, varlen, block_size cal_diff(out_flash, out_torch, "out") cal_diff(lse_flash, lse_torch, "lse") - t = triton.testing.do_bench(flash_mla) - FLOPS = s_q * total_seqlens * h_q * (d + dv) * 2 - bytes = (total_seqlens * h_kv * d + b * s_q * h_q * d + b * s_q * h_q * dv) * (torch.finfo(dtype).bits // 8) - print(f"{b}, {s_q}, {mean_sk}, {h_q}, {h_kv}, {d}, {dv}, {causal}, {varlen}, {t:.3f}, {FLOPS / 10 ** 9 / t:.0f}, {bytes / 10 ** 6 / t:.0f}") @pytest.mark.parametrize("b",[128]) @pytest.mark.parametrize("s_q",[1,2]) -@pytest.mark.parametrize("mean_sk",[4096, 8192]) +@pytest.mark.parametrize("mean_sk",[4096, 8192, 32768]) @pytest.mark.parametrize("h_q",[16, 32, 64, 128]) @pytest.mark.parametrize("h_kv",[1]) @pytest.mark.parametrize("d",[576]) @@ -109,7 +104,14 @@ def test_flash_mla(b, s_q, mean_sk, h_q, h_kv, d, dv, causal, varlen, block_size @pytest.mark.parametrize("causal",[True]) @pytest.mark.parametrize("varlen",[True,False]) @pytest.mark.parametrize("block_size",[1,4,16,64]) -def test_flash_mla_checkin(b, s_q, mean_sk, h_q, h_kv, d, dv, causal, varlen, block_size): +@pytest.mark.parametrize("dtype",[torch.bfloat16, torch.float16]) +def test_flash_mla_checkin(b, s_q, mean_sk, h_q, h_kv, d, dv, causal, varlen, block_size, dtype): + device = torch.device("cuda:0") + torch.set_default_dtype(dtype) + torch.set_default_device(device) + torch.cuda.set_device(device) + torch.manual_seed(0) + random.seed(0) test_flash_mla(b, s_q, mean_sk, h_q, h_kv, d, dv, causal, varlen, block_size) if __name__ == "__main__": @@ -120,14 +122,13 @@ if __name__ == "__main__": torch.cuda.set_device(device) torch.manual_seed(0) random.seed(0) - print(f"batch, seqlen_q, mean_sk, h_q, h_kv, d, dv, causal, varlen, time(ms), TFLOPS, bandwith(GB/s)") h_kv = 1 d, dv = 576, 512 causal = True for block_size in [1,4,16,64]: for b in [128]: - for s in [4096, 8192]: + for s in [4096, 8192, 32768]: for h_q in [16, 32, 64, 128]: # TP = 8, 4, 2, 1 for s_q in [1, 2]: # MTP =1, 2 for varlen in [False, True]: