diff --git a/3rdparty/composable_kernel b/3rdparty/composable_kernel index 10cb6916c34..207a95d5e40 160000 --- a/3rdparty/composable_kernel +++ b/3rdparty/composable_kernel @@ -1 +1 @@ -Subproject commit 10cb6916c34f957e81e8472c085603b4427baab9 +Subproject commit 207a95d5e4081316f3fb18a035b3918c118367c4 diff --git a/csrc/cpp_itfs/mha_bwd.cu b/csrc/cpp_itfs/mha_bwd.cu index 69f55db0e49..fba1c4d075f 100644 --- a/csrc/cpp_itfs/mha_bwd.cu +++ b/csrc/cpp_itfs/mha_bwd.cu @@ -3,6 +3,7 @@ #include "asm_fmha_v3_bwd_configs.hpp" #include #include +#include namespace aiter { std::tuple get_padded_hdim(int hdim_q, int hdim_v, std::string arch_id) @@ -133,6 +134,36 @@ float mha_bwd(mha_bwd_args a, const ck_tile::stream_config& s) #if ONLY_FAV3 return asm_ret; #else // !ONLY_FAV3 + if(asm_ret != -1) + return asm_ret; + + // For group mode, the new launcher needs seqstart arrays on host during construction. + // Locally D2H-copy them; the host buffers stay alive only for traits/launcher + // construction below, then are released — they are not retained by the launcher. + std::vector seqstart_q_host, seqstart_k_host; + const int* seqstart_qs_ptr = nullptr; + const int* seqstart_ks_ptr = nullptr; + if(a.is_group_mode) + { + if(a.seqstart_q_ptr == nullptr || a.seqstart_k_ptr == nullptr) + { + AITER_LOG_ERROR("mha_bwd: group mode requires seqstart_q_ptr and seqstart_k_ptr"); + return -1; + } + seqstart_q_host.resize(a.batch + 1); + seqstart_k_host.resize(a.batch + 1); + HIP_CALL(hipMemcpy(seqstart_q_host.data(), + a.seqstart_q_ptr, + sizeof(int) * (a.batch + 1), + hipMemcpyDeviceToHost)); + HIP_CALL(hipMemcpy(seqstart_k_host.data(), + a.seqstart_k_ptr, + sizeof(int) * (a.batch + 1), + hipMemcpyDeviceToHost)); + seqstart_qs_ptr = seqstart_q_host.data(); + seqstart_ks_ptr = seqstart_k_host.data(); + } + const fmha_bwd_traits traits{ a.seqlen_q, a.seqlen_k, @@ -151,8 +182,14 @@ float mha_bwd(mha_bwd_args a, const ck_tile::stream_config& s) a.has_dropout, a.is_store_randval, a.is_deterministic, + seqstart_qs_ptr, + seqstart_ks_ptr, }; + const fmha_bwd_launcher launcher(traits); + void* workspace_ptr = a.workspace_alloc(launcher.workspace_size, /*zero_init=*/false); + launcher.prepare_workspace(workspace_ptr); + fmha_bwd_args ck_args{ /* q_ptr */ a.q_ptr, /* k_ptr */ a.k_ptr, @@ -167,7 +204,7 @@ float mha_bwd(mha_bwd_args a, const ck_tile::stream_config& s) /* dk_ptr */ a.dk_ptr, /* dv_ptr */ a.dv_ptr, /* dbias_ptr */ a.dbias_ptr, - /* dq_acc_ptr */ a.dq_acc_ptr, + /* workspace_ptr */ workspace_ptr, /* sink_ptr */ a.sink_ptr, /* d_sink_ptr */ a.d_sink_ptr, @@ -196,7 +233,6 @@ float mha_bwd(mha_bwd_args a, const ck_tile::stream_config& s) /* stride_o */ a.stride_o, /* stride_randval */ a.stride_randval, /* stride_do */ a.stride_do, - /* stride_dq_acc */ a.stride_dq_acc, /* stride_dq */ a.stride_dq, /* stride_dk */ a.stride_dk, /* stride_dv */ a.stride_dv, @@ -210,7 +246,6 @@ float mha_bwd(mha_bwd_args a, const ck_tile::stream_config& s) /* nhead_stride_randval*/ a.nhead_stride_randval, /* nhead_stride_do */ a.nhead_stride_do, /* nhead_stride_lsed */ a.nhead_stride_lsed, - /* nhead_stride_dq_acc*/ static_cast(a.nhead_stride_dq_acc), /* nhead_stride_dq */ a.nhead_stride_dq, /* nhead_stride_dk */ a.nhead_stride_dk, /* nhead_stride_dv */ a.nhead_stride_dv, @@ -224,13 +259,11 @@ float mha_bwd(mha_bwd_args a, const ck_tile::stream_config& s) /* batch_stride_randval*/ a.batch_stride_randval, /* batch_stride_do */ a.batch_stride_do, /* batch_stride_lsed */ a.batch_stride_lsed, - /* batch_stride_dq_acc*/ static_cast(a.batch_stride_dq_acc), /* batch_stride_dq */ a.batch_stride_dq, /* batch_stride_dk */ a.batch_stride_dk, /* batch_stride_dv */ a.batch_stride_dv, /* batch_stride_dbias */ a.batch_stride_dbias, - /* split_stride_dq_acc*/ a.split_stride_dq_acc, /* window_size_left */ a.window_size_left, /* window_size_right */ a.window_size_right, /* mask_type */ a.mask_type, @@ -239,21 +272,12 @@ float mha_bwd(mha_bwd_args a, const ck_tile::stream_config& s) /* drop_seed_offset */ a.drop_seed_offset, }; - if(asm_ret == -1) - { - return fmha_bwd(traits, ck_args, s); - } - return asm_ret; + return launcher.run(ck_args, s); #endif } float fmha_v3_bwd(mha_bwd_args a, const ck_tile::stream_config& s) { - if(a.nhead_stride_dq_acc < a.stride_dq_acc) - { - return -1; // dq_acc only support BHSD layout - } - std::string arch_id = get_gpu_arch(); if((!a.use_asm_v3) || (a.hdim_q % 8 != 0) || (a.hdim_v % 8 != 0) || (a.has_dbias) || (a.bias_type != 0) || (a.has_dropout) || (a.is_deterministic) || @@ -457,8 +481,40 @@ float fmha_v3_bwd(mha_bwd_args a, const ck_tile::stream_config& s) impl_ptr_pre->launch_kernel({&odo_args, &arg_size, gdx, gdy, gdz, bdx, 1, 1, s.stream_id_}); }; + // ASM dq_accum layout (a.seqlen_q is total_q in group mode): + // atomic_fp32 batch: (1, batch, nhead_q, seqlen_q, hdim_q) [fp32] + // atomic16 batch: (1, batch, nhead_q, ceil16(seqlen_q), pad_hdim_q) [q-dtype] + // atomic_fp32 group: (1, nhead_q, total_q, hdim_q) [fp32] + // atomic16 group: (1, batch, nhead_q, ceil16(max_seqlen_q), 128) [q-dtype] + void* dq_acc_ptr = nullptr; + size_t stride_dq_acc = 0; + size_t nhead_stride_dq_acc = 0; + size_t batch_stride_dq_acc = 0; + size_t dq_acc_element_size = 0; + + if(need_post_processing) + { + dq_acc_element_size = a.v3_atomic_fp32 ? 4 : 2; + + const size_t a16_pad_seq = (a.max_seqlen_q + 15) / 16 * 16; + const size_t a16_pad_hdim = a.hdim_q == 192 ? 192 : 128; + const size_t dq_acc_seq = a.v3_atomic_fp32 ? a.seqlen_q : a16_pad_seq; + const size_t dq_acc_hdim = a.v3_atomic_fp32 ? a.hdim_q : a16_pad_hdim; + + const size_t per_batch_elems = static_cast(a.nhead_q) * dq_acc_seq * dq_acc_hdim; + const size_t effective_batch = (a.is_group_mode && a.v3_atomic_fp32) ? 1 : a.batch; + + stride_dq_acc = dq_acc_hdim; + nhead_stride_dq_acc = dq_acc_seq * dq_acc_hdim; + batch_stride_dq_acc = (effective_batch == 1) ? 0 : per_batch_elems; + const size_t dq_acc_bytes = effective_batch * per_batch_elems * dq_acc_element_size; + + // ASM kernel atomically accumulates into dq_accum; require zero-init. + dq_acc_ptr = a.workspace_alloc(dq_acc_bytes, /*zero_init=*/true); + } + fmha_bwd_dqdkdv_args dqdkdv_args; - dqdkdv_args.ptr_dq = need_post_processing ? a.dq_acc_ptr : a.dq_ptr; + dqdkdv_args.ptr_dq = need_post_processing ? dq_acc_ptr : a.dq_ptr; dqdkdv_args.ptr_dk = a.dk_ptr; dqdkdv_args.ptr_dv = a.dv_ptr; dqdkdv_args.ptr_q = a.q_ptr; @@ -568,14 +624,13 @@ float fmha_v3_bwd(mha_bwd_args a, const ck_tile::stream_config& s) [=](const ck_tile::stream_config& s_) { dqdkdv_kernel_launch(); }); } - int dq_acc_element_size = a.v3_atomic_fp32 ? 4 : 2; fmha_bwd_post_kernel_args post_args; - post_args.ptr_dq_acc = a.dq_acc_ptr; + post_args.ptr_dq_acc = dq_acc_ptr; post_args.ptr_dq = a.dq_ptr; - post_args.Hs_dq_acc = a.nhead_stride_dq_acc * dq_acc_element_size; - post_args.BAs_dq_acc = a.batch_stride_dq_acc * dq_acc_element_size; - post_args.Seqs_dq_acc = a.stride_dq_acc * dq_acc_element_size; + post_args.Hs_dq_acc = nhead_stride_dq_acc * dq_acc_element_size; + post_args.BAs_dq_acc = batch_stride_dq_acc * dq_acc_element_size; + post_args.Seqs_dq_acc = stride_dq_acc * dq_acc_element_size; post_args.Hs_dq = a.nhead_stride_dq * 2; post_args.BAs_dq = a.batch_stride_dq * 2; post_args.Seqs_dq = a.stride_dq * 2; diff --git a/csrc/include/mha_bwd.h b/csrc/include/mha_bwd.h index 0afaeadd221..05dcc855328 100644 --- a/csrc/include/mha_bwd.h +++ b/csrc/include/mha_bwd.h @@ -8,6 +8,7 @@ #if ENABLE_CK #include "fmha_bwd.hpp" #endif +#include #include namespace aiter { @@ -46,7 +47,6 @@ struct mha_bwd_args void* dk_ptr; void* dv_ptr; void* dbias_ptr; - void* dq_acc_ptr; const void* sink_ptr = nullptr; // sink scores [batch, nhead] log-space (LSEDataType=float); nullptr disables sink void* d_sink_ptr = nullptr; // sink gradient accumulator [nhead] (LSEDataType=float); nullptr disables sink grad // Usage notes for sequence length pointer parameters: @@ -108,7 +108,6 @@ struct mha_bwd_args int stride_o; int stride_randval; int stride_do; - int stride_dq_acc; int stride_dq; int stride_dk; int stride_dv; @@ -121,7 +120,6 @@ struct mha_bwd_args int nhead_stride_randval; int nhead_stride_do; int nhead_stride_lsed; - int64_t nhead_stride_dq_acc; int nhead_stride_dq; int nhead_stride_dk; int nhead_stride_dv; @@ -134,18 +132,21 @@ struct mha_bwd_args int batch_stride_randval; int batch_stride_do; int batch_stride_lsed; - int64_t batch_stride_dq_acc; int batch_stride_dq; int batch_stride_dk; int batch_stride_dv; int batch_stride_dbias; - int split_stride_dq_acc; int window_size_left; int window_size_right; float p_drop; float p_undrop; std::variant, std::pair> drop_seed_offset; + + // Per-call device-buffer allocator. Caller keeps the returned pointer alive + // until aiter::mha_bwd returns. If zero_init is true the bytes must be zero + // by the time the kernel reads them. + std::function workspace_alloc{}; }; struct __attribute__((packed)) fmha_bwd_dqdkdv_args diff --git a/csrc/py_itfs_ck/mha_bwd_kernels.cu b/csrc/py_itfs_ck/mha_bwd_kernels.cu index 2615ed667f7..69ce6b85ca1 100644 --- a/csrc/py_itfs_ck/mha_bwd_kernels.cu +++ b/csrc/py_itfs_ck/mha_bwd_kernels.cu @@ -138,14 +138,14 @@ mha_bwd(const at::Tensor &dout, // [b, sq, hq, d_v] auto stream = at::hip::getCurrentHIPStream(); auto softmax_d = torch::empty({batch_size, num_heads, seqlen_q}, opts.dtype(at::kFloat)); - // nsplits: deterministic mode splits dK into ceil(seqlen_k/16) pieces for atomic-free accumulation. - constexpr ck_tile::index_t kN0 = 16; - const ck_tile::index_t nsplits = deterministic - ? ck_tile::integer_divide_ceil(seqlen_k, kN0) - : 1; - // Always zero dq_accum: the dq_dk_dv kernel writes via atomicAdd regardless of - // deterministic mode, so an uninitialized accumulator would corrupt dQ. - at::Tensor dq_accum = torch::zeros({batch_size, num_heads, nsplits, seqlen_q, head_size_q}, opts.dtype(at::kFloat)); + + at::Tensor workspace; + auto workspace_alloc = [&workspace, opts](size_t bytes, bool zero_init) -> void* { + workspace = zero_init + ? torch::zeros({static_cast(bytes)}, opts.dtype(at::kByte)) + : torch::empty({static_cast(bytes)}, opts.dtype(at::kByte)); + return workspace.data_ptr(); + }; at::Tensor dk_expanded, dv_expanded; if (num_heads_k != num_heads) { // MQA / GQA @@ -240,12 +240,6 @@ mha_bwd(const at::Tensor &dout, // [b, sq, hq, d_v] ck_tile::index_t stride_dv = dv_expanded.stride(1); ck_tile::index_t nhead_stride_dv = dv_expanded.stride(2); - // dq_acc: (batch_size, nheads, split, seqlen_q, hdim_q) - ck_tile::long_index_t batch_stride_dq_acc = dq_accum.stride(0); - ck_tile::long_index_t nhead_stride_dq_acc = dq_accum.stride(1); - ck_tile::index_t split_stride_dq_acc = dq_accum.stride(2); - ck_tile::index_t stride_dq_acc = dq_accum.stride(3); - float p_undrop = 1.0 - p_dropout; void *bias_ptr = nullptr; @@ -338,7 +332,6 @@ mha_bwd(const at::Tensor &dout, // [b, sq, hq, d_v] dk_expanded.data_ptr(), dv_expanded.data_ptr(), dbias_ptr, - dq_accum.data_ptr(), sink_data_ptr, // sink_ptr [b, hq] d_sink_data_ptr, // d_sink_ptr [hq] nullptr, // seqstart_q_ptr @@ -362,7 +355,6 @@ mha_bwd(const at::Tensor &dout, // [b, sq, hq, d_v] stride_o, 0, // stride_randval stride_do, - stride_dq_acc, stride_dq, stride_dk, stride_dv, @@ -375,7 +367,6 @@ mha_bwd(const at::Tensor &dout, // [b, sq, hq, d_v] 0, // nhead_stride_randval nhead_stride_do, nhead_stride_lse, - nhead_stride_dq_acc, nhead_stride_dq, nhead_stride_dk, nhead_stride_dv, @@ -388,17 +379,16 @@ mha_bwd(const at::Tensor &dout, // [b, sq, hq, d_v] 0, // batch_stride_randval batch_stride_do, batch_stride_lse, - batch_stride_dq_acc, batch_stride_dq, batch_stride_dk, batch_stride_dv, batch_stride_dbias, - split_stride_dq_acc, mask.left, mask.right, p_dropout, p_undrop, - drop_seed_offset}; + drop_seed_offset, + workspace_alloc}; }(); float t = aiter::mha_bwd(args, stream_config); diff --git a/csrc/py_itfs_ck/mha_varlen_bwd_kernels.cu b/csrc/py_itfs_ck/mha_varlen_bwd_kernels.cu index 90f7b57b438..419e131b6cf 100644 --- a/csrc/py_itfs_ck/mha_varlen_bwd_kernels.cu +++ b/csrc/py_itfs_ck/mha_varlen_bwd_kernels.cu @@ -150,19 +150,19 @@ mha_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v] bias_enum bias_type = alibi_slopes_.has_value() ? bias_enum::alibi : bias_enum::no_bias; auto opts = q.options(); - // nsplits: deterministic mode splits dK into ceil(max_seqlen_k/16) pieces for atomic-free accumulation. - constexpr ck_tile::index_t kN0 = 16; - const ck_tile::index_t nsplits = deterministic - ? ck_tile::integer_divide_ceil(max_seqlen_k, kN0) - : 1; const at::hip::OptionalHIPGuardMasqueradingAsCUDA device_guard{q.device()}; auto stream = at::hip::getCurrentHIPStream(); auto softmax_d = torch::empty({batch_size, num_heads, total_q}, opts.dtype(at::kFloat)); - // Always zero dq_accum: the dq_dk_dv kernel writes via atomicAdd regardless of - // deterministic mode, so an uninitialized accumulator would corrupt dQ. - at::Tensor dq_accum = torch::zeros({num_heads, nsplits, total_q, head_size_q}, opts.dtype(at::kFloat)); + + at::Tensor workspace; + auto workspace_alloc = [&workspace, opts](size_t bytes, bool zero_init) -> void* { + workspace = zero_init + ? torch::zeros({static_cast(bytes)}, opts.dtype(at::kByte)) + : torch::empty({static_cast(bytes)}, opts.dtype(at::kByte)); + return workspace.data_ptr(); + }; at::Tensor dk_expanded, dv_expanded; if (num_heads_k != num_heads) { // MQA / GQA @@ -254,12 +254,6 @@ mha_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v] ck_tile::index_t stride_dv = dv_expanded.stride(0); ck_tile::index_t nhead_stride_dv = dv_expanded.stride(1); - // dq_acc: (nheads, split, total_q, hdim_v) - ck_tile::long_index_t batch_stride_dq_acc = 0; - ck_tile::long_index_t nhead_stride_dq_acc = dq_accum.stride(0); - ck_tile::index_t split_stride_dq_acc = dq_accum.stride(1); - ck_tile::index_t stride_dq_acc = dq_accum.stride(2); - float p_undrop = 1.0 - p_dropout; void *alibi_slopes_ptr = nullptr; @@ -346,7 +340,6 @@ mha_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v] dk_expanded.data_ptr(), dv_expanded.data_ptr(), nullptr, // dbias - dq_accum.data_ptr(), // dq_acc sink_data_ptr, // sink_ptr [b, hq] d_sink_data_ptr, // d_sink_ptr [hq] seqstart_q_ptr, // seqstart_q_ptr (physical cumulative) @@ -370,7 +363,6 @@ mha_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v] stride_o, 0, // stride_randval stride_do, - stride_dq_acc, stride_dq, stride_dk, stride_dv, @@ -383,7 +375,6 @@ mha_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v] 0, // nhead_stride_randval nhead_stride_do, nhead_stride_lse, - nhead_stride_dq_acc, nhead_stride_dq, nhead_stride_dk, nhead_stride_dv, @@ -396,17 +387,16 @@ mha_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v] 0, // batch_stride_randval batch_stride_do, batch_stride_lse, - batch_stride_dq_acc, batch_stride_dq, batch_stride_dk, batch_stride_dv, 0 , // batch_stride_dbias, FA without dbias - split_stride_dq_acc, mask.left, mask.right, p_dropout, p_undrop, - drop_seed_offset}; + drop_seed_offset, + workspace_alloc}; }(); float t = aiter::mha_bwd(args, stream_config); diff --git a/csrc/py_itfs_cu/asm_mha_bwd.cu b/csrc/py_itfs_cu/asm_mha_bwd.cu index 415640c694a..a4522104585 100644 --- a/csrc/py_itfs_cu/asm_mha_bwd.cu +++ b/csrc/py_itfs_cu/asm_mha_bwd.cu @@ -130,18 +130,14 @@ std::vector fmha_v3_bwd(const at::Tensor &dout, // [b, sq, h auto opts = q.options(); auto softmax_d = torch::empty({batch_size, num_heads, seqlen_q}, opts.dtype(at::kFloat)); - at::Tensor dq_accum; - - if (!deterministic) { - if (is_v3_atomic_fp32) { - dq_accum = torch::zeros({1, batch_size, num_heads, seqlen_q, head_size_q}, opts.dtype(at::kFloat)); - } else { - // When atomic16, padding dq_accum seqlen to 16x, head dim to 128/192 - // In this case, dq_accum could have any layout, we set it to be `bhsd` - int padded_head_size_q = head_size_q == 192? 192: 128; - dq_accum = torch::zeros({1, batch_size, num_heads, (seqlen_q + 15) / 16 * 16, padded_head_size_q}, opts.dtype(q_dtype)); - } - } + + at::Tensor workspace; + auto workspace_alloc = [&workspace, opts](size_t bytes, bool zero_init) -> void* { + workspace = zero_init + ? torch::zeros({static_cast(bytes)}, opts.dtype(at::kByte)) + : torch::empty({static_cast(bytes)}, opts.dtype(at::kByte)); + return workspace.data_ptr(); + }; at::Tensor dk_expanded, dv_expanded; if (num_heads_k != num_heads) { // MQA / GQA @@ -225,13 +221,6 @@ std::vector fmha_v3_bwd(const at::Tensor &dout, // [b, sq, h ck_tile::index_t stride_dv = dv_expanded.stride(1); ck_tile::index_t nhead_stride_dv = dv_expanded.stride(2); - // TODO: if dq_acc layout do no harm to performance consider reuse this api - // dq_acc: (split, batch_size, nheads, seqlen_q, hdim_q) - ck_tile::index_t split_stride_dq_acc = dq_accum.stride(0); - ck_tile::long_index_t batch_stride_dq_acc = dq_accum.stride(1); - ck_tile::long_index_t nhead_stride_dq_acc = dq_accum.stride(2); - ck_tile::index_t stride_dq_acc = dq_accum.stride(3); - float p_undrop = 1.0 - p_dropout; void *alibi_slopes_ptr = nullptr; @@ -276,7 +265,6 @@ std::vector fmha_v3_bwd(const at::Tensor &dout, // [b, sq, h dk_expanded.data_ptr(), dv_expanded.data_ptr(), nullptr, // dbias - dq_accum.data_ptr(), nullptr, // sink_ptr (not used in v3 asm path) nullptr, // d_sink_ptr (not used in v3 asm path) nullptr, // seqstart_q_ptr (batch mode) @@ -300,7 +288,6 @@ std::vector fmha_v3_bwd(const at::Tensor &dout, // [b, sq, h stride_o, 0, // stride_randval stride_do, - stride_dq_acc, stride_dq, stride_dk, stride_dv, @@ -313,7 +300,6 @@ std::vector fmha_v3_bwd(const at::Tensor &dout, // [b, sq, h 0, // nhead_stride_randval nhead_stride_do, nhead_stride_lse, - nhead_stride_dq_acc, nhead_stride_dq, nhead_stride_dk, nhead_stride_dv, @@ -326,17 +312,16 @@ std::vector fmha_v3_bwd(const at::Tensor &dout, // [b, sq, h 0, // batch_stride_randval batch_stride_do, batch_stride_lse, - batch_stride_dq_acc, batch_stride_dq, batch_stride_dk, batch_stride_dv, 0 , // batch_stride_dbias, FA without dbias - split_stride_dq_acc, mask.left, mask.right, p_dropout, p_undrop, - drop_seed_offset}; + drop_seed_offset, + workspace_alloc}; }(); float t = aiter::mha_bwd(args, stream_config); diff --git a/csrc/py_itfs_cu/asm_mha_varlen_bwd.cu b/csrc/py_itfs_cu/asm_mha_varlen_bwd.cu index c705c9962c7..234c57341a1 100644 --- a/csrc/py_itfs_cu/asm_mha_varlen_bwd.cu +++ b/csrc/py_itfs_cu/asm_mha_varlen_bwd.cu @@ -152,17 +152,14 @@ fmha_v3_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v auto opts = q.options(); auto softmax_d = torch::empty({batch_size, num_heads, total_q}, opts.dtype(at::kFloat)); - at::Tensor dq_accum; - - if (!deterministic) { - if (is_v3_atomic_fp32) { - dq_accum = torch::zeros({1, num_heads, total_q, head_size_q}, opts.dtype(at::kFloat)); - } else { - // When atomic16, padding dq_accum seqlen to 16x of max_seqlen_q, head dim to 128 - // In this case, dq_accum could have any layout, we set it to be `bhsd` - dq_accum = torch::zeros({1, batch_size, num_heads, (max_seqlen_q + 15) / 16 * 16, 128}, opts.dtype(q_dtype)); - } - } + + at::Tensor workspace; + auto workspace_alloc = [&workspace, opts](size_t bytes, bool zero_init) -> void* { + workspace = zero_init + ? torch::zeros({static_cast(bytes)}, opts.dtype(at::kByte)) + : torch::empty({static_cast(bytes)}, opts.dtype(at::kByte)); + return workspace.data_ptr(); + }; at::Tensor dk_expanded, dv_expanded; if (num_heads_k != num_heads) { // MQA / GQA @@ -256,24 +253,6 @@ fmha_v3_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v ck_tile::index_t stride_dv = dv_expanded.stride(0); ck_tile::index_t nhead_stride_dv = dv_expanded.stride(1); - ck_tile::index_t split_stride_dq_acc; - ck_tile::long_index_t batch_stride_dq_acc; - ck_tile::long_index_t nhead_stride_dq_acc; - ck_tile::index_t stride_dq_acc; - // For atomic32, dq_acc layout is (1, num_heads, total_q, head_size_q) - // For atomic16, dq_acc layout is (1, batch_size, num_heads, (max_seqlen_q + 15) / 16 * 16, 128) - if (is_v3_atomic_fp32) { - split_stride_dq_acc = dq_accum.stride(0); - batch_stride_dq_acc = 0; - nhead_stride_dq_acc = dq_accum.stride(1); - stride_dq_acc = dq_accum.stride(2); - } else { - split_stride_dq_acc = dq_accum.stride(0); - batch_stride_dq_acc = dq_accum.stride(1); - nhead_stride_dq_acc = dq_accum.stride(2); - stride_dq_acc = dq_accum.stride(3); - } - float p_undrop = 1.0 - p_dropout; void *alibi_slopes_ptr = nullptr; @@ -337,7 +316,6 @@ fmha_v3_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v dk_expanded.data_ptr(), dv_expanded.data_ptr(), nullptr, // dbias - dq_accum.data_ptr(), // dq_acc nullptr, // sink_ptr (not used in v3 asm path) nullptr, // d_sink_ptr (not used in v3 asm path) seqstart_q_ptr, // seqstart_q @@ -361,7 +339,6 @@ fmha_v3_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v stride_o, 0, // stride_randval stride_do, - stride_dq_acc, stride_dq, stride_dk, stride_dv, @@ -374,7 +351,6 @@ fmha_v3_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v 0, // nhead_stride_randval nhead_stride_do, nhead_stride_lse, - nhead_stride_dq_acc, nhead_stride_dq, nhead_stride_dk, nhead_stride_dv, @@ -387,17 +363,16 @@ fmha_v3_varlen_bwd(const at::Tensor &dout, // [total_q, hq, d_v 0, // batch_stride_randval batch_stride_do, batch_stride_lse, - batch_stride_dq_acc, batch_stride_dq, batch_stride_dk, batch_stride_dv, 0 , // batch_stride_dbias, FA without dbias - split_stride_dq_acc, mask.left, mask.right, p_dropout, p_undrop, - drop_seed_offset}; + drop_seed_offset, + workspace_alloc}; }(); float t = aiter::mha_bwd(args, stream_config); diff --git a/op_tests/cpp/mha/benchmark_mha_bwd.cpp b/op_tests/cpp/mha/benchmark_mha_bwd.cpp index 0fbbf45c714..c3deab1e693 100644 --- a/op_tests/cpp/mha/benchmark_mha_bwd.cpp +++ b/op_tests/cpp/mha/benchmark_mha_bwd.cpp @@ -4,6 +4,7 @@ #include "mha_bwd.h" #include "utils.hpp" +#include #include #include #include @@ -351,6 +352,32 @@ bool run(const ck_tile::ArgParser& arg_parser) (mode == mode_enum::batch ? seqlen_q : seqstart_q_host.back()); const ck_tile::index_t shape_seqlen_k = (mode == mode_enum::batch ? seqlen_k : seqstart_k_host.back()); + // For group mode, the new launcher needs seqstart on host during construction. + // seqstart_q_host / seqstart_k_host are already host std::vector from earlier. + const fmha_bwd_traits traits{ + shape_seqlen_q, + shape_seqlen_k, + batch, + max_seqlen_q, + max_seqlen_k, + hdim_q, + hdim_v, + nhead, + nhead_k, + data_type, + mode == mode_enum::group, + mask.type, + bias.type, + use_dbias, + p_drop > 0.0f, + s_randval, + deterministic, + (mode == mode_enum::group) ? seqstart_q_host.data() : nullptr, + (mode == mode_enum::group) ? seqstart_k_host.data() : nullptr, + }; + fmha_bwd_launcher launcher(traits); + // nsplits is still needed for the ASM atomic16 dq_accum tensor shape below. + // Recompute it independently since the new launcher no longer exposes dq_acc_splits. const ck_tile::index_t kN0 = (hdim_q <= 128) ? 128 : 64; const ck_tile::index_t nsplits = deterministic ? ck_tile::integer_divide_ceil(max_seqlen_k, kN0) : 1; @@ -463,9 +490,24 @@ bool run(const ck_tile::ArgParser& arg_parser) ck_tile::DeviceMem drop_seed_buf(drop_prefs ? sizeof(uint64_t) : 0); ck_tile::DeviceMem drop_offset_buf(drop_prefs ? sizeof(uint64_t) : 0); ck_tile::DeviceMem alibi_slope_buf(alibi_slope_host.get_element_space_size_in_bytes()); - ck_tile::DeviceMem dq_acc_buf(v3_atomic_fp32 - ? dq_acc_host.get_element_space_size_in_bytes() - : dq_acc_host_a16.get_element_space_size_in_bytes()); + const std::size_t asm_dq_acc_bytes = + v3_atomic_fp32 ? dq_acc_host.get_element_space_size_in_bytes() + : dq_acc_host_a16.get_element_space_size_in_bytes(); + const std::size_t ck_workspace_bytes = launcher.workspace_size; + ck_tile::DeviceMem dq_acc_buf(std::max(asm_dq_acc_bytes, ck_workspace_bytes)); + + // Pre-allocated dq_acc_buf serves the workspace_alloc callback (sized to max of + // ASM dq_accum and CK launcher workspace requirements). + auto workspace_alloc = [&dq_acc_buf](size_t bytes, bool zero_init) -> void* { + AITER_CHECK(bytes <= dq_acc_buf.GetBufferSize(), + "benchmark workspace_alloc: requested ", bytes, + " bytes but dq_acc_buf is ", dq_acc_buf.GetBufferSize()); + if(zero_init) + { + HIP_CHECK_ERROR(hipMemset(dq_acc_buf.GetDeviceBuffer(), 0, bytes)); + } + return dq_acc_buf.GetDeviceBuffer(); + }; q_buf.ToDevice(q_host.data()); k_buf.ToDevice(k_host.data()); @@ -519,7 +561,6 @@ bool run(const ck_tile::ArgParser& arg_parser) const ck_tile::index_t stride_o = (o_perm ? hdim_v : nhead * hdim_v); const ck_tile::index_t stride_randval = (max_seqlen_k); const ck_tile::index_t stride_do = (o_perm ? hdim_v : nhead * hdim_v); - const ck_tile::index_t stride_dq_acc = a16_dq_acc_hdim; const ck_tile::index_t stride_dk = (i_perm ? hdim_q : nhead * hdim_q); const ck_tile::index_t stride_dv = (i_perm ? hdim_v : nhead * hdim_v); const ck_tile::index_t stride_dbias = (i_perm ? max_seqlen_k : nhead * max_seqlen_k); @@ -532,7 +573,6 @@ bool run(const ck_tile::ArgParser& arg_parser) const ck_tile::index_t nhead_stride_randval = (shape_seqlen_q * max_seqlen_k); const ck_tile::index_t nhead_stride_do = (o_perm ? shape_seqlen_q * hdim_v : hdim_v); const ck_tile::index_t nhead_stride_lsed = shape_seqlen_q; - const ck_tile::long_index_t nhead_stride_dq_acc = a16_dq_acc_seq * a16_dq_acc_hdim; const ck_tile::index_t nhead_stride_dbias = (i_perm ? shape_seqlen_q * max_seqlen_k : max_seqlen_k); // setup batch_stride_* arguments @@ -547,10 +587,6 @@ bool run(const ck_tile::ArgParser& arg_parser) const ck_tile::index_t batch_stride_dk = (nhead * shape_seqlen_k * hdim_q); const ck_tile::index_t batch_stride_dv = (nhead * shape_seqlen_k * hdim_v); const ck_tile::index_t batch_stride_dbias = (nhead * shape_seqlen_q * max_seqlen_k); - const ck_tile::long_index_t batch_stride_dq_acc = - (nhead * a16_dq_acc_seq * a16_dq_acc_hdim); - const ck_tile::index_t split_stride_dq_acc = - (shape_batch * nhead * shape_seqlen_q * hdim_q); const auto drop_seed_offset = [&]() -> decltype(fmha_bwd_args::drop_seed_offset) { if(drop_prefs) @@ -594,7 +630,6 @@ bool run(const ck_tile::ArgParser& arg_parser) dk_buf.GetDeviceBuffer(), dv_buf.GetDeviceBuffer(), dbias_buf.GetDeviceBuffer(), - dq_acc_buf.GetDeviceBuffer(), nullptr, // sink_ptr nullptr, // d_sink_ptr seqstart_q.GetDeviceBuffer(), @@ -619,7 +654,6 @@ bool run(const ck_tile::ArgParser& arg_parser) stride_o, stride_randval, stride_do, - stride_dq_acc, stride_q, // stride_dq stride_dk, stride_dv, @@ -632,7 +666,6 @@ bool run(const ck_tile::ArgParser& arg_parser) nhead_stride_randval, nhead_stride_do, nhead_stride_lsed, - nhead_stride_dq_acc, nhead_stride_q, // nhead_stride_dq nhead_stride_k, // nhead_stride_dk nhead_stride_v, // nhead_stride_dv @@ -645,17 +678,16 @@ bool run(const ck_tile::ArgParser& arg_parser) batch_stride_randval, batch_stride_do, batch_stride_lsed, - batch_stride_dq_acc, // batch_stride_dq_acc batch_stride_q, // batch_stride_dq batch_stride_dk, batch_stride_dv, batch_stride_dbias, - split_stride_dq_acc, mask.left, mask.right, p_drop, p_undrop, - drop_seed_offset}; + drop_seed_offset, + workspace_alloc}; }(); float ave_time = aiter::mha_bwd(mha_args, stream_config); @@ -891,7 +923,6 @@ bool run(const ck_tile::ArgParser& arg_parser) lse_buf.ToDevice(lse_host.data()); dq_buf.SetZero(); dbias_buf.SetZero(); - dq_acc_buf.SetZero(); ck_tile::stream_config stream_config_v{ nullptr, true, 0, 0, 1, arg_parser.get_str("timer") == std::string("gpu")};