Skip to content
Merged
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
31 changes: 31 additions & 0 deletions LeanBandits/UCB.lean
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ 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
-/
import LeanBandits.AlgorithmBuilding
import LeanBandits.Regret

/-!
Expand All @@ -18,6 +19,36 @@ namespace Bandits

variable {α : Type*} {mα : MeasurableSpace α} {ν : Kernel α ℝ} {k : ℕ → α} {t : ℕ} {a : α}

section Algorithm

variable [Nonempty α] [DecidableEq α] [Finite α] [Encodable α] [MeasurableSingletonClass α]

/-- The exploration bonus of the UCB algorithm, which corresponds to the width of
a confidence interval. -/
noncomputable def ucbWidth' (c : ℝ) (n : ℕ) (h : Iic n → α × ℝ) (a : α) : ℝ :=
√(c * log (n + 1) / (pullCount' n h a))

open Classical in
/-- Arm pulled by the UCB algorithm at time `n + 1`. -/
noncomputable
def ucbNextArm (c : ℝ) (n : ℕ) (h : Iic n → α × ℝ) : α :=
measurableArgmax (fun h a ↦ empMean' n h a + ucbWidth' c n h a) h

@[fun_prop]
lemma measurable_ucbNextArm (c : ℝ) (n : ℕ) : Measurable (ucbNextArm c n (α := α)) := by
classical
refine measurable_measurableArgmax fun a ↦ ?_
unfold ucbWidth'
fun_prop

/-- The UCB algorithm. -/
noncomputable
def ucbAlgorithm (c : ℝ) : Algorithm α ℝ where
policy n := Kernel.deterministic (ucbNextArm c n) (by fun_prop)
p0 := Measure.dirac (Classical.arbitrary α)

end Algorithm

variable [Fintype α] [Nonempty α] {c : ℝ} {μ : α → ℝ} {N : α → ℕ} {a : α}

/-- The exploration bonus of the UCB algorithm, which corresponds to the width of
Expand Down