Skip to content

Commit 012ba5e

Browse files
committed
randInt and uniform
1 parent 1801a5e commit 012ba5e

5 files changed

Lines changed: 196 additions & 3 deletions

File tree

‎RandomDo/NumLean/Distributions.lean‎

Lines changed: 81 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,15 +7,19 @@ module
77

88
public import RandomDo.NumLean.PCG64
99
public meta import RandomDo.NumLean.PCG64
10+
public import Batteries.Data.Float.Basic
1011

1112
/-!
1213
# Sample from specific distributions using the PCG-64 generator.
1314
14-
This files provides samplers for specific distributions using the PCG-64 generator.
15+
This file provides samplers for specific distributions using the PCG-64 generator. Each one draws
16+
exactly what its numpy counterpart draws, so that a `RandPCG` program and a numpy `Generator`
17+
seeded alike produce the same values.
1518
1619
## Main definitions
17-
* `randUInt64`: sample a `UInt64`.
20+
* `randUInt64` / `randUInt32`: sample a `UInt64`, or a `UInt32` as numpy's `next_uint32` does.
1821
* `random`: sample a `Float` in `[0, 1)`.
22+
* `randInt`: sample an integer in `[low, high)` or `[low, high]`, as numpy's `Generator.integers`.
1923
-/
2024

2125
@[expose] public section
@@ -28,9 +32,84 @@ def randUInt64 : RandPCG IO UInt64 := do
2832
set (ULift.up g)
2933
return x
3034

35+
/-- Sample a `UInt32` from a PCG-64 generator, as numpy's `next_uint32` does for `PCG64`: the two
36+
halves of each 64-bit output are handed out in turn, see `PCG64.nextUInt32`. -/
37+
def randUInt32 : RandPCG IO UInt32 := do
38+
let (x, g) := (← get).down.nextUInt32
39+
set (ULift.up g)
40+
return x
41+
42+
/-- Sample a `UInt32` uniformly in `[0, rng]` by Lemire's nearly divisionless rejection method, as
43+
numpy's `buffered_bounded_lemire_uint32`. -/
44+
def randLemireUInt32 (rng : UInt32) : RandPCG IO UInt32 := do
45+
let rngExcl := rng + 1
46+
let mut m := (← randUInt32).toUInt64 * rngExcl.toUInt64
47+
if m.toUInt32 < rngExcl then
48+
let threshold := (0xFFFFFFFF - rng) % rngExcl
49+
while m.toUInt32 < threshold do
50+
m := (← randUInt32).toUInt64 * rngExcl.toUInt64
51+
return (m >>> 32).toUInt32
52+
53+
/-- Sample a `UInt64` uniformly in `[0, rng]` by Lemire's method, as numpy's
54+
`bounded_lemire_uint64`. -/
55+
def randLemireUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
56+
let rngExcl := rng + 1
57+
let mut x ← randUInt64
58+
let mut leftover := x * rngExcl
59+
if leftover < rngExcl then
60+
let threshold := (0xFFFFFFFFFFFFFFFF - rng) % rngExcl
61+
while leftover < threshold do
62+
x ← randUInt64
63+
leftover := x * rngExcl
64+
return PCG64.mulHi x rngExcl
65+
66+
/-- Sample a `UInt64` uniformly in `[0, rng]`, as numpy's `random_bounded_uint64` with Lemire's
67+
rejection. -/
68+
def randBoundedUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
69+
if rng == 0 then return 0
70+
else if rng == 0xFFFFFFFF then return (← randUInt32).toUInt64
71+
else if rng < 0xFFFFFFFF then return (← randLemireUInt32 rng.toUInt32).toUInt64
72+
else if rng == 0xFFFFFFFFFFFFFFFF then randUInt64
73+
else randLemireUInt64 rng
74+
75+
/-- Sample an integer uniformly in `[low, high)`, or in `[low, high]` when `endpoint` is set, as
76+
numpy's `Generator.integers`. -/
77+
def randInt₀ (low high : Int) (endpoint : Bool := false) : RandPCG IO Int := do
78+
let high := if endpoint then high else high - 1
79+
if low < Int64.minValue.toInt then throw <| IO.userError "low is out of bounds for int64"
80+
if high > Int64.maxValue.toInt then throw <| IO.userError "high is out of bounds for int64"
81+
if low > high then throw <| IO.userError (if endpoint then "low > high" else "low >= high")
82+
let x ← randBoundedUInt64 (high - low).toNat.toUInt64
83+
return low + x.toNat
84+
85+
/-- Sample an integer uniformly in `[0, high)`, or in `[0, high]` when `endpoint` is set. -/
86+
def randInt (high : Int) (endpoint : Bool := false) : RandPCG IO Int := randInt₀ 0 high endpoint
87+
88+
/-- Sample an integer uniformly in `[low, high)`, or in `[low, high]` when `endpoint` is set. -/
89+
def randBoundedInt (low high : Int) (endpoint : Bool := false) : RandPCG IO Int :=
90+
randInt₀ low high endpoint
91+
3192
/-- Sample a `Float` in `[0, 1)` from a PCG-64 generator. -/
3293
def random : RandPCG IO Float := do
3394
let x ← randUInt64
3495
return (x >>> 11).toFloat * (Float.ofBits <| 0x3CA <<< (52 : UInt64))
3596

