Skip to content

Commit 451a720

Browse files
committed
Support for bernoulli
1 parent 14c26e6 commit 451a720

3 files changed

Lines changed: 93 additions & 6 deletions

File tree

‎RandomDo/Tactic/Computable/Polymorphic.lean‎

Lines changed: 61 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,11 +6,19 @@ Authors: Gaëtan Serré
66
module
77

88
public import Mathlib.Probability.Distributions.Gaussian.Real
9+
public import Mathlib.Probability.Distributions.Bernoulli
10+
public import Mathlib.MeasureTheory.Function.SpecialFunctions.Basic
911
public import RandomDo.Monad.Instances
1012
public import RandomDo.NumLean.Distributions
1113

1214
/-!
15+
# Polymorphic `rdo` programs
1316
17+
A program written over an arbitrary `MeasurableSpaceMonad` `m`, drawing through the classes of this
18+
file, is read at `m := Measure` to prove things about it and run at `m := RandM` to sample from it.
19+
Each class has an instance of each kind: the distribution of Mathlib on `ℝ`, and the sampler of
20+
`NumLean` on `Float`. The scalar classes `HasExp`, `HasLog` and `HasSqrt` do the same for the
21+
functions a program computes with.
1422
-/
1523

1624
@[expose] public section
@@ -35,3 +43,56 @@ instance instMeasurableSpaceFloat : MeasurableSpace Float := ⊤
3543

3644
instance : HasGaussian RandM Float Float Float where
3745
gaussian μ v := normal' μ v
46+
47+
/-- A typeclass for monads that can draw from a Bernoulli distribution. -/
48+
class HasBernoulli (m : (α : Type) → [MeasurableSpace α] → Type v) (P : Type) where
49+
/-- Draw `true` with probability `p`, and `false` otherwise. -/
50+
bernoulli : P → m Bool
51+
52+
/-- A probability outside `[0, 1]` is clamped to it. -/
53+
noncomputable instance : HasBernoulli Measure ℝ where
54+
bernoulli p := bernoulliMeasure true false (Set.projIcc 0 1 zero_le_one p)
55+
56+
/-- A probability outside `[0, 1]` is clamped to it, as in the instance on `Measure`. -/
57+
instance : HasBernoulli RandM Float where
58+
bernoulli p := show RandPCG IO Bool from return (← bernoulli (max 0 (min 1 p))) == 1
59+
60+
/-- A typeclass for scalars with an exponential. -/
61+
class HasExp (R : Type) where
62+
/-- The exponential. -/
63+
exp : R → R
64+
65+
noncomputable instance : HasExp ℝ := ⟨Real.exp⟩
66+
67+
@[fun_prop]
68+
lemma HasExp.measurable_exp_real : Measurable (HasExp.exp : ℝ → ℝ) := Real.measurable_exp
69+
70+
instance : HasExp Float := ⟨Float.exp⟩
71+
72+
/-- A typeclass for scalars with a logarithm. -/
73+
class HasLog (R : Type) where
74+
/-- The logarithm. -/
75+
log : R → R
76+
77+
noncomputable instance : HasLog ℝ := ⟨Real.log⟩
78+
79+
@[fun_prop]
80+
lemma HasLog.measurable_log_real : Measurable (HasLog.log : ℝ → ℝ) := Real.measurable_log
81+
82+
instance : HasLog Float := ⟨Float.log⟩
83+
84+
/-- A typeclass for scalars with a square root. -/
85+
class HasSqrt (R : Type) where
86+
/-- The square root. -/
87+
sqrt : R → R
88+
89+
noncomputable instance : HasSqrt ℝ := ⟨Real.sqrt⟩
90+
91+
@[fun_prop]
92+
lemma HasSqrt.measurable_sqrt_real : Measurable (HasSqrt.sqrt : ℝ → ℝ) :=
93+
Real.continuous_sqrt.measurable
94+
95+
instance : HasSqrt Float := ⟨Float.sqrt⟩
96+
97+
/-- A natural number as a `Float`, so that programs polymorphic in the scalars can cast counts. -/
98+
instance instNatCastFloat : NatCast Float := ⟨Nat.toFloat⟩

‎RandomDo/Tactic/IsMarkov/Elab.lean‎

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -136,6 +136,9 @@ def closeLeaf (g : MVarId) : MetaM (List MVarId) := do
136136
if let some gs ← observing? (g.applyConst ``IsMarkov.gaussianReal) then
137137
trace[is_markov] "`gaussianReal` leaf: handing back the measurability of its parameters"
138138
return gs
139+
if let some gs ← observing? (g.applyConst ``IsMarkov.bernoulliMeasure) then
140+
trace[is_markov] "`bernoulliMeasure` leaf: handing back the measurability of its parameters"
141+
return gs
139142
return [g]
140143

