Skip to content

Commit c6b596f

Browse files
committed
Binomial distribution
1 parent c9192fd commit c6b596f

7 files changed

Lines changed: 270 additions & 15 deletions

File tree

‎RandomDo.lean‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ 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.Binomial
910
public import RandomDo.NumLean.Distributions
1011
public import RandomDo.NumLean.PCG64
1112
public import RandomDo.NumLean.SeedSequence

‎RandomDo/NumLean/Binomial.lean‎

Lines changed: 194 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,194 @@
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+
10+
/-!
11+
# The binomial distribution
12+
13+
`binomial n p` draws exactly what numpy's `Generator.binomial` draws: the inversion of the
14+
cumulative distribution where the mean `n * p` is at most `30`, the BTPE algorithm of
15+
Kachitvichyanukul and Schmeiser beyond, and the mirror image of either when `p > 1 / 2`.
16+
17+
`binomial`, which chooses between the two and mirrors them, is in
18+
`RandomDo.NumLean.Distributions`, with the other distributions.
19+
20+
## Main definitions
21+
22+
* `Inversion`, `inversionSetup`, `inversionDraw`: numpy's `random_binomial_inversion`.
23+
* `Btpe`, `btpeSetup`, `Btpe.accept`, `btpeDraw`: numpy's `random_binomial_btpe`.
24+
25+
## References
26+
27+
* V. Kachitvichyanukul and B. W. Schmeiser, *Binomial random variate generation*, Communications
28+
of the ACM 31 (1988), 216-222.
29+
* numpy's `numpy/random/src/distributions/distributions.c`.
30+
-/
31+
32+
@[expose] public section
33+
34+
namespace NumLean
35+
36+
/-! ## Inversion -/
37+
38+
/-- The constants the inversion reads a draw against. -/
39+
structure Inversion where
40+
/-- The number of trials. -/
41+
n : Float
42+
/-- The probability of a success, at most one half. -/
43+
p : Float
44+
/-- The probability of a failure, `1 - p`. -/
45+
q : Float
46+
/-- `q ^ n`, the probability that no trial succeeds. -/
47+
qn : Float
48+
/-- The number of successes past which the walk gives up and starts over. -/
49+
bound : Float
50+
51+
/-- The constants of `random_binomial_inversion`, for `p ≤ 0.5` and `n * p ≤ 30`. -/
52+
@[inline] def inversionSetup (n p : Float) : Inversion :=
53+
let q := 1.0 - p
54+
let np := n * p
55+
let b := np + 10.0 * Float.sqrt (np * q + 1)
56+
{ n, p, q, qn := Float.exp (n * Float.log q), bound := if n < b then n else b }
57+
58+
/-- Walk up the cumulative distribution from zero until it passes `u`, as the loop of
59+
`random_binomial_inversion`. Answers `-1` where the walk runs past `bound`, which the tail beyond
60+
it is too thin to reach and where numpy starts the draw over. -/
61+
partial def inversionWalk (s : Inversion) (x px u : Float) : Float :=
62+
if u > px then
63+
let x := x + 1
64+
if x > s.bound then -1
65+
else inversionWalk s x (((s.n - x + 1) * s.p * px) / (x * s.q)) (u - px)
66+
else x
67+
68+
/-- Sample by inverting the cumulative distribution, as numpy's `random_binomial_inversion`. -/
69+
partial def inversionDraw (s : Inversion) : RandPCG IO Float := do
70+
let x := inversionWalk s 0 s.qn (← random)
71+
if x < 0 then inversionDraw s else return x
72+
73+
/-! ## BTPE -/
74+
75+
/-- The constants BTPE reads a draw against. -/
76+
structure Btpe where
77+
/-- The number of trials. -/
78+
n : Float
79+
/-- The probability of a success, at most one half. -/
80+
r : Float
81+
/-- The probability of a failure, `1 - r`. -/
82+
q : Float
83+
/-- The mode of the distribution. -/
84+
m : Float
85+
/-- The middle of the triangle, `m + 1 / 2`. -/
86+
xm : Float
87+
/-- The half-width of the triangle. -/
88+
p1 : Float
89+
/-- The left end of the parallelogram. -/
90+
xl : Float
91+
/-- The right end of the parallelogram. -/
92+
xr : Float
93+
/-- The height of the parallelogram, relative to the triangle. -/
94+
c : Float
95+
/-- The rate of the left exponential tail. -/
96+
laml : Float
97+
/-- The rate of the right exponential tail. -/
98+
lamr : Float
99+
/-- The area of the triangle and the parallelogram. -/
100+
p2 : Float
101+
/-- The area of the triangle, the parallelogram and the left tail. -/
102+
p3 : Float
103+
/-- The area of all four regions, which a draw is scaled by. -/
104+
p4 : Float
105+
/-- The variance `n * r * q`. -/
106+
nrq : Float
107+
108+
/-- The constants of `random_binomial_btpe`, for `p ≤ 0.5` and `n * p > 30`. -/
109+
@[inline] def btpeSetup (n p : Float) : Btpe :=
110+
let r := if p < 1.0 - p then p else 1.0 - p
111+
let q := 1.0 - r
112+
let fm := n * r + r
113+
let m := Float.floor fm
114+
let p1 := Float.floor (2.195 * Float.sqrt (n * r * q) - 4.6 * q) + 0.5
115+
let xm := m + 0.5
116+
let xl := xm - p1
117+
let xr := xm + p1
118+
let c := 0.134 + 20.5 / (15.3 + m)
119+
let al := (fm - xl) / (fm - xl * r)
120+
let ar := (xr - fm) / (xr * q)
121+
let laml := al * (1.0 + al / 2.0)
122+
let lamr := ar * (1.0 + ar / 2.0)
123+
let p2 := p1 * (1.0 + 2.0 * c)
124+
let p3 := p2 + c / laml
125+
{ n, r, q, m, xm, p1, xl, xr, c, laml, lamr, p2, p3, p4 := p3 + c / lamr, nrq := n * r * q }
126+
127+
/-- One term of the Stirling series bounding `log` of a factorial, as the last test of BTPE spells
128+
it out. `u2` is `u * u`. -/
129+
@[inline] def btpeStirling (u u2 : Float) : Float :=
130+
(13680.0 - (462.0 - (132.0 - (99.0 - 140.0 / u2) / u2) / u2) / u2) / u / 166320.0
131+
132+
/-- The ratios of the probabilities from the mode up to `y`, multiplied into `f` one at a time as
133+
the step 50 of `random_binomial_btpe` takes them. -/
134+
partial def btpeUp (a s f i y : Float) : Float :=
135+
if i ≤ y then btpeUp a s (f * (a / i - s)) (i + 1) y else f
136+
137+
/-- The ratios from `y` up to the mode, divided out of `f` one at a time. Dividing the running
138+
value and dividing by the product do not round alike, and BTPE reads the first. -/
139+
partial def btpeDown (a s f i m : Float) : Float :=
140+
if i ≤ m then btpeDown a s (f / (a / i - s)) (i + 1) m else f
141+
142+
/-- Whether BTPE accepts the candidate `y` drawn with `v`, as the steps 50 and 52 of
143+
`random_binomial_btpe`: by the explicit product of the ratios of the probabilities between the mode
144+
and `y` when the two are close, and by a squeeze then the Stirling bound otherwise. -/
145+
def Btpe.accept (b : Btpe) (y v : Float) : Bool := Id.run do
146+
let k := Float.abs (y - b.m)
147+
unless k > 20 && k < b.nrq / 2.0 - 1 do
148+
let s := b.r / b.q
149+
let a := s * (b.n + 1)
150+
if b.m < y then return !(v > btpeUp a s 1.0 (b.m + 1) y)
151+
if b.m > y then return !(v > btpeDown a s 1.0 (y + 1) b.m)
152+
return !(v > 1.0)
153+
let rho := (k / b.nrq) * ((k * (k / 3.0 + 0.625) + 0.16666666666666666) / b.nrq + 0.5)
154+
let t := -k * k / (2 * b.nrq)
155+
let a := Float.log v
156+
if a < t - rho then return true
157+
if a > t + rho then return false
158+
let x1 := y + 1
159+
let f1 := b.m + 1
160+
let z := b.n + 1 - b.m
161+
let w := b.n - y + 1
162+
return !(a > b.xm * Float.log (f1 / x1) + (b.n - b.m + 0.5) * Float.log (z / w)
163+
+ (y - b.m) * Float.log (w * b.r / (x1 * b.q))
164+
+ btpeStirling f1 (f1 * f1) + btpeStirling z (z * z)
165+
+ btpeStirling x1 (x1 * x1) + btpeStirling w (w * w))
166+
167+
/-- Draw a candidate from the triangle, the parallelogram or one of the two exponential tails, and
168+
start over until one is accepted, as the steps 10 to 60 of `random_binomial_btpe`. -/
169+
partial def btpeDraw (b : Btpe) : RandPCG IO Float := do
170+
let u := (← random) * b.p4
171+
let v ← random
172+
if u ≤ b.p1 then
173+
return Float.floor (b.xm - b.p1 * v + u)
174+
else if u ≤ b.p2 then
175+
let x := b.xl + (u - b.p1) / b.c
176+
let v := v * b.c + 1.0 - Float.abs (b.m - x + 0.5) / b.p1
177+
if v > 1.0 then btpeDraw b else
178+
let y := Float.floor x
179+
if b.accept y v then return y else btpeDraw b
180+
else if u ≤ b.p3 then
181+
let y := Float.floor (b.xl + Float.log v / b.laml)
182+
-- `v` can be zero, and the floor of the resulting infinity is no candidate.
183+
if y < 0 || v == 0.0 then btpeDraw b else
184+
let v := v * (u - b.p2) * b.laml
185+
if b.accept y v then return y else btpeDraw b
186+
else
187+
let y := Float.floor (b.xr - Float.log v / b.lamr)
188+
if y > b.n || v == 0.0 then btpeDraw b else
189+
let v := v * (u - b.p3) * b.lamr
190+
if b.accept y v then return y else btpeDraw b
191+
192+
end NumLean
193+
194+
end

