Skip to content
Merged
Show file tree
Hide file tree
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
91 changes: 63 additions & 28 deletions LeanBandits/Algorithm.lean
Original file line number Diff line number Diff line change
Expand Up @@ -44,21 +44,6 @@ structure Environment (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] wh
instance (env : Environment α R) (n : ℕ) : IsMarkovKernel (env.feedback n) := env.h_feedback n
instance (env : Environment α R) : IsMarkovKernel env.ν0 := env.hp0

/-- A deterministic algorithm. -/
noncomputable
def detAlgorithm (nextaction : (n : ℕ) → (Iic n → α × R) → α)
(h_next : ∀ n, Measurable (nextaction n)) (action0 : α) :
Algorithm α R where
policy n := Kernel.deterministic (nextaction n) (h_next n)
p0 := Measure.dirac action0

/-- A stationary environment, in which the distribution of the next reward depends only on the last
action. -/
@[simps]
def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R where
feedback _ := ν.prodMkLeft _
ν0 := ν

/-- Kernel describing the distribution of the next action-reward pair given the history
up to `n`. -/
noncomputable
Expand Down Expand Up @@ -200,33 +185,83 @@ lemma condDistrib_reward_zero [StandardBorelSpace R] [Nonempty R]

section DetAlgorithm

/-- A deterministic algorithm. -/
@[simps]
noncomputable
def detAlgorithm (nextaction : (n : ℕ) → (Iic n → α × R) → α)
(h_next : ∀ n, Measurable (nextaction n)) (action0 : α) :
Algorithm α R where
policy n := Kernel.deterministic (nextaction n) (h_next n)
p0 := Measure.dirac action0

variable {nextaction : (n : ℕ) → (Iic n → α × R) → α} {h_next : ∀ n, Measurable (nextaction n)}
{action0 : α} {env : Environment α R}

lemma HasLaw_action_zero_detAlgorithm :
HasLaw (action 0) (Measure.dirac action0)
(trajMeasure (detAlgorithm nextaction h_next action0) env) where
local notation "𝔓" => trajMeasure (detAlgorithm nextaction h_next action0) env

lemma HasLaw_action_zero_detAlgorithm : HasLaw (action 0) (Measure.dirac action0) 𝔓 where
map_eq := (hasLaw_action_zero _ _).map_eq

lemma action_zero_detAlgorithm [MeasurableSingletonClass α] :
action 0 =ᵐ[trajMeasure (detAlgorithm nextaction h_next action0) env] fun _ ↦ action0 := by
have h_eq : ∀ᵐ x ∂((trajMeasure (detAlgorithm nextaction h_next action0) env).map (action 0)), x
= action0 := by
lemma action_zero_detAlgorithm [MeasurableSingletonClass α] : action 0 =ᵐ[𝔓] fun _ ↦ action0 := by
have h_eq : ∀ᵐ x ∂((𝔓).map (action 0)), x = action0 := by
rw [(hasLaw_action_zero _ _).map_eq]
simp [detAlgorithm]
exact ae_of_ae_map (by fun_prop) h_eq

