diff --git a/README.md b/README.md index 5ba5618..6336611 100644 --- a/README.md +++ b/README.md @@ -1 +1,49 @@ -# Random-do notation \ No newline at end of file +# Random-do notation + +Write a probability program once using `rdo`, then interpret it in different measurable-space +monads. `Measure` gives its distribution; `RandomM Ω P` samples from a probability source while +preserving fresh source state. `SampleM Ω P` uses an infinite stream of independent `P` draws. + +To relate the two interpretations, add `rdo_program` to a polymorphic definition: + +```lean +import RandomDo + +open MeasureTheory +universe v + +def sumDraws {m : (α : Type) → [MeasurableSpace α] → Type v} + [MeasurableSpaceMonad m] (coin : m Bool) : ℕ → m ℕ + | 0 => rdo return 0 + | n + 1 => rdo + let b ← coin + let s ← sumDraws coin n + return b.toNat + s + +attribute [rdo_program] sumDraws + +example {Ω : Type*} [MeasurableSpace Ω] {P : Measure Ω} + [IsProbabilityMeasure P] (coin : RandomM Ω P Bool) (n : ℕ) : + (sumDraws coin n).law = sumDraws (m := Measure) coin.law n := + sumDraws.law coin n +``` + +The attribute leaves the original definition unchanged. It generates a recorded program +(`.program`), a certificate (`.valid` and `.certified`), bridges to the original interpretations +(`.sample_bridge` and `.measure_bridge`), and the resulting `.law` theorem. The proof uses +`RDo.Program.Certified.law`, which holds for every certified program. The `rdo_valid` tactic can +also construct certificates directly, leaving any unresolved measurability conditions as goals. + +The current automation supports returns, binds, measurable conditionals, sampler arguments, +and one-argument sampler families. A family argument adds a joint-measurability hypothesis to +the generated certificate and law theorem. `SampleM.ofKernel` provides this property for Markov +kernels; `SampleM.ofMeasure` supplies independent draws from probability measures on standard +Borel spaces. Both constructors use a stream of uniform draws from the unit interval. + +The attribute currently requires leading `{m} [MeasurableSpaceMonad m]` parameters and an +independent universe parameter for the monad's output, as above. It attempts induction on the +last explicit `Nat` argument. General recursion and loops need further support. Certification +fails if a proof obligation remains; no global measurable-evaluation assumption is required. + +See `Test/Program.lean` for continuous and kernel examples, and `Test/SampleM.lean` for independent +draws and the proof that a sum of `n` Bernoulli samples has the binomial distribution. diff --git a/RandomDo.lean b/RandomDo.lean index 1b94f05..6fdf66b 100644 --- a/RandomDo.lean +++ b/RandomDo.lean @@ -6,6 +6,8 @@ public import RandomDo.Monad.ForInInstances public import RandomDo.Monad.Instances public import RandomDo.Monad.MeasurableSpace public import RandomDo.Monad.Notation +public import RandomDo.Monad.Program +public import RandomDo.Monad.Sample public import RandomDo.NumLean.Distributions public import RandomDo.NumLean.PCG64 public import RandomDo.NumLean.SeedSequence @@ -16,3 +18,4 @@ public import RandomDo.Tactic.Elab public import RandomDo.Tactic.ForInStep public import RandomDo.Tactic.IsMarkov public import RandomDo.Tactic.Lemmas +public import RandomDo.Tactic.Program diff --git a/RandomDo/Monad/Instances.lean b/RandomDo/Monad/Instances.lean index fd064d9..ad4ceb8 100644 --- a/RandomDo/Monad/Instances.lean +++ b/RandomDo/Monad/Instances.lean @@ -11,7 +11,15 @@ public import Mathlib.Probability.ProductMeasure /-! # Instances for `MeasurableSpaceMonad` -**TODO** +`Measure` gives distribution semantics, and `PseudoRandomM` gives executable pseudorandom sampling. +`RandomM Ω P α` pairs a sampler with a proof that its returned value is independent of the remaining +source state, whose distribution is still `P`. + +For a probability source `P`, `RandomM` supports `pure`, measurable `map`, and `bindOfMeasurable` +for jointly measurable sampler families. Its `MeasurableSpaceMonad` instance checks this condition +inside `bind`. The `LawfulMeasurableSpaceMonad` instance additionally assumes +`RandomM.MeasurableEvalDomain Ω P`. Countable sources with measurable singletons satisfy this +assumption. -/ @@ -54,16 +62,445 @@ section RandomM open Function -/-- A monad for random number generation. -/ +/-- A monad for random variables. -/ structure RandomM (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) (α : Type u) [MeasurableSpace α] where - /-- Draws a value from a state of the source of randomness, and hands back the state left for the - next draw. -/ + /-- Return a value together with the remaining source state. -/ sample : Ω → α × Ω + /-- The remaining state has law `P` and is independent of the returned value. -/ measurePreserving : MeasurePreserving sample P ((Measure.map (Prod.fst ∘ sample) P).prod P) -/-- TODO -/ -abbrev SampleM (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) := - RandomM (ℕ → Ω) (Measure.infinitePi fun _ : ℕ ↦ P) +namespace RandomM + +variable {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} + {α β γ : Type u} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +@[ext] +theorem ext {x y : RandomM Ω P α} (h : ∀ ω, x.sample ω = y.sample ω) : x = y := by + cases x + cases y + congr + exact funext h + +theorem sample_injective : Injective (sample : RandomM Ω P α → Ω → α × Ω) := + fun _ _ h ↦ ext fun ω ↦ congrFun h ω + +/-- The distribution of the value returned by a sampler. -/ +noncomputable def law (x : RandomM Ω P α) : Measure α := + P.map (Prod.fst ∘ x.sample) + +@[fun_prop] +theorem measurable_sample (x : RandomM Ω P α) : Measurable x.sample := + x.measurePreserving.measurable + +@[fun_prop] +theorem measurable_fst_sample (x : RandomM Ω P α) : Measurable (Prod.fst ∘ x.sample) := + measurable_fst.comp x.measurable_sample + +@[fun_prop] +theorem measurable_snd_sample (x : RandomM Ω P α) : Measurable (Prod.snd ∘ x.sample) := + measurable_snd.comp x.measurable_sample + +/-- A sampler family indexed by a countable space with measurable singletons is jointly +measurable, even when the source of randomness is uncountable. -/ +@[fun_prop] +theorem measurable_sample_uncurry_of_countable {δ : Type*} [MeasurableSpace δ] + [Countable δ] [MeasurableSingletonClass δ] (f : δ → RandomM Ω P α) : + Measurable (fun p : δ × Ω ↦ (f p.1).sample p.2) := + measurable_from_prod_countable_right fun a ↦ (f a).measurable_sample + +@[simp] +theorem map_sample (x : RandomM Ω P α) : P.map x.sample = x.law.prod P := + x.measurePreserving.map_eq + +theorem measurePreserving_fst (x : RandomM Ω P α) : + MeasurePreserving (Prod.fst ∘ x.sample) P x.law := + ⟨x.measurable_fst_sample, rfl⟩ + +/-- The measurable structure induced by evaluating samplers at each fixed source state. -/ +instance : MeasurableSpace (RandomM Ω P α) := + MeasurableSpace.comap sample inferInstance + +theorem measurable_iff {δ : Type*} [MeasurableSpace δ] {f : δ → RandomM Ω P α} : + Measurable f ↔ ∀ ω, Measurable fun d ↦ (f d).sample ω := by + change Measurable[_, MeasurableSpace.comap sample inferInstance] f ↔ _ + rw [measurable_comap_iff, measurable_pi_iff] + rfl + +@[fun_prop] +theorem measurable_sample_apply (ω : Ω) : + Measurable (fun x : RandomM Ω P α ↦ x.sample ω) := + measurable_iff.1 measurable_id ω + +/-- Evaluation of a sampler at a varying source state is jointly measurable. -/ +class MeasurableEval (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) + (α : Type u) [MeasurableSpace α] : Prop where + /-- Evaluating a varying sampler at a varying source state is measurable. -/ + measurable_eval : Measurable (fun p : RandomM Ω P α × Ω ↦ p.1.sample p.2) + +/-- A source admits jointly measurable evaluation for every result space in universe `u`. + +This is an additional assumption on the source, not a consequence of the measurability of each +sampler. In particular, the pointwise measurable structure does not supply it for arbitrary +uncountable sources. -/ +class MeasurableEvalDomain (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) : Prop where + /-- Joint measurability of evaluation for each result space. -/ + hasMeasurableEval (α : Type u) [MeasurableSpace α] : MeasurableEval Ω P α + +instance [h : MeasurableEvalDomain.{u, w} Ω P] : MeasurableEval Ω P α := + h.hasMeasurableEval α + +instance [Countable Ω] [MeasurableSingletonClass Ω] : MeasurableEvalDomain Ω P where + hasMeasurableEval _ _ := + ⟨measurable_from_prod_countable_left fun ω ↦ measurable_sample_apply ω⟩ + +attribute [fun_prop] MeasurableEval.measurable_eval + +@[fun_prop] +theorem measurable_sample_uncurry [MeasurableEval Ω P α] + {δ : Type*} [MeasurableSpace δ] {f : δ → RandomM Ω P α} (hf : Measurable f) : + Measurable (fun p : δ × Ω ↦ (f p.1).sample p.2) := + MeasurableEval.measurable_eval.comp (hf.prodMap measurable_id) + +theorem measurable_iff_uncurry [MeasurableEval Ω P α] + {δ : Type*} [MeasurableSpace δ] {f : δ → RandomM Ω P α} : + Measurable f ↔ Measurable (fun p : δ × Ω ↦ (f p.1).sample p.2) := + ⟨measurable_sample_uncurry, fun hf ↦ measurable_iff.2 fun _ ↦ + hf.comp measurable_prodMk_right⟩ + +variable [IsProbabilityMeasure P] + +instance (x : RandomM Ω P α) : IsProbabilityMeasure x.law := + Measure.isProbabilityMeasure_map x.measurable_fst_sample.aemeasurable + +theorem measurePreserving_snd (x : RandomM Ω P α) : + MeasurePreserving (Prod.snd ∘ x.sample) P P := + (MeasureTheory.measurePreserving_snd (μ := x.law) (ν := P)).comp x.measurePreserving + +/-- Return a value without consuming any of the source of randomness. -/ +def pure (a : α) : RandomM Ω P α where + sample ω := (a, ω) + measurePreserving := by + refine ⟨measurable_const.prodMk measurable_id, ?_⟩ + simp [Function.comp_def, Measure.dirac_prod] + +@[simp] +theorem sample_pure (a : α) (ω : Ω) : (pure (P := P) a).sample ω = (a, ω) := rfl + +@[simp] +theorem law_pure (a : α) : (pure (P := P) a).law = Measure.dirac a := by + simp [law, pure, Function.comp_def] + +@[fun_prop] +theorem measurable_pure : Measurable (pure : α → RandomM Ω P α) := + measurable_iff.2 fun _ ↦ measurable_id.prodMk measurable_const + +section Map + +variable {β : Type v} {γ : Type*} [MeasurableSpace β] [MeasurableSpace γ] + +/-- Apply a measurable function to the returned value, retaining the remaining source state. -/ +def map (f : α → β) (hf : Measurable f) (x : RandomM Ω P α) : RandomM Ω P β where + sample := Prod.map f id ∘ x.sample + measurePreserving := by + have h := ((hf.measurePreserving x.law).prod (MeasurePreserving.id P)).comp + x.measurePreserving + convert h using 1 + rw [law, Measure.map_map hf x.measurable_fst_sample] + rfl + +@[simp] +theorem sample_map (f : α → β) (hf : Measurable f) (x : RandomM Ω P α) (ω : Ω) : + (map f hf x).sample ω = (f (x.sample ω).1, (x.sample ω).2) := rfl + +@[simp] +theorem law_map (f : α → β) (hf : Measurable f) (x : RandomM Ω P α) : + (map f hf x).law = x.law.map f := by + rw [law, law, Measure.map_map hf x.measurable_fst_sample] + rfl + +@[fun_prop] +theorem measurable_map (f : α → β) (hf : Measurable f) : + Measurable (map f hf : RandomM Ω P α → RandomM Ω P β) := by + apply measurable_iff.2 + intro ω + exact (hf.comp (measurable_sample_apply ω).fst).prodMk (measurable_sample_apply ω).snd + +@[simp] +theorem map_id (x : RandomM Ω P α) : map id measurable_id x = x := + ext fun _ ↦ rfl + +theorem map_map (f : α → β) (hf : Measurable f) (g : β → γ) (hg : Measurable g) + (x : RandomM Ω P α) : map g hg (map f hf x) = map (g ∘ f) (hg.comp hf) x := + ext fun _ ↦ rfl + +@[simp] +theorem map_pure (f : α → β) (hf : Measurable f) (a : α) : + map f hf (pure (P := P) a) = pure (f a) := + ext fun _ ↦ rfl + +end Map + +theorem measurable_law_of_uncurry {δ : Type*} [MeasurableSpace δ] + {f : δ → RandomM Ω P α} + (hf : Measurable (fun p : δ × Ω ↦ (f p.1).sample p.2)) : + Measurable (fun d ↦ (f d).law) := by + refine Measure.measurable_of_measurable_coe _ fun s hs ↦ ?_ + simp only [law, Measure.map_apply (measurable_fst_sample _) hs] + exact measurable_measure_prodMk_left (hf.fst hs) + +@[fun_prop] +theorem measurable_law [MeasurableEval Ω P α] : + Measurable (law : RandomM Ω P α → Measure α) := + measurable_law_of_uncurry MeasurableEval.measurable_eval + +private theorem map_bind_sample_prod (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) + {s : Set β} {t : Set Ω} (hs : MeasurableSet s) (ht : MeasurableSet t) : + P.map (fun ω ↦ (f (x.sample ω).1).sample (x.sample ω).2) (s ×ˢ t) = + (∫⁻ a, (f a).law s ∂x.law) * P t := by + change P.map ((fun p : α × Ω ↦ (f p.1).sample p.2) ∘ x.sample) (s ×ˢ t) = _ + rw [← Measure.map_map hf x.measurable_sample, x.map_sample, + Measure.map_apply hf (hs.prod ht), Measure.prod_apply (hf (hs.prod ht))] + have h (a : α) : + P (Prod.mk a ⁻¹' ((fun p : α × Ω ↦ (f p.1).sample p.2) ⁻¹' (s ×ˢ t))) = + (f a).law s * P t := by + change P ((f a).sample ⁻¹' (s ×ˢ t)) = _ + rw [← Measure.map_apply (f a).measurable_sample (hs.prod ht), (f a).map_sample, + Measure.prod_prod] + simp_rw [h] + exact lintegral_mul_const _ ((Measure.measurable_coe hs).comp + (measurable_law_of_uncurry hf)) + +private theorem map_fst_bind_sample (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) + {s : Set β} (hs : MeasurableSet s) : + P.map (fun ω ↦ ((f (x.sample ω).1).sample (x.sample ω).2).1) s = + ∫⁻ a, (f a).law s ∂x.law := by + change P.map (Prod.fst ∘ ((fun p : α × Ω ↦ (f p.1).sample p.2) ∘ x.sample)) s = _ + rw [← Measure.map_map measurable_fst (hf.comp x.measurable_sample), + Measure.map_apply measurable_fst hs] + simpa only [Function.comp_def, Set.prod_univ, measure_univ, mul_one] using + map_bind_sample_prod x f hf hs MeasurableSet.univ + +/-- Sequence samplers whose sampling functions are jointly measurable in the value and state. +This construction does not require a `MeasurableEvalDomain` instance. -/ +def bindOfMeasurable (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) : RandomM Ω P β where + sample ω := (f (x.sample ω).1).sample (x.sample ω).2 + measurePreserving := by + refine ⟨hf.comp x.measurable_sample, ?_⟩ + apply Eq.symm + apply Measure.prod_eq + intro s t hs ht + rw [map_bind_sample_prod x f hf hs ht] + exact congrArg (· * P t) (map_fst_bind_sample x f hf hs).symm + +@[simp] +theorem sample_bindOfMeasurable (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) (ω : Ω) : + (bindOfMeasurable x f hf).sample ω = (f (x.sample ω).1).sample (x.sample ω).2 := rfl + +@[simp] +theorem law_bindOfMeasurable (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) : + (bindOfMeasurable x f hf).law = x.law.bind (fun a ↦ (f a).law) := by + ext s hs + rw [Measure.bind_apply hs (measurable_law_of_uncurry hf).aemeasurable] + exact map_fst_bind_sample x f hf hs + +@[fun_prop] +theorem measurable_bindOfMeasurable (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) : + Measurable (fun x ↦ bindOfMeasurable x f hf) := + measurable_iff.2 fun ω ↦ hf.comp (measurable_sample_apply ω) + +/-- Sequence samplers when the continuation's sampling function is jointly measurable. +Otherwise select one of its samplers arbitrarily. The joint-measurability condition is checked +inside this definition, so constructing a bind does not require a `MeasurableEval` instance. -/ +noncomputable def bind (x : RandomM Ω P α) (f : α → RandomM Ω P β) : RandomM Ω P β := by + classical + exact if hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2) then bindOfMeasurable x f hf + else f (x.sample (Classical.choice (nonempty_of_isProbabilityMeasure P))).1 + +theorem bind_eq_bindOfMeasurable (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) : + bind x f = bindOfMeasurable x f hf := by + classical + simp only [bind, dite_eq_left hf] + +/-- The sampling equation needs only joint measurability of this particular continuation. -/ +theorem sample_bind_of_measurable (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) (ω : Ω) : + (bind x f).sample ω = (f (x.sample ω).1).sample (x.sample ω).2 := by + rw [bind_eq_bindOfMeasurable x f hf, sample_bindOfMeasurable] + +theorem law_bind_of_measurable (x : RandomM Ω P α) (f : α → RandomM Ω P β) + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) : + (bind x f).law = x.law.bind (fun a ↦ (f a).law) := by + rw [bind_eq_bindOfMeasurable x f hf, law_bindOfMeasurable] + +@[simp] +theorem sample_bind_of_countable [Countable α] [MeasurableSingletonClass α] + (x : RandomM Ω P α) (f : α → RandomM Ω P β) (ω : Ω) : + (bind x f).sample ω = (f (x.sample ω).1).sample (x.sample ω).2 := + sample_bind_of_measurable x f (measurable_sample_uncurry_of_countable f) ω + +@[simp] +theorem law_bind_of_countable [Countable α] [MeasurableSingletonClass α] + (x : RandomM Ω P α) (f : α → RandomM Ω P β) : + (bind x f).law = x.law.bind (fun a ↦ (f a).law) := + law_bind_of_measurable x f (measurable_sample_uncurry_of_countable f) + +@[fun_prop] +theorem measurable_bind_of_uncurry {f : α → RandomM Ω P β} + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) : + Measurable (fun x ↦ bind x f) := by + simpa only [bind_eq_bindOfMeasurable _ f hf] using measurable_bindOfMeasurable f hf + +@[simp] +theorem sample_bind [MeasurableEval Ω P β] (x : RandomM Ω P α) + (f : α → RandomM Ω P β) (hf : Measurable f) (ω : Ω) : + (bind x f).sample ω = (f (x.sample ω).1).sample (x.sample ω).2 := + sample_bind_of_measurable x f (measurable_sample_uncurry hf) ω + +@[simp] +theorem law_bind [MeasurableEval Ω P β] (x : RandomM Ω P α) + (f : α → RandomM Ω P β) (hf : Measurable f) : + (bind x f).law = x.law.bind (fun a ↦ (f a).law) := + law_bind_of_measurable x f (measurable_sample_uncurry hf) + +@[fun_prop] +theorem measurable_bind [MeasurableEval Ω P β] {f : α → RandomM Ω P β} + (hf : Measurable f) : Measurable (fun x ↦ bind x f) := + measurable_bind_of_uncurry (measurable_sample_uncurry hf) + +/-- Bind a measurable family of samplers to a jointly measurable family of continuations. -/ +@[fun_prop] +theorem measurable_bind₂ [MeasurableEval Ω P β] {δ : Type*} [MeasurableSpace δ] + {x : δ → RandomM Ω P α} {f : δ → α → RandomM Ω P β} + (hx : Measurable x) (hf : Measurable (uncurry f)) : + Measurable (fun d ↦ bind (x d) (f d)) := by + apply measurable_iff.2 + intro ω + have hs := (measurable_sample_apply ω).comp hx + have h := MeasurableEval.measurable_eval.comp + ((hf.comp (measurable_id.prodMk hs.fst)).prodMk hs.snd) + convert h using 1 + funext d + exact sample_bind (x d) (f d) hf.of_uncurry_left ω + +@[simp] +theorem pure_bind (a : α) (f : α → RandomM Ω P β) : bind (pure a) f = f a := by + classical + by_cases hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2) + · apply ext + intro ω + simp only [sample_bind_of_measurable _ _ hf, sample_pure] + · simp only [bind, dite_eq_right hf, sample_pure] + +@[simp] +theorem bind_pure (x : RandomM Ω P α) : bind x pure = x := by + apply ext + intro ω + exact sample_bind_of_measurable x pure measurable_id ω + +theorem bind_assoc_of_measurable (x : RandomM Ω P α) + {f : α → RandomM Ω P β} {g : β → RandomM Ω P γ} + (hf : Measurable (fun p : α × Ω ↦ (f p.1).sample p.2)) + (hg : Measurable (fun p : β × Ω ↦ (g p.1).sample p.2)) : + bind (bind x f) g = bind x (fun a ↦ bind (f a) g) := by + have hfg : Measurable (fun p : α × Ω ↦ (bind (f p.1) g).sample p.2) := by + simp_rw [sample_bind_of_measurable _ _ hg] + exact hg.comp hf + apply ext + intro ω + simp only [sample_bind_of_measurable _ _ hf, sample_bind_of_measurable _ _ hg, + sample_bind_of_measurable x (fun a ↦ bind (f a) g) hfg] + +theorem bind_assoc [MeasurableEval Ω P β] [MeasurableEval Ω P γ] (x : RandomM Ω P α) + {f : α → RandomM Ω P β} {g : β → RandomM Ω P γ} (hf : Measurable f) (hg : Measurable g) : + bind (bind x f) g = bind x (fun a ↦ bind (f a) g) := + bind_assoc_of_measurable x (measurable_sample_uncurry hf) (measurable_sample_uncurry hg) + +theorem bind_pure_comp (x : RandomM Ω P α) + {f : α → β} (hf : Measurable f) : bind x (fun a ↦ pure (f a)) = map f hf x := by + have hpf : Measurable (fun p : α × Ω ↦ (pure (P := P) (f p.1)).sample p.2) := + hf.prodMap measurable_id + apply ext + intro ω + simp only [sample_bind_of_measurable x (fun a ↦ pure (f a)) hpf, sample_pure, sample_map] + +noncomputable instance : + MeasurableSpaceMonad (RandomM Ω P : (α : Type u) → [MeasurableSpace α] → Type (max u w)) where + mPure := pure + mBind := bind + +instance [MeasurableEvalDomain.{u, w} Ω P] : + LawfulMeasurableSpaceMonad + (RandomM Ω P : (α : Type u) → [MeasurableSpace α] → Type (max u w)) where + mMap_const := rfl + id_mMap := bind_pure + measurable_mPure {α} {_} := @measurable_pure Ω _ P α _ _ + measurable_mBind := measurable_bind + mBind_mPure_comp _ _ := rfl + mPure_mBind a f _ := pure_bind a f + mBind_assoc x _ _ hf hg := bind_assoc x hf hg + +@[simp] +theorem sample_mPure (a : α) (ω : Ω) : (mPure a : RandomM Ω P α).sample ω = (a, ω) := rfl + +@[simp] +theorem law_mPure (a : α) : (mPure a : RandomM Ω P α).law = Measure.dirac a := law_pure a + +@[fun_prop] +theorem measurable_mBind₂ [MeasurableEval Ω P β] {δ : Type*} [MeasurableSpace δ] + {x : δ → RandomM Ω P α} {f : δ → α → RandomM Ω P β} + (hx : Measurable x) (hf : Measurable (uncurry f)) : + Measurable (fun d ↦ x d >>=ₘ f d) := measurable_bind₂ hx hf + +@[simp] +theorem sample_mBind [MeasurableEval Ω P β] (x : RandomM Ω P α) + (f : α → RandomM Ω P β) (hf : Measurable f) + (ω : Ω) : (x >>=ₘ f).sample ω = (f (x.sample ω).1).sample (x.sample ω).2 := + sample_bind x f hf ω + +@[simp] +theorem law_mBind [MeasurableEval Ω P β] (x : RandomM Ω P α) + (f : α → RandomM Ω P β) (hf : Measurable f) : + (x >>=ₘ f).law = x.law.bind (fun a ↦ (f a).law) := law_bind x f hf + +@[simp] +theorem sample_mBind_of_countable [Countable α] [MeasurableSingletonClass α] + (x : RandomM Ω P α) (f : α → RandomM Ω P β) (ω : Ω) : + (x >>=ₘ f).sample ω = (f (x.sample ω).1).sample (x.sample ω).2 := + sample_bind_of_countable x f ω + +@[simp] +theorem law_mBind_of_countable [Countable α] [MeasurableSingletonClass α] + (x : RandomM Ω P α) (f : α → RandomM Ω P β) : + (x >>=ₘ f).law = x.law.bind (fun a ↦ (f a).law) := law_bind_of_countable x f + +theorem mMap_eq_map (f : α → β) (hf : Measurable f) (x : RandomM Ω P α) : + f <$>ₘ x = map f hf x := bind_pure_comp x hf + +@[simp] +theorem sample_mMap (f : α → β) (hf : Measurable f) (x : RandomM Ω P α) (ω : Ω) : + (f <$>ₘ x).sample ω = (f (x.sample ω).1, (x.sample ω).2) := by + rw [mMap_eq_map f hf x, sample_map] + +@[simp] +theorem law_mMap (f : α → β) (hf : Measurable f) (x : RandomM Ω P α) : + (f <$>ₘ x).law = x.law.map f := by + rw [mMap_eq_map f hf x, law_map] + +end RandomM + +/-- Samplers driven by a sequence of independent source values with common law `P`. +For a probability measure `P`, the monad operations are available without a domain assumption. +The lawful instance still requires measurable evaluation on the full stream space `ℕ → Ω`. -/ +abbrev SampleM (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) + (α : Type u) [MeasurableSpace α] : Type (max u w) := + RandomM (ℕ → Ω) (Measure.infinitePi fun _ : ℕ ↦ P) α end RandomM diff --git a/RandomDo/Monad/Program.lean b/RandomDo/Monad/Program.lean new file mode 100644 index 0000000..f002822 --- /dev/null +++ b/RandomDo/Monad/Program.lean @@ -0,0 +1,159 @@ +/- +Copyright (c) 2026 David Ledvinka. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: David Ledvinka +-/ +module + +public import RandomDo.Monad.Sample + +/-! +# Certified representations of random programs + +`RDo.Program` records returns, primitive samplers, and binds. A polymorphic `rdo` definition +can be interpreted in this monad without changing its source. `RDo.Program.Valid` certifies +families of programs compositionally, including the joint measurability needed by bind. +Its soundness theorem relates the actual `RandomM` and `Measure` interpretations. +-/ + +@[expose] public section + +open MeasureTheory MeasurableSpacePure MeasurableSpaceBind + +namespace RDo + +universe u w + +/-- A program tree whose primitive operations are samplers on a fixed probability source. -/ +inductive Program (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) : + (α : Type u) → [MeasurableSpace α] → Type (max (u + 1) w) + /-- Return a value. -/ + | pure {α : Type u} [MeasurableSpace α] (a : α) : Program Ω P α + /-- Perform a primitive sampler. -/ + | sample {α : Type u} [MeasurableSpace α] (x : RandomM Ω P α) : Program Ω P α + /-- Sequence two programs. -/ + | bind {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + (x : Program Ω P α) (f : α → Program Ω P β) : Program Ω P β + +namespace Program + +variable {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} + +instance : MeasurableSpaceMonad (Program Ω P : + (α : Type u) → [MeasurableSpace α] → Type (max (u + 1) w)) where + mPure := pure + mBind := bind + +variable [IsProbabilityMeasure P] + +/-- Interpret a program using the original sampler operations. -/ +noncomputable def sampleSemantics {α : Type u} [MeasurableSpace α] : + Program Ω P α → RandomM Ω P α + | .pure a => mPure a + | .sample x => x + | @Program.bind _ _ _ _ _ _ _ x f => + RandomM.bind (sampleSemantics x) fun a ↦ sampleSemantics (f a) + +/-- Interpret a program using the laws of its primitives and the measure monad. -/ +noncomputable def measureSemantics {α : Type u} [MeasurableSpace α] : + Program Ω P α → Measure α + | .pure a => mPure a + | .sample x => x.law + | @Program.bind _ _ _ _ _ _ _ x f => + (measureSemantics x).bind fun a ↦ measureSemantics (f a) + +/-- A compositional certificate for a family of program trees. The input parameter includes +values returned by previous draws, so measurability is checked jointly through nested binds. -/ +inductive Valid : {γ α : Type u} → [MeasurableSpace γ] → [MeasurableSpace α] → + (γ → Program Ω P α) → Prop + /-- A measurable return expression. -/ + | pure {γ α : Type u} [MeasurableSpace γ] [MeasurableSpace α] + (f : γ → α) (hf : Measurable f) : Valid (fun c ↦ .pure (f c)) + /-- A jointly measurable family of primitive samplers. -/ + | sample {γ α : Type u} [MeasurableSpace γ] [MeasurableSpace α] + (f : γ → RandomM Ω P α) + (hf : Measurable (fun p : γ × Ω ↦ (f p.1).sample p.2)) : + Valid (fun c ↦ .sample (f c)) + /-- Bind certificates compose, carrying the previous input along with the returned value. -/ + | bind {γ α β : Type u} [MeasurableSpace γ] [MeasurableSpace α] [MeasurableSpace β] + (x : γ → Program Ω P α) (f : γ → α → Program Ω P β) + (hx : Valid x) (hf : Valid (Function.uncurry f)) : + Valid (fun c ↦ .bind (x c) (f c)) + /-- Reparameterize a certified family by a measurable function. -/ + | comp {γ δ α : Type u} [MeasurableSpace γ] [MeasurableSpace δ] [MeasurableSpace α] + (p : γ → Program Ω P α) (g : δ → γ) (hp : Valid p) (hg : Measurable g) : + Valid (fun c ↦ p (g c)) + /-- Branch on a measurable predicate. -/ + | ite {γ α : Type u} [MeasurableSpace γ] [MeasurableSpace α] + (p : γ → Prop) [DecidablePred p] (x y : γ → Program Ω P α) + (hp : Measurable p) (hx : Valid x) (hy : Valid y) : + Valid (fun c ↦ if p c then x c else y c) + +variable {γ α : Type u} [MeasurableSpace γ] [MeasurableSpace α] + +/-- Soundness of the certificate: sampling is jointly measurable, and its pushforward agrees +with the measure interpretation. No global measurable-evaluation assumption is needed. -/ +theorem Valid.sound {p : γ → Program Ω P α} (hp : Valid p) : + Measurable (fun z : γ × Ω ↦ (sampleSemantics (p z.1)).sample z.2) ∧ + ∀ c, (sampleSemantics (p c)).law = measureSemantics (p c) := by + classical + induction hp with + | pure f hf => + exact ⟨hf.prodMap measurable_id, fun _ ↦ RandomM.law_mPure _⟩ + | sample f hf => exact ⟨hf, fun _ ↦ rfl⟩ + | @bind γ α β _ _ _ x f hx hf ihx ihf => + have hfc (c : γ) : Measurable (fun z : α × Ω ↦ (sampleSemantics (f c z.1)).sample z.2) := + ihf.1.comp ((measurable_const.prodMk measurable_fst).prodMk measurable_snd) + constructor + · have h := ihf.1.comp ((measurable_fst.prodMk ihx.1.fst).prodMk ihx.1.snd) + convert h using 1 + funext z + exact RandomM.sample_bind_of_measurable _ _ (hfc z.1) z.2 + · intro c + change (RandomM.bind _ _).law = (measureSemantics (x c)).bind _ + rw [RandomM.law_bind_of_measurable _ _ (hfc c), ihx.2 c] + exact congrArg (Measure.bind _) (funext fun a ↦ ihf.2 (c, a)) + | comp p g hp hg ih => + exact ⟨ih.1.comp (hg.prodMap measurable_id), fun c ↦ ih.2 (g c)⟩ + | ite p x y hp hx hy ihx ihy => + constructor + · have h : Measurable (fun z : _ × Ω ↦ + if p z.1 then (sampleSemantics (x z.1)).sample z.2 + else (sampleSemantics (y z.1)).sample z.2) := + ihx.1.ite ((hp.comp measurable_fst).setOf) ihy.1 + convert h using 1 + funext z + by_cases h : p z.1 <;> simp only [h, ite_true, ite_false] + · intro c + by_cases h : p c + · simpa only [h, ite_true] using ihx.2 c + · simpa only [h, ite_false] using ihy.2 c + +omit [IsProbabilityMeasure P] in +/-- A closed certificate can be used in a larger input context. -/ +theorem Valid.const {p : Program Ω P α} (hp : Valid (fun _ : PUnit ↦ p)) : + Valid (fun _ : γ ↦ p) := + .comp (fun _ : PUnit ↦ p) (fun _ ↦ PUnit.unit) hp measurable_const + +/-- A program together with its compositional certificate. -/ +structure Certified (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) + [IsProbabilityMeasure P] (α : Type u) [MeasurableSpace α] where + /-- The recorded program. -/ + program : Program Ω P α + /-- Joint measurability is certified at every bind. -/ + valid : Valid (fun _ : PUnit ↦ program) + +/-- The universal sampler-to-measure theorem for certified program representations. -/ +theorem Certified.law (p : Certified Ω P α) : + (sampleSemantics p.program).law = measureSemantics p.program := + p.valid.sound.2 PUnit.unit + +/-- Transport the universal theorem through bridges to existing program definitions. -/ +theorem Certified.law_of_bridge (p : Certified Ω P α) (x : RandomM Ω P α) (μ : Measure α) + (hs : sampleSemantics p.program = x) (hm : measureSemantics p.program = μ) : x.law = μ := by + rw [← hs, ← hm] + exact p.law + +end Program + +end RDo diff --git a/RandomDo/Monad/Sample.lean b/RandomDo/Monad/Sample.lean new file mode 100644 index 0000000..9e0d6f3 --- /dev/null +++ b/RandomDo/Monad/Sample.lean @@ -0,0 +1,160 @@ +/- +Copyright (c) 2026 David Ledvinka. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: David Ledvinka +-/ +module + +public import RandomDo.Monad.Instances +public import Mathlib.Probability.Independence.InfinitePi +public import Mathlib.Probability.Kernel.Representation + +/-! +# Sampling measures and kernels + +`SampleM.draw P` consumes one entry of an IID stream with marginal law `P`. +`SampleM.ofMeasure` and `SampleM.ofKernel` use a common stream of uniform values in `[0,1]` +to sample probability measures and Markov kernels with standard Borel outputs. These constructors +are noncomputable mathematical samplers. Their distribution and measurability theorems do not +require a global `RandomM.MeasurableEvalDomain` instance. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory +open MeasurableSpaceBind MeasurableSpaceFunctor + +universe u v w + +namespace RandomM + +variable {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + +/-- The value returned by a sampler is independent of the remaining source state. -/ +theorem indepFun_fst_snd (x : RandomM Ω P α) : + IndepFun (fun ω ↦ (x.sample ω).1) (fun ω ↦ (x.sample ω).2) P := by + apply (indepFun_iff_map_prod_eq_prod_map_map x.measurable_fst_sample.aemeasurable + x.measurable_snd_sample.aemeasurable).2 + change P.map x.sample = x.law.prod (P.map (Prod.snd ∘ x.sample)) + rw [x.measurePreserving_snd.map_eq, x.map_sample] + +/-- Sequencing two fixed samplers gives independent results, regardless of earlier draws. -/ +theorem law_mBind_mMap_pair (x : RandomM Ω P α) (y : RandomM Ω P β) : + (x >>=ₘ fun a ↦ Prod.mk a <$>ₘ y).law = x.law.prod y.law := by + have hf : Measurable (fun p : α × Ω ↦ (Prod.mk p.1 <$>ₘ y).sample p.2) := by + simp only [sample_mMap, measurable_prodMk_left] + have hy : Measurable (fun p : α × Ω ↦ y.sample p.2) := + y.measurable_sample.comp measurable_snd + exact (measurable_fst.prodMk hy.fst).prodMk hy.snd + change (bind x _).law = _ + rw [law_bind_of_measurable _ _ hf] + simp only [law_mMap, measurable_prodMk_left] + ext s hs + rw [Measure.bind_apply hs Measurable.map_prodMk_left.aemeasurable, Measure.prod_apply hs] + congr 1 + funext a + rw [Measure.map_apply measurable_prodMk_left hs] + +end RandomM + +namespace SampleM + +variable {Ω : Type w} [MeasurableSpace Ω] + +/-- Consume the next entry of an IID stream with common law `P`. +The returned value is independent of the remaining stream, which has its original law. -/ +noncomputable def draw (P : Measure Ω) [IsProbabilityMeasure P] : SampleM Ω P Ω where + sample ω := (ω 0, fun n ↦ ω (n + 1)) + measurePreserving := by + have hi := iIndepFun_infinitePi (P := fun _ : ℕ ↦ P) + (X := fun _ ω ↦ ω) (fun _ ↦ measurable_id) + have h := indep_iSup_of_disjoint (fun i ↦ (measurable_pi_apply i).comap_le) hi.iIndep + (S := {0}) (T := Set.range Nat.succ) (by simp) + have hind : IndepFun (fun ω : ℕ → Ω ↦ ω 0) (fun ω : ℕ → Ω ↦ fun n ↦ ω (n + 1)) + (Measure.infinitePi fun _ : ℕ ↦ P) := by + rw [IndepFun_iff_Indep, MeasurableSpace.comap_process_pi] + simpa only [Set.mem_singleton_iff, iSup_iSup_eq_left, iSup_range] using h + refine ⟨by fun_prop, ?_⟩ + change (Measure.infinitePi fun _ : ℕ ↦ P).map (fun ω ↦ (ω 0, fun n ↦ ω (n + 1))) = + ((Measure.infinitePi fun _ : ℕ ↦ P).map (fun ω ↦ ω 0)).prod _ + rw [hind.map_prod_eq_prod_map_map (Measurable.aemeasurable (by fun_prop)) + (Measurable.aemeasurable (by fun_prop)), + Measure.map_infinitePi_infinitePi_of_inj Nat.succ_injective] + +@[simp] +theorem sample_draw (P : Measure Ω) [IsProbabilityMeasure P] (ω : ℕ → Ω) : + (draw P).sample ω = (ω 0, fun n ↦ ω (n + 1)) := rfl + +@[simp] +theorem law_draw (P : Measure Ω) [IsProbabilityMeasure P] : (draw P).law = P := + Measure.infinitePi_map_eval (fun _ : ℕ ↦ P) 0 + +variable {α : Type u} [MeasurableSpace α] [StandardBorelSpace α] + +/-- Sample a probability measure on a standard Borel space using one fresh uniform value. +All such samplers use the same source, so draws from different measures can be sequenced. -/ +noncomputable def ofMeasure (μ : Measure α) [IsProbabilityMeasure μ] : + SampleM unitInterval volume α := by + letI := nonempty_of_isProbabilityMeasure μ + exact RandomM.map μ.exists_measurable_map_eq.choose μ.exists_measurable_map_eq.choose_spec.1 + (draw volume) + +@[simp] +theorem snd_sample_ofMeasure (μ : Measure α) [IsProbabilityMeasure μ] + (ω : ℕ → unitInterval) : ((ofMeasure μ).sample ω).2 = fun n ↦ ω (n + 1) := rfl + +@[simp] +theorem law_ofMeasure (μ : Measure α) [IsProbabilityMeasure μ] : (ofMeasure μ).law = μ := by + let := nonempty_of_isProbabilityMeasure μ + simpa only [ofMeasure, RandomM.law_map, law_draw] using μ.exists_measurable_map_eq.choose_spec.2 + +variable {δ : Type v} [MeasurableSpace δ] [Nonempty α] + +/-- Sample a Markov kernel using one fresh uniform value. The chosen realization is jointly +measurable in the kernel's input and the source value, so it supports dependent sequencing. -/ +noncomputable def ofKernel (κ : Kernel δ α) [IsMarkovKernel κ] (a : δ) : + SampleM unitInterval volume α := + RandomM.map (κ.exists_measurable_map_eq_unitInterval.choose a) + κ.exists_measurable_map_eq_unitInterval.choose_spec.1.of_uncurry_left (draw volume) + +@[simp] +theorem snd_sample_ofKernel (κ : Kernel δ α) [IsMarkovKernel κ] (a : δ) + (ω : ℕ → unitInterval) : ((ofKernel κ a).sample ω).2 = fun n ↦ ω (n + 1) := rfl + +@[simp] +theorem law_ofKernel (κ : Kernel δ α) [IsMarkovKernel κ] (a : δ) : + (ofKernel κ a).law = κ a := by + simpa only [ofKernel, RandomM.law_map, law_draw] using + κ.exists_measurable_map_eq_unitInterval.choose_spec.2 a + +@[fun_prop] +theorem measurable_sample_ofKernel (κ : Kernel δ α) [IsMarkovKernel κ] : + Measurable (fun p : δ × (ℕ → unitInterval) ↦ (ofKernel κ p.1).sample p.2) := by + change Measurable (fun p : δ × (ℕ → unitInterval) ↦ + (κ.exists_measurable_map_eq_unitInterval.choose p.1 (p.2 0), fun n ↦ p.2 (n + 1))) + exact (κ.exists_measurable_map_eq_unitInterval.choose_spec.1.comp + (measurable_fst.prodMk ((measurable_pi_apply 0).comp measurable_snd))).prodMk (by fun_prop) + +@[fun_prop] +theorem measurable_ofKernel (κ : Kernel δ α) [IsMarkovKernel κ] : Measurable (ofKernel κ) := + RandomM.measurable_iff.2 fun _ ↦ (measurable_sample_ofKernel κ).comp measurable_prodMk_right + +variable {β : Type u} [MeasurableSpace β] + +@[simp] +theorem sample_mBind_ofKernel (x : SampleM unitInterval volume β) + (κ : Kernel β α) [IsMarkovKernel κ] (ω : ℕ → unitInterval) : + (x >>=ₘ ofKernel κ).sample ω = (ofKernel κ (x.sample ω).1).sample (x.sample ω).2 := + RandomM.sample_bind_of_measurable x (ofKernel κ) (measurable_sample_ofKernel κ) ω + +/-- A dependent kernel draw has the usual distribution semantics of kernel composition. -/ +@[simp] +theorem law_mBind_ofKernel (x : SampleM unitInterval volume β) + (κ : Kernel β α) [IsMarkovKernel κ] : + (x >>=ₘ ofKernel κ).law = x.law.bind κ := by + change (RandomM.bind x _).law = _ + rw [RandomM.law_bind_of_measurable _ _ (measurable_sample_ofKernel κ)] + simp only [law_ofKernel] + +end SampleM diff --git a/RandomDo/Tactic/Program.lean b/RandomDo/Tactic/Program.lean new file mode 100644 index 0000000..dfa2dfc --- /dev/null +++ b/RandomDo/Tactic/Program.lean @@ -0,0 +1,269 @@ +/- +Copyright (c) 2026 David Ledvinka. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: David Ledvinka +-/ +module + +public import RandomDo.Monad.Program +public meta import Lean.Elab.Tactic.Basic + +/-! +# Certificates for existing polymorphic `rdo` definitions + +`attribute [rdo_program] foo` keeps `foo` unchanged and generates: + +* `foo.program`: the definition interpreted in `RDo.Program`, with sampler primitives recorded; +* `foo.sample_bridge` and `foo.measure_bridge`: equalities with the original interpretations; +* `foo.valid` and `foo.certified`: a compositional measurability certificate; +* `foo.law`: the original sampler's pushforward equals the original measure interpretation. + +For example, if `sumDraws coin n` is polymorphic in its measurable-space monad, the attribute +generates `sumDraws.law coin n : (sumDraws coin n).law = sumDraws coin.law n`. +This follows from `RDo.Program.Certified.law`, the universal theorem about certified programs. + +`rdo_valid` constructs compositional validity certificates for recorded program families. +It recognizes returns, primitive samplers, binds, measurable branches, and reparameterization. +Definitions are unfolded as needed, and existing certificates can be supplied as hypotheses. + +The attribute currently expects leading `{m} [MeasurableSpaceMonad m]` parameters and an +independent universe parameter for the monad's output. Primitive arguments can be samplers +or one-argument sampler families with a fixed result type. For each family, the generated +certificate and law theorem take an additional joint-measurability hypothesis. Automatic +induction uses the last explicit `Nat` argument. General recursion and loops are not yet +handled automatically; unsupported constructs or unresolved measurability obligations fail. +-/ + +public meta section + +open Lean Meta Elab Tactic + +namespace RDo.Tactic + +/-- Inspect a program family, unfolding its head while preserving the program constructors. -/ +def programBody {α : Type} (p : Expr) (k : Expr → Expr → MetaM α) : MetaM α := do + let p ← etaExpand p + lambdaBoundedTelescope p 1 fun cs body ↦ do + let body ← withTransparency .default <| whnfHeadPred body fun e ↦ + return !e.isAppOf ``Program.pure && !e.isAppOf ``Program.sample && + !e.isAppOf ``Program.bind && !e.isAppOf ``ite + k cs[0]! body + +/-- Construct the structural part of a program certificate, leaving analytic side conditions. -/ +partial def programValidCore (g : MVarId) : MetaM (List MVarId) := g.withContext do + if ← g.assumptionCore then return [] + let target ← instantiateMVars (← g.getType) + unless target.isAppOf ``Program.Valid do return [g] + let rule ← programBody target.appArg! fun c body ↦ do + if body.isAppOf ``Program.pure then return some ``Program.Valid.pure + if body.isAppOf ``Program.sample then return some ``Program.Valid.sample + if body.isAppOf ``Program.bind then return some ``Program.Valid.bind + if body.isAppOf ``ite then return some ``Program.Valid.ite + if !body.containsFVar c.fvarId! && !(← inferType c).isConstOf ``PUnit then + return some ``Program.Valid.const + if body.isApp && !body.appFn!.containsFVar c.fvarId! + && body.appArg!.containsFVar c.fvarId! && body.appArg! != c then + return some ``Program.Valid.comp + return none + let some rule := rule | return [g] + let gs ← g.applyConst rule + return (← gs.mapM programValidCore).flatten + +/-- Reuse a joint-measurability hypothesis after changing the family parameter and source state. +Keeping these two arguments paired avoids asking `fun_prop` to prove separate measurability. -/ +def programSampleFamilyMeasurable (g : MVarId) : MetaM (List MVarId) := g.withContext do + let target ← instantiateMVars (← g.getType) + unless target.isAppOf ``Measurable do return [g] + lambdaBoundedTelescope (← etaExpand target.appArg!) 1 fun xs body ↦ do + let body ← whnfHeadPred body fun e ↦ return !e.isAppOf ``RandomM.sample + unless body.isAppOfArity ``RandomM.sample 7 do return [g] + let sampler := body.appFn!.appArg! + unless sampler.isApp && !sampler.appFn!.containsFVar xs[0]!.fvarId! do return [g] + let family ← withLocalDeclD `a (← inferType sampler.appArg!) fun a ↦ do + mkLambdaFVars #[a] (← mkAppM ``RandomM.sample #[mkApp sampler.appFn! a]) + let condition ← mkAppM ``Measurable #[← mkAppM ``Function.uncurry #[family]] + for decl in ← getLCtx do + if ← isDefEq decl.type condition then + let change ← mkLambdaFVars xs (← mkAppM ``Prod.mk #[sampler.appArg!, body.appArg!]) + let hg ← mkFreshExprMVar (← mkAppM ``Measurable #[change]) + g.assign (← mkAppM ``Measurable.comp #[decl.toExpr, hg]) + return [hg.mvarId!] + return [g] + +/-- Prove the compositional validity of a recorded program family. Unproved measurability +conditions and unsupported constructs are left as goals. -/ +syntax (name := rdoValidTac) "rdo_valid" : tactic + +/-- Elaborate a validity proof by composing program certificates and proving measurability. -/ +@[tactic rdoValidTac] +def elabRdoValid : Tactic := fun _ ↦ do + liftMetaTactic fun g ↦ do + return (← (← programValidCore g).mapM programSampleFamilyMeasurable).flatten + evalTactic (← `(tactic| all_goals try first | assumption | fun_prop)) + +/-- Add an exposed mathematical companion definition, without requiring executable code. -/ +def addProgramDef (name : Name) (levels : List Name) (args : Array Expr) (body : Expr) : + TermElabM Expr := do + let value ← instantiateMVars (← mkLambdaFVars args body) + addDecl (.defnDecl { + name, levelParams := levels, type := ← inferType value, value, + hints := .abbrev, safety := .safe + }) (forceExpose := true) + modifyEnv (Lean.addNoncomputable · name) + enableRealizationsForConst name + return mkAppN (mkConst name (levels.map Level.param)) args + +/-- Add a kernel-checked companion theorem with the same explicit arguments. -/ +def addProgramTheorem (name : Name) (levels : List Name) (args : Array Expr) (proof : Expr) : + TermElabM Expr := do + Term.synthesizeSyntheticMVarsNoPostponing + let value ← instantiateMVars (← mkLambdaFVars args proof) + if value.hasSorry || value.hasMVar then + throwError "rdo_program: the certificate for {name} has unresolved proof obligations" + addDecl (.thmDecl { name, levelParams := levels, type := ← inferType value, value }) + return mkAppN (mkConst name (levels.map Level.param)) args + +/-- Prove a companion obligation, using induction on the last explicit natural-number argument. +This handles structural counting programs without unrolling a fixed number of draws. -/ +def proveProgramObligation (args : Array Expr) (target : Expr) (tac : TSyntax `tactic) : + TermElabM Expr := Term.withoutErrToSorry do + let n? ← args.findSomeRevM? fun a ↦ do + let decl ← a.fvarId!.getDecl + return if decl.binderInfo.isExplicit && decl.type.isConstOf ``Nat then + some (mkIdent decl.userName) else none + let tacticCode ← match n? with + | none => `(by $tac:tactic) + | some n => + let n ← `(Lean.Parser.Tactic.elimTarget| $n:term) + `(by induction $n <;> $tac:tactic) + let proof ← mkFreshExprSyntheticOpaqueMVar target + -- Fail before Lean abstracts a public proof into an auxiliary theorem: checking `hasSorry` + -- on the resulting theorem reference would miss an unresolved goal in its body. + Term.runTactic proof.mvarId! tacticCode .term (report := false) + Term.synthesizeSyntheticMVarsNoPostponing + let proof ← instantiateMVars proof + if proof.hasSorry || proof.hasMVar then + throwError "rdo_program: could not certify this program; an unsupported construct or \ + a measurability obligation remains" + return proof + +/-- Build a recording interpretation, its certificate, bridges to both original interpretations, +and the resulting law theorem for a polymorphic definition. -/ +def addProgramCompanions (declName : Name) : TermElabM Unit := do + let info ← getConstInfo declName + unless info.hasValue do + throwError "rdo_program: expected a definition" + let .forallE _ _ (.forallE _ instType _ _) _ := info.type | + throwError "rdo_program: expected a leading monad parameter \ + and its MeasurableSpaceMonad instance" + unless instType.isAppOfArity ``MeasurableSpaceMonad 1 && instType.appArg! == .bvar 0 do + throwError "rdo_program: expected a leading monad parameter \ + and its MeasurableSpaceMonad instance" + let [u, .param v] := instType.getAppFn.constLevels! | + throwError "rdo_program: the monad's output universe must be a universe parameter" + if (.param v : Level).occurs u then + throwError "rdo_program: the monad's output universe must be independent of its input universe" + let mut wName := `rdo_w + while info.levelParams.contains wName do + wName := wName.appendAfter "_" + let w := Level.param wName + let levels := info.levelParams.filter (· != v) ++ [wName] + let root (m : Expr) (out : Level) : MetaM Expr := do + let inst ← synthInstance (mkApp (mkConst ``MeasurableSpaceMonad [u, out]) m) + return mkApp2 (mkConst declName (info.levelParams.map fun n ↦ + if n == v then out else .param n)) m inst + withLocalDecl `Ω .implicit (mkSort (.succ w)) fun Ω ↦ do + withLocalDecl `instΩ .instImplicit (mkApp (mkConst ``MeasurableSpace [w]) Ω) fun mΩ ↦ do + withLocalDecl `P .implicit (mkApp2 (mkConst ``MeasureTheory.Measure [w]) Ω mΩ) fun P ↦ do + withLocalDecl `prob .instImplicit + (mkApp3 (mkConst ``MeasureTheory.IsProbabilityMeasure [w]) Ω mΩ P) fun prob ↦ do + let samplerM := mkApp3 (mkConst ``RandomM [u, w]) Ω mΩ P + let programM := mkApp3 (mkConst ``Program [u, w]) Ω mΩ P + let samplerRoot ← root samplerM (Level.max u w).normalize + let programRoot ← root programM (Level.max (.succ u) w).normalize + let measureRoot ← root (mkConst ``MeasureTheory.Measure [u]) u + forallTelescope (← inferType samplerRoot) fun args resultType ↦ do + unless resultType.isAppOfArity ``RandomM 5 do + throwError "rdo_program: expected the definition to return m α" + let mut programArgs := #[] + let mut measureArgs := #[] + let mut familyConditions : Array (Name × Expr) := #[] + for a in args do + let ty ← inferType a + if ty.isAppOfArity ``RandomM 5 then + programArgs := programArgs.push (← mkAppM ``Program.sample #[a]) + measureArgs := measureArgs.push (← mkAppM ``RandomM.law #[a]) + else if ty.isForall then + let lifted? ← forallBoundedTelescope ty (some 1) fun xs result ↦ do + unless result.isAppOfArity ``RandomM 5 do return none + if result.containsFVar xs[0]!.fvarId! then + throwError "rdo_program: dependent result types in primitive families are unsupported" + let value := mkApp a xs[0]! + let program ← mkLambdaFVars xs (← mkAppM ``Program.sample #[value]) + let measure ← mkLambdaFVars xs (← mkAppM ``RandomM.law #[value]) + let samples ← mkLambdaFVars xs (← mkAppM ``RandomM.sample #[value]) + let condition ← mkAppM ``Measurable #[← mkAppM ``Function.uncurry #[samples]] + return some (program, measure, condition) + if let some (program, measure, condition) := lifted? then + programArgs := programArgs.push program + measureArgs := measureArgs.push measure + let name := (← a.fvarId!.getDecl).userName.appendAfter "_measurable" + familyConditions := familyConditions.push (name, condition) + else + programArgs := programArgs.push a + measureArgs := measureArgs.push a + else + programArgs := programArgs.push a + measureArgs := measureArgs.push a + let recorded := mkAppN programRoot programArgs + let sampled := mkAppN samplerRoot args + let measured := mkAppN measureRoot measureArgs + check recorded + check measured + let allArgs := #[Ω, mΩ, P, prob] ++ args + let program ← addProgramDef (declName ++ `program) levels allArgs recorded + let decl := mkIdent declName + let programDecl := mkIdent (declName ++ `program) + let bridgeTac ← `(tactic| simp_all only [$programDecl:ident, $decl:ident, + MeasurableSpacePure.mPure, MeasurableSpaceBind.mBind, + Program.sampleSemantics, Program.measureSemantics, apply_ite]) + let sampleEq ← mkEq (← mkAppM ``Program.sampleSemantics #[program]) sampled + let sampleBridge ← addProgramTheorem (declName ++ `sample_bridge) levels allArgs + (← proveProgramObligation args sampleEq bridgeTac) + let measureEq ← mkEq (← mkAppM ``Program.measureSemantics #[program]) measured + let measureBridge ← addProgramTheorem (declName ++ `measure_bridge) levels allArgs + (← proveProgramObligation args measureEq bridgeTac) + withLocalDeclsD (familyConditions.map fun (name, ty) ↦ (name, fun _ ↦ pure ty)) fun hs ↦ do + let certArgs := allArgs ++ hs + let family ← withLocalDeclD `unit (mkConst ``PUnit [.succ u]) fun unit ↦ + mkLambdaFVars #[unit] recorded + let validType ← mkAppM ``Program.Valid #[family] + let validProof ← proveProgramObligation args validType (← `(tactic| rdo_valid)) + let valid ← addProgramTheorem (declName ++ `valid) levels certArgs validProof + let certified ← addProgramDef (declName ++ `certified) levels certArgs + (← mkAppM ``Program.Certified.mk #[program, valid]) + let law ← mkAppM ``Program.Certified.law_of_bridge + #[certified, sampled, measured, sampleBridge, measureBridge] + discard <| addProgramTheorem (declName ++ `law) levels certArgs law + +/-- Generate a certified representation and sampler–measure bridges for a polymorphic program. +The original definition is unchanged. Primitive sampler arguments are interpreted by their laws. +Primitive families add joint-measurability hypotheses to the generated certificate and law. +The first version supports leading `{m} [MeasurableSpaceMonad m]` parameters, an independent +output-universe parameter, and structural induction on the last explicit `Nat` argument. -/ +initialize registerBuiltinAttribute { + name := `rdo_program + descr := "generate a certified program representation, interpretation bridges, and a law theorem" + applicationTime := .afterCompilation + add := fun declName _stx kind ↦ do + unless kind == AttributeKind.global do + throwError "rdo_program must be a global attribute" + let env ← getEnv + try + (addProgramCompanions declName).run'.run' + catch e => + setEnv env + throw e +} + +end RDo.Tactic diff --git a/Test.lean b/Test.lean index 7862fd2..8a8022b 100644 --- a/Test.lean +++ b/Test.lean @@ -8,3 +8,6 @@ public import Test.Instances public import Test.IsMarkov public import Test.Loops public import Test.MonadLaws +public import Test.Program +public import Test.RandomM +public import Test.SampleM diff --git a/Test/MonadLaws.lean b/Test/MonadLaws.lean index 7787395..0f84603 100644 --- a/Test/MonadLaws.lean +++ b/Test/MonadLaws.lean @@ -7,9 +7,9 @@ set_option linter.style.header false /-! # The `MeasurableSpaceMonad` laws at `Measure` -`Measure` is the one `LawfulMeasurableSpaceMonad` instance the library provides. Each law is -guarded by measurability hypotheses, which is what makes the Giry monad fit the class at all, so -these tests also record the exact shape each law is stated in. +Each law is guarded by measurability hypotheses, which is what makes the Giry monad fit the class +at all, so these tests also record the exact shape each law is stated in. The sampler instance is +covered in `Test.RandomM`. -/ open MeasureTheory ProbabilityTheory MeasurableSpacePure MeasurableSpaceBind diff --git a/Test/Program.lean b/Test/Program.lean new file mode 100644 index 0000000..3979dd5 --- /dev/null +++ b/Test/Program.lean @@ -0,0 +1,120 @@ +module + +public import Test.Common + +set_option linter.style.header false + +/-! +# Certified interpretations of unchanged polymorphic programs + +The generated bridges relate the original definitions, interpreted in `RandomM` and `Measure`. +These tests include continuous outputs and dependent kernel draws, without assuming joint +evaluation on the entire sampler space. The recursive Bernoulli example is in `Test.SampleM`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory + +namespace Test.Program + +universe u v w + +/-- A program with two continuous draws. -/ +@[rdo_program] +def addDraws {m : (α : Type) → [MeasurableSpace α] → Type v} + [MeasurableSpaceMonad m] (x y : m ℝ) : m ℝ := rdo + let a ← x + let b ← y + return a + b + +example {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (x y : RandomM Ω P ℝ) : (addDraws x y).law = addDraws x.law y.law := + addDraws.law x y + +example {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (x y : RandomM Ω P ℝ) : (addDraws.program x y).sampleSemantics = addDraws x y := + addDraws.sample_bridge x y + +example {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (x y : RandomM Ω P ℝ) : (addDraws.program x y).measureSemantics = addDraws x.law y.law := + addDraws.measure_bridge x y + +example {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (x y : RandomM Ω P ℝ) : (addDraws x y).law = addDraws x.law y.law := by + have h := (addDraws.certified x y).law + change (addDraws.program x y).sampleSemantics.law = + (addDraws.program x y).measureSemantics at h + simpa only [addDraws.sample_bridge, addDraws.measure_bridge] using h + +/-- A measurable branch depending on a sampled real number. -/ +@[rdo_program] +noncomputable def positivePart {m : (α : Type) → [MeasurableSpace α] → Type v} + [MeasurableSpaceMonad m] (x : m ℝ) : m ℝ := rdo + let a ← x + if a ≤ 0 then return 0 else return a + +example {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + (x : RandomM Ω P ℝ) : (positivePart x).law = positivePart x.law := + positivePart.law x + +/-- A primitive draw can depend on a preceding continuous result. -/ +@[rdo_program] +def dependentDraws {m : (α : Type) → [MeasurableSpace α] → Type v} + [MeasurableSpaceMonad m] (x : m ℝ) (k : ℝ → m ℝ) : m ℝ := rdo + let a ← x + let b ← k a + return a + b + +example (x : SampleM unitInterval volume ℝ) (κ : Kernel ℝ ℝ) [IsMarkovKernel κ] : + (dependentDraws x (SampleM.ofKernel κ)).law = dependentDraws x.law κ := by + simpa only [SampleM.law_ofKernel] using + dependentDraws.law x (SampleM.ofKernel κ) (by fun_prop) + +/-- Result types may be universe polymorphic; supplied measurability proofs remain usable. -/ +@[rdo_program] +def mapDraw {m : (α : Type u) → [MeasurableSpace α] → Type v} + [MeasurableSpaceMonad m] {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + (x : m α) (f : α → β) (_hf : Measurable f) : m β := rdo + let a ← x + return f a + +example {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + (x : RandomM Ω P α) (f : α → β) (hf : Measurable f) : + (mapDraw x f hf).law = mapDraw x.law f hf := + mapDraw.law x f hf + +/-- An arbitrary real function need not be measurable, so it cannot be silently certified. -/ +def uncertifiedMap {m : (α : Type) → [MeasurableSpace α] → Type v} + [MeasurableSpaceMonad m] (x : m ℝ) (f : ℝ → ℝ) : m ℝ := rdo + let a ← x + return f a + +/-- +error: unsolved goals +case hf +Ω : Type rdo_w +instΩ : MeasurableSpace Ω +P : Measure Ω +prob : IsProbabilityMeasure P +x : RandomM Ω P ℝ +f : ℝ → ℝ +⊢ Measurable fun c ↦ f c.2 +-/ +#guard_msgs in +attribute [rdo_program] uncertifiedMap + +/-- error: Unknown constant `Test.Program.uncertifiedMap.program` -/ +#guard_msgs in +#check Test.Program.uncertifiedMap.program + +/-- The recording interpretation needs an independent output-universe parameter. -/ +def fixedUniverse {m : (α : Type) → [MeasurableSpace α] → Type} + [MeasurableSpaceMonad m] (x : m ℝ) : m ℝ := x + +/-- error: rdo_program: the monad's output universe must be a universe parameter -/ +#guard_msgs in +attribute [rdo_program] fixedUniverse + +end Test.Program diff --git a/Test/RandomM.lean b/Test/RandomM.lean new file mode 100644 index 0000000..5b4baea --- /dev/null +++ b/Test/RandomM.lean @@ -0,0 +1,143 @@ +module + +public import Test.Common + +set_option linter.style.header false + +/-! +# The random-variable monad + +Check evaluation-domain inference across universes, measurability automation, and both the +operational and distribution semantics of an `rdo` program that changes its source state. +-/ + +open MeasureTheory MeasurableSpacePure MeasurableSpaceBind MeasurableSpaceFunctor + +@[expose] public section + +namespace Test.RandomM + +universe u w + +section Abstract + +variable {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + [RandomM.MeasurableEvalDomain.{u, w} Ω P] + {α β γ : Type u} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +example : RandomM.MeasurableEval Ω P α := inferInstance + +example {f : α → RandomM Ω P β} (hf : Measurable f) : + Measurable (fun p : α × Ω ↦ (f p.1).sample p.2) := by fun_prop + +example {f : α → RandomM Ω P β} (hf : Measurable f) : + Measurable (fun x : RandomM Ω P α ↦ (x >>=ₘ f).law) := by fun_prop + +example (x : RandomM Ω P α) {f : α → RandomM Ω P β} {g : β → RandomM Ω P γ} + (hf : Measurable f) (hg : Measurable g) : + x >>=ₘ f >>=ₘ g = x >>=ₘ fun a ↦ f a >>=ₘ g := + LawfulMeasurableSpaceMonad.mBind_assoc x hf hg + +example (x : RandomM Ω P α) : x >>=ₘ mPure = x := by simp + +example (x : RandomM Ω P α) {f : α → β} (hf : Measurable f) : + (f <$>ₘ x).law = f <$>ₘ x.law := by + rw [RandomM.law_mMap f hf, ← Measure.bind_dirac_eq_map x.law hf] + rfl + +end Abstract + +/-! The countable-source instance also works for result spaces in higher universes. -/ + +example : RandomM.MeasurableEvalDomain.{u, 0} ℕ (Measure.dirac 0) := inferInstance + +/-- Return the current state and halve it. Zero is preserved, so `dirac 0` is an invariant source. +Running at a nonzero state checks the exact state-passing semantics, even away from the support. -/ +def halve : RandomM ℕ (Measure.dirac 0) ℕ where + sample n := (n, n / 2) + measurePreserving := by + refine ⟨by fun_prop, ?_⟩ + rw [Measure.map_dirac' (by fun_prop), Measure.map_dirac' (by fun_prop)] + simp [Measure.dirac_prod_dirac] + +/-- The second sampler sees the state left by the first. -/ +noncomputable def twiceHalve : RandomM ℕ (Measure.dirac 0) ℕ := rdo + let a ← halve + let b ← halve + return a + b + +example : twiceHalve.sample 8 = (12, 2) := by + simp (disch := fun_prop) [twiceHalve, halve] + +example : twiceHalve.law = Measure.dirac 0 := by + have h : halve.law = Measure.dirac 0 := by + simp [RandomM.law, halve, Function.comp_def] + simp (disch := fun_prop) [twiceHalve, h] + +/-! Explicitly jointly measurable families can still be bound on an arbitrary source. -/ + +example {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + {α : Type u} [MeasurableSpace α] (x : RandomM Ω P α) : + RandomM.bindOfMeasurable x (fun a ↦ RandomM.pure a) (by fun_prop) = x := by + apply RandomM.ext + intro ω + rfl + +section SampleM + +variable {Ω : Type w} [MeasurableSpace Ω] {P : Measure Ω} [IsProbabilityMeasure P] + {α β γ : Type u} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +/-! Constructing stream programs needs no evaluation-domain assumption. -/ + +noncomputable example : + MeasurableSpaceMonad (SampleM Ω P : (α : Type u) → [MeasurableSpace α] → Type (max u w)) := + inferInstance + +/-- Return a value without consuming any stream entries. -/ +noncomputable def streamReturn (a : α) : SampleM Ω P α := rdo + return a + +example (a : α) (ω : ℕ → Ω) : (streamReturn (P := P) a).sample ω = (a, ω) := by + simp [streamReturn] + +example (a : α) : (streamReturn (P := P) a).law = Measure.dirac a := by + simp [streamReturn] + +/-- Bind stream samplers without requiring a global measurable-evaluation instance. -/ +noncomputable def streamBind (x : SampleM Ω P α) (f : α → SampleM Ω P β) : SampleM Ω P β := rdo + let a ← x + f a + +example (x : SampleM Ω P α) : streamBind x streamReturn = x := RandomM.bind_pure x + +example (x : SampleM Ω P α) {f : α → β} (hf : Measurable f) : + (f <$>ₘ x).law = x.law.map f := by simp [hf] + +example (x : SampleM Ω P α) (f : α → SampleM Ω P β) + (hf : Measurable (fun p : α × (ℕ → Ω) ↦ (f p.1).sample p.2)) (ω : ℕ → Ω) : + (streamBind x f).sample ω = (f (x.sample ω).1).sample (x.sample ω).2 := + RandomM.sample_bind_of_measurable x f hf ω + +example (x : SampleM Ω P α) (f : α → SampleM Ω P β) + (hf : Measurable (fun p : α × (ℕ → Ω) ↦ (f p.1).sample p.2)) : + (streamBind x f).law = x.law.bind (fun a ↦ (f a).law) := + RandomM.law_bind_of_measurable x f hf + +example (x : SampleM Ω P α) (f : α → SampleM Ω P β) (g : β → SampleM Ω P γ) + (hf : Measurable (fun p : α × (ℕ → Ω) ↦ (f p.1).sample p.2)) + (hg : Measurable (fun p : β × (ℕ → Ω) ↦ (g p.1).sample p.2)) : + streamBind (streamBind x f) g = streamBind x (fun a ↦ streamBind (f a) g) := + RandomM.bind_assoc_of_measurable x hf hg + +/-! The lawful instance still records the hypothesis on the full stream space. -/ + +example [RandomM.MeasurableEvalDomain.{u, w} (ℕ → Ω) (Measure.infinitePi fun _ : ℕ ↦ P)] : + LawfulMeasurableSpaceMonad + (SampleM Ω P : (α : Type u) → [MeasurableSpace α] → Type (max u w)) := inferInstance + +end SampleM + +end Test.RandomM + +end diff --git a/Test/SampleM.lean b/Test/SampleM.lean new file mode 100644 index 0000000..32d44bf --- /dev/null +++ b/Test/SampleM.lean @@ -0,0 +1,177 @@ +module + +public import Test.Common +public import Mathlib.MeasureTheory.MeasurableSpace.NCard +public import Mathlib.Probability.Distributions.Binomial + +set_option linter.style.header false + +/-! +# Independent draws and Bernoulli sums + +These tests use uncountable stream spaces without a `MeasurableEvalDomain` assumption. +They check exact state consumption, independent draws from measures, dependent kernel draws, +and the binomial law of a sum of `n` Bernoulli samples. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory MeasurableSpacePure MeasurableSpaceBind MeasurableSpaceFunctor +open scoped BigOperators + +namespace Test.SampleM + +/-- Draw two fair coins directly from a stream of fair coins. -/ +noncomputable def twoFairCoins : SampleM Bool fairCoin (Bool × Bool) := rdo + let a ← SampleM.draw fairCoin + let b ← SampleM.draw fairCoin + return (a, b) + +example (ω : ℕ → Bool) : + twoFairCoins.sample ω = ((ω 0, ω 1), fun n ↦ ω (n + 2)) := by + simp [twoFairCoins, Nat.add_assoc] + +theorem law_twoFairCoins : twoFairCoins.law = fairCoin.prod fairCoin := by + change (SampleM.draw fairCoin >>=ₘ fun a ↦ Prod.mk a <$>ₘ SampleM.draw fairCoin).law = _ + rw [RandomM.law_mBind_mMap_pair] + simp + +example (a b : Bool) : twoFairCoins.law {(a, b)} = (1 : ENNReal) / 4 := by + have h (b : Bool) : fairCoin {b} = (1 : ENNReal) / 2 := by + cases b <;> norm_num [fairCoin, bernoulliMeasure_apply, unitInterval.toNNReal] <;> + change ((1 / 2 : NNReal) : ENNReal) = _ <;> norm_num + rw [law_twoFairCoins, ← Set.singleton_prod_singleton, Measure.prod_prod, h, h] + norm_num [← ENNReal.mul_inv] + +section Constructors + +universe u + +variable {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + +/-- Different probability measures can be sampled within the same computation. -/ +noncomputable def independentDraws [StandardBorelSpace α] [StandardBorelSpace β] + (μ : Measure α) (ν : Measure β) [IsProbabilityMeasure μ] [IsProbabilityMeasure ν] : + SampleM unitInterval volume (α × β) := rdo + let a ← SampleM.ofMeasure μ + let b ← SampleM.ofMeasure ν + return (a, b) + +example [StandardBorelSpace α] [StandardBorelSpace β] + (μ : Measure α) (ν : Measure β) [IsProbabilityMeasure μ] [IsProbabilityMeasure ν] : + (independentDraws μ ν).law = μ.prod ν := by + change (SampleM.ofMeasure μ >>=ₘ fun a ↦ Prod.mk a <$>ₘ SampleM.ofMeasure ν).law = _ + rw [RandomM.law_mBind_mMap_pair] + simp + +variable [StandardBorelSpace β] [Nonempty β] + +/-- A kernel draw may depend on the previous result, using fresh randomness. -/ +noncomputable def dependentDraws (x : SampleM unitInterval volume α) + (κ : Kernel α β) [IsMarkovKernel κ] : SampleM unitInterval volume β := rdo + let a ← x + SampleM.ofKernel κ a + +example (x : SampleM unitInterval volume α) (κ : Kernel α β) [IsMarkovKernel κ] : + (dependentDraws x κ).law = x.law.bind κ := by simp [dependentDraws] + +example (κ : Kernel α β) [IsMarkovKernel κ] : Measurable (SampleM.ofKernel κ) := by fun_prop + +example (κ : Kernel α β) [IsMarkovKernel κ] : + Measurable (fun p : α × (ℕ → unitInterval) ↦ (SampleM.ofKernel κ p.1).sample p.2) := by + fun_prop + +end Constructors + +universe v + +/-- One polymorphic program for counting successes, interpreted as either a sampler or a measure. -/ +def sumDraws {m : (α : Type) → [MeasurableSpace α] → Type v} + [MeasurableSpaceMonad m] (coin : m Bool) : ℕ → m ℕ + | 0 => rdo return 0 + | n + 1 => rdo + let b ← coin + let s ← sumDraws coin n + return b.toNat + s + +attribute [rdo_program] sumDraws + +/-- Taking the sampler's law agrees with interpreting the same program in the measure monad. -/ +theorem law_sumDraws {Ω : Type*} [MeasurableSpace Ω] {P : Measure Ω} + [IsProbabilityMeasure P] (coin : RandomM Ω P Bool) (n : ℕ) : + (sumDraws coin n).law = sumDraws (m := Measure) coin.law n := + sumDraws.law coin n + +theorem law_sumDraws_congr {Ω Ω' : Type*} [MeasurableSpace Ω] [MeasurableSpace Ω'] + {P : Measure Ω} {P' : Measure Ω'} [IsProbabilityMeasure P] [IsProbabilityMeasure P'] + (coin : RandomM Ω P Bool) (coin' : RandomM Ω' P' Bool) (h : coin.law = coin'.law) + (n : ℕ) : (sumDraws coin n).law = (sumDraws coin' n).law := by + induction n with + | zero => simp [sumDraws] + | succ n ih => simp only [sumDraws, RandomM.law_mBind_of_countable, RandomM.law_mPure, h, ih] + +theorem sample_sumDraws_draw (P : Measure Bool) [IsProbabilityMeasure P] + (n : ℕ) (ω : ℕ → Bool) : + (sumDraws (SampleM.draw P) n).sample ω = + (∑ i ∈ Finset.range n, (ω i).toNat, fun i ↦ ω (i + n)) := by + induction n generalizing ω with + | zero => simp [sumDraws] + | succ n ih => + simp [sumDraws, ih, Finset.sum_range_succ', Nat.add_comm, Nat.add_left_comm, Nat.add_assoc] + +private theorem ncard_successes (n : ℕ) (ω : ℕ → Bool) : + {i | i < n ∧ ω i = true}.ncard = ∑ i ∈ Finset.range n, (ω i).toNat := by + have hs : {i | i < n ∧ ω i = true} = ↑((Finset.range n).filter fun i ↦ ω i = true) := by + ext i + simp + rw [hs, Set.ncard_coe_finset, Finset.card_eq_sum_ones, Finset.sum_filter] + apply Finset.sum_congr rfl + intro i _ + cases ω i <;> rfl + +private theorem law_successes (n : ℕ) (p : unitInterval) : + (Measure.infinitePi fun _ : ℕ ↦ bernoulliMeasure true false p).map + (fun ω : ℕ → Bool ↦ {i | i < n ∧ ω i = true}) = setBernoulli (Set.Iio n) p := by + rw [setBernoulli_eq_map] + have h : (Measure.infinitePi fun _ : ℕ ↦ bernoulliMeasure true false p).map + (fun ω : ℕ → Bool ↦ fun i ↦ i < n ∧ ω i = true) = + Measure.infinitePi (fun i : ℕ ↦ bernoulliMeasure (i < n) False p) := by + rw [Measure.infinitePi_map_pi _ (f := fun i (b : Bool) ↦ i < n ∧ b = true) (by fun_prop)] + congr 1 + funext i + rw [map_bernoulliMeasure' _ _ (by fun_prop)] + simp + have h' := congrArg (Measure.map (fun f : ℕ → Prop ↦ {i | f i})) h + rw [Measure.map_map (by fun_prop) (by fun_prop)] at h' + simpa only [Function.comp_def, Set.mem_Iio] using h' + +theorem law_sumDraws_draw (n : ℕ) (p : unitInterval) : + (sumDraws (SampleM.draw (bernoulliMeasure true false p)) n).law = binomial n p := by + have hs : + (fun ω : ℕ → Bool ↦ + ((sumDraws (SampleM.draw (bernoulliMeasure true false p)) n).sample ω).1) = + fun ω ↦ {i | i < n ∧ ω i = true}.ncard := by + funext ω + simp [sample_sumDraws_draw, ncard_successes] + change (Measure.infinitePi fun _ : ℕ ↦ bernoulliMeasure true false p).map + (fun ω ↦ ((sumDraws (SampleM.draw (bernoulliMeasure true false p)) n).sample ω).1) = _ + rw [hs] + change (Measure.infinitePi fun _ : ℕ ↦ bernoulliMeasure true false p).map + (Set.ncard ∘ fun ω ↦ {i | i < n ∧ ω i = true}) = _ + rw [← Measure.map_map (by fun_prop) (by fun_prop), law_successes] + rfl + +/-- Sum `n` independent Bernoulli draws using the general measure-to-sampler constructor. -/ +noncomputable def sumBernoulli (p : unitInterval) (n : ℕ) : SampleM unitInterval volume ℕ := + sumDraws (SampleM.ofMeasure (bernoulliMeasure true false p)) n + +theorem law_sumBernoulli_eq_measure (p : unitInterval) (n : ℕ) : + (sumBernoulli p n).law = sumDraws (m := Measure) (bernoulliMeasure true false p) n := by + rw [sumBernoulli, law_sumDraws, SampleM.law_ofMeasure] + +theorem law_sumBernoulli (p : unitInterval) (n : ℕ) : + (sumBernoulli p n).law = binomial n p := by + rw [sumBernoulli, law_sumDraws_congr _ (SampleM.draw (bernoulliMeasure true false p)) (by simp), + law_sumDraws_draw] + +end Test.SampleM