diff --git a/LMLTutorial.lean b/LMLTutorial.lean index 6872a6d1..605cd0d5 100644 --- a/LMLTutorial.lean +++ b/LMLTutorial.lean @@ -1,8 +1,10 @@ -import LMLTutorial.Front -import LMLTutorial.Pages.BasicProbability -import LMLTutorial.Pages.DefiningAlgorithm -import LMLTutorial.Pages.Installation -import LMLTutorial.Pages.MarkovKernels -import LMLTutorial.Pages.Martingales -import LMLTutorial.References -import LMLTutorial.Tutorial +module -- shake: keep-all --deprecated_module: ignore + +public import LMLTutorial.Front +public import LMLTutorial.Pages.BasicProbability +public import LMLTutorial.Pages.DefiningAlgorithm +public import LMLTutorial.Pages.Installation +public import LMLTutorial.Pages.MarkovKernels +public import LMLTutorial.Pages.Martingales +public import LMLTutorial.References +public import LMLTutorial.Tutorial diff --git a/LMLTutorial/Pages/DefiningAlgorithm.lean b/LMLTutorial/Pages/DefiningAlgorithm.lean index bc9387d4..bcf9b60c 100644 --- a/LMLTutorial/Pages/DefiningAlgorithm.lean +++ b/LMLTutorial/Pages/DefiningAlgorithm.lean @@ -105,7 +105,7 @@ It starts by choosing each action once and then chooses $`\arg\max_a (\hat{\mu}_ To define the algorithm, we first define the exploration bonus and the next action function, and then we use `detAlgorithm` to build the algorithm. We also need to prove that the next action function is measurable, which is done by the `measurable_nextArm` lemma. -Note that we are careful to use a measurable version of the argmax function, `measurableArgmax`. +Note that we are careful to use a measurable version of the argmax function, `argmax`. {docstring Bandits.ucbWidth'} diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 12be32f8..ce81269a 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -1,9 +1,10 @@ module -- shake: keep-all --deprecated_module: ignore -public import LeanMachineLearning.ForMathlib.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax public import LeanMachineLearning.ForMathlib.MeasureTheory.Constructions.Polish.StandardBorel public import LeanMachineLearning.ForMathlib.MeasureTheory.Measurable public import LeanMachineLearning.ForMathlib.MeasureTheory.Measure.AbsolutelyContinuous +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.Lattice +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg public import LeanMachineLearning.ForMathlib.MeasureTheory.OuterMeasure.Basic public import LeanMachineLearning.ForMathlib.Probability.HasCondDistrib public import LeanMachineLearning.ForMathlib.Probability.Independence.CondDistrib diff --git a/LeanMachineLearning/ForMathlib/MeasureTheory/Constructions/BorelSpace/MeasurableArgMax.lean b/LeanMachineLearning/ForMathlib/MeasureTheory/Constructions/BorelSpace/MeasurableArgMax.lean deleted file mode 100644 index 96399e2b..00000000 --- a/LeanMachineLearning/ForMathlib/MeasureTheory/Constructions/BorelSpace/MeasurableArgMax.lean +++ /dev/null @@ -1,88 +0,0 @@ -/- -Copyright (c) 2025 Rémy Degenne. All rights reserved. -Released under Apache 2.0 license as described in the file LICENSE. -Authors: Rémy Degenne --/ -module - -public import Mathlib.MeasureTheory.Constructions.BorelSpace.Order - -/-! # Measurable argmax function - --/ - -@[expose] public section - -open MeasureTheory Finset -open scoped ENNReal NNReal - -section MeasurableArgmax -- copied from PR #27579 (and changed from argmin to argmax) - -lemma measurable_encode {α : Type*} {_ : MeasurableSpace α} [Encodable α] - [MeasurableSingletonClass α] : - Measurable (Encodable.encode (α := α)) := by - refine measurable_to_nat fun a ↦ ?_ - have : Encodable.encode ⁻¹' {Encodable.encode a} = {a} := by ext; simp - rw [this] - exact measurableSet_singleton _ - -lemma measurableEmbedding_encode (α : Type*) {_ : MeasurableSpace α} [Encodable α] - [MeasurableSingletonClass α] : - MeasurableEmbedding (Encodable.encode (α := α)) where - injective := Encodable.encode_injective - measurable := measurable_encode - measurableSet_image' _ _ := .of_discrete - -section Finite - -variable {𝓧 𝓨 α : Type*} {m𝓧 : MeasurableSpace 𝓧} {m𝓨 : MeasurableSpace 𝓨} - {mα : MeasurableSpace α} [TopologicalSpace α] [LinearOrder α] - [OpensMeasurableSpace α] [OrderClosedTopology α] [SecondCountableTopology α] - -lemma measurableSet_isMax [Countable 𝓨] - {f : 𝓧 → 𝓨 → α} (hf : ∀ y, Measurable (fun x ↦ f x y)) (y : 𝓨) : - MeasurableSet {x | ∀ z, f x z ≤ f x y} := by - rw [show {x | ∀ y', f x y' ≤ f x y} = ⋂ y', {x | f x y' ≤ f x y} by ext; simp] - exact MeasurableSet.iInter fun z ↦ measurableSet_le (by fun_prop) (by fun_prop) - -lemma exists_isMaxOn' {α : Type*} [LinearOrder α] - [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] (f : 𝓧 → 𝓨 → α) (x : 𝓧) : - ∃ n : ℕ, ∃ y, n = Encodable.encode y ∧ ∀ z, f x z ≤ f x y := by - obtain ⟨y, h⟩ := Finite.exists_max (f x) - exact ⟨Encodable.encode y, y, rfl, h⟩ - -/-- A measurable argmax function. -/ -noncomputable -def measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] - (f : 𝓧 → 𝓨 → α) - [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y] - (x : 𝓧) : - 𝓨 := - (measurableEmbedding_encode 𝓨).invFun (Nat.find (exists_isMaxOn' f x)) - -lemma measurable_measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] - {f : 𝓧 → 𝓨 → α} - [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y] - (hf : ∀ y, Measurable (fun x ↦ f x y)) : - Measurable (measurableArgmax f) := by - refine (MeasurableEmbedding.measurable_invFun (measurableEmbedding_encode 𝓨)).comp ?_ - refine measurable_find _ fun n ↦ ?_ - have : {x | ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y} - = ⋃ y, ({x | n = Encodable.encode y} ∩ {x | ∀ z, f x z ≤ f x y}) := by ext; simp - rw [this] - refine MeasurableSet.iUnion fun y ↦ (MeasurableSet.inter (by simp) ?_) - exact measurableSet_isMax (by fun_prop) y - -lemma isMaxOn_measurableArgmax {α : Type*} [LinearOrder α] - [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨] - (f : 𝓧 → 𝓨 → α) - [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y] - (x : 𝓧) (z : 𝓨) : - f x z ≤ f x (measurableArgmax f x) := by - obtain ⟨y, h_eq, h_le⟩ := Nat.find_spec (exists_isMaxOn' f x) - refine le_trans (h_le z) (le_of_eq ?_) - rw [measurableArgmax, h_eq, - MeasurableEmbedding.leftInverse_invFun (measurableEmbedding_encode 𝓨) y] - -end Finite -end MeasurableArgmax diff --git a/LeanMachineLearning/ForMathlib/MeasureTheory/Order/Lattice.lean b/LeanMachineLearning/ForMathlib/MeasureTheory/Order/Lattice.lean new file mode 100644 index 00000000..2552702e --- /dev/null +++ b/LeanMachineLearning/ForMathlib/MeasureTheory/Order/Lattice.lean @@ -0,0 +1,27 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import Mathlib.MeasureTheory.Order.Lattice + +/-! # Measurable inf of a finite set + +-/ + +@[expose] public section + +open Finset + +variable {α δ : Type*} [MeasurableSpace δ] [SemilatticeInf α] {m : MeasurableSpace α} + [MeasurableInf₂ α] + +attribute [to_dual existing] MeasurableInf₂ + +/-- Dual version of `Finset.measurable_sup'`. -/ +@[to_dual existing (attr := fun_prop)] +theorem Finset.measurable_inf' {ι : Type*} {s : Finset ι} (hs : s.Nonempty) {f : ι → δ → α} + (hf : ∀ n ∈ s, Measurable (f n)) : Measurable (s.inf' hs f) := + Finset.inf'_induction hs _ (fun _f hf _g hg => hf.inf hg) fun n hn => hf n hn diff --git a/LeanMachineLearning/ForMathlib/MeasureTheory/Order/MeasurableArg.lean b/LeanMachineLearning/ForMathlib/MeasureTheory/Order/MeasurableArg.lean new file mode 100644 index 00000000..7c7bf9af --- /dev/null +++ b/LeanMachineLearning/ForMathlib/MeasureTheory/Order/MeasurableArg.lean @@ -0,0 +1,92 @@ +/- +Copyright (c) 2026 Gaëtan Serré. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Gaëtan Serré +-/ +module + +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.Lattice +public import Mathlib.CategoryTheory.Countable +public import Mathlib.MeasureTheory.Constructions.Polish.Basic +public import Mathlib.Order.CompletePartialOrder + +/-! # Argmax and argmin functions on finite sets + +We prove in particular that those functions are measurable. + +-/ + +@[expose] public section + +open Finset + +variable {ι α : Type*} [LinearOrder α] [Fintype ι] [Nonempty ι] (f : ι → α) + +namespace Function + +/-- The maximum value of a tuple. -/ +@[to_dual /-- The minimum value of a tuple. -/] +abbrev max : α := univ.sup' univ_nonempty f + +@[to_dual min_le] +lemma le_max (x : ι) : f x ≤ max f := le_sup' _ (by simp) + +end Function + +section Argmax + +@[to_dual exists_argmin] +lemma exists_argmax : ∃ i, f i = f.max := by + obtain ⟨i, -, hi⟩ := Finset.exists_mem_eq_sup' (by simp : Finset.univ.Nonempty) f + exact ⟨i, hi.symm⟩ + +/-- The index of the maximum value of a tuple. -/ +@[to_dual argmin /-- The index of the minimum value of a tuple. -/] +noncomputable def argmax := (exists_argmax f).choose + +@[to_dual argmin_spec] +lemma argmax_spec : f (argmax f) = f.max := (exists_argmax f).choose_spec + +@[to_dual isMinOn_argmin] +lemma isMaxOn_argmax (x : ι) : f x ≤ f (argmax f) := by + rw [argmax_spec f] + exact f.le_max x + +variable [MeasurableSpace α] + +@[to_dual (attr := fun_prop)] +lemma measurable_max [MeasurableSup₂ α] : Measurable (fun (t : ι → α) => t.max) := by + suffices (fun f : ι → α ↦ f.max) = (univ.sup' univ_nonempty fun i f => f i) by + rw [this] + exact measurable_sup' univ_nonempty (fun i _ => measurable_pi_apply i) + ext + simp [Function.max] + +@[to_dual (attr := fun_prop) measurable_argmin] +lemma measurable_argmax [MeasurableSpace ι] [MeasurableEq α] [MeasurableSup₂ α] : + Measurable fun f : ι → α ↦ argmax f := by + refine measurable_to_countable' fun i ↦ ?_ + simp only [Set.preimage, Set.mem_singleton_iff] + let Maximizers (f : ι → α) : Set ι := {i | f i = f.max} + suffices {f : ι → α | argmax f = i} = ⋃ (S) + (hS : ∀ x, Maximizers x = S → argmax x = i), {f | Maximizers f = S} by + rw [this] + refine MeasurableSet.iUnion fun S ↦ (.iUnion fun hS ↦ ?_) + exact measurableSet_eq_fun (by fun_prop) measurable_const + ext f + simp only [Set.mem_setOf_eq, Set.mem_iUnion, exists_prop, exists_eq_right'] + constructor + · intro hf x hx + rw [← hf] + exact Classical.choose.congr_simp hx (exists_argmax x) + · intro h + exact h f rfl + +end Argmax + +lemma neg_max_eq_min_neg [AddGroup α] [AddLeftMono α] [AddRightMono α] : -f.max = (-f).min := by + refine le_antisymm ?_ ?_ + · simp; grind + · simp only [inf'_le_iff, mem_univ, Pi.neg_apply, neg_le_neg_iff, sup'_le_iff, forall_const, + true_and] + exact ⟨argmax f, isMaxOn_argmax f⟩ diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean b/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean index 9bdda650..da305e65 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean @@ -7,7 +7,7 @@ module public import LeanMachineLearning.Online.Bandit.SumRewards public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin -public import LeanMachineLearning.ForMathlib.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg /-! # The Explore-Then-Commit Algorithm @@ -33,7 +33,7 @@ def ETC.nextArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK if hn : n < K * m - 1 then RoundRobin.nextAction hK n else - if hn_eq : n = K * m - 1 then measurableArgmax (empMean' n) h + if hn_eq : n = K * m - 1 then argmax (empMean' n h) else (h ⟨n, by simp⟩).1 /-- The next arm pulled by ETC is chosen in a measurable way. -/ @@ -44,7 +44,7 @@ lemma ETC.measurable_nextArm (hK : 0 < K) (m n : ℕ) : Measurable (nextArm hK m simp only [dite_eq_ite] refine Measurable.ite (by simp) (by fun_prop) ?_ refine Measurable.ite (by simp) ?_ (by fun_prop) - exact measurable_measurableArgmax fun a ↦ by fun_prop + fun_prop /-- The Explore-Then-Commit algorithm: deterministic algorithm that chooses the next arm according to `ETC.nextArm`. -/ @@ -100,8 +100,8 @@ lemma arm_of_lt [Nonempty (Fin K)] phase. -/ lemma arm_mul [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (etcAlgorithm hK m) (stationaryEnv ν) P) (hm : m ≠ 0) : - A (K * m) =ᵐ[P] fun ω ↦ measurableArgmax (empMean' (K * m - 1)) - (history A R (K * m - 1) ω) := by + A (K * m) =ᵐ[P] + fun ω ↦ argmax (empMean' (K * m - 1) (history A R (K * m - 1) ω)) := by have : K * m = (K * m - 1) + 1 := by have : 0 < K * m := Nat.mul_pos hK hm.bot_lt grind @@ -176,8 +176,8 @@ lemma sumRewards_bestArm_le_of_arm_mul_eq [Nonempty (Fin K)] sumRewards A R a (K * m) h := by filter_upwards [arm_mul h hm, pullCount_mul h a, pullCount_mul h (bestArm ν)] with h h_arm ha h_best h_eq - have h_max := isMaxOn_measurableArgmax (empMean' (K * m - 1)) (history A R (K * m - 1) h) - (bestArm ν) + have h_max := isMaxOn_argmax + (empMean' (K * m - 1) (history A R (K * m - 1) h)) (bestArm ν) rw [← h_arm, h_eq] at h_max rw [sumRewards_eq_pullCount_mul_empMean, sumRewards_eq_pullCount_mul_empMean, ha, h_best] · gcongr diff --git a/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean b/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean index 2d1e2fb7..8183b0f8 100644 --- a/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean +++ b/LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean @@ -7,7 +7,7 @@ module public import LeanMachineLearning.Online.Bandit.SumRewards public import LeanMachineLearning.SequentialLearning.Algorithms.RoundRobin -public import LeanMachineLearning.ForMathlib.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg /-! # UCB algorithm @@ -37,13 +37,12 @@ noncomputable def UCB.nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K := have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK if n < K - 1 then RoundRobin.nextAction hK n else - measurableArgmax (fun h a ↦ empMean' n h a + ucbWidth' c n h a) h + argmax (fun a ↦ empMean' n h a + ucbWidth' c n h a) @[fun_prop] lemma UCB.measurable_nextArm (hK : 0 < K) (c : ℝ) (n : ℕ) : Measurable (nextArm hK c n) := by refine Measurable.ite (by simp) (by fun_prop) ?_ have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK - refine measurable_measurableArgmax fun a ↦ ?_ unfold ucbWidth' fun_prop @@ -125,8 +124,8 @@ lemma ucbIndex_le_ucbIndex_arm [Nonempty (Fin K)] have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK simp_rw [h_arm, empMean_eq_empMean' (by grind : n ≠ 0), ucbWidth_eq_ucbWidth' (A := A) (R := R) _ _ _ _ (by grind : n ≠ 0)] - exact isMaxOn_measurableArgmax (fun h a ↦ empMean' (n - 1) h a + ucbWidth' c (n - 1) h a) - (history A R (n - 1) h) a + exact isMaxOn_argmax (fun a ↦ empMean' (n - 1) (history A R (n - 1) h) a + + ucbWidth' c (n - 1) (history A R (n - 1) h) a) _ lemma forall_arm_eq_mod_of_lt [Nonempty (Fin K)] (h : IsAlgEnvSeq A R (ucbAlgorithm hK c) (stationaryEnv ν) P) : diff --git a/LeanMachineLearning/Online/Bandit/BayesRegret.lean b/LeanMachineLearning/Online/Bandit/BayesRegret.lean index 981f1c70..00210b41 100644 --- a/LeanMachineLearning/Online/Bandit/BayesRegret.lean +++ b/LeanMachineLearning/Online/Bandit/BayesRegret.lean @@ -5,7 +5,7 @@ Authors: Paulo Rauber, Rémy Degenne -/ module -public import LeanMachineLearning.ForMathlib.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax +public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg public import LeanMachineLearning.Online.Bandit.Regret /-! @@ -68,14 +68,14 @@ lemma integrable_uncurry_actionMean_comp [Countable 𝓐] [MeasurableSingletonCl /-- A random variable that gives the action with the highest mean feedback. -/ noncomputable -def bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] - (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (ω : Ω) : 𝓐 := - measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω +def bestAction [Nonempty 𝓐] [Fintype 𝓐] (κ : Kernel (𝓔 × 𝓐) ℝ) (E : Ω → 𝓔) (ω : Ω) : 𝓐 := + argmax (fun a ↦ actionMean κ E a ω) @[fun_prop] -lemma measurable_bestAction [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] - {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} (hE : Measurable E) : Measurable (bestAction κ E) := - measurable_measurableArgmax (by fun_prop) +lemma measurable_bestAction [Nonempty 𝓐] [Fintype 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} + (hE : Measurable E) : Measurable (bestAction κ E) := by + unfold bestAction + fun_prop /-- A random variable that gives the gap at time `n`. -/ noncomputable @@ -94,13 +94,13 @@ lemma gap_le_of_mem_Icc [Nonempty 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω Bandits.gap_le_of_mem_Icc (h (E ω)) omit [MeasurableSpace Ω] in -lemma gap_eq_sub [Nonempty 𝓐] [Fintype 𝓐] [Encodable 𝓐] [MeasurableSingletonClass 𝓐] - {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} {n : ℕ} {ω : Ω} : - gap κ E A n ω = actionMean κ E (bestAction κ E ω) ω - actionMean κ E (A n ω) ω := by +lemma gap_eq_sub [Nonempty 𝓐] [Fintype 𝓐] {κ : Kernel (𝓔 × 𝓐) ℝ} {E : Ω → 𝓔} {A : ℕ → Ω → 𝓐} + {n : ℕ} {ω : Ω} : gap κ E A n ω = + actionMean κ E (bestAction κ E ω) ω - actionMean κ E (A n ω) ω := by rw [gap, Bandits.gap] congr apply le_antisymm - · exact ciSup_le (isMaxOn_measurableArgmax (fun ω' a ↦ actionMean κ E a ω') ω) + · exact ciSup_le <| isMaxOn_argmax (fun a ↦ actionMean κ E a ω) · exact Finite.le_ciSup (fun a ↦ actionMean κ E a ω) _ @[fun_prop]