Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions src/gfloat/encode.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
16 changes: 14 additions & 2 deletions src/gfloat/round.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand All @@ -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
Expand Down Expand Up @@ -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:
Expand Down
16 changes: 15 additions & 1 deletion src/gfloat/round_ndarray.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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}
Expand All @@ -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
Expand Down Expand Up @@ -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)

Expand Down
30 changes: 30 additions & 0 deletions test/test_encode.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
141 changes: 141 additions & 0 deletions test/test_round_underflow.py
Original file line number Diff line number Diff line change
@@ -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))
Loading