Skip to content

Commit 0cf260d

Browse files
committed
Cleanup and tests for polymorphic programs
1 parent 0f3af93 commit 0cf260d

8 files changed

Lines changed: 125 additions & 49 deletions

File tree

‎Compute.lean‎

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

‎Polymorphic.lean‎

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

‎RandomDo.lean‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,8 @@ public import RandomDo.Tactic.Computable.Counterparts
1616
public import RandomDo.Tactic.Computable.Defs
1717
public import RandomDo.Tactic.Computable.Deriving
1818
public import RandomDo.Tactic.Computable.Example
19-
public import RandomDo.Tactic.Computable.Polymorphic
19+
public import RandomDo.Tactic.Computable.Polymorphic.Polymorphic
20+
public import RandomDo.Tactic.Computable.Polymorphic.Scalar
2021
public import RandomDo.Tactic.IsMarkov.Defs
2122
public import RandomDo.Tactic.IsMarkov.Deriving
2223
public import RandomDo.Tactic.IsMarkov.Elab

‎RandomDo/NumLean/Distributions.lean‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -145,8 +145,8 @@ deviation. -/
145145
/-- Draw samples from a binomial distribution. -/
146146
def binomial (n : Nat) (p : Float) : RandPCG IO Nat := do
147147
-- The comparisons are the `Bool` ones: through `Decidable`, each costs more than a draw.
148-
if p.lt 0.0 || Float.lt 1.0 p || p.isNaN then
149-
throw <| IO.userError "p < 0, p > 1 or p is NaN"
148+
if p < 0.0 || 1.0 < p || p.isNaN then
149+
throw <| IO.userError s!"expected 0 ≤ p ≤ 1, got {p}"
150150
let n := n.toUInt64.toFloat
151151
if n == 0 || p == 0.0 then return 0
152152
if Float.le p 0.5 then

RandomDo/Tactic/Computable/Polymorphic.lean renamed to RandomDo/Tactic/Computable/Polymorphic/Polymorphic.lean

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,7 @@ noncomputable instance : HasBernoulli Measure ℝ where
5555

5656
/-- A probability outside `[0, 1]` is clamped to it, as in the instance on `Measure`. -/
5757
instance : HasBernoulli RandM Float where
58-
bernoulli p := show RandPCG IO Bool from return (← bernoulli (max 0 (min 1 p))) == 1
58+
bernoulli p := return (← NumLean.bernoulli p) == 1
5959

6060
/-- A typeclass for scalars with an exponential. -/
6161
class HasExp (R : Type) where

‎Test.lean‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,4 +9,5 @@ public import Test.Instances
99
public import Test.IsMarkov
1010
public import Test.Loops
1111
public import Test.MonadLaws
12+
public import Test.Polymorphic
1213
public import Test.Sample

‎Test/Polymorphic.lean‎

