Skip to content

Commit db6a226

Browse files
committed
mk_all
2 parents 11eaad3 + b06c1ed commit db6a226

21 files changed

Lines changed: 911 additions & 232 deletions

‎.github/workflows/build.yml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ jobs:
4747
# the single fetched commit.
4848
fetch-depth: 0 # Fetch all history for all branches and tags
4949

50-
- name: Build and lint project
50+
- name: Build, lint and test project
5151
uses: leanprover/lean-action@38fbc41a8c28c4cbaec22d7f7de508ec2e7c0dd9 # v1.5.0
5252
with:
5353
build: true

‎RandomDo.lean‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@ module -- shake: keep-all --deprecated_module: ignore
22

33
public import RandomDo.ForMathlib.MeasureTheory.MeasurableSpace.Embedding
44
public import RandomDo.Measurable
5-
public import RandomDo.Monad.Examples
65
public import RandomDo.Monad.ForInInstances
76
public import RandomDo.Monad.Instances
87
public import RandomDo.Monad.MeasurableSpace
@@ -16,8 +15,8 @@ public import RandomDo.Probability.Tactic
1615
public import RandomDo.Probability.Thompson
1716
public import RandomDo.Probability.Trace
1817
public import RandomDo.Probability.Transfer
18+
public import RandomDo.Tactic.Deriving
1919
public import RandomDo.Tactic.Elab
20-
public import RandomDo.Tactic.Examples
2120
public import RandomDo.Tactic.ForInStep
2221
public import RandomDo.Tactic.IsMarkov
2322
public import RandomDo.Tactic.Lemmas

‎RandomDo/Measurable.lean‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ open MeasureTheory Set
3131

3232
variable {α : Type*} [MeasurableSpace α]
3333

34+
/-- Lists are measurably equivalent to the sigma type of tuples of a given length. -/
3435
def List.measurableEquivSigmaTuple : List α ≃ᵐ Σ n, Fin n → α where
3536
toFun := List.equivSigmaTuple
3637
invFun := List.equivSigmaTuple.symm
@@ -102,6 +103,7 @@ lemma Vector.measurableSpace_eq_comap {n : ℕ} :
102103
(MeasurableSpace.comap List.ofFn inferInstance) := MeasurableSpace.comap_comp.symm
103104
_ = _ := by rw [(measurableEmbedding_ofFn n).comap_eq]
104105

106+
/-- Vectors are equivalent to tuples of a given length. -/
105107
def Vector.measurableEquivTuple {n : ℕ} : Vector α n ≃ᵐ (Fin n → α) where
106108
toFun v := fun i ↦ v[i]
107109
invFun := .ofFn
@@ -115,7 +117,8 @@ def Vector.measurableEquivTuple {n : ℕ} : Vector α n ≃ᵐ (Fin n → α) wh
115117
ext
116118
simp
117119

118-
instance : MeasurableSpace (Option α) := MeasurableSpace.map some inferInstance
120+
instance instMeasurableSpaceOption : MeasurableSpace (Option α) :=
121+
MeasurableSpace.map some inferInstance
119122

