Ines-1 / mini_v41 /execution.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
7.26 kB
"""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)