Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions Test.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
module -- shake: keep-all --deprecated_module: ignore

public import Test.Bind
public import Test.Common
public import Test.Control
public import Test.Gaps
public import Test.Instances
public import Test.IsMarkov
public import Test.Loops
public import Test.MonadLaws
84 changes: 84 additions & 0 deletions Test/Bind.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
module

public import Test.Common

set_option linter.style.header false

/-!
# `rdo`: binding, sequencing and `return`

`rdo` elaborates to `MeasurableSpacePure.mPure` and `MeasurableSpaceBind.mBind` rather than to
`pure` and `bind`, so every shape a `do` block can take has to be re-checked here.
-/

open MeasureTheory ProbabilityTheory

@[expose] public section

namespace Test.Bind

/-- A single bind. -/
def one (xs : List ℕ) : IdM ℕ := rdo
let x ← (xs.headD 0 : IdM ℕ)
return x + 1

example : IdM.run (one [4, 5]) = 5 := rfl

/-- A chain of binds. -/
def chain (a b : ℕ) : IdM ℕ := rdo
let x ← (a : IdM ℕ)
let y ← (b : IdM ℕ)
return x * y

example : IdM.run (chain 6 7) = 42 := rfl

/-- `←` nested inside a larger expression is lifted out of it. -/
def nested (a b : ℕ) : IdM ℕ := rdo
return (← (a : IdM ℕ)) + (← (b : IdM ℕ)) * 2

example : IdM.run (nested 3 4) = 11 := rfl

/-- A pure `let`, and a `have`, alongside the monadic ones. -/
def mixedLets (a : ℕ) : IdM ℕ := rdo
let x ← (a : IdM ℕ)
let y := x + 1
have : 0 < y + 1 := Nat.succ_pos y
return y * 2

example : IdM.run (mixedLets 5) = 12 := rfl

/-- Destructuring a bound pair. -/
def destructure (p : ℕ × ℕ) : IdM ℕ := rdo
let (a, b) ← (p : IdM (ℕ × ℕ))
return a + b

example : IdM.run (destructure (3, 4)) = 7 := rfl

/-- A statement in the middle of a block is sequenced, not dropped. -/
def sequenced (a : ℕ) : IdM ℕ := rdo
let mut s := 0
(pure () : IdM PUnit)
s := s + a
return s

example : IdM.run (sequenced 9) = 9 := rfl

/-- `return` in straight-line code drops the rest of the block. -/
def earlyReturn (a : ℕ) : IdM ℕ := rdo
if a = 0 then
return 100
return a

example : IdM.run (earlyReturn 0) = 100 := rfl

example : IdM.run (earlyReturn 7) = 7 := rfl

/-- The same shapes at `Measure`, where the binds are genuine integrals. -/
noncomputable def twoCoins : Measure Bool := rdo
let x ← fairCoin
let y ← fairCoin
return x && y

end Test.Bind

end
40 changes: 40 additions & 0 deletions Test/Common.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
module

public import RandomDo

set_option linter.style.header false

/-!
# Shared scaffolding for the `rdo` test suite

The files in this library exercise the features of `rdo` that are currently implemented. They are
`module` files, like the library they test.

Where a program reduces, a test states the value it computes rather than only that it elaborates.
One family does not reduce here: a `for` loop over several collections streams the ones past
the first through `Std.Stream`, and for an `Array` or a `Vector` that goes through the
`Array → Subarray` conversion, which core marks `@[no_expose]`. Those are checked by elaborating.
-/

open MeasureTheory ProbabilityTheory

@[expose] public section

universe u

/-- A deterministic `MeasurableSpaceMonad`. An `rdo` program written at `IdM` denotes a value, so
the tests can state what a program computes and not merely that it typechecks. -/
abbrev IdM := Monad.toMeasurableSpaceMonad Id

/-- Read the value out of a deterministic `rdo` program. `IdM α` is definitionally `α` but not
reducibly so, which is what otherwise stops numerals and `rfl` from seeing through it. -/
def IdM.run {α : Type u} [MeasurableSpace α] (x : IdM α) : α := x

/-- The fair coin, as a probability measure on `Bool`. -/
noncomputable def fairCoin : Measure Bool := bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩

instance : IsProbabilityMeasure fairCoin := by
unfold fairCoin
infer_instance

end
122 changes: 122 additions & 0 deletions Test/Control.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
module

public import Test.Common

set_option linter.style.header false

/-!
# `rdo`: branching

Every branching `do` element goes through core's elaborators, which build the block out of the
`DoOps` that `rdo` supplies. Each shape is checked here to make sure the substituted `mPure` and
`mBind` reach it.
-/

open MeasureTheory ProbabilityTheory

@[expose] public section

namespace Test.Control

/-- `if … then … else`, branching the whole rest of the block. -/
def branch (n : ℕ) : IdM ℕ := rdo
if n = 0 then
return 100
else
return n

example : IdM.run (branch 0) = 100 := rfl

example : IdM.run (branch 5) = 5 := rfl

/-- `if` with no `else`: the block carries on afterwards. -/
def clampZero (n : ℕ) : IdM ℕ := rdo
let mut s := n
if n = 0 then
s := 100
return s

example : IdM.run (clampZero 0) = 100 := rfl

example : IdM.run (clampZero 5) = 5 := rfl

/-- Nested branches. -/
def nestedIf (a b : ℕ) : IdM ℕ := rdo
if a = 0 then
if b = 0 then
return 0
else
return 1
else
return 2

example : IdM.run (nestedIf 0 0) = 0 := rfl

example : IdM.run (nestedIf 0 1) = 1 := rfl

example : IdM.run (nestedIf 1 0) = 2 := rfl

/-- A dependent `if`, whose branch uses the proof it introduces. -/
def headOr (xs : List ℕ) : IdM ℕ := rdo
if h : 0 < xs.length then
return xs[0]'h
else
return 0

example : IdM.run (headOr [7, 8]) = 7 := rfl

example : IdM.run (headOr []) = 0 := rfl

/-- A `match` on a value bound by `←`. -/
def matchArrow (o : Option ℕ) : IdM ℕ := rdo
match ← (o : IdM (Option ℕ)) with
| none => return 0
| some n => return n + 1

example : IdM.run (matchArrow (some 4)) = 5 := rfl

example : IdM.run (matchArrow none) = 0 := rfl

/-- A pure `match` inside the block. -/
def matchPure (o : Option ℕ) : IdM ℕ := rdo
let v ← (o : IdM (Option ℕ))
match v with
| none => return 0
| some n => return n + 1

example : IdM.run (matchPure (some 4)) = 5 := rfl

/-- `if let`. -/
def ifLet (o : Option ℕ) : IdM ℕ := rdo
let v ← (o : IdM (Option ℕ))
if let some n := v then
return n + 1
else
return 0

example : IdM.run (ifLet (some 4)) = 5 := rfl

example : IdM.run (ifLet none) = 0 := rfl

/-- `unless`. -/
def unlessFlag (b : Bool) : IdM ℕ := rdo
let mut s := 0
unless b do
s := 1
return s

example : IdM.run (unlessFlag true) = 0 := rfl

example : IdM.run (unlessFlag false) = 1 := rfl

/-- Branching on a value drawn from a genuine distribution. -/
noncomputable def fairCoinBranch : Measure ℕ := rdo
let b ← fairCoin
if b then
return 1
else
return 0

end Test.Control

end
Loading
Loading