"""Layer execution plans. Today every plan is sequential and this module changes nothing about how the model behaves. It exists to remove one structural assumption that would otherwise be expensive to undo later: for i, block in enumerate(self.encoder): past = cache.enc_kv[i] # <-- cache slot IS the layer index That line ties *which weights run* to *which cache slot they write*. A recurrent core that visits layer 2 twice needs one set of weights but two cache slots, because the two visits see different inputs and therefore produce different K/V. With the indices conflated, the second visit silently overwrites the first visit's cache and every subsequent decode step reads corrupted history -- a bug that would not show up in a single forward pass, only in cached generation. So: `layer_index` selects weights, `step_index` selects cache. A sequential plan makes them equal, which is why nothing changes today. **Nothing here implements looping.** No loop embeddings, no adaptive depth, no PoLar, no weight tying beyond what `layer_index` already expresses. See docs/RECURRENT_DEPTH_PLAN.md for what would come next and what constraints any future plan has to respect (CED boundary, KV sources, Engram injection points). """ from __future__ import annotations from dataclasses import dataclass __all__ = ["LayerStep", "LayerExecutionPlan"] @dataclass(frozen=True) class LayerStep: """One execution of one layer. `layer_index` which layer's parameters to use `step_index` which cache slot this execution owns `visit` 0 for the first time this layer runs in the plan, 1 for the second, ... Carried so a future recurrent core can condition on depth without the plan having to be re-derived. """ layer_index: int step_index: int visit: int = 0 #: Factor applied to this execution's residual contribution. 1.0 leaves the #: block's output untouched and is the only value the baseline ever sees -- #: the model short-circuits on it rather than computing `h + 1.0 * (out - h)`, #: which is algebraically the same and not bitwise the same. residual_scale: float = 1.0 @dataclass(frozen=True) class LayerExecutionPlan: """An ordered list of layer executions. `n_layers` is the number of distinct weight sets; `len(steps)` is how many executions happen, and therefore how many cache slots are needed. For a sequential plan the two coincide. """ steps: tuple[LayerStep, ...] n_layers: int @classmethod def sequential(cls, n_layers: int) -> LayerExecutionPlan: """The only plan currently used: each layer once, in order.""" return cls( steps=tuple(LayerStep(i, i, 0) for i in range(n_layers)), n_layers=n_layers, ) @classmethod def from_sequence(cls, order, n_layers: int | None = None) -> LayerExecutionPlan: """Build a plan from a layer-index sequence, e.g. ``[0,1,2,3,2,3,4,5]``. Cache slots are assigned in execution order, so a repeated layer gets a fresh slot for each visit. Not used by the model yet; it exists so the plan abstraction can be tested for the property that actually matters. """ order = list(order) if not order: raise ValueError("execution plan cannot be empty") if n_layers is None: n_layers = max(order) + 1 bad = [i for i in order if not 0 <= i < n_layers] if bad: raise ValueError(f"layer indices {bad} outside 0..{n_layers - 1}") seen: dict[int, int] = {} steps = [] for slot, layer in enumerate(order): visit = seen.get(layer, 0) seen[layer] = visit + 1 steps.append(LayerStep(layer, slot, visit)) return cls(steps=tuple(steps), n_layers=n_layers) @classmethod def recurrent( cls, n_layers: int, span: tuple[int, int], loop_count: int, residual_scale: float | None = None, forbidden_layers: set[int] | None = None, ) -> LayerExecutionPlan: """A sequential plan with one contiguous span executed `loop_count` times. `span` is a half-open `[start, stop)` slice of layer indices, so `(2, 5)` means E2 E3 E4. With `loop_count=1` this returns exactly `sequential(n_layers)` -- not an equivalent plan, the same one -- so the control arm runs the baseline code path rather than a path that happens to agree with it. `forbidden_layers` are refused inside the body. That is where the Engram injection points go: running one twice applies two gated writes from the same table rows to a stream that changed in between, which is a semantic change rather than a depth change and makes the gate statistics uninterpretable (docs/RECURRENT_DEPTH_PLAN.md, constraint b). """ start, stop = span if not 0 <= start < stop <= n_layers: raise ValueError( f"recurrent span [{start}, {stop}) outside 0..{n_layers}" ) if loop_count < 1: raise ValueError(f"loop_count must be >= 1, got {loop_count}") if loop_count == 1: return cls.sequential(n_layers) body = list(range(start, stop)) clash = sorted(set(body) & set(forbidden_layers or ())) if clash: raise ValueError( f"recurrent span [{start}, {stop}) contains layer(s) {clash}, " f"which must not be executed more than once. Engram injection " f"points and KV sources belong here; see " f"docs/RECURRENT_DEPTH_PLAN.md constraints (b) and (c)." ) order = list(range(start)) + body * loop_count + list(range(stop, n_layers)) if residual_scale is None: residual_scale = 1.0 / loop_count seen: dict[int, int] = {} steps = [] for slot, layer in enumerate(order): visit = seen.get(layer, 0) seen[layer] = visit + 1 steps.append(LayerStep( layer_index=layer, step_index=slot, visit=visit, # Only the looped body is rescaled. Everything outside it runs # exactly as the baseline does. residual_scale=residual_scale if start <= layer < stop else 1.0, )) return cls(steps=tuple(steps), n_layers=n_layers) @property def n_steps(self) -> int: """Cache slots required.""" return len(self.steps) @property def is_sequential(self) -> bool: return all( s.layer_index == i and s.step_index == i and s.visit == 0 for i, s in enumerate(self.steps) ) and len(self.steps) == self.n_layers def layers_visited_more_than_once(self) -> list[int]: counts: dict[int, int] = {} for s in self.steps: counts[s.layer_index] = counts.get(s.layer_index, 0) + 1 return sorted(k for k, v in counts.items() if v > 1) def __iter__(self): return iter(self.steps) def __len__(self) -> int: return len(self.steps)