Skip to content

Commit ef5c373

Browse files
committed
Inline and separate Ziggurat
1 parent 4f0aeb0 commit ef5c373

5 files changed

Lines changed: 127 additions & 60 deletions

File tree

‎RandomDo.lean‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ public import RandomDo.NumLean.Distributions
1010
public import RandomDo.NumLean.PCG64
1111
public import RandomDo.NumLean.SeedSequence
1212
public import RandomDo.NumLean.Ziggurat
13+
public import RandomDo.NumLean.ZigguratSampler
1314
public import RandomDo.Tactic.Deriving
1415
public import RandomDo.Tactic.Elab
1516
public import RandomDo.Tactic.ForInStep

‎RandomDo/NumLean/Distributions.lean‎

Lines changed: 17 additions & 58 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ public import RandomDo.NumLean.PCG64
99
public meta import RandomDo.NumLean.PCG64
1010
public import FFI.Float
1111
public 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
1718
exactly what its numpy counterpart draws, so that a `RandPCG` program and a numpy `Generator`
1819
seeded 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

3335
namespace 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
4938
numpy'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
7362
rejection. -/
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
8271
numpy'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
13695
numpy's `random_standard_normal`: draw from an exponential tail until the pair of draws falls under
13796
the 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
147106
the strip in its low byte, then the sign, then a 52-bit abscissa; the tail takes its sign from a
148107
further 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
159118
output is first shifted by three, then gives the strip in its low byte and a 53-bit abscissa; the
160119
tail 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

‎RandomDo/NumLean/PCG64.lean‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,8 @@ native machine arithmetic.
2929
hands out the two halves of a 64-bit output in turn
3030
* `PCG64.seedWords` / `PCG64.seed` / `mkPCG64`: seeding, following the reference
3131
`pcg_setseq_128_srandom_r`
32+
* `randUInt64` / `randUInt32` / `random`: the three outputs a `RandPCG` computation draws straight
33+
from the generator.
3234
3335
## References
3436
@@ -200,6 +202,25 @@ instance : RandomGen PCG64 where
200202
`RandPCG m α` should be thought of a random value in `m α`. -/
201203
abbrev RandPCG := RandGT PCG64
202204

205+
/-- Sample a `UInt64` from a PCG-64 generator. -/
206+
@[inline] def randUInt64 : RandPCG IO UInt64 := do
207+
let (x, g) := (← get).down.nextUInt64
208+
set (ULift.up g)
209+
return x
210+
211+
/-- Sample a `UInt32` from a PCG-64 generator, as numpy's `next_uint32` does for `PCG64`: the two
212+
halves of each 64-bit output are handed out in turn, see `PCG64.nextUInt32`. -/
213+
@[inline] def randUInt32 : RandPCG IO UInt32 := do
214+
let (x, g) := (← get).down.nextUInt32
215+
set (ULift.up g)
216+
return x
217+
218+
/-- Sample a `Float` in `[0, 1)` from a PCG-64 generator, as numpy's `next_double`: the top 53 bits
219+
of a 64-bit output, scaled by `2 ^ (-53)`. -/
220+
@[inline] def random : RandPCG IO Float := do
221+
let x ← randUInt64
222+
return (x >>> 11).toFloat * (Float.ofBits <| 0x3CA <<< (52 : UInt64))
223+
203224
end NumLean
204225

