Skip to content

Commit e42eafa

Browse files
committed
Add prototype
1 parent 3909646 commit e42eafa

1 file changed

Lines changed: 377 additions & 0 deletions

File tree

‎notes/prototypes/TraceM.lean‎

Lines changed: 377 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,377 @@
1+
import RandomDo.Monad.Instances
2+
import RandomDo.Measurable
3+
import RandomDo.Tactic.IsMarkov.Elab
4+
import RandomDo.Tactic.Computable.Polymorphic
5+
import RandomDo.ForMathlib.Probability.Kernel.Composition.MeasureComp
6+
import Mathlib.Probability.Distributions.Gaussian.Real
7+
8+
/-!
9+
# Prototype: the graded trace monad, with a proving and a running instance
10+
11+
A program is written once over a `GradedMonad m`, drawing through `HasGaussianG`. Every
12+
`let x ← …` records `x` under its own name: the grade `l` of `m l α` lists the recorded variables
13+
with their names and is inferred from the program. Two instances:
14+
15+
* `TraceM l α := Measure (Rec l × α)`, the joint law of the recorded variables and the result,
16+
on the nested product `Rec l`. That product is internal: `gen_projections p` moves to the
17+
`Fin`-indexed space Mathlib's probability library speaks, defining `p.Ω := Fin n → T`,
18+
`p.toFin : Rec l → p.Ω`, `p.P : Measure p.Ω`, and `p.x : p.Ω → T` for each recorded `x`
19+
(`fun ω ↦ ω i`). The laws are computed by `simp` with the monad laws at `Measure`.
20+
* `SamplerM l α := RandPCG IO α`, which ignores the grade and samples.
21+
22+
Not part of any library: open it in the editor, it elaborates with the project's imports.
23+
-/
24+
25+
open Lean Meta Elab Term Command Do MeasureTheory ProbabilityTheory NumLean
26+
open MeasurableSpaceBind MeasurableSpacePure LawfulMeasurableSpaceMonad
27+
28+
/-! ## Grades and trace spaces -/
29+
30+
/-- A type with a measurable structure, as a grade entry. -/
31+
structure MType where
32+
carrier : Type
33+
[inst : MeasurableSpace carrier]
34+
35+
attribute [instance] MType.inst
36+
37+
/-- Print a grade entry as its carrier. -/
38+
@[app_unexpander MType.mk] def unexpandMType : PrettyPrinter.Unexpander
39+
| `($_ $T $_) => `($T)
40+
| `($_ $T) => `($T)
41+
| _ => throw ()
42+
43+
/-- A grade: the recorded variables, in order. -/
44+
abbrev Grade := List (String × MType)
45+
46+
/-- The trace space of a grade: one factor per recorded variable. -/
47+
@[reducible] def Rec : Grade → Type
48+
| [] => PUnit
49+
| (_, T) :: l => T.carrier × Rec l
50+
51+
instance Rec.instMeasurableSpace : (l : Grade) → MeasurableSpace (Rec l)
52+
| [] => inferInstance
53+
| _ :: l => letI := Rec.instMeasurableSpace l; inferInstance
54+
55+
def Rec.append : {l₁ l₂ : Grade} → Rec l₁ → Rec l₂ → Rec (l₁ ++ l₂)
56+
| [], _, _, r₂ => r₂
57+
| _ :: _, _, (a, r₁), r₂ => (a, Rec.append r₁ r₂)
58+
59+
lemma Rec.measurable_append_uncurry : ∀ {l₁ l₂ : Grade},
60+
Measurable fun p : Rec l₁ × Rec l₂ ↦ Rec.append p.1 p.2
61+
| [], _ => measurable_snd
62+
| _ :: l₁, l₂ => by
63+
change Measurable fun p : (_ × Rec l₁) × Rec l₂ ↦ (p.1.1, Rec.append p.1.2 p.2)
64+
exact measurable_fst.fst.prodMk
65+
(Rec.measurable_append_uncurry.comp (measurable_fst.snd.prodMk measurable_snd))
66+
67+
@[fun_prop]
68+
lemma Rec.measurable_append {l₁ l₂ : Grade} {X : Type*} [MeasurableSpace X]
69+
{f : X → Rec l₁} {g : X → Rec l₂} (hf : Measurable f) (hg : Measurable g) :
70+
Measurable fun x ↦ Rec.append (f x) (g x) :=
71+
Rec.measurable_append_uncurry.comp (hf.prodMk hg)
72+
73+
/-! ## The graded monad, and the classes a program draws through -/
74+
75+
/-- A monad graded by the recorded variables. -/
76+
class GradedMonad (m : Grade → (α : Type) → [MeasurableSpace α] → Type) where
77+
gPure {α : Type} [MeasurableSpace α] : α → m [] α
78+
gBind {l₁ l₂ : Grade} {α β : Type} [MeasurableSpace α] [MeasurableSpace β] :
79+
m l₁ α → (α → m l₂ β) → m (l₁ ++ l₂) β
80+
/-- Record `a` under the name `n`: one more coordinate of the trace. Inserted by the
81+
elaborator after every `let x ← …`. -/
82+
grecord (n : String) {α : Type} [MeasurableSpace α] (a : α) : m [(n, MType.mk α)] Unit
83+
84+
export GradedMonad (gPure gBind grecord)
85+
86+
/-- A graded monad that can draw from a Gaussian with values in `R`. -/
87+
class HasGaussianG (m : Grade → (α : Type) → [MeasurableSpace α] → Type) (R : Type)
88+
[MeasurableSpace R] where
89+
gaussian : R → R → m [] R
90+
91+
/-! ### The proving instance: the joint law -/
92+
93+
/-- The trace monad: the joint law of the recorded variables and the result. -/
94+
def TraceM (l : Grade) (α : Type) [MeasurableSpace α] : Type := Measure (Rec l × α)
95+
96+
/- The operations are `rdo`-shaped programs at `Measure`, so that `is_markov` reads them. -/
97+
98+
noncomputable instance TraceM.gradedMonad : GradedMonad TraceM where
99+
gPure a := (mPure ((), a) : Measure (Rec [] × _))
100+
gBind {l₁ l₂ _ _ _ _} x f :=
101+
((x : Measure (Rec l₁ × _)) >>=ₘ fun p ↦
102+
(f p.2 : Measure (Rec l₂ × _)) >>=ₘ fun q ↦ mPure (Rec.append p.1 q.1, q.2)
103+
: Measure (Rec (l₁ ++ l₂) × _))
104+
grecord n α _ a := (mPure ((a, ()), ()) : Measure (Rec [(n, MType.mk α)] × Unit))
105+
106+
noncomputable instance TraceM.hasGaussian : HasGaussianG TraceM ℝ where
107+
gaussian μ v := (gaussianReal μ (Real.toNNReal v) >>=ₘ fun x ↦ mPure ((), x) : Measure (Rec [] × ℝ))
108+
109+
/-! ### The running instance: a sampler that ignores the grade -/
110+
111+
/-- The sampling monad, graded trivially. -/
112+
def SamplerM (_l : Grade) (α : Type) [MeasurableSpace α] : Type := RandPCG IO α
113+
114+
instance SamplerM.gradedMonad : GradedMonad SamplerM where
115+
gPure a := (pure a : RandPCG IO _)
116+
gBind x f := (bind (x : RandPCG IO _) f : RandPCG IO _)
117+
grecord _ _ _ _ := (pure () : RandPCG IO Unit)
118+
119+
instance SamplerM.hasGaussian : HasGaussianG SamplerM Float where
120+
gaussian μ v := (normal' μ v : RandPCG IO Float)
121+
122+
/-! ## The `DoOps`: `gPure`/`gBind` of whatever `m` the expected type names -/
123+
124+
def gradeElem : Expr := mkApp2 (mkConst ``Prod [.zero, .one]) (mkConst ``String) (mkConst ``MType)
125+
126+
def mkGM (m l α σ : Expr) : Expr := mkApp3 m l α σ
127+
128+
def gradedOps : DoOps := { DoOps.default with
129+
mkPureApp α e := do
130+
let m := (← read).monadInfo.m
131+
let e ← Term.ensureHasType α e
132+
let σ ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α))
133+
let inst ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``GradedMonad) m))
134+
return mkAppN (mkConst ``GradedMonad.gPure) #[m, inst, α, σ, e]
135+
mkBindApp α β e k := do
136+
let m := (← read).monadInfo.m
137+
Term.synthesizeSyntheticMVarsNoPostponing
138+
let σα ← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α)
139+
let σβ ← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) β)
140+
let eType ← instantiateMVars (← inferType e)
141+
let .app (.app (.app _ l₁) _) _ := eType.consumeMData | throwError "graded bind: {e} : {eType}"
142+
let kType ← instantiateMVars (← inferType k)
143+
let .forallE _ _ body _ := kType.consumeMData | throwError "graded bind: {k} : {kType}"
144+
let .app (.app (.app _ l₂) _) _ := body.consumeMData | throwError "graded bind: {k} : {kType}"
145+
if body.hasLooseBVars then throwError "graded bind: the grade {l₂} depends on the bound value"
146+
let e ← Term.ensureHasType (mkGM m l₁ α σα) e
147+
let k ← Term.ensureHasType (← mkArrow α (mkGM m l₂ β σβ)) k
148+
let σα ← instantiateMVars σα
149+
let σβ ← instantiateMVars σβ
150+
let inst ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``GradedMonad) m))
151+
let l ← reduce (mkApp3 (mkConst ``List.append [.one]) gradeElem l₁ l₂) (skipTypes := false)
152+
mkExpectedTypeHint (mkAppN (mkConst ``GradedMonad.gBind) #[m, inst, l₁, l₂, α, β, σα, σβ, e, k])
153+
(mkGM m l β σβ)
154+
isPureApp? e := if e.isAppOfArity ``GradedMonad.gPure 5 then some (e.getArg! 4) else none
155+
splitMonadApp? type := do
156+
let .app mα _ := type.consumeMData | return none
157+
let .app ml resultType := mα.consumeMData | return none
158+
let .app m _ := ml.consumeMData | return none
159+
unless ← isType resultType do return none
160+
return some ({ m := m, u := 0, v := 0 }, resultType)
161+
mkMonadApp α := do
162+
let m := (← read).monadInfo.m
163+
let l ← mkFreshExprMVar (mkConst ``Grade)
164+
let σ ← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α)
165+
return mkGM m l α σ }
166+
167+
/-! ### Recording: every `let x ← …` records `x`
168+
169+
The binder name is known where the bind is built, so recording is one more step in `mkBindApp`:
170+
`let x ← e; k` becomes `gBind e (fun x ↦ gBind (grecord "x" x) (fun _ ↦ k x))`, for every binder
171+
the user wrote. The elaborator's own binders (`__do_lift`, `__r`, `_`) are not recorded. -/
172+
173+
def recordingOps : DoOps := { gradedOps with
174+
mkBindApp α β e k := do
175+
let k ← instantiateMVars k
176+
let .lam x _ _ _ := k | gradedOps.mkBindApp α β e k
177+
if x.hasMacroScopes || x.isInternal || x == `_ then return ← gradedOps.mkBindApp α β e k
178+
let σα ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α))
179+
let k' ← withLocalDeclD x α fun xv ↦ do
180+
let recd := mkAppN (mkConst ``GradedMonad.grecord)
181+
#[(← read).monadInfo.m, ← mkInstMVar (mkApp (mkConst ``GradedMonad) (← read).monadInfo.m),
182+
mkStrLit x.toString, α, σα, xv]
183+
let rest ← withLocalDeclD `__r (mkConst ``Unit) fun u ↦ mkLambdaFVars #[u] (k.beta #[xv])
184+
let inner ← gradedOps.mkBindApp (mkConst ``Unit) β recd rest
185+
mkLambdaFVars #[xv] inner
186+
gradedOps.mkBindApp α β e k' }
187+
188+
syntax (name := gdoKind) "gdo" doSeq : term
189+
@[term_elab gdoKind] def elabGdo : TermElab := fun stx et? => do
190+
let `(gdo $doSeq) := stx | throwUnsupportedSyntax
191+
elabDoWith recordingOps doSeq et?
192+
193+
/-- `rdef p : m α := …` defines the program `p`, of type `m l α` for the grade `l` inferred from
194+
the body. A `def` cannot infer a hole in its header from its body, so this expands to a `def`
195+
with the ascription `(gdo … : m _ α)` in the body. -/
196+
macro "rdef " n:ident " : " m:ident α:term:max " := " body:doSeq : command =>
197+
`(def $n := (gdo $body : $m _ $α))
198+
199+
/-! ## Vectors: from the nested product to `Fin n → T` -/
200+
201+
@[fun_prop]
202+
lemma Measurable.vecCons {X α : Type*} [MeasurableSpace X] [MeasurableSpace α] {n : ℕ}
203+
{f : X → α} {g : X → Fin n → α} (hf : Measurable f) (hg : Measurable g) :
204+
Measurable fun x ↦ Matrix.vecCons (f x) (g x) :=
205+
measurable_finCons.comp (hf.prodMk hg)
206+
207+
@[fun_prop]
208+
lemma measurable_vecEmpty {X α : Type*} [MeasurableSpace X] [MeasurableSpace α] :
209+
Measurable fun _ : X ↦ (Matrix.vecEmpty : Fin 0 → α) :=
210+
measurable_const
211+
212+
/-! ## Generating the named projections and the trace measure -/
213+
214+
partial def readGrade (l : Expr) : MetaM (List (String × Expr)) := do
215+
match_expr l with
216+
| List.nil _ => return []
217+
| List.cons _ hd tl =>
218+
let_expr Prod.mk _ _ n T := hd | throwError "not a grade entry: {hd}"
219+
let .lit (.strVal s) := n | throwError "not a name literal: {n}"
220+
let_expr MType.mk T _ := T | throwError "not a measurable type: {T}"
221+
return (s, T) :: (← readGrade tl)
222+
| _ => throwError "not a literal grade: {l}"
223+
224+
/-- `gen_projections p`, for a program whose recorded variables all have the same type `T`,
225+
defines the `Fin`-indexed trace space and everything on it:
226+
227+
* `p.Ω := Fin n → T`, and `p.toFin : Rec l → p.Ω` with `p.measurable_toFin`;
228+
* `p.P : Measure p.Ω`, the joint law of the recorded variables;
229+
* `p.x : p.Ω → T`, `fun ω ↦ ω i`, with `p.measurable_x`, for each recorded `x` at position `i`. -/
230+
elab "gen_projections " n:ident : command => liftTermElabM do
231+
let c ← realizeGlobalConstNoOverload n
232+
let ty ← instantiateMVars (← getConstInfo c).type
233+
let .app (.app (.app _ l) _) _ := ty | throwError "{c} is not a graded program: {ty}"
234+
let l ← reduce l (skipTypes := false)
235+
let ΩRec := mkApp (mkConst ``Rec) l
236+
let entries ← readGrade l
237+
let some (_, T) := entries.head? | throwError "{c} records nothing"
238+
for (_, T') in entries do
239+
unless ← isDefEq T T' do throwError "heterogeneous grade {l}: not supported by this prototype"
240+
let k := entries.length
241+
let finK := mkApp (mkConst ``Fin) (mkNatLit k)
242+
let Ω ← mkArrow finK T
243+
let define (name : Name) (type value : Expr) (compile := true) : TermElabM Unit := do
244+
let decl := .defnDecl <| mkDefinitionValEx (c ++ name) [] type value .abbrev .safe []
245+
-- `P` is a measure, hence noncomputable: add it without compiling it.
246+
if compile then addAndCompile decl else addDecl decl
247+
enableRealizationsForConst (c ++ name)
248+
logInfo m!"{c ++ name} : {type}"
249+
let prove (name : Name) (type : Expr) (tac : TSyntax ``Lean.Parser.Tactic.tacticSeq) :
250+
TermElabM Unit := do
251+
let prf ← Term.elabTermAndSynthesize (← `(by $tac)) type
252+
addDecl <| .thmDecl <| mkTheoremValEx (c ++ name) [] type (← instantiateMVars prf) []
253+
define `Ω (mkSort .one) Ω
254+
-- `toFin ω = ![ω.1, ω.2.1, …]`
255+
let toFin ← withLocalDeclD `ω ΩRec fun ω ↦ do
256+
let mut coords := #[]
257+
for i in [0:k] do
258+
let mut e := ω
259+
for _ in [0:i] do e ← mkAppM ``Prod.snd #[e]
260+
coords := coords.push (← mkAppM ``Prod.fst #[e])
261+
let mut v ← mkAppOptM ``Matrix.vecEmpty #[T]
262+
for e in coords.reverse do v ← mkAppM ``Matrix.vecCons #[e, v]
263+
mkLambdaFVars #[ω] v
264+
define `toFin (← mkArrow ΩRec Ω) toFin
265+
prove `measurable_toFin (← mkAppM ``Measurable #[mkConst (c ++ `toFin)])
266+
(← `(tacticSeq| unfold $(mkIdent (c ++ `toFin)):ident; fun_prop))
267+
let P ← Term.elabTermAndSynthesize (← `(MeasureTheory.Measure.map $(mkIdent (c ++ `toFin))
268+
(MeasureTheory.Measure.map Prod.fst ($(mkIdent c) : MeasureTheory.Measure _)))) none
269+
define `P (← inferType P) (← instantiateMVars P) (compile := false)
270+
for (name, _) in entries, i in [0:k] do
271+
let idx ← Term.elabTermAndSynthesize
272+
(← `(($(Syntax.mkNumLit (toString i)) : Fin $(Syntax.mkNumLit (toString k))))) finK
273+
let proj ← withLocalDeclD `ω Ω fun ω ↦ mkLambdaFVars #[ω] (mkApp ω idx)
274+
define name.toName (← mkArrow Ω T) proj
275+
prove (Name.mkSimple ("measurable_" ++ name)) (← mkAppM ``Measurable #[mkConst (c ++ name.toName)])
276+
(← `(tacticSeq| unfold $(mkIdent (c ++ name.toName)):ident; fun_prop))
277+
278+
/-! ## The program, written once -/
279+
280+
section
281+
variable {m : Grade → (α : Type) → [MeasurableSpace α] → Type} [GradedMonad m]
282+
{R : Type} [MeasurableSpace R] [Add R] [OfNat R 0] [OfNat R 1] [HasGaussianG m R]
283+
284+
-- Draw `x ∼ 𝒩(0, 1)`, then `y ∼ 𝒩(x, 1)`, and return their sum. Both draws are recorded.
285+
rdef sum2 : m R :=
286+
let x ← HasGaussianG.gaussian (m := m) (0 : R) 1
287+
let y ← HasGaussianG.gaussian (m := m) x 1
288+
return x + y
289+
290+
end
291+
292+
#check @sum2
293+
294+
/-! ## Running it -/
295+
296+
#eval IO.runRandPCGWith 42 (sum2 (m := SamplerM) (R := Float) : RandPCG IO Float)
297+
298+
/-! ## Proving about it -/
299+
300+
/-- The program read as a joint law. -/
301+
noncomputable def sum2T := sum2 (m := TraceM) (R := ℝ)
302+
303+
#check sum2T
304+
305+
gen_projections sum2T
306+
307+
/-- `Measure.bind` and `Measure.dirac` are the monad's `mBind` and `mPure`, syntactically. -/
308+
lemma Measure.bind_eq_mBind {α β : Type} [MeasurableSpace α] [MeasurableSpace β] (μ : Measure α)
309+
(f : α → Measure β) : μ.bind f = μ >>=ₘ f := rfl
310+
311+
lemma Measure.dirac_eq_mPure {α : Type} [MeasurableSpace α] (a : α) :
312+
Measure.dirac a = (mPure a : Measure α) := rfl
313+
314+
/-- A draw that is not used afterwards integrates out. -/
315+
lemma mBind_const {α β : Type} [MeasurableSpace α] [MeasurableSpace β] (μ : Measure α)
316+
[IsProbabilityMeasure μ] (ν : Measure β) : (μ >>=ₘ fun _ ↦ ν) = ν := by
317+
change μ.bind _ = ν
318+
rw [Measure.bind_const, measure_univ, one_smul]
319+
320+
/-- The side conditions of the monad laws at `Measure`: measurability of a continuation, which is
321+
the Markov property of the program it is. -/
322+
macro "markov_side" : tactic =>
323+
`(tactic| first
324+
| fun_prop
325+
| (apply (config := { allowSynthFailures := true }) IsMarkov.measurable; is_markov))
326+
327+
/-- Unfold the program at `TraceM` down to `>>=ₘ`/`mPure`, and normalise with the monad laws. -/
328+
macro "trace_normalize" : tactic =>
329+
`(tactic| (
330+
-- `delta`, not `simp`: the unfolded grades are `[] ++ l`, equal to `l` only by unfolding
331+
-- `List.append`, which `simp`'s congruence closure does not do.
332+
delta sum2T sum2 TraceM.gradedMonad TraceM.hasGaussian
333+
dsimp only [id]
334+
simp only [Real.toNNReal_one]
335+
simp (disch := fun_prop) only [← Measure.bind_dirac_eq_map]
336+
simp only [Measure.bind_eq_mBind, Measure.dirac_eq_mPure]
337+
simp (disch := markov_side) only [mBind_assoc, mPure_mBind, Rec.append]))
338+
339+
/-- The internal joint law, on the nested product, in composition-product form. -/
340+
instance : IsMarkov fun x : ℝ ↦ gaussianReal x 1 := by is_markov
341+
342+
/-- The kernel `x ↦ 𝒩(x, 1)`. -/
343+
noncomputable def gk : Kernel ℝ ℝ := IsMarkov.toKernel fun x : ℝ ↦ gaussianReal x 1
344+
345+
instance : IsMarkovKernel gk := by unfold gk; infer_instance
346+
347+
@[simp] lemma gk_apply (x : ℝ) : gk x = gaussianReal x 1 := rfl
348+
349+
lemma sum2T.rec_eq :
350+
(sum2T : Measure _).map Prod.fst = (gaussianReal 0 1 ⊗ₘ gk).map fun p ↦ (p.1, (p.2, ())) := by
351+
rw [Measure.map_compProd_eq_bind _ _ (by fun_prop)]
352+
simp only [gk_apply]
353+
trace_normalize
354+
355+
/-- **The joint law of `(x, y)`**, on `Fin 2 → ℝ`: `x ∼ 𝒩(0, 1)`, then `y ∼ 𝒩(x, 1)`. -/
356+
theorem sum2T.P_eq : sum2T.P = (gaussianReal 0 1 ⊗ₘ gk).map fun p ↦ ![p.1, p.2] := by
357+
delta sum2T.P
358+
rw [sum2T.rec_eq, Measure.map_map sum2T.measurable_toFin (by fun_prop)]
359+
rfl
360+
361+
/-- The marginal law of `x`. -/
362+
theorem sum2T.map_x : sum2T.P.map sum2T.x = gaussianReal 0 1 := by
363+
rw [sum2T.P_eq, Measure.map_map sum2T.measurable_x (by fun_prop)]
364+
simp only [Function.comp_def, sum2T.x, Matrix.cons_val_zero]
365+
rw [Measure.map_compProd_eq_bind _ _ (by fun_prop)]
366+
simp only [gk_apply]
367+
simp (disch := fun_prop) only [← Measure.bind_dirac_eq_map]
368+
simp only [Measure.bind_eq_mBind, Measure.dirac_eq_mPure]
369+
simp only [mBind_const, mBind_mPure]
370+
371+
/-- `x` has law `𝒩(0, 1)` on the trace space. -/
372+
theorem sum2T.hasLaw_x : HasLaw sum2T.x (gaussianReal 0 1) sum2T.P :=
373+
⟨sum2T.measurable_x.aemeasurable, sum2T.map_x⟩
374+
375+
#check sum2T.P
376+
#check sum2T.hasLaw_x
377+
#check sum2T.P_eq

0 commit comments

Comments
 (0)