CUTLASS 3.1 (#915)

Co-authored-by: Aniket Shivam <ashivam@nvidia.com>
This commit is contained in:
ANIKET SHIVAM
2023-04-14 23:19:34 -04:00
committed by GitHub
co-authored by Aniket Shivam
parent 9b8166e3f0
commit d572cc1aab
482 changed files with 37175 additions and 16410 deletions
+64
View File
@@ -443,6 +443,8 @@ NumericTypeID_enumerants[] = {
{"s16", "S16", NumericTypeID::kS16},
{"s32", "S32", NumericTypeID::kS32},
{"s64", "S64", NumericTypeID::kS64},
{"fe4m3", "FE4M3", NumericTypeID::kFE4M3},
{"fe5m2", "FE5M2", NumericTypeID::kFE5M2},
{"f16", "F16", NumericTypeID::kF16},
{"bf16", "BF16", NumericTypeID::kBF16},
{"f32", "F32", NumericTypeID::kF32},
@@ -504,6 +506,8 @@ NumericTypeID from_string<NumericTypeID>(std::string const &str) {
/// Returns the size of a data type in bits
int sizeof_bits(NumericTypeID type) {
switch (type) {
case NumericTypeID::kFE4M3: return 8;
case NumericTypeID::kFE5M2: return 8;
case NumericTypeID::kF16: return 16;
case NumericTypeID::kBF16: return 16;
case NumericTypeID::kTF32: return 32;
@@ -581,6 +585,8 @@ bool is_integer_type(NumericTypeID type) {
/// Returns true if numeric type is signed
bool is_signed_type(NumericTypeID type) {
switch (type) {
case NumericTypeID::kFE4M3: return true;
case NumericTypeID::kFE5M2: return true;
case NumericTypeID::kF16: return true;
case NumericTypeID::kBF16: return true;
case NumericTypeID::kTF32: return true;
@@ -610,6 +616,8 @@ bool is_unsigned_integer(NumericTypeID type) {
/// Returns true if numeric type is floating-point type
bool is_float_type(NumericTypeID type) {
switch (type) {
case NumericTypeID::kFE4M3: return true;
case NumericTypeID::kFE5M2: return true;
case NumericTypeID::kF16: return true;
case NumericTypeID::kBF16: return true;
case NumericTypeID::kTF32: return true;
@@ -1050,6 +1058,20 @@ bool lexical_cast(std::vector<uint8_t> &bytes, NumericTypeID type, std::string c
ss >> *reinterpret_cast<int64_t *>(bytes.data());
}
break;
case NumericTypeID::kFE4M3:
{
float tmp;
ss >> tmp;
*reinterpret_cast<float_e4m3_t *>(bytes.data()) = static_cast<float_e4m3_t>(tmp);
}
break;
case NumericTypeID::kFE5M2:
{
float tmp;
ss >> tmp;
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(tmp);
}
break;
case NumericTypeID::kF16:
{
float tmp;
@@ -1187,6 +1209,18 @@ std::string lexical_cast(std::vector<uint8_t> &bytes, NumericTypeID type) {
ss << *reinterpret_cast<int64_t *>(bytes.data());
}
break;
case NumericTypeID::kFE4M3:
{
float tmp = *reinterpret_cast<float_e4m3_t *>(bytes.data());
ss << tmp;
}
break;
case NumericTypeID::kFE5M2:
{
float tmp = *reinterpret_cast<float_e5m2_t *>(bytes.data());
ss << tmp;
}
break;
case NumericTypeID::kF16:
{
float tmp = *reinterpret_cast<half_t *>(bytes.data());
@@ -1329,6 +1363,16 @@ bool cast_from_int64(std::vector<uint8_t> &bytes, NumericTypeID type, int64_t sr
*reinterpret_cast<int64_t *>(bytes.data()) = static_cast<int64_t>(src);
}
break;
case NumericTypeID::kFE4M3:
{
*reinterpret_cast<float_e4m3_t *>(bytes.data()) = static_cast<float_e4m3_t>(float(src));
}
break;
case NumericTypeID::kFE5M2:
{
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(float(src));
}
break;
case NumericTypeID::kF16:
{
*reinterpret_cast<half_t *>(bytes.data()) = static_cast<half_t>(float(src));
@@ -1429,6 +1473,16 @@ bool cast_from_uint64(std::vector<uint8_t> &bytes, NumericTypeID type, uint64_t
*reinterpret_cast<int64_t *>(bytes.data()) = static_cast<int64_t>(src);
}
break;
case NumericTypeID::kFE4M3:
{
*reinterpret_cast<float_e4m3_t *>(bytes.data()) = static_cast<float_e4m3_t>(float(src));
}
break;
case NumericTypeID::kFE5M2:
{
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(float(src));
}
break;
case NumericTypeID::kF16:
{
*reinterpret_cast<half_t *>(bytes.data()) = static_cast<half_t>(float(src));
@@ -1530,6 +1584,16 @@ bool cast_from_double(std::vector<uint8_t> &bytes, NumericTypeID type, double sr
*reinterpret_cast<int64_t *>(bytes.data()) = static_cast<int64_t>(src);
}
break;
case NumericTypeID::kFE4M3:
{
*reinterpret_cast<float_e4m3_t *>(bytes.data()) = static_cast<float_e4m3_t>(float(src));
}
break;
case NumericTypeID::kFE5M2:
{
*reinterpret_cast<float_e5m2_t *>(bytes.data()) = static_cast<float_e5m2_t>(float(src));
}
break;
case NumericTypeID::kF16:
{
*reinterpret_cast<half_t *>(bytes.data()) = static_cast<half_t>(float(src));