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
1 change: 1 addition & 0 deletions LeanBandits.lean
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import LeanBandits.AlgorithmBuilding
import LeanBandits.Bandit
import LeanBandits.ETC
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.Regret
import LeanBandits.RewardByCountMeasure
import LeanBandits.UCB
15 changes: 14 additions & 1 deletion LeanBandits/Bandit.lean
Original file line number Diff line number Diff line change
Expand Up @@ -130,7 +130,7 @@ lemma measurable_reward_prod : Measurable (fun p : ℕ × (ℕ → α × R) ↦
lemma measurable_hist (n : ℕ) : Measurable (hist n (α := α) (R := R)) := by unfold hist; fun_prop

/-- Filtration of the bandit process. -/
def ℱ (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] :
protected def filtration (α R : Type*) [MeasurableSpace α] [MeasurableSpace R] :
Filtration ℕ (inferInstance : MeasurableSpace (ℕ → α × R)) :=
MeasureTheory.Filtration.piLE (X := fun _ ↦ α × R)

Expand All @@ -152,6 +152,19 @@ lemma condDistrib_arm [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace
=ᵐ[(Bandit.trajMeasure alg ν).map (hist n)] alg.policy n := by
sorry

lemma hasLaw_arm_zero [StandardBorelSpace α] [Nonempty α] [StandardBorelSpace R] [Nonempty R]
(alg : Algorithm α R) (ν : Kernel α R) [IsMarkovKernel ν] :
HasLaw (arm 0) alg.p0 (Bandit.trajMeasure alg ν) where
map_eq := by
sorry

/-- The reward at time `n+1` is independent of the history up to time `n` given the arm at `n+1`. -/
lemma condIndepFun_reward_hist_arm [StandardBorelSpace α] [StandardBorelSpace R]
{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 ν) := by
sorry

end MeasureSpace

end Bandits
14 changes: 6 additions & 8 deletions LeanBandits/ETC.lean
Original file line number Diff line number Diff line change
Expand Up @@ -43,14 +43,12 @@ def etcAlgorithm (hK : 0 < K) (m : ℕ) : Algorithm (Fin K) ℝ where
p0 := Measure.dirac ⟨0, hK⟩

lemma ETC.arm_zero (hK : 0 < K) (m : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν] :
arm 0 =ᵐ[Bandit.trajMeasure (etcAlgorithm hK m) ν] fun h ↦ ⟨0, hK⟩ := by
suffices h : (Bandit.trajMeasure (etcAlgorithm hK m) ν).map (arm 0) = (etcAlgorithm hK m).p0 by
have h_eq : ∀ᵐ x ∂((Bandit.trajMeasure (etcAlgorithm hK m) ν).map (arm 0)), x = ⟨0, hK⟩ := by
rw [h]
simp [etcAlgorithm]
exact ae_of_ae_map (by fun_prop) h_eq
-- extract lemma
sorry
arm 0 =ᵐ[Bandit.trajMeasure (etcAlgorithm hK m) ν] fun _ ↦ ⟨0, hK⟩ := by
have h_eq : ∀ᵐ x ∂((Bandit.trajMeasure (etcAlgorithm hK m) ν).map (arm 0)), x = ⟨0, hK⟩ := by
have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
rw [(hasLaw_arm_zero _ _).map_eq]
simp [etcAlgorithm]
exact ae_of_ae_map (by fun_prop) h_eq

lemma ETC.arm_ae_eq_etcNextArm (hK : 0 < K) (m : ℕ) (ν : Kernel (Fin K) ℝ) [IsMarkovKernel ν]
(n : ℕ) :
Expand Down
457 changes: 457 additions & 0 deletions LeanBandits/ForMathlib/CondDistrib.lean

Large diffs are not rendered by default.

29 changes: 29 additions & 0 deletions LeanBandits/Regret.lean
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,27 @@ lemma arm_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount (arm · h) a (s
rwa [← pullCount_eq_pullCount]
exact h_ne

lemma arm_eq_of_stepsUntil_eq_coe {ω : ℕ → α × ℝ} (hm : m ≠ 0)
(h : stepsUntil (arm · ω) a m = n) :
arm n ω = a := by
have : n = (stepsUntil (fun x ↦ arm x ω) a m).toNat := by simp [h]
rw [this, arm_stepsUntil hm]
by_contra! h_contra
rw [← stepsUntil_eq_top_iff] at h_contra
simp [h_contra] at h

lemma stepsUntil_eq_congr {k' : ℕ → α} (h : ∀ i ≤ n, k i = k' i) :
stepsUntil k a m = n ↔ stepsUntil k' a m = n := by
sorry

lemma pullCount_stepsUntil_add_one (h_exists : ∃ s, pullCount k a (s + 1) = m) :
pullCount k a (stepsUntil k a m + 1).toNat = m := by
sorry

lemma pullCount_stepsUntil (hm : m ≠ 0) (h_exists : ∃ s, pullCount k a (s + 1) = m) :
pullCount k a (stepsUntil k a m).toNat = m - 1 := by
sorry

/-- Reward obtained when pulling arm `a` for the `m`-th time. -/
noncomputable
def rewardByCount (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) : ℝ :=
Expand All @@ -144,6 +165,14 @@ lemma rewardByCount_eq_ite (a : α) (m : ℕ) (h : ℕ → α × ℝ) (z : ℕ
unfold rewardByCount
cases stepsUntil (arm · h) a m <;> simp

lemma rewardByCount_of_stepsUntil_eq_top {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)}
(h : stepsUntil (arm · ω.1) a m = ⊤) :
rewardByCount a m ω.1 ω.2 = ω.2 m a := by simp [rewardByCount_eq_ite, h]

lemma rewardByCount_of_stepsUntil_eq_coe {ω : (ℕ → α × ℝ) × (ℕ → α → ℝ)}
(h : stepsUntil (arm · ω.1) a m = n) :
rewardByCount a m ω.1 ω.2 = reward n ω.1 := by simp [rewardByCount_eq_ite, h]

lemma rewardByCount_pullCount_add_one_eq_reward (t : ℕ) (h : ℕ → α × ℝ) (z : ℕ → α → ℝ) :
rewardByCount (arm t h) (pullCount (arm · h) (arm t h) t + 1) h z = reward t h := by
rw [rewardByCount, ← pullCount_eq_pullCount_add_one, stepsUntil_pullCount_eq]
Expand Down
249 changes: 190 additions & 59 deletions LeanBandits/RewardByCountMeasure.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ Released under Apache 2.0 license as described in the file LICENSE.
Authors: Rémy Degenne
-/
import LeanBandits.Bandit
import LeanBandits.ForMathlib.CondDistrib
import LeanBandits.Regret

/-! # Laws of `stepsUntil` and `rewardByCount`
Expand All @@ -12,58 +13,6 @@ import LeanBandits.Regret
open MeasureTheory ProbabilityTheory Finset
open scoped ENNReal NNReal

section Aux

variable {α β γ Ω Ω' : Type*} [MeasurableSpace Ω] [StandardBorelSpace Ω] [Nonempty Ω]
{mα : MeasurableSpace α} {μ : Measure α} {mβ : MeasurableSpace β} {mγ : MeasurableSpace γ}
[MeasurableSpace Ω'] [StandardBorelSpace Ω'] [Nonempty Ω']
{X : α → β} {Y : α → Ω} {Z : α → Ω'}

lemma MeasureTheory.Measure.comp_congr {κ η : Kernel α β} (h : ∀ᵐ a ∂μ, κ a = η a) :
κ ∘ₘ μ = η ∘ₘ μ :=
Measure.bind_congr_right h

lemma MeasureTheory.Measure.copy_comp_map (hX : AEMeasurable X μ) :
Kernel.copy β ∘ₘ (μ.map X) = μ.map (fun a ↦ (X a, X a)) := by
rw [Kernel.copy, deterministic_comp_eq_map, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX]
congr

lemma MeasureTheory.Measure.compProd_deterministic [SFinite μ] (hX : Measurable X) :
μ ⊗ₘ (Kernel.deterministic X hX) = μ.map (fun a ↦ (a, X a)) := by
rw [Measure.compProd_eq_comp_prod, Kernel.id, Kernel.deterministic_prod_deterministic,
Measure.deterministic_comp_eq_map]
rfl

lemma ProbabilityTheory.condDistrib_comp_map [IsFiniteMeasure μ]
(hX : AEMeasurable X μ) (hY : AEMeasurable Y μ) :
condDistrib Y X μ ∘ₘ (μ.map X) = μ.map Y := by
rw [← Measure.snd_compProd, compProd_map_condDistrib hY, Measure.snd_map_prodMk₀ hX]

lemma ProbabilityTheory.condDistrib_comp [IsFiniteMeasure μ]
(hX : AEMeasurable X μ) {f : β → Ω} (hf : Measurable f) :
condDistrib (f ∘ X) X μ =ᵐ[μ.map X] Kernel.deterministic f hf := by
rw [← Kernel.compProd_eq_iff, compProd_map_condDistrib (by fun_prop),
Measure.compProd_deterministic, AEMeasurable.map_map_of_aemeasurable (by fun_prop) hX]
congr

lemma ProbabilityTheory.condDistrib_const [IsFiniteMeasure μ]
(hX : AEMeasurable X μ) (c : Ω) :
condDistrib (fun _ ↦ c) X μ =ᵐ[μ.map X] Kernel.deterministic (fun _ ↦ c) (by fun_prop) := by
have : (fun _ : α ↦ c) = (fun _ : β ↦ c) ∘ X := rfl
conv_lhs => rw [this]
filter_upwards [condDistrib_comp hX (by fun_prop : Measurable (fun _ ↦ c))] with b hb
rw [hb]

@[fun_prop]
lemma Measurable.coe_nat_enat {f : α → ℕ} (hf : Measurable f) :
Measurable (fun a ↦ (f a : ℕ∞)) := Measurable.comp (by fun_prop) hf

@[fun_prop]
lemma Measurable.toNat {f : α → ℕ∞} (hf : Measurable f) : Measurable (fun a ↦ (f a).toNat) :=
Measurable.comp (by fun_prop) hf

end Aux

namespace Bandits

variable {α : Type*} {mα : MeasurableSpace α} [DecidableEq α] [MeasurableSingletonClass α]
Expand Down Expand Up @@ -126,16 +75,199 @@ lemma measurable_rewardByCount (a : α) (m : ℕ) :
(measurable_stepsUntil' a m).toNat.prodMk (by fun_prop)
exact Measurable.comp (by fun_prop) this

lemma condDistrib_rewardByCount_stepsUntil [StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0) :
variable {alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν]

omit [DecidableEq α] [MeasurableSingletonClass α] in
lemma hasLaw_Z (a : α) (m : ℕ) :
HasLaw (fun ω ↦ ω.2 m a) (ν a) (Bandit.measure alg ν) where
map_eq := by
calc ((Bandit.trajMeasure alg ν).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.2 m a)
_ = (((Bandit.trajMeasure alg ν).prod (Bandit.streamMeasure ν)).map (fun ω ↦ ω.2)).map
(fun ω ↦ ω m a) := by
rw [Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = (Bandit.streamMeasure ν).map (fun ω ↦ ω m a) := by simp [Measure.map_snd_prod]
_ = ((Measure.infinitePi fun _ ↦ Measure.infinitePi ν).map (fun ω ↦ ω m)).map
(fun ω ↦ ω a) := by
rw [Bandit.streamMeasure, Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = ν a := by simp_rw [(measurePreserving_eval_infinitePi _ _).map_eq]

/-- Law of `Y` conditioned on the event `s`.-/
notation "𝓛[" Y " | " s "; " μ "]" => Measure.map Y (μ[|s])
/-- Law of `Y` conditioned on the event that `X` is in `s`. -/
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}])
/-- Law of `Y` conditioned on `X`. -/
notation "𝓛[" Y " | " X "; " μ "]" => condDistrib Y X μ

omit [DecidableEq α] [MeasurableSingletonClass α] in
lemma condDistrib_reward' (n : ℕ) :
𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; Bandit.measure alg ν]
=ᵐ[(Bandit.measure alg ν).map (fun ω ↦ arm n ω.1)] ν := by
let μ := Bandit.measure alg ν
have h_ra' : 𝓛[reward n | arm n; Bandit.trajMeasure alg ν]
=ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] ν := condDistrib_reward alg ν n
have h_law : μ.map (fun ω ↦ arm n ω.1) = (Bandit.trajMeasure alg ν).map (arm n) := by
calc μ.map (fun ω ↦ arm n ω.1)
_ = (μ.map (fun ω ↦ ω.1)).map (fun ω ↦ arm n ω) := by
rw [Measure.map_map (by fun_prop) (by fun_prop)]
rfl
_ = _ := by unfold μ Bandit.measure; simp [Measure.map_fst_prod]
rw [h_law]
have h_prod : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; μ]
=ᵐ[(Bandit.trajMeasure alg ν).map (arm n)] 𝓛[reward n | arm n; Bandit.trajMeasure alg ν] :=
condDistrib_fst_prod (by fun_prop) (by fun_prop) _
filter_upwards [h_ra', h_prod] with ω h_eq h_prod
rw [h_prod, h_eq]

omit [DecidableEq α] in
lemma reward_cond_arm [Countable α] (a : α) (n : ℕ)
(hμa : (Bandit.measure alg ν).map (fun ω ↦ arm n ω.1) {a} ≠ 0) :
𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; Bandit.measure alg ν] = ν a := by
let μ := Bandit.measure alg ν
have h_ra : 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1; μ] =ᵐ[μ.map (fun ω ↦ arm n ω.1)] ν :=
condDistrib_reward' n
have h_eq := condDistrib_ae_eq_cond (μ := μ)
(X := fun ω ↦ arm n ω.1) (Y := fun ω ↦ reward n ω.1) (by fun_prop) (by fun_prop)
rw [Filter.EventuallyEq, ae_iff_of_countable] at h_ra h_eq
specialize h_ra a hμa
specialize h_eq a hμa
rw [h_ra] at h_eq
exact h_eq.symm

lemma condIndepFun_reward_stepsUntil_arm [StandardBorelSpace α] [Countable α] [Nonempty α]
(a : α) (m n : ℕ) (hm : m ≠ 0) :
CondIndepFun (mα.comap (fun ω ↦ arm n ω.1)) ((measurable_arm n).comp measurable_fst).comap_le
(fun ω ↦ reward n ω.1) ({ω | stepsUntil (arm · ω.1) a m = ↑n}.indicator (fun _ ↦ 1))
(Bandit.measure alg ν) := by
-- first restrict to the `trajMeasure` side
suffices h_indep :
CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le
(reward n) ({ω | stepsUntil (arm · ω) a m = ↑n}.indicator (fun _ ↦ 1))
(Bandit.trajMeasure alg ν) by
sorry
-- Now prove the independence : the indicator of `stepsUntil ... = n` is a function of
-- `hist (n-1)` and `arm n`.
-- It thus suffices to prove the independence of `reward n` and `hist (n-1)` conditionally
-- on `arm n`.
have hn : n ≠ 0 := by
sorry -- assume it?
have h_indep : CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le (reward n)
(hist (n - 1)) (Bandit.trajMeasure alg ν) := by
convert condIndepFun_reward_hist_arm (alg := alg) (ν := ν) (n - 1)
<;> rw [Nat.sub_add_cancel (by grind)]
have h_indep' : CondIndepFun (mα.comap (arm n)) (measurable_arm n).comap_le (reward n)
(fun ω ↦ (hist (n - 1) ω, arm n ω)) (Bandit.trajMeasure alg ν) :=
h_indep.prod_right (by fun_prop) (by fun_prop) (by fun_prop)
suffices ∃ φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ, Measurable φ ∧
({ω : ℕ → α × ℝ | stepsUntil (arm · ω) a m = ↑n}.indicator (fun _ ↦ 1))
= φ ∘ (fun ω : ℕ → α × ℝ ↦ (hist (n - 1) ω, arm n ω)) by
obtain ⟨φ, hφ_meas, h_eq⟩ := this
rw [h_eq]
exact h_indep'.comp measurable_id hφ_meas
-- it would follow from measurability wrt the sigma-algebra generated by
-- `hist (n-1)` and `arm n`, but we can also give an explicit function
let k : ((Iic (n - 1) → α × ℝ) × α) → (ℕ → α) := fun x i ↦
if hi : i ∈ Iic (n - 1) then (x.1 ⟨i, hi⟩).1 else if i = n then x.2 else a -- a is arbitrary
let φ : ((Iic (n - 1) → α × ℝ) × α) → ℕ := fun x ↦ if stepsUntil (k x) a m = ↑n then 1 else 0
classical
have hφ_meas : Measurable φ := by
refine Measurable.ite ?_ (by fun_prop) (by fun_prop)
refine (measurableSet_singleton _).preimage ?_
refine (measurable_stepsUntil a m).comp ?_
unfold k
rw [measurable_pi_iff]
intro i
split_ifs <;> fun_prop
refine ⟨φ, hφ_meas, ?_⟩
ext ω
classical
simp only [Set.indicator_apply, Set.mem_setOf_eq, Function.comp_apply, φ]
congr 1
rw [stepsUntil_eq_congr]
intro i hin
simp only [arm, mem_Iic, hist, dite_eq_ite, left_eq_ite_iff, not_le, k]
intro hni
have : i = n := by grind
simp [this]

lemma reward_cond_stepsUntil [StandardBorelSpace α] [Countable α] [Nonempty α] (a : α) (m n : ℕ)
(hm : m ≠ 0)
(hμn : (Bandit.measure alg ν) ((fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n}) ≠ 0) :
𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil (fun x ↦ arm x ω.1) a m ← (n : ℕ∞);
Bandit.measure alg ν] = ν a := by
let μ := Bandit.measure alg ν
have hμna :
μ ((fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}) ≠ 0 := by
suffices ((fun ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) ↦
stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a})
= (fun ω ↦ stepsUntil (arm · ω.1) a m) ⁻¹' {↑n} by simpa [this] using hμn
ext ω
simp only [Set.mem_inter_iff, Set.mem_preimage, Set.mem_singleton_iff, and_iff_left_iff_imp]
exact arm_eq_of_stepsUntil_eq_coe hm
have hμa : μ.map (fun ω ↦ arm n ω.1) {a} ≠ 0 := by
rw [Measure.map_apply (by fun_prop) (measurableSet_singleton _)]
refine fun h_zero ↦ hμn (measure_mono_null (fun ω ↦ ?_) h_zero)
simp only [Set.mem_preimage, Set.mem_singleton_iff]
exact arm_eq_of_stepsUntil_eq_coe hm
calc 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ stepsUntil (arm · ω.1) a m ← (n : ℕ∞); μ]
_ = (μ[|(fun ω ↦ stepsUntil (fun x ↦ arm x ω.1) a m) ⁻¹' {↑n} ∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map
(fun ω ↦ reward n ω.1) := by
congr with ω
simp only [Set.mem_preimage, Set.mem_singleton_iff, Set.mem_inter_iff, iff_self_and]
exact arm_eq_of_stepsUntil_eq_coe hm
_ = (μ[|{ω : (ℕ → α × ℝ) × (ℕ → α → ℝ) | stepsUntil (arm · ω.1) a m = ↑n}.indicator 1 ⁻¹' {1}
∩ (fun ω ↦ arm n ω.1) ⁻¹' {a}]).map (fun ω ↦ reward n ω.1) := by
congr 3 with ω
simp [Set.indicator_apply]
_ = 𝓛[fun ω ↦ reward n ω.1 | fun ω ↦ arm n ω.1 ← a; μ] := by
rw [cond_of_condIndepFun (by fun_prop)]
· exact condIndepFun_reward_stepsUntil_arm a m n hm
· refine measurable_one.indicator ?_
exact measurableSet_eq_fun' (by fun_prop) (by fun_prop)
· fun_prop
· convert hμna
ext ω
simp [Set.indicator_apply]
_ = ν a := reward_cond_arm a n hμa

lemma condDistrib_rewardByCount_stepsUntil [Countable α] [StandardBorelSpace α] [Nonempty α]
(a : α) (m : ℕ) (hm : m ≠ 0) :
condDistrib (fun ω ↦ rewardByCount a m ω.1 ω.2) (fun ω ↦ stepsUntil (arm · ω.1) a m)
(Bandit.measure alg ν)
=ᵐ[(Bandit.measure alg ν).map (fun ω ↦ stepsUntil (arm · ω.1) a m)] Kernel.const _ (ν a) := by
sorry
let μ := Bandit.measure alg ν
refine (condDistrib_ae_eq_cond (μ := μ)
(X := fun ω ↦ stepsUntil (arm · ω.1) a m) (by fun_prop) (by fun_prop)).trans ?_
rw [Filter.EventuallyEq, ae_iff_of_countable]
intro n hn
simp only [Kernel.const_apply]
cases n with
| top =>
rw [Measure.map_congr (g := fun ω ↦ ω.2 m a)]
swap
· refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_
simp only [Set.mem_preimage, Set.mem_singleton_iff]
exact fun ω ↦ rewardByCount_of_stepsUntil_eq_top
rw [cond_of_indepFun _ (by fun_prop) (by fun_prop) (measurableSet_singleton _)]
· exact (hasLaw_Z a m).map_eq
· rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn
· exact indepFun_prod (X := fun ω : ℕ → α × ℝ ↦ stepsUntil (arm · ω) a m)
(Y := fun ω : ℕ → α → ℝ ↦ ω m a) (by fun_prop) (by fun_prop)
| coe n =>
rw [Measure.map_congr (g := fun ω ↦ reward n ω.1)]
swap
· refine ae_cond_of_forall_mem ((measurableSet_singleton _).preimage (by fun_prop)) ?_
simp only [Set.mem_preimage, Set.mem_singleton_iff]
exact fun ω ↦ rewardByCount_of_stepsUntil_eq_coe
refine reward_cond_stepsUntil a m n hm ?_
rwa [Measure.map_apply (by fun_prop) (measurableSet_singleton _)] at hn

/-- The reward received at the `m`-th pull of arm `a` has law `ν a`. -/
lemma hasLaw_rewardByCount [StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (m : ℕ) (hm : m ≠ 0):
lemma hasLaw_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α]
(a : α) (m : ℕ) (hm : m ≠ 0) :
HasLaw (fun ω ↦ rewardByCount a m ω.1 ω.2) (ν a) (Bandit.measure alg ν) where
map_eq := by
have h_condDistrib :
Expand All @@ -157,8 +289,7 @@ lemma hasLaw_rewardByCount [StandardBorelSpace α] [Nonempty α]
isProbabilityMeasure_map (by fun_prop)
simp

lemma identDistrib_rewardByCount [StandardBorelSpace α] [Nonempty α]
{alg : Algorithm α ℝ} {ν : Kernel α ℝ} [IsMarkovKernel ν] (a : α) (n m : ℕ)
lemma identDistrib_rewardByCount [Countable α] [StandardBorelSpace α] [Nonempty α] (a : α) (n m : ℕ)
(hn : n ≠ 0) (hm : m ≠ 0) :
IdentDistrib (fun ω ↦ rewardByCount a n ω.1 ω.2) (fun ω ↦ rewardByCount a m ω.1 ω.2)
(Bandit.measure alg ν) (Bandit.measure alg ν) where
Expand Down
Loading