iinfo and scalar overflow detection (#2009)

This commit is contained in:
Awni Hannun
2025-03-27 19:54:56 -07:00
committed by GitHub
parent bc62932984
commit 5580b47291
6 changed files with 112 additions and 0 deletions

View File

@@ -2,6 +2,7 @@
#include "python/src/utils.h"
#include "mlx/ops.h"
#include "mlx/utils.h"
#include "python/src/convert.h"
mx::array to_array(
@@ -16,6 +17,16 @@ mx::array to_array(
? mx::int64
: mx::int32;
auto out_t = dtype.value_or(default_type);
if (mx::issubdtype(out_t, mx::integer) && out_t.size() < 8) {
auto info = mx::iinfo(out_t);
if (val < info.min || val > static_cast<int64_t>(info.max)) {
std::ostringstream msg;
msg << "Converting " << val << " to " << out_t
<< " would result in overflow.";
throw std::invalid_argument(msg.str());
}
}
// bool_ is an exception and is always promoted
return mx::array(val, (out_t == mx::bool_) ? mx::int32 : out_t);
} else if (auto pv = std::get_if<nb::float_>(&v); pv) {