@@ -6,7 +6,8 @@ Authors: Rémy Degenne
66module
77
88public import RandomDo.Probability.Tactic
9- public import RandomDo.Tactic.Examples
9+ public import RandomDo.Tactic.Elab
10+ public import LeanMachineLearning.ForMathlib.MeasureTheory.Order.MeasurableArg
1011
1112set_option linter.style.header false
1213
@@ -56,6 +57,8 @@ namespace RDo.Thompson
5657
5758variable {K n : ℕ}
5859
60+ attribute[fun_prop] Measurable.ite
61+
5962/-! ## The two stages -/
6063
6164/-- Stage one: fold the history into the per-arm pull counts `N` (started at one) and reward
@@ -96,6 +99,20 @@ instance : IsMarkovKernel (sampleK (K := K)) := by unfold sampleK; infer_instanc
9699
97100@[simp] lemma sampleK_apply (NS : (Fin K → ℝ) × (Fin K → ℝ)) : sampleK NS = sample NS := rfl
98101
102+ def thompson {K n : ℕ} (hK : 0 < K) (hist : Vector (Fin K × ℝ) n) :
103+ Measure (Fin K) := rdo
104+ let mut N : Fin K → ℝ := fun _ ↦ 1
105+ let mut S : Fin K → ℝ := fun _ ↦ 0
106+ for (a, r) in hist rdo
107+ N := fun j ↦ if j = a then N j + 1 else N j
108+ S := fun j ↦ if j = a then S j + r else S j
109+ let mut θ : Fin K → ℝ := fun _ ↦ 0
110+ for j in List.finRange K rdo
111+ let z ← gaussianReal (S j / N j) (Real.toNNReal (1 / N j))
112+ θ := fun k ↦ if k = j then z else θ k
113+ have : Nonempty (Fin K) := Fin.pos_iff_nonempty.mp hK
114+ return argmax θ
115+
99116/-- `thompson` is exactly: fold the history into `(N, S)`, draw the posterior sample `θ` given
100117them, play `argmax θ`. Both sides elaborate to the same two loops; all that separates them is the
101118`return` at the end of each stage. -/
0 commit comments