diff --git a/LeanMachineLearning.lean b/LeanMachineLearning.lean index ef058a9c..71a759e8 100644 --- a/LeanMachineLearning.lean +++ b/LeanMachineLearning.lean @@ -50,6 +50,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.ReinforcementLearning.MDP.Basic public import LeanMachineLearning.SequentialLearning.ActionIndicator public import LeanMachineLearning.SequentialLearning.Algorithm public import LeanMachineLearning.SequentialLearning.AlgorithmDensity diff --git a/LeanMachineLearning/ReinforcementLearning/MDP/Basic.lean b/LeanMachineLearning/ReinforcementLearning/MDP/Basic.lean new file mode 100644 index 00000000..faebfff4 --- /dev/null +++ b/LeanMachineLearning/ReinforcementLearning/MDP/Basic.lean @@ -0,0 +1,92 @@ +/- +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.IonescuTulceaSpace +public import LeanMachineLearning.SequentialLearning.Deterministic + +/-! +# Markov decision processes + +-/ + +@[expose] public section + +open MeasureTheory ProbabilityTheory Learning + +/-- Markov decision process with state space `𝓒`, action space `𝓐`, and reward space `𝓑`, described +by a transition kernel `P : Kernel (𝓒 Γ— 𝓐) 𝓒` and a reward kernel `R : Kernel (𝓒 Γ— 𝓐) 𝓑`. +See `MDP.env` for the environment associated with the MDP. -/ +structure MDP (𝓒 𝓐 𝓑 : Type*) [MeasurableSpace 𝓒] [MeasurableSpace 𝓐] [MeasurableSpace 𝓑] where + P : Kernel (𝓒 Γ— 𝓐) 𝓒 + [hP : IsMarkovKernel P] + R : Kernel (𝓒 Γ— 𝓐) 𝓑 + [hR : IsMarkovKernel R] + +namespace Learning.MDP + +variable {𝓒 𝓐 𝓑 : Type*} {m𝓒 : MeasurableSpace 𝓒} {m𝓐 : MeasurableSpace 𝓐} {m𝓑 : MeasurableSpace 𝓑} + +instance (M : MDP 𝓒 𝓐 𝓑) : IsMarkovKernel M.P := M.hP +instance (M : MDP 𝓒 𝓐 𝓑) : IsMarkovKernel M.R := M.hR + +/-! ### The environment -/ + +open Classical in +protected noncomputable def env [h𝓒 : Nonempty 𝓒] (M : MDP 𝓒 𝓐 𝓑) (ΞΌβ‚€ : Measure 𝓒) : + Environment 𝓒 𝓐 𝓑 where + obs + | 0 => Kernel.const _ (if IsProbabilityMeasure ΞΌβ‚€ then ΞΌβ‚€ else Measure.dirac h𝓒.some) + | n + 1 => M.P.comap (fun h ↦ ((h (Fin.last n)).obs, (h (Fin.last n)).action)) (by fun_prop) + feedback := fun _ ↦ M.R.comap (fun p ↦ (p.1.2, p.2)) (by fun_prop) + isMarkovKernel_obs n := by + cases n + Β· split_ifs <;> infer_instance + Β· infer_instance + +lemma measurable_env [Nonempty 𝓒] (M : MDP 𝓒 𝓐 𝓑) : Measurable M.env := by + rw [measurable_environment_iff] + refine fun n ↦ ⟨?_, by fun_prop⟩ + cases n with + | zero => + refine Kernel.measurable_const.comp ?_ + exact Measurable.ite ProbabilityMeasure.measurableSet_isProbabilityMeasure + (by fun_prop) (by fun_prop) + | succ n => fun_prop + +/-! ### Stationary policies and their trajectory laws -/ + +-- todo: this is a generic definition of a stationary policy, not specific to MDPs. +/-- The stationary deterministic policy `Ο€` as an algorithm: it plays `Ο€ s` in the current +state `s`. -/ +noncomputable def policyAlg (Ο€ : 𝓒 β†’ 𝓐) (hΟ€ : Measurable Ο€) : Algorithm 𝓒 𝓐 𝓑 := + detAlgorithm (fun _ p ↦ Ο€ p.2) fun _ ↦ hΟ€.comp measurable_snd + +variable [Nonempty 𝓒] + +/-- The law of the trajectory of the policy `Ο€` started at the state `s`: `E^Ο€_s`. -/ +noncomputable def policyMeasure (M : MDP 𝓒 𝓐 𝓑) (Ο€ : 𝓒 β†’ 𝓐) (hΟ€ : Measurable Ο€) (s : 𝓒) : + Measure (β„• β†’ Round 𝓒 𝓐 𝓑) := + trajMeasure (policyAlg Ο€ hΟ€) (M.env (Measure.dirac s)) +deriving IsProbabilityMeasure + +/-- The law of the state-action trajectory of the policy `Ο€` from the state `s` (no rewards). -/ +noncomputable def stateLaw (M : MDP 𝓒 𝓐 𝓑) (Ο€ : 𝓒 β†’ 𝓐) (hΟ€ : Measurable Ο€) (s : 𝓒) : + Measure (β„• β†’ 𝓒 Γ— 𝓐) := + (policyMeasure M Ο€ hΟ€ s).map (fun h n ↦ ((h n).obs, (h n).action)) + +instance (M : MDP 𝓒 𝓐 𝓑) (Ο€ : 𝓒 β†’ 𝓐) (hΟ€ : Measurable Ο€) (s : 𝓒) : + IsProbabilityMeasure (stateLaw M Ο€ hΟ€ s) := by unfold stateLaw; infer_instance + +/-! ### Mean rewards -/ + + +variable [NormedAddCommGroup 𝓑] [NormedSpace ℝ 𝓑] + +/-- The mean reward `r(s, a) = 𝔼[R (s, a)]` of a state-action pair. -/ +noncomputable def meanReward (R : Kernel (𝓒 Γ— 𝓐) 𝓑) [IsMarkovKernel R] (p : 𝓒 Γ— 𝓐) : 𝓑 := (R p)[id] + +end Learning.MDP