v3.9 (#2185)
* v3.8 update x * fix blackwell gg * doc change * doc change * doc change --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com> Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com>
This commit is contained in:
co-authored by
yuzhai
Haicheng Wu
Haicheng Wu
parent
8c4d1dc47d
commit
62750a2b75
@@ -32,6 +32,11 @@
|
||||
|
||||
#include <cutlass/arch/config.h> // CUTLASS_ARCH_MMA_SMxx_ENABLED
|
||||
|
||||
// MMA SM90A
|
||||
#if defined(CUTLASS_ARCH_MMA_SM90A_ENABLED)
|
||||
# define CUTE_ARCH_MMA_SM90A_ENABLED
|
||||
#endif
|
||||
|
||||
// TMA instructions
|
||||
#if defined(CUTLASS_ARCH_MMA_SM90_ENABLED)
|
||||
# define CUTE_ARCH_TMA_SM90_ENABLED
|
||||
@@ -48,41 +53,59 @@
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED))
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED) || defined(CUTLASS_ARCH_MMA_SM101A_ENABLED) ||\
|
||||
defined(CUTLASS_ARCH_MMA_SM120A_ENABLED))
|
||||
# define CUTE_ARCH_TMA_SM90_ENABLED
|
||||
# define CUTE_ARCH_DEVICE_MODIFIABLE_TMA_SM90_ENABLED
|
||||
# define CUTE_ARCH_STSM_SM90_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100A_ENABLED)
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED) || defined(CUTLASS_ARCH_MMA_SM101A_ENABLED))
|
||||
# define CUTE_ARCH_TCGEN05_TF32_MMA_ENABLED
|
||||
# define CUTE_ARCH_TCGEN05_F16F32_MMA_ENABLED
|
||||
# define CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED
|
||||
# define CUTE_ARCH_TCGEN05_MXF4_MMA_ENABLED
|
||||
# define CUTE_ARCH_TCGEN05_MXF4NVF4_MMA_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100A_ENABLED)
|
||||
# define CUTE_ARCH_TCGEN05_F16BF16_MMA_SCALED_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100A_ENABLED)
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED) || defined(CUTLASS_ARCH_MMA_SM101A_ENABLED))
|
||||
# define CUTE_ARCH_TCGEN05_S8_MMA_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100A_ENABLED)
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED) || defined(CUTLASS_ARCH_MMA_SM101A_ENABLED) ||\
|
||||
defined(CUTLASS_ARCH_MMA_SM120A_ENABLED))
|
||||
# define CUTE_ARCH_LDSM_SM100A_ENABLED
|
||||
# define CUTE_ARCH_STSM_SM100A_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100A_ENABLED)
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED) || defined(CUTLASS_ARCH_MMA_SM101A_ENABLED))
|
||||
# define CUTE_ARCH_TCGEN05_TMEM_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100A_ENABLED)
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED) || defined(CUTLASS_ARCH_MMA_SM101A_ENABLED))
|
||||
# define CUTE_ARCH_TMA_SM100_ENABLED
|
||||
#endif
|
||||
|
||||
// {add, mul, fma}.f32x2 PTX
|
||||
#if defined(CUTLASS_ARCH_MMA_SM100A_ENABLED)
|
||||
#if (defined(CUTLASS_ARCH_MMA_SM100A_ENABLED))
|
||||
#define CUTE_ARCH_FLOAT2_MATH_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM120_ENABLED) || defined(CUTLASS_ARCH_MMA_SM120A_ENABLED)
|
||||
# define CUTE_ARCH_MMA_SM120_ENABLED
|
||||
# define CUTE_ARCH_TMA_SM120_ENABLED
|
||||
#endif
|
||||
|
||||
#if defined(CUTLASS_ARCH_MMA_SM120_ENABLED) || defined(CUTLASS_ARCH_MMA_SM120A_ENABLED)
|
||||
# if (__CUDACC_VER_MAJOR__ > 12 || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 8))
|
||||
# define CUTE_ARCH_F8F6F4_MMA_ENABLED
|
||||
# define CUTE_ARCH_MXF8F6F4_MMA_ENABLED
|
||||
# define CUTE_ARCH_MXF4NVF4_2X_UE8M0_MMA_ENABLED
|
||||
# define CUTE_ARCH_MXF4NVF4_4X_UE4M3_MMA_ENABLED
|
||||
# endif
|
||||
#endif
|
||||
|
||||
|
||||
@@ -208,7 +208,7 @@ to_CUtensorMapDataType() {
|
||||
if constexpr (is_same_v<T, uint8_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT8; } else
|
||||
if constexpr (is_same_v<T, float_e4m3_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT8; } else
|
||||
if constexpr (is_same_v<T, float_e5m2_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT8; } else
|
||||
if constexpr (is_same_v<T, float_ue8m0_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT8; } else
|
||||
if constexpr (is_same_v<T, float_ue8m0_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT8; } else
|
||||
if constexpr (is_same_v<T, type_erased_dynamic_float8_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT8;} else
|
||||
if constexpr (is_same_v<T, uint16_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT16; } else
|
||||
if constexpr (is_same_v<T, uint32_t>) { return CU_TENSOR_MAP_DATA_TYPE_UINT32; } else
|
||||
@@ -221,18 +221,18 @@ to_CUtensorMapDataType() {
|
||||
if constexpr (is_same_v<T, bfloat16_t>) { return CU_TENSOR_MAP_DATA_TYPE_BFLOAT16; } else
|
||||
if constexpr (is_same_v<T, tfloat32_t>) { return CU_TENSOR_MAP_DATA_TYPE_TFLOAT32; } else
|
||||
#if ((__CUDACC_VER_MAJOR__ > 12) || ((__CUDACC_VER_MAJOR__ == 12) && (__CUDACC_VER_MINOR__ > 6)))
|
||||
if constexpr (is_same_v<T, float_e2m3_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, float_e3m2_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, float_e2m1_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B;} else
|
||||
if constexpr (is_same_v<T, cutlass::detail::float_e2m1_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, cutlass::detail::float_e2m3_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, cutlass::detail::float_e3m2_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, detail::type_erased_dynamic_float6_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, type_erased_dynamic_float6_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, detail::type_erased_dynamic_float4_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN16B;} else
|
||||
if constexpr (is_same_v<T, type_erased_dynamic_float4_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B; } else
|
||||
if constexpr (is_same_v<T, float_e2m1_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B; } else
|
||||
if constexpr (is_same_v<T, float_e2m3_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B; } else
|
||||
if constexpr (is_same_v<T, float_e3m2_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B; } else
|
||||
if constexpr (is_same_v<T, type_erased_dynamic_float4_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B; } else
|
||||
if constexpr (is_same_v<T, type_erased_dynamic_float6_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B; } else
|
||||
if constexpr (is_same_v<T, detail::float_e2m1_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN16B; } else
|
||||
if constexpr (is_same_v<T, detail::float_e2m3_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B; } else
|
||||
if constexpr (is_same_v<T, detail::float_e3m2_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B; } else
|
||||
if constexpr (is_same_v<T, detail::type_erased_dynamic_float4_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN16B; } else
|
||||
if constexpr (is_same_v<T, detail::type_erased_dynamic_float6_unpacksmem_t>) { return CU_TENSOR_MAP_DATA_TYPE_16U6_ALIGN16B; } else
|
||||
#endif
|
||||
|
||||
|
||||
{ static_assert(sizeof(T) < 0, "Unknown TMA Format!"); }
|
||||
}
|
||||
|
||||
@@ -258,7 +258,6 @@ to_CUtensorMapSwizzle(SmemSwizzleBits const& t, SmemSwizzleBase const& b) {
|
||||
case SmemSwizzleBase::SWIZZLE_BASE_32B: return CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B;
|
||||
case SmemSwizzleBase::SWIZZLE_BASE_64B: return CU_TENSOR_MAP_SWIZZLE_128B_ATOM_64B;
|
||||
#endif
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -56,6 +56,15 @@ 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);
|
||||
cutlass::arch::synclog_emit_tma_load(__LINE__, gmem_int_desc, smem_int_mbar, smem_int_ptr);
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.1d.shared::cta.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), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.1d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3}], [%2], %4;"
|
||||
@@ -63,6 +72,7 @@ struct SM90_TMA_LOAD_1D
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "l"(cache_hint)
|
||||
: "memory");
|
||||
#endif
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
@@ -102,6 +112,15 @@ 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);
|
||||
cutlass::arch::synclog_emit_tma_load(__LINE__, gmem_int_desc, smem_int_mbar, smem_int_ptr);
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.2d.shared::cta.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), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4}], [%2], %5;"
|
||||
@@ -109,6 +128,7 @@ struct SM90_TMA_LOAD_2D
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1), "l"(cache_hint)
|
||||
: "memory");
|
||||
#endif
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
@@ -148,6 +168,15 @@ 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);
|
||||
cutlass::arch::synclog_emit_tma_load(__LINE__, gmem_int_desc, smem_int_mbar, smem_int_ptr);
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.3d.shared::cta.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), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4, %5}], [%2], %6;"
|
||||
@@ -155,6 +184,7 @@ struct SM90_TMA_LOAD_3D
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "l"(cache_hint)
|
||||
: "memory");
|
||||
#endif
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
@@ -194,6 +224,15 @@ 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);
|
||||
cutlass::arch::synclog_emit_tma_load(__LINE__, gmem_int_desc, smem_int_mbar, smem_int_ptr);
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.4d.shared::cta.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), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.4d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4, %5, %6}], [%2], %7;"
|
||||
@@ -201,6 +240,7 @@ struct SM90_TMA_LOAD_4D
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "l"(cache_hint)
|
||||
: "memory");
|
||||
#endif
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
@@ -240,6 +280,15 @@ 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);
|
||||
cutlass::arch::synclog_emit_tma_load(__LINE__, gmem_int_desc, smem_int_mbar, smem_int_ptr);
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.5d.shared::cta.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), "l"(cache_hint)
|
||||
: "memory");
|
||||
#else
|
||||
asm volatile (
|
||||
"cp.async.bulk.tensor.5d.shared::cluster.global.mbarrier::complete_tx::bytes.L2::cache_hint"
|
||||
" [%0], [%1, {%3, %4, %5, %6, %7}], [%2], %8;"
|
||||
@@ -247,6 +296,7 @@ struct SM90_TMA_LOAD_5D
|
||||
: "r"(smem_int_ptr), "l"(gmem_int_desc), "r"(smem_int_mbar),
|
||||
"r"(crd0), "r"(crd1), "r"(crd2), "r"(crd3), "r"(crd4), "l"(cache_hint)
|
||||
: "memory");
|
||||
#endif
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
@@ -581,6 +631,9 @@ struct SM90_TMA_LOAD_MULTICAST_1D
|
||||
int32_t const& crd0)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
@@ -607,6 +660,9 @@ struct SM90_TMA_LOAD_MULTICAST_2D
|
||||
int32_t const& crd0, int32_t const& crd1)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
@@ -633,6 +689,9 @@ struct SM90_TMA_LOAD_MULTICAST_3D
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
@@ -659,6 +718,9 @@ struct SM90_TMA_LOAD_MULTICAST_4D
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
@@ -685,6 +747,9 @@ struct SM90_TMA_LOAD_MULTICAST_5D
|
||||
int32_t const& crd0, int32_t const& crd1, int32_t const& crd2, int32_t const& crd3, int32_t const& crd4)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
@@ -757,6 +822,9 @@ struct SM90_TMA_LOAD_IM2COL_MULTICAST_3D
|
||||
uint16_t const& offset_w)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
@@ -786,6 +854,9 @@ struct SM90_TMA_LOAD_IM2COL_MULTICAST_4D
|
||||
uint16_t const& offset_w, uint16_t const& offset_h)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
@@ -815,6 +886,9 @@ struct SM90_TMA_LOAD_IM2COL_MULTICAST_5D
|
||||
uint16_t const& offset_w, uint16_t const& offset_h, uint16_t const& offset_d)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TMA_SM90_ENABLED)
|
||||
#if defined(CUTE_ARCH_TMA_SM120_ENABLED)
|
||||
CUTE_INVALID_CONTROL_PATH("Trying to use tma without CUTE_ARCH_TMA_SM90_ENABLED.");
|
||||
#endif
|
||||
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(desc_ptr);
|
||||
uint32_t smem_int_mbar = cast_smem_ptr_to_uint(mbar_ptr);
|
||||
uint32_t smem_int_ptr = cast_smem_ptr_to_uint(smem_ptr);
|
||||
|
||||
@@ -552,7 +552,8 @@ make_runtime_instr_desc(UMMA::InstrDescriptor desc_i, uint16_t sparse_id2 = 0u,
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One,
|
||||
bool is_sparse = false>
|
||||
bool is_sparse = false
|
||||
>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
UMMA::InstrDescriptorBlockScaled
|
||||
make_instr_desc_block_scaled()
|
||||
|
||||
@@ -28,9 +28,6 @@
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
//
|
||||
|
||||
//
|
||||
|
||||
#pragma once
|
||||
|
||||
@@ -303,6 +300,92 @@ struct SM100_MMA_F16BF16_TS_SCALED
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_TF32_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 64 || M == 128, "SM100_MMA_TF32_SS_SPARSE M-mode size should be 64 or 128 for 1 CTA cluster MMA.");
|
||||
static_assert((M == 64 && (N % 8 == 0) && (8 <= N) && (N <= 256)) ||
|
||||
(M == 128 && (N % 16 == 0) && (16 <= N) && (N <= 256)),
|
||||
"SM100_MMA_TF32_SS_SPARSE N-mode size should be a multiple of 8 between 8 and 256 for M=64,\
|
||||
or a multiple of 16 between 16 and 256 for M=128.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_TF32_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[4] = {0, 0, 0, 0};
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::1.kind::tf32 [%0], %1, %2, [%9], %3, {%5, %6, %7, %8}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_TF32_SS_SPARSE without CUTE_ARCH_TCGEN05_TF32_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_F16BF16_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 64 || M == 128, "SM100_MMA_F16BF16_SS_SPARSE M-mode size should be 64 or 128 for 1 CTA cluster MMA.");
|
||||
static_assert((M == 64 && (N % 8 == 0) && (8 <= N) && (N <= 256)) ||
|
||||
(M == 128 && (N % 16 == 0) && (16 <= N) && (N <= 256)),
|
||||
"SM100_MMA_F16BF16_SS_SPARSE N-mode size should be a multiple of 8 between 8 and 256 for M=64,\
|
||||
or a multiple of 16 between 16 and 256 for M=128.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_F16F32_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[4] = {0, 0, 0, 0};
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::1.kind::f16 [%0], %1, %2, [%9], %3, {%5, %6, %7, %8}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_F16BF16_SS_SPARSE without CUTE_ARCH_TCGEN05_F16F32_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
@@ -551,6 +634,88 @@ struct SM100_MMA_F16BF16_2x1SM_TS_SCALED
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_TF32_2x1SM_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 128 || M == 256, "SM100_MMA_TF32_2x1SM_SS_SPARSE M-mode size should be 128 or 256 for 2 CTA cluster MMA.");
|
||||
static_assert((N % 32 == 0) && (32 <= N) && (N <= 256), "SM100_MMA_TF32_2x1SM_SS_SPARSE N-mode size should be a multiple of 32 between 32 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_TF32_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[8] = {0, 0, 0, 0, 0, 0, 0, 0};
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::2.kind::tf32 [%0], %1, %2, [%13], %3, {%5, %6, %7, %8, %9, %10, %11, %12}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]),
|
||||
"r"(mask[4]), "r"(mask[5]), "r"(mask[6]), "r"(mask[7]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_TF32_2x1SM_SS_SPARSE without CUTE_ARCH_TCGEN05_TF32_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_F16BF16_2x1SM_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 128 || M == 256, "SM100_MMA_F16BF16_2x1SM_SS_SPARSE M-mode size should be 128 or 256 for 2 CTA cluster MMA.");
|
||||
static_assert((N % 32 == 0) && (32 <= N) && (N <= 256), "SM100_MMA_F16BF16_2x1SM_SS_SPARSE N-mode size should be a multiple of 32 between 32 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_F16F32_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[8] = {0, 0, 0, 0, 0, 0, 0, 0};
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::2.kind::f16 [%0], %1, %2, [%13], %3, {%5, %6, %7, %8, %9, %10, %11, %12}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]),
|
||||
"r"(mask[4]), "r"(mask[5]), "r"(mask[6]), "r"(mask[7]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_F16BF16_2x1SM_SS_SPARSE without CUTE_ARCH_TCGEN05_F16F32_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::Saturate c_sat = UMMA::Saturate::False>
|
||||
@@ -632,6 +797,47 @@ struct SM100_MMA_S8_TS
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::Saturate c_sat = UMMA::Saturate::False>
|
||||
struct SM100_MMA_S8_SS_SPARSE
|
||||
{
|
||||
static_assert(is_same_v<c_type, int32_t>, "SM100_MMA_S8_SS_SPARSE result type can only be int32_t.");
|
||||
static_assert(M == 64 || M == 128, "SM100_MMA_S8_SS_SPARSE M-mode size should be 64 or 128 for 1 CTA cluster MMA.");
|
||||
static_assert(N == 8 || ((N % 16 == 0) && (16 <= N) && (N <= 256)), "SM100_MMA_S8_SS_SPARSE N-mode size should be 8 or a multiple of 16 between 16 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_S8_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[4] = {0, 0, 0, 0};
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::1.kind::i8 [%0], %1, %2, [%9], %3, {%5, %6, %7, %8}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_S8_SS_SPARSE without CUTE_ARCH_TCGEN05_S8_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::Saturate c_sat = UMMA::Saturate::False>
|
||||
@@ -714,10 +920,49 @@ struct SM100_MMA_S8_2x1SM_TS
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::Saturate c_sat = UMMA::Saturate::False>
|
||||
struct SM100_MMA_S8_2x1SM_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 128 || M == 256, "SM100_MMA_S8 M-mode size should be 128 or 256 for 2 CTA cluster MMA.");
|
||||
static_assert((N % 32 == 0) && (32 <= N) && (N <= 256), "SM100_MMA_S8 N-mode size should be a multiple of 32 between 32 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_S8_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[8] = {0, 0, 0, 0, 0, 0, 0, 0};
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::2.kind::i8 [%0], %1, %2, [%13], %3, {%5, %6, %7, %8, %9, %10, %11, %12}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]),
|
||||
"r"(mask[4]), "r"(mask[5]), "r"(mask[6]), "r"(mask[7]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_S8_2x1SM_SS_SPARSE without CUTE_ARCH_TCGEN05_S8_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
struct SM100_MMA_F8F6F4_SS
|
||||
{
|
||||
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
@@ -876,6 +1121,91 @@ struct SM100_MMA_F8F6F4_2x1SM_TS
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_F8F6F4_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 64 || M == 128, "SM100_MMA_F8F6F4_SS_SPARSE M-mode size should be 64 or 128 for 1 CTA cluster MMA.");
|
||||
static_assert((M == 64 && (N % 8 == 0) && (8 <= N) && (N <= 256)) ||
|
||||
(M == 128 && (N % 16 == 0) && (16 <= N) && (N <= 256)),
|
||||
"SM100_MMA_F8F6F4_SS_SPARSE N-mode size should be a multiple of 8 between 8 and 256 for M=64,\
|
||||
or a multiple of 16 between 16 and 256 for M=128.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[4] = {0, 0, 0, 0}; // %5, %6, %7, %8
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::1.kind::f8f6f4 [%0], %1, %2, [%9], %3, {%5, %6, %7, %8}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_F8F6F4_SS_SPARSE without CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_MXF8F6F4_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 128, "SM100_MMA_MXF8F6F4_SS_SPARSE M-mode size should be 128 for 1 CTA cluster MMA.");
|
||||
static_assert((N % 8 == 0) && (8 <= N) && (N <= 256), "SM100_MMA_MXF8F6F4_SS_SPARSE N-mode size should be a multiple of 8 between 8 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
using SFARegisters = uint32_t[1];
|
||||
using SFBRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tsfa_addr,
|
||||
uint32_t const& tsfb_addr,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::1.kind::mxf8f6f4.block_scale [%0], %1, %2, [%7], %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC), "r"(tsfa_addr), "r"(tsfb_addr), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_MXF8F6F4_SS_SPARSE without CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
struct SM100_MMA_F8F6F4_2x1SM_SS
|
||||
{
|
||||
using DRegisters = void;
|
||||
@@ -910,6 +1240,47 @@ struct SM100_MMA_F8F6F4_2x1SM_SS
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_MXF8F6F4_2x1SM_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 256, "SM100_MMA_MXF8F6F4_2x1SM_SS_SPARSE M-mode size should be 256 for 2 CTA cluster MMA.");
|
||||
static_assert((N % 16 == 0) && (16 <= N) && (N <= 256), "SM100_MMA_MXF8F6F4_2x1SM_SS_SPARSE N-mode size should be a multiple of 16 between 16 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tsfa_addr,
|
||||
uint32_t const& tsfb_addr,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::2.kind::mxf8f6f4.block_scale [%0], %1, %2, [%7], %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(tsfa_addr), "r"(tsfb_addr), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_MXF8F6F4_2x1SM_SS_SPARSE without CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
@@ -950,6 +1321,46 @@ struct SM100_MMA_MXF8F6F4_2x1SM_SS
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type,
|
||||
int M, int N, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_F8F6F4_2x1SM_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 128 || M == 256, "SM100_MMA_F8F6F4_2x1SM_SS_SPARSE M-mode size should be 128 or 256 for 2 CTA cluster MMA.");
|
||||
static_assert((N % 32 == 0) && (32 <= N) && (N <= 256), "SM100_MMA_F8F6F4_2x1SM_SS_SPARSE N-mode size should be a multiple of 32 between 32 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
uint32_t mask[8] = {0, 0, 0, 0, 0, 0, 0, 0};
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::2.kind::f8f6f4 [%0], %1, %2, [%13], %3, {%5, %6, %7, %8, %9, %10, %11, %12}, p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(mask[0]), "r"(mask[1]), "r"(mask[2]), "r"(mask[3]),
|
||||
"r"(mask[4]), "r"(mask[5]), "r"(mask[6]), "r"(mask[7]), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_F8F6F4_2x1SM_SS_SPARSE without CUTE_ARCH_TCGEN05_MXF8F6F4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, int VS, UMMA::Major a_major, UMMA::Major b_major,
|
||||
@@ -1014,7 +1425,68 @@ struct SM100_MMA_MXF4_SS
|
||||
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, int VS, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_MXF4NVF4_SS_SPARSE
|
||||
{
|
||||
static_assert(M == 128, "SM100_MMA_MXF4NVF4_SS_SPARSE M-mode size should be 128 for 1 CTA cluster MMA.");
|
||||
static_assert((N % 8 == 0) && (8 <= N) && (N <= 256), "SM100_MMA_MXF4NVF4_SS_SPARSE N-mode size should be a multiple of 8 between 8 and 256.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
using SFARegisters = uint32_t[1];
|
||||
using SFBRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tsfa_addr,
|
||||
uint32_t const& tsfb_addr,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
if constexpr (VS == 32) {
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF4NVF4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::1.kind::mxf4nvf4.block_scale.scale_vec::4X [%0], %1, %2, [%7], %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(tsfa_addr), "r"(tsfb_addr), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_MXF4NVF4_SS_SPARSE (VS = 32) without CUTE_ARCH_TCGEN05_MXF4NVF4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
|
||||
if constexpr (VS == 64) {
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::1.kind::mxf4.block_scale.scale_vec::2X [%0], %1, %2, [%7], %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(tsfa_addr), "r"(tsfb_addr), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_MXF4NVF4_SS_SPARSE (VS = 64) without CUTE_ARCH_TCGEN05_MXF4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, int VS, UMMA::Major a_major, UMMA::Major b_major,
|
||||
@@ -1078,5 +1550,67 @@ struct SM100_MMA_MXF4_2x1SM_SS
|
||||
}
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type,
|
||||
int M, int N, int VS, UMMA::Major a_major, UMMA::Major b_major,
|
||||
UMMA::ScaleIn a_neg = UMMA::ScaleIn::One, UMMA::ScaleIn b_neg = UMMA::ScaleIn::One>
|
||||
struct SM100_MMA_MXF4NVF4_2x1SM_SS_SPARSE
|
||||
{
|
||||
static_assert((N % 16 == 0) && (16 <= N) && (N <= 256), "SM100_MMA_MXF4NVF4_2x1SM_SS_SPARSE N-mode size should be a multiple of 16 between 16 and 256.");
|
||||
static_assert((VS == 32) || (VS == 64), "SM100_MMA_MXF4NVF4_2x1SM_SS_SPARSE Vector size can only be 32 or 64.");
|
||||
|
||||
using DRegisters = void;
|
||||
using ARegisters = uint64_t[1];
|
||||
using BRegisters = uint64_t[1];
|
||||
using CRegisters = uint32_t[1];
|
||||
using SFARegisters = uint32_t[1];
|
||||
using SFBRegisters = uint32_t[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(uint64_t const& desc_a,
|
||||
uint64_t const& desc_b,
|
||||
uint32_t const& tmem_c,
|
||||
uint32_t const& scaleC,
|
||||
uint64_t const& idescE,
|
||||
uint32_t const& tsfa_addr,
|
||||
uint32_t const& tsfb_addr,
|
||||
uint32_t const& tmem_e)
|
||||
{
|
||||
if constexpr (VS == 32) {
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF4NVF4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::2.kind::mxf4nvf4.block_scale.scale_vec::4X [%0], %1, %2, [%7], %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(tsfa_addr), "r"(tsfb_addr), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_MXF4NVF4_2x1SM_SS_SPARSE (VS = 32) without CUTE_ARCH_TCGEN05_MXF4NVF4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
|
||||
if constexpr (VS == 64) {
|
||||
#if defined(CUTE_ARCH_TCGEN05_MXF4_MMA_ENABLED)
|
||||
if (cute::elect_one_sync()) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"setp.ne.b32 p, %4, 0;\n\t"
|
||||
"tcgen05.mma.sp.cta_group::2.kind::mxf4.block_scale.scale_vec::2X [%0], %1, %2, [%7], %3, [%5], [%6], p; \n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(uint32_t(idescE>>32)), "r"(scaleC),
|
||||
"r"(tsfa_addr), "r"(tsfb_addr), "r"(tmem_e));
|
||||
}
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM100_MMA_MXF4NVF4_2x1SM_SS_SPARSE (VS = 64) without CUTE_ARCH_TCGEN05_MXF4_MMA_ENABLED");
|
||||
#endif
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -31,15 +31,10 @@
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cute/config.hpp>
|
||||
#include <cute/arch/config.hpp>
|
||||
|
||||
#include <cute/arch/mma.hpp>
|
||||
|
||||
// Config
|
||||
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) && defined(__CUDA_ARCH_FEAT_SM90_ALL))
|
||||
# define CUTE_ARCH_MMA_SM90A_ENABLED
|
||||
#endif
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cute {
|
||||
|
||||
@@ -626,7 +626,7 @@ make_tma_atom_im2col(CopyOp,
|
||||
auto tma_layout_trunc = take<0,smem_tma_rank>(tma_layout_full);
|
||||
|
||||
// Split according to the portion each multicast CTA will be responsible for
|
||||
auto tma_layout_vt = logical_divide(tma_layout_trunc, shape_div(size(tma_layout_trunc), num_multicast));
|
||||
auto tma_layout_vt = logical_divide(tma_layout_trunc, safe_div(size(tma_layout_trunc), num_multicast));
|
||||
|
||||
#if 0
|
||||
print("glayout_basis : "); print(glayout_basis); print("\n");
|
||||
@@ -748,7 +748,7 @@ make_tma_copy_im2col(CopyOp const& copy_op,
|
||||
// Scale that up to cover all of the smem_coords
|
||||
auto layout_V = tile_to_shape(make_layout(layout_v), size(cta_v_map));
|
||||
// CTA T -> smem idx
|
||||
auto layout_t = make_layout(cosize(cta_t_map), shape_div(num_elems_per_tma, cosize(cta_t_map)));
|
||||
auto layout_t = make_layout(cosize(cta_t_map), safe_div(num_elems_per_tma, cosize(cta_t_map)));
|
||||
// CTA TID -> smem coord
|
||||
auto layout_T = composition(inv_smem_layout, composition(layout_t, cta_t_map));
|
||||
// Combine with the T mapping
|
||||
|
||||
@@ -1165,7 +1165,7 @@ make_tma_copy_tiled(CopyOp const& copy_op,
|
||||
// Scale that up to cover all of the smem_coords
|
||||
auto layout_V = tile_to_shape(make_layout(layout_v), size(cta_v_map));
|
||||
// CTA T -> smem idx
|
||||
auto layout_t = make_layout(cosize(cta_t_map), shape_div(num_elems_per_tma, cosize(cta_t_map)));
|
||||
auto layout_t = make_layout(cosize(cta_t_map), safe_div(num_elems_per_tma, cosize(cta_t_map)));
|
||||
// CTA TID -> smem coord
|
||||
auto layout_T = composition(inv_smem_layout, composition(layout_t, cta_t_map));
|
||||
// Combine with the T mapping
|
||||
@@ -1400,16 +1400,19 @@ tma_partition(Copy_Atom<Args...> const& copy_atom,
|
||||
}
|
||||
|
||||
// TMA Multicast Masks Calculation
|
||||
template <int Mode, class CtaLayout, class CtaCoord>
|
||||
template <class CtaLayout, class CtaCoord>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
uint16_t
|
||||
create_tma_multicast_mask(CtaLayout const& cta_layout_vmnk,
|
||||
CtaCoord const& cta_coord_vmnk)
|
||||
{
|
||||
auto cta_coord_slicer = replace<Mode>(cta_coord_vmnk, _);
|
||||
auto [cta_layout, elected_cta] = slice_and_offset(cta_coord_slicer, cta_layout_vmnk);
|
||||
auto [cta_layout, elected_cta] = slice_and_offset(cta_coord_vmnk, cta_layout_vmnk);
|
||||
|
||||
uint16_t mcast_mask = 0;
|
||||
if constexpr (rank_v<decltype(cta_layout)> == 0) {
|
||||
// Trivial case with no additional ctas
|
||||
mcast_mask = uint16_t(1);
|
||||
} else
|
||||
if constexpr (rank_v<decltype(cta_layout)> == 1 and depth_v<decltype(cta_layout)> <= 1 and
|
||||
not is_static<decltype(cta_layout)>::value) {
|
||||
// Get the instruction code -- optimized for dynamic flat-rank-1 cta_layout
|
||||
@@ -1432,6 +1435,16 @@ create_tma_multicast_mask(CtaLayout const& cta_layout_vmnk,
|
||||
return mcast_mask;
|
||||
}
|
||||
|
||||
// Projections multicast mask
|
||||
template <int Mode, int... Modes, class CtaLayout, class CtaCoord>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
uint16_t
|
||||
create_tma_multicast_mask(CtaLayout const& cta_layout_vmnk,
|
||||
CtaCoord const& cta_coord_vmnk)
|
||||
{
|
||||
return create_tma_multicast_mask<Modes...>(cta_layout_vmnk, replace<Mode>(cta_coord_vmnk, _));
|
||||
}
|
||||
|
||||
////////////////////////////////////
|
||||
// Make TMA copy A/B/C
|
||||
///////////////////////////////////
|
||||
|
||||
@@ -154,9 +154,10 @@ struct MMA_Atom<MMA_Traits<MMAOperation, Args...>>
|
||||
if constexpr (has_dereference<FrgTypeA>::value) {
|
||||
// If the intended FrgTypeA is a view (of the current tensor), forward the whole
|
||||
static_assert(is_same<ValTypeA, typename remove_cvref_t<ATensor>::value_type>::value
|
||||
|
||||
|| (sizeof_bits_v<typename remove_cvref_t<ATensor>::value_type> == 8 &&
|
||||
(sizeof_bits_v<ValTypeA> == 8 || sizeof_bits_v<ValTypeA> == 6 || sizeof_bits_v<ValTypeA> == 4))
|
||||
|| (sizeof_bits_v<typename remove_cvref_t<ATensor>::value_type> == 4 &&
|
||||
(sizeof_bits_v<ValTypeA> == 4 || sizeof_bits_v<ValTypeA> == 3 || sizeof_bits_v<ValTypeA> == 2))
|
||||
, "Expecting ValTypeA type");
|
||||
return make_tensor<FrgTypeA>(static_cast<ATensor&&>(atensor));
|
||||
} else {
|
||||
@@ -1117,4 +1118,7 @@ print_svg(TiledMMA<Args...> const &mma) {
|
||||
#include <cute/atom/mma_traits_sm90.hpp>
|
||||
#include <cute/atom/mma_traits_sm90_gmma.hpp>
|
||||
#include <cute/atom/mma_traits_sm100.hpp>
|
||||
#include <cute/atom/mma_traits_sm120.hpp>
|
||||
#include <cute/atom/mma_traits_sm120_sparse.hpp>
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,262 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cute/arch/mma_sm120.hpp>
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
#include <cute/atom/mma_traits_sm80.hpp>
|
||||
#include <cute/layout.hpp>
|
||||
#include <cute/numeric/numeric_types.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
namespace SM120::BLOCKSCALED {
|
||||
|
||||
template <class MMAOp,
|
||||
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<MMAOp> const& traits,
|
||||
Tensor<TD, DLayout> & D,
|
||||
Tensor<TA, ALayout> const& A_zipped,
|
||||
Tensor<TB, BLayout> const& B_zipped,
|
||||
Tensor<TC, CLayout> const& C)
|
||||
{
|
||||
static_assert(is_rmem<TD>::value, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem<TA>::value, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem<TB>::value, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem<TC>::value, "Expected registers in MMA_Atom::call");
|
||||
|
||||
// Register value types from the MMA_Operation register arrays
|
||||
using RegTypeD = typename remove_extent<typename MMAOp::DRegisters>::type;
|
||||
using RegTypeA = typename remove_extent<typename MMAOp::ARegisters>::type;
|
||||
using RegTypeB = typename remove_extent<typename MMAOp::BRegisters>::type;
|
||||
using RegTypeC = typename remove_extent<typename MMAOp::CRegisters>::type;
|
||||
using RegTypeSFA = typename remove_extent<typename MMAOp::SFARegisters>::type;
|
||||
using RegTypeSFB = typename remove_extent<typename MMAOp::SFBRegisters>::type;
|
||||
|
||||
constexpr int RegNumD = extent<typename MMAOp::DRegisters>::value;
|
||||
constexpr int RegNumA = extent<typename MMAOp::ARegisters>::value;
|
||||
constexpr int RegNumB = extent<typename MMAOp::BRegisters>::value;
|
||||
constexpr int RegNumC = extent<typename MMAOp::CRegisters>::value;
|
||||
constexpr int RegNumSFA = extent<typename MMAOp::SFARegisters>::value;
|
||||
constexpr int RegNumSFB = extent<typename MMAOp::SFBRegisters>::value;
|
||||
|
||||
auto [A, SFA] = unzip_tensor(A_zipped);
|
||||
auto [B, SFB] = unzip_tensor(B_zipped);
|
||||
|
||||
using Shape_MNK = typename MMA_Traits<MMAOp>::Shape_MNK;
|
||||
constexpr int SFVecSize = MMA_Traits<MMAOp>::SFVecSize;
|
||||
|
||||
// Assert logical size
|
||||
CUTE_STATIC_ASSERT_V(size(SFA) == size<2>(Shape_MNK{}));
|
||||
CUTE_STATIC_ASSERT_V(size(SFB) == size<2>(Shape_MNK{}));
|
||||
|
||||
// Assert physical size
|
||||
CUTE_STATIC_ASSERT(decltype(cosize(layout(SFA))){} == size<2>(Shape_MNK{}) / SFVecSize);
|
||||
CUTE_STATIC_ASSERT(decltype(cosize(layout(SFB))){} == size<2>(Shape_MNK{}) / SFVecSize);
|
||||
|
||||
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>{});
|
||||
|
||||
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>{});
|
||||
|
||||
Tensor rSFA = recast<RegTypeSFA>(filter_zeros(SFA));
|
||||
Tensor rSFB = recast<RegTypeSFB>(filter_zeros(SFB));
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size(rSFA) == Int<RegNumSFA>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rSFB) == Int<RegNumSFB>{});
|
||||
|
||||
detail::explode(MMAOp::fma,
|
||||
rD, make_int_sequence<RegNumD>{},
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{},
|
||||
rSFA, make_int_sequence<RegNumSFA>{},
|
||||
rSFB, make_int_sequence<RegNumSFB>{});
|
||||
}
|
||||
} // namespace SM120::BLOCKSCALED
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// MMA F8F6F4 16x8x32 TN
|
||||
template <class a_type, class b_type, class c_type>
|
||||
struct MMA_Traits<SM120_16x8x32_TN<a_type, b_type, c_type>>
|
||||
: MMA_Traits<SM80_16x8x32_S32S8S8S32_TN>
|
||||
{
|
||||
// The MMA accepts 8-bit inputs regardless of the types for A and B
|
||||
using ValTypeA = uint8_t;
|
||||
using ValTypeB = uint8_t;
|
||||
|
||||
using ValTypeD = c_type;
|
||||
using ValTypeC = c_type;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// MMA MXF8F6F4 16x8x64 TN
|
||||
template <class a_type, class b_type, class c_type, class sf_type, int VS>
|
||||
struct MMA_Traits<SM120::BLOCKSCALED::SM120_16x8x64_TN_VS<a_type, b_type, c_type, sf_type, VS>>
|
||||
{
|
||||
// The MMA accepts 4-bit inputs regardless of the types for A and B
|
||||
using ValTypeA = uint4_t;
|
||||
using ValTypeB = uint4_t;
|
||||
|
||||
using ValTypeD = c_type;
|
||||
using ValTypeC = c_type;
|
||||
|
||||
using ValTypeSF = sf_type;
|
||||
constexpr static int SFVecSize = VS;
|
||||
|
||||
using Shape_MNK = Shape<_16,_8,_64>;
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// (T32,V32) -> (M16,K64)
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _8,_2, _2>>,
|
||||
Stride<Stride<_128,_1>,Stride<_16,_8,_512>>>;
|
||||
// (T32,V16) -> (M16,K64)
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_8, _2>>,
|
||||
Stride<Stride<_64,_1>,Stride<_8,_256>>>;
|
||||
// (T32,V64) -> (M16,K64)
|
||||
using SFALayout = Layout<Shape <Shape <_2,_2,_8>,_64>, // Effectively 16 threads due to the 2:0 mode
|
||||
Stride<Stride<_8,_0,_1>,_16>>;
|
||||
// (T32,V64) -> (N8,K64)
|
||||
using SFBLayout = Layout<Shape <Shape <_4,_8>,_64>, // Effectively 8 threads due to the 4:0 mode
|
||||
Stride<Stride<_0,_1>, _8>>;
|
||||
// (T32,V4) -> (M16,N8)
|
||||
using CLayout = SM80_16x8_Row;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// MMA MXF8F6F4 16x8x32 TN
|
||||
template <class a_type, class b_type, class c_type, class sf_type, int VS>
|
||||
struct MMA_Traits<SM120::BLOCKSCALED::SM120_16x8x32_TN_VS<a_type, b_type, c_type, sf_type, VS>>
|
||||
{
|
||||
using UnderlyingTraits = MMA_Traits<SM120_16x8x32_TN<a_type, b_type, c_type>>;
|
||||
|
||||
// The MMA accepts 8-bit inputs regardless of the types for A and B
|
||||
using ValTypeA = typename UnderlyingTraits::ValTypeA;
|
||||
using ValTypeB = typename UnderlyingTraits::ValTypeB;
|
||||
|
||||
using ValTypeD = typename UnderlyingTraits::ValTypeD;
|
||||
using ValTypeC = typename UnderlyingTraits::ValTypeC;
|
||||
|
||||
using Shape_MNK = typename UnderlyingTraits::Shape_MNK;
|
||||
using ThrID = typename UnderlyingTraits::ThrID;
|
||||
|
||||
using ALayout = typename UnderlyingTraits::ALayout;
|
||||
using BLayout = typename UnderlyingTraits::BLayout;
|
||||
using CLayout = typename UnderlyingTraits::CLayout;
|
||||
|
||||
// Scaling factor
|
||||
using ValTypeSF = sf_type;
|
||||
constexpr static int SFVecSize = VS;
|
||||
|
||||
// (T32,V32) -> (M16,K32)
|
||||
using SFALayout = Layout<Shape <Shape <_2,_2,_8>,_32>, // Effectively 16 threads due to the 2:0 mode
|
||||
Stride<Stride<_8,_0,_1>,_16>>;
|
||||
// (T32,V32) -> (N8,K32)
|
||||
using SFBLayout = Layout<Shape <Shape <_4,_8>,_32>, // Effectively 8 threads due to the 4:0 mode
|
||||
Stride<Stride<_0,_1>, _8>>;
|
||||
};
|
||||
|
||||
// Transform if needed
|
||||
template<class MMA_Op, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_A(MMA_Op const& op, Tensor&& tensor) {
|
||||
}
|
||||
template<class MMA_Op, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_B(MMA_Op const& op, Tensor&& tensor) {
|
||||
}
|
||||
|
||||
// For SM120 MMA F8F6F4 input fp4, the operand A/B are load from ld.matrix.
|
||||
// ld.matrix b4x16_p64 places FP4 data at the first four bits in each
|
||||
// eight-bit container, whereas MMA F8F6F4 expects the four-bit data to be in
|
||||
// the middle of the eight-bit container. Thus, e2m1 operands being fed
|
||||
// to MMA F8F6F4 must be shifted left by two bits.
|
||||
// 0b0000ABCD --> 0b00ABCD00
|
||||
// NOTE: Same transformation is NOT needed for FP6 and FP8.
|
||||
template<class AType, class BType, class... MMAArgs, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_A(SM120_16x8x32_TN<AType, BType, MMAArgs ...> const&, Tensor&& tensor) {
|
||||
using RegisterTypeA = typename remove_extent<typename
|
||||
SM120_16x8x32_TN<AType, BType, MMAArgs ...>::ARegisters>::type;
|
||||
if constexpr (cute::is_same_v<AType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeA>(tensor), [](RegisterTypeA& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
template<class AType, class BType, class... MMAArgs, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_B(SM120_16x8x32_TN<AType, BType, MMAArgs ...> const&, Tensor&& tensor) {
|
||||
using RegisterTypeB = typename remove_extent<typename
|
||||
SM120_16x8x32_TN<AType, BType, MMAArgs ...>::BRegisters>::type;
|
||||
if constexpr (cute::is_same_v<BType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeB>(tensor), [](RegisterTypeB& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
|
||||
namespace SM120::BLOCKSCALED {
|
||||
|
||||
// Template function with scale factor needs to enmuerate types one by one, as template
|
||||
// arguments contatins two variadic lists, which cannot be deduced in one shot.
|
||||
template<class AType, class BType, class CType, class SFType, int VS, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_A(SM120::BLOCKSCALED::SM120_16x8x32_TN_VS<AType, BType, CType, SFType, VS> const&, Tensor&& tensor) {
|
||||
using RegisterTypeA = typename remove_extent<typename
|
||||
SM120::BLOCKSCALED::SM120_16x8x32_TN_VS<AType, BType, CType, SFType, VS>::ARegisters>::type;
|
||||
if constexpr (cute::is_same_v<AType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeA>(tensor), [](RegisterTypeA& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
template<class AType, class BType, class CType, class SFType, int VS, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_B(SM120::BLOCKSCALED::SM120_16x8x32_TN_VS<AType, BType, CType, SFType, VS> const&, Tensor&& tensor) {
|
||||
using RegisterTypeB = typename remove_extent<typename
|
||||
SM120::BLOCKSCALED::SM120_16x8x32_TN_VS<AType, BType, CType, SFType, VS>::BRegisters>::type;
|
||||
if constexpr (cute::is_same_v<BType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeB>(tensor), [](RegisterTypeB& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
} // end namespace cute
|
||||
@@ -0,0 +1,326 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2025 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cute/arch/mma_sm120.hpp>
|
||||
#include <cute/arch/mma_sm120_sparse.hpp>
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
#include <cute/layout.hpp>
|
||||
#include <cute/numeric/numeric_types.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
namespace {
|
||||
|
||||
// (T32,V4) -> (M16,N8)
|
||||
using SM120_16x8_Row = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
|
||||
}
|
||||
|
||||
namespace SM120::BLOCKSCALED::SPARSE
|
||||
{
|
||||
|
||||
// Unpack explode/mma call with sparse and block scalaring inputs.
|
||||
template <class MMAOp,
|
||||
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<MMAOp> const&,
|
||||
Tensor<TD, DLayout> & D,
|
||||
Tensor<TA, ALayout> const& A,
|
||||
Tensor<TB, BLayout> const& B,
|
||||
Tensor<TC, CLayout> const& C)
|
||||
{
|
||||
static_assert(is_rmem_v<TD>, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem_v<TA>, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem_v<TB>, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem_v<TC>, "Expected registers in MMA_Atom::call");
|
||||
using DRegisters = typename MMAOp::DRegisters;
|
||||
using ARegisters = typename MMAOp::ARegisters;
|
||||
using ERegisters = typename MMAOp::ERegisters;
|
||||
using BRegisters = typename MMAOp::BRegisters;
|
||||
using CRegisters = typename MMAOp::CRegisters;
|
||||
using SFARegisters = typename MMAOp::SFARegisters;
|
||||
using SFBRegisters = typename MMAOp::SFBRegisters;
|
||||
// Register value types from the MMAOp register arrays
|
||||
using RegTypeD = typename remove_extent<DRegisters>::type;
|
||||
using RegTypeA = typename remove_extent<ARegisters>::type;
|
||||
using RegTypeE = typename remove_extent<ERegisters>::type;
|
||||
using RegTypeB = typename remove_extent<BRegisters>::type;
|
||||
using RegTypeC = typename remove_extent<CRegisters>::type;
|
||||
using RegTypeSFA = typename remove_extent<SFARegisters>::type;
|
||||
using RegTypeSFB = typename remove_extent<SFBRegisters>::type;
|
||||
constexpr int RegNumD = extent<DRegisters>::value;
|
||||
constexpr int RegNumA = extent<ARegisters>::value;
|
||||
constexpr int RegNumE = extent<ERegisters>::value;
|
||||
constexpr int RegNumB = extent<BRegisters>::value;
|
||||
constexpr int RegNumC = extent<CRegisters>::value;
|
||||
constexpr int RegNumSFA = extent<SFARegisters>::value;
|
||||
constexpr int RegNumSFB = extent<SFBRegisters>::value;
|
||||
|
||||
auto [tA, tSFA, tE] = unzip_tensor(A);
|
||||
auto [tB, tSFB ] = unzip_tensor(B);
|
||||
Tensor rA = recast<RegTypeA>(tA);
|
||||
Tensor rE = recast<RegTypeE>(tE);
|
||||
Tensor rB = recast<RegTypeB>(tB);
|
||||
Tensor rD = recast<RegTypeD>(D);
|
||||
Tensor rC = recast<RegTypeC>(C);
|
||||
Tensor rSFA = recast<RegTypeSFA>(tSFA);
|
||||
Tensor rSFB = recast<RegTypeSFB>(tSFB);
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size(rA) == Int<RegNumA>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rE) == Int<RegNumE>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rB) == Int<RegNumB>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumD>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rC) == Int<RegNumC>{});
|
||||
CUTE_STATIC_ASSERT_V(size(filter_zeros(rSFA)) == Int<RegNumSFA>{});
|
||||
CUTE_STATIC_ASSERT_V(size(filter_zeros(rSFB)) == Int<RegNumSFB>{});
|
||||
|
||||
detail::explode(MMAOp::fma,
|
||||
rD, make_int_sequence<RegNumD>{},
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{},
|
||||
rE, make_int_sequence<RegNumE>{},
|
||||
rSFA, make_int_sequence<RegNumSFA>{},
|
||||
rSFB, make_int_sequence<RegNumSFB>{});
|
||||
}
|
||||
|
||||
} // end namespace SM120::BLOCKSCALED::SPARSE
|
||||
|
||||
|
||||
namespace SM120::SPARSE
|
||||
{
|
||||
|
||||
template <class MMAOp,
|
||||
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<MMAOp> const&,
|
||||
Tensor<TD, DLayout> & D,
|
||||
Tensor<TA, ALayout> const& A,
|
||||
Tensor<TB, BLayout> const& B,
|
||||
Tensor<TC, CLayout> const& C)
|
||||
{
|
||||
static_assert(is_rmem_v<TD>, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem_v<TA>, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem_v<TB>, "Expected registers in MMA_Atom::call");
|
||||
static_assert(is_rmem_v<TC>, "Expected registers in MMA_Atom::call");
|
||||
using DRegisters = typename MMAOp::DRegisters;
|
||||
using ARegisters = typename MMAOp::ARegisters;
|
||||
using ERegisters = typename MMAOp::ERegisters;
|
||||
using BRegisters = typename MMAOp::BRegisters;
|
||||
using CRegisters = typename MMAOp::CRegisters;
|
||||
// Register value types from the MMAOp register arrays
|
||||
using RegTypeD = typename remove_extent<DRegisters>::type;
|
||||
using RegTypeA = typename remove_extent<ARegisters>::type;
|
||||
using RegTypeE = typename remove_extent<ERegisters>::type;
|
||||
using RegTypeB = typename remove_extent<BRegisters>::type;
|
||||
using RegTypeC = typename remove_extent<CRegisters>::type;
|
||||
constexpr int RegNumD = extent<DRegisters>::value;
|
||||
constexpr int RegNumA = extent<ARegisters>::value;
|
||||
constexpr int RegNumE = extent<ERegisters>::value;
|
||||
constexpr int RegNumB = extent<BRegisters>::value;
|
||||
constexpr int RegNumC = extent<CRegisters>::value;
|
||||
|
||||
auto [tA, tE] = unzip_tensor(A);
|
||||
Tensor rA = recast<RegTypeA>(tA);
|
||||
Tensor rE = recast<RegTypeE>(tE);
|
||||
Tensor rB = recast<RegTypeB>(B);
|
||||
Tensor rD = recast<RegTypeD>(D);
|
||||
Tensor rC = recast<RegTypeC>(C);
|
||||
CUTE_STATIC_ASSERT_V(size(rA) == Int<RegNumA>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rE) == Int<RegNumE>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rB) == Int<RegNumB>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumD>{});
|
||||
CUTE_STATIC_ASSERT_V(size(rC) == Int<RegNumC>{});
|
||||
|
||||
detail::explode(MMAOp::fma,
|
||||
rD, make_int_sequence<RegNumD>{},
|
||||
rA, make_int_sequence<RegNumA>{},
|
||||
rB, make_int_sequence<RegNumB>{},
|
||||
rC, make_int_sequence<RegNumC>{},
|
||||
rE, make_int_sequence<RegNumE>{});
|
||||
}
|
||||
|
||||
} // end namespace SM120::SPARSE
|
||||
|
||||
// sparse F8F6F4 without block-scaling
|
||||
template <class a_type, class b_type, class c_type>
|
||||
struct MMA_Traits<SM120::SPARSE::SM120_SPARSE_16x8x64_TN<a_type, b_type, c_type>>
|
||||
{
|
||||
using ValTypeA = sparse_elem<2, a_type>;
|
||||
using ValTypeE = sparse_elem<8, uint8_t>;
|
||||
using ValTypeB = uint8_t;
|
||||
using FrgTypeA = sparse_elem<2, uint8_t>;
|
||||
using FrgTypeE = sparse_elem<8, uint8_t>;
|
||||
|
||||
using ValTypeC = c_type;
|
||||
using ValTypeD = c_type;
|
||||
|
||||
using Shape_MNK = Shape<_16, _8, _64>;
|
||||
using ThrID = Layout<_32>;
|
||||
// (T32,V32) -> (M16,K64)
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _8,_2, _2>>,
|
||||
Stride<Stride<_128,_1>,Stride<_16,_8,_512>>>;
|
||||
// (T32,V16) -> (N8,K64)
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_4, _4>>,
|
||||
Stride<Stride<_32,_1>,Stride<_8,_128>>>;
|
||||
// (T32,V4) -> (M16,N8)
|
||||
using CLayout = SM120_16x8_Row;
|
||||
|
||||
// (T32, V32) -> (M16, K64)
|
||||
using ELayout = Layout<Shape <Shape <_2, _2,_8>, _32>,
|
||||
Stride<Stride<_8,_512,_1>,_16>>;
|
||||
};
|
||||
|
||||
// sparse MXF8F6F4 with block-scaling.
|
||||
template <class a_type, class b_type, class c_type, class sf_type, int VS>
|
||||
struct MMA_Traits<SM120::BLOCKSCALED::SPARSE::SM120_SPARSE_16x8x64_TN_VS<a_type, b_type, c_type, sf_type, VS>>
|
||||
: MMA_Traits<SM120::SPARSE::SM120_SPARSE_16x8x64_TN<a_type, b_type, c_type>>
|
||||
{
|
||||
using ValTypeA = sparse_elem<2, a_type>;
|
||||
using ValTypeE = sparse_elem<8, uint8_t>;
|
||||
using ValTypeB = uint8_t;
|
||||
using FrgTypeA = sparse_elem<2, uint8_t>;
|
||||
using FrgTypeE = sparse_elem<8, uint8_t>;
|
||||
|
||||
using ValTypeD = c_type;
|
||||
using ValTypeC = c_type;
|
||||
|
||||
using ValTypeSF = sf_type;
|
||||
constexpr static int SFVecSize = VS;
|
||||
|
||||
using UnderlyingSFTraits = MMA_Traits<SM120::BLOCKSCALED::SM120_16x8x64_TN_VS<a_type, b_type, c_type, sf_type, VS>>;
|
||||
using SFALayout = typename UnderlyingSFTraits::SFALayout;
|
||||
using SFBLayout = typename UnderlyingSFTraits::SFBLayout;
|
||||
};
|
||||
|
||||
template <class a_type, class b_type, class c_type, class sf_type, int VS>
|
||||
struct MMA_Traits<SM120::BLOCKSCALED::SPARSE::SM120_SPARSE_16x8x128_TN_VS<a_type, b_type, c_type, sf_type, VS>>
|
||||
{
|
||||
using ValTypeA = sparse_elem<4, uint8_t>;
|
||||
using ValTypeE = sparse_elem<16, uint8_t>;
|
||||
using ValTypeB = uint4_t;
|
||||
using FrgTypeA = sparse_elem<4, uint8_t>;
|
||||
using FrgTypeE = sparse_elem<16, uint8_t>;
|
||||
|
||||
using ValTypeC = c_type;
|
||||
using ValTypeD = c_type;
|
||||
|
||||
using ValTypeSF = sf_type;
|
||||
|
||||
constexpr static int SFVecSize = VS;
|
||||
|
||||
using Shape_MNK = Shape<_16, _8, _128>;
|
||||
using ThrID = Layout<_32>;
|
||||
// (T32,V64) -> (M16,K128)
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape <_16,_2, _2>>,
|
||||
Stride<Stride<_256,_1>,Stride<_16,_8,_1024>>>;
|
||||
// (T32,V32) -> (N8,K128)
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_8, _4>>,
|
||||
Stride<Stride<_64,_1>,Stride<_8,_256>>>;
|
||||
// (T32,V128) -> (M16,K128)
|
||||
using SFALayout = Layout<Shape <Shape <_2,_2,_8>,_128>,
|
||||
Stride<Stride<_8,_0,_1>, _16>>;
|
||||
// (T32,V128) -> (N8,K128)
|
||||
using SFBLayout = Layout<Shape <Shape <_4,_8>,_128>,
|
||||
Stride<Stride<_0,_1>, _8>>;
|
||||
// (T32,V4) -> (M16,N8)
|
||||
using CLayout = SM120_16x8_Row;
|
||||
// (T32, V64) -> (M16, K128)
|
||||
using ELayout = Layout<Shape <Shape <_2, _2,_8>, Shape< _64>>,
|
||||
Stride<Stride<_8,_1024,_1>,Stride<_16>>>;
|
||||
};
|
||||
|
||||
namespace SM120::SPARSE {
|
||||
|
||||
// For SM120 MMA F8F6F4 input fp4, the operand A/B are load from ld.matrix.
|
||||
// ld.matrix b4x16_p64 places FP4 data at the first four bits in each
|
||||
// eight-bit container, whereas MMA F8F6F4 expects the four-bit data to be in
|
||||
// the middle of the eight-bit container. Thus, e2m1 operands being fed
|
||||
// to MMA F8F6F4 must be shifted left by two bits.
|
||||
// 0b0000ABCD --> 0b00ABCD00
|
||||
// NOTE: Same transformation is NOT needed for FP6 and FP8.
|
||||
template<class AType, class BType, class... MMAArgs, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_A(SM120_SPARSE_16x8x64_TN<AType, BType, MMAArgs ...> const&, Tensor&& tensor) {
|
||||
using RegisterTypeA = typename remove_extent<typename
|
||||
SM120_SPARSE_16x8x64_TN<AType, BType, MMAArgs ...>::ARegisters>::type;
|
||||
if constexpr (cute::is_same_v<AType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeA>(tensor), [](RegisterTypeA& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
template<class AType, class BType, class... MMAArgs, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_B(SM120_SPARSE_16x8x64_TN<AType, BType, MMAArgs ...> const&, Tensor&& tensor) {
|
||||
using RegisterTypeB = typename remove_extent<typename
|
||||
SM120_SPARSE_16x8x64_TN<AType, BType, MMAArgs ...>::BRegisters>::type;
|
||||
if constexpr (cute::is_same_v<BType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeB>(tensor), [](RegisterTypeB& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
|
||||
} // end namespace SM120::SPARSE
|
||||
|
||||
namespace SM120::BLOCKSCALED::SPARSE {
|
||||
|
||||
// Template function with scale factor needs to enmuerate types one by one, as template
|
||||
// arguments contatins two variadic lists, which cannot be deduced in one shot.
|
||||
template<class AType, class BType, class CType, class SFType, int VS, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_A(SM120_SPARSE_16x8x64_TN_VS<AType, BType, CType, SFType, VS> const&, Tensor&& tensor) {
|
||||
using RegisterTypeA = typename remove_extent<typename
|
||||
SM120_SPARSE_16x8x64_TN_VS<AType, BType, CType, SFType, VS>::ARegisters>::type;
|
||||
if constexpr (cute::is_same_v<AType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeA>(tensor), [](RegisterTypeA& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
template<class AType, class BType, class CType, class SFType, int VS, class Tensor>
|
||||
CUTLASS_DEVICE void
|
||||
fp4_shift_B(SM120_SPARSE_16x8x64_TN_VS<AType, BType, CType, SFType, VS> const&, Tensor&& tensor) {
|
||||
using RegisterTypeB = typename remove_extent<typename
|
||||
SM120_SPARSE_16x8x64_TN_VS<AType, BType, CType, SFType, VS>::BRegisters>::type;
|
||||
if constexpr (cute::is_same_v<BType, cutlass::float_e2m1_t>) {
|
||||
cute::transform(recast<RegisterTypeB>(tensor), [](RegisterTypeB& v){ return v << 2; });
|
||||
}
|
||||
}
|
||||
|
||||
} // end namespace SM120::BLOCKSCALED::SPARSE
|
||||
|
||||
} // end namespace cute
|
||||
@@ -239,11 +239,10 @@ make_gmma_desc(Tensor<TEngine,TLayout> const& tensor)
|
||||
"Not a canonical GMMA_MN Layout: Expected K-size 256/sizeof_bits<T> for dense or (128|512)/sizeof_bits<T> for sparse.");
|
||||
|
||||
// Construct the canonical GMMA T Layout with shape ((W,n),(8,2))
|
||||
Layout canonical_layout = logical_divide(layout(u128_tensor), make_tile(Layout<Int<W>,_1>{}, Layout<Int<8>,_1>{}));
|
||||
Layout canonical_layout = logical_divide(layout(u128_tensor), Tile<Layout<Int<W>,_1>,Layout<Int<8>,_1>>{});
|
||||
|
||||
// Check ranks of canonical
|
||||
CUTE_STATIC_ASSERT_V(rank<0>(canonical_layout) == Int<2>{}, "Not a canonical GMMA_MN Layout: No flat offset mode");
|
||||
CUTE_STATIC_ASSERT_V(rank<1>(canonical_layout) == Int<2>{}, "Not a canonical GMMA_MN Layout: No flat offset mode");
|
||||
// Check profile of canonical
|
||||
CUTE_STATIC_ASSERT_V(congruent(canonical_layout, Shape<Shape<_1,_1>,Shape<_1,_1>>{}), "Not a canonical GMMA_MN Layout: Expected profile failure.");
|
||||
// Check canonical mode strides
|
||||
constexpr uint32_t stride_00 = stride<0,0>(canonical_layout);
|
||||
constexpr uint32_t expected_stride_00 = LAYOUT_TYPE == LayoutType::INTERLEAVE ? stride<0,0>(canonical_layout) : 1;
|
||||
@@ -274,11 +273,10 @@ make_gmma_desc(Tensor<TEngine,TLayout> const& tensor)
|
||||
"Not a canonical GMMA_K Layout: Expected K-size 2 for dense or 4 for sparse (in units of uint128_t).");
|
||||
|
||||
// Construct the canonical GMMA N Layout with shape ((8,n),(2,1))
|
||||
Layout canonical_layout = logical_divide(layout(u128_tensor), make_tile(Layout<_8,_1>{}, Layout<_2,_1>{}));
|
||||
Layout canonical_layout = logical_divide(layout(u128_tensor), Tile<Layout<_8,_1>,Layout<_2,_1>>{});
|
||||
|
||||
// Check ranks of canonical
|
||||
CUTE_STATIC_ASSERT_V(rank<0>(canonical_layout) == Int<2>{}, "Not a canonical GMMA_K Layout: No flat offset mode");
|
||||
CUTE_STATIC_ASSERT_V(rank<1>(canonical_layout) == Int<2>{}, "Not a canonical GMMA_K Layout: No flat offset mode");
|
||||
// Check profile of canonical
|
||||
CUTE_STATIC_ASSERT_V(congruent(canonical_layout, Shape<Shape<_1,_1>,Shape<_1,_1>>{}), "Not a canonical GMMA_K Layout: Expected profile failure.");
|
||||
// Check canonical mode strides
|
||||
constexpr uint32_t stride_00 = stride<0,0>(canonical_layout);
|
||||
constexpr uint32_t expected_stride_00 = W;
|
||||
|
||||
+80
-27
@@ -34,6 +34,7 @@
|
||||
#include <cute/container/array.hpp> // cute::array
|
||||
#include <cute/container/tuple.hpp> // cute::is_tuple
|
||||
#include <cute/numeric/integral_constant.hpp> // cute::Int
|
||||
#include <cute/numeric/integer_sequence.hpp> // cute::seq
|
||||
#include <cute/algorithm/tuple_algorithms.hpp> // cute::transform
|
||||
|
||||
/** IntTuple is an integer or a tuple of IntTuples.
|
||||
@@ -349,7 +350,6 @@ ceil_div(IntTupleA const& a, IntTupleB const& b)
|
||||
//
|
||||
// round_up
|
||||
// Round @a a up to the nearest multiple of @a b.
|
||||
// For negative numbers, rounds away from zero.
|
||||
//
|
||||
|
||||
template <class IntTupleA, class IntTupleB>
|
||||
@@ -378,7 +378,7 @@ round_up(IntTupleA const& a, IntTupleB const& b)
|
||||
* Return shape_div(a, product(b))
|
||||
* Case Int Int:
|
||||
* Enforce the divisibility condition a % b == 0 || b % a == 0 when possible
|
||||
* Return a / b with rounding away from 0 (that is, 1 or -1 when a < b)
|
||||
* Return ceil_div(a, b)
|
||||
*/
|
||||
template <class IntTupleA, class IntTupleB>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -399,32 +399,19 @@ shape_div(IntTupleA const& a, IntTupleB const& b)
|
||||
} else
|
||||
if constexpr (is_tuple<IntTupleB>::value) { // int tuple
|
||||
return shape_div(a, product(b));
|
||||
} else
|
||||
if constexpr (is_static<IntTupleA>::value && is_static<IntTupleB>::value) {
|
||||
static_assert(IntTupleA::value % IntTupleB::value == 0 || IntTupleB::value % IntTupleA::value == 0, "Static shape_div failure");
|
||||
return C<shape_div(IntTupleA::value, IntTupleB::value)>{};
|
||||
} else { // int int
|
||||
//assert(a % b == 0 || b % a == 0); // Waive dynamic assertion
|
||||
return a / b != 0 ? a / b : signum(a) * signum(b); // Division with rounding away from zero
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
/** Minimum for Shapes
|
||||
*/
|
||||
template <class IntTupleA, class IntTupleB>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
shape_min(IntTupleA const& a, IntTupleB const& b)
|
||||
{
|
||||
if constexpr (is_tuple<IntTupleA>::value || is_tuple<IntTupleB>::value) {
|
||||
static_assert(dependent_false<IntTupleA>, "Not implemented.");
|
||||
} else
|
||||
if constexpr (is_constant<1, IntTupleA>::value || is_constant<1, IntTupleB>::value) {
|
||||
return Int<1>{}; // _1 is less than all other shapes, preserve static
|
||||
} else {
|
||||
return cute::min(a, b);
|
||||
// Strong divisibility condition
|
||||
//static_assert((IntTupleA::value % IntTupleB::value == 0) or (IntTupleB::value % IntTupleA::value == 0), "Divisibility Condition");
|
||||
|
||||
// Weak divisibility condition
|
||||
if constexpr (is_static<IntTupleA>::value and is_static<IntTupleB>::value) {
|
||||
static_assert(((IntTupleA::value % IntTupleB::value) == 0) or ((IntTupleB::value % IntTupleA::value) == 0), "Divisibility Condition");
|
||||
} else {
|
||||
// DEBUG assert can cause extra registers and inappropriate compile-time/run-time failure
|
||||
//assert((((a % b) == 0) or ((a % b) == 0)) && "Divisibility Condition");
|
||||
}
|
||||
|
||||
return (a + b - Int<1>{}) / b;
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
@@ -572,6 +559,72 @@ filter_zeros(Tuple const& t)
|
||||
return filter_zeros(t, t);
|
||||
}
|
||||
|
||||
//
|
||||
// Static sorting utilities in detail::
|
||||
//
|
||||
|
||||
namespace detail {
|
||||
|
||||
// Some compilers fail to constexpr evaluate quick_sort
|
||||
// template <class T, size_t N>
|
||||
// constexpr cute::array<T,N> quick_sort(cute::array<T,N> a, int lo = 0, int hi = N-1) {
|
||||
// if (hi <= lo) return;
|
||||
// int p = lo;
|
||||
// for (int i = lo; i < hi; ++i) {
|
||||
// if (a[i] < a[hi]) {
|
||||
// T tmp = a[p]; a[p] = a[i]; a[i] = tmp;
|
||||
// ++p;
|
||||
// }
|
||||
// }
|
||||
// T tmp = a[p]; a[p] = a[hi]; a[hi] = tmp;
|
||||
// a = quick_sort(a, lo, p-1);
|
||||
// a = quick_sort(a, p+1, hi);
|
||||
// return a;
|
||||
// }
|
||||
|
||||
template <class T, size_t N>
|
||||
constexpr cute::array<T,N> exchange_sort(cute::array<T,N> a) {
|
||||
for (size_t i = 0; i < N; ++i) {
|
||||
for (size_t j = i+1; j < N; ++j) {
|
||||
if (a[j] < a[i]) {
|
||||
T tmp = a[j]; a[j] = a[i]; a[i] = tmp;
|
||||
}
|
||||
}
|
||||
}
|
||||
return a;
|
||||
}
|
||||
|
||||
template <class V, class I = cute::make_int_sequence<cute::tuple_size_v<V>>>
|
||||
struct Sort : Sort<to_seq_t<V>, to_seq_t<I>> {};
|
||||
|
||||
template <int... Vs, int... Is>
|
||||
struct Sort<seq<Vs...>, seq<Is...>> {
|
||||
static_assert(sizeof...(Vs) == sizeof...(Is));
|
||||
static constexpr cute::array<int,sizeof...(Is)> orig_array = {Vs...};
|
||||
static constexpr cute::array<int,sizeof...(Is)> sort_array = exchange_sort(orig_array);
|
||||
using type = seq<sort_array[Is]...>;
|
||||
};
|
||||
|
||||
struct kvpair {
|
||||
int key, val;
|
||||
constexpr bool operator<(kvpair const& o) const { return key < o.key; };
|
||||
};
|
||||
|
||||
template <class K, class V, class I = cute::make_int_sequence<cute::tuple_size_v<K>>>
|
||||
struct SortByKey : SortByKey<to_seq_t<K>, to_seq_t<V>, to_seq_t<I>> {};
|
||||
|
||||
template <int... Ks, int... Vs, int... Is>
|
||||
struct SortByKey<seq<Ks...>, seq<Vs...>, seq<Is...>> {
|
||||
static_assert(sizeof...(Ks) == sizeof...(Vs));
|
||||
static_assert(sizeof...(Ks) == sizeof...(Is));
|
||||
static constexpr cute::array<kvpair,sizeof...(Is)> orig_array = {kvpair{Ks,Vs}...};
|
||||
static constexpr cute::array<kvpair,sizeof...(Is)> sort_array = exchange_sort(orig_array);
|
||||
using key_type = seq<sort_array[Is].key...>;
|
||||
using val_type = seq<sort_array[Is].val...>;
|
||||
};
|
||||
|
||||
} // end namespace detail
|
||||
|
||||
//
|
||||
// Converters and constructors with arrays and params
|
||||
//
|
||||
|
||||
+164
-104
@@ -627,27 +627,37 @@ depth(Layout<Shape,Stride> const& layout)
|
||||
return depth(shape<Is...>(layout));
|
||||
}
|
||||
|
||||
// Return the coprofile of a mode as a tuple of _0s
|
||||
// @post congruent(coprofile(@a layout), @a layout(i)) for any i
|
||||
// @return T Tuple that is congruent with the codomain of @a a.
|
||||
template <int... Is, class Shape, class Stride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
coprofile(Layout<Shape,Stride> const& layout)
|
||||
{
|
||||
return repeat_like(as_arithmetic_tuple(sum(stride<Is...>(layout))), Int<0>{});
|
||||
}
|
||||
|
||||
// Return the codomain shape of a mode
|
||||
// @post size(coshape(@a a)) == cosize(@a a)
|
||||
// @post size(coshape(@a layout)) == cosize(@a layout)
|
||||
// @return C Coordinate with smallest elements such that
|
||||
// @a elem_less(sub_layout(c), C) for all c < size(@a sub_layout)
|
||||
// where sub_layout = get<Is...>(layout).
|
||||
// elem_less(@a sub_layout(c), C) for all c < size(@a sub_layout)
|
||||
// where @a sub_layout = get<Is...>(layout).
|
||||
template <int... Is, class Shape, class Stride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
coshape(Layout<Shape,Stride> const& layout)
|
||||
{
|
||||
// Protect against negative strides
|
||||
auto abs_sub_layout = make_layout(shape<Is...>(layout),
|
||||
transform_leaf(stride<Is...>(layout), abs_fn{}));
|
||||
auto co_coord = as_arithmetic_tuple(abs_sub_layout(size(abs_sub_layout) - Int<1>{}));
|
||||
return co_coord + repeat_like(co_coord, Int<1>{});
|
||||
auto m1_shapes = transform_leaf( shape<Is...>(layout), [](auto s) { return s - Int<1>{}; });
|
||||
auto abs_strides = transform_leaf(stride<Is...>(layout), abs_fn{});
|
||||
auto co_coord = as_arithmetic_tuple(inner_product(m1_shapes, abs_strides));
|
||||
return transform_leaf(co_coord, [](auto c) { return c + Int<1>{}; });
|
||||
}
|
||||
|
||||
// Return the codomain size of a mode
|
||||
// @return M smallest integer such that
|
||||
// @a sub_layout(c) < M for all c < size(@a sub_layout)
|
||||
// where sub_layout = get<Is...>(layout).
|
||||
// size(@a sub_layout(c)) < M for all c < size(@a sub_layout)
|
||||
// where @a sub_layout = get<Is...>(layout).
|
||||
template <int... Is, class Shape, class Stride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
@@ -1019,61 +1029,93 @@ auto
|
||||
composition_impl(LShape const& lhs_shape, LStride const& lhs_stride,
|
||||
RShape const& rhs_shape, RStride const& rhs_stride)
|
||||
{
|
||||
if constexpr (is_tuple<RShape>::value) {
|
||||
// Apply the right-distributivity of Layout composition
|
||||
if constexpr (is_tuple<RShape>::value) { // Right-distributivity of Layout composition for RHS tuple
|
||||
return transform_layout(rhs_shape, rhs_stride, [&](auto const& s, auto const& d) {
|
||||
return composition_impl(lhs_shape, lhs_stride, s, d);
|
||||
});
|
||||
} else
|
||||
if constexpr (is_scaled_basis<RStride>::value) {
|
||||
// Special case for a ScaledBasis stride
|
||||
if constexpr (is_scaled_basis<RStride>::value) { // Special case for a RHS ScaledBasis stride
|
||||
return composition_impl(basis_get(rhs_stride, lhs_shape), basis_get(rhs_stride, lhs_stride),
|
||||
rhs_shape, basis_value(rhs_stride));
|
||||
} else
|
||||
if constexpr (is_constant<0, RStride>::value) {
|
||||
// Special case shortcut for any static stride-0
|
||||
if constexpr (is_constant<0, RStride>::value) { // Special case shortcut for any RHS static stride-0
|
||||
return Layout<RShape, RStride>{rhs_shape, rhs_stride};
|
||||
} else
|
||||
if constexpr (is_integral<decltype(lhs_shape)>::value) {
|
||||
// Special case shortcut for any integral LShape
|
||||
if constexpr (is_integral<LShape>::value) { // Special case shortcut for any LHS integral shape
|
||||
return Layout{rhs_shape, rhs_stride * lhs_stride};
|
||||
} else
|
||||
if constexpr (is_constant<1, RStride>::value) {
|
||||
// Special case shortcut for any static stride-1
|
||||
constexpr int R = rank_v<LShape>;
|
||||
auto result_shape_0 = take<0,R-1>(lhs_shape);
|
||||
} else { // General case: LHS tuple, RHS integral
|
||||
constexpr int R = tuple_size<LShape>::value;
|
||||
|
||||
// Mod out the rhs_shape from the lhs_shape
|
||||
auto [result_shape_1, rest_shape] = fold(result_shape_0, cute::make_tuple(cute::make_tuple(), rhs_shape),
|
||||
[] (auto const& init, auto const& si) {
|
||||
return cute::make_tuple(append(get<0>(init), shape_min(abs(si), get<1>(init))), shape_div(get<1>(init), abs(si)));
|
||||
});
|
||||
auto [result_shape, result_stride, rest_shape, rest_stride] =
|
||||
cute::fold(make_seq<R-1>{}, // t = [0,1,2,...,R-1)
|
||||
cute::make_tuple(cute::tuple<>{}, // v = (result_shape,
|
||||
cute::tuple<>{}, // result_stride,
|
||||
rhs_shape, // rest_shape:Integral,
|
||||
rhs_stride), // rest_stride:Integral)
|
||||
[&](auto const& init, auto curr_i) { // f(v,t) -> v'
|
||||
// Can ICE on some compilers
|
||||
//auto [result_shape, result_stride, rest_shape, rest_stride] = init;
|
||||
//auto [curr_shape, curr_stride] = curr;
|
||||
// Unpack inputs
|
||||
auto result_shape = get<0>(init);
|
||||
auto result_stride = get<1>(init);
|
||||
auto rest_shape = get<2>(init);
|
||||
auto rest_stride = get<3>(init);
|
||||
|
||||
// Jump into coalesce and append (rest_shape, get<R-1>(lhs_stride))
|
||||
return detail::bw_coalesce<R-2>(result_shape_1, lhs_stride, rest_shape, get<R-1>(lhs_stride));
|
||||
} else {
|
||||
// General case: integral RShape and RStride, tuple LShape and LStride
|
||||
constexpr int R = rank_v<LShape>;
|
||||
auto result_shape_0 = take<0,R-1>(lhs_shape);
|
||||
auto result_stride_0 = take<0,R-1>(lhs_stride);
|
||||
auto curr_shape = get<curr_i>(lhs_shape);
|
||||
auto curr_stride = get<curr_i>(lhs_stride);
|
||||
|
||||
// Divide out the rhs_stride from the lhs_shape
|
||||
auto [result_shape_1, rest_stride] = fold(result_shape_0, cute::make_tuple(cute::make_tuple(), rhs_stride),
|
||||
[] (auto const& init, auto const& di) {
|
||||
return cute::make_tuple(append(get<0>(init), shape_div(di, get<1>(init))), shape_div(get<1>(init), di));
|
||||
});
|
||||
// Strong divisibility condition -- requires composition to be statically verifiable.
|
||||
//CUTE_STATIC_ASSERT_V(((rest_stride % curr_shape) == Int<0>{}) or (rest_stride < curr_shape), "Stride Divisibility Condition");
|
||||
|
||||
// Apply any lhs_shape changes to the stride
|
||||
auto result_stride_1 = elem_scale(result_stride_0, shape_div(result_shape_0, result_shape_1));
|
||||
// Weak divisibility condition -- verify the divisibility condition whenever possible
|
||||
if constexpr (is_static<decltype(curr_shape)>::value and is_static<decltype(rest_stride)>::value) {
|
||||
CUTE_STATIC_ASSERT_V(((rest_stride % curr_shape) == Int<0>{}) or (rest_stride < curr_shape), "Stride Divisibility Condition");
|
||||
} else {
|
||||
// DEBUG assert can cause extra registers and inappropriate compile-time/run-time failure
|
||||
//assert((((rest_stride % curr_shape) == 0) or (rest_stride < curr_shape)) && "Stride Divisibility Condition");
|
||||
}
|
||||
|
||||
// Mod out the rhs_shape from the lhs_shape
|
||||
auto [result_shape_2, rest_shape] = fold(result_shape_1, cute::make_tuple(cute::make_tuple(), rhs_shape),
|
||||
[] (auto const& init, auto const& si) {
|
||||
return cute::make_tuple(append(get<0>(init), shape_min(abs(si), get<1>(init))), shape_div(get<1>(init), abs(si)));
|
||||
});
|
||||
// next_shape: ceil(exclusive_prefix_product<r>(lhs_shape) / rhs_stride)
|
||||
[[maybe_unused]] auto next_shape = cute::ceil_div(curr_shape, abs(rest_stride));
|
||||
// next_stride: ceil(rhs_stride / exclusive_prefix_product<r>(lhs_shape))
|
||||
[[maybe_unused]] auto next_stride = cute::ceil_div(abs(rest_stride), curr_shape) * signum(rest_stride);
|
||||
|
||||
// Jump into coalesce and append (rest_shape, rest_stride * get<R-1>(lhs_stride))
|
||||
return detail::bw_coalesce<R-2>(result_shape_2, result_stride_1, rest_shape, rest_stride * get<R-1>(lhs_stride));
|
||||
if constexpr (is_constant<1, decltype(next_shape)>::value or is_constant<1, decltype(rest_shape)>::value) {
|
||||
return cute::make_tuple(result_shape,
|
||||
result_stride,
|
||||
rest_shape,
|
||||
next_stride);
|
||||
} else {
|
||||
auto new_shape = cute::min(next_shape, rest_shape);
|
||||
|
||||
// Strong divisibility condition
|
||||
//CUTE_STATIC_ASSERT_V(((rest_shape % new_shape) == Int<0>{}), "Shape Divisibility Condition");
|
||||
|
||||
// Weak divisibility condition
|
||||
if constexpr (is_static<decltype(new_shape)>::value and is_static<decltype(rest_shape)>::value) {
|
||||
CUTE_STATIC_ASSERT_V(((rest_shape % new_shape) == Int<0>{}), "Shape Divisibility Condition");
|
||||
} else {
|
||||
// DEBUG assert can cause extra registers and inappropriate compile-time/run-time failure
|
||||
//assert(((rest_shape % new_shape) == 0) && "Shape Divisibility Condition");
|
||||
}
|
||||
|
||||
return cute::make_tuple(append(result_shape, new_shape),
|
||||
append(result_stride, rest_stride * curr_stride),
|
||||
rest_shape / new_shape,
|
||||
next_stride);
|
||||
}
|
||||
});
|
||||
|
||||
if constexpr (tuple_size<decltype(result_shape)>::value == 0) {
|
||||
return Layout{rest_shape, rest_stride * get<R-1>(lhs_stride)};
|
||||
} else
|
||||
if constexpr (is_constant<1, decltype(rest_shape)>::value) {
|
||||
return Layout{unwrap(result_shape), unwrap(result_stride)};
|
||||
} else {
|
||||
return Layout{append(result_shape, rest_shape),
|
||||
append(result_stride, rest_stride * get<R-1>(lhs_stride))};
|
||||
}
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
@@ -1088,8 +1130,7 @@ auto
|
||||
composition(Layout<LShape,LStride> const& lhs,
|
||||
Layout<RShape,RStride> const& rhs)
|
||||
{
|
||||
auto coprofile = repeat_like(decltype(coshape(rhs)){}, Int<0>{});
|
||||
auto flat_lhs = detail::coalesce_x(lhs, coprofile);
|
||||
auto flat_lhs = detail::coalesce_x(lhs, coprofile(rhs));
|
||||
return detail::composition_impl(flat_lhs.shape(), flat_lhs.stride(), rhs.shape(), rhs.stride());
|
||||
}
|
||||
|
||||
@@ -1203,37 +1244,6 @@ complement(Layout<Shape,Stride> const& layout)
|
||||
// Right-Inverse and Left-Inverse
|
||||
//
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <int NextStride, class Shape, class Stride, int... Is>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
inverse_seq(Shape const& shape, Stride const& stride, seq<Is...>)
|
||||
{
|
||||
auto next_I = cute::find_if(stride, [](auto a) { return is_constant<NextStride, decltype(a)>{}; });
|
||||
|
||||
if constexpr (next_I == decltype(rank(stride))::value) {
|
||||
// If not found, return current seq
|
||||
return seq<Is...>{};
|
||||
} else {
|
||||
// auto next_stride = get<next_I>(shape) * get<next_I>(stride);
|
||||
// NOTE: Needed for g++-7
|
||||
using next_stride = decltype(get<next_I>(shape) * get<next_I>(stride));
|
||||
|
||||
if constexpr (is_static<next_stride>::value && !is_constant<NextStride, next_stride>::value) {
|
||||
// If next_stride is static and unique, then continue
|
||||
return inverse_seq<next_stride::value>(shape, stride, seq<Is..., next_I>{});
|
||||
} else {
|
||||
// Else return current seq + next_I
|
||||
return seq<Is..., next_I>{};
|
||||
}
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
} // end namespace detail
|
||||
|
||||
//
|
||||
// Build the right-inverse of a layout
|
||||
// @pre is_static<Layout>
|
||||
@@ -1248,22 +1258,40 @@ CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
right_inverse(Layout<Shape,Stride> const& layout)
|
||||
{
|
||||
auto flat_layout = coalesce(layout);
|
||||
auto astride = transform_leaf(flat_layout.stride(), abs_fn{});
|
||||
// Flatten and filter shape-1
|
||||
auto clayout = coalesce(layout);
|
||||
auto lstride = wrap(clayout.stride());
|
||||
auto lshape = wrap(clayout.shape());
|
||||
|
||||
// Find Int<1>{}, the starting stride, and follow the strides to gen inverse_seq
|
||||
[[maybe_unused]] auto iseq = detail::inverse_seq<1>(flat_layout.shape(), astride, seq<>{});
|
||||
// Prefix product of the shape
|
||||
auto preprod_shape = cute::fold(lshape, cute::tuple<_1>{}, [](auto c, auto vi) { return append(c, vi*back(c)); });
|
||||
|
||||
if constexpr (iseq.size() == 0) {
|
||||
return Layout<_1,_0>{}; // Empty case, nothing found
|
||||
} else {
|
||||
// Generate the corresponding new strides and construct
|
||||
auto rstride = compact_major<LayoutLeft>(flat_layout.shape());
|
||||
return make_layout(unwrap(transform(iseq, [&](auto i) { return shape<i>(flat_layout); })),
|
||||
unwrap(transform(iseq, [&](auto i) { return signum(stride<i>(flat_layout)) * get<i>(rstride); })));
|
||||
}
|
||||
// Filter out any dynamic strides
|
||||
[[maybe_unused]] auto filtered_seq = filter_tuple(make_seq<rank(lstride)>{}, lstride, [](auto i, auto d) {
|
||||
return conditional_return<is_static_v<decltype(d)>>(cute::tuple{i}, cute::tuple<>{}); });
|
||||
[[maybe_unused]] auto filtered_stride = transform(filtered_seq, [&](auto i) { return get<i>(lstride); });
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
// Sort by strides
|
||||
using Sorted = detail::SortByKey<decltype(filtered_stride), decltype(filtered_seq)>;
|
||||
auto sorted_seq = typename Sorted::val_type{};
|
||||
//auto sorted_stride = typename Sorted::key_type{};
|
||||
|
||||
auto [result_shape, result_stride, curr] = cute::fold(sorted_seq, tuple<tuple<_1>,tuple<_0>,_1>{},
|
||||
[&](auto const& init, auto i) {
|
||||
[[maybe_unused]] auto ishape = get<i>(lshape);
|
||||
[[maybe_unused]] auto istride = get<i>(lstride);
|
||||
[[maybe_unused]] auto curr_stride = get<2>(init);
|
||||
|
||||
if constexpr (is_constant<decltype(istride)::value, decltype(curr_stride)>::value) {
|
||||
return make_tuple(append(get<0>(init), ishape), // result_shape
|
||||
append(get<1>(init), get<i>(preprod_shape)), // result_stride
|
||||
ishape * istride);
|
||||
} else {
|
||||
return init;
|
||||
}
|
||||
});
|
||||
|
||||
return coalesce(make_layout(result_shape, result_stride));
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -1274,13 +1302,12 @@ right_inverse(Underscore const& _)
|
||||
}
|
||||
|
||||
//
|
||||
// Build the left-inverse of a layout
|
||||
// Build the quasi-inverse of a layout (left-inverse when layout is injective)
|
||||
// @pre is_static<Layout>
|
||||
// @pre @a layout is an injective function
|
||||
// @result A layout @a result such that
|
||||
// @a result(@a layout(i)) == i for all i < size(@a layout)
|
||||
// @a layout(@a result(@a layout(i))) == @a layout(i) for all i < size(@a layout)
|
||||
// @result A layout @a result such that
|
||||
// composition(@a result, @a layout) is identical to make_layout(shape(layout))
|
||||
// composition(@layout, composition(@a result, @a layout)) is identical to @a layout
|
||||
//
|
||||
|
||||
template <class Shape, class Stride>
|
||||
@@ -1288,7 +1315,39 @@ CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
left_inverse(Layout<Shape,Stride> const& layout)
|
||||
{
|
||||
return right_inverse(make_layout(layout, complement(layout)));
|
||||
// Flatten and filter shape-1
|
||||
auto clayout = coalesce(layout);
|
||||
auto lstride = wrap(clayout.stride());
|
||||
auto lshape = wrap(clayout.shape());
|
||||
|
||||
// Prefix product of the shape
|
||||
auto preprod_shape = cute::fold(lshape, cute::tuple<_1>{}, [](auto c, auto vi) { return append(c, vi*back(c)); });
|
||||
|
||||
// Sort by strides
|
||||
static_assert(is_static<decltype(lstride)>::value, "Left inverse requires static strides.");
|
||||
using Sorted = detail::SortByKey<decltype(lstride), tuple_seq<decltype(lstride)>>;
|
||||
auto sorted_seq = typename Sorted::val_type{};
|
||||
//auto sorted_stride = typename Sorted::key_type{};
|
||||
|
||||
auto [result_shape, result_stride] = cute::fold(sorted_seq, tuple<tuple<>,tuple<_0>>{},
|
||||
[&](auto const& init, auto i) {
|
||||
[[maybe_unused]] auto istride = get<i>(lstride);
|
||||
|
||||
if constexpr (is_constant<0, decltype(istride)>::value) {
|
||||
return init;
|
||||
} else {
|
||||
auto result_shape = get<0>(init);
|
||||
auto result_stride = get<1>(init);
|
||||
|
||||
CUTE_STATIC_ASSERT_V((istride % size(result_shape)) == Int<0>{}, "Left inverse divisibility condition");
|
||||
|
||||
return make_tuple(append(result_shape, istride / size(result_shape)),
|
||||
append(result_stride, get<i>(preprod_shape)));
|
||||
}
|
||||
});
|
||||
|
||||
return coalesce(make_layout(append(result_shape, get<back(sorted_seq)>(lshape)),
|
||||
result_stride));
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -1506,7 +1565,7 @@ auto
|
||||
logical_divide(Layout<LShape,LStride> const& layout,
|
||||
Layout<TShape,TStride> const& tiler)
|
||||
{
|
||||
return composition(layout, make_layout(tiler, complement(tiler, shape(layout))));
|
||||
return composition(layout, make_layout(tiler, complement(tiler, shape(coalesce(layout)))));
|
||||
}
|
||||
|
||||
template <class LShape, class LStride, class Tiler>
|
||||
@@ -1760,10 +1819,11 @@ upcast(Shape const& shape, Stride const& stride)
|
||||
} else if constexpr (is_constant<0, Stride>::value) { // static-0 stride
|
||||
return Layout<Shape,Stride>{shape,stride};
|
||||
} else if constexpr (is_static<Stride>::value) { // static stride
|
||||
return make_layout(shape_div(shape, shape_div(Int<N>{}, abs(stride))),
|
||||
shape_div(stride, Int<N>{}));
|
||||
static_assert(Stride::value % N == 0 or N % Stride::value == 0, "Divisibility condition");
|
||||
return make_layout(ceil_div(shape, ceil_div(Int<N>{}, abs(stride))),
|
||||
signum(stride) * ceil_div(abs(stride), Int<N>{}));
|
||||
} else { // dynamic stride
|
||||
// assume dynamic strides are larger than N and divisible
|
||||
// Assume dynamic strides are larger than N and divisible
|
||||
// assert(stride % N == 0);
|
||||
return make_layout(shape, safe_div(stride, Int<N>{}));
|
||||
}
|
||||
|
||||
@@ -37,7 +37,7 @@
|
||||
/* This implements a ComposedLayout of the form
|
||||
* LayoutA o Offset o LayoutB
|
||||
* and is useful in cases where composition() does not or cannot apply to LayoutA and LayoutB.
|
||||
* For example, when the "divisibility condition" in shape_div is violated in composition(LayoutA, LayoutB).
|
||||
* For example, when the "divisibility condition" is violated in composition(LayoutA, LayoutB).
|
||||
*
|
||||
* This ComposedLayout provides similar functionality to Layout including tiling, partitioning,
|
||||
* coordinate-to-index mapping and layout manipulations, but is not considered a "normal" layout.
|
||||
|
||||
@@ -370,12 +370,21 @@ safe_div(ScaledBasis<T,M> const& b, U const& u)
|
||||
template <class T, int M, class U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
shape_div(ScaledBasis<T,M> const& b, U const& u)
|
||||
ceil_div(ScaledBasis<T,M> const& b, U const& u)
|
||||
{
|
||||
auto t = shape_div(b.value(), u);
|
||||
auto t = ceil_div(b.value(), u);
|
||||
return ScaledBasis<decltype(t),M>{t};
|
||||
}
|
||||
|
||||
template <class T, int N>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
abs(ScaledBasis<T,N> const& e)
|
||||
{
|
||||
auto t = abs(e.value());
|
||||
return ScaledBasis<decltype(t),N>{t};
|
||||
}
|
||||
|
||||
// Equality
|
||||
template <class T, int N, class U, int M>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -399,14 +408,6 @@ operator==(T const&, ScaledBasis<U,M> const&) {
|
||||
return {};
|
||||
}
|
||||
|
||||
// Abs
|
||||
template <class T, int N>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
abs(ScaledBasis<T,N> const& e) {
|
||||
return ScaledBasis<decltype(abs(e.value())),N>{abs(e.value())};
|
||||
}
|
||||
|
||||
// Multiplication
|
||||
template <class A, class T, int N>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
|
||||
@@ -124,6 +124,31 @@ using tuple_seq = make_seq<tuple_size<remove_cvref_t<Tuple>>::value>;
|
||||
template <class Tuple>
|
||||
using tuple_rseq = make_rseq<tuple_size<remove_cvref_t<Tuple>>::value>;
|
||||
|
||||
//
|
||||
// Convert a parameter pack to an int sequence
|
||||
//
|
||||
|
||||
template <class T>
|
||||
struct to_seq;
|
||||
|
||||
template <>
|
||||
struct to_seq<integer_sequence<int>> {
|
||||
using type = seq<>;
|
||||
};
|
||||
|
||||
template <int I, int... Is>
|
||||
struct to_seq<integer_sequence<int, I, Is...>> {
|
||||
using type = seq<I, Is...>;
|
||||
};
|
||||
|
||||
template <template <class...> class TupleLike, class... Ts>
|
||||
struct to_seq<TupleLike<Ts...>> {
|
||||
using type = seq<Ts::value...>;
|
||||
};
|
||||
|
||||
template <class T>
|
||||
using to_seq_t = typename to_seq<T>::type;
|
||||
|
||||
//
|
||||
// Specialize cute::tuple-traits for std::integer_sequence
|
||||
//
|
||||
|
||||
@@ -74,29 +74,40 @@ struct integral_constant : C<v> {
|
||||
|
||||
// Use cute::is_std_integral<T> to match built-in integral types (int, int64_t, unsigned, etc)
|
||||
// Use cute::is_integral<T> to match both built-in integral types AND static integral types.
|
||||
|
||||
template <class T>
|
||||
struct is_integral : bool_constant<is_std_integral<T>::value> {};
|
||||
template <auto v>
|
||||
struct is_integral<C<v> > : true_type {};
|
||||
template <class T, T v>
|
||||
struct is_integral<integral_constant<T,v>> : true_type {};
|
||||
template <class T>
|
||||
constexpr bool is_integral_v = is_integral<T>::value;
|
||||
|
||||
// Register FastDivmod as the integral type
|
||||
// Register FastDivmod as integral type
|
||||
template<>
|
||||
struct is_integral<cutlass::FastDivmod> : true_type {};
|
||||
|
||||
// is_static detects if an (abstract) value is defined completely by its type (no members)
|
||||
template <class T>
|
||||
struct is_static : bool_constant<is_empty<remove_cvref_t<T>>::value> {};
|
||||
|
||||
struct is_static : bool_constant<is_empty<T>::value> {};
|
||||
template <class T>
|
||||
struct is_static<T const > : is_static<T> {};
|
||||
template <class T>
|
||||
struct is_static<T const&> : is_static<T> {};
|
||||
template <class T>
|
||||
struct is_static<T &> : is_static<T> {};
|
||||
template <class T>
|
||||
struct is_static<T &&> : is_static<T> {};
|
||||
template <class T>
|
||||
constexpr bool is_static_v = is_static<T>::value;
|
||||
|
||||
// is_constant detects if a type is a static integral type and if v is equal to a value
|
||||
|
||||
template <auto n, class T>
|
||||
struct is_constant : false_type {};
|
||||
template <auto n, auto v>
|
||||
struct is_constant<n, C<v> > : bool_constant<v == n> {};
|
||||
template <auto n, class T, T v>
|
||||
struct is_constant<n, integral_constant<T,v>> : bool_constant<v == n> {};
|
||||
template <auto n, class T>
|
||||
struct is_constant<n, T const > : is_constant<n,T> {};
|
||||
template <auto n, class T>
|
||||
@@ -105,10 +116,8 @@ template <auto n, class T>
|
||||
struct is_constant<n, T &> : is_constant<n,T> {};
|
||||
template <auto n, class T>
|
||||
struct is_constant<n, T &&> : is_constant<n,T> {};
|
||||
template <auto n, auto v>
|
||||
struct is_constant<n, C<v> > : bool_constant<v == n> {};
|
||||
template <auto n, class T, T v>
|
||||
struct is_constant<n, integral_constant<T,v>> : bool_constant<v == n> {};
|
||||
template <auto n, class T>
|
||||
constexpr bool is_constant_v = is_constant<n,T>::value;
|
||||
|
||||
//
|
||||
// Specializations
|
||||
|
||||
@@ -573,8 +573,8 @@ logical_product(Layout<Shape,Stride> const& layout,
|
||||
auto active_Y = swizzle_active_bits & typename Swizzle<B,M,S>::yyy_msk{};
|
||||
|
||||
// Pass the identifiers through the old layout and new layout to make a new swizzle identifier, L*(L[(P o L)(c*)])
|
||||
auto new_active_Z = new_layout(Int<0>{}, tiler.layout_b()[active_Z]);
|
||||
auto new_active_Y = new_layout(Int<0>{}, tiler.layout_b()[active_Y]);
|
||||
auto new_active_Z = new_layout(Int<0>{}, tiler.layout_b()(active_Z));
|
||||
auto new_active_Y = new_layout(Int<0>{}, tiler.layout_b()(active_Y));
|
||||
|
||||
// Use this new swizzle identifier to construxt the new swizzle for new_layout
|
||||
// (this also makes sure it's a "valid" swizzle that Swizzle can represent)
|
||||
|
||||
@@ -481,7 +481,7 @@ CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_counting_tensor(Layout const& layout)
|
||||
{
|
||||
return make_tensor(make_inttuple_iter(repeat_like(coshape(layout), Int<0>{})), layout);
|
||||
return make_tensor(make_inttuple_iter(coprofile(layout)), layout);
|
||||
}
|
||||
|
||||
//
|
||||
@@ -788,7 +788,7 @@ recast(Tensor&& tensor)
|
||||
* vectorization should be attempted.
|
||||
*
|
||||
* Note that the return value does NOT include alignment concerns such as the pointer value and
|
||||
* the divisbility of dynamic strides.
|
||||
* the divisibility of dynamic strides.
|
||||
*/
|
||||
template <class SrcEngine, class SrcLayout,
|
||||
class DstEngine, class DstLayout>
|
||||
@@ -828,7 +828,7 @@ max_common_vector(Tensor<SrcEngine,SrcLayout> const& a,
|
||||
* are both identity Layouts.
|
||||
*
|
||||
* Note that the returned layout does NOT include alignment concerns such as the pointer value and
|
||||
* the divisbility of dynamic strides.
|
||||
* the divisibility of dynamic strides.
|
||||
*/
|
||||
template <class SrcEngine, class SrcLayout,
|
||||
class DstEngine, class DstLayout>
|
||||
|
||||
@@ -154,20 +154,23 @@ using CUTE_STL_NAMESPACE::is_pointer_v;
|
||||
using CUTE_STL_NAMESPACE::declval;
|
||||
|
||||
template <class T>
|
||||
constexpr T&& forward(remove_reference_t<T>& t) noexcept
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
T&& forward(remove_reference_t<T>& t) noexcept
|
||||
{
|
||||
return static_cast<T&&>(t);
|
||||
}
|
||||
|
||||
template <class T>
|
||||
constexpr T&& forward(remove_reference_t<T>&& t) noexcept
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
T&& forward(remove_reference_t<T>&& t) noexcept
|
||||
{
|
||||
static_assert(! is_lvalue_reference_v<T>, "T cannot be an lvalue reference (e.g., U&).");
|
||||
return static_cast<T&&>(t);
|
||||
}
|
||||
|
||||
template <class T>
|
||||
constexpr remove_reference_t<T>&& move(T&& t) noexcept
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
remove_reference_t<T>&& move(T&& t) noexcept
|
||||
{
|
||||
return static_cast<remove_reference_t<T>&&>(t);
|
||||
}
|
||||
@@ -220,8 +223,6 @@ struct tuple_size;
|
||||
template <class T>
|
||||
struct tuple_size<T,void_t<typename CUTE_STL_NAMESPACE::tuple_size<T>::type>> : CUTE_STL_NAMESPACE::integral_constant<size_t, CUTE_STL_NAMESPACE::tuple_size<T>::value> {};
|
||||
|
||||
// S = : std::integral_constant<std::size_t, std::tuple_size<T>::value> {};
|
||||
|
||||
template <class T>
|
||||
constexpr size_t tuple_size_v = tuple_size<T>::value;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user