Skip to content

Commit 2cfd628

Browse files
committed
feat(optimizer): add loop unrolling pass (topic 10)
Implement LoopUnroller that expands FOR/ENDFOR loops with constant trip counts by replicating the loop body. Two modes: - Full unroll: trip_count <= unroll_factor removes FOR/ENDFOR and inlines the body with induction variable replaced by compile-time constants. - Partial unroll: trip_count > unroll_factor replicates the body unroll_factor times per iteration with offset ADDs, plus a fully-unrolled remainder for leftover iterations. Changes: - Add scratchv/optimizer/loop_unroll.py (LoopUnroller class) - Export LoopUnroller in scratchv/optimizer/__init__.py - Register loop-unroll pass in compiler.py after LICM - Add tests/test_loop_unroll.py with 13 test cases covering: full unroll, partial unroll, nested loops, edge cases
1 parent 997d2aa commit 2cfd628

4 files changed

Lines changed: 637 additions & 0 deletions

File tree

‎scratchv/compiler.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -374,10 +374,12 @@ def _run_optimizations(self, program) -> PassResult:
374374
from scratchv.optimizer.peephole import IRPeepholeOptimizer
375375
from scratchv.optimizer.muladd_fusion import MulAddFusion
376376
from scratchv.optimizer.licm import LICM
377+
from scratchv.optimizer.loop_unroll import LoopUnroller
377378

378379
pm.add(_PassAdapter("ir-peephole", IRPeepholeOptimizer(program)))
379380
pm.add(_PassAdapter("muladd-fusion", MulAddFusion(program)))
380381
pm.add(_PassAdapter("licm", LICM(program)))
382+
pm.add(_PassAdapter("loop-unroll", LoopUnroller(program)))
381383

382384
return pm.run(program)
383385

‎scratchv/optimizer/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,11 +3,13 @@
33
from .peephole import IRPeepholeOptimizer
44
from .muladd_fusion import MulAddFusion
55
from .licm import LICM
6+
from .loop_unroll import LoopUnroller
67

78
__all__ = [
89
"ConstantFolder",
910
"DeadCodeEliminator",
1011
"IRPeepholeOptimizer",
1112
"MulAddFusion",
1213
"LICM",
14+
"LoopUnroller",
1315
]

‎scratchv/optimizer/loop_unroll.py‎

