File size: 41,714 Bytes
30a4470
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
"""Looped-MoE modeling code (port of the seedvar `loop-lm` architecture).

Provenance
----------
Ported from ``modeling_loop_lm.py`` as published with
``ml-ryanlee/seedvar-looped-moe-1e18-d704-seed42..47`` (arXiv 2605.09165,
*Sparse Layers are Critical to Scaling Looped Language Models*). That file is the
architecture specification: the published checkpoints are its training product,
and the diagnostics pipeline is already validated against it.

This is a **port, not a copy**. The numerics are kept faithful (see "Faithful to
the original" below) because the baseline's whole job is to reproduce the
published architecture; the deviations are all in service of three requirements
from HANDOVER §4.3 that the original file does not meet:

  1. **Semantic parameters have no defaults.** The original ``LoopLMConfig``
     defaults ``d_model=1024``, ``num_experts=8`` and so on. Defaults are how a
     silently-wrong run happens: a typo'd key name falls back to a plausible
     number and the run looks fine. Every shape/semantic field here is required
     and a missing one raises.
  2. **Every loop step is hookable, by explicit index.** The loop counter is
     threaded down to each block, so a consumer never recovers the loop axis by
     reshaping a flattened layer axis (which silently yields transposed
     semantics).
  3. **Router logits are exposed with their semantics labelled.** Both the
     pre-softmax logits and the post-softmax probabilities are handed out, tagged
     via ``src.model.loop_trace.ROUTER_LOGITS_ARE_PRESOFTMAX``.

Scope
-----
Only the ``looped-moe`` variant is implemented. The original file carries four
variants (base / looped / moe / looped-moe); all four architectures this project
trains -- baseline and DVF-a/b/c -- are looped-moe, differing only in the (L, R, E)
triple. Porting the unused three would be dead code (project rule: no entities
beyond necessity). They remain available in the original file if ever needed.

Shape parameters, and what the experiment varies
------------------------------------------------
====================  ======  =========================================
config field          symbol  meaning
====================  ======  =========================================
num_layers_in_stack   L       physical layers in the shared stack
num_stacks            R       times the stack is called (loop count)
num_experts           E       experts per MoE layer
num_active            k       experts activated per token
====================  ======  =========================================

The DVF ("dual vector foil") series holds L*R = 16 and E*L = 64 fixed and only
moves capacity around: baseline (8,2,8), DVF-a (4,4,16), DVF-b (2,8,32),
DVF-c (1,16,64).

Faithful to the original (do not "fix" these -- they are muP, not bugs)
----------------------------------------------------------------------
* ``RMSNorm`` has **no** gain parameter.
* Attention is scaled by ``1/d_k``, not ``1/sqrt(d_k)``.
* Softmax upcasts to float32.
* Expert FFN width is ``d_ff // num_active`` -- divided by k, *not* by E. This is
  what makes E*L=64 hold total expert parameters constant across the DVF series.
* Initialisation is muP: ``std = std_base / sqrt(width_ratio)`` with
  ``std_base = sqrt(2/(fan_in_base + fan_out_base))`` against a d_base=128 proxy.

Deliberate deviation
--------------------
``RotaryPositionalEmbedding`` builds its rotation table with vectorised torch ops
instead of the original's ``max_seq_len * d_k/2`` nested Python loop, which costs
minutes at seq_len 4096. ``tests/test_rope_equivalence.py`` pins the vectorised
table against a literal transcription of the original loop.
"""

from __future__ import annotations

import math
from typing import Any, Optional

import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import einsum, rearrange, reduce, repeat
from torch import Tensor
from torch.nn.functional import grouped_mm, silu
from transformers import PretrainedConfig, PreTrainedModel
from transformers.generation import GenerationMixin
from transformers.modeling_outputs import CausalLMOutputWithPast

# This file is loaded in three different ways, and the trace types have to
# resolve in all of them:
#   1. as part of this repo             -> `src.model.loop_trace`
#   2. via HuggingFace trust_remote_code -> copied into a generated package under
#      `transformers_modules/<ckpt>/`, where the sibling is a *relative* import
#   3. as a loose script with the checkpoint directory on sys.path
# Case 2 is the one that matters for HANDOVER §4.7: the diagnostics pipeline is a
# separate repository that has never heard of `pretrain`, and `save_trajectory`
# bundles `loop_trace.py` next to this file so the checkpoint stands alone.
try:
    from src.model.loop_trace import ACTIVE_TRACE_SINK, TraceMeta, TracePoint, TraceSink
except ImportError:  # pragma: no cover - covered by the cold-load test
    try:
        from .loop_trace import ACTIVE_TRACE_SINK, TraceMeta, TracePoint, TraceSink  # type: ignore[no-redef]
    except ImportError:
        from loop_trace import ACTIVE_TRACE_SINK, TraceMeta, TracePoint, TraceSink  # type: ignore[no-redef]

__all__ = ["LoopMoEConfig", "LoopMoEForCausalLM", "LoopedMoETransformer"]

