co-authored by
Aniket Shivam
parent
9b8166e3f0
commit
d572cc1aab
@@ -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));
|
||||
|
||||
Reference in New Issue
Block a user