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
90 changes: 18 additions & 72 deletions LeanBandits/Bandit/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,6 @@ import LeanBandits.ForMathlib.IndepFun
import LeanBandits.ForMathlib.IndepInfinitePi
import LeanBandits.ForMathlib.KernelRepresentation
import LeanBandits.ForMathlib.StandardBorel
import LeanBandits.SequentialLearning.Deterministic
import LeanBandits.SequentialLearning.FiniteActions
import LeanBandits.SequentialLearning.StationaryEnv

Expand All @@ -26,32 +25,6 @@ variable {α R : Type*} {mα : MeasurableSpace α} {mR : MeasurableSpace R}

section MeasureSpace

namespace Bandit

/-- Kernel describing the distribution of the next action-reward pair given the history up to
time `n`. -/
noncomputable
def stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
Kernel (Iic n → α × R) (α × R) :=
Learning.stepKernel alg (stationaryEnv ν) n
deriving IsMarkovKernel

@[simp]
lemma fst_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
(stepKernel alg ν n).fst = alg.policy n := by
rw [stepKernel, Learning.fst_stepKernel]

@[simp]
lemma snd_stepKernel (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
(stepKernel alg ν n).snd = ν ∘ₖ alg.policy n := by
rw [stepKernel, Learning.stepKernel, stationaryEnv_feedback, Kernel.snd_compProd_prodMkLeft]

/-- Measure on the sequence of actions pulled and rewards observed generated by the bandit. -/
noncomputable
def trajMeasure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] : Measure (ℕ → α × R) :=
Learning.trajMeasure alg (stationaryEnv ν)
deriving IsProbabilityMeasure