Lines changed: 317 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,317 @@
1+
"""Loop Unrolling optimization pass.
2+
3+
Expands FOR/ENDFOR loops with constant trip counts by replicating the
4+
loop body. Two modes:
5+
6+
* **Full unroll** – trip_count <= unroll_factor: the FOR/ENDFOR wrapper
7+
is removed entirely and the body is inlined ``trip_count`` times with
8+
the induction variable replaced by compile-time constants.
9+
10+
* **Partial unroll** – trip_count > unroll_factor: the loop is kept but
11+
each iteration now executes ``unroll_factor`` copies of the original
12+
body. A remainder loop handles the leftover iterations.
13+
14+
Only loops whose ``start``, ``end`` and ``step`` attributes are all
15+
compile-time integers are eligible.
16+
"""
17+
18+
from __future__ import annotations
19+
20+
from typing import Optional
21+
22+
from scratchv.ir.types import (
23+
DataType,
24+
OpCode,
25+
Value,
26+
Instruction,
27+
BasicBlock,
28+
Program,
29+
)
30+
31+
32+
class LoopUnroller:
33+
"""Unroll FOR/ENDFOR loops with constant trip counts."""
34+
35+
def __init__(self, program: Program, unroll_factor: int = 4):
36+
self.program = program
37+
self.unroll_factor = unroll_factor
38+
self._changes = 0
39+
self._name_counter = 0
40+
41+
def run(self) -> int:
42+
"""Run loop unrolling on all functions.
43+
44+
Returns the number of loops unrolled.
45+
"""
46+
for func in self.program.functions:
47+
for block in func.blocks:
48+
self._process_block(block)
49+
return self._changes
50+
51+
# ── helpers ────────────────────────────────────────────────────────
52+
53+
def _fresh_name(self, prefix: str = "lu") -> str:
54+
self._name_counter += 1
55+
return f"{prefix}_{self._name_counter}"
56+
57+
def _fresh_value(self, orig: Value) -> Value:
58+
"""Create a fresh Value with the same type/shape but a new name."""
59+
return Value(
60+
name=self._fresh_name(orig.name),
61+
dtype=orig.dtype,
62+
is_constant=orig.is_constant,
63+
const_value=orig.const_value,
64+
shape=orig.shape,
65+
)
66+
67+
# ── core logic ─────────────────────────────────────────────────────
68+
69+
def _process_block(self, block: BasicBlock) -> None:
70+
"""Scan *block* for FOR/ENDFOR pairs and unroll them."""
71+
instrs = block.instructions
72+
i = 0
73+
while i < len(instrs):
74+
if instrs[i].opcode != OpCode.FOR:
75+
i += 1
76+
continue
77+
78+
for_idx = i
79+
endfor_idx = self._find_matching_endfor(instrs, for_idx)
80+
if endfor_idx is None:
81+
i += 1
82+
continue
83+
84+
for_instr = instrs[for_idx]
85+
start = for_instr.attrs.get("start")
86+
end = for_instr.attrs.get("end")
87+
step = for_instr.attrs.get("step", 1)
88+
89+
if not self._is_const_int(start, end, step):
90+
i += 1
91+
continue
92+
93+
trip_count = self._compute_trip_count(start, end, step)
94+
if trip_count <= 0:
95+
i += 1
96+
continue
97+
98+
body = instrs[for_idx + 1 : endfor_idx]
99+
if not body:
100+
# empty loop – just remove
101+
del instrs[for_idx : endfor_idx + 1]
102+
self._changes += 1
103+
i = for_idx
104+
continue
105+
106+
iv = for_instr.dest
107+
factor = self.unroll_factor
108+
109+
if trip_count <= factor:
110+
self._full_unroll(instrs, for_idx, endfor_idx, body, iv,
111+
start, step, trip_count)
112+
self._changes += 1
113+
i = 0 # restart: unrolled body may contain new FORs
114+
else:
115+
self._partial_unroll(instrs, for_idx, endfor_idx,
116+
body, iv, start, end, step,
117+
trip_count, factor)
118+
self._changes += 1
119+
# Don't restart: the loop itself is still eligible but
120+
# was already partially unrolled. Advance past it.
121+
i = for_idx + 1
122+
123+
# ── full unroll ────────────────────────────────────────────────────
124+
125+
def _full_unroll(
126+
self,
127+
instrs: list[Instruction],
128+
for_idx: int,
129+
endfor_idx: int,
130+
body: list[Instruction],
131+
iv: Optional[Value],
132+
start: int,
133+
step: int,
134+
trip_count: int,
135+
) -> None:
136+
"""Remove FOR/ENDFOR and inline body *trip_count* times."""
137+
replacements: list[Instruction] = []
138+
for k in range(trip_count):
139+
iv_const = start + k * step
140+
for orig_instr in body:
141+
cloned = self._clone_instr(orig_instr)
142+
if iv is not None:
143+
self._subst_iv(cloned, iv, iv_const)
144+
replacements.append(cloned)
145+
146+
instrs[for_idx : endfor_idx + 1] = replacements
147+
148+
# ── partial unroll ─────────────────────────────────────────────────
149+
150+
def _partial_unroll(
151+
self,
152+
instrs: list[Instruction],
153+
for_idx: int,
154+
endfor_idx: int,
155+
body: list[Instruction],
156+
iv: Optional[Value],
157+
start: int,
158+
end: int,
159+
step: int,
160+
trip_count: int,
161+
factor: int,
162+
) -> int:
163+
"""Partial unroll: replicate body *factor* times per iteration.
164+
165+
Returns the number of instructions inserted (for caller to update
166+
its index).
167+
"""
168+
full_iters = trip_count // factor
169+
remainder = trip_count % factor
170+
171+
# Build the unrolled body (factor copies)
172+
unrolled_body: list[Instruction] = []
173+
for k in range(factor):
174+
for orig_instr in body:
175+
cloned = self._clone_instr(orig_instr)
176+
# Replace iv with (iv + k*step) – but only for k>0
177+
# k==0 keeps the original iv reference
178+
if iv is not None and k > 0:
179+
self._subst_iv_offset(cloned, iv, k * step, unrolled_body)
180+
unrolled_body.append(cloned)
181+
182+
# Adjust FOR bounds: new_end = start + full_iters * factor * step
183+
new_end = start + full_iters * factor * step
184+
instrs[for_idx].attrs["end"] = new_end
185+
186+
# Replace old body with unrolled body
187+
instrs[for_idx + 1 : endfor_idx] = unrolled_body
188+
new_endfor_idx = for_idx + 1 + len(unrolled_body)
189+
# Ensure ENDFOR is still there
190+
if new_endfor_idx < len(instrs) and instrs[new_endfor_idx].opcode == OpCode.ENDFOR:
191+
pass # already there
192+
else:
193+
instrs.insert(new_endfor_idx, Instruction(opcode=OpCode.ENDFOR))
194+
195+
insert_point = new_endfor_idx + 1
196+
197+
# Remainder: fully unroll the leftover iterations
198+
if remainder > 0:
199+
remainder_instrs: list[Instruction] = []
200+
base_iv_val = new_end
201+
for k in range(remainder):
202+
iv_const = base_iv_val + k * step
203+
for orig_instr in body:
204+
cloned = self._clone_instr(orig_instr)
205+
if iv is not None:
206+
self._subst_iv(cloned, iv, iv_const)
207+
remainder_instrs.append(cloned)
208+
instrs[insert_point:insert_point] = remainder_instrs
209+
insert_point += len(remainder_instrs)
210+
211+
return insert_point - for_idx
212+
213+
# ── instruction cloning ────────────────────────────────────────────
214+
215+
def _clone_instr(self, instr: Instruction) -> Instruction:
216+
"""Deep-clone an instruction, giving dest a fresh name."""
217+
new_dest = self._fresh_value(instr.dest) if instr.dest else None
218+
# Operands keep their names (we fixup iv separately)
219+
new_operands = list(instr.operands)
220+
return Instruction(
221+
opcode=instr.opcode,
222+
dest=new_dest,
223+
operands=new_operands,
224+
attrs=dict(instr.attrs),
225+
target=instr.target,
226+
)
227+
228+
# ── induction variable substitution ────────────────────────────────
229+
230+
def _subst_iv(self, instr: Instruction, iv: Value, const_val: int) -> None:
231+
"""Replace all references to *iv* in *instr* with a constant."""
232+
if instr.dest and instr.dest.name == iv.name:
233+
instr.dest = Value(
234+
name=instr.dest.name,
235+
dtype=instr.dest.dtype,
236+
is_constant=True,
237+
const_value=const_val,
238+
shape=instr.dest.shape,
239+
)
240+
new_ops: list[Value] = []
241+
for op in instr.operands:
242+
if op.name == iv.name:
243+
new_ops.append(Value(
244+
name=op.name,
245+
dtype=op.dtype,
246+
is_constant=True,
247+
const_value=const_val,
248+
shape=op.shape,
249+
))
250+
else:
251+
new_ops.append(op)
252+
instr.operands = new_ops
253+
254+
def _subst_iv_offset(
255+
self,
256+
instr: Instruction,
257+
iv: Value,
258+
offset: int,
259+
prefix_instrs: list[Instruction],
260+
) -> None:
261+
"""Replace references to *iv* with ``(iv + offset)``.
262+
263+
For partial unrolling, we cannot replace iv with a constant because
264+
the loop variable is still live. Instead we insert a fresh
265+
``ADD iv, offset`` before this instruction and rewrite the operand
266+
reference to the new value.
267+
"""
268+
iv_offset_val = Value(
269+
name=self._fresh_name(f"{iv.name}_off"),
270+
dtype=iv.dtype,
271+
)
272+
273+
# Emit: iv_offset_val = ADD iv, <offset_const>
274+
offset_const = Value(
275+
name=self._fresh_name("off"),
276+
dtype=DataType.INT32,
277+
is_constant=True,
278+
const_value=offset,
279+
)
280+
prefix_instrs.append(Instruction(
281+
opcode=OpCode.ADD,
282+
dest=iv_offset_val,
283+
operands=[iv, offset_const],
284+
))
285+
286+
# Rewrite operands
287+
new_ops: list[Value] = []
288+
for op in instr.operands:
289+
if op.name == iv.name:
290+
new_ops.append(iv_offset_val)
291+
else:
292+
new_ops.append(op)
293+
instr.operands = new_ops
294+
295+
# ── utilities ──────────────────────────────────────────────────────
296+
297+
@staticmethod
298+
def _find_matching_endfor(instrs: list[Instruction], start: int) -> Optional[int]:
299+
depth = 0
300+
for i in range(start, len(instrs)):
301+
if instrs[i].opcode == OpCode.FOR:
302+
depth += 1
303+
elif instrs[i].opcode == OpCode.ENDFOR:
304+
depth -= 1
305+
if depth == 0:
306+
return i
307+
return None
308+
309+
@staticmethod
310+
def _is_const_int(*vals) -> bool:
311+
return all(isinstance(v, int) for v in vals)
312+
313+
@staticmethod
314+
def _compute_trip_count(start: int, end: int, step: int) -> int:
315+
if step <= 0 or end <= start:
316+
return 0
317+
return (end - start + step - 1) // step # ceiling division

0 commit comments

Comments
 (0)