lemma action_detAlgorithm_ae_eq (n : ℕ) :
action (n + 1) =ᵐ[trajMeasure (detAlgorithm nextaction h_next action0) env]
fun h ↦ nextaction n (fun i ↦ h i) := by
lemma action_detAlgorithm_ae_eq
[StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(n : ℕ) :
action (n + 1) =ᵐ[𝔓] fun h ↦ nextaction n (fun i ↦ h i) := by
have h := condDistrib_action (detAlgorithm nextaction h_next action0) env n
simp only [detAlgorithm_policy] at h
sorry

example [MeasurableSingletonClass α] :
∀ᵐ h ∂(trajMeasure (detAlgorithm nextaction h_next action0) env),
action 0 h = action0 ∧ ∀ n, action (n + 1) h = nextaction n (fun i ↦ h i) := by
example [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R] :
∀ᵐ h ∂𝔓, action 0 h = action0 ∧ ∀ n, action (n + 1) h = nextaction n (fun i ↦ h i) := by
rw [eventually_and, ae_all_iff]
exact ⟨action_zero_detAlgorithm, action_detAlgorithm_ae_eq⟩

end DetAlgorithm

section stationaryEnv

/-- A stationary environment, in which the distribution of the next reward depends only on the last
action. -/
@[simps]
def stationaryEnv (ν : Kernel α R) [IsMarkovKernel ν] : Environment α R where
feedback _ := ν.prodMkLeft _
ν0 := ν

variable {alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν]

local notation "𝔓" => trajMeasure alg (stationaryEnv ν)

lemma condDistrib_reward_stationaryEnv [StandardBorelSpace α] [Nonempty α]
[StandardBorelSpace R] [Nonempty R] (n : ℕ) :
condDistrib (reward n) (action n) 𝔓 =ᵐ[(𝔓).map (action n)] ν := by
cases n with
| zero =>
rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (by fun_prop)]
change (𝔓).map (step 0) = (𝔓).map (action 0) ⊗ₘ ν
rw [(hasLaw_action_zero alg (stationaryEnv ν)).map_eq,
(hasLaw_step_zero alg (stationaryEnv ν)).map_eq, stationaryEnv_ν0]
| succ n =>
have h_eq := condDistrib_reward alg (stationaryEnv ν) n
rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (by fun_prop)] at h_eq ⊢
have : (𝔓).map (action (n + 1)) = ((𝔓).map (fun x ↦ (hist n x, action (n + 1) x))).snd := by
rw [Measure.snd_map_prodMk (by fun_prop)]
simp only [stationaryEnv_feedback] at h_eq
rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq,
Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)]
congr

lemma condIndepFun_reward_hist_action [StandardBorelSpace α] [Nonempty α]
[StandardBorelSpace R] [Nonempty R] (n : ℕ) :
CondIndepFun (MeasurableSpace.comap (action (n + 1)) inferInstance)
(measurable_action _).comap_le (reward (n + 1)) (hist n) (𝔓) :=
condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft
(by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward alg (stationaryEnv ν) n)

end stationaryEnv

end Learning
21 changes: 3 additions & 18 deletions LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -138,22 +138,8 @@ lemma condDistrib_reward' [StandardBorelSpace α] [Nonempty α] [StandardBorelSp
lemma condDistrib_reward [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
condDistrib (reward n) (arm n) (Bandit.trajMeasure alg ν)
=ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν := by
cases n with
| zero =>
rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (by fun_prop)]
change (Bandit.trajMeasure alg ν).map (fun h ↦ h 0)
= (Bandit.trajMeasure alg ν).map (arm 0) ⊗ₘ ν
rw [(hasLaw_arm_zero alg ν).map_eq, (hasLaw_step_zero alg ν).map_eq]
| succ n =>
have h_eq := condDistrib_reward' alg ν n
rw [condDistrib_ae_eq_iff_measure_eq_compProd₀ (by fun_prop) (by fun_prop)] at h_eq ⊢
have : (Bandit.trajMeasure alg ν).map (arm (n + 1))
= ((Bandit.trajMeasure alg ν).map (fun x ↦ (hist n x, arm (n + 1) x))).snd := by
rw [Measure.snd_map_prodMk (by fun_prop)]
rw [this, ← Measure.snd_prodAssoc_compProd_prodMkLeft, ← h_eq,
Measure.snd_map_prodMk (by fun_prop), Measure.map_map (by fun_prop) (by fun_prop)]
congr
=ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν :=
Learning.condDistrib_reward_stationaryEnv n

lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
Expand All @@ -167,8 +153,7 @@ lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α R} {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ) :
CondIndepFun (MeasurableSpace.comap (arm (n + 1)) inferInstance)
(measurable_arm _).comap_le (reward (n + 1)) (hist n) (Bandit.trajMeasure alg ν) :=
condIndepFun_of_exists_condDistrib_prod_ae_eq_prodMkLeft
(by fun_prop) (by fun_prop) (by fun_prop) (condDistrib_reward' alg ν n)
Learning.condIndepFun_reward_hist_action n

section DetAlgorithm

Expand Down
15 changes: 7 additions & 8 deletions LeanBandits/ETC.lean
Original file line number Diff line number Diff line change
Expand Up @@ -59,11 +59,11 @@ lemma arm_ae_eq_etcNextArm (n : ℕ) :
exact arm_detAlgorithm_ae_eq n

lemma pullCount_mul (a : Fin K) :
(fun ω ↦ pullCount (arm · ω) a (K * m)) =ᵐ[𝔓b] fun _ ↦ m := by
pullCount a (K * m) =ᵐ[𝔓b] fun _ ↦ m := by
sorry

lemma pullCount_of_ge (a : Fin K) {n : ℕ} (hn : K * m ≤ n) :
(fun ω ↦ pullCount (arm · ω) a n)
pullCount a n
=ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by
sorry

Expand All @@ -84,9 +84,9 @@ lemma prob_arm_mul_eq_le (a : Fin K) :
_ ≤ (𝔓).real {ω | ∑ s ∈ range (K * m), (if (arm s ω.1) = bestArm ν then (reward s ω.1) else 0)
≤ ∑ s ∈ range (K * m), if (arm s ω.1) = a then (reward s ω.1) else 0} := by
sorry
_ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (arm · ω.1) (bestArm ν) (K * m)),
_ = (𝔓).real {ω | ∑ s ∈ Icc 1 (pullCount (bestArm ν) (K * m) ω.1),
rewardByCount (bestArm ν) s ω.1 ω.2
≤ ∑ s ∈ Icc 1 (pullCount (arm · ω.1) a (K * m)), rewardByCount a s ω.1 ω.2} := by
≤ ∑ s ∈ Icc 1 (pullCount a (K * m) ω.1), rewardByCount a s ω.1 ω.2} := by
sorry
_ = (𝔓).real {ω | ∑ s ∈ Icc 1 m, rewardByCount (bestArm ν) s ω.1 ω.2
≤ ∑ s ∈ Icc 1 m, rewardByCount a s ω.1 ω.2} := by
Expand Down Expand Up @@ -118,9 +118,9 @@ lemma prob_arm_mul_eq_le (a : Fin K) :
norm_num

lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) :
𝔓b[fun ω ↦ (pullCount (arm · ω) a n : ℝ)]
𝔓b[fun ω ↦ (pullCount a n ω : ℝ)]
≤ m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4) := by
have : (fun ω ↦ (pullCount (arm · ω) a n : ℝ))
have : (fun ω ↦ (pullCount a n ω : ℝ))
=ᵐ[𝔓b] fun ω ↦ m + (n - K * m) * {ω' | arm (K * m) ω' = a}.indicator (fun _ ↦ 1) ω := by
filter_upwards [pullCount_of_ge a hn] with ω h
simp only [h, Set.indicator_apply, Set.mem_setOf_eq, mul_ite, mul_one, mul_zero, Nat.cast_add,
Expand All @@ -143,8 +143,7 @@ lemma expectation_pullCount_le (a : Fin K) {n : ℕ} (hn : K * m ≤ n) :
· exact (measurableSet_singleton _).preimage (by fun_prop)

lemma regret_le (n : ℕ) (hn : K * m ≤ n) :
𝔓b[fun ω ↦ regret ν (arm · ω) n]
≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by
𝔓b[regret ν n] ≤ ∑ a, gap ν a * (m + (n - K * m) * Real.exp (- (m : ℝ) * gap ν a ^ 2 / 4)) := by
simp_rw [regret_eq_sum_pullCount_mul_gap]
rw [integral_finset_sum]
swap
Expand Down
Loading