Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 25 additions & 18 deletions src/relay/transforms/pattern_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -468,43 +468,43 @@ inline bool IsEqualScalar(const Expr& a, const Expr& b) {
* \param i element index
* \return Converted scalar value, or None if conversion failed
*/
static inline std::optional<long double> TryToScalar(const runtime::NDArray& array, size_t i = 0) {
template <typename T>
static inline std::optional<T> TryToScalar(const runtime::NDArray& array, size_t i = 0) {
if (array->dtype.code == kDLInt) {
if (array->dtype.bits == 8) {
return std::optional<long double>(reinterpret_cast<int8_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<int8_t*>(array->data)[i]);
} else if (array->dtype.bits == 16) {
return std::optional<long double>(reinterpret_cast<int16_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<int16_t*>(array->data)[i]);
} else if (array->dtype.bits == 32) {
return std::optional<long double>(reinterpret_cast<int32_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<int32_t*>(array->data)[i]);
} else if (array->dtype.bits == 64) {
return std::optional<long double>(reinterpret_cast<int64_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<int64_t*>(array->data)[i]);
}
} else if (array->dtype.code == kDLUInt) {
if (array->dtype.bits == 1) { // bool
return std::optional<long double>(reinterpret_cast<uint8_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<uint8_t*>(array->data)[i]);
} else if (array->dtype.bits == 8) {
return std::optional<long double>(reinterpret_cast<uint8_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<uint8_t*>(array->data)[i]);
} else if (array->dtype.bits == 16) {
return std::optional<long double>(reinterpret_cast<uint16_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<uint16_t*>(array->data)[i]);
} else if (array->dtype.bits == 32) {
return std::optional<long double>(reinterpret_cast<uint32_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<uint32_t*>(array->data)[i]);
} else if (array->dtype.bits == 64) {
return std::optional<long double>(reinterpret_cast<uint64_t*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<uint64_t*>(array->data)[i]);
}
} else if (array->dtype.code == kDLFloat) {
if (array->dtype.bits == 16) {
return std::optional<long double>(
__extendXfYf2__<uint16_t, uint16_t, 10, float, uint32_t, 23>(
reinterpret_cast<uint16_t*>(array->data)[i]));
return std::optional<T>(__extendXfYf2__<uint16_t, uint16_t, 10, float, uint32_t, 23>(
reinterpret_cast<uint16_t*>(array->data)[i]));
}
if (array->dtype.bits == 32) {
return std::optional<long double>(reinterpret_cast<float*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<float*>(array->data)[i]);
} else if (array->dtype.bits == 64) {
return std::optional<long double>(reinterpret_cast<double*>(array->data)[i]);
return std::optional<T>(reinterpret_cast<double*>(array->data)[i]);
}
} else if (array->dtype.code == kDLBfloat) {
if (array->dtype.bits == 16) {
return std::optional<long double>(__extendXfYf2__<uint16_t, uint16_t, 7, float, uint32_t, 23>(
return std::optional<T>(__extendXfYf2__<uint16_t, uint16_t, 7, float, uint32_t, 23>(
reinterpret_cast<uint16_t*>(array->data)[i]));
}
}
Expand All @@ -517,8 +517,15 @@ static inline std::optional<long double> TryToScalar(const runtime::NDArray& arr
* \param i element index
* \return Converted scalar value
*/
template <typename T>
static inline T ToScalar(const runtime::NDArray& array, size_t i = 0) {
auto try_value = TryToScalar<T>(array, i);
ICHECK(try_value) << "Unknown data type: " << tvm::runtime::DLDataType2String(array->dtype);
return try_value.value();
}

static inline long double ToScalar(const runtime::NDArray& array, size_t i = 0) {
auto try_value = TryToScalar(array, i);
auto try_value = TryToScalar<long double>(array, i);
ICHECK(try_value) << "Unknown data type: " << tvm::runtime::DLDataType2String(array->dtype);
return try_value.value();
}
Expand All @@ -534,7 +541,7 @@ static inline Array<Integer> ToVector(const runtime::NDArray& array) {
size_t len = array.Shape().front();
Array<Integer> out;
for (size_t i = 0; i < len; ++i) {
long double elem_val = ToScalar(array, i);
uint64_t elem_val = ToScalar<uint64_t>(array, i);
out.push_back(Integer(IntImm(DataType::Int(32), static_cast<int64_t>(elem_val))));
}
return out;
Expand Down
2 changes: 1 addition & 1 deletion src/relay/transforms/simplify_expr.cc
Original file line number Diff line number Diff line change
Expand Up @@ -794,7 +794,7 @@ class EliminateIdentityRewrite : public DFPatternRewrite {
if (!IsScalar(GetRef<Expr>(constant))) {
return false;
}
auto value = TryToScalar(constant->data, 0);
auto value = TryToScalar<long double>(constant->data, 0);
if (!value) {
// unsupported dtype
return false;
Expand Down