From da7bf5241c45cb52a7914ccf239d58f66f9ec9ff Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Mon, 31 Aug 2026 15:05:55 +0200 Subject: [PATCH 01/34] extract random variables --- RandomDo.lean | 6 + RandomDo/Probability/AlgTrace.lean | 541 ++++++++++++++++++++++++++++ RandomDo/Probability/Examples.lean | 286 +++++++++++++++ RandomDo/Probability/Record.lean | 113 ++++++ RandomDo/Probability/Tactic.lean | 550 +++++++++++++++++++++++++++++ RandomDo/Probability/Thompson.lean | 221 ++++++++++++ RandomDo/Probability/Trace.lean | 440 +++++++++++++++++++++++ notes/TRACE_SEMANTICS.md | 338 ++++++++++++++++++ 8 files changed, 2495 insertions(+) create mode 100644 RandomDo/Probability/AlgTrace.lean create mode 100644 RandomDo/Probability/Examples.lean create mode 100644 RandomDo/Probability/Record.lean create mode 100644 RandomDo/Probability/Tactic.lean create mode 100644 RandomDo/Probability/Thompson.lean create mode 100644 RandomDo/Probability/Trace.lean create mode 100644 notes/TRACE_SEMANTICS.md diff --git a/RandomDo.lean b/RandomDo.lean index 342175a..e1610ff 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -7,6 +7,12 @@ 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.Record +public import RandomDo.Probability.Tactic +public import RandomDo.Probability.Thompson +public import RandomDo.Probability.Trace public import RandomDo.Tactic.Elab public import RandomDo.Tactic.Examples public import RandomDo.Tactic.ForInStep diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean new file mode 100644 index 0000000..72fedfd --- /dev/null +++ b/RandomDo/Probability/AlgTrace.lean @@ -0,0 +1,541 @@ +/- +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 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 + +namespace RDo + +universe 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. -/ +theorem exists_isAlgEnvSeq_trace [MeasurableEq 𝓐] + {A₀ : ℕ → Ω₀ → 𝓐} {Y₀ : ℕ → Ω₀ → 𝓨} (h₀ : IsAlgEnvSeq A₀ Y₀ alg env P) : + ∃ (Ω' : Type (max uA uY uW)) (_ : MeasurableSpace Ω') (P' : Measure Ω') + (_ : IsProbabilityMeasure P') (A : ℕ → Ω' → 𝓐) (Y : ℕ → Ω' → 𝓨) (T : ℕ → Ω' → Ω), + IsAlgEnvSeq A Y alg env 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 + have h := IT.isAlgEnvSeq_trajMeasure tr.algorithm (env.withTrace Ω) + refine ⟨ℕ → (Ω × 𝓐) × 𝓨, inferInstance, trajMeasure tr.algorithm (env.withTrace Ω), + inferInstance, fun n (ω : ℕ → (Ω × 𝓐) × 𝓨) ↦ (IT.action n ω).2, IT.feedback, + fun n (ω : ℕ → (Ω × 𝓐) × 𝓨) ↦ (IT.action n ω).1, + tr.isAlgEnvSeq_snd h, ?_, tr.hasLaw_trace_zero h, tr.hasCondDistrib_trace h, + tr.action_zero_ae_eq h, tr.action_ae_eq h⟩ + exact isAlgEnvSeq_unique (tr.isAlgEnvSeq_snd h) h₀ + +/-- **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. -/ +theorem wlog_trace [MeasurableEq 𝓐] + {motive : (Ω₀ : Type (max uA uY uW)) → [MeasurableSpace Ω₀] → (P : Measure Ω₀) → + [IsProbabilityMeasure P] → (ℕ → Ω₀ → 𝓐) → (ℕ → Ω₀ → 𝓨) → Prop} + (traced : ∀ (Ω' : Type (max uA uY uW)) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (A' : ℕ → Ω' → 𝓐) (Y' : ℕ → Ω' → 𝓨) (T : ℕ → Ω' → Ω), + IsAlgEnvSeq A' Y' alg env 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 uA uY uW)) [MeasurableSpace Ω₁] (P₁ : Measure Ω₁) + [IsProbabilityMeasure P₁] (A₁ : ℕ → Ω₁ → 𝓐) (Y₁ : ℕ → Ω₁ → 𝓨) + (Ω₂ : Type (max 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 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 + obtain ⟨Ω', mΩ', P', hP', A', Y', T, hseq, hlaw, hT0, hT, hA0, hA⟩ := + tr.exists_isAlgEnvSeq_trace h + exact transfer Ω₀ P A Y Ω' P' A' Y' h hseq hlaw (traced Ω' P' A' Y' T hseq 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 + +/-- 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 mentioning the probability 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 `T`'s law, its + conditional law given the history, and the equations expressing 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`. + +* `alg_env_trace tr using h` names the hypothesis to use rather than searching for one. + +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. -/ +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 + -- Everything else that mentions the space has to travel with the goal, or it would be lost. + let deps ← do + let mut deps : Array FVarId := #[] + for d in ← getLCtx do + if !d.isImplementationDetail && !spaceFVars.contains d.fvarId then + let dty ← instantiateMVars d.type + if spaceFVars.any fun f ↦ dty.containsFVar f then + deps := deps.push d.fvarId + pure deps + let (_, g) ← g.revert deps + let (_, g) ← g.revert spaceFVars (preserveOrder := true) + -- 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) ← 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 + 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" + unless ← isDefEq concl (← g.getType) do + throwError "alg_env_trace: the goal does not have the expected shape{indentExpr + (← g.getType)}" + 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) + -- Introduce the traced space and its properties, then whatever travelled with the goal. + let given := (names?.map (·.map (·.getId))).getD #[] + let defaults : Array Name := #[`Ω, `P, `A, `Y, `T, `hseq, `hT₀, `hT, `hA₀, `hA] + 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] + let (_, traced) ← traced.introN intros.size intros.toList + let (_, traced) ← traced.introNP deps.size + replaceMainGoal [traced, transfer] + +end RDo.Tactic + +end + +@[expose] public section + +open MeasureTheory ProbabilityTheory Finset Learning RDo + +noncomputable section + +/-! ## An example: an algorithm whose policy is an `rdo` program + +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. + +To do the same for `thompson` one needs the measurable equivalence between `Iic n → 𝓐 × 𝓨` and +`Vector (𝓐 × 𝓨) (n + 1)` that turns it into a policy — `Vector.v_equiv` in +`RandomDo.Tactic.Examples`, still a `sorry` there (and stated one element short). Everything after +that point is what follows below. +-/ + +namespace RDo.Example + +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, and leaves the obligation that the statement only +depends on the law of the trajectory. Any hypothesis mentioning the space 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 hZ₀ hZ hA₀ hA + case traced => + -- `Z`, `hZ₀`, `hZ` and `hA` are the algorithm's draws and their laws, now available. + exact hseq.hasLaw_action_zero.map_eq + case transfer => + intro Ω₁ _ P₁ _ A₁ Y₁ Ω₂ _ P₂ _ A₂ Y₂ h₁ h₂ hlaw h₀ + have e₁ : A₁ 0 = (fun t ↦ (t 0).1) ∘ trajectory A₁ Y₁ := rfl + have e₂ : A₂ 0 = (fun t ↦ (t 0).1) ∘ trajectory A₂ Y₂ := rfl + rw [e₁, ← Measure.map_map (by fun_prop) + (measurable_trajectory h₁.measurable_action h₁.measurable_feedback), ← hlaw, + Measure.map_map (by fun_prop) + (measurable_trajectory h₂.measurable_action h₂.measurable_feedback), ← e₂, h₀] + +end RDo.Example + +end + +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/Record.lean b/RandomDo/Probability/Record.lean new file mode 100644 index 0000000..965a1a1 --- /dev/null +++ b/RandomDo/Probability/Record.lean @@ -0,0 +1,113 @@ +/- +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. -/ +@[simp] +lemma record_bind (x : α) {f : α → Measure β} (hf : Measurable f) : record x >>=ₘ 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..3cd7467 --- /dev/null +++ b/RandomDo/Probability/Tactic.lean @@ -0,0 +1,550 @@ +/- +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 + suggested : Name + 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..46e4dbc --- /dev/null +++ b/RandomDo/Probability/Thompson.lean @@ -0,0 +1,221 @@ +/- +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.Examples + +set_option linter.style.header false + +/-! +# Thompson sampling as random variables + +`thompson` in `RandomDo.Tactic.Examples` 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 : ℕ} + +/-! ## 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` 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/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. From 719d74f489e69d2ef9aa3f9556b5e59a950b9a65 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 31 Aug 2026 17:13:52 +0200 Subject: [PATCH 02/34] Update RandomDo.lean --- RandomDo.lean | 9 +-------- 1 file changed, 1 insertion(+), 8 deletions(-) diff --git a/RandomDo.lean b/RandomDo.lean index e1610ff..e04384d 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -2,19 +2,12 @@ module -- shake: keep-all --deprecated_module: ignore public import RandomDo.ForMathlib.MeasureTheory.MeasurableSpace.Embedding public import RandomDo.Measurable -public import RandomDo.Monad.Examples 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.Record -public import RandomDo.Probability.Tactic -public import RandomDo.Probability.Thompson -public import RandomDo.Probability.Trace +public import RandomDo.Tactic.Deriving public import RandomDo.Tactic.Elab -public import RandomDo.Tactic.Examples public import RandomDo.Tactic.ForInStep public import RandomDo.Tactic.IsMarkov public import RandomDo.Tactic.Lemmas From c876e383931797acd136126e6d01a5a668458d06 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 31 Aug 2026 17:14:24 +0200 Subject: [PATCH 03/34] Update RandomDo.lean --- RandomDo.lean | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/RandomDo.lean b/RandomDo.lean index e04384d..1a20f1c 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -6,6 +6,12 @@ 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.Record +public import RandomDo.Probability.Tactic +public import RandomDo.Probability.Thompson +public import RandomDo.Probability.Trace public import RandomDo.Tactic.Deriving public import RandomDo.Tactic.Elab public import RandomDo.Tactic.ForInStep From b06c1edea8eb77a72745d371aa767411c6650f88 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 31 Aug 2026 17:31:48 +0200 Subject: [PATCH 04/34] Fix `Thompson.lean` --- RandomDo/Probability/Thompson.lean | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/RandomDo/Probability/Thompson.lean b/RandomDo/Probability/Thompson.lean index 46e4dbc..1df4afe 100644 --- a/RandomDo/Probability/Thompson.lean +++ b/RandomDo/Probability/Thompson.lean @@ -6,7 +6,8 @@ Authors: Rémy Degenne module public import RandomDo.Probability.Tactic -public import RandomDo.Tactic.Examples +public import RandomDo.Tactic.Elab +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg set_option linter.style.header false @@ -56,6 +57,8 @@ 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 @@ -96,6 +99,20 @@ instance : IsMarkovKernel (sampleK (K := K)) := by unfold sampleK; infer_instanc @[simp] lemma sampleK_apply (NS : (Fin K → ℝ) × (Fin K → ℝ)) : sampleK NS = sample NS := rfl +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. -/ From 11eaad365f7369577ac60237992093bbc896c765 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 4 Sep 2026 15:17:40 +0200 Subject: [PATCH 05/34] add extend_space tactic --- RandomDo.lean | 3 + RandomDo/Probability/AlgTrace.lean | 72 +++ RandomDo/Probability/Extend.lean | 555 +++++++++++++++++++++++ RandomDo/Probability/ExtendExamples.lean | 268 +++++++++++ RandomDo/Probability/Transfer.lean | 144 ++++++ 5 files changed, 1042 insertions(+) create mode 100644 RandomDo/Probability/Extend.lean create mode 100644 RandomDo/Probability/ExtendExamples.lean create mode 100644 RandomDo/Probability/Transfer.lean diff --git a/RandomDo.lean b/RandomDo.lean index e1610ff..82d8bfd 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -9,10 +9,13 @@ 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.ExtendExamples 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.Elab public import RandomDo.Tactic.Examples public import RandomDo.Tactic.ForInStep diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean index 72fedfd..6413e2f 100644 --- a/RandomDo/Probability/AlgTrace.lean +++ b/RandomDo/Probability/AlgTrace.lean @@ -6,6 +6,7 @@ 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 @@ -57,6 +58,48 @@ open MeasureTheory ProbabilityTheory Finset Learning noncomputable section +/-- 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. -/ +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; +`h.measurable_action` and `h.measurable_feedback` provide the side conditions when an +`IsAlgEnvSeq` hypothesis `h` is around. -/ +@[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 uA uY uW @@ -534,6 +577,35 @@ example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] { Measure.map_map (by fun_prop) (measurable_trajectory h₂.measurable_action h₂.measurable_feedback), ← e₂, h₀] +/-- **`extend_space` alongside an algorithm-environment sequence.** The sequence pulls back along +the projection `f` by `IsAlgEnvSeq.comp_measurePreserving`, and the larger space also carries a +Gaussian `U` independent of the whole trajectory. The statement does not mention the original +space, so the `transfer` obligation is trivial and `extend_space` closes it. -/ +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 + extend_space (gaussianReal 0 1) using P with Ω' P' f hf U hU hind + exact ⟨Ω', inferInstance, P', inferInstance, fun n ω ↦ A n (f ω), fun n ω ↦ Y n (f ω), U, + h.comp_measurePreserving hf, hU, + hind.comp (measurable_trajectory h.measurable_action h.measurable_feedback) measurable_id⟩ + +/-- **The `transfer` tactic with an algorithm-environment sequence.** 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 (gaussianReal 0 1) with Ω' P' f hf U hU hind + transfer hf at h + exact h.hasLaw_action_zero.map_eq + end RDo.Example end diff --git a/RandomDo/Probability/Extend.lean b/RandomDo/Probability/Extend.lean new file mode 100644 index 0000000..51811f1 --- /dev/null +++ b/RandomDo/Probability/Extend.lean @@ -0,0 +1,555 @@ +/- +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 +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 μ` replaces the probability space `(Ω, P)` of the goal by one that also carries 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". The new space is `Ω × E` +with the product measure, but the goal is stated on an abstract space `Ω'` related to `Ω` by a +measure-preserving map `f : Ω' → Ω`, which is all the extended goal may use. + +Every random variable `X : Ω → α` the goal mentions becomes `fun ω ↦ X (f ω)`, every event +`s : Set Ω` becomes `f ⁻¹' s`, and the measure becomes `P'`. Hypotheses about the old space stay +in the context untouched, since they are still true: they are pulled back along `f` on demand +(`hX.comp hf.hasLaw`, `hind.comp hX measurable_id`, `h.comp_measurePreserving hf`, …). A +hypothesis the goal itself depends on, such as a measurability proof inside a `Kernel.comap`, is +generalized and reintroduced under a primed name. + +Two goals are left: + +* `extended`: the same statement on `(Ω', P')`, with `f`, `hf : MeasurePreserving f P' P`, `Z`, + `hZ : HasLaw Z μ P'` and `hind : IndepFun f Z P'` in the context; +* `transfer`: the obligation that the statement pulls back along any measure-preserving map. This + is what makes the replacement sound. The `transfer` tactic discharges it for laws, events, + integrals, independence and conditional laws, using the `@[transfer]` lemmas of this file, and + `extend_space` runs it: the goal is only left when the tactic fails. In the extended goal, + `transfer hf at h` pulls a hypothesis `h` about the old space back to the new one. + +`extend_space κ` for a Markov kernel `κ : Kernel Ω E` does the same with a draw whose conditional +law given the old space is `κ`: it provides `hZ : HasCondDistrib Z f κ P'` instead of a law and an +independence. For a draw conditional on a random variable `X`, use `κ.comap X hX` and read the +result through `HasCondDistrib.comp_right`. + +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. +* `MeasureTheory.MeasurePreserving.map_fun_comp`, `hasLaw_fun_comp_iff`, `indepFun_fun_comp_iff`, + `hasCondDistrib_fun_comp_iff`: pulling statements back along a measure-preserving map. +* `ProbabilityTheory.HasCondDistrib.comp_measurePreserving`: the forward direction, without any + measurability assumption. +-/ + +@[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_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 + +end MeasureTheory.MeasurePreserving + +/-- A conditional law pulls back along a measure-preserving map. -/ +lemma ProbabilityTheory.HasCondDistrib.comp_measurePreserving {Ω Ω' 𝓧 𝓨 : Type*} + {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} {m𝓧 : MeasurableSpace 𝓧} + {m𝓨 : MeasurableSpace 𝓨} {P : Measure Ω} {P' : Measure Ω'} {f : Ω' → Ω} {X : Ω → 𝓧} + {Y : Ω → 𝓨} {κ : 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 + +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 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), + 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 + (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 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), + 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 + 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 + +/-- 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 + let mentions (e : Expr) : Bool := e.hasAnyFVar dep.contains + if mentions d.type || (d.value?.map mentions).getD false then dep := dep.insert d.fvarId + return dep + +/-- Whether a local hypothesis of type `ty` can be transported along `f : Ω' → Ω`: its type has to +be `ι₁ → ⋯ → ιₖ → Ω → α` or `ι₁ → ⋯ → ιₖ → Set Ω`, with `Ω` appearing nowhere else. -/ +partial def transportable (Ω : FVarId) (ty : Expr) : Bool := + if ty.isAppOfArity ``Set 1 then ty.appArg! == .fvar Ω + else match ty with + | .forallE _ d b _ => + if d == .fvar Ω then !b.containsFVar Ω + else !d.containsFVar Ω && transportable Ω b + | _ => false + +/-- 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 transportable sp.Ω ty then return ← go todo seen acc + 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" + 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 + +/-- `extend_space μ` replaces the probability space `(Ω, P)` of the goal by one that also carries +a random variable `Z` with law `μ`, independent of everything defined on `Ω`. `μ` is a probability +measure on some `E`. Two goals are left: + +* `extended`: the same statement on a space `(Ω', P')` with a measure-preserving map + `f : Ω' → Ω`. Every random variable `X : Ω → α` of the goal becomes `fun ω ↦ X (f ω)` and every + event `s : Set Ω` becomes `f ⁻¹' s`; the context gains `hf : MeasurePreserving f P' P`, + `hZ : HasLaw Z μ P'` and `hind : IndepFun f Z P'`. Hypotheses about the old space stay as they + are and may be pulled back along `f`, for instance by `hX.comp hf.hasLaw` for a law, + `hind.comp hX measurable_id` for 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. +* `transfer`: the obligation that the statement pulls back along a measure-preserving map, which is + what makes the replacement sound. See `MeasurePreserving.map_fun_comp` and the + `MeasurePreserving.*_fun_comp_iff` lemmas. + +* `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`. +* `extend_space μ using P` names the measure to extend rather than reading it off the goal. +* `extend_space μ with Ω' P' f hf Z hZ hind` names what is introduced. + +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 + +elab_rules : tactic + | `(tactic| extend_space $μ $[using $P?]? $[with $names?*]?) => withMainContext do + 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 + "extend_space: {μE} is neither a measure nor a kernel; it has type{indentExpr μty}" + -- 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 + "extend_space: the goal mentions no measure on a local space; name one with `using`" + | _ => throwError + "extend_space: the goal mentions several measures, {cands.map Expr.fvar}; choose one \ + with `using`" + let sp ← ProbSpace.ofMeasure PE + let ΩE := Expr.fvar sp.Ω + if let some dom := dom? then + unless ← isDefEq dom ΩE do + throwError "extend_space: 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 "extend_space: {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₀) + 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 + "extend_space: 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 "extend_space: `{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 "extend_space: 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 "extend_space: 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, then what was generalized, under primed names. + let extended := args[← idx `extended]!.mvarId! + extended.setKind .syntheticOpaque + extended.setTag `extended + let given := (names?.map (·.map (·.getId))).getD #[] + let defaults : Array Name := #[.mkSimple "Ω'", .mkSimple "P'", `f, `hf, `Z, `hZ, `hind] + let nNames := if isKernel then 6 else 7 + if given.size > nNames then + throwError "extend_space: at most {nNames} names may be given" + let pick (i : Nat) : Name := if h : i < given.size then given[i] else defaults[i]! + let mut intros : Array Name := #[pick 0, `inst, pick 1, `inst, pick 2, pick 3, pick 4, pick 5] + if !isKernel then intros := intros.push (pick 6) + let (_, extended) ← extended.introN intros.size intros.toList + let primed ← gens.toList.mapM fun f ↦ do pure ((← f.getUserName).appendAfter "'") + let (_, extended) ← extended.introN gens.size primed + -- 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 ← saveState + -- `tryCatchRuntimeEx`: a failure inside `measurability` may be a maximum recursion depth error, + -- which `try … catch` lets through. + let transferLeft ← tryCatchRuntimeEx + (do + setGoals [transfer'.mvarId!] + evalTactic (← `(tactic| transfer)) + unless (← getUnsolvedGoals).isEmpty do throwError "transfer left goals" + pure []) + (fun _ ↦ do + s.restore + pure [transfer'.mvarId!]) + setGoals ([extended] ++ transferLeft ++ rest) + +end RDo.Tactic + +end diff --git a/RandomDo/Probability/ExtendExamples.lean b/RandomDo/Probability/ExtendExamples.lean new file mode 100644 index 0000000..ee8e4ae --- /dev/null +++ b/RandomDo/Probability/ExtendExamples.lean @@ -0,0 +1,268 @@ +/- +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.Extend +public import Mathlib.Probability.Independence.InfinitePi + +set_option linter.style.header false + +/-! +# Tests and examples for `extend_space` and `transfer` + +The first section pins down what `extend_space` produces: which hypotheses appear, what the goal +becomes, when the `transfer` obligation is closed automatically and when it is left, and how +hypotheses the goal depends on are handled. The next ones show `transfer` at work on laws, +independence, conditional laws, events, integrals and almost-everywhere statements, both on the +obligation and to pull hypotheses back to the new space; then a draw with a conditional law, an +i.i.d. sequence, and the errors the tactics report. + +Throughout, `Ω` lives in `Type u` and `E` in `Type`: the tactic lifts the product to the universe +of `Ω`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory RDo + +noncomputable section + +namespace RDo.Example.Extend + +universe u + +variable {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + {E : Type} [MeasurableSpace E] (μ : Measure E) [IsProbabilityMeasure μ] + +include μ + +/-! ## The shape of the goals -/ + +/-- Random variables, families of random variables and events are transported along `f`, and the +measure becomes `P'`. With the measurability hypotheses around, the `transfer` obligation is +discharged by `extend_space` itself, and only the extended goal is left. -/ +example (X : Ω → ℝ) (A : ℕ → Ω → ℝ) (s : Set Ω) (hX : Measurable X) (hA : ∀ n, Measurable (A n)) + (hs : MeasurableSet s) : + P.map X = P.map X ∧ (∀ n, P.map (A n) = P.map (A n)) ∧ P s = P s := by + extend_space μ with Ω' P' f hf Z hZ hind + guard_hyp hf : MeasurePreserving f P' P + guard_hyp hZ : HasLaw Z μ P' + guard_hyp hind : IndepFun f Z P' + guard_target =ₐ P'.map (fun ω ↦ X (f ω)) = P'.map (fun ω ↦ X (f ω)) + ∧ (∀ n, P'.map (fun ω ↦ A n (f ω)) = P'.map (fun ω ↦ A n (f ω))) + ∧ P' (f ⁻¹' s) = P' (f ⁻¹' s) + exact ⟨rfl, fun _ ↦ rfl, rfl⟩ + +/-- A statement `transfer` has no lemma for, here `IsProbabilityMeasure P`: the obligation is left. +It is stated for an arbitrary measure-preserving map, with the original goal as its conclusion. +Without `with`, the names are `Ω' P' f hf Z hZ hind`. -/ +example : IsProbabilityMeasure P := by + extend_space μ + case extended => + guard_hyp hf : MeasurePreserving f P' P + guard_hyp hZ : HasLaw Z μ P' + guard_hyp hind : IndepFun f Z P' + infer_instance + case transfer => + guard_target =ₐ ∀ (Ω' : Type u) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → + IsProbabilityMeasure P' → IsProbabilityMeasure P + intro Ω' _ P' _ f hf h + infer_instance + +/-- An operation `transfer` has no lemma for, here `Measure.restrict`: the obligation is left. -/ +example (s : Set Ω) : P.restrict s Set.univ = P s := by + extend_space μ + case extended => + guard_target =ₐ P'.restrict (f ⁻¹' s) Set.univ = P' (f ⁻¹' s) + exact Measure.restrict_apply_univ _ + case transfer => + intro Ω' _ P' _ f hf h + exact Measure.restrict_apply_univ _ + +/-- The goal depends on the measurability proof `hX`: it is generalized and reintroduced as `hX'`, +now about `X ∘ f`, while `hX` itself stays. The goal mentions no measure, so `using P` says which +space to extend. `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 with Ω' P' f hf Z hZ hind + 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 + +/-! ## Transferring statements + +`extend_space` closes the `transfer` obligation, and `transfer hf at h` pulls a hypothesis about +the old space back to the new one. Both rewrite with the `@[transfer]` lemmas, whose side +conditions are measurability statements found by `assumption`, `fun_prop` or `measurability`. -/ + +/-- A law. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space μ with Ω' P' f hf Z hZ hind + transfer hf at hXν + guard_hyp hXν : HasLaw (fun ω ↦ X (f ω)) ν P' + exact hXν.map_eq + +/-- The same with `HasLaw` as the goal, and the hypothesis pulled back by hand. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : + HasLaw X ν P := by + extend_space μ + exact hXν.comp hf.hasLaw + +/-- Independence. -/ +example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) (hXY : IndepFun X Y P) : + IndepFun X Y P := by + extend_space μ + transfer hf at hXY + guard_hyp hXY : IndepFun (fun ω ↦ X (f ω)) (fun ω ↦ Y (f ω)) P' + exact hXY + +/-- 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 μ + transfer hf at hXY + exact hXY + +/-- An event, given by a set-builder expression: `measurability` proves it measurable. -/ +example (X : Ω → ℝ) (hX : Measurable X) (h : P {ω | 0 < X ω} = 1 / 2) : + P {ω | 0 < X ω} = 1 / 2 := by + extend_space μ + transfer hf at h + guard_hyp h : P' {ω | 0 < X (f ω)} = 1 / 2 + exact h + +/-- The real-valued measure of an event. -/ +example (s : Set Ω) (hs : MeasurableSet s) (r : ℝ) (h : P.real s = r) : P.real s = r := by + extend_space μ + transfer hf at h + exact h + +/-- An integral, with an almost-everywhere hypothesis. -/ +example (X : Ω → ℝ) (hX : Measurable X) (h : ∀ᵐ ω ∂P, 0 ≤ X ω) : 0 ≤ ∫ ω, X ω ∂P := by + extend_space μ + transfer hf at h + guard_hyp h : ∀ᵐ ω ∂P', 0 ≤ X (f ω) + exact integral_nonneg_of_ae h + +/-- A Lebesgue integral. -/ +example (X : Ω → ℝ) (hX : Measurable X) (c : ENNReal) (h : ∫⁻ ω, ‖X ω‖ₑ ∂P = c) : + ∫⁻ ω, ‖X ω‖ₑ ∂P = c := by + extend_space μ + transfer hf at h + 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 μ + transfer hf at h + guard_hyp h : (fun ω ↦ X (f ω)) =ᵐ[P'] fun ω ↦ Y (f ω) + exact h + +/-- `transfer hf` on a goal, outside of `extend_space`: the goal is moved to the new space and +closed by the hypothesis. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) {Ω' : Type u} [MeasurableSpace Ω'] + {P' : Measure Ω'} (f : Ω' → Ω) (hf : MeasurePreserving f P' P) + (h : HasLaw (fun ω ↦ X (f ω)) ν P') : + HasLaw X ν P := by + transfer hf + +/-! ## 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: +`X ∘ f` and `Z`. -/ +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 with Ω' P' f hf Z hZ hind + exact ⟨Ω', inferInstance, P', inferInstance, fun ω ↦ X (f ω), Z, hXν.comp hf.hasLaw, hZ, + hind.comp hX measurable_id⟩ + +/-- 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 with Ω' P' f hf Z hZ hind + have hZn (n : ℕ) : HasLaw (fun ω ↦ Z ω n) μ P' := + (measurePreserving_eval_infinitePi _ n).hasLaw.comp hZ + exact ⟨Ω', inferInstance, P', inferInstance, fun ω ↦ X (f ω), fun n ω ↦ Z ω n, + hXν.comp hf.hasLaw, hZn, (iIndepFun_iff_hasLaw_Pi_infinitePi hZn hZ.aemeasurable).2 hZ, + hind.comp hX measurable_id⟩ + +/-! ## A draw with a conditional law -/ + +/-- `extend_space κ` for a kernel on `Ω`: the draw has conditional law `κ` given the old space, and +there is no independence hypothesis. For a law conditional on `X`, extend with `κ.comap X hX` and +read `hZ` through `HasCondDistrib.comp_right`. -/ +example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel ℝ E) [IsMarkovKernel κ] (ν : Measure ℝ) + (hXν : HasLaw X ν P) : + P.map X = ν := by + extend_space (κ.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) : P.map X = P.map X := by + extend_space μ' + rfl + +/-! ## 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: 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 + +end RDo.Example.Extend + +end + +end diff --git a/RandomDo/Probability/Transfer.lean b/RandomDo/Probability/Transfer.lean new file mode 100644 index 0000000..d515092 --- /dev/null +++ b/RandomDo/Probability/Transfer.lean @@ -0,0 +1,144 @@ +/- +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.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. A lemma tagged `@[transfer]` states one such 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. + +`transfer hf` rewrites the goal with every such 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` rewrites hypotheses instead: this pulls a fact about the old space back +to the new one. `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. The `transfer` tactic +rewrites with all of them. -/ +register_label_attr transfer + +namespace RDo.Tactic + +/-- 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. -/ +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).getAppFn.constName? + let tac ← if head.any funProps.contains then `(tactic| first | assumption | fun_prop) + else `(tactic| first | assumption | measurability) + tryCatchRuntimeEx (evalTactic tac) fun e ↦ + throwError "transfer_discharger: {e.toMessageData}" + +/-- The `@[transfer]` lemmas instantiated at `hf`, as `simp` arguments, together with +`Set.preimage_ofPred_eq`, which puts the transferred events in the same form as `extend_space`. -/ +def transferSimpArgs (hf : Term) : CoreM (Array (TSyntax ``Lean.Parser.Tactic.simpLemma)) := do + let args ← (← labelled `transfer).mapM fun n ↦ + `(Lean.Parser.Tactic.simpLemma| $(mkIdent n):ident $hf) + return args.push (← `(Lean.Parser.Tactic.simpLemma| Set.preimage_ofPred_eq)) + +/-- 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) + +/-- `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. + +* `transfer hf at h₁ h₂` rewrites hypotheses instead: a fact about the old space becomes the + corresponding fact about the new one. +* `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, _ => + let args ← transferSimpArgs hf + evalTactic (← `(tactic| + simp -failIfUnchanged (disch := transfer_discharger) only [$args,*] $(loc?)?)) + if loc?.isNone then + unless (← getUnsolvedGoals).isEmpty do + evalTactic (← `(tactic| try assumption)) + | 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 From feff0e013830f3f054b3adb37582e0fe56d86a3c Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 4 Sep 2026 16:01:28 +0200 Subject: [PATCH 06/34] better extend_space: remove the old space --- RandomDo/Probability/AlgTrace.lean | 35 +- RandomDo/Probability/Extend.lean | 816 +++++++++++++++++------ RandomDo/Probability/ExtendExamples.lean | 303 ++++++--- RandomDo/Probability/Transfer.lean | 149 ++++- 4 files changed, 960 insertions(+), 343 deletions(-) diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean index 6413e2f..b242ecb 100644 --- a/RandomDo/Probability/AlgTrace.lean +++ b/RandomDo/Probability/AlgTrace.lean @@ -59,7 +59,9 @@ open MeasureTheory ProbabilityTheory Finset Learning noncomputable section /-- 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. -/ +`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 : ℕ → Ω → 𝓨} @@ -577,10 +579,12 @@ example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] { Measure.map_map (by fun_prop) (measurable_trajectory h₂.measurable_action h₂.measurable_feedback), ← e₂, h₀] -/-- **`extend_space` alongside an algorithm-environment sequence.** The sequence pulls back along -the projection `f` by `IsAlgEnvSeq.comp_measurePreserving`, and the larger space also carries a -Gaussian `U` independent of the whole trajectory. The statement does not mention the original -space, so the `transfer` obligation is trivial and `extend_space` closes it. -/ +/-- **`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) : @@ -588,21 +592,24 @@ example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] { (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 - extend_space (gaussianReal 0 1) using P with Ω' P' f hf U hU hind - exact ⟨Ω', inferInstance, P', inferInstance, fun n ω ↦ A n (f ω), fun n ω ↦ Y n (f ω), U, - h.comp_measurePreserving hf, hU, - hind.comp (measurable_trajectory h.measurable_action h.measurable_feedback) measurable_id⟩ - -/-- **The `transfer` tactic with an algorithm-environment sequence.** 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. -/ + 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 (gaussianReal 0 1) with Ω' P' f hf U hU hind + extend_space_map (gaussianReal 0 1) with Ω' P' f hf U hU hind transfer hf at h exact h.hasLaw_action_zero.map_eq diff --git a/RandomDo/Probability/Extend.lean b/RandomDo/Probability/Extend.lean index 51811f1..e36efd9 100644 --- a/RandomDo/Probability/Extend.lean +++ b/RandomDo/Probability/Extend.lean @@ -18,33 +18,35 @@ set_option linter.style.header false /-! # Extending a probability space by a new random variable -`extend_space μ` replaces the probability space `(Ω, P)` of the goal by one that also carries 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". The new space is `Ω × E` -with the product measure, but the goal is stated on an abstract space `Ω'` related to `Ω` by a -measure-preserving map `f : Ω' → Ω`, which is all the extended goal may use. - -Every random variable `X : Ω → α` the goal mentions becomes `fun ω ↦ X (f ω)`, every event -`s : Set Ω` becomes `f ⁻¹' s`, and the measure becomes `P'`. Hypotheses about the old space stay -in the context untouched, since they are still true: they are pulled back along `f` on demand -(`hX.comp hf.hasLaw`, `hind.comp hX measurable_id`, `h.comp_measurePreserving hf`, …). A -hypothesis the goal itself depends on, such as a measurability proof inside a `Kernel.comap`, is -generalized and reintroduced under a primed name. - -Two goals are left: - -* `extended`: the same statement on `(Ω', P')`, with `f`, `hf : MeasurePreserving f P' P`, `Z`, - `hZ : HasLaw Z μ P'` and `hind : IndepFun f Z P'` in the context; -* `transfer`: the obligation that the statement pulls back along any measure-preserving map. This - is what makes the replacement sound. The `transfer` tactic discharges it for laws, events, - integrals, independence and conditional laws, using the `@[transfer]` lemmas of this file, and - `extend_space` runs it: the goal is only left when the tactic fails. In the extended goal, - `transfer hf at h` pulls a hypothesis `h` about the old space back to the new one. - -`extend_space κ` for a Markov kernel `κ : Kernel Ω E` does the same with a draw whose conditional -law given the old space is `κ`: it provides `hZ : HasCondDistrib Z f κ P'` instead of a law and an -independence. For a draw conditional on a random variable `X`, use `κ.comap X hX` and read the -result through `HasCondDistrib.comp_right`. +`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`, `hZ : HasLaw Z μ P` and + `hind`, the independence of `Z` from the transported random variables, as a tuple. +* `extend_space! μ` does the same and clears the old space, the map, and everything that + mentions them. +* `extend_space_map μ` is the explicit form: nothing is renamed, the new space is `Ω'` with + `P'`, `f : Ω' → Ω`, `hf`, `Z`, `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. + +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 hypothesis the goal itself depends on, such as a measurability proof +inside a `Kernel.comap`, is generalized along with the goal. + +`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 +a random variable `X`, extend with `κ.comap X hX`: `extend_space` then states `hZ` as +`HasCondDistrib Z X κ P`. 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 @@ -55,8 +57,9 @@ statement by statement. Here the projection `f` is measure preserving, which is * `RDo.wlog_extend`, `RDo.wlog_extend_kernel`: the principles behind the tactic. * `MeasureTheory.MeasurePreserving.map_fun_comp`, `hasLaw_fun_comp_iff`, `indepFun_fun_comp_iff`, `hasCondDistrib_fun_comp_iff`: pulling statements back along a measure-preserving map. -* `ProbabilityTheory.HasCondDistrib.comp_measurePreserving`: the forward direction, without any - measurability assumption. +* The `@[transfer]` lemmas `MeasurePreserving.transfer_*` and the `@[transfer_forward]` lemmas + `*.comp_measurePreserving`, `*.preimage_measurePreserving`: the same facts in the forms the + `transfer` tactic uses. -/ @[expose] public section @@ -163,19 +166,80 @@ lemma transfer_ae_eq (hf : MeasurePreserving f P' P) {X Y : Ω → 𝓧} 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 : Ω → 𝓨} + +@[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. -/ -lemma ProbabilityTheory.HasCondDistrib.comp_measurePreserving {Ω Ω' 𝓧 𝓨 : Type*} - {mΩ : MeasurableSpace Ω} {mΩ' : MeasurableSpace Ω'} {m𝓧 : MeasurableSpace 𝓧} - {m𝓨 : MeasurableSpace 𝓨} {P : Measure Ω} {P' : Measure Ω'} {f : Ω' → Ω} {X : Ω → 𝓧} - {Y : Ω → 𝓨} {κ : Kernel 𝓧 𝓨} (h : HasCondDistrib Y X κ P) (hf : MeasurePreserving f P' P) : +@[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 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 + namespace RDo /-! ### The product extension -/ @@ -268,6 +332,8 @@ open MeasureTheory ProbabilityTheory namespace RDo.Tactic +/-! ### 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. -/ @@ -318,19 +384,27 @@ def spaceDependents (sp : ProbSpace) : MetaM FVarIdSet := do 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 - let mentions (e : Expr) : Bool := e.hasAnyFVar dep.contains - if mentions d.type || (d.value?.map mentions).getD false then dep := dep.insert d.fvarId + -- 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 -/-- Whether a local hypothesis of type `ty` can be transported along `f : Ω' → Ω`: its type has to -be `ι₁ → ⋯ → ιₖ → Ω → α` or `ι₁ → ⋯ → ιₖ → Set Ω`, with `Ω` appearing nowhere else. -/ -partial def transportable (Ω : FVarId) (ty : Expr) : Bool := - if ty.isAppOfArity ``Set 1 then ty.appArg! == .fvar Ω +/-- 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 !b.containsFVar Ω - else !d.containsFVar Ω && transportable Ω b - | _ => false + 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ₖ`. -/ @@ -385,170 +459,520 @@ partial def forallArity : Expr → Nat | .forallE _ _ b _ => forallArity b + 1 | _ => 0 -/-- `extend_space μ` replaces the probability space `(Ω, P)` of the goal by one that also carries -a random variable `Z` with law `μ`, independent of everything defined on `Ω`. `μ` is a probability -measure on some `E`. Two goals are left: - -* `extended`: the same statement on a space `(Ω', P')` with a measure-preserving map - `f : Ω' → Ω`. Every random variable `X : Ω → α` of the goal becomes `fun ω ↦ X (f ω)` and every - event `s : Set Ω` becomes `f ⁻¹' s`; the context gains `hf : MeasurePreserving f P' P`, - `hZ : HasLaw Z μ P'` and `hind : IndepFun f Z P'`. Hypotheses about the old space stay as they - are and may be pulled back along `f`, for instance by `hX.comp hf.hasLaw` for a law, - `hind.comp hX measurable_id` for 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. -* `transfer`: the obligation that the statement pulls back along a measure-preserving map, which is - what makes the replacement sound. See `MeasurePreserving.map_fun_comp` and the - `MeasurePreserving.*_fun_comp_iff` lemmas. +/-! ### 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 + /-- 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 + +/-- The old-space version of a transported random variable `x : ι₁ → ⋯ → ιₖ → Ω → α`, as a +random variable `fun ω ↦ fun i₁ … iₖ ↦ x i₁ … iₖ ω` for the independence statement. -/ +partial def oldComponent (Ω ω x ty : Expr) : MetaM Expr := + 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] (← oldComponent Ω ω (mkApp x i) (b.instantiate1 i)) + | _ => 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, as the independence of `Z` +and the tuple of those whose measurability `fun_prop` can prove. -/ +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 mut comps : Array (Expr × Expr) := #[] + for (x, _, _, isSet) in transported do + if isSet then continue + let ty ← instantiateMVars (← x.getType) + let c ← withLocalDecl `ω .default ΩE fun ω ↦ do + mkLambdaFVars #[ω] (← oldComponent ΩE ω (.fvar x) ty) + let c := c.eta + let goal ← mkFreshExprSyntheticOpaqueMVar (← mkAppM ``Measurable #[c]) + if let some hc ← tryTactic? goal.mvarId! (← `(tactic| fun_prop)) then + comps := comps.push (c, hc) + if comps.isEmpty then return none + let rec mkTuple : List (Expr × Expr) → MetaM (Expr × Expr) + | [] => throwError "extend_space: internal error, empty tuple" + | [ch] => pure ch + | (c, hc) :: rest => do + let (r, hr) ← mkTuple rest + let φ ← withLocalDecl `ω .default ΩE fun ω ↦ do + mkLambdaFVars #[ω] (← mkAppM ``Prod.mk #[(mkApp c ω).headBeta, (mkApp r ω).headBeta]) + pure (φ, ← mkAppM ``Measurable.prodMk #[hc, hr]) + let (φ, hφ) ← mkTuple comps.toList + let fE := Expr.fvar new.f + let tuple ← withLocalDecl `ω .default (.fvar new.Ω) fun ω ↦ do + mkLambdaFVars #[ω] (← Core.betaReduce (mkApp φ (mkApp fE ω))) + 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 X hX` with `X` a transported variable, the conditional law of `Z` given +`X` on the new space, `HasCondDistrib Z (fun ω ↦ X (f ω)) κ P'`. -/ +def deriveCondDistrib (g : MVarId) (new : NewSpace) (κ : Expr) (transported : FVarIdSet) : + TacticM (Option (FVarId × MVarId)) := g.withContext do + unless κ.isAppOfArity ``ProbabilityTheory.Kernel.comap 9 do return none + let X := κ.getArg! 7 + let .fvar x := X | return none + unless transported.contains x do return none + let κ₀ := κ.getArg! 6 + let fE := Expr.fvar new.f + let Xf ← withLocalDecl `ω .default (.fvar new.Ω) fun ω ↦ + mkLambdaFVars #[ω] (mkApp X (mkApp fE ω)) + 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) + +/-- 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, for a kernel `κ.comap X hX`, the conditional law of `Z` is stated given `X`. +With `clearOld`, the old space, the map and everything mentioning them are cleared. -/ +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. + let mut g := g + for (x, n) in olds do + unless n.hasMacroScopes do g ← g.rename 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 ty ← instantiateMVars (← x.getType) + 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 κ transportedSet 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 := toClear ++ hdefs ++ #[new.f, new.hf] ++ olds.map (·.1) + 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}" + -- 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 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₀) + 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" + let pick (defaults : Array Name) (i : Nat) : Name := + if h : i < given.size then given[i] else defaults[i]! + 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, 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, 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]!, if isKernel then none else some fvs[8]!⟩ + 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 ← saveState + -- `tryCatchRuntimeEx`: a failure inside `measurability` may be a maximum recursion depth error, + -- which `try … catch` lets through. + let transferLeft ← tryCatchRuntimeEx + (do + setGoals [transfer'.mvarId!] + evalTactic (← `(tactic| transfer)) + unless (← getUnsolvedGoals).isEmpty do throwError "transfer left goals" + pure []) + (fun _ ↦ do + s.restore + 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`, `hZ : HasLaw Z μ P`, and `hind`, the independence of `Z` from the transported + random variables, as a tuple; +* 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 Lean does not let a tactic clear: hypotheses introduced by `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`. + `κ` given the old space, `hZ : HasCondDistrib Z f κ P`, and no `hind`. For `κ.comap X hX` with + `X` a random variable, `hZ` is stated as `HasCondDistrib Z X κ P`. * `extend_space μ using P` names the measure to extend rather than reading it off the goal. -* `extend_space μ with Ω' P' f hf Z hZ hind` names what is introduced. +* `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`, `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?*]?) => withMainContext do - 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 - "extend_space: {μE} is neither a measure nor a kernel; it has type{indentExpr μty}" - -- 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 - "extend_space: the goal mentions no measure on a local space; name one with `using`" - | _ => throwError - "extend_space: the goal mentions several measures, {cands.map Expr.fvar}; choose one \ - with `using`" - let sp ← ProbSpace.ofMeasure PE - let ΩE := Expr.fvar sp.Ω - if let some dom := dom? then - unless ← isDefEq dom ΩE do - throwError "extend_space: 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 "extend_space: {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₀) - 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 - "extend_space: 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 "extend_space: `{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 "extend_space: 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 "extend_space: 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, then what was generalized, under primed names. - let extended := args[← idx `extended]!.mvarId! - extended.setKind .syntheticOpaque - extended.setTag `extended - let given := (names?.map (·.map (·.getId))).getD #[] - let defaults : Array Name := #[.mkSimple "Ω'", .mkSimple "P'", `f, `hf, `Z, `hZ, `hind] - let nNames := if isKernel then 6 else 7 - if given.size > nNames then - throwError "extend_space: at most {nNames} names may be given" - let pick (i : Nat) : Name := if h : i < given.size then given[i] else defaults[i]! - let mut intros : Array Name := #[pick 0, `inst, pick 1, `inst, pick 2, pick 3, pick 4, pick 5] - if !isKernel then intros := intros.push (pick 6) - let (_, extended) ← extended.introN intros.size intros.toList - let primed ← gens.toList.mapM fun f ↦ do pure ((← f.getUserName).appendAfter "'") - let (_, extended) ← extended.introN gens.size primed - -- 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 ← saveState - -- `tryCatchRuntimeEx`: a failure inside `measurability` may be a maximum recursion depth error, - -- which `try … catch` lets through. - let transferLeft ← tryCatchRuntimeEx - (do - setGoals [transfer'.mvarId!] - evalTactic (← `(tactic| transfer)) - unless (← getUnsolvedGoals).isEmpty do throwError "transfer left goals" - pure []) - (fun _ ↦ do - s.restore - pure [transfer'.mvarId!]) - setGoals ([extended] ++ transferLeft ++ rest) + | `(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 diff --git a/RandomDo/Probability/ExtendExamples.lean b/RandomDo/Probability/ExtendExamples.lean index ee8e4ae..34f674e 100644 --- a/RandomDo/Probability/ExtendExamples.lean +++ b/RandomDo/Probability/ExtendExamples.lean @@ -13,12 +13,11 @@ set_option linter.style.header false /-! # Tests and examples for `extend_space` and `transfer` -The first section pins down what `extend_space` produces: which hypotheses appear, what the goal -becomes, when the `transfer` obligation is closed automatically and when it is left, and how -hypotheses the goal depends on are handled. The next ones show `transfer` at work on laws, -independence, conditional laws, events, integrals and almost-everywhere statements, both on the -obligation and to pull hypotheses back to the new space; then a draw with a conditional law, an -i.i.d. sequence, and the errors the tactics report. +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, and what `extend_space!` clears. Then come the explicit form +`extend_space_map`, the `transfer` tactic on its own, a draw with a conditional law, an i.i.d. +sequence, and the errors the tactics report. Throughout, `Ω` lives in `Type u` and `E` in `Type`: the tactic lifts the product to the universe of `Ω`. @@ -39,192 +38,284 @@ variable {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasu include μ -/-! ## The shape of the goals -/ +/-! ## The context after `extend_space` -/ -/-- Random variables, families of random variables and events are transported along `f`, and the -measure becomes `P'`. With the measurability hypotheses around, the `transfer` obligation is -discharged by `extend_space` itself, and only the extended goal is left. -/ +/-- 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. 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) : - P.map X = P.map X ∧ (∀ n, P.map (A n) = P.map (A n)) ∧ P s = P s := by - extend_space μ with Ω' P' f hf Z hZ hind - guard_hyp hf : MeasurePreserving f P' P - guard_hyp hZ : HasLaw Z μ P' - guard_hyp hind : IndepFun f Z P' - guard_target =ₐ P'.map (fun ω ↦ X (f ω)) = P'.map (fun ω ↦ X (f ω)) - ∧ (∀ n, P'.map (fun ω ↦ A n (f ω)) = P'.map (fun ω ↦ A n (f ω))) - ∧ P' (f ⁻¹' s) = P' (f ⁻¹' s) - exact ⟨rfl, fun _ ↦ rfl, rfl⟩ - -/-- A statement `transfer` has no lemma for, here `IsProbabilityMeasure P`: the obligation is left. -It is stated for an arbitrary measure-preserving map, with the original goal as its conclusion. -Without `with`, the names are `Ω' P' f hf Z hZ hind`. -/ -example : IsProbabilityMeasure P := by + (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 hZ : HasLaw Z μ P + guard_hyp hind : IndepFun (fun ω ↦ (X ω, fun n ↦ A n ω)) 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 μ - case extended => - guard_hyp hf : MeasurePreserving f P' P - guard_hyp hZ : HasLaw Z μ P' - guard_hyp hind : IndepFun f Z P' - infer_instance - case transfer => - guard_target =ₐ ∀ (Ω' : Type u) [MeasurableSpace Ω'] (P' : Measure Ω') - [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → - IsProbabilityMeasure P' → IsProbabilityMeasure P - intro Ω' _ P' _ f hf h - infer_instance + 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₀ -/-- An operation `transfer` has no lemma for, here `Measure.restrict`: the obligation is left. -/ +/-- 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 (f ⁻¹' s) Set.univ = P' (f ⁻¹' s) + 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 and reintroduced as `hX'`, -now about `X ∘ f`, while `hX` itself stays. The goal mentions no measure, so `using P` says which -space to extend. `transfer` cannot rewrite under a binder the goal depends on, so the obligation -is left. -/ +/-- 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 with Ω' P' f hf Z hZ hind + extend_space μ using P case extended => guard_hyp hX : Measurable X - guard_hyp hX' : Measurable fun ω ↦ X (f ω) - guard_target =ₐ IsMarkovKernel (κ.comap (fun ω ↦ X (f ω)) hX') + guard_hyp hX₀ : Measurable X₀ + guard_target =ₐ IsMarkovKernel (κ.comap X hX) infer_instance case transfer => intro Ω' _ P' _ f hf h hX infer_instance -/-! ## Transferring statements +/-! ## `extend_space!` -/ -`extend_space` closes the `transfer` obligation, and `transfer hf at h` pulls a hypothesis about -the old space back to the new one. Both rewrite with the `@[transfer]` lemmas, whose side -conditions are measurability statements found by `assumption`, `fun_prop` or `measurability`. -/ - -/-- A law. -/ -example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : +/-- 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 μ with Ω' P' f hf Z hZ hind - transfer hf at hXν - guard_hyp hXν : HasLaw (fun ω ↦ X (f ω)) ν P' + 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 -/-- The same with `HasLaw` as the goal, and the hypothesis pulled back by hand. -/ -example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : - HasLaw X ν P := by - extend_space μ - exact hXν.comp hf.hasLaw +/-! ## 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 μ - transfer hf at hXY - guard_hyp hXY : IndepFun (fun ω ↦ X (f ω)) (fun ω ↦ Y (f ω)) P' + 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 μ - transfer hf at hXY exact hXY -/-- An event, given by a set-builder expression: `measurability` proves it measurable. -/ -example (X : Ω → ℝ) (hX : Measurable X) (h : P {ω | 0 < X ω} = 1 / 2) : - P {ω | 0 < X ω} = 1 / 2 := by +/-- 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 μ - transfer hf at h - guard_hyp h : P' {ω | 0 < X (f ω)} = 1 / 2 - exact h - -/-- The real-valued measure of an event. -/ -example (s : Set Ω) (hs : MeasurableSet s) (r : ℝ) (h : P.real s = r) : P.real s = r := by - extend_space μ - transfer hf at h - exact h + guard_hyp h : P {ω | 0 < X ω} = 1 / 2 + exact ⟨h, h'⟩ -/-- An integral, with an almost-everywhere hypothesis. -/ -example (X : Ω → ℝ) (hX : Measurable X) (h : ∀ᵐ ω ∂P, 0 ≤ X ω) : 0 ≤ ∫ ω, X ω ∂P := by +/-- 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 μ - transfer hf at h - guard_hyp h : ∀ᵐ ω ∂P', 0 ≤ X (f ω) - exact integral_nonneg_of_ae h + 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 μ - transfer hf at h 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 μ - transfer hf at h - guard_hyp h : (fun ω ↦ X (f ω)) =ᵐ[P'] fun ω ↦ Y (f ω) exact h -/-- `transfer hf` on a goal, outside of `extend_space`: the goal is moved to the new space and -closed by the hypothesis. -/ -example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) {Ω' : Type u} [MeasurableSpace Ω'] - {P' : Measure Ω'} (f : Ω' → Ω) (hf : MeasurePreserving f P' P) - (h : HasLaw (fun ω ↦ X (f ω)) ν P') : - HasLaw X ν P := by - transfer hf - /-! ## 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: -`X ∘ f` and `Z`. -/ +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 with Ω' P' f hf Z hZ hind - exact ⟨Ω', inferInstance, P', inferInstance, fun ω ↦ X (f ω), Z, hXν.comp hf.hasLaw, hZ, - hind.comp hX measurable_id⟩ + 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 with Ω' P' f hf Z hZ hind - have hZn (n : ℕ) : HasLaw (fun ω ↦ Z ω n) μ P' := + 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, fun ω ↦ X (f ω), fun n ω ↦ Z ω n, - hXν.comp hf.hasLaw, hZn, (iIndepFun_iff_hasLaw_Pi_infinitePi hZn hZ.aemeasurable).2 hZ, - hind.comp hX measurable_id⟩ + 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 κ` for a kernel on `Ω`: the draw has conditional law `κ` given the old space, and -there is no independence hypothesis. For a law conditional on `X`, extend with `κ.comap X hX` and -read `hZ` through `HasCondDistrib.comp_right`. -/ +/-- `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 + +/-! ## 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 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 (κ.comap X hX) with Ω' P' f hf Z hZ + 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 +/-! ## The `transfer` tactic on its own -/ + +/-- `transfer hf` on a goal: the goal is moved to the new space and closed by the hypothesis. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) {Ω' : Type u} [MeasurableSpace Ω'] + {P' : Measure Ω'} (f : Ω' → Ω) (hf : MeasurePreserving f P' P) + (h : HasLaw (fun ω ↦ X (f ω)) ν P') : + HasLaw X ν P := by + transfer hf + +/-- `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) + {Ω' : Type u} [MeasurableSpace Ω'] {P' : Measure Ω'} (f : Ω' → Ω) + (hf : MeasurePreserving f P' P) : + 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⟩ + /-! ## Universes -/ /-- `Ω` and `E` in the same universe. -/ example {E' : Type u} [MeasurableSpace E'] (μ' : Measure E') [IsProbabilityMeasure μ'] - (X : Ω → ℝ) (hX : Measurable X) : P.map X = P.map X := by + (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) (hXν : HasLaw X ν P) : P.map X = ν := by extend_space μ' - rfl + exact hXν.map_eq /-! ## Errors -/ diff --git a/RandomDo/Probability/Transfer.lean b/RandomDo/Probability/Transfer.lean index d515092..cb53be4 100644 --- a/RandomDo/Probability/Transfer.lean +++ b/RandomDo/Probability/Transfer.lean @@ -19,19 +19,25 @@ set_option linter.style.header false 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. A lemma tagged `@[transfer]` states one such 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. - -`transfer hf` rewrites the goal with every such 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` rewrites hypotheses instead: this pulls a fact about the old space back -to the new one. `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 +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: this pulls a fact about the old space back to +the new one. `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. -/ @@ -41,10 +47,15 @@ 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. The `transfer` tactic -rewrites with all of them. -/ +`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 discharger for the side conditions of `@[transfer]` lemmas: `assumption`, then `fun_prop` @@ -63,12 +74,61 @@ elab_rules : tactic tryCatchRuntimeEx (evalTactic tac) fun e ↦ throwError "transfer_discharger: {e.toMessageData}" -/-- The `@[transfer]` lemmas instantiated at `hf`, as `simp` arguments, together with -`Set.preimage_ofPred_eq`, which puts the transferred events in the same form as `extend_space`. -/ +/-- The `@[transfer]` lemmas instantiated at `hf`, as `simp` arguments, together with the lemmas +pushing a preimage through set operations, which put the transferred events in the same form as +`extend_space`. -/ def transferSimpArgs (hf : Term) : CoreM (Array (TSyntax ``Lean.Parser.Tactic.simpLemma)) := do let args ← (← labelled `transfer).mapM fun n ↦ `(Lean.Parser.Tactic.simpLemma| $(mkIdent n):ident $hf) - return args.push (← `(Lean.Parser.Tactic.simpLemma| Set.preimage_ofPred_eq)) + let extra ← #[``Set.preimage_ofPred_eq, ``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. -/ +def tryTactic? (g : MVarId) (tac : Syntax) : TacticM (Option Expr) := do + let s ← saveState + tryCatchRuntimeEx + (do + let gs ← Tactic.run g (evalTactic tac) + if gs.isEmpty then return some (← instantiateMVars (.mvar g)) + s.restore + return none) + (fun _ ↦ do + s.restore + 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`. Returns the +statement and proof of the transported hypothesis. -/ +def transferForward? (h hf : Expr) : TacticM (Option (Expr × Expr)) := do + let hStx ← Term.exprToSyntax h + let hfStx ← Term.exprToSyntax hf + for n in ← labelled `transfer_forward do + let s ← saveState + 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 (concl, 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 @@ -89,14 +149,42 @@ where return (← whnfR (← fv.getType)).isAppOf ``MeasureTheory.MeasurePreserving go g (if isMap then some fv else none) +/-- Transfer the goal along `hf`: rewrite it with the `@[transfer]` lemmas, then close it by +`assumption` if possible. -/ +def transferGoal (hfStx : Term) : TacticM Unit := do + let args ← transferSimpArgs hfStx + evalTactic (← `(tactic| simp -failIfUnchanged (disch := transfer_discharger) only [$args,*])) + unless (← getUnsolvedGoals).isEmpty do + evalTactic (← `(tactic| try assumption)) + +/-- 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. -/ +def transferHyps (hfStx : Term) (hf : Expr) (hs : Array FVarId) : TacticM Unit := do + let args ← transferSimpArgs hfStx + 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)) + for h in hs do + withMainContext do + -- `simp` replaces the hypothesis when it rewrites it; otherwise it is still there. + let some d := (← getLCtx).find? h | return + let some (ty, pf) ← transferForward? d.toExpr hf | return + let g ← (← getMainGoal).assert d.userName ty pf + let (_, g) ← g.intro1P + replaceMainGoal [← g.tryClear h] + /-- `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. -* `transfer hf at h₁ h₂` rewrites hypotheses instead: a fact about the old space becomes the - corresponding fact about the new one. +* `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. * `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. -/ @@ -108,13 +196,20 @@ elab_rules : tactic match hf?, loc? with | none, some _ => throwError "transfer: `at` needs the map to transfer along, as in `transfer hf at h`" - | some hf, _ => - let args ← transferSimpArgs hf - evalTactic (← `(tactic| - simp -failIfUnchanged (disch := transfer_discharger) only [$args,*] $(loc?)?)) - if loc?.isNone then - unless (← getUnsolvedGoals).isEmpty do - evalTactic (← `(tactic| try assumption)) + | some hf, some loc => + let hfE ← Tactic.elabTerm hf none + match expandLocation loc with + | .wildcard => + let hs ← withMainContext do + (← getLCtx).foldlM (init := #[]) fun hs d ↦ do + if d.isImplementationDetail || !(← isProp d.type) then pure hs + else pure (hs.push d.fvarId) + transferHyps hf hfE hs + transferGoal hf + | .targets hyps type => + transferHyps hf hfE (← withMainContext do hyps.mapM getFVarId) + if type then transferGoal hf + | some hf, none => transferGoal hf | none, none => let (g, hf, h) ← introTransferObligation (← getMainGoal) replaceMainGoal [g] From 969f657b3d1e88b57d6ae202d50a7b531cdd9970 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 4 Sep 2026 16:28:43 +0200 Subject: [PATCH 07/34] use transfer tactic in alg_env_trace --- RandomDo/Probability/AlgTrace.lean | 74 +++++++++++++++++++++++------- RandomDo/Probability/Extend.lean | 10 ++-- RandomDo/Probability/Transfer.lean | 59 ++++++++++++++++++++---- 3 files changed, 114 insertions(+), 29 deletions(-) diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean index b242ecb..1fc4a1b 100644 --- a/RandomDo/Probability/AlgTrace.lean +++ b/RandomDo/Probability/AlgTrace.lean @@ -58,6 +58,8 @@ open MeasureTheory ProbabilityTheory Finset Learning noncomputable section +attribute [fun_prop] Learning.measurable_history + /-- 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. -/ @@ -342,6 +344,8 @@ 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. -/ @@ -403,7 +407,9 @@ sequences, is abstracted away from that space and two goals are left: 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`. + 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. * `alg_env_trace tr using h` names the hypothesis to use rather than searching for one. @@ -437,11 +443,13 @@ elab_rules : tactic let e ← Term.elabTerm tr none Term.synthesizeSyntheticMVarsNoPostponing instantiateMVars e - let (traced, transfer) ← g.withContext do + 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 @@ -459,7 +467,7 @@ elab_rules : tactic transfer.setKind .syntheticOpaque traced.setTag `traced transfer.setTag `transfer - return (traced, transfer) + return (traced, transfer, args[iMotive]!) -- Introduce the traced space and its properties, then whatever travelled with the goal. let given := (names?.map (·.map (·.getId))).getD #[] let defaults : Array Name := #[`Ω, `P, `A, `Y, `T, `hseq, `hT₀, `hT, `hA₀, `hA] @@ -469,7 +477,37 @@ elab_rules : tactic pick 8, pick 9] let (_, traced) ← traced.introN intros.size intros.toList let (_, traced) ← traced.introNP deps.size - replaceMainGoal [traced, transfer] + -- 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) + -- 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 + generalize hν : Measure.map (trajectory A₁ Y₁) P₁ = ν at hlaw + have hf₁ : MeasurePreserving (trajectory A₁ Y₁) P₁ ν := + ⟨measurable_trajectory h₁.measurable_action h₁.measurable_feedback, hν⟩ + have hf₂ : MeasurePreserving (trajectory A₂ Y₂) P₂ ν := + ⟨measurable_trajectory h₂.measurable_action h₂.measurable_feedback, hlaw⟩ + have : IsProbabilityMeasure ν := hν ▸ Measure.isProbabilityMeasure_map + (measurable_trajectory h₁.measurable_action h₁.measurable_feedback).aemeasurable + exact (fun hS : $motiveStx _ inferInstance ν inferInstance + (fun n t ↦ (t n).1) (fun n t ↦ (t 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 @@ -559,25 +597,29 @@ theorem exists_noise (env : Environment (Fin K) ℝ) {Ω₀ : Type*} [Measurable 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, and leaves the obligation that the statement only -depends on the law of the trajectory. Any hypothesis mentioning the space travels with the goal, so -nothing is silently lost. -/ +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 mentioning the space 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 hZ₀ hZ hA₀ hA - case traced => - -- `Z`, `hZ₀`, `hZ` and `hA` are the algorithm's draws and their laws, now available. - exact hseq.hasLaw_action_zero.map_eq + -- `Z`, `hZ₀`, `hZ` and `hA` are the algorithm's draws and their laws, now available. + 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₀ - have e₁ : A₁ 0 = (fun t ↦ (t 0).1) ∘ trajectory A₁ Y₁ := rfl - have e₂ : A₂ 0 = (fun t ↦ (t 0).1) ∘ trajectory A₂ Y₂ := rfl - rw [e₁, ← Measure.map_map (by fun_prop) - (measurable_trajectory h₁.measurable_action h₁.measurable_feedback), ← hlaw, - Measure.map_map (by fun_prop) - (measurable_trajectory h₂.measurable_action h₂.measurable_feedback), ← e₂, h₀] + infer_instance /-- **`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 diff --git a/RandomDo/Probability/Extend.lean b/RandomDo/Probability/Extend.lean index e36efd9..aa1902a 100644 --- a/RandomDo/Probability/Extend.lean +++ b/RandomDo/Probability/Extend.lean @@ -332,6 +332,8 @@ 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 @@ -903,17 +905,19 @@ def extendSpace (mode : ExtendMode) (μ : Term) (P? : Option Ident) (given : Arr transfer.assign transfer' -- Discharge the transfer obligation with the `transfer` tactic when it can; leave it otherwise. let rest := (← getGoals).drop 1 - let s ← saveState + 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!] - evalTactic (← `(tactic| transfer)) + Term.withoutErrToSorry <| evalTactic (← `(tactic| transfer)) unless (← getUnsolvedGoals).isEmpty do throwError "transfer left goals" pure []) - (fun _ ↦ do + (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) diff --git a/RandomDo/Probability/Transfer.lean b/RandomDo/Probability/Transfer.lean index cb53be4..4996997 100644 --- a/RandomDo/Probability/Transfer.lean +++ b/RandomDo/Probability/Transfer.lean @@ -58,6 +58,23 @@ 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 } + /-- 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 @@ -68,18 +85,39 @@ elab_rules : tactic | `(tactic| transfer_discharger) => withMainContext do let funProps : Array Name := #[``Measurable, ``AEMeasurable, `MeasureTheory.AEStronglyMeasurable, `MeasureTheory.StronglyMeasurable] - let head := (← getMainTarget).getAppFn.constName? - let tac ← if head.any funProps.contains then `(tactic| first | assumption | fun_prop) - else `(tactic| first | assumption | measurability) + 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)) tryCatchRuntimeEx (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, which put the transferred events in the same form as -`extend_space`. -/ -def transferSimpArgs (hf : Term) : CoreM (Array (TSyntax ``Lean.Parser.Tactic.simpLemma)) := do - let args ← (← labelled `transfer).mapM fun n ↦ - `(Lean.Parser.Tactic.simpLemma| $(mkIdent n):ident $hf) +`extend_space`. 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_inter, ``Set.preimage_union, ``Set.preimage_compl, ``Set.preimage_sdiff].mapM fun n ↦ `(Lean.Parser.Tactic.simpLemma| $(mkIdent n):ident) @@ -88,10 +126,11 @@ def transferSimpArgs (hf : Term) : CoreM (Array (TSyntax ``Lean.Parser.Tactic.si /-- 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. -/ def tryTactic? (g : MVarId) (tac : Syntax) : TacticM (Option Expr) := do - let s ← saveState + let s ← saveFullState tryCatchRuntimeEx (do - let gs ← Tactic.run g (evalTactic tac) + -- Without error recovery, a failure inside a nested `by` is a failure, not a `sorry`. + let gs ← Term.withoutErrToSorry <| Tactic.run g (evalTactic tac) if gs.isEmpty then return some (← instantiateMVars (.mvar g)) s.restore return none) @@ -106,7 +145,7 @@ def transferForward? (h hf : Expr) : TacticM (Option (Expr × Expr)) := do let hStx ← Term.exprToSyntax h let hfStx ← Term.exprToSyntax hf for n in ← labelled `transfer_forward do - let s ← saveState + let s ← saveFullState let r ← tryCatchRuntimeEx (do let e ← Term.withoutErrToSorry <| From ab21c238d155b3efbf3d4a93189e7dfdbeff6b4a Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 4 Sep 2026 16:41:11 +0200 Subject: [PATCH 08/34] move tests to Test --- RandomDo.lean | 2 +- RandomDo/Probability/AlgTrace.lean | 148 ------------ RandomDo/Probability/Extend.lean | 186 +-------------- RandomDo/Probability/MeasurePreserving.lean | 213 ++++++++++++++++++ Test.lean | 3 + Test/AlgTrace.lean | 154 +++++++++++++ .../ExtendExamples.lean => Test/Extend.lean | 52 +---- Test/Transfer.lean | 62 +++++ 8 files changed, 445 insertions(+), 375 deletions(-) create mode 100644 RandomDo/Probability/MeasurePreserving.lean create mode 100644 Test/AlgTrace.lean rename RandomDo/Probability/ExtendExamples.lean => Test/Extend.lean (88%) create mode 100644 Test/Transfer.lean diff --git a/RandomDo.lean b/RandomDo.lean index 27071c9..98c6e73 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -9,7 +9,7 @@ public import RandomDo.Monad.Notation public import RandomDo.Probability.AlgTrace public import RandomDo.Probability.Examples public import RandomDo.Probability.Extend -public import RandomDo.Probability.ExtendExamples +public import RandomDo.Probability.MeasurePreserving public import RandomDo.Probability.Record public import RandomDo.Probability.Tactic public import RandomDo.Probability.Thompson diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean index 1fc4a1b..565fe59 100644 --- a/RandomDo/Probability/AlgTrace.lean +++ b/RandomDo/Probability/AlgTrace.lean @@ -512,151 +512,3 @@ elab_rules : tactic end RDo.Tactic end - -@[expose] public section - -open MeasureTheory ProbabilityTheory Finset Learning RDo - -noncomputable section - -/-! ## An example: an algorithm whose policy is an `rdo` program - -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. - -To do the same for `thompson` one needs the measurable equivalence between `Iic n → 𝓐 × 𝓨` and -`Vector (𝓐 × 𝓨) (n + 1)` that turns it into a policy — `Vector.v_equiv` in -`RandomDo.Tactic.Examples`, still a `sorry` there (and stated one element short). Everything after -that point is what follows below. --/ - -namespace RDo.Example - -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 mentioning the space 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 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 - -/-- 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 - -/-- **`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 RDo.Example - -end - -end diff --git a/RandomDo/Probability/Extend.lean b/RandomDo/Probability/Extend.lean index aa1902a..ecc661f 100644 --- a/RandomDo/Probability/Extend.lean +++ b/RandomDo/Probability/Extend.lean @@ -5,11 +5,7 @@ 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 +public import RandomDo.Probability.MeasurePreserving public import Mathlib.Probability.Kernel.Composition.MeasureCompProd public meta import Lean.Elab.Tactic.Basic @@ -55,11 +51,9 @@ statement by statement. Here the projection `f` is measure preserving, which is ## Main results * `RDo.wlog_extend`, `RDo.wlog_extend_kernel`: the principles behind the tactic. -* `MeasureTheory.MeasurePreserving.map_fun_comp`, `hasLaw_fun_comp_iff`, `indepFun_fun_comp_iff`, - `hasCondDistrib_fun_comp_iff`: pulling statements back along a measure-preserving map. -* The `@[transfer]` lemmas `MeasurePreserving.transfer_*` and the `@[transfer_forward]` lemmas - `*.comp_measurePreserving`, `*.preimage_measurePreserving`: the same facts in the forms the - `transfer` tactic uses. +* `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 @@ -68,178 +62,6 @@ 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_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 : Ω → 𝓨} - -@[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 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 - namespace RDo /-! ### The product extension -/ diff --git a/RandomDo/Probability/MeasurePreserving.lean b/RandomDo/Probability/MeasurePreserving.lean new file mode 100644 index 0000000..e306cc8 --- /dev/null +++ b/RandomDo/Probability/MeasurePreserving.lean @@ -0,0 +1,213 @@ +/- +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, 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_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 : Ω → 𝓨} + +@[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 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/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..75e2dc3 --- /dev/null +++ b/Test/AlgTrace.lean @@ -0,0 +1,154 @@ +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 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 + +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 mentioning the space 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 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 + +/-- 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 + +/-- **`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/RandomDo/Probability/ExtendExamples.lean b/Test/Extend.lean similarity index 88% rename from RandomDo/Probability/ExtendExamples.lean rename to Test/Extend.lean index 34f674e..1a7b6a6 100644 --- a/RandomDo/Probability/ExtendExamples.lean +++ b/Test/Extend.lean @@ -1,35 +1,30 @@ -/- -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.Extend +public import Test.Common public import Mathlib.Probability.Independence.InfinitePi set_option linter.style.header false /-! -# Tests and examples for `extend_space` and `transfer` +# 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, and what `extend_space!` clears. Then come the explicit form -`extend_space_map`, the `transfer` tactic on its own, a draw with a conditional law, an i.i.d. -sequence, and the errors the tactics report. +`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 `Ω`. -/ -@[expose] public section - open MeasureTheory ProbabilityTheory RDo +@[expose] public section + noncomputable section -namespace RDo.Example.Extend +namespace Test.Extend universe u @@ -288,27 +283,6 @@ example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel ℝ E) [IsMarkovKernel transfer hf at hXν exact hXν.map_eq -/-! ## The `transfer` tactic on its own -/ - -/-- `transfer hf` on a goal: the goal is moved to the new space and closed by the hypothesis. -/ -example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) {Ω' : Type u} [MeasurableSpace Ω'] - {P' : Measure Ω'} (f : Ω' → Ω) (hf : MeasurePreserving f P' P) - (h : HasLaw (fun ω ↦ X (f ω)) ν P') : - HasLaw X ν P := by - transfer hf - -/-- `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) - {Ω' : Type u} [MeasurableSpace Ω'] {P' : Measure Ω'} (f : Ω' → Ω) - (hf : MeasurePreserving f P' P) : - 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⟩ - /-! ## Universes -/ /-- `Ω` and `E` in the same universe. -/ @@ -342,17 +316,7 @@ example {E' : Type (u + 1)} [MeasurableSpace E'] (μ' : Measure E') [IsProbabili (X : Ω → ℝ) : P.map X = P.map X := by extend_space μ' -/-- -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 - -end RDo.Example.Extend +end Test.Extend end diff --git a/Test/Transfer.lean b/Test/Transfer.lean new file mode 100644 index 0000000..492c9fb --- /dev/null +++ b/Test/Transfer.lean @@ -0,0 +1,62 @@ +module + +public import Test.Common + +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: on a goal, and on hypotheses, where the rewriting +by the `@[transfer]` lemmas and the fallback on the `@[transfer_forward]` lemmas both show. +-/ + +open MeasureTheory ProbabilityTheory RDo + +@[expose] public section + +noncomputable section + +namespace Test.Transfer + +universe u + +variable {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} + +/-- `transfer hf` on a goal: the goal is moved to the new space and closed by the hypothesis. -/ +example (X : Ω → ℝ) (hX : Measurable X) (ν : Measure ℝ) {Ω' : Type u} [MeasurableSpace Ω'] + {P' : Measure Ω'} (f : Ω' → Ω) (hf : MeasurePreserving f P' P) + (h : HasLaw (fun ω ↦ X (f ω)) ν P') : + HasLaw X ν P := by + transfer hf + +/-- `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) + {Ω' : Type u} [MeasurableSpace Ω'] {P' : Measure Ω'} (f : Ω' → Ω) + (hf : MeasurePreserving f P' P) : + 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⟩ + +/-! ## 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 + +end Test.Transfer + +end + +end From 950a0ce26a7ae27fb1a29e234b06ff3729f2cd38 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 4 Sep 2026 16:52:16 +0200 Subject: [PATCH 09/34] minor --- RandomDo/Probability/Thompson.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/RandomDo/Probability/Thompson.lean b/RandomDo/Probability/Thompson.lean index 1df4afe..dae6539 100644 --- a/RandomDo/Probability/Thompson.lean +++ b/RandomDo/Probability/Thompson.lean @@ -14,7 +14,7 @@ set_option linter.style.header false /-! # Thompson sampling as random variables -`thompson` in `RandomDo.Tactic.Examples` is an `rdo` program: a loop folding the history into +`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. From ae5f52fa5bb69f4fbf464a7cd45b0c1b960bdf71 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 4 Sep 2026 16:58:37 +0200 Subject: [PATCH 10/34] fix lint --- RandomDo/Probability/Record.lean | 6 ++++-- RandomDo/Probability/Tactic.lean | 2 ++ RandomDo/Probability/Thompson.lean | 3 +++ Test/Gaps.lean | 5 +++++ Test/Instances.lean | 2 ++ Test/IsMarkov.lean | 8 ++++++++ 6 files changed, 24 insertions(+), 2 deletions(-) diff --git a/RandomDo/Probability/Record.lean b/RandomDo/Probability/Record.lean index 965a1a1..d7257e1 100644 --- a/RandomDo/Probability/Record.lean +++ b/RandomDo/Probability/Record.lean @@ -69,9 +69,11 @@ 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. -/ +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 >>=ₘ f = f x := +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: diff --git a/RandomDo/Probability/Tactic.lean b/RandomDo/Probability/Tactic.lean index 3cd7467..76a8fbd 100644 --- a/RandomDo/Probability/Tactic.lean +++ b/RandomDo/Probability/Tactic.lean @@ -444,7 +444,9 @@ partial def constValue? (K ι : Expr) : MetaM (Option Expr) := do /-- 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 diff --git a/RandomDo/Probability/Thompson.lean b/RandomDo/Probability/Thompson.lean index dae6539..db0c1a4 100644 --- a/RandomDo/Probability/Thompson.lean +++ b/RandomDo/Probability/Thompson.lean @@ -99,6 +99,9 @@ instance : IsMarkovKernel (sampleK (K := K)) := by unfold sampleK; infer_instanc @[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 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 From 5efde9a58a150ca7f6640e7e924bf4b2affffb28 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 7 Sep 2026 16:03:42 +0200 Subject: [PATCH 11/34] Computable tactic --- RandomDo.lean | 14 ++- RandomDo/Tactic/Computable/Counterparts.lean | 34 ++++++ RandomDo/Tactic/Computable/Defs.lean | 52 ++++++++ RandomDo/Tactic/Computable/Deriving.lean | 115 ++++++++++++++++++ RandomDo/Tactic/Computable/Example.lean | 40 ++++++ .../{IsMarkov.lean => IsMarkov/Defs.lean} | 0 RandomDo/Tactic/{ => IsMarkov}/Deriving.lean | 2 +- RandomDo/Tactic/{ => IsMarkov}/Elab.lean | 2 +- RandomDo/Tactic/{ => IsMarkov}/ForInStep.lean | 2 +- RandomDo/Tactic/{ => IsMarkov}/Lemmas.lean | 2 +- 10 files changed, 254 insertions(+), 9 deletions(-) create mode 100644 RandomDo/Tactic/Computable/Counterparts.lean create mode 100644 RandomDo/Tactic/Computable/Defs.lean create mode 100644 RandomDo/Tactic/Computable/Deriving.lean create mode 100644 RandomDo/Tactic/Computable/Example.lean rename RandomDo/Tactic/{IsMarkov.lean => IsMarkov/Defs.lean} (100%) rename RandomDo/Tactic/{ => IsMarkov}/Deriving.lean (98%) rename RandomDo/Tactic/{ => IsMarkov}/Elab.lean (99%) rename RandomDo/Tactic/{ => IsMarkov}/ForInStep.lean (98%) rename RandomDo/Tactic/{ => IsMarkov}/Lemmas.lean (99%) diff --git a/RandomDo.lean b/RandomDo.lean index 1b94f05..ee6fc15 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -11,8 +11,12 @@ public import RandomDo.NumLean.PCG64 public import RandomDo.NumLean.SeedSequence public import RandomDo.NumLean.Ziggurat public import RandomDo.NumLean.ZigguratSampler -public import RandomDo.Tactic.Deriving -public import RandomDo.Tactic.Elab -public import RandomDo.Tactic.ForInStep -public import RandomDo.Tactic.IsMarkov -public import RandomDo.Tactic.Lemmas +public import RandomDo.Tactic.Computable.Counterparts +public import RandomDo.Tactic.Computable.Defs +public import RandomDo.Tactic.Computable.Deriving +public import RandomDo.Tactic.Computable.Example +public import RandomDo.Tactic.IsMarkov.Defs +public import RandomDo.Tactic.IsMarkov.Deriving +public import RandomDo.Tactic.IsMarkov.Elab +public import RandomDo.Tactic.IsMarkov.ForInStep +public import RandomDo.Tactic.IsMarkov.Lemmas diff --git a/RandomDo/Tactic/Computable/Counterparts.lean b/RandomDo/Tactic/Computable/Counterparts.lean new file mode 100644 index 0000000..75384fa --- /dev/null +++ b/RandomDo/Tactic/Computable/Counterparts.lean @@ -0,0 +1,34 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import RandomDo.Tactic.Computable.Defs +public import RandomDo.NumLean.Distributions +public import Mathlib.Probability.Distributions.Gaussian.Real + +/-! +# Computable counterparts of the pieces an `rdo` program is made of + +The `@[computable]` attribute translates a program by replacing each piece it is made of by the +counterpart recorded here through `@[computable_as]`. There is one entry per piece the programs of +`RandomDo.Tactic.Computable.Example` are built from: the two types their values live in, and the +one distribution they draw from. +-/ + +public meta section + +/-! ## Types -/ + +attribute [computable_as Float] Real +attribute [computable_as Float] NNReal + +/-! ## Distributions -/ + +/- `gaussianReal` reads its second argument as a variance and `normal` reads it as a standard +deviation: the two agree on the `1` the example draws with, not in general. -/ +attribute [computable_as NumLean.normal] ProbabilityTheory.gaussianReal + +end diff --git a/RandomDo/Tactic/Computable/Defs.lean b/RandomDo/Tactic/Computable/Defs.lean new file mode 100644 index 0000000..ffbc027 --- /dev/null +++ b/RandomDo/Tactic/Computable/Defs.lean @@ -0,0 +1,52 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public meta import Batteries.Lean.NameMapAttribute +public meta import Lean.ReservedNameAction + +/-! +# The `@[computable_as]` attribute + +An `rdo` program is written over the Giry monad: it draws from measures on `ℝ`, which no machine +samples. Turning it into a program that runs, which the `@[computable]` attribute of +`RandomDo.Tactic.Computable.Deriving` does, asks for a counterpart of each piece the program is +built from: `Float` for `ℝ`, `NumLean.normal` for `gaussianReal`. This file holds the attribute +recording them; `RandomDo.Tactic.Computable.Counterparts` holds the counterparts themselves. + +`@[computable_as f]` on a declaration `d` reads: `f` is what `d` becomes in a translated program. +Only the pieces denoting something the program computes with need one. The scaffolding around +them — numerals, arithmetic — is polymorphic, and the translation keeps it as it is, at the +translated types. +-/ + +public meta section + +open Lean + +namespace RDo.Tactic + +/-- The counterparts recorded by `@[computable_as]`, keyed by the declaration they translate. -/ +initialize computableAsExt : NameMapExtension Name ← + registerNameMapAttribute { + name := `computable_as + descr := "record the computable counterpart of this declaration" + /- `@[computable_as f]` is read by `Lean.Parser.Attr.simple`, the parser an attribute that + declares no syntax of its own gets: `f` is the single child of `stx[1]`. -/ + add := fun _ stx ↦ do + let f := stx[1][0] + unless f.isIdent do + throwError "`computable_as` takes the name of one declaration" + realizeGlobalConstNoOverload f + } + +/-- The computable counterpart of `declName`, when `@[computable_as]` recorded one. -/ +def computableAs? (declName : Name) : CoreM (Option Name) := + return computableAsExt.find? (← getEnv) declName + +end RDo.Tactic + +end diff --git a/RandomDo/Tactic/Computable/Deriving.lean b/RandomDo/Tactic/Computable/Deriving.lean new file mode 100644 index 0000000..dc44ee4 --- /dev/null +++ b/RandomDo/Tactic/Computable/Deriving.lean @@ -0,0 +1,115 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import RandomDo.Tactic.Computable.Counterparts +public import RandomDo.Monad.MeasurableSpace +public meta import Lean.Elab.Tactic.Basic +public meta import Batteries.Tactic.Lint.Basic + +/-! +# The `@[computable]` attribute + +Writing an `rdo` program and writing the program that samples from it are two separate steps, and +the second is mechanical: every construct of the first has a counterpart in the second. This +attribute walks the program and writes that counterpart down, so a program carries the sampler it +denotes: + +``` +@[computable] +noncomputable def shifted : Measure ℝ := rdo + let x ← gaussianReal 0 1 + return x + 1 +``` + +adds `shifted_computable : RandPCG IO Float`, which draws from `NumLean.normal 0 1` and adds one. + +## How the program is read + +The two constructs `rdo` is made of become the two of `do`: `return e` becomes `pure e`, and +`let x ← p; q` becomes `p >>= q`, at `RandPCG IO`. Everything else — a distribution, a numeral, an +operation on the values the program computes with — is rebuilt from the counterpart +`@[computable_as]` records for its head, applied to the translation of its arguments; the +instances it asks for are synthesized anew, for the translated types. +-/ + +public meta section + +open Lean Meta NumLean + +namespace RDo.Tactic + +/-- The monad a translated program lives in, `RandPCG IO`: the counterpart of the Giry monad an +`rdo` program is written over. -/ +def computableMonad : MetaM Expr := mkAppM ``RandPCG #[mkConst ``IO] + +/-- Translate `e`, a piece of an `rdo` program, into its computable counterpart. `σ` sends the +variables the program binds to the ones the translated program binds in their place, which the +change of types makes necessary. -/ +partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := do + -- `MeasurableSpacePure.mPure` takes five arguments, the last of which is the value returned. + if e.isAppOfArity ``MeasurableSpacePure.mPure 5 then + mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ (e.getArg! 4)] + -- `MeasurableSpaceBind.mBind` takes eight, the last two being the program and the continuation. + else if e.isAppOfArity ``MeasurableSpaceBind.mBind 8 then + mkAppOptM ``Bind.bind #[← computableMonad, none, none, none, + ← translate σ (e.getArg! 6), ← translate σ (e.getArg! 7)] + else match e with + | .fvar fvarId => return σ.get fvarId + | .sort .. | .lit .. => return e + | .mdata _ b => translate σ b + | .lam .. => + lambdaBoundedTelescope e 1 fun xs body ↦ do + let x := xs[0]! + let t ← translate σ (← x.fvarId!.getType) + withLocalDeclD (← x.fvarId!.getUserName) t fun y ↦ do + mkLambdaFVars #[y] (← translate (σ.insert x.fvarId! y) body) + | _ => + let .const declName _ := e.getAppFn + | throwError "`computable`: cannot translate{indentExpr e}" + let counterpart := (← computableAs? declName).getD declName + let mut f ← mkConstWithFreshMVarLevels counterpart + for arg in e.getAppArgs do + let .forallE _ t _ bi ← whnf (← inferType f) + | throwError "`computable`: {f} does not take the argument{indentExpr arg}" + let arg ← if bi.isInstImplicit then synthInstance t else translate σ arg + unless ← isDefEq (← inferType arg) t do + throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}" + f := mkApp f arg + return f + +/-- Translate the `rdo` program `declName` and add the translation to the environment, under the +name `declName` followed by `_computable`. -/ +def addComputableDecl (declName : Name) : MetaM Unit := do + let info ← getConstInfo declName + let some value := info.value? + | throwError "`computable` can only be derived for a definition, but {declName} has no value" + let value ← instantiateMVars (← translate {} value) + let type ← instantiateMVars (← inferType value) + let translated := declName.appendAfter "_computable" + addAndCompile <| .defnDecl <| ← mkDefinitionValInferringUnsafe translated info.levelParams type + value (.regular (getMaxHeight (← getEnv) value + 1)) + addDocStringCore translated s!"The program that samples from `{declName}`, written by the \ + `@[computable]` attribute." + /- The name is one the attribute picks and not one the user wrote, so the underscore in it is + reported for every program translated unless it is exempted here. -/ + setEnv (← ofExcept (Batteries.Tactic.Lint.nolintAttr.setParam (← getEnv) translated + #[`defsWithUnderscore])) + +/-- The `@[computable]` attribute. -/ +initialize registerBuiltinAttribute { + name := `computable + descr := "translate this `rdo` program into the program that samples from it" + applicationTime := .afterCompilation + add := fun declName _stx kind ↦ do + unless kind == AttributeKind.global do + throwError "`computable` must be a global attribute" + (addComputableDecl declName).run' +} + +end RDo.Tactic + +end diff --git a/RandomDo/Tactic/Computable/Example.lean b/RandomDo/Tactic/Computable/Example.lean new file mode 100644 index 0000000..13f7e54 --- /dev/null +++ b/RandomDo/Tactic/Computable/Example.lean @@ -0,0 +1,40 @@ +module +/- /- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import RandomDo.Monad.Instances +public import RandomDo.Monad.Notation +public import RandomDo.NumLean.Distributions +public import RandomDo.Tactic.Computable.Deriving +public import Mathlib +meta import RandomDo.NumLean.Distributions +meta import Batteries.Data.Float.Basic + +/-! +# `@[computable]` on an `rdo` program + +`test` denotes a distribution: a Gaussian draw, shifted by one. The attribute reads it and writes +the program that samples from it, drawing from `NumLean.normal` instead and shifting the draw the +same way. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory NumLean + +/-- Test -/ +@[computable] +noncomputable def test (m : ℝ) : Measure ℝ := rdo + let x ← gaussianReal m 1 + return x + 1 + +/- And it runs: seeded alike, it draws what numpy 2.3.4 draws from `default_rng(42).normal() + 1`, +down to the last bit. -/ +run_cmd do IO.runRandPCG do + let x ← (IO.runRandPCG <| test_computable 10 : IO Float) + Lean.logInfo m!"`test_computable` drew {x.toStringFull}" + -/ diff --git a/RandomDo/Tactic/IsMarkov.lean b/RandomDo/Tactic/IsMarkov/Defs.lean similarity index 100% rename from RandomDo/Tactic/IsMarkov.lean rename to RandomDo/Tactic/IsMarkov/Defs.lean diff --git a/RandomDo/Tactic/Deriving.lean b/RandomDo/Tactic/IsMarkov/Deriving.lean similarity index 98% rename from RandomDo/Tactic/Deriving.lean rename to RandomDo/Tactic/IsMarkov/Deriving.lean index 6129c23..5ebbff4 100644 --- a/RandomDo/Tactic/Deriving.lean +++ b/RandomDo/Tactic/IsMarkov/Deriving.lean @@ -5,7 +5,7 @@ Authors: Rémy Degenne -/ module -public import RandomDo.Tactic.Elab +public import RandomDo.Tactic.IsMarkov.Elab /-! # The `@[is_markov]` attribute diff --git a/RandomDo/Tactic/Elab.lean b/RandomDo/Tactic/IsMarkov/Elab.lean similarity index 99% rename from RandomDo/Tactic/Elab.lean rename to RandomDo/Tactic/IsMarkov/Elab.lean index 4f7c8ce..e7e7244 100644 --- a/RandomDo/Tactic/Elab.lean +++ b/RandomDo/Tactic/IsMarkov/Elab.lean @@ -5,7 +5,7 @@ Authors: Gaëtan Serré -/ module -public import RandomDo.Tactic.Lemmas +public import RandomDo.Tactic.IsMarkov.Lemmas public meta import Lean.Elab.Tactic.Basic /-! diff --git a/RandomDo/Tactic/ForInStep.lean b/RandomDo/Tactic/IsMarkov/ForInStep.lean similarity index 98% rename from RandomDo/Tactic/ForInStep.lean rename to RandomDo/Tactic/IsMarkov/ForInStep.lean index e1cce94..6351f45 100644 --- a/RandomDo/Tactic/ForInStep.lean +++ b/RandomDo/Tactic/IsMarkov/ForInStep.lean @@ -5,7 +5,7 @@ Authors: Gaëtan Serré -/ module -public import RandomDo.Tactic.IsMarkov +public import RandomDo.Tactic.IsMarkov.Defs public import RandomDo.Monad.MeasurableSpace /-! diff --git a/RandomDo/Tactic/Lemmas.lean b/RandomDo/Tactic/IsMarkov/Lemmas.lean similarity index 99% rename from RandomDo/Tactic/Lemmas.lean rename to RandomDo/Tactic/IsMarkov/Lemmas.lean index 1284d31..61c21f0 100644 --- a/RandomDo/Tactic/Lemmas.lean +++ b/RandomDo/Tactic/IsMarkov/Lemmas.lean @@ -8,7 +8,7 @@ module public import RandomDo.Monad.Instances public import RandomDo.Monad.ForInInstances public import RandomDo.Measurable -public import RandomDo.Tactic.ForInStep +public import RandomDo.Tactic.IsMarkov.ForInStep public import Mathlib.MeasureTheory.Measure.ProbabilityMeasure public import Mathlib.Data.List.OfFn public import Mathlib.Probability.Distributions.Gaussian.Real From 069f7444a8da103566dbcb3738973dccdad03993 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 7 Sep 2026 16:46:06 +0200 Subject: [PATCH 12/34] Trace --- RandomDo/Tactic/Computable/Deriving.lean | 81 +++++++++++++++--------- 1 file changed, 50 insertions(+), 31 deletions(-) diff --git a/RandomDo/Tactic/Computable/Deriving.lean b/RandomDo/Tactic/Computable/Deriving.lean index dc44ee4..375dd97 100644 --- a/RandomDo/Tactic/Computable/Deriving.lean +++ b/RandomDo/Tactic/Computable/Deriving.lean @@ -34,6 +34,10 @@ The two constructs `rdo` is made of become the two of `do`: `return e` becomes ` operation on the values the program computes with — is rebuilt from the counterpart `@[computable_as]` records for its head, applied to the translation of its arguments; the instances it asks for are synthesized anew, for the translated types. + +Setting `set_option trace.computable true` prints the tree of pieces the attribute walked through, +each with what it became, the instances synthesized along the way, and the declaration written at +the end. -/ public meta section @@ -42,6 +46,8 @@ open Lean Meta NumLean namespace RDo.Tactic +initialize registerTraceClass `computable + /-- The monad a translated program lives in, `RandPCG IO`: the counterpart of the Giry monad an `rdo` program is written over. -/ def computableMonad : MetaM Expr := mkAppM ``RandPCG #[mkConst ``IO] @@ -49,37 +55,49 @@ def computableMonad : MetaM Expr := mkAppM ``RandPCG #[mkConst ``IO] /-- Translate `e`, a piece of an `rdo` program, into its computable counterpart. `σ` sends the variables the program binds to the ones the translated program binds in their place, which the change of types makes necessary. -/ -partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := do - -- `MeasurableSpacePure.mPure` takes five arguments, the last of which is the value returned. - if e.isAppOfArity ``MeasurableSpacePure.mPure 5 then - mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ (e.getArg! 4)] - -- `MeasurableSpaceBind.mBind` takes eight, the last two being the program and the continuation. - else if e.isAppOfArity ``MeasurableSpaceBind.mBind 8 then - mkAppOptM ``Bind.bind #[← computableMonad, none, none, none, - ← translate σ (e.getArg! 6), ← translate σ (e.getArg! 7)] - else match e with - | .fvar fvarId => return σ.get fvarId - | .sort .. | .lit .. => return e - | .mdata _ b => translate σ b - | .lam .. => - lambdaBoundedTelescope e 1 fun xs body ↦ do - let x := xs[0]! - let t ← translate σ (← x.fvarId!.getType) - withLocalDeclD (← x.fvarId!.getUserName) t fun y ↦ do - mkLambdaFVars #[y] (← translate (σ.insert x.fvarId! y) body) - | _ => - let .const declName _ := e.getAppFn - | throwError "`computable`: cannot translate{indentExpr e}" - let counterpart := (← computableAs? declName).getD declName - let mut f ← mkConstWithFreshMVarLevels counterpart - for arg in e.getAppArgs do - let .forallE _ t _ bi ← whnf (← inferType f) - | throwError "`computable`: {f} does not take the argument{indentExpr arg}" - let arg ← if bi.isInstImplicit then synthInstance t else translate σ arg - unless ← isDefEq (← inferType arg) t do - throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}" - f := mkApp f arg - return f +partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := + withTraceNode `computable + (fun + | .ok e' => return m!"{e} ↦ {e'}" + | .error _ => return m!"{e}: not translated") do + -- `MeasurableSpacePure.mPure` takes five arguments, the last of which is the value returned. + if e.isAppOfArity ``MeasurableSpacePure.mPure 5 then + mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ (e.getArg! 4)] + -- `MeasurableSpaceBind.mBind` takes eight, the last two the program and the continuation. + else if e.isAppOfArity ``MeasurableSpaceBind.mBind 8 then + mkAppOptM ``Bind.bind #[← computableMonad, none, none, none, + ← translate σ (e.getArg! 6), ← translate σ (e.getArg! 7)] + else match e with + | .fvar fvarId => return σ.get fvarId + | .sort .. | .lit .. => return e + | .mdata _ b => translate σ b + | .lam .. => + lambdaBoundedTelescope e 1 fun xs body ↦ do + let x := xs[0]! + let t ← translate σ (← x.fvarId!.getType) + withLocalDeclD (← x.fvarId!.getUserName) t fun y ↦ do + mkLambdaFVars #[y] (← translate (σ.insert x.fvarId! y) body) + | _ => + let .const declName _ := e.getAppFn + | throwError "`computable`: cannot translate{indentExpr e}" + let counterpart := (← computableAs? declName).getD declName + let mut f ← mkConstWithFreshMVarLevels counterpart + for arg in e.getAppArgs do + let .forallE _ t _ bi ← whnf (← inferType f) + | throwError "`computable`: {f} does not take the argument{indentExpr arg}" + /- An instance is the one argument that is not translated but synthesized anew, so the + trace above it says nothing: it is reported here instead. -/ + let arg ← + if bi.isInstImplicit then do + let inst ← synthInstance t + trace[computable] "instance: {inst}" + pure inst + else + translate σ arg + unless ← isDefEq (← inferType arg) t do + throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}" + f := mkApp f arg + return f /-- Translate the `rdo` program `declName` and add the translation to the environment, under the name `declName` followed by `_computable`. -/ @@ -92,6 +110,7 @@ def addComputableDecl (declName : Name) : MetaM Unit := do let translated := declName.appendAfter "_computable" addAndCompile <| .defnDecl <| ← mkDefinitionValInferringUnsafe translated info.levelParams type value (.regular (getMaxHeight (← getEnv) value + 1)) + trace[computable] "wrote {translated} :{indentExpr type}" addDocStringCore translated s!"The program that samples from `{declName}`, written by the \ `@[computable]` attribute." /- The name is one the attribute picks and not one the user wrote, so the underscore in it is From f66de1db9b5a4047ceda782b9193f03bbe907d6c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 7 Sep 2026 16:50:05 +0200 Subject: [PATCH 13/34] Add `NumLean.normal'` --- RandomDo/NumLean/Distributions.lean | 6 ++++++ RandomDo/Tactic/Computable/Counterparts.lean | 2 +- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/RandomDo/NumLean/Distributions.lean b/RandomDo/NumLean/Distributions.lean index 73640d4..7d2fe19 100644 --- a/RandomDo/NumLean/Distributions.lean +++ b/RandomDo/NumLean/Distributions.lean @@ -130,6 +130,12 @@ tail is the exponential's own, memoryless, so one draw beyond `Ziggurat.expR` su if scale < 0 then throw <| IO.userError "scale < 0" return Float.fma scale (← standardNormal) loc +/-- Draw random samples from a normal (Gaussian) distribution, with variance instead of standard +deviation. -/ +@[inline] def normal' (loc : Float := 0) (var : Float := 1) : RandPCG IO Float := do + if var < 0 then throw <| IO.userError "var < 0" + return Float.fma (Float.sqrt var) (← standardNormal) loc + /-- Draw samples from an exponential distribution. -/ @[inline] def exponential (scale : Float := 1) : RandPCG IO Float := do if scale < 0 then throw <| IO.userError "scale < 0" diff --git a/RandomDo/Tactic/Computable/Counterparts.lean b/RandomDo/Tactic/Computable/Counterparts.lean index 75384fa..b5a0abb 100644 --- a/RandomDo/Tactic/Computable/Counterparts.lean +++ b/RandomDo/Tactic/Computable/Counterparts.lean @@ -29,6 +29,6 @@ attribute [computable_as Float] NNReal /- `gaussianReal` reads its second argument as a variance and `normal` reads it as a standard deviation: the two agree on the `1` the example draws with, not in general. -/ -attribute [computable_as NumLean.normal] ProbabilityTheory.gaussianReal +attribute [computable_as NumLean.normal'] ProbabilityTheory.gaussianReal end From f52cc70a41c9e12c5bb18aedda0360fa1a7602aa Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 7 Sep 2026 16:59:29 +0200 Subject: [PATCH 14/34] Add support for `sqrt`, `log` and `expr` --- RandomDo/Tactic/Computable/Counterparts.lean | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/RandomDo/Tactic/Computable/Counterparts.lean b/RandomDo/Tactic/Computable/Counterparts.lean index b5a0abb..00549a1 100644 --- a/RandomDo/Tactic/Computable/Counterparts.lean +++ b/RandomDo/Tactic/Computable/Counterparts.lean @@ -27,8 +27,12 @@ attribute [computable_as Float] NNReal /-! ## Distributions -/ -/- `gaussianReal` reads its second argument as a variance and `normal` reads it as a standard -deviation: the two agree on the `1` the example draws with, not in general. -/ attribute [computable_as NumLean.normal'] ProbabilityTheory.gaussianReal +attribute [computable_as Float.sqrt] Real.sqrt + +attribute [computable_as Float.log] Real.log + +attribute [computable_as Float.exp] Real.exp + end From 8a7666e449722d7b862c73cc74ae2b076d0ec3f2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Mon, 7 Sep 2026 17:10:21 +0200 Subject: [PATCH 15/34] Refactor --- RandomDo/NumLean/Distributions.lean | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/RandomDo/NumLean/Distributions.lean b/RandomDo/NumLean/Distributions.lean index 7d2fe19..81b5f23 100644 --- a/RandomDo/NumLean/Distributions.lean +++ b/RandomDo/NumLean/Distributions.lean @@ -134,7 +134,7 @@ tail is the exponential's own, memoryless, so one draw beyond `Ziggurat.expR` su deviation. -/ @[inline] def normal' (loc : Float := 0) (var : Float := 1) : RandPCG IO Float := do if var < 0 then throw <| IO.userError "var < 0" - return Float.fma (Float.sqrt var) (← standardNormal) loc + normal loc (Float.sqrt var) /-- Draw samples from an exponential distribution. -/ @[inline] def exponential (scale : Float := 1) : RandPCG IO Float := do From c9192fd0b6e6c4456ce6335e1ec2c2c893e4327c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Tue, 8 Sep 2026 15:30:04 +0200 Subject: [PATCH 16/34] Better support --- RandomDo/Tactic/Computable/Counterparts.lean | 2 + RandomDo/Tactic/Computable/Defs.lean | 11 +- RandomDo/Tactic/Computable/Deriving.lean | 167 +++++++++++-------- RandomDo/Tactic/IsMarkov/Deriving.lean | 2 +- Test.lean | 1 + Test/Computable.lean | 45 +++++ 6 files changed, 146 insertions(+), 82 deletions(-) create mode 100644 Test/Computable.lean diff --git a/RandomDo/Tactic/Computable/Counterparts.lean b/RandomDo/Tactic/Computable/Counterparts.lean index 00549a1..7814874 100644 --- a/RandomDo/Tactic/Computable/Counterparts.lean +++ b/RandomDo/Tactic/Computable/Counterparts.lean @@ -29,6 +29,8 @@ attribute [computable_as Float] NNReal attribute [computable_as NumLean.normal'] ProbabilityTheory.gaussianReal +/-! ## Classical functions -/ + attribute [computable_as Float.sqrt] Real.sqrt attribute [computable_as Float.log] Real.log diff --git a/RandomDo/Tactic/Computable/Defs.lean b/RandomDo/Tactic/Computable/Defs.lean index ffbc027..66b5c81 100644 --- a/RandomDo/Tactic/Computable/Defs.lean +++ b/RandomDo/Tactic/Computable/Defs.lean @@ -11,16 +11,11 @@ public meta import Lean.ReservedNameAction /-! # The `@[computable_as]` attribute -An `rdo` program is written over the Giry monad: it draws from measures on `ℝ`, which no machine -samples. Turning it into a program that runs, which the `@[computable]` attribute of -`RandomDo.Tactic.Computable.Deriving` does, asks for a counterpart of each piece the program is -built from: `Float` for `ℝ`, `NumLean.normal` for `gaussianReal`. This file holds the attribute +An `rdo` program is written over the Giry monad: it draws from measures on a measurable space, which no machine samples. Turning it into a program that runs, which the `@[computable]` attribute of `RandomDo.Tactic.Computable.Deriving` does, asks for a counterpart of each piece the program is +built from, e.g, `NumLean.normal'` for `gaussianReal`. This file holds the attribute recording them; `RandomDo.Tactic.Computable.Counterparts` holds the counterparts themselves. `@[computable_as f]` on a declaration `d` reads: `f` is what `d` becomes in a translated program. -Only the pieces denoting something the program computes with need one. The scaffolding around -them — numerals, arithmetic — is polymorphic, and the translation keeps it as it is, at the -translated types. -/ public meta section @@ -34,8 +29,6 @@ initialize computableAsExt : NameMapExtension Name ← registerNameMapAttribute { name := `computable_as descr := "record the computable counterpart of this declaration" - /- `@[computable_as f]` is read by `Lean.Parser.Attr.simple`, the parser an attribute that - declares no syntax of its own gets: `f` is the single child of `stx[1]`. -/ add := fun _ stx ↦ do let f := stx[1][0] unless f.isIdent do diff --git a/RandomDo/Tactic/Computable/Deriving.lean b/RandomDo/Tactic/Computable/Deriving.lean index 375dd97..41b706f 100644 --- a/RandomDo/Tactic/Computable/Deriving.lean +++ b/RandomDo/Tactic/Computable/Deriving.lean @@ -8,15 +8,11 @@ module public import RandomDo.Tactic.Computable.Counterparts public import RandomDo.Monad.MeasurableSpace public meta import Lean.Elab.Tactic.Basic -public meta import Batteries.Tactic.Lint.Basic /-! # The `@[computable]` attribute -Writing an `rdo` program and writing the program that samples from it are two separate steps, and -the second is mechanical: every construct of the first has a counterpart in the second. This -attribute walks the program and writes that counterpart down, so a program carries the sampler it -denotes: +`@[computable]` reads an `rdo` program and adds the program that samples from it: ``` @[computable] @@ -25,19 +21,19 @@ noncomputable def shifted : Measure ℝ := rdo return x + 1 ``` -adds `shifted_computable : RandPCG IO Float`, which draws from `NumLean.normal 0 1` and adds one. +adds `shiftedComputable : RandPCG IO Float`, which draws from `NumLean.normal' 0 1` and adds one. -## How the program is read +The Giry monad and its two operations become `RandPCG IO`, `pure` and `bind`. Anything else is +rebuilt from the counterpart `@[computable_as]` records for its head, with its arguments translated +in turn and its instances synthesized anew. A term translates into a term of the translation of its +type; where the rebuilt one does not, its head is a definition nothing is known about, and its body +is read in its place. `@[computable]` records the program it writes, so a program drawing from +another translates into one calling that other's translation. -The two constructs `rdo` is made of become the two of `do`: `return e` becomes `pure e`, and -`let x ← p; q` becomes `p >>= q`, at `RandPCG IO`. Everything else — a distribution, a numeral, an -operation on the values the program computes with — is rebuilt from the counterpart -`@[computable_as]` records for its head, applied to the translation of its arguments; the -instances it asks for are synthesized anew, for the translated types. +Two things extend the attribute: an `@[computable_as]` entry, and an alternative of `translate` for +a construct of `rdo` it has not been taught. -Setting `set_option trace.computable true` prints the tree of pieces the attribute walked through, -each with what it became, the instances synthesized along the way, and the declaration written at -the end. +`set_option trace.computable true` prints what each piece became. -/ public meta section @@ -48,77 +44,104 @@ namespace RDo.Tactic initialize registerTraceClass `computable -/-- The monad a translated program lives in, `RandPCG IO`: the counterpart of the Giry monad an -`rdo` program is written over. -/ +/-- The monad the translated programs live in. -/ def computableMonad : MetaM Expr := mkAppM ``RandPCG #[mkConst ``IO] -/-- Translate `e`, a piece of an `rdo` program, into its computable counterpart. `σ` sends the -variables the program binds to the ones the translated program binds in their place, which the -change of types makes necessary. -/ +mutual + +/-- Translate a piece of an `rdo` program. `σ` maps the variables the program binds to the ones the +translated program binds in their place. -/ partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := - withTraceNode `computable - (fun - | .ok e' => return m!"{e} ↦ {e'}" - | .error _ => return m!"{e}: not translated") do - -- `MeasurableSpacePure.mPure` takes five arguments, the last of which is the value returned. - if e.isAppOfArity ``MeasurableSpacePure.mPure 5 then - mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ (e.getArg! 4)] - -- `MeasurableSpaceBind.mBind` takes eight, the last two the program and the continuation. - else if e.isAppOfArity ``MeasurableSpaceBind.mBind 8 then + withTraceNode `computable (fun + | .ok e' => return m!"{e} ↦ {e'}" + | .error _ => return m!"{e}: not translated") do + match_expr e with + | MeasurableSpacePure.mPure _ _ _ _ a => + mkAppOptM ``Pure.pure #[← computableMonad, none, none, ← translate σ a] + | MeasurableSpaceBind.mBind _ _ _ _ _ _ p k => mkAppOptM ``Bind.bind #[← computableMonad, none, none, none, - ← translate σ (e.getArg! 6), ← translate σ (e.getArg! 7)] - else match e with - | .fvar fvarId => return σ.get fvarId - | .sort .. | .lit .. => return e - | .mdata _ b => translate σ b - | .lam .. => - lambdaBoundedTelescope e 1 fun xs body ↦ do - let x := xs[0]! - let t ← translate σ (← x.fvarId!.getType) - withLocalDeclD (← x.fvarId!.getUserName) t fun y ↦ do - mkLambdaFVars #[y] (← translate (σ.insert x.fvarId! y) body) - | _ => - let .const declName _ := e.getAppFn - | throwError "`computable`: cannot translate{indentExpr e}" - let counterpart := (← computableAs? declName).getD declName - let mut f ← mkConstWithFreshMVarLevels counterpart - for arg in e.getAppArgs do - let .forallE _ t _ bi ← whnf (← inferType f) - | throwError "`computable`: {f} does not take the argument{indentExpr arg}" - /- An instance is the one argument that is not translated but synthesized anew, so the - trace above it says nothing: it is reported here instead. -/ - let arg ← - if bi.isInstImplicit then do - let inst ← synthInstance t - trace[computable] "instance: {inst}" - pure inst - else - translate σ arg - unless ← isDefEq (← inferType arg) t do - throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}" - f := mkApp f arg - return f + ← translate σ p, ← translate σ k] + | MeasureTheory.Measure α _ => return mkApp (← computableMonad) (← translate σ α) + | _ => match e with + | .fvar x => return σ.get x + | .sort .. | .lit .. => return e + | .mdata _ b => translate σ b + | .lam .. => lambdaBoundedTelescope e 1 fun xs body ↦ do + let x := xs[0]!.fvarId! + withLocalDeclD (← x.getUserName) (← translate σ (← x.getType)) fun y ↦ do + mkLambdaFVars #[y] (← translate (σ.insert x y) body) + | _ => translateApp σ e + +/-- Rebuild an application from the counterpart of its head; where nothing known about that head +fits, look through it and read its body in its place. -/ +partial def translateApp (σ : FVarSubst) (e : Expr) : MetaM Expr := do + if e.getAppFn.isLambda then return ← translate σ e.headBeta + let .const declName _ := e.getAppFn | throwError "`computable`: cannot translate{indentExpr e}" + let counterpart? ← computableAs? declName + let head := counterpart?.getD declName + /- We save the state of metavariables, so that a failed rebuild does not leave them in a + half-built state. -/ + let s ← saveState + try + /- A term translates into a term of the translation of its type. A head nothing is known about + rebuilds into itself, of the type it had, and might fail. We check that the rebuilt term has + the translation of the type, and if not, we look through it. -/ + let f ← rebuild σ head e + -- Useless for a type, whose type is a sort either way. + if (← inferType f).isSort then return f + let expected ← translate σ (← inferType e) + unless ← isDefEq (← inferType f) expected do + throwError "`computable`: {f} is of type{indentExpr (← inferType f)}\n\ + where the translation asks for{indentExpr expected}" + return f + catch ex => + restoreState s + /- The rebuild failed, we try to look through the head, and read its body in its place. -/ + if counterpart?.isNone then + if let some e' ← unfoldDefinition? e then + trace[computable] "nothing known about {declName}, looking through it" + return ← translate σ e' + throw ex + +/-- Apply `head` to the arguments of `e`, translated in turn, and check that the result has the +translation of the type of `e`. -/ +partial def rebuild (σ : FVarSubst) (head : Name) (e : Expr) : MetaM Expr := do + let mut f ← mkConstWithFreshMVarLevels head + for arg in e.getAppArgs do + let .forallE _ t _ bi ← whnf (← inferType f) + | throwError "`computable`: {f} does not take the argument{indentExpr arg}" + let arg ← + if bi.isInstImplicit then do + -- An instance is not translated: it is asked for anew, at the translated types. + let inst ← synthInstance t + trace[computable] "instance: {inst}" + pure inst + else + translate σ arg + unless ← isDefEq (← inferType arg) t do + throwError "`computable`: {arg} does not fit the argument of {f}, of type{indentExpr t}" + f := mkApp f arg + return f + +end /-- Translate the `rdo` program `declName` and add the translation to the environment, under the -name `declName` followed by `_computable`. -/ +name `declName` followed by `Computable`. -/ def addComputableDecl (declName : Name) : MetaM Unit := do let info ← getConstInfo declName let some value := info.value? | throwError "`computable` can only be derived for a definition, but {declName} has no value" let value ← instantiateMVars (← translate {} value) let type ← instantiateMVars (← inferType value) - let translated := declName.appendAfter "_computable" + let translated := declName.appendAfter "Computable" addAndCompile <| .defnDecl <| ← mkDefinitionValInferringUnsafe translated info.levelParams type value (.regular (getMaxHeight (← getEnv) value + 1)) - trace[computable] "wrote {translated} :{indentExpr type}" - addDocStringCore translated s!"The program that samples from `{declName}`, written by the \ - `@[computable]` attribute." - /- The name is one the attribute picks and not one the user wrote, so the underscore in it is - reported for every program translated unless it is exempted here. -/ - setEnv (← ofExcept (Batteries.Tactic.Lint.nolintAttr.setParam (← getEnv) translated - #[`defsWithUnderscore])) - -/-- The `@[computable]` attribute. -/ + addDocStringCore translated s!"The computable program that samples from `{declName}` \ + (automatically generated by the `@[computable]` attribute)." + computableAsExt.add declName translated + trace[computable] "wrote {translated}:{indentExpr type}" + +@[inherit_doc addComputableDecl] initialize registerBuiltinAttribute { name := `computable descr := "translate this `rdo` program into the program that samples from it" diff --git a/RandomDo/Tactic/IsMarkov/Deriving.lean b/RandomDo/Tactic/IsMarkov/Deriving.lean index 5ebbff4..2c90e5e 100644 --- a/RandomDo/Tactic/IsMarkov/Deriving.lean +++ b/RandomDo/Tactic/IsMarkov/Deriving.lean @@ -63,7 +63,7 @@ def addIsMarkovInstance (declName : Name) : TermElabM Unit := do value := ← instantiateMVars proof }) Meta.addInstance instName .global 1000 -/-- The `@[is_markov]` attribute. -/ +@[inherit_doc isMarkovStatement] initialize registerBuiltinAttribute { name := `is_markov descr := "prove that this `rdo` program is a Markov kernel, and register it as an instance" diff --git a/Test.lean b/Test.lean index 7862fd2..4009bb6 100644 --- a/Test.lean +++ b/Test.lean @@ -2,6 +2,7 @@ module -- shake: keep-all --deprecated_module: ignore public import Test.Bind public import Test.Common +public import Test.Computable public import Test.Control public import Test.Gaps public import Test.Instances diff --git a/Test/Computable.lean b/Test/Computable.lean new file mode 100644 index 0000000..ea31b4d --- /dev/null +++ b/Test/Computable.lean @@ -0,0 +1,45 @@ +module + +public import Test.IsMarkov +import Batteries.Data.Float.Basic +/- A `run_cmd` runs at elaboration time, so what it calls has to be imported as `meta` too: the +sampler it draws with, and `Float.toStringFull` it prints with. -/ +meta import RandomDo.NumLean.Distributions +meta import Batteries.Data.Float.Basic + +set_option linter.style.header false + +set_option trace.computable true + +namespace Test.Computable + +open Test.IsMarkov NumLean Lean.Elab.Command + +def logComputable (prog : RandPCG IO Float) : CommandElabM Unit := do + let x ← (IO.runRandPCG prog : IO Float) + let y ← (IO.runRandPCGWith 42 prog : IO Float) + Lean.logInfo m!"x = {x.toStringFull}" + Lean.logInfo m!"y (seed 42) = {y.toStringFull}" + +--attribute [computable] sumTwo + +--run_cmd do logComputable sumTwoComputable + +@[computable] +noncomputable +def test : MeasureTheory.Measure ℝ := rdo + let y ← sumTwo + let x ← ProbabilityTheory.gaussianReal 0 1 + return x + y + +attribute [computable] centred + +run_cmd do logComputable (centredComputable 20) + +attribute [computable] branchOn + +run_cmd do logComputable (branchOnComputable 20) + +run_cmd do logComputable (branchOnComputable (-1)) + +end Test.Computable From c6b596f24647b25b070449e93b64114958abf086 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Tue, 8 Sep 2026 16:57:57 +0200 Subject: [PATCH 17/34] Binomial distribution --- RandomDo.lean | 1 + RandomDo/NumLean/Binomial.lean | 194 +++++++++++++++++++ RandomDo/NumLean/Distributions.lean | 19 ++ RandomDo/Tactic/Computable/Counterparts.lean | 11 +- RandomDo/Tactic/Computable/Deriving.lean | 4 + Test/Computable.lean | 29 +-- scripts/check_binomial.py | 27 +++ 7 files changed, 270 insertions(+), 15 deletions(-) create mode 100644 RandomDo/NumLean/Binomial.lean create mode 100644 scripts/check_binomial.py diff --git a/RandomDo.lean b/RandomDo.lean index ee6fc15..cbabc80 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -6,6 +6,7 @@ public import RandomDo.Monad.ForInInstances public import RandomDo.Monad.Instances public import RandomDo.Monad.MeasurableSpace public import RandomDo.Monad.Notation +public import RandomDo.NumLean.Binomial public import RandomDo.NumLean.Distributions public import RandomDo.NumLean.PCG64 public import RandomDo.NumLean.SeedSequence diff --git a/RandomDo/NumLean/Binomial.lean b/RandomDo/NumLean/Binomial.lean new file mode 100644 index 0000000..0f1d7cc --- /dev/null +++ b/RandomDo/NumLean/Binomial.lean @@ -0,0 +1,194 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import RandomDo.NumLean.PCG64 + +/-! +# The binomial distribution + +`binomial n p` draws exactly what numpy's `Generator.binomial` draws: the inversion of the +cumulative distribution where the mean `n * p` is at most `30`, the BTPE algorithm of +Kachitvichyanukul and Schmeiser beyond, and the mirror image of either when `p > 1 / 2`. + +`binomial`, which chooses between the two and mirrors them, is in +`RandomDo.NumLean.Distributions`, with the other distributions. + +## Main definitions + +* `Inversion`, `inversionSetup`, `inversionDraw`: numpy's `random_binomial_inversion`. +* `Btpe`, `btpeSetup`, `Btpe.accept`, `btpeDraw`: numpy's `random_binomial_btpe`. + +## References + +* V. Kachitvichyanukul and B. W. Schmeiser, *Binomial random variate generation*, Communications + of the ACM 31 (1988), 216-222. +* numpy's `numpy/random/src/distributions/distributions.c`. +-/ + +@[expose] public section + +namespace NumLean + +/-! ## Inversion -/ + +/-- The constants the inversion reads a draw against. -/ +structure Inversion where + /-- The number of trials. -/ + n : Float + /-- The probability of a success, at most one half. -/ + p : Float + /-- The probability of a failure, `1 - p`. -/ + q : Float + /-- `q ^ n`, the probability that no trial succeeds. -/ + qn : Float + /-- The number of successes past which the walk gives up and starts over. -/ + bound : Float + +/-- The constants of `random_binomial_inversion`, for `p ≤ 0.5` and `n * p ≤ 30`. -/ +@[inline] def inversionSetup (n p : Float) : Inversion := + let q := 1.0 - p + let np := n * p + let b := np + 10.0 * Float.sqrt (np * q + 1) + { n, p, q, qn := Float.exp (n * Float.log q), bound := if n < b then n else b } + +/-- Walk up the cumulative distribution from zero until it passes `u`, as the loop of +`random_binomial_inversion`. Answers `-1` where the walk runs past `bound`, which the tail beyond +it is too thin to reach and where numpy starts the draw over. -/ +partial def inversionWalk (s : Inversion) (x px u : Float) : Float := + if u > px then + let x := x + 1 + if x > s.bound then -1 + else inversionWalk s x (((s.n - x + 1) * s.p * px) / (x * s.q)) (u - px) + else x + +/-- Sample by inverting the cumulative distribution, as numpy's `random_binomial_inversion`. -/ +partial def inversionDraw (s : Inversion) : RandPCG IO Float := do + let x := inversionWalk s 0 s.qn (← random) + if x < 0 then inversionDraw s else return x + +/-! ## BTPE -/ + +/-- The constants BTPE reads a draw against. -/ +structure Btpe where + /-- The number of trials. -/ + n : Float + /-- The probability of a success, at most one half. -/ + r : Float + /-- The probability of a failure, `1 - r`. -/ + q : Float + /-- The mode of the distribution. -/ + m : Float + /-- The middle of the triangle, `m + 1 / 2`. -/ + xm : Float + /-- The half-width of the triangle. -/ + p1 : Float + /-- The left end of the parallelogram. -/ + xl : Float + /-- The right end of the parallelogram. -/ + xr : Float + /-- The height of the parallelogram, relative to the triangle. -/ + c : Float + /-- The rate of the left exponential tail. -/ + laml : Float + /-- The rate of the right exponential tail. -/ + lamr : Float + /-- The area of the triangle and the parallelogram. -/ + p2 : Float + /-- The area of the triangle, the parallelogram and the left tail. -/ + p3 : Float + /-- The area of all four regions, which a draw is scaled by. -/ + p4 : Float + /-- The variance `n * r * q`. -/ + nrq : Float + +/-- The constants of `random_binomial_btpe`, for `p ≤ 0.5` and `n * p > 30`. -/ +@[inline] def btpeSetup (n p : Float) : Btpe := + let r := if p < 1.0 - p then p else 1.0 - p + let q := 1.0 - r + let fm := n * r + r + let m := Float.floor fm + let p1 := Float.floor (2.195 * Float.sqrt (n * r * q) - 4.6 * q) + 0.5 + let xm := m + 0.5 + let xl := xm - p1 + let xr := xm + p1 + let c := 0.134 + 20.5 / (15.3 + m) + let al := (fm - xl) / (fm - xl * r) + let ar := (xr - fm) / (xr * q) + let laml := al * (1.0 + al / 2.0) + let lamr := ar * (1.0 + ar / 2.0) + let p2 := p1 * (1.0 + 2.0 * c) + let p3 := p2 + c / laml + { n, r, q, m, xm, p1, xl, xr, c, laml, lamr, p2, p3, p4 := p3 + c / lamr, nrq := n * r * q } + +/-- One term of the Stirling series bounding `log` of a factorial, as the last test of BTPE spells +it out. `u2` is `u * u`. -/ +@[inline] def btpeStirling (u u2 : Float) : Float := + (13680.0 - (462.0 - (132.0 - (99.0 - 140.0 / u2) / u2) / u2) / u2) / u / 166320.0 + +/-- The ratios of the probabilities from the mode up to `y`, multiplied into `f` one at a time as +the step 50 of `random_binomial_btpe` takes them. -/ +partial def btpeUp (a s f i y : Float) : Float := + if i ≤ y then btpeUp a s (f * (a / i - s)) (i + 1) y else f + +/-- The ratios from `y` up to the mode, divided out of `f` one at a time. Dividing the running +value and dividing by the product do not round alike, and BTPE reads the first. -/ +partial def btpeDown (a s f i m : Float) : Float := + if i ≤ m then btpeDown a s (f / (a / i - s)) (i + 1) m else f + +/-- Whether BTPE accepts the candidate `y` drawn with `v`, as the steps 50 and 52 of +`random_binomial_btpe`: by the explicit product of the ratios of the probabilities between the mode +and `y` when the two are close, and by a squeeze then the Stirling bound otherwise. -/ +def Btpe.accept (b : Btpe) (y v : Float) : Bool := Id.run do + let k := Float.abs (y - b.m) + unless k > 20 && k < b.nrq / 2.0 - 1 do + let s := b.r / b.q + let a := s * (b.n + 1) + if b.m < y then return !(v > btpeUp a s 1.0 (b.m + 1) y) + if b.m > y then return !(v > btpeDown a s 1.0 (y + 1) b.m) + return !(v > 1.0) + let rho := (k / b.nrq) * ((k * (k / 3.0 + 0.625) + 0.16666666666666666) / b.nrq + 0.5) + let t := -k * k / (2 * b.nrq) + let a := Float.log v + if a < t - rho then return true + if a > t + rho then return false + let x1 := y + 1 + let f1 := b.m + 1 + let z := b.n + 1 - b.m + let w := b.n - y + 1 + return !(a > b.xm * Float.log (f1 / x1) + (b.n - b.m + 0.5) * Float.log (z / w) + + (y - b.m) * Float.log (w * b.r / (x1 * b.q)) + + btpeStirling f1 (f1 * f1) + btpeStirling z (z * z) + + btpeStirling x1 (x1 * x1) + btpeStirling w (w * w)) + +/-- Draw a candidate from the triangle, the parallelogram or one of the two exponential tails, and +start over until one is accepted, as the steps 10 to 60 of `random_binomial_btpe`. -/ +partial def btpeDraw (b : Btpe) : RandPCG IO Float := do + let u := (← random) * b.p4 + let v ← random + if u ≤ b.p1 then + return Float.floor (b.xm - b.p1 * v + u) + else if u ≤ b.p2 then + let x := b.xl + (u - b.p1) / b.c + let v := v * b.c + 1.0 - Float.abs (b.m - x + 0.5) / b.p1 + if v > 1.0 then btpeDraw b else + let y := Float.floor x + if b.accept y v then return y else btpeDraw b + else if u ≤ b.p3 then + let y := Float.floor (b.xl + Float.log v / b.laml) + -- `v` can be zero, and the floor of the resulting infinity is no candidate. + if y < 0 || v == 0.0 then btpeDraw b else + let v := v * (u - b.p2) * b.laml + if b.accept y v then return y else btpeDraw b + else + let y := Float.floor (b.xr - Float.log v / b.lamr) + if y > b.n || v == 0.0 then btpeDraw b else + let v := v * (u - b.p3) * b.lamr + if b.accept y v then return y else btpeDraw b + +end NumLean + +end diff --git a/RandomDo/NumLean/Distributions.lean b/RandomDo/NumLean/Distributions.lean index 81b5f23..9ac50df 100644 --- a/RandomDo/NumLean/Distributions.lean +++ b/RandomDo/NumLean/Distributions.lean @@ -8,6 +8,7 @@ module public import RandomDo.NumLean.PCG64 public meta import RandomDo.NumLean.PCG64 public import FFI.Float +public import RandomDo.NumLean.Binomial public import RandomDo.NumLean.Ziggurat public import RandomDo.NumLean.ZigguratSampler @@ -141,4 +142,22 @@ deviation. -/ if scale < 0 then throw <| IO.userError "scale < 0" return scale * (← standardExponential) +/-- Draw samples from a binomial distribution. -/ +def binomial (n : Nat) (p : Float) : RandPCG IO Nat := do + -- The comparisons are the `Bool` ones: through `Decidable`, each costs more than a draw. + if p.lt 0.0 || Float.lt 1.0 p || p.isNaN then + throw <| IO.userError "p < 0, p > 1 or p is NaN" + let n := n.toUInt64.toFloat + if n == 0 || p == 0.0 then return 0 + if Float.le p 0.5 then + if Float.le (p * n) 30.0 then return (← inversionDraw (inversionSetup n p)).toUInt64.toNat + else return (← btpeDraw (btpeSetup n p)).toUInt64.toNat + else + let q := 1.0 - p + if Float.le (q * n) 30.0 then return (n - (← inversionDraw (inversionSetup n q))).toUInt64.toNat + else return (n - (← btpeDraw (btpeSetup n q))).toUInt64.toNat + +/-- Draw samples from a Bernoulli distribution. -/ +@[inline] def bernoulli (p : Float) := binomial 1 p + end NumLean diff --git a/RandomDo/Tactic/Computable/Counterparts.lean b/RandomDo/Tactic/Computable/Counterparts.lean index 7814874..7243973 100644 --- a/RandomDo/Tactic/Computable/Counterparts.lean +++ b/RandomDo/Tactic/Computable/Counterparts.lean @@ -8,6 +8,7 @@ module public import RandomDo.Tactic.Computable.Defs public import RandomDo.NumLean.Distributions public import Mathlib.Probability.Distributions.Gaussian.Real +public import Mathlib.Probability.Distributions.Bernoulli /-! # Computable counterparts of the pieces an `rdo` program is made of @@ -18,7 +19,7 @@ counterpart recorded here through `@[computable_as]`. There is one entry per pie one distribution they draw from. -/ -public meta section +@[expose] public section /-! ## Types -/ @@ -29,6 +30,12 @@ attribute [computable_as Float] NNReal attribute [computable_as NumLean.normal'] ProbabilityTheory.gaussianReal +def bernoulliChoice (α : Type) [MeasurableSpace α] (x y : α) (p : Float) : + NumLean.RandPCG IO α := do + return if (← NumLean.bernoulli p) == 1 then x else y + +attribute [computable_as bernoulliChoice] ProbabilityTheory.bernoulliMeasure + /-! ## Classical functions -/ attribute [computable_as Float.sqrt] Real.sqrt @@ -36,5 +43,3 @@ attribute [computable_as Float.sqrt] Real.sqrt attribute [computable_as Float.log] Real.log attribute [computable_as Float.exp] Real.exp - -end diff --git a/RandomDo/Tactic/Computable/Deriving.lean b/RandomDo/Tactic/Computable/Deriving.lean index 41b706f..e9353ca 100644 --- a/RandomDo/Tactic/Computable/Deriving.lean +++ b/RandomDo/Tactic/Computable/Deriving.lean @@ -62,6 +62,10 @@ partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := mkAppOptM ``Bind.bind #[← computableMonad, none, none, none, ← translate σ p, ← translate σ k] | MeasureTheory.Measure α _ => return mkApp (← computableMonad) (← translate σ α) + /- A subtype is its carrier, and one of its values is the value it carries: the constraint and + the proof of it are what a computable counterpart does not have. -/ + | Subtype α _ => translate σ α + | Subtype.mk _ _ v _ => translate σ v | _ => match e with | .fvar x => return σ.get x | .sort .. | .lit .. => return e diff --git a/Test/Computable.lean b/Test/Computable.lean index ea31b4d..3ef04f0 100644 --- a/Test/Computable.lean +++ b/Test/Computable.lean @@ -1,11 +1,8 @@ module public import Test.IsMarkov +public import Test.Bind import Batteries.Data.Float.Basic -/- A `run_cmd` runs at elaboration time, so what it calls has to be imported as `meta` too: the -sampler it draws with, and `Float.toStringFull` it prints with. -/ -meta import RandomDo.NumLean.Distributions -meta import Batteries.Data.Float.Basic set_option linter.style.header false @@ -15,19 +12,19 @@ namespace Test.Computable open Test.IsMarkov NumLean Lean.Elab.Command -def logComputable (prog : RandPCG IO Float) : CommandElabM Unit := do - let x ← (IO.runRandPCG prog : IO Float) - let y ← (IO.runRandPCGWith 42 prog : IO Float) - Lean.logInfo m!"x = {x.toStringFull}" - Lean.logInfo m!"y (seed 42) = {y.toStringFull}" +def logComputable {α : Type} [Lean.ToMessageData α] (prog : RandPCG IO α) : CommandElabM Unit := do + let x ← (IO.runRandPCG prog : IO α) + let y ← (IO.runRandPCGWith 42 prog : IO α) + Lean.logInfo m!"x = {x}" + Lean.logInfo m!"y (seed 42) = {y}" ---attribute [computable] sumTwo +attribute [computable] sumTwo ---run_cmd do logComputable sumTwoComputable +run_cmd do logComputable sumTwoComputable @[computable] noncomputable -def test : MeasureTheory.Measure ℝ := rdo +def unfoldSumTwo : MeasureTheory.Measure ℝ := rdo let y ← sumTwo let x ← ProbabilityTheory.gaussianReal 0 1 return x + y @@ -42,4 +39,12 @@ run_cmd do logComputable (branchOnComputable 20) run_cmd do logComputable (branchOnComputable (-1)) +attribute [computable] fairCoin + +run_cmd do logComputable (fairCoinComputable) + +attribute [computable] Bind.twoCoins + +run_cmd do logComputable (Bind.twoCoinsComputable) + end Test.Computable diff --git a/scripts/check_binomial.py b/scripts/check_binomial.py new file mode 100644 index 0000000..a26caf1 --- /dev/null +++ b/scripts/check_binomial.py @@ -0,0 +1,27 @@ +import numpy as np +from common import compare + +print("Checking binomial distribution...") + +# One pair per branch of numpy's `random_binomial`: inversion and BTPE, each on both sides of +# `p = 1/2`, on both sides of the `n * p = 30` threshold, and the three degenerate cases. +PARAMS = [(10, 0.3), (60, 0.5), (100, 0.31), (1000, 0.5), (1000, 0.9), (5, 0.99), + (0, 0.5), (10, 0.0), (10, 1.0)] + +LEAN = """import RandomDo +def params : List (Nat × Float) := + [%s] +def main (args : List String) : IO Unit := do + for s in args do + IO.FS.withFile (System.FilePath.mk s!"@DIR@/pcg64-{s}.txt") .write fun h ↦ + IO.runRandPCGWith s.toNat! do + for _ in List.range (@N@ / %d) do + for (n, p) in params do + h.putStrLn (toString (← NumLean.binomial n p)) +""" % (",\n ".join(f"({n}, {p})" for n, p in PARAMS), len(PARAMS)) + +def distrib(seed, N): + g = np.random.default_rng(seed) + return [g.binomial(n, p) for _ in range(N // len(PARAMS)) for n, p in PARAMS] + +compare(LEAN, distrib) From dfeb1876e946462c9caab73501e90649a35ff712 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Wed, 9 Sep 2026 11:04:19 +0200 Subject: [PATCH 18/34] Sample examples --- Test/Computable.lean | 10 +++++----- Test/Sample.lean | 31 +++++++++++++++++++++++++++++++ 2 files changed, 36 insertions(+), 5 deletions(-) create mode 100644 Test/Sample.lean diff --git a/Test/Computable.lean b/Test/Computable.lean index 3ef04f0..22c879c 100644 --- a/Test/Computable.lean +++ b/Test/Computable.lean @@ -20,7 +20,7 @@ def logComputable {α : Type} [Lean.ToMessageData α] (prog : RandPCG IO α) : C attribute [computable] sumTwo -run_cmd do logComputable sumTwoComputable +run_cmd logComputable sumTwoComputable @[computable] noncomputable @@ -31,20 +31,20 @@ def unfoldSumTwo : MeasureTheory.Measure ℝ := rdo attribute [computable] centred -run_cmd do logComputable (centredComputable 20) +run_cmd logComputable (centredComputable 20) attribute [computable] branchOn run_cmd do logComputable (branchOnComputable 20) -run_cmd do logComputable (branchOnComputable (-1)) +run_cmd logComputable (branchOnComputable (-1)) attribute [computable] fairCoin -run_cmd do logComputable (fairCoinComputable) +run_cmd logComputable (fairCoinComputable) attribute [computable] Bind.twoCoins -run_cmd do logComputable (Bind.twoCoinsComputable) +run_cmd logComputable (Bind.twoCoinsComputable) end Test.Computable diff --git a/Test/Sample.lean b/Test/Sample.lean new file mode 100644 index 0000000..1d1a133 --- /dev/null +++ b/Test/Sample.lean @@ -0,0 +1,31 @@ +module + +public import RandomDo + +set_option linter.style.header false + +open NumLean IO Lean + +run_cmd do + let x ← rand 0 1000 + logInfo m!"Random number between 0 and 1000: {x}" + +run_cmd do + let x ← (runRandPCG <| randInt 1000 : IO Int) + logInfo m!"Random number between 0 and 1000 (PCG): {x}" + +run_cmd do + let x ← (runRandPCG <| normal 10 2 : IO Float) + logInfo m!"{x}" + +run_cmd do + let x ← (runRandPCG <| exponential 10 : IO Float) + logInfo m!"{x}" + +run_cmd do + let x ← (runRandPCG <| binomial 10 0.5 : IO Nat) + logInfo m!"{x}" + +run_cmd do + let x ← (runRandPCG <| bernoulli 0.5 : IO Nat) + logInfo m!"{x}" From 52eeab6b13669dea36ff4fe775bac233bbcd6d2b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Wed, 9 Sep 2026 16:36:45 +0200 Subject: [PATCH 19/34] `mk_all` --- Test.lean | 1 + 1 file changed, 1 insertion(+) diff --git a/Test.lean b/Test.lean index 4009bb6..9a561fd 100644 --- a/Test.lean +++ b/Test.lean @@ -9,3 +9,4 @@ public import Test.Instances public import Test.IsMarkov public import Test.Loops public import Test.MonadLaws +public import Test.Sample From cdcbf76b41c248acacfa5e66d2b2eeee8de80568 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Wed, 9 Sep 2026 16:59:05 +0200 Subject: [PATCH 20/34] Support for mutable expr and for loop --- RandomDo/Tactic/Computable/Deriving.lean | 20 ++++++++++++++++---- Test/Computable.lean | 18 +++++++++++++++--- 2 files changed, 31 insertions(+), 7 deletions(-) diff --git a/RandomDo/Tactic/Computable/Deriving.lean b/RandomDo/Tactic/Computable/Deriving.lean index e9353ca..551d299 100644 --- a/RandomDo/Tactic/Computable/Deriving.lean +++ b/RandomDo/Tactic/Computable/Deriving.lean @@ -5,7 +5,7 @@ Authors: Gaëtan Serré -/ module -public import RandomDo.Tactic.Computable.Counterparts +public meta import RandomDo.Tactic.Computable.Counterparts public import RandomDo.Monad.MeasurableSpace public meta import Lean.Elab.Tactic.Basic @@ -23,9 +23,10 @@ noncomputable def shifted : Measure ℝ := rdo adds `shiftedComputable : RandPCG IO Float`, which draws from `NumLean.normal' 0 1` and adds one. -The Giry monad and its two operations become `RandPCG IO`, `pure` and `bind`. Anything else is -rebuilt from the counterpart `@[computable_as]` records for its head, with its arguments translated -in turn and its instances synthesized anew. A term translates into a term of the translation of its +The Giry monad and its two operations become `RandPCG IO`, `pure` and `bind`, the `for` loop of +`rdo` becomes the `for` loop of that monad, and a `let` stays a `let`. Anything else is rebuilt +from the counterpart `@[computable_as]` records for its head, with its arguments translated in turn +and its instances synthesized anew. A term translates into a term of the translation of its type; where the rebuilt one does not, its head is a definition nothing is known about, and its body is read in its place. `@[computable]` records the program it writes, so a program drawing from another translates into one calling that other's translation. @@ -33,6 +34,11 @@ another translates into one calling that other's translation. Two things extend the attribute: an `@[computable_as]` entry, and an alternative of `translate` for a construct of `rdo` it has not been taught. +The counterparts are imported `meta` as well as publicly, and so reach every module the attribute +does. A translated program is written to be run, and a `run_cmd` or an `#eval` runs it in the very +module that writes it: the interpreter then asks for the code of what it calls, which a plain +`import` does not carry. + `set_option trace.computable true` prints what each piece became. -/ @@ -61,6 +67,9 @@ partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := | MeasurableSpaceBind.mBind _ _ _ _ _ _ p k => mkAppOptM ``Bind.bind #[← computableMonad, none, none, none, ← translate σ p, ← translate σ k] + | MeasurableSpaceForIn.forIn _ _ _ _ _ _ xs init body => + mkAppOptM ``ForIn.forIn #[← computableMonad, none, none, none, none, + ← translate σ xs, ← translate σ init, ← translate σ body] | MeasureTheory.Measure α _ => return mkApp (← computableMonad) (← translate σ α) /- A subtype is its carrier, and one of its values is the value it carries: the constraint and the proof of it are what a computable counterpart does not have. -/ @@ -74,6 +83,9 @@ partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := let x := xs[0]!.fvarId! withLocalDeclD (← x.getUserName) (← translate σ (← x.getType)) fun y ↦ do mkLambdaFVars #[y] (← translate (σ.insert x y) body) + | .letE n t v b _ => withLetDecl n t v fun x ↦ do + withLetDecl n (← translate σ t) (← translate σ v) fun y ↦ do + mkLetFVars #[y] (← translate (σ.insert x.fvarId! y) (b.instantiate1 x)) | _ => translateApp σ e /-- Rebuild an application from the counterpart of its head; where nothing known about that head diff --git a/Test/Computable.lean b/Test/Computable.lean index 22c879c..0b7752b 100644 --- a/Test/Computable.lean +++ b/Test/Computable.lean @@ -2,6 +2,7 @@ module public import Test.IsMarkov public import Test.Bind +public meta import RandomDo import Batteries.Data.Float.Basic set_option linter.style.header false @@ -10,7 +11,7 @@ set_option trace.computable true namespace Test.Computable -open Test.IsMarkov NumLean Lean.Elab.Command +open Test.IsMarkov NumLean Lean.Elab.Command MeasureTheory ProbabilityTheory def logComputable {α : Type} [Lean.ToMessageData α] (prog : RandPCG IO α) : CommandElabM Unit := do let x ← (IO.runRandPCG prog : IO α) @@ -24,9 +25,9 @@ run_cmd logComputable sumTwoComputable @[computable] noncomputable -def unfoldSumTwo : MeasureTheory.Measure ℝ := rdo +def unfoldSumTwo : Measure ℝ := rdo let y ← sumTwo - let x ← ProbabilityTheory.gaussianReal 0 1 + let x ← gaussianReal 0 1 return x + y attribute [computable] centred @@ -47,4 +48,15 @@ attribute [computable] Bind.twoCoins run_cmd logComputable (Bind.twoCoinsComputable) +@[computable] +noncomputable +def ex1 : Measure ℝ := rdo + let mut x := 0 + for _ in List.range 1000 rdo + let y ← gaussianReal 0 1 + x := x + y + return x + +run_cmd logComputable (ex1Computable) + end Test.Computable From 71605d27ef4aefc736b0f336c993a727f6456b9c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Wed, 9 Sep 2026 17:11:44 +0200 Subject: [PATCH 21/34] Lint --- RandomDo/Tactic/Computable/Counterparts.lean | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/RandomDo/Tactic/Computable/Counterparts.lean b/RandomDo/Tactic/Computable/Counterparts.lean index 7243973..ee61542 100644 --- a/RandomDo/Tactic/Computable/Counterparts.lean +++ b/RandomDo/Tactic/Computable/Counterparts.lean @@ -30,7 +30,9 @@ attribute [computable_as Float] NNReal attribute [computable_as NumLean.normal'] ProbabilityTheory.gaussianReal -def bernoulliChoice (α : Type) [MeasurableSpace α] (x y : α) (p : Float) : +/-- Draw from a Bernoulli distribution with probability `p` of returning `x` and `1 - p` of +returning `y`. -/ +def bernoulliChoice (α : Type) [_h : MeasurableSpace α] (x y : α) (p : Float) : NumLean.RandPCG IO α := do return if (← NumLean.bernoulli p) == 1 then x else y From b8113f6016275b809d047e19328c40ad81ebc6ea Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Wed, 9 Sep 2026 17:59:09 +0200 Subject: [PATCH 22/34] Test `HasGaussian` --- Test/Computable_test.lean | 77 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 77 insertions(+) create mode 100644 Test/Computable_test.lean diff --git a/Test/Computable_test.lean b/Test/Computable_test.lean new file mode 100644 index 0000000..63f3933 --- /dev/null +++ b/Test/Computable_test.lean @@ -0,0 +1,77 @@ +module + +public import RandomDo +public meta import RandomDo + +set_option linter.style.header false + +open MeasureTheory ProbabilityTheory NumLean + +universe v + +variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] + +/-! ## Une capacité dont le type de valeur est le même des deux côtés + +Le crochet de la `MeasurableSpace` peut rester implicite dans le paramètre de la classe : Lean la +synthétise alors dans le type du champ, et `(by infer_instance)` devient inutile. -/ + +/-- Tirer un bit. -/ +class HasBit (m : (α : Type) → [MeasurableSpace α] → Type v) where + /-- Le tirage. -/ + bit : m Bool + +noncomputable instance : HasBit Measure where + bit := bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩ + +/-- Un programme qui ne dit pas dans quelle monade il vit. -/ +def sampleBitsArray [HasBit m] (n : ℕ) : m (Array Bool) := rdo + let mut xs : Array Bool := #[] + for _ in List.range n rdo + let b ← HasBit.bit (m := m) + xs := xs.push b + return xs + +/-! ## La même chose pour la gaussienne + +`Bool` est le même objet dans les deux mondes, mais les réels ne le sont pas : une mesure vit sur +`ℝ`, un échantillonneur rend un `Float`. La classe laisse donc le type des scalaires libre, et c'est +chaque instance qui le fixe. -/ + +/-- Un espace mesurable sur `Float` : il est fini, donc toutes ses parties sont mesurables. -/ +instance : MeasurableSpace Float := ⊤ + +/-- Tirer une gaussienne de moyenne et de variance données, à valeurs dans `R`. -/ +class HasGaussian (m : (α : Type) → [MeasurableSpace α] → Type v) + (R : Type) [MeasurableSpace R] where + /-- Le tirage, de moyenne le premier argument et de variance le second. -/ + gaussian : R → R → m R + +noncomputable instance : HasGaussian Measure ℝ where + gaussian μ v := gaussianReal μ v.toNNReal + +/-- La monade qui échantillonne, vue comme une `MeasurableSpaceMonad` comme les autres. -/ +abbrev RandM := Monad.toMeasurableSpaceMonad (RandPCG IO) + +instance : HasGaussian RandM Float where + gaussian μ v := normal' μ v + +/- Les scalaires sur lesquels un programme compte : de quoi écrire `0` et `+`. Les lois ne sont +pas demandées, seulement les opérations, ce qui laisse `Float` passer. -/ +variable {R : Type} [MeasurableSpace R] [Add R] [OfNat R 0] [OfNat R 1] + +/-- Un seul programme, écrit une fois. -/ +def ex1 [HasGaussian m R] : m R := rdo + let mut x : R := 0 + for _ in List.range 10 rdo + let y ← HasGaussian.gaussian (m := m) 0 1 + x := x + y + return x + +/-- Lu comme une mesure. -/ +noncomputable example : Measure ℝ := ex1 + +/- Lu comme un échantillonneur, et il tourne. -/ +run_cmd do + let x ← (IO.runRandPCGWith 42 (ex1 (m := RandM) (R := Float)) : IO Float) + Lean.logInfo m!"ex1 échantillonné (seed 42) = {x}" From dd3e68f82e0775b060dc4679a888b8eac241dc63 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 16 Sep 2026 10:31:15 +0200 Subject: [PATCH 23/34] improve transfer and extend_space tactics --- RandomDo/Probability/AlgTrace.lean | 7 +- RandomDo/Probability/Extend.lean | 249 +++++++++++++++----- RandomDo/Probability/MeasurePreserving.lean | 32 ++- RandomDo/Probability/Transfer.lean | 237 +++++++++++++------ Test/Extend.lean | 164 ++++++++++++- Test/Transfer.lean | 163 ++++++++++++- 6 files changed, 701 insertions(+), 151 deletions(-) diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean index 565fe59..385fabe 100644 --- a/RandomDo/Probability/AlgTrace.lean +++ b/RandomDo/Probability/AlgTrace.lean @@ -78,9 +78,10 @@ lemma _root_.Learning.IsAlgEnvSeq.comp_measurePreserving {𝓐 𝓨 Ω Ω' : Typ 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; -`h.measurable_action` and `h.measurable_feedback` provide the side conditions when an -`IsAlgEnvSeq` hypothesis `h` is around. -/ +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 Ω'} diff --git a/RandomDo/Probability/Extend.lean b/RandomDo/Probability/Extend.lean index ecc661f..32a7010 100644 --- a/RandomDo/Probability/Extend.lean +++ b/RandomDo/Probability/Extend.lean @@ -25,24 +25,35 @@ product measure, presented as an abstract space related to the old one by a meas 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`, `hZ : HasLaw Z μ P` and - `hind`, the independence of `Z` from the transported random variables, as a tuple. + 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. + 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`, `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. + `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 hypothesis the goal itself depends on, such as a measurability proof -inside a `Kernel.comap`, is generalized along with the goal. +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 -a random variable `X`, extend with `κ.comap X hX`: `extend_space` then states `hZ` as -`HasCondDistrib Z X κ P`. +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 @@ -95,7 +106,8 @@ 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 draw `Z` with conditional law `κ` given `f`, *provided* +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: @@ -108,18 +120,19 @@ theorem wlog_extend_kernel {Ω : Type (max u v)} [mΩ : MeasurableSpace Ω] {E : (κ : Kernel Ω E) [IsMarkovKernel κ] (extended : ∀ (Ω' : Type (max u v)) [MeasurableSpace Ω'] (P' : Measure Ω') [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → ∀ (Z : Ω' → E), - HasCondDistrib Z f κ P' → motive Ω' P' f) + 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 - (hasCondDistrib_snd_fst_compProd κ)) + 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 draw `Z` with law `μ`, independent of `f`, *provided* +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: @@ -132,13 +145,13 @@ theorem wlog_extend {Ω : Type (max u v)} [mΩ : MeasurableSpace Ω] {E : Type v (μ : Measure E) [IsProbabilityMeasure μ] (extended : ∀ (Ω' : Type (max u v)) [MeasurableSpace Ω'] (P' : Measure Ω') [IsProbabilityMeasure P'] (f : Ω' → Ω), MeasurePreserving f P' P → ∀ (Z : Ω' → E), - HasLaw Z μ P' → IndepFun f Z P' → motive Ω' P' f) + 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 + (extended (Ω × E) (P.prod μ) Prod.fst measurePreserving_fst Prod.snd measurable_snd measurePreserving_snd.hasLaw (indepFun_fst_snd_prod μ)) end RDo @@ -261,10 +274,10 @@ where if !deps.contains f || special.contains f then return ← go todo seen acc let d ← f.getDecl let ty ← instantiateMVars d.type - if transportable sp.Ω ty then return ← go todo seen acc 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" + 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. -/ @@ -301,6 +314,8 @@ structure NewSpace where 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. -/ @@ -340,16 +355,18 @@ partial def foldTransports (f : FVarId) (subst : Array (FVarId × Nat × Bool × | _ => none | _ => none -/-- The old-space version of a transported random variable `x : ι₁ → ⋯ → ιₖ → Ω → α`, as a -random variable `fun ω ↦ fun i₁ … iₖ ↦ x i₁ … iₖ ω` for the independence statement. -/ -partial def oldComponent (Ω ω x ty : Expr) : MetaM Expr := - match ty with +/-- 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] (← oldComponent Ω ω (mkApp x i) (b.instantiate1 i)) + 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)` @@ -408,35 +425,47 @@ def proveTransported (h : FVarId) (φ' : Expr) (new : NewSpace) : TacticM (Optio if let some pf ← tryTactic? goal.mvarId! tac then return some pf return none -/-- The independence of `Z` from the transported random variables, as the independence of `Z` -and the tuple of those whose measurability `fun_prop` can prove. -/ +/-- 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 mut comps : Array (Expr × Expr) := #[] - for (x, _, _, isSet) in transported do - if isSet then continue + 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 #[ω] (← oldComponent ΩE ω (.fvar x) ty) - let c := c.eta + mkLambdaFVars #[ω] (← component ΩE ω (.fvar x) ty isSet) + let c := if isSet then c else c.eta let goal ← mkFreshExprSyntheticOpaqueMVar (← mkAppM ``Measurable #[c]) - if let some hc ← tryTactic? goal.mvarId! (← `(tactic| fun_prop)) then - comps := comps.push (c, hc) + -- 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) → MetaM (Expr × Expr) + let rec mkTuple : List (Expr × Expr × Expr) → MetaM (Expr × Expr × Expr) | [] => throwError "extend_space: internal error, empty tuple" | [ch] => pure ch - | (c, hc) :: rest => do - let (r, hr) ← mkTuple rest + | (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]) - pure (φ, ← mkAppM ``Measurable.prodMk #[hc, hr]) - let (φ, hφ) ← mkTuple comps.toList - let fE := Expr.fvar new.f - let tuple ← withLocalDecl `ω .default (.fvar new.Ω) fun ω ↦ do - mkLambdaFVars #[ω] (← Core.betaReduce (mkApp φ (mkApp fE ω))) + 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) @@ -448,18 +477,18 @@ def deriveIndep (g : MVarId) (sp : ProbSpace) (new : NewSpace) (hind : FVarId) let (hind', g) ← g.intro1P return some (hind', g) -/-- For a kernel `κ.comap X hX` with `X` a transported variable, the conditional law of `Z` given -`X` on the new space, `HasCondDistrib Z (fun ω ↦ X (f ω)) κ P'`. -/ -def deriveCondDistrib (g : MVarId) (new : NewSpace) (κ : Expr) (transported : FVarIdSet) : +/-- 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 .fvar x := X | return none - unless transported.contains x do return none let κ₀ := κ.getArg! 6 let fE := Expr.fvar new.f let Xf ← withLocalDecl `ω .default (.fvar new.Ω) fun ω ↦ - mkLambdaFVars #[ω] (mkApp X (mkApp fE ω)) + 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) @@ -469,12 +498,65 @@ def deriveCondDistrib (g : MVarId) (new : NewSpace) (κ : Expr) (transported : F 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, for a kernel `κ.comap X hX`, the conditional law of `Z` is stated given `X`. -With `clearOld`, the old space, the map and everything mentioning them are cleared. -/ +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.Ω @@ -489,10 +571,11 @@ def hidePresentation (g : MVarId) (sp : ProbSpace) (deps special : FVarIdSet) (n 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. + -- 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 ← g.rename x (n.appendAfter "₀") + 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Ω) @@ -509,7 +592,10 @@ def hidePresentation (g : MVarId) (sp : ProbSpace) (deps special : FVarIdSet) (n let mut out : Array (FVarId × Expr × Nat × Bool) := #[] for (x, _) in olds do if special.contains x then continue - let ty ← instantiateMVars (← x.getType) + 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) @@ -548,7 +634,7 @@ def hidePresentation (g : MVarId) (sp : ProbSpace) (deps special : FVarIdSet) (n 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 κ transportedSet then + if let some (hZ', g') ← deriveCondDistrib g new κ then g := g' moved := moved.push hZ' toClear := toClear.push new.hZ @@ -563,7 +649,8 @@ def hidePresentation (g : MVarId) (sp : ProbSpace) (deps special : FVarIdSet) (n hdefs := hdefs' -- 8. Clean up. if clearOld then - toClear := toClear ++ hdefs ++ #[new.f, new.hf] ++ olds.map (·.1) + 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 @@ -599,6 +686,12 @@ def extendSpace (mode : ExtendMode) (μ : Term) (P? : Option Ident) (given : Arr 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)) @@ -617,6 +710,11 @@ def extendSpace (mode : ExtendMode) (μ : Term) (P? : Option Ident) (given : Arr "{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}" @@ -632,6 +730,14 @@ def extendSpace (mode : ExtendMode) (μ : Term) (P? : Option Ident) (given : Arr 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 ← @@ -693,17 +799,31 @@ def extendSpace (mode : ExtendMode) (μ : Term) (P? : Option Ident) (given : Arr | _, 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 defaults[i]! + 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, pick d 5] + #[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, pick d 1] + #[.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 @@ -714,7 +834,7 @@ def extendSpace (mode : ExtendMode) (μ : Term) (P? : Option Ident) (given : Arr | .map => pure extended | _ => let new : NewSpace := ⟨fvs[0]!, fvs[1]!, fvs[2]!, fvs[3]!, fvs[4]!, fvs[5]!, fvs[6]!, - fvs[7]!, if isKernel then none else some fvs[8]!⟩ + 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`. @@ -749,8 +869,9 @@ are kept: `Ω`, `P`, every random variable `X : Ω → α` and every event `s : objects on the extended space, hypotheses about them are transported, and the goal reads as before. The context gains -* `Z : Ω → E`, `hZ : HasLaw Z μ P`, and `hind`, the independence of `Z` from the transported - random variables, as a tuple; +* `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. @@ -760,10 +881,13 @@ the statement pulls back along a measure-preserving map, which makes the replace `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 X hX` with - `X` a random variable, `hZ` is stated as `HasCondDistrib Z X κ P`. + `κ` 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). @@ -778,7 +902,8 @@ syntax (name := extendSpaceClearTac) "extend_space!" ppSpace term (" using " ide /-- `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`, `hZ : HasLaw Z μ P'` and `hind : IndepFun f Z P'`. +`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 diff --git a/RandomDo/Probability/MeasurePreserving.lean b/RandomDo/Probability/MeasurePreserving.lean index e306cc8..ef3cfa9 100644 --- a/RandomDo/Probability/MeasurePreserving.lean +++ b/RandomDo/Probability/MeasurePreserving.lean @@ -18,8 +18,9 @@ set_option linter.style.header false 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, conditional -laws, integrability. This file collects these facts in the forms the `transfer` tactic uses. +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. @@ -102,6 +103,17 @@ lemma transfer_indepFun (hf : MeasurePreserving f P' P) (hX : Measurable X) (hY 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 𝓧 𝓨} : @@ -153,6 +165,15 @@ variable {Ω Ω' 𝓧 𝓨 : Type*} {mΩ : MeasurableSpace Ω} {mΩ' : Measurabl {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 ω) := @@ -200,6 +221,13 @@ lemma ProbabilityTheory.IndepFun.comp_measurePreserving (h : IndepFun X Y P) 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) : diff --git a/RandomDo/Probability/Transfer.lean b/RandomDo/Probability/Transfer.lean index 4996997..aae75a2 100644 --- a/RandomDo/Probability/Transfer.lean +++ b/RandomDo/Probability/Transfer.lean @@ -6,6 +6,7 @@ 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 @@ -35,10 +36,12 @@ on. Two kinds of lemmas record this. `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: this pulls a fact about the old space back to -the new one. `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. +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 @@ -84,7 +87,7 @@ syntax (name := transferDischarger) "transfer_discharger" : tactic elab_rules : tactic | `(tactic| transfer_discharger) => withMainContext do let funProps : Array Name := #[``Measurable, ``AEMeasurable, - `MeasureTheory.AEStronglyMeasurable, `MeasureTheory.StronglyMeasurable] + ``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 @@ -97,13 +100,16 @@ elab_rules : tactic | (apply Measurable.aemeasurable; fun_prop) | (apply Measurable.aestronglyMeasurable; fun_prop))) else `(tactic| first | assumption | (intros; first | assumption | measurability)) - tryCatchRuntimeEx (evalTactic tac) fun e ↦ + -- 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 (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, which put the transferred events in the same form as -`extend_space`. 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. -/ +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 @@ -118,8 +124,8 @@ def transferSimpArgs (hf : Term) : TacticM (Array (TSyntax ``Lean.Parser.Tactic. (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_inter, ``Set.preimage_union, - ``Set.preimage_compl, ``Set.preimage_sdiff].mapM fun n ↦ + 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 @@ -138,36 +144,70 @@ def tryTactic? (g : MVarId) (tac : Syntax) : TacticM (Option Expr) := 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`. Returns the -statement and proof of the transported hypothesis. -/ +`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 hStx ← Term.exprToSyntax h let hfStx ← Term.exprToSyntax hf - for n in ← labelled `transfer_forward do - 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 (concl, pf))) - (fun _ ↦ do - s.restore - pure none) - if r.isSome then return r - return none + 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 @@ -188,42 +228,98 @@ where return (← whnfR (← fv.getType)).isAppOf ``MeasureTheory.MeasurePreserving go g (if isMap then some fv else none) -/-- Transfer the goal along `hf`: rewrite it with the `@[transfer]` lemmas, then close it by -`assumption` if possible. -/ -def transferGoal (hfStx : Term) : TacticM Unit := do - let args ← transferSimpArgs hfStx +/-- 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,*])) - unless (← getUnsolvedGoals).isEmpty do - evalTactic (← `(tactic| try assumption)) + 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. -/ -def transferHyps (hfStx : Term) (hf : Expr) (hs : Array FVarId) : TacticM Unit := do - let args ← transferSimpArgs hfStx +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)) - for h in hs do - withMainContext do - -- `simp` replaces the hypothesis when it rewrites it; otherwise it is still there. - let some d := (← getLCtx).find? h | return - let some (ty, pf) ← transferForward? d.toExpr hf | return - let g ← (← getMainGoal).assert d.userName ty pf - let (_, g) ← g.intro1P - replaceMainGoal [← g.tryClear h] + 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. +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. + 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. -/ @@ -237,18 +333,27 @@ elab_rules : tactic 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 ← withMainContext do - (← getLCtx).foldlM (init := #[]) fun hs d ↦ do - if d.isImplementationDetail || !(← isProp d.type) then pure hs - else pure (hs.push d.fvarId) - transferHyps hf hfE hs - transferGoal hf + 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 => - transferHyps hf hfE (← withMainContext do hyps.mapM getFVarId) - if type then transferGoal hf - | some hf, none => transferGoal hf + 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] diff --git a/Test/Extend.lean b/Test/Extend.lean index 1a7b6a6..fbfe50b 100644 --- a/Test/Extend.lean +++ b/Test/Extend.lean @@ -1,7 +1,8 @@ module -public import Test.Common public import Mathlib.Probability.Independence.InfinitePi +public import RandomDo.Probability.Extend +public import RandomDo.Probability.MeasurePreserving set_option linter.style.header false @@ -10,9 +11,9 @@ set_option linter.style.header false 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, and what `extend_space!` clears. 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. +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 `Ω`. @@ -37,16 +38,17 @@ include μ /-- 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. The goal reads as -before, and its `transfer` obligation is discharged. -/ +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 ω)) 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 @@ -126,6 +128,68 @@ example {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasur 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 @@ -184,6 +248,21 @@ example (X Y : Ω → ℝ) (hX : Measurable X) (hY : Measurable Y) (h : X =ᵐ[P 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 @@ -227,6 +306,13 @@ example (X : Ω → ℝ) (hX : Measurable X) (κ : Kernel Ω E) [IsMarkovKernel 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`, @@ -238,6 +324,7 @@ example (X : Ω → ℝ) (A : ℕ → Ω → ℝ) (s : Set Ω) (hX : Measurable 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 = ν @@ -316,6 +403,69 @@ example {E' : Type (u + 1)} [MeasurableSpace E'] (μ' : Measure E') [IsProbabili (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 diff --git a/Test/Transfer.lean b/Test/Transfer.lean index 492c9fb..215eec7 100644 --- a/Test/Transfer.lean +++ b/Test/Transfer.lean @@ -1,6 +1,7 @@ module -public import Test.Common +public import RandomDo.Probability.MeasurePreserving +public import RandomDo.Probability.Transfer set_option linter.style.header false @@ -8,11 +9,13 @@ 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: on a goal, and on hypotheses, where the rewriting -by the `@[transfer]` lemmas and the fallback on the `@[transfer_forward]` lemmas both show. +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 RDo +open MeasureTheory ProbabilityTheory @[expose] public section @@ -22,20 +25,41 @@ namespace Test.Transfer universe u -variable {Ω : Type u} [MeasurableSpace Ω] {P : Measure Ω} +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 ℝ) {Ω' : Type u} [MeasurableSpace Ω'] - {P' : Measure Ω'} (f : Ω' → Ω) (hf : MeasurePreserving f P' P) - (h : HasLaw (fun ω ↦ X (f ω)) ν P') : +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) - {Ω' : Type u} [MeasurableSpace Ω'] {P' : Measure Ω'} (f : Ω' → Ω) - (hf : MeasurePreserving f P' P) : +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 ω) @@ -43,6 +67,74 @@ example (X : Ω → ℝ) (hX : Measurable X) (s : Set Ω) (hs : MeasurableSet 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 -/ /-- @@ -55,6 +147,55 @@ but is 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 From d983f5d3d87528a8dd696df9e3dc7f54786de3a6 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Wed, 16 Sep 2026 10:56:35 +0200 Subject: [PATCH 24/34] improve algtrace --- RandomDo/Probability/AlgTrace.lean | 199 +++++++++++++++++++++-------- RandomDo/Probability/Transfer.lean | 21 ++- Test/AlgTrace.lean | 152 +++++++++++++++++++++- 3 files changed, 313 insertions(+), 59 deletions(-) diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean index 385fabe..c49f21f 100644 --- a/RandomDo/Probability/AlgTrace.lean +++ b/RandomDo/Probability/AlgTrace.lean @@ -58,7 +58,7 @@ open MeasureTheory ProbabilityTheory Finset Learning noncomputable section -attribute [fun_prop] Learning.measurable_history +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 @@ -107,7 +107,7 @@ lemma _root_.MeasureTheory.MeasurePreserving.transfer_isAlgEnvSeq {𝓐 𝓨 Ω namespace RDo -universe uA uY uW +universe u uA uY uW variable {𝓐 : Type uA} {𝓨 : Type uY} {Ω : Type uW} [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] @@ -279,55 +279,73 @@ end Projection `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. -/ +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 uA uY uW)) (_ : MeasurableSpace Ω') (P' : Measure Ω') + ∃ (Ω' : 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 - have h := IT.isAlgEnvSeq_trajMeasure tr.algorithm (env.withTrace Ω) - refine ⟨ℕ → (Ω × 𝓐) × 𝓨, inferInstance, trajMeasure tr.algorithm (env.withTrace Ω), - inferInstance, fun n (ω : ℕ → (Ω × 𝓐) × 𝓨) ↦ (IT.action n ω).2, IT.feedback, - fun n (ω : ℕ → (Ω × 𝓐) × 𝓨) ↦ (IT.action n ω).1, - tr.isAlgEnvSeq_snd h, ?_, tr.hasLaw_trace_zero h, tr.hasCondDistrib_trace h, - tr.action_zero_ae_eq h, tr.action_ae_eq h⟩ - exact isAlgEnvSeq_unique (tr.isAlgEnvSeq_snd h) h₀ + -- 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. -/ +`isAlgEnvSeq_unique` gives. The space may live in any universe at least those of `𝓐`, `𝓨` and +`Ω`. -/ theorem wlog_trace [MeasurableEq 𝓐] - {motive : (Ω₀ : Type (max uA uY uW)) → [MeasurableSpace Ω₀] → (P : Measure Ω₀) → + {motive : (Ω₀ : Type (max u uA uY uW)) → [MeasurableSpace Ω₀] → (P : Measure Ω₀) → [IsProbabilityMeasure P] → (ℕ → Ω₀ → 𝓐) → (ℕ → Ω₀ → 𝓨) → Prop} - (traced : ∀ (Ω' : Type (max uA uY uW)) [MeasurableSpace Ω'] (P' : Measure Ω') + (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 uA uY uW)) [MeasurableSpace Ω₁] (P₁ : Measure Ω₁) + (transfer : ∀ (Ω₁ : Type (max u uA uY uW)) [MeasurableSpace Ω₁] (P₁ : Measure Ω₁) [IsProbabilityMeasure P₁] (A₁ : ℕ → Ω₁ → 𝓐) (Y₁ : ℕ → Ω₁ → 𝓨) - (Ω₂ : Type (max uA uY uW)) [MeasurableSpace Ω₂] (P₂ : Measure Ω₂) + (Ω₂ : 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 uA uY uW)) [MeasurableSpace Ω₀] (P : Measure Ω₀) [IsProbabilityMeasure P] + ∀ (Ω₀ : 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 - obtain ⟨Ω', mΩ', P', hP', A', Y', T, hseq, hlaw, hT0, hT, hA0, hA⟩ := - tr.exists_isAlgEnvSeq_trace h - exact transfer Ω₀ P A Y Ω' P' A' Y' h hseq hlaw (traced Ω' P' A' Y' T hseq hT0 hT hA0 hA) + -- 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 @@ -400,22 +418,29 @@ def findAlgEnvSeq? : MetaM (Option FVarId) := do 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 mentioning the probability space, the measure or the two -sequences, is abstracted away from that space and two goals are left: +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 `T`'s law, its - conditional law given the history, and the equations expressing each action as the readout of the - history and the draws; +* `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. -/ +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 @@ -428,17 +453,71 @@ elab_rules : tactic | some f => pure f | none => throwError "alg_env_trace: no `IsAlgEnvSeq` hypothesis in the context" let spaceFVars ← algEnvSpaceFVars hFVar - -- Everything else that mentions the space has to travel with the goal, or it would be lost. - let deps ← do - let mut deps : Array FVarId := #[] - for d in ← getLCtx do - if !d.isImplementationDetail && !spaceFVars.contains d.fvarId then - let dty ← instantiateMVars d.type - if spaceFVars.any fun f ↦ dty.containsFVar f then - deps := deps.push d.fvarId - pure deps - let (_, g) ← g.revert deps - let (_, g) ← g.revert spaceFVars (preserveOrder := true) + 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 @@ -455,9 +534,14 @@ elab_rules : tactic 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)}" + (← 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))) @@ -470,14 +554,12 @@ elab_rules : tactic transfer.setTag `transfer return (traced, transfer, args[iMotive]!) -- Introduce the traced space and its properties, then whatever travelled with the goal. - let given := (names?.map (·.map (·.getId))).getD #[] - let defaults : Array Name := #[`Ω, `P, `A, `Y, `T, `hseq, `hT₀, `hT, `hA₀, `hA] 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 8, pick 9, pick 10] let (_, traced) ← traced.introN intros.size intros.toList - let (_, traced) ← traced.introNP deps.size + 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. @@ -488,18 +570,35 @@ elab_rules : tactic 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 - generalize hν : Measure.map (trajectory A₁ Y₁) P₁ = ν at hlaw - have hf₁ : MeasurePreserving (trajectory A₁ Y₁) P₁ ν := - ⟨measurable_trajectory h₁.measurable_action h₁.measurable_feedback, hν⟩ - have hf₂ : MeasurePreserving (trajectory A₂ Y₂) P₂ ν := - ⟨measurable_trajectory h₂.measurable_action h₂.measurable_feedback, hlaw⟩ - have : IsProbabilityMeasure ν := hν ▸ Measure.isProbabilityMeasure_map - (measurable_trajectory h₁.measurable_action h₁.measurable_feedback).aemeasurable + 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 ↦ (t n).1) (fun n t ↦ (t n).2) ↦ (by transfer hf₁ at hS; exact hS)) + (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 []) diff --git a/RandomDo/Probability/Transfer.lean b/RandomDo/Probability/Transfer.lean index aae75a2..eca63b3 100644 --- a/RandomDo/Probability/Transfer.lean +++ b/RandomDo/Probability/Transfer.lean @@ -78,10 +78,21 @@ 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. -/ +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 @@ -102,7 +113,7 @@ elab_rules : tactic 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 (Tactic.withoutRecover (evalTactic tac)) fun e ↦ + 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 @@ -130,13 +141,15 @@ def transferSimpArgs (hf : Term) : TacticM (Array (TSyntax ``Lean.Parser.Tactic. 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. -/ +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 ← Term.withoutErrToSorry <| Tactic.run g (evalTactic tac) + 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) diff --git a/Test/AlgTrace.lean b/Test/AlgTrace.lean index 75e2dc3..c2055a4 100644 --- a/Test/AlgTrace.lean +++ b/Test/AlgTrace.lean @@ -9,8 +9,10 @@ set_option linter.style.header false 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 same algorithm then -exercises `extend_space` alongside an algorithm-environment sequence. +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 @@ -25,6 +27,8 @@ 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 @@ -84,23 +88,84 @@ theorem exists_noise (env : Environment (Fin K) ℝ) {Ω₀ : Type*} [Measurable ∧ 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⟩ := + 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 mentioning the space travels with the goal, so nothing is +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 hZ₀ hZ hA₀ hA + 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 Ω₀} @@ -113,6 +178,83 @@ example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] { 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 From cebedb49a36dce2a60f3204f5731b02a0cd6354f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Wed, 16 Sep 2026 14:54:38 +0200 Subject: [PATCH 25/34] Polymorphic test --- Compute.lean | 15 ++++ Polymorphic.lean | 22 ++++++ RandomDo.lean | 1 + RandomDo/Tactic/Computable/Polymorphic.lean | 36 ++++++++++ Test/Computable_test.lean | 77 --------------------- lakefile.toml | 10 ++- 6 files changed, 83 insertions(+), 78 deletions(-) create mode 100644 Compute.lean create mode 100644 Polymorphic.lean create mode 100644 RandomDo/Tactic/Computable/Polymorphic.lean delete mode 100644 Test/Computable_test.lean diff --git a/Compute.lean b/Compute.lean new file mode 100644 index 0000000..a04cb6c --- /dev/null +++ b/Compute.lean @@ -0,0 +1,15 @@ +import RandomDo + +open MeasureTheory ProbabilityTheory NumLean + +@[computable] +noncomputable def ex1 : Measure ℝ := rdo + let mut x := 0 + for _ in List.range 1000000 rdo + let y ← gaussianReal 0 1 + x := x + y + return x + +def main : IO Unit := do + let x ← (IO.runRandPCGWith 42 ex1Computable : IO Float) + IO.println s!"x = {x}" diff --git a/Polymorphic.lean b/Polymorphic.lean new file mode 100644 index 0000000..fa77014 --- /dev/null +++ b/Polymorphic.lean @@ -0,0 +1,22 @@ +import RandomDo + +open MeasureTheory ProbabilityTheory NumLean + +universe v + +variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] +variable {α β : Type*} +-- We could define a typeclass for type that have "classical" behavior, inheriting from `Add`, etc. +variable {R : Type} [MeasurableSpace R] [Add R] [OfNat R 0] [OfNat α 0] [OfNat β 1] + +def ex1 [HasGaussian m α β R] : m R := rdo + let mut x : R := 0 + for _ in List.range 1000000 rdo + let y ← HasGaussian.gaussian (m := m) (α := α) (β := β) 0 1 + x := x + y + return x + +def main : IO Unit := do + let x ← (IO.runRandPCGWith 42 + (ex1 (m := RandM) (α := Float) (β := Float) (R := Float)) : IO Float) + IO.println s!"x = {x}" diff --git a/RandomDo.lean b/RandomDo.lean index cbabc80..e37619e 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -16,6 +16,7 @@ public import RandomDo.Tactic.Computable.Counterparts public import RandomDo.Tactic.Computable.Defs public import RandomDo.Tactic.Computable.Deriving public import RandomDo.Tactic.Computable.Example +public import RandomDo.Tactic.Computable.Polymorphic public import RandomDo.Tactic.IsMarkov.Defs public import RandomDo.Tactic.IsMarkov.Deriving public import RandomDo.Tactic.IsMarkov.Elab diff --git a/RandomDo/Tactic/Computable/Polymorphic.lean b/RandomDo/Tactic/Computable/Polymorphic.lean new file mode 100644 index 0000000..85279d7 --- /dev/null +++ b/RandomDo/Tactic/Computable/Polymorphic.lean @@ -0,0 +1,36 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.Probability.Distributions.Gaussian.Real +public import RandomDo.Monad.Instances +public import RandomDo.NumLean.Distributions + +/-! + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory NumLean + +universe v + +/-- A typeclass for monads that can draw from a Gaussian distribution. -/ +class HasGaussian (m : (α : Type) → [MeasurableSpace α] → Type v) + (α β : Type*) (R : Type) [MeasurableSpace R] where + gaussian : α → β → m R + +noncomputable instance : HasGaussian Measure ℝ NNReal ℝ where + gaussian μ v := gaussianReal μ v + +/-- The monad that samples, seen as a `MeasurableSpaceMonad`. -/ +abbrev RandM := Monad.toMeasurableSpaceMonad (RandPCG IO) + +instance : MeasurableSpace Float := ⊤ + +instance : HasGaussian RandM Float Float Float where + gaussian μ v := normal' μ v diff --git a/Test/Computable_test.lean b/Test/Computable_test.lean deleted file mode 100644 index 63f3933..0000000 --- a/Test/Computable_test.lean +++ /dev/null @@ -1,77 +0,0 @@ -module - -public import RandomDo -public meta import RandomDo - -set_option linter.style.header false - -open MeasureTheory ProbabilityTheory NumLean - -universe v - -variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] - -/-! ## Une capacité dont le type de valeur est le même des deux côtés - -Le crochet de la `MeasurableSpace` peut rester implicite dans le paramètre de la classe : Lean la -synthétise alors dans le type du champ, et `(by infer_instance)` devient inutile. -/ - -/-- Tirer un bit. -/ -class HasBit (m : (α : Type) → [MeasurableSpace α] → Type v) where - /-- Le tirage. -/ - bit : m Bool - -noncomputable instance : HasBit Measure where - bit := bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩ - -/-- Un programme qui ne dit pas dans quelle monade il vit. -/ -def sampleBitsArray [HasBit m] (n : ℕ) : m (Array Bool) := rdo - let mut xs : Array Bool := #[] - for _ in List.range n rdo - let b ← HasBit.bit (m := m) - xs := xs.push b - return xs - -/-! ## La même chose pour la gaussienne - -`Bool` est le même objet dans les deux mondes, mais les réels ne le sont pas : une mesure vit sur -`ℝ`, un échantillonneur rend un `Float`. La classe laisse donc le type des scalaires libre, et c'est -chaque instance qui le fixe. -/ - -/-- Un espace mesurable sur `Float` : il est fini, donc toutes ses parties sont mesurables. -/ -instance : MeasurableSpace Float := ⊤ - -/-- Tirer une gaussienne de moyenne et de variance données, à valeurs dans `R`. -/ -class HasGaussian (m : (α : Type) → [MeasurableSpace α] → Type v) - (R : Type) [MeasurableSpace R] where - /-- Le tirage, de moyenne le premier argument et de variance le second. -/ - gaussian : R → R → m R - -noncomputable instance : HasGaussian Measure ℝ where - gaussian μ v := gaussianReal μ v.toNNReal - -/-- La monade qui échantillonne, vue comme une `MeasurableSpaceMonad` comme les autres. -/ -abbrev RandM := Monad.toMeasurableSpaceMonad (RandPCG IO) - -instance : HasGaussian RandM Float where - gaussian μ v := normal' μ v - -/- Les scalaires sur lesquels un programme compte : de quoi écrire `0` et `+`. Les lois ne sont -pas demandées, seulement les opérations, ce qui laisse `Float` passer. -/ -variable {R : Type} [MeasurableSpace R] [Add R] [OfNat R 0] [OfNat R 1] - -/-- Un seul programme, écrit une fois. -/ -def ex1 [HasGaussian m R] : m R := rdo - let mut x : R := 0 - for _ in List.range 10 rdo - let y ← HasGaussian.gaussian (m := m) 0 1 - x := x + y - return x - -/-- Lu comme une mesure. -/ -noncomputable example : Measure ℝ := ex1 - -/- Lu comme un échantillonneur, et il tourne. -/ -run_cmd do - let x ← (IO.runRandPCGWith 42 (ex1 (m := RandM) (R := Float)) : IO Float) - Lean.logInfo m!"ex1 échantillonné (seed 42) = {x}" diff --git a/lakefile.toml b/lakefile.toml index 1dd6684..db861c8 100644 --- a/lakefile.toml +++ b/lakefile.toml @@ -1,5 +1,5 @@ name = "RandomDo" -defaultTargets = ["RandomDo"] +defaultTargets = ["RandomDo", "compute", "polymorphic"] lintDriver = "batteries/runLinter" testDriver = "Test" @@ -29,3 +29,11 @@ name = "Test" name = "dump" root = "Dump" srcDir = "test_data" + +[[lean_exe]] +name = "compute" +root = "Compute" + +[[lean_exe]] +name = "polymorphic" +root = "Polymorphic" From 14c26e6dfae65f25facf48f3b571d9a99754b4d7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ga=C3=ABtan=20Serr=C3=A9?= Date: Wed, 16 Sep 2026 15:51:31 +0200 Subject: [PATCH 26/34] Lint --- RandomDo/Tactic/Computable/Polymorphic.lean | 3 ++- lakefile.toml | 2 +- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/RandomDo/Tactic/Computable/Polymorphic.lean b/RandomDo/Tactic/Computable/Polymorphic.lean index 85279d7..f890657 100644 --- a/RandomDo/Tactic/Computable/Polymorphic.lean +++ b/RandomDo/Tactic/Computable/Polymorphic.lean @@ -22,6 +22,7 @@ universe v /-- A typeclass for monads that can draw from a Gaussian distribution. -/ class HasGaussian (m : (α : Type) → [MeasurableSpace α] → Type v) (α β : Type*) (R : Type) [MeasurableSpace R] where + /-- Draw a sample from a Gaussian distribution with mean `μ` and variance `v`. -/ gaussian : α → β → m R noncomputable instance : HasGaussian Measure ℝ NNReal ℝ where @@ -30,7 +31,7 @@ noncomputable instance : HasGaussian Measure ℝ NNReal ℝ where /-- The monad that samples, seen as a `MeasurableSpaceMonad`. -/ abbrev RandM := Monad.toMeasurableSpaceMonad (RandPCG IO) -instance : MeasurableSpace Float := ⊤ +instance instMeasurableSpaceFloat : MeasurableSpace Float := ⊤ instance : HasGaussian RandM Float Float Float where gaussian μ v := normal' μ v diff --git a/lakefile.toml b/lakefile.toml index db861c8..8655135 100644 --- a/lakefile.toml +++ b/lakefile.toml @@ -1,5 +1,5 @@ name = "RandomDo" -defaultTargets = ["RandomDo", "compute", "polymorphic"] +defaultTargets = ["RandomDo"] lintDriver = "batteries/runLinter" testDriver = "Test" From 3294a53fe09a1415e0fd3d8b569e0f663c0d7d8e Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Fri, 18 Sep 2026 16:02:24 +0200 Subject: [PATCH 27/34] MetropolisHastings example --- .gitignore | 3 +- MetropolisHastings.lean | 31 ++ MetropolisHastings/Computable.lean | 32 ++ MetropolisHastings/Defs.lean | 69 +++++ MetropolisHastings/Main.lean | 142 +++++++++ MetropolisHastings/Polymorphic.lean | 112 +++++++ MetropolisHastings/Targets.lean | 80 +++++ MetropolisHastings/Theory.lean | 271 +++++++++++++++++ RandomDo/Measurable.lean | 10 + RandomDo/Tactic/Computable/Deriving.lean | 9 + RandomDo/Tactic/Computable/Polymorphic.lean | 45 +++ RandomDo/Tactic/IsMarkov/Elab.lean | 14 +- RandomDo/Tactic/IsMarkov/Lemmas.lean | 24 +- Test/IsMarkov.lean | 6 + lakefile.toml | 7 + scripts/mh_plot.py | 315 ++++++++++++++++++++ 16 files changed, 1163 insertions(+), 7 deletions(-) create mode 100644 MetropolisHastings.lean create mode 100644 MetropolisHastings/Computable.lean create mode 100644 MetropolisHastings/Defs.lean create mode 100644 MetropolisHastings/Main.lean create mode 100644 MetropolisHastings/Polymorphic.lean create mode 100644 MetropolisHastings/Targets.lean create mode 100644 MetropolisHastings/Theory.lean create mode 100644 scripts/mh_plot.py diff --git a/.gitignore b/.gitignore index e4f6948..a68d3eb 100644 --- a/.gitignore +++ b/.gitignore @@ -29,4 +29,5 @@ *.synctex.gz(busy) *.pdfsync test_data/ -__pycache__/ \ No newline at end of file +__pycache__/ +mh_output/ diff --git a/MetropolisHastings.lean b/MetropolisHastings.lean new file mode 100644 index 0000000..3e554b0 --- /dev/null +++ b/MetropolisHastings.lean @@ -0,0 +1,31 @@ +/- +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 -- shake: keep-all --deprecated_module: ignore + +public import MetropolisHastings.Computable +public import MetropolisHastings.Defs +public import MetropolisHastings.Polymorphic +public import MetropolisHastings.Targets +public import MetropolisHastings.Theory + +/-! +# Random-walk Metropolis–Hastings, written in `rdo`, proved and run + +* `MetropolisHastings.Defs`: the algorithm, as `rdo` programs over the Giry monad. +* `MetropolisHastings.Theory`: they are Markov kernels, satisfy detailed balance with respect to the + target, and leave it invariant after any number of steps. +* `MetropolisHastings.Computable`: the samplers `@[computable]` writes from them. +* `MetropolisHastings.Polymorphic`: the same algorithm, polymorphic in the monad; at `Measure` it + is the one of `Defs`, so the theorems carry over, and at `RandM` it samples. +* `MetropolisHastings.Targets`: two targets to run the chain on, for both routes. + +To run it and draw the plots, from the root of the repository: + +``` +lake exe mh # runs both samplers, checks they agree, writes mh_output/ +python3 scripts/mh_plot.py # checks them against numpy, draws mh_output/*.png +``` +-/ diff --git a/MetropolisHastings/Computable.lean b/MetropolisHastings/Computable.lean new file mode 100644 index 0000000..320cdc5 --- /dev/null +++ b/MetropolisHastings/Computable.lean @@ -0,0 +1,32 @@ +/- +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 MetropolisHastings.Defs +public meta import MetropolisHastings.Defs + +/-! +# Random-walk Metropolis–Hastings, translated into a sampler + +The `@[computable]` attribute reads the programs of `MetropolisHastings.Defs`, written over the +Giry monad, and writes the programs that sample from them: + +* `mhStepComputable : (Float → Float) → Float → Float → RandPCG IO Float`, +* `mhChainComputable : (Float → Float) → Float → ℕ → Float → RandPCG IO Float`. + +The Gaussian proposal becomes `NumLean.normal'`, the Bernoulli draw `bernoulliChoice`, `ℝ` and `ℝ≥0` +become `Float`, and the log-density becomes a function on `Float`. The theorems of +`MetropolisHastings.Theory` are about the programs read here; what is to be trusted is the +translation. +-/ + +@[expose] public section + +namespace MetropolisHastings + +attribute [computable] mhStep mhChain + +end MetropolisHastings diff --git a/MetropolisHastings/Defs.lean b/MetropolisHastings/Defs.lean new file mode 100644 index 0000000..30c9103 --- /dev/null +++ b/MetropolisHastings/Defs.lean @@ -0,0 +1,69 @@ +/- +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 + +/-! +# The random-walk Metropolis–Hastings algorithm, as an `rdo` program + +The target is the measure `π(dx) = exp (logπ x) dx` on `ℝ`, given through its log-density, which +need not be normalised. From the current state `x`, one step of the algorithm proposes +`y ∼ 𝒩(x, s)` and accepts it with probability `min 1 (exp (logπ y - logπ x))`; on rejection the +chain stays at `x`. The chain runs `n` such steps from `x₀`. + +Both programs denote Markov kernels, written over the Giry monad. `MetropolisHastings.Theory` +proves what they satisfy, `MetropolisHastings.Computable` and `MetropolisHastings.Polymorphic` turn +them into programs that run. + +## Main definitions + +* `MetropolisHastings.target logπ`: the measure with density `exp ∘ logπ` against Lebesgue. +* `MetropolisHastings.acceptProb logπ x y`: the probability of accepting a move from `x` to `y`. +* `MetropolisHastings.mhStep logπ s x`: the law of one step of the chain started at `x`. +* `MetropolisHastings.mhChain logπ s n x₀`: the law of the state after `n` steps from `x₀`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory +open scoped NNReal ENNReal + +namespace MetropolisHastings + +variable (logπ : ℝ → ℝ) (s : ℝ≥0) + +/-- The target measure, with density `exp ∘ logπ` against the Lebesgue measure. -/ +noncomputable def target : Measure ℝ := + volume.withDensity fun x ↦ ENNReal.ofReal (Real.exp (logπ x)) + +/-- The probability of accepting a move from `x` to `y`: the ratio of the target densities, capped +at one. -/ +noncomputable def acceptProb (x y : ℝ) : unitInterval := + ⟨min 1 (Real.exp (logπ y - logπ x)), + le_min zero_le_one (Real.exp_nonneg _), min_le_left _ _⟩ + +@[fun_prop] +lemma measurable_acceptProb {γ : Type*} [MeasurableSpace γ] {logπ : ℝ → ℝ} {f g : γ → ℝ} + (hπ : Measurable logπ) (hf : Measurable f) (hg : Measurable g) : + Measurable fun c ↦ acceptProb logπ (f c) (g c) := + Measurable.subtype_mk (by fun_prop) + +/-- One step of random-walk Metropolis–Hastings from `x`, with proposal variance `s`. -/ +noncomputable def mhStep (x : ℝ) : Measure ℝ := rdo + let y ← gaussianReal x s + let accept ← bernoulliMeasure true false (acceptProb logπ x y) + return if accept then y else x + +/-- `n` steps of random-walk Metropolis–Hastings from `x₀`, with proposal variance `s`. -/ +noncomputable def mhChain (n : ℕ) (x₀ : ℝ) : Measure ℝ := rdo + let mut x := x₀ + for _ in List.range n rdo + let y ← mhStep logπ s x + x := y + return x + +end MetropolisHastings diff --git a/MetropolisHastings/Main.lean b/MetropolisHastings/Main.lean new file mode 100644 index 0000000..fa202f8 --- /dev/null +++ b/MetropolisHastings/Main.lean @@ -0,0 +1,142 @@ +/- +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 +-/ +import MetropolisHastings.Targets + +/-! +# Running random-walk Metropolis–Hastings + +`lake exe mh` runs the chain on the targets of `MetropolisHastings.Targets`, through both routes +from the theory to a sampler: the program `@[computable]` wrote, and the polymorphic program run at +`RandM`. Seeded alike, the two draw the same numbers in the same order, so they must produce the +same trajectory, bit for bit; the executable checks it, and that the one-shot chains +`mhChainComputable` and `mhChainPoly` land on the last state of that trajectory. + +It writes, in `mh_output/`: +* `.csv`: the trajectory of each run, each state as the hexadecimal bits of its `Float`, one + column per route; +* `runs.csv`: the parameters of each run, for `scripts/mh_plot.py` to replay them with numpy; +* `acceptance.csv`: the acceptance rate against the proposal standard deviation, for both targets. +-/ + +open MetropolisHastings NumLean + +/-- A run of the chain. -/ +structure Run where + /-- The name of the run, which names its file. -/ + name : String + /-- The name of the target. -/ + target : String + /-- The log-density, as `@[computable]` translated it. -/ + logπA : Float → Float + /-- The log-density, polymorphic, read at `Float`. -/ + logπB : Float → Float + /-- The proposal variance. -/ + var : Float + /-- The number of steps. -/ + steps : Nat + /-- The initial state. -/ + x₀ : Float := 0 + /-- The seed of the generator. -/ + seed : Nat := 42 + +/-- One step of the chain, through the `@[computable]` route. -/ +def Run.stepA (r : Run) (x : Float) : RandPCG IO Float := mhStepComputable r.logπA r.var x + +/-- One step of the chain, through the polymorphic route. -/ +def Run.stepB (r : Run) (x : Float) : RandPCG IO Float := + (mhStepPoly (m := RandM) (R := Float) (V := Float) r.logπB r.var x : RandM Float) + +/-- The whole trajectory of `n` steps of a chain from `x₀`, one state per step. -/ +def trajectory (step : Float → RandPCG IO Float) (n : Nat) (x₀ : Float) : + RandPCG IO (Array Float) := do + let mut x := x₀ + let mut xs := #[x₀] + for _ in [0:n] do + x ← step x + xs := xs.push x + return xs + +/-- The fraction of steps that moved: a proposal is almost surely not the current state, so this is +the fraction of accepted proposals. -/ +def acceptanceRate (xs : Array Float) : Float := Id.run do + let mut moves := 0 + for i in [1:xs.size] do + if xs[i]!.toBits != xs[i - 1]!.toBits then moves := moves + 1 + return moves.toFloat / (xs.size - 1).toFloat + +/-- The bits of a `Float`, in hexadecimal: they are read back exactly. -/ +def hexBits (x : Float) : String := String.ofList (Nat.toDigits 16 x.toBits.toNat) + +/-- The sample mean of the states after `burnIn`. -/ +def mean (xs : Array Float) (burnIn : Nat) : Float := + let ys := xs.extract burnIn xs.size + ys.foldl (· + ·) 0 / ys.size.toFloat + +/-- Run `r` through both routes, check that they agree, and write its trajectory. Returns whether +the checks passed. -/ +def Run.go (r : Run) : IO Bool := do + let a ← (IO.runRandPCGWith r.seed (trajectory r.stepA r.steps r.x₀) : IO (Array Float)) + let b ← (IO.runRandPCGWith r.seed (trajectory r.stepB r.steps r.x₀) : IO (Array Float)) + let chainA ← (IO.runRandPCGWith r.seed (mhChainComputable r.logπA r.var r.steps r.x₀) : IO Float) + let chainB ← (IO.runRandPCGWith r.seed + (mhChainPoly (m := RandM) (R := Float) (V := Float) r.logπB r.var r.steps r.x₀ : RandM Float) : + IO Float) + let sameTrajectory := a.map (·.toBits) == b.map (·.toBits) + let last := a.back!.toBits + let sameChain := chainA.toBits == last && chainB.toBits == last + IO.FS.withFile s!"mh_output/{r.name}.csv" .write fun h ↦ do + h.putStrLn "step,computable,polymorphic" + for i in [0:a.size] do + h.putStrLn s!"{i},{hexBits a[i]!},{hexBits b[i]!}" + IO.println s!"{r.name}: {r.steps} steps, proposal std {r.var.sqrt}, \ + acceptance {acceptanceRate a}, mean after burn-in {mean a (r.steps / 10)}" + IO.println s!" computable = polymorphic, all {a.size} states: {sameTrajectory}; \ + one-shot chains land on the last state: {sameChain}" + return sameTrajectory && sameChain + +/-- The acceptance rate of `steps` steps of the chain, through both routes, for each proposal +standard deviation in `stds`. -/ +def sweep (target : String) (logπA logπB : Float → Float) (stds : List Float) (steps : Nat) : + IO (List String × Bool) := do + let mut lines := [] + let mut ok := true + for std in stds do + let r : Run := { name := "", target, logπA, logπB, var := std * std, steps, seed := 7 } + let a ← (IO.runRandPCGWith r.seed (trajectory r.stepA r.steps r.x₀) : IO (Array Float)) + let b ← (IO.runRandPCGWith r.seed (trajectory r.stepB r.steps r.x₀) : IO (Array Float)) + ok := ok && a.map (·.toBits) == b.map (·.toBits) + lines := lines ++ [s!"{target},{std},{acceptanceRate a},{acceptanceRate b}"] + return (lines, ok) + +def main : IO UInt32 := do + IO.FS.createDirAll "mh_output" + let runs : List Run := [ + { name := "stdnormal", target := "stdNormal", logπA := stdNormalComputable, + logπB := stdNormalPoly, var := 2.4 * 2.4, steps := 50000 }, + { name := "bimodal_small", target := "bimodal", logπA := bimodalComputable, + logπB := bimodalPoly, var := 0.25 * 0.25, steps := 50000 }, + { name := "bimodal_good", target := "bimodal", logπA := bimodalComputable, + logπB := bimodalPoly, var := 3 * 3, steps := 50000 }, + { name := "bimodal_large", target := "bimodal", logπA := bimodalComputable, + logπB := bimodalPoly, var := 30 * 30, steps := 50000 }] + let mut ok := true + for r in runs do + ok := (← r.go) && ok + IO.FS.withFile "mh_output/runs.csv" .write fun h ↦ do + h.putStrLn "name,target,var,steps,x0,seed" + for r in runs do + h.putStrLn s!"{r.name},{r.target},{hexBits r.var},{r.steps},{hexBits r.x₀},{r.seed}" + let stds := (List.range 25).map fun k ↦ 0.05 * Float.pow 1000 (k.toFloat / 24) + let (normalLines, okNormal) ← sweep "stdNormal" stdNormalComputable stdNormalPoly stds 5000 + let (bimodalLines, okBimodal) ← sweep "bimodal" bimodalComputable bimodalPoly stds 5000 + IO.FS.withFile "mh_output/acceptance.csv" .write fun h ↦ do + h.putStrLn "target,std,computable,polymorphic" + for l in normalLines ++ bimodalLines do h.putStrLn l + IO.println s!"acceptance sweep over {stds.length} proposal stds: \ + computable = polymorphic: {okNormal && okBimodal}" + ok := ok && okNormal && okBimodal + IO.println (if ok then "all checks passed" else "SOME CHECKS FAILED") + return if ok then 0 else 1 diff --git a/MetropolisHastings/Polymorphic.lean b/MetropolisHastings/Polymorphic.lean new file mode 100644 index 0000000..e5cb6f8 --- /dev/null +++ b/MetropolisHastings/Polymorphic.lean @@ -0,0 +1,112 @@ +/- +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 MetropolisHastings.Theory + +/-! +# Random-walk Metropolis–Hastings, polymorphic in the monad + +The algorithm of `MetropolisHastings.Defs`, written once over an arbitrary `MeasurableSpaceMonad` +`m` and scalar type `R`. It draws through `HasGaussian` and `HasBernoulli`, and computes with the +operations of `R`. Read at `m := Measure` and `R := ℝ`, it is the program of +`MetropolisHastings.Defs` (`mhStepPoly_measure`, `mhChainPoly_measure`), so it inherits all of +`MetropolisHastings.Theory`. Run at `m := RandM` and `R := Float`, it samples. + +Unlike the `@[computable]` route of `MetropolisHastings.Computable`, no program is written for us +here: the definition that runs is the one the theorems are about, and the only thing to trust is +that the instances at `RandM` sample from the distributions of the instances at `Measure`. + +## Main definitions + +* `mhStepPoly logπ s x`: one step of the chain from `x`. +* `mhChainPoly logπ s n x₀`: `n` steps of the chain from `x₀`. + +## Main results + +* `mhStepPoly_measure`, `mhChainPoly_measure`: at `Measure`, they are `mhStep` and `mhChain`. +* `isReversible_mhStepPoly`, `invariant_mhChainPoly`: detailed balance and stationarity, for the + polymorphic programs. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory +open scoped NNReal + +namespace MetropolisHastings + +universe v + +variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] + {R V : Type} [MeasurableSpace R] [Sub R] [Min R] [One R] [HasExp R] + +/-- One step of random-walk Metropolis–Hastings from `x`, with proposal variance `s`, in any monad +that can draw from a Gaussian and a Bernoulli distribution. -/ +def mhStepPoly [HasGaussian m R V R] [HasBernoulli m R] (logπ : R → R) (s : V) (x : R) : m R := rdo + let y ← HasGaussian.gaussian (m := m) x s + let accept ← HasBernoulli.bernoulli (m := m) (min 1 (HasExp.exp (logπ y - logπ x))) + return if accept then y else x + +/-- `n` steps of random-walk Metropolis–Hastings from `x₀`, with proposal variance `s`, in any +monad that can draw from a Gaussian and a Bernoulli distribution. -/ +def mhChainPoly [HasGaussian m R V R] [HasBernoulli m R] (logπ : R → R) (s : V) (n : ℕ) + (x₀ : R) : m R := rdo + let mut x := x₀ + for _ in List.range n rdo + let y ← mhStepPoly (m := m) logπ s x + x := y + return x + +section Measure + +variable (logπ : ℝ → ℝ) (s : ℝ≥0) + +/-- At `Measure`, the polymorphic step is `mhStep`. -/ +theorem mhStepPoly_measure : mhStepPoly (m := Measure) logπ s = mhStep logπ s := by + ext1 x + change (gaussianReal x s).bind (fun y ↦ (Ber(true, false, + Set.projIcc 0 1 zero_le_one (min 1 (Real.exp (logπ y - logπ x))))).bind + fun b ↦ Measure.dirac (if b = true then y else x)) = _ + congr with y : 1 + have h : Set.projIcc (0 : ℝ) 1 zero_le_one (min 1 (Real.exp (logπ y - logπ x))) + = acceptProb logπ x y := + Set.projIcc_of_mem _ (acceptProb logπ x y).2 + rw [h] + rfl + +/-- At `Measure`, the polymorphic chain is `mhChain`. -/ +theorem mhChainPoly_measure : mhChainPoly (m := Measure) logπ s = mhChain logπ s := by + ext1 n + unfold mhChainPoly mhChain + rw [mhStepPoly_measure] + +variable [hπ : Fact (Measurable logπ)] + +/- `is_markov` reads the polymorphic programs directly, through the instances at `Measure`. -/ + +instance : IsMarkov (mhStepPoly (m := Measure) logπ s) := by + have := hπ.out + is_markov + +instance (n : ℕ) : IsMarkov (mhChainPoly (m := Measure) logπ s n) := by + have := hπ.out + is_markov + +/-- **Detailed balance** of the polymorphic step, read at `Measure`. -/ +theorem isReversible_mhStepPoly : + (IsMarkov.toKernel (mhStepPoly (m := Measure) logπ s)).IsReversible (target logπ) := by + convert isReversible_mhStep logπ s using 2 + exact mhStepPoly_measure logπ s + +/-- **Stationarity** of the polymorphic chain, read at `Measure`. -/ +theorem invariant_mhChainPoly (n : ℕ) : + (target logπ).bind (mhChainPoly (m := Measure) logπ s n) = target logπ := by + rw [mhChainPoly_measure, invariant_mhChain] + +end Measure + +end MetropolisHastings diff --git a/MetropolisHastings/Targets.lean b/MetropolisHastings/Targets.lean new file mode 100644 index 0000000..2b45938 --- /dev/null +++ b/MetropolisHastings/Targets.lean @@ -0,0 +1,80 @@ +/- +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 MetropolisHastings.Computable +public import MetropolisHastings.Polymorphic +public meta import MetropolisHastings.Computable + +/-! +# Two targets to run the chain on + +Each target is written twice, once for each route from the theory to a sampler: on `ℝ`, where +`@[computable]` translates it into a function on `Float`, and polymorphically in the scalars, to be +read at `ℝ` for the theorems and at `Float` for running. The two agree at `ℝ` by definition. + +* `stdNormal`: the standard Gaussian, `logπ x = -x² / 2`. +* `bimodal`: the mixture `0.3 𝒩(-3, 1) + 0.7 𝒩(3, 1)`, up to normalisation. Its modes are far + apart, so how well the chain moves between them depends on the proposal variance. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory +open scoped NNReal + +namespace MetropolisHastings + +/-! ## On `ℝ`, translated by `@[computable]` -/ + +/-- The log-density of the standard Gaussian, up to an additive constant. -/ +@[computable] +noncomputable def stdNormal (x : ℝ) : ℝ := -(x * x) / 2 + +/-- The log-density of the mixture `0.3 𝒩(-3, 1) + 0.7 𝒩(3, 1)`, up to an additive constant. -/ +@[computable] +noncomputable def bimodal (x : ℝ) : ℝ := + Real.log (0.3 * Real.exp (-((x + 3) * (x + 3)) / 2) + 0.7 * Real.exp (-((x - 3) * (x - 3)) / 2)) + +instance : Fact (Measurable stdNormal) := ⟨by unfold stdNormal; fun_prop⟩ + +instance : Fact (Measurable bimodal) := ⟨by unfold bimodal; fun_prop⟩ + +/-! ## Polymorphic in the scalars -/ + +section Polymorphic + +variable {R : Type} [Add R] [Sub R] [Mul R] [Div R] [Neg R] [OfNat R 2] [OfNat R 3] + [OfScientific R] [HasExp R] [HasLog R] + +/-- `stdNormal`, polymorphic in the scalars. -/ +def stdNormalPoly (x : R) : R := -(x * x) / 2 + +/-- `bimodal`, polymorphic in the scalars. -/ +def bimodalPoly (x : R) : R := + HasLog.log (0.3 * HasExp.exp (-((x + 3) * (x + 3)) / 2) + + 0.7 * HasExp.exp (-((x - 3) * (x - 3)) / 2)) + +theorem stdNormalPoly_real : stdNormalPoly (R := ℝ) = stdNormal := rfl + +theorem bimodalPoly_real : bimodalPoly (R := ℝ) = bimodal := rfl + +end Polymorphic + +/-! ## What the theory says about the chains we run -/ + +/-- Started from the bimodal target, the chain run on it has that target as its law after any +number of steps, whatever the proposal variance. -/ +example (s : ℝ≥0) (n : ℕ) : (target bimodal).bind (mhChain bimodal s n) = target bimodal := + invariant_mhChain bimodal s n + +/-- The same, for the polymorphic chain run on the polymorphic target. -/ +example (s : ℝ≥0) (n : ℕ) : + (target bimodal).bind (mhChainPoly (m := Measure) bimodalPoly s n) = target bimodal := by + rw [bimodalPoly_real] + exact invariant_mhChainPoly bimodal s n + +end MetropolisHastings diff --git a/MetropolisHastings/Theory.lean b/MetropolisHastings/Theory.lean new file mode 100644 index 0000000..ce68270 --- /dev/null +++ b/MetropolisHastings/Theory.lean @@ -0,0 +1,271 @@ +/- +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 MetropolisHastings.Defs +public import Mathlib.Probability.Kernel.Invariance + +/-! +# Random-walk Metropolis–Hastings leaves its target invariant + +We prove that one step of the chain, `mhStep logπ s`, satisfies detailed balance with respect to +the target `π(dx) = exp (logπ x) dx`, hence leaves it invariant, and that so does the whole chain +`mhChain logπ s n` for every number of steps `n`: started from the target, the chain has the target +as its law at every step. + +The log-density is assumed measurable, through the instance `Fact (Measurable logπ)`, so that the +programs are Markov kernels: that is the instance `isMarkov_mhStep`, proved by `is_markov`. + +## The proof + +A step proposes `y ∼ Q x` and moves there with probability `p x y`. The mass it sends from a set +`A` into a set `B` is then the *flow* of accepted moves from `A` to `B`, plus the mass of `A ∩ B` +that stays put on rejection. The second term is symmetric in `A` and `B`, so the kernel satisfies +detailed balance as soon as the flow is: that is `isReversible_acceptReject`, for any proposal. + +For the Gaussian proposal, the flow has density `min (π x) (π y) · φ_s(y - x)` against Lebesgue on +`ℝ × ℝ` (`flow_eq`), which is symmetric in `x` and `y`. Swapping the two integrals (Tonelli) gives +the symmetry of the flow (`flow_symm`). + +## Main results + +* `isReversible_acceptReject`: a proposal followed by an accept-reject step satisfies detailed + balance as soon as its flow of accepted moves is symmetric. +* `isReversible_mhStep`: **detailed balance** of `mhStep logπ s` with respect to `target logπ`. +* `invariant_mhStep`: `target logπ` is invariant for `mhStep logπ s`. +* `mhChain_zero`, `mhChain_succ`: the `for` loop of `mhChain` unrolled, one step at a time. +* `invariant_mhChain`: `target logπ` is invariant for `mhChain logπ s n`, for every `n`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory unitInterval Function +open MeasurableSpacePure LawfulMeasurableSpaceMonad +open scoped ENNReal NNReal + +namespace MetropolisHastings + +section AcceptReject + +variable {α : Type*} [MeasurableSpace α] + +lemma bernoulliMeasure_apply_eq (y x : α) (q : I) {B : Set α} (hB : MeasurableSet B) : + Ber(y, x, q) B = toNNReal q * B.indicator 1 y + toNNReal (σ q) * B.indicator 1 x := by + simp only [bernoulliMeasure_def, Measure.add_apply, Measure.smul_apply, + Measure.dirac_apply' _ hB, ENNReal.smul_def, smul_eq_mul] + +/-- The mass a proposal from `Q x` followed by an accept-reject step sends into `B`: the accepted +moves into `B`, and, when `x ∈ B`, the rejected ones. -/ +lemma acceptReject_apply (Q : Kernel α α) {p : α → α → I} (hp : Measurable (uncurry p)) + (x : α) {B : Set α} (hB : MeasurableSet B) : + ((Q x).bind fun y ↦ Ber(y, x, p x y)) B + = ∫⁻ y in B, (toNNReal (p x y) : ℝ≥0∞) ∂Q x + + B.indicator (fun x ↦ ∫⁻ y, (toNNReal (σ (p x y)) : ℝ≥0∞) ∂Q x) x := by + have hpx : Measurable (p x) := hp.comp (measurable_const.prodMk measurable_id) + have hmeas : Measurable fun y ↦ Ber(y, x, p x y) := + (IsMarkov.bernoulliMeasure measurable_id measurable_const hpx).measurable + rw [Measure.bind_apply hB hmeas.aemeasurable] + simp_rw [bernoulliMeasure_apply_eq _ _ _ hB] + rw [lintegral_add_left (by fun_prop), lintegral_mul_const _ (by fun_prop), + ← lintegral_indicator hB] + congr 1 + · congr with y + by_cases hy : y ∈ B <;> simp [hy] + · by_cases hx : x ∈ B <;> simp [hx] + +/-- A proposal from `Q` followed by an accept-reject step with acceptance probability `p` satisfies +detailed balance with respect to `π` as soon as its flow of accepted moves is symmetric. -/ +lemma isReversible_acceptReject {κ : Kernel α α} (Q : Kernel α α) [IsSFiniteKernel Q] + {p : α → α → I} (hp : Measurable (uncurry p)) + (hκ : ∀ x, κ x = (Q x).bind fun y ↦ Ber(y, x, p x y)) {π : Measure α} + (hflow : ∀ ⦃A B⦄, MeasurableSet A → MeasurableSet B → + ∫⁻ x in A, ∫⁻ y in B, (toNNReal (p x y) : ℝ≥0∞) ∂Q x ∂π + = ∫⁻ x in B, ∫⁻ y in A, (toNNReal (p x y) : ℝ≥0∞) ∂Q x ∂π) : + κ.IsReversible π := by + intro A B hA hB + simp_rw [hκ, acceptReject_apply Q hp _ hB, acceptReject_apply Q hp _ hA] + have hrej : Measurable fun x ↦ ∫⁻ y, (toNNReal (σ (p x y)) : ℝ≥0∞) ∂Q x := + Measurable.lintegral_kernel_prod_right' (f := fun z ↦ (toNNReal (σ (p z.1 z.2)) : ℝ≥0∞)) + (by fun_prop) + rw [lintegral_add_right _ (hrej.indicator hB), lintegral_add_right _ (hrej.indicator hA), + hflow hA hB, lintegral_indicator hB, lintegral_indicator hA, Measure.restrict_restrict hB, + Measure.restrict_restrict hA, Set.inter_comm] + +end AcceptReject + +section Gaussian + +variable {logπ : ℝ → ℝ} {s : ℝ≥0} + +instance : IsMarkov fun x : ℝ ↦ gaussianReal x s := + IsMarkov.gaussianReal measurable_id measurable_const + +/-- A step of the chain is a Gaussian proposal followed by an accept-reject step. -/ +lemma mhStep_eq_bind (logπ : ℝ → ℝ) (s : ℝ≥0) (x : ℝ) : + mhStep logπ s x = (gaussianReal x s).bind fun y ↦ Ber(y, x, acceptProb logπ x y) := by + change (gaussianReal x s).bind (fun y ↦ (Ber(true, false, acceptProb logπ x y)).bind + fun b ↦ Measure.dirac (if b = true then y else x)) = _ + congr with y : 1 + rw [Measure.bind_dirac_eq_map _ (by fun_prop), map_bernoulliMeasure] + simp + +lemma coe_toNNReal_eq_ofReal (q : I) : (toNNReal q : ℝ≥0∞) = ENNReal.ofReal q := by + rw [← ENNReal.ofReal_coe_nnreal, coe_toNNReal] + +/-- The density of the flow of accepted moves, `π(dx) 𝒩(x, s)(dy) a(x, y)`, against Lebesgue on +`ℝ × ℝ`: `min (π x) (π y)` times the Gaussian density. -/ +noncomputable def flowDensity (logπ : ℝ → ℝ) (s : ℝ≥0) (x y : ℝ) : ℝ≥0∞ := + ENNReal.ofReal (min (Real.exp (logπ x)) (Real.exp (logπ y))) * gaussianPDF x s y + +lemma flowDensity_comm (x y : ℝ) : flowDensity logπ s x y = flowDensity logπ s y x := by + simp only [flowDensity, gaussianPDF, gaussianPDFReal, min_comm] + congr 4 + ring + +lemma measurable_flowDensity (hπ : Measurable logπ) : Measurable (uncurry (flowDensity logπ s)) := + (by fun_prop : Measurable fun z : ℝ × ℝ ↦ + ENNReal.ofReal (min (Real.exp (logπ z.1)) (Real.exp (logπ z.2)))).mul + (measurable_uncurry_gaussianPDF.comp (measurable_fst.prodMk (measurable_const.prodMk + measurable_snd))) + +/-- The heart of the matter: the density of the target at `x`, times the probability of accepting a +move from `x` to `y`, is the smaller of the densities at `x` and `y`. -/ +lemma exp_mul_acceptProb (x y : ℝ) : + ENNReal.ofReal (Real.exp (logπ x)) * (toNNReal (acceptProb logπ x y) : ℝ≥0∞) + = ENNReal.ofReal (min (Real.exp (logπ x)) (Real.exp (logπ y))) := by + rw [coe_toNNReal_eq_ofReal, ← ENNReal.ofReal_mul (Real.exp_nonneg _)] + congr 1 + simp only [acceptProb] + rw [mul_min_of_nonneg _ _ (Real.exp_nonneg _), mul_one, ← Real.exp_add, add_sub_cancel] + +lemma flow_eq (hπ : Measurable logπ) (hs : s ≠ 0) {A B : Set ℝ} (hA : MeasurableSet A) + (hB : MeasurableSet B) : + ∫⁻ x in A, ∫⁻ y in B, (toNNReal (acceptProb logπ x y) : ℝ≥0∞) ∂gaussianReal x s ∂target logπ + = ∫⁻ x in A, ∫⁻ y in B, flowDensity logπ s x y := by + rw [target, restrict_withDensity hA, lintegral_withDensity_eq_lintegral_mul_non_measurable _ + (by fun_prop) (by simp)] + congr with x + rw [Pi.mul_apply, gaussianReal_of_var_ne_zero _ hs, restrict_withDensity hB, + lintegral_withDensity_eq_lintegral_mul_non_measurable _ (measurable_gaussianPDF _ _) + (by simp [gaussianPDF_lt_top]), ← lintegral_const_mul' _ _ (by simp)] + congr with y + rw [flowDensity, ← exp_mul_acceptProb, Pi.mul_apply] + ring + +/-- With zero variance the proposal is the current state, and the flow from `A` to `B` is the +mass of `A ∩ B`. -/ +lemma flow_zero_var {A B : Set ℝ} (hB : MeasurableSet B) : + ∫⁻ x in A, ∫⁻ y in B, (toNNReal (acceptProb logπ x y) : ℝ≥0∞) ∂gaussianReal x 0 ∂target logπ + = ∫⁻ x in A ∩ B, (toNNReal (acceptProb logπ x x) : ℝ≥0∞) ∂target logπ := by + classical + simp_rw [gaussianReal_zero_var, setLIntegral_dirac] + rw [Set.inter_comm, ← Measure.restrict_restrict hB, ← lintegral_indicator hB] + simp only [Set.indicator_apply] + +lemma flow_symm (hπ : Measurable logπ) {A B : Set ℝ} (hA : MeasurableSet A) + (hB : MeasurableSet B) : + ∫⁻ x in A, ∫⁻ y in B, (toNNReal (acceptProb logπ x y) : ℝ≥0∞) ∂gaussianReal x s ∂target logπ + = ∫⁻ x in B, ∫⁻ y in A, (toNNReal (acceptProb logπ x y) : ℝ≥0∞) ∂gaussianReal x s + ∂target logπ := by + rcases eq_or_ne s 0 with rfl | hs + · rw [flow_zero_var hB, flow_zero_var hA, Set.inter_comm] + rw [flow_eq hπ hs hA hB, flow_eq hπ hs hB hA, + lintegral_lintegral_swap (measurable_flowDensity hπ).aemeasurable] + exact lintegral_congr fun y ↦ lintegral_congr fun x ↦ flowDensity_comm x y + +end Gaussian + +section Theorems + +variable (logπ : ℝ → ℝ) (s : ℝ≥0) [hπ : Fact (Measurable logπ)] + +instance isMarkov_mhStep : IsMarkov (mhStep logπ s) := by + have := hπ.out + is_markov + +instance isMarkov_mhChain (n : ℕ) : IsMarkov (mhChain logπ s n) := by + have := hπ.out + is_markov + +/-- **Detailed balance** of random-walk Metropolis–Hastings: the flow of mass from `A` to `B` under +the target equals the flow from `B` to `A`. -/ +theorem isReversible_mhStep : + (IsMarkov.toKernel (mhStep logπ s)).IsReversible (target logπ) := + isReversible_acceptReject (IsMarkov.toKernel fun x ↦ gaussianReal x s) (p := acceptProb logπ) + (measurable_acceptProb hπ.out measurable_fst measurable_snd) (mhStep_eq_bind logπ s) + fun _ _ hA hB ↦ flow_symm hπ.out hA hB + +/-- The target is invariant for one step of random-walk Metropolis–Hastings. -/ +theorem invariant_mhStep : (target logπ).bind (mhStep logπ s) = target logπ := + (isReversible_mhStep logπ s).invariant + +omit hπ in +/-- A `for` loop whose body does not look at the index only sees the length of the list. -/ +lemma forIn_congr_length {ι : Type} (g : ℝ → Measure (ForInStep ℝ)) : + ∀ (l l' : List ι), l.length = l'.length → ∀ x : ℝ, + MeasurableSpaceForIn.forIn (m := Measure) l x (fun _ ↦ g) + = MeasurableSpaceForIn.forIn (m := Measure) l' x (fun _ ↦ g) + | [], [], _, _ => rfl + | _ :: l, _ :: l', h, x => by + rw [IsMarkov.forIn_cons, IsMarkov.forIn_cons] + refine MeasurableSpaceBind.bind_congr fun step ↦ ?_ + cases step with + | done => rfl + | yield y => exact forIn_congr_length g l l' (by simpa using h) y + +omit hπ in +lemma mhChain_eq_forIn (n : ℕ) (x₀ : ℝ) : + mhChain logπ s n x₀ = MeasurableSpaceForIn.forIn (m := Measure) (List.range n) x₀ + (fun _ z ↦ mhStep logπ s z >>=ₘ fun y ↦ mPure (ForInStep.yield y)) := + mBind_mPure _ + +omit hπ in +lemma mhChain_zero (x₀ : ℝ) : mhChain logπ s 0 x₀ = Measure.dirac x₀ := by + rw [mhChain_eq_forIn, List.range_zero, IsMarkov.forIn_nil] + rfl + +/-- `n + 1` steps of the chain are one step, followed by `n` steps. -/ +lemma mhChain_succ (n : ℕ) (x₀ : ℝ) : + mhChain logπ s (n + 1) x₀ = (mhStep logπ s x₀).bind (mhChain logπ s n) := by + rw [mhChain_eq_forIn, List.range_succ_eq_map, IsMarkov.forIn_cons] + have hk : Measurable fun step : ForInStep ℝ ↦ ForInStep.casesOn (motive := fun _ ↦ Measure ℝ) + step mPure (mhChain logπ s n) := + (ForInStep.measurable_CasesOn (done := fun (_ : Unit) (b : ℝ) ↦ (mPure b : Measure ℝ)) + (yield := fun _ b ↦ mhChain logπ s n b) (by fun_prop) + ((isMarkov_mhChain logπ s n).measurable.comp measurable_snd)).comp + (measurable_const.prodMk measurable_id : Measurable fun step : ForInStep ℝ ↦ ((), step)) + have hloop : ∀ step : ForInStep ℝ, ForInStep.casesOn (motive := fun _ ↦ Measure ℝ) step mPure + (fun b' ↦ MeasurableSpaceForIn.forIn (m := Measure) (List.map Nat.succ (List.range n)) b' + (fun _ z ↦ mhStep logπ s z >>=ₘ fun y ↦ mPure (ForInStep.yield y))) + = ForInStep.casesOn (motive := fun _ ↦ Measure ℝ) step mPure (mhChain logπ s n) := by + rintro (b | b) + · rfl + · exact (forIn_congr_length _ _ _ (by simp) b).trans (mhChain_eq_forIn logπ s n b).symm + rw [MeasurableSpaceBind.bind_congr hloop, mBind_assoc _ (by fun_prop) hk] + refine MeasurableSpaceBind.bind_congr fun y ↦ ?_ + rw [mPure_mBind _ hk] + +/-- **Stationarity**: the target is invariant for `n` steps of random-walk Metropolis–Hastings. -/ +theorem invariant_mhChain (n : ℕ) : (target logπ).bind (mhChain logπ s n) = target logπ := by + induction n with + | zero => + rw [funext (mhChain_zero logπ s)] + exact Measure.bind_dirac + | succ n ih => + rw [funext (mhChain_succ logπ s n), ← Measure.bind_bind + (isMarkov_mhStep logπ s).measurable.aemeasurable + (isMarkov_mhChain logπ s n).measurable.aemeasurable, invariant_mhStep, ih] + +/-- Stationarity, for any multiple of the target: in particular for the normalised target, a +probability measure when the total mass of `target logπ` is finite and positive. Started from it, +the chain has it as its law at every step. -/ +theorem invariant_mhChain_smul (c : ℝ≥0∞) (n : ℕ) : + (c • target logπ).bind (mhChain logπ s n) = c • target logπ := by + rw [Measure.bind_smul, invariant_mhChain] + +end Theorems + +end MetropolisHastings diff --git a/RandomDo/Measurable.lean b/RandomDo/Measurable.lean index 262db34..017c837 100644 --- a/RandomDo/Measurable.lean +++ b/RandomDo/Measurable.lean @@ -23,6 +23,8 @@ This file contains results on the measurable structure of lists, arrays and vect are measurable. * `measurable_of_prodList`: a map out of `δ × List α` is measurable as soon as it is measurable on every stratum, which is how one reasons about a program taking a list as an argument. +* `Measurable.ite_bool`: `if b a then f a else g a` is measurable, for a measurable `b` valued in + `Bool`. -/ @[expose] public section @@ -144,4 +146,12 @@ lemma measurable_isSome : Measurable (Option.isSome : Option α → Bool) := lemma measurable_getD (a : α) : Measurable (fun o : Option α ↦ o.getD a) := measurable_option_iff.2 measurable_id +/-- A choice between two measurable functions on a measurable `Bool`, as in `if b then y else x` +after drawing `b` from a Bernoulli distribution in an `rdo` program. -/ +@[fun_prop] +lemma Measurable.ite_bool {β : Type*} [MeasurableSpace β] {b : α → Bool} {f g : α → β} + (hb : Measurable b) (hf : Measurable f) (hg : Measurable g) : + Measurable fun a ↦ if b a = true then f a else g a := + Measurable.ite (hb (measurableSet_singleton true)) hf hg + end diff --git a/RandomDo/Tactic/Computable/Deriving.lean b/RandomDo/Tactic/Computable/Deriving.lean index 551d299..8ce3af5 100644 --- a/RandomDo/Tactic/Computable/Deriving.lean +++ b/RandomDo/Tactic/Computable/Deriving.lean @@ -83,6 +83,11 @@ partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := let x := xs[0]!.fvarId! withLocalDeclD (← x.getUserName) (← translate σ (← x.getType)) fun y ↦ do mkLambdaFVars #[y] (← translate (σ.insert x y) body) + -- The type of a function a program takes, e.g. a log-density `ℝ → ℝ`. + | .forallE .. => forallBoundedTelescope e (some 1) fun xs body ↦ do + let x := xs[0]!.fvarId! + withLocalDeclD (← x.getUserName) (← translate σ (← x.getType)) fun y ↦ do + mkForallFVars #[y] (← translate (σ.insert x y) body) | .letE n t v b _ => withLetDecl n t v fun x ↦ do withLetDecl n (← translate σ t) (← translate σ v) fun y ↦ do mkLetFVars #[y] (← translate (σ.insert x.fvarId! y) (b.instantiate1 x)) @@ -92,6 +97,10 @@ partial def translate (σ : FVarSubst) (e : Expr) : MetaM Expr := fits, look through it and read its body in its place. -/ partial def translateApp (σ : FVarSubst) (e : Expr) : MetaM Expr := do if e.getAppFn.isLambda then return ← translate σ e.headBeta + -- A function the program takes as an argument, e.g. `logπ y`: its translation is applied to the + -- translated arguments. + if e.getAppFn.isFVar then + return mkAppN (← translate σ e.getAppFn) (← e.getAppArgs.mapM (translate σ)) let .const declName _ := e.getAppFn | throwError "`computable`: cannot translate{indentExpr e}" let counterpart? ← computableAs? declName let head := counterpart?.getD declName diff --git a/RandomDo/Tactic/Computable/Polymorphic.lean b/RandomDo/Tactic/Computable/Polymorphic.lean index f890657..266a3ca 100644 --- a/RandomDo/Tactic/Computable/Polymorphic.lean +++ b/RandomDo/Tactic/Computable/Polymorphic.lean @@ -6,11 +6,19 @@ Authors: Gaëtan Serré module public import Mathlib.Probability.Distributions.Gaussian.Real +public import Mathlib.Probability.Distributions.Bernoulli +public import Mathlib.MeasureTheory.Function.SpecialFunctions.Basic public import RandomDo.Monad.Instances public import RandomDo.NumLean.Distributions /-! +# Polymorphic `rdo` programs +A program written over an arbitrary `MeasurableSpaceMonad` `m`, drawing through the classes of this +file, is read at `m := Measure` to prove things about it and run at `m := RandM` to sample from it. +Each class has an instance of each kind: the distribution of Mathlib on `ℝ`, and the sampler of +`NumLean` on `Float`. The scalar classes `HasExp` and `HasLog` do the same for the functions a +program computes with. -/ @[expose] public section @@ -35,3 +43,40 @@ instance instMeasurableSpaceFloat : MeasurableSpace Float := ⊤ instance : HasGaussian RandM Float Float Float where gaussian μ v := normal' μ v + +/-- A typeclass for monads that can draw from a Bernoulli distribution. -/ +class HasBernoulli (m : (α : Type) → [MeasurableSpace α] → Type v) (P : Type) where + /-- Draw `true` with probability `p`, and `false` otherwise. -/ + bernoulli : P → m Bool + +/-- A probability outside `[0, 1]` is clamped to it. -/ +noncomputable instance : HasBernoulli Measure ℝ where + bernoulli p := bernoulliMeasure true false (Set.projIcc 0 1 zero_le_one p) + +/-- A probability outside `[0, 1]` is clamped to it, as in the instance on `Measure`. -/ +instance : HasBernoulli RandM Float where + bernoulli p := show RandPCG IO Bool from return (← bernoulli (max 0 (min 1 p))) == 1 + +/-- A typeclass for scalars with an exponential. -/ +class HasExp (R : Type) where + /-- The exponential. -/ + exp : R → R + +noncomputable instance : HasExp ℝ := ⟨Real.exp⟩ + +@[fun_prop] +lemma HasExp.measurable_exp_real : Measurable (HasExp.exp : ℝ → ℝ) := Real.measurable_exp + +instance : HasExp Float := ⟨Float.exp⟩ + +/-- A typeclass for scalars with a logarithm. -/ +class HasLog (R : Type) where + /-- The logarithm. -/ + log : R → R + +noncomputable instance : HasLog ℝ := ⟨Real.log⟩ + +@[fun_prop] +lemma HasLog.measurable_log_real : Measurable (HasLog.log : ℝ → ℝ) := Real.measurable_log + +instance : HasLog Float := ⟨Float.log⟩ diff --git a/RandomDo/Tactic/IsMarkov/Elab.lean b/RandomDo/Tactic/IsMarkov/Elab.lean index e7e7244..fa013d8 100644 --- a/RandomDo/Tactic/IsMarkov/Elab.lean +++ b/RandomDo/Tactic/IsMarkov/Elab.lean @@ -136,6 +136,9 @@ def closeLeaf (g : MVarId) : MetaM (List MVarId) := do if let some gs ← observing? (g.applyConst ``IsMarkov.gaussianReal) then trace[is_markov] "`gaussianReal` leaf: handing back the measurability of its parameters" return gs + if let some gs ← observing? (g.applyConst ``IsMarkov.bernoulliMeasure) then + trace[is_markov] "`bernoulliMeasure` leaf: handing back the measurability of its parameters" + return gs return [g] /-- The constant heading the body of `κ`, when it is a definition the tactic could look through. -/ @@ -172,7 +175,9 @@ def abstractLoopVars (vars : Array FVarId) (g : MVarId) : MetaM MVarId := do /-- Turn a goal `IsMarkov κ` into the list of goals the user is left with. -/ partial def isMarkovCore (g : MVarId) (fuel : Nat) : MetaM (List MVarId) := g.withContext do - let target ← instantiateMVars (← g.getType) + /- The annotations a goal may carry, e.g. the one a tactic `have` leaves on the goal after it, + would hide the head of the statement. -/ + let target := (← instantiateMVars (← g.getType)).cleanupAnnotations -- `IsMarkov` takes five arguments: `γ`, `α`, their `MeasurableSpace` instances, and `κ`. unless target.isAppOfArity ``IsMarkov 5 do trace[is_markov] "not an `IsMarkov` goal, handed back: {target}" @@ -251,7 +256,8 @@ partial def isMarkovCore (g : MVarId) (fuel : Nat) : MetaM (List MVarId) := g.wi let mut goals := [] let mut side := [] for g' in gs do - if (← instantiateMVars (← g'.getType)).isAppOfArity ``IsMarkov 5 then + let t := (← instantiateMVars (← g'.getType)).cleanupAnnotations + if t.isAppOfArity ``IsMarkov 5 then goals := goals ++ (← isMarkovCore g' fuel) else side := side ++ [g'] @@ -280,7 +286,7 @@ partial def isMarkovCore (g : MVarId) (fuel : Nat) : MetaM (List MVarId) := g.wi if ← g.isAssigned then return leftover /- The goal was not closed, so we try to unfold names in the head of the program until we reach a known shape. If that fails, we leave the goal to the user. -/ - match ← unfoldToKnownShape (← instantiateMVars (← g.getType)) fuel with + match ← unfoldToKnownShape (← instantiateMVars (← g.getType)).cleanupAnnotations fuel with | some target => trace[is_markov] "unfolded the head definition to: {target.appArg!}" isMarkovCore (← g.change target) (fuel - 1) @@ -316,7 +322,7 @@ lemma _root_.isProbabilityMeasure_of_isMarkov {α : Type*} [MeasurableSpace α] /-- Bring a goal of the form `IsProbabilityMeasure μ` into the form `IsMarkov fun _ : Unit ↦ μ`, so that `isMarkovCore` can be applied. -/ def toIsMarkovGoal (g : MVarId) : MetaM MVarId := do - let target ← instantiateMVars (← g.getType) + let target := (← instantiateMVars (← g.getType)).cleanupAnnotations unless target.isAppOfArity ``IsProbabilityMeasure 3 do return g match ← g.applyConst ``isProbabilityMeasure_of_isMarkov with diff --git a/RandomDo/Tactic/IsMarkov/Lemmas.lean b/RandomDo/Tactic/IsMarkov/Lemmas.lean index 61c21f0..38cc0d0 100644 --- a/RandomDo/Tactic/IsMarkov/Lemmas.lean +++ b/RandomDo/Tactic/IsMarkov/Lemmas.lean @@ -12,6 +12,7 @@ public import RandomDo.Tactic.IsMarkov.ForInStep public import Mathlib.MeasureTheory.Measure.ProbabilityMeasure public import Mathlib.Data.List.OfFn public import Mathlib.Probability.Distributions.Gaussian.Real +public import Mathlib.Probability.Distributions.Bernoulli /-! # Markov property of `rdo` programs @@ -31,6 +32,8 @@ complex program to the Markov property/measurability of its underlying mathemati the bound variable is Markovian in the parameter. * `gaussianReal`: A Gaussian distribution whose mean and variance depend measurably on the parameter is Markovian in the parameter. +* `bernoulliMeasure`: A Bernoulli distribution whose two outcomes and probability depend measurably + on the parameter is Markovian in the parameter. * `comp`: Composing a Markov kernel `κ` with a measurable function `g` is Markovian in the parameter. * `ite`: A conditional `rdo` program that chooses between two Markov kernels `κ` and `η` based on a @@ -43,6 +46,7 @@ complex program to the Markov property/measurability of its underlying mathemati * `forInList_comp`, `forInArray_comp`, `forInVector_comp`: The same three, for a loop over a collection the program takes as an argument. The body is then asked to be Markovian jointly in the parameter and in the element, which the fixed collections do not need. +* `forIn_nil`, `forIn_cons`: A `for` loop over a list, unrolled one element at a time. * `breakRunK`: The case analysis a program performs after a loop that returns early, on the `Option` slot holding the returned value, is Markovian as soon as both of its branches are. -/ @@ -51,6 +55,7 @@ complex program to the Markov property/measurability of its underlying mathemati open MeasureTheory ProbabilityTheory Function open MeasurableSpacePure +open scoped ENNReal namespace IsMarkov @@ -91,6 +96,18 @@ lemma gaussianReal {m : γ → ℝ} {v : γ → NNReal} (hm : Measurable m) (hv IsMarkov fun c ↦ ProbabilityTheory.gaussianReal (m c) (v c) := ⟨ProbabilityTheory.measurable_gaussianReal.comp (hm.prodMk hv), fun _ ↦ inferInstance⟩ +lemma bernoulliMeasure {x y : γ → α} {p : γ → unitInterval} (hx : Measurable x) + (hy : Measurable y) (hp : Measurable p) : + IsMarkov fun c ↦ ProbabilityTheory.bernoulliMeasure (x c) (y c) (p c) := by + refine ⟨Measure.measurable_of_measurable_coe _ fun s hs ↦ ?_, fun _ ↦ inferInstance⟩ + simp only [bernoulliMeasure_def, Measure.add_apply, Measure.smul_apply, + Measure.dirac_apply' _ hs, ENNReal.smul_def, smul_eq_mul] + have hp' : Measurable fun c ↦ ((unitInterval.toNNReal (p c) : ℝ≥0∞)) := by fun_prop + have hq' : Measurable fun c ↦ ((unitInterval.toNNReal (unitInterval.symm (p c)) : ℝ≥0∞)) := by + fun_prop + exact (hp'.mul ((measurable_one.indicator hs).comp hx)).add + (hq'.mul ((measurable_one.indicator hs).comp hy)) + lemma comp {κ : γ → Measure α} (hκ : IsMarkov κ) {g : σ → γ} (hg : Measurable g) : IsMarkov fun c ↦ κ (g c) := ⟨hκ.measurable.comp hg, fun _ ↦ hκ.isProbabilityMeasure _⟩ @@ -228,11 +245,14 @@ private lemma forIn_eq_listLoop (l : List ι) (b : σ) (g : ι → σ → Measur MeasurableSpaceForIn.forIn (m := Measure) l b g = listLoop g l b := loop_eq_listLoop g l b l _ (fun _ _ _ ↦ rfl) ⟨[], rfl⟩ -private lemma forIn_nil (b : σ) (g : ι → σ → Measure (ForInStep σ)) : +/-- A `for` loop over the empty list returns its initial state. -/ +lemma forIn_nil (b : σ) (g : ι → σ → Measure (ForInStep σ)) : MeasurableSpaceForIn.forIn (m := Measure) ([] : List ι) b g = mPure b := forIn_eq_listLoop _ _ _ -private lemma forIn_cons (a : ι) (l : List ι) (b : σ) (g : ι → σ → Measure (ForInStep σ)) : +/-- A `for` loop over `a :: l` runs its body on `a`, then stops or carries on with the loop over +`l`. -/ +lemma forIn_cons (a : ι) (l : List ι) (b : σ) (g : ι → σ → Measure (ForInStep σ)) : MeasurableSpaceForIn.forIn (m := Measure) (a :: l) b g = g a b >>=ₘ fun step ↦ ForInStep.casesOn (motive := fun _ ↦ Measure σ) step mPure fun b' ↦ MeasurableSpaceForIn.forIn (m := Measure) l b' g := by diff --git a/Test/IsMarkov.lean b/Test/IsMarkov.lean index a63e188..c110265 100644 --- a/Test/IsMarkov.lean +++ b/Test/IsMarkov.lean @@ -103,6 +103,12 @@ example : IsProbabilityMeasure layerTwo := by is_markov example : IsProbabilityMeasure layerTwo := by is_markov (fuel := 3) +/-! ## A goal left behind by a tactic `have`, which annotates it -/ + +example (c : ℝ) (hc : 0 < c) : IsMarkov fun x : ℝ ↦ gaussianReal (c * x) 1 := by + have _hc' : 0 ≤ c := hc.le + is_markov + /-! ## The resulting instance is a `Kernel` -/ instance : IsMarkov centred := by is_markov diff --git a/lakefile.toml b/lakefile.toml index 8655135..cf64f1f 100644 --- a/lakefile.toml +++ b/lakefile.toml @@ -24,6 +24,9 @@ name = "RandomDo" [[lean_lib]] name = "Test" +[[lean_lib]] +name = "MetropolisHastings" + # Used to run the tests in `scripts` [[lean_exe]] name = "dump" @@ -37,3 +40,7 @@ root = "Compute" [[lean_exe]] name = "polymorphic" root = "Polymorphic" + +[[lean_exe]] +name = "mh" +root = "MetropolisHastings.Main" diff --git a/scripts/mh_plot.py b/scripts/mh_plot.py new file mode 100644 index 0000000..08db64e --- /dev/null +++ b/scripts/mh_plot.py @@ -0,0 +1,315 @@ +"""Check and plot the Metropolis–Hastings runs of `lake exe mh`. + +Usage, from the root of the repository: + + python3 scripts/mh_plot.py # reads mh_output/, written by `lake exe mh` + python3 scripts/mh_plot.py --run # runs `lake exe mh` first + +It checks that +* the two routes from the theory to a sampler, the program `@[computable]` wrote and the polymorphic + program run at `RandM`, produced the same trajectory, bit for bit; +* a reference implementation in numpy, seeded alike, produces that trajectory too. It draws with + `default_rng(seed).standard_normal` and `.binomial(1, p)`, which `NumLean` reproduces, and computes + in doubles as the Lean program does: the proposal is `fma(sqrt(var), z, x)`, as in + `NumLean.normal`, and `exp` and `log` are libm's, as `Float.exp` and `Float.log`. + +and draws, in mh_output/: +* bimodal_traces.png: the first steps of the chain on the bimodal target, for three proposal sizes; +* bimodal_histograms.png: the states visited against the target density, for the same three; +* bimodal_diagnostics.png: autocorrelation and running mean, for the same three; +* acceptance.png: acceptance rate against proposal size, for both targets; +* stdnormal.png: the chain on the standard Gaussian. +""" + +import csv +import math +import struct +import subprocess +import sys +from fractions import Fraction +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np + +OUT = Path("mh_output") + +# The reference palette of the data-viz guidelines: surfaces, ink, and the first three categorical +# slots, which stay distinguishable pairwise under colour-vision deficiencies. +SURFACE, INK, INK_2, MUTED, GRID, AXIS = ( + "#fcfcfb", "#0b0b0b", "#52514e", "#898781", "#e1e0d9", "#c3c2b7") +BLUE, ORANGE, AQUA = "#2a78d6", "#eb6834", "#1baf7a" + +# The three bimodal runs, keyed by the colour that stands for them in every figure. +BIMODAL_RUNS = [("bimodal_small", BLUE), ("bimodal_good", ORANGE), ("bimodal_large", AQUA)] + + +def from_hex(h): + return struct.unpack(" 0 else (-math.inf if v == 0 else math.nan) + + +def std_normal(x): + return -(x * x) / 2 + + +def bimodal(x): + return flog(0.3 * fexp(-((x + 3) * (x + 3)) / 2) + 0.7 * fexp(-((x - 3) * (x - 3)) / 2)) + + +TARGETS = {"stdNormal": std_normal, "bimodal": bimodal} + + +def std_normal_density(x): + return np.exp(-x * x / 2) / math.sqrt(2 * math.pi) + + +def bimodal_density(x): + phi = lambda m: np.exp(-(x - m) ** 2 / 2) / math.sqrt(2 * math.pi) + return 0.3 * phi(-3) + 0.7 * phi(3) + + +# -- The reference implementation ------------------------------------------------------------------ + +def fma(a, b, c): + """`a * b + c`, rounded once, as `Float.fma`: exact in rationals, then rounded to a double.""" + return float(Fraction(a) * Fraction(b) + Fraction(c)) + + +def mh_reference(logpi, var, steps, x0, seed): + rng = np.random.default_rng(seed) + sd = math.sqrt(var) + x, xs = x0, [x0] + for _ in range(steps): + y = fma(sd, rng.standard_normal(), x) + # `1 ⊓ v` on `Float` is `if 1 ≤ v then 1 else v`, which is Python's `min(1.0, v)`. + p = min(1.0, fexp(logpi(y) - logpi(x))) + if rng.binomial(1, p) == 1: + x = y + xs.append(x) + return xs + + +# -- Loading and checking -------------------------------------------------------------------------- + +def load_runs(): + with open(OUT / "runs.csv") as f: + runs = list(csv.DictReader(f)) + for r in runs: + r["var"], r["x0"] = from_hex(r["var"]), from_hex(r["x0"]) + r["steps"], r["seed"] = int(r["steps"]), int(r["seed"]) + with open(OUT / f"{r['name']}.csv") as f: + rows = list(csv.DictReader(f)) + r["computable"] = [from_hex(row["computable"]) for row in rows] + r["polymorphic"] = [from_hex(row["polymorphic"]) for row in rows] + return runs + + +def bits(xs): + return [struct.pack(" Date: Fri, 18 Sep 2026 16:53:56 +0200 Subject: [PATCH 28/34] Update LeanMachineLearning and port to its new algorithm API MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `lake update LeanMachineLearning` moves LML to dde3322, and Mathlib to 217ba06 with it. LML's rounds now have an observation before the action: `Algorithm 𝓞 𝓐 𝓨`, `Environment 𝓞 𝓐 𝓨`, `IsAlgEnvSeq O A Y alg env P`, with histories `Hist 𝓞 𝓐 𝓨 n = Fin n → Round 𝓞 𝓐 𝓨` and the action at step `n` drawn given the history and the observation `O n`. `AlgTrace` and `alg_env_trace` are ported to it. The first action is now the policy at the empty history, so the separate `K0`/`out0` fields of `AlgTrace` and the step-zero lemmas go: `hasCondDistrib_trace` and `action_ae_eq` hold uniformly in `n`, and the tactic introduces `Ω P O A Y T hseq htr hT hA`. The toy bandit of `Test/AlgTrace.lean` is an `Algorithm Unit (Fin K) ℝ`. Mathlib's `Measure.map` along a non-measurable function is now a Dirac mass, so `IsProbabilityMeasure (μ.map f)` is an instance (`isProbabilityMeasure_map` is gone) and `Measure.bind_smul` asks for measurability; `measurable_pi_lambda` is now `Measurable.of_eval`. An `extend_space` error test relied on `P.map Y` not being known to be a probability measure, and now uses `μ + μ`. Co-Authored-By: Claude Opus 5 (1M context) --- MetropolisHastings/Theory.lean | 3 +- RandomDo/Measurable.lean | 2 +- RandomDo/Probability/AlgTrace.lean | 414 ++++++++++++++--------------- RandomDo/Probability/Trace.lean | 3 +- Test/AlgTrace.lean | 266 +++++++++--------- Test/Extend.lean | 6 +- lake-manifest.json | 10 +- notes/TRACE_SEMANTICS.md | 78 +++--- 8 files changed, 399 insertions(+), 383 deletions(-) diff --git a/MetropolisHastings/Theory.lean b/MetropolisHastings/Theory.lean index ce68270..adefea3 100644 --- a/MetropolisHastings/Theory.lean +++ b/MetropolisHastings/Theory.lean @@ -264,7 +264,8 @@ probability measure when the total mass of `target logπ` is finite and positive the chain has it as its law at every step. -/ theorem invariant_mhChain_smul (c : ℝ≥0∞) (n : ℕ) : (c • target logπ).bind (mhChain logπ s n) = c • target logπ := by - rw [Measure.bind_smul, invariant_mhChain] + rw [Measure.bind_smul _ _ (isMarkov_mhChain logπ s n).measurable.aemeasurable, + invariant_mhChain] end Theorems diff --git a/RandomDo/Measurable.lean b/RandomDo/Measurable.lean index 017c837..2a4a1f5 100644 --- a/RandomDo/Measurable.lean +++ b/RandomDo/Measurable.lean @@ -55,7 +55,7 @@ lemma measurable_ofFn (n : ℕ) : Measurable (List.ofFn : (Fin n → α) → Lis @[fun_prop] lemma measurable_finCons {n : ℕ} : Measurable fun q : α × (Fin n → α) ↦ (Fin.cons q.1 q.2 : Fin (n + 1) → α) := by - refine measurable_pi_lambda _ fun i ↦ ?_ + refine Measurable.of_eval fun i ↦ ?_ refine Fin.cases ?_ (fun j ↦ ?_) i · simp only [Fin.cons_zero]; fun_prop · simp only [Fin.cons_succ]; fun_prop diff --git a/RandomDo/Probability/AlgTrace.lean b/RandomDo/Probability/AlgTrace.lean index c49f21f..8c81518 100644 --- a/RandomDo/Probability/AlgTrace.lean +++ b/RandomDo/Probability/AlgTrace.lean @@ -14,20 +14,20 @@ 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. +`IsAlgEnvSeq O A Y alg env P` says that `O`, `A` and `Y` are the observations, 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: +`K n` for their law at step `n` given the history and the observation, 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`). +* the draws have the conditional law `K n` given the history and the observation, and the action is + `out n` of those 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 @@ -44,9 +44,9 @@ space is existentially quantified precisely because it does not matter. * `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.hasCondDistrib_trace`: the law of the draws. +* `RDo.AlgTrace.action_ae_eq`: the action is the readout of the history, the observation 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. @@ -58,125 +58,130 @@ open MeasureTheory ProbabilityTheory Finset Learning noncomputable section -attribute [fun_prop] Learning.measurable_history measurable_up measurable_down +attribute [fun_prop] 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 +lemma _root_.Learning.IsAlgEnvSeq.comp_measurePreserving {𝓞 𝓐 𝓨 Ω Ω' : Type*} + [MeasurableSpace 𝓞] [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] {_ : MeasurableSpace Ω} + {_ : MeasurableSpace Ω'} {P : Measure Ω} [IsFiniteMeasure P] {P' : Measure Ω'} + [IsFiniteMeasure P'] {O : ℕ → Ω → 𝓞} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} + {alg : Algorithm 𝓞 𝓐 𝓨} {env : Environment 𝓞 𝓐 𝓨} {f : Ω' → Ω} + (h : IsAlgEnvSeq O A Y alg env P) (hf : MeasurePreserving f P' P) : + IsAlgEnvSeq (fun n ω ↦ O n (f ω)) (fun n ω ↦ A n (f ω)) (fun n ω ↦ Y n (f ω)) alg env P' + where + measurable_obs n := (h.measurable_obs n).comp hf.measurable 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_obs n := (h.hasCondDistrib_obs n).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. -/ +around, put `h.measurable_obs`, `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' +lemma _root_.MeasureTheory.MeasurePreserving.transfer_isAlgEnvSeq {𝓞 𝓐 𝓨 Ω Ω' : Type*} + [MeasurableSpace 𝓞] [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] {_ : MeasurableSpace Ω} + {_ : MeasurableSpace Ω'} {P : Measure Ω} [IsFiniteMeasure P] {P' : Measure Ω'} + [IsFiniteMeasure P'] {f : Ω' → Ω} (hf : MeasurePreserving f P' P) {O : ℕ → Ω → 𝓞} + {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} (hO : ∀ n, Measurable (O n)) (hA : ∀ n, Measurable (A n)) + (hY : ∀ n, Measurable (Y n)) {alg : Algorithm 𝓞 𝓐 𝓨} {env : Environment 𝓞 𝓐 𝓨} : + IsAlgEnvSeq O A Y alg env P ↔ + IsAlgEnvSeq (fun n ω ↦ O n (f ω)) (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_obs := hO + 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_obs n := + (hf.hasCondDistrib_fun_comp_iff (measurable_history hO hA hY n) (hO n)).1 + (h.hasCondDistrib_obs n) hasCondDistrib_action n := - (hf.hasCondDistrib_fun_comp_iff (measurable_history hA hY n) (hA (n + 1))).1 - (h.hasCondDistrib_action n) + (hf.hasCondDistrib_fun_comp_iff ((measurable_history hO hA hY n).prodMk (hO n)) + (hA n)).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) } + (hf.hasCondDistrib_fun_comp_iff + (((measurable_history hO hA hY n).prodMk (hO n)).prodMk (hA n)) (hY n)).1 + (h.hasCondDistrib_feedback n) } namespace RDo -universe u uA uY uW +universe u uO uA uY uW -variable {𝓐 : Type uA} {𝓨 : Type uY} {Ω : Type uW} - [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] +variable {𝓞 : Type uO} {𝓐 : Type uA} {𝓨 : Type uY} {Ω : Type uW} + [MeasurableSpace 𝓞] [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) +def forgetTrace (n : ℕ) (h : Hist 𝓞 (Ω × 𝓐) 𝓨 n) : Hist 𝓞 𝓐 𝓨 n := + fun i ↦ ((h i).1, (h i).2.1.2, (h i).2.2) @[fun_prop] lemma measurable_forgetTrace (n : ℕ) : - Measurable (forgetTrace (Ω := Ω) (𝓐 := 𝓐) (𝓨 := 𝓨) n) := - measurable_pi_lambda _ fun _ ↦ by unfold forgetTrace; fun_prop + Measurable (forgetTrace (𝓞 := 𝓞) (Ω := Ω) (𝓐 := 𝓐) (𝓨 := 𝓨) n) := + Measurable.of_eval fun _ ↦ by unfold forgetTrace; fun_prop + +/-- What a traced policy sees at step `n` — the history and the current observation — with the +draws forgotten. -/ +def forgetTraceObs (n : ℕ) (p : Hist 𝓞 (Ω × 𝓐) 𝓨 n × 𝓞) : Hist 𝓞 𝓐 𝓨 n × 𝓞 := + (forgetTrace n p.1, p.2) -omit [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] in +@[fun_prop] +lemma measurable_forgetTraceObs (n : ℕ) : + Measurable (forgetTraceObs (𝓞 := 𝓞) (Ω := Ω) (𝓐 := 𝓐) (𝓨 := 𝓨) n) := by + unfold forgetTraceObs + fun_prop + +omit [MeasurableSpace 𝓞] [MeasurableSpace 𝓐] [MeasurableSpace 𝓨] [MeasurableSpace Ω] in @[simp] -lemma forgetTrace_history {Ω₀ : Type*} {_ : MeasurableSpace Ω₀} (A : ℕ → Ω₀ → Ω × 𝓐) +lemma forgetTrace_history {Ω₀ : Type*} (O : ℕ → Ω₀ → 𝓞) (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 → 𝓐 × 𝓨) Ω + forgetTrace n ∘ history O A Y n = history O (fun n ω ↦ (A n ω).2) Y n := rfl + +/-- A *trace* of an algorithm: one space `Ω` of internal draws, a kernel giving their law at each +step given the history and the observation, 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 and the observation. -/ + K : (n : ℕ) → Kernel (Hist 𝓞 𝓐 𝓨 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 → 𝓐 × 𝓨) × Ω → 𝓐 + /-- The action at step `n`, read off the history, the observation and the draws. -/ + out : (n : ℕ) → (Hist 𝓞 𝓐 𝓨 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 + +attribute [instance] AlgTrace.markov /-- 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 + (env : Environment 𝓞 𝓐 𝓨) : Environment 𝓞 (Ω × 𝓐) 𝓨 where + obs n := (env.obs n).comap (forgetTrace n) (by fun_prop) 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 + (fun p : (Hist 𝓞 (Ω × 𝓐) 𝓨 n × 𝓞) × (Ω × 𝓐) ↦ (forgetTraceObs n p.1, p.2.2)) (by fun_prop) namespace AlgTrace -variable {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} (tr : AlgTrace alg Ω) +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) := + Measurable fun p : (Hist 𝓞 (Ω × 𝓐) 𝓨 n × 𝓞) × Ω ↦ tr.out n (forgetTraceObs n p.1, p.2) := (tr.hasTrace n).measurable_out.comp - (((measurable_forgetTrace n).comp measurable_fst).prodMk measurable_snd) + (((measurable_forgetTraceObs 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 +def algorithm : Algorithm 𝓞 (Ω × 𝓐) 𝓨 where policy n := - (tr.K n).comap (forgetTrace n) (by fun_prop) + (tr.K n).comap (forgetTraceObs 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 α] @@ -186,166 +191,143 @@ lemma _root_.ProbabilityTheory.Kernel.comap_deterministic {α β γ : Type*} [Me 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) := +/-- The traced policy at a given history and observation, with the draws forgotten, is the +original policy. -/ +lemma map_snd_policy_apply (n : ℕ) (h : Hist 𝓞 (Ω × 𝓐) 𝓨 n × 𝓞) : + (tr.algorithm.policy n h).map Prod.snd = alg.policy n (forgetTraceObs n h) := by + have hg : Measurable fun a ↦ tr.out n (forgetTraceObs 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) + have hp : Measurable fun a ↦ (a, tr.out n (forgetTraceObs n h, a)) := measurable_id.prodMk hg + change (((tr.K n).comap (forgetTraceObs n) (measurable_forgetTraceObs 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) + exact (tr.hasTrace n).map_eq (forgetTraceObs 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) := + (tr.algorithm.policy n).snd = (alg.policy n).comap (forgetTraceObs 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 : ℕ → Ω₀ → 𝓨} + {O : ℕ → Ω₀ → 𝓞} {A : ℕ → Ω₀ → Ω × 𝓐} {Y : ℕ → Ω₀ → 𝓨} section Projection -variable (h : IsAlgEnvSeq A Y tr.algorithm (env.withTrace Ω) P) +variable (h : IsAlgEnvSeq O 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 +lemma isAlgEnvSeq_snd : IsAlgEnvSeq O (fun n ω ↦ (A n ω).2) Y alg env P where + measurable_obs n := h.measurable_obs n 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_obs n := + HasCondDistrib.comp_right (hf := measurable_forgetTrace n) (h.hasCondDistrib_obs n) 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 + exact HasCondDistrib.comp_right (hf := measurable_forgetTraceObs 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. -/ +/-- **The conditional law of the draws.** Given the history and the observation, 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 + HasCondDistrib (fun ω ↦ (A n ω).1) + (fun ω ↦ (history O (fun n ω ↦ (A n ω).2) Y n ω, O 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 + exact HasCondDistrib.comp_right (hf := measurable_forgetTraceObs n) h1 -/-- **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 +/-- **The action is the readout of the history, the observation and the draws.** -/ +lemma action_ae_eq [MeasurableEq 𝓐] (n : ℕ) : + (fun ω ↦ (A n ω).2) =ᵐ[P] + fun ω ↦ tr.out n ((history O (fun n ω ↦ (A n ω).2) Y n ω, O n ω), (A n ω).1) := by + have hO := h.measurable_obs 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 + exact ae_eq_of_hasCondDistrib_deterministic (tr.measurable_out n) + (X := fun ω ↦ ((history O A Y n ω, O n ω), (A n ω).1)) + (by fun_prop) (measurable_snd.comp (h.measurable_action n)).aemeasurable h1 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 `Ω`. +the same trajectory law — so anything proved about the law of the observations, 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 +theorem exists_isAlgEnvSeq_trace [MeasurableEq 𝓐] {O₀ : ℕ → Ω₀ → 𝓞} {A₀ : ℕ → Ω₀ → 𝓐} + {Y₀ : ℕ → Ω₀ → 𝓨} (h₀ : IsAlgEnvSeq O₀ A₀ Y₀ alg env P) : + ∃ (Ω' : Type (max u uO uA uY uW)) (_ : MeasurableSpace Ω') (P' : Measure Ω') + (_ : IsProbabilityMeasure P') (O : ℕ → Ω' → 𝓞) (A : ℕ → Ω' → 𝓐) (Y : ℕ → Ω' → 𝓨) + (T : ℕ → Ω' → Ω), + IsAlgEnvSeq O A Y alg env P' + ∧ IsAlgEnvSeq O (fun n ω ↦ (T n ω, A n ω)) Y tr.algorithm (env.withTrace Ω) P' + ∧ P'.map (trajectory O A Y) = P.map (trajectory O₀ A₀ Y₀) + ∧ (∀ n, HasCondDistrib (T n) (fun ω ↦ (history O A Y n ω, O n ω)) (tr.K n) P') + ∧ (∀ n, A n =ᵐ[P'] fun ω ↦ tr.out n ((history O A Y n ω, O n ω), T n ω)) := 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 + let e : ULift.{u} (ℕ → Round 𝓞 (Ω × 𝓐) 𝓨) ≃ᵐ (ℕ → Round 𝓞 (Ω × 𝓐) 𝓨) := + 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) := + have htr : IsAlgEnvSeq (fun n ω ↦ IT.obs n (e ω)) (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 ω), + exact ⟨ULift (ℕ → Round 𝓞 (Ω × 𝓐) 𝓨), inferInstance, base.map e.symm, inferInstance, + fun n ω ↦ IT.obs n (e ω), 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⟩ + isAlgEnvSeq_unique (tr.isAlgEnvSeq_snd htr) h₀, tr.hasCondDistrib_trace 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 -`Ω`. -/ +`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₁) + {motive : (Ω₀ : Type (max u uO uA uY uW)) → [MeasurableSpace Ω₀] → (P : Measure Ω₀) → + [IsProbabilityMeasure P] → (ℕ → Ω₀ → 𝓞) → (ℕ → Ω₀ → 𝓐) → (ℕ → Ω₀ → 𝓨) → Prop} + (traced : ∀ (Ω' : Type (max u uO uA uY uW)) [MeasurableSpace Ω'] (P' : Measure Ω') + [IsProbabilityMeasure P'] (O' : ℕ → Ω' → 𝓞) (A' : ℕ → Ω' → 𝓐) (Y' : ℕ → Ω' → 𝓨) + (T : ℕ → Ω' → Ω), + IsAlgEnvSeq O' A' Y' alg env P' → + IsAlgEnvSeq O' (fun n ω ↦ (T n ω, A' n ω)) Y' tr.algorithm (env.withTrace Ω) P' → + (∀ n, HasCondDistrib (T n) (fun ω ↦ (history O' A' Y' n ω, O' n ω)) (tr.K n) P') → + (∀ n, A' n =ᵐ[P'] fun ω ↦ tr.out n ((history O' A' Y' n ω, O' n ω), T n ω)) → + motive Ω' P' O' A' Y') + (transfer : ∀ (Ω₁ : Type (max u uO uA uY uW)) [MeasurableSpace Ω₁] (P₁ : Measure Ω₁) + [IsProbabilityMeasure P₁] (O₁ : ℕ → Ω₁ → 𝓞) (A₁ : ℕ → Ω₁ → 𝓐) (Y₁ : ℕ → Ω₁ → 𝓨) + (Ω₂ : Type (max u uO uA uY uW)) [MeasurableSpace Ω₂] (P₂ : Measure Ω₂) + [IsProbabilityMeasure P₂] (O₂ : ℕ → Ω₂ → 𝓞) (A₂ : ℕ → Ω₂ → 𝓐) (Y₂ : ℕ → Ω₂ → 𝓨), + IsAlgEnvSeq O₁ A₁ Y₁ alg env P₁ → IsAlgEnvSeq O₂ A₂ Y₂ alg env P₂ → + P₂.map (trajectory O₂ A₂ Y₂) = P₁.map (trajectory O₁ A₁ Y₁) → + motive Ω₂ P₂ O₂ A₂ Y₂ → motive Ω₁ P₁ O₁ 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 + ∀ (Ω₀ : Type (max u uO uA uY uW)) [MeasurableSpace Ω₀] (P : Measure Ω₀) + [IsProbabilityMeasure P] (O : ℕ → Ω₀ → 𝓞) (A : ℕ → Ω₀ → 𝓐) (Y : ℕ → Ω₀ → 𝓨), + IsAlgEnvSeq O A Y alg env P → motive Ω₀ P O A Y := by + intro Ω₀ _ P _ O 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) + obtain ⟨Ω', mΩ', P', hP', O', A', Y', T, hseq, htr, hlaw, hT, hA⟩ := + tr.exists_isAlgEnvSeq_trace.{u, _, _, _, _, _} h + exact transfer Ω₀ P O A Y Ω' P' O' A' Y' h hseq hlaw (traced Ω' P' O' A' Y' T hseq htr hT hA) end AlgTrace @@ -366,13 +348,14 @@ 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 +its σ-algebra, the measure, the `IsProbabilityMeasure` instance, and the three 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 O := as[as.size - 7]! let A := as[as.size - 6]! let Y := as[as.size - 5]! let P := as[as.size - 2]! @@ -392,9 +375,10 @@ def algEnvSpaceFVars (hFVar : FVarId) : MetaM (Array FVarId) := do return none) | throwError "alg_env_trace: no `IsProbabilityMeasure` hypothesis for {P} in the context" out := out.push hP - for e in #[A, Y] do + for e in #[O, A, Y] do let .fvar f := e | throwError - "alg_env_trace: the action and feedback sequences must be local hypotheses, but {e} is not" + "alg_env_trace: the observation, action and feedback sequences must be local hypotheses, \ + but {e} is not" out := out.push f return out.push hFVar @@ -422,9 +406,9 @@ The goal, together with every hypothesis about the sequence — a statement ment 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; + sequence again, `htr`, actions and draws together as a sequence for the traced algorithm, `hT`, + the conditional law of the draws given the history and the observation, and `hA`, each action as + the readout of those 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 @@ -436,11 +420,12 @@ counterpart on the traced space: the goal may not depend on it, and it is cleare 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. +* `alg_env_trace tr with Ω P O A Y T hseq htr hT 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. -/ +three 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 observations, the actions, the feedbacks and +the draws. -/ syntax (name := algEnvTraceTac) "alg_env_trace" ppSpace term (" using " ident)? (" with " (ppSpace colGt ident)+)? : tactic @@ -456,16 +441,17 @@ elab_rules : tactic 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. + -- The trajectory space `ℕ → Round 𝓞 𝓐 𝓨`, 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 #[𝓐, 𝓨]) + let 𝓞 := (← instantiateMVars (← inferType (.fvar spaceFVars[4]!))).getForallBody + let 𝓐 := (← instantiateMVars (← inferType (.fvar spaceFVars[5]!))).getForallBody + let 𝓨 := (← instantiateMVars (← inferType (.fvar spaceFVars[6]!))).getForallBody + let trajTy ← mkArrow (mkConst ``Nat) (← mkAppM ``Prod #[𝓞, ← 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] + let defaults : Array Name := #[`Ω, `P, `O, `A, `Y, `T, `hseq, `htr, `hT, `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 @@ -541,7 +527,7 @@ elab_rules : tactic 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." + the observations, 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))) @@ -557,12 +543,12 @@ elab_rules : tactic 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] + pick 8, pick 9] 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. + -- sequences are measure-preserving maps onto `(ℕ → Round 𝓞 𝓐 𝓨, ν)`, 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 @@ -573,31 +559,31 @@ elab_rules : tactic 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)) + intro Ω₁ _ P₁ _ O₁ A₁ Y₁ Ω₂ _ P₂ _ O₂ A₂ Y₂ h₁ h₂ hlaw h + have hlaw' : (Measure.map (trajectory O₂ A₂ Y₂) P₂).map (ULift.up : _ → $liftStx) + = (Measure.map (trajectory O₁ A₁ Y₁) P₁).map ULift.up := by rw [hlaw] + generalize hν : (Measure.map (trajectory O₁ A₁ Y₁) P₁).map + (ULift.up : _ → $liftStx) = ν at hlaw' + have hf₁ : MeasurePreserving (fun ω ↦ (ULift.up (trajectory O₁ 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)) + MeasurePreserving ULift.up (Measure.map (trajectory O₁ A₁ Y₁) P₁) _).comp + ⟨measurable_trajectory h₁.measurable_obs h₁.measurable_action + h₁.measurable_feedback, rfl⟩ + have hf₂ : MeasurePreserving (fun ω ↦ (ULift.up (trajectory O₂ 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 + MeasurePreserving ULift.up (Measure.map (trajectory O₂ A₂ Y₂) P₂) _).comp + ⟨measurable_trajectory h₂.measurable_obs h₂.measurable_action + h₂.measurable_feedback, rfl⟩ have : IsProbabilityMeasure ν := by rw [← hν] - exact Measure.isProbabilityMeasure_map measurable_up.aemeasurable + infer_instance exact (fun hS : $motiveStx _ inferInstance ν inferInstance - (fun n (t : $liftStx) ↦ (t.down n).1) (fun n (t : $liftStx) ↦ (t.down n).2) ↦ + (fun n (t : $liftStx) ↦ (t.down n).1) (fun n (t : $liftStx) ↦ (t.down n).2.1) + (fun n (t : $liftStx) ↦ (t.down n).2.2) ↦ (by transfer hf₁ at hS; exact hS)) (by beta_reduce; transfer hf₂)))) unless (← getUnsolvedGoals).isEmpty do throwError "transfer left goals" diff --git a/RandomDo/Probability/Trace.lean b/RandomDo/Probability/Trace.lean index 89abcd8..2f7fe2b 100644 --- a/RandomDo/Probability/Trace.lean +++ b/RandomDo/Probability/Trace.lean @@ -143,8 +143,7 @@ lemma isMarkov (h : HasTrace prog P out) [IsMarkovKernel P] : IsMarkov prog wher 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 + infer_instance lemma congr {prog' : γ → Measure β} (h : HasTrace prog P out) (h' : ∀ c, prog' c = prog c) : HasTrace prog' P out := diff --git a/Test/AlgTrace.lean b/Test/AlgTrace.lean index c2055a4..7968783 100644 --- a/Test/AlgTrace.lean +++ b/Test/AlgTrace.lean @@ -14,9 +14,7 @@ down what the tactic does with the rest of the context, what it introduces, and 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. +The algorithm is a bandit algorithm: it sees no observations, so its observation space is `Unit`. -/ open MeasureTheory ProbabilityTheory Finset Learning RDo @@ -31,27 +29,45 @@ universe u variable {K : ℕ} (hK : 0 < K) +/-- The action of the last round of a history, or arm `0` before the first round. -/ +def lastAction (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : Fin K := + if hn : 0 < n then (h ⟨n - 1, by omega⟩).action else ⟨0, hK⟩ + +/-- The feedback of the last round of a history, or `0` before the first round. -/ +def lastFeedback (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : ℝ := + if hn : 0 < n then (h ⟨n - 1, by omega⟩).feedback else 0 + +@[fun_prop] +lemma measurable_lastAction (n : ℕ) : Measurable (lastAction hK n) := by + unfold lastAction + split_ifs <;> fun_prop + +@[fun_prop] +lemma measurable_lastFeedback (n : ℕ) : Measurable (lastFeedback (K := K) n) := by + unfold lastFeedback + split_ifs <;> fun_prop + /-- 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 +def readout (n : ℕ) (p : (Hist Unit (Fin K) ℝ n × Unit) × ℝ) : Fin K := + if 0 < p.2 then ⟨0, hK⟩ else lastAction hK n p.1.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)) + ((measurable_lastAction hK n).comp (measurable_fst.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) +/-- The policy: perturb the last feedback by Gaussian noise, then read the action off it. -/ +def policy (n : ℕ) (p : Hist Unit (Fin K) ℝ n × Unit) : Measure (Fin K) := rdo + let z ← gaussianReal (lastFeedback n p.1) 1 + return readout hK n (p, 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) +def noise (n : ℕ) : Kernel (Hist Unit (Fin K) ℝ n × Unit) ℝ := + markovKernel (fun p ↦ gaussianReal (lastFeedback n p.1) 1) (IsMarkov.gaussianReal (by fun_prop) measurable_const) instance (n : ℕ) : IsMarkovKernel (noise (K := K) n) := by unfold noise; infer_instance @@ -61,92 +77,94 @@ lemma hasTrace_policy (n : ℕ) : HasTrace (policy hK n) (noise n) (readout hK n exact h /-- The algorithm. -/ -def alg : Algorithm (Fin K) ℝ where +def alg : Algorithm Unit (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) : +given the history and the observation, and the action is `readout` of those and the noise. The +trajectory keeps the same law, so anything proved there about the observations, actions and +feedbacks holds of the original sequence. -/ +theorem exists_noise (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type*} [MeasurableSpace Ω₀] + {P : Measure Ω₀} [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} + {Y : ℕ → Ω₀ → ℝ} (h : IsAlgEnvSeq O 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⟩ := + (O' : ℕ → Ω' → Unit) (A' : ℕ → Ω' → Fin K) (Y' : ℕ → Ω' → ℝ) (Z : ℕ → Ω' → ℝ), + IsAlgEnvSeq O' A' Y' (alg hK) env P' + ∧ P'.map (trajectory O' A' Y') = P.map (trajectory O A Y) + ∧ (∀ n, HasCondDistrib (Z n) (fun ω ↦ (history O' A' Y' n ω, O' n ω)) (noise n) P') + ∧ (∀ n, A' n =ᵐ[P'] fun ω ↦ readout hK n ((history O' A' Y' n ω, O' n ω), Z n ω)) := by + obtain ⟨Ω', mΩ', P', hP', O', A', Y', Z, hseq, -, hlaw, hZ, hA⟩ := (trace hK).exists_isAlgEnvSeq_trace h - exact ⟨Ω', mΩ', P', hP', A', Y', Z, hseq, hlaw, hZ, hA⟩ + exact ⟨Ω', mΩ', P', hP', O', 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 +that also carries the noise `Z` the policy draws. The first action is arm `0`: whatever the noise, +the readout is arm `0` before the first round. 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. -/ +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A Y (alg hK) env P) : + ∀ᵐ ω ∂P, A 0 ω = ⟨0, hK⟩ := by + alg_env_trace (trace hK) with Ω P O A Y Z hseq htr hZ hA + -- `Z`, `hZ` and `hA` are the algorithm's draws and their laws, now available. + filter_upwards [hA 0] with ω hω + rw [hω] + simp [trace, readout, lastAction] + +/-- Without `with`, the names are `Ω P O A Y T hseq htr hT hA`. -/ +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A Y (alg hK) env P) : + ∀ᵐ ω ∂P, A 0 ω = ⟨0, hK⟩ := by alg_env_trace (trace hK) using h - guard_hyp hseq : IsAlgEnvSeq A Y (alg hK) env P + guard_hyp hseq : IsAlgEnvSeq O 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 + IsAlgEnvSeq O (fun n ω ↦ (T n ω, A n ω)) Y (trace hK).algorithm (env.withTrace ℝ) P + guard_hyp hT : + ∀ n, HasCondDistrib (T n) (fun ω ↦ (history O A Y n ω, O n ω)) ((trace hK).K n) P + guard_hyp hA : + ∀ n, A n =ᵐ[P] fun ω ↦ (trace hK).out n ((history O A Y n ω, O n ω), T n ω) + filter_upwards [hA 0] with ω hω + rw [hω] + simp [trace, readout, lastAction] /-- **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) : +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O 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ω + filter_upwards [hA 1] with ω hω rw [hω] - by_cases h0 : 0 < T (0 + 1) ω <;> simp [trace, readout, history, h0] + by_cases h0 : 0 < T 1 ω <;> simp [trace, readout, lastAction, 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 +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type u} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A Y (alg hK) env P) : + ∀ᵐ ω ∂P, A 0 ω = ⟨0, hK⟩ := by alg_env_trace (trace hK) - exact hseq.hasLaw_action_zero.map_eq + filter_upwards [hA 0] with ω hω + rw [hω] + simp [trace, readout, lastAction] /-! ## 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) = ν) : +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O 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) = ν @@ -154,36 +172,37 @@ example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] { /-- 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 : Ω₀) +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A Y (alg hK) env P) (X : Ω₀ → ℝ) (_hX : Measurable X) (x : Ω₀) (_hx : X x = 0) : - P.map (A 0) = Measure.dirac ⟨0, hK⟩ := by + ∀ᵐ ω ∂P, A 0 ω = ⟨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 + filter_upwards [hA 0] with ω hω + rw [hω] + simp [trace, readout, lastAction] /-- 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) : +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O 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₀ + intro Ω₁ _ P₁ _ O₁ A₁ Y₁ Ω₂ _ P₂ _ O₂ 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 +def alg2 : Algorithm Unit (Fin K) ℝ where policy _ := Kernel.const _ (Measure.dirac ⟨0, hK⟩) - p0 := Measure.dirac ⟨0, hK⟩ /-- error: alg_env_trace: the goal depends on @@ -193,16 +212,16 @@ of type 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 +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O 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 Ω₀} +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} [IsProbabilityMeasure P] : P Set.univ = 1 := by alg_env_trace (trace hK) @@ -210,7 +229,7 @@ example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] { error: alg_env_trace: hP is not an `IsAlgEnvSeq` hypothesis -/ #guard_msgs in -example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} +example (env : Environment Unit (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 @@ -218,76 +237,79 @@ example (env : Environment (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] { 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 +example (env : Environment Unit (Fin K) ℝ) {P : Measure (ℕ → ℝ)} [IsProbabilityMeasure P] + {O : ℕ → (ℕ → ℝ) → Unit} {A : ℕ → (ℕ → ℝ) → Fin K} {Y : ℕ → (ℕ → ℝ) → ℝ} + (h : IsAlgEnvSeq O A Y (alg hK) env P) : + ∀ᵐ ω ∂P, A 0 ω = ⟨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 +error: alg_env_trace: the observation, 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 +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A (fun n ω ↦ Y n ω + 0) (alg hK) env P) : + ∀ᵐ ω ∂P, A 0 ω = ⟨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 +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A Y (alg2 hK) env P) : + ∀ᵐ ω ∂P, A 0 ω = ⟨0, hK⟩ := by alg_env_trace (trace hK) /-- -error: alg_env_trace: at most 11 names may be given +error: alg_env_trace: at most 10 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 +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A Y (alg hK) env P) : + ∀ᵐ ω ∂P, A 0 ω = ⟨0, hK⟩ := by + alg_env_trace (trace hK) with a b c d e f g i j k l /-! ## `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 +`O`, `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) : +independence statement `hind` covers `O`, `A` and `Y`. -/ +example (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O 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 + (O' : ℕ → Ω' → Unit) (A' : ℕ → Ω' → Fin K) (Y' : ℕ → Ω' → ℝ) (U : Ω' → ℝ), + IsAlgEnvSeq O' A' Y' (alg hK) env P' ∧ HasLaw U (gaussianReal 0 1) P' + ∧ IndepFun (trajectory O' A' Y') U P' := by + have hO := h.measurable_obs 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 hOAY : IndepFun (trajectory O A Y) U P := + hind.comp (φ := fun (p : (ℕ → Unit) × (ℕ → Fin K) × (ℕ → ℝ)) (n : ℕ) ↦ + (p.1 n, p.2.1 n, p.2.2 n)) (by fun_prop) measurable_id + exact ⟨Ω₀, inferInstance, P, inferInstance, O, A, Y, U, h, hU, hOAY⟩ + +/-- **The explicit form, `extend_space_map`.** The goal mentions the space through `P`, `O 0` 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 Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] {P : Measure Ω₀} + [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} {Y : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A Y (alg hK) env P) : + HasCondDistrib (A 0) (O 0) (alg hK).p0 P := by + have hO := h.measurable_obs 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 + exact h.hasCondDistrib_action_zero end Test.AlgTrace diff --git a/Test/Extend.lean b/Test/Extend.lean index fbfe50b..5d66032 100644 --- a/Test/Extend.lean +++ b/Test/Extend.lean @@ -460,11 +460,11 @@ example {Q : Measure Ω} (X : Ω → ℝ) (ν : Measure ℝ) : Q.map X = ν := b extend_space μ /-- -error: extend_space: Measure.map Y P is not known to be a probability measure: no `IsProbabilityMeasure` instance was found +error: extend_space: μ + μ 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) +example (X : Ω → ℝ) (ν : Measure ℝ) : P.map X = ν := by + extend_space (μ + μ) end Test.Extend diff --git a/lake-manifest.json b/lake-manifest.json index 3a05c19..23a3ad9 100644 --- a/lake-manifest.json +++ b/lake-manifest.json @@ -5,7 +5,7 @@ "type": "git", "subDir": null, "scope": "", - "rev": "f86702d123689bc1739b65862132d4625a051843", + "rev": "dde3322b0a3978367029a100be43446ac8d04053", "name": "LeanMachineLearning", "manifestFile": "lake-manifest.json", "inputRev": "main", @@ -15,10 +15,10 @@ "type": "git", "subDir": null, "scope": "leanprover-community", - "rev": "cf65d43b4f5e1a79482e8c488d121853b9d7ca05", + "rev": "217ba069a556ad9de837366b35c2ec1bec1de384", "name": "mathlib", "manifestFile": "lake-manifest.json", - "inputRev": "cf65d43b4f5e1a79482e8c488d121853b9d7ca05", + "inputRev": "217ba069a556ad9de837366b35c2ec1bec1de384", "inherited": true, "configFile": "lakefile.lean"}, {"url": "https://github.com/leanprover/verso", @@ -55,7 +55,7 @@ "type": "git", "subDir": null, "scope": "leanprover-community", - "rev": "d8823026ac7ef130c253089d95685f9877b95323", + "rev": "1681d78dd6e65e38b143f9740d829c826673807c", "name": "importGraph", "manifestFile": "lake-manifest.json", "inputRev": "main", @@ -95,7 +95,7 @@ "type": "git", "subDir": null, "scope": "leanprover-community", - "rev": "d54dddc581e08be364c278052863524bff7a99a9", + "rev": "4cac2177c37f5530c4da76aa8e4307f3fc9e4dcb", "name": "batteries", "manifestFile": "lake-manifest.json", "inputRev": "main", diff --git a/notes/TRACE_SEMANTICS.md b/notes/TRACE_SEMANTICS.md index 217d649..ce6d353 100644 --- a/notes/TRACE_SEMANTICS.md +++ b/notes/TRACE_SEMANTICS.md @@ -249,12 +249,12 @@ trace type cannot be a fixed nest of products. Two shapes work: 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`. + exist as `IsMarkov.forIn_nil` / `IsMarkov.forIn_cons` in + [Lemmas.lean](../RandomDo/Tactic/IsMarkov/Lemmas.lean). 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 +* **`Fin n → α`**, a fixed number `n` of iterations. It is the shape of LML's histories, + `Hist 𝓞 𝓐 𝓨 n = Fin n → Round 𝓞 𝓐 𝓨`, 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 @@ -268,20 +268,25 @@ 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 +LML. `IsAlgEnvSeq O 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. +In LML a round is an observation, then an action, then a feedback, and the policy at step `n` is a +kernel from `Hist 𝓞 𝓐 𝓨 n × 𝓞` — the `n` rounds so far and the current observation — to `𝓐`. +Step `0` is the policy at the empty history; there is no separate initial distribution. + 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: +draws, a kernel `K n` for their law at step `n` given the history and the observation, 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. +* `AlgTrace.hasCondDistrib_trace` — the draws have conditional law `K n` given the history and + the observation; +* `AlgTrace.action_ae_eq` — and the action is `out n` of those 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 @@ -294,12 +299,13 @@ every conclusion about the actions and feedbacks back. `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 — +mentioning the probability space, the measure or the three 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; +* **`traced`**: the same statement on a space that also carries the draws `T`, with `hT` their + conditional law given the history and the observation, and `hA` the equations expressing each + action as the readout of those 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). @@ -307,32 +313,34 @@ That second goal is what makes the move sound rather than a hole: the traced seq 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 => … +example … (h : IsAlgEnvSeq O A Y (alg hK) env P) : ∀ᵐ ω ∂P, A 0 ω = ⟨0, hK⟩ := by + alg_env_trace (trace hK) with Ω P O A Y Z hseq htr hZ hA + filter_upwards [hA 0] with ω hω + … ``` 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) ω) +Ω : Type P : Measure Ω O : ℕ → Ω → Unit A : ℕ → Ω → Fin K Y Z : ℕ → Ω → ℝ +hseq : IsAlgEnvSeq O A Y (alg hK) env P +hZ : ∀ n, HasCondDistrib (Z n) (fun ω ↦ (history O A Y n ω, O n ω)) ((trace hK).K n) P +hA : ∀ n, A n =ᵐ[P] fun ω ↦ (trace hK).out n ((history O A Y n ω, O n ω), Z n ω) ``` `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. +and the three 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. +[Test/AlgTrace.lean](../Test/AlgTrace.lean) runs the whole thing end to end on a toy bandit policy +written as an `rdo` program: `rdo_trace` gives the trace, the `AlgTrace` packages it, +`exists_noise` hands back the noise the policy draws at each step, and the examples drive the +tactic. + +To do the same for `thompson`, its history has to be read as an LML history. `thompson` takes a +`Vector (Fin K × ℝ) n`, and a bandit history is a +`Hist Unit (Fin K) ℝ n = Fin n → Unit × Fin K × ℝ`: both hold `n` action-reward pairs, so the +policy is `thompson` precomposed with the measurable map +`h ↦ Vector.ofFn fun i ↦ ((h i).action, (h i).feedback)`. Everything downstream of that point is +done. From b3414709540af006edfb7a1c02294c5aa7a600c0 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 19 Sep 2026 08:33:17 +0200 Subject: [PATCH 29/34] bandit examples --- .gitignore | 1 + Bandits.lean | 25 + Bandits/Defs.lean | 113 +++++ Bandits/Main.lean | 125 +++++ Bandits/Theory.lean | 482 ++++++++++++++++++++ RandomDo/Tactic/Computable/Polymorphic.lean | 20 +- lakefile.toml | 7 + scripts/bandit_plot.py | 247 ++++++++++ 8 files changed, 1018 insertions(+), 2 deletions(-) create mode 100644 Bandits.lean create mode 100644 Bandits/Defs.lean create mode 100644 Bandits/Main.lean create mode 100644 Bandits/Theory.lean create mode 100644 scripts/bandit_plot.py diff --git a/.gitignore b/.gitignore index a68d3eb..aa8e3ac 100644 --- a/.gitignore +++ b/.gitignore @@ -31,3 +31,4 @@ test_data/ __pycache__/ mh_output/ +bandit_output/ diff --git a/Bandits.lean b/Bandits.lean new file mode 100644 index 0000000..6e436cb --- /dev/null +++ b/Bandits.lean @@ -0,0 +1,25 @@ +/- +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 -- shake: keep-all --deprecated_module: ignore + +public import Bandits.Defs +public import Bandits.Theory + +/-! +# A Gaussian bandit, written in `rdo`, proved and run + +* `Bandits.Defs`: one round of interaction, and `n` rounds, as `rdo` programs polymorphic in the + monad and the scalars; explore-then-commit and UCB as the algorithms choosing the arm. +* `Bandits.Theory`: read at `Measure`, the programs have the law of LeanMachineLearning's + interaction, so its regret bounds hold for them. + +To run them and draw the regret against the bounds, from the root of the repository: + +``` +lake exe bandits # 300 seeds × 5000 rounds; writes bandit_output/ +python3 scripts/bandit_plot.py # checks them against numpy, draws bandit_output/*.png +``` +-/ diff --git a/Bandits/Defs.lean b/Bandits/Defs.lean new file mode 100644 index 0000000..f98e2d8 --- /dev/null +++ b/Bandits/Defs.lean @@ -0,0 +1,113 @@ +/- +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 +public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg + +/-! +# A Gaussian bandit, as an `rdo` program + +`K` arms; pulling arm `a` returns a reward drawn from `𝒩(μ a, σ2)`. A bandit algorithm chooses the +arm to pull from what it has seen so far. The algorithms here only need a summary of it, a +`State`: how many times each arm was pulled, the sum of the rewards each arm returned, and the last +arm pulled. One round of interaction is `banditStep`, and `banditRun n` plays `n` rounds. + +The programs are polymorphic in the monad and in the scalars, as those of +`RandomDo.Tactic.Computable.Polymorphic`: read at `Measure` and `ℝ`, they are what +`Bandits.Theory` proves things about; run at `RandM` and `Float`, they sample. + +The two algorithms, `etcArm` (explore-then-commit) and `ucbArm` (upper confidence bound), mirror +the definitions of `Bandits.ETC.nextArm` and `Bandits.UCB.nextArm` in LeanMachineLearning, with the +statistics of the history in place of sums over it. + +## Main definitions + +* `RDoBandit.HasArgmax`: scalars in which a tuple has an index of its maximum. +* `RDoBandit.State`, `RDoBandit.State.update`: what the algorithms keep of the history. +* `RDoBandit.etcArm`, `RDoBandit.ucbArm`: the arm the two algorithms pull. +* `RDoBandit.banditStep`, `RDoBandit.banditRun`: one round, and `n` rounds, of interaction. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Learning MeasurableSpacePure + +/-- A typeclass for scalars in which a nonempty tuple has an index of its maximum. -/ +class RDoBandit.HasArgmax (R : Type) where + /-- An index at which the tuple is maximal. -/ + argmax {K : ℕ} [NeZero K] : (Fin K → R) → Fin K + +namespace RDoBandit + +/-- At `ℝ`, the `argmax` of LeanMachineLearning: some index at which the tuple is maximal. -/ +noncomputable instance : HasArgmax ℝ := ⟨fun f ↦ argmax f⟩ + +/-- At `Float`, the first index at which the tuple is maximal. It agrees with the instance at `ℝ` +when the maximum is attained once. -/ +instance : HasArgmax Float where + argmax {K} _ f := (List.finRange K).foldl (fun best a ↦ if f best < f a then a else best) 0 + +/-- What a bandit algorithm keeps of the history: the number of pulls of each arm, the sum of the +rewards each arm returned, and the last arm pulled (arm `0` before the first round). The first two +are arrays, so that updating them in a long run stays cheap. -/ +abbrev State (K : ℕ) (R : Type) := Vector ℕ K × Vector R K × Fin K + +variable {K : ℕ} [NeZero K] {R : Type} + +/-- The state before the first round. -/ +def State.init [Zero R] : State K R := (Vector.replicate K 0, Vector.replicate K 0, 0) + +/-- The state after pulling arm `a` and receiving reward `r`. -/ +def State.update [Add R] (s : State K R) (a : Fin K) (r : R) : State K R := + (s.1.set a (s.1[a] + 1), s.2.1.set a (s.2.1[a] + r), a) + +/-- Arm `n % K`: pulling the arms in turn. This is `RoundRobin.nextAction` of LeanMachineLearning, +which is not compiled there. -/ +def roundRobin (K : ℕ) [NeZero K] (n : ℕ) : Fin K := ⟨n % K, Nat.mod_lt _ (NeZero.pos K)⟩ + +variable [Div R] [NatCast R] [HasArgmax R] + +/-- **Explore-then-commit** with `m` pulls of each arm: the arm pulled at time `n`. Arms are +pulled in turn for the first `K * m` rounds; at round `K * m` the algorithm commits to the arm with +the best empirical mean, and pulls it from then on. -/ +def etcArm (m : ℕ) (n : ℕ) (s : State K R) : Fin K := + if n < K * m then roundRobin K n + else if n = K * m then HasArgmax.argmax fun a ↦ s.2.1[a] / (s.1[a] : R) + else s.2.2 + +/-- **UCB** with exploration parameter `c`: the arm pulled at time `n`. Arms are pulled in turn for +the first `K` rounds; afterwards the algorithm pulls the arm maximizing the empirical mean plus +`√(2 c log (n + 1) / N)`, where `N` is the number of pulls of the arm. -/ +def ucbArm [Add R] [Mul R] [One R] [OfNat R 2] [HasLog R] [HasSqrt R] (c : R) (n : ℕ) + (s : State K R) : Fin K := + if n < K then roundRobin K n + else HasArgmax.argmax fun a ↦ + s.2.1[a] / (s.1[a] : R) + HasSqrt.sqrt (2 * c * HasLog.log ((n : R) + 1) / (s.1[a] : R)) + +universe v + +variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] + [Zero R] [Add R] [MeasurableSpace R] {V : Type} [HasGaussian m R V R] + +/-- **One round.** The algorithm `arm` chooses an arm from the state, the arm returns a reward drawn +from `𝒩(μ a, σ2)`, and the state records it. -/ +def banditStep (arm : ℕ → State K R → Fin K) (μ : Fin K → R) (σ2 : V) (n : ℕ) (s : State K R) : + m (State K R) := rdo + let a := arm n s + let r ← HasGaussian.gaussian (m := m) (μ a) σ2 + return s.update a r + +/-- **`n` rounds**, from the initial state. -/ +def banditRun (arm : ℕ → State K R → Fin K) (μ : Fin K → R) (σ2 : V) : ℕ → m (State K R) + | 0 => mPure State.init + | n + 1 => rdo + let s ← banditRun arm μ σ2 n + let s' ← banditStep (m := m) arm μ σ2 n s + return s' + +end RDoBandit diff --git a/Bandits/Main.lean b/Bandits/Main.lean new file mode 100644 index 0000000..42bf16b --- /dev/null +++ b/Bandits/Main.lean @@ -0,0 +1,125 @@ +/- +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 +-/ +import Bandits.Defs + +/-! +# Running the Gaussian bandit + +`lake exe bandits [rounds] [seeds]` runs explore-then-commit and UCB on a three-armed Gaussian +bandit, through the program `banditStep` of `Bandits.Defs` read at `RandM` and `Float`, for many +seeds, and writes the cumulative pseudo-regret in `bandit_output/`: + +* `.csv`: for each round, the mean regret over the seeds, its standard error, and its + 10% and 90% quantiles; +* `_paths.csv`: the regret of the first few seeds, round by round, as the hexadecimal + bits of each `Float`, for `scripts/bandit_plot.py` to replay them with numpy; +* `config.csv`: the parameters, for the same purpose. + +It also checks that the one-shot program `banditRun n` lands on the state the round-by-round run +reaches after `n` rounds. +-/ + +open RDoBandit NumLean + +/-- The number of arms. -/ +def numArms : ℕ := 3 + +instance : NeZero numArms := ⟨by decide⟩ + +/-- The mean reward of each arm. -/ +def means (a : Fin numArms) : Float := #[1.0, 0.5, 0.0][a.val]! + +/-- The gap of each arm: how much is lost, in expectation, by pulling it rather than arm `0`. -/ +def gap (a : Fin numArms) : Float := 1.0 - means a + +/-- The variance of the rewards. -/ +def variance : Float := 1.0 + +/-- A bandit algorithm to run. -/ +structure Algo where + /-- Its name, which names its files. -/ + name : String + /-- The arm it pulls at each round, given the state. -/ + arm : ℕ → State numArms Float → Fin numArms + +/-- One round, at `RandM` and `Float`. -/ +def Algo.step (alg : Algo) (n : ℕ) (s : State numArms Float) : + RandPCG IO (State numArms Float) := + (banditStep (m := RandM) (R := Float) (V := Float) alg.arm means variance n s : RandM _) + +/-- `T` rounds, recording the cumulative pseudo-regret after each. -/ +def Algo.run (alg : Algo) (T : ℕ) : RandPCG IO (Array Float × State numArms Float) := do + let mut s := State.init + let mut regret := 0.0 + let mut curve := Array.mkEmpty T + for n in [0:T] do + s ← alg.step n s + regret := regret + gap s.2.2 + curve := curve.push regret + return (curve, s) + +/-- The bits of a `Float`, in hexadecimal: they are read back exactly. -/ +def hexBits (x : Float) : String := String.ofList (Nat.toDigits 16 x.toBits.toNat) + +/-- The `q`-quantile of a sorted array, by the nearest-rank method. -/ +def quantile (xs : Array Float) (q : Float) : Float := + xs[(q * (xs.size - 1).toFloat).round.toUInt64.toNat]! + +/-- Whether two states are the same, bit for bit. -/ +def sameState (s t : State numArms Float) : Bool := + (List.finRange numArms).all fun a ↦ + s.1[a] == t.1[a] && s.2.1[a].toBits == t.2.1[a].toBits && s.2.2 == t.2.2 + +/-- Run `alg` for `T` rounds on each of `reps` seeds, write its files, and check the one-shot +program on the first seed. -/ +def Algo.go (alg : Algo) (T reps : ℕ) (maxPaths : ℕ := 5) : IO Bool := do + let paths := min maxPaths reps + let mut curves : Array (Array Float) := #[] + for r in [0:reps] do + let (c, _) ← (IO.runRandPCGWith (r + 1) (alg.run T) : IO _) + curves := curves.push c + IO.FS.withFile s!"bandit_output/{alg.name}.csv" .write fun h ↦ do + h.putStrLn "round,mean,se,q10,q90" + for n in [0:T] do + let xs := (curves.map (·[n]!)).qsort (· < ·) + let mean := xs.foldl (· + ·) 0 / reps.toFloat + let var := xs.foldl (fun acc x ↦ acc + (x - mean) * (x - mean)) 0 / (reps - 1).toFloat + h.putStrLn s!"{n + 1},{mean},{(var / reps.toFloat).sqrt},{quantile xs 0.1},{quantile xs 0.9}" + IO.FS.withFile s!"bandit_output/{alg.name}_paths.csv" .write fun h ↦ do + h.putStrLn ("round," ++ ",".intercalate ((List.range paths).map (s!"seed{· + 1}"))) + for n in [0:T] do + h.putStrLn (s!"{n + 1}," ++ ",".intercalate + ((List.range paths).map fun r ↦ hexBits (curves[r]!)[n]!)) + -- The one-shot program, on the first seed: the same draws, in the same order. + let nCheck := min 500 T + let (_, sDriver) ← (IO.runRandPCGWith 1 (alg.run nCheck) : IO _) + let sOneShot ← (IO.runRandPCGWith 1 + (banditRun (m := RandM) (R := Float) (V := Float) alg.arm means variance nCheck : RandM _) : + IO (State numArms Float)) + let ok := sameState sDriver sOneShot + let final := curves.map (·.back!) + IO.println s!"{alg.name}: {reps} seeds × {T} rounds, mean final regret \ + {final.foldl (· + ·) 0 / reps.toFloat}; banditRun {nCheck} = {nCheck} steps: {ok}" + return ok + +/-- `lake exe bandits [rounds] [seeds]`, by default 5000 rounds and 300 seeds. -/ +def main (args : List String) : IO UInt32 := do + IO.FS.createDirAll "bandit_output" + let T := (args[0]?.bind String.toNat?).getD 5000 + let reps := (args[1]?.bind String.toNat?).getD 300 + let algos : List (Algo × String) := [ + ({ name := "etc_m10", arm := etcArm 10 }, "etc,10"), + ({ name := "etc_m50", arm := etcArm 50 }, "etc,50"), + ({ name := "ucb_c3", arm := ucbArm 3 }, "ucb,3")] + IO.FS.withFile "bandit_output/config.csv" .write fun h ↦ do + h.putStrLn "name,algorithm,parameter,rounds,seeds,means,variance" + for (alg, desc) in algos do + h.putStrLn s!"{alg.name},{desc},{T},{reps},1.0;0.5;0.0,{variance}" + let mut ok := true + for (alg, _) in algos do + ok := (← alg.go T reps) && ok + IO.println (if ok then "all checks passed" else "SOME CHECKS FAILED") + return if ok then 0 else 1 diff --git a/Bandits/Theory.lean b/Bandits/Theory.lean new file mode 100644 index 0000000..d735f4b --- /dev/null +++ b/Bandits/Theory.lean @@ -0,0 +1,482 @@ +/- +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 Bandits.Defs +public import LeanMachineLearning.Online.Bandit.Algorithms.Regret.ETC +public import LeanMachineLearning.Online.Bandit.Algorithms.Regret.UCB + +/-! +# The regret of the `rdo` bandit programs + +LeanMachineLearning proves regret bounds for explore-then-commit and UCB, as statements about any +algorithm-environment sequence (`IsAlgEnvSeq`) for its algorithms `etcAlgorithm` and +`ucbAlgorithm`. This file carries them over to the programs of `Bandits.Defs`, read at `Measure`: +the expected pseudo-regret of the state `banditRun n` reaches is bounded as LeanMachineLearning +bounds the expected regret. + +## The bridge + +LeanMachineLearning's algorithms read the whole history; the programs read its *state*, the +number of pulls and the sum of the rewards of each arm and the last arm pulled (`histState`). + +* `etcArm_histState`, `ucbArm_histState`: at `ℝ`, the arm the programs pull from the state of a + history is the arm LeanMachineLearning's algorithms pull from the history. +* `histState_snoc`: the state of a history extended by one round is the state updated by it, + which is what `banditStep` computes. +* `banditRun_eq_map`: **the law of the program**. For any algorithm reading the history through + its state, `banditRun n` is the law of the state of the history of `n` rounds of the + interaction, on LeanMachineLearning's trajectory space. The proof is an induction on `n`: the + history of `n + 1` rounds is that of `n` rounds followed by a round drawn from the step kernel + (`map_hist_succ`), and a round of a deterministic algorithm against Gaussian arms is the arm it + chooses and a Gaussian reward (`map_stepKernel`). +* `integral_pseudoRegret_banditRun`: hence the expected pseudo-regret of the program is the + expected regret of the interaction. + +## Main results + +* `integral_regret_etc_le`: the regret bound of explore-then-commit. +* `integral_regret_ucb_le`, `integral_regret_ucb_le_of_gt`: the regret bounds of UCB, the second + logarithmic in the number of rounds when `c > 2 σ2`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Learning Bandits Finset +open scoped NNReal ENNReal + +namespace RDoBandit + +variable {K : ℕ} [NeZero K] + +section Vector + +variable {α : Type*} [MeasurableSpace α] {n : ℕ} + +@[fun_prop] +lemma measurable_vector_getElem (i : Fin n) : Measurable fun v : Vector α n ↦ v[i] := + (measurable_pi_apply i).comp Vector.measurableEquivTuple.measurable + +lemma measurable_vector_iff {β : Type*} [MeasurableSpace β] {f : β → Vector α n} : + Measurable f ↔ ∀ i : Fin n, Measurable fun b ↦ (f b)[i] := + ⟨fun hf i ↦ (measurable_vector_getElem i).comp hf, + fun h ↦ by + have h' : Measurable fun b ↦ Vector.ofFn fun i : Fin n ↦ (f b)[i] := + Vector.measurableEquivTuple.symm.measurable.comp (measurable_pi_iff.2 h) + simpa using h'⟩ + +@[fun_prop] +lemma measurable_vector_ofFn {β : Type*} [MeasurableSpace β] {f : β → Fin n → α} + (hf : ∀ i, Measurable fun b ↦ f b i) : Measurable fun b ↦ Vector.ofFn (f b) := + measurable_vector_iff.2 fun i ↦ by simpa using hf i + +end Vector + +/-- The last arm of a history, arm `0` before the first round. -/ +noncomputable def lastArm (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : Fin K := + if hn : 0 < n then (h ⟨n - 1, by omega⟩).action else 0 + +/-- The state of a history: the number of pulls of each arm, the sum of its rewards, and the last +arm pulled. -/ +noncomputable def histState (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : State K ℝ := + (Vector.ofFn (pullCount' n h), Vector.ofFn (sumRewards' n h), lastArm n h) + +@[fun_prop] +lemma measurable_lastArm (n : ℕ) : Measurable (lastArm (K := K) n) := by + unfold lastArm + split_ifs <;> fun_prop + +@[fun_prop] +lemma measurable_histState (n : ℕ) : Measurable (histState (K := K) n) := by + unfold histState + refine Measurable.prodMk ?_ (Measurable.prodMk ?_ (measurable_lastArm n)) + · exact measurable_vector_ofFn fun a ↦ measurable_pullCount' n a + · exact measurable_vector_ofFn fun a ↦ measurable_sumRewards' n a + +lemma histState_zero (h : Hist Unit (Fin K) ℝ 0) : histState 0 h = State.init := by + simp only [histState, State.init, lastArm, lt_self_iff_false, dite_false, Prod.mk.injEq, + and_true] + constructor <;> ext <;> simp [pullCount'_eq_sum, sumRewards'] + +omit [NeZero K] in +lemma pullCount'_snoc (n : ℕ) (h : Hist Unit (Fin K) ℝ n) (a b : Fin K) (r : ℝ) : + pullCount' (n + 1) (Fin.snoc h ((), a, r)) b = pullCount' n h b + if a = b then 1 else 0 := by + rw [pullCount'_eq_sum, pullCount'_eq_sum, Fin.sum_univ_castSucc] + simp [Fin.snoc_castSucc, Fin.snoc_last] + +omit [NeZero K] in +lemma sumRewards'_snoc (n : ℕ) (h : Hist Unit (Fin K) ℝ n) (a b : Fin K) (r : ℝ) : + sumRewards' (n + 1) (Fin.snoc h ((), a, r)) b = sumRewards' n h b + if a = b then r else 0 := by + rw [sumRewards', sumRewards', Fin.sum_univ_castSucc] + simp [Fin.snoc_castSucc, Fin.snoc_last] + +lemma histState_snoc (n : ℕ) (h : Hist Unit (Fin K) ℝ n) (a : Fin K) (r : ℝ) : + histState (n + 1) (Fin.snoc h ((), a, r)) = (histState n h).update a r := by + simp only [histState, State.update, Prod.mk.injEq] + refine ⟨?_, ?_, ?_⟩ + · ext i hi + rw [Vector.getElem_ofFn, pullCount'_snoc, Vector.getElem_set] + by_cases hai : (a : ℕ) = i + · obtain rfl : a = ⟨i, hi⟩ := Fin.ext hai + simp + · have : a ≠ ⟨i, hi⟩ := fun h ↦ hai (congrArg Fin.val h) + simp [hai, this] + · ext i hi + rw [Vector.getElem_ofFn, sumRewards'_snoc, Vector.getElem_set] + by_cases hai : (a : ℕ) = i + · obtain rfl : a = ⟨i, hi⟩ := Fin.ext hai + simp + · have : a ≠ ⟨i, hi⟩ := fun h ↦ hai (congrArg Fin.val h) + simp [hai, this] + · simp [lastArm, Fin.snoc, Fin.last] + +lemma etcArm_histState (m n : ℕ) (h : Hist Unit (Fin K) ℝ n) : + etcArm m n (histState n h) = ETC.nextArm K m n h := by + unfold etcArm ETC.nextArm + split_ifs with h1 h2 + · rfl + · simp only [histState, Fin.getElem_fin, Vector.getElem_ofFn, Fin.eta] + rfl + · simp [histState, lastArm, show 0 < n by omega] + +lemma ucbArm_histState (c : ℝ) (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : + ucbArm c n (histState n h) = UCB.nextArm K c n h := by + unfold ucbArm UCB.nextArm + split_ifs + · rfl + · simp only [histState, Fin.getElem_fin, Vector.getElem_ofFn, Fin.eta] + rfl + +section Measures + +open MeasurableSpacePure MeasurableSpaceBind + +variable {α β γ : Type*} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +/-- Mapping a composition-product is binding the kernel, mapped along the section. -/ +lemma map_compProd_eq_bind (ρ : Measure α) [SFinite ρ] (κ : Kernel α β) [IsSFiniteKernel κ] + {F : α × β → γ} (hF : Measurable F) : + (ρ ⊗ₘ κ).map F = ρ.bind fun a ↦ (κ a).map fun b ↦ F (a, b) := by + have hmap (a : α) {t : Set γ} (ht : MeasurableSet t) : + ((κ a).map fun b ↦ F (a, b)) t = κ a (Prod.mk a ⁻¹' (F ⁻¹' t)) := + Measure.map_apply (hF.comp measurable_prodMk_left) ht + have hmeas : Measurable fun a ↦ (κ a).map fun b ↦ F (a, b) := by + refine Measure.measurable_of_measurable_coe _ fun t ht ↦ ?_ + simp_rw [hmap _ ht] + exact Kernel.measurable_kernel_prodMk_left (hF ht) + ext s hs + rw [Measure.map_apply hF hs, Measure.compProd_apply (hF hs), + Measure.bind_apply hs hmeas.aemeasurable] + simp_rw [hmap _ hs] + +lemma dirac_compProd_eq_map [MeasurableSingletonClass α] (a : α) (κ : Kernel α β) + [IsSFiniteKernel κ] : Measure.dirac a ⊗ₘ κ = (κ a).map (Prod.mk a) := by + ext s hs + rw [Measure.dirac_compProd_apply hs, Measure.map_apply measurable_prodMk_left hs] + +/-- Binding after mapping is binding the composite. -/ +lemma bind_map_eq (ρ : Measure α) {f : α → β} (hf : Measurable f) {g : β → Measure γ} + (hg : Measurable g) : (ρ.map f).bind g = ρ.bind (g ∘ f) := by + ext s hs + have hgs : Measurable fun b ↦ g b s := (Measure.measurable_coe hs).comp hg + rw [Measure.bind_apply hs hg.aemeasurable, Measure.bind_apply hs (hg.comp hf).aemeasurable, + lintegral_map hgs hf] + rfl + +end Measures + +section Arms + +variable (μ : Fin K → ℝ) (σ2 : ℝ≥0) + +/-- The arms, as a Markov kernel: arm `a` returns rewards drawn from `𝒩(μ a, σ2)`. -/ +noncomputable def arms : Kernel (Fin K) ℝ where + toFun a := gaussianReal (μ a) σ2 + measurable' := measurable_of_countable _ + +omit [NeZero K] in +lemma arms_apply (a : Fin K) : arms μ σ2 a = gaussianReal (μ a) σ2 := rfl + +instance : IsMarkovKernel (arms μ σ2) := ⟨fun a ↦ by rw [arms_apply]; infer_instance⟩ + +omit [NeZero K] in +lemma integral_arms (a : Fin K) : ∫ x, x ∂(arms μ σ2 a) = μ a := by + simp [arms_apply, integral_id_gaussianReal] + +/-- The gap of an arm: how much less its mean is than the best one. -/ +noncomputable def gapOf (a : Fin K) : ℝ := (⨆ i, μ i) - μ a + +omit [NeZero K] in +lemma gap_arms (a : Fin K) : gap (arms μ σ2) a = gapOf μ a := by + simp [gap, gapOf, integral_arms] + +omit [NeZero K] in +/-- Gaussian rewards are sub-Gaussian, with variance proxy their variance. -/ +lemma hasSubgaussianMGF_arms (a : Fin K) : + HasSubgaussianMGF (fun x ↦ x - (arms μ σ2 a)[id]) σ2 (arms μ σ2 a) := by + rw [show (arms μ σ2 a)[id] = μ a from integral_arms μ σ2 a, arms_apply] + refine ⟨fun t ↦ ?_, fun t ↦ ?_⟩ + · have := (integrable_exp_mul_gaussianReal (μ := μ a) (v := σ2) t).const_mul + (Real.exp (-(t * μ a))) + refine this.congr (Filter.Eventually.of_forall fun x ↦ ?_) + simp only + rw [← Real.exp_add] + ring_nf + · rw [mgf_gaussianReal ⟨by fun_prop, gaussianReal_map_sub_const (μ a)⟩ t] + simp + +end Arms + +section Law + +open MeasurableSpacePure MeasurableSpaceBind + +variable (μ : Fin K → ℝ) (σ2 : ℝ≥0) + +omit [NeZero K] in +lemma getElem_set_fin {α : Type*} (v : Vector α K) (a b : Fin K) (x : α) : + (v.set a x)[b] = if a = b then x else v[b] := by + simp only [Fin.getElem_fin, Vector.getElem_set, Fin.ext_iff] + +omit [NeZero K] in +/-- For a fixed arm, updating the state is measurable in the state and the reward. -/ +lemma measurable_update_const (a : Fin K) : + Measurable fun p : State K ℝ × ℝ ↦ p.1.update a p.2 := by + refine Measurable.prodMk (measurable_vector_iff.2 fun i ↦ ?_) + (Measurable.prodMk (measurable_vector_iff.2 fun i ↦ ?_) measurable_const) + · simp only [getElem_set_fin] + split_ifs + · exact (measurable_of_countable (· + 1 : ℕ → ℕ)).comp (by fun_prop) + · fun_prop + · simp only [getElem_set_fin] + split_ifs <;> fun_prop + +omit [NeZero K] in +@[fun_prop] +lemma Measurable.stateUpdate {X : Type*} [MeasurableSpace X] {f : X → State K ℝ} + {g : X → Fin K} {r : X → ℝ} (hf : Measurable f) (hg : Measurable g) (hr : Measurable r) : + Measurable fun x ↦ (f x).update (g x) (r x) := by + have h : Measurable fun q : (State K ℝ × ℝ) × Fin K ↦ q.1.1.update q.2 q.1.2 := + measurable_from_prod_countable_left fun a ↦ measurable_update_const a + exact h.comp ((hf.prodMk hr).prodMk hg) + +variable {μ σ2} + +omit [NeZero K] in +/-- At `Measure`, one round is the law of the reward of the chosen arm, mapped by the update. -/ +lemma banditStep_eq (arm : ℕ → State K ℝ → Fin K) (n : ℕ) (s : State K ℝ) : + banditStep (m := Measure) arm μ σ2 n s + = (gaussianReal (μ (arm n s)) σ2).map (s.update (arm n s)) := by + change (gaussianReal (μ (arm n s)) σ2).bind (fun r ↦ Measure.dirac (s.update (arm n s) r)) = _ + rw [Measure.bind_dirac_eq_map _ (by fun_prop)] + +lemma banditRun_succ (arm : ℕ → State K ℝ → Fin K) (n : ℕ) : + banditRun (m := Measure) arm μ σ2 (n + 1) + = (banditRun (m := Measure) arm μ σ2 n).bind (banditStep (m := Measure) arm μ σ2 n) := by + rw [banditRun] + rfl + +omit [NeZero K] in +lemma measurable_banditStep {arm : ℕ → State K ℝ → Fin K} {n : ℕ} (harm : Measurable (arm n)) : + Measurable (banditStep (m := Measure) arm μ σ2 n) := by + have hk : IsMarkov fun s ↦ gaussianReal (μ (arm n s)) σ2 := + IsMarkov.gaussianReal ((measurable_of_countable μ).comp harm) measurable_const + have hu : Measurable fun p : State K ℝ × ℝ ↦ p.1.update (arm n p.1) p.2 := by fun_prop + refine Measure.measurable_of_measurable_coe _ fun t ht ↦ ?_ + have hmap (s : State K ℝ) : (banditStep (m := Measure) arm μ σ2 n s) t + = IsMarkov.toKernel (fun s ↦ gaussianReal (μ (arm n s)) σ2) s + (Prod.mk s ⁻¹' ((fun p : State K ℝ × ℝ ↦ p.1.update (arm n p.1) p.2) ⁻¹' t)) := by + rw [banditStep_eq, Measure.map_apply (by fun_prop) ht] + rfl + simp_rw [hmap] + exact Kernel.measurable_kernel_prodMk_left (hu ht) + +omit [NeZero K] in +lemma measurable_snoc {X : Type*} [MeasurableSpace X] (n : ℕ) : + Measurable fun p : (Fin n → X) × X ↦ (Fin.snoc p.1 p.2 : Fin (n + 1) → X) := by + refine Measurable.of_eval fun i ↦ ?_ + refine Fin.lastCases ?_ (fun j ↦ ?_) i + · simp only [Fin.snoc_last] + exact measurable_snd + · simp only [Fin.snoc_castSucc] + exact (measurable_pi_apply j).comp measurable_fst + +omit [NeZero K] in +lemma hist_succ_eq (n : ℕ) : + IT.hist (𝓞 := Unit) (𝓐 := Fin K) (𝓨 := ℝ) (n + 1) + = fun ω ↦ Fin.snoc (IT.hist n ω) (IT.step n ω) := by + funext ω i + refine Fin.lastCases ?_ (fun j ↦ ?_) i + · simp [IT.hist, IT.step] + · simp [IT.hist] + +omit [NeZero K] in +/-- The history after `n + 1` rounds is the history after `n` rounds, followed by one round drawn +from the step kernel. -/ +lemma map_hist_succ {γ : Type*} [MeasurableSpace γ] (alg : Algorithm Unit (Fin K) ℝ) + (env : Environment Unit (Fin K) ℝ) (n : ℕ) {F : Hist Unit (Fin K) ℝ (n + 1) → γ} + (hF : Measurable F) : + (trajMeasure alg env).map (F ∘ IT.hist (n + 1)) + = ((trajMeasure alg env).map (IT.hist n)).bind + fun h ↦ (stepKernel alg env n h).map fun x ↦ F (Fin.snoc h x) := by + have e : F ∘ IT.hist (n + 1) + = (fun p ↦ F (Fin.snoc p.1 p.2)) ∘ (fun ω ↦ (IT.hist n ω, IT.step n ω)) := by + rw [hist_succ_eq] + rfl + rw [e, ← Measure.map_map (g := fun p ↦ F (Fin.snoc p.1 p.2)) + (f := fun ω ↦ (IT.hist n ω, IT.step n ω)) (hF.comp (measurable_snoc n)) (by fun_prop), + (IT.hasCondDistrib_step alg env n).map_eq, + map_compProd_eq_bind (F := fun p ↦ F (Fin.snoc p.1 p.2)) _ _ (hF.comp (measurable_snoc n))] + +variable {nextA : (n : ℕ) → Hist Unit (Fin K) ℝ n × Unit → Fin K} + {hnext : ∀ n, Measurable (nextA n)} + +omit [NeZero K] in +/-- One round of a deterministic algorithm against the Gaussian arms: the chosen arm, then its +reward. -/ +lemma map_stepKernel {γ : Type*} [MeasurableSpace γ] (n : ℕ) (h : Hist Unit (Fin K) ℝ n) + {G : Round Unit (Fin K) ℝ → γ} (hG : Measurable G) : + (stepKernel (detAlgorithm nextA hnext) (stationaryEnv (arms μ σ2)) n h).map G + = (gaussianReal (μ (nextA n (h, ()))) σ2).map fun r ↦ G ((), nextA n (h, ()), r) := by + rw [stepKernel_stationaryEnv, Kernel.compProd_apply_eq_compProd_sectR, Kernel.const_apply, + Measure.dirac_unit_compProd, Kernel.sectR_apply, Kernel.compProd_apply_eq_compProd_sectR, + detAlgorithm_policy, Kernel.deterministic_apply, dirac_compProd_eq_map, Kernel.sectR_apply, + Kernel.prodMkLeft_apply, arms_apply, Measure.map_map hG measurable_prodMk_left, + Measure.map_map (hG.comp measurable_prodMk_left) measurable_prodMk_left] + rfl + +/-- **The law of the program.** For an algorithm that reads the history only through its state, +the state after `n` rounds of `banditRun` has the law of the state of the history of `n` rounds of +LeanMachineLearning's algorithm-environment interaction. -/ +theorem banditRun_eq_map (arm : ℕ → State K ℝ → Fin K) (harm_meas : ∀ n, Measurable (arm n)) + (harm : ∀ n h, arm n (histState n h) = nextA n (h, ())) (n : ℕ) : + banditRun (m := Measure) arm μ σ2 n + = (trajMeasure (detAlgorithm nextA hnext) (stationaryEnv (arms μ σ2))).map + (histState n ∘ IT.hist n) := by + induction n with + | zero => + have : histState 0 ∘ IT.hist (𝓞 := Unit) (𝓐 := Fin K) (𝓨 := ℝ) 0 = fun _ ↦ State.init := by + funext ω + exact histState_zero _ + rw [this, Measure.map_const, measure_univ, one_smul] + rfl + | succ n ih => + rw [banditRun_succ, ih, map_hist_succ _ _ n (measurable_histState (n + 1)), + ← Measure.map_map (measurable_histState n) (IT.measurable_hist n), + bind_map_eq _ (measurable_histState n) (measurable_banditStep (harm_meas n))] + congr 1 + funext h + have hG : Measurable fun x ↦ histState (n + 1) (Fin.snoc h x) := + (measurable_histState (n + 1)).comp ((measurable_snoc n).comp + (measurable_const.prodMk measurable_id)) + rw [Function.comp_apply, banditStep_eq, map_stepKernel n h hG, harm] + simp_rw [histState_snoc] + +end Law + +section Regret + +variable (μ : Fin K → ℝ) (σ2 : ℝ≥0) + +/-- The pseudo-regret of a state: the number of pulls of each arm, times its gap. -/ +noncomputable def pseudoRegret (s : State K ℝ) : ℝ := ∑ a, ((s.1[a] : ℕ) : ℝ) * gapOf μ a + +omit [NeZero K] in +@[fun_prop] +lemma measurable_pseudoRegret : Measurable (pseudoRegret (K := K) μ) := by + unfold pseudoRegret + refine Finset.measurable_sum _ fun a _ ↦ ?_ + exact ((measurable_of_countable (fun k : ℕ ↦ (k : ℝ))).comp + ((measurable_vector_getElem a).comp measurable_fst)).mul_const _ + +/-- The regret of the interaction is the pseudo-regret of the state of its history. -/ +lemma regret_eq_pseudoRegret (n : ℕ) (ω : ℕ → Round Unit (Fin K) ℝ) : + regret (arms μ σ2) IT.action n ω = pseudoRegret μ (histState n (IT.hist n ω)) := by + rw [regret_eq_sum_pullCount_mul_gap] + simp only [pseudoRegret, histState, Fin.getElem_fin, Vector.getElem_ofFn, gap_arms] + congr with a + rw [pullCount_eq_pullCount' (O := IT.obs) (R' := IT.feedback), IT.history_obs_action_feedback] + +variable {nextA : (n : ℕ) → Hist Unit (Fin K) ℝ n × Unit → Fin K} + {hnext : ∀ n, Measurable (nextA n)} + +/-- **The expected pseudo-regret of the program is the expected regret of the interaction.** -/ +theorem integral_pseudoRegret_banditRun (arm : ℕ → State K ℝ → Fin K) + (harm_meas : ∀ n, Measurable (arm n)) (harm : ∀ n h, arm n (histState n h) = nextA n (h, ())) + (n : ℕ) : + ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) arm μ σ2 n) + = (trajMeasure (detAlgorithm nextA hnext) (stationaryEnv (arms μ σ2)))[ + regret (arms μ σ2) IT.action n] := by + rw [banditRun_eq_map (hnext := hnext) arm harm_meas harm n, + integral_map (by fun_prop) (measurable_pseudoRegret μ).aestronglyMeasurable] + congr with ω + exact (regret_eq_pseudoRegret μ σ2 n ω).symm + +omit [NeZero K] in +@[fun_prop] +lemma measurable_hasArgmax_real [NeZero K] : + Measurable (HasArgmax.argmax : (Fin K → ℝ) → Fin K) := measurable_argmax + +omit [NeZero K] in +@[fun_prop] +lemma measurable_count (a : Fin K) : Measurable fun s : State K ℝ ↦ ((s.1[a] : ℕ) : ℝ) := + (measurable_of_countable (fun k : ℕ ↦ (k : ℝ))).comp + ((measurable_vector_getElem a).comp measurable_fst) + +lemma measurable_etcArm (m n : ℕ) : Measurable (etcArm (K := K) (R := ℝ) m n) := by + unfold etcArm + split_ifs + · exact measurable_const + · exact measurable_hasArgmax_real.comp (measurable_pi_iff.2 fun a ↦ by fun_prop) + · fun_prop + +lemma measurable_ucbArm (c : ℝ) (n : ℕ) : Measurable (ucbArm (K := K) (R := ℝ) c n) := by + unfold ucbArm + split_ifs + · exact measurable_const + · exact measurable_hasArgmax_real.comp (measurable_pi_iff.2 fun a ↦ by fun_prop) + +/-- **Regret of explore-then-commit.** The bound of `Bandits.ETC.regret_le`, for the program. -/ +theorem integral_regret_etc_le {m : ℕ} (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : + ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) (etcArm m) μ σ2 n) + ≤ ∑ a, gapOf μ a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gapOf μ a ^ 2 / (4 * σ2))) := by + have h := ETC.regret_le + (IT.isAlgEnvSeq_trajMeasure (etcAlgorithm K m) (stationaryEnv (arms μ σ2))) + (hasSubgaussianMGF_arms μ σ2) hm n hn + simp only [gap_arms] at h + rw [integral_pseudoRegret_banditRun μ σ2 (hnext := fun n ↦ ETC.measurable_nextArm m n |>.comp + measurable_fst) (etcArm m) (measurable_etcArm m) (fun n h ↦ etcArm_histState m n h) n] + exact h + +/-- **Regret of UCB.** The bound of `Bandits.UCB.regret_le'`, for the program. -/ +theorem integral_regret_ucb_le {c : ℝ} (hc : 0 < c) (hσ2 : σ2 ≠ 0) (n : ℕ) : + ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) (ucbArm c) μ σ2 n) + ≤ ∑ a, (8 * c * Real.log (n + 1) / gapOf μ a + + gapOf μ a * (2 + 2 * UCB.constSum (c / σ2) n)) := by + have h := UCB.regret_le' + (IT.isAlgEnvSeq_trajMeasure (ucbAlgorithm K c) (stationaryEnv (arms μ σ2))) + (hasSubgaussianMGF_arms μ σ2) hσ2 hc n + simp only [gap_arms] at h + rw [integral_pseudoRegret_banditRun μ σ2 (hnext := fun n ↦ UCB.measurable_nextArm c n |>.comp + measurable_fst) (ucbArm c) (measurable_ucbArm c) (fun n h ↦ ucbArm_histState c n h) n] + exact h + +/-- **Regret of UCB**, for `c > 2 σ2`: logarithmic in the number of rounds. -/ +theorem integral_regret_ucb_le_of_gt {c : ℝ} (hc : 2 * σ2 < c) (hσ2 : σ2 ≠ 0) (n : ℕ) : + ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) (ucbArm c) μ σ2 n) + ≤ ∑ a, (8 * c * Real.log (n + 1) / gapOf μ a + + gapOf μ a * (4 + 2 * σ2 / (c - 2 * σ2))) := by + have h := UCB.regret_le_of_gt_two' + (IT.isAlgEnvSeq_trajMeasure (ucbAlgorithm K c) (stationaryEnv (arms μ σ2))) + (hasSubgaussianMGF_arms μ σ2) hσ2 hc n + simp only [gap_arms] at h + rw [integral_pseudoRegret_banditRun μ σ2 (hnext := fun n ↦ UCB.measurable_nextArm c n |>.comp + measurable_fst) (ucbArm c) (measurable_ucbArm c) (fun n h ↦ ucbArm_histState c n h) n] + exact h + +end Regret + +end RDoBandit diff --git a/RandomDo/Tactic/Computable/Polymorphic.lean b/RandomDo/Tactic/Computable/Polymorphic.lean index 266a3ca..f036302 100644 --- a/RandomDo/Tactic/Computable/Polymorphic.lean +++ b/RandomDo/Tactic/Computable/Polymorphic.lean @@ -17,8 +17,8 @@ public import RandomDo.NumLean.Distributions A program written over an arbitrary `MeasurableSpaceMonad` `m`, drawing through the classes of this file, is read at `m := Measure` to prove things about it and run at `m := RandM` to sample from it. Each class has an instance of each kind: the distribution of Mathlib on `ℝ`, and the sampler of -`NumLean` on `Float`. The scalar classes `HasExp` and `HasLog` do the same for the functions a -program computes with. +`NumLean` on `Float`. The scalar classes `HasExp`, `HasLog` and `HasSqrt` do the same for the +functions a program computes with. -/ @[expose] public section @@ -80,3 +80,19 @@ noncomputable instance : HasLog ℝ := ⟨Real.log⟩ lemma HasLog.measurable_log_real : Measurable (HasLog.log : ℝ → ℝ) := Real.measurable_log instance : HasLog Float := ⟨Float.log⟩ + +/-- A typeclass for scalars with a square root. -/ +class HasSqrt (R : Type) where + /-- The square root. -/ + sqrt : R → R + +noncomputable instance : HasSqrt ℝ := ⟨Real.sqrt⟩ + +@[fun_prop] +lemma HasSqrt.measurable_sqrt_real : Measurable (HasSqrt.sqrt : ℝ → ℝ) := + Real.continuous_sqrt.measurable + +instance : HasSqrt Float := ⟨Float.sqrt⟩ + +/-- A natural number as a `Float`, so that programs polymorphic in the scalars can cast counts. -/ +instance instNatCastFloat : NatCast Float := ⟨Nat.toFloat⟩ diff --git a/lakefile.toml b/lakefile.toml index cf64f1f..e5a9ab7 100644 --- a/lakefile.toml +++ b/lakefile.toml @@ -27,6 +27,9 @@ name = "Test" [[lean_lib]] name = "MetropolisHastings" +[[lean_lib]] +name = "Bandits" + # Used to run the tests in `scripts` [[lean_exe]] name = "dump" @@ -44,3 +47,7 @@ root = "Polymorphic" [[lean_exe]] name = "mh" root = "MetropolisHastings.Main" + +[[lean_exe]] +name = "bandits" +root = "Bandits.Main" diff --git a/scripts/bandit_plot.py b/scripts/bandit_plot.py new file mode 100644 index 0000000..604fa0f --- /dev/null +++ b/scripts/bandit_plot.py @@ -0,0 +1,247 @@ +"""Check and plot the Gaussian bandit runs of `lake exe bandits`. + +Usage, from the root of the repository: + + python3 scripts/bandit_plot.py # reads bandit_output/, written by `lake exe bandits` + python3 scripts/bandit_plot.py --run # runs `lake exe bandits` first + +It checks that a reference implementation in numpy, seeded alike, produces the regret of the first +seeds of each algorithm, bit for bit. It draws with `default_rng(seed).standard_normal`, which +`NumLean` reproduces, and computes in doubles as the Lean program does: the reward is +`fma(sqrt(variance), z, mean)`, as in `NumLean.normal`, the argmax is the first maximal index, as +the `Float` instance of `HasArgmax`, and `sqrt` and `log` are libm's. + +It draws, in bandit_output/: +* bandit_regret.png: the regret of each algorithm against the bound proved for it in + `Bandits.Theory`; +* bandit_compare.png: the three algorithms together, and single runs of explore-then-commit. +""" + +import csv +import math +import struct +import subprocess +import sys +from fractions import Fraction +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np + +OUT = Path("bandit_output") + +SURFACE, INK, INK_2, MUTED, GRID, AXIS = ( + "#fcfcfb", "#0b0b0b", "#52514e", "#898781", "#e1e0d9", "#c3c2b7") +COLORS = {"etc_m10": "#2a78d6", "etc_m50": "#eb6834", "ucb_c3": "#1baf7a"} +LABELS = {"etc_m10": "explore-then-commit, m = 10", "etc_m50": "explore-then-commit, m = 50", + "ucb_c3": "UCB, c = 3"} + + +def from_hex(h): + return struct.unpack(" 0) + + +# -- Loading and checking -------------------------------------------------------------------------- + +def load(): + with open(OUT / "config.csv") as f: + config = list(csv.DictReader(f)) + runs = {} + for c in config: + name = c["name"] + with open(OUT / f"{name}.csv") as f: + rows = list(csv.DictReader(f)) + with open(OUT / f"{name}_paths.csv") as f: + prows = list(csv.DictReader(f)) + seeds = [k for k in prows[0] if k.startswith("seed")] + runs[name] = dict( + alg=c["algorithm"], param=float(c["parameter"]) if c["algorithm"] == "ucb" + else int(c["parameter"]), T=int(c["rounds"]), seeds=int(c["seeds"]), + means=[float(x) for x in c["means"].split(";")], var=float(c["variance"]), + round=np.array([int(r["round"]) for r in rows]), + mean=np.array([float(r["mean"]) for r in rows]), + se=np.array([float(r["se"]) for r in rows]), + q10=np.array([float(r["q10"]) for r in rows]), + q90=np.array([float(r["q90"]) for r in rows]), + paths={int(k[4:]): [from_hex(r[k]) for r in prows] for k in seeds}) + return runs + + +def check(runs): + ok = True + for name, r in runs.items(): + for seed, path in r["paths"].items(): + ref = run_reference(r["alg"], r["param"], r["means"], r["var"], seed, r["T"]) + same = [bits(x) for x in ref] == [bits(x) for x in path] + ok = ok and same + if not same: + i = next(i for i, (x, y) in enumerate(zip(ref, path)) if bits(x) != bits(y)) + print(f"{name} seed {seed}: DIVERGES at round {i + 1}") + print(f"{name}: numpy reference identical on seeds {sorted(r['paths'])}, " + f"all {r['T']} rounds" if ok else f"{name}: MISMATCH") + return ok + + +# -- Plotting -------------------------------------------------------------------------------------- + +def style(): + plt.rcParams.update({ + "figure.facecolor": SURFACE, "axes.facecolor": SURFACE, "savefig.facecolor": SURFACE, + "font.family": "sans-serif", "font.size": 9, + "text.color": INK, "axes.titlecolor": INK, "axes.titlesize": 10, + "axes.titleweight": "bold", "axes.titlelocation": "left", + "axes.labelcolor": INK_2, "xtick.color": MUTED, "ytick.color": MUTED, + "xtick.labelcolor": INK_2, "ytick.labelcolor": INK_2, + "axes.edgecolor": AXIS, "axes.linewidth": 0.8, + "axes.spines.top": False, "axes.spines.right": False, + "axes.grid": True, "grid.color": GRID, "grid.linewidth": 0.6, "grid.linestyle": "-", + "axes.axisbelow": True, "legend.frameon": False, "legend.labelcolor": INK_2, + "lines.linewidth": 1.8, "lines.solid_capstyle": "round", + }) + + +def bound_curve(r): + means, var = r["means"], r["var"] + gaps = [max(means) - m for m in means] + n = r["round"] + if r["alg"] == "etc": + m = r["param"] + start = len(means) * m + ns = n[n >= start] + return ns, np.array([etc_bound(gaps, m, var, k) for k in ns]) + c = r["param"] + const = np.cumsum([1.0 / (s + 1) ** (c / var - 1) for s in range(int(n[-1]))]) + return n, np.array([ucb_bound(gaps, c, var, k, const[k - 1]) for k in n]) + + +def plot_regret(runs): + fig, axes = plt.subplots(1, 3, figsize=(11, 3.7)) + for ax, name in zip(axes, ["etc_m10", "etc_m50", "ucb_c3"]): + r, color = runs[name], COLORS[name] + ns, b = bound_curve(r) + ax.plot(ns, b, color=INK, linewidth=1.5, label="bound proved in Lean") + ax.fill_between(r["round"], r["q10"], r["q90"], color=color, alpha=0.16, linewidth=0, + label="10%–90% of the runs") + ax.plot(r["round"], r["mean"], color=color, label=f"mean over {r['seeds']} runs") + ax.set_title(LABELS[name]) + ax.set_xlabel("round") + ax.set_ylim(0, max(b[-1], r["q90"][-1]) * 1.08) + ax.annotate(f"{b[-1]:.0f}", (ns[-1], b[-1]), xytext=(-4, 4), textcoords="offset points", + ha="right", color=INK_2, fontsize=8) + ax.annotate(f"{r['mean'][-1]:.0f}", (r["round"][-1], r["mean"][-1]), xytext=(-4, 4), + textcoords="offset points", ha="right", color=INK_2, fontsize=8) + axes[0].set_ylabel("cumulative pseudo-regret") + axes[0].legend(loc="upper left", fontsize=8) + fig.suptitle("Regret of the rdo bandit programs against the bounds proved for them " + "(3 Gaussian arms, means 1, 0.5, 0, variance 1)", x=0.01, ha="left", + fontsize=11, fontweight="bold") + fig.tight_layout() + fig.savefig(OUT / "bandit_regret.png", dpi=150) + plt.close(fig) + + +def plot_compare(runs): + fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 3.7)) + for name in ["etc_m10", "etc_m50", "ucb_c3"]: + r = runs[name] + ax1.plot(r["round"], r["mean"], color=COLORS[name], label=LABELS[name]) + ax1.set_title("Mean regret") + ax1.set_xlabel("round") + ax1.set_ylabel("cumulative pseudo-regret") + ax1.legend(loc="upper left", fontsize=8) + r = runs["etc_m10"] + for seed, path in sorted(r["paths"].items()): + ax2.plot(r["round"], path, color=COLORS["etc_m10"], linewidth=1.1, alpha=0.9) + ax2.axvline(len(r["means"]) * r["param"], color=AXIS, linewidth=0.8, zorder=0) + ax2.set_title("Five runs of explore-then-commit, m = 10") + ax2.set_xlabel("round") + ax2.annotate("commits after 30 rounds", (len(r["means"]) * r["param"], 0), xytext=(4, 6), + textcoords="offset points", color=MUTED, fontsize=8) + fig.tight_layout() + fig.savefig(OUT / "bandit_compare.png", dpi=150) + plt.close(fig) + + +def main(): + if "--run" in sys.argv: + subprocess.run(["lake", "exe", "bandits"], check=True) + runs = load() + ok = check(runs) + style() + plot_regret(runs) + plot_compare(runs) + print(f"plots written to {OUT}/") + print("all checks passed" if ok else "SOME CHECKS FAILED") + sys.exit(0 if ok else 1) + + +if __name__ == "__main__": + main() From 389e0f758632f0d52b051a280dfcf2a0d9e99448 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 19 Sep 2026 10:02:41 +0200 Subject: [PATCH 30/34] =?UTF-8?q?=CE=B5-greedy=20bandit,=20with=20alg=5Fen?= =?UTF-8?q?v=5Ftrace?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ε-greedy's policy is an rdo program drawing a coin of bias ε and a uniform arm. With `rdo_trace` and `alg_env_trace`, its draws give that every arm is pulled with probability at least ε / K at every round (`le_map_action`), hence a linear lower bound on the regret (`le_integral_regret`). The bandit programs gain a randomized step, `banditStepRand` and `banditRunRand`, where the arm is itself a program, and `banditRunRand_eq_map` extends the law of the program to it, so that the lower bound holds for the program that runs (`le_integral_regret_banditRunRand`). The greedy arm pulls the arms never pulled first, so that it never compares undefined means. `lake exe bandits` runs ε-greedy with ε = 0.1, and `scripts/bandit_plot.py` replays it with numpy and draws its regret against the lower bound. Co-Authored-By: Claude Opus 5 (1M context) --- Bandits.lean | 3 + Bandits/Defs.lean | 64 +++++++++++- Bandits/EpsGreedy.lean | 221 +++++++++++++++++++++++++++++++++++++++++ Bandits/Main.lean | 48 +++++---- Bandits/Theory.lean | 93 +++++++++++++++++ scripts/bandit_plot.py | 75 +++++++++----- 6 files changed, 460 insertions(+), 44 deletions(-) create mode 100644 Bandits/EpsGreedy.lean diff --git a/Bandits.lean b/Bandits.lean index 6e436cb..874242d 100644 --- a/Bandits.lean +++ b/Bandits.lean @@ -6,6 +6,7 @@ Authors: Rémy Degenne module -- shake: keep-all --deprecated_module: ignore public import Bandits.Defs +public import Bandits.EpsGreedy public import Bandits.Theory /-! @@ -15,6 +16,8 @@ public import Bandits.Theory monad and the scalars; explore-then-commit and UCB as the algorithms choosing the arm. * `Bandits.Theory`: read at `Measure`, the programs have the law of LeanMachineLearning's interaction, so its regret bounds hold for them. +* `Bandits.EpsGreedy`: ε-greedy, a randomized algorithm whose policy is an `rdo` program; with + `alg_env_trace`, its internal draws give an exploration bound and a linear regret lower bound. To run them and draw the regret against the bounds, from the root of the repository: diff --git a/Bandits/Defs.lean b/Bandits/Defs.lean index f98e2d8..ac93ba1 100644 --- a/Bandits/Defs.lean +++ b/Bandits/Defs.lean @@ -8,6 +8,7 @@ module public import RandomDo public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg +public import Mathlib.Probability.Distributions.Uniform /-! # A Gaussian bandit, as an `rdo` program @@ -23,7 +24,8 @@ The programs are polymorphic in the monad and in the scalars, as those of The two algorithms, `etcArm` (explore-then-commit) and `ucbArm` (upper confidence bound), mirror the definitions of `Bandits.ETC.nextArm` and `Bandits.UCB.nextArm` in LeanMachineLearning, with the -statistics of the history in place of sums over it. +statistics of the history in place of sums over it. A third, `epsGreedyArm` (ε-greedy), is +randomized: it is itself an `rdo` program, played by `banditStepRand` and `banditRunRand`. ## Main definitions @@ -31,12 +33,22 @@ statistics of the history in place of sums over it. * `RDoBandit.State`, `RDoBandit.State.update`: what the algorithms keep of the history. * `RDoBandit.etcArm`, `RDoBandit.ucbArm`: the arm the two algorithms pull. * `RDoBandit.banditStep`, `RDoBandit.banditRun`: one round, and `n` rounds, of interaction. +* `RDoBandit.HasUniformFin`, `RDoBandit.epsGreedyArm`: ε-greedy, drawing its arm. +* `RDoBandit.banditStepRand`, `RDoBandit.banditRunRand`: the interaction for an algorithm that + draws its arm. -/ @[expose] public section open MeasureTheory ProbabilityTheory Learning MeasurableSpacePure +universe v + +/-- A typeclass for monads that can draw an arm uniformly. -/ +class RDoBandit.HasUniformFin (m : (α : Type) → [MeasurableSpace α] → Type v) where + /-- Draw an element of `Fin K` uniformly. -/ + uniformFin (K : ℕ) [NeZero K] : m (Fin K) + /-- A typeclass for scalars in which a nonempty tuple has an index of its maximum. -/ class RDoBandit.HasArgmax (R : Type) where /-- An index at which the tuple is maximal. -/ @@ -52,6 +64,19 @@ when the maximum is attained once. -/ instance : HasArgmax Float where argmax {K} _ f := (List.finRange K).foldl (fun best a ↦ if f best < f a then a else best) 0 +noncomputable instance : HasUniformFin Measure where + uniformFin K _ := (PMF.uniformOfFintype (Fin K)).toMeasure + +instance (K : ℕ) [NeZero K] : + IsProbabilityMeasure (HasUniformFin.uniformFin (m := Measure) K) := by + change IsProbabilityMeasure (PMF.uniformOfFintype (Fin K)).toMeasure + infer_instance + +/-- The uniform draw of `NumLean`, as numpy's `Generator.integers`. -/ +instance : HasUniformFin RandM where + uniformFin K _ := show NumLean.RandPCG IO (Fin K) from do + return ⟨(← NumLean.randInt K).toNat % K, Nat.mod_lt _ (NeZero.pos K)⟩ + /-- What a bandit algorithm keeps of the history: the number of pulls of each arm, the sum of the rewards each arm returned, and the last arm pulled (arm `0` before the first round). The first two are arrays, so that updating them in a long run stays cheap. -/ @@ -89,7 +114,17 @@ def ucbArm [Add R] [Mul R] [One R] [OfNat R 2] [HasLog R] [HasSqrt R] (c : R) (n else HasArgmax.argmax fun a ↦ s.2.1[a] / (s.1[a] : R) + HasSqrt.sqrt (2 * c * HasLog.log ((n : R) + 1) / (s.1[a] : R)) -universe v +/-- The first arm never pulled, if any. -/ +def firstUnpulled (s : State K R) : Option (Fin K) := + (List.finRange K).find? fun a ↦ s.1[a] == 0 + +/-- The greedy arm: the first arm never pulled if there is one, and otherwise the arm with the best +empirical mean. Pulling the unpulled arms first keeps the empirical means defined when they are +compared. -/ +def greedyArm (s : State K R) : Fin K := + match firstUnpulled s with + | some a => a + | none => HasArgmax.argmax fun a ↦ s.2.1[a] / (s.1[a] : R) variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] [Zero R] [Add R] [MeasurableSpace R] {V : Type} [HasGaussian m R V R] @@ -102,6 +137,14 @@ def banditStep (arm : ℕ → State K R → Fin K) (μ : Fin K → R) (σ2 : V) let r ← HasGaussian.gaussian (m := m) (μ a) σ2 return s.update a r +/-- **ε-greedy**, drawing the arm it pulls: a coin of bias `ε`; on heads an arm drawn uniformly, on +tails the greedy arm. Both draws are made every round. -/ +def epsGreedyArm [HasBernoulli m R] [HasUniformFin m] (ε : R) (_n : ℕ) (s : State K R) : + m (Fin K) := rdo + let explore ← HasBernoulli.bernoulli (m := m) ε + let u ← HasUniformFin.uniformFin (m := m) K + return if explore then u else greedyArm s + /-- **`n` rounds**, from the initial state. -/ def banditRun (arm : ℕ → State K R → Fin K) (μ : Fin K → R) (σ2 : V) : ℕ → m (State K R) | 0 => mPure State.init @@ -110,4 +153,21 @@ def banditRun (arm : ℕ → State K R → Fin K) (μ : Fin K → R) (σ2 : V) : let s' ← banditStep (m := m) arm μ σ2 n s return s' +/-- **One round**, for an algorithm that draws its arm: the program `arm` draws the arm from the +state, the arm returns a reward drawn from `𝒩(μ a, σ2)`, and the state records it. -/ +def banditStepRand (arm : ℕ → State K R → m (Fin K)) (μ : Fin K → R) (σ2 : V) (n : ℕ) + (s : State K R) : m (State K R) := rdo + let a ← arm n s + let r ← HasGaussian.gaussian (m := m) (μ a) σ2 + return s.update a r + +/-- **`n` rounds**, for an algorithm that draws its arm. -/ +def banditRunRand (arm : ℕ → State K R → m (Fin K)) (μ : Fin K → R) (σ2 : V) : + ℕ → m (State K R) + | 0 => mPure State.init + | n + 1 => rdo + let s ← banditRunRand arm μ σ2 n + let s' ← banditStepRand (m := m) arm μ σ2 n s + return s' + end RDoBandit diff --git a/Bandits/EpsGreedy.lean b/Bandits/EpsGreedy.lean new file mode 100644 index 0000000..55afb20 --- /dev/null +++ b/Bandits/EpsGreedy.lean @@ -0,0 +1,221 @@ +/- +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 Bandits.Theory + +/-! +# ε-greedy, and the draws inside its algorithm-environment sequence + +ε-greedy is a randomized bandit algorithm: at each round it tosses a coin of bias `ε`, and on heads +pulls an arm drawn uniformly, on tails the greedy arm (the first arm never pulled if there is one, +and otherwise the one with the best empirical mean). Its policy is written here as an `rdo` +program, and turned into a LeanMachineLearning `Algorithm`. + +A statement such as *every arm is pulled with probability at least `ε / K`* is about the coin and +the uniform draw, but an algorithm-environment sequence `IsAlgEnvSeq O A R alg env P` only has the +actions: the draws of the policy are not random variables of `(Ω, P)`. `alg_env_trace` supplies +them. From the trace `rdo_trace` finds for the policy (`hasTrace_policy`), it moves the goal to a +space that also carries the draws `T n`, with their conditional law given the history and the +action as a readout of them, and discharges the obligation of transferring the result back. + +There, the draws have the law `Ber(ε) ⊗ uniform` whatever the history (`draws_eq_const`); on the +event where the coin says explore and the uniform draw is `a`, the action is `a`. Hence the +exploration bound, for any environment, and a linear lower bound on the regret. + +## Main results + +* `le_map_action`: at every round, every arm is pulled with probability at least `ε / K`. +* `le_integral_regret`: the expected regret of ε-greedy after `n` rounds is at least + `n ε / K ∑ₐ Δₐ`. +* `le_integral_regret_banditRunRand`: the same lower bound for the program `banditRunRand` of + `Bandits.Defs` running `epsGreedyArm`, which is ε-greedy on the statistics of the history + (`policy_eq_epsGreedyArm`, `Bandits.Theory.banditRunRand_eq_map`). This is the program that runs. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Learning RDo Bandits +open scoped ENNReal NNReal + +namespace RDoBandit.EpsGreedy + +variable {K : ℕ} [NeZero K] + +/-- The uniform distribution on the arms. -/ +noncomputable def uniformArm : Measure (Fin K) := (PMF.uniformOfFintype (Fin K)).toMeasure + +instance : IsProbabilityMeasure (uniformArm (K := K)) := by unfold uniformArm; infer_instance + +lemma uniformArm_singleton (a : Fin K) : uniformArm {a} = (K : ℝ≥0∞)⁻¹ := by + rw [uniformArm, PMF.toMeasure_apply_singleton _ _ (measurableSet_singleton a), + PMF.uniformOfFintype_apply, Fintype.card_fin] + +/-- The greedy arm of the statistics is measurable: the first unpulled arm is read off counts in a +countable set, and the argmax is measurable. -/ +@[fun_prop] +lemma measurable_greedyArm_state : Measurable (RDoBandit.greedyArm (K := K) (R := ℝ)) := by + let F : (Fin K → ℕ) × Fin K → Fin K := fun p ↦ + match (List.finRange K).find? fun a ↦ p.1 a == 0 with + | some a => a + | none => p.2 + have hF : Measurable F := by + refine measurable_from_prod_countable_right fun c ↦ ?_ + simp only [F] + split <;> fun_prop + have hc : Measurable fun s : State K ℝ ↦ fun a : Fin K ↦ s.1[a] := + measurable_pi_iff.2 fun a ↦ (measurable_vector_getElem a).comp measurable_fst + have hg : Measurable fun s : State K ℝ ↦ + (HasArgmax.argmax fun a ↦ s.2.1[a] / (s.1[a] : ℝ) : Fin K) := + measurable_hasArgmax_real.comp (measurable_pi_iff.2 fun a ↦ by fun_prop) + exact hF.comp (hc.prodMk hg) + +/-- The greedy arm of a history: the greedy arm of its statistics, the first arm never pulled if +there is one, and otherwise the one with the best empirical mean. -/ +noncomputable def greedyArm (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : Fin K := + RDoBandit.greedyArm (histState n h) + +@[fun_prop] +lemma measurable_greedyArm (n : ℕ) : Measurable (greedyArm (K := K) n) := + measurable_greedyArm_state.comp (measurable_histState n) + +variable (ε : unitInterval) + +/-- **The ε-greedy policy**, as an `rdo` program: toss a coin of bias `ε`; on heads explore an arm +drawn uniformly, on tails exploit the greedy arm. -/ +noncomputable def policy (n : ℕ) (p : Hist Unit (Fin K) ℝ n × Unit) : Measure (Fin K) := rdo + let explore ← bernoulliMeasure true false ε + let u ← uniformArm + return if explore then u else greedyArm n p.1 + +instance (n : ℕ) : IsMarkov (policy (K := K) ε n) := by unfold policy; is_markov + +/-- ε-greedy, as a LeanMachineLearning algorithm. -/ +noncomputable def alg : Algorithm Unit (Fin K) ℝ where + policy n := markovKernel (policy ε n) inferInstance + +/-- The draws the policy makes at round `n`: the coin, then the uniform arm. This is the trace +kernel `rdo_trace` finds; neither draw reads the history. -/ +noncomputable def draws (n : ℕ) : Kernel (Hist Unit (Fin K) ℝ n × Unit) (Bool × Fin K) := + Kernel.const _ (bernoulliMeasure true false ε) + ⊗ₖ Kernel.prodMkRight Bool (Kernel.const _ uniformArm) + +/-- The action, read off the history and the draws. -/ +noncomputable def readout (n : ℕ) (p : (Hist Unit (Fin K) ℝ n × Unit) × (Bool × Fin K)) : Fin K := + if p.2.1 = true then p.2.2 else greedyArm n p.1.1 + +instance (n : ℕ) : IsMarkovKernel (draws (K := K) ε n) := by unfold draws; infer_instance + +/-- `rdo_trace` finds the trace of the policy. -/ +lemma hasTrace_policy (n : ℕ) : HasTrace (policy (K := K) ε n) (draws ε n) (readout n) := by + rdo_trace (policy (K := K) ε n) with h + exact h + +/-- The trace of ε-greedy, for `alg_env_trace`. -/ +noncomputable def trace : AlgTrace (alg (K := K) ε) (Bool × Fin K) where + K := draws ε + out := readout + hasTrace n := hasTrace_policy ε n + +/-- The draws do not depend on the history: they are a coin and an independent uniform arm. -/ +lemma draws_eq_const (n : ℕ) : + draws (K := K) ε n = Kernel.const _ ((bernoulliMeasure true false ε).prod uniformArm) := by + ext1 p + rw [draws, Kernel.compProd_apply_eq_compProd_sectR, Kernel.const_apply, Kernel.const_apply] + have : (Kernel.prodMkRight Bool + (Kernel.const (Hist Unit (Fin K) ℝ n × Unit) (uniformArm (K := K)))).sectR p + = Kernel.const Bool (uniformArm (K := K)) := by ext1 b; rfl + rw [this, Measure.compProd_const] + +/-- **Every arm is explored.** Whatever the environment, at every round, ε-greedy pulls each arm +with probability at least `ε / K`. -/ +theorem le_map_action (env : Environment Unit (Fin K) ℝ) {Ω₀ : Type} [MeasurableSpace Ω₀] + {P : Measure Ω₀} [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} {A : ℕ → Ω₀ → Fin K} + {R : ℕ → Ω₀ → ℝ} (h : IsAlgEnvSeq O A R (alg (K := K) ε) env P) (n : ℕ) (a : Fin K) : + ENNReal.ofReal (ε / K) ≤ P.map (A n) {a} := by + alg_env_trace (trace (K := K) ε) with Ω P O A R T hseq htr hT hA + have hlaw : HasLaw (T n) ((bernoulliMeasure true false ε).prod uniformArm) P := by + have h1 := hT n + simp only [trace, draws_eq_const] at h1 + exact h1.hasLaw_of_const + -- When the coin says explore and the uniform draw is `a`, the action is `a`. + have hsub : T n ⁻¹' {(true, a)} ≤ᵐ[P] A n ⁻¹' {a} := by + filter_upwards [hA n] with ω hω hTω + simp only [Set.mem_preimage, Set.mem_singleton_iff] at hTω ⊢ + rw [hω] + simp [trace, readout, hTω] + calc ENNReal.ofReal (ε / K) + = ((bernoulliMeasure true false ε).prod uniformArm) {(true, a)} := by + rw [← Set.singleton_prod_singleton, Measure.prod_prod, + bernoulliMeasure_apply_of_mem_of_notMem _ (measurableSet_singleton _) (by simp) + (by simp), uniformArm_singleton, + ENNReal.ofReal_div_of_pos (by exact_mod_cast NeZero.pos K), ENNReal.ofReal_natCast, + div_eq_mul_inv, ← ENNReal.ofReal_coe_nnreal, unitInterval.coe_toNNReal] + _ = P (T n ⁻¹' {(true, a)}) := by + rw [← hlaw.map_eq, Measure.map_apply_of_aemeasurable hlaw.aemeasurable + (measurableSet_singleton _)] + _ ≤ P (A n ⁻¹' {a}) := measure_mono_ae hsub + _ = P.map (A n) {a} := + (Measure.map_apply (hseq.measurable_action n) (measurableSet_singleton a)).symm + +/-- **ε-greedy has linear regret.** Against any stationary environment, the expected regret after +`n` rounds is at least `n ε / K` times the sum of the gaps: every round, each arm is explored with +probability at least `ε / K`. -/ +theorem le_integral_regret (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] {Ω₀ : Type} + [MeasurableSpace Ω₀] {P : Measure Ω₀} [IsProbabilityMeasure P] {O : ℕ → Ω₀ → Unit} + {A : ℕ → Ω₀ → Fin K} {R : ℕ → Ω₀ → ℝ} + (h : IsAlgEnvSeq O A R (alg (K := K) ε) (stationaryEnv ν) P) (n : ℕ) : + (n : ℝ) * ((ε : ℝ) / K) * ∑ a, gap ν a ≤ P[regret ν A n] := by + have hA := h.measurable_action + have hint (s : ℕ) : Integrable (fun ω ↦ gap ν (A s ω)) P := + (Integrable.of_finite (μ := P.map (A s)) (f := gap ν)).comp_measurable (hA s) + have hround (s : ℕ) : (ε : ℝ) / K * ∑ a, gap ν a ≤ ∫ ω, gap ν (A s ω) ∂P := by + rw [← integral_map (hA s).aemeasurable (measurable_of_countable _).aestronglyMeasurable, + integral_fintype (Integrable.of_finite (f := gap ν)), Finset.mul_sum] + refine Finset.sum_le_sum fun a _ ↦ ?_ + rw [smul_eq_mul] + refine mul_le_mul_of_nonneg_right ?_ (gap_nonneg (ν := ν) (a := a)) + rw [measureReal_def] + exact (ENNReal.ofReal_le_iff_le_toReal (measure_ne_top _ _)).1 + (le_map_action ε (stationaryEnv ν) h s a) + simp_rw [regret_eq_sum_gap] + rw [integral_finsetSum _ fun s _ ↦ hint s] + calc (n : ℝ) * ((ε : ℝ) / K) * ∑ a, gap ν a + = ∑ _s ∈ Finset.range n, (ε : ℝ) / K * ∑ a, gap ν a := by + simp [Finset.sum_const]; ring + _ ≤ ∑ s ∈ Finset.range n, ∫ ω, gap ν (A s ω) ∂P := Finset.sum_le_sum fun s _ ↦ hround s + +/-! ### The program that runs -/ + +/-- At `Measure`, the program `epsGreedyArm` is a Markov kernel in the state. -/ +instance isMarkov_epsGreedyArm (n : ℕ) : + IsMarkov (epsGreedyArm (m := Measure) (K := K) (ε : ℝ) n) := by + unfold epsGreedyArm + is_markov + +/-- **The program is ε-greedy.** At `Measure`, the program `epsGreedyArm` on the statistics of a +history is the policy of ε-greedy on that history. -/ +lemma policy_eq_epsGreedyArm (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : + (alg (K := K) ε).policy n (h, ()) = epsGreedyArm (m := Measure) (ε : ℝ) n (histState n h) := by + change policy ε n (h, ()) = (bernoulliMeasure true false (Set.projIcc 0 1 zero_le_one ε)).bind + fun b ↦ (uniformArm (K := K)).bind fun u ↦ + Measure.dirac (if b = true then u else RDoBandit.greedyArm (histState n h)) + rw [Set.projIcc_of_mem _ ε.2] + rfl + +/-- **ε-greedy has linear regret**, for the program that runs: the expected pseudo-regret of +`banditRunRand` with `epsGreedyArm`, against Gaussian arms, is at least `n ε / K ∑ₐ Δₐ`. -/ +theorem le_integral_regret_banditRunRand (μ : Fin K → ℝ) (σ2 : ℝ≥0) (n : ℕ) : + (n : ℝ) * ((ε : ℝ) / K) * ∑ a, gapOf μ a + ≤ ∫ s, pseudoRegret μ s + ∂(banditRunRand (m := Measure) (epsGreedyArm (m := Measure) (ε : ℝ)) μ σ2 n) := by + rw [integral_pseudoRegret_banditRunRand μ σ2 (alg ε) _ (isMarkov_epsGreedyArm ε) + (policy_eq_epsGreedyArm ε) n] + have h := le_integral_regret ε (arms μ σ2) + (IT.isAlgEnvSeq_trajMeasure (alg (K := K) ε) (stationaryEnv (arms μ σ2))) n + simpa only [gap_arms] using h + +end RDoBandit.EpsGreedy diff --git a/Bandits/Main.lean b/Bandits/Main.lean index 42bf16b..0c15d1a 100644 --- a/Bandits/Main.lean +++ b/Bandits/Main.lean @@ -8,9 +8,9 @@ import Bandits.Defs /-! # Running the Gaussian bandit -`lake exe bandits [rounds] [seeds]` runs explore-then-commit and UCB on a three-armed Gaussian -bandit, through the program `banditStep` of `Bandits.Defs` read at `RandM` and `Float`, for many -seeds, and writes the cumulative pseudo-regret in `bandit_output/`: +`lake exe bandits [rounds] [seeds]` runs explore-then-commit, UCB and ε-greedy on a three-armed +Gaussian bandit, through the programs `banditStep` and `banditStepRand` of `Bandits.Defs` read at +`RandM` and `Float`, for many seeds, and writes the cumulative pseudo-regret in `bandit_output/`: * `.csv`: for each round, the mean regret over the seeds, its standard error, and its 10% and 90% quantiles; @@ -18,8 +18,8 @@ seeds, and writes the cumulative pseudo-regret in `bandit_output/`: bits of each `Float`, for `scripts/bandit_plot.py` to replay them with numpy; * `config.csv`: the parameters, for the same purpose. -It also checks that the one-shot program `banditRun n` lands on the state the round-by-round run -reaches after `n` rounds. +It also checks that the one-shot program, `banditRun n` or `banditRunRand n`, lands on the state the +round-by-round run reaches after `n` rounds. -/ open RDoBandit NumLean @@ -42,13 +42,24 @@ def variance : Float := 1.0 structure Algo where /-- Its name, which names its files. -/ name : String - /-- The arm it pulls at each round, given the state. -/ - arm : ℕ → State numArms Float → Fin numArms - -/-- One round, at `RandM` and `Float`. -/ -def Algo.step (alg : Algo) (n : ℕ) (s : State numArms Float) : - RandPCG IO (State numArms Float) := - (banditStep (m := RandM) (R := Float) (V := Float) alg.arm means variance n s : RandM _) + /-- One round, at `RandM` and `Float`. -/ + step : ℕ → State numArms Float → RandPCG IO (State numArms Float) + /-- The one-shot program: `n` rounds from the initial state. -/ + oneShot : ℕ → RandPCG IO (State numArms Float) + +/-- An algorithm choosing its arm deterministically from the state, run by `banditStep`. -/ +def Algo.det (name : String) (arm : ℕ → State numArms Float → Fin numArms) : Algo where + name := name + step n s := (banditStep (m := RandM) (R := Float) (V := Float) arm means variance n s : RandM _) + oneShot n := (banditRun (m := RandM) (R := Float) (V := Float) arm means variance n : RandM _) + +/-- An algorithm drawing its arm, run by `banditStepRand`. -/ +def Algo.rand (name : String) (arm : ℕ → State numArms Float → RandM (Fin numArms)) : Algo where + name := name + step n s := + (banditStepRand (m := RandM) (R := Float) (V := Float) arm means variance n s : RandM _) + oneShot n := + (banditRunRand (m := RandM) (R := Float) (V := Float) arm means variance n : RandM _) /-- `T` rounds, recording the cumulative pseudo-regret after each. -/ def Algo.run (alg : Algo) (T : ℕ) : RandPCG IO (Array Float × State numArms Float) := do @@ -96,13 +107,11 @@ def Algo.go (alg : Algo) (T reps : ℕ) (maxPaths : ℕ := 5) : IO Bool := do -- The one-shot program, on the first seed: the same draws, in the same order. let nCheck := min 500 T let (_, sDriver) ← (IO.runRandPCGWith 1 (alg.run nCheck) : IO _) - let sOneShot ← (IO.runRandPCGWith 1 - (banditRun (m := RandM) (R := Float) (V := Float) alg.arm means variance nCheck : RandM _) : - IO (State numArms Float)) + let sOneShot ← (IO.runRandPCGWith 1 (alg.oneShot nCheck) : IO (State numArms Float)) let ok := sameState sDriver sOneShot let final := curves.map (·.back!) IO.println s!"{alg.name}: {reps} seeds × {T} rounds, mean final regret \ - {final.foldl (· + ·) 0 / reps.toFloat}; banditRun {nCheck} = {nCheck} steps: {ok}" + {final.foldl (· + ·) 0 / reps.toFloat}; one-shot {nCheck} = {nCheck} steps: {ok}" return ok /-- `lake exe bandits [rounds] [seeds]`, by default 5000 rounds and 300 seeds. -/ @@ -111,9 +120,10 @@ def main (args : List String) : IO UInt32 := do let T := (args[0]?.bind String.toNat?).getD 5000 let reps := (args[1]?.bind String.toNat?).getD 300 let algos : List (Algo × String) := [ - ({ name := "etc_m10", arm := etcArm 10 }, "etc,10"), - ({ name := "etc_m50", arm := etcArm 50 }, "etc,50"), - ({ name := "ucb_c3", arm := ucbArm 3 }, "ucb,3")] + (.det "etc_m10" (etcArm 10), "etc,10"), + (.det "etc_m50" (etcArm 50), "etc,50"), + (.det "ucb_c3" (ucbArm 3), "ucb,3"), + (.rand "epsgreedy_0.1" (epsGreedyArm (m := RandM) 0.1), "epsgreedy,0.1")] IO.FS.withFile "bandit_output/config.csv" .write fun h ↦ do h.putStrLn "name,algorithm,parameter,rounds,seeds,means,variance" for (alg, desc) in algos do diff --git a/Bandits/Theory.lean b/Bandits/Theory.lean index d735f4b..8173909 100644 --- a/Bandits/Theory.lean +++ b/Bandits/Theory.lean @@ -377,6 +377,87 @@ theorem banditRun_eq_map (arm : ℕ → State K ℝ → Fin K) (harm_meas : ∀ end Law +section RandomizedLaw + +/-! ### Algorithms drawing their arm + +The same statement for an algorithm whose policy, at `Measure`, is a program drawing the arm from +the state: the step kernel then binds over the arm the policy draws, instead of taking the one it +chooses. -/ + +variable {μ : Fin K → ℝ} {σ2 : ℝ≥0} + +instance (s : ℝ≥0) : IsMarkov fun x : ℝ ↦ gaussianReal x s := + IsMarkov.gaussianReal measurable_id measurable_const + +omit [NeZero K] in +/-- At `Measure`, one round is the arm drawn by the policy, then its reward, mapped by the +update. -/ +lemma banditStepRand_eq (arm : ℕ → State K ℝ → Measure (Fin K)) (n : ℕ) (s : State K ℝ) : + banditStepRand (m := Measure) arm μ σ2 n s + = (arm n s).bind fun a ↦ (gaussianReal (μ a) σ2).map (s.update a) := by + change (arm n s).bind (fun a ↦ (gaussianReal (μ a) σ2).bind + (fun r ↦ Measure.dirac (s.update a r))) = _ + congr with a : 1 + rw [Measure.bind_dirac_eq_map _ (by fun_prop)] + +lemma banditRunRand_succ (arm : ℕ → State K ℝ → Measure (Fin K)) (n : ℕ) : + banditRunRand (m := Measure) arm μ σ2 (n + 1) + = (banditRunRand (m := Measure) arm μ σ2 n).bind + (banditStepRand (m := Measure) arm μ σ2 n) := by + rw [banditRunRand] + rfl + +omit [NeZero K] in +lemma isMarkov_banditStepRand {arm : ℕ → State K ℝ → Measure (Fin K)} {n : ℕ} + (harm : IsMarkov (arm n)) : IsMarkov (banditStepRand (m := Measure) arm μ σ2 n) := by + have hμ : Measurable μ := measurable_of_countable μ + unfold banditStepRand + is_markov + +omit [NeZero K] in +/-- One round of an algorithm against the Gaussian arms: the arm its policy draws, then the reward +of that arm. -/ +lemma map_stepKernel_bind {γ : Type*} [MeasurableSpace γ] (alg : Algorithm Unit (Fin K) ℝ) + (n : ℕ) (h : Hist Unit (Fin K) ℝ n) {G : Round Unit (Fin K) ℝ → γ} (hG : Measurable G) : + (stepKernel alg (stationaryEnv (arms μ σ2)) n h).map G + = (alg.policy n (h, ())).bind fun a ↦ (gaussianReal (μ a) σ2).map fun r ↦ G ((), a, r) := by + rw [stepKernel_stationaryEnv, Kernel.compProd_apply_eq_compProd_sectR, Kernel.const_apply, + Measure.dirac_unit_compProd, Kernel.sectR_apply, Kernel.compProd_apply_eq_compProd_sectR, + Measure.map_map hG measurable_prodMk_left, + map_compProd_eq_bind _ _ (hG.comp measurable_prodMk_left)] + rfl + +/-- **The law of the program**, for an algorithm drawing its arm. If, at every round, the policy +of `alg` on a history is the program `arm` on the state of that history, then the state after `n` +rounds of `banditRunRand` has the law of the state of the history of `n` rounds of the +interaction of `alg` with the Gaussian arms. -/ +theorem banditRunRand_eq_map (alg : Algorithm Unit (Fin K) ℝ) + (arm : ℕ → State K ℝ → Measure (Fin K)) (harm_markov : ∀ n, IsMarkov (arm n)) + (harm : ∀ n h, alg.policy n (h, ()) = arm n (histState n h)) (n : ℕ) : + banditRunRand (m := Measure) arm μ σ2 n + = (trajMeasure alg (stationaryEnv (arms μ σ2))).map (histState n ∘ IT.hist n) := by + induction n with + | zero => + have : histState 0 ∘ IT.hist (𝓞 := Unit) (𝓐 := Fin K) (𝓨 := ℝ) 0 = fun _ ↦ State.init := by + funext ω + exact histState_zero _ + rw [this, Measure.map_const, measure_univ, one_smul] + rfl + | succ n ih => + rw [banditRunRand_succ, ih, map_hist_succ _ _ n (measurable_histState (n + 1)), + ← Measure.map_map (measurable_histState n) (IT.measurable_hist n), + bind_map_eq _ (measurable_histState n) (isMarkov_banditStepRand (harm_markov n)).measurable] + congr 1 + funext h + have hG : Measurable fun x ↦ histState (n + 1) (Fin.snoc h x) := + (measurable_histState (n + 1)).comp ((measurable_snoc n).comp + (measurable_const.prodMk measurable_id)) + rw [Function.comp_apply, banditStepRand_eq, map_stepKernel_bind alg n h hG, harm] + simp_rw [histState_snoc] + +end RandomizedLaw + section Regret variable (μ : Fin K → ℝ) (σ2 : ℝ≥0) @@ -477,6 +558,18 @@ theorem integral_regret_ucb_le_of_gt {c : ℝ} (hc : 2 * σ2 < c) (hσ2 : σ2 measurable_fst) (ucbArm c) (measurable_ucbArm c) (fun n h ↦ ucbArm_histState c n h) n] exact h +/-- **The expected pseudo-regret of the program is the expected regret of the interaction**, for an +algorithm drawing its arm. -/ +theorem integral_pseudoRegret_banditRunRand (alg : Algorithm Unit (Fin K) ℝ) + (arm : ℕ → State K ℝ → Measure (Fin K)) (harm_markov : ∀ n, IsMarkov (arm n)) + (harm : ∀ n h, alg.policy n (h, ()) = arm n (histState n h)) (n : ℕ) : + ∫ s, pseudoRegret μ s ∂(banditRunRand (m := Measure) arm μ σ2 n) + = (trajMeasure alg (stationaryEnv (arms μ σ2)))[regret (arms μ σ2) IT.action n] := by + rw [banditRunRand_eq_map alg arm harm_markov harm n, + integral_map (by fun_prop) (measurable_pseudoRegret μ).aestronglyMeasurable] + congr with ω + exact (regret_eq_pseudoRegret μ σ2 n ω).symm + end Regret end RDoBandit diff --git a/scripts/bandit_plot.py b/scripts/bandit_plot.py index 604fa0f..7b3c177 100644 --- a/scripts/bandit_plot.py +++ b/scripts/bandit_plot.py @@ -6,15 +6,17 @@ python3 scripts/bandit_plot.py --run # runs `lake exe bandits` first It checks that a reference implementation in numpy, seeded alike, produces the regret of the first -seeds of each algorithm, bit for bit. It draws with `default_rng(seed).standard_normal`, which -`NumLean` reproduces, and computes in doubles as the Lean program does: the reward is -`fma(sqrt(variance), z, mean)`, as in `NumLean.normal`, the argmax is the first maximal index, as +seeds of each algorithm, bit for bit. It draws with `default_rng(seed)`: `standard_normal` for the +rewards, and for ε-greedy `binomial(1, ε)` for the coin and `integers(K)` for the uniform arm, in +that order, which `NumLean` reproduces. It computes in doubles as the Lean program does: the reward +is `fma(sqrt(variance), z, mean)`, as in `NumLean.normal`, the argmax is the first maximal index, as the `Float` instance of `HasArgmax`, and `sqrt` and `log` are libm's. It draws, in bandit_output/: -* bandit_regret.png: the regret of each algorithm against the bound proved for it in - `Bandits.Theory`; -* bandit_compare.png: the three algorithms together, and single runs of explore-then-commit. +* bandit_regret.png: the regret of each algorithm against the bound proved for it, an upper bound in + `Bandits.Theory` for explore-then-commit and UCB, a lower bound in `Bandits.EpsGreedy` for + ε-greedy; +* bandit_compare.png: the four algorithms together, and single runs of explore-then-commit. """ import csv @@ -35,9 +37,11 @@ SURFACE, INK, INK_2, MUTED, GRID, AXIS = ( "#fcfcfb", "#0b0b0b", "#52514e", "#898781", "#e1e0d9", "#c3c2b7") -COLORS = {"etc_m10": "#2a78d6", "etc_m50": "#eb6834", "ucb_c3": "#1baf7a"} +COLORS = {"etc_m10": "#2a78d6", "etc_m50": "#eb6834", "ucb_c3": "#1baf7a", + "epsgreedy_0.1": "#eda100"} LABELS = {"etc_m10": "explore-then-commit, m = 10", "etc_m50": "explore-then-commit, m = 50", - "ucb_c3": "UCB, c = 3"} + "ucb_c3": "UCB, c = 3", "epsgreedy_0.1": "ε-greedy, ε = 0.1"} +ORDER = ["etc_m10", "etc_m50", "ucb_c3", "epsgreedy_0.1"] def from_hex(h): @@ -63,6 +67,14 @@ def argmax_first(xs): return best +def greedy(N, S): + """The first arm never pulled, and otherwise the best empirical mean: `greedyArm`.""" + for a in range(len(N)): + if N[a] == 0: + return a + return argmax_first([S[a] / float(N[a]) for a in range(len(N))]) + + def choose(alg, param, n, N, S, last): K = len(N) if alg == "etc": @@ -86,7 +98,12 @@ def run_reference(alg, param, means, var, seed, T): sd = math.sqrt(var) regret, curve = 0.0, [] for n in range(T): - a = choose(alg, param, n, N, S, last) + if alg == "epsgreedy": + coin = rng.binomial(1, max(0.0, min(1.0, param))) + u = int(rng.integers(K)) + a = u if coin == 1 else greedy(N, S) + else: + a = choose(alg, param, n, N, S, last) r = fma(sd, rng.standard_normal(), means[a]) N[a] += 1 S[a] = S[a] + r @@ -104,6 +121,11 @@ def etc_bound(gaps, m, var, n): return sum(g * (m + (n - K * m) * math.exp(-m * g * g / (4 * var))) for g in gaps) +def epsgreedy_lower_bound(gaps, eps, n): + """`le_integral_regret_banditRunRand`: a lower bound.""" + return n * eps / len(gaps) * sum(gaps) + + def ucb_bound(gaps, c, var, n, const_sum): """`integral_regret_ucb_le`; the arm with no gap contributes nothing.""" return sum(8 * c * math.log(n + 1) / g + g * (2 + 2 * const_sum) for g in gaps if g > 0) @@ -123,8 +145,8 @@ def load(): prows = list(csv.DictReader(f)) seeds = [k for k in prows[0] if k.startswith("seed")] runs[name] = dict( - alg=c["algorithm"], param=float(c["parameter"]) if c["algorithm"] == "ucb" - else int(c["parameter"]), T=int(c["rounds"]), seeds=int(c["seeds"]), + alg=c["algorithm"], param=int(c["parameter"]) if c["algorithm"] == "etc" + else float(c["parameter"]), T=int(c["rounds"]), seeds=int(c["seeds"]), means=[float(x) for x in c["means"].split(";")], var=float(c["variance"]), round=np.array([int(r["round"]) for r in rows]), mean=np.array([float(r["mean"]) for r in rows]), @@ -172,6 +194,8 @@ def bound_curve(r): means, var = r["means"], r["var"] gaps = [max(means) - m for m in means] n = r["round"] + if r["alg"] == "epsgreedy": + return n, np.array([epsgreedy_lower_bound(gaps, r["param"], k) for k in n]) if r["alg"] == "etc": m = r["param"] start = len(means) * m @@ -183,34 +207,39 @@ def bound_curve(r): def plot_regret(runs): - fig, axes = plt.subplots(1, 3, figsize=(11, 3.7)) - for ax, name in zip(axes, ["etc_m10", "etc_m50", "ucb_c3"]): + fig, axes = plt.subplots(2, 2, figsize=(10, 7)) + axes = axes.ravel() + for ax, name in zip(axes, ORDER): r, color = runs[name], COLORS[name] ns, b = bound_curve(r) - ax.plot(ns, b, color=INK, linewidth=1.5, label="bound proved in Lean") + kind = "lower" if r["alg"] == "epsgreedy" else "upper" + ax.plot(ns, b, color=INK, linewidth=1.5, label=f"{kind} bound proved in Lean") ax.fill_between(r["round"], r["q10"], r["q90"], color=color, alpha=0.16, linewidth=0, label="10%–90% of the runs") ax.plot(r["round"], r["mean"], color=color, label=f"mean over {r['seeds']} runs") ax.set_title(LABELS[name]) ax.set_xlabel("round") ax.set_ylim(0, max(b[-1], r["q90"][-1]) * 1.08) - ax.annotate(f"{b[-1]:.0f}", (ns[-1], b[-1]), xytext=(-4, 4), textcoords="offset points", - ha="right", color=INK_2, fontsize=8) - ax.annotate(f"{r['mean'][-1]:.0f}", (r["round"][-1], r["mean"][-1]), xytext=(-4, 4), + below = kind == "lower" + ax.annotate(f"{b[-1]:.0f}", (ns[-1], b[-1]), xytext=(-4, -11 if below else 4), textcoords="offset points", ha="right", color=INK_2, fontsize=8) + ax.annotate(f"{r['mean'][-1]:.0f}", (r["round"][-1], r["mean"][-1]), + xytext=(-4, 4), textcoords="offset points", ha="right", color=INK_2, fontsize=8) + ax.legend(loc="upper left", fontsize=8) axes[0].set_ylabel("cumulative pseudo-regret") - axes[0].legend(loc="upper left", fontsize=8) - fig.suptitle("Regret of the rdo bandit programs against the bounds proved for them " - "(3 Gaussian arms, means 1, 0.5, 0, variance 1)", x=0.01, ha="left", - fontsize=11, fontweight="bold") - fig.tight_layout() + axes[2].set_ylabel("cumulative pseudo-regret") + fig.suptitle("Regret of the rdo bandit programs against the bounds proved for them", + x=0.01, ha="left", fontsize=11, fontweight="bold") + fig.text(0.01, 0.945, "3 Gaussian arms, means 1, 0.5 and 0, variance 1; 300 runs of 5,000 " + "rounds each", color=INK_2, fontsize=9) + fig.tight_layout(rect=(0, 0, 1, 0.95)) fig.savefig(OUT / "bandit_regret.png", dpi=150) plt.close(fig) def plot_compare(runs): fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 3.7)) - for name in ["etc_m10", "etc_m50", "ucb_c3"]: + for name in ORDER: r = runs[name] ax1.plot(r["round"], r["mean"], color=COLORS[name], label=LABELS[name]) ax1.set_title("Mean regret") From 1a498ac9efe9ef413aeb12ddf581cbd6db732974 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 19 Sep 2026 10:09:49 +0200 Subject: [PATCH 31/34] Fix the mk_all check: executables out of the libraries `lake exe mk_all --check` wants the root files of `MetropolisHastings` and `Bandits` to be the import lists it generates. They carried a docstring, and the executables `MetropolisHastings/Main.lean` and `Bandits/Main.lean`, being in the library directories, were picked up as modules of the libraries. The executables move to the top level, as `RunMetropolisHastings.lean` and `RunBandits.lean`, like the other executables of the repository; the root files are regenerated by `mk_all`, and their overview moves to the docstrings of `MetropolisHastings.Defs` and `Bandits.Defs`. Co-Authored-By: Claude Opus 5 (1M context) --- Bandits.lean | 23 ------------------ Bandits/Defs.lean | 8 ++++++- MetropolisHastings.lean | 24 ------------------- MetropolisHastings/Defs.lean | 8 ++++++- Bandits/Main.lean => RunBandits.lean | 0 .../Main.lean => RunMetropolisHastings.lean | 0 lakefile.toml | 4 ++-- 7 files changed, 16 insertions(+), 51 deletions(-) rename Bandits/Main.lean => RunBandits.lean (100%) rename MetropolisHastings/Main.lean => RunMetropolisHastings.lean (100%) diff --git a/Bandits.lean b/Bandits.lean index 874242d..0c12b48 100644 --- a/Bandits.lean +++ b/Bandits.lean @@ -1,28 +1,5 @@ -/- -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 -- shake: keep-all --deprecated_module: ignore public import Bandits.Defs public import Bandits.EpsGreedy public import Bandits.Theory - -/-! -# A Gaussian bandit, written in `rdo`, proved and run - -* `Bandits.Defs`: one round of interaction, and `n` rounds, as `rdo` programs polymorphic in the - monad and the scalars; explore-then-commit and UCB as the algorithms choosing the arm. -* `Bandits.Theory`: read at `Measure`, the programs have the law of LeanMachineLearning's - interaction, so its regret bounds hold for them. -* `Bandits.EpsGreedy`: ε-greedy, a randomized algorithm whose policy is an `rdo` program; with - `alg_env_trace`, its internal draws give an exploration bound and a linear regret lower bound. - -To run them and draw the regret against the bounds, from the root of the repository: - -``` -lake exe bandits # 300 seeds × 5000 rounds; writes bandit_output/ -python3 scripts/bandit_plot.py # checks them against numpy, draws bandit_output/*.png -``` --/ diff --git a/Bandits/Defs.lean b/Bandits/Defs.lean index ac93ba1..74a67d2 100644 --- a/Bandits/Defs.lean +++ b/Bandits/Defs.lean @@ -20,7 +20,13 @@ arm pulled. One round of interaction is `banditStep`, and `banditRun n` plays `n The programs are polymorphic in the monad and in the scalars, as those of `RandomDo.Tactic.Computable.Polymorphic`: read at `Measure` and `ℝ`, they are what -`Bandits.Theory` proves things about; run at `RandM` and `Float`, they sample. +`Bandits.Theory` and `Bandits.EpsGreedy` prove things about; run at `RandM` and `Float`, they +sample. To run them and draw the regret against the bounds, from the root of the repository: + +``` +lake exe bandits # 300 seeds × 5000 rounds; writes bandit_output/ +python3 scripts/bandit_plot.py # checks them against numpy, draws bandit_output/*.png +``` The two algorithms, `etcArm` (explore-then-commit) and `ucbArm` (upper confidence bound), mirror the definitions of `Bandits.ETC.nextArm` and `Bandits.UCB.nextArm` in LeanMachineLearning, with the diff --git a/MetropolisHastings.lean b/MetropolisHastings.lean index 3e554b0..5dd4535 100644 --- a/MetropolisHastings.lean +++ b/MetropolisHastings.lean @@ -1,8 +1,3 @@ -/- -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 -- shake: keep-all --deprecated_module: ignore public import MetropolisHastings.Computable @@ -10,22 +5,3 @@ public import MetropolisHastings.Defs public import MetropolisHastings.Polymorphic public import MetropolisHastings.Targets public import MetropolisHastings.Theory - -/-! -# Random-walk Metropolis–Hastings, written in `rdo`, proved and run - -* `MetropolisHastings.Defs`: the algorithm, as `rdo` programs over the Giry monad. -* `MetropolisHastings.Theory`: they are Markov kernels, satisfy detailed balance with respect to the - target, and leave it invariant after any number of steps. -* `MetropolisHastings.Computable`: the samplers `@[computable]` writes from them. -* `MetropolisHastings.Polymorphic`: the same algorithm, polymorphic in the monad; at `Measure` it - is the one of `Defs`, so the theorems carry over, and at `RandM` it samples. -* `MetropolisHastings.Targets`: two targets to run the chain on, for both routes. - -To run it and draw the plots, from the root of the repository: - -``` -lake exe mh # runs both samplers, checks they agree, writes mh_output/ -python3 scripts/mh_plot.py # checks them against numpy, draws mh_output/*.png -``` --/ diff --git a/MetropolisHastings/Defs.lean b/MetropolisHastings/Defs.lean index 30c9103..e60eab3 100644 --- a/MetropolisHastings/Defs.lean +++ b/MetropolisHastings/Defs.lean @@ -17,7 +17,13 @@ chain stays at `x`. The chain runs `n` such steps from `x₀`. Both programs denote Markov kernels, written over the Giry monad. `MetropolisHastings.Theory` proves what they satisfy, `MetropolisHastings.Computable` and `MetropolisHastings.Polymorphic` turn -them into programs that run. +them into programs that run, and `MetropolisHastings.Targets` gives two targets to run them on. +To run them and draw the plots, from the root of the repository: + +``` +lake exe mh # runs both samplers, checks they agree, writes mh_output/ +python3 scripts/mh_plot.py # checks them against numpy, draws mh_output/*.png +``` ## Main definitions diff --git a/Bandits/Main.lean b/RunBandits.lean similarity index 100% rename from Bandits/Main.lean rename to RunBandits.lean diff --git a/MetropolisHastings/Main.lean b/RunMetropolisHastings.lean similarity index 100% rename from MetropolisHastings/Main.lean rename to RunMetropolisHastings.lean diff --git a/lakefile.toml b/lakefile.toml index e5a9ab7..b5bc89c 100644 --- a/lakefile.toml +++ b/lakefile.toml @@ -46,8 +46,8 @@ root = "Polymorphic" [[lean_exe]] name = "mh" -root = "MetropolisHastings.Main" +root = "RunMetropolisHastings" [[lean_exe]] name = "bandits" -root = "Bandits.Main" +root = "RunBandits" From dcf6664e6e144ea3df04efde124953d098a799e0 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 19 Sep 2026 10:28:27 +0200 Subject: [PATCH 32/34] Build and lint the examples in CI `MetropolisHastings` and `Bandits` become default targets, so that `lake build` and `lake lint`, and hence CI, build and lint them along with `RandomDo`. Co-Authored-By: Claude Opus 5 (1M context) --- lakefile.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lakefile.toml b/lakefile.toml index b5bc89c..99149eb 100644 --- a/lakefile.toml +++ b/lakefile.toml @@ -1,5 +1,5 @@ name = "RandomDo" -defaultTargets = ["RandomDo"] +defaultTargets = ["RandomDo", "MetropolisHastings", "Bandits"] lintDriver = "batteries/runLinter" testDriver = "Test" From 3909646635dba62d39564177e9f5714b0e4e6e30 Mon Sep 17 00:00:00 2001 From: Remy Degenne Date: Sat, 19 Sep 2026 10:57:21 +0200 Subject: [PATCH 33/34] Bandits: state regret with LML's gap, move general lemmas to ForMathlib/ForLML MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit pseudoRegret takes the reward kernel, as LML's regret does, and the ETC, UCB and ε-greedy bounds are stated with LML's gap, replacing gapOf. General lemmas move out of Bandits.Theory into files mirroring their upstream homes, generalized where the statement allows: RandomDo/ForMathlib (map_compProd_eq_bind, bind_map, Measurable.finSnoc, hasSubgaussianMGF_gaussianReal) and RandomDo/ForLML (pullCount'_snoc, sumRewards'_snoc, IT.hist_succ_eq_snoc, IT.map_hist_succ). The Vector measurability lemmas join RandomDo.Measurable. Duplicates are removed: RDo.map_compProd and RDo.bind_map of Trace.lean are now the ForMathlib lemmas, and LML's Measure.dirac_compProd replaces dirac_compProd_eq_map. Co-Authored-By: Claude Opus 5 (1M context) --- Bandits/EpsGreedy.lean | 6 +- Bandits/Theory.lean | 202 ++++-------------- RandomDo.lean | 7 + .../SequentialLearning/FiniteActions.lean | 29 +++ .../IonescuTulceaSpace.lean | 56 +++++ .../ForLML/SequentialLearning/SumRewards.lean | 29 +++ .../MeasurableSpace/Constructions.lean | 31 +++ .../MeasureTheory/Measure/GiryMonad.lean | 31 +++ .../Kernel/Composition/MeasureComp.lean | 37 ++++ .../Probability/Moments/SubGaussian.lean | 39 ++++ RandomDo/Measurable.lean | 20 ++ RandomDo/Probability/Trace.lean | 29 +-- 12 files changed, 329 insertions(+), 187 deletions(-) create mode 100644 RandomDo/ForLML/SequentialLearning/FiniteActions.lean create mode 100644 RandomDo/ForLML/SequentialLearning/IonescuTulceaSpace.lean create mode 100644 RandomDo/ForLML/SequentialLearning/SumRewards.lean create mode 100644 RandomDo/ForMathlib/MeasureTheory/MeasurableSpace/Constructions.lean create mode 100644 RandomDo/ForMathlib/MeasureTheory/Measure/GiryMonad.lean create mode 100644 RandomDo/ForMathlib/Probability/Kernel/Composition/MeasureComp.lean create mode 100644 RandomDo/ForMathlib/Probability/Moments/SubGaussian.lean diff --git a/Bandits/EpsGreedy.lean b/Bandits/EpsGreedy.lean index 55afb20..b39fbbc 100644 --- a/Bandits/EpsGreedy.lean +++ b/Bandits/EpsGreedy.lean @@ -209,13 +209,13 @@ lemma policy_eq_epsGreedyArm (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : /-- **ε-greedy has linear regret**, for the program that runs: the expected pseudo-regret of `banditRunRand` with `epsGreedyArm`, against Gaussian arms, is at least `n ε / K ∑ₐ Δₐ`. -/ theorem le_integral_regret_banditRunRand (μ : Fin K → ℝ) (σ2 : ℝ≥0) (n : ℕ) : - (n : ℝ) * ((ε : ℝ) / K) * ∑ a, gapOf μ a - ≤ ∫ s, pseudoRegret μ s + (n : ℝ) * ((ε : ℝ) / K) * ∑ a, gap (arms μ σ2) a + ≤ ∫ s, pseudoRegret (arms μ σ2) s ∂(banditRunRand (m := Measure) (epsGreedyArm (m := Measure) (ε : ℝ)) μ σ2 n) := by rw [integral_pseudoRegret_banditRunRand μ σ2 (alg ε) _ (isMarkov_epsGreedyArm ε) (policy_eq_epsGreedyArm ε) n] have h := le_integral_regret ε (arms μ σ2) (IT.isAlgEnvSeq_trajMeasure (alg (K := K) ε) (stationaryEnv (arms μ σ2))) n - simpa only [gap_arms] using h + exact h end RDoBandit.EpsGreedy diff --git a/Bandits/Theory.lean b/Bandits/Theory.lean index 8173909..f2a8f78 100644 --- a/Bandits/Theory.lean +++ b/Bandits/Theory.lean @@ -6,6 +6,11 @@ Authors: Rémy Degenne module public import Bandits.Defs +public import RandomDo.ForLML.SequentialLearning.FiniteActions +public import RandomDo.ForLML.SequentialLearning.IonescuTulceaSpace +public import RandomDo.ForLML.SequentialLearning.SumRewards +public import RandomDo.ForMathlib.MeasureTheory.Measure.GiryMonad +public import RandomDo.ForMathlib.Probability.Moments.SubGaussian public import LeanMachineLearning.Online.Bandit.Algorithms.Regret.ETC public import LeanMachineLearning.Online.Bandit.Algorithms.Regret.UCB @@ -52,29 +57,6 @@ namespace RDoBandit variable {K : ℕ} [NeZero K] -section Vector - -variable {α : Type*} [MeasurableSpace α] {n : ℕ} - -@[fun_prop] -lemma measurable_vector_getElem (i : Fin n) : Measurable fun v : Vector α n ↦ v[i] := - (measurable_pi_apply i).comp Vector.measurableEquivTuple.measurable - -lemma measurable_vector_iff {β : Type*} [MeasurableSpace β] {f : β → Vector α n} : - Measurable f ↔ ∀ i : Fin n, Measurable fun b ↦ (f b)[i] := - ⟨fun hf i ↦ (measurable_vector_getElem i).comp hf, - fun h ↦ by - have h' : Measurable fun b ↦ Vector.ofFn fun i : Fin n ↦ (f b)[i] := - Vector.measurableEquivTuple.symm.measurable.comp (measurable_pi_iff.2 h) - simpa using h'⟩ - -@[fun_prop] -lemma measurable_vector_ofFn {β : Type*} [MeasurableSpace β] {f : β → Fin n → α} - (hf : ∀ i, Measurable fun b ↦ f b i) : Measurable fun b ↦ Vector.ofFn (f b) := - measurable_vector_iff.2 fun i ↦ by simpa using hf i - -end Vector - /-- The last arm of a history, arm `0` before the first round. -/ noncomputable def lastArm (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : Fin K := if hn : 0 < n then (h ⟨n - 1, by omega⟩).action else 0 @@ -101,18 +83,6 @@ lemma histState_zero (h : Hist Unit (Fin K) ℝ 0) : histState 0 h = State.init and_true] constructor <;> ext <;> simp [pullCount'_eq_sum, sumRewards'] -omit [NeZero K] in -lemma pullCount'_snoc (n : ℕ) (h : Hist Unit (Fin K) ℝ n) (a b : Fin K) (r : ℝ) : - pullCount' (n + 1) (Fin.snoc h ((), a, r)) b = pullCount' n h b + if a = b then 1 else 0 := by - rw [pullCount'_eq_sum, pullCount'_eq_sum, Fin.sum_univ_castSucc] - simp [Fin.snoc_castSucc, Fin.snoc_last] - -omit [NeZero K] in -lemma sumRewards'_snoc (n : ℕ) (h : Hist Unit (Fin K) ℝ n) (a b : Fin K) (r : ℝ) : - sumRewards' (n + 1) (Fin.snoc h ((), a, r)) b = sumRewards' n h b + if a = b then r else 0 := by - rw [sumRewards', sumRewards', Fin.sum_univ_castSucc] - simp [Fin.snoc_castSucc, Fin.snoc_last] - lemma histState_snoc (n : ℕ) (h : Hist Unit (Fin K) ℝ n) (a : Fin K) (r : ℝ) : histState (n + 1) (Fin.snoc h ((), a, r)) = (histState n h).update a r := by simp only [histState, State.update, Prod.mk.injEq] @@ -150,44 +120,6 @@ lemma ucbArm_histState (c : ℝ) (n : ℕ) (h : Hist Unit (Fin K) ℝ n) : · simp only [histState, Fin.getElem_fin, Vector.getElem_ofFn, Fin.eta] rfl -section Measures - -open MeasurableSpacePure MeasurableSpaceBind - -variable {α β γ : Type*} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] - -/-- Mapping a composition-product is binding the kernel, mapped along the section. -/ -lemma map_compProd_eq_bind (ρ : Measure α) [SFinite ρ] (κ : Kernel α β) [IsSFiniteKernel κ] - {F : α × β → γ} (hF : Measurable F) : - (ρ ⊗ₘ κ).map F = ρ.bind fun a ↦ (κ a).map fun b ↦ F (a, b) := by - have hmap (a : α) {t : Set γ} (ht : MeasurableSet t) : - ((κ a).map fun b ↦ F (a, b)) t = κ a (Prod.mk a ⁻¹' (F ⁻¹' t)) := - Measure.map_apply (hF.comp measurable_prodMk_left) ht - have hmeas : Measurable fun a ↦ (κ a).map fun b ↦ F (a, b) := by - refine Measure.measurable_of_measurable_coe _ fun t ht ↦ ?_ - simp_rw [hmap _ ht] - exact Kernel.measurable_kernel_prodMk_left (hF ht) - ext s hs - rw [Measure.map_apply hF hs, Measure.compProd_apply (hF hs), - Measure.bind_apply hs hmeas.aemeasurable] - simp_rw [hmap _ hs] - -lemma dirac_compProd_eq_map [MeasurableSingletonClass α] (a : α) (κ : Kernel α β) - [IsSFiniteKernel κ] : Measure.dirac a ⊗ₘ κ = (κ a).map (Prod.mk a) := by - ext s hs - rw [Measure.dirac_compProd_apply hs, Measure.map_apply measurable_prodMk_left hs] - -/-- Binding after mapping is binding the composite. -/ -lemma bind_map_eq (ρ : Measure α) {f : α → β} (hf : Measurable f) {g : β → Measure γ} - (hg : Measurable g) : (ρ.map f).bind g = ρ.bind (g ∘ f) := by - ext s hs - have hgs : Measurable fun b ↦ g b s := (Measure.measurable_coe hs).comp hg - rw [Measure.bind_apply hs hg.aemeasurable, Measure.bind_apply hs (hg.comp hf).aemeasurable, - lintegral_map hgs hf] - rfl - -end Measures - section Arms variable (μ : Fin K → ℝ) (σ2 : ℝ≥0) @@ -206,27 +138,17 @@ omit [NeZero K] in lemma integral_arms (a : Fin K) : ∫ x, x ∂(arms μ σ2 a) = μ a := by simp [arms_apply, integral_id_gaussianReal] -/-- The gap of an arm: how much less its mean is than the best one. -/ -noncomputable def gapOf (a : Fin K) : ℝ := (⨆ i, μ i) - μ a - omit [NeZero K] in -lemma gap_arms (a : Fin K) : gap (arms μ σ2) a = gapOf μ a := by - simp [gap, gapOf, integral_arms] +/-- The gap of an arm: how much less its mean is than the best one. -/ +lemma gap_arms (a : Fin K) : gap (arms μ σ2) a = (⨆ i, μ i) - μ a := by + simp [gap, integral_arms] omit [NeZero K] in /-- Gaussian rewards are sub-Gaussian, with variance proxy their variance. -/ lemma hasSubgaussianMGF_arms (a : Fin K) : HasSubgaussianMGF (fun x ↦ x - (arms μ σ2 a)[id]) σ2 (arms μ σ2 a) := by rw [show (arms μ σ2 a)[id] = μ a from integral_arms μ σ2 a, arms_apply] - refine ⟨fun t ↦ ?_, fun t ↦ ?_⟩ - · have := (integrable_exp_mul_gaussianReal (μ := μ a) (v := σ2) t).const_mul - (Real.exp (-(t * μ a))) - refine this.congr (Filter.Eventually.of_forall fun x ↦ ?_) - simp only - rw [← Real.exp_add] - ring_nf - · rw [mgf_gaussianReal ⟨by fun_prop, gaussianReal_map_sub_const (μ a)⟩ t] - simp + exact hasSubgaussianMGF_gaussianReal (μ a) σ2 end Arms @@ -294,43 +216,6 @@ lemma measurable_banditStep {arm : ℕ → State K ℝ → Fin K} {n : ℕ} (har simp_rw [hmap] exact Kernel.measurable_kernel_prodMk_left (hu ht) -omit [NeZero K] in -lemma measurable_snoc {X : Type*} [MeasurableSpace X] (n : ℕ) : - Measurable fun p : (Fin n → X) × X ↦ (Fin.snoc p.1 p.2 : Fin (n + 1) → X) := by - refine Measurable.of_eval fun i ↦ ?_ - refine Fin.lastCases ?_ (fun j ↦ ?_) i - · simp only [Fin.snoc_last] - exact measurable_snd - · simp only [Fin.snoc_castSucc] - exact (measurable_pi_apply j).comp measurable_fst - -omit [NeZero K] in -lemma hist_succ_eq (n : ℕ) : - IT.hist (𝓞 := Unit) (𝓐 := Fin K) (𝓨 := ℝ) (n + 1) - = fun ω ↦ Fin.snoc (IT.hist n ω) (IT.step n ω) := by - funext ω i - refine Fin.lastCases ?_ (fun j ↦ ?_) i - · simp [IT.hist, IT.step] - · simp [IT.hist] - -omit [NeZero K] in -/-- The history after `n + 1` rounds is the history after `n` rounds, followed by one round drawn -from the step kernel. -/ -lemma map_hist_succ {γ : Type*} [MeasurableSpace γ] (alg : Algorithm Unit (Fin K) ℝ) - (env : Environment Unit (Fin K) ℝ) (n : ℕ) {F : Hist Unit (Fin K) ℝ (n + 1) → γ} - (hF : Measurable F) : - (trajMeasure alg env).map (F ∘ IT.hist (n + 1)) - = ((trajMeasure alg env).map (IT.hist n)).bind - fun h ↦ (stepKernel alg env n h).map fun x ↦ F (Fin.snoc h x) := by - have e : F ∘ IT.hist (n + 1) - = (fun p ↦ F (Fin.snoc p.1 p.2)) ∘ (fun ω ↦ (IT.hist n ω, IT.step n ω)) := by - rw [hist_succ_eq] - rfl - rw [e, ← Measure.map_map (g := fun p ↦ F (Fin.snoc p.1 p.2)) - (f := fun ω ↦ (IT.hist n ω, IT.step n ω)) (hF.comp (measurable_snoc n)) (by fun_prop), - (IT.hasCondDistrib_step alg env n).map_eq, - map_compProd_eq_bind (F := fun p ↦ F (Fin.snoc p.1 p.2)) _ _ (hF.comp (measurable_snoc n))] - variable {nextA : (n : ℕ) → Hist Unit (Fin K) ℝ n × Unit → Fin K} {hnext : ∀ n, Measurable (nextA n)} @@ -343,7 +228,7 @@ lemma map_stepKernel {γ : Type*} [MeasurableSpace γ] (n : ℕ) (h : Hist Unit = (gaussianReal (μ (nextA n (h, ()))) σ2).map fun r ↦ G ((), nextA n (h, ()), r) := by rw [stepKernel_stationaryEnv, Kernel.compProd_apply_eq_compProd_sectR, Kernel.const_apply, Measure.dirac_unit_compProd, Kernel.sectR_apply, Kernel.compProd_apply_eq_compProd_sectR, - detAlgorithm_policy, Kernel.deterministic_apply, dirac_compProd_eq_map, Kernel.sectR_apply, + detAlgorithm_policy, Kernel.deterministic_apply, Measure.dirac_compProd, Kernel.sectR_apply, Kernel.prodMkLeft_apply, arms_apply, Measure.map_map hG measurable_prodMk_left, Measure.map_map (hG.comp measurable_prodMk_left) measurable_prodMk_left] rfl @@ -364,15 +249,14 @@ theorem banditRun_eq_map (arm : ℕ → State K ℝ → Fin K) (harm_meas : ∀ rw [this, Measure.map_const, measure_univ, one_smul] rfl | succ n ih => - rw [banditRun_succ, ih, map_hist_succ _ _ n (measurable_histState (n + 1)), + rw [banditRun_succ, ih, IT.map_hist_succ _ _ n (measurable_histState (n + 1)), ← Measure.map_map (measurable_histState n) (IT.measurable_hist n), - bind_map_eq _ (measurable_histState n) (measurable_banditStep (harm_meas n))] + Measure.bind_map _ (measurable_histState n) (measurable_banditStep (harm_meas n))] congr 1 funext h have hG : Measurable fun x ↦ histState (n + 1) (Fin.snoc h x) := - (measurable_histState (n + 1)).comp ((measurable_snoc n).comp - (measurable_const.prodMk measurable_id)) - rw [Function.comp_apply, banditStep_eq, map_stepKernel n h hG, harm] + (measurable_histState (n + 1)).comp (measurable_const.finSnoc measurable_id) + rw [banditStep_eq, map_stepKernel n h hG, harm] simp_rw [histState_snoc] end Law @@ -425,7 +309,7 @@ lemma map_stepKernel_bind {γ : Type*} [MeasurableSpace γ] (alg : Algorithm Uni rw [stepKernel_stationaryEnv, Kernel.compProd_apply_eq_compProd_sectR, Kernel.const_apply, Measure.dirac_unit_compProd, Kernel.sectR_apply, Kernel.compProd_apply_eq_compProd_sectR, Measure.map_map hG measurable_prodMk_left, - map_compProd_eq_bind _ _ (hG.comp measurable_prodMk_left)] + Measure.map_compProd_eq_bind _ _ (hG.comp measurable_prodMk_left)] rfl /-- **The law of the program**, for an algorithm drawing its arm. If, at every round, the policy @@ -445,15 +329,15 @@ theorem banditRunRand_eq_map (alg : Algorithm Unit (Fin K) ℝ) rw [this, Measure.map_const, measure_univ, one_smul] rfl | succ n ih => - rw [banditRunRand_succ, ih, map_hist_succ _ _ n (measurable_histState (n + 1)), + rw [banditRunRand_succ, ih, IT.map_hist_succ _ _ n (measurable_histState (n + 1)), ← Measure.map_map (measurable_histState n) (IT.measurable_hist n), - bind_map_eq _ (measurable_histState n) (isMarkov_banditStepRand (harm_markov n)).measurable] + Measure.bind_map _ (measurable_histState n) + (isMarkov_banditStepRand (harm_markov n)).measurable] congr 1 funext h have hG : Measurable fun x ↦ histState (n + 1) (Fin.snoc h x) := - (measurable_histState (n + 1)).comp ((measurable_snoc n).comp - (measurable_const.prodMk measurable_id)) - rw [Function.comp_apply, banditStepRand_eq, map_stepKernel_bind alg n h hG, harm] + (measurable_histState (n + 1)).comp (measurable_const.finSnoc measurable_id) + rw [banditStepRand_eq, map_stepKernel_bind alg n h hG, harm] simp_rw [histState_snoc] end RandomizedLaw @@ -462,22 +346,24 @@ section Regret variable (μ : Fin K → ℝ) (σ2 : ℝ≥0) -/-- The pseudo-regret of a state: the number of pulls of each arm, times its gap. -/ -noncomputable def pseudoRegret (s : State K ℝ) : ℝ := ∑ a, ((s.1[a] : ℕ) : ℝ) * gapOf μ a +/-- The pseudo-regret of a state against the rewards `ν`: the number of pulls of each arm, times its +gap. -/ +noncomputable def pseudoRegret (ν : Kernel (Fin K) ℝ) (s : State K ℝ) : ℝ := + ∑ a, ((s.1[a] : ℕ) : ℝ) * gap ν a omit [NeZero K] in @[fun_prop] -lemma measurable_pseudoRegret : Measurable (pseudoRegret (K := K) μ) := by +lemma measurable_pseudoRegret (ν : Kernel (Fin K) ℝ) : Measurable (pseudoRegret ν) := by unfold pseudoRegret refine Finset.measurable_sum _ fun a _ ↦ ?_ exact ((measurable_of_countable (fun k : ℕ ↦ (k : ℝ))).comp ((measurable_vector_getElem a).comp measurable_fst)).mul_const _ /-- The regret of the interaction is the pseudo-regret of the state of its history. -/ -lemma regret_eq_pseudoRegret (n : ℕ) (ω : ℕ → Round Unit (Fin K) ℝ) : - regret (arms μ σ2) IT.action n ω = pseudoRegret μ (histState n (IT.hist n ω)) := by +lemma regret_eq_pseudoRegret (ν : Kernel (Fin K) ℝ) (n : ℕ) (ω : ℕ → Round Unit (Fin K) ℝ) : + regret ν IT.action n ω = pseudoRegret ν (histState n (IT.hist n ω)) := by rw [regret_eq_sum_pullCount_mul_gap] - simp only [pseudoRegret, histState, Fin.getElem_fin, Vector.getElem_ofFn, gap_arms] + simp only [pseudoRegret, histState, Fin.getElem_fin, Vector.getElem_ofFn] congr with a rw [pullCount_eq_pullCount' (O := IT.obs) (R' := IT.feedback), IT.history_obs_action_feedback] @@ -488,13 +374,13 @@ variable {nextA : (n : ℕ) → Hist Unit (Fin K) ℝ n × Unit → Fin K} theorem integral_pseudoRegret_banditRun (arm : ℕ → State K ℝ → Fin K) (harm_meas : ∀ n, Measurable (arm n)) (harm : ∀ n h, arm n (histState n h) = nextA n (h, ())) (n : ℕ) : - ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) arm μ σ2 n) + ∫ s, pseudoRegret (arms μ σ2) s ∂(banditRun (m := Measure) arm μ σ2 n) = (trajMeasure (detAlgorithm nextA hnext) (stationaryEnv (arms μ σ2)))[ regret (arms μ σ2) IT.action n] := by rw [banditRun_eq_map (hnext := hnext) arm harm_meas harm n, - integral_map (by fun_prop) (measurable_pseudoRegret μ).aestronglyMeasurable] + integral_map (by fun_prop) (measurable_pseudoRegret _).aestronglyMeasurable] congr with ω - exact (regret_eq_pseudoRegret μ σ2 n ω).symm + exact (regret_eq_pseudoRegret _ n ω).symm omit [NeZero K] in @[fun_prop] @@ -522,38 +408,36 @@ lemma measurable_ucbArm (c : ℝ) (n : ℕ) : Measurable (ucbArm (K := K) (R := /-- **Regret of explore-then-commit.** The bound of `Bandits.ETC.regret_le`, for the program. -/ theorem integral_regret_etc_le {m : ℕ} (hm : m ≠ 0) {n : ℕ} (hn : K * m ≤ n) : - ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) (etcArm m) μ σ2 n) - ≤ ∑ a, gapOf μ a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gapOf μ a ^ 2 / (4 * σ2))) := by + ∫ s, pseudoRegret (arms μ σ2) s ∂(banditRun (m := Measure) (etcArm m) μ σ2 n) + ≤ ∑ a, gap (arms μ σ2) a + * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap (arms μ σ2) a ^ 2 / (4 * σ2))) := by have h := ETC.regret_le (IT.isAlgEnvSeq_trajMeasure (etcAlgorithm K m) (stationaryEnv (arms μ σ2))) (hasSubgaussianMGF_arms μ σ2) hm n hn - simp only [gap_arms] at h rw [integral_pseudoRegret_banditRun μ σ2 (hnext := fun n ↦ ETC.measurable_nextArm m n |>.comp measurable_fst) (etcArm m) (measurable_etcArm m) (fun n h ↦ etcArm_histState m n h) n] exact h /-- **Regret of UCB.** The bound of `Bandits.UCB.regret_le'`, for the program. -/ theorem integral_regret_ucb_le {c : ℝ} (hc : 0 < c) (hσ2 : σ2 ≠ 0) (n : ℕ) : - ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) (ucbArm c) μ σ2 n) - ≤ ∑ a, (8 * c * Real.log (n + 1) / gapOf μ a - + gapOf μ a * (2 + 2 * UCB.constSum (c / σ2) n)) := by + ∫ s, pseudoRegret (arms μ σ2) s ∂(banditRun (m := Measure) (ucbArm c) μ σ2 n) + ≤ ∑ a, (8 * c * Real.log (n + 1) / gap (arms μ σ2) a + + gap (arms μ σ2) a * (2 + 2 * UCB.constSum (c / σ2) n)) := by have h := UCB.regret_le' (IT.isAlgEnvSeq_trajMeasure (ucbAlgorithm K c) (stationaryEnv (arms μ σ2))) (hasSubgaussianMGF_arms μ σ2) hσ2 hc n - simp only [gap_arms] at h rw [integral_pseudoRegret_banditRun μ σ2 (hnext := fun n ↦ UCB.measurable_nextArm c n |>.comp measurable_fst) (ucbArm c) (measurable_ucbArm c) (fun n h ↦ ucbArm_histState c n h) n] exact h /-- **Regret of UCB**, for `c > 2 σ2`: logarithmic in the number of rounds. -/ theorem integral_regret_ucb_le_of_gt {c : ℝ} (hc : 2 * σ2 < c) (hσ2 : σ2 ≠ 0) (n : ℕ) : - ∫ s, pseudoRegret μ s ∂(banditRun (m := Measure) (ucbArm c) μ σ2 n) - ≤ ∑ a, (8 * c * Real.log (n + 1) / gapOf μ a - + gapOf μ a * (4 + 2 * σ2 / (c - 2 * σ2))) := by + ∫ s, pseudoRegret (arms μ σ2) s ∂(banditRun (m := Measure) (ucbArm c) μ σ2 n) + ≤ ∑ a, (8 * c * Real.log (n + 1) / gap (arms μ σ2) a + + gap (arms μ σ2) a * (4 + 2 * σ2 / (c - 2 * σ2))) := by have h := UCB.regret_le_of_gt_two' (IT.isAlgEnvSeq_trajMeasure (ucbAlgorithm K c) (stationaryEnv (arms μ σ2))) (hasSubgaussianMGF_arms μ σ2) hσ2 hc n - simp only [gap_arms] at h rw [integral_pseudoRegret_banditRun μ σ2 (hnext := fun n ↦ UCB.measurable_nextArm c n |>.comp measurable_fst) (ucbArm c) (measurable_ucbArm c) (fun n h ↦ ucbArm_histState c n h) n] exact h @@ -563,12 +447,12 @@ algorithm drawing its arm. -/ theorem integral_pseudoRegret_banditRunRand (alg : Algorithm Unit (Fin K) ℝ) (arm : ℕ → State K ℝ → Measure (Fin K)) (harm_markov : ∀ n, IsMarkov (arm n)) (harm : ∀ n h, alg.policy n (h, ()) = arm n (histState n h)) (n : ℕ) : - ∫ s, pseudoRegret μ s ∂(banditRunRand (m := Measure) arm μ σ2 n) + ∫ s, pseudoRegret (arms μ σ2) s ∂(banditRunRand (m := Measure) arm μ σ2 n) = (trajMeasure alg (stationaryEnv (arms μ σ2)))[regret (arms μ σ2) IT.action n] := by rw [banditRunRand_eq_map alg arm harm_markov harm n, - integral_map (by fun_prop) (measurable_pseudoRegret μ).aestronglyMeasurable] + integral_map (by fun_prop) (measurable_pseudoRegret _).aestronglyMeasurable] congr with ω - exact (regret_eq_pseudoRegret μ σ2 n ω).symm + exact (regret_eq_pseudoRegret _ n ω).symm end Regret diff --git a/RandomDo.lean b/RandomDo.lean index f50bfc4..032f45a 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -1,6 +1,13 @@ module -- shake: keep-all --deprecated_module: ignore +public import RandomDo.ForLML.SequentialLearning.FiniteActions +public import RandomDo.ForLML.SequentialLearning.IonescuTulceaSpace +public import RandomDo.ForLML.SequentialLearning.SumRewards +public import RandomDo.ForMathlib.MeasureTheory.MeasurableSpace.Constructions public import RandomDo.ForMathlib.MeasureTheory.MeasurableSpace.Embedding +public import RandomDo.ForMathlib.MeasureTheory.Measure.GiryMonad +public import RandomDo.ForMathlib.Probability.Kernel.Composition.MeasureComp +public import RandomDo.ForMathlib.Probability.Moments.SubGaussian public import RandomDo.Measurable public import RandomDo.Monad.ForInInstances public import RandomDo.Monad.Instances diff --git a/RandomDo/ForLML/SequentialLearning/FiniteActions.lean b/RandomDo/ForLML/SequentialLearning/FiniteActions.lean new file mode 100644 index 0000000..077ccd8 --- /dev/null +++ b/RandomDo/ForLML/SequentialLearning/FiniteActions.lean @@ -0,0 +1,29 @@ +/- +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 LeanMachineLearning.SequentialLearning.FiniteActions + +/-! +# The number of pulls of an action, one round later + +## Main results + +* `pullCount'_snoc`: the number of pulls of an action in a history extended by one round. +-/ + +@[expose] public section + +namespace Learning + +variable {𝓞 𝓐 𝓨 : Type*} [DecidableEq 𝓐] + +lemma pullCount'_snoc (n : ℕ) (h : Hist 𝓞 𝓐 𝓨 n) (o : 𝓞) (a b : 𝓐) (y : 𝓨) : + pullCount' (n + 1) (Fin.snoc h (o, a, y)) b = pullCount' n h b + if a = b then 1 else 0 := by + rw [pullCount'_eq_sum, pullCount'_eq_sum, Fin.sum_univ_castSucc] + simp [Fin.snoc_castSucc, Fin.snoc_last] + +end Learning diff --git a/RandomDo/ForLML/SequentialLearning/IonescuTulceaSpace.lean b/RandomDo/ForLML/SequentialLearning/IonescuTulceaSpace.lean new file mode 100644 index 0000000..2f84442 --- /dev/null +++ b/RandomDo/ForLML/SequentialLearning/IonescuTulceaSpace.lean @@ -0,0 +1,56 @@ +/- +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 LeanMachineLearning.SequentialLearning.IonescuTulceaSpace +public import RandomDo.ForMathlib.MeasureTheory.MeasurableSpace.Constructions +public import RandomDo.ForMathlib.Probability.Kernel.Composition.MeasureComp + +/-! +# The history of the interaction, one round later + +## Main results + +* `IT.hist_succ_eq_snoc`: the history before time `n + 1` is the history before time `n`, followed + by the round at time `n`. +* `IT.map_hist_succ`: under `trajMeasure alg env`, the law of a function of the history before + time `n + 1` is the law of the history before time `n`, bound to one round drawn from the step + kernel. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory + +namespace Learning.IT + +variable {𝓞 𝓐 𝓨 : Type*} {m𝓞 : MeasurableSpace 𝓞} {m𝓐 : MeasurableSpace 𝓐} + {m𝓨 : MeasurableSpace 𝓨} + +lemma hist_succ_eq_snoc (n : ℕ) : + hist (𝓞 := 𝓞) (𝓐 := 𝓐) (𝓨 := 𝓨) (n + 1) = fun ω ↦ Fin.snoc (hist n ω) (step n ω) := by + funext ω i + refine Fin.lastCases ?_ (fun j ↦ ?_) i + · simp [hist, step] + · simp [hist] + +/-- The history before time `n + 1` is the history before time `n`, followed by one round drawn +from the step kernel. -/ +lemma map_hist_succ {γ : Type*} [MeasurableSpace γ] (alg : Algorithm 𝓞 𝓐 𝓨) + (env : Environment 𝓞 𝓐 𝓨) (n : ℕ) {F : Hist 𝓞 𝓐 𝓨 (n + 1) → γ} (hF : Measurable F) : + (trajMeasure alg env).map (F ∘ hist (n + 1)) + = ((trajMeasure alg env).map (hist n)).bind + fun h ↦ (stepKernel alg env n h).map fun x ↦ F (Fin.snoc h x) := by + have hsnoc : Measurable fun p : Hist 𝓞 𝓐 𝓨 n × Round 𝓞 𝓐 𝓨 ↦ F (Fin.snoc p.1 p.2) := + hF.comp (measurable_fst.finSnoc measurable_snd) + have e : F ∘ hist (n + 1) + = (fun p ↦ F (Fin.snoc p.1 p.2)) ∘ (fun ω ↦ (hist n ω, step n ω)) := by + rw [hist_succ_eq_snoc] + rfl + rw [e, ← Measure.map_map hsnoc (by fun_prop), (hasCondDistrib_step alg env n).map_eq, + Measure.map_compProd_eq_bind _ _ hsnoc] + +end Learning.IT diff --git a/RandomDo/ForLML/SequentialLearning/SumRewards.lean b/RandomDo/ForLML/SequentialLearning/SumRewards.lean new file mode 100644 index 0000000..2abfbf3 --- /dev/null +++ b/RandomDo/ForLML/SequentialLearning/SumRewards.lean @@ -0,0 +1,29 @@ +/- +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 LeanMachineLearning.SequentialLearning.SumRewards + +/-! +# The sum of the rewards of an action, one round later + +## Main results + +* `sumRewards'_snoc`: the sum of the rewards of an action in a history extended by one round. +-/ + +@[expose] public section + +namespace Learning + +variable {𝓞 𝓐 𝓨 : Type*} [DecidableEq 𝓐] [AddCommGroup 𝓨] + +lemma sumRewards'_snoc (n : ℕ) (h : Hist 𝓞 𝓐 𝓨 n) (o : 𝓞) (a b : 𝓐) (y : 𝓨) : + sumRewards' (n + 1) (Fin.snoc h (o, a, y)) b = sumRewards' n h b + if a = b then y else 0 := by + rw [sumRewards', sumRewards', Fin.sum_univ_castSucc] + simp [Fin.snoc_castSucc, Fin.snoc_last] + +end Learning diff --git a/RandomDo/ForMathlib/MeasureTheory/MeasurableSpace/Constructions.lean b/RandomDo/ForMathlib/MeasureTheory/MeasurableSpace/Constructions.lean new file mode 100644 index 0000000..c5ec9e4 --- /dev/null +++ b/RandomDo/ForMathlib/MeasureTheory/MeasurableSpace/Constructions.lean @@ -0,0 +1,31 @@ +/- +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.MeasureTheory.MeasurableSpace.Constructions + +/-! +# Measurability of `Fin.snoc` + +## Main results + +* `Measurable.finSnoc`: appending an element to a tuple is measurable, the measurable counterpart + of `Continuous.finSnoc`. +-/ + +@[expose] public section + +variable {α : Type*} [MeasurableSpace α] {n : ℕ} {X : Fin (n + 1) → Type*} + [∀ i, MeasurableSpace (X i)] + +@[fun_prop] +lemma Measurable.finSnoc {f : α → ∀ j : Fin n, X j.castSucc} {g : α → X (Fin.last n)} + (hf : Measurable f) (hg : Measurable g) : Measurable fun a ↦ Fin.snoc (f a) (g a) := by + refine measurable_pi_iff.2 fun i ↦ ?_ + refine Fin.lastCases ?_ (fun j ↦ ?_) i + · simpa only [Fin.snoc_last] using hg + · simp only [Fin.snoc_castSucc] + exact (measurable_pi_apply j).comp hf diff --git a/RandomDo/ForMathlib/MeasureTheory/Measure/GiryMonad.lean b/RandomDo/ForMathlib/MeasureTheory/Measure/GiryMonad.lean new file mode 100644 index 0000000..711ca7d --- /dev/null +++ b/RandomDo/ForMathlib/MeasureTheory/Measure/GiryMonad.lean @@ -0,0 +1,31 @@ +/- +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.MeasureTheory.Measure.GiryMonad + +/-! +# Binding a pushforward measure + +## Main results + +* `Measure.bind_map`: binding after mapping is binding the composite, the measure counterpart of + `PMF.bind_map`. +-/ + +@[expose] public section + +namespace MeasureTheory.Measure + +variable {α β γ : Type*} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +/-- `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, map_map hk hf] + rfl + +end MeasureTheory.Measure diff --git a/RandomDo/ForMathlib/Probability/Kernel/Composition/MeasureComp.lean b/RandomDo/ForMathlib/Probability/Kernel/Composition/MeasureComp.lean new file mode 100644 index 0000000..22c1d06 --- /dev/null +++ b/RandomDo/ForMathlib/Probability/Kernel/Composition/MeasureComp.lean @@ -0,0 +1,37 @@ +/- +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.Probability.Kernel.Composition.MeasureComp + +/-! +# Pushforwards of the composition-product of a measure and a kernel + +## Main results + +* `Measure.map_compProd_eq_bind`: mapping `μ ⊗ₘ κ` along `g` is binding `κ a` mapped along the + section `g (a, ·)`. +-/ + +@[expose] public section + +open ProbabilityTheory + +namespace MeasureTheory.Measure + +variable {α β γ : Type*} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +/-- Mapping a composition-product is binding the kernel, mapped along the section. -/ +lemma map_compProd_eq_bind (μ : Measure α) [SFinite μ] (κ : Kernel α β) [IsSFiniteKernel κ] + {g : α × β → γ} (hg : Measurable g) : + (μ ⊗ₘ κ).map g = μ.bind fun a ↦ (κ a).map fun b ↦ g (a, b) := by + rw [compProd_eq_comp_prod, map_comp _ _ hg] + refine bind_congr_right (.of_forall fun a ↦ ?_) + rw [Kernel.map_apply _ hg, Kernel.prod_apply, Kernel.id_apply, dirac_prod, + map_map hg measurable_prodMk_left] + rfl + +end MeasureTheory.Measure diff --git a/RandomDo/ForMathlib/Probability/Moments/SubGaussian.lean b/RandomDo/ForMathlib/Probability/Moments/SubGaussian.lean new file mode 100644 index 0000000..968fae5 --- /dev/null +++ b/RandomDo/ForMathlib/Probability/Moments/SubGaussian.lean @@ -0,0 +1,39 @@ +/- +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.Probability.Distributions.Gaussian.Real +public import Mathlib.Probability.Moments.SubGaussian + +/-! +# Gaussian distributions are sub-Gaussian + +## Main results + +* `hasSubgaussianMGF_gaussianReal`: centered, a Gaussian is sub-Gaussian with variance proxy its + variance. +-/ + +@[expose] public section + +open MeasureTheory +open scoped NNReal + +namespace ProbabilityTheory + +/-- Centered, a Gaussian is sub-Gaussian with variance proxy its variance. -/ +lemma hasSubgaussianMGF_gaussianReal (μ : ℝ) (v : ℝ≥0) : + HasSubgaussianMGF (fun x ↦ x - μ) v (gaussianReal μ v) := by + refine ⟨fun t ↦ ?_, fun t ↦ ?_⟩ + · have := (integrable_exp_mul_gaussianReal (μ := μ) (v := v) t).const_mul (Real.exp (-(t * μ))) + refine this.congr (Filter.Eventually.of_forall fun x ↦ ?_) + simp only + rw [← Real.exp_add] + ring_nf + · rw [mgf_gaussianReal ⟨by fun_prop, gaussianReal_map_sub_const μ⟩ t] + simp + +end ProbabilityTheory diff --git a/RandomDo/Measurable.lean b/RandomDo/Measurable.lean index 2a4a1f5..9a4d1dd 100644 --- a/RandomDo/Measurable.lean +++ b/RandomDo/Measurable.lean @@ -23,6 +23,8 @@ This file contains results on the measurable structure of lists, arrays and vect are measurable. * `measurable_of_prodList`: a map out of `δ × List α` is measurable as soon as it is measurable on every stratum, which is how one reasons about a program taking a list as an argument. +* `Vector.measurableEquivTuple`, `measurable_vector_iff`: a vector is measurably a tuple, and a map + into vectors is measurable when each of its coordinates is. * `Measurable.ite_bool`: `if b a then f a else g a` is measurable, for a measurable `b` valued in `Bool`. -/ @@ -119,6 +121,24 @@ def Vector.measurableEquivTuple {n : ℕ} : Vector α n ≃ᵐ (Fin n → α) wh ext simp +@[fun_prop] +lemma measurable_vector_getElem {n : ℕ} (i : Fin n) : Measurable fun v : Vector α n ↦ v[i] := + (measurable_pi_apply i).comp Vector.measurableEquivTuple.measurable + +/-- A map into vectors is measurable when each of its coordinates is. -/ +lemma measurable_vector_iff {n : ℕ} {β : Type*} [MeasurableSpace β] {f : β → Vector α n} : + Measurable f ↔ ∀ i : Fin n, Measurable fun b ↦ (f b)[i] := + ⟨fun hf i ↦ (measurable_vector_getElem i).comp hf, + fun h ↦ by + have h' : Measurable fun b ↦ Vector.ofFn fun i : Fin n ↦ (f b)[i] := + Vector.measurableEquivTuple.symm.measurable.comp (measurable_pi_iff.2 h) + simpa using h'⟩ + +@[fun_prop] +lemma measurable_vector_ofFn {n : ℕ} {β : Type*} [MeasurableSpace β] {f : β → Fin n → α} + (hf : ∀ i, Measurable fun b ↦ f b i) : Measurable fun b ↦ Vector.ofFn (f b) := + measurable_vector_iff.2 fun i ↦ by simpa using hf i + instance instMeasurableSpaceOption : MeasurableSpace (Option α) := MeasurableSpace.map some inferInstance diff --git a/RandomDo/Probability/Trace.lean b/RandomDo/Probability/Trace.lean index 2f7fe2b..561e00f 100644 --- a/RandomDo/Probability/Trace.lean +++ b/RandomDo/Probability/Trace.lean @@ -5,6 +5,8 @@ Authors: Rémy Degenne -/ module +public import RandomDo.ForMathlib.MeasureTheory.Measure.GiryMonad +public import RandomDo.ForMathlib.Probability.Kernel.Composition.MeasureComp public import RandomDo.Probability.Record public import RandomDo.Tactic.IsMarkov.Defs public import RandomDo.Monad.Instances @@ -67,29 +69,6 @@ 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 δ] @@ -256,11 +235,11 @@ protected lemma bind {prog : γ → Measure α} {P : Kernel γ Ω} [IsSFiniteKer 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] + rw [Kernel.compProd_apply_eq_compProd_sectR, Measure.map_compProd_eq_bind _ _ 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 + exact (Measure.bind_map (P c) hout hcont').symm end HasTrace From e42eafa22022aa92227c61f4c16d01c4e1e77b38 Mon Sep 17 00:00:00 2001 From: Paulo Rauber Date: Wed, 23 Sep 2026 09:47:54 +0100 Subject: [PATCH 34/34] Add prototype --- notes/prototypes/TraceM.lean | 377 +++++++++++++++++++++++++++++++++++ 1 file changed, 377 insertions(+) create mode 100644 notes/prototypes/TraceM.lean diff --git a/notes/prototypes/TraceM.lean b/notes/prototypes/TraceM.lean new file mode 100644 index 0000000..e05b071 --- /dev/null +++ b/notes/prototypes/TraceM.lean @@ -0,0 +1,377 @@ +import RandomDo.Monad.Instances +import RandomDo.Measurable +import RandomDo.Tactic.IsMarkov.Elab +import RandomDo.Tactic.Computable.Polymorphic +import RandomDo.ForMathlib.Probability.Kernel.Composition.MeasureComp +import Mathlib.Probability.Distributions.Gaussian.Real + +/-! +# Prototype: the graded trace monad, with a proving and a running instance + +A program is written once over a `GradedMonad m`, drawing through `HasGaussianG`. Every +`let x ← …` records `x` under its own name: the grade `l` of `m l α` lists the recorded variables +with their names and is inferred from the program. Two instances: + +* `TraceM l α := Measure (Rec l × α)`, the joint law of the recorded variables and the result, + on the nested product `Rec l`. That product is internal: `gen_projections p` moves to the + `Fin`-indexed space Mathlib's probability library speaks, defining `p.Ω := Fin n → T`, + `p.toFin : Rec l → p.Ω`, `p.P : Measure p.Ω`, and `p.x : p.Ω → T` for each recorded `x` + (`fun ω ↦ ω i`). The laws are computed by `simp` with the monad laws at `Measure`. +* `SamplerM l α := RandPCG IO α`, which ignores the grade and samples. + +Not part of any library: open it in the editor, it elaborates with the project's imports. +-/ + +open Lean Meta Elab Term Command Do MeasureTheory ProbabilityTheory NumLean +open MeasurableSpaceBind MeasurableSpacePure LawfulMeasurableSpaceMonad + +/-! ## Grades and trace spaces -/ + +/-- A type with a measurable structure, as a grade entry. -/ +structure MType where + carrier : Type + [inst : MeasurableSpace carrier] + +attribute [instance] MType.inst + +/-- Print a grade entry as its carrier. -/ +@[app_unexpander MType.mk] def unexpandMType : PrettyPrinter.Unexpander + | `($_ $T $_) => `($T) + | `($_ $T) => `($T) + | _ => throw () + +/-- A grade: the recorded variables, in order. -/ +abbrev Grade := List (String × MType) + +/-- The trace space of a grade: one factor per recorded variable. -/ +@[reducible] def Rec : Grade → Type + | [] => PUnit + | (_, T) :: l => T.carrier × Rec l + +instance Rec.instMeasurableSpace : (l : Grade) → MeasurableSpace (Rec l) + | [] => inferInstance + | _ :: l => letI := Rec.instMeasurableSpace l; inferInstance + +def Rec.append : {l₁ l₂ : Grade} → Rec l₁ → Rec l₂ → Rec (l₁ ++ l₂) + | [], _, _, r₂ => r₂ + | _ :: _, _, (a, r₁), r₂ => (a, Rec.append r₁ r₂) + +lemma Rec.measurable_append_uncurry : ∀ {l₁ l₂ : Grade}, + Measurable fun p : Rec l₁ × Rec l₂ ↦ Rec.append p.1 p.2 + | [], _ => measurable_snd + | _ :: l₁, l₂ => by + change Measurable fun p : (_ × Rec l₁) × Rec l₂ ↦ (p.1.1, Rec.append p.1.2 p.2) + exact measurable_fst.fst.prodMk + (Rec.measurable_append_uncurry.comp (measurable_fst.snd.prodMk measurable_snd)) + +@[fun_prop] +lemma Rec.measurable_append {l₁ l₂ : Grade} {X : Type*} [MeasurableSpace X] + {f : X → Rec l₁} {g : X → Rec l₂} (hf : Measurable f) (hg : Measurable g) : + Measurable fun x ↦ Rec.append (f x) (g x) := + Rec.measurable_append_uncurry.comp (hf.prodMk hg) + +/-! ## The graded monad, and the classes a program draws through -/ + +/-- A monad graded by the recorded variables. -/ +class GradedMonad (m : Grade → (α : Type) → [MeasurableSpace α] → Type) where + gPure {α : Type} [MeasurableSpace α] : α → m [] α + gBind {l₁ l₂ : Grade} {α β : Type} [MeasurableSpace α] [MeasurableSpace β] : + m l₁ α → (α → m l₂ β) → m (l₁ ++ l₂) β + /-- Record `a` under the name `n`: one more coordinate of the trace. Inserted by the + elaborator after every `let x ← …`. -/ + grecord (n : String) {α : Type} [MeasurableSpace α] (a : α) : m [(n, MType.mk α)] Unit + +export GradedMonad (gPure gBind grecord) + +/-- A graded monad that can draw from a Gaussian with values in `R`. -/ +class HasGaussianG (m : Grade → (α : Type) → [MeasurableSpace α] → Type) (R : Type) + [MeasurableSpace R] where + gaussian : R → R → m [] R + +/-! ### The proving instance: the joint law -/ + +/-- The trace monad: the joint law of the recorded variables and the result. -/ +def TraceM (l : Grade) (α : Type) [MeasurableSpace α] : Type := Measure (Rec l × α) + +/- The operations are `rdo`-shaped programs at `Measure`, so that `is_markov` reads them. -/ + +noncomputable instance TraceM.gradedMonad : GradedMonad TraceM where + gPure a := (mPure ((), a) : Measure (Rec [] × _)) + gBind {l₁ l₂ _ _ _ _} x f := + ((x : Measure (Rec l₁ × _)) >>=ₘ fun p ↦ + (f p.2 : Measure (Rec l₂ × _)) >>=ₘ fun q ↦ mPure (Rec.append p.1 q.1, q.2) + : Measure (Rec (l₁ ++ l₂) × _)) + grecord n α _ a := (mPure ((a, ()), ()) : Measure (Rec [(n, MType.mk α)] × Unit)) + +noncomputable instance TraceM.hasGaussian : HasGaussianG TraceM ℝ where + gaussian μ v := (gaussianReal μ (Real.toNNReal v) >>=ₘ fun x ↦ mPure ((), x) : Measure (Rec [] × ℝ)) + +/-! ### The running instance: a sampler that ignores the grade -/ + +/-- The sampling monad, graded trivially. -/ +def SamplerM (_l : Grade) (α : Type) [MeasurableSpace α] : Type := RandPCG IO α + +instance SamplerM.gradedMonad : GradedMonad SamplerM where + gPure a := (pure a : RandPCG IO _) + gBind x f := (bind (x : RandPCG IO _) f : RandPCG IO _) + grecord _ _ _ _ := (pure () : RandPCG IO Unit) + +instance SamplerM.hasGaussian : HasGaussianG SamplerM Float where + gaussian μ v := (normal' μ v : RandPCG IO Float) + +/-! ## The `DoOps`: `gPure`/`gBind` of whatever `m` the expected type names -/ + +def gradeElem : Expr := mkApp2 (mkConst ``Prod [.zero, .one]) (mkConst ``String) (mkConst ``MType) + +def mkGM (m l α σ : Expr) : Expr := mkApp3 m l α σ + +def gradedOps : DoOps := { DoOps.default with + mkPureApp α e := do + let m := (← read).monadInfo.m + let e ← Term.ensureHasType α e + let σ ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α)) + let inst ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``GradedMonad) m)) + return mkAppN (mkConst ``GradedMonad.gPure) #[m, inst, α, σ, e] + mkBindApp α β e k := do + let m := (← read).monadInfo.m + Term.synthesizeSyntheticMVarsNoPostponing + let σα ← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α) + let σβ ← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) β) + let eType ← instantiateMVars (← inferType e) + let .app (.app (.app _ l₁) _) _ := eType.consumeMData | throwError "graded bind: {e} : {eType}" + let kType ← instantiateMVars (← inferType k) + let .forallE _ _ body _ := kType.consumeMData | throwError "graded bind: {k} : {kType}" + let .app (.app (.app _ l₂) _) _ := body.consumeMData | throwError "graded bind: {k} : {kType}" + if body.hasLooseBVars then throwError "graded bind: the grade {l₂} depends on the bound value" + let e ← Term.ensureHasType (mkGM m l₁ α σα) e + let k ← Term.ensureHasType (← mkArrow α (mkGM m l₂ β σβ)) k + let σα ← instantiateMVars σα + let σβ ← instantiateMVars σβ + let inst ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``GradedMonad) m)) + let l ← reduce (mkApp3 (mkConst ``List.append [.one]) gradeElem l₁ l₂) (skipTypes := false) + mkExpectedTypeHint (mkAppN (mkConst ``GradedMonad.gBind) #[m, inst, l₁, l₂, α, β, σα, σβ, e, k]) + (mkGM m l β σβ) + isPureApp? e := if e.isAppOfArity ``GradedMonad.gPure 5 then some (e.getArg! 4) else none + splitMonadApp? type := do + let .app mα _ := type.consumeMData | return none + let .app ml resultType := mα.consumeMData | return none + let .app m _ := ml.consumeMData | return none + unless ← isType resultType do return none + return some ({ m := m, u := 0, v := 0 }, resultType) + mkMonadApp α := do + let m := (← read).monadInfo.m + let l ← mkFreshExprMVar (mkConst ``Grade) + let σ ← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α) + return mkGM m l α σ } + +/-! ### Recording: every `let x ← …` records `x` + +The binder name is known where the bind is built, so recording is one more step in `mkBindApp`: +`let x ← e; k` becomes `gBind e (fun x ↦ gBind (grecord "x" x) (fun _ ↦ k x))`, for every binder +the user wrote. The elaborator's own binders (`__do_lift`, `__r`, `_`) are not recorded. -/ + +def recordingOps : DoOps := { gradedOps with + mkBindApp α β e k := do + let k ← instantiateMVars k + let .lam x _ _ _ := k | gradedOps.mkBindApp α β e k + if x.hasMacroScopes || x.isInternal || x == `_ then return ← gradedOps.mkBindApp α β e k + let σα ← instantiateMVars (← mkInstMVar (mkApp (mkConst ``MeasurableSpace [0]) α)) + let k' ← withLocalDeclD x α fun xv ↦ do + let recd := mkAppN (mkConst ``GradedMonad.grecord) + #[(← read).monadInfo.m, ← mkInstMVar (mkApp (mkConst ``GradedMonad) (← read).monadInfo.m), + mkStrLit x.toString, α, σα, xv] + let rest ← withLocalDeclD `__r (mkConst ``Unit) fun u ↦ mkLambdaFVars #[u] (k.beta #[xv]) + let inner ← gradedOps.mkBindApp (mkConst ``Unit) β recd rest + mkLambdaFVars #[xv] inner + gradedOps.mkBindApp α β e k' } + +syntax (name := gdoKind) "gdo" doSeq : term +@[term_elab gdoKind] def elabGdo : TermElab := fun stx et? => do + let `(gdo $doSeq) := stx | throwUnsupportedSyntax + elabDoWith recordingOps doSeq et? + +/-- `rdef p : m α := …` defines the program `p`, of type `m l α` for the grade `l` inferred from +the body. A `def` cannot infer a hole in its header from its body, so this expands to a `def` +with the ascription `(gdo … : m _ α)` in the body. -/ +macro "rdef " n:ident " : " m:ident α:term:max " := " body:doSeq : command => + `(def $n := (gdo $body : $m _ $α)) + +/-! ## Vectors: from the nested product to `Fin n → T` -/ + +@[fun_prop] +lemma Measurable.vecCons {X α : Type*} [MeasurableSpace X] [MeasurableSpace α] {n : ℕ} + {f : X → α} {g : X → Fin n → α} (hf : Measurable f) (hg : Measurable g) : + Measurable fun x ↦ Matrix.vecCons (f x) (g x) := + measurable_finCons.comp (hf.prodMk hg) + +@[fun_prop] +lemma measurable_vecEmpty {X α : Type*} [MeasurableSpace X] [MeasurableSpace α] : + Measurable fun _ : X ↦ (Matrix.vecEmpty : Fin 0 → α) := + measurable_const + +/-! ## Generating the named projections and the trace measure -/ + +partial def readGrade (l : Expr) : MetaM (List (String × Expr)) := do + match_expr l with + | List.nil _ => return [] + | List.cons _ hd tl => + let_expr Prod.mk _ _ n T := hd | throwError "not a grade entry: {hd}" + let .lit (.strVal s) := n | throwError "not a name literal: {n}" + let_expr MType.mk T _ := T | throwError "not a measurable type: {T}" + return (s, T) :: (← readGrade tl) + | _ => throwError "not a literal grade: {l}" + +/-- `gen_projections p`, for a program whose recorded variables all have the same type `T`, +defines the `Fin`-indexed trace space and everything on it: + +* `p.Ω := Fin n → T`, and `p.toFin : Rec l → p.Ω` with `p.measurable_toFin`; +* `p.P : Measure p.Ω`, the joint law of the recorded variables; +* `p.x : p.Ω → T`, `fun ω ↦ ω i`, with `p.measurable_x`, for each recorded `x` at position `i`. -/ +elab "gen_projections " n:ident : command => liftTermElabM do + let c ← realizeGlobalConstNoOverload n + let ty ← instantiateMVars (← getConstInfo c).type + let .app (.app (.app _ l) _) _ := ty | throwError "{c} is not a graded program: {ty}" + let l ← reduce l (skipTypes := false) + let ΩRec := mkApp (mkConst ``Rec) l + let entries ← readGrade l + let some (_, T) := entries.head? | throwError "{c} records nothing" + for (_, T') in entries do + unless ← isDefEq T T' do throwError "heterogeneous grade {l}: not supported by this prototype" + let k := entries.length + let finK := mkApp (mkConst ``Fin) (mkNatLit k) + let Ω ← mkArrow finK T + let define (name : Name) (type value : Expr) (compile := true) : TermElabM Unit := do + let decl := .defnDecl <| mkDefinitionValEx (c ++ name) [] type value .abbrev .safe [] + -- `P` is a measure, hence noncomputable: add it without compiling it. + if compile then addAndCompile decl else addDecl decl + enableRealizationsForConst (c ++ name) + logInfo m!"{c ++ name} : {type}" + let prove (name : Name) (type : Expr) (tac : TSyntax ``Lean.Parser.Tactic.tacticSeq) : + TermElabM Unit := do + let prf ← Term.elabTermAndSynthesize (← `(by $tac)) type + addDecl <| .thmDecl <| mkTheoremValEx (c ++ name) [] type (← instantiateMVars prf) [] + define `Ω (mkSort .one) Ω + -- `toFin ω = ![ω.1, ω.2.1, …]` + let toFin ← withLocalDeclD `ω ΩRec fun ω ↦ do + let mut coords := #[] + for i in [0:k] do + let mut e := ω + for _ in [0:i] do e ← mkAppM ``Prod.snd #[e] + coords := coords.push (← mkAppM ``Prod.fst #[e]) + let mut v ← mkAppOptM ``Matrix.vecEmpty #[T] + for e in coords.reverse do v ← mkAppM ``Matrix.vecCons #[e, v] + mkLambdaFVars #[ω] v + define `toFin (← mkArrow ΩRec Ω) toFin + prove `measurable_toFin (← mkAppM ``Measurable #[mkConst (c ++ `toFin)]) + (← `(tacticSeq| unfold $(mkIdent (c ++ `toFin)):ident; fun_prop)) + let P ← Term.elabTermAndSynthesize (← `(MeasureTheory.Measure.map $(mkIdent (c ++ `toFin)) + (MeasureTheory.Measure.map Prod.fst ($(mkIdent c) : MeasureTheory.Measure _)))) none + define `P (← inferType P) (← instantiateMVars P) (compile := false) + for (name, _) in entries, i in [0:k] do + let idx ← Term.elabTermAndSynthesize + (← `(($(Syntax.mkNumLit (toString i)) : Fin $(Syntax.mkNumLit (toString k))))) finK + let proj ← withLocalDeclD `ω Ω fun ω ↦ mkLambdaFVars #[ω] (mkApp ω idx) + define name.toName (← mkArrow Ω T) proj + prove (Name.mkSimple ("measurable_" ++ name)) (← mkAppM ``Measurable #[mkConst (c ++ name.toName)]) + (← `(tacticSeq| unfold $(mkIdent (c ++ name.toName)):ident; fun_prop)) + +/-! ## The program, written once -/ + +section +variable {m : Grade → (α : Type) → [MeasurableSpace α] → Type} [GradedMonad m] + {R : Type} [MeasurableSpace R] [Add R] [OfNat R 0] [OfNat R 1] [HasGaussianG m R] + +-- Draw `x ∼ 𝒩(0, 1)`, then `y ∼ 𝒩(x, 1)`, and return their sum. Both draws are recorded. +rdef sum2 : m R := + let x ← HasGaussianG.gaussian (m := m) (0 : R) 1 + let y ← HasGaussianG.gaussian (m := m) x 1 + return x + y + +end + +#check @sum2 + +/-! ## Running it -/ + +#eval IO.runRandPCGWith 42 (sum2 (m := SamplerM) (R := Float) : RandPCG IO Float) + +/-! ## Proving about it -/ + +/-- The program read as a joint law. -/ +noncomputable def sum2T := sum2 (m := TraceM) (R := ℝ) + +#check sum2T + +gen_projections sum2T + +/-- `Measure.bind` and `Measure.dirac` are the monad's `mBind` and `mPure`, syntactically. -/ +lemma Measure.bind_eq_mBind {α β : Type} [MeasurableSpace α] [MeasurableSpace β] (μ : Measure α) + (f : α → Measure β) : μ.bind f = μ >>=ₘ f := rfl + +lemma Measure.dirac_eq_mPure {α : Type} [MeasurableSpace α] (a : α) : + Measure.dirac a = (mPure a : Measure α) := rfl + +/-- A draw that is not used afterwards integrates out. -/ +lemma mBind_const {α β : Type} [MeasurableSpace α] [MeasurableSpace β] (μ : Measure α) + [IsProbabilityMeasure μ] (ν : Measure β) : (μ >>=ₘ fun _ ↦ ν) = ν := by + change μ.bind _ = ν + rw [Measure.bind_const, measure_univ, one_smul] + +/-- The side conditions of the monad laws at `Measure`: measurability of a continuation, which is +the Markov property of the program it is. -/ +macro "markov_side" : tactic => + `(tactic| first + | fun_prop + | (apply (config := { allowSynthFailures := true }) IsMarkov.measurable; is_markov)) + +/-- Unfold the program at `TraceM` down to `>>=ₘ`/`mPure`, and normalise with the monad laws. -/ +macro "trace_normalize" : tactic => + `(tactic| ( + -- `delta`, not `simp`: the unfolded grades are `[] ++ l`, equal to `l` only by unfolding + -- `List.append`, which `simp`'s congruence closure does not do. + delta sum2T sum2 TraceM.gradedMonad TraceM.hasGaussian + dsimp only [id] + simp only [Real.toNNReal_one] + simp (disch := fun_prop) only [← Measure.bind_dirac_eq_map] + simp only [Measure.bind_eq_mBind, Measure.dirac_eq_mPure] + simp (disch := markov_side) only [mBind_assoc, mPure_mBind, Rec.append])) + +/-- The internal joint law, on the nested product, in composition-product form. -/ +instance : IsMarkov fun x : ℝ ↦ gaussianReal x 1 := by is_markov + +/-- The kernel `x ↦ 𝒩(x, 1)`. -/ +noncomputable def gk : Kernel ℝ ℝ := IsMarkov.toKernel fun x : ℝ ↦ gaussianReal x 1 + +instance : IsMarkovKernel gk := by unfold gk; infer_instance + +@[simp] lemma gk_apply (x : ℝ) : gk x = gaussianReal x 1 := rfl + +lemma sum2T.rec_eq : + (sum2T : Measure _).map Prod.fst = (gaussianReal 0 1 ⊗ₘ gk).map fun p ↦ (p.1, (p.2, ())) := by + rw [Measure.map_compProd_eq_bind _ _ (by fun_prop)] + simp only [gk_apply] + trace_normalize + +/-- **The joint law of `(x, y)`**, on `Fin 2 → ℝ`: `x ∼ 𝒩(0, 1)`, then `y ∼ 𝒩(x, 1)`. -/ +theorem sum2T.P_eq : sum2T.P = (gaussianReal 0 1 ⊗ₘ gk).map fun p ↦ ![p.1, p.2] := by + delta sum2T.P + rw [sum2T.rec_eq, Measure.map_map sum2T.measurable_toFin (by fun_prop)] + rfl + +/-- The marginal law of `x`. -/ +theorem sum2T.map_x : sum2T.P.map sum2T.x = gaussianReal 0 1 := by + rw [sum2T.P_eq, Measure.map_map sum2T.measurable_x (by fun_prop)] + simp only [Function.comp_def, sum2T.x, Matrix.cons_val_zero] + rw [Measure.map_compProd_eq_bind _ _ (by fun_prop)] + simp only [gk_apply] + simp (disch := fun_prop) only [← Measure.bind_dirac_eq_map] + simp only [Measure.bind_eq_mBind, Measure.dirac_eq_mPure] + simp only [mBind_const, mBind_mPure] + +/-- `x` has law `𝒩(0, 1)` on the trace space. -/ +theorem sum2T.hasLaw_x : HasLaw sum2T.x (gaussianReal 0 1) sum2T.P := + ⟨sum2T.measurable_x.aemeasurable, sum2T.map_x⟩ + +#check sum2T.P +#check sum2T.hasLaw_x +#check sum2T.P_eq