This commit is contained in:
nickfeeney 2020-02-24 17:07:38 -08:00
parent f407de6da1
commit 92d8d90d24
3 changed files with 21 additions and 19 deletions

View file

@ -30,6 +30,24 @@ DML_TENSOR_DATA_TYPE GetDmlDataTypeFromMlDataTypeNoThrow(MLOperatorTensorDataTyp
};
}
bool IsSigned(DML_TENSOR_DATA_TYPE dataType)
{
switch (dataType)
{
case DML_TENSOR_DATA_TYPE_FLOAT32: return true;
case DML_TENSOR_DATA_TYPE_FLOAT16: return true;
case DML_TENSOR_DATA_TYPE_UINT32: return false;
case DML_TENSOR_DATA_TYPE_UINT16: return false;
case DML_TENSOR_DATA_TYPE_UINT8: return false;
case DML_TENSOR_DATA_TYPE_INT32: return true;
case DML_TENSOR_DATA_TYPE_INT16: return true;
case DML_TENSOR_DATA_TYPE_INT8: return true;
}
assert(false);
return false;
}
DML_TENSOR_DATA_TYPE GetDmlDataTypeFromMlDataType(MLOperatorTensorDataType tensorDataType)
{
DML_TENSOR_DATA_TYPE dmlTensorDataType = GetDmlDataTypeFromMlDataTypeNoThrow(tensorDataType);

View file

@ -18,6 +18,8 @@ namespace Dml
size_t ComputeByteSizeFromDimensions(gsl::span<const DimensionType> dimensions, MLOperatorTensorDataType tensorDataType);
size_t ComputeByteSizeFromTensor(IMLOperatorTensor& tensor);
bool IsSigned(DML_TENSOR_DATA_TYPE dataType);
/** Calculates the minimum number of bytes required to store a buffer tensor with the specified type, sizes, and
strides. The formula can be expressed as the following:

View file

@ -215,22 +215,4 @@ private:
// allocated memory if the fixed stack array is exhausted.
FixedBucket m_fixed;
std::deque<DynamicBucket> m_dynamic;
};
inline bool IsSigned(DML_TENSOR_DATA_TYPE dataType)
{
switch (dataType)
{
case DML_TENSOR_DATA_TYPE_FLOAT32: return true;
case DML_TENSOR_DATA_TYPE_FLOAT16: return true;
case DML_TENSOR_DATA_TYPE_UINT32: return false;
case DML_TENSOR_DATA_TYPE_UINT16: return false;
case DML_TENSOR_DATA_TYPE_UINT8: return false;
case DML_TENSOR_DATA_TYPE_INT32: return true;
case DML_TENSOR_DATA_TYPE_INT16: return true;
case DML_TENSOR_DATA_TYPE_INT8: return true;
}
assert(false);
return false;
}
};