CUTLASS 3.8 Release (#2059)

* CUTLASS 3.8 Release

* update

* Update README.md

* Revert "Update README.md"

This reverts commit b353e36fe83e0815f99b44e46c0c95494c44726b.

* update

* update

---------

Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com>
Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
mihir-awatramani
2025-01-25 02:44:06 -05:00
committed by GitHub
co-authored by Haicheng Wu Haicheng Wu
parent 9eb01fa0b0
commit 389e493055
290 changed files with 91222 additions and 291 deletions
+267
View File
@@ -334,6 +334,7 @@ static struct {
OperationKind_enumerants[] = {
{"eq_gemm", "EqGemm", OperationKind::kEqGemm},
{"gemm", "Gemm", OperationKind::kGemm},
{"block_scaled_gemm", "blockScaledGemm", OperationKind::kBlockScaledGemm},
{"rank_k", "RankK", OperationKind::kRankK},
{"rank_2k", "Rank2K", OperationKind::kRank2K},
{"trmm", "Trmm", OperationKind::kTrmm},
@@ -422,6 +423,53 @@ Status from_string<Status>(std::string const &str) {
///////////////////////////////////////////////////////////////////////////////////////////////////
static struct {
char const *text;
char const *pretty;
RuntimeDatatype enumerant;
}
RuntimeDatatype_enumerants[] = {
{"e4m3", "<e4m3>", RuntimeDatatype::kE4M3},
{"e5m2", "<e5m2>", RuntimeDatatype::kE5M2},
{"e3m2", "<e3m2>", RuntimeDatatype::kE3M2},
{"e2m3", "<e2m3>", RuntimeDatatype::kE2M3},
{"e2m1", "<e2m1>", RuntimeDatatype::kE2M1}
};
/// Converts a RuntimeDatatype enumerant to a string
char const *to_string(RuntimeDatatype type, bool pretty) {
for (auto const & possible : RuntimeDatatype_enumerants) {
if (type == possible.enumerant) {
if (pretty) {
return possible.pretty;
}
else {
return possible.text;
}
}
}
return pretty ? "Invalid" : "invalid";
}
/// Converts a RuntimeDatatype enumerant from a string
template <>
RuntimeDatatype from_string<RuntimeDatatype>(std::string const &str) {
for (auto const & possible : RuntimeDatatype_enumerants) {
if ((str.compare(possible.text) == 0) ||
(str.compare(possible.pretty) == 0)) {
return possible.enumerant;
}
}
return RuntimeDatatype::kInvalid;
}
///////////////////////////////////////////////////////////////////////////////////////////////////
static struct {
@@ -447,6 +495,16 @@ NumericTypeID_enumerants[] = {
{"s64", "S64", NumericTypeID::kS64},
{"fe4m3", "FE4M3", NumericTypeID::kFE4M3},
{"fe5m2", "FE5M2", NumericTypeID::kFE5M2},
{"f8", "F8", NumericTypeID::kF8},
{"f6", "F6", NumericTypeID::kF6},
{"f4", "F4", NumericTypeID::kF4},
{"fe2m3", "FE2M3", NumericTypeID::kFE2M3},
{"fe3m2", "FE3M2", NumericTypeID::kFE3M2},
{"fe2m1", "FE2M1", NumericTypeID::kFE2M1},
{"fue8m0", "FUE8M0", NumericTypeID::kFUE8M0},
{"fue4m3", "FUE4M3", NumericTypeID::kFUE4M3},
{"f16", "F16", NumericTypeID::kF16},
{"bf16", "BF16", NumericTypeID::kBF16},
{"f32", "F32", NumericTypeID::kF32},
@@ -510,6 +568,16 @@ int sizeof_bits(NumericTypeID type) {
switch (type) {
case NumericTypeID::kFE4M3: return 8;
case NumericTypeID::kFE5M2: return 8;
case NumericTypeID::kF8: return 8;
case NumericTypeID::kF6: return 6;
case NumericTypeID::kF4: return 4;
case NumericTypeID::kFE2M3: return 6;
case NumericTypeID::kFE3M2: return 6;
case NumericTypeID::kFE2M1: return 4;
case NumericTypeID::kFUE8M0: return 8;
case NumericTypeID::kFUE4M3: return 8;
case NumericTypeID::kF16: return 16;
case NumericTypeID::kBF16: return 16;
case NumericTypeID::kTF32: return 32;
@@ -589,6 +657,16 @@ bool is_signed_type(NumericTypeID type) {
switch (type) {
case NumericTypeID::kFE4M3: return true;
case NumericTypeID::kFE5M2: return true;
case NumericTypeID::kF8: return true;
case NumericTypeID::kF6: return true;
case NumericTypeID::kF4: return true;
case NumericTypeID::kFE2M3: return true;
case NumericTypeID::kFE3M2: return true;
case NumericTypeID::kFE2M1: return true;
case NumericTypeID::kFUE8M0: return false;
case NumericTypeID::kFUE4M3: return false;
case NumericTypeID::kF16: return true;
case NumericTypeID::kBF16: return true;
case NumericTypeID::kTF32: return true;
@@ -620,6 +698,16 @@ bool is_float_type(NumericTypeID type) {
switch (type) {
case NumericTypeID::kFE4M3: return true;
case NumericTypeID::kFE5M2: return true;
case NumericTypeID::kF8: return true;
case NumericTypeID::kF6: return true;
case NumericTypeID::kF4: return true;
case NumericTypeID::kFE2M3: return true;
case NumericTypeID::kFE3M2: return true;
case NumericTypeID::kFE2M1: return true;
case NumericTypeID::kFUE8M0: return true;
case NumericTypeID::kFUE4M3: return true;
case NumericTypeID::kF16: return true;
case NumericTypeID::kBF16: return true;
case NumericTypeID::kTF32: return true;
@@ -1168,6 +1256,43 @@ bool lexical_cast(std::vector<uint8_t> &bytes, NumericTypeID type, std::string c
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(tmp);
}
break;
case NumericTypeID::kFE2M3:
{
float tmp;
ss >> tmp;
*reinterpret_cast<float_e2m3_t *>(bytes.data()) = static_cast<float_e2m3_t>(tmp);
}
break;
case NumericTypeID::kFE3M2:
{
float tmp;
ss >> tmp;
*reinterpret_cast<float_e3m2_t *>(bytes.data()) = static_cast<float_e3m2_t>(tmp);
}
break;
case NumericTypeID::kFE2M1:
{
float tmp;
ss >> tmp;
*reinterpret_cast<float_e2m1_t *>(bytes.data()) = static_cast<float_e2m1_t>(tmp);
}
break;
case NumericTypeID::kFUE8M0:
{
float tmp;
ss >> tmp;
*reinterpret_cast<float_ue8m0_t *>(bytes.data()) = static_cast<float_ue8m0_t>(tmp);
}
break;
case NumericTypeID::kFUE4M3:
{
float tmp;
ss >> tmp;
*reinterpret_cast<float_ue4m3_t *>(bytes.data()) = static_cast<float_ue4m3_t>(tmp);
}
break;
case NumericTypeID::kF16:
{
float tmp;
@@ -1317,6 +1442,38 @@ std::string lexical_cast(std::vector<uint8_t> &bytes, NumericTypeID type) {
ss << tmp;
}
break;
case NumericTypeID::kFE2M3:
{
float tmp = *reinterpret_cast<float_e2m3_t *>(bytes.data());
ss << tmp;
}
break;
case NumericTypeID::kFE3M2:
{
float tmp = *reinterpret_cast<float_e3m2_t *>(bytes.data());
ss << tmp;
}
break;
case NumericTypeID::kFE2M1:
{
float tmp = *reinterpret_cast<float_e2m1_t *>(bytes.data());
ss << tmp;
}
break;
case NumericTypeID::kFUE8M0:
{
float tmp = *reinterpret_cast<float_ue8m0_t *>(bytes.data());
ss << tmp;
}
break;
case NumericTypeID::kFUE4M3:
{
float tmp = *reinterpret_cast<float_ue4m3_t *>(bytes.data());
ss << tmp;
}
break;
case NumericTypeID::kF16:
{
float tmp = *reinterpret_cast<half_t *>(bytes.data());
@@ -1469,6 +1626,33 @@ bool cast_from_int64(std::vector<uint8_t> &bytes, NumericTypeID type, int64_t sr
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(float(src));
}
break;
case NumericTypeID::kFE2M3:
{
*reinterpret_cast<float_e2m3_t *>(bytes.data()) = static_cast<float_e2m3_t>(float(src));
}
break;
case NumericTypeID::kFE3M2:
{
*reinterpret_cast<float_e3m2_t *>(bytes.data()) = static_cast<float_e3m2_t>(float(src));
}
break;
case NumericTypeID::kFE2M1:
{
*reinterpret_cast<float_e2m1_t *>(bytes.data()) = static_cast<float_e2m1_t>(float(src));
}
break;
case NumericTypeID::kFUE8M0:
{
*reinterpret_cast<float_ue8m0_t *>(bytes.data()) = static_cast<float_ue8m0_t>(float(src));
}
break;
case NumericTypeID::kFUE4M3:
{
*reinterpret_cast<float_ue4m3_t *>(bytes.data()) = static_cast<float_ue4m3_t>(float(src));
}
break;
case NumericTypeID::kF16:
{
*reinterpret_cast<half_t *>(bytes.data()) = static_cast<half_t>(float(src));
@@ -1579,6 +1763,33 @@ bool cast_from_uint64(std::vector<uint8_t> &bytes, NumericTypeID type, uint64_t
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(float(src));
}
break;
case NumericTypeID::kFE2M3:
{
*reinterpret_cast<float_e2m3_t *>(bytes.data()) = static_cast<float_e2m3_t>(float(src));
}
break;
case NumericTypeID::kFE3M2:
{
*reinterpret_cast<float_e3m2_t *>(bytes.data()) = static_cast<float_e3m2_t>(float(src));
}
break;
case NumericTypeID::kFE2M1:
{
*reinterpret_cast<float_e2m1_t *>(bytes.data()) = static_cast<float_e2m1_t>(float(src));
}
break;
case NumericTypeID::kFUE8M0:
{
*reinterpret_cast<float_ue8m0_t *>(bytes.data()) = static_cast<float_ue8m0_t>(float(src));
}
break;
case NumericTypeID::kFUE4M3:
{
*reinterpret_cast<float_ue4m3_t *>(bytes.data()) = static_cast<float_ue4m3_t>(float(src));
}
break;
case NumericTypeID::kF16:
{
*reinterpret_cast<half_t *>(bytes.data()) = static_cast<half_t>(float(src));
@@ -1690,6 +1901,33 @@ bool cast_from_double(std::vector<uint8_t> &bytes, NumericTypeID type, double sr
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(float(src));
}
break;
case NumericTypeID::kFE2M3:
{
*reinterpret_cast<float_e2m3_t *>(bytes.data()) = static_cast<float_e2m3_t>(float(src));
}
break;
case NumericTypeID::kFE3M2:
{
*reinterpret_cast<float_e3m2_t *>(bytes.data()) = static_cast<float_e3m2_t>(float(src));
}
break;
case NumericTypeID::kFE2M1:
{
*reinterpret_cast<float_e2m1_t *>(bytes.data()) = static_cast<float_e2m1_t>(float(src));
}
break;
case NumericTypeID::kFUE8M0:
{
*reinterpret_cast<float_ue8m0_t *>(bytes.data()) = static_cast<float_ue8m0_t>(float(src));
}
break;
case NumericTypeID::kFUE4M3:
{
*reinterpret_cast<float_ue4m3_t *>(bytes.data()) = static_cast<float_ue4m3_t>(float(src));
}
break;
case NumericTypeID::kF16:
{
*reinterpret_cast<half_t *>(bytes.data()) = static_cast<half_t>(float(src));
@@ -1751,6 +1989,35 @@ bool cast_from_double(std::vector<uint8_t> &bytes, NumericTypeID type, double sr
return true;
}
NumericTypeID dynamic_datatype_to_id(RuntimeDatatype type) {
NumericTypeID element{};
switch (type) {
case RuntimeDatatype::kE4M3:
element = NumericTypeID::kFE4M3;
break;
case RuntimeDatatype::kE5M2:
element = NumericTypeID::kFE5M2;
break;
case RuntimeDatatype::kE2M3:
element = NumericTypeID::kFE2M3;
break;
case RuntimeDatatype::kE3M2:
element = NumericTypeID::kFE3M2;
break;
case RuntimeDatatype::kE2M1:
element = NumericTypeID::kFE2M1;
break;
default:
assert("illegal runtime datatype!");
break;
}
return element;
}
///////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace library