ahiok commited on
Commit
4f0fb70
·
verified ·
1 Parent(s): 2c00109

Upload folder using huggingface_hub

Browse files
Files changed (7) hide show
  1. README.md +79 -0
  2. model.py +619 -0
  3. model_config.json +33 -0
  4. pytorch_model.bin +3 -0
  5. summary.json +153 -0
  6. tokenizer.json +0 -0
  7. training_config.json +103 -0
README.md ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ language: [en]
4
+ datasets: [HuggingFaceFW/fineweb]
5
+ tags: [looped-transformer, recurrent-depth, latent-reasoning, test-time-compute, small-lm]
6
+ pipeline_tag: text-generation
7
+ ---
8
+
9
+ # ahiok/looped-fineweb-10m
10
+
11
+ A **looped** decoder-only language model trained under a hard budget of
12
+ **9,441,152 parameters** (6,295,424 non-embedding) and **100,000,000 training
13
+ tokens** of [FineWeb](https://huggingface.co/datasets/HuggingFaceFW/fineweb).
14
+
15
+ The block that gets looped is Qwen3-style (RMSNorm pre-norm, GQA, QK-norm,
16
+ SwiGLU, RoPE). The same 2 layers are applied `R` times; `R` is chosen at
17
+ inference, so the same weights can be run cheap or deep.
18
+
19
+ ## Results
20
+
21
+ | eval loops R | val loss | perplexity | bits/byte |
22
+ |---|---|---|---|
23
+ | 1 | 7.4054 | 1644.79 | 2.8394 |
24
+ | 2 | 6.4121 | 609.16 | 2.4586 |
25
+ | 4 | 5.2127 | 183.59 | 1.9987 |
26
+ | 8 | 4.1340 | 62.43 | 1.5851 |
27
+ | 16 | 3.7898 | 44.25 | 1.4531 |
28
+ | 24 | 3.8698 | 47.93 | 1.4838 |
29
+ | 32 | 3.9913 | 54.13 | 1.5304 |
30
+ | 48 | 4.1892 | 65.97 | 1.6063 |
31
+ | 64 | 4.3275 | 75.76 | 1.6593 |
32
+ | 96 | 4.5078 | 90.72 | 1.7284 |
33
+ | 128 | 4.6195 | 101.44 | 1.7712 |
34
+
35
+ Validation is a held-out document split of the same FineWeb shard, 0
36
+ tokens, tokenised with the 8192-entry byte-level BPE included in this repo.
37
+ Bits-per-byte is reported alongside perplexity because perplexity alone is not
38
+ comparable across tokenizers.
39
+
40
+ The recurrence is `s <- Block(s + e)`: the embedded input is added back into the state at the start of every iteration. That one tensor add is the entire difference from an unlooped model of **identical parameter count**, and it is worth 0.10 nats here. Run it at **R=16**, the depth it was trained at: this variant buys quality rather than depth robustness and degrades sharply on either side (3.79 at R=16, 4.13 at R=8, 3.99 at R=32).
41
+
42
+ For reference, an **unlooped** 4-layer model of the same size trained on the same
43
+ 100M tokens reaches 3.8965 / 49.23 / 1.4940, and a plain looped model with no
44
+ update rule reaches 3.8315 / 46.13 / 1.4691.
45
+
46
+ This is a research artefact for studying test-time depth scaling under a hard
47
+ budget, not a usable text generator. At 9.4M parameters and 100M tokens it
48
+ produces the statistics of English, not sentences you would want to read.
49
+
50
+ ## Usage
51
+
52
+ ```python
53
+ import importlib.util, json, sys, torch
54
+ from huggingface_hub import hf_hub_download
55
+ from tokenizers import Tokenizer
56
+
57
+ repo = "ahiok/looped-fineweb-10m"
58
+ src = hf_hub_download(repo, "model.py")
59
+ spec = importlib.util.spec_from_file_location("loopllm_model", src)
60
+ mod = importlib.util.module_from_spec(spec)
61
+ sys.modules["loopllm_model"] = mod # required: @dataclass resolves via sys.modules
62
+ spec.loader.exec_module(mod)
63
+
64
+ cfg = mod.ModelConfig(**json.load(open(hf_hub_download(repo, "model_config.json"))))
65
+ model = mod.LoopedLM(cfg)
66
+ model.load_state_dict(torch.load(hf_hub_download(repo, "pytorch_model.bin"), map_location="cpu"))
67
+ model.eval()
68
+
69
+ tok = Tokenizer.from_file(hf_hub_download(repo, "tokenizer.json"))
70
+ ids = torch.tensor([tok.encode("The capital of France is").ids])
71
+ out = model(ids, n_loops=32) # spend more or less compute here
72
+ print(tok.decode([int(out["logits"][0, -1].argmax())]))
73
+ ```
74
+
75
+ ## Training
76
+
77
+ Code, full ablations and the report: https://github.com/ahiokk/looped-models
78
+
79
+
model.py ADDED
@@ -0,0 +1,619 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Looped decoder-only LM with a Qwen3-style block.
2
+
3
+ Layout follows the prelude / recurrent / coda decomposition:
4
+
5
+ x -> embed -> [prelude L_p layers] -> e
6
+ s_0 = e
7
+ s_r = Block(s_{r-1}, e, r, R) for r = 1..R (shared weights)
8
+ logits = head(norm(coda(s_R)))
9
+
10
+ Every research knob is a config flag so that one binary can produce the whole
11
+ ablation ladder and every run is described by its config dict alone.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import math
17
+ from dataclasses import asdict, dataclass, field
18
+ from typing import Optional
19
+
20
+ import torch
21
+ import torch.nn as nn
22
+ import torch.nn.functional as F
23
+ import torch.utils.checkpoint
24
+
25
+
26
+ # --------------------------------------------------------------------------------------
27
+ # config
28
+ # --------------------------------------------------------------------------------------
29
+ @dataclass
30
+ class ModelConfig:
31
+ # --- Qwen3-style backbone -----------------------------------------------------
32
+ vocab_size: int = 8192
33
+ d_model: int = 384
34
+ n_heads: int = 6
35
+ n_kv_heads: int = 2
36
+ head_dim: int = 64
37
+ d_ff: int = 1024
38
+ max_seq_len: int = 512
39
+ rope_theta: float = 10_000.0
40
+ rms_eps: float = 1e-6
41
+ tie_embeddings: bool = True
42
+ # pre = Qwen3 default. sandwich = Huginn's block, which normalises after each
43
+ # residual add as well; costs 2d per layer and bounds the residual stream.
44
+ block_norm: str = "pre" # pre | sandwich
45
+
46
+ # --- depth layout -------------------------------------------------------------
47
+ n_prelude: int = 1
48
+ n_recurrent: int = 2
49
+ n_coda: int = 1
50
+
51
+ # --- looping ------------------------------------------------------------------
52
+ n_loops: int = 8 # R used at train time (mean of the distribution if sampled)
53
+ max_loops: int = 256 # size of the precomputed depth-embedding table
54
+ state_init: str = "prelude" # prelude | randn
55
+ state_init_std: float = 0.4 # only for state_init == "randn"
56
+
57
+ input_injection: str = "add" # none | add | adapter
58
+ state_norm: str = "none" # none | rms (normalise s at loop entry)
59
+ # residual : s <- Block(s) (the usual looped transformer)
60
+ # convex : s <- (1-a) s + a Block(s) (learned step size)
61
+ # flow : s <- s + (gain/R) * Delta(s, r/R) (explicit Euler step of a learned flow)
62
+ update_rule: str = "residual"
63
+ # pre-sigmoid init of the convex step size. +3 starts at ~0.95, i.e. almost a
64
+ # full replacement (the usual looped behaviour); -3 starts at ~0.05, so the
65
+ # loop begins as a near-identity and has to earn its depth, which is what
66
+ # makes very deep shared stacks trainable at all
67
+ update_gate_init: float = 3.0
68
+ depth_cond: str = "none" # none | film
69
+ depth_cond_input: str = "progress" # absolute | progress | both
70
+ depth_cond_dim: int = 64
71
+
72
+ loop_noise: float = 0.0 # std of exploration noise injected at loop entry
73
+ noise_schedule: str = "linear" # linear | const | cosine (annealed towards 0 at r=R)
74
+
75
+ # learned halting, PonderNet style: a per-token probability of stopping after
76
+ # each iteration, trained jointly with the language-model loss. Costs d + 1
77
+ # parameters and is independent of R, so the maximum depth stays a runtime knob.
78
+ halting: str = "none" # none | ponder
79
+ halt_prior: float = 0.1 # geometric prior on the halting step
80
+ halt_kl_weight: float = 0.01
81
+
82
+ # --- init ---------------------------------------------------------------------
83
+ init_std: float = 0.02
84
+ depth_scaled_init: bool = True
85
+
86
+ def __post_init__(self) -> None:
87
+ assert self.n_heads % self.n_kv_heads == 0
88
+ assert self.state_init in {"prelude", "randn"}
89
+ assert self.input_injection in {"none", "add", "adapter"}
90
+ assert self.block_norm in {"pre", "sandwich"}
91
+ assert self.state_norm in {"none", "rms"}
92
+ assert self.update_rule in {"residual", "convex", "flow"}
93
+ assert self.depth_cond in {"none", "film"}
94
+ assert self.depth_cond_input in {"absolute", "progress", "both"}
95
+ assert self.noise_schedule in {"linear", "const", "cosine"}
96
+ assert self.halting in {"none", "ponder"}
97
+
98
+ def to_dict(self) -> dict:
99
+ return asdict(self)
100
+
101
+
102
+ # --------------------------------------------------------------------------------------
103
+ # primitives
104
+ # --------------------------------------------------------------------------------------
105
+ class RMSNorm(nn.Module):
106
+ def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = True):
107
+ super().__init__()
108
+ self.eps = eps
109
+ self.weight = nn.Parameter(torch.ones(dim)) if elementwise_affine else None
110
+
111
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
112
+ dtype = x.dtype
113
+ x = x.float()
114
+ x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
115
+ x = x.to(dtype)
116
+ return x * self.weight if self.weight is not None else x
117
+
118
+
119
+ def build_rope_cache(seq_len: int, head_dim: int, theta: float, device, dtype=torch.float32):
120
+ inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim))
121
+ t = torch.arange(seq_len, device=device, dtype=torch.float32)
122
+ freqs = torch.outer(t, inv_freq) # (T, hd/2)
123
+ emb = torch.cat((freqs, freqs), dim=-1) # (T, hd)
124
+ return emb.cos().to(dtype), emb.sin().to(dtype)
125
+
126
+
127
+ def rotate_half(x: torch.Tensor) -> torch.Tensor:
128
+ x1, x2 = x.chunk(2, dim=-1)
129
+ return torch.cat((-x2, x1), dim=-1)
130
+
131
+
132
+ def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
133
+ # x: (B, H, T, hd); cos/sin: (T, hd)
134
+ cos = cos[None, None, :, :]
135
+ sin = sin[None, None, :, :]
136
+ return x * cos + rotate_half(x) * sin
137
+
138
+
139
+ class Attention(nn.Module):
140
+ """Qwen3 attention: GQA, no qkv bias, RMSNorm on q and k heads."""
141
+
142
+ def __init__(self, cfg: ModelConfig):
143
+ super().__init__()
144
+ self.n_heads = cfg.n_heads
145
+ self.n_kv_heads = cfg.n_kv_heads
146
+ self.head_dim = cfg.head_dim
147
+ self.q_proj = nn.Linear(cfg.d_model, cfg.n_heads * cfg.head_dim, bias=False)
148
+ self.k_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False)
149
+ self.v_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False)
150
+ self.o_proj = nn.Linear(cfg.n_heads * cfg.head_dim, cfg.d_model, bias=False)
151
+ self.q_norm = RMSNorm(cfg.head_dim, cfg.rms_eps)
152
+ self.k_norm = RMSNorm(cfg.head_dim, cfg.rms_eps)
153
+
154
+ def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
155
+ B, T, _ = x.shape
156
+ q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
157
+ k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
158
+ v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
159
+
160
+ q = self.q_norm(q)
161
+ k = self.k_norm(k)
162
+ q = apply_rope(q, cos, sin)
163
+ k = apply_rope(k, cos, sin)
164
+
165
+ o = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True)
166
+ o = o.transpose(1, 2).contiguous().view(B, T, self.n_heads * self.head_dim)
167
+ return self.o_proj(o)
168
+
169
+
170
+ class MLP(nn.Module):
171
+ def __init__(self, cfg: ModelConfig):
172
+ super().__init__()
173
+ self.gate_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
174
+ self.up_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
175
+ self.down_proj = nn.Linear(cfg.d_ff, cfg.d_model, bias=False)
176
+
177
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
178
+ return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
179
+
180
+
181
+ class DecoderLayer(nn.Module):
182
+ """Qwen3 pre-norm layer, optionally with Huginn's sandwich norm.
183
+
184
+ Pre-norm (`block_norm="pre"`) is the Qwen3 default: the stream is normalised
185
+ on the way *into* each sublayer and the residual add is left alone, so the
186
+ stream is free to grow. Section 4.2 measures that growth and identifies it as
187
+ the reason late iterations stop mattering.
188
+
189
+ Sandwich (`block_norm="sandwich"`) is what the Huginn recurrent block
190
+ actually does: it normalises again *after* each residual add, which bounds
191
+ the stream without removing the residual path itself. That distinction is the
192
+ whole reason the loop-entry normalisation of 5.2 failed and this does not.
193
+ """
194
+
195
+ def __init__(self, cfg: ModelConfig):
196
+ super().__init__()
197
+ self.input_layernorm = RMSNorm(cfg.d_model, cfg.rms_eps)
198
+ self.self_attn = Attention(cfg)
199
+ self.post_attention_layernorm = RMSNorm(cfg.d_model, cfg.rms_eps)
200
+ self.mlp = MLP(cfg)
201
+ self.sandwich = cfg.block_norm == "sandwich"
202
+ if self.sandwich:
203
+ self.post_attn_residual_norm = RMSNorm(cfg.d_model, cfg.rms_eps)
204
+ self.post_mlp_residual_norm = RMSNorm(cfg.d_model, cfg.rms_eps)
205
+
206
+ def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
207
+ x = x + self.self_attn(self.input_layernorm(x), cos, sin)
208
+ if self.sandwich:
209
+ x = self.post_attn_residual_norm(x)
210
+ x = x + self.mlp(self.post_attention_layernorm(x))
211
+ if self.sandwich:
212
+ x = self.post_mlp_residual_norm(x)
213
+ return x
214
+
215
+
216
+ # --------------------------------------------------------------------------------------
217
+ # recurrent block
218
+ # --------------------------------------------------------------------------------------
219
+ class RecurrentBlock(nn.Module):
220
+ """The shared block applied R times.
221
+
222
+ Everything that makes iteration r behave differently from iteration r+1 has
223
+ to enter here, because the weights themselves are identical across r.
224
+ """
225
+
226
+ def __init__(self, cfg: ModelConfig):
227
+ super().__init__()
228
+ self.cfg = cfg
229
+ self.layers = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_recurrent)])
230
+
231
+ if cfg.input_injection == "adapter":
232
+ self.adapter = nn.Linear(2 * cfg.d_model, cfg.d_model, bias=False)
233
+
234
+ if cfg.state_norm == "rms":
235
+ self.entry_norm = RMSNorm(cfg.d_model, cfg.rms_eps)
236
+
237
+ if cfg.depth_cond == "film":
238
+ # sinusoidal features -> (scale, shift). Cost is O(d), independent of R,
239
+ # which is what keeps this usable at larger scale.
240
+ self.film = nn.Linear(cfg.depth_cond_dim, 2 * cfg.d_model, bias=True)
241
+ nn.init.zeros_(self.film.weight)
242
+ nn.init.zeros_(self.film.bias)
243
+
244
+ if cfg.update_rule == "convex":
245
+ # learned per-channel step size, sigmoid-gated, initialised near 1.0 so the
246
+ # untouched model starts out identical to the plain residual update
247
+ self.alpha = nn.Parameter(torch.full((cfg.d_model,), float(cfg.update_gate_init)))
248
+ elif cfg.update_rule == "flow":
249
+ # learned per-channel speed of the flow; the 1/R factor lives in forward()
250
+ self.flow_gain = nn.Parameter(torch.ones(cfg.d_model))
251
+
252
+ def forward(self, s, e, cos, sin, depth_feat: Optional[torch.Tensor] = None,
253
+ noise_std: Optional[torch.Tensor] = None, step_scale: Optional[torch.Tensor] = None):
254
+ # noise_std and step_scale arrive as 0-dim tensors on purpose: as python
255
+ # floats dynamo specialises the graph on their value and recompiles the
256
+ # block for every distinct loop count and noise level.
257
+ cfg = self.cfg
258
+ h = s
259
+
260
+ if cfg.input_injection == "add":
261
+ h = h + e
262
+ elif cfg.input_injection == "adapter":
263
+ h = self.adapter(torch.cat([h, e], dim=-1))
264
+
265
+ if cfg.state_norm == "rms":
266
+ h = self.entry_norm(h)
267
+
268
+ if cfg.depth_cond == "film" and depth_feat is not None:
269
+ mod = self.film(depth_feat) # (2d,)
270
+ scale, shift = mod.chunk(2, dim=-1)
271
+ h = h * (1.0 + scale) + shift
272
+
273
+ if cfg.loop_noise > 0.0 and noise_std is not None:
274
+ h = h + noise_std * torch.randn_like(h)
275
+
276
+ inner = h
277
+ for layer in self.layers:
278
+ inner = layer(inner, cos, sin)
279
+
280
+ if cfg.update_rule == "convex":
281
+ a = torch.sigmoid(self.alpha)
282
+ return (1.0 - a) * s + a * inner
283
+ if cfg.update_rule == "flow":
284
+ # explicit Euler step: s' = s + h_step * g(s, r/R). The loop count then
285
+ # sets the integration resolution rather than the amount of drift, so
286
+ # raising R at inference refines the same trajectory instead of
287
+ # walking further along it.
288
+ return s + (step_scale * self.flow_gain) * (inner - h)
289
+ return inner
290
+
291
+
292
+ # --------------------------------------------------------------------------------------
293
+ # full model
294
+ # --------------------------------------------------------------------------------------
295
+ class LoopedLM(nn.Module):
296
+ def __init__(self, cfg: ModelConfig):
297
+ super().__init__()
298
+ self.cfg = cfg
299
+ self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.d_model)
300
+ self.prelude = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_prelude)])
301
+ self.block = RecurrentBlock(cfg)
302
+ self.coda = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_coda)])
303
+ self.norm = RMSNorm(cfg.d_model, cfg.rms_eps)
304
+ if cfg.halting == "ponder":
305
+ self.halt_head = nn.Linear(cfg.d_model, 1)
306
+ self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
307
+ if cfg.tie_embeddings:
308
+ self.lm_head.weight = self.embed_tokens.weight
309
+
310
+ cos, sin = build_rope_cache(cfg.max_seq_len, cfg.head_dim, cfg.rope_theta, device="cpu")
311
+ self.register_buffer("rope_cos", cos, persistent=False)
312
+ self.register_buffer("rope_sin", sin, persistent=False)
313
+
314
+ self._depth_cache: dict[tuple, torch.Tensor] = {}
315
+
316
+ self.apply(self._init_weights)
317
+ if cfg.depth_scaled_init:
318
+ self._rescale_residual_projections()
319
+ if cfg.depth_cond == "film":
320
+ # zero-init the modulation so an untrained depth-conditioned model is
321
+ # bit-identical to the unconditioned one at step 0
322
+ nn.init.zeros_(self.block.film.weight)
323
+ nn.init.zeros_(self.block.film.bias)
324
+
325
+ # -- init ------------------------------------------------------------------------
326
+ def _init_weights(self, module: nn.Module) -> None:
327
+ std = self.cfg.init_std
328
+ if isinstance(module, nn.Linear):
329
+ nn.init.normal_(module.weight, mean=0.0, std=std)
330
+ if module.bias is not None:
331
+ nn.init.zeros_(module.bias)
332
+ elif isinstance(module, nn.Embedding):
333
+ nn.init.normal_(module.weight, mean=0.0, std=std)
334
+
335
+ def _rescale_residual_projections(self) -> None:
336
+ """GPT-2 style 1/sqrt(2 * depth) scaling, with depth counted through the loop.
337
+
338
+ The flow update already divides every step by R, so counting the loop
339
+ twice would leave the block effectively dead at initialisation.
340
+ """
341
+ loops = 1 if self.cfg.update_rule == "flow" else self.cfg.n_loops
342
+ depth = self.cfg.n_prelude + self.cfg.n_recurrent * loops + self.cfg.n_coda
343
+ scale = 1.0 / math.sqrt(2.0 * max(depth, 1))
344
+ for mod in self.modules():
345
+ if isinstance(mod, DecoderLayer):
346
+ mod.self_attn.o_proj.weight.data.mul_(scale)
347
+ mod.mlp.down_proj.weight.data.mul_(scale)
348
+
349
+ def _depth_table(self, R: int, device, dtype) -> torch.Tensor:
350
+ """Sinusoidal encodings of every loop index, shape (R + 1, depth_cond_dim).
351
+
352
+ Two things can be encoded: the absolute index r (tells the block how much
353
+ work has been done) and the progress r/R (tells it how much is left).
354
+ Which one matters is an experiment, not an assumption, hence the flag.
355
+ Cached per (R, device, dtype) so the table is built once per run.
356
+ """
357
+ cfg = self.cfg
358
+ key = (R, str(device), str(dtype))
359
+ if key in self._depth_cache:
360
+ return self._depth_cache[key]
361
+
362
+ r = torch.arange(R + 1, device=device, dtype=torch.float32)
363
+ vals = []
364
+ if cfg.depth_cond_input in {"absolute", "both"}:
365
+ vals.append(r)
366
+ if cfg.depth_cond_input in {"progress", "both"}:
367
+ vals.append(r / max(R, 1) * 32.0) # rescale so low frequencies stay informative
368
+
369
+ per = cfg.depth_cond_dim // (2 * len(vals))
370
+ idx = torch.arange(per, device=device, dtype=torch.float32)
371
+ freq = torch.exp(-math.log(10_000.0) * idx / max(per - 1, 1))
372
+ feats = []
373
+ for v in vals:
374
+ ang = v[:, None] * freq[None, :]
375
+ feats.append(torch.cat([torch.sin(ang), torch.cos(ang)], dim=-1))
376
+ out = torch.cat(feats, dim=-1)
377
+ if out.shape[-1] < cfg.depth_cond_dim:
378
+ out = F.pad(out, (0, cfg.depth_cond_dim - out.shape[-1]))
379
+ out = out.to(dtype)
380
+ self._depth_cache[key] = out
381
+ return out
382
+
383
+ def _noise_std(self, r: int, R: int) -> float:
384
+ cfg = self.cfg
385
+ if cfg.loop_noise <= 0.0 or not self.training:
386
+ return 0.0
387
+ if cfg.noise_schedule == "const":
388
+ return cfg.loop_noise
389
+ frac = (r - 1) / max(R - 1, 1)
390
+ if cfg.noise_schedule == "linear":
391
+ return cfg.loop_noise * (1.0 - frac)
392
+ return cfg.loop_noise * 0.5 * (1.0 + math.cos(math.pi * frac))
393
+
394
+ # -- forward ----------------------------------------------------------------------
395
+ def _readout_hidden(self, s: torch.Tensor, cos, sin) -> torch.Tensor:
396
+ """Coda output after the final norm; shared by the LM head and the halting head."""
397
+ h = s
398
+ for layer in self.coda:
399
+ h = layer(h, cos, sin)
400
+ return self.norm(h)
401
+
402
+ def _readout(self, s: torch.Tensor, cos, sin) -> torch.Tensor:
403
+ return self.lm_head(self._readout_hidden(s, cos, sin))
404
+
405
+ def forward(
406
+ self,
407
+ idx: torch.Tensor,
408
+ targets: Optional[torch.Tensor] = None,
409
+ n_loops: Optional[int] = None,
410
+ backprop_loops: int = 0,
411
+ readout_loops: Optional[list[int]] = None,
412
+ return_states: bool = False,
413
+ grad_checkpoint: bool = False,
414
+ readout_mode: str = "logits",
415
+ ):
416
+ """Run the model.
417
+
418
+ Args:
419
+ n_loops: R for this call (defaults to cfg.n_loops).
420
+ backprop_loops: if > 0, only the last k iterations carry gradient.
421
+ readout_loops: loop indices (1-based) whose intermediate logits are
422
+ also returned, used for deep supervision and for the coda lens.
423
+ return_states: also return the per-loop hidden states (diagnostics).
424
+ grad_checkpoint: recompute each iteration's internals in the backward
425
+ pass. Activation memory then stops growing with R, so a *full*
426
+ backward through 32 or 64 loops fits, which truncation does not
427
+ achieve without also changing what is being optimised.
428
+ readout_mode: "logits" keeps every intermediate logit tensor, which is
429
+ what deep supervision needs. "stats" reduces each one to per-token
430
+ loss and confidence immediately and throws the logits away; a
431
+ (B, T, 8192) tensor per loop is ~130 MB, so reading out all 32
432
+ loops for diagnostics costs gigabytes otherwise. "grad_stats" is
433
+ the same reduction but keeps the graph, which is what the ponder
434
+ objective needs: it weights every loop's loss by a learned halting
435
+ probability and so requires all of them to be differentiable.
436
+ """
437
+ cfg = self.cfg
438
+ B, T = idx.shape
439
+ R = n_loops if n_loops is not None else cfg.n_loops
440
+ cos = self.rope_cos[:T].to(idx.device)
441
+ sin = self.rope_sin[:T].to(idx.device)
442
+
443
+ h = self.embed_tokens(idx)
444
+ for layer in self.prelude:
445
+ h = layer(h, cos, sin)
446
+ e = h
447
+
448
+ if cfg.state_init == "randn":
449
+ s = torch.randn_like(e) * cfg.state_init_std
450
+ else:
451
+ s = e
452
+
453
+ readout_set = set(readout_loops or [])
454
+ aux_logits: dict[int, torch.Tensor] = {}
455
+ aux_stats: dict[int, dict] = {}
456
+ states = [s.detach()] if return_states else None
457
+
458
+ halt_logits: dict[int, torch.Tensor] = {}
459
+
460
+ def record(loop_idx: int, hidden: torch.Tensor) -> None:
461
+ hn = self._readout_hidden(hidden, cos, sin)
462
+ if cfg.halting == "ponder":
463
+ halt_logits[loop_idx] = self.halt_head(hn).squeeze(-1).reshape(-1)
464
+ lg = self.lm_head(hn)
465
+ if readout_mode == "logits":
466
+ aux_logits[loop_idx] = lg
467
+ return
468
+ flat = lg.float().view(-1, lg.size(-1))
469
+ tl = F.cross_entropy(flat, targets.reshape(-1), reduction="none")
470
+ aux_stats[loop_idx] = {
471
+ "token_loss": tl if readout_mode == "grad_stats" else tl.detach(),
472
+ "confidence": F.softmax(flat, dim=-1).max(dim=-1).values.detach(),
473
+ }
474
+
475
+ no_grad_until = 0
476
+ if backprop_loops and backprop_loops < R:
477
+ no_grad_until = R - backprop_loops
478
+
479
+ depth_table = self._depth_table(R, s.device, s.dtype) if cfg.depth_cond == "film" else None
480
+ step_scale = (
481
+ torch.tensor(1.0 / R, device=s.device, dtype=s.dtype)
482
+ if cfg.update_rule == "flow" else None
483
+ )
484
+ noise_table = (
485
+ torch.tensor([self._noise_std(r, R) for r in range(R + 1)], device=s.device, dtype=s.dtype)
486
+ if cfg.loop_noise > 0.0 else None
487
+ )
488
+
489
+ for r in range(1, R + 1):
490
+ depth_feat = depth_table[r] if depth_table is not None else None
491
+ noise = noise_table[r] if noise_table is not None else None
492
+ if r <= no_grad_until:
493
+ with torch.no_grad():
494
+ s = self.block(s, e, cos, sin, depth_feat, noise, step_scale)
495
+ s = s.detach()
496
+ elif grad_checkpoint and self.training:
497
+ s = torch.utils.checkpoint.checkpoint(
498
+ self.block, s, e, cos, sin, depth_feat, noise, step_scale,
499
+ use_reentrant=False,
500
+ )
501
+ else:
502
+ s = self.block(s, e, cos, sin, depth_feat, noise, step_scale)
503
+ if return_states:
504
+ states.append(s.detach())
505
+ if r in readout_set and r != R:
506
+ record(r, s)
507
+
508
+ final_hidden = self._readout_hidden(s, cos, sin)
509
+ logits = self.lm_head(final_hidden)
510
+ if cfg.halting == "ponder":
511
+ halt_logits[R] = self.halt_head(final_hidden).squeeze(-1).reshape(-1)
512
+ if readout_mode in {"stats", "grad_stats"} and R in readout_set:
513
+ flat = logits.float().view(-1, logits.size(-1))
514
+ tl = F.cross_entropy(flat, targets.reshape(-1), reduction="none")
515
+ aux_stats[R] = {
516
+ "token_loss": tl if readout_mode == "grad_stats" else tl.detach(),
517
+ "confidence": F.softmax(flat, dim=-1).max(dim=-1).values.detach(),
518
+ }
519
+
520
+ loss = None
521
+ if targets is not None:
522
+ loss = F.cross_entropy(
523
+ logits.float().view(-1, logits.size(-1)), targets.reshape(-1), ignore_index=-1
524
+ )
525
+
526
+ out = {"logits": logits, "loss": loss, "aux_logits": aux_logits,
527
+ "aux_stats": aux_stats, "halt_logits": halt_logits, "n_loops": R}
528
+ if return_states:
529
+ out["states"] = states
530
+ out["e"] = e.detach()
531
+ return out
532
+
533
+ # -- bookkeeping ------------------------------------------------------------------
534
+ def param_counts(self) -> dict:
535
+ total = sum(p.numel() for p in self.parameters())
536
+ emb = self.embed_tokens.weight.numel()
537
+ if not self.cfg.tie_embeddings:
538
+ emb += self.lm_head.weight.numel()
539
+ return {"total": total, "embedding": emb, "non_embedding": total - emb}
540
+
541
+ def flops_per_token(self, n_loops: Optional[int] = None) -> float:
542
+ """Forward FLOPs per token, counting matmuls only (attention scores included)."""
543
+ cfg = self.cfg
544
+ R = n_loops if n_loops is not None else cfg.n_loops
545
+ d, hd = cfg.d_model, cfg.head_dim
546
+ proj = 2 * d * (cfg.n_heads * hd) + 2 * d * (cfg.n_kv_heads * hd) # q,o and k,v
547
+ mlp = 3 * d * cfg.d_ff
548
+ attn_scores = 2 * cfg.n_heads * hd * cfg.max_seq_len / 2 # causal, averaged
549
+ per_layer = 2 * (proj + mlp) + 2 * attn_scores
550
+ n_layers = cfg.n_prelude + cfg.n_coda + cfg.n_recurrent * R
551
+ return per_layer * n_layers + 2 * d * cfg.vocab_size
552
+
553
+
554
+ def halting_distribution(halt_logits: torch.Tensor) -> torch.Tensor:
555
+ """Per-token distribution over the halting step, PonderNet style.
556
+
557
+ ``halt_logits`` is (R, N) pre-sigmoid. With ``lam_r`` the probability of
558
+ stopping at r given that r was reached,
559
+
560
+ p_r = lam_r * prod_{j<r} (1 - lam_j),
561
+
562
+ and all remaining mass is forced onto r = R, since the loop cannot run
563
+ further. Returns (R, N) summing to one along the loop axis.
564
+ """
565
+ # float32 and a loose clamp on purpose: under bf16 autocast, 1 - 1e-6 rounds
566
+ # to exactly 1.0, log1p(-1.0) is -inf, and the whole objective becomes NaN
567
+ # within a few hundred steps.
568
+ lam = torch.sigmoid(halt_logits.float()).clamp(1e-4, 1 - 1e-4)
569
+ log_not = torch.log1p(-lam)
570
+ # exclusive cumulative sum: log prod_{j<r} (1 - lam_j)
571
+ cum = torch.cumsum(log_not, dim=0) - log_not
572
+ p = lam * cum.exp()
573
+ leftover = (cum[-1] + log_not[-1]).exp()
574
+ return torch.cat([p[:-1], p[-1:] + leftover.unsqueeze(0)], dim=0)
575
+
576
+
577
+ def ponder_loss(token_losses: torch.Tensor, halt_logits: torch.Tensor,
578
+ prior: float = 0.1, kl_weight: float = 0.01):
579
+ """PonderNet objective: expected loss under the halting distribution, plus a
580
+ KL pull towards a geometric prior that sets the expected number of loops.
581
+
582
+ ``token_losses`` and ``halt_logits`` are both (R, N).
583
+ """
584
+ R = token_losses.shape[0]
585
+ p = halting_distribution(halt_logits)
586
+ token_losses = token_losses.float()
587
+ expected = (p * token_losses).sum(0).mean()
588
+
589
+ steps = torch.arange(1, R + 1, device=p.device, dtype=p.dtype).unsqueeze(1)
590
+ prior_p = prior * (1.0 - prior) ** (steps - 1)
591
+ prior_p = prior_p / prior_p.sum(0, keepdim=True)
592
+ kl = (p * (p.clamp_min(1e-9).log() - prior_p.log())).sum(0).mean()
593
+
594
+ expected_steps = (p * steps).sum(0).mean()
595
+ return expected + kl_weight * kl, {
596
+ "expected_loss": float(expected.detach()),
597
+ "kl": float(kl.detach()),
598
+ "expected_steps": float(expected_steps.detach()),
599
+ }
600
+
601
+
602
+ @torch.no_grad()
603
+ def q_exit(halt_logits: torch.Tensor, tau: float = 0.5) -> torch.Tensor:
604
+ """Deterministic exit step per token: the first r whose cumulative halting
605
+ probability reaches ``tau``. This is the PALBERT criterion, chosen over
606
+ sampling from the halting distribution because sampling adds variance to the
607
+ exit index for no benefit at inference.
608
+
609
+ Returns a (N,) tensor of 0-based loop indices.
610
+ """
611
+ p = halting_distribution(halt_logits)
612
+ reached = p.cumsum(0) >= tau
613
+ R = p.shape[0]
614
+ return torch.where(reached.any(0), reached.float().argmax(0),
615
+ torch.full((p.shape[1],), R - 1, device=p.device, dtype=torch.long))
616
+
617
+
618
+ def build_model(cfg: ModelConfig) -> LoopedLM:
619
+ return LoopedLM(cfg)
model_config.json ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "vocab_size": 8192,
3
+ "d_model": 384,
4
+ "n_heads": 6,
5
+ "n_kv_heads": 2,
6
+ "head_dim": 64,
7
+ "d_ff": 1024,
8
+ "max_seq_len": 512,
9
+ "rope_theta": 10000.0,
10
+ "rms_eps": 1e-06,
11
+ "tie_embeddings": true,
12
+ "n_prelude": 1,
13
+ "n_recurrent": 2,
14
+ "n_coda": 1,
15
+ "n_loops": 16,
16
+ "max_loops": 256,
17
+ "state_init": "prelude",
18
+ "state_init_std": 0.4,
19
+ "input_injection": "add",
20
+ "state_norm": "none",
21
+ "update_rule": "residual",
22
+ "update_gate_init": 3.0,
23
+ "depth_cond": "none",
24
+ "depth_cond_input": "progress",
25
+ "depth_cond_dim": 64,
26
+ "loop_noise": 0.0,
27
+ "noise_schedule": "linear",
28
+ "halting": "none",
29
+ "halt_prior": 0.1,
30
+ "halt_kl_weight": 0.01,
31
+ "init_std": 0.02,
32
+ "depth_scaled_init": true
33
+ }
pytorch_model.bin ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1aa66b8c90bc4e1c35e8760e8bc847c74dfd460ee8d4c9dd37abfca5aadae2f2
3
+ size 37782865
summary.json ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "run_name": "F_inject_add_s1_4090",
3
+ "params": {
4
+ "total": 9441152,
5
+ "embedding": 3145728,
6
+ "non_embedding": 6295424
7
+ },
8
+ "tokens_trained": 99999744,
9
+ "wall_clock_s": 3477.89067530632,
10
+ "val_loss_by_loops": {
11
+ "1": 7.405364990234375,
12
+ "2": 6.412084007263184,
13
+ "4": 5.2127032995224,
14
+ "8": 4.134036874771118,
15
+ "16": 3.7898231625556944,
16
+ "24": 3.8697949051856995,
17
+ "32": 3.9913308143615724,
18
+ "48": 4.1892451524734495,
19
+ "64": 4.327530193328857,
20
+ "96": 4.507761836051941,
21
+ "128": 4.6194603681564335
22
+ },
23
+ "val_ppl_by_loops": {
24
+ "1": 1644.7850570163394,
25
+ "2": 609.1618568928884,
26
+ "4": 183.5896858870107,
27
+ "8": 62.429434742451754,
28
+ "16": 44.248574779205214,
29
+ "24": 47.9325543543014,
30
+ "32": 54.12687427505751,
31
+ "48": 65.97297262494666,
32
+ "64": 75.75695030099443,
33
+ "96": 90.7185481353149,
34
+ "128": 101.43927749287803
35
+ },
36
+ "val_bpb_by_loops": {
37
+ "1": 2.839405015475947,
38
+ "2": 2.4585558596889587,
39
+ "4": 1.9986828349946704,
40
+ "8": 1.585094732247102,
41
+ "16": 1.4531144527944435,
42
+ "24": 1.4837776500061182,
43
+ "32": 1.5303776042999957,
44
+ "48": 1.6062629880738588,
45
+ "64": 1.6592849848407885,
46
+ "96": 1.7283903741052271,
47
+ "128": 1.7712184281845387
48
+ },
49
+ "best_val_loss": 3.7898231625556944,
50
+ "config": {
51
+ "model": {
52
+ "vocab_size": 8192,
53
+ "d_model": 384,
54
+ "n_heads": 6,
55
+ "n_kv_heads": 2,
56
+ "head_dim": 64,
57
+ "d_ff": 1024,
58
+ "max_seq_len": 512,
59
+ "rope_theta": 10000.0,
60
+ "rms_eps": 1e-06,
61
+ "tie_embeddings": true,
62
+ "n_prelude": 1,
63
+ "n_recurrent": 2,
64
+ "n_coda": 1,
65
+ "n_loops": 16,
66
+ "max_loops": 256,
67
+ "state_init": "prelude",
68
+ "state_init_std": 0.4,
69
+ "input_injection": "add",
70
+ "state_norm": "none",
71
+ "update_rule": "residual",
72
+ "update_gate_init": 3.0,
73
+ "depth_cond": "none",
74
+ "depth_cond_input": "progress",
75
+ "depth_cond_dim": 64,
76
+ "loop_noise": 0.0,
77
+ "noise_schedule": "linear",
78
+ "halting": "none",
79
+ "halt_prior": 0.1,
80
+ "halt_kl_weight": 0.01,
81
+ "init_std": 0.02,
82
+ "depth_scaled_init": true
83
+ },
84
+ "train": {
85
+ "data_dir": "/root/looped/data",
86
+ "out_dir": "/root/looped/runs",
87
+ "run_name": "F_inject_add_s1_4090",
88
+ "notes": "",
89
+ "max_tokens": 100000000,
90
+ "seq_len": 512,
91
+ "micro_batch": 16,
92
+ "grad_accum": 1,
93
+ "lr": 0.0015,
94
+ "min_lr_ratio": 0.1,
95
+ "warmup_frac": 0.02,
96
+ "weight_decay": 0.1,
97
+ "beta1": 0.9,
98
+ "beta2": 0.95,
99
+ "grad_clip": 1.0,
100
+ "seed": 1,
101
+ "loop_schedule": "fixed",
102
+ "loop_mean": 16,
103
+ "loop_min": 1,
104
+ "loop_max": 16,
105
+ "backprop_loops": 0,
106
+ "grad_checkpoint": false,
107
+ "aux_loss_weight": 0.0,
108
+ "aux_readouts": 1,
109
+ "device": "cuda",
110
+ "dtype": "bfloat16",
111
+ "compile": true,
112
+ "eval_every": 500,
113
+ "eval_batches": 20,
114
+ "log_every": 25,
115
+ "eval_loops": [
116
+ 1,
117
+ 2,
118
+ 4,
119
+ 8,
120
+ 16,
121
+ 24,
122
+ 32,
123
+ 48,
124
+ 64,
125
+ 96,
126
+ 128
127
+ ],
128
+ "save_checkpoint": true
129
+ },
130
+ "params": {
131
+ "total": 9441152,
132
+ "embedding": 3145728,
133
+ "non_embedding": 6295424
134
+ },
135
+ "tokens_per_step": 8192,
136
+ "total_steps": 12207,
137
+ "corpus": {
138
+ "tokenizer_vocab_size": 8192,
139
+ "train_tokens": 110197899,
140
+ "val_tokens": 1674551,
141
+ "train_docs": 133732,
142
+ "val_docs": 2098,
143
+ "train_bytes": 414483859,
144
+ "val_bytes": 6300747,
145
+ "train_bytes_per_token": 3.7612682524918193,
146
+ "val_bytes_per_token": 3.7626486144644145,
147
+ "shard": "sample/10BT/000_00000.parquet",
148
+ "val_every": 500
149
+ },
150
+ "torch": "2.9.0+cu129",
151
+ "gpu": "NVIDIA GeForce RTX 4090"
152
+ }
153
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
training_config.json ADDED
@@ -0,0 +1,103 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model": {
3
+ "vocab_size": 8192,
4
+ "d_model": 384,
5
+ "n_heads": 6,
6
+ "n_kv_heads": 2,
7
+ "head_dim": 64,
8
+ "d_ff": 1024,
9
+ "max_seq_len": 512,
10
+ "rope_theta": 10000.0,
11
+ "rms_eps": 1e-06,
12
+ "tie_embeddings": true,
13
+ "n_prelude": 1,
14
+ "n_recurrent": 2,
15
+ "n_coda": 1,
16
+ "n_loops": 16,
17
+ "max_loops": 256,
18
+ "state_init": "prelude",
19
+ "state_init_std": 0.4,
20
+ "input_injection": "add",
21
+ "state_norm": "none",
22
+ "update_rule": "residual",
23
+ "update_gate_init": 3.0,
24
+ "depth_cond": "none",
25
+ "depth_cond_input": "progress",
26
+ "depth_cond_dim": 64,
27
+ "loop_noise": 0.0,
28
+ "noise_schedule": "linear",
29
+ "halting": "none",
30
+ "halt_prior": 0.1,
31
+ "halt_kl_weight": 0.01,
32
+ "init_std": 0.02,
33
+ "depth_scaled_init": true
34
+ },
35
+ "train": {
36
+ "data_dir": "/root/looped/data",
37
+ "out_dir": "/root/looped/runs",
38
+ "run_name": "F_inject_add_s1_4090",
39
+ "notes": "",
40
+ "max_tokens": 100000000,
41
+ "seq_len": 512,
42
+ "micro_batch": 16,
43
+ "grad_accum": 1,
44
+ "lr": 0.0015,
45
+ "min_lr_ratio": 0.1,
46
+ "warmup_frac": 0.02,
47
+ "weight_decay": 0.1,
48
+ "beta1": 0.9,
49
+ "beta2": 0.95,
50
+ "grad_clip": 1.0,
51
+ "seed": 1,
52
+ "loop_schedule": "fixed",
53
+ "loop_mean": 16,
54
+ "loop_min": 1,
55
+ "loop_max": 16,
56
+ "backprop_loops": 0,
57
+ "grad_checkpoint": false,
58
+ "aux_loss_weight": 0.0,
59
+ "aux_readouts": 1,
60
+ "device": "cuda",
61
+ "dtype": "bfloat16",
62
+ "compile": true,
63
+ "eval_every": 500,
64
+ "eval_batches": 20,
65
+ "log_every": 25,
66
+ "eval_loops": [
67
+ 1,
68
+ 2,
69
+ 4,
70
+ 8,
71
+ 16,
72
+ 24,
73
+ 32,
74
+ 48,
75
+ 64,
76
+ 96,
77
+ 128
78
+ ],
79
+ "save_checkpoint": true
80
+ },
81
+ "params": {
82
+ "total": 9441152,
83
+ "embedding": 3145728,
84
+ "non_embedding": 6295424
85
+ },
86
+ "tokens_per_step": 8192,
87
+ "total_steps": 12207,
88
+ "corpus": {
89
+ "tokenizer_vocab_size": 8192,
90
+ "train_tokens": 110197899,
91
+ "val_tokens": 1674551,
92
+ "train_docs": 133732,
93
+ "val_docs": 2098,
94
+ "train_bytes": 414483859,
95
+ "val_bytes": 6300747,
96
+ "train_bytes_per_token": 3.7612682524918193,
97
+ "val_bytes_per_token": 3.7626486144644145,
98
+ "shard": "sample/10BT/000_00000.parquet",
99
+ "val_every": 500
100
+ },
101
+ "torch": "2.9.0+cu129",
102
+ "gpu": "NVIDIA GeForce RTX 4090"
103
+ }