205226
namespace IO
Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
public import RandomDo.NumLean.PCG64
9+
public import FFI.Float
10+
11+
/-!
12+
# The ziggurat sampler
13+
14+
The sampler shape numpy's `random_standard_normal` and `random_standard_exponential` share, the
15+
ziggurat of Marsaglia and Tsang. The density is covered by 256 strips of equal area; a draw picks a
16+
strip and a point in it, and is accepted at once when the point falls in the strip's rectangular
17+
part, which is where about 99% of the draws end.
18+
19+
The tables the two samplers read live in `RandomDo.NumLean.Ziggurat`, and the samplers themselves
20+
in `RandomDo.NumLean.Distributions`; this file holds only the shape they share, which is generic in
21+
the tables `k`, `w` and `f` and in the three functions that say how a 64-bit output is cut up, what
22+
the density is, and how the base strip's tail is sampled.
23+
24+
## Main definitions
25+
26+
* `ziggurat`: the sampler, whose fast path is inlined into its caller
27+
* `zigguratSlow`: the branches the remaining ~1% of the draws take
28+
* `zigguratWedge`: the rejection test on a strip that sticks out of the curve
29+
30+
## References
31+
32+
* G. Marsaglia and W. W. Tsang, *The Ziggurat Method for Generating Random Variables*, Journal of
33+
Statistical Software, 2000.
34+
* numpy's samplers: `numpy/random/src/distributions/distributions.c`
35+
-/
36+
37+
@[expose] public section
38+
39+
namespace NumLean
40+
41+
/-- The rejection test on a strip that sticks out of the curve: the point drawn at height `u`
42+
between the density at the strip's two edges lies under the curve. -/
43+
@[inline] def zigguratWedge (f : Array Float) (idx : Nat) (u density : Float) : Bool :=
44+
Float.fma (f[idx - 1]! - f[idx]!) u f[idx]! < density
45+
46+
/-- The rare branches of `ziggurat`, reached by about 1% of the draws: the strip either is the base
47+
one, whose unbounded part `tail` samples from the integer drawn, or sticks out of the curve, and
48+
then the point is tested against `density` and the whole draw is started over on rejection.
49+
50+
`idx`, `ri` and `x` are the strip, the integer and the abscissa `ziggurat` has already drawn and
51+
found not to land in the rectangular part of its strip. Each redraw retries the fast path here
52+
rather than returning to `ziggurat`, so the two together run exactly the loop of numpy's samplers.
53+
54+
This is kept apart from `ziggurat` so that the fast path can be inlined into its caller: a
55+
recursive function cannot be, and behind a call boundary every draw would have to box the generator
56+
state and its result, which costs several times the draw itself. -/
57+
@[specialize] partial def zigguratSlow (k : Array UInt64) (w f : Array Float)
58+
(split : UInt64 → Nat × UInt64 × Bool) (density : Float → Float)
59+
(tail : UInt64 → RandPCG IO Float) (idx : Nat) (ri : UInt64) (x : Float) :
60+
RandPCG IO Float := do
61+
if idx == 0 then tail ri
62+
else if zigguratWedge f idx (← random) (density x) then return x
63+
else
64+
let (idx, ri, negate) := split (← randUInt64)
65+
let x := ri.toFloat * w[idx]!
66+
let x := if negate then -x else x
67+
if ri < k[idx]! then return x
68+
else zigguratSlow k w f split density tail idx ri x
69+
70+
/-- The sampler shape shared by numpy's `random_standard_normal` and
71+
`random_standard_exponential`, the ziggurat of Marsaglia and Tsang: the density is covered by 256
72+
strips of equal area, and one 64-bit output supplies at once the strip `idx` and the integer `ri`
73+
that `w` scales to an abscissa, `split` saying how those bits are laid out and whether the deviate
74+
comes out negated.
75+
76+
The draw is returned as it stands when `ri` falls below `k[idx]`, which is where about 99% of the
77+
draws end; `zigguratSlow` takes over the remaining ones. -/
78+
@[inline] def ziggurat (k : Array UInt64) (w f : Array Float)
79+
(split : UInt64 → Nat × UInt64 × Bool) (density : Float → Float)
80+
(tail : UInt64 → RandPCG IO Float) : RandPCG IO Float := do
81+
let (idx, ri, negate) := split (← randUInt64)
82+
let x := ri.toFloat * w[idx]!
83+
let x := if negate then -x else x
84+
if ri < k[idx]! then return x
85+
else zigguratSlow k w f split density tail idx ri x
86+
87+
end NumLean

‎lakefile.toml‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,7 @@ name = "RandomDo"
2424
[[lean_lib]]
2525
name = "Test"
2626

27-
# The check scripts generate `test_data/Dump.lean` and run it through this target, so that the
28-
# sampler runs as compiled code: `lake env lean --run` would interpret it, which is far slower.
27+
# Used to run the tests in `scripts`
2928
[[lean_exe]]
3029
name = "dump"
3130
root = "Dump"

0 commit comments

Comments
 (0)