Fix double rounding in conversions to the custom float types (float64, integers, Python and NumPy scalars); fix a constructor crash on non-numeric strings - #401
Conversation
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
|
Thanks for your pull request! It looks like this may be your first contribution to a Google open source project. Before we can look at your pull request, you'll need to sign a Contributor License Agreement (CLA). View this failed invocation of the CLA check for more information. For the most up to date status, view the checks section at the bottom of the pull request. |
|
@googlebot I signed it! |
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
70f04a0 to
24253eb
Compare
24253eb to
023a1f0
Compare
…ypes The NumPy cast kernel NPyCast<From, To> converted every source value through `float` first (CastToFloat) and only then to the target type. For a float64 or longdouble source that is two roundings: the first can land exactly on a midpoint of the target format, and the second then resolves that spurious tie to even instead of picking the true nearest value. For example the float64 -431.99999999999994 is nearest to the float8_e4m3fn value -416, but rounds to -432.0f (a tie) and then to -448; and 1 + 2^-8 + 2^-30 is nearest to the bfloat16 value 1.0078125 but became 1.0. Every ml_dtypes float type was affected through `.astype`. (The C++ scalar conversion for the float8 types is already exact because it converts from the double's bits directly; bfloat16 double-rounds on both paths because Eigen's bfloat16(double) also narrows to float first.) Narrow the wide source to float with round-to-odd before the final rounding. This makes the subsequent round-to-nearest-even exact for any target with at least two fewer significand bits than float (Boldo & Melquiond, "When double rounding is odd", 2005), which holds for every type here (bfloat16: 8, float8: <= 4). It is the same trick float8_base::ConvertFrom already applies for long double. Add a regression test that, for each float type, takes adjacent representable values, nudges a float64 one ulp off their midpoint (invisible to float32), and checks that both directions round to the true nearest neighbour. It fails for every type without the fix. Signed-off-by: Eylon Krause <eylon1909@gmail.com>
…ctors The NumPy cast kernel still narrowed 32- and 64-bit integers through a plain float, and the scalar constructor narrowed every source through float: T(double) for Eigen's bfloat16, static_cast<float>(long) for Python ints, and a float32 setitem for NumPy scalars. All of them double-round for the same reason as the float64 array cast. Route every source through one NarrowToFloat helper: exact when the value fits in a float, round-to-odd otherwise (a new IntegerToFloatRoundToOdd for wide integers: keep the top 24 bits, fold the dropped bits into the lowest kept bit, scale back exactly). NumPy scalars are read at full width (64-bit integer or long double by scalar kind) instead of float32, and a failing setitem now returns false instead of leaving an exception set and returning garbage (a SystemError for e.g. a datetime64 argument before). Tests cover int64/uint64 array casts and construction from Python int/float and NumPy int32/uint32/int64/uint64/float64/longdouble scalars, next to a midpoint of every float type.
bfloat16("abc") (and any custom float constructed from a string that is
not a number) dereferenced the null result of PyFloat_FromString and
crashed the interpreter. The parsed float was also never released.
Propagate the parser's ValueError, own the parsed float, and return the
conversion's own error when CastToCustomFloat fails with one set (the
NumPy-scalar path now reports those instead of returning garbage), both in
the constructor and in the array setitem, mirroring ints.cc.
023a1f0 to
566ac19
Compare
hawkinsp
left a comment
There was a problem hiding this comment.
Sorry for the slow review, this seems like a good fix.
Please squash your commits.
|
|
||
| #include <array> | ||
| #include <cmath> | ||
| #include <cmath> |
| return false; | ||
| } | ||
| c = NarrowToFloat(value); | ||
| } else { |
There was a problem hiding this comment.
I gather the long double parsing path is actually a bit quirky. I would have a check for long double specifically and then have the else case go via double.
| } | ||
| } | ||
| constexpr int kFloatDigits = std::numeric_limits<float>::digits; | ||
| int shift = 0; |
There was a problem hiding this comment.
I think you could add a fast path of this form
if ((magnitude >> kFloatDigits) == 0) return static_cast(from);
| self.assertEqual( | ||
| float(np.array(down).astype(float_type).astype(np.float64)), lo) | ||
|
|
||
| def testCastFromInt64DoesNotDoubleRound(self, float_type): |
There was a problem hiding this comment.
Test negative integers as well?
| self.assertEqual(float(float_type("0.5")), 0.5) | ||
| self.assertEqual(float(float_type(np.str_("0.5"))), 0.5) | ||
|
|
||
| def testConstructFromScalarDoesNotDoubleRound(self, float_type): |
Problem
Converting to any of the ml_dtypes float types double-rounds whenever the source is wider than
float. Every conversion path narrows the source value tofloatfirst and only then to the target:np.array(x, np.float64).astype(T)(alsolongdouble,complex128)static_cast<To>(static_cast<float>(from))np.array(n, np.int64).astype(T)(alsouint64,int32,uint32)static_cast<float>T(python_float)T(double): exact for the float8 types, but Eigen'sbfloat16(double)narrows tofloatfirstT(python_int)T(static_cast<float>(long))T(np.int64(n)),T(np.float64(x)),T(np.longdouble(x)), …setitemThe first rounding (to
float) can land exactly on a midpoint of the target format; the second rounding then resolves that spurious tie to even instead of picking the true nearest value.-431.99999999999994→float8_e4m3fn(neighbours −416 / −448)1 + 2⁻⁸ + 2⁻⁴⁰→bfloat16(neighbours 1 / 1.0078125)16973823→bfloat16(neighbours 16908288 / 17039360, midpoint 16973824)3·2³⁰ − 1→float8_e8m0fnu(neighbours 2³¹ / 2³²)bfloat16(16973823),bfloat16(np.int64(16973823))bfloat16(1 + 2⁻⁸ + 2⁻⁴⁰),bfloat16(np.longdouble(…))NumPy's own
int64 → float16andfloat64 → float16casts are single-rounded on the same inputs, which is what users expect from a dtype.Fix
Narrow the wide source to
floatwith round-to-odd before the final rounding. This makes the subsequent round-to-nearest-even exact for any target with at least two fewer significand bits thanfloat(Boldo & Melquiond, When double rounding is odd, 2005), which holds for every type here (bfloat16: 8 bits, float8: ≤ 4). It is the trickfloat8_base::ConvertFromalready applies forlong double.NarrowToFloatRoundToOdd(float64 / long double) — as in the first commit.IntegerToFloatRoundToOdd(32/64-bit integers): keep the top 24 bits, fold any dropped bit into the lowest kept bit, scale back with an exactldexp.NarrowToFloatentry point (exact when the value fits in afloat, round-to-odd otherwise) used by the NumPy cast kernel and by all three scalar-constructor paths. NumPy scalars are now read at full width (long long/unsigned long long/long doubleby scalar kind) instead of through the float32setitem, soT(np.float64(x))andT(np.array([x]).astype(T)[0])agree.setitemresult is now checked: before, a failing conversion (e.g.bfloat16(np.datetime64(...))) returned garbage with an exception left set, which surfaced asSystemError: returned a result with an exception set; the conversion's own error propagates now.bfloat16("abc")(any custom float built from a string that is not a number) segfaults onmain— the constructor passes the null result ofPyFloat_FromStringstraight intoCastToCustomFloat, which dereferences it (and on success the parsed float was never released). It now raises the parser'sValueError, and the constructor and arraysetitemreturn a conversion's own error instead of replacing it with a genericTypeError, mirroringints.cc.Validation
float64 → Tover all float16 values plus every midpoint ± 1 float64 ulp — before: e4m3fn 252, e5m2 246, e4m3b11fnuz 254, e4m3fnuz 254, e5m2fnuz 254, e3m4 222, e2m1fn 14, e2m3fn 62, e3m2fn 62, bfloat16 65 278 mismatches; after: 0 for every type.int64/uint64 → bfloat16over integers next to every midpoint in[2²⁴, 2⁶²): before 156/468 mismatches, after 0.float8_e8m0fnumismatches are exact ties / its lowest binade, i.e. e8m0's ownfloat → e8m0rounding, which Round exact ties to nearest-even when converting to float8_e8m0fnu #398 fixes separately; the tests restrict themselves to float32-normal midpoints so they do not depend on Round exact ties to nearest-even when converting to float8_e8m0fnu #398.testCastFromFloat64DoesNotDoubleRound,testCastFromInt64DoesNotDoubleRound,testConstructFromScalarDoesNotDoubleRound(Pythonint/float, NumPyint32/uint32/int64/uint64/float64/longdoublescalars), for every parameterized type. Each fails without the corresponding fix; the fullcustom_float_test.pypasses with-W error(2195 passed, 3 runs) and so does the rest of the suite (2752 passed).testConstructFromInvalidValueRaises:str/bytes/np.str_garbage raisesValueError(was a segfault / aSystemError),np.datetime64raises, numeric strings still parse.INT64_MIN,UINT64_MAX,-0.0sign, NaN/±inf, overflow (float8_e4m3fn(1e300)still NaN),complex128(real part,ComplexWarningkept),np.bool_,np.str_, scalars of the other ml_dtypes types.Disclosure: this contribution was authored with an AI coding assistant (Claude) and reviewed and validated before submission.
🤖 Generated with Claude Code