@@ -0,0 +1,671 @@
|
||||
/***************************************************************************************************
|
||||
* 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 <type_traits>
|
||||
|
||||
#include <cute/config.hpp>
|
||||
|
||||
#include <cute/arch/copy.hpp>
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
namespace cute {
|
||||
|
||||
// Generic copy_unpack for any Copy_Traits
|
||||
template <class Operation, class... Args,
|
||||
class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
copy_unpack(Copy_Traits<Operation, Args...> const&,
|
||||
Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst)
|
||||
{
|
||||
// Specializations can generalize on these checks
|
||||
//static_assert(is_smem<TS>::value, "Expected smem for this Copy_Traits<Operation>");
|
||||
//static_assert(is_rmem<TD>::value, "Expected rmem for this Copy_Traits<Operation>");
|
||||
|
||||
using RegistersSrc = typename Operation::SRegisters;
|
||||
using RegistersDst = typename Operation::DRegisters;
|
||||
using RegTypeSrc = typename std::remove_extent<RegistersSrc>::type;
|
||||
using RegTypeDst = typename std::remove_extent<RegistersDst>::type;
|
||||
constexpr int RegNumSrc = std::extent<RegistersSrc>::value;
|
||||
constexpr int RegNumDst = std::extent<RegistersDst>::value;
|
||||
|
||||
Tensor rS = recast<RegTypeSrc>(src);
|
||||
Tensor rD = recast<RegTypeDst>(dst);
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size(rS) == Int<RegNumSrc>{},
|
||||
"In CopyAtom, src layout doesn't vectorize into registers. This src layout is incompatible with this tiled copy.");
|
||||
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumDst>{},
|
||||
"In CopyAtom, dst layout doesn't vectorize into registers. This dst layout is incompatible with this tiled copy.");
|
||||
|
||||
detail::explode(Operation::copy,
|
||||
rS, make_int_sequence<RegNumSrc>{},
|
||||
rD, make_int_sequence<RegNumDst>{});
|
||||
}
|
||||
|
||||
|
||||
template <class... Args>
|
||||
struct Copy_Atom;
|
||||
|
||||
template <class CopyOperation, class T>
|
||||
struct Copy_Atom<CopyOperation, T> : Copy_Atom<Copy_Traits<CopyOperation>, T>
|
||||
{};
|
||||
|
||||
template <class... Args, class T>
|
||||
struct Copy_Atom<Copy_Traits<Args...>, T>
|
||||
: Copy_Traits<Args...>
|
||||
{
|
||||
using Traits = Copy_Traits<Args...>;
|
||||
|
||||
// Bit and Thr layouts from the Copy_Traits
|
||||
using ThrID = typename Traits::ThrID;
|
||||
using BitLayoutSrc = typename Traits::SrcLayout;
|
||||
using BitLayoutDst = typename Traits::DstLayout;
|
||||
using BitLayoutRef = typename Traits::RefLayout;
|
||||
|
||||
using ValType = T;
|
||||
|
||||
using ValLayoutSrc = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutSrc{}));
|
||||
using ValLayoutDst = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutDst{}));
|
||||
using ValLayoutRef = decltype(upcast<sizeof_bits<ValType>::value>(BitLayoutRef{}));
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size<0>(ValLayoutSrc{}) == size(ThrID{}), "CopyOperation is not valid for Src of ValType.");
|
||||
CUTE_STATIC_ASSERT_V(size<0>(ValLayoutDst{}) == size(ThrID{}), "CopyOperation is not valid for Dst of ValType.");
|
||||
CUTE_STATIC_ASSERT_V(size<0>(ValLayoutRef{}) == size(ThrID{}), "CopyOperation is not valid for Ref of ValType.");
|
||||
|
||||
static constexpr int NumValSrc = size<1>(ValLayoutSrc{});
|
||||
static constexpr int NumValDst = size<1>(ValLayoutDst{});
|
||||
|
||||
// Additional Trait parameters/transformations
|
||||
template <class... TraitsArgs>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
with(TraitsArgs&&... args) const {
|
||||
auto traits = Traits::with(std::forward<TraitsArgs>(args)...);
|
||||
return Copy_Atom<decltype(traits), T>{traits};
|
||||
}
|
||||
|
||||
// Print thread and data layouts for debugging
|
||||
CUTE_HOST_DEVICE static
|
||||
void
|
||||
print_all()
|
||||
{
|
||||
print("ThrID: "); print(ThrID{}); print("\n");
|
||||
print("BitLayoutSrc: "); print(BitLayoutSrc{}); print("\n");
|
||||
print("BitLayoutDst: "); print(BitLayoutDst{}); print("\n");
|
||||
print("BitLayoutRef: "); print(BitLayoutRef{}); print("\n");
|
||||
print("ValLayoutSrc: "); print(ValLayoutSrc{}); print("\n");
|
||||
print("ValLayoutDst: "); print(ValLayoutDst{}); print("\n");
|
||||
print("ValLayoutRef: "); print(ValLayoutRef{}); print("\n");
|
||||
print("ValueType: %db", sizeof_bits<ValType>::value); print("\n");
|
||||
}
|
||||
|
||||
//
|
||||
// Tensor call interfaces
|
||||
//
|
||||
|
||||
// Cast, check, and call
|
||||
template <class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
call(Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,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
|
||||
return copy_unpack(*this, src, dst);
|
||||
} else {
|
||||
// Recurse if needed by peeling the tensor mode
|
||||
return copy(*this, tensor<0>(src), tensor<0>(dst));
|
||||
}
|
||||
}
|
||||
|
||||
// Accept mutable temporaries
|
||||
template <class SEngine, class SLayout,
|
||||
class DEngine, class DLayout>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
call(Tensor<SEngine,SLayout> const& src,
|
||||
Tensor<DEngine,DLayout> && dst) const
|
||||
{
|
||||
return call(src, dst);
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// A tiling of copy atoms
|
||||
//
|
||||
|
||||
template <class Copy_Atom,
|
||||
class LayoutCopy_TV, // (tid,vid) -> coord [Need not be 2D...]
|
||||
class ShapeTile_MN> // coord space
|
||||
struct TiledCopy : Copy_Atom
|
||||
{
|
||||
// Layout information from the CopyAtom
|
||||
using AtomThrID = typename Copy_Atom::ThrID; // thrid -> thr_idx
|
||||
using AtomLayoutSrc = typename Copy_Atom::ValLayoutSrc; // (thr,val) -> offset
|
||||
using AtomLayoutDst = typename Copy_Atom::ValLayoutDst; // (thr,val) -> offset
|
||||
using AtomLayoutRef = typename Copy_Atom::ValLayoutRef; // (thr,val) -> offset
|
||||
|
||||
using AtomNumThr = decltype(size<0>(AtomLayoutRef{}));
|
||||
using AtomNumVal = decltype(size<1>(AtomLayoutRef{}));
|
||||
|
||||
// Layout information for the TiledCopy
|
||||
using Tiler_MN = ShapeTile_MN;
|
||||
using TiledShape_MN = decltype(shape(ShapeTile_MN{}));
|
||||
using TiledLayout_TV = LayoutCopy_TV;
|
||||
using TiledNumThr = decltype(size<0>(TiledLayout_TV{}));
|
||||
using TiledNumVal = decltype(size<1>(TiledLayout_TV{}));
|
||||
|
||||
CUTE_STATIC_ASSERT_V(TiledNumThr{} % AtomNumThr{} == Int<0>{}, "TiledCopy uses too few thrs for selected CopyAtom");
|
||||
CUTE_STATIC_ASSERT_V(TiledNumVal{} % AtomNumVal{} == Int<0>{}, "TiledCopy uses too few vals for selected CopyAtom");
|
||||
|
||||
// Tile a tensor or a layout from shape
|
||||
// (M,N,...)
|
||||
// to shape
|
||||
// ((ThrV,ThrX),FrgV,(RestM,RestN,...))
|
||||
// where
|
||||
// ThrV: The threads local to a COPY_ATOM Src.
|
||||
// ThrX: The threads tiled across COPY_ATOMs Src.
|
||||
// FrgV: The values local to a COPY_ATOM Src.
|
||||
// RestM: The values tiled in M.
|
||||
// RestN: The values tiled in N.
|
||||
template <class STensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
tidfrg_S(STensor&& stensor)
|
||||
{
|
||||
return thrfrg(stensor, right_inverse(AtomLayoutRef{}).compose(AtomLayoutSrc{}));
|
||||
}
|
||||
|
||||
// Tile a tensor or a layout from shape
|
||||
// (M,N,...)
|
||||
// to shape
|
||||
// ((ThrV,ThrX),FrgV,(RestM,RestN,...))
|
||||
// where
|
||||
// ThrV: The threads local to a COPY_ATOM Dst.
|
||||
// ThrX: The threads tiled across COPY_ATOMs Dst.
|
||||
// FrgV: The values local to a COPY_ATOM Dst.
|
||||
// RestM: The values tiled in M.
|
||||
// RestN: The values tiled in N.
|
||||
template <class DTensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
tidfrg_D(DTensor&& dtensor)
|
||||
{
|
||||
return thrfrg(dtensor, right_inverse(AtomLayoutRef{}).compose(AtomLayoutDst{}));
|
||||
}
|
||||
|
||||
template <class Tensor, class Ref2TrgLayout>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
thrfrg(Tensor&& tensor, Ref2TrgLayout const& ref2trg)
|
||||
{
|
||||
constexpr int R = remove_cvref_t<Tensor>::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>{});
|
||||
|
||||
// Take the thrs/vals that the atom is interested in
|
||||
// NOTE: Assumes the AtomNumThr are contiguous and identity within TiledThrID
|
||||
auto atom_layout_TV = zipped_divide(TiledLayout_TV{}, make_shape(AtomNumThr{}, AtomNumVal{}));
|
||||
// ((atom_tid,atom_val),(rest_tid,rest_val)) -> (m,n)
|
||||
|
||||
// Transform to the trg layout
|
||||
auto trg_layout_TV = atom_layout_TV.compose(ref2trg, _);
|
||||
// ((trg_tid,trg_val),(rest_tid,rest_val)) -> (m,n)
|
||||
|
||||
// Transform the thrs mode from thrid to thr_idx
|
||||
// NOTE: Assumes the AtomNumThr are contiguous and identity within TiledThrID
|
||||
auto thrval2mn = coalesce(zip(trg_layout_TV), Shape<_1,Shape<_1,_1>>{});
|
||||
// ((trg_tid,rest_tid),(trg_val,rest_val)) -> (m,n)
|
||||
|
||||
/// ==================
|
||||
|
||||
// Tile the tensor for TiledLayout
|
||||
auto t_tensor = zipped_divide(tensor, Tiler_MN{});
|
||||
// ((TileM,TileN,...),(RestM,RestN,...))
|
||||
|
||||
// Transform the tile mode
|
||||
auto tv_tensor = t_tensor.compose(thrval2mn, _);
|
||||
// ((thrid,val),(RM,RN,...))
|
||||
|
||||
// Unfold and return
|
||||
return tv_tensor(make_coord(_,_), _);
|
||||
}
|
||||
|
||||
// retile_S and retile_D assume they are working with the reference layout -- they are the same
|
||||
template <class Tensor>
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
retile(Tensor&& tensor)
|
||||
{
|
||||
constexpr int R = remove_cvref_t<Tensor>::rank;
|
||||
// Assert that AtomLayoutSrc|Dst is identity so we can skip the Ref transformation
|
||||
|
||||
// Assume the first size<0>(tensor) elements are the first val_ids in TiledLayout_TV.
|
||||
// Then, we only need the shape+layout of those size<0>(tensor) elements in TiledLayout_TV
|
||||
// and that shape is what we gather from the other modes of tensor
|
||||
|
||||
auto V = size<0>(tensor);
|
||||
|
||||
auto frg_layout_mn = upcast<TiledNumThr{} * V>(right_inverse(TiledLayout_TV{}).with_shape(TiledShape_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{}));
|
||||
// (atom_vals,rest_vals) -> (v,m,n)
|
||||
|
||||
/// =======
|
||||
|
||||
// Tile the tensor for TileFrg
|
||||
auto t_tensor = zipped_divide(tensor, prepend(product_each(shape(frg_layout_mn)), V));
|
||||
// ((TileV,TileM,TileN,...),(1,RestM,RestN,...))
|
||||
|
||||
// Transform the tile mode
|
||||
auto v_tensor = t_tensor.compose(frg_layout_v, _);
|
||||
// ((atom_vals,rest_vals),(1,RM,RN,...))
|
||||
|
||||
// Unfold and return
|
||||
return v_tensor(_, append<R>(Int<0>{},_));
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
get_layoutS_MN()
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_S = make_layout(TiledShape_MN{});
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
auto layoutS_TV = tidfrg_S(ref_S);
|
||||
// (M,K) -> (thr_idx,val_idx)
|
||||
auto layoutS_MK = right_inverse(layoutS_TV).with_shape(shape(ref_S));
|
||||
|
||||
// athrid = (v,m,k) -> thr_idx
|
||||
auto thrID_S = make_layout(size<0>(TiledLayout_TV{}));
|
||||
|
||||
return cute::make_tuple(layoutS_MK, thrID_S);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
get_layoutS_TV()
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_S = make_layout(TiledShape_MN{});
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
return tidfrg_S(ref_S)(_,_,Int<0>{});
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
get_layoutD_MN()
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_D = make_layout(TiledShape_MN{});
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
auto layoutD_TV = tidfrg_D(ref_D);
|
||||
// (M,K) -> (thr_idx,val_idx)
|
||||
auto layoutD_MK = right_inverse(layoutD_TV).with_shape(shape(ref_D));
|
||||
|
||||
// athrid = (v,m,k) -> thr_idx
|
||||
auto thrID_D = make_layout(size<0>(TiledLayout_TV{}));
|
||||
|
||||
return cute::make_tuple(layoutD_MK, thrID_D);
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr static
|
||||
auto
|
||||
get_layoutD_TV()
|
||||
{
|
||||
// (M,N) -> (M,N)
|
||||
auto ref_D = make_layout(TiledShape_MN{});
|
||||
// (thr_idx,val_idx) -> (M,N)
|
||||
return tidfrg_D(ref_D)(_,_,Int<0>{});
|
||||
}
|
||||
|
||||
template <class ThrIdx>
|
||||
struct ThrCopy : Copy_Atom
|
||||
{
|
||||
ThrIdx thr_idx_;
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
ThrCopy(ThrIdx const& thr_idx) : thr_idx_(thr_idx) {}
|
||||
|
||||
template <class STensor>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
partition_S(STensor&& stensor) {
|
||||
//static_assert(sizeof(typename remove_cvref_t<STensor>::value_type) == sizeof(typename Copy_Atom::ValType),
|
||||
// "Expected ValType for tiling SrcTensor.");
|
||||
auto thr_tensor = make_tensor(std::forward<STensor>(stensor).data(), tidfrg_S(stensor.layout()));
|
||||
return thr_tensor(thr_idx_, _, repeat<rank_v<STensor>>(_));
|
||||
}
|
||||
|
||||
template <class DTensor>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
partition_D(DTensor&& dtensor) {
|
||||
//static_assert(sizeof(typename remove_cvref_t<DTensor>::value_type) == sizeof(typename Copy_Atom::ValType),
|
||||
// "Expected ValType for tiling DstTensor.");
|
||||
auto thr_tensor = make_tensor(std::forward<DTensor>(dtensor).data(), tidfrg_D(dtensor.layout()));
|
||||
return thr_tensor(thr_idx_, _, repeat<rank_v<DTensor>>(_));
|
||||
}
|
||||
|
||||
template <class STensor>
|
||||
CUTE_HOST_DEVICE static
|
||||
auto
|
||||
retile_S(STensor&& stensor) {
|
||||
static_assert(sizeof(typename remove_cvref_t<STensor>::value_type) == sizeof(typename Copy_Atom::ValType),
|
||||
"Expected ValType for tiling SrcTensor.");
|
||||
return make_tensor(std::forward<STensor>(stensor).data(), TiledCopy::retile(stensor.layout()));
|
||||
}
|
||||
|
||||
template <class DTensor>
|
||||
CUTE_HOST_DEVICE static
|
||||
auto
|
||||
retile_D(DTensor&& dtensor) {
|
||||
static_assert(sizeof(typename remove_cvref_t<DTensor>::value_type) == sizeof(typename Copy_Atom::ValType),
|
||||
"Expected ValType for tiling DstTensor.");
|
||||
return make_tensor(std::forward<DTensor>(dtensor).data(), TiledCopy::retile(dtensor.layout()));
|
||||
}
|
||||
};
|
||||
|
||||
template <class ThrIdx,
|
||||
__CUTE_REQUIRES(is_integral<ThrIdx>::value)>
|
||||
CUTE_HOST_DEVICE static
|
||||
auto
|
||||
get_slice(ThrIdx const& thr_idx)
|
||||
{
|
||||
return ThrCopy<ThrIdx>(thr_idx);
|
||||
}
|
||||
|
||||
template <class ThrIdx,
|
||||
__CUTE_REQUIRES(is_integral<ThrIdx>::value)>
|
||||
CUTE_HOST_DEVICE static
|
||||
auto
|
||||
get_thread_slice(ThrIdx const& thr_idx)
|
||||
{
|
||||
return get_slice(thr_idx);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
template <class... Args,
|
||||
class LayoutCopy_TV,
|
||||
class... TLayout>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_impl(Copy_Atom<Args...> const& atom,
|
||||
LayoutCopy_TV const&,
|
||||
Tile<TLayout...> const&)
|
||||
{
|
||||
return TiledCopy<Copy_Atom<Args...>, LayoutCopy_TV, Tile<TLayout...>>{atom};
|
||||
}
|
||||
|
||||
//
|
||||
// These tile the Copy_Atom as a whole
|
||||
//
|
||||
|
||||
template <class... Args,
|
||||
class TiledMMA>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_A(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledMMA const& tiled_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{})));
|
||||
}
|
||||
|
||||
template <class... Args,
|
||||
class TiledMMA>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_B(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledMMA const& tiled_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{})));
|
||||
}
|
||||
|
||||
template <class... Args,
|
||||
class TiledMMA>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_C(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledMMA const& tiled_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{})));
|
||||
}
|
||||
|
||||
template <class... Args,
|
||||
class ThrLayout,
|
||||
class ValLayout = Layout<_1>>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy(Copy_Atom<Args...> const& copy_atom,
|
||||
ThrLayout const& thr_layout = {}, // (m,n) -> thr_idx
|
||||
ValLayout const& val_layout = {})
|
||||
{
|
||||
constexpr int R = cute::max(rank_v<ThrLayout>, rank_v<ValLayout>);
|
||||
|
||||
auto thr_layout_mn = append<R>(thr_layout, Layout<_1>{});
|
||||
auto val_layout_mn = append<R>(val_layout, Layout<_1>{});
|
||||
|
||||
// Take the raked_products to compute the Layout_MN
|
||||
auto layout_mn = raked_product(thr_layout_mn, val_layout_mn);
|
||||
auto layout_tv = right_inverse(layout_mn).with_shape(make_shape(size(thr_layout), size(val_layout)));
|
||||
|
||||
//print("thr_layout: "); print(thr_layout_mn); print("\n");
|
||||
//print("val_layout: "); print(val_layout_mn); print("\n");
|
||||
//print("layout_mn : "); print(layout_mn); print("\n");
|
||||
//print("layout_tv : "); print(layout_tv); print("\n");
|
||||
|
||||
return make_tiled_copy_impl(copy_atom, layout_tv, product_each(shape(layout_mn)));
|
||||
}
|
||||
|
||||
// Make a TiledCopy out of the copy_atom that matches the Src-Layout of tiled_copy
|
||||
template <class... Args,
|
||||
class TiledCopy>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_S(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledCopy const& tiled_copy)
|
||||
{
|
||||
return make_tiled_copy_impl(copy_atom, tiled_copy.get_layoutS_TV(), typename TiledCopy::Tiler_MN{});
|
||||
}
|
||||
|
||||
// Make a TiledCopy out of the copy_atom that matches the Dst-Layout of tiled_copy
|
||||
template <class... Args,
|
||||
class TiledCopy>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
make_tiled_copy_D(Copy_Atom<Args...> const& copy_atom,
|
||||
TiledCopy const& tiled_copy)
|
||||
{
|
||||
return make_tiled_copy_impl(copy_atom, tiled_copy.get_layoutD_TV(), typename TiledCopy::Tiler_MN{});
|
||||
}
|
||||
|
||||
//
|
||||
// Size
|
||||
//
|
||||
|
||||
// The logical size of a TileCopy
|
||||
template <int... I, class... Args>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
tile_size(TiledCopy<Args...> const&)
|
||||
{
|
||||
return size<I...>(typename TiledCopy<Args...>::TiledShape_MN{});
|
||||
}
|
||||
|
||||
// The number of threads involved in a TiledCopy
|
||||
template <class... Args>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
size(TiledCopy<Args...> const&)
|
||||
{
|
||||
return typename TiledCopy<Args...>::TiledNumThr{};
|
||||
}
|
||||
|
||||
//
|
||||
// Display utilities
|
||||
//
|
||||
|
||||
template <class... Args>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
print_latex(TiledCopy<Args...> const& copy)
|
||||
{
|
||||
auto [layoutS_MN, thrID_S] = copy.get_layoutS_MN();
|
||||
auto [layoutD_MN, thrID_D] = copy.get_layoutD_MN();
|
||||
|
||||
print_latex_copy(layoutS_MN, thrID_S,
|
||||
layoutD_MN, thrID_D);
|
||||
}
|
||||
|
||||
// MNK Copy Layout to Latex TIKZ -- 8-value color coded by thread
|
||||
template <class LayoutS, class ThrIDS,
|
||||
class LayoutD, class ThrIDD>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print_latex_copy(LayoutS const& S, ThrIDS const& TS, // (m,n) -> (tid,vid) and tid -> thr_idx
|
||||
LayoutD const& D, ThrIDD const& TD) // (m,n) -> (tid,vid) and tid -> thr_idx
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(S) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(D) == Int<2>{});
|
||||
|
||||
assert(size<0>(S) == size<0>(D));
|
||||
assert(size<1>(S) == size<1>(D));
|
||||
|
||||
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("%% LayoutS: "); print(S); printf("\n");
|
||||
printf("%% ThrIDS : "); print(TS); printf("\n");
|
||||
printf("%% LayoutD: "); print(D); printf("\n");
|
||||
printf("%% ThrIDD : "); print(TD); printf("\n\n");
|
||||
|
||||
printf(latex_header);
|
||||
|
||||
// S starting at 0,0
|
||||
for (int i = 0; i < size<0>(S); ++i) {
|
||||
for (int j = 0; j < size<1>(S); ++j) {
|
||||
int thrid = S(i,j) % size(TS);
|
||||
int val_idx = S(i,j) / size(TS);
|
||||
int thr_idx = TS(thrid);
|
||||
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[thr_idx % 8],
|
||||
i, j,
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
|
||||
// D starting at 0,size<1>(S)+3
|
||||
for (int i = 0; i < size<0>(D); ++i) {
|
||||
for (int j = 0; j < size<1>(D); ++j) {
|
||||
int thrid = D(i,j) % size(TD);
|
||||
int val_idx = D(i,j) / size(TD);
|
||||
int thr_idx = TD(thrid);
|
||||
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[thr_idx % 8],
|
||||
i, j + size<1>(S) + 3,
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
|
||||
// S Labels
|
||||
for (int i = 0, j = -1; i < size<0>(S); ++i) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j, i);
|
||||
}
|
||||
for (int j = 0, i = -1; j < size<1>(S); ++j) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j, j);
|
||||
}
|
||||
// D Labels
|
||||
for (int i = 0, j = size<1>(D); i < size<0>(S); ++i) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j + size<1>(S) + 3, i);
|
||||
}
|
||||
for (int j = 0, i = -1; j < size<1>(D); ++j) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j + size<1>(S) + 3, j);
|
||||
}
|
||||
|
||||
// Footer
|
||||
printf(latex_footer);
|
||||
}
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
#include <cute/atom/copy_traits_sm75.hpp>
|
||||
#include <cute/atom/copy_traits_sm80.hpp>
|
||||
#include <cute/atom/copy_traits_sm90.hpp>
|
||||
// Config
|
||||
#if (__CUDACC_VER_MAJOR__ >= 12)
|
||||
# define CUTE_COPY_ATOM_TMA_SM90_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTE_COPY_ATOM_TMA_SM90_ENABLED)
|
||||
#include <cute/atom/copy_traits_sm90_tma.hpp>
|
||||
#endif
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,76 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/copy.hpp>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <class CopyOperation, class... CopyOpArgs>
|
||||
struct Copy_Traits
|
||||
{
|
||||
static_assert(sizeof(CopyOperation) == 0, "Copy_Traits not implemented for this Copy_Operation.");
|
||||
};
|
||||
|
||||
template <class S, class D>
|
||||
struct Copy_Traits<UniversalCopy<S,D>>
|
||||
{
|
||||
// Logical thread id to thread idx (one-thread)
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,Int<sizeof_bits<S>::value>>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,Int<sizeof_bits<D>::value>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<DefaultCopy>
|
||||
{
|
||||
// Logical thread id to thread idx (one-thread)
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,_1>, Stride<_0,_0>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,_1>, Stride<_0,_0>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
} // end namespace cute
|
||||
@@ -0,0 +1,143 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/copy_sm75.hpp>
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM75_U32x1_LDSM_N>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape <Shape < _8,_4>,_128>,
|
||||
Stride<Stride<_128,_0>, _1>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape <_32,_32>,
|
||||
Stride<_32, _1>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = DstLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM75_U32x2_LDSM_N>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape <Shape < _16,_2>,_128>,
|
||||
Stride<Stride<_128,_0>, _1>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape <_32,Shape <_32, _2>>,
|
||||
Stride<_32,Stride< _1,_1024>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = DstLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM75_U32x4_LDSM_N>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape < _32,_128>,
|
||||
Stride<_128, _1>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape <_32,Shape <_32, _4>>,
|
||||
Stride<_32,Stride< _1,_1024>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = DstLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM75_U16x2_LDSM_T>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape <Shape < _8,_4>,_128>,
|
||||
Stride<Stride<_128,_0>, _1>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape <Shape < _4, _8>,Shape <_16, _2>>,
|
||||
Stride<Stride<_256,_16>,Stride< _1,_128>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = DstLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM75_U16x4_LDSM_T>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape <Shape < _16,_2>,_128>,
|
||||
Stride<Stride<_128,_0>, _1>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape <Shape < _4, _8>,Shape <_16, _2, _2>>,
|
||||
Stride<Stride<_256,_16>,Stride< _1,_128,_1024>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = DstLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM75_U16x8_LDSM_T>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape < _32,_128>,
|
||||
Stride<_128, _1>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape <Shape < _4, _8>,Shape <_16, _2, _4>>,
|
||||
Stride<Stride<_256,_16>,Stride< _1,_128,_1024>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = DstLayout;
|
||||
};
|
||||
|
||||
} // end namespace cute
|
||||
@@ -0,0 +1,98 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/copy_sm80.hpp>
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <class S, class D>
|
||||
struct Copy_Traits<SM80_CP_ASYNC_CACHEALWAYS<S,D>>
|
||||
{
|
||||
// Logical thread id to thread idx (one-thread)
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,Int<sizeof_bits<S>::value>>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,Int<sizeof_bits<D>::value>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <class S, class D>
|
||||
struct Copy_Traits<SM80_CP_ASYNC_CACHEGLOBAL<S,D>>
|
||||
{
|
||||
// Logical thread id to thread idx (one-thread)
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,Int<sizeof_bits<S>::value>>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,Int<sizeof_bits<D>::value>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Element copy selector
|
||||
template <class SrcTensor, class DstTensor>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
select_elementwise_copy(SrcTensor const&, DstTensor const&)
|
||||
{
|
||||
using SrcType = typename SrcTensor::value_type;
|
||||
using DstType = typename DstTensor::value_type;
|
||||
|
||||
#if defined(CUTE_ARCH_CP_ASYNC_SM80_ENABLED)
|
||||
if constexpr (is_gmem<SrcTensor>::value && is_smem<DstTensor>::value &&
|
||||
sizeof(SrcType) == sizeof(DstType) &&
|
||||
(sizeof(SrcType) == 4 || sizeof(SrcType) == 8 || sizeof(SrcType) == 16))
|
||||
{
|
||||
return SM80_CP_ASYNC_CACHEALWAYS<SrcType,DstType>{};
|
||||
} else {
|
||||
return UniversalCopy<SrcType,DstType>{};
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
#else
|
||||
return UniversalCopy<SrcType,DstType>{};
|
||||
#endif
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/copy_sm90.hpp>
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
#include <cute/atom/copy_traits_sm75.hpp>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM90_U32x1_STSM_N>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = typename Copy_Traits<SM75_U32x1_LDSM_N>::DstLayout;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = typename Copy_Traits<SM75_U32x1_LDSM_N>::SrcLayout;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM90_U32x2_STSM_N>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = typename Copy_Traits<SM75_U32x2_LDSM_N>::DstLayout;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = typename Copy_Traits<SM75_U32x2_LDSM_N>::SrcLayout;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM90_U32x4_STSM_N>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = typename Copy_Traits<SM75_U32x4_LDSM_N>::DstLayout;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = typename Copy_Traits<SM75_U32x4_LDSM_N>::SrcLayout;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM90_U16x2_STSM_T>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = typename Copy_Traits<SM75_U16x2_LDSM_T>::DstLayout;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = typename Copy_Traits<SM75_U16x2_LDSM_T>::SrcLayout;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM90_U16x4_STSM_T>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = typename Copy_Traits<SM75_U16x4_LDSM_T>::DstLayout;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = typename Copy_Traits<SM75_U16x4_LDSM_T>::SrcLayout;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM90_U16x8_STSM_T>
|
||||
{
|
||||
// Logical thread id to thread idx (warp)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = typename Copy_Traits<SM75_U16x8_LDSM_T>::DstLayout;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = typename Copy_Traits<SM75_U16x8_LDSM_T>::SrcLayout;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
} // end namespace cute
|
||||
@@ -0,0 +1,795 @@
|
||||
/***************************************************************************************************
|
||||
* 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 <cuda.h>
|
||||
|
||||
#include <cute/arch/copy_sm90_desc.hpp>
|
||||
#include <cute/arch/copy_sm90_tma.hpp>
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
///////////////////////////// TMA_LOAD ///////////////////////////////////////
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct SM90_TMA_LOAD_OP : SM90_TMA_LOAD {};
|
||||
|
||||
// The executable SM90_TMA_LOAD with tma_desc and tma_mbar
|
||||
template <class NumBits>
|
||||
struct Copy_Traits<SM90_TMA_LOAD_OP, NumBits>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,NumBits>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,NumBits>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
|
||||
// SM90_TMA_LOAD arguments
|
||||
TmaDescriptor const& tma_desc_;
|
||||
uint64_t& tma_load_mbar_;
|
||||
|
||||
template <class Coord, int... Is>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
copy_unpack_(void const* const dst_ptr,
|
||||
Coord const& src_coord, seq<Is...>) const
|
||||
{
|
||||
#if 0
|
||||
print("THR (%d,%d,%d) BLK (%d,%d,%d)\n",
|
||||
threadIdx.x, threadIdx.y, threadIdx.z,
|
||||
blockIdx.x, blockIdx.y, blockIdx.z);
|
||||
print(" TMA Coord "); print(src_coord); print("\n");
|
||||
print(" TMA Shape "); print(make_tuple(uint64_t(tma_desc_.size0_),
|
||||
uint64_t(tma_desc_.size1_),
|
||||
uint64_t(tma_desc_.size2_),
|
||||
uint64_t(tma_desc_.size3_))); print("\n");
|
||||
#endif
|
||||
|
||||
SM90_TMA_LOAD::copy(&tma_desc_,
|
||||
tma_load_mbar_,
|
||||
dst_ptr,
|
||||
get<Is>(src_coord)...);
|
||||
}
|
||||
|
||||
// This is the copy_unpack dispatch for this Copy_Traits
|
||||
// Src needs to be a gmem tensor with TmaCoordIterator .data()
|
||||
// Dst needs to be a smem tensor
|
||||
template <class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE friend constexpr
|
||||
void
|
||||
copy_unpack(Copy_Traits const& traits,
|
||||
Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst)
|
||||
{
|
||||
//static_assert(is_gmem<TS>::value, "Expected gmem src for SM90_TMA_LOAD"); // TMA spoofed src tensor
|
||||
static_assert(is_smem<TD>::value, "Expected smem dst for SM90_TMA_LOAD");
|
||||
|
||||
traits.copy_unpack_(dst.data().get(), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
|
||||
}
|
||||
};
|
||||
|
||||
// The non-executable SM90_TMA_LOAD with tma_desc and no tma_mbar
|
||||
// Use .with(tma_mbar) to construct an executable version
|
||||
template <class NumBits, class GmemStrides>
|
||||
struct Copy_Traits<SM90_TMA_LOAD, NumBits, GmemStrides>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,NumBits>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,NumBits>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
|
||||
// SM90_TMA_LOAD arguments
|
||||
TmaDescriptor tma_desc_;
|
||||
GmemStrides g_stride_;
|
||||
|
||||
// Return TmaDescriptor/TensorMap
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
TmaDescriptor const*
|
||||
get_tma_descriptor() const {
|
||||
return &tma_desc_;
|
||||
}
|
||||
|
||||
// Construct an executable SM90_TMA_LOAD with tma_mbar
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_OP, NumBits>
|
||||
with(uint64_t& tma_mbar, 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};
|
||||
}
|
||||
|
||||
// Generate the TMA coord tensor
|
||||
template <class GShape>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_tma_tensor(GShape const& g_shape) const {
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
|
||||
constexpr int tma_rank = decltype(cute::min(rank(flatten(g_stride_)), Int<5>{}))::value;
|
||||
return make_tensor(ArithmeticTupleIterator(as_arithmetic_tuple(repeat<tma_rank>(Int<0>{}))),
|
||||
g_shape,
|
||||
g_stride_);
|
||||
}
|
||||
|
||||
// Don't try to execute a copy with SM90_TMA_LOAD before calling .with()
|
||||
template <class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE friend constexpr void
|
||||
copy_unpack(Copy_Traits const& traits,
|
||||
Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst) = delete;
|
||||
};
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
///////////////////////////// TMA_LOAD_MULTICAST /////////////////////////////
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct SM90_TMA_LOAD_MULTICAST_OP : SM90_TMA_LOAD_MULTICAST {};
|
||||
|
||||
template <class NumBits>
|
||||
struct Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBits>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,NumBits>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,NumBits>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
|
||||
// SM90_TMA_LOAD_MULTICAST arguments
|
||||
TmaDescriptor const& tma_desc_;
|
||||
uint64_t& tma_load_mbar_;
|
||||
uint16_t const& multicast_mask_;
|
||||
|
||||
template <class Coord, int... Is>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
copy_unpack_(void const* const dst_ptr,
|
||||
Coord const& src_coord, seq<Is...>) const
|
||||
{
|
||||
#if 0
|
||||
print("THR (%d,%d,%d) BLK (%d,%d,%d)\n",
|
||||
threadIdx.x, threadIdx.y, threadIdx.z,
|
||||
blockIdx.x, blockIdx.y, blockIdx.z);
|
||||
print(" TMA Coord "); print(src_coord); print("\n");
|
||||
print(" TMA Shape "); print(make_tuple(uint64_t(tma_desc_.size0_),
|
||||
uint64_t(tma_desc_.size1_),
|
||||
uint64_t(tma_desc_.size2_),
|
||||
uint64_t(tma_desc_.size3_))); print("\n");
|
||||
#endif
|
||||
|
||||
SM90_TMA_LOAD_MULTICAST::copy(&tma_desc_,
|
||||
tma_load_mbar_,
|
||||
multicast_mask_,
|
||||
dst_ptr,
|
||||
get<Is>(src_coord)...);
|
||||
}
|
||||
|
||||
template <class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE friend constexpr
|
||||
void
|
||||
copy_unpack(Copy_Traits const& traits,
|
||||
Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst)
|
||||
{
|
||||
//static_assert(is_gmem<TS>::value, "Expected gmem src for SM90_TMA_LOAD"); // TMA spoofed src tensor
|
||||
static_assert(is_smem<TD>::value, "Expected smem dst for SM90_TMA_LOAD_MULTICAST");
|
||||
|
||||
traits.copy_unpack_(dst.data().get(), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
|
||||
}
|
||||
};
|
||||
|
||||
template <class NumBits, class GmemStrides>
|
||||
struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBits, GmemStrides>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,NumBits>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,NumBits>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
|
||||
// SM90_TMA_LOAD_MULTICAST arguments
|
||||
TmaDescriptor tma_desc_;
|
||||
GmemStrides g_stride_;
|
||||
|
||||
// Return TmaDescriptor/TensorMap
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
TmaDescriptor const*
|
||||
get_tma_descriptor() const {
|
||||
return &tma_desc_;
|
||||
}
|
||||
|
||||
// Construct an executable SM90_TMA_LOAD_MULTICAST with tma_mbar
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBits>
|
||||
with(uint64_t& tma_load_mbar, uint16_t const& multicast_mask) const {
|
||||
return {tma_desc_, tma_load_mbar, multicast_mask};
|
||||
}
|
||||
|
||||
// Generate the TMA coord tensor
|
||||
template <class GShape>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_tma_tensor(GShape const& g_shape) const {
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
|
||||
constexpr int tma_rank = decltype(cute::min(rank(flatten(g_stride_)), Int<5>{}))::value;
|
||||
return make_tensor(ArithmeticTupleIterator(as_arithmetic_tuple(repeat<tma_rank>(Int<0>{}))),
|
||||
g_shape,
|
||||
g_stride_);
|
||||
}
|
||||
|
||||
// Don't try to execute a copy with SM90_TMA_LOAD_MULTICAST before calling .with()
|
||||
template <class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE friend constexpr void
|
||||
copy_unpack(Copy_Traits const& traits,
|
||||
Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst) = delete;
|
||||
};
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
///////////////////////////// TMA_STORE //////////////////////////////////////
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// The executable SM90_TMA_STORE with tma_desc
|
||||
template <class NumBits, class GmemStrides>
|
||||
struct Copy_Traits<SM90_TMA_STORE, NumBits, GmemStrides>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,NumBits>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,NumBits>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
|
||||
// SM90_TMA_STORE arguments
|
||||
TmaDescriptor tma_desc_;
|
||||
GmemStrides g_stride_;
|
||||
|
||||
// Generate the TMA coord tensor
|
||||
template <class GShape>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_tma_tensor(GShape const& g_shape) const {
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
|
||||
constexpr int tma_rank = decltype(cute::min(rank(flatten(g_stride_)), Int<5>{}))::value;
|
||||
return make_tensor(ArithmeticTupleIterator(as_arithmetic_tuple(repeat<tma_rank>(Int<0>{}))),
|
||||
g_shape,
|
||||
g_stride_);
|
||||
}
|
||||
|
||||
template <class Coord, int... Is>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
copy_unpack_(void const* const src_ptr,
|
||||
Coord const& dst_coord, seq<Is...>) const
|
||||
{
|
||||
#if 0
|
||||
print("THR (%d,%d,%d) BLK (%d,%d,%d)\n",
|
||||
threadIdx.x, threadIdx.y, threadIdx.z,
|
||||
blockIdx.x, blockIdx.y, blockIdx.z);
|
||||
print(" TMA Coord "); print(dst_coord); print("\n");
|
||||
print(" TMA Shape "); print(make_tuple(uint64_t(tma_desc_.size0_),
|
||||
uint64_t(tma_desc_.size1_),
|
||||
uint64_t(tma_desc_.size2_),
|
||||
uint64_t(tma_desc_.size3_))); print("\n");
|
||||
#endif
|
||||
|
||||
SM90_TMA_STORE::copy(&tma_desc_,
|
||||
src_ptr,
|
||||
get<Is>(dst_coord)...);
|
||||
}
|
||||
|
||||
// This is the copy_unpack dispatch for this Copy_Traits
|
||||
// Src needs to be a smem tensor
|
||||
// Dst needs to be a gmem tensor with TmaCoordIterator .data()
|
||||
template <class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE friend constexpr
|
||||
void
|
||||
copy_unpack(Copy_Traits const& traits,
|
||||
Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst)
|
||||
{
|
||||
static_assert(is_smem<TS>::value, "Expected smem src for SM90_TMA_STORE");
|
||||
//static_assert(is_gmem<TD>::value, "Expected gmem dst for SM90_TMA_STORE"); // TMA spoofed src tensor
|
||||
|
||||
traits.copy_unpack_(src.data().get(), dst.data().coord_, tuple_seq<decltype(dst.data().coord_)>{});
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// MAKE_TMA_COPY and related
|
||||
//
|
||||
|
||||
template <int B, int M, int S, class Offset, class SLayout>
|
||||
TMA::SmemSwizzleBits
|
||||
get_tma_swizzle_bits(ComposedLayout<Swizzle<B,M,S>,Offset,SLayout>)
|
||||
{
|
||||
static_assert(M == 4, "Expected 128b=16B=(2^4)B base swizzle.");
|
||||
static_assert(S == 3, "Unsupported layout swizzle");
|
||||
|
||||
switch (B) {
|
||||
default: static_assert(0 <= B && B <= 3, "Expected B = 0,1,2, or 3. Unsupported layout swizzle.");
|
||||
case 3: return TMA::SmemSwizzleBits::B128;
|
||||
case 2: return TMA::SmemSwizzleBits::B64;
|
||||
case 1: return TMA::SmemSwizzleBits::B32;
|
||||
case 0: return TMA::SmemSwizzleBits::DISABLE;
|
||||
}
|
||||
}
|
||||
|
||||
template <class Shape, class Stride>
|
||||
TMA::SmemSwizzleBits
|
||||
get_tma_swizzle_bits(Layout<Shape,Stride>)
|
||||
{
|
||||
return TMA::SmemSwizzleBits::DISABLE;
|
||||
}
|
||||
|
||||
template <int B, int M, int S, class Offset, class SLayout>
|
||||
auto
|
||||
get_nonswizzle_layout(ComposedLayout<Swizzle<B,M,S>,Offset,SLayout> const& slayout)
|
||||
{
|
||||
return slayout.layout_fn();
|
||||
}
|
||||
|
||||
template <class Shape, class Stride>
|
||||
auto
|
||||
get_nonswizzle_layout(Layout<Shape,Stride> const& slayout)
|
||||
{
|
||||
return slayout;
|
||||
}
|
||||
|
||||
/** Make a CuTe CTA-collective TiledCopy for a TMA operation.
|
||||
*
|
||||
* @param CopyOp The target copy operation: SM90_TMA_LOAD, SM90_TMA_LOAD_MULTICAST, SM90_TMA_STORE
|
||||
* @param gtensor The GMEM Tensor to be involved in the TMA.
|
||||
* @param slayout The SMEM Layout to be involved in the TMA.
|
||||
* @param cta_tile The CTA-local tile that each CTA will be tiling GMEM with.
|
||||
* This is often the blk_shape that is used to tile the GMEM for CTAs:
|
||||
* local_tile(gtensor, blk_shape, blk_coord) -> CTA-local tile of gtensor
|
||||
* @param cluster_size When using SM90_TMA_LOAD_MULTICAST, this can be a (static) power-of-2 <= 16
|
||||
* defining the multicast size (used to further partition the SMEM)
|
||||
* Else, static-1
|
||||
*
|
||||
* This code attempts to maximize the TMA box size. It does this by tracing
|
||||
* the SMEM "vector" -- the inverse of the smem layout -- to find the largest
|
||||
* contiguous array of smem that can be written to/from global memory given
|
||||
* the constraints that the TMA instruction imposes.
|
||||
*
|
||||
* This is accomplished by assigning "basis" strides to the GMEM to track which
|
||||
* modes of SMEM map to which modes of GMEM, then reorder the modes of GMEM according
|
||||
* to the SMEM vector, and then using those GMEM/SMEM modes to fill in the desc.
|
||||
*
|
||||
* Examples:
|
||||
using T = float;
|
||||
T* gptr = nullptr;
|
||||
|
||||
{
|
||||
// Simple 2D
|
||||
Tensor gtensor = make_tensor(gptr, make_shape(1024, 256), GenRowMajor{}); // K-Major GMEM
|
||||
auto slayout = make_layout(make_shape(_64{}, _32{}), GenRowMajor{}); // K-Major SMEM
|
||||
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout);
|
||||
}
|
||||
|
||||
{
|
||||
// GMMA 2D
|
||||
Tensor gtensor = make_tensor(gptr, make_shape(1024, 256)); // MN-Major GMEM
|
||||
auto slayout = tile_to_shape(GMMA::Layout_MN_SW128_Atom<T>{}, make_shape(_128{},_64{})); // MN-Major Swizzled+Tiled 128x64 SMEM
|
||||
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout);
|
||||
}
|
||||
|
||||
{
|
||||
// 3D
|
||||
Tensor gtensor = make_tensor(gptr, make_shape(1024, 32, 512), make_stride(64, Int<1>{}, 65536)); // GMEM
|
||||
auto slayout = make_layout(make_shape(_16{}, _8{}, _2{}), make_stride(_16{}, _1{}, _8{})); // SMEM w/ same major-mode
|
||||
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout);
|
||||
}
|
||||
|
||||
{
|
||||
// cuTENSOR 4D
|
||||
auto layout = make_shape(make_shape(32,40),make_shape(make_shape(8,8),656)); // GMEM
|
||||
auto cta_tile = make_shape(_128{},make_shape(_32{},_2{})); // GMEM Tiling:
|
||||
// Take 128-elem from m: m0 must divide 128,
|
||||
// m-last may be predicated
|
||||
// Take 32-elem from k0, 2-elem from k1
|
||||
auto slayout = make_layout(cta_tile); // Col-Major SMEM
|
||||
auto tma = make_tma_copy(SM90_TMA_LOAD{}, gtensor, slayout, cta_tile, Int<1>{});
|
||||
}
|
||||
*
|
||||
* Check the TMA box size and desc:
|
||||
print("TMA Box size: "); print(typename decltype(tma)::Tiler_MN{}); print("\n");
|
||||
print("TMA desc : "); print(tma.tma_desc_); print("\n");
|
||||
*
|
||||
* Usage:
|
||||
Tensor mA = tma_a.get_tma_tensor(make_shape(M,N)); // (M,N) TMA coord tensor
|
||||
Tensor gA = local_tile(mA, cta_tile, cta_coord); // (BLK_M,BLK_N) TMA coord tensor for this CTA
|
||||
Tensor sA = make_tensor(make_smem_ptr<T>(sptr), slayout); // (BLK_M,BLK_N) SMEM tensor
|
||||
|
||||
auto cta_tma = tma.get_slice(cta_idx_in_cluster); // Slice for multicast partitioning
|
||||
Tensor tAgA = cta_tma.partition_S(gA); // Partition for src
|
||||
Tensor tAsA = cta_tma.partition_D(sA); // Partition for dst
|
||||
|
||||
copy(tma.with(barrier, mcast_mask), tAgA, tAsA); // copy with supporting TMA params
|
||||
*/
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tile,
|
||||
class Cluster_Size>
|
||||
CUTE_HOST
|
||||
auto
|
||||
make_tma_copy(CopyOp,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tile const& cta_tile,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
static_assert((std::is_same<CopyOp, SM90_TMA_LOAD>::value && is_constant<1, Cluster_Size>::value) ||
|
||||
(std::is_same<CopyOp, SM90_TMA_LOAD_MULTICAST>::value) ||
|
||||
(std::is_same<CopyOp, SM90_TMA_STORE>::value && is_constant<1, Cluster_Size>::value));
|
||||
|
||||
using T = typename Tensor<GEngine,GLayout>::value_type;
|
||||
|
||||
//
|
||||
// TMA parameter checking
|
||||
//
|
||||
|
||||
auto flat_glayout = flatten(gtensor.layout());
|
||||
|
||||
CUTE_STATIC_ASSERT_V(rank(flatten(cta_tile)) <= Int<5>{},
|
||||
"CTA_Tile cannot have more than five modes, TMA arch restriction.");
|
||||
CUTE_STATIC_ASSERT_V(rank(flat_glayout) <= Int<5>{} || rank(flatten(cta_tile)) <= Int<4>{},
|
||||
"If GTensor has more than five modes, then CTA_Tile cannot have more than four modes. TMA multimode.");
|
||||
CUTE_STATIC_ASSERT_V(compatible(product_each(shape(slayout)), shape(cta_tile)),
|
||||
"CTA_Tile must be compatible with SLayout.");
|
||||
CUTE_STATIC_ASSERT_V(is_integral<Cluster_Size>{} && has_single_bit(cluster_size) && cluster_size <= Int<16>{},
|
||||
"Expecting a pow2 integral Cluster_Size leq 16.");
|
||||
CUTE_STATIC_ASSERT_V(size(slayout) % cluster_size == Int<0>{},
|
||||
"ClusterShape must divide domain size of slayout.");
|
||||
|
||||
//
|
||||
// TMA slayout manipulation
|
||||
//
|
||||
|
||||
auto tma_multimode = rank(flat_glayout) > Int<5>{};
|
||||
|
||||
// Invert the smem to get the largest contiguous vector in the smem layout
|
||||
auto inv_smem_layout = right_inverse(get_nonswizzle_layout(slayout));
|
||||
// trunc_smem_idx -> trunc_smem_coord
|
||||
|
||||
// Map from smem idx to a gmem mode
|
||||
auto sidx_to_gmode = flatten(composition(make_identity_layout(cta_tile), inv_smem_layout));
|
||||
|
||||
// Truncate any incompatibilities
|
||||
auto smem_rank = find_if(stride(sidx_to_gmode), [](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 smem-gmem vectorization for TMA.");
|
||||
constexpr int smem_tma_rank = cute::min(int(smem_rank), (tma_multimode ? 4 : 5));
|
||||
|
||||
// Keep only the static-1 basis modes into gmem
|
||||
auto sidx_to_gmode_cluster_trunc = take<0,smem_tma_rank>(sidx_to_gmode);
|
||||
// Keep only the portion each multicast CTA will be responsible for
|
||||
auto sidx_to_gmode_cta_trunc = composition(sidx_to_gmode_cluster_trunc, shape_div(size(sidx_to_gmode_cluster_trunc), cluster_size));
|
||||
|
||||
//
|
||||
// TMA gtensor manipulation
|
||||
//
|
||||
|
||||
// Generate a TupleBasis for the gtensor
|
||||
auto flat_gbasis = make_basis_like(shape(flat_glayout));
|
||||
|
||||
// Fold the flat_gbasis into the glayout
|
||||
auto glayout_basis = make_layout(shape(gtensor),
|
||||
stride(composition(make_layout(repeat_like(shape(flat_glayout), Int<2>{}), flat_gbasis),
|
||||
make_layout(repeat_like(shape(gtensor), Int<2>{})))));
|
||||
|
||||
// Tile the modes of gtensor with cta_tile
|
||||
auto cta_glayout_basis = composition(glayout_basis, cta_tile);
|
||||
|
||||
// Check that the cta_tile selects modes from gtensor properly
|
||||
for_each(flatten(stride(cta_glayout_basis)), [](auto d) {
|
||||
static_assert(is_constant<1, decltype(d.value())>::value,
|
||||
"CTA_Tile does not faithfully partition the GMEM, it should select the number of elements from each mode of glayout.");
|
||||
});
|
||||
|
||||
// Tile the modes of gtensor again with the truncated cta_tile o inv_smem_layout
|
||||
auto tma_layout_cta_trunc = flatten(composition(glayout_basis, sidx_to_gmode_cta_trunc));
|
||||
|
||||
// Append any missing basis on the end as size-1 modes b/c they got truncated
|
||||
auto missing_basis = fold(stride(tma_layout_cta_trunc), flat_gbasis, [](auto init, auto e){
|
||||
auto k = find(init, e);
|
||||
return remove<k>(init);
|
||||
});
|
||||
|
||||
// The appended map from truncated smem codomain to gmem mode: trunc_smem_idx -> gmem_mode
|
||||
auto tma_layout_cta = flatten(make_layout(tma_layout_cta_trunc,
|
||||
make_layout(repeat<rank(missing_basis)>(Int<1>{}), missing_basis)));
|
||||
|
||||
#if 0
|
||||
print("g_layout : "); print(gtensor.layout()); print("\n");
|
||||
print("s_layout : "); print(slayout); print("\n");
|
||||
print("cta_tile : "); print(cta_tile); print("\n");
|
||||
print("cluster_size : "); print(cluster_size); print("\n");
|
||||
print("flat_gbasis : "); print(flat_gbasis); print("\n");
|
||||
print("cta_glayout : "); print(cta_glayout_basis); print("\n");
|
||||
print("inv_smem : "); print(inv_smem_layout); print("\n");
|
||||
print("sidx_to_gmode : "); print(sidx_to_gmode); print("\n");
|
||||
print("missing_b : "); print(missing_basis); print("\n");
|
||||
print("tma_layout_cta: "); print(tma_layout_cta); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// TMA gmem desc info
|
||||
//
|
||||
|
||||
constexpr int TmaRANK = cute::min(rank(flat_glayout), 5);
|
||||
void* gmem_address = (void*) gtensor.data();
|
||||
|
||||
cute::array<cuuint64_t, 5> gmem_prob_shape = {1,1,1,1,1};
|
||||
cute::array<cuuint64_t, 5> gmem_prob_stride = {0,0,0,0,0};
|
||||
for_each(make_seq<rank(tma_layout_cta)>{}, [&](auto i) {
|
||||
// NOTE : WAR g++-7.3.5, let it deduce e rather than fuse with below
|
||||
auto e = stride<i>(tma_layout_cta);
|
||||
constexpr int j = decltype(e.mode())::value;
|
||||
constexpr int tma_i = i < 5 ? i : 4;
|
||||
|
||||
// Problem stride
|
||||
uint64_t stride_j = stride<j>(flat_glayout) * sizeof(T);
|
||||
uint64_t old_stride = gmem_prob_stride[tma_i];
|
||||
gmem_prob_stride[tma_i] = gcd(gmem_prob_stride[tma_i], stride_j);
|
||||
|
||||
// Problem shape
|
||||
uint64_t shape_j = shape<j>(flat_glayout);
|
||||
if (gmem_prob_stride[tma_i] != 0) {
|
||||
// We're "resetting" this TMA mode and using it as a "multimode"
|
||||
// Recurrence: g_shape = (s_i - 1) * (d_i / gcd_j d_j) + 1
|
||||
gmem_prob_shape[tma_i] = (gmem_prob_shape[tma_i]-1) * (old_stride / gmem_prob_stride[tma_i])
|
||||
+ (shape_j-1) * (stride_j / gmem_prob_stride[tma_i])
|
||||
+ 1;
|
||||
} else {
|
||||
gmem_prob_shape[tma_i] = shape_j;
|
||||
}
|
||||
});
|
||||
|
||||
assert((reinterpret_cast<uint64_t>(gmem_address) & 0b1111) == 0); // Address must be 16B-aligned
|
||||
|
||||
assert(gmem_prob_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(gmem_prob_shape[0] <= (uint64_t(1) << 32)); // Size must be max 2^32
|
||||
assert(gmem_prob_shape[1] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(gmem_prob_shape[1] <= (uint64_t(1) << 32)); // Size must be max 2^32
|
||||
assert(gmem_prob_shape[2] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(gmem_prob_shape[2] <= (uint64_t(1) << 32)); // Size must be max 2^32
|
||||
assert(gmem_prob_shape[3] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(gmem_prob_shape[3] <= (uint64_t(1) << 32)); // Size must be max 2^32
|
||||
assert(gmem_prob_shape[4] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(gmem_prob_shape[4] <= (uint64_t(1) << 32)); // Size must be max 2^32
|
||||
|
||||
assert((gmem_prob_stride[0]) == sizeof(T)); // First stride is implicitly 1
|
||||
assert((gmem_prob_stride[1]) < (uint64_t(1) << 40)); // Stride must be max 2^40
|
||||
assert((gmem_prob_stride[1] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
|
||||
assert((gmem_prob_stride[2]) < (uint64_t(1) << 40)); // Stride must be max 2^40
|
||||
assert((gmem_prob_stride[2] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
|
||||
assert((gmem_prob_stride[3]) < (uint64_t(1) << 40)); // Stride must be max 2^40
|
||||
assert((gmem_prob_stride[3] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
|
||||
assert((gmem_prob_stride[4]) < (uint64_t(1) << 40)); // Stride must be max 2^40
|
||||
assert((gmem_prob_stride[4] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
|
||||
|
||||
//
|
||||
// TMA smem desc info
|
||||
//
|
||||
|
||||
// TMA smem box size
|
||||
cute::array<cuuint32_t, 5> smem_box_shape = {1,1,1,1,1};
|
||||
for_each(make_seq<rank(tma_layout_cta)>{}, [&](auto i) {
|
||||
uint32_t shape_i = shape<i>(tma_layout_cta);
|
||||
constexpr int tma_i = i < 5 ? i : 4;
|
||||
if (tma_multimode && tma_i == 4) {
|
||||
// We're "reusing" this TMA mode and using it as a "multimode"
|
||||
smem_box_shape[tma_i] = 1;
|
||||
} else {
|
||||
smem_box_shape[tma_i] = shape_i;
|
||||
}
|
||||
});
|
||||
|
||||
// TMA smem mode strides
|
||||
[[maybe_unused]] cute::array<cuuint32_t, 5> smem_box_stride = {1,1,1,1,1};
|
||||
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8
|
||||
|
||||
assert(smem_box_stride[0] >= (uint32_t(1))); // Stride must be min 1
|
||||
assert(smem_box_stride[0] <= (uint32_t(8))); // Stride must be max 2^3
|
||||
assert(smem_box_stride[1] >= (uint32_t(1))); // Stride must be min 1
|
||||
assert(smem_box_stride[1] <= (uint32_t(8))); // Stride must be max 2^3
|
||||
assert(smem_box_stride[2] >= (uint32_t(1))); // Stride must be min 1
|
||||
assert(smem_box_stride[2] <= (uint32_t(8))); // Stride must be max 2^3
|
||||
assert(smem_box_stride[3] >= (uint32_t(1))); // Stride must be min 1
|
||||
assert(smem_box_stride[3] <= (uint32_t(8))); // Stride must be max 2^3
|
||||
assert(smem_box_stride[4] >= (uint32_t(1))); // Stride must be min 1
|
||||
assert(smem_box_stride[4] <= (uint32_t(8))); // Stride must be max 2^3
|
||||
|
||||
//
|
||||
// Construct the descriptor
|
||||
//
|
||||
|
||||
TmaDescriptor tma_desc = {0};
|
||||
|
||||
#if (__CUDACC_VER_MAJOR__ >= 12)
|
||||
|
||||
//
|
||||
// TMA general info
|
||||
//
|
||||
|
||||
cuuint32_t tma_dim = TmaRANK;
|
||||
CUtensorMapDataType tma_format = TMA::to_CUtensorMapDataType<T>();
|
||||
CUtensorMapInterleave tma_interleave = CU_TENSOR_MAP_INTERLEAVE_NONE;
|
||||
CUtensorMapL2promotion tma_l2Promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE;
|
||||
CUtensorMapFloatOOBfill tma_oobFill = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE;
|
||||
|
||||
// TMA smem swizzle type
|
||||
CUtensorMapSwizzle smem_swizzle = TMA::to_CUtensorMapSwizzle(get_tma_swizzle_bits(slayout));
|
||||
|
||||
CUresult result = cuTensorMapEncodeTiled(
|
||||
&tma_desc,
|
||||
tma_format,
|
||||
tma_dim,
|
||||
gmem_address,
|
||||
gmem_prob_shape.data(),
|
||||
gmem_prob_stride.data() + 1, // gmem_prob_stride[0] implicitly 1
|
||||
smem_box_shape.data(),
|
||||
smem_box_stride.data(),
|
||||
tma_interleave,
|
||||
smem_swizzle,
|
||||
tma_l2Promotion,
|
||||
tma_oobFill);
|
||||
|
||||
if (result != CUDA_SUCCESS) {
|
||||
std::cerr << "TMA Desc Addr: " << &tma_desc
|
||||
<< "\nformat " << tma_format
|
||||
<< "\ndim " << tma_dim
|
||||
<< "\ngmem_address " << gmem_address
|
||||
<< "\nglobalDim " << gmem_prob_shape
|
||||
<< "\nglobalStrides " << gmem_prob_stride
|
||||
<< "\nboxDim " << smem_box_shape
|
||||
<< "\nelementStrides " << smem_box_stride
|
||||
<< "\ninterleave " << tma_interleave
|
||||
<< "\nswizzle " << smem_swizzle
|
||||
<< "\nl2Promotion " << tma_l2Promotion
|
||||
<< "\noobFill " << tma_oobFill << std::endl;
|
||||
std::cerr << "Error: Failed to intialize the TMA descriptor " << result << std::endl;
|
||||
assert(false);
|
||||
}
|
||||
#endif // (__CUDACC_VER_MAJOR__ >= 12)
|
||||
|
||||
//
|
||||
// Construct the Copy_Traits
|
||||
//
|
||||
|
||||
// Finally, get the inverse permutation of the E<i> bases for the mocked gmem stride
|
||||
auto gmem_stride_bases_flat = transform(make_seq<rank(tma_layout_cta)>{}, [&](auto i) {
|
||||
auto k = find(stride(tma_layout_cta), E<i>{});
|
||||
// NOTE: gcc 7.3.5 WAR -- avoid if constexpr
|
||||
int32_t tma_coord_stride = int32_t(stride<i>(flat_glayout) * sizeof(T) / (gmem_prob_stride[4] != 0 ? gmem_prob_stride[4] : 16));
|
||||
return conditional_return(tma_multimode && (k >= Int<4>{}),
|
||||
E<4>{} * tma_coord_stride, // The 4th TMA mode is the multimode, use int32_t coord stride
|
||||
E<k>{});
|
||||
});
|
||||
|
||||
// Give that the profile of gtensor and fold it
|
||||
auto gmem_stride_bases = stride(composition(make_layout(repeat_like(shape(flat_glayout), Int<2>{}), gmem_stride_bases_flat),
|
||||
make_layout(repeat_like(shape(gtensor), Int<2>{}))));
|
||||
|
||||
constexpr int num_bits = size(sidx_to_gmode_cta_trunc) * sizeof(T) * 8;
|
||||
using Traits = Copy_Traits<CopyOp, Int<num_bits>, decltype(gmem_stride_bases)>;
|
||||
|
||||
#if 0
|
||||
print("num_bits : "); print(num_bits); print("\n");
|
||||
print("g_stride_bases: "); print(gmem_stride_bases); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// Construct the TiledCopy
|
||||
//
|
||||
|
||||
// The ThrVal layout for 1 TMA instruction within cta_tile
|
||||
auto layout_tv_1 = composition(inv_smem_layout, make_layout(make_shape(cluster_size, size(sidx_to_gmode_cta_trunc)), GenRowMajor{}));
|
||||
// The ThrVal layout for N TMA instructions within cta_tile
|
||||
auto layout_tv = tile_to_shape(layout_tv_1, make_shape(cluster_size, size(cta_tile)/cluster_size));
|
||||
|
||||
#if 0
|
||||
print("layout_tv : "); print(layout_tv); print("\n");
|
||||
#endif
|
||||
|
||||
return TiledCopy<Copy_Atom<Traits,T>, decltype(layout_tv), decltype(cta_tile)>{tma_desc, gmem_stride_bases};
|
||||
}
|
||||
|
||||
// Explicit defaulting
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout>
|
||||
CUTE_HOST
|
||||
auto
|
||||
make_tma_copy(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout)
|
||||
{
|
||||
return make_tma_copy(copy_op, gtensor, slayout, product_each(shape(slayout)), Int<1>{});
|
||||
}
|
||||
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class Cluster_Size>
|
||||
CUTE_HOST
|
||||
auto
|
||||
make_tma_copy(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
return make_tma_copy(copy_op, gtensor, slayout, product_each(shape(slayout)), cluster_size);
|
||||
}
|
||||
|
||||
} // end namespace cute
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,70 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/mma.hpp>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <class MMAOperation, class... MMAOpArgs>
|
||||
struct MMA_Traits
|
||||
{
|
||||
static_assert(sizeof(MMAOperation) == 0, "MMA_Traits not implemented for this MMA_Operation.");
|
||||
};
|
||||
|
||||
template <class D, class A, class B, class C>
|
||||
struct MMA_Traits<UniversalFMA<D,A,B,C>>
|
||||
{
|
||||
using ElementDVal = D;
|
||||
using ElementAVal = A;
|
||||
using ElementBVal = B;
|
||||
using ElementCVal = C;
|
||||
|
||||
// Logical shape of the MMA
|
||||
using Shape_MNK = Shape<_1,_1,_1>;
|
||||
|
||||
// Logical thread id (tid) -> tidx
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
// (Logical thread id (tid), Logical value id (vid)) -> coord
|
||||
|
||||
// (tid,vid) -> (m,k)
|
||||
using ALayout = Layout<Shape<_1,_1>>;
|
||||
// (tid,vid) -> (n,k)
|
||||
using BLayout = Layout<Shape<_1,_1>>;
|
||||
// (tid,vid) -> (m,n)
|
||||
using CLayout = Layout<Shape<_1,_1>>;
|
||||
};
|
||||
|
||||
} // namespace cute
|
||||
@@ -0,0 +1,73 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/mma_sm61.hpp>
|
||||
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM61_DP4A>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
|
||||
using Shape_MNK = Shape<_1,_1,_4>;
|
||||
using ThrID = Layout<_1>;
|
||||
using ALayout = Layout<Shape<_1,_4>>;
|
||||
using BLayout = Layout<Shape<_1,_4>>;
|
||||
using CLayout = Layout<Shape<_1,_1>>;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM61_DP2A>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int16_t;
|
||||
using ElementBVal = int16_t;
|
||||
using ElementCVal = int32_t;
|
||||
|
||||
using Shape_MNK = Shape<_1,_1,_2>;
|
||||
using ThrID = Layout<_1>;
|
||||
using ALayout = Layout<Shape<_1,_2>>;
|
||||
using BLayout = Layout<Shape<_1,_2>>;
|
||||
using CLayout = Layout<Shape<_1,_1>>;
|
||||
};
|
||||
|
||||
} // namespace cute
|
||||
@@ -0,0 +1,198 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/mma_sm70.hpp>
|
||||
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
namespace {
|
||||
|
||||
// Logical thread id to thread idx (quadpair)
|
||||
using SM70_QuadPair = Layout<Shape <_4, _2>,
|
||||
Stride<_1,_16>>;
|
||||
// (T8,V4) -> (M8,K4)
|
||||
using SM70_8x4_Row = Layout<Shape <_8,_4>,
|
||||
Stride<_1,_8>>;
|
||||
// (T8,V4) -> (M8,K4)
|
||||
using SM70_8x4_Col = Layout<Shape <Shape <_4,_2>,_4>,
|
||||
Stride<Stride<_8,_4>,_1>>;
|
||||
// (T8,V8) -> (M8,N8)
|
||||
using SM70_8x8_16b = Layout<Shape <_8,_8>,
|
||||
Stride<_1,_8>>;
|
||||
// (T8,V8) -> (M8,N8)
|
||||
using SM70_8x8_32b = Layout<Shape <Shape <_2, _2,_2>,Shape <_2,_2, _2>>,
|
||||
Stride<Stride<_1,_16,_4>,Stride<_8,_2,_32>>>;
|
||||
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F16F16F16F16_TN>
|
||||
{
|
||||
using ElementDVal = half_t;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = half_t;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Row;
|
||||
using BLayout = SM70_8x4_Row;
|
||||
using CLayout = SM70_8x8_16b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F16F16F16F16_NT>
|
||||
{
|
||||
using ElementDVal = half_t;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = half_t;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Col;
|
||||
using BLayout = SM70_8x4_Col;
|
||||
using CLayout = SM70_8x8_16b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F16F16F16F16_NN>
|
||||
{
|
||||
using ElementDVal = half_t;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = half_t;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Col;
|
||||
using BLayout = SM70_8x4_Row;
|
||||
using CLayout = SM70_8x8_16b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F16F16F16F16_TT>
|
||||
{
|
||||
using ElementDVal = half_t;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = half_t;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Row;
|
||||
using BLayout = SM70_8x4_Col;
|
||||
using CLayout = SM70_8x8_16b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F32F16F16F32_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = float;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Row;
|
||||
using BLayout = SM70_8x4_Row;
|
||||
using CLayout = SM70_8x8_32b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F32F16F16F32_NT>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = float;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Col;
|
||||
using BLayout = SM70_8x4_Col;
|
||||
using CLayout = SM70_8x8_32b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F32F16F16F32_NN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = float;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Col;
|
||||
using BLayout = SM70_8x4_Row;
|
||||
using CLayout = SM70_8x8_32b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM70_8x8x4_F32F16F16F32_TT>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = float;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = SM70_QuadPair;
|
||||
using ALayout = SM70_8x4_Row;
|
||||
using BLayout = SM70_8x4_Col;
|
||||
using CLayout = SM70_8x8_32b;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
} // namespace cute
|
||||
@@ -0,0 +1,81 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/mma_sm75.hpp>
|
||||
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM75_16x8x8_F32F16F16F32_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = float;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_8>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_2>,Stride<_16,_1>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,_2>,
|
||||
Stride<Stride<_16,_1>,_8>>;
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_2>,Stride<_16,_1>>>;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM75_8x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_16>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,_4>,
|
||||
Stride<Stride<_32,_1>,_8>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,_4>,
|
||||
Stride<Stride<_32,_1>,_8>>;
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,_2>,
|
||||
Stride<Stride<_16,_1>,_8>>;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cute
|
||||
@@ -0,0 +1,446 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/mma_sm80.hpp>
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
#include <cute/numeric/integer_subbyte.hpp>
|
||||
|
||||
#include <cutlass/numeric_types.h>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
namespace {
|
||||
|
||||
// (T32,V1) -> (M8,N8)
|
||||
using SM80_8x4 = Layout<Shape <Shape < _4,_8>,_1>,
|
||||
Stride<Stride< _8,_1>,_0>>;
|
||||
// (T32,V2) -> (M8,N8)
|
||||
using SM80_8x8_Row = Layout<Shape <Shape < _4,_8>,_2>,
|
||||
Stride<Stride<_16,_1>,_8>>;
|
||||
// (T32,V4) -> (M8,N16)
|
||||
using SM80_8x16_Row = Layout<Shape <Shape < _4,_8>,_4>,
|
||||
Stride<Stride<_32,_1>,_8>>;
|
||||
// (T32,V4) -> (M16,N8)
|
||||
using SM80_16x8_Row = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
|
||||
}
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////// fp16 = fp16 * fp16 + fp16 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x8_F16F16F16F16_TN>
|
||||
{
|
||||
using ElementDVal = half_t;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = half_t;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_8>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = SM80_16x8_Row;
|
||||
using BLayout = SM80_8x8_Row;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_F16F16F16F16_TN>
|
||||
{
|
||||
using ElementDVal = half_t;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = half_t;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_16>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2, _2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8,_128>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _2>>,
|
||||
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////// fp32 = fp16 * fp16 + fp32 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x8_F32F16F16F32_TN>
|
||||
: MMA_Traits<SM80_16x8x8_F16F16F16F16_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = float;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_F32F16F16F32_TN>
|
||||
: MMA_Traits<SM80_16x8x16_F16F16F16F16_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = half_t;
|
||||
using ElementBVal = half_t;
|
||||
using ElementCVal = float;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////// fp32 = bf16 * bf16 + fp32 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x8_F32BF16BF16F32_TN>
|
||||
: MMA_Traits<SM80_16x8x8_F16F16F16F16_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = bfloat16_t;
|
||||
using ElementBVal = bfloat16_t;
|
||||
using ElementCVal = float;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_F32BF16BF16F32_TN>
|
||||
: MMA_Traits<SM80_16x8x16_F16F16F16F16_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = bfloat16_t;
|
||||
using ElementBVal = bfloat16_t;
|
||||
using ElementCVal = float;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////// fp32 = tf32 * tf32 + fp32 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x4_F32TF32TF32F32_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = cutlass::tfloat32_t;
|
||||
using ElementBVal = cutlass::tfloat32_t;
|
||||
using ElementCVal = float;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_4>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,_2>,
|
||||
Stride<Stride<_16,_1>,_8>>;
|
||||
using BLayout = SM80_8x4;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x8_F32TF32TF32F32_TN>
|
||||
{
|
||||
using ElementDVal = float;
|
||||
using ElementAVal = cutlass::tfloat32_t;
|
||||
using ElementBVal = cutlass::tfloat32_t;
|
||||
using ElementCVal = float;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_8>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _2>>,
|
||||
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
|
||||
using BLayout = Layout<Shape <Shape <_4,_8>, _2>,
|
||||
Stride<Stride<_8,_1>,_32>>;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////// fp64 = fp64 * fp64 + fp64 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x4_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = double;
|
||||
using ElementAVal = double;
|
||||
using ElementBVal = double;
|
||||
using ElementCVal = double;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_4>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = SM80_8x4;
|
||||
using BLayout = SM80_8x4;
|
||||
using CLayout = SM80_8x8_Row;
|
||||
};
|
||||
|
||||
// Custom complex fp64 MMA composed of 4 fp64 MMAs -- same layouts
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x4_C64C64C64C64_TN>
|
||||
: MMA_Traits<SM80_8x8x4_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = complex<double>;
|
||||
using ElementAVal = complex<double>;
|
||||
using ElementBVal = complex<double>;
|
||||
using ElementCVal = complex<double>;
|
||||
};
|
||||
|
||||
// Custom complex fp64 MMA composed of 3 fp64 MMAs -- same layouts
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x4_GC64C64C64GC64_TN>
|
||||
: MMA_Traits<SM80_8x8x4_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = typename SM80_8x8x4_GC64C64C64GC64_TN::GaussComplex;
|
||||
using ElementAVal = complex<double>;
|
||||
using ElementBVal = complex<double>;
|
||||
using ElementCVal = typename SM80_8x8x4_GC64C64C64GC64_TN::GaussComplex;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////// s32 = s8 * s8 + s32 ///////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
|
||||
using Shape_MNK = Shape<_8,_8,_16>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = SM80_8x16_Row;
|
||||
using BLayout = SM80_8x16_Row;
|
||||
using CLayout = SM80_8x8_Row;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32S8S8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_16>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2>>,
|
||||
Stride<Stride<_64,_1>,Stride<_16,_8>>>;
|
||||
using BLayout = SM80_8x16_Row;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32S8S8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_32>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2, _2>>,
|
||||
Stride<Stride<_64,_1>,Stride<_16,_8,_256>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>, Shape <_4, _2>>,
|
||||
Stride<Stride<_32,_1>, Stride<_8,_128>>>;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32S8S8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN> {};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////// s32 = s8 * u8 + s32 ///////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32S8U8S32_TN>
|
||||
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = uint8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32S8U8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_8x8x16_S32S8U8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32S8U8S32_TN>
|
||||
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = uint8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32S8U8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x16_S32S8U8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32S8U8S32_TN>
|
||||
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = int8_t;
|
||||
using ElementBVal = uint8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32S8U8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x32_S32S8U8S32_TN> {};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////// s32 = u8 * s8 + s32 ///////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32U8S8S32_TN>
|
||||
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = uint8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32U8S8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_8x8x16_S32U8S8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32U8S8S32_TN>
|
||||
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = uint8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32U8S8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x16_S32U8S8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32U8S8S32_TN>
|
||||
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = uint8_t;
|
||||
using ElementBVal = int8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32U8S8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x32_S32U8S8S32_TN> {};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////// s32 = u8 * u8 + s32 ///////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32U8U8S32_TN>
|
||||
: MMA_Traits<SM80_8x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = uint8_t;
|
||||
using ElementBVal = uint8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_8x8x16_S32U8U8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_8x8x16_S32U8U8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32U8U8S32_TN>
|
||||
: MMA_Traits<SM80_16x8x16_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = uint8_t;
|
||||
using ElementBVal = uint8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x16_S32U8U8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x16_S32U8U8S32_TN> {};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32U8U8S32_TN>
|
||||
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = uint8_t;
|
||||
using ElementBVal = uint8_t;
|
||||
using ElementCVal = int32_t;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x32_S32U8U8S32_TN_SATURATE>
|
||||
: MMA_Traits<SM80_16x8x32_S32U8U8S32_TN> {};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
/////////////////////////// s32 = b1 ^ b1 + s32 ///////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM80_16x8x256_S32U1U1S32_TN_XORPOPC>
|
||||
{
|
||||
using ElementDVal = int32_t;
|
||||
using ElementAVal = cute::uint1b_t;
|
||||
using ElementBVal = cute::uint1b_t;
|
||||
using ElementCVal = int32_t;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_256>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <_32,Shape < _8, _4,_2, _2>>,
|
||||
Stride<_64,Stride<_64,_16,_8,_2048>>>;
|
||||
using BLayout = Layout<Shape <_32,Shape <_32, _2>>,
|
||||
Stride<_32,Stride< _1,_1024>>>;
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
} // end namespace cute
|
||||
@@ -0,0 +1,132 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/mma_sm90.hpp>
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
|
||||
#include <cute/layout.hpp>
|
||||
|
||||
namespace cute {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////// fp64 = fp64 * fp64 + fp64 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = double;
|
||||
using ElementAVal = double;
|
||||
using ElementBVal = double;
|
||||
using ElementCVal = double;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_4>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,_2>,
|
||||
Stride<Stride<_16,_1>,_8>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,_1>,
|
||||
Stride<Stride< _8,_1>,_0>>;
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = double;
|
||||
using ElementAVal = double;
|
||||
using ElementBVal = double;
|
||||
using ElementCVal = double;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_8>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _2>>,
|
||||
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>, _2>,
|
||||
Stride<Stride< _8,_1>,_32>>;
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = double;
|
||||
using ElementAVal = double;
|
||||
using ElementBVal = double;
|
||||
using ElementCVal = double;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_16>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape <_2, _4>>,
|
||||
Stride<Stride<_16,_1>,Stride<_8,_64>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>, _4>,
|
||||
Stride<Stride< _8,_1>,_32>>;
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////
|
||||
//////////////////////// cfp64 = cfp64 * cfp64 + cfp64 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x4_C64C64C64C64_TN>
|
||||
: MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = complex<double>;
|
||||
using ElementAVal = complex<double>;
|
||||
using ElementBVal = complex<double>;
|
||||
using ElementCVal = complex<double>;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x8_C64C64C64C64_TN>
|
||||
: MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = complex<double>;
|
||||
using ElementAVal = complex<double>;
|
||||
using ElementBVal = complex<double>;
|
||||
using ElementCVal = complex<double>;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x16_C64C64C64C64_TN>
|
||||
: MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
|
||||
{
|
||||
using ElementDVal = complex<double>;
|
||||
using ElementAVal = complex<double>;
|
||||
using ElementBVal = complex<double>;
|
||||
using ElementCVal = complex<double>;
|
||||
};
|
||||
|
||||
} // end namespace cute
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user