Skip to content

Commit 33f7a5b

Browse files
author
Cong Zhang
committed
[cl] Support nested conditionals and break_loop in task-scheduling domain loops
1 parent 66964ac commit 33f7a5b

9 files changed

Lines changed: 1284 additions & 31 deletions

File tree

‎experimental/task_scheduling/README.md‎

Lines changed: 76 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,82 @@ The manager also computes a warp-group-rounded initial register budget and threa
3030

3131
Automatic value routing hides storage positions from callback authors. Each callback receives its routed inputs after the `StageInfo` argument and returns the values declared by its captured work method. The frozen `DeviceStep` performs the corresponding value-stack operations, while `DeviceTask.make_context(tasks_inputs)` creates the immutable execution context. The tuple representation remains an internal compiler detail; no placeholder or named routing slot is exposed to resource authors.
3232

33-
`outputs=N` creates `N` independent lexical `ScheduleValue` instances. A work call returns one value directly or a tuple that can be unpacked normally, and its device callback returns the same number of runtime values. A value produced inside a conditional or domain loop cannot escape that scope accidentally. A domain loop(don't support while senmantic and contional early break) that needs state across iterations uses `domain_loop(start, end, step, body, *initial_values)`, with one callback parameter per initial value. The schedule builder invokes `body(*iter_values)` once while capturing the loop body, so every resource call remains a schedule step. The callback must return one work-call output per input; those outputs become both the next-iteration values and the post-loop results. The device adapter preserves zero-trip pass-through and materializes the backedge without exposing routing positions to callbacks. Work methods read the current device offset from `stage_info.loop_offset`; it is not a callback parameter or a routed value. Loop bodies may use `first_iter()`, `last_iter()`, and `every(period, start=...)` to capture iteration-guarded regions.
33+
`outputs=N` creates `N` independent lexical `ScheduleValue` instances. A work call returns one value directly or a tuple that can be unpacked normally, and its device callback returns the same number of runtime values. A value produced inside a conditional or domain loop cannot escape that scope accidentally. A domain loop that needs state across iterations uses `domain_loop(start, end, step, body, *initial_values)`, with one callback parameter per initial value. The schedule builder invokes `body(*iter_values)` once while capturing the loop body, so every resource call remains a schedule step. The callback must return one work-call output per input; those outputs become both the next-iteration values and the post-loop results. The device adapter preserves zero-trip pass-through and materializes the backedge without exposing routing positions to callbacks. Work methods read the current device offset from `stage_info.loop_offset`; it is not a callback parameter or a routed value. Loop bodies may use `first_iter()`, `last_iter()`, and `every(period, start=...)` to capture iteration-guarded regions.
34+
35+
`when_true()` and `when_false()` regions may nest, including inside iteration
36+
guards. An inner region executes only when every enclosing guard is active.
37+
Predicates and other work outputs produced inside a branch may be used by its
38+
nested regions, but cannot escape to a parent or sibling scope. Reusing the same
39+
routed predicate preserves true/false correlation in host analysis.
40+
41+
The runnable `tutorial/01_copy_basics/04_copy_tma_nested_conditional.py` example
42+
extends the conditional TMA copy with all four two-level true/false combinations
43+
and a third-level condition. It checks the copied tensor and exact per-row branch
44+
markers, including the absence of writes from inactive branches:
45+
46+
```bash
47+
python experimental/task_scheduling/tutorial/01_copy_basics/04_copy_tma_nested_conditional.py --rows-cols 9,512
48+
```
49+
50+
`break_loop()` exits the enclosing domain loop immediately, including from a
51+
nested conditional. It skips the rest of that iteration and resumes after the
52+
loop. The context-manager handle provides the same operation:
53+
54+
```python
55+
with ts.domain_loop(num_tiles) as loop:
56+
done = resource.is_done()
57+
with ts.when_true(done):
58+
loop.break_loop()
59+
resource.process()
60+
```
61+
62+
Use `ts.break_loop()` in a functional loop body. A bare break skips the normal
63+
backedge: returned loop-carried results retain their values from the last
64+
completed iteration (or the initial values if the first iteration breaks).
65+
Pass explicit results to return values computed during the interrupted iteration:
66+
67+
```python
68+
initial = resource.init_state()
69+
70+
def body(state):
71+
updated = resource.advance_state(state)
72+
done = resource.is_done(updated)
73+
with ts.when_true(done):
74+
ts.break_loop(updated)
75+
return updated
76+
77+
result = ts.domain_loop(0, num_tiles, 1, body, initial)
78+
resource.consume_state(result)
79+
```
80+
81+
For multiple carried inputs, use `ts.break_loop(next_count, next_total)` in the
82+
same order as the functional loop's inputs. Supply all carried results or none.
83+
Each explicit argument must be a visible routed `ScheduleValue` with the same
84+
pipeline-stage provenance and a compatible CUDA type as its carried input;
85+
ordinary constants must first be returned by a work method. Values created
86+
inside a nested branch may be returned directly by a break in that branch.
87+
The exit values are captured before branch-local routes are discarded. Wrong
88+
arity, foreign-schedule values, and values escaping a sibling/child scope are
89+
rejected during capture; CUDA types are checked during device compilation.
90+
An untaken break uses the normal body return values, and a zero-trip loop still
91+
returns its initial values.
92+
93+
Side effects and pipeline-state updates performed before the break are retained.
94+
A break exits only the domain loop, so an enclosing `work_tile_loop()` continues.
95+
`last_iter()` still means the last iteration of the declared range; it is not an
96+
exit hook. Breaks outside a domain loop and expired loop handles are rejected.
97+
Unbounded while loops remain unsupported.
98+
99+
The nested-copy tutorial also accepts `--stop-row 5`. Both producer and consumer
100+
break before acquiring/waiting for row 5, and the example verifies that the
101+
uncopied output and trace remain zero. A pipeline schedule must ensure that its
102+
producer and consumer exit consistently and finish any acquired work before
103+
exiting; `break_loop()` does not implicitly release or drain an interrupted step.
104+
105+
Host expansion honors breaks for its selected opaque-condition assignment and
106+
representative loop bounds, including zero-trip loops. Opaque assignments remain
107+
fixed during each explored execution; these checks do not prove safety for all
108+
possible iteration-varying runtime predicate sequences.
34109

35110
`StageInfo` contains the current `stage_idx`, `phase`, selected full `barrier`, zero-based iteration count, loop offset/bounds, work label, and owning `ExecutionContext`; loop offset/bounds are `None` for peeled work outside a domain loop. This lets pipeline payload work use `stage_info.barrier` without knowing how barrier arrays are stored by the device manager.
36111

‎experimental/task_scheduling/src/task_scheduling/__init__.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
)
2323
from .exhaustive_checker import check_all_interleavings, expand_task
2424
from .ir import (
25+
BreakLoopIR,
2526
ConditionalIR,
2627
DependencyEdgeIR,
2728
DomainLoopIR,
@@ -58,6 +59,8 @@
5859
producer_work,
5960
)
6061
from .schedule_builder import (
62+
BreakLoop,
63+
break_loop,
6164
ConditionalBlock,
6265
DomainLoop,
6366
DynamicDomainBound,
@@ -96,6 +99,9 @@
9699

97100

98101
__all__ = [
102+
"BreakLoop",
103+
"BreakLoopIR",
104+
"break_loop",
99105
"BarrierAllocation",
100106
"BarrierAllocator",
101107
"check_all_interleavings",

‎experimental/task_scheduling/src/task_scheduling/exhaustive_checker.py‎

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,11 +18,13 @@
1818
guard_fires,
1919
)
2020
from .schedule_builder import (
21+
BreakLoop,
2122
ConditionalBlock,
2223
DomainLoop,
2324
Step,
2425
WorkTileLoop,
2526
_iter_nodes,
27+
validate_break_placement,
2628
)
2729
from .resources import WorkQueue
2830
from .task import Task
@@ -63,11 +65,14 @@ def expand_task(
6365
dynamic_domain_fallback: int = 1,
6466
representative_domain: bool = False,
6567
) -> list[FlatScheduleOp]:
68+
validate_break_placement(task.schedule.body)
6669
result = []
6770
domain_entry = itertools.count()
6871

6972
def emit(nodes, iteration=None, count=None, phase="O", skipped=False):
7073
for node in nodes:
74+
if isinstance(node, BreakLoop):
75+
return True
7176
if isinstance(node, Step):
7277
# Static work-queue stages only update task-local scheduler
7378
# state; unlike a CLC queue, they do not synchronize tasks.
@@ -96,10 +101,11 @@ def emit(nodes, iteration=None, count=None, phase="O", skipped=False):
96101
stride = _bound(node.step, task, 1)
97102
if stride <= 0:
98103
raise ValueError("exhaustive checker requires positive domain step")
99-
values = tuple(range(start, end, stride)) or (start,)
104+
values = tuple(range(start, end, stride))
100105
next(domain_entry)
101106
for index, _ in enumerate(values):
102-
emit(node.body, index, len(values), "", skipped)
107+
if emit(node.body, index, len(values), "", skipped):
108+
break
103109
elif isinstance(node, WorkTileLoop):
104110
for _ in range(num_tiles):
105111
emit(node.body, iteration, count, phase, skipped_tile)
@@ -126,7 +132,9 @@ def emit(nodes, iteration=None, count=None, phase="O", skipped=False):
126132
if guard is LAST_ITER
127133
else phase
128134
)
129-
emit(node.body, iteration, count, guard_phase, skipped)
135+
if emit(node.body, iteration, count, guard_phase, skipped):
136+
return True
137+
return False
130138