# muP proxy-model widths. The initialisation std of every weight is derived from
# a d_base=128 model and rescaled by width_ratio = d_model / 128.
HEAD_TAIL_NUM_EXPERTS = 8
"""Experts in a head/tail layer, fixed across every ablation configuration.

Recipe 2026-09-12 section 2.1: head and tail are "identical in all configurations"
(1 layer, 8 experts, top-2). It is deliberately independent of the loop block's E --
S4 gives the loop 64 experts per layer and its head still has 8 -- so the loop block's
count must not be reused here.
"""

BASE_D_MODEL = 128
BASE_D_FF = 384


def softmax(logits: Tensor, dim: int) -> Tensor:
    """Max-shifted softmax in float32 (verbatim semantics from the original)."""
    logits = logits.float()
    max_values = torch.max(logits, dim=dim, keepdim=True).values
    shifted = logits - max_values
    shifted_exps = torch.exp(shifted)
    shifted_exp_sums = torch.sum(shifted_exps, dim=dim, keepdim=True)
    return shifted_exps / shifted_exp_sums


class Linear(nn.Module):
    """Bias-free linear layer with muP initialisation."""

    def __init__(self, in_features, out_features, width_ratio, std_base, device=None, dtype=None):
        super().__init__()
        # Registered before init so the shape exists under HF meta-device loading.
        self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype, device=device))
        # Kept so the init can be replayed: torchtitan builds the model on the
        # meta device and then calls `init_weights()` on materialised (but
        # uninitialised) storage, so constructor-time init alone leaves the
        # model full of garbage.
        self._init_std = std_base / math.sqrt(width_ratio)
        self.reset_parameters()

    def reset_parameters(self) -> None:
        std = self._init_std
        nn.init.trunc_normal_(self.weight, mean=0.0, std=std, a=-3 * std, b=3 * std)

    def forward(self, x: Tensor) -> Tensor:
        return einsum(self.weight, x, "d_out d_in, ... d_in -> ... d_out")


class Embedding(nn.Module):
    def __init__(self, num_embeddings, embedding_dim, device=None, dtype=None):
        super().__init__()
        self.weight = nn.Parameter(torch.empty(num_embeddings, embedding_dim, dtype=dtype, device=device))
        self.reset_parameters()

    def reset_parameters(self) -> None:
        nn.init.trunc_normal_(self.weight, mean=0.0, std=1.0, a=-3, b=3)

    def forward(self, token_ids: Tensor) -> Tensor:
        return self.weight[token_ids]


class RMSNorm(nn.Module):
    """RMS norm **without** a gain parameter (muP convention)."""

    def __init__(self, d_model: int, eps: float = 1e-5, device=None, dtype=None):
        super().__init__()
        self.d_model = d_model
        self.eps = eps

    def forward(self, x: Tensor) -> Tensor:
        in_dtype = x.dtype
        x = x.to(torch.float32)
        mean_squared_sum = (1 / self.d_model) * einsum(x, x, "... seq d, ... seq d -> ... seq")
        rms = torch.sqrt(mean_squared_sum + self.eps)
        rms_norm = einsum(x, 1 / rms, "... seq d, ... seq -> ... seq d")
        return rms_norm.to(in_dtype)


class PositionwiseFeedforward(nn.Module):
    """SwiGLU: W2(SiLU(W1 x) * W3 x)."""

    def __init__(self, d_model: int, d_ff: int, width_ratio: float, device=None, dtype=None):
        super().__init__()
        w_std_base = math.sqrt(2 / (BASE_D_MODEL + BASE_D_FF))
        self.w1 = Linear(d_model, d_ff, width_ratio, w_std_base, device=device, dtype=dtype)
        self.w2 = Linear(d_ff, d_model, width_ratio, w_std_base, device=device, dtype=dtype)
        self.w3 = Linear(d_model, d_ff, width_ratio, w_std_base, device=device, dtype=dtype)

    def forward(self, x: Tensor) -> Tensor:
        return self.w2(silu(self.w1(x)) * self.w3(x))


