Skip to content

Commit 81eb277

Browse files
committed
use measurable argmax
1 parent f1ab9cc commit 81eb277

3 files changed

Lines changed: 109 additions & 8 deletions

File tree

‎LeanBandits/Bandit.lean‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
/-
22
Copyright (c) 2025 Rémy Degenne. All rights reserved.
33
Released under Apache 2.0 license as described in the file LICENSE.
4-
Authors: Rémy Degenne
4+
Authors: Rémy Degenne, Paulo Rauber
55
-/
66
import Mathlib
77

‎LeanBandits/ETC.lean‎

Lines changed: 107 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,8 @@
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+
-/
16
import Mathlib.Probability.Moments.SubGaussian
27
import LeanBandits.Bandit
38
import LeanBandits.Regret
@@ -9,28 +14,124 @@ import LeanBandits.Regret
914
open MeasureTheory ProbabilityTheory Finset
1015
open 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+
1288
namespace 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. -/
32133
noncomputable
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

36137
end Bandits

‎LeanBandits/Regret.lean‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
/-
22
Copyright (c) 2025 Rémy Degenne. All rights reserved.
33
Released under Apache 2.0 license as described in the file LICENSE.
4-
Authors: Rémy Degenne
4+
Authors: Rémy Degenne, Paulo Rauber
55
-/
66
import Mathlib
77
import LeanBandits.Bandit

0 commit comments

Comments
 (0)