diff --git a/RandomDo.lean b/RandomDo.lean index e04384d..98c6e73 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -6,6 +6,15 @@ public import RandomDo.Monad.ForInInstances public import RandomDo.Monad.Instances public import RandomDo.Monad.MeasurableSpace public import RandomDo.Monad.Notation +public import RandomDo.Probability.AlgTrace +public import RandomDo.Probability.Examples +public import RandomDo.Probability.Extend +public import RandomDo.Probability.MeasurePreserving +public import RandomDo.Probability.Record +public import RandomDo.Probability.Tactic +public import RandomDo.Probability.Thompson +public import RandomDo.Probability.Trace +public import RandomDo.Probability.Transfer public import RandomDo.Tactic.Deriving public import RandomDo.Tactic.Elab public import RandomDo.Tactic.ForInStep diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean new file mode 100644 index 0000000..c49f21f --- /dev/null +++ b/RandomDo/Probability/AlgTrace.lean @@ -0,0 +1,614 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Probability.Tactic +public import RandomDo.Probability.Extend +public import LeanMachineLearning.SequentialLearning.IonescuTulceaSpace + +set_option linter.style.header false + +/-! +# The internal draws of an algorithm, as random variables of an algorithm-environment sequence + +`IsAlgEnvSeq A Y alg env P` says that `A` and `Y` are the actions and feedbacks generated by `alg` +interacting with `env`. It says nothing about *how* the algorithm produced its actions: when `alg` +comes from an `rdo` program, the draws that program makes are not random variables of `(Ω, P)` at +all. + +This file adds them. Given a trace of the algorithm — one space `Ω` of internal draws, a kernel +`K n` for their law at step `n`, and a readout `out n` reconstructing the action from them, which +is what `rdo_trace` produces — it builds an algorithm whose actions are pairs `(draws, action)`, +and shows that: + +* projecting away the draws turns an algorithm-environment sequence for it into one for `alg` + (`AlgTrace.isAlgEnvSeq_snd`); +* the draws have the conditional law `K n` given the history, and the action is `out n` of the + history and the draws (`AlgTrace.hasCondDistrib_trace`, `AlgTrace.action_ae_eq`). + +Since the traced algorithm interacts with the same environment, LML's `isAlgEnvSeq_unique` says +its trajectory has the same law as the original. So `AlgTrace.exists_isAlgEnvSeq_trace` may be used +to replace an arbitrary `IsAlgEnvSeq` hypothesis by one on a space that also carries the draws: the +space is existentially quantified precisely because it does not matter. + +## Main definitions + +* `RDo.AlgTrace alg Ω`: a trace of the algorithm `alg`, with draws in `Ω`. +* `RDo.AlgTrace.algorithm`: the algorithm whose actions carry their own draws. +* `Learning.Environment.withTrace`: the environment that ignores them. + +## Main results + +* `RDo.AlgTrace.isAlgEnvSeq_snd`: forgetting the draws recovers an algorithm-environment sequence + for the original algorithm. +* `RDo.AlgTrace.hasLaw_trace_zero`, `RDo.AlgTrace.hasCondDistrib_trace`: the law of the draws. +* `RDo.AlgTrace.action_zero_ae_eq`, `RDo.AlgTrace.action_ae_eq`: the action is the readout of the + history and the draws. +* `RDo.AlgTrace.exists_isAlgEnvSeq_trace`: any algorithm-environment sequence can be replaced by + one that also carries the draws, with the same trajectory law. + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Finset Learning + +noncomputable section + +attribute [fun_prop] Learning.measurable_history measurable_up measurable_down + +/-- An algorithm-environment sequence pulls back along a measure-preserving map. With +`extend_space`, this lets one add independent randomness to a space carrying such a sequence: as +a `@[transfer_forward]` lemma, it is how the hypothesis is transported to the extended space. -/ +@[transfer_forward] +lemma _root_.Learning.IsAlgEnvSeq.comp_measurePreserving {𝓐 𝓨 Ω Ω' : Type*} [MeasurableSpace 𝓐] + [MeasurableSpace 𝓨] {_ : MeasurableSpace Ω} {_ : MeasurableSpace Ω'} {P : Measure Ω} + [IsFiniteMeasure P] {P' : Measure Ω'} [IsFiniteMeasure P'] {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} + {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} {f : Ω' → Ω} + (h : IsAlgEnvSeq A Y alg env P) (hf : MeasurePreserving f P' P) : + IsAlgEnvSeq (fun n ω ↦ A n (f ω)) (fun n ω ↦ Y n (f ω)) alg env P' where + measurable_action n := (h.measurable_action n).comp hf.measurable + measurable_feedback n := (h.measurable_feedback n).comp hf.measurable + hasLaw_action_zero := h.hasLaw_action_zero.comp hf.hasLaw + hasCondDistrib_feedback_zero := h.hasCondDistrib_feedback_zero.comp_measurePreserving hf + hasCondDistrib_action n := (h.hasCondDistrib_action n).comp_measurePreserving hf + hasCondDistrib_feedback n := (h.hasCondDistrib_feedback n).comp_measurePreserving hf + +/-- Being an algorithm-environment sequence is invariant under pulling back along a +measure-preserving map, for measurable sequences. This is the form the `transfer` tactic uses. The +side conditions are the measurability of the sequences: when an `IsAlgEnvSeq` hypothesis `h` is +around, put `h.measurable_action` and `h.measurable_feedback` in the context for the discharger to +find them. -/ +@[transfer] +lemma _root_.MeasureTheory.MeasurePreserving.transfer_isAlgEnvSeq {𝓐 𝓨 Ω Ω' : Type*} + [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] {_ : MeasurableSpace Ω} {_ : MeasurableSpace Ω'} + {P : Measure Ω} [IsFiniteMeasure P] {P' : Measure Ω'} [IsFiniteMeasure P'] {f : Ω' → Ω} + (hf : MeasurePreserving f P' P) {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} + (hA : ∀ n, Measurable (A n)) (hY : ∀ n, Measurable (Y n)) {alg : Algorithm 𝓐 𝓨} + {env : Environment 𝓐 𝓨} : + IsAlgEnvSeq A Y alg env P ↔ IsAlgEnvSeq (fun n ω ↦ A n (f ω)) (fun n ω ↦ Y n (f ω)) alg env P' + where + mp h := h.comp_measurePreserving hf + mpr h := + { measurable_action := hA + measurable_feedback := hY + hasLaw_action_zero := (hf.hasLaw_fun_comp_iff (hA 0)).1 h.hasLaw_action_zero + hasCondDistrib_feedback_zero := + (hf.hasCondDistrib_fun_comp_iff (hA 0) (hY 0)).1 h.hasCondDistrib_feedback_zero + hasCondDistrib_action n := + (hf.hasCondDistrib_fun_comp_iff (measurable_history hA hY n) (hA (n + 1))).1 + (h.hasCondDistrib_action n) + hasCondDistrib_feedback n := + (hf.hasCondDistrib_fun_comp_iff ((measurable_history hA hY n).prodMk (hA (n + 1))) + (hY (n + 1))).1 (h.hasCondDistrib_feedback n) } + +namespace RDo + +universe u uA uY uW + +variable {𝓐 : Type uA} {𝓨 : Type uY} {Ω : Type uW} + [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] + +/-- Forget the draws carried alongside each action of a traced history. -/ +def forgetTrace (n : ℕ) (h : Iic n → (Ω × 𝓐) × 𝓨) : Iic n → 𝓐 × 𝓨 := + fun i ↦ ((h i).1.2, (h i).2) + +@[fun_prop] +lemma measurable_forgetTrace (n : ℕ) : + Measurable (forgetTrace (Ω := Ω) (𝓐 := 𝓐) (𝓨 := 𝓨) n) := + measurable_pi_lambda _ fun _ ↦ by unfold forgetTrace; fun_prop + +omit [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] in +@[simp] +lemma forgetTrace_history {Ω₀ : Type*} {_ : MeasurableSpace Ω₀} (A : ℕ → Ω₀ → Ω × 𝓐) + (Y : ℕ → Ω₀ → 𝓨) (n : ℕ) : + forgetTrace n ∘ history A Y n = history (fun n ω ↦ (A n ω).2) Y n := rfl + +/-- A *trace* of an algorithm: one space `Ω` of internal draws, a kernel giving their joint law at +each step given the history, and a readout reconstructing the action from them. This is what +`rdo_trace` produces for an algorithm whose policy is an `rdo` program. -/ +structure AlgTrace (alg : Algorithm 𝓐 𝓨) (Ω : Type*) [MeasurableSpace Ω] where + /-- The law of the draws the policy makes at step `n`, given the history. -/ + K : (n : ℕ) → Kernel (Iic n → 𝓐 × 𝓨) Ω + /-- Those are Markov kernels. -/ + [markov : ∀ n, IsMarkovKernel (K n)] + /-- The action at step `n + 1`, read off the history and the draws. -/ + out : (n : ℕ) → (Iic n → 𝓐 × 𝓨) × Ω → 𝓐 + /-- `K n` and `out n` trace the policy at step `n`. -/ + hasTrace (n : ℕ) : HasTrace (⇑(alg.policy n)) (K n) (out n) + /-- The law of the draws made before the first action. -/ + K0 : Measure Ω + /-- It is a probability measure. -/ + [markov0 : IsProbabilityMeasure K0] + /-- The first action, read off those draws. -/ + out0 : Ω → 𝓐 + /-- The first readout is measurable. -/ + measurable_out0 : Measurable out0 + /-- `K0` and `out0` trace the initial distribution. -/ + map_out0 : K0.map out0 = alg.p0 + +attribute [instance] AlgTrace.markov AlgTrace.markov0 + +/-- The environment a traced algorithm interacts with: the same one, reading only the action +component of each action-with-draws. -/ +def _root_.Learning.Environment.withTrace (Ω : Type*) [MeasurableSpace Ω] + (env : Environment 𝓐 𝓨) : Environment (Ω × 𝓐) 𝓨 where + feedback n := (env.feedback n).comap + (fun p : (Iic n → (Ω × 𝓐) × 𝓨) × (Ω × 𝓐) ↦ (forgetTrace n p.1, p.2.2)) (by fun_prop) + ν0 := env.ν0.comap Prod.snd measurable_snd + +namespace AlgTrace + +variable {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} (tr : AlgTrace alg Ω) + +@[fun_prop] +lemma measurable_out (n : ℕ) : + Measurable fun p : (Iic n → (Ω × 𝓐) × 𝓨) × Ω ↦ tr.out n (forgetTrace n p.1, p.2) := + (tr.hasTrace n).measurable_out.comp + (((measurable_forgetTrace n).comp measurable_fst).prodMk measurable_snd) + +/-- The traced algorithm: it draws the policy's internal randomness, then reads the action off it. +Its actions are pairs `(draws, action)`. -/ +def algorithm : Algorithm (Ω × 𝓐) 𝓨 where + policy n := + (tr.K n).comap (forgetTrace n) (by fun_prop) + ⊗ₖ Kernel.deterministic _ (tr.measurable_out n) + p0 := tr.K0 ⊗ₘ Kernel.deterministic tr.out0 tr.measurable_out0 + +/-- Comap turns a deterministic kernel into the deterministic kernel of the composite. -/ +lemma _root_.ProbabilityTheory.Kernel.comap_deterministic {α β γ : Type*} [MeasurableSpace α] + [MeasurableSpace β] [MeasurableSpace γ] {f : α → β} (hf : Measurable f) {g : γ → α} + (hg : Measurable g) : + (Kernel.deterministic f hf).comap g hg = Kernel.deterministic (f ∘ g) (hf.comp hg) := by + ext c s hs + simp [Kernel.comap_apply, Kernel.deterministic_apply] + +/-- The traced policy at a given history, with the draws forgotten, is the original policy. -/ +lemma map_snd_policy_apply (n : ℕ) (h : Iic n → (Ω × 𝓐) × 𝓨) : + (tr.algorithm.policy n h).map Prod.snd = alg.policy n (forgetTrace n h) := by + have hg : Measurable fun a ↦ tr.out n (forgetTrace n h, a) := + (tr.hasTrace n).measurable_out.comp (measurable_const.prodMk measurable_id) + have hp : Measurable fun a ↦ (a, tr.out n (forgetTrace n h, a)) := measurable_id.prodMk hg + change (((tr.K n).comap (forgetTrace n) (measurable_forgetTrace n) + ⊗ₖ Kernel.deterministic _ (tr.measurable_out n)) h).map Prod.snd = _ + rw [Kernel.compProd_apply_eq_compProd_sectR, Kernel.sectR, Kernel.comap_deterministic, + Measure.compProd_deterministic, Kernel.comap_apply] + simp only [Function.comp_apply] + rw [Measure.map_map measurable_snd hp] + exact (tr.hasTrace n).map_eq (forgetTrace n h) + +/-- Forgetting the draws of the traced policy gives back the original policy. -/ +lemma snd_policy (n : ℕ) : + (tr.algorithm.policy n).snd = (alg.policy n).comap (forgetTrace n) (by fun_prop) := + Kernel.ext fun h ↦ by + rw [Kernel.snd_apply, Kernel.comap_apply, tr.map_snd_policy_apply n h] + +/-- Forgetting the draws of the traced initial distribution gives back the original one. -/ +lemma map_snd_p0 : tr.algorithm.p0.map Prod.snd = alg.p0 := by + have hp : Measurable fun a ↦ (a, tr.out0 a) := measurable_id.prodMk tr.measurable_out0 + change (tr.K0 ⊗ₘ Kernel.deterministic tr.out0 tr.measurable_out0).map Prod.snd = alg.p0 + rw [Measure.compProd_deterministic, Measure.map_map measurable_snd hp, ← tr.map_out0] + rfl + +/-- `Prod.snd` of the traced initial distribution has the original law. -/ +lemma hasLaw_snd_p0 : HasLaw Prod.snd alg.p0 tr.algorithm.p0 := + ⟨measurable_snd.aemeasurable, tr.map_snd_p0⟩ + +variable {Ω₀ : Type*} {_ : MeasurableSpace Ω₀} {P : Measure Ω₀} [IsProbabilityMeasure P] + {A : ℕ → Ω₀ → Ω × 𝓐} {Y : ℕ → Ω₀ → 𝓨} + +section Projection + +variable (h : IsAlgEnvSeq A Y tr.algorithm (env.withTrace Ω) P) +include h + +/-- **Forgetting the draws.** An algorithm-environment sequence for the traced algorithm is, after +dropping the draws from each action, one for the original algorithm. -/ +lemma isAlgEnvSeq_snd : IsAlgEnvSeq (fun n ω ↦ (A n ω).2) Y alg env P where + measurable_action n := measurable_snd.comp (h.measurable_action n) + measurable_feedback n := h.measurable_feedback n + hasLaw_action_zero := tr.hasLaw_snd_p0.comp h.hasLaw_action_zero + hasCondDistrib_feedback_zero := + HasCondDistrib.comp_right (hf := measurable_snd) h.hasCondDistrib_feedback_zero + hasCondDistrib_action n := by + have h1 := (h.hasCondDistrib_action n).snd + rw [tr.snd_policy n] at h1 + exact HasCondDistrib.comp_right (hf := measurable_forgetTrace n) h1 + hasCondDistrib_feedback n := + HasCondDistrib.comp_right (hf := by fun_prop) (h.hasCondDistrib_feedback n) + +/-- The draws made before the first action have law `K0`. -/ +lemma hasLaw_trace_zero : HasLaw (fun ω ↦ (A 0 ω).1) tr.K0 P := + h.hasLaw_action_zero.compProd_fst + +/-- **The conditional law of the draws.** Given the history, the draws the policy makes at step +`n` are distributed as `K n` — the kernel the trace of the `rdo` program produced. -/ +lemma hasCondDistrib_trace (n : ℕ) : + HasCondDistrib (fun ω ↦ (A (n + 1) ω).1) (history (fun n ω ↦ (A n ω).2) Y n) (tr.K n) P := by + have h1 := (h.hasCondDistrib_action n).fst + rw [algorithm, Kernel.fst_compProd] at h1 + exact HasCondDistrib.comp_right (hf := measurable_forgetTrace n) h1 + +variable [MeasurableEq 𝓐] + +/-- **The action is the readout of the draws.** -/ +lemma action_zero_ae_eq : (fun ω ↦ (A 0 ω).2) =ᵐ[P] fun ω ↦ tr.out0 ((A 0 ω).1) := by + refine ae_eq_of_hasCondDistrib_deterministic tr.measurable_out0 ?_ ?_ ?_ + · exact (measurable_fst.comp (h.measurable_action 0)).aemeasurable + · exact (measurable_snd.comp (h.measurable_action 0)).aemeasurable + · exact h.hasLaw_action_zero.compProd_snd + +/-- **The action is the readout of the history and the draws.** -/ +lemma action_ae_eq (n : ℕ) : + (fun ω ↦ (A (n + 1) ω).2) + =ᵐ[P] fun ω ↦ tr.out n (history (fun n ω ↦ (A n ω).2) Y n ω, (A (n + 1) ω).1) := by + have hA := h.measurable_action + have hY := h.measurable_feedback + have h1 := (h.hasCondDistrib_action n).compProd_snd + have h2 := ae_eq_of_hasCondDistrib_deterministic (tr.measurable_out n) + (X := fun ω ↦ (history A Y n ω, (A (n + 1) ω).1)) + (by fun_prop) (measurable_snd.comp (h.measurable_action (n + 1))).aemeasurable h1 + exact h2 + +end Projection + +/-- **The draws may be assumed to be there.** Any algorithm-environment sequence for `alg` and +`env` can be replaced by one on a space that also carries the algorithm's internal draws `T`, with +the same trajectory law — so anything proved about the law of the actions and feedbacks there holds +of the original. The space is existentially quantified because, by `isAlgEnvSeq_unique`, it does +not matter; it may be taken in any universe at least those of `𝓐`, `𝓨` and `Ω`. + +Besides the laws of the draws and the readout equations, the actions and draws together form an +algorithm-environment sequence for the traced algorithm. This says more than the rest: the draws +at a step are conditionally independent of the earlier draws given the history, and the feedback +does not read the draws. -/ +theorem exists_isAlgEnvSeq_trace [MeasurableEq 𝓐] + {A₀ : ℕ → Ω₀ → 𝓐} {Y₀ : ℕ → Ω₀ → 𝓨} (h₀ : IsAlgEnvSeq A₀ Y₀ alg env P) : + ∃ (Ω' : Type (max u uA uY uW)) (_ : MeasurableSpace Ω') (P' : Measure Ω') + (_ : IsProbabilityMeasure P') (A : ℕ → Ω' → 𝓐) (Y : ℕ → Ω' → 𝓨) (T : ℕ → Ω' → Ω), + IsAlgEnvSeq A Y alg env P' + ∧ IsAlgEnvSeq (fun n ω ↦ (T n ω, A n ω)) Y tr.algorithm (env.withTrace Ω) P' + ∧ P'.map (trajectory A Y) = P.map (trajectory A₀ Y₀) + ∧ HasLaw (T 0) tr.K0 P' + ∧ (∀ n, HasCondDistrib (T (n + 1)) (history A Y n) (tr.K n) P') + ∧ A 0 =ᵐ[P'] (fun ω ↦ tr.out0 (T 0 ω)) + ∧ (∀ n, A (n + 1) =ᵐ[P'] fun ω ↦ tr.out n (history A Y n ω, T (n + 1) ω)) := by + -- The trajectory space of the traced algorithm, lifted to the universe asked for. + let base := trajMeasure tr.algorithm (env.withTrace Ω) + let e : ULift.{u} (ℕ → (Ω × 𝓐) × 𝓨) ≃ᵐ (ℕ → (Ω × 𝓐) × 𝓨) := MeasurableEquiv.ulift + have hup : MeasurePreserving e.symm base (base.map e.symm) := ⟨e.symm.measurable, rfl⟩ + have hf : MeasurePreserving e (base.map e.symm) base := hup.symm e.symm + have htr : IsAlgEnvSeq (fun n ω ↦ IT.action n (e ω)) (fun n ω ↦ IT.feedback n (e ω)) + tr.algorithm (env.withTrace Ω) (base.map e.symm) := + (IT.isAlgEnvSeq_trajMeasure tr.algorithm (env.withTrace Ω)).comp_measurePreserving hf + have : IsProbabilityMeasure (base.map e.symm) := + Measure.isProbabilityMeasure_map e.symm.measurable.aemeasurable + exact ⟨ULift (ℕ → (Ω × 𝓐) × 𝓨), inferInstance, base.map e.symm, inferInstance, + fun n ω ↦ (IT.action n (e ω)).2, fun n ω ↦ IT.feedback n (e ω), + fun n ω ↦ (IT.action n (e ω)).1, tr.isAlgEnvSeq_snd htr, htr, + isAlgEnvSeq_unique (tr.isAlgEnvSeq_snd htr) h₀, tr.hasLaw_trace_zero htr, + tr.hasCondDistrib_trace htr, tr.action_zero_ae_eq htr, tr.action_ae_eq htr⟩ + +/-- **The principle behind the `alg_env_trace` tactic.** To prove a statement `motive` about an +algorithm-environment sequence it is enough to prove it on a space that also carries the +algorithm's internal draws, *provided* the statement only depends on the law of the trajectory — +which is what the `transfer` hypothesis asks for, and which is exactly the freedom +`isAlgEnvSeq_unique` gives. The space may live in any universe at least those of `𝓐`, `𝓨` and +`Ω`. -/ +theorem wlog_trace [MeasurableEq 𝓐] + {motive : (Ω₀ : Type (max u uA uY uW)) → [MeasurableSpace Ω₀] → (P : Measure Ω₀) → + [IsProbabilityMeasure P] → (ℕ → Ω₀ → 𝓐) → (ℕ → Ω₀ → 𝓨) → Prop} + (traced : ∀ (Ω' : Type (max u uA uY uW)) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (A' : ℕ → Ω' → 𝓐) (Y' : ℕ → Ω' → 𝓨) (T : ℕ → Ω' → Ω), + IsAlgEnvSeq A' Y' alg env P' → + IsAlgEnvSeq (fun n ω ↦ (T n ω, A' n ω)) Y' tr.algorithm (env.withTrace Ω) P' → + HasLaw (T 0) tr.K0 P' → + (∀ n, HasCondDistrib (T (n + 1)) (history A' Y' n) (tr.K n) P') → + A' 0 =ᵐ[P'] (fun ω ↦ tr.out0 (T 0 ω)) → + (∀ n, A' (n + 1) =ᵐ[P'] fun ω ↦ tr.out n (history A' Y' n ω, T (n + 1) ω)) → + motive Ω' P' A' Y') + (transfer : ∀ (Ω₁ : Type (max u uA uY uW)) [MeasurableSpace Ω₁] (P₁ : Measure Ω₁) + [IsProbabilityMeasure P₁] (A₁ : ℕ → Ω₁ → 𝓐) (Y₁ : ℕ → Ω₁ → 𝓨) + (Ω₂ : Type (max u uA uY uW)) [MeasurableSpace Ω₂] (P₂ : Measure Ω₂) + [IsProbabilityMeasure P₂] (A₂ : ℕ → Ω₂ → 𝓐) (Y₂ : ℕ → Ω₂ → 𝓨), + IsAlgEnvSeq A₁ Y₁ alg env P₁ → IsAlgEnvSeq A₂ Y₂ alg env P₂ → + P₂.map (trajectory A₂ Y₂) = P₁.map (trajectory A₁ Y₁) → + motive Ω₂ P₂ A₂ Y₂ → motive Ω₁ P₁ A₁ Y₁) + : + ∀ (Ω₀ : Type (max u uA uY uW)) [MeasurableSpace Ω₀] (P : Measure Ω₀) [IsProbabilityMeasure P] + (A : ℕ → Ω₀ → 𝓐) (Y : ℕ → Ω₀ → 𝓨), IsAlgEnvSeq A Y alg env P → motive Ω₀ P A Y := by + intro Ω₀ _ P _ A Y h + -- The universe of the traced space is that of `Ω₀`; Lean does not solve it on its own. + obtain ⟨Ω', mΩ', P', hP', A', Y', T, hseq, htr, hlaw, hT0, hT, hA0, hA⟩ := + tr.exists_isAlgEnvSeq_trace.{u, _, _, _, _} h + exact transfer Ω₀ P A Y Ω' P' A' Y' h hseq hlaw + (traced Ω' P' A' Y' T hseq htr hT0 hT hA0 hA) + +end AlgTrace + + +end RDo + +end + +end + +public meta section + +open Lean Lean.Meta Lean.Elab Lean.Elab.Tactic +open MeasureTheory ProbabilityTheory Learning + +namespace RDo.Tactic + +initialize registerTraceClass `alg_env_trace + +/-- The free variables carrying the probability space of an `IsAlgEnvSeq` hypothesis: the space, +its σ-algebra, the measure, the `IsProbabilityMeasure` instance, and the two sequences. They have +to be local hypotheses, since the tactic abstracts the goal over them. -/ +def algEnvSpaceFVars (hFVar : FVarId) : MetaM (Array FVarId) := do + let ty ← instantiateMVars (← hFVar.getType) + unless ty.isAppOf ``IsAlgEnvSeq do + throwError "alg_env_trace: {Expr.fvar hFVar} is not an `IsAlgEnvSeq` hypothesis" + let as := ty.getAppArgs + let A := as[as.size - 6]! + let Y := as[as.size - 5]! + let P := as[as.size - 2]! + let mΩ := (← whnfR (← inferType P)).getAppArgs[1]! + let Ω := (← whnfR (← inferType P)).getAppArgs[0]! + let mut out : Array FVarId := #[] + for e in #[Ω, mΩ, P] do + let .fvar f := e | throwError + "alg_env_trace: the probability space must be given by local hypotheses, but {e} is not" + out := out.push f + -- the `IsProbabilityMeasure` hypothesis, which `isAlgEnvSeq_unique` needs + let some hP ← (do + for d in ← getLCtx do + if !d.isImplementationDetail then + if (← instantiateMVars d.type) == (← mkAppM ``IsProbabilityMeasure #[P]) then + return some d.fvarId + return none) | throwError + "alg_env_trace: no `IsProbabilityMeasure` hypothesis for {P} in the context" + out := out.push hP + for e in #[A, Y] do + let .fvar f := e | throwError + "alg_env_trace: the action and feedback sequences must be local hypotheses, but {e} is not" + out := out.push f + return out.push hFVar + +/-- How many binders of `wlog_trace` come before its conclusion: up to and including its third +explicit argument, `transfer`. Everything after that belongs to the `∀`-shaped conclusion, which is +what the goal is unified with. -/ +partial def preConclusionArity : Expr → Nat → Nat → Nat + | .forallE _ _ b bi, n, e => + let e := if bi.isExplicit then e + 1 else e + if e == 3 then n + 1 else preConclusionArity b (n + 1) e + | _, n, _ => n + +/-- Find an `IsAlgEnvSeq` hypothesis in the context. -/ +def findAlgEnvSeq? : MetaM (Option FVarId) := do + for d in ← getLCtx do + if !d.isImplementationDetail then + if (← instantiateMVars d.type).isAppOf ``IsAlgEnvSeq then return some d.fvarId + return none + +/-- `alg_env_trace tr` replaces an `IsAlgEnvSeq` situation by one in which the algorithm's internal +draws are present. `tr` is an `RDo.AlgTrace` for the algorithm — the trace of its policy, as +produced by `rdo_trace`. + +The goal, together with every hypothesis about the sequence — a statement mentioning the space, +the measure or the two sequences — is abstracted away from that space and two goals are left: + +* `traced`: the same statement on a space that also carries the draws `T`, with `hseq`, the + sequence again, `htr`, actions and draws together as a sequence for the traced algorithm, `hT₀` + and `hT`, the law of the draws and their conditional law given the history, and `hA₀` and `hA`, + each action as the readout of the history and the draws; +* `transfer`: the obligation that the statement only depends on the law of the trajectory. This is + what makes the replacement sound — the traced sequence lives on a different space, and all that + relates it to the original is `isAlgEnvSeq_unique`. The `transfer` tactic discharges it through + the trajectory space, onto which both sequences are measure-preserving maps, and the goal is only + left when that fails. + +Data on the space that is not the sequence — a random variable, an event, a point — has no +counterpart on the traced space: the goal may not depend on it, and it is cleared, together with +the hypotheses about it, before the change of space. + +* `alg_env_trace tr using h` names the hypothesis to use rather than searching for one. +* `alg_env_trace tr with Ω P A Y T hseq htr hT₀ hT hA₀ hA` names what is introduced. + +The probability space, its σ-algebra, the measure, the `IsProbabilityMeasure` hypothesis and the +two sequences all have to be local hypotheses, since the goal is abstracted over them, and the +space has to live in a universe at least those of the actions, the feedbacks and the draws. -/ +syntax (name := algEnvTraceTac) "alg_env_trace" ppSpace term (" using " ident)? + (" with " (ppSpace colGt ident)+)? : tactic + +elab_rules : tactic + | `(tactic| alg_env_trace $tr $[using $h?]? $[with $names?*]?) => withMainContext do + let g ← getMainGoal + let hFVar ← match h? with + | some h => getFVarId h + | none => match ← findAlgEnvSeq? with + | some f => pure f + | none => throwError "alg_env_trace: no `IsAlgEnvSeq` hypothesis in the context" + let spaceFVars ← algEnvSpaceFVars hFVar + let spaceSet : FVarIdSet := spaceFVars.foldl (·.insert ·) {} + let hTy ← instantiateMVars (← hFVar.getType) + let algE := hTy.getAppArgs[hTy.getAppNumArgs - 4]! + -- The trajectory space `ℕ → 𝓐 × 𝓨`, lifted to the universe of the space: the obligation is + -- discharged through it. + let liftTy ← do + let 𝓐 := (← instantiateMVars (← inferType (.fvar spaceFVars[4]!))).getForallBody + let 𝓨 := (← instantiateMVars (← inferType (.fvar spaceFVars[5]!))).getForallBody + let trajTy ← mkArrow (mkConst ``Nat) (← mkAppM ``Prod #[𝓐, 𝓨]) + pure (mkApp (mkConst ``ULift [← getDecLevel (.fvar spaceFVars[0]!), ← getDecLevel trajTy]) + trajTy) + let given := (names?.map (·.map (·.getId))).getD #[] + let defaults : Array Name := #[`Ω, `P, `A, `Y, `T, `hseq, `htr, `hT₀, `hT, `hA₀, `hA] + if given.size > defaults.size then + throwError "alg_env_trace: at most {defaults.size} names may be given" + -- What else mentions the space. A statement about the sequence travels with the goal, and is + -- pulled back to the traced space. Data on the space — a random variable, an event, a point — + -- has no counterpart there: the goal may not depend on it, and it is cleared, together with + -- what is about it. + let lctx ← getLCtx + let fvarsOf (d : LocalDecl) : MetaM (Array FVarId) := do + let mut st := Lean.collectFVars {} (← instantiateMVars d.type) + if let some v := d.value? then st := Lean.collectFVars st (← instantiateMVars v) + pure st.fvarIds + let mut props : Array FVarId := #[] + let mut data : FVarIdSet := {} + for d in lctx do + if d.isImplementationDetail || spaceSet.contains d.fvarId then continue + unless (← fvarsOf d).any spaceSet.contains do continue + if !d.isLet && (← isProp d.type) then props := props.push d.fvarId + else data := data.insert d.fvarId + let mut changed := true + while changed do + changed := false + for d in lctx do + if d.isImplementationDetail || data.contains d.fvarId || spaceSet.contains d.fvarId then + continue + if (← fvarsOf d).any data.contains then + data := data.insert d.fvarId + changed := true + let mut goalDeps : FVarIdSet := {} + let mut todo := (Lean.collectFVars {} (← instantiateMVars (← g.getType))).fvarIds + while !todo.isEmpty do + let x := todo.back! + todo := todo.pop + if goalDeps.contains x then continue + goalDeps := goalDeps.insert x + todo := todo ++ (← fvarsOf (← x.getDecl)) + for d in lctx do + if data.contains d.fvarId && goalDeps.contains d.fvarId then + throwError "alg_env_trace: the goal depends on{indentExpr d.toExpr}\nof type{indentExpr + (← instantiateMVars d.type)}\nwhich lives on the space of the sequence without being \ + part of it. Only statements about the actions, the feedbacks and the measure survive \ + the change of space." + let propsKept := props.filter (!data.contains ·) + let toClear ← sortFVarIds data.toArray + unless toClear.isEmpty do + trace[alg_env_trace] "cleared, as data on the space that is not the sequence: \ + {toClear.map Expr.fvar}" + let g ← g.tryClearMany toClear + let (propsReverted, g) ← g.revert propsKept + let (spaceReverted, g) ← g.revert spaceFVars (preserveOrder := true) + -- What was reverted beyond the space itself: the statements, and anything `revert` had to + -- take along. + let nTravelling := spaceReverted.size - spaceFVars.size + propsReverted.size + -- Build `wlog_trace tr ?traced ?transfer` and check it proves the abstracted goal. + let trE ← g.withContext do + let e ← Term.elabTerm tr none + Term.synthesizeSyntheticMVarsNoPostponing + instantiateMVars e + let (traced, transfer, motiveE) ← g.withContext do + let c ← mkConstWithFreshMVarLevels ``RDo.AlgTrace.wlog_trace + let cty ← inferType c + let (args, bis, concl) ← forallMetaBoundedTelescope cty (preConclusionArity cty 0 0) + let explicits := (args.zip bis).filterMap fun (a, b) ↦ if b.isExplicit then some a else none + let some iMotive := binderIndex? cty `motive + | throwError "alg_env_trace: `wlog_trace` no longer has a `motive` binder" + unless explicits.size == 3 do + throwError "alg_env_trace: `wlog_trace` no longer has the expected shape" + unless ← isDefEq explicits[0]! trE do + throwError "alg_env_trace: {trE} is not a trace of the algorithm of the hypothesis" + let some iAlg := binderIndex? cty `alg + | throwError "alg_env_trace: `wlog_trace` no longer has an `alg` binder" + unless ← isDefEq (← instantiateMVars args[iAlg]!) algE do + throwError "alg_env_trace: {trE} is not a trace of the algorithm of the hypothesis" + unless ← isDefEq concl (← g.getType) do + throwError "alg_env_trace: the goal does not have the expected shape{indentExpr + (← g.getType)}\nThe space has to live in a universe at least those of the actions, \ + the feedbacks and the draws." + for (a, b) in args.zip bis do + if b.isInstImplicit && !(← a.mvarId!.isAssigned) then + a.mvarId!.assign (← synthInstance (← instantiateMVars (← a.mvarId!.getType))) + g.assign (mkAppN c args) + let traced := explicits[1]!.mvarId! + let transfer := explicits[2]!.mvarId! + traced.setKind .syntheticOpaque + transfer.setKind .syntheticOpaque + traced.setTag `traced + transfer.setTag `transfer + return (traced, transfer, args[iMotive]!) + -- Introduce the traced space and its properties, then whatever travelled with the goal. + let pick (i : Nat) : Name := if h : i < given.size then given[i] else defaults[i]! + let intros : Array Name := + #[pick 0, `inst, pick 1, `inst, pick 2, pick 3, pick 4, pick 5, pick 6, pick 7, + pick 8, pick 9, pick 10] + let (_, traced) ← traced.introN intros.size intros.toList + let (_, traced) ← traced.introNP nTravelling + -- Discharge the transfer obligation through the trajectory space when `transfer` can: both + -- sequences are measure-preserving maps onto `(ℕ → 𝓐 × 𝓨, ν)`, on which the statement is + -- proved from the second sequence, then pulled back to the first. + let rest := (← getGoals).drop 1 + let s ← saveFullState + let transferLeft ← tryCatchRuntimeEx + (do + setGoals [transfer] + withMainContext do + let motiveStx ← Term.exprToSyntax (← instantiateMVars motiveE) + let liftStx ← Term.exprToSyntax liftTy + -- Without error recovery, a failure inside a nested `by` is a failure, not a `sorry`. + Term.withoutErrToSorry <| evalTactic (← `(tactic| ( + intro Ω₁ _ P₁ _ A₁ Y₁ Ω₂ _ P₂ _ A₂ Y₂ h₁ h₂ hlaw h + have hlaw' : (Measure.map (trajectory A₂ Y₂) P₂).map (ULift.up : _ → $liftStx) + = (Measure.map (trajectory A₁ Y₁) P₁).map ULift.up := by rw [hlaw] + generalize hν : (Measure.map (trajectory A₁ Y₁) P₁).map (ULift.up : _ → $liftStx) = ν + at hlaw' + have hf₁ : MeasurePreserving (fun ω ↦ (ULift.up (trajectory A₁ Y₁ ω) : $liftStx)) + P₁ ν := by + rw [← hν] + exact (⟨measurable_up, rfl⟩ : + MeasurePreserving ULift.up (Measure.map (trajectory A₁ Y₁) P₁) _).comp + ⟨measurable_trajectory h₁.measurable_action h₁.measurable_feedback, rfl⟩ + have hf₂ : MeasurePreserving (fun ω ↦ (ULift.up (trajectory A₂ Y₂ ω) : $liftStx)) + P₂ ν := by + rw [← hlaw'] + exact (⟨measurable_up, rfl⟩ : + MeasurePreserving ULift.up (Measure.map (trajectory A₂ Y₂) P₂) _).comp + ⟨measurable_trajectory h₂.measurable_action h₂.measurable_feedback, rfl⟩ + have : IsProbabilityMeasure (Measure.map (trajectory A₁ Y₁) P₁) := + Measure.isProbabilityMeasure_map + (measurable_trajectory h₁.measurable_action h₁.measurable_feedback).aemeasurable + have : IsProbabilityMeasure ν := by + rw [← hν] + exact Measure.isProbabilityMeasure_map measurable_up.aemeasurable + exact (fun hS : $motiveStx _ inferInstance ν inferInstance + (fun n (t : $liftStx) ↦ (t.down n).1) (fun n (t : $liftStx) ↦ (t.down n).2) ↦ + (by transfer hf₁ at hS; exact hS)) + (by beta_reduce; transfer hf₂)))) + unless (← getUnsolvedGoals).isEmpty do throwError "transfer left goals" + pure []) + (fun e ↦ do + let msg ← (← addMessageContextFull e.toMessageData).toString + s.restore + trace[alg_env_trace] "the transfer obligation is left, because: {msg}" + pure [transfer]) + setGoals ([traced] ++ transferLeft ++ rest) + +end RDo.Tactic + +end diff --git a/RandomDo/Probability/Examples.lean b/RandomDo/Probability/Examples.lean new file mode 100644 index 0000000..e62fa55 --- /dev/null +++ b/RandomDo/Probability/Examples.lean @@ -0,0 +1,286 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Probability.Tactic +public import RandomDo.Tactic.Elab + +set_option linter.style.header false + +/-! +# Reading probabilistic statements off an `rdo` program + +Worked examples of `RDo.HasTrace`. Nothing here is specific to a particular distribution: the +programs draw from arbitrary Markov kernels, which is exactly what an `rdo` program does once +`is_markov` has run on its leaves. + +The pattern is always the same. + +1. Build the trace of the program bottom-up with the combinators of `RandomDo.Probability.Trace`, + one per `rdo` construct. The trace kernel comes out as a right-nested `⊗ₖ`, one factor per `←`. +2. Fix the parameters `c`. The program's probability space is then `(Ω, P c)`, and the draws are + the coordinates of `Ω`. +3. Peel the `⊗ₖ` factors off one at a time with `HasLaw.compProd_snd` / `Kernel.sectR_compProd` / + `HasCondDistrib.compProd_fst` / `HasCondDistrib.compProd_snd`. Step `k` of the peeling hands + you the conditional distribution of the `k`-th draw given the `k-1` draws before it. + +Steps 1 and 3 are entirely mechanical, which is the point: `rdo_trace` and `rdo_peel` do them. +The `Automation` section at the end of this file re-derives, in two tactic calls, everything the +first two sections prove by hand. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory RDo +open MeasurableSpacePure MeasurableSpaceBind + +noncomputable section + +/-! ## Two independent draws + +``` +rdo + let x ← μ + let y ← μ + return x + y +``` +-/ + +section Independent + +variable (μ : Measure ℝ) [IsProbabilityMeasure μ] + +/-- A program that draws twice from the same distribution. It takes no parameter, so its trace is +a kernel from `Unit`. -/ +def sum2 : Measure ℝ := rdo + let x ← μ + let y ← μ + return x + y + +/-- The trace space of `sum2` is `ℝ × ℝ`, one coordinate per `←`, and the trace kernel is the +composition-product of the two draws. The second factor is a `Kernel.prodMkRight`, which records +in the *type* that the second draw does not look at the first. -/ +lemma hasTrace_sum2 : + HasTrace (fun _ : Unit ↦ sum2 μ) + (Kernel.const Unit μ ⊗ₖ Kernel.prodMkRight ℝ (Kernel.const Unit μ)) + (fun p : Unit × (ℝ × ℝ) ↦ p.2.1 + p.2.2) := by + have hfirst : HasTrace (fun _ : Unit ↦ μ) (Kernel.const Unit μ) Prod.snd := + (HasTrace.sample (Kernel.const Unit μ)).congr fun _ ↦ rfl + have htail : HasTrace (fun p : Unit × ℝ ↦ μ >>=ₘ fun y ↦ mPure (p.2 + y)) + (Kernel.prodMkRight ℝ (Kernel.const Unit μ)) + (fun q : (Unit × ℝ) × ℝ ↦ q.1.2 + q.2) := + (hfirst.prodMkRight ℝ).bindPure (f := fun q : (Unit × ℝ) × ℝ ↦ q.1.2 + q.2) (by fun_prop) + have hcont : Measurable fun p : Unit × ℝ ↦ μ >>=ₘ fun y ↦ mPure (p.2 + y) := by + have : IsMarkov fun p : Unit × ℝ ↦ μ >>=ₘ fun y ↦ mPure (p.2 + y) := by is_markov + exact this.measurable + exact (hfirst.bind hcont htail).congr fun _ ↦ rfl + +/-- The joint law of the two draws. -/ +local notation "P₂" => (Kernel.const Unit μ ⊗ₖ Kernel.prodMkRight ℝ (Kernel.const Unit μ)) () + +/-- The first draw has law `μ`. -/ +example : HasLaw (Prod.fst : ℝ × ℝ → ℝ) μ P₂ := + hasLaw_fst_compProd (Kernel.const Unit μ) (Kernel.prodMkRight ℝ (Kernel.const Unit μ)) () + +/-- The second draw has law `μ` too. -/ +example : HasLaw (Prod.snd : ℝ × ℝ → ℝ) μ P₂ := + hasLaw_snd_compProd_prodMkRight (P := Kernel.const Unit μ) (R := Kernel.const Unit μ) () + +/-- And the two are independent: this is read off the `Kernel.prodMkRight` in the trace kernel, +which is itself read off the fact that the second `←` does not mention `x`. -/ +example : IndepFun (Prod.fst : ℝ × ℝ → ℝ) Prod.snd P₂ := + indepFun_snd_compProd_prodMkRight (Kernel.const Unit μ) () + +/-- The program returns the sum of the two draws. Together with the three statements above, this +says exactly: `sum2 μ` is the law of `X + Y` for `X`, `Y` independent with law `μ`. -/ +example : HasLaw (fun ω : ℝ × ℝ ↦ ω.1 + ω.2) (sum2 μ) P₂ := (hasTrace_sum2 μ).hasLaw_out () + +end Independent + +/-! ## A dependent chain + +``` +rdo + let x ← κ c + let y ← η (c, x) + let z ← θ ((c, x), y) + return x + y + z +``` + +Every draw may read the parameter and all the draws before it. This is the general shape of a +straight-line `rdo` program: `κ`, `η`, `θ` stand for whatever `is_markov` produced at each `←`. +-/ + +section Chain + +variable (κ : Kernel ℝ ℝ) [IsMarkovKernel κ] (η : Kernel (ℝ × ℝ) ℝ) [IsMarkovKernel η] + (θ : Kernel ((ℝ × ℝ) × ℝ) ℝ) [IsMarkovKernel θ] + +/-- Three draws, each depending on everything before it. -/ +def chain (c : ℝ) : Measure ℝ := rdo + let x ← κ c + let y ← η (c, x) + let z ← θ ((c, x), y) + return x + y + z + +lemma hasTrace_chain : + HasTrace (chain κ η θ) (κ ⊗ₖ (η ⊗ₖ θ)) + (fun p : ℝ × (ℝ × (ℝ × ℝ)) ↦ p.2.1 + p.2.2.1 + p.2.2.2) := by + have h3 : HasTrace (fun q : (ℝ × ℝ) × ℝ ↦ θ q >>=ₘ fun z ↦ mPure (q.1.2 + q.2 + z)) θ + (fun r : ((ℝ × ℝ) × ℝ) × ℝ ↦ r.1.1.2 + r.1.2 + r.2) := + (HasTrace.sample θ).bindPure (f := fun r : ((ℝ × ℝ) × ℝ) × ℝ ↦ r.1.1.2 + r.1.2 + r.2) + (by fun_prop) + have hcont2 : Measurable fun q : (ℝ × ℝ) × ℝ ↦ θ q >>=ₘ fun z ↦ mPure (q.1.2 + q.2 + z) := by + have : IsMarkov fun q : (ℝ × ℝ) × ℝ ↦ θ q >>=ₘ fun z ↦ mPure (q.1.2 + q.2 + z) := by + is_markov + exact this.measurable + have h2 := (HasTrace.sample η).bind hcont2 h3 + have hcont1 : Measurable fun p : ℝ × ℝ ↦ + η p >>=ₘ fun y ↦ θ (p, y) >>=ₘ fun z ↦ mPure (p.2 + y + z) := by + have : IsMarkov fun p : ℝ × ℝ ↦ + η p >>=ₘ fun y ↦ θ (p, y) >>=ₘ fun z ↦ mPure (p.2 + y + z) := by is_markov + exact this.measurable + exact ((HasTrace.sample κ).bind hcont1 h2).congr fun _ ↦ rfl + +variable (c : ℝ) + +/-- The joint law of the three draws, given the parameter `c`. -/ +local notation "P₃" => (κ ⊗ₖ (η ⊗ₖ θ)) c + +/-- **The mechanical peeling.** One `compProd_*` step per `←` in the program: `X` has law `κ c`, +`Y` given `X` has law `η (c, X)`, and `Z` given `(X, Y)` has law `θ ((c, X), Y)`. The kernels on +the right-hand sides are literally the ones written in the program. -/ +example : + HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1) (κ c) P₃ + ∧ HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.1) (fun ω ↦ ω.1) (Kernel.sectR η c) P₃ + ∧ HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.2) (fun ω ↦ (ω.1, ω.2.1)) + (θ.comap (fun p : ℝ × ℝ ↦ ((c, p.1), p.2)) (by fun_prop)) P₃ := by + have h0 := hasLaw_id_compProd κ (η ⊗ₖ θ) c + have htail := h0.compProd_snd + rw [Kernel.sectR_compProd] at htail + exact ⟨h0.compProd_fst, htail.compProd_fst, htail.compProd_snd⟩ + +/-- Unfolding what those kernels are: `Kernel.sectR η c x = η (c, x)`, so the middle statement +above really is "given `X = x`, the second draw is distributed as `η (c, x)`". -/ +example (x : ℝ) : Kernel.sectR η c x = η (c, x) := Kernel.sectR_apply η x c + +/-- The program's result, as a random variable on the trace space. -/ +example : HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1 + ω.2.1 + ω.2.2) (chain κ η θ c) P₃ := + (hasTrace_chain κ η θ).hasLaw_out c + +end Chain + +/-! ## A loop, at coarse granularity + +A `for` loop is a Markov kernel, so `HasTrace.of_isMarkov` (here in its `HasTrace.sample` form) +makes the whole loop *one* trace coordinate holding its result. That is enough to state the +conditional law of everything that comes after the loop given what the loop produced; it does not +decompose the loop's own iterations, which needs the extra machinery described in +`notes/TRACE_SEMANTICS.md`. +-/ + +section Loop + +variable (μ : Measure ℝ) [IsProbabilityMeasure μ] (η : Kernel ℝ ℝ) [IsMarkovKernel η] + +/-- A loop summing `l.length` independent draws. -/ +def loopPart (l : List ℕ) : Measure ℝ := rdo + let mut S := 0 + for _ in l rdo + let x ← μ + S := S + x + return S + +instance (l : List ℕ) : IsProbabilityMeasure (loopPart μ l) := by + unfold loopPart + is_markov + +/-- The loop, then a draw whose distribution depends on what the loop returned. -/ +def loopThen (l : List ℕ) : Measure ℝ := rdo + let S ← loopPart μ l + let y ← η S + return y + +/-- The second `←` reads only the value the loop returned, i.e. the trace so far. -/ +def afterLoopK : Kernel (Unit × ℝ) ℝ := η.comap Prod.snd measurable_snd + +instance : IsMarkovKernel (afterLoopK η) := by + unfold afterLoopK + infer_instance + +variable (l : List ℕ) + +lemma hasTrace_loopThen : + HasTrace (fun _ : Unit ↦ loopThen μ η l) + (Kernel.const Unit (loopPart μ l) ⊗ₖ afterLoopK η) + (fun p : Unit × (ℝ × ℝ) ↦ p.2.2) := by + have hloop : HasTrace (fun _ : Unit ↦ loopPart μ l) + (Kernel.const Unit (loopPart μ l)) Prod.snd := + (HasTrace.sample (Kernel.const Unit (loopPart μ l))).congr fun _ ↦ rfl + have htail : HasTrace (fun p : Unit × ℝ ↦ η p.2) (afterLoopK η) + (Prod.snd : (Unit × ℝ) × ℝ → ℝ) := + (HasTrace.sample (afterLoopK η)).congr fun _ ↦ rfl + have hcont : Measurable fun p : Unit × ℝ ↦ η p.2 := + (Kernel.measurable η).comp measurable_snd + exact (hloop.bind hcont htail).congr fun _ ↦ by unfold loopThen; rfl + +/-- The loop's result has the loop's law. -/ +example : HasLaw (Prod.fst : ℝ × ℝ → ℝ) (loopPart μ l) + ((Kernel.const Unit (loopPart μ l) ⊗ₖ afterLoopK η) ()) := + hasLaw_fst_compProd (Kernel.const Unit (loopPart μ l)) (afterLoopK η) () + +/-- And given it, the draw that follows the loop has law `η S`. -/ +example : HasCondDistrib (Prod.snd : ℝ × ℝ → ℝ) Prod.fst (Kernel.sectR (afterLoopK η) ()) + ((Kernel.const Unit (loopPart μ l) ⊗ₖ afterLoopK η) ()) := + hasCondDistrib_snd_compProd _ _ () + +example (S : ℝ) : Kernel.sectR (afterLoopK η) () S = η S := rfl + +end Loop + +/-! ## The same, automatically + +`rdo_trace` walks the program and builds the trace; `rdo_peel` reads the laws off it. The +hypotheses they leave are exactly the statements proved by hand above. +-/ + +section Automation + +variable (μ : Measure ℝ) [IsProbabilityMeasure μ] (κ : Kernel ℝ ℝ) [IsMarkovKernel κ] + (η : Kernel (ℝ × ℝ) ℝ) [IsMarkovKernel η] (θ : Kernel ((ℝ × ℝ) × ℝ) ℝ) [IsMarkovKernel θ] + +/-- Two independent draws: the tactics find the trace, both laws, the independence, and the law of +the result. -/ +example : True := by + rdo_trace (sum2 μ) with h + rdo_peel h () with hX hY hY' hindep hout + -- `h : HasTrace (fun _ ↦ sum2 μ) + -- (Kernel.const Unit μ ⊗ₖ Kernel.prodMkRight ℝ (Kernel.const Unit μ)) + -- (fun p ↦ p.2.1 + p.2.2)` + have _ : HasLaw (Prod.fst : ℝ × ℝ → ℝ) μ _ := hX + have _ : HasCondDistrib (Prod.snd : ℝ × ℝ → ℝ) Prod.fst _ _ := hY + have _ : HasLaw (Prod.snd : ℝ × ℝ → ℝ) μ _ := hY' + have _ : IndepFun (Prod.fst : ℝ × ℝ → ℝ) Prod.snd _ := hindep + have _ : HasLaw (fun ω : ℝ × ℝ ↦ ω.1 + ω.2) (sum2 μ) _ := hout + trivial + +/-- The dependent chain: one conditional law per `←`, with the kernel written at that `←`. -/ +example (c : ℝ) : True := by + rdo_trace (chain κ η θ) with h + rdo_peel h c with hX hY hZ hout + have _ : HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1) (κ c) _ := hX + have _ : HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.1) (fun ω ↦ ω.1) + (η.comap (fun ω ↦ (c, ω)) (by fun_prop)) _ := hY + have _ : HasCondDistrib (fun ω : ℝ × (ℝ × ℝ) ↦ ω.2.2) (fun ω ↦ (ω.1, ω.2.1)) + (θ.comap (fun p : ℝ × ℝ ↦ ((c, p.1), p.2)) (by fun_prop)) _ := hZ + have _ : HasLaw (fun ω : ℝ × (ℝ × ℝ) ↦ ω.1 + ω.2.1 + ω.2.2) (chain κ η θ c) _ := hout + trivial + +end Automation + +end + +end diff --git a/RandomDo/Probability/Extend.lean b/RandomDo/Probability/Extend.lean new file mode 100644 index 0000000..32a7010 --- /dev/null +++ b/RandomDo/Probability/Extend.lean @@ -0,0 +1,930 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Probability.MeasurePreserving +public import Mathlib.Probability.Kernel.Composition.MeasureCompProd +public meta import Lean.Elab.Tactic.Basic + +set_option linter.style.header false + +/-! +# Extending a probability space by a new random variable + +`extend_space μ` adds to the probability space `(Ω, P)` of the goal a random variable `Z` with law +`μ`, independent of everything defined on `Ω`. It is the tactic form of "without loss of +generality, let `Z ~ μ` be independent of the rest". Underneath, the new space is `Ω × E` with the +product measure, presented as an abstract space related to the old one by a measure-preserving map +`f`, which is all the extended goal may use. Three presentations are available. + +* `extend_space μ` keeps the names: `Ω`, `P`, every random variable `X : Ω → α` and every event + `s : Set Ω` now denote objects on the extended space, hypotheses about them are transported, and + the goal reads as before. The old space and its objects are still there, renamed `Ω₀`, `P₀`, + `X₀`, …, together with the map `f : Ω → Ω₀`, `hf : MeasurePreserving f P P₀` and the defining + equations `hX_def : ∀ ω, X₀ (f ω) = X ω`. A hypothesis that cannot be transported stays about the + old space, under its `₀` name. The context gains `Z : Ω → E`, `hZm : Measurable Z`, + `hZ : HasLaw Z μ P` and `hind`, the independence of `Z` from the transported random variables + and events, as a tuple, an event `s` entering as `fun ω ↦ ω ∈ s`. A random variable whose + measurability cannot be proved is left out of the tuple; when nothing is left, `hind` is the + independence of `Z` from the map `f`. +* `extend_space! μ` does the same and clears the old space, the map, and everything that + mentions them, except what the hypotheses that stay need: when `hZ` or `hind` is stated in terms + of `f`, the map stays with `hf`, and the old space with its instances. +* `extend_space_map μ` is the explicit form: nothing is renamed, the new space is `Ω'` with + `P'`, `f : Ω' → Ω`, `hf`, `Z`, `hZm`, `hZ : HasLaw Z μ P'` and `hind : IndepFun f Z P'`, and the + goal is restated with `fun ω ↦ X (f ω)` for `X` and `f ⁻¹' s` for `s`. Hypotheses about the old + space stay as they are and are pulled back on demand, by `transfer hf at h` or by hand. + +Extending twice works as expected: the first draw `Z` and its hypotheses are transported like any +other random variable, and the second extension names its objects `Z'`, `hZ'`, `f'`, … so as not +to shadow them. The spaces left behind are `Ω₀`, the one just left, then `Ω₀₀`. + +In every case a `transfer` goal is left when the `transfer` tactic cannot discharge it: the +obligation that the statement pulls back along any measure-preserving map, which is what makes +the replacement sound. A statement about the points of `Ω`, such as `∀ ω, 0 ≤ X ω`, does not pull +back, and its obligation is left. A hypothesis the goal itself depends on, such as a measurability +proof inside a `Kernel.comap`, is generalized along with the goal. The goal may not depend on data +on `Ω` other than random variables and events, such as a second measure, nor on a local +definition: neither can be transported, and the tactic fails. + +`extend_space κ` for a Markov kernel `κ : Kernel Ω E` gives instead a draw with conditional law +`κ` given the old space, `hZ : HasCondDistrib Z f κ P'`, and no `hind`. For a draw conditional on +random variables, extend with `κ.comap g hg`: `extend_space` then states `hZ` as +`HasCondDistrib Z g κ P`, with `g` read in terms of the transported variables. + +Compare `alg_env_trace`: there the new space is not an extension of the old one, only a space with +the same trajectory law, so its `transfer` obligation is about laws and has to be proved +statement by statement. Here the projection `f` is measure preserving, which is a uniform principle. + +## Main results + +* `RDo.wlog_extend`, `RDo.wlog_extend_kernel`: the principles behind the tactic. +* `RDo.indepFun_fst_snd_prod`, `RDo.hasCondDistrib_snd_fst_compProd`: the product extension. + +The lemmas the `transfer` tactic rewrites with are in `RandomDo.Probability.MeasurePreserving`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory + +noncomputable section + +namespace RDo + +/-! ### The product extension -/ + +section Product + +variable {Ω E : Type*} [MeasurableSpace Ω] [MeasurableSpace E] {P : Measure Ω} + +/-- Under a product of probability measures, the two coordinates are independent. -/ +lemma indepFun_fst_snd_prod [IsProbabilityMeasure P] (μ : Measure E) [IsProbabilityMeasure μ] : + IndepFun (Prod.fst : Ω × E → Ω) Prod.snd (P.prod μ) := by + rw [indepFun_iff_map_prod_eq_prod_map_map measurable_fst.aemeasurable + measurable_snd.aemeasurable] + simp [Measure.map_id', measurePreserving_fst.map_eq, measurePreserving_snd.map_eq] + +/-- The first coordinate of `P ⊗ₘ κ` has law `P`. -/ +lemma measurePreserving_fst_compProd [SFinite P] (κ : Kernel Ω E) [IsMarkovKernel κ] : + MeasurePreserving Prod.fst (P ⊗ₘ κ) P := + ⟨measurable_fst, Measure.fst_compProd P κ⟩ + +/-- Under `P ⊗ₘ κ`, the second coordinate has conditional law `κ` given the first. -/ +lemma hasCondDistrib_snd_fst_compProd [SFinite P] (κ : Kernel Ω E) [IsMarkovKernel κ] : + HasCondDistrib Prod.snd Prod.fst κ (P ⊗ₘ κ) := by + change HasLaw _ ((P ⊗ₘ κ).fst ⊗ₘ κ) _ + rw [Measure.fst_compProd] + exact HasLaw.id + +end Product + +universe u v + +/-- **The principle behind `extend_space κ`.** To prove a statement `motive` about the probability +space `(Ω, P)`, it is enough to prove it on a space `(Ω', P')` that projects onto `Ω` by a +measure-preserving map `f` and carries a measurable draw `Z` with conditional law `κ` given `f`, +*provided* +the statement pulls back along measure-preserving maps, which is what `transfer` asks for. + +The space `Ω'` is the product `Ω × E` with the measure `P ⊗ₘ κ`, but `extended` may not use that: +all it knows of `Ω'` is `f`, `Z` and their laws. The universe of `E` may not exceed that of `Ω`, +since the product has to live in the universe of `Ω`. -/ +theorem wlog_extend_kernel {Ω : Type (max u v)} [mΩ : MeasurableSpace Ω] {E : Type v} + [MeasurableSpace E] {P : Measure Ω} [hP : IsProbabilityMeasure P] + {motive : (Ω' : Type (max u v)) → [MeasurableSpace Ω'] → (P' : Measure Ω') → + [IsProbabilityMeasure P'] → (Ω' → Ω) → Prop} + (κ : Kernel Ω E) [IsMarkovKernel κ] + (extended : ∀ (Ω' : Type (max u v)) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → ∀ (Z : Ω' → E), + Measurable Z → HasCondDistrib Z f κ P' → motive Ω' P' f) + (transfer : ∀ (Ω' : Type (max u v)) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → + motive Ω' P' f → motive Ω P id) : + motive Ω P id := + transfer (Ω × E) (P ⊗ₘ κ) Prod.fst (measurePreserving_fst_compProd κ) + (extended (Ω × E) (P ⊗ₘ κ) Prod.fst (measurePreserving_fst_compProd κ) Prod.snd + measurable_snd (hasCondDistrib_snd_fst_compProd κ)) + +/-- **The principle behind `extend_space μ`.** To prove a statement `motive` about the probability +space `(Ω, P)`, it is enough to prove it on a space `(Ω', P')` that projects onto `Ω` by a +measure-preserving map `f` and carries a measurable draw `Z` with law `μ`, independent of `f`, +*provided* +the statement pulls back along measure-preserving maps, which is what `transfer` asks for. + +The space `Ω'` is the product `Ω × E` with the product measure, but `extended` may not use that: +all it knows of `Ω'` is `f`, `Z` and their laws. The universe of `E` may not exceed that of `Ω`, +since the product has to live in the universe of `Ω`. -/ +theorem wlog_extend {Ω : Type (max u v)} [mΩ : MeasurableSpace Ω] {E : Type v} + [MeasurableSpace E] {P : Measure Ω} [hP : IsProbabilityMeasure P] + {motive : (Ω' : Type (max u v)) → [MeasurableSpace Ω'] → (P' : Measure Ω') → + [IsProbabilityMeasure P'] → (Ω' → Ω) → Prop} + (μ : Measure E) [IsProbabilityMeasure μ] + (extended : ∀ (Ω' : Type (max u v)) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → ∀ (Z : Ω' → E), + Measurable Z → HasLaw Z μ P' → IndepFun f Z P' → motive Ω' P' f) + (transfer : ∀ (Ω' : Type (max u v)) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → + motive Ω' P' f → motive Ω P id) : + motive Ω P id := + transfer (Ω × E) (P.prod μ) Prod.fst measurePreserving_fst + (extended (Ω × E) (P.prod μ) Prod.fst measurePreserving_fst Prod.snd measurable_snd + measurePreserving_snd.hasLaw (indepFun_fst_snd_prod μ)) + +end RDo + +end + +end + +public meta section + +open Lean Lean.Meta Lean.Elab Lean.Elab.Tactic +open MeasureTheory ProbabilityTheory + +namespace RDo.Tactic + +initialize registerTraceClass `extend_space + +/-! ### The space and what depends on it -/ + +/-- The probability space being extended, as local hypotheses: the space, its σ-algebra, the +measure and, when it is a local hypothesis, its `IsProbabilityMeasure` instance. They have to be +local hypotheses, since the goal is abstracted over them. -/ +structure ProbSpace where + /-- The space. -/ + Ω : FVarId + /-- Its σ-algebra. -/ + mΩ : FVarId + /-- The measure. -/ + P : FVarId + /-- The `IsProbabilityMeasure` hypothesis, when it is a local one. -/ + hP : Option FVarId + +/-- The `IsProbabilityMeasure` hypothesis for `P`, if there is a local one. -/ +def findIsProbabilityMeasure? (P : Expr) : MetaM (Option FVarId) := do + let target ← mkAppM ``IsProbabilityMeasure #[P] + for d in ← getLCtx do + if !d.isImplementationDetail then + if (← instantiateMVars d.type) == target then return some d.fvarId + return none + +/-- The space of a measure `P`, read off its type `Measure Ω`. -/ +def ProbSpace.ofMeasure (P : Expr) : MetaM ProbSpace := do + let .fvar Pf := P + | throwError "extend_space: the measure must be a local hypothesis, but {P} is not" + let ty ← whnfR (← inferType P) + unless ty.isAppOfArity ``Measure 2 do throwError "extend_space: {P} is not a measure" + let Ω := ty.appFn!.appArg! + let mΩ := ty.appArg! + let .fvar Ωf := Ω + | throwError "extend_space: the space must be a local hypothesis, but {Ω} is not" + let .fvar mf := mΩ + | throwError "extend_space: the σ-algebra on {Ω} must be a local hypothesis, but {mΩ} is not" + return { Ω := Ωf, mΩ := mf, P := Pf, hP := ← findIsProbabilityMeasure? P } + +/-- The local hypotheses of type `Measure Ω`, with `Ω` a local hypothesis, that `e` mentions. -/ +def measureFVars (e : Expr) : MetaM (Array FVarId) := do + let mut out := #[] + for f in (Lean.collectFVars {} e).fvarIds do + let ty ← whnfR (← instantiateMVars (← f.getType)) + if ty.isAppOfArity ``Measure 2 && ty.appFn!.appArg!.isFVar then out := out.push f + return out + +/-- The local hypotheses depending on the space, directly or through other hypotheses. The space +itself, its σ-algebra, the measure and its `IsProbabilityMeasure` hypothesis are included. -/ +def spaceDependents (sp : ProbSpace) : MetaM FVarIdSet := do + let mut dep : FVarIdSet := {} + for f in #[sp.Ω, sp.mΩ, sp.P] ++ sp.hP.toArray do dep := dep.insert f + for d in ← getLCtx do + if d.isImplementationDetail || dep.contains d.fvarId then continue + -- A hypothesis introduced by `have` may still carry assigned metavariables in its type. + let mentions (e : Expr) : MetaM Bool := return (← instantiateMVars e).hasAnyFVar dep.contains + if (← mentions d.type) || (← (d.value?.mapM mentions)).getD false then + dep := dep.insert d.fvarId + return dep + +/-- The shape of a transportable type, `ι₁ → ⋯ → ιₖ → Ω → α` or `ι₁ → ⋯ → ιₖ → Set Ω` with `Ω` +appearing nowhere else: the number `k` of leading binders, and whether it is a family of sets. -/ +partial def transportShape (Ω : FVarId) (ty : Expr) (k : Nat := 0) : Option (Nat × Bool) := + if ty.isAppOfArity ``Set 1 then + if ty.appArg! == .fvar Ω then some (k, true) else none + else match ty with + | .forallE _ d b _ => + if d == .fvar Ω then (if b.containsFVar Ω then none else some (k, false)) + else if d.containsFVar Ω then none + else transportShape Ω b (k + 1) + | _ => none + +/-- Whether a local hypothesis of type `ty` can be transported along `f : Ω' → Ω`. -/ +def transportable (Ω : FVarId) (ty : Expr) : Bool := + (transportShape Ω ty).isSome + +/-- Transport `x : ι₁ → ⋯ → ιₖ → Ω → α` along `f : Ω' → Ω` to `fun i₁ … iₖ ω ↦ x i₁ … iₖ (f ω)`, +and `s : ι₁ → ⋯ → ιₖ → Set Ω` to `fun i₁ … iₖ ↦ f ⁻¹' s i₁ … iₖ`. -/ +partial def transportAlong (Ω Ω' f : Expr) (x ty : Expr) : MetaM Expr := do + if ty.isAppOfArity ``Set 1 then + return ← mkAppM ``Set.preimage #[f, x] + match ty with + | .forallE n d b bi => + if d == Ω then + withLocalDecl `ω bi Ω' fun ω ↦ mkLambdaFVars #[ω] (mkApp x (mkApp f ω)) + else + let n := if n.hasMacroScopes || n.isAnonymous then `i else n + withLocalDecl n bi d fun i ↦ do + mkLambdaFVars #[i] (← transportAlong Ω Ω' f (mkApp x i) (b.instantiate1 i)) + | _ => throwError "extend_space: cannot transport {x} : {ty} to the extended space" + +/-- The local hypotheses that `e` mentions, that depend on the space but cannot be transported: +these are generalized. The set is closed under the hypotheses their types mention, so that it is +compatible with `mkForallFVars`. -/ +partial def toGeneralize (sp : ProbSpace) (deps special : FVarIdSet) (e : Expr) : + MetaM (Array FVarId) := + go (Lean.collectFVars {} e).fvarIds.toList {} #[] +where + go (todo : List FVarId) (seen : FVarIdSet) (acc : Array FVarId) : MetaM (Array FVarId) := do + match todo with + | [] => return acc + | f :: todo => + if seen.contains f then return ← go todo seen acc + let seen := seen.insert f + if !deps.contains f || special.contains f then return ← go todo seen acc + let d ← f.getDecl + let ty ← instantiateMVars d.type + if d.isLet then throwError + "extend_space: the goal depends on the local definition {Expr.fvar f}, which cannot be \ + transported to the extended space; unfold it or `clear_value` it first" + if transportable sp.Ω ty then return ← go todo seen acc + go ((Lean.collectFVars {} ty).fvarIds.toList ++ todo) seen (acc.push f) + +/-- Index of the binder named `n` in a `∀`-telescope. -/ +partial def binderIndex? (ty : Expr) (n : Name) (i : Nat := 0) : Option Nat := + match ty with + | .forallE m _ b _ => if m == n then some i else binderIndex? b n (i + 1) + | _ => none + +/-- The domain of the `i`-th binder of a `∀`-telescope. -/ +def binderDomain! (ty : Expr) : Nat → Expr + | 0 => ty.bindingDomain! + | i + 1 => binderDomain! ty.bindingBody! i + +/-- The number of leading `∀`s. -/ +partial def forallArity : Expr → Nat + | .forallE _ _ b _ => forallArity b + 1 + | _ => 0 + +/-! ### The hidden presentation -/ + +/-- The new space and what comes with it, once the extended goal has been introduced. -/ +structure NewSpace where + /-- The new space. -/ + Ω : FVarId + /-- Its σ-algebra. -/ + mΩ : FVarId + /-- The new measure. -/ + P : FVarId + /-- Its `IsProbabilityMeasure` instance. -/ + hP : FVarId + /-- The projection onto the old space. -/ + f : FVarId + /-- `MeasurePreserving f P' P`. -/ + hf : FVarId + /-- The new draw. -/ + Z : FVarId + /-- `Measurable Z`. -/ + hZm : FVarId + /-- Its law, or its conditional law given `f`. -/ + hZ : FVarId + /-- `IndepFun f Z P'`, in the independent case. -/ + hind? : Option FVarId + +/-- Eta-reduce every subterm. -/ +def etaAll (e : Expr) : CoreM Expr := + Core.transform e (post := fun e ↦ return .done e.eta) + +/-- Replace, in `e`, `X i₁ … iₖ (f t)` by `X' i₁ … iₖ t` and `f ⁻¹' s i₁ … iₖ` by `s' i₁ … iₖ`, +for every transported variable `(X, k, isSet, X')`. -/ +partial def foldTransports (f : FVarId) (subst : Array (FVarId × Nat × Bool × Expr)) (e : Expr) : + Expr := + e.replace fun e ↦ + if e.isAppOfArity ``Set.preimage 4 && e.getArg! 2 == .fvar f then + let s := e.appArg! + match s.getAppFn with + | .fvar x => + match subst.find? (·.1 == x) with + | some (_, k, true, s') => + let args := s.getAppArgs + if args.size == k then some (mkAppN s' (args.map (foldTransports f subst))) else none + | _ => none + | _ => none + else match e.getAppFn with + | .fvar x => + match subst.find? (·.1 == x) with + | some (_, k, false, X') => + let args := e.getAppArgs + if h : k < args.size then + let a := args[k] + if a.isApp && a.appFn! == .fvar f then + let args' := (args.extract 0 k).push a.appArg! ++ args.extract (k + 1) args.size + some (mkAppN X' (args'.map (foldTransports f subst))) + else none + else none + | _ => none + | _ => none + +/-- A transported variable `x` at a point `ω` of its space, as a component of the tuple of the +independence statement: `fun i₁ … iₖ ↦ x i₁ … iₖ ω` for a random variable +`x : ι₁ → ⋯ → ιₖ → Ω → α`, and `fun i₁ … iₖ ↦ ω ∈ x i₁ … iₖ` for a family of events. -/ +partial def component (Ω ω x ty : Expr) (isSet : Bool) : MetaM Expr := + if isSet && ty.isAppOfArity ``Set 1 then mkAppM ``Membership.mem #[x, ω] + else match ty with + | .forallE n d b bi => + if d == Ω then pure (mkApp x ω) + else + let n := if n.hasMacroScopes || n.isAnonymous then `i else n + withLocalDecl n bi d fun i ↦ do + mkLambdaFVars #[i] (← component Ω ω (mkApp x i) (b.instantiate1 i) isSet) + | _ => throwError "extend_space: internal error, {x} : {ty} is not a random variable" + +/-- Generalize the transports in the goal: for each transported variable `(X, v, k, isSet, name)` +with transport `v`, the goal `G` becomes `∀ (X' : _) (hX'_def : ∀ i… ω, X i… (f ω) = X' i… ω), G'` +where `G'` is `G` with `v` folded into `X'`. Returns the goal with those introduced, the new +variables and the defining equations. -/ +def generalizeTransports (g : MVarId) (f : FVarId) + (transported : Array (FVarId × Expr × Nat × Bool × Name)) : + MetaM (MVarId × Array FVarId × Array FVarId) := g.withContext do + let G ← instantiateMVars (← g.getType) + let n := transported.size + let decls : Array (Name × BinderInfo × (Array Expr → MetaM Expr)) ← + transported.mapM fun (_, v, _, _, name) ↦ do + let ty ← inferType v + pure (name, .default, fun _ ↦ pure ty) + let (ty, args) ← withLocalDecls decls fun X's ↦ do + let defDecls : Array (Name × BinderInfo × (Array Expr → MetaM Expr)) ← + (transported.zip X's).mapM fun ((_, v, _, _, name), X') ↦ do + let dty ← forallTelescope (← inferType v) fun bs _ ↦ do + mkForallFVars bs (← mkEq (mkAppN v bs).headBeta (mkAppN X' bs)) + pure (Name.mkSimple s!"h{name}_def", .default, fun _ ↦ pure dty) + withLocalDecls defDecls fun hdefs ↦ do + let subst := (transported.zip X's).map fun ((x, _, k, isSet, _), X') ↦ (x, k, isSet, X') + let G' ← etaAll (foldTransports f subst G) + let ty ← mkForallFVars (X's ++ hdefs) G' + let rfls ← transported.mapM fun (_, v, _, _, _) ↦ do + forallTelescope (← inferType v) fun bs _ ↦ do + mkLambdaFVars bs (← mkEqRefl (mkAppN v bs).headBeta) + pure (ty, transported.map (·.2.1) ++ rfls) + let g₂ ← mkFreshExprSyntheticOpaqueMVar ty (← g.getTag) + g.assign (mkAppN g₂ args) + let (fvs, g₂) ← g₂.mvarId!.introNP (2 * n) + return (g₂, fvs.extract 0 n, fvs.extract n (2 * n)) + +/-- Prove the transported statement `φ'` of a hypothesis `h` about the old space: by a +`@[transfer_forward]` lemma, by transferring a copy of `h`, or, for a measurability statement or +an instance, from scratch. -/ +def proveTransported (h : FVarId) (φ' : Expr) (new : NewSpace) : TacticM (Option Expr) := do + let hE := Expr.fvar h + let hfE := Expr.fvar new.hf + if let some (ty, pf) ← transferForward? hE hfE then + if ← isDefEq ty φ' then return some pf + let hStx ← Term.exprToSyntax hE + let hfStx ← Term.exprToSyntax hfE + let mut tacs : Array (TSyntax `tactic) := + #[← `(tactic| (have h' := $hStx; transfer $hfStx at h'; exact h'))] + if (← isClass? φ').isSome then tacs := tacs.push (← `(tactic| infer_instance)) + let funProps : Array Name := #[``Measurable, ``AEMeasurable, ``AEStronglyMeasurable, + ``StronglyMeasurable] + let setProps : Array Name := #[``MeasurableSet, ``NullMeasurableSet] + if let some c := φ'.getForallBody.getAppFn.constName? then + if funProps.contains c then tacs := tacs.push (← `(tactic| (intros; fun_prop))) + if setProps.contains c then tacs := tacs.push (← `(tactic| (intros; measurability))) + for tac in tacs do + let goal ← mkFreshExprSyntheticOpaqueMVar φ' + if let some pf ← tryTactic? goal.mvarId! tac then return some pf + return none + +/-- The independence of `Z` from the transported random variables and events, as the independence +of `Z` and the tuple of those whose measurability `fun_prop` or `measurability` can prove, or a +`MeasurePreserving` hypothesis gives, an event `s` entering as `fun ω ↦ ω ∈ s`. The tuple is +written on the new space with the transports, as `X (f ω)` and `ω ∈ f ⁻¹' s`, which is what +`generalizeTransports` folds into `X ω` and `ω ∈ s`. When no component is left, `hind` stays as it +is, the independence of `Z` and `f`. -/ +def deriveIndep (g : MVarId) (sp : ProbSpace) (new : NewSpace) (hind : FVarId) + (transported : Array (FVarId × Expr × Nat × Bool)) : TacticM (Option (FVarId × MVarId)) := + g.withContext do + let ΩE := Expr.fvar sp.Ω + let Ω'E := Expr.fvar new.Ω + -- Each component on the old space, with its measurability, and on the new space. + let mut comps : Array (Expr × Expr × Expr) := #[] + for (x, v, _, isSet) in transported do + let ty ← instantiateMVars (← x.getType) + let c ← withLocalDecl `ω .default ΩE fun ω ↦ do + mkLambdaFVars #[ω] (← component ΩE ω (.fvar x) ty isSet) + let c := if isSet then c else c.eta + let goal ← mkFreshExprSyntheticOpaqueMVar (← mkAppM ``Measurable #[c]) + -- The map of a previous extension is measurable by its `MeasurePreserving` hypothesis. + let some hc ← tryTactic? goal.mvarId! (← `(tactic| first + | fun_prop + | (exact MeasureTheory.MeasurePreserving.measurable ‹_›) + | measurability)) + | continue + let c' ← withLocalDecl `ω .default Ω'E fun ω ↦ do + mkLambdaFVars #[ω] (← component Ω'E ω v (← inferType v) isSet) + comps := comps.push (c, hc, c') + if comps.isEmpty then return none + let rec mkTuple : List (Expr × Expr × Expr) → MetaM (Expr × Expr × Expr) + | [] => throwError "extend_space: internal error, empty tuple" + | [ch] => pure ch + | (c, hc, c') :: rest => do + let (r, hr, r') ← mkTuple rest + let φ ← withLocalDecl `ω .default ΩE fun ω ↦ do + mkLambdaFVars #[ω] (← mkAppM ``Prod.mk #[(mkApp c ω).headBeta, (mkApp r ω).headBeta]) + let φ' ← withLocalDecl `ω .default Ω'E fun ω ↦ do + mkLambdaFVars #[ω] (← mkAppM ``Prod.mk #[(mkApp c' ω).headBeta, (mkApp r' ω).headBeta]) + pure (φ, ← mkAppM ``Measurable.prodMk #[hc, hr], φ') + let (_, hφ, tuple) ← mkTuple comps.toList + let tuple ← Core.betaReduce tuple + let ty ← mkAppM ``ProbabilityTheory.IndepFun #[tuple, .fvar new.Z, .fvar new.P] + let goal ← mkFreshExprSyntheticOpaqueMVar ty + let hindStx ← Term.exprToSyntax (.fvar hind) + let hφStx ← Term.exprToSyntax hφ + let some pf ← tryTactic? goal.mvarId! + (← `(tactic| exact ProbabilityTheory.IndepFun.comp $hindStx $hφStx measurable_id)) + | return none + let g ← g.assert (← hind.getUserName) ty pf + let (hind', g) ← g.intro1P + return some (hind', g) + +/-- For a kernel `κ.comap g hg`, the conditional law of `Z` given `g` on the new space, +`HasCondDistrib Z (fun ω ↦ g (f ω)) κ P'`. Once `generalizeTransports` has folded the transports, +`g (f ω)` reads as `g` of the transported variables: `X ω` for `κ.comap X hX`, `(X ω, Y ω)` for a +kernel conditioned on a pair. -/ +def deriveCondDistrib (g : MVarId) (new : NewSpace) (κ : Expr) : + TacticM (Option (FVarId × MVarId)) := g.withContext do + unless κ.isAppOfArity ``ProbabilityTheory.Kernel.comap 9 do return none + let X := κ.getArg! 7 + let κ₀ := κ.getArg! 6 + let fE := Expr.fvar new.f + let Xf ← withLocalDecl `ω .default (.fvar new.Ω) fun ω ↦ + mkLambdaFVars #[ω] (mkApp X (mkApp fE ω)).headBeta + let ty ← mkAppM ``ProbabilityTheory.HasCondDistrib #[.fvar new.Z, Xf, κ₀, .fvar new.P] + let goal ← mkFreshExprSyntheticOpaqueMVar ty + let hZStx ← Term.exprToSyntax (.fvar new.hZ) + let some pf ← tryTactic? goal.mvarId! + (← `(tactic| exact ProbabilityTheory.HasCondDistrib.comp_right $hZStx)) | return none + let g ← g.assert (← new.hZ.getUserName) ty pf + let (hZ', g) ← g.intro1P + return some (hZ', g) + +/-- `n`, bumped by `bump` until it is not in `taken`. -/ +partial def freshName (bump : Name → Name) (n : Name) (taken : NameSet) : Name := + if taken.contains n then freshName bump (bump n) taken else n + +/-- Rename `x` to `n`. A hypothesis already named `n` is first renamed `n₀`, recursively, so that +`Ω` becomes `Ω₀`, the `Ω₀` of a previous extension becomes `Ω₀₀`, and so on. -/ +partial def renameBumping (g : MVarId) (x : FVarId) (n : Name) : MetaM MVarId := do + let g ← g.withContext do + match (← getLCtx).findFromUserName? n with + | some d => + if d.fvarId != x && !d.isImplementationDetail then + renameBumping g d.fvarId (n.appendAfter "₀") + else pure g + | none => pure g + g.rename x n + +/-- Remove from `toClear` what the hypotheses that stay need: everything their types mention, +transitively, and, about those objects, the instances and the hypotheses listed in `about`. With +`hZ : HasCondDistrib Z f κ₀ P'` staying, this keeps `f` with `hf`, `κ₀` with its Markov instance, +and the old space with its σ-algebra and probability instance. -/ +def keepNeeded (g : MVarId) (toClear : Array FVarId) (about : Array FVarId) : + MetaM (Array FVarId) := g.withContext do + let clearSet : FVarIdSet := toClear.foldl (·.insert ·) {} + let lctx ← getLCtx + let fvarsOf (d : LocalDecl) : MetaM (Array FVarId) := do + let mut st := Lean.collectFVars {} (← instantiateMVars d.type) + if let some v := d.value? then st := Lean.collectFVars st (← instantiateMVars v) + pure st.fvarIds + let mut keep : FVarIdSet := {} + let mut todo := (Lean.collectFVars {} (← instantiateMVars (← g.getType))).fvarIds + for d in lctx do + if d.isImplementationDetail || clearSet.contains d.fvarId then continue + todo := todo ++ (← fvarsOf d) + let mut progress := true + while progress do + progress := false + while !todo.isEmpty do + let x := todo.back! + todo := todo.pop + if keep.contains x || !clearSet.contains x then continue + keep := keep.insert x + progress := true + todo := todo ++ (← fvarsOf (← x.getDecl)) + for d in lctx do + if !clearSet.contains d.fvarId || keep.contains d.fvarId then continue + if (← isClass? d.type).isSome || about.contains d.fvarId then + if (← instantiateMVars d.type).hasAnyFVar keep.contains then + keep := keep.insert d.fvarId + progress := true + todo := todo ++ (← fvarsOf d) + return toClear.filter (!keep.contains ·) + +/-- Hide the extension. The new space and its objects take the names of the old ones, which are +renamed with `₀`; every hypothesis about the old space is transported when possible; the +transports `fun ω ↦ X (f ω)` become fresh variables `X` with defining equations +`hX_def : ∀ ω, X₀ (f ω) = X ω`; the independence of `Z` is restated against the transported random +variables and events; and, for a kernel `κ.comap g hg`, the conditional law of `Z` is stated given +`g`. With `clearOld`, the old space, the map and everything mentioning them are cleared, except +what the remaining hypotheses need. -/ +def hidePresentation (g : MVarId) (sp : ProbSpace) (deps special : FVarIdSet) (new : NewSpace) + (gensOld gensNew : Array FVarId) (κ? : Option Expr) (clearOld : Bool) : TacticM MVarId := do + let ΩE := Expr.fvar sp.Ω + let Ω'E := Expr.fvar new.Ω + let fE := Expr.fvar new.f + let hfE := Expr.fvar new.hf + -- The old objects, in context order, with their names. + let olds ← g.withContext do + let mut out : Array (FVarId × Name) := #[] + for d in ← getLCtx do + if !d.isImplementationDetail && deps.contains d.fvarId then + out := out.push (d.fvarId, d.userName) + pure out + let origName (x : FVarId) : Name := ((olds.find? (·.1 == x)).map (·.2)).getD .anonymous + -- 1. The old objects are renamed with `₀`, and the new space takes the old names. A hypothesis + -- already named `Ω₀`, left by a previous extension, becomes `Ω₀₀` first. + let mut g := g + for (x, n) in olds do + unless n.hasMacroScopes do g ← renameBumping g x (n.appendAfter "₀") + g ← g.rename new.Ω (origName sp.Ω) + g ← g.rename new.P (origName sp.P) + unless (origName sp.mΩ).hasMacroScopes do g ← g.rename new.mΩ (origName sp.mΩ) + if let some hP := sp.hP then + unless (origName hP).hasMacroScopes do g ← g.rename new.hP (origName hP) + -- 2. `Measurable f`, for `fun_prop` and `measurability`. + let (hfm, g') ← g.withContext do + let g ← g.assert `hfm (← mkAppM ``Measurable #[fE]) + (← mkAppM ``MeasurePreserving.measurable #[hfE]) + g.intro1P + g := g' + -- 3. The transportable variables and their transports. + let transported ← g.withContext do + let mut out : Array (FVarId × Expr × Nat × Bool) := #[] + for (x, _) in olds do + if special.contains x then continue + let d ← x.getDecl + -- A local definition is not transported: its transport would not have its value. + if d.isLet then continue + let ty ← instantiateMVars d.type + if ← isProp ty then continue + if let some (k, isSet) := transportShape sp.Ω ty then + out := out.push (x, ← transportAlong ΩE Ω'E fE (.fvar x) ty, k, isSet) + pure out + let transportedSet : FVarIdSet := transported.foldl (fun s t ↦ s.insert t.1) {} + let xs := #[ΩE, .fvar sp.mΩ, .fvar sp.P] ++ sp.hP.toArray.map Expr.fvar + ++ transported.map (Expr.fvar ·.1) + let vs := #[Ω'E, .fvar new.mΩ, .fvar new.P] ++ (if sp.hP.isSome then #[.fvar new.hP] else #[]) + ++ transported.map (·.2.1) + -- 4. The hypotheses about the old space are transported when possible. + let mut moved : Array FVarId := #[] + for (h, n) in olds do + if special.contains h || transportedSet.contains h || gensOld.contains h then continue + let r ← g.withContext do + let φ ← instantiateMVars (← h.getType) + unless ← isProp φ do return none + if φ.hasAnyFVar (fun x ↦ deps.contains x && !special.contains x + && !transportedSet.contains x) then + return none + let φ' ← Core.betaReduce (φ.replaceFVars xs vs) + let ok ← try check φ'; pure true catch _ => pure false + unless ok do return none + let some pf ← proveTransported h φ' new | return none + let g ← g.assert n φ' pf + let (h', g) ← g.intro1P + return some (h', g) + if let some (h', g') := r then + g := g' + moved := moved.push h' + -- 5. The independence of `Z`, against the transported random variables. + let mut toClear : Array FVarId := #[hfm] + if let some hind := new.hind? then + if let some (hind', g') ← deriveIndep g sp new hind transported then + g := g' + moved := moved.push hind' + toClear := toClear.push hind + -- 6. The conditional law of `Z` given the conditioning variable, for `κ.comap X hX`. + if let some κ := κ? then + if let some (hZ', g') ← deriveCondDistrib g new κ then + g := g' + moved := moved.push hZ' + toClear := toClear.push new.hZ + -- 7. The transports become fresh variables, named as the old ones. + let mut hdefs : Array FVarId := #[] + unless transported.isEmpty do + let (reverted, g') ← g.revert (moved ++ gensNew) + let (g', _, hdefs') ← generalizeTransports g' new.f + (transported.map fun (x, v, k, isSet) ↦ (x, v, k, isSet, origName x)) + let (_, g') ← g'.introNP reverted.size + g := g' + hdefs := hdefs' + -- 8. Clean up. + if clearOld then + toClear ← keepNeeded g (toClear ++ hdefs ++ #[new.f, new.hf] ++ olds.map (·.1)) + #[new.hf, hfm] + let sorted ← g.withContext do sortFVarIds toClear + g.tryClearMany sorted + +/-! ### The tactics -/ + +/-- How the extended space is presented. -/ +inductive ExtendMode where + /-- `extend_space_map`: the new space is `Ω'`, with the map `f : Ω' → Ω` explicit. -/ + | map + /-- `extend_space`: the new space takes the names of the old one, which is renamed with `₀`. -/ + | hidden + /-- `extend_space!`: as `hidden`, and the old space is cleared. -/ + | clear + +/-- The common implementation of `extend_space`, `extend_space!` and `extend_space_map`. -/ +def extendSpace (mode : ExtendMode) (μ : Term) (P? : Option Ident) (given : Array Name) : + TacticM Unit := withMainContext do + let tac := match mode with + | .map => "extend_space_map" + | .hidden => "extend_space" + | .clear => "extend_space!" + let g ← getMainGoal + let T₀ ← instantiateMVars (← g.getType) + -- The measure or kernel to extend with. + let μE ← Term.elabTerm μ none + Term.synthesizeSyntheticMVarsNoPostponing + let μE ← instantiateMVars μE + let μty ← whnfR (← inferType μE) + let (isKernel, E, dom?) ← + if μty.isAppOfArity ``Measure 2 then pure (false, μty.appFn!.appArg!, none) + else if μty.isAppOfArity ``Kernel 4 then + let as := μty.getAppArgs + pure (true, as[1]!, some as[0]!) + else throwError + "{tac}: {μE} is neither a measure nor a kernel; it has type{indentExpr μty}" + let (instName, what, cls) := + if isKernel then (``IsMarkovKernel, "a Markov kernel", "IsMarkovKernel") + else (``IsProbabilityMeasure, "a probability measure", "IsProbabilityMeasure") + try discard <| synthInstance (← mkAppM instName #[μE]) + catch _ => throwError + "{tac}: {μE} is not known to be {what}: no `{cls}` instance was found" + -- The measure to extend. + let PE ← match P? with + | some P => pure (Expr.fvar (← getFVarId P)) + | none => do + let mut cands ← measureFVars T₀ + if let .fvar m := μE then cands := cands.erase m + if let some dom := dom? then + cands ← cands.filterM fun P ↦ do + let ty ← whnfR (← inferType (.fvar P)) + isDefEq ty.appFn!.appArg! dom + match cands with + | #[P] => pure (Expr.fvar P) + | #[] => throwError + "{tac}: the goal mentions no measure on a local space; name one with `using`" + | _ => throwError + "{tac}: the goal mentions several measures, {cands.map Expr.fvar}; choose one with `using`" + let sp ← ProbSpace.ofMeasure PE + let ΩE := Expr.fvar sp.Ω + if sp.hP.isNone then + try discard <| synthInstance (← mkAppM ``IsProbabilityMeasure #[PE]) + catch _ => throwError + "{tac}: {PE} is not known to be a probability measure: no `IsProbabilityMeasure` instance \ + was found" + if let some dom := dom? then + unless ← isDefEq dom ΩE do + throwError "{tac}: the kernel {μE} is on {dom}, not on {ΩE}" + -- The product `Ω × E` has to live in the universe of `Ω`. + let lvlΩ ← getDecLevel ΩE + let lvlE ← getDecLevel E + unless ← isLevelDefEq (mkLevelMax lvlΩ lvlE) lvlΩ do + throwError "{tac}: {E} lives in universe {toString lvlE} and {ΩE} in universe \ + {toString lvlΩ},\nso the product {ΩE} × {E} does not live in the universe of {ΩE}: lift \ + {E} with `ULift`" + -- Hypotheses the goal depends on that cannot be transported are generalized. + let deps ← spaceDependents sp + let mut special : FVarIdSet := {} + for f in #[sp.Ω, sp.mΩ, sp.P] ++ sp.hP.toArray do special := special.insert f + let gens ← sortFVarIds (← toGeneralize sp deps special T₀) + -- Generalizing data on `Ω`, a second measure say, would make the extended goal quantify over + -- all such data on the new space: not what was asked, and not provable in general. + for x in gens do + let ty ← instantiateMVars (← x.getType) + unless ← isProp ty do + throwError "{tac}: the goal depends on{indentExpr (Expr.fvar x)}\nof type{indentExpr ty}\n\ + which is neither a random variable nor an event on {ΩE}, so it cannot be transported to \ + the extended space" + let T ← mkForallFVars (gens.map Expr.fvar) T₀ + -- The motive: the goal on a space `Ω'` with a map `f : Ω' → Ω`. + let motive ← + withLocalDecl (.mkSimple "Ω'") .implicit (← inferType ΩE) fun Ω' ↦ do + withLocalDecl (.mkSimple "mΩ'") .instImplicit (← mkAppM ``MeasurableSpace #[Ω']) fun mΩ' ↦ do + withLocalDecl (.mkSimple "P'") .default (← mkAppOptM ``Measure #[Ω', mΩ']) fun P' ↦ do + withLocalDecl (.mkSimple "hP'") .instImplicit + (← mkAppOptM ``IsProbabilityMeasure #[Ω', mΩ', P']) fun hP' ↦ do + withLocalDecl `f .default (← mkArrow Ω' ΩE) fun f ↦ do + let mut xs := #[ΩE, .fvar sp.mΩ, PE] + let mut vs := #[Ω', mΩ', P'] + if let some h := sp.hP then + xs := xs.push (.fvar h) + vs := vs.push hP' + for x in (Lean.collectFVars {} T).fvarIds do + if deps.contains x && !special.contains x then + let ty ← instantiateMVars (← x.getType) + xs := xs.push (.fvar x) + vs := vs.push (← transportAlong ΩE Ω' f (.fvar x) ty) + let T' ← Core.betaReduce (T.replaceFVars xs vs) + try check T' + catch e => throwError + "{tac}: the goal does not survive the change of space:{indentExpr T'}\n\ + {e.toMessageData}" + mkLambdaFVars #[Ω', mΩ', P', hP', f] T' + -- Apply the principle. + let thm := if isKernel then ``RDo.wlog_extend_kernel else ``RDo.wlog_extend + let c := mkConst thm [lvlΩ, lvlE] + let cty ← inferType c + let idx (n : Name) : TacticM Nat := do + let some i := binderIndex? cty n + | throwError "{tac}: `{thm}` no longer has the expected shape" + pure i + let iTransfer ← idx `transfer + let (args, bis, concl) ← forallMetaBoundedTelescope cty (iTransfer + 1) + let assign (n : Name) (e : Expr) : TacticM Unit := do + unless ← isDefEq args[← idx n]! e do + throwError "{tac}: cannot use {e} as `{n}` of `{thm}`" + assign `Ω ΩE + assign `mΩ (.fvar sp.mΩ) + assign `P PE + if let some h := sp.hP then assign `hP (.fvar h) + assign `motive motive + assign (if isKernel then `κ else `μ) μE + unless ← isDefEq concl T do + throwError "{tac}: the goal does not have the expected shape{indentExpr T}" + for (a, b) in args.zip bis do + if b.isInstImplicit && !(← a.mvarId!.isAssigned) then + a.mvarId!.assign (← synthInstance (← instantiateMVars (← a.mvarId!.getType))) + g.assign (mkAppN (mkAppN c args) (gens.map Expr.fvar)) + -- The extended goal: introduce the new space and what was generalized. + let extended := args[← idx `extended]!.mvarId! + extended.setKind .syntheticOpaque + extended.setTag `extended + let nNames := match mode, isKernel with + | .map, false => 7 + | .map, true => 6 + | _, false => 5 + | _, true => 4 + if given.size > nNames then + throwError "{tac}: at most {nNames} names may be given" + -- A default name that a hypothesis of the old space bears, and that its transport will bear + -- again, is primed: a second extension names its draw `Z'`. In the explicit form, where nothing + -- is renamed, any hypothesis counts. + let taken : NameSet ← (← getLCtx).foldlM (init := {}) fun taken d ↦ do + if d.isImplementationDetail then return taken + let counts := match mode with + | .map => true + | _ => deps.contains d.fvarId + return if counts then taken.insert d.userName else taken + let fresh (n : Name) : Name := freshName (·.appendAfter "'") n taken + let pick (defaults : Array Name) (i : Nat) : Name := + if h : i < given.size then given[i] else fresh defaults[i]! + -- `hZm` for a draw `Z`, `hUm` for a draw named `U`; primed like the rest when taken. + let hZmName (defaults : Array Name) (i : Nat) : Name := + let z := if h : i < given.size then given[i] else defaults[i]! + fresh (Name.mkSimple s!"h{z}m") + let intros : Array Name := match mode with + | .map => + let d := #[.mkSimple "Ω'", .mkSimple "P'", `f, `hf, `Z, `hZ, `hind] + #[pick d 0, `inst, pick d 1, `inst, pick d 2, pick d 3, pick d 4, hZmName d 4, pick d 5] + ++ (if isKernel then #[] else #[pick d 6]) + | _ => + let d := if isKernel then #[`Z, `hZ, `f, `hf] else #[`Z, `hZ, `hind, `f, `hf] + let (f, hf) := if isKernel then (pick d 2, pick d 3) else (pick d 3, pick d 4) + #[.mkSimple "Ω'", `inst, .mkSimple "P'", `inst, f, hf, pick d 0, hZmName d 0, pick d 1] + ++ (if isKernel then #[] else #[pick d 2]) + let (fvs, extended) ← extended.introN intros.size intros.toList + let gensNames ← gens.toList.mapM fun x ↦ do + let n ← x.getUserName + pure (match mode with | .map => n.appendAfter "'" | _ => n) + let (gensNew, extended) ← extended.introN gens.size gensNames + let extended ← match mode with + | .map => pure extended + | _ => + let new : NewSpace := ⟨fvs[0]!, fvs[1]!, fvs[2]!, fvs[3]!, fvs[4]!, fvs[5]!, fvs[6]!, + fvs[7]!, fvs[8]!, if isKernel then none else some fvs[9]!⟩ + hidePresentation extended sp deps special new gens gensNew + (if isKernel then some μE else none) (match mode with | .clear => true | _ => false) + -- The transfer goal, with the original goal as its conclusion rather than `motive Ω P id`. + let transfer := args[iTransfer]!.mvarId! + let tty ← instantiateMVars (← transfer.getType) + -- Only the binders of the theorem's hypothesis: its conclusion may itself be a `∀`. + let nBinders := forallArity (binderDomain! cty iTransfer) + let tty' ← forallBoundedTelescope tty (some nBinders) fun xs _ ↦ mkForallFVars xs T + let transfer' ← mkFreshExprSyntheticOpaqueMVar tty' (tag := `transfer) + transfer.assign transfer' + -- Discharge the transfer obligation with the `transfer` tactic when it can; leave it otherwise. + let rest := (← getGoals).drop 1 + let s ← saveFullState + -- `tryCatchRuntimeEx`: a failure inside `measurability` may be a maximum recursion depth error, + -- which `try … catch` lets through. + let transferLeft ← tryCatchRuntimeEx + (do + setGoals [transfer'.mvarId!] + Term.withoutErrToSorry <| evalTactic (← `(tactic| transfer)) + unless (← getUnsolvedGoals).isEmpty do throwError "transfer left goals" + pure []) + (fun e ↦ do + let msg ← (← addMessageContextFull e.toMessageData).toString + s.restore + trace[extend_space] "the transfer obligation is left, because: {msg}" + pure [transfer'.mvarId!]) + setGoals ([extended] ++ transferLeft ++ rest) + +/-- `extend_space μ` adds to the probability space `(Ω, P)` of the goal a random variable `Z` with +law `μ`, a probability measure on some `E`, independent of everything defined on `Ω`. The names +are kept: `Ω`, `P`, every random variable `X : Ω → α` and every event `s : Set Ω` now denote +objects on the extended space, hypotheses about them are transported, and the goal reads as +before. The context gains + +* `Z : Ω → E`, `hZm : Measurable Z`, `hZ : HasLaw Z μ P`, and `hind`, the independence of `Z` + from the transported random variables and events, as a tuple, or from the map `f` when none is + provably measurable; +* the old space and its objects, renamed `Ω₀`, `P₀`, `X₀`, …, with `f : Ω → Ω₀`, + `hf : MeasurePreserving f P P₀`, and the defining equations `hX_def : ∀ ω, X₀ (f ω) = X ω`. + A hypothesis that cannot be transported stays about the old space, under its `₀` name. + +Two goals may be left: `extended`, the goal on the new space, and `transfer`, the obligation that +the statement pulls back along a measure-preserving map, which makes the replacement sound. The +`transfer` tactic is run on it, and it is only left when that fails. + +* `extend_space! μ` also clears the old space, the map and everything mentioning them, except + what the hypotheses that stay need, such as the map when `hZ` or `hind` is stated with it, and + what Lean does not let a tactic clear: hypotheses introduced by `variable`. +* Extending twice is fine: the second extension names its objects `Z'`, `hZ'`, `f'`, … so as not + to shadow the first draw, which is transported like any other random variable. +* `extend_space κ` for a Markov kernel `κ : Kernel Ω E` gives instead a draw with conditional law + `κ` given the old space, `hZ : HasCondDistrib Z f κ P`, and no `hind`. For `κ.comap g hg`, `hZ` + is stated as `HasCondDistrib Z g κ P`, with `g` read in terms of the transported variables. +* `extend_space μ using P` names the measure to extend rather than reading it off the goal. +* `extend_space μ with Z hZ hind f hf` names what is introduced (`with Z hZ f hf` for a kernel). + +The space, its σ-algebra and the measure have to be local hypotheses, since the goal is abstracted +over them, and `E` has to live in the universe of `Ω` or in a smaller one. -/ +syntax (name := extendSpaceTac) "extend_space" ppSpace term (" using " ident)? + (" with " (ppSpace colGt ident)+)? : tactic + +@[inherit_doc extendSpaceTac] +syntax (name := extendSpaceClearTac) "extend_space!" ppSpace term (" using " ident)? + (" with " (ppSpace colGt ident)+)? : tactic + +/-- `extend_space_map μ` is the explicit form of `extend_space μ`: nothing is renamed, the goal is +restated on a new space `Ω'` with a measure-preserving map `f : Ω' → Ω`, and the context gains +`hf : MeasurePreserving f P' P`, `Z : Ω' → E`, `hZm : Measurable Z`, `hZ : HasLaw Z μ P'` and +`hind : IndepFun f Z P'`. +Every random variable `X : Ω → α` of the goal becomes `fun ω ↦ X (f ω)` and every event +`s : Set Ω` becomes `f ⁻¹' s`. Hypotheses about the old space stay as they are and are pulled back +on demand, by `transfer hf at h`, `hX.comp hf.hasLaw` for a law, `hind.comp hX measurable_id` for +the independence of `X ∘ f` and `Z`, or `h.comp_measurePreserving hf` for a conditional law. A +hypothesis the goal itself depends on is generalized and reintroduced under a primed name. + +* `extend_space_map κ` for a Markov kernel `κ : Kernel Ω E` gives instead + `hZ : HasCondDistrib Z f κ P'`, and no `hind`. +* `extend_space_map μ using P` names the measure to extend. +* `extend_space_map μ with Ω' P' f hf Z hZ hind` names what is introduced. -/ +syntax (name := extendSpaceMapTac) "extend_space_map" ppSpace term (" using " ident)? + (" with " (ppSpace colGt ident)+)? : tactic + +elab_rules : tactic + | `(tactic| extend_space $μ $[using $P?]? $[with $names?*]?) => + extendSpace .hidden μ P? ((names?.map (·.map (·.getId))).getD #[]) + | `(tactic| extend_space! $μ $[using $P?]? $[with $names?*]?) => + extendSpace .clear μ P? ((names?.map (·.map (·.getId))).getD #[]) + | `(tactic| extend_space_map $μ $[using $P?]? $[with $names?*]?) => + extendSpace .map μ P? ((names?.map (·.map (·.getId))).getD #[]) + +end RDo.Tactic + +end diff --git a/RandomDo/Probability/MeasurePreserving.lean b/RandomDo/Probability/MeasurePreserving.lean new file mode 100644 index 0000000..ef3cfa9 --- /dev/null +++ b/RandomDo/Probability/MeasurePreserving.lean @@ -0,0 +1,241 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Probability.Transfer +public import Mathlib.MeasureTheory.Integral.Bochner.Basic +public import Mathlib.MeasureTheory.Measure.Real +public import Mathlib.Probability.HasCondDistrib +public import Mathlib.Probability.Independence.Basic + +set_option linter.style.header false + +/-! +# Pulling probabilistic statements back along a measure-preserving map + +For a measure-preserving map `f : Ω' → Ω` from `(Ω', P')` to `(Ω, P)`, a statement about random +variables on `Ω` is equivalent to the same statement about their compositions with `f` on `Ω'`: +laws, events and their measure, integrals, almost-everywhere statements, independence of two +random variables or of a family, conditional laws, integrability. This file collects these facts +in the forms the `transfer` tactic uses. + +* `MeasurePreserving.map_fun_comp` and the `MeasurePreserving.*_fun_comp_iff` lemmas: the + statement on `Ω'` on the left. +* The `@[transfer]` lemmas `MeasurePreserving.transfer_*`: the statement on `Ω` on the left, with + `hf` as first explicit argument, which is what `transfer` rewrites with. +* The `@[transfer_forward]` lemmas `*.comp_measurePreserving` and `*.preimage_measurePreserving`: + one-way transport of a hypothesis, for statements that do not go both ways, such as + measurability. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory + +noncomputable section + +/-! ### Pulling statements back along a measure-preserving map -/ + +namespace MeasureTheory.MeasurePreserving + +variable {Ω Ω' 𝓧 𝓨 : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + {m𝓧 : MeasurableSpace 𝓧} {m𝓨 : MeasurableSpace 𝓨} {P : Measure Ω} {P' : Measure Ω'} + {f : Ω' → Ω} {X : Ω → 𝓧} {Y : Ω → 𝓨} + +/-- The law of `X ∘ f` under `P'` is the law of `X` under `P`. -/ +lemma map_fun_comp (hf : MeasurePreserving f P' P) (hX : AEMeasurable X P) : + P'.map (fun ω ↦ X (f ω)) = P.map X := by + rw [← hf.map_eq] at hX ⊢ + exact (AEMeasurable.map_map_of_aemeasurable hX hf.measurable.aemeasurable).symm + +lemma hasLaw_fun_comp_iff (hf : MeasurePreserving f P' P) (hX : Measurable X) {ν : Measure 𝓧} : + HasLaw (fun ω ↦ X (f ω)) ν P' ↔ HasLaw X ν P where + mp h := ⟨hX.aemeasurable, by rw [← hf.map_fun_comp hX.aemeasurable]; exact h.map_eq⟩ + mpr h := h.comp hf.hasLaw + +lemma indepFun_fun_comp_iff (hf : MeasurePreserving f P' P) (hX : Measurable X) + (hY : Measurable Y) : + IndepFun (fun ω ↦ X (f ω)) (fun ω ↦ Y (f ω)) P' ↔ IndepFun X Y P := by + simp only [indepFun_iff_measure_inter_preimage_eq_mul] + refine forall₄_congr fun s t hs ht ↦ ?_ + change P' (f ⁻¹' (X ⁻¹' s) ∩ f ⁻¹' (Y ⁻¹' t)) + = P' (f ⁻¹' (X ⁻¹' s)) * P' (f ⁻¹' (Y ⁻¹' t)) ↔ _ + rw [← Set.preimage_inter, hf.measure_preimage ((hX hs).inter (hY ht)).nullMeasurableSet, + hf.measure_preimage (hX hs).nullMeasurableSet, hf.measure_preimage (hY ht).nullMeasurableSet] + +lemma hasCondDistrib_fun_comp_iff (hf : MeasurePreserving f P' P) (hX : Measurable X) + (hY : Measurable Y) {κ : Kernel 𝓧 𝓨} : + HasCondDistrib (fun ω ↦ Y (f ω)) (fun ω ↦ X (f ω)) κ P' ↔ HasCondDistrib Y X κ P := by + unfold HasCondDistrib + rw [hf.map_fun_comp hX.aemeasurable] + exact hf.hasLaw_fun_comp_iff (hX.prodMk hY) + +/-! ### The `@[transfer]` lemmas: from the old space to the new one + +The same facts, stated with the old space on the left and `hf` as the first explicit argument, +which is what the `transfer` tactic rewrites with. -/ + +@[transfer] +lemma transfer_map (hf : MeasurePreserving f P' P) (hX : AEMeasurable X P) : + P.map X = P'.map (fun ω ↦ X (f ω)) := + (hf.map_fun_comp hX).symm + +@[transfer] +lemma transfer_measure (hf : MeasurePreserving f P' P) {s : Set Ω} (hs : NullMeasurableSet s P) : + P s = P' (f ⁻¹' s) := + (hf.measure_preimage hs).symm + +@[transfer] +lemma transfer_real (hf : MeasurePreserving f P' P) {s : Set Ω} (hs : NullMeasurableSet s P) : + P.real s = P'.real (f ⁻¹' s) := by + simp only [measureReal_def, hf.measure_preimage hs] + +@[transfer] +lemma transfer_hasLaw (hf : MeasurePreserving f P' P) (hX : Measurable X) {ν : Measure 𝓧} : + HasLaw X ν P ↔ HasLaw (fun ω ↦ X (f ω)) ν P' := + (hf.hasLaw_fun_comp_iff hX).symm + +@[transfer] +lemma transfer_indepFun (hf : MeasurePreserving f P' P) (hX : Measurable X) (hY : Measurable Y) : + IndepFun X Y P ↔ IndepFun (fun ω ↦ X (f ω)) (fun ω ↦ Y (f ω)) P' := + (hf.indepFun_fun_comp_iff hX hY).symm + +@[transfer] +lemma transfer_iIndepFun {ι : Type*} {β : ι → Type*} [∀ i, MeasurableSpace (β i)] + (hf : MeasurePreserving f P' P) {X : ∀ i, Ω → β i} (hX : ∀ i, Measurable (X i)) : + iIndepFun X P ↔ iIndepFun (fun i ω ↦ X i (f ω)) P' := by + simp only [iIndepFun_iff_measure_inter_preimage_eq_mul] + refine forall_congr' fun S ↦ forall_congr' fun sets ↦ imp_congr_right fun hsets ↦ ?_ + have hm : ∀ i ∈ S, MeasurableSet (X i ⁻¹' sets i) := fun i hi ↦ hX i (hsets i hi) + rw [← hf.measure_preimage (S.measurableSet_biInter hm).nullMeasurableSet, + Finset.prod_congr rfl fun i hi ↦ (hf.measure_preimage (hm i hi).nullMeasurableSet).symm] + simp only [Set.preimage_iInter₂, Set.preimage_preimage] + +@[transfer] +lemma transfer_hasCondDistrib (hf : MeasurePreserving f P' P) (hX : Measurable X) + (hY : Measurable Y) {κ : Kernel 𝓧 𝓨} : + HasCondDistrib Y X κ P ↔ HasCondDistrib (fun ω ↦ Y (f ω)) (fun ω ↦ X (f ω)) κ P' := + (hf.hasCondDistrib_fun_comp_iff hX hY).symm + +@[transfer] +lemma transfer_integral {G : Type*} [NormedAddCommGroup G] [NormedSpace ℝ G] + (hf : MeasurePreserving f P' P) {g : Ω → G} (hg : AEStronglyMeasurable g P) : + ∫ ω, g ω ∂P = ∫ ω, g (f ω) ∂P' := by + rw [← hf.map_eq] at hg ⊢ + exact integral_map hf.measurable.aemeasurable hg + +@[transfer] +lemma transfer_lintegral (hf : MeasurePreserving f P' P) {g : Ω → ENNReal} + (hg : AEMeasurable g P) : + ∫⁻ ω, g ω ∂P = ∫⁻ ω, g (f ω) ∂P' := by + rw [← hf.map_eq] at hg ⊢ + exact lintegral_map' hg hf.measurable.aemeasurable + +@[transfer] +lemma transfer_ae (hf : MeasurePreserving f P' P) {p : Ω → Prop} + (hp : NullMeasurableSet {ω | p ω} P) : + (∀ᵐ ω ∂P, p ω) ↔ ∀ᵐ ω ∂P', p (f ω) := by + rw [ae_iff, ae_iff, ← hf.measure_preimage (s := {ω | ¬ p ω}) hp.compl, Set.preimage_ofPred_eq] + +@[transfer] +lemma transfer_ae_eq (hf : MeasurePreserving f P' P) {X Y : Ω → 𝓧} + (h : NullMeasurableSet {ω | X ω = Y ω} P) : + X =ᵐ[P] Y ↔ (fun ω ↦ X (f ω)) =ᵐ[P'] fun ω ↦ Y (f ω) := + hf.transfer_ae h + +@[transfer] +lemma transfer_integrable {G : Type*} [NormedAddCommGroup G] (hf : MeasurePreserving f P' P) + {g : Ω → G} (hg : AEStronglyMeasurable g P) : + Integrable g P ↔ Integrable (fun ω ↦ g (f ω)) P' := + (hf.integrable_comp hg).symm + +end MeasureTheory.MeasurePreserving + +/-! ### Forward transport of hypotheses + +A hypothesis about the old space gives one about the new space. These are the +`@[transfer_forward]` lemmas: the hypothesis first, then `hf`, then side conditions. -/ + +section Forward + +variable {Ω Ω' 𝓧 𝓨 : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} + {m𝓧 : MeasurableSpace 𝓧} {m𝓨 : MeasurableSpace 𝓨} {P : Measure Ω} {P' : Measure Ω'} + {f : Ω' → Ω} {X : Ω → 𝓧} {Y : Ω → 𝓨} + +/-- A measure-preserving map out of the old space, composed with the map to the new one: this +transports the map of a previous `extend_space` along a new one. -/ +@[transfer_forward] +lemma MeasureTheory.MeasurePreserving.comp_measurePreserving {𝓩 : Type*} + {m𝓩 : MeasurableSpace 𝓩} {ν : Measure 𝓩} {g : Ω → 𝓩} (h : MeasurePreserving g P ν) + (hf : MeasurePreserving f P' P) : + MeasurePreserving (fun ω ↦ g (f ω)) P' ν := + h.comp hf + +@[transfer_forward] +lemma Measurable.comp_measurePreserving (hX : Measurable X) (hf : MeasurePreserving f P' P) : + Measurable fun ω ↦ X (f ω) := + hX.comp hf.measurable + +@[transfer_forward] +lemma AEMeasurable.comp_measurePreserving (hX : AEMeasurable X P) + (hf : MeasurePreserving f P' P) : + AEMeasurable (fun ω ↦ X (f ω)) P' := + hX.comp_quasiMeasurePreserving hf.quasiMeasurePreserving + +attribute [transfer_forward] MeasureTheory.AEStronglyMeasurable.comp_measurePreserving + +@[transfer_forward] +lemma MeasurableSet.preimage_measurePreserving {s : Set Ω} (hs : MeasurableSet s) + (hf : MeasurePreserving f P' P) : + MeasurableSet (f ⁻¹' s) := + hf.measurable hs + +@[transfer_forward] +lemma MeasureTheory.NullMeasurableSet.preimage_measurePreserving {s : Set Ω} + (hs : NullMeasurableSet s P) (hf : MeasurePreserving f P' P) : + NullMeasurableSet (f ⁻¹' s) P' := + hs.preimage hf.quasiMeasurePreserving + +@[transfer_forward] +lemma ProbabilityTheory.HasLaw.comp_measurePreserving {ν : Measure 𝓧} (hX : HasLaw X ν P) + (hf : MeasurePreserving f P' P) : + HasLaw (fun ω ↦ X (f ω)) ν P' := + hX.comp hf.hasLaw + +/-- A conditional law pulls back along a measure-preserving map. -/ +@[transfer_forward] +lemma ProbabilityTheory.HasCondDistrib.comp_measurePreserving {κ : Kernel 𝓧 𝓨} + (h : HasCondDistrib Y X κ P) (hf : MeasurePreserving f P' P) : + HasCondDistrib (fun ω ↦ Y (f ω)) (fun ω ↦ X (f ω)) κ P' := by + have hX := h.aemeasurable_fst + unfold HasCondDistrib at h ⊢ + rw [hf.map_fun_comp hX] + exact h.comp hf.hasLaw + +@[transfer_forward] +lemma ProbabilityTheory.IndepFun.comp_measurePreserving (h : IndepFun X Y P) + (hf : MeasurePreserving f P' P) (hX : Measurable X) (hY : Measurable Y) : + IndepFun (fun ω ↦ X (f ω)) (fun ω ↦ Y (f ω)) P' := + (hf.indepFun_fun_comp_iff hX hY).2 h + +@[transfer_forward] +lemma ProbabilityTheory.iIndepFun.comp_measurePreserving {ι : Type*} {β : ι → Type*} + [∀ i, MeasurableSpace (β i)] {X : ∀ i, Ω → β i} (h : iIndepFun X P) + (hf : MeasurePreserving f P' P) (hX : ∀ i, Measurable (X i)) : + iIndepFun (fun i ω ↦ X i (f ω)) P' := + (hf.transfer_iIndepFun hX).1 h + +@[transfer_forward] +lemma MeasureTheory.Integrable.comp_measurePreserving {G : Type*} [NormedAddCommGroup G] + {g : Ω → G} (hg : Integrable g P) (hf : MeasurePreserving f P' P) : + Integrable (fun ω ↦ g (f ω)) P' := + (hf.integrable_comp hg.aestronglyMeasurable).2 hg + +end Forward + +end + +end diff --git a/RandomDo/Probability/Record.lean b/RandomDo/Probability/Record.lean new file mode 100644 index 0000000..d7257e1 --- /dev/null +++ b/RandomDo/Probability/Record.lean @@ -0,0 +1,115 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Monad.Notation +public import RandomDo.Monad.Instances +public import RandomDo.Tactic.IsMarkov +public meta import Lean.Elab.Do + +/-! +# Recording a value of an `rdo` program as a random variable + +The trace of an `rdo` program (`RandomDo.Probability.Trace`) has one coordinate per `←`. Values the +program computes without drawing them — a `let mut` accumulator, an intermediate quantity — are +therefore *not* coordinates: they can be spoken about only through whatever the program did draw. + +`record` is the annotation that changes that. Writing + +``` +record N +``` + +on a line of an `rdo` block turns `N` into a coordinate of the trace from that point on, hence into +a random variable one can state laws about and condition on. It denotes the one-point distribution +at `N`, so it never changes what the program computes (`record_bind`); all it does is put a `←` in +the program text where there was none, which is exactly what the trace is built from. + +## Main definitions + +* `RDo.record x`: the one-point distribution at `x`, marked so that the trace keeps a coordinate + for it. +* the `record x, y, …` `doElem`, sugar for `x ← RDo.record x` on each of them. + +## Main results + +* `RDo.record_bind`, `RDo.record_bind_of_isMarkov`: drawing from `record x` and carrying on is the + same as carrying on with `x`. Recording never changes the measure a program denotes; the second + form is the one to feed `simp (disch := is_markov)` to erase every `record` from a program. + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory +open MeasurableSpacePure MeasurableSpaceBind + +namespace RDo + +universe u + +variable {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + +/-- `record x` is the one-point distribution at `x`. Drawing from it changes nothing — +`record x = mPure x` — but the `←` it is drawn with gives the program's trace a coordinate holding +the value of `x` at that point, so that `x` becomes a random variable of the program. + +Inside an `rdo` block, write `record x` (sugar for `x ← RDo.record x`). -/ +noncomputable def record (x : α) : Measure α := Measure.dirac x + +lemma record_eq_mPure (x : α) : record x = mPure x := rfl + +instance (x : α) : IsProbabilityMeasure (record x) := by + rw [record]; infer_instance + +instance : IsMarkov (record : α → Measure α) := + ⟨Measure.measurable_dirac, fun _ ↦ inferInstance⟩ + +/-- **Recording is transparent.** Drawing `x` from `record x` and carrying on is the same as +carrying on with `x`: inserting a `record` never changes the measure an `rdo` program denotes. +Stated with `Measure.bind`, the `simp` normal form of `>>=ₘ` at `Measure`; the `>>=ₘ` form is +`record_bind_of_isMarkov`. -/ +@[simp] +lemma record_bind (x : α) {f : α → Measure β} (hf : Measurable f) : (record x).bind f = f x := + Measure.dirac_bind hf x + +/-- `record_bind` in the shape a `simp` call can use with `is_markov` as its discharger: +`simp (disch := is_markov) only [record_bind_of_isMarkov]` erases every `record` from a program. -/ +lemma record_bind_of_isMarkov (x : α) (f : α → Measure β) (h : IsMarkov f) : + record x >>=ₘ f = f x := + Measure.dirac_bind h.measurable x + +end RDo + +end + +public meta section + +open Lean Lean.Parser Lean.Parser.Term Lean.Elab Lean.Elab.Do + +namespace RDo + +/-- `record x, y, …` inside an `rdo` block records the variables `x`, `y`, … as random variables of +the program at that point: each becomes a coordinate of the program's trace, so that laws and +conditional laws can be stated about it. It is sugar for `x ← RDo.record x`, and changes nothing +about what the program computes — see `RDo.record_bind`. + +The variables must be `let mut` variables, since each is reassigned from its own recording. To +record an arbitrary expression, write `let y ← RDo.record e` instead. + +Note that this declaration reserves `record` as a token, so `RDo.record` has to be written +qualified (or escaped as `«record»`) once this module is imported. -/ +syntax (name := rdoRecord) "record " Lean.Parser.ident,+ : doElem + +/-- Expand `record x, y` into one reassignment per variable. -/ +@[macro rdoRecord] def expandRdoRecord : Macro := fun stx => do + let `(doElem| record $xs:ident,*) := stx | Macro.throwUnsupported + let items ← xs.getElems.mapM fun x => `(doSeqItem| $x:ident ← RDo.record $x) + `(doElem| do $items*) + +end RDo + +end diff --git a/RandomDo/Probability/Tactic.lean b/RandomDo/Probability/Tactic.lean new file mode 100644 index 0000000..76a8fbd --- /dev/null +++ b/RandomDo/Probability/Tactic.lean @@ -0,0 +1,552 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Probability.Trace +public import RandomDo.Tactic.Elab +public meta import Lean.Elab.Tactic.Basic + +/-! +# The `rdo_trace` and `rdo_peel` tactics + +`rdo_trace` builds the trace of an `rdo` program: it walks the same program tree that `is_markov` +walks, applying at each node the combinator of `RandomDo.Probability.Trace` for that construct, and +adds `h : RDo.HasTrace prog P out` to the context with `P` and `out` computed. The measurability +and `IsMarkov` obligations the combinators leave behind become goals, and are attacked with +`is_markov` and `fun_prop`. + +`rdo_peel` then reads the probabilistic statements off such a hypothesis: the law of the first +draw, the conditional law of each later draw given the draws before it, the law of the program's +result — and, for every draw whose kernel does not depend on the draws before it, an +unconditional law together with an independence statement. + +Both keep the local context readable: an `rdo` program's text ends up inside its trace kernel, and +every statement a peeled trace adds mentions the trace measure built from those kernels. Anything +too wide to print is given a local definition of its own — `κ₁`, `κ₂`, … for the kernels, `P` for +the trace measure — so the hypotheses stay one line each. + +-/ + +public meta section + +open Lean Lean.Meta Lean.Elab Lean.Elab.Tactic +open MeasureTheory ProbabilityTheory + +namespace RDo.Tactic + +initialize registerTraceClass `rdo_trace + +/-! ## Small expression utilities -/ + +/-- Beta-reduce, and reduce `(a, b).1` and `(a, b).2`. This is what turns the continuation of a +`mBind`, which the tactic builds by substituting a pair for the parameter, back into a readable +term. -/ +partial def betaProj (e : Expr) : MetaM Expr := do + let step (e : Expr) : Expr := + let e := e.headBeta + if e.isAppOfArity ``Prod.fst 3 then + let p := e.appArg! + if p.isAppOfArity ``Prod.mk 4 then p.appFn!.appArg! else e + else if e.isAppOfArity ``Prod.snd 3 then + let p := e.appArg! + if p.isAppOfArity ``Prod.mk 4 then p.appArg! else e + else if e.isAppOfArity ``DFunLike.coe 6 then + -- `(Kernel.const γ ν) a = ν` and `(markovKernel prog h) a = prog a`, both by `rfl`. + let as := e.getAppArgs + let f := as[4]! + if f.isAppOf ``ProbabilityTheory.Kernel.const then f.appArg! + else if f.isAppOfArity ``markovKernel 6 then mkApp f.appFn!.appArg! as[5]! + else e + else e + let e ← Meta.transform e (post := fun e ↦ return .done (step e)) + Meta.transform e (post := fun e ↦ return .done (step e)) + +/-! ### Keeping the context readable + +An `rdo` program's own text ends up inside its trace kernel, and every hypothesis a peeled trace +adds mentions the trace measure built from those kernels. Printed in full, once per hypothesis, +that is unreadable. So the kernels that are big are given local definitions of their own, `κ₁`, +`κ₂`, …, and the trace measure a local definition `P`; the hypotheses then mention only those. +-/ + +/-- Wider than this, printed, and a term is pulled out into a local definition rather than shown in +every hypothesis that mentions it. Printed width is the right measure here rather than the number +of subterms: a kernel's implicit type arguments are large but are not shown. -/ +def abstractionWidth : Nat := 60 + +/-- Is `e` too wide to print in every hypothesis? -/ +def tooWide (e : Expr) : MetaM Bool := do + return (← ppExpr e).pretty.length > abstractionWidth + +/-- Unfold a head that is a local `let` variable — the ones introduced below to keep the context +readable — so that the shape of a kernel stays visible to the tactic. -/ +partial def unfoldLetHead (e : Expr) : MetaM Expr := do + let .fvar fid := e | return e + match ← fid.getDecl with + | .ldecl (value := v) .. => unfoldLetHead v + | _ => return e + +/-- Subscript digits, for naming the definitions `κ₁`, `κ₂`, … -/ +def subscript (i : Nat) : String := + (toString i).map fun c => Char.ofNat (0x2080 + (c.toNat - '0'.toNat)) + +/-- Substitute the terms of `subst` throughout `e`. -/ +def substTerms (subst : Array (Expr × Expr)) (e : Expr) : Expr := + if subst.isEmpty then e + else e.replace fun s => (subst.find? fun p ↦ p.1 == s).map Prod.snd + +/-- Apply the instance arguments a term is still waiting for. `mkAppM` leaves those that come +after the last explicit argument it was given. -/ +partial def saturateInstances (e : Expr) : MetaM Expr := do + match ← whnf (← inferType e) with + | .forallE _ d _ bi => + if bi.isInstImplicit then saturateInstances (mkApp e (← synthInstance d)) else return e + | _ => return e + +/-- The last `n` arguments of an application. -/ +def lastArgs (e : Expr) (n : Nat) : Array Expr := + let as := e.getAppArgs + as[(as.size - n)...*] + +/-- The kernels at the leaves of the `⊗ₖ` spine of a trace kernel: one per `←` of the program. -/ +partial def compProdLeaves (P : Expr) : Array Expr := + if P.isAppOf ``ProbabilityTheory.Kernel.compProd then + let as := lastArgs P 2 + compProdLeaves as[0]! ++ compProdLeaves as[1]! + else #[P] + +/-- Introduce `base := value` as a local definition, reusing one already in the context if it has +that same value. -/ +def localDefFor (g : MVarId) (base : Name) (value : Expr) : MetaM (MVarId × Expr) := + g.withContext do + let value ← instantiateMVars value + for d in ← getLCtx do + if let .ldecl (fvarId := fid) (value := v) .. := d then + if (← instantiateMVars v) == value then return (g, .fvar fid) + let name := (← getLCtx).getUnusedName base + let g ← g.define name (← inferType value) value + let (fid, g) ← g.intro1P + return (g, .fvar fid) + +/-- Give every kernel of the `⊗ₖ` spine of `P` that is too big to read a local definition of its +own. Returns the new goal and the substitution to apply to the statements being added. -/ +def abstractKernels (g : MVarId) (P : Expr) : MetaM (MVarId × Array (Expr × Expr)) := do + let mut g := g + let mut subst : Array (Expr × Expr) := #[] + let mut i := 1 + for κ in compProdLeaves (← instantiateMVars P) do + if (← g.withContext <| tooWide κ) && !subst.any (·.1 == κ) then + let (g', fv) ← localDefFor g (Name.mkSimple ("κ" ++ subscript i)) κ + g := g' + subst := subst.push (κ, fv) + i := i + 1 + return (g, subst) + +/-- `prog`, `P` and `out` of a proof of `RDo.HasTrace prog P out`. -/ +def traceParts (h : Expr) : MetaM (Expr × Expr × Expr) := do + let ty ← instantiateMVars (← inferType h) + unless ty.isAppOfArity ``HasTrace 9 do + throwError "rdo_trace: not a `HasTrace` statement: {ty}" + let as := lastArgs ty 3 + return (as[0]!, as[1]!, as[2]!) + +/-- Try to read `prog` as the coercion of a `Kernel`, so that the trace kernel is the user's own +kernel rather than a wrapper. -/ +def asKernel? (prog : Expr) : MetaM (Option Expr) := do + let (γ, α) ← forallBoundedTelescope (← inferType prog) (some 1) fun cs body ↦ do + return (← inferType cs[0]!, (← whnfR body).getAppArgs[0]!) + let κ ← mkFreshExprMVar (← saturateInstances (← mkAppM ``ProbabilityTheory.Kernel #[γ, α])) + unless ← isDefEq (← mkAppM ``DFunLike.coe #[κ]) prog do return none + let κ ← instantiateMVars κ + if κ.hasExprMVar then return none + unless (← trySynthInstance (← mkAppM ``IsMarkovKernel #[κ])) matches .some _ do return none + return some κ + +/-! ## Building the trace -/ + +/-- The tactic's state: the side conditions produced so far. -/ +abbrev TraceM := StateRefT (Array MVarId) MetaM + +/-- Record a side condition and return the proof term standing for it. -/ +def sideGoal (type : Expr) : TraceM Expr := do + let m ← mkFreshExprSyntheticOpaqueMVar type + modify (·.push m.mvarId!) + return m + +/-- Look through the definitions heading a program until one of the constructs the tactic knows +about shows up. -/ +partial def unfoldProgram (prog : Expr) (fuel : Nat) : MetaM (Option Expr) := do + if fuel == 0 then return none + let some n ← programHeadDef? prog | return none + -- Never unfold the `Kernel` coercion: a kernel is a leaf, not a definition to look through. + if n == ``DFunLike.coe then return none + let prog' ← deltaExpand prog (· == n) + if prog' == prog then return none + match ← shapeOf prog' with + | .leaf | .const => unfoldProgram prog' (fuel - 1) + | _ => return some prog' + +/-- Look through the definitions heading a program until an `rdo` construct — not merely a known +shape — shows up. This is what tells an `rdo` program apart from a distribution such as +`gaussianReal`, whose body happens to be an `ite`. -/ +partial def unfoldToRdoProgram (prog : Expr) (fuel : Nat) : MetaM (Option Expr) := do + if fuel == 0 then return none + let some n ← programHeadDef? prog | return none + if n == ``DFunLike.coe then return none + let prog' ← deltaExpand prog (· == n) + if prog' == prog then return none + match ← shapeOf prog' with + | .mBind | .mPure | .forIn .. | .breakRunK => return some prog' + | .leaf | .const => unfoldToRdoProgram prog' (fuel - 1) + | _ => return none + +/-- Build the trace of `prog : γ → Measure β`. + +`depth` counts how many trace coordinates the ambient parameter `γ` has accumulated from the binds +above; it is what licenses the `HasTrace.prodMkRight` step, which is the step that records in the +*type* of the trace kernel that a draw does not read the draws before it. + +`root` marks the program the user asked about. A subprogram carrying an `IsMarkov` instance is a +leaf — that is the granularity knob — but the program at the root is the one being traced, so +there its own instance must not stop the traversal. -/ +partial def traceCore (prog : Expr) (depth : Nat) (fuel : Nat) (root : Bool := false) : + TraceM Expr := do + let prog ← instantiateMVars prog + if depth > 0 then + if let some h ← tryWeaken prog depth fuel then return h + if let some h ← tryRecord prog then return h + let shape ← shapeOf prog + let prf ← withTraceNode `rdo_trace (fun _ ↦ return m!"{shape}: {prog}") do + match shape with + | .mPure => tracePure prog + | .mBind => traceBind prog depth fuel + | _ => traceLeaf prog depth fuel root + normalizeOut prf +where + /-- Reduce the projections the combinators pile up in the readout, so that the trace stays + readable and the continuations built from it stay small. -/ + normalizeOut (prf : Expr) : TraceM Expr := do + let (prog, P, out) ← traceParts prf + let out' ← betaProj out + let P' ← betaProj P + if out' == out && P' == P then return prf + mkExpectedTypeHint prf (← mkAppM ``HasTrace #[prog, P', out']) + /-- `fun q : Γ × Ω ↦ prog' q.1`: a subprogram that does not read the last draw. -/ + tryWeaken (prog : Expr) (depth fuel : Nat) : TraceM (Option Expr) := do + let dom ← whnfR (← inferType prog).bindingDomain! + unless dom.isAppOfArity ``Prod 2 do return none + let Γ := dom.appFn!.appArg! + let Ω := dom.appArg! + let prog'? ← withLocalDeclD `c Γ fun c ↦ withLocalDeclD `w Ω fun w ↦ do + let body ← betaProj (mkApp prog (← mkAppM ``Prod.mk #[c, w])) + if body.containsFVar w.fvarId! then return none + return some (← mkLambdaFVars #[c] body) + let some prog' := prog'? | return none + trace[rdo_trace] "weakening: the draw does not read the last coordinate" + let h ← traceCore prog' (depth - 1) fuel + let prf ← mkAppM ``HasTrace.prodMkRight #[Ω, h] + let (prog'', _, _) ← traceParts prf + unless ← isDefEq prog'' prog do return none + return some prf + /-- A value the program marked with `record`: a coordinate of its own, holding a deterministic + function of everything drawn before it. -/ + tryRecord (prog : Expr) : TraceM (Option Expr) := do + let progE ← etaExpand (← whnfR prog) + let f? ← lambdaBoundedTelescope progE 1 fun cs body ↦ do + let body ← whnfR body + unless body.isAppOfArity ``RDo.record 3 do return none + return some (← mkLambdaFVars cs (← betaProj body.appArg!)) + let some f := f? | return none + trace[rdo_trace] "recorded value" + let hf ← sideGoal (← mkAppM ``Measurable #[f]) + return some (← mkAppM ``HasTrace.record #[hf]) + /-- `return e`. -/ + tracePure (prog : Expr) : TraceM Expr := do + let prog ← etaExpand (← whnfR prog) + let g ← lambdaBoundedTelescope prog 1 fun cs body ↦ do + let body ← whnfR body + unless body.isAppOfArity ``MeasurableSpacePure.mPure 5 do + throwError "rdo_trace: expected `mPure`, got {body}" + mkLambdaFVars cs body.getAppArgs[4]! + let hg ← sideGoal (← mkAppM ``Measurable #[g]) + mkAppM ``HasTrace.pure #[hg] + /-- `let x ← p; q`. -/ + traceBind (prog : Expr) (depth fuel : Nat) : TraceM Expr := do + let progE ← etaExpand (← whnfR prog) + let (p, cont) ← lambdaBoundedTelescope progE 1 fun cs body ↦ do + let c := cs[0]! + let body ← whnfR body + unless body.isAppOfArity ``MeasurableSpaceBind.mBind 8 do + throwError "rdo_trace: expected `mBind`, got {body}" + let as := body.getAppArgs + let α := as[2]! + -- Eta-reduce, so that a named subprogram is recognised as itself and its `IsMarkov` + -- instance — the granularity knob — can stop the traversal there. + let p := (← mkLambdaFVars #[c] as[6]!).eta + let kAbs ← mkLambdaFVars #[c] as[7]! + let γ ← inferType c + let cont ← withLocalDeclD `q (← mkAppM ``Prod #[γ, α]) fun q ↦ do + let q1 ← mkAppM ``Prod.fst #[q] + let q2 ← mkAppM ``Prod.snd #[q] + mkLambdaFVars #[q] (← betaProj (mkAppN kAbs #[q1, q2])) + return (p, cont) + -- `let x ← p; return f x` adds no coordinate: only the readout changes. + if let some f ← pureBody? cont then + trace[rdo_trace] "deterministic tail" + let hP ← traceCore p depth fuel + let hf ← sideGoal (← mkAppM ``Measurable #[f]) + return ← mkAppM ``HasTrace.bindPure #[hP, hf] + let hP ← traceCore p depth fuel + let (_, _, out) ← traceParts hP + let Ω := (← whnfR (← inferType out)).bindingDomain!.appArg! + let γ := (← whnfR (← inferType prog)).bindingDomain! + let cont' ← withLocalDeclD `q (← mkAppM ``Prod #[γ, Ω]) fun q ↦ do + let q1 ← mkAppM ``Prod.fst #[q] + mkLambdaFVars #[q] (← mkAppM' cont #[← mkAppM ``Prod.mk #[q1, ← mkAppM' out #[q]]]) + let hQ ← traceCore cont' (depth + 1) fuel + let hm ← sideGoal (← mkAppM ``IsMarkov #[cont]) + let hcont ← mkAppOptM ``IsMarkov.measurable #[none, none, none, none, cont, hm] + mkAppM ``HasTrace.bind #[hcont, hP, hQ] + /-- `fun q ↦ mPure (f q)`, if that is what `cont` is. -/ + pureBody? (cont : Expr) : TraceM (Option Expr) := do + let contE ← etaExpand (← whnfR cont) + lambdaBoundedTelescope contE 1 fun qs body ↦ do + let body ← whnfR body + unless body.isAppOfArity ``MeasurableSpacePure.mPure 5 do return none + return some (← mkLambdaFVars qs body.getAppArgs[4]!) + /-- A leaf: a distribution, or a subprogram treated as one atomic draw. -/ + traceLeaf (prog : Expr) (depth fuel : Nat) (root : Bool) : TraceM Expr := do + -- At the root, look inside the program before considering it atomic. + if root then + if let some prog' ← unfoldToRdoProgram prog fuel then + trace[rdo_trace] "unfolded the program under trace to {prog'}" + let h ← traceCore prog' depth (fuel - 1) + let (_, P, out) ← traceParts h + return ← mkExpectedTypeHint h (← mkAppM ``HasTrace #[prog, P, out]) + if let some κ ← asKernel? prog then + trace[rdo_trace] "leaf: the kernel {κ} itself" + return ← mkAppM ``HasTrace.sample #[κ] + -- A constant family is a `Kernel.const`, which reads better than the generic wrapper. + if (← shapeOf prog) matches .const then + let μ? ← lambdaBoundedTelescope (← etaExpand (← whnfR prog)) 1 fun cs body ↦ do + let body ← betaProj body + return if body.containsFVar cs[0]!.fvarId! then none else some body + if let some μ := μ? then + if (← trySynthInstance (← mkAppM ``IsProbabilityMeasure #[μ])) matches .some _ then + let γ := (← whnfR (← inferType prog)).bindingDomain! + trace[rdo_trace] "leaf: the constant kernel at {μ}" + let prf ← mkAppM ``HasTrace.sample #[← mkAppM ``ProbabilityTheory.Kernel.const #[γ, μ]] + let (prog', _, _) ← traceParts prf + if ← isDefEq prog' prog then return prf + -- Something already known to be Markov? Then it is one atomic draw. + let g ← mkFreshExprSyntheticOpaqueMVar (← mkAppM ``IsMarkov #[prog]) + let leftover ← closeLeaf g.mvarId! + if ← g.mvarId!.isAssigned then + trace[rdo_trace] "leaf: an atomic draw" + modify (· ++ leftover.toArray) + return ← mkAppM ``HasTrace.leaf #[prog, g] + -- Otherwise look through the definition heading the program and carry on. Only a genuine + -- leaf is unfolded: `ite`, `for` and friends are constructs, not definitions to look through. + let unfolded? ← + if (← shapeOf prog) matches .leaf | .const then unfoldProgram prog fuel else pure none + if let some prog' := unfolded? then + trace[rdo_trace] "unfolded to {prog'}" + let h ← traceCore prog' depth (fuel - 1) + let (_, P, out) ← traceParts h + return ← mkExpectedTypeHint h (← mkAppM ``HasTrace #[prog, P, out]) + trace[rdo_trace] "leaf: `IsMarkov` handed back to the user" + modify (·.push g.mvarId!) + mkAppM ``HasTrace.leaf #[prog, g] + +/-- Discharge a side condition `rdo_trace` produced, leaving whatever it cannot prove. -/ +def dischargeSide (g : MVarId) : TacticM (List MVarId) := do + if ← g.isAssigned then return [] + let ty ← instantiateMVars (← g.getType) + let tac ← + if ty.isAppOfArity ``IsMarkov 5 then `(tactic| try is_markov) + else `(tactic| try (fun_prop (disch := measurability))) + Lean.Elab.Tactic.run g (evalTactic tac) + +/-- `rdo_trace prog` computes the trace of the `rdo` program `prog` — a family of measures, or a +single measure — and adds `htrace : RDo.HasTrace prog P out` to the context, with the trace kernel +`P` and the readout `out` built from the shape of the program: one `Kernel.compProd` factor per +`←`. + +* `rdo_trace prog with h` names the hypothesis `h`. +* `rdo_trace prog (fuel := n)` looks through up to `n` definitions heading the program. + +The measurability and `IsMarkov` obligations the construction leaves behind are attacked with +`is_markov` and `fun_prop (disch := measurability)`, and whatever survives is handed back. + +A program's own text ends up inside its trace kernel, so any kernel too wide to print is given a +local definition `κ₁`, `κ₂`, … of its own and the hypothesis mentions only those. + +Use `rdo_peel` on the resulting hypothesis to get the laws of the individual draws. +`set_option trace.rdo_trace true` prints the tree of constructs the tactic walked through. -/ +syntax (name := rdoTraceTac) "rdo_trace" ppSpace term (" (" &"fuel" " := " num ")")? + (" with " ident)? : tactic + +elab_rules : tactic + | `(tactic| rdo_trace $prog $[(fuel := $fuel?)]? $[with $name?]?) => classical do + let fuel := (fuel?.map (·.getNat)).getD defaultUnfoldFuel + let name := (name?.map (·.getId)).getD `htrace + let g ← getMainGoal + let (prf, sides) ← g.withContext do + let e ← Term.elabTerm prog none + Term.synthesizeSyntheticMVarsNoPostponing + let e ← instantiateMVars e + -- A program with no parameter is a constant family over `Unit`. + let ty ← whnfR (← inferType e) + let prog ← + if ty.isForall then pure e + else withLocalDeclD `u (mkConst ``Unit) fun u ↦ mkLambdaFVars #[u] e + let (prf, sides) ← (traceCore prog 0 fuel (root := true)).run #[] + let (prog', P, out) ← traceParts prf + let prf ← + if ← isDefEq prog' prog then + mkExpectedTypeHint prf (← mkAppM ``HasTrace #[prog, P, out]) + else pure prf + return (← instantiateMVars prf, sides) + let (_, P, _) ← g.withContext (traceParts prf) + let (g, subst) ← abstractKernels g P + let type ← g.withContext do return substTerms subst (← instantiateMVars (← inferType prf)) + let (_, g) ← (← g.assert name type prf).intro1P + let mut remaining := [] + for s in sides do + remaining := remaining ++ (← dischargeSide s) + replaceMainGoal (g :: remaining) + +/-! ## Peeling the trace -/ + +/-- The measure a kernel `K` reparametrised by `ι` is constant at, when the shapes `rdo_trace` +produces make it constant: a `Kernel.const`, or a `Kernel.prodMkRight` whose base is reached +through a constant reparametrisation. Everything here is definitional, so the caller can retype a +`HasCondDistrib` at `Kernel.const _ ν` without a rewrite. -/ +partial def constValue? (K ι : Expr) : MetaM (Option Expr) := do + let K ← unfoldLetHead (← whnfR K) + if K.isAppOf ``ProbabilityTheory.Kernel.const then + return some (lastArgs K 1)[0]! + if K.isAppOf ``ProbabilityTheory.Kernel.prodMkRight then + let R := (lastArgs K 1)[0]! + let ιTy ← whnfR (← inferType ι) + let ι' ← withLocalDeclD `d ιTy.bindingDomain! fun d ↦ do + mkLambdaFVars #[d] (← betaProj (← mkAppM ``Prod.fst #[mkApp ι d])) + return ← constValue? R ι' + -- Not a recognised wrapper: constant only if the reparametrisation is. + let ιE ← etaExpand (← whnfR ι) + lambdaBoundedTelescope ιE 1 fun ds body ↦ do + let body ← betaProj body + if body.containsFVar ds[0]!.fvarId! then return none + return some (← mkAppM ``DFunLike.coe #[K, body]) + +/-- One fact produced by `rdo_peel`, with a suggested name. -/ +structure PeelFact where + /-- The name suggested for the hypothesis. -/ + suggested : Name + /-- Its proof. -/ + proof : Expr + +/-- Peel the trace kernel `P` at the parameter `c` into the law of each draw given the ones before +it, then the law of the program's result. -/ +def peelFacts (h : Expr) (c : Expr) : MetaM (Array PeelFact) := do + let (_, P, _) ← traceParts h + let mut facts : Array PeelFact := #[] + let P ← whnfR (← instantiateMVars P) + if P.isAppOf ``ProbabilityTheory.Kernel.compProd then + let as := lastArgs P 2 + let κ := as[0]! + let Q := as[1]! + facts := facts.push ⟨`law, ← saturateInstances (← mkAppM ``hasLaw_fst_compProd #[κ, Q, c])⟩ + let mut t ← saturateInstances (← mkAppM ``hasCondDistrib_snd_compProd_comap #[κ, Q, c]) + repeat + match ← observing? (mkAppM ``HasCondDistrib.comap_compProd_fst #[t]) with + | none => break + | some f => + facts := facts.push ⟨`law, f⟩ + facts := facts ++ (← indepFacts f) + t ← mkAppM ``HasCondDistrib.comap_compProd_snd #[t] + facts := facts.push ⟨`law, t⟩ + facts := facts ++ (← indepFacts t) + facts := facts.push ⟨`law_out, ← mkAppM ``HasTrace.hasLaw_out #[h, c]⟩ + facts.mapM fun f ↦ do + let prf ← instantiateMVars f.proof + let ty ← betaProj (← instantiateMVars (← inferType prf)) + return { f with proof := ← mkExpectedTypeHint prf ty } +where + /-- When the kernel of a conditional law does not depend on what it is conditioned on, add the + unconditional law and the independence statement. -/ + indepFacts (t : Expr) : MetaM (Array PeelFact) := do + let ty ← instantiateMVars (← inferType t) + unless ty.isAppOf ``HasCondDistrib do return #[] + let as := lastArgs ty 4 + let K ← whnfR as[2]! + unless K.isAppOf ``ProbabilityTheory.Kernel.comap do return #[] + let ks := lastArgs K 3 + let some ν ← constValue? ks[0]! ks[1]! | return #[] + let δ := (← whnfR (← inferType ks[1]!)).bindingDomain! + let const ← mkAppM ``ProbabilityTheory.Kernel.const #[δ, ν] + let tc ← observing? do + mkExpectedTypeHint t (← mkAppM ``HasCondDistrib #[as[0]!, as[1]!, const, as[3]!]) + let some tc := tc | return #[] + let mut out := #[] + if let some f ← observing? (mkAppM ``HasCondDistrib.hasLaw_of_const #[tc]) then + out := out.push ⟨`law_indep, f⟩ + if let some f ← observing? (mkAppM ``HasCondDistrib.indepFun_of_const #[tc]) then + out := out.push ⟨`indep, f⟩ + return out + +/-- `rdo_peel h c` reads the probabilistic content off a trace `h : RDo.HasTrace prog P out`, at +the parameter `c`: the law of the first draw, the conditional law of every later draw given the +draws before it, and the law of the program's result. For a draw whose kernel does not read the +draws before it, it also adds the unconditional law of that draw and its independence from them. + +* `rdo_peel h c with h₁ h₂ …` names the facts in the order they are produced. + +The trace measure is shared by every statement, so it is given a local definition `P` rather than +repeated; wide kernels get definitions `κ₁`, `κ₂`, … the same way (reusing any `rdo_trace` already +introduced). + +For a program with no parameter, pass `()`. -/ +syntax (name := rdoPeelTac) "rdo_peel" ppSpace ident ppSpace term + (" with " (ppSpace colGt ident)+)? : tactic + +elab_rules : tactic + | `(tactic| rdo_peel $h $c $[with $names?*]?) => classical do + let g ← getMainGoal + let (facts, P) ← g.withContext do + let hE ← Term.elabTerm h none + let (_, P, _) ← traceParts hE + let γ := (← whnfR (← inferType P)).getAppArgs[0]! + let cE ← Term.elabTerm c γ + Term.synthesizeSyntheticMVarsNoPostponing + return (← peelFacts hE (← instantiateMVars cE), ← instantiateMVars P) + -- Name the kernels that are too big to read, then the trace measure built from them: without + -- this every statement below repeats the whole program. + let (g₁, subst) ← abstractKernels g P + let types ← g₁.withContext <| facts.mapM fun f ↦ do + return substTerms subst (← instantiateMVars (← inferType f.proof)) + -- The trace measure appears in every statement, so it always gets a name of its own. + let (g₂, subst') ← do + if h : 0 < types.size then + let μ := (lastArgs types[0] 1)[0]! + if μ.isFVar then pure (g₁, #[]) + else do + let (g', fv) ← localDefFor g₁ `P μ + pure (g', #[(μ, fv)]) + else pure (g₁, #[]) + let names := (names?.map (·.map (·.getId))).getD #[] + let mut g := g₂ + let mut i := 0 + for f in facts do + let name := if h : i < names.size then names[i] else Name.mkSimple s!"{f.suggested}{i + 1}" + let prf ← instantiateMVars f.proof + g := (← (← g.assert name (substTerms subst' types[i]!) prf).intro1P).2 + i := i + 1 + replaceMainGoal [g] + +end RDo.Tactic + +end diff --git a/RandomDo/Probability/Thompson.lean b/RandomDo/Probability/Thompson.lean new file mode 100644 index 0000000..db0c1a4 --- /dev/null +++ b/RandomDo/Probability/Thompson.lean @@ -0,0 +1,241 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Probability.Tactic +public import RandomDo.Tactic.Elab +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg + +set_option linter.style.header false + +/-! +# Thompson sampling as random variables + +`thompson`, defined below, is an `rdo` program: a loop folding the history into +per-arm pull counts `N` and reward sums `S`, a loop drawing one Gaussian posterior sample per arm +into a vector `θ`, and `return argmax θ`. As a measure on `Fin K` it has no random variables — +there is no `θ` to talk about. This file gives it some. + +The trace `rdo_trace` finds has one coordinate per stage: + +* `ω.1 : (Fin K → ℝ) × (Fin K → ℝ)` — the sufficient statistics `(N, S)` read off the history; +* `ω.2 : Fin K → ℝ` — the vector `θ` of posterior samples; + +with readout `argmax ω.2`. Peeling it gives the three statements of `hasLaw_thompson`: the +statistics have the law of the folding stage, `θ` given them has the law of the sampling stage — +so `θ` depends on the history *only through* `(N, S)` — and the action Thompson sampling plays is +`argmax θ`. + +Both loops are traced coarsely, as one draw each: see `notes/TRACE_SEMANTICS.md` for what +per-iteration granularity inside the sampling loop would take. + +## Main definitions + +* `RDo.Thompson.stats`, `RDo.Thompson.sample`: the two stages, as `rdo` programs of their own. +* `RDo.Thompson.statsK`, `RDo.Thompson.sampleK`: the same, as Markov kernels. +* `RDo.Thompson.traceMeasure`: the joint law of the two stages given the history. + +## Main results + +* `RDo.Thompson.thompson_eq`: `thompson` *is* `stats`, then `sample`, then `argmax`. +* `RDo.Thompson.hasTrace_thompson`: the trace of `thompson`. +* `RDo.Thompson.hasLaw_thompson`: the three statements above. + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory +open MeasurableSpacePure MeasurableSpaceBind + +noncomputable section + +namespace RDo.Thompson + +variable {K n : ℕ} + +attribute[fun_prop] Measurable.ite + +/-! ## The two stages -/ + +/-- Stage one: fold the history into the per-arm pull counts `N` (started at one) and reward +sums `S`. It draws nothing. -/ +def stats (hist : Vector (Fin K × ℝ) n) : Measure ((Fin K → ℝ) × (Fin K → ℝ)) := rdo + let mut N : Fin K → ℝ := fun _ ↦ 1 + let mut S : Fin K → ℝ := fun _ ↦ 0 + for (a, r) in hist rdo + N := fun j ↦ if j = a then N j + 1 else N j + S := fun j ↦ if j = a then S j + r else S j + return (N, S) + +instance : IsMarkov (stats (K := K) (n := n)) := by is_markov + +/-- Stage two: one Gaussian posterior draw per arm, with mean and variance determined by the +statistics, collected into the vector `θ`. -/ +def sample (NS : (Fin K → ℝ) × (Fin K → ℝ)) : Measure (Fin K → ℝ) := rdo + let mut θ : Fin K → ℝ := fun _ ↦ 0 + for j in List.finRange K rdo + let z ← gaussianReal (NS.2 j / NS.1 j) (Real.toNNReal (1 / NS.1 j)) + θ := fun k ↦ if k = j then z else θ k + return θ + +instance : IsMarkov (sample (K := K)) := by is_markov + +/-- Stage one as a Markov kernel from the history. -/ +def statsK : Kernel (Vector (Fin K × ℝ) n) ((Fin K → ℝ) × (Fin K → ℝ)) := + markovKernel stats inferInstance + +instance : IsMarkovKernel (statsK (K := K) (n := n)) := by unfold statsK; infer_instance + +/-- Stage two as a Markov kernel from the statistics. -/ +def sampleK : Kernel ((Fin K → ℝ) × (Fin K → ℝ)) (Fin K → ℝ) := markovKernel sample inferInstance + +instance : IsMarkovKernel (sampleK (K := K)) := by unfold sampleK; infer_instance + +@[simp] lemma statsK_apply (hist : Vector (Fin K × ℝ) n) : statsK hist = stats hist := rfl + +@[simp] lemma sampleK_apply (NS : (Fin K → ℝ) × (Fin K → ℝ)) : sampleK NS = sample NS := rfl + +/-- Thompson sampling for `K` arms with Gaussian rewards, as an `rdo` program: fold the history +into per-arm pull counts `N`, starting at `1`, and reward sums `S`; draw for each arm one sample +of the posterior `gaussianReal (S j / N j) (1 / N j)`; play the arm with the largest sample. -/ +def thompson {K n : ℕ} (hK : 0 < K) (hist : Vector (Fin K × ℝ) n) : + Measure (Fin K) := rdo + let mut N : Fin K → ℝ := fun _ ↦ 1 + let mut S : Fin K → ℝ := fun _ ↦ 0 + for (a, r) in hist rdo + N := fun j ↦ if j = a then N j + 1 else N j + S := fun j ↦ if j = a then S j + r else S j + let mut θ : Fin K → ℝ := fun _ ↦ 0 + for j in List.finRange K rdo + let z ← gaussianReal (S j / N j) (Real.toNNReal (1 / N j)) + θ := fun k ↦ if k = j then z else θ k + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + return argmax θ + +/-- `thompson` is exactly: fold the history into `(N, S)`, draw the posterior sample `θ` given +them, play `argmax θ`. Both sides elaborate to the same two loops; all that separates them is the +`return` at the end of each stage. -/ +theorem thompson_eq (hK : 0 < K) (hist : Vector (Fin K × ℝ) n) : + haveI : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + thompson hK hist = stats hist >>=ₘ fun NS ↦ sample NS >>=ₘ fun θ ↦ mPure (argmax θ) := by + unfold thompson stats sample + simp only [Prod.mk.eta, mBind_mPure] + +/-! ## The trace -/ + +/-- The joint law of the two stages: the statistics, then the posterior sample given them. -/ +def traceMeasure (hist : Vector (Fin K × ℝ) n) : + Measure (((Fin K → ℝ) × (Fin K → ℝ)) × (Fin K → ℝ)) := + (statsK ⊗ₖ sampleK.comap Prod.snd measurable_snd) hist + +instance (hist : Vector (Fin K × ℝ) n) : IsProbabilityMeasure (traceMeasure hist) := by + unfold traceMeasure + infer_instance + +/-- **The trace of `thompson`.** Found by `rdo_trace`, which stops at `stats` and at `sample` +because each carries an `IsMarkov` instance: that is the granularity knob, and it is what keeps +the trace readable rather than exposing the two raw loops. -/ +theorem hasTrace_thompson (hK : 0 < K) : + haveI : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + HasTrace (thompson (n := n) hK) (statsK ⊗ₖ sampleK.comap Prod.snd measurable_snd) + (fun p : Vector (Fin K × ℝ) n × (((Fin K → ℝ) × (Fin K → ℝ)) × (Fin K → ℝ)) ↦ + argmax p.2.2) := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + rdo_trace (fun hist : Vector (Fin K × ℝ) n ↦ + stats hist >>=ₘ fun NS ↦ sample NS >>=ₘ fun θ ↦ mPure (argmax θ)) with h + exact h.congr fun hist ↦ thompson_eq hK hist + +/-! ## The random variables -/ + +/-- **Thompson sampling, as random variables.** On the trace space, writing `NS := ω.1` for the +sufficient statistics and `θ := ω.2` for the vector of posterior samples: + +* `NS` has the law of the folding stage; +* `θ` has, given `NS`, the law of the sampling stage — so `θ` depends on the history only through + `NS`, which is the whole content of "Thompson sampling is a function of the sufficient + statistics"; +* the action the algorithm plays is `argmax θ`. + +Everything below the `rdo_trace` line is produced by `rdo_peel`. -/ +theorem hasLaw_thompson (hK : 0 < K) (hist : Vector (Fin K × ℝ) n) : + haveI : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + HasLaw (fun ω ↦ ω.1) (stats hist) (traceMeasure hist) + ∧ HasCondDistrib (fun ω ↦ ω.2) (fun ω ↦ ω.1) sampleK (traceMeasure hist) + ∧ HasLaw (fun ω ↦ argmax ω.2) (thompson hK hist) (traceMeasure hist) := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + have h := hasTrace_thompson (n := n) hK + rdo_peel h hist with hNS hθ hout + exact ⟨hNS, hθ, hout⟩ + +/-! ## Recording a value the program only computes + +The trace above has a coordinate for each of the two *stages*, so `N` and `S` can only be spoken +about together, as `ω.1`. They are not draws — the program computes them — so nothing makes them +random variables of their own. + +`record N, S` is the annotation that does. Written on a line of the program, it turns `N` and `S` +into coordinates of the trace from that point on. It denotes the one-point distribution at each, so +the program still computes the same measure (`thompsonRecord_eq`); all it does is put a `←` in the +program text where there was none, and a `←` is what the trace is built from. +-/ + +section Record + +/-- The trace space of the annotated program: what the folding loop returned, then `N`, then `S`, +then `θ`. -/ +abbrev RecordTrace (K : ℕ) : Type := + ((Fin K → ℝ) × (Fin K → ℝ)) × ((Fin K → ℝ) × ((Fin K → ℝ) × (Fin K → ℝ))) + +/-- `thompson` again, with `N` and `S` recorded as random variables just before the sampling +loop. -/ +def thompsonRecord (hK : 0 < K) (hist : Vector (Fin K × ℝ) n) : Measure (Fin K) := rdo + let mut N : Fin K → ℝ := fun _ ↦ 1 + let mut S : Fin K → ℝ := fun _ ↦ 0 + for (a, r) in hist rdo + N := fun j ↦ if j = a then N j + 1 else N j + S := fun j ↦ if j = a then S j + r else S j + record N, S + let mut θ : Fin K → ℝ := fun _ ↦ 0 + for j in List.finRange K rdo + let z ← gaussianReal (S j / N j) (Real.toNNReal (1 / N j)) + θ := fun k ↦ if k = j then z else θ k + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + return argmax θ + +/-- **Recording is invisible to the semantics.** The annotated program is the same measure as +`thompson`; only its trace is finer. -/ +theorem thompsonRecord_eq (hK : 0 < K) (hist : Vector (Fin K × ℝ) n) : + thompsonRecord hK hist = thompson hK hist := by + unfold thompsonRecord thompson + simp (disch := is_markov) only [record_bind_of_isMarkov] + +/-- The trace of the annotated program has four coordinates: what the folding loop returned, then +`N`, then `S`, then `θ`. `N` and `S` are deterministic given the fold — `rdo_peel` reports their +conditional law as a `Kernel.deterministic` — and the sampling loop's kernel is now a function of +*those* coordinates, which is what lets one condition on `N` and `S`. -/ +example (hK : 0 < K) (hist : Vector (Fin K × ℝ) n) : True := by + have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK + rdo_trace (thompsonRecord (n := n) hK) with h + rdo_peel h hist with hfold hN hS hθ hout + -- `N` and `S`, now random variables, read off what the folding loop returned + have _ : HasCondDistrib (fun ω : RecordTrace K ↦ ω.2.1) (fun ω ↦ ω.1) _ _ := hN + have _ : HasCondDistrib (fun ω : RecordTrace K ↦ ω.2.2.1) (fun ω ↦ (ω.1, ω.2.1)) _ _ := hS + -- `θ` drawn given them + have _ : HasCondDistrib (fun ω : RecordTrace K ↦ ω.2.2.2) + (fun ω ↦ ((ω.1, ω.2.1), ω.2.2.1)) _ _ := hθ + -- and the action played is still Thompson sampling's + rw [thompsonRecord_eq hK hist] at hout + have _ : HasLaw (fun ω : RecordTrace K ↦ argmax ω.2.2.2) (thompson hK hist) _ := hout + trivial + +end Record + +end RDo.Thompson + +end + +end diff --git a/RandomDo/Probability/Trace.lean b/RandomDo/Probability/Trace.lean new file mode 100644 index 0000000..2e1f6f8 --- /dev/null +++ b/RandomDo/Probability/Trace.lean @@ -0,0 +1,440 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import RandomDo.Probability.Record +public import RandomDo.Tactic.IsMarkov +public import RandomDo.Monad.Instances +public import LeanMachineLearning.ForMathlib.Probability.HasCondDistrib + +/-! +# Trace semantics for `rdo` programs + +An `rdo` program denotes a measure, and a measure alone has no random variables: there is nothing +to condition on and nothing to be independent of. This file adds the missing layer. + +To every `rdo` program we attach a *trace*: a measurable space `Ω` recording the value drawn at +each `←` in the program, a Markov kernel `P : Kernel γ Ω` giving the joint law of those draws as +a function of the program's parameters, and a *deterministic* readout `out : γ × Ω → β` +reconstructing the program's result from its parameters and its draws. The defining property is + +`(P c).map (fun ω ↦ out (c, ω)) = prog c`, + +that is, running the program is the same as drawing a trace and reading the answer off it. +This is `HasTrace prog P out`. + +The point of the construction is that `P` is built by `Kernel.compProd`, one factor per `←` +in the program: the *shape of the program is the shape of `P`*. All the probabilistic content +then comes from two facts about `κ ⊗ₖ η`, proved once in the `Coordinates` section below: + +* the first component of the trace has law `κ` (`HasLaw.fst_compProd`); +* the second has conditional distribution `η` given the first + (`hasCondDistrib_snd_compProd`), and is *independent* of it when `η` does not read the + first (`indepFun_snd_compProd_prodMkRight`). + +Iterating those along the nesting of `⊗ₖ` reads the conditional law of every draw given the +draws before it straight off the program text. + +## Main definitions + +* `RDo.HasTrace prog P out`: `P` and `out` are a trace representation of the program `prog`. + +## Main results + +* `RDo.HasTrace.bind`: the trace of `let x ← p; q` is the composition-product of the traces of + `p` and of `q`. This is the rule that turns program structure into kernel structure. +* `RDo.HasTrace.sample`, `RDo.HasTrace.pure`, `RDo.HasTrace.bindPure`: the remaining `rdo` + constructs. +* `RDo.HasTrace.hasLaw_out`: the program's result, as a random variable on the trace space, has + the program's law. +* `ProbabilityTheory.HasLaw.fst_compProd`, `ProbabilityTheory.HasCondDistrib.fst_compProd`: laws + and conditional laws of earlier draws survive appending a further draw. +* `RDo.hasCondDistrib_snd_compProd`, `RDo.indepFun_snd_compProd_prodMkRight`: the conditional law + of the last draw given all the previous ones, and its independence in the case where the + program does not look at them. + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory +open MeasurableSpacePure MeasurableSpaceBind + +namespace RDo + +universe u + +section Prerequisites + +variable {α β δ : Type*} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace δ] + +/-- Pushing a measure through a kernel and then through a map, in the form used to peel one `←` +off a program: the first component of the product is the value the rest of the program sees. -/ +lemma map_compProd (μ : Measure α) [SFinite μ] (κ : Kernel α β) [IsSFiniteKernel κ] + {g : α × β → δ} (hg : Measurable g) : + (μ ⊗ₘ κ).map g = μ.bind fun a ↦ (κ a).map fun b ↦ g (a, b) := by + rw [Measure.compProd_eq_comp_prod, Measure.map_comp _ _ hg] + refine Measure.bind_congr_right (.of_forall fun a ↦ ?_) + rw [Kernel.map_apply _ hg, Kernel.prod_apply, Kernel.id_apply, Measure.dirac_prod, + Measure.map_map hg measurable_prodMk_left] + rfl + +/-- `Measure.bind` sees through a `Measure.map` on the left. -/ +lemma bind_map (μ : Measure α) {f : α → β} (hf : Measurable f) {k : β → Measure δ} + (hk : Measurable k) : (μ.map f).bind k = μ.bind fun a ↦ k (f a) := by + rw [Measure.bind, Measure.bind, Measure.map_map hk hf] + rfl + +end Prerequisites + +variable {γ Ω Ω' δ : Type*} {α β : Type u} + [MeasurableSpace γ] [MeasurableSpace Ω] [MeasurableSpace Ω'] + [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace δ] + +/-- `HasTrace prog P out` states that the family of measures `prog : γ → Measure β` is realised by +drawing a *trace* `ω : Ω` from the Markov kernel `P` and reading the deterministic value +`out (c, ω)` off it. + +`Ω` is meant to be the space of all the values the program draws at its `←`s, and `P` an iterated +`Kernel.compProd`, one factor per `←`. See the module docstring. -/ +structure HasTrace (prog : γ → Measure β) (P : Kernel γ Ω) (out : γ × Ω → β) : Prop where + /-- The readout is deterministic: all the randomness of the program sits in the trace. -/ + measurable_out : Measurable out + /-- Drawing a trace and reading the answer off it runs the program. -/ + map_eq (c : γ) : (P c).map (fun ω ↦ out (c, ω)) = prog c + +/-- The kernel attached to a subprogram known to be Markov. Unlike `IsMarkov.toKernel`, the proof +is an ordinary argument rather than an instance, so `IsMarkovKernel` can be found by unification on +the kernel term alone, whatever the proof is — including while it is still an unsolved goal. This +is what `rdo_trace` emits at the leaves of a program. -/ +def markovKernel (prog : γ → Measure α) (h : IsMarkov prog) : Kernel γ α := ⟨prog, h.measurable⟩ + +@[simp] +lemma markovKernel_apply (prog : γ → Measure α) (h : IsMarkov prog) (c : γ) : + markovKernel prog h c = prog c := rfl + +instance (prog : γ → Measure α) (h : IsMarkov prog) : IsMarkovKernel (markovKernel prog h) := + ⟨fun c ↦ h.isProbabilityMeasure c⟩ + +namespace HasTrace + +variable {prog : γ → Measure β} {P : Kernel γ Ω} {out : γ × Ω → β} + +/-- The program's result, as a random variable on the trace space, has the program's law. -/ +lemma hasLaw_out (h : HasTrace prog P out) (c : γ) : + HasLaw (fun ω ↦ out (c, ω)) (prog c) (P c) where + aemeasurable := (h.measurable_out.comp measurable_prodMk_left).aemeasurable + map_eq := h.map_eq c + +/-- A traced program is the kernel obtained by reading `out` off `P`. -/ +lemma eq_map [IsSFiniteKernel P] (h : HasTrace prog P out) : + prog = ⇑((Kernel.id ×ₖ P).map out) := by + funext c + rw [← h.map_eq c, Kernel.map_apply _ h.measurable_out, Kernel.prod_apply, Kernel.id_apply, + Measure.dirac_prod, Measure.map_map h.measurable_out measurable_prodMk_left] + rfl + +/-- A program with a Markov trace is a Markov kernel. This is the `is_markov` property, obtained +for free from the trace. -/ +lemma isMarkov (h : HasTrace prog P out) [IsMarkovKernel P] : IsMarkov prog where + measurable' := by rw [h.eq_map]; exact Kernel.measurable _ + isProbabilityMeasure c := by + rw [← h.map_eq c] + exact Measure.isProbabilityMeasure_map + (h.measurable_out.comp measurable_prodMk_left).aemeasurable + +lemma congr {prog' : γ → Measure β} (h : HasTrace prog P out) (h' : ∀ c, prog' c = prog c) : + HasTrace prog' P out := + ⟨h.measurable_out, fun c ↦ (h.map_eq c).trans (h' c).symm⟩ + +/-- A single draw `let x ← κ`: the trace is the value drawn, and the readout is that value. -/ +protected lemma sample (κ : Kernel γ α) : HasTrace (⇑κ) κ Prod.snd := + ⟨measurable_snd, fun _ ↦ Measure.map_id⟩ + +/-- Any subprogram already known to be Markov can be treated as a single atomic draw. This is the +knob controlling how finely a program is traced: stop here and the subprogram's own `←`s stay +hidden inside one coordinate. -/ +protected lemma of_isMarkov (prog : γ → Measure α) [IsMarkov prog] : + HasTrace prog (IsMarkov.toKernel prog) Prod.snd := + ⟨measurable_snd, fun _ ↦ Measure.map_id⟩ + + +/-- The leaf rule as `rdo_trace` uses it: a Markov subprogram is one atomic draw, with the +`IsMarkov` proof carried in the kernel so that `IsMarkovKernel` stays available downstream. -/ +protected lemma leaf (prog : γ → Measure α) (h : IsMarkov prog) : + HasTrace prog (markovKernel prog h) Prod.snd := + ⟨measurable_snd, fun _ ↦ Measure.map_id⟩ + +/-- A value marked with `record`: a coordinate of the trace holding a deterministic function of +the parameter — which, inside a program, means of everything drawn before it. The program itself is +unchanged (`RDo.record_bind`); the trace gains a coordinate. -/ +protected lemma «record» {f : γ → α} (hf : Measurable f) : + HasTrace (fun c ↦ RDo.record (f c)) (Kernel.deterministic f hf) Prod.snd := + ⟨measurable_snd, fun c ↦ by + rw [Kernel.deterministic_apply] + exact Measure.map_id⟩ + +/-- `return e` draws nothing: its trace space is `PUnit`. -/ +protected lemma pure {g : γ → β} (hg : Measurable g) : + HasTrace (fun c ↦ mPure (g c)) (Kernel.const γ (Measure.dirac PUnit.unit)) + (fun p : γ × PUnit ↦ g p.1) := + ⟨hg.comp measurable_fst, fun c ↦ by + rw [Kernel.const_apply] + change (Measure.dirac PUnit.unit).map (fun _ : PUnit ↦ g c) = mPure (g c) + rw [Measure.map_dirac' measurable_const] + rfl⟩ + +/-- A program that does not look at its parameter. -/ +protected lemma const {μ : Measure β} {Q : Measure Ω} {o : Ω → β} (ho : Measurable o) + (h : Q.map o = μ) : HasTrace (fun _ : γ ↦ μ) (Kernel.const γ Q) (fun p ↦ o p.2) := + ⟨ho.comp measurable_snd, fun _ ↦ h⟩ + +/-- Weakening: a subprogram that does not read the trace drawn so far, seen as a program over the +enlarged parameter `γ × Ω`. Presenting its trace kernel as `Kernel.prodMkRight` is what makes +`indepFun_snd_compProd_prodMkRight` applicable afterwards. -/ +protected lemma prodMkRight (Ω₀ : Type*) [MeasurableSpace Ω₀] {prog : γ → Measure α} + {R : Kernel γ Ω'} {o : γ × Ω' → α} (h : HasTrace prog R o) : + HasTrace (fun p : γ × Ω₀ ↦ prog p.1) (Kernel.prodMkRight Ω₀ R) + (fun q : (γ × Ω₀) × Ω' ↦ o (q.1.1, q.2)) := + ⟨h.measurable_out.comp ((measurable_fst.comp measurable_fst).prodMk measurable_snd), + fun p ↦ h.map_eq p.1⟩ + +/-- Reparametrisation: a traced program read through a measurable change of parameters. -/ +protected lemma comp (h : HasTrace prog P out) {g : δ → γ} (hg : Measurable g) : + HasTrace (fun d ↦ prog (g d)) (P.comap g hg) (fun p ↦ out (g p.1, p.2)) := + ⟨h.measurable_out.comp ((hg.comp measurable_fst).prodMk measurable_snd), fun d ↦ by + rw [Kernel.comap_apply]; exact h.map_eq (g d)⟩ + +/-- `let x ← p; return f x`: a deterministic tail draws nothing, so it only changes the readout. -/ +protected lemma bindPure {prog : γ → Measure α} {P : Kernel γ Ω} {out : γ × Ω → α} + (h : HasTrace prog P out) {f : γ × α → β} (hf : Measurable f) : + HasTrace (fun c ↦ prog c >>=ₘ fun a ↦ mPure (f (c, a))) P (fun p ↦ f (p.1, out p)) where + measurable_out := hf.comp (measurable_fst.prodMk h.measurable_out) + map_eq c := by + have hf' : Measurable fun a ↦ f (c, a) := hf.comp (measurable_const.prodMk measurable_id) + have hout : Measurable fun ω ↦ out (c, ω) := h.measurable_out.comp measurable_prodMk_left + change (P c).map (fun ω ↦ f (c, out (c, ω))) + = (prog c).bind fun a ↦ Measure.dirac (f (c, a)) + rw [Measure.bind_dirac_eq_map _ hf', ← h.map_eq c, Measure.map_map hf' hout] + rfl + +/-- `if p c then … else …`. The two branches have to draw in the same space; the trace kernel is +then the piecewise kernel and the readout branches on the same condition. Branches drawing in +different spaces are handled by tracing each into the sum of the two spaces first. -/ +protected lemma ite {p : γ → Prop} [DecidablePred p] (hp : MeasurableSet {c | p c}) + {prog₁ prog₂ : γ → Measure β} {P₁ P₂ : Kernel γ Ω} {out₁ out₂ : γ × Ω → β} + (h₁ : HasTrace prog₁ P₁ out₁) (h₂ : HasTrace prog₂ P₂ out₂) : + HasTrace (fun c ↦ if p c then prog₁ c else prog₂ c) (Kernel.piecewise hp P₁ P₂) + (fun q : γ × Ω ↦ if p q.1 then out₁ q else out₂ q) where + measurable_out := Measurable.ite (measurable_fst hp) h₁.measurable_out h₂.measurable_out + map_eq c := by + rw [Kernel.piecewise_apply] + by_cases hc : p c <;> simp only [Set.mem_ofPred_eq, hc, ite_true, ite_false] + · exact h₁.map_eq c + · exact h₂.map_eq c + +/-- **The bind rule.** The trace of `let x ← p; q x` is the composition-product of the trace of `p` +with the trace of `q`, the latter being allowed to depend on the whole trace of `p` and not only +on the value `x` that `p` returned. + +This is the rule that turns the sequential structure of an `rdo` program into the `⊗ₖ` structure of +its trace kernel, and hence into conditional laws: see `hasCondDistrib_snd_compProd`. -/ +protected lemma bind {prog : γ → Measure α} {P : Kernel γ Ω} [IsSFiniteKernel P] {out : γ × Ω → α} + {cont : γ × α → Measure β} (hcont : Measurable cont) + {Q : Kernel (γ × Ω) Ω'} [IsSFiniteKernel Q] {out' : (γ × Ω) × Ω' → β} + (h : HasTrace prog P out) (h' : HasTrace (fun p ↦ cont (p.1, out p)) Q out') : + HasTrace (fun c ↦ prog c >>=ₘ fun a ↦ cont (c, a)) (P ⊗ₖ Q) + (fun p : γ × (Ω × Ω') ↦ out' ((p.1, p.2.1), p.2.2)) where + measurable_out := + h'.measurable_out.comp ((measurable_fst.prodMk (measurable_fst.comp measurable_snd)).prodMk + (measurable_snd.comp measurable_snd)) + map_eq c := by + have hg : Measurable fun w : Ω × Ω' ↦ out' ((c, w.1), w.2) := + h'.measurable_out.comp ((measurable_const.prodMk measurable_fst).prodMk measurable_snd) + have hcont' : Measurable fun a ↦ cont (c, a) := + hcont.comp (measurable_const.prodMk measurable_id) + have hout : Measurable fun ω ↦ out (c, ω) := h.measurable_out.comp measurable_prodMk_left + rw [Kernel.compProd_apply_eq_compProd_sectR, map_compProd _ _ hg] + change (Measure.bind (P c) fun ω ↦ (Kernel.sectR Q c ω).map fun ω' ↦ out' ((c, ω), ω')) = _ + rw [show (fun ω ↦ (Kernel.sectR Q c ω).map fun ω' ↦ out' ((c, ω), ω')) + = fun ω ↦ cont (c, out (c, ω)) from funext fun ω ↦ h'.map_eq (c, ω), ← h.map_eq c] + exact (bind_map (P c) hout hcont').symm + +end HasTrace + + +/-! ## Reading the random variables off the trace + +The trace kernel of a program is a right-nested `⊗ₖ`, one factor per `←`. This section provides +the single rule needed to exploit that: *peeling* a `⊗ₖ` off a law or a conditional law splits it +into the law of the next draw given the history, and a conditional law for everything after it. +Applying it once per `←` walks down the program, and produces, for the `k`-th draw, its +conditional distribution given the first `k-1` draws — read straight off the program text. + +When the peeled kernel does not depend on the history (`Kernel.prodMkRight`, or a constant), +that conditional law is an unconditional law together with an independence statement: +`hasLaw_of_const` and `indepFun_of_const`. +-/ + +section Peel + +variable {Ω α β δ : Type*} [MeasurableSpace Ω] [MeasurableSpace α] [MeasurableSpace β] + [MeasurableSpace δ] {P : Measure Ω} {W : Ω → α × β} + +/-- The first half of a random variable whose law is a composition-product. -/ +lemma _root_.ProbabilityTheory.HasLaw.compProd_fst {ν : Measure α} [SFinite ν] + {η : Kernel α β} [IsMarkovKernel η] (h : HasLaw W (ν ⊗ₘ η) P) : + HasLaw (fun ω ↦ (W ω).1) ν P where + aemeasurable := measurable_fst.comp_aemeasurable h.aemeasurable + map_eq := by + rw [show (fun ω ↦ (W ω).1) = Prod.fst ∘ W from rfl, + ← AEMeasurable.map_map_of_aemeasurable measurable_fst.aemeasurable h.aemeasurable, h.map_eq, + ← Measure.fst, Measure.fst_compProd] + +/-- **The peeling rule, unconditional form.** If the joint law of a pair is `ν ⊗ₘ η`, then the +second component has conditional distribution `η` given the first. -/ +lemma _root_.ProbabilityTheory.HasLaw.compProd_snd {ν : Measure α} [SFinite ν] + {η : Kernel α β} [IsMarkovKernel η] (h : HasLaw W (ν ⊗ₘ η) P) : + HasCondDistrib (fun ω ↦ (W ω).2) (fun ω ↦ (W ω).1) η P := by + refine ⟨h.aemeasurable, ?_⟩ + rw [h.compProd_fst.map_eq] + exact h.map_eq + +variable {H : Ω → δ} {κ : Kernel δ α} {η : Kernel (δ × α) β} {W : Ω → α × β} + +/-- The first half of a random variable whose conditional law is a composition-product. -/ +lemma _root_.ProbabilityTheory.HasCondDistrib.compProd_fst [SFinite P] [IsSFiniteKernel κ] + [IsMarkovKernel η] (h : HasCondDistrib W H (κ ⊗ₖ η) P) : + HasCondDistrib (fun ω ↦ (W ω).1) H κ P := by + have := h.fst + rwa [Kernel.fst_compProd] at this + +/-- **The peeling rule.** If, given the history `H`, the rest of the trace has conditional law +`κ ⊗ₖ η`, then the next draw has conditional law `κ` given `H`, and everything after it has +conditional law `η` given `H` *and* that draw. Iterating this walks down an `rdo` program one +`←` at a time. -/ +lemma _root_.ProbabilityTheory.HasCondDistrib.compProd_snd [SFinite P] [IsSFiniteKernel κ] + [IsMarkovKernel η] (h : HasCondDistrib W H (κ ⊗ₖ η) P) : + HasCondDistrib (fun ω ↦ (W ω).2) (fun ω ↦ (H ω, (W ω).1)) η P := + HasCondDistrib.of_compProd (Y := fun ω ↦ (W ω).1) (Z := fun ω ↦ (W ω).2) h + +end Peel + +section Coordinates + +variable {γ Ω Ω' Ω'' : Type*} [MeasurableSpace γ] [MeasurableSpace Ω] [MeasurableSpace Ω'] + [MeasurableSpace Ω''] + +/-- A section of a composition-product of kernels is the composition-product of the sections. +This is what lets the peeling rule be applied again to the tail of a trace. -/ +lemma _root_.ProbabilityTheory.Kernel.sectR_compProd (κ : Kernel (γ × Ω) Ω') + (η : Kernel ((γ × Ω) × Ω') Ω'') [IsSFiniteKernel κ] [IsSFiniteKernel η] (c : γ) : + Kernel.sectR (κ ⊗ₖ η) c + = Kernel.sectR κ c ⊗ₖ η.comap (fun p : Ω × Ω' ↦ ((c, p.1), p.2)) (by fun_prop) := by + ext ω s hs + rw [Kernel.sectR_apply, Kernel.compProd_apply hs, Kernel.compProd_apply hs] + simp [Kernel.comap_apply] + +/-- The whole trace of a program whose last construct is a `←`, as a random variable on its own +space: its law is a composition-product, ready for the peeling rule. This is the entry point of +the recursion. -/ +lemma hasLaw_id_compProd (P : Kernel γ Ω) (Q : Kernel (γ × Ω) Ω') [IsSFiniteKernel P] + [IsSFiniteKernel Q] (c : γ) : + HasLaw (id : Ω × Ω' → Ω × Ω') (P c ⊗ₘ Kernel.sectR Q c) ((P ⊗ₖ Q) c) := + ⟨measurable_id.aemeasurable, by + rw [Measure.map_id, Kernel.compProd_apply_eq_compProd_sectR]⟩ + +/-- A draw that does not read the history is drawn from a constant kernel. Combined with +`HasCondDistrib.hasLaw_of_const` and `HasCondDistrib.indepFun_of_const`, this turns a conditional +law into a law plus an independence statement. -/ +@[simp] +lemma _root_.ProbabilityTheory.Kernel.sectR_prodMkRight_eq_const {R : Kernel γ Ω'} (c : γ) : + Kernel.sectR (Kernel.prodMkRight Ω R) c = Kernel.const Ω (R c) := + Kernel.ext fun _ ↦ by + rw [Kernel.sectR_apply, Kernel.prodMkRight_apply, Kernel.const_apply] + +variable (P : Kernel γ Ω) (Q : Kernel (γ × Ω) Ω') (c : γ) + +@[simp] +lemma map_fst_compProd [IsSFiniteKernel P] [IsMarkovKernel Q] : + ((P ⊗ₖ Q) c).map Prod.fst = P c := by + rw [← Kernel.fst_apply, Kernel.fst_compProd] + +/-- The draws made before the last `←` have the law given by the earlier factor. -/ +lemma hasLaw_fst_compProd [IsSFiniteKernel P] [IsMarkovKernel Q] : + HasLaw (Prod.fst : Ω × Ω' → Ω) (P c) ((P ⊗ₖ Q) c) := + ⟨measurable_fst.aemeasurable, map_fst_compProd P Q c⟩ + +/-- The last draw of a program has conditional distribution `Q` given all the draws before it. -/ +lemma hasCondDistrib_snd_compProd [IsSFiniteKernel P] [IsMarkovKernel Q] : + HasCondDistrib Prod.snd Prod.fst (Kernel.sectR Q c) ((P ⊗ₖ Q) c) := + (hasLaw_id_compProd P Q c).compProd_snd + +/-- The same statement in the `Kernel.comap` form that the peeling chain keeps: `Kernel.sectR Q c` +*is* `Q.comap (fun ω ↦ (c, ω))`, and every later step produces a `comap` too. -/ +lemma hasCondDistrib_snd_compProd_comap [IsSFiniteKernel P] [IsMarkovKernel Q] : + HasCondDistrib Prod.snd Prod.fst + (Q.comap (fun ω ↦ (c, ω)) (measurable_const.prodMk measurable_id)) ((P ⊗ₖ Q) c) := + hasCondDistrib_snd_compProd P Q c + +/-- When a draw does not read the earlier ones, it is independent of them. -/ +lemma indepFun_snd_compProd_prodMkRight [IsMarkovKernel P] {R : Kernel γ Ω'} [IsMarkovKernel R] : + IndepFun (Prod.fst : Ω × Ω' → Ω) Prod.snd ((P ⊗ₖ Kernel.prodMkRight Ω R) c) := by + have h := hasCondDistrib_snd_compProd P (Kernel.prodMkRight Ω R) c + rw [Kernel.sectR_prodMkRight_eq_const] at h + exact h.indepFun_of_const + +/-- When a draw does not read the earlier ones, its law is the kernel it was drawn from. -/ +lemma hasLaw_snd_compProd_prodMkRight [IsMarkovKernel P] {R : Kernel γ Ω'} [IsMarkovKernel R] : + HasLaw (Prod.snd : Ω × Ω' → Ω') (R c) ((P ⊗ₖ Kernel.prodMkRight Ω R) c) := by + have h := hasCondDistrib_snd_compProd P (Kernel.prodMkRight Ω R) c + rw [Kernel.sectR_prodMkRight_eq_const] at h + exact h.hasLaw_of_const + +end Coordinates + +/-! ## Peeling through a reparametrisation + +After the first step the tail kernel of the peeling chain is a `Kernel.comap`, and stays one. These +are the lemmas the `rdo_peel` tactic iterates. +-/ + +section Peel' + +variable {Γ Ωk Ωr Ω δ : Type*} [MeasurableSpace Γ] [MeasurableSpace Ωk] [MeasurableSpace Ωr] + [MeasurableSpace Ω] [MeasurableSpace δ] + +/-- Comap distributes over the composition-product. -/ +lemma _root_.ProbabilityTheory.Kernel.comap_compProd (κ : Kernel Γ Ωk) (η : Kernel (Γ × Ωk) Ωr) + [IsSFiniteKernel κ] [IsSFiniteKernel η] {ι : δ → Γ} (hι : Measurable ι) : + (κ ⊗ₖ η).comap ι hι = κ.comap ι hι ⊗ₖ + η.comap (fun p : δ × Ωk ↦ (ι p.1, p.2)) ((hι.comp measurable_fst).prodMk measurable_snd) := by + ext d s hs + rw [Kernel.comap_apply, Kernel.compProd_apply hs, Kernel.compProd_apply hs] + simp [Kernel.comap_apply] + +variable {P : Measure Ω} {H : Ω → δ} {W : Ω → Ωk × Ωr} {κ : Kernel Γ Ωk} {η : Kernel (Γ × Ωk) Ωr} + {ι : δ → Γ} {hι : Measurable ι} [SFinite P] [IsSFiniteKernel κ] [IsMarkovKernel η] + +/-- Peel the next draw off a reparametrised tail. -/ +lemma _root_.ProbabilityTheory.HasCondDistrib.comap_compProd_fst + (h : HasCondDistrib W H ((κ ⊗ₖ η).comap ι hι) P) : + HasCondDistrib (fun ω ↦ (W ω).1) H (κ.comap ι hι) P := by + rw [Kernel.comap_compProd] at h + exact h.compProd_fst + +/-- ... and keep the rest of the tail, still a `comap`, for the next step. -/ +lemma _root_.ProbabilityTheory.HasCondDistrib.comap_compProd_snd + (h : HasCondDistrib W H ((κ ⊗ₖ η).comap ι hι) P) : + HasCondDistrib (fun ω ↦ (W ω).2) (fun ω ↦ (H ω, (W ω).1)) + (η.comap (fun p : δ × Ωk ↦ (ι p.1, p.2)) + ((hι.comp measurable_fst).prodMk measurable_snd)) P := by + rw [Kernel.comap_compProd] at h + exact h.compProd_snd + +end Peel' + +end RDo diff --git a/RandomDo/Probability/Transfer.lean b/RandomDo/Probability/Transfer.lean new file mode 100644 index 0000000..eca63b3 --- /dev/null +++ b/RandomDo/Probability/Transfer.lean @@ -0,0 +1,396 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import Mathlib.Dynamics.Ergodic.MeasurePreserving +public import Mathlib.MeasureTheory.Function.StronglyMeasurable.AEStronglyMeasurable +public import Mathlib.Tactic.FunProp +public import Mathlib.Tactic.Measurability +public meta import Lean.LabelAttribute +public import Lean.LabelAttribute + +set_option linter.style.header false + +/-! +# The `transfer` tactic + +A statement about random variables on a probability space `(Ω, P)` pulls back along a +measure-preserving map `f : Ω' → Ω`: the law of `X` under `P` is the law of `fun ω ↦ X (f ω)` under +`P'`, two random variables are independent if and only if their compositions with `f` are, and so +on. Two kinds of lemmas record this. + +* A `@[transfer]` lemma states an invariance, + ``` + S X P ↔ S (fun ω ↦ X (f ω)) P' or t X P = t (fun ω ↦ X (f ω)) P' + ``` + with the hypothesis `hf : MeasurePreserving f P' P` as its first explicit argument, possibly + followed by side conditions such as the measurability of `X`. The old space is on the left, so + that rewriting with the lemma only ever needs first-order unification. +* A `@[transfer_forward]` lemma states an implication `S X P → S (fun ω ↦ X (f ω)) P'`, as + `(h : S X P) (hf : MeasurePreserving f P' P) (side conditions…) : S (fun ω ↦ X (f ω)) P'`, for + statements that only go one way, such as measurability. + +`transfer hf` rewrites the goal with every `@[transfer]` lemma instantiated at `hf`, discharging the +side conditions by `assumption`, `fun_prop` and `measurability`, and closes the goal by +`assumption` if it can. `transfer hf at h` transports a hypothesis instead, by rewriting or, when +nothing rewrites, by a `@[transfer_forward]` lemma, under the binders of `h` if it has any: this +pulls a fact about the old space back to the new one. `transfer hf at h ⊢` does both, the goal +first, since transporting a measurability hypothesis destroys the fact the goal's side conditions +need. A named target that mentions the old space and is left as it was is an error. `transfer` +alone is for the obligation left by `extend_space`: it introduces the new space, the map and the +statement on the new space, transfers the goal, and closes it with that statement. +-/ + +public meta section + +open Lean Meta Elab Tactic + +/-- A lemma stating that a probabilistic statement pulls back along a measure-preserving map +`hf : MeasurePreserving f P' P`, in the form `S X P ↔ S (fun ω ↦ X (f ω)) P'` or +`t X P = t (fun ω ↦ X (f ω)) P'`, with `hf` as its first explicit argument, possibly followed by +side conditions. The `transfer` tactic rewrites with all of them. -/ +register_label_attr transfer + +/-- A lemma transporting a hypothesis along a measure-preserving map, in the form +`(h : S X P) (hf : MeasurePreserving f P' P) (side conditions…) : S (fun ω ↦ X (f ω)) P'`. The +`transfer` tactic uses them on hypotheses that no `@[transfer]` lemma rewrites. -/ +register_label_attr transfer_forward + +namespace RDo.Tactic + +/-- The tactic state together with the message log, which `SavedState.restore` keeps: an attempt +that fails must not leave the errors it logged behind. -/ +structure FullState where + /-- The tactic state. -/ + state : Tactic.SavedState + /-- The message log. -/ + messages : MessageLog + +/-- Save the tactic state and the message log. -/ +def saveFullState : TacticM FullState := + return ⟨← saveState, (← getThe Core.State).messages⟩ + +/-- Restore the tactic state and the message log. -/ +def FullState.restore (s : FullState) : TacticM Unit := do + s.state.restore + modifyThe Core.State fun st ↦ { st with messages := s.messages } + +/-- Run `x` with a budget of `n` thousand heartbeats of its own, or the ambient budget if that is +smaller. An attempt that fails, such as `measurability` on a set that is not measurable, must not +exhaust the budget of the declaration for what comes after it. -/ +def withHeartbeatBudget {m : Type → Type} {α : Type} [Monad m] [MonadControlT CoreM m] + [MonadReaderOf Core.Context m] [MonadWithReaderOf Core.Context m] (n : Nat) (x : m α) : + m α := do + let ambient := (← readThe Core.Context).maxHeartbeats + let budget := if ambient == 0 then n * 1000 else min ambient (n * 1000) + withCurrHeartbeats <| withTheReader Core.Context (fun c ↦ { c with maxHeartbeats := budget }) x + +/-- The discharger for the side conditions of `@[transfer]` lemmas: `assumption`, then `fun_prop` +for the measurability of a function and `measurability` for that of a set. A maximum recursion +depth error inside `measurability`, which happens on unprovable goals, is turned into a plain +failure so that it only makes the rewrite fail, and the search runs on a budget of its own so that +such a failure does not exhaust the heartbeats of the declaration. -/ +syntax (name := transferDischarger) "transfer_discharger" : tactic + +elab_rules : tactic + | `(tactic| transfer_discharger) => withMainContext do + let funProps : Array Name := #[``Measurable, ``AEMeasurable, + ``MeasureTheory.AEStronglyMeasurable, ``MeasureTheory.StronglyMeasurable] + let head := (← getMainTarget).getForallBody.getAppFn.constName? + -- `fun_prop` may fail on `AEMeasurable` where it succeeds on `Measurable`. + let tac ← if head.any funProps.contains then + `(tactic| first + | assumption + | (intros + first + | assumption + | fun_prop + | (apply Measurable.aemeasurable; fun_prop) + | (apply Measurable.aestronglyMeasurable; fun_prop))) + else `(tactic| first | assumption | (intros; first | assumption | measurability)) + -- Without error recovery, an alternative that fails to elaborate fails instead of logging an + -- error and going on with `sorry`: nothing a failed discharge tried leaks into the messages. + tryCatchRuntimeEx (withHeartbeatBudget 10000 <| Tactic.withoutRecover (evalTactic tac)) fun e ↦ + throwError "transfer_discharger: {e.toMessageData}" + +/-- The `@[transfer]` lemmas instantiated at `hf`, as `simp` arguments, together with the lemmas +pushing a preimage through set operations and through a random variable, which put the transferred +events in the form `extend_space` uses: `f ⁻¹' s` for an event `s`, and `(fun ω ↦ X (f ω)) ⁻¹' t` +for an event `X ⁻¹' t`. A lemma that does not elaborate at `hf`, for want of an instance on the +measure for example, is left out rather than making the whole rewrite fail. -/ +def transferSimpArgs (hf : Term) : TacticM (Array (TSyntax ``Lean.Parser.Tactic.simpLemma)) := do + let mut args := #[] + for n in ← labelled `transfer do + let s ← saveFullState + let ok ← tryCatchRuntimeEx + (do + -- Postponing keeps a lemma whose instances depend on yet unknown types, such as the + -- codomain of an integrand, and rejects one whose instances fail outright. + discard <| Term.withoutErrToSorry <| + Tactic.elabTerm (← `($(mkIdent n):ident $hf)) none (mayPostpone := true) + pure true) + (fun _ ↦ pure false) + s.restore + if ok then args := args.push (← `(Lean.Parser.Tactic.simpLemma| $(mkIdent n):ident $hf)) + let extra ← #[``Set.preimage_ofPred_eq, ``Set.preimage_preimage, ``Set.preimage_inter, + ``Set.preimage_union, ``Set.preimage_compl, ``Set.preimage_sdiff].mapM fun n ↦ + `(Lean.Parser.Tactic.simpLemma| $(mkIdent n):ident) + return args ++ extra + +/-- Try to close the goal `g` with `tac`, returning its proof. The state is restored on failure, +and a runtime error such as a maximum recursion depth counts as a failure. The attempt runs on a +heartbeat budget of its own. -/ +def tryTactic? (g : MVarId) (tac : Syntax) : TacticM (Option Expr) := do + let s ← saveFullState + tryCatchRuntimeEx + (do + -- Without error recovery, a failure inside a nested `by` is a failure, not a `sorry`. + let gs ← withHeartbeatBudget (m := TermElabM) 20000 <| Term.withoutErrToSorry <| + Tactic.run g (evalTactic tac) + if gs.isEmpty then return some (← instantiateMVars (.mvar g)) + s.restore + return none) + (fun _ ↦ do + s.restore + return none) + +/-- The old space of `hf : MeasurePreserving f P' P`: the type `Ω` and the measure `P`. -/ +def oldSpace? (hf : Expr) : MetaM (Option (Expr × Expr)) := do + let ty ← whnfR (← instantiateMVars (← inferType hf)) + unless ty.isAppOfArity ``MeasureTheory.MeasurePreserving 7 do return none + let as := ty.getAppArgs + return some (as[1]!, as[6]!) + +/-- Whether `ty` mentions the old space of `hf`, its type or its measure. A statement that does +not has nothing for `transfer` to do, so leaving it as it is is not a failure. When the old space +cannot be read off `hf`, every statement counts as mentioning it. -/ +def mentionsOldSpace (hf ty : Expr) : MetaM Bool := do + let some (Ω, P) ← oldSpace? hf | return true + let ty ← instantiateMVars ty + return Ω.occurs ty || P.occurs ty + +/-- The head constant of a statement, under its binders. -/ +def headUnderBinders (ty : Expr) : MetaM (Option Name) := + forallTelescope ty fun _ body ↦ return body.getAppFn.constName? + +/-- The head constant of the statement a `@[transfer_forward]` lemma transports: that of the type +of its first explicit argument. -/ +def forwardLemmaHead? (n : Name) : MetaM (Option Name) := do + forallTelescope (← getConstInfo n).type fun xs _ ↦ do + for x in xs do + if (← x.fvarId!.getBinderInfo).isExplicit then + return ← headUnderBinders (← x.fvarId!.getType) + return none + +/-- Transport the hypothesis `h` forward along `hf` with a `@[transfer_forward]` lemma, as +`lemma h hf side…`, the side conditions being discharged by `transfer_discharger`. A hypothesis +with binders, `∀ n, S (X n) P`, is transported under them. Only the lemmas about the head constant +of `h` are tried, so that a definition is transported as itself and not unfolded to what it is +defined as. Returns the statement and proof of the transported hypothesis. -/ +def transferForward? (h hf : Expr) : TacticM (Option (Expr × Expr)) := do + let hfStx ← Term.exprToSyntax hf + forallTelescope (← instantiateMVars (← inferType h)) fun xs body ↦ do + let hStx ← Term.exprToSyntax (mkAppN h xs) + let head := body.getAppFn.constName? + for n in ← labelled `transfer_forward do + if let some lemmaHead ← forwardLemmaHead? n then + unless head == some lemmaHead do continue + let s ← saveFullState + let r ← tryCatchRuntimeEx + (do + let e ← Term.withoutErrToSorry <| + Tactic.elabTerm (← `($(mkIdent n):ident $hStx $hfStx)) none + let (args, bis, concl) ← forallMetaTelescope (← inferType e) + for (a, bi) in args.zip bis do + if bi.isInstImplicit then + a.mvarId!.assign (← synthInstance (← instantiateMVars (← inferType a))) + else if bi.isExplicit then + let some _ ← tryTactic? a.mvarId! (← `(tactic| transfer_discharger)) + | throwError "side condition" + let pf ← instantiateMVars (mkAppN e args) + -- `g ∘ f` is put in the form `fun ω ↦ g (f ω)`. + let concl ← instantiateMVars concl + let concl ← Core.betaReduce (← deltaExpand concl (· == ``Function.comp)) + if pf.hasExprMVar || concl.hasExprMVar then throwError "metavariables" + pure (some (← mkForallFVars xs concl, ← mkLambdaFVars xs pf))) + (fun _ ↦ do + s.restore + pure none) + if r.isSome then return r + return none + +/-- Introduce the binders of a `transfer` obligation: everything up to and including the statement +on the new space, which is the binder after the `MeasurePreserving` hypothesis. Returns the new +goal, the map hypothesis and that statement. -/ +partial def introTransferObligation (g : MVarId) : MetaM (MVarId × FVarId × FVarId) := + go g none +where + go (g : MVarId) (hf? : Option FVarId) : MetaM (MVarId × FVarId × FVarId) := do + unless (← instantiateMVars (← g.getType)).isForall do + throwError "transfer: the goal is not a `transfer` obligation. It should have the form\n \ + ∀ Ω' [MeasurableSpace Ω'] (P' : Measure Ω') [IsProbabilityMeasure P'] (f : Ω' → Ω), \ + MeasurePreserving f P' P → T' → T\nbut is{indentExpr (← g.getType)}" + let (fv, g) ← g.intro1P + match hf? with + | some hf => return (g, hf, fv) + | none => + let isMap ← g.withContext do + return (← whnfR (← fv.getType)).isAppOf ``MeasureTheory.MeasurePreserving + go g (if isMap then some fv else none) + +/-- The `@[transfer]` lemmas at `hf`, as computed by `transferSimpArgs`. -/ +abbrev SimpArgs := Array (TSyntax ``Lean.Parser.Tactic.simpLemma) + +/-- Rewrite the goal with the `@[transfer]` lemmas. Returns whether the goal changed or was +closed. -/ +def rewriteGoal (args : SimpArgs) : TacticM Bool := do + let g ← getMainGoal + let before ← instantiateMVars (← g.getType) + evalTactic (← `(tactic| simp -failIfUnchanged (disch := transfer_discharger) only [$args,*])) + let gs ← getUnsolvedGoals + if gs.isEmpty then return true + return gs[0]! != g || (← instantiateMVars (← gs[0]!.getType)) != before + +/-- Close the goal by `assumption` if it can be. -/ +def closeByAssumption : TacticM Unit := do + unless (← getUnsolvedGoals).isEmpty do evalTactic (← `(tactic| try assumption)) + +/-- After the goal has been transferred with `changed` reporting whether it was rewritten: fail if +it is still there, was not rewritten, and mentions the old space. -/ +def checkGoalTransferred (hf : Expr) (changed : Bool) : TacticM Unit := do + if changed || (← getUnsolvedGoals).isEmpty then return + withMainContext do + let ty ← getMainTarget + if ← mentionsOldSpace hf ty then + throwError "transfer: no `@[transfer]` lemma rewrites the goal{indentExpr ty}\n\ + A side condition, such as the measurability of a random variable, may not have been \ + discharged." + +/-- Transfer the goal along `hf`: rewrite it with the `@[transfer]` lemmas, then close it by +`assumption` if possible. Fails if the goal mentions the old space and nothing rewrote it. -/ +def transferGoal (args : SimpArgs) (hf : Expr) : TacticM Unit := do + let changed ← rewriteGoal args + closeByAssumption + checkGoalTransferred hf changed + +/-- Forward-transport the hypothesis `h`, named `name`, if the rewriting left it as it was. Returns +whether `h` was rewritten or transported. When `strict`, a hypothesis that mentions the old space +and could be neither rewritten nor transported is an error. -/ +def forwardHyp (hf : Expr) (h : FVarId) (name : Name) (strict : Bool) : TacticM Bool := + withMainContext do + -- `simp` replaces the hypothesis when it rewrites it, so that either `h` is gone or its name + -- now denotes a newer hypothesis; otherwise it is still there, as it was. + let lctx ← getLCtx + let some d := lctx.find? h | return true + if lctx.findFromUserName? name |>.any (·.fvarId != h) then return true + match ← transferForward? d.toExpr hf with + | some (ty, pf) => + let g ← (← getMainGoal).assert d.userName ty pf + let (_, g) ← g.intro1P + replaceMainGoal [← g.tryClear h] + return true + | none => + if strict && (← mentionsOldSpace hf d.type) then + throwError "transfer: nothing transfers the hypothesis {d.toExpr} :{indentExpr d.type}\n\ + No `@[transfer]` lemma rewrites it and no `@[transfer_forward]` lemma transports it.\n\ + A side condition, such as the measurability of a random variable, may not have been \ + discharged." + return false + +/-- Transfer the hypotheses `hs` along `hf`: rewrite each with the `@[transfer]` lemmas, then +replace each one that did not change by its forward transport by a `@[transfer_forward]` lemma. +All the rewriting comes first, since a forward transport destroys the measurability facts the +rewriting may need. When `strict`, a hypothesis that mentions the old space and is left as it was +is an error. Returns whether some hypothesis changed. -/ +def transferHyps (args : SimpArgs) (hf : Expr) (hs : Array FVarId) (strict : Bool) : + TacticM Bool := do + if (← getUnsolvedGoals).isEmpty then return false + let names ← withMainContext do hs.mapM (·.getUserName) + for h in hs do + let hStx ← withMainContext do Term.exprToSyntax (.fvar h) + evalTactic (← `(tactic| + simp -failIfUnchanged (disch := transfer_discharger) only [$args,*] at $hStx:term)) + let mut changed := false + for h in hs, name in names do + if (← getUnsolvedGoals).isEmpty then return true + changed := (← forwardHyp hf h name strict) || changed + return changed + +/-- `transfer hf`, for `hf : MeasurePreserving f P' P`, rewrites the goal with every `@[transfer]` +lemma instantiated at `hf`: the law of `X` under `P` becomes the law of `fun ω ↦ X (f ω)` under +`P'`, and likewise for events, integrals, independence and conditional laws. Side conditions, which +are measurability statements, are discharged by `assumption`, `fun_prop` and `measurability`. The +goal is then closed by `assumption` if possible. It is an error if the goal mentions the old space +and nothing rewrote it. + +* `transfer hf at h₁ h₂` transports hypotheses instead: a fact about the old space becomes the + corresponding fact about the new one, by the same rewriting or, for a hypothesis nothing + rewrites, by a `@[transfer_forward]` lemma, under the binders of the hypothesis if it has any. + It is an error if a named hypothesis mentions the old space and is left as it was. +* `transfer hf at h ⊢` and `transfer hf at *` do both, the goal first: transporting a + measurability hypothesis destroys the fact the goal's side conditions may need. With `*`, the + only error is when nothing at all changes. +* `transfer` alone discharges the `transfer` obligation of `extend_space`: it introduces the new + space, the map and the statement on the new space, transfers the goal and closes it with that + statement. -/ +syntax (name := transferTac) "transfer" (ppSpace colGt term)? + (Lean.Parser.Tactic.location)? : tactic + +elab_rules : tactic + | `(tactic| transfer $[$hf?]? $[$loc?]?) => withMainContext do + match hf?, loc? with + | none, some _ => + throwError "transfer: `at` needs the map to transfer along, as in `transfer hf at h`" + | some hf, some loc => + let hfE ← Tactic.elabTerm hf none + let args ← transferSimpArgs hf + match expandLocation loc with + | .wildcard => + let hs ← (← getLCtx).foldlM (init := #[]) fun hs d ↦ do + if d.isImplementationDetail || !(← isProp d.type) then pure hs + else pure (hs.push d.fvarId) + let goalChanged ← rewriteGoal args + let hypsChanged ← transferHyps args hfE hs (strict := false) + closeByAssumption + unless goalChanged || hypsChanged do + throwError "transfer: nothing to transfer" + | .targets hyps type => + let hs ← hyps.mapM getFVarId + let goalChanged ← if type then rewriteGoal args else pure false + discard <| transferHyps args hfE hs (strict := true) + if type then + closeByAssumption + checkGoalTransferred hfE goalChanged + | some hf, none => + let hfE ← Tactic.elabTerm hf none + transferGoal (← transferSimpArgs hf) hfE + | none, none => + let (g, hf, h) ← introTransferObligation (← getMainGoal) + replaceMainGoal [g] + withMainContext do + let args ← transferSimpArgs (← Term.exprToSyntax (.fvar hf)) + -- The statement on the new space is simplified too, so that both sides are normalized + -- the same way: `simp` turns `a = a` into `True` on its own, for instance. + let hName ← h.getUserName + let hStx ← Term.exprToSyntax (.fvar h) + evalTactic (← `(tactic| + simp -failIfUnchanged (disch := transfer_discharger) only [$args,*] at $hStx:term ⊢)) + if (← getUnsolvedGoals).isEmpty then return + let g ← getMainGoal + g.withContext do + -- `simp at h` replaces `h` by a new hypothesis of the same name. + let some hDecl := (← getLCtx).findFromUserName? hName + | throwError "transfer: could not close the goal after transferring it:{indentExpr + (← g.getType)}" + unless ← isDefEq (← g.getType) hDecl.type do + throwError "transfer: could not close the goal after transferring it:{indentExpr + (← g.getType)}\nwith the statement on the new space:{indentExpr hDecl.type}" + g.assign hDecl.toExpr + replaceMainGoal [] + +end RDo.Tactic + +end diff --git a/Test.lean b/Test.lean index 7862fd2..3effc40 100644 --- a/Test.lean +++ b/Test.lean @@ -1,10 +1,13 @@ module -- shake: keep-all --deprecated_module: ignore +public import Test.AlgTrace public import Test.Bind public import Test.Common public import Test.Control +public import Test.Extend public import Test.Gaps public import Test.Instances public import Test.IsMarkov public import Test.Loops public import Test.MonadLaws +public import Test.Transfer diff --git a/Test/AlgTrace.lean b/Test/AlgTrace.lean new file mode 100644 index 0000000..c2055a4 --- /dev/null +++ b/Test/AlgTrace.lean @@ -0,0 +1,296 @@ +module + +public import Test.Common + +set_option linter.style.header false + +/-! +# The algorithm-environment tactics on a toy algorithm + +A toy sequential algorithm, to show the pipeline end to end: write the policy as an `rdo` program, +get its trace from `rdo_trace`, package it as an `AlgTrace`, and then read the algorithm's internal +draws off any algorithm-environment sequence, with `alg_env_trace`. The sections after that pin +down what the tactic does with the rest of the context, what it introduces, and the errors it +reports. The same algorithm then exercises `extend_space` alongside an algorithm-environment +sequence. + +To do the same for `thompson` one needs the measurable equivalence between `Iic n → 𝓐 × 𝓨` and +`Vector (𝓐 × 𝓨) (n + 1)` that turns it into a policy, which is not available yet. Everything after +that point is what follows below. +-/ + +open MeasureTheory ProbabilityTheory Finset Learning RDo + +@[expose] public section + +noncomputable section + +namespace Test.AlgTrace + +universe u + +variable {K : ℕ} (hK : 0 < K) + +/-- The action, read off the history and the noise: depending on the sign of the noise, either +switch to arm `0` or repeat the last action. -/ +def readout (n : ℕ) (p : (Iic n → Fin K × ℝ) × ℝ) : Fin K := + if 0 < p.2 then ⟨0, hK⟩ else (p.1 ⟨n, by simp⟩).1 + +@[fun_prop] +lemma measurable_readout (n : ℕ) : Measurable (readout hK n) := by + unfold readout + exact Measurable.ite (measurableSet_lt measurable_const measurable_snd) measurable_const + (measurable_fst.comp ((measurable_pi_apply _).comp measurable_fst)) + +/-- The policy: perturb the last reward by Gaussian noise, then read the action off it. -/ +def policy (n : ℕ) (h : Iic n → Fin K × ℝ) : Measure (Fin K) := rdo + let z ← gaussianReal (h ⟨n, by simp⟩).2 1 + return readout hK n (h, z) + +instance (n : ℕ) : IsMarkov (policy hK n) := by unfold policy; is_markov + +/-- The noise the policy draws at step `n`, as a kernel: the one coordinate of its trace. -/ +def noise (n : ℕ) : Kernel (Iic n → Fin K × ℝ) ℝ := + markovKernel (fun h ↦ gaussianReal (h ⟨n, by simp⟩).2 1) + (IsMarkov.gaussianReal (by fun_prop) measurable_const) + +instance (n : ℕ) : IsMarkovKernel (noise (K := K) n) := by unfold noise; infer_instance + +lemma hasTrace_policy (n : ℕ) : HasTrace (policy hK n) (noise n) (readout hK n) := by + rdo_trace (policy hK n) with h + exact h + +/-- The algorithm. -/ +def alg : Algorithm (Fin K) ℝ where + policy n := markovKernel (policy hK n) inferInstance + p0 := Measure.dirac ⟨0, hK⟩ + +/-- Its trace: one Gaussian draw per step. -/ +def trace : AlgTrace (alg hK) ℝ where + K := noise + out := readout hK + hasTrace n := hasTrace_policy hK n + K0 := gaussianReal 0 1 + out0 := fun _ ↦ ⟨0, hK⟩ + measurable_out0 := measurable_const + map_out0 := by rw [Measure.map_const]; simp [alg] + +/-- **The payoff.** Given any algorithm-environment sequence for this algorithm, one may assume the +space also carries the noise `Z` the policy draws at each step: it has the conditional law `noise n` +given the history, and the action is `readout` of the history and it. The trajectory keeps the same +law, so anything proved there about the actions and feedbacks holds of the original sequence. -/ +theorem exists_noise (env : Environment (Fin K) ℝ) {Ω₀ : Type*} [MeasurableSpace Ω₀] + {P : Measure Ω₀} [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + ∃ (Ω' : Type) (_ : MeasurableSpace Ω') (P' : Measure Ω') (_ : IsProbabilityMeasure P') + (A' : ℕ → Ω' → Fin K) (Y' : ℕ → Ω' → ℝ) (Z : ℕ → Ω' → ℝ), + IsAlgEnvSeq A' Y' (alg hK) env P' + ∧ P'.map (trajectory A' Y') = P.map (trajectory A Y) + ∧ (∀ n, HasCondDistrib (Z (n + 1)) (history A' Y' n) (noise n) P') + ∧ (∀ n, A' (n + 1) =ᵐ[P'] fun ω ↦ readout hK n (history A' Y' n ω, Z (n + 1) ω)) := by + obtain ⟨Ω', mΩ', P', hP', A', Y', Z, hseq, -, hlaw, -, hZ, -, hA⟩ := + (trace hK).exists_isAlgEnvSeq_trace h + exact ⟨Ω', mΩ', P', hP', A', Y', Z, hseq, hlaw, hZ, hA⟩ + +/-- **The tactic at work.** `alg_env_trace` replaces the context and the goal by ones on a space +that also carries the noise `Z` the policy draws. The obligation that the statement only depends +on the law of the trajectory is discharged by `transfer` through the trajectory space, so only the +traced goal is left. Any hypothesis about the sequence travels with the goal, so nothing is +silently lost. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) with Ω P A Y Z hseq htr hZ₀ hZ hA₀ hA + -- `Z`, `hZ₀`, `hZ` and `hA` are the algorithm's draws and their laws, now available. + exact hseq.hasLaw_action_zero.map_eq + +/-- Without `with`, the names are `Ω P A Y T hseq htr hT₀ hT hA₀ hA`. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) using h + guard_hyp hseq : IsAlgEnvSeq A Y (alg hK) env P + guard_hyp htr : + IsAlgEnvSeq (fun n ω ↦ (T n ω, A n ω)) Y (trace hK).algorithm (env.withTrace ℝ) P + guard_hyp hT₀ : HasLaw (T 0) (trace hK).K0 P + guard_hyp hT : ∀ n, HasCondDistrib (T (n + 1)) (history A Y n) ((trace hK).K n) P + guard_hyp hA₀ : A 0 =ᵐ[P] fun ω ↦ (trace hK).out0 (T 0 ω) + guard_hyp hA : ∀ n, A (n + 1) =ᵐ[P] fun ω ↦ (trace hK).out n (history A Y n ω, T (n + 1) ω) + exact hseq.hasLaw_action_zero.map_eq + +/-- **Using the draws.** The second action is either arm `0` or the first action, since it is the +readout of the noise and the history: a statement about the actions, proved from `hA` on the traced +space and transferred back to the original one. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + ∀ᵐ ω ∂P, A 1 ω = ⟨0, hK⟩ ∨ A 1 ω = A 0 ω := by + alg_env_trace (trace hK) + filter_upwards [hA 0] with ω hω + rw [hω] + by_cases h0 : 0 < T (0 + 1) ω <;> simp [trace, readout, history, h0] + +/-- The space may live in any universe. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type u} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) + exact hseq.hasLaw_action_zero.map_eq + +/-! ## What travels with the goal, and what does not -/ + +/-- A hypothesis about the sequence travels with the goal, is available on the traced space, and +the obligation is still discharged. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) (ν : Measure (Fin K)) (hA1 : P.map (A 1) = ν) : + P.map (A 1) = ν := by + alg_env_trace (trace hK) + guard_hyp hA1 : P.map (A 1) = ν + exact hA1 + +/-- Data on the space that the goal does not depend on — a random variable, a point, and what is +about them — is cleared: it has no counterpart on the traced space. The obligation is discharged. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) (X : Ω₀ → ℝ) (_hX : Measurable X) (x : Ω₀) + (_hx : X x = 0) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) + fail_if_success guard_hyp X + fail_if_success guard_hyp _hX + fail_if_success guard_hyp x + fail_if_success guard_hyp _hx + exact hseq.hasLaw_action_zero.map_eq + +/-- A statement `transfer` has no lemma for leaves the obligation, which is then proved by hand, +here trivially. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + IsProbabilityMeasure P := by + alg_env_trace (trace hK) + case traced => infer_instance + case transfer => + intro Ω₁ _ P₁ _ A₁ Y₁ Ω₂ _ P₂ _ A₂ Y₂ h₁ h₂ hlaw h₀ + infer_instance + +/-! ## Errors -/ + +/-- Another algorithm, to check that a trace is matched against the algorithm of the hypothesis. -/ +def alg2 : Algorithm (Fin K) ℝ where + policy _ := Kernel.const _ (Measure.dirac ⟨0, hK⟩) + p0 := Measure.dirac ⟨0, hK⟩ + +/-- +error: alg_env_trace: the goal depends on + s +of type + Set Ω₀ +which lives on the space of the sequence without being part of it. Only statements about the actions, the feedbacks and the measure survive the change of space. +-/ +#guard_msgs in +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) (s : Set Ω₀) (hs : P s = 1 / 2) : P s = 1 / 2 := by + alg_env_trace (trace hK) + +/-- +error: alg_env_trace: no `IsAlgEnvSeq` hypothesis in the context +-/ +#guard_msgs in +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] : P Set.univ = 1 := by + alg_env_trace (trace hK) + +/-- +error: alg_env_trace: hP is not an `IsAlgEnvSeq` hypothesis +-/ +#guard_msgs in +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] (hP : P Set.univ = 1) : P Set.univ = 1 := by + alg_env_trace (trace hK) using hP + +/-- +error: alg_env_trace: the probability space must be given by local hypotheses, but ℕ → ℝ is not +-/ +#guard_msgs in +example (env : Environment (Fin K) ℝ) {P : Measure (ℕ → ℝ)} [IsProbabilityMeasure P] + {A : ℕ → (ℕ → ℝ) → Fin K} {Y : ℕ → (ℕ → ℝ) → ℝ} (h : IsAlgEnvSeq A Y (alg hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) + +/-- +error: alg_env_trace: the action and feedback sequences must be local hypotheses, but fun n ω ↦ Y n ω + 0 is not +-/ +#guard_msgs in +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A (fun n ω ↦ Y n ω + 0) (alg hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) + +/-- +error: alg_env_trace: trace hK is not a trace of the algorithm of the hypothesis +-/ +#guard_msgs in +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg2 hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) + +/-- +error: alg_env_trace: at most 11 names may be given +-/ +#guard_msgs in +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) with a b c d e f g i j k l m + +/-! ## `extend_space` alongside an algorithm-environment sequence -/ + +/-- **`extend_space` alongside an algorithm-environment sequence.** After the extension, `Ω`, `P`, +`A` and `Y` live on a larger space that also carries a Gaussian `U` independent of the whole +trajectory, and `h` has been transported by `IsAlgEnvSeq.comp_measurePreserving`. The statement +does not mention the original space, so the `transfer` obligation is trivial and `extend_space` +closes it. The measurability of the sequence is put in the context first, so that the +independence statement `hind` covers `A` and `Y`. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + ∃ (Ω' : Type) (_ : MeasurableSpace Ω') (P' : Measure Ω') (_ : IsProbabilityMeasure P') + (A' : ℕ → Ω' → Fin K) (Y' : ℕ → Ω' → ℝ) (U : Ω' → ℝ), + IsAlgEnvSeq A' Y' (alg hK) env P' ∧ HasLaw U (gaussianReal 0 1) P' + ∧ IndepFun (trajectory A' Y') U P' := by + have hA := h.measurable_action + have hY := h.measurable_feedback + extend_space! (gaussianReal 0 1) using P with U hU hind + have hAY : IndepFun (trajectory A Y) U P := + hind.comp (φ := fun (p : (ℕ → Fin K) × (ℕ → ℝ)) (n : ℕ) ↦ (p.1 n, p.2 n)) (by fun_prop) + measurable_id + exact ⟨Ω₀, inferInstance, P, inferInstance, A, Y, U, h, hU, hAY⟩ + +/-- **The explicit form, `extend_space_map`.** The goal mentions the space through `P` and `A 0`; +`transfer` moves it to the new space, with the measurability of the sequence taken from `h`. In +the extended goal, `transfer hf at h` pulls the sequence back. -/ +example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq A Y (alg hK) env P) : + P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + have hA := h.measurable_action + have hY := h.measurable_feedback + extend_space_map (gaussianReal 0 1) with Ω' P' f hf U hU hind + transfer hf at h + exact h.hasLaw_action_zero.map_eq + +end Test.AlgTrace + +end + +end diff --git a/Test/Extend.lean b/Test/Extend.lean new file mode 100644 index 0000000..fbfe50b --- /dev/null +++ b/Test/Extend.lean @@ -0,0 +1,473 @@ +module + +public import Mathlib.Probability.Independence.InfinitePi +public import RandomDo.Probability.Extend +public import RandomDo.Probability.MeasurePreserving + +set_option linter.style.header false + +/-! +# The `extend_space` tactic + +The first sections pin down what `extend_space` produces: the context after the extension, what +is transported and what is left about the old space, when the `transfer` obligation is closed +automatically and when it is left, what `extend_space!` clears and what it has to keep, and what a +second extension looks like. Then come the explicit form `extend_space_map`, a draw with a +conditional law, an i.i.d. sequence, and the errors the tactic reports. + +Throughout, `Ω` lives in `Type u` and `E` in `Type`: the tactic lifts the product to the universe +of `Ω`. +-/ + +open MeasureTheory ProbabilityTheory RDo + +@[expose] public section + +noncomputable section + +namespace Test.Extend + +universe u + +variable {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + {E : Type} [MeasurableSpace E] (μ : Measure E) [IsProbabilityMeasure μ] + +include μ + +/-! ## The context after `extend_space` -/ + +/-- The names are kept: `Ω`, `P`, `X`, `A` and `s` now live on the extended space, related to the +old `Ω₀`, `P₀`, `X₀`, … by the map `f` and the defining equations. Hypotheses about them are +transported, and `Z` is independent of the tuple of the random variables and events. The goal +reads as before, and its `transfer` obligation is discharged. -/ +example (X : Ω → ℝ) (A : ℕ → Ω → ℝ) (s : Set Ω) (hX : Measurable X) (hA : ∀ n, Measurable (A n)) + (hs : MeasurableSet s) (ν : Measure ℝ) (c : ENNReal) (h1 : P.map X = ν) + (h2 : ∀ n, P.map (A n) = ν) (h3 : P s = c) : + P.map X = ν ∧ (∀ n, P.map (A n) = ν) ∧ P s = c := by + extend_space μ with Z hZ hind f hf + guard_hyp hf : MeasurePreserving f P P₀ + guard_hyp hZm : Measurable Z + guard_hyp hZ : HasLaw Z μ P + guard_hyp hind : IndepFun (fun ω ↦ (X ω, fun n ↦ A n ω, ω ∈ s)) Z P + guard_hyp hX_def : ∀ ω, X₀ (f ω) = X ω + guard_hyp hA_def : ∀ n ω, A₀ n (f ω) = A n ω + guard_hyp hs_def : f ⁻¹' s₀ = s + guard_hyp hX : Measurable X + guard_hyp hA : ∀ n, Measurable (A n) + guard_hyp hs : MeasurableSet s + guard_hyp h1 : P.map X = ν + guard_hyp h2 : ∀ n, P.map (A n) = ν + guard_hyp h3 : P s = c + guard_hyp h1₀ : P₀.map X₀ = ν + guard_target =ₐ P.map X = ν ∧ (∀ n, P.map (A n) = ν) ∧ P s = c + exact ⟨h1, h2, h3⟩ + +/-- Without `with`, the names are `Z hZ hind f hf`. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space μ + guard_hyp hZ : HasLaw Z μ P + guard_hyp hind : IndepFun X Z P + guard_hyp hf : MeasurePreserving f P P₀ + exact hXν.map_eq + +/-- A hypothesis nothing transports, here about `Measure.restrict`, stays about the old space, +under its `₀` name. -/ +example (s : Set Ω) (hs : MeasurableSet s) (h : P.restrict s Set.univ = 1) : + P s = 1 := by + extend_space μ + guard_hyp h₀ : P₀.restrict s₀ Set.univ = 1 + guard_hyp hs : MeasurableSet s + rw [← hs_def, hf.measure_preimage hs₀.nullMeasurableSet, ← Measure.restrict_apply_univ] + exact h₀ + +/-- A statement `transfer` has no lemma for, here `Measure.restrict`, leaves the obligation. It is +stated for an arbitrary measure-preserving map, with the original goal as its conclusion. -/ +example (s : Set Ω) : P.restrict s Set.univ = P s := by + extend_space μ + case extended => + guard_target =ₐ P.restrict s Set.univ = P s + exact Measure.restrict_apply_univ _ + case transfer => + guard_target =ₐ ∀ (Ω' : Type u) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → + P'.restrict (f ⁻¹' s) Set.univ = P' (f ⁻¹' s) → P.restrict s Set.univ = P s + intro Ω' _ P' _ f hf h + exact Measure.restrict_apply_univ _ + +/-- The goal depends on the measurability proof `hX`: it is generalized along with the goal, and +comes back about the new `X`. The goal mentions no measure, so `using P` says which space to +extend, and `transfer` cannot rewrite under a binder the goal depends on, so the obligation is +left. -/ +example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel ℝ E) [IsMarkovKernel κ] : + IsMarkovKernel (κ.comap X hX) := by + extend_space μ using P + case extended => + guard_hyp hX : Measurable X + guard_hyp hX₀ : Measurable X₀ + guard_target =ₐ IsMarkovKernel (κ.comap X hX) + infer_instance + case transfer => + intro Ω' _ P' _ f hf h hX + infer_instance + +/-! ## `extend_space!` -/ + +/-- The old space is cleared: nothing mentions `Ω₀`, `f` or the defining equations any more. The +space is a binder of the statement here rather than a `variable`: Lean does not let a tactic clear +a `variable`, so those would stay. -/ +example {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space! μ + fail_if_success guard_hyp Ω₀ : Type u + fail_if_success guard_hyp f : Ω → Ω₀ + fail_if_success guard_hyp hX_def : ∀ ω, X₀ (f ω) = X ω + guard_hyp hZ : HasLaw Z μ P + guard_hyp hind : IndepFun X Z P + guard_hyp hXν : HasLaw X ν P + exact hXν.map_eq + +/-- `extend_space! κ` for a kernel on `Ω`: `hZ` is stated given the map, so the map stays with +`hf`, and the old space with its instances. The rest is cleared. -/ +example {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel Ω E) [IsMarkovKernel κ] (ν : Measure ℝ) + (hXν : HasLaw X ν P) : P.map X = ν := by + extend_space! κ + guard_hyp hZ : HasCondDistrib Z f κ₀ P + guard_hyp hf : MeasurePreserving f P P₀ + fail_if_success guard_hyp hX₀ : Measurable X₀ + fail_if_success guard_hyp hX_def : ∀ ω, X₀ (f ω) = X ω + have hκ : IsMarkovKernel κ₀ := inferInstance + have hP₀ : IsProbabilityMeasure P₀ := inferInstance + clear hκ hP₀ + exact hXν.map_eq + +/-- With only an event in the goal, `hind` is the independence of `Z` and `fun ω ↦ ω ∈ s`, and +the old space is cleared entirely. -/ +example {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (s : Set Ω) (hs : MeasurableSet s) (c : ENNReal) (h : P s = c) : P s = c := by + extend_space! μ + guard_hyp hind : IndepFun (fun ω ↦ ω ∈ s) Z P + fail_if_success guard_hyp f : Ω → Ω₀ + fail_if_success guard_hyp Ω₀ : Type u + exact h + +/-! ## Extending twice + +The first draw `Z` and its hypotheses `hZm`, `hZ` and `hind` are transported like any other random +variable, and the second extension names its objects with a prime so as not to shadow them. The +space left by the first extension is `Ω₀`, the original one `Ω₀₀`; `hf'` relates `P` to `P₀` +through the new map `f'`, and `hf` to `P₀₀` through the transported map `f`, which is now the +composite. -/ + +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space μ + extend_space μ + guard_hyp hZm : Measurable Z + guard_hyp hZ : HasLaw Z μ P + guard_hyp hind : IndepFun X Z P + guard_hyp hZm' : Measurable Z' + guard_hyp hZ' : HasLaw Z' μ P + guard_hyp hind' : IndepFun (fun ω ↦ (f ω, Z ω, X ω)) Z' P + guard_hyp hf' : MeasurePreserving f' P P₀ + guard_hyp hf : MeasurePreserving f P P₀₀ + guard_hyp hX_def : ∀ ω, X₀ (f' ω) = X ω + guard_hyp hX_def₀ : ∀ ω, X₀₀ (f₀ ω) = X₀ ω + exact hXν.map_eq + +/-- With `extend_space!`, nothing of the intermediate space is left. -/ +example {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space! μ + extend_space! μ + guard_hyp hZ : HasLaw Z μ P + guard_hyp hZ' : HasLaw Z' μ P + guard_hyp hind : IndepFun X Z P + guard_hyp hind' : IndepFun (fun ω ↦ (Z ω, X ω)) Z' P + fail_if_success guard_hyp Ω₀ : Type u + exact hXν.map_eq + +/-! ## Transported hypotheses + +Laws, independence, conditional laws, events, integrals and almost-everywhere statements are +transported by the `@[transfer]` lemmas, measurability and the like by the `@[transfer_forward]` +lemmas. -/ + +/-- Independence. -/ +example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) (hXY : IndepFun X Y P) : + IndepFun X Y P := by + extend_space μ + guard_hyp hXY : IndepFun X Y P + guard_hyp hind : IndepFun (fun ω ↦ (X ω, Y ω)) Z P + exact hXY + +/-- A hypothesis introduced by `have`, whose type may still carry metavariables, is transported +like any other. -/ +example (X : Ω → ℝ) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + P.map X = ν := by + have hXae := hXν.aemeasurable + extend_space μ + guard_hyp hXae : AEMeasurable X P + guard_hyp hXae₀ : AEMeasurable X₀ P₀ + exact hXν.map_eq + +/-- A conditional law. -/ +example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) (κ : Kernel ℝ ℝ) + (hXY : HasCondDistrib Y X κ P) : + HasCondDistrib Y X κ P := by + extend_space μ + exact hXY + +/-- An event given by a set-builder expression, and its real-valued measure. -/ +example (X : Ω → ℝ) (hX : Measurable X) (h : P {ω | 0 < X ω} = 1 / 2) + (h' : P.real {ω | 0 < X ω} = 1 / 2) : + P {ω | 0 < X ω} = 1 / 2 ∧ P.real {ω | 0 < X ω} = 1 / 2 := by + extend_space μ + guard_hyp h : P {ω | 0 < X ω} = 1 / 2 + exact ⟨h, h'⟩ + +/-- An integral, with an integrability and an almost-everywhere hypothesis. -/ +example (X : Ω → ℝ) (hX : Measurable X) (hint : Integrable X P) (h : ∀ᵐ ω ∂P, 0 ≤ X ω) : + 0 ≤ ∫ ω, X ω ∂P ∧ Integrable X P := by + extend_space μ + guard_hyp hint : Integrable X P + guard_hyp h : ∀ᵐ ω ∂P, 0 ≤ X ω + exact ⟨integral_nonneg_of_ae h, hint⟩ + +/-- A Lebesgue integral. -/ +example (X : Ω → ℝ) (hX : Measurable X) (c : ENNReal) (h : ∫⁻ ω, ‖X ω‖ₑ ∂P = c) : + ∫⁻ ω, ‖X ω‖ₑ ∂P = c := by + extend_space μ + exact h + +/-- An almost-everywhere equality. -/ +example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) (h : X =ᵐ[P] Y) : X =ᵐ[P] Y := by + extend_space μ + exact h + +/-- Mutual independence of a family. -/ +example (A : ℕ → Ω → ℝ) (hA : ∀ n, Measurable (A n)) (h : iIndepFun A P) (ν : Measure ℝ) + (h0 : HasLaw (A 0) ν P) : P.map (A 0) = ν ∧ iIndepFun A P := by + extend_space μ + guard_hyp h : iIndepFun A P + guard_hyp hind : IndepFun (fun ω n ↦ A n ω) Z P + exact ⟨h0.map_eq, h⟩ + +/-- A random variable whose measurability is not at hand is left out of the tuple of `hind`. When +nothing is left, `hind` is the independence of `Z` and the map. -/ +example (X : Ω → ℝ) (ν : Measure ℝ) (hXν : HasLaw X ν P) : P.map X = ν := by + extend_space μ + guard_hyp hind : IndepFun f Z P + exact hXν.map_eq + +/-! ## Using the new draw + +A statement that does not mention the space has a trivial `transfer` obligation: this is the +existential form of the tactic. -/ + +/-- Any random variable has an independent companion with any prescribed law, on a larger space: +after the extension, that space is `Ω` itself. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + ∃ (Ω' : Type u) (_ : MeasurableSpace Ω') (P' : Measure Ω') (_ : IsProbabilityMeasure P') + (X' : Ω' → ℝ) (Z : Ω' → E), HasLaw X' ν P' ∧ HasLaw Z μ P' ∧ IndepFun X' Z P' := by + extend_space μ using P + exact ⟨Ω, inferInstance, P, inferInstance, X, Z, hXν, hZ, hind⟩ + +/-- An i.i.d. sequence, by extending with `Measure.infinitePi`: `Z ω : ℕ → E`. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + ∃ (Ω' : Type u) (_ : MeasurableSpace Ω') (P' : Measure Ω') (_ : IsProbabilityMeasure P') + (X' : Ω' → ℝ) (Z : ℕ → Ω' → E), HasLaw X' ν P' ∧ (∀ n, HasLaw (Z n) μ P') + ∧ iIndepFun Z P' ∧ IndepFun X' (fun ω n ↦ Z n ω) P' := by + extend_space Measure.infinitePi (fun _ : ℕ ↦ μ) using P + have hZn (n : ℕ) : HasLaw (fun ω ↦ Z ω n) μ P := + (measurePreserving_eval_infinitePi _ n).hasLaw.comp hZ + exact ⟨Ω, inferInstance, P, inferInstance, X, fun n ω ↦ Z ω n, hXν, hZn, + (iIndepFun_iff_hasLaw_Pi_infinitePi hZn hZ.aemeasurable).2 hZ, hind⟩ + +/-! ## A draw with a conditional law -/ + +/-- `extend_space (κ.comap X hX)`: the draw has conditional law `κ` given `X`. -/ +example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel ℝ E) [IsMarkovKernel κ] (ν : Measure ℝ) + (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space (κ.comap X hX) with Z hZ + guard_hyp hZ : HasCondDistrib Z X κ P + exact hXν.map_eq + +/-- `extend_space κ` for a kernel on `Ω` itself: the kernel is on the old space, now `κ₀`, and +the conditional law is given the map `f`. -/ +example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel Ω E) [IsMarkovKernel κ] (ν : Measure ℝ) + (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space κ with Z hZ f hf + guard_hyp hZ : HasCondDistrib Z f κ₀ P + exact hXν.map_eq + +/-- A kernel conditioned on a pair of random variables: `hZ` is stated given that pair. -/ +example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) (κ : Kernel (ℝ × ℝ) E) + [IsMarkovKernel κ] (ν : Measure ℝ) (hXν : HasLaw X ν P) : P.map X = ν := by + extend_space (κ.comap (fun ω ↦ (X ω, Y ω)) (hX.prodMk hY)) + guard_hyp hZ : HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) κ P + exact hXν.map_eq + +/-! ## The explicit form, `extend_space_map` -/ + +/-- Nothing is renamed: the goal is restated on `Ω'`, with `X ∘ f` for `X` and `f ⁻¹' s` for `s`, +and the hypotheses stay about the old space. With the measurability hypotheses around, the +`transfer` obligation is discharged. -/ +example (X : Ω → ℝ) (A : ℕ → Ω → ℝ) (s : Set Ω) (hX : Measurable X) (hA : ∀ n, Measurable (A n)) + (hs : MeasurableSet s) (ν : Measure ℝ) (c : ENNReal) (h1 : P.map X = ν) + (h2 : ∀ n, P.map (A n) = ν) (h3 : P s = c) : + P.map X = ν ∧ (∀ n, P.map (A n) = ν) ∧ P s = c := by + extend_space_map μ with Ω' P' f hf Z hZ hind + guard_hyp hf : MeasurePreserving f P' P + guard_hyp hZm : Measurable Z + guard_hyp hZ : HasLaw Z μ P' + guard_hyp hind : IndepFun f Z P' + guard_hyp h1 : P.map X = ν + guard_target =ₐ P'.map (fun ω ↦ X (f ω)) = ν ∧ (∀ n, P'.map (fun ω ↦ A n (f ω)) = ν) + ∧ P' (f ⁻¹' s) = c + transfer hf at h1 h2 h3 + guard_hyp h1 : P'.map (fun ω ↦ X (f ω)) = ν + exact ⟨h1, h2, h3⟩ + +/-- Without `with`, the names are `Ω' P' f hf Z hZ hind`. A hypothesis is pulled back by +`transfer hf at h`, or by hand: `hXν.comp hf.hasLaw` for a law, `hind.comp hX measurable_id` for +independence. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space_map μ + have hXν' : HasLaw (fun ω ↦ X (f ω)) ν P' := hXν.comp hf.hasLaw + have hind' : IndepFun (fun ω ↦ X (f ω)) Z P' := hind.comp hX measurable_id + transfer hf at hX + guard_hyp hX : Measurable fun ω ↦ X (f ω) + exact hXν'.map_eq + +/-- The goal depends on `hX`: it is generalized and reintroduced as `hX'`, about `X ∘ f`, while +`hX` itself stays. -/ +example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel ℝ E) [IsMarkovKernel κ] : + IsMarkovKernel (κ.comap X hX) := by + extend_space_map μ using P + case extended => + guard_hyp hX : Measurable X + guard_hyp hX' : Measurable fun ω ↦ X (f ω) + guard_target =ₐ IsMarkovKernel (κ.comap (fun ω ↦ X (f ω)) hX') + infer_instance + case transfer => + intro Ω' _ P' _ f hf h hX + infer_instance + +/-- A draw with a conditional law: `HasCondDistrib.comp_right` reads `hZ` given `X ∘ f`. -/ +example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel ℝ E) [IsMarkovKernel κ] (ν : Measure ℝ) + (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space_map (κ.comap X hX) with Ω' P' f hf Z hZ + guard_hyp hZ : HasCondDistrib Z f (κ.comap X hX) P' + have hZ' : HasCondDistrib Z (fun ω ↦ X (f ω)) κ P' := hZ.comp_right + transfer hf at hXν + exact hXν.map_eq + +/-! ## Universes -/ + +/-- `Ω` and `E` in the same universe. -/ +example {E' : Type u} [MeasurableSpace E'] (μ' : Measure E') [IsProbabilityMeasure μ'] + (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : P.map X = ν := by + extend_space μ' + exact hXν.map_eq + +/-! ## Errors -/ + +/-- +error: extend_space: the goal mentions no measure on a local space; name one with `using` +-/ +#guard_msgs in +example (X : Ω → ℝ) : X = X := by + extend_space μ + +/-- +error: extend_space: the goal mentions several measures, [P, Q]; choose one with `using` +-/ +#guard_msgs in +example (Q : Measure Ω) (X : Ω → ℝ) : P.map X = Q.map X := by + extend_space μ + +/-- +error: extend_space: E' lives in universe u + 1 and Ω in universe u, +so the product Ω × E' does not live in the universe of Ω: lift E' with `ULift` +-/ +#guard_msgs in +example {E' : Type (u + 1)} [MeasurableSpace E'] (μ' : Measure E') [IsProbabilityMeasure μ'] + (X : Ω → ℝ) : P.map X = P.map X := by + extend_space μ' + +/-- +error: extend_space: the goal depends on + Q +of type + Measure Ω +which is neither a random variable nor an event on Ω, so it cannot be transported to the extended space +-/ +#guard_msgs in +example (Q : Measure Ω) (X : Ω → ℝ) : P.map X = Q.map X := by + extend_space μ using P + +/-- +error: extend_space: the goal depends on the local definition Y, which cannot be transported to the extended space; unfold it or `clear_value` it first +-/ +#guard_msgs in +example (X : Ω → ℝ) (ν : Measure ℝ) : ∃ Y : Ω → ℝ, P.map Y = ν := by + let Y : Ω → ℝ := X + refine ⟨Y, ?_⟩ + extend_space μ + +/-- +error: extend_space: the kernel κ is on ℝ, not on Ω +-/ +#guard_msgs in +example (X : Ω → ℝ) (κ : Kernel ℝ E) [IsMarkovKernel κ] (ν : Measure ℝ) : P.map X = ν := by + extend_space κ using P + +/-- +error: extend_space: X is not a measure +-/ +#guard_msgs in +example (X : Ω → ℝ) (ν : Measure ℝ) : P.map X = ν := by + extend_space μ using X + +/-- +error: extend_space: at most 5 names may be given +-/ +#guard_msgs in +example (X : Ω → ℝ) (ν : Measure ℝ) : P.map X = ν := by + extend_space μ with a b c d e g + +/-- +error: extend_space: the space must be a local hypothesis, but ℕ → ℝ is not +-/ +#guard_msgs in +example (Q : Measure (ℕ → ℝ)) [IsProbabilityMeasure Q] (X : (ℕ → ℝ) → ℝ) (ν : Measure ℝ) : + Q.map X = ν := by + extend_space μ using Q + +/-- +error: extend_space: Q is not known to be a probability measure: no `IsProbabilityMeasure` instance was found +-/ +#guard_msgs in +example {Q : Measure Ω} (X : Ω → ℝ) (ν : Measure ℝ) : Q.map X = ν := by + extend_space μ + +/-- +error: extend_space: Measure.map Y P is not known to be a probability measure: no `IsProbabilityMeasure` instance was found +-/ +#guard_msgs in +example (X Y : Ω → ℝ) (ν : Measure ℝ) : P.map X = ν := by + extend_space (P.map Y) + +end Test.Extend + +end + +end diff --git a/Test/Gaps.lean b/Test/Gaps.lean index 39aeb1c..c272fad 100644 --- a/Test/Gaps.lean +++ b/Test/Gaps.lean @@ -31,6 +31,7 @@ error: No `ControlInfo` inference handler found for `RDo.rdoFor` in syntax Register a handler with `@[doElem_control_info RDo.rdoFor]`. -/ #guard_msgs (whitespace := lax) in +/-- Two nested `for` loops. -/ def nestedLoops (xs ys : List ℕ) : IdM ℕ := rdo let mut s := 0 for x in xs rdo @@ -52,6 +53,7 @@ error: failed to synthesize instance of type class Hint: Type class instance resolution failures can be inspected with the `set_option trace.Meta.synthInstance true` command. -/ #guard_msgs in +/-- A `while` loop. -/ def whileLoop : IdM ℕ := rdo let mut i := 0 while i < 3 do @@ -70,6 +72,7 @@ error: failed to synthesize instance of type class Hint: Type class instance resolution failures can be inspected with the `set_option trace.Meta.synthInstance true` command. -/ #guard_msgs in +/-- A `try … catch` block. -/ noncomputable def tryCatch : Measure Bool := rdo try let x ← fairCoin @@ -90,6 +93,7 @@ error: failed to synthesize instance of type class Hint: Type class instance resolution failures can be inspected with the `set_option trace.Meta.synthInstance true` command. -/ #guard_msgs in +/-- A `for` loop over a range. -/ def overRange : IdM ℕ := rdo let mut s := 0 for _ in [0:3] rdo @@ -103,6 +107,7 @@ error: failed to synthesize instance of type class Hint: Type class instance resolution failures can be inspected with the `set_option trace.Meta.synthInstance true` command. -/ #guard_msgs in +/-- A `for` loop over a `Finset`. -/ def overFinset : IdM ℕ := rdo let mut s := 0 for _ in Finset.range 3 rdo diff --git a/Test/Instances.lean b/Test/Instances.lean index 1105bdb..ac12b12 100644 --- a/Test/Instances.lean +++ b/Test/Instances.lean @@ -68,11 +68,13 @@ def sumOver (xs : List ℕ) (f : ℕ → m ℕ) : m ℕ := rdo example : IdM.run (sumOver [1, 2, 3] (fun x ↦ ((x * 2 : ℕ) : IdM ℕ))) = 12 := rfl +/-- At `Measure`, summing terms that are each `x` or `0` on a fair coin. -/ noncomputable def sumOverMeasure : Measure ℕ := sumOver (m := Measure) [1, 2, 3] (fun x ↦ rdo let b ← fairCoin return (if b then x else 0)) +/-- At `PseudoRandomM`, the same program as an executable sampler. -/ def sumOverRandom : PseudoRandomM ℕ := sumOver (m := PseudoRandomM) [1, 2, 3] (fun x ↦ rdo let b ← Random.randBool diff --git a/Test/IsMarkov.lean b/Test/IsMarkov.lean index a63e188..f93c274 100644 --- a/Test/IsMarkov.lean +++ b/Test/IsMarkov.lean @@ -19,6 +19,7 @@ namespace Test.IsMarkov /-! ## `return` -/ +/-- A `return` alone. -/ noncomputable def shiftBy (c : ℝ) : Measure ℝ := rdo return c + 1 @@ -26,6 +27,7 @@ example : IsMarkov shiftBy := by is_markov /-! ## `let x ← _` -/ +/-- Two Gaussian draws, summed. -/ noncomputable def sumTwo : Measure ℝ := rdo let x ← gaussianReal 0 1 let y ← gaussianReal 0 1 @@ -50,6 +52,7 @@ example (μ : Measure ℝ) [IsProbabilityMeasure μ] : IsMarkov fun _ : ℝ ↦ /-! ## `if … then … else` between two families -/ +/-- One of two families, chosen on the sign of the argument. -/ noncomputable def branchOn (c : ℝ) : Measure ℝ := rdo if 0 < c then let x ← gaussianReal c 1 @@ -62,6 +65,7 @@ example : IsMarkov branchOn := by is_markov /-! ## `for` over a fixed collection -/ +/-- A sum of Gaussian draws accumulated by a loop. -/ noncomputable def sumLoop : Measure ℝ := rdo let mut s : ℝ := 0 for _ in List.range 3 rdo @@ -73,6 +77,7 @@ example : IsProbabilityMeasure sumLoop := by is_markov /-! ## `for` with an early `return`, which goes through `Break.runK` -/ +/-- The first positive draw among three, or `0`: a loop with an early `return`. -/ noncomputable def firstPositive : Measure ℝ := rdo for _ in List.range 3 rdo let x ← gaussianReal 0 1 @@ -84,6 +89,7 @@ example : IsProbabilityMeasure firstPositive := by is_markov /-! ## `for` over a collection read off the argument -/ +/-- A random walk driven by the argument: a loop over a collection read off it. -/ noncomputable def overList (xs : List ℝ) : Measure ℝ := rdo let mut s : ℝ := 0 for x in xs rdo @@ -95,8 +101,10 @@ example : IsMarkov overList := by is_markov /-! ## Looking through definitions, and the `fuel` argument -/ +/-- `sumTwo` behind one definition. -/ noncomputable def layerOne : Measure ℝ := sumTwo +/-- `sumTwo` behind two definitions. -/ noncomputable def layerTwo : Measure ℝ := layerOne example : IsProbabilityMeasure layerTwo := by is_markov diff --git a/Test/Transfer.lean b/Test/Transfer.lean new file mode 100644 index 0000000..215eec7 --- /dev/null +++ b/Test/Transfer.lean @@ -0,0 +1,203 @@ +module + +public import RandomDo.Probability.MeasurePreserving +public import RandomDo.Probability.Transfer + +set_option linter.style.header false + +/-! +# The `transfer` tactic on its own + +`transfer` is mostly run by `extend_space` and `alg_env_trace` on the obligations they leave, and +is tested with them. Here it is used directly, in each of its forms: on the goal, on hypotheses, +where the rewriting by the `@[transfer]` lemmas and the fallback on the `@[transfer_forward]` +lemmas both show, on both at once, and on a `transfer` obligation. Then come the errors: a target +about the old space that the tactic cannot move is an error, not a silent no-op. +-/ + +open MeasureTheory ProbabilityTheory + +@[expose] public section + +noncomputable section + +namespace Test.Transfer + +universe u + +variable {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} {Ω' : Type u} [MeasurableSpace Ω'] + {P' : Measure Ω'} {f : Ω' → Ω} + +section Map + +variable (hf : MeasurePreserving f P' P) +include hf + +/-! ## On the goal -/ + +/-- `transfer hf` on a goal: the goal is moved to the new space and closed by the hypothesis. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (h : HasLaw (fun ω ↦ X (f ω)) ν P') : + HasLaw X ν P := by + transfer hf + +/-- The side conditions are discharged by `fun_prop`, here the strong measurability of a product. +Once rewritten, the two sides agree and `simp` closes the goal. -/ +example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) : + ∫ ω, X ω * Y ω ∂P = ∫ ω, X (f ω) * Y (f ω) ∂P' := by + transfer hf + +/-- Events: an event `X ⁻¹' t` given by a random variable becomes `(fun ω ↦ X (f ω)) ⁻¹' t`, as +`extend_space` writes it, and a set-builder event has the map pushed inside. -/ +example (X : Ω → ℝ) (hX : Measurable X) (t : Set ℝ) (ht : MeasurableSet t) (c c' : ENNReal) + (h : P' ((fun ω ↦ X (f ω)) ⁻¹' t) = c) (h' : P' {ω | 0 < X (f ω)} = c') : + P (X ⁻¹' t) = c ∧ P {ω | 0 < X ω} = c' := by + transfer hf + guard_target =ₐ P' ((fun ω ↦ X (f ω)) ⁻¹' t) = c ∧ P' {ω | 0 < X (f ω)} = c' + exact ⟨h, h'⟩ + +/-! ## On hypotheses -/ + +/-- `transfer hf at h` rewrites with the `@[transfer]` lemmas when it can, and falls back on the +`@[transfer_forward]` lemmas otherwise. -/ +example (X : Ω → ℝ) (hX : Measurable X) (s : Set Ω) (hs : MeasurableSet s) (h : P s = 1) : + P' (f ⁻¹' s) = 1 ∧ MeasurableSet (f ⁻¹' s) ∧ Measurable fun ω ↦ X (f ω) := by + transfer hf at hX hs h + guard_hyp hX : Measurable fun ω ↦ X (f ω) + guard_hyp hs : MeasurableSet (f ⁻¹' s) + guard_hyp h : P' (f ⁻¹' s) = 1 + exact ⟨h, hs, hX⟩ + +/-- Independence, rewritten with the measurability hypotheses as side conditions. -/ +example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) (hXY : IndepFun X Y P) : + IndepFun (fun ω ↦ X (f ω)) (fun ω ↦ Y (f ω)) P' := by + transfer hf at hXY + exact hXY + +/-- A definition is transported as itself. Without measurability hypotheses no `@[transfer]` lemma +rewrites a conditional law, and the `@[transfer_forward]` lemma that transports it is the one about +`HasCondDistrib`, not the one about the `HasLaw` it unfolds to. -/ +example (X Y : Ω → ℝ) (κ : Kernel ℝ ℝ) (h : HasCondDistrib Y X κ P) : + HasCondDistrib (fun ω ↦ Y (f ω)) (fun ω ↦ X (f ω)) κ P' := by + transfer hf at h + guard_hyp h : HasCondDistrib (fun ω ↦ Y (f ω)) (fun ω ↦ X (f ω)) κ P' + exact h + +/-- A hypothesis with binders is transported under them, whether by rewriting or by a +`@[transfer_forward]` lemma. -/ +example (A : ℕ → Ω → ℝ) (hA : ∀ n, Measurable (A n)) (ν : Measure ℝ) + (h : ∀ n, HasLaw (A n) ν P) : + ∀ n, HasLaw (fun ω ↦ A n (f ω)) ν P' := by + transfer hf at h hA + guard_hyp hA : ∀ n, Measurable fun ω ↦ A n (f ω) + guard_hyp h : ∀ n, HasLaw (fun ω ↦ A n (f ω)) ν P' + exact h + +/-- Mutual independence of a family. -/ +example (A : ℕ → Ω → ℝ) (hA : ∀ n, Measurable (A n)) (h : iIndepFun A P) : + iIndepFun (fun n ω ↦ A n (f ω)) P' := by + transfer hf at h + exact h + +/-! ## On the goal and hypotheses at once + +The goal goes first: transporting `hX` destroys the measurability of `X`, which the side condition +of the goal needs. -/ + +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (h : HasLaw (fun ω ↦ X (f ω)) ν P') : + P.map X = ν := by + transfer hf at hX ⊢ + guard_hyp hX : Measurable fun ω ↦ X (f ω) + guard_target =ₐ P'.map (fun ω ↦ X (f ω)) = ν + exact h.map_eq + +/-- `transfer hf at *` transfers what it can and leaves alone what is not about the old space: a +hypothesis about the new space, the map itself, a hypothesis about neither. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (h : HasLaw X ν P) (Z : Ω' → ℝ) + (hZ : HasLaw Z ν P') (n m : ℕ) (hn : n < m) : + P.map X = ν ∧ HasLaw Z ν P' ∧ n < m := by + transfer hf at * + guard_hyp hX : Measurable fun ω ↦ X (f ω) + guard_hyp h : HasLaw (fun ω ↦ X (f ω)) ν P' + guard_hyp hZ : HasLaw Z ν P' + guard_hyp hf : MeasurePreserving f P' P + guard_hyp hn : n < m + guard_target =ₐ P'.map (fun ω ↦ X (f ω)) = ν ∧ HasLaw Z ν P' ∧ n < m + exact ⟨h.map_eq, hZ, hn⟩ + +end Map + +/-! ## The obligation of `extend_space` -/ + +/-- `transfer` alone: the new space, the map and the statement on the new space are introduced, +the goal is transferred and closed with that statement. -/ +example (s : Set Ω) (hs : MeasurableSet s) : + ∀ (Ω' : Type u) [MeasurableSpace Ω'] (P' : Measure Ω') [IsProbabilityMeasure P'] (f : Ω' → Ω), + MeasurePreserving f P' P → P' (f ⁻¹' s) = 1 → P s = 1 := by + transfer + +/-! ## Errors -/ + +/-- +error: transfer: the goal is not a `transfer` obligation. It should have the form + ∀ Ω' [MeasurableSpace Ω'] (P' : Measure Ω') [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → T' → T +but is + True +-/ +#guard_msgs in +example : True := by + transfer + +/-- +error: transfer: could not close the goal after transferring it: + P' (f ⁻¹' s) = 1 +with the statement on the new space: + P' (f ⁻¹' s) = 2 +-/ +#guard_msgs in +example (s : Set Ω) (hs : MeasurableSet s) : + ∀ (Ω' : Type u) [MeasurableSpace Ω'] (P' : Measure Ω') [IsProbabilityMeasure P'] (f : Ω' → Ω), + MeasurePreserving f P' P → P' (f ⁻¹' s) = 2 → P s = 1 := by + transfer + +/-- +error: transfer: `at` needs the map to transfer along, as in `transfer hf at h` +-/ +#guard_msgs in +example : True → True := by + intro _h + transfer at _h + +-- A hypothesis about the old space that nothing transfers is an error, not a silent no-op. +/-- +error: transfer: nothing transfers the hypothesis h : + (P.restrict s) Set.univ = 1 +No `@[transfer]` lemma rewrites it and no `@[transfer_forward]` lemma transports it. +A side condition, such as the measurability of a random variable, may not have been discharged. +-/ +#guard_msgs in +example (s : Set Ω) (h : P.restrict s Set.univ = 1) (hf : MeasurePreserving f P' P) : True := by + transfer hf at h + +-- Without the measurability of `X`, the side condition fails and the goal cannot be rewritten: +-- an error, and nothing the discharger tried leaks into the messages. +/-- +error: transfer: no `@[transfer]` lemma rewrites the goal + HasLaw X ν P +A side condition, such as the measurability of a random variable, may not have been discharged. +-/ +#guard_msgs in +example (X : Ω → ℝ) (ν : Measure ℝ) (hf : MeasurePreserving f P' P) : HasLaw X ν P := by + transfer hf + +/-- +error: transfer: nothing to transfer +-/ +#guard_msgs in +example (hf : MeasurePreserving f P' P) (n m : ℕ) (hn : n < m) : n < m := by + transfer hf at * + +end Test.Transfer + +end + +end diff --git a/notes/TRACE_SEMANTICS.md b/notes/TRACE_SEMANTICS.md new file mode 100644 index 0000000..217d649 --- /dev/null +++ b/notes/TRACE_SEMANTICS.md @@ -0,0 +1,338 @@ +# Random variables for `rdo` programs + +A way to get from "this `rdo` program denotes a measure" to "these are its random variables, this +is their joint law, here is the conditional law of each draw given the ones before it, and here is +which of them are independent". + +Implemented in [RandomDo/Probability/Trace.lean](../RandomDo/Probability/Trace.lean) (the theory) +and [RandomDo/Probability/Tactic.lean](../RandomDo/Probability/Tactic.lean) (the `rdo_trace` and +`rdo_peel` tactics), demonstrated in +[RandomDo/Probability/Examples.lean](../RandomDo/Probability/Examples.lean) and, on `thompson`, in +[RandomDo/Probability/Thompson.lean](../RandomDo/Probability/Thompson.lean). Everything below that +is described as *done* compiles with no `sorry`. + +## The idea + +A measure has no random variables. So give the program a probability space of its own: the space of +its **traces**, one coordinate per `←`. + +```lean +structure HasTrace (prog : γ → Measure β) (P : Kernel γ Ω) (out : γ × Ω → β) : Prop where + measurable_out : Measurable out + map_eq (c : γ) : (P c).map (fun ω ↦ out (c, ω)) = prog c +``` + +* `γ` — the program's parameters (`hist` in `thompson`, `Unit` for a closed program). +* `Ω` — the trace space: the product of the types drawn at each `←`. +* `P` — the joint law of all the draws, as a Markov kernel in the parameters. +* `out` — a **deterministic** readout reconstructing the program's result from parameters + draws. + +`map_eq` says: running the program = drawing a trace and reading the answer off it. Every random +variable of interest is now a function on `(Ω, P c)`, and the program's own result is one of them +(`HasTrace.hasLaw_out`). + +The whole payoff comes from *how* `P` is built. Each `←` contributes one `Kernel.compProd` factor, +so the trace kernel is a right-nested `⊗ₖ` whose shape is literally the shape of the program: + +``` +rdo P = κ ⊗ₖ (η ⊗ₖ θ) on Ω = ℝ × (ℝ × ℝ) + let x ← κ c + let y ← η (c, x) + let z ← θ ((c, x), y) + return x + y + z out (c, ω) = ω.1 + ω.2.1 + ω.2.2 +``` + +## Program construct → trace combinator + +| `rdo` construct | elaborated head | rule | trace kernel built | +|---|---|---|---| +| `let x ← p; q` | `mBind` | `HasTrace.bind` | `P ⊗ₖ Q` | +| `let x ← p; return f x` | `mBind`/`mPure` | `HasTrace.bindPure` | `P` (no new draw) | +| `return e` | `mPure` | `HasTrace.pure` | `Kernel.const _ (dirac ())` | +| a leaf distribution | — | `HasTrace.sample` / `.of_isMarkov` | the kernel itself | +| a subprogram not reading the trace so far | — | `HasTrace.prodMkRight` | `Kernel.prodMkRight` | +| `if p c then … else …` | `ite` | `HasTrace.ite` | `Kernel.piecewise` | +| reparametrisation `κ (g c)` | `comp` | `HasTrace.comp` | `Kernel.comap` | +| `for … rdo …` | `forIn` | `.of_isMarkov` (coarse) — see below | one atomic coordinate | +| `record x` | `RDo.record` | `HasTrace.record` | `Kernel.deterministic` | + +`HasTrace.bind` asks for `Measurable cont` on the continuation. That is exactly the obligation +`is_markov` already discharges, so the two tactics compose: `is_markov` first, then the trace. + +`HasTrace.of_isMarkov` is the granularity knob. Any subprogram already known to be Markov can be +made a *single* trace coordinate, hiding its internal `←`s. Trace only as deep as the statement you +want to prove needs. + +## Recording a value as a random variable + +A value the program *computes* rather than draws — a `let mut` accumulator, an intermediate +quantity — is not a coordinate, so nothing makes it a random variable. The `record` annotation +([Record.lean](../RandomDo/Probability/Record.lean)) fixes that. Writing + +``` +record N, S +``` + +on a line of an `rdo` block gives the trace a coordinate for `N` and one for `S` from that point +on. It is sugar for `N ← RDo.record N`, and `RDo.record x` is the one-point distribution at `x`, so +the program still denotes the same measure — `RDo.record_bind` says so once and for all, and +`simp (disch := is_markov) only [record_bind_of_isMarkov]` erases every `record` from a program. +All the annotation does is put a `←` in the program text where there was none, and a `←` is what +the trace is built from. + +`rdo_trace` gives such a coordinate the kernel `Kernel.deterministic f`, so `rdo_peel` reports it +as being a deterministic function of everything drawn before it — and, more to the point, every +*later* draw's kernel is now written in terms of that coordinate, which is what lets you condition +on it. (`record` is a reserved token once the module is imported, so the underlying definition has +to be written `RDo.record`.) + +## Reading the statements off: the peeling rule + +Given the trace, the probabilistic content is extracted by one rule applied once per `←`. +`hasLaw_id_compProd` starts it, and then: + +```lean +HasLaw.compProd_fst -- law of the next draw +HasLaw.compProd_snd -- conditional law of everything after it, given it +Kernel.sectR_compProd -- re-expose the tail as a ⊗ₖ, so the rule applies again +HasCondDistrib.compProd_fst -- conditional law of the next draw given the history +HasCondDistrib.compProd_snd -- ... and of the rest, given the history *and* that draw +``` + +Step *k* of the peeling hands you the conditional distribution of the *k*-th draw given the first +*k−1*, and the kernel it hands you is the one written at that `←` in the program. For the chain +above, four lines of proof give + +```lean +HasLaw (fun ω ↦ ω.1) (κ c) P₃ +HasCondDistrib (fun ω ↦ ω.2.1) (fun ω ↦ ω.1) (Kernel.sectR η c) P₃ +HasCondDistrib (fun ω ↦ ω.2.2) (fun ω ↦ (ω.1, ω.2.1)) (θ.comap …) P₃ +``` + +**Independence is a special case.** When a draw does not mention the earlier ones, its factor is a +`Kernel.prodMkRight`, whose section is a constant kernel +(`Kernel.sectR_prodMkRight_eq_const`); `HasCondDistrib.indepFun_of_const` and +`.hasLaw_of_const` (both already in LML) then turn the conditional law into an unconditional law +plus an `IndepFun`. So *the syntactic fact that the second `←` does not mention `x` is what +produces the independence proof* — which is the property asked for. + +`RandomDo/Probability/Examples.lean` carries both halves: two independent draws with +`IndepFun` + `HasLaw` for each, and the dependent three-step chain with its conditional laws. + +## The tactics + +Both steps are mechanical, and both are automated. + +**`rdo_trace prog with h`** walks the program — reusing `RDo.Tactic.shapeOf`, the same classifier +`is_markov` walks — applying the combinator of the table above at each node, and adds +`h : HasTrace prog P out` to the context with `P` and `out` computed. Its side conditions +(`Measurable`, `IsMarkov`) are attacked with `fun_prop (disch := measurability)` and `is_markov`, +and whatever survives is handed back as a goal. `(fuel := n)` bounds how many definitions it looks +through; `set_option trace.rdo_trace true` prints the tree it walked. + +It recognises the weakening opportunity itself: when a draw does not mention the draws before it, +the factor it emits is a `Kernel.prodMkRight`, which is what makes the independence statement +available downstream. `mBind`, `mPure`, `let x ← p; return f x` (no coordinate) and leaves are +traced structurally; `ite`, `for` and early returns are traced coarsely, as one atomic draw — the +`HasTrace.ite` rule is there to be applied by hand when finer branching is wanted. + +**`rdo_peel h c`** iterates the peeling rule on such a hypothesis at the parameter `c`, adding one +`HasLaw`/`HasCondDistrib` per `←` plus `HasLaw` for the program's result, and — for every draw +whose kernel turns out not to read the draws before it — the unconditional law of that draw and its +`IndepFun` from them. `with h₁ h₂ …` names the facts in order. + +**Readability.** A program's own text ends up inside its trace kernel, and every statement a peeled +trace adds mentions the trace measure built from those kernels — printed in full, once per +hypothesis, that is unreadable. So anything too wide to print gets a local definition of its own: +`κ₁`, `κ₂`, … for the kernels (`rdo_trace`), `P` for the trace measure (`rdo_peel`, which reuses +the definitions `rdo_trace` already made). Each statement is then one line: + +```lean +κ₁ : Kernel (Vector (Fin K × ℝ) n) ((Fin K → ℝ) × (Fin K → ℝ)) := markovKernel … +κ₂ : Kernel … := markovKernel … +h : HasTrace (thompsonRecord hK) (κ₁ ⊗ₖ (Kernel.deterministic … ⊗ₖ (κ₂ ⊗ₖ κ₃))) fun p ↦ argmax p.2.2.2.2 +P : Measure … := (κ₁ ⊗ₖ …) hist +law1 : HasLaw Prod.fst (κ₁ hist) P +law2 : HasCondDistrib (fun ω ↦ ω.2.1) Prod.fst (… .comap (fun ω ↦ (hist, ω)) ⋯) P +… +``` + +The cut-off is printed width rather than subterm count, since a kernel's implicit type arguments +are large but never shown; small kernels such as `Kernel.const`, `Kernel.prodMkRight` and +`Kernel.deterministic` stay inline, which also keeps the independence detection above able to see +their shape. + +Together, on the two examples above: + +```lean +example : True := by + rdo_trace (sum2 μ) with h + rdo_peel h () with hX hY hY' hindep hout + -- hX : HasLaw Prod.fst μ P hY' : HasLaw Prod.snd μ P + -- hY : HasCondDistrib Prod.snd Prod.fst … P + -- hindep : Prod.fst ⟂ᵢ[P] Prod.snd + -- hout : HasLaw (fun ω ↦ ω.1 + ω.2) (sum2 μ) P + trivial +``` + +The `Automation` section of `Examples.lean` pins those statements with `have _ : … := hX`, so the +tests fail if the tactics ever produce something else. + +## Worked example: Thompson sampling + +[RandomDo/Probability/Thompson.lean](../RandomDo/Probability/Thompson.lean) runs the whole thing on +`thompson`. That program is two loops — fold the history into per-arm pull counts `N` and reward +sums `S`, then draw one Gaussian posterior sample per arm into a vector `θ` — followed by +`return argmax θ`. + +Naming the two stages as `rdo` programs of their own, `stats` and `sample`, the decomposition is an +equation both sides of which elaborate to the same two loops; only the `return` at the end of each +stage separates them, so it falls to `simp only [Prod.mk.eta, mBind_mPure]`: + +```lean +thompson hK hist = stats hist >>=ₘ fun NS ↦ sample NS >>=ₘ fun θ ↦ mPure (argmax θ) +``` + +Giving `stats` and `sample` their own `IsMarkov` instances then makes `rdo_trace` stop at each — +the granularity knob — so the trace it finds is the readable two-coordinate one, +`statsK ⊗ₖ sampleK.comap Prod.snd`, rather than the two raw loop terms. (Without those instances +`rdo_trace` still succeeds, and still splits the program in two; it just carries the unfolded loops +in the kernels.) `rdo_peel` then delivers, on the trace measure: + +```lean +HasLaw (fun ω ↦ ω.1) (stats hist) -- the statistics +HasCondDistrib (fun ω ↦ ω.2) (fun ω ↦ ω.1) sampleK -- θ given them, and nothing else +HasLaw (fun ω ↦ argmax ω.2) (thompson hK hist) -- the action played +``` + +The middle line is the substantive one: it says `θ` depends on the history *only through* +`(N, S)`, which is the whole content of "Thompson sampling is a function of the sufficient +statistics" — and it is read off the program text, not proved by hand. + +The `Record` section of the same file shows the annotation at work on the monolithic program: +`record N, S` just before the sampling loop turns `N` and `S` into coordinates two and three of a +four-coordinate trace, without restructuring the program into stages, and `thompsonRecord_eq` shows +the annotated program is the same measure as `thompson`. The sampling loop's kernel is then written +in terms of those coordinates, so `θ` can be conditioned on `N` and `S` individually. + +What is missing there is inside the sampling loop: `θ`'s `K` coordinates are drawn independently, +and that is exactly the statement the coarse loop trace cannot make. See below. + +## What is left + +### 1. Loops at per-iteration granularity + +Today a `for` loop is traced coarsely: it is Markov, so `HasTrace.of_isMarkov` makes the whole loop +one coordinate holding its final accumulator. That already gives the conditional law of *the loop's +result* given everything before it — often enough — but not the law of the individual iterations. +The `Loop` section of `Examples.lean` does exactly this: a loop summing `l.length` draws, followed +by a draw distributed as `η S` given the loop's result `S`. + +For per-iteration granularity the obstacle is that the number of `⊗ₖ` factors is symbolic, so the +trace type cannot be a fixed nest of products. Two shapes work: + +* **`List α`**, the values drawn in order. Fits this repo well: `RandomDo/Measurable.lean` already + provides the σ-algebra on `List α` and `measurable_cons`. The definitions are + + ```lean + noncomputable def loopTrace (K : ι → σ → Measure α) (upd : ι → σ → α → σ) : + List ι → σ → Measure (List α) + | [], _ => Measure.dirac [] + | i :: l, s => (K i s).bind fun z ↦ (loopTrace K upd l (upd i s z)).map (z :: ·) + + def loopOut (upd : ι → σ → α → σ) : List ι → σ → List α → σ + | [], s, _ => s + | _ :: _, s, [] => s + | i :: l, s, z :: zs => loopOut upd l (upd i s z) zs + ``` + + and the two theorems needed are `IsMarkov (loopTrace K upd l)` and + `forIn l s (fun i s ↦ (K i s).map fun z ↦ .yield (upd i s z)) = (loopTrace K upd l s).map + (loopOut upd l s)`, both by induction on `l`. The unrolling equations for the induction already + exist as `forIn_nil` / `forIn_cons` in [Lemmas.lean](../RandomDo/Tactic/Lemmas.lean) — they are + `private` and would need exposing. A third theorem, peeling the head off `loopTrace`, then gives + the conditional law of iteration `k` given iterations `< k`. + +* **`Π i : Iic n, α`** via Mathlib's `Kernel.partialTraj`. More machinery, but it is the *same* + history type as `Learning.Algorithm.policy`, so a loop traced this way plugs straight into LML's + `IsAlgEnvSeq` filtration and conditional-distribution API. + +Either shape would then get its own `traceCore` case in +[Tactic.lean](../RandomDo/Probability/Tactic.lean), next to `mBind`. + +Loop bodies with `break` / early `return` change the number of draws per iteration and need the +`Break.runK` construct handled too; a `Sum`-shaped trace space is the natural target, exactly as for +`ite` branches drawing in different spaces. The same `Sum` construction is what `rdo_trace` would +need to trace an `ite` finely rather than as one atomic draw. + +## The algorithm's draws inside an `IsAlgEnvSeq` + +[RandomDo/Probability/AlgTrace.lean](../RandomDo/Probability/AlgTrace.lean) closes the loop with +LML. `IsAlgEnvSeq A Y alg env P` says nothing about *how* the algorithm produced its actions: when +`alg` comes from an `rdo` program, the draws that program makes are not random variables of +`(Ω, P)` at all. This file makes them available. + +A `RDo.AlgTrace alg Ω` bundles what `rdo_trace` produces for a policy: one space `Ω` of internal +draws, a kernel `K n` for their law at step `n` given the history, and a readout `out n` +reconstructing the action. From it, `AlgTrace.algorithm` is an algorithm whose actions are pairs +`(draws, action)` — the draws first, then the action *deterministically* read off them, which is +what makes the two halves fall straight out of the peeling rule: + +* `AlgTrace.isAlgEnvSeq_snd` — forgetting the draws turns an algorithm-environment sequence for the + traced algorithm into one for `alg`; +* `AlgTrace.hasCondDistrib_trace` — the draws have conditional law `K n` given the history; +* `AlgTrace.action_ae_eq` — and the action is `out n` of the history and the draws. + +Since the traced algorithm faces the same environment, LML's `isAlgEnvSeq_unique` gives the +punchline, `AlgTrace.exists_isAlgEnvSeq_trace`: **any** algorithm-environment sequence may be +replaced by one on a space that also carries the draws, with the same trajectory law. The space is +existentially quantified precisely because it does not matter — no extension theorem is needed, +since the traced sequence is built on LML's canonical Ionescu–Tulcea space and uniqueness transfers +every conclusion about the actions and feedbacks back. + +### The `alg_env_trace` tactic + +`alg_env_trace tr` does the replacement in one step. Given an `IsAlgEnvSeq` hypothesis in the +context and an `AlgTrace tr` for its algorithm, it abstracts the goal — and every hypothesis +mentioning the probability space, the measure or the two sequences, so nothing is silently lost — +away from that space, and leaves two goals: + +* **`traced`**: the same statement on a space that also carries the draws `T`, with `hT₀` its law, + `hT` its conditional law given the history, and `hA₀`/`hA` the equations expressing each action + as the readout of the history and the draws; +* **`transfer`**: the obligation that the statement depends only on the law of the trajectory, with + both sequences' `IsAlgEnvSeq` available (so measurability is at hand). + +That second goal is what makes the move sound rather than a hole: the traced sequence lives on a +different space, and all that relates it to the original is `isAlgEnvSeq_unique`. + +```lean +example … (h : IsAlgEnvSeq A Y (alg hK) env P) : P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + alg_env_trace (trace hK) with Ω P A Y Z hseq hZ₀ hZ hA₀ hA + case traced => exact hseq.hasLaw_action_zero.map_eq + case transfer => … +``` + +after which the context reads + +``` +Ω : Type P : Measure Ω A : ℕ → Ω → Fin K Y Z : ℕ → Ω → ℝ +hseq : IsAlgEnvSeq A Y (alg hK) env P +hZ₀ : HasLaw (Z 0) (trace hK).K0 P +hZ : ∀ n, HasCondDistrib (Z (n+1)) (history A Y n) ((trace hK).K n) P +hA : ∀ n, A (n+1) =ᵐ[P] fun ω ↦ (trace hK).out n (history A Y n ω, Z (n+1) ω) +``` + +`alg_env_trace tr using h` names the hypothesis rather than searching for one; `with` names the +introduced variables. The space, its σ-algebra, the measure, the `IsProbabilityMeasure` hypothesis +and the two sequences all have to be local hypotheses, since the goal is abstracted over them. +`AlgTrace.wlog_trace` is the principle behind it, usable directly. + +`RDo.Example` at the end of the file runs the whole thing end to end on a toy policy written as an +`rdo` program: `rdo_trace` gives the trace, the `AlgTrace` packages it, `Example.exists_noise` +hands back the noise the policy draws at each step, and the last example drives the tactic. + +To do the same for `thompson` one still needs the measurable equivalence between `Iic n → 𝓐 × 𝓨` +and `Vector (𝓐 × 𝓨) (n + 1)` that turns it into a policy — `Vector.v_equiv` in +`RandomDo/Tactic/Examples.lean`, which is a `sorry` there (and stated one element short of the +right cardinality). Everything downstream of that point is done.