1+ /-
2+ Copyright (c) 2025 Rémy Degenne. All rights reserved.
3+ Released under Apache 2.0 license as described in the file LICENSE.
4+ Authors: Rémy Degenne
5+ -/
16import Mathlib.Probability.Moments.SubGaussian
27import LeanBandits.Bandit
38import LeanBandits.Regret
@@ -9,28 +14,124 @@ import LeanBandits.Regret
914open MeasureTheory ProbabilityTheory Finset
1015open scoped ENNReal NNReal
1116
17+ section MeasurableArgmax -- copied from PR #27579 (and changed from argmin to argmax)
18+
19+ lemma measurable_encode {α : Type *} {_ : MeasurableSpace α} [Encodable α]
20+ [MeasurableSingletonClass α] :
21+ Measurable (Encodable.encode (α := α)) := by
22+ refine measurable_to_nat fun a ↦ ?_
23+ have : Encodable.encode ⁻¹' {Encodable.encode a} = {a} := by ext; simp
24+ rw [this]
25+ exact measurableSet_singleton _
26+
27+ lemma measurableEmbedding_encode (α : Type *) {_ : MeasurableSpace α} [Encodable α]
28+ [MeasurableSingletonClass α] :
29+ MeasurableEmbedding (Encodable.encode (α := α)) where
30+ injective := Encodable.encode_injective
31+ measurable := measurable_encode
32+ measurableSet_image' _ _ := .of_discrete
33+
34+ section Finite
35+
36+ variable {𝓧 𝓨 α : Type *} {m𝓧 : MeasurableSpace 𝓧} {m𝓨 : MeasurableSpace 𝓨}
37+ {mα : MeasurableSpace α} [TopologicalSpace α] [LinearOrder α]
38+ [OpensMeasurableSpace α] [OrderClosedTopology α] [SecondCountableTopology α]
39+
40+ lemma measurableSet_isMax [Countable 𝓨]
41+ {f : 𝓧 → 𝓨 → α} (hf : ∀ y, Measurable (fun x ↦ f x y)) (y : 𝓨) :
42+ MeasurableSet {x | ∀ z, f x z ≤ f x y} := by
43+ rw [show {x | ∀ y', f x y' ≤ f x y} = ⋂ y', {x | f x y' ≤ f x y} by ext; simp]
44+ exact MeasurableSet.iInter fun z ↦ measurableSet_le (by fun_prop) (by fun_prop)
45+
46+ lemma exists_isMaxOn' {α : Type *} [LinearOrder α]
47+ [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] (f : 𝓧 → 𝓨 → α) (x : 𝓧) :
48+ ∃ n : ℕ, ∃ y, n = Encodable.encode y ∧ ∀ z, f x z ≤ f x y := by
49+ obtain ⟨y, h⟩ := Finite.exists_max (f x)
50+ exact ⟨Encodable.encode y, y, rfl, h⟩
51+
52+ /-- A measurable argmax function. -/
53+ noncomputable
54+ def measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨]
55+ (f : 𝓧 → 𝓨 → α)
56+ [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y]
57+ (x : 𝓧) :
58+ 𝓨 :=
59+ (measurableEmbedding_encode 𝓨).invFun (Nat.find (exists_isMaxOn' f x))
60+
61+ lemma measurable_measurableArgmax [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨]
62+ {f : 𝓧 → 𝓨 → α}
63+ [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y]
64+ (hf : ∀ y, Measurable (fun x ↦ f x y)) :
65+ Measurable (measurableArgmax f) := by
66+ refine (MeasurableEmbedding.measurable_invFun (measurableEmbedding_encode 𝓨)).comp ?_
67+ refine measurable_find _ fun n ↦ ?_
68+ have : {x | ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y}
69+ = ⋃ y, ({x | n = Encodable.encode y} ∩ {x | ∀ z, f x z ≤ f x y}) := by ext; simp
70+ rw [this]
71+ refine MeasurableSet.iUnion fun y ↦ (MeasurableSet.inter (by simp) ?_)
72+ exact measurableSet_isMax (by fun_prop) y
73+
74+ lemma isMaxOn_measurableArgmax {α : Type *} [LinearOrder α]
75+ [Nonempty 𝓨] [Finite 𝓨] [Encodable 𝓨] [MeasurableSingletonClass 𝓨]
76+ (f : 𝓧 → 𝓨 → α)
77+ [∀ x, DecidablePred fun n ↦ ∃ y, n = Encodable.encode y ∧ ∀ (z : 𝓨), f x z ≤ f x y]
78+ (x : 𝓧) (z : 𝓨) :
79+ f x z ≤ f x (measurableArgmax f x) := by
80+ obtain ⟨y, h_eq, h_le⟩ := Nat.find_spec (exists_isMaxOn' f x)
81+ refine le_trans (h_le z) (le_of_eq ?_)
82+ rw [measurableArgmax, h_eq,
83+ MeasurableEmbedding.leftInverse_invFun (measurableEmbedding_encode 𝓨) y]
84+
85+ end Finite
86+ end MeasurableArgmax
87+
1288namespace Bandits
1389
14- def etcArm {K : ℕ} (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K :=
90+ variable {K : ℕ}
91+
92+ /-- The empirical mean of arm `a` at time `n` but weighted by `m`.
93+ We will use it only for `n = K * m - 1`, time for which there is indeed `m` samples for each arm,
94+ but for reasons that have to do with type equalities we define it for arbitrary `n`. -/
95+ noncomputable
96+ def empMeanETC (m n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fin K) :=
97+ (∑ s : Iic n, if (h s).1 = a then (h s).2 else 0 ) / m
98+
99+ @[fun_prop]
100+ lemma measurable_empMeanETC (m n : ℕ) (a : Fin K) :
101+ Measurable (fun h : Iic n → Fin K × ℝ ↦ empMeanETC m n h a) := by
102+ simp only [empMeanETC]
103+ have h_meas s :
104+ Measurable (fun (h : Iic n → Fin K × ℝ) ↦ if (h s).1 = a then (h s).2 else 0 ) := by
105+ refine Measurable.ite ?_ (by fun_prop) (by fun_prop)
106+ change MeasurableSet ((fun h : Iic n → Fin K × ℝ ↦ (h s).1 ) ⁻¹' {a})
107+ exact MeasurableSet.preimage (measurableSet_singleton _) (by fun_prop)
108+ fun_prop
109+
110+ /-- Arm pulled by the ETC algorithm. -/
111+ noncomputable
112+ def etcArm (hK : 0 < K) (m n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K :=
113+ have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
15114 if hn : n < K * m - 1 then
16115 ⟨n % K, Nat.mod_lt _ hK⟩
17116 else
18117 if hn_eq : n = K * m - 1 then
19- let hatμ a := (∑ s : Iic n, if (h s).1 = a then (h s).2 else 0 ) / m
20- ⟨0 , hK⟩ -- TODO placeholder, replace by the argmax
118+ measurableArgmax (empMeanETC m n) h
21119 else
22120 have : 0 < n := by grind
23121 (h ⟨n - 1 , by simp⟩).1
24122
25123@[fun_prop]
26- lemma measurable_etcArm {K : ℕ} (hK : 0 < K) (m n : ℕ) : Measurable (etcArm hK m n) := by
124+ lemma measurable_etcArm (hK : 0 < K) (m n : ℕ) : Measurable (etcArm hK m n) := by
125+ have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
27126 unfold etcArm
28127 simp only [dite_eq_ite]
29- fun_prop
128+ refine Measurable.ite (by simp) (by fun_prop) ?_
129+ refine Measurable.ite (by simp) ?_ (by fun_prop)
130+ exact measurable_measurableArgmax fun a ↦ by fun_prop
30131
31132/-- The Explore-Then-Commit Kernel, which describes the arm pulled by the ETC algorithm. -/
32133noncomputable
33- def etcKernel {K : ℕ} (hK : 0 < K) (m n : ℕ) : Kernel (Iic n → Fin K × ℝ) (Fin K) :=
134+ def etcKernel (hK : 0 < K) (m n : ℕ) : Kernel (Iic n → Fin K × ℝ) (Fin K) :=
34135 Kernel.deterministic (etcArm hK m n) (by fun_prop)
35136
36137end Bandits
0 commit comments