"""Verify the conversion from a 64-bit draw to a float32 uniform in [0, 1).

Goualard (2022), section 2.2, states the rule for drawing a float uniformly from
the unit interval: draw from the set S = {k * 2**-p | 0 <= k < 2**p}, where p is
the width of the significand.  A set smaller than 2**p wastes representable
values, and a set larger than 2**p destroys uniformity.  For binary32, p = 24.

This script checks four claims about the two conversions in play.

Claim 1.  Taking the top 24 bits of the draw and scaling by 2**-24 produces
          exactly the set S with p = 24, so the values are equally spaced and
          the largest of them is the largest float32 below one.
Claim 2.  Both operands of that product are exactly representable in binary32,
          so the product introduces no rounding.
Claim 3.  Keeping 53 bits and narrowing the result to binary32 admits 2**53
          values, which is more than 2**24, and the top of that range rounds up
          to exactly 1.0 under round-to-nearest-even.
Claim 4.  The rounding in claim 3 reaches 1.0 for every draw at or above
          2**64 - 2**39, a fraction 2**-25 of the range.

Ref: goualard2022drawing (section 2.2), blackman2021scrambled (section 5.3).
"""
import numpy as np
import sympy as sp

P32 = 24        # binary32 significand width
P64 = 53        # binary64 significand width
W = 64          # draw width

# ---------------------------------------------------------------------------
# Claim 1: the top-24-bit conversion produces Goualard's set S with p = 24.

k = sp.symbols("k", integer=True, nonnegative=True)
scale = sp.Rational(1, 2 ** P32)

# The conversion is (x >> 40) * 2**-24.  Over x in [0, 2**64) the shift takes
# every value in [0, 2**24), so the output set is {k * 2**-24 | 0 <= k < 2**24}.
shift = W - P32
assert shift == 40
low, high = 0, 2 ** W - 1
assert low >> shift == 0
assert high >> shift == 2 ** P32 - 1

# Equal spacing: consecutive members differ by exactly 2**-24.
spacing = sp.simplify((k + 1) * scale - k * scale)
assert spacing == scale, spacing

# The largest member is the largest binary32 below one.
largest = sp.Rational(2 ** P32 - 1, 2 ** P32)
assert largest == 1 - scale
assert largest < 1
assert np.float32(largest) == np.nextafter(np.float32(1.0), np.float32(0.0))

# The count is exactly 2**p, so the set is neither wasteful nor over-full.
assert (2 ** P32 - 1) - 0 + 1 == 2 ** P32

# ---------------------------------------------------------------------------
# Claim 2: the product is exact.

# A non-negative integer below 2**24 is exactly representable in binary32, and
# 2**-24 is a power of two, so the product is a significand scaled by an
# exponent and carries no rounding error.
for value in (0, 1, 2 ** 12, 2 ** P32 - 2, 2 ** P32 - 1):
    exact = sp.Rational(value, 2 ** P32)
    computed = np.float32(value) * np.float32(2.0 ** -P32)
    assert np.float32(value) == value
    assert sp.Rational(float(computed)) == exact, (value, computed)

# ---------------------------------------------------------------------------
# Claim 3: keeping 53 bits and narrowing admits more than 2**24 values.

assert 2 ** P64 > 2 ** P32

# The largest binary64 the 53-bit conversion produces sits below one.
largest64 = sp.Rational(2 ** P64 - 1, 2 ** P64)
assert largest64 < 1

# It is above the midpoint between the largest binary32 below one and one, so
# narrowing rounds it up to exactly 1.0.
midpoint = 1 - sp.Rational(1, 2 ** (P32 + 1))
assert largest64 > midpoint
assert np.float32(np.float64(largest64)) == np.float32(1.0)

# ---------------------------------------------------------------------------
# Claim 4: the rounding starts at 2**64 - 2**39 and covers 2**-25 of the range.

# The 53-bit conversion is (x >> 11) * 2**-53.  Solve for the smallest x whose
# image is at or above the midpoint.  Round-to-nearest-even carries the midpoint
# itself up, because 1.0 has an even significand and the float below it does not.
y = sp.symbols("y", integer=True, nonnegative=True)
threshold = sp.solve(sp.Eq(y * sp.Rational(1, 2 ** P64), midpoint), y)
assert len(threshold) == 1
first_bad_shifted = threshold[0]
assert first_bad_shifted == 2 ** P64 - 2 ** 28, first_bad_shifted

# The shift discards the low 11 bits, so the smallest draw with that image is
# the shifted value scaled back up.
first_bad = first_bad_shifted * 2 ** (W - P64)
assert first_bad == 2 ** W - 2 ** 39, first_bad

for draw in (int(first_bad), int(first_bad) + 1, 2 ** W - 1):
    u64 = np.float64(draw >> 11) * np.float64(2.0 ** -P64)
    assert u64 < 1.0
    assert np.float32(u64) == np.float32(1.0), draw

# One draw below the threshold still narrows to a value below one.
below = int(first_bad) - 2 ** 11
u64 = np.float64(below >> 11) * np.float64(2.0 ** -P64)
assert np.float32(u64) < np.float32(1.0)

# The affected fraction of the draw range.
fraction = sp.Rational(2 ** 39, 2 ** W)
assert fraction == sp.Rational(1, 2 ** 25)
assert abs(float(1 / fraction) - 3.3554432e7) < 1.0

# The top-24-bit conversion has no such threshold: no draw reaches one.
worst = np.float32(np.uint64(2 ** W - 1) >> np.uint64(shift)) * np.float32(2.0 ** -P32)
assert worst < 1.0
assert worst == np.float32(largest)

print("largest top-24-bit value  : %.9f = 1 - 2**-24" % float(largest))
print("narrowing threshold       : 2**64 - 2**39 = %d" % int(first_bad))
print("affected fraction         : 2**-25 = 1 in %.4g" % float(1 / fraction))
print("all claims verified")
