77
88public import RandomDo.NumLean.PCG64
99public meta import RandomDo.NumLean.PCG64
10- public import Batteries.Data.Float.Basic
10+ public import FFI.Float
11+ public import RandomDo.NumLean.Ziggurat
1112
1213/-!
1314# Sample from specific distributions using the PCG-64 generator.
@@ -20,6 +21,11 @@ seeded alike produce the same values.
2021* `randUInt64` / `randUInt32`: sample a `UInt64`, or a `UInt32` as numpy's `next_uint32` does.
2122* `random`: sample a `Float` in `[0, 1)`.
2223* `randInt`: sample an integer in `[low, high)` or `[low, high]`, as numpy's `Generator.integers`.
24+ * `uniform`: sample a `Float` in `[low, high)`, as numpy's `Generator.uniform`.
25+ * `standardNormal` / `normal`: sample a normal deviate, as numpy's `Generator.normal`.
26+ * `standardExponential` / `exponential`: sample an exponential deviate, as numpy's
27+ `Generator.exponential`.
28+ * `ziggurat`: the sampler shape both of those share.
2329 -/
2430
2531@[expose] public section
@@ -94,22 +100,80 @@ def random : RandPCG IO Float := do
94100 let x ← randUInt64
95101 return (x >>> 11 ).toFloat * (Float.ofBits <| 0x3CA <<< (52 : UInt64))
96102
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-
109103/-- Sample a `Float` uniformly in `[low, high)`. -/
110104def uniform (low high : Float) : RandPCG IO Float := do
111105 if low > high then throw <| IO.userError "low > high"
112106 let x ← random
113- return fma x (high - low) low
107+ return Float.fma x (high - low) low
108+
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+
135+ /-- Sample the tail of the standard normal beyond `Ziggurat.norR`, as the `idx == 0` branch of
136+ numpy's `random_standard_normal`: draw from an exponential tail until the pair of draws falls under
137+ the normal's, which is Marsaglia's method for the tail. `negate` carries the sign numpy reads off
138+ the integer already drawn, which the loop does not redraw. -/
139+ partial def normalTail (negate : Bool) : RandPCG IO Float := do
140+ let xx := -Ziggurat.norInvR * Float.log1p (-(← random))
141+ let yy := -Float.log1p (-(← random))
142+ if yy + yy > xx * xx then
143+ return if negate then -(Ziggurat.norR + xx) else Ziggurat.norR + xx
144+ normalTail negate
145+
146+ /-- Sample from the standard normal, as numpy's `random_standard_normal`. The 64-bit output gives
147+ the strip in its low byte, then the sign, then a 52-bit abscissa; the tail takes its sign from a
148+ further bit of that same abscissa, as numpy does. -/
149+ def standardNormal : RandPCG IO Float :=
150+ ziggurat Ziggurat.ki Ziggurat.wi Ziggurat.fi
151+ (fun r =>
152+ let idx := (r &&& 0xFF ).toNat
153+ let r := r >>> 8
154+ (idx, (r >>> 1 ) &&& 0x000FFFFFFFFFFFFF , (r &&& 1 ) == 1 ))
155+ (fun x => Float.exp ((-0 .5 ) * x * x))
156+ (fun rabs => normalTail (((rabs >>> 8 ) &&& 1 ) == 1 ))
157+
158+ /-- Sample from the standard exponential, as numpy's `random_standard_exponential`. The 64-bit
159+ output is first shifted by three, then gives the strip in its low byte and a 53-bit abscissa; the
160+ tail is the exponential's own, memoryless, so one draw beyond `Ziggurat.expR` suffices. -/
161+ def standardExponential : RandPCG IO Float :=
162+ ziggurat Ziggurat.ke Ziggurat.we Ziggurat.fe
163+ (fun r =>
164+ let r := r >>> 3
165+ ((r &&& 0xFF ).toNat, r >>> 8 , false ))
166+ (fun x => Float.exp (-x))
167+ (fun _ => do return Ziggurat.expR - Float.log1p (-(← random)))
168+
169+ /-- Draw random samples from a normal (Gaussian) distribution. -/
170+ def normal (loc : Float := 0 ) (scale : Float := 1 ) : RandPCG IO Float := do
171+ if scale < 0 then throw <| IO.userError "scale < 0"
172+ return Float.fma scale (← standardNormal) loc
173+
174+ /-- Draw samples from an exponential distribution. -/
175+ def exponential (scale : Float := 1 ) : RandPCG IO Float := do
176+ if scale < 0 then throw <| IO.userError "scale < 0"
177+ return scale * (← standardExponential)
114178
115179end NumLean
0 commit comments