Skip to content

Commit 3b817db

Browse files
committed
fix for loop over several collections
1 parent 0c2b643 commit 3b817db

3 files changed

Lines changed: 130 additions & 1 deletion

File tree

‎RandomDo/Monad/Examples.lean‎

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ module
22

33
public import RandomDo.Monad.ForInInstances
44
public import RandomDo.Monad.Instances
5+
public import RandomDo.Measurable
56
public import Mathlib.Algebra.Ring.BooleanRing
67
public import Mathlib.Probability.Distributions.Bernoulli
78

@@ -56,4 +57,125 @@ def sampleBitsArray [HasBit m] (n : ℕ) : m (Array Bool) := rdo
5657
xs := xs.push b
5758
return xs
5859

60+
/- # `for` over several collections
61+
62+
A `for` loop over several collections is expanded into a loop over the first one whose body reads
63+
the remaining ones off a `Std.Stream` held in a mutable variable. That body has to stay part of the
64+
surrounding `rdo` block: expanding it into a fresh `rdo` block instead puts the mutable variables
65+
declared before the loop out of scope, and reassigning one of them is then rejected.
66+
67+
The loops below are run in a deterministic monad so that the tests can pin the value the loop
68+
computes, not merely that it elaborates. -/
69+
70+
section MultiCollectionFor
71+
72+
/-- A deterministic `MeasurableSpaceMonad`, so that `rdo` programs reduce to a value. -/
73+
abbrev IdM := Monad.toMeasurableSpaceMonad Id
74+
75+
/-- Read the value out of a deterministic `rdo` program. `IdM α` is definitionally `α`, but it is
76+
not reducibly so, so the tests below go through this to state the value a loop computes. -/
77+
def IdM.run {α : Type u} [MeasurableSpace α] (x : IdM α) : α := x
78+
79+
def zipDot (xs ys : List ℕ) : IdM ℕ := rdo
80+
let mut s := 0
81+
for x in xs, y in ys rdo
82+
s := s + x * y
83+
return s
84+
85+
example : IdM.run (zipDot [1, 2, 3] [10, 20, 30]) = 140 := rfl
86+
87+
/-- Iteration stops with the shorter collection, whichever one that is. -/
88+
example : IdM.run (zipDot [1, 2, 3] [10, 20]) = 50 := rfl
89+
90+
example : IdM.run (zipDot [1, 2] [10, 20, 30]) = 50 := rfl
91+
92+
example : IdM.run (zipDot [] [10, 20]) = 0 := rfl
93+
94+
def zipTriple (xs ys zs : List ℕ) : IdM ℕ := rdo
95+
let mut s := 0
96+
for x in xs, y in ys, z in zs rdo
97+
s := s + x * y * z
98+
return s
99+
100+
example : IdM.run (zipTriple [1, 2] [3, 4] [5, 6]) = 63 := rfl
101+
102+
/-- `break` in the body of a loop over several collections. -/
103+
def zipUntilZero (xs ys : List ℕ) : IdM ℕ := rdo
104+
let mut s := 0
105+
for x in xs, y in ys rdo
106+
if x = 0 then
107+
break
108+
s := s + y
109+
return s
110+
111+
example : IdM.run (zipUntilZero [1, 0, 1] [10, 20, 30]) = 10 := rfl
112+
113+
/-- `continue` in the body of a loop over several collections. -/
114+
def zipSkipZero (xs ys : List ℕ) : IdM ℕ := rdo
115+
let mut s := 0
116+
for x in xs, y in ys rdo
117+
if x = 0 then
118+
continue
119+
s := s + y
120+
return s
121+
122+
example : IdM.run (zipSkipZero [1, 0, 1] [10, 20, 30]) = 40 := rfl
123+
124+
/-- Early `return` out of a loop over several collections. -/
125+
def firstAgreement (xs ys : List ℕ) : IdM (Option ℕ) := rdo
126+
for x in xs, y in ys rdo
127+
if x = y then
128+
return some x
129+
return none
130+
131+
example : IdM.run (firstAgreement [1, 2, 3] [3, 2, 1]) = some 2 := rfl
132+
133+
example : IdM.run (firstAgreement [1, 2] [3, 4]) = none := rfl
134+
135+
/-- The collections a loop ranges over need not have the same type, nor need an `Array` be the
136+
leading one: the collections past the first are iterated through `Std.toStream`, and an `Array`
137+
streams as a `Subarray`, which is measurable.
138+
139+
The two definitions here are checked by elaborating: without `MeasurableSpace (Subarray _)` they
140+
are rejected. Their value is not pinned the way the loops above are, because `Std.Slice`, which
141+
`Subarray` is built from, does not reduce inside a `module` file. -/
142+
def zipMixed (xs : List ℕ) (ys : Array Bool) : IdM ℕ := rdo
143+
let mut s := 0
144+
for x in xs, y in ys rdo
145+
s := s + (if y then x else 0)
146+
return s
147+
148+
def zipArrays (xs ys : Array ℕ) : IdM ℕ := rdo
149+
let mut s := 0
150+
for x in xs, y in ys rdo
151+
s := s + x * y
152+
return s
153+
154+
/-- A `Vector` leading a loop over several collections, where it is consumed by
155+
`MeasurableSpaceForIn'`. -/
156+
def zipVectorFirst (xs : Vector ℕ 3) (ys : List ℕ) : IdM ℕ := rdo
157+
let mut s := 0
158+
for x in xs, y in ys rdo
159+
s := s + x * y
160+
return s
161+
162+
example : IdM.run (zipVectorFirst #v[1, 2, 3] [10, 20, 30]) = 140 := rfl
163+
164+
/-- A `Vector` following one, where it streams as a `Subarray` just as an `Array` does. -/
165+
def zipVectorSecond (xs : List ℕ) (ys : Vector ℕ 3) : IdM ℕ := rdo
166+
let mut s := 0
167+
for x in xs, y in ys rdo
168+
s := s + x * y
169+
return s
170+
171+
/-- The same loop shape, in a genuinely probabilistic monad. -/
172+
noncomputable def zipBernoulli (xs ys : List Bool) : Measure Bool := rdo
173+
let mut acc := false
174+
for x in xs, y in ys rdo
175+
let b ← bernoulliMeasure true false ⟨(1 : ℝ) / 2, by norm_num⟩
176+
acc := acc || (b && x && y)
177+
return acc
178+
179+
end MultiCollectionFor
180+
59181
end

‎RandomDo/Monad/ForInInstances.lean‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,10 +29,17 @@ instance : MeasurableSpace (Array α) :=
2929
instance {n : ℕ} : MeasurableSpace (Vector α n) :=
3030
MeasurableSpace.comap Vector.toArray inferInstance
3131

32+
instance : 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

‎RandomDo/Monad/Notation.lean‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -138,7 +138,7 @@ def rdoForDecl := leading_parser
138138
| none => break
139139
| some ($y, s') =>
140140
$s:ident := s'
141-
rdo $body)
141+
do $body)
142142
doElems := doElems.push (← `(doSeqItem| for%$tk $[$h? : ]? $x:ident in $xs rdo $body))
143143
`(doElem| do $doElems*)
144144
| _ => Macro.throwUnsupported

0 commit comments

Comments
 (0)