Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 10 additions & 8 deletions LMLTutorial.lean
Original file line number Diff line number Diff line change
@@ -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
2 changes: 1 addition & 1 deletion LMLTutorial/Pages/DefiningAlgorithm.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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'}

Expand Down
3 changes: 2 additions & 1 deletion LeanMachineLearning.lean
Original file line number Diff line number Diff line change
@@ -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
Expand Down

This file was deleted.

27 changes: 27 additions & 0 deletions LeanMachineLearning/ForMathlib/MeasureTheory/Order/Lattice.lean
Original file line number Diff line number Diff line change
@@ -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 : ι → δ → α}
Comment thread
gaetanserre marked this conversation as resolved.
(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
Original file line number Diff line number Diff line change
@@ -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⟩
14 changes: 7 additions & 7 deletions LeanMachineLearning/Online/Bandit/Algorithms/ETC.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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. -/
Expand All @@ -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`. -/
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
9 changes: 4 additions & 5 deletions LeanMachineLearning/Online/Bandit/Algorithms/UCB.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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) :
Expand Down
22 changes: 11 additions & 11 deletions LeanMachineLearning/Online/Bandit/BayesRegret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -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

/-!
Expand Down Expand Up @@ -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
Expand All @@ -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]
Expand Down