77
88public import RandomDo.NumLean.PCG64
99public 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. -/
3293def 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+
36115end NumLean
0 commit comments