SYNKA 27938: verified step_00028000
Browse files- step_00028000/config.json +31 -0
- step_00028000/loop_trace.py +539 -0
- step_00028000/manifest.json +96 -0
- step_00028000/modeling_loop_lm.py +936 -0
- step_00028000/pytorch_model.bin +3 -0
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
|