v4.1 release update v2. (#2481)
This commit is contained in:
@@ -505,8 +505,12 @@ struct FwdRunner {
|
||||
Tensor mLSE = make_tensor(make_gmem_ptr(buffer.block_ref_LSE.get()),
|
||||
select<0,3>(problem_shape),
|
||||
stride_LSE);
|
||||
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
|
||||
fmha_reference(problem_shape, mQ, mK, mV, mO, mLSE, ActiveMask{});
|
||||
auto problem_shape_ref = cute::make_tuple(Q, K, D, D, HB);
|
||||
|
||||
fmha_reference(problem_shape_ref, mQ, mK, mV, mO, mLSE, ActiveMask{});
|
||||
|
||||
cudaError_t result = cudaDeviceSynchronize();
|
||||
if (result != cudaSuccess) {
|
||||
|
||||
@@ -32,7 +32,7 @@
|
||||
\brief Example implementation of fused multi-head attention for Blackwell using CUTLASS 3.
|
||||
|
||||
This example showcases the use of CUTLASS to build backward fused
|
||||
multi-head attantion (FMHA) collectives from existing CUTLASS collectives targeting
|
||||
multi-head attention (FMHA) collectives from existing CUTLASS collectives targeting
|
||||
the NVIDIA Blackwell architecture.
|
||||
|
||||
Background and motivation
|
||||
@@ -117,6 +117,7 @@ struct Options {
|
||||
std::vector<int> varlen_q;
|
||||
std::vector<int> varlen_k;
|
||||
int d = 128;
|
||||
int d_vo = 128;
|
||||
int iterations = 3;
|
||||
bool verify = false;
|
||||
bool verbose = false;
|
||||
@@ -178,6 +179,7 @@ struct Options {
|
||||
}
|
||||
|
||||
cmd.get_cmd_line_argument("d", d, defaults.d);
|
||||
cmd.get_cmd_line_argument("d_vo", d_vo, d);
|
||||
cmd.get_cmd_line_argument("h", h, -1);
|
||||
if (h == -1) h = 2048 / d;
|
||||
|
||||
@@ -301,6 +303,7 @@ struct Options {
|
||||
<< " --varlen-q=<int>:<int...> Sets the variable Q extent per batch (colon separated)\n"
|
||||
<< " --varlen-k=<int>:<int...> Sets the variable K extent per batch (colon separated)\n"
|
||||
<< " --d=<int> Sets the D extent\n"
|
||||
<< " --d_vo=<int> Sets the D_VO extent\n"
|
||||
<< " --iterations=<int> Benchmarking iterations\n"
|
||||
<< " --verify Verify results\n"
|
||||
<< " --verbose Print smem and execution time per kernel\n"
|
||||
@@ -387,6 +390,7 @@ struct ExampleResult {
|
||||
|
||||
template<
|
||||
bool kIsVarlen,
|
||||
bool kIsMla,
|
||||
class TileShape,
|
||||
class DispatchPolicy,
|
||||
class ActiveMask,
|
||||
@@ -404,8 +408,8 @@ struct BwdRunner {
|
||||
// Q K D (H B)
|
||||
using ProblemShape = std::conditional_t<
|
||||
kIsVarlen,
|
||||
cute::tuple<VariableLength, VariableLength, int, cute::tuple<int, int>>,
|
||||
cute::tuple<int, int, int, cute::tuple<int, int>>
|
||||
cute::tuple<VariableLength, VariableLength, int, int, cute::tuple<int, int>>,
|
||||
cute::tuple<int, int, int, int, cute::tuple<int, int>>
|
||||
>;
|
||||
|
||||
using TensorStride = Stride<int, _1, Stride<int, int>>; // Seq D (H B)
|
||||
@@ -461,45 +465,45 @@ struct BwdRunner {
|
||||
// Methods
|
||||
//
|
||||
bool verify(const ProblemShape& problem_shape) {
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
auto [Q, K, D, D_VO, HB] = problem_shape;
|
||||
auto [H, B] = HB;
|
||||
|
||||
Tensor mQ = make_tensor(make_gmem_ptr(block_Q.get()),
|
||||
select<0,2,3>(problem_shape),
|
||||
select<0,2,4>(problem_shape),
|
||||
stride_Q);
|
||||
|
||||
Tensor mK = make_tensor(make_gmem_ptr(block_K.get()),
|
||||
select<1,2,3>(problem_shape),
|
||||
select<1,2,4>(problem_shape),
|
||||
stride_K);
|
||||
|
||||
Tensor mV = make_tensor(make_gmem_ptr(block_V.get()),
|
||||
select<1,2,3>(problem_shape),
|
||||
select<1,3,4>(problem_shape),
|
||||
stride_V);
|
||||
|
||||
Tensor mO = make_tensor(make_gmem_ptr(block_O.get()),
|
||||
select<0,2,3>(problem_shape),
|
||||
select<0,3,4>(problem_shape),
|
||||
stride_O);
|
||||
|
||||
// keep going here! (this might be better in cursor)
|
||||
|
||||
Tensor mLSE = make_tensor(make_gmem_ptr(block_LSE.get()),
|
||||
select<0,3>(problem_shape),
|
||||
select<0,4>(problem_shape),
|
||||
stride_LSE);
|
||||
|
||||
Tensor mDQ = make_tensor(make_gmem_ptr(block_ref_dQ.get()),
|
||||
select<0,2,3>(problem_shape),
|
||||
select<0,2,4>(problem_shape),
|
||||
stride_dQ);
|
||||
|
||||
Tensor mDK = make_tensor(make_gmem_ptr(block_ref_dK.get()),
|
||||
select<1,2,3>(problem_shape),
|
||||
select<1,2,4>(problem_shape),
|
||||
stride_dK);
|
||||
|
||||
Tensor mDV = make_tensor(make_gmem_ptr(block_ref_dV.get()),
|
||||
select<1,2,3>(problem_shape),
|
||||
select<1,3,4>(problem_shape),
|
||||
stride_dV);
|
||||
|
||||
Tensor mDO = make_tensor(make_gmem_ptr(block_dO.get()),
|
||||
select<0,2,3>(problem_shape),
|
||||
select<0,3,4>(problem_shape),
|
||||
stride_dO);
|
||||
|
||||
fmha_bwd_reference(problem_shape, mQ, mK, mV, mO, mLSE, mDO, mDQ, mDK, mDV, ActiveMask{});
|
||||
@@ -595,14 +599,14 @@ struct BwdRunner {
|
||||
ProblemShape problem_shape{
|
||||
{max_seqlen_q, block_cumulative_seqlen_q.get(), total_seqlen_q},
|
||||
{max_seqlen_kv, block_cumulative_seqlen_kv.get(), total_seqlen_kv},
|
||||
options.d, {options.h, options.b}
|
||||
options.d, options.d_vo, {options.h, options.b}
|
||||
};
|
||||
auto tensor_shape = make_shape(total_seqlen_q, total_seqlen_kv, options.d, make_shape(options.h, 1));
|
||||
auto tensor_shape = make_shape(total_seqlen_q, total_seqlen_kv, options.d, options.d_vo, make_shape(options.h, 1));
|
||||
|
||||
return cute::make_tuple(problem_shape, tensor_shape);
|
||||
}
|
||||
else {
|
||||
ProblemShape problem_shape{options.q, options.k, options.d, {options.h, options.b}};
|
||||
ProblemShape problem_shape{options.q, options.k, options.d, options.d_vo, {options.h, options.b}};
|
||||
return cute::make_tuple(problem_shape, problem_shape);
|
||||
}
|
||||
}
|
||||
@@ -610,24 +614,25 @@ struct BwdRunner {
|
||||
/// Initialize operands to be used in the GEMM and reference GEMM
|
||||
ProblemShape initialize(Options const& options) {
|
||||
auto [problem_shape, tensor_shape] = initialize_problem_shape(options);
|
||||
auto [Q, K, D, HB] = tensor_shape;
|
||||
auto [Q, K, D, D_VO, HB] = tensor_shape;
|
||||
auto [H, B] = HB;
|
||||
D = cutlass::round_up(D, 8); // Alignment
|
||||
|
||||
// for varlen, Q == total_Q, K == total_K, B = 1
|
||||
// but in problem_shape, they've got to be max_Q/max_K, and B = B
|
||||
|
||||
auto shape_QO = make_shape(Q, D, make_shape(H, B));
|
||||
auto shape_KV = make_shape(K, D, make_shape(H, B));
|
||||
auto shape_Q = make_shape(Q, D, make_shape(H, B));
|
||||
auto shape_O = make_shape(Q, D_VO, make_shape(H, B));
|
||||
auto shape_K = make_shape(K, D, make_shape(H, B));
|
||||
auto shape_V = make_shape(K, D_VO, make_shape(H, B));
|
||||
auto shape_LSE = make_shape(Q, make_shape(H, B));
|
||||
|
||||
stride_Q = make_stride(D, _1{}, make_stride(D*Q, B == 1 ? 0 : D*Q*H));
|
||||
stride_K = make_stride(D, _1{}, make_stride(D*K, B == 1 ? 0 : D*K*H));
|
||||
stride_V = make_stride(D_VO, _1{}, make_stride(D_VO*K, B == 1 ? 0 : D_VO*K*H));
|
||||
stride_O = make_stride(D_VO, _1{}, make_stride(D_VO*Q, B == 1 ? 0 : D_VO*Q*H));
|
||||
stride_LSE = make_stride(_1{}, make_stride(Q, B == 1 ? 0 : Q*H));
|
||||
|
||||
stride_V = stride_K;
|
||||
stride_O = stride_Q;
|
||||
|
||||
stride_dQ = stride_Q;
|
||||
stride_dK = stride_K;
|
||||
stride_dV = stride_V;
|
||||
@@ -637,20 +642,20 @@ struct BwdRunner {
|
||||
return size(make_shape(1ull, shape));
|
||||
};
|
||||
|
||||
block_Q.reset(lsize(shape_QO));
|
||||
block_K.reset(lsize(shape_KV));
|
||||
block_V.reset(lsize(shape_KV));
|
||||
block_O.reset(lsize(shape_QO));
|
||||
block_Q.reset(lsize(shape_Q));
|
||||
block_K.reset(lsize(shape_K));
|
||||
block_V.reset(lsize(shape_V));
|
||||
block_O.reset(lsize(shape_O));
|
||||
block_LSE.reset(lsize(shape_LSE));
|
||||
|
||||
block_dQ.reset(lsize(shape_QO));
|
||||
block_dK.reset(lsize(shape_KV));
|
||||
block_dV.reset(lsize(shape_KV));
|
||||
block_dO.reset(lsize(shape_QO));
|
||||
block_dQ.reset(lsize(shape_Q));
|
||||
block_dK.reset(lsize(shape_K));
|
||||
block_dV.reset(lsize(shape_V));
|
||||
block_dO.reset(lsize(shape_O));
|
||||
|
||||
block_ref_dQ.reset(lsize(shape_QO));
|
||||
block_ref_dK.reset(lsize(shape_KV));
|
||||
block_ref_dV.reset(lsize(shape_KV));
|
||||
block_ref_dQ.reset(lsize(shape_Q));
|
||||
block_ref_dK.reset(lsize(shape_K));
|
||||
block_ref_dV.reset(lsize(shape_V));
|
||||
|
||||
initialize_block(block_Q, seed + 2023, options.init_style_q);
|
||||
initialize_block(block_K, seed + 2022, options.init_style_k);
|
||||
@@ -665,23 +670,23 @@ struct BwdRunner {
|
||||
initialize_block(block_ref_dV, seed + 2035);
|
||||
|
||||
Tensor mQ = make_tensor(make_gmem_ptr(block_Q.get()),
|
||||
select<0,2,3>(problem_shape),
|
||||
select<0,2,4>(problem_shape),
|
||||
stride_Q);
|
||||
|
||||
Tensor mK = make_tensor(make_gmem_ptr(block_K.get()),
|
||||
select<1,2,3>(problem_shape),
|
||||
select<1,2,4>(problem_shape),
|
||||
stride_K);
|
||||
|
||||
Tensor mV = make_tensor(make_gmem_ptr(block_V.get()),
|
||||
select<1,2,3>(problem_shape),
|
||||
select<1,3,4>(problem_shape),
|
||||
stride_V);
|
||||
|
||||
Tensor mO = make_tensor(make_gmem_ptr(block_O.get()),
|
||||
select<0,2,3>(problem_shape),
|
||||
select<0,3,4>(problem_shape),
|
||||
stride_O);
|
||||
|
||||
Tensor mLSE = make_tensor(make_gmem_ptr(block_LSE.get()),
|
||||
select<0,3>(problem_shape),
|
||||
select<0,4>(problem_shape),
|
||||
stride_LSE);
|
||||
|
||||
if (! options.skip_reference) {
|
||||
@@ -698,7 +703,7 @@ struct BwdRunner {
|
||||
|
||||
ExampleResult example_result;
|
||||
|
||||
using Operation = cutlass::fmha::device::Sm100FmhaBwd<ProblemShape, Element, ElementAccumulator, TileShape, ActiveMask>;
|
||||
using Operation = cutlass::fmha::device::Sm100FmhaBwd<ProblemShape, Element, ElementAccumulator, TileShape, kIsMla, ActiveMask>;
|
||||
|
||||
typename Operation::Arguments arguments{
|
||||
problem_shape,
|
||||
@@ -811,12 +816,12 @@ struct BwdRunner {
|
||||
|
||||
runtime_ms /= static_cast<float>(options.iterations);
|
||||
|
||||
double flops = 10.0 * (std::is_same_v<ActiveMask, CausalForBackwardMask> ? 0.5 : 1.0);
|
||||
double flops = 2.0 * (std::is_same_v<ActiveMask, CausalForBackwardMask> ? 0.5 : 1.0);
|
||||
flops *= static_cast<double>(get<0>(problem_shape));
|
||||
flops *= static_cast<double>(get<1>(problem_shape));
|
||||
flops *= static_cast<double>(get<2>(problem_shape));
|
||||
flops *= static_cast<double>(get<3,0>(problem_shape));
|
||||
flops *= static_cast<double>(get<3,1>(problem_shape));
|
||||
flops *= (3 * static_cast<double>(get<2>(problem_shape)) + 2 * static_cast<double>(get<3>(problem_shape)));
|
||||
flops *= static_cast<double>(get<4,0>(problem_shape));
|
||||
flops *= static_cast<double>(get<4,1>(problem_shape));
|
||||
double tflops_s = flops * 1e-12 /*tera*/ / (runtime_ms * 1e-3 /*ms*/);
|
||||
example_result.tflops_tc_s = tflops_s;
|
||||
example_result.runtime_ms = runtime_ms;
|
||||
@@ -892,7 +897,7 @@ template<class Mask>
|
||||
void run_bwd_64(Mask fusion, Options const & options, cutlass::KernelHardwareInfo const& hw_info) {
|
||||
auto run = [&](auto shape, auto kernel, const char* name, auto... kernel_options) {
|
||||
dispatch_bool(options.varlen, [&](auto is_varlen) {
|
||||
BwdRunner<decltype(is_varlen)::value, decltype(shape), decltype(kernel), Mask, decltype(kernel_options)...> runner;
|
||||
BwdRunner<decltype(is_varlen)::value, false,decltype(shape), decltype(kernel), Mask, decltype(kernel_options)...> runner;
|
||||
auto result = runner.run(options, hw_info);
|
||||
print_result(name, result, options.verbose);
|
||||
});
|
||||
@@ -900,7 +905,7 @@ void run_bwd_64(Mask fusion, Options const & options, cutlass::KernelHardwareInf
|
||||
|
||||
using HeadDim = _64;
|
||||
|
||||
run(Shape<_128, _128, HeadDim>{}, KernelCoop{}, "tma");
|
||||
run(Shape<_128, _128, HeadDim, HeadDim>{}, KernelCoop{}, "tma");
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -909,7 +914,7 @@ template<class Mask>
|
||||
void run_bwd_128(Mask fusion, Options const & options, cutlass::KernelHardwareInfo const& hw_info) {
|
||||
auto run = [&](auto shape, auto kernel, const char* name, auto... kernel_options) {
|
||||
dispatch_bool(options.varlen, [&](auto is_varlen) {
|
||||
BwdRunner<decltype(is_varlen)::value, decltype(shape), decltype(kernel), Mask, decltype(kernel_options)...> runner;
|
||||
BwdRunner<decltype(is_varlen)::value, false, decltype(shape), decltype(kernel), Mask, decltype(kernel_options)...> runner;
|
||||
auto result = runner.run(options, hw_info);
|
||||
print_result(name, result, options.verbose);
|
||||
});
|
||||
@@ -917,7 +922,22 @@ void run_bwd_128(Mask fusion, Options const & options, cutlass::KernelHardwareIn
|
||||
|
||||
using HeadDim = _128;
|
||||
|
||||
run(Shape<_128, _128, HeadDim>{}, KernelCoop{}, "tma");
|
||||
run(Shape<_128, _128, HeadDim, HeadDim>{}, KernelCoop{}, "tma");
|
||||
}
|
||||
|
||||
template<class Mask>
|
||||
void run_bwd_mla_192(Mask fusion, Options const & options, cutlass::KernelHardwareInfo const& hw_info) {
|
||||
auto run = [&](auto shape, auto kernel, const char* name, auto... kernel_options) {
|
||||
dispatch_bool(options.varlen, [&](auto is_varlen) {
|
||||
BwdRunner<decltype(is_varlen)::value, true, decltype(shape), decltype(kernel), Mask, decltype(kernel_options)...> runner;
|
||||
auto result = runner.run(options, hw_info);
|
||||
print_result(name, result, options.verbose);
|
||||
});
|
||||
};
|
||||
|
||||
using HeadDim = _192;
|
||||
|
||||
run(Shape<_64, _128, HeadDim, _128>{}, KernelCoop{}, "tma");
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -981,7 +1001,7 @@ int main_single(int argc, char const **args) {
|
||||
hw_info.sm_count = options.sm_count;
|
||||
}
|
||||
|
||||
std::cout << "###### B " << options.b << " H " << options.h << " Q " << options.q << " K " << options.k << " D " << options.d << " ";
|
||||
std::cout << "###### B " << options.b << " H " << options.h << " Q " << options.q << " K " << options.k << " D " << options.d << " D_VO " << options.d_vo << " ";
|
||||
std::cout << "Backward" << " " << (options.causal ? "Causal" : "Full") << " ";
|
||||
std::cout << "#SM " << hw_info.sm_count << std::endl;
|
||||
|
||||
@@ -998,12 +1018,15 @@ int main_single(int argc, char const **args) {
|
||||
};
|
||||
|
||||
with_causal([&](auto fusion) {
|
||||
if (options.d <= 64) {
|
||||
if (options.d <= 64 && options.d_vo == options.d) {
|
||||
run_bwd_64(fusion, options, hw_info);
|
||||
}
|
||||
else if (options.d <= 128) {
|
||||
else if (options.d <= 128 && options.d_vo == options.d) {
|
||||
run_bwd_128(fusion, options, hw_info);
|
||||
}
|
||||
else if (options.d == 192 && options.d_vo == 128) {
|
||||
run_bwd_mla_192(fusion, options, hw_info);
|
||||
}
|
||||
else {
|
||||
std::cout << "No kernel instantiated for d=" << options.d << std::endl;
|
||||
}
|
||||
|
||||
@@ -485,7 +485,11 @@ struct MlaFwdRunner {
|
||||
select<0,3>(problem_shape),
|
||||
stride_LSE);
|
||||
|
||||
fmha_reference(problem_shape, mQ, mK, mV, mO, mLSE, ActiveMask{});
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
|
||||
auto problem_shape_ref = cute::make_tuple(Q, K, D, D, HB);
|
||||
|
||||
fmha_reference(problem_shape_ref, mQ, mK, mV, mO, mLSE, ActiveMask{});
|
||||
|
||||
cudaError_t result = cudaDeviceSynchronize();
|
||||
if (result != cudaSuccess) {
|
||||
|
||||
@@ -84,6 +84,8 @@ set(TEST_GEN_REMAP --b=2 --h=4 --h_k=2 --k=512 --d=128 --verify --remap)
|
||||
set(TEST_GEN_CACHEONLY --b=2 --h=4 --h_k=2 --k=512 --d=128 --verify --cache-only)
|
||||
|
||||
set(TEST_MLA_BASIC --b=1 --k=512 --page=128 --verify)
|
||||
set(TEST_BWD_MLA_BASIC --b=1 --h=4 --q=512 --k=512 --d=192 --d_vo=128 --verify --mask=no)
|
||||
set(TEST_BWD_MLA_VARLEN --b=1 --h=4 --q=512 --k=512 --d=192 --d_vo=128 --verify --mask=residual --varlen)
|
||||
|
||||
if(NOT WIN32 AND (NOT (CMAKE_CXX_COMPILER_ID MATCHES "Clang")) AND (CUTLASS_NVCC_ARCHS MATCHES 100a))
|
||||
|
||||
@@ -174,6 +176,8 @@ if(NOT WIN32 AND (NOT (CMAKE_CXX_COMPILER_ID MATCHES "Clang")) AND (CUTLASS_NVCC
|
||||
TEST_VARLEN_12
|
||||
TEST_VARLEN_13
|
||||
TEST_VARLEN_14
|
||||
TEST_BWD_MLA_BASIC
|
||||
TEST_BWD_MLA_VARLEN
|
||||
)
|
||||
target_include_directories(77_blackwell_fmha_bwd_${PREC} PRIVATE ${CMAKE_CURRENT_SOURCE_DIR})
|
||||
target_compile_definitions(77_blackwell_fmha_bwd_${PREC} PRIVATE ${PREC_MACRO})
|
||||
|
||||
@@ -37,13 +37,19 @@ There are three kernels to compute backwards:
|
||||
|
||||
`Sm100FmhaBwdKernelTmaWarpSpecialized` is the main point of this sample, as it demonstrates how to use tensor cores to achieve a high performance fused kernel.
|
||||
|
||||
## MLA Blackwell Backward
|
||||
|
||||
The sample also provides the feature of MLA backward(d=192, d_vo=128). To enable MLA backward, please specify `--d=192 --d_vo=128` when running the bwd sample.
|
||||
|
||||
`Sm100FmhaBwdMlaKernelTmaWarpSpecialized`is the main point for MLA backward. The MLA approach is slightly different from the original one to enable high performance with the MLA shape.
|
||||
|
||||
# MLA Inference for Blackwell
|
||||
|
||||
This sample provides code for fused multi-head latent attention inference in
|
||||
the weight-absorbed regime, i.e. for latent head dim 512, and rope head dim 64.
|
||||
It supports fp16, bf16, and fp8 input and output types.
|
||||
|
||||
To accomodate the large output accumulator due to the large latent head dimension,
|
||||
To accommodate the large output accumulator due to the large latent head dimension,
|
||||
the sample demonstrates how to leverage 2Sm Blackwell tensor cores.
|
||||
|
||||
Loading can be done via TMA (either without paging or with page size 128), or using `cp.async`
|
||||
|
||||
@@ -39,6 +39,7 @@
|
||||
|
||||
#include "../device/fmha.hpp"
|
||||
#include "../kernel/sm100_fmha_bwd_kernel_tma_warpspecialized.hpp"
|
||||
#include "../kernel/sm100_fmha_bwd_mla_kernel_tma_warpspecialized.hpp"
|
||||
#include "../kernel/fmha_kernel_bwd_sum_OdO.hpp"
|
||||
#include "../kernel/fmha_kernel_bwd_convert.hpp"
|
||||
|
||||
@@ -55,13 +56,14 @@ template<
|
||||
class Element,
|
||||
class ElementAccumulator,
|
||||
class TileShape,
|
||||
bool IsMla,
|
||||
class Mask
|
||||
>
|
||||
class Sm100FmhaBwd {
|
||||
public:
|
||||
/// Argument structure: User API
|
||||
struct Arguments {
|
||||
// Q K D HB
|
||||
// Q K D D_VO HB
|
||||
ProblemShape problem_shape;
|
||||
|
||||
const Element* ptr_Q;
|
||||
@@ -98,11 +100,20 @@ public:
|
||||
cutlass::fmha::kernel::FmhaKernelBwdConvert<ProblemShape, Element, ElementAccumulator>
|
||||
>;
|
||||
|
||||
using Operation = cutlass::fmha::device::FMHA<
|
||||
using OperationNormal= cutlass::fmha::device::FMHA<
|
||||
cutlass::fmha::kernel::Sm100FmhaBwdKernelTmaWarpSpecialized<
|
||||
ProblemShape, Element, ElementAccumulator, TileShape, Mask
|
||||
>
|
||||
>;
|
||||
|
||||
using OperationMla = cutlass::fmha::device::FMHA<
|
||||
cutlass::fmha::kernel::Sm100FmhaBwdMlaKernelTmaWarpSpecialized<
|
||||
ProblemShape, Element, ElementAccumulator, TileShape, Mask
|
||||
>
|
||||
>;
|
||||
|
||||
using Operation = std::conditional_t<IsMla, OperationMla, OperationNormal>;
|
||||
|
||||
using Kernel = typename Operation::Kernel;
|
||||
|
||||
struct Params {
|
||||
@@ -121,7 +132,7 @@ private:
|
||||
ElementAccumulator* sum_odo = nullptr,
|
||||
ElementAccumulator* scaled_lse = nullptr) {
|
||||
using namespace cute;
|
||||
auto [Q_, K, D, HB] = args.problem_shape;
|
||||
auto [Q_, K, D, D_VO, HB] = args.problem_shape;
|
||||
auto [H, B] = HB;
|
||||
D = cutlass::round_up(D, 8); // Alignment
|
||||
int Q = cutlass::round_up(static_cast<int>(Q_), 8); // Alignment
|
||||
@@ -141,7 +152,7 @@ private:
|
||||
|
||||
static typename OperationConvert::Arguments to_convert_arguments(Arguments const& args, ElementAccumulator* src = nullptr) {
|
||||
using namespace cute;
|
||||
auto [Q_, K, D, HB] = args.problem_shape;
|
||||
auto [Q_, K, D, D_VO, HB] = args.problem_shape;
|
||||
auto [H, B] = HB;
|
||||
D = cutlass::round_up(D, 8); // Alignment
|
||||
int Q = cutlass::round_up(static_cast<int>(Q_), 8); // Alignment
|
||||
@@ -163,6 +174,7 @@ private:
|
||||
ElementAccumulator* sum_OdO = nullptr, cute::tuple<cute::_1, cute::tuple<int, int>> const& stride_sum_OdO = {},
|
||||
ElementAccumulator* scaled_lse = nullptr, cute::tuple<cute::_1, cute::tuple<int, int>> const& stride_scaled_lse = {},
|
||||
ElementAccumulator* dQ_acc = nullptr, cute::tuple<int, cute::_1, cute::tuple<int, int>> const& stride_dQ = {}) {
|
||||
|
||||
return typename Operation::Arguments{
|
||||
args.problem_shape,
|
||||
{ args.ptr_Q, args.stride_Q,
|
||||
@@ -207,7 +219,7 @@ public:
|
||||
/// Gets the workspace size
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args) {
|
||||
auto [Q_, K, D, HB] = args.problem_shape;
|
||||
auto [Q_, K, D, D_VO, HB] = args.problem_shape;
|
||||
auto [H, B] = HB;
|
||||
D = cutlass::round_up(D, 8); // Alignment
|
||||
int Q = cutlass::round_up(static_cast<int>(Q_), 8); // Alignment
|
||||
@@ -227,7 +239,7 @@ public:
|
||||
CUTLASS_TRACE_HOST("Universal::initialize_split() - workspace_dQ="
|
||||
<< workspace_dQ << ", workspace_sum_OdO=" << workspace_sum_OdO << "stream: " << (stream ? "non-null" : "null"));
|
||||
|
||||
auto [Q_, K, D, HB] = args.problem_shape;
|
||||
auto [Q_, K, D, D_VO, HB] = args.problem_shape;
|
||||
auto [H, B] = HB;
|
||||
D = cutlass::round_up(D, 8); // Alignment
|
||||
int Q = cutlass::round_up(static_cast<int>(Q_), 8); // Alignment
|
||||
@@ -256,7 +268,7 @@ public:
|
||||
CUTLASS_TRACE_HOST("Universal::initialize() - workspace "
|
||||
<< workspace << ", stream: " << (stream ? "non-null" : "null"));
|
||||
|
||||
auto [Q_, K, D, HB] = args.problem_shape;
|
||||
auto [Q_, K, D, D_VO, HB] = args.problem_shape;
|
||||
auto [H, B] = HB;
|
||||
D = cutlass::round_up(D, 8); // Alignment
|
||||
int Q = cutlass::round_up(static_cast<int>(Q_), 8); // Alignment
|
||||
|
||||
@@ -85,11 +85,11 @@ struct FmhaKernelBwdConvert {
|
||||
static const int kIterationsSeq = kBlockSeq / kNumThreadsSeq;
|
||||
|
||||
static bool can_implement(Arguments const& args) {
|
||||
return get<2>(args.problem_shape) % kElementsPerLoad == 0;
|
||||
return get<2>(args.problem_shape) % kElementsPerLoad == 0 && get<3>(args.problem_shape) % kElementsPerLoad == 0;
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
dim3 grid(size<3,0>(params.problem_shape), size<3,1>(params.problem_shape), ceil_div(std::max(size<0>(params.problem_shape), size<1>(params.problem_shape)), kBlockSeq));
|
||||
dim3 grid(size<4,0>(params.problem_shape), size<4,1>(params.problem_shape), ceil_div(std::max(size<0>(params.problem_shape), size<1>(params.problem_shape)), kBlockSeq));
|
||||
return grid;
|
||||
}
|
||||
|
||||
@@ -103,7 +103,7 @@ struct FmhaKernelBwdConvert {
|
||||
}
|
||||
|
||||
template<class StrideSrc, class StrideDest, class Count>
|
||||
CUTLASS_DEVICE void copy(Params const& params, const ElementAcc* ptr_src, StrideSrc const& stride_src, Element* ptr_dest, StrideDest const& stride_dest, Count const& count) {
|
||||
CUTLASS_DEVICE void copy(Params const& params, const ElementAcc* ptr_src, StrideSrc const& stride_src, Element* ptr_dest, StrideDest const& stride_dest, Count const& count, int d_dim) {
|
||||
auto ptr_src_bh = ptr_src + get<2,0>(stride_src) * blockIdx.x + get<2,1>(stride_src) * blockIdx.y;
|
||||
auto ptr_dest_bh = ptr_dest + get<2,0>(stride_dest) * blockIdx.x + get<2,1>(stride_dest) * blockIdx.y;
|
||||
|
||||
@@ -120,7 +120,7 @@ struct FmhaKernelBwdConvert {
|
||||
auto ptr_src_bhs = ptr_src_bh + idx_s * get<0>(stride_src);
|
||||
auto ptr_dest_bhs = ptr_dest_bh + idx_s * get<0>(stride_dest);
|
||||
|
||||
for (int idx_d = threadIdx.x * kElementsPerLoad; idx_d < get<2>(params.problem_shape); idx_d += kElementsPerLoad * kNumThreadsD) {
|
||||
for (int idx_d = threadIdx.x * kElementsPerLoad; idx_d < d_dim; idx_d += kElementsPerLoad * kNumThreadsD) {
|
||||
ElementAcc value_src[kElementsPerLoad];
|
||||
Element value_dest[kElementsPerLoad];
|
||||
|
||||
@@ -139,13 +139,13 @@ struct FmhaKernelBwdConvert {
|
||||
|
||||
CUTLASS_DEVICE void operator()(const Params ¶ms, char* smem) {
|
||||
if (params.ptr_src_dQ != nullptr) {
|
||||
copy(params, params.ptr_src_dQ, params.stride_src_dQ, params.ptr_dest_dQ, params.stride_dest_dQ, get<0>(params.problem_shape));
|
||||
copy(params, params.ptr_src_dQ, params.stride_src_dQ, params.ptr_dest_dQ, params.stride_dest_dQ, get<0>(params.problem_shape), get<2>(params.problem_shape));
|
||||
}
|
||||
if (params.ptr_src_dK != nullptr) {
|
||||
copy(params, params.ptr_src_dK, params.stride_src_dK, params.ptr_dest_dK, params.stride_dest_dK, get<1>(params.problem_shape));
|
||||
copy(params, params.ptr_src_dK, params.stride_src_dK, params.ptr_dest_dK, params.stride_dest_dK, get<1>(params.problem_shape), get<2>(params.problem_shape));
|
||||
}
|
||||
if (params.ptr_src_dV != nullptr) {
|
||||
copy(params, params.ptr_src_dV, params.stride_src_dV, params.ptr_dest_dV, params.stride_dest_dV, get<1>(params.problem_shape));
|
||||
copy(params, params.ptr_src_dV, params.stride_src_dV, params.ptr_dest_dV, params.stride_dest_dV, get<1>(params.problem_shape), get<3>(params.problem_shape));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
@@ -86,11 +86,11 @@ struct FmhaKernelBwdSumOdO {
|
||||
static const int kIterationsQ = kBlockQ / kNumThreadsQ;
|
||||
|
||||
static bool can_implement(Arguments const& args) {
|
||||
return get<2>(args.problem_shape) % kElementsPerLoad == 0;
|
||||
return get<2>(args.problem_shape) % kElementsPerLoad == 0 && get<3>(args.problem_shape) % kElementsPerLoad == 0;
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
dim3 grid(ceil_div(size<0>(params.problem_shape), kBlockQ), size<3,0>(params.problem_shape), size<3,1>(params.problem_shape));
|
||||
dim3 grid(ceil_div(size<0>(params.problem_shape), kBlockQ), size<4,0>(params.problem_shape), size<4,1>(params.problem_shape));
|
||||
return grid;
|
||||
}
|
||||
|
||||
@@ -131,7 +131,7 @@ struct FmhaKernelBwdSumOdO {
|
||||
auto ptr_lse_bhq = ptr_lse_bh + idx_q * get<0>(params.stride_lse);
|
||||
auto ptr_scaled_lse_bhq = ptr_scaled_lse_bh + idx_q * get<0>(params.stride_scaled_lse);
|
||||
|
||||
for (int idx_d = threadIdx.x * kElementsPerLoad; idx_d < get<2>(params.problem_shape); idx_d += kElementsPerLoad * kNumThreadsD) {
|
||||
for (int idx_d = threadIdx.x * kElementsPerLoad; idx_d < get<3>(params.problem_shape); idx_d += kElementsPerLoad * kNumThreadsD) {
|
||||
Element value_O[kElementsPerLoad];
|
||||
Element value_dO[kElementsPerLoad];
|
||||
|
||||
|
||||
@@ -344,12 +344,12 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
|
||||
|
||||
static bool can_implement(Arguments const& args) {
|
||||
auto [Q, K, D, HB] = args.problem_shape;
|
||||
auto [Q, K, D, D_VO, HB] = args.problem_shape;
|
||||
auto [H, B] = HB;
|
||||
if (Q <= 0 || K <= 0 || D <= 0 || H <= 0 || B <= 0) {
|
||||
if (Q <= 0 || K <= 0 || D <= 0 || D_VO <= 0 || H <= 0 || B <= 0) {
|
||||
return false;
|
||||
}
|
||||
if (D % Alignment != 0) {
|
||||
if (D % Alignment != 0 || D_VO % Alignment != 0) {
|
||||
return false;
|
||||
}
|
||||
return true;
|
||||
@@ -362,7 +362,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args, void*) {
|
||||
auto [Q_, K_, D, HB] = args.problem_shape;
|
||||
auto [Q_, K_, D, D_VO, HB] = args.problem_shape;
|
||||
int Q = Q_;
|
||||
int K = K_;
|
||||
|
||||
@@ -381,7 +381,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
}, /*workspace=*/nullptr);
|
||||
|
||||
auto params_vdo = CollectiveMmaVDO::to_underlying_arguments(
|
||||
make_shape(K, Q, D, HB),
|
||||
make_shape(K, Q, D_VO, HB),
|
||||
typename CollectiveMmaVDO::Arguments {
|
||||
args.mainloop.ptr_v, args.mainloop.stride_v,
|
||||
args.mainloop.ptr_do, args.mainloop.stride_do,
|
||||
@@ -446,21 +446,21 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
PipelineLoadComputeSumOdO& pipeline_load_compute_sum_odo,
|
||||
typename PipelineLoadComputeSumOdO::PipelineState& pipeline_load_compute_sum_odo_producer_state) {
|
||||
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
auto [Q, K, D, D_VO, HB] = problem_shape;
|
||||
|
||||
using X = Underscore;
|
||||
|
||||
uint16_t mcast_mask = 0;
|
||||
|
||||
auto mK_in = mainloop_params.tma_load_k.get_tma_tensor(make_shape(K, D, HB));
|
||||
auto mV_in = mainloop_params.tma_load_v.get_tma_tensor(make_shape(K, D, HB));
|
||||
auto mV_in = mainloop_params.tma_load_v.get_tma_tensor(make_shape(K, D_VO, HB));
|
||||
auto mQ_in = mainloop_params.tma_load_q.get_tma_tensor(make_shape(Q, D, HB));
|
||||
auto mDO_in = mainloop_params.tma_load_do.get_tma_tensor(make_shape(Q, D, HB));
|
||||
auto mDO_in = mainloop_params.tma_load_do.get_tma_tensor(make_shape(Q, D_VO, HB));
|
||||
|
||||
auto mK = domain_offset(select<1,2,3>(blk_offset), mK_in);
|
||||
auto mV = domain_offset(select<1,2,3>(blk_offset), mV_in);
|
||||
auto mQ = domain_offset(select<0,2,3>(blk_offset), mQ_in);
|
||||
auto mDO = domain_offset(select<0,2,3>(blk_offset), mDO_in);
|
||||
auto mK = domain_offset(select<1,2,4>(blk_offset), mK_in);
|
||||
auto mV = domain_offset(select<1,3,4>(blk_offset), mV_in);
|
||||
auto mQ = domain_offset(select<0,2,4>(blk_offset), mQ_in);
|
||||
auto mDO = domain_offset(select<0,3,4>(blk_offset), mDO_in);
|
||||
|
||||
auto gK = local_tile(mK, TileShapeKQ{}, make_coord(_,_,_), Step<_1, X, _1>{});
|
||||
auto gQ = local_tile(mQ, TileShapeKQ{}, make_coord(_,_,_), Step<X, _1, _1>{});
|
||||
@@ -495,7 +495,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
|
||||
// set up lse and sum_odo
|
||||
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_batch] = blk_coord;
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_dv, blk_coord_batch] = blk_coord;
|
||||
|
||||
pipeline_load_mma_q.producer_acquire(pipeline_load_mma_q_producer_state);
|
||||
auto tma_barrier = pipeline_load_mma_q.producer_get_barrier(pipeline_load_mma_q_producer_state);
|
||||
@@ -681,7 +681,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
PipelineMmaComputeDKDV& pipeline_mma_compute_dkdv,
|
||||
typename PipelineMmaComputeDKDV::PipelineState& pipeline_mma_compute_dkdv_producer_state) {
|
||||
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
auto [Q, K, D, D_VO, HB] = problem_shape;
|
||||
|
||||
auto sQ = make_tensor(make_smem_ptr(shared_tensors.smem_q.begin()), SmemLayoutQ{});
|
||||
auto sK = make_tensor(make_smem_ptr(shared_tensors.smem_k.begin()), SmemLayoutK{});
|
||||
@@ -974,11 +974,11 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
MainloopArguments const& mainloop_args,
|
||||
EpilogueArguments const& epilogue_args) {
|
||||
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_batch] = blk_coord;
|
||||
auto [Q, K, D, D_VO, HB] = problem_shape;
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_dv, blk_coord_batch] = blk_coord;
|
||||
|
||||
auto mDK_in = make_tensor(make_gmem_ptr(epilogue_args.ptr_dk), make_shape(K, TileShapeDQK{}, HB), epilogue_args.stride_dk);
|
||||
auto mDK = domain_offset(select<1,2,3>(blk_offset), mDK_in);
|
||||
auto mDK = domain_offset(select<1,2,4>(blk_offset), mDK_in);
|
||||
auto gDK = local_tile(mDK, TileShapeDSQ{}, make_coord(_,_,_), Step<_1, _1, X>{})
|
||||
(_, _, blk_coord_k, _0{}, blk_coord_batch);
|
||||
|
||||
@@ -988,7 +988,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
);
|
||||
|
||||
auto mDV_in = make_tensor(make_gmem_ptr(epilogue_args.ptr_dv), make_shape(K, TileShapeDVO{}, HB), epilogue_args.stride_dv);
|
||||
auto mDV = domain_offset(select<1,2,3>(blk_offset), mDV_in);
|
||||
auto mDV = domain_offset(select<1,3,4>(blk_offset), mDV_in);
|
||||
auto gDV = local_tile(mDV, TileShapePDO{}, make_coord(_,_,_), Step<_1, _1, X>{})
|
||||
(_, _, blk_coord_k, _0{}, blk_coord_batch);
|
||||
|
||||
@@ -1003,7 +1003,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
}
|
||||
}
|
||||
for (int i = threadIdx.x; i < size(gDV); i += blockDim.x) {
|
||||
if (elem_less(cDV(i), select<1,2>(problem_shape))) {
|
||||
if (elem_less(cDV(i), select<1,3>(problem_shape))) {
|
||||
gDV(i) = Element(0);
|
||||
}
|
||||
}
|
||||
@@ -1020,8 +1020,8 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
PipelineMmaComputeDKDV& pipeline_mma_compute_dkdv,
|
||||
typename PipelineMmaComputeDKDV::PipelineState& pipeline_mma_compute_dkdv_consumer_state) {
|
||||
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_batch] = blk_coord;
|
||||
auto [Q, K, D, D_VO, HB] = problem_shape;
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_dv, blk_coord_batch] = blk_coord;
|
||||
|
||||
auto load_op = SM100_TMEM_LOAD_32dp32b16x{};
|
||||
|
||||
@@ -1029,7 +1029,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
tDKtDK.data() = TmemAllocation::kDK;
|
||||
|
||||
auto mDK_in = make_tensor(make_gmem_ptr(epilogue_args.ptr_dk), make_shape(K, TileShapeDQK{}, HB), epilogue_args.stride_dk);
|
||||
auto mDK = domain_offset(select<1,2,3>(blk_offset), mDK_in);
|
||||
auto mDK = domain_offset(select<1,2,4>(blk_offset), mDK_in);
|
||||
auto gDK = local_tile(mDK, TileShapeDSQ{}, make_coord(_,_,_), Step<_1, _1, X>{})
|
||||
(_, _, blk_coord_k, _0{}, blk_coord_batch);
|
||||
|
||||
@@ -1065,7 +1065,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
tDVtDV.data() = TmemAllocation::kDV;
|
||||
|
||||
auto mDV_in = make_tensor(make_gmem_ptr(epilogue_args.ptr_dv), make_shape(K, TileShapeDVO{}, HB), epilogue_args.stride_dv);
|
||||
auto mDV = domain_offset(select<1,2,3>(blk_offset), mDV_in);
|
||||
auto mDV = domain_offset(select<1,3,4>(blk_offset), mDV_in);
|
||||
auto gDV = local_tile(mDV, TileShapePDO{}, make_coord(_,_,_), Step<_1, _1, X>{})
|
||||
(_, _, blk_coord_k, _0{}, blk_coord_batch);
|
||||
|
||||
@@ -1088,7 +1088,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
cute::copy(tiled_t2r_dv, tTR_tDV, tTR_rDV);
|
||||
|
||||
// store tDVgDV
|
||||
store(tTR_gDV, tTR_rDV, tTR_cDV, select<1,2>(problem_shape));
|
||||
store(tTR_gDV, tTR_rDV, tTR_cDV, select<1,3>(problem_shape));
|
||||
|
||||
cutlass::arch::fence_view_async_tmem_load();
|
||||
pipeline_mma_compute_dkdv.consumer_release(pipeline_mma_compute_dkdv_consumer_state);
|
||||
@@ -1140,7 +1140,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
typename PipelineMmaComputeDKDV::PipelineState& pipeline_mma_compute_dkdv_consumer_state) {
|
||||
|
||||
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
auto [Q, K, D, D_VO, HB] = problem_shape;
|
||||
|
||||
// in tmem, S & P overlap
|
||||
// and dP and dQ overlap
|
||||
@@ -1396,9 +1396,9 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
|
||||
using X = Underscore;
|
||||
|
||||
auto [Q, K, D, HB] = problem_shape;
|
||||
auto [Q, K, D, D_VO, HB] = problem_shape;
|
||||
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_batch] = blk_coord;
|
||||
auto [blk_coord_q, blk_coord_k, blk_coord_d, blk_coord_dv, blk_coord_batch] = blk_coord;
|
||||
|
||||
// must match TileShapeDQ
|
||||
auto load_op = SM100_TMEM_LOAD_32dp32b32x{};
|
||||
@@ -1676,7 +1676,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
|
||||
pipeline_init_wait(size(ClusterShape{}));
|
||||
|
||||
auto blk_coord = make_coord(_0{}, blockIdx.x, _0{}, make_coord(blockIdx.y, blockIdx.z));
|
||||
auto blk_coord = make_coord(_0{}, blockIdx.x, _0{}, _0{}, make_coord(blockIdx.y, blockIdx.z));
|
||||
auto [problem_shape, blk_offset] = apply_variable_length_offset(
|
||||
params.problem_shape,
|
||||
blk_coord
|
||||
@@ -1809,7 +1809,7 @@ struct Sm100FmhaBwdKernelTmaWarpSpecialized {
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
auto [Q, K, D, HB] = params.problem_shape;
|
||||
auto [Q, K, D, D_VO, HB] = params.problem_shape;
|
||||
auto [H, B] = HB;
|
||||
dim3 grid(ceil_div(K, TileShapeK{}), H, B);
|
||||
return grid;
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -33,7 +33,9 @@
|
||||
#pragma once
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
#include "collective/fmha_fusion.hpp"
|
||||
|
||||
using namespace cutlass::fmha::collective;
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<
|
||||
@@ -61,20 +63,20 @@ void __global__ fmha_bwd_reference_dQ_kernel(
|
||||
|
||||
ElementAccumulator softmax_scale = 1.0 / sqrt(ElementAccumulator(size<2>(problem_shape_in)));
|
||||
|
||||
for (int idx_L = blockIdx.y; idx_L < size<3>(problem_shape_in); idx_L += gridDim.y) {
|
||||
for (int idx_L = blockIdx.y; idx_L < size<4>(problem_shape_in); idx_L += gridDim.y) {
|
||||
auto [problem_shape, offset] = apply_variable_length_offset(
|
||||
problem_shape_in,
|
||||
make_coord(_0{}, _0{}, _0{}, idx2crd(idx_L, get<3>(problem_shape_in)))
|
||||
make_coord(_0{}, _0{}, _0{}, _0{},idx2crd(idx_L, get<4>(problem_shape_in)))
|
||||
);
|
||||
// problem_shape = problem_shape_in;
|
||||
// offset = repeat_like(problem_shape_in, _0{});
|
||||
auto mQ = domain_offset(select<0,2,3>(offset), mQ_in);
|
||||
auto mK = domain_offset(select<1,2,3>(offset), mK_in);
|
||||
auto mV = domain_offset(select<1,2,3>(offset), mV_in);
|
||||
auto mO = domain_offset(select<0,2,3>(offset), mO_in);
|
||||
auto mLSE = domain_offset(select<0,3>(offset), mLSE_in);
|
||||
auto mDO = domain_offset(select<0,2,3>(offset), mDO_in);
|
||||
auto mDQ = domain_offset(select<0,2,3>(offset), mDQ_in);
|
||||
auto mQ = domain_offset(select<0,2,4>(offset), mQ_in);
|
||||
auto mK = domain_offset(select<1,2,4>(offset), mK_in);
|
||||
auto mV = domain_offset(select<1,3,4>(offset), mV_in);
|
||||
auto mO = domain_offset(select<0,3,4>(offset), mO_in);
|
||||
auto mLSE = domain_offset(select<0,4>(offset), mLSE_in);
|
||||
auto mDO = domain_offset(select<0,3,4>(offset), mDO_in);
|
||||
auto mDQ = domain_offset(select<0,2,4>(offset), mDQ_in);
|
||||
for (int idx_Q = blockIdx.x; idx_Q < size<0>(problem_shape); idx_Q += gridDim.x) {
|
||||
for (int idx_K = threadIdx.x; idx_K < size<1>(problem_shape); idx_K += blockDim.x) {
|
||||
ElementAccumulator acc_qk = 0;
|
||||
@@ -82,10 +84,15 @@ void __global__ fmha_bwd_reference_dQ_kernel(
|
||||
ElementAccumulator acc_doo = 0;
|
||||
for (int idx_D0 = 0; idx_D0 < size<2>(problem_shape); idx_D0++) {
|
||||
acc_qk += mQ(idx_Q, idx_D0, idx_L) * mK(idx_K, idx_D0, idx_L);
|
||||
acc_dov += mDO(idx_Q, idx_D0, idx_L) * mV(idx_K, idx_D0, idx_L);
|
||||
acc_doo += mDO(idx_Q, idx_D0, idx_L) * mO(idx_Q, idx_D0, idx_L);
|
||||
// acc_dov += mDO(idx_Q, idx_D0, idx_L) * mV(idx_K, idx_D0, idx_L);
|
||||
// acc_doo += mDO(idx_Q, idx_D0, idx_L) * mO(idx_Q, idx_D0, idx_L);
|
||||
} // for idx_D0
|
||||
|
||||
for (int idx_D1 = 0; idx_D1 < size<3>(problem_shape); idx_D1++) {
|
||||
acc_dov += mDO(idx_Q, idx_D1, idx_L) * mV(idx_K, idx_D1, idx_L);
|
||||
acc_doo += mDO(idx_Q, idx_D1, idx_L) * mO(idx_Q, idx_D1, idx_L);
|
||||
}
|
||||
|
||||
auto id = make_identity_tensor(make_shape(1, 1));
|
||||
auto frag = make_tensor<ElementAccumulator>(Shape<_1, _1>{});
|
||||
frag(0) = acc_qk;
|
||||
@@ -135,20 +142,20 @@ void __global__ fmha_bwd_reference_dK_kernel(
|
||||
|
||||
ElementAccumulator softmax_scale = 1.0 / sqrt(ElementAccumulator(size<2>(problem_shape_in)));
|
||||
|
||||
for (int idx_L = blockIdx.y; idx_L < size<3>(problem_shape_in); idx_L += gridDim.y) {
|
||||
for (int idx_L = blockIdx.y; idx_L < size<4>(problem_shape_in); idx_L += gridDim.y) {
|
||||
auto [problem_shape, offset] = apply_variable_length_offset(
|
||||
problem_shape_in,
|
||||
make_coord(_0{}, _0{}, _0{}, idx2crd(idx_L, get<3>(problem_shape_in)))
|
||||
make_coord(_0{}, _0{}, _0{}, _0{}, idx2crd(idx_L, get<4>(problem_shape_in)))
|
||||
);
|
||||
// problem_shape = problem_shape_in;
|
||||
// offset = repeat_like(problem_shape_in, _0{});
|
||||
auto mQ = domain_offset(select<0,2,3>(offset), mQ_in);
|
||||
auto mK = domain_offset(select<1,2,3>(offset), mK_in);
|
||||
auto mV = domain_offset(select<1,2,3>(offset), mV_in);
|
||||
auto mO = domain_offset(select<0,2,3>(offset), mO_in);
|
||||
auto mLSE = domain_offset(select<0,3>(offset), mLSE_in);
|
||||
auto mDO = domain_offset(select<0,2,3>(offset), mDO_in);
|
||||
auto mDK = domain_offset(select<1,2,3>(offset), mDK_in);
|
||||
auto mQ = domain_offset(select<0,2,4>(offset), mQ_in);
|
||||
auto mK = domain_offset(select<1,2,4>(offset), mK_in);
|
||||
auto mV = domain_offset(select<1,3,4>(offset), mV_in);
|
||||
auto mO = domain_offset(select<0,3,4>(offset), mO_in);
|
||||
auto mLSE = domain_offset(select<0,4>(offset), mLSE_in);
|
||||
auto mDO = domain_offset(select<0,3,4>(offset), mDO_in);
|
||||
auto mDK = domain_offset(select<1,2,4>(offset), mDK_in);
|
||||
for (int idx_K = blockIdx.x; idx_K < size<1>(problem_shape); idx_K += gridDim.x) {
|
||||
for (int idx_Q = threadIdx.x; idx_Q < size<0>(problem_shape); idx_Q += blockDim.x) {
|
||||
ElementAccumulator acc_qk = 0;
|
||||
@@ -156,10 +163,14 @@ void __global__ fmha_bwd_reference_dK_kernel(
|
||||
ElementAccumulator acc_doo = 0;
|
||||
for (int idx_D0 = 0; idx_D0 < size<2>(problem_shape); idx_D0++) {
|
||||
acc_qk += mQ(idx_Q, idx_D0, idx_L) * mK(idx_K, idx_D0, idx_L);
|
||||
acc_dov += mDO(idx_Q, idx_D0, idx_L) * mV(idx_K, idx_D0, idx_L);
|
||||
acc_doo += mDO(idx_Q, idx_D0, idx_L) * mO(idx_Q, idx_D0, idx_L);
|
||||
// acc_dov += mDO(idx_Q, idx_D0, idx_L) * mV(idx_K, idx_D0, idx_L);
|
||||
// acc_doo += mDO(idx_Q, idx_D0, idx_L) * mO(idx_Q, idx_D0, idx_L);
|
||||
} // for idx_D0
|
||||
|
||||
|
||||
for (int idx_D1 = 0; idx_D1 < size<3>(problem_shape); idx_D1++) {
|
||||
acc_dov += mDO(idx_Q, idx_D1, idx_L) * mV(idx_K, idx_D1, idx_L);
|
||||
acc_doo += mDO(idx_Q, idx_D1, idx_L) * mO(idx_Q, idx_D1, idx_L);
|
||||
}
|
||||
auto id = make_identity_tensor(make_shape(1, 1));
|
||||
auto frag = make_tensor<ElementAccumulator>(Shape<_1, _1>{});
|
||||
frag(0) = acc_qk;
|
||||
@@ -209,20 +220,20 @@ void __global__ fmha_bwd_reference_dV_kernel(
|
||||
|
||||
ElementAcc softmax_scale = 1.0 / sqrt(ElementAcc(size<2>(problem_shape_in)));
|
||||
|
||||
for (int idx_L = blockIdx.y; idx_L < size<3>(problem_shape_in); idx_L += gridDim.y) {
|
||||
for (int idx_L = blockIdx.y; idx_L < size<4>(problem_shape_in); idx_L += gridDim.y) {
|
||||
auto [problem_shape, offset] = apply_variable_length_offset(
|
||||
problem_shape_in,
|
||||
make_coord(_0{}, _0{}, _0{}, idx2crd(idx_L, get<3>(problem_shape_in)))
|
||||
make_coord(_0{}, _0{}, _0{}, _0{}, idx2crd(idx_L, get<4>(problem_shape_in)))
|
||||
);
|
||||
// problem_shape = problem_shape_in;
|
||||
// offset = repeat_like(problem_shape_in, _0{});
|
||||
auto mQ = domain_offset(select<0,2,3>(offset), mQ_in);
|
||||
auto mK = domain_offset(select<1,2,3>(offset), mK_in);
|
||||
auto mV = domain_offset(select<1,2,3>(offset), mV_in);
|
||||
auto mO = domain_offset(select<0,2,3>(offset), mO_in);
|
||||
auto mLSE = domain_offset(select<0,3>(offset), mLSE_in);
|
||||
auto mDO = domain_offset(select<0,2,3>(offset), mDO_in);
|
||||
auto mDV = domain_offset(select<1,2,3>(offset), mDV_in);
|
||||
auto mQ = domain_offset(select<0,2,4>(offset), mQ_in);
|
||||
auto mK = domain_offset(select<1,2,4>(offset), mK_in);
|
||||
auto mV = domain_offset(select<1,3,4>(offset), mV_in);
|
||||
auto mO = domain_offset(select<0,3,4>(offset), mO_in);
|
||||
auto mLSE = domain_offset(select<0,4>(offset), mLSE_in);
|
||||
auto mDO = domain_offset(select<0,3,4>(offset), mDO_in);
|
||||
auto mDV = domain_offset(select<1,3,4>(offset), mDV_in);
|
||||
for (int idx_K = blockIdx.x; idx_K < size<1>(problem_shape); idx_K += gridDim.x) {
|
||||
for (int idx_Q = threadIdx.x; idx_Q < size<0>(problem_shape); idx_Q += blockDim.x) {
|
||||
ElementAcc acc_qk = 0;
|
||||
@@ -244,7 +255,7 @@ void __global__ fmha_bwd_reference_dV_kernel(
|
||||
|
||||
__syncthreads();
|
||||
|
||||
for (int idx_D = threadIdx.x; idx_D < size<2>(problem_shape); idx_D += blockDim.x) {
|
||||
for (int idx_D = threadIdx.x; idx_D < size<3>(problem_shape); idx_D += blockDim.x) {
|
||||
ElementAcc acc = 0;
|
||||
for (int idx_Q = 0; idx_Q < size<0>(problem_shape); idx_Q++) {
|
||||
ElementAcc rS = static_cast<Element>(mS[idx_Q]);
|
||||
|
||||
@@ -62,19 +62,20 @@ void __global__ fmha_reference_kernel(
|
||||
ElementAccumulator softmax_scale = static_cast<ElementAccumulator>(1.0 / sqrt(1.0 * size<1>(mQ)));
|
||||
|
||||
auto id = make_identity_tensor(make_shape(1, 1));
|
||||
for (int idx_L = blockIdx.y; idx_L < size<3>(problem_shape_in); idx_L += gridDim.y) {
|
||||
|
||||
for (int idx_L = blockIdx.y; idx_L < size<4>(problem_shape_in); idx_L += gridDim.y) {
|
||||
for (int idx_Q = blockIdx.x; idx_Q < size<0>(problem_shape_in); idx_Q += gridDim.x) {
|
||||
|
||||
auto coord_L = idx2crd(idx_L, shape<3>(problem_shape_in));
|
||||
auto coord_L = idx2crd(idx_L, shape<4>(problem_shape_in));
|
||||
auto get_coord_in = [&]() {
|
||||
if constexpr (rank_v<decltype(get<2>(ProblemShapeIn{}))> == 2) {
|
||||
return cute::make_tuple(idx_Q, _0{}, cute::make_tuple(_0{}, _0{}), coord_L);
|
||||
return cute::make_tuple(idx_Q, _0{}, cute::make_tuple(_0{}, _0{}), cute::make_tuple(_0{}, _0{}), coord_L);
|
||||
} else {
|
||||
return cute::make_tuple(idx_Q, _0{}, _0{}, coord_L);
|
||||
return cute::make_tuple(idx_Q, _0{}, _0{}, _0{}, coord_L);
|
||||
}
|
||||
};
|
||||
auto coord_in = get_coord_in();
|
||||
auto [problem_shape, coord] = apply_variable_length(problem_shape_in, coord_in, get<3,1>(coord_in));
|
||||
auto [problem_shape, coord] = apply_variable_length(problem_shape_in, coord_in, get<4,1>(coord_in));
|
||||
|
||||
int head_qk = 0;
|
||||
int head_v = 0;
|
||||
@@ -83,7 +84,7 @@ void __global__ fmha_reference_kernel(
|
||||
head_qk = size<2, 0>(problem_shape) + size<2, 1>(problem_shape);
|
||||
head_v = size<2, 0>(problem_shape);
|
||||
} else {
|
||||
head_qk = size<2>(problem_shape);
|
||||
head_qk = size<3>(problem_shape);
|
||||
head_v = head_qk;
|
||||
}
|
||||
|
||||
@@ -157,6 +158,7 @@ void __global__ fmha_reference_kernel(
|
||||
mO(idx_Q + offset_Q, idx_D, idx_L) = static_cast<typename TensorO::value_type>(acc * scale);
|
||||
}
|
||||
|
||||
|
||||
if (threadIdx.x == 0 && mLSE.data() != nullptr) {
|
||||
mLSE(idx_Q + offset_Q, idx_L) = log(sum) + softmax_scale * maxS;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user