|
| 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