97+
/-- `x * y + z`, rounded once, as the C `fma`: the product and the sum are computed exactly, as
98+
integers scaled by a power of two, and only the final value is rounded to a `Float`. -/
99+
def fma (x y z : Float) : Float :=
100+
match x.toRatParts, y.toRatParts, z.toRatParts with
101+
| some (vx, ex), some (vy, ey), some (vz, ez) =>
102+
-- the product is `vx * vy * 2 ^ (ex + ey)`, so the sum is exact over the common exponent `e`
103+
let ep := ex + ey
104+
let e := min ep ez
105+
let n := vx * vy * 2 ^ (ep - e).toNat + vz * 2 ^ (ez - e).toNat
106+
if e ≥ 0 then Int.divFloat (n * 2 ^ e.toNat) 1 else Int.divFloat n (2 ^ (-e).toNat)
107+
| _, _, _ => x * y + z
108+
109+
/-- Sample a `Float` uniformly in `[low, high)`. -/
110+
def uniform (low high : Float) : RandPCG IO Float := do
111+
if low > high then throw <| IO.userError "low > high"
112+
let x ← random
113+
return fma x (high - low) low
114+
36115
end NumLean

‎RandomDo/NumLean/PCG64.lean‎

Lines changed: 22 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ module
77

88
public import Mathlib.Control.Random
99
public import RandomDo.NumLean.SeedSequence
10+
public import Batteries.Lean.LawfulMonad
1011

1112
/-!
1213
# The PCG-64 pseudo-random number generator
@@ -24,6 +25,8 @@ native machine arithmetic.
2425
## Main definitions
2526
2627
* `PCG64`: the generator state, and its `RandomGen` instance
28+
* `PCG64.nextUInt64` / `PCG64.nextUInt32`: the 64-bit output, and numpy's 32-bit output, which
29+
hands out the two halves of a 64-bit output in turn
2730
* `PCG64.seedWords` / `PCG64.seed` / `mkPCG64`: seeding, following the reference
2831
`pcg_setseq_128_srandom_r`
2932
@@ -40,7 +43,8 @@ native machine arithmetic.
4043
namespace NumLean
4144

4245
/-- The state of a PCG-64 generator: a 128-bit LCG state and a 128-bit increment (which selects the
43-
stream and must be odd), each stored as a pair of 64-bit words. -/
46+
stream and must be odd), each stored as a pair of 64-bit words, together with the one-word buffer
47+
numpy's bit generator keeps for its 32-bit outputs, see `nextUInt32`. -/
4448
structure PCG64 where
4549
/-- High 64 bits of the LCG state. -/
4650
stateHi : UInt64
@@ -50,6 +54,12 @@ structure PCG64 where
5054
incHi : UInt64
5155
/-- Low 64 bits of the increment; it is odd for any generator built through the API below. -/
5256
incLo : UInt64
57+
/-- Whether `uinteger` holds the unused upper half of the last output drawn through `nextUInt32`.
58+
Mirrors the `has_uint32` field of numpy's `pcg64_state`. -/
59+
hasUInt32 : Bool := false
60+
/-- The upper 32 bits of the last output drawn through `nextUInt32`, meaningful only while
61+
`hasUInt32` is set. Mirrors the `uinteger` field of numpy's `pcg64_state`. -/
62+
uinteger : UInt32 := 0
5363
deriving Repr, DecidableEq
5464

5565
namespace PCG64
@@ -97,6 +107,17 @@ implementation, the state is stepped before the output permutation is applied. -
97107
let g := g.step
98108
(g.output, g)
99109

110+
/-- The next 32-bit output, together with the advanced generator, as numpy's `pcg64_next32`: a
111+
64-bit output is drawn and its low half returned, its high half being kept in `uinteger` to be
112+
returned by the next call, which then does not step the generator. The buffer survives
113+
intermediate `nextUInt64` calls, as it does in numpy. -/
114+
@[inline] def nextUInt32 (g : PCG64) : UInt32 × PCG64 :=
115+
if g.hasUInt32 then
116+
(g.uinteger, { g with hasUInt32 := false })
117+
else
118+
let (x, g) := g.nextUInt64
119+
(x.toUInt32, { g with hasUInt32 := true, uinteger := (x >>> 32).toUInt32 })
120+
100121
/-- Seed a generator from four 64-bit words, following the reference
101122
`pcg_setseq_128_srandom_r`: `stateHi:stateLo` is the initial state and `seqHi:seqLo` selects the
102123
stream, whose increment is `2 * seq + 1`. -/
Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,9 @@ def main (args : List String) : IO Unit := do
1616
shutil.rmtree(DIR, ignore_errors=True)
1717
os.makedirs(DIR)
1818
open(f"{DIR}/dump.lean", "w").write(LEAN.replace("@DIR@", DIR).replace("@N@", str(N)))
19+
subprocess.run(
20+
["lake", "build"], check=True
21+
)
1922
subprocess.run(
2023
["lake", "env", "lean", "--run", f"{DIR}/dump.lean", *map(str, SEEDS)], check=True
2124
)

