|
| 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 meta import RandomDo.NumLean.PCG64 |
| 10 | +public import FFI.Float |
| 11 | +public import RandomDo.NumLean.Ziggurat |
| 12 | +public import RandomDo.NumLean.ZigguratSampler |
| 13 | + |
| 14 | +/-! |
| 15 | +# Sample from specific distributions using the PCG-64 generator. |
| 16 | +
|
| 17 | +This file provides samplers for specific distributions using the PCG-64 generator. Each one draws |
| 18 | +exactly what its numpy counterpart draws, so that a `RandPCG` program and a numpy `Generator` |
| 19 | +seeded alike produce the same values. |
| 20 | +
|
| 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 | +
|
| 25 | +## Main definitions |
| 26 | +* `randInt`: sample an integer in `[low, high)` or `[low, high]`, as numpy's `Generator.integers`. |
| 27 | +* `uniform`: sample a `Float` in `[low, high)`, as numpy's `Generator.uniform`. |
| 28 | +* `standardNormal` / `normal`: sample a normal deviate, as numpy's `Generator.normal`. |
| 29 | +* `standardExponential` / `exponential`: sample an exponential deviate, as numpy's |
| 30 | + `Generator.exponential`. |
| 31 | +-/ |
| 32 | + |
| 33 | +@[expose] public section |
| 34 | + |
| 35 | +namespace NumLean |
| 36 | + |
| 37 | +/-- Sample a `UInt32` uniformly in `[0, rng]` by Lemire's nearly divisionless rejection method, as |
| 38 | +numpy's `buffered_bounded_lemire_uint32`. -/ |
| 39 | +@[inline] def randLemireUInt32 (rng : UInt32) : RandPCG IO UInt32 := do |
| 40 | + let rngExcl := rng + 1 |
| 41 | + let mut m := (← randUInt32).toUInt64 * rngExcl.toUInt64 |
| 42 | + if m.toUInt32 < rngExcl then |
| 43 | + let threshold := (0xFFFFFFFF - rng) % rngExcl |
| 44 | + while m.toUInt32 < threshold do |
| 45 | + m := (← randUInt32).toUInt64 * rngExcl.toUInt64 |
| 46 | + return (m >>> 32).toUInt32 |
| 47 | + |
| 48 | +/-- Sample a `UInt64` uniformly in `[0, rng]` by Lemire's method, as numpy's |
| 49 | +`bounded_lemire_uint64`. -/ |
| 50 | +@[inline] def randLemireUInt64 (rng : UInt64) : RandPCG IO UInt64 := do |
| 51 | + let rngExcl := rng + 1 |
| 52 | + let mut x ← randUInt64 |
| 53 | + let mut leftover := x * rngExcl |
| 54 | + if leftover < rngExcl then |
| 55 | + let threshold := (0xFFFFFFFFFFFFFFFF - rng) % rngExcl |
| 56 | + while leftover < threshold do |
| 57 | + x ← randUInt64 |
| 58 | + leftover := x * rngExcl |
| 59 | + return PCG64.mulHi x rngExcl |
| 60 | + |
| 61 | +/-- Sample a `UInt64` uniformly in `[0, rng]`, as numpy's `random_bounded_uint64` with Lemire's |
| 62 | +rejection. -/ |
| 63 | +@[inline] def randBoundedUInt64 (rng : UInt64) : RandPCG IO UInt64 := do |
| 64 | + if rng == 0 then return 0 |
| 65 | + else if rng == 0xFFFFFFFF then return (← randUInt32).toUInt64 |
| 66 | + else if rng < 0xFFFFFFFF then return (← randLemireUInt32 rng.toUInt32).toUInt64 |
| 67 | + else if rng == 0xFFFFFFFFFFFFFFFF then randUInt64 |
| 68 | + else randLemireUInt64 rng |
| 69 | + |
| 70 | +/-- Sample an integer uniformly in `[low, high)`, or in `[low, high]` when `endpoint` is set, as |
| 71 | +numpy's `Generator.integers`. -/ |
| 72 | +@[inline] def randInt₀ (low high : Int) (endpoint : Bool := false) : RandPCG IO Int := do |
| 73 | + let high := if endpoint then high else high - 1 |
| 74 | + if low < Int64.minValue.toInt then throw <| IO.userError "low is out of bounds for int64" |
| 75 | + if high > Int64.maxValue.toInt then throw <| IO.userError "high is out of bounds for int64" |
| 76 | + if low > high then throw <| IO.userError (if endpoint then "low > high" else "low >= high") |
| 77 | + let x ← randBoundedUInt64 (high - low).toNat.toUInt64 |
| 78 | + return low + x.toNat |
| 79 | + |
| 80 | +/-- Sample an integer uniformly in `[0, high)`, or in `[0, high]` when `endpoint` is set. -/ |
| 81 | +@[inline] def randInt (high : Int) (endpoint : Bool := false) : RandPCG IO Int := |
| 82 | + randInt₀ 0 high endpoint |
| 83 | + |
| 84 | +/-- Sample an integer uniformly in `[low, high)`, or in `[low, high]` when `endpoint` is set. -/ |
| 85 | +@[inline] def randBoundedInt (low high : Int) (endpoint : Bool := false) : RandPCG IO Int := |
| 86 | + randInt₀ low high endpoint |
| 87 | + |
| 88 | +/-- Sample a `Float` uniformly in `[low, high)`. -/ |
| 89 | +@[inline] def uniform (low high : Float) : RandPCG IO Float := do |
| 90 | + if low > high then throw <| IO.userError "low > high" |
| 91 | + let x ← random |
| 92 | + return Float.fma x (high - low) low |
| 93 | + |
| 94 | +/-- Sample the tail of the standard normal beyond `Ziggurat.norR`, as the `idx == 0` branch of |
| 95 | +numpy's `random_standard_normal`: draw from an exponential tail until the pair of draws falls under |
| 96 | +the normal's, which is Marsaglia's method for the tail. `negate` carries the sign numpy reads off |
| 97 | +the integer already drawn, which the loop does not redraw. -/ |
| 98 | +partial def normalTail (negate : Bool) : RandPCG IO Float := do |
| 99 | + let xx := -Ziggurat.norInvR * Float.log1p (-(← random)) |
| 100 | + let yy := -Float.log1p (-(← random)) |
| 101 | + if yy + yy > xx * xx then |
| 102 | + return if negate then -(Ziggurat.norR + xx) else Ziggurat.norR + xx |
| 103 | + normalTail negate |
| 104 | + |
| 105 | +/-- Sample from the standard normal, as numpy's `random_standard_normal`. The 64-bit output gives |
| 106 | +the strip in its low byte, then the sign, then a 52-bit abscissa; the tail takes its sign from a |
| 107 | +further bit of that same abscissa, as numpy does. -/ |
| 108 | +@[inline] def standardNormal : RandPCG IO Float := |
| 109 | + ziggurat Ziggurat.ki Ziggurat.wi Ziggurat.fi |
| 110 | + (fun r => |
| 111 | + let idx := (r &&& 0xFF).toNat |
| 112 | + let r := r >>> 8 |
| 113 | + (idx, (r >>> 1) &&& 0x000FFFFFFFFFFFFF, (r &&& 1) == 1)) |
| 114 | + (fun x => Float.exp ((-0.5) * x * x)) |
| 115 | + (fun rabs => normalTail (((rabs >>> 8) &&& 1) == 1)) |
| 116 | + |
| 117 | +/-- Sample from the standard exponential, as numpy's `random_standard_exponential`. The 64-bit |
| 118 | +output is first shifted by three, then gives the strip in its low byte and a 53-bit abscissa; the |
| 119 | +tail is the exponential's own, memoryless, so one draw beyond `Ziggurat.expR` suffices. -/ |
| 120 | +@[inline] def standardExponential : RandPCG IO Float := |
| 121 | + ziggurat Ziggurat.ke Ziggurat.we Ziggurat.fe |
| 122 | + (fun r => |
| 123 | + let r := r >>> 3 |
| 124 | + ((r &&& 0xFF).toNat, r >>> 8, false)) |
| 125 | + (fun x => Float.exp (-x)) |
| 126 | + (fun _ => do return Ziggurat.expR - Float.log1p (-(← random))) |
| 127 | + |
| 128 | +/-- Draw random samples from a normal (Gaussian) distribution. -/ |
| 129 | +@[inline] def normal (loc : Float := 0) (scale : Float := 1) : RandPCG IO Float := do |
| 130 | + if scale < 0 then throw <| IO.userError "scale < 0" |
| 131 | + return Float.fma scale (← standardNormal) loc |
| 132 | + |
| 133 | +/-- Draw samples from an exponential distribution. -/ |
| 134 | +@[inline] def exponential (scale : Float := 1) : RandPCG IO Float := do |
| 135 | + if scale < 0 then throw <| IO.userError "scale < 0" |
| 136 | + return scale * (← standardExponential) |
| 137 | + |
| 138 | +end NumLean |
0 commit comments