From 71326b1d9dad999e8bac6ca10331e8521a3f94de Mon Sep 17 00:00:00 2001 From: David Ledvinka Date: Fri, 31 Jul 2026 16:12:36 -0400 Subject: [PATCH] Add RDo monad infrastructure --- LeanMachineLearning.lean | 5 + LeanMachineLearning/RDo/Examples.lean | 57 ++++ LeanMachineLearning/RDo/ForInInstances.lean | 110 +++++++ .../RDo/MeasurableSpaceMonad.lean | 259 +++++++++++++++++ LeanMachineLearning/RDo/MonadInstances.lean | 64 +++++ LeanMachineLearning/RDo/RDo.lean | 270 ++++++++++++++++++ 6 files changed, 765 insertions(+) create mode 100644 LeanMachineLearning/RDo/Examples.lean create mode 100644 LeanMachineLearning/RDo/ForInInstances.lean create mode 100644 LeanMachineLearning/RDo/MeasurableSpaceMonad.lean create mode 100644 LeanMachineLearning/RDo/MonadInstances.lean create mode 100644 LeanMachineLearning/RDo/RDo.lean diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 8bf9da4e..196035c5 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -30,6 +30,11 @@ public import LeanMachineLearning.Online.Bandit.BayesRegret public import LeanMachineLearning.Online.Bandit.Regret public import LeanMachineLearning.Online.Bandit.RewardByCountMeasure public import LeanMachineLearning.Online.Bandit.SumRewards +public import LeanMachineLearning.RDo.Examples +public import LeanMachineLearning.RDo.ForInInstances +public import LeanMachineLearning.RDo.MeasurableSpaceMonad +public import LeanMachineLearning.RDo.MonadInstances +public import LeanMachineLearning.RDo.RDo public import LeanMachineLearning.SequentialLearning.Algorithm public import LeanMachineLearning.SequentialLearning.AlgorithmDensity public import LeanMachineLearning.SequentialLearning.AlgorithmDensityBayes diff --git a/LeanMachineLearning/RDo/Examples.lean b/LeanMachineLearning/RDo/Examples.lean new file mode 100644 index 00000000..a618e9f7 --- /dev/null +++ b/LeanMachineLearning/RDo/Examples.lean @@ -0,0 +1,57 @@ +module + +public import LeanMachineLearning.RDo.MonadInstances +public import LeanMachineLearning.RDo.ForInInstances +public import Mathlib.Probability.Distributions.Bernoulli +public import Mathlib.Algebra.Ring.BooleanRing + +set_option linter.style.header false + +@[expose] public section + +open MeasureTheory ProbabilityTheory Measure + +/- # Nonpolymorphic examples -/ + +universe u v + +noncomputable def measureSample : Measure Bool := rdo + let x ← bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩ + let y ← bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩ + return x + y + +def pseudoSample : PseudoRandomM Bool := rdo + let x ← Random.randBool + let y ← Random.randBool + return x + y + +/- # Polymorphic examples -/ + +variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] + +class HasBit (m : (α : Type) → MeasurableSpace α → Type v) where + bit : m Bool (by infer_instance) + +noncomputable instance : HasBit Measure where + bit := bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩ + +instance : HasBit PseudoRandomM where + bit := Random.randBool + +def indepAnd [HasBit m] : m Bool := rdo + let x ← HasBit.bit + let y ← HasBit.bit + return x && y + +noncomputable def indepAndMeasure : Measure Bool := indepAnd (m := Measure) + +def indepAndGen : PseudoRandomM Bool := indepAnd (m := PseudoRandomM) + +variable {α : Type*} [MeasurableSpace α] + +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 diff --git a/LeanMachineLearning/RDo/ForInInstances.lean b/LeanMachineLearning/RDo/ForInInstances.lean new file mode 100644 index 00000000..01f87828 --- /dev/null +++ b/LeanMachineLearning/RDo/ForInInstances.lean @@ -0,0 +1,110 @@ +/- +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 LeanMachineLearning.RDo.RDo + +/-! +# Instances for `MeasurableSpaceForIn` + +**TODO** + +-/ + +@[expose] public section + +section MeasurableSpace + +variable {α : Type*} [MeasurableSpace α] + +instance : MeasurableSpace (List α) := + MeasurableSpace.comap List.equivSigmaTuple inferInstance + +instance : MeasurableSpace (Array α) := + MeasurableSpace.comap Array.toList inferInstance + +end MeasurableSpace + +universe u v + +open MeasurableSpacePure + +variable {α : Type u} [mα : MeasurableSpace α] {m : (α : Type u) → [MeasurableSpace α] → Type v} + +section Array + +/-- Compiler implementation for `forIn` -/ +@[inline] unsafe def Array.measurableSpaceForIn'Unsafe [MeasurableSpaceMonad m] + {β : Type u} [mβ : MeasurableSpace β] + (as : Array α) (b : β) (f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β := + let sz := as.usize + let rec @[specialize] loop (i : USize) (b : β) : m β := rdo + if i < sz then + let a := as.uget i lcProof + match (← f a lcProof b) with + | ForInStep.done b => mPure b + | ForInStep.yield b => loop (i+1) b + else + mPure b + loop 0 b + +/-- Reference implementation for `forIn'` -/ +@[implemented_by Array.measurableSpaceForIn'Unsafe] +protected def Array.measurableSpaceForIn' [MeasurableSpaceMonad m] + {β : Type u} [mβ : MeasurableSpace β] + (as : Array α) (b : β) (f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β := + let rec loop (i : Nat) (h : i ≤ as.size) (b : β) : m β := rdo + match i, h with + | 0, _ => mPure b + | i+1, h => + have h' : i < as.size := Nat.lt_of_lt_of_le (Nat.lt_succ_self i) h + have : as.size - 1 < as.size := Nat.sub_lt (Nat.zero_lt_of_lt h') (by decide) + have : as.size - 1 - i < as.size := Nat.lt_of_le_of_lt (Nat.sub_le (as.size - 1) i) this + match (← f as[as.size - 1 - i] (getElem_mem this) b) with + | ForInStep.done b => mPure b + | ForInStep.yield b => loop i (Nat.le_of_lt h') b + loop as.size (Nat.le_refl _) b + +instance [MeasurableSpaceMonad m] : MeasurableSpaceForIn' m (Array α) α inferInstance where + forIn' := Array.measurableSpaceForIn' + +instance {n : ℕ} [MeasurableSpaceMonad m] : + MeasurableSpaceForIn' m (Vector α n) α inferInstance where + forIn' xs b f := Array.measurableSpaceForIn' xs.toArray b (fun a h b => f a (by simpa using h) b) + +end Array + +section List + +variable {α β : Type*} [MeasurableSpace α] [MeasurableSpace β] [Ring α] + +/-- Implimentation for `forIn'` -/ +@[inline] +protected def List.measurableSpaceForIn' [MeasurableSpaceMonad m] + {β : Type u} [mβ : MeasurableSpace β] (as : @& List α) (init : β) + (f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β := + let rec @[specialize] + loop : (as' : @& List α) → (b : β) → Exists (fun bs => bs ++ as' = as) → m β + | [], b, _ => mPure b + | a::as', b, h => rdo + have : a ∈ as := by + clear f + have ⟨bs, h⟩ := h + subst h + exact mem_append_right _ (Mem.head ..) + match (← f a this b) with + | ForInStep.done b => mPure b + | ForInStep.yield b => + have : Exists (fun bs => bs ++ as' = as) := + have ⟨bs, h⟩ := h + ⟨bs ++ [a], by rw [← h, append_cons (bs := as')]⟩ + loop as' b this + loop as init ⟨[], rfl⟩ + +instance [MeasurableSpaceMonad m] : MeasurableSpaceForIn' m (List α) α inferInstance where + forIn' := List.measurableSpaceForIn' + +end List diff --git a/LeanMachineLearning/RDo/MeasurableSpaceMonad.lean b/LeanMachineLearning/RDo/MeasurableSpaceMonad.lean new file mode 100644 index 00000000..a2cd268a --- /dev/null +++ b/LeanMachineLearning/RDo/MeasurableSpaceMonad.lean @@ -0,0 +1,259 @@ +/- +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 Mathlib.Control.Random +public import Mathlib.MeasureTheory.Measure.GiryMonad + +/-! +# Measurable Space Monad + +**TODO** + +## Main definitions + +* `MeasurableSpaceFunctor f`: **TODO** +* `MeasurableSpaceMonad m`: **TODO** +* `LawfulMeasurableSpaceFunctor f`: **TODO** +* `LawfulMeasurableSpaceMonad m`: **TODO** +* `MeasurableSpaceForIn`: **TODO** +* `MeasurableSpaceForIn'`: **TODO** + +-/ + +@[expose] public section + +open Function + +section + +universe u v + +section MeasurableSpaceMonad + +/-- A functor on types with a `MeasurableSpace` instance. The `mMap` operator `<$>ₘ` is overloaded +via instances of this class. This class does not require proofs of the `MeasurableSpaceFunctor` +axioms. Proofs may be provided or required via the `LawfulMeasurableSpaceFunctor` class. -/ +class MeasurableSpaceFunctor (f : (α : Type u) → [MeasurableSpace α] → Type v) : + Type (max (u+1) v) where + /-- + Applies a function inside a functor on measurable spaces. This is used to overload the + `<$>ₘ` operator. + + When mapping a constant function (if one cares about executing the code), use + `Functor.mMapConst` instead, because it may be more efficient. + -/ + mMap {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] : (α → β) → f α → f β + /-- + Mapping a constant function. + + Given `a : α` and `v : f β`, `mMapConst a v` is equivalent to `(fun _ => a) <$>ₘ v`. For some + functors, this can be implemented more efficiently; for all other functors, the default + implementation may be used. + -/ + mMapConst {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] : α → f β → f α := + mMap ∘ (const _) + +@[inherit_doc] infixr:100 " <$>ₘ " => MeasurableSpaceFunctor.mMap + +/-- +The `mPure` function is overloaded via `MeasurableSpacePure` instances. + +`MeasurableSpacePure` is typically accessed via `MeasurableSpaceMonad` instances, which extend it. +-/ +class MeasurableSpacePure (f : (α : Type u) → [MeasurableSpace α] → Type v) where + mPure {α : Type u} [MeasurableSpace α] : α → f α + +/-- +The `>>=ₘ` operator is overloaded via instances of `MeasurableSpaceBind`. + +`MeasurableSpaceBind` is typically used via `MeasurableSpaceMonad`, which extends it. +-/ +class MeasurableSpaceBind (m : (α : Type u) → [MeasurableSpace α] → Type v) where + /-- + Sequences two computations, allowing the second to depend on the value computed by the first. + + If `x : m α` and `f : α → m β`, then `x >>=ₘ f : m β` represents the result of executing `x` + to get a value of type `α` and then passing it to `f`. + -/ + mBind {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] : m α → (α → m β) → m β + +@[inherit_doc] infixl:55 " >>=ₘ " => MeasurableSpaceBind.mBind + +/-- A monad on types with a `MeasurableSpace` instance. This abstraction allows the user to write +code that depends on "random" side-effects -/ +class MeasurableSpaceMonad (m : (α : Type u) → [MeasurableSpace α] → Type v) : + Type (max (u+1) v) + extends MeasurableSpaceFunctor m, MeasurableSpacePure m, MeasurableSpaceBind m where + mMap f μ := mBind μ (Function.comp mPure f) + +variable {m : (α : Type u) → [MeasurableSpace α] → Type v} {α β : Type u} + [MeasurableSpace α] [MeasurableSpace β] + +theorem MeasurableSpaceBind.bind_congr [MeasurableSpaceBind m] {x : m α} {f g : α → m β} + (h : ∀ a, f a = g a) : x >>=ₘ f = x >>=ₘ g := by + simp [funext h] + +theorem MeasurableSpaceFunctor.map_congr [MeasurableSpaceFunctor m] {x : m α} {f g : α → β} + (h : ∀ a, f a = g a) : (f <$>ₘ x : m β) = g <$>ₘ x := by + simp [funext h] + +end MeasurableSpaceMonad + +section Lawful + +open MeasurableSpaceFunctor + +/-- A `MeasurableSpaceFunctor` satisfies the measurable space functor laws. -/ +class LawfulMeasurableSpaceFunctor + (f : (α : Type u) → [MeasurableSpace α] → Type v) [MeasurableSpaceFunctor f] + [∀ α, [MeasurableSpace α] → MeasurableSpace (f α)] : Prop where + /-- `mMap` of a measurable function is a measurable function. -/ + measurable_mMap {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + {g : α → β} (hg : Measurable g) : Measurable ((g <$>ₘ ·) : f α → f β) + /-- The `mMapConst` implimentation is equivalent to the default implimentation. -/ + mMap_const {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] : + (mMapConst : α → f β → f α) = mMap ∘ const β + /-- `mMap` preserves identity. -/ + id_mMap {α : Type u} [MeasurableSpace α] (x : f α) : id <$>ₘ x = x + /-- `mMap` preserves function composition of measurable functions. -/ + comp_mMap {α β γ : Type u} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + {g₀ : α → β} {g₁ : β → γ} (x : f α) (hg₀ : Measurable g₀) (hg₁ : Measurable g₁) : + (g₁ ∘ g₀) <$>ₘ x = g₁ <$>ₘ g₀ <$>ₘ x + +open LawfulMeasurableSpaceFunctor + +attribute [fun_prop] measurable_mMap +attribute [simp] id_mMap + +variable {f : (α : Type u) → [MeasurableSpace α] → Type v} [MeasurableSpaceFunctor f] + [∀ α, [MeasurableSpace α] → MeasurableSpace (f α)] [LawfulMeasurableSpaceFunctor f] + {α β γ : Type u} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +@[simp] theorem id_mMap' (x : f α) : (fun a => a) <$>ₘ x = x := id_mMap x + +@[simp] theorem mMap_mMap {g₀ : α → β} {g₁ : β → γ} (x : f α) + (hg₀ : Measurable g₀) (hg₁ : Measurable g₁) : + g₁ <$>ₘ g₀ <$>ₘ x = (fun a => g₁ (g₀ a)) <$>ₘ x := + (comp_mMap x hg₀ hg₁).symm + +@[simp] theorem mMap_unit {a : f PUnit} : (fun _ => PUnit.unit) <$>ₘ a = a := by simp + +open MeasurableSpaceBind MeasurableSpacePure MeasurableSpaceFunctor + +/-- A `MeasurableSpaceMonad` satisfies the measurable space monad laws. -/ +class LawfulMeasurableSpaceMonad + (m : (α : Type u) → [MeasurableSpace α] → Type v) [MeasurableSpaceMonad m] + [∀ α, [MeasurableSpace α] → MeasurableSpace (m α)] : Prop + extends LawfulMeasurableSpaceFunctor m where + /-- `mPure` is a measurable function. -/ + measurable_mPure {α : Type u} [MeasurableSpace α] : Measurable (mPure : α → m α) + /-- `mBind` of a measurable function is a measurable function. -/ + measurable_mBind {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + {f : α → m β} (hf : Measurable f) : Measurable (fun x : m α => x >>=ₘ f) + /-- A `mBind` followed by `mPure` composed with a measurable function is equivalent to a + functorial map. -/ + mBind_mPure_comp {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + {f : α → β} (hf : Measurable f) (x : m α) : + x >>=ₘ (fun a => mPure (f a)) = f <$>ₘ x + /-- `mPure` followed by `mBind` of a function application is equivalent to function + application. -/ + mPure_mBind {α β : Type u} [MeasurableSpace α] [MeasurableSpace β] + (x : α) {f : α → m β} (hf : Measurable f) : + mPure x >>=ₘ f = f x + /-- `mBind` is associative on measurable functions. -/ + mBind_assoc {α β γ : Type u} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + (x : m α) {f : α → m β} {g : β → m γ} (hf : Measurable f) (hg : Measurable g) : + x >>=ₘ f >>=ₘ g = x >>=ₘ fun x => f x >>=ₘ g + measurable_mMap hg := (by + convert measurable_mBind (measurable_mPure.comp hg) + exact (mBind_mPure_comp hg _).symm) + comp_mMap x g_meas h_meas := (by + rw [← mBind_mPure_comp (by fun_prop), ← mBind_mPure_comp h_meas, + ← mBind_mPure_comp g_meas, mBind_assoc _ (by fun_prop) (by fun_prop)] + congr with _ + exact (mPure_mBind _ (measurable_mPure.comp h_meas)).symm) + +open LawfulMeasurableSpaceMonad + +attribute [fun_prop] measurable_mPure measurable_mBind +attribute [simp] pure_bind bind_assoc bind_pure_comp +attribute [grind <=] pure_bind + +variable {m : (α : Type u) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m] + [∀ α, [MeasurableSpace α] → MeasurableSpace (m α)] [LawfulMeasurableSpaceMonad m] + {α β γ : Type u} [MeasurableSpace α] [MeasurableSpace β] [MeasurableSpace γ] + +@[simp] theorem mMap_mPure {f : α → β} (hf : Measurable f) (x : α) : + f <$>ₘ (mPure x : m α) = mPure (f x) := by + rw [← mBind_mPure_comp hf, mPure_mBind _ (by fun_prop)] + +@[simp] theorem mBind_mPure (x : m α) : x >>=ₘ mPure = x := by + change x >>=ₘ (fun a => (mPure (id a))) = x + rw [mBind_mPure_comp (by fun_prop), id_mMap] + +theorem mMap_eq_mPure_mBind {f : α → β} (hf : Measurable f) (x : m α) : + f <$>ₘ x = x >>=ₘ fun a => mPure (f a) := by + rw [← mBind_mPure_comp hf x] + +theorem mBind_mPure_unit {x : m PUnit} : (x >>=ₘ fun _ => mPure ⟨⟩) = x := by rw [mBind_mPure] + +@[simp] theorem mMap_mBind {f : β → γ} (hf : Measurable f) (x : m α) + {g : α → m β} (hg : Measurable g) : + f <$>ₘ (x >>=ₘ g) = x >>=ₘ fun a => f <$>ₘ g a := by + rw [← mBind_mPure_comp hf, mBind_assoc _ hg (by fun_prop)] + simp (disch := fun_prop) [mBind_mPure_comp] + +@[simp] theorem mBind_mMap_left {f : α → β} (hf : Measurable f) (x : m α) + {g : β → m γ} (hg : Measurable g) : + ((f <$>ₘ x) >>=ₘ fun b => g b) = (x >>=ₘ fun a => g (f a)) := by + rw [← mBind_mPure_comp hf] + simp (disch := fun_prop) [mBind_assoc, mPure_mBind] + +end Lawful + +end + +section MeasurableSpaceFor + +universe uρ uα u v + +variable {α : Type u} [mα : MeasurableSpace α] (m : (α : Type u) → [MeasurableSpace α] → Type v) + {m' : (α : Type u) → Type v} + +instance instMeasurableSpace {β : Type u} [mβ : MeasurableSpace β] : + MeasurableSpace (ForInStep β) := mβ.map ForInStep.yield ⊓ mβ.map ForInStep.done + +/-- +Monadic iteration in `rdo`-blocks, using the `for x in xs` notation. +-/ +class MeasurableSpaceForIn (ρ : Type uρ) (α : outParam (Type uα)) where + /-- + Monadically iterates over the contents of a collection `xs`, with a local state `b` and the + possibility of early termination. + -/ + forIn {β : Type u} [MeasurableSpace β] (xs : ρ) (b : β) + (f : α → β → m (ForInStep β)) : m β + +/-- +Monadic iteration in `rdo`-blocks with a membership proof, using the `for h : x in xs` notation. +-/ +class MeasurableSpaceForIn' (ρ : Type uρ) (α : outParam (Type uα)) + (d : outParam (Membership α ρ)) where + /-- + Monadically iterates over the contents of a collection `xs`, with a local state `b` and the + possibility of early termination. At each iteration, the body of the loop is provided with a proof + that the current element is in the collection. + -/ + forIn' {β : Type u} [MeasurableSpace β] (xs : ρ) (b : β) + (f : (a : α) → a ∈ xs → β → m (ForInStep β)) : m β + +instance (priority := 500) instMeasurableSpaceForInOfForIn' + {ρ : Type uρ} {α : Type uα} {d : Membership α ρ} [MeasurableSpaceForIn' m ρ α d] : + MeasurableSpaceForIn m ρ α where + forIn x b f := MeasurableSpaceForIn'.forIn' x b fun a _ s => f a s + +end MeasurableSpaceFor diff --git a/LeanMachineLearning/RDo/MonadInstances.lean b/LeanMachineLearning/RDo/MonadInstances.lean new file mode 100644 index 00000000..62f89424 --- /dev/null +++ b/LeanMachineLearning/RDo/MonadInstances.lean @@ -0,0 +1,64 @@ +/- +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 LeanMachineLearning.RDo.MeasurableSpaceMonad +public import Mathlib.Probability.ProductMeasure + +/-! +# Instances for `MeasurableSpaceMonad` + +**TODO** + +-/ + +@[expose] public section + +universe u v w + +/-- A (core) monad automatically defines a (not necessarily lawful) measurable space monad by +forgetting the measurable space argument. -/ +def Monad.toMeasurableSpaceMonad (m : Type u → Type v) [Monad m] (α : Type u) [MeasurableSpace α] : + Type v := m α + +instance {m : Type u → Type v} [Monad m] : + MeasurableSpaceMonad (Monad.toMeasurableSpaceMonad m) where + mPure := pure + mBind := bind + +/-- A measurable space monad for pseudo random number generation. -/ +abbrev PseudoRandomM := Monad.toMeasurableSpaceMonad Rand + +open MeasureTheory + +open MeasurableSpacePure MeasurableSpaceBind MeasurableSpaceFunctor MeasurableSpaceMonad + +noncomputable instance : MeasurableSpaceMonad Measure where + mPure := Measure.dirac + mBind := Measure.bind + +instance : LawfulMeasurableSpaceMonad Measure where + mMap_const := by simp [mMapConst, mMap] + id_mMap μ := by simp [mMap] + measurable_mPure := by unfold mPure; fun_prop + measurable_mBind := by unfold mBind; fun_prop + mBind_mPure_comp _ _ := by rfl + mPure_mBind x _ hf := Measure.dirac_bind hf x + mBind_assoc _ _ _ hf hg := Measure.bind_bind hf.aemeasurable hg.aemeasurable + +section RandomM + +open Function + +structure RandomM (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) + (α : Type u) [MeasurableSpace α] where + sample : Ω → α × Ω + measurePreserving : MeasurePreserving sample P ((Measure.map (Prod.fst ∘ sample) P).prod P) + +abbrev SampleM (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) := + RandomM (ℕ → Ω) (Measure.infinitePi fun _ : ℕ ↦ P) + +end RandomM diff --git a/LeanMachineLearning/RDo/RDo.lean b/LeanMachineLearning/RDo/RDo.lean new file mode 100644 index 00000000..74d269c5 --- /dev/null +++ b/LeanMachineLearning/RDo/RDo.lean @@ -0,0 +1,270 @@ +/- +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 LeanMachineLearning.RDo.MeasurableSpaceMonad +public meta import Lean.Meta.ProdN + +/-! +# `rdo` Notation + +**TODO** + +-/ + +open Lean Lean.Parser Lean.Parser.Term Lean.Meta Lean.Elab Lean.Elab.Do Lean.Elab.Term Std Meta + +open MeasureTheory Measure MeasurableSpacePure MeasurableSpaceBind + +public meta section + +namespace RDo + +/-- Define `DoOp`s for `rdo` notation. -/ +def randOps : DoOps := { DoOps.default with + mkPureApp α e := do + let ⟨m, u, v, _, _⟩ := (← read).monadInfo + let α ← Term.ensureHasType (mkSort (mkLevelSucc u)) α + let e ← Term.ensureHasType α e + let instPure ← mkInstMVar (.app (mkConst ``MeasurableSpacePure [u,v]) m) + let instPure ← instantiateMVars instPure + let σ ← mkInstMVar (.app (mkConst ``MeasurableSpace [u]) α) + let σ ← instantiateMVars σ + return mkAppN (mkConst ``mPure [u, v]) #[m, instPure, α, σ, e] + mkBindApp α β e k := do + let ⟨m, u, v, _, _⟩ := (← read).monadInfo + let α ← Term.ensureHasType (mkSort (mkLevelSucc u)) α + let σα ← mkInstMVar (.app (mkConst ``MeasurableSpace [u]) α) + let σβ ← mkInstMVar (.app (mkConst ``MeasurableSpace [u]) β) + let mα := mkApp2 m α σα + let mβ := mkApp2 m β σβ + let e ← Term.ensureHasType mα e + let k ← Term.ensureHasType (← mkArrow α mβ) k + let σα ← instantiateMVars σα + let σβ ← instantiateMVars σβ + let instBind ← mkInstMVar (.app (mkConst ``MeasurableSpaceBind [u, v]) m) + let instBind ← instantiateMVars instBind + return mkAppN (mkConst ``mBind [u, v]) #[m, instBind, α, β, σα, σβ, e, k] + isPureApp? e := + if e.isAppOfArity ``mPure 5 then some (e.getArg! 4) else none + splitMonadApp? type := do + let .app m _ := type.consumeMData | return none + let .app m resultType := m.consumeMData | return none + unless ← isType resultType do return none + let u ← getDecLevel resultType + let v ← getDecLevel type + return some ({ m := m, u := u.normalize, v := v.normalize }, resultType) + mkMonadApp α := do + let ⟨m, u, _, _, _⟩ := (← read).monadInfo + let σ ← mkInstMVar (mkApp (mkConst ``MeasurableSpace [u]) α) + return mkApp2 m α σ + } + +syntax (name := randKind) "rdo" doSeq : term + +/-- Define `rdo` notation elaborator. -/ +@[term_elab randKind] def elabRand : Term.TermElab := fun stx et? => do + let `(rdo $doSeq) := stx | throwUnsupportedSyntax + elabDoWith randOps doSeq et? + +section LoopElab + +/-- parses `_ : _ in` for `rdo` for loops -/ +def rdoForDecl := leading_parser + Lean.Parser.optional (atomic (Term.ident >> " : ")) >> termParser >> " in " >> + withForbidden "rdo" termParser + +/-- parser for `rdo` for loops -/ +@[doElem_parser] def rdoFor := leading_parser + "for " >> sepBy1 rdoForDecl ", " >> "rdo " >> doSeq + +/-- Define expander for loops in `rdo` notation. Note this code mirrors core's implimentation for +`do` notation as much as possible. -/ +@[macro rdoFor] def expandRDoFor : Macro := fun stx => do + match stx with + | `(rdoFor| for $[$_ : ]? $_:ident in $_ rdo $_) => + -- This is the target form of the expander, handled by `elabRDoFor` below. + Macro.throwUnsupported + | `(rdoFor| for%$tk $decls:rdoForDecl,* rdo $body) => + let decls := decls.getElems + let `(rdoForDecl| $[$h? : ]? $pattern in $xs) := decls[0]! | Macro.throwUnsupported + let mut doElems := #[] + let mut body := body + -- Expand `pattern` into an `Ident` `x`: + let x ← + if pattern.raw.isIdent then + pure ⟨pattern⟩ + else if pattern.raw.isOfKind ``Lean.Parser.Term.hole then + Term.mkFreshIdent pattern + else + -- This case is a last resort, because it introduces a `match` and that will cause eager + -- defaulting. In practice this means that `mut` vars default to `Nat` too often. + -- Hence we try to only generate a `match` if we absolutely must. + let x ← Term.mkFreshIdent pattern + body ← `(doSeq| match $x:term with | $pattern => $body) + pure x + -- Expand the remaining `rdoForDecl`s: + for rdoForDecl in decls[1...*] do + /- + Expand + ``` + for x in xs, y in ys rdo + body + ``` + into + ``` + let mut s := Std.toStream ys + for x in xs rdo + match Std.Stream.next? s with + | none => break + | some (y, s') => + s := s' + body + ``` + -/ + let `(rdoForDecl| $[$h? : ]? $y in $ys) := rdoForDecl | Macro.throwUnsupported + if let some h := h? then + Macro.throwErrorAt h "The proof annotation here has not been implemented yet." + /- Recall that `@` (explicit) disables `coeAtOutParam`. + We used `@` at `Stream` functions to make sure `resultIsOutParamSupport` is not used. -/ + let toStreamApp ← withRef ys `(@Std.toStream _ _ _ $ys) + let s := mkIdentFrom ys (← withFreshMacroScope <| MonadQuotation.addMacroScope `__s) + doElems := doElems.push (← `(doSeqItem| let mut $s := $toStreamApp:term)) + body ← `(doSeq| + match @Std.Stream.next? _ _ _ $s with + | none => break + | some ($y, s') => + $s:ident := s' + rdo $body) + doElems := doElems.push (← `(doSeqItem| for%$tk $[$h? : ]? $x:ident in $xs rdo $body)) + `(doElem| do $doElems*) + | _ => Macro.throwUnsupported + +/-- Define loop elaborator for `rdo` notation. Note this code mirrors core's implimentation for +`do` notation as much as possible. -/ +@[doElem_elab rdoFor] def elabRDoFor : DoElab := fun stx dec => do + let `(rdoFor| for%$tk $[$h? : ]? $x:ident in $xs rdo $body) := stx | throwUnsupportedSyntax + let dec ← dec.ensureUnitAt tk + checkMutVarsForShadowing #[x] + let uα ← mkFreshLevelMVar + let uρ ← mkFreshLevelMVar + let α ← mkFreshExprMVar (mkSort (uα.succ)) (userName := `α) -- assigned by outParam + let ρ ← mkFreshExprMVar (mkSort (uρ.succ)) (userName := `ρ) -- assigned in the next line + let xs ← Term.elabTermEnsuringType xs ρ + let mi := (← read).monadInfo + let mutVars := (← read).mutVars + + let info ← inferControlInfoSeq body + let oldReturnCont ← getReturnCont + let returnVarName ← mkFreshUserName `__r + let loopMutVars := mutVars.filter fun x => info.reassigns.contains x.getId + let loopMutVarNames := + if info.returnsEarly then + returnVarName :: (loopMutVars.map (·.getId)).toList + else + (loopMutVars.map (·.getId)).toList + let useLoopMutVars (e : Option Expr) : TermElabM (Array Expr) := do + let mut defs := #[] + unless e.isNone || info.returnsEarly do + throwError "Early returning {e} but the info said there is no early return" + if info.returnsEarly then + let returnVar ← + match e with + | none => mkNone oldReturnCont.resultType + | some e => mkSome oldReturnCont.resultType e + defs := defs.push returnVar + for x in loopMutVars do + let defn ← getLocalDeclFromUserName x.getId + Term.addTermInfo' x.ident defn.toExpr + -- ForIn forces the mut tuple into the universe mi.u: that of the do block result type. + -- If we don't do this, then we are stuck on solving constraints such as + -- `max ?u.46 ?u.47 =?= max (max ?u.22 ?u.46) ?u.47` + -- It's important we do this as a separate isLevelDefEq check on the decremented level because + -- otherwise (`ensureHasType (mkSort mi.u.succ)`) we are stuck on constraints like + -- `max (?u+1) (?v+1) =?= ?u+1` + let u ← getDecLevel defn.type + discard <| isLevelDefEq u mi.u + defs := defs.push defn.toExpr + if info.returnsEarly && loopMutVars.isEmpty then + defs := defs.push (mkConst ``Unit.unit) + return defs + + let (preS, σ) ← mkProdMkN (← useLoopMutVars none) mi.u + + let mσ ← Term.mkInstMVar <| mkApp (mkConst ``MeasurableSpace [mi.u]) σ + let (app, p?) ← match h? with + | none => + let instForIn ← Term.mkInstMVar <| + mkApp3 (mkConst ``MeasurableSpaceForIn [uρ, uα, mi.u, mi.v]) mi.m ρ α + let app := Lean.mkConst ``MeasurableSpaceForIn.forIn [uρ, uα, mi.u, mi.v] + let app := mkApp8 app mi.m ρ α instForIn σ mσ xs preS + pure (app, none) + | some _ => + let d ← mkFreshExprMVar (mkApp2 (mkConst ``Membership [uα, uρ]) α ρ) (userName := `d) + let instForIn ← Term.mkInstMVar <| + mkApp4 (mkConst ``MeasurableSpaceForIn' [uρ, uα, mi.u, mi.v]) mi.m ρ α d + let app := Lean.mkConst ``MeasurableSpaceForIn'.forIn' [uρ, uα, mi.u, mi.v] + let app := mkApp9 app mi.m ρ α d instForIn σ mσ xs preS + pure (app, some d) + let s ← mkFreshUserName `__s + let xh : Array (Name × (Array Expr → DoElabM Expr)) := match h?, p? with + | some h, some d => + #[(x.getId, fun _ => pure α), + (h.getId, fun x => pure (mkApp5 (mkConst ``Membership.mem [uα, uρ]) α ρ d xs x[0]!))] + | _, _ => + #[(x.getId, fun _ => pure α)] + + let body ← + withLocalDeclsD xh fun xh => do + Term.addLocalVarInfo x xh[0]! + if let some h := h? then + Term.addLocalVarInfo h xh[1]! + withLocalDecl s .default σ (kind := .implDetail) fun loopS => do + mkLambdaFVars (xh.push loopS) <| ← do + bindMutVarsFromTuple loopMutVarNames loopS.fvarId! do + let newDoBlockResultType := mkApp (mkConst ``ForInStep [mi.u]) σ + withDoBlockResultType newDoBlockResultType do + let continueCont := do + let (tuple, _tupleTy) ← mkProdMkN (← useLoopMutVars none) mi.u + let yield := mkApp2 (mkConst ``ForInStep.yield [mi.u]) σ tuple + mkPureApp newDoBlockResultType yield + let breakCont := do + let (tuple, _tupleTy) ← mkProdMkN (← useLoopMutVars none) mi.u + let done := mkApp2 (mkConst ``ForInStep.done [mi.u]) σ tuple + mkPureApp newDoBlockResultType done + let returnCont := { oldReturnCont with k := fun e => do + let (tuple, _tupleTy) ← mkProdMkN (← useLoopMutVars (some e)) mi.u + let done := mkApp2 (mkConst ``ForInStep.done [mi.u]) σ tuple + mkPureApp newDoBlockResultType done + } + enterLoopBody breakCont continueCont returnCont do + -- Elaborate the loop body, which must have result type `PUnit`, just like the whole `for` loop. + elabDoSeq body { dec with k := continueCont, kind := .duplicable } + + let forIn := mkApp app body + + let γ := (← read).doBlockResultType + let rest ← + withLocalDeclD s σ fun postS => do mkLambdaFVars #[postS] <| ← do + bindMutVarsFromTuple loopMutVarNames postS.fvarId! do + if info.returnsEarly then + let ret ← getFVarFromUserName returnVarName + let ret ← if loopMutVars.isEmpty then mkAppM ``Prod.fst #[ret] else pure ret + let motive := mkLambda `_ .default (← inferType ret) (← mkMonadApp γ) + let app := mkApp3 (mkConst ``Break.runK.match_1 [mi.u, mi.v.succ]) + oldReturnCont.resultType motive ret + let none := mkSimpleThunk (← dec.continueWithUnit) + let some ← withLocalDeclD (← mkFreshUserName `r) oldReturnCont.resultType fun r => do + mkLambdaFVars #[r] (← oldReturnCont.k r) + return mkApp2 app some none + else + dec.continueWithUnit + + mkBindApp σ γ forIn rest + +end LoopElab + +end RDo