‎scripts/check_int.py‎

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,45 @@
1+
import os, shutil, subprocess, sys
2+
from decimal import Decimal
3+
import numpy as np
4+
5+
N, SEEDS, DIR = 1_000_000, np.random.randint(1_000_000_000, size=5), "test_data"
6+
7+
LEAN = """import RandomDo
8+
import Batteries.Data.Float.Basic
9+
def main (args : List String) : IO Unit := do
10+
for s in args do
11+
IO.FS.withFile (System.FilePath.mk s!"@DIR@/pcg64-{s}.txt") .write fun h ↦
12+
IO.runRandPCGWith s.toNat! do
13+
for _ in List.range @N@ do h.putStrLn <| toString (← NumLean.randInt 1000000)
14+
"""
15+
16+
shutil.rmtree(DIR, ignore_errors=True)
17+
os.makedirs(DIR)
18+
open(f"{DIR}/dump.lean", "w").write(LEAN.replace("@DIR@", DIR).replace("@N@", str(N)))
19+
subprocess.run(
20+
["lake", "build"], check=True
21+
)
22+
subprocess.run(
23+
["lake", "env", "lean", "--run", f"{DIR}/dump.lean", *map(str, SEEDS)], check=True
24+
)
25+
26+
27+
def check_distrib(path, xs):
28+
with open(path) as f:
29+
for i, (line, x) in enumerate(zip(f, xs)):
30+
if line.rstrip("\n") != format(Decimal(float(x)), "f"):
31+
return i
32+
return None
33+
34+
35+
ok = True
36+
for seed in SEEDS:
37+
xs = np.random.default_rng(seed).integers(1000000, size=N)
38+
bad_line = check_distrib(f"{DIR}/pcg64-{seed}.txt", xs)
39+
if bad_line is None:
40+
print(f"seed {seed} {N} identical draws")
41+
else:
42+
print(f"seed {seed} DIVERGENCE at line {bad_line}")
43+
ok = ok and bad_line is None
44+
45+
sys.exit(0 if ok else 1)

‎scripts/check_uniform.py‎

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,45 @@
1+
import os, shutil, subprocess, sys
2+
from decimal import Decimal
3+
import numpy as np
4+
5+
N, SEEDS, DIR = 1_000_000, np.random.randint(1_000_000_000, size=5), "test_data"
6+
7+
LEAN = """import RandomDo
8+
import Batteries.Data.Float.Basic
9+
def main (args : List String) : IO Unit := do
10+
for s in args do
11+
IO.FS.withFile (System.FilePath.mk s!"@DIR@/pcg64-{s}.txt") .write fun h ↦
12+
IO.runRandPCGWith s.toNat! do
13+
for _ in List.range @N@ do h.putStrLn (← NumLean.uniform (-1000000) 1000000).toStringFull
14+
"""
15+
16+
shutil.rmtree(DIR, ignore_errors=True)
17+
os.makedirs(DIR)
18+
open(f"{DIR}/dump.lean", "w").write(LEAN.replace("@DIR@", DIR).replace("@N@", str(N)))
19+
subprocess.run(
20+
["lake", "build"], check=True
21+
)
22+
subprocess.run(
23+
["lake", "env", "lean", "--run", f"{DIR}/dump.lean", *map(str, SEEDS)], check=True
24+
)
25+
26+
27+
def check_distrib(path, xs):
28+
with open(path) as f:
29+
for i, (line, x) in enumerate(zip(f, xs)):
30+
if line.rstrip("\n") != format(Decimal(float(x)), "f"):
31+
return i
32+
return None
33+
34+
35+
ok = True
36+
for seed in SEEDS:
37+
xs = np.random.default_rng(seed).uniform(-1000000, 1000000, size=N)
38+
bad_line = check_distrib(f"{DIR}/pcg64-{seed}.txt", xs)
39+
if bad_line is None:
40+
print(f"seed {seed} {N} identical draws")
41+
else:
42+
print(f"seed {seed} DIVERGENCE at line {bad_line}")
43+
ok = ok and bad_line is None
44+
45+
sys.exit(0 if ok else 1)

0 commit comments

Comments
 (0)