‎RandomDo/NumLean/Distributions.lean‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@ module
88
public import RandomDo.NumLean.PCG64
99
public meta import RandomDo.NumLean.PCG64
1010
public import FFI.Float
11+
public import RandomDo.NumLean.Binomial
1112
public import RandomDo.NumLean.Ziggurat
1213
public import RandomDo.NumLean.ZigguratSampler
1314

@@ -141,4 +142,22 @@ deviation. -/
141142
if scale < 0 then throw <| IO.userError "scale < 0"
142143
return scale * (← standardExponential)
143144

145+
/-- Draw samples from a binomial distribution. -/
146+
def binomial (n : Nat) (p : Float) : RandPCG IO Nat := do
147+
-- The comparisons are the `Bool` ones: through `Decidable`, each costs more than a draw.
148+
if p.lt 0.0 || Float.lt 1.0 p || p.isNaN then
149+
throw <| IO.userError "p < 0, p > 1 or p is NaN"
150+
let n := n.toUInt64.toFloat
151+
if n == 0 || p == 0.0 then return 0
152+
if Float.le p 0.5 then
153+
if Float.le (p * n) 30.0 then return (← inversionDraw (inversionSetup n p)).toUInt64.toNat
154+
else return (← btpeDraw (btpeSetup n p)).toUInt64.toNat
155+
else
156+
let q := 1.0 - p
157+
if Float.le (q * n) 30.0 then return (n - (← inversionDraw (inversionSetup n q))).toUInt64.toNat
158+
else return (n - (← btpeDraw (btpeSetup n q))).toUInt64.toNat
159+
160+
/-- Draw samples from a Bernoulli distribution. -/
161+
@[inline] def bernoulli (p : Float) := binomial 1 p
162+
144163
end NumLean

