File size: 7,257 Bytes
61b6fb9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
"""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)