ShourenWSR commited on
Commit
dfffebc
·
verified ·
1 Parent(s): 9c09460

SYNKA 27938: verified step_00028000

Browse files
step_00028000/config.json ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "LoopMoEForCausalLM"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "modeling_loop_lm.LoopMoEConfig",
7
+ "AutoModelForCausalLM": "modeling_loop_lm.LoopMoEForCausalLM"
8
+ },
9
+ "context_length": 4096,
10
+ "d_ff": 4608,
11
+ "d_model": 1792,
12
+ "lb_loss_factor": 0.01,
13
+ "lz_loss_factor": 0.001,
14
+ "model_type": "loop-moe",
15
+ "model_variant": "looped-moe",
16
+ "num_active": 2,
17
+ "num_experts": 16,
18
+ "num_head_layers": 1,
19
+ "num_heads": 28,
20
+ "num_layers": 16,
21
+ "num_layers_in_stack": 4,
22
+ "num_stacks": 4,
23
+ "num_tail_layers": 1,
24
+ "per_pass_attention": true,
25
+ "rope_theta": 10000.0,
26
+ "tie_word_embeddings": false,
27
+ "transformers_version": "5.15.0",
28
+ "unrolled_depth": 18,
29
+ "vocab_size": 49152,
30
+ "width_ratio": 14.0
31
+ }
step_00028000/loop_trace.py ADDED
@@ -0,0 +1,539 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Per-loop-step trace contract between the model and the metrics layer.
2
+
3
+ This module is the *interface* between `pretrain.loopmoe` (which produces raw
4
+ tensors) and `pretrain.metrics` (which turns them into numbers). It deliberately
5
+ imports **nothing** from torchtitan or from any trainer, so that:
6
+
7
+ * the metrics layer can be developed and tested without a trainer, and
8
+ * swapping the training backend (see HANDOVER §4.1' item 1) touches nothing here.
9
+
10
+ Division of labour (agreed with infra-metrics, 2026-08-11):
11
+
12
+ the model / trainer -> *when* to collect and *how* to persist
13
+ the metrics layer -> *what* to collect and *how* to compute it
14
+
15
+ Hence the model hands over **raw tensors, never reduced scalars**. A norm
16
+ computed in two places is two definitions of that norm; the project has already
17
+ paid for that mistake once (HANDOVER §4.4: "same-name metrics have several
18
+ non-equivalent definitions"). The single definition lives in the metrics layer.
19
+
20
+ Tensor lifetime contract (IMPORTANT)
21
+ ------------------------------------
22
+ `TracePoint` tensors are ``detach()``-ed **views of live activations**, not
23
+ copies. Passing them costs ~0 extra memory, which is why the model can afford to
24
+ hand over all R x L residual tensors instead of subsampling layers. The price is
25
+ a rule the consumer MUST obey:
26
+
27
+ 1. The collector callback is **synchronous**. When it returns, the tensors are
28
+ considered dead.
29
+ 2. The collector MUST NOT store a TracePoint (or any tensor inside one) in any
30
+ container that outlives the call -- no ``self.foo = point``, no appending to
31
+ a list, no closure capture.
32
+ 3. Anything needed across steps must first be reduced to a Python scalar.
33
+
34
+ Violating this pins whole activation graphs in memory and turns a constant-memory
35
+ probe into a leak that grows with training length. `assert_trace_released` below
36
+ turns that failure into a loud test failure instead of a slow OOM at step 40k.
37
+ """
38
+
39
+ from __future__ import annotations
40
+
41
+ import random
42
+ import weakref
43
+ from dataclasses import dataclass
44
+ from typing import Any, Protocol, runtime_checkable
45
+
46
+ import torch
47
+ import torch.nn.functional as F
48
+ from torch import Tensor
49
+
50
+ __all__ = [
51
+ "ACTIVE_TRACE_SINK",
52
+ "ROUTER_LOGITS_ARE_PRESOFTMAX",
53
+ "LOOP_AXIS_PROVENANCE",
54
+ "TracePoint",
55
+ "LossComponents",
56
+ "TraceMeta",
57
+ "TraceSink",
58
+ "LoopTraceCollector",
59
+ "NullCollector",
60
+ "run_probe_forward",
61
+ "assert_trace_released",
62
+ ]
63
+
64
+
65
+ # --- semantic markers -------------------------------------------------------
66
+ #
67
+ # HANDOVER §4.3: the project has been burned by mislabelling sigmoid/logits,
68
+ # which produced results that "looked entirely reasonable but were all wrong".
69
+ # The collector schema has a `router_logits_are_presoftmax` field that pairs with
70
+ # this constant; it is a constant rather than a runtime flag because the model
71
+ # has exactly one behaviour and a flag that can only take one value is a lie
72
+ # waiting to happen.
73
+
74
+ ROUTER_LOGITS_ARE_PRESOFTMAX: bool = True
75
+ """`TracePoint.router_logits` is the raw router linear output (pre-softmax)."""
76
+
77
+ LOOP_AXIS_PROVENANCE: str = "hook_call_index"
78
+ """How `TracePoint.loop_step` was determined.
79
+
80
+ The value means the loop index was threaded down from the Python `for` loop that
81
+ drives the stack -- a per-module call counter. It was **not** recovered by
82
+ reshaping a flattened layer axis, which silently yields transposed semantics
83
+ (HANDOVER §4.3).
84
+
85
+ **The exact string matters.** It is not a description; it is a value the
86
+ diagnostics pipeline validates. `looped_diag/collect/schema.py` accepts only::
87
+
88
+ allowed = ("hook_call_index", "hook_call_index+external")
89
+
90
+ and raises otherwise (schema.py:398-401), with `probe_*.py` scripts asserting the
91
+ same. An earlier draft of this constant read ``"explicit_call_counter"`` -- the
92
+ same claim in different words, and it would have made every §4.7 ingestion of our
93
+ traces raise. Do not "improve" the wording.
94
+ """
95
+
96
+
97
+ @dataclass(frozen=True)
98
+ class TracePoint:
99
+ """Raw per-(loop_step, layer) tensors. See the lifetime contract above.
100
+
101
+ All tensors are detached views; none require grad. Shapes use
102
+ B=batch, S=sequence, E=num_experts, k=num_active, d=d_model.
103
+ """
104
+
105
+ block: str
106
+ """Which part of the network produced this row: "head", "loop" or "tail".
107
+
108
+ Head and tail are single MoE layers run once, outside the loop (ablation recipe
109
+ 2026-09-12 section 2.1). `loop_step` and `layer_idx` describe a position inside the
110
+ shared stack and are meaningless for them; they are recorded as 0/0 and must not be
111
+ read as coordinates unless `block == "loop"`.
112
+ """
113
+
114
+ unrolled_pos: int
115
+ """Position in the unrolled network, 0-based: THE depth coordinate.
116
+
117
+ head = 0, the loop's layers run 1 .. num_stacks*num_layers_in_stack in execution
118
+ order, tail = 1 + num_stacks*num_layers_in_stack (and correspondingly higher when a
119
+ configuration has more than one head/tail layer). This is the sink's key, so it is
120
+ also the only ordering that is guaranteed unique: sorting by (loop_step, layer_idx)
121
+ puts head, tail and the loop's first layer on top of each other.
122
+ """
123
+
124
+ loop_step: int
125
+ """Which pass through the shared stack, 0-based. Range: [0, num_stacks).
126
+
127
+ Only meaningful when `block == "loop"`; 0 for head and tail.
128
+ """
129
+
130
+ layer_idx: int
131
+ """Which physical layer inside the stack, 0-based. Range: [0, num_layers_in_stack).
132
+
133
+ Only meaningful when `block == "loop"`; 0 for head and tail.
134
+ """
135
+
136
+ num_experts: int
137
+ """E for THIS layer, taken from the router tensor's own shape.
138
+
139
+ Per row, not per model, because head/tail layers keep 8 experts whatever the loop
140
+ block uses (ablation recipe 2.1) -- S4 gives the loop block 64. A consumer that reads
141
+ the model-level `TraceMeta.num_experts` and applies it to a head row computes, for
142
+ example, L2 >= 1 - 8/64 = 0.875 and reads a perfectly healthy head layer as heavily
143
+ collapsed. Nothing raises: the number is in range and looks plausible, and head/tail
144
+ behaviour is exactly what the ablation is there to measure.
145
+ """
146
+
147
+ top_k: int
148
+ """k for THIS layer, from the selection tensor's own shape. Same reasoning as above."""
149
+
150
+ router_logits: Tensor
151
+ """[B, S, E] raw router output, pre-softmax. See ROUTER_LOGITS_ARE_PRESOFTMAX."""
152
+
153
+ router_probs: Tensor
154
+ """[B, S, E] softmax over **all** E experts (not renormalised over top-k)."""
155
+
156
+ topk_idx: Tensor
157
+ """[B, S, k] indices of the selected experts."""
158
+
159
+ topk_weights: Tensor
160
+ """[B, S, k] combine weights: top-k of `router_probs`, renormalised to sum to 1.
161
+
162
+ This is the quantity the MoE combine step already computes; exposing it adds
163
+ no arithmetic.
164
+ """
165
+
166
+ residual: Tensor
167
+ """[B, S, d] the block's output tensor, with **no reduction applied**.
168
+
169
+ Semantically this is the diagnostics repo's KIND_RESIDUAL / CAPTURE_BLOCK_OUTPUT:
170
+ the return value of the decoder block's forward. The metrics layer derives
171
+ both the within-step and the step-boundary norms from this one tensor, so
172
+ that both come from a single definition.
173
+ """
174
+
175
+
176
+ @dataclass(frozen=True)
177
+ class LossComponents:
178
+ """Per-step loss breakdown. Cheap enough to emit on *every* step.
179
+
180
+ HANDOVER §4.4 calls the aux losses the "blood-pressure monitor" and asks for
181
+ them every step. The *factors* are included alongside the values because
182
+ milestone M3.5 has to be able to reconstruct, from the logs alone, which
183
+ coefficient was in force at any step (HANDOVER §3.3'').
184
+ """
185
+
186
+ task_loss: float
187
+ """Cross-entropy on next-token prediction, before any aux term."""
188
+
189
+ lb_loss: float
190
+ """Switch-style load-balancing loss, E * sum_i f_i * p_i, averaged over all
191
+ unrolled layers (num_stacks * num_layers_in_stack)."""
192
+
193
+ lb_loss_factor: float
194
+ """Coefficient multiplying `lb_loss` in `total_loss`."""
195
+
196
+ z_loss: float
197
+ """Router z-loss, mean((logsumexp logits)^2), averaged the same way."""
198
+
199
+ z_loss_factor: float
200
+ """Coefficient multiplying `z_loss` in `total_loss`."""
201
+
202
+ total_loss: float
203
+ """task_loss + lb_loss_factor * lb_loss + z_loss_factor * z_loss."""
204
+
205
+
206
+ @dataclass(frozen=True)
207
+ class TraceMeta:
208
+ """Shape/provenance context so a collector never has to infer structure."""
209
+
210
+ num_stacks: int
211
+ """R: how many times the shared stack is called (loop count)."""
212
+
213
+ num_layers_in_stack: int
214
+ """L: physical layers per stack."""
215
+
216
+ num_experts: int
217
+ """E: experts per MoE layer."""
218
+
219
+ num_active: int
220
+ """k: experts activated per token."""
221
+
222
+ router_logits_are_presoftmax: bool = ROUTER_LOGITS_ARE_PRESOFTMAX
223
+ loop_axis_provenance: str = LOOP_AXIS_PROVENANCE
224
+
225
+
226
+ # A sink is just the dict the model fills in during a traced forward. The
227
+ # trainer creates one, passes it to the model, hands it to the collector, and
228
+ # drops it -- so no hook registration/deregistration dance, and no state living
229
+ # on the modules between steps.
230
+ TraceSink = dict[int, TracePoint]
231
+ """Keyed by ``unrolled_pos`` -- explicit, never shape-derived.
232
+
233
+ It was keyed by ``(loop_step, layer_idx)`` until the head/tail layers arrived (ablation
234
+ recipe 2026-09-12). Those run outside the loop, so they have no meaningful value for
235
+ either coordinate; recording them as 0/0 -- which is what they are -- would have put
236
+ head, the loop's very first layer, and tail on the same key, and a dict assignment does
237
+ not complain. Two of the three rows would have vanished with nothing in the output
238
+ saying so. `unrolled_pos` is the depth coordinate Research defined, and keying on it
239
+ makes the key unique by construction rather than by a sentinel convention that a later
240
+ reader has to know about.
241
+ """
242
+
243
+
244
+ class _ActiveTraceSink:
245
+ """The sink the blocks write into, reached as a module-level object.
246
+
247
+ Why this exists rather than passing the dict down as a keyword argument:
248
+ `fully_shard` is applied per block (`src/pretrain/train_spec.py`), and FSDP2
249
+ repacks kwargs across that boundary -- a dict passed by keyword arrives as a
250
+ fresh copy at every sharded block. The blocks were writing faithfully into
251
+ copies that were then discarded, so multi-GPU runs produced no router rows
252
+ at all while every intermediate layer looked correct. Diagnosed 2026-09-05
253
+ by printing `id()` on both sides: identical single-process, different at
254
+ every block under two ranks.
255
+
256
+ A module-level object does not cross that boundary, which is exactly why
257
+ `AUX_LOSS_STATE` never had the problem -- the aux losses travel up through
258
+ return values and are published to a global. This is the same pattern.
259
+
260
+ Process-local by construction: each rank has its own interpreter and so its
261
+ own instance, which is what makes the collected values rank-local. That is a
262
+ property to label in the output, not to hide -- see `aggregation_scope` in
263
+ the router rows.
264
+ """
265
+
266
+ def __init__(self) -> None:
267
+ self.sink: TraceSink | None = None
268
+
269
+ def arm(self) -> TraceSink:
270
+ """Start collecting; returns the dict that will be filled."""
271
+ self.sink = {}
272
+ return self.sink
273
+
274
+ def disarm(self) -> TraceSink | None:
275
+ """Stop collecting and hand back whatever was gathered."""
276
+ sink, self.sink = self.sink, None
277
+ return sink
278
+
279
+
280
+ ACTIVE_TRACE_SINK = _ActiveTraceSink()
281
+
282
+
283
+ @runtime_checkable
284
+ class LoopTraceCollector(Protocol):
285
+ """What the metrics layer implements; what the training loop calls."""
286
+
287
+ def on_metrics_step(
288
+ self, step: int, loss: LossComponents, grad_norm: float | None = None
289
+ ) -> dict[str, float]:
290
+ """Called on **every** optimizer step. Returns flat scalars to log.
291
+
292
+ `grad_norm` is passed only on the steps where the training loop
293
+ already has it as a host float; on all other steps it is `None` and
294
+ the field is omitted from the row rather than guessed at.
295
+ """
296
+ ...
297
+
298
+ def on_trace_step(
299
+ self, step: int, sink: TraceSink, meta: TraceMeta
300
+ ) -> dict[str, float]:
301
+ """Called every N steps with the raw tensors.
302
+
303
+ MUST be synchronous and MUST NOT retain any tensor from `sink`
304
+ (see the lifetime contract at the top of this module).
305
+ """
306
+ ...
307
+
308
+ def on_probe_step(
309
+ self, step: int, sink: TraceSink, meta: TraceMeta, per_token_loss: Tensor
310
+ ) -> dict[str, float]:
311
+ """Called every M steps with a forward over the fixed probe corpus.
312
+
313
+ Deliberately a separate method rather than `on_trace_step` with a
314
+ `source="probe"` flag: probe data answers different questions
315
+ (cross-loop-step overlap, repetition-stratified loss) and lands in a
316
+ different file, and a flag is something a downstream analysis can forget
317
+ to filter on.
318
+
319
+ `sink` is structurally identical to `on_trace_step`'s. `per_token_loss`
320
+ is [B, S] and obeys the same no-retention contract.
321
+ """
322
+ ...
323
+
324
+
325
+ class NullCollector:
326
+ """No-op collector, so training runs with metrics switched off."""
327
+
328
+ def on_metrics_step(
329
+ self, step: int, loss: LossComponents, grad_norm: float | None = None
330
+ ) -> dict[str, float]:
331
+ return {}
332
+
333
+ def on_trace_step(
334
+ self, step: int, sink: TraceSink, meta: TraceMeta
335
+ ) -> dict[str, float]:
336
+ return {}
337
+
338
+ def on_probe_step(
339
+ self, step: int, sink: TraceSink, meta: TraceMeta, per_token_loss: Tensor
340
+ ) -> dict[str, float]:
341
+ return {}
342
+
343
+
344
+ def _snapshot_rng() -> dict[str, Any]:
345
+ """Capture every RNG stream a probe forward could disturb."""
346
+ state: dict[str, Any] = {
347
+ "python": random.getstate(),
348
+ "torch": torch.get_rng_state(),
349
+ }
350
+ try:
351
+ import numpy as np
352
+
353
+ state["numpy"] = np.random.get_state()
354
+ except ImportError: # pragma: no cover - numpy is a hard dependency in practice
355
+ pass
356
+ if torch.cuda.is_available():
357
+ state["cuda"] = torch.cuda.get_rng_state_all()
358
+ return state
359
+
360
+
361
+ def _restore_rng(state: dict[str, Any]) -> None:
362
+ random.setstate(state["python"])
363
+ torch.set_rng_state(state["torch"])
364
+ if "numpy" in state:
365
+ import numpy as np
366
+
367
+ np.random.set_state(state["numpy"])
368
+ if "cuda" in state:
369
+ torch.cuda.set_rng_state_all(state["cuda"])
370
+
371
+
372
+ def run_probe_forward(
373
+ model: Any,
374
+ input_ids: Tensor,
375
+ *,
376
+ collector: LoopTraceCollector,
377
+ step: int,
378
+ ) -> dict[str, float]:
379
+ """Run the fixed probe corpus through the model and hand it to the collector.
380
+
381
+ Three isolation properties, each of which exists because violating it would
382
+ corrupt something silently:
383
+
384
+ 1. **No gradients, eval mode.** LT2 disabled probing outright because
385
+ activation memory multiplies by the loop count (HANDOVER §3.4); at R=16
386
+ this is 4x the pressure they measured, so a probe that built a graph would
387
+ OOM at exactly the configurations we most want to measure.
388
+ 2. **RNG is restored on exit** (python / numpy / torch / cuda). Without this,
389
+ whether a probe ran would change the training trajectory, and a resumed
390
+ run would silently diverge from an uninterrupted one -- breaking the §4.6
391
+ requirement that the two be statistically indistinguishable.
392
+ 3. **The training dataloader is never touched.** The probe corpus is a fixed
393
+ tensor held separately, so probing does not advance the data position that
394
+ checkpoints record.
395
+
396
+ The model's training/eval mode is restored even if the collector raises.
397
+
398
+ Returns whatever scalars the collector produced.
399
+ """
400
+ was_training = model.training
401
+ rng_state = _snapshot_rng()
402
+ sink: TraceSink = {}
403
+ try:
404
+ model.eval()
405
+ with torch.no_grad():
406
+ out = model(input_ids, trace_sink=sink)
407
+ logits = out.logits if hasattr(out, "logits") else out
408
+
409
+ # Next-token targets. The model does not shift internally (the
410
+ # dataloader supplies pre-shifted labels during training), so the
411
+ # shift is explicit here. The final position has no next token and is
412
+ # masked out; cross_entropy returns 0.0 there.
413
+ labels = input_ids.new_full(input_ids.shape, -100)
414
+ labels[:, :-1] = input_ids[:, 1:]
415
+ per_token_loss = F.cross_entropy(
416
+ logits.transpose(1, 2).float(), labels,
417
+ reduction="none", ignore_index=-100,
418
+ ) # [B, S]
419
+
420
+ meta = model.config.trace_meta
421
+ return collector.on_probe_step(step, sink, meta, per_token_loss)
422
+ finally:
423
+ sink.clear()
424
+ _restore_rng(rng_state)
425
+ if was_training:
426
+ model.train()
427
+
428
+
429
+ def assert_trace_released(sink: TraceSink) -> None:
430
+ """Raise if a collector retained tensors from `sink` past its callback.
431
+
432
+ Call this immediately after the collector returns and after clearing local
433
+ references. It weak-references every tensor, drops the sink, and checks that
434
+ nothing kept them alive.
435
+
436
+ This is the enforcement half of the lifetime contract. It is meant to be
437
+ wired into the toy/CI run rather than the production hot path.
438
+ """
439
+ refs: list[weakref.ref] = []
440
+ for point in sink.values():
441
+ for tensor in (
442
+ point.router_logits,
443
+ point.router_probs,
444
+ point.topk_idx,
445
+ point.topk_weights,
446
+ point.residual,
447
+ ):
448
+ try:
449
+ refs.append(weakref.ref(tensor))
450
+ except TypeError: # pragma: no cover - torch tensors are weakref-able
451
+ continue
452
+ # After a `for` loop the loop variables stay bound in this frame, so `point`
453
+ # and `tensor` would still reference the *last* TracePoint when we collect --
454
+ # and this function would report itself as the leaker. Drop them explicitly.
455
+ point = tensor = None # noqa: F841 - rebinding to release references
456
+ sink.clear()
457
+
458
+ import gc
459
+
460
+ gc.collect()
461
+ leaked = sum(1 for ref in refs if ref() is not None)
462
+ if leaked:
463
+ raise RuntimeError(
464
+ f"{leaked}/{len(refs)} trace tensors are still alive after the collector "
465
+ "returned. A collector must not store TracePoint tensors beyond the "
466
+ "callback; reduce to scalars first. See the lifetime contract in "
467
+ "src/model/loop_trace.py."
468
+ )
469
+
470
+ def unrolled_position(loop_step: int, layer_idx: int, num_layers_in_stack: int) -> int:
471
+ """Depth coordinate of a loop-block layer: `loop_step * n_layers + layer_idx`.
472
+
473
+ THE definition of that arithmetic. It is needed in two places -- the legacy
474
+ checkpoint adapter in the cross-architecture probe, and the `loop_only`
475
+ branch of `cross_arch_v2_offline` that reads pre-head/tail artifacts -- and
476
+ the two must agree, because one writes the coordinate and the other reads it
477
+ back. Two copies of a formula that indexes into a dump do not fail loudly
478
+ when they drift; they silently address different cells.
479
+
480
+ Only meaningful for `block == "loop"`. Head and tail run once and are their
481
+ own physical layers, so their position comes from the network's order, not
482
+ from this arithmetic.
483
+ """
484
+ if num_layers_in_stack <= 0:
485
+ raise ValueError(
486
+ f"num_layers_in_stack must be positive, got {num_layers_in_stack}. "
487
+ "A zero or negative stack size collapses every pass onto the same "
488
+ "position, which reads as a valid coordinate."
489
+ )
490
+ return loop_step * num_layers_in_stack + layer_idx
491
+
492
+
493
+ def legacy_trace_point_factory(num_layers_in_stack: int):
494
+ """A `TracePoint` stand-in accepting the PRE-head/tail keyword set.
495
+
496
+ Checkpoints trained before the head/tail change bundle their own
497
+ `modeling_loop_lm.py`, which constructs `TracePoint(loop_step=..., layer_idx=...,
498
+ <tensors>)`. Those four fields are now required, so such a checkpoint cannot
499
+ be loaded at all -- the failure is a `TypeError` inside the bundled code,
500
+ raised before any forward pass completes.
501
+
502
+ The four fields are filled rather than defaulted, and deliberately NOT given
503
+ defaults on `TracePoint` itself: `num_experts` being required per cell is
504
+ what stops a head layer's balanced 8 experts being divided by the loop
505
+ block's 64, which would put a floor of 0.875 under `L2` and read as severe
506
+ collapse in every configuration. That protection is worth keeping for new
507
+ checkpoints even though old ones need this adapter.
508
+
509
+ A pre-head/tail model is all loop block by construction, so `block` is
510
+ `"loop"` for every row and the depth coordinate is the unrolled arithmetic.
511
+ """
512
+
513
+ # Bound HERE, not looked up inside `make`. The adapter is installed by
514
+ # REPLACING the module-level `TracePoint` name, so a lookup at call time
515
+ # would find this factory instead of the class -- the adapter would call
516
+ # itself. Capturing the class before installation is what makes the
517
+ # substitution safe.
518
+ cls = TracePoint
519
+
520
+ def make(**kwargs: Any) -> "TracePoint":
521
+ logits = kwargs.get("router_logits")
522
+ topk_idx = kwargs.get("topk_idx")
523
+ if logits is None or topk_idx is None:
524
+ raise ValueError(
525
+ "legacy TracePoint adapter needs router_logits and topk_idx to "
526
+ f"infer num_experts and top_k; got keys {sorted(kwargs)}"
527
+ )
528
+ return cls(
529
+ block="loop",
530
+ unrolled_pos=unrolled_position(
531
+ int(kwargs["loop_step"]), int(kwargs["layer_idx"]), num_layers_in_stack
532
+ ),
533
+ # Read off the tensors, which are the only source that exists here.
534
+ num_experts=int(logits.shape[-1]),
535
+ top_k=int(topk_idx.shape[-1]),
536
+ **kwargs,
537
+ )
538
+
539
+ return make
step_00028000/manifest.json ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation_checkpoint_mode": "selective",
3
+ "arch": "UL2",
4
+ "autocast": {
5
+ "enabled": false
6
+ },
7
+ "code_diff_sha256": null,
8
+ "code_dirty": false,
9
+ "code_revision_source": "job_start_env",
10
+ "compute_dtype": "bfloat16",
11
+ "data_position": {
12
+ "doc_idx_in_shard": 227593,
13
+ "dp_rank": 0,
14
+ "dp_world_size": 8,
15
+ "epoch": 0,
16
+ "global_sample_idx": 2688000,
17
+ "offset_in_shard_bytes": 952172140,
18
+ "seed": 42,
19
+ "shard_id": "shard_002",
20
+ "shard_sha256": "1f87dc91a903b6b1dd6994d5187a63d091d37384ed2be60eacc475ad282b15b6"
21
+ },
22
+ "epochs_completed": null,
23
+ "git_commit": "bfb53c6ae536fcc8159ac351d3ababe00eb6b4c5",
24
+ "kind": "trajectory",
25
+ "learning_rate": null,
26
+ "lr_schedule": {
27
+ "decay_ratio": 0.1,
28
+ "decay_type": "sqrt",
29
+ "min_lr_factor": 0.05,
30
+ "peak_lr": 0.0003,
31
+ "warmup_steps": 1000
32
+ },
33
+ "max_seq_len": null,
34
+ "model_config": {
35
+ "_name_or_path": "",
36
+ "architectures": [
37
+ "LoopMoEForCausalLM"
38
+ ],
39
+ "auto_map": {
40
+ "AutoConfig": "modeling_loop_lm.LoopMoEConfig",
41
+ "AutoModelForCausalLM": "modeling_loop_lm.LoopMoEForCausalLM"
42
+ },
43
+ "chunk_size_feed_forward": 0,
44
+ "context_length": 4096,
45
+ "d_ff": 4608,
46
+ "d_model": 1792,
47
+ "dtype": null,
48
+ "id2label": {
49
+ "0": "LABEL_0",
50
+ "1": "LABEL_1"
51
+ },
52
+ "is_encoder_decoder": false,
53
+ "label2id": {
54
+ "LABEL_0": 0,
55
+ "LABEL_1": 1
56
+ },
57
+ "lb_loss_factor": 0.01,
58
+ "lz_loss_factor": 0.001,
59
+ "model_type": "loop-moe",
60
+ "model_variant": "looped-moe",
61
+ "num_active": 2,
62
+ "num_experts": 16,
63
+ "num_head_layers": 1,
64
+ "num_heads": 28,
65
+ "num_layers": 16,
66
+ "num_layers_in_stack": 4,
67
+ "num_stacks": 4,
68
+ "num_tail_layers": 1,
69
+ "output_attentions": false,
70
+ "output_hidden_states": false,
71
+ "per_pass_attention": true,
72
+ "problem_type": null,
73
+ "return_dict": true,
74
+ "rope_theta": 10000.0,
75
+ "tie_word_embeddings": false,
76
+ "transformers_version": "5.15.0",
77
+ "unrolled_depth": 18,
78
+ "vocab_size": 49152,
79
+ "width_ratio": 14.0
80
+ },
81
+ "perturbation_probe": {
82
+ "runs_offline": true
83
+ },
84
+ "resume_checkpoint_interval_steps": 500,
85
+ "rng_state_saved": false,
86
+ "run_name": "UL2",
87
+ "seed": 42,
88
+ "step": 28000,
89
+ "steps_overridden": false,
90
+ "storage_dtype": "bfloat16",
91
+ "tokenizer_name": "smollm2",
92
+ "tokens_consumed": 11010048000,
93
+ "torch_version": "2.12.0+cu126",
94
+ "training_steps": 50000,
95
+ "transformers_version": "5.15.0"
96
+ }
step_00028000/modeling_loop_lm.py ADDED
@@ -0,0 +1,936 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Looped-MoE modeling code (port of the seedvar `loop-lm` architecture).
2
+
3
+ Provenance
4
+ ----------
5
+ Ported from ``modeling_loop_lm.py`` as published with
6
+ ``ml-ryanlee/seedvar-looped-moe-1e18-d704-seed42..47`` (arXiv 2605.09165,
7
+ *Sparse Layers are Critical to Scaling Looped Language Models*). That file is the
8
+ architecture specification: the published checkpoints are its training product,
9
+ and the diagnostics pipeline is already validated against it.
10
+
11
+ This is a **port, not a copy**. The numerics are kept faithful (see "Faithful to
12
+ the original" below) because the baseline's whole job is to reproduce the
13
+ published architecture; the deviations are all in service of three requirements
14
+ from HANDOVER §4.3 that the original file does not meet:
15
+
16
+ 1. **Semantic parameters have no defaults.** The original ``LoopLMConfig``
17
+ defaults ``d_model=1024``, ``num_experts=8`` and so on. Defaults are how a
18
+ silently-wrong run happens: a typo'd key name falls back to a plausible
19
+ number and the run looks fine. Every shape/semantic field here is required
20
+ and a missing one raises.
21
+ 2. **Every loop step is hookable, by explicit index.** The loop counter is
22
+ threaded down to each block, so a consumer never recovers the loop axis by
23
+ reshaping a flattened layer axis (which silently yields transposed
24
+ semantics).
25
+ 3. **Router logits are exposed with their semantics labelled.** Both the
26
+ pre-softmax logits and the post-softmax probabilities are handed out, tagged
27
+ via ``src.model.loop_trace.ROUTER_LOGITS_ARE_PRESOFTMAX``.
28
+
29
+ Scope
30
+ -----
31
+ Only the ``looped-moe`` variant is implemented. The original file carries four
32
+ variants (base / looped / moe / looped-moe); all four architectures this project
33
+ trains -- baseline and DVF-a/b/c -- are looped-moe, differing only in the (L, R, E)
34
+ triple. Porting the unused three would be dead code (project rule: no entities
35
+ beyond necessity). They remain available in the original file if ever needed.
36
+
37
+ Shape parameters, and what the experiment varies
38
+ ------------------------------------------------
39
+ ==================== ====== =========================================
40
+ config field symbol meaning
41
+ ==================== ====== =========================================
42
+ num_layers_in_stack L physical layers in the shared stack
43
+ num_stacks R times the stack is called (loop count)
44
+ num_experts E experts per MoE layer
45
+ num_active k experts activated per token
46
+ ==================== ====== =========================================
47
+
48
+ The DVF ("dual vector foil") series holds L*R = 16 and E*L = 64 fixed and only
49
+ moves capacity around: baseline (8,2,8), DVF-a (4,4,16), DVF-b (2,8,32),
50
+ DVF-c (1,16,64).
51
+
52
+ Faithful to the original (do not "fix" these -- they are muP, not bugs)
53
+ ----------------------------------------------------------------------
54
+ * ``RMSNorm`` has **no** gain parameter.
55
+ * Attention is scaled by ``1/d_k``, not ``1/sqrt(d_k)``.
56
+ * Softmax upcasts to float32.
57
+ * Expert FFN width is ``d_ff // num_active`` -- divided by k, *not* by E. This is
58
+ what makes E*L=64 hold total expert parameters constant across the DVF series.
59
+ * Initialisation is muP: ``std = std_base / sqrt(width_ratio)`` with
60
+ ``std_base = sqrt(2/(fan_in_base + fan_out_base))`` against a d_base=128 proxy.
61
+
62
+ Deliberate deviation
63
+ --------------------
64
+ ``RotaryPositionalEmbedding`` builds its rotation table with vectorised torch ops
65
+ instead of the original's ``max_seq_len * d_k/2`` nested Python loop, which costs
66
+ minutes at seq_len 4096. ``tests/test_rope_equivalence.py`` pins the vectorised
67
+ table against a literal transcription of the original loop.
68
+ """
69
+
70
+ from __future__ import annotations
71
+
72
+ import math
73
+ from typing import Any, Optional
74
+
75
+ import torch
76
+ import torch.nn as nn
77
+ import torch.nn.functional as F
78
+ from einops import einsum, rearrange, reduce, repeat
79
+ from torch import Tensor
80
+ from torch.nn.functional import grouped_mm, silu
81
+ from transformers import PretrainedConfig, PreTrainedModel
82
+ from transformers.generation import GenerationMixin
83
+ from transformers.modeling_outputs import CausalLMOutputWithPast
84
+
85
+ # This file is loaded in three different ways, and the trace types have to
86
+ # resolve in all of them:
87
+ # 1. as part of this repo -> `src.model.loop_trace`
88
+ # 2. via HuggingFace trust_remote_code -> copied into a generated package under
89
+ # `transformers_modules/<ckpt>/`, where the sibling is a *relative* import
90
+ # 3. as a loose script with the checkpoint directory on sys.path
91
+ # Case 2 is the one that matters for HANDOVER §4.7: the diagnostics pipeline is a
92
+ # separate repository that has never heard of `pretrain`, and `save_trajectory`
93
+ # bundles `loop_trace.py` next to this file so the checkpoint stands alone.
94
+ try:
95
+ from src.model.loop_trace import ACTIVE_TRACE_SINK, TraceMeta, TracePoint, TraceSink
96
+ except ImportError: # pragma: no cover - covered by the cold-load test
97
+ try:
98
+ from .loop_trace import ACTIVE_TRACE_SINK, TraceMeta, TracePoint, TraceSink # type: ignore[no-redef]
99
+ except ImportError:
100
+ from loop_trace import ACTIVE_TRACE_SINK, TraceMeta, TracePoint, TraceSink # type: ignore[no-redef]
101
+
102
+ __all__ = ["LoopMoEConfig", "LoopMoEForCausalLM", "LoopedMoETransformer"]
103
+
104
+ # muP proxy-model widths. The initialisation std of every weight is derived from
105
+ # a d_base=128 model and rescaled by width_ratio = d_model / 128.
106
+ HEAD_TAIL_NUM_EXPERTS = 8
107
+ """Experts in a head/tail layer, fixed across every ablation configuration.
108
+
109
+ Recipe 2026-09-12 section 2.1: head and tail are "identical in all configurations"
110
+ (1 layer, 8 experts, top-2). It is deliberately independent of the loop block's E --
111
+ S4 gives the loop 64 experts per layer and its head still has 8 -- so the loop block's
112
+ count must not be reused here.
113
+ """
114
+
115
+ BASE_D_MODEL = 128
116
+ BASE_D_FF = 384
117
+
118
+
119
+ def softmax(logits: Tensor, dim: int) -> Tensor:
120
+ """Max-shifted softmax in float32 (verbatim semantics from the original)."""
121
+ logits = logits.float()
122
+ max_values = torch.max(logits, dim=dim, keepdim=True).values
123
+ shifted = logits - max_values
124
+ shifted_exps = torch.exp(shifted)
125
+ shifted_exp_sums = torch.sum(shifted_exps, dim=dim, keepdim=True)
126
+ return shifted_exps / shifted_exp_sums
127
+
128
+
129
+ class Linear(nn.Module):
130
+ """Bias-free linear layer with muP initialisation."""
131
+
132
+ def __init__(self, in_features, out_features, width_ratio, std_base, device=None, dtype=None):
133
+ super().__init__()
134
+ # Registered before init so the shape exists under HF meta-device loading.
135
+ self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype, device=device))
136
+ # Kept so the init can be replayed: torchtitan builds the model on the
137
+ # meta device and then calls `init_weights()` on materialised (but
138
+ # uninitialised) storage, so constructor-time init alone leaves the
139
+ # model full of garbage.
140
+ self._init_std = std_base / math.sqrt(width_ratio)
141
+ self.reset_parameters()
142
+
143
+ def reset_parameters(self) -> None:
144
+ std = self._init_std
145
+ nn.init.trunc_normal_(self.weight, mean=0.0, std=std, a=-3 * std, b=3 * std)
146
+
147
+ def forward(self, x: Tensor) -> Tensor:
148
+ return einsum(self.weight, x, "d_out d_in, ... d_in -> ... d_out")
149
+
150
+
151
+ class Embedding(nn.Module):
152
+ def __init__(self, num_embeddings, embedding_dim, device=None, dtype=None):
153
+ super().__init__()
154
+ self.weight = nn.Parameter(torch.empty(num_embeddings, embedding_dim, dtype=dtype, device=device))
155
+ self.reset_parameters()
156
+
157
+ def reset_parameters(self) -> None:
158
+ nn.init.trunc_normal_(self.weight, mean=0.0, std=1.0, a=-3, b=3)
159
+
160
+ def forward(self, token_ids: Tensor) -> Tensor:
161
+ return self.weight[token_ids]
162
+
163
+
164
+ class RMSNorm(nn.Module):
165
+ """RMS norm **without** a gain parameter (muP convention)."""
166
+
167
+ def __init__(self, d_model: int, eps: float = 1e-5, device=None, dtype=None):
168
+ super().__init__()
169
+ self.d_model = d_model
170
+ self.eps = eps
171
+
172
+ def forward(self, x: Tensor) -> Tensor:
173
+ in_dtype = x.dtype
174
+ x = x.to(torch.float32)
175
+ mean_squared_sum = (1 / self.d_model) * einsum(x, x, "... seq d, ... seq d -> ... seq")
176
+ rms = torch.sqrt(mean_squared_sum + self.eps)
177
+ rms_norm = einsum(x, 1 / rms, "... seq d, ... seq -> ... seq d")
178
+ return rms_norm.to(in_dtype)
179
+
180
+
181
+ class PositionwiseFeedforward(nn.Module):
182
+ """SwiGLU: W2(SiLU(W1 x) * W3 x)."""
183
+
184
+ def __init__(self, d_model: int, d_ff: int, width_ratio: float, device=None, dtype=None):
185
+ super().__init__()
186
+ w_std_base = math.sqrt(2 / (BASE_D_MODEL + BASE_D_FF))
187
+ self.w1 = Linear(d_model, d_ff, width_ratio, w_std_base, device=device, dtype=dtype)
188
+ self.w2 = Linear(d_ff, d_model, width_ratio, w_std_base, device=device, dtype=dtype)
189
+ self.w3 = Linear(d_model, d_ff, width_ratio, w_std_base, device=device, dtype=dtype)
190
+
191
+ def forward(self, x: Tensor) -> Tensor:
192
+ return self.w2(silu(self.w1(x)) * self.w3(x))
193
+
194
+
195
+ class RotaryPositionalEmbedding(nn.Module):
196
+ """RoPE with a precomputed [seq, d_k/2, 2, 2] rotation table.
197
+
198
+ Vectorised rebuild of the original's nested Python loop; see the module
199
+ docstring and ``tests/test_rope_equivalence.py``.
200
+ """
201
+
202
+ def __init__(self, theta: float, d_k: int, max_seq_len: int, device=None, dtype=None):
203
+ super().__init__()
204
+ # Retained so the table can be rebuilt: `to_empty()` replaces buffer
205
+ # storage with uninitialised memory just as it does for parameters, so a
206
+ # meta-device build leaves the rotation table as garbage unless
207
+ # `reset_parameters()` regenerates it.
208
+ self._rope_theta, self._rope_d_k = theta, d_k
209
+ self._rope_max_seq_len, self._rope_dtype = max_seq_len, dtype
210
+ rotations = self._build_table(theta, d_k, max_seq_len, device, dtype)
211
+ self.register_buffer("rotations", rotations, persistent=True)
212
+
213
+ @staticmethod
214
+ def _build_table(theta: float, d_k: int, max_seq_len: int, device, dtype) -> Tensor:
215
+ """[seq, d_k/2, 2, 2] rotation table.
216
+
217
+ Angles are built in float64 and only then cast down. At seq_len 4096 the
218
+ largest angle is ~4096 rad, where float32 spacing is ~2.4e-4; computing
219
+ cos/sin at float32 there loses ~4 decimal digits. The original does this
220
+ implicitly (Python floats are float64), so float64 here is both more
221
+ accurate and what keeps the table equal to the reference.
222
+ """
223
+ positions = torch.arange(max_seq_len, device=device, dtype=torch.float64)
224
+ pair_idx = torch.arange(d_k // 2, device=device, dtype=torch.float64)
225
+ inv_freq = theta ** (2 * pair_idx / d_k)
226
+ angles = positions[:, None] / inv_freq[None, :]
227
+ cos, sin = torch.cos(angles), torch.sin(angles)
228
+ # rows of the 2x2 rotation: [[cos, -sin], [sin, cos]]
229
+ table = torch.stack(
230
+ [torch.stack([cos, -sin], dim=-1), torch.stack([sin, cos], dim=-1)], dim=-2
231
+ )
232
+ return table.to(dtype if dtype is not None else torch.float32)
233
+
234
+ @torch.no_grad()
235
+ def reset_parameters(self) -> None:
236
+ """Regenerate the rotation table in place (buffers survive nothing)."""
237
+ self.rotations.copy_(
238
+ self._build_table(
239
+ self._rope_theta, self._rope_d_k, self._rope_max_seq_len,
240
+ self.rotations.device, self.rotations.dtype,
241
+ )
242
+ )
243
+
244
+ def forward(self, x: Tensor, token_positions: Tensor) -> Tensor:
245
+ rot = self.rotations[token_positions].to(dtype=x.dtype)
246
+ x_pairs = rearrange(x, "... seq_dim (feature_dim i) -> ... seq_dim feature_dim i", i=2)
247
+ y_pairs = einsum(
248
+ rot,
249
+ x_pairs,
250
+ "... seq_dim feature_dim i j, ... seq_dim feature_dim j -> ... seq_dim feature_dim i",
251
+ )
252
+ return rearrange(y_pairs, "... seq_dim feature_dim i -> ... seq_dim (feature_dim i)")
253
+
254
+
255
+ class MultiheadSelfAttention(nn.Module):
256
+ """Causal MHSA with RoPE. muP: attention logits scaled by 1/d_k."""
257
+
258
+ def __init__(self, d_model: int, num_heads: int, max_seq_len: int, theta: float,
259
+ width_ratio: float, device=None, dtype=None):
260
+ super().__init__()
261
+ if d_model % num_heads != 0:
262
+ raise ValueError(f"d_model ({d_model}) must be divisible by num_heads ({num_heads})")
263
+ self.d_model = d_model
264
+ self.num_heads = num_heads
265
+
266
+ attn_std_base = math.sqrt(2 / (BASE_D_MODEL + BASE_D_MODEL))
267
+ self.q_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
268
+ self.k_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
269
+ self.v_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
270
+ self.output_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
271
+ self.rope = RotaryPositionalEmbedding(theta, d_model // num_heads, max_seq_len, device, dtype)
272
+
273
+ def forward(self, x: Tensor, token_positions: Optional[Tensor] = None) -> Tensor:
274
+ d_k = self.d_model // self.num_heads
275
+ q_heads = rearrange(self.q_proj(x), "... seq (heads d_k) -> ... heads seq d_k", d_k=d_k)
276
+ k_heads = rearrange(self.k_proj(x), "... seq (heads d_k) -> ... heads seq d_k", d_k=d_k)
277
+ v_heads = rearrange(self.v_proj(x), "... seq (heads d_v) -> ... heads seq d_v", d_v=d_k)
278
+
279
+ if token_positions is None:
280
+ token_positions = rearrange(torch.arange(x.shape[-2], device=x.device), "seq -> 1 seq")
281
+ q_heads = self.rope(q_heads, token_positions)
282
+ k_heads = self.rope(k_heads, token_positions)
283
+
284
+ mha_heads = F.scaled_dot_product_attention(
285
+ q_heads, k_heads, v_heads, is_causal=True, scale=1.0 / d_k
286
+ )
287
+ return self.output_proj(rearrange(mha_heads, "... heads seq d_v -> ... seq (heads d_v)"))
288
+
289
+
290
+ class Router(nn.Module):
291
+ """Top-k softmax router. Returns pre-softmax logits *and* probabilities.
292
+
293
+ The two are returned side by side, and the caller labels which is which via
294
+ ``ROUTER_LOGITS_ARE_PRESOFTMAX``. There is no jitter noise and no temperature;
295
+ routing is deterministic given the input (matching the original).
296
+ """
297
+
298
+ def __init__(self, d_model: int, num_experts: int, num_active: int, width_ratio: float,
299
+ device=None, dtype=None):
300
+ super().__init__()
301
+ std_base = math.sqrt(2 / (BASE_D_MODEL + num_experts))
302
+ self.gate = Linear(d_model, num_experts, width_ratio, std_base, device=device, dtype=dtype)
303
+ self.num_active = num_active
304
+
305
+ def forward(self, x: Tensor) -> tuple[Tensor, Tensor, Tensor, Tensor]:
306
+ logits = self.gate(x) # [B, S, E] -- pre-softmax
307
+ probs = softmax(logits, dim=-1) # [B, S, E] -- over all E experts
308
+ top_scores, top_experts = torch.topk(probs, k=self.num_active, dim=-1)
309
+ # Renormalise within the selected set so the combine weights sum to 1.
310
+ top_scores = top_scores / torch.sum(top_scores, dim=-1, keepdim=True)
311
+ return logits, probs, top_scores, top_experts
312
+
313
+
314
+ class GroupedMoEPrenormBlock(nn.Module):
315
+ """Pre-norm block whose FFN is a grouped top-k MoE.
316
+
317
+ Layout: x -> +attn(ln1(x)) -> +moe(ln2(.)). Aux losses are returned rather
318
+ than stashed on the module, so nothing has to be reset between loop steps.
319
+ """
320
+
321
+ @staticmethod
322
+ def _init_expert_weights(num_experts, in_features, out_features, width_ratio, std_base,
323
+ device, dtype) -> nn.Parameter:
324
+ w = torch.empty(num_experts, in_features, out_features, device=device, dtype=dtype)
325
+ std_scaled = std_base / math.sqrt(width_ratio)
326
+ nn.init.trunc_normal_(w, mean=0.0, std=std_scaled, a=-3 * std_scaled, b=3 * std_scaled)
327
+ return nn.Parameter(w)
328
+
329
+ @torch.no_grad()
330
+ def reset_parameters(self) -> None:
331
+ """Re-init the grouped expert weights (see Linear.reset_parameters)."""
332
+ std = self._expert_init_std
333
+ for w in (self.experts_w1, self.experts_w2, self.experts_w3):
334
+ nn.init.trunc_normal_(w, mean=0.0, std=std, a=-3 * std, b=3 * std)
335
+
336
+ def __init__(self, d_model: int, num_heads: int, d_ff: int, num_experts: int, num_active: int,
337
+ max_seq_len: int, theta: float, width_ratio: float, device=None, dtype=None,
338
+ num_attention_sets: int = 1):
339
+ """`num_attention_sets` = how many independent attention parameter sets
340
+ this block holds, one per pass over the looped stack.
341
+
342
+ 1 is the shared-attention architecture every run before 2026-09-18 used.
343
+ It keeps `self.attn` a single module, so the state dict key stays
344
+ `attn.q_proj.weight` and existing checkpoints load unchanged -- a
345
+ ModuleList of one would rename every key to `attn.0.*` and silently
346
+ invalidate every archive we hold.
347
+
348
+ >1 makes `self.attn` a ModuleList of that many attention modules, each
349
+ initialised independently under the same muP standard deviation, chosen
350
+ at forward time by `loop_step`. Everything else in the block -- norms,
351
+ router, experts -- remains shared across passes, which is the point of
352
+ the comparison.
353
+ """
354
+ super().__init__()
355
+ if num_attention_sets < 1:
356
+ raise ValueError(f"num_attention_sets must be >= 1, got {num_attention_sets}")
357
+ self.num_attention_sets = num_attention_sets
358
+ self.ln1 = RMSNorm(d_model, device=device, dtype=dtype)
359
+ _make_attn = lambda: MultiheadSelfAttention( # noqa: E731 -- one expression, used twice
360
+ d_model, num_heads, max_seq_len, theta, width_ratio, device, dtype
361
+ )
362
+ self.attn = (
363
+ _make_attn() if num_attention_sets == 1
364
+ else nn.ModuleList([_make_attn() for _ in range(num_attention_sets)])
365
+ )
366
+ self.ln2 = RMSNorm(d_model, device=device, dtype=dtype)
367
+ self.router = Router(d_model, num_experts, num_active, width_ratio, device=device, dtype=dtype)
368
+
369
+ self.num_experts = num_experts
370
+ self.num_active = num_active
371
+
372
+ # NOTE: divided by num_active (k), not by num_experts (E). This is what
373
+ # keeps total expert parameters constant across the DVF series.
374
+ d_ff_expert = d_ff // num_active
375
+ w_std_base = math.sqrt(2 / (BASE_D_MODEL + BASE_D_FF))
376
+ self._expert_init_std = w_std_base / math.sqrt(width_ratio)
377
+ self.experts_w1 = self._init_expert_weights(num_experts, d_model, d_ff_expert, width_ratio, w_std_base, device, dtype)
378
+ self.experts_w2 = self._init_expert_weights(num_experts, d_ff_expert, d_model, width_ratio, w_std_base, device, dtype)
379
+ self.experts_w3 = self._init_expert_weights(num_experts, d_model, d_ff_expert, width_ratio, w_std_base, device, dtype)
380
+
381
+ def attention_for(self, loop_step: Optional[int]):
382
+ """The attention module this pass uses.
383
+
384
+ `loop_step` is the explicit loop variable of
385
+ `LoopedMoETransformer.forward` -- the index of the current pass over the
386
+ shared stack -- threaded down as an argument. This block keeps **no
387
+ cross-call state**: nothing here remembers which pass ran last, and two
388
+ calls with the same `loop_step` select the same module. That is why the
389
+ looped model can be traced, resumed and re-entered in any order.
390
+
391
+ There is deliberately **no loop-step embedding** (owner's ruling
392
+ 2026-09-18): the passes differ only by their attention parameters, which
393
+ is what keeps U1-U4 a controlled comparison against S1-S4.
394
+
395
+ Refuses rather than defaults when `loop_step` is missing on a per-pass
396
+ block: falling back to set 0 would run every pass through the first
397
+ pass's attention, which trains, converges and reports nothing unusual
398
+ while being a different architecture from the one on the config.
399
+ """
400
+ if self.num_attention_sets == 1:
401
+ return self.attn
402
+ assert self.num_attention_sets == len(self.attn), (
403
+ f"{self.num_attention_sets} sets declared but {len(self.attn)} modules held; "
404
+ "the count and the container have drifted apart"
405
+ )
406
+ if loop_step is None:
407
+ raise ValueError(
408
+ "this block has per-pass attention "
409
+ f"({self.num_attention_sets} sets) and needs loop_step to choose one; "
410
+ "got None. Every caller in this file passes it -- a new caller must too."
411
+ )
412
+ if not 0 <= loop_step < len(self.attn):
413
+ raise IndexError(
414
+ f"loop_step {loop_step} is outside the {self.num_attention_sets} attention "
415
+ "sets this block holds. The block is built with one set per pass, so a "
416
+ "loop_step beyond that means the model and the config disagree about R."
417
+ )
418
+ return self.attn[loop_step]
419
+
420
+ def forward(
421
+ self,
422
+ x: Tensor,
423
+ token_positions: Optional[Tensor] = None,
424
+ *,
425
+ loop_step: Optional[int] = None,
426
+ layer_idx: Optional[int] = None,
427
+ block: Optional[str] = None,
428
+ unrolled_pos: Optional[int] = None,
429
+ trace_sink: Optional[TraceSink] = None,
430
+ ) -> tuple[Tensor, Tensor, Tensor]:
431
+ batch, seq, dim = x.shape
432
+ total_tokens = batch * seq
433
+
434
+ norm1_out = self.ln1(x)
435
+ attn_out = self.attention_for(loop_step)(norm1_out, token_positions)
436
+ assert x.shape == attn_out.shape
437
+ resid1_out = attn_out + x
438
+
439
+ norm2_out = self.ln2(resid1_out)
440
+ logits, probs, top_scores, top_experts = self.router(norm2_out)
441
+
442
+ # `softmax` computes in float32 and does not cast back, so `top_scores`
443
+ # is float32 regardless of the activation dtype. The combine weights get
444
+ # multiplied into bf16 expert outputs below, so they must match, or
445
+ # einsum raises "expected m1 and m2 to have the same dtype".
446
+ #
447
+ # This is invisible in the original: seedvar's published checkpoints are
448
+ # float32, where the cast is a no-op. Under the bf16 training this
449
+ # project uses (LT2 runs pure bf16, HANDOVER §4.1) it is a hard failure
450
+ # on the first forward.
451
+ #
452
+ # Only the combine weights are cast. `probs` and `logits` stay float32
453
+ # for the aux-loss and z-loss reductions, which is where the extra
454
+ # precision is worth having.
455
+ top_scores = top_scores.to(x.dtype)
456
+
457
+ # Flatten and sort by expert so grouped_mm can run one matmul per expert.
458
+ x_flat = rearrange(norm2_out, "b s d -> (b s) d")
459
+ flat_expert_ids = rearrange(top_experts, "b s k -> (b s k)")
460
+ flat_scores = rearrange(top_scores, "b s k -> (b s k)")
461
+ flat_positions = torch.arange(total_tokens, device=x.device)
462
+ flat_token_ids = repeat(flat_positions, "n -> (n k)", k=self.num_active)
463
+
464
+ sort_indices = flat_expert_ids.argsort(stable=True)
465
+ sorted_expert_ids = flat_expert_ids[sort_indices]
466
+ sorted_token_ids = flat_token_ids[sort_indices]
467
+ sorted_scores = flat_scores[sort_indices]
468
+ sorted_x = x_flat[sorted_token_ids]
469
+
470
+ counts = torch.bincount(sorted_expert_ids, minlength=self.num_experts)
471
+ offs = counts.cumsum(0).to(torch.int32)
472
+
473
+ h1 = grouped_mm(sorted_x, self.experts_w1, offs=offs)
474
+ h3 = grouped_mm(sorted_x, self.experts_w3, offs=offs)
475
+ gated = silu(h1) * h3
476
+ expert_out = grouped_mm(gated, self.experts_w2, offs=offs)
477
+
478
+ expert_out = einsum(expert_out, sorted_scores, "n d, n -> n d")
479
+ output_flat = torch.zeros(total_tokens, dim, device=x.device, dtype=expert_out.dtype)
480
+ output_flat.index_add_(0, sorted_token_ids, expert_out)
481
+ experts_out = rearrange(output_flat, "(b s) d -> b s d", b=batch, s=seq)
482
+
483
+ # Aux losses, per HANDOVER §3.3':
484
+ # L_LB = E * sum_i f_i * p_i (switch-style load balancing)
485
+ # L_RZ = mean( (logsumexp logits)^2 ) (router z-loss)
486
+ # Both are computed per layer per loop step; the caller averages over the
487
+ # unrolled depth (num_stacks * num_layers_in_stack).
488
+ fi = counts.float() / (total_tokens * self.num_active)
489
+ pi = reduce(probs, "b s e -> e", "mean")
490
+ lb = self.num_experts * einsum(fi, pi, "e, e ->")
491
+
492
+ logsumexp = torch.logsumexp(logits.float(), dim=-1)
493
+ lz = reduce(logsumexp**2, "... -> ", "mean")
494
+
495
+ assert experts_out.shape == resid1_out.shape
496
+ final_out = resid1_out + experts_out
497
+
498
+ # Write into every sink that is armed. Under FSDP2 the `trace_sink`
499
+ # keyword arrives as a per-block COPY (see `_ActiveTraceSink`), so the
500
+ # module-level one is the only sink the trainer can actually read back;
501
+ # the keyword remains for callers that pass their own dict directly
502
+ # (the toy launcher, the offline probes), where it is the same object.
503
+ # Writing to both is harmless: each is first-write-wins.
504
+ sinks = [s for s in (ACTIVE_TRACE_SINK.sink, trace_sink) if s is not None]
505
+ if sinks:
506
+ if loop_step is None or layer_idx is None or block is None or unrolled_pos is None:
507
+ raise ValueError(
508
+ "trace_sink was provided but loop_step/layer_idx/block/unrolled_pos "
509
+ "were not. Both axes must be explicit counters, never inferred."
510
+ )
511
+ if block not in ("head", "loop", "tail"):
512
+ raise ValueError(f"block must be head/loop/tail, got {block!r}")
513
+ # Keyed by the depth coordinate: head, the loop's first layer and tail all
514
+ # carry loop_step=layer_idx=0, so the old pair-key silently collapsed them.
515
+ key = unrolled_pos
516
+ # First write wins. Under selective activation checkpointing the
517
+ # block's forward runs a second time during backward to recompute
518
+ # activations, so every key is legitimately visited twice -- that is
519
+ # how AC works, not a bug. The recomputed values are identical by
520
+ # construction, so keeping the first and ignoring the rest is both
521
+ # correct and cheap.
522
+ #
523
+ # An earlier version raised on the second visit. That guard was aimed
524
+ # at double-*counting*, which first-write-wins prevents directly; as
525
+ # written it instead killed every AC-enabled run at the first traced
526
+ # step. `test_ac_recomputation_does_not_disturb_the_trace` pins the
527
+ # property that actually matters: same keys, same values, with AC on.
528
+ # Detached views, not copies -- see the lifetime contract in
529
+ # src/model/loop_trace.py.
530
+ point = TracePoint(
531
+ block=block,
532
+ unrolled_pos=unrolled_pos,
533
+ # From the tensors themselves, never from the model config: this layer's
534
+ # E and k are what produced these numbers, and head/tail differ from the
535
+ # loop block.
536
+ num_experts=int(logits.shape[-1]),
537
+ top_k=int(top_experts.shape[-1]),
538
+ loop_step=loop_step,
539
+ layer_idx=layer_idx,
540
+ router_logits=logits.detach(),
541
+ router_probs=probs.detach(),
542
+ topk_idx=top_experts.detach(),
543
+ topk_weights=top_scores.detach(),
544
+ residual=final_out.detach(),
545
+ )
546
+ for sink in sinks:
547
+ sink.setdefault(key, point)
548
+
549
+ return final_out, lb, lz
550
+
551
+
552
+ class LoopedStack(nn.Module):
553
+ """The stack of L MoE blocks that gets called R times."""
554
+
555
+ def __init__(self, context_length: int, d_model: int, num_layers_in_stack: int, num_heads: int,
556
+ d_ff: int, rope_theta: float, width_ratio: float, num_experts: int,
557
+ num_active: int, device=None, dtype=None, num_attention_sets: int = 1):
558
+ super().__init__()
559
+ #: One attention set per pass when per-pass attention is on, 1 otherwise.
560
+ #: Only the LOOPED layers get this: head and tail run once, so "per pass"
561
+ #: has no meaning for them and they keep a single set in every variant.
562
+ self.num_attention_sets = num_attention_sets
563
+ self.layers = nn.ModuleList(
564
+ [
565
+ GroupedMoEPrenormBlock(
566
+ d_model, num_heads, d_ff, num_experts, num_active,
567
+ context_length, rope_theta, width_ratio, device, dtype,
568
+ num_attention_sets=num_attention_sets,
569
+ )
570
+ for _ in range(num_layers_in_stack)
571
+ ]
572
+ )
573
+
574
+ def forward(
575
+ self,
576
+ x: Tensor,
577
+ *,
578
+ loop_step: int,
579
+ unrolled_pos_start: int,
580
+ trace_sink: Optional[TraceSink] = None,
581
+ ) -> tuple[Tensor, Tensor, Tensor]:
582
+ """`unrolled_pos_start` is the depth coordinate this call's first layer occupies.
583
+
584
+ Passed in rather than recomputed from `loop_step`, so the caller owns the depth
585
+ axis in one place: the stack does not need to know how many layers ran before it.
586
+ """
587
+ lb_total = x.new_zeros(())
588
+ lz_total = x.new_zeros(())
589
+ for layer_idx, layer in enumerate(self.layers):
590
+ x, lb, lz = layer(
591
+ x, loop_step=loop_step, layer_idx=layer_idx, block="loop",
592
+ unrolled_pos=unrolled_pos_start + layer_idx, trace_sink=trace_sink,
593
+ )
594
+ lb_total = lb_total + lb
595
+ lz_total = lz_total + lz
596
+ return x, lb_total, lz_total
597
+
598
+
599
+ class LoopedMoETransformer(nn.Module):
600
+ """Looped MoE transformer: one shared stack applied ``num_stacks`` times.
601
+
602
+ The loop is an explicit Python ``for``; ``loop_step`` is the loop variable and
603
+ is threaded all the way down to each block. Nothing downstream ever has to
604
+ recover it from tensor shapes.
605
+ """
606
+
607
+ def __init__(self, vocab_size: int, context_length: int, d_model: int, num_layers_in_stack: int,
608
+ num_stacks: int, num_heads: int, d_ff: int, num_experts: int, num_active: int,
609
+ rope_theta: float, width_ratio: float, num_head_layers: int = 0,
610
+ num_tail_layers: int = 0, head_tail_num_experts: int = HEAD_TAIL_NUM_EXPERTS,
611
+ device=None, dtype=None, per_pass_attention: bool = False):
612
+ super().__init__()
613
+ self.num_stacks = num_stacks
614
+ self.num_layers_in_stack = num_layers_in_stack
615
+ self.total_layers = num_stacks * num_layers_in_stack
616
+ self.num_head_layers = num_head_layers
617
+ self.num_tail_layers = num_tail_layers
618
+ # The unrolled depth, and the denominator the aux losses are averaged over.
619
+ # Written as head + loop + tail rather than the recipe's shorthand "2 + D*L":
620
+ # S15 has two head and two tail layers, so a literal 2 would divide by the wrong
621
+ # number there -- and it would not fail, it would just make the auxiliary losses
622
+ # quietly larger than intended.
623
+ self.unrolled_depth = num_head_layers + self.total_layers + num_tail_layers
624
+
625
+ self.token_embeddings = Embedding(vocab_size, d_model, device=device, dtype=dtype)
626
+ # Head and tail are ordinary MoE layers that run once. They keep E=8/top-2
627
+ # regardless of the loop block's expert count (recipe section 2.1: "identical in
628
+ # every configuration"), so the loop block's E is deliberately not passed here.
629
+ make_outer = lambda: GroupedMoEPrenormBlock(
630
+ d_model, num_heads, d_ff, head_tail_num_experts, num_active,
631
+ context_length, rope_theta, width_ratio, device, dtype,
632
+ )
633
+ self.head_layers = nn.ModuleList([make_outer() for _ in range(num_head_layers)])
634
+ self.tail_layers = nn.ModuleList([make_outer() for _ in range(num_tail_layers)])
635
+ self.per_pass_attention = per_pass_attention
636
+ self.stack = LoopedStack(
637
+ context_length, d_model, num_layers_in_stack, num_heads, d_ff, rope_theta,
638
+ width_ratio, num_experts, num_active, device=device, dtype=dtype,
639
+ # R sets when on: the stack is entered `num_stacks` times, and each
640
+ # entry is what "a pass" means here.
641
+ num_attention_sets=num_stacks if per_pass_attention else 1,
642
+ )
643
+ self.ln_final = RMSNorm(d_model, device=device, dtype=dtype)
644
+ std_base_lm_head = math.sqrt(2 / (BASE_D_MODEL + vocab_size))
645
+ self.lm_head = Linear(d_model, vocab_size, width_ratio, std_base_lm_head, device=device, dtype=dtype)
646
+
647
+ @classmethod
648
+ def from_config(cls, config: "LoopMoEConfig", *, device=None, dtype=None) -> "LoopedMoETransformer":
649
+ """THE way to build this model from a config. Both call sites use it.
650
+
651
+ There were two: the HF wrapper and `pretrain/train_spec.py`, each with its own
652
+ hand-written keyword list. When head/tail layers were added, the training path's
653
+ list was not updated, so it silently built a model with no head or tail while its
654
+ config said otherwise -- it trained, the loss looked plausible, and every artifact
655
+ recorded the config's depth rather than the depth that ran. Nothing could raise,
656
+ because a shorter model is a perfectly valid model.
657
+
658
+ A single entry point makes that class of drift impossible rather than merely
659
+ tested-for: a field added to the config is read here once, and both paths get it.
660
+ """
661
+ return cls(
662
+ vocab_size=config.vocab_size,
663
+ context_length=config.context_length,
664
+ d_model=config.d_model,
665
+ num_layers_in_stack=config.num_layers_in_stack,
666
+ num_stacks=config.num_stacks,
667
+ num_heads=config.num_heads,
668
+ d_ff=config.d_ff,
669
+ num_experts=config.num_experts,
670
+ num_active=config.num_active,
671
+ rope_theta=config.rope_theta,
672
+ width_ratio=config.width_ratio,
673
+ num_head_layers=config.num_head_layers,
674
+ num_tail_layers=config.num_tail_layers,
675
+ # `getattr` with the config's own default, not a bare False: a config
676
+ # object built by older code has no such attribute, and this is the
677
+ # one entry point, so a silent False here would be the only place the
678
+ # variant could be lost.
679
+ per_pass_attention=getattr(config, "per_pass_attention", False),
680
+ device=device,
681
+ dtype=dtype,
682
+ )
683
+
684
+ def forward(
685
+ self,
686
+ x: Tensor,
687
+ *,
688
+ trace_sink: Optional[TraceSink] = None,
689
+ ) -> tuple[Tensor, Tensor, Tensor]:
690
+ lb_total = None
691
+ lz_total = None
692
+
693
+ x = self.token_embeddings(x)
694
+
695
+ def run_outer(layers, block: str, pos: int, x, lb_total, lz_total):
696
+ for i, layer in enumerate(layers):
697
+ x, lb, lz = layer(
698
+ x, loop_step=0, layer_idx=0, block=block, unrolled_pos=pos + i,
699
+ trace_sink=trace_sink,
700
+ )
701
+ lb_total = lb if lb_total is None else lb_total + lb
702
+ lz_total = lz if lz_total is None else lz_total + lz
703
+ return x, lb_total, lz_total
704
+
705
+ # head -> loop x num_stacks -> tail, with one running depth coordinate. head and
706
+ # tail record loop_step=layer_idx=0 because neither coordinate means anything
707
+ # outside the loop; `block` and `unrolled_pos` are what identifies them.
708
+ x, lb_total, lz_total = run_outer(self.head_layers, "head", 0, x, lb_total, lz_total)
709
+ pos = self.num_head_layers
710
+ for loop_step in range(self.num_stacks):
711
+ x, lb, lz = self.stack(
712
+ x, loop_step=loop_step, unrolled_pos_start=pos, trace_sink=trace_sink,
713
+ )
714
+ pos += self.num_layers_in_stack
715
+ lb_total = lb if lb_total is None else lb_total + lb
716
+ lz_total = lz if lz_total is None else lz_total + lz
717
+ x, lb_total, lz_total = run_outer(self.tail_layers, "tail", pos, x, lb_total, lz_total)
718
+
719
+ x = self.lm_head(self.ln_final(x))
720
+
721
+ # Averaged over the *unrolled* depth: every physical layer contributes once per
722
+ # loop step, and head/tail contribute once each (recipe section 2.1). Equals
723
+ # `total_layers` exactly when there are no head/tail layers, which is every
724
+ # pre-ablation configuration -- so their published losses are unchanged.
725
+ return x, lb_total / self.unrolled_depth, lz_total / self.unrolled_depth
726
+
727
+
728
+ def _require(kwargs: dict[str, Any], name: str) -> Any:
729
+ """Fetch a required config field or raise.
730
+
731
+ Project rule (HANDOVER §4.3 / §5.6): semantic parameters get no defaults.
732
+ A default is a silent-wrong-answer generator -- a mistyped or dropped key
733
+ becomes a plausible number instead of an error.
734
+ """
735
+ if name not in kwargs or kwargs[name] is None:
736
+ raise ValueError(
737
+ f"LoopMoEConfig: required field {name!r} is missing. Semantic "
738
+ "parameters have no defaults in this project; state it explicitly."
739
+ )
740
+ return kwargs.pop(name)
741
+
742
+
743
+ class LoopMoEConfig(PretrainedConfig):
744
+ """Config for the looped-MoE architecture. **Every field is required.**
745
+
746
+ Compatible with ``save_pretrained``/``from_pretrained``: a config.json written
747
+ by this class round-trips, and one is rejected loudly if a field is absent.
748
+ """
749
+
750
+ model_type = "loop-moe"
751
+
752
+ # Tells transformers not to introspect defaults by constructing `cls()` with
753
+ # no arguments -- which this class deliberately rejects. Without it,
754
+ # `save_pretrained` fails inside `_get_generation_parameters`. This is the
755
+ # supported escape hatch for configs whose fields are all required.
756
+ has_no_defaults_at_init = True
757
+
758
+ def __init__(self, **kwargs: Any):
759
+ # `from_pretrained` on a *torch-saved* config, and some HF-internal paths,
760
+ # construct with no arguments at all; only a fully-specified call is valid.
761
+ self.vocab_size = _require(kwargs, "vocab_size")
762
+ self.context_length = _require(kwargs, "context_length")
763
+ self.d_model = _require(kwargs, "d_model")
764
+ self.num_heads = _require(kwargs, "num_heads")
765
+ self.d_ff = _require(kwargs, "d_ff")
766
+ self.rope_theta = _require(kwargs, "rope_theta")
767
+ self.width_ratio = _require(kwargs, "width_ratio")
768
+ self.num_layers_in_stack = _require(kwargs, "num_layers_in_stack") # L
769
+ self.num_stacks = _require(kwargs, "num_stacks") # R
770
+ self.num_experts = _require(kwargs, "num_experts") # E
771
+ self.num_active = _require(kwargs, "num_active") # k
772
+ self.lb_loss_factor = _require(kwargs, "lb_loss_factor")
773
+ # Head/tail layers: MoE layers run once, outside the loop (ablation recipe 2026-09-12
774
+ # section 2.1). 0 means the architecture has none, which is not a guess -- it is what
775
+ # every configuration built before this recipe actually is, and
776
+ # `test_config_registry` pins their parameter counts as unchanged. Real ablation
777
+ # configs never rely on the fallback: `config_registry.build_ablation` states both
778
+ # counts for every entry, and a test asserts it does.
779
+ self.num_head_layers = int(kwargs.pop("num_head_layers", 0))
780
+ self.num_tail_layers = int(kwargs.pop("num_tail_layers", 0))
781
+
782
+ # Per-pass attention: each of the R passes over the shared stack gets its
783
+ # OWN attention parameters, while the MoE, the norms and the router stay
784
+ # shared. Default False, and defaulted rather than required precisely
785
+ # because every config written before 2026-09-18 lacks the key: those
786
+ # files must keep loading, and they must keep meaning what they meant.
787
+ # False reproduces the previous model bit for bit, including the state
788
+ # dict's key names -- see GroupedMoEPrenormBlock.
789
+ self.per_pass_attention = bool(kwargs.pop("per_pass_attention", False))
790
+ self.lz_loss_factor = _require(kwargs, "lz_loss_factor")
791
+
792
+ # Stated, not inherited. `PretrainedConfig` defaults this to True, and the
793
+ # only reason the embedding and the LM head are not already sharing storage
794
+ # is that this model never implemented `get_output_embeddings()`. The day
795
+ # someone adds it for tool compatibility, every configuration would start
796
+ # tying weights -- a different model, trained to a different loss, with
797
+ # nothing in any artifact saying so. The architecture uses untied weights
798
+ # (the parameter counts in the recipe assume it), so the config says so.
799
+ kwargs.pop("tie_word_embeddings", None)
800
+ self.tie_word_embeddings = False
801
+
802
+ self._validate()
803
+
804
+ # Derived, for readers; never an input. The unrolled depth now includes the
805
+ # layers that run once outside the loop.
806
+ self.num_layers = self.num_stacks * self.num_layers_in_stack
807
+ self.unrolled_depth = self.num_head_layers + self.num_layers + self.num_tail_layers
808
+
809
+ # The original config mirrored `context_length` into `max_length` for
810
+ # lm-evaluation-harness. transformers >=5 classifies `max_length` as a
811
+ # generation parameter and refuses to serialise a config carrying one
812
+ # (the check is `hasattr`, so even a property trips it). `context_length`
813
+ # is therefore the single source of truth for sequence length; pass
814
+ # `max_length` to the harness explicitly at eval time instead.
815
+ # Popped so that loading a seedvar-era config.json cannot reintroduce it.
816
+ kwargs.pop("max_length", None)
817
+
818
+ super().__init__(**kwargs)
819
+
820
+ def _validate(self) -> None:
821
+ """Reject out-of-domain values loudly rather than failing deep in a kernel."""
822
+ positive = (
823
+ "vocab_size", "context_length", "d_model", "num_heads", "d_ff",
824
+ "num_layers_in_stack", "num_stacks", "num_experts", "num_active",
825
+ )
826
+ for name in positive:
827
+ value = getattr(self, name)
828
+ if not isinstance(value, int) or value < 1:
829
+ raise ValueError(f"LoopMoEConfig.{name} must be a positive int, got {value!r}")
830
+ if self.d_model % self.num_heads != 0:
831
+ raise ValueError(
832
+ f"d_model ({self.d_model}) must be divisible by num_heads ({self.num_heads})"
833
+ )
834
+ if self.num_active > self.num_experts:
835
+ raise ValueError(
836
+ f"num_active ({self.num_active}) cannot exceed num_experts ({self.num_experts})"
837
+ )
838
+ if self.d_ff % self.num_active != 0:
839
+ raise ValueError(
840
+ f"d_ff ({self.d_ff}) must be divisible by num_active ({self.num_active}); "
841
+ "expert width is d_ff // num_active and truncation would silently "
842
+ "change the parameter count."
843
+ )
844
+ for name in ("num_head_layers", "num_tail_layers"):
845
+ value = getattr(self, name)
846
+ if not isinstance(value, int) or value < 0:
847
+ raise ValueError(f"LoopMoEConfig.{name} must be a non-negative int, got {value!r}")
848
+ if (self.d_model // self.num_heads) % 2 != 0:
849
+ raise ValueError(
850
+ f"head dim ({self.d_model // self.num_heads}) must be even for RoPE"
851
+ )
852
+
853
+ @property
854
+ def trace_meta(self) -> TraceMeta:
855
+ """Shape/provenance block handed to the metrics collector."""
856
+ return TraceMeta(
857
+ num_stacks=self.num_stacks,
858
+ num_layers_in_stack=self.num_layers_in_stack,
859
+ num_experts=self.num_experts,
860
+ num_active=self.num_active,
861
+ )
862
+
863
+
864
+ class LoopMoEForCausalLM(PreTrainedModel, GenerationMixin):
865
+ """HF-compatible causal LM wrapper.
866
+
867
+ Kept HF-shaped on purpose: HANDOVER §4.7 makes "the diagnostics pipeline
868
+ ingests our checkpoints unchanged" an acceptance criterion, and that pipeline
869
+ loads models through ``from_pretrained``.
870
+ """
871
+
872
+ config_class = LoopMoEConfig
873
+
874
+ def __init__(self, config: LoopMoEConfig):
875
+ super().__init__(config)
876
+ self.model = LoopedMoETransformer.from_config(config)
877
+ self.post_init()
878
+
879
+ def get_input_embeddings(self):
880
+ return self.model.token_embeddings
881
+
882
+ def set_input_embeddings(self, value):
883
+ self.model.token_embeddings = value
884
+
885
+ def forward(
886
+ self,
887
+ input_ids: torch.LongTensor,
888
+ attention_mask: Optional[Tensor] = None, # unused: the mask is built in
889
+ labels: Optional[torch.LongTensor] = None,
890
+ trace_sink: Optional[TraceSink] = None,
891
+ **kwargs: Any,
892
+ ) -> CausalLMOutputWithPast:
893
+ """Forward pass.
894
+
895
+ Returns a ``CausalLMOutputWithPast`` whose ``loss`` is the *total* loss
896
+ (CE + weighted aux). The unweighted components are attached as
897
+ ``task_loss`` / ``lb_loss`` / ``z_loss`` so the training loop can log the
898
+ breakdown without recomputing anything.
899
+
900
+ **Label contract (documented here 2026-08-21, `ABCI_ERR_20260821_0405_
901
+ gate1_label_shift_root_cause.md`)**: ``labels`` must already be
902
+ next-token-shifted by the caller -- ``labels[..., t] == input_ids[..., t+1]``,
903
+ with the last position set to ``-100`` (no target exists after it). This
904
+ method does **not** shift internally; it passes ``labels`` to
905
+ ``F.cross_entropy`` exactly as given. Before this date the only written
906
+ record of this contract was a comment in
907
+ ``src/model/loop_trace.py`` ("the dataloader supplies pre-shifted
908
+ labels during training, so the shift is explicit here") -- not here, at
909
+ the definition itself. That gap let four separate call sites
910
+ (the pretraining dataloader path aside, which was correct) independently
911
+ get this wrong the same way, rather than it being four unrelated
912
+ mistakes. Any caller not shifting first -- e.g. ``model(input_ids=ids,
913
+ labels=ids)`` -- silently trains/evaluates on the trivial
914
+ copy-the-current-token target instead of next-token prediction.
915
+ """
916
+ logits, lb, lz = self.model(input_ids, trace_sink=trace_sink)
917
+
918
+ loss = task_loss = None
919
+ if labels is not None:
920
+ task_loss = F.cross_entropy(logits.view(-1, logits.size(-1)), labels.view(-1))
921
+ loss = (
922
+ task_loss
923
+ + self.config.lb_loss_factor * lb
924
+ + self.config.lz_loss_factor * lz
925
+ )
926
+
927
+ out = CausalLMOutputWithPast(loss=loss, logits=logits)
928
+ # Unweighted components; the trainer pairs them with the factors from
929
+ # config to build LossComponents.
930
+ out.task_loss = task_loss
931
+ out.lb_loss = lb
932
+ out.z_loss = lz
933
+ return out
934
+
935
+ def prepare_inputs_for_generation(self, input_ids, **kwargs):
936
+ return {"input_ids": input_ids}
step_00028000/pytorch_model.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f470809cd2e0f4b5009b468ae669932b2a1090214d2727ee0459f9e2e5a25958
3
+ size 2826236439