diff --git a/src/gfloat/encode.py b/src/gfloat/encode.py index dd739c0..5f08af4 100644 --- a/src/gfloat/encode.py +++ b/src/gfloat/encode.py @@ -53,6 +53,10 @@ def encode_float(fi: FormatInfo, v: float) -> int: sign = fi.is_signed and np.signbit(v) vpos = -v if sign else v + # Zero has its own code point even when the format has no subnormals. + if fi.has_zero and vpos == 0: + return fi.code_of_negzero if sign and fi.has_nz else fi.code_of_zero + if fi.has_subnormals and vpos <= fi.smallest_subnormal / 2: isig = 0 biased_exp = 0 diff --git a/src/gfloat/round.py b/src/gfloat/round.py index a48fdc2..32be3a8 100644 --- a/src/gfloat/round.py +++ b/src/gfloat/round.py @@ -48,7 +48,10 @@ def round_float( An input NaN will convert to a NaN in the target. An input Infinity will convert to the largest float if :paramref:`sat`, otherwise to an Inf, if present, otherwise to a NaN. - Negative zero will be returned if the format has negative zero, otherwise zero. + If the format has zero, negative zero will be returned if it has negative zero, + otherwise zero. Formats without zero clamp finite values below the smallest + magnitude to that magnitude, preserving the sign for signed formats, regardless + of :paramref:`sat` and :paramref:`rnd`. Args: fi (FormatInfo): Describes the target format @@ -87,6 +90,9 @@ def round_float( sign = np.signbit([v]).item() and fi.is_signed vpos = -v if sign else v + if not fi.has_zero and math.isfinite(vpos) and vpos < fi.smallest: + return -fi.smallest if sign else fi.smallest + if math.isinf(vpos): result = np.inf @@ -107,6 +113,12 @@ def round_float( # use ldexp instead of vpos*2**-expval to avoid overflow fsignificand = math.ldexp(vpos, -expval) + # Without subnormals, code points 0 and 1 may be separated by + # more than one significand step. Round across that whole gap. + zero_gap = not fi.has_subnormals and fi.has_zero and vpos < fi.smallest + if zero_gap: + fsignificand = vpos / fi.smallest + # Round isignificand = math.floor(fsignificand) delta = fsignificand - isignificand @@ -160,7 +172,7 @@ def round_float( isignificand += 1 # Reconstruct rounded result to float - result = isignificand * (2.0**expval) + result = isignificand * (fi.smallest if zero_gap else 2.0**expval) if result == 0: if sign and fi.has_nz: diff --git a/src/gfloat/round_ndarray.py b/src/gfloat/round_ndarray.py index 957fb77..d64372a 100644 --- a/src/gfloat/round_ndarray.py +++ b/src/gfloat/round_ndarray.py @@ -90,7 +90,10 @@ def round_ndarray( Input NaNs will convert to NaNs in the target, not necessarily preserving payload. An input Infinity will convert to the largest float if :paramref:`sat`, otherwise to an Inf, if present, otherwise to a NaN. - Negative zero will be returned if the format has negative zero, otherwise zero. + If the format has zero, negative zero will be returned if it has negative zero, + otherwise zero. Formats without zero clamp finite values below the smallest + magnitude to that magnitude, preserving the sign for signed formats, regardless + of :paramref:`sat` and :paramref:`rnd`. Args: fi (FormatInfo): Describes the target format @@ -120,6 +123,9 @@ def round_ndarray( is_negative = xp.signbit(v) & fi.is_signed absv = xp_where(is_negative, -v, v) + if not fi.has_zero: + absv = xp_where(xp.isfinite(absv) & (absv < fi.smallest), fi.smallest, absv) + finite_nonzero = ~(xp.isnan(v) | xp.isinf(v) | (v == 0)) # Place 1.0 where finite_nonzero is False, to avoid log of {0,inf,nan} @@ -136,6 +142,10 @@ def round_ndarray( expval = expval - p + 1 fsignificand = _ldexp(absv_masked, -expval) + if not fi.has_subnormals and fi.has_zero: + zero_gap = absv_masked < fi.smallest + fsignificand = xp_where(zero_gap, absv_masked / fi.smallest, fsignificand) + floorfsignificand = xp.floor(fsignificand) isignificand = xp.astype(floorfsignificand, int_type) delta = fsignificand - floorfsignificand @@ -187,6 +197,10 @@ def round_ndarray( isignificand = xp_where(should_round_away, isignificand + 1, isignificand) fresult = _ldexp(xp.astype(isignificand, v.dtype), expval) + if not fi.has_subnormals and fi.has_zero: + fresult = xp_where( + zero_gap, xp.astype(isignificand, v.dtype) * fi.smallest, fresult + ) result = xp_where(finite_nonzero, fresult, absv) diff --git a/test/test_encode.py b/test/test_encode.py index 8ab344f..381a4c2 100644 --- a/test/test_encode.py +++ b/test/test_encode.py @@ -79,3 +79,33 @@ def test_encode_binary64_signed_zero(values: npt.ArrayLike) -> None: assert codes.dtype == np.dtype(np.uint64) assert codes.shape == v.shape np.testing.assert_array_equal(codes, v.view(np.uint64)) + + +@pytest.mark.parametrize("bias", [-3, 0, 3]) +@pytest.mark.parametrize("precision", [2, 3]) +@pytest.mark.parametrize( + "signed, negative_zero", [(False, False), (True, False), (True, True)] +) +def test_encode_zero_without_subnormals( + bias: int, precision: int, signed: bool, negative_zero: bool +) -> None: + fi = FormatInfo( + name="no_subnormals", + k=5, + precision=precision, + bias=bias, + is_signed=signed, + domain=Domain.Finite, + has_nz=negative_zero, + num_high_nans=0, + has_subnormals=False, + is_twos_complement=False, + ) + values = np.array([0.0, -0.0]) + expected = [fi.code_of_zero, fi.code_of_negzero if negative_zero else fi.code_of_zero] + assert [encode_float(fi, value) for value in values] == expected + np.testing.assert_array_equal(encode_ndarray(fi, values), expected) + for code, value in zip(expected, values): + decoded = decode_float(fi, code).fval + assert decoded == 0.0 + assert np.signbit(decoded) == (negative_zero and np.signbit(value)) diff --git a/test/test_round_underflow.py b/test/test_round_underflow.py new file mode 100644 index 0000000..1c6194e --- /dev/null +++ b/test/test_round_underflow.py @@ -0,0 +1,141 @@ +# Copyright (c) 2024 Graphcore Ltd. All rights reserved. + +from typing import Callable + +import numpy as np +import pytest + +from gfloat import FormatInfo, RoundMode, decode_float, encode_float +from gfloat import round_float, round_ndarray +from gfloat.formats import Domain, format_info_ocp_e8m0 + + +def scalar_round( + fi: FormatInfo, value: float, rnd: RoundMode, sat: bool, bits: int +) -> float: + return round_float(fi, value, rnd, sat, srbits=bits, srnumbits=8) + + +def array_round( + fi: FormatInfo, value: float, rnd: RoundMode, sat: bool, bits: int +) -> float: + return float( + round_ndarray( + fi, np.array([value]), rnd, sat, srbits=np.array([bits]), srnumbits=8 + )[0] + ) + + +def no_subnormal_format(precision: int, signed: bool = False) -> FormatInfo: + return FormatInfo( + name="no_subnormals", + k=5, + precision=precision, + bias=3, + is_signed=signed, + domain=Domain.Finite, + has_nz=signed, + num_high_nans=0, + has_subnormals=False, + is_twos_complement=False, + ) + + +@pytest.mark.parametrize("rounder", [scalar_round, array_round]) +@pytest.mark.parametrize("rnd", RoundMode) +@pytest.mark.parametrize("sat", [False, True]) +@pytest.mark.parametrize("sign", [1, -1]) +@pytest.mark.parametrize("fraction", [0.25, 0.5, 0.75]) +@pytest.mark.parametrize("bits", [0, 255]) +def test_round_no_subnormals_zero_gap( + rounder: Callable, rnd: RoundMode, sat: bool, sign: int, fraction: float, bits: int +) -> None: + fi = no_subnormal_format(precision=2, signed=True) + if rnd == RoundMode.TowardZero: + up = False + elif rnd == RoundMode.TowardPositive: + up = sign > 0 + elif rnd == RoundMode.TowardNegative: + up = sign < 0 + elif rnd == RoundMode.TiesToEven: + up = fraction > 0.5 + elif rnd == RoundMode.TiesToAway: + up = fraction >= 0.5 + elif rnd == RoundMode.ToOdd: + up = True + else: + up = bits == 255 + expected = sign * (fi.smallest if up else 0.0) + actual = rounder(fi, sign * fraction * fi.smallest, rnd, sat, bits) + assert actual == expected + assert np.signbit(actual) == np.signbit(expected) + assert actual in [decode_float(fi, code).fval for code in range(2**fi.bits)] + assert decode_float(fi, encode_float(fi, actual)).fval == actual + + +@pytest.mark.parametrize("rounder", [scalar_round, array_round]) +@pytest.mark.parametrize("rnd", RoundMode) +@pytest.mark.parametrize("sat", [False, True]) +@pytest.mark.parametrize("bits", [0, 255]) +@pytest.mark.parametrize( + "fi", [format_info_ocp_e8m0, no_subnormal_format(1), no_subnormal_format(1, True)] +) +def test_round_no_zero_underflow( + rounder: Callable, rnd: RoundMode, sat: bool, bits: int, fi: FormatInfo +) -> None: + values = [ + 0.0, + -0.0, + np.nextafter(0.0, 1.0), + fi.smallest / 8, + fi.smallest / 2, + 0.75 * fi.smallest, + np.nextafter(fi.smallest, 0.0), + fi.smallest, + ] + if fi.is_signed: + values += [-v for v in values] + for value in values: + actual = rounder(fi, value, rnd, sat, bits) + expected = -fi.smallest if fi.is_signed and np.signbit(value) else fi.smallest + assert actual == expected + assert decode_float(fi, encode_float(fi, actual)).fval == expected + + +@pytest.mark.parametrize("strict", [False, True]) +@pytest.mark.parametrize("sat", [False, True]) +def test_round_e8m0_mixed_array(strict: bool, sat: bool) -> None: + import array_api_strict + + xp = array_api_strict if strict else np + fi = format_info_ocp_e8m0 + values = xp.asarray( + [[0.0, fi.smallest / 2, fi.smallest], [1.0, np.inf, np.nan]], dtype=xp.float64 + ) + result = round_ndarray(fi, values, sat=sat) + assert result.shape == values.shape + assert result.dtype == values.dtype + np.testing.assert_equal( + np.asarray(result), + [ + [fi.smallest, fi.smallest, fi.smallest], + [1.0, fi.max if sat else np.nan, np.nan], + ], + ) + empty = xp.asarray([], dtype=xp.float64) + assert round_ndarray(fi, empty, sat=sat).shape == (0,) + + +@pytest.mark.parametrize("rnd", RoundMode) +@pytest.mark.parametrize("bits", [0, 255]) +def test_round_zero_gap_array_api(rnd: RoundMode, bits: int) -> None: + import array_api_strict as xp + + fi = no_subnormal_format(precision=2, signed=True) + values = np.array([-0.75, -0.5, -0.25, 0.0, 0.25, 0.5, 0.75]) * fi.smallest + result = round_ndarray( + fi, xp.asarray(values), rnd, srbits=xp.asarray([bits] * len(values)), srnumbits=8 # type: ignore[arg-type] + ) + expected = [scalar_round(fi, v, rnd, False, bits) for v in values] + np.testing.assert_equal(np.asarray(result), expected) + np.testing.assert_array_equal(np.signbit(np.asarray(result)), np.signbit(expected))