Lines changed: 119 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,119 @@
1+
module
2+
3+
public import Test.IsMarkov
4+
public meta import RandomDo
5+
import Batteries.Data.Float.Basic
6+
7+
set_option linter.style.header false
8+
9+
/-!
10+
# Polymorphic `rdo` programs
11+
12+
The programs of `Test.Computable`, written once over an arbitrary `MeasurableSpaceMonad` `m` and
13+
drawing through `HasGaussian` and `HasBernoulli`. Read at `m := Measure`, each one is a probability
14+
measure, checked by `is_markov`, and is the program of `Test.IsMarkov` when there is one. Run at
15+
`m := RandM`, it samples.
16+
-/
17+
18+
@[expose] public section
19+
20+
namespace Test.Polymorphic
21+
22+
open Test.IsMarkov NumLean Lean.Elab.Command MeasureTheory ProbabilityTheory
23+
24+
universe v
25+
26+
variable {m : (α : Type) → [MeasurableSpace α] → Type v} [MeasurableSpaceMonad m]
27+
{R V : Type} [Scalar R] [Scalar V]
28+
29+
def logPolymorphic {α : Type} [MeasurableSpace α] [Lean.ToMessageData α] (prog : RandM α) :
30+
CommandElabM Unit := do
31+
let x ← (IO.runRandPCG prog : IO α)
32+
let y ← (IO.runRandPCGWith 42 prog : IO α)
33+
Lean.logInfo m!"x = {x}"
34+
Lean.logInfo m!"y (seed 42) = {y}"
35+
36+
def sumTwo [HasGaussian m R V R] : m R := rdo
37+
let x ← HasGaussian.gaussian (m := m) (0 : R) (1 : V)
38+
let y ← HasGaussian.gaussian (m := m) (0 : R) (1 : V)
39+
return x + y
40+
41+
example : sumTwo (m := Measure) (R := ℝ) (V := NNReal) = Test.IsMarkov.sumTwo := rfl
42+
43+
run_cmd logPolymorphic (sumTwo (m := RandM) (R := Float) (V := Float))
44+
45+
def unfoldSumTwo [HasGaussian m R V R] : m R := rdo
46+
let y ← sumTwo (m := m) (V := V)
47+
let x ← HasGaussian.gaussian (m := m) (0 : R) (1 : V)
48+
return x + y
49+
50+
example : IsProbabilityMeasure (unfoldSumTwo (m := Measure) (R := ℝ) (V := NNReal)) := by
51+
is_markov
52+
53+
run_cmd logPolymorphic (unfoldSumTwo (m := RandM) (R := Float) (V := Float))
54+
55+
def centred [HasGaussian m R V R] (c : R) : m R := rdo
56+
let x ← HasGaussian.gaussian (m := m) c (1 : V)
57+
return x
58+
59+
example : centred (m := Measure) (R := ℝ) (V := NNReal) = Test.IsMarkov.centred := rfl
60+
61+
run_cmd logPolymorphic (centred (m := RandM) (R := Float) (V := Float) 20)
62+
63+
def branchOn [LT R] [DecidableLT R] [HasGaussian m R V R] (c : R) : m R := rdo
64+
if 0 < c then
65+
let x ← HasGaussian.gaussian (m := m) c (1 : V)
66+
return x
67+
else
68+
let x ← HasGaussian.gaussian (m := m) (0 : R) (1 : V)
69+
return x
70+
71+
example : branchOn (m := Measure) (R := ℝ) (V := NNReal) = Test.IsMarkov.branchOn := rfl
72+
73+
run_cmd logPolymorphic (branchOn (m := RandM) (R := Float) (V := Float) 20)
74+
75+
run_cmd logPolymorphic (branchOn (m := RandM) (R := Float) (V := Float) (-1))
76+
77+
def coin [HasBernoulli m R] (p : R) : m Bool := rdo
78+
let b ← HasBernoulli.bernoulli (m := m) p
79+
return b
80+
81+
example : IsProbabilityMeasure (coin (m := Measure) (1 / 2 : ℝ)) := by is_markov
82+
83+
run_cmd logPolymorphic (coin (m := RandM) (0.5 : Float))
84+
85+
/--
86+
error: expected 0 ≤ p ≤ 1, got 2.000000
87+
-/
88+
#guard_msgs in
89+
run_cmd logPolymorphic (coin (m := RandM) (2 : Float))
90+
91+
/--
92+
error: expected 0 ≤ p ≤ 1, got -1.000000
93+
-/
94+
#guard_msgs in
95+
run_cmd logPolymorphic (coin (m := RandM) (-1 : Float))
96+
97+
def twoCoins [HasBernoulli m R] (p : R) : m Bool := rdo
98+
let x ← coin (m := m) p
99+
let y ← coin (m := m) p
100+
return x && y
101+
102+
example : IsProbabilityMeasure (twoCoins (m := Measure) (1 / 2 : ℝ)) := by is_markov
103+
104+
run_cmd logPolymorphic (twoCoins (m := RandM) (0.5 : Float))
105+
106+
def ex1 [HasGaussian m R V R] : m R := rdo
107+
let mut x : R := 0
108+
for _ in List.range 1000 rdo
109+
let y ← HasGaussian.gaussian (m := m) (0 : R) (1 : V)
110+
x := x + y
111+
return x
112+
113+
example : IsProbabilityMeasure (ex1 (m := Measure) (R := ℝ) (V := NNReal)) := by is_markov
114+
115+
run_cmd logPolymorphic (ex1 (m := RandM) (R := Float) (V := Float))
116+
117+
end Test.Polymorphic
118+
119+
end

‎lakefile.toml‎

Lines changed: 0 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -29,11 +29,3 @@ name = "Test"
2929
name = "dump"
3030
root = "Dump"
3131
srcDir = "test_data"
32-
33-
[[lean_exe]]
34-
name = "compute"
35-
root = "Compute"
36-
37-
[[lean_exe]]
38-
name = "polymorphic"
39-
root = "Polymorphic"

0 commit comments

Comments
 (0)