131139
emit(task.schedule.body)
132140
return result

‎experimental/task_scheduling/src/task_scheduling/ir.py‎

Lines changed: 14 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -68,6 +68,13 @@ class StepIR:
6868
stage_resource_id: int | None = None
6969

7070

71+
@dataclass(frozen=True)
72+
class BreakLoopIR:
73+
"""Exit the enclosing DomainLoopIR without taking its backedge."""
74+
75+
exit_values: tuple[int, ...] = ()
76+
77+
7178
@dataclass(frozen=True)
7279
class ConditionalIR:
7380
body: tuple["NodeIR", ...]
@@ -95,7 +102,7 @@ class WorkTileLoopIR:
95102
skip_if: object | None = None
96103

97104

98-
NodeIR = StepIR | ConditionalIR | DomainLoopIR | WorkTileLoopIR
105+
NodeIR = StepIR | BreakLoopIR | ConditionalIR | DomainLoopIR | WorkTileLoopIR
99106

100107

101108
@dataclass(frozen=True)
@@ -136,7 +143,7 @@ def resource_by_id(self) -> Mapping[int, ResourceIR]:
136143
def _iter_ir_nodes(nodes: tuple[NodeIR, ...]):
137144
for node in nodes:
138145
yield node
139-
if not isinstance(node, StepIR):
146+
if not isinstance(node, (StepIR, BreakLoopIR)):
140147
yield from _iter_ir_nodes(node.body)
141148