‎RandomDo/Tactic/Computable/Counterparts.lean‎

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@ module
88
public import RandomDo.Tactic.Computable.Defs
99
public import RandomDo.NumLean.Distributions
1010
public import Mathlib.Probability.Distributions.Gaussian.Real
11+
public import Mathlib.Probability.Distributions.Bernoulli
1112

1213
/-!
1314
# Computable counterparts of the pieces an `rdo` program is made of
@@ -18,7 +19,7 @@ counterpart recorded here through `@[computable_as]`. There is one entry per pie
1819
one distribution they draw from.
1920
-/
2021

21-
public meta section
22+
@[expose] public section
2223

2324
/-! ## Types -/
2425

@@ -29,12 +30,16 @@ attribute [computable_as Float] NNReal
2930

3031
attribute [computable_as NumLean.normal'] ProbabilityTheory.gaussianReal
3132

33+
def bernoulliChoice (α : Type) [MeasurableSpace α] (x y : α) (p : Float) :
34+
NumLean.RandPCG IO α := do
35+
return if (← NumLean.bernoulli p) == 1 then x else y
36+
37+
attribute [computable_as bernoulliChoice] ProbabilityTheory.bernoulliMeasure
38+
3239
/-! ## Classical functions -/
3340

3441
attribute [computable_as Float.sqrt] Real.sqrt
3542

3643
attribute [computable_as Float.log] Real.log
3744

3845
attribute [computable_as Float.exp] Real.exp
39-
40-
end

‎RandomDo/Tactic/Computable/Deriving.lean‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -62,6 +62,10 @@ partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr :=
6262
mkAppOptM ``Bind.bind #[← computableMonad, none, none, none,
6363
← translate σ p, ← translate σ k]
6464
| MeasureTheory.Measure α _ => return mkApp (← computableMonad) (← translate σ α)
65+
/- A subtype is its carrier, and one of its values is the value it carries: the constraint and
66+
the proof of it are what a computable counterpart does not have. -/
67+
| Subtype α _ => translate σ α
68+
| Subtype.mk _ _ v _ => translate σ v
6569
| _ => match e with
6670
| .fvar x => return σ.get x
6771
| .sort .. | .lit .. => return e

