Skip to content

Commit 7d16e75

Browse files
authored
Merge pull request #7 from LeanMachineLearning/NumLean
Add Numpy's RNG
2 parents 1eee785 + ef5c373 commit 7d16e75

21 files changed

Lines changed: 1481 additions & 0 deletions

‎.gitignore‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,3 +28,5 @@
2828
*.synctex.gz
2929
*.synctex.gz(busy)
3030
*.pdfsync
31+
test_data/
32+
__pycache__/

‎FFI.lean‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,7 @@
1+
module -- shake: keep-all --deprecated_module: ignore
2+
3+
public import FFI.FMA
4+
public import FFI.Float
5+
public import FFI.Float32
6+
public import FFI.Model.Float
7+
public import FFI.Model.Float32

‎FFI/FMA.lean‎

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,57 @@
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+
prelude
9+
public import Init.Data.Float.Model.Unpacked.Round
10+
11+
-- This file is part of the logical model for floats which authors of float libraries
12+
-- need to rely on.
13+
@[expose] public section
14+
15+
namespace Float.Model.UnpackedFloat
16+
17+
/--
18+
Computes the fused multiply-add `x * y + z` of three floating point numbers and rounds the
19+
result according to the given specification. The product `x * y` is exact and never rounded
20+
on its own; only the final sum is rounded.
21+
-/
22+
def fma (spec : Format) : UnpackedFloat → UnpackedFloat → UnpackedFloat → UnpackedFloat
23+
| .notANumber, _, _ => .notANumber
24+
| _, .notANumber, _ => .notANumber
25+
| _, _, .notANumber => .notANumber
26+
| .zero _, .infinity _, _ => .notANumber
27+
| .infinity _, .zero _, _ => .notANumber
28+
| .infinity sign₁, .infinity sign₂, .infinity sign₃ =>
29+
if sign₁ * sign₂ == sign₃ then .infinity sign₃ else .notANumber
30+
| .infinity sign₁, .finite sign₂ .., .infinity sign₃ =>
31+
if sign₁ * sign₂ == sign₃ then .infinity sign₃ else .notANumber
32+
| .finite sign₁ .., .infinity sign₂, .infinity sign₃ =>
33+
if sign₁ * sign₂ == sign₃ then .infinity sign₃ else .notANumber
34+
| .infinity sign₁, .infinity sign₂, _ => .infinity (sign₁ * sign₂)
35+
| .infinity sign₁, .finite sign₂ .., _ => .infinity (sign₁ * sign₂)
36+
| .finite sign₁ .., .infinity sign₂, _ => .infinity (sign₁ * sign₂)
37+
| _, _, .infinity sign₃ => .infinity sign₃
38+
| .zero sign₁, .zero sign₂, .zero sign₃ =>
39+
if sign₁ * sign₂ == sign₃ then .zero sign₃ else .zero .positive
40+
| .zero sign₁, .finite sign₂ .., .zero sign₃ =>
41+
if sign₁ * sign₂ == sign₃ then .zero sign₃ else .zero .positive
42+
| .finite sign₁ .., .zero sign₂, .zero sign₃ =>
43+
if sign₁ * sign₂ == sign₃ then .zero sign₃ else .zero .positive
44+
| .zero _, _, z => z
45+
| _, .zero _, z => z
46+
| .finite s₁ m₁ e₁ _, .finite s₂ m₂ e₂ _, .zero _ =>
47+
roundWithAccuracy spec (s₁ * s₂) (m₁ * m₂) (e₁ + e₂) .exact
48+
| .finite s₁ m₁ e₁ _, .finite s₂ m₂ e₂ _, .finite s₃ m₃ e₃ _ =>
49+
let productMantissa := m₁ * m₂
50+
let productExponent := e₁ + e₂
51+
let smallerExponent := min productExponent e₃
52+
let (productMantissa, _) := decreaseExponent productMantissa productExponent smallerExponent
53+
let (m₃, _) := decreaseExponent m₃ e₃ smallerExponent
54+
let mantissa := (s₁ * s₂).apply productMantissa + s₃.apply m₃
55+
normalize spec mantissa smallerExponent .positive
56+
57+
end Float.Model.UnpackedFloat