141144
/-- The constant heading the body of `κ`, when it is a definition the tactic could look through. -/
@@ -172,7 +175,9 @@ def abstractLoopVars (vars : Array FVarId) (g : MVarId) : MetaM MVarId := do
172175

173176
/-- Turn a goal `IsMarkov κ` into the list of goals the user is left with. -/
174177
partial def isMarkovCore (g : MVarId) (fuel : Nat) : MetaM (List MVarId) := g.withContext do
175-
let target ← instantiateMVars (← g.getType)
178+
/- The annotations a goal may carry, e.g. the one a tactic `have` leaves on the goal after it,
179+
would hide the head of the statement. -/
180+
let target := (← instantiateMVars (← g.getType)).cleanupAnnotations
176181
-- `IsMarkov` takes five arguments: `γ`, `α`, their `MeasurableSpace` instances, and `κ`.
177182
unless target.isAppOfArity ``IsMarkov 5 do
178183
trace[is_markov] "not an `IsMarkov` goal, handed back: {target}"
@@ -251,7 +256,8 @@ partial def isMarkovCore (g : MVarId) (fuel : Nat) : MetaM (List MVarId) := g.wi
251256
let mut goals := []
252257
let mut side := []
253258
for g' in gs do
254-
if (← instantiateMVars (← g'.getType)).isAppOfArity ``IsMarkov 5 then
259+
let t := (← instantiateMVars (← g'.getType)).cleanupAnnotations
260+
if t.isAppOfArity ``IsMarkov 5 then
255261
goals := goals ++ (← isMarkovCore g' fuel)
256262
else
257263
side := side ++ [g']
@@ -280,7 +286,7 @@ partial def isMarkovCore (g : MVarId) (fuel : Nat) : MetaM (List MVarId) := g.wi
280286
if ← g.isAssigned then return leftover
281287
/- The goal was not closed, so we try to unfold names in the head of the program until we
282288
reach a known shape. If that fails, we leave the goal to the user. -/
283-
match ← unfoldToKnownShape (← instantiateMVars (← g.getType)) fuel with
289+
match ← unfoldToKnownShape (← instantiateMVars (← g.getType)).cleanupAnnotations fuel with
284290
| some target =>
285291
trace[is_markov] "unfolded the head definition to: {target.appArg!}"
286292
isMarkovCore (← g.change target) (fuel - 1)
@@ -316,7 +322,7 @@ lemma _root_.isProbabilityMeasure_of_isMarkov {α : Type*} [MeasurableSpace α]
316322
/-- Bring a goal of the form `IsProbabilityMeasure μ` into the form `IsMarkov fun _ : Unit ↦ μ`, so
317323
that `isMarkovCore` can be applied. -/
318324
def toIsMarkovGoal (g : MVarId) : MetaM MVarId := do
319-
let target ← instantiateMVars (← g.getType)
325+
let target := (← instantiateMVars (← g.getType)).cleanupAnnotations
320326
unless target.isAppOfArity ``IsProbabilityMeasure 3 do
321327
return g
322328
match ← g.applyConst ``isProbabilityMeasure_of_isMarkov with

‎RandomDo/Tactic/IsMarkov/Lemmas.lean‎

Lines changed: 22 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ public import RandomDo.Tactic.IsMarkov.ForInStep
1212
public import Mathlib.MeasureTheory.Measure.ProbabilityMeasure
1313
public import Mathlib.Data.List.OfFn
1414
public import Mathlib.Probability.Distributions.Gaussian.Real
15+
public import Mathlib.Probability.Distributions.Bernoulli
1516

1617
/-!
1718
# Markov property of `rdo` programs
@@ -31,6 +32,8 @@ complex program to the Markov property/measurability of its underlying mathemati
3132
the bound variable is Markovian in the parameter.
3233
* `gaussianReal`: A Gaussian distribution whose mean and variance depend measurably on the parameter
3334
is Markovian in the parameter.
35+
* `bernoulliMeasure`: A Bernoulli distribution whose two outcomes and probability depend measurably
36+
on the parameter is Markovian in the parameter.
3437
* `comp`: Composing a Markov kernel `κ` with a measurable function `g` is Markovian in the
3538
parameter.
3639
* `ite`: A conditional `rdo` program that chooses between two Markov kernels `κ` and `η` based on a
@@ -43,6 +46,7 @@ complex program to the Markov property/measurability of its underlying mathemati
4346
* `forInList_comp`, `forInArray_comp`, `forInVector_comp`: The same three, for a loop over a
4447
collection the program takes as an argument. The body is then asked to be Markovian jointly in the
4548
parameter and in the element, which the fixed collections do not need.
49+
* `forIn_nil`, `forIn_cons`: A `for` loop over a list, unrolled one element at a time.
4650
* `breakRunK`: The case analysis a program performs after a loop that returns early, on the `Option`
4751
slot holding the returned value, is Markovian as soon as both of its branches are.
4852
-/
@@ -51,6 +55,7 @@ complex program to the Markov property/measurability of its underlying mathemati
5155

5256
open MeasureTheory ProbabilityTheory Function
5357
open MeasurableSpacePure
58+
open scoped ENNReal
5459

5560
namespace IsMarkov
5661

@@ -91,6 +96,18 @@ lemma gaussianReal {m : γ → ℝ} {v : γ → NNReal} (hm : Measurable m) (hv
9196
IsMarkov fun c ↦ ProbabilityTheory.gaussianReal (m c) (v c) :=
9297
⟨ProbabilityTheory.measurable_gaussianReal.comp (hm.prodMk hv), fun _ ↦ inferInstance⟩
9398

99+
lemma bernoulliMeasure {x y : γ → α} {p : γ → unitInterval} (hx : Measurable x)
100+
(hy : Measurable y) (hp : Measurable p) :
101+
IsMarkov fun c ↦ ProbabilityTheory.bernoulliMeasure (x c) (y c) (p c) := by
102+
refine ⟨Measure.measurable_of_measurable_coe _ fun s hs ↦ ?_, fun _ ↦ inferInstance⟩
103+
simp only [bernoulliMeasure_def, Measure.add_apply, Measure.smul_apply,
104+
Measure.dirac_apply' _ hs, ENNReal.smul_def, smul_eq_mul]
105+
have hp' : Measurable fun c ↦ ((unitInterval.toNNReal (p c) : ℝ≥0∞)) := by fun_prop
106+
have hq' : Measurable fun c ↦ ((unitInterval.toNNReal (unitInterval.symm (p c)) : ℝ≥0∞)) := by
107+
fun_prop
108+
exact (hp'.mul ((measurable_one.indicator hs).comp hx)).add
109+
(hq'.mul ((measurable_one.indicator hs).comp hy))
110+
94111
lemma comp {κ : γ → Measure α} (hκ : IsMarkov κ) {g : σ → γ} (hg : Measurable g) :
95112
IsMarkov fun c ↦ κ (g c) := ⟨hκ.measurable.comp hg, fun _ ↦ hκ.isProbabilityMeasure _⟩
96113

@@ -228,11 +245,14 @@ private lemma forIn_eq_listLoop (l : List ι) (b : σ) (g : ι → σ → Measur
228245
MeasurableSpaceForIn.forIn (m := Measure) l b g = listLoop g l b :=
229246
loop_eq_listLoop g l b l _ (fun _ _ _ ↦ rfl) ⟨[], rfl⟩
230247

231-
private lemma forIn_nil (b : σ) (g : ι → σ → Measure (ForInStep σ)) :
248+
/-- A `for` loop over the empty list returns its initial state. -/
249+
lemma forIn_nil (b : σ) (g : ι → σ → Measure (ForInStep σ)) :
232250
MeasurableSpaceForIn.forIn (m := Measure) ([] : List ι) b g = mPure b :=
233251
forIn_eq_listLoop _ _ _
234252

235-
private lemma forIn_cons (a : ι) (l : List ι) (b : σ) (g : ι → σ → Measure (ForInStep σ)) :
253+
/-- A `for` loop over `a :: l` runs its body on `a`, then stops or carries on with the loop over
254+
`l`. -/
255+
lemma forIn_cons (a : ι) (l : List ι) (b : σ) (g : ι → σ → Measure (ForInStep σ)) :
236256
MeasurableSpaceForIn.forIn (m := Measure) (a :: l) b g
237257
= g a b >>=ₘ fun step ↦ ForInStep.casesOn (motive := fun _ ↦ Measure σ) step mPure
238258
fun b' ↦ MeasurableSpaceForIn.forIn (m := Measure) l b' g := by

0 commit comments

Comments
 (0)