|
| 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