CUTLASS 3.4.0 (#1286)
* CUTLASS 3.4.0 * Update CHANGELOG.md --------- Co-authored-by: Pradeep Ramani <prramani@nvidia.com>
This commit is contained in:
co-authored by
Pradeep Ramani
parent
b7508e3379
commit
8236f30675
@@ -95,6 +95,7 @@ CUTE_DEVICE dim3 cluster_grid_dims()
|
||||
return gridDim;
|
||||
#elif defined(_MSC_VER)
|
||||
CUTE_RUNTIME_ASSERT("cluster_grid_dims() can only be called on device");
|
||||
return {0, 0, 0};
|
||||
#else
|
||||
return {0, 0, 0};
|
||||
#endif
|
||||
@@ -114,6 +115,7 @@ CUTE_DEVICE dim3 cluster_id_in_grid()
|
||||
return blockIdx;
|
||||
#elif defined(_MSC_VER)
|
||||
CUTE_RUNTIME_ASSERT("cluster_id_in_grid() can only be called on device");
|
||||
return {0, 0, 0};
|
||||
#else
|
||||
return {0, 0, 0};
|
||||
#endif
|
||||
|
||||
@@ -40,6 +40,11 @@
|
||||
# define CUTE_ARCH_TMA_SM90_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED) && \
|
||||
((__CUDACC_VER_MAJOR__ > 12) || ((__CUDACC_VER_MAJOR__ == 12) && (__CUDACC_VER_MINOR__ >= 3)))
|
||||
# define CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED
|
||||
#endif
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
|
||||
@@ -42,6 +42,7 @@
|
||||
|
||||
#include <cute/container/alignment.hpp>
|
||||
#include <cute/container/bit_field.hpp>
|
||||
#include <cute/container/array.hpp>
|
||||
#include <cute/numeric/int.hpp> // to_Format<[u]intX>
|
||||
#include <cute/numeric/half.hpp> // to_Format<half_t>
|
||||
|
||||
@@ -200,6 +201,141 @@ prefetch_tma_descriptor(TmaDescriptor const* desc_ptr)
|
||||
#endif
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Perform a TensorMap modification (by each field)
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Replace tensor pointer directly in GMEM
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
tma_descriptor_replace_addr_in_global_mem(TmaDescriptor const* desc_ptr,
|
||||
void const* const new_tensor_ptr)
|
||||
{
|
||||
#if defined(CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED)
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint64_t const new_desc_addr = reinterpret_cast<uint64_t>(new_tensor_ptr);
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_address.global.b1024.b64 [%0], %1;"
|
||||
:: "l"(gmem_int_desc), "l"(new_desc_addr));
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Using TMA Descriptor modification without CUTE_ARCH_TMA_SM90_ENABLED and CUDA 12.3");
|
||||
#endif
|
||||
}
|
||||
|
||||
// Replace tensor pointer by bringing the tensormap from GMEM into the shared memory
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
tma_descriptor_replace_addr_in_shared_mem(TmaDescriptor& smem_desc,
|
||||
void const* const new_tensor_ptr)
|
||||
{
|
||||
#if defined(CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED)
|
||||
uint32_t smem_int_desc = cast_smem_ptr_to_uint(&smem_desc);
|
||||
uint64_t const new_desc_addr = reinterpret_cast<uint64_t>(new_tensor_ptr);
|
||||
uint64_t const smem_int64_desc = 0;
|
||||
asm volatile (
|
||||
"cvt.u64.u32 %0, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(smem_int_desc));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_address.shared::cta.b1024.b64 [%0], %1;"
|
||||
:: "l"(smem_int64_desc), "l"(new_desc_addr));
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Using TMA Descriptor modification without CUTE_ARCH_TMA_SM90_ENABLED and CUDA 12.3");
|
||||
#endif
|
||||
}
|
||||
|
||||
// Replace tensor dims and strides for GEMMs by bringing the tensormap from GMEM into the shared memory
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
tma_descriptor_replace_dims_strides_in_shared_mem(TmaDescriptor & smem_desc,
|
||||
cute::array<uint32_t, 3> const& prob_shape,
|
||||
cute::array<uint64_t, 3> const& prob_stride)
|
||||
{
|
||||
#if defined(CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED)
|
||||
uint32_t smem_int_desc = cast_smem_ptr_to_uint(&smem_desc);
|
||||
uint64_t const smem_int64_desc = 0;
|
||||
asm volatile (
|
||||
"cvt.u64.u32 %0, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(smem_int_desc));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 0, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[0]));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 1, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[1]));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 2, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[2]));
|
||||
// Strides must be a multiple of 16. Also, stride for the intermost dimension is implicitly 1
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 0, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[1] >> 4));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 1, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[2] >> 4));
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Using TMA Descriptor modification without CUTE_ARCH_TMA_SM90_ENABLED and CUDA 12.3");
|
||||
#endif
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Perform a fused copy and fence operation (needed when modifying tensormap in shared memory)
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
tma_descriptor_cp_fence_release(TmaDescriptor const* gmem_desc_ptr, TmaDescriptor& smem_desc)
|
||||
{
|
||||
#if defined(CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED)
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(gmem_desc_ptr);
|
||||
uint32_t smem_int_desc = cast_smem_ptr_to_uint(&smem_desc);
|
||||
asm volatile (
|
||||
"tensormap.cp_fenceproxy.global.shared::cta.tensormap::generic.release.gpu.sync.aligned [%0], [%1], 128;"
|
||||
:: "l"(gmem_int_desc), "r"(smem_int_desc));
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Using TMA Descriptor modification without CUTE_ARCH_TMA_SM90_ENABLED and CUDA 12.3");
|
||||
#endif
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Perform a release fence operation (needed when modifying tensormap directly in GMEM)
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
tma_descriptor_fence_release()
|
||||
{
|
||||
#if defined(CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED)
|
||||
asm volatile ("fence.proxy.tensormap::generic.release.gpu;");
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Using TMA Descriptor modification without CUTE_ARCH_TMA_SM90_ENABLED and CUDA 12.3");
|
||||
#endif
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Perform a acquire fence operation
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
tma_descriptor_fence_acquire(TmaDescriptor const* desc_ptr)
|
||||
{
|
||||
#if defined(CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED)
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
asm volatile (
|
||||
"fence.proxy.tensormap::generic.acquire.gpu [%0], 128;"
|
||||
:
|
||||
: "l"(gmem_int_desc)
|
||||
: "memory");
|
||||
asm volatile (
|
||||
"cvta.global.u64 %0, %0;"
|
||||
:
|
||||
: "l"(gmem_int_desc), "l"(gmem_int_desc)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Using TMA Descriptor modification without CUTE_ARCH_TMA_SM90_ENABLED and CUDA 12.3");
|
||||
#endif
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
@@ -775,6 +775,104 @@ struct SM90_TMA_STORE
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// TMA_STORE im2col: Initiates a TMA copy, in im2col mode, from shared memory to global memory
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct SM90_TMA_STORE_IM2COL_3D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* const desc_ptr,
|
||||
void const* const smem_ptr,
|
||||
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_n)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.3d.global.shared::cta.im2col_no_offs.bulk_group"
|
||||
" [%0, {%2, %3, %4}], [%1];"
|
||||
:
|
||||
: "l"(gmem_int_desc), "r"(smem_int_ptr),
|
||||
"r"(coord_c), "r"(coord_w), "r"(coord_n)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
struct SM90_TMA_STORE_IM2COL_4D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* const desc_ptr,
|
||||
void const* const smem_ptr,
|
||||
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_n)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.4d.global.shared::cta.im2col_no_offs.bulk_group"
|
||||
" [%0, {%2, %3, %4, %5}], [%1];"
|
||||
:
|
||||
: "l"(gmem_int_desc), "r"(smem_int_ptr),
|
||||
"r"(coord_c), "r"(coord_w), "r"(coord_h), "r"(coord_n)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
struct SM90_TMA_STORE_IM2COL_5D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* const desc_ptr,
|
||||
void const* const smem_ptr,
|
||||
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_d, int32_t const& coord_n)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.5d.global.shared::cta.im2col_no_offs.bulk_group"
|
||||
" [%0, {%2, %3, %4, %5, %6}], [%1];"
|
||||
:
|
||||
: "l"(gmem_int_desc), "r"(smem_int_ptr),
|
||||
"r"(coord_c), "r"(coord_w), "r"(coord_h), "r"(coord_d), "r"(coord_n)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_RUNTIME_ASSERT("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
struct SM90_TMA_STORE_IM2COL
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* const desc_ptr,
|
||||
void const* const smem_ptr,
|
||||
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_n)
|
||||
{
|
||||
return SM90_TMA_STORE_IM2COL_3D::copy(desc_ptr, smem_ptr, coord_c, coord_w, coord_n);
|
||||
}
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* const desc_ptr,
|
||||
void const* const smem_ptr,
|
||||
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_n)
|
||||
{
|
||||
return SM90_TMA_STORE_IM2COL_4D::copy(desc_ptr, smem_ptr, coord_c, coord_w, coord_h, coord_n);
|
||||
}
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* const desc_ptr,
|
||||
void const* const smem_ptr,
|
||||
int32_t const& coord_c, int32_t const& coord_w, int32_t const& coord_h, int32_t const& coord_d, int32_t const& coord_n)
|
||||
{
|
||||
return SM90_TMA_STORE_IM2COL_5D::copy(desc_ptr, smem_ptr, coord_c, coord_w, coord_h, coord_d, coord_n);
|
||||
}
|
||||
};
|
||||
|
||||
// Indicate arrival of warp issuing TMA_STORE
|
||||
CUTE_HOST_DEVICE static void
|
||||
tma_store_arrive() {
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
/**************************************************************************************************
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2023 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
|
||||
@@ -31,26 +31,28 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/config.hpp>
|
||||
|
||||
#include <cute/arch/copy.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
#include <cute/atom/mma_atom.hpp>
|
||||
|
||||
#include <cute/util/type_traits.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <class... Args>
|
||||
struct Copy_Atom;
|
||||
|
||||
template <class CopyOperation, class T>
|
||||
struct Copy_Atom<CopyOperation, T> : Copy_Atom<Copy_Traits<CopyOperation>, T>
|
||||
template <class CopyOperation, class CopyInternalType>
|
||||
struct Copy_Atom<CopyOperation, CopyInternalType> : Copy_Atom<Copy_Traits<CopyOperation>, CopyInternalType>
|
||||
{};
|
||||
|
||||
template <class... Args, class T>
|
||||
struct Copy_Atom<Copy_Traits<Args...>, T>
|
||||
template <class... Args, class CopyInternalType>
|
||||
struct Copy_Atom<Copy_Traits<Args...>, CopyInternalType>
|
||||
: Copy_Traits<Args...>
|
||||
{
|
||||
using Traits = Copy_Traits<Args...>;
|
||||
@@ -61,7 +63,7 @@ struct Copy_Atom<Copy_Traits<Args...>, T>
|
||||
using BitLayoutDst = typename Traits::DstLayout;
|
||||
using BitLayoutRef = typename Traits::RefLayout;
|
||||
|
||||
using ValType = T;
|
||||
using ValType = CopyInternalType;
|
||||
|
||||
using ValLayoutSrc = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutSrc{}));
|
||||
using ValLayoutDst = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutDst{}));
|
||||
@@ -80,7 +82,7 @@ struct Copy_Atom<Copy_Traits<Args...>, T>
|
||||
auto
|
||||
with(TraitsArgs&&... args) const {
|
||||
auto traits = Traits::with(std::forward<TraitsArgs>(args)...);
|
||||
return Copy_Atom<decltype(traits), T>{traits};
|
||||
return Copy_Atom<decltype(traits), CopyInternalType>{traits};
|
||||
}
|
||||
|
||||
//
|
||||
@@ -88,19 +90,19 @@ struct Copy_Atom<Copy_Traits<Args...>, T>
|
||||
//
|
||||
|
||||
// Check and call instruction, or recurse
|
||||
template <class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
template <class SEngine, class SLayout,
|
||||
class DEngine, class DLayout>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
call(Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst) const
|
||||
call(Tensor<SEngine,SLayout> const& src,
|
||||
Tensor<DEngine,DLayout> & dst) const
|
||||
{
|
||||
static_assert(SLayout::rank == 1, "Expected rank-1 src tensor");
|
||||
static_assert(DLayout::rank == 1, "Expected rank-1 dst tensor");
|
||||
|
||||
if constexpr (is_constant<NumValSrc, decltype(size(src))>::value ||
|
||||
is_constant<NumValDst, decltype(size(dst))>::value) {
|
||||
// Dispatch to unpack for instruction
|
||||
// Dispatch to unpack to execute instruction
|
||||
return copy_unpack(*this, src, dst);
|
||||
} else
|
||||
if constexpr (is_tuple<decltype(shape(src))>::value &&
|
||||
@@ -110,7 +112,7 @@ struct Copy_Atom<Copy_Traits<Args...>, T>
|
||||
// ((A,B,C,...)) -> (A,B,C,...)
|
||||
return copy(*this, tensor<0>(src), tensor<0>(dst));
|
||||
} else {
|
||||
static_assert(sizeof(TS) < 0, "No instruction match and no recursion possible.");
|
||||
static_assert(dependent_false<SEngine>, "No instruction match and no recursion possible.");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -135,7 +137,7 @@ struct ThrCopy;
|
||||
|
||||
template <class Copy_Atom,
|
||||
class LayoutCopy_TV, // (tid,vid) -> coord [Need not be 2D...]
|
||||
class ShapeTile_MN> // coord space
|
||||
class ShapeTiler_MN> // coord space
|
||||
struct TiledCopy : Copy_Atom
|
||||
{
|
||||
// Layout information from the CopyAtom
|
||||
@@ -148,8 +150,7 @@ struct TiledCopy : Copy_Atom
|
||||
using AtomNumVal = decltype(size<1>(AtomLayoutRef{}));
|
||||
|
||||
// Layout information for the TiledCopy
|
||||
using Tiler_MN = ShapeTile_MN;
|
||||
using TiledShape_MN = decltype(shape(ShapeTile_MN{}));
|
||||
using Tiler_MN = ShapeTiler_MN;
|
||||
using TiledLayout_TV = LayoutCopy_TV;
|
||||
using TiledNumThr = decltype(size<0>(TiledLayout_TV{}));
|
||||
using TiledNumVal = decltype(size<1>(TiledLayout_TV{}));
|
||||
@@ -172,12 +173,9 @@ struct TiledCopy : Copy_Atom
|
||||
auto
|
||||
tidfrg_S(STensor&& stensor)
|
||||
{
|
||||
constexpr int R = remove_cvref_t<STensor>::rank;
|
||||
static_assert(R >= rank_v<TiledShape_MN>, "Rank of tensor to be partitioned too small.");
|
||||
// Generalize the dimension checks for arbitrary rank
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(stensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(stensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(stensor) >= rank(Tiler_MN{}), "Rank of tensor to be partitioned too small.");
|
||||
|
||||
// Tile the stensor and compute the (src-thr, src-val) -> (ref-thr, ref-val) layout
|
||||
return tile2thrfrg(zipped_divide(stensor,Tiler_MN{}), right_inverse(AtomLayoutRef{}).compose(AtomLayoutSrc{}));
|
||||
}
|
||||
|
||||
@@ -196,17 +194,14 @@ struct TiledCopy : Copy_Atom
|
||||
auto
|
||||
tidfrg_D(DTensor&& dtensor)
|
||||
{
|
||||
constexpr int R = remove_cvref_t<DTensor>::rank;
|
||||
static_assert(R >= rank_v<TiledShape_MN>, "Rank of tensor to be partitioned too small.");
|
||||
// Generalize the dimension checks for arbitrary rank
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(stensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(stensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(dtensor) >= rank(Tiler_MN{}), "Rank of tensor to be partitioned too small.");
|
||||
|
||||
// Tile the dtensor and compute the (dst-thr, dst-val) -> (ref-thr, ref-val) layout
|
||||
return tile2thrfrg(zipped_divide(dtensor,Tiler_MN{}), right_inverse(AtomLayoutRef{}).compose(AtomLayoutDst{}));
|
||||
}
|
||||
|
||||
// Tile a tensor or a layout from shape
|
||||
// (Tile,(RestM,RestN,...))
|
||||
// ((TileM,TileN,...), (RestM,RestN,...))
|
||||
// to shape
|
||||
// ((ThrV,ThrX),FrgV,(RestM,RestN,...))
|
||||
template <class Tensor, class Ref2TrgLayout>
|
||||
@@ -232,7 +227,7 @@ struct TiledCopy : Copy_Atom
|
||||
|
||||
// Transform the tile mode
|
||||
auto tv_tensor = tensor.compose(thrval2mn, _);
|
||||
// ((thrid,val),(RM,RN,...))
|
||||
// ((thrid,val),(RestM,RestN,...))
|
||||
|
||||
// Unfold and return
|
||||
return tv_tensor(make_coord(_,_), _);
|
||||
@@ -253,7 +248,7 @@ struct TiledCopy : Copy_Atom
|
||||
|
||||
auto V = size<0>(tensor);
|
||||
|
||||
auto frg_layout_mn = upcast<TiledNumThr{} * V>(right_inverse(TiledLayout_TV{}).with_shape(TiledShape_MN{}));
|
||||
auto frg_layout_mn = upcast<TiledNumThr{} * V>(right_inverse(TiledLayout_TV{}).with_shape(shape(Tiler_MN{})));
|
||||
// (m,n) -> v_idx -- The shape and order of the V inside of TiledLayout_TV
|
||||
|
||||
auto frg_layout_v = zipped_divide(logical_product(make_layout(V), right_inverse(frg_layout_mn)), make_layout(AtomNumVal{}));
|
||||
@@ -278,7 +273,7 @@ struct TiledCopy : Copy_Atom
|
||||
get_layoutS_TV()
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_S = make_layout(make_shape(TiledShape_MN{}, Int<1>{}));
|
||||
auto ref_S = make_layout(make_shape(shape(Tiler_MN{}), Int<1>{}));
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
return tile2thrfrg(ref_S, right_inverse(AtomLayoutRef{}).compose(AtomLayoutSrc{}))(_,_,Int<0>{});
|
||||
}
|
||||
@@ -290,7 +285,7 @@ struct TiledCopy : Copy_Atom
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
auto layoutS_TV = get_layoutS_TV();
|
||||
// (M,K) -> (thr_idx,val_idx)
|
||||
auto layoutS_MK = right_inverse(layoutS_TV).with_shape(TiledShape_MN{});
|
||||
auto layoutS_MK = right_inverse(layoutS_TV).with_shape(shape(Tiler_MN{}));
|
||||
|
||||
// athrid = (v,m,k) -> thr_idx
|
||||
auto thrID_S = make_layout(size<0>(TiledLayout_TV{}));
|
||||
@@ -303,7 +298,7 @@ struct TiledCopy : Copy_Atom
|
||||
get_layoutD_TV()
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_D = make_layout(make_shape(TiledShape_MN{}, Int<1>{}));
|
||||
auto ref_D = make_layout(make_shape(shape(Tiler_MN{}), Int<1>{}));
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
return tile2thrfrg(ref_D, right_inverse(AtomLayoutRef{}).compose(AtomLayoutDst{}))(_,_,Int<0>{});
|
||||
}
|
||||
@@ -315,7 +310,7 @@ struct TiledCopy : Copy_Atom
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
auto layoutD_TV = get_layoutD_TV();
|
||||
// (M,K) -> (thr_idx,val_idx)
|
||||
auto layoutD_MK = right_inverse(layoutD_TV).with_shape(TiledShape_MN{});
|
||||
auto layoutD_MK = right_inverse(layoutD_TV).with_shape(shape(Tiler_MN{}));
|
||||
|
||||
// athrid = (v,m,k) -> thr_idx
|
||||
auto thrID_D = make_layout(size<0>(TiledLayout_TV{}));
|
||||
@@ -406,51 +401,44 @@ make_tiled_copy_impl(Copy_Atom<Args...> const& atom,
|
||||
// These tile the Copy_Atom as a whole
|
||||
//
|
||||
|
||||
template <class... Args,
|
||||
class TiledMMA>
|
||||
template <class... CArgs, class... MArgs>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_A(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledMMA const& tiled_mma)
|
||||
make_tiled_copy_A(Copy_Atom<CArgs...> const& copy_atom,
|
||||
TiledMMA<MArgs...> const& mma)
|
||||
{
|
||||
using MNK = typename TiledMMA::TiledShape_MNK;
|
||||
return make_tiled_copy_impl(copy_atom, tiled_mma.get_layoutA_TV(), make_shape(size<0>(MNK{}),size<2>(MNK{})));
|
||||
return make_tiled_copy_impl(copy_atom, mma.get_layoutA_TV(), make_shape(tile_size<0>(mma),tile_size<2>(mma)));
|
||||
}
|
||||
|
||||
template <class... Args,
|
||||
class TiledMMA>
|
||||
template <class... CArgs, class... MArgs>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_B(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledMMA const& tiled_mma)
|
||||
make_tiled_copy_B(Copy_Atom<CArgs...> const& copy_atom,
|
||||
TiledMMA<MArgs...> const& mma)
|
||||
{
|
||||
using MNK = typename TiledMMA::TiledShape_MNK;
|
||||
return make_tiled_copy_impl(copy_atom, tiled_mma.get_layoutB_TV(), make_shape(size<1>(MNK{}),size<2>(MNK{})));
|
||||
return make_tiled_copy_impl(copy_atom, mma.get_layoutB_TV(), make_shape(tile_size<1>(mma),tile_size<2>(mma)));
|
||||
}
|
||||
|
||||
template <class... Args,
|
||||
class TiledMMA>
|
||||
template <class... CArgs, class... MArgs>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_C(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledMMA const& tiled_mma)
|
||||
make_tiled_copy_C(Copy_Atom<CArgs...> const& copy_atom,
|
||||
TiledMMA<MArgs...> const& mma)
|
||||
{
|
||||
using MNK = typename TiledMMA::TiledShape_MNK;
|
||||
return make_tiled_copy_impl(copy_atom, tiled_mma.get_layoutC_TV(), make_shape(size<0>(MNK{}),size<1>(MNK{})));
|
||||
return make_tiled_copy_impl(copy_atom, mma.get_layoutC_TV(), make_shape(tile_size<0>(mma),tile_size<1>(mma)));
|
||||
}
|
||||
|
||||
// returns the smallest tiled copy that can retile LayoutC_TV
|
||||
// for use with pipelined epilogues with subtiled stores
|
||||
template <class... Args,
|
||||
class TiledMMA>
|
||||
template <class... CArgs, class... MArgs>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_C_atom(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledMMA const& tiled_mma)
|
||||
make_tiled_copy_C_atom(Copy_Atom<CArgs...> const& copy_atom,
|
||||
TiledMMA<MArgs...> const& mma)
|
||||
{
|
||||
// Truncate the V-layout to just the Copy_Atom, keep the V-order
|
||||
auto layoutC_TV = tiled_mma.get_layoutC_TV();
|
||||
auto copy_V = Int<Copy_Atom<Args...>::NumValSrc>{};
|
||||
auto layoutC_TV = mma.get_layoutC_TV();
|
||||
auto copy_V = Int<Copy_Atom<CArgs...>::NumValSrc>{};
|
||||
CUTE_STATIC_ASSERT_V(copy_V <= size<1>(layoutC_TV));
|
||||
auto layout_TV = composition(layoutC_TV, make_layout(make_shape(size<0>(layoutC_TV), copy_V)));
|
||||
|
||||
@@ -458,8 +446,7 @@ make_tiled_copy_C_atom(Copy_Atom<Args...> const& copy_atom,
|
||||
|
||||
// Tiler -- Find the active elements in the MMA tensor and generate a tiler to extract them
|
||||
// Convert to the awkward by-mode tiler to preserve the modes of the tiled MMA
|
||||
using MNK = typename TiledMMA::TiledShape_MNK;
|
||||
auto mma_tiler = make_shape(size<0>(MNK{}),size<1>(MNK{}));
|
||||
auto mma_tiler = make_shape(tile_size<0>(mma),tile_size<1>(mma));
|
||||
auto mma_zeros = repeat_like(mma_tiler, Int<0>{});
|
||||
|
||||
auto tiler = transform(make_seq<rank(mma_tiler)>{}, [&](auto i) {
|
||||
@@ -474,8 +461,6 @@ make_tiled_copy_C_atom(Copy_Atom<Args...> const& copy_atom,
|
||||
// (tid,vid) -> tile_coord
|
||||
auto layout_tv = composition(left_inverse(tile2mma), layout_TV);
|
||||
|
||||
|
||||
using MNK = typename TiledMMA::TiledShape_MNK;
|
||||
return make_tiled_copy_impl(copy_atom, layout_tv, tiler);
|
||||
}
|
||||
|
||||
@@ -655,8 +640,10 @@ print(TiledCopy<Atom, Args...> const& copy, char const* pad = "")
|
||||
template <class TiledCopy, class ThrIdx>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print(ThrCopy<TiledCopy, ThrIdx> const&)
|
||||
print(ThrCopy<TiledCopy, ThrIdx> const& thr_copy)
|
||||
{
|
||||
print("ThrCopy\n");
|
||||
print(" ThrIdx: "); print(thr_copy.thr_idx_); print("\n");
|
||||
print(TiledCopy{});
|
||||
}
|
||||
|
||||
|
||||
@@ -43,10 +43,14 @@
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <class GmemStrides_, class TmaGBasis_, class TmaSwizzle_>
|
||||
template <class GmemTmaBasisStrides_, class TmaGmemBasis_, class TmaSwizzle_>
|
||||
struct AuxTmaParams {
|
||||
using GmemStrides = GmemStrides_;
|
||||
using GmemStrides = GmemTmaBasisStrides_; // Strides for Gmem mode -> Tma coord mode, may be dynamic
|
||||
GmemStrides g_stride_;
|
||||
using TmaGmemBasis = TmaGmemBasis_; // Layout for Tma box shape -> Gmem mode(s), always static
|
||||
static_assert(is_static<TmaGmemBasis>::value);
|
||||
using TmaSwizzle = TmaSwizzle_; // Tma swizzle, always Swizzle<B,M,S>
|
||||
static_assert(is_static<TmaSwizzle>::value);
|
||||
};
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
@@ -138,13 +142,19 @@ struct Copy_Traits<SM90_TMA_LOAD, NumBitsPerTMA, AuxParams_>
|
||||
// Construct an executable SM90_TMA_LOAD with tma_mbar
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_OP, NumBitsPerTMA>
|
||||
with(uint64_t& tma_mbar, uint16_t const& multicast_mask = 0) const {
|
||||
with(uint64_t& tma_mbar, [[maybe_unused]] uint16_t const& multicast_mask = 0) const {
|
||||
// We accept multicast_mask here to keep the API for both atoms consistent
|
||||
// assert(multicast_mask == 0);
|
||||
(void) multicast_mask;
|
||||
return {tma_desc_, tma_mbar};
|
||||
}
|
||||
|
||||
// Construct an executable SM90_TMA_LOAD with tma_mbar (temp. overloaded for grouped gemm/ptr array gemm)
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_OP, NumBitsPerTMA>
|
||||
with(TmaDescriptor const* new_tma_desc, uint64_t& tma_mbar, [[maybe_unused]] uint16_t const& multicast_mask = 0) const {
|
||||
// We accept multicast_mask here to keep the API for both atoms consistent
|
||||
return {*new_tma_desc, tma_mbar};
|
||||
}
|
||||
|
||||
// Generate the TMA coord tensor
|
||||
template <class GShape>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -251,6 +261,13 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBitsPerTMA, AuxParams_>
|
||||
return {tma_desc_, tma_load_mbar, multicast_mask};
|
||||
}
|
||||
|
||||
// Construct an executable SM90_TMA_LOAD_MULTICAST_OP with tma_mbar (temp. overloaded for grouped gemm/ptr array gemm)
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBitsPerTMA>
|
||||
with(TmaDescriptor const* new_tma_desc, uint64_t& tma_load_mbar, uint16_t const& multicast_mask) const {
|
||||
return {*new_tma_desc, tma_load_mbar, multicast_mask};
|
||||
}
|
||||
|
||||
// Generate the TMA coord tensor
|
||||
template <class GShape>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -508,51 +525,91 @@ coalesce_256(Layout<Shape,Stride> const& layout)
|
||||
return coalesce_256_impl<1>(flat_shape, flat_stride, get<0>(flat_shape), get<0>(flat_stride));
|
||||
}
|
||||
|
||||
template <class Engine, class Layout>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
coalesce_256(Tensor<Engine,Layout> const& tensor)
|
||||
{
|
||||
return make_tensor(tensor.data(), coalesce_256(tensor.layout()));
|
||||
}
|
||||
|
||||
|
||||
// Use a smem_inv_h to read through the GMEM tensor
|
||||
// and construct a TMA Descriptor for the resulting instruction
|
||||
// At the same time, construct the Tma Tensor's Stride to generate
|
||||
// the TMA coordinates that the instruction consumes.
|
||||
//
|
||||
template <class TmaInternalType,
|
||||
class GEngine, class GLayout,
|
||||
class SShape, class SStride,
|
||||
int B, int M, int S>
|
||||
CUTE_HOST_RTC
|
||||
class VShape, class VStride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GMEM Tensor
|
||||
Layout<SShape,SStride> const& smem_inv_h, // smem_idx to hier gmode
|
||||
Swizzle<B,M,S> const& swizzle) // Swizzle fn on smem_idx
|
||||
construct_tma_gbasis(Tensor<GEngine,GLayout> const& gtensor, // The original GMEM Tensor
|
||||
Layout<SShape,SStride> const& slayout, // The layout of SMEM
|
||||
Layout<VShape,VStride> const& cta_v_map) // smem_idx to hier gmode
|
||||
{
|
||||
//
|
||||
// TMA parameter checking
|
||||
//
|
||||
|
||||
CUTE_STATIC_ASSERT_V(product_each(shape(slayout)) == product_each(shape(cta_v_map)),
|
||||
"TMA requires CTA_Tile and SLayout top-level shape equivalence.");
|
||||
|
||||
#if 0
|
||||
print("gtensor : "); print(gtensor); print("\n");
|
||||
print("slayout : "); print(slayout); print("\n");
|
||||
print("cta_v_map : "); print(cta_v_map); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// TMA slayout manipulation
|
||||
//
|
||||
|
||||
// Invert the smem to get the largest contiguous vector in the smem layout
|
||||
// smem idx -> smem coord
|
||||
auto inv_smem_layout = right_inverse(get_nonswizzle_portion(slayout));
|
||||
|
||||
// Compose with the V-Map to convert smem coord (CTA val idx) to gmem mode
|
||||
// smem idx -> gmem mode
|
||||
auto sidx2gmode_full = coalesce(composition(cta_v_map, inv_smem_layout));
|
||||
|
||||
#if 0
|
||||
print("inv_smem_layout : "); print(inv_smem_layout); print("\n");
|
||||
print("sidx2gmode_full : "); print(sidx2gmode_full); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// TMA gtensor truncation
|
||||
//
|
||||
|
||||
// Truncate any incompatibilities -- no starting in the middle of gmodes
|
||||
auto smem_rank = find_if(stride(sidx2gmode_full), [](auto e) {
|
||||
[[maybe_unused]] auto v = basis_value(e);
|
||||
return not is_constant<1,decltype(v)>{};
|
||||
});
|
||||
static_assert(smem_rank > 0, "Could not find a common tile-gmem vectorization. Does the Tile select out major GMEM modes?");
|
||||
|
||||
// Keep only the static-1 basis modes into gmem
|
||||
auto sidx2gmode = take<0,smem_rank>(sidx2gmode_full);
|
||||
|
||||
#if 0
|
||||
print("smem_rank : "); print(smem_rank); print("\n");
|
||||
print("sidx2gmode : "); print(sidx2gmode); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// TMA gtensor manipulation
|
||||
//
|
||||
|
||||
// The smem vector is the same units as gtensor, so compose first and then recast
|
||||
// tma_val_idx:gmem_strides
|
||||
Tensor tile_gstride = recast<TmaInternalType>(gtensor.compose(smem_inv_h));
|
||||
auto tile_gstride = recast<TmaInternalType>(gtensor.compose(sidx2gmode)).layout();
|
||||
// Coalesce modes up to size-256 (the maximum TMA box extent in units of TmaInternalType)
|
||||
// tma_box_shape:gmem_strides
|
||||
Tensor tma_gstride = coalesce_256(tile_gstride);
|
||||
auto tma_gstride = coalesce_256(tile_gstride);
|
||||
|
||||
// Perform the tiling to the gmem vector again, but with indirections to the gtensor modes
|
||||
// Perform the tiling, recast, and coalesce to the gmem vector again, but with indirections to the gtensor modes
|
||||
auto gbasis = make_identity_layout(shape(gtensor));
|
||||
auto tile_gbasis_tmp = gbasis.compose(smem_inv_h);
|
||||
auto tile_gbasis_tmp = gbasis.compose(sidx2gmode);
|
||||
|
||||
// Instead of the recast (gbasis doesn't have type info), replace the shape with the already-recasted shape
|
||||
// tma_box_shape:gmem_mode
|
||||
auto tile_gbasis = make_layout(shape(tile_gstride), stride(tile_gbasis_tmp));
|
||||
|
||||
// Recast the original tensor for shape inspections
|
||||
auto gtensor_T = recast<TmaInternalType>(gtensor);
|
||||
// "Coalesce" the tile basis into a compatible shape with the tma_gstride
|
||||
auto tma_gbasis_tile = tile_gbasis.compose(make_layout(wrap(shape(tma_gstride))));
|
||||
|
||||
// Recast the original tensor for shape/stride inspections
|
||||
Tensor gtensor_T = recast<TmaInternalType>(gtensor);
|
||||
|
||||
// Find missing bases that don't appear in tile_gbasis
|
||||
// NOTE This is essentially ArithmeticTuple complement...
|
||||
// NOTE in pursuit of implementing an ArithmeticTuple logical_divide for smem_inv_h
|
||||
auto tile_gbasis_remaining_stride = filter_tuple(flatten(shape (gtensor_T)), flatten(stride(gtensor_T)),
|
||||
flatten(stride(gbasis)),
|
||||
[&](auto s, auto d, auto e)
|
||||
@@ -561,7 +618,7 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
return cute::tuple<>{}; // If size-1 or stride-0, then don't append
|
||||
} else {
|
||||
using E = decltype(e);
|
||||
auto has_e = any_of(stride(tile_gbasis), [] (auto tb) { return tb == E{}; });
|
||||
auto has_e = any_of(flatten(stride(tma_gbasis_tile)), [] (auto tb) { return tb == E{}; });
|
||||
if constexpr (decltype(has_e)::value) {
|
||||
return cute::tuple<>{}; // If d was found, then don't append
|
||||
} else {
|
||||
@@ -569,13 +626,10 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
}
|
||||
}
|
||||
});
|
||||
auto tile_gbasis_remaining_rank = rank(tile_gbasis_remaining_stride);
|
||||
|
||||
// "Coalesce" the tile basis into a compatible shape with the tma
|
||||
auto tma_gbasis_tile = tile_gbasis.compose(make_layout(wrap(shape(tma_gstride))));
|
||||
|
||||
// Append the remaining basis modes that contribute to the TMA with size-1
|
||||
auto tma_gbasis_full = make_layout(tuple_cat(wrap( shape(tma_gbasis_tile)), wrap(repeat<tile_gbasis_remaining_rank>(Int<1>{}))),
|
||||
auto tile_gbasis_remaining_shape = repeat<rank(tile_gbasis_remaining_stride)>(Int<1>{});
|
||||
auto tma_gbasis_full = make_layout(tuple_cat(wrap( shape(tma_gbasis_tile)), wrap(tile_gbasis_remaining_shape )),
|
||||
tuple_cat(wrap(stride(tma_gbasis_tile)), wrap(tile_gbasis_remaining_stride)));
|
||||
|
||||
// Group the trailing modes to make this max rank-5 -- TMA rank limitation
|
||||
@@ -583,15 +637,98 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
auto tma_gbasis = group<cute::min(rank(tma_gbasis_full),4),-1>(tma_gbasis_full);
|
||||
|
||||
#if 0
|
||||
print("smem_inv_h : "); print(smem_inv_h); print("\n");
|
||||
print("gtensor : "); print(gtensor); print("\n");
|
||||
print("tile_gstride : "); print(tile_gstride); print("\n");
|
||||
print("tma_gstride : "); print(tma_gstride); print("\n");
|
||||
print("gbasis : "); print(gbasis); print("\n");
|
||||
print("tile_gbasis : "); print(tile_gbasis); print("\n");
|
||||
print("tile_gbasis : "); print(tma_gbasis_tile); print("\n");
|
||||
print("tma_gbasis : "); print(tma_gbasis); print("\n");
|
||||
#endif
|
||||
|
||||
return tma_gbasis;
|
||||
}
|
||||
|
||||
template <class GEngine, class GLayout,
|
||||
class TmaGmemBasisStride,
|
||||
class ShapeT, size_t TmaRank>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
fill_tma_gmem_shape_stride(Tensor<GEngine,GLayout> const& gtensor, // Gmem Shapes and Strides, in units of TmaInternalType
|
||||
TmaGmemBasisStride const& tma_gbasis_stride, // Map Tma mode idx -> Gmem mode(s)
|
||||
cute::array<ShapeT, TmaRank> & gmem_prob_shape, // Tma Shapes, uint32_t or uin64_t
|
||||
cute::array<uint64_t, TmaRank> & gmem_prob_stride) // Tma Strides
|
||||
{
|
||||
static_assert(is_tuple<TmaGmemBasisStride>::value);
|
||||
static_assert(is_same<uint32_t, ShapeT>::value || is_same<uint64_t, ShapeT>::value);
|
||||
|
||||
using TmaInternalType = typename GEngine::value_type;
|
||||
constexpr int tma_rank = decltype(rank(tma_gbasis_stride))::value;
|
||||
static_assert(TmaRank >= tma_rank);
|
||||
|
||||
auto gmem_shape = shape(gtensor);
|
||||
auto gmem_stride = stride(gtensor);
|
||||
// Use the indirections in tma_gbasis_stride into gtensor to construct the tma gmem shapes/strides
|
||||
for_each(make_seq<tma_rank>{}, [&](auto i) {
|
||||
constexpr int tma_i_rank = decltype(rank<i>(tma_gbasis_stride))::value;
|
||||
if constexpr (tma_i_rank == 1) {
|
||||
// Trivial contribution of this gmem mode to this tma mode
|
||||
auto ej = unwrap(get<i>(tma_gbasis_stride));
|
||||
gmem_prob_shape[i] = basis_get(ej, gmem_shape);
|
||||
gmem_prob_stride[i] = basis_get(ej, gmem_stride) * sizeof_bits_v<TmaInternalType> / 8;
|
||||
} else {
|
||||
// Apply a recurrence to each gmem mode that contributes to this tma mode
|
||||
for_each(get<i>(tma_gbasis_stride), [&](auto ej) {
|
||||
// Problem shape
|
||||
uint64_t shape_j = basis_get(ej, gmem_shape);
|
||||
// Problem stride (in bytes)
|
||||
uint64_t stride_j = basis_get(ej, gmem_stride) * sizeof_bits_v<TmaInternalType> / 8;
|
||||
|
||||
uint64_t old_stride = gmem_prob_stride[i];
|
||||
gmem_prob_stride[i] = gcd(gmem_prob_stride[i], stride_j);
|
||||
|
||||
if (gmem_prob_stride[i] != 0) {
|
||||
// Recurrence: g_shape = (s_i - 1) * (d_i / gcd_j d_j) + 1
|
||||
gmem_prob_shape[i] = (gmem_prob_shape[i]-1) * (old_stride / gmem_prob_stride[i])
|
||||
+ (shape_j-1) * (stride_j / gmem_prob_stride[i])
|
||||
+ 1;
|
||||
} else {
|
||||
gmem_prob_shape[i] = shape_j;
|
||||
}
|
||||
});
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
// Overload for an existing Copy_Traits
|
||||
template <class GEngine, class GLayout,
|
||||
class Op, class Bits, class Aux,
|
||||
class ShapeT, size_t TmaRank>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
fill_tma_gmem_shape_stride(Copy_Traits<Op,Bits,Aux> const& tma_traits,
|
||||
Tensor<GEngine,GLayout> const& gtensor, // Gmem Shapes and Strides, value_type = TmaInternalType
|
||||
cute::array<ShapeT, TmaRank> & gmem_prob_shape, // Tma Shapes, uint32_t or uin64_t
|
||||
cute::array<uint64_t, TmaRank> & gmem_prob_stride) // Tma Strides
|
||||
{
|
||||
return fill_tma_gmem_shape_stride(gtensor, stride(typename Aux::TmaGmemBasis{}),
|
||||
gmem_prob_shape, gmem_prob_stride);
|
||||
}
|
||||
|
||||
// Use a sidx2gmode to read through the GMEM tensor
|
||||
// and construct a TMA Descriptor for the resulting instruction
|
||||
// At the same time, construct the Tma Tensor's Stride to generate
|
||||
// the TMA coordinates that the instruction consumes.
|
||||
//
|
||||
template <class TmaInternalType,
|
||||
class GEngine, class GLayout,
|
||||
class TShape, class TStride,
|
||||
int B, int M, int S>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GMEM Tensor
|
||||
Layout<TShape,TStride> const& tma_gbasis, // TMA mode -> GMEM mode mapping
|
||||
Swizzle<B,M,S> const& swizzle, // Swizzle fn on smem_idx
|
||||
uint32_t num_multicast) // The number of CTAs in multicasting
|
||||
{
|
||||
//
|
||||
// TMA desc creation
|
||||
//
|
||||
@@ -602,31 +739,16 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
// TMA gmem desc info
|
||||
//
|
||||
|
||||
// Recast the original tensor for shape/stride inspections
|
||||
Tensor gtensor_T = recast<TmaInternalType>(gtensor);
|
||||
|
||||
void* gmem_address = (void*) raw_pointer_cast(gtensor_T.data());
|
||||
auto gmem_layout = gtensor_T.layout();
|
||||
|
||||
cute::array<uint64_t, 5> gmem_prob_shape = {1,1,1,1,1};
|
||||
cute::array<uint64_t, 5> gmem_prob_stride = {0,0,0,0,0};
|
||||
// Use the indirections in tma_gbasis in the values of flat_glayout to construct the gmem shapes/strides
|
||||
for_each(make_seq<tma_dim>{}, [&](auto i) {
|
||||
for_each(stride<i>(tma_gbasis), [&](auto ej) {
|
||||
// Problem stride
|
||||
uint64_t stride_j = ceil_div(basis_get(ej, stride(gmem_layout)) * sizeof_bits_v<TmaInternalType>, 8);
|
||||
uint64_t old_stride = gmem_prob_stride[i];
|
||||
gmem_prob_stride[i] = gcd(gmem_prob_stride[i], stride_j);
|
||||
|
||||
// Problem shape
|
||||
uint64_t shape_j = basis_get(ej, shape(gmem_layout));
|
||||
if (gmem_prob_stride[i] != 0) {
|
||||
// Recurrence: g_shape = (s_i - 1) * (d_i / gcd_j d_j) + 1
|
||||
gmem_prob_shape[i] = (gmem_prob_shape[i]-1) * (old_stride / gmem_prob_stride[i])
|
||||
+ (shape_j-1) * (stride_j / gmem_prob_stride[i])
|
||||
+ 1;
|
||||
} else {
|
||||
gmem_prob_shape[i] = shape_j;
|
||||
}
|
||||
});
|
||||
});
|
||||
fill_tma_gmem_shape_stride(gtensor_T, stride(tma_gbasis), gmem_prob_shape, gmem_prob_stride);
|
||||
|
||||
assert((reinterpret_cast<uint64_t>(gmem_address) & 0b1111) == 0); // Address must be 16B-aligned
|
||||
|
||||
@@ -663,6 +785,13 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
for_each(make_seq<tma_dim>{}, [&](auto i) {
|
||||
smem_box_shape[i] *= size<i>(tma_gbasis);
|
||||
});
|
||||
// Finally, truncate the tma box by the num_multicast
|
||||
for (uint32_t i = tma_dim-1, multicast = num_multicast; multicast > 1; --i) {
|
||||
assert(smem_box_shape[i] % multicast == 0 || multicast % smem_box_shape[i] == 0);
|
||||
uint32_t new_mult = ceil_div(multicast, smem_box_shape[i]);
|
||||
smem_box_shape[i] = ceil_div(smem_box_shape[i], multicast);
|
||||
multicast = new_mult;
|
||||
}
|
||||
|
||||
assert(smem_box_shape[0] >= (uint32_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint32_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
@@ -740,26 +869,27 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
auto recast_ratio = cute::ratio(Int<sizeof_bits<typename GEngine::value_type>::value>{},
|
||||
Int<sizeof_bits< TmaInternalType>::value>{});
|
||||
|
||||
auto gbasis = make_basis_like(shape(gtensor));
|
||||
|
||||
// Finally, get the inverse permutation of the E<i> bases for the mocked gmem stride
|
||||
// NOTE This is essentially ArithmeticTuple inverse...
|
||||
auto gmem_stride_bases = transform_leaf(stride(gbasis), [&](auto ei) {
|
||||
auto gmem_tma_basis_stride = transform_leaf(gbasis, [&](auto ei) {
|
||||
auto si = basis_get(ei, shape(gmem_layout));
|
||||
auto di = basis_get(ei, stride(gmem_layout));
|
||||
if constexpr (is_constant<1, decltype(si)>::value || is_constant<0, decltype(di)>::value) {
|
||||
return Int<0>{}; // If size-1 or stride-0, return arithmetic identity -- no contribution to the TMA
|
||||
} else {
|
||||
auto tma_gbasis_stride = stride(tma_gbasis);
|
||||
auto tma_gmem_basis_stride = stride(tma_gbasis);
|
||||
// Find j such that E<i> is in stride<j>(tma_gbasis)
|
||||
using EI = decltype(ei);
|
||||
[[maybe_unused]] auto j = find_if(tma_gbasis_stride, [&](auto tma_stride_j) { return any_of(tma_stride_j, [&](auto dj) { return dj == EI{}; }); });
|
||||
if constexpr (decltype(j == rank(tma_gbasis_stride))::value) {
|
||||
[[maybe_unused]] auto j = find_if(tma_gmem_basis_stride, [&](auto tma_stride_j) { return any_of(tma_stride_j, [&](auto dj) { return dj == EI{}; }); });
|
||||
if constexpr (decltype(j == rank(tma_gmem_basis_stride))::value) {
|
||||
return Int<0>{}; // If not-found, return arithmetic identity -- no contribution to the TMA
|
||||
} else
|
||||
if constexpr (decltype(j == Int<0>{})::value) {
|
||||
auto scale = recast_ratio * basis_get(ei, stride(gtensor));
|
||||
return E<j>{} * scale; // Return TMA Coord basis -- with a recast scale factor
|
||||
} else
|
||||
if constexpr (decltype(rank<j>(tma_gbasis_stride) == Int<1>{})::value) {
|
||||
if constexpr (decltype(rank<j>(tma_gmem_basis_stride) == Int<1>{})::value) {
|
||||
return E<j>{}; // Return TMA Coord basis -- known scale of Int<1>{}
|
||||
} else {
|
||||
int32_t scale = ceil_div(int32_t(di * sizeof_bits_v<TmaInternalType> / cute::max(gmem_prob_stride[j], 16)), 8);
|
||||
@@ -768,14 +898,64 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
}
|
||||
});
|
||||
|
||||
#if 0
|
||||
print("tma_gbasis : "); print(gmem_stride_bases); print("\n");
|
||||
#endif
|
||||
#if 0
|
||||
print("gmem_tma_basis_stride : "); print(gmem_tma_basis_stride); print("\n");
|
||||
#endif
|
||||
|
||||
using AuxParams = AuxTmaParams<decltype(gmem_stride_bases),
|
||||
using AuxParams = AuxTmaParams<decltype(gmem_tma_basis_stride),
|
||||
decltype(tma_gbasis),
|
||||
decltype(swizzle)>;
|
||||
return cute::make_tuple(tma_desc, AuxParams{gmem_stride_bases});
|
||||
return cute::make_tuple(tma_desc, AuxParams{gmem_tma_basis_stride});
|
||||
}
|
||||
|
||||
template <class TmaInternalType,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class VShape, class VStride>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy_atom(CopyOp,
|
||||
Tensor<GEngine,GLayout> const& gtensor, // Full GMEM Tensor
|
||||
SLayout const& slayout, // CTA Tile of SMEM, potentially swizzled
|
||||
uint32_t const& num_multicast, // The number of CTAs involved in multicasting
|
||||
Layout<VShape,VStride> const& cta_v_map) // V: CTA val idx -> gmem mode
|
||||
{
|
||||
//
|
||||
// TMA truncated layout
|
||||
//
|
||||
|
||||
auto smem_swizzle = get_swizzle_portion(slayout);
|
||||
auto smem_layout = get_nonswizzle_portion(slayout);
|
||||
|
||||
auto tma_gbasis = detail::construct_tma_gbasis<TmaInternalType>(gtensor, smem_layout, cta_v_map);
|
||||
|
||||
//
|
||||
// Construct the TMA Desc and the strides of the TMA Tensor
|
||||
//
|
||||
|
||||
auto [tma_desc, aux_params] = detail::make_tma_copy_desc<TmaInternalType>(gtensor,
|
||||
tma_gbasis,
|
||||
smem_swizzle,
|
||||
num_multicast);
|
||||
|
||||
//
|
||||
// Construct the Copy_Traits
|
||||
//
|
||||
|
||||
constexpr int num_bits_per_tma = decltype(size(tma_gbasis))::value * sizeof_bits_v<TmaInternalType>;
|
||||
using Traits = Copy_Traits<CopyOp, cute::C<num_bits_per_tma>, decltype(aux_params)>;
|
||||
using Atom = Copy_Atom<Traits, typename GEngine::value_type>;
|
||||
|
||||
Traits tma_traits{tma_desc, aux_params};
|
||||
|
||||
#if 0
|
||||
print("num_bits_per_tma : "); print(num_bits_per_tma); print("\n");
|
||||
print("g_stride_bases : "); print(tma_traits.aux_params_.g_stride_); print("\n");
|
||||
#endif
|
||||
|
||||
// Return the Copy_Atom
|
||||
return Atom{tma_traits};
|
||||
}
|
||||
|
||||
// The "logical TMA tid" is a map from the CTA rank to its logical id
|
||||
@@ -790,122 +970,46 @@ template <class TmaInternalType,
|
||||
class VShape, class VStride>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy_tiled(CopyOp,
|
||||
make_tma_copy_tiled(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor, // Full GMEM Tensor
|
||||
SLayout const& slayout, // CTA Tile of SMEM
|
||||
Layout<TShape,TStride> const& cta_t_map, // T: CTA thr idx -> logical TMA tid
|
||||
Layout<VShape,VStride> const& cta_v_map) // V: CTA val idx -> gmem mode
|
||||
{
|
||||
//
|
||||
// TMA parameter checking
|
||||
//
|
||||
|
||||
CUTE_STATIC_ASSERT_V(product_each(shape(slayout)) == product_each(shape(cta_v_map)),
|
||||
"TMA requires CTA_Tile and SLayout top-level shape equivalence.");
|
||||
CUTE_STATIC_ASSERT_V(size(slayout) % cosize(cta_t_map) == Int<0>{},
|
||||
"Number of active CTAs in TMA must divide domain size of slayout.");
|
||||
|
||||
//
|
||||
// TMA slayout manipulation
|
||||
//
|
||||
|
||||
// Invert the smem to get the largest contiguous vector in the smem layout
|
||||
// smem idx -> smem coord
|
||||
auto inv_smem_layout = right_inverse(get_nonswizzle_portion(slayout));
|
||||
|
||||
// Compose with the V-Map to convert smem coord (CTA val idx) to gmem mode
|
||||
// smem idx -> gmem mode
|
||||
auto sidx_to_gmode = coalesce(composition(cta_v_map, inv_smem_layout));
|
||||
|
||||
#if 0
|
||||
print("g_tensor : "); print(gtensor); print("\n");
|
||||
print("s_layout : "); print(slayout); print("\n");
|
||||
print("cta_t_map : "); print(cta_t_map); print("\n");
|
||||
print("cta_v_map : "); print(cta_v_map); print("\n");
|
||||
print("inv_s_layout : "); print(inv_smem_layout); print("\n");
|
||||
print("sidx_to_gmode : "); print(sidx_to_gmode); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// TMA gtensor manipulation
|
||||
//
|
||||
|
||||
// Generate a TupleBasis for the gtensor
|
||||
// gmem coord -> gmem coord
|
||||
auto glayout_basis = make_identity_layout(shape(gtensor));
|
||||
|
||||
// Tile the modes of gtensor with the truncated cta_v_map o inv_smem_layout_trunc
|
||||
// smem idx -> gmem coord
|
||||
auto tma_layout_full = flatten(composition(glayout_basis, sidx_to_gmode));
|
||||
|
||||
// Truncate any incompatibilities -- no starting in the middle of gmodes
|
||||
auto smem_rank = find_if(stride(tma_layout_full), [](auto e) {
|
||||
[[maybe_unused]] auto v = basis_value(e);
|
||||
return not is_constant<1,decltype(v)>{};
|
||||
});
|
||||
static_assert(smem_rank > 0, "Could not find a common tile-gmem vectorization. Does the Tile select out major GMEM modes?");
|
||||
|
||||
// Keep only the static-1 basis modes into gmem
|
||||
auto tma_layout_trunc = take<0,smem_rank>(tma_layout_full);
|
||||
|
||||
// Keep only the portion each multicast CTA will be responsible for
|
||||
auto tma_layout_v = composition(tma_layout_trunc, shape_div(size(tma_layout_trunc), cosize(cta_t_map)));
|
||||
|
||||
#if 0
|
||||
print("glayout_basis : "); print(glayout_basis); print("\n");
|
||||
print("tma_layout_full : "); print(tma_layout_full); print("\n");
|
||||
|
||||
print("tma_layout_trunc: "); print(tma_layout_trunc); print("\n");
|
||||
print("tma_layout_v : "); print(tma_layout_v); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// Construct the TMA Desc and the strides of the TMA Tensor
|
||||
//
|
||||
|
||||
auto [tma_desc, aux_params] = detail::make_tma_copy_desc<TmaInternalType>(gtensor,
|
||||
tma_layout_v,
|
||||
get_swizzle_portion(slayout));
|
||||
|
||||
//
|
||||
// Construct the Copy_Traits
|
||||
//
|
||||
|
||||
using T = typename GEngine::value_type;
|
||||
constexpr int num_bits_per_tma = decltype(size(tma_layout_trunc))::value * sizeof_bits_v<T>;
|
||||
using Traits = Copy_Traits<CopyOp, cute::C<num_bits_per_tma>, decltype(aux_params)>;
|
||||
using Atom = Copy_Atom<Traits, T>;
|
||||
|
||||
Traits tma_traits{tma_desc, aux_params};
|
||||
|
||||
#if 0
|
||||
print("num_bits_per_tma : "); print(num_bits_per_tma); print("\n");
|
||||
print("g_stride_bases : "); print(tma_traits.aux_params_.g_stride_); print("\n");
|
||||
#endif
|
||||
Copy_Atom atom = make_tma_copy_atom<TmaInternalType>(copy_op, gtensor, slayout,
|
||||
cosize(cta_t_map), cta_v_map);
|
||||
|
||||
//
|
||||
// Construct the TiledCopy
|
||||
//
|
||||
|
||||
auto cta_tiler = product_each(shape(cta_v_map));
|
||||
[[maybe_unused]] auto cta_tiler = product_each(shape(cta_v_map));
|
||||
|
||||
auto num_elems_per_tma = size<1>(typename decltype(atom)::RefLayout{}) / Int<sizeof_bits_v<typename GEngine::value_type>>{};
|
||||
|
||||
// smem idx -> smem coord
|
||||
auto inv_smem_layout = right_inverse(get_nonswizzle_portion(slayout));
|
||||
// CTA V -> smem_coord
|
||||
auto layout_v = composition(inv_smem_layout, size(tma_layout_trunc));
|
||||
auto layout_v = composition(inv_smem_layout, num_elems_per_tma);
|
||||
// Scale that up to cover all of the smem_coords
|
||||
auto layout_V = tile_to_shape(make_layout(layout_v), size(cta_v_map));
|
||||
// CTA T -> smem idx
|
||||
auto layout_t = make_layout(cosize(cta_t_map), shape_div(size(tma_layout_trunc), cosize(cta_t_map)));
|
||||
auto layout_t = make_layout(cosize(cta_t_map), shape_div(num_elems_per_tma, cosize(cta_t_map)));
|
||||
// CTA TID -> smem coord
|
||||
auto layout_T = composition(inv_smem_layout, composition(layout_t, cta_t_map));
|
||||
// Combine with the T mapping
|
||||
auto layout_TV = make_layout(layout_T, layout_V);
|
||||
[[maybe_unused]] auto layout_TV = make_layout(layout_T, layout_V);
|
||||
|
||||
#if 0
|
||||
print("cta_tiler : "); print(cta_tiler); print("\n");
|
||||
print("layout_VT : "); print(layout_VT); print("\n");
|
||||
print("layout_v : "); print(layout_v); print("\n");
|
||||
print("layout_V : "); print(layout_V); print("\n");
|
||||
print("layout_t : "); print(layout_t); print("\n");
|
||||
print("layout_T : "); print(layout_T); print("\n");
|
||||
print("layout_TV : "); print(layout_TV); print("\n");
|
||||
#endif
|
||||
|
||||
return TiledCopy<Atom, decltype(layout_TV), decltype(cta_tiler)>{tma_traits};
|
||||
return TiledCopy<decltype(atom), decltype(layout_TV), decltype(cta_tiler)>{atom};
|
||||
}
|
||||
|
||||
} // end namespace detail
|
||||
@@ -982,7 +1086,7 @@ make_tma_copy_tiled(CopyOp,
|
||||
|
||||
copy(tma.with(barrier, mcast_mask), tAgA, tAsA); // copy with supporting TMA params
|
||||
*/
|
||||
template <class TmaInternalType,
|
||||
template <class TmaInternalType = void,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
@@ -996,37 +1100,16 @@ make_tma_copy(CopyOp const& copy_op,
|
||||
CTA_Tiler const& cta_tiler,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
auto cta_v_tile = make_identity_layout(shape(gtensor)).compose(cta_tiler);
|
||||
auto cta_t_tile = make_layout(cluster_size);
|
||||
return detail::make_tma_copy_tiled<TmaInternalType>(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
cta_t_tile,
|
||||
cta_v_tile);
|
||||
auto cta_v_tile = make_identity_layout(shape(gtensor)).compose(cta_tiler);
|
||||
auto cta_t_tile = make_layout(cluster_size);
|
||||
// Prefer TmaInternalType if specified. Fallback to GEngine::value_type
|
||||
using TmaType = conditional_t<is_same<void, TmaInternalType>::value, typename GEngine::value_type, TmaInternalType>;
|
||||
return detail::make_tma_copy_tiled<TmaType>(copy_op,
|
||||
gtensor, slayout,
|
||||
cta_t_tile, cta_v_tile);
|
||||
}
|
||||
|
||||
// Explicit defaulting
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tile,
|
||||
class Cluster_Size>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tile const& cta_tile,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
using TmaInternalType = typename GEngine::value_type;
|
||||
return make_tma_copy<TmaInternalType>(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
cta_tile,
|
||||
cluster_size);
|
||||
}
|
||||
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout>
|
||||
@@ -1039,6 +1122,7 @@ make_tma_copy(CopyOp const& copy_op,
|
||||
return make_tma_copy(copy_op, gtensor, slayout, product_each(shape(slayout)), Int<1>{});
|
||||
}
|
||||
|
||||
// Explicit defaulting
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
@@ -1053,4 +1137,86 @@ make_tma_copy(CopyOp const& copy_op,
|
||||
return make_tma_copy(copy_op, gtensor, slayout, product_each(shape(slayout)), cluster_size);
|
||||
}
|
||||
|
||||
////////////////////////////////////
|
||||
// Experimental Make TMA Atom and Partitioner
|
||||
///////////////////////////////////
|
||||
|
||||
template <class TmaInternalType = void,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tiler,
|
||||
class Cluster_Size>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_atom(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tiler const& cta_tiler,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
auto cta_v_tile = make_identity_layout(shape(gtensor)).compose(cta_tiler);
|
||||
// Prefer TmaInternalType if specified. Fallback to GEngine::value_type
|
||||
using TmaType = conditional_t<is_same<void, TmaInternalType>::value, typename GEngine::value_type, TmaInternalType>;
|
||||
return detail::make_tma_copy_atom<TmaType>(copy_op,
|
||||
gtensor, slayout,
|
||||
size(cluster_size), cta_v_tile);
|
||||
}
|
||||
|
||||
// The "VectorCopy Partitioner" for TMA
|
||||
template <class... Args,
|
||||
class CtaCoord,
|
||||
class TShape, class TStride,
|
||||
class SEngine, class SLayout,
|
||||
class GEngine, class GLayout>
|
||||
CUTE_DEVICE
|
||||
auto
|
||||
tma_partition(Copy_Atom<Args...> const& copy_atom,
|
||||
CtaCoord const& cta_coord,
|
||||
Layout<TShape,TStride> const& cta_layout, // T: CTA coord -> logical multicast id
|
||||
Tensor<SEngine,SLayout> const& stensor, // SMEM Tensor (TMATile, Iter)
|
||||
Tensor<GEngine,GLayout> const& gtensor) // GMEM Tensor (TMATile, Iter)
|
||||
{
|
||||
// Invert the smem to get the largest contiguous vector in the smem layout
|
||||
Layout inv_smem_layout = right_inverse(get_nonswizzle_portion(layout<0>(stensor)));
|
||||
// Scale that up to cover all of the smem_coords
|
||||
Layout layout_v = tile_to_shape(make_layout(inv_smem_layout), size<0>(stensor));
|
||||
|
||||
// Factor out the single-instrucion portion
|
||||
Layout tma_layout_v = make_layout(Int<Copy_Atom<Args...>::NumValSrc>{});
|
||||
Layout layout_V = logical_divide(layout_v, tma_layout_v);
|
||||
|
||||
// Transform tile mode and coalesce
|
||||
Tensor gtensor_v = coalesce(gtensor.compose(layout_V, _), Shape<Shape<_1,_1>,_1>{}); // ((TMA,TMA_Iter),Iter)
|
||||
Tensor stensor_v = coalesce(stensor.compose(layout_V, _), Shape<Shape<_1,_1>,_1>{}); // ((TMA,TMA_Iter),Iter)
|
||||
|
||||
#if 0
|
||||
if (thread0()) {
|
||||
print("layout_V : "); print(layout_V); print("\n");
|
||||
print("gtensor_v : "); print(gtensor_v); print("\n");
|
||||
print("stensor_v : "); print(stensor_v); print("\n");
|
||||
}
|
||||
#endif
|
||||
|
||||
// Restride the cta-into-tma-instr layout
|
||||
Layout tma_layout_t = composition(make_layout(Int<1>{}, shape_div(size(tma_layout_v), cosize(cta_layout))), cta_layout);
|
||||
Layout tma_layout_tv = make_layout(tma_layout_t, tma_layout_v);
|
||||
|
||||
// Transform TMA mode
|
||||
Tensor gtensor_tv = gtensor_v.compose(make_tile(tma_layout_tv, _), _); // (((Thr,Frg),TMA_Iter),Iter)
|
||||
Tensor stensor_tv = stensor_v.compose(make_tile(tma_layout_tv, _), _); // (((Thr,Frg),TMA_Iter),Iter)
|
||||
|
||||
#if 0
|
||||
if (thread0()) {
|
||||
print("tma_layout_tv : "); print(tma_layout_tv); print("\n");
|
||||
print("gtensor_tv : "); print(gtensor_tv); print("\n");
|
||||
print("stensor_tv : "); print(stensor_tv); print("\n");
|
||||
}
|
||||
#endif
|
||||
|
||||
// Slice and group Frg,TMA_Iter and return
|
||||
auto c = make_coord(make_coord(make_coord(cta_coord, _), _), _);
|
||||
return cute::make_tuple(group_modes<0,2>(gtensor_tv(c)), group_modes<0,2>(stensor_tv(c)));
|
||||
}
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
+150
-282
@@ -31,12 +31,12 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/config.hpp>
|
||||
#include <cute/arch/mma.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cute/arch/mma.hpp>
|
||||
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cute/util/type_traits.hpp>
|
||||
|
||||
namespace cute {
|
||||
@@ -196,39 +196,37 @@ struct MMA_Atom<MMA_Traits<Args...>>
|
||||
template <class TiledMMA, class ThrCoord>
|
||||
struct ThrMMA;
|
||||
|
||||
// @tparam MMA_Atom The MMA_Atom to use in the TiledMMA
|
||||
// @tparam AtomLayoutMNK The MNK-tiling of the Atom to be performed.
|
||||
// @tparam PermuationsMNK Permutations to apply to each MNK-mode before tiling for the Atom.
|
||||
template <class MMA_Atom,
|
||||
class AtomLayoutMNK = Layout<Shape<_1,_1,_1>>,
|
||||
class ValLayoutMNK = Layout<Shape<_1,_1,_1>>,
|
||||
class PermutationsMNK = Tile<Underscore,Underscore,Underscore>>
|
||||
class AtomLayoutMNK,
|
||||
class PermutationMNK = Tile<Underscore,Underscore,Underscore>>
|
||||
struct TiledMMA : MMA_Atom
|
||||
{
|
||||
static_assert(rank_v<AtomLayoutMNK> == 3, "TiledMMA requires rank-3 AtomLayoutMNK");
|
||||
static_assert(rank_v<ValLayoutMNK> == 3, "TiledMMA requires rank-3 ValLayoutMNK");
|
||||
static_assert(rank_v<PermutationsMNK> == 3, "TiledMMA requires rank-3 PermutationsMNK");
|
||||
|
||||
using Atom = MMA_Atom;
|
||||
using AtomShape_MNK = typename MMA_Atom::Shape_MNK;
|
||||
|
||||
using AtomThrID = typename MMA_Atom::ThrID;
|
||||
using AtomLayoutC_TV = typename MMA_Atom::LayoutC_TV;
|
||||
using AtomLayoutA_TV = typename MMA_Atom::LayoutA_TV;
|
||||
using AtomLayoutB_TV = typename MMA_Atom::LayoutB_TV;
|
||||
|
||||
// ThrV -> thread_idx
|
||||
using AtomThrID = typename MMA_Atom::ThrID;
|
||||
static_assert( rank_v<AtomLayoutMNK> == 3, "TiledMMA requires rank-3 AtomLayoutMNK");
|
||||
static_assert( rank_v<PermutationMNK> == 3, "TiledMMA requires rank-3 PermutationMNK");
|
||||
static_assert( is_tile<PermutationMNK>::value, "TiledMMA requires independent permutations of MNK.");
|
||||
static_assert(is_static<PermutationMNK>::value, "TiledMMA requires static permutations of MNK.");
|
||||
|
||||
// (M,N,K)
|
||||
using TiledShape_MNK = decltype(make_shape(size<0>(AtomShape_MNK{})*size<0>(AtomLayoutMNK{})*size<0>(ValLayoutMNK{}),
|
||||
size<1>(AtomShape_MNK{})*size<1>(AtomLayoutMNK{})*size<1>(ValLayoutMNK{}),
|
||||
size<2>(AtomShape_MNK{})*size<2>(AtomLayoutMNK{})*size<2>(ValLayoutMNK{})));
|
||||
|
||||
// thrid = (ThrV,ThrM,ThrN,ThrK) -> thr_idx
|
||||
using ThrLayoutVMNK = decltype(tiled_product(AtomThrID{}, AtomLayoutMNK{}));
|
||||
ThrLayoutVMNK thr_layout_vmnk_;
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
using TidLayout = decltype(right_inverse(ThrLayoutVMNK{}));
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
TiledMMA(MMA_Atom const& mma_atom = {}, AtomLayoutMNK const& thr_layout_mnk = {})
|
||||
: MMA_Atom(mma_atom),
|
||||
thr_layout_vmnk_(tiled_product(AtomThrID{}, thr_layout_mnk)) {}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr auto
|
||||
get_thr_layout_vmnk() const {
|
||||
return ThrLayoutVMNK{};
|
||||
return thr_layout_vmnk_;
|
||||
}
|
||||
|
||||
// Tile a tensor or a layout from shape
|
||||
@@ -243,17 +241,17 @@ struct TiledMMA : MMA_Atom
|
||||
// RestM: The values tiled in M.
|
||||
// RestN: The values tiled in N.
|
||||
template <class CTensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_C(CTensor&& ctensor)
|
||||
thrfrg_C(CTensor&& ctensor) const
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(ctensor) >= Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(size<0>(ctensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
|
||||
CUTE_STATIC_ASSERT_V(size<1>(ctensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(ctensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(ctensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(left_inverse(get<0>(PermutationsMNK{})),
|
||||
left_inverse(get<1>(PermutationsMNK{})));
|
||||
auto t_tile = make_tile(get<0>(PermutationMNK{}),
|
||||
get<1>(PermutationMNK{}));
|
||||
auto t_tensor = logical_divide(ctensor, t_tile); // (PermM,PermN)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
@@ -266,25 +264,13 @@ struct TiledMMA : MMA_Atom
|
||||
|
||||
// Tile the tensor for the C-threads
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<1>(ThrLayoutVMNK{})),
|
||||
make_layout(size<2>(ThrLayoutVMNK{}))));
|
||||
make_tile(make_layout(size<1>(thr_layout_vmnk_)),
|
||||
make_layout(size<2>(thr_layout_vmnk_))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrM,ThrN)),(FrgV,(RestM,RestN)))
|
||||
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
// Tile from (M,N,...)
|
||||
// to (thr_idx,(FrgV,(RestM,RestN,...)))
|
||||
template <class CTensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
tidfrg_C(CTensor&& ctensor)
|
||||
{
|
||||
// Don't need a ctile composition because ThrK is last mode in TidLayout
|
||||
|
||||
return thrfrg_C(ctensor).compose(TidLayout{}, _);
|
||||
}
|
||||
|
||||
// Tile a tensor or a layout from shape
|
||||
// (M,K,...)
|
||||
// to shape
|
||||
@@ -297,17 +283,17 @@ struct TiledMMA : MMA_Atom
|
||||
// RestM: The values tiled in M.
|
||||
// RestK: The values tiled in K.
|
||||
template <class ATensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_A(ATensor&& atensor)
|
||||
thrfrg_A(ATensor&& atensor) const
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(atensor) >= Int<2>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(atensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
|
||||
//UTE_STATIC_ASSERT_V(size<1>(atensor) % size<2>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(atensor) % size<2>(TiledShape_MNK{}) == Int<0>{});
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(left_inverse(get<0>(PermutationsMNK{})),
|
||||
left_inverse(get<2>(PermutationsMNK{})));
|
||||
auto t_tile = make_tile(get<0>(PermutationMNK{}),
|
||||
get<2>(PermutationMNK{}));
|
||||
auto t_tensor = logical_divide(atensor, t_tile); // (PermM,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
@@ -320,29 +306,13 @@ struct TiledMMA : MMA_Atom
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<1>(ThrLayoutVMNK{})),
|
||||
make_layout(size<3>(ThrLayoutVMNK{}))));
|
||||
make_tile(make_layout(size<1>(thr_layout_vmnk_)),
|
||||
make_layout(size<3>(thr_layout_vmnk_))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrM,ThrK)),(FrgV,(RestM,RestK)))
|
||||
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
// Tile from (M,K,...)
|
||||
// to (thr_idx,(FrgV,(RestM,RestK,...)))
|
||||
template <class ATensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
tidfrg_A(ATensor&& atensor)
|
||||
{
|
||||
auto atile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(ThrLayoutVMNK{}), size<2>(ThrLayoutVMNK{})),
|
||||
make_stride( Int<1>{} , Int<0>{} )),
|
||||
_));
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
|
||||
return thrfrg_A(atensor).compose(atile, _).compose(TidLayout{}, _);
|
||||
}
|
||||
|
||||
// Tile a tensor or a layout from shape
|
||||
// (N,K,...)
|
||||
// to shape
|
||||
@@ -355,17 +325,17 @@ struct TiledMMA : MMA_Atom
|
||||
// RestN: The values tiled in N.
|
||||
// RestK: The values tiled in K.
|
||||
template <class BTensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_B(BTensor&& btensor)
|
||||
thrfrg_B(BTensor&& btensor) const
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(btensor) >= Int<2>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(btensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(btensor) % size<2>(TiledShape_MNK{}) == Int<0>{});
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(left_inverse(get<1>(PermutationsMNK{})),
|
||||
left_inverse(get<2>(PermutationsMNK{})));
|
||||
auto t_tile = make_tile(get<1>(PermutationMNK{}),
|
||||
get<2>(PermutationMNK{}));
|
||||
auto t_tensor = logical_divide(btensor, t_tile); // (PermN,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
@@ -378,44 +348,28 @@ struct TiledMMA : MMA_Atom
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<2>(ThrLayoutVMNK{})),
|
||||
make_layout(size<3>(ThrLayoutVMNK{}))));
|
||||
make_tile(make_layout(size<2>(thr_layout_vmnk_)),
|
||||
make_layout(size<3>(thr_layout_vmnk_))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrN,ThrK)),(FrgV,(RestN,RestK)))
|
||||
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
// Tile from (N,K,...)
|
||||
// to (thr_idx,(FrgV,(RestN,RestK,...)))
|
||||
template <class BTensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
template <class ThrIdx,
|
||||
__CUTE_REQUIRES(is_integral<ThrIdx>::value)>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
tidfrg_B(BTensor&& btensor)
|
||||
get_slice(ThrIdx const& thr_idx) const
|
||||
{
|
||||
auto btile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(ThrLayoutVMNK{}), size<2>(ThrLayoutVMNK{})),
|
||||
make_stride( Int<0>{} , Int<1>{} )),
|
||||
_));
|
||||
// (ThrV,(ThrN,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
|
||||
return thrfrg_B(btensor).compose(btile, _).compose(TidLayout{}, _);
|
||||
auto thr_vmnk = thr_layout_vmnk_.get_flat_coord(thr_idx);
|
||||
return ThrMMA<TiledMMA, decltype(thr_vmnk)>{*this, thr_vmnk};
|
||||
}
|
||||
|
||||
template <class ThrIdx,
|
||||
__CUTE_REQUIRES(is_integral<ThrIdx>::value)>
|
||||
CUTE_HOST_DEVICE static constexpr
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_slice(ThrIdx const& thr_idx)
|
||||
{
|
||||
auto thr_vmnk = ThrLayoutVMNK{}.get_flat_coord(thr_idx);
|
||||
return ThrMMA<TiledMMA, decltype(thr_vmnk)>(thr_vmnk);
|
||||
}
|
||||
|
||||
template <class ThrIdx,
|
||||
__CUTE_REQUIRES(is_integral<ThrIdx>::value)>
|
||||
CUTE_HOST_DEVICE static constexpr
|
||||
auto
|
||||
get_thread_slice(ThrIdx const& thr_idx)
|
||||
get_thread_slice(ThrIdx const& thr_idx) const
|
||||
{
|
||||
return get_slice(thr_idx);
|
||||
}
|
||||
@@ -424,104 +378,144 @@ struct TiledMMA : MMA_Atom
|
||||
// Utility for printing and visualization
|
||||
//
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
// The size of the MNK-mode
|
||||
template <int I>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutC_MN()
|
||||
tile_size_mnk() const {
|
||||
static_assert(0 <= I && I < 3);
|
||||
auto core_size = size<I>(AtomShape_MNK{}) * size<I+1>(get_thr_layout_vmnk());
|
||||
[[maybe_unused]] auto perm_size = size<I>(PermutationMNK{});
|
||||
if constexpr (is_underscore<decltype(perm_size)>::value) {
|
||||
return core_size;
|
||||
} else {
|
||||
return cute::max(core_size, perm_size);
|
||||
}
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutC_MN() const
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_C = make_layout(make_shape(size<0>(TiledShape_MNK{}), size<1>(TiledShape_MNK{})));
|
||||
auto ref_C = make_layout(make_shape(tile_size_mnk<0>(), tile_size_mnk<1>()));
|
||||
// (cthrid,val) -> (M,N)
|
||||
auto layoutC_TV = thrfrg_C(ref_C);
|
||||
// (M,N) -> (cthrid,frg)
|
||||
auto layoutC_MN = right_inverse(layoutC_TV).with_shape(shape(ref_C));
|
||||
|
||||
// cthrid = (v,m,n) -> thr_idx
|
||||
auto thrID_C = ThrLayoutVMNK{}(_,_,_,Int<0>{});
|
||||
auto thrID_C = thr_layout_vmnk_(_,_,_,Int<0>{});
|
||||
|
||||
return cute::make_tuple(layoutC_MN, thrID_C);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutC_TV()
|
||||
get_layoutC_TV() const
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_C = make_layout(make_shape(size<0>(TiledShape_MNK{}), size<1>(TiledShape_MNK{})));
|
||||
auto ref_C = make_layout(make_shape(tile_size_mnk<0>(), tile_size_mnk<1>()));
|
||||
// (cthrid,val) -> (M,N)
|
||||
auto layoutC_TV = thrfrg_C(ref_C);
|
||||
|
||||
return tidfrg_C(ref_C);
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk_);
|
||||
|
||||
// (thr_idx,val) -> (M,N)
|
||||
return layoutC_TV.compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutA_MK()
|
||||
get_layoutA_MK() const
|
||||
{
|
||||
// (M,K) -> (M,K)
|
||||
auto ref_A = make_layout(make_shape(size<0>(TiledShape_MNK{}), size<2>(TiledShape_MNK{})));
|
||||
auto ref_A = make_layout(make_shape(tile_size_mnk<0>(), tile_size_mnk<2>()));
|
||||
// (athrid,val) -> (M,K)
|
||||
auto layoutA_TV = thrfrg_A(ref_A);
|
||||
// (M,K) -> (athrid,frg)
|
||||
auto layoutA_MK = right_inverse(layoutA_TV).with_shape(shape(ref_A));
|
||||
|
||||
// athrid = (v,m,k) -> thr_idx
|
||||
auto thrID_A = ThrLayoutVMNK{}(_,_,Int<0>{},_);
|
||||
auto thrID_A = thr_layout_vmnk_(_,_,Int<0>{},_);
|
||||
|
||||
return cute::make_tuple(layoutA_MK, thrID_A);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutA_TV()
|
||||
get_layoutA_TV() const
|
||||
{
|
||||
// (M,K) -> (M,K)
|
||||
auto ref_A = make_layout(make_shape(size<0>(TiledShape_MNK{}), size<2>(TiledShape_MNK{})));
|
||||
auto ref_A = make_layout(make_shape(tile_size_mnk<0>(), tile_size_mnk<2>()));
|
||||
// (athrid,val) -> (M,K)
|
||||
auto layoutA_TV = thrfrg_A(ref_A);
|
||||
|
||||
return tidfrg_A(ref_A);
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto atile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk_), size<2>(thr_layout_vmnk_)),
|
||||
make_stride( Int<1>{} , Int<0>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk_);
|
||||
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_A(ref_A).compose(atile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutB_NK()
|
||||
get_layoutB_NK() const
|
||||
{
|
||||
// (N,K) -> (N,K)
|
||||
auto ref_B = make_layout(make_shape(size<1>(TiledShape_MNK{}), size<2>(TiledShape_MNK{})));
|
||||
auto ref_B = make_layout(make_shape(tile_size_mnk<1>(), tile_size_mnk<2>()));
|
||||
// (bthrid,val) -> (N,K)
|
||||
auto layoutB_TV = thrfrg_B(ref_B);
|
||||
// (N,K) -> (bthrid,frg)
|
||||
auto layoutB_NK = right_inverse(layoutB_TV).with_shape(shape(ref_B));
|
||||
|
||||
// bthrid = (v,n,k) -> thr_idx
|
||||
auto thrID_B = ThrLayoutVMNK{}(_,Int<0>{},_,_);
|
||||
auto thrID_B = thr_layout_vmnk_(_,Int<0>{},_,_);
|
||||
|
||||
return cute::make_tuple(layoutB_NK, thrID_B);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutB_TV()
|
||||
get_layoutB_TV() const
|
||||
{
|
||||
// (N,K) -> (N,K)
|
||||
auto ref_B = make_layout(make_shape(size<1>(TiledShape_MNK{}), size<2>(TiledShape_MNK{})));
|
||||
auto ref_B = make_layout(make_shape(tile_size_mnk<1>(), tile_size_mnk<2>()));
|
||||
// (bthrid,val) -> (N,K)
|
||||
auto layoutB_TV = thrfrg_B(ref_B);
|
||||
|
||||
return tidfrg_B(ref_B);
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto btile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk_), size<2>(thr_layout_vmnk_)),
|
||||
make_stride( Int<0>{} , Int<1>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk_);
|
||||
|
||||
// (thr_idx,val) -> (N,K)
|
||||
return thrfrg_B(ref_B).compose(btile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
};
|
||||
|
||||
template <class TiledMMA, class ThrVMNK>
|
||||
struct ThrMMA : TiledMMA
|
||||
{
|
||||
// Use ThrVMNK and thrfrg rather than thr_idx and tidfrg
|
||||
// to support swizzled threads partitioning dynamic layouts
|
||||
ThrVMNK thr_vmnk_;
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ThrMMA(ThrVMNK const& thr_vmnk) : thr_vmnk_(thr_vmnk) {}
|
||||
|
||||
template <class CTensor>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_C(CTensor&& ctensor) const
|
||||
{
|
||||
auto thr_tensor = make_tensor(std::forward<CTensor>(ctensor).data(), TiledMMA::thrfrg_C(ctensor.layout()));
|
||||
auto thr_tensor = make_tensor(std::forward<CTensor>(ctensor).data(), this->thrfrg_C(ctensor.layout()));
|
||||
|
||||
auto thr_vmn = make_coord(get<0>(thr_vmnk_), make_coord(get<1>(thr_vmnk_), get<2>(thr_vmnk_)));
|
||||
return thr_tensor(thr_vmn, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
@@ -532,7 +526,7 @@ struct ThrMMA : TiledMMA
|
||||
auto
|
||||
partition_A(ATensor&& atensor) const
|
||||
{
|
||||
auto thr_tensor = make_tensor(std::forward<ATensor>(atensor).data(), TiledMMA::thrfrg_A(atensor.layout()));
|
||||
auto thr_tensor = make_tensor(std::forward<ATensor>(atensor).data(), this->thrfrg_A(atensor.layout()));
|
||||
|
||||
auto thr_vmk = make_coord(get<0>(thr_vmnk_), make_coord(get<1>(thr_vmnk_), get<3>(thr_vmnk_)));
|
||||
return thr_tensor(thr_vmk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
@@ -543,7 +537,7 @@ struct ThrMMA : TiledMMA
|
||||
auto
|
||||
partition_B(BTensor&& btensor) const
|
||||
{
|
||||
auto thr_tensor = make_tensor(std::forward<BTensor>(btensor).data(), TiledMMA::thrfrg_B(btensor.layout()));
|
||||
auto thr_tensor = make_tensor(std::forward<BTensor>(btensor).data(), this->thrfrg_B(btensor.layout()));
|
||||
|
||||
auto thr_vnk = make_coord(get<0>(thr_vmnk_), make_coord(get<2>(thr_vmnk_), get<3>(thr_vmnk_)));
|
||||
return thr_tensor(thr_vnk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
@@ -580,38 +574,32 @@ struct ThrMMA : TiledMMA
|
||||
|
||||
template <class MMA_Op,
|
||||
class MMAThrLayout = Layout<Shape<_1,_1,_1>>,
|
||||
class MMAValLayout = Layout<Shape<_1,_1,_1>>,
|
||||
class Permutations = Tile<Underscore,Underscore,Underscore>>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_tiled_mma(MMA_Atom<MMA_Op> const&,
|
||||
make_tiled_mma(MMA_Atom<MMA_Op> const& mma_atom,
|
||||
MMAThrLayout const& thr_layout = {},
|
||||
MMAValLayout const& val_layout = {},
|
||||
Permutations const& permutations = {})
|
||||
{
|
||||
auto thr_layout_mnk = append<3>(thr_layout, Layout<_1,_0>{});
|
||||
auto val_layout_mnk = append<3>(val_layout, Layout<_1,_0>{});
|
||||
auto permutation_mnk = append<3>(permutations, _);
|
||||
|
||||
return TiledMMA<MMA_Atom<MMA_Op>,
|
||||
decltype(thr_layout_mnk),
|
||||
decltype(val_layout_mnk),
|
||||
decltype(permutation_mnk)>{};
|
||||
decltype(permutation_mnk)>{mma_atom, thr_layout_mnk};
|
||||
}
|
||||
|
||||
template <class MMA_Op,
|
||||
class MMAThrLayout = Layout<Shape<_1,_1,_1>>,
|
||||
class MMAValLayout = Layout<Shape<_1,_1,_1>>,
|
||||
class Permutations = Tile<Underscore,Underscore,Underscore>>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_tiled_mma(MMA_Op const&,
|
||||
MMAThrLayout const& thr_layout = {},
|
||||
MMAValLayout const& val_layout = {},
|
||||
Permutations const& permutations = {})
|
||||
{
|
||||
// Attempt to wrap in an MMA_Atom<> and forward
|
||||
return make_tiled_mma(MMA_Atom<MMA_Op>{}, thr_layout, val_layout, permutations);
|
||||
return make_tiled_mma(MMA_Atom<MMA_Op>{}, thr_layout, permutations);
|
||||
}
|
||||
|
||||
//
|
||||
@@ -680,28 +668,38 @@ partition_shape_B(TiledMMA<Args...> const& mma, Shape_NK const& shape_NK)
|
||||
// Size
|
||||
//
|
||||
|
||||
template <int... I, class... Args>
|
||||
template <int I, class... Args>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
tile_size(TiledMMA<Args...> const& mma)
|
||||
{
|
||||
return size<I...>(typename TiledMMA<Args...>::TiledShape_MNK{});
|
||||
return mma.template tile_size_mnk<I>();
|
||||
}
|
||||
|
||||
template <int... I, class... Args>
|
||||
template <class... Args>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
tile_shape(TiledMMA<Args...> const& mma)
|
||||
{
|
||||
return shape<I...>(typename TiledMMA<Args...>::TiledShape_MNK{});
|
||||
return make_shape(tile_size<0>(mma), tile_size<1>(mma), tile_size<2>(mma));
|
||||
}
|
||||
|
||||
// Deprecate?
|
||||
template <int... I, class... Args>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
size(TiledMMA<Args...> const& mma)
|
||||
{
|
||||
return size<I...>(typename TiledMMA<Args...>::ThrLayoutVMNK{});
|
||||
return size<I...>(mma.get_thr_layout_vmnk());
|
||||
}
|
||||
|
||||
// Alias
|
||||
template <int... I, class... Args>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thr_size(TiledMMA<Args...> const& mma)
|
||||
{
|
||||
return size<I...>(mma.get_thr_layout_vmnk());
|
||||
}
|
||||
|
||||
//
|
||||
@@ -715,33 +713,31 @@ print(MMA_Atom<MMA_Traits<Args...>> const&)
|
||||
{
|
||||
using Atom = MMA_Atom<MMA_Traits<Args...>>;
|
||||
print("MMA_Atom\n");
|
||||
print(" ThrID: "); print(typename Atom::ThrID{}); print("\n");
|
||||
print(" LayoutA_TV: "); print(typename Atom::LayoutA_TV{}); print("\n");
|
||||
print(" LayoutB_TV: "); print(typename Atom::LayoutB_TV{}); print("\n");
|
||||
print(" LayoutC_TV: "); print(typename Atom::LayoutC_TV{}); print("\n");
|
||||
print(" ThrID: "); print(typename Atom::ThrID{}); print("\n");
|
||||
print(" LayoutA_TV: "); print(typename Atom::LayoutA_TV{}); print("\n");
|
||||
print(" LayoutB_TV: "); print(typename Atom::LayoutB_TV{}); print("\n");
|
||||
print(" LayoutC_TV: "); print(typename Atom::LayoutC_TV{}); print("\n");
|
||||
}
|
||||
|
||||
template <class Atom, class TiledThr, class TiledVal, class TiledPerm>
|
||||
template <class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print(TiledMMA<Atom, TiledThr, TiledVal, TiledPerm> const& mma)
|
||||
print(TiledMMA<Atom, TiledThr, TiledPerm> const& mma)
|
||||
{
|
||||
using MMA = TiledMMA<Atom, TiledThr, TiledVal, TiledPerm>;
|
||||
print("TiledMMA\n");
|
||||
print(" TiledThr: "); print(TiledThr{}); print("\n");
|
||||
print(" TiledVal: "); print(TiledVal{}); print("\n");
|
||||
print(" TiledPerm: "); print(TiledPerm{}); print("\n");
|
||||
print(" TiledShape_MNK: "); print(typename MMA::TiledShape_MNK{}); print("\n");
|
||||
print(" ThrLayoutVMNK: "); print(typename MMA::ThrLayoutVMNK{}); print("\n");
|
||||
print(" ThrLayoutVMNK: "); print(mma.get_thr_layout_vmnk()); print("\n");
|
||||
print(" PermutationMNK: "); print(TiledPerm{}); print("\n");
|
||||
print(static_cast<Atom const&>(mma));
|
||||
}
|
||||
|
||||
template <class TiledMMA, class ThrVMNK>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print(ThrMMA<TiledMMA, ThrVMNK> const&)
|
||||
print(ThrMMA<TiledMMA, ThrVMNK> const& thr_mma)
|
||||
{
|
||||
print(TiledMMA{});
|
||||
print("ThrMMA\n");
|
||||
print(" Thr VMNK: "); print(thr_mma.thr_vmnk_); print("\n");
|
||||
print(static_cast<TiledMMA>(thr_mma));
|
||||
}
|
||||
|
||||
template <class... Args>
|
||||
@@ -766,18 +762,6 @@ print_latex(TiledMMA<Args...> const& mma)
|
||||
layoutB_NK, thrID_B);
|
||||
}
|
||||
|
||||
// EXPERIMENTAL -- Doesn't work with Swizzled Thr TileMMAs...
|
||||
template <class... Args>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
print_latex_2(TiledMMA<Args...> const& mma)
|
||||
{
|
||||
print_latex_mma(typename TiledMMA<Args...>::TiledShape_MNK{},
|
||||
mma.get_layoutC_TV(),
|
||||
mma.get_layoutA_TV(),
|
||||
mma.get_layoutB_TV());
|
||||
}
|
||||
|
||||
// MNK MMA Layout to console printer -- 8-value color coded by thread
|
||||
template <class LayoutC, class ThrIDC,
|
||||
class LayoutA, class ThrIDA,
|
||||
@@ -943,122 +927,6 @@ print_latex_mma(LayoutC const& C, ThrIDC const& TC, // (m,n) -> (tid,vid) and
|
||||
printf(latex_footer);
|
||||
}
|
||||
|
||||
// ThrVal MMA Layout to Latex TIKZ -- 8-value color coded by thread
|
||||
template <class Shape_MNK,
|
||||
class LayoutC, class LayoutA, class LayoutB>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print_latex_mma(Shape_MNK const& shape_mnk,
|
||||
LayoutC const& C, // (thr_idx,vid) -> (m,n)
|
||||
LayoutA const& A, // (thr_idx,vid) -> (m,k)
|
||||
LayoutB const& B) // (thr_idx,vid) -> (n,k)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(C) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(A) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(B) == Int<2>{});
|
||||
|
||||
char const* latex_header =
|
||||
"\\documentclass{standalone}\n"
|
||||
"\\usepackage{tikz}\n"
|
||||
"\\usetikzlibrary{external}\n"
|
||||
"\\tikzexternalize\n"
|
||||
"\\begin{document}\n"
|
||||
"\\begin{tikzpicture}[x={(0cm,-1cm)},y={(1cm,0cm)},box/.style={rectangle,draw=black,thick,minimum size=1cm,anchor=center}]\n\n";
|
||||
char const* latex_footer =
|
||||
"\\end{tikzpicture}\n"
|
||||
"\\end{document}\n";
|
||||
|
||||
char const* color_map[8] = {"{rgb,255:red,175;green,175;blue,255}",
|
||||
"{rgb,255:red,175;green,255;blue,175}",
|
||||
"{rgb,255:red,255;green,255;blue,175}",
|
||||
"{rgb,255:red,255;green,175;blue,175}",
|
||||
"{rgb,255:red,210;green,210;blue,255}",
|
||||
"{rgb,255:red,210;green,255;blue,210}",
|
||||
"{rgb,255:red,255;green,255;blue,210}",
|
||||
"{rgb,255:red,255;green,210;blue,210}"};
|
||||
|
||||
// Header
|
||||
printf("%% Shape_MNK: "); print(shape_mnk); printf("\n");
|
||||
printf("%% LayoutC : "); print(C); printf("\n");
|
||||
printf("%% LayoutA : "); print(A); printf("\n");
|
||||
printf("%% LayoutB : "); print(B); printf("\n\n");
|
||||
|
||||
printf(latex_header);
|
||||
|
||||
auto M = size<0>(shape_mnk);
|
||||
auto N = size<1>(shape_mnk);
|
||||
auto K = size<2>(shape_mnk);
|
||||
|
||||
// C starting at 0,0
|
||||
bool c_filled[M][N] = {};
|
||||
for (int t = 0; t < size<0>(C); ++t) {
|
||||
for (int v = 0; v < size<1>(C); ++v) {
|
||||
int m = C(t,v) % M;
|
||||
int n = C(t,v) / M;
|
||||
|
||||
if (not c_filled[m][n]) {
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[t % 8],
|
||||
m, n,
|
||||
t, v);
|
||||
c_filled[m][n] = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A starting at 0,-size<1>(A)-1
|
||||
bool a_filled[M][K] = {};
|
||||
for (int t = 0; t < size<0>(A); ++t) {
|
||||
for (int v = 0; v < size<1>(A); ++v) {
|
||||
int m = A(t,v) % M;
|
||||
int k = A(t,v) / M;
|
||||
|
||||
if (not a_filled[m][k]) {
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[t % 8],
|
||||
m, k - 1 - K,
|
||||
t, v);
|
||||
a_filled[m][k] = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// B starting at -size<1>(B)-1,0
|
||||
bool b_filled[N][K] = {};
|
||||
for (int t = 0; t < size<0>(B); ++t) {
|
||||
for (int v = 0; v < size<1>(B); ++v) {
|
||||
int n = B(t,v) % N;
|
||||
int k = B(t,v) / N;
|
||||
|
||||
if (not b_filled[n][k]) {
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[t % 8],
|
||||
k - 1 - K, n,
|
||||
t, v);
|
||||
b_filled[n][k] = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A labels
|
||||
for (int m = 0, k = -1; m < M; ++m) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", m, k - 1 - K, m);
|
||||
}
|
||||
for (int k = 0, m = -1; k < K; ++k) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", m, k - 1 - K, k);
|
||||
}
|
||||
// B labels
|
||||
for (int n = 0, k = -1; n < N; ++n) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", k - 1 - K, n, n);
|
||||
}
|
||||
for (int k = 0, n = -1; k < K; ++k) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", k - 1 - K, n, k);
|
||||
}
|
||||
|
||||
// Footer
|
||||
printf(latex_footer);
|
||||
}
|
||||
|
||||
} // namespace cute
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -47,7 +47,7 @@
|
||||
#endif
|
||||
|
||||
#if !defined(__CUDACC_RTC__) && !defined(__clang__) && \
|
||||
(defined(__CUDA_ARCH__) || defined(_NVHPC_CUDA))
|
||||
(defined(__CUDA_ARCH__) || defined(_NVHPC_CUDA))
|
||||
# define CUTE_UNROLL #pragma unroll
|
||||
# define CUTE_NO_UNROLL #pragma unroll 1
|
||||
#elif defined(__CUDACC_RTC__) || defined(__clang__)
|
||||
@@ -120,6 +120,8 @@
|
||||
#include <cassert>
|
||||
#endif
|
||||
|
||||
#define CUTE_STATIC_V(x) decltype(x)::value
|
||||
|
||||
#define CUTE_STATIC_ASSERT static_assert
|
||||
#define CUTE_STATIC_ASSERT_V(x,...) static_assert(decltype(x)::value, ##__VA_ARGS__)
|
||||
|
||||
|
||||
@@ -91,7 +91,7 @@ private:
|
||||
// Flag for fast branching on straddled elements
|
||||
static constexpr bool is_storage_unaligned = ((sizeof_bits_v<storage_type> % sizeof_bits_v<element_type>) != 0);
|
||||
|
||||
friend class subbyte_iterator<T>;
|
||||
friend struct subbyte_iterator<T>;
|
||||
|
||||
// Pointer to storage element
|
||||
storage_type* ptr_ = nullptr;
|
||||
@@ -208,7 +208,7 @@ struct subbyte_iterator
|
||||
|
||||
private:
|
||||
|
||||
template <class, class> friend class swizzle_ptr;
|
||||
template <class, class> friend struct swizzle_ptr;
|
||||
|
||||
// Pointer to storage element
|
||||
storage_type* ptr_ = nullptr;
|
||||
|
||||
@@ -327,7 +327,7 @@ ceil_div(IntTupleA const& a, IntTupleB const& b)
|
||||
{
|
||||
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
|
||||
static_assert(tuple_size<IntTupleA>::value >= tuple_size<IntTupleB>::value, "Mismatched ranks");
|
||||
constexpr int R = tuple_size<IntTupleA>::value; // Missing ranks in TupleB are implictly 1
|
||||
constexpr int R = tuple_size<IntTupleA>::value; // Missing ranks in TupleB are implicitly 1
|
||||
return transform(a, append<R>(b,Int<1>{}), [](auto const& x, auto const& y) { return ceil_div(x,y); });
|
||||
} else {
|
||||
return (a + b - Int<1>{}) / b;
|
||||
@@ -336,6 +336,28 @@ ceil_div(IntTupleA const& a, IntTupleB const& b)
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
//
|
||||
// round_up
|
||||
// Round @a a up to the nearest multiple of @a b.
|
||||
// For negative numbers, rounds away from zero.
|
||||
//
|
||||
|
||||
template <class IntTupleA, class IntTupleB>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
round_up(IntTupleA const& a, IntTupleB const& b)
|
||||
{
|
||||
if constexpr (is_tuple<IntTupleA>::value && is_tuple<IntTupleB>::value) {
|
||||
static_assert(tuple_size<IntTupleA>::value >= tuple_size<IntTupleB>::value, "Mismatched ranks");
|
||||
constexpr int R = tuple_size<IntTupleA>::value; // Missing ranks in TupleB are implicitly 1
|
||||
return transform(a, append<R>(b,Int<1>{}), [](auto const& x, auto const& y) { return round_up(x,y); });
|
||||
} else {
|
||||
return ((a + b - Int<1>{}) / b) * b;
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
/** Division for Shapes
|
||||
* Case Tuple Tuple:
|
||||
* Perform shape_div element-wise
|
||||
@@ -429,6 +451,7 @@ template <class A, class B>
|
||||
using is_congruent = decltype(congruent(declval<A>(), declval<B>()));
|
||||
|
||||
/** Test if two IntTuple have the similar profiles up to Shape A (hierarchical rank division)
|
||||
* weakly_congruent is a partial order on A and B: A <= B
|
||||
*/
|
||||
template <class IntTupleA, class IntTupleB>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -458,7 +481,7 @@ using is_weakly_congruent = decltype(weakly_congruent(declval<A>(), declval<B>()
|
||||
|
||||
/** Test if Shape B is compatible with Shape A:
|
||||
* Any coordinate into A can also be used as a coordinate into B
|
||||
* A <= B is a partially ordered set of factored shapes
|
||||
* compatible is a partial order on A and B: A <= B
|
||||
*/
|
||||
template <class IntTupleA, class IntTupleB>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -487,7 +510,8 @@ template <class A, class B>
|
||||
using is_compatible = decltype(compatible(declval<A>(), declval<B>()));
|
||||
|
||||
/** Test if Shape B is weakly compatible with Shape A:
|
||||
* Shape B divides Shape A at some level of refinement
|
||||
* Shape B is a multiple of a shape that is compatible with Shape A
|
||||
* weakly_compatible is a partial order on A and B: A <= B
|
||||
*/
|
||||
template <class IntTupleA, class IntTupleB>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -502,7 +526,7 @@ weakly_compatible(IntTupleA const& a, IntTupleB const& b)
|
||||
[](auto const&... z) { return (true_type{} && ... && z); });
|
||||
}
|
||||
} else if constexpr (is_integral<IntTupleA>::value) {
|
||||
return a % size(b) == Int<0>{};
|
||||
return size(b) % a == Int<0>{};
|
||||
} else if constexpr (is_integral<IntTupleB>::value) {
|
||||
return false_type{};
|
||||
} else {
|
||||
|
||||
+51
-57
@@ -981,7 +981,6 @@ auto
|
||||
composition(Layout<LShape,LStride> const& lhs,
|
||||
Layout<RShape,RStride> const& rhs)
|
||||
{
|
||||
//return detail::composition_impl(flatten(lhs), rhs.shape(), rhs.stride());
|
||||
return detail::composition_impl(lhs, rhs.shape(), rhs.stride());
|
||||
}
|
||||
|
||||
@@ -997,8 +996,8 @@ composition(Layout<LShape,LStride> const& lhs,
|
||||
return detail::transform_layout(lhs, rhs, [](auto const& l, auto const& r) { return composition(l,r); }, make_seq<tuple_size<IntTuple>::value>{}, seq<>{}, seq<>{});
|
||||
} else if constexpr (is_underscore<IntTuple>::value) {
|
||||
return lhs;
|
||||
} else {
|
||||
return composition(lhs, make_layout(rhs));
|
||||
} else if constexpr (is_integral<IntTuple>::value) {
|
||||
return detail::composition_impl(lhs, rhs, Int<1>{});
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
@@ -1097,15 +1096,18 @@ inverse_seq(Shape const& shape, Stride const& stride, seq<Is...>)
|
||||
auto next_I = cute::find_if(stride, [](auto a) { return is_constant<NextStride, decltype(a)>{}; });
|
||||
|
||||
if constexpr (next_I == decltype(rank(stride))::value) {
|
||||
// If not found, return current seq
|
||||
return seq<Is...>{};
|
||||
} else {
|
||||
// auto next_stride = get<next_I>(shape) * get<next_I>(stride);
|
||||
// NOTE: Needed for g++-7
|
||||
using next_stride = decltype(get<next_I>(shape) * get<next_I>(stride));
|
||||
|
||||
if constexpr (is_static<next_stride>::value) {
|
||||
if constexpr (is_static<next_stride>::value && !is_constant<NextStride, next_stride>::value) {
|
||||
// If next_stride is static and unique, then continue
|
||||
return inverse_seq<next_stride::value>(shape, stride, seq<Is..., next_I>{});
|
||||
} else {
|
||||
// Else return current seq + next_I
|
||||
return seq<Is..., next_I>{};
|
||||
}
|
||||
}
|
||||
@@ -1340,28 +1342,24 @@ template <class LShape, class LStride,
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
logical_divide(Layout<LShape,LStride> const& layout,
|
||||
Layout<TShape,TStride> const& tile)
|
||||
Layout<TShape,TStride> const& tiler)
|
||||
{
|
||||
//CUTE_STATIC_ASSERT_V(size(layout) % size(tile) == Int<0>{},
|
||||
// "Tiling does not evenly divide the block");
|
||||
// NOTE: With tiles that have stride-0, this doesn't have to be true
|
||||
|
||||
return composition(layout, make_layout(tile, complement(tile, size(layout))));
|
||||
return composition(layout, make_layout(tiler, complement(tiler, size(layout))));
|
||||
}
|
||||
|
||||
template <class LShape, class LStride, class IntTuple>
|
||||
template <class LShape, class LStride, class Tiler>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
logical_divide(Layout<LShape,LStride> const& layout,
|
||||
IntTuple const& tile)
|
||||
Tiler const& tiler)
|
||||
{
|
||||
if constexpr (is_tuple<IntTuple>::value) {
|
||||
static_assert(tuple_size<IntTuple>::value <= Layout<LShape,LStride>::rank, "logical_divide: Too many modes in tile.");
|
||||
return transform_layout(layout, tile, [](auto const& l, auto const& t) { return logical_divide(l,t); });
|
||||
} else if constexpr (is_underscore<IntTuple>::value) {
|
||||
if constexpr (is_tuple<Tiler>::value) {
|
||||
static_assert(tuple_size<Tiler>::value <= Layout<LShape,LStride>::rank, "logical_divide: Too many modes in tiler.");
|
||||
return transform_layout(layout, tiler, [](auto const& l, auto const& t) { return logical_divide(l,t); });
|
||||
} else if constexpr (is_underscore<Tiler>::value) {
|
||||
return layout;
|
||||
} else if constexpr (is_integral<IntTuple>::value) {
|
||||
return logical_divide(layout, make_layout(tile));
|
||||
} else if constexpr (is_integral<Tiler>::value) {
|
||||
return logical_divide(layout, make_layout(tiler));
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
@@ -1374,24 +1372,24 @@ logical_divide(Layout<LShape,LStride> const& layout,
|
||||
//
|
||||
|
||||
template <class LShape, class LStride,
|
||||
class Tile>
|
||||
class Tiler>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
zipped_divide(Layout<LShape,LStride> const& layout,
|
||||
Tile const& tile)
|
||||
Tiler const& tiler)
|
||||
{
|
||||
return tile_unzip(logical_divide(layout, tile), tile);
|
||||
return tile_unzip(logical_divide(layout, tiler), tiler);
|
||||
}
|
||||
|
||||
// Same as zipped_divide, but unpacks the second mode: ((BLK_A,BLK_B,...),a,b,...,x,y)
|
||||
template <class LShape, class LStride,
|
||||
class Tile>
|
||||
class Tiler>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
tiled_divide(Layout<LShape,LStride> const& layout,
|
||||
Tile const& tile)
|
||||
Tiler const& tiler)
|
||||
{
|
||||
auto div = zipped_divide(layout, tile);
|
||||
auto div = zipped_divide(layout, tiler);
|
||||
|
||||
auto R = rank<1>(div);
|
||||
return div(_, repeat<R>(_));
|
||||
@@ -1399,13 +1397,13 @@ tiled_divide(Layout<LShape,LStride> const& layout,
|
||||
|
||||
// Same as zipped_divide, but unpacks both modes: (BLK_A,BLK_B,...,a,b,...,x,y)
|
||||
template <class LShape, class LStride,
|
||||
class Tile>
|
||||
class Tiler>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
flat_divide(Layout<LShape,LStride> const& layout,
|
||||
Tile const& tile)
|
||||
Tiler const& tiler)
|
||||
{
|
||||
auto div = zipped_divide(layout, tile);
|
||||
auto div = zipped_divide(layout, tiler);
|
||||
|
||||
auto R0 = rank<0>(div);
|
||||
auto R1 = rank<1>(div);
|
||||
@@ -1421,24 +1419,24 @@ template <class LShape, class LStride,
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
logical_product(Layout<LShape,LStride> const& layout,
|
||||
Layout<TShape,TStride> const& tile)
|
||||
Layout<TShape,TStride> const& tiler)
|
||||
{
|
||||
return make_layout(layout, composition(complement(layout, size(layout)*cosize(tile)), tile));
|
||||
return make_layout(layout, composition(complement(layout, size(layout)*cosize(tiler)), tiler));
|
||||
}
|
||||
|
||||
template <class LShape, class LStride, class IntTuple>
|
||||
template <class LShape, class LStride, class Tiler>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
logical_product(Layout<LShape,LStride> const& layout,
|
||||
IntTuple const& tile)
|
||||
Tiler const& tiler)
|
||||
{
|
||||
if constexpr (is_tuple<IntTuple>::value) {
|
||||
static_assert(tuple_size<IntTuple>::value <= Layout<LShape,LStride>::rank);
|
||||
return transform_layout(layout, tile, [](auto const& l, auto const& t) { return logical_product(l,t); });
|
||||
} else if constexpr (is_underscore<IntTuple>::value) {
|
||||
if constexpr (is_tuple<Tiler>::value) {
|
||||
static_assert(tuple_size<Tiler>::value <= Layout<LShape,LStride>::rank, "logical_product: Too many modes in tiler.");
|
||||
return transform_layout(layout, tiler, [](auto const& l, auto const& t) { return logical_product(l,t); });
|
||||
} else if constexpr (is_underscore<Tiler>::value) {
|
||||
return layout;
|
||||
} else if constexpr (is_integral<IntTuple>::value) {
|
||||
return logical_product(layout, make_layout(tile));
|
||||
} else if constexpr (is_integral<Tiler>::value) {
|
||||
return logical_product(layout, make_layout(tiler));
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
@@ -1451,45 +1449,43 @@ logical_product(Layout<LShape,LStride> const& layout,
|
||||
//
|
||||
|
||||
template <class LShape, class LStride,
|
||||
class Tile>
|
||||
class Tiler>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
zipped_product(Layout<LShape,LStride> const& layout,
|
||||
Tile const& tile)
|
||||
Tiler const& tiler)
|
||||
{
|
||||
return tile_unzip(logical_product(layout, tile), tile);
|
||||
return tile_unzip(logical_product(layout, tiler), tiler);
|
||||
}
|
||||
|
||||
// Same as zipped_product, but unpacks the second mode: ((BLK_A,BLK_B,...),a,b,...,x,y)
|
||||
template <class LShape, class LStride,
|
||||
class Tile>
|
||||
class Tiler>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
tiled_product(Layout<LShape,LStride> const& layout,
|
||||
Tile const& tile)
|
||||
Tiler const& tiler)
|
||||
{
|
||||
auto div = zipped_product(layout, tile);
|
||||
auto div = zipped_product(layout, tiler);
|
||||
|
||||
auto R = rank(tile);
|
||||
auto R = rank<1>(div);
|
||||
return div(_, repeat<R>(_));
|
||||
}
|
||||
|
||||
// Attempts to reproduce layout "block" over layout "layout"
|
||||
// That is, think of every element of "layout" as a "block"
|
||||
// Attempts to reproduce a layout over a tiler
|
||||
// That is, think of every element of "tiler" as a "layout"
|
||||
// and return the layout of the resulting structure
|
||||
template <class TShape, class TStride,
|
||||
class UShape, class UStride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
blocked_product(Layout<TShape,TStride> const& block,
|
||||
Layout<UShape,UStride> const& layout)
|
||||
blocked_product(Layout<TShape,TStride> const& layout,
|
||||
Layout<UShape,UStride> const& tiler)
|
||||
{
|
||||
constexpr int R = cute::max(rank_v<TShape>, rank_v<UShape>);
|
||||
auto padded_block = append<R>(block);
|
||||
auto padded_layout = append<R>(layout);
|
||||
|
||||
auto result = logical_product(padded_block, padded_layout);
|
||||
|
||||
auto result = logical_product(append<R>(layout), append<R>(tiler));
|
||||
|
||||
return coalesce(zip(get<0>(result), get<1>(result)), repeat<R>(Int<1>{}));
|
||||
}
|
||||
|
||||
@@ -1497,14 +1493,12 @@ template <class TShape, class TStride,
|
||||
class UShape, class UStride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
raked_product(Layout<TShape,TStride> const& block,
|
||||
Layout<UShape,UStride> const& layout)
|
||||
raked_product(Layout<TShape,TStride> const& layout,
|
||||
Layout<UShape,UStride> const& tiler)
|
||||
{
|
||||
constexpr int R = cute::max(rank_v<TShape>, rank_v<UShape>);
|
||||
auto padded_block = append<R>(block);
|
||||
auto padded_layout = append<R>(layout);
|
||||
|
||||
auto result = logical_product(padded_block, padded_layout);
|
||||
auto result = logical_product(append<R>(layout), append<R>(tiler));
|
||||
|
||||
return coalesce(zip(get<1>(result), get<0>(result)), repeat<R>(Int<1>{}));
|
||||
}
|
||||
|
||||
@@ -473,6 +473,16 @@ zipped_divide(ComposedLayout<A,O,B> const& a,
|
||||
return composition(a.layout_a(), a.offset(), zipped_divide(a.layout_b(), b));
|
||||
}
|
||||
|
||||
template <class A, class O, class B,
|
||||
class Tile>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
flat_divide(ComposedLayout<A,O,B> const& a,
|
||||
Tile const& b)
|
||||
{
|
||||
return composition(a.layout_a(), a.offset(), flat_divide(a.layout_b(), b));
|
||||
}
|
||||
|
||||
template <class A, class O, class B,
|
||||
class Tile>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
|
||||
@@ -181,7 +181,7 @@ template <class T0, class T1, class... Ts>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_inttuple_iter(T0 const& t0, T1 const& t1, Ts const&... ts) {
|
||||
return make_tuple_iter(cute::make_tuple(t0, t1, ts...));
|
||||
return make_inttuple_iter(cute::make_tuple(t0, t1, ts...));
|
||||
}
|
||||
|
||||
//
|
||||
|
||||
@@ -148,7 +148,9 @@ using _96 = Int<96>;
|
||||
using _128 = Int<128>;
|
||||
using _192 = Int<192>;
|
||||
using _256 = Int<256>;
|
||||
using _384 = Int<384>;
|
||||
using _512 = Int<512>;
|
||||
using _768 = Int<768>;
|
||||
using _1024 = Int<1024>;
|
||||
using _2048 = Int<2048>;
|
||||
using _4096 = Int<4096>;
|
||||
|
||||
@@ -97,7 +97,7 @@ template <class T, class U,
|
||||
__CUTE_REQUIRES(is_std_integral<T>::value &&
|
||||
is_std_integral<U>::value)>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
cute::common_type_t<T, U>
|
||||
gcd(T t, U u) {
|
||||
while (true) {
|
||||
if (t == 0) { return u; }
|
||||
@@ -112,7 +112,7 @@ template <class T, class U,
|
||||
__CUTE_REQUIRES(is_std_integral<T>::value &&
|
||||
is_std_integral<U>::value)>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
cute::common_type_t<T, U>
|
||||
lcm(T const& t, U const& u) {
|
||||
return (t / gcd(t,u)) * u;
|
||||
}
|
||||
|
||||
@@ -233,14 +233,14 @@ CUTE_HOST_DEVICE void print(T const* const ptr)
|
||||
template <class T>
|
||||
CUTE_HOST_DEVICE void print(counting_iterator<T> ptr)
|
||||
{
|
||||
printf("counting_iter_"); print(ptr.n_);
|
||||
printf("counting_iter("); print(ptr.n_); printf(")");
|
||||
}
|
||||
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
template <class T>
|
||||
CUTE_HOST std::ostream& operator<<(std::ostream& os, counting_iterator<T> ptr)
|
||||
{
|
||||
return os << "counting_iter_" << ptr.n_;
|
||||
return os << "counting_iter(" << ptr.n_ << ")";
|
||||
}
|
||||
#endif // !defined(__CUDACC_RTC__)
|
||||
|
||||
|
||||
@@ -990,13 +990,10 @@ CUTE_HOST_DEVICE void print_tensor(Tensor<Engine,Layout> const& tensor)
|
||||
{
|
||||
print(tensor); print(":\n");
|
||||
|
||||
auto format = get_format(tensor(0));
|
||||
using type = typename decltype(format)::type;
|
||||
|
||||
if constexpr (Layout::rank == 1)
|
||||
{
|
||||
for (int m = 0; m < size(tensor); ++m) {
|
||||
printf(format.format, format.digits, type(tensor(m)));
|
||||
pretty_print(tensor(m));
|
||||
printf("\n");
|
||||
}
|
||||
} else
|
||||
@@ -1004,7 +1001,7 @@ CUTE_HOST_DEVICE void print_tensor(Tensor<Engine,Layout> const& tensor)
|
||||
{
|
||||
for (int m = 0; m < size<0>(tensor); ++m) {
|
||||
for (int n = 0; n < size<1>(tensor); ++n) {
|
||||
printf(format.format, format.digits, type(tensor(m,n)));
|
||||
pretty_print(tensor(m,n));
|
||||
}
|
||||
printf("\n");
|
||||
}
|
||||
@@ -1013,7 +1010,7 @@ CUTE_HOST_DEVICE void print_tensor(Tensor<Engine,Layout> const& tensor)
|
||||
{
|
||||
print_tensor(tensor(_,_,0));
|
||||
for (int k = 1; k < size<2>(tensor); ++k) {
|
||||
for (int i = 0; i < format.digits*size<1>(tensor); ++i) { print("-"); } print("\n");
|
||||
for (int i = 0; i < 5*size<1>(tensor); ++i) { print("-"); } print("\n");
|
||||
print_tensor(tensor(_,_,k));
|
||||
}
|
||||
} else
|
||||
@@ -1021,7 +1018,7 @@ CUTE_HOST_DEVICE void print_tensor(Tensor<Engine,Layout> const& tensor)
|
||||
{
|
||||
print_tensor(tensor(_,_,_,0));
|
||||
for (int p = 1; p < size<3>(tensor); ++p) {
|
||||
for (int i = 0; i < format.digits*size<1>(tensor); ++i) { print("="); } print("\n");
|
||||
for (int i = 0; i < 5*size<1>(tensor); ++i) { print("="); } print("\n");
|
||||
print_tensor(tensor(_,_,_,p));
|
||||
}
|
||||
}
|
||||
|
||||
+50
-55
@@ -57,61 +57,6 @@ num_digits(int x)
|
||||
10)))))))));
|
||||
}
|
||||
|
||||
template <class T>
|
||||
struct format_and_size {
|
||||
using type = T;
|
||||
char const* format;
|
||||
int digits;
|
||||
};
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<int>
|
||||
get_format(bool) {
|
||||
return {"%*d", 3};
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<int32_t>
|
||||
get_format(int32_t) {
|
||||
return {"%*d", 5};
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<uint32_t>
|
||||
get_format(uint32_t) {
|
||||
return {"%*d", 5};
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<int64_t>
|
||||
get_format(int64_t) {
|
||||
return {"%*d", 5};
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<uint64_t>
|
||||
get_format(uint64_t) {
|
||||
return {"%*d", 5};
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<float>
|
||||
get_format(half_t) {
|
||||
return {"%*.2f", 8};
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<float>
|
||||
get_format(float) {
|
||||
return {"%*.2e", 10};
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
format_and_size<double>
|
||||
get_format(double) {
|
||||
return {"%*.3e", 11};
|
||||
}
|
||||
|
||||
//
|
||||
// print dispatcher
|
||||
//
|
||||
@@ -195,4 +140,54 @@ print(char const* format) {
|
||||
printf("%s", format);
|
||||
}
|
||||
|
||||
//
|
||||
// pretty printing
|
||||
//
|
||||
|
||||
template <class T>
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(T const& v) {
|
||||
printf(" "); print(v);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(bool const& v) {
|
||||
printf("%*d", 3, int(v));
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(int32_t const& v) {
|
||||
printf("%*d", 5, v);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(uint32_t const& v) {
|
||||
printf("%*d", 5, v);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(int64_t const& v) {
|
||||
printf("%*lld", 5, static_cast<long long>(v));
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(uint64_t const& v) {
|
||||
printf("%*llu", 5, static_cast<unsigned long long>(v));
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(half_t const& v) {
|
||||
printf("%*.2f", 8, float(v));
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(float const& v) {
|
||||
printf("%*.2e", 10, v);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(double const& v) {
|
||||
printf("%*.3e", 11, v);
|
||||
}
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
@@ -122,6 +122,9 @@ using CUTE_STL_NAMESPACE::is_empty_v;
|
||||
|
||||
using CUTE_STL_NAMESPACE::invoke_result_t;
|
||||
|
||||
using CUTE_STL_NAMESPACE::common_type;
|
||||
using CUTE_STL_NAMESPACE::common_type_t;
|
||||
|
||||
// <utility>
|
||||
using CUTE_STL_NAMESPACE::declval;
|
||||
|
||||
|
||||
@@ -47,6 +47,18 @@ namespace cutlass {
|
||||
namespace arch {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// Enumerates the reserved named barriers to avoid potential conflicts
|
||||
// This enum class specifies the NamedBarriers reserved by CUTLASS.
|
||||
enum class ReservedNamedBarriers {
|
||||
EpilogueBarrier = 0,
|
||||
TransposeBarrier = 1,
|
||||
TransformBarrier = 2,
|
||||
StreamkBarrier0 = 3,
|
||||
StreamkBarrier1 = 4
|
||||
, FirstUserBarrier = StreamkBarrier1 + 1
|
||||
};
|
||||
|
||||
|
||||
class NamedBarrier {
|
||||
|
||||
// Data Members:
|
||||
@@ -60,9 +72,19 @@ class NamedBarrier {
|
||||
|
||||
public:
|
||||
|
||||
// Constructor for CUTLASS developers:
|
||||
// effective barrier ID starts from 0
|
||||
CUTLASS_DEVICE
|
||||
NamedBarrier(uint32_t num_threads, ReservedNamedBarriers reserved_named_barriers)
|
||||
: num_threads_(num_threads), id_(static_cast<uint32_t>(reserved_named_barriers)) {}
|
||||
|
||||
// Constructor for CUTLASS users:
|
||||
// effective barrier ID starts from ReservedNamedBarrierCount
|
||||
CUTLASS_DEVICE
|
||||
NamedBarrier(uint32_t num_threads, uint32_t id = 0)
|
||||
: num_threads_(num_threads), id_(id) {}
|
||||
: num_threads_(num_threads), id_(id + ReservedNamedBarrierCount) {
|
||||
CUTLASS_ASSERT(id + ReservedNamedBarrierCount <= HardwareMaxNumNamedBarriers && "Effective barrier_id should not exceed 16.");
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void arrive_and_wait() const {
|
||||
@@ -80,8 +102,52 @@ class NamedBarrier {
|
||||
}
|
||||
|
||||
// Static variants
|
||||
|
||||
// Calling interface for CUTLASS users:
|
||||
// effective barrier ID starts from ReservedNamedBarrierCount
|
||||
CUTLASS_DEVICE
|
||||
static void arrive_and_wait(uint32_t num_threads, uint32_t barrier_id) {
|
||||
arrive_and_wait_internal(num_threads, barrier_id + ReservedNamedBarrierCount);
|
||||
}
|
||||
|
||||
// Calling interface for CUTLASS developers:
|
||||
// effective barrier ID starts from 0
|
||||
CUTLASS_DEVICE
|
||||
static void arrive_and_wait(uint32_t num_threads, ReservedNamedBarriers reserved_named_barriers) {
|
||||
arrive_and_wait_internal(num_threads, static_cast<int>(reserved_named_barriers));
|
||||
}
|
||||
|
||||
// Calling interface for CUTLASS users:
|
||||
// effective barrier ID starts from ReservedNamedBarrierCount
|
||||
CUTLASS_DEVICE
|
||||
static void arrive(uint32_t num_threads, uint32_t barrier_id) {
|
||||
arrive_internal(num_threads, barrier_id + ReservedNamedBarrierCount);
|
||||
}
|
||||
|
||||
// Calling interface for CUTLASS developers:
|
||||
// effective barrier ID starts from 0
|
||||
CUTLASS_DEVICE
|
||||
static void arrive(uint32_t num_threads, ReservedNamedBarriers reserved_named_barriers) {
|
||||
arrive_internal(num_threads, static_cast<int>(reserved_named_barriers));
|
||||
}
|
||||
|
||||
// Calling interface for CUTLASS users:
|
||||
// effective barrier ID starts from ReservedNamedBarrierCount
|
||||
CUTLASS_DEVICE
|
||||
static void sync(uint32_t num_threads, uint32_t barrier_id) {
|
||||
sync_internal(num_threads, barrier_id + ReservedNamedBarrierCount);
|
||||
}
|
||||
|
||||
// Calling interface for CUTLASS developers:
|
||||
// effective barrier ID starts from 0
|
||||
CUTLASS_DEVICE
|
||||
static void sync(uint32_t num_threads, ReservedNamedBarriers reserved_named_barriers) {
|
||||
sync_internal(num_threads, static_cast<int>(reserved_named_barriers));
|
||||
}
|
||||
|
||||
private:
|
||||
CUTLASS_DEVICE
|
||||
static void arrive_and_wait_internal(uint32_t num_threads, uint32_t barrier_id) {
|
||||
#if CUDA_BARRIER_ENABLED
|
||||
asm volatile("bar.sync %0, %1;" : : "r"(barrier_id), "r"(num_threads));
|
||||
#elif defined(__CUDA_ARCH__)
|
||||
@@ -90,7 +156,7 @@ class NamedBarrier {
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static void arrive(uint32_t num_threads, uint32_t barrier_id) {
|
||||
static void arrive_internal(uint32_t num_threads, uint32_t barrier_id) {
|
||||
#if CUDA_BARRIER_ENABLED
|
||||
asm volatile("bar.arrive %0, %1;" : : "r"(barrier_id), "r"(num_threads));
|
||||
#elif defined(__CUDA_ARCH__)
|
||||
@@ -99,9 +165,16 @@ class NamedBarrier {
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static void sync(uint32_t num_threads, uint32_t barrier_id) {
|
||||
NamedBarrier::arrive_and_wait(num_threads, barrier_id);
|
||||
static void sync_internal(uint32_t num_threads, uint32_t barrier_id) {
|
||||
NamedBarrier::arrive_and_wait_internal(num_threads, barrier_id);
|
||||
}
|
||||
|
||||
public:
|
||||
// Currently we reserve 8 NamedBarriers for CUTLASS' own use cases,
|
||||
// while leaving the renaming for general users.
|
||||
static const uint32_t ReservedNamedBarrierCount = static_cast<uint32_t>(ReservedNamedBarriers::FirstUserBarrier);
|
||||
static const uint32_t HardwareMaxNumNamedBarriers = 16;
|
||||
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -80,7 +80,7 @@ public:
|
||||
static int const kElementsPerStoredItem = int(sizeof(Storage) * 8) / sizeof_bits<T>::value;
|
||||
|
||||
/// Number of storage elements
|
||||
static size_t const kStorageElements = N / kElementsPerStoredItem;
|
||||
static size_t const kStorageElements = (N + kElementsPerStoredItem - 1) / kElementsPerStoredItem;
|
||||
|
||||
/// Number of logical elements
|
||||
static size_t const kElements = N;
|
||||
|
||||
@@ -68,7 +68,7 @@ template <
|
||||
struct NamedBarrierSync {
|
||||
CUTLASS_DEVICE
|
||||
static void sync() {
|
||||
cutlass::arch::NamedBarrier::sync(ThreadCount, BarrierId);
|
||||
cutlass::arch::NamedBarrier::sync(ThreadCount, static_cast<arch::ReservedNamedBarriers>(BarrierId));
|
||||
}
|
||||
};
|
||||
|
||||
@@ -227,9 +227,9 @@ template <
|
||||
uint32_t MaxNumNamedBarriers = 16
|
||||
>
|
||||
struct NamedBarrierManager {
|
||||
static constexpr uint32_t HardwareMaxNumNamedBarriers = 16;
|
||||
static_assert(MaxNumNamedBarriers <= HardwareMaxNumNamedBarriers);
|
||||
static_assert(MaxNumNamedBarriers + Offset <= HardwareMaxNumNamedBarriers, "Barrier IDs cannot exceed 15");
|
||||
|
||||
static_assert(MaxNumNamedBarriers <= arch::NamedBarrier::HardwareMaxNumNamedBarriers);
|
||||
static_assert(MaxNumNamedBarriers + Offset <= arch::NamedBarrier::HardwareMaxNumNamedBarriers, "Barrier IDs cannot exceed 15");
|
||||
|
||||
// Number of threads participating in the barrier
|
||||
static constexpr uint32_t ThreadCount = ThreadCount_;
|
||||
|
||||
+26
-11
@@ -55,6 +55,7 @@
|
||||
#include <cstring>
|
||||
#endif
|
||||
|
||||
#include <cuda_bf16.h>
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
namespace cutlass {
|
||||
@@ -83,6 +84,28 @@ struct alignas(2) bfloat16_t {
|
||||
return h;
|
||||
}
|
||||
|
||||
private:
|
||||
struct from_32_bit_integer_t {};
|
||||
static constexpr from_32_bit_integer_t from_32_bit_integer{};
|
||||
|
||||
template<class T>
|
||||
CUTLASS_HOST_DEVICE
|
||||
explicit bfloat16_t(from_32_bit_integer_t, T x) {
|
||||
static_assert(cutlass::platform::is_integral<T>::value && sizeof(T) == 4, "Requires 32-bit integer");
|
||||
|
||||
float flt = static_cast<float>(x);
|
||||
uint32_t bits;
|
||||
|
||||
#if defined(__CUDA_ARCH__)
|
||||
bits = reinterpret_cast<uint32_t &>(flt);
|
||||
#else
|
||||
std::memcpy(&bits, &flt, sizeof(bits));
|
||||
#endif
|
||||
|
||||
storage = uint16_t(bits >> 16);
|
||||
}
|
||||
|
||||
public:
|
||||
/// Default constructor
|
||||
bfloat16_t() = default;
|
||||
|
||||
@@ -129,18 +152,10 @@ struct alignas(2) bfloat16_t {
|
||||
|
||||
/// Integer conversion - round toward nearest
|
||||
CUTLASS_HOST_DEVICE
|
||||
explicit bfloat16_t(int x) {
|
||||
float flt = static_cast<float>(x);
|
||||
uint32_t bits;
|
||||
explicit bfloat16_t(int x) : bfloat16_t(from_32_bit_integer, x) {}
|
||||
|
||||
#if defined(__CUDA_ARCH__)
|
||||
bits = reinterpret_cast<uint32_t &>(flt);
|
||||
#else
|
||||
std::memcpy(&bits, &flt, sizeof(bits));
|
||||
#endif
|
||||
|
||||
storage = uint16_t(bits >> 16);
|
||||
}
|
||||
CUTLASS_HOST_DEVICE
|
||||
explicit bfloat16_t(uint32_t x) : bfloat16_t(from_32_bit_integer, x) {}
|
||||
|
||||
/// Converts to float
|
||||
CUTLASS_HOST_DEVICE
|
||||
|
||||
@@ -35,7 +35,6 @@
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdio>
|
||||
#include <cuda_runtime_api.h>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
@@ -28,6 +28,16 @@
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*
|
||||
Note: CUTLASS 3x increases the host compiler requirements to C++17. However, certain
|
||||
existing integrations of CUTLASS require C++11 host compilers.
|
||||
|
||||
Until this requirement can be lifted, certain headers with this annotation are required
|
||||
to be remain consistent with C++11 syntax.
|
||||
|
||||
C++11 compatibility is enforced by this unit test: `cutlass_test_unit_core_cpp11`.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuComplex.h>
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2017 - 2023 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Interface betweeen a CUTLASS device-wide operator and CUDA.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda_runtime_api.h>
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
#include "cutlass/platform/platform.h"
|
||||
#if ! defined(__CUDACC_RTC__)
|
||||
#include <cstdio>
|
||||
#endif
|
||||
|
||||
#if ((__CUDACC_VER_MAJOR__ >= 12) || ((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ >= 8)))
|
||||
# define CUTLASS_SM90_CLUSTER_LAUNCH_ENABLED
|
||||
#endif
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//
|
||||
// Macro-level guard for CUDA Host Adapter
|
||||
//
|
||||
#if !defined(CUTLASS_ENABLE_CUDA_HOST_ADAPTER)
|
||||
#define CUTLASS_ENABLE_CUDA_HOST_ADAPTER false
|
||||
#endif
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// This class defines an object which abstracts interactions between the CUTLASS device-wide GEMM and
|
||||
/// CUDA. The intention is to enable CUTLASS to be used with both the CUDA Runtime API and CUDA Driver API.
|
||||
struct CudaHostAdapter {
|
||||
|
||||
/// Limit the number of kernels
|
||||
static constexpr int32_t kMaximumKernelCount = 4;
|
||||
|
||||
/// Maximum cluster size
|
||||
static constexpr int MaxClusterSize = 32;
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
/// Handles
|
||||
void *kernel_handles[kMaximumKernelCount];
|
||||
int32_t kernel_count = 0;
|
||||
|
||||
CudaHostAdapter() = default;
|
||||
|
||||
/// Dtor
|
||||
virtual ~CudaHostAdapter() {}
|
||||
|
||||
/// Copy Ctor deleted
|
||||
CudaHostAdapter(const CudaHostAdapter&) = delete;
|
||||
|
||||
/// Copy Assignment deleted
|
||||
CudaHostAdapter& operator=(const CudaHostAdapter&) = delete;
|
||||
|
||||
/// Move ctor deleted
|
||||
CudaHostAdapter(CudaHostAdapter&&) = delete;
|
||||
|
||||
/// Move assignment deleted
|
||||
CudaHostAdapter& operator=(CudaHostAdapter&&) = delete;
|
||||
|
||||
|
||||
/// Ctor
|
||||
inline CudaHostAdapter(
|
||||
void **kernel_handles_,
|
||||
int32_t kernel_count_
|
||||
):
|
||||
kernel_count(kernel_count_)
|
||||
{
|
||||
CUTLASS_ASSERT(kernel_count >= 0);
|
||||
for (int32_t i = 0; i < kernel_count && i < kMaximumKernelCount; ++i) {
|
||||
kernel_handles[i] = kernel_handles_[i];
|
||||
}
|
||||
}
|
||||
|
||||
/// Queries the occupancy of a kernel
|
||||
virtual Status query_occupancy(
|
||||
int32_t *device_sms,
|
||||
int32_t *sm_occupancy,
|
||||
int32_t kernel_index,
|
||||
int32_t thread_count,
|
||||
int32_t smem_size) = 0;
|
||||
|
||||
/// Launches a kernel without using Threadblock Clusters.
|
||||
virtual Status launch(
|
||||
dim3 const grid_dims,
|
||||
dim3 const block_dims,
|
||||
size_t const smem_size,
|
||||
cudaStream_t cuda_stream,
|
||||
void** kernel_params,
|
||||
int32_t kernel_index) = 0;
|
||||
|
||||
/// Launches a kernel using the CUDA Extensible Launch API and Threadblock Clusters.
|
||||
virtual Status launch(
|
||||
dim3 const grid_dims,
|
||||
dim3 const cluster_dims,
|
||||
dim3 const block_dims,
|
||||
size_t const smem_size,
|
||||
cudaStream_t cuda_stream,
|
||||
void** kernel_params,
|
||||
int32_t kernel_index) = 0;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,64 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2023 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 "cute/container/tuple.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <size_t I, class Tuple>
|
||||
struct deduce_mixed_width_dtype {
|
||||
static_assert(I >= 0u && I <= 2u, "Valid indices are 0, 1, and 2, which represent Operand, Scale, and Bias, respectively.");
|
||||
|
||||
private:
|
||||
using underlying_tuple = cute::conditional_t<cute::is_tuple<Tuple>::value, Tuple, cute::tuple<Tuple>>;
|
||||
static constexpr size_t valid_index = cute::min(I, cute::tuple_size_v<underlying_tuple> - 1);
|
||||
|
||||
public:
|
||||
using type = cute::conditional_t<(I < cute::tuple_size_v<underlying_tuple>),
|
||||
cute::tuple_element_t<valid_index, underlying_tuple>,
|
||||
void>;
|
||||
};
|
||||
|
||||
template <size_t I, class Tuple>
|
||||
using deduce_mixed_width_dtype_t = typename deduce_mixed_width_dtype<I, Tuple>::type;
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
@@ -187,6 +187,7 @@ constexpr bool is_tma_copy_engine() {
|
||||
|| cute::is_base_of_v<cute::SM90_TMA_LOAD_IM2COL, GmemTiledCopy>
|
||||
|| cute::is_base_of_v<cute::SM90_TMA_LOAD_IM2COL_MULTICAST, GmemTiledCopy>
|
||||
|| cute::is_base_of_v<cute::SM90_TMA_STORE, GmemTiledCopy>
|
||||
|| cute::is_base_of_v<cute::SM90_TMA_STORE_IM2COL, GmemTiledCopy>
|
||||
) {
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -104,19 +104,22 @@ sm90_compute_tile_shape_or_override() {
|
||||
if constexpr (cute::is_same_v<EpilogueTileType, EpilogueTileAuto>) {
|
||||
|
||||
if constexpr (detail::sm90_is_cooperative_v<Schedule>) {
|
||||
using N_tile = decltype(cute::min(_32{}, get<1>(TileShape_MNK{})));
|
||||
if constexpr (size<0>(TileShape_MNK{}) >= 128) {
|
||||
return Shape<_128,_32>{};
|
||||
return Shape<_128, N_tile>{};
|
||||
}
|
||||
else {
|
||||
return Shape<_64,_32>{};
|
||||
return Shape<_64, N_tile>{};
|
||||
}
|
||||
}
|
||||
else if constexpr (detail::sm90_is_warp_specialized_v<Schedule>) {
|
||||
if constexpr (sizeof_bits_v<ElementD> == 8) {
|
||||
return Shape<_64,_64>{};
|
||||
using N_tile = decltype(cute::min(_64{}, get<1>(TileShape_MNK{})));
|
||||
return Shape<_64, N_tile>{};
|
||||
}
|
||||
else {
|
||||
return Shape<_64,_32>{};
|
||||
using N_tile = decltype(cute::min(_32{}, get<1>(TileShape_MNK{})));
|
||||
return Shape<_64,N_tile>{};
|
||||
}
|
||||
}
|
||||
else {
|
||||
@@ -265,6 +268,13 @@ struct Sm90TmaBuilderImpl {
|
||||
using GmemStrideTypeC = cutlass::detail::TagToStrideC_t<GmemLayoutTagC>;
|
||||
using GmemStrideTypeD = cutlass::detail::TagToStrideC_t<GmemLayoutTagD>;
|
||||
|
||||
using CopyOpS2G =
|
||||
SM90_TMA_STORE
|
||||
;
|
||||
using CopyOpG2S =
|
||||
SM90_TMA_LOAD
|
||||
;
|
||||
|
||||
// TMA builder allows for passing callbacks directly, which is either a fusion::FusionCallbacks
|
||||
// instance or a direct visitor implementation, e.g. fusion::Sm90LinearCombination
|
||||
using FusionCallbacks =
|
||||
@@ -285,10 +295,10 @@ struct Sm90TmaBuilderImpl {
|
||||
ElementD,
|
||||
GmemStrideTypeD,
|
||||
FusionCallbacks,
|
||||
SM90_TMA_LOAD,
|
||||
CopyOpG2S,
|
||||
decltype(detail::sm90_get_epilogue_smem_swizzle_layout_atom<GmemStrideTypeC, ElementC, EpilogueTile_MN>()),
|
||||
decltype(detail::sm90_get_smem_load_op_for_source<GmemStrideTypeC, ElementC>()),
|
||||
SM90_TMA_STORE,
|
||||
CopyOpS2G,
|
||||
decltype(detail::sm90_get_epilogue_smem_swizzle_layout_atom<GmemStrideTypeD, ElementD, EpilogueTile_MN>()),
|
||||
decltype(detail::sm90_get_smem_store_op_for_accumulator<GmemStrideTypeD, ElementD>())
|
||||
>;
|
||||
@@ -400,6 +410,7 @@ template <
|
||||
class ElementD,
|
||||
class GmemLayoutTagD,
|
||||
int AlignmentD,
|
||||
class Schedule,
|
||||
FloatRoundStyle RoundStyle
|
||||
>
|
||||
struct CollectiveBuilder<
|
||||
@@ -416,9 +427,11 @@ struct CollectiveBuilder<
|
||||
ElementD,
|
||||
GmemLayoutTagD,
|
||||
AlignmentD,
|
||||
NoSmemWarpSpecialized,
|
||||
fusion::LinearCombination<ElementD,ElementCompute,ElementCompute,RoundStyle>,
|
||||
void> {
|
||||
Schedule,
|
||||
fusion::LinearCombination<ElementD,ElementCompute,ElementC_,ElementCompute,RoundStyle>,
|
||||
cute::enable_if_t<cute::is_same_v<Schedule, NoSmemWarpSpecialized> ||
|
||||
cute::is_same_v<Schedule, NoSmemWarpSpecializedArray> ||
|
||||
cute::is_same_v<Schedule, NoSmemWarpSpecializedGroup> >> {
|
||||
|
||||
// Passing void C disables source load
|
||||
using ElementC = cute::conditional_t<cute::is_void_v<ElementC_>,
|
||||
@@ -433,12 +446,21 @@ struct CollectiveBuilder<
|
||||
ElementD, FragmentSize, ElementAccumulator, ElementCompute,
|
||||
ScaleType, RoundStyle, ElementC>;
|
||||
|
||||
using CollectiveOp = cutlass::epilogue::collective::detail::Sm90TmaWarpSpecializedAdapter<
|
||||
cutlass::epilogue::collective::DefaultEpilogue<
|
||||
cutlass::detail::TagToStrideC_t<GmemLayoutTagC>,
|
||||
cutlass::detail::TagToStrideC_t<GmemLayoutTagD>,
|
||||
ThreadOp,
|
||||
cutlass::gemm::EpilogueDefault>
|
||||
using CollectiveOp = cute::conditional_t<
|
||||
cute::is_same_v<Schedule, NoSmemWarpSpecialized>,
|
||||
cutlass::epilogue::collective::detail::Sm90TmaWarpSpecializedAdapter<
|
||||
cutlass::epilogue::collective::DefaultEpilogue<
|
||||
cutlass::detail::TagToStrideC_t<GmemLayoutTagC>,
|
||||
cutlass::detail::TagToStrideC_t<GmemLayoutTagD>,
|
||||
ThreadOp,
|
||||
cutlass::gemm::EpilogueDefault>>,
|
||||
// Epilogue for Ptr-Array and Grouped Gemm
|
||||
cutlass::epilogue::collective::detail::Sm90TmaWarpSpecializedAdapter<
|
||||
cutlass::epilogue::collective::DefaultEpilogueArray<
|
||||
cutlass::detail::TagToStrideC_t<GmemLayoutTagC>,
|
||||
cutlass::detail::TagToStrideC_t<GmemLayoutTagD>,
|
||||
ThreadOp,
|
||||
Schedule>>
|
||||
>;
|
||||
};
|
||||
|
||||
@@ -533,6 +555,9 @@ struct CollectiveBuilder<
|
||||
FusionOperation,
|
||||
void> {
|
||||
private:
|
||||
static_assert(cute::is_same_v<FusionOperation, fusion::LinearCombination<ElementD,ElementCompute,ElementC,ElementCompute>>,
|
||||
"Auto schedule doesn't support fusion. Use one of the TmaWarpSpecialized schedules instead.");
|
||||
|
||||
// Pick No-Smem epilogue as the Auto Epilogue Schedule (Auto schedules do not guarantee best performance)
|
||||
// since TMA epilogues are not compatible with non-TMA non-WS mainloops
|
||||
using EpilogueSchedule = NoSmemWarpSpecialized;
|
||||
@@ -595,7 +620,7 @@ CollectiveBuilder<
|
||||
cute::is_base_of_v<TmaWarpSpecializedCooperativeElementwiseBase, Schedule> >> {
|
||||
private:
|
||||
using FusionOp =
|
||||
fusion::LinCombEltAct<Schedule::template ActivationFunctor, ElementD, ElementCompute, ElementCompute, Schedule::Round>;
|
||||
fusion::LinCombEltAct<Schedule::template ActivationFunctor, ElementD, ElementCompute, ElementC, ElementCompute, Schedule::Round>;
|
||||
using ImplSchedule =
|
||||
cute::conditional_t<cute::is_base_of_v<TmaWarpSpecializedElementwiseBase, Schedule>,
|
||||
TmaWarpSpecialized, TmaWarpSpecializedCooperative>;
|
||||
@@ -677,7 +702,7 @@ private:
|
||||
GmemStrideTypeAux, typename Schedule::ElementT>());
|
||||
using FusionOperationAux = fusion::LinCombPerRowBiasEltActAux<
|
||||
GmemLayoutTagD, Schedule::template ActivationFunctor, ElementD, ElementCompute,
|
||||
typename Schedule::ElementT, typename Schedule::ElementBias, ElementCompute
|
||||
typename Schedule::ElementT, typename Schedule::ElementBias, ElementC_, ElementCompute
|
||||
>;
|
||||
using FusionCallbacksAux = fusion::FusionCallbacks<
|
||||
DispatchPolicy, FusionOperationAux, TileShape_MNK, EpilogueTile_MN, SmemLayoutAtomAux, SmemCopyOpAux
|
||||
@@ -685,7 +710,7 @@ private:
|
||||
|
||||
using FusionOperationNoAux = fusion::LinCombPerRowBiasEltAct<
|
||||
Schedule::template ActivationFunctor, ElementD, ElementCompute,
|
||||
typename Schedule::ElementBias, ElementCompute
|
||||
typename Schedule::ElementBias, ElementC_, ElementCompute
|
||||
>;
|
||||
using FusionCallbacksNoAux = fusion::FusionCallbacks<
|
||||
DispatchPolicy, FusionOperationNoAux, TileShape_MNK, EpilogueTile_MN
|
||||
@@ -750,7 +775,7 @@ struct CollectiveBuilder<
|
||||
GmemLayoutTagD,
|
||||
AlignmentD,
|
||||
cutlass::gemm::EpilogueTransposed,
|
||||
fusion::LinearCombination<ElementD,ElementCompute,ElementCompute,RoundStyle>,
|
||||
fusion::LinearCombination<ElementD,ElementCompute,ElementC_,ElementCompute,RoundStyle>,
|
||||
void> {
|
||||
// Passing void C disables source load
|
||||
using ElementC = cute::conditional_t<cute::is_void_v<ElementC_>,
|
||||
|
||||
@@ -62,7 +62,7 @@ template <
|
||||
class GmemLayoutTagD,
|
||||
int AlignmentD,
|
||||
class Schedule,
|
||||
class FusionOpOrCallbacks = cutlass::epilogue::fusion::LinearCombination<ElementD,ElementCompute>,
|
||||
class FusionOpOrCallbacks = cutlass::epilogue::fusion::LinearCombination<ElementD,ElementCompute,ElementC,ElementCompute>,
|
||||
class Enable = void
|
||||
>
|
||||
struct CollectiveBuilder {
|
||||
|
||||
@@ -54,6 +54,7 @@ class CollectiveEpilogue {
|
||||
|
||||
#include "detail.hpp"
|
||||
#include "default_epilogue.hpp"
|
||||
#include "default_epilogue_array.hpp"
|
||||
#include "epilogue_tensor_broadcast.hpp"
|
||||
#include "sm70_epilogue_vectorized.hpp"
|
||||
#include "sm90_epilogue_tma_warpspecialized.hpp"
|
||||
|
||||
@@ -131,8 +131,9 @@ public:
|
||||
return true;
|
||||
}
|
||||
|
||||
// Note: SharedStorage is unused for DefaultEpilogue
|
||||
CUTLASS_HOST_DEVICE
|
||||
DefaultEpilogue(Params const& params_)
|
||||
DefaultEpilogue(Params const& params_, SharedStorage const& shared_storage = SharedStorage())
|
||||
: params(params_), epilogue_op(params_.thread) { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2023 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief Functor performing elementwise operations used by epilogues.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
#include "cutlass/epilogue/collective/detail.hpp"
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/numeric/int.hpp"
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Applies an element wise operation to all elements within the fragment
|
||||
// and writes them out to destination storage.
|
||||
template <
|
||||
class StrideC_,
|
||||
class StrideD_,
|
||||
class ThreadEpilogueOp_,
|
||||
class EpilogueSchedule_
|
||||
>
|
||||
class DefaultEpilogueArray {
|
||||
public:
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using EpilogueSchedule = EpilogueSchedule_;
|
||||
|
||||
// derived types of output thread level operator
|
||||
using ThreadEpilogueOp = ThreadEpilogueOp_;
|
||||
using ElementOutput = typename ThreadEpilogueOp::ElementOutput;
|
||||
using ElementAccumulator = typename ThreadEpilogueOp::ElementAccumulator;
|
||||
using ElementCompute = typename ThreadEpilogueOp::ElementCompute;
|
||||
using ElementScalar = ElementCompute;
|
||||
using ElementC = typename ThreadEpilogueOp::ElementC;
|
||||
using StrideC = StrideC_;
|
||||
using ElementD = typename ThreadEpilogueOp::ElementD;
|
||||
using StrideD = StrideD_;
|
||||
using StridesC = cute::conditional_t<cute::is_same_v<EpilogueSchedule, NoSmemWarpSpecializedGroup>,
|
||||
StrideC const*, StrideC>;
|
||||
using StridesD = cute::conditional_t<cute::is_same_v<EpilogueSchedule, NoSmemWarpSpecializedGroup>,
|
||||
StrideD const*, StrideD>;
|
||||
|
||||
using GmemTiledCopyC = void;
|
||||
using GmemTiledCopyD = void;
|
||||
|
||||
static const int kOutputAlignment = ThreadEpilogueOp::kCount;
|
||||
using AlignmentType = typename cute::uint_bit<sizeof_bits<ElementOutput>::value * kOutputAlignment>::type;
|
||||
|
||||
static_assert(cute::is_same_v<EpilogueSchedule, NoSmemWarpSpecializedGroup> ||
|
||||
cute::is_same_v<EpilogueSchedule, NoSmemWarpSpecializedArray>, "Incompatible epilogue schedule.");
|
||||
static_assert(rank(StrideC{}) == 3, "StrideCD must be rank-3: [M, N, L]");
|
||||
static_assert(rank(StrideD{}) == 3, "StrideCD must be rank-3: [M, N, L]");
|
||||
|
||||
struct SharedStorage { };
|
||||
|
||||
// Host side epilogue arguments
|
||||
struct Arguments {
|
||||
typename ThreadEpilogueOp::Params thread{};
|
||||
ElementC const** ptr_C = nullptr;
|
||||
StridesC dC{};
|
||||
ElementD** ptr_D = nullptr;
|
||||
StridesD dD{};
|
||||
};
|
||||
|
||||
// Device side epilogue params
|
||||
using Params = Arguments;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
template <class ProblemShape>
|
||||
static constexpr Params
|
||||
to_underlying_arguments(
|
||||
ProblemShape const&,
|
||||
Arguments const& args,
|
||||
[[maybe_unused]] void* workspace) {
|
||||
return args;
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static size_t
|
||||
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static cutlass::Status
|
||||
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
template<class ProblemShape>
|
||||
CUTLASS_HOST_DEVICE static bool
|
||||
can_implement(
|
||||
[[maybe_unused]] ProblemShape const& problem_shape,
|
||||
[[maybe_unused]] Arguments const& args) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
DefaultEpilogueArray(Params const& params_)
|
||||
: params(params_), epilogue_op(params_.thread) { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
is_source_needed() {
|
||||
return epilogue_op.is_source_needed();
|
||||
}
|
||||
|
||||
template<
|
||||
class ProblemShapeMNKL,
|
||||
class BlockShapeMNK,
|
||||
class BlockCoordMNKL,
|
||||
class FrgEngine, class FrgLayout,
|
||||
class TiledMma,
|
||||
class ResidueMNK
|
||||
>
|
||||
CUTLASS_HOST_DEVICE void
|
||||
operator()(
|
||||
ProblemShapeMNKL problem_shape_mnkl,
|
||||
BlockShapeMNK blk_shape_MNK,
|
||||
BlockCoordMNKL blk_coord_mnkl,
|
||||
cute::Tensor<FrgEngine, FrgLayout> const& accumulators,
|
||||
TiledMma tiled_mma,
|
||||
ResidueMNK residue_mnk,
|
||||
int thread_idx,
|
||||
[[maybe_unused]] char* smem_buf)
|
||||
{
|
||||
using namespace cute;
|
||||
using X = Underscore;
|
||||
|
||||
static_assert(rank(ProblemShapeMNKL{}) == 4, "ProblemShapeMNKL must be rank 4");
|
||||
static_assert(is_static<BlockShapeMNK>::value, "ThreadBlock tile shape must be static");
|
||||
static_assert(rank(BlockShapeMNK{}) == 3, "BlockShapeMNK must be rank 3");
|
||||
static_assert(rank(BlockCoordMNKL{}) == 4, "BlockCoordMNKL must be rank 3");
|
||||
|
||||
// Separate out problem shape for convenience
|
||||
auto M = get<0>(problem_shape_mnkl);
|
||||
auto N = get<1>(problem_shape_mnkl);
|
||||
auto L = get<3>(problem_shape_mnkl);
|
||||
// Batches are managed by using appropriate pointers to C and D matrices
|
||||
const int32_t mock_L = 1;
|
||||
const int32_t mock_l_coord = 0;
|
||||
// Slice to get the tile this CTA is responsible for
|
||||
auto [m_coord, n_coord, k_coord, l_coord] = blk_coord_mnkl;
|
||||
|
||||
StrideC stride_c;
|
||||
StrideD stride_d;
|
||||
if constexpr (cute::is_same_v<EpilogueSchedule, NoSmemWarpSpecializedGroup>) {
|
||||
stride_c = detail::get_epilogue_stride<EpilogueSchedule>(params.dC[l_coord]);
|
||||
stride_d = detail::get_epilogue_stride<EpilogueSchedule>(params.dD[l_coord]);
|
||||
}
|
||||
else {
|
||||
stride_c = detail::get_epilogue_stride<EpilogueSchedule>(params.dC);
|
||||
stride_d = detail::get_epilogue_stride<EpilogueSchedule>(params.dD);
|
||||
}
|
||||
|
||||
// Represent the full output tensor
|
||||
Tensor mC_mnl = make_tensor(make_gmem_ptr(params.ptr_C[l_coord]), make_shape(M,N,mock_L), stride_c); // (m,n,l)
|
||||
Tensor mD_mnl = make_tensor(make_gmem_ptr(params.ptr_D[l_coord]), make_shape(M,N,mock_L), stride_d); // (m,n,l)
|
||||
Tensor gC_mnl = local_tile(mC_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
|
||||
Tensor gD_mnl = local_tile(mD_mnl, blk_shape_MNK, make_coord(_,_,_), Step<_1,_1, X>{}); // (BLK_M,BLK_N,m,n,l)
|
||||
|
||||
Tensor gC = gC_mnl(_,_,m_coord,n_coord, mock_l_coord); // (BLK_M,BLK_N)
|
||||
Tensor gD = gD_mnl(_,_,m_coord,n_coord, mock_l_coord); // (BLK_M,BLK_N)
|
||||
|
||||
// Partition source and destination tiles to match the accumulator partitioning
|
||||
auto thr_mma = tiled_mma.get_thread_slice(thread_idx);
|
||||
Tensor tCgD = thr_mma.partition_C(gD); // (VEC,THR_M,THR_N)
|
||||
Tensor tCgC = thr_mma.partition_C(gC); // (VEC,THR_M,THR_N)
|
||||
|
||||
static_assert(is_static<FrgLayout>::value, "Accumulator layout must be static");
|
||||
CUTE_STATIC_ASSERT_V(size(tCgC) == size(tCgD),
|
||||
"Source and destination must have the same number of elements.");
|
||||
CUTE_STATIC_ASSERT_V(size(tCgD) == size(accumulators),
|
||||
"Accumulator count must have the same destination element count.");
|
||||
|
||||
// Make an identity coordinate tensor for predicating our output MN tile
|
||||
auto cD = make_identity_tensor(make_shape(unwrap(shape<0>(gD)), unwrap(shape<1>(gD))));
|
||||
Tensor tCcD = thr_mma.partition_C(cD);
|
||||
|
||||
// source is needed
|
||||
if (epilogue_op.is_source_needed()) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(accumulators); ++i) {
|
||||
if (elem_less(tCcD(i), make_coord(get<0>(residue_mnk), get<1>(residue_mnk)))) {
|
||||
tCgD(i) = epilogue_op(accumulators(i), tCgC(i));
|
||||
}
|
||||
}
|
||||
}
|
||||
// source is not needed, avoid load
|
||||
else {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(accumulators); ++i) {
|
||||
if (elem_less(tCcD(i), make_coord(get<0>(residue_mnk), get<1>(residue_mnk)))) {
|
||||
tCgD(i) = epilogue_op(accumulators(i));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
Params params;
|
||||
ThreadEpilogueOp epilogue_op;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace collective
|
||||
} // namespace epilogue
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -170,7 +170,8 @@ public:
|
||||
[[maybe_unused]] TileCoordMNKL tile_coord_mnkl,
|
||||
[[maybe_unused]] TiledMma tiled_mma,
|
||||
[[maybe_unused]] int thread_idx,
|
||||
[[maybe_unused]] TensorStorage& shared_tensors)
|
||||
[[maybe_unused]] TensorStorage& shared_tensors,
|
||||
[[maybe_unused]] int subtile_idx=-1)
|
||||
{
|
||||
return load_pipe_producer_state;
|
||||
}
|
||||
@@ -202,7 +203,8 @@ public:
|
||||
cute::Tensor<AccEngine,AccLayout> accumulators,
|
||||
TiledMma tiled_mma,
|
||||
int thread_idx,
|
||||
TensorStorage& shared_tensors)
|
||||
TensorStorage& shared_tensors,
|
||||
int subtile_index = -1)
|
||||
{
|
||||
constexpr int BLK_M_RANK = cute::rank<0>(tile_shape_MNK);
|
||||
auto m_max_coord = unwrap(cute::transform(make_seq<BLK_M_RANK>{}, [&](auto i) {
|
||||
|
||||
@@ -85,6 +85,7 @@ public:
|
||||
using CopyAtomR2G = CopyAtomR2G_;
|
||||
|
||||
static const int kOutputAlignment = ThreadEpilogueOp::kCount;
|
||||
|
||||
using AlignmentType = typename cute::uint_bit<sizeof_bits<ElementOutput>::value * kOutputAlignment>::type;
|
||||
|
||||
static_assert(cute::rank(StrideC{}) == 3, "StrideCD must be rank-3: [M, N, L]");
|
||||
@@ -179,7 +180,7 @@ public:
|
||||
|
||||
// synchronizing function for smem reads/writes
|
||||
#if CUDA_BARRIER_ENABLED
|
||||
auto synchronize = [] () { cutlass::arch::NamedBarrier::sync(typename TiledCopyS2R::TiledNumThr{}, 0); };
|
||||
auto synchronize = [] () { cutlass::arch::NamedBarrier::sync(typename TiledCopyS2R::TiledNumThr{}, cutlass::arch::ReservedNamedBarriers::EpilogueBarrier); };
|
||||
#else
|
||||
auto synchronize = [] () { __syncthreads(); };
|
||||
#endif
|
||||
|
||||
@@ -109,8 +109,8 @@ public:
|
||||
using CopyOpR2S = CopyOpR2S_;
|
||||
|
||||
using ThreadEpilogueOp = typename epilogue::fusion::FusionCallbacksTraits<FusionCallbacks>::Operation;
|
||||
using GmemTiledCopyC = SM90_TMA_LOAD;
|
||||
using GmemTiledCopyD = SM90_TMA_STORE;
|
||||
using GmemTiledCopyC = CopyOpG2S;
|
||||
using GmemTiledCopyD = CopyOpS2G;
|
||||
|
||||
static_assert(!is_layout<EpilogueTile>::value && is_tuple<EpilogueTile>::value, "EpilogueTile must be a cute::Tile or cute::Shape");
|
||||
static_assert(cute::rank(CtaTileMNK{}) == 3, "CtaTileMNK must be rank-3: [CTA_M, CTA_N, CTA_K]");
|
||||
@@ -198,12 +198,12 @@ public:
|
||||
struct Params {
|
||||
using TMA_C = decltype(make_tma_copy(
|
||||
CopyOpG2S{},
|
||||
make_tensor(static_cast<SmemElementC const*>(nullptr),
|
||||
make_tensor(make_gmem_ptr(static_cast<SmemElementC const*>(nullptr)),
|
||||
repeat_like(StrideC{}, int32_t(0)), StrideC{}),
|
||||
SmemLayoutC{}(_,_,0)));
|
||||
using TMA_D = decltype(make_tma_copy(
|
||||
CopyOpS2G{},
|
||||
make_tensor(static_cast<ElementD const*>(nullptr),
|
||||
make_tensor(make_gmem_ptr(static_cast<ElementD const*>(nullptr)),
|
||||
repeat_like(StrideD{}, int32_t(0)), StrideD{}),
|
||||
SmemLayoutD{}(_,_,0)));
|
||||
|
||||
@@ -225,14 +225,20 @@ public:
|
||||
// Optionally append 1s until problem shape is rank-4 in case its is only rank-3 (MNK)
|
||||
auto problem_shape_MNKL = append<4>(problem_shape, 1);
|
||||
auto [M, N, K, L] = problem_shape_MNKL;
|
||||
auto M_C =
|
||||
size(M)
|
||||
;
|
||||
auto M_D =
|
||||
size(M)
|
||||
;
|
||||
|
||||
typename Params::TMA_C tma_load_c;
|
||||
if constexpr (not cute::is_void_v<ElementC>) {
|
||||
Tensor tensor_c = make_tensor(args.ptr_C, make_layout(make_shape(M,N,L), args.dC));
|
||||
Tensor tensor_c = make_tensor(make_gmem_ptr(args.ptr_C), make_layout(make_shape(M_C,N,L), args.dC));
|
||||
tma_load_c = make_tma_copy(CopyOpG2S{}, tensor_c, SmemLayoutC{}(_,_,0));
|
||||
}
|
||||
|
||||
Tensor tensor_d = make_tensor(args.ptr_D, make_layout(make_shape(M,N,L), args.dD));
|
||||
Tensor tensor_d = make_tensor(make_gmem_ptr(args.ptr_D), make_layout(make_shape(M_D,N,L), args.dD));
|
||||
typename Params::TMA_D tma_store_d = make_tma_copy(
|
||||
CopyOpS2G{},
|
||||
tensor_d,
|
||||
@@ -332,25 +338,31 @@ public:
|
||||
TileCoordMNKL tile_coord_mnkl,
|
||||
TiledMma tiled_mma,
|
||||
int thread_idx,
|
||||
TensorStorage& shared_tensors) {
|
||||
TensorStorage& shared_tensors,
|
||||
int subtile_idx=-1) {
|
||||
using namespace cute;
|
||||
|
||||
// Indexing variables
|
||||
auto [M, N, K, L] = problem_shape_mnkl;
|
||||
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
|
||||
|
||||
auto coord_shape =
|
||||
make_coord(m_coord, n_coord, l_coord)
|
||||
;
|
||||
|
||||
// Tile residue
|
||||
auto m_max_coord = unwrap(cute::transform(make_seq<cute::rank<0>(tile_shape_MNK)>{}, [&](auto i) {
|
||||
auto m_max_coord = unwrap(cute::transform(make_seq<rank<0>(tile_shape_MNK)>{}, [&](auto i) {
|
||||
return get<0,i>(problem_shape_mnkl) - get<0,i>(tile_shape_MNK) * get<0,i>(tile_coord_mnkl);
|
||||
}));
|
||||
auto n_max_coord = unwrap(cute::transform(make_seq<cute::rank<1>(tile_shape_MNK)>{}, [&](auto i) {
|
||||
auto n_max_coord = unwrap(cute::transform(make_seq<rank<1>(tile_shape_MNK)>{}, [&](auto i) {
|
||||
return get<1,i>(problem_shape_mnkl) - get<1,i>(tile_shape_MNK) * get<1,i>(tile_coord_mnkl);
|
||||
}));
|
||||
auto residue_mn = make_coord(m_max_coord, n_max_coord);
|
||||
|
||||
// Represent the full source tensor, slice to get the tile this CTA is currently responsible for
|
||||
Tensor mC = params.tma_load_c.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
|
||||
Tensor gC = local_tile(mC, take<0,2>(CtaTileMNK{}), make_coord(m_coord,n_coord,l_coord)); // (CTA_M,CTA_N)
|
||||
Tensor mC_mn = params.tma_load_c.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
|
||||
Tensor mC = coalesce(mC_mn, take<0,2>(CtaTileMNK{}));
|
||||
Tensor gC = local_tile(mC, take<0,2>(CtaTileMNK{}), coord_shape); // (CTA_M,CTA_N)
|
||||
|
||||
// Apply epilogue subtile, get matching smem tensor
|
||||
SmemElementC* ptr_sC = reinterpret_cast<SmemElementC*>(shared_tensors.smem_D.data());
|
||||
@@ -391,6 +403,9 @@ public:
|
||||
for (int epi_n = 0; epi_n < size<3>(gC_epi); ++epi_n) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int epi_m = 0; epi_m < size<2>(gC_epi); ++epi_m) {
|
||||
if (subtile_idx != -1 && (epi_n * static_cast<int>(size<2>(gC_epi)) + epi_m) != subtile_idx) {
|
||||
continue;
|
||||
}
|
||||
// Acquire the lock for this stage
|
||||
constexpr uint16_t mcast_mask = 0;
|
||||
uint64_t* tma_barrier = load_pipeline.producer_get_barrier(load_pipe_producer_state);
|
||||
@@ -449,31 +464,36 @@ public:
|
||||
cute::Tensor<AccEngine,AccLayout> accumulators,
|
||||
TiledMma tiled_mma,
|
||||
int thread_idx,
|
||||
TensorStorage& shared_tensors) {
|
||||
TensorStorage& shared_tensors,
|
||||
int subtile_idx=-1) {
|
||||
using namespace cute;
|
||||
using ElementAccumulator = typename AccEngine::value_type;
|
||||
using ElementCompute_ = typename epilogue::fusion::FusionCallbacksTraits<FusionCallbacks>::ElementCompute;
|
||||
using ElementCompute = cute::conditional_t<cute::is_void_v<ElementCompute_>,ElementAccumulator,ElementCompute_>;
|
||||
|
||||
static_assert(is_rmem<AccEngine>::value, "Accumulator must be RF resident.");
|
||||
static_assert(cute::rank(AccLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA,MMA_M,MMA_N)");
|
||||
static_assert(cute::rank(ProblemShapeMNKL{}) == 4, "ProblemShapeMNKL must be rank 4");
|
||||
static_assert(rank(AccLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA,MMA_M,MMA_N)");
|
||||
static_assert(rank(ProblemShapeMNKL{}) == 4, "ProblemShapeMNKL must be rank 4");
|
||||
static_assert(is_static<TileShapeMNK>::value, "TileShapeMNK must be static");
|
||||
static_assert(cute::rank(TileShapeMNK{}) == 3, "TileShapeMNK must be rank 3");
|
||||
static_assert(cute::rank(TileCoordMNKL{}) == 4, "TileCoordMNKL must be rank 4");
|
||||
static_assert(rank(TileShapeMNK{}) == 3, "TileShapeMNK must be rank 3");
|
||||
static_assert(rank(TileCoordMNKL{}) == 4, "TileCoordMNKL must be rank 4");
|
||||
|
||||
// Indexing variables
|
||||
auto [M, N, K, L] = problem_shape_mnkl;
|
||||
auto [m_coord, n_coord, k_coord, l_coord] = tile_coord_mnkl;
|
||||
auto mma_tile_m = size<0>(typename TiledMma::TiledShape_MNK{});
|
||||
auto mma_tile_n = size<1>(typename TiledMma::TiledShape_MNK{});
|
||||
auto mma_tile_m = tile_size<0>(tiled_mma);
|
||||
auto mma_tile_n = tile_size<1>(tiled_mma);
|
||||
auto epi_tile_m = size<0>(EpilogueTile{});
|
||||
auto epi_tile_n = size<1>(EpilogueTile{});
|
||||
|
||||
auto coord_shape =
|
||||
make_coord(m_coord, n_coord, l_coord)
|
||||
;
|
||||
|
||||
// Represent the full output tensor, slice to get the tile this CTA is responsible for
|
||||
Tensor mD = params.tma_store_d.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
|
||||
Tensor gD = local_tile(mD, take<0,2>(CtaTileMNK{}), make_coord(m_coord,n_coord,l_coord)); // (CTA_M,CTA_N)
|
||||
Tensor mD_mn = params.tma_store_d.get_tma_tensor(make_shape(M,N,L)); // (M,N,L)
|
||||
Tensor mD = coalesce(mD_mn, take<0,2>(CtaTileMNK{}));
|
||||
Tensor gD = local_tile(mD, take<0,2>(CtaTileMNK{}), coord_shape); // (CTA_M,CTA_N)
|
||||
|
||||
// Apply epilogue subtiling
|
||||
Tensor gD_epi = flat_divide(gD, EpilogueTile{}); // (EPI_TILE_M,EPI_TILE_N,EPI_M,EPI_N)
|
||||
@@ -530,11 +550,11 @@ public:
|
||||
Tensor bSG_gD = thrblk_s2g.partition_D(gD_epi); // (S2G,S2G_M,S2G_N,EPI_M,EPI_N)
|
||||
|
||||
// Coordinate tensors and residue for tile quantization
|
||||
auto m_max_coord = unwrap(cute::transform(make_seq<cute::rank<0>(CtaTileMNK{})>{}, [&](auto i) {
|
||||
auto m_max_coord = unwrap(cute::transform(make_seq<rank<0>(CtaTileMNK{})>{}, [&](auto i) {
|
||||
auto c_m = get<0,i>(problem_shape_mnkl) - get<0,i>(CtaTileMNK{}) * get<0,i>(tile_coord_mnkl);
|
||||
return cute::max(0, c_m);
|
||||
}));
|
||||
auto n_max_coord = unwrap(cute::transform(make_seq<cute::rank<1>(CtaTileMNK{})>{}, [&](auto i) {
|
||||
auto n_max_coord = unwrap(cute::transform(make_seq<rank<1>(CtaTileMNK{})>{}, [&](auto i) {
|
||||
auto c_n = get<1,i>(problem_shape_mnkl) - get<1,i>(CtaTileMNK{}) * get<1,i>(tile_coord_mnkl);
|
||||
return cute::max(0, c_n);
|
||||
}));
|
||||
@@ -559,13 +579,13 @@ public:
|
||||
tRS_cD,
|
||||
tRS_rC
|
||||
};
|
||||
auto cst_callbacks = fusion_callbacks.template get_consumer_store_callbacks<RefSrc>(cst_args);
|
||||
auto cst_callbacks = fusion_callbacks.get_consumer_store_callbacks<RefSrc>(cst_args);
|
||||
bool is_producer_load_needed = fusion_callbacks.is_producer_load_needed();
|
||||
bool is_C_load_needed = is_source_supported && fusion_callbacks.is_C_load_needed();
|
||||
|
||||
// Thread synchronizer for previously issued waits or fences
|
||||
// to ensure visibility of smem reads/writes to threads or TMA unit
|
||||
auto synchronize = [&] () { cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 0); };
|
||||
auto synchronize = [&] () { cutlass::arch::NamedBarrier::sync(size(TiledMma{}), cutlass::arch::ReservedNamedBarriers::EpilogueBarrier); };
|
||||
|
||||
// Predication for TMA store (one warp issues TMA store)
|
||||
bool issue_tma_store = (thread_idx / NumThreadsPerWarp) == 0;
|
||||
@@ -594,9 +614,12 @@ public:
|
||||
for (int epi_m = 0; epi_m < size<2>(gD_epi); ++epi_m) {
|
||||
bool is_last_iteration = epi_m == size<2>(gD_epi)-1 && epi_n == size<3>(gD_epi)-1;
|
||||
|
||||
if (subtile_idx != -1 && (epi_n * static_cast<int>(size<2>(gD_epi)) + epi_m) != subtile_idx) {
|
||||
continue;
|
||||
}
|
||||
// The current tile in accumulator
|
||||
int mma_m = epi_m;
|
||||
int mma_n = (epi_n * epi_tile_n) / mma_tile_n;
|
||||
int mma_n = (epi_n * size<1>(EpilogueTile{})) / mma_tile_n;
|
||||
Tensor tRS_rAcc_frg_mn = tRS_rAcc_frg(_,mma_m,mma_n);
|
||||
|
||||
if (is_producer_load_needed) {
|
||||
|
||||
@@ -46,6 +46,8 @@ namespace cutlass::epilogue {
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct NoSmemWarpSpecialized {};
|
||||
struct NoSmemWarpSpecializedArray {};
|
||||
struct NoSmemWarpSpecializedGroup {};
|
||||
struct TmaWarpSpecialized {};
|
||||
struct TmaWarpSpecializedCooperative {};
|
||||
// DEPRECATED schedules, will be removed in next release
|
||||
|
||||
@@ -50,6 +50,8 @@ struct FusionOperation {
|
||||
// metadata types/queries that can be overrided
|
||||
using ElementOutput = void;
|
||||
using ElementCompute = void;
|
||||
|
||||
using ElementSource = void;
|
||||
static constexpr bool IsSourceSupported = false;
|
||||
|
||||
using ElementScalar = void;
|
||||
@@ -96,11 +98,13 @@ struct ScaledAcc : FusionOperation {
|
||||
template<
|
||||
class ElementOutput_,
|
||||
class ElementCompute_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
struct LinearCombination
|
||||
: ScaledAcc<ElementOutput_, ElementCompute_, ElementScalar_, RoundStyle_> {
|
||||
using ElementSource = ElementSource_;
|
||||
static constexpr bool IsSourceSupported = true;
|
||||
};
|
||||
|
||||
@@ -109,11 +113,12 @@ template<
|
||||
template <class> class ActivationFn_,
|
||||
class ElementOutput_,
|
||||
class ElementCompute_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
struct LinCombEltAct
|
||||
: LinearCombination<ElementOutput_, ElementCompute_, ElementScalar_, RoundStyle_> {
|
||||
: LinearCombination<ElementOutput_, ElementCompute_, ElementSource_, ElementScalar_, RoundStyle_> {
|
||||
using ActivationFn = ActivationFn_<ElementCompute_>;
|
||||
static constexpr bool IsEltActSupported = true;
|
||||
};
|
||||
@@ -123,12 +128,13 @@ template<
|
||||
class ElementOutput_,
|
||||
class ElementCompute_,
|
||||
class ElementBias_ = ElementOutput_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
|
||||
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
struct LinCombPerRowBias
|
||||
: LinearCombination<ElementOutput_, ElementCompute_, ElementScalar_, RoundStyle_> {
|
||||
: LinearCombination<ElementOutput_, ElementCompute_, ElementSource_, ElementScalar_, RoundStyle_> {
|
||||
using ElementBias = ElementBias_;
|
||||
static constexpr int AlignmentBias = AlignmentBias_;
|
||||
static constexpr bool IsPerRowBiasSupported = true;
|
||||
@@ -140,13 +146,14 @@ template<
|
||||
class ElementOutput_,
|
||||
class ElementCompute_,
|
||||
class ElementBias_ = ElementOutput_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
|
||||
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
struct LinCombPerRowBiasEltAct
|
||||
: LinCombPerRowBias<ElementOutput_, ElementCompute_,
|
||||
ElementBias_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
ElementBias_, ElementSource_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
using ActivationFn = ActivationFn_<ElementCompute_>;
|
||||
static constexpr bool IsEltActSupported = true;
|
||||
};
|
||||
@@ -160,6 +167,7 @@ template<
|
||||
class ElementCompute_,
|
||||
class ElementAux_ = ElementOutput_,
|
||||
class ElementBias_ = ElementOutput_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
int AlignmentAux_ = 128 / sizeof_bits_v<ElementAux_>,
|
||||
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
|
||||
@@ -167,7 +175,7 @@ template<
|
||||
>
|
||||
struct LinCombPerRowBiasEltActAux
|
||||
: LinCombPerRowBiasEltAct<ActivationFn_, ElementOutput_, ElementCompute_,
|
||||
ElementBias_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
ElementBias_, ElementSource_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
using ElementAux = ElementAux_;
|
||||
using GmemLayoutTagAux = GmemLayoutTagAux_;
|
||||
static constexpr int AlignmentAux = AlignmentAux_;
|
||||
@@ -180,6 +188,7 @@ template<
|
||||
class ElementOutput_,
|
||||
class ElementCompute_,
|
||||
class ElementBias_ = ElementOutput_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_, // per-row alpha/beta
|
||||
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
|
||||
int AlignmentScalar_ = 128 / sizeof_bits_v<ElementScalar_>,
|
||||
@@ -187,7 +196,7 @@ template<
|
||||
>
|
||||
struct PerRowLinCombPerRowBiasEltAct
|
||||
: LinCombPerRowBiasEltAct<ActivationFn_, ElementOutput_, ElementCompute_,
|
||||
ElementBias_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
ElementBias_, ElementSource_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
static constexpr int AlignmentScalar = AlignmentScalar_;
|
||||
static constexpr bool IsPerRowScaleSupported = true;
|
||||
};
|
||||
@@ -202,13 +211,14 @@ template<
|
||||
class ElementOutput_,
|
||||
class ElementCompute_,
|
||||
class ElementBias_ = ElementOutput_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
|
||||
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
struct ScaledLinCombPerRowBiasEltAct
|
||||
: LinCombPerRowBiasEltAct<ActivationFn_, ElementOutput_, ElementCompute_,
|
||||
ElementBias_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
ElementBias_, ElementSource_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
static constexpr bool IsScaleFactorSupported = true;
|
||||
};
|
||||
|
||||
@@ -231,6 +241,7 @@ template<
|
||||
class ElementAux_ = ElementOutput_,
|
||||
class ElementAmax_ = ElementCompute_,
|
||||
class ElementBias_ = ElementOutput_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
int AlignmentAux_ = 128 / sizeof_bits_v<ElementAux_>,
|
||||
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
|
||||
@@ -238,7 +249,7 @@ template<
|
||||
>
|
||||
struct ScaledLinCombPerRowBiasEltActAmaxAux
|
||||
: ScaledLinCombPerRowBiasEltAct<ActivationFn_, ElementOutput_, ElementCompute_,
|
||||
ElementBias_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
ElementBias_, ElementSource_, ElementScalar_, AlignmentBias_, RoundStyle_> {
|
||||
using ElementAmax = ElementAmax_;
|
||||
static constexpr bool IsAbsMaxSupported = true;
|
||||
|
||||
@@ -257,12 +268,16 @@ template<
|
||||
class ElementOutput_,
|
||||
class ElementCompute_,
|
||||
class ElementAux_ = ElementOutput_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
int AlignmentAux_ = 128 / sizeof_bits_v<ElementAux_>,
|
||||
FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
struct LinCombDeEltAct
|
||||
: LinCombEltAct<ActivationFn_, ElementOutput_, ElementCompute_, ElementScalar_, RoundStyle_> {
|
||||
: LinearCombination<ElementOutput_, ElementCompute_, ElementSource_, ElementScalar_, RoundStyle_> {
|
||||
using ActivationFn = ActivationFn_<ElementCompute_>;
|
||||
static constexpr bool IsDeEltActSupported = true;
|
||||
|
||||
using ElementAux = ElementAux_;
|
||||
using GmemLayoutTagAux = GmemLayoutTagAux_;
|
||||
static constexpr int AlignmentAux = AlignmentAux_;
|
||||
@@ -280,6 +295,7 @@ template<
|
||||
class ElementCompute_,
|
||||
class ElementAux_ = ElementOutput_,
|
||||
class ElementBias_ = ElementCompute_,
|
||||
class ElementSource_ = ElementOutput_,
|
||||
class ElementScalar_ = ElementCompute_,
|
||||
int AlignmentAux_ = 128 / sizeof_bits_v<ElementAux_>,
|
||||
int AlignmentBias_ = 128 / sizeof_bits_v<ElementBias_>,
|
||||
@@ -287,7 +303,7 @@ template<
|
||||
>
|
||||
struct LinCombDeEltActDePerRowBias
|
||||
: LinCombDeEltAct<GmemLayoutTagAux_, ActivationFn_, ElementOutput_, ElementCompute_,
|
||||
ElementAux_, ElementScalar_, AlignmentAux_, RoundStyle_> {
|
||||
ElementAux_, ElementSource_, ElementScalar_, AlignmentAux_, RoundStyle_> {
|
||||
using ElementBias = ElementBias_;
|
||||
static constexpr int AlignmentBias = AlignmentBias_;
|
||||
static constexpr bool IsDePerRowBiasSupported = true;
|
||||
|
||||
@@ -113,13 +113,14 @@ struct FusionCallbacks<
|
||||
template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
using Sm90LinearCombination =
|
||||
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc)
|
||||
Sm90ScalarBroadcast<ElementScalar>, // beta
|
||||
Sm90SrcFetch, // C
|
||||
Sm90SrcFetch<ElementSource>, // C
|
||||
Sm90EVT<Sm90Compute<multiplies, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc
|
||||
Sm90ScalarBroadcast<ElementScalar>, // alpha
|
||||
Sm90AccFetch // acc
|
||||
@@ -133,6 +134,7 @@ template <
|
||||
bool ReuseSmemC,
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
FloatRoundStyle RoundStyle,
|
||||
class CtaTileShapeMNK,
|
||||
@@ -140,13 +142,13 @@ template <
|
||||
>
|
||||
struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::LinearCombination<ElementOutput, ElementCompute, ElementScalar, RoundStyle>,
|
||||
fusion::LinearCombination<ElementOutput, ElementCompute, ElementSource, ElementScalar, RoundStyle>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile
|
||||
> : Sm90LinearCombination<typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementScalar, RoundStyle> {
|
||||
> : Sm90LinearCombination<typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementSource, ElementScalar, RoundStyle> {
|
||||
|
||||
using Impl = Sm90LinearCombination<typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementScalar, RoundStyle>;
|
||||
using Operation = fusion::LinearCombination<ElementOutput, ElementCompute, ElementScalar, RoundStyle>;
|
||||
using Impl = Sm90LinearCombination<typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementSource, ElementScalar, RoundStyle>;
|
||||
using Operation = fusion::LinearCombination<ElementOutput, ElementCompute, ElementSource, ElementScalar, RoundStyle>;
|
||||
|
||||
struct Arguments {
|
||||
ElementScalar alpha = ElementScalar(1);
|
||||
@@ -180,12 +182,13 @@ template<
|
||||
template <class> class ActivationFn,
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
using Sm90LinCombEltAct =
|
||||
Sm90EVT<Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>, // activation(beta * C + (alpha * acc))
|
||||
Sm90LinearCombination<ElementCompute, ElementCompute, ElementScalar, RoundStyle> // beta * C + (alpha * acc)
|
||||
Sm90LinearCombination<ElementCompute, ElementCompute, ElementSource, ElementScalar, RoundStyle> // beta * C + (alpha * acc)
|
||||
>;
|
||||
|
||||
template <
|
||||
@@ -196,6 +199,7 @@ template <
|
||||
template <class> class ActivationFn,
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
FloatRoundStyle RoundStyle,
|
||||
class CtaTileShapeMNK,
|
||||
@@ -203,13 +207,13 @@ template <
|
||||
>
|
||||
struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementScalar, RoundStyle>,
|
||||
fusion::LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementSource, ElementScalar, RoundStyle>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile
|
||||
> : Sm90LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementScalar, RoundStyle> {
|
||||
> : Sm90LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementSource, ElementScalar, RoundStyle> {
|
||||
|
||||
using Impl = Sm90LinCombEltAct<ActivationFn, typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementScalar, RoundStyle>;
|
||||
using Operation = fusion::LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementScalar, RoundStyle>;
|
||||
using Impl = Sm90LinCombEltAct<ActivationFn, typename cutlass::detail::get_unpacked_element_type<ElementOutput>::type, ElementCompute, ElementSource, ElementScalar, RoundStyle>;
|
||||
using Operation = fusion::LinCombEltAct<ActivationFn, ElementOutput, ElementCompute, ElementSource, ElementScalar, RoundStyle>;
|
||||
|
||||
struct Arguments {
|
||||
ElementScalar alpha = ElementScalar(1);
|
||||
@@ -250,6 +254,7 @@ template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
|
||||
@@ -257,7 +262,7 @@ template<
|
||||
using Sm90LinCombPerRowBias =
|
||||
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
|
||||
Sm90ScalarBroadcast<ElementScalar>, // beta
|
||||
Sm90SrcFetch, // C
|
||||
Sm90SrcFetch<ElementSource>, // C
|
||||
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
|
||||
Sm90ScalarBroadcast<ElementScalar>, // alpha
|
||||
Sm90AccFetch, // acc
|
||||
@@ -273,6 +278,7 @@ template <
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentBias,
|
||||
FloatRoundStyle RoundStyle,
|
||||
@@ -281,15 +287,15 @@ template <
|
||||
>
|
||||
struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::LinCombPerRowBias<ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle>,
|
||||
fusion::LinCombPerRowBias<ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile
|
||||
> : Sm90LinCombPerRowBias<
|
||||
CtaTileShapeMNK, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle> {
|
||||
CtaTileShapeMNK, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle> {
|
||||
using Impl = Sm90LinCombPerRowBias<
|
||||
CtaTileShapeMNK, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle>;
|
||||
CtaTileShapeMNK, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle>;
|
||||
using Operation = fusion::LinCombPerRowBias<
|
||||
ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle>;
|
||||
ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle>;
|
||||
|
||||
struct Arguments {
|
||||
ElementScalar alpha = ElementScalar(1);
|
||||
@@ -330,13 +336,14 @@ template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
using Sm90LinCombPerRowBiasEltAct =
|
||||
Sm90EVT<Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>,
|
||||
Sm90LinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle>
|
||||
Sm90LinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle>
|
||||
>;
|
||||
|
||||
template <
|
||||
@@ -348,6 +355,7 @@ template <
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentBias,
|
||||
FloatRoundStyle RoundStyle,
|
||||
@@ -357,21 +365,21 @@ template <
|
||||
struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::LinCombPerRowBiasEltAct<
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile
|
||||
> : Sm90LinCombPerRowBiasEltAct<
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
> {
|
||||
|
||||
using Impl =
|
||||
Sm90LinCombPerRowBiasEltAct<
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
>;
|
||||
using Operation =
|
||||
fusion::LinCombPerRowBiasEltAct<
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
>;
|
||||
|
||||
struct Arguments {
|
||||
@@ -426,6 +434,7 @@ template<
|
||||
class ElementCompute,
|
||||
class ElementAux = ElementOutput,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentAux = 128 / sizeof_bits_v<ElementAux>,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
@@ -434,7 +443,7 @@ template<
|
||||
using Sm90LinCombPerRowBiasEltActAux =
|
||||
Sm90EVT<Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>,
|
||||
Sm90EVT<Sm90AuxStore<Stages, EpilogueTile, ElementAux, RoundStyle, StrideAux, SmemLayoutAtom, CopyOpR2S, AlignmentAux>,
|
||||
Sm90LinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle>
|
||||
Sm90LinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle>
|
||||
>
|
||||
>;
|
||||
|
||||
@@ -449,6 +458,7 @@ template <
|
||||
class ElementCompute,
|
||||
class ElementAux,
|
||||
class ElementBias,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentAux,
|
||||
int AlignmentBias,
|
||||
@@ -462,7 +472,7 @@ struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::LinCombPerRowBiasEltActAux<
|
||||
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
|
||||
ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile,
|
||||
@@ -470,18 +480,18 @@ struct FusionCallbacks<
|
||||
CopyOpR2S
|
||||
> : Sm90LinCombPerRowBiasEltActAux<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesD, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpR2S, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
> {
|
||||
|
||||
using Impl =
|
||||
Sm90LinCombPerRowBiasEltActAux<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesD, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpR2S, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>;
|
||||
using Operation =
|
||||
fusion::LinCombPerRowBiasEltActAux<
|
||||
GmemLayoutTagAux, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>;
|
||||
|
||||
struct Arguments {
|
||||
@@ -535,6 +545,7 @@ template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
int AlignmentScalar = 128 / sizeof_bits_v<ElementScalar>,
|
||||
@@ -543,7 +554,7 @@ template<
|
||||
using Sm90PerRowLinCombPerRowBias =
|
||||
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
|
||||
Sm90ColBroadcast<0, CtaTileShapeMNK, ElementScalar, Stride<_1,_0,_0>, AlignmentScalar>, // beta
|
||||
Sm90SrcFetch, // C
|
||||
Sm90SrcFetch<ElementSource>, // C
|
||||
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
|
||||
Sm90ColBroadcast<0, CtaTileShapeMNK, ElementScalar, Stride<_1,_0,_0>, AlignmentScalar>, // alpha
|
||||
Sm90AccFetch, // acc
|
||||
@@ -558,6 +569,7 @@ template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
int AlignmentScalar = 128 / sizeof_bits_v<ElementScalar>,
|
||||
@@ -566,7 +578,7 @@ template<
|
||||
using Sm90PerRowLinCombPerRowBiasEltAct =
|
||||
Sm90EVT<Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>,
|
||||
Sm90PerRowLinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute,
|
||||
ElementBias, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle>
|
||||
ElementBias, ElementSource, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle>
|
||||
>;
|
||||
|
||||
template <
|
||||
@@ -578,6 +590,7 @@ template <
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentBias,
|
||||
int AlignmentScalar,
|
||||
@@ -588,21 +601,21 @@ template <
|
||||
struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile
|
||||
> : Sm90PerRowLinCombPerRowBiasEltAct<
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
> {
|
||||
|
||||
using Impl =
|
||||
Sm90PerRowLinCombPerRowBiasEltAct<
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
>;
|
||||
using Operation =
|
||||
fusion::PerRowLinCombPerRowBiasEltAct<
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, AlignmentScalar, RoundStyle
|
||||
>;
|
||||
|
||||
struct Arguments {
|
||||
@@ -664,6 +677,7 @@ template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
|
||||
@@ -671,7 +685,7 @@ template<
|
||||
using Sm90ScaledLinCombPerRowBias =
|
||||
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>, // beta * C + (alpha * acc + bias)
|
||||
Sm90ScalarBroadcast<ElementScalar, Stride<_0,_0,_0>, 2>, // scale_c * beta
|
||||
Sm90SrcFetch, // C
|
||||
Sm90SrcFetch<ElementSource>, // C
|
||||
Sm90EVT<Sm90Compute<homogeneous_multiply_add, ElementCompute, ElementCompute, RoundStyle>, // alpha * acc + bias
|
||||
Sm90ScalarBroadcast<ElementScalar, Stride<_0,_0,_0>, 3>, // scale_a * scale_b * alpha
|
||||
Sm90AccFetch, // acc
|
||||
@@ -690,6 +704,7 @@ template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
|
||||
@@ -698,7 +713,7 @@ using Sm90ScaledLinCombPerRowBiasEltAct =
|
||||
Sm90EVT<Sm90Compute<detail::ScaleOutOp<ElementOutput>::template Op, ElementOutput, ElementCompute, RoundStyle>, // activation(Z) * scale_d
|
||||
Sm90EVT<Sm90Compute<ActivationFn, ElementCompute, ElementCompute, RoundStyle>, // activation(Z)
|
||||
// Z = scale_a * scale_b * alpha * acc + beta * scale_c * C + per-row bias
|
||||
Sm90ScaledLinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle>
|
||||
Sm90ScaledLinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle>
|
||||
>,
|
||||
Sm90ScalarBroadcast<ElementScalar> // scale_d
|
||||
>;
|
||||
@@ -712,6 +727,7 @@ template <
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementBias,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentBias,
|
||||
FloatRoundStyle RoundStyle,
|
||||
@@ -721,21 +737,21 @@ template <
|
||||
struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::ScaledLinCombPerRowBiasEltAct<
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile
|
||||
> : Sm90ScaledLinCombPerRowBiasEltAct<
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
> {
|
||||
|
||||
using Impl =
|
||||
Sm90ScaledLinCombPerRowBiasEltAct<
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
CtaTileShapeMNK, ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
>;
|
||||
using Operation =
|
||||
fusion::ScaledLinCombPerRowBiasEltAct<
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle
|
||||
ActivationFn, ElementOutput, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle
|
||||
>;
|
||||
|
||||
struct Arguments {
|
||||
@@ -819,6 +835,7 @@ template<
|
||||
class ElementAux = ElementOutput,
|
||||
class ElementAmax = ElementCompute,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentAux = 128 / sizeof_bits_v<ElementAux>,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
@@ -827,7 +844,7 @@ template<
|
||||
using Sm90ScaledLinCombPerRowBiasEltActAmaxAux =
|
||||
Sm90SplitTreeVisitor<
|
||||
// Z = scale_a * scale_b * alpha * acc + scale_c * beta * C + per-row bias
|
||||
Sm90ScaledLinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementScalar, AlignmentBias, RoundStyle>,
|
||||
Sm90ScaledLinCombPerRowBias<CtaTileShapeMNK, ElementCompute, ElementCompute, ElementBias, ElementSource, ElementScalar, AlignmentBias, RoundStyle>,
|
||||
// D = activation(Z) * scale_d, amax_d = max(abs(elements in D))
|
||||
Sm90EVT<Sm90Compute<detail::ScaleOutOp<ElementOutput>::template Op, ElementOutput, ElementCompute, RoundStyle>, // activation(Z) * scale_d
|
||||
Sm90EVT<Sm90ScalarReduction<detail::amax, atomic_maximum, ElementAmax, ElementCompute, RoundStyle>, // amax_d
|
||||
@@ -860,6 +877,7 @@ template <
|
||||
class ElementAux,
|
||||
class ElementAmax,
|
||||
class ElementBias,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentAux,
|
||||
int AlignmentBias,
|
||||
@@ -873,7 +891,7 @@ struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::ScaledLinCombPerRowBiasEltActAmaxAux<
|
||||
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
|
||||
ElementAux, ElementAmax, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementAux, ElementAmax, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile,
|
||||
@@ -882,19 +900,19 @@ struct FusionCallbacks<
|
||||
> : Sm90ScaledLinCombPerRowBiasEltActAmaxAux<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesD, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>,
|
||||
SmemLayoutAtom, CopyOpR2S, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementAmax, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementAmax, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
> {
|
||||
|
||||
using Impl =
|
||||
Sm90ScaledLinCombPerRowBiasEltActAmaxAux<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesD, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>,
|
||||
SmemLayoutAtom, CopyOpR2S, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementAmax, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementAmax, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>;
|
||||
using Operation =
|
||||
fusion::ScaledLinCombPerRowBiasEltActAmaxAux<
|
||||
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
|
||||
ElementAux, ElementAmax, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementAux, ElementAmax, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>;
|
||||
|
||||
struct Arguments {
|
||||
@@ -1014,13 +1032,14 @@ template<
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementAux = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentAux = 128 / sizeof_bits_v<ElementAux>,
|
||||
FloatRoundStyle RoundStyle = FloatRoundStyle::round_to_nearest
|
||||
>
|
||||
using Sm90LinCombDeEltAct =
|
||||
Sm90EVT<Sm90Compute<ActivationFn, ElementOutput, ElementCompute, RoundStyle>, // activation(beta * C + (alpha * acc), aux)
|
||||
Sm90LinearCombination<ElementCompute, ElementCompute, ElementScalar, RoundStyle>, // beta * C + (alpha * acc)
|
||||
Sm90LinearCombination<ElementCompute, ElementCompute, ElementSource, ElementScalar, RoundStyle>, // beta * C + (alpha * acc)
|
||||
Sm90AuxLoad<Stages, EpilogueTile, ElementAux, StrideAux, SmemLayoutAtom, CopyOpS2R, AlignmentAux> // aux
|
||||
>;
|
||||
|
||||
@@ -1034,6 +1053,7 @@ template <
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
class ElementAux,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentAux,
|
||||
FloatRoundStyle RoundStyle,
|
||||
@@ -1046,7 +1066,7 @@ struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::LinCombDeEltAct<
|
||||
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
|
||||
ElementAux, ElementScalar, AlignmentAux, RoundStyle
|
||||
ElementAux, ElementSource, ElementScalar, AlignmentAux, RoundStyle
|
||||
>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile,
|
||||
@@ -1054,18 +1074,18 @@ struct FusionCallbacks<
|
||||
CopyOpS2R
|
||||
> : Sm90LinCombDeEltAct<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementScalar, AlignmentAux, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementSource, ElementScalar, AlignmentAux, RoundStyle
|
||||
> {
|
||||
|
||||
using Impl =
|
||||
Sm90LinCombDeEltAct<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementScalar, AlignmentAux, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementSource, ElementScalar, AlignmentAux, RoundStyle
|
||||
>;
|
||||
using Operation =
|
||||
fusion::LinCombDeEltAct<
|
||||
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
|
||||
ElementAux, ElementScalar, AlignmentAux, RoundStyle
|
||||
ElementAux, ElementSource, ElementScalar, AlignmentAux, RoundStyle
|
||||
>;
|
||||
|
||||
struct Arguments {
|
||||
@@ -1118,6 +1138,7 @@ template<
|
||||
class ElementCompute,
|
||||
class ElementAux = ElementOutput,
|
||||
class ElementBias = ElementOutput,
|
||||
class ElementSource = ElementOutput,
|
||||
class ElementScalar = ElementCompute,
|
||||
int AlignmentAux = 128 / sizeof_bits_v<ElementAux>,
|
||||
int AlignmentBias = 128 / sizeof_bits_v<ElementBias>,
|
||||
@@ -1128,7 +1149,7 @@ using Sm90LinCombDeEltActDePerRowBias =
|
||||
Sm90EVT<Sm90ColReduction<plus, plus, 0, CtaTileShapeMNK,
|
||||
ElementBias, ElementCompute, RoundStyle, Stride<_1,_0,int>, AlignmentBias>,
|
||||
Sm90LinCombDeEltAct<CtaTileShapeMNK, EpilogueTile, Stages, StrideAux, SmemLayoutAtom, CopyOpS2R, ActivationFn,
|
||||
ElementCompute, ElementCompute, ElementAux, ElementScalar, AlignmentAux, RoundStyle>
|
||||
ElementCompute, ElementCompute, ElementAux, ElementSource, ElementScalar, AlignmentAux, RoundStyle>
|
||||
>
|
||||
>;
|
||||
|
||||
@@ -1143,6 +1164,7 @@ template <
|
||||
class ElementCompute,
|
||||
class ElementAux,
|
||||
class ElementBias,
|
||||
class ElementSource,
|
||||
class ElementScalar,
|
||||
int AlignmentAux,
|
||||
int AlignmentBias,
|
||||
@@ -1156,7 +1178,7 @@ struct FusionCallbacks<
|
||||
epilogue::Sm90TmaWarpSpecialized<StagesC, StagesD, FragmentSize, ReuseSmemC>,
|
||||
fusion::LinCombDeEltActDePerRowBias<
|
||||
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
|
||||
ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>,
|
||||
CtaTileShapeMNK,
|
||||
EpilogueTile,
|
||||
@@ -1164,18 +1186,18 @@ struct FusionCallbacks<
|
||||
CopyOpS2R
|
||||
> : Sm90LinCombDeEltActDePerRowBias<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
> {
|
||||
|
||||
using Impl =
|
||||
Sm90LinCombDeEltActDePerRowBias<
|
||||
CtaTileShapeMNK, EpilogueTile, StagesC, cutlass::gemm::TagToStrideC_t<GmemLayoutTagAux>, SmemLayoutAtom, CopyOpS2R, ActivationFn,
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementOutput, ElementCompute, ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>;
|
||||
using Operation =
|
||||
fusion::LinCombDeEltActDePerRowBias<
|
||||
GmemLayoutTagAux, ActivationFn, ElementOutput, ElementCompute,
|
||||
ElementAux, ElementBias, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
ElementAux, ElementBias, ElementSource, ElementScalar, AlignmentAux, AlignmentBias, RoundStyle
|
||||
>;
|
||||
|
||||
struct Arguments {
|
||||
|
||||
@@ -215,16 +215,17 @@ template <
|
||||
class StrideScalar,
|
||||
int ScalarCount,
|
||||
template <class> class ScalarReduceFn,
|
||||
class ElementSource,
|
||||
class InputAddOp // Z
|
||||
>
|
||||
struct Sm90TreeVisitor<
|
||||
Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>,
|
||||
Sm90ScalarBroadcast<ElementScalar, StrideScalar, ScalarCount, ScalarReduceFn>,
|
||||
Sm90SrcFetch,
|
||||
Sm90SrcFetch<ElementSource>,
|
||||
InputAddOp
|
||||
> : Sm90VisitorImpl<
|
||||
Sm90ScalarBroadcast<ElementScalar, StrideScalar, ScalarCount, ScalarReduceFn>,
|
||||
Sm90SrcFetch,
|
||||
Sm90SrcFetch<ElementSource>,
|
||||
InputAddOp,
|
||||
Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>
|
||||
>
|
||||
@@ -232,11 +233,10 @@ struct Sm90TreeVisitor<
|
||||
using Impl =
|
||||
Sm90VisitorImpl<
|
||||
Sm90ScalarBroadcast<ElementScalar, StrideScalar, ScalarCount, ScalarReduceFn>,
|
||||
Sm90SrcFetch,
|
||||
Sm90SrcFetch<ElementSource>,
|
||||
InputAddOp,
|
||||
Sm90Compute<homogeneous_multiply_add, ElementOutput, ElementCompute, RoundStyle>
|
||||
>;
|
||||
|
||||
using Params = typename Impl::Params;
|
||||
using SharedStorage = typename Impl::SharedStorage;
|
||||
|
||||
@@ -260,8 +260,9 @@ struct Sm90TreeVisitor<
|
||||
CUTLASS_DEVICE bool
|
||||
is_C_load_needed() const {
|
||||
auto const& bcast_op = get<0>(Impl::ops);
|
||||
auto const& src_op = get<1>(Impl::ops);
|
||||
auto const& added_op = get<2>(Impl::ops);
|
||||
return bcast_op.scalar != 0 || added_op.is_C_load_needed();
|
||||
return (bcast_op.scalar != 0 && src_op.is_C_load_needed()) || added_op.is_C_load_needed();
|
||||
}
|
||||
|
||||
template <class CallbacksImpl>
|
||||
@@ -320,17 +321,9 @@ struct Sm90TreeVisitor<
|
||||
// ReLU with aux bit tensor dReLU/dZ
|
||||
// Aux(i) = Z(i) >= 0 ? 1 : 0
|
||||
namespace detail {
|
||||
template <
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
FloatRoundStyle RoundStyle,
|
||||
class StrideMNL,
|
||||
int Alignment,
|
||||
bool EnableNullptr
|
||||
>
|
||||
struct Sm90ReLUAuxStore {
|
||||
static_assert(Alignment % 128 == 0, "sub-16B alignment not supported yet");
|
||||
|
||||
// Placeholder node so we can retain standard EVT structure
|
||||
template <class StrideMNL>
|
||||
struct Sm90ReLUAuxStore : Sm90VisitorImpl<> {
|
||||
struct SharedStorage {};
|
||||
|
||||
struct Arguments {
|
||||
@@ -362,41 +355,90 @@ struct Sm90ReLUAuxStore {
|
||||
Sm90ReLUAuxStore() { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Sm90ReLUAuxStore(Params const& params, SharedStorage const& shared_storage)
|
||||
: params(params) { }
|
||||
Sm90ReLUAuxStore(Params const& params, SharedStorage const& shared_storage) { }
|
||||
};
|
||||
} // namespace detail
|
||||
|
||||
Params const params;
|
||||
// Specialization on the generic compute+aux EVT
|
||||
template <
|
||||
// Compute node
|
||||
template <class> class Activation,
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
FloatRoundStyle RoundStyle,
|
||||
// Aux node
|
||||
int Stages,
|
||||
class EpilogueTile,
|
||||
class StrideMNL,
|
||||
class SmemLayoutAtom,
|
||||
class CopyOpR2S,
|
||||
int Alignment,
|
||||
bool EnableNullptr,
|
||||
// Input node
|
||||
class InputOp
|
||||
>
|
||||
struct Sm90TreeVisitor<
|
||||
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle,
|
||||
enable_if_t<is_same_v<Activation<ElementCompute>, cutlass::epilogue::thread::ReLu<ElementCompute>> ||
|
||||
is_same_v<Activation<ElementCompute>, cutlass::epilogue::thread::Clamp<ElementCompute>> >>,
|
||||
Sm90TreeVisitor<
|
||||
Sm90AuxStore<
|
||||
Stages,
|
||||
EpilogueTile,
|
||||
cutlass::uint1b_t,
|
||||
RoundStyle,
|
||||
StrideMNL,
|
||||
SmemLayoutAtom,
|
||||
CopyOpR2S,
|
||||
Alignment,
|
||||
EnableNullptr
|
||||
>,
|
||||
InputOp
|
||||
>
|
||||
> : Sm90VisitorImpl<
|
||||
Sm90VisitorImpl<
|
||||
InputOp,
|
||||
detail::Sm90ReLUAuxStore<StrideMNL>
|
||||
>,
|
||||
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle>
|
||||
>
|
||||
{
|
||||
using Impl =
|
||||
Sm90VisitorImpl<
|
||||
Sm90VisitorImpl<
|
||||
InputOp,
|
||||
detail::Sm90ReLUAuxStore<StrideMNL>
|
||||
>,
|
||||
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle>
|
||||
>;
|
||||
using Params = typename Impl::Params;
|
||||
using SharedStorage = typename Impl::SharedStorage;
|
||||
|
||||
CUTLASS_DEVICE bool
|
||||
is_producer_load_needed() const {
|
||||
return false;
|
||||
}
|
||||
CUTLASS_HOST_DEVICE
|
||||
Sm90TreeVisitor() {}
|
||||
|
||||
CUTLASS_DEVICE bool
|
||||
is_C_load_needed() const {
|
||||
return false;
|
||||
}
|
||||
CUTLASS_HOST_DEVICE
|
||||
Sm90TreeVisitor(Params const& params_, SharedStorage const& shared_storage)
|
||||
: params(params_), Impl(params_, shared_storage) {}
|
||||
|
||||
template <class... Args>
|
||||
CUTLASS_DEVICE auto
|
||||
get_producer_load_callbacks(ProducerLoadArgs<Args...> const& args) {
|
||||
return EmptyProducerLoadCallbacks{};
|
||||
}
|
||||
Params const& params;
|
||||
|
||||
template <class RTensor, class GTensor, class CTensor, class ResidueMN>
|
||||
struct ConsumerStoreCallbacks : EmptyConsumerStoreCallbacks {
|
||||
template <class RTensor, class GTensor, class CTensor, class ResidueMN, class CallbacksImpl>
|
||||
struct ConsumerStoreCallbacks : CallbacksImpl {
|
||||
CUTLASS_DEVICE
|
||||
ConsumerStoreCallbacks(
|
||||
RTensor&& tC_rAux,
|
||||
GTensor&& tC_gAux,
|
||||
CTensor tC_cAux,
|
||||
ResidueMN residue_mn,
|
||||
Params const& params)
|
||||
Params const& params,
|
||||
CallbacksImpl&& impl)
|
||||
: tC_rAux(cute::forward<RTensor>(tC_rAux)),
|
||||
tC_gAux(cute::forward<GTensor>(tC_gAux)),
|
||||
tC_cAux(tC_cAux),
|
||||
residue_mn(residue_mn),
|
||||
params(params) {}
|
||||
params(params),
|
||||
CallbacksImpl(cute::forward<CallbacksImpl>(impl)) {}
|
||||
|
||||
RTensor tC_rAux; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
|
||||
GTensor tC_gAux; // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
|
||||
@@ -404,13 +446,23 @@ struct Sm90ReLUAuxStore {
|
||||
ResidueMN residue_mn;
|
||||
Params const& params;
|
||||
|
||||
template <typename ElementAccumulator, typename ElementInput, int FragmentSize>
|
||||
CUTLASS_DEVICE auto
|
||||
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n,
|
||||
Array<ElementInput, FragmentSize> const& frg_input) {
|
||||
template <typename ElementAccumulator, int FragmentSize>
|
||||
CUTLASS_DEVICE Array<ElementOutput, FragmentSize>
|
||||
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n) {
|
||||
// Unpack callbacks + params
|
||||
auto& [callbacks_input_aux, callbacks_compute] = CallbacksImpl::callbacks_tuple;
|
||||
auto& [callbacks_input, callbacks_aux] = callbacks_input_aux.callbacks_tuple;
|
||||
auto const& [params_input_aux, params_compute] = params;
|
||||
auto const& [params_input, params_aux] = params_input_aux;
|
||||
|
||||
// Visit the input node
|
||||
Array frg_input = callbacks_input.visit(frg_acc, epi_v, epi_m, epi_n);
|
||||
|
||||
// Compute activation + aux
|
||||
using ElementInput = typename decltype(frg_input)::Element;
|
||||
using ConvertInput = NumericArrayConverter<ElementCompute, ElementInput, FragmentSize, RoundStyle>;
|
||||
using ConvertAux = PackPredicates<FragmentSize>;
|
||||
using ComputeOutput = cutlass::epilogue::thread::ReLu<ElementCompute>;
|
||||
using ComputeOutput = Activation<ElementCompute>;
|
||||
using ConvertOutput = NumericArrayConverter<ElementOutput, ElementCompute, FragmentSize, RoundStyle>;
|
||||
ConvertInput convert_input{};
|
||||
ComputeOutput relu{};
|
||||
@@ -422,7 +474,12 @@ struct Sm90ReLUAuxStore {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < FragmentSize; ++i) {
|
||||
ElementCompute pre_relu = frg_compute[i];
|
||||
frg_compute[i] = relu(frg_compute[i]);
|
||||
if constexpr (is_same_v<Activation<ElementCompute>, cutlass::epilogue::thread::Clamp<ElementCompute>>) {
|
||||
frg_compute[i] = relu(frg_compute[i], params_compute);
|
||||
}
|
||||
else {
|
||||
frg_compute[i] = relu(frg_compute[i]);
|
||||
}
|
||||
frg_aux[i] = frg_compute[i] == pre_relu;
|
||||
}
|
||||
|
||||
@@ -435,8 +492,18 @@ struct Sm90ReLUAuxStore {
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
end() {
|
||||
// Unpack callbacks + params
|
||||
auto& [callbacks_input_aux, callbacks_compute] = CallbacksImpl::callbacks_tuple;
|
||||
auto& [callbacks_input, callbacks_aux] = callbacks_input_aux.callbacks_tuple;
|
||||
auto const& [params_input_aux, params_compute] = params;
|
||||
auto const& [params_input, params_aux] = params_input_aux;
|
||||
|
||||
// Visit the input node
|
||||
callbacks_input.end();
|
||||
|
||||
// Nullptr is no-op
|
||||
if constexpr (EnableNullptr) {
|
||||
if (params.ptr_aux == nullptr) {
|
||||
if (params_aux.ptr_aux == nullptr) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
@@ -473,114 +540,25 @@ struct Sm90ReLUAuxStore {
|
||||
>
|
||||
CUTLASS_DEVICE auto
|
||||
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
|
||||
// Unpack params
|
||||
auto const& [params_input_aux, params_compute] = params;
|
||||
auto const& [params_input, params_aux] = params_input_aux;
|
||||
|
||||
auto [M, N, K, L] = args.problem_shape_mnkl;
|
||||
auto [m, n, k, l] = args.tile_coord_mnkl;
|
||||
gmem_ptr ptr_aux = make_gmem_ptr(subbyte_iterator<cutlass::uint1b_t>(params.ptr_aux));
|
||||
Tensor mAux = make_tensor(ptr_aux, make_layout(make_shape(M,N,L), params.dAux)); // (M,N,L)
|
||||
gmem_ptr ptr_aux = make_gmem_ptr(subbyte_iterator<cutlass::uint1b_t>(params_aux.ptr_aux));
|
||||
Tensor mAux = make_tensor(ptr_aux, make_layout(make_shape(M,N,L), params_aux.dAux)); // (M,N,L)
|
||||
Tensor gAux = local_tile(mAux, take<0,2>(args.tile_shape_mnk), make_coord(m,n,l)); // (CTA_M,CTA_N)
|
||||
|
||||
Tensor tC_gAux = sm90_partition_for_epilogue<ReferenceSrc>( // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
|
||||
gAux, args.epi_tile, args.tiled_copy, args.thread_idx);
|
||||
Tensor tC_rAux = make_tensor<cutlass::uint1b_t>(shape(tC_gAux)); // (CPY,CPY_M,CPY_N,EPI_M,EPI_N)
|
||||
|
||||
return ConsumerStoreCallbacks<decltype(tC_rAux), decltype(tC_gAux), decltype(args.tCcD), decltype(args.residue_mn)>(
|
||||
cute::move(tC_rAux), cute::move(tC_gAux), args.tCcD, args.residue_mn, params);
|
||||
auto callbacks_impl = Impl::template get_consumer_store_callbacks<ReferenceSrc>(args);
|
||||
return ConsumerStoreCallbacks<decltype(tC_rAux), decltype(tC_gAux), decltype(args.tCcD), decltype(args.residue_mn), decltype(callbacks_impl)>(
|
||||
cute::move(tC_rAux), cute::move(tC_gAux), args.tCcD, args.residue_mn, params, cute::move(callbacks_impl));
|
||||
}
|
||||
};
|
||||
} // namespace detail
|
||||
|
||||
// Specialization on the generic compute+aux EVT
|
||||
template <
|
||||
// Compute node
|
||||
template <class> class Activation,
|
||||
class ElementOutput,
|
||||
class ElementCompute,
|
||||
FloatRoundStyle RoundStyle,
|
||||
// Aux node
|
||||
int Stages,
|
||||
class EpilogueTile,
|
||||
class StrideMNL,
|
||||
class SmemLayoutAtom,
|
||||
class CopyOpR2S,
|
||||
int Alignment,
|
||||
bool EnableNullptr,
|
||||
// Input node
|
||||
class InputOp
|
||||
>
|
||||
struct Sm90TreeVisitor<
|
||||
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle,
|
||||
enable_if_t<is_same_v<Activation<ElementCompute>, cutlass::epilogue::thread::ReLu<ElementCompute>>, void>>,
|
||||
Sm90TreeVisitor<
|
||||
Sm90AuxStore<
|
||||
Stages,
|
||||
EpilogueTile,
|
||||
cutlass::uint1b_t,
|
||||
RoundStyle,
|
||||
StrideMNL,
|
||||
SmemLayoutAtom,
|
||||
CopyOpR2S,
|
||||
Alignment,
|
||||
EnableNullptr
|
||||
>,
|
||||
InputOp
|
||||
>
|
||||
> : Sm90VisitorImpl<
|
||||
Sm90VisitorImpl<
|
||||
InputOp,
|
||||
detail::Sm90ReLUAuxStore<ElementOutput, ElementCompute, RoundStyle, StrideMNL, Alignment, EnableNullptr>
|
||||
>,
|
||||
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle>
|
||||
>
|
||||
{
|
||||
using Impl =
|
||||
Sm90VisitorImpl<
|
||||
Sm90VisitorImpl<
|
||||
InputOp,
|
||||
detail::Sm90ReLUAuxStore<ElementOutput, ElementCompute, RoundStyle, StrideMNL, Alignment, EnableNullptr>
|
||||
>,
|
||||
Sm90Compute<Activation, ElementOutput, ElementCompute, RoundStyle>
|
||||
>;
|
||||
|
||||
using Params = typename Impl::Params;
|
||||
using SharedStorage = typename Impl::SharedStorage;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Sm90TreeVisitor() {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Sm90TreeVisitor(
|
||||
Params const& params,
|
||||
SharedStorage const& shared_storage)
|
||||
: Impl(params, shared_storage) {}
|
||||
|
||||
template <class CallbacksImpl>
|
||||
struct ConsumerStoreCallbacks : CallbacksImpl {
|
||||
CUTLASS_DEVICE
|
||||
ConsumerStoreCallbacks(CallbacksImpl&& impl)
|
||||
: CallbacksImpl(cute::forward<CallbacksImpl>(impl)) { }
|
||||
|
||||
template <typename ElementAccumulator, int FragmentSize>
|
||||
CUTLASS_DEVICE Array<ElementOutput, FragmentSize>
|
||||
visit(Array<ElementAccumulator, FragmentSize> const& frg_acc, int epi_v, int epi_m, int epi_n) {
|
||||
auto& [callbacks_input, callbacks_relu_aux] = get<0>(CallbacksImpl::callbacks_tuple).callbacks_tuple;
|
||||
|
||||
Array frg_input = callbacks_input.visit(frg_acc, epi_v, epi_m, epi_n);
|
||||
return callbacks_relu_aux.visit(frg_acc, epi_v, epi_m, epi_n, frg_input);
|
||||
}
|
||||
};
|
||||
|
||||
template <
|
||||
bool ReferenceSrc, // do register tensors reference the src or dst layout of the tiled copy
|
||||
class... Args
|
||||
>
|
||||
CUTLASS_DEVICE auto
|
||||
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
|
||||
auto callbacks_tuple = Impl::template get_consumer_store_callbacks<ReferenceSrc>(args);
|
||||
return ConsumerStoreCallbacks<decltype(callbacks_tuple)>(std::move(callbacks_tuple));
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
// Aux load for uint1b_t
|
||||
template <
|
||||
|
||||
@@ -85,16 +85,17 @@ using Sm90SplitTreeFetch = Sm90AccFetch;
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// returns C
|
||||
template <class Element>
|
||||
struct Sm90SrcFetch : Sm90VisitorImpl<> {
|
||||
|
||||
CUTLASS_DEVICE bool
|
||||
is_producer_load_needed() const {
|
||||
return true;
|
||||
return is_C_load_needed();
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE bool
|
||||
is_C_load_needed() const {
|
||||
return true;
|
||||
return not is_void_v<Element>;
|
||||
}
|
||||
|
||||
using Sm90VisitorImpl<>::Sm90VisitorImpl;
|
||||
@@ -105,7 +106,6 @@ struct Sm90SrcFetch : Sm90VisitorImpl<> {
|
||||
ConsumerStoreCallbacks(SrcTensor const& tCrC)
|
||||
: tCrC(tCrC) {}
|
||||
|
||||
// make this a pointer if we need default ctor for generic tuple of visitors
|
||||
SrcTensor const& tCrC; // (CPY,CPY_M,CPY_N)
|
||||
|
||||
template <typename ElementAccumulator, int FragmentSize>
|
||||
@@ -122,7 +122,7 @@ struct Sm90SrcFetch : Sm90VisitorImpl<> {
|
||||
>
|
||||
CUTLASS_DEVICE auto
|
||||
get_consumer_store_callbacks(ConsumerStoreArgs<Args...> const& args) {
|
||||
|
||||
// register type may differ from logical type so we can't assert matching types here
|
||||
return ConsumerStoreCallbacks(args.tCrC);
|
||||
}
|
||||
};
|
||||
@@ -344,7 +344,6 @@ struct Sm90AuxLoad {
|
||||
make_tensor(make_smem_ptr(smem_aux), SmemLayout{})); // (EPI_TILE_M,EPI_TILE_N,PIPE)
|
||||
auto tSR_sAux = tiled_s2r.get_slice(args.thread_idx).partition_S(sAux_epi); // (S2R,S2R_M,S2R_N,PIPE)
|
||||
|
||||
|
||||
return ConsumerStoreCallbacks<decltype(tC_rAux), decltype(tiled_s2r), decltype(tSR_sAux)>(
|
||||
cute::move(tC_rAux), tiled_s2r, cute::move(tSR_sAux), params_ptr);
|
||||
}
|
||||
|
||||
@@ -303,7 +303,7 @@ private:
|
||||
static constexpr bool IsAtomic = is_atomic<GmemReduceFn<ElementCompute>>::value;
|
||||
static_assert(IsAtomic, "non-atomic scalar reduction not supported yet");
|
||||
|
||||
public:
|
||||
public:
|
||||
struct SharedStorage { };
|
||||
|
||||
struct Arguments {
|
||||
@@ -814,7 +814,7 @@ public:
|
||||
}
|
||||
|
||||
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
|
||||
Tensor tCrCol_mn = tCrCol(_,_,_,epi_m,epi_n);
|
||||
Tensor tCcCol_mn = tCcCol(_,_,_,epi_m,epi_n);
|
||||
@@ -843,8 +843,8 @@ public:
|
||||
return;
|
||||
}
|
||||
|
||||
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
|
||||
auto [m, n, k, l] = tile_coord_mnkl;
|
||||
constexpr bool ReferenceSrc = decltype(ref_src)::value;
|
||||
@@ -1002,15 +1002,15 @@ public:
|
||||
return;
|
||||
}
|
||||
|
||||
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
auto& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
|
||||
|
||||
using ReduceOutput = GmemReduceFn<ElementCompute>;
|
||||
using ConvertOutput = NumericConverter<ElementOutput, ElementCompute, RoundStyle>;
|
||||
ReduceOutput reduce_output{};
|
||||
ConvertOutput convert_output{};
|
||||
|
||||
|
||||
// Reduction over batches
|
||||
if (size<2>(stride(gCol_l)) == 0) {
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
@@ -1051,8 +1051,8 @@ public:
|
||||
|
||||
CUTLASS_DEVICE bool
|
||||
is_reduction_buffer_needed(int epi_m, int epi_n, bool is_last_iteration) const {
|
||||
auto const& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
auto const& [ref_src, tCrCol, tCcCol, gCol_l, cCol, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
tile_coord_mnkl, residue_mn, epi_tile, tiled_copy, thread_idx] = args_tuple;
|
||||
|
||||
return (not IsAtomic && // atomic reduction doesn't use smem
|
||||
@@ -1111,7 +1111,7 @@ public:
|
||||
|
||||
auto args_tuple = make_tuple(
|
||||
bool_constant<ReferenceSrc>{}, cute::move(tCrCol), args.tCcD, gCol_l, args.cD, gBuf_nl, sBuf_layout,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
lane_layout_MN, lane_mn, warp_layout_MN, warp_mn,
|
||||
args.tile_coord_mnkl, args.residue_mn, args.epi_tile, args.tiled_copy, args.thread_idx);
|
||||
return ConsumerStoreCallbacks<decltype(args_tuple)>(std::move(args_tuple), params);
|
||||
}
|
||||
|
||||
@@ -563,7 +563,6 @@ struct Sm90TreeVisitor : Sm90VisitorImpl<ChildOps..., NodeOp> {
|
||||
template get_consumer_store_callbacks<ReferenceSrc>(args);
|
||||
return ConsumerStoreCallbacks<decltype(callbacks_tuple)>(std::move(callbacks_tuple));
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -614,7 +613,6 @@ struct Sm90SplitTreeVisitor : Sm90VisitorImpl<InputTree, AuxOutTrees..., OutputT
|
||||
template get_consumer_store_callbacks<ReferenceSrc>(args);
|
||||
return ConsumerStoreCallbacks<decltype(callbacks_tuple)>(std::move(callbacks_tuple));
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -692,7 +690,6 @@ struct Sm90TopologicalVisitor : Sm90VisitorImpl<Ops...> {
|
||||
template get_consumer_store_callbacks<ReferenceSrc>(args);
|
||||
return ConsumerStoreCallbacks<decltype(callbacks_tuple)>(std::move(callbacks_tuple));
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -49,7 +49,6 @@ namespace cutlass {
|
||||
namespace epilogue {
|
||||
namespace thread {
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
@@ -66,13 +65,65 @@ struct ArrayMaximum {
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < ElementsPerAccess; ++i) {
|
||||
result[i] = fmax(lhs[i], rhs[i]);
|
||||
result[i] = platform::max(lhs[i].get(), rhs[i]);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<Element, ElementsPerAccess> operator()(
|
||||
Array<Element, ElementsPerAccess> const &lhs,
|
||||
Element rhs) const {
|
||||
|
||||
Array<Element, ElementsPerAccess> result;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < ElementsPerAccess; ++i) {
|
||||
result[i] = platform::max(lhs[i].get(), rhs);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
/// Partial specialization: Element=float
|
||||
template <int ElementsPerAccess>
|
||||
struct ArrayMaximum<float, ElementsPerAccess> {
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<float, ElementsPerAccess> operator()(
|
||||
Array<float, ElementsPerAccess> const &lhs,
|
||||
Array<float, ElementsPerAccess> const &rhs) const {
|
||||
|
||||
Array<float, ElementsPerAccess> result;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < ElementsPerAccess; ++i) {
|
||||
result[i] = fmax(lhs[i], rhs[i]);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<float, ElementsPerAccess> operator()(
|
||||
Array<float, ElementsPerAccess> const &lhs,
|
||||
float rhs) const {
|
||||
|
||||
Array<float, ElementsPerAccess> result;
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < ElementsPerAccess; ++i) {
|
||||
result[i] = fmax(lhs[i], rhs);
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization: Element=half
|
||||
template <int ElementsPerAccess>
|
||||
struct ArrayMaximum<half_t, ElementsPerAccess> {
|
||||
|
||||
@@ -96,6 +147,8 @@ struct ArrayMaximum<half_t, ElementsPerAccess> {
|
||||
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_ptr[i]);
|
||||
}
|
||||
|
||||
static_assert(!(ElementsPerAccess % 2), "Output array must be divisible by vector length.");
|
||||
|
||||
#else
|
||||
__half const *lhs_ptr = reinterpret_cast<__half const *>(lhs.raw_data());
|
||||
__half const *rhs_ptr = reinterpret_cast<__half const *>(rhs.raw_data());
|
||||
@@ -133,6 +186,8 @@ struct ArrayMaximum<half_t, ElementsPerAccess> {
|
||||
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_pair);
|
||||
}
|
||||
|
||||
static_assert(!(ElementsPerAccess % 2), "Output array must be divisible by vector length.");
|
||||
|
||||
#else
|
||||
|
||||
__half const *lhs_ptr = reinterpret_cast<__half const *>(lhs.raw_data());
|
||||
@@ -150,6 +205,90 @@ struct ArrayMaximum<half_t, ElementsPerAccess> {
|
||||
}
|
||||
};
|
||||
|
||||
/// Partial specialization: Element=bfloat16_t
|
||||
template <int ElementsPerAccess>
|
||||
struct ArrayMaximum<bfloat16_t, ElementsPerAccess> {
|
||||
|
||||
using NvType = __nv_bfloat16;
|
||||
using NvTypeV2 = __nv_bfloat162;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
Array<bfloat16_t, ElementsPerAccess> operator()(
|
||||
Array<bfloat16_t, ElementsPerAccess> const &lhs,
|
||||
Array<bfloat16_t, ElementsPerAccess> const &rhs) const {
|
||||
|
||||
Array<bfloat16_t, ElementsPerAccess> result;
|
||||
|
||||
#if __CUDA_ARCH__ >= 800
|
||||
int const kVectorCount = ElementsPerAccess / 2;
|
||||
|
||||
|
||||
NvTypeV2 const *lhs_ptr = reinterpret_cast<NvTypeV2 const *>(lhs.raw_data());
|
||||
NvTypeV2 const *rhs_ptr = reinterpret_cast<NvTypeV2 const *>(rhs.raw_data());
|
||||
NvTypeV2 *res_ptr = reinterpret_cast<NvTypeV2 *>(result.raw_data());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kVectorCount; ++i) {
|
||||
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_ptr[i]);
|
||||
}
|
||||
|
||||
#else
|
||||
NvType const *lhs_ptr = reinterpret_cast<NvType const *>(lhs.raw_data());
|
||||
NvType const *rhs_ptr = reinterpret_cast<NvType const *>(rhs.raw_data());
|
||||
NvType *res_ptr = reinterpret_cast<NvType *>(result.raw_data());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < ElementsPerAccess; ++i) {
|
||||
res_ptr[i] = ((lhs_ptr[i] < rhs_ptr[i]) ? rhs_ptr[i] : lhs_ptr[i]);
|
||||
}
|
||||
|
||||
#endif
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
Array<bfloat16_t, ElementsPerAccess> operator()(
|
||||
Array<bfloat16_t, ElementsPerAccess> const &lhs,
|
||||
bfloat16_t rhs) const {
|
||||
|
||||
Array<bfloat16_t, ElementsPerAccess> result;
|
||||
|
||||
#if __CUDA_ARCH__ >= 800
|
||||
int const kVectorCount = ElementsPerAccess / 2;
|
||||
|
||||
|
||||
NvType rhs_raw = reinterpret_cast<NvType const &>(rhs);
|
||||
NvTypeV2 rhs_pair = __bfloat162bfloat162(rhs_raw);
|
||||
|
||||
NvTypeV2 const *lhs_ptr = reinterpret_cast<NvTypeV2 const *>(lhs.raw_data());
|
||||
NvTypeV2 *res_ptr = reinterpret_cast<NvTypeV2 *>(result.raw_data());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < kVectorCount; ++i) {
|
||||
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_pair);
|
||||
}
|
||||
|
||||
static_assert(!(ElementsPerAccess % 2), "Output array must be divisible by vector length.");
|
||||
|
||||
#else
|
||||
|
||||
NvType const *lhs_ptr = reinterpret_cast<NvType const *>(lhs.raw_data());
|
||||
NvType const rhs_raw = reinterpret_cast<NvType const &>(rhs);
|
||||
NvType *res_ptr = reinterpret_cast<NvType *>(result.raw_data());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < ElementsPerAccess; ++i) {
|
||||
res_ptr[i] = ((lhs_ptr[i] < rhs_raw) ? rhs_raw : lhs_ptr[i]);
|
||||
}
|
||||
|
||||
#endif
|
||||
|
||||
return result;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename Element, int ElementsPerAccess>
|
||||
@@ -187,6 +326,25 @@ struct ReluConditional<half_t, ElementsPerAccess> {
|
||||
}
|
||||
};
|
||||
|
||||
template <int ElementsPerAccess>
|
||||
struct ReluConditional<bfloat16_t, ElementsPerAccess> {
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void operator()(
|
||||
bool conditional[],
|
||||
Array<bfloat16_t, ElementsPerAccess> const &fragment,
|
||||
bfloat16_t threshold) const {
|
||||
|
||||
__nv_bfloat16 y = reinterpret_cast<__nv_bfloat16 const &>(threshold);
|
||||
__nv_bfloat16 const *x = reinterpret_cast<__nv_bfloat16 const *>(fragment.raw_data());
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < ElementsPerAccess; ++i) {
|
||||
conditional[i] = !__hlt(x[i], y);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -92,9 +92,9 @@ public:
|
||||
|
||||
ElementCompute alpha; ///< scales accumulators
|
||||
ElementCompute beta; ///< scales source tensor
|
||||
ElementCompute threshold; ///< minimum value that is output
|
||||
ElementCompute const *alpha_ptr; ///< pointer to accumulator scalar - if not null, loads it from memory
|
||||
ElementCompute const *beta_ptr; ///< pointer to source scalar - if not null, loads it from memory
|
||||
ElementCompute threshold; ///< minimum value that is output
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
@@ -87,9 +87,9 @@ public:
|
||||
|
||||
ElementCompute alpha; ///< scales accumulators
|
||||
ElementCompute beta; ///< scales source tensor
|
||||
ElementCompute threshold; ///< minimum value that is output
|
||||
ElementCompute const *alpha_ptr; ///< pointer to accumulator scalar - if not null, loads it from memory
|
||||
ElementCompute const *beta_ptr; ///< pointer to source scalar - if not null, loads it from memory
|
||||
ElementCompute threshold; ///< minimum value that is output
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
@@ -1698,13 +1698,13 @@ private:
|
||||
//
|
||||
|
||||
if (OutputOp::kStoreZ) {
|
||||
destination_iterator += reduce_fragment_idx;
|
||||
destination_iterator.store(frag_Z);
|
||||
++destination_iterator;
|
||||
}
|
||||
|
||||
if (OutputOp::kStoreT) {
|
||||
tensor_iterator += reduce_fragment_idx;
|
||||
tensor_iterator.store(frag_T);
|
||||
++tensor_iterator;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
@@ -187,7 +187,7 @@ CUTLASS_CONSTEXPR_IF_CXX17
|
||||
value_t lcm(value_t a, value_t b) {
|
||||
value_t temp = gcd(a, b);
|
||||
|
||||
return temp ? (cutlass::abs_for_integer(a) / temp * cutlass::abs_for_integer(b)) : 0;
|
||||
return temp ? (cutlass::abs_for_integer(a) / temp * cutlass::abs_for_integer(b)) : value_t();
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -207,8 +207,11 @@ template <typename value_t>
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr
|
||||
value_t lcm_cxx11(value_t a, value_t b) {
|
||||
return gcd_cxx11(a, b) ? (cutlass::abs_for_integer(a) / gcd_cxx11(a, b) * cutlass::abs_for_integer(b)) : 0;
|
||||
return gcd_cxx11(a, b) ? (cutlass::abs_for_integer(a) / gcd_cxx11(a, b) *
|
||||
cutlass::abs_for_integer(b))
|
||||
: value_t();
|
||||
}
|
||||
|
||||
/// Returns the smallest value in the half-open range [a, a+b) that is a multiple of b
|
||||
CUTLASS_HOST_DEVICE
|
||||
CUTLASS_CONSTEXPR_IF_CXX17
|
||||
|
||||
@@ -31,6 +31,7 @@
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/detail/layout.hpp"
|
||||
#include "cutlass/detail/collective.hpp"
|
||||
|
||||
#include "cute/atom/mma_traits_sm90_gmma.hpp"
|
||||
#include "cute/atom/copy_traits_sm90_tma.hpp"
|
||||
@@ -322,13 +323,21 @@ is_input_fp8() {
|
||||
(cute::is_same_v<ElementB, float_e4m3_t> || cute::is_same_v<ElementB, float_e5m2_t>));
|
||||
}
|
||||
|
||||
template <class ElementA, class LayoutA, class ElementB, class LayoutB>
|
||||
// We need to handle the tuples in this function since it is used in SFINAE dispatch in the CollectiveBuilder.
|
||||
// At that point, it is not guaranteed that the tuples have been split out into the required parts.
|
||||
template <class MaybeTupleElementA, class LayoutA, class MaybeTupleElementB, class LayoutB>
|
||||
constexpr bool
|
||||
is_use_rmem_A() {
|
||||
|
||||
using ElementA = detail::deduce_mixed_width_dtype_t<0, MaybeTupleElementA>;
|
||||
using ElementB = detail::deduce_mixed_width_dtype_t<0, MaybeTupleElementB>;
|
||||
|
||||
constexpr bool IsABDifferentWidth = cute::sizeof_bits_v<ElementA> != cute::sizeof_bits_v<ElementB>;
|
||||
constexpr bool HasScales = cute::is_tuple<MaybeTupleElementA>::value ^ cute::is_tuple<MaybeTupleElementB>::value;
|
||||
constexpr bool IsInputSizeTwoBytes = is_input_size_two_bytes<ElementA, ElementB>();
|
||||
constexpr bool IsLayoutAkBk = cutlass::gemm::detail::is_k_major_A<LayoutA>() &&
|
||||
cutlass::gemm::detail::is_k_major_B<LayoutB>();
|
||||
constexpr bool IsUseRmemA = !IsInputSizeTwoBytes && !IsLayoutAkBk;
|
||||
constexpr bool IsUseRmemA = (!IsInputSizeTwoBytes && !IsLayoutAkBk) || IsABDifferentWidth || HasScales;
|
||||
return IsUseRmemA;
|
||||
}
|
||||
|
||||
|
||||
@@ -79,6 +79,50 @@ compute_stage_count_or_override(StageCountAutoCarveout<carveout_bytes> stage_cou
|
||||
return (CapacityBytes - carveout_bytes) / stage_bytes;
|
||||
}
|
||||
|
||||
// Returns the maximum number of smem tiles that can be used with a given smem capacity (with an optional scale matrix), or overrides with manual count.
|
||||
template<int CapacityBytes, class ElementA, class ElementB, class ElementScale, class ElementZero, class TileShapeMNK, int stages>
|
||||
constexpr int
|
||||
compute_stage_count_or_override_single_affine_transformed_input(StageCount<stages> stage_count) {
|
||||
return stages;
|
||||
}
|
||||
|
||||
template <class Element>
|
||||
constexpr int get_bits_for_possibly_void_element() {
|
||||
if constexpr (cute::is_same_v<Element, void>) {
|
||||
return 0;
|
||||
}
|
||||
else {
|
||||
return sizeof_bits<Element>::value;
|
||||
}
|
||||
}
|
||||
|
||||
// Returns the maximum number of smem tiles that can be used with a given smem capacity (with an optional scale matrix), or overrides with manual count.
|
||||
template<int CapacityBytes, class ElementA, class ElementB, class ElementScale, class ElementZero, class TileShapeMNK, int carveout_bytes>
|
||||
constexpr int
|
||||
compute_stage_count_or_override_single_affine_transformed_input(StageCountAutoCarveout<carveout_bytes> stage_count) {
|
||||
|
||||
// 32 bytes to account for barriers etc.
|
||||
constexpr int stage_barrier_bytes = 32;
|
||||
constexpr int scale_zero_k_tile = 1;
|
||||
constexpr int a_bits = static_cast<int>(sizeof_bits<ElementA>::value);
|
||||
constexpr int b_bits = static_cast<int>(sizeof_bits<ElementB>::value);
|
||||
constexpr int s_bits = get_bits_for_possibly_void_element<ElementScale>();
|
||||
constexpr int z_bits = get_bits_for_possibly_void_element<ElementZero>();
|
||||
|
||||
constexpr int scale_bytes = (s_bits * size<0>(TileShapeMNK{}) * scale_zero_k_tile) / 8;
|
||||
constexpr int zero_bytes = (z_bits * size<0>(TileShapeMNK{}) * scale_zero_k_tile) / 8;
|
||||
static_assert(scale_bytes % 128 == 0, "Scale bytes must be a multiple of 128");
|
||||
static_assert(zero_bytes % 128 == 0, "Zero bytes must be a multiple of 128");
|
||||
|
||||
// When scales are void, s_bits will be 0 so no smem will be allocated for scales.
|
||||
constexpr int stage_bytes =
|
||||
(a_bits * size<0>(TileShapeMNK{}) * size<2>(TileShapeMNK{})) / 8 +
|
||||
(b_bits * size<1>(TileShapeMNK{}) * size<2>(TileShapeMNK{})) / 8 +
|
||||
scale_bytes + zero_bytes + stage_barrier_bytes;
|
||||
|
||||
return (CapacityBytes - carveout_bytes) / stage_bytes;
|
||||
}
|
||||
|
||||
template <class ElementA, class LayoutA, class ElementB, class LayoutB>
|
||||
constexpr bool
|
||||
is_swapAB(){
|
||||
@@ -105,29 +149,6 @@ is_warpspecialized_transpose_B(){
|
||||
return IsWarpSpecializedTransposeB;
|
||||
}
|
||||
|
||||
template <typename ElementA, typename ElementB>
|
||||
struct Sm90TypeWidths {
|
||||
static constexpr bool IsElementALarger = (cute::sizeof_bits_v<ElementA>) > cute::sizeof_bits_v<ElementB>;
|
||||
using WideType = cute::conditional_t<IsElementALarger, ElementA, ElementB>;
|
||||
using NarrowType = cute::conditional_t<IsElementALarger, ElementB, ElementA>;
|
||||
};
|
||||
|
||||
|
||||
template <class ElementA, class LayoutA, class ElementB, class LayoutB>
|
||||
constexpr bool
|
||||
sm90_is_narrow_type_k_major() {
|
||||
using Widths = Sm90TypeWidths<ElementA, ElementB>;
|
||||
using NarrowType = typename Widths::NarrowType;
|
||||
using WideType = typename Widths::WideType;
|
||||
|
||||
constexpr bool IsANarrow = cute::is_same_v<NarrowType, ElementA>;
|
||||
constexpr cute::GMMA::Major NarrowGmmaMajor = IsANarrow ? detail::gmma_rs_tag_to_major_A<LayoutA>() :
|
||||
detail::gmma_rs_tag_to_major_B<LayoutB>();
|
||||
|
||||
constexpr bool IsNarrowLayoutKMajor = NarrowGmmaMajor == cute::GMMA::Major::K;
|
||||
return IsNarrowLayoutKMajor;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -163,18 +184,24 @@ struct CollectiveBuilder<
|
||||
cute::enable_if_t<
|
||||
(cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecialized> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedPingpong> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperative>) &&
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperative> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelArrayTmaWarpSpecializedCooperative> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelGroupTmaWarpSpecializedCooperative>) &&
|
||||
not detail::is_use_rmem_A<ElementA, GmemLayoutA, ElementB, GmemLayoutB>()>
|
||||
> {
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
static_assert(detail::is_aligned<ElementA, AlignmentA, ElementB, AlignmentB, detail::tma_alignment_bytes>(),
|
||||
"Should meet TMA alignment requirement\n");
|
||||
|
||||
static constexpr bool IsFP8Input = detail::is_input_fp8<ElementA, ElementB>();
|
||||
static constexpr bool IsArrayOfPointersGemm = (cute::is_same_v<KernelScheduleType, KernelArrayTmaWarpSpecializedCooperative> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelGroupTmaWarpSpecializedCooperative>);
|
||||
static constexpr bool IsFP8Input = detail::is_input_fp8<ElementA, ElementB>();
|
||||
static_assert(!IsFP8Input || (IsFP8Input && !IsArrayOfPointersGemm),
|
||||
"Kernel[Array/Group]TmaWarpSpecializedCooperative is only compatible with FP8 FastAccum version right now\n");
|
||||
|
||||
// For fp32 types, map to tf32 MMA value type
|
||||
using MmaElementA = cute::conditional_t<cute::is_same_v<ElementA, float>, tfloat32_t, ElementA>;
|
||||
@@ -183,7 +210,8 @@ static constexpr bool IsFP8Input = detail::is_input_fp8<ElementA, ElementB>();
|
||||
static constexpr cute::GMMA::Major GmmaMajorA = detail::gmma_ss_tag_to_major_A<MmaElementA, GmemLayoutA>();
|
||||
static constexpr cute::GMMA::Major GmmaMajorB = detail::gmma_ss_tag_to_major_B<MmaElementB, GmemLayoutB>();
|
||||
|
||||
using AtomLayoutMNK = cute::conditional_t<cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperative>,
|
||||
using AtomLayoutMNK = cute::conditional_t<
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperative> || IsArrayOfPointersGemm,
|
||||
Layout<Shape<_2,_1,_1>>, Layout<Shape<_1,_1,_1>>>;
|
||||
|
||||
using TiledMma = decltype(cute::make_tiled_mma(cute::GMMA::ss_op_selector<
|
||||
@@ -199,10 +227,12 @@ static constexpr bool IsFP8Input = detail::is_input_fp8<ElementA, ElementB>();
|
||||
|
||||
static constexpr int PipelineStages = detail::compute_stage_count_or_override<detail::sm90_smem_capacity_bytes,
|
||||
MmaElementA, MmaElementB, TileShape_MNK>(StageCountType{});
|
||||
/* For FP8 use a separate mainloop compared to other datatypes */
|
||||
using DispatchPolicy = cute::conditional_t<IsFP8Input,
|
||||
MainloopSm90TmaGmmaWarpSpecializedFP8<PipelineStages, ClusterShape_MNK, KernelScheduleType>,
|
||||
MainloopSm90TmaGmmaWarpSpecialized<PipelineStages, ClusterShape_MNK, KernelScheduleType>>;
|
||||
using DispatchPolicy = cute::conditional_t<IsArrayOfPointersGemm,
|
||||
MainloopSm90ArrayTmaGmmaWarpSpecialized<PipelineStages, ClusterShape_MNK, KernelScheduleType>,
|
||||
/* For FP8 use a separate mainloop compared to other datatypes */
|
||||
cute::conditional_t<IsFP8Input,
|
||||
MainloopSm90TmaGmmaWarpSpecializedFP8<PipelineStages, ClusterShape_MNK, KernelScheduleType>,
|
||||
MainloopSm90TmaGmmaWarpSpecialized<PipelineStages, ClusterShape_MNK, KernelScheduleType>>>;
|
||||
|
||||
using SmemCopyAtomA = void;
|
||||
using SmemCopyAtomB = void;
|
||||
@@ -267,7 +297,7 @@ struct CollectiveBuilder<
|
||||
static_assert(detail::is_aligned<ElementA, AlignmentA, ElementB, AlignmentB, detail::tma_alignment_bytes>(),
|
||||
"Should meet TMA alignment requirement\n");
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
static constexpr cute::GMMA::Major GmmaMajorA = detail::gmma_rs_tag_to_major_A<GmemLayoutA>();
|
||||
static constexpr cute::GMMA::Major GmmaMajorB = detail::gmma_rs_tag_to_major_B<GmemLayoutB>();
|
||||
@@ -323,13 +353,13 @@ struct CollectiveBuilder<
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// GMMA_TMA_WS_RS Mixed GEMM
|
||||
// GMMA_TMA_WS_RS Mixed Scaled GEMM
|
||||
template <
|
||||
class ElementPairA_,
|
||||
class GmemLayoutPairA_,
|
||||
class GmemLayoutA_,
|
||||
int AlignmentA,
|
||||
class ElementPairB_,
|
||||
class GmemLayoutPairB_,
|
||||
class GmemLayoutB_,
|
||||
int AlignmentB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
@@ -341,10 +371,10 @@ struct CollectiveBuilder<
|
||||
arch::Sm90,
|
||||
arch::OpClassTensorOp,
|
||||
ElementPairA_,
|
||||
GmemLayoutPairA_,
|
||||
GmemLayoutA_,
|
||||
AlignmentA,
|
||||
ElementPairB_,
|
||||
GmemLayoutPairB_,
|
||||
GmemLayoutB_,
|
||||
AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK,
|
||||
@@ -357,22 +387,38 @@ struct CollectiveBuilder<
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperativeMixedInput>)>
|
||||
> {
|
||||
|
||||
private:
|
||||
using ScaleA = detail::deduce_mixed_width_dtype_t<1, ElementPairA_>;
|
||||
using ScaleB = detail::deduce_mixed_width_dtype_t<1, ElementPairB_>;
|
||||
using ZeroA = detail::deduce_mixed_width_dtype_t<2, ElementPairA_>;
|
||||
using ZeroB = detail::deduce_mixed_width_dtype_t<2, ElementPairB_>;
|
||||
static constexpr bool NeitherIsTuple = !cute::is_tuple<ElementPairA_>::value && !cute::is_tuple<ElementPairB_>::value;
|
||||
|
||||
public:
|
||||
static constexpr bool IsATransformed = cute::sizeof_bits_v<ElementPairA_> < cute::sizeof_bits_v<ElementPairB_>;
|
||||
using ElementA = detail::deduce_mixed_width_dtype_t<0, ElementPairA_>;
|
||||
using ElementB = detail::deduce_mixed_width_dtype_t<0, ElementPairB_>;
|
||||
static_assert(cute::is_tuple<ElementPairA_>::value ^ cute::is_tuple<ElementPairB_>::value ||
|
||||
(NeitherIsTuple && (sizeof_bits<ElementA>::value != sizeof_bits<ElementB>::value)),
|
||||
"Either A OR B must be a tuple or the widths of A and B must be different.");
|
||||
|
||||
// Split out items for processessing, no splitting for now since scales aren't supported.
|
||||
using ElementA = ElementPairA_;
|
||||
using ElementB = ElementPairB_;
|
||||
static constexpr bool IsANarrow = sizeof_bits<ElementA>::value < sizeof_bits<ElementB>::value;
|
||||
|
||||
using GmemLayoutA = GmemLayoutPairA_;
|
||||
using GmemLayoutB = GmemLayoutPairB_;
|
||||
using GmemLayoutA = GmemLayoutA_;
|
||||
using GmemLayoutB = GmemLayoutB_;
|
||||
|
||||
using ElementPairA = cute::conditional_t<IsANarrow && NeitherIsTuple, cute::tuple<ElementA>, ElementPairA_>;
|
||||
using ElementPairB = cute::conditional_t<!IsANarrow && NeitherIsTuple, cute::tuple<ElementB>, ElementPairB_>;
|
||||
|
||||
static constexpr bool IsATransformed = cute::is_tuple<ElementPairA>::value;
|
||||
using ElementScale = cute::conditional_t<IsATransformed, ScaleA, ScaleB>;
|
||||
using ElementZero = cute::conditional_t<IsATransformed, ZeroA, ZeroB>;
|
||||
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
static_assert(detail::is_aligned<ElementA, AlignmentA, ElementB, AlignmentB, detail::tma_alignment_bytes>(),
|
||||
"Should meet TMA alignment requirement\n");
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
static constexpr cute::GMMA::Major GmmaMajorA = detail::gmma_rs_tag_to_major_A<GmemLayoutA>();
|
||||
static constexpr cute::GMMA::Major GmmaMajorB = detail::gmma_rs_tag_to_major_B<GmemLayoutB>();
|
||||
@@ -382,12 +428,7 @@ public:
|
||||
|
||||
// If A is scaled, then we don't need to swap. Otherwise, we must ensure B goes to RF and we must swap the operands.
|
||||
static constexpr bool SwapAB = !IsATransformed;
|
||||
static_assert(detail::sm90_is_narrow_type_k_major<ElementA, GmemLayoutA, ElementB, GmemLayoutB>(), "The narrow type must be K-major.");
|
||||
|
||||
static_assert((IsATransformed && (cute::sizeof_bits_v<ElementA> <= 8) && (sizeof(ElementB) == 2)) ||
|
||||
(!IsATransformed && (cute::sizeof_bits_v<ElementB> <= 8) && (sizeof(ElementA) == 2)) ||
|
||||
(GmmaMajorA == cute::GMMA::Major::K && GmmaMajorB == cute::GMMA::Major::K),
|
||||
"The unscaled element must be 2 bytes OR both inputs must be K-major");
|
||||
// When we relax the above assertion, we must handle setting the tile mma GmmaMajorB correctly.
|
||||
static constexpr cute::GMMA::Major TiledMmaGmmaMajorB = SwapAB ? GmmaMajorA : GmmaMajorB;
|
||||
|
||||
@@ -400,6 +441,7 @@ public:
|
||||
|
||||
using GmemTiledCopyA = decltype(detail::sm90_cluster_shape_to_tma_atom(shape<1>(ClusterShape_MNK{})));
|
||||
using GmemTiledCopyB = decltype(detail::sm90_cluster_shape_to_tma_atom(shape<0>(ClusterShape_MNK{})));
|
||||
|
||||
using SmemLayoutAtomA = decltype(detail::rs_smem_selector<GmmaMajorA, ElementA,
|
||||
decltype(cute::get<0>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{})), IsWarpSpecializedTransposeB>());
|
||||
using SmemLayoutAtomB = decltype(detail::rs_smem_selector<GmmaMajorB, ElementB,
|
||||
@@ -407,8 +449,8 @@ public:
|
||||
|
||||
using RealElementA = cute::conditional_t<SwapAB, ElementB, ElementA>;
|
||||
using RealElementB = cute::conditional_t<SwapAB, ElementA, ElementB>;
|
||||
static constexpr int PipelineStages = detail::compute_stage_count_or_override<detail::sm90_smem_capacity_bytes,
|
||||
RealElementA, RealElementB, TileShape_MNK>(StageCountType{});
|
||||
static constexpr int PipelineStages = detail::compute_stage_count_or_override_single_affine_transformed_input<detail::sm90_smem_capacity_bytes,
|
||||
RealElementA, RealElementB, ElementScale, ElementZero, TileShape_MNK>(StageCountType{});
|
||||
|
||||
using SmemCopyAtomA = cute::conditional_t<SwapAB, void, Copy_Atom<cute::DefaultCopy, ElementA>>;
|
||||
using SmemCopyAtomB = cute::conditional_t<SwapAB, Copy_Atom<cute::DefaultCopy, ElementB>, void>;
|
||||
@@ -416,35 +458,24 @@ public:
|
||||
using DispatchPolicy = MainloopSm90TmaGmmaRmemAWarpSpecializedMixedInput<PipelineStages, ClusterShape_MNK, KernelScheduleType>;
|
||||
|
||||
// We pack the scale data with the operand that will be optionally scaled and converted before MMA.
|
||||
using StrideAPair = TagToStrideA_t<GmemLayoutA>;
|
||||
using StrideBPair = TagToStrideB_t<GmemLayoutB>;
|
||||
using StrideA = TagToStrideA_t<GmemLayoutA>;
|
||||
using StrideB = TagToStrideB_t<GmemLayoutB>;
|
||||
|
||||
using GmemTiledCopyAPair = GmemTiledCopyA;
|
||||
using SmemLayoutAtomAPair = SmemLayoutAtomA;
|
||||
using SmemCopyAtomAPair = SmemCopyAtomA;
|
||||
|
||||
using GmemTiledCopyBPair = GmemTiledCopyB;
|
||||
using SmemLayoutAtomBPair = SmemLayoutAtomB;
|
||||
using SmemCopyAtomBPair = SmemCopyAtomB;
|
||||
|
||||
|
||||
// If the src type of the converter is the same as ElementA,
|
||||
// interpret this as if the user wanted to apply the scale to the A matrix.
|
||||
using CollectiveOp = CollectiveMma<
|
||||
DispatchPolicy,
|
||||
TileShape_MNK,
|
||||
ElementPairA_,
|
||||
StrideAPair,
|
||||
ElementPairB_,
|
||||
StrideBPair,
|
||||
ElementPairA,
|
||||
StrideA,
|
||||
ElementPairB,
|
||||
StrideB,
|
||||
TiledMma,
|
||||
GmemTiledCopyAPair,
|
||||
SmemLayoutAtomAPair,
|
||||
SmemCopyAtomAPair,
|
||||
GmemTiledCopyA,
|
||||
SmemLayoutAtomA,
|
||||
SmemCopyAtomA,
|
||||
cute::identity,
|
||||
GmemTiledCopyBPair,
|
||||
SmemLayoutAtomBPair,
|
||||
SmemCopyAtomBPair,
|
||||
GmemTiledCopyB,
|
||||
SmemLayoutAtomB,
|
||||
SmemCopyAtomB,
|
||||
cute::identity
|
||||
>;
|
||||
|
||||
@@ -483,7 +514,9 @@ struct CollectiveBuilder<
|
||||
cute::enable_if_t<
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedFP8FastAccum> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedPingpongFP8FastAccum> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperativeFP8FastAccum>>
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperativeFP8FastAccum> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelArrayTmaWarpSpecializedCooperativeFP8FastAccum> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelGroupTmaWarpSpecializedCooperativeFP8FastAccum>>
|
||||
> {
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
@@ -495,13 +528,16 @@ struct CollectiveBuilder<
|
||||
static_assert(!detail::is_use_rmem_A<ElementA, GmemLayoutA, ElementB, GmemLayoutB>(),
|
||||
"Not supported for fp8 non-TN warp specialized kernels yet\n");
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
|
||||
static constexpr cute::GMMA::Major GmmaMajorA = detail::gmma_ss_tag_to_major_A<ElementA, GmemLayoutA>();
|
||||
static constexpr cute::GMMA::Major GmmaMajorB = detail::gmma_ss_tag_to_major_B<ElementB, GmemLayoutB>();
|
||||
|
||||
using AtomLayoutMNK = cute::conditional_t<cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperativeFP8FastAccum>,
|
||||
static constexpr bool IsArrayOfPointersGemm = (cute::is_same_v<KernelScheduleType, KernelArrayTmaWarpSpecializedCooperativeFP8FastAccum> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelGroupTmaWarpSpecializedCooperativeFP8FastAccum>);
|
||||
using AtomLayoutMNK = cute::conditional_t<cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedCooperativeFP8FastAccum> ||
|
||||
IsArrayOfPointersGemm,
|
||||
Layout<Shape<_2,_1,_1>>, Layout<Shape<_1,_1,_1>>>;
|
||||
|
||||
using TiledMma = decltype(cute::make_tiled_mma(cute::GMMA::ss_op_selector<
|
||||
@@ -517,8 +553,9 @@ struct CollectiveBuilder<
|
||||
|
||||
static constexpr int PipelineStages = detail::compute_stage_count_or_override<detail::sm90_smem_capacity_bytes,
|
||||
ElementA, ElementB, TileShape_MNK>(StageCountType{});
|
||||
using DispatchPolicy = MainloopSm90TmaGmmaWarpSpecialized<
|
||||
PipelineStages, ClusterShape_MNK, KernelScheduleType>;
|
||||
using DispatchPolicy = cute::conditional_t<IsArrayOfPointersGemm,
|
||||
MainloopSm90ArrayTmaGmmaWarpSpecialized<PipelineStages, ClusterShape_MNK, KernelScheduleType>,
|
||||
MainloopSm90TmaGmmaWarpSpecialized<PipelineStages, ClusterShape_MNK, KernelScheduleType>>;
|
||||
|
||||
using SmemCopyAtomA = void;
|
||||
using SmemCopyAtomB = void;
|
||||
@@ -580,7 +617,7 @@ struct CollectiveBuilder<
|
||||
static_assert(detail::is_aligned<ElementA, AlignmentA, ElementB, AlignmentB, detail::tma_alignment_bytes>(),
|
||||
"Should meet TMA alignment requirement\n");
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
|
||||
// For fp32 types, map to tf32 MMA value type
|
||||
@@ -721,7 +758,7 @@ struct CollectiveBuilder<
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
|
||||
// For fp32 types, map to tf32 MMA value type
|
||||
@@ -819,7 +856,7 @@ struct CollectiveBuilder<
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
|
||||
// For fp32 types, map to tf32 MMA value type
|
||||
@@ -918,13 +955,20 @@ struct CollectiveBuilder<
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
|
||||
static constexpr bool IsTmaCompatible = detail::is_aligned<
|
||||
ElementA, AlignmentA, ElementB, AlignmentB, detail::tma_alignment_bytes>();
|
||||
using ExtractedElementA = detail::deduce_mixed_width_dtype_t<0, ElementA>;
|
||||
using ExtractedElementB = detail::deduce_mixed_width_dtype_t<0, ElementB>;
|
||||
|
||||
static constexpr bool IsMixedWidthInput = cute::sizeof_bits_v<ElementA> != cute::sizeof_bits_v<ElementB>;
|
||||
static constexpr bool IsTmaCompatible = detail::is_aligned<
|
||||
ExtractedElementA, AlignmentA, ExtractedElementB, AlignmentB, detail::tma_alignment_bytes>();
|
||||
|
||||
// Users opt into scales via the builder by passing a tuple of Elements for the input that will be scaled. We detect
|
||||
// scale support if ONLY one of the inputs have tuples to describe them.
|
||||
static constexpr bool OnlyOneIsTuple = cute::is_tuple<ElementA>::value ^ cute::is_tuple<ElementB>::value;
|
||||
static constexpr bool IsDifferentWidth = sizeof_bits<ExtractedElementA>::value != sizeof_bits<ExtractedElementB>::value;
|
||||
static constexpr bool IsMixedWidthInput = IsDifferentWidth || (IsDifferentWidth && OnlyOneIsTuple);
|
||||
|
||||
#if ((__CUDACC_VER_MAJOR__ > 12) || ((__CUDACC_VER_MAJOR__ == 12) && (__CUDACC_VER_MINOR__ >= 1)))
|
||||
// Persistent schedules perform best for CUDA Toolkits with version >= 12.1
|
||||
|
||||
@@ -56,7 +56,7 @@ template <
|
||||
class TransformB
|
||||
>
|
||||
struct CollectiveMma {
|
||||
static_assert(cutlass::detail::dependent_false<ElementA> == 0, "Could not find a mainloop specialization.");
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Could not find a mainloop specialization.");
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -73,5 +73,6 @@ struct CollectiveMma {
|
||||
#include "cutlass/gemm/collective/sm90_mma_tma_gmma_rs_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_tma_gmma_rs_warpspecialized_mixed_input.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_tma_gmma_ss_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_tma_gmma_ss_warpspecialized_fp8.hpp"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -0,0 +1,741 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2023 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/arch/cluster_sm90.hpp"
|
||||
#include "cute/arch/copy_sm90.hpp"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "cute/algorithm/functional.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cute/algorithm/gemm.hpp"
|
||||
#include "cute/tensor_predicate.hpp"
|
||||
#include "cute/numeric/arithmetic_tuple.hpp"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
using namespace cute;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// WarpSpecialized Mainloop
|
||||
template <
|
||||
int Stages,
|
||||
class ClusterShape,
|
||||
class KernelSchedule,
|
||||
class TileShape_,
|
||||
class ElementA_,
|
||||
class StrideA_,
|
||||
class ElementB_,
|
||||
class StrideB_,
|
||||
class TiledMma_,
|
||||
class GmemTiledCopyA_,
|
||||
class SmemLayoutAtomA_,
|
||||
class SmemCopyAtomA_,
|
||||
class TransformA_,
|
||||
class GmemTiledCopyB_,
|
||||
class SmemLayoutAtomB_,
|
||||
class SmemCopyAtomB_,
|
||||
class TransformB_>
|
||||
struct CollectiveMma<
|
||||
MainloopSm90ArrayTmaGmmaWarpSpecialized<Stages, ClusterShape, KernelSchedule>,
|
||||
TileShape_,
|
||||
ElementA_,
|
||||
StrideA_,
|
||||
ElementB_,
|
||||
StrideB_,
|
||||
TiledMma_,
|
||||
GmemTiledCopyA_,
|
||||
SmemLayoutAtomA_,
|
||||
SmemCopyAtomA_,
|
||||
TransformA_,
|
||||
GmemTiledCopyB_,
|
||||
SmemLayoutAtomB_,
|
||||
SmemCopyAtomB_,
|
||||
TransformB_>
|
||||
{
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using DispatchPolicy = MainloopSm90ArrayTmaGmmaWarpSpecialized<Stages, ClusterShape, KernelSchedule>;
|
||||
using TileShape = TileShape_;
|
||||
using ElementA = ElementA_;
|
||||
using StrideA = StrideA_;
|
||||
using ElementB = ElementB_;
|
||||
using StrideB = StrideB_;
|
||||
using TiledMma = TiledMma_;
|
||||
using ElementAccumulator = typename TiledMma::ValTypeC;
|
||||
using GmemTiledCopyA = GmemTiledCopyA_;
|
||||
using GmemTiledCopyB = GmemTiledCopyB_;
|
||||
using SmemLayoutAtomA = SmemLayoutAtomA_;
|
||||
using SmemLayoutAtomB = SmemLayoutAtomB_;
|
||||
using SmemCopyAtomA = SmemCopyAtomA_;
|
||||
using SmemCopyAtomB = SmemCopyAtomB_;
|
||||
using TransformA = TransformA_;
|
||||
using TransformB = TransformB_;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<DispatchPolicy::Stages>;
|
||||
using PipelineState = cutlass::PipelineState<DispatchPolicy::Stages>;
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
|
||||
static_assert(rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)");
|
||||
static_assert((size<0>(TileShape{}) % size<0>(SmemLayoutAtomA{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert((size<2>(TileShape{}) % size<1>(SmemLayoutAtomA{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
|
||||
static_assert(rank(SmemLayoutAtomB{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)");
|
||||
static_assert((size<1>(TileShape{}) % size<0>(SmemLayoutAtomB{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert((size<2>(TileShape{}) % size<1>(SmemLayoutAtomB{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
|
||||
// Tile along modes in a way that maximizes the TMA box size.
|
||||
using SmemLayoutA = decltype(tile_to_shape(
|
||||
SmemLayoutAtomA{},
|
||||
make_shape(shape<0>(TileShape{}), shape<2>(TileShape{}), Int<DispatchPolicy::Stages>{}),
|
||||
conditional_t< ::cutlass::gemm::detail::is_major<0,StrideA>(), Step<_2,_1,_3>, Step<_1,_2,_3>>{}));
|
||||
using SmemLayoutB = decltype(tile_to_shape(
|
||||
SmemLayoutAtomB{},
|
||||
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{}), Int<DispatchPolicy::Stages>{}),
|
||||
conditional_t< ::cutlass::gemm::detail::is_major<0,StrideB>(), Step<_2,_1,_3>, Step<_1,_2,_3>>{}));
|
||||
|
||||
static_assert(DispatchPolicy::Stages >= 2, "Specialization requires Stages set to value 2 or more.");
|
||||
static_assert(cute::is_base_of<cute::GMMA::DescriptorIterator, typename TiledMma::FrgTypeA>::value &&
|
||||
cute::is_base_of<cute::GMMA::DescriptorIterator, typename TiledMma::FrgTypeB>::value,
|
||||
"MMA atom must source both A and B operand from smem_desc for this mainloop.");
|
||||
static_assert(cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_MULTICAST>,
|
||||
"GmemTiledCopy - invalid SM90 TMA copy atom specified.");
|
||||
static_assert(cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>,
|
||||
"GmemTiledCopy - invalid SM90 TMA copy atom specified.");
|
||||
|
||||
// TMA converts f32 input to tf32 when copying from GMEM to SMEM
|
||||
// For all other types, cast to size equivalent uint type to avoid any rounding by TMA.
|
||||
static constexpr bool ConvertF32toTF32A = cute::is_same_v<float, ElementA>;
|
||||
static constexpr bool ConvertF32toTF32B = cute::is_same_v<float, ElementB>;
|
||||
using InternalElementA = cute::conditional_t<ConvertF32toTF32A, tfloat32_t, uint_bit_t<sizeof_bits_v<ElementA>>>;
|
||||
using InternalElementB = cute::conditional_t<ConvertF32toTF32B, tfloat32_t, uint_bit_t<sizeof_bits_v<ElementB>>>;
|
||||
|
||||
// Assumption: StrideA is congruent with Problem_MK
|
||||
using TMA_A = decltype(make_tma_copy(
|
||||
GmemTiledCopyA{},
|
||||
make_tensor(static_cast<InternalElementA const*>(nullptr), repeat_like(StrideA{}, int32_t(0)), StrideA{}),
|
||||
SmemLayoutA{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<0>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<1>(ClusterShape{}))); // mcast along N mode for this M load, if any
|
||||
// Assumption: StrideB is congruent with Problem_NK
|
||||
using TMA_B = decltype(make_tma_copy(
|
||||
GmemTiledCopyB{},
|
||||
make_tensor(static_cast<InternalElementB const*>(nullptr), repeat_like(StrideB{}, int32_t(0)), StrideB{}),
|
||||
SmemLayoutB{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<0>(ClusterShape{}))); // mcast along M mode for this N load, if any
|
||||
|
||||
struct SharedStorage {
|
||||
struct TensorStorage : cute::aligned_struct<128> {
|
||||
cute::array_aligned<typename TiledMma::ValTypeA, cute::cosize_v<SmemLayoutA>> smem_A;
|
||||
cute::array_aligned<typename TiledMma::ValTypeB, cute::cosize_v<SmemLayoutB>> smem_B;
|
||||
} tensors;
|
||||
|
||||
struct TensorMapStorage : cute::aligned_struct<128> {
|
||||
cute::TmaDescriptor smem_tensormap_A;
|
||||
cute::TmaDescriptor smem_tensormap_B;
|
||||
} tensormaps;
|
||||
|
||||
using PipelineStorage = typename MainloopPipeline::SharedStorage;
|
||||
PipelineStorage pipeline;
|
||||
};
|
||||
using TensorStorage = typename SharedStorage::TensorStorage;
|
||||
using TensorMapStorage = typename SharedStorage::TensorMapStorage;
|
||||
using PipelineStorage = typename SharedStorage::PipelineStorage;
|
||||
|
||||
static constexpr bool IsGroupedGemmKernel = cute::is_base_of_v<KernelGroupTmaWarpSpecializedCooperative, KernelSchedule>;
|
||||
using StridesA = cute::conditional_t<IsGroupedGemmKernel, StrideA const*, StrideA>;
|
||||
using StridesB = cute::conditional_t<IsGroupedGemmKernel, StrideB const*, StrideB>;
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
ElementA const** ptr_A;
|
||||
StridesA dA;
|
||||
ElementB const** ptr_B;
|
||||
StridesB dB;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
TMA_A tma_load_a;
|
||||
TMA_B tma_load_b;
|
||||
void* tensormaps;
|
||||
InternalElementA const** ptr_A;
|
||||
StridesA dA;
|
||||
InternalElementB const** ptr_B;
|
||||
StridesB dB;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
template <class ProblemShape>
|
||||
static constexpr Params
|
||||
to_underlying_arguments(
|
||||
ProblemShape problem_shapes,
|
||||
Arguments const& args,
|
||||
void* workspace) {
|
||||
// Optionally append 1s until problem shape is rank-4 (MNKL), in case it is only rank-3 (MNK)
|
||||
auto problem_shape_MNKL = append<4>(problem_shapes.get_host_problem_shape(0), 1);
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
const uint32_t mock_L = 1;
|
||||
|
||||
// These tensor pointers are only used to create tensormap/tma desc.
|
||||
// This address to the tensor will be replaced with correct address before the initial tma load
|
||||
InternalElementA const* ptr_A_first_batch = reinterpret_cast<InternalElementA const*>(args.ptr_A);
|
||||
InternalElementB const* ptr_B_first_batch = reinterpret_cast<InternalElementA const*>(args.ptr_B);
|
||||
cudaError_t cuda_error = cudaGetLastError(); // clear previous error
|
||||
|
||||
StrideA stride_a;
|
||||
StrideB stride_b;
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
// Strides for Grouped Gemm will be replaced prior to the first access regardless
|
||||
stride_a = StrideA{};
|
||||
stride_b = StrideB{};
|
||||
}
|
||||
else {
|
||||
stride_a = args.dA;
|
||||
stride_b = args.dB;
|
||||
}
|
||||
Tensor tensor_a = make_tensor(ptr_A_first_batch, make_layout(make_shape(M,K,mock_L), stride_a));
|
||||
Tensor tensor_b = make_tensor(ptr_B_first_batch, make_layout(make_shape(N,K,mock_L), stride_b));
|
||||
TMA_A tma_load_a = make_tma_copy(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
SmemLayoutA{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<0>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<1>(ClusterShape{})); // mcast along N mode for this M load, if any
|
||||
TMA_B tma_load_b = make_tma_copy(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<0>(ClusterShape{})); // mcast along M mode for this N load, if any
|
||||
|
||||
void* tensormaps = workspace;
|
||||
|
||||
return {
|
||||
tma_load_a,
|
||||
tma_load_b,
|
||||
tensormaps,
|
||||
reinterpret_cast<InternalElementA const**>(args.ptr_A),
|
||||
args.dA,
|
||||
reinterpret_cast<InternalElementB const**>(args.ptr_B),
|
||||
args.dB
|
||||
};
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static size_t
|
||||
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args, int sm_count) {
|
||||
constexpr uint32_t NumInputTensors = 2;
|
||||
constexpr size_t SizeOfCuTensorMap = sizeof(cute::TmaDescriptor);
|
||||
// Allocate gmem space for input tensormaps per each SM, A tensormap copies followed by B tensormap copies
|
||||
return (NumInputTensors * SizeOfCuTensorMap * sm_count);
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static cutlass::Status
|
||||
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
template<class ProblemShape>
|
||||
CUTLASS_HOST_DEVICE static bool
|
||||
can_implement(
|
||||
ProblemShape problem_shapes,
|
||||
Arguments const& args) {
|
||||
constexpr int tma_alignment_bits = 128;
|
||||
constexpr int min_tma_aligned_elements_A = tma_alignment_bits / cutlass::sizeof_bits<ElementA>::value;
|
||||
constexpr int min_tma_aligned_elements_B = tma_alignment_bits / cutlass::sizeof_bits<ElementB>::value;
|
||||
|
||||
bool implementable = true;
|
||||
// Check alignment for all problem sizes
|
||||
for (int i = 0; i < problem_shapes.groups(); i++) {
|
||||
auto problem_shape_MNKL = append<4>(problem_shapes.get_host_problem_shape(i), 1);
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_A>(cute::make_shape(M,K,L), StrideA{});
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_B>(cute::make_shape(N,K,L), StrideB{});
|
||||
}
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n");
|
||||
}
|
||||
return implementable;
|
||||
}
|
||||
|
||||
static constexpr int K_PIPE_MAX = DispatchPolicy::Stages;
|
||||
static constexpr int K_PIPE_MMAS = 1;
|
||||
static constexpr uint32_t TmaTransactionBytes =
|
||||
(size<0>(SmemLayoutA{}) * size<1>(SmemLayoutA{}) * static_cast<uint32_t>(sizeof_bits<ElementA>::value)) / 8+
|
||||
(size<0>(SmemLayoutB{}) * size<1>(SmemLayoutB{}) * static_cast<uint32_t>(sizeof_bits<ElementB>::value)) / 8;
|
||||
|
||||
|
||||
// Set up the data needed by this collective for load and mma.
|
||||
// Returns a tuple of tensors. The collective and the kernel layer have the contract that the
|
||||
// returned tuple must contain at least two elements, with the first two elements being:
|
||||
// gA_mkl - The tma tensor, A after a local tile so it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// gB_nkl - The tma tensor, B after a local tile so it has shape (BLK_N,BLK_K,n,k,l)
|
||||
// The rest of the tensors can be specified as needed by this collective.
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE auto
|
||||
load_init(ProblemShape_MNKL const& problem_shape_MNKL, Params const& mainloop_params) const {
|
||||
using X = Underscore;
|
||||
// Separate out problem shape for convenience
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
const int32_t mock_L = 1;
|
||||
|
||||
// TMA requires special handling of strides to deal with coord codomain mapping
|
||||
// Represent the full tensors -- get these from TMA
|
||||
Tensor mA_mkl = mainloop_params.tma_load_a.get_tma_tensor(make_shape(M,K,mock_L)); // (m,k,l)
|
||||
Tensor mB_nkl = mainloop_params.tma_load_b.get_tma_tensor(make_shape(N,K,mock_L)); // (n,k,l)
|
||||
|
||||
// Make tiled views, defer the slice
|
||||
Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k,l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
|
||||
|
||||
return cute::make_tuple(gA_mkl, gB_nkl);
|
||||
}
|
||||
|
||||
// Perform a collective-scoped matrix multiply-accumulate
|
||||
// Producer Perspective
|
||||
template <
|
||||
class TensorA, class TensorB,
|
||||
class TensorMapA, class TensorMapB,
|
||||
class KTileIterator, class BlockCoord
|
||||
>
|
||||
CUTLASS_DEVICE void
|
||||
load(
|
||||
Params const& mainloop_params,
|
||||
MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_write,
|
||||
cute::tuple<TensorA, TensorB> const& load_inputs,
|
||||
cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps,
|
||||
BlockCoord const& blk_coord,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thread_idx,
|
||||
uint32_t block_rank_in_cluster,
|
||||
TensorStorage& shared_tensors) {
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
if (warp_idx_in_warp_group == 0 and lane_predicate) {
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
|
||||
//
|
||||
// Prepare the TMA loads for A and B
|
||||
//
|
||||
|
||||
constexpr uint32_t cluster_shape_x = get<0>(DispatchPolicy::ClusterShape());
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
auto block_tma_a = mainloop_params.tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = mainloop_params.tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
// Partition the inputs based on the current block coordinates.
|
||||
auto [m_coord, n_coord, k_coord, l_coord] = blk_coord;
|
||||
Tensor gA = gA_mkl(_,_,m_coord,_,l_coord); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nkl(_,_,n_coord,_,l_coord); // (BLK_N,BLK_K,k)
|
||||
|
||||
// Applies the mapping from block_tma_a
|
||||
Tensor tAgA = block_tma_a.partition_S(gA); // (TMA,TMA_M,TMA_K,k)
|
||||
Tensor tAsA = block_tma_a.partition_D(sA); // (TMA,TMA_M,TMA_K,PIPE)
|
||||
|
||||
Tensor tBgB = block_tma_b.partition_S(gB); // (TMA,TMA_N,TMA_K,k)
|
||||
Tensor tBsB = block_tma_b.partition_D(sB); // (TMA,TMA_N,TMA_K,PIPE)
|
||||
|
||||
uint16_t mcast_mask_a = 0;
|
||||
uint16_t mcast_mask_b = 0;
|
||||
|
||||
// Issue TmaLoads
|
||||
// Maps the tile -> block, value
|
||||
if constexpr (cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_MULTICAST>) {
|
||||
auto block_layout = Layout<typename DispatchPolicy::ClusterShape>{}; // (m,n) -> block_id
|
||||
for (int n = 0; n < size<1>(block_layout); ++n) {
|
||||
mcast_mask_a |= (uint16_t(1) << block_layout(cluster_local_block_id.x,n,Int<0>{}));
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>) {
|
||||
auto block_layout = Layout<typename DispatchPolicy::ClusterShape>{}; // (m,n) -> block_id
|
||||
for (int m = 0; m < size<0>(block_layout); ++m) {
|
||||
mcast_mask_b |= (uint16_t(1) << block_layout(m,cluster_local_block_id.y,Int<0>{}));
|
||||
}
|
||||
}
|
||||
|
||||
// Mainloop
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for ( ; k_tile_count > 0; --k_tile_count)
|
||||
{
|
||||
// LOCK smem_pipe_write for _writing_
|
||||
pipeline.producer_acquire(smem_pipe_write);
|
||||
|
||||
//
|
||||
// Copy gmem to smem for *k_tile_iter
|
||||
//
|
||||
|
||||
using BarrierType = typename MainloopPipeline::ProducerBarrierType;
|
||||
BarrierType* tma_barrier = pipeline.producer_get_barrier(smem_pipe_write);
|
||||
|
||||
int write_stage = smem_pipe_write.index();
|
||||
copy(mainloop_params.tma_load_a.with(get<0>(input_tensormaps), *tma_barrier, mcast_mask_a), tAgA(_,_,_,*k_tile_iter), tAsA(_,_,_,write_stage));
|
||||
copy(mainloop_params.tma_load_b.with(get<1>(input_tensormaps), *tma_barrier, mcast_mask_b), tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,write_stage));
|
||||
++k_tile_iter;
|
||||
|
||||
// Advance smem_pipe_write
|
||||
++smem_pipe_write;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Perform a Producer Epilogue to prevent early exit of blocks in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(MainloopPipeline pipeline, PipelineState smem_pipe_write) {
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
// Issue the epilogue waits
|
||||
if (warp_idx_in_warp_group == 0 and lane_predicate) {
|
||||
// This helps avoid early exit of blocks in Cluster.
|
||||
// Waits for all stages to either be released (all
|
||||
// Consumer UNLOCKs), or if the stage was never used
|
||||
// then it would just be acquired since the phase was
|
||||
// still inverted from make_producer_start_state.
|
||||
pipeline.producer_tail(smem_pipe_write);
|
||||
}
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Consumer Perspective
|
||||
template <
|
||||
class FrgTensorC
|
||||
>
|
||||
CUTLASS_DEVICE void
|
||||
mma(MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_read,
|
||||
FrgTensorC& accum,
|
||||
int k_tile_count,
|
||||
int thread_idx,
|
||||
TensorStorage& shared_tensors,
|
||||
Params const& mainloop_params) {
|
||||
static_assert(is_rmem<FrgTensorC>::value, "C tensor must be rmem resident.");
|
||||
static_assert(rank(SmemLayoutA{}) == 3, "Smem layout must be rank 3.");
|
||||
static_assert(rank(SmemLayoutB{}) == 3, "Smem layout must be rank 3.");
|
||||
static_assert(cute::is_void_v<SmemCopyAtomA>,
|
||||
"SM90 GMMA mainloops cannot have a non-void copy atom for smem sourced instructions.");
|
||||
static_assert(cute::is_void_v<SmemCopyAtomB>,
|
||||
"SM90 GMMA mainloops cannot have a non-void copy atom for smem sourced instructions.");
|
||||
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
|
||||
//
|
||||
// Define C accumulators and A/B partitioning
|
||||
//
|
||||
|
||||
TiledMma tiled_mma;
|
||||
auto thread_mma = tiled_mma.get_thread_slice(thread_idx);
|
||||
|
||||
Tensor tCsA = thread_mma.partition_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor tCsB = thread_mma.partition_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
// Allocate "fragments/descriptors"
|
||||
Tensor tCrA = thread_mma.make_fragment_A(tCsA); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor tCrB = thread_mma.make_fragment_B(tCsB); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCsA) == size<1>(accum)); // M
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCsB) == size<2>(accum)); // N
|
||||
CUTE_STATIC_ASSERT_V(size<2>(tCsA) == size<2>(tCsB)); // K
|
||||
CUTE_STATIC_ASSERT_V(size<3>(tCsA) == size<3>(tCsB)); // PIPE
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<2>(sA)); // PIPE
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<2>(sB)); // PIPE
|
||||
|
||||
//
|
||||
// PIPELINED MAIN LOOP
|
||||
//
|
||||
static_assert((0 <= K_PIPE_MMAS) && (K_PIPE_MMAS < K_PIPE_MAX),
|
||||
"ERROR : Incorrect number of MMAs in flight");
|
||||
|
||||
// We release buffers to producer warps(dma load) with some mmas in flight
|
||||
PipelineState smem_pipe_release = smem_pipe_read;
|
||||
|
||||
// Prologue GMMAs
|
||||
int prologue_mma_count = min(K_PIPE_MMAS, k_tile_count);
|
||||
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::Zero;
|
||||
|
||||
warpgroup_fence_operand(accum);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_tile_prologue = prologue_mma_count; k_tile_prologue > 0; --k_tile_prologue)
|
||||
{
|
||||
// WAIT on smem_pipe_read until its data are available (phase bit flips from rdPhaseBit value)
|
||||
auto barrier_token = pipeline.consumer_try_wait(smem_pipe_read);
|
||||
pipeline.consumer_wait(smem_pipe_read, barrier_token);
|
||||
|
||||
int read_stage = smem_pipe_read.index();
|
||||
warpgroup_arrive();
|
||||
// Unroll the K mode manually to set scale D to 1
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
|
||||
// (V,M,K) x (V,N,K) => (V,M,N)
|
||||
cute::gemm(tiled_mma, tCrA(_,_,k_block,read_stage), tCrB(_,_,k_block,read_stage), accum);
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
|
||||
}
|
||||
|
||||
warpgroup_commit_batch();
|
||||
|
||||
++smem_pipe_read;
|
||||
}
|
||||
|
||||
warpgroup_fence_operand(accum);
|
||||
// Mainloop GMMAs
|
||||
k_tile_count -= prologue_mma_count;
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for ( ; k_tile_count > 0; --k_tile_count)
|
||||
{
|
||||
// WAIT on smem_pipe_read until its data are available (phase bit flips from rdPhaseBit value)
|
||||
auto barrier_token = pipeline.consumer_try_wait(smem_pipe_read);
|
||||
pipeline.consumer_wait(smem_pipe_read, barrier_token);
|
||||
|
||||
//
|
||||
// Compute on k_tile
|
||||
//
|
||||
|
||||
int read_stage = smem_pipe_read.index();
|
||||
warpgroup_fence_operand(accum);
|
||||
warpgroup_arrive();
|
||||
// Unroll the K mode manually to set scale D to 1
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
|
||||
// (V,M,K) x (V,N,K) => (V,M,N)
|
||||
cute::gemm(tiled_mma, tCrA(_,_,k_block,read_stage), tCrB(_,_,k_block,read_stage), accum);
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
|
||||
}
|
||||
warpgroup_commit_batch();
|
||||
|
||||
/// Wait on the GMMA barrier for K_PIPE_MMAS (or fewer) outstanding to ensure smem_pipe_write is consumed
|
||||
warpgroup_wait<K_PIPE_MMAS>();
|
||||
warpgroup_fence_operand(accum);
|
||||
|
||||
// UNLOCK smem_pipe_release, done _computing_ on it
|
||||
pipeline.consumer_release(smem_pipe_release);
|
||||
|
||||
// Advance smem_pipe_read and smem_pipe_release
|
||||
++smem_pipe_read;
|
||||
++smem_pipe_release;
|
||||
}
|
||||
|
||||
warpgroup_fence_operand(accum);
|
||||
}
|
||||
|
||||
/// Perform a Consumer Epilogue to release all buffers
|
||||
CUTLASS_DEVICE void
|
||||
mma_tail(MainloopPipeline pipeline, PipelineState smem_pipe_release, int k_tile_count) {
|
||||
// Prologue GMMAs
|
||||
int prologue_mma_count = min(K_PIPE_MMAS, k_tile_count);
|
||||
k_tile_count -= prologue_mma_count;
|
||||
|
||||
smem_pipe_release.advance(k_tile_count);
|
||||
|
||||
// Wait on all GMMAs to complete
|
||||
warpgroup_wait<0>();
|
||||
|
||||
for (int count = 0; count < prologue_mma_count; ++count) {
|
||||
pipeline.consumer_release(smem_pipe_release); // UNLOCK smem_pipe_release, done _computing_ on it
|
||||
++smem_pipe_release;
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// Methods to perform different parts of TMA/Tensormap modifications
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE auto
|
||||
tensormaps_init(Params const& mainloop_params, int32_t const sm_count, int32_t const sm_idx) const {
|
||||
cute::TmaDescriptor* gmem_tensormap = reinterpret_cast<cute::TmaDescriptor*>(mainloop_params.tensormaps);
|
||||
|
||||
cute::TmaDescriptor* tma_desc_a = &gmem_tensormap[sm_idx];
|
||||
cute::TmaDescriptor* tma_desc_b = &gmem_tensormap[sm_idx + sm_count];
|
||||
|
||||
if (cute::elect_one_sync()) {
|
||||
// Bringing tensormaps from params to gmem for modification later
|
||||
Tensor pA_tensormap = make_tensor(mainloop_params.tma_load_a.get_tma_descriptor(), Int<1>{}, Int<1>{});
|
||||
Tensor gA_tensormap = make_tensor(tma_desc_a, Int<1>{}, Int<1>{});
|
||||
Tensor pB_tensormap = make_tensor(mainloop_params.tma_load_b.get_tma_descriptor(), Int<1>{}, Int<1>{});
|
||||
Tensor gB_tensormap = make_tensor(tma_desc_b, Int<1>{}, Int<1>{});
|
||||
|
||||
copy(recast<uint128_t>(pA_tensormap), recast<uint128_t>(gA_tensormap));
|
||||
copy(recast<uint128_t>(pB_tensormap), recast<uint128_t>(gB_tensormap));
|
||||
}
|
||||
|
||||
return cute::make_tuple(tma_desc_a, tma_desc_b);
|
||||
}
|
||||
|
||||
// Bringing tensormaps to smem (to be done by single thread)
|
||||
template <class TensorMapA, class TensorMapB>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_fetch_to_smem(
|
||||
TensorMapStorage& shared_tensormap,
|
||||
cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps) const {
|
||||
Tensor gA_tensormap = make_tensor(make_gmem_ptr(get<0>(input_tensormaps)), Int<1>{}, Int<1>{});
|
||||
Tensor sA_tensormap = make_tensor(make_smem_ptr(&shared_tensormap.smem_tensormap_A), Int<1>{}, Int<1>{});
|
||||
Tensor gB_tensormap = make_tensor(make_gmem_ptr(get<1>(input_tensormaps)), Int<1>{}, Int<1>{});
|
||||
Tensor sB_tensormap = make_tensor(make_smem_ptr(&shared_tensormap.smem_tensormap_B), Int<1>{}, Int<1>{});
|
||||
|
||||
copy(recast<uint128_t>(gA_tensormap), recast<uint128_t>(sA_tensormap));
|
||||
copy(recast<uint128_t>(gB_tensormap), recast<uint128_t>(sB_tensormap));
|
||||
|
||||
cp_async_fence();
|
||||
cp_async_wait<0>();
|
||||
}
|
||||
|
||||
// Replace address for the global tensor (to be done by single thread)
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_replace_global_address(
|
||||
TensorMapStorage& shared_tensormap,
|
||||
Params const& mainloop_params,
|
||||
int32_t next_batch) {
|
||||
// Replacing global_address for the next batch
|
||||
cute::tma_descriptor_replace_addr_in_shared_mem(shared_tensormap.smem_tensormap_A,
|
||||
mainloop_params.ptr_A[next_batch]);
|
||||
cute::tma_descriptor_replace_addr_in_shared_mem(shared_tensormap.smem_tensormap_B,
|
||||
mainloop_params.ptr_B[next_batch]);
|
||||
}
|
||||
|
||||
// Replace dim and strides for the global tensor - used only for Grouped GEMM (to be done by single thread)
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_replace_global_tensor_properties(
|
||||
TensorMapStorage& shared_tensormap,
|
||||
Params const& mainloop_params,
|
||||
int32_t next_group,
|
||||
ProblemShape_MNKL problem_shape_mnkl) {
|
||||
const uint32_t M = get<0>(problem_shape_mnkl);
|
||||
const uint32_t N = get<1>(problem_shape_mnkl);
|
||||
const uint32_t K = get<2>(problem_shape_mnkl);
|
||||
// Only consider dimensions and strides that we need to recalculate and replace for each group
|
||||
constexpr int TensorRank = rank(ProblemShape_MNKL{}) - 1; // excluding either M or N
|
||||
static_assert(TensorRank == Int<3>{},
|
||||
"Descriptor modification for global dims & strides expects rank as 3.");
|
||||
cute::array<uint32_t, TensorRank> prob_shape_A = {1,1,1};
|
||||
cute::array<uint64_t, TensorRank> prob_stride_A = {0,0,0};
|
||||
cute::array<uint32_t, TensorRank> prob_shape_B = {1,1,1};
|
||||
cute::array<uint64_t, TensorRank> prob_stride_B = {0,0,0};
|
||||
|
||||
InternalElementA const* ptr_A = nullptr;
|
||||
Tensor tensor_a = make_tensor(ptr_A, make_shape(M,K,Int<1>{}), mainloop_params.dA[next_group]);
|
||||
|
||||
InternalElementB const* ptr_B = nullptr;
|
||||
Tensor tensor_b = make_tensor(ptr_B, make_shape(N,K,Int<1>{}), mainloop_params.dB[next_group]);
|
||||
|
||||
cute::detail::fill_tma_gmem_shape_stride(mainloop_params.tma_load_a, tensor_a,
|
||||
prob_shape_A, prob_stride_A);
|
||||
cute::detail::fill_tma_gmem_shape_stride(mainloop_params.tma_load_b, tensor_b,
|
||||
prob_shape_B, prob_stride_B);
|
||||
|
||||
cute::tma_descriptor_replace_dims_strides_in_shared_mem(shared_tensormap.smem_tensormap_A,
|
||||
prob_shape_A,
|
||||
prob_stride_A);
|
||||
cute::tma_descriptor_replace_dims_strides_in_shared_mem(shared_tensormap.smem_tensormap_B,
|
||||
prob_shape_B,
|
||||
prob_stride_B);
|
||||
}
|
||||
|
||||
template <class TensorMapA, class TensorMapB, class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_perform_update(
|
||||
TensorMapStorage& shared_tensormap,
|
||||
Params const& mainloop_params,
|
||||
cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps,
|
||||
ProblemShape_MNKL problem_shape_mnkl,
|
||||
int32_t next_batch) {
|
||||
if (cute::elect_one_sync()) {
|
||||
// Bringing tensormaps to smem
|
||||
tensormaps_fetch_to_smem(shared_tensormap, input_tensormaps);
|
||||
|
||||
// Replacing global_address for the next batch
|
||||
tensormaps_replace_global_address(shared_tensormap, mainloop_params, next_batch);
|
||||
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
// Replacing global dims and strides for the next batch
|
||||
tensormaps_replace_global_tensor_properties(shared_tensormap,
|
||||
mainloop_params, next_batch, problem_shape_mnkl);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <class TensorMapA, class TensorMapB>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_cp_fence_release (
|
||||
TensorMapStorage& shared_tensormap,
|
||||
cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps) {
|
||||
// Entire warp must do this (ie its aligned)
|
||||
tma_descriptor_cp_fence_release(get<0>(input_tensormaps), shared_tensormap.smem_tensormap_A);
|
||||
tma_descriptor_cp_fence_release(get<1>(input_tensormaps), shared_tensormap.smem_tensormap_B);
|
||||
}
|
||||
|
||||
template <class TensorMapA, class TensorMapB>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_fence_acquire(cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps) {
|
||||
cute::tma_descriptor_fence_acquire(get<0>(input_tensormaps));
|
||||
cute::tma_descriptor_fence_acquire(get<1>(input_tensormaps));
|
||||
}
|
||||
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -134,9 +134,7 @@ struct CollectiveMma<
|
||||
using TransformB = TransformB_;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<
|
||||
DispatchPolicy::Stages,
|
||||
typename DispatchPolicy::ClusterShape>;
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<DispatchPolicy::Stages>;
|
||||
using PipelineState = cutlass::PipelineState<DispatchPolicy::Stages>;
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
@@ -336,22 +334,20 @@ struct CollectiveMma<
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params)
|
||||
{
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params) {
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_a.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b.get_tma_descriptor());
|
||||
}
|
||||
|
||||
/// Set up the data needed by this collective for load and mma.
|
||||
/// Returns a tuple of tensors. The collective and the kernel layer have the contract
|
||||
/// that the tuple must contain at least two elements, with the first two elements being:
|
||||
/// Returned tuple must contain at least two elements, with the first two elements being:
|
||||
/// gA_mkl - The tma tensor, A after a local tile so it has shape (BLK_M,BLK_K,m,k,l)
|
||||
/// gB_nkl - The tma tensor, B after a local tile so it has shape (BLK_N,BLK_K,n,k,l)
|
||||
/// The rest of the tensors can be specified as needed by this collective.
|
||||
template <class ProblemShape_MNKL,
|
||||
class TileShapeMNK>
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE auto
|
||||
tile_input_tensors(ProblemShape_MNKL const& problem_shape_MNKL, Params const& mainloop_params, TileShapeMNK const& tileshape_mnk) {
|
||||
load_init(ProblemShape_MNKL const& problem_shape_MNKL, Params const& mainloop_params) const {
|
||||
using X = Underscore;
|
||||
// Separate out problem shape for convenience
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
@@ -362,8 +358,8 @@ struct CollectiveMma<
|
||||
Tensor mB_nkl = mainloop_params.tma_load_b.get_tma_tensor(make_shape(N,K,L)); // (n,k,l)
|
||||
|
||||
// Make tiled views, defer the slice
|
||||
Tensor gA_mkl = local_tile(mA_mkl, tileshape_mnk, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k,l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, tileshape_mnk, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
|
||||
Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k,l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
|
||||
|
||||
return cute::make_tuple(gA_mkl, gB_nkl);
|
||||
}
|
||||
@@ -379,7 +375,7 @@ struct CollectiveMma<
|
||||
Params const& mainloop_params,
|
||||
MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_write,
|
||||
cute::tuple<TensorA, TensorB> const& tiled_tensors,
|
||||
cute::tuple<TensorA, TensorB> const& load_inputs,
|
||||
BlockCoord const& blk_coord,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thread_idx,
|
||||
@@ -405,16 +401,16 @@ struct CollectiveMma<
|
||||
constexpr uint32_t cluster_shape_x = get<0>(ClusterShape());
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
|
||||
Tensor gA_mkl = get<0>(tiled_tensors);
|
||||
Tensor gB_nkl = get<1>(tiled_tensors);
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
auto block_tma_a = mainloop_params.tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = mainloop_params.tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
// Partition the inputs based on the current block coordinates.
|
||||
auto [m_coord, n_coord, k_coord, l_coord] = blk_coord;
|
||||
Tensor gA = gA_mkl(_,_,m_coord,_,l_coord); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nkl(_,_,n_coord,_,l_coord); // (BLK_N,BLK_K,k)
|
||||
Tensor gA = gA_mkl(_,_,m_coord,_,l_coord); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nkl(_,_,n_coord,_,l_coord); // (BLK_N,BLK_K,k)
|
||||
|
||||
// Applies the mapping from block_tma_a
|
||||
Tensor tAgA = block_tma_a.partition_S(gA); // (TMA,TMA_M,TMA_K,k)
|
||||
|
||||
+705
-120
File diff suppressed because it is too large
Load Diff
@@ -106,9 +106,7 @@ struct CollectiveMma<
|
||||
using TransformB = TransformB_;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<
|
||||
DispatchPolicy::Stages,
|
||||
typename DispatchPolicy::ClusterShape>;
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<DispatchPolicy::Stages>;
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename cutlass::PipelineState<DispatchPolicy::Stages>;
|
||||
@@ -147,8 +145,7 @@ struct CollectiveMma<
|
||||
using InternalElementA = cute::conditional_t<ConvertF32toTF32A, tfloat32_t, uint_bit_t<sizeof_bits_v<ElementA>>>;
|
||||
using InternalElementB = cute::conditional_t<ConvertF32toTF32B, tfloat32_t, uint_bit_t<sizeof_bits_v<ElementB>>>;
|
||||
|
||||
struct SharedStorage
|
||||
{
|
||||
struct SharedStorage {
|
||||
cute::array_aligned<typename TiledMma::ValTypeA, cute::cosize_v<SmemLayoutA>> smem_A;
|
||||
cute::array_aligned<typename TiledMma::ValTypeB, cute::cosize_v<SmemLayoutB>> smem_B;
|
||||
|
||||
@@ -244,13 +241,13 @@ struct CollectiveMma<
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params)
|
||||
{
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params) {
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_a.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b.get_tma_descriptor());
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Producer Perspective
|
||||
template <
|
||||
class TensorA, class TMA_LOAD_A,
|
||||
class TensorB, class TMA_LOAD_B,
|
||||
@@ -281,8 +278,8 @@ struct CollectiveMma<
|
||||
"SM90 GMMA mainloops cannot have a non-void copy atom for smem sourced instructions.");
|
||||
|
||||
SharedStorage& storage = *reinterpret_cast<SharedStorage*>(shared_memory);
|
||||
Tensor sA = make_tensor(make_smem_ptr(storage.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(storage.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
Tensor sA = make_tensor(make_smem_ptr(storage.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(storage.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
|
||||
//
|
||||
// Prepare the TMA loads for A and B
|
||||
@@ -295,11 +292,11 @@ struct CollectiveMma<
|
||||
auto block_tma_b = tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
// Applies the mapping from block_tma_a
|
||||
Tensor tAgA = block_tma_a.partition_S(gA); // (TMA,TMA_M,TMA_K,k)
|
||||
Tensor tAsA = block_tma_a.partition_D(sA); // (TMA,TMA_M,TMA_K,PIPE)
|
||||
Tensor tAgA = block_tma_a.partition_S(gA); // (TMA,TMA_M,TMA_K,k)
|
||||
Tensor tAsA = block_tma_a.partition_D(sA); // (TMA,TMA_M,TMA_K,PIPE)
|
||||
|
||||
Tensor tBgB = block_tma_b.partition_S(gB); // (TMA,TMA_N,TMA_K,k)
|
||||
Tensor tBsB = block_tma_b.partition_D(sB); // (TMA,TMA_N,TMA_K,PIPE)
|
||||
Tensor tBgB = block_tma_b.partition_S(gB); // (TMA,TMA_N,TMA_K,k)
|
||||
Tensor tBsB = block_tma_b.partition_D(sB); // (TMA,TMA_N,TMA_K,PIPE)
|
||||
|
||||
//
|
||||
// Prepare TMA membars and PREFETCH
|
||||
@@ -335,9 +332,7 @@ struct CollectiveMma<
|
||||
params.is_leader = warp_group_thread_idx == 0;
|
||||
params.num_consumers = NumThreadsPerWarpGroup;
|
||||
|
||||
MainloopPipeline pipeline(
|
||||
storage.pipeline_storage,
|
||||
params);
|
||||
MainloopPipeline pipeline(storage.pipeline_storage, params, ClusterShape{});
|
||||
|
||||
// State variables used for iterating the circular buffer
|
||||
// smem_pipe_read / release is used by the consumer of SMEM data - i.e MMA
|
||||
|
||||
@@ -107,9 +107,7 @@ struct CollectiveMma<
|
||||
using TransformB = TransformB_;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<
|
||||
DispatchPolicy::Stages,
|
||||
typename DispatchPolicy::ClusterShape>;
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<DispatchPolicy::Stages>;
|
||||
using PipelineState = cutlass::PipelineState<DispatchPolicy::Stages>;
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
@@ -255,22 +253,20 @@ struct CollectiveMma<
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params)
|
||||
{
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params) {
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_a.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b.get_tma_descriptor());
|
||||
}
|
||||
|
||||
/// Set up the data needed by this collective for load and mma.
|
||||
/// Returns a tuple of tensors. The collective and the kernel layer have the contract
|
||||
/// that the tuple must contain at least two elements, with the first two elements being:
|
||||
/// Returned tuple must contain at least two elements, with the first two elements being:
|
||||
/// gA_mkl - The tma tensor, A after a local tile so it has shape (BLK_M,BLK_K,m,k,l)
|
||||
/// gB_nkl - The tma tensor, B after a local tile so it has shape (BLK_N,BLK_K,n,k,l)
|
||||
/// The rest of the tensors can be specified as needed by this collective.
|
||||
template <class ProblemShape_MNKL,
|
||||
class TileShapeMNK>
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE auto
|
||||
tile_input_tensors(ProblemShape_MNKL const& problem_shape_MNKL, Params const& mainloop_params, TileShapeMNK const& tileshape_mnk) {
|
||||
load_init(ProblemShape_MNKL const& problem_shape_MNKL, Params const& mainloop_params) const {
|
||||
using X = Underscore;
|
||||
// Separate out problem shape for convenience
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
@@ -281,8 +277,8 @@ struct CollectiveMma<
|
||||
Tensor mB_nkl = mainloop_params.tma_load_b.get_tma_tensor(make_shape(N,K,L)); // (n,k,l)
|
||||
|
||||
// Make tiled views, defer the slice
|
||||
Tensor gA_mkl = local_tile(mA_mkl, tileshape_mnk, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k,l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, tileshape_mnk, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
|
||||
Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k,l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
|
||||
|
||||
return cute::make_tuple(gA_mkl, gB_nkl);
|
||||
}
|
||||
@@ -298,14 +294,12 @@ struct CollectiveMma<
|
||||
Params const& mainloop_params,
|
||||
MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_write,
|
||||
cute::tuple<TensorA, TensorB> const& tiled_tensors,
|
||||
cute::tuple<TensorA, TensorB> const& load_inputs,
|
||||
BlockCoord const& blk_coord,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thread_idx,
|
||||
uint32_t block_rank_in_cluster,
|
||||
TensorStorage& shared_tensors)
|
||||
{
|
||||
|
||||
TensorStorage& shared_tensors) {
|
||||
using namespace cute;
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
@@ -322,16 +316,16 @@ struct CollectiveMma<
|
||||
constexpr uint32_t cluster_shape_x = get<0>(typename DispatchPolicy::ClusterShape());
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
|
||||
Tensor gA_mkl = get<0>(tiled_tensors);
|
||||
Tensor gB_nkl = get<1>(tiled_tensors);
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
auto block_tma_a = mainloop_params.tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = mainloop_params.tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
// Partition the inputs based on the current block coordinates.
|
||||
auto [m_coord, n_coord, k_coord, l_coord] = blk_coord;
|
||||
Tensor gA = gA_mkl(_,_,m_coord,_,l_coord); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nkl(_,_,n_coord,_,l_coord); // (BLK_N,BLK_K,k)
|
||||
Tensor gA = gA_mkl(_,_,m_coord,_,l_coord); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nkl(_,_,n_coord,_,l_coord); // (BLK_N,BLK_K,k)
|
||||
|
||||
// Applies the mapping from block_tma_a
|
||||
Tensor tAgA = block_tma_a.partition_S(gA); // (TMA,TMA_M,TMA_K,k)
|
||||
@@ -386,10 +380,7 @@ struct CollectiveMma<
|
||||
|
||||
/// Perform a Producer Epilogue to prevent early exit of blocks in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(
|
||||
MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_write)
|
||||
{
|
||||
load_tail(MainloopPipeline pipeline, PipelineState smem_pipe_write) {
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
@@ -418,8 +409,7 @@ struct CollectiveMma<
|
||||
int k_tile_count,
|
||||
int thread_idx,
|
||||
TensorStorage& shared_tensors,
|
||||
Params const& mainloop_params)
|
||||
{
|
||||
Params const& mainloop_params) {
|
||||
using namespace cute;
|
||||
|
||||
static_assert(is_rmem<FrgTensorC>::value, "C tensor must be rmem resident.");
|
||||
|
||||
@@ -108,9 +108,7 @@ struct CollectiveMma<
|
||||
using TransformB = TransformB_;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<
|
||||
DispatchPolicy::Stages,
|
||||
typename DispatchPolicy::ClusterShape>;
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<DispatchPolicy::Stages>;
|
||||
using PipelineState = cutlass::PipelineState<DispatchPolicy::Stages>;
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
@@ -261,14 +259,12 @@ struct CollectiveMma<
|
||||
|
||||
/// Set up the data needed by this collective for load and mma.
|
||||
/// Returns a tuple of tensors. The collective and the kernel layer have the contract
|
||||
/// that the tuple must contain at least two elements, with the first two elements being:
|
||||
/// Returned tuple must contain at least two elements, with the first two elements being:
|
||||
/// gA_mkl - The tma tensor, A after a local tile so it has shape (BLK_M,BLK_K,m,k,l)
|
||||
/// gB_nkl - The tma tensor, B after a local tile so it has shape (BLK_N,BLK_K,n,k,l)
|
||||
/// The rest of the tensors can be specified as needed by this collective.
|
||||
template <class ProblemShape_MNKL,
|
||||
class TileShapeMNK>
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE auto
|
||||
tile_input_tensors(ProblemShape_MNKL const& problem_shape_MNKL, Params const& mainloop_params, TileShapeMNK const& tileshape_mnk) {
|
||||
load_init(ProblemShape_MNKL const& problem_shape_MNKL, Params const& mainloop_params) const {
|
||||
using X = Underscore;
|
||||
// Separate out problem shape for convenience
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
@@ -279,8 +275,8 @@ struct CollectiveMma<
|
||||
Tensor mB_nkl = mainloop_params.tma_load_b.get_tma_tensor(make_shape(N,K,L)); // (n,k,l)
|
||||
|
||||
// Make tiled views, defer the slice
|
||||
Tensor gA_mkl = local_tile(mA_mkl, tileshape_mnk, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k,l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, tileshape_mnk, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
|
||||
Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k,l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
|
||||
|
||||
return cute::make_tuple(gA_mkl, gB_nkl);
|
||||
}
|
||||
@@ -296,7 +292,7 @@ struct CollectiveMma<
|
||||
Params const& mainloop_params,
|
||||
MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_write,
|
||||
cute::tuple<TensorA, TensorB> const& tiled_tensors,
|
||||
cute::tuple<TensorA, TensorB> const& load_inputs,
|
||||
BlockCoord const& blk_coord,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thread_idx,
|
||||
@@ -320,8 +316,8 @@ struct CollectiveMma<
|
||||
constexpr uint32_t cluster_shape_x = get<0>(ClusterShape());
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
|
||||
Tensor gA_mkl = get<0>(tiled_tensors);
|
||||
Tensor gB_nkl = get<1>(tiled_tensors);
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
auto block_tma_a = mainloop_params.tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = mainloop_params.tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
@@ -42,6 +42,7 @@
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/detail/layout.hpp"
|
||||
#include "cutlass/detail/mma.hpp"
|
||||
#include "cutlass/cuda_host_adapter.hpp"
|
||||
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
#include "cutlass/cluster_launch.hpp"
|
||||
@@ -106,9 +107,12 @@ public:
|
||||
using LayoutC = gemm::detail::StrideToLayoutTagC_t<typename GemmKernel::StrideC>;
|
||||
using LayoutD = gemm::detail::StrideToLayoutTagC_t<typename GemmKernel::StrideD>;
|
||||
|
||||
// NOTE: 3.0 kernels do not support complex transforms for now ...
|
||||
static ComplexTransform const kTransformA = ComplexTransform::kNone;
|
||||
static ComplexTransform const kTransformB = ComplexTransform::kNone;
|
||||
static bool const kEnableCudaHostAdapter = CUTLASS_ENABLE_CUDA_HOST_ADAPTER;
|
||||
|
||||
static ComplexTransform const kTransformA = cute::is_same_v<typename GemmKernel::CollectiveMainloop::TransformA, cute::conjugate> ?
|
||||
ComplexTransform::kConjugate : ComplexTransform::kNone;
|
||||
static ComplexTransform const kTransformB = cute::is_same_v<typename GemmKernel::CollectiveMainloop::TransformB, cute::conjugate> ?
|
||||
ComplexTransform::kConjugate : ComplexTransform::kNone;
|
||||
|
||||
// Legacy: Assume MultiplyAdd only since we do not use this tag type in 3.0
|
||||
using MathOperator = cutlass::arch::OpMultiplyAdd;
|
||||
@@ -141,17 +145,17 @@ public:
|
||||
static int const kThreadCount = GemmKernel::MaxThreadsPerBlock;
|
||||
|
||||
// Warp shape is not a primary API type in 3.x
|
||||
// But we can best approximate it by inspecting the TiledMma::TiledShape_MNK
|
||||
// But we can best approximate it by inspecting the TiledMma
|
||||
// For this, we make the assumption that we always have 4 warps along M, and rest along N, none along K
|
||||
// We also always round up the warp count to 4 if the tiled mma is smaller than 128 threads
|
||||
static constexpr int WarpsInMma = cute::max(4, cute::size(typename GemmKernel::TiledMma{}) / 32);
|
||||
static constexpr int WarpsInMma = cute::max(4, CUTE_STATIC_V(cute::size(typename GemmKernel::TiledMma{})) / 32);
|
||||
static constexpr int WarpsInMmaM = 4;
|
||||
static constexpr int WarpsInMmaN = cute::ceil_div(WarpsInMma, WarpsInMmaM);
|
||||
using WarpCount = cutlass::gemm::GemmShape<WarpsInMmaM, WarpsInMmaN, 1>;
|
||||
using WarpShape = cutlass::gemm::GemmShape<
|
||||
cute::size<0>(typename CollectiveMainloop::TiledMma::TiledShape_MNK{}) / WarpsInMmaM,
|
||||
cute::size<1>(typename CollectiveMainloop::TiledMma::TiledShape_MNK{}) / WarpsInMmaN,
|
||||
cute::size<2>(typename CollectiveMainloop::TiledMma::TiledShape_MNK{})>;
|
||||
CUTE_STATIC_V(cute::tile_size<0>(typename CollectiveMainloop::TiledMma{})) / WarpsInMmaM,
|
||||
CUTE_STATIC_V(cute::tile_size<1>(typename CollectiveMainloop::TiledMma{})) / WarpsInMmaN,
|
||||
CUTE_STATIC_V(cute::tile_size<2>(typename CollectiveMainloop::TiledMma{}))>;
|
||||
|
||||
static int constexpr kStages = CollectiveMainloop::DispatchPolicy::Stages;
|
||||
|
||||
@@ -270,7 +274,12 @@ public:
|
||||
|
||||
/// Initializes GEMM state from arguments.
|
||||
Status
|
||||
initialize(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr) {
|
||||
initialize(
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::initialize() - workspace "
|
||||
<< workspace << ", stream: " << (stream ? "non-null" : "null"));
|
||||
|
||||
@@ -283,20 +292,33 @@ public:
|
||||
// Initialize the Params structure
|
||||
params_ = GemmKernel::to_underlying_arguments(args, workspace);
|
||||
|
||||
// account for dynamic smem capacity if needed
|
||||
int smem_size = GemmKernel::SharedStorageSize;
|
||||
if (smem_size >= (48 << 10)) {
|
||||
CUTLASS_TRACE_HOST(" Setting smem size to " << smem_size);
|
||||
cudaError_t result = cudaFuncSetAttribute(
|
||||
device_kernel<GemmKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
if (cudaSuccess != result) {
|
||||
result = cudaGetLastError(); // to clear the error bit
|
||||
CUTLASS_TRACE_HOST(" cudaFuncSetAttribute() returned error: " << cudaGetErrorString(result));
|
||||
return Status::kErrorInternal;
|
||||
// Don't set the function attributes - require the CudaHostAdapter to set it.
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else {
|
||||
//
|
||||
// Account for dynamic smem capacity if needed
|
||||
//
|
||||
int smem_size = GemmKernel::SharedStorageSize;
|
||||
|
||||
CUTLASS_ASSERT(cuda_adapter == nullptr);
|
||||
|
||||
if (smem_size >= (48 << 10)) {
|
||||
CUTLASS_TRACE_HOST(" Setting smem size to " << smem_size);
|
||||
cudaError_t result = cudaFuncSetAttribute(
|
||||
device_kernel<GemmKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
if (cudaSuccess != result) {
|
||||
result = cudaGetLastError(); // to clear the error bit
|
||||
CUTLASS_TRACE_HOST(" cudaFuncSetAttribute() returned error: " << cudaGetErrorString(result));
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
@@ -317,7 +339,7 @@ public:
|
||||
/// Primary run() entry point API that is static allowing users to create and manage their own params.
|
||||
/// Supplied params struct must be construct by calling GemmKernel::to_underling_arguments()
|
||||
static Status
|
||||
run(Params& params, cudaStream_t stream = nullptr) {
|
||||
run(Params& params, cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::run()");
|
||||
dim3 const block = GemmKernel::get_block_shape();
|
||||
dim3 const grid = get_grid_shape(params);
|
||||
@@ -331,13 +353,54 @@ public:
|
||||
dim3 cluster(cute::size<0>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<1>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<2>(typename GemmKernel::DispatchPolicy::ClusterShape{}));
|
||||
void const* kernel = (void const*) device_kernel<GemmKernel>;
|
||||
|
||||
void* kernel_params[] = {¶ms};
|
||||
launch_result = ClusterLauncher::launch(grid, cluster, block, smem_size, stream, kernel, kernel_params);
|
||||
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
//
|
||||
// Use the cuda host adapter
|
||||
//
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
if (cuda_adapter) {
|
||||
|
||||
launch_result = cuda_adapter->launch(
|
||||
grid, cluster, block, smem_size, stream, kernel_params, 0
|
||||
);
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
else {
|
||||
|
||||
CUTLASS_ASSERT(cuda_adapter == nullptr);
|
||||
void const* kernel = (void const*) device_kernel<GemmKernel>;
|
||||
|
||||
launch_result = ClusterLauncher::launch(
|
||||
grid, cluster, block, smem_size, stream, kernel, kernel_params);
|
||||
|
||||
}
|
||||
}
|
||||
else {
|
||||
launch_result = Status::kSuccess;
|
||||
device_kernel<GemmKernel><<<grid, block, smem_size, stream>>>(params);
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
if (cuda_adapter) {
|
||||
void* kernel_params[] = {¶ms};
|
||||
|
||||
launch_result = cuda_adapter->launch(
|
||||
grid, block, smem_size, stream, kernel_params, 0
|
||||
);
|
||||
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
else {
|
||||
CUTLASS_ASSERT(cuda_adapter == nullptr);
|
||||
device_kernel<GemmKernel><<<grid, block, smem_size, stream>>>(params);
|
||||
}
|
||||
}
|
||||
|
||||
cudaError_t result = cudaGetLastError();
|
||||
@@ -356,18 +419,27 @@ public:
|
||||
|
||||
/// Launches the kernel after first constructing Params internal state from supplied arguments.
|
||||
Status
|
||||
run(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr) {
|
||||
run(
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr
|
||||
) {
|
||||
Status status = initialize(args, workspace, stream);
|
||||
if (Status::kSuccess == status) {
|
||||
status = run(params_, stream);
|
||||
status = run(params_, stream, cuda_adapter);
|
||||
}
|
||||
return status;
|
||||
}
|
||||
|
||||
/// Launches the kernel after first constructing Params internal state from supplied arguments.
|
||||
Status
|
||||
operator()(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr) {
|
||||
return run(args, workspace, stream);
|
||||
operator()(
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
return run(args, workspace, stream, cuda_adapter);
|
||||
}
|
||||
|
||||
/// Overload that allows a user to re-launch the same kernel without updating internal params struct.
|
||||
@@ -387,7 +459,7 @@ public:
|
||||
////////////////////////////// CUTLASS 2.x API /////////////////////////////////
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename GemmKernel_>
|
||||
template <class GemmKernel_>
|
||||
class GemmUniversalAdapter<
|
||||
GemmKernel_,
|
||||
cute::enable_if_t<not gemm::detail::IsCutlass3GemmKernel<GemmKernel_>::value>>
|
||||
@@ -501,9 +573,14 @@ public:
|
||||
}
|
||||
|
||||
/// Initializes GEMM state from arguments.
|
||||
Status initialize(Arguments const &args, void *workspace = nullptr, cudaStream_t stream = nullptr) {
|
||||
Status initialize(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr
|
||||
) {
|
||||
|
||||
return underlying_operator_.initialize(to_underlying_arguments(args), workspace, stream);
|
||||
return underlying_operator_.initialize(to_underlying_arguments(args), workspace, stream, cuda_adapter);
|
||||
}
|
||||
|
||||
/// Lightweight update given a subset of arguments.
|
||||
@@ -513,13 +590,18 @@ public:
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status run(cudaStream_t stream = nullptr) {
|
||||
Status run(
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
return underlying_operator_.run(stream);
|
||||
return underlying_operator_.run(stream, cuda_adapter);
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status operator()(cudaStream_t stream = nullptr) {
|
||||
Status operator()(
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
return run(stream);
|
||||
}
|
||||
|
||||
@@ -527,12 +609,13 @@ public:
|
||||
Status operator()(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) {
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
Status status = initialize(args, workspace, stream);
|
||||
Status status = initialize(args, workspace, stream, cuda_adapter);
|
||||
|
||||
if (status == Status::kSuccess) {
|
||||
status = run(stream);
|
||||
status = run(stream, cuda_adapter);
|
||||
}
|
||||
|
||||
return status;
|
||||
|
||||
@@ -46,6 +46,7 @@
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/arch/arch.h"
|
||||
#include "cutlass/device_kernel.h"
|
||||
#include "cutlass/cuda_host_adapter.hpp"
|
||||
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/kernel/gemm_universal.h"
|
||||
@@ -69,6 +70,8 @@ class GemmUniversalBase {
|
||||
public:
|
||||
|
||||
using GemmKernel = GemmKernel_;
|
||||
static bool const kEnableCudaHostAdapter = CUTLASS_ENABLE_CUDA_HOST_ADAPTER;
|
||||
|
||||
using ThreadblockShape = typename GemmKernel::Mma::Shape;
|
||||
|
||||
using ElementA = typename GemmKernel::ElementA;
|
||||
@@ -295,7 +298,8 @@ public:
|
||||
Status initialize(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr)
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr)
|
||||
{
|
||||
CUTLASS_TRACE_HOST("GemmUniversalBase::initialize() - workspace "
|
||||
<< workspace << ", stream: " << (stream ? "non-null" : "null"));
|
||||
@@ -323,9 +327,8 @@ public:
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status run(cudaStream_t stream = nullptr)
|
||||
Status run(cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr)
|
||||
{
|
||||
CUTLASS_TRACE_HOST("GemmUniversalBase::run()");
|
||||
|
||||
@@ -339,13 +342,27 @@ public:
|
||||
"block: (" << block << "), "
|
||||
"SMEM: (" << smem_size_ << ")");
|
||||
|
||||
Kernel2<GemmKernel><<<grid, block, smem_size_, stream>>>(params_);
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
if (cuda_adapter) {
|
||||
void* kernel_params[] = {¶ms_};
|
||||
return cuda_adapter->launch(grid, block, smem_size_, stream, kernel_params, 0);
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
else {
|
||||
CUTLASS_ASSERT(cuda_adapter == nullptr);
|
||||
|
||||
// Query for errors
|
||||
cudaError_t result = cudaGetLastError();
|
||||
if (result != cudaSuccess) {
|
||||
CUTLASS_TRACE_HOST(" grid launch failed with error " << cudaGetErrorString(result));
|
||||
return Status::kErrorInternal;
|
||||
Kernel2<GemmKernel><<<grid, block, smem_size_, stream>>>(params_);
|
||||
|
||||
// Query for errors
|
||||
cudaError_t result = cudaGetLastError();
|
||||
if (result != cudaSuccess) {
|
||||
CUTLASS_TRACE_HOST(" grid launch failed with error " << cudaGetErrorString(result));
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
return Status::kSuccess;
|
||||
@@ -363,12 +380,13 @@ public:
|
||||
Status operator()(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr)
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr)
|
||||
{
|
||||
Status status = initialize(args, workspace, stream);
|
||||
|
||||
if (status == Status::kSuccess) {
|
||||
status = run(stream);
|
||||
status = run(stream, cuda_adapter);
|
||||
}
|
||||
|
||||
return status;
|
||||
|
||||
@@ -53,6 +53,8 @@ struct KernelTma { };
|
||||
struct KernelTmaWarpSpecialized { };
|
||||
struct KernelTmaWarpSpecializedPingpong { };
|
||||
struct KernelTmaWarpSpecializedCooperative { };
|
||||
struct KernelArrayTmaWarpSpecializedCooperative { };
|
||||
struct KernelGroupTmaWarpSpecializedCooperative { };
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -65,6 +67,8 @@ struct KernelTmaWarpSpecializedCooperative { };
|
||||
struct KernelTmaWarpSpecializedFP8FastAccum : KernelTmaWarpSpecialized { };
|
||||
struct KernelTmaWarpSpecializedPingpongFP8FastAccum : KernelTmaWarpSpecializedPingpong { };
|
||||
struct KernelTmaWarpSpecializedCooperativeFP8FastAccum: KernelTmaWarpSpecializedCooperative { };
|
||||
struct KernelArrayTmaWarpSpecializedCooperativeFP8FastAccum : KernelArrayTmaWarpSpecializedCooperative { };
|
||||
struct KernelGroupTmaWarpSpecializedCooperativeFP8FastAccum : KernelGroupTmaWarpSpecializedCooperative { };
|
||||
|
||||
// Policies to opt into mixed type GEMMs
|
||||
struct KernelTmaWarpSpecializedMixedInput : KernelTmaWarpSpecialized { };
|
||||
@@ -225,6 +229,23 @@ struct MainloopSm90TmaGmmaWarpSpecializedFP8
|
||||
"KernelSchedule must be one of the warp specialized policies");
|
||||
};
|
||||
|
||||
// n-buffer in smem (Hopper TMA), pipelined with Hopper GMMA and TMA, Warp specialized dynamic schedule for Ptr-Array and Grouped Gemm
|
||||
template<
|
||||
int Stages_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
class KernelSchedule = KernelGroupTmaWarpSpecializedCooperative
|
||||
>
|
||||
struct MainloopSm90ArrayTmaGmmaWarpSpecialized {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm90;
|
||||
using Schedule = KernelSchedule;
|
||||
static_assert(
|
||||
cute::is_base_of_v<KernelArrayTmaWarpSpecializedCooperative, KernelSchedule> ||
|
||||
cute::is_base_of_v<KernelGroupTmaWarpSpecializedCooperative, KernelSchedule>,
|
||||
"KernelSchedule must be one of the Ptr-Array or Grouped Gemm TMA Warp Specialized Cooperative policies");
|
||||
};
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm
|
||||
|
||||
@@ -69,6 +69,7 @@ enum class GemmUniversalMode {
|
||||
kGemmSplitKParallel,
|
||||
kBatched,
|
||||
kArray,
|
||||
kGrouped,
|
||||
kInvalid
|
||||
};
|
||||
|
||||
|
||||
@@ -0,0 +1,111 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2023 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief This file contains definitions and utility functions for describing problem shapes
|
||||
for 3.x Ptr-Array GEMMs and Grouped GEMMs.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/tensor_coord.h"
|
||||
|
||||
#include "cute/container/array.hpp"
|
||||
|
||||
#if ! defined(__CUDACC_RTC__)
|
||||
#include <initializer_list>
|
||||
#endif
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <class ProblemShape_>
|
||||
struct GroupProblemShape {
|
||||
using UnderlyingProblemShape = ProblemShape_;
|
||||
int32_t num_groups = 1;
|
||||
UnderlyingProblemShape* problem_shapes = nullptr;
|
||||
UnderlyingProblemShape const* host_problem_shapes = nullptr;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t groups() const { return num_groups; }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
UnderlyingProblemShape const
|
||||
get_problem_shape(int32_t group_idx) const {
|
||||
return problem_shapes[group_idx];
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
UnderlyingProblemShape const
|
||||
get_host_problem_shape(int32_t group_idx) const {
|
||||
return host_problem_shapes[group_idx];
|
||||
}
|
||||
};
|
||||
|
||||
template <class ProblemShape_>
|
||||
class ArrayProblemShape {
|
||||
public:
|
||||
using UnderlyingProblemShape = ProblemShape_;
|
||||
|
||||
ArrayProblemShape() = default;
|
||||
ArrayProblemShape(UnderlyingProblemShape ps) : problem_shape_(ps) {}
|
||||
|
||||
// Num of groups for Ptr-Array GEMM always remain one, just the number of batches (l) can vary
|
||||
// This is just to maintain uniformity with GroupProblemShape
|
||||
constexpr int32_t groups() const { return 1; }
|
||||
|
||||
UnderlyingProblemShape* problem_shapes() const {
|
||||
return &problem_shape_;
|
||||
}
|
||||
UnderlyingProblemShape const* host_problem_shapes() const {
|
||||
return &problem_shape_;
|
||||
}
|
||||
|
||||
// This is just to maintain uniformity with GroupProblemShape
|
||||
CUTLASS_HOST_DEVICE
|
||||
UnderlyingProblemShape const
|
||||
get_problem_shape(int32_t /* unused */ = 0) const {
|
||||
return problem_shape_;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
UnderlyingProblemShape const
|
||||
get_host_problem_shape(int32_t /* unused */ = 0) const {
|
||||
return problem_shape_;
|
||||
}
|
||||
private:
|
||||
UnderlyingProblemShape problem_shape_{};
|
||||
};
|
||||
|
||||
} // namespace cutlass::gemm
|
||||
@@ -851,7 +851,9 @@ protected:
|
||||
}
|
||||
ptr_D += tile_work.tiled_coord.k() * params.batch_stride_D;
|
||||
if (ptr_Tensor) {
|
||||
ptr_Tensor += tile_work.tiled_coord.k() * params.batch_stride_Tensor;
|
||||
ptr_Tensor = ReferenceFactory<typename Epilogue::ElementTensor>::add_pointer_offset(
|
||||
ptr_Tensor,
|
||||
tile_work.tiled_coord.k() * params.batch_stride_Tensor);
|
||||
}
|
||||
if (ptr_Vector) {
|
||||
ptr_Vector += tile_work.tiled_coord.k() * params.batch_stride_Vector;
|
||||
@@ -2024,7 +2026,9 @@ protected:
|
||||
ptr_C += tile_work.tiled_coord.k() * params.batch_stride_C;
|
||||
ptr_D += tile_work.tiled_coord.k() * params.batch_stride_D;
|
||||
if (ptr_Tensor) {
|
||||
ptr_Tensor += tile_work.tiled_coord.k() * params.batch_stride_Tensor;
|
||||
ptr_Tensor = ReferenceFactory<typename Epilogue::ElementTensor>::add_pointer_offset(
|
||||
ptr_Tensor,
|
||||
tile_work.tiled_coord.k() * params.batch_stride_Tensor);
|
||||
}
|
||||
if (ptr_Vector) {
|
||||
ptr_Vector += tile_work.tiled_coord.k() * params.batch_stride_Vector;
|
||||
|
||||
@@ -67,9 +67,9 @@ class GemmUniversal<
|
||||
Epilogue_,
|
||||
ThreadblockSwizzle_,
|
||||
void,
|
||||
// 3.x kernels use the first template argument to define the ProblemShape tuple
|
||||
// 3.x kernels use the first template argument to define the ProblemShape
|
||||
// We use this invariant to SFINAE dispatch against either the 2.x API or the 3.x API
|
||||
cute::enable_if_t<not cute::is_tuple<Mma_>::value>
|
||||
cute::enable_if_t<not (cute::is_tuple<Mma_>::value || IsCutlass3ArrayKernel<Mma_>::value)>
|
||||
> {
|
||||
public:
|
||||
|
||||
|
||||
@@ -61,6 +61,19 @@ template <
|
||||
>
|
||||
class GemmUniversal;
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// In cases where ProblemShape is not a tuple, this is used to check if the
|
||||
// underlying problem shape type is aliased within or not.
|
||||
// Used for dispatching GemmUniversal to 2.x API or 3.x API
|
||||
template <class ProblemShape, class = void>
|
||||
struct IsCutlass3ArrayKernel : cute::false_type { };
|
||||
|
||||
template <typename ProblemShape>
|
||||
struct IsCutlass3ArrayKernel<ProblemShape, cute::void_t<typename ProblemShape::UnderlyingProblemShape>>
|
||||
: cute::true_type { };
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::kernel
|
||||
@@ -75,4 +88,5 @@ class GemmUniversal;
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_tma_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_tma_warpspecialized_pingpong.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_tma_warpspecialized_cooperative.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_array_tma_warpspecialized_cooperative.hpp"
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -42,7 +42,7 @@
|
||||
#include "cutlass/complex.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
#include "cutlass/gemm/kernel/params_universal_base.h"
|
||||
|
||||
#include "cutlass/subbyte_reference.h"
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -676,7 +676,9 @@ public:
|
||||
}
|
||||
ptr_D += threadblock_tile_offset.k() * params.batch_stride_D;
|
||||
if (ptr_Tensor) {
|
||||
ptr_Tensor += threadblock_tile_offset.k() * params.batch_stride_Tensor;
|
||||
ptr_Tensor = ReferenceFactory<typename Epilogue::ElementTensor>::add_pointer_offset(
|
||||
ptr_Tensor,
|
||||
threadblock_tile_offset.k() * params.batch_stride_Tensor);
|
||||
}
|
||||
if (ptr_Vector) {
|
||||
ptr_Vector += threadblock_tile_offset.k() * params.batch_stride_Vector;
|
||||
@@ -1387,7 +1389,9 @@ public:
|
||||
ptr_C += threadblock_tile_offset.k() * params.batch_stride_C;
|
||||
ptr_D += threadblock_tile_offset.k() * params.batch_stride_D;
|
||||
if (ptr_Tensor) {
|
||||
ptr_Tensor += threadblock_tile_offset.k() * params.batch_stride_Tensor;
|
||||
ptr_Tensor = ReferenceFactory<typename Epilogue::ElementTensor>::add_pointer_offset(
|
||||
ptr_Tensor,
|
||||
threadblock_tile_offset.k() * params.batch_stride_Tensor);
|
||||
}
|
||||
if (ptr_Vector) {
|
||||
ptr_Vector += threadblock_tile_offset.k() * params.batch_stride_Vector;
|
||||
|
||||
@@ -217,6 +217,8 @@ public:
|
||||
// Only used by device-level operator
|
||||
GemmCoord *host_problem_sizes;
|
||||
|
||||
bool allow_early_exit;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -235,7 +237,8 @@ public:
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr),
|
||||
host_problem_sizes(nullptr)
|
||||
host_problem_sizes(nullptr),
|
||||
allow_early_exit(false)
|
||||
{
|
||||
|
||||
}
|
||||
@@ -256,7 +259,8 @@ public:
|
||||
typename LayoutB::Stride::LongIndex *ldb,
|
||||
typename LayoutC::Stride::LongIndex *ldc,
|
||||
typename LayoutC::Stride::LongIndex *ldd,
|
||||
GemmCoord *host_problem_sizes=nullptr
|
||||
GemmCoord *host_problem_sizes=nullptr,
|
||||
bool allow_early_exit=false
|
||||
):
|
||||
mode(mode),
|
||||
problem_sizes(problem_sizes),
|
||||
@@ -271,7 +275,8 @@ public:
|
||||
ldb(ldb),
|
||||
ldc(ldc),
|
||||
ldd(ldd),
|
||||
host_problem_sizes(host_problem_sizes)
|
||||
host_problem_sizes(host_problem_sizes),
|
||||
allow_early_exit(allow_early_exit)
|
||||
{
|
||||
|
||||
}
|
||||
@@ -303,6 +308,7 @@ public:
|
||||
typename LayoutC::Stride::LongIndex *ldc;
|
||||
typename LayoutC::Stride::LongIndex *ldd;
|
||||
|
||||
bool allow_early_exit;
|
||||
|
||||
//
|
||||
// Methods
|
||||
@@ -318,7 +324,8 @@ public:
|
||||
lda(nullptr),
|
||||
ldb(nullptr),
|
||||
ldc(nullptr),
|
||||
ldd(nullptr)
|
||||
ldd(nullptr),
|
||||
allow_early_exit(false)
|
||||
{ }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -333,7 +340,8 @@ public:
|
||||
lda(args.lda),
|
||||
ldb(args.ldb),
|
||||
ldc(args.ldc),
|
||||
ldd(args.ldd)
|
||||
ldd(args.ldd),
|
||||
allow_early_exit(args.allow_early_exit)
|
||||
{
|
||||
|
||||
}
|
||||
@@ -388,6 +396,12 @@ public:
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
// Early exit following LAPACK's definition
|
||||
if (params.allow_early_exit &&
|
||||
(params.output_op.alpha == ElementC(0)) && (params.output_op.beta == ElementC(1))) {
|
||||
return;
|
||||
}
|
||||
|
||||
//
|
||||
// Problem visitor.
|
||||
//
|
||||
|
||||
@@ -140,6 +140,8 @@ public:
|
||||
typename LayoutC::Stride::Index ldc;
|
||||
typename LayoutC::Stride::Index ldd;
|
||||
|
||||
bool allow_early_exit;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -147,7 +149,8 @@ public:
|
||||
Arguments():
|
||||
mode(GemmUniversalMode::kGemm),
|
||||
batch_count(1),
|
||||
ptr_A(nullptr), ptr_B(nullptr), ptr_C(nullptr), ptr_D(nullptr) { }
|
||||
ptr_A(nullptr), ptr_B(nullptr), ptr_C(nullptr), ptr_D(nullptr),
|
||||
allow_early_exit(false) { }
|
||||
|
||||
/// constructs an arguments structure
|
||||
Arguments(
|
||||
@@ -166,7 +169,8 @@ public:
|
||||
typename LayoutA::Stride::Index lda,
|
||||
typename LayoutB::Stride::Index ldb,
|
||||
typename LayoutC::Stride::Index ldc,
|
||||
typename LayoutC::Stride::Index ldd
|
||||
typename LayoutC::Stride::Index ldd,
|
||||
bool allow_early_exit = false
|
||||
):
|
||||
mode(mode),
|
||||
problem_size(problem_size),
|
||||
@@ -174,7 +178,8 @@ public:
|
||||
epilogue(epilogue),
|
||||
ptr_A(ptr_A), ptr_B(ptr_B), ptr_C(ptr_C), ptr_D(ptr_D),
|
||||
batch_stride_A(batch_stride_A), batch_stride_C(batch_stride_C), batch_stride_D(batch_stride_D),
|
||||
lda(lda), ldb(ldb), ldc(ldc), ldd(ldd) {
|
||||
lda(lda), ldb(ldb), ldc(ldc), ldd(ldd),
|
||||
allow_early_exit(allow_early_exit) {
|
||||
|
||||
}
|
||||
|
||||
@@ -231,6 +236,8 @@ public:
|
||||
|
||||
int *semaphore;
|
||||
|
||||
bool allow_early_exit;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -255,7 +262,8 @@ public:
|
||||
batch_stride_B(0),
|
||||
batch_stride_C(0),
|
||||
batch_stride_D(0),
|
||||
semaphore(nullptr) { }
|
||||
semaphore(nullptr),
|
||||
allow_early_exit(false) { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(
|
||||
@@ -285,7 +293,8 @@ public:
|
||||
batch_stride_B(args.batch_stride_B),
|
||||
batch_stride_C(args.batch_stride_C),
|
||||
batch_stride_D(args.batch_stride_D),
|
||||
semaphore(static_cast<int *>(workspace)) {
|
||||
semaphore(static_cast<int *>(workspace)),
|
||||
allow_early_exit(args.allow_early_exit) {
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -347,6 +356,12 @@ public:
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
// Early exit following LAPACK's definition
|
||||
if (params.allow_early_exit &&
|
||||
(params.output_op.alpha == ElementC(0)) && (params.output_op.beta == ElementC(1))) {
|
||||
return;
|
||||
}
|
||||
|
||||
// Compute threadblock location
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
|
||||
@@ -125,6 +125,8 @@ public:
|
||||
typename LayoutC::Stride::Index ldc;
|
||||
typename LayoutC::Stride::Index ldd;
|
||||
|
||||
bool allow_early_exit;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -132,7 +134,8 @@ public:
|
||||
Arguments():
|
||||
mode(GemmUniversalMode::kGemm),
|
||||
batch_count(1),
|
||||
ptr_A(nullptr), ptr_C(nullptr), ptr_D(nullptr) { }
|
||||
ptr_A(nullptr), ptr_C(nullptr), ptr_D(nullptr),
|
||||
allow_early_exit(false) { }
|
||||
|
||||
/// constructs an arguments structure
|
||||
Arguments(
|
||||
@@ -148,7 +151,8 @@ public:
|
||||
int64_t batch_stride_D,
|
||||
typename LayoutA::Stride::Index lda,
|
||||
typename LayoutC::Stride::Index ldc,
|
||||
typename LayoutC::Stride::Index ldd
|
||||
typename LayoutC::Stride::Index ldd,
|
||||
bool allow_early_exit = false
|
||||
):
|
||||
mode(mode),
|
||||
problem_size(problem_size),
|
||||
@@ -156,7 +160,8 @@ public:
|
||||
epilogue(epilogue),
|
||||
ptr_A(ptr_A), ptr_C(ptr_C), ptr_D(ptr_D),
|
||||
batch_stride_A(batch_stride_A), batch_stride_C(batch_stride_C), batch_stride_D(batch_stride_D),
|
||||
lda(lda), ldb(ldb), ldc(ldc), ldd(ldd) {
|
||||
lda(lda), ldb(ldb), ldc(ldc), ldd(ldd),
|
||||
allow_early_exit(allow_early_exit) {
|
||||
|
||||
}
|
||||
|
||||
@@ -196,6 +201,8 @@ public:
|
||||
|
||||
int *semaphore;
|
||||
|
||||
bool allow_early_exit;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -218,7 +225,8 @@ public:
|
||||
batch_stride_B(0),
|
||||
batch_stride_C(0),
|
||||
batch_stride_D(0),
|
||||
semaphore(nullptr) { }
|
||||
semaphore(nullptr),
|
||||
allow_early_exit(false) { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(
|
||||
@@ -246,7 +254,8 @@ public:
|
||||
batch_stride_B(args.batch_stride_A),
|
||||
batch_stride_C(args.batch_stride_C),
|
||||
batch_stride_D(args.batch_stride_D),
|
||||
semaphore(static_cast<int *>(workspace)) {
|
||||
semaphore(static_cast<int *>(workspace)),
|
||||
allow_early_exit(args.allow_early_exit) {
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -313,6 +322,12 @@ public:
|
||||
cutlass::gemm::GemmCoord threadblock_tile_offset =
|
||||
threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
// Early exit following LAPACK's definition
|
||||
if (params.allow_early_exit &&
|
||||
(params.output_op.alpha == ElementC(0)) && (params.output_op.beta == ElementC(1))) {
|
||||
return;
|
||||
}
|
||||
|
||||
// Early exit if CTA is out of range
|
||||
if (params.grid_tiled_shape.m() <= threadblock_tile_offset.m() ||
|
||||
params.grid_tiled_shape.n() <= threadblock_tile_offset.n()) {
|
||||
|
||||
@@ -60,7 +60,7 @@ public:
|
||||
//
|
||||
using ProblemShape = ProblemShape_;
|
||||
|
||||
static_assert(cute::rank(ProblemShape{}) == 3 or cute::rank(ProblemShape{}) == 4,
|
||||
static_assert(rank(ProblemShape{}) == 3 or rank(ProblemShape{}) == 4,
|
||||
"ProblemShape{} should be <M,N,K> or <M,N,K,L>");
|
||||
|
||||
// Mainloop derived types
|
||||
@@ -101,7 +101,7 @@ public:
|
||||
sizeof(typename CollectiveMainloop::SharedStorage),
|
||||
sizeof(typename CollectiveEpilogue::SharedStorage)));
|
||||
|
||||
static constexpr uint32_t MaxThreadsPerBlock = cute::size(TiledMma{});
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(cute::size(TiledMma{}));
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
// Device side arguments
|
||||
@@ -141,8 +141,9 @@ public:
|
||||
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
return args.mode == GemmUniversalMode::kGemm or
|
||||
(args.mode == GemmUniversalMode::kBatched && cute::rank(ProblemShape{}) == 4);
|
||||
bool mode_implementable = args.mode == GemmUniversalMode::kGemm or
|
||||
(args.mode == GemmUniversalMode::kBatched && rank(ProblemShape{}) == 4);
|
||||
return mode_implementable && TileScheduler::can_implement(args.scheduler);
|
||||
}
|
||||
|
||||
static int
|
||||
|
||||
@@ -0,0 +1,749 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2023 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/workspace.h"
|
||||
#include "cutlass/fast_math.h"
|
||||
#include "cutlass/kernel_hardware_info.hpp"
|
||||
#include "cute/arch/cluster_sm90.hpp"
|
||||
#include "cutlass/arch/reg_reconfig.h"
|
||||
#include "cutlass/arch/mma_sm90.h"
|
||||
#include "cutlass/epilogue/collective/detail.hpp"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
#include "cutlass/gemm/group_array_problem_shape.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler.hpp"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/trace.h"
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::kernel {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class ProblemShape_,
|
||||
class CollectiveMainloop_,
|
||||
class CollectiveEpilogue_,
|
||||
class TileScheduler_
|
||||
>
|
||||
class GemmUniversal<
|
||||
ProblemShape_,
|
||||
CollectiveMainloop_,
|
||||
CollectiveEpilogue_,
|
||||
TileScheduler_,
|
||||
cute::enable_if_t<cute::is_base_of_v<KernelArrayTmaWarpSpecializedCooperative, typename CollectiveMainloop_::DispatchPolicy::Schedule> ||
|
||||
cute::is_base_of_v<KernelGroupTmaWarpSpecializedCooperative, typename CollectiveMainloop_::DispatchPolicy::Schedule>>
|
||||
>
|
||||
{
|
||||
public:
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using ProblemShape = ProblemShape_;
|
||||
static_assert(rank(typename ProblemShape::UnderlyingProblemShape{}) == 3 or rank(typename ProblemShape::UnderlyingProblemShape{}) == 4,
|
||||
"ProblemShape{} should be <M,N,K> or <M,N,K,L>");
|
||||
|
||||
// Mainloop derived types
|
||||
using CollectiveMainloop = CollectiveMainloop_;
|
||||
using TileShape = typename CollectiveMainloop::TileShape;
|
||||
using TiledMma = typename CollectiveMainloop::TiledMma;
|
||||
using ArchTag = typename CollectiveMainloop::ArchTag;
|
||||
using ElementA = typename CollectiveMainloop::ElementA;
|
||||
using StrideA = typename CollectiveMainloop::StrideA;
|
||||
using ElementB = typename CollectiveMainloop::ElementB;
|
||||
using StrideB = typename CollectiveMainloop::StrideB;
|
||||
using DispatchPolicy = typename CollectiveMainloop::DispatchPolicy;
|
||||
using Schedule = typename DispatchPolicy::Schedule;
|
||||
using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator;
|
||||
using ClusterShape = typename DispatchPolicy::ClusterShape;
|
||||
using MainloopArguments = typename CollectiveMainloop::Arguments;
|
||||
using MainloopParams = typename CollectiveMainloop::Params;
|
||||
|
||||
// Epilogue derived types
|
||||
using CollectiveEpilogue = CollectiveEpilogue_;
|
||||
using ElementC = typename CollectiveEpilogue::ElementC;
|
||||
using StrideC = typename CollectiveEpilogue::StrideC;
|
||||
using ElementD = typename CollectiveEpilogue::ElementD;
|
||||
using StrideD = typename CollectiveEpilogue::StrideD;
|
||||
using EpilogueArguments = typename CollectiveEpilogue::Arguments;
|
||||
using EpilogueParams = typename CollectiveEpilogue::Params;
|
||||
|
||||
static_assert(ArchTag::kMinComputeCapability >= 90);
|
||||
static_assert(cute::is_void_v<TileScheduler_>,
|
||||
"Ptr-Array Cooperative and Grouped Gemm Cooperative kernel only supports the default scheduler.");
|
||||
|
||||
static constexpr bool IsGroupedGemmKernel = cute::is_base_of_v<KernelGroupTmaWarpSpecializedCooperative, Schedule>;
|
||||
|
||||
using TileScheduler = cute::conditional_t<IsGroupedGemmKernel,
|
||||
typename detail::TileSchedulerSelector<
|
||||
GroupScheduler, ArchTag,
|
||||
TileShape, ClusterShape,
|
||||
ProblemShape>::Scheduler,
|
||||
typename detail::TileSchedulerSelector<
|
||||
void, ArchTag, TileShape, ClusterShape>::Scheduler>;
|
||||
using TileSchedulerArguments = typename TileScheduler::Arguments;
|
||||
using TileSchedulerParams = typename TileScheduler::Params;
|
||||
|
||||
static constexpr uint32_t NumLoadWarpGroups = 1;
|
||||
static constexpr uint32_t NumMmaWarpGroups = CUTE_STATIC_V(size(TiledMma{})) / NumThreadsPerWarpGroup;
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(size(TiledMma{})) + (NumLoadWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
/// Register requirement for Load and Math WGs
|
||||
static constexpr uint32_t LoadRegisterRequirement = 40;
|
||||
static constexpr uint32_t MmaRegisterRequirement = 232;
|
||||
|
||||
// 1 stage ordered sequence between mainloop and epilogue producer load threads
|
||||
using LoadWarpOrderBarrier = cutlass::OrderedSequenceBarrier<1,2>;
|
||||
|
||||
// Kernel level shared memory storage
|
||||
struct SharedStorage {
|
||||
struct TensorStorage : cute::aligned_struct<128> {
|
||||
using MainloopTensorStorage = typename CollectiveMainloop::TensorStorage;
|
||||
using EpilogueTensorStorage = typename CollectiveEpilogue::TensorStorage;
|
||||
|
||||
MainloopTensorStorage mainloop;
|
||||
EpilogueTensorStorage epilogue;
|
||||
} tensors;
|
||||
|
||||
struct PipelineStorage : cute::aligned_struct<16> {
|
||||
using MainloopPipelineStorage = typename CollectiveMainloop::PipelineStorage;
|
||||
using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage;
|
||||
|
||||
alignas(16) MainloopPipelineStorage mainloop;
|
||||
alignas(16) EpiLoadPipelineStorage epi_load;
|
||||
alignas(16) typename LoadWarpOrderBarrier::SharedStorage load_order;
|
||||
} pipelines;
|
||||
|
||||
struct TensorMapStorage : cute::aligned_struct<128> {
|
||||
using MainloopTensorMapStorage = typename CollectiveMainloop::TensorMapStorage;
|
||||
alignas(128) MainloopTensorMapStorage mainloop;
|
||||
} tensormaps;
|
||||
};
|
||||
|
||||
static constexpr int SharedStorageSize = sizeof(SharedStorage);
|
||||
|
||||
// Device side arguments
|
||||
struct Arguments {
|
||||
GemmUniversalMode mode{};
|
||||
ProblemShape problem_shape{};
|
||||
MainloopArguments mainloop{};
|
||||
EpilogueArguments epilogue{};
|
||||
KernelHardwareInfo hw_info{};
|
||||
TileSchedulerArguments scheduler{};
|
||||
};
|
||||
|
||||
// Kernel entry point API
|
||||
struct Params {
|
||||
GemmUniversalMode mode;
|
||||
ProblemShape problem_shape;
|
||||
MainloopParams mainloop;
|
||||
EpilogueParams epilogue;
|
||||
KernelHardwareInfo hw_info;
|
||||
TileSchedulerParams scheduler;
|
||||
void* workspace;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
// Convert to underlying arguments. In this case, a simple copy for the aliased type.
|
||||
static
|
||||
Params
|
||||
to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
CUTLASS_TRACE_HOST("to_underlying_arguments():");
|
||||
|
||||
ProblemShape problem_shapes = args.problem_shape;
|
||||
|
||||
// Get SM count if needed, otherwise use user supplied SM count
|
||||
int sm_count = args.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(args.hw_info.device_id);
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST("to_underlying_arguments(): Setting persistent grid SM count to " << sm_count);
|
||||
|
||||
KernelHardwareInfo hw_info{args.hw_info.device_id, sm_count};
|
||||
|
||||
// Calculate workspace pointers
|
||||
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
|
||||
size_t workspace_offset = 0;
|
||||
|
||||
void* scheduler_workspace = workspace_ptr;
|
||||
workspace_offset += TileScheduler::template get_workspace_size<typename ProblemShape::UnderlyingProblemShape, ElementAccumulator>(
|
||||
args.scheduler, problem_shapes.get_host_problem_shape(0), args.hw_info, NumMmaWarpGroups);
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
|
||||
void* epilogue_workspace = workspace_ptr + workspace_offset;
|
||||
workspace_offset += CollectiveEpilogue::get_workspace_size(problem_shapes, args.epilogue);
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
|
||||
void* mainloop_workspace = workspace_ptr + workspace_offset;
|
||||
workspace_offset += CollectiveMainloop::get_workspace_size(problem_shapes, args.mainloop, args.hw_info.sm_count);
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
|
||||
// Precompute the sub tiles numbers in epilogue, pass into tile scheduler. Therefore it will be used
|
||||
// in separate reduction scheme for streamk case, NumEpilogueSubTiles default value is 1, which means
|
||||
// subtile will not be used, therefore separate reduction will not be enabled.
|
||||
constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_store_pipe_increment(TileShape{});
|
||||
TileSchedulerParams scheduler;
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
scheduler = TileScheduler::to_underlying_arguments(
|
||||
problem_shapes, TileShape{}, ClusterShape{}, hw_info, args.scheduler, scheduler_workspace, NumEpilogueSubTiles);
|
||||
}
|
||||
else {
|
||||
scheduler = TileScheduler::to_underlying_arguments(
|
||||
problem_shapes.get_host_problem_shape(), TileShape{}, ClusterShape{}, hw_info, args.scheduler, scheduler_workspace, NumEpilogueSubTiles);
|
||||
}
|
||||
|
||||
return {
|
||||
args.mode,
|
||||
problem_shapes,
|
||||
CollectiveMainloop::to_underlying_arguments(problem_shapes, args.mainloop, mainloop_workspace),
|
||||
CollectiveEpilogue::to_underlying_arguments(problem_shapes, args.epilogue, epilogue_workspace),
|
||||
hw_info,
|
||||
scheduler,
|
||||
workspace
|
||||
};
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE static
|
||||
bool
|
||||
can_implement(Arguments const& args) {
|
||||
bool implementable = true;
|
||||
if constexpr (cute::is_base_of_v<KernelArrayTmaWarpSpecializedCooperative, Schedule>) {
|
||||
implementable &= (args.mode == GemmUniversalMode::kArray && rank(typename ProblemShape::UnderlyingProblemShape{}) == 4);
|
||||
} else if constexpr (IsGroupedGemmKernel) {
|
||||
// Group GEMM currently only supports rank-3 problem shapes
|
||||
implementable &= (args.mode == GemmUniversalMode::kGrouped && rank(typename ProblemShape::UnderlyingProblemShape{}) == 3);
|
||||
}
|
||||
else {
|
||||
implementable = false;
|
||||
}
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Arguments or Problem Shape don't meet the requirements for Ptr Array Gemm or Grouped Gemm.\n");
|
||||
return implementable;
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
return implementable;
|
||||
}
|
||||
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args) {
|
||||
size_t workspace_size = 0;
|
||||
constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_store_pipe_increment(TileShape{});
|
||||
|
||||
workspace_size += TileScheduler::template get_workspace_size<typename ProblemShape::UnderlyingProblemShape, ElementAccumulator>(
|
||||
args.scheduler, args.problem_shape.get_host_problem_shape(0), args.hw_info, NumMmaWarpGroups, NumEpilogueSubTiles);
|
||||
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
|
||||
|
||||
workspace_size += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
|
||||
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
|
||||
|
||||
// Get SM count if needed, otherwise use user supplied SM count
|
||||
int sm_count = args.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(args.hw_info.device_id);
|
||||
}
|
||||
|
||||
workspace_size += CollectiveMainloop::get_workspace_size(args.problem_shape, args.mainloop, sm_count);
|
||||
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
|
||||
|
||||
return workspace_size;
|
||||
}
|
||||
|
||||
static cutlass::Status
|
||||
initialize_workspace(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr) {
|
||||
Status status = Status::kSuccess;
|
||||
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
|
||||
size_t workspace_offset = 0;
|
||||
constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_store_pipe_increment(TileShape{});
|
||||
|
||||
status = TileScheduler::template initialize_workspace<typename ProblemShape::UnderlyingProblemShape, ElementAccumulator>(
|
||||
args.scheduler, workspace_ptr + workspace_offset, stream, args.problem_shape.get_host_problem_shape(0), args.hw_info, NumMmaWarpGroups, NumEpilogueSubTiles);
|
||||
workspace_offset += TileScheduler::template get_workspace_size<typename ProblemShape::UnderlyingProblemShape, ElementAccumulator>(
|
||||
args.scheduler, args.problem_shape.get_host_problem_shape(0), args.hw_info, NumMmaWarpGroups, NumEpilogueSubTiles);
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
status = CollectiveEpilogue::initialize_workspace(args.problem_shape, args.epilogue, workspace_ptr + workspace_offset, stream);
|
||||
workspace_offset += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
|
||||
status = CollectiveMainloop::initialize_workspace(args.problem_shape, args.mainloop, workspace_ptr + workspace_offset, stream);
|
||||
workspace_offset += CollectiveMainloop::get_workspace_size(args.problem_shape, args.mainloop, args.hw_info.sm_count);
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
}
|
||||
|
||||
return status;
|
||||
}
|
||||
|
||||
// Computes the kernel launch grid shape based on runtime parameters
|
||||
static dim3
|
||||
get_grid_shape(Params const& params) {
|
||||
// Given device SM count, set grid size s.t. we do not launch more thread blocks than we can run concurrently
|
||||
TileSchedulerArguments args{};
|
||||
if constexpr (!std::is_const_v<decltype(args.max_swizzle_size)>) {
|
||||
args.max_swizzle_size = 1 << params.scheduler.log_swizzle_size_;
|
||||
}
|
||||
args.raster_order = params.scheduler.raster_order_ == TileScheduler::RasterOrder::AlongN ? TileScheduler::RasterOrderOptions::AlongN : TileScheduler::RasterOrderOptions::AlongM;
|
||||
dim3 grid_shape;
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
grid_shape = TileScheduler::get_grid_shape(params.problem_shape, TileShape{}, ClusterShape{}, params.hw_info, args);
|
||||
}
|
||||
else {
|
||||
grid_shape = TileScheduler::get_grid_shape(params.problem_shape.get_host_problem_shape(), TileShape{}, ClusterShape{}, params.hw_info, args);
|
||||
}
|
||||
return grid_shape;
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_block_shape() {
|
||||
return dim3(MaxThreadsPerBlock, 1, 1);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
operator()(Params const& params, char* smem_buf) {
|
||||
using namespace cute;
|
||||
using X = Underscore;
|
||||
|
||||
// Any Tensor Op MMA Atom in the WGMMA ISA is arch conditional to sm90a.
|
||||
#if ! defined(__CUDA_ARCH_FEAT_SM90_ALL)
|
||||
if constexpr(size<0>(typename TiledMma::AtomShape_MNK{}) == 64) {
|
||||
printf("ERROR : Arch conditional MMA instruction used without targeting sm90a compute capability. Aborting.\n");
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
|
||||
// Preconditions
|
||||
static_assert(size(TiledMma{}) == 256, "Cooperative kernel must have TiledMMA operating using 256 threads.");
|
||||
static_assert(size<0>(TileShape{}) >= 128,
|
||||
"Cooperative kernel requires Tile Size to be greater than or equal to 128 along the M-dimension.");
|
||||
|
||||
static_assert(cute::rank(StrideA{}) == 3, "StrideA must be rank-3: [M, K, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
static_assert(cute::rank(StrideB{}) == 3, "StrideB must be rank-3: [N, K, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
static_assert(cute::rank(StrideC{}) == 3, "StrideC must be rank-3: [M, N, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
static_assert(cute::rank(StrideD{}) == 3, "StrideD must be rank-3: [M, N, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
|
||||
/* In the Cooperative kernel, Consumer0 and Consumer1 collaborate on the same tile */
|
||||
enum class WarpGroupRole {
|
||||
Producer = 0,
|
||||
Consumer0 = 1,
|
||||
Consumer1 = 2
|
||||
};
|
||||
enum class ProducerWarpRole {
|
||||
Mainloop = 0,
|
||||
Warp1 = 1,
|
||||
Epilogue = 2,
|
||||
Warp3 = 3
|
||||
};
|
||||
|
||||
// Kernel level shared memory storage
|
||||
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
int lane_idx = canonical_lane_idx();
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % NumWarpsPerWarpGroup;
|
||||
int warp_group_thread_idx = thread_idx % NumThreadsPerWarpGroup;
|
||||
int mma_thread_idx = thread_idx % size(TiledMma{});
|
||||
auto warp_group_role = WarpGroupRole(canonical_warp_group_idx());
|
||||
auto producer_warp_role = ProducerWarpRole(warp_idx_in_warp_group);
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
uint32_t block_rank_in_cluster = cute::block_rank_in_cluster();
|
||||
|
||||
// Note: Tma Descriptor Prefetch (from either const or param) is not applicable here
|
||||
|
||||
// Mainloop Load pipeline
|
||||
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
|
||||
typename MainloopPipeline::Params mainloop_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::Mainloop) {
|
||||
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer0 || warp_group_role == WarpGroupRole::Consumer1) {
|
||||
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
mainloop_pipeline_params.num_consumers = size(TiledMma{});
|
||||
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params, ClusterShape{});
|
||||
|
||||
// Epilogue Load pipeline
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
typename EpiLoadPipeline::Params epi_load_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer && producer_warp_role == ProducerWarpRole::Epilogue) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer0 || warp_group_role == WarpGroupRole::Consumer1) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
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 = size(TiledMma{});
|
||||
epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes;
|
||||
EpiLoadPipeline epi_load_pipeline(shared_storage.pipelines.epi_load, epi_load_pipeline_params);
|
||||
|
||||
// Epilogue Store pipeline
|
||||
using EpiStorePipeline = typename CollectiveEpilogue::StorePipeline;
|
||||
typename EpiStorePipeline::Params epi_store_pipeline_params;
|
||||
epi_store_pipeline_params.always_wait = true;
|
||||
EpiStorePipeline epi_store_pipeline(epi_store_pipeline_params);
|
||||
|
||||
typename LoadWarpOrderBarrier::Params params_load_order_barrier;
|
||||
params_load_order_barrier.group_id = producer_warp_role == ProducerWarpRole::Mainloop ? 0 : 1;
|
||||
params_load_order_barrier.group_size = NumThreadsPerWarp;
|
||||
LoadWarpOrderBarrier load_order_barrier(shared_storage.pipelines.load_order, params_load_order_barrier);
|
||||
|
||||
// Initialize starting pipeline states for the collectives
|
||||
// Epilogue store pipe is producer-only (consumer is TMA unit, waits via scoreboarding)
|
||||
typename CollectiveMainloop::PipelineState mainloop_pipe_consumer_state;
|
||||
typename CollectiveEpilogue::LoadPipelineState epi_load_pipe_consumer_state;
|
||||
// Purpose of maintaining this pipeline state is to make sure TMA loads have finished before doing descriptor updates
|
||||
typename CollectiveMainloop::PipelineState mainloop_pipe_tma_consumer_state;
|
||||
|
||||
// For the DMA Load (producer) we start with an opposite phase
|
||||
// i.e., we skip all waits since we know that the buffer is indeed empty
|
||||
PipelineState mainloop_pipe_producer_state = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
PipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state<EpiLoadPipeline>();
|
||||
PipelineState epi_store_pipe_producer_state = cutlass::make_producer_start_state<EpiStorePipeline>();
|
||||
|
||||
auto cluster_wait_fn = [] () {
|
||||
// We need this to guarantee that the Pipeline init is visible
|
||||
// To all producers and consumer thread blocks in the Cluster
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
cute::cluster_arrive_relaxed();
|
||||
return [] () { cute::cluster_wait(); };
|
||||
}
|
||||
else {
|
||||
__syncthreads();
|
||||
return [] () {}; // do nothing
|
||||
}
|
||||
} ();
|
||||
|
||||
// Get the appropriate blocks for this thread block -- potential for thread block locality
|
||||
TiledMma tiled_mma;
|
||||
auto blk_shape = TileShape{}; // (BLK_M,BLK_N,BLK_K)
|
||||
|
||||
TileScheduler scheduler{params.scheduler};
|
||||
auto work_tile_info = scheduler.get_current_work();
|
||||
|
||||
// Optionally append 1s until problem shape is rank-4 in case it is only rank-3 (MNK)
|
||||
auto problem_shape_MNKL = append<4>(params.problem_shape.get_problem_shape(work_tile_info.L_idx), Int<1>{});
|
||||
|
||||
// In a warp specialized kernel, collectives expose data movement and compute operations separately
|
||||
CollectiveMainloop collective_mainloop;
|
||||
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
|
||||
|
||||
// Prepare and partition the input tensors. Expects a tuple of tensors where:
|
||||
// get<0>(load_inputs) is the tma tensor A after local tiling so that it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// get<1>(load_inputs) is the tma tensor B after local tiling so that it has shape (BLK_N,BLK_K,n,k,l)
|
||||
auto load_inputs = collective_mainloop.load_init(problem_shape_MNKL, params.mainloop);
|
||||
static_assert(cute::tuple_size_v<decltype(load_inputs)> >= 2, "Output of load_init must have at least two elements (A, B)");
|
||||
|
||||
// Extract out partitioned A and B.
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
// Get pipeline stage increments from tensor shapes
|
||||
auto k_tile_count = size<3>(gA_mkl);
|
||||
|
||||
// Wait for all thread blocks in the Cluster
|
||||
cluster_wait_fn();
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
cutlass::arch::warpgroup_reg_dealloc<LoadRegisterRequirement>();
|
||||
|
||||
// Mainloop Producer Warp
|
||||
if (producer_warp_role == ProducerWarpRole::Mainloop) {
|
||||
int32_t curr_batch = idx2crd(work_tile_info.L_idx, shape<4>(gB_nkl)); // Usually just returns work_tile_info.L_idx;
|
||||
int32_t next_batch = curr_batch;
|
||||
int32_t const mock_l_coord = 0;
|
||||
int32_t const sm_idx = blockIdx.x + (blockIdx.y * gridDim.x);
|
||||
int32_t const sm_count = params.hw_info.sm_count;
|
||||
|
||||
// Fetch a copy of tensormaps for the CTA
|
||||
auto input_tensormaps = collective_mainloop.tensormaps_init(params.mainloop, sm_count, sm_idx);
|
||||
|
||||
// Update tensormap for the initial batch for the CTA
|
||||
if (work_tile_info.is_valid()) {
|
||||
collective_mainloop.tensormaps_perform_update(
|
||||
shared_storage.tensormaps.mainloop,
|
||||
params.mainloop,
|
||||
input_tensormaps,
|
||||
problem_shape_MNKL,
|
||||
next_batch
|
||||
);
|
||||
// Ensure warp is converged before issuing tensor replace
|
||||
__syncwarp();
|
||||
// Entire warp must do this (ie its aligned)
|
||||
collective_mainloop.tensormaps_cp_fence_release(shared_storage.tensormaps.mainloop, input_tensormaps);
|
||||
}
|
||||
|
||||
bool do_load_order_arrive = true;
|
||||
while (work_tile_info.is_valid()) {
|
||||
if (!TileScheduler::valid_warpgroup_in_work_tile(work_tile_info)) {
|
||||
work_tile_info = fetch_next_work(work_tile_info, scheduler);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Compute m_coord, n_coord, l_coord with the post-tiled m-shape and n-shape
|
||||
auto m_coord = idx2crd(work_tile_info.M_idx, shape<2>(gA_mkl));
|
||||
auto n_coord = idx2crd(work_tile_info.N_idx, shape<2>(gB_nkl));
|
||||
auto blk_coord = make_coord(m_coord, n_coord, _, mock_l_coord);
|
||||
|
||||
// Get the number of K tiles to compute for this work as well as the starting K tile offset of the work.
|
||||
auto work_k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, blk_shape);
|
||||
auto work_k_tile_start = TileScheduler::get_work_k_tile_start(work_tile_info);
|
||||
auto k_tile_iter = cute::make_coord_iterator(idx2crd(work_k_tile_start, shape<3>(gA_mkl)), shape<3>(gA_mkl));
|
||||
|
||||
collective_mainloop.tensormaps_fence_acquire(input_tensormaps);
|
||||
|
||||
collective_mainloop.load(
|
||||
params.mainloop,
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
load_inputs,
|
||||
input_tensormaps,
|
||||
blk_coord,
|
||||
k_tile_iter, work_k_tile_count,
|
||||
lane_idx,
|
||||
block_rank_in_cluster,
|
||||
shared_storage.tensors.mainloop
|
||||
);
|
||||
// Update starting pipeline state for the next tile
|
||||
mainloop_pipe_producer_state.advance(work_k_tile_count);
|
||||
|
||||
// Signal for the epilogue load warp to begin
|
||||
if (do_load_order_arrive) {
|
||||
load_order_barrier.arrive();
|
||||
do_load_order_arrive = false;
|
||||
}
|
||||
|
||||
// Get next work tile
|
||||
work_tile_info = fetch_next_work(work_tile_info, scheduler);
|
||||
next_batch = idx2crd(work_tile_info.L_idx, shape<4>(gB_nkl)); // Usually just returns work_tile_info.L_idx
|
||||
|
||||
if (work_tile_info.is_valid() && next_batch != curr_batch ) {
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
problem_shape_MNKL = append<4>(params.problem_shape.get_problem_shape(next_batch), Int<1>{});
|
||||
}
|
||||
// Wait for the last TMA stage to complete loading, before issuing tensormap updates
|
||||
mainloop_pipe_tma_consumer_state.advance(work_k_tile_count-1);
|
||||
mainloop_pipeline.consumer_wait(mainloop_pipe_tma_consumer_state);
|
||||
collective_mainloop.tensormaps_perform_update(
|
||||
shared_storage.tensormaps.mainloop,
|
||||
params.mainloop,
|
||||
input_tensormaps,
|
||||
problem_shape_MNKL,
|
||||
next_batch
|
||||
);
|
||||
// Ensure warp is converged before issuing tensor replace
|
||||
__syncwarp();
|
||||
// Entire warp must do this (ie its aligned)
|
||||
collective_mainloop.tensormaps_cp_fence_release(shared_storage.tensormaps.mainloop, input_tensormaps);
|
||||
curr_batch = next_batch;
|
||||
// Advance the TMA consumer state for the last remaining stage that was being waited for above
|
||||
mainloop_pipe_tma_consumer_state.advance(1);
|
||||
}
|
||||
else if (work_tile_info.is_valid()) { // case where batch/group didn't change between tiles
|
||||
// Advance the TMA consumer state for all the stages to be in sync
|
||||
mainloop_pipe_tma_consumer_state.advance(work_k_tile_count);
|
||||
}
|
||||
} // Scheduler work fetch loop
|
||||
|
||||
// Make sure all Consumer Warp Groups have been waited upon
|
||||
collective_mainloop.load_tail(mainloop_pipeline, mainloop_pipe_producer_state);
|
||||
} // Mainloop Producer Warp End
|
||||
|
||||
// Epilogue Producer Warp
|
||||
else if (producer_warp_role == ProducerWarpRole::Epilogue && collective_epilogue.is_producer_load_needed()) {
|
||||
while (work_tile_info.is_valid()) {
|
||||
if (!TileScheduler::requires_separate_reduction(params.scheduler)) {
|
||||
load_order_barrier.wait();
|
||||
}
|
||||
if (TileScheduler::compute_epilogue(work_tile_info, params.scheduler)) {
|
||||
// Compute m_coord, n_coord, l_coord with the post-tiled m-shape and n-shape
|
||||
auto m_coord = idx2crd(work_tile_info.M_idx, shape<2>(gA_mkl));
|
||||
auto n_coord = idx2crd(work_tile_info.N_idx, shape<2>(gB_nkl));
|
||||
auto l_coord = idx2crd(work_tile_info.L_idx, shape<4>(gB_nkl));
|
||||
auto blk_coord = make_coord(m_coord, n_coord, _, l_coord);
|
||||
|
||||
epi_load_pipe_producer_state =
|
||||
collective_epilogue.load(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_producer_state,
|
||||
problem_shape_MNKL,
|
||||
blk_shape,
|
||||
blk_coord,
|
||||
tiled_mma,
|
||||
lane_idx,
|
||||
shared_storage.tensors.epilogue,
|
||||
work_tile_info.reduction_subtile_idx()
|
||||
);
|
||||
}
|
||||
|
||||
// Get next work tile
|
||||
work_tile_info = fetch_next_work(work_tile_info, scheduler);
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
problem_shape_MNKL = append<4>(params.problem_shape.get_problem_shape(work_tile_info.L_idx), Int<1>{});
|
||||
}
|
||||
} // Scheduler work fetch loop
|
||||
|
||||
// Make sure all Consumer Warp Groups have been waited upon
|
||||
collective_epilogue.load_tail(epi_load_pipeline, epi_load_pipe_producer_state);
|
||||
} // Epilogue Producer Warp End
|
||||
} // Producer Warp Group End
|
||||
|
||||
else if (warp_group_role == WarpGroupRole::Consumer0 || warp_group_role == WarpGroupRole::Consumer1) {
|
||||
cutlass::arch::warpgroup_reg_alloc<MmaRegisterRequirement>();
|
||||
|
||||
// Do we potentially issue tail arrives for TMA stores, if epilogue load is waiting for it
|
||||
bool do_store_tail = false;
|
||||
while (work_tile_info.is_valid()) {
|
||||
// Compute m_coord, n_coord, l_coord with the post-tiled m-shape and n-shape
|
||||
auto m_coord = idx2crd(work_tile_info.M_idx, shape<2>(gA_mkl));
|
||||
auto n_coord = idx2crd(work_tile_info.N_idx, shape<2>(gB_nkl));
|
||||
auto l_coord = idx2crd(work_tile_info.L_idx, shape<4>(gB_nkl));
|
||||
auto blk_coord = make_coord(m_coord, n_coord, _, l_coord);
|
||||
auto work_k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, blk_shape);
|
||||
|
||||
// Allocate the accumulators for the (M,N) blk_shape
|
||||
//
|
||||
// MSVC CTAD breaks if we say "Tensor" here, so we use "auto" instead.
|
||||
auto accumulators = partition_fragment_C(tiled_mma, take<0,2>(blk_shape)); // (MMA,MMA_M,MMA_N)
|
||||
if(TileScheduler::valid_warpgroup_in_work_tile(work_tile_info)) {
|
||||
collective_mainloop.mma(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
accumulators,
|
||||
work_k_tile_count,
|
||||
mma_thread_idx,
|
||||
shared_storage.tensors.mainloop,
|
||||
params.mainloop
|
||||
);
|
||||
|
||||
// Make sure the math instructions are done and free buffers before entering the epilogue
|
||||
collective_mainloop.mma_tail(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
work_k_tile_count
|
||||
);
|
||||
|
||||
// Update starting mainloop pipeline state for the next tile
|
||||
mainloop_pipe_consumer_state.advance(work_k_tile_count);
|
||||
}
|
||||
// Index of warp group within consumer warp groups
|
||||
int consumer_warp_group_idx = canonical_warp_group_idx() - NumLoadWarpGroups;
|
||||
|
||||
// Perform reduction across splits, if needed
|
||||
TileScheduler::fixup(
|
||||
params.scheduler, work_tile_info, accumulators, NumMmaWarpGroups, consumer_warp_group_idx);
|
||||
|
||||
if (TileScheduler::compute_epilogue(work_tile_info, params.scheduler)) {
|
||||
// Epilogue and write to gD
|
||||
auto [epi_load_pipe_consumer_state_next, epi_store_pipe_producer_state_next] =
|
||||
collective_epilogue.store(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_consumer_state,
|
||||
epi_store_pipeline,
|
||||
epi_store_pipe_producer_state,
|
||||
problem_shape_MNKL,
|
||||
blk_shape,
|
||||
blk_coord,
|
||||
accumulators,
|
||||
tiled_mma,
|
||||
mma_thread_idx,
|
||||
shared_storage.tensors.epilogue,
|
||||
work_tile_info.reduction_subtile_idx()
|
||||
);
|
||||
epi_load_pipe_consumer_state = epi_load_pipe_consumer_state_next;
|
||||
epi_store_pipe_producer_state = epi_store_pipe_producer_state_next;
|
||||
do_store_tail = true;
|
||||
}
|
||||
|
||||
// Get next work tile
|
||||
work_tile_info = fetch_next_work(work_tile_info, scheduler);
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
problem_shape_MNKL = append<4>(params.problem_shape.get_problem_shape(work_tile_info.L_idx), Int<1>{});
|
||||
}
|
||||
} // Scheduler work fetch loop
|
||||
|
||||
if (do_store_tail) {
|
||||
collective_epilogue.store_tail(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_consumer_state,
|
||||
epi_store_pipeline,
|
||||
epi_store_pipe_producer_state
|
||||
);
|
||||
}
|
||||
} // Consumer Warp Groups End
|
||||
}
|
||||
|
||||
private:
|
||||
// Kernel helper function to get next work unit
|
||||
CUTLASS_DEVICE
|
||||
typename TileScheduler::WorkTileInfo
|
||||
fetch_next_work(
|
||||
typename TileScheduler::WorkTileInfo& work_tile_info,
|
||||
TileScheduler& scheduler) const {
|
||||
// Check whether we should continue on with the current work unit. If this is the case,
|
||||
// the work unit will have been updated in continue_current_work to reflect the new
|
||||
// tile to be computed.
|
||||
if (scheduler.continue_current_work(work_tile_info)) {
|
||||
return work_tile_info;
|
||||
}
|
||||
|
||||
// Get next work tile
|
||||
scheduler.advance_to_next_work();
|
||||
return scheduler.get_current_work();
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::kernel
|
||||
@@ -121,7 +121,7 @@ public:
|
||||
sizeof(typename CollectiveMainloop::SharedStorage),
|
||||
sizeof(typename CollectiveEpilogue::SharedStorage)));
|
||||
|
||||
static constexpr uint32_t MaxThreadsPerBlock = size(TiledMma{});
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(size(TiledMma{}));
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
// Device side arguments
|
||||
@@ -176,6 +176,8 @@ public:
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
|
||||
return implementable;
|
||||
}
|
||||
|
||||
|
||||
@@ -128,7 +128,7 @@ public:
|
||||
|
||||
static constexpr uint32_t NumLoadWarpGroups = 1;
|
||||
static constexpr uint32_t NumMmaWarpGroups = 1;
|
||||
static constexpr uint32_t MaxThreadsPerBlock = size(TiledMma{}) + (NumLoadWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(size(TiledMma{})) + (NumLoadWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
// Device side arguments
|
||||
@@ -183,6 +183,8 @@ public:
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
|
||||
return implementable;
|
||||
}
|
||||
|
||||
@@ -270,7 +272,7 @@ public:
|
||||
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
mainloop_pipeline_params.num_consumers = NumThreadsPerWarpGroup;
|
||||
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params);
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params, ClusterShape{});
|
||||
|
||||
// Epilogue Load pipeline
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
@@ -335,14 +337,14 @@ public:
|
||||
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
|
||||
|
||||
// Prepare and partition the input tensors. Expects a tuple of tensors where:
|
||||
// get<0>(tiled_tensors) is the tma tensor A after local tiling so that it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// get<1>(tiled_tensors) is the tma tensor B after local tiling so that it has shape (BLK_N,BLK_K,n,k,l)
|
||||
auto tiled_tensors = collective_mainloop.tile_input_tensors(problem_shape_MNKL, params.mainloop, blk_shape);
|
||||
static_assert(cute::tuple_size_v<decltype(tiled_tensors)> >= 2, "Output of tile_input_tensors must have at least two elements (A, B)");
|
||||
// get<0>(load_inputs) is the tma tensor A after local tiling so that it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// get<1>(load_inputs) is the tma tensor B after local tiling so that it has shape (BLK_N,BLK_K,n,k,l)
|
||||
auto load_inputs = collective_mainloop.load_init(problem_shape_MNKL, params.mainloop);
|
||||
static_assert(cute::tuple_size_v<decltype(load_inputs)> >= 2, "Output of load_init must have at least two elements (A, B)");
|
||||
|
||||
// Extract out partitioned A and B.
|
||||
Tensor gA_mkl = get<0>(tiled_tensors);
|
||||
Tensor gB_nkl = get<1>(tiled_tensors);
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
// Compute m_coord, n_coord, and l_coord with their post-tiled shapes
|
||||
auto m_coord = idx2crd(int(blockIdx.x), shape<2>(gA_mkl));
|
||||
@@ -363,7 +365,7 @@ public:
|
||||
params.mainloop,
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
tiled_tensors,
|
||||
load_inputs,
|
||||
blk_coord,
|
||||
k_tile_iter, k_tile_count,
|
||||
lane_idx,
|
||||
@@ -378,8 +380,7 @@ public:
|
||||
if (collective_epilogue.is_producer_load_needed()) {
|
||||
// Ensure warp is converged before issuing epilogue loads
|
||||
__syncwarp();
|
||||
epi_load_pipe_producer_state =
|
||||
collective_epilogue.load(
|
||||
epi_load_pipe_producer_state = collective_epilogue.load(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_producer_state,
|
||||
problem_shape_MNKL,
|
||||
|
||||
@@ -105,8 +105,8 @@ public:
|
||||
using TileSchedulerParams = typename TileScheduler::Params;
|
||||
|
||||
static constexpr uint32_t NumLoadWarpGroups = 1;
|
||||
static constexpr uint32_t NumMmaWarpGroups = size(TiledMma{}) / NumThreadsPerWarpGroup;
|
||||
static constexpr uint32_t MaxThreadsPerBlock = size(TiledMma{}) + (NumLoadWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t NumMmaWarpGroups = CUTE_STATIC_V(size(TiledMma{})) / NumThreadsPerWarpGroup;
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(size(TiledMma{})) + (NumLoadWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
/// Register requirement for Load and Math WGs
|
||||
@@ -203,6 +203,12 @@ public:
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
|
||||
void* mainloop_workspace = nullptr;
|
||||
// Precompute the sub tiles numbers in epilogue, pass into tile scheduler. Therefore it will be used
|
||||
// in separate reduction scheme for streamk case, NumEpilogueSubTiles default value is 1, which means
|
||||
// subtile will not be used, therefore separate reduction will not be enabled.
|
||||
constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_store_pipe_increment(TileShape{});
|
||||
TileSchedulerParams scheduler = TileScheduler::to_underlying_arguments(
|
||||
problem_shape_MNKL, TileShape{}, ClusterShape{}, hw_info, args.scheduler, scheduler_workspace, NumEpilogueSubTiles);
|
||||
|
||||
return {
|
||||
args.mode,
|
||||
@@ -210,7 +216,7 @@ public:
|
||||
CollectiveMainloop::to_underlying_arguments(args.problem_shape, args.mainloop, mainloop_workspace),
|
||||
CollectiveEpilogue::to_underlying_arguments(args.problem_shape, args.epilogue, epilogue_workspace),
|
||||
hw_info,
|
||||
TileScheduler::to_underlying_arguments(problem_shape_MNKL, TileShape{}, ClusterShape{}, hw_info, args.scheduler, scheduler_workspace),
|
||||
scheduler,
|
||||
workspace
|
||||
};
|
||||
}
|
||||
@@ -226,14 +232,17 @@ public:
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
return implementable;
|
||||
}
|
||||
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args) {
|
||||
size_t workspace_size = 0;
|
||||
constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_store_pipe_increment(TileShape{});
|
||||
|
||||
workspace_size += TileScheduler::template get_workspace_size<ProblemShape, ElementAccumulator>(
|
||||
args.scheduler, args.problem_shape, args.hw_info, NumMmaWarpGroups);
|
||||
args.scheduler, args.problem_shape, args.hw_info, NumMmaWarpGroups, NumEpilogueSubTiles);
|
||||
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
|
||||
|
||||
workspace_size += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
|
||||
@@ -247,11 +256,12 @@ public:
|
||||
Status status = Status::kSuccess;
|
||||
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
|
||||
size_t workspace_offset = 0;
|
||||
constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_store_pipe_increment(TileShape{});
|
||||
|
||||
status = TileScheduler::template initialize_workspace<ProblemShape, ElementAccumulator>(
|
||||
args.scheduler, workspace_ptr + workspace_offset, stream, args.problem_shape, args.hw_info, NumMmaWarpGroups);
|
||||
args.scheduler, workspace_ptr + workspace_offset, stream, args.problem_shape, args.hw_info, NumMmaWarpGroups, NumEpilogueSubTiles);
|
||||
workspace_offset += TileScheduler::template get_workspace_size<ProblemShape, ElementAccumulator>(
|
||||
args.scheduler, args.problem_shape, args.hw_info, NumMmaWarpGroups);
|
||||
args.scheduler, args.problem_shape, args.hw_info, NumMmaWarpGroups, NumEpilogueSubTiles);
|
||||
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
|
||||
if (status != Status::kSuccess) {
|
||||
return status;
|
||||
@@ -353,7 +363,7 @@ public:
|
||||
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
mainloop_pipeline_params.num_consumers = size(TiledMma{});
|
||||
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params);
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params, ClusterShape{});
|
||||
|
||||
// Epilogue Load pipeline
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
@@ -392,7 +402,7 @@ public:
|
||||
PipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state<EpiLoadPipeline>();
|
||||
PipelineState epi_store_pipe_producer_state = cutlass::make_producer_start_state<EpiStorePipeline>();
|
||||
|
||||
auto cluster_wait_fn = [&] () {
|
||||
auto cluster_wait_fn = [] () {
|
||||
// We need this to guarantee that the Pipeline init is visible
|
||||
// To all producers and consumer thread blocks in the Cluster
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
@@ -420,14 +430,14 @@ public:
|
||||
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
|
||||
|
||||
// Prepare and partition the input tensors. Expects a tuple of tensors where:
|
||||
// get<0>(tiled_tensors) is the tma tensor A after local tiling so that it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// get<1>(tiled_tensors) is the tma tensor B after local tiling so that it has shape (BLK_N,BLK_K,n,k,l)
|
||||
auto tiled_tensors = collective_mainloop.tile_input_tensors(problem_shape_MNKL, params.mainloop, blk_shape);
|
||||
static_assert(cute::tuple_size_v<decltype(tiled_tensors)> >= 2, "Output of tile_input_tensors must have at least two elements (A, B)");
|
||||
// get<0>(load_inputs) is the tma tensor A after local tiling so that it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// get<1>(load_inputs) is the tma tensor B after local tiling so that it has shape (BLK_N,BLK_K,n,k,l)
|
||||
auto load_inputs = collective_mainloop.load_init(problem_shape_MNKL, params.mainloop);
|
||||
static_assert(cute::tuple_size_v<decltype(load_inputs)> >= 2, "Output of load_init must have at least two elements (A, B)");
|
||||
|
||||
// Extract out partitioned A and B.
|
||||
Tensor gA_mkl = get<0>(tiled_tensors);
|
||||
Tensor gB_nkl = get<1>(tiled_tensors);
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
// Get pipeline stage increments from tensor shapes
|
||||
auto k_tile_count = size<3>(gA_mkl);
|
||||
@@ -442,6 +452,11 @@ public:
|
||||
if (producer_warp_role == ProducerWarpRole::Mainloop) {
|
||||
bool do_load_order_arrive = true;
|
||||
while (work_tile_info.is_valid()) {
|
||||
if (!TileScheduler::valid_warpgroup_in_work_tile(work_tile_info)) {
|
||||
work_tile_info = fetch_next_work(work_tile_info, scheduler);
|
||||
continue;
|
||||
}
|
||||
|
||||
// Compute m_coord, n_coord, l_coord with the post-tiled m-shape and n-shape
|
||||
auto m_coord = idx2crd(work_tile_info.M_idx, shape<2>(gA_mkl));
|
||||
auto n_coord = idx2crd(work_tile_info.N_idx, shape<2>(gB_nkl));
|
||||
@@ -457,7 +472,7 @@ public:
|
||||
params.mainloop,
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
tiled_tensors,
|
||||
load_inputs,
|
||||
blk_coord,
|
||||
k_tile_iter, work_k_tile_count,
|
||||
lane_idx,
|
||||
@@ -483,8 +498,10 @@ public:
|
||||
|
||||
// Epilogue Producer Warp
|
||||
else if (producer_warp_role == ProducerWarpRole::Epilogue && collective_epilogue.is_producer_load_needed()) {
|
||||
load_order_barrier.wait();
|
||||
while (work_tile_info.is_valid()) {
|
||||
if (!TileScheduler::requires_separate_reduction(params.scheduler)) {
|
||||
load_order_barrier.wait();
|
||||
}
|
||||
if (TileScheduler::compute_epilogue(work_tile_info, params.scheduler)) {
|
||||
// Compute m_coord, n_coord, l_coord with the post-tiled m-shape and n-shape
|
||||
auto m_coord = idx2crd(work_tile_info.M_idx, shape<2>(gA_mkl));
|
||||
@@ -501,7 +518,8 @@ public:
|
||||
blk_coord,
|
||||
tiled_mma,
|
||||
lane_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
shared_storage.tensors.epilogue,
|
||||
work_tile_info.reduction_subtile_idx()
|
||||
);
|
||||
}
|
||||
|
||||
@@ -531,27 +549,27 @@ public:
|
||||
//
|
||||
// MSVC CTAD breaks if we say "Tensor" here, so we use "auto" instead.
|
||||
auto accumulators = partition_fragment_C(tiled_mma, take<0,2>(blk_shape)); // (MMA,MMA_M,MMA_N)
|
||||
if(TileScheduler::valid_warpgroup_in_work_tile(work_tile_info)) {
|
||||
collective_mainloop.mma(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
accumulators,
|
||||
work_k_tile_count,
|
||||
mma_thread_idx,
|
||||
shared_storage.tensors.mainloop,
|
||||
params.mainloop
|
||||
);
|
||||
|
||||
collective_mainloop.mma(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
accumulators,
|
||||
work_k_tile_count,
|
||||
mma_thread_idx,
|
||||
shared_storage.tensors.mainloop,
|
||||
params.mainloop
|
||||
);
|
||||
|
||||
// Make sure the math instructions are done and free buffers before entering the epilogue
|
||||
collective_mainloop.mma_tail(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
work_k_tile_count
|
||||
);
|
||||
|
||||
// Update starting mainloop pipeline state for the next tile
|
||||
mainloop_pipe_consumer_state.advance(work_k_tile_count);
|
||||
// Make sure the math instructions are done and free buffers before entering the epilogue
|
||||
collective_mainloop.mma_tail(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
work_k_tile_count
|
||||
);
|
||||
|
||||
// Update starting mainloop pipeline state for the next tile
|
||||
mainloop_pipe_consumer_state.advance(work_k_tile_count);
|
||||
}
|
||||
// Index of warp group within consumer warp groups
|
||||
int consumer_warp_group_idx = canonical_warp_group_idx() - NumLoadWarpGroups;
|
||||
|
||||
@@ -573,7 +591,8 @@ public:
|
||||
accumulators,
|
||||
tiled_mma,
|
||||
mma_thread_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
shared_storage.tensors.epilogue,
|
||||
work_tile_info.reduction_subtile_idx()
|
||||
);
|
||||
epi_load_pipe_consumer_state = epi_load_pipe_consumer_state_next;
|
||||
epi_store_pipe_producer_state = epi_store_pipe_producer_state_next;
|
||||
|
||||
@@ -107,7 +107,7 @@ public:
|
||||
|
||||
static constexpr uint32_t NumLoadWarpGroups = 1;
|
||||
static constexpr uint32_t NumMmaWarpGroups = 2;
|
||||
static constexpr uint32_t MaxThreadsPerBlock = size(TiledMma{}) + (NumMmaWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(size(TiledMma{})) + (NumMmaWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
/// Register requirement for Load and Math WGs
|
||||
@@ -232,6 +232,8 @@ public:
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
|
||||
return implementable;
|
||||
}
|
||||
|
||||
@@ -324,7 +326,7 @@ public:
|
||||
|
||||
// Kernel level shared memory storage
|
||||
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
|
||||
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
int lane_idx = canonical_lane_idx();
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
@@ -353,7 +355,7 @@ public:
|
||||
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
mainloop_pipeline_params.num_consumers = NumThreadsPerWarpGroup;
|
||||
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params);
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params, ClusterShape{});
|
||||
|
||||
// Epilogue Load pipeline
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
@@ -424,14 +426,14 @@ public:
|
||||
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
|
||||
|
||||
// Prepare and partition the input tensors. Expects a tuple of tensors where:
|
||||
// get<0>(tiled_tensors) is the tma tensor A after local tiling so that it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// get<1>(tiled_tensors) is the tma tensor B after local tiling so that it has shape (BLK_N,BLK_K,n,k,l)
|
||||
auto tiled_tensors = collective_mainloop.tile_input_tensors(problem_shape_MNKL, params.mainloop, blk_shape);
|
||||
static_assert(cute::tuple_size_v<decltype(tiled_tensors)> >= 2, "Output of tile_input_tensors must have at least two elements (A, B)");
|
||||
// get<0>(load_inputs) is the tma tensor A after local tiling so that it has shape (BLK_M,BLK_K,m,k,l)
|
||||
// get<1>(load_inputs) is the tma tensor B after local tiling so that it has shape (BLK_N,BLK_K,n,k,l)
|
||||
auto load_inputs = collective_mainloop.load_init(problem_shape_MNKL, params.mainloop);
|
||||
static_assert(cute::tuple_size_v<decltype(load_inputs)> >= 2, "Output of load_init must have at least two elements (A, B)");
|
||||
|
||||
// Extract out partitioned A and B.
|
||||
Tensor gA_mkl = get<0>(tiled_tensors);
|
||||
Tensor gB_nkl = get<1>(tiled_tensors);
|
||||
Tensor gA_mkl = get<0>(load_inputs);
|
||||
Tensor gB_nkl = get<1>(load_inputs);
|
||||
|
||||
// Get pipeline stage increments from tensor shapes
|
||||
auto k_tile_count = size<3>(gA_mkl);
|
||||
@@ -472,7 +474,7 @@ public:
|
||||
params.mainloop,
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
tiled_tensors,
|
||||
load_inputs,
|
||||
blk_coord,
|
||||
k_tile_iter, k_tile_count,
|
||||
lane_idx,
|
||||
|
||||
@@ -187,6 +187,8 @@ public:
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
|
||||
return implementable;
|
||||
}
|
||||
|
||||
|
||||
@@ -207,6 +207,8 @@ public:
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
|
||||
return implementable;
|
||||
}
|
||||
|
||||
|
||||
@@ -219,6 +219,8 @@ public:
|
||||
}
|
||||
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
|
||||
return implementable;
|
||||
}
|
||||
|
||||
|
||||
@@ -76,6 +76,12 @@ public:
|
||||
is_final_split(uint32_t k_tiles_per_output_tile) const {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t
|
||||
reduction_subtile_idx() const {
|
||||
return -1;
|
||||
}
|
||||
};
|
||||
|
||||
using Params = PersistentTileSchedulerSm90Params;
|
||||
@@ -101,7 +107,8 @@ public:
|
||||
ClusterShape cluster_shape,
|
||||
[[maybe_unused]] KernelHardwareInfo const& hw_info,
|
||||
Arguments const& arguments,
|
||||
[[maybe_unused]] void* workspace=nullptr) {
|
||||
[[maybe_unused]] void* workspace=nullptr,
|
||||
[[maybe_unused]] const uint32_t epilogue_subtile = 1) {
|
||||
|
||||
// We only need the tile and cluster shape during scheduler setup, so let FTAD do the magic
|
||||
static_assert(cute::is_static<TileShape>::value);
|
||||
@@ -114,13 +121,19 @@ public:
|
||||
problem_blocks,
|
||||
to_gemm_coord(cluster_shape),
|
||||
hw_info,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.raster_order
|
||||
);
|
||||
|
||||
return params;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
PersistentTileSchedulerSm90() { };
|
||||
|
||||
@@ -164,7 +177,7 @@ public:
|
||||
scheduler_params.divmod_cluster_shape_major_,
|
||||
scheduler_params.divmod_cluster_shape_minor_,
|
||||
scheduler_params.divmod_cluster_blk_major_,
|
||||
scheduler_params.log_swizzle_size_,
|
||||
scheduler_params.log_swizzle_size_,
|
||||
scheduler_params.raster_order_);
|
||||
|
||||
return {work_idx_m, work_idx_n, static_cast<int32_t>(work_idx_l), true};
|
||||
@@ -180,11 +193,11 @@ public:
|
||||
static CUTLASS_DEVICE
|
||||
cute::tuple<int32_t, int32_t>
|
||||
get_work_idx_m_and_n(
|
||||
uint64_t blk_per_grid_dim,
|
||||
uint64_t blk_per_grid_dim,
|
||||
FastDivmodU64Pow2 const& divmod_cluster_shape_major,
|
||||
FastDivmodU64Pow2 const& divmod_cluster_shape_minor,
|
||||
FastDivmodU64 const& divmod_cluster_blk_major,
|
||||
int32_t log_swizzle_size,
|
||||
int32_t log_swizzle_size,
|
||||
RasterOrder raster_order) {
|
||||
|
||||
uint64_t cluster_id, cluster_major_offset = 0, cluster_minor_offset = 0;
|
||||
@@ -199,26 +212,26 @@ public:
|
||||
}
|
||||
|
||||
uint64_t cluster_idx_minor, cluster_idx_major;
|
||||
|
||||
|
||||
uint64_t cluster_idx_minor_div_swizzle, extra, offset;
|
||||
|
||||
offset = cluster_id & ((1 << log_swizzle_size) - 1);
|
||||
extra = cluster_id >> log_swizzle_size;
|
||||
|
||||
|
||||
divmod_cluster_blk_major(cluster_idx_minor_div_swizzle, cluster_idx_major, extra);
|
||||
|
||||
cluster_idx_minor = cluster_idx_minor_div_swizzle * (1 << log_swizzle_size) + offset;
|
||||
|
||||
auto minor_work_idx = static_cast<int32_t>(cluster_idx_minor * divmod_cluster_shape_minor.divisor +
|
||||
auto minor_work_idx = static_cast<int32_t>(cluster_idx_minor * divmod_cluster_shape_minor.divisor +
|
||||
cluster_minor_offset);
|
||||
auto major_work_idx = static_cast<int32_t>(cluster_idx_major * divmod_cluster_shape_major.divisor +
|
||||
auto major_work_idx = static_cast<int32_t>(cluster_idx_major * divmod_cluster_shape_major.divisor +
|
||||
cluster_major_offset);
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
return {minor_work_idx, major_work_idx};
|
||||
}
|
||||
else {
|
||||
return {major_work_idx, minor_work_idx};
|
||||
return {major_work_idx, minor_work_idx};
|
||||
}
|
||||
|
||||
}
|
||||
@@ -331,13 +344,14 @@ public:
|
||||
// The basic tile scheduler does not require any additional workspace
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static int
|
||||
get_workspace_size(Arguments const&, ProblemShape, KernelHardwareInfo const&, uint32_t) {
|
||||
get_workspace_size(Arguments const&, ProblemShape, KernelHardwareInfo const&, uint32_t, const uint32_t = 1) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static cutlass::Status
|
||||
initialize_workspace(Arguments const&, void*, cudaStream_t, ProblemShape, KernelHardwareInfo const&, uint32_t) {
|
||||
initialize_workspace(Arguments const&, void*, cudaStream_t, ProblemShape, KernelHardwareInfo const&,
|
||||
uint32_t, const uint32_t = 1) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
@@ -353,8 +367,61 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint32_t
|
||||
get_work_k_tile_start(WorkTileInfo const&) {
|
||||
// All work units returned by this scheduler start from K tile 0
|
||||
return 0u;
|
||||
// All work units returned by this scheduler start from K tile 0
|
||||
return 0u;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
need_separate_reduction(Params const& params) {
|
||||
return false;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
is_work_tile_for_reduction(WorkTileInfo const& work_tile_info, Params const& params) {
|
||||
return false;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
uint32_t
|
||||
epilgoue_subtile_idx(WorkTileInfo const& work_tile_info, Params const& params) const {
|
||||
return 0;
|
||||
}
|
||||
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
separate_reduction(
|
||||
Params const& params,
|
||||
WorkTileInfo const& work_tile_info,
|
||||
FrgTensorC& accumulators,
|
||||
uint32_t num_barriers,
|
||||
uint32_t barrier_idx) {
|
||||
}
|
||||
|
||||
// Shares the accumulator set with peers in the global workspace
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
static void
|
||||
share(
|
||||
Params const& params,
|
||||
WorkTileInfo const& work_tile_info,
|
||||
FrgTensorC& accumulators,
|
||||
uint32_t num_barriers,
|
||||
uint32_t barrier_idx) {
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
valid_warpgroup_in_work_tile(WorkTileInfo const& work_tile_info) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
requires_separate_reduction(Params const& params) {
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -0,0 +1,431 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2023 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/fast_math.h"
|
||||
#include "cutlass/gemm_coord.hpp"
|
||||
#include "cutlass/kernel_hardware_info.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler_params.h"
|
||||
#include "cute/layout.hpp"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cute/arch/cluster_sm90.hpp"
|
||||
|
||||
namespace cutlass::gemm::kernel::detail {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Persistent Thread Block (TB) scheduler
|
||||
template <class GroupProblemShape>
|
||||
class PersistentTileSchedulerSm90Group {
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
private:
|
||||
uint64_t current_work_linear_idx_ = 0;
|
||||
uint64_t total_grid_size_ = 0;
|
||||
|
||||
// Tracking current group, its starting linear idx and total tiles
|
||||
struct GroupInfo {
|
||||
uint64_t group = 0;
|
||||
uint64_t start_linear_idx = 0;
|
||||
uint64_t total_tiles = 0;
|
||||
} current_group_info_;
|
||||
|
||||
public:
|
||||
struct WorkTileInfo {
|
||||
int32_t M_idx = 0;
|
||||
int32_t N_idx = 0;
|
||||
int32_t L_idx = 0;
|
||||
bool is_valid_tile = false;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_valid() const {
|
||||
return is_valid_tile;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static WorkTileInfo
|
||||
invalid_work_tile() {
|
||||
return {-1, -1, -1, false};
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_final_split(uint32_t k_tiles_per_output_tile) const {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t
|
||||
reduction_subtile_idx() const {
|
||||
return -1;
|
||||
}
|
||||
};
|
||||
|
||||
using ProblemShape = typename GroupProblemShape::UnderlyingProblemShape;
|
||||
using Params = PersistentTileSchedulerSm90GroupParams<ProblemShape>;
|
||||
using RasterOrder = typename Params::RasterOrder;
|
||||
using RasterOrderOptions = typename Params::RasterOrderOptions;
|
||||
struct Arguments {
|
||||
int max_swizzle_size = 1;
|
||||
// Not applying Heuristics for Grouped problems, since largest dimension can change per group
|
||||
RasterOrderOptions raster_order = RasterOrderOptions::AlongM;
|
||||
};
|
||||
|
||||
// Sink scheduler params as a member
|
||||
Params scheduler_params;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
template <class TileShape, class ClusterShape>
|
||||
static Params
|
||||
to_underlying_arguments(
|
||||
GroupProblemShape problem_shapes,
|
||||
TileShape tile_shape,
|
||||
ClusterShape cluster_shape,
|
||||
[[maybe_unused]] KernelHardwareInfo const& hw_info,
|
||||
Arguments const& arguments,
|
||||
[[maybe_unused]] void* workspace=nullptr,
|
||||
[[maybe_unused]] const uint32_t epilogue_subtile = 1) {
|
||||
|
||||
// We only need the tile and cluster shape during scheduler setup, so let FTAD do the magic
|
||||
static_assert(cute::is_static<TileShape>::value);
|
||||
static_assert(cute::is_static<ClusterShape>::value);
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(
|
||||
problem_shapes.groups(),
|
||||
reinterpret_cast<ProblemShape const*>(problem_shapes.host_problem_shapes),
|
||||
tile_shape, cluster_shape);
|
||||
|
||||
Params params;
|
||||
params.initialize(
|
||||
problem_blocks,
|
||||
problem_shapes.groups(),
|
||||
reinterpret_cast<ProblemShape*>(problem_shapes.problem_shapes),
|
||||
to_gemm_coord(tile_shape),
|
||||
to_gemm_coord(cluster_shape),
|
||||
hw_info,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.raster_order
|
||||
);
|
||||
|
||||
return params;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
return true;
|
||||
}
|
||||
|
||||
PersistentTileSchedulerSm90Group() = default;
|
||||
|
||||
CUTLASS_DEVICE explicit PersistentTileSchedulerSm90Group(Params const& params_) : scheduler_params(params_) {
|
||||
// MSVC requires protecting use of CUDA-specific nonstandard syntax,
|
||||
// like blockIdx and gridDim, with __CUDA_ARCH__.
|
||||
#if defined(__CUDA_ARCH__)
|
||||
if (params_.raster_order_ == RasterOrder::AlongN) {
|
||||
current_work_linear_idx_ = uint64_t(blockIdx.x) + uint64_t(blockIdx.y) * uint64_t(gridDim.x);
|
||||
}
|
||||
else {
|
||||
current_work_linear_idx_ = uint64_t(blockIdx.x) * uint64_t(gridDim.y) + uint64_t(blockIdx.y);
|
||||
}
|
||||
|
||||
total_grid_size_ = uint64_t(gridDim.x) * uint64_t(gridDim.y) * uint64_t(gridDim.z);
|
||||
|
||||
auto cta_m = cute::size(cute::ceil_div(cute::shape<0>(params_.problem_shapes_[0]), params_.cta_shape_.m()));
|
||||
auto cta_n = cute::size(cute::ceil_div(cute::shape<1>(params_.problem_shapes_[0]), params_.cta_shape_.n()));
|
||||
current_group_info_.total_tiles = cta_m * cta_n;
|
||||
#else
|
||||
CUTLASS_ASSERT(false && "This line should never be reached");
|
||||
#endif
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work() {
|
||||
return get_current_work_for_linear_idx(current_work_linear_idx_);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work_for_linear_idx(uint64_t linear_idx) {
|
||||
if (linear_idx >= scheduler_params.blocks_per_problem_) {
|
||||
return WorkTileInfo::invalid_work_tile();
|
||||
}
|
||||
|
||||
uint64_t blk_per_grid_dim = scheduler_params.divmod_cluster_shape_minor_.divide(linear_idx);
|
||||
|
||||
auto [work_idx_m, work_idx_n, new_group_info, valid_tile] = get_work_idx_m_and_n(blk_per_grid_dim,
|
||||
current_group_info_,
|
||||
scheduler_params.groups_,
|
||||
scheduler_params.problem_shapes_,
|
||||
scheduler_params.cta_shape_,
|
||||
scheduler_params.divmod_cluster_shape_major_,
|
||||
scheduler_params.divmod_cluster_shape_minor_,
|
||||
scheduler_params.log_swizzle_size_,
|
||||
scheduler_params.raster_order_);
|
||||
|
||||
current_group_info_ = new_group_info;
|
||||
return {work_idx_m, work_idx_n, static_cast<int>(current_group_info_.group), valid_tile};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
advance_to_next_work(uint32_t advance_count = 1) {
|
||||
current_work_linear_idx_ += total_grid_size_ * uint64_t(advance_count);
|
||||
}
|
||||
|
||||
// get work_idx_m, work_idx_n from blk_per_grid_dim while applying swizzle
|
||||
static CUTLASS_DEVICE
|
||||
cute::tuple<int32_t, int32_t, struct GroupInfo, bool>
|
||||
get_work_idx_m_and_n(
|
||||
uint64_t blk_per_grid_dim,
|
||||
struct GroupInfo group_info,
|
||||
int32_t total_problem_groups,
|
||||
ProblemShape* problem_shapes,
|
||||
GemmCoord cta_shape,
|
||||
FastDivmodU64Pow2 const& divmod_cluster_shape_major,
|
||||
FastDivmodU64Pow2 const& divmod_cluster_shape_minor,
|
||||
int32_t log_swizzle_size,
|
||||
RasterOrder raster_order) {
|
||||
|
||||
bool valid_tile = true;
|
||||
int cta_m = cute::size(cute::ceil_div(cute::shape<0>(problem_shapes[group_info.group]), cta_shape.m()));
|
||||
int cta_n = cute::size(cute::ceil_div(cute::shape<1>(problem_shapes[group_info.group]), cta_shape.n()));
|
||||
|
||||
while (group_info.start_linear_idx + group_info.total_tiles <= blk_per_grid_dim) {
|
||||
group_info.group++;
|
||||
group_info.start_linear_idx += group_info.total_tiles;
|
||||
cta_m = cute::size(cute::ceil_div(cute::shape<0>(problem_shapes[group_info.group]), cta_shape.m()));
|
||||
cta_n = cute::size(cute::ceil_div(cute::shape<1>(problem_shapes[group_info.group]), cta_shape.n()));
|
||||
group_info.total_tiles = cta_m * cta_n;
|
||||
}
|
||||
|
||||
uint64_t cluster_id, cluster_major_offset = 0, cluster_minor_offset = 0;
|
||||
divmod_cluster_shape_major(cluster_id, cluster_major_offset, blk_per_grid_dim - group_info.start_linear_idx);
|
||||
|
||||
auto [cta_m_in_cluster, cta_n_in_cluster, _] = cute::block_id_in_cluster();
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
cluster_minor_offset = cta_m_in_cluster;
|
||||
}
|
||||
else {
|
||||
cluster_minor_offset = cta_n_in_cluster;
|
||||
}
|
||||
|
||||
uint64_t cluster_idx_minor, cluster_idx_major;
|
||||
|
||||
uint64_t cluster_idx_minor_div_swizzle, extra, offset;
|
||||
|
||||
offset = cluster_id & ((1 << log_swizzle_size) - 1);
|
||||
extra = cluster_id >> log_swizzle_size;
|
||||
|
||||
uint64_t curr_group_cluster_blk_major, remainder;
|
||||
divmod_cluster_shape_major(curr_group_cluster_blk_major, remainder, cta_m);
|
||||
cluster_idx_minor_div_swizzle = extra / curr_group_cluster_blk_major;
|
||||
cluster_idx_major = extra % curr_group_cluster_blk_major;
|
||||
|
||||
cluster_idx_minor = cluster_idx_minor_div_swizzle * (1 << log_swizzle_size) + offset;
|
||||
|
||||
auto minor_work_idx = static_cast<int32_t>(cluster_idx_minor * divmod_cluster_shape_minor.divisor +
|
||||
cluster_minor_offset);
|
||||
auto major_work_idx = static_cast<int32_t>(cluster_idx_major * divmod_cluster_shape_major.divisor +
|
||||
cluster_major_offset);
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
return {minor_work_idx, major_work_idx, group_info, valid_tile};
|
||||
}
|
||||
else {
|
||||
return {major_work_idx, minor_work_idx, group_info, valid_tile};
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
// Given the inputs, computes the total number of output blocks this problem will compute over
|
||||
// Note that this is only the logical size of our grid, not the physical grid we will actually launch.
|
||||
template<class BlockShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE static
|
||||
dim3
|
||||
get_tiled_cta_shape_mnl(int groups, ProblemShape const* problem_shapes, BlockShape cta_shape, ClusterShape cluster_shape) {
|
||||
uint32_t total_ctas = 0;
|
||||
uint32_t cta_in_N_dim = 1; // We linearize the blocks across all the problems here
|
||||
for (int group = 0; group < groups; group++) {
|
||||
auto cta_m = cute::size(cute::ceil_div(cute::shape<0>(problem_shapes[group]), cute::shape<0>(cta_shape)));
|
||||
auto cta_n = cute::size(cute::ceil_div(cute::shape<1>(problem_shapes[group]), cute::shape<1>(cta_shape)));
|
||||
total_ctas += cta_m * cta_n;
|
||||
}
|
||||
|
||||
return Params::get_tiled_cta_shape_mnl(
|
||||
to_gemm_coord(cluster_shape),
|
||||
total_ctas, cta_in_N_dim
|
||||
);
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
template<class BlockShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE static
|
||||
dim3
|
||||
get_grid_shape(
|
||||
GroupProblemShape problem_shapes,
|
||||
BlockShape cta_shape,
|
||||
ClusterShape cluster_shape,
|
||||
KernelHardwareInfo hw_info,
|
||||
Arguments arguments,
|
||||
bool truncate_by_problem_size=true) {
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(
|
||||
problem_shapes.groups(),
|
||||
reinterpret_cast<ProblemShape const*>(problem_shapes.host_problem_shapes),
|
||||
cta_shape, cluster_shape);
|
||||
|
||||
return Params::get_grid_shape(
|
||||
problem_blocks,
|
||||
to_gemm_coord(cluster_shape),
|
||||
hw_info,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.raster_order,
|
||||
/* truncate_by_problem_size = */true
|
||||
);
|
||||
}
|
||||
|
||||
// Returns whether the block assigned this work should compute the epilogue for the corresponding
|
||||
// output tile. For the basic tile scheduler, this is always true.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
compute_epilogue(WorkTileInfo const&, Params const&) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Performs the reduction across splits for a given output tile. Since this scheduler does
|
||||
// not split output tiles, no reduction is needed.
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
static void
|
||||
fixup(Params const&, WorkTileInfo const&, FrgTensorC&, uint32_t, uint32_t) {}
|
||||
|
||||
// Returns whether the current WorkTileInfo passed in should continue to be used. Since
|
||||
// this scheduler only schedules work in units of single, full output tiles, the WorkTileInfo
|
||||
// passed in should not be used after having been processed.
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
continue_current_work(WorkTileInfo&) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// The basic tile scheduler does not require any additional workspace
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static int
|
||||
get_workspace_size(Arguments const&, ProblemShape, KernelHardwareInfo const&, uint32_t, const uint32_t = 1) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static cutlass::Status
|
||||
initialize_workspace(Arguments const&, void*, cudaStream_t, ProblemShape, KernelHardwareInfo const&,
|
||||
uint32_t, const uint32_t = 1) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
template <class ProblemShape_MNKL, class TileShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int
|
||||
get_work_k_tile_count(WorkTileInfo const& work_tile_info, ProblemShape_MNKL problem_shape, TileShape tile_shape) {
|
||||
// All work units returned by this scheduler cover the entire K iteration
|
||||
// space of the output tile assigned to the work unit.
|
||||
return cute::size(cute::ceil_div(cute::get<2>(problem_shape), cute::get<2>(tile_shape)));
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint32_t
|
||||
get_work_k_tile_start(WorkTileInfo const&) {
|
||||
// All work units returned by this scheduler start from K tile 0
|
||||
return 0u;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
need_separate_reduction(Params const& params) {
|
||||
return false;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
is_work_tile_for_reduction(WorkTileInfo const& work_tile_info, Params const& params) {
|
||||
return false;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
uint32_t
|
||||
epilgoue_subtile_idx(WorkTileInfo const& work_tile_info, Params const& params) const {
|
||||
return 0;
|
||||
}
|
||||
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
separate_reduction(
|
||||
Params const& params,
|
||||
WorkTileInfo const& work_tile_info,
|
||||
FrgTensorC& accumulators,
|
||||
uint32_t num_barriers,
|
||||
uint32_t barrier_idx) {
|
||||
}
|
||||
|
||||
// Shares the accumulator set with peers in the global workspace
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
static void
|
||||
share(
|
||||
Params const& params,
|
||||
WorkTileInfo const& work_tile_info,
|
||||
FrgTensorC& accumulators,
|
||||
uint32_t num_barriers,
|
||||
uint32_t barrier_idx) {
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
valid_warpgroup_in_work_tile(WorkTileInfo const& work_tile_info) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
requires_separate_reduction(Params const& params) {
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cutlass::gemm::kernel::detail
|
||||
@@ -69,6 +69,7 @@ public:
|
||||
|
||||
using Params = PersistentTileSchedulerSm90StreamKParams;
|
||||
using ReductionMode = Params::ReductionMode;
|
||||
using DecompositionMode = Params::DecompositionMode;
|
||||
|
||||
struct WorkTileInfo {
|
||||
int32_t M_idx = 0;
|
||||
@@ -83,11 +84,43 @@ public:
|
||||
// Number of k tiles remaining for the work unit as a whole
|
||||
uint32_t k_tile_remaining = 0;
|
||||
|
||||
// Whether this unit of work is the final split for the given tile
|
||||
bool is_separate_reduction = false;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_valid() const {
|
||||
// Use negative indices to denote invalid work
|
||||
return M_idx >= 0;
|
||||
// A work tile that computes no K tiles is invalid unless it is a separate-reduction work tile
|
||||
// (which only performs reduction and epilogue)
|
||||
return k_tile_count > 0 || is_separate_reduction;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_reduction_unit() const {
|
||||
return is_separate_reduction;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
int32_t
|
||||
reduction_subtile_idx() const {
|
||||
// For separate reduction units, the K_idx of the work tile is unused.
|
||||
// Therefore, we override it to contain the subtile of that the reduction
|
||||
// unit operates on.
|
||||
return is_reduction_unit() ? K_idx : -1;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
void
|
||||
setup_separate_reduction(int32_t epilogue_subtile_idx) {
|
||||
// Set the epilogue subtile in the K_idx, since this is otherwise unused
|
||||
// by separate reduction units.
|
||||
K_idx = epilogue_subtile_idx;
|
||||
|
||||
is_separate_reduction = true;
|
||||
k_tile_count = 0;
|
||||
// Clean up remaining k tiles
|
||||
k_tile_remaining = 0;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -113,7 +146,10 @@ public:
|
||||
Arguments&
|
||||
operator=(Arguments const& args) {
|
||||
splits = args.splits;
|
||||
max_swizzle_size = args.max_swizzle_size;
|
||||
raster_order = args.raster_order;
|
||||
reduction_mode = args.reduction_mode;
|
||||
decomposition_mode = args.decomposition_mode;
|
||||
return *this;
|
||||
}
|
||||
|
||||
@@ -121,7 +157,10 @@ public:
|
||||
Arguments&
|
||||
operator=(Arguments&& args) noexcept {
|
||||
splits = args.splits;
|
||||
max_swizzle_size = args.max_swizzle_size;
|
||||
raster_order = args.raster_order;
|
||||
reduction_mode = args.reduction_mode;
|
||||
decomposition_mode = args.decomposition_mode;
|
||||
return *this;
|
||||
}
|
||||
|
||||
@@ -129,18 +168,20 @@ public:
|
||||
Arguments(int splits_) : splits(splits_) {}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(int splits_, int max_swizzle_size_, RasterOrderOptions raster_order_) :
|
||||
Arguments(int splits_, int max_swizzle_size_, RasterOrderOptions raster_order_, DecompositionMode decomposition_mode_) :
|
||||
splits(splits_),
|
||||
max_swizzle_size(max_swizzle_size_),
|
||||
raster_order(raster_order_) {}
|
||||
raster_order(raster_order_),
|
||||
decomposition_mode(decomposition_mode_) {}
|
||||
|
||||
// The splitting factor to be used in a split-K decomposition of the problem.
|
||||
// If this is set to a value greater than 1, stream-K decomposition logic
|
||||
// is bypassed in favor of a split-K decomposition.
|
||||
int splits = 1;
|
||||
const int max_swizzle_size = 1;
|
||||
int max_swizzle_size = 1;
|
||||
RasterOrderOptions raster_order = RasterOrderOptions::Heuristic;
|
||||
ReductionMode reduction_mode = ReductionMode::Deterministic;
|
||||
DecompositionMode decomposition_mode = DecompositionMode::Heuristic;
|
||||
};
|
||||
|
||||
// Sink scheduler params as a member
|
||||
@@ -158,7 +199,8 @@ public:
|
||||
ClusterShape cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
Arguments const& args,
|
||||
void* workspace) {
|
||||
void* workspace,
|
||||
const uint32_t epilogue_subtile = 1) {
|
||||
|
||||
static_assert(cute::is_static<TileShape>::value);
|
||||
static_assert(cute::is_static<ClusterShape>::value);
|
||||
@@ -177,11 +219,22 @@ public:
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
args.reduction_mode,
|
||||
workspace
|
||||
args.decomposition_mode,
|
||||
workspace,
|
||||
epilogue_subtile
|
||||
);
|
||||
return params;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
// Split count > 1 is only valid for heuristic and split-K decomposition modes
|
||||
return (args.splits == 1 ||
|
||||
args.decomposition_mode == DecompositionMode::Heuristic ||
|
||||
args.decomposition_mode == DecompositionMode::SplitK);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
PersistentTileSchedulerSm90StreamK() { };
|
||||
|
||||
@@ -210,7 +263,7 @@ public:
|
||||
// for the fact that we have splits_ peers per output tile, we multiply this
|
||||
// value by splits_. For stream-K, this multiplication ends up being a no-op
|
||||
// because splits_ is set to 1 for stream-K.
|
||||
if (linear_idx >= params.units_per_problem_ * params.splits_) {
|
||||
if(linear_idx >= (params.units_per_problem_ * params.splits_ + params.separate_reduction_units_)) {
|
||||
// Invalid work. Return an empty result.
|
||||
return WorkTileInfo::invalid_work_tile();
|
||||
}
|
||||
@@ -231,8 +284,8 @@ public:
|
||||
current_work_linear_idx_, work_tile_info, scheduler_params);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE static
|
||||
bool
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
continue_current_work_for_linear_idx(
|
||||
uint64_t linear_idx,
|
||||
WorkTileInfo& work_tile_info,
|
||||
@@ -243,9 +296,8 @@ public:
|
||||
if (work_tile_info.k_tile_remaining == 0) {
|
||||
return false;
|
||||
}
|
||||
|
||||
assign_work(params, linear_idx, work_tile_info);
|
||||
return true;
|
||||
return work_tile_info.is_valid();
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
@@ -281,7 +333,7 @@ public:
|
||||
problem_blocks,
|
||||
to_gemm_coord(cluster_shape),
|
||||
hw_info,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.max_swizzle_size,
|
||||
arguments.raster_order
|
||||
);
|
||||
}
|
||||
@@ -290,8 +342,22 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
requires_fixup(Params const& params, WorkTileInfo const& work_tile_info) {
|
||||
// Fixup is not needed for data-parallel tiles
|
||||
return work_tile_info.k_tile_count != params.divmod_tiles_per_output_tile_.divisor;
|
||||
// Fixup is not needed for invalid or data-parallel tiles
|
||||
return work_tile_info.is_valid() && work_tile_info.k_tile_count != params.divmod_tiles_per_output_tile_.divisor;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
requires_separate_reduction(Params const& params) {
|
||||
return params.requires_separate_reduction();
|
||||
}
|
||||
|
||||
// When the work tile is not special for reduction, it's valid. Otherwise need to skip
|
||||
// global loading that producer warpgroup do, also math computation that consumer warpgroup do.
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
valid_warpgroup_in_work_tile(WorkTileInfo const& work_tile_info) {
|
||||
return !work_tile_info.is_reduction_unit();
|
||||
}
|
||||
|
||||
// Performs the reduction across splits for a given output tile.
|
||||
@@ -304,7 +370,7 @@ public:
|
||||
FrgTensorC& accumulators,
|
||||
uint32_t num_barriers,
|
||||
uint32_t barrier_idx) {
|
||||
static constexpr uint32_t Offset = 2;
|
||||
static constexpr uint32_t Offset = static_cast<int>(cutlass::arch::ReservedNamedBarriers::StreamkBarrier0);
|
||||
static constexpr uint32_t MaxNumNamedBarriers = 2;
|
||||
using BarrierManager = NamedBarrierManager<NumThreadsPerWarpGroup, Offset, MaxNumNamedBarriers>;
|
||||
return fixup_helper<FrgTensorC, BarrierManager>(
|
||||
@@ -327,16 +393,27 @@ public:
|
||||
if (!requires_fixup(params, work_tile_info)) {
|
||||
return;
|
||||
}
|
||||
|
||||
auto tile_idx = output_tile_index(params, work_tile_info);
|
||||
|
||||
// Index of the lock on which to wait
|
||||
auto lock_idx = (tile_idx * num_barriers) + barrier_idx;
|
||||
|
||||
auto reduction_tile_idx = tile_idx;
|
||||
auto [first_peer_id, my_peer_id, last_peer_id] = tile_peer_range(params, tile_idx, static_cast<uint32_t>(work_tile_info.K_idx));
|
||||
auto reduction_peer_offset = 0;
|
||||
if (params.requires_separate_reduction()) {
|
||||
// If separate reduction is to be performed, each stream-K unit writes its partials
|
||||
// to a separate portion of the workspace. There are as many of these portions as there
|
||||
// are peers for a given output tile, so we multiply the tile index by the maximum peer count.
|
||||
reduction_tile_idx *= Params::max_peers_per_tile(params.sk_units_, params.sk_tiles_);
|
||||
reduction_peer_offset = my_peer_id * cute::size<0>(TileShape{}) * cute::size<1>(TileShape{});
|
||||
}
|
||||
|
||||
// Reductions use BlockStripedReduce with a width of BarrierManager::ThreadCount under the hood.
|
||||
// Thus, the start of the reduction space is the same across all threads in a warp group.
|
||||
int reduction_offset =
|
||||
(cute::size<0>(TileShape{}) * cute::size<1>(TileShape{}) * tile_idx) +
|
||||
(cute::size<0>(TileShape{}) * cute::size<1>(TileShape{}) * reduction_tile_idx) +
|
||||
reduction_peer_offset +
|
||||
(size(accumulators) * barrier_idx * BarrierManager::ThreadCount);
|
||||
|
||||
ElementAccumulator* group_reduction_workspace = reinterpret_cast<ElementAccumulator*>(params.reduction_workspace_) + reduction_offset;
|
||||
@@ -344,56 +421,86 @@ public:
|
||||
using AccumulatorArrayT = Array<typename FrgTensorC::value_type, size(FrgTensorC{})>;
|
||||
using BlockStripedReduceT = BlockStripedReduce<BarrierManager::ThreadCount, AccumulatorArrayT>;
|
||||
|
||||
AccumulatorArrayT* reduction_workspace_array = reinterpret_cast<AccumulatorArrayT*>(group_reduction_workspace);
|
||||
AccumulatorArrayT* accumulator_array = reinterpret_cast<AccumulatorArrayT*>(&accumulators);
|
||||
|
||||
int barrier_group_thread_idx = threadIdx.x % BarrierManager::ThreadCount;
|
||||
|
||||
// The number of tiles for which reduction is required is either:
|
||||
// (a) the total number of output tiles (in the case of split-K)
|
||||
// (b) the number of stream-K tiles
|
||||
// (b) the number of stream-K tiles (potentially multiplied by peer count if using separate reduction)
|
||||
// To calculate the total number of output tiles in the split-K case, we
|
||||
// note that, in the split-K case, the units_per_problem_ member of Params will be
|
||||
// the total number of output tiles.
|
||||
auto reduction_tiles = params.splits_ > 1 ? params.units_per_problem_ : params.sk_tiles_;
|
||||
uint32_t reduction_tiles = 0;
|
||||
if (params.splits_ > 1) {
|
||||
reduction_tiles = params.units_per_problem_;
|
||||
}
|
||||
else if (params.requires_separate_reduction()) {
|
||||
reduction_tiles = params.sk_tiles_ * Params::max_peers_per_tile(params.sk_units_, params.sk_tiles_);
|
||||
}
|
||||
else {
|
||||
reduction_tiles = params.sk_tiles_;
|
||||
}
|
||||
|
||||
auto reduction_workspace_size = Params::get_reduction_workspace_size(
|
||||
reduction_tiles, to_gemm_coord(TileShape{}), sizeof_bits<ElementAccumulator>::value);
|
||||
BarrierType* lock_workspace = reinterpret_cast<BarrierType*>(
|
||||
reinterpret_cast<uint8_t*>(params.reduction_workspace_) + reduction_workspace_size);
|
||||
|
||||
AccumulatorArrayT* reduction_workspace_array = reinterpret_cast<AccumulatorArrayT*>(group_reduction_workspace);
|
||||
AccumulatorArrayT* accumulator_array = reinterpret_cast<AccumulatorArrayT*>(&accumulators);
|
||||
int barrier_group_thread_idx = threadIdx.x % BarrierManager::ThreadCount;
|
||||
if (work_tile_info.is_reduction_unit()) {
|
||||
plus<AccumulatorArrayT> add_fragments;
|
||||
auto peer_offset = size(accumulators) * num_barriers * BarrierManager::ThreadCount;
|
||||
|
||||
if (!work_tile_info.is_final_split(params.divmod_tiles_per_output_tile_.divisor)) {
|
||||
if (work_tile_info.K_idx == 0) {
|
||||
// First peer initializes the workspace partials
|
||||
// Wait until the peers collaborating on this output tile have all written
|
||||
// their accumulators to workspace.
|
||||
uint32_t num_peers = last_peer_id - first_peer_id + 1;
|
||||
BarrierManager::wait_eq(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, num_peers);
|
||||
|
||||
// Load the first peer's data
|
||||
BlockStripedReduceT::load(*accumulator_array, reduction_workspace_array, barrier_group_thread_idx);
|
||||
|
||||
for (int i = 1; i < num_peers; ++i) {
|
||||
// Load peer fragment
|
||||
AccumulatorArrayT addend_fragment;
|
||||
auto peer_reduction_workspace = reinterpret_cast<AccumulatorArrayT*>(group_reduction_workspace + (i * peer_offset));
|
||||
|
||||
BlockStripedReduceT::load(addend_fragment, peer_reduction_workspace, barrier_group_thread_idx);
|
||||
|
||||
// Add peer fragment
|
||||
*accumulator_array = add_fragments(*accumulator_array, addend_fragment);
|
||||
}
|
||||
}
|
||||
else if (!compute_epilogue(work_tile_info, params)) {
|
||||
if (params.requires_separate_reduction() || work_tile_info.K_idx == 0) {
|
||||
// The first peer initializes the workspace partials in the non-separate-reduction case,
|
||||
// and all peers write to their own location in workspace when using separate reduction
|
||||
BlockStripedReduceT::store(reduction_workspace_array, *accumulator_array, barrier_group_thread_idx);
|
||||
}
|
||||
else {
|
||||
if (params.reduction_mode_ == ReductionMode::Deterministic) {
|
||||
// Wait until the preceding split added its accumulators
|
||||
BarrierManager::wait_eq(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, work_tile_info.K_idx);
|
||||
}
|
||||
else {
|
||||
// Wait until the first split has stored its accumulators. Note that the first split will have
|
||||
// accumulated a value into the lock potentially greater than one (since the locked value is
|
||||
// incremented by work_tile_info.k_tile_count below for both the deterministic and non-deterministic)
|
||||
// cases. For non-deterministic reductions, all that non-first or last splits care about is whether
|
||||
// the first split has been written, so we only wait while the locked value is less than 1. This
|
||||
// avoids having to add logic to determine the work_tile_info.k_tile_count for the first split.
|
||||
BarrierManager::wait_lt(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, 1);
|
||||
}
|
||||
// Wait until the preceding split added its accumulators
|
||||
BarrierManager::wait_eq(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, work_tile_info.K_idx);
|
||||
|
||||
// Perform reduction in workspace
|
||||
BlockStripedReduceT::reduce(reduction_workspace_array, *accumulator_array, barrier_group_thread_idx);
|
||||
}
|
||||
|
||||
// If separate reduction is being performed, each participating stream-K unit increments the barrier
|
||||
// by only 1. Otherwise, increment by the K tile count that this unit has processed.
|
||||
int32_t increment = params.requires_separate_reduction() ? 1 : work_tile_info.k_tile_count;
|
||||
|
||||
// Signal our arrival
|
||||
BarrierManager::arrive_inc(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, work_tile_info.k_tile_count);
|
||||
BarrierManager::arrive_inc(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, increment);
|
||||
}
|
||||
else {
|
||||
// Wait until the preceding split added its accumulators.
|
||||
// For both the deterministic and non-deterministic case, each preceding split will have incremented
|
||||
// the locked value by work_tile_info.k_tile_count. Thus, the final split konws that it can begin
|
||||
// loading the partially-reduced value when the locked value reaches its starting K tile index (i.e.,
|
||||
// work_tile_info.K_idx).
|
||||
BarrierManager::wait_eq(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, work_tile_info.K_idx);
|
||||
if (params.reduction_mode_ == ReductionMode::Deterministic) {
|
||||
// Wait until the preceding split added its accumulators
|
||||
BarrierManager::wait_eq(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, work_tile_info.K_idx);
|
||||
}
|
||||
else {
|
||||
// Wait unitl the first split has stored its accumulators
|
||||
BarrierManager::wait_lt(barrier_idx, lock_workspace, barrier_group_thread_idx, lock_idx, 1);
|
||||
}
|
||||
|
||||
// The block computing the final split for the tile adds previously-reduced partials
|
||||
// to its accumulators and computes the epilogue.
|
||||
@@ -406,7 +513,13 @@ public:
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
compute_epilogue(WorkTileInfo const& work_tile_info, Params const& params) {
|
||||
return work_tile_info.is_final_split(params.divmod_tiles_per_output_tile_.divisor);
|
||||
// `is_final_split` will be set to `true` for the following scenarios, all of which must compute the epilogue:
|
||||
// 1. The tile is computed in data-parallel mode
|
||||
// 2. The tile is computed in split-/stream-K mode and this work unit represents the final split of the tile
|
||||
// 3. The tile is computed in split-/stream-K mode and separate reduction is used, and this is a separate reduction unit
|
||||
return work_tile_info.is_valid() &&
|
||||
(work_tile_info.is_final_split(params.divmod_tiles_per_output_tile_.divisor) &&
|
||||
!params.requires_separate_reduction()) || work_tile_info.is_separate_reduction;
|
||||
}
|
||||
|
||||
// Returns the linearized index of the output tile corresponding to the tile with offset [L, M, K]
|
||||
@@ -432,7 +545,8 @@ public:
|
||||
Arguments const& args,
|
||||
ProblemShape problem_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t mma_warp_groups) {
|
||||
uint32_t mma_warp_groups,
|
||||
const uint32_t epilogue_subtile = 1) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
@@ -451,9 +565,11 @@ public:
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
args.decomposition_mode,
|
||||
mma_warp_groups,
|
||||
sizeof_bits<BarrierType>::value,
|
||||
sizeof_bits<ElementAccumulator>::value
|
||||
sizeof_bits<ElementAccumulator>::value,
|
||||
epilogue_subtile
|
||||
);
|
||||
}
|
||||
|
||||
@@ -464,8 +580,9 @@ public:
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
ProblemShape const& problem_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t mma_warp_groups) {
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t mma_warp_groups,
|
||||
const uint32_t epilogue_subtile = 1) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
@@ -486,9 +603,11 @@ public:
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
args.decomposition_mode,
|
||||
mma_warp_groups,
|
||||
sizeof_bits<BarrierType>::value,
|
||||
sizeof_bits<ElementAccumulator>::value
|
||||
sizeof_bits<ElementAccumulator>::value,
|
||||
epilogue_subtile
|
||||
);
|
||||
}
|
||||
|
||||
@@ -505,6 +624,7 @@ public:
|
||||
return work_tile_info.K_idx;
|
||||
}
|
||||
|
||||
private:
|
||||
// Sets the current stream-K work to compute within work_tile_info. If new_unit is true, work_tile_info
|
||||
// is populated as a new unit of work. Otherwise, state existing in work_tile_info (e.g., remaining
|
||||
// iterations) is used to find the next tile in the current work unit.
|
||||
@@ -515,10 +635,22 @@ public:
|
||||
uint64_t linear_idx,
|
||||
WorkTileInfo& work_tile_info) {
|
||||
|
||||
uint64_t true_tile_id = linear_idx;
|
||||
if (linear_idx >= params.sk_units_ && params.splits_ == 1) {
|
||||
uint64_t output_tile_id = linear_idx;
|
||||
if (linear_idx >= params.units_per_problem_ * params.splits_) {
|
||||
// Separate-reduction work
|
||||
auto cluster_size = params.get_cluster_size();
|
||||
// Divide up the linearized separate reduction units into clusters
|
||||
auto cluster_linear_reduction_unit_idx = params.div_cluster_size((linear_idx - params.units_per_problem_));
|
||||
uint64_t cluster_tile_idx, epi_subtile_idx;
|
||||
params.divmod_epilogue_subtile_(cluster_tile_idx, epi_subtile_idx, cluster_linear_reduction_unit_idx);
|
||||
// Bring the linearized tile ID back into the space of tiles, rather than clusters
|
||||
output_tile_id = cluster_tile_idx * cluster_size;
|
||||
|
||||
work_tile_info.setup_separate_reduction(epi_subtile_idx);
|
||||
}
|
||||
else if (linear_idx >= params.sk_units_ && params.splits_ == 1) {
|
||||
// Data-parallel work
|
||||
true_tile_id = linear_idx - params.sk_units_ + params.sk_tiles_;
|
||||
output_tile_id = linear_idx - params.sk_units_ + params.sk_tiles_;
|
||||
work_tile_info.K_idx = 0;
|
||||
work_tile_info.k_tile_count = params.divmod_tiles_per_output_tile_.divisor;
|
||||
work_tile_info.k_tile_remaining = params.divmod_tiles_per_output_tile_.divisor;
|
||||
@@ -540,48 +672,114 @@ public:
|
||||
// To do so, we divide up the linearized stream-K units into clusters and share the same K
|
||||
// offsets for work within clusters.
|
||||
|
||||
// Equivalent to linear_idx / cluster_size
|
||||
auto cluster_linear_work_idx = params.divmod_cluster_shape_minor_.divide(
|
||||
params.divmod_cluster_shape_major_.divide(linear_idx)
|
||||
);
|
||||
auto cluster_linear_work_idx = params.div_cluster_size(linear_idx);
|
||||
|
||||
uint64_t group_idx;
|
||||
params.divmod_sk_groups_(cluster_linear_work_idx, group_idx, cluster_linear_work_idx);
|
||||
|
||||
// Determine whether we are in a "big group" that will process an additional
|
||||
// stream-K cluster tile.
|
||||
auto sk_cluster_tiles = params.div_cluster_size(params.sk_tiles_);
|
||||
auto sk_cluster_tiles_in_group = params.divmod_sk_groups_.divide(sk_cluster_tiles);
|
||||
if (group_idx < params.big_groups_) {
|
||||
++sk_cluster_tiles_in_group;
|
||||
}
|
||||
|
||||
// Determine whether we are in a "big unit" within the group, that will process
|
||||
// an additional K chunk in the group.
|
||||
auto sk_tiles_in_group = sk_cluster_tiles_in_group * params.get_cluster_size();
|
||||
auto k_tiles_in_group = sk_tiles_in_group * params.divmod_tiles_per_output_tile_.divisor;
|
||||
auto k_tiles_per_unit_in_group = params.divmod_sk_units_per_group_.divide(k_tiles_in_group);
|
||||
auto big_units_in_group = params.div_cluster_size(
|
||||
k_tiles_in_group - (k_tiles_per_unit_in_group * params.divmod_sk_units_per_group_.divisor));
|
||||
|
||||
uint64_t split;
|
||||
params.divmod_clusters_mnl_(split, cluster_linear_work_idx, cluster_linear_work_idx);
|
||||
auto big_unit_cmp = params.splits_ > 1 ? split : cluster_linear_work_idx;
|
||||
auto linear_idx_mult = params.splits_ > 1 ? params.divmod_tiles_per_output_tile_.divisor : params.k_tiles_per_sk_unit_;
|
||||
|
||||
bool is_split_k = params.splits_ > 1;
|
||||
auto big_unit_cmp_lhs = is_split_k ? split : cluster_linear_work_idx;
|
||||
auto big_unit_cmp_rhs = is_split_k ? params.big_units_ : big_units_in_group;
|
||||
auto linear_idx_mult = is_split_k ? params.divmod_tiles_per_output_tile_.divisor : k_tiles_per_unit_in_group;
|
||||
auto k_tiles_per_split = is_split_k ? params.k_tiles_per_sk_unit_ : k_tiles_per_unit_in_group;
|
||||
|
||||
// Determine the starting k iteration computed by this stream-K work unit
|
||||
uint32_t unit_iter_start = (linear_idx_mult * cluster_linear_work_idx) + (params.k_tiles_per_sk_unit_ * split);
|
||||
uint32_t unit_iter_start = (linear_idx_mult * cluster_linear_work_idx) +
|
||||
(k_tiles_per_split * split);
|
||||
|
||||
// Adjust the starting position and number of k iterations for "big units," which
|
||||
// compute one extra iteration. These are the first big_units_ units in the
|
||||
// linearized ID space.
|
||||
bool is_big_unit = big_unit_cmp < params.big_units_;
|
||||
if (is_big_unit) {
|
||||
// compute one extra iteration. If there are any big units, they will be the first
|
||||
// in the linearized ID space.
|
||||
auto k_tiles_in_my_split = k_tiles_per_split;
|
||||
if (big_unit_cmp_lhs < big_unit_cmp_rhs) {
|
||||
// Since the "big units" are the first units in the linearized ID space, each
|
||||
// of the units preceding this big unit computed one extra iteration. Thus,
|
||||
// we must offset our start iteration by the number of units that precede
|
||||
// the current unit in the linearized ID space.
|
||||
unit_iter_start += big_unit_cmp;
|
||||
unit_iter_start += big_unit_cmp_lhs;
|
||||
++k_tiles_in_my_split;
|
||||
}
|
||||
else {
|
||||
// Increment by one for each of the big clusters (since all big units precede this unit)
|
||||
unit_iter_start += params.big_units_;
|
||||
unit_iter_start += big_unit_cmp_rhs;
|
||||
}
|
||||
|
||||
if (!is_split_k) {
|
||||
// Adjust the unit starting position and number of tiles to avoid
|
||||
// computing splits of size less than min_iters_per_sk_unit_
|
||||
int unused, start_tile_k_tile;
|
||||
params.divmod_tiles_per_output_tile_(unused, start_tile_k_tile, unit_iter_start);
|
||||
if (start_tile_k_tile < Params::min_iters_per_sk_unit_) {
|
||||
// Starting K tile is in range [0, Params::min_iters_per_sk_unit_), which means that another
|
||||
// stream-K unit will be computing a split with fewer than Params::min_iters_per_sk_unit_ K tiles.
|
||||
// Adjust our work to take over these K tiles.
|
||||
unit_iter_start -= start_tile_k_tile;
|
||||
k_tiles_in_my_split += start_tile_k_tile;
|
||||
}
|
||||
else if (start_tile_k_tile > (params.divmod_tiles_per_output_tile_.divisor - Params::min_iters_per_sk_unit_)) {
|
||||
// Starting K tile is within the final Params::min_iters_per_sk_unit_ K tiles of some output tile,
|
||||
// which means that this unit will compute a split with fewer than Params::min_iters_per_sk_unit_ K tiles.
|
||||
// Adjust our work to shed these K tiles to a neighboring stream-K unit that will compute more consecutive K tiles.
|
||||
auto adjustment_tiles = (params.divmod_tiles_per_output_tile_.divisor - start_tile_k_tile);
|
||||
unit_iter_start += adjustment_tiles;
|
||||
k_tiles_in_my_split -= adjustment_tiles;
|
||||
}
|
||||
}
|
||||
|
||||
if (work_tile_info.k_tile_count == 0) {
|
||||
// This is a new unit
|
||||
work_tile_info.k_tile_remaining = params.k_tiles_per_sk_unit_;
|
||||
|
||||
// Only adjust iteration count for big unit if we are initializing this
|
||||
// work unit. For existing work units, the extra iteration for big units
|
||||
// has already been accounted for in k_tiles_reamaining
|
||||
if (is_big_unit) {
|
||||
++work_tile_info.k_tile_remaining;
|
||||
if (!is_split_k) {
|
||||
//
|
||||
// Adjust the unit ending position and number of tiles to avoid
|
||||
// computing splits of size less than min_iters_per_sk_unit_
|
||||
//
|
||||
|
||||
// Begin by assuming that no adjustment is needed
|
||||
auto initial_unit_iter_end = unit_iter_start + k_tiles_in_my_split;
|
||||
|
||||
int unused, end_tile_k_tile;
|
||||
params.divmod_tiles_per_output_tile_(unused, end_tile_k_tile, initial_unit_iter_end);
|
||||
|
||||
if (end_tile_k_tile < Params::min_iters_per_sk_unit_) {
|
||||
// Ending K tile is within the first Params::min_iters_per_sk_unit_ K tiles of some output tile,
|
||||
// which means that this unit will compute a split with fewer than Params::min_iters_per_sk_unit_ K tiles.
|
||||
// Adjust our work to shed these K tiles to a neighboring stream-K unit that will compute more consecutive K tiles.
|
||||
k_tiles_in_my_split -= end_tile_k_tile;
|
||||
}
|
||||
else if (end_tile_k_tile > (params.divmod_tiles_per_output_tile_.divisor - Params::min_iters_per_sk_unit_)) {
|
||||
// Ending K tile is within the final Params::min_iters_per_sk_unit_ K tiles of some output tile,
|
||||
// which means that some other unit will compute a split with fewer than Params::min_iters_per_sk_unit_ K tiles.
|
||||
// Adjust our work to take on these K tiles.
|
||||
k_tiles_in_my_split += (params.divmod_tiles_per_output_tile_.divisor - end_tile_k_tile);
|
||||
}
|
||||
}
|
||||
|
||||
work_tile_info.k_tile_remaining = k_tiles_in_my_split;
|
||||
}
|
||||
|
||||
// Find the output tile corresponding to the final k iteration covered by this
|
||||
uint32_t unit_iter_end = unit_iter_start + work_tile_info.k_tile_remaining - 1;
|
||||
|
||||
// Find the output tile corresponding to the final k tile covered by this
|
||||
// work unit. Stream-K work units will work backwards in terms of the tiles they
|
||||
// are responsible computing. This is beneficial because the final (partial)
|
||||
// tile computed by a stream-K block is typically the beginning of the output
|
||||
@@ -590,43 +788,45 @@ public:
|
||||
// other work units computing portions of that output tile, it is preferable
|
||||
// for them to be computed later, so as to reduce the likelihood of blocking
|
||||
// on other work.
|
||||
uint32_t unit_iter_end = unit_iter_start + work_tile_info.k_tile_remaining - 1;
|
||||
|
||||
true_tile_id = params.divmod_tiles_per_output_tile_.divide(unit_iter_end);
|
||||
uint32_t true_tile_iter_start = true_tile_id * params.divmod_tiles_per_output_tile_.divisor;
|
||||
uint32_t true_tile_iter_end = true_tile_iter_start + params.divmod_tiles_per_output_tile_.divisor;
|
||||
auto output_tile_id_in_group = params.divmod_tiles_per_output_tile_.divide(unit_iter_end);
|
||||
uint32_t output_tile_iter_start = output_tile_id_in_group * params.divmod_tiles_per_output_tile_.divisor;
|
||||
uint32_t output_tile_iter_end = output_tile_iter_start + params.divmod_tiles_per_output_tile_.divisor;
|
||||
|
||||
// Convert the output tile from the linearized space within each group to the
|
||||
// overall linearized space.
|
||||
output_tile_id = (output_tile_id_in_group * params.divmod_sk_groups_.divisor) + group_idx;
|
||||
|
||||
// Bring the linearized tile ID back into the space of tiles, rather than clusters
|
||||
true_tile_id *= params.divmod_cluster_shape_major_.divisor * params.divmod_cluster_shape_minor_.divisor;
|
||||
output_tile_id *= params.get_cluster_size();
|
||||
|
||||
auto [cta_m_in_cluster, cta_n_in_cluster, _] = cute::block_id_in_cluster();
|
||||
|
||||
// The final linearized tile ID is in units of the cluster dimension over which we rasterize.
|
||||
if (params.raster_order_ == RasterOrder::AlongN) {
|
||||
true_tile_id += cta_n_in_cluster * params.divmod_cluster_shape_minor_.divisor;
|
||||
output_tile_id += cta_n_in_cluster * params.divmod_cluster_shape_minor_.divisor;
|
||||
}
|
||||
else {
|
||||
true_tile_id += cta_m_in_cluster * params.divmod_cluster_shape_minor_.divisor;
|
||||
output_tile_id += cta_m_in_cluster * params.divmod_cluster_shape_minor_.divisor;
|
||||
}
|
||||
|
||||
// The unit's starting k iteration in the current tile is either the starting
|
||||
// iteration for the tile as a whole, or the starting k iteration for the unit
|
||||
// as a whole (if the latter is greater than the former).
|
||||
uint32_t tile_iter_start = max(true_tile_iter_start, unit_iter_start);
|
||||
uint32_t tile_iter_start = max(output_tile_iter_start, unit_iter_start);
|
||||
|
||||
// Similarly, the unit's ending k iteration (exclusive) is either the end of
|
||||
// the current tile it is assigned, or the ending iteration of the unit as a whole
|
||||
// (if the latter is less than the former).
|
||||
uint32_t tile_iter_end = min(true_tile_iter_end, unit_iter_end + 1);
|
||||
uint32_t tile_iter_end = min(output_tile_iter_end, unit_iter_end + 1);
|
||||
|
||||
// Set the k offset to be the starting k tile for this output tile
|
||||
work_tile_info.K_idx = static_cast<int32_t>(tile_iter_start - true_tile_iter_start);
|
||||
|
||||
work_tile_info.K_idx = static_cast<int32_t>(tile_iter_start - output_tile_iter_start);
|
||||
work_tile_info.k_tile_count = tile_iter_end - tile_iter_start;
|
||||
}
|
||||
|
||||
uint64_t work_idx_l, remainder;
|
||||
params.divmod_batch_(work_idx_l, remainder, true_tile_id);
|
||||
params.divmod_batch_(work_idx_l, remainder, output_tile_id);
|
||||
|
||||
uint64_t cta_per_grid_dim = params.divmod_cluster_shape_minor_.divide(remainder);
|
||||
|
||||
@@ -642,7 +842,57 @@ public:
|
||||
work_tile_info.M_idx = work_idx_m;
|
||||
work_tile_info.N_idx = work_idx_n;
|
||||
work_tile_info.L_idx = static_cast<int32_t>(work_idx_l);
|
||||
}
|
||||
|
||||
// Returns the starting and ending peer ID of this tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
static auto
|
||||
tile_peer_range(Params const& params, uint32_t tile_idx, uint32_t cur_k_tile) {
|
||||
auto tile_idx_in_cluster_path = params.div_cluster_size(tile_idx);
|
||||
auto start_k_tile = params.divmod_tiles_per_output_tile_.divisor * tile_idx_in_cluster_path;
|
||||
auto end_k_tile = start_k_tile + params.divmod_tiles_per_output_tile_.divisor - 1;
|
||||
auto big_unit_k_tiles = params.big_units_ * (params.k_tiles_per_sk_unit_ + 1);
|
||||
|
||||
auto adjust_unit = [&](uint32_t k_tile, uint32_t unit_idx, uint32_t k_tiles_per_unit) {
|
||||
auto unit_k_start = unit_idx * k_tiles_per_unit;
|
||||
auto unit_k_end = unit_k_start + k_tiles_per_unit;
|
||||
if (k_tile - start_k_tile < Params::min_iters_per_sk_unit_ &&
|
||||
unit_k_end - start_k_tile < Params::min_iters_per_sk_unit_) {
|
||||
// k_tile is within the first min_iters_per_sk_unit_ K tiles of this output tile,
|
||||
// and the stream-K unit computes fewer than min_iters_per_sk_unit_ K tiles for this
|
||||
// output tile. This work will thus be subsumed by the next stream-K unit.
|
||||
++unit_idx;
|
||||
}
|
||||
|
||||
if (end_k_tile + 1 - k_tile < Params::min_iters_per_sk_unit_ &&
|
||||
end_k_tile + 1 - unit_k_start < Params::min_iters_per_sk_unit_) {
|
||||
// k_tile is within the last min_iters_per_sk_unit_ K tiles of this output tile,
|
||||
// and the stream-K unit computes fewer than min_iters_per_sk_unit_ K tiles for this
|
||||
// output tile. This work will thus be subsumed by the previous stream-K unit.
|
||||
--unit_idx;
|
||||
}
|
||||
|
||||
return unit_idx;
|
||||
};
|
||||
|
||||
// Lambda to find the ID of the stream-K unit that computes this K tile
|
||||
auto find_unit = [&](uint32_t k_tile) {
|
||||
if (k_tile < big_unit_k_tiles) {
|
||||
// The tile is within the "big unit range"
|
||||
auto k_tiles_per_unit = params.k_tiles_per_sk_unit_ + 1;
|
||||
auto unit_idx = k_tile / k_tiles_per_unit;
|
||||
return static_cast<uint64_t>(adjust_unit(k_tile, unit_idx, k_tiles_per_unit));
|
||||
}
|
||||
else {
|
||||
// The tile is after the "big unit range." Account for this by finding the "normal unit"
|
||||
// that it belongs to, and then offsetting by the number of big units
|
||||
auto k_tiles_per_unit = params.k_tiles_per_sk_unit_;
|
||||
auto unit_idx = ((k_tile - big_unit_k_tiles) / params.k_tiles_per_sk_unit_) + (params.big_units_);
|
||||
return static_cast<uint64_t>(adjust_unit(k_tile, unit_idx, k_tiles_per_unit));
|
||||
}
|
||||
};
|
||||
|
||||
return cute::make_tuple(find_unit(start_k_tile), find_unit(cur_k_tile), find_unit(end_k_tile));
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -38,6 +38,7 @@
|
||||
#include "cutlass/detail/dependent_false.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler_stream_k.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler_group.hpp"
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm {
|
||||
@@ -52,6 +53,8 @@ struct PersistentScheduler { };
|
||||
|
||||
struct StreamKScheduler { };
|
||||
|
||||
struct GroupScheduler { }; // Only used for Grouped GEMMs
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm
|
||||
@@ -69,6 +72,7 @@ template <
|
||||
class ArchTag,
|
||||
class TileShape,
|
||||
class ClusterShape
|
||||
, class ProblemShapeType = void
|
||||
>
|
||||
struct TileSchedulerSelector {
|
||||
static_assert(cutlass::detail::dependent_false<ArchTag>,
|
||||
@@ -122,6 +126,21 @@ struct TileSchedulerSelector<
|
||||
using Scheduler = PersistentTileSchedulerSm90StreamK<TileShape, ClusterShape>;
|
||||
};
|
||||
|
||||
template <
|
||||
class TileShape,
|
||||
class ClusterShape
|
||||
, class GroupProblemShape
|
||||
>
|
||||
struct TileSchedulerSelector<
|
||||
GroupScheduler,
|
||||
arch::Sm90,
|
||||
TileShape,
|
||||
ClusterShape
|
||||
, GroupProblemShape
|
||||
> {
|
||||
using Scheduler = PersistentTileSchedulerSm90Group<GroupProblemShape>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::kernel::detail
|
||||
|
||||
@@ -218,7 +218,7 @@ struct PersistentTileSchedulerSm90Params {
|
||||
|
||||
auto possibly_truncate = [&](int x, int y) {
|
||||
if (truncate_by_problem_size) {
|
||||
return cutlass::platform::min(x, y);
|
||||
return platform::min(x, y);
|
||||
}
|
||||
else {
|
||||
return x;
|
||||
@@ -272,7 +272,7 @@ struct PersistentTileSchedulerSm90Params {
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t
|
||||
get_log_swizzle_size(int problem_ctas_m, int problem_ctas_n, int max_swizzle_size) {
|
||||
int min_cta_dim = cutlass::platform::min(problem_ctas_m, problem_ctas_n);
|
||||
int min_cta_dim = platform::min(problem_ctas_m, problem_ctas_n);
|
||||
if (max_swizzle_size >= 8 && min_cta_dim >= 6) {
|
||||
return 3;
|
||||
}
|
||||
@@ -370,6 +370,18 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
Nondeterministic
|
||||
};
|
||||
|
||||
// Strategies for decomposing the problem
|
||||
enum class DecompositionMode {
|
||||
// Use a heuristic to determine whether data-parallel, split-K, or stream-K decomposition should be performed
|
||||
Heuristic,
|
||||
// Force a data-parallel decomposition
|
||||
DataParallel,
|
||||
// Force a split-K decomposition. This should be paired with setting the `splits` parameter
|
||||
SplitK,
|
||||
// Force a stream-K decomposition
|
||||
StreamK
|
||||
};
|
||||
|
||||
using UnderlyingParams = PersistentTileSchedulerSm90Params;
|
||||
using RasterOrder = UnderlyingParams::RasterOrder;
|
||||
using RasterOrderOptions = UnderlyingParams::RasterOrderOptions;
|
||||
@@ -387,6 +399,17 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// and may be overridden in other decompositions.
|
||||
FastDivmodU64 divmod_clusters_mnl_{};
|
||||
|
||||
// We divide up the number of stream-K tiles amongst G groups of stream-K units.
|
||||
// The stream-K units within a group collaborate to comptue over the `sk_tiles / G`
|
||||
// tiles assigned to that group. Non-unit group sizes can help to preserve L2 locality of
|
||||
// partial chunks computed by stream-K units -- units 0 in each group will compute identical K extents
|
||||
// of tiles that would be assigned in the same wave according to the rasterization order of the
|
||||
// data-parallel formulation of the problem.
|
||||
FastDivmodU64 divmod_sk_groups_{};
|
||||
|
||||
// Number of stream-K units in each group
|
||||
FastDivmodU64 divmod_sk_units_per_group_{};
|
||||
|
||||
uint64_t units_per_problem_ = 0;
|
||||
FastDivmod divmod_tiles_per_output_tile_{};
|
||||
int32_t log_swizzle_size_ = 0;
|
||||
@@ -403,6 +426,9 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// at the granularity of a cluster, we store only the number of big clusters.
|
||||
uint32_t big_units_ = 0;
|
||||
|
||||
// The number of groups of stream-K units that will process an extra stream-K tile cluster.
|
||||
uint32_t big_groups_ = 0;
|
||||
|
||||
// Workspace for holding partial accumulators to be reduced across stream-K/split-K units
|
||||
void* reduction_workspace_ = nullptr;
|
||||
|
||||
@@ -419,8 +445,53 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// Strategy to use when reducing between collaborating CTAs
|
||||
ReductionMode reduction_mode_ = ReductionMode::Deterministic;
|
||||
|
||||
// Minimum number of tiled k that can be assigned to a stream-K unit
|
||||
static constexpr uint32_t min_iters_per_sk_unit_ = 4u;
|
||||
// The number of sub blocks in the kernel epilogue
|
||||
FastDivmodU64 divmod_epilogue_subtile_{};
|
||||
|
||||
// The number of blocks that launched for doing separate reduction
|
||||
uint32_t separate_reduction_units_ = 0;
|
||||
|
||||
// Minimum number of k tiles that can be assigned to a stream-K unit
|
||||
static constexpr uint32_t min_iters_per_sk_unit_ = 8u;
|
||||
|
||||
// Maximum number of groups of stream-K units
|
||||
static constexpr uint32_t max_sk_groups_ = 8u;
|
||||
|
||||
// Divides dividend by the cluster size
|
||||
CUTLASS_HOST_DEVICE
|
||||
uint64_t
|
||||
div_cluster_size(uint64_t dividend) const {
|
||||
// Use each underlying fast divmod rather than performing integer division
|
||||
// by the multiplication of major.divisor * minor.divisor
|
||||
return divmod_cluster_shape_minor_.divide(
|
||||
divmod_cluster_shape_major_.divide(dividend)
|
||||
);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
uint64_t
|
||||
get_cluster_size() const {
|
||||
return divmod_cluster_shape_minor_.divisor * divmod_cluster_shape_major_.divisor;
|
||||
}
|
||||
|
||||
// Returns whether the kernel uses separate reduction
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
requires_separate_reduction() const {
|
||||
return separate_reduction_units_ > 0;
|
||||
}
|
||||
|
||||
// Returns the maximum number of peers that can collaborate on a given output tile
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint32_t
|
||||
max_peers_per_tile(uint64_t sk_units, uint64_t sk_tiles) {
|
||||
// When we can divide up our SK units to SK tiles evenly, the number of peers
|
||||
// per SK tile is exactly (sk_units_ / sk_tiles_). In cases where this division
|
||||
// is not exact, some tiles will need to be covered by additional SK units. Because
|
||||
// the extra work can occur at both the beginning and the end of the SK tile, at
|
||||
// most 2 extra peers will be needed.
|
||||
return static_cast<uint32_t>(sk_units / sk_tiles + 2);
|
||||
}
|
||||
|
||||
// Initializes members. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
@@ -434,7 +505,9 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
ReductionMode reduction_mode,
|
||||
void* workspace
|
||||
DecompositionMode decomposition_mode,
|
||||
void* workspace,
|
||||
const uint32_t epilogue_subtile = 1
|
||||
) {
|
||||
dim3 problem_blocks = UnderlyingParams::get_tiled_cta_shape_mnl(
|
||||
problem_shape, tile_shape, cluster_shape);
|
||||
@@ -451,7 +524,9 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
reduction_mode,
|
||||
workspace
|
||||
decomposition_mode,
|
||||
workspace,
|
||||
epilogue_subtile
|
||||
);
|
||||
}
|
||||
|
||||
@@ -468,7 +543,9 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
ReductionMode reduction_mode,
|
||||
void* workspace
|
||||
DecompositionMode decomposition_mode,
|
||||
void* workspace,
|
||||
const uint32_t epilogue_subtile = 1
|
||||
) {
|
||||
UnderlyingParams underlying_params;
|
||||
underlying_params.initialize(
|
||||
@@ -488,7 +565,8 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// Reduction workspace is at the beginning of the workspace. Lock workspace follows.
|
||||
void* reduction_workspace = workspace;
|
||||
|
||||
if (splits > 1) {
|
||||
if (decomposition_mode == DecompositionMode::SplitK ||
|
||||
(decomposition_mode == DecompositionMode::Heuristic && splits > 1)) {
|
||||
// Short circuit to basic split-K decomposition
|
||||
|
||||
// Don't split by more than the available number of SMs
|
||||
@@ -531,24 +609,9 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
uint64_t ctas_per_wave = grid.x * grid.y;
|
||||
|
||||
// The number of output tiles to be computed in stream-K and data-parallel fashion, respectively.
|
||||
uint32_t sk_tiles = get_num_sk_tiles(output_tiles, ctas_per_wave, k_tiles_per_output_tile);
|
||||
uint32_t sk_tiles = get_num_sk_tiles(output_tiles, ctas_per_wave, k_tiles_per_output_tile, decomposition_mode);
|
||||
uint64_t dp_tiles = output_tiles - sk_tiles;
|
||||
|
||||
if (sk_tiles == 0) {
|
||||
// Short circuit to basic data-parallel decomposition
|
||||
set_params_basic(
|
||||
underlying_params,
|
||||
problem_blocks_m,
|
||||
problem_blocks_n,
|
||||
problem_blocks_l,
|
||||
/* splits = */ 1,
|
||||
k_tiles_per_output_tile,
|
||||
reduction_workspace,
|
||||
reduction_mode
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
// Calculate the number of work units covering the data-parallel and stream-K tiles.
|
||||
// A "work unit" is a single index in the linearized ID space used by the scheduler.
|
||||
// We distinguish it from a "block," which is typically tied to a hardware unit
|
||||
@@ -576,12 +639,127 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
uint64_t min_sized_sk_units = (k_tiles_sk_total / min_iters_per_sk_unit_);
|
||||
min_sized_sk_units = (min_sized_sk_units / cluster_size) * cluster_size;
|
||||
|
||||
uint64_t sk_units = cutlass::platform::min(ctas_per_wave, min_sized_sk_units);
|
||||
uint64_t sk_units = platform::min(ctas_per_wave, min_sized_sk_units);
|
||||
|
||||
// If the number of stream-K units is a multiple of the number of stream-K tiles, then
|
||||
// the problem can leverage a basic split-K decomposition for the stream-K tiles.
|
||||
if (sk_tiles < sk_units && sk_units % sk_tiles == 0) {
|
||||
// Short circuit to basic split-K decomposition
|
||||
if (decomposition_mode == DecompositionMode::DataParallel ||
|
||||
(decomposition_mode == DecompositionMode::Heuristic && sk_tiles == 0) ||
|
||||
sk_units == 0) {
|
||||
// Short circuit to basic data-parallel decomposition
|
||||
set_params_basic(
|
||||
underlying_params,
|
||||
problem_blocks_m,
|
||||
problem_blocks_n,
|
||||
problem_blocks_l,
|
||||
/* splits = */ 1,
|
||||
k_tiles_per_output_tile,
|
||||
reduction_workspace,
|
||||
reduction_mode
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
bool do_separate_reduction = should_perform_separate_reduction(
|
||||
epilogue_subtile, sk_units, sk_tiles, dp_tiles, ctas_per_wave);
|
||||
|
||||
// Determine the number of stream-K groups that will be used. We currently use
|
||||
// max_sk_groups_ unless this extends beyond the extent of the dimension over
|
||||
// which the problem is rasterized. For example, if the tiled problem shape
|
||||
// (in CTA_M x CTA_N representation) when using 1x1 clusters is 4x16,
|
||||
// and we rasterize along the M dimension, we choose 4 groups, rather than 8.
|
||||
// If the cluster shape is 2x1, we choose 2 groups (CTA_M / CLUSTER_M).
|
||||
uint32_t max_groups_problem;
|
||||
if (underlying_params.raster_order_ == RasterOrder::AlongM) {
|
||||
max_groups_problem = problem_blocks_m / cluster_shape.m();
|
||||
}
|
||||
else {
|
||||
max_groups_problem = problem_blocks_n / cluster_shape.n();
|
||||
}
|
||||
|
||||
// Select the number of groups that will be use. We start with the maximum
|
||||
// number of potential groups, and iterate down looking for a group size that
|
||||
// evenly divides the stream-K units and tiles, and for which the resulting
|
||||
// number of K tiles per stream-K unit remains above min_iters_per_sk_unit_
|
||||
|
||||
uint32_t groups = platform::min(max_groups_problem, uint32_t(max_sk_groups_));
|
||||
|
||||
// Grouping is disabled when separate reduction is used
|
||||
if (do_separate_reduction) {
|
||||
groups = 1;
|
||||
}
|
||||
|
||||
uint32_t fallback_groups = 0;
|
||||
auto sk_cluster_tiles = sk_tiles / cluster_size;
|
||||
auto sk_cluster_units = sk_units / cluster_size;
|
||||
|
||||
auto sk_splits_too_small = [&](uint32_t g) {
|
||||
// Check whether the number of K tiles computed per stream-K unit is less
|
||||
// than min_iters_per_sk_unit_
|
||||
auto total_sk_k_tiles = (sk_tiles / g) * k_tiles_per_output_tile;
|
||||
auto k_tiles_per_sk_unit = total_sk_k_tiles / (sk_units / g);
|
||||
return k_tiles_per_sk_unit < min_iters_per_sk_unit_;
|
||||
};
|
||||
|
||||
auto is_ideal_grouping = [&](uint32_t g) {
|
||||
// An ideal grouping will evenly divide stream-K clusters, evenly divide
|
||||
// stream-K tiles, and not result in stream-K splits that are too small.
|
||||
return (sk_cluster_units % g == 0) && (sk_cluster_tiles % g == 0) && !sk_splits_too_small(g);
|
||||
};
|
||||
|
||||
auto is_valid_grouping = [&](uint32_t g) {
|
||||
// A grouping is valid, but not ideal, if it evenly divides the
|
||||
// stream-K clusters and does not result in stream-K splits that are
|
||||
// too small. Such a setting can be used as a fallback option in the
|
||||
// case that an ideal grouping is not achievable
|
||||
return sk_cluster_units % g == 0 && !sk_splits_too_small(g);
|
||||
};
|
||||
|
||||
while (groups > 1 && !is_ideal_grouping(groups)) {
|
||||
if (fallback_groups == 0 && is_valid_grouping(groups)) {
|
||||
// Set fallback groups once in preference for a larger number of groups.
|
||||
fallback_groups = groups;
|
||||
}
|
||||
--groups;
|
||||
}
|
||||
|
||||
// If groups == 1, we did not find a group count that satisfies all criteria. If we have
|
||||
// found a fallback group count, use this instead.
|
||||
if (groups == 1 && fallback_groups > 0) {
|
||||
groups = fallback_groups;
|
||||
}
|
||||
|
||||
auto sk_units_per_group = sk_units / groups;
|
||||
|
||||
// sk_tiles is guaranteed to be divisible by cluster_size because it is calculated as:
|
||||
// sk_tiles = (waves <= 2) ? total_tiles : (sm_count + (total_tiles % sm_count))
|
||||
// Both total_tiles and sm_count are multiples of cluster size due to padding added
|
||||
// prior to kernel launch.
|
||||
uint64_t sk_clustered_tiles = sk_tiles / cluster_size;
|
||||
uint64_t sk_clustered_tiles_per_group = sk_clustered_tiles / groups;
|
||||
uint64_t sk_tiles_per_group = sk_clustered_tiles_per_group * cluster_size;
|
||||
|
||||
// Groups that will process an extra stream-K tile cluster. These differ from "big_units," which
|
||||
// are stream-K units within a group that process an extra K chunk.
|
||||
uint64_t sk_big_groups = sk_clustered_tiles % groups;
|
||||
|
||||
uint64_t k_tiles_per_group = k_tiles_per_output_tile * sk_tiles_per_group;
|
||||
|
||||
// Number of k tiles computed per stream-K unit
|
||||
uint64_t k_tiles_per_sk_unit = k_tiles_per_group / sk_units_per_group;
|
||||
|
||||
uint32_t reduction_units = 0;
|
||||
|
||||
// Use separate reduction when we have less than one wave of output tiles (dp_tiles == 0)
|
||||
// and when each tile will be operated on by at least two stream-K units (sk_units > 2 * sk_tiles)
|
||||
if (do_separate_reduction) {
|
||||
// Each reduction unit will reduce the partials of an epilogue subtile for
|
||||
// a given output tile and compute the epilogue. Thus, there are as many reduction
|
||||
// units as there are epilogue subtiles.
|
||||
reduction_units = sk_tiles * epilogue_subtile;
|
||||
}
|
||||
else if (decomposition_mode == DecompositionMode::Heuristic && sk_tiles < sk_units && sk_units % sk_tiles == 0) {
|
||||
// If the number of stream-K units is a multiple of the number of stream-K tiles, then
|
||||
// the problem can leverage a basic split-K decomposition for the stream-K tiles.
|
||||
// This case happens when separate reduction is disable.
|
||||
uint32_t sk_splits = static_cast<uint32_t>(sk_units / sk_tiles);
|
||||
set_params_basic(
|
||||
underlying_params,
|
||||
@@ -595,37 +773,13 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
// Number of k iterations computed per stream-K units
|
||||
uint64_t k_tiles_per_sk_unit = k_tiles_sk_total / sk_units;
|
||||
|
||||
// Number of stream-K units that need to compute extra iterations in order to cover
|
||||
// the residual k iterations. This assumes that each such unit computes one additional
|
||||
// iteration.
|
||||
uint64_t sk_big_units = k_tiles_sk_total - (k_tiles_per_sk_unit * sk_units);
|
||||
|
||||
// The division below is guaranteed to be exact because sk_big_units is guaranteed
|
||||
// to be a multiple of cluster_size. This is useful because
|
||||
// it allows us to use a block's linearized cluster ID to determine whether it is
|
||||
// a big block. The reasoning behind this guarnatee is explained as follows:
|
||||
// sk_big_units = k_tiles_sk_total - (k_tiles_per_sk_unit * sk_units);
|
||||
//
|
||||
// - k_tiles_sk_total is a multiple of cluster_size because it is the product
|
||||
// of number of tail tiles and the number of k iterations per tile. Because
|
||||
// both the number of output tiles and number of available SMs are rounded
|
||||
// to be multiples of cluster shape, the number of tail tiles
|
||||
// (output_tiles % avail_sms) is a multpile of cluster_size.
|
||||
//
|
||||
// - sk_units is a multiple of cluster_size because it is either blocks_per_wave
|
||||
// or 0, and blocks_per_wave is a multiple of the cluster_size due to the grid-planning
|
||||
// logic rounding to multiples of cluster dimensions
|
||||
uint64_t sk_big_units_per_cluster = sk_big_units / cluster_size;
|
||||
|
||||
divmod_cluster_shape_major_ = underlying_params.divmod_cluster_shape_major_;
|
||||
divmod_cluster_shape_minor_ = underlying_params.divmod_cluster_shape_minor_;
|
||||
divmod_batch_ = underlying_params.divmod_batch_;
|
||||
divmod_tiles_per_output_tile_ = FastDivmod(k_tiles_per_output_tile);
|
||||
divmod_cluster_blk_major_ = underlying_params.divmod_cluster_blk_major_;
|
||||
divmod_sk_groups_ = FastDivmodU64(static_cast<uint64_t>(groups));
|
||||
divmod_sk_units_per_group_ = FastDivmodU64(static_cast<uint64_t>(sk_units / groups));
|
||||
|
||||
// Override divmod_clusters_mnl_ to be the number of cluster-sized stream-K units.
|
||||
// This setting ensures that the use of this divmod for stream-K decompositions
|
||||
@@ -635,12 +789,19 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
log_swizzle_size_ = underlying_params.log_swizzle_size_;
|
||||
units_per_problem_ = static_cast<uint32_t>(dp_units + sk_units);
|
||||
raster_order_ = underlying_params.raster_order_;
|
||||
big_units_ = static_cast<uint32_t>(sk_big_units_per_cluster);
|
||||
|
||||
// Assign big_units_ assuming that group count == 1. This is unused by stream-K
|
||||
// when group count > 1.
|
||||
big_units_ = static_cast<uint32_t>(k_tiles_per_group % k_tiles_per_sk_unit);
|
||||
|
||||
big_groups_ = static_cast<uint32_t>(sk_big_groups);
|
||||
reduction_workspace_ = reduction_workspace;
|
||||
sk_tiles_ = sk_tiles;
|
||||
sk_units_ = static_cast<uint32_t>(sk_units);
|
||||
k_tiles_per_sk_unit_ = static_cast<uint32_t>(k_tiles_per_sk_unit);
|
||||
reduction_mode_ = reduction_mode;
|
||||
divmod_epilogue_subtile_ = FastDivmodU64(epilogue_subtile);
|
||||
separate_reduction_units_ = reduction_units;
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
@@ -696,23 +857,28 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// Returns the number of stream-K tiles that will be computed amongst `output_tiles` total
|
||||
// output tiles on a device with `ctas_per_wave` CTAs in each wave.
|
||||
static uint32_t
|
||||
get_num_sk_tiles(uint64_t output_tiles, uint64_t ctas_per_wave, uint32_t k_tiles_per_output_tile) {
|
||||
get_num_sk_tiles(uint64_t output_tiles, uint64_t ctas_per_wave, uint32_t k_tiles_per_output_tile, DecompositionMode decomposition_mode) {
|
||||
uint32_t full_waves = static_cast<uint32_t>(output_tiles / ctas_per_wave);
|
||||
uint32_t total_waves = static_cast<uint32_t>((output_tiles + ctas_per_wave - 1) / ctas_per_wave);
|
||||
|
||||
if (full_waves == total_waves || k_tiles_per_output_tile <= min_iters_per_sk_unit_) {
|
||||
// All tiles will be data-parallel tiles if there is either no quantization
|
||||
// or if there is no work to be split.
|
||||
if (decomposition_mode == DecompositionMode::DataParallel ||
|
||||
decomposition_mode == DecompositionMode::SplitK) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
//
|
||||
// The final wave is not full. Perform some stream-K work.
|
||||
//
|
||||
if (decomposition_mode == DecompositionMode::Heuristic) {
|
||||
if (full_waves == total_waves || k_tiles_per_output_tile <= min_iters_per_sk_unit_) {
|
||||
// All tiles will be data-parallel tiles if there is either no quantization
|
||||
// or if there is no work to be split.
|
||||
return 0;
|
||||
}
|
||||
|
||||
// Rudimentary heuristic: prefer data-parallel decomposition if we have more than
|
||||
// one wave and the tail wave is more than half full. This is subject to change.
|
||||
if (full_waves != 0) {
|
||||
//
|
||||
// The final wave is not full. Perform some stream-K work.
|
||||
//
|
||||
|
||||
// Rudimentary heuristic: prefer data-parallel decomposition if we have more than
|
||||
// one wave and the tail wave is more than half full. This is subject to change.
|
||||
uint64_t tail_tiles = output_tiles - (full_waves * ctas_per_wave);
|
||||
if (tail_tiles >= (ctas_per_wave / 2)) {
|
||||
return 0;
|
||||
@@ -729,6 +895,22 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
return static_cast<uint32_t>(output_tiles - dp_tiles);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint64_t
|
||||
get_num_sk_units(GemmCoord cluster_shape, uint64_t ctas_per_wave, uint32_t sk_tiles, uint32_t k_tiles_per_output_tile) {
|
||||
// Number of k iterations computed by the stream-K units as a whole
|
||||
uint64_t k_tiles_sk_total = k_tiles_per_output_tile * sk_tiles;
|
||||
|
||||
// Calculate the number of stream-K units that would be needed if each stream-K unit
|
||||
// computed the minimum allowable k iterations. Truncate this to be in units of clusters.
|
||||
auto cluster_size = cluster_shape.m() * cluster_shape.n();
|
||||
uint64_t min_sized_sk_units = (k_tiles_sk_total / min_iters_per_sk_unit_);
|
||||
min_sized_sk_units = (min_sized_sk_units / cluster_size) * cluster_size;
|
||||
|
||||
uint64_t sk_units = platform::min(ctas_per_wave, min_sized_sk_units);
|
||||
return sk_units;
|
||||
}
|
||||
|
||||
// Calculates the size of the workspace needed for holding reduction barriers
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int
|
||||
@@ -759,9 +941,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
uint32_t mma_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t accumulator_bits) {
|
||||
uint32_t accumulator_bits,
|
||||
uint32_t epilogue_subtile = 1) {
|
||||
|
||||
auto log_swizzle_size = UnderlyingParams::get_log_swizzle_size(problem_blocks.x, problem_blocks.y, max_swizzle);
|
||||
problem_blocks.x = round_up(problem_blocks.x, (1 << log_swizzle_size) * cluster_shape.m());
|
||||
@@ -771,7 +955,12 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// of output tiles that will be split, and then calculate the workspace needed to cover these.
|
||||
uint64_t output_tiles = problem_blocks.x * problem_blocks.y * problem_blocks.z;
|
||||
|
||||
if (splits > 1) {
|
||||
if (decomposition_mode == DecompositionMode::DataParallel) {
|
||||
barrier_workspace_size = 0;
|
||||
reduction_workspace_size = 0;
|
||||
}
|
||||
else if (decomposition_mode == DecompositionMode::SplitK ||
|
||||
(decomposition_mode == DecompositionMode::Heuristic && splits > 1)) {
|
||||
// Basic split-K variant requires workspace for all output tiles
|
||||
barrier_workspace_size = get_barrier_workspace_size(output_tiles, mma_warp_groups, barrier_bits);
|
||||
reduction_workspace_size = get_reduction_workspace_size(output_tiles, tile_shape, accumulator_bits);
|
||||
@@ -794,14 +983,39 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
raster_order_option
|
||||
);
|
||||
uint64_t ctas_per_wave = grid.x * grid.y;
|
||||
uint32_t sk_tiles = get_num_sk_tiles(output_tiles, ctas_per_wave, static_cast<uint32_t>(k_tiles_per_output_tile));
|
||||
uint32_t sk_tiles = get_num_sk_tiles(output_tiles, ctas_per_wave, static_cast<uint32_t>(k_tiles_per_output_tile), decomposition_mode);
|
||||
uint64_t sk_units = get_num_sk_units(cluster_shape, ctas_per_wave, sk_tiles, k_tiles_per_output_tile);
|
||||
uint64_t dp_tiles = output_tiles - sk_tiles;
|
||||
|
||||
uint64_t reduction_tiles = sk_tiles;
|
||||
if (should_perform_separate_reduction(epilogue_subtile, sk_units, sk_tiles, dp_tiles, ctas_per_wave)) {
|
||||
// In separate reduction, each peer writes to its own location in scratch space.
|
||||
// Thus, for separate reduction, we need as many reduction tiles per output tile
|
||||
// as there are the maximum number of peers that can collaborate on an output tile.
|
||||
reduction_tiles *= max_peers_per_tile(sk_units, sk_tiles);
|
||||
}
|
||||
|
||||
// Though separate reduction requires a larger reduction workspace, only one barrier
|
||||
// is needed per output tile. Each peer will increment the barrier by one once the peer has
|
||||
// written its accumulator to scratch space. The separate reduction unit will only begin
|
||||
// performing the reduction when the barrier has reached the number of peers for the output tile.
|
||||
barrier_workspace_size = get_barrier_workspace_size(sk_tiles, mma_warp_groups, barrier_bits);
|
||||
reduction_workspace_size = get_reduction_workspace_size(sk_tiles, tile_shape, accumulator_bits);
|
||||
reduction_workspace_size = get_reduction_workspace_size(reduction_tiles, tile_shape, accumulator_bits);
|
||||
}
|
||||
}
|
||||
#endif // !defined(__CUDACC_RTC__)
|
||||
|
||||
// Returns whether the kernel is configured in a manner for which separate reduction should be used
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
should_perform_separate_reduction(uint32_t epilogue_subtile, uint64_t sk_units, uint64_t sk_tiles, uint64_t dp_tiles, uint64_t ctas_per_wave) {
|
||||
// We perform separate reduction if we have fewer than one wave of output tiles
|
||||
// and each output tile is covered by at least to stream-K units. When sk_units is
|
||||
// multiple of sk_tiles, will choose basic split-k path instead of separate reduction for now.
|
||||
return (epilogue_subtile != 1) && (dp_tiles == 0) && (sk_units > 2u * sk_tiles) &&
|
||||
(sk_units + sk_tiles * epilogue_subtile <= ctas_per_wave);
|
||||
}
|
||||
|
||||
// Get the amount of scratch workspace needed for the kernel. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
static int
|
||||
@@ -813,9 +1027,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
uint32_t mma_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits) {
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t epilogue_subtile) {
|
||||
|
||||
dim3 problem_blocks = UnderlyingParams::get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
uint32_t k_tiles_per_output_tile = (problem_shape.k() + tile_shape.k() - 1) / tile_shape.k();
|
||||
@@ -829,9 +1045,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
mma_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits
|
||||
element_accumulator_bits,
|
||||
epilogue_subtile
|
||||
);
|
||||
}
|
||||
|
||||
@@ -848,9 +1066,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
uint32_t mma_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits) {
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t epilogue_subtile = 1) {
|
||||
|
||||
int barrier_workspace_size = 0;
|
||||
int reduction_workspace_size = 0;
|
||||
@@ -867,9 +1087,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
mma_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits
|
||||
element_accumulator_bits,
|
||||
epilogue_subtile
|
||||
);
|
||||
#endif
|
||||
|
||||
@@ -889,9 +1111,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
uint32_t mma_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits) {
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t epilogue_subtile) {
|
||||
|
||||
dim3 problem_blocks = UnderlyingParams::get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
uint32_t k_tiles_per_output_tile = (problem_shape.k() + tile_shape.k() - 1) / tile_shape.k();
|
||||
@@ -907,9 +1131,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
mma_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits
|
||||
element_accumulator_bits,
|
||||
epilogue_subtile
|
||||
);
|
||||
}
|
||||
|
||||
@@ -928,9 +1154,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
uint32_t mma_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits) {
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t epilogue_subtile = 1) {
|
||||
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
int barrier_workspace_size = 0;
|
||||
@@ -947,9 +1175,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
mma_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits
|
||||
element_accumulator_bits,
|
||||
epilogue_subtile
|
||||
);
|
||||
|
||||
if (barrier_workspace_size > 0) {
|
||||
@@ -982,6 +1212,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
divmod_cluster_shape_minor_ = underlying_params.divmod_cluster_shape_minor_;
|
||||
divmod_batch_ = FastDivmodU64(blocks_m * blocks_n);
|
||||
divmod_tiles_per_output_tile_ = FastDivmod(k_tiles_per_output_tile);
|
||||
divmod_sk_groups_ = FastDivmodU64(1u);
|
||||
auto cluster_size = underlying_params.divmod_cluster_shape_major_.divisor * underlying_params.divmod_cluster_shape_minor_.divisor;
|
||||
divmod_clusters_mnl_ = FastDivmodU64((blocks_m * blocks_n * blocks_l) / cluster_size);
|
||||
splits_ = splits;
|
||||
@@ -997,9 +1228,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// No stream-K work is performed for "basic" data-parallel and split-K decompositions
|
||||
sk_tiles_ = 0;
|
||||
sk_units_ = 0;
|
||||
divmod_sk_units_per_group_ = FastDivmodU64(1u);
|
||||
separate_reduction_units_ = 0;
|
||||
}
|
||||
|
||||
private:
|
||||
private:
|
||||
// Round up number of bytes to the nearest multiple of L2 cache line alignment
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int
|
||||
@@ -1009,6 +1242,236 @@ private:
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Parameters for SM90 persistent group scheduler (only used for Grouped Gemms)
|
||||
template<class ProblemShape>
|
||||
struct PersistentTileSchedulerSm90GroupParams {
|
||||
|
||||
enum class RasterOrder {
|
||||
AlongM,
|
||||
AlongN
|
||||
};
|
||||
|
||||
enum class RasterOrderOptions {
|
||||
Heuristic,
|
||||
AlongM,
|
||||
AlongN
|
||||
};
|
||||
|
||||
FastDivmodU64Pow2 divmod_cluster_shape_major_{};
|
||||
FastDivmodU64Pow2 divmod_cluster_shape_minor_{};
|
||||
FastDivmodU64 divmod_batch_{};
|
||||
|
||||
uint64_t blocks_per_problem_ = 0;
|
||||
int32_t log_swizzle_size_ = 0;
|
||||
RasterOrder raster_order_ = RasterOrder::AlongN;
|
||||
|
||||
int32_t groups_ = 0;
|
||||
ProblemShape* problem_shapes_ = nullptr;
|
||||
GemmCoord cta_shape_;
|
||||
|
||||
// Version of initialize that takes in as input the number of CTAs in the M and N and L dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
void
|
||||
initialize(
|
||||
dim3 problem_blocks,
|
||||
int32_t groups,
|
||||
ProblemShape* problem_shapes,
|
||||
GemmCoord cta_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
|
||||
CUTLASS_UNUSED(hw_info);
|
||||
|
||||
// Round up to nearest multiple of swizzle_size along each mode
|
||||
auto log_swizzle_size = get_log_swizzle_size(problem_blocks.x, problem_blocks.y, max_swizzle_size);
|
||||
auto problem_blocks_m = round_up(problem_blocks.x, (1 << log_swizzle_size) * cluster_shape.m());
|
||||
auto problem_blocks_n = round_up(problem_blocks.y, (1 << log_swizzle_size) * cluster_shape.n());
|
||||
|
||||
RasterOrder raster_order = get_rasterization_order(
|
||||
problem_blocks_m,
|
||||
problem_blocks_n,
|
||||
raster_order_option
|
||||
);
|
||||
|
||||
//
|
||||
// Set members
|
||||
//
|
||||
groups_ = groups;
|
||||
problem_shapes_ = problem_shapes;
|
||||
cta_shape_ = cta_shape;
|
||||
|
||||
blocks_per_problem_ = problem_blocks_m * problem_blocks_n * problem_blocks.z;
|
||||
log_swizzle_size_ = log_swizzle_size;
|
||||
raster_order_ = raster_order;
|
||||
divmod_batch_ = FastDivmodU64(problem_blocks_m * problem_blocks_n);
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
divmod_cluster_shape_major_ = FastDivmodU64Pow2(cluster_shape.n());
|
||||
divmod_cluster_shape_minor_ = FastDivmodU64Pow2(cluster_shape.m());
|
||||
}
|
||||
else {
|
||||
divmod_cluster_shape_major_ = FastDivmodU64Pow2(cluster_shape.m());
|
||||
divmod_cluster_shape_minor_ = FastDivmodU64Pow2(cluster_shape.n());
|
||||
}
|
||||
}
|
||||
|
||||
// Version of get_tiled_cta_shape_mnl that takes in as input the number of CTAs in the M and N dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_tiled_cta_shape_mnl(GemmCoord cluster_shape, uint32_t cta_m, uint32_t cta_n) {
|
||||
// Round up to nearest multiple of cluster dim along each mode
|
||||
auto problem_blocks_m = ((cta_m + cluster_shape.m() - 1) / cluster_shape.m()) * cluster_shape.m();
|
||||
auto problem_blocks_n = ((cta_n + cluster_shape.n() - 1) / cluster_shape.n()) * cluster_shape.n();
|
||||
|
||||
return {
|
||||
static_cast<uint32_t>(problem_blocks_m),
|
||||
static_cast<uint32_t>(problem_blocks_n),
|
||||
static_cast<uint32_t>(1) // Only a single batch per group is currently supported
|
||||
};
|
||||
}
|
||||
|
||||
// Version of get_grid_shape that takes in as input the number of CTAs in the M and N and L dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
CUTLASS_HOST_DEVICE static
|
||||
dim3
|
||||
get_grid_shape(
|
||||
dim3 problem_blocks,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo hw_info,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option,
|
||||
bool truncate_by_problem_size=true) {
|
||||
|
||||
int const sm_count = hw_info.sm_count;
|
||||
|
||||
// Round up to nearest multiple of swizzle_size along each mode
|
||||
auto log_swizzle_size = get_log_swizzle_size(problem_blocks.x, problem_blocks.y, max_swizzle_size);
|
||||
auto problem_blocks_m = round_up(problem_blocks.x, (1 << log_swizzle_size) * cluster_shape.m());
|
||||
auto problem_blocks_n = round_up(problem_blocks.y, (1 << log_swizzle_size) * cluster_shape.n());
|
||||
|
||||
int problem_blocks_total = problem_blocks_m * problem_blocks_n * problem_blocks.z;
|
||||
|
||||
RasterOrder raster_order = get_rasterization_order(
|
||||
problem_blocks_m,
|
||||
problem_blocks_n,
|
||||
raster_order_option
|
||||
);
|
||||
|
||||
dim3 launch_grid;
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid = dim3(cluster_shape.m(), 1, 1);
|
||||
}
|
||||
else {
|
||||
launch_grid = dim3(1, cluster_shape.n(), 1);
|
||||
}
|
||||
|
||||
auto possibly_truncate = [&](int x, int y) {
|
||||
if (truncate_by_problem_size) {
|
||||
return platform::min(x, y);
|
||||
}
|
||||
else {
|
||||
return x;
|
||||
}
|
||||
};
|
||||
|
||||
// The else path is generic, however, we can avoid some divs if we know cluster size is 1
|
||||
auto cluster_size = cluster_shape.m() * cluster_shape.n();
|
||||
if (cluster_size == 1) {
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = possibly_truncate(sm_count, problem_blocks_total);
|
||||
}
|
||||
else {
|
||||
launch_grid.x = possibly_truncate(sm_count, problem_blocks_total);
|
||||
}
|
||||
}
|
||||
else {
|
||||
// Optimal grid size calculation is based on
|
||||
// GH100: 8 GPCs, 72 TPCs (9 TPCs/GPC), 2 SMs/TPC, 144 SMs per full GPU
|
||||
// Hence, maximum SMs per GPC = 18
|
||||
constexpr int max_sm_per_gpc = 18;
|
||||
// Provided SM count could possibly be less than the assumed maximum SMs per GPC
|
||||
auto cluster_size = cluster_shape.m() * cluster_shape.n();
|
||||
int const min_num_gpc = sm_count < max_sm_per_gpc ? 1 : sm_count / max_sm_per_gpc;
|
||||
int const max_cta_occupancy_per_gpc = max_sm_per_gpc - (max_sm_per_gpc % cluster_size);
|
||||
int cta_per_device = min_num_gpc * max_cta_occupancy_per_gpc;
|
||||
|
||||
// The calculation below allows for larger grid size launch for different GPUs.
|
||||
int const num_gpc_residual = sm_count < max_sm_per_gpc ? 0 : sm_count % max_sm_per_gpc;
|
||||
int const max_cta_occupancy_per_residual_gpc = num_gpc_residual - (num_gpc_residual % cluster_size);
|
||||
cta_per_device += max_cta_occupancy_per_residual_gpc;
|
||||
|
||||
cta_per_device = sm_count < cta_per_device ? sm_count : cta_per_device;
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = possibly_truncate(
|
||||
cta_per_device / cluster_shape.m(),
|
||||
problem_blocks_total / cluster_shape.m());
|
||||
}
|
||||
else {
|
||||
launch_grid.x = possibly_truncate(
|
||||
cta_per_device / cluster_shape.n(),
|
||||
problem_blocks_total / cluster_shape.n());
|
||||
}
|
||||
}
|
||||
return launch_grid;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t
|
||||
get_log_swizzle_size(int problem_ctas_m, int problem_ctas_n, int max_swizzle_size) {
|
||||
int min_cta_dim = platform::min(problem_ctas_m, problem_ctas_n);
|
||||
if (max_swizzle_size >= 8 && min_cta_dim >= 6) {
|
||||
return 3;
|
||||
}
|
||||
else if (max_swizzle_size >= 4 && min_cta_dim >= 3) {
|
||||
return 2;
|
||||
}
|
||||
else if (max_swizzle_size >= 2 && min_cta_dim >= 2) {
|
||||
return 1;
|
||||
}
|
||||
else {
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static RasterOrder
|
||||
get_rasterization_order(
|
||||
uint32_t tiles_m,
|
||||
uint32_t tiles_n,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
|
||||
if (raster_order_option == RasterOrderOptions::Heuristic) {
|
||||
if (tiles_n > tiles_m) {
|
||||
return RasterOrder::AlongM;
|
||||
}
|
||||
else {
|
||||
return RasterOrder::AlongN;
|
||||
}
|
||||
}
|
||||
else {
|
||||
switch (raster_order_option) {
|
||||
case RasterOrderOptions::AlongN:
|
||||
return RasterOrder::AlongN;
|
||||
break;
|
||||
default:
|
||||
return RasterOrder::AlongM;
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
} // namespace detail
|
||||
} // namespace kernel
|
||||
|
||||
@@ -114,6 +114,9 @@ struct MmaGeneric {
|
||||
|
||||
static bool const kMultipleOf2 = ((Shape::kM % 2 == 0) && (Shape::kN % 2 == 0));
|
||||
|
||||
static bool const kAllFp32 = platform::is_same<ElementA, float>::value &&
|
||||
platform::is_same<ElementB, float>::value &&
|
||||
platform::is_same<ElementC, float>::value;
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
@@ -144,11 +147,7 @@ struct MmaGeneric {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k = 0; k < Shape::kK; ++k) {
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 860)
|
||||
if (kMultipleOf2 &&
|
||||
platform::is_same<ElementA, float>::value &&
|
||||
platform::is_same<ElementB, float>::value &&
|
||||
platform::is_same<ElementC, float>::value) {
|
||||
|
||||
if (kMultipleOf2 && kAllFp32) {
|
||||
//2x2 zigzag - m and n loops to increment by 2. Inner loop to process 4 multiply-adds in a 2x2 tile.
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < Shape::kN; n+=2) {
|
||||
@@ -157,7 +156,7 @@ struct MmaGeneric {
|
||||
for (int m = 0; m < Shape::kM; m+=2) {
|
||||
|
||||
int m_serpentine = (n % 4) ? (Shape::kM - 2 - m) : m;
|
||||
|
||||
|
||||
//top-left element in 2x2 tile
|
||||
{
|
||||
MatrixCoord mn(m_serpentine, n);
|
||||
|
||||
@@ -60,9 +60,6 @@ namespace layout {
|
||||
//
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tag used for 3-D NWC tensors for 1D conv, only used in 3.x API
|
||||
class TensorNWC {};
|
||||
|
||||
/// Mapping function for 4-D NHWC tensors.
|
||||
class TensorNHWC {
|
||||
public:
|
||||
@@ -632,6 +629,14 @@ public:
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Tag used for linearized tensors with shape (NW, C) for 1D conv, only used in 3.x API
|
||||
class TensorLinearizedNWC {};
|
||||
/// Tag used for linearized tensors with shape (NHW, C) for 2D conv, only used in 3.x API
|
||||
class TensorLinearizedNHWC : public TensorNHWC {};
|
||||
/// Tag used for linearized tensors with shape (NDHW, C) for 3D conv, only used in 3.x API
|
||||
class TensorLinearizedNDHWC : public TensorNDHWC {};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
+1079
-222
File diff suppressed because it is too large
Load Diff
@@ -48,7 +48,8 @@ using namespace cute;
|
||||
|
||||
enum class BarrierStatus : uint32_t {
|
||||
WaitAgain = 0u,
|
||||
WaitDone = 1u
|
||||
WaitDone = 1u,
|
||||
WaitOnly = 2u
|
||||
};
|
||||
|
||||
class ArrivalToken {
|
||||
@@ -81,6 +82,16 @@ private:
|
||||
friend bool operator==(const BarrierStatus& left, const ArrivalToken& right) {
|
||||
return left == right.get();
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
friend bool operator!=(const ArrivalToken& left, const BarrierStatus& right) {
|
||||
return left.get() != right;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
friend bool operator!=(const BarrierStatus& left, const ArrivalToken& right) {
|
||||
return left != right.get();
|
||||
}
|
||||
};
|
||||
|
||||
class ProducerToken : public ArrivalToken {
|
||||
@@ -188,13 +199,9 @@ PipelineState<Pipeline::Stages> make_producer_start_state() {
|
||||
// Assumptions : Constructor is visible Cluster-wide (as it needs a Cluster-Sync)
|
||||
// We have exactly one thread elected in the Producer as the "leader"
|
||||
// Currently, it is optional to elect a leader for the Consumers
|
||||
template <
|
||||
int Stages_,
|
||||
class ClusterShape_
|
||||
>
|
||||
template <int Stages_>
|
||||
class PipelineTmaAsync {
|
||||
public :
|
||||
using ClusterShape = ClusterShape_;
|
||||
using FullBarrier = cutlass::arch::ClusterTransactionBarrier;
|
||||
using EmptyBarrier = cutlass::arch::ClusterBarrier;
|
||||
using ProducerBarrierType = FullBarrier::ValueType;
|
||||
@@ -222,15 +229,15 @@ public :
|
||||
};
|
||||
|
||||
// Constructor
|
||||
template<typename ClusterShape>
|
||||
CUTLASS_DEVICE
|
||||
PipelineTmaAsync(SharedStorage& storage, Params params)
|
||||
PipelineTmaAsync(SharedStorage& storage, Params params, ClusterShape cluster_shape)
|
||||
: params_(params)
|
||||
, full_barrier_ptr_(&storage.full_barrier_[0])
|
||||
, empty_barrier_ptr_(&storage.empty_barrier_[0]) {
|
||||
|
||||
int warp_idx = canonical_warp_idx();
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
auto cluster_shape = ClusterShape{};
|
||||
|
||||
if (warp_idx == 0 && lane_predicate == 1) {
|
||||
// Barrier FULL init
|
||||
@@ -283,7 +290,8 @@ public :
|
||||
is_signalling_thread_ &= dst_blockid_ < cluster_size;
|
||||
is_signalling_thread_ &= is_same_row_or_col(dst_blockid_, block_id, cluster_shape);
|
||||
}
|
||||
|
||||
|
||||
template <typename ClusterShape>
|
||||
CUTLASS_DEVICE
|
||||
bool is_same_row_or_col(int dst_block_id, dim3 block_id, ClusterShape cluster_shape) {
|
||||
return (((dst_block_id % cute::size<0>(cluster_shape)) == block_id.x) ||
|
||||
@@ -332,7 +340,7 @@ public :
|
||||
CUTLASS_DEVICE
|
||||
void producer_tail(PipelineState state) {
|
||||
for (int count = 0; count < Stages; ++count) {
|
||||
producer_acquire(state);
|
||||
producer_acquire(state, {BarrierStatus::WaitOnly});
|
||||
++state;
|
||||
}
|
||||
}
|
||||
@@ -388,9 +396,12 @@ private :
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void producer_acquire(uint32_t stage, uint32_t phase, ProducerToken barrier_token) {
|
||||
if (barrier_token == BarrierStatus::WaitAgain) {
|
||||
if (barrier_token != BarrierStatus::WaitDone) {
|
||||
empty_barrier_ptr_[stage].wait(phase);
|
||||
}
|
||||
if (barrier_token == BarrierStatus::WaitOnly) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (params_.is_leader) {
|
||||
full_barrier_ptr_[stage].arrive_and_expect_tx(params_.transaction_bytes);
|
||||
@@ -417,7 +428,7 @@ private :
|
||||
full_barrier_ptr_[stage].complete_transaction(bytes);
|
||||
|
||||
// STEP 2 : Commit to other blocks in our cluster
|
||||
auto cluster_shape = ClusterShape{};
|
||||
auto cluster_shape = cute::cluster_shape();
|
||||
Layout block_layout_in_cluster = make_layout(cluster_shape);
|
||||
dim3 local_block_id = cute::block_id_in_cluster();
|
||||
|
||||
|
||||
@@ -394,7 +394,13 @@ public:
|
||||
/// Unpacks an element from memory
|
||||
CUTLASS_HOST_DEVICE
|
||||
Element get() const {
|
||||
Storage item = Storage((*ptr_ >> (offset_ * sizeof_bits<Element>::value)) & kMask);
|
||||
uint8_t const* byte_ptr = reinterpret_cast<uint8_t const*>(ptr_);
|
||||
// Convert offset in elements to offset in bytes
|
||||
constexpr int elements_per_byte = cutlass::sizeof_bits<uint8_t>::value / cutlass::sizeof_bits<Element>::value;
|
||||
byte_ptr += offset_ / elements_per_byte;
|
||||
// Offset of element within a byte
|
||||
int byte_offset = offset_ % elements_per_byte;
|
||||
uint8_t item = uint8_t((*byte_ptr >> (byte_offset * cutlass::sizeof_bits<Element>::value)) & kMask);
|
||||
return reinterpret_cast<Element const &>(item);
|
||||
}
|
||||
|
||||
@@ -607,6 +613,7 @@ public:
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename T> using _war = T;
|
||||
template <
|
||||
typename Element_, /// CUTLASS numeric element type.
|
||||
typename Storage_ /// Underlying basic storage type.
|
||||
@@ -647,7 +654,7 @@ private:
|
||||
StorageUnit const kMask = (StorageUnit(1) << sizeof_bits<Element>::value) - StorageUnit(1);
|
||||
|
||||
/// Pointer to array containing element
|
||||
StorageVecPointer ptr_;
|
||||
_war<StorageVecPointer> ptr_;
|
||||
|
||||
/// Offset (in units of elements) from pointer.
|
||||
///
|
||||
@@ -979,6 +986,7 @@ public:
|
||||
}
|
||||
};
|
||||
|
||||
template<typename T> using _war = T;
|
||||
template <
|
||||
typename Element_, /// CUTLASS numeric element type.
|
||||
typename Storage_ /// Underlying storage type. Must be able to hold an integer
|
||||
@@ -1019,7 +1027,7 @@ private:
|
||||
StorageUnit const kMask = (StorageUnit(1) << sizeof_bits<Element>::value) - StorageUnit(1);
|
||||
|
||||
/// Pointer to array containing element
|
||||
StorageVecPointer ptr_;
|
||||
_war<StorageVecPointer> ptr_;
|
||||
|
||||
/// Offset (in units of elements) from pointer.
|
||||
///
|
||||
@@ -1276,6 +1284,10 @@ struct ReferenceFactory;
|
||||
|
||||
template <typename Element>
|
||||
struct ReferenceFactory<Element, false> {
|
||||
|
||||
///! Number of elements per storage vector
|
||||
static int const kElementsPerVector = 1;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Element &get(Element *ptr, int64_t offset) {
|
||||
return ptr[offset];
|
||||
@@ -1285,10 +1297,25 @@ struct ReferenceFactory<Element, false> {
|
||||
static Element const &get(Element const *ptr, int64_t offset) {
|
||||
return ptr[offset];
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Element *add_pointer_offset(Element *ptr, int64_t offset) {
|
||||
return ptr + offset;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Element const *add_pointer_offset(Element const *ptr, int64_t offset) {
|
||||
return ptr + offset;
|
||||
}
|
||||
};
|
||||
|
||||
template <typename Element>
|
||||
struct ReferenceFactory<Element, true> {
|
||||
|
||||
//
|
||||
// Static methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static SubbyteReference<Element> get(Element *ptr, int64_t offset) {
|
||||
return SubbyteReference<Element>(ptr, offset);
|
||||
@@ -1299,6 +1326,22 @@ struct ReferenceFactory<Element, true> {
|
||||
int64_t offset) {
|
||||
return ConstSubbyteReference<Element>(ptr, offset);
|
||||
}
|
||||
|
||||
/// Helper to add an offset in number of elements, assuming this offset is divisible
|
||||
/// by the vector size.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Element *add_pointer_offset(Element *ptr, int64_t offset_in_elements) {
|
||||
|
||||
return ptr + offset_in_elements * sizeof_bits<Element>::value / sizeof(Element) / 8;
|
||||
}
|
||||
|
||||
/// Helper to add an offset in number of elements, assuming this offset is divisible
|
||||
/// by the vector size.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Element const *add_pointer_offset(Element const *ptr, int64_t offset_in_elements) {
|
||||
|
||||
return ptr + offset_in_elements * sizeof_bits<Element>::value / sizeof(Element) / 8;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -103,14 +103,14 @@ public:
|
||||
using SmemLayoutB = SmemLayoutB_;
|
||||
using SmemLayoutAtomB = SmemLayoutAtomB_;
|
||||
using ElementB = ElementB_;
|
||||
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
NoTranspositionOperandB(
|
||||
int,
|
||||
int,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
int,
|
||||
int,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB) { }
|
||||
|
||||
template <
|
||||
@@ -148,12 +148,12 @@ public:
|
||||
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
UniversalTranspositionOperandB(
|
||||
int warp_idx_,
|
||||
int warp_group_thread_idx_,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB)
|
||||
int warp_idx_,
|
||||
int warp_group_thread_idx_,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB)
|
||||
: warp_idx(warp_idx_)
|
||||
, warp_group_thread_idx(warp_group_thread_idx_) { }
|
||||
|
||||
@@ -168,9 +168,9 @@ public:
|
||||
return;
|
||||
}
|
||||
|
||||
constexpr int NumMathWarpGroup = size(TiledMma{}) / NumThreadsPerWarpGroup;
|
||||
static_assert(NumMathWarpGroup == 1 ||
|
||||
(!detail::use_universal_transposition<SmemLayoutAtomB, ElementB>() && NumMathWarpGroup == 2),
|
||||
constexpr int NumMathWarpGroup = CUTE_STATIC_V(size(TiledMma{})) / NumThreadsPerWarpGroup;
|
||||
static_assert(NumMathWarpGroup == 1 ||
|
||||
(!detail::use_universal_transposition<SmemLayoutAtomB, ElementB>() && NumMathWarpGroup == 2),
|
||||
"Wrong math warp group number for TransposeB");
|
||||
constexpr int WarpgroupTileSize = size<1>(SmemLayoutB{}); // A warp group tile would process entire Smem K.
|
||||
|
||||
@@ -234,14 +234,14 @@ public:
|
||||
if (step == 0) {
|
||||
// SMEM fence to make sure B is transposed before math
|
||||
cutlass::arch::fence_view_async_shared();
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 1);
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), cutlass::arch::ReservedNamedBarriers::TransposeBarrier);
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void synchronize() {
|
||||
// SMEM fence to make sure B is transposed before math
|
||||
cutlass::arch::fence_view_async_shared();
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 1);
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), cutlass::arch::ReservedNamedBarriers::TransposeBarrier);
|
||||
}
|
||||
|
||||
template <
|
||||
@@ -251,7 +251,7 @@ public:
|
||||
TensorSmemB const& sB,
|
||||
TensorTransposedSmemB const& gmma_sB,
|
||||
int read_stage) {
|
||||
|
||||
|
||||
this->operator()(sB, gmma_sB, read_stage, 0);
|
||||
synchronize();
|
||||
|
||||
@@ -267,7 +267,7 @@ template<
|
||||
class SmemLayoutB_,
|
||||
class SmemLayoutAtomB_,
|
||||
class ElementB_>
|
||||
class AsyncTranspositionOperandB {
|
||||
class AsyncTranspositionOperandB {
|
||||
public:
|
||||
|
||||
using TiledMma = TiledMma_;
|
||||
@@ -276,9 +276,9 @@ public:
|
||||
using ElementB = ElementB_;
|
||||
|
||||
static constexpr int Steps = 2;
|
||||
static constexpr int NumMathWarpGroup = size(TiledMma{}) / NumThreadsPerWarpGroup;
|
||||
static constexpr int NumMathWarpGroup = CUTE_STATIC_V(size(TiledMma{})) / NumThreadsPerWarpGroup;
|
||||
static constexpr int StepsPerWarpGroup = Steps / NumMathWarpGroup;
|
||||
static_assert(NumMathWarpGroup <= 2,
|
||||
static_assert(NumMathWarpGroup <= 2,
|
||||
"Wrong math warp group number for TransposeB");
|
||||
static constexpr int WarpgroupTileSize = size<1>(SmemLayoutB{}); // A warp group tile would process entire Smem K.
|
||||
static constexpr int NumWarpsPerWarpGroup = NumThreadsPerWarpGroup / NumThreadsPerWarp;
|
||||
@@ -303,23 +303,23 @@ public:
|
||||
"Copy size must evenly divide SMEM tile.");
|
||||
static constexpr int WarpgroupTileNum = size<0>(SmemLayoutB{}) / WarpgroupTileSize;
|
||||
|
||||
static_assert(size<2>(typename TiledMma::AtomShape_MNK{}) <= WarpThreadShapeK,
|
||||
static_assert(size<2>(typename TiledMma::AtomShape_MNK{}) <= WarpThreadShapeK,
|
||||
"Need to be able to transpose first k-block in the first step");
|
||||
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
AsyncTranspositionOperandB(
|
||||
int warp_idx_,
|
||||
int warp_group_thread_idx_,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB)
|
||||
int warp_idx_,
|
||||
int warp_group_thread_idx_,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB)
|
||||
: warp_idx(warp_idx_)
|
||||
, warp_group_thread_idx(warp_group_thread_idx_)
|
||||
, warp_idx_in_warp_group(warp_idx_ % NumWarpsPerWarpGroup)
|
||||
, current_warp_tile_n_coord_LUT((WarpTileNCoordLUT >> ((warp_idx_
|
||||
, current_warp_tile_n_coord_LUT((WarpTileNCoordLUT >> ((warp_idx_
|
||||
% NumWarpsPerWarpGroup) * NumBitsPerWarp)) & MaskPerWarp)
|
||||
, current_warp_tile_k_coord_LUT((WarpTileKCoordLUT >> ((warp_idx_
|
||||
, current_warp_tile_k_coord_LUT((WarpTileKCoordLUT >> ((warp_idx_
|
||||
% NumWarpsPerWarpGroup) * NumBitsPerWarp)) & MaskPerWarp) { }
|
||||
|
||||
template <
|
||||
@@ -328,12 +328,12 @@ public:
|
||||
CUTLASS_DEVICE void operator()(
|
||||
TensorSmemB const& sB,
|
||||
TensorTransposedSmemB const& gmma_sB,
|
||||
int read_stage, int current_step)
|
||||
int read_stage, int current_step)
|
||||
{
|
||||
if (current_step >= StepsPerWarpGroup) {
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
static constexpr auto WarpThreadLayout = make_layout(make_shape(Int<WarpThreadShapeN>{}, Int<WarpThreadShapeK>{}));
|
||||
//////////////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// A warp group uses 2 steps to transpose the whole WarpgroupTileSize x WarpgroupTileSize.
|
||||
@@ -384,11 +384,11 @@ public:
|
||||
};
|
||||
|
||||
[[maybe_unused]] int step = current_step * NumMathWarpGroup;
|
||||
if constexpr (NumMathWarpGroup == 2) {
|
||||
// For 2 math warpgroup, warp idx4~7 is 1st warp group and 8~9 is 2nd, so decide if 2nd warpgroup need warp idx divide 8.
|
||||
if constexpr (NumMathWarpGroup == 2) {
|
||||
// For 2 math warpgroup, warp idx4~7 is 1st warp group and 8~9 is 2nd, so decide if 2nd warpgroup need warp idx divide 8.
|
||||
step += warp_idx / (NumWarpsPerWarpGroup * 2);
|
||||
}
|
||||
|
||||
|
||||
int tmp_warp_tile_n_coord_LUT = current_warp_tile_n_coord_LUT >> (NumBitsPerStep * current_step);
|
||||
int tmp_warp_tile_k_coord_LUT = current_warp_tile_k_coord_LUT >> (NumBitsPerStep * current_step);
|
||||
|
||||
@@ -396,7 +396,7 @@ public:
|
||||
tmp_warp_tile_n_coord_LUT >>= NumBitsPerStep * (warp_idx / (NumWarpsPerWarpGroup * 2));
|
||||
tmp_warp_tile_k_coord_LUT >>= NumBitsPerStep * (warp_idx / (NumWarpsPerWarpGroup * 2));
|
||||
}
|
||||
|
||||
|
||||
// decoding the warp tile coord.
|
||||
int warp_tile0_n, warp_tile0_k;
|
||||
if constexpr (StepsPerWarpGroup <= NumStepsEncoded) {
|
||||
@@ -412,7 +412,7 @@ public:
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int warp_group_tile = 0; warp_group_tile < WarpgroupTileNum; ++warp_group_tile) {
|
||||
|
||||
|
||||
static_assert(TilesPerWarp == 2);
|
||||
|
||||
// [warp_tile][n/k]
|
||||
@@ -427,7 +427,7 @@ public:
|
||||
Tensor tCsB = sB_thr_copy.partition_S(
|
||||
flatten(s_tile(_, make_coord(warp_tile_coord[warp_tile][0], warp_tile_coord[warp_tile][1])))
|
||||
); // (CPY, CPY_N, CPY_K)
|
||||
|
||||
|
||||
copy(sB_tiled_copy, tCsB, transpose_fragments[warp_tile]);
|
||||
}
|
||||
|
||||
@@ -442,20 +442,20 @@ public:
|
||||
copy(sB_tiled_copy, transpose_fragments[warp_tile], tCsB_transposed);
|
||||
}
|
||||
|
||||
} // loop warp_group_tile
|
||||
} // loop warp_group_tile
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void synchronize(int step) {
|
||||
if (step < StepsPerWarpGroup) {
|
||||
// SMEM fence to make sure B is transposed before math
|
||||
cutlass::arch::fence_view_async_shared();
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 1);
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), cutlass::arch::ReservedNamedBarriers::TransposeBarrier);
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void synchronize() {
|
||||
cutlass::arch::fence_view_async_shared();
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 1);
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), cutlass::arch::ReservedNamedBarriers::TransposeBarrier);
|
||||
}
|
||||
|
||||
template <
|
||||
@@ -465,7 +465,7 @@ public:
|
||||
TensorSmemB const& sB,
|
||||
TensorTransposedSmemB const& gmma_sB,
|
||||
int read_stage) {
|
||||
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < StepsPerWarpGroup; ++i) {
|
||||
this->operator()(sB, gmma_sB, read_stage, i);
|
||||
@@ -486,7 +486,7 @@ template<
|
||||
class SmemLayoutB_,
|
||||
class SmemLayoutAtomB_,
|
||||
class ElementB_>
|
||||
class AsyncTranspositionOperandB_1BElementB {
|
||||
class AsyncTranspositionOperandB_1BElementB {
|
||||
public:
|
||||
|
||||
static_assert(sizeof(ElementB_) == 1);
|
||||
@@ -495,11 +495,11 @@ public:
|
||||
using SmemLayoutB = SmemLayoutB_;
|
||||
using SmemLayoutAtomB = SmemLayoutAtomB_;
|
||||
using ElementB = ElementB_;
|
||||
|
||||
|
||||
static constexpr int Steps = 8;
|
||||
static constexpr int NumMathWarpGroup = size(TiledMma{}) / NumThreadsPerWarpGroup;
|
||||
static constexpr int NumMathWarpGroup = CUTE_STATIC_V(size(TiledMma{})) / NumThreadsPerWarpGroup;
|
||||
static constexpr int StepsPerWarpGroup = Steps / NumMathWarpGroup;
|
||||
static_assert(NumMathWarpGroup <= 2,
|
||||
static_assert(NumMathWarpGroup <= 2,
|
||||
"Wrong math warp group number for TransposeB");
|
||||
static constexpr int WarpgroupTileSize = size<1>(SmemLayoutB{}); // A warp group tile would process entire Smem K.
|
||||
static constexpr int NumWarpsPerWarpGroup = NumThreadsPerWarpGroup / NumThreadsPerWarp;
|
||||
@@ -524,21 +524,20 @@ public:
|
||||
"Copy size must evenly divide SMEM tile.");
|
||||
static constexpr int WarpgroupTileNum = size<0>(SmemLayoutB{}) / WarpgroupTileSize;
|
||||
|
||||
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
AsyncTranspositionOperandB_1BElementB(
|
||||
int warp_idx_,
|
||||
int warp_group_thread_idx_,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB)
|
||||
int warp_idx_,
|
||||
int warp_group_thread_idx_,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB)
|
||||
: warp_idx(warp_idx_)
|
||||
, warp_group_thread_idx(warp_group_thread_idx_)
|
||||
, warp_idx_in_warp_group(warp_idx_ % NumWarpsPerWarpGroup)
|
||||
, current_warp_tile_n_coord_LUT((WarpTileNCoordLUT >> ((warp_idx_
|
||||
, current_warp_tile_n_coord_LUT((WarpTileNCoordLUT >> ((warp_idx_
|
||||
% NumWarpsPerWarpGroup) * NumBitsPerWarp)) & MaskPerWarp)
|
||||
, current_warp_tile_k_coord_LUT((WarpTileKCoordLUT >> ((warp_idx_
|
||||
, current_warp_tile_k_coord_LUT((WarpTileKCoordLUT >> ((warp_idx_
|
||||
% NumWarpsPerWarpGroup) * NumBitsPerWarp)) & MaskPerWarp) { }
|
||||
|
||||
template <
|
||||
@@ -547,7 +546,7 @@ public:
|
||||
CUTLASS_DEVICE void operator()(
|
||||
TensorSmemB const& sB,
|
||||
TensorTransposedSmemB const& gmma_sB,
|
||||
int read_stage, int current_step)
|
||||
int read_stage, int current_step)
|
||||
{
|
||||
if (current_step > 0) {
|
||||
return;
|
||||
@@ -628,7 +627,7 @@ public:
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for (int step_per_warp_group = 0; step_per_warp_group < StepsPerWarpGroup; ++step_per_warp_group) {
|
||||
// For 2 math warpgroup, warp idx4~7 is 1st warp group and 8~9 is 2nd, so decide if 2nd warpgroup need warp idx divide 8.
|
||||
// For 2 math warpgroup, warp idx4~7 is 1st warp group and 8~9 is 2nd, so decide if 2nd warpgroup need warp idx divide 8.
|
||||
int step = step_per_warp_group * NumMathWarpGroup + warp_idx / (NumWarpsPerWarpGroup * 2);
|
||||
// decoding the warp tile coord.
|
||||
int warp_tile0_n = step < NumStepsEncoded ? (tmp_warp_tile_n_coord_LUT & MaskPerStep) : 4 + warp_idx_in_warp_group;
|
||||
@@ -653,7 +652,7 @@ public:
|
||||
Tensor tCsB = sB_thr_copy.partition_S(
|
||||
flatten(s_tile(_, make_coord(warp_tile_coord[warp_tile][0], warp_tile_coord[warp_tile][1])))
|
||||
); // (CPY, CPY_N, CPY_K)
|
||||
|
||||
|
||||
copy(sB_tiled_copy, tCsB, transpose_fragments[warp_tile]);
|
||||
}
|
||||
|
||||
@@ -675,13 +674,13 @@ public:
|
||||
if (step == 0) {
|
||||
// SMEM fence to make sure B is transposed before math
|
||||
cutlass::arch::fence_view_async_shared();
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 1);
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), cutlass::arch::ReservedNamedBarriers::TransposeBarrier);
|
||||
}
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void synchronize() {
|
||||
cutlass::arch::fence_view_async_shared();
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), 1);
|
||||
cutlass::arch::NamedBarrier::sync(size(TiledMma{}), cutlass::arch::ReservedNamedBarriers::TransposeBarrier);
|
||||
}
|
||||
|
||||
template <
|
||||
@@ -711,35 +710,35 @@ template<
|
||||
class ElementB,
|
||||
bool TransposeB
|
||||
>
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
auto
|
||||
constexpr CUTLASS_HOST_DEVICE
|
||||
auto
|
||||
make_transpose_operand_b(
|
||||
int warp_idx,
|
||||
int warp_group_thread_idx,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
int warp_idx,
|
||||
int warp_group_thread_idx,
|
||||
TiledMma,
|
||||
SmemLayoutB,
|
||||
SmemLayoutAtomB,
|
||||
ElementB,
|
||||
cute::bool_constant<TransposeB>)
|
||||
{
|
||||
if constexpr (!TransposeB) {
|
||||
return NoTranspositionOperandB(
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
SmemLayoutB{}, SmemLayoutAtomB{}, ElementB{});
|
||||
}
|
||||
else if constexpr (use_universal_transposition<SmemLayoutAtomB, ElementB>()) {
|
||||
return UniversalTranspositionOperandB(
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
SmemLayoutB{}, SmemLayoutAtomB{}, ElementB{});
|
||||
}
|
||||
else if constexpr (sizeof(ElementB) == 1) {
|
||||
return AsyncTranspositionOperandB_1BElementB(
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
SmemLayoutB{}, SmemLayoutAtomB{}, ElementB{});
|
||||
}
|
||||
else {
|
||||
return AsyncTranspositionOperandB(
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
warp_idx, warp_group_thread_idx, TiledMma{},
|
||||
SmemLayoutB{}, SmemLayoutAtomB{}, ElementB{});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -71,7 +71,6 @@ struct PredicatedTileAccessIteratorDesc {
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIteratorDesc() = default;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
@@ -279,7 +278,6 @@ struct PredicatedTileAccessIteratorParams {
|
||||
return initialize(LongIndex(stride), desc);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
PredicatedTileAccessIteratorParams() = default;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
|
||||
@@ -56,9 +56,9 @@
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
/// Optionally enable GCC's built-in type
|
||||
#if (defined(__x86_64) || defined (__aarch64__)) && !defined(__CUDA_ARCH__) && defined(__GNUC__)
|
||||
#if (defined(__x86_64) || defined (__aarch64__)) && !(defined(__CUDA_ARCH__) && ((__CUDACC_VER_MAJOR__ <= 10) || ((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ <= 4)))) && defined(__GNUC__)
|
||||
#define CUTLASS_UINT128_NATIVE
|
||||
#elif defined(_MSC_VER) && defined(_M_AMD64) && !defined(__CUDA_ARCH__)
|
||||
#elif defined(_MSC_VER) && defined(_M_AMD64) && !(defined(__CUDA_ARCH__) && ((__CUDACC_VER_MAJOR__ <= 10) || ((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ <= 4))))
|
||||
#define CUTLASS_INT128_ARITHMETIC
|
||||
#include <intrin.h>
|
||||
#if _MSC_VER >= 1920
|
||||
|
||||
Reference in New Issue
Block a user