Skip to content

Commit 718ea04

Browse files
committed
Nested loops
1 parent 3a835f9 commit 718ea04

3 files changed

Lines changed: 104 additions & 23 deletions

File tree

‎RandomDo/Monad/Notation.lean‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -260,6 +260,11 @@ def rdoForDecl := leading_parser
260260
dec.continueWithUnit
261261
mkBindApp σ γ forIn rest
262262

263+
/-- Infer the `ControlInfo` of an `rdo` loop as that of the core `for` loop with the same body. -/
264+
@[doElem_control_info rdoFor] def controlInfoRDoFor : ControlInfoHandler := fun stx => do
265+
let `(rdoFor| for $_:rdoForDecl,* rdo $body) := stx | throwUnsupportedSyntax
266+
inferControlInfoElem (← `(doElem| for _ in #[()] do $body))
267+
263268
end LoopElab
264269

265270
end RDo

‎Test/Gaps.lean‎

Lines changed: 0 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -18,26 +18,6 @@ open MeasureTheory ProbabilityTheory
1818

1919
namespace Test.Gaps
2020

21-
/-! ## Nested loops
22-
23-
TODO: register a `ControlInfo` inference handler for `RDo.rdoFor`, mirroring the rule core states
24-
inline for `doFor` in `Lean/Elab/Do/InferControlInfo.lean`.
25-
-/
26-
27-
/--
28-
error: No `ControlInfo` inference handler found for `RDo.rdoFor` in syntax
29-
for y in ys rdo
30-
s := s + x * y
31-
Register a handler with `@[doElem_control_info RDo.rdoFor]`.
32-
-/
33-
#guard_msgs (whitespace := lax) in
34-
def nestedLoops (xs ys : List ℕ) : IdM ℕ := rdo
35-
let mut s := 0
36-
for x in xs rdo
37-
for y in ys rdo
38-
s := s + x * y
39-
return s
40-
4121
/-! ## Unbounded and conditional iteration
4222
4323
TODO: `while`, `repeat` and `repeat … until` all expand to `for _ in Loop.mk do …`, which reaches

‎Test/Loops.lean‎

Lines changed: 99 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -10,9 +10,9 @@ set_option linter.style.header false
1010
`rdo` has its own `for … rdo …` parser, expander and elaborator, mirroring core's but emitting
1111
`MeasurableSpaceForIn.forIn`. Instances exist for `List`, `Array` and `Vector`.
1212
13-
There is no test for a loop nested inside another: `rdoFor` has no registered `ControlInfo`
14-
inference handler, so the outer loop cannot work out what the inner one does to the control flow,
15-
and such a program is rejected before elaboration.
13+
A loop can sit under another construct, including another loop: the enclosing one learns what the
14+
loop does to the control flow from the `ControlInfo` handler of `rdoFor`, which is that of core's
15+
`for` loop with the same body.
1616
-/
1717

1818
open MeasureTheory ProbabilityTheory
@@ -139,6 +139,102 @@ noncomputable def countHeads (n : ℕ) : Measure ℕ := rdo
139139
c := c + 1
140140
return c
141141

142+
/-! ## Loops under other constructs -/
143+
144+
/-- A loop nested inside another, reassigning a variable of the enclosing block. -/
145+
def nestedLoops (xs ys : List ℕ) : IdM ℕ := rdo
146+
let mut s := 0
147+
for x in xs rdo
148+
for y in ys rdo
149+
s := s + x * y
150+
return s
151+
152+
example : IdM.run (nestedLoops [1, 2] [3, 4]) = 21 := rfl
153+
154+
example : IdM.run (nestedLoops [1, 2] []) = 0 := rfl
155+
156+
/-- `break` in the inner loop leaves the inner loop only. -/
157+
def innerBreak (xs ys : List ℕ) : IdM ℕ := rdo
158+
let mut s := 0
159+
for x in xs rdo
160+
for y in ys rdo
161+
if y = 0 then
162+
break
163+
s := s + x * y
164+
return s
165+
166+
example : IdM.run (innerBreak [1, 2] [3, 0, 5]) = 9 := rfl
167+
168+
/-- `continue` in the inner loop skips to the next inner iteration. -/
169+
def innerContinue (xs ys : List ℕ) : IdM ℕ := rdo
170+
let mut s := 0
171+
for x in xs rdo
172+
for y in ys rdo
173+
if y = 0 then
174+
continue
175+
s := s + x * y
176+
return s
177+
178+
example : IdM.run (innerContinue [1, 2] [3, 0, 5]) = 24 := rfl
179+
180+
/-- An early `return` in the inner loop leaves the whole program. -/
181+
def firstProductOver (xs ys : List ℕ) (limit : ℕ) : IdM ℕ := rdo
182+
for x in xs rdo
183+
for y in ys rdo
184+
if x * y > limit then
185+
return x * y
186+
return 0
187+
188+
example : IdM.run (firstProductOver [1, 2, 3] [1, 2] 3) = 4 := rfl
189+
190+
example : IdM.run (firstProductOver [1, 2] [1, 2] 10) = 0 := rfl
191+
192+
/-- An inner loop over several collections, which the expander rewrites first. -/
193+
def nestedZip (xs ys zs : List ℕ) : IdM ℕ := rdo
194+
let mut s := 0
195+
for x in xs rdo
196+
for y in ys, z in zs rdo
197+
s := s + x * y * z
198+
return s
199+
200+
example : IdM.run (nestedZip [1, 2] [1, 2] [3, 4]) = 33 := rfl
201+
202+
/-- A loop in a branch of an `if`. -/
203+
def sumIf (b : Bool) (xs : List ℕ) : IdM ℕ := rdo
204+
let mut s := 0
205+
if b then
206+
for x in xs rdo
207+
s := s + x
208+
return s
209+
210+
example : IdM.run (sumIf true [1, 2, 3]) = 6 := rfl
211+
212+
example : IdM.run (sumIf false [1, 2, 3]) = 0 := rfl
213+
214+
/-- A loop in an arm of a `match`. -/
215+
def sumHead (xss : List (List ℕ)) : IdM ℕ := rdo
216+
let mut s := 0
217+
match xss with
218+
| [] => pure ()
219+
| xs :: _ =>
220+
for x in xs rdo
221+
s := s + x
222+
return s
223+
224+
example : IdM.run (sumHead [[1, 2], [10]]) = 3 := rfl
225+
226+
example : IdM.run (sumHead []) = 0 := rfl
227+
228+
/-- Nested loops whose body binds monadically, at `Measure`. -/
229+
noncomputable def countPairsOfHeads (n : ℕ) : Measure ℕ := rdo
230+
let mut c := 0
231+
for _ in List.range n rdo
232+
for _ in List.range n rdo
233+
let b ← fairCoin
234+
if b then
235+
c := c + 1
236+
return c
237+
142238
end Test.Loops
143239

144240
end

0 commit comments

Comments
 (0)