v4.0 update. (#2371)
This commit is contained in:
@@ -0,0 +1,158 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "../collective/fmha_collective_tma.hpp"
|
||||
#include "../collective/fmha_collective_tma_warpspecialized.hpp"
|
||||
#include "../collective/fmha_epilogue.hpp"
|
||||
#include "../kernel/fmha_kernel_tma.hpp"
|
||||
#include "../kernel/fmha_kernel_tma_warpspecialized.hpp"
|
||||
#include "../kernel/fmha_options.hpp"
|
||||
|
||||
namespace cutlass::fmha::kernel {
|
||||
|
||||
template<
|
||||
class Element_,
|
||||
class ElementAccumulatorQK_,
|
||||
class ElementAccumulatorPV_,
|
||||
class TileShape_, // BlockQO, BlockKV, BlockHead
|
||||
class LayoutQ_,
|
||||
class LayoutK_,
|
||||
class LayoutV_,
|
||||
class Fusion,
|
||||
class DispatchPolicy,
|
||||
class... Options
|
||||
>
|
||||
struct FmhaBuilder;
|
||||
|
||||
template<
|
||||
class Element,
|
||||
class ElementAccumulator,
|
||||
class TileShape, // BlockQO, BlockKV, BlockHead
|
||||
class Fusion,
|
||||
class... Options
|
||||
>
|
||||
struct FmhaBuilder<
|
||||
Element,
|
||||
ElementAccumulator,
|
||||
ElementAccumulator,
|
||||
TileShape,
|
||||
cute::tuple<int, _1, cute::tuple<int, int>>,
|
||||
cute::tuple<int, _1, cute::tuple<int, int>>,
|
||||
cute::tuple<int, _1, cute::tuple<int, int>>,
|
||||
Fusion,
|
||||
cutlass::gemm::KernelTma,
|
||||
Options...
|
||||
> {
|
||||
|
||||
using CollectiveMainloop = cutlass::fmha::collective::FmhaMainloopTma<Element, ElementAccumulator, TileShape, Fusion, Options...>;
|
||||
|
||||
using CollectiveEpilogue = cutlass::fmha::collective::FmhaFwdEpilogue<
|
||||
Element, ElementAccumulator, typename CollectiveMainloop::TileShapePV>;
|
||||
|
||||
using Kernel = cutlass::fmha::kernel::FmhaKernelTma<CollectiveMainloop, CollectiveEpilogue, Options...>;
|
||||
};
|
||||
|
||||
template<
|
||||
class Element,
|
||||
class ElementAccumulatorQK,
|
||||
class ElementAccumulatorPV,
|
||||
class TileShape, // BlockQO, BlockKV, BlockHead
|
||||
class LayoutQ,
|
||||
class LayoutK,
|
||||
class LayoutV,
|
||||
class Fusion,
|
||||
class... Options
|
||||
>
|
||||
struct FmhaBuilder<
|
||||
Element,
|
||||
ElementAccumulatorQK,
|
||||
ElementAccumulatorPV,
|
||||
TileShape,
|
||||
LayoutQ,
|
||||
LayoutK,
|
||||
LayoutV,
|
||||
Fusion,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedCooperative,
|
||||
Options...
|
||||
> {
|
||||
|
||||
using CollectiveMainloop = cutlass::fmha::collective::FmhaMainloopTmaWarpSpecialized<
|
||||
Element, ElementAccumulatorQK, ElementAccumulatorPV,
|
||||
TileShape, LayoutQ, LayoutK, LayoutV,
|
||||
Fusion, Options...>;
|
||||
|
||||
using CollectiveEpilogue = cutlass::fmha::collective::FmhaFwdEpilogue<
|
||||
Element, ElementAccumulatorPV, typename CollectiveMainloop::TileShapePV>;
|
||||
|
||||
static constexpr bool kIsPersistent = find_option_t<Tag::kIsPersistent, false_type, Options...>::value;
|
||||
using TileScheduler = std::conditional_t<kIsPersistent, cutlass::fmha::kernel::PersistentTileScheduler, cutlass::fmha::kernel::IndividualTileScheduler>;
|
||||
|
||||
using Kernel = cutlass::fmha::kernel::FmhaKernelTmaWarpSpecialized<CollectiveMainloop, CollectiveEpilogue, TileScheduler, Options...>;
|
||||
};
|
||||
|
||||
template<
|
||||
class Element,
|
||||
class ElementAccumulatorQK,
|
||||
class ElementAccumulatorPV,
|
||||
class TileShape, // BlockQO, BlockKV, BlockHead
|
||||
class LayoutQ,
|
||||
class LayoutK,
|
||||
class LayoutV,
|
||||
class Fusion,
|
||||
class... Options
|
||||
>
|
||||
struct FmhaBuilder<
|
||||
Element,
|
||||
ElementAccumulatorQK,
|
||||
ElementAccumulatorPV,
|
||||
TileShape,
|
||||
LayoutQ,
|
||||
LayoutK,
|
||||
LayoutV,
|
||||
Fusion,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedPingpong,
|
||||
Options...
|
||||
> {
|
||||
using Kernel = typename FmhaBuilder<
|
||||
Element, ElementAccumulatorQK, ElementAccumulatorPV,
|
||||
TileShape,
|
||||
LayoutQ, LayoutK, LayoutV,
|
||||
Fusion,
|
||||
cutlass::gemm::KernelTmaWarpSpecializedCooperative,
|
||||
Options...,
|
||||
Option<Tag::kIsPersistent, true_type>,
|
||||
Option<Tag::kLoadsQSeparately, true_type>
|
||||
>::Kernel;
|
||||
};
|
||||
|
||||
} // namespace cutlass::fmha::kernel
|
||||
@@ -0,0 +1,143 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/layout.hpp"
|
||||
|
||||
namespace cutlass::fmha::kernel {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template<class Element, class ElementAccumulator>
|
||||
struct FmhaKernelBwdConvert {
|
||||
|
||||
struct Arguments {
|
||||
tuple<int, int, int, int, int> problem_size;
|
||||
|
||||
const ElementAccumulator* ptr_src_dQ;
|
||||
tuple<int, int, int, _1> stride_src_dQ;
|
||||
const ElementAccumulator* ptr_src_dK;
|
||||
tuple<int, int, int, _1> stride_src_dK;
|
||||
const ElementAccumulator* ptr_src_dV;
|
||||
tuple<int, int, int, _1> stride_src_dV;
|
||||
|
||||
Element* ptr_dest_dQ;
|
||||
tuple<int, int, int, _1> stride_dest_dQ;
|
||||
Element* ptr_dest_dK;
|
||||
tuple<int, int, int, _1> stride_dest_dK;
|
||||
Element* ptr_dest_dV;
|
||||
tuple<int, int, int, _1> stride_dest_dV;
|
||||
};
|
||||
|
||||
using Params = Arguments;
|
||||
|
||||
using ClusterShape = Shape<_1, _1, _1>;
|
||||
static constexpr int SharedStorageSize = 0;
|
||||
|
||||
static const int MinBlocksPerMultiprocessor = 1;
|
||||
static const int MaxThreadsPerBlock = 128;
|
||||
using ArchTag = cutlass::arch::Sm90;
|
||||
|
||||
static const int kBlockSeq = 8;
|
||||
|
||||
static size_t get_workspace_size(Arguments const& args) { return 0; }
|
||||
static cutlass::Status initialize_workspace(Arguments const&, void*, cudaStream_t) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
static const int kNumThreadsD = 16;
|
||||
static const int kNumThreadsSeq = MaxThreadsPerBlock / kNumThreadsD;
|
||||
static const int kElementsPerLoad = 4;
|
||||
|
||||
static const int kIterationsSeq = kBlockSeq / kNumThreadsSeq;
|
||||
|
||||
static bool can_implement(Arguments const& args) {
|
||||
return get<4>(args.problem_size) % kElementsPerLoad == 0;
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
dim3 grid(size<0>(params.problem_size), size<1>(params.problem_size), ceil_div(std::max(size<2>(params.problem_size), size<3>(params.problem_size)), kBlockSeq));
|
||||
return grid;
|
||||
}
|
||||
|
||||
static dim3 get_block_shape() {
|
||||
dim3 block(kNumThreadsD, kNumThreadsSeq, 1);
|
||||
return block;
|
||||
}
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
return args;
|
||||
}
|
||||
|
||||
template<class StrideSrc, class StrideDest>
|
||||
CUTLASS_DEVICE void copy(Params const& params, const ElementAccumulator* ptr_src, StrideSrc const& stride_src, Element* ptr_dest, StrideDest const& stride_dest, int count) {
|
||||
auto ptr_src_bh = ptr_src + get<0>(stride_src) * blockIdx.x + get<1>(stride_src) * blockIdx.y;
|
||||
auto ptr_dest_bh = ptr_dest + get<0>(stride_dest) * blockIdx.x + get<1>(stride_dest) * blockIdx.y;
|
||||
|
||||
for (int idx_s_t = threadIdx.y; idx_s_t < kBlockSeq; idx_s_t += kNumThreadsSeq) {
|
||||
int idx_s = idx_s_t + kBlockSeq * blockIdx.z;
|
||||
if (idx_s >= count) continue;
|
||||
auto ptr_src_bhs = ptr_src_bh + idx_s * get<2>(stride_src);
|
||||
auto ptr_dest_bhs = ptr_dest_bh + idx_s * get<2>(stride_dest);
|
||||
|
||||
for (int idx_d = threadIdx.x * kElementsPerLoad; idx_d < get<4>(params.problem_size); idx_d += kElementsPerLoad * kNumThreadsD) {
|
||||
ElementAccumulator value_src[kElementsPerLoad];
|
||||
Element value_dest[kElementsPerLoad];
|
||||
|
||||
using VecSrc = uint_bit_t<sizeof_bits_v<ElementAccumulator> * kElementsPerLoad>;
|
||||
using VecDest = uint_bit_t<sizeof_bits_v<Element> * kElementsPerLoad>;
|
||||
*reinterpret_cast<VecSrc*>(value_src) = *reinterpret_cast<const VecSrc*>(&ptr_src_bhs[idx_d]);
|
||||
|
||||
for (int v = 0; v < kElementsPerLoad; v++) {
|
||||
value_dest[v] = value_src[v];
|
||||
}
|
||||
|
||||
*reinterpret_cast<VecDest*>(&ptr_dest_bhs[idx_d]) = *reinterpret_cast<const VecDest*>(value_dest);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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<2>(params.problem_size));
|
||||
}
|
||||
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<3>(params.problem_size));
|
||||
}
|
||||
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<3>(params.problem_size));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cutlass::fmha::kernel
|
||||
@@ -0,0 +1,134 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cute/layout.hpp"
|
||||
|
||||
namespace cutlass::fmha::kernel {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template<class Element, class ElementAccumulator>
|
||||
struct FmhaKernelBwdSumOdO {
|
||||
|
||||
struct Arguments {
|
||||
cute::tuple<int, int, int, int, int> problem_size;
|
||||
|
||||
const Element* ptr_O;
|
||||
cute::tuple<int, int, int, cute::_1> stride_O;
|
||||
const Element* ptr_dO;
|
||||
cute::tuple<int, int, int, cute::_1> stride_dO;
|
||||
|
||||
ElementAccumulator* ptr_sum_OdO;
|
||||
cute::tuple<int, int, _1> stride_sum_OdO;
|
||||
};
|
||||
|
||||
using Params = Arguments;
|
||||
|
||||
using ClusterShape = Shape<_1, _1, _1>;
|
||||
static constexpr int SharedStorageSize = 0;
|
||||
|
||||
static const int MinBlocksPerMultiprocessor = 1;
|
||||
static const int MaxThreadsPerBlock = 128;
|
||||
using ArchTag = cutlass::arch::Sm90;
|
||||
|
||||
static size_t get_workspace_size(Arguments const& args) { return 0; }
|
||||
static cutlass::Status initialize_workspace(Arguments const&, void*, cudaStream_t) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
static const int kBlockQ = 16;
|
||||
|
||||
static const int kNumThreadsD = 8;
|
||||
static const int kNumThreadsQ = MaxThreadsPerBlock / kNumThreadsD;
|
||||
static const int kElementsPerLoad = 2;
|
||||
|
||||
static const int kIterationsQ = kBlockQ / kNumThreadsQ;
|
||||
|
||||
static bool can_implement(Arguments const& args) {
|
||||
return get<4>(args.problem_size) % kElementsPerLoad == 0;
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
dim3 grid(ceil_div(size<2>(params.problem_size), kBlockQ), size<1>(params.problem_size), size<0>(params.problem_size));
|
||||
return grid;
|
||||
}
|
||||
|
||||
static dim3 get_block_shape() {
|
||||
dim3 block(kNumThreadsD, kNumThreadsQ, 1);
|
||||
return block;
|
||||
}
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
return args;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void operator()(const Params ¶ms, char* smem) {
|
||||
auto ptr_O_bh = params.ptr_O + blockIdx.y * get<1>(params.stride_O) + blockIdx.z * get<0>(params.stride_O);
|
||||
auto ptr_dO_bh = params.ptr_dO + blockIdx.y * get<1>(params.stride_dO) + blockIdx.z * get<0>(params.stride_dO);
|
||||
auto ptr_sum_OdO_bh = params.ptr_sum_OdO + blockIdx.y * get<1>(params.stride_sum_OdO) + blockIdx.z * get<0>(params.stride_sum_OdO);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int idx_q_t = threadIdx.y; idx_q_t < kBlockQ; idx_q_t += kNumThreadsQ) {
|
||||
int idx_q = idx_q_t + kBlockQ * blockIdx.x;
|
||||
if (idx_q >= get<2>(params.problem_size)) continue;
|
||||
ElementAccumulator acc = 0;
|
||||
auto ptr_O_bhq = ptr_O_bh + idx_q * get<2>(params.stride_O);
|
||||
auto ptr_dO_bhq = ptr_dO_bh + idx_q * get<2>(params.stride_dO);
|
||||
auto ptr_sum_OdO_bhq = ptr_sum_OdO_bh + idx_q * get<2>(params.stride_sum_OdO);
|
||||
|
||||
for (int idx_d = threadIdx.x * kElementsPerLoad; idx_d < get<4>(params.problem_size); idx_d += kElementsPerLoad * kNumThreadsD) {
|
||||
Element value_O[kElementsPerLoad];
|
||||
Element value_dO[kElementsPerLoad];
|
||||
|
||||
using Vec = uint_bit_t<sizeof_bits_v<Element> * kElementsPerLoad>;
|
||||
*reinterpret_cast<Vec*>(value_O) = *reinterpret_cast<const Vec*>(&ptr_O_bhq[idx_d]);
|
||||
*reinterpret_cast<Vec*>(value_dO) = *reinterpret_cast<const Vec*>(&ptr_dO_bhq[idx_d]);
|
||||
|
||||
for (int v = 0; v < kElementsPerLoad; v++) {
|
||||
acc += value_O[v] * value_dO[v];
|
||||
}
|
||||
}
|
||||
|
||||
for (int i = 1; i < kNumThreadsD; i *= 2) {
|
||||
acc += __shfl_xor_sync((uint32_t)-1, acc, i, kNumThreadsD);
|
||||
}
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
*ptr_sum_OdO_bhq = acc;
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cutlass::fmha::kernel
|
||||
@@ -0,0 +1,222 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/arch/arch.h"
|
||||
|
||||
#include "../kernel/fmha_tile_scheduler.hpp"
|
||||
#include "../kernel/fmha_options.hpp"
|
||||
|
||||
namespace cutlass::fmha::kernel {
|
||||
|
||||
template<
|
||||
class CollectiveMainloop,
|
||||
class CollectiveEpilogue,
|
||||
class... Options
|
||||
>
|
||||
struct FmhaKernelTma {
|
||||
|
||||
// Options
|
||||
static constexpr int kBlocksPerSM = find_option_t<Tag::kBlocksPerSM, Int<2>, Options...>::value;
|
||||
|
||||
using Element = typename CollectiveMainloop::Element;
|
||||
using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator;
|
||||
|
||||
using TileScheduler = IndividualTileScheduler;
|
||||
|
||||
using StagesQ = typename CollectiveMainloop::StagesQ;
|
||||
using Stages = typename CollectiveMainloop::Stages;
|
||||
|
||||
using TileShape = typename CollectiveMainloop::TileShape;
|
||||
using ClusterShape = typename CollectiveMainloop::ClusterShape;
|
||||
|
||||
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
|
||||
using MainloopPipelineQ = typename CollectiveMainloop::MainloopPipelineQ;
|
||||
|
||||
using SmemLayoutQ = typename CollectiveMainloop::SmemLayoutQ;
|
||||
using SmemLayoutK = typename CollectiveMainloop::SmemLayoutK;
|
||||
|
||||
struct SharedStorage {
|
||||
union {
|
||||
typename CollectiveMainloop::SharedStorage mainloop;
|
||||
typename CollectiveEpilogue::TensorStorage epilogue;
|
||||
};
|
||||
|
||||
using PipelineStorage = typename MainloopPipeline::SharedStorage;
|
||||
using PipelineStorageQ = typename MainloopPipelineQ::SharedStorage;
|
||||
alignas(16) PipelineStorage pipeline_storage;
|
||||
alignas(16) PipelineStorageQ pipeline_storage_q;
|
||||
|
||||
using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage;
|
||||
alignas(16) EpiLoadPipelineStorage epi_load;
|
||||
};
|
||||
|
||||
static constexpr int SharedStorageSize = sizeof(SharedStorage);
|
||||
|
||||
using ProblemShape = cute::tuple<int, int, int, int, int>;
|
||||
|
||||
struct Arguments {
|
||||
ProblemShape problem_size;
|
||||
typename CollectiveMainloop::Arguments mainloop;
|
||||
typename CollectiveEpilogue::Arguments epilogue;
|
||||
KernelHardwareInfo hw_info;
|
||||
};
|
||||
|
||||
struct Params {
|
||||
ProblemShape problem_size;
|
||||
typename CollectiveMainloop::Params mainloop;
|
||||
typename CollectiveEpilogue::Params epilogue;
|
||||
typename TileScheduler::Params tile_scheduler;
|
||||
};
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename cutlass::PipelineState<MainloopPipeline::Stages>;
|
||||
using PipelineParamsQ = typename MainloopPipelineQ::Params;
|
||||
using PipelineStateQ = typename cutlass::PipelineState<MainloopPipelineQ::Stages>;
|
||||
|
||||
static const int MinBlocksPerMultiprocessor = kBlocksPerSM;
|
||||
static const int MaxThreadsPerBlock = CollectiveMainloop::MaxThreadsPerBlock;
|
||||
using ArchTag = cutlass::arch::Sm90;
|
||||
|
||||
static size_t get_workspace_size(Arguments const& args) { return 0; }
|
||||
static cutlass::Status initialize_workspace(Arguments const&, void*, cudaStream_t) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
static bool can_implement(Arguments const& args) {
|
||||
return CollectiveMainloop::can_implement(args.problem_size, args.mainloop);
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
return TileScheduler::get_grid_shape(params.tile_scheduler);
|
||||
}
|
||||
|
||||
static dim3 get_block_shape() {
|
||||
dim3 block(MaxThreadsPerBlock, 1, 1);
|
||||
return block;
|
||||
}
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
return Params{
|
||||
args.problem_size,
|
||||
CollectiveMainloop::to_underlying_arguments(args.problem_size, args.mainloop, workspace),
|
||||
CollectiveEpilogue::to_underlying_arguments(args.problem_size, args.epilogue, workspace),
|
||||
TileScheduler::to_underlying_arguments(args.problem_size, args.hw_info, ClusterShape{}, TileShape{})
|
||||
};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void operator()(const Params ¶ms, char* smem) {
|
||||
TileScheduler tile_scheduler{params.tile_scheduler};
|
||||
|
||||
// Shared memory.
|
||||
auto& storage = *reinterpret_cast<SharedStorage*>(smem);
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
|
||||
uint32_t block_rank_in_cluster = cute::block_rank_in_cluster();
|
||||
|
||||
int warp_idx = cutlass::canonical_warp_idx_sync();
|
||||
int warp_group_thread_idx = thread_idx % cutlass::NumThreadsPerWarpGroup;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
// Issue Tma Descriptor Prefetch from a single thread
|
||||
if ((warp_idx == 0) && lane_predicate) {
|
||||
CollectiveMainloop::prefetch_tma_descriptors(params.mainloop);
|
||||
}
|
||||
|
||||
|
||||
PipelineParamsQ pipeline_params_q;
|
||||
pipeline_params_q.transaction_bytes = size(SmemLayoutQ{}(_,_,_0{})) * sizeof(Element); // Q
|
||||
pipeline_params_q.role = MainloopPipelineQ::ThreadCategory::ProducerConsumer;
|
||||
pipeline_params_q.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params_q.num_consumers = cutlass::NumThreadsPerWarpGroup;
|
||||
|
||||
PipelineParams pipeline_params;
|
||||
pipeline_params.transaction_bytes = size(SmemLayoutK{}(_,_,_0{})) * sizeof(Element); // KV
|
||||
pipeline_params.role = MainloopPipeline::ThreadCategory::ProducerConsumer;
|
||||
pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params.num_consumers = cutlass::NumThreadsPerWarpGroup;
|
||||
|
||||
MainloopPipelineQ pipeline_q(storage.pipeline_storage_q, pipeline_params_q, Shape<_1, _1, _1>{});
|
||||
MainloopPipeline pipeline(storage.pipeline_storage, pipeline_params, ClusterShape{});
|
||||
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
typename EpiLoadPipeline::Params epi_load_pipeline_params;
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::ProducerConsumer;
|
||||
epi_load_pipeline_params.dst_blockid = cute::block_rank_in_cluster();
|
||||
epi_load_pipeline_params.producer_arv_count = NumThreadsPerWarp;
|
||||
epi_load_pipeline_params.consumer_arv_count = NumThreadsPerWarpGroup;
|
||||
epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes;
|
||||
EpiLoadPipeline epi_load_pipeline(storage.epi_load, epi_load_pipeline_params);
|
||||
|
||||
// State variables used for iterating the circular buffer
|
||||
// smem_pipe_read / release is used by the consumer of SMEM data - i.e MMA
|
||||
// smem_pipe_write is used by the producer of SMEM data - i.e TMA
|
||||
PipelineState smem_pipe_read;
|
||||
PipelineState smem_pipe_write = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
|
||||
PipelineStateQ smem_pipe_read_q;
|
||||
PipelineStateQ smem_pipe_write_q = cutlass::make_producer_start_state<MainloopPipelineQ>();
|
||||
|
||||
// We need this to guarantee that the Pipeline init is visible
|
||||
// To all producers and consumer blocks in the Cluster
|
||||
// and to finish smem init
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
cute::cluster_arrive_relaxed();
|
||||
cute::cluster_wait();
|
||||
}
|
||||
else {
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
auto blk_coord = tile_scheduler.get_block_coord();
|
||||
|
||||
CollectiveMainloop collective_mainloop;
|
||||
auto result = collective_mainloop.compute(
|
||||
block_rank_in_cluster,
|
||||
blk_coord, params.mainloop, params.problem_size,
|
||||
pipeline, smem_pipe_read, smem_pipe_write,
|
||||
pipeline_q, smem_pipe_read_q, smem_pipe_write_q,
|
||||
storage.mainloop
|
||||
);
|
||||
|
||||
CollectiveEpilogue epilogue;
|
||||
epilogue(typename CollectiveMainloop::TileShapePV{}, blk_coord,
|
||||
result, typename CollectiveMainloop::TiledMmaPV{},
|
||||
params.problem_size, params.epilogue,
|
||||
epi_load_pipeline, storage.epilogue);
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cutlass::fmha::kernel
|
||||
@@ -0,0 +1,418 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/arch/reg_reconfig.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/arch/arch.h"
|
||||
|
||||
#include "../kernel/fmha_options.hpp"
|
||||
|
||||
namespace cutlass::fmha::kernel {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template<
|
||||
class CollectiveMainloop,
|
||||
class CollectiveEpilogue,
|
||||
class TileScheduler,
|
||||
class... Options
|
||||
>
|
||||
struct FmhaKernelTmaWarpSpecialized {
|
||||
|
||||
// Options
|
||||
static constexpr bool kIsEpilogueLocked = find_option_t<Tag::kIsEpilogueLocked, false_type, Options...>::value;
|
||||
static constexpr bool kLoadsQSeparately = find_option_t<Tag::kLoadsQSeparately, false_type, Options...>::value;
|
||||
|
||||
|
||||
static const int NumLoadWarpGroups = 1;
|
||||
static constexpr int NumMmaWarpGroups = CollectiveMainloop::NumMmaWarpGroups;
|
||||
|
||||
using TileShape = typename CollectiveMainloop::TileShape;
|
||||
using ClusterShape = typename CollectiveMainloop::ClusterShape;
|
||||
|
||||
using MainloopPipelineOuter = typename CollectiveMainloop::MainloopPipelineQ;
|
||||
using MainloopPipelineInner = typename CollectiveMainloop::MainloopPipeline;
|
||||
using MainloopPipelineReducer = cutlass::PipelineAsync<2>;
|
||||
|
||||
static constexpr uint32_t StagesPerMathWarpGroup = 2;
|
||||
using MathWarpGroupOrderBarrier = cutlass::OrderedSequenceBarrier<
|
||||
StagesPerMathWarpGroup, NumMmaWarpGroups>;
|
||||
|
||||
struct TensorStorageStruct {
|
||||
typename CollectiveMainloop::SharedStorage mainloop;
|
||||
typename CollectiveEpilogue::TensorStorage epilogue[NumMmaWarpGroups];
|
||||
};
|
||||
union TensorStorageUnion {
|
||||
typename CollectiveMainloop::SharedStorage mainloop;
|
||||
typename CollectiveEpilogue::TensorStorage epilogue[NumMmaWarpGroups];
|
||||
};
|
||||
using TensorStorage = std::conditional_t<CollectiveMainloop::kIsPersistent, TensorStorageStruct, TensorStorageUnion>;
|
||||
|
||||
struct SharedStorage {
|
||||
TensorStorage tensors;
|
||||
|
||||
using PipelineStorageInner = typename MainloopPipelineInner::SharedStorage;
|
||||
using PipelineStorageOuter = typename MainloopPipelineOuter::SharedStorage;
|
||||
using PipelineStorageReducer = typename MainloopPipelineReducer::SharedStorage;
|
||||
|
||||
alignas(16) PipelineStorageInner pipeline_storage_inner;
|
||||
alignas(16) PipelineStorageOuter pipeline_storage_outer;
|
||||
alignas(16) PipelineStorageReducer pipeline_storage_reducer;
|
||||
|
||||
using MathWarpGroupOrderBarrierStorage = typename MathWarpGroupOrderBarrier::SharedStorage;
|
||||
alignas(16) MathWarpGroupOrderBarrierStorage math_wg_order;
|
||||
|
||||
alignas(16) cutlass::arch::ClusterBarrier load_warp_barrier;
|
||||
|
||||
using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage;
|
||||
alignas(16) EpiLoadPipelineStorage epi_load;
|
||||
};
|
||||
|
||||
static constexpr int SharedStorageSize = sizeof(SharedStorage);
|
||||
|
||||
using ProblemShape = cute::tuple<int, int, int, int, int>;
|
||||
|
||||
struct Arguments {
|
||||
ProblemShape problem_size;
|
||||
typename CollectiveMainloop::Arguments mainloop;
|
||||
typename CollectiveEpilogue::Arguments epilogue;
|
||||
KernelHardwareInfo hw_info;
|
||||
};
|
||||
|
||||
struct Params {
|
||||
ProblemShape problem_size;
|
||||
typename CollectiveMainloop::Params mainloop;
|
||||
typename CollectiveEpilogue::Params epilogue;
|
||||
typename TileScheduler::Params tile_scheduler;
|
||||
};
|
||||
|
||||
using PipelineParamsInner = typename MainloopPipelineInner::Params;
|
||||
using PipelineStateInner = typename cutlass::PipelineState<MainloopPipelineInner::Stages>;
|
||||
using PipelineParamsOuter = typename MainloopPipelineOuter::Params;
|
||||
using PipelineStateOuter = typename cutlass::PipelineState<MainloopPipelineOuter::Stages>;
|
||||
using PipelineParamsReducer = typename MainloopPipelineReducer::Params;
|
||||
using PipelineStateReducer = typename cutlass::PipelineState<MainloopPipelineReducer::Stages>;
|
||||
|
||||
static const int MinBlocksPerMultiprocessor = 1;
|
||||
static const int MaxThreadsPerBlock = (NumMmaWarpGroups + NumLoadWarpGroups) * cutlass::NumThreadsPerWarpGroup;
|
||||
using ArchTag = cutlass::arch::Sm90;
|
||||
|
||||
static constexpr uint32_t LoadRegisterRequirement = 40 - 2 * 8;
|
||||
static constexpr uint32_t TotalRegisterSupply = (64*1024 / MaxThreadsPerBlock / MinBlocksPerMultiprocessor / 8) * 8 * MaxThreadsPerBlock / cutlass::NumThreadsPerWarpGroup;
|
||||
static constexpr uint32_t MmaRegisterRequirement = ((TotalRegisterSupply - LoadRegisterRequirement) / NumMmaWarpGroups / 8) * 8;
|
||||
|
||||
static size_t get_workspace_size(Arguments const& args) { return 0; }
|
||||
static cutlass::Status initialize_workspace(Arguments const&, void*, cudaStream_t) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
static bool can_implement(Arguments const& args) {
|
||||
return CollectiveMainloop::can_implement(args.problem_size, args.mainloop);
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
return TileScheduler::get_grid_shape(params.tile_scheduler);
|
||||
}
|
||||
|
||||
static dim3 get_block_shape() {
|
||||
dim3 block(MaxThreadsPerBlock, 1, 1);
|
||||
return block;
|
||||
}
|
||||
|
||||
static Params to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
return Params{
|
||||
args.problem_size,
|
||||
CollectiveMainloop::to_underlying_arguments(args.problem_size, args.mainloop, workspace),
|
||||
CollectiveEpilogue::to_underlying_arguments(args.problem_size, args.epilogue, workspace),
|
||||
TileScheduler::to_underlying_arguments(args.problem_size, args.hw_info, ClusterShape{}, TileShape{})
|
||||
};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void operator()(const Params ¶ms, char* smem) {
|
||||
|
||||
enum class WarpGroupRole {
|
||||
Producer = 0,
|
||||
Consumer0 = 1,
|
||||
Consumer1 = 2,
|
||||
Consumer2 = 3,
|
||||
Consumer3 = 4,
|
||||
};
|
||||
enum class ProducerWarpRole {
|
||||
LoadKV = 1,
|
||||
Reducer = 0,
|
||||
MaybeLoadQ = 2, // is kLoadsQSeparately is true, this warp loads Q (otherwise warp 0 does it)
|
||||
MainloopEpilogue = 3,
|
||||
};
|
||||
|
||||
static constexpr ProducerWarpRole WarpRoleLoadQ = kLoadsQSeparately ? ProducerWarpRole::MaybeLoadQ : ProducerWarpRole::LoadKV;
|
||||
|
||||
TileScheduler tile_scheduler{params.tile_scheduler};
|
||||
|
||||
// Shared memory.
|
||||
auto& storage = *reinterpret_cast<SharedStorage*>(smem);
|
||||
|
||||
int lane_idx = cutlass::canonical_lane_idx();
|
||||
int warp_idx = cutlass::canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % cutlass::NumWarpsPerWarpGroup;
|
||||
int warp_group_idx = cutlass::canonical_warp_group_idx();
|
||||
auto warp_group_role = WarpGroupRole(warp_group_idx);
|
||||
auto producer_warp_role = ProducerWarpRole(warp_idx_in_warp_group);
|
||||
int consumer_warp_group_idx = warp_group_idx - (int) WarpGroupRole::Consumer0;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
uint32_t block_rank_in_cluster = cute::block_rank_in_cluster();
|
||||
|
||||
// Issue Tma Descriptor Prefetch from a single thread
|
||||
if ((warp_idx == 0) && lane_predicate) {
|
||||
CollectiveMainloop::prefetch_tma_descriptors(params.mainloop);
|
||||
}
|
||||
|
||||
PipelineParamsOuter pipeline_params_outer;
|
||||
pipeline_params_outer.transaction_bytes = CollectiveMainloop::kOuterLoadBytes;
|
||||
pipeline_params_outer.is_leader = lane_predicate && (producer_warp_role == WarpRoleLoadQ);
|
||||
pipeline_params_outer.num_consumers = cutlass::NumThreadsPerWarpGroup;
|
||||
|
||||
PipelineParamsInner pipeline_params_inner;
|
||||
pipeline_params_inner.transaction_bytes = CollectiveMainloop::kInnerLoadBytes;
|
||||
pipeline_params_inner.is_leader = lane_predicate && (producer_warp_role == ProducerWarpRole::LoadKV);
|
||||
pipeline_params_inner.num_consumers = NumMmaWarpGroups * cutlass::NumThreadsPerWarpGroup;
|
||||
|
||||
PipelineParamsReducer pipeline_params_reducer;
|
||||
pipeline_params_reducer.producer_arv_count = NumMmaWarpGroups * cutlass::NumThreadsPerWarpGroup;
|
||||
pipeline_params_reducer.consumer_arv_count = cutlass::NumThreadsPerWarp;
|
||||
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
typename EpiLoadPipeline::Params epi_load_pipeline_params;
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::MainloopEpilogue) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::LoadKV) {
|
||||
pipeline_params_inner.role = MainloopPipelineInner::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == WarpRoleLoadQ) {
|
||||
pipeline_params_outer.role = MainloopPipelineOuter::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::Reducer) {
|
||||
pipeline_params_reducer.role = MainloopPipelineReducer::ThreadCategory::Consumer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer0 ||
|
||||
warp_group_role == WarpGroupRole::Consumer1 ||
|
||||
warp_group_role == WarpGroupRole::Consumer2 ||
|
||||
warp_group_role == WarpGroupRole::Consumer3
|
||||
) {
|
||||
pipeline_params_inner.role = MainloopPipelineInner::ThreadCategory::Consumer;
|
||||
pipeline_params_outer.role = MainloopPipelineOuter::ThreadCategory::Consumer;
|
||||
pipeline_params_reducer.role = MainloopPipelineReducer::ThreadCategory::Producer;
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
|
||||
MainloopPipelineOuter pipeline_outer(storage.pipeline_storage_outer, pipeline_params_outer, Shape<_1, _1, _1>{});
|
||||
MainloopPipelineInner pipeline_inner(storage.pipeline_storage_inner, pipeline_params_inner, ClusterShape{});
|
||||
MainloopPipelineReducer pipeline_reducer(storage.pipeline_storage_reducer, pipeline_params_reducer);
|
||||
|
||||
// State variables used for iterating the circular buffer
|
||||
// smem_pipe_read / release is used by the consumer of SMEM data - i.e MMA
|
||||
// smem_pipe_write is used by the producer of SMEM data - i.e TMA
|
||||
PipelineStateInner smem_pipe_read_inner;
|
||||
PipelineStateInner smem_pipe_write_inner = cutlass::make_producer_start_state<MainloopPipelineInner>();
|
||||
|
||||
PipelineStateOuter smem_pipe_read_outer;
|
||||
PipelineStateOuter smem_pipe_write_outer = cutlass::make_producer_start_state<MainloopPipelineOuter>();
|
||||
|
||||
PipelineStateReducer smem_pipe_read_reducer;
|
||||
PipelineStateReducer smem_pipe_write_reducer = cutlass::make_producer_start_state<MainloopPipelineReducer>();
|
||||
|
||||
typename MathWarpGroupOrderBarrier::Params params_math_wg_order_barrier;
|
||||
// DMA Load WG will not participate in these Ordered Barrier syncs
|
||||
params_math_wg_order_barrier.group_id = consumer_warp_group_idx;
|
||||
params_math_wg_order_barrier.group_size = cutlass::NumThreadsPerWarpGroup; // Number of threads / participants in a group
|
||||
MathWarpGroupOrderBarrier math_wg_order_barrier(storage.math_wg_order, params_math_wg_order_barrier);
|
||||
|
||||
// Epilogue Load pipeline
|
||||
epi_load_pipeline_params.dst_blockid = cute::block_rank_in_cluster();
|
||||
epi_load_pipeline_params.producer_arv_count = NumThreadsPerWarp;
|
||||
epi_load_pipeline_params.consumer_arv_count = NumThreadsPerWarpGroup;
|
||||
epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes;
|
||||
EpiLoadPipeline epi_load_pipeline(storage.epi_load, epi_load_pipeline_params);
|
||||
|
||||
if constexpr (kLoadsQSeparately) {
|
||||
if ((warp_idx == 0) && lane_predicate) {
|
||||
storage.load_warp_barrier.init(2 * cutlass::NumThreadsPerWarp);
|
||||
}
|
||||
cutlass::arch::fence_barrier_init();
|
||||
}
|
||||
|
||||
// We need this to guarantee that the Pipeline init is visible
|
||||
// To all producers and consumer blocks in the Cluster
|
||||
// and to finish smem init
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
cute::cluster_arrive_relaxed();
|
||||
cute::cluster_wait();
|
||||
}
|
||||
else {
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
CollectiveMainloop collective_mainloop;
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
cutlass::arch::warpgroup_reg_dealloc<LoadRegisterRequirement>();
|
||||
if (producer_warp_role == ProducerWarpRole::LoadKV) {
|
||||
bool do_barrier = kLoadsQSeparately;
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for (; tile_scheduler.is_valid(); ++tile_scheduler) {
|
||||
auto blk_coord = tile_scheduler.get_block_coord();
|
||||
collective_mainloop.template load_kv_maybe_q<!kLoadsQSeparately>(
|
||||
block_rank_in_cluster,
|
||||
blk_coord, params.mainloop, params.problem_size,
|
||||
pipeline_inner, smem_pipe_write_inner,
|
||||
pipeline_outer, smem_pipe_write_outer,
|
||||
storage.tensors.mainloop,
|
||||
storage.load_warp_barrier, do_barrier
|
||||
);
|
||||
do_barrier = false;
|
||||
}
|
||||
}
|
||||
else if (kLoadsQSeparately && (producer_warp_role == ProducerWarpRole::MaybeLoadQ)) {
|
||||
bool do_barrier = true;
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for (; tile_scheduler.is_valid(); ++tile_scheduler) {
|
||||
auto blk_coord = tile_scheduler.get_block_coord();
|
||||
collective_mainloop.load_maybe_q(
|
||||
blk_coord, params.mainloop, params.problem_size,
|
||||
pipeline_outer, smem_pipe_write_outer,
|
||||
storage.tensors.mainloop,
|
||||
storage.load_warp_barrier, do_barrier
|
||||
);
|
||||
do_barrier = false;
|
||||
}
|
||||
} else if (producer_warp_role == ProducerWarpRole::Reducer) {
|
||||
for (; tile_scheduler.is_valid(); ++tile_scheduler) {
|
||||
auto blk_coord = tile_scheduler.get_block_coord();
|
||||
collective_mainloop.reduce(
|
||||
blk_coord, params.mainloop, params.problem_size,
|
||||
pipeline_reducer, smem_pipe_read_reducer,
|
||||
storage.tensors.mainloop
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
else if (
|
||||
warp_group_role == WarpGroupRole::Consumer0 ||
|
||||
warp_group_role == WarpGroupRole::Consumer1 ||
|
||||
warp_group_role == WarpGroupRole::Consumer2 ||
|
||||
warp_group_role == WarpGroupRole::Consumer3
|
||||
) {
|
||||
cutlass::arch::warpgroup_reg_alloc<MmaRegisterRequirement>();
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for (; tile_scheduler.is_valid(); ++tile_scheduler) {
|
||||
auto blk_coord = tile_scheduler.get_block_coord();
|
||||
auto wg_coord = blk_coord;
|
||||
|
||||
constexpr int kOuterLoads = CollectiveMainloop::kOuterLoads;
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Consumer0) {
|
||||
smem_pipe_read_outer.advance(0 * kOuterLoads);
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer1) {
|
||||
smem_pipe_read_outer.advance(1 * kOuterLoads);
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer2) {
|
||||
smem_pipe_read_outer.advance(2 * kOuterLoads);
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer3) {
|
||||
smem_pipe_read_outer.advance(3 * kOuterLoads);
|
||||
}
|
||||
|
||||
constexpr int wg_dim = is_constant<0, decltype(get<1>(wg_coord))>::value ? 0 : 1;
|
||||
auto& wg_block = get<wg_dim>(wg_coord);
|
||||
if (warp_group_role == WarpGroupRole::Consumer0) {
|
||||
wg_block = NumMmaWarpGroups * wg_block + 0;
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer1) {
|
||||
wg_block = NumMmaWarpGroups * wg_block + 1;
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer2) {
|
||||
wg_block = NumMmaWarpGroups * wg_block + 2;
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer3) {
|
||||
wg_block = NumMmaWarpGroups * wg_block + 3;
|
||||
}
|
||||
|
||||
auto result = collective_mainloop.compute(
|
||||
blk_coord, wg_coord,
|
||||
params.mainloop, params.problem_size,
|
||||
pipeline_inner, smem_pipe_read_inner,
|
||||
pipeline_outer, smem_pipe_read_outer,
|
||||
pipeline_reducer, smem_pipe_write_reducer,
|
||||
storage.tensors.mainloop,
|
||||
math_wg_order_barrier
|
||||
);
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Consumer0) {
|
||||
smem_pipe_read_outer.advance(kOuterLoads * (NumMmaWarpGroups - 0));
|
||||
}
|
||||
if constexpr (NumMmaWarpGroups >= 2) {
|
||||
if (warp_group_role == WarpGroupRole::Consumer1) {
|
||||
smem_pipe_read_outer.advance(kOuterLoads * (NumMmaWarpGroups - 1));
|
||||
}
|
||||
}
|
||||
if constexpr (NumMmaWarpGroups >= 3) {
|
||||
if (warp_group_role == WarpGroupRole::Consumer2) {
|
||||
smem_pipe_read_outer.advance(kOuterLoads * (NumMmaWarpGroups - 2));
|
||||
}
|
||||
}
|
||||
if constexpr (NumMmaWarpGroups >= 4) {
|
||||
if (warp_group_role == WarpGroupRole::Consumer3) {
|
||||
smem_pipe_read_outer.advance(kOuterLoads * (NumMmaWarpGroups - 3));
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (kIsEpilogueLocked) ; math_wg_order_barrier.wait();
|
||||
|
||||
CollectiveEpilogue epilogue;
|
||||
epilogue(typename CollectiveMainloop::TileShapePV{}, wg_coord,
|
||||
result, typename CollectiveMainloop::TiledMmaPV{},
|
||||
params.problem_size, params.epilogue,
|
||||
epi_load_pipeline, storage.tensors.epilogue[consumer_warp_group_idx]);
|
||||
|
||||
if constexpr (kIsEpilogueLocked) ; math_wg_order_barrier.arrive();
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cutlass::fmha::kernel
|
||||
@@ -0,0 +1,83 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
namespace cutlass::fmha::kernel {
|
||||
|
||||
template<auto kTag, typename Default, typename... Options>
|
||||
struct find_option;
|
||||
|
||||
template<auto kTag, typename Default>
|
||||
struct find_option<kTag, Default> {
|
||||
using option_value = Default;
|
||||
};
|
||||
|
||||
template<auto kTag, typename Default, typename Option, typename... Options>
|
||||
struct find_option<kTag, Default, Option, Options...> :
|
||||
std::conditional_t<
|
||||
Option::tag == kTag,
|
||||
Option,
|
||||
find_option<kTag, Default, Options...>
|
||||
>
|
||||
{};
|
||||
|
||||
template<auto kTag, typename Default, typename... Options>
|
||||
using find_option_t = typename find_option<kTag, Default, Options...>::option_value;
|
||||
|
||||
enum class Tag {
|
||||
kIsPersistent,
|
||||
kNumMmaWarpGroups,
|
||||
kLoadsQSeparately,
|
||||
|
||||
kIsMainloopLocked,
|
||||
kIsEpilogueLocked,
|
||||
|
||||
kStagesQ,
|
||||
kStagesKV,
|
||||
|
||||
kEpilogueKind,
|
||||
|
||||
kBlocksPerSM,
|
||||
kClusterM,
|
||||
|
||||
kAccQK
|
||||
};
|
||||
|
||||
template<auto kTag, class Value>
|
||||
struct Option {
|
||||
static constexpr auto tag = kTag;
|
||||
using option_value = Value;
|
||||
};
|
||||
|
||||
} // namespace cutlass::fmha::kernel
|
||||
@@ -0,0 +1,204 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/kernel_hardware_info.h"
|
||||
|
||||
namespace cutlass::fmha::kernel {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct IndividualTileScheduler {
|
||||
|
||||
struct Params {
|
||||
dim3 grid;
|
||||
};
|
||||
|
||||
bool valid_ = true;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
IndividualTileScheduler(Params const&) {}
|
||||
|
||||
template<class ProblemSize, class ClusterShape, class TileShape>
|
||||
static Params to_underlying_arguments(
|
||||
ProblemSize const& problem_size, KernelHardwareInfo hw_info,
|
||||
ClusterShape const& cluster_shape, TileShape const& tile_shape)
|
||||
{
|
||||
using namespace cute;
|
||||
dim3 grid(round_up(ceil_div(size<2>(problem_size), size<0>(tile_shape)), size<0>(cluster_shape)), size<0>(problem_size), size<1>(problem_size));
|
||||
return Params{ grid };
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
return params.grid;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool is_valid() {
|
||||
return valid_;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
auto get_block_coord() {
|
||||
using namespace cute;
|
||||
return make_coord(blockIdx.x, _0{}, make_coord(blockIdx.y, blockIdx.z));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
IndividualTileScheduler& operator++() {
|
||||
valid_ = false;
|
||||
return *this;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct PersistentTileScheduler {
|
||||
|
||||
struct Params {
|
||||
int num_blocks;
|
||||
FastDivmod divmod_m_block;
|
||||
FastDivmod divmod_b;
|
||||
FastDivmod divmod_h;
|
||||
|
||||
KernelHardwareInfo hw_info;
|
||||
};
|
||||
|
||||
int block_idx = 0;
|
||||
Params params;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileScheduler(Params const& params) : block_idx(blockIdx.x), params(params) {}
|
||||
|
||||
template<class ProblemSize, class ClusterShape, class TileShape>
|
||||
static Params to_underlying_arguments(
|
||||
ProblemSize const& problem_size, KernelHardwareInfo hw_info,
|
||||
ClusterShape const& cluster_shape, TileShape const& tile_shape)
|
||||
{
|
||||
using namespace cute;
|
||||
// Get SM count if needed, otherwise use user supplied SM count
|
||||
int sm_count = hw_info.sm_count;
|
||||
if (sm_count <= 0) {
|
||||
CUTLASS_TRACE_HOST(" WARNING: Arguments do not include a valid SM count.\n"
|
||||
" For optimal performance, populate the arguments KernelHardwareInfo struct with the SM count.");
|
||||
sm_count = KernelHardwareInfo::query_device_multiprocessor_count(hw_info.device_id);
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST("to_underlying_arguments(): Setting persistent grid SM count to " << sm_count);
|
||||
hw_info.sm_count = sm_count;
|
||||
|
||||
int num_m_blocks = cutlass::round_up(ceil_div(size<2>(problem_size), size<0>(tile_shape)), size<0>(cluster_shape));
|
||||
int num_blocks = num_m_blocks * size<0>(problem_size) * size<1>(problem_size);
|
||||
|
||||
return Params {
|
||||
num_blocks,
|
||||
{ num_m_blocks}, { size<0>(problem_size) }, { size<1>(problem_size) },
|
||||
hw_info
|
||||
};
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
dim3 grid(std::min(params.num_blocks, params.hw_info.sm_count), 1, 1);
|
||||
return grid;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool is_valid() {
|
||||
return block_idx < params.num_blocks;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
auto get_block_coord() {
|
||||
using namespace cute;
|
||||
int block_decode = block_idx;
|
||||
int m_block, bidb, bidh;
|
||||
params.divmod_m_block(block_decode, m_block, block_decode);
|
||||
params.divmod_b(block_decode, bidb, block_decode);
|
||||
params.divmod_h(block_decode, bidh, block_decode);
|
||||
return make_coord(m_block, _0{}, make_coord(bidb, bidh));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileScheduler& operator++() {
|
||||
block_idx += gridDim.x;
|
||||
return *this;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename Base>
|
||||
struct TileSchedulerBwdAdapter {
|
||||
|
||||
using Params = typename Base::Params;
|
||||
|
||||
Base base_;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
TileSchedulerBwdAdapter(Params const& params) : base_(params) {}
|
||||
|
||||
template<class ProblemSize, class ClusterShape, class TileShape>
|
||||
static Params to_underlying_arguments(
|
||||
ProblemSize const& problem_size, KernelHardwareInfo hw_info,
|
||||
ClusterShape const& cluster_shape, TileShape const& tile_shape)
|
||||
{
|
||||
using namespace cute;
|
||||
return Base::to_underlying_arguments(select<0,1,3,2,4>(problem_size), hw_info, select<1,0,2>(cluster_shape), select<1,0,2>(tile_shape));
|
||||
}
|
||||
|
||||
static dim3 get_grid_shape(Params const& params) {
|
||||
return Base::get_grid_shape(params);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool is_valid() {
|
||||
return base_.is_valid();
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
auto get_block_coord() {
|
||||
using namespace cute;
|
||||
return select<1,0,2>(base_.get_block_coord());
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
TileSchedulerBwdAdapter& operator++() {
|
||||
++base_;
|
||||
return *this;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::fmha::kernel
|
||||
Reference in New Issue
Block a user