‎Test/Computable.lean‎

Lines changed: 17 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,8 @@
11
module
22

33
public import Test.IsMarkov
4+
public import Test.Bind
45
import Batteries.Data.Float.Basic
5-
/- A `run_cmd` runs at elaboration time, so what it calls has to be imported as `meta` too: the
6-
sampler it draws with, and `Float.toStringFull` it prints with. -/
7-
meta import RandomDo.NumLean.Distributions
8-
meta import Batteries.Data.Float.Basic
96

107
set_option linter.style.header false
118

@@ -15,19 +12,19 @@ namespace Test.Computable
1512

1613
open Test.IsMarkov NumLean Lean.Elab.Command
1714

18-
def logComputable (prog : RandPCG IO Float) : CommandElabM Unit := do
19-
let x ← (IO.runRandPCG prog : IO Float)
20-
let y ← (IO.runRandPCGWith 42 prog : IO Float)
21-
Lean.logInfo m!"x = {x.toStringFull}"
22-
Lean.logInfo m!"y (seed 42) = {y.toStringFull}"
15+
def logComputable {α : Type} [Lean.ToMessageData α] (prog : RandPCG IO α) : CommandElabM Unit := do
16+
let x ← (IO.runRandPCG prog : IO α)
17+
let y ← (IO.runRandPCGWith 42 prog : IO α)
18+
Lean.logInfo m!"x = {x}"
19+
Lean.logInfo m!"y (seed 42) = {y}"
2320

24-
--attribute [computable] sumTwo
21+
attribute [computable] sumTwo
2522

26-
--run_cmd do logComputable sumTwoComputable
23+
run_cmd do logComputable sumTwoComputable
2724

2825
@[computable]
2926
noncomputable
30-
def test : MeasureTheory.Measure ℝ := rdo
27+
def unfoldSumTwo : MeasureTheory.Measure ℝ := rdo
3128
let y ← sumTwo
3229
let x ← ProbabilityTheory.gaussianReal 0 1
3330
return x + y
@@ -42,4 +39,12 @@ run_cmd do logComputable (branchOnComputable 20)
4239

4340
run_cmd do logComputable (branchOnComputable (-1))
4441

42+
attribute [computable] fairCoin
43+
44+
run_cmd do logComputable (fairCoinComputable)
45+
46+
attribute [computable] Bind.twoCoins
47+
48+
run_cmd do logComputable (Bind.twoCoinsComputable)
49+
4550
end Test.Computable

‎scripts/check_binomial.py‎

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,27 @@
1+
import numpy as np
2+
from common import compare
3+
4+
print("Checking binomial distribution...")
5+
6+
# One pair per branch of numpy's `random_binomial`: inversion and BTPE, each on both sides of
7+
# `p = 1/2`, on both sides of the `n * p = 30` threshold, and the three degenerate cases.
8+
PARAMS = [(10, 0.3), (60, 0.5), (100, 0.31), (1000, 0.5), (1000, 0.9), (5, 0.99),
9+
(0, 0.5), (10, 0.0), (10, 1.0)]
10+
11+
LEAN = """import RandomDo
12+
def params : List (Nat × Float) :=
13+
[%s]
14+
def main (args : List String) : IO Unit := do
15+
for s in args do
16+
IO.FS.withFile (System.FilePath.mk s!"@DIR@/pcg64-{s}.txt") .write fun h ↦
17+
IO.runRandPCGWith s.toNat! do
18+
for _ in List.range (@N@ / %d) do
19+
for (n, p) in params do
20+
h.putStrLn (toString (← NumLean.binomial n p))
21+
""" % (",\n ".join(f"({n}, {p})" for n, p in PARAMS), len(PARAMS))
22+
23+
def distrib(seed, N):
24+
g = np.random.default_rng(seed)
25+
return [g.binomial(n, p) for _ in range(N // len(PARAMS)) for n, p in PARAMS]
26+
27+
compare(LEAN, distrib)

0 commit comments

Comments
 (0)