142149

@@ -155,7 +162,7 @@ def freeze_program_ir(
155162
"""Normalize capture objects into an identity-free, frozen program view."""
156163

157164
from .enums import Every, FIRST_ITER, LAST_ITER, OpaqueCondition, SKIPPABLE
158-
from .schedule_builder import ConditionalBlock, DomainLoop, Step, WorkTileLoop
165+
from .schedule_builder import BreakLoop, ConditionalBlock, DomainLoop, Step, WorkTileLoop
159166

160167
resource_ids = {resource: index for index, resource in enumerate(resources)}
161168
resource_irs = tuple(
@@ -193,6 +200,8 @@ def freeze_guard(guard) -> GuardIR:
193200
raise TypeError(f"unsupported captured guard {type(guard).__name__}")
194201

195202
def freeze_node(node) -> NodeIR:
203+
if isinstance(node, BreakLoop):
204+
return BreakLoopIR(tuple(value.value_id for value in node.exit_values))
196205
if isinstance(node, Step):
197206
return StepIR(
198207
resource_id=resource_ids[node.memory_resource],
@@ -263,6 +272,8 @@ def freeze_node(node) -> NodeIR:
263272
if isinstance(node, DomainLoopIR):
264273
consumed.update(node.initial_values)
265274
consumed.update(node.yield_values)
275+
elif isinstance(node, BreakLoopIR):
276+
consumed.update(node.exit_values)
266277
task_irs.append(
267278
TaskIR(
268279
task_id=task_id,

‎experimental/task_scheduling/src/task_scheduling/schedule_builder.py‎

Lines changed: 69 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -46,6 +46,13 @@ class Step:
4646
unique_id: int
4747

4848

49+
@dataclass(kw_only=True, eq=False, frozen=True)
50+
class BreakLoop:
51+
"""Exit the enclosing domain loop, skipping the remaining iteration body."""
52+
53+
exit_values: tuple[ScheduleValue, ...] = ()
54+
55+
4956
@dataclass(kw_only=True, eq=False, frozen=True)
5057
class ConditionalBlock:
5158
body: tuple["Node", ...] | list["Node"]
@@ -125,7 +132,7 @@ class WorkTileLoop:
125132
skip_if: Callable[..., object] | None = None
126133

127134

128-
Node = Step | ConditionalBlock | DomainLoop | WorkTileLoop
135+
Node = Step | BreakLoop | ConditionalBlock | DomainLoop | WorkTileLoop
129136

130137

131138
@dataclass(kw_only=True, frozen=True)
@@ -151,7 +158,7 @@ def _visit_node(node: Node, visitor: Callable[[object], None]) -> None:
151158

152159

153160
def _child_nodes(node: Node) -> tuple[Node, ...] | list[Node]:
154-
if isinstance(node, Step):
161+
if isinstance(node, (Step, BreakLoop)):
155162
return ()
156163
return node.body
157164

@@ -186,7 +193,26 @@ def validate_queue_advance_placement(
186193
elif isinstance(node, WorkTileLoop):
187194
validate_queue_advance_placement(node.body, parent_is_work_tile_loop=True)
188195
else:
189-
validate_queue_advance_placement(node.body)
196+
validate_queue_advance_placement(_child_nodes(node))
197+
198+
199+
def validate_break_placement(nodes, *, in_domain_loop=False) -> None:
200+
for node in nodes:
201+
if isinstance(node, BreakLoop) and not in_domain_loop:
202+
raise ScheduleError("break_loop() must be called inside domain_loop()")
203+
validate_break_placement(
204+
_child_nodes(node),
205+
in_domain_loop=in_domain_loop or isinstance(node, DomainLoop),
206+
)
207+
208+
209+
def contains_break(nodes) -> bool:
210+
"""Whether this scope can exit its enclosing domain loop."""
211+
return any(
212+
isinstance(node, BreakLoop)
213+
or (isinstance(node, ConditionalBlock) and contains_break(node.body))
214+
for node in nodes
215+
)
190216

191217

192218
def _format_guard(guard: BlockGuard) -> str:
@@ -212,6 +238,11 @@ def _format_schedule(schedule: Schedule, *, show_routes: bool = True) -> str:
212238

213239
def emit(node: Node, indent: int) -> None:
214240
pad = " " * indent
241+
if isinstance(node, BreakLoop):
242+
values = ", ".join(f"%{value.value_id}" for value in node.exit_values)
243+
suffix = f"values=[{values}]" if values and show_routes else ""
244+
lines.append(f"{pad}BreakLoop({suffix})")
245+
return
215246
if isinstance(node, Step):
216247
label = f", label={node.label!r}" if node.label else ""
217248
lines.append(
@@ -354,13 +385,6 @@ def make_value(
354385
def _inside(self, kind: type) -> bool:
355386
return any(isinstance(item, kind) for item in self.stack)
356387

357-
def _in_condition(self) -> bool:
358-
return any(
359-
isinstance(item, ConditionalBlock)
360-
and isinstance(item.condition, (IterationPredicate, OpaqueCondition))
361-
for item in self.stack
362-
)
363-
364388
def open_scope(self, node: ConditionalBlock | DomainLoop | WorkTileLoop) -> None:
365389
if self.after_work_tile_loop:
366390
raise ScheduleError("no blocks may follow work_tile_loop()")
@@ -383,8 +407,6 @@ def open_scope(self, node: ConditionalBlock | DomainLoop | WorkTileLoop) -> None
383407
DomainLoop
384408
):
385409
raise ScheduleError("iteration predicates require domain_loop()")
386-
if self._in_condition():
387-
raise ScheduleError("conditional scheduling blocks cannot be nested")
388410
self.stack[-1].body.append(node)
389411
self.stack.append(node)
390412

@@ -417,8 +439,11 @@ def finalize(self) -> Schedule:
417439
if len(self.stack) != 1:
418440
raise ScheduleError("schedule ended with an unclosed block")
419441
validate_queue_advance_placement(self.schedule.body)
442+
validate_break_placement(self.schedule.body)
420443

421444
def freeze(node: Node) -> Node:
445+
if isinstance(node, BreakLoop):
446+
return node
422447
if isinstance(node, Step):
423448
return replace(
424449
node,
@@ -589,10 +614,12 @@ def _conditional(cond: object, key: Hashable | None, negated: bool):
589614

590615

591616
def when_true(cond: object, *, key: Hashable | None = None):
617+
"""Capture a guarded region, evaluated only when all enclosing guards fire."""
592618
return _conditional(cond, key, False)
593619

594620

595621
def when_false(cond: object, *, key: Hashable | None = None):
622+
"""Capture a negated guarded region; nested values retain lexical scope."""
596623
return _conditional(cond, key, True)
597624

598625

@@ -608,6 +635,30 @@ def first_iter():
608635
return when_true(FIRST_ITER)
609636

610637

638+
def break_loop(*values: ScheduleValue) -> None:
639+
"""Exit the enclosing domain loop when execution reaches this operation.
640+
641+
Remaining work and the functional loop's backedge are skipped. Explicit
642+
values become the loop's results, in carried-input order. With no arguments,
643+
results retain their values at entry to the interrupted iteration.
644+
"""
645+
_require_active_domain_loop("break_loop")
646+
builder = _require_active_builder("break_loop()")
647+
loop = next(node for node in reversed(builder.stack) if isinstance(node, DomainLoop))
648+
if values and len(values) != len(loop.initial_values):
649+
raise ScheduleError(
650+
"break_loop() must supply one value for every carried input; "
651+
f"expected {len(loop.initial_values)}, got {len(values)}"
652+
)
653+
for (name, initial), value in zip(loop.initial_values.items(), values):
654+
if not isinstance(value, ScheduleValue):
655+
raise ScheduleError(f"break value {name!r} must be a routed schedule value")
656+
builder.validate_value_use(value, name)
657+
if value.stage_resource is not initial.stage_resource:
658+
raise ScheduleError(f"break value {name!r} changes pipeline stage provenance")
659+
builder.stack[-1].body.append(BreakLoop(exit_values=values))
660+
661+
611662
def last_iter():
612663
"""Capture a block that runs only on the last domain-loop iteration."""
613664
_require_active_domain_loop("last_iter")
@@ -680,6 +731,12 @@ def first_iter(self):
680731
self._check()
681732
return first_iter()
682733

734+
def break_loop(self, *values: ScheduleValue) -> None:
735+
self._check()
736+
if _active_builder.get() is not self.builder:
737+
raise ScheduleError("domain-loop handle belongs to another schedule")
738+
break_loop(*values)
739+
683740
def last_iter(self):
684741
self._check()
685742
return last_iter()

0 commit comments

Comments
 (0)