@@ -9,6 +9,7 @@ public import RandomDo.NumLean.PCG64
99public meta import RandomDo.NumLean.PCG64
1010public import FFI.Float
1111public import RandomDo.NumLean.Ziggurat
12+ public import RandomDo.NumLean.ZigguratSampler
1213
1314/-!
1415# Sample from specific distributions using the PCG-64 generator.
@@ -17,37 +18,25 @@ This file provides samplers for specific distributions using the PCG-64 generato
1718exactly what its numpy counterpart draws, so that a `RandPCG` program and a numpy `Generator`
1819seeded alike produce the same values.
1920
21+ The draws taken straight from the generator, `randUInt64`, `randUInt32` and `random`, are in
22+ `RandomDo.NumLean.PCG64`; the ziggurat the normal and the exponential share is in
23+ `RandomDo.NumLean.ZigguratSampler` and its tables in `RandomDo.NumLean.Ziggurat`.
24+
2025## Main definitions
21- * `randUInt64` / `randUInt32`: sample a `UInt64`, or a `UInt32` as numpy's `next_uint32` does.
22- * `random`: sample a `Float` in `[0, 1)`.
2326* `randInt`: sample an integer in `[low, high)` or `[low, high]`, as numpy's `Generator.integers`.
2427* `uniform`: sample a `Float` in `[low, high)`, as numpy's `Generator.uniform`.
2528* `standardNormal` / `normal`: sample a normal deviate, as numpy's `Generator.normal`.
2629* `standardExponential` / `exponential`: sample an exponential deviate, as numpy's
2730 `Generator.exponential`.
28- * `ziggurat`: the sampler shape both of those share.
2931 -/
3032
3133@[expose] public section
3234
3335namespace NumLean
3436
35- /-- Sample a `UInt64` from a PCG-64 generator. -/
36- def randUInt64 : RandPCG IO UInt64 := do
37- let (x, g) := (← get).down.nextUInt64
38- set (ULift.up g)
39- return x
40-
41- /-- Sample a `UInt32` from a PCG-64 generator, as numpy's `next_uint32` does for `PCG64`: the two
42- halves of each 64-bit output are handed out in turn, see `PCG64.nextUInt32`. -/
43- def randUInt32 : RandPCG IO UInt32 := do
44- let (x, g) := (← get).down.nextUInt32
45- set (ULift.up g)
46- return x
47-
4837/-- Sample a `UInt32` uniformly in `[0, rng]` by Lemire's nearly divisionless rejection method, as
4938numpy's `buffered_bounded_lemire_uint32`. -/
50- def randLemireUInt32 (rng : UInt32) : RandPCG IO UInt32 := do
39+ @[inline] def randLemireUInt32 (rng : UInt32) : RandPCG IO UInt32 := do
5140 let rngExcl := rng + 1
5241 let mut m := (← randUInt32).toUInt64 * rngExcl.toUInt64
5342 if m.toUInt32 < rngExcl then
@@ -58,7 +47,7 @@ def randLemireUInt32 (rng : UInt32) : RandPCG IO UInt32 := do
5847
5948/-- Sample a `UInt64` uniformly in `[0, rng]` by Lemire's method, as numpy's
6049`bounded_lemire_uint64`. -/
61- def randLemireUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
50+ @[inline] def randLemireUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
6251 let rngExcl := rng + 1
6352 let mut x ← randUInt64
6453 let mut leftover := x * rngExcl
@@ -71,7 +60,7 @@ def randLemireUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
7160
7261/-- Sample a `UInt64` uniformly in `[0, rng]`, as numpy's `random_bounded_uint64` with Lemire's
7362rejection. -/
74- def randBoundedUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
63+ @[inline] def randBoundedUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
7564 if rng == 0 then return 0
7665 else if rng == 0xFFFFFFFF then return (← randUInt32).toUInt64
7766 else if rng < 0xFFFFFFFF then return (← randLemireUInt32 rng.toUInt32).toUInt64
@@ -80,7 +69,7 @@ def randBoundedUInt64 (rng : UInt64) : RandPCG IO UInt64 := do
8069
8170/-- Sample an integer uniformly in `[low, high)`, or in `[low, high]` when `endpoint` is set, as
8271numpy's `Generator.integers`. -/
83- def randInt₀ (low high : Int) (endpoint : Bool := false ) : RandPCG IO Int := do
72+ @[inline] def randInt₀ (low high : Int) (endpoint : Bool := false ) : RandPCG IO Int := do
8473 let high := if endpoint then high else high - 1
8574 if low < Int64.minValue.toInt then throw <| IO.userError "low is out of bounds for int64"
8675 if high > Int64.maxValue.toInt then throw <| IO.userError "high is out of bounds for int64"
@@ -89,49 +78,19 @@ def randInt₀ (low high : Int) (endpoint : Bool := false) : RandPCG IO Int := d
8978 return low + x.toNat
9079
9180/-- Sample an integer uniformly in `[0, high)`, or in `[0, high]` when `endpoint` is set. -/
92- def randInt (high : Int) (endpoint : Bool := false ) : RandPCG IO Int := randInt₀ 0 high endpoint
81+ @[inline] def randInt (high : Int) (endpoint : Bool := false ) : RandPCG IO Int :=
82+ randInt₀ 0 high endpoint
9383
9484/-- Sample an integer uniformly in `[low, high)`, or in `[low, high]` when `endpoint` is set. -/
95- def randBoundedInt (low high : Int) (endpoint : Bool := false ) : RandPCG IO Int :=
85+ @[inline] def randBoundedInt (low high : Int) (endpoint : Bool := false ) : RandPCG IO Int :=
9686 randInt₀ low high endpoint
9787
98- /-- Sample a `Float` in `[0, 1)` from a PCG-64 generator. -/
99- def random : RandPCG IO Float := do
100- let x ← randUInt64
101- return (x >>> 11 ).toFloat * (Float.ofBits <| 0x3CA <<< (52 : UInt64))
102-
10388/-- Sample a `Float` uniformly in `[low, high)`. -/
104- def uniform (low high : Float) : RandPCG IO Float := do
89+ @[inline] def uniform (low high : Float) : RandPCG IO Float := do
10590 if low > high then throw <| IO.userError "low > high"
10691 let x ← random
10792 return Float.fma x (high - low) low
10893
109- /-- The rejection test on a strip that sticks out of the curve: the point drawn at height `u`
110- between the density at the strip's two edges lies under the curve. -/
111- @[inline] def zigguratWedge (f : Array Float) (idx : Nat) (u density : Float) : Bool :=
112- Float.fma (f[idx - 1 ]! - f[idx]!) u f[idx]! < density
113-
114- /-- The sampler shape shared by numpy's `random_standard_normal` and
115- `random_standard_exponential`, the ziggurat of Marsaglia and Tsang: the density is covered by 256
116- strips of equal area, and one 64-bit output supplies at once the strip `idx` and the integer `ri`
117- that `w` scales to an abscissa, `split` saying how those bits are laid out and whether the deviate
118- comes out negated.
119-
120- The draw is returned as it stands when `ri` falls below `k[idx]`, which is where about 99% of the
121- draws end. Otherwise the strip either is the base one, whose unbounded part `tail` samples from the
122- integer drawn, or sticks out of the curve, and then the point is tested against `density` and the
123- whole draw is started over on rejection. -/
124- @[specialize] partial def ziggurat (k : Array UInt64) (w f : Array Float)
125- (split : UInt64 → Nat × UInt64 × Bool) (density : Float → Float)
126- (tail : UInt64 → RandPCG IO Float) : RandPCG IO Float := do
127- let (idx, ri, negate) := split (← randUInt64)
128- let x := ri.toFloat * w[idx]!
129- let x := if negate then -x else x
130- if ri < k[idx]! then return x
131- if idx == 0 then tail ri
132- else if zigguratWedge f idx (← random) (density x) then return x
133- else ziggurat k w f split density tail
134-
13594/-- Sample the tail of the standard normal beyond `Ziggurat.norR`, as the `idx == 0` branch of
13695numpy's `random_standard_normal`: draw from an exponential tail until the pair of draws falls under
13796the normal's, which is Marsaglia's method for the tail. `negate` carries the sign numpy reads off
@@ -146,7 +105,7 @@ partial def normalTail (negate : Bool) : RandPCG IO Float := do
146105/-- Sample from the standard normal, as numpy's `random_standard_normal`. The 64-bit output gives
147106the strip in its low byte, then the sign, then a 52-bit abscissa; the tail takes its sign from a
148107further bit of that same abscissa, as numpy does. -/
149- def standardNormal : RandPCG IO Float :=
108+ @[inline] def standardNormal : RandPCG IO Float :=
150109 ziggurat Ziggurat.ki Ziggurat.wi Ziggurat.fi
151110 (fun r =>
152111 let idx := (r &&& 0xFF ).toNat
@@ -158,7 +117,7 @@ def standardNormal : RandPCG IO Float :=
158117/-- Sample from the standard exponential, as numpy's `random_standard_exponential`. The 64-bit
159118output is first shifted by three, then gives the strip in its low byte and a 53-bit abscissa; the
160119tail is the exponential's own, memoryless, so one draw beyond `Ziggurat.expR` suffices. -/
161- def standardExponential : RandPCG IO Float :=
120+ @[inline] def standardExponential : RandPCG IO Float :=
162121 ziggurat Ziggurat.ke Ziggurat.we Ziggurat.fe
163122 (fun r =>
164123 let r := r >>> 3
@@ -167,12 +126,12 @@ def standardExponential : RandPCG IO Float :=
167126 (fun _ => do return Ziggurat.expR - Float.log1p (-(← random)))
168127
169128/-- Draw random samples from a normal (Gaussian) distribution. -/
170- def normal (loc : Float := 0 ) (scale : Float := 1 ) : RandPCG IO Float := do
129+ @[inline] def normal (loc : Float := 0 ) (scale : Float := 1 ) : RandPCG IO Float := do
171130 if scale < 0 then throw <| IO.userError "scale < 0"
172131 return Float.fma scale (← standardNormal) loc
173132
174133/-- Draw samples from an exponential distribution. -/
175- def exponential (scale : Float := 1 ) : RandPCG IO Float := do
134+ @[inline] def exponential (scale : Float := 1 ) : RandPCG IO Float := do
176135 if scale < 0 then throw <| IO.userError "scale < 0"
177136 return scale * (← standardExponential)
178137
0 commit comments