class RotaryPositionalEmbedding(nn.Module):
    """RoPE with a precomputed [seq, d_k/2, 2, 2] rotation table.

    Vectorised rebuild of the original's nested Python loop; see the module
    docstring and ``tests/test_rope_equivalence.py``.
    """

    def __init__(self, theta: float, d_k: int, max_seq_len: int, device=None, dtype=None):
        super().__init__()
        # Retained so the table can be rebuilt: `to_empty()` replaces buffer
        # storage with uninitialised memory just as it does for parameters, so a
        # meta-device build leaves the rotation table as garbage unless
        # `reset_parameters()` regenerates it.
        self._rope_theta, self._rope_d_k = theta, d_k
        self._rope_max_seq_len, self._rope_dtype = max_seq_len, dtype
        rotations = self._build_table(theta, d_k, max_seq_len, device, dtype)
        self.register_buffer("rotations", rotations, persistent=True)

    @staticmethod
    def _build_table(theta: float, d_k: int, max_seq_len: int, device, dtype) -> Tensor:
        """[seq, d_k/2, 2, 2] rotation table.

        Angles are built in float64 and only then cast down. At seq_len 4096 the
        largest angle is ~4096 rad, where float32 spacing is ~2.4e-4; computing
        cos/sin at float32 there loses ~4 decimal digits. The original does this
        implicitly (Python floats are float64), so float64 here is both more
        accurate and what keeps the table equal to the reference.
        """
        positions = torch.arange(max_seq_len, device=device, dtype=torch.float64)
        pair_idx = torch.arange(d_k // 2, device=device, dtype=torch.float64)
        inv_freq = theta ** (2 * pair_idx / d_k)
        angles = positions[:, None] / inv_freq[None, :]
        cos, sin = torch.cos(angles), torch.sin(angles)
        # rows of the 2x2 rotation: [[cos, -sin], [sin, cos]]
        table = torch.stack(
            [torch.stack([cos, -sin], dim=-1), torch.stack([sin, cos], dim=-1)], dim=-2
        )
        return table.to(dtype if dtype is not None else torch.float32)

    @torch.no_grad()
    def reset_parameters(self) -> None:
        """Regenerate the rotation table in place (buffers survive nothing)."""
        self.rotations.copy_(
            self._build_table(
                self._rope_theta, self._rope_d_k, self._rope_max_seq_len,
                self.rotations.device, self.rotations.dtype,
            )
        )

    def forward(self, x: Tensor, token_positions: Tensor) -> Tensor:
        rot = self.rotations[token_positions].to(dtype=x.dtype)
        x_pairs = rearrange(x, "... seq_dim (feature_dim i) -> ... seq_dim feature_dim i", i=2)
        y_pairs = einsum(
            rot,
            x_pairs,
            "... seq_dim feature_dim i j, ... seq_dim feature_dim j -> ... seq_dim feature_dim i",
        )
        return rearrange(y_pairs, "... seq_dim feature_dim i -> ... seq_dim (feature_dim i)")


class MultiheadSelfAttention(nn.Module):
    """Causal MHSA with RoPE. muP: attention logits scaled by 1/d_k."""

    def __init__(self, d_model: int, num_heads: int, max_seq_len: int, theta: float,
                 width_ratio: float, device=None, dtype=None):
        super().__init__()
        if d_model % num_heads != 0:
            raise ValueError(f"d_model ({d_model}) must be divisible by num_heads ({num_heads})")
        self.d_model = d_model
        self.num_heads = num_heads

        attn_std_base = math.sqrt(2 / (BASE_D_MODEL + BASE_D_MODEL))
        self.q_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
        self.k_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
        self.v_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
        self.output_proj = Linear(d_model, d_model, width_ratio, attn_std_base, device=device, dtype=dtype)
        self.rope = RotaryPositionalEmbedding(theta, d_model // num_heads, max_seq_len, device, dtype)

    def forward(self, x: Tensor, token_positions: Optional[Tensor] = None) -> Tensor:
        d_k = self.d_model // self.num_heads
        q_heads = rearrange(self.q_proj(x), "... seq (heads d_k) -> ... heads seq d_k", d_k=d_k)
        k_heads = rearrange(self.k_proj(x), "... seq (heads d_k) -> ... heads seq d_k", d_k=d_k)
        v_heads = rearrange(self.v_proj(x), "... seq (heads d_v) -> ... heads seq d_v", d_v=d_k)

        if token_positions is None:
            token_positions = rearrange(torch.arange(x.shape[-2], device=x.device), "seq -> 1 seq")
        q_heads = self.rope(q_heads, token_positions)
        k_heads = self.rope(k_heads, token_positions)

        mha_heads = F.scaled_dot_product_attention(
            q_heads, k_heads, v_heads, is_causal=True, scale=1.0 / d_k
        )
        return self.output_proj(rearrange(mha_heads, "... heads seq d_v -> ... seq (heads d_v)"))


class Router(nn.Module):
    """Top-k softmax router. Returns pre-softmax logits *and* probabilities.

    The two are returned side by side, and the caller labels which is which via
    ``ROUTER_LOGITS_ARE_PRESOFTMAX``. There is no jitter noise and no temperature;
    routing is deterministic given the input (matching the original).
    """

    def __init__(self, d_model: int, num_experts: int, num_active: int, width_ratio: float,
                 device=None, dtype=None):
        super().__init__()
        std_base = math.sqrt(2 / (BASE_D_MODEL + num_experts))
        self.gate = Linear(d_model, num_experts, width_ratio, std_base, device=device, dtype=dtype)
        self.num_active = num_active

    def forward(self, x: Tensor) -> tuple[Tensor, Tensor, Tensor, Tensor]:
        logits = self.gate(x)  # [B, S, E] -- pre-softmax
        probs = softmax(logits, dim=-1)  # [B, S, E] -- over all E experts
        top_scores, top_experts = torch.topk(probs, k=self.num_active, dim=-1)
        # Renormalise within the selected set so the combine weights sum to 1.
        top_scores = top_scores / torch.sum(top_scores, dim=-1, keepdim=True)
        return logits, probs, top_scores, top_experts


class GroupedMoEPrenormBlock(nn.Module):
    """Pre-norm block whose FFN is a grouped top-k MoE.

    Layout: x -> +attn(ln1(x)) -> +moe(ln2(.)). Aux losses are returned rather
    than stashed on the module, so nothing has to be reset between loop steps.
    """

    @staticmethod
    def _init_expert_weights(num_experts, in_features, out_features, width_ratio, std_base,
                             device, dtype) -> nn.Parameter:
        w = torch.empty(num_experts, in_features, out_features, device=device, dtype=dtype)
        std_scaled = std_base / math.sqrt(width_ratio)
        nn.init.trunc_normal_(w, mean=0.0, std=std_scaled, a=-3 * std_scaled, b=3 * std_scaled)
        return nn.Parameter(w)

    @torch.no_grad()
    def reset_parameters(self) -> None:
        """Re-init the grouped expert weights (see Linear.reset_parameters)."""
        std = self._expert_init_std
        for w in (self.experts_w1, self.experts_w2, self.experts_w3):
            nn.init.trunc_normal_(w, mean=0.0, std=std, a=-3 * std, b=3 * std)

    def __init__(self, d_model: int, num_heads: int, d_ff: int, num_experts: int, num_active: int,
                 max_seq_len: int, theta: float, width_ratio: float, device=None, dtype=None):
        super().__init__()
        self.ln1 = RMSNorm(d_model, device=device, dtype=dtype)
        self.attn = MultiheadSelfAttention(d_model, num_heads, max_seq_len, theta, width_ratio, device, dtype)
        self.ln2 = RMSNorm(d_model, device=device, dtype=dtype)
        self.router = Router(d_model, num_experts, num_active, width_ratio, device=device, dtype=dtype)

        self.num_experts = num_experts
        self.num_active = num_active

        # NOTE: divided by num_active (k), not by num_experts (E). This is what
        # keeps total expert parameters constant across the DVF series.
        d_ff_expert = d_ff // num_active
        w_std_base = math.sqrt(2 / (BASE_D_MODEL + BASE_D_FF))
        self._expert_init_std = w_std_base / math.sqrt(width_ratio)
        self.experts_w1 = self._init_expert_weights(num_experts, d_model, d_ff_expert, width_ratio, w_std_base, device, dtype)
        self.experts_w2 = self._init_expert_weights(num_experts, d_ff_expert, d_model, width_ratio, w_std_base, device, dtype)
        self.experts_w3 = self._init_expert_weights(num_experts, d_model, d_ff_expert, width_ratio, w_std_base, device, dtype)

    def forward(
        self,
        x: Tensor,
        token_positions: Optional[Tensor] = None,
        *,
        loop_step: Optional[int] = None,
        layer_idx: Optional[int] = None,
        block: Optional[str] = None,
        unrolled_pos: Optional[int] = None,
        trace_sink: Optional[TraceSink] = None,
    ) -> tuple[Tensor, Tensor, Tensor]:
        batch, seq, dim = x.shape
        total_tokens = batch * seq

        norm1_out = self.ln1(x)
        attn_out = self.attn(norm1_out, token_positions)
        assert x.shape == attn_out.shape
        resid1_out = attn_out + x

        norm2_out = self.ln2(resid1_out)
        logits, probs, top_scores, top_experts = self.router(norm2_out)

        # `softmax` computes in float32 and does not cast back, so `top_scores`
        # is float32 regardless of the activation dtype. The combine weights get
        # multiplied into bf16 expert outputs below, so they must match, or
        # einsum raises "expected m1 and m2 to have the same dtype".
        #
        # This is invisible in the original: seedvar's published checkpoints are
        # float32, where the cast is a no-op. Under the bf16 training this
        # project uses (LT2 runs pure bf16, HANDOVER §4.1) it is a hard failure
        # on the first forward.
        #
        # Only the combine weights are cast. `probs` and `logits` stay float32
        # for the aux-loss and z-loss reductions, which is where the extra
        # precision is worth having.
        top_scores = top_scores.to(x.dtype)

        # Flatten and sort by expert so grouped_mm can run one matmul per expert.
        x_flat = rearrange(norm2_out, "b s d -> (b s) d")
        flat_expert_ids = rearrange(top_experts, "b s k -> (b s k)")
        flat_scores = rearrange(top_scores, "b s k -> (b s k)")
        flat_positions = torch.arange(total_tokens, device=x.device)
        flat_token_ids = repeat(flat_positions, "n -> (n k)", k=self.num_active)

        sort_indices = flat_expert_ids.argsort(stable=True)
        sorted_expert_ids = flat_expert_ids[sort_indices]
        sorted_token_ids = flat_token_ids[sort_indices]
        sorted_scores = flat_scores[sort_indices]
        sorted_x = x_flat[sorted_token_ids]

        counts = torch.bincount(sorted_expert_ids, minlength=self.num_experts)
        offs = counts.cumsum(0).to(torch.int32)

        h1 = grouped_mm(sorted_x, self.experts_w1, offs=offs)
        h3 = grouped_mm(sorted_x, self.experts_w3, offs=offs)
        gated = silu(h1) * h3
        expert_out = grouped_mm(gated, self.experts_w2, offs=offs)

        expert_out = einsum(expert_out, sorted_scores, "n d, n -> n d")
        output_flat = torch.zeros(total_tokens, dim, device=x.device, dtype=expert_out.dtype)
        output_flat.index_add_(0, sorted_token_ids, expert_out)
        experts_out = rearrange(output_flat, "(b s) d -> b s d", b=batch, s=seq)

        # Aux losses, per HANDOVER §3.3':
        #   L_LB = E * sum_i f_i * p_i   (switch-style load balancing)
        #   L_RZ = mean( (logsumexp logits)^2 )   (router z-loss)
        # Both are computed per layer per loop step; the caller averages over the
        # unrolled depth (num_stacks * num_layers_in_stack).
        fi = counts.float() / (total_tokens * self.num_active)
        pi = reduce(probs, "b s e -> e", "mean")
        lb = self.num_experts * einsum(fi, pi, "e, e ->")

        logsumexp = torch.logsumexp(logits.float(), dim=-1)
        lz = reduce(logsumexp**2, "... -> ", "mean")

        assert experts_out.shape == resid1_out.shape
        final_out = resid1_out + experts_out

        # Write into every sink that is armed. Under FSDP2 the `trace_sink`
        # keyword arrives as a per-block COPY (see `_ActiveTraceSink`), so the
        # module-level one is the only sink the trainer can actually read back;
        # the keyword remains for callers that pass their own dict directly
        # (the toy launcher, the offline probes), where it is the same object.
        # Writing to both is harmless: each is first-write-wins.
        sinks = [s for s in (ACTIVE_TRACE_SINK.sink, trace_sink) if s is not None]
        if sinks:
            if loop_step is None or layer_idx is None or block is None or unrolled_pos is None:
                raise ValueError(
                    "trace_sink was provided but loop_step/layer_idx/block/unrolled_pos "
                    "were not. Both axes must be explicit counters, never inferred."
                )
            if block not in ("head", "loop", "tail"):
                raise ValueError(f"block must be head/loop/tail, got {block!r}")
            # Keyed by the depth coordinate: head, the loop's first layer and tail all
            # carry loop_step=layer_idx=0, so the old pair-key silently collapsed them.
            key = unrolled_pos
            # First write wins. Under selective activation checkpointing the
            # block's forward runs a second time during backward to recompute
            # activations, so every key is legitimately visited twice -- that is
            # how AC works, not a bug. The recomputed values are identical by
            # construction, so keeping the first and ignoring the rest is both
            # correct and cheap.
            #
            # An earlier version raised on the second visit. That guard was aimed
            # at double-*counting*, which first-write-wins prevents directly; as
            # written it instead killed every AC-enabled run at the first traced
            # step. `test_ac_recomputation_does_not_disturb_the_trace` pins the
            # property that actually matters: same keys, same values, with AC on.
            # Detached views, not copies -- see the lifetime contract in
            # src/model/loop_trace.py.
            point = TracePoint(
                block=block,
                unrolled_pos=unrolled_pos,
                # From the tensors themselves, never from the model config: this layer's
                # E and k are what produced these numbers, and head/tail differ from the
                # loop block.
                num_experts=int(logits.shape[-1]),
                top_k=int(top_experts.shape[-1]),
                loop_step=loop_step,
                layer_idx=layer_idx,
                router_logits=logits.detach(),
                router_probs=probs.detach(),
                topk_idx=top_experts.detach(),
                topk_weights=top_scores.detach(),
                residual=final_out.detach(),
            )
            for sink in sinks:
                sink.setdefault(key, point)

        return final_out, lb, lz


class LoopedStack(nn.Module):
    """The stack of L MoE blocks that gets called R times."""

    def __init__(self, context_length: int, d_model: int, num_layers_in_stack: int, num_heads: int,
                 d_ff: int, rope_theta: float, width_ratio: float, num_experts: int,
                 num_active: int, device=None, dtype=None):
        super().__init__()
        self.layers = nn.ModuleList(
            [
                GroupedMoEPrenormBlock(
                    d_model, num_heads, d_ff, num_experts, num_active,
                    context_length, rope_theta, width_ratio, device, dtype,
                )
                for _ in range(num_layers_in_stack)
            ]
        )

    def forward(
        self,
        x: Tensor,
        *,
        loop_step: int,
        unrolled_pos_start: int,
        trace_sink: Optional[TraceSink] = None,
    ) -> tuple[Tensor, Tensor, Tensor]:
        """`unrolled_pos_start` is the depth coordinate this call's first layer occupies.

        Passed in rather than recomputed from `loop_step`, so the caller owns the depth
        axis in one place: the stack does not need to know how many layers ran before it.
        """
        lb_total = x.new_zeros(())
        lz_total = x.new_zeros(())
        for layer_idx, layer in enumerate(self.layers):
            x, lb, lz = layer(
                x, loop_step=loop_step, layer_idx=layer_idx, block="loop",
                unrolled_pos=unrolled_pos_start + layer_idx, trace_sink=trace_sink,
            )
            lb_total = lb_total + lb
            lz_total = lz_total + lz
        return x, lb_total, lz_total


class LoopedMoETransformer(nn.Module):
    """Looped MoE transformer: one shared stack applied ``num_stacks`` times.

    The loop is an explicit Python ``for``; ``loop_step`` is the loop variable and
    is threaded all the way down to each block. Nothing downstream ever has to
    recover it from tensor shapes.
    """

    def __init__(self, vocab_size: int, context_length: int, d_model: int, num_layers_in_stack: int,
                 num_stacks: int, num_heads: int, d_ff: int, num_experts: int, num_active: int,
                 rope_theta: float, width_ratio: float, num_head_layers: int = 0,
                 num_tail_layers: int = 0, head_tail_num_experts: int = HEAD_TAIL_NUM_EXPERTS,
                 device=None, dtype=None):
        super().__init__()
        self.num_stacks = num_stacks
        self.num_layers_in_stack = num_layers_in_stack
        self.total_layers = num_stacks * num_layers_in_stack
        self.num_head_layers = num_head_layers
        self.num_tail_layers = num_tail_layers
        # The unrolled depth, and the denominator the aux losses are averaged over.
        # Written as head + loop + tail rather than the recipe's shorthand "2 + D*L":
        # S15 has two head and two tail layers, so a literal 2 would divide by the wrong
        # number there -- and it would not fail, it would just make the auxiliary losses
        # quietly larger than intended.
        self.unrolled_depth = num_head_layers + self.total_layers + num_tail_layers

        self.token_embeddings = Embedding(vocab_size, d_model, device=device, dtype=dtype)
        # Head and tail are ordinary MoE layers that run once. They keep E=8/top-2
        # regardless of the loop block's expert count (recipe section 2.1: "identical in
        # every configuration"), so the loop block's E is deliberately not passed here.
        make_outer = lambda: GroupedMoEPrenormBlock(
            d_model, num_heads, d_ff, head_tail_num_experts, num_active,
            context_length, rope_theta, width_ratio, device, dtype,
        )
        self.head_layers = nn.ModuleList([make_outer() for _ in range(num_head_layers)])
        self.tail_layers = nn.ModuleList([make_outer() for _ in range(num_tail_layers)])
        self.stack = LoopedStack(
            context_length, d_model, num_layers_in_stack, num_heads, d_ff, rope_theta,
            width_ratio, num_experts, num_active, device=device, dtype=dtype,
        )
        self.ln_final = RMSNorm(d_model, device=device, dtype=dtype)
        std_base_lm_head = math.sqrt(2 / (BASE_D_MODEL + vocab_size))
        self.lm_head = Linear(d_model, vocab_size, width_ratio, std_base_lm_head, device=device, dtype=dtype)

    @classmethod
    def from_config(cls, config: "LoopMoEConfig", *, device=None, dtype=None) -> "LoopedMoETransformer":
        """THE way to build this model from a config. Both call sites use it.

        There were two: the HF wrapper and `pretrain/train_spec.py`, each with its own
        hand-written keyword list. When head/tail layers were added, the training path's
        list was not updated, so it silently built a model with no head or tail while its
        config said otherwise -- it trained, the loss looked plausible, and every artifact
        recorded the config's depth rather than the depth that ran. Nothing could raise,
        because a shorter model is a perfectly valid model.

        A single entry point makes that class of drift impossible rather than merely
        tested-for: a field added to the config is read here once, and both paths get it.
        """
        return cls(
            vocab_size=config.vocab_size,
            context_length=config.context_length,
            d_model=config.d_model,
            num_layers_in_stack=config.num_layers_in_stack,
            num_stacks=config.num_stacks,
            num_heads=config.num_heads,
            d_ff=config.d_ff,
            num_experts=config.num_experts,
            num_active=config.num_active,
            rope_theta=config.rope_theta,
            width_ratio=config.width_ratio,
            num_head_layers=config.num_head_layers,
            num_tail_layers=config.num_tail_layers,
            device=device,
            dtype=dtype,
        )

    def forward(
        self,
        x: Tensor,
        *,
        trace_sink: Optional[TraceSink] = None,
    ) -> tuple[Tensor, Tensor, Tensor]:
        lb_total = None
        lz_total = None

        x = self.token_embeddings(x)

        def run_outer(layers, block: str, pos: int, x, lb_total, lz_total):
            for i, layer in enumerate(layers):
                x, lb, lz = layer(
                    x, loop_step=0, layer_idx=0, block=block, unrolled_pos=pos + i,
                    trace_sink=trace_sink,
                )
                lb_total = lb if lb_total is None else lb_total + lb
                lz_total = lz if lz_total is None else lz_total + lz
            return x, lb_total, lz_total

        # head -> loop x num_stacks -> tail, with one running depth coordinate. head and
        # tail record loop_step=layer_idx=0 because neither coordinate means anything
        # outside the loop; `block` and `unrolled_pos` are what identifies them.
        x, lb_total, lz_total = run_outer(self.head_layers, "head", 0, x, lb_total, lz_total)
        pos = self.num_head_layers
        for loop_step in range(self.num_stacks):
            x, lb, lz = self.stack(
                x, loop_step=loop_step, unrolled_pos_start=pos, trace_sink=trace_sink,
            )
            pos += self.num_layers_in_stack
            lb_total = lb if lb_total is None else lb_total + lb
            lz_total = lz if lz_total is None else lz_total + lz
        x, lb_total, lz_total = run_outer(self.tail_layers, "tail", pos, x, lb_total, lz_total)

        x = self.lm_head(self.ln_final(x))

        # Averaged over the *unrolled* depth: every physical layer contributes once per
        # loop step, and head/tail contribute once each (recipe section 2.1). Equals
        # `total_layers` exactly when there are no head/tail layers, which is every
        # pre-ablation configuration -- so their published losses are unchanged.
        return x, lb_total / self.unrolled_depth, lz_total / self.unrolled_depth


def _require(kwargs: dict[str, Any], name: str) -> Any:
    """Fetch a required config field or raise.

    Project rule (HANDOVER §4.3 / §5.6): semantic parameters get no defaults.
    A default is a silent-wrong-answer generator -- a mistyped or dropped key
    becomes a plausible number instead of an error.
    """
    if name not in kwargs or kwargs[name] is None:
        raise ValueError(
            f"LoopMoEConfig: required field {name!r} is missing. Semantic "
            "parameters have no defaults in this project; state it explicitly."
        )
    return kwargs.pop(name)


class LoopMoEConfig(PretrainedConfig):
    """Config for the looped-MoE architecture. **Every field is required.**

    Compatible with ``save_pretrained``/``from_pretrained``: a config.json written
    by this class round-trips, and one is rejected loudly if a field is absent.
    """

    model_type = "loop-moe"

    # Tells transformers not to introspect defaults by constructing `cls()` with
    # no arguments -- which this class deliberately rejects. Without it,
    # `save_pretrained` fails inside `_get_generation_parameters`. This is the
    # supported escape hatch for configs whose fields are all required.
    has_no_defaults_at_init = True

    def __init__(self, **kwargs: Any):
        # `from_pretrained` on a *torch-saved* config, and some HF-internal paths,
        # construct with no arguments at all; only a fully-specified call is valid.
        self.vocab_size = _require(kwargs, "vocab_size")
        self.context_length = _require(kwargs, "context_length")
        self.d_model = _require(kwargs, "d_model")
        self.num_heads = _require(kwargs, "num_heads")
        self.d_ff = _require(kwargs, "d_ff")
        self.rope_theta = _require(kwargs, "rope_theta")
        self.width_ratio = _require(kwargs, "width_ratio")
        self.num_layers_in_stack = _require(kwargs, "num_layers_in_stack")  # L
        self.num_stacks = _require(kwargs, "num_stacks")  # R
        self.num_experts = _require(kwargs, "num_experts")  # E
        self.num_active = _require(kwargs, "num_active")  # k
        self.lb_loss_factor = _require(kwargs, "lb_loss_factor")
        # Head/tail layers: MoE layers run once, outside the loop (ablation recipe 2026-09-12
        # section 2.1). 0 means the architecture has none, which is not a guess -- it is what
        # every configuration built before this recipe actually is, and
        # `test_config_registry` pins their parameter counts as unchanged. Real ablation
        # configs never rely on the fallback: `config_registry.build_ablation` states both
        # counts for every entry, and a test asserts it does.
        self.num_head_layers = int(kwargs.pop("num_head_layers", 0))
        self.num_tail_layers = int(kwargs.pop("num_tail_layers", 0))
        self.lz_loss_factor = _require(kwargs, "lz_loss_factor")

        # Stated, not inherited. `PretrainedConfig` defaults this to True, and the
        # only reason the embedding and the LM head are not already sharing storage
        # is that this model never implemented `get_output_embeddings()`. The day
        # someone adds it for tool compatibility, every configuration would start
        # tying weights -- a different model, trained to a different loss, with
        # nothing in any artifact saying so. The architecture uses untied weights
        # (the parameter counts in the recipe assume it), so the config says so.
        kwargs.pop("tie_word_embeddings", None)
        self.tie_word_embeddings = False

        self._validate()

        # Derived, for readers; never an input. The unrolled depth now includes the
        # layers that run once outside the loop.
        self.num_layers = self.num_stacks * self.num_layers_in_stack
        self.unrolled_depth = self.num_head_layers + self.num_layers + self.num_tail_layers

        # The original config mirrored `context_length` into `max_length` for
        # lm-evaluation-harness. transformers >=5 classifies `max_length` as a
        # generation parameter and refuses to serialise a config carrying one
        # (the check is `hasattr`, so even a property trips it). `context_length`
        # is therefore the single source of truth for sequence length; pass
        # `max_length` to the harness explicitly at eval time instead.
        # Popped so that loading a seedvar-era config.json cannot reintroduce it.
        kwargs.pop("max_length", None)

        super().__init__(**kwargs)

    def _validate(self) -> None:
        """Reject out-of-domain values loudly rather than failing deep in a kernel."""
        positive = (
            "vocab_size", "context_length", "d_model", "num_heads", "d_ff",
            "num_layers_in_stack", "num_stacks", "num_experts", "num_active",
        )
        for name in positive:
            value = getattr(self, name)
            if not isinstance(value, int) or value < 1:
                raise ValueError(f"LoopMoEConfig.{name} must be a positive int, got {value!r}")
        if self.d_model % self.num_heads != 0:
            raise ValueError(
                f"d_model ({self.d_model}) must be divisible by num_heads ({self.num_heads})"
            )
        if self.num_active > self.num_experts:
            raise ValueError(
                f"num_active ({self.num_active}) cannot exceed num_experts ({self.num_experts})"
            )
        if self.d_ff % self.num_active != 0:
            raise ValueError(
                f"d_ff ({self.d_ff}) must be divisible by num_active ({self.num_active}); "
                "expert width is d_ff // num_active and truncation would silently "
                "change the parameter count."
            )
        for name in ("num_head_layers", "num_tail_layers"):
            value = getattr(self, name)
            if not isinstance(value, int) or value < 0:
                raise ValueError(f"LoopMoEConfig.{name} must be a non-negative int, got {value!r}")
        if (self.d_model // self.num_heads) % 2 != 0:
            raise ValueError(
                f"head dim ({self.d_model // self.num_heads}) must be even for RoPE"
            )

    @property
    def trace_meta(self) -> TraceMeta:
        """Shape/provenance block handed to the metrics collector."""
        return TraceMeta(
            num_stacks=self.num_stacks,
            num_layers_in_stack=self.num_layers_in_stack,
            num_experts=self.num_experts,
            num_active=self.num_active,
        )


class LoopMoEForCausalLM(PreTrainedModel, GenerationMixin):
    """HF-compatible causal LM wrapper.

    Kept HF-shaped on purpose: HANDOVER §4.7 makes "the diagnostics pipeline
    ingests our checkpoints unchanged" an acceptance criterion, and that pipeline
    loads models through ``from_pretrained``.
    """

    config_class = LoopMoEConfig

    def __init__(self, config: LoopMoEConfig):
        super().__init__(config)
        self.model = LoopedMoETransformer.from_config(config)
        self.post_init()

    def get_input_embeddings(self):
        return self.model.token_embeddings

    def set_input_embeddings(self, value):
        self.model.token_embeddings = value

    def forward(
        self,
        input_ids: torch.LongTensor,
        attention_mask: Optional[Tensor] = None,  # unused: the mask is built in
        labels: Optional[torch.LongTensor] = None,
        trace_sink: Optional[TraceSink] = None,
        **kwargs: Any,
    ) -> CausalLMOutputWithPast:
        """Forward pass.

        Returns a ``CausalLMOutputWithPast`` whose ``loss`` is the *total* loss
        (CE + weighted aux). The unweighted components are attached as
        ``task_loss`` / ``lb_loss`` / ``z_loss`` so the training loop can log the
        breakdown without recomputing anything.

        **Label contract (documented here 2026-08-21, `ABCI_ERR_20260821_0405_
        gate1_label_shift_root_cause.md`)**: ``labels`` must already be
        next-token-shifted by the caller -- ``labels[..., t] == input_ids[..., t+1]``,
        with the last position set to ``-100`` (no target exists after it). This
        method does **not** shift internally; it passes ``labels`` to
        ``F.cross_entropy`` exactly as given. Before this date the only written
        record of this contract was a comment in
        ``src/model/loop_trace.py`` ("the dataloader supplies pre-shifted
        labels during training, so the shift is explicit here") -- not here, at
        the definition itself. That gap let four separate call sites
        (the pretraining dataloader path aside, which was correct) independently
        get this wrong the same way, rather than it being four unrelated
        mistakes. Any caller not shifting first -- e.g. ``model(input_ids=ids,
        labels=ids)`` -- silently trains/evaluates on the trivial
        copy-the-current-token target instead of next-token prediction.
        """
        logits, lb, lz = self.model(input_ids, trace_sink=trace_sink)

        loss = task_loss = None
        if labels is not None:
            task_loss = F.cross_entropy(logits.view(-1, logits.size(-1)), labels.view(-1))
            loss = (
                task_loss
                + self.config.lb_loss_factor * lb
                + self.config.lz_loss_factor * lz
            )

        out = CausalLMOutputWithPast(loss=loss, logits=logits)
        # Unweighted components; the trainer pairs them with the factors from
        # config to build LossComponents.
        out.task_loss = task_loss
        out.lb_loss = lb
        out.z_loss = lz
        return out

    def prepare_inputs_for_generation(self, input_ids, **kwargs):
        return {"input_ids": input_ids}