120123
theorem measurableSet_option_iff {s : Set (Option α)} :
121124
MeasurableSet s ↔ MeasurableSet (some ⁻¹' s) := Iff.rfl

‎RandomDo/Monad/Examples.lean‎

Lines changed: 0 additions & 59 deletions
This file was deleted.

‎RandomDo/Monad/ForInInstances.lean‎

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -20,19 +20,26 @@ section MeasurableSpace
2020

2121
variable {α : Type*} [MeasurableSpace α]
2222

23-
instance : MeasurableSpace (List α) :=
23+
instance instMeasurableSpaceList : MeasurableSpace (List α) :=
2424
MeasurableSpace.comap List.equivSigmaTuple inferInstance
2525

26-
instance : MeasurableSpace (Array α) :=
26+
instance instMeasurableSpaceArray : MeasurableSpace (Array α) :=
2727
MeasurableSpace.comap Array.toList inferInstance
2828

29-
instance {n : ℕ} : MeasurableSpace (Vector α n) :=
29+
instance instMeasurableSpaceVector {n : ℕ} : MeasurableSpace (Vector α n) :=
3030
MeasurableSpace.comap Vector.toArray inferInstance
3131

32+
instance instMeasurableSpaceSubarray : MeasurableSpace (Subarray α) :=
33+
MeasurableSpace.comap (fun s : Subarray α ↦ s.toList) inferInstance
34+
3235
@[fun_prop]
3336
lemma measurable_toList : Measurable (Array.toList : Array α → List α) :=
3437
Measurable.of_comap_le fun _ a ↦ a
3538

39+
@[fun_prop]
40+
lemma measurable_subarray_toList : Measurable (fun s : Subarray α ↦ s.toList) :=
41+
Measurable.of_comap_le fun _ a ↦ a
42+
3643
@[fun_prop]
3744
lemma measurable_toArray {n : ℕ} : Measurable (Vector.toArray : Vector α n → Array α) :=
3845
Measurable.of_comap_le fun _ a ↦ a
@@ -52,7 +59,7 @@ section Array
5259
{β : Type u} [mβ : MeasurableSpace β]
5360
(as : Array α) (b : β) (f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β :=
5461
let sz := as.usize
55-
let rec @[specialize] loop (i : USize) (b : β) : m β := rdo
62+
let rec @[specialize, nolint docBlame] loop (i : USize) (b : β) : m β := rdo
5663
if i < sz then
5764
let a := as.uget i lcProof
5865
match (← f a lcProof b) with
@@ -67,7 +74,7 @@ section Array
6774
protected def Array.measurableSpaceForIn' [MeasurableSpaceMonad m]
6875
{β : Type u} [mβ : MeasurableSpace β]
6976
(as : Array α) (b : β) (f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β :=
70-
let rec loop (i : Nat) (h : i ≤ as.size) (b : β) : m β := rdo
77+
let rec @[nolint docBlame] loop (i : Nat) (h : i ≤ as.size) (b : β) : m β := rdo
7178
match i, h with
7279
| 0, _ => mPure b
7380
| i+1, h =>
@@ -97,7 +104,7 @@ variable {α β : Type*} [MeasurableSpace α] [MeasurableSpace β] [Ring α]
97104
protected def List.measurableSpaceForIn' [MeasurableSpaceMonad m]
98105
{β : Type u} [mβ : MeasurableSpace β] (as : @& List α) (init : β)
99106
(f : (a : α) → a ∈ as → β → m (ForInStep β)) : m β :=
100-
let rec @[specialize]
107+
let rec @[specialize, nolint docBlame]
101108
loop : (as' : @& List α) → (b : β) → Exists (fun bs => bs ++ as' = as) → m β
102109
| [], b, _ => mPure b
103110
| a::as', b, h => rdo

‎RandomDo/Monad/Instances.lean‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@ universe u v w
2121

2222
/-- A (core) monad automatically defines a (not necessarily lawful) measurable space monad by
2323
forgetting the measurable space argument. -/
24-
def Monad.toMeasurableSpaceMonad (m : Type u → Type v) [Monad m] (α : Type u) [MeasurableSpace α] :
24+
def Monad.toMeasurableSpaceMonad (m : Type u → Type v) (α : Type u) [_mα : MeasurableSpace α] :
2525
Type v := m α
2626

2727
instance {m : Type u → Type v} [Monad m] :
@@ -54,11 +54,15 @@ section RandomM
5454

5555
open Function
5656

57+
/-- A monad for random number generation. -/
5758
structure RandomM (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω)
5859
(α : Type u) [MeasurableSpace α] where
60+
/-- Draws a value from a state of the source of randomness, and hands back the state left for the
61+
next draw. -/
5962
sample : Ω → α × Ω
6063
measurePreserving : MeasurePreserving sample P ((Measure.map (Prod.fst ∘ sample) P).prod P)
6164

65+
/-- TODO -/
6266
abbrev SampleM (Ω : Type w) [MeasurableSpace Ω] (P : Measure Ω) :=
6367
RandomM (ℕ → Ω) (Measure.infinitePi fun _ : ℕ ↦ P)
6468

‎RandomDo/Monad/MeasurableSpace.lean‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,8 @@ The `mPure` function is overloaded via `MeasurableSpacePure` instances.
6565
`MeasurableSpacePure` is typically accessed via `MeasurableSpaceMonad` instances, which extend it.
6666
-/
6767
class MeasurableSpacePure (f : (α : Type u) → [MeasurableSpace α] → Type v) where
68+
/-- Given `a : α` where `α` has a `MeasurableSpace` instance, `mPure a : f α` represents an
69+
action that does nothing and returns a -/
6870
mPure {α : Type u} [MeasurableSpace α] : α → f α
6971

7072
/--
@@ -140,8 +142,6 @@ variable {f : (α : Type u) → [MeasurableSpace α] → Type v} [MeasurableSpac
140142
g₁ <$>ₘ g₀ <$>ₘ x = (fun a => g₁ (g₀ a)) <$>ₘ x :=
141143
(comp_mMap x hg₀ hg₁).symm
142144

143-
@[simp] theorem mMap_unit {a : f PUnit} : (fun _ => PUnit.unit) <$>ₘ a = a := by simp
144-
145145
open MeasurableSpaceBind MeasurableSpacePure MeasurableSpaceFunctor
146146

147147
/-- A `MeasurableSpaceMonad` satisfies the measurable space monad laws. -/

‎RandomDo/Monad/Notation.lean‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -63,6 +63,7 @@ def randOps : DoOps := { DoOps.default with
6363
return mkApp2 m α σ
6464
}
6565

66+
/-- The `do` notation for writing monadic programs depending on a `MeasurableSpace` instance. -/
6667
syntax (name := randKind) "rdo" doSeq : term
6768

6869
/-- Define `rdo` notation elaborator. -/
@@ -138,7 +139,7 @@ def rdoForDecl := leading_parser
138139
| none => break
139140
| some ($y, s') =>
140141
$s:ident := s'
141-
rdo $body)
142+
do $body)
142143
doElems := doElems.push (← `(doSeqItem| for%$tk $[$h? : ]? $x:ident in $xs rdo $body))
143144
`(doElem| do $doElems*)
144145
| _ => Macro.throwUnsupported

‎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. -/

‎RandomDo/Tactic/Deriving.lean‎

Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,79 @@
1+
/-
2+
Copyright (c) 2026 Rémy Degenne. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Rémy Degenne
5+
-/
6+
module
7+
8+
public import RandomDo.Tactic.Elab
9+
10+
/-!
11+
# The `@[is_markov]` attribute
12+
13+
Writing an `rdo` program and then stating that it is a Markov kernel are two separate steps, and
14+
the second is mechanical: it is what `is_markov` does. This attribute runs it at the declaration,
15+
so a program carries its own instance:
16+
17+
```
18+
@[is_markov]
19+
noncomputable def centred (c : ℝ) : Measure ℝ := rdo
20+
let x ← gaussianReal c 1
21+
return x
22+
```
23+
24+
## Which statement is generated
25+
26+
A program's *last* argument is read as the kernel's parameter when it is explicit, giving
27+
`IsMarkov`. Otherwise the program denotes one fixed distribution and the statement is
28+
`IsProbabilityMeasure`. So `centred` above yields `IsMarkov centred`, while a parameterless program
29+
yields `IsProbabilityMeasure` of it, and a program whose trailing arguments are instance-implicit —
30+
`(μ : Measure ℝ) [IsProbabilityMeasure μ]` — yields `IsProbabilityMeasure` of it too, with those
31+
arguments bound.
32+
-/
33+
34+
public meta section
35+
36+
open Lean Meta Elab Term MeasureTheory
37+
38+
namespace RDo.Tactic
39+
40+
/-- The statement to prove for `declName`, as described in the module docstring. -/
41+
def isMarkovStatement (declName : Name) : MetaM Expr := do
42+
let info ← getConstInfo declName
43+
forallTelescope info.type fun args body ↦ do
44+
unless body.isAppOfArity ``MeasureTheory.Measure 2 do
45+
throwError "`IsMarkov` can only be derived for a declaration valued in `Measure`, but \
46+
{declName} is valued in{indentExpr body}"
47+
let f := mkAppN (mkConst declName (info.levelParams.map .param)) args
48+
if h : 0 < args.size then
49+
let last := args[args.size - 1]
50+
if (← last.fvarId!.getDecl).binderInfo.isExplicit then
51+
return ← mkForallFVars args.pop (← mkAppM ``IsMarkov #[← mkLambdaFVars #[last] f])
52+
mkForallFVars args (← mkAppM ``MeasureTheory.IsProbabilityMeasure #[f])
53+
54+
/-- Prove `isMarkovStatement declName` with the `is_markov` tactic and register it as an
55+
instance. -/
56+
def addIsMarkovInstance (declName : Name) : TermElabM Unit := do
57+
let goal ← isMarkovStatement declName
58+
let proof ← Term.elabTerm (← `(by intros; is_markov)) (some goal)
59+
Term.synthesizeSyntheticMVarsNoPostponing
60+
let info ← getConstInfo declName
61+
let instName := declName ++ `isMarkov
62+
addDecl (.thmDecl { name := instName, levelParams := info.levelParams, type := goal,
63+
value := ← instantiateMVars proof })
64+
Meta.addInstance instName .global 1000
65+
66+
/-- The `@[is_markov]` attribute. -/
67+
initialize registerBuiltinAttribute {
68+
name := `is_markov
69+
descr := "prove that this `rdo` program is a Markov kernel, and register it as an instance"
70+
applicationTime := .afterCompilation
71+
add := fun declName _stx kind ↦ do
72+
unless kind == AttributeKind.global do
73+
throwError "`is_markov` must be a global attribute"
74+
(addIsMarkovInstance declName).run'.run'
75+
}
76+
77+
end RDo.Tactic
78+
79+
end

0 commit comments

Comments
 (0)