CUTLASS 3.6.0 (#1850)
* v3.6 * update changelog * update readme * fix typo * fixing typos * hopper gemm with weight prefetch --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
co-authored by
yuzhai
Haicheng Wu
parent
0837a2a00a
commit
cc3c29a81a
@@ -30,16 +30,13 @@
|
||||
**************************************************************************************************/
|
||||
#pragma once
|
||||
|
||||
#include <cute/config.hpp>
|
||||
|
||||
#include <cute/arch/copy.hpp>
|
||||
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
#include <cute/atom/mma_atom.hpp>
|
||||
|
||||
#include <cute/util/type_traits.hpp>
|
||||
|
||||
#include <cute/tensor_impl.hpp>
|
||||
#include <cute/config.hpp> // CUTE_HOST_DEVICE
|
||||
#include <cute/tensor_impl.hpp> // cute::Tensor
|
||||
#include <cute/util/type_traits.hpp> // cute::__CUTE_REQUIRES
|
||||
#include <cute/container/tuple.hpp> // cute::is_tuple
|
||||
#include <cute/numeric/integral_constant.hpp> // cute::is_constant, cute::is_integral
|
||||
#include <cute/atom/copy_traits.hpp> // cute::Copy_Traits
|
||||
#include <cute/atom/mma_atom.hpp> // cute::TiledMMA
|
||||
|
||||
namespace cute
|
||||
{
|
||||
@@ -651,10 +648,12 @@ print(ThrCopy<TiledCopy, ThrIdx> const& thr_copy)
|
||||
print(TiledCopy{});
|
||||
}
|
||||
|
||||
template <class... Args>
|
||||
// TiledCopy to LaTeX TikZ
|
||||
template <class... Args, class TikzColorFn = TikzColor_TV>
|
||||
CUTE_HOST_DEVICE
|
||||
auto
|
||||
print_latex(TiledCopy<Args...> const& copy)
|
||||
print_latex(TiledCopy<Args...> const& copy,
|
||||
TikzColorFn color = {}) // lambda(thr_idx,val_idx) -> tikz color string
|
||||
{
|
||||
auto [layoutS_MN, thrID_S] = copy.get_layoutS_MN();
|
||||
auto [layoutD_MN, thrID_D] = copy.get_layoutD_MN();
|
||||
@@ -663,13 +662,15 @@ print_latex(TiledCopy<Args...> const& copy)
|
||||
layoutD_MN, thrID_D);
|
||||
}
|
||||
|
||||
// MNK Copy Layout to Latex TIKZ -- 8-value color coded by thread
|
||||
// MNK Copy Layout to LaTeX TikZ
|
||||
template <class LayoutS, class ThrIDS,
|
||||
class LayoutD, class ThrIDD>
|
||||
class LayoutD, class ThrIDD,
|
||||
class TikzColorFn = TikzColor_TV>
|
||||
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
|
||||
LayoutD const& D, ThrIDD const& TD, // (m,n) -> (tid,vid) and tid -> thr_idx
|
||||
TikzColorFn color = {}) // lambda(thr_idx,val_idx) -> tikz color string
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(S) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(D) == Int<2>{});
|
||||
@@ -677,33 +678,17 @@ print_latex_copy(LayoutS const& S, ThrIDS const& TS, // (m,n) -> (tid,vid) and
|
||||
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
|
||||
// Commented prints
|
||||
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);
|
||||
// Header
|
||||
printf("\\documentclass[convert]{standalone}\n"
|
||||
"\\usepackage{tikz}\n\n"
|
||||
"\\begin{document}\n"
|
||||
"\\begin{tikzpicture}[x={(0cm,-1cm)},y={(1cm,0cm)},every node/.style={minimum size=1cm, outer sep=0pt}]\n\n");
|
||||
|
||||
// S starting at 0,0
|
||||
for (int i = 0; i < size<0>(S); ++i) {
|
||||
@@ -712,12 +697,22 @@ print_latex_copy(LayoutS const& S, ThrIDS const& TS, // (m,n) -> (tid,vid) and
|
||||
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],
|
||||
printf("\\node[fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color(thr_idx, val_idx),
|
||||
i, j,
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
// Grid
|
||||
printf("\\draw[color=black,thick,shift={(-0.5,-0.5)}] (%d,%d) grid (%d,%d);\n\n",
|
||||
0, 0, int(size<0>(S)), int(size<1>(S)));
|
||||
// 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 i = -1, j = 0; j < size<1>(S); ++j) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", i, j, j);
|
||||
}
|
||||
|
||||
// D starting at 0,size<1>(S)+3
|
||||
for (int i = 0; i < size<0>(D); ++i) {
|
||||
@@ -726,30 +721,26 @@ print_latex_copy(LayoutS const& S, ThrIDS const& TS, // (m,n) -> (tid,vid) and
|
||||
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],
|
||||
printf("\\node[fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color(thr_idx, val_idx),
|
||||
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);
|
||||
}
|
||||
// Grid
|
||||
printf("\\draw[color=black,thick,shift={(-0.5,-0.5)}] (%d,%d) grid (%d,%d);\n\n",
|
||||
0, int(size<1>(S)+3), int(size<0>(D)), int(size<1>(D)+size<1>(S)+3));
|
||||
// D Labels
|
||||
for (int i = 0, j = size<1>(D); i < size<0>(S); ++i) {
|
||||
for (int i = 0, j = size<1>(D); i < size<0>(D); ++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) {
|
||||
for (int i = -1, j = 0; 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);
|
||||
printf("\\end{tikzpicture}\n"
|
||||
"\\end{document}\n");
|
||||
}
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
@@ -39,7 +39,7 @@ namespace cute
|
||||
{
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM50_Shuffle_U32_2x2Trans>
|
||||
struct Copy_Traits<SM50_Shuffle_U32_2x2Trans_XOR1>
|
||||
{
|
||||
// Logical thread id to thread idx (one-thread)
|
||||
using ThrID = Layout<_32>;
|
||||
@@ -55,4 +55,21 @@ struct Copy_Traits<SM50_Shuffle_U32_2x2Trans>
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct Copy_Traits<SM50_Shuffle_U32_2x2Trans_XOR4>
|
||||
{
|
||||
// Logical thread id to thread idx (one-thread)
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape <_32,_64>,
|
||||
Stride<_64, _1>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape <Shape < _4, _2, _4>, Shape<_32, _2>>,
|
||||
Stride<Stride<_64, _32, _512>,Stride< _1, _256>>>;
|
||||
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
};
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
@@ -450,7 +450,9 @@ make_im2col_tma_copy_desc(
|
||||
CUtensorMapInterleave tma_interleave = CU_TENSOR_MAP_INTERLEAVE_NONE;
|
||||
CUtensorMapL2promotion tma_l2Promotion = to_CUtensorMapL2promotion(aux_params.l2promo_);
|
||||
CUtensorMapFloatOOBfill tma_oob_fill = to_CUtensorMapFloatOOBfill(aux_params.oobfill_);
|
||||
CUtensorMapSwizzle tma_swizzle = TMA::to_CUtensorMapSwizzle(detail::get_tma_swizzle_bits(smem_swizzle));
|
||||
TMA::SmemSwizzleBits swizzle_bits = detail::get_tma_swizzle_bits(smem_swizzle);
|
||||
TMA::SmemSwizzleBase swizzle_base = detail::get_tma_swizzle_base(smem_swizzle);
|
||||
CUtensorMapSwizzle tma_swizzle = TMA::to_CUtensorMapSwizzle(swizzle_bits, swizzle_base);
|
||||
|
||||
CUresult encode_result = CUTLASS_CUDA_DRIVER_WRAPPER_CALL(cuTensorMapEncodeIm2col)(
|
||||
&tma_desc,
|
||||
@@ -636,11 +638,11 @@ make_tma_atom_im2col(CopyOp,
|
||||
|
||||
auto range_c = size<0,0>(tma_layout_vt);
|
||||
auto range_whdn = size<0,1>(tma_layout_vt);
|
||||
|
||||
Tensor gtensor_cwhdn = make_tensor(gtensor.data(),
|
||||
flatten(make_layout(basis_get(stride<0,0>(tma_layout_vt), gtensor.layout()),
|
||||
basis_get(stride<0,1>(tma_layout_vt), gtensor.layout()))));
|
||||
|
||||
flatten(make_layout(make_layout(basis_get(stride<0,0>(tma_layout_vt), gtensor.shape()),
|
||||
basis_get(stride<0,0>(tma_layout_vt), gtensor.stride())),
|
||||
make_layout(basis_get(stride<0,1>(tma_layout_vt), gtensor.shape()),
|
||||
basis_get(stride<0,1>(tma_layout_vt), gtensor.stride())))));
|
||||
auto [tma_desc, tma_tensor] = make_im2col_tma_copy_desc(
|
||||
gtensor_cwhdn,
|
||||
range_c,
|
||||
|
||||
@@ -41,6 +41,7 @@
|
||||
#include <cute/algorithm/prefetch.hpp>
|
||||
|
||||
#include <cute/numeric/integral_ratio.hpp>
|
||||
|
||||
#include <cutlass/cuda_host_adapter.hpp>
|
||||
|
||||
namespace cute
|
||||
@@ -241,15 +242,22 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBitsPerTMA, AuxParams_>
|
||||
// Construct an executable SM90_TMA_LOAD_MULTICAST with tma_mbar
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBitsPerTMA>
|
||||
with(uint64_t& tma_load_mbar, uint16_t const& multicast_mask) const {
|
||||
return {{}, {&tma_desc_, &tma_load_mbar, multicast_mask}};
|
||||
with(
|
||||
uint64_t& tma_load_mbar,
|
||||
uint16_t const& multicast_mask,
|
||||
TMA::CacheHintSm90 const& cache_hint = TMA::CacheHintSm90::EVICT_NORMAL) const {
|
||||
return {{}, {&tma_desc_, &tma_load_mbar, multicast_mask, static_cast<uint64_t>(cache_hint)}};
|
||||
}
|
||||
|
||||
// Construct an executable SM90_TMA_LOAD_MULTICAST_OP with tma_mbar (temp. overloaded for grouped gemm/ptr array gemm)
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBitsPerTMA>
|
||||
with(TmaDescriptor const* new_tma_desc, uint64_t& tma_load_mbar, uint16_t const& multicast_mask) const {
|
||||
return {{}, {new_tma_desc, &tma_load_mbar, multicast_mask}};
|
||||
with(
|
||||
TmaDescriptor const* new_tma_desc,
|
||||
uint64_t& tma_load_mbar,
|
||||
uint16_t const& multicast_mask,
|
||||
TMA::CacheHintSm90 const& cache_hint = TMA::CacheHintSm90::EVICT_NORMAL) const {
|
||||
return {{}, {new_tma_desc, &tma_load_mbar, multicast_mask, static_cast<uint64_t>(cache_hint)}};
|
||||
}
|
||||
|
||||
// Generate the TMA coord tensor
|
||||
@@ -287,7 +295,8 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBitsPerTMA>
|
||||
tuple<
|
||||
TmaDescriptor const*,
|
||||
uint64_t*, // smem mbarrier
|
||||
uint16_t // multicast mask
|
||||
uint16_t, // multicast mask
|
||||
uint64_t // cache hint
|
||||
> const opargs_;
|
||||
};
|
||||
|
||||
@@ -684,8 +693,10 @@ construct_tma_gbasis(Tensor<GEngine,GLayout> const& gtensor, // The origin
|
||||
// TMA parameter checking
|
||||
//
|
||||
|
||||
CUTE_STATIC_ASSERT_V(product_each(shape(slayout)) == product_each(shape(cta_v_map)),
|
||||
"TMA requires CTA_Tile and SLayout top-level shape equivalence.");
|
||||
// CUTE_STATIC_ASSERT_V(product_each(shape(slayout)) == product_each(shape(cta_v_map)),
|
||||
// "TMA requires CTA_Tile and SLayout top-level shape equivalence.");
|
||||
CUTE_STATIC_ASSERT_V(size(slayout) == size(cta_v_map),
|
||||
"TMA requires CTA_Tile and SLayout top-level size equivalence.");
|
||||
|
||||
#if 0
|
||||
print("gtensor : "); print(gtensor); print("\n");
|
||||
@@ -983,7 +994,9 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The origin
|
||||
CUtensorMapFloatOOBfill tma_oobFill = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE;
|
||||
|
||||
// TMA smem swizzle type
|
||||
CUtensorMapSwizzle smem_swizzle = TMA::to_CUtensorMapSwizzle(get_tma_swizzle_bits(swizzle));
|
||||
TMA::SmemSwizzleBits swizzle_bits = get_tma_swizzle_bits(swizzle);
|
||||
TMA::SmemSwizzleBase swizzle_base = get_tma_swizzle_base(swizzle);
|
||||
CUtensorMapSwizzle smem_swizzle = TMA::to_CUtensorMapSwizzle(swizzle_bits, swizzle_base);
|
||||
CUresult result = CUTLASS_CUDA_DRIVER_WRAPPER_CALL(cuTensorMapEncodeTiled)(
|
||||
&tma_desc,
|
||||
tma_format,
|
||||
|
||||
@@ -68,4 +68,26 @@ get_tma_swizzle_bits(Layout const& layout)
|
||||
return get_tma_swizzle_bits(get_swizzle_portion(layout));
|
||||
}
|
||||
|
||||
template <int B, int M, int S>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
TMA::SmemSwizzleBase
|
||||
get_tma_swizzle_base(Swizzle<B,M,S>)
|
||||
{
|
||||
if constexpr (M == 4) {
|
||||
static_assert(0 <= B && B <= 3, "Expected B = 0,1,2, or 3 when M == 4. Unsupported layout swizzle.");
|
||||
static_assert(S == 3, "Expected S = 3 when M == 4. Unsupported layout swizzle.");
|
||||
return TMA::SmemSwizzleBase::SWIZZLE_BASE_16B;
|
||||
}
|
||||
else {
|
||||
static_assert(M == 4, "Expected 128b=16B=(2^4)B base swizzle.");
|
||||
}
|
||||
}
|
||||
|
||||
template <class Layout>
|
||||
TMA::SmemSwizzleBase
|
||||
get_tma_swizzle_base(Layout const& layout)
|
||||
{
|
||||
return get_tma_swizzle_base(get_swizzle_portion(layout));
|
||||
}
|
||||
|
||||
} // namespace cute::detail
|
||||
|
||||
+117
-118
@@ -45,11 +45,12 @@ template <class MMAOperation>
|
||||
struct MMA_Atom<MMAOperation> : MMA_Atom<MMA_Traits<MMAOperation>>
|
||||
{};
|
||||
|
||||
template <class... Args>
|
||||
struct MMA_Atom<MMA_Traits<Args...>>
|
||||
: MMA_Traits<Args...>
|
||||
template <class MMAOperation, class... Args>
|
||||
struct MMA_Atom<MMA_Traits<MMAOperation, Args...>>
|
||||
: MMA_Traits<MMAOperation, Args...>
|
||||
{
|
||||
using Traits = MMA_Traits<Args...>;
|
||||
using MMA_Op = MMAOperation;
|
||||
using Traits = MMA_Traits<MMAOperation, Args...>;
|
||||
|
||||
// Element value types from the MMA_Traits
|
||||
using ValTypeD = typename Traits::ValTypeD;
|
||||
@@ -331,7 +332,7 @@ struct TiledMMA : MMA_Atom
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto b_tensor = zipped_divide(t_tensor, b_tile); // ((AtomN,AtomK),(RestN,RestK))
|
||||
|
||||
// Transform the Atom mode from (N,K) to (Thr,Val)
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = b_tensor.compose(AtomLayoutB_TV{},_); // ((ThrV,FrgV),(RestN,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
@@ -733,18 +734,22 @@ print(ThrMMA<TiledMMA, ThrVMNK> const& thr_mma)
|
||||
print(static_cast<TiledMMA>(thr_mma));
|
||||
}
|
||||
|
||||
template <class... Args>
|
||||
// MMA Atom to LaTeX TikZ
|
||||
template <class... Args, class TikzColorFn = TikzColor_TV>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print_latex(MMA_Atom<Args...> const& mma_atom)
|
||||
print_latex(MMA_Atom<Args...> const& mma_atom,
|
||||
TikzColorFn color = {}) // lambda(thr_idx,val_idx) -> tikz color string
|
||||
{
|
||||
print_latex(make_tiled_mma(mma_atom));
|
||||
}
|
||||
|
||||
template <class... Args>
|
||||
// TiledMMA to LaTeX TikZ
|
||||
template <class... Args, class TikzColorFn = TikzColor_TV>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print_latex(TiledMMA<Args...> const& mma)
|
||||
print_latex(TiledMMA<Args...> const& mma,
|
||||
TikzColorFn color = {}) // lambda(thr_idx,val_idx) -> tikz color string
|
||||
{
|
||||
auto layout_and_thrid_C = mma.get_layoutC_MN();
|
||||
auto layoutC_MN = get<0>(layout_and_thrid_C);
|
||||
@@ -763,6 +768,109 @@ print_latex(TiledMMA<Args...> const& mma)
|
||||
layoutB_NK, thrID_B);
|
||||
}
|
||||
|
||||
// MNK MMA Layout to LaTeX TikZ
|
||||
template <class LayoutC, class ThrIDC,
|
||||
class LayoutA, class ThrIDA,
|
||||
class LayoutB, class ThrIDB,
|
||||
class TikzColorFn = TikzColor_TV>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print_latex_mma(LayoutC const& C, ThrIDC const& TC, // (m,n) -> (tid,vid) and tid -> thr_idx
|
||||
LayoutA const& A, ThrIDA const& TA, // (m,k) -> (tid,vid) and tid -> thr_idx
|
||||
LayoutB const& B, ThrIDB const& TB, // (n,k) -> (tid,vid) and tid -> thr_idx
|
||||
TikzColorFn color = {}) // lambda(thr_idx,val_idx) -> tikz color string
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(C) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(A) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(B) == Int<2>{});
|
||||
|
||||
assert(size<0>(A) == size<0>(C));
|
||||
assert(size<0>(B) == size<1>(C));
|
||||
assert(size<1>(A) == size<1>(B));
|
||||
|
||||
// Commented prints
|
||||
printf("%% LayoutC: "); print(C); printf("\n");
|
||||
printf("%% ThrIDC : "); print(TC); printf("\n");
|
||||
printf("%% LayoutA: "); print(A); printf("\n");
|
||||
printf("%% ThrIDA : "); print(TA); printf("\n");
|
||||
printf("%% LayoutB: "); print(B); printf("\n");
|
||||
printf("%% ThrIDB : "); print(TB); printf("\n\n");
|
||||
// Header
|
||||
printf("\\documentclass[convert]{standalone}\n"
|
||||
"\\usepackage{tikz}\n\n"
|
||||
"\\begin{document}\n"
|
||||
"\\begin{tikzpicture}[x={(0cm,-1cm)},y={(1cm,0cm)},every node/.style={minimum size=1cm, outer sep=0pt}]\n\n");
|
||||
|
||||
// C starting at 0,0
|
||||
for (int m = 0; m < size<0>(C); ++m) {
|
||||
for (int n = 0; n < size<1>(C); ++n) {
|
||||
int thrid = C(m,n) % size(TC);
|
||||
int val_idx = C(m,n) / size(TC);
|
||||
int thr_idx = TC(thrid);
|
||||
|
||||
printf("\\node[fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color(thr_idx, val_idx),
|
||||
m, n,
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
// Grid
|
||||
printf("\\draw[color=black,thick,shift={(-0.5,-0.5)}] (%d,%d) grid (%d,%d);\n\n",
|
||||
0, 0, int(size<0>(C)), int(size<1>(C)));
|
||||
|
||||
// A starting at 0,-size<1>(A)-1
|
||||
for (int m = 0; m < size<0>(A); ++m) {
|
||||
for (int k = 0; k < size<1>(A); ++k) {
|
||||
int thrid = A(m,k) % size(TA);
|
||||
int val_idx = A(m,k) / size(TA);
|
||||
int thr_idx = TA(thrid);
|
||||
|
||||
printf("\\node[fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color(thr_idx, val_idx),
|
||||
m, k-1-size<1>(A),
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
// Grid
|
||||
printf("\\draw[color=black,thick,shift={(-0.5,-0.5)}] (%d,%d) grid (%d,%d);\n\n",
|
||||
0, int(-size<1>(A)-1), int(size<0>(A)), -1);
|
||||
// A labels
|
||||
for (int m = 0, k = -1; m < size<0>(A); ++m) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", m, k-1-size<1>(A), m);
|
||||
}
|
||||
for (int m = -1, k = 0; k < size<1>(A); ++k) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", m, k-1-size<1>(A), k);
|
||||
}
|
||||
|
||||
// B starting at -size<1>(B)-1,0
|
||||
for (int n = 0; n < size<0>(B); ++n) {
|
||||
for (int k = 0; k < size<1>(B); ++k) {
|
||||
int thrid = B(n,k) % size(TB);
|
||||
int val_idx = B(n,k) / size(TB);
|
||||
int thr_idx = TB(thrid);
|
||||
|
||||
printf("\\node[fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color(thr_idx, val_idx),
|
||||
k-1-size<1>(B), n,
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
// Grid
|
||||
printf("\\draw[color=black,thick,shift={(-0.5,-0.5)}] (%d,%d) grid (%d,%d);\n\n",
|
||||
int(-size<1>(B)-1), 0, -1, int(size<0>(B)));
|
||||
// B labels
|
||||
for (int n = 0, k = -1; n < size<0>(B); ++n) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", k-1-size<1>(B), n, n);
|
||||
}
|
||||
for (int n = -1, k = 0; k < size<1>(B); ++k) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", k-1-size<1>(B), n, k);
|
||||
}
|
||||
|
||||
// Footer
|
||||
printf("\\end{tikzpicture}\n"
|
||||
"\\end{document}\n");
|
||||
}
|
||||
|
||||
// MNK MMA Layout to console printer
|
||||
template <class LayoutC, class ThrIDC,
|
||||
class LayoutA, class ThrIDA,
|
||||
@@ -819,115 +927,6 @@ print_layout_mma(LayoutC const& C, ThrIDC const& TC, // (m,n) -> (tid,vid) and
|
||||
printf("+\n");
|
||||
}
|
||||
|
||||
// MNK MMA Layout to Latex TIKZ -- 8-value color coded by thread
|
||||
template <class LayoutC, class ThrIDC,
|
||||
class LayoutA, class ThrIDA,
|
||||
class LayoutB, class ThrIDB>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print_latex_mma(LayoutC const& C, ThrIDC const& TC, // (m,n) -> (tid,vid) and tid -> thr_idx
|
||||
LayoutA const& A, ThrIDA const& TA, // (m,k) -> (tid,vid) and tid -> thr_idx
|
||||
LayoutB const& B, ThrIDB const& TB) // (n,k) -> (tid,vid) and tid -> thr_idx
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(C) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(A) == Int<2>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(B) == Int<2>{});
|
||||
|
||||
assert(size<0>(A) == size<0>(C));
|
||||
assert(size<0>(B) == size<1>(C));
|
||||
assert(size<1>(A) == size<1>(B));
|
||||
|
||||
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("%% LayoutC: "); print(C); printf("\n");
|
||||
printf("%% ThrIDC : "); print(TC); printf("\n");
|
||||
printf("%% LayoutA: "); print(A); printf("\n");
|
||||
printf("%% ThrIDA : "); print(TA); printf("\n");
|
||||
printf("%% LayoutB: "); print(B); printf("\n");
|
||||
printf("%% ThrIDB : "); print(TB); printf("\n\n");
|
||||
|
||||
printf(latex_header);
|
||||
|
||||
// C starting at 0,0
|
||||
for (int m = 0; m < size<0>(C); ++m) {
|
||||
for (int n = 0; n < size<1>(C); ++n) {
|
||||
int thrid = C(m,n) % size(TC);
|
||||
int val_idx = C(m,n) / size(TC);
|
||||
int thr_idx = TC(thrid);
|
||||
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[thr_idx % 8],
|
||||
m, n,
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
|
||||
// A starting at 0,-size<1>(A)-1
|
||||
for (int m = 0; m < size<0>(A); ++m) {
|
||||
for (int k = 0; k < size<1>(A); ++k) {
|
||||
int thrid = A(m,k) % size(TA);
|
||||
int val_idx = A(m,k) / size(TA);
|
||||
int thr_idx = TA(thrid);
|
||||
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[thr_idx % 8],
|
||||
m, k-1-size<1>(A),
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
|
||||
// B starting at -size<1>(B)-1,0
|
||||
for (int n = 0; n < size<0>(B); ++n) {
|
||||
for (int k = 0; k < size<1>(B); ++k) {
|
||||
int thrid = B(n,k) % size(TB);
|
||||
int val_idx = B(n,k) / size(TB);
|
||||
int thr_idx = TB(thrid);
|
||||
|
||||
printf("\\node[box,fill=%s] at (%d,%d) {\\shortstack{T%d \\\\ V%d}};\n",
|
||||
color_map[thr_idx % 8],
|
||||
k-1-size<1>(B), n,
|
||||
thr_idx, val_idx);
|
||||
}
|
||||
}
|
||||
|
||||
// A labels
|
||||
for (int m = 0, k = -1; m < size<0>(A); ++m) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", m, k-1-size<1>(A), m);
|
||||
}
|
||||
for (int k = 0, m = -1; k < size<1>(A); ++k) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", m, k-1-size<1>(A), k);
|
||||
}
|
||||
// B labels
|
||||
for (int n = 0, k = -1; n < size<0>(B); ++n) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", k-1-size<1>(B), n, n);
|
||||
}
|
||||
for (int k = 0, n = -1; k < size<1>(B); ++k) {
|
||||
printf("\\node at (%d,%d) {\\Large{\\texttt{%d}}};\n", k-1-size<1>(B), n, k);
|
||||
}
|
||||
|
||||
// Footer
|
||||
printf(latex_footer);
|
||||
}
|
||||
|
||||
// MNK MMA Layout to SVG -- 8-value color coded by thread
|
||||
template <class LayoutC, class ThrIDC,
|
||||
class LayoutA, class ThrIDA,
|
||||
|
||||
@@ -30,23 +30,14 @@
|
||||
**************************************************************************************************/
|
||||
#pragma once
|
||||
|
||||
#include <cute/arch/mma.hpp>
|
||||
|
||||
#include <cute/tensor_impl.hpp>
|
||||
#include <cute/tensor_impl.hpp> // cute::Tensor
|
||||
#include <cute/pointer.hpp> // cute::is_rmem
|
||||
#include <cute/arch/mma.hpp> // cute::UniversalFMA
|
||||
#include <cute/arch/util.hpp> // cute::detail::explode
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <class X, class = void>
|
||||
struct supports_output_scaling { static constexpr bool value = false; };
|
||||
|
||||
template <class X>
|
||||
struct supports_output_scaling<X, void_t<decltype(declval<X>().accumulate_)>> { static constexpr bool value = true; };
|
||||
|
||||
} // end namespace detail
|
||||
|
||||
/**
|
||||
* concept MMA_Traits
|
||||
* {
|
||||
@@ -99,17 +90,27 @@ struct MMA_Traits<UniversalFMA<D,A,B,C>>
|
||||
using CLayout = Layout<Shape<_1,_1>>;
|
||||
};
|
||||
|
||||
// Extract an MMA_Op from an MMA_Traits
|
||||
template <class MMA_Traits>
|
||||
struct MMA_Op {};
|
||||
|
||||
template <class MMA_Op_Arg, class... Args>
|
||||
struct MMA_Op<MMA_Traits<MMA_Op_Arg, Args...>> {
|
||||
using type = MMA_Op_Arg;
|
||||
};
|
||||
|
||||
//
|
||||
// Generic mma_unpack for any MMA_Traits
|
||||
//
|
||||
template <class MMA_Op, class... MMA_Args,
|
||||
|
||||
template <class AnyMMATraits,
|
||||
class TD, class DLayout,
|
||||
class TA, class ALayout,
|
||||
class TB, class BLayout,
|
||||
class TC, class CLayout>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
mma_unpack(MMA_Traits<MMA_Op, MMA_Args...> const& traits,
|
||||
mma_unpack(AnyMMATraits const& traits,
|
||||
Tensor<TD, DLayout> & D,
|
||||
Tensor<TA, ALayout> const& A,
|
||||
Tensor<TB, BLayout> const& B,
|
||||
@@ -121,87 +122,47 @@ mma_unpack(MMA_Traits<MMA_Op, MMA_Args...> const& traits,
|
||||
static_assert(is_rmem<TC>::value, "Expected registers in MMA_Atom::call");
|
||||
|
||||
// Register value types from the MMA_Operation register arrays
|
||||
using MMA_Op = typename MMA_Op<AnyMMATraits>::type;
|
||||
using RegTypeD = typename remove_extent<typename MMA_Op::DRegisters>::type;
|
||||
using RegTypeA = typename remove_extent<typename MMA_Op::ARegisters>::type;
|
||||
using RegTypeB = typename remove_extent<typename MMA_Op::BRegisters>::type;
|
||||
using RegTypeC = typename remove_extent<typename MMA_Op::CRegisters>::type;
|
||||
using MMATraits = MMA_Traits<MMA_Op, MMA_Args...>;
|
||||
|
||||
[[maybe_unused]] constexpr int RegNumD = extent<typename MMA_Op::DRegisters>::value;
|
||||
Tensor rA = recast<RegTypeA>(A);
|
||||
Tensor rB = recast<RegTypeB>(B);
|
||||
Tensor rD = recast<RegTypeD>(D);
|
||||
Tensor rC = recast<RegTypeC>(C);
|
||||
|
||||
constexpr int RegNumD = extent<typename MMA_Op::DRegisters>::value;
|
||||
constexpr int RegNumA = extent<typename MMA_Op::ARegisters>::value;
|
||||
constexpr int RegNumB = extent<typename MMA_Op::BRegisters>::value;
|
||||
constexpr int RegNumC = extent<typename MMA_Op::CRegisters>::value;
|
||||
|
||||
Tensor rA = recast<RegTypeA>(A);
|
||||
Tensor rB = recast<RegTypeB>(B);
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size(rA) == Int<RegNumA>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rB) == Int<RegNumB>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumD>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rC) == Int<RegNumC>{});
|
||||
|
||||
if constexpr (is_same<RegTypeD, void>::value)
|
||||
{
|
||||
static_assert(is_same<typename TD::value_type, typename TC::value_type>::value, "GMMA C and D value_type must match.");
|
||||
static_assert(is_same<DLayout, CLayout>::value, "GMMA C and D layouts must match.");
|
||||
// assert((void*)&C == (void*)&D);
|
||||
|
||||
Tensor rC = recast<RegTypeC>(D); // NOTE: D and C are same, so use mutable D
|
||||
|
||||
//CUTE_STATIC_ASSERT_V(size(rC) == Int<RegNumC>{});
|
||||
|
||||
if constexpr (detail::supports_output_scaling<MMATraits>::value) {
|
||||
detail::explode(MMA_Op::fma,
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{},
|
||||
&(traits.accumulate_), seq<0>{});
|
||||
}
|
||||
else {
|
||||
detail::explode(MMA_Op::fma,
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{});
|
||||
}
|
||||
}
|
||||
else {
|
||||
Tensor rD = recast<RegTypeD>(D);
|
||||
Tensor rC = recast<RegTypeC>(C);
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumD>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rC) == Int<RegNumC>{});
|
||||
if constexpr (detail::supports_output_scaling<MMATraits>::value) {
|
||||
detail::explode(MMA_Op::fma,
|
||||
rD, make_int_sequence<RegNumD>{},
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{},
|
||||
&(traits.accumulate_), seq<0>{});
|
||||
}
|
||||
else {
|
||||
detail::explode(MMA_Op::fma,
|
||||
rD, make_int_sequence<RegNumD>{},
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{});
|
||||
}
|
||||
}
|
||||
detail::explode(MMA_Op::fma,
|
||||
rD, make_int_sequence<RegNumD>{},
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{});
|
||||
}
|
||||
|
||||
//
|
||||
// Accept mutable temporaries
|
||||
//
|
||||
|
||||
template <class MMA_Op, class... MMA_Args,
|
||||
template <class AnyMMATraits,
|
||||
class TD, class DLayout,
|
||||
class TA, class ALayout,
|
||||
class TB, class BLayout,
|
||||
class TC, class CLayout>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
mma_unpack(MMA_Traits<MMA_Op, MMA_Args...> const& traits,
|
||||
Tensor<TD, DLayout> && D,
|
||||
Tensor<TA, ALayout> const& A,
|
||||
Tensor<TB, BLayout> const& B,
|
||||
Tensor<TC, CLayout> const& C)
|
||||
mma_unpack(AnyMMATraits const& traits,
|
||||
Tensor<TD, DLayout> && D,
|
||||
Tensor<TA, ALayout> const& A,
|
||||
Tensor<TB, BLayout> const& B,
|
||||
Tensor<TC, CLayout> const& C)
|
||||
{
|
||||
mma_unpack(traits, D, A, B, C);
|
||||
}
|
||||
|
||||
@@ -41,6 +41,8 @@ namespace cute {
|
||||
//////////////////////// fp64 = fp64 * fp64 + fp64 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
using SM90_16x8x4_F64F64F64F64_TN = SM90::MMA_16x8x4_F64F64F64F64_TN;
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
|
||||
{
|
||||
@@ -59,6 +61,8 @@ struct MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
};
|
||||
|
||||
using SM90_16x8x8_F64F64F64F64_TN = SM90::MMA_16x8x8_F64F64F64F64_TN;
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
|
||||
{
|
||||
@@ -77,6 +81,8 @@ struct MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
};
|
||||
|
||||
using SM90_16x8x16_F64F64F64F64_TN = SM90::MMA_16x8x16_F64F64F64F64_TN;
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
|
||||
{
|
||||
@@ -99,9 +105,11 @@ struct MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
|
||||
//////////////////////// cfp64 = cfp64 * cfp64 + cfp64 ////////////////////////////
|
||||
///////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
using SM90_16x8x4_C64C64C64C64_TN = SM90::MMA_16x8x4_C64C64C64C64_TN;
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x4_C64C64C64C64_TN>
|
||||
: MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
|
||||
: MMA_Traits<SM90_16x8x4_F64F64F64F64_TN>
|
||||
{
|
||||
using ValTypeD = complex<double>;
|
||||
using ValTypeA = complex<double>;
|
||||
@@ -109,9 +117,11 @@ struct MMA_Traits<SM90_16x8x4_C64C64C64C64_TN>
|
||||
using ValTypeC = complex<double>;
|
||||
};
|
||||
|
||||
using SM90_16x8x8_C64C64C64C64_TN = SM90::MMA_16x8x8_C64C64C64C64_TN;
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x8_C64C64C64C64_TN>
|
||||
: MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
|
||||
: MMA_Traits<SM90_16x8x8_F64F64F64F64_TN>
|
||||
{
|
||||
using ValTypeD = complex<double>;
|
||||
using ValTypeA = complex<double>;
|
||||
@@ -119,9 +129,11 @@ struct MMA_Traits<SM90_16x8x8_C64C64C64C64_TN>
|
||||
using ValTypeC = complex<double>;
|
||||
};
|
||||
|
||||
using SM90_16x8x16_C64C64C64C64_TN = SM90::MMA_16x8x16_C64C64C64C64_TN;
|
||||
|
||||
template <>
|
||||
struct MMA_Traits<SM90_16x8x16_C64C64C64C64_TN>
|
||||
: MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
|
||||
: MMA_Traits<SM90_16x8x16_F64F64F64F64_TN>
|
||||
{
|
||||
using ValTypeD = complex<double>;
|
||||
using ValTypeA = complex<double>;
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user