Fix: fix large input length offset overflow
This commit is contained in:
parent
48b347b794
commit
2e747f0a42
|
|
@ -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});
|
||||
|
|
|
|||
|
|
@ -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<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, false, HeaddimVO, Num_Stages>(params, stream);
|
||||
} else {
|
||||
constexpr static int kBlockM = 32;
|
||||
constexpr static int kBlockN = 16;
|
||||
constexpr static int kNWarps = 4;
|
||||
run_flash_splitkv_fwd_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, false, HeaddimVO, Num_Stages>(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<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, Is_splits, HeaddimVO, Num_Stages>(params, stream);
|
||||
} else {
|
||||
constexpr static int kBlockM = 32;
|
||||
constexpr static int kBlockN = 16;
|
||||
constexpr static int kNWarps = 4;
|
||||
run_flash_splitkv_fwd_template<HeaddimQK, kBlockM, kBlockN, kNWarps, true, true, elem_type, Is_splits, HeaddimVO, Num_Stages>(params, stream);
|
||||
}
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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<ElementO *>(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<ElementO *>(smem_) + warp_offset + thread_offset;
|
||||
ElementO *Osmem_ptr_sts = reinterpret_cast<ElementO *>(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<ElementO *>(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<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
|
||||
// Shape<Int<kBlockM>, Int<kHeadDimV>>{},
|
||||
// make_stride(kHeadDimV, _1{}));
|
||||
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
|
||||
Shape<Int<kBlockM>, Int<kHeadDimV>>{},
|
||||
make_stride(kHeadDimV, _1{}));
|
||||
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lseaccum),
|
||||
Shape<Int<kBlockM>>{}, 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<ElementO *>(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<ElementO *>(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<ElementO *>(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<kHeadDimV/64>>{},
|
||||
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<ElementO>(shape(tOgOaccum));
|
||||
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
|
||||
|
||||
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (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<Kernel_traits, Is_even_MN, Is_even_K>(
|
||||
taccOrOaccum, taccOgOaccum, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
|
||||
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
|
||||
tOrOaccum, tOgOaccum, tOcaccO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<ElementO *>(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<ElementO *>(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<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
|
||||
// Shape<Int<kBlockM>, Int<kHeadDimV>>{},
|
||||
// make_stride(kHeadDimV, _1{}));
|
||||
Tensor gOaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementO *>(params.oaccum_ptr) + row_offset_oaccum),
|
||||
Shape<Int<kBlockM>, Int<kHeadDimV/2>>{},
|
||||
make_stride(kHeadDimV, _1{}));
|
||||
Tensor gLSEaccum = make_tensor(make_gmem_ptr(reinterpret_cast<ElementAccum *>(params.softmax_lseaccum_ptr) + row_offset_lseaccum),
|
||||
Shape<Int<kBlockM>>{}, 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<ElementO *>(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<ElementO *>(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<ElementO *>(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<kHeadDimV/2/64>>{},
|
||||
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<ElementO>(shape(tOgOaccum));
|
||||
cute::copy(gmem_tiled_copy_Oaccum, tOsOaccum, tOrOaccum);
|
||||
flash::sync_threads();
|
||||
Tensor caccO = make_identity_tensor(Shape<Int<kBlockM>, Int<kHeadDimV>>{}); // (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<Kernel_traits, Is_even_MN, Is_even_K>(
|
||||
taccOrOaccum, taccOgOaccum, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
|
||||
flash::copy_reg_to_global<Is_even_MN, Is_even_K>(
|
||||
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<Is_even_MN, Is_even_K>(
|
||||
tOrOaccum, tOgOaccum, tOcaccO, params.d_v, binfo.actual_seqlen_q - m_block * kBlockM
|
||||
);
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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<DefaultCopy, ElementAccum>{},
|
||||
make_tiled_copy(Copy_Atom<UniversalCopy<uint128_t>, ElementAccum>{},
|
||||
GmemLayoutAtomOaccum{},
|
||||
Layout<Shape < _1, _4>>{})); // Val layout, 4 vals per store
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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<Engine, Layout> &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<Engine0, Layout0> 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<Engine0, Layout0> 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<Engine0, Layout0> 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 <typename Kernel_traits>
|
||||
__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 <typename Engine, typename Layout>
|
||||
__forceinline__ __device__ decltype(auto) permute_4x4_b16(Tensor<Engine, Layout> &t) {
|
||||
using data_type = typename Engine::value_type;
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
Loading…
Reference in New Issue