brkdrd commited on
Commit
c6c8db8
·
verified ·
1 Parent(s): c8448d2

Upload folder using huggingface_hub

Browse files
README.md CHANGED
@@ -1,3 +1,98 @@
1
  ---
2
- license: mit
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ license: apache-2.0
3
+ datasets:
4
+ - HuggingFaceFW/fineweb
5
+ language:
6
+ - en
7
+ library_name: pytorch
8
+ tags:
9
+ - looped-transformer
10
+ - weight-tying
11
+ - recurrent-depth
12
+ - small-language-model
13
  ---
14
+
15
+ # SimLoop 1+loop x1+1 (10.0M)
16
+
17
+ A **looped** (weight-tied recurrent) Qwen3-style transformer trained from
18
+ scratch on FineWeb under a hard budget of **<= 10M parameters** and
19
+ **<= 100M training tokens**.
20
+
21
+ One middle block is applied **K = 1** times; the layers around it are
22
+ ordinary unlooped layers:
23
+
24
+ ```
25
+ embed -> [1 layer] -> ( 1 looped layer ) x K -> [1 layer] -> RMSNorm -> tied head
26
+ ```
27
+
28
+ | | |
29
+ |---|---|
30
+ | parameters | **9,962,496** (incl. embeddings, input/output tied) |
31
+ | d_model / d_mlp | 384 / 720 |
32
+ | heads (GQA) | 6 query / 2 key-value, head_dim 64 |
33
+ | context | 512 tokens |
34
+ | vocabulary | 16,384 byte-level BPE trained on FineWeb |
35
+ | training tokens | 60,014,592 |
36
+ | primitives | RMSNorm, RoPE, SwiGLU, GQA, QK-norm |
37
+
38
+ ## Results
39
+
40
+ | k | CE (nats) | perplexity | bits-per-byte |
41
+ |---|---|---|---|
42
+ | 0 | 5.4088 | 223.35 | 1.8913 |
43
+ | 1 | 4.3966 | 81.17 | 1.5374 **<- best** |
44
+ | 2 | 4.5995 | 99.43 | 1.6083 |
45
+ | 3 | 4.9413 | 139.95 | 1.7278 |
46
+ | 4 | 5.2659 | 193.62 | 1.8413 |
47
+ | 5 | 5.5476 | 256.63 | 1.9399 |
48
+ | 6 | 5.7900 | 327.01 | 2.0246 |
49
+ | 7 | 6.0000 | 403.41 | 2.0980 |
50
+ | 8 | 6.1835 | 484.69 | 2.1622 |
51
+ | 9 | 6.3452 | 569.77 | 2.2188 |
52
+ | 10 | 6.4887 | 657.65 | 2.2689 |
53
+ | 11 | 6.6166 | 747.41 | 2.3137 |
54
+ | 12 | 6.7313 | 838.20 | 2.3537 |
55
+
56
+ Reference points on the same validation split: a context-free **unigram** model scores CE 7.5476 (ppl 1896.10); **uniform** over the vocabulary scores CE 9.7041.
57
+
58
+
59
+ `k` is the number of applications of the looped block at inference. The model
60
+ is weight-tied, so **any k can be run**; the table is a single checkpoint
61
+ evaluated at every depth.
62
+
63
+ **Measured behaviour:** TRAINED UNLOOPED (K=1): the k>1 rows show how a model trained at one application behaves when looped anyway, not whether looping pays.
64
+
65
+ - first application of the looped block buys **+1.0122** nats (k=0 -> k=1)
66
+ - all further applications buy **+0.0000** nats (k=1 -> k=1)
67
+
68
+
69
+ ## Usage
70
+
71
+ ```python
72
+ import torch
73
+ from simloop.stack import StackConfig, StackedLoop
74
+
75
+ ck = torch.load("model.pt", map_location="cpu", weights_only=False)
76
+ model = StackedLoop(StackConfig(**ck["model_cfg"]))
77
+ model.load_state_dict(ck["model"]); model.eval()
78
+
79
+ ids = torch.tensor([[1, 2, 3]]) # from tokenizer.json
80
+ logits = model(ids, K=1) # try other K: the block is tied
81
+ ```
82
+
83
+ See `load_example.py`. The tokenizer is a `tokenizers` BPE:
84
+ `Tokenizer.from_file("tokenizer.json")`.
85
+
86
+ ## Honest limitations
87
+
88
+ - Trained on 60,014,592 tokens at ~10M parameters. It is a research artifact
89
+ for studying looped depth, **not** a useful general-purpose language model.
90
+ - Perplexity is tokenizer-dependent; **bits-per-byte** is the comparable
91
+ number and is reported above.
92
+ - The looped block **saturates**: past the depth listed as best above, extra
93
+ applications make cross-entropy worse, not better. This is measured, not
94
+ assumed, and is the central finding of the project.
95
+ - English-only, no instruction tuning, no safety filtering beyond FineWeb's.
96
+
97
+ Full experimental record, including every failed experiment:
98
+ https://github.com/brkdrd/SimLoop
config.json ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "arch": "stack",
3
+ "model": {
4
+ "vocab_size": 16384,
5
+ "d_model": 384,
6
+ "n_pre": 1,
7
+ "n_loop": 1,
8
+ "n_post": 1,
9
+ "n_heads": 6,
10
+ "n_kv_heads": 2,
11
+ "head_dim": 64,
12
+ "d_mlp": 720,
13
+ "max_seq_len": 512,
14
+ "rope_theta": 10000.0,
15
+ "dropout": 0.0,
16
+ "norm_eps": 1e-06,
17
+ "init_std": 0.02,
18
+ "scale_resid_init": true
19
+ },
20
+ "train": {
21
+ "K": 1,
22
+ "k_schedule": "const",
23
+ "k_end": 1,
24
+ "k_decay_start_frac": 0.5,
25
+ "k_decay_end_frac": 0.9,
26
+ "k_min": 1,
27
+ "drop_eps": 0.01,
28
+ "eval_k_max": 8,
29
+ "batch_size": 16,
30
+ "grad_accum": 2,
31
+ "seq_len": 512,
32
+ "max_tokens": 60000000,
33
+ "lr": 0.001,
34
+ "min_lr_ratio": 0.1,
35
+ "warmup_steps": 100,
36
+ "weight_decay": 0.1,
37
+ "beta1": 0.9,
38
+ "beta2": 0.95,
39
+ "grad_clip": 1.0,
40
+ "ce_weight": 1.0,
41
+ "teacher_run": "",
42
+ "teacher_k": 0,
43
+ "init_from_run": "",
44
+ "distill_weight": 0.0,
45
+ "distill_space": "hidden",
46
+ "distill_depths": "first",
47
+ "distill_start_frac": 0.0,
48
+ "distill_tau": 1.0,
49
+ "anytime_ce_weight": 0.0,
50
+ "grad_checkpoint": false,
51
+ "budget_locked": true,
52
+ "log_every": 20,
53
+ "eval_every": 150,
54
+ "val_batches": 32,
55
+ "ckpt_every": 1000,
56
+ "best_metric": "val_ce",
57
+ "seed": 1337
58
+ },
59
+ "K": 1,
60
+ "params": 9962496,
61
+ "training_tokens": 60014592
62
+ }
eval_results.json ADDED
@@ -0,0 +1,181 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "ckpt": "/kaggle/working/out/stack_k1_60m/best.pt",
3
+ "arch": "stack",
4
+ "K_train": 1,
5
+ "k_max": 12,
6
+ "val_tokens": 1539584,
7
+ "params": 9962496,
8
+ "uniform_ce": 9.704060527839234,
9
+ "unigram_ce": 7.547556945473469,
10
+ "best_k": 1,
11
+ "gain_first_loop": 1.012203666907749,
12
+ "gain_extra_loops": 0.0,
13
+ "mean_rel_step": 0.28043451470633346,
14
+ "verdict": "TRAINED UNLOOPED (K=1): the k>1 rows show how a model trained at one application behaves when looped anyway, not whether looping pays.",
15
+ "rows": [
16
+ {
17
+ "k": 0,
18
+ "ce": 5.408761565218267,
19
+ "ppl": 223.35480596749395,
20
+ "bpb": 1.8913016483504717,
21
+ "delta_ce": null,
22
+ "flops_per_token": 18259968
23
+ },
24
+ {
25
+ "k": 1,
26
+ "ce": 4.396557898310518,
27
+ "ppl": 81.17098845673196,
28
+ "bpb": 1.5373606508401174,
29
+ "delta_ce": 1.012203666907749,
30
+ "flops_per_token": 21098496,
31
+ "std_ratio": 0.9512100219726562,
32
+ "erank": 172.63084411621094,
33
+ "step_cos": 0.7777038216590881,
34
+ "rel_step": 0.8602083921432495,
35
+ "state_norm": 15.174788475036621
36
+ },
37
+ {
38
+ "k": 2,
39
+ "ce": 4.5994899281640365,
40
+ "ppl": 99.4335844337849,
41
+ "bpb": 1.6083206437954694,
42
+ "delta_ce": -0.20293202985351844,
43
+ "flops_per_token": 23937024,
44
+ "std_ratio": 0.9564117789268494,
45
+ "erank": 152.0776138305664,
46
+ "step_cos": 0.9423013627529144,
47
+ "rel_step": 0.5801940560340881,
48
+ "state_norm": 21.41031837463379
49
+ },
50
+ {
51
+ "k": 3,
52
+ "ce": 4.941250652976713,
53
+ "ppl": 139.94516299387544,
54
+ "bpb": 1.727825379655199,
55
+ "delta_ce": -0.3417607248126764,
56
+ "flops_per_token": 26775552,
57
+ "std_ratio": 0.9550213515758514,
58
+ "erank": 126.67342376708984,
59
+ "step_cos": 0.9775733947753906,
60
+ "rel_step": 0.3683442771434784,
61
+ "state_norm": 27.632020950317383
62
+ },
63
+ {
64
+ "k": 4,
65
+ "ce": 5.265907488037353,
66
+ "ppl": 193.6219386774754,
67
+ "bpb": 1.8413493351660124,
68
+ "delta_ce": -0.32465683506064025,
69
+ "flops_per_token": 29614080,
70
+ "std_ratio": 0.9511829912662506,
71
+ "erank": 109.17454528808594,
72
+ "step_cos": 0.9875146746635437,
73
+ "rel_step": 0.27501752972602844,
74
+ "state_norm": 33.76785087585449
75
+ },
76
+ {
77
+ "k": 5,
78
+ "ce": 5.547618621125902,
79
+ "ppl": 256.625704638506,
80
+ "bpb": 1.9398563083325306,
81
+ "delta_ce": -0.28171113308854867,
82
+ "flops_per_token": 32452608,
83
+ "std_ratio": 0.945938229560852,
84
+ "erank": 88.04708099365234,
85
+ "step_cos": 0.9916830360889435,
86
+ "rel_step": 0.2244347706437111,
87
+ "state_norm": 39.811710357666016
88
+ },
89
+ {
90
+ "k": 6,
91
+ "ce": 5.789976763812497,
92
+ "ppl": 327.00542592830055,
93
+ "bpb": 2.024602575888849,
94
+ "delta_ce": -0.24235814268659528,
95
+ "flops_per_token": 35291136,
96
+ "std_ratio": 0.9405999779701233,
97
+ "erank": 86.94815063476562,
98
+ "step_cos": 0.9938283264636993,
99
+ "rel_step": 0.19342376291751862,
100
+ "state_norm": 45.64400100708008
101
+ },
102
+ {
103
+ "k": 7,
104
+ "ce": 5.99996355549616,
105
+ "ppl": 403.4140909984358,
106
+ "bpb": 2.0980294334891267,
107
+ "delta_ce": -0.2099867916836633,
108
+ "flops_per_token": 38129664,
109
+ "std_ratio": 0.9329713881015778,
110
+ "erank": 64.25496482849121,
111
+ "step_cos": 0.9951082766056061,
112
+ "rel_step": 0.1726725548505783,
113
+ "state_norm": 53.134456634521484
114
+ },
115
+ {
116
+ "k": 8,
117
+ "ce": 6.18351374157727,
118
+ "ppl": 484.6920503676553,
119
+ "bpb": 2.1622121054934955,
120
+ "delta_ce": -0.18355018608110996,
121
+ "flops_per_token": 40968192,
122
+ "std_ratio": 0.9266401827335358,
123
+ "erank": 56.75149154663086,
124
+ "step_cos": 0.9959644377231598,
125
+ "rel_step": 0.157696433365345,
126
+ "state_norm": 59.62197303771973
127
+ },
128
+ {
129
+ "k": 9,
130
+ "ce": 6.345236095432273,
131
+ "ppl": 569.7718943785347,
132
+ "bpb": 2.218762158723424,
133
+ "delta_ce": -0.16172235385500233,
134
+ "flops_per_token": 43806720,
135
+ "std_ratio": 0.9184004366397858,
136
+ "erank": 48.32929611206055,
137
+ "step_cos": 0.9965897798538208,
138
+ "rel_step": 0.14616148173809052,
139
+ "state_norm": 67.28807830810547
140
+ },
141
+ {
142
+ "k": 10,
143
+ "ce": 6.488679977838169,
144
+ "ppl": 657.6546712541515,
145
+ "bpb": 2.268920711280938,
146
+ "delta_ce": -0.14344388240589634,
147
+ "flops_per_token": 46645248,
148
+ "std_ratio": 0.9087829887866974,
149
+ "erank": 38.35821723937988,
150
+ "step_cos": 0.9970792233943939,
151
+ "rel_step": 0.13675487786531448,
152
+ "state_norm": 75.25280380249023
153
+ },
154
+ {
155
+ "k": 11,
156
+ "ce": 6.616616099706154,
157
+ "ppl": 747.4116465669175,
158
+ "bpb": 2.3136566078914447,
159
+ "delta_ce": -0.1279361218679851,
160
+ "flops_per_token": 49483776,
161
+ "std_ratio": 0.8985040187835693,
162
+ "erank": 32.420021057128906,
163
+ "step_cos": 0.9974798560142517,
164
+ "rel_step": 0.12871047109365463,
165
+ "state_norm": 83.94036483764648
166
+ },
167
+ {
168
+ "k": 12,
169
+ "ce": 6.731254219889919,
170
+ "ppl": 838.1978914258499,
171
+ "bpb": 2.3537425431010153,
172
+ "delta_ce": -0.11463812018376451,
173
+ "flops_per_token": 52322304,
174
+ "std_ratio": 0.8866880834102631,
175
+ "erank": 26.89052963256836,
176
+ "step_cos": 0.9978146255016327,
177
+ "rel_step": 0.12159556895494461,
178
+ "state_norm": 93.48854446411133
179
+ }
180
+ ]
181
+ }
load_example.py ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Minimal loader. Requires: torch, tokenizers."""
2
+ import torch
3
+ from tokenizers import Tokenizer
4
+ from simloop.stack import StackConfig, StackedLoop
5
+
6
+ ck = torch.load("model.pt", map_location="cpu", weights_only=False)
7
+ model = StackedLoop(StackConfig(**ck["model_cfg"]))
8
+ model.load_state_dict(ck["model"])
9
+ model.eval()
10
+
11
+ tok = Tokenizer.from_file("tokenizer.json")
12
+ ids = torch.tensor([tok.encode("The capital of France is").ids])
13
+ with torch.no_grad():
14
+ logits = model(ids, K=1)
15
+ print(tok.decode([int(logits[0, -1].argmax())]))
model.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e4ee0d100eb055908914969dc5f291a500eb073e8c30a41ee632f49316ac4c2a
3
+ size 39862696
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0e1ecc44d47c74e4a057f099f3095b18e2236665637563710345728b30e98086
3
+ size 39853464
run_config.json ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "arch": "stack",
3
+ "model": {
4
+ "vocab_size": 16384,
5
+ "d_model": 384,
6
+ "n_pre": 1,
7
+ "n_loop": 1,
8
+ "n_post": 1,
9
+ "n_heads": 6,
10
+ "n_kv_heads": 2,
11
+ "head_dim": 64,
12
+ "d_mlp": 720,
13
+ "max_seq_len": 512,
14
+ "rope_theta": 10000.0,
15
+ "dropout": 0.0,
16
+ "norm_eps": 1e-06,
17
+ "init_std": 0.02,
18
+ "scale_resid_init": true
19
+ },
20
+ "train": {
21
+ "K": 1,
22
+ "k_schedule": "const",
23
+ "k_end": 1,
24
+ "k_decay_start_frac": 0.5,
25
+ "k_decay_end_frac": 0.9,
26
+ "k_min": 1,
27
+ "drop_eps": 0.01,
28
+ "eval_k_max": 8,
29
+ "batch_size": 16,
30
+ "grad_accum": 2,
31
+ "seq_len": 512,
32
+ "max_tokens": 60000000,
33
+ "lr": 0.001,
34
+ "min_lr_ratio": 0.1,
35
+ "warmup_steps": 100,
36
+ "weight_decay": 0.1,
37
+ "beta1": 0.9,
38
+ "beta2": 0.95,
39
+ "grad_clip": 1.0,
40
+ "distill_weight": 0.0,
41
+ "distill_space": "hidden",
42
+ "distill_depths": "first",
43
+ "distill_start_frac": 0.0,
44
+ "distill_tau": 1.0,
45
+ "anytime_ce_weight": 0.0,
46
+ "grad_checkpoint": false,
47
+ "log_every": 20,
48
+ "eval_every": 150,
49
+ "val_batches": 32,
50
+ "ckpt_every": 1000,
51
+ "best_metric": "val_ce",
52
+ "seed": 1337,
53
+ "budget_locked": true
54
+ }
55
+ }
simloop/__init__.py ADDED
File without changes
simloop/model.py ADDED
@@ -0,0 +1,205 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SimLoop model: Qwen3-style looped transformer.
2
+
3
+ A single block of cfg.n_layers transformer layers (pre-RMSNorm, RoPE, SwiGLU,
4
+ GQA, QK-norm) is applied K times to the embedded input. Input/output
5
+ embeddings are tied. Dropout inside the block doubles as the stochastic
6
+ augmentation for the siamese consistency objective: two forward passes of the
7
+ same input in train mode give two "views" (SimCSE-style).
8
+
9
+ Loop states are exposed via forward_loops(..., collect=True) so a future
10
+ early-exit head can read intermediate states.
11
+ """
12
+
13
+ from dataclasses import dataclass, replace
14
+
15
+ import torch
16
+ import torch.nn as nn
17
+ import torch.nn.functional as F
18
+ from torch.utils.checkpoint import checkpoint
19
+
20
+
21
+ @dataclass
22
+ class SimLoopConfig:
23
+ vocab_size: int = 16384
24
+ d_model: int = 384
25
+ n_layers: int = 2 # layers inside the looped block
26
+ n_heads: int = 6
27
+ n_kv_heads: int = 2 # GQA
28
+ head_dim: int = 64
29
+ d_mlp: int = 1024
30
+ max_seq_len: int = 512
31
+ rope_theta: float = 10000.0
32
+ dropout: float = 0.1 # the augmentation for the siamese objective
33
+ # Where the augmentation noise enters:
34
+ # "input" - dropout applied ONCE to the embedding, per branch; the
35
+ # loop itself is deterministic. Branch divergence is then a
36
+ # fixed input perturbation, so consistency across depth
37
+ # measures whether the loop CONTRACTS it.
38
+ # "loop" - standard transformer dropout inside every layer, resampled
39
+ # at every loop application. Divergence accumulates with
40
+ # depth and contraction cannot be measured.
41
+ aug_mode: str = "loop"
42
+ norm_eps: float = 1e-6
43
+ predictor_hidden: int = 0 # >0: SimSiam-style MLP predictor (train-only)
44
+
45
+
46
+ class RMSNorm(nn.Module):
47
+ def __init__(self, dim, eps=1e-6):
48
+ super().__init__()
49
+ self.weight = nn.Parameter(torch.ones(dim))
50
+ self.eps = eps
51
+
52
+ def forward(self, x):
53
+ dtype = x.dtype
54
+ x = x.float()
55
+ x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
56
+ return (x * self.weight.float()).to(dtype)
57
+
58
+
59
+ def _rope_tables(head_dim, max_len, theta):
60
+ inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
61
+ t = torch.arange(max_len, dtype=torch.float32)
62
+ freqs = torch.outer(t, inv_freq) # [T, head_dim/2]
63
+ return freqs.cos(), freqs.sin()
64
+
65
+
66
+ def _apply_rope(x, cos, sin):
67
+ # x: [B, H, T, hd]; cos/sin: [max_T, hd/2]
68
+ T = x.shape[-2]
69
+ cos = cos[:T].view(1, 1, T, -1)
70
+ sin = sin[:T].view(1, 1, T, -1)
71
+ x1, x2 = x[..., 0::2], x[..., 1::2]
72
+ out = torch.stack((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1)
73
+ return out.flatten(-2)
74
+
75
+
76
+ class Attention(nn.Module):
77
+ def __init__(self, cfg):
78
+ super().__init__()
79
+ self.n_heads, self.n_kv, self.hd = cfg.n_heads, cfg.n_kv_heads, cfg.head_dim
80
+ self.q_proj = nn.Linear(cfg.d_model, cfg.n_heads * cfg.head_dim, bias=False)
81
+ self.k_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False)
82
+ self.v_proj = nn.Linear(cfg.d_model, cfg.n_kv_heads * cfg.head_dim, bias=False)
83
+ self.o_proj = nn.Linear(cfg.n_heads * cfg.head_dim, cfg.d_model, bias=False)
84
+ self.q_norm = RMSNorm(cfg.head_dim, cfg.norm_eps) # Qwen3 QK-norm
85
+ self.k_norm = RMSNorm(cfg.head_dim, cfg.norm_eps)
86
+ self.dropout = cfg.dropout
87
+
88
+ def forward(self, x, cos, sin):
89
+ B, T, _ = x.shape
90
+ q = self.q_norm(self.q_proj(x).view(B, T, self.n_heads, self.hd))
91
+ k = self.k_norm(self.k_proj(x).view(B, T, self.n_kv, self.hd))
92
+ v = self.v_proj(x).view(B, T, self.n_kv, self.hd)
93
+ q, k, v = (t.transpose(1, 2) for t in (q, k, v)) # [B, H, T, hd]
94
+ q, k = _apply_rope(q, cos, sin), _apply_rope(k, cos, sin)
95
+ rep = self.n_heads // self.n_kv
96
+ k, v = k.repeat_interleave(rep, dim=1), v.repeat_interleave(rep, dim=1)
97
+ o = F.scaled_dot_product_attention(
98
+ q, k, v, is_causal=True,
99
+ dropout_p=self.dropout if self.training else 0.0)
100
+ return self.o_proj(o.transpose(1, 2).reshape(B, T, -1))
101
+
102
+
103
+ class SwiGLU(nn.Module):
104
+ def __init__(self, cfg):
105
+ super().__init__()
106
+ self.gate = nn.Linear(cfg.d_model, cfg.d_mlp, bias=False)
107
+ self.up = nn.Linear(cfg.d_model, cfg.d_mlp, bias=False)
108
+ self.down = nn.Linear(cfg.d_mlp, cfg.d_model, bias=False)
109
+
110
+ def forward(self, x):
111
+ return self.down(F.silu(self.gate(x)) * self.up(x))
112
+
113
+
114
+ class Layer(nn.Module):
115
+ def __init__(self, cfg):
116
+ super().__init__()
117
+ self.ln_attn = RMSNorm(cfg.d_model, cfg.norm_eps)
118
+ self.attn = Attention(cfg)
119
+ self.ln_mlp = RMSNorm(cfg.d_model, cfg.norm_eps)
120
+ self.mlp = SwiGLU(cfg)
121
+ self.drop = nn.Dropout(cfg.dropout)
122
+
123
+ def forward(self, x, cos, sin):
124
+ x = x + self.drop(self.attn(self.ln_attn(x), cos, sin))
125
+ x = x + self.drop(self.mlp(self.ln_mlp(x)))
126
+ return x
127
+
128
+
129
+ class PredictorMLP(nn.Module):
130
+ """SimSiam-style bottleneck predictor h, applied to post-norm states.
131
+ Training-time only — discarded at inference, exactly like SimSiam's h.
132
+ Restores the predictor asymmetry the extra-loop predictor lacks (g=f)."""
133
+
134
+ def __init__(self, d_model, hidden):
135
+ super().__init__()
136
+ self.fc1 = nn.Linear(d_model, hidden)
137
+ self.ln = nn.LayerNorm(hidden)
138
+ self.fc2 = nn.Linear(hidden, d_model)
139
+
140
+ def forward(self, x):
141
+ return self.fc2(F.gelu(self.ln(self.fc1(x))))
142
+
143
+
144
+ class SimLoop(nn.Module):
145
+ def __init__(self, cfg: SimLoopConfig):
146
+ super().__init__()
147
+ self.cfg = cfg
148
+ self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model)
149
+ # input mode: noise lives at the embedding, the loop is deterministic
150
+ loop_cfg = replace(cfg, dropout=0.0) if cfg.aug_mode == "input" else cfg
151
+ self.in_drop = nn.Dropout(cfg.dropout if cfg.aug_mode == "input" else 0.0)
152
+ self.layers = nn.ModuleList(Layer(loop_cfg) for _ in range(cfg.n_layers))
153
+ self.norm_out = RMSNorm(cfg.d_model, cfg.norm_eps)
154
+ self.predictor = (PredictorMLP(cfg.d_model, cfg.predictor_hidden)
155
+ if cfg.predictor_hidden > 0 else None)
156
+ cos, sin = _rope_tables(cfg.head_dim, cfg.max_seq_len, cfg.rope_theta)
157
+ self.register_buffer("rope_cos", cos, persistent=False)
158
+ self.register_buffer("rope_sin", sin, persistent=False)
159
+ self.apply(self._init)
160
+
161
+ @staticmethod
162
+ def _init(m):
163
+ if isinstance(m, (nn.Linear, nn.Embedding)):
164
+ nn.init.normal_(m.weight, std=0.02)
165
+
166
+ def embed_aug(self, x):
167
+ """Embedded input with the branch augmentation applied. Call once PER
168
+ BRANCH so each draws its own mask. A no-op unless aug_mode="input",
169
+ so aug_mode="loop" behaves exactly as before."""
170
+ return self.in_drop(self.embed(x))
171
+
172
+ def step(self, z):
173
+ """One loop iteration: the whole block applied once."""
174
+ for layer in self.layers:
175
+ z = layer(z, self.rope_cos, self.rope_sin)
176
+ return z
177
+
178
+ def forward_loops(self, z0, n_steps, grad_checkpoint=False, collect=False):
179
+ """Apply the block n_steps times.
180
+
181
+ Returns (z_final, states); states = [z0, z1, ..., z_n] if collect,
182
+ else None. Checkpointing preserves RNG state, so recomputed dropout
183
+ masks match the original forward.
184
+ """
185
+ states = [z0] if collect else None
186
+ z = z0
187
+ for _ in range(n_steps):
188
+ if grad_checkpoint and self.training and torch.is_grad_enabled():
189
+ z = checkpoint(self.step, z, use_reentrant=False)
190
+ else:
191
+ z = self.step(z)
192
+ if collect:
193
+ states.append(z)
194
+ return z, states
195
+
196
+ def post(self, z):
197
+ """Final RMSNorm — the space where consistency is matched (per token)
198
+ and where the LM head reads."""
199
+ return self.norm_out(z)
200
+
201
+ def logits(self, z):
202
+ return F.linear(self.post(z), self.embed.weight) # tied head
203
+
204
+ def num_params(self):
205
+ return sum(p.numel() for p in self.parameters())
simloop/stack.py ADDED
@@ -0,0 +1,158 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SimLoop-Stack (E9+): an ordinary Qwen3-style transformer stack in which
2
+ ONE middle block is applied K times (weight-tied); every other layer is a
3
+ plain unlooped layer.
4
+
5
+ embed -> [pre layers] -> ( loop block ) x K -> [post layers] -> norm -> tied head
6
+
7
+ This is deliberately NOT the phase-1 SimLoop architecture (`simloop/model.py`,
8
+ where every layer lives inside the loop and there is no readout stack). It is
9
+ kept as a separate class so every phase-1 result stays bit-reproducible.
10
+
11
+ The point of the arrangement is that K is the ONLY thing that differs between
12
+ the baseline and the looped run: same parameters, same initial values (init
13
+ draws do not depend on K), same data order. So any difference in the learning
14
+ curve is attributable to the extra loop applications and nothing else.
15
+
16
+ Two properties the experiment depends on:
17
+
18
+ * k = 0 is a real forward pass. `enter()` -> `readout()` skips the looped
19
+ block entirely, so "the loops stopped mattering" can be told apart from
20
+ "the loop learned to be the identity" (the trivial minimiser of any
21
+ deep-to-shallow distillation term).
22
+ * Weight tying means one checkpoint can be evaluated at any k, including k
23
+ larger than the K it was trained at.
24
+ """
25
+
26
+ import math
27
+ from dataclasses import dataclass
28
+
29
+ import torch
30
+ import torch.nn as nn
31
+ import torch.nn.functional as F
32
+ from torch.utils.checkpoint import checkpoint
33
+
34
+ from simloop.model import RMSNorm, Layer, _rope_tables
35
+
36
+
37
+ @dataclass
38
+ class StackConfig:
39
+ vocab_size: int = 16384
40
+ d_model: int = 384
41
+ n_pre: int = 1 # unlooped layers before the loop
42
+ n_loop: int = 1 # layers inside the looped (weight-tied) block
43
+ n_post: int = 1 # unlooped layers after the loop
44
+ n_heads: int = 6
45
+ n_kv_heads: int = 2 # GQA
46
+ head_dim: int = 64
47
+ d_mlp: int = 720 # 1.875x d_model: the largest that fits 3 blocks
48
+ # + a tied 16k embedding under the 10M cap
49
+ max_seq_len: int = 512
50
+ rope_theta: float = 10000.0
51
+ dropout: float = 0.0 # 0: <1 epoch of FineWeb, and it would confound K
52
+ norm_eps: float = 1e-6
53
+ init_std: float = 0.02
54
+ # Scale residual-branch output init by 1/sqrt(2*depth). depth counts
55
+ # PARAMETER blocks, not loop applications, on purpose: making it depend
56
+ # on K would give the baseline and the looped run different initial
57
+ # weights and destroy the controlled comparison.
58
+ scale_resid_init: bool = True
59
+
60
+
61
+ class StackedLoop(nn.Module):
62
+ def __init__(self, cfg: StackConfig):
63
+ super().__init__()
64
+ self.cfg = cfg
65
+ self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model)
66
+ self.pre_layers = nn.ModuleList(Layer(cfg) for _ in range(cfg.n_pre))
67
+ self.loop_layers = nn.ModuleList(Layer(cfg) for _ in range(cfg.n_loop))
68
+ self.post_layers = nn.ModuleList(Layer(cfg) for _ in range(cfg.n_post))
69
+ self.norm_out = RMSNorm(cfg.d_model, cfg.norm_eps)
70
+ cos, sin = _rope_tables(cfg.head_dim, cfg.max_seq_len, cfg.rope_theta)
71
+ self.register_buffer("rope_cos", cos, persistent=False)
72
+ self.register_buffer("rope_sin", sin, persistent=False)
73
+ self._init_weights()
74
+
75
+ # ---------------------------------------------------------------- init
76
+ def _init_weights(self):
77
+ for m in self.modules():
78
+ if isinstance(m, (nn.Linear, nn.Embedding)):
79
+ nn.init.normal_(m.weight, std=self.cfg.init_std)
80
+ if not self.cfg.scale_resid_init:
81
+ return
82
+ depth = self.cfg.n_pre + self.cfg.n_loop + self.cfg.n_post
83
+ std = self.cfg.init_std / math.sqrt(2.0 * depth)
84
+ for layer in self.all_layers():
85
+ nn.init.normal_(layer.attn.o_proj.weight, std=std)
86
+ nn.init.normal_(layer.mlp.down.weight, std=std)
87
+
88
+ def all_layers(self):
89
+ return list(self.pre_layers) + list(self.loop_layers) + list(self.post_layers)
90
+
91
+ # ------------------------------------------------------------- forward
92
+ def enter(self, x):
93
+ """Token ids -> the state that enters the looped block (k = 0)."""
94
+ h = self.embed(x)
95
+ for layer in self.pre_layers:
96
+ h = layer(h, self.rope_cos, self.rope_sin)
97
+ return h
98
+
99
+ def loop_once(self, h):
100
+ """One application of the looped block."""
101
+ for layer in self.loop_layers:
102
+ h = layer(h, self.rope_cos, self.rope_sin)
103
+ return h
104
+
105
+ def readout(self, h):
106
+ """Loop-exit state -> final normalised hidden state."""
107
+ for layer in self.post_layers:
108
+ h = layer(h, self.rope_cos, self.rope_sin)
109
+ return self.norm_out(h)
110
+
111
+ def logits_at(self, h):
112
+ """Logits from a loop-exit state at ANY k (tied head)."""
113
+ return F.linear(self.readout(h), self.embed.weight)
114
+
115
+ def run(self, x, K, collect=True, grad_checkpoint=False):
116
+ """Returns (h_K, states) with states = [h_0, h_1, ..., h_K], h_k the
117
+ loop-exit state after k applications. Collecting is free: autograd
118
+ already holds these tensors."""
119
+ h = self.enter(x)
120
+ states = [h] if collect else None
121
+ for _ in range(K):
122
+ if grad_checkpoint and self.training and torch.is_grad_enabled():
123
+ h = checkpoint(self.loop_once, h, use_reentrant=False)
124
+ else:
125
+ h = self.loop_once(h)
126
+ if collect:
127
+ states.append(h)
128
+ return h, states
129
+
130
+ def forward(self, x, K):
131
+ h, _ = self.run(x, K, collect=False)
132
+ return self.logits_at(h)
133
+
134
+ # --------------------------------------------------------------- meta
135
+ def num_params(self):
136
+ return sum(p.numel() for p in self.parameters())
137
+
138
+ def flops_per_token(self, K, seq_len=None, backward=True):
139
+ """Analytic FLOPs/token. Matmuls counted as 2*params (multiply+add);
140
+ attention scores counted causally (half the full T x T grid).
141
+ backward=True applies the usual 3x (fwd + bwd ~ 2x fwd)."""
142
+ c = self.cfg
143
+ seq_len = seq_len or c.max_seq_len
144
+ qo = 2 * c.d_model * c.n_heads * c.head_dim # q_proj + o_proj
145
+ kv = 2 * c.d_model * c.n_kv_heads * c.head_dim # k_proj + v_proj
146
+ mlp = 3 * c.d_model * c.d_mlp # gate + up + down
147
+ per_layer = 2 * (qo + kv + mlp) + 2 * seq_len * c.n_heads * c.head_dim
148
+ n_applied = c.n_pre + K * c.n_loop + c.n_post
149
+ head = 2 * c.d_model * c.vocab_size
150
+ f = n_applied * per_layer + head
151
+ return f * (3 if backward else 1)
152
+
153
+ def describe(self, K):
154
+ c = self.cfg
155
+ return (f"{c.n_pre}pre + {c.n_loop}loop x K={K} + {c.n_post}post "
156
+ f"= {c.n_pre + K * c.n_loop + c.n_post} layer applications, "
157
+ f"{c.n_pre + c.n_loop + c.n_post} parameter blocks, "
158
+ f"d_model={c.d_model} d_mlp={c.d_mlp} vocab={c.vocab_size}")
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff