Skip to content

Commit fbbfe2a

Browse files
committed
feat : initial algorithm foundation for linUCB
1 parent 395fbb9 commit fbbfe2a

2 files changed

Lines changed: 231 additions & 0 deletions

File tree

‎LeanMachineLearning.lean‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@ public import LeanMachineLearning.MeasureTheory.Constructions.BorelSpace.Measura
44
public import LeanMachineLearning.MeasureTheory.Constructions.Polish.StandardBorel
55
public import LeanMachineLearning.MeasureTheory.Measurable
66
public import LeanMachineLearning.Online.Bandit.Algorithms.ETC
7+
public import LeanMachineLearning.Online.Bandit.Algorithms.LinUCB
78
public import LeanMachineLearning.Online.Bandit.Algorithms.UCB
89
public import LeanMachineLearning.Online.Bandit.ArrayProbSpace
910
public import LeanMachineLearning.Online.Bandit.Regret
Lines changed: 230 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,230 @@
1+
/-
2+
Copyright (c) 2026. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: OpenAI, Fawad Haider
5+
-/
6+
module
7+
8+
public import LeanMachineLearning.Online.Bandit.SumRewards
9+
public import LeanMachineLearning.SequentialLearning.Deterministic
10+
public import LeanMachineLearning.MeasureTheory.Constructions.BorelSpace.MeasurableArgMax
11+
public import Mathlib.LinearAlgebra.Matrix.NonsingularInverse
12+
13+
/-!
14+
# LinUCB for finite-action linear bandits
15+
Chapter 19 of *Bandit Algorithms*:
16+
-/
17+
18+
@[expose] public section
19+
20+
open MeasureTheory ProbabilityTheory Filter Real Finset Learning
21+
22+
open scoped ENNReal NNReal Matrix
23+
24+
namespace Bandits
25+
26+
variable {K d : ℕ}
27+
28+
section Algorithm
29+
30+
namespace LinUCB
31+
32+
abbrev Feature (d : ℕ) := Fin d → ℝ
33+
34+
noncomputable def designMatrix' (reg : ℝ) (x : Fin K → Feature d)
35+
(n : ℕ) (h : Iic n → Fin K × ℝ) : Matrix (Fin d) (Fin d) ℝ :=
36+
reg • 1 + ∑ s : Iic n, Matrix.vecMulVec (x (h s).1) (x (h s).1)
37+
38+
noncomputable def responseVector' (x : Fin K → Feature d)
39+
(n : ℕ) (h : Iic n → Fin K × ℝ) : Feature d :=
40+
∑ s : Iic n, (h s).2 • x (h s).1
41+
42+
noncomputable def thetaHat' (reg : ℝ) (x : Fin K → Feature d)
43+
(n : ℕ) (h : Iic n → Fin K × ℝ) : Feature d :=
44+
Matrix.mulVec (designMatrix' reg x n h)⁻¹ (responseVector' x n h)
45+
46+
noncomputable def estimatedReward' (reg : ℝ) (x : Fin K → Feature d)
47+
(n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fin K) : ℝ :=
48+
dotProduct (thetaHat' reg x n h) (x a)
49+
50+
noncomputable def width' (reg : ℝ) (x : Fin K → Feature d)
51+
(n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fin K) : ℝ :=
52+
√(dotProduct (x a) (Matrix.mulVec (designMatrix' reg x n h)⁻¹ (x a)))
53+
54+
/-- LinUCB optimistic index of an arm.
55+
56+
The parameter `β` is a confidence-radius schedule. Since `h : Iic n → Fin K × ℝ`
57+
contains the observations through time `n`, this index is used to choose the arm
58+
at time `n + 1`, and we evaluate the schedule at `n + 2`
59+
-/
60+
noncomputable def index' (reg : ℝ) (β : ℕ → ℝ) (x : Fin K → Feature d)
61+
(n : ℕ) (h : Iic n → Fin K × ℝ) (a : Fin K) : ℝ :=
62+
estimatedReward' reg x n h a + √(β (n + 2)) * width' reg x n h a
63+
64+
open Classical in
65+
/-- Arm pulled by finite-action LinUCB at time `n + 1`. -/
66+
noncomputable def nextArm (hK : 0 < K) (reg : ℝ) (β : ℕ → ℝ)
67+
(x : Fin K → Feature d)
68+
(_h_index : ∀ n a, Measurable (fun h ↦ index' reg β x n h a))
69+
(n : ℕ) (h : Iic n → Fin K × ℝ) : Fin K :=
70+
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
71+
measurableArgmax (fun h a ↦ index' reg β x n h a) h
72+
73+
@[fun_prop]
74+
lemma measurable_nextArm (hK : 0 < K) (reg : ℝ) (β : ℕ → ℝ)
75+
(x : Fin K → Feature d)
76+
(h_index : ∀ n a, Measurable (fun h ↦ index' reg β x n h a))
77+
(n : ℕ) :
78+
Measurable (nextArm hK reg β x h_index n) := by
79+
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
80+
exact measurable_measurableArgmax fun a ↦ h_index n a
81+
82+
end LinUCB
83+
84+
/-- The finite-action LinUCB algorithm. -/
85+
noncomputable def linUCBAlgorithm (hK : 0 < K) (reg : ℝ) (β : ℕ → ℝ)
86+
(x : Fin K → LinUCB.Feature d)
87+
(h_index : ∀ n a, Measurable (fun h ↦ LinUCB.index' reg β x n h a)) :
88+
Algorithm (Fin K) ℝ :=
89+
detAlgorithm (LinUCB.nextArm hK reg β x h_index) (by fun_prop) ⟨0, hK⟩
90+
91+
end Algorithm
92+
93+
namespace LinUCB
94+
95+
variable {hK : 0 < K} {reg : ℝ} {β : ℕ → ℝ} {x : Fin K → Feature d}
96+
{h_index : ∀ n a, Measurable (fun h ↦ index' reg β x n h a)}
97+
{ν : Kernel (Fin K) ℝ} [IsMarkovKernel ν]
98+
{Ω : Type*} {mΩ : MeasurableSpace Ω}
99+
{P : Measure Ω} [IsProbabilityMeasure P]
100+
{A : ℕ → Ω → Fin K} {R : ℕ → Ω → ℝ}
101+
{n : ℕ} {ω : Ω}
102+
103+
section AlgorithmBehavior
104+
105+
/-- The process-level design matrix built from actions up to time `n` excluded. -/
106+
noncomputable def designMatrix (A : ℕ → Ω → Fin K) (reg : ℝ)
107+
(x : Fin K → Feature d) (n : ℕ) (ω : Ω) : Matrix (Fin d) (Fin d) ℝ :=
108+
reg • 1 + ∑ s ∈ range n, Matrix.vecMulVec (x (A s ω)) (x (A s ω))
109+
110+
/-- The process-level reward-feature vector built from history up to time `n` excluded. -/
111+
noncomputable def responseVector (A : ℕ → Ω → Fin K) (R : ℕ → Ω → ℝ)
112+
(x : Fin K → Feature d) (n : ℕ) (ω : Ω) : Feature d :=
113+
∑ s ∈ range n, R s ω • x (A s ω)
114+
115+
/-- The process-level regularized least-squares estimate. -/
116+
noncomputable def thetaHat (A : ℕ → Ω → Fin K) (R : ℕ → Ω → ℝ)
117+
(reg : ℝ) (x : Fin K → Feature d) (n : ℕ) (ω : Ω) : Feature d :=
118+
Matrix.mulVec (designMatrix A reg x n ω)⁻¹ (responseVector A R x n ω)
119+
120+
/-- The process-level estimated linear reward. -/
121+
noncomputable def estimatedReward (A : ℕ → Ω → Fin K) (R : ℕ → Ω → ℝ)
122+
(reg : ℝ) (x : Fin K → Feature d) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ :=
123+
dotProduct (thetaHat A R reg x n ω) (x a)
124+
125+
/-- The process-level elliptical confidence width. -/
126+
noncomputable def width (A : ℕ → Ω → Fin K) (reg : ℝ)
127+
(x : Fin K → Feature d) (a : Fin K) (n : ℕ) (ω : Ω) : ℝ :=
128+
√(dotProduct (x a) (Matrix.mulVec (designMatrix A reg x n ω)⁻¹ (x a)))
129+
130+
/-- The process-level LinUCB optimistic index. -/
131+
noncomputable def index (A : ℕ → Ω → Fin K) (R : ℕ → Ω → ℝ)
132+
(reg : ℝ) (β : ℕ → ℝ) (x : Fin K → Feature d) (a : Fin K)
133+
(n : ℕ) (ω : Ω) : ℝ :=
134+
estimatedReward A R reg x a n ω + √(β (n + 1)) * width A reg x a n ω
135+
136+
lemma designMatrix_eq_designMatrix' (reg : ℝ) (x : Fin K → Feature d) (n : ℕ)
137+
(ω : Ω) (hn : n ≠ 0) :
138+
designMatrix A reg x n ω =
139+
designMatrix' reg x (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) := by
140+
cases n with
141+
| zero => exact absurd rfl hn
142+
| succ n =>
143+
simp only [designMatrix, designMatrix', IsAlgEnvSeq.hist]
144+
rw [Nat.range_succ_eq_Iic]
145+
exact congrArg (fun S ↦ reg • 1 + S) <|
146+
(Finset.sum_coe_sort (Iic n)
147+
(fun s ↦ Matrix.vecMulVec (x (A s ω)) (x (A s ω)))).symm
148+
149+
lemma responseVector_eq_responseVector' (x : Fin K → Feature d)
150+
(n : ℕ) (ω : Ω) (hn : n ≠ 0) :
151+
responseVector A R x n ω = responseVector' x (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) := by
152+
cases n with
153+
| zero => exact absurd rfl hn
154+
| succ n =>
155+
simp only [responseVector, responseVector', IsAlgEnvSeq.hist]
156+
rw [Nat.range_succ_eq_Iic]
157+
exact (Finset.sum_coe_sort (Iic n) (fun s ↦ R s ω • x (A s ω))).symm
158+
159+
lemma thetaHat_eq_thetaHat' (reg : ℝ) (x : Fin K → Feature d)
160+
(n : ℕ) (ω : Ω) (hn : n ≠ 0) :
161+
thetaHat A R reg x n ω = thetaHat' reg x (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) := by
162+
simp [thetaHat, thetaHat', designMatrix_eq_designMatrix' (A := A) (R := R) reg x n ω hn,
163+
responseVector_eq_responseVector' (A := A) (R := R) x n ω hn]
164+
165+
lemma estimatedReward_eq_estimatedReward' (reg : ℝ) (x : Fin K → Feature d)
166+
(a : Fin K) (n : ℕ) (ω : Ω) (hn : n ≠ 0) :
167+
estimatedReward A R reg x a n ω =
168+
estimatedReward' reg x (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) a := by
169+
simp [estimatedReward, estimatedReward', thetaHat_eq_thetaHat' (A := A) (R := R) reg x n ω hn]
170+
171+
lemma width_eq_width' (reg : ℝ) (x : Fin K → Feature d)
172+
(a : Fin K) (n : ℕ) (ω : Ω) (hn : n ≠ 0) :
173+
width A reg x a n ω = width' reg x (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) a := by
174+
simp [width, width', designMatrix_eq_designMatrix' (A := A) (R := R) reg x n ω hn]
175+
176+
lemma index_eq_index' (reg : ℝ) (β : ℕ → ℝ) (x : Fin K → Feature d)
177+
(a : Fin K) (n : ℕ) (ω : Ω) (hn : n ≠ 0) :
178+
index A R reg β x a n ω =
179+
index' reg β x (n - 1) (IsAlgEnvSeq.hist A R (n - 1) ω) a := by
180+
have htime : n + 1 = n - 1 + 2 := by grind
181+
simp [index, index', estimatedReward_eq_estimatedReward' (A := A) (R := R) reg x a n ω hn,
182+
width_eq_width' (A := A) (R := R) reg x a n ω hn, htime]
183+
184+
/-- The action at time `n + 1` is the finite-action LinUCB argmax for the observed history. -/
185+
lemma arm_ae_eq_linUCBNextArm [Nonempty (Fin K)]
186+
(h : IsAlgEnvSeq A R (linUCBAlgorithm hK reg β x h_index) (stationaryEnv ν) P)
187+
(n : ℕ) :
188+
A (n + 1) =ᵐ[P]
189+
fun ω ↦ nextArm hK reg β x h_index n (IsAlgEnvSeq.hist A R n ω) := by
190+
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
191+
exact h.action_detAlgorithm_ae_eq n
192+
193+
/-- Almost surely, every positive-time action is the finite-action LinUCB argmax. -/
194+
lemma arm_ae_all_eq [Nonempty (Fin K)]
195+
(h : IsAlgEnvSeq A R (linUCBAlgorithm hK reg β x h_index) (stationaryEnv ν) P) :
196+
∀ᵐ ω ∂P,
197+
∀ n, A (n + 1) ω =
198+
nextArm hK reg β x h_index n (IsAlgEnvSeq.hist A R n ω) := by
199+
simp_rw [ae_all_iff]
200+
exact fun n ↦ arm_ae_eq_linUCBNextArm h n
201+
202+
/-- Finite-action LinUCB chooses an arm maximizing the LinUCB index. -/
203+
lemma index_le_index_arm [Nonempty (Fin K)]
204+
(h : IsAlgEnvSeq A R (linUCBAlgorithm hK reg β x h_index) (stationaryEnv ν) P)
205+
(a : Fin K) (hn : n ≠ 0) :
206+
∀ᵐ ω ∂P, index A R reg β x a n ω ≤ index A R reg β x (A n ω) n ω := by
207+
filter_upwards [arm_ae_eq_linUCBNextArm h (n - 1)] with ω h_arm
208+
have hn_succ : n - 1 + 1 = n := by grind
209+
simp only [hn_succ] at h_arm
210+
rw [index_eq_index' (A := A) (R := R) reg β x a n ω hn,
211+
index_eq_index' (A := A) (R := R) reg β x (A n ω) n ω hn]
212+
rw [h_arm]
213+
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
214+
exact isMaxOn_measurableArgmax (fun h a ↦ index' reg β x (n - 1) h a)
215+
(IsAlgEnvSeq.hist A R (n - 1) ω) a
216+
217+
/-- Almost surely, the selected arm maximizes the LinUCB index at every positive time. -/
218+
lemma forall_index_le_index_arm [Nonempty (Fin K)]
219+
(h : IsAlgEnvSeq A R (linUCBAlgorithm hK reg β x h_index) (stationaryEnv ν) P)
220+
(a : Fin K) :
221+
∀ᵐ ω ∂P, ∀ n, n ≠ 0 →
222+
index A R reg β x a n ω ≤ index A R reg β x (A n ω) n ω := by
223+
simp_rw [ae_all_iff]
224+
exact fun n hn ↦ index_le_index_arm h a hn
225+
226+
end AlgorithmBehavior
227+
228+
end LinUCB
229+
230+
end Bandits

0 commit comments

Comments
 (0)