|
| 1 | +/- |
| 2 | +Copyright (c) 2026 Rémy Degenne. All rights reserved. |
| 3 | +Released under Apache 2.0 license as described in the file LICENSE. |
| 4 | +Authors: Rémy Degenne |
| 5 | +-/ |
| 6 | +module |
| 7 | + |
| 8 | +public import RandomDo.Probability.Tactic |
| 9 | +public import RandomDo.Tactic.Elab |
| 10 | + |
| 11 | +set_option linter.style.header false |
| 12 | + |
| 13 | +/-! |
| 14 | +# Reading probabilistic statements off an `rdo` program |
| 15 | +
|
| 16 | +Worked examples of `RDo.HasTrace`. Nothing here is specific to a particular distribution: the |
| 17 | +programs draw from arbitrary Markov kernels, which is exactly what an `rdo` program does once |
| 18 | +`is_markov` has run on its leaves. |
| 19 | +
|
| 20 | +The pattern is always the same. |
| 21 | +
|
| 22 | +1. Build the trace of the program bottom-up with the combinators of `RandomDo.Probability.Trace`, |
| 23 | + one per `rdo` construct. The trace kernel comes out as a right-nested `⊗ₖ`, one factor per `←`. |
| 24 | +2. Fix the parameters `c`. The program's probability space is then `(Ω, P c)`, and the draws are |
| 25 | + the coordinates of `Ω`. |
| 26 | +3. Peel the `⊗ₖ` factors off one at a time with `HasLaw.compProd_snd` / `Kernel.sectR_compProd` / |
| 27 | + `HasCondDistrib.compProd_fst` / `HasCondDistrib.compProd_snd`. Step `k` of the peeling hands |
| 28 | + you the conditional distribution of the `k`-th draw given the `k-1` draws before it. |
| 29 | +
|
| 30 | +Steps 1 and 3 are entirely mechanical, which is the point: `rdo_trace` and `rdo_peel` do them. |
| 31 | +The `Automation` section at the end of this file re-derives, in two tactic calls, everything the |
| 32 | +first two sections prove by hand. |
| 33 | +-/ |
| 34 | + |
| 35 | +@[expose] public section |
| 36 | + |
| 37 | +open MeasureTheory ProbabilityTheory RDo |
| 38 | +open MeasurableSpacePure MeasurableSpaceBind |
| 39 | + |
| 40 | +noncomputable section |
| 41 | + |
| 42 | +/-! ## Two independent draws |
| 43 | +
|
| 44 | +``` |
| 45 | +rdo |
| 46 | + let x ← μ |
| 47 | + let y ← μ |
| 48 | + return x + y |
| 49 | +``` |
| 50 | +-/ |
| 51 | + |
| 52 | +section Independent |
| 53 | + |
| 54 | +variable (μ : Measure ℝ) [IsProbabilityMeasure μ] |
| 55 | + |
| 56 | +/-- A program that draws twice from the same distribution. It takes no parameter, so its trace is |
| 57 | +a kernel from `Unit`. -/ |
| 58 | +def sum2 : Measure ℝ := rdo |
| 59 | + let x ← μ |
| 60 | + let y ← μ |
| 61 | + return x + y |
| 62 | + |
| 63 | +/-- The trace space of `sum2` is `ℝ × ℝ`, one coordinate per `←`, and the trace kernel is the |
| 64 | +composition-product of the two draws. The second factor is a `Kernel.prodMkRight`, which records |
| 65 | +in the *type* that the second draw does not look at the first. -/ |
| 66 | +lemma hasTrace_sum2 : |
| 67 | + HasTrace (fun _ : Unit ↦ sum2 μ) |
| 68 | + (Kernel.const Unit μ ⊗ₖ Kernel.prodMkRight ℝ (Kernel.const Unit μ)) |
| 69 | + (fun p : Unit × (ℝ × ℝ) ↦ p.2.1 + p.2.2) := by |
| 70 | + have hfirst : HasTrace (fun _ : Unit ↦ μ) (Kernel.const Unit μ) Prod.snd := |
| 71 | + (HasTrace.sample (Kernel.const Unit μ)).congr fun _ ↦ rfl |
| 72 | + have htail : HasTrace (fun p : Unit × ℝ ↦ μ >>=ₘ fun y ↦ mPure (p.2 + y)) |
| 73 | + (Kernel.prodMkRight ℝ (Kernel.const Unit μ)) |
| 74 | + (fun q : (Unit × ℝ) × ℝ ↦ q.1.2 + q.2) := |
| 75 | + (hfirst.prodMkRight ℝ).bindPure (f := fun q : (Unit × ℝ) × ℝ ↦ q.1.2 + q.2) (by fun_prop) |
| 76 | + have hcont : Measurable fun p : Unit × ℝ ↦ μ >>=ₘ fun y ↦ mPure (p.2 + y) := by |
| 77 | + have : IsMarkov fun p : Unit × ℝ ↦ μ >>=ₘ fun y ↦ mPure (p.2 + y) := by is_markov |
| 78 | + exact this.measurable |
| 79 | + exact (hfirst.bind hcont htail).congr fun _ ↦ rfl |
| 80 | + |
| 81 | +/-- The joint law of the two draws. -/ |
| 82 | +local notation "P₂" => (Kernel.const Unit μ ⊗ₖ Kernel.prodMkRight ℝ (Kernel.const Unit μ)) () |
| 83 | + |
| 84 | +/-- The first draw has law `μ`. -/ |
| 85 | +example : HasLaw (Prod.fst : ℝ × ℝ → ℝ) μ P₂ := |
| 86 | + hasLaw_fst_compProd (Kernel.const Unit μ) (Kernel.prodMkRight ℝ (Kernel.const Unit μ)) () |
| 87 | + |
| 88 | +/-- The second draw has law `μ` too. -/ |
| 89 | +example : HasLaw (Prod.snd : ℝ × ℝ → ℝ) μ P₂ := |
| 90 | + hasLaw_snd_compProd_prodMkRight (P := Kernel.const Unit μ) (R := Kernel.const Unit μ) () |
| 91 | + |
| 92 | +/-- And the two are independent: this is read off the `Kernel.prodMkRight` in the trace kernel, |
| 93 | +which is itself read off the fact that the second `←` does not mention `x`. -/ |
| 94 | +example : IndepFun (Prod.fst : ℝ × ℝ → ℝ) Prod.snd P₂ := |
| 95 | + indepFun_snd_compProd_prodMkRight (Kernel.const Unit μ) () |
| 96 | + |
| 97 | +/-- The program returns the sum of the two draws. Together with the three statements above, this |
| 98 | +says exactly: `sum2 μ` is the law of `X + Y` for `X`, `Y` independent with law `μ`. -/ |
| 99 | +example : HasLaw (fun ω : ℝ × ℝ ↦ ω.1 + ω.2) (sum2 μ) P₂ := (hasTrace_sum2 μ).hasLaw_out () |
| 100 | + |
| 101 | +end Independent |
| 102 | + |
| 103 | +/-! ## A dependent chain |
| 104 | +
|
| 105 | +``` |
| 106 | +rdo |
| 107 | + let x ← κ c |
| 108 | + let y ← η (c, x) |
| 109 | + let z ← θ ((c, x), y) |
| 110 | + return x + y + z |
| 111 | +``` |
| 112 | +
|
| 113 | +Every draw may read the parameter and all the draws before it. This is the general shape of a |
| 114 | +straight-line `rdo` program: `κ`, `η`, `θ` stand for whatever `is_markov` produced at each `←`. |
| 115 | +-/ |
| 116 | + |
| 117 | +section Chain |
| 118 | + |
| 119 | +variable (κ : Kernel ℝ ℝ) [IsMarkovKernel κ] (η : Kernel (ℝ × ℝ) ℝ) [IsMarkovKernel η] |
| 120 | + (θ : Kernel ((ℝ × ℝ) × ℝ) ℝ) [IsMarkovKernel θ] |
| 121 | + |
| 122 | +/-- Three draws, each depending on everything before it. -/ |
| 123 | +def chain (c : ℝ) : Measure ℝ := rdo |
| 124 | + let x ← κ c |
| 125 | + let y ← η (c, x) |
| 126 | + let z ← θ ((c, x), y) |
| 127 | + return x + y + z |
| 128 | + |
| 129 | +lemma hasTrace_chain : |
| 130 | + HasTrace (chain κ η θ) (κ ⊗ₖ (η ⊗ₖ θ)) |
| 131 | + (fun p : ℝ × (ℝ × (ℝ × ℝ)) ↦ p.2.1 + p.2.2.1 + p.2.2.2) := by |
| 132 | + have h3 : HasTrace (fun q : (ℝ × ℝ) × ℝ ↦ θ q >>=ₘ fun z ↦ mPure (q.1.2 + q.2 + z)) θ |
| 133 | + (fun r : ((ℝ × ℝ) × ℝ) × ℝ ↦ r.1.1.2 + r.1.2 + r.2) := |
| 134 | + (HasTrace.sample θ).bindPure (f := fun r : ((ℝ × ℝ) × ℝ) × ℝ ↦ r.1.1.2 + r.1.2 + r.2) |
| 135 | + (by fun_prop) |
| 136 | + have hcont2 : Measurable fun q : (ℝ × ℝ) × ℝ ↦ θ q >>=ₘ fun z ↦ mPure (q.1.2 + q.2 + z) := by |
| 137 | + have : IsMarkov fun q : (ℝ × ℝ) × ℝ ↦ θ q >>=ₘ fun z ↦ mPure (q.1.2 + q.2 + z) := by |
| 138 | + is_markov |
| 139 | + exact this.measurable |
| 140 | + have h2 := (HasTrace.sample η).bind hcont2 h3 |
| 141 | + have hcont1 : Measurable fun p : ℝ × ℝ ↦ |
| 142 | + η p >>=ₘ fun y ↦ θ (p, y) >>=ₘ fun z ↦ mPure (p.2 + y + z) := by |
| 143 | + have : IsMarkov fun p : ℝ × ℝ ↦ |
| 144 | + η p >>=ₘ fun y ↦ θ (p, y) >>=ₘ fun z ↦ mPure (p.2 + y + z) := by is_markov |
| 145 | + exact this.measurable |
| 146 | + exact ((HasTrace.sample κ).bind hcont1 h2).congr fun _ ↦ rfl |
| 147 | + |
| 148 | +variable (c : ℝ) |
| 149 | + |
| 150 | +/-- The joint law of the three draws, given the parameter `c`. -/ |
| 151 | +local notation "P₃" => (κ ⊗ₖ (η ⊗ₖ θ)) c |
| 152 | + |
| 153 | +/-- **The mechanical peeling.** One `compProd_*` step per `←` in the program: `X` has law `κ c`, |
| 154 | +`Y` given `X` has law `η (c, X)`, and `Z` given `(X, Y)` has law `θ ((c, X), Y)`. The kernels on |
| 155 | +the right-hand sides are literally the ones written in the program. -/ |
| 156 | +example : |
| 157 | + HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1) (κ c) P₃ |
| 158 | + ∧ HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.1) (fun ω ↦ ω.1) (Kernel.sectR η c) P₃ |
| 159 | + ∧ HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.2) (fun ω ↦ (ω.1, ω.2.1)) |
| 160 | + (θ.comap (fun p : ℝ × ℝ ↦ ((c, p.1), p.2)) (by fun_prop)) P₃ := by |
| 161 | + have h0 := hasLaw_id_compProd κ (η ⊗ₖ θ) c |
| 162 | + have htail := h0.compProd_snd |
| 163 | + rw [Kernel.sectR_compProd] at htail |
| 164 | + exact ⟨h0.compProd_fst, htail.compProd_fst, htail.compProd_snd⟩ |
| 165 | + |
| 166 | +/-- Unfolding what those kernels are: `Kernel.sectR η c x = η (c, x)`, so the middle statement |
| 167 | +above really is "given `X = x`, the second draw is distributed as `η (c, x)`". -/ |
| 168 | +example (x : ℝ) : Kernel.sectR η c x = η (c, x) := Kernel.sectR_apply η x c |
| 169 | + |
| 170 | +/-- The program's result, as a random variable on the trace space. -/ |
| 171 | +example : HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1 + ω.2.1 + ω.2.2) (chain κ η θ c) P₃ := |
| 172 | + (hasTrace_chain κ η θ).hasLaw_out c |
| 173 | + |
| 174 | +end Chain |
| 175 | + |
| 176 | +/-! ## A loop, at coarse granularity |
| 177 | +
|
| 178 | +A `for` loop is a Markov kernel, so `HasTrace.of_isMarkov` (here in its `HasTrace.sample` form) |
| 179 | +makes the whole loop *one* trace coordinate holding its result. That is enough to state the |
| 180 | +conditional law of everything that comes after the loop given what the loop produced; it does not |
| 181 | +decompose the loop's own iterations, which needs the extra machinery described in |
| 182 | +`notes/TRACE_SEMANTICS.md`. |
| 183 | +-/ |
| 184 | + |
| 185 | +section Loop |
| 186 | + |
| 187 | +variable (μ : Measure ℝ) [IsProbabilityMeasure μ] (η : Kernel ℝ ℝ) [IsMarkovKernel η] |
| 188 | + |
| 189 | +/-- A loop summing `l.length` independent draws. -/ |
| 190 | +def loopPart (l : List ℕ) : Measure ℝ := rdo |
| 191 | + let mut S := 0 |
| 192 | + for _ in l rdo |
| 193 | + let x ← μ |
| 194 | + S := S + x |
| 195 | + return S |
| 196 | + |
| 197 | +instance (l : List ℕ) : IsProbabilityMeasure (loopPart μ l) := by |
| 198 | + unfold loopPart |
| 199 | + is_markov |
| 200 | + |
| 201 | +/-- The loop, then a draw whose distribution depends on what the loop returned. -/ |
| 202 | +def loopThen (l : List ℕ) : Measure ℝ := rdo |
| 203 | + let S ← loopPart μ l |
| 204 | + let y ← η S |
| 205 | + return y |
| 206 | + |
| 207 | +/-- The second `←` reads only the value the loop returned, i.e. the trace so far. -/ |
| 208 | +def afterLoopK : Kernel (Unit × ℝ) ℝ := η.comap Prod.snd measurable_snd |
| 209 | + |
| 210 | +instance : IsMarkovKernel (afterLoopK η) := by |
| 211 | + unfold afterLoopK |
| 212 | + infer_instance |
| 213 | + |
| 214 | +variable (l : List ℕ) |
| 215 | + |
| 216 | +lemma hasTrace_loopThen : |
| 217 | + HasTrace (fun _ : Unit ↦ loopThen μ η l) |
| 218 | + (Kernel.const Unit (loopPart μ l) ⊗ₖ afterLoopK η) |
| 219 | + (fun p : Unit × (ℝ × ℝ) ↦ p.2.2) := by |
| 220 | + have hloop : HasTrace (fun _ : Unit ↦ loopPart μ l) |
| 221 | + (Kernel.const Unit (loopPart μ l)) Prod.snd := |
| 222 | + (HasTrace.sample (Kernel.const Unit (loopPart μ l))).congr fun _ ↦ rfl |
| 223 | + have htail : HasTrace (fun p : Unit × ℝ ↦ η p.2) (afterLoopK η) |
| 224 | + (Prod.snd : (Unit × ℝ) × ℝ → ℝ) := |
| 225 | + (HasTrace.sample (afterLoopK η)).congr fun _ ↦ rfl |
| 226 | + have hcont : Measurable fun p : Unit × ℝ ↦ η p.2 := |
| 227 | + (Kernel.measurable η).comp measurable_snd |
| 228 | + exact (hloop.bind hcont htail).congr fun _ ↦ by unfold loopThen; rfl |
| 229 | + |
| 230 | +/-- The loop's result has the loop's law. -/ |
| 231 | +example : HasLaw (Prod.fst : ℝ × ℝ → ℝ) (loopPart μ l) |
| 232 | + ((Kernel.const Unit (loopPart μ l) ⊗ₖ afterLoopK η) ()) := |
| 233 | + hasLaw_fst_compProd (Kernel.const Unit (loopPart μ l)) (afterLoopK η) () |
| 234 | + |
| 235 | +/-- And given it, the draw that follows the loop has law `η S`. -/ |
| 236 | +example : HasCondDistrib (Prod.snd : ℝ × ℝ → ℝ) Prod.fst (Kernel.sectR (afterLoopK η) ()) |
| 237 | + ((Kernel.const Unit (loopPart μ l) ⊗ₖ afterLoopK η) ()) := |
| 238 | + hasCondDistrib_snd_compProd _ _ () |
| 239 | + |
| 240 | +example (S : ℝ) : Kernel.sectR (afterLoopK η) () S = η S := rfl |
| 241 | + |
| 242 | +end Loop |
| 243 | + |
| 244 | +/-! ## The same, automatically |
| 245 | +
|
| 246 | +`rdo_trace` walks the program and builds the trace; `rdo_peel` reads the laws off it. The |
| 247 | +hypotheses they leave are exactly the statements proved by hand above. |
| 248 | +-/ |
| 249 | + |
| 250 | +section Automation |
| 251 | + |
| 252 | +variable (μ : Measure ℝ) [IsProbabilityMeasure μ] (κ : Kernel ℝ ℝ) [IsMarkovKernel κ] |
| 253 | + (η : Kernel (ℝ × ℝ) ℝ) [IsMarkovKernel η] (θ : Kernel ((ℝ × ℝ) × ℝ) ℝ) [IsMarkovKernel θ] |
| 254 | + |
| 255 | +/-- Two independent draws: the tactics find the trace, both laws, the independence, and the law of |
| 256 | +the result. -/ |
| 257 | +example : True := by |
| 258 | + rdo_trace (sum2 μ) with h |
| 259 | + rdo_peel h () with hX hY hY' hindep hout |
| 260 | + -- `h : HasTrace (fun _ ↦ sum2 μ) |
| 261 | + -- (Kernel.const Unit μ ⊗ₖ Kernel.prodMkRight ℝ (Kernel.const Unit μ)) |
| 262 | + -- (fun p ↦ p.2.1 + p.2.2)` |
| 263 | + have _ : HasLaw (Prod.fst : ℝ × ℝ → ℝ) μ _ := hX |
| 264 | + have _ : HasCondDistrib (Prod.snd : ℝ × ℝ → ℝ) Prod.fst _ _ := hY |
| 265 | + have _ : HasLaw (Prod.snd : ℝ × ℝ → ℝ) μ _ := hY' |
| 266 | + have _ : IndepFun (Prod.fst : ℝ × ℝ → ℝ) Prod.snd _ := hindep |
| 267 | + have _ : HasLaw (fun ω : ℝ × ℝ ↦ ω.1 + ω.2) (sum2 μ) _ := hout |
| 268 | + trivial |
| 269 | + |
| 270 | +/-- The dependent chain: one conditional law per `←`, with the kernel written at that `←`. -/ |
| 271 | +example (c : ℝ) : True := by |
| 272 | + rdo_trace (chain κ η θ) with h |
| 273 | + rdo_peel h c with hX hY hZ hout |
| 274 | + have _ : HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1) (κ c) _ := hX |
| 275 | + have _ : HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.1) (fun ω ↦ ω.1) |
| 276 | + (η.comap (fun ω ↦ (c, ω)) (by fun_prop)) _ := hY |
| 277 | + have _ : HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.2) (fun ω ↦ (ω.1, ω.2.1)) |
| 278 | + (θ.comap (fun p : ℝ × ℝ ↦ ((c, p.1), p.2)) (by fun_prop)) _ := hZ |
| 279 | + have _ : HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1 + ω.2.1 + ω.2.2) (chain κ η θ c) _ := hout |
| 280 | + trivial |
| 281 | + |
| 282 | +end Automation |
| 283 | + |
| 284 | +end |
| 285 | + |
| 286 | +end |
0 commit comments