Skip to content

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

Open
EylonKrause wants to merge 4 commits into
jax-ml:mainfrom
EylonKrause:fix/npy-cast-double-rounding
Open

EylonKrause wants to merge 4 commits into
jax-ml:mainfrom
EylonKrause:fix/npy-cast-double-rounding

Conversation

@EylonKrause

@EylonKrause EylonKrause commented Sep 2, 2026 •

Copy link
Copy Markdown

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 to float first and only then to the target:

path today
np.array(x, np.float64).astype(T) (also longdouble, complex128) static_cast<To>(static_cast<float>(from))
np.array(n, np.int64).astype(T) (also uint64, int32, uint32) same, via static_cast<float>
T(python_float) T(double): exact for the float8 types, but Eigen's bfloat16(double) narrows to float first
T(python_int) T(static_cast<float>(long))
T(np.int64(n)), T(np.float64(x)), T(np.longdouble(x)), … the NumPy scalar is read through the float32 setitem

The 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.

input true nearest today
float64 -431.99999999999994 → float8_e4m3fn (neighbours −416 / −448) −416 −448
float64 1 + 2⁻⁸ + 2⁻⁴⁰ → bfloat16 (neighbours 1 / 1.0078125) 1.0078125 1.0
int64 16973823 → bfloat16 (neighbours 16908288 / 17039360, midpoint 16973824) 16908288 17039360
int64 3·2³⁰ − 1 → float8_e8m0fnu (neighbours 2³¹ / 2³²) 2³¹ 2³²
bfloat16(16973823), bfloat16(np.int64(16973823)) 16908288 17039360
bfloat16(1 + 2⁻⁸ + 2⁻⁴⁰), bfloat16(np.longdouble(…)) 1.0078125 1.0

NumPy's own int64 → float16 and float64 → float16 casts are single-rounded on the same inputs, which is what users expect from a dtype.

Fix

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 bits, float8: ≤ 4). It is the trick float8_base::ConvertFrom already applies for long 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 exact ldexp.
  • One NarrowToFloat entry point (exact when the value fits in a float, 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 double by scalar kind) instead of through the float32 setitem, so T(np.float64(x)) and T(np.array([x]).astype(T)[0]) agree.
  • The setitem result is now checked: before, a failing conversion (e.g. bfloat16(np.datetime64(...))) returned garbage with an exception left set, which surfaced as SystemError: returned a result with an exception set; the conversion's own error propagates now.
  • Second commit, found while reviewing that path: bfloat16("abc") (any custom float built from a string that is not a number) segfaults on main — the constructor passes the null result of PyFloat_FromString straight into CastToCustomFloat, which dereferences it (and on success the parsed float was never released). It now raises the parser's ValueError, and the constructor and array setitem return a conversion's own error instead of replacing it with a generic TypeError, mirroring ints.cc.

Validation

  • Exhaustive oracle sweep (nearest representable value, ties to even by encoding) for every type: float64 → T over 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 → bfloat16 over integers next to every midpoint in [2²⁴, 2⁶²): before 156/468 mismatches, after 0.
  • The only remaining float8_e8m0fnu mismatches are exact ties / its lowest binade, i.e. e8m0's own float → e8m0 rounding, 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.
  • Tests: testCastFromFloat64DoesNotDoubleRound, testCastFromInt64DoesNotDoubleRound, testConstructFromScalarDoesNotDoubleRound (Python int/float, NumPy int32/uint32/int64/uint64/float64/longdouble scalars), for every parameterized type. Each fails without the corresponding fix; the full custom_float_test.py passes with -W error (2195 passed, 3 runs) and so does the rest of the suite (2752 passed).
  • testConstructFromInvalidValueRaises: str/bytes/np.str_ garbage raises ValueError (was a segfault / a SystemError), np.datetime64 raises, numeric strings still parse.
  • Edge cases checked by hand: INT64_MIN, UINT64_MAX, -0.0 sign, NaN/±inf, overflow (float8_e4m3fn(1e300) still NaN), complex128 (real part, ComplexWarning kept), 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

EylonKrause added a commit to EylonKrause/ml_dtypes that referenced this pull request Sep 2, 2026
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
@google-cla

google-cla Bot commented Sep 2, 2026

Copy link
Copy Markdown

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.

@EylonKrause

Copy link
Copy Markdown
Author

@googlebot I signed it!

EylonKrause added a commit to EylonKrause/ml_dtypes that referenced this pull request Sep 2, 2026
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
@EylonKrause
EylonKrause force-pushed the fix/npy-cast-double-rounding branch from 70f04a0 to 24253eb Compare September 2, 2026 17:07
EylonKrause added a commit to EylonKrause/ml_dtypes that referenced this pull request Sep 2, 2026
@EylonKrause
EylonKrause force-pushed the fix/npy-cast-double-rounding branch from 24253eb to 023a1f0 Compare September 2, 2026 17:09
…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.
@EylonKrause
EylonKrause force-pushed the fix/npy-cast-double-rounding branch from 023a1f0 to 566ac19 Compare September 11, 2026 12:23
@EylonKrause EylonKrause changed the title Fix double rounding in NumPy casts from float64 to the custom float types Fix double rounding in conversions to the custom float types (float64, integers, Python and NumPy scalars); fix a constructor crash on non-numeric strings Sep 11, 2026

@hawkinsp hawkinsp left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Sorry for the slow review, this seems like a good fix.

Please squash your commits.

Comment thread ml_dtypes/_src/floats.cc

#include <array>
#include <cmath>
#include <cmath>

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Duplicate cmath include.

Comment thread ml_dtypes/_src/floats.cc
return false;
}
c = NarrowToFloat(value);
} else {

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread ml_dtypes/_src/floats.cc
}
}
constexpr int kFloatDigits = std::numeric_limits<float>::digits;
int shift = 0;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Test negative integers?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants