File size: 28,074 Bytes
4f0fb70
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
"""Looped decoder-only LM with a Qwen3-style block.



Layout follows the prelude / recurrent / coda decomposition:



    x -> embed -> [prelude L_p layers] -> e

    s_0 = e

    s_r = Block(s_{r-1}, e, r, R)          for r = 1..R      (shared weights)

    logits = head(norm(coda(s_R)))



Every research knob is a config flag so that one binary can produce the whole

ablation ladder and every run is described by its config dict alone.

"""

from __future__ import annotations

import math
from dataclasses import asdict, dataclass, field
from typing import Optional

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint


# --------------------------------------------------------------------------------------
# config
# --------------------------------------------------------------------------------------
@dataclass
class ModelConfig:
    # --- Qwen3-style backbone -----------------------------------------------------
    vocab_size: int = 8192
    d_model: int = 384
    n_heads: int = 6
    n_kv_heads: int = 2
    head_dim: int = 64
    d_ff: int = 1024
    max_seq_len: int = 512
    rope_theta: float = 10_000.0
    rms_eps: float = 1e-6
    tie_embeddings: bool = True
    # pre = Qwen3 default. sandwich = Huginn's block, which normalises after each
    # residual add as well; costs 2d per layer and bounds the residual stream.
    block_norm: str = "pre"           # pre | sandwich

    # --- depth layout -------------------------------------------------------------
    n_prelude: int = 1
    n_recurrent: int = 2
    n_coda: int = 1

    # --- looping ------------------------------------------------------------------
    n_loops: int = 8                 # R used at train time (mean of the distribution if sampled)
    max_loops: int = 256             # size of the precomputed depth-embedding table
    state_init: str = "prelude"      # prelude | randn
    state_init_std: float = 0.4      # only for state_init == "randn"

    input_injection: str = "add"     # none | add | adapter
    state_norm: str = "none"         # none | rms          (normalise s at loop entry)
    # residual : s <- Block(s)                        (the usual looped transformer)
    # convex   : s <- (1-a) s + a Block(s)            (learned step size)
    # flow     : s <- s + (gain/R) * Delta(s, r/R)    (explicit Euler step of a learned flow)
    update_rule: str = "residual"
    # pre-sigmoid init of the convex step size. +3 starts at ~0.95, i.e. almost a
    # full replacement (the usual looped behaviour); -3 starts at ~0.05, so the
    # loop begins as a near-identity and has to earn its depth, which is what
    # makes very deep shared stacks trainable at all
    update_gate_init: float = 3.0
    depth_cond: str = "none"         # none | film
    depth_cond_input: str = "progress"   # absolute | progress | both
    depth_cond_dim: int = 64

    loop_noise: float = 0.0          # std of exploration noise injected at loop entry
    noise_schedule: str = "linear"   # linear | const | cosine   (annealed towards 0 at r=R)

    # learned halting, PonderNet style: a per-token probability of stopping after
    # each iteration, trained jointly with the language-model loss. Costs d + 1
    # parameters and is independent of R, so the maximum depth stays a runtime knob.
    halting: str = "none"            # none | ponder
    halt_prior: float = 0.1          # geometric prior on the halting step
    halt_kl_weight: float = 0.01

    # --- init ---------------------------------------------------------------------
    init_std: float = 0.02
    depth_scaled_init: bool = True

    def __post_init__(self) -> None:
        assert self.n_heads % self.n_kv_heads == 0
        assert self.state_init in {"prelude", "randn"}
        assert self.input_injection in {"none", "add", "adapter"}
        assert self.block_norm in {"pre", "sandwich"}
        assert self.state_norm in {"none", "rms"}
        assert self.update_rule in {"residual", "convex", "flow"}
        assert self.depth_cond in {"none", "film"}
        assert self.depth_cond_input in {"absolute", "progress", "both"}
        assert self.noise_schedule in {"linear", "const", "cosine"}
        assert self.halting in {"none", "ponder"}

    def to_dict(self) -> dict:
        return asdict(self)


# --------------------------------------------------------------------------------------
# primitives
# --------------------------------------------------------------------------------------
class RMSNorm(nn.Module):
    def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = True):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim)) if elementwise_affine else None

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        dtype = x.dtype
        x = x.float()
        x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
        x = x.to(dtype)
        return x * self.weight if self.weight is not None else x


def build_rope_cache(seq_len: int, head_dim: int, theta: float, device, dtype=torch.float32):
    inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim))
    t = torch.arange(seq_len, device=device, dtype=torch.float32)
    freqs = torch.outer(t, inv_freq)              # (T, hd/2)
    emb = torch.cat((freqs, freqs), dim=-1)       # (T, hd)
    return emb.cos().to(dtype), emb.sin().to(dtype)


def rotate_half(x: torch.Tensor) -> torch.Tensor:
    x1, x2 = x.chunk(2, dim=-1)
    return torch.cat((-x2, x1), dim=-1)


def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
    # x: (B, H, T, hd); cos/sin: (T, hd)
    cos = cos[None, None, :, :]
    sin = sin[None, None, :, :]
    return x * cos + rotate_half(x) * sin


class Attention(nn.Module):
    """Qwen3 attention: GQA, no qkv bias, RMSNorm on q and k heads."""

    def __init__(self, cfg: ModelConfig):
        super().__init__()
        self.n_heads = cfg.n_heads
        self.n_kv_heads = cfg.n_kv_heads
        self.head_dim = cfg.head_dim
        self.q_proj = nn.Linear(cfg.d_model, cfg.n_heads * cfg.head_dim, bias=False)
        self.k_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False)
        self.v_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False)
        self.o_proj = nn.Linear(cfg.n_heads * cfg.head_dim, cfg.d_model, bias=False)
        self.q_norm = RMSNorm(cfg.head_dim, cfg.rms_eps)
        self.k_norm = RMSNorm(cfg.head_dim, cfg.rms_eps)

    def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
        B, T, _ = x.shape
        q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
        k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
        v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)

        q = self.q_norm(q)
        k = self.k_norm(k)
        q = apply_rope(q, cos, sin)
        k = apply_rope(k, cos, sin)

        o = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=True)
        o = o.transpose(1, 2).contiguous().view(B, T, self.n_heads * self.head_dim)
        return self.o_proj(o)


class MLP(nn.Module):
    def __init__(self, cfg: ModelConfig):
        super().__init__()
        self.gate_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
        self.up_proj = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
        self.down_proj = nn.Linear(cfg.d_ff, cfg.d_model, bias=False)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))


class DecoderLayer(nn.Module):
    """Qwen3 pre-norm layer, optionally with Huginn's sandwich norm.



    Pre-norm (`block_norm="pre"`) is the Qwen3 default: the stream is normalised

    on the way *into* each sublayer and the residual add is left alone, so the

    stream is free to grow. Section 4.2 measures that growth and identifies it as

    the reason late iterations stop mattering.



    Sandwich (`block_norm="sandwich"`) is what the Huginn recurrent block

    actually does: it normalises again *after* each residual add, which bounds

    the stream without removing the residual path itself. That distinction is the

    whole reason the loop-entry normalisation of 5.2 failed and this does not.

    """

    def __init__(self, cfg: ModelConfig):
        super().__init__()
        self.input_layernorm = RMSNorm(cfg.d_model, cfg.rms_eps)
        self.self_attn = Attention(cfg)
        self.post_attention_layernorm = RMSNorm(cfg.d_model, cfg.rms_eps)
        self.mlp = MLP(cfg)
        self.sandwich = cfg.block_norm == "sandwich"
        if self.sandwich:
            self.post_attn_residual_norm = RMSNorm(cfg.d_model, cfg.rms_eps)
            self.post_mlp_residual_norm = RMSNorm(cfg.d_model, cfg.rms_eps)

    def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
        x = x + self.self_attn(self.input_layernorm(x), cos, sin)
        if self.sandwich:
            x = self.post_attn_residual_norm(x)
        x = x + self.mlp(self.post_attention_layernorm(x))
        if self.sandwich:
            x = self.post_mlp_residual_norm(x)
        return x


# --------------------------------------------------------------------------------------
# recurrent block
# --------------------------------------------------------------------------------------
class RecurrentBlock(nn.Module):
    """The shared block applied R times.



    Everything that makes iteration r behave differently from iteration r+1 has

    to enter here, because the weights themselves are identical across r.

    """

    def __init__(self, cfg: ModelConfig):
        super().__init__()
        self.cfg = cfg
        self.layers = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_recurrent)])

        if cfg.input_injection == "adapter":
            self.adapter = nn.Linear(2 * cfg.d_model, cfg.d_model, bias=False)

        if cfg.state_norm == "rms":
            self.entry_norm = RMSNorm(cfg.d_model, cfg.rms_eps)

        if cfg.depth_cond == "film":
            # sinusoidal features -> (scale, shift). Cost is O(d), independent of R,
            # which is what keeps this usable at larger scale.
            self.film = nn.Linear(cfg.depth_cond_dim, 2 * cfg.d_model, bias=True)
            nn.init.zeros_(self.film.weight)
            nn.init.zeros_(self.film.bias)

        if cfg.update_rule == "convex":
            # learned per-channel step size, sigmoid-gated, initialised near 1.0 so the
            # untouched model starts out identical to the plain residual update
            self.alpha = nn.Parameter(torch.full((cfg.d_model,), float(cfg.update_gate_init)))
        elif cfg.update_rule == "flow":
            # learned per-channel speed of the flow; the 1/R factor lives in forward()
            self.flow_gain = nn.Parameter(torch.ones(cfg.d_model))

    def forward(self, s, e, cos, sin, depth_feat: Optional[torch.Tensor] = None,

                noise_std: Optional[torch.Tensor] = None, step_scale: Optional[torch.Tensor] = None):
        # noise_std and step_scale arrive as 0-dim tensors on purpose: as python
        # floats dynamo specialises the graph on their value and recompiles the
        # block for every distinct loop count and noise level.
        cfg = self.cfg
        h = s

        if cfg.input_injection == "add":
            h = h + e
        elif cfg.input_injection == "adapter":
            h = self.adapter(torch.cat([h, e], dim=-1))

        if cfg.state_norm == "rms":
            h = self.entry_norm(h)

        if cfg.depth_cond == "film" and depth_feat is not None:
            mod = self.film(depth_feat)                    # (2d,)
            scale, shift = mod.chunk(2, dim=-1)
            h = h * (1.0 + scale) + shift

        if cfg.loop_noise > 0.0 and noise_std is not None:
            h = h + noise_std * torch.randn_like(h)

        inner = h
        for layer in self.layers:
            inner = layer(inner, cos, sin)

        if cfg.update_rule == "convex":
            a = torch.sigmoid(self.alpha)
            return (1.0 - a) * s + a * inner
        if cfg.update_rule == "flow":
            # explicit Euler step: s' = s + h_step * g(s, r/R). The loop count then
            # sets the integration resolution rather than the amount of drift, so
            # raising R at inference refines the same trajectory instead of
            # walking further along it.
            return s + (step_scale * self.flow_gain) * (inner - h)
        return inner


# --------------------------------------------------------------------------------------
# full model
# --------------------------------------------------------------------------------------
class LoopedLM(nn.Module):
    def __init__(self, cfg: ModelConfig):
        super().__init__()
        self.cfg = cfg
        self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.d_model)
        self.prelude = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_prelude)])
        self.block = RecurrentBlock(cfg)
        self.coda = nn.ModuleList([DecoderLayer(cfg) for _ in range(cfg.n_coda)])
        self.norm = RMSNorm(cfg.d_model, cfg.rms_eps)
        if cfg.halting == "ponder":
            self.halt_head = nn.Linear(cfg.d_model, 1)
        self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
        if cfg.tie_embeddings:
            self.lm_head.weight = self.embed_tokens.weight

        cos, sin = build_rope_cache(cfg.max_seq_len, cfg.head_dim, cfg.rope_theta, device="cpu")
        self.register_buffer("rope_cos", cos, persistent=False)
        self.register_buffer("rope_sin", sin, persistent=False)

        self._depth_cache: dict[tuple, torch.Tensor] = {}

        self.apply(self._init_weights)
        if cfg.depth_scaled_init:
            self._rescale_residual_projections()
        if cfg.depth_cond == "film":
            # zero-init the modulation so an untrained depth-conditioned model is
            # bit-identical to the unconditioned one at step 0
            nn.init.zeros_(self.block.film.weight)
            nn.init.zeros_(self.block.film.bias)

    # -- init ------------------------------------------------------------------------
    def _init_weights(self, module: nn.Module) -> None:
        std = self.cfg.init_std
        if isinstance(module, nn.Linear):
            nn.init.normal_(module.weight, mean=0.0, std=std)
            if module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.Embedding):
            nn.init.normal_(module.weight, mean=0.0, std=std)

    def _rescale_residual_projections(self) -> None:
        """GPT-2 style 1/sqrt(2 * depth) scaling, with depth counted through the loop.



        The flow update already divides every step by R, so counting the loop

        twice would leave the block effectively dead at initialisation.

        """
        loops = 1 if self.cfg.update_rule == "flow" else self.cfg.n_loops
        depth = self.cfg.n_prelude + self.cfg.n_recurrent * loops + self.cfg.n_coda
        scale = 1.0 / math.sqrt(2.0 * max(depth, 1))
        for mod in self.modules():
            if isinstance(mod, DecoderLayer):
                mod.self_attn.o_proj.weight.data.mul_(scale)
                mod.mlp.down_proj.weight.data.mul_(scale)

    def _depth_table(self, R: int, device, dtype) -> torch.Tensor:
        """Sinusoidal encodings of every loop index, shape (R + 1, depth_cond_dim).



        Two things can be encoded: the absolute index r (tells the block how much

        work has been done) and the progress r/R (tells it how much is left).

        Which one matters is an experiment, not an assumption, hence the flag.

        Cached per (R, device, dtype) so the table is built once per run.

        """
        cfg = self.cfg
        key = (R, str(device), str(dtype))
        if key in self._depth_cache:
            return self._depth_cache[key]

        r = torch.arange(R + 1, device=device, dtype=torch.float32)
        vals = []
        if cfg.depth_cond_input in {"absolute", "both"}:
            vals.append(r)
        if cfg.depth_cond_input in {"progress", "both"}:
            vals.append(r / max(R, 1) * 32.0)  # rescale so low frequencies stay informative

        per = cfg.depth_cond_dim // (2 * len(vals))
        idx = torch.arange(per, device=device, dtype=torch.float32)
        freq = torch.exp(-math.log(10_000.0) * idx / max(per - 1, 1))
        feats = []
        for v in vals:
            ang = v[:, None] * freq[None, :]
            feats.append(torch.cat([torch.sin(ang), torch.cos(ang)], dim=-1))
        out = torch.cat(feats, dim=-1)
        if out.shape[-1] < cfg.depth_cond_dim:
            out = F.pad(out, (0, cfg.depth_cond_dim - out.shape[-1]))
        out = out.to(dtype)
        self._depth_cache[key] = out
        return out

    def _noise_std(self, r: int, R: int) -> float:
        cfg = self.cfg
        if cfg.loop_noise <= 0.0 or not self.training:
            return 0.0
        if cfg.noise_schedule == "const":
            return cfg.loop_noise
        frac = (r - 1) / max(R - 1, 1)
        if cfg.noise_schedule == "linear":
            return cfg.loop_noise * (1.0 - frac)
        return cfg.loop_noise * 0.5 * (1.0 + math.cos(math.pi * frac))

    # -- forward ----------------------------------------------------------------------
    def _readout_hidden(self, s: torch.Tensor, cos, sin) -> torch.Tensor:
        """Coda output after the final norm; shared by the LM head and the halting head."""
        h = s
        for layer in self.coda:
            h = layer(h, cos, sin)
        return self.norm(h)

    def _readout(self, s: torch.Tensor, cos, sin) -> torch.Tensor:
        return self.lm_head(self._readout_hidden(s, cos, sin))

    def forward(

        self,

        idx: torch.Tensor,

        targets: Optional[torch.Tensor] = None,

        n_loops: Optional[int] = None,

        backprop_loops: int = 0,

        readout_loops: Optional[list[int]] = None,

        return_states: bool = False,

        grad_checkpoint: bool = False,

        readout_mode: str = "logits",

    ):
        """Run the model.



        Args:

            n_loops: R for this call (defaults to cfg.n_loops).

            backprop_loops: if > 0, only the last k iterations carry gradient.

            readout_loops: loop indices (1-based) whose intermediate logits are

                also returned, used for deep supervision and for the coda lens.

            return_states: also return the per-loop hidden states (diagnostics).

            grad_checkpoint: recompute each iteration's internals in the backward

                pass. Activation memory then stops growing with R, so a *full*

                backward through 32 or 64 loops fits, which truncation does not

                achieve without also changing what is being optimised.

            readout_mode: "logits" keeps every intermediate logit tensor, which is

                what deep supervision needs. "stats" reduces each one to per-token

                loss and confidence immediately and throws the logits away; a

                (B, T, 8192) tensor per loop is ~130 MB, so reading out all 32

                loops for diagnostics costs gigabytes otherwise. "grad_stats" is

                the same reduction but keeps the graph, which is what the ponder

                objective needs: it weights every loop's loss by a learned halting

                probability and so requires all of them to be differentiable.

        """
        cfg = self.cfg
        B, T = idx.shape
        R = n_loops if n_loops is not None else cfg.n_loops
        cos = self.rope_cos[:T].to(idx.device)
        sin = self.rope_sin[:T].to(idx.device)

        h = self.embed_tokens(idx)
        for layer in self.prelude:
            h = layer(h, cos, sin)
        e = h

        if cfg.state_init == "randn":
            s = torch.randn_like(e) * cfg.state_init_std
        else:
            s = e

        readout_set = set(readout_loops or [])
        aux_logits: dict[int, torch.Tensor] = {}
        aux_stats: dict[int, dict] = {}
        states = [s.detach()] if return_states else None

        halt_logits: dict[int, torch.Tensor] = {}

        def record(loop_idx: int, hidden: torch.Tensor) -> None:
            hn = self._readout_hidden(hidden, cos, sin)
            if cfg.halting == "ponder":
                halt_logits[loop_idx] = self.halt_head(hn).squeeze(-1).reshape(-1)
            lg = self.lm_head(hn)
            if readout_mode == "logits":
                aux_logits[loop_idx] = lg
                return
            flat = lg.float().view(-1, lg.size(-1))
            tl = F.cross_entropy(flat, targets.reshape(-1), reduction="none")
            aux_stats[loop_idx] = {
                "token_loss": tl if readout_mode == "grad_stats" else tl.detach(),
                "confidence": F.softmax(flat, dim=-1).max(dim=-1).values.detach(),
            }

        no_grad_until = 0
        if backprop_loops and backprop_loops < R:
            no_grad_until = R - backprop_loops

        depth_table = self._depth_table(R, s.device, s.dtype) if cfg.depth_cond == "film" else None
        step_scale = (
            torch.tensor(1.0 / R, device=s.device, dtype=s.dtype)
            if cfg.update_rule == "flow" else None
        )
        noise_table = (
            torch.tensor([self._noise_std(r, R) for r in range(R + 1)], device=s.device, dtype=s.dtype)
            if cfg.loop_noise > 0.0 else None
        )

        for r in range(1, R + 1):
            depth_feat = depth_table[r] if depth_table is not None else None
            noise = noise_table[r] if noise_table is not None else None
            if r <= no_grad_until:
                with torch.no_grad():
                    s = self.block(s, e, cos, sin, depth_feat, noise, step_scale)
                s = s.detach()
            elif grad_checkpoint and self.training:
                s = torch.utils.checkpoint.checkpoint(
                    self.block, s, e, cos, sin, depth_feat, noise, step_scale,
                    use_reentrant=False,
                )
            else:
                s = self.block(s, e, cos, sin, depth_feat, noise, step_scale)
            if return_states:
                states.append(s.detach())
            if r in readout_set and r != R:
                record(r, s)

        final_hidden = self._readout_hidden(s, cos, sin)
        logits = self.lm_head(final_hidden)
        if cfg.halting == "ponder":
            halt_logits[R] = self.halt_head(final_hidden).squeeze(-1).reshape(-1)
        if readout_mode in {"stats", "grad_stats"} and R in readout_set:
            flat = logits.float().view(-1, logits.size(-1))
            tl = F.cross_entropy(flat, targets.reshape(-1), reduction="none")
            aux_stats[R] = {
                "token_loss": tl if readout_mode == "grad_stats" else tl.detach(),
                "confidence": F.softmax(flat, dim=-1).max(dim=-1).values.detach(),
            }

        loss = None
        if targets is not None:
            loss = F.cross_entropy(
                logits.float().view(-1, logits.size(-1)), targets.reshape(-1), ignore_index=-1
            )

        out = {"logits": logits, "loss": loss, "aux_logits": aux_logits,
               "aux_stats": aux_stats, "halt_logits": halt_logits, "n_loops": R}
        if return_states:
            out["states"] = states
            out["e"] = e.detach()
        return out

    # -- bookkeeping ------------------------------------------------------------------
    def param_counts(self) -> dict:
        total = sum(p.numel() for p in self.parameters())
        emb = self.embed_tokens.weight.numel()
        if not self.cfg.tie_embeddings:
            emb += self.lm_head.weight.numel()
        return {"total": total, "embedding": emb, "non_embedding": total - emb}

    def flops_per_token(self, n_loops: Optional[int] = None) -> float:
        """Forward FLOPs per token, counting matmuls only (attention scores included)."""
        cfg = self.cfg
        R = n_loops if n_loops is not None else cfg.n_loops
        d, hd = cfg.d_model, cfg.head_dim
        proj = 2 * d * (cfg.n_heads * hd) + 2 * d * (cfg.n_kv_heads * hd)   # q,o and k,v
        mlp = 3 * d * cfg.d_ff
        attn_scores = 2 * cfg.n_heads * hd * cfg.max_seq_len / 2            # causal, averaged
        per_layer = 2 * (proj + mlp) + 2 * attn_scores
        n_layers = cfg.n_prelude + cfg.n_coda + cfg.n_recurrent * R
        return per_layer * n_layers + 2 * d * cfg.vocab_size


def halting_distribution(halt_logits: torch.Tensor) -> torch.Tensor:
    """Per-token distribution over the halting step, PonderNet style.



    ``halt_logits`` is (R, N) pre-sigmoid. With ``lam_r`` the probability of

    stopping at r given that r was reached,



        p_r = lam_r * prod_{j<r} (1 - lam_j),



    and all remaining mass is forced onto r = R, since the loop cannot run

    further. Returns (R, N) summing to one along the loop axis.

    """
    # float32 and a loose clamp on purpose: under bf16 autocast, 1 - 1e-6 rounds
    # to exactly 1.0, log1p(-1.0) is -inf, and the whole objective becomes NaN
    # within a few hundred steps.
    lam = torch.sigmoid(halt_logits.float()).clamp(1e-4, 1 - 1e-4)
    log_not = torch.log1p(-lam)
    # exclusive cumulative sum: log prod_{j<r} (1 - lam_j)
    cum = torch.cumsum(log_not, dim=0) - log_not
    p = lam * cum.exp()
    leftover = (cum[-1] + log_not[-1]).exp()
    return torch.cat([p[:-1], p[-1:] + leftover.unsqueeze(0)], dim=0)


def ponder_loss(token_losses: torch.Tensor, halt_logits: torch.Tensor,

                prior: float = 0.1, kl_weight: float = 0.01):
    """PonderNet objective: expected loss under the halting distribution, plus a

    KL pull towards a geometric prior that sets the expected number of loops.



    ``token_losses`` and ``halt_logits`` are both (R, N).

    """
    R = token_losses.shape[0]
    p = halting_distribution(halt_logits)
    token_losses = token_losses.float()
    expected = (p * token_losses).sum(0).mean()

    steps = torch.arange(1, R + 1, device=p.device, dtype=p.dtype).unsqueeze(1)
    prior_p = prior * (1.0 - prior) ** (steps - 1)
    prior_p = prior_p / prior_p.sum(0, keepdim=True)
    kl = (p * (p.clamp_min(1e-9).log() - prior_p.log())).sum(0).mean()

    expected_steps = (p * steps).sum(0).mean()
    return expected + kl_weight * kl, {
        "expected_loss": float(expected.detach()),
        "kl": float(kl.detach()),
        "expected_steps": float(expected_steps.detach()),
    }


@torch.no_grad()
def q_exit(halt_logits: torch.Tensor, tau: float = 0.5) -> torch.Tensor:
    """Deterministic exit step per token: the first r whose cumulative halting

    probability reaches ``tau``. This is the PALBERT criterion, chosen over

    sampling from the halting distribution because sampling adds variance to the

    exit index for no benefit at inference.



    Returns a (N,) tensor of 0-based loop indices.

    """
    p = halting_distribution(halt_logits)
    reached = p.cumsum(0) >= tau
    R = p.shape[0]
    return torch.where(reached.any(0), reached.float().argmax(0),
                       torch.full((p.shape[1],), R - 1, device=p.device, dtype=torch.long))


def build_model(cfg: ModelConfig) -> LoopedLM:
    return LoopedLM(cfg)