@@ -68,7 +68,7 @@ struct UniversalCopy
|
||||
|
||||
//
|
||||
// Placeholder for the copy algorithm's stronger auto-vectorizing behavior
|
||||
// that assumes alignment of dynamic layouts up to MaxVecBits
|
||||
// that assumes alignment of pointers and dynamic layouts up to MaxVecBits
|
||||
//
|
||||
|
||||
template <int MaxVecBits = 128>
|
||||
@@ -80,15 +80,17 @@ struct AutoVectorizingCopyWithAssumedAlignment
|
||||
};
|
||||
|
||||
//
|
||||
// Placeholder for the copy algorithm's default auto-vectorizing behavior
|
||||
// that does not assume alignment of dynamic layouts
|
||||
// AutoVectorizingCopy alias assumes maximal alignment of pointers and dynamic strides.
|
||||
// If this is not the case then AutoVectorizingCopyWithAssumedAlignment should be used instead
|
||||
//
|
||||
|
||||
using AutoVectorizingCopy = AutoVectorizingCopyWithAssumedAlignment<8>;
|
||||
using AutoVectorizingCopy = AutoVectorizingCopyWithAssumedAlignment<128>;
|
||||
|
||||
// Alias
|
||||
using DefaultCopy = AutoVectorizingCopy;
|
||||
//
|
||||
// DefaultCopy alias does not assume alignment of pointers or dynamic strides.
|
||||
//
|
||||
|
||||
using DefaultCopy = AutoVectorizingCopyWithAssumedAlignment<8>;
|
||||
|
||||
//
|
||||
// Global memory prefetch into L2
|
||||
|
||||
@@ -95,8 +95,8 @@ wait_barrier(uint64_t& smem_barrier, // 64 bits user-mange
|
||||
".reg .pred P1;\n"
|
||||
"LAB_WAIT:\n"
|
||||
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
|
||||
"@P1 bra.uni DONE;\n"
|
||||
"bra.uni LAB_WAIT;\n"
|
||||
"@P1 bra DONE;\n"
|
||||
"bra LAB_WAIT;\n"
|
||||
"DONE:\n"
|
||||
"}\n"
|
||||
:: "r"(smem_int_ptr),
|
||||
@@ -134,6 +134,48 @@ enum class SmemSwizzleBits : uint8_t {
|
||||
B128 = 3,
|
||||
};
|
||||
|
||||
enum class OOBFill : uint8_t {
|
||||
ZERO = 0,
|
||||
CONSTANT = 1,
|
||||
};
|
||||
|
||||
CUTE_HOST_DEVICE char const* to_string(OOBFill const& t) {
|
||||
switch (t) {
|
||||
case OOBFill::ZERO: return "ZERO";
|
||||
case OOBFill::CONSTANT: return "CONSTANT";
|
||||
}
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
enum class L2Promotion : uint8_t {
|
||||
DISABLE = 0,
|
||||
B64 = 1,
|
||||
B128 = 2,
|
||||
B256 = 3,
|
||||
};
|
||||
|
||||
CUTE_HOST_DEVICE char const* to_string(L2Promotion const& t) {
|
||||
switch (t) {
|
||||
case L2Promotion::DISABLE: return "DISABLE";
|
||||
case L2Promotion::B64: return "B64";
|
||||
case L2Promotion::B128: return "B128";
|
||||
case L2Promotion::B256: return "B256";
|
||||
}
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
// Aux parameters which are independent with the problem size
|
||||
struct DescriptorAuxParams {
|
||||
OOBFill oobfill_ = OOBFill::ZERO;
|
||||
L2Promotion l2promo_ = L2Promotion::DISABLE;
|
||||
};
|
||||
|
||||
enum class CacheHintSm90 : uint64_t {
|
||||
EVICT_NORMAL = 0x1000000000000000,
|
||||
EVICT_FIRST = 0x12F0000000000000,
|
||||
EVICT_LAST = 0x14F0000000000000,
|
||||
};
|
||||
|
||||
#if (__CUDACC_VER_MAJOR__ >= 12)
|
||||
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
@@ -168,6 +210,27 @@ to_CUtensorMapSwizzle(SmemSwizzleBits const& t) {
|
||||
case SmemSwizzleBits::B128: return CU_TENSOR_MAP_SWIZZLE_128B;
|
||||
}
|
||||
}
|
||||
|
||||
inline CUtensorMapFloatOOBfill
|
||||
to_CUtensorMapFloatOOBfill(OOBFill const& t) {
|
||||
switch(t) {
|
||||
default: assert(false && "Unknown OOBFill!");
|
||||
case OOBFill::ZERO: return CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE;
|
||||
case OOBFill::CONSTANT: return CU_TENSOR_MAP_FLOAT_OOB_FILL_NAN_REQUEST_ZERO_FMA;
|
||||
}
|
||||
}
|
||||
|
||||
inline CUtensorMapL2promotion
|
||||
to_CUtensorMapL2promotion(L2Promotion const& t) {
|
||||
switch(t) {
|
||||
default: assert(false && "Unknown L2Promotion!");
|
||||
case L2Promotion::DISABLE: return CU_TENSOR_MAP_L2_PROMOTION_NONE;
|
||||
case L2Promotion::B64: return CU_TENSOR_MAP_L2_PROMOTION_L2_64B;
|
||||
case L2Promotion::B128: return CU_TENSOR_MAP_L2_PROMOTION_L2_128B;
|
||||
case L2Promotion::B256: return CU_TENSOR_MAP_L2_PROMOTION_L2_256B;
|
||||
}
|
||||
}
|
||||
|
||||
#endif // !defined(__CUDACC_RTC__)
|
||||
|
||||
#endif // (__CUDACC_VER_MAJOR__ >= 12)
|
||||
@@ -257,22 +320,32 @@ tma_descriptor_replace_dims_strides_in_shared_mem(TmaDescriptor
|
||||
asm volatile (
|
||||
"cvt.u64.u32 %0, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(smem_int_desc));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 0, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[0]));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 1, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[1]));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 2, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[2]));
|
||||
// Strides must be a multiple of 16. Also, stride for the intermost dimension is implicitly 1
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 0, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[1] >> 4));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 1, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[2] >> 4));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 0, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[0]));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 1, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[1]));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 2, %1;"
|
||||
:: "l"(smem_int64_desc), "r"(prob_shape[2]));
|
||||
// Strides must be a multiple of 16. Also, stride for the intermost dimension is implicitly 1
|
||||
#if ((__CUDACC_VER_MAJOR__ > 12) || ((__CUDACC_VER_MAJOR__ == 12) && (__CUDACC_VER_MINOR__ >= 5)))
|
||||
// 4 LSBs are not included
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 0, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[1]));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 1, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[2]));
|
||||
#else
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 0, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[1] >> 4));
|
||||
asm volatile (
|
||||
"tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 1, %1;"
|
||||
:: "l"(smem_int64_desc), "l"(prob_stride[2] >> 4));
|
||||
#endif
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Using TMA Descriptor modification without CUTE_ARCH_TMA_SM90_ENABLED and CUDA 12.3");
|
||||
#endif
|
||||
|
||||
@@ -44,7 +44,7 @@ namespace cute
|
||||
struct SM90_TMA_LOAD_1D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0)
|
||||
{
|
||||
@@ -53,11 +53,11 @@ struct SM90_TMA_LOAD_1D
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.1d.shared::cluster.global.mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3}], [%2];"
|
||||
"cp.async.bulk.tensor.1d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3}], [%2], %4;"
|
||||
:
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0)
|
||||
"r"(crd0), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
@@ -89,7 +89,7 @@ struct SM90_TMA_LOAD_1D
|
||||
struct SM90_TMA_LOAD_2D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1)
|
||||
{
|
||||
@@ -98,11 +98,11 @@ struct SM90_TMA_LOAD_2D
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3, %4}], [%2];"
|
||||
"cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4}], [%2], %5;"
|
||||
:
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1)
|
||||
"r"(crd0), "r"(crd1), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
@@ -134,7 +134,7 @@ struct SM90_TMA_LOAD_2D
|
||||
struct SM90_TMA_LOAD_3D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
|
||||
{
|
||||
@@ -143,11 +143,11 @@ struct SM90_TMA_LOAD_3D
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3, %4, %5}], [%2];"
|
||||
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4, %5}], [%2], %6;"
|
||||
:
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1), "r"(crd2)
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
@@ -179,7 +179,7 @@ struct SM90_TMA_LOAD_3D
|
||||
struct SM90_TMA_LOAD_4D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
|
||||
{
|
||||
@@ -188,11 +188,11 @@ struct SM90_TMA_LOAD_4D
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.4d.shared::cluster.global.mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3, %4, %5, %6}], [%2];"
|
||||
"cp.async.bulk.tensor.4d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4, %5, %6}], [%2], %7;"
|
||||
:
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3)
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
@@ -224,7 +224,7 @@ struct SM90_TMA_LOAD_4D
|
||||
struct SM90_TMA_LOAD_5D
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
|
||||
{
|
||||
@@ -233,11 +233,11 @@ struct SM90_TMA_LOAD_5D
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.5d.shared::cluster.global.mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3, %4, %5, %6, %7}], [%2];"
|
||||
"cp.async.bulk.tensor.5d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4, %5, %6, %7}], [%2], %8;"
|
||||
:
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "r"(crd4)
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "r"(crd4), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
@@ -269,39 +269,39 @@ struct SM90_TMA_LOAD_5D
|
||||
struct SM90_TMA_LOAD
|
||||
{
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0)
|
||||
{
|
||||
return SM90_TMA_LOAD_1D::copy(desc_ptr, mbar_ptr, smem_ptr, crd0);
|
||||
return SM90_TMA_LOAD_1D::copy(desc_ptr, mbar_ptr, cache_hint, smem_ptr, crd0);
|
||||
}
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1)
|
||||
{
|
||||
return SM90_TMA_LOAD_2D::copy(desc_ptr, mbar_ptr, smem_ptr, crd0, crd1);
|
||||
return SM90_TMA_LOAD_2D::copy(desc_ptr, mbar_ptr, cache_hint, smem_ptr, crd0, crd1);
|
||||
}
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
|
||||
{
|
||||
return SM90_TMA_LOAD_3D::copy(desc_ptr, mbar_ptr, smem_ptr, crd0, crd1, crd2);
|
||||
return SM90_TMA_LOAD_3D::copy(desc_ptr, mbar_ptr, cache_hint, smem_ptr, crd0, crd1, crd2);
|
||||
}
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
|
||||
{
|
||||
return SM90_TMA_LOAD_4D::copy(desc_ptr, mbar_ptr, smem_ptr, crd0, crd1, crd2, crd3);
|
||||
return SM90_TMA_LOAD_4D::copy(desc_ptr, mbar_ptr, cache_hint, smem_ptr, crd0, crd1, crd2, crd3);
|
||||
}
|
||||
CUTE_HOST_DEVICE static void
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr,
|
||||
copy(void const* desc_ptr, uint64_t* mbar_ptr, uint64_t cache_hint,
|
||||
void * smem_ptr,
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
|
||||
{
|
||||
return SM90_TMA_LOAD_5D::copy(desc_ptr, mbar_ptr, smem_ptr, crd0, crd1, crd2, crd3, crd4);
|
||||
return SM90_TMA_LOAD_5D::copy(desc_ptr, mbar_ptr, cache_hint, smem_ptr, crd0, crd1, crd2, crd3, crd4);
|
||||
}
|
||||
|
||||
struct PREFETCH
|
||||
|
||||
@@ -85,7 +85,6 @@ CUTE_HOST std::ostream& operator<<(std::ostream& os, LayoutType const& t) {
|
||||
|
||||
union GmmaDescriptor
|
||||
{
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
GmmaDescriptor() noexcept : desc_(0) {}
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -135,21 +134,22 @@ union GmmaDescriptor
|
||||
// Decay to a uint64_t
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
operator uint64_t() const noexcept { return desc_; }
|
||||
|
||||
// Printer
|
||||
CUTE_HOST_DEVICE friend void print(GmmaDescriptor const& t)
|
||||
{
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
printf("GmmaDescriptor: 0x%016llx\n", static_cast<unsigned long long>(t.desc_));
|
||||
printf(" start_addr : 0x%04x\n", t.bitfield.start_address_);
|
||||
printf(" leading_off: 0x%04x (%d)\n", t.bitfield.leading_byte_offset_, t.bitfield.leading_byte_offset_);
|
||||
printf(" stride_off : 0x%04x (%d)\n", t.bitfield.stride_byte_offset_, t.bitfield.stride_byte_offset_);
|
||||
printf(" base_offset: 0x%01x\n", t.bitfield.base_offset_);
|
||||
printf(" layout_type: 0x%01x (%s)\n", t.bitfield.layout_type_, to_string(static_cast<GMMA::LayoutType>(t.bitfield.layout_type_)));
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
// Printer
|
||||
CUTE_HOST_DEVICE void
|
||||
print(GmmaDescriptor const& t)
|
||||
{
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
printf("GmmaDescriptor: 0x%016llx\n", static_cast<unsigned long long>(t.desc_));
|
||||
printf(" start_addr : 0x%04x\n", t.bitfield.start_address_);
|
||||
printf(" leading_off: 0x%04x (%d)\n", t.bitfield.leading_byte_offset_, t.bitfield.leading_byte_offset_);
|
||||
printf(" stride_off : 0x%04x (%d)\n", t.bitfield.stride_byte_offset_, t.bitfield.stride_byte_offset_);
|
||||
printf(" base_offset: 0x%01x\n", t.bitfield.base_offset_);
|
||||
printf(" layout_type: 0x%01x (%s)\n", t.bitfield.layout_type_, to_string(static_cast<GMMA::LayoutType>(t.bitfield.layout_type_)));
|
||||
#endif // !defined(__CUDACC_RTC__)
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cute
|
||||
|
||||
+14
-13
@@ -235,24 +235,25 @@ explode(Fn fn,
|
||||
}
|
||||
|
||||
template <class Fn,
|
||||
class PtrD, int... Id,
|
||||
class PtrA, int... Ia,
|
||||
class PtrB, int... Ib,
|
||||
class PtrC, int... Ic,
|
||||
class PtrSFA, int... Isfa,
|
||||
class PtrSFB, int... Isfb>
|
||||
class PtrD, int... Id,
|
||||
class PtrA, int... Ia,
|
||||
class PtrB, int... Ib,
|
||||
class PtrC, int... Ic,
|
||||
class PtrE, int... Ie,
|
||||
class PtrF, int... If>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
void
|
||||
explode(Fn fn,
|
||||
PtrD&& d, int_sequence<Id...>,
|
||||
PtrA&& a, int_sequence<Ia...>,
|
||||
PtrB&& b, int_sequence<Ib...>,
|
||||
PtrC&& c, int_sequence<Ic...>,
|
||||
PtrSFA&& sfa, int_sequence<Isfa...>,
|
||||
PtrSFB&& sfb, int_sequence<Isfb...>)
|
||||
PtrD&& d, int_sequence<Id...>,
|
||||
PtrA&& a, int_sequence<Ia...>,
|
||||
PtrB&& b, int_sequence<Ib...>,
|
||||
PtrC&& c, int_sequence<Ic...>,
|
||||
PtrE&& e, int_sequence<Ie...>,
|
||||
PtrF&& f, int_sequence<If...>)
|
||||
{
|
||||
return fn(d[Id]..., a[Ia]..., b[Ib]..., c[Ic]..., sfa[Isfa]..., sfb[Isfb]...);
|
||||
return fn(d[Id]..., a[Ia]..., b[Ib]..., c[Ic]..., e[Ie]..., f[If]...);
|
||||
}
|
||||
|
||||
//
|
||||
// Utility for exploding tuples into functions
|
||||
//
|
||||
|
||||
Reference in New Issue
Block a user