Skip to content

Commit b06c1ed

Browse files
committed
Fix Thompson.lean
1 parent c876e38 commit b06c1ed

1 file changed

Lines changed: 18 additions & 1 deletion

File tree

‎RandomDo/Probability/Thompson.lean‎

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,8 @@ Authors: Rémy Degenne
66
module
77

88
public 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

1112
set_option linter.style.header false
1213

@@ -56,6 +57,8 @@ namespace RDo.Thompson
5657

5758
variable {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
100117
them, 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

Comments
 (0)