Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
97 changes: 76 additions & 21 deletions csrc/cpp_itfs/mha_bwd.cu
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include "asm_fmha_v3_bwd_configs.hpp"
#include <memory>
#include <string>
#include <vector>

namespace aiter {
std::tuple<int, int> get_padded_hdim(int hdim_q, int hdim_v, std::string arch_id)
Expand Down Expand Up @@ -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<int> 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,
Expand All @@ -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);
Comment thread
DDEle marked this conversation as resolved.
launcher.prepare_workspace(workspace_ptr);

fmha_bwd_args ck_args{
/* q_ptr */ a.q_ptr,
/* k_ptr */ a.k_ptr,
Expand All @@ -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,

Expand Down Expand Up @@ -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,
Expand All @@ -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<ck_tile::long_index_t>(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,
Expand All @@ -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<ck_tile::long_index_t>(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,
Expand All @@ -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) ||
Expand Down Expand Up @@ -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<size_t>(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);
}
Comment thread
DDEle marked this conversation as resolved.

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;
Expand Down Expand Up @@ -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;
Expand Down
11 changes: 6 additions & 5 deletions csrc/include/mha_bwd.h
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
#if ENABLE_CK
#include "fmha_bwd.hpp"
#endif
#include <functional>
#include <variant>

namespace aiter {
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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;
Expand All @@ -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;
Expand All @@ -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<uint64_t, uint64_t>, std::pair<const void*, const void*>>
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<void*(size_t bytes, bool zero_init)> workspace_alloc{};
Comment thread
DDEle marked this conversation as resolved.
};

struct __attribute__((packed)) fmha_bwd_dqdkdv_args
Expand Down
30 changes: 10 additions & 20 deletions csrc/py_itfs_ck/mha_bwd_kernels.cu
Original file line number Diff line number Diff line change
Expand Up @@ -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<int64_t>(bytes)}, opts.dtype(at::kByte))
: torch::empty({static_cast<int64_t>(bytes)}, opts.dtype(at::kByte));
return workspace.data_ptr();
};

at::Tensor dk_expanded, dv_expanded;
if (num_heads_k != num_heads) { // MQA / GQA
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -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);
Expand Down
Loading
Loading