Visitar URL original
BUG: fix overflow in np.round of float16 by tuan2k33 · Pull Request #32916 · numpy/numpy · GitHub
Skip to content

BUG: fix overflow in np.round of float16 - #32916

Open
tuan2k33 wants to merge 2 commits into
numpy:mainfrom
tuan2k33:fix/round-float16
Open

tuan2k33 wants to merge 2 commits into
numpy:mainfrom
tuan2k33:fix/round-float16

Conversation

@tuan2k33

@tuan2k33 tuan2k33 commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

PR summary

Fixes #13699.

np.round on a float16 array ran multiply, rint and true_divide as three float16 ufunc passes. Since float16 has no arithmetic of its own, each pass converted every element float16 -> float32, computed, and converted the result float32 -> float16 again, so an element went through three round trips per call. Each of those conversions stores the result as float16, so
x * 10**decimals becomes inf above 65504, and the scale factor 10**decimals is already inf in float16 for decimals >= 5.

This PR fix float16 case: it rounds float16 arrays in float32, using the existing float32 code, and converts back once. This leads to decimal >= 5 coverage (no longer warn overflow) and a small speedup. Others are unchanged.

Checks I ran locally:

  • The NumPy test suite passes. New tests are in TestMethods in test_multiarray.py.
  • float16 np.round is about 3x faster than on main; float32 and float64 are unchanged.
Benchmark script

Run it once per NumPy build, pinned to one core, and compare:

taskset -c 6 python bench_round_pr.py main1.json   # build of main
taskset -c 6 python bench_round_pr.py pr1.json     # build of this PR
# repeat both once more (main2.json, pr2.json), then:
python compare_round_pr.py main1.json main2.json -- pr1.json pr2.json

bench_round_pr.py:

"""Time np.round for float16/float32/float64.

usage: python bench_round_pr.py out.json
Run it once per NumPy build (e.g. with PYTHONPATH pointing at each build,
pinned to one core: taskset -c 6 python bench_round_pr.py out.json).
"""
import json
import sys
import time

import numpy as np


def seconds_per_call(f):
    f()
    k = 1
    while True:  # repeat until one timing takes >= 20 ms
        t0 = time.perf_counter()
        for _ in range(k):
            f()
        if time.perf_counter() - t0 >= 0.02 or k >= 100_000:
            break
        k *= 2

    def one():
        t0 = time.perf_counter()
        for _ in range(k):
            f()
        return (time.perf_counter() - t0) / k

    return min(one() for _ in range(9))


res = {}
for dtype in (np.float16, np.float32, np.float64):
    for decimals in (-2, 2, 6):
        for n in (1000, 100_000, 1_000_000):
            rng = np.random.default_rng(12345)
            a = (rng.random(n) * 200 - 100).astype(dtype)
            out = np.empty_like(a)
            cases = {
                "round": lambda: np.round(a, decimals),
                "round_out": lambda: np.round(a, decimals, out=out),
                "round_strided": lambda: np.round(a[::2], decimals),
            }
            for name, f in cases.items():
                key = f"{dtype.__name__}|decimals={decimals}|n={n}|{name}"
                res[key] = seconds_per_call(f)

with open(sys.argv[1], "w") as fh:
    json.dump(res, fh)

compare_round_pr.py:

"""usage: python compare_round_pr.py base1.json [base2.json ...] -- new1.json [new2.json ...]
Takes the min over the runs of each build and prints the speedup (base / new)."""
import collections
import json
import statistics
import sys

i = sys.argv.index("--")
base = [json.load(open(f)) for f in sys.argv[1:i]]
new = [json.load(open(f)) for f in sys.argv[i + 1:]]
groups = collections.defaultdict(list)
for key in base[0]:
    speedup = min(r[key] for r in base) / min(r[key] for r in new)
    dtype, _, _, name = key.split("|")
    groups[(dtype, name)].append(speedup)
for (dtype, name), v in sorted(groups.items()):
    print(f"{dtype:8s} {name:14s} median {statistics.median(v):6.2f}x  "
          f"min {min(v):6.2f}x  max {max(v):6.2f}x")

Speedup (main time / PR time) on my machine (AVX2/FMA/F16C, one pinned core, 2 interleaved runs, min of each):

float16  round          median   3.58x  min   2.82x  max   4.23x
float16  round_out      median   3.40x  min   2.94x  max   4.27x
float16  round_strided  median   3.47x  min   2.41x  max   4.46x
float32  round          median   0.99x  min   0.96x  max   1.05x
float32  round_out      median   1.03x  min   0.96x  max   1.08x
float32  round_strided  median   1.03x  min   1.00x  max   1.04x
float64  round          median   1.01x  min   0.95x  max   1.06x
float64  round_out      median   1.01x  min   0.88x  max   1.04x
float64  round_strided  median   1.04x  min   0.95x  max   1.12x

AI Disclosure

Claude Code was used.

round() ran multiply, rint and true_divide on the float16 array.  Each pass
stores its result as float16, so ``x * 10**decimals`` became inf above 65504,
and the scale factor itself is already inf in float16 for decimals >= 5, which
gave nan/inf for values such as ``np.round(np.float16(2.0), 5)`` (numpygh-13699).
Intermediate results were also rounded to float16 before ``rint``.

Round exact float16 arrays in float32 and convert back once, using the
existing float32 code.  Subclasses and an ``out`` of another dtype or shape
keep using the old path.
The test converts every float16 bit pattern to float32, including signalling
NaNs.  On some platforms (ARM64 CI) that conversion raises "invalid", which
the test configuration turns into an error.  Do the conversion inside the
existing errstate block.

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Surprising overflows in np.round of float16.

1 participant