From 0982d1e33e4b95b0d40a884ee224c6404ede870e Mon Sep 17 00:00:00 2001 From: "wenjian.ma" Date: Mon, 1 Jan 2024 18:16:21 +0800 Subject: [PATCH] [Relay] make "ToScalar" support directly obtaining "int64_t" Because on Windows, "long double" is 64 bits instead of 128 bits like on Linux, to avoid overflow from "long double" to "int64_t" --- src/relay/transforms/pattern_utils.h | 43 ++++++++++++++++----------- src/relay/transforms/simplify_expr.cc | 2 +- 2 files changed, 26 insertions(+), 19 deletions(-) diff --git a/src/relay/transforms/pattern_utils.h b/src/relay/transforms/pattern_utils.h index 50c2e0029885..b26bd7649630 100644 --- a/src/relay/transforms/pattern_utils.h +++ b/src/relay/transforms/pattern_utils.h @@ -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 TryToScalar(const runtime::NDArray& array, size_t i = 0) { +template +static inline std::optional TryToScalar(const runtime::NDArray& array, size_t i = 0) { if (array->dtype.code == kDLInt) { if (array->dtype.bits == 8) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 16) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 32) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 64) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } } else if (array->dtype.code == kDLUInt) { if (array->dtype.bits == 1) { // bool - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 8) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 16) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 32) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 64) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } } else if (array->dtype.code == kDLFloat) { if (array->dtype.bits == 16) { - return std::optional( - __extendXfYf2__( - reinterpret_cast(array->data)[i])); + return std::optional(__extendXfYf2__( + reinterpret_cast(array->data)[i])); } if (array->dtype.bits == 32) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } else if (array->dtype.bits == 64) { - return std::optional(reinterpret_cast(array->data)[i]); + return std::optional(reinterpret_cast(array->data)[i]); } } else if (array->dtype.code == kDLBfloat) { if (array->dtype.bits == 16) { - return std::optional(__extendXfYf2__( + return std::optional(__extendXfYf2__( reinterpret_cast(array->data)[i])); } } @@ -517,8 +517,15 @@ static inline std::optional TryToScalar(const runtime::NDArray& arr * \param i element index * \return Converted scalar value */ +template +static inline T ToScalar(const runtime::NDArray& array, size_t i = 0) { + auto try_value = TryToScalar(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(array, i); ICHECK(try_value) << "Unknown data type: " << tvm::runtime::DLDataType2String(array->dtype); return try_value.value(); } @@ -534,7 +541,7 @@ static inline Array ToVector(const runtime::NDArray& array) { size_t len = array.Shape().front(); Array out; for (size_t i = 0; i < len; ++i) { - long double elem_val = ToScalar(array, i); + uint64_t elem_val = ToScalar(array, i); out.push_back(Integer(IntImm(DataType::Int(32), static_cast(elem_val)))); } return out; diff --git a/src/relay/transforms/simplify_expr.cc b/src/relay/transforms/simplify_expr.cc index 208c9821b670..8036d301e191 100644 --- a/src/relay/transforms/simplify_expr.cc +++ b/src/relay/transforms/simplify_expr.cc @@ -794,7 +794,7 @@ class EliminateIdentityRewrite : public DFPatternRewrite { if (!IsScalar(GetRef(constant))) { return false; } - auto value = TryToScalar(constant->data, 0); + auto value = TryToScalar(constant->data, 0); if (!value) { // unsupported dtype return false;