diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index 945a88b6..97d0683c 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -13,6 +13,7 @@ public import LeanMachineLearning.ForMathlib.Probability.Independence.IndepFun public import LeanMachineLearning.ForMathlib.Probability.Independence.IndepInfinitePi public import LeanMachineLearning.ForMathlib.Probability.Integrable public import LeanMachineLearning.ForMathlib.Probability.Kernel.Basic +public import LeanMachineLearning.ForMathlib.Probability.Kernel.Composition.IntegralCompProd public import LeanMachineLearning.ForMathlib.Probability.Kernel.Composition.MapComap public import LeanMachineLearning.ForMathlib.Probability.Kernel.Composition.MeasureCompProd public import LeanMachineLearning.ForMathlib.Probability.Kernel.IonescuTulcea.Traj @@ -29,6 +30,7 @@ public import LeanMachineLearning.Online.Bandit.BayesRegret public import LeanMachineLearning.Online.Bandit.Regret public import LeanMachineLearning.Online.Bandit.RewardByCountMeasure public import LeanMachineLearning.Online.Bandit.SumRewards +public import LeanMachineLearning.SequentialLearning.ActionIndicator public import LeanMachineLearning.SequentialLearning.Algorithm public import LeanMachineLearning.SequentialLearning.AlgorithmDensity public import LeanMachineLearning.SequentialLearning.AlgorithmDensityBayes @@ -39,9 +41,12 @@ public import LeanMachineLearning.SequentialLearning.Algorithms.Uniform public import LeanMachineLearning.SequentialLearning.BayesStationaryEnv public import LeanMachineLearning.SequentialLearning.Deterministic public import LeanMachineLearning.SequentialLearning.EvaluationEnv +public import LeanMachineLearning.SequentialLearning.FeedbackMartingale public import LeanMachineLearning.SequentialLearning.FiniteActions public import LeanMachineLearning.SequentialLearning.IonescuTulceaSpace +public import LeanMachineLearning.SequentialLearning.Means public import LeanMachineLearning.SequentialLearning.StationaryEnv +public import LeanMachineLearning.SequentialLearning.SumRewards public import LeanMachineLearning.Tactic.EqLift public import LeanMachineLearning.Tactic.EqLift.ForMathlib.Kernel public import LeanMachineLearning.Tactic.EqLift.ForMathlib.MeasurableEquiv diff --git a/LeanMachineLearning/ForMathlib/Probability/HasCondDistrib.lean b/LeanMachineLearning/ForMathlib/Probability/HasCondDistrib.lean index 9d6da919..87a686a6 100644 --- a/LeanMachineLearning/ForMathlib/Probability/HasCondDistrib.lean +++ b/LeanMachineLearning/ForMathlib/Probability/HasCondDistrib.lean @@ -129,6 +129,13 @@ lemma HasLaw.prod_of_hasCondDistrib {P : Measure β} HasLaw (fun ω ↦ (X ω, Y ω)) (P ⊗ₘ κ) μ := ⟨by fun_prop, by rw [h2.map_eq, h1.map_eq]⟩ +lemma HasCondDistrib.hasLaw_comp [SFinite μ] [IsSFiniteKernel κ] (h : HasCondDistrib Y X κ μ) : + HasLaw Y (κ ∘ₘ (μ.map X)) μ := by + refine ⟨by fun_prop, ?_⟩ + rw [← Measure.snd_compProd, ← h.map_eq, Measure.snd, + AEMeasurable.map_map_of_aemeasurable (by fun_prop) (by fun_prop)] + rfl + lemma HasCondDistrib.prod {Z : α → Ω'} {η : Kernel (β × Ω) Ω'} (h1 : HasCondDistrib Y X κ μ) (h2 : HasCondDistrib Z (fun ω ↦ (X ω, Y ω)) η μ) : HasCondDistrib (fun ω ↦ (Y ω, Z ω)) X (κ ⊗ₖ η) μ := by diff --git a/LeanMachineLearning/ForMathlib/Probability/Kernel/Composition/IntegralCompProd.lean b/LeanMachineLearning/ForMathlib/Probability/Kernel/Composition/IntegralCompProd.lean new file mode 100644 index 00000000..86056b63 --- /dev/null +++ b/LeanMachineLearning/ForMathlib/Probability/Kernel/Composition/IntegralCompProd.lean @@ -0,0 +1,68 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import Mathlib.Probability.Kernel.Composition.IntegralCompProd + +import Mathlib.Analysis.Convex.Integral + +/-! +# Lp functions with respect to a composition of kernels and measures +-/ + +@[expose] public section + +open ProbabilityTheory +open scoped ENNReal + +namespace MeasureTheory + +protected lemma Measure.memLp_comp_iff + {α β E : Type*} {mα : MeasurableSpace α} {mβ : MeasurableSpace β} [NormedAddCommGroup E] + {κ : Kernel α β} {μ : Measure α} {f : β → E} {p : ℝ≥0∞} (hp0 : p ≠ 0) (hp_top : p ≠ ∞) + (hf : AEStronglyMeasurable f (κ ∘ₘ μ)) : + MemLp f p (κ ∘ₘ μ) + ↔ (∀ᵐ x ∂μ, MemLp f p (κ x)) ∧ Integrable (fun x ↦ ∫ y, ‖f y‖ ^ p.toReal ∂κ x) μ := by + rw [← integrable_norm_rpow_iff (by fun_prop) hp0 hp_top, Measure.integrable_comp_iff] + swap; · exact (hf.norm.aemeasurable.pow_const p.toReal).aestronglyMeasurable + -- todo extract + unfold AEStronglyMeasurable at hf + obtain ⟨g, hg, hfg⟩ := hf + obtain hfg' := Measure.ae_ae_of_ae_comp hfg + have hf' : ∀ᵐ ω ∂μ, AEStronglyMeasurable f (κ ω) := by + filter_upwards [hfg'] with ω hω using ⟨g, hg, hω⟩ + -- + congr! 1 + · suffices ∀ᵐ x ∂μ, Integrable (fun x ↦ ‖f x‖ ^ p.toReal) (κ x) ↔ MemLp f p (κ x) by + refine ⟨fun h ↦ ?_, fun h ↦ ?_⟩ + <;> filter_upwards [h, this] with x hx h_iff + · rwa [h_iff] at hx + · rwa [← h_iff] at hx + filter_upwards [hf'] with ω hω + rw [integrable_norm_rpow_iff hω hp0 hp_top] + · congr! 4 with y + simp only [Real.norm_eq_abs, abs_eq_self] + positivity + +/-- **Jensen's inequality** for the convex function `x ↦ ‖x‖ ^ p`, `1 ≤ p`. -/ +lemma norm_integral_rpow_le_integral_norm_rpow + {α E : Type*} {mα : MeasurableSpace α} {μ : Measure α} [IsProbabilityMeasure μ] + [NormedAddCommGroup E] [NormedSpace ℝ E] {f : α → E} {p : ℝ≥0∞} + (hp1 : 1 ≤ p) (hp_top : p ≠ ∞) (hf : MemLp f p μ) : + ‖∫ x, f x ∂μ‖ ^ p.toReal ≤ ∫ x, ‖f x‖ ^ p.toReal ∂μ := by + have hp0 : p ≠ 0 := by positivity + have hp1' : 1 ≤ p.toReal := by simpa using ENNReal.toReal_mono hp_top hp1 + calc ‖∫ x, f x ∂μ‖ ^ p.toReal + _ ≤ (∫ x, ‖f x‖ ∂μ) ^ p.toReal := by + gcongr + exact norm_integral_le_integral_norm _ + _ ≤ ∫ x, ‖f x‖ ^ p.toReal ∂μ := + ConvexOn.map_integral_le (convexOn_rpow hp1') + (Real.continuous_rpow_const (by positivity)).continuousOn isClosed_Ici + (ae_of_all _ fun x ↦ norm_nonneg _) (hf.integrable hp1).norm + ((integrable_norm_rpow_iff hf.1 hp0 hp_top).mpr hf) + +end MeasureTheory diff --git a/LeanMachineLearning/Online/Bandit/ArrayProbSpace.lean b/LeanMachineLearning/Online/Bandit/ArrayProbSpace.lean index a3ff3bf9..291fdbae 100644 --- a/LeanMachineLearning/Online/Bandit/ArrayProbSpace.lean +++ b/LeanMachineLearning/Online/Bandit/ArrayProbSpace.lean @@ -9,7 +9,7 @@ public import LeanMachineLearning.ForMathlib.Probability.Independence.CondIndepF public import LeanMachineLearning.ForMathlib.Probability.Independence.IndepFun public import LeanMachineLearning.ForMathlib.Probability.Independence.IndepInfinitePi public import LeanMachineLearning.ForMathlib.Probability.Integrable -public import LeanMachineLearning.SequentialLearning.FiniteActions +public import LeanMachineLearning.SequentialLearning.SumRewards public import LeanMachineLearning.SequentialLearning.StationaryEnv public import Mathlib.Probability.Independence.Integration public import Mathlib.Probability.Kernel.Representation diff --git a/LeanMachineLearning/SequentialLearning/ActionIndicator.lean b/LeanMachineLearning/SequentialLearning/ActionIndicator.lean new file mode 100644 index 00000000..6c17a2d6 --- /dev/null +++ b/LeanMachineLearning/SequentialLearning/ActionIndicator.lean @@ -0,0 +1,110 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import LeanMachineLearning.SequentialLearning.SumRewards + +/-! +# The action indicator + +`actionIndicator A k n ω = 𝟙{A n ω = k}` is the `{0,1}`-valued indicator that action `k` was chosen +at round `n`. It is the increment weight of every per-action sum attached to an action process: +`pullCount A k n` is its partial sum (`sum_range_actionIndicator_eq_pullCount`) and +`sumRewards A Y k n` is its reward-weighted partial sum (`sum_actionIndicator_mul`). + + +## Main definitions + +* `Learning.actionIndicator` + +## Main results + +* `Learning.sum_range_actionIndicator_eq_pullCount`, `Learning.sum_actionIndicator_mul` — the two + partial-sum identities. +* `Learning.adapted_actionIndicator`, `Learning.integrable_actionIndicator`. +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Filter Finset + +namespace Learning + +variable {Ω 𝓐 𝓨 : Type*} {mΩ : MeasurableSpace Ω} {m𝓐 : MeasurableSpace 𝓐} {m𝓨 : MeasurableSpace 𝓨} + [MeasurableSingletonClass 𝓐] {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} {P : Measure Ω} + +/-- The `{0,1}`-valued assignment indicator of action `k`: +`actionIndicator A k n ω = 𝟙{A n ω = k}`. -/ +noncomputable def actionIndicator (A : ℕ → Ω → 𝓐) (k : 𝓐) (n : ℕ) (ω : Ω) : ℝ := + {ω | A n ω = k}.indicator (fun _ ↦ (1 : ℝ)) ω + +/-- `actionIndicator A k n ω = 1` exactly when action `k` is chosen at time `n`. -/ +lemma actionIndicator_eq_one_iff {k : 𝓐} {n : ℕ} {ω : Ω} : + actionIndicator A k n ω = 1 ↔ A n ω = k := by simp [actionIndicator] + +lemma actionIndicator_nonneg (A : ℕ → Ω → 𝓐) (k : 𝓐) (n : ℕ) (ω : Ω) : + 0 ≤ actionIndicator A k n ω := + Set.indicator_apply_nonneg fun _ ↦ zero_le_one + +lemma actionIndicator_le_one (A : ℕ → Ω → 𝓐) (k : 𝓐) (n : ℕ) (ω : Ω) : + actionIndicator A k n ω ≤ 1 := by + unfold actionIndicator + by_cases h : A n ω = k <;> simp [h] + +/-- Exactly one arm is pulled at each round, so the indicators sum to `1`. -/ +lemma sum_actionIndicator [Fintype 𝓐] (A : ℕ → Ω → 𝓐) (j : ℕ) (ω : Ω) : + ∑ k, actionIndicator A k j ω = 1 := by + classical + simp [actionIndicator, Set.indicator_apply] + +lemma sum_actionIndicator_eq_pullCount [DecidableEq 𝓐] (A : ℕ → Ω → 𝓐) (k : 𝓐) (n : ℕ) + (ω : Ω) : + ∑ j ∈ range n, actionIndicator A k j ω = (pullCount A k n ω : ℝ) := by + classical + rw [pullCount_eq_sum] + push_cast + refine Finset.sum_congr rfl fun j _ ↦ ?_ + simp only [actionIndicator, Set.indicator_apply, Set.mem_ofPred_eq] + +lemma sum_actionIndicator_smul [DecidableEq 𝓐] [AddCommGroup 𝓨] [Module ℝ 𝓨] + (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) (k : 𝓐) (t : ℕ) (ω : Ω) : + ∑ j ∈ range t, actionIndicator A k j ω • Y j ω = sumRewards A Y k t ω := by + rw [sumRewards] + refine Finset.sum_congr rfl fun j _ ↦ ?_ + simp only [actionIndicator, Set.indicator_apply, Set.mem_ofPred_eq] + split_ifs <;> simp + +lemma sum_actionIndicator_mul [DecidableEq 𝓐] (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → ℝ) (k : 𝓐) (t : ℕ) + (ω : Ω) : + ∑ j ∈ range t, actionIndicator A k j ω * Y j ω = sumRewards A Y k t ω := + sum_actionIndicator_smul A Y k t ω + +lemma measurable_actionIndicator (k : 𝓐) {n : ℕ} (hA : Measurable (A n)) : + Measurable (actionIndicator A k n) := + measurable_const.indicator (hA (measurableSet_singleton k)) + +lemma integrable_actionIndicator (P : Measure Ω) [IsFiniteMeasure P] + (k : 𝓐) {n : ℕ} (hA : Measurable (A n)) : + Integrable (actionIndicator A k n) P := + (integrable_const (1 : ℝ)).indicator (hA (measurableSet_singleton k)) + +/-- The action indicator is adapted to the history filtration: whether action `k` was chosen at `n` +is known at time `n`. -/ +lemma IsAlgEnvSeq.adapted_actionIndicator {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} + [IsFiniteMeasure P] (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) : + Adapted h.filtration (actionIndicator A k) := + fun _ ↦ Measurable.indicator measurable_const (h.adapted_action _ (measurableSet_singleton k)) + +/-- The action indicator is adapted to the history+action filtration: whether action `k` was chosen +at `n` is known once we know the action at `n`. -/ +lemma IsAlgEnvSeq.adapted_actionIndicator_filtrationAction + {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} + [IsFiniteMeasure P] (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) : + Adapted h.filtrationAction (actionIndicator A k) := + fun _ ↦ Measurable.indicator measurable_const + (h.adapted_action_filtrationAction _ (measurableSet_singleton k)) + +end Learning diff --git a/LeanMachineLearning/SequentialLearning/Algorithm.lean b/LeanMachineLearning/SequentialLearning/Algorithm.lean index 46ad1617..c3091abf 100644 --- a/LeanMachineLearning/SequentialLearning/Algorithm.lean +++ b/LeanMachineLearning/SequentialLearning/Algorithm.lean @@ -250,6 +250,18 @@ lemma IsAlgEnvSeq.hasLaw_history_zero (h : IsAlgEnvSeq A Y alg env P) : HasLaw ( have hY := h.measurable_feedback exact (Measure.map_map (by fun_prop) (by fun_prop)).symm +lemma IsAlgEnvSeq.hasLaw_action_comp (h : IsAlgEnvSeq A Y alg env P) (n : ℕ) : + HasLaw (A (n + 1)) (alg.policy n ∘ₘ (P.map (history A Y n))) P := + HasCondDistrib.hasLaw_comp (h.hasCondDistrib_action n) + +lemma IsAlgEnvSeq.hasLaw_feedback_comp (h : IsAlgEnvSeq A Y alg env P) (n : ℕ) : + HasLaw (Y (n + 1)) ((env.feedback n) ∘ₘ (P.map fun ω ↦ (history A Y n ω, A (n + 1) ω))) P := + HasCondDistrib.hasLaw_comp (h.hasCondDistrib_feedback n) + +lemma IsAlgEnvSeq.hasLaw_feedback_zero_comp (h : IsAlgEnvSeq A Y alg env P) : + HasLaw (Y 0) (env.ν0 ∘ₘ (P.map (A 0))) P := + HasCondDistrib.hasLaw_comp (h.hasCondDistrib_feedback_zero) + section Filtration namespace IsAlgEnvSeq diff --git a/LeanMachineLearning/SequentialLearning/FeedbackMartingale.lean b/LeanMachineLearning/SequentialLearning/FeedbackMartingale.lean new file mode 100644 index 00000000..34aedb10 --- /dev/null +++ b/LeanMachineLearning/SequentialLearning/FeedbackMartingale.lean @@ -0,0 +1,221 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import LeanMachineLearning.SequentialLearning.ActionIndicator +public import LeanMachineLearning.SequentialLearning.Means + +/-! +# Martingale decomposition of the sum of rewards + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Finset Learning + +open scoped ENNReal + +namespace Learning + +variable {Ω 𝓐 𝓨 : Type*} {mΩ : MeasurableSpace Ω} {m𝓐 : MeasurableSpace 𝓐} {m𝓨 : MeasurableSpace 𝓨} + [NormedAddCommGroup 𝓨] [NormedSpace ℝ 𝓨] + {P : Measure Ω} [IsFiniteMeasure P] + {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} + +/-- The sum of noise terms for action `k`. +This is the martingale part of `sumRewards A Y k` for the filtration +`IsAlgEnvSeq.filtrationAction`. -/ +noncomputable +def noiseSum (env : Environment 𝓐 𝓨) (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) (k : 𝓐) (n : ℕ) (ω : Ω) : 𝓨 := + ∑ m ∈ range n, {ω | A m ω = k}.indicator (fun ω ↦ Y m ω - env.means A Y (A m ω) m ω) ω + +/-- The sum of mean terms for action `k`. +This is the predictable part of `sumRewards A Y k` for the filtration +`IsAlgEnvSeq.filtrationAction`. -/ +noncomputable +def meanSum (env : Environment 𝓐 𝓨) (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) (k : 𝓐) (n : ℕ) (ω : Ω) : 𝓨 := + ∑ m ∈ range n, {ω | A m ω = k}.indicator (fun ω ↦ env.means A Y (A m ω) m ω) ω + +lemma noiseSum_add_meanSum' (k : 𝓐) (n : ℕ) (ω : Ω) : + noiseSum env A Y k n ω + meanSum env A Y k n ω = + ∑ m ∈ range n, {ω | A m ω = k}.indicator (Y m) ω := by + simp only [noiseSum, meanSum, ← sum_add_distrib] + congr with m + by_cases h : A m ω = k <;> simp [h] + +lemma noiseSum_add_meanSum [DecidableEq 𝓐] (k : 𝓐) (n : ℕ) (ω : Ω) : + noiseSum env A Y k n ω + meanSum env A Y k n ω = sumRewards A Y k n ω := by + unfold sumRewards + rw [noiseSum_add_meanSum' k n ω] + congr with m + by_cases h : A m ω = k <;> simp [h] + +@[simp] +lemma noiseSum_zero (k : 𝓐) : noiseSum env A Y k 0 = fun _ ↦ 0 := by unfold noiseSum; simp + +@[simp] +lemma meanSum_zero (k : 𝓐) : meanSum env A Y k 0 = fun _ ↦ 0 := by unfold meanSum; simp + +lemma noiseSum_succ (k : 𝓐) (n : ℕ) : + noiseSum env A Y k (n + 1) = noiseSum env A Y k n + + {ω | A n ω = k}.indicator (fun ω ↦ Y n ω - env.means A Y (A n ω) n ω) := by + ext ω + simp [noiseSum, Finset.sum_range_succ] + +lemma noiseSum_succ_sub (k : 𝓐) (n : ℕ) (ω : Ω) : + noiseSum env A Y k (n + 1) ω - noiseSum env A Y k n ω + = {ω | A n ω = k}.indicator (fun ω ↦ Y n ω - env.means A Y (A n ω) n ω) ω := by + simp [noiseSum_succ] + +lemma meanSum_succ (k : 𝓐) (n : ℕ) : + meanSum env A Y k (n + 1) = meanSum env A Y k n + + {ω | A n ω = k}.indicator (fun ω ↦ env.means A Y (A n ω) n ω) := by + ext ω + simp [meanSum, Finset.sum_range_succ] + +lemma meanSum_succ_sub (k : 𝓐) (n : ℕ) (ω : Ω) : + meanSum env A Y k (n + 1) ω - meanSum env A Y k n ω + = {ω | A n ω = k}.indicator (fun ω ↦ env.means A Y (A n ω) n ω) ω := by + simp [meanSum_succ] + +variable [MeasurableSingletonClass 𝓐] [SecondCountableTopology 𝓨] + +@[fun_prop] +lemma IsAlgEnvSeq.integrable_noiseSum_increment [OpensMeasurableSpace 𝓨] + {m : ℕ} (h : IsAlgEnvSeq A Y alg env P) (hint : Integrable (Y m) P) (k : 𝓐) : + Integrable (fun ω ↦ {ω | A m ω = k}.indicator + (fun ω ↦ Y m ω - env.means A Y (A m ω) m ω) ω) P := by + exact (hint.sub (h.integrable_means_action hint)).indicator + (h.measurable_action _ (measurableSet_singleton k)) + +@[fun_prop] +lemma IsAlgEnvSeq.integrable_meanSum_increment [OpensMeasurableSpace 𝓨] + {m : ℕ} (h : IsAlgEnvSeq A Y alg env P) (hint : Integrable (Y m) P) (k : 𝓐) : + Integrable (fun ω ↦ {ω | A m ω = k}.indicator (fun ω ↦ env.means A Y (A m ω) m ω) ω) P := by + exact (h.integrable_means_action hint).indicator + (h.measurable_action _ (measurableSet_singleton k)) + +@[fun_prop] +lemma IsAlgEnvSeq.integrable_noiseSum [OpensMeasurableSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (hint : ∀ n, Integrable (Y n) P) (k : 𝓐) (n : ℕ) : + Integrable (noiseSum env A Y k n) P := + integrable_finsetSum _ fun m _ ↦ h.integrable_noiseSum_increment (hint m) k + +@[fun_prop] +lemma IsAlgEnvSeq.integrable_meanSum [OpensMeasurableSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (hint : ∀ n, Integrable (Y n) P) (k : 𝓐) (n : ℕ) : + Integrable (meanSum env A Y k n) P := + integrable_finsetSum _ fun m _ ↦ h.integrable_meanSum_increment (hint m) k + +lemma IsAlgEnvSeq.memLp_noiseSum_increment [BorelSpace 𝓨] + {m : ℕ} (k : 𝓐) (h : IsAlgEnvSeq A Y alg env P) {p : ℝ≥0∞} (hp1 : 1 ≤ p) (hp_top : p ≠ ∞) + (hY : MemLp (Y m) p P) : + MemLp ({ω | A m ω = k}.indicator (fun ω ↦ Y m ω - env.means A Y (A m ω) m ω)) p P := by + refine (hY.sub ?_).indicator (h.measurable_action _ (measurableSet_singleton k)) + exact h.memLp_means_action hp1 hp_top hY + +lemma IsAlgEnvSeq.memLp_meanSum_increment [BorelSpace 𝓨] + {m : ℕ} (k : 𝓐) (h : IsAlgEnvSeq A Y alg env P) {p : ℝ≥0∞} (hp1 : 1 ≤ p) (hp_top : p ≠ ∞) + (hY : MemLp (Y m) p P) : + MemLp ({ω | A m ω = k}.indicator (fun ω ↦ env.means A Y (A m ω) m ω)) p P := by + exact (h.memLp_means_action hp1 hp_top hY).indicator + (h.measurable_action _ (measurableSet_singleton k)) + +lemma IsAlgEnvSeq.memLp_noiseSum [BorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) {p : ℝ≥0∞} (hp1 : 1 ≤ p) (hp_top : p ≠ ∞) + (hY : ∀ n, MemLp (Y n) p P) (k : 𝓐) (n : ℕ) : + MemLp (noiseSum env A Y k n) p P := + memLp_finsetSum _ fun m _ ↦ memLp_noiseSum_increment k h hp1 hp_top (hY m) + +lemma IsAlgEnvSeq.memLp_meanSum [BorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) {p : ℝ≥0∞} (hp1 : 1 ≤ p) (hp_top : p ≠ ∞) + (hY : ∀ n, MemLp (Y n) p P) (k : 𝓐) (n : ℕ) : + MemLp (meanSum env A Y k n) p P := + memLp_finsetSum _ fun m _ ↦ memLp_meanSum_increment k h hp1 hp_top (hY m) + +section Martingale + +variable [BorelSpace 𝓨] + +lemma IsAlgEnvSeq.adapted_noiseSum (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) : + Adapted h.filtrationAction (noiseSum env A Y k) := by + refine fun n ↦ Finset.measurable_fun_sum _ fun m hm ↦ ?_ + have hAm : Measurable[h.filtrationAction n] (A m) := + h.adapted_action_filtrationAction.measurable_le (by grind) + have hYm : Measurable[h.filtrationAction n] (Y m) := + h.measurable_feedback_filtrationAction_of_lt (by grind) + refine (hYm.sub ?_).indicator (hAm (measurableSet_singleton k)) + exact h.adapted_means_filtrationAction.measurable_le (by grind) + +lemma IsAlgEnvSeq.stronglyAdapted_noiseSum (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) : + StronglyAdapted h.filtrationAction (noiseSum env A Y k) := + (adapted_noiseSum h k).stronglyAdapted + +lemma IsAlgEnvSeq.isStronglyPredictable_meanSum (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) : + IsStronglyPredictable h.filtrationAction (meanSum env A Y k) := by + refine .of_measurable_add_one ?_ fun n ↦ ?_ + · simp only [meanSum_zero] + fun_prop + · refine Finset.stronglyMeasurable_fun_sum _ fun m hm ↦ ?_ + have hAm : Measurable[h.filtrationAction n] (A m) := + h.adapted_action_filtrationAction.measurable_le (by grind) + refine StronglyMeasurable.indicator ?_ (hAm (measurableSet_singleton k)) + exact (h.stronglyAdapted_means_filtrationAction m).mono (h.filtrationAction.mono (by grind)) + +lemma IsAlgEnvSeq.condExp_noiseSum_increment [CompleteSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) (i : ℕ) (hint : Integrable (Y i) P) : + P[{ω | A i ω = k}.indicator (fun ω ↦ Y i ω - env.means A Y (A i ω) i ω) | h.filtrationAction i] + =ᵐ[P] 0 := by + let c : Ω → ℝ := actionIndicator A k i + let g : Ω → 𝓨 := fun ω ↦ Y i ω - env.means A Y (A i ω) i ω + have h_smul : c • g = {ω | A i ω = k}.indicator (fun ω ↦ Y i ω - env.means A Y (A i ω) i ω) := by + ext ω + by_cases hω : A i ω = k <;> simp [c, g, actionIndicator, hω] + have hAG : Measurable[h.filtrationAction i] (A i) := h.adapted_action_filtrationAction i + have hcG : StronglyMeasurable[h.filtrationAction i] c := + (h.adapted_actionIndicator_filtrationAction k i).stronglyMeasurable + have hgint : Integrable g P := hint.sub (h.integrable_means_action hint) + have hcint : Integrable (c • g) P := by + rw [h_smul] + exact integrable_noiseSum_increment h hint k + have hcondg : P[g | h.filtrationAction i] =ᵐ[P] 0 := by + refine (condExp_sub hint (h.integrable_means_action hint) _).trans ?_ + have h1 := h.condExp_feedback i hint + grw [h1] + rw [condExp_of_stronglyMeasurable] + · simp + · exact h.adapted_means_filtrationAction.stronglyAdapted i + · exact h.integrable_means_action hint + have hpull := condExp_smul_of_aestronglyMeasurable_left hcG.aestronglyMeasurable hcint hgint + filter_upwards [hpull, hcondg] with ω hp hcg + rw [← h_smul, hp] + simp only [Pi.smul_apply', hcg, Pi.ofNat_apply, smul_eq_zero] + rcases eq_or_ne (A i ω) k with hak | hak + · simp + · simp [c, actionIndicator, hak] + +lemma IsAlgEnvSeq.martingale_noiseSum [CompleteSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (hint : ∀ n, Integrable (Y n) P) (k : 𝓐) : + Martingale (noiseSum env A Y k) h.filtrationAction P := by + have hInt : ∀ n, Integrable (noiseSum env A Y k n) P := h.integrable_noiseSum (hint) k + refine martingale_nat (h.stronglyAdapted_noiseSum k) hInt fun i ↦ ?_ + rw [noiseSum_succ] + symm + have hadd := condExp_add (hInt i) + (integrable_noiseSum_increment h (hint i) k) (h.filtrationAction i) + have hself : P[noiseSum env A Y k i | h.filtrationAction i] = noiseSum env A Y k i := + condExp_of_stronglyMeasurable (h.filtrationAction.le i) + (h.stronglyAdapted_noiseSum k i) (hInt i) + have hincr := condExp_noiseSum_increment h k i (hint i) + filter_upwards [hadd, hincr] with ω ha hin + rw [ha, Pi.add_apply, congrFun hself ω] + simp only [add_eq_left] + rw [hin, Pi.zero_apply] + +end Martingale + +end Learning diff --git a/LeanMachineLearning/SequentialLearning/FiniteActions.lean b/LeanMachineLearning/SequentialLearning/FiniteActions.lean index a1bb588d..e5c8644b 100644 --- a/LeanMachineLearning/SequentialLearning/FiniteActions.lean +++ b/LeanMachineLearning/SequentialLearning/FiniteActions.lean @@ -238,7 +238,7 @@ lemma stronglyAdapted_pullCount_add_one [MeasurableSingletonClass 𝓐] StronglyAdapted h.filtration (fun n ↦ pullCount A a (n + 1)) := (adapted_pullCount_add_one h a).stronglyAdapted -lemma isPredictable_pullCount [MeasurableSingletonClass 𝓐] +lemma isStronglyPredictable_pullCount [MeasurableSingletonClass 𝓐] (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : IsStronglyPredictable h.filtration (pullCount A a) := by rw [IsStronglyPredictable.iff_measurable_add_one] @@ -311,7 +311,6 @@ lemma stepsUntil_eq_dite (a : 𝓐) (m : ℕ) (ω : Ω) simpa using (h' s) set_option backward.isDefEq.respectTransparency false in --- todo: this is in ℝ because of the limited def of leastGE lemma stepsUntil_eq_leastGE (a : 𝓐) (hm : m ≠ 0) : stepsUntil A a m = leastGE (fun n (ω : Ω) ↦ pullCount A a (n + 1) ω) m := by classical @@ -774,260 +773,4 @@ lemma sum_pullCount' [Fintype 𝓐] (n : ℕ) (h : Iic n → 𝓐 × ℝ) : ∑ simp [Finset.sum_ite_eq univ (h s).1 (fun _ ↦ (1 : ℕ))] simp [hcol] -section SumRewards - -/-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ -def sumRewards (A : ℕ → Ω → 𝓐) (R' : ℕ → Ω → ℝ) (a : 𝓐) (t : ℕ) (ω : Ω) : ℝ := - ∑ s ∈ range t, if A s ω = a then R' s ω else 0 - -/-- Sum of rewards of arm `a` up to (and including) time `n`. -/ -noncomputable -def sumRewards' (n : ℕ) (h : Iic n → 𝓐 × ℝ) (a : 𝓐) := - ∑ s, if (h s).1 = a then (h s).2 else 0 - -/-- Empirical mean reward obtained when pulling action `a` up to time `t` (exclusive). -/ -noncomputable -def empMean (A : ℕ → Ω → 𝓐) (R' : ℕ → Ω → ℝ) (a : 𝓐) (t : ℕ) (ω : Ω) : ℝ := - sumRewards A R' a t ω / pullCount A a t ω - -/-- Empirical mean of arm `a` at time `n`. -/ -noncomputable -def empMean' (n : ℕ) (h : Iic n → 𝓐 × ℝ) (a : 𝓐) := - (sumRewards' n h a) / (pullCount' n h a) - -@[simp] -lemma sumRewards_zero {R' : ℕ → Ω → ℝ} : sumRewards A R' a 0 = 0 := by ext; simp [sumRewards] - -lemma sumRewards_add_one {R' : ℕ → Ω → ℝ} : - sumRewards A R' a (t + 1) ω = sumRewards A R' a t ω + if A t ω = a then R' t ω else 0 := by - unfold sumRewards - rw [sum_range_succ] - -lemma sumRewards_eq_of_pullCount_eq {R' : ℕ → Ω → ℝ} {s t : ℕ} - (h_eq : pullCount A a s ω = pullCount A a t ω) : - sumRewards A R' a s ω = sumRewards A R' a t ω := by - wlog hst : s ≤ t - · have hts : t ≤ s := by lia - exact (this h_eq.symm hts).symm - induction t, hst using Nat.le_induction with - | base => rfl - | succ t hst' ih => - have h_mono' : pullCount A a t ω ≤ pullCount A a (t + 1) ω := pullCount_mono a (Nat.le_succ t) ω - have h_eq_t : pullCount A a s ω = pullCount A a t ω := - le_antisymm (pullCount_mono a hst' ω) (h_eq ▸ h_mono') - have hne : A t ω ≠ a := by - intro ha - have h1 := ha ▸ pullCount_action_eq_pullCount_add_one (A := A) t ω - lia - rw [sumRewards_add_one, ite_eq_right hne, add_zero, ih h_eq_t] - -lemma sumRewards_eq_pullCount_mul_empMean {R' : ℕ → Ω → ℝ} {ω : Ω} - (h_pull : pullCount A a t ω ≠ 0) : - sumRewards A R' a t ω = pullCount A a t ω * empMean A R' a t ω := by unfold empMean; field_simp - -lemma sum_rewardByCount_eq_sumRewards {R' : ℕ → Ω → ℝ} (a : 𝓐) (t : ℕ) (ω : Ω × (ℕ → 𝓐 → ℝ)) : - ∑ m ∈ Icc 1 (pullCount A a t ω.1), rewardByCount A R' a m ω = sumRewards A R' a t ω.1 := by - induction t with - | zero => simp [pullCount, sumRewards] - | succ t ht => - by_cases hta : A t ω.1 = a - · rw [← hta] at ht ⊢ - rw [pullCount_action_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] - unfold sumRewards - rw [sum_range_succ, ite_eq_left rfl, rewardByCount_pullCount_add_one_eq_reward] - · unfold sumRewards - rwa [pullCount_eq_pullCount_of_action_ne hta, sum_range_succ, ite_eq_right hta, add_zero] - -lemma sumRewards_add_one_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : - sumRewards A R' a (n + 1) ω = sumRewards' n (fun i ↦ (A i ω, R' i ω)) a := by - unfold sumRewards sumRewards' - rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then R' s ω else 0) (Iic n)] - congr with m - simp only [mem_range, mem_Iic] - grind - -lemma sumRewards_eq_sumRewards' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : - sumRewards A R' a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by - cases n with - | zero => exact absurd rfl hn - | succ n => simp [sumRewards_add_one_eq_sumRewards'] - -lemma empMean_add_one_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : - empMean A R' a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R' i ω)) a := by - unfold empMean empMean' - rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] - -lemma empMean_eq_empMean' {R' : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : - empMean A R' a n ω = empMean' (n - 1) (fun i ↦ (A i ω, R' i ω)) a := by - unfold empMean empMean' - rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] - -lemma sumRewards_sub_pullCount_mul_eq_sum {R' : ℕ → Ω → ℝ} (c : 𝓐 → ℝ) : - sumRewards A R' a (n + 1) ω - pullCount A a (n + 1) ω * c a = - ∑ i ∈ range (n + 1), (if A i ω = a then R' i ω - c a else 0) := by - induction n with - | zero => - simp_rw [sumRewards_add_one, pullCount_add_one] - simp only [sumRewards_zero, Pi.zero_apply, zero_add, pullCount_zero, Nat.cast_ite, Nat.cast_one, - CharP.cast_eq_zero, ite_mul, one_mul, zero_mul, range_one, sum_singleton] - grind - | succ n hn => - simp_rw [sumRewards_add_one (t := n + 1), pullCount_add_one (t := n + 1)] - split_ifs with ha - · conv_rhs => rw [sum_range_succ] - simp only [Nat.cast_add, Nat.cast_one, ha, ↓reduceIte, add_mul, one_mul] - grind - · simp only [add_zero, hn] - conv_rhs => rw [sum_range_succ] - simp [ha] - -@[fun_prop] -lemma measurable_sumRewards [MeasurableSingletonClass 𝓐] {R' : ℕ → Ω → ℝ} - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) (a : 𝓐) (t : ℕ) : - Measurable (sumRewards A R' a t) := by - unfold sumRewards - have h_meas s : Measurable (fun h : Ω ↦ if A s h = a then R' s h else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - fun_prop - -@[fun_prop] -lemma measurable_uncurry_sumRewards_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] - {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) {f : Ω → 𝓐} - (hf : Measurable f) {g : Ω → ℕ} (hg : Measurable g) : - Measurable (fun ω ↦ sumRewards A R' (f ω) (g ω) ω) := by - change Measurable ((fun aω ↦ sumRewards A R' aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) - apply Measurable.comp _ (by fun_prop) - refine measurable_from_prod_countable_right fun a ↦ ?_ - change Measurable ((fun tω ↦ sumRewards A R' a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) - apply Measurable.comp _ (by fun_prop) - exact measurable_from_prod_countable_right (fun t ↦ measurable_sumRewards hA hR' a t) - -@[fun_prop] -lemma measurable_empMean [MeasurableSingletonClass 𝓐] {R' : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) - (hR' : ∀ n, Measurable (R' n)) (a : 𝓐) (n : ℕ) : - Measurable (empMean A R' a n) := by - unfold empMean - fun_prop - -@[fun_prop] -lemma measurable_uncurry_empMean_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] {R' : ℕ → Ω → ℝ} - (hA : ∀ n, Measurable (A n)) (hR' : ∀ n, Measurable (R' n)) {f : Ω → 𝓐} (hf : Measurable f) - {g : Ω → ℕ} (hg : Measurable g) : - Measurable (fun ω ↦ empMean A R' (f ω) (g ω) ω) := by - unfold empMean - fun_prop - -@[fun_prop] -lemma measurable_sumRewards' [MeasurableSingletonClass 𝓐] (n : ℕ) (a : 𝓐) : - Measurable (fun h ↦ sumRewards' n h a) := by - simp_rw [sumRewards'] - have h_meas s : Measurable (fun (h : Iic n → 𝓐 × ℝ) ↦ if (h s).1 = a then (h s).2 else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact (measurableSet_singleton _).preimage (by fun_prop) - fun_prop - -@[fun_prop] -lemma measurable_uncurry_sumRewards' [MeasurableEq 𝓐] (n : ℕ) : - Measurable (fun p : (Iic n → 𝓐 × ℝ) × 𝓐 ↦ sumRewards' n p.1 p.2) := by - simp_rw [sumRewards'] - have h_meas s : Measurable (fun p : (Iic n → 𝓐 × ℝ) × 𝓐 ↦ - if (p.1 s).1 = p.2 then (p.1 s).2 else 0) := by - refine Measurable.ite ?_ (by fun_prop) (by fun_prop) - exact measurableSet_eq_fun (by fun_prop) (by fun_prop) - fun_prop - -@[fun_prop] -lemma measurable_empMean' [MeasurableSingletonClass 𝓐] (n : ℕ) (a : 𝓐) : - Measurable (fun h ↦ empMean' n h a) := by - unfold empMean' - fun_prop - -@[fun_prop] -lemma measurable_uncurry_empMean' [MeasurableEq 𝓐] (n : ℕ) : - Measurable (fun p : (Iic n → 𝓐 × ℝ) × 𝓐 ↦ empMean' n p.1 p.2) := by - unfold empMean' - fun_prop - -lemma IsAlgEnvSeq.isPredictable_sumRewards [StandardBorelSpace 𝓐] {R' : ℕ → Ω → ℝ} - {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} - (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : - IsStronglyPredictable h.filtration (sumRewards A R' a) := by - rw [IsStronglyPredictable.iff_measurable_add_one] - constructor - · simp only [sumRewards_zero] - fun_prop - refine fun n ↦ Measurable.stronglyMeasurable ?_ - refine measurable_fun_sum _ fun i hi ↦ Measurable.ite ?_ ?_ (by fun_prop) - · refine (measurableSet_singleton a).preimage ?_ - have h_meas_i := h.adapted_action i - simp only [mem_range] at hi - exact h_meas_i.mono (h.filtration.mono (by lia)) le_rfl - · have h_meas_i := h.adapted_feedback i - simp only [mem_range] at hi - exact h_meas_i.mono (h.filtration.mono (by lia)) le_rfl - -lemma IsAlgEnvSeq.stronglyAdapted_sumRewards_add_one [StandardBorelSpace 𝓐] - {R' : ℕ → Ω → ℝ} {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} - (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : - StronglyAdapted h.filtration (fun n ↦ sumRewards A R' a (n + 1)) := by - have h_predictable := h.isPredictable_sumRewards a - rw [IsStronglyPredictable.iff_measurable_add_one] at h_predictable - exact h_predictable.2 - -lemma IsAlgEnvSeq.adapted_sumRewards_add_one [StandardBorelSpace 𝓐] {R' : ℕ → Ω → ℝ} - {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} - (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : - Adapted h.filtration (fun n ↦ sumRewards A R' a (n + 1)) := - (h.stronglyAdapted_sumRewards_add_one a).adapted - -section CopiedFromPR - -open Set - -lemma _root_.MeasureTheory.StronglyMeasurable.div₀' {𝓐 β : Type*} - {m𝓐 : MeasurableSpace 𝓐} [TopologicalSpace β] - [GroupWithZero β] [ContinuousMul β] [ContinuousInv₀ β] - [TopologicalSpace.PseudoMetrizableSpace β] - [MeasurableSpace β] [BorelSpace β] [MeasurableSingletonClass β] - {f g : 𝓐 → β} (hf : StronglyMeasurable f) (hg : StronglyMeasurable g) : - StronglyMeasurable (f / g) := by - refine ⟨fun n => hf.approx n / (hg.approx n).restrict {x | g x ≠ 0}, fun x => ?_⟩ - have : MeasurableSet {x | g x ≠ 0} := ((MeasurableSet.singleton 0).preimage hg.measurable).compl - by_cases h : g x = 0 - · simp_all only [ne_eq, SimpleFunc.coe_div, SimpleFunc.coe_restrict, Pi.div_apply, mem_ofPred_eq, - not_true_eq_false, not_false_eq_true, indicator_of_notMem, _root_.div_zero] - exact tendsto_const_nhds - · simp_all only [ne_eq, SimpleFunc.coe_div, SimpleFunc.coe_restrict, - Pi.div_apply, mem_ofPred_eq, not_false_eq_true, indicator_of_mem] - exact (hf.tendsto_approx x).div (hg.tendsto_approx x) h - -end CopiedFromPR - -lemma IsAlgEnvSeq.isPredictable_empMean [StandardBorelSpace 𝓐] {R' : ℕ → Ω → ℝ} - {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} - (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : - IsStronglyPredictable h.filtration (empMean A R' a) := by - unfold empMean - refine StronglyMeasurable.div₀' ?_ ?_ - · exact h.isPredictable_sumRewards a - · have h_meas := (isPredictable_pullCount h a).measurable - fun_prop - -lemma IsAlgEnvSeq.stronglyAdapted_empMean_add_one [StandardBorelSpace 𝓐] - {R' : ℕ → Ω → ℝ} {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} - (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : - StronglyAdapted h.filtration (fun n ↦ empMean A R' a (n + 1)) := by - have h_predictable := h.isPredictable_empMean a - rw [IsStronglyPredictable.iff_measurable_add_one] at h_predictable - exact h_predictable.2 - -lemma IsAlgEnvSeq.adapted_empMean_add_one [StandardBorelSpace 𝓐] {R' : ℕ → Ω → ℝ} - {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} - (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : - Adapted h.filtration (fun n ↦ empMean A R' a (n + 1)) := - (h.stronglyAdapted_empMean_add_one a).adapted - -end SumRewards - end Learning diff --git a/LeanMachineLearning/SequentialLearning/Means.lean b/LeanMachineLearning/SequentialLearning/Means.lean new file mode 100644 index 00000000..7dad5aee --- /dev/null +++ b/LeanMachineLearning/SequentialLearning/Means.lean @@ -0,0 +1,257 @@ +/- +Copyright (c) 2026 Rémy Degenne. All rights reserved. +Released under Apache 2.0 license as described in the file LICENSE. +Authors: Rémy Degenne +-/ +module + +public import LeanMachineLearning.SequentialLearning.StationaryEnv +public import LeanMachineLearning.ForMathlib.Probability.Kernel.Composition.IntegralCompProd + +/-! +# The means of the feedback distribution + +## Main definitions + +* `Environment.means` + +## Main results + +* +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Filter Finset + +open scoped ENNReal + +namespace ProbabilityTheory + +variable {Ω β 𝓨 : Type*} {mΩ : MeasurableSpace Ω} {mβ : MeasurableSpace β} + {m𝓨 : MeasurableSpace 𝓨} [StandardBorelSpace 𝓨] [Nonempty 𝓨] + {P : Measure Ω} [IsFiniteMeasure P] {X : Ω → β} {Y : Ω → 𝓨} + {κ : Kernel β 𝓨} [IsFiniteKernel κ] + +lemma HasCondDistrib.condExp_comp_eq {F : Type*} [NormedAddCommGroup F] [NormedSpace ℝ F] + [CompleteSpace F] (h : HasCondDistrib Y X κ P) (hX : Measurable X) + {g : 𝓨 → F} (hg : StronglyMeasurable g) (hint : Integrable (fun ω ↦ g (Y ω)) P) : + P[fun ω ↦ g (Y ω) | mβ.comap X] =ᵐ[P] fun ω ↦ ∫ y, g y ∂(κ (X ω)) := by + refine (condExp_ae_eq_integral_condDistrib hX h.aemeasurable_snd hg hint).trans ?_ + filter_upwards [ae_of_ae_map hX.aemeasurable h.condDistrib_eq] with ω hω + rw [hω] + +end ProbabilityTheory + +namespace Learning + +variable {Ω 𝓐 𝓨 : Type*} {mΩ : MeasurableSpace Ω} {m𝓐 : MeasurableSpace 𝓐} {m𝓨 : MeasurableSpace 𝓨} + [NormedAddCommGroup 𝓨] [NormedSpace ℝ 𝓨] + {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} {P : Measure Ω} [IsFiniteMeasure P] + {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} + +/-- The kernel that gives the measure of the feedback distribution as a function of the action +chosen at time `n`. -/ +noncomputable def Environment.measure (env : Environment 𝓐 𝓨) (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) + (n : ℕ) (ω : Ω) : Kernel 𝓐 𝓨 := + if n = 0 then env.ν0 else (env.feedback (n - 1)).sectR (history A Y (n - 1) ω) + +/-- The means of the feedback distribution as a function of the action chosen at time `n`. -/ +noncomputable def Environment.means (env : Environment 𝓐 𝓨) (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) + (k : 𝓐) (n : ℕ) (ω : Ω) : 𝓨 := + (env.measure A Y n ω k)[id] + +@[simp] +lemma means_zero (env : Environment 𝓐 𝓨) (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) + (k : 𝓐) (ω : Ω) : + env.means A Y k 0 ω = (env.ν0 k)[id] := by simp [Environment.means, Environment.measure] + +@[simp] +lemma means_of_isObliviousEnv [IsObliviousEnv env] (A : ℕ → Ω → 𝓐) (Y : ℕ → Ω → 𝓨) + (k : 𝓐) (n : ℕ) (ω : Ω) : + env.means A Y k n ω = (feedbackCondAction env n k)[id] := by + simp only [Environment.means, Environment.measure, ν0_eq_feedbackCondAction, id_eq, + feedback_eq_feedbackCondAction] + split_ifs with hn + · simp [hn] + · simp [Nat.sub_add_cancel (by grind : 1 ≤ n)] + +lemma means_obliviousEnv (ν : ℕ → Kernel 𝓐 𝓨) [∀ n, IsMarkovKernel (ν n)] + (k : 𝓐) (n : ℕ) (ω : Ω) : + (obliviousEnv ν).means A Y k n ω = (ν n k)[id] := by simp + +lemma means_stationaryEnv (ν : Kernel 𝓐 𝓨) [IsMarkovKernel ν] (k : 𝓐) (n : ℕ) (ω : Ω) : + (stationaryEnv ν).means A Y k n ω = (ν k)[id] := by simp + +@[fun_prop] +lemma IsAlgEnvSeq.stronglyMeasurable_means [SecondCountableTopology 𝓨] [OpensMeasurableSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) (n : ℕ) : + StronglyMeasurable (env.means A Y k n) := by + unfold Environment.means + have h_eq ω : env.measure A Y n ω k = + (if n = 0 then env.ν0 ∘ₖ (Kernel.deterministic (fun _ ↦ k) (by fun_prop)) + else (env.feedback (n - 1)) ∘ₖ (Kernel.deterministic (fun ω ↦ (history A Y (n - 1) ω, k)) + ((h.measurable_history (n - 1)).prodMk (by fun_prop)))) ω := by + split_ifs with hn <;> simp [hn, Environment.measure, Kernel.comp_deterministic_eq_comap] + simp_rw [h_eq] + fun_prop + +@[fun_prop] +lemma IsAlgEnvSeq.measurable_means [SecondCountableTopology 𝓨] [BorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (k : 𝓐) (n : ℕ) : + Measurable (env.means A Y k n) := + (h.stronglyMeasurable_means k n).measurable + +lemma IsAlgEnvSeq.adapted_means_filtrationAction [SecondCountableTopology 𝓨] [BorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) : + Adapted h.filtrationAction (fun n ω ↦ env.means A Y (A n ω) n ω) := by + intro n + cases n with + | zero => exact measurable_comp_comap _ stronglyMeasurable_id.integral_kernel.measurable + | succ n => + simp only [Environment.means, Environment.measure, Nat.add_eq_zero_iff, one_ne_zero, and_false, + ↓reduceIte, Nat.add_one_sub_one, Kernel.sectR_apply, id_eq] + change Measurable[h.filtrationAction (n + 1)] + ((fun ω ↦ ∫ x, x ∂(env.feedback n ω)) ∘ (fun ω ↦ (history A Y n ω, A (n + 1) ω))) + rw [IsAlgEnvSeq.filtrationAction_eq_comap _ _ (by grind)] + exact measurable_comp_comap _ stronglyMeasurable_id.integral_kernel.measurable + +lemma IsAlgEnvSeq.stronglyAdapted_means_filtrationAction [SecondCountableTopology 𝓨] [BorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) : + StronglyAdapted h.filtrationAction (fun n ω ↦ env.means A Y (A n ω) n ω) := + (h.adapted_means_filtrationAction).stronglyAdapted + +lemma IsAlgEnvSeq.adapted_means [SecondCountableTopology 𝓨] [BorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) : + Adapted h.filtration (fun n ω ↦ env.means A Y (A n ω) n ω) := + fun n ↦ (h.adapted_means_filtrationAction n).mono (h.filtrationAction_le_filtration n) le_rfl + +omit [NormedSpace ℝ 𝓨] in +lemma IsAlgEnvSeq.condExp_feedback_zero_comp {𝓩 : Type*} [NormedAddCommGroup 𝓩] [NormedSpace ℝ 𝓩] + [CompleteSpace 𝓩] [StandardBorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) + {g : 𝓨 → 𝓩} (hg : StronglyMeasurable g) (hint : Integrable (fun ω ↦ g (Y 0 ω)) P) : + P[fun ω ↦ g (Y 0 ω) | h.filtrationAction 0] =ᵐ[P] fun ω ↦ (env.ν0 (A 0 ω))[g] := by + have hX : Measurable (fun ω ↦ (history A Y 0 ω, A 0 ω)) := + (h.measurable_history 0).prodMk (h.measurable_action 0) + rw [h.filtrationAction_zero_eq_comap] + exact h.hasCondDistrib_feedback_zero.condExp_comp_eq (h.measurable_action 0) hg hint + +omit [NormedSpace ℝ 𝓨] in +lemma IsAlgEnvSeq.condExp_feedback_comp {𝓩 : Type*} [NormedAddCommGroup 𝓩] [NormedSpace ℝ 𝓩] + [CompleteSpace 𝓩] [StandardBorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (n : ℕ) + {g : 𝓨 → 𝓩} (hg : StronglyMeasurable g) (hint : Integrable (fun ω ↦ g (Y (n + 1) ω)) P) : + P[fun ω ↦ g (Y (n + 1) ω) | h.filtrationAction (n + 1)] =ᵐ[P] + fun ω ↦ (env.feedback n (history A Y n ω, A (n + 1) ω))[g] := by + have hX : Measurable (fun ω ↦ (history A Y n ω, A (n + 1) ω)) := + (h.measurable_history n).prodMk (h.measurable_action (n + 1)) + rw [h.filtrationAction_eq_comap (n + 1) (by simp)] + exact (h.hasCondDistrib_feedback n).condExp_comp_eq hX hg hint + +lemma IsAlgEnvSeq.condExp_feedback [BorelSpace 𝓨] [SecondCountableTopology 𝓨] [CompleteSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) (n : ℕ) + (hint : Integrable (Y n) P) : + P[Y n | h.filtrationAction n] =ᵐ[P] fun ω ↦ env.means A Y (A n ω) n ω := by + cases n with + | zero => exact condExp_feedback_zero_comp h stronglyMeasurable_id hint + | succ n => exact condExp_feedback_comp h n stronglyMeasurable_id hint + +lemma IsAlgEnvSeq.memLp_means_action [SecondCountableTopology 𝓨] [BorelSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) {n : ℕ} {p : ℝ≥0∞} (hp1 : 1 ≤ p) (hp_top : p ≠ ∞) + (hint : MemLp (Y n) p P) : + MemLp (fun ω ↦ env.means A Y (A n ω) n ω) p P := by + have hp0 : p ≠ 0 := by positivity + have hA := h.measurable_action + have h_hist := h.measurable_history + have hint' : MemLp id p (P.map (Y n)) := by + rwa [memLp_map_measure_iff (by fun_prop) (h.measurable_feedback _).aemeasurable] + unfold Environment.means Environment.measure + cases n with + | zero => + simp only [↓reduceIte, id_eq] + rw [h.hasLaw_feedback_zero_comp.map_eq, Measure.memLp_comp_iff hp0 hp_top (by fun_prop)] + at hint' + have hint'' := hint'.2.comp_aemeasurable (by fun_prop) + have h_eq ω : env.ν0 (A 0 ω) = (env.ν0 ∘ₖ Kernel.deterministic (A 0) (by fun_prop)) ω := by + simp [Kernel.comp_deterministic_eq_comap] + rw [← integrable_norm_rpow_iff _ hp0 hp_top] + swap + · refine StronglyMeasurable.aestronglyMeasurable ?_ + simp_rw [h_eq] + exact StronglyMeasurable.integral_kernel (by fun_prop) + simp only [id_eq] at hint'' + refine Integrable.mono' hint'' ?_ ?_ + · refine ((AEMeasurable.norm ?_).pow_const _).aestronglyMeasurable + refine (StronglyMeasurable.measurable ?_).aemeasurable + simp_rw [h_eq] + exact StronglyMeasurable.integral_kernel (by fun_prop) + · simp only [Real.norm_eq_abs, Function.comp_apply] + filter_upwards [ae_of_ae_map (hA 0).aemeasurable hint'.1] with ω hω + rw [abs_of_nonneg (by positivity)] + exact norm_integral_rpow_le_integral_norm_rpow hp1 hp_top hω + | succ n => + simp only [Nat.add_eq_zero_iff, one_ne_zero, and_false, ↓reduceIte, Nat.add_one_sub_one, id_eq] + rw [(h.hasLaw_feedback_comp n).map_eq, Measure.memLp_comp_iff hp0 hp_top (by fun_prop)] at hint' + have hint'' := hint'.2.comp_aemeasurable (by fun_prop) + have h_eq ω : (env.feedback n) (history A Y n ω, A (n + 1) ω) = + (env.feedback n ∘ₖ + Kernel.deterministic (fun ω ↦ (history A Y n ω, A (n + 1) ω)) (by fun_prop)) ω := by + simp [Kernel.comp_deterministic_eq_comap] + rw [← integrable_norm_rpow_iff _ hp0 hp_top] + swap + · refine StronglyMeasurable.aestronglyMeasurable ?_ + simp_rw [Kernel.sectR_apply, h_eq] + exact StronglyMeasurable.integral_kernel (by fun_prop) + simp only [id_eq] at hint'' + refine Integrable.mono' hint'' ?_ ?_ + · refine ((AEMeasurable.norm ?_).pow_const _).aestronglyMeasurable + refine (StronglyMeasurable.measurable ?_).aemeasurable + simp_rw [Kernel.sectR_apply, h_eq] + exact StronglyMeasurable.integral_kernel (by fun_prop) + · simp only [Real.norm_eq_abs, Function.comp_apply, Kernel.sectR_apply] + filter_upwards [ae_of_ae_map ((h_hist n).prodMk (hA (n + 1))).aemeasurable hint'.1] + with ω hω + rw [abs_of_nonneg (by positivity)] + exact norm_integral_rpow_le_integral_norm_rpow hp1 hp_top hω + +lemma IsAlgEnvSeq.integrable_means_action [SecondCountableTopology 𝓨] [OpensMeasurableSpace 𝓨] + (h : IsAlgEnvSeq A Y alg env P) {n : ℕ} (hint : Integrable (Y n) P) : + Integrable (fun ω ↦ env.means A Y (A n ω) n ω) P := by + have hA := h.measurable_action + have h_hist := h.measurable_history + have hint' : Integrable id (P.map (Y n)) := by + rwa [integrable_map_measure (by fun_prop) (h.measurable_feedback _).aemeasurable] + unfold Environment.means Environment.measure + cases n with + | zero => + simp only [↓reduceIte, id_eq] + rw [h.hasLaw_feedback_zero_comp.map_eq, Measure.integrable_comp_iff (by fun_prop)] at hint' + have hint'' := hint'.2.comp_aemeasurable (by fun_prop) + simp only [id_eq] at hint'' + refine Integrable.mono' hint'' ?_ ?_ + · refine StronglyMeasurable.aestronglyMeasurable ?_ + have h_eq ω : env.ν0 (A 0 ω) = + (env.ν0 ∘ₖ Kernel.deterministic (A 0) (by fun_prop)) ω := by + simp [Kernel.comp_deterministic_eq_comap] + simp_rw [h_eq] + exact StronglyMeasurable.integral_kernel (by fun_prop) + · simp only [Function.comp_apply] + filter_upwards with ω using norm_integral_le_integral_norm _ + | succ n => + simp only [Nat.add_eq_zero_iff, one_ne_zero, and_false, ↓reduceIte, Nat.add_one_sub_one, id_eq] + rw [(h.hasLaw_feedback_comp n).map_eq, Measure.integrable_comp_iff (by fun_prop)] at hint' + have hint'' := hint'.2.comp_aemeasurable (by fun_prop) + simp only [id_eq] at hint'' + refine Integrable.mono' hint'' ?_ ?_ + · refine StronglyMeasurable.aestronglyMeasurable ?_ + have h_eq ω : (env.feedback n) (history A Y n ω, A (n + 1) ω) = + (env.feedback n ∘ₖ + Kernel.deterministic (fun ω ↦ (history A Y n ω, A (n + 1) ω)) (by fun_prop)) ω := by + simp [Kernel.comp_deterministic_eq_comap] + simp_rw [Kernel.sectR_apply, h_eq] + exact StronglyMeasurable.integral_kernel (by fun_prop) + · simp only [Function.comp_apply] + filter_upwards with ω using norm_integral_le_integral_norm _ + +end Learning diff --git a/LeanMachineLearning/SequentialLearning/StationaryEnv.lean b/LeanMachineLearning/SequentialLearning/StationaryEnv.lean index d707a586..1d47b818 100644 --- a/LeanMachineLearning/SequentialLearning/StationaryEnv.lean +++ b/LeanMachineLearning/SequentialLearning/StationaryEnv.lean @@ -77,6 +77,16 @@ variable {Ω : Type*} {mΩ : MeasurableSpace Ω} {A : ℕ → Ω → 𝓐} {Y : ℕ → Ω → 𝓨} {n N : ℕ} {ν : ℕ → Kernel 𝓐 𝓨} [∀ n, IsMarkovKernel (ν n)] +lemma hasCondDistrib_feedback_history_action [IsObliviousEnv env] + (h : IsAlgEnvSeq A Y alg env P) (n : ℕ) : + HasCondDistrib (Y (n + 1)) (fun ω ↦ (history A Y n ω, A (n + 1) ω)) + ((feedbackCondAction env (n + 1)).prodMkLeft _) P := by + have hA := h.measurable_action + have hR' := h.measurable_feedback + refine ⟨by fun_prop, ?_⟩ + have h_eq := (h.hasCondDistrib_feedback n).map_eq + simpa only [feedback_eq_feedbackCondAction] using h_eq + lemma hasCondDistrib_feedback [IsObliviousEnv env] (h : IsAlgEnvSeq A Y alg env P) (n : ℕ) : HasCondDistrib (Y n) (A n) (feedbackCondAction env n) P := by have hA := h.measurable_action diff --git a/LeanMachineLearning/SequentialLearning/SumRewards.lean b/LeanMachineLearning/SequentialLearning/SumRewards.lean new file mode 100644 index 00000000..804e62f3 --- /dev/null +++ b/LeanMachineLearning/SequentialLearning/SumRewards.lean @@ -0,0 +1,253 @@ +/- +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, Paulo Rauber +-/ +module + +public import LeanMachineLearning.SequentialLearning.FiniteActions + +/-! +# Sums of rewards +-/ + +@[expose] public section + +open MeasureTheory Finset Learning + +namespace Learning + +variable {𝓐 𝓨 Ω : Type*} {m𝓐 : MeasurableSpace 𝓐} {m𝓨 : MeasurableSpace 𝓨} {mΩ : MeasurableSpace Ω} + [DecidableEq 𝓐] [AddCommGroup 𝓨] + {P : Measure Ω} [IsProbabilityMeasure P] + {A : ℕ → Ω → 𝓐} {R : ℕ → Ω → 𝓨} + {a : 𝓐} {m n t : ℕ} {ω : Ω} + +/-- Sum of rewards obtained when pulling action `a` up to time `t` (exclusive). -/ +noncomputable def sumRewards (A : ℕ → Ω → 𝓐) (R : ℕ → Ω → 𝓨) (a : 𝓐) (t : ℕ) (ω : Ω) : 𝓨 := + ∑ s ∈ range t, if A s ω = a then (R s) ω else 0 + +/-- Sum of rewards of arm `a` up to (and including) time `n`. -/ +noncomputable +def sumRewards' (n : ℕ) (h : Iic n → 𝓐 × 𝓨) (a : 𝓐) := + ∑ s, if (h s).1 = a then (h s).2 else 0 + +/-- Empirical mean reward obtained when pulling action `a` up to time `t` (exclusive). -/ +noncomputable +def empMean (A : ℕ → Ω → 𝓐) (R : ℕ → Ω → ℝ) (a : 𝓐) (t : ℕ) (ω : Ω) : ℝ := + sumRewards A R a t ω / pullCount A a t ω + +/-- Empirical mean of arm `a` at time `n`. -/ +noncomputable +def empMean' (n : ℕ) (h : Iic n → 𝓐 × ℝ) (a : 𝓐) := + sumRewards' n h a / pullCount' n h a + +@[simp] +lemma sumRewards_zero {R : ℕ → Ω → 𝓨} : sumRewards A R a 0 = 0 := by ext; simp [sumRewards] + +lemma sumRewards_add_one {R : ℕ → Ω → 𝓨} : + sumRewards A R a (t + 1) ω = sumRewards A R a t ω + if A t ω = a then R t ω else 0 := by + unfold sumRewards + rw [sum_range_succ] + +lemma sumRewards_eq_of_pullCount_eq {R : ℕ → Ω → 𝓨} {s t : ℕ} + (h_eq : pullCount A a s ω = pullCount A a t ω) : + sumRewards A R a s ω = sumRewards A R a t ω := by + wlog hst : s ≤ t + · have hts : t ≤ s := by lia + exact (this h_eq.symm hts).symm + induction t, hst using Nat.le_induction with + | base => rfl + | succ t hst' ih => + have h_mono' : pullCount A a t ω ≤ pullCount A a (t + 1) ω := pullCount_mono a (Nat.le_succ t) ω + have h_eq_t : pullCount A a s ω = pullCount A a t ω := + le_antisymm (pullCount_mono a hst' ω) (h_eq ▸ h_mono') + have hne : A t ω ≠ a := by + intro ha + have h1 := ha ▸ pullCount_action_eq_pullCount_add_one (A := A) t ω + lia + rw [sumRewards_add_one, ite_eq_right hne, add_zero, ih h_eq_t] + +lemma sumRewards_eq_pullCount_mul_empMean {R : ℕ → Ω → ℝ} {ω : Ω} + (h_pull : pullCount A a t ω ≠ 0) : + sumRewards A R a t ω = pullCount A a t ω * empMean A R a t ω := by unfold empMean; field_simp + +lemma sum_rewardByCount_eq_sumRewards {R : ℕ → Ω → 𝓨} (a : 𝓐) (t : ℕ) (ω : Ω × (ℕ → 𝓐 → 𝓨)) : + ∑ m ∈ Icc 1 (pullCount A a t ω.1), rewardByCount A R a m ω = sumRewards A R a t ω.1 := by + induction t with + | zero => simp [pullCount, sumRewards] + | succ t ht => + by_cases hta : A t ω.1 = a + · rw [← hta] at ht ⊢ + rw [pullCount_action_eq_pullCount_add_one, sum_Icc_succ_top (Nat.le_add_left 1 _), ht] + unfold sumRewards + rw [sum_range_succ, ite_eq_left rfl, rewardByCount_pullCount_add_one_eq_reward] + · unfold sumRewards + rwa [pullCount_eq_pullCount_of_action_ne hta, sum_range_succ, ite_eq_right hta, add_zero] + +lemma sumRewards_add_one_eq_sumRewards' {R : ℕ → Ω → 𝓨} {n : ℕ} {ω : Ω} : + sumRewards A R a (n + 1) ω = sumRewards' n (fun i ↦ (A i ω, R i ω)) a := by + unfold sumRewards sumRewards' + rw [Finset.sum_coe_sort (f := fun s ↦ if A s ω = a then R s ω else 0) (Iic n)] + congr with m + simp only [mem_range, mem_Iic] + grind + +lemma sumRewards_eq_sumRewards' {R : ℕ → Ω → 𝓨} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + sumRewards A R a n ω = sumRewards' (n - 1) (fun i ↦ (A i ω, R i ω)) a := by + cases n with + | zero => exact absurd rfl hn + | succ n => simp [sumRewards_add_one_eq_sumRewards'] + +lemma empMean_add_one_eq_empMean' {R : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} : + empMean A R a (n + 1) ω = empMean' n (fun i ↦ (A i ω, R i ω)) a := by + unfold empMean empMean' + rw [sumRewards_add_one_eq_sumRewards', pullCount_add_one_eq_pullCount'] + +lemma empMean_eq_empMean' {R : ℕ → Ω → ℝ} {n : ℕ} {ω : Ω} (hn : n ≠ 0) : + empMean A R a n ω = empMean' (n - 1) (fun i ↦ (A i ω, R i ω)) a := by + unfold empMean empMean' + rw [sumRewards_eq_sumRewards' hn, pullCount_eq_pullCount' hn] + +lemma sumRewards_sub_pullCount_smul_eq_sum {R : ℕ → Ω → 𝓨} (c : 𝓐 → 𝓨) : + sumRewards A R a (n + 1) ω - pullCount A a (n + 1) ω • c a = + ∑ i ∈ range (n + 1), (if A i ω = a then R i ω - c a else 0) := by + induction n with + | zero => simp_rw [sumRewards_add_one, pullCount_add_one]; simp; grind + | succ n hn => + simp_rw [sumRewards_add_one (t := n + 1), pullCount_add_one (t := n + 1)] + split_ifs with ha + · conv_rhs => rw [sum_range_succ] + simp only [ha, ↓reduceIte] + rw [add_smul] + grind + · simp only [add_zero, hn] + conv_rhs => rw [sum_range_succ] + simp [ha] + +@[fun_prop] +lemma measurable_sumRewards [MeasurableSingletonClass 𝓐] [MeasurableAdd₂ 𝓨] {R : ℕ → Ω → 𝓨} + (hA : ∀ n, Measurable (A n)) (hR : ∀ n, Measurable (R n)) (a : 𝓐) (t : ℕ) : + Measurable (sumRewards A R a t) := by + unfold sumRewards + have h_meas s : Measurable (fun h : Ω ↦ if A s h = a then R s h else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + fun_prop + +@[fun_prop] +lemma measurable_uncurry_sumRewards_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] + [MeasurableAdd₂ 𝓨] + {R : ℕ → Ω → 𝓨} (hA : ∀ n, Measurable (A n)) (hR : ∀ n, Measurable (R n)) {f : Ω → 𝓐} + (hf : Measurable f) {g : Ω → ℕ} (hg : Measurable g) : + Measurable (fun ω ↦ sumRewards A R (f ω) (g ω) ω) := by + change Measurable ((fun aω ↦ sumRewards A R aω.1 (g aω.2) aω.2) ∘ fun ω ↦ (f ω, ω)) + apply Measurable.comp _ (by fun_prop) + refine measurable_from_prod_countable_right fun a ↦ ?_ + change Measurable ((fun tω ↦ sumRewards A R a tω.1 tω.2) ∘ fun ω ↦ (g ω, ω)) + apply Measurable.comp _ (by fun_prop) + exact measurable_from_prod_countable_right (fun t ↦ measurable_sumRewards hA hR a t) + +@[fun_prop] +lemma measurable_empMean [MeasurableSingletonClass 𝓐] {R : ℕ → Ω → ℝ} (hA : ∀ n, Measurable (A n)) + (hR : ∀ n, Measurable (R n)) (a : 𝓐) (n : ℕ) : + Measurable (empMean A R a n) := by unfold empMean; fun_prop + +@[fun_prop] +lemma measurable_uncurry_empMean_comp [Countable 𝓐] [MeasurableSingletonClass 𝓐] {R : ℕ → Ω → ℝ} + (hA : ∀ n, Measurable (A n)) (hR : ∀ n, Measurable (R n)) {f : Ω → 𝓐} (hf : Measurable f) + {g : Ω → ℕ} (hg : Measurable g) : + Measurable (fun ω ↦ empMean A R (f ω) (g ω) ω) := by unfold empMean; fun_prop + +@[fun_prop] +lemma measurable_sumRewards' [MeasurableSingletonClass 𝓐] [MeasurableAdd₂ 𝓨] (n : ℕ) (a : 𝓐) : + Measurable (sumRewards' (𝓨 := 𝓨) n · a) := by + simp_rw [sumRewards'] + have h_meas s : Measurable (fun (h : Iic n → 𝓐 × ℝ) ↦ if (h s).1 = a then (h s).2 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact (measurableSet_singleton _).preimage (by fun_prop) + refine Finset.measurable_fun_sum _ fun s hs ↦ ?_ + exact Measurable.ite (by measurability) (by fun_prop) (by fun_prop) + +@[fun_prop] +lemma measurable_uncurry_sumRewards' [MeasurableEq 𝓐] [MeasurableAdd₂ 𝓨] (n : ℕ) : + Measurable (fun p : (Iic n → 𝓐 × 𝓨) × 𝓐 ↦ sumRewards' n p.1 p.2) := by + simp_rw [sumRewards'] + have h_meas s : Measurable (fun p : (Iic n → 𝓐 × ℝ) × 𝓐 ↦ + if (p.1 s).1 = p.2 then (p.1 s).2 else 0) := by + refine Measurable.ite ?_ (by fun_prop) (by fun_prop) + exact measurableSet_eq_fun (by fun_prop) (by fun_prop) + refine Finset.measurable_fun_sum _ fun s hs ↦ ?_ + exact Measurable.ite (by measurability) (by fun_prop) (by fun_prop) + +@[fun_prop] +lemma measurable_empMean' [MeasurableSingletonClass 𝓐] (n : ℕ) (a : 𝓐) : + Measurable (empMean' n · a) := by unfold empMean'; fun_prop + +@[fun_prop] +lemma measurable_uncurry_empMean' [MeasurableEq 𝓐] (n : ℕ) : + Measurable (fun p : (Iic n → 𝓐 × ℝ) × 𝓐 ↦ empMean' n p.1 p.2) := by unfold empMean'; fun_prop + +variable [MeasurableSingletonClass 𝓐] + +lemma IsAlgEnvSeq.isStronglyPredictable_sumRewards {𝓨 : Type*} {_ : MeasurableSpace 𝓨} + [NormedAddCommGroup 𝓨] [OpensMeasurableSpace 𝓨] [SecondCountableTopology 𝓨] + {R : ℕ → Ω → 𝓨} {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} + (h : IsAlgEnvSeq A R alg env P) (a : 𝓐) : + IsStronglyPredictable h.filtration (sumRewards A R a) := by + rw [IsStronglyPredictable.iff_measurable_add_one] + constructor + · simp only [sumRewards_zero] + fun_prop + refine fun n ↦ Finset.stronglyMeasurable_fun_sum _ + fun i hi ↦ (Measurable.ite ?_ ?_ (by fun_prop)).stronglyMeasurable + · refine (measurableSet_singleton a).preimage ?_ + have h_meas_i := h.adapted_action i + simp only [mem_range] at hi + exact h_meas_i.mono (h.filtration.mono (by lia)) le_rfl + · have h_meas_i := h.adapted_feedback i + simp only [mem_range] at hi + exact h_meas_i.mono (h.filtration.mono (by lia)) le_rfl + +lemma IsAlgEnvSeq.stronglyAdapted_sumRewards_add_one {𝓨 : Type*} {_ : MeasurableSpace 𝓨} + [NormedAddCommGroup 𝓨] [OpensMeasurableSpace 𝓨] [SecondCountableTopology 𝓨] + {R : ℕ → Ω → 𝓨} {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} + (h : IsAlgEnvSeq A R alg env P) (a : 𝓐) : + StronglyAdapted h.filtration (fun n ↦ sumRewards A R a (n + 1)) := by + have h_predictable := h.isStronglyPredictable_sumRewards a + rw [IsStronglyPredictable.iff_measurable_add_one] at h_predictable + exact h_predictable.2 + +-- TODO: give a direct proof, without a topology +lemma IsAlgEnvSeq.adapted_sumRewards_add_one {𝓨 : Type*} {_ : MeasurableSpace 𝓨} + [NormedAddCommGroup 𝓨] [BorelSpace 𝓨] [SecondCountableTopology 𝓨] + {R : ℕ → Ω → 𝓨} {alg : Algorithm 𝓐 𝓨} {env : Environment 𝓐 𝓨} + (h : IsAlgEnvSeq A R alg env P) (a : 𝓐) : + Adapted h.filtration (fun n ↦ sumRewards A R a (n + 1)) := + (h.stronglyAdapted_sumRewards_add_one a).adapted + +lemma IsAlgEnvSeq.isStronglyPredictable_empMean {R' : ℕ → Ω → ℝ} + {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} + (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : + IsStronglyPredictable h.filtration (empMean A R' a) := by + unfold empMean + refine StronglyMeasurable.div ?_ ?_ + · exact h.isStronglyPredictable_sumRewards a + · have h_meas := (isStronglyPredictable_pullCount h a).measurable + fun_prop + +lemma IsAlgEnvSeq.stronglyAdapted_empMean_add_one + {R' : ℕ → Ω → ℝ} {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} + (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : + StronglyAdapted h.filtration (fun n ↦ empMean A R' a (n + 1)) := by + have h_predictable := h.isStronglyPredictable_empMean a + rw [IsStronglyPredictable.iff_measurable_add_one] at h_predictable + exact h_predictable.2 + +lemma IsAlgEnvSeq.adapted_empMean_add_one {R' : ℕ → Ω → ℝ} + {alg : Algorithm 𝓐 ℝ} {env : Environment 𝓐 ℝ} + (h : IsAlgEnvSeq A R' alg env P) (a : 𝓐) : + Adapted h.filtration (fun n ↦ empMean A R' a (n + 1)) := + (h.stronglyAdapted_empMean_add_one a).adapted + +end Learning