11/-
22Copyright (c) 2025 Rémy Degenne. All rights reserved.
33Released under Apache 2.0 license as described in the file LICENSE.
4- Authors: Rémy Degenne
4+ Authors: Rémy Degenne, Paulo Rauber
55-/
66import Mathlib
77
@@ -22,115 +22,122 @@ def MeasurableEquiv.piIicZero (α : Type*) [MeasurableSpace α] :
2222
2323namespace Bandits
2424
25- variable {α : Type *} {mα : MeasurableSpace α}
25+ variable {α R : Type *} {mα : MeasurableSpace α} {mR : MeasurableSpace R }
2626
2727section MeasureSpace
2828
29- /-- A bandit interaction between an agent described by a policy and an environment given by
30- reward distributions. -/
31- structure Bandit (α : Type *) [MeasurableSpace α] where
32- /-- Conditional distribution of the rewards given the arm pulled. -/
33- ν : Kernel α ℝ
34- hν : IsMarkovKernel ν
29+ /-- A stochastic, sequential algorithm. -/
30+ structure Algorithm (α R : Type *) [MeasurableSpace α] [MeasurableSpace R] where
3531 /-- Policy or sampling rule: distribution of the next pull. -/
36- policy : (n : ℕ) → Kernel (Iic n → α × ℝ ) α
37- h_policy n : IsMarkovKernel (policy n)
32+ policy : (n : ℕ) → Kernel (Iic n → α × R ) α
33+ [ h_policy : ∀ n, IsMarkovKernel (policy n)]
3834 /-- Distribution of the first pull. -/
3935 p0 : Measure α
40- hp0 : IsProbabilityMeasure p0
36+ [ hp0 : IsProbabilityMeasure p0]
4137
42- instance (b : Bandit α) : IsMarkovKernel b.ν := b.hν
43- instance (b : Bandit α) (n : ℕ) : IsMarkovKernel (b.policy n) := b.h_policy n
44- instance (b : Bandit α) : IsProbabilityMeasure b.p0 := b.hp0
38+ instance (alg : Algorithm α R) (n : ℕ) : IsMarkovKernel (alg.policy n) := alg.h_policy n
39+ instance (alg : Algorithm α R) : IsProbabilityMeasure alg.p0 := alg.hp0
4540
4641namespace Bandit
4742
4843/-- Kernel describing the distribution of the next arm-reward pair given the history up to `n`. -/
4944noncomputable
50- def stepKernel (b : Bandit α ) (n : ℕ) : Kernel (Iic n → α × ℝ ) (α × ℝ ) :=
51- (b .policy n) ⊗ₖ b. ν.prodMkLeft (Iic n → α × ℝ )
45+ def stepKernel (alg : Algorithm α R ) (ν : Kernel α R) ( n : ℕ) : Kernel (Iic n → α × R ) (α × R ) :=
46+ (alg .policy n) ⊗ₖ ν.prodMkLeft (Iic n → α × R )
5247
53- instance (b : Bandit α) (n : ℕ) : IsMarkovKernel (b.stepKernel n) := by
48+ instance (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
49+ IsMarkovKernel (stepKernel alg ν n) := by
5450 rw [stepKernel]
5551 infer_instance
5652
5753@[simp]
58- lemma fst_stepKernel (b : Bandit α) (n : ℕ) : (b.stepKernel n).fst = b.policy n := by
54+ lemma fst_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
55+ (stepKernel alg ν n).fst = alg.policy n := by
5956 rw [stepKernel, Kernel.fst_compProd]
6057
6158@[simp]
62- lemma snd_stepKernel (b : Bandit α) (n : ℕ) : (b.stepKernel n).snd = b.ν ∘ₖ b.policy n := by
59+ lemma snd_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
60+ (stepKernel alg ν n).snd = ν ∘ₖ alg.policy n := by
6361 rw [stepKernel, Kernel.snd_compProd_prodMkLeft]
6462
6563/-- Kernel sending a partial trajectory of the bandit interaction `Iic n → α × ℝ` to a measure
6664on `ℕ → α × ℝ`, supported on full trajectories that start with the partial one. -/
67- noncomputable def traj (b : Bandit α) (n : ℕ) : Kernel (Iic n → α × ℝ) (ℕ → α × ℝ) :=
68- ProbabilityTheory.Kernel.traj (X := fun _ ↦ α × ℝ) b.stepKernel n
69-
70- instance (b : Bandit α) (n : ℕ) : IsMarkovKernel (b.traj n) := by
71- rw [traj]
72- infer_instance
65+ noncomputable def traj (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
66+ Kernel (Iic n → α × R) (ℕ → α × R) :=
67+ ProbabilityTheory.Kernel.traj (X := fun _ ↦ α × R) (stepKernel alg ν) n
68+ deriving IsMarkovKernel
7369
7470/-- Measure on the sequence of arms pulled and rewards observed generated by the bandit. -/
7571noncomputable
76- def trajMeasure (b : Bandit α) : Measure (ℕ → α × ℝ ) :=
77- (b. traj 0 ) ∘ₘ ((b .p0 ⊗ₘ b. ν).map (MeasurableEquiv.piIicZero _).symm)
72+ def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R ) :=
73+ (traj alg ν 0 ) ∘ₘ ((alg .p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _).symm)
7874
7975/-- Measure of an infinite stream of rewards from each arm. -/
8076noncomputable
81- def streamMeasure (b : Bandit α) : Measure (ℕ → α → ℝ) :=
82- Measure.infinitePi fun _ ↦ Measure.infinitePi b.ν
77+ def streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α → R) :=
78+ Measure.infinitePi fun _ ↦ Measure.infinitePi ν
79+ deriving IsProbabilityMeasure
80+
81+ instance (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
82+ IsProbabilityMeasure (trajMeasure alg ν) := by
83+ rw [trajMeasure]
84+ have : IsProbabilityMeasure ((alg.p0 ⊗ₘ ν).map (MeasurableEquiv.piIicZero _).symm) :=
85+ isProbabilityMeasure_map <| by fun_prop
86+ infer_instance
8387
8488/-- Joint distribution of the sequence of arm pulled and rewards, and a stream of independent
8589rewards from all arms. -/
8690noncomputable
87- def measure (b : Bandit α) : Measure ((ℕ → α × ℝ) × (ℕ → α → ℝ)) :=
88- (b.trajMeasure).prod (b.streamMeasure)
89-
90- instance (b : Bandit α) : IsProbabilityMeasure b.trajMeasure := by
91- rw [Bandit.trajMeasure]
92- have : IsProbabilityMeasure ((b.p0 ⊗ₘ b.ν).map (MeasurableEquiv.piIicZero _).symm) :=
93- isProbabilityMeasure_map <| by fun_prop
94- infer_instance
95-
96- instance (b : Bandit α) : IsProbabilityMeasure b.streamMeasure := by
97- rw [streamMeasure]
98- infer_instance
99-
100- instance (b : Bandit α) : IsProbabilityMeasure b.measure := by
101- rw [measure]
102- infer_instance
91+ def measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
92+ Measure ((ℕ → α × R) × (ℕ → α → R)) :=
93+ (trajMeasure alg ν).prod (streamMeasure ν)
94+ deriving IsProbabilityMeasure
10395
10496end Bandit
10597
10698/-- `arm n` is the arm pulled at time `n`. This is a random variable on the measurable space
10799`ℕ → α × ℝ`. -/
108- def arm (n : ℕ) (h : ℕ → α × ℝ ) : α := (h n).1
100+ def arm (n : ℕ) (h : ℕ → α × R ) : α := (h n).1
109101
110102/-- `reward n` is the reward at time `n`. This is a random variable on the measurable space
111- `ℕ → α × ℝ `. -/
112- def reward (n : ℕ) (h : ℕ → α × ℝ ) : ℝ := (h n).2
103+ `ℕ → α × R `. -/
104+ def reward (n : ℕ) (h : ℕ → α × R ) : R := (h n).2
113105
114106/-- `hist n` is the history up to time `n`. This is a random variable on the measurable space
115- `ℕ → α × ℝ`. -/
116- def hist (n : ℕ) (h : ℕ → α × ℝ) : Iic n → α × ℝ := fun i ↦ h i
107+ `ℕ → α × R`. -/
108+ def hist (n : ℕ) (h : ℕ → α × R) : Iic n → α × R := fun i ↦ h i
109+
110+ @[fun_prop]
111+ lemma measurable_arm (n : ℕ) : Measurable (arm n (α := α) (R := R)) := by unfold arm; fun_prop
112+
113+ @[fun_prop]
114+ lemma measurable_reward (n : ℕ) : Measurable (reward n (α := α) (R := R)) := by
115+ unfold reward; fun_prop
116+
117+ @[fun_prop]
118+ lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop
117119
118120/-- Filtration of the bandit process. -/
119121def ℱ (α : Type *) [MeasurableSpace α] :
120- Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × ℝ)) :=
121- MeasureTheory.Filtration.piLE (X := fun _ ↦ α × ℝ)
122-
123- lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α] (b : Bandit α) (n : ℕ) :
124- condDistrib (fun h ↦ (arm n h, reward n h)) (hist n) b.trajMeasure = b.stepKernel n := by
122+ Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) :=
123+ MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R)
124+
125+ lemma condDistrib_arm_reward [StandardBorelSpace α] [Nonempty α]
126+ [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν]
127+ (n : ℕ) :
128+ condDistrib (fun h ↦ (arm n h, reward n h)) (hist n) (Bandit.trajMeasure alg ν)
129+ = Bandit.stepKernel alg ν n := by
125130 sorry
126131
127- lemma condDistrib_reward (b : Bandit α) (n : ℕ) :
128- condDistrib (reward n) (arm n) b.trajMeasure = b.ν := by
132+ lemma condDistrib_reward [StandardBorelSpace R] [Nonempty R] (alg : Algorithm α R)
133+ (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
134+ condDistrib (reward n) (arm n) (Bandit.trajMeasure alg ν) = ν := by
129135 sorry
130136
131- lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] (b : Bandit α) (n : ℕ) :
132- condDistrib (arm n) (hist n) b.trajMeasure = b.policy n := by
133- rw [← b.fst_stepKernel, ← condDistrib_arm_reward]
137+ lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
138+ (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
139+ condDistrib (arm n) (hist n) (Bandit.trajMeasure alg ν) = alg.policy n := by
140+ rw [← Bandit.fst_stepKernel alg ν n, ← condDistrib_arm_reward alg ν n]
134141 sorry
135142
136143end MeasureSpace
0 commit comments