‎FFI/Float.lean‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,31 @@
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+
prelude
9+
public import Init.Data.Float.Float
10+
public import FFI.Model.Float
11+
12+
@[expose] public section
13+
14+
namespace Float
15+
16+
/--
17+
Computes the fused multiply-add `x * y + z` of three floating-point numbers. This operation is
18+
performed with a single rounding, which can be more accurate than performing the multiplication and
19+
addition separately.
20+
21+
This function has a logical model in terms of `Float.Model`. It is implemented in compiled code by
22+
the C function `fma`.
23+
-/
24+
@[extern "fma"] def fma : Float → Float → Float → Float :=
25+
fun x y z => .ofModel (x.toModel.fma y.toModel z.toModel)
26+
27+
/-- `log (1 + x)`, the C99 `log1p`, accurate for small `x` where `log (1 + x)` loses the leading
28+
digits of the result to the rounding of `1 + x`. -/
29+
@[extern "log1p"] opaque log1p (x : Float) : Float
30+
31+
end Float

‎FFI/Float32.lean‎

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,29 @@
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+
prelude
9+
public import Init.Data.Float.Float32
10+
public import FFI.Model.Float32
11+
12+
@[expose] public section
13+
14+
namespace Float32
15+
16+
/--
17+
Computes the fused multiply-add `x * y + z` of three floating-point numbers. This operation is performed with a single rounding, which can be more accurate than performing the multiplication and addition separately.
18+
19+
This function has a logical model in terms of `Float32.Model`. It is implemented in compiled code
20+
by the C function `fmaf`.
21+
-/
22+
@[extern "fmaf"] def fma : Float32 → Float32 → Float32 → Float32 :=
23+
fun x y z => .ofModel (x.toModel.fma y.toModel z.toModel)
24+
25+
/-- `log (1 + x)`, the C99 `log1p`, accurate for small `x` where `log (1 + x)` loses the leading
26+
digits of the result to the rounding of `1 + x`. -/
27+
@[extern "log1pf"] opaque log1p (x : Float32) : Float32
28+
29+
end Float32

‎FFI/Model/Float.lean‎

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,24 @@
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+
prelude
9+
public import Init.Data.Float.Model.Float
10+
public import FFI.FMA
11+
12+
-- This file is part of the logical model for floats which authors of float libraries
13+
-- need to rely on.
14+
@[expose] public section
15+
16+
namespace Float.Model
17+
18+
/--
19+
Compute the fused multiply-add `a * b + c` of three `Float.Model`, with a single rounding.
20+
-/
21+
def fma (a b c : Float.Model) : Float.Model :=
22+
pack (UnpackedFloat.fma Format.binary64 a.unpack b.unpack c.unpack)
23+
24+
end Float.Model

‎FFI/Model/Float32.lean‎

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,26 @@
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+
prelude
9+
public import Init.Data.Float.Model.Float32
10+
public import FFI.FMA
11+
12+
-- This file is part of the logical model for floats which authors of float libraries
13+
-- need to rely on.
14+
@[expose] public section
15+
16+
namespace Float32.Model
17+
18+
open Float.Model (Format UnpackedFloat)
19+
20+
/--
21+
Compute the fused multiply-add `a * b + c` of three `Float32.Model`, with a single rounding.
22+
-/
23+
def fma (a b c : Float32.Model) : Float32.Model :=
24+
pack (UnpackedFloat.fma Format.binary32 a.unpack b.unpack c.unpack)
25+
26+
end Float32.Model

‎RandomDo.lean‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,11 @@ public import RandomDo.Monad.ForInInstances
66
public import RandomDo.Monad.Instances
77
public import RandomDo.Monad.MeasurableSpace
88
public import RandomDo.Monad.Notation
9+
public import RandomDo.NumLean.Distributions
10+
public import RandomDo.NumLean.PCG64
11+
public import RandomDo.NumLean.SeedSequence
12+
public import RandomDo.NumLean.Ziggurat
13+
public import RandomDo.NumLean.ZigguratSampler
914
public import RandomDo.Tactic.Deriving
1015
public import RandomDo.Tactic.Elab
1116
public import RandomDo.Tactic.ForInStep
Lines changed: 138 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,138 @@
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

Comments
 (0)