miguelcsx commited on
Commit
fe09835
·
verified ·
1 Parent(s): b4d8ba8

select chck_100M

Browse files
README.md ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ language:
3
+ - en
4
+ license: other
5
+ library_name: transformers
6
+ pipeline_tag: fill-mask
7
+ tags:
8
+ - babylm
9
+ - strict-small
10
+ - research-release
11
+ ---
12
+
13
+ # Structured Direct-Sum syntax plus lexical
14
+
15
+ This repository preserves an already-trained checkpoint from the controlled
16
+ BabyLM research tournament. No training or evaluation was run for this release.
17
+
18
+ ## Selected revision
19
+
20
+ `main` is identical to `chck_100M`, selected from the existing local tournament
21
+ record. Other revisions, when present, are archived checkpoints rather than new
22
+ experiments.
23
+
24
+ ## Existing evaluation record
25
+
26
+ | Evaluation | Score |
27
+ |---|---:|
28
+ | BLiMP | 65.51 |
29
+ | Supp | 59.60 |
30
+ | EWoK | 49.09 |
31
+ | ET | 19.23 |
32
+ | COMPS | 52.15 |
33
+ | GlobalPIQA | 39.58 |
34
+
35
+ The canonical code is maintained at
36
+ [miguelcsx/tolm](https://github.com/miguelcsx/tolm). Remote code is required to
37
+ load this custom Transformers model; review `tolm.py` before use.
config.json ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_type": "tolm",
3
+ "architectures": [
4
+ "TOLMForMaskedLM"
5
+ ],
6
+ "auto_map": {
7
+ "AutoConfig": "tolm.TOLMConfig",
8
+ "AutoModel": "tolm.TOLMModel",
9
+ "AutoModelForMaskedLM": "tolm.TOLMForMaskedLM",
10
+ "AutoModelForCausalLM": "tolm.TOLMForCausalLM"
11
+ },
12
+ "vocab_size": 16000,
13
+ "max_seq_len": 512,
14
+ "hidden_size": 512,
15
+ "num_hidden_layers": 12,
16
+ "num_attention_heads": 8,
17
+ "intermediate_size": 1536,
18
+ "position_buckets": 32,
19
+ "dropout": 0.1,
20
+ "attention_dropout": 0.1,
21
+ "initializer_range": 0.03227486121839514,
22
+ "value_gating": true,
23
+ "residual_mixing": true,
24
+ "pad_token_id": 1,
25
+ "bos_token_id": 2,
26
+ "eos_token_id": 3,
27
+ "mask_token_id": 4,
28
+ "absolute_positions": false,
29
+ "use_rope": false,
30
+ "use_alibi": false,
31
+ "recurrent_steps": 1,
32
+ "num_experts": 1,
33
+ "experts_per_token": 1,
34
+ "expert_intermediate_size": null,
35
+ "future_offsets": [],
36
+ "state_mixer_kernel": 0,
37
+ "rtd_auxiliary": false,
38
+ "geometry_lexical_dim": 0,
39
+ "geometry_curvature": 1.0,
40
+ "cognitive_readout_layer": 0,
41
+ "cognitive_readout_weight": 0.0,
42
+ "direct_sum_dims": [
43
+ 96,
44
+ 128,
45
+ 288
46
+ ],
47
+ "direct_sum_heads": [
48
+ 3,
49
+ 4,
50
+ 6
51
+ ],
52
+ "direct_sum_intermediate_sizes": [
53
+ 512,
54
+ 640,
55
+ 1344
56
+ ]
57
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:441db679c7fb37b7894281d9768ac0f13f36d916772c91d7173749ed72cdfed0
3
+ size 136454508
release_manifest.json ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "checkpoint_revisions": [
3
+ "chck_100M"
4
+ ],
5
+ "checkpoints": {
6
+ "chck_100M": {
7
+ "config.json": "f2255ff25207a54f8e5925f874032e6be49e6576d9796f24006e7f87a3a527d1",
8
+ "model.safetensors": "441db679c7fb37b7894281d9768ac0f13f36d916772c91d7173749ed72cdfed0",
9
+ "special_tokens_map.json": "dbd124d177a0bae11b6740c061b6b9e4772887f9f068867381691db10c88a3c2",
10
+ "tokenizer.json": "abfd627e173d48addf924e32fd6c38a3d485734c496d12916e559e654ba15c51",
11
+ "tokenizer_config.json": "8f4a998564d87fc85b27652f0b295677a1e03744bf400bfacc342b700cc3653a",
12
+ "tolm.py": "c10e51279e01cd758e860194455a953ad9698c5e225da4a1a8af5dc432207d55",
13
+ "training_manifest.json": "ae8f0213c749cf9046e3ab3cde292a517a3f7fee9b745c64224e302f63c97379"
14
+ }
15
+ },
16
+ "schema_version": 1,
17
+ "selected_revision": "chck_100M",
18
+ "selection_scores": {
19
+ "BLiMP": 65.51,
20
+ "COMPS": 52.15,
21
+ "ET": 19.23,
22
+ "EWoK": 49.09,
23
+ "GlobalPIQA": 39.58,
24
+ "Supp": 59.6
25
+ },
26
+ "source_root": "../tolm_structured/runs/structured_ds_syntax_lexical",
27
+ "variant": "structured_ds_syntax_lexical"
28
+ }
special_tokens_map.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token": "<s>",
3
+ "eos_token": "</s>",
4
+ "unk_token": "<unk>",
5
+ "pad_token": "<pad>",
6
+ "mask_token": "<mask>"
7
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token": "<s>",
3
+ "eos_token": "</s>",
4
+ "mask_token": "<mask>",
5
+ "model_max_length": 512,
6
+ "pad_token": "<pad>",
7
+ "tokenizer_class": "PreTrainedTokenizerFast",
8
+ "unk_token": "<unk>"
9
+ }
tolm.py ADDED
@@ -0,0 +1,842 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import copy
2
+ import math
3
+
4
+ import torch
5
+ from torch import nn
6
+ from torch.nn import functional as F
7
+ from transformers import PretrainedConfig, PreTrainedModel
8
+ from transformers.modeling_outputs import (
9
+ BaseModelOutput,
10
+ CausalLMOutput,
11
+ MaskedLMOutput,
12
+ )
13
+
14
+
15
+ class TOLMConfig(PretrainedConfig):
16
+ model_type = "tolm"
17
+
18
+ def __init__(
19
+ self,
20
+ vocab_size=16000,
21
+ max_seq_len=512,
22
+ max_position_embeddings=None,
23
+ hidden_size=256,
24
+ num_hidden_layers=4,
25
+ num_attention_heads=4,
26
+ intermediate_size=1024,
27
+ position_buckets=32,
28
+ dropout=0.1,
29
+ hidden_dropout_prob=None,
30
+ attention_dropout=0.1,
31
+ attention_probs_dropout_prob=None,
32
+ initializer_range=0.03952847075210474,
33
+ layer_norm_eps=1.0e-5,
34
+ lm_head_gelu_approximate="tanh",
35
+ shared_relative_embeddings=False,
36
+ feedforward_dropout_after_projection=False,
37
+ attention_output_dropout=False,
38
+ embedding_padding_idx=True,
39
+ value_gating=True,
40
+ residual_mixing=True,
41
+ pad_token_id=1,
42
+ bos_token_id=2,
43
+ eos_token_id=3,
44
+ mask_token_id=4,
45
+ absolute_positions=False,
46
+ use_rope=False,
47
+ use_alibi=False,
48
+ recurrent_steps=1,
49
+ num_experts=1,
50
+ experts_per_token=1,
51
+ expert_intermediate_size=None,
52
+ future_offsets=None,
53
+ state_mixer_kernel=0,
54
+ geometry_lexical_dim=0,
55
+ geometry_curvature=1.0,
56
+ cognitive_readout_layer=0,
57
+ cognitive_readout_weight=0.0,
58
+ direct_sum_dims=None,
59
+ direct_sum_heads=None,
60
+ direct_sum_intermediate_sizes=None,
61
+ **kwargs,
62
+ ):
63
+ super().__init__(
64
+ pad_token_id=pad_token_id,
65
+ bos_token_id=bos_token_id,
66
+ eos_token_id=eos_token_id,
67
+ mask_token_id=mask_token_id,
68
+ **kwargs,
69
+ )
70
+ self.vocab_size = vocab_size
71
+ self.max_seq_len = max_position_embeddings or max_seq_len
72
+ self.max_position_embeddings = self.max_seq_len
73
+ self.hidden_size = hidden_size
74
+ self.num_hidden_layers = num_hidden_layers
75
+ self.num_attention_heads = num_attention_heads
76
+ self.intermediate_size = intermediate_size
77
+ self.position_buckets = position_buckets
78
+ self.dropout = (
79
+ hidden_dropout_prob if hidden_dropout_prob is not None else dropout
80
+ )
81
+ self.hidden_dropout_prob = self.dropout
82
+ self.attention_dropout = (
83
+ attention_probs_dropout_prob
84
+ if attention_probs_dropout_prob is not None
85
+ else attention_dropout
86
+ )
87
+ self.attention_probs_dropout_prob = self.attention_dropout
88
+ self.initializer_range = initializer_range
89
+ self.layer_norm_eps = layer_norm_eps
90
+ self.lm_head_gelu_approximate = lm_head_gelu_approximate
91
+ self.shared_relative_embeddings = shared_relative_embeddings
92
+ self.feedforward_dropout_after_projection = (
93
+ feedforward_dropout_after_projection
94
+ )
95
+ self.attention_output_dropout = attention_output_dropout
96
+ self.embedding_padding_idx = embedding_padding_idx
97
+ self.value_gating = value_gating
98
+ self.residual_mixing = residual_mixing
99
+ self.absolute_positions = absolute_positions
100
+ self.use_rope = use_rope
101
+ self.use_alibi = use_alibi
102
+ self.recurrent_steps = recurrent_steps
103
+ self.num_experts = num_experts
104
+ self.experts_per_token = experts_per_token
105
+ self.expert_intermediate_size = expert_intermediate_size
106
+ self.future_offsets = future_offsets or []
107
+ self.state_mixer_kernel = state_mixer_kernel
108
+ self.geometry_lexical_dim = geometry_lexical_dim
109
+ self.geometry_curvature = geometry_curvature
110
+ self.cognitive_readout_layer = cognitive_readout_layer
111
+ self.cognitive_readout_weight = cognitive_readout_weight
112
+ self.direct_sum_dims = direct_sum_dims or []
113
+ self.direct_sum_heads = direct_sum_heads or []
114
+ self.direct_sum_intermediate_sizes = direct_sum_intermediate_sizes or []
115
+
116
+
117
+ def _valid_tokens(input_ids, attention_mask):
118
+ if attention_mask is None:
119
+ return torch.ones_like(input_ids, dtype=torch.bool)
120
+ return attention_mask.to(torch.bool)
121
+
122
+
123
+ def _bidirectional_mask(valid):
124
+ return valid[:, None, None, :] & valid[:, None, :, None]
125
+
126
+
127
+ def _causal_mask(valid):
128
+ length = valid.size(1)
129
+ causal = torch.ones((length, length), dtype=torch.bool, device=valid.device).tril()
130
+ return _bidirectional_mask(valid) & causal[None, None, :, :]
131
+
132
+
133
+ class RotaryPositionEncoding(nn.Module):
134
+ def __init__(self, head_width, max_length, *, base=10_000.0):
135
+ super().__init__()
136
+ if head_width % 2:
137
+ raise ValueError("RoPE head width must be even")
138
+ inv_freq = 1.0 / (
139
+ base ** (torch.arange(0, head_width, 2, dtype=torch.float32) / head_width)
140
+ )
141
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
142
+ frequencies = self._frequencies(max_length, inv_freq.device)
143
+ self.register_buffer("cos", frequencies.cos(), persistent=False)
144
+ self.register_buffer("sin", frequencies.sin(), persistent=False)
145
+
146
+ def _frequencies(self, length, device):
147
+ positions = torch.arange(length, dtype=torch.float32, device=device)
148
+ return torch.outer(positions, self.inv_freq.to(device=device))
149
+
150
+ def _rotate(self, value):
151
+ length = value.size(-2)
152
+ if length > self.cos.size(0):
153
+ frequencies = self._frequencies(length, value.device)
154
+ self.cos = frequencies.cos()
155
+ self.sin = frequencies.sin()
156
+ even, odd = value[..., 0::2], value[..., 1::2]
157
+ cos = self.cos[:length].to(device=value.device, dtype=value.dtype)
158
+ sin = self.sin[:length].to(device=value.device, dtype=value.dtype)
159
+ rotated = torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1)
160
+ return rotated.flatten(-2)
161
+
162
+ def encode(self, query, key):
163
+ return self._rotate(query), self._rotate(key)
164
+
165
+
166
+ class RelativeLogBucketSelfAttention(nn.Module):
167
+ def __init__(self, config):
168
+ super().__init__()
169
+ self.n_heads = config.num_attention_heads
170
+ self.d_head = config.hidden_size // config.num_attention_heads
171
+ self.max_seq_len = config.max_seq_len
172
+ self.buckets = config.position_buckets
173
+ self.value_gating = config.value_gating
174
+ self.use_rope = config.use_rope
175
+ self.use_alibi = config.use_alibi
176
+ self.shared_relative_embeddings = config.shared_relative_embeddings
177
+ self.attention_output_dropout = config.attention_output_dropout
178
+ if self.use_rope and self.use_alibi:
179
+ raise ValueError("RoPE and ALiBi are mutually exclusive")
180
+ self.qk = nn.Linear(config.hidden_size, 2 * config.hidden_size)
181
+ self.value = nn.Linear(
182
+ config.hidden_size,
183
+ 2 * config.hidden_size if config.value_gating else config.hidden_size,
184
+ )
185
+ self.out = nn.Linear(config.hidden_size, config.hidden_size)
186
+ self.dropout = nn.Dropout(config.attention_dropout)
187
+ self.relative_embedding = (
188
+ None
189
+ if self.use_rope or self.use_alibi or self.shared_relative_embeddings
190
+ else nn.Parameter(
191
+ torch.empty(2 * config.position_buckets - 1, config.hidden_size)
192
+ )
193
+ )
194
+ self.relative_norm = (
195
+ None
196
+ if self.relative_embedding is None
197
+ else nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
198
+ )
199
+ self.value_gate_norm = (
200
+ nn.LayerNorm(
201
+ config.hidden_size,
202
+ eps=config.layer_norm_eps,
203
+ elementwise_affine=False,
204
+ )
205
+ if config.value_gating
206
+ else None
207
+ )
208
+ self.rope = (
209
+ RotaryPositionEncoding(self.d_head, config.max_seq_len)
210
+ if config.use_rope
211
+ else None
212
+ )
213
+ self.scale = 1.0 / math.sqrt(
214
+ self.d_head if config.use_rope or config.use_alibi else 3.0 * self.d_head
215
+ )
216
+ if self.relative_embedding is not None:
217
+ nn.init.trunc_normal_(
218
+ self.relative_embedding,
219
+ mean=0.0,
220
+ std=config.initializer_range,
221
+ a=-2 * config.initializer_range,
222
+ b=2 * config.initializer_range,
223
+ )
224
+ self.register_buffer(
225
+ "position_indices",
226
+ self._position_indices(config.max_seq_len, torch.device("cpu")),
227
+ persistent=False,
228
+ )
229
+ self.register_buffer(
230
+ "alibi_bias",
231
+ self._alibi_bias(config.max_seq_len, torch.device("cpu")),
232
+ persistent=False,
233
+ )
234
+
235
+ def _alibi_bias(self, length, device):
236
+ positions = torch.arange(length, device=device)
237
+ distance = (positions[:, None] - positions[None, :]).abs().float()
238
+ slopes = torch.pow(
239
+ 2.0,
240
+ -8.0
241
+ * (torch.arange(self.n_heads, device=device).float() + 1.0)
242
+ / self.n_heads,
243
+ )
244
+ return -slopes[None, :, None, None] * distance[None, None, :, :]
245
+
246
+ def _position_indices(self, length, device):
247
+ positions = torch.arange(length, device=device)
248
+ relative = positions[:, None] - positions[None, :]
249
+ sign = torch.sign(relative)
250
+ half = self.buckets // 2
251
+ absolute = relative.abs().clamp(max=max(half + 1, self.max_seq_len - 1))
252
+ near = absolute <= half
253
+ safe = absolute.clamp_min(half)
254
+ denominator = math.log(max((self.max_seq_len - 1) / half, 1.0001))
255
+ logged = (
256
+ torch.ceil(torch.log(safe / half) / denominator * (half - 1)).long() + half
257
+ )
258
+ bucketed = torch.where(near, relative, logged * sign)
259
+ return (
260
+ bucketed.long().clamp(-self.buckets + 1, self.buckets - 1)
261
+ + self.buckets
262
+ - 1
263
+ )
264
+
265
+ def forward(self, x, mask, relative_embedding=None):
266
+ batch, length, width = x.shape
267
+ if length > self.position_indices.size(0):
268
+ self.position_indices = self._position_indices(length, x.device)
269
+ q, k = self.qk(x).chunk(2, dim=-1)
270
+ if self.value_gating:
271
+ v, gate = self.value(x).chunk(2, dim=-1)
272
+ gate = F.gelu(gate)
273
+ else:
274
+ v, gate = self.value(x), None
275
+ q = q.view(batch, length, self.n_heads, self.d_head).transpose(1, 2)
276
+ k = k.view(batch, length, self.n_heads, self.d_head).transpose(1, 2)
277
+ v = v.view(batch, length, self.n_heads, self.d_head).transpose(1, 2)
278
+
279
+ if self.rope is not None:
280
+ q, k = self.rope.encode(q, k)
281
+ scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
282
+ elif self.use_alibi:
283
+ if length > self.alibi_bias.size(-1):
284
+ self.alibi_bias = self._alibi_bias(length, x.device)
285
+ scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
286
+ scores = scores + self.alibi_bias[:, :, :length, :length].to(
287
+ device=x.device, dtype=scores.dtype
288
+ )
289
+ else:
290
+ if relative_embedding is None:
291
+ assert self.relative_embedding is not None
292
+ assert self.relative_norm is not None
293
+ relative_embedding = self.relative_norm(self.relative_embedding)
294
+ relative = self.qk(self.dropout(relative_embedding))
295
+ relative = relative[self.position_indices[:length, :length].to(x.device)]
296
+ q_pos, k_pos = relative.chunk(2, dim=-1)
297
+ q_pos = q_pos.view(length, length, self.n_heads, self.d_head)
298
+ k_pos = k_pos.view(length, length, self.n_heads, self.d_head)
299
+ scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
300
+ scores = scores + torch.einsum("bhqd,qkhd->bhqk", q, k_pos) * self.scale
301
+ scores = scores + torch.einsum("bhkd,qkhd->bhqk", k, q_pos) * self.scale
302
+
303
+ probs = torch.softmax(
304
+ scores.masked_fill(~mask, torch.finfo(scores.dtype).min), dim=-1
305
+ )
306
+ probs = self.dropout(probs) if self.training else probs
307
+ output = (
308
+ torch.matmul(probs, v)
309
+ .transpose(1, 2)
310
+ .contiguous()
311
+ .view(batch, length, width)
312
+ )
313
+ if gate is not None and self.value_gate_norm is not None:
314
+ output = self.value_gate_norm(output * gate)
315
+ output = self.out(output)
316
+ return self.dropout(output) if self.attention_output_dropout else output
317
+
318
+
319
+ class GeGLU(nn.Module):
320
+ def __init__(self, config, width=None):
321
+ super().__init__()
322
+ width = width or config.intermediate_size
323
+ self.up = nn.Linear(config.hidden_size, 2 * width, bias=False)
324
+ self.post_activation_norm = nn.LayerNorm(
325
+ width, eps=config.layer_norm_eps, elementwise_affine=False
326
+ )
327
+ self.down = nn.Linear(width, config.hidden_size, bias=False)
328
+ self.dropout = nn.Dropout(config.dropout)
329
+ self.dropout_after_projection = config.feedforward_dropout_after_projection
330
+
331
+ def forward(self, x):
332
+ value, gate = self.up(x).chunk(2, dim=-1)
333
+ hidden = value * F.gelu(gate, approximate="tanh")
334
+ hidden = self.post_activation_norm(hidden)
335
+ if self.dropout_after_projection:
336
+ return self.dropout(self.down(hidden))
337
+ return self.down(self.dropout(hidden))
338
+
339
+
340
+ class RoutedGeGLU(nn.Module):
341
+ def __init__(self, config):
342
+ super().__init__()
343
+ if not 1 <= config.experts_per_token <= config.num_experts:
344
+ raise ValueError("experts_per_token must be in [1, num_experts]")
345
+ width = config.expert_intermediate_size or max(
346
+ 1, config.intermediate_size // config.num_experts
347
+ )
348
+ self.top_k = config.experts_per_token
349
+ self.router = nn.Linear(config.hidden_size, config.num_experts, bias=False)
350
+ self.experts = nn.ModuleList(
351
+ GeGLU(config, width) for _ in range(config.num_experts)
352
+ )
353
+
354
+ def forward(self, x):
355
+ probabilities = self.router(x).softmax(dim=-1)
356
+ weights, indices = probabilities.topk(self.top_k, dim=-1)
357
+ weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-8)
358
+ gates = torch.zeros_like(probabilities).scatter(-1, indices, weights)
359
+ outputs = torch.stack([expert(x) for expert in self.experts], dim=-2)
360
+ return (outputs * gates.unsqueeze(-1)).sum(dim=-2)
361
+
362
+
363
+ class CausalStateMixer(nn.Module):
364
+ def __init__(self, config):
365
+ super().__init__()
366
+ kernel = int(config.state_mixer_kernel)
367
+ self.input = nn.Linear(config.hidden_size, 2 * config.hidden_size, bias=False)
368
+ self.state = nn.Conv1d(
369
+ config.hidden_size,
370
+ config.hidden_size,
371
+ kernel,
372
+ groups=config.hidden_size,
373
+ padding=kernel - 1,
374
+ bias=False,
375
+ )
376
+ self.output = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
377
+ self.gate = nn.Parameter(torch.tensor(-2.0))
378
+
379
+ def forward(self, hidden):
380
+ value, gate = self.input(hidden).chunk(2, dim=-1)
381
+ state = self.state(value.transpose(1, 2))[..., : hidden.size(1)].transpose(1, 2)
382
+ return self.output(state * F.silu(gate)) * self.gate.sigmoid()
383
+
384
+
385
+ class DynamicWeightedAverage(nn.Module):
386
+ def __init__(self, n_sublayers):
387
+ super().__init__()
388
+ self.alphas = nn.ParameterList(
389
+ nn.Parameter(torch.cat([torch.zeros(i + 1), torch.ones(1)]))
390
+ for i in range(int(n_sublayers))
391
+ )
392
+ self._states = None
393
+
394
+ def initialize(self, hidden):
395
+ self._states = [hidden]
396
+
397
+ def forward(self, hidden, sublayer_index):
398
+ self._states.append(hidden)
399
+ return torch.tensordot(
400
+ self.alphas[sublayer_index], torch.stack(self._states), dims=1
401
+ )
402
+
403
+
404
+ class GPTBertBlock(nn.Module):
405
+ def __init__(self, config):
406
+ super().__init__()
407
+ self.attention_norm = nn.LayerNorm(
408
+ config.hidden_size,
409
+ eps=config.layer_norm_eps,
410
+ elementwise_affine=False,
411
+ )
412
+ self.attention = RelativeLogBucketSelfAttention(config)
413
+ self.state_mixer = (
414
+ CausalStateMixer(config) if config.state_mixer_kernel else None
415
+ )
416
+ self.feedforward_norm = nn.LayerNorm(
417
+ config.hidden_size,
418
+ eps=config.layer_norm_eps,
419
+ elementwise_affine=False,
420
+ )
421
+ self.feedforward = (
422
+ RoutedGeGLU(config) if config.num_experts > 1 else GeGLU(config)
423
+ )
424
+
425
+ def attend(self, hidden, mask, relative_embedding=None):
426
+ normalized = self.attention_norm(hidden)
427
+ attention = self.attention(normalized, mask, relative_embedding)
428
+ return (
429
+ attention
430
+ if self.state_mixer is None
431
+ else attention + self.state_mixer(normalized)
432
+ )
433
+
434
+ def transform(self, hidden):
435
+ return self.feedforward(self.feedforward_norm(hidden))
436
+
437
+
438
+ class GPTBertBackbone(nn.Module):
439
+ def __init__(self, config):
440
+ super().__init__()
441
+ self.config = config
442
+ self.embed_tokens = nn.Embedding(
443
+ config.vocab_size,
444
+ config.hidden_size,
445
+ padding_idx=config.pad_token_id if config.embedding_padding_idx else None,
446
+ )
447
+ self.relative_embedding = (
448
+ nn.Parameter(
449
+ torch.empty(
450
+ 2 * config.position_buckets - 1, config.hidden_size
451
+ )
452
+ )
453
+ if config.shared_relative_embeddings
454
+ else None
455
+ )
456
+ self.relative_norm = (
457
+ nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
458
+ if config.shared_relative_embeddings
459
+ else None
460
+ )
461
+ if self.relative_embedding is not None:
462
+ nn.init.trunc_normal_(
463
+ self.relative_embedding,
464
+ std=config.initializer_range,
465
+ a=-2 * config.initializer_range,
466
+ b=2 * config.initializer_range,
467
+ )
468
+ self.geometry_lexical_dim = int(config.geometry_lexical_dim)
469
+ if not 0 <= self.geometry_lexical_dim < config.hidden_size:
470
+ raise ValueError("geometry_lexical_dim must be in [0, hidden_size)")
471
+ self.geometry_curvature = float(config.geometry_curvature)
472
+ if self.geometry_curvature <= 0:
473
+ raise ValueError("geometry_curvature must be positive")
474
+ self.lexical_angle = None
475
+ self.lexical_radius = None
476
+ if self.geometry_lexical_dim:
477
+ self.lexical_angle = nn.Embedding(
478
+ config.vocab_size, self.geometry_lexical_dim, config.pad_token_id
479
+ )
480
+ self.lexical_radius = nn.Embedding(config.vocab_size, 1, config.pad_token_id)
481
+ self.embed_positions = (
482
+ nn.Embedding(config.max_seq_len, config.hidden_size)
483
+ if getattr(config, "absolute_positions", False)
484
+ else None
485
+ )
486
+ self.embed_norm = nn.LayerNorm(
487
+ config.hidden_size,
488
+ eps=config.layer_norm_eps,
489
+ elementwise_affine=False,
490
+ )
491
+ self.dropout = nn.Dropout(config.dropout)
492
+ self.blocks = nn.ModuleList(
493
+ GPTBertBlock(config) for _ in range(config.num_hidden_layers)
494
+ )
495
+ self.recurrent_steps = max(1, int(config.recurrent_steps))
496
+ self.future_projections = nn.ModuleDict(
497
+ {
498
+ str(offset): nn.Linear(
499
+ config.hidden_size, config.hidden_size, bias=False
500
+ )
501
+ for offset in config.future_offsets
502
+ }
503
+ )
504
+ self.residual_mixer = (
505
+ DynamicWeightedAverage(config.num_hidden_layers * self.recurrent_steps * 2)
506
+ if config.residual_mixing
507
+ else None
508
+ )
509
+ self.cognitive_readout_layer = int(config.cognitive_readout_layer)
510
+ self.cognitive_readout_weight = float(config.cognitive_readout_weight)
511
+ maximum_depth = config.num_hidden_layers * self.recurrent_steps
512
+ if self.cognitive_readout_layer > maximum_depth:
513
+ raise ValueError("cognitive_readout_layer exceeds the executed depth")
514
+
515
+ def lexical_geometry(self, token_ids):
516
+ if self.lexical_angle is None or self.lexical_radius is None:
517
+ raise RuntimeError("lexical geometry is disabled")
518
+ direction = F.normalize(self.lexical_angle(token_ids), dim=-1)
519
+ radius = F.softplus(self.lexical_radius(token_ids)).squeeze(-1)
520
+ scale = math.sqrt(self.geometry_curvature)
521
+ point = torch.tanh(scale * radius / 2).unsqueeze(-1) * direction / scale
522
+ return point, radius
523
+
524
+ def forward(self, input_ids, mask):
525
+ embedded = self.embed_tokens(input_ids)
526
+ if self.geometry_lexical_dim:
527
+ _, radius = self.lexical_geometry(input_ids)
528
+ direction = F.normalize(self.lexical_angle(input_ids), dim=-1)
529
+ tangent = radius.unsqueeze(-1) * direction
530
+ embedded = torch.cat((embedded[..., :-self.geometry_lexical_dim], tangent), -1)
531
+ if self.embed_positions is not None:
532
+ positions = torch.arange(input_ids.size(1), device=input_ids.device)
533
+ embedded = embedded + self.embed_positions(positions)
534
+ hidden = self.dropout(self.embed_norm(embedded))
535
+ mixer = self.residual_mixer
536
+ if mixer is not None:
537
+ mixer.initialize(hidden)
538
+ sublayer = 0
539
+ cognitive_hidden = None
540
+ layer_index = 0
541
+ relative = (
542
+ self.relative_norm(self.relative_embedding)
543
+ if self.relative_norm is not None and self.relative_embedding is not None
544
+ else None
545
+ )
546
+ for _ in range(self.recurrent_steps):
547
+ for block in self.blocks:
548
+ hidden = hidden + block.attend(hidden, mask, relative)
549
+ if mixer is not None:
550
+ hidden = mixer(hidden, sublayer)
551
+ sublayer += 1
552
+ hidden = hidden + block.transform(hidden)
553
+ if mixer is not None:
554
+ hidden = mixer(hidden, sublayer)
555
+ sublayer += 1
556
+ layer_index += 1
557
+ if layer_index == self.cognitive_readout_layer:
558
+ cognitive_hidden = hidden
559
+ if cognitive_hidden is not None and self.cognitive_readout_weight > 0:
560
+ weight = self.cognitive_readout_weight
561
+ hidden = (1.0 - weight) * hidden + weight * cognitive_hidden
562
+ return hidden
563
+
564
+
565
+ class DirectSumStream(nn.Module):
566
+ def __init__(self, config):
567
+ super().__init__()
568
+ self.relative_embedding = (
569
+ nn.Parameter(torch.empty(2 * config.position_buckets - 1, config.hidden_size))
570
+ if config.shared_relative_embeddings
571
+ else None
572
+ )
573
+ self.relative_norm = (
574
+ nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
575
+ if config.shared_relative_embeddings
576
+ else None
577
+ )
578
+ if self.relative_embedding is not None:
579
+ nn.init.trunc_normal_(
580
+ self.relative_embedding,
581
+ std=config.initializer_range,
582
+ a=-2 * config.initializer_range,
583
+ b=2 * config.initializer_range,
584
+ )
585
+ self.embed_positions = (
586
+ nn.Embedding(config.max_seq_len, config.hidden_size)
587
+ if config.absolute_positions
588
+ else None
589
+ )
590
+ self.embed_norm = nn.LayerNorm(
591
+ config.hidden_size,
592
+ eps=config.layer_norm_eps,
593
+ elementwise_affine=False,
594
+ )
595
+ self.dropout = nn.Dropout(config.dropout)
596
+ self.blocks = nn.ModuleList(
597
+ GPTBertBlock(config) for _ in range(config.num_hidden_layers)
598
+ )
599
+ self.recurrent_steps = max(1, int(config.recurrent_steps))
600
+ self.residual_mixer = (
601
+ DynamicWeightedAverage(config.num_hidden_layers * self.recurrent_steps * 2)
602
+ if config.residual_mixing
603
+ else None
604
+ )
605
+
606
+ def forward(self, embedded, mask):
607
+ if self.embed_positions is not None:
608
+ positions = torch.arange(embedded.size(1), device=embedded.device)
609
+ embedded = embedded + self.embed_positions(positions)
610
+ hidden = self.dropout(self.embed_norm(embedded))
611
+ mixer = self.residual_mixer
612
+ if mixer is not None:
613
+ mixer.initialize(hidden)
614
+ relative = (
615
+ self.relative_norm(self.relative_embedding)
616
+ if self.relative_norm is not None and self.relative_embedding is not None
617
+ else None
618
+ )
619
+ sublayer = 0
620
+ for _ in range(self.recurrent_steps):
621
+ for block in self.blocks:
622
+ hidden = hidden + block.attend(hidden, mask, relative)
623
+ if mixer is not None:
624
+ hidden = mixer(hidden, sublayer)
625
+ sublayer += 1
626
+ hidden = hidden + block.transform(hidden)
627
+ if mixer is not None:
628
+ hidden = mixer(hidden, sublayer)
629
+ sublayer += 1
630
+ return hidden
631
+
632
+
633
+ class DirectSumBackbone(nn.Module):
634
+ def __init__(self, config):
635
+ super().__init__()
636
+ dims = tuple(int(value) for value in config.direct_sum_dims)
637
+ heads = tuple(int(value) for value in config.direct_sum_heads)
638
+ widths = tuple(int(value) for value in config.direct_sum_intermediate_sizes)
639
+ if len(dims) != 3 or len(heads) != 3 or len(widths) != 3:
640
+ raise ValueError("direct sum requires three dims, heads and FFN widths")
641
+ if sum(dims) != config.hidden_size:
642
+ raise ValueError("direct_sum_dims must sum to hidden_size")
643
+ if any(dim % head for dim, head in zip(dims, heads)):
644
+ raise ValueError("each direct-sum dimension must divide its head count")
645
+ self.dims = dims
646
+ self.embed_tokens = nn.Embedding(
647
+ config.vocab_size,
648
+ config.hidden_size,
649
+ padding_idx=config.pad_token_id if config.embedding_padding_idx else None,
650
+ )
651
+ streams = []
652
+ for dim, head, width in zip(dims, heads, widths):
653
+ stream_config = copy.copy(config)
654
+ stream_config.hidden_size = dim
655
+ stream_config.num_attention_heads = head
656
+ stream_config.intermediate_size = width
657
+ stream_config.direct_sum_dims = []
658
+ stream_config.direct_sum_heads = []
659
+ stream_config.direct_sum_intermediate_sizes = []
660
+ stream_config.geometry_lexical_dim = 0
661
+ stream_config.future_offsets = []
662
+ stream_config.cognitive_readout_layer = 0
663
+ stream_config.cognitive_readout_weight = 0.0
664
+ streams.append(DirectSumStream(stream_config))
665
+ self.streams = nn.ModuleList(streams)
666
+ self.concept_radius = nn.Embedding(
667
+ config.vocab_size, 1, padding_idx=config.pad_token_id
668
+ )
669
+
670
+ @property
671
+ def factor_slices(self):
672
+ syntax, lexical, conceptual = self.dims
673
+ return (
674
+ slice(0, syntax),
675
+ slice(syntax, syntax + lexical),
676
+ slice(syntax + lexical, syntax + lexical + conceptual),
677
+ )
678
+
679
+ def conceptual_geometry(self, token_ids):
680
+ conceptual = self.embed_tokens(token_ids)[..., self.factor_slices[2]]
681
+ direction = F.normalize(conceptual, dim=-1)
682
+ radius = (1.0 - 1.0e-4) * torch.sigmoid(
683
+ self.concept_radius(token_ids).squeeze(-1)
684
+ )
685
+ if self.concept_radius.padding_idx is not None:
686
+ radius = radius.masked_fill(
687
+ token_ids.eq(self.concept_radius.padding_idx), 0.0
688
+ )
689
+ return radius.unsqueeze(-1) * direction, radius
690
+
691
+ def forward(self, input_ids, mask):
692
+ embedded = self.embed_tokens(input_ids)
693
+ conceptual, _ = self.conceptual_geometry(input_ids)
694
+ parts = list(embedded.split(self.dims, dim=-1))
695
+ parts[2] = conceptual
696
+ return torch.cat(
697
+ [stream(part, mask) for stream, part in zip(self.streams, parts)], dim=-1
698
+ )
699
+
700
+
701
+ def _backbone(config):
702
+ return DirectSumBackbone(config) if config.direct_sum_dims else GPTBertBackbone(config)
703
+
704
+
705
+ class GPTBertLMHead(nn.Module):
706
+ def __init__(self, config):
707
+ super().__init__()
708
+ self.norm = nn.LayerNorm(
709
+ config.hidden_size,
710
+ eps=config.layer_norm_eps,
711
+ elementwise_affine=False,
712
+ )
713
+ self.dense = nn.Linear(config.hidden_size, config.hidden_size)
714
+ self.post_norm = nn.LayerNorm(
715
+ config.hidden_size,
716
+ eps=config.layer_norm_eps,
717
+ elementwise_affine=False,
718
+ )
719
+ self.dropout = nn.Dropout(config.dropout)
720
+ self._approximate = config.lm_head_gelu_approximate
721
+ self.bias = nn.Parameter(torch.zeros(config.vocab_size))
722
+
723
+ def forward(self, hidden):
724
+ projected = self.dropout(
725
+ self.post_norm(
726
+ F.gelu(
727
+ self.dense(self.norm(hidden)), approximate=self._approximate
728
+ )
729
+ )
730
+ )
731
+ return F.linear(projected, self.weight, self.bias)
732
+
733
+
734
+ class TOLMModel(PreTrainedModel):
735
+ config_class = TOLMConfig
736
+ base_model_prefix = "tolm"
737
+ _no_split_modules = ["GPTBertBlock"]
738
+
739
+ def __init__(self, config):
740
+ super().__init__(config)
741
+ self.backbone = _backbone(config)
742
+ self.post_init()
743
+
744
+ def get_input_embeddings(self):
745
+ return self.backbone.embed_tokens
746
+
747
+ def forward(self, input_ids=None, attention_mask=None, **kwargs):
748
+ if input_ids is None:
749
+ raise ValueError("input_ids is required")
750
+ hidden = self.backbone(
751
+ input_ids, _bidirectional_mask(_valid_tokens(input_ids, attention_mask))
752
+ )
753
+ return BaseModelOutput(
754
+ last_hidden_state=hidden, hidden_states=None, attentions=None
755
+ )
756
+
757
+
758
+ class TOLMForMaskedLM(PreTrainedModel):
759
+ config_class = TOLMConfig
760
+ base_model_prefix = "tolm"
761
+ _no_split_modules = ["GPTBertBlock"]
762
+ _tied_weights_keys = ["heads.lm.weight"]
763
+
764
+ def __init__(self, config):
765
+ super().__init__(config)
766
+ self.backbone = _backbone(config)
767
+ head = GPTBertLMHead(config)
768
+ self.heads = nn.ModuleDict({"lm": head})
769
+ if config.direct_sum_dims:
770
+ self.factor_dual_lambdas = nn.Parameter(
771
+ torch.ones(3), requires_grad=False
772
+ )
773
+ self.post_init()
774
+
775
+ def get_input_embeddings(self):
776
+ return self.backbone.embed_tokens
777
+
778
+ def get_output_embeddings(self):
779
+ return self.heads["lm"]
780
+
781
+ def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
782
+ if input_ids is None:
783
+ raise ValueError("input_ids is required")
784
+ hidden = self.backbone(
785
+ input_ids, _bidirectional_mask(_valid_tokens(input_ids, attention_mask))
786
+ )
787
+ logits = self.heads["lm"](hidden)
788
+ loss = (
789
+ None
790
+ if labels is None
791
+ else F.cross_entropy(
792
+ logits.view(-1, logits.size(-1)), labels.view(-1), ignore_index=-100
793
+ )
794
+ )
795
+ return MaskedLMOutput(
796
+ loss=loss, logits=logits, hidden_states=None, attentions=None
797
+ )
798
+
799
+
800
+ class TOLMForCausalLM(PreTrainedModel):
801
+ config_class = TOLMConfig
802
+ base_model_prefix = "tolm"
803
+ _no_split_modules = ["GPTBertBlock"]
804
+ _tied_weights_keys = ["heads.lm.weight"]
805
+
806
+ def __init__(self, config):
807
+ super().__init__(config)
808
+ self.backbone = _backbone(config)
809
+ head = GPTBertLMHead(config)
810
+ self.heads = nn.ModuleDict({"lm": head})
811
+ if config.direct_sum_dims:
812
+ self.factor_dual_lambdas = nn.Parameter(
813
+ torch.ones(3), requires_grad=False
814
+ )
815
+ self.post_init()
816
+
817
+ def get_input_embeddings(self):
818
+ return self.backbone.embed_tokens
819
+
820
+ def get_output_embeddings(self):
821
+ return self.heads["lm"]
822
+
823
+ def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **kwargs):
824
+ return {"input_ids": input_ids, "attention_mask": attention_mask}
825
+
826
+ def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
827
+ if input_ids is None:
828
+ raise ValueError("input_ids is required")
829
+ hidden = self.backbone(
830
+ input_ids, _causal_mask(_valid_tokens(input_ids, attention_mask))
831
+ )
832
+ logits = self.heads["lm"](hidden)
833
+ loss = None
834
+ if labels is not None:
835
+ loss = F.cross_entropy(
836
+ logits[:, :-1].contiguous().view(-1, logits.size(-1)),
837
+ labels[:, 1:].contiguous().view(-1),
838
+ ignore_index=-100,
839
+ )
840
+ return CausalLMOutput(
841
+ loss=loss, logits=logits, hidden_states=None, attentions=None
842
+ )
training_manifest.json ADDED
@@ -0,0 +1,218 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "artifact_provenance": {
3
+ "corpus_manifest_sha256": "404b1a6db939378f3066f0fbbaf7bfd3c0044bc2e5688cad0e688d3e8c5663cc",
4
+ "corpus_source_sha256": "d6f7db41b8a8cb48321ab7192d884bc8fc9275be3e0022bd76ec9ed776ca0895",
5
+ "corpus_words": 10000000,
6
+ "factor_manifest_sha256": "5ebd31d5d4e1ea69c3a202431ae8c0b3bc6cc249307a95b80edfa50c0b579f68",
7
+ "tokenizer_sha256": "abfd627e173d48addf924e32fd6c38a3d485734c496d12916e559e654ba15c51"
8
+ },
9
+ "compliance": {
10
+ "competition_status": "strict-safe",
11
+ "external_teacher": false,
12
+ "generated_words": 0,
13
+ "same_run_checkpoint_distribution": false,
14
+ "teacher_queries": 0,
15
+ "total_exposure_words": 100000000,
16
+ "track": "strict-small",
17
+ "unique_corpus_words": 10000000
18
+ },
19
+ "factor_priors": {
20
+ "angular_margin": 0.1,
21
+ "dual": {
22
+ "enabled": false,
23
+ "learning_rate": 0.001,
24
+ "max_weight": 1.0,
25
+ "targets": [
26
+ 1.0,
27
+ 4.0,
28
+ 0.1
29
+ ]
30
+ },
31
+ "enabled": true,
32
+ "lambda_conceptual": 0.0,
33
+ "lambda_lexical": 0.02,
34
+ "lambda_syntax": 0.05,
35
+ "max_positions": 256,
36
+ "radial_margin": 0.02,
37
+ "radial_weight": 1.0,
38
+ "randomize": false,
39
+ "randomize_seed": 42,
40
+ "syntax_depth_weight": 0.1,
41
+ "syntax_distance_weight": 0.1,
42
+ "syntax_window": 32,
43
+ "temperature": 0.07,
44
+ "warmup_words": 1000000
45
+ },
46
+ "geometry": {
47
+ "curvature": 1.0,
48
+ "distance_margin": 0.1,
49
+ "enabled": false,
50
+ "lambda_radial": 0.02,
51
+ "lambda_related": 0.02,
52
+ "max_tokens": 256,
53
+ "radial_margin": 0.05,
54
+ "warmup_words": 5000000
55
+ },
56
+ "model": {
57
+ "absolute_positions": false,
58
+ "attention_dropout": 0.1,
59
+ "bos_token_id": 2,
60
+ "cognitive_readout_layer": 0,
61
+ "cognitive_readout_weight": 0.0,
62
+ "direct_sum_dims": [
63
+ 96,
64
+ 128,
65
+ 288
66
+ ],
67
+ "direct_sum_heads": [
68
+ 3,
69
+ 4,
70
+ 6
71
+ ],
72
+ "direct_sum_intermediate_sizes": [
73
+ 512,
74
+ 640,
75
+ 1344
76
+ ],
77
+ "dropout": 0.1,
78
+ "eos_token_id": 3,
79
+ "expert_intermediate_size": null,
80
+ "experts_per_token": 1,
81
+ "future_offsets": [],
82
+ "geometry_curvature": 1.0,
83
+ "geometry_lexical_dim": 0,
84
+ "hidden_size": 512,
85
+ "initializer_range": 0.03227486121839514,
86
+ "intermediate_size": 1536,
87
+ "mask_token_id": 4,
88
+ "max_seq_len": 512,
89
+ "num_attention_heads": 8,
90
+ "num_experts": 1,
91
+ "num_hidden_layers": 12,
92
+ "pad_token_id": 1,
93
+ "position_buckets": 32,
94
+ "recurrent_steps": 1,
95
+ "residual_mixing": true,
96
+ "rtd_auxiliary": false,
97
+ "state_mixer_kernel": 0,
98
+ "use_alibi": false,
99
+ "use_rope": false,
100
+ "value_gating": true,
101
+ "vocab_size": 16000
102
+ },
103
+ "release": "TOLM",
104
+ "structured_priors": {
105
+ "decay_end_words": 7000000,
106
+ "enabled": false,
107
+ "hold_until_words": 2000000,
108
+ "init_scale": 0.0,
109
+ "lambda_lexical": 0.0,
110
+ "lambda_orth": 0.0,
111
+ "lambda_syntax": 0.0,
112
+ "lexical_dim": 128,
113
+ "max_positions": 256,
114
+ "prior_mode": "contrastive",
115
+ "syntax_dim": 128,
116
+ "temperature": 0.07,
117
+ "warmup_words": 100000
118
+ },
119
+ "training": {
120
+ "adaptive_masking": {
121
+ "enabled": false,
122
+ "max_mask_prob": 4.0,
123
+ "min_mask_prob": 0.25,
124
+ "momentum": 0.99
125
+ },
126
+ "batch_size": 16,
127
+ "beta1": 0.9,
128
+ "beta2": 0.98,
129
+ "causal_fraction": 0.0,
130
+ "causal_noise_kind": "random",
131
+ "causal_noise_probability": 0.0,
132
+ "causal_unit": "segment",
133
+ "cooldown_fraction": 0.016,
134
+ "data2vec_layers": 4,
135
+ "data2vec_weight": 0.5,
136
+ "device": "xpu:0",
137
+ "ema_decay": 0.9998,
138
+ "epsilon": 1e-08,
139
+ "exposure_words": 100000000,
140
+ "final_lr_ratio": 0.1,
141
+ "frequency_aware_masking": {
142
+ "enabled": false,
143
+ "max_mask_prob": 4.0,
144
+ "min_mask_prob": 0.25,
145
+ "temperature": 1.0
146
+ },
147
+ "future_loss_weight": 0.0,
148
+ "gradient_clip": 2.0,
149
+ "label_smoothing": 0.0,
150
+ "learning_progress": {
151
+ "buckets_per_axis": 4,
152
+ "enabled": false,
153
+ "fast_momentum": 0.9,
154
+ "forgetting_weight": 1.0,
155
+ "slow_momentum": 0.99,
156
+ "uniform_floor": 0.3,
157
+ "window_documents": 4000
158
+ },
159
+ "learning_rate": 0.0035,
160
+ "learning_rate_schedule": "cosine",
161
+ "log_interval": 50,
162
+ "mask_probability_end": 0.5,
163
+ "mask_probability_schedule": "uniform",
164
+ "mask_probability_start": 0.15,
165
+ "mask_replace_probability": 0.8,
166
+ "mask_schedule": "complementary",
167
+ "masked_fraction": 1.0,
168
+ "masking_unit": "word",
169
+ "microbatch_tokens": 8192,
170
+ "mixed_precision": "bf16",
171
+ "num_workers": 0,
172
+ "objective_period": 16,
173
+ "objective_rng_reset_words": [],
174
+ "objective_sanity_window_steps": 512,
175
+ "objective_schedule": "coverage",
176
+ "optimizer": "lamb",
177
+ "packing_strategy": "dense",
178
+ "pin_memory": false,
179
+ "random_replace_probability": 0.1,
180
+ "router_aux_weight": 0.0,
181
+ "save_steps_words": [
182
+ 1000000,
183
+ 2000000,
184
+ 3000000,
185
+ 4000000,
186
+ 5000000,
187
+ 6000000,
188
+ 7000000,
189
+ 8000000,
190
+ 9000000,
191
+ 10000000,
192
+ 20000000,
193
+ 30000000,
194
+ 40000000,
195
+ 50000000,
196
+ 60000000,
197
+ 70000000,
198
+ 80000000,
199
+ 90000000,
200
+ 100000000
201
+ ],
202
+ "schedule_total_words": 100000000,
203
+ "span_max_length": 3,
204
+ "stable_fraction": 0.9,
205
+ "telemetry": {
206
+ "enabled": true,
207
+ "gradient_interval_words": 2000000,
208
+ "sampler_trace": true
209
+ },
210
+ "threads_per_process": 8,
211
+ "tokens_per_update": 16384,
212
+ "warmup_fraction": 0.016,
213
+ "weight_decay": 0.1,
214
+ "z_loss_weight": 0.0001
215
+ },
216
+ "variant": "structured_ds_syntax_lexical",
217
+ "words_seen": 100000000
218
+ }