/-- Measure of an infinite stream of rewards from each action. -/
noncomputable
def streamMeasure (ν : Kernel α R) : Measure (ℕ → α → R) :=
Expand All @@ -61,26 +34,6 @@ instance (ν : Kernel α R) [IsMarkovKernel ν] : IsProbabilityMeasure (streamMe
unfold streamMeasure
infer_instance

/-- Joint distribution of the sequence of action pulled and rewards, and a stream of independent
rewards from all actions. -/
noncomputable
def measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
Measure ((ℕ → α × R) × (ℕ → α → R)) :=
(trajMeasure alg ν).prod (streamMeasure ν)
deriving IsProbabilityMeasure

@[simp]
lemma fst_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
(measure alg ν).fst = trajMeasure alg ν := by
rw [measure, Measure.fst_prod]

@[simp]
lemma snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
(measure alg ν).snd = streamMeasure ν := by
rw [measure, Measure.snd_prod]

end Bandit

section StreamMeasure

lemma _root_.hasLaw_eval_infinitePi {ι : Type*} {X : ι → Type*} {mX : ∀ i, MeasurableSpace (X i)}
Expand All @@ -90,15 +43,15 @@ lemma _root_.hasLaw_eval_infinitePi {ι : Type*} {X : ι → Type*} {mX : ∀ i,
map_eq := by exact (measurePreserving_eval_infinitePi μ i).map_eq

lemma hasLaw_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
HasLaw (fun h : ℕ → α → R ↦ h n) (Measure.infinitePi ν) (Bandit.streamMeasure ν) :=
HasLaw (fun h : ℕ → α → R ↦ h n) (Measure.infinitePi ν) (streamMeasure ν) :=
hasLaw_eval_infinitePi (fun _ ↦ Measure.infinitePi ν) n

lemma hasLaw_eval_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) :
HasLaw (fun h : ℕ → α → R ↦ h n a) (ν a) (Bandit.streamMeasure ν) :=
HasLaw (fun h : ℕ → α → R ↦ h n a) (ν a) (streamMeasure ν) :=
(hasLaw_eval_infinitePi ν a).comp (hasLaw_eval_streamMeasure ν n)

lemma identDistrib_eval_eval_id_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) (a : α) :
IdentDistrib (fun h : ℕ → α → R ↦ h n a) id (Bandit.streamMeasure ν) (ν a) where
IdentDistrib (fun h : ℕ → α → R ↦ h n a) id (streamMeasure ν) (ν a) where
aemeasurable_fst := Measurable.aemeasurable (by fun_prop)
aemeasurable_snd := Measurable.aemeasurable (by fun_prop)
map_eq := by
Expand All @@ -118,47 +71,40 @@ lemma Integrable.congr_identDistrib {Ω Ω' : Type*}

lemma integrable_eval_streamMeasure (ν : Kernel α ℝ) [IsMarkovKernel ν] (n : ℕ) (a : α)
(h_int : Integrable id (ν a)) :
Integrable (fun h : ℕ → α → ℝ ↦ h n a) (Bandit.streamMeasure ν) :=
Integrable (fun h : ℕ → α → ℝ ↦ h n a) (streamMeasure ν) :=
Integrable.congr_identDistrib h_int (identDistrib_eval_eval_id_streamMeasure ν n a).symm

lemma integral_eval_streamMeasure (ν : Kernel α ℝ) [IsMarkovKernel ν] (n : ℕ) (a : α) :
∫ h, h n a ∂(Bandit.streamMeasure ν) = (ν a)[id] := by
calc ∫ h, h n a ∂(Bandit.streamMeasure ν)
_ = ∫ x, x ∂((Bandit.streamMeasure ν).map (fun h ↦ h n a)) := by
∫ h, h n a ∂(streamMeasure ν) = (ν a)[id] := by
calc ∫ h, h n a ∂(streamMeasure ν)
_ = ∫ x, x ∂((streamMeasure ν).map (fun h ↦ h n a)) := by
rw [integral_map (Measurable.aemeasurable (by fun_prop)) (by fun_prop)]
_ = (ν a)[id] := by simp [(hasLaw_eval_eval_streamMeasure ν n a).map_eq]

lemma iIndepFun_eval_streamMeasure' (ν : Kernel α R) [IsMarkovKernel ν] :
iIndepFun (fun n ω ↦ ω n) (Bandit.streamMeasure ν) :=
iIndepFun (fun n ω ↦ ω n) (streamMeasure ν) :=
iIndepFun_infinitePi (P := fun (_ : ℕ) ↦ Measure.infinitePi ν) (Ω := fun _ ↦ α → R)
(X := fun i u ↦ u) (fun i ↦ by fun_prop)

lemma iIndepFun_eval_streamMeasure'' (ν : Kernel α R) [IsMarkovKernel ν] (a : α) :
iIndepFun (fun n ω ↦ ω n a) (Bandit.streamMeasure ν) :=
iIndepFun (fun n ω ↦ ω n a) (streamMeasure ν) :=
(iIndepFun_eval_streamMeasure' ν).comp (g := fun i ω ↦ ω a) (by fun_prop)

lemma iIndepFun_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] :
iIndepFun (fun (p : ℕ × α) ω ↦ ω p.1 p.2) (Bandit.streamMeasure ν) :=
iIndepFun (fun (p : ℕ × α) ω ↦ ω p.1 p.2) (streamMeasure ν) :=
iIndepFun_uncurry_infinitePi' (X := fun _ _ ↦ id) (fun _ ↦ ν) (by fun_prop)

lemma indepFun_eval_streamMeasure (ν : Kernel α R) [IsMarkovKernel ν] {n m : ℕ} {a b : α}
(h : n ≠ m ∨ a ≠ b) :
IndepFun (fun ω ↦ ω n a) (fun ω ↦ ω m b) (Bandit.streamMeasure ν) := by
IndepFun (fun ω ↦ ω n a) (fun ω ↦ ω m b) (streamMeasure ν) := by
change IndepFun (fun ω ↦ ω (n, a).1 (n, a).2) (fun ω ↦ ω (m, b).1 (m, b).2)
(Bandit.streamMeasure ν)
(streamMeasure ν)
exact (iIndepFun_eval_streamMeasure ν).indepFun (by grind)

lemma indepFun_eval_streamMeasure' (ν : Kernel α R) [IsMarkovKernel ν] {a b : α} (h : a ≠ b) :
IndepFun (fun ω n ↦ ω n a) (fun ω n ↦ ω n b) (Bandit.streamMeasure ν) :=
IndepFun (fun ω n ↦ ω n a) (fun ω n ↦ ω n b) (streamMeasure ν) :=
indepFun_proj_infinitePi_infinitePi h

lemma indepFun_eval_snd_measure (alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν]
{a b : α} (h : a ≠ b) :
IndepFun (fun ω n ↦ ω.2 n a) (fun ω n ↦ ω.2 n b) (Bandit.measure alg ν) := by
refine indepFun_snd_prod ?_ ?_ (indepFun_eval_streamMeasure' ν h) (Bandit.trajMeasure alg ν)
· exact Measurable.aemeasurable (by fun_prop)
· exact Measurable.aemeasurable (by fun_prop)

end StreamMeasure

namespace ArrayModel
Expand All @@ -181,7 +127,7 @@ instance {α R : Type*} [Countable α] [MeasurableSpace R] [StandardBorelSpace R
/-- Probability measure for the array model of stochastic bandits. -/
noncomputable
def arrayMeasure (ν : Kernel α R) : Measure (probSpace α R) :=
(Measure.infinitePi fun _ ↦ volume).prod (Bandit.streamMeasure ν)
(Measure.infinitePi fun _ ↦ volume).prod (streamMeasure ν)

instance (ν : Kernel α R) [IsMarkovKernel ν] : IsProbabilityMeasure (arrayMeasure ν) :=
Measure.prod.instIsProbabilityMeasure _ _
Expand Down Expand Up @@ -606,7 +552,7 @@ lemma map_snd_apply_arrayMeasure {ν : Kernel α R} [IsMarkovKernel ν] (n : ℕ
rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = ν a := by
rw [arrayMeasure, Measure.snd_prod, Bandit.streamMeasure]
rw [arrayMeasure, Measure.snd_prod, streamMeasure]
have : (fun ω ↦ ω n a) = (fun h : α → R ↦ h a) ∘ (fun ω : ℕ → α → R ↦ ω n) := rfl
rw [this, ← Measure.map_map (by fun_prop) (by fun_prop), Measure.infinitePi_map_eval,
Measure.infinitePi_map_eval]
Expand All @@ -629,7 +575,7 @@ omit [DecidableEq α] [Nonempty α] [StandardBorelSpace α] in
lemma indepFun_fst_add_one_aux (ν : Kernel α R) [IsMarkovKernel ν] (n : ℕ) :
(fun ω ↦ ω.1 (n + 1)) ⟂ᵢ[arrayMeasure ν] (fun ω ↦ (fun (i : Iic n) ↦ ω.1 i, ω.2)) := by
let μ₁ : Measure (ℕ → I) := Measure.infinitePi fun _ ↦ volume
let μ₂ : Measure (ℕ → α → R) := Bandit.streamMeasure ν
let μ₂ : Measure (ℕ → α → R) := streamMeasure ν
-- Coordinates of μ₁ are independent
have h_indep : iIndepFun (fun i (ω : ℕ → I) ↦ ω i) μ₁ :=
iIndepFun_infinitePi (fun _ ↦ measurable_id)
Expand Down Expand Up @@ -1016,10 +962,10 @@ lemma hasCondDistrib_action' (alg : Algorithm α R) (ν : Kernel α R) [IsMarkov
rw [h_indep']
congr
simp only [arrayMeasure]
calc ((Measure.infinitePi fun x ↦ ℙ).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.1 (n + 1))
calc ((Measure.infinitePi fun x ↦ ℙ).prod (streamMeasure ν)).map (fun ω ↦ ω.1 (n + 1))
_ = (Measure.infinitePi fun x ↦ ℙ).map (Function.eval (n + 1)) := by
nth_rw 2 [← Measure.fst_prod (μ := Measure.infinitePi fun x ↦ ℙ)
(ν := Bandit.streamMeasure ν)]
(ν := streamMeasure ν)]
rw [Measure.fst, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = ℙ := by rw [Measure.infinitePi_map_eval]
Expand Down
102 changes: 5 additions & 97 deletions LeanBandits/Bandit/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ variable {α Ω : Type*} {mα : MeasurableSpace α} {mΩ : MeasurableSpace Ω} [
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν]
{h_inter : IsAlgEnvSeq A R alg (stationaryEnv ν) P}

local notation "𝔓'" => P.prod (Bandit.streamMeasure ν)
local notation "𝔓'" => P.prod (streamMeasure ν)

omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in
lemma hasLaw_Z (a : α) (m : ℕ) :
Expand All @@ -30,10 +30,10 @@ lemma hasLaw_Z (a : α) (m : ℕ) :
_ = ((𝔓').snd).map (fun ω ↦ ω m a) := by
rw [Measure.snd, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp
_ = (streamMeasure ν).map (fun ω ↦ ω m a) := by simp
_ = ((Measure.infinitePi fun _ ↦ Measure.infinitePi ν).map (fun ω ↦ ω m)).map
(fun ω ↦ ω a) := by
rw [Bandit.streamMeasure, Measure.map_map (by fun_prop) (by fun_prop)]
rw [streamMeasure, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = ν a := by simp_rw [(measurePreserving_eval_infinitePi _ _).map_eq]

Expand All @@ -44,9 +44,6 @@ notation "𝓛[" Y " | " X " in " s "; " μ "]" => Measure.map Y (μ[|X ⁻¹' s
/-- Law of `Y` conditioned on the event that `X` equals `x`. -/
notation "𝓛[" Y " | " X " ← " x "; " μ "]" => Measure.map Y (μ[|X ⁻¹' {x}])

local notation "𝔓t" => Bandit.trajMeasure alg ν
local notation "𝔓" => Bandit.measure alg ν

omit [DecidableEq α] in
lemma condDistrib_reward'' [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (n : ℕ) :
Expand Down Expand Up @@ -107,7 +104,7 @@ lemma condIndepFun_reward_stepsUntil_action [StandardBorelSpace Ω] [Countable
(fun ω ↦ R n ω.1) ({ω | stepsUntil A a m ω.1 = ↑n}.indicator (fun _ ↦ 1)) 𝔓' := by
have hA := h.measurable_A
have hR := h.measurable_R
exact condIndepFun_fst_prod (ν := Bandit.streamMeasure ν)
exact condIndepFun_fst_prod (ν := streamMeasure ν)
(measurable_indicator_stepsUntil_eq hA hR a m n) (by fun_prop) (by fun_prop)
(condIndepFun_reward_stepsUntil_action' h a m n)

Expand Down Expand Up @@ -229,97 +226,8 @@ lemma identDistrib_rewardByCount_id [StandardBorelSpace Ω] [Countable α]

lemma identDistrib_rewardByCount_eval [StandardBorelSpace Ω] [Countable α]
(h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) (n m : ℕ) (hn : n ≠ 0) :
IdentDistrib (rewardByCount A R a n) (fun ω ↦ ω m a) 𝔓' (Bandit.streamMeasure ν) :=
IdentDistrib (rewardByCount A R a n) (fun ω ↦ ω m a) 𝔓' (streamMeasure ν) :=
(identDistrib_rewardByCount_id h a n hn).trans
(identDistrib_eval_eval_id_streamMeasure ν m a).symm

-- lemma indepFun_rewardByCount_Iic [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α)
-- (n : ℕ) :
-- (rewardByCount A R a (n + 1)) ⟂ᵢ[𝔓'] fun ω (i : Iic n) ↦ rewardByCount A R a i ω := by
-- sorry

-- lemma iIndepFun_rewardByCount' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) :
-- iIndepFun (rewardByCount A R a) 𝔓' := by
-- have hA := h.measurable_A
-- have hR := h.measurable_R
-- rw [iIndepFun_nat_iff_forall_indepFun (by fun_prop)]
-- exact indepFun_rewardByCount_Iic h a

-- lemma iIndepFun_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) :
-- iIndepFun (fun (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1)) 𝔓' := by
-- sorry

-- lemma identDistrib_rewardByCount_stream_all [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) :
-- IdentDistrib (fun ω (p : α × ℕ) ↦ rewardByCount A R p.1 (p.2 + 1) ω)
-- (fun ω p ↦ ω p.2 p.1) 𝔓' (Bandit.streamMeasure ν) := by
-- refine IdentDistrib.pi (fun p ↦ ?_) ?_ ?_
-- · refine identDistrib_rewardByCount_eval h p.1 (p.2 + 1) p.2 (by simp) (ν := ν)
-- · exact iIndepFun_rewardByCount h
-- · sorry

-- lemma identDistrib_rewardByCount_stream' [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) :
-- IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω n a)
-- 𝔓' (Bandit.streamMeasure ν) := by
-- refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_
-- · refine identDistrib_rewardByCount_eval h a (n + 1) n (by simp) (ν := ν)
-- · have h_indep := iIndepFun_rewardByCount' h a
-- exact iIndepFun.precomp (g := fun n ↦ n + 1) (fun i j hij ↦ by grind) h_indep
-- · exact iIndepFun_eval_streamMeasure'' ν a

omit [DecidableEq α] [StandardBorelSpace α] [Nonempty α] in
lemma identDistrib_eval_streamMeasure_measure (a : α) :
IdentDistrib (fun ω n ↦ ω n a) (fun ω n ↦ ω.2 n a)
(Bandit.streamMeasure ν) 𝔓 := by
refine IdentDistrib.pi (fun n ↦ ?_) ?_ ?_
· rw [← Bandit.snd_measure alg ν, Measure.snd,
identDistrib_map_left_iff (by fun_prop) (by fun_prop)
(Measurable.aemeasurable <| by fun_prop)]
exact IdentDistrib.refl (by fun_prop)
· exact iIndepFun_eval_streamMeasure'' ν a
· change iIndepFun (fun n ↦ ((fun ω ↦ ω n a) ∘ Prod.snd)) 𝔓
rw [← iIndepFun_map_iff (by fun_prop) (fun _ ↦ Measurable.aemeasurable (by fun_prop))]
rw [← Measure.snd, Bandit.snd_measure]
exact iIndepFun_eval_streamMeasure'' ν a

-- lemma identDistrib_rewardByCount_stream [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (a : α) :
-- IdentDistrib (fun ω n ↦ rewardByCount A R a (n + 1) ω) (fun ω n ↦ ω.2 n a) 𝔓' 𝔓 :=
-- (identDistrib_rewardByCount_stream' h a).trans (identDistrib_eval_streamMeasure_measure a)

-- lemma indepFun_rewardByCount_of_ne [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) {a b : α} (hab : a ≠ b) :
-- IndepFun (fun ω s ↦ rewardByCount A R a s ω) (fun ω s ↦ rewardByCount A R b s ω) 𝔓' := by
-- sorry

-- lemma identDistrib_sum_Icc_rewardByCount [StandardBorelSpace Ω] [Nonempty Ω] [Countable α]
-- (h : IsAlgEnvSeq A R alg (stationaryEnv ν) P) (m : ℕ) (a : α) :
-- IdentDistrib (fun ω ↦ ∑ s ∈ Icc 1 m, rewardByCount A R a s ω)
-- (fun ω ↦ ∑ s ∈ range m, ω.2 s a) 𝔓' 𝔓 := by
-- have h1 (a : α) :
-- IdentDistrib (fun ω s ↦ rewardByCount A R a (s + 1) ω) (fun ω s ↦ ω.2 s a) 𝔓' 𝔓 :=
-- identDistrib_rewardByCount_stream h a
-- have h_eq (ω : Ω × (ℕ → α → ℝ)) : ∑ s ∈ Icc 1 m, rewardByCount A R a s ω
-- = ∑ s ∈ range m, rewardByCount A R a (s + 1) ω := by
-- let e : Icc 1 m ≃ range m :=
-- { toFun x := ⟨x - 1, by have h := x.2; simp only [mem_Icc] at h; simp; grind⟩
-- invFun x := ⟨x + 1, by
-- have h := x.2
-- simp only [mem_Icc, le_add_iff_nonneg_left, zero_le, true_and, ge_iff_le]
-- simp only [mem_range] at h
-- grind⟩
-- left_inv x := by have h := x.2; simp only [mem_Icc] at h; grind
-- right_inv x := by have h := x.2; grind }
-- rw [← sum_coe_sort (Icc 1 m), ← sum_coe_sort (range m), sum_equiv e]
-- · simp
-- · simp only [univ_eq_attach, mem_attach, forall_const, Subtype.forall, mem_Icc,
-- forall_and_index]
-- grind
-- simp_rw [h_eq]
-- exact IdentDistrib.comp (h1 a) (u := fun p ↦ ∑ s ∈ range m, p s) (by fun_prop)

end Bandits
Loading
Loading