miguelcsx commited on
Commit
d545a34
·
verified ·
1 Parent(s): 5ba27fe

main chck_30M

Browse files
.complete.json ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "files": {
3
+ "config.json": "d46dca7c612cb8e054389d73b073d8ddd9d3aa23b480e2bc5fa75830b8f55fe8",
4
+ "dial_release.json": "e4ccf1959de8364ecf1e08c481b8c2ceaef041ae1602f163c472928ad545a92a",
5
+ "dial_scaffold.safetensors": "653fded2bd45b6d5c46cdaaa71f53c454a874fba4649097061baa6c431da5543",
6
+ "model.safetensors": "9f801287db7ad1e077ad9cdf433a3f9cdc2badb3ae98b961ac3078bc81f342a4",
7
+ "special_tokens_map.json": "2be97b602cd0e4a6c2874cbf03c0d1025b3af5474666da931bc04f6c23a9a39d",
8
+ "tokenizer.json": "2b8d1b3f51c0d8f276a64ea6a69efa50b4d9899780f1b282b4d0dcb7be5bbb0f",
9
+ "tokenizer_config.json": "8f4a998564d87fc85b27652f0b295677a1e03744bf400bfacc342b700cc3653a",
10
+ "tolm.py": "604eb005807b93b3a0c2145851c78f8d61a6f2e27d0fedf0889de844ee8e1713",
11
+ "training_manifest.json": "3c31f8afe5286b75e6f5478be6dd80ae081b43c0fd74c18e3f2f25f9acd0c084"
12
+ },
13
+ "schema_version": 1,
14
+ "words_seen": 30000000
15
+ }
config.json ADDED
@@ -0,0 +1,61 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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": 16384,
13
+ "max_seq_len": 512,
14
+ "hidden_size": 384,
15
+ "num_hidden_layers": 12,
16
+ "num_attention_heads": 6,
17
+ "intermediate_size": 1280,
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": 3,
25
+ "bos_token_id": 1,
26
+ "eos_token_id": 2,
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": true,
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
+ "direct_sum_heads": [],
44
+ "direct_sum_intermediate_sizes": [],
45
+ "lexical_residual_buckets": 0,
46
+ "lexical_residual_dim": 0,
47
+ "lexical_residual_scale": 1.0,
48
+ "structured_projection_dim": 0,
49
+ "factor_readout_mode": "",
50
+ "factor_readout_dims": [
51
+ 128,
52
+ 128,
53
+ 128
54
+ ],
55
+ "factor_readout_reflectors": 384,
56
+ "factor_readout_seed": 314159,
57
+ "tie_word_embeddings": false,
58
+ "dial_mode": "shuffled",
59
+ "dial_rank": 128,
60
+ "dial_projector_seed": 20260719
61
+ }
dial_release.json ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "auxiliary_active_at_checkpoint": false,
3
+ "auxiliary_coefficient_at_checkpoint": 0.0,
4
+ "calibration_minibatch": 0,
5
+ "checkpoint_words": 30000000,
6
+ "condition": "shuffled",
7
+ "factor_scales": {
8
+ "associative": 0.5741430521011353,
9
+ "hierarchical": 7.806242942810059,
10
+ "order": 0.8848216533660889
11
+ },
12
+ "initial_aux_gradient_ratio": 0.25,
13
+ "input_output_embeddings_tied": false,
14
+ "lambda0": 0.1258963644504547,
15
+ "projectors": {
16
+ "concatenated_condition_number": 1.0,
17
+ "cross_block_overlap_frobenius": [
18
+ 0.0,
19
+ 0.0,
20
+ 0.0
21
+ ],
22
+ "mode": "shuffled",
23
+ "ranks": [
24
+ 128,
25
+ 128,
26
+ 128
27
+ ],
28
+ "within_block_orthonormal_error": [
29
+ 0.0,
30
+ 0.0,
31
+ 0.0
32
+ ]
33
+ },
34
+ "schema_version": 1,
35
+ "variant": "dial_shuffled_s3",
36
+ "withdrawal_words": 10000000
37
+ }
dial_scaffold.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:653fded2bd45b6d5c46cdaaa71f53c454a874fba4649097061baa6c431da5543
3
+ size 983497
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9f801287db7ad1e077ad9cdf433a3f9cdc2badb3ae98b961ac3078bc81f342a4
3
+ size 158531528
special_tokens_map.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token": "<s>",
3
+ "eos_token": "</s>",
4
+ "mask_token": "<mask>",
5
+ "pad_token": "<pad>",
6
+ "unk_token": "<unk>"
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,922 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ lexical_residual_buckets=0,
62
+ lexical_residual_dim=0,
63
+ lexical_residual_scale=1.0,
64
+ structured_projection_dim=0,
65
+ tie_word_embeddings=True,
66
+ dial_mode="",
67
+ dial_rank=128,
68
+ dial_projector_seed=20260719,
69
+ **kwargs,
70
+ ):
71
+ super().__init__(
72
+ pad_token_id=pad_token_id,
73
+ bos_token_id=bos_token_id,
74
+ eos_token_id=eos_token_id,
75
+ mask_token_id=mask_token_id,
76
+ tie_word_embeddings=tie_word_embeddings,
77
+ **kwargs,
78
+ )
79
+ self.vocab_size = vocab_size
80
+ self.max_seq_len = max_position_embeddings or max_seq_len
81
+ self.max_position_embeddings = self.max_seq_len
82
+ self.hidden_size = hidden_size
83
+ self.num_hidden_layers = num_hidden_layers
84
+ self.num_attention_heads = num_attention_heads
85
+ self.intermediate_size = intermediate_size
86
+ self.position_buckets = position_buckets
87
+ self.dropout = (
88
+ hidden_dropout_prob if hidden_dropout_prob is not None else dropout
89
+ )
90
+ self.hidden_dropout_prob = self.dropout
91
+ self.attention_dropout = (
92
+ attention_probs_dropout_prob
93
+ if attention_probs_dropout_prob is not None
94
+ else attention_dropout
95
+ )
96
+ self.attention_probs_dropout_prob = self.attention_dropout
97
+ self.initializer_range = initializer_range
98
+ self.layer_norm_eps = layer_norm_eps
99
+ self.lm_head_gelu_approximate = lm_head_gelu_approximate
100
+ self.shared_relative_embeddings = shared_relative_embeddings
101
+ self.feedforward_dropout_after_projection = (
102
+ feedforward_dropout_after_projection
103
+ )
104
+ self.attention_output_dropout = attention_output_dropout
105
+ self.embedding_padding_idx = embedding_padding_idx
106
+ self.value_gating = value_gating
107
+ self.residual_mixing = residual_mixing
108
+ self.absolute_positions = absolute_positions
109
+ self.use_rope = use_rope
110
+ self.use_alibi = use_alibi
111
+ self.recurrent_steps = recurrent_steps
112
+ self.num_experts = num_experts
113
+ self.experts_per_token = experts_per_token
114
+ self.expert_intermediate_size = expert_intermediate_size
115
+ self.future_offsets = future_offsets or []
116
+ self.state_mixer_kernel = state_mixer_kernel
117
+ self.geometry_lexical_dim = geometry_lexical_dim
118
+ self.geometry_curvature = geometry_curvature
119
+ self.cognitive_readout_layer = cognitive_readout_layer
120
+ self.cognitive_readout_weight = cognitive_readout_weight
121
+ self.direct_sum_dims = direct_sum_dims or []
122
+ self.direct_sum_heads = direct_sum_heads or []
123
+ self.direct_sum_intermediate_sizes = direct_sum_intermediate_sizes or []
124
+ self.lexical_residual_buckets = lexical_residual_buckets
125
+ self.lexical_residual_dim = lexical_residual_dim
126
+ self.lexical_residual_scale = lexical_residual_scale
127
+ self.structured_projection_dim = structured_projection_dim
128
+ self.tie_word_embeddings = tie_word_embeddings
129
+ self.dial_mode = dial_mode
130
+ self.dial_rank = dial_rank
131
+ self.dial_projector_seed = dial_projector_seed
132
+
133
+
134
+ def _valid_tokens(input_ids, attention_mask):
135
+ if attention_mask is None:
136
+ return torch.ones_like(input_ids, dtype=torch.bool)
137
+ return attention_mask.to(torch.bool)
138
+
139
+
140
+ def _bidirectional_mask(valid):
141
+ return valid[:, None, None, :] & valid[:, None, :, None]
142
+
143
+
144
+ def _causal_mask(valid):
145
+ length = valid.size(1)
146
+ causal = torch.ones((length, length), dtype=torch.bool, device=valid.device).tril()
147
+ return _bidirectional_mask(valid) & causal[None, None, :, :]
148
+
149
+
150
+ class RotaryPositionEncoding(nn.Module):
151
+ def __init__(self, head_width, max_length, *, base=10_000.0):
152
+ super().__init__()
153
+ if head_width % 2:
154
+ raise ValueError("RoPE head width must be even")
155
+ inv_freq = 1.0 / (
156
+ base ** (torch.arange(0, head_width, 2, dtype=torch.float32) / head_width)
157
+ )
158
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
159
+ frequencies = self._frequencies(max_length, inv_freq.device)
160
+ self.register_buffer("cos", frequencies.cos(), persistent=False)
161
+ self.register_buffer("sin", frequencies.sin(), persistent=False)
162
+
163
+ def _frequencies(self, length, device):
164
+ positions = torch.arange(length, dtype=torch.float32, device=device)
165
+ return torch.outer(positions, self.inv_freq.to(device=device))
166
+
167
+ def _rotate(self, value):
168
+ length = value.size(-2)
169
+ if length > self.cos.size(0):
170
+ frequencies = self._frequencies(length, value.device)
171
+ self.cos = frequencies.cos()
172
+ self.sin = frequencies.sin()
173
+ even, odd = value[..., 0::2], value[..., 1::2]
174
+ cos = self.cos[:length].to(device=value.device, dtype=value.dtype)
175
+ sin = self.sin[:length].to(device=value.device, dtype=value.dtype)
176
+ rotated = torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1)
177
+ return rotated.flatten(-2)
178
+
179
+ def encode(self, query, key):
180
+ return self._rotate(query), self._rotate(key)
181
+
182
+
183
+ class RelativeLogBucketSelfAttention(nn.Module):
184
+ def __init__(self, config):
185
+ super().__init__()
186
+ self.n_heads = config.num_attention_heads
187
+ self.d_head = config.hidden_size // config.num_attention_heads
188
+ self.max_seq_len = config.max_seq_len
189
+ self.buckets = config.position_buckets
190
+ self.value_gating = config.value_gating
191
+ self.use_rope = config.use_rope
192
+ self.use_alibi = config.use_alibi
193
+ self.shared_relative_embeddings = config.shared_relative_embeddings
194
+ self.attention_output_dropout = config.attention_output_dropout
195
+ if self.use_rope and self.use_alibi:
196
+ raise ValueError("RoPE and ALiBi are mutually exclusive")
197
+ self.qk = nn.Linear(config.hidden_size, 2 * config.hidden_size)
198
+ self.value = nn.Linear(
199
+ config.hidden_size,
200
+ 2 * config.hidden_size if config.value_gating else config.hidden_size,
201
+ )
202
+ self.out = nn.Linear(config.hidden_size, config.hidden_size)
203
+ self.dropout = nn.Dropout(config.attention_dropout)
204
+ self.relative_embedding = (
205
+ None
206
+ if self.use_rope or self.use_alibi or self.shared_relative_embeddings
207
+ else nn.Parameter(
208
+ torch.empty(2 * config.position_buckets - 1, config.hidden_size)
209
+ )
210
+ )
211
+ self.relative_norm = (
212
+ None
213
+ if self.relative_embedding is None
214
+ else nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
215
+ )
216
+ self.value_gate_norm = (
217
+ nn.LayerNorm(
218
+ config.hidden_size,
219
+ eps=config.layer_norm_eps,
220
+ elementwise_affine=False,
221
+ )
222
+ if config.value_gating
223
+ else None
224
+ )
225
+ self.rope = (
226
+ RotaryPositionEncoding(self.d_head, config.max_seq_len)
227
+ if config.use_rope
228
+ else None
229
+ )
230
+ self.scale = 1.0 / math.sqrt(
231
+ self.d_head if config.use_rope or config.use_alibi else 3.0 * self.d_head
232
+ )
233
+ if self.relative_embedding is not None:
234
+ nn.init.trunc_normal_(
235
+ self.relative_embedding,
236
+ mean=0.0,
237
+ std=config.initializer_range,
238
+ a=-2 * config.initializer_range,
239
+ b=2 * config.initializer_range,
240
+ )
241
+ self.register_buffer(
242
+ "position_indices",
243
+ self._position_indices(config.max_seq_len, torch.device("cpu")),
244
+ persistent=False,
245
+ )
246
+ self.register_buffer(
247
+ "alibi_bias",
248
+ self._alibi_bias(config.max_seq_len, torch.device("cpu")),
249
+ persistent=False,
250
+ )
251
+
252
+ def _alibi_bias(self, length, device):
253
+ positions = torch.arange(length, device=device)
254
+ distance = (positions[:, None] - positions[None, :]).abs().float()
255
+ slopes = torch.pow(
256
+ 2.0,
257
+ -8.0
258
+ * (torch.arange(self.n_heads, device=device).float() + 1.0)
259
+ / self.n_heads,
260
+ )
261
+ return -slopes[None, :, None, None] * distance[None, None, :, :]
262
+
263
+ def _position_indices(self, length, device):
264
+ positions = torch.arange(length, device=device)
265
+ relative = positions[:, None] - positions[None, :]
266
+ sign = torch.sign(relative)
267
+ half = self.buckets // 2
268
+ absolute = relative.abs().clamp(max=max(half + 1, self.max_seq_len - 1))
269
+ near = absolute <= half
270
+ safe = absolute.clamp_min(half)
271
+ denominator = math.log(max((self.max_seq_len - 1) / half, 1.0001))
272
+ logged = (
273
+ torch.ceil(torch.log(safe / half) / denominator * (half - 1)).long() + half
274
+ )
275
+ bucketed = torch.where(near, relative, logged * sign)
276
+ return (
277
+ bucketed.long().clamp(-self.buckets + 1, self.buckets - 1)
278
+ + self.buckets
279
+ - 1
280
+ )
281
+
282
+ def forward(self, x, mask, relative_embedding=None):
283
+ batch, length, width = x.shape
284
+ if length > self.position_indices.size(0):
285
+ self.position_indices = self._position_indices(length, x.device)
286
+ q, k = self.qk(x).chunk(2, dim=-1)
287
+ if self.value_gating:
288
+ v, gate = self.value(x).chunk(2, dim=-1)
289
+ gate = F.gelu(gate)
290
+ else:
291
+ v, gate = self.value(x), None
292
+ q = q.view(batch, length, self.n_heads, self.d_head).transpose(1, 2)
293
+ k = k.view(batch, length, self.n_heads, self.d_head).transpose(1, 2)
294
+ v = v.view(batch, length, self.n_heads, self.d_head).transpose(1, 2)
295
+
296
+ if self.rope is not None:
297
+ q, k = self.rope.encode(q, k)
298
+ scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
299
+ elif self.use_alibi:
300
+ if length > self.alibi_bias.size(-1):
301
+ self.alibi_bias = self._alibi_bias(length, x.device)
302
+ scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
303
+ scores = scores + self.alibi_bias[:, :, :length, :length].to(
304
+ device=x.device, dtype=scores.dtype
305
+ )
306
+ else:
307
+ if relative_embedding is None:
308
+ assert self.relative_embedding is not None
309
+ assert self.relative_norm is not None
310
+ relative_embedding = self.relative_norm(self.relative_embedding)
311
+ relative = self.qk(self.dropout(relative_embedding))
312
+ relative = relative[self.position_indices[:length, :length].to(x.device)]
313
+ q_pos, k_pos = relative.chunk(2, dim=-1)
314
+ q_pos = q_pos.view(length, length, self.n_heads, self.d_head)
315
+ k_pos = k_pos.view(length, length, self.n_heads, self.d_head)
316
+ scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale
317
+ scores = scores + torch.einsum("bhqd,qkhd->bhqk", q, k_pos) * self.scale
318
+ scores = scores + torch.einsum("bhkd,qkhd->bhqk", k, q_pos) * self.scale
319
+
320
+ probs = torch.softmax(
321
+ scores.masked_fill(~mask, torch.finfo(scores.dtype).min), dim=-1
322
+ )
323
+ probs = self.dropout(probs) if self.training else probs
324
+ output = (
325
+ torch.matmul(probs, v)
326
+ .transpose(1, 2)
327
+ .contiguous()
328
+ .view(batch, length, width)
329
+ )
330
+ if gate is not None and self.value_gate_norm is not None:
331
+ output = self.value_gate_norm(output * gate)
332
+ output = self.out(output)
333
+ return self.dropout(output) if self.attention_output_dropout else output
334
+
335
+
336
+ class GeGLU(nn.Module):
337
+ def __init__(self, config, width=None):
338
+ super().__init__()
339
+ width = width or config.intermediate_size
340
+ self.up = nn.Linear(config.hidden_size, 2 * width, bias=False)
341
+ self.post_activation_norm = nn.LayerNorm(
342
+ width, eps=config.layer_norm_eps, elementwise_affine=False
343
+ )
344
+ self.down = nn.Linear(width, config.hidden_size, bias=False)
345
+ self.dropout = nn.Dropout(config.dropout)
346
+ self.dropout_after_projection = config.feedforward_dropout_after_projection
347
+
348
+ def forward(self, x):
349
+ value, gate = self.up(x).chunk(2, dim=-1)
350
+ hidden = value * F.gelu(gate, approximate="tanh")
351
+ hidden = self.post_activation_norm(hidden)
352
+ if self.dropout_after_projection:
353
+ return self.dropout(self.down(hidden))
354
+ return self.down(self.dropout(hidden))
355
+
356
+
357
+ class RoutedGeGLU(nn.Module):
358
+ def __init__(self, config):
359
+ super().__init__()
360
+ if not 1 <= config.experts_per_token <= config.num_experts:
361
+ raise ValueError("experts_per_token must be in [1, num_experts]")
362
+ width = config.expert_intermediate_size or max(
363
+ 1, config.intermediate_size // config.num_experts
364
+ )
365
+ self.top_k = config.experts_per_token
366
+ self.router = nn.Linear(config.hidden_size, config.num_experts, bias=False)
367
+ self.experts = nn.ModuleList(
368
+ GeGLU(config, width) for _ in range(config.num_experts)
369
+ )
370
+
371
+ def forward(self, x):
372
+ probabilities = self.router(x).softmax(dim=-1)
373
+ weights, indices = probabilities.topk(self.top_k, dim=-1)
374
+ weights = weights / weights.sum(dim=-1, keepdim=True).clamp_min(1e-8)
375
+ gates = torch.zeros_like(probabilities).scatter(-1, indices, weights)
376
+ outputs = torch.stack([expert(x) for expert in self.experts], dim=-2)
377
+ return (outputs * gates.unsqueeze(-1)).sum(dim=-2)
378
+
379
+
380
+ class CausalStateMixer(nn.Module):
381
+ def __init__(self, config):
382
+ super().__init__()
383
+ kernel = int(config.state_mixer_kernel)
384
+ self.input = nn.Linear(config.hidden_size, 2 * config.hidden_size, bias=False)
385
+ self.state = nn.Conv1d(
386
+ config.hidden_size,
387
+ config.hidden_size,
388
+ kernel,
389
+ groups=config.hidden_size,
390
+ padding=kernel - 1,
391
+ bias=False,
392
+ )
393
+ self.output = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
394
+ self.gate = nn.Parameter(torch.tensor(-2.0))
395
+
396
+ def forward(self, hidden):
397
+ value, gate = self.input(hidden).chunk(2, dim=-1)
398
+ state = self.state(value.transpose(1, 2))[..., : hidden.size(1)].transpose(1, 2)
399
+ return self.output(state * F.silu(gate)) * self.gate.sigmoid()
400
+
401
+
402
+ class DynamicWeightedAverage(nn.Module):
403
+ def __init__(self, n_sublayers):
404
+ super().__init__()
405
+ self.alphas = nn.ParameterList(
406
+ nn.Parameter(torch.cat([torch.zeros(i + 1), torch.ones(1)]))
407
+ for i in range(int(n_sublayers))
408
+ )
409
+ self._states = None
410
+
411
+ def initialize(self, hidden):
412
+ self._states = [hidden]
413
+
414
+ def forward(self, hidden, sublayer_index):
415
+ self._states.append(hidden)
416
+ return torch.tensordot(
417
+ self.alphas[sublayer_index], torch.stack(self._states), dims=1
418
+ )
419
+
420
+
421
+ class GPTBertBlock(nn.Module):
422
+ def __init__(self, config):
423
+ super().__init__()
424
+ self.attention_norm = nn.LayerNorm(
425
+ config.hidden_size,
426
+ eps=config.layer_norm_eps,
427
+ elementwise_affine=False,
428
+ )
429
+ self.attention = RelativeLogBucketSelfAttention(config)
430
+ self.state_mixer = (
431
+ CausalStateMixer(config) if config.state_mixer_kernel else None
432
+ )
433
+ self.feedforward_norm = nn.LayerNorm(
434
+ config.hidden_size,
435
+ eps=config.layer_norm_eps,
436
+ elementwise_affine=False,
437
+ )
438
+ self.feedforward = (
439
+ RoutedGeGLU(config) if config.num_experts > 1 else GeGLU(config)
440
+ )
441
+
442
+ def attend(self, hidden, mask, relative_embedding=None):
443
+ normalized = self.attention_norm(hidden)
444
+ attention = self.attention(normalized, mask, relative_embedding)
445
+ return (
446
+ attention
447
+ if self.state_mixer is None
448
+ else attention + self.state_mixer(normalized)
449
+ )
450
+
451
+ def transform(self, hidden):
452
+ return self.feedforward(self.feedforward_norm(hidden))
453
+
454
+
455
+ class LexicalResidualEmbedding(nn.Module):
456
+ def __init__(self, config):
457
+ super().__init__()
458
+ buckets = int(config.lexical_residual_buckets)
459
+ width = int(config.lexical_residual_dim) or config.hidden_size
460
+ self.buckets = buckets
461
+ self.pad_token_id = config.pad_token_id
462
+ self.scale = float(config.lexical_residual_scale)
463
+ self.embedding = nn.Embedding(buckets, width, padding_idx=0)
464
+ self.projection = (
465
+ nn.Identity()
466
+ if width == config.hidden_size
467
+ else nn.Linear(width, config.hidden_size, bias=False)
468
+ )
469
+ self.register_buffer(
470
+ "word_start_vocab_mask",
471
+ torch.zeros(config.vocab_size, dtype=torch.bool),
472
+ )
473
+
474
+ def forward(self, token_ids):
475
+ batch, length = token_ids.shape
476
+ positions = torch.arange(length, device=token_ids.device).expand(batch, -1)
477
+ starts = self.word_start_vocab_mask[token_ids].clone()
478
+ starts[:, 0] = True
479
+ start_positions = torch.where(starts, positions, 0).cummax(dim=1).values
480
+ relative = positions - start_positions
481
+ ordinal = relative.long() + 1
482
+ mixed = (token_ids.long() + 1) * ordinal * 1_000_003 + ordinal * 97_409
483
+ cumulative = mixed.cumsum(dim=1)
484
+ before = F.pad(cumulative[:, :-1], (1, 0))
485
+ prefix_hashes = cumulative - before.gather(1, start_positions)
486
+ buckets = prefix_hashes.remainder(self.buckets - 1) + 1
487
+ buckets = buckets.masked_fill(token_ids.eq(self.pad_token_id), 0)
488
+ return self.scale * self.projection(self.embedding(buckets))
489
+
490
+
491
+ class GPTBertBackbone(nn.Module):
492
+ def __init__(self, config):
493
+ super().__init__()
494
+ self.config = config
495
+ self.embed_tokens = nn.Embedding(
496
+ config.vocab_size,
497
+ config.hidden_size,
498
+ padding_idx=config.pad_token_id if config.embedding_padding_idx else None,
499
+ )
500
+ self.relative_embedding = (
501
+ nn.Parameter(
502
+ torch.empty(
503
+ 2 * config.position_buckets - 1, config.hidden_size
504
+ )
505
+ )
506
+ if config.shared_relative_embeddings
507
+ else None
508
+ )
509
+ self.relative_norm = (
510
+ nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
511
+ if config.shared_relative_embeddings
512
+ else None
513
+ )
514
+ if self.relative_embedding is not None:
515
+ nn.init.trunc_normal_(
516
+ self.relative_embedding,
517
+ std=config.initializer_range,
518
+ a=-2 * config.initializer_range,
519
+ b=2 * config.initializer_range,
520
+ )
521
+ self.geometry_lexical_dim = int(config.geometry_lexical_dim)
522
+ if not 0 <= self.geometry_lexical_dim < config.hidden_size:
523
+ raise ValueError("geometry_lexical_dim must be in [0, hidden_size)")
524
+ self.geometry_curvature = float(config.geometry_curvature)
525
+ if self.geometry_curvature <= 0:
526
+ raise ValueError("geometry_curvature must be positive")
527
+ self.lexical_angle = None
528
+ self.lexical_radius = None
529
+ if self.geometry_lexical_dim:
530
+ self.lexical_angle = nn.Embedding(
531
+ config.vocab_size, self.geometry_lexical_dim, config.pad_token_id
532
+ )
533
+ self.lexical_radius = nn.Embedding(config.vocab_size, 1, config.pad_token_id)
534
+ self.lexical_residual = (
535
+ LexicalResidualEmbedding(config)
536
+ if config.lexical_residual_buckets
537
+ else None
538
+ )
539
+ self.embed_positions = (
540
+ nn.Embedding(config.max_seq_len, config.hidden_size)
541
+ if getattr(config, "absolute_positions", False)
542
+ else None
543
+ )
544
+ self.embed_norm = nn.LayerNorm(
545
+ config.hidden_size,
546
+ eps=config.layer_norm_eps,
547
+ elementwise_affine=False,
548
+ )
549
+ self.dropout = nn.Dropout(config.dropout)
550
+ self.blocks = nn.ModuleList(
551
+ GPTBertBlock(config) for _ in range(config.num_hidden_layers)
552
+ )
553
+ self.recurrent_steps = max(1, int(config.recurrent_steps))
554
+ self.future_projections = nn.ModuleDict(
555
+ {
556
+ str(offset): nn.Linear(
557
+ config.hidden_size, config.hidden_size, bias=False
558
+ )
559
+ for offset in config.future_offsets
560
+ }
561
+ )
562
+ self.residual_mixer = (
563
+ DynamicWeightedAverage(config.num_hidden_layers * self.recurrent_steps * 2)
564
+ if config.residual_mixing
565
+ else None
566
+ )
567
+ self.cognitive_readout_layer = int(config.cognitive_readout_layer)
568
+ self.cognitive_readout_weight = float(config.cognitive_readout_weight)
569
+ maximum_depth = config.num_hidden_layers * self.recurrent_steps
570
+ if self.cognitive_readout_layer > maximum_depth:
571
+ raise ValueError("cognitive_readout_layer exceeds the executed depth")
572
+
573
+ def lexical_geometry(self, token_ids):
574
+ if self.lexical_angle is None or self.lexical_radius is None:
575
+ raise RuntimeError("lexical geometry is disabled")
576
+ direction = F.normalize(self.lexical_angle(token_ids), dim=-1)
577
+ radius = F.softplus(self.lexical_radius(token_ids)).squeeze(-1)
578
+ scale = math.sqrt(self.geometry_curvature)
579
+ point = torch.tanh(scale * radius / 2).unsqueeze(-1) * direction / scale
580
+ return point, radius
581
+
582
+ def forward(self, input_ids, mask):
583
+ embedded = self.embed_tokens(input_ids)
584
+ if self.lexical_residual is not None:
585
+ embedded = embedded + self.lexical_residual(input_ids)
586
+ if self.geometry_lexical_dim:
587
+ _, radius = self.lexical_geometry(input_ids)
588
+ direction = F.normalize(self.lexical_angle(input_ids), dim=-1)
589
+ tangent = radius.unsqueeze(-1) * direction
590
+ embedded = torch.cat((embedded[..., :-self.geometry_lexical_dim], tangent), -1)
591
+ if self.embed_positions is not None:
592
+ positions = torch.arange(input_ids.size(1), device=input_ids.device)
593
+ embedded = embedded + self.embed_positions(positions)
594
+ hidden = self.dropout(self.embed_norm(embedded))
595
+ mixer = self.residual_mixer
596
+ if mixer is not None:
597
+ mixer.initialize(hidden)
598
+ sublayer = 0
599
+ cognitive_hidden = None
600
+ layer_index = 0
601
+ relative = (
602
+ self.relative_norm(self.relative_embedding)
603
+ if self.relative_norm is not None and self.relative_embedding is not None
604
+ else None
605
+ )
606
+ for _ in range(self.recurrent_steps):
607
+ for block in self.blocks:
608
+ hidden = hidden + block.attend(hidden, mask, relative)
609
+ if mixer is not None:
610
+ hidden = mixer(hidden, sublayer)
611
+ sublayer += 1
612
+ hidden = hidden + block.transform(hidden)
613
+ if mixer is not None:
614
+ hidden = mixer(hidden, sublayer)
615
+ sublayer += 1
616
+ layer_index += 1
617
+ if layer_index == self.cognitive_readout_layer:
618
+ cognitive_hidden = hidden
619
+ if cognitive_hidden is not None and self.cognitive_readout_weight > 0:
620
+ weight = self.cognitive_readout_weight
621
+ hidden = (1.0 - weight) * hidden + weight * cognitive_hidden
622
+ return hidden
623
+
624
+
625
+ class DirectSumStream(nn.Module):
626
+ def __init__(self, config):
627
+ super().__init__()
628
+ self.relative_embedding = (
629
+ nn.Parameter(torch.empty(2 * config.position_buckets - 1, config.hidden_size))
630
+ if config.shared_relative_embeddings
631
+ else None
632
+ )
633
+ self.relative_norm = (
634
+ nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
635
+ if config.shared_relative_embeddings
636
+ else None
637
+ )
638
+ if self.relative_embedding is not None:
639
+ nn.init.trunc_normal_(
640
+ self.relative_embedding,
641
+ std=config.initializer_range,
642
+ a=-2 * config.initializer_range,
643
+ b=2 * config.initializer_range,
644
+ )
645
+ self.embed_positions = (
646
+ nn.Embedding(config.max_seq_len, config.hidden_size)
647
+ if config.absolute_positions
648
+ else None
649
+ )
650
+ self.embed_norm = nn.LayerNorm(
651
+ config.hidden_size,
652
+ eps=config.layer_norm_eps,
653
+ elementwise_affine=False,
654
+ )
655
+ self.dropout = nn.Dropout(config.dropout)
656
+ self.blocks = nn.ModuleList(
657
+ GPTBertBlock(config) for _ in range(config.num_hidden_layers)
658
+ )
659
+ self.recurrent_steps = max(1, int(config.recurrent_steps))
660
+ self.residual_mixer = (
661
+ DynamicWeightedAverage(config.num_hidden_layers * self.recurrent_steps * 2)
662
+ if config.residual_mixing
663
+ else None
664
+ )
665
+
666
+ def forward(self, embedded, mask):
667
+ if self.embed_positions is not None:
668
+ positions = torch.arange(embedded.size(1), device=embedded.device)
669
+ embedded = embedded + self.embed_positions(positions)
670
+ hidden = self.dropout(self.embed_norm(embedded))
671
+ mixer = self.residual_mixer
672
+ if mixer is not None:
673
+ mixer.initialize(hidden)
674
+ relative = (
675
+ self.relative_norm(self.relative_embedding)
676
+ if self.relative_norm is not None and self.relative_embedding is not None
677
+ else None
678
+ )
679
+ sublayer = 0
680
+ for _ in range(self.recurrent_steps):
681
+ for block in self.blocks:
682
+ hidden = hidden + block.attend(hidden, mask, relative)
683
+ if mixer is not None:
684
+ hidden = mixer(hidden, sublayer)
685
+ sublayer += 1
686
+ hidden = hidden + block.transform(hidden)
687
+ if mixer is not None:
688
+ hidden = mixer(hidden, sublayer)
689
+ sublayer += 1
690
+ return hidden
691
+
692
+
693
+ class DirectSumBackbone(nn.Module):
694
+ def __init__(self, config):
695
+ super().__init__()
696
+ dims = tuple(int(value) for value in config.direct_sum_dims)
697
+ heads = tuple(int(value) for value in config.direct_sum_heads)
698
+ widths = tuple(int(value) for value in config.direct_sum_intermediate_sizes)
699
+ if len(dims) != 3 or len(heads) != 3 or len(widths) != 3:
700
+ raise ValueError("direct sum requires three dims, heads and FFN widths")
701
+ if sum(dims) != config.hidden_size:
702
+ raise ValueError("direct_sum_dims must sum to hidden_size")
703
+ if any(dim % head for dim, head in zip(dims, heads)):
704
+ raise ValueError("each direct-sum dimension must divide its head count")
705
+ self.dims = dims
706
+ self.embed_tokens = nn.Embedding(
707
+ config.vocab_size,
708
+ config.hidden_size,
709
+ padding_idx=config.pad_token_id if config.embedding_padding_idx else None,
710
+ )
711
+ streams = []
712
+ for dim, head, width in zip(dims, heads, widths):
713
+ stream_config = copy.copy(config)
714
+ stream_config.hidden_size = dim
715
+ stream_config.num_attention_heads = head
716
+ stream_config.intermediate_size = width
717
+ stream_config.direct_sum_dims = []
718
+ stream_config.direct_sum_heads = []
719
+ stream_config.direct_sum_intermediate_sizes = []
720
+ stream_config.geometry_lexical_dim = 0
721
+ stream_config.future_offsets = []
722
+ stream_config.cognitive_readout_layer = 0
723
+ stream_config.cognitive_readout_weight = 0.0
724
+ streams.append(DirectSumStream(stream_config))
725
+ self.streams = nn.ModuleList(streams)
726
+ self.concept_radius = nn.Embedding(
727
+ config.vocab_size, 1, padding_idx=config.pad_token_id
728
+ )
729
+
730
+ @property
731
+ def factor_slices(self):
732
+ syntax, lexical, conceptual = self.dims
733
+ return (
734
+ slice(0, syntax),
735
+ slice(syntax, syntax + lexical),
736
+ slice(syntax + lexical, syntax + lexical + conceptual),
737
+ )
738
+
739
+ def conceptual_geometry(self, token_ids):
740
+ conceptual = self.embed_tokens(token_ids)[..., self.factor_slices[2]]
741
+ direction = F.normalize(conceptual, dim=-1)
742
+ radius = (1.0 - 1.0e-4) * torch.sigmoid(
743
+ self.concept_radius(token_ids).squeeze(-1)
744
+ )
745
+ if self.concept_radius.padding_idx is not None:
746
+ radius = radius.masked_fill(
747
+ token_ids.eq(self.concept_radius.padding_idx), 0.0
748
+ )
749
+ return radius.unsqueeze(-1) * direction, radius
750
+
751
+ def forward(self, input_ids, mask):
752
+ embedded = self.embed_tokens(input_ids)
753
+ conceptual, _ = self.conceptual_geometry(input_ids)
754
+ parts = list(embedded.split(self.dims, dim=-1))
755
+ parts[2] = conceptual
756
+ return torch.cat(
757
+ [stream(part, mask) for stream, part in zip(self.streams, parts)], dim=-1
758
+ )
759
+
760
+
761
+ def _backbone(config):
762
+ return DirectSumBackbone(config) if config.direct_sum_dims else GPTBertBackbone(config)
763
+
764
+
765
+ def _structured_projections(config):
766
+ width = int(config.structured_projection_dim)
767
+ return (
768
+ nn.ModuleDict(
769
+ {
770
+ "syntax": nn.Linear(config.hidden_size, width, bias=False),
771
+ "lexical": nn.Linear(config.hidden_size, width, bias=False),
772
+ }
773
+ )
774
+ if width
775
+ else nn.ModuleDict()
776
+ )
777
+
778
+
779
+ class GPTBertLMHead(nn.Module):
780
+ def __init__(self, config):
781
+ super().__init__()
782
+ self.norm = nn.LayerNorm(
783
+ config.hidden_size,
784
+ eps=config.layer_norm_eps,
785
+ elementwise_affine=False,
786
+ )
787
+ self.dense = nn.Linear(config.hidden_size, config.hidden_size)
788
+ self.post_norm = nn.LayerNorm(
789
+ config.hidden_size,
790
+ eps=config.layer_norm_eps,
791
+ elementwise_affine=False,
792
+ )
793
+ self.dropout = nn.Dropout(config.dropout)
794
+ self._approximate = config.lm_head_gelu_approximate
795
+ self.bias = nn.Parameter(torch.zeros(config.vocab_size))
796
+ self.weight = nn.Parameter(
797
+ torch.empty(config.vocab_size, config.hidden_size)
798
+ )
799
+
800
+ def forward(self, hidden):
801
+ projected = self.dropout(
802
+ self.post_norm(
803
+ F.gelu(
804
+ self.dense(self.norm(hidden)), approximate=self._approximate
805
+ )
806
+ )
807
+ )
808
+ return F.linear(projected, self.weight, self.bias)
809
+
810
+
811
+ class TOLMModel(PreTrainedModel):
812
+ config_class = TOLMConfig
813
+ base_model_prefix = "tolm"
814
+ _no_split_modules = ["GPTBertBlock"]
815
+
816
+ def __init__(self, config):
817
+ super().__init__(config)
818
+ self.backbone = _backbone(config)
819
+ self.structured_projections = _structured_projections(config)
820
+ self.post_init()
821
+
822
+ def get_input_embeddings(self):
823
+ return self.backbone.embed_tokens
824
+
825
+ def forward(self, input_ids=None, attention_mask=None, **kwargs):
826
+ if input_ids is None:
827
+ raise ValueError("input_ids is required")
828
+ hidden = self.backbone(
829
+ input_ids, _bidirectional_mask(_valid_tokens(input_ids, attention_mask))
830
+ )
831
+ return BaseModelOutput(
832
+ last_hidden_state=hidden, hidden_states=None, attentions=None
833
+ )
834
+
835
+
836
+ class TOLMForMaskedLM(PreTrainedModel):
837
+ config_class = TOLMConfig
838
+ base_model_prefix = "tolm"
839
+ _no_split_modules = ["GPTBertBlock"]
840
+ _tied_weights_keys = ["heads.lm.weight"]
841
+
842
+ def __init__(self, config):
843
+ super().__init__(config)
844
+ self.backbone = _backbone(config)
845
+ head = GPTBertLMHead(config)
846
+ self.heads = nn.ModuleDict({"lm": head})
847
+ self.structured_projections = _structured_projections(config)
848
+ if config.direct_sum_dims:
849
+ self.factor_dual_lambdas = nn.Parameter(
850
+ torch.ones(3), requires_grad=False
851
+ )
852
+ self.post_init()
853
+
854
+ def get_input_embeddings(self):
855
+ return self.backbone.embed_tokens
856
+
857
+ def get_output_embeddings(self):
858
+ return self.heads["lm"]
859
+
860
+ def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
861
+ if input_ids is None:
862
+ raise ValueError("input_ids is required")
863
+ hidden = self.backbone(
864
+ input_ids, _bidirectional_mask(_valid_tokens(input_ids, attention_mask))
865
+ )
866
+ logits = self.heads["lm"](hidden)
867
+ loss = (
868
+ None
869
+ if labels is None
870
+ else F.cross_entropy(
871
+ logits.view(-1, logits.size(-1)), labels.view(-1), ignore_index=-100
872
+ )
873
+ )
874
+ return MaskedLMOutput(
875
+ loss=loss, logits=logits, hidden_states=None, attentions=None
876
+ )
877
+
878
+
879
+ class TOLMForCausalLM(PreTrainedModel):
880
+ config_class = TOLMConfig
881
+ base_model_prefix = "tolm"
882
+ _no_split_modules = ["GPTBertBlock"]
883
+ _tied_weights_keys = ["heads.lm.weight"]
884
+
885
+ def __init__(self, config):
886
+ super().__init__(config)
887
+ self.backbone = _backbone(config)
888
+ head = GPTBertLMHead(config)
889
+ self.heads = nn.ModuleDict({"lm": head})
890
+ self.structured_projections = _structured_projections(config)
891
+ if config.direct_sum_dims:
892
+ self.factor_dual_lambdas = nn.Parameter(
893
+ torch.ones(3), requires_grad=False
894
+ )
895
+ self.post_init()
896
+
897
+ def get_input_embeddings(self):
898
+ return self.backbone.embed_tokens
899
+
900
+ def get_output_embeddings(self):
901
+ return self.heads["lm"]
902
+
903
+ def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **kwargs):
904
+ return {"input_ids": input_ids, "attention_mask": attention_mask}
905
+
906
+ def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
907
+ if input_ids is None:
908
+ raise ValueError("input_ids is required")
909
+ hidden = self.backbone(
910
+ input_ids, _causal_mask(_valid_tokens(input_ids, attention_mask))
911
+ )
912
+ logits = self.heads["lm"](hidden)
913
+ loss = None
914
+ if labels is not None:
915
+ loss = F.cross_entropy(
916
+ logits[:, :-1].contiguous().view(-1, logits.size(-1)),
917
+ labels[:, 1:].contiguous().view(-1),
918
+ ignore_index=-100,
919
+ )
920
+ return CausalLMOutput(
921
+ loss=loss, logits=logits, hidden_states=None, attentions=None
922
+ )
training_manifest.json ADDED
@@ -0,0 +1,311 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "artifact_provenance": {
3
+ "corpus_manifest_sha256": "d065f2065712e9c83d73d3aa58e7453205e4ccbea8dcc232b71a57d54dd17d97",
4
+ "corpus_offsets_sha256": "c0fef973b028cb4927e64182d07328d320aa9984357704051e2a1c8705843e0d",
5
+ "corpus_source_sha256": "25bd27eec0d21e0ec3e33a0467c91f9f3d1d6f78adc25369720cc0289e191bf3",
6
+ "corpus_tokens_sha256": "e5d4fc1f685dbec12ce458e257e4d9bfb39fc084cbbcedc5963229800ee8ac52",
7
+ "corpus_words": 8900000,
8
+ "corpus_words_sha256": "4036ddbe4921b35fb732d9cffb46f968c58ed8463def22bb846a35f5250790c3",
9
+ "dial_manifest_sha256": "5d72f99411ae628ece4a63b79ae77c00e711a6c66439dd20d0d83459c36aaa5e",
10
+ "document_order_sha256": null,
11
+ "factor_manifest_sha256": null,
12
+ "factorization_manifest_sha256": null,
13
+ "influence_preflight_sha256": null,
14
+ "repositories": {
15
+ "dial_priors": {
16
+ "repo": "miguelcsx/dial-lexical-scaffold-priors",
17
+ "repo_type": "dataset",
18
+ "revision": "e2991729096d52b1dbb9c9a2d838646cdb256dca"
19
+ },
20
+ "encoded_corpus": {
21
+ "repo": "miguelcsx/babylm-bbpe16k-512-encoded",
22
+ "repo_type": "dataset",
23
+ "revision": "a62a6f274f3d86c05c8dd967a0c7200573ae9ef6"
24
+ },
25
+ "encoded_factorized_real": {
26
+ "repo": "miguelcsx/factorized-strict-small-real-corpus",
27
+ "repo_type": "dataset",
28
+ "revision": "b86f852224ea8f3c150f395209e447f5223c7aa0"
29
+ },
30
+ "tokenizer": {
31
+ "repo": "miguelcsx/causal-focus-bbpe16k-tokenizer",
32
+ "repo_type": "model",
33
+ "revision": "5c95ae8fe9cc0f745f8d643192775fbceff5fedb"
34
+ },
35
+ "tokenizer_factorized": {
36
+ "repo": "miguelcsx/factorized-strict-small-bpe16k-tokenizer",
37
+ "repo_type": "model",
38
+ "revision": "a1f87b2b439ea988ba2b94b62898ed44780fb5e5"
39
+ }
40
+ },
41
+ "structured_priors_manifest_sha256": null,
42
+ "tokenizer_sha256": "2b8d1b3f51c0d8f276a64ea6a69efa50b4d9899780f1b282b4d0dcb7be5bbb0f",
43
+ "tokenizer_word_start_sha256": "42d9a905c49df4c54dbe3bc93e567c9bbe491ec6fbee5fcbed70e8d1ad533d49"
44
+ },
45
+ "compliance": {
46
+ "auxiliary_induction_words": 1000000,
47
+ "competition_status": "not_submitted_to_leaderboard",
48
+ "conservative_counted_words": 30100000,
49
+ "evaluation_holdout_words": 100000,
50
+ "external_teacher": false,
51
+ "input_only_scaffold": true,
52
+ "model_exposure_words": 30000000,
53
+ "model_training_corpus_words": 8900000,
54
+ "output_embeddings_tied": false,
55
+ "parser_used": false,
56
+ "schedule_total_words": 100000000,
57
+ "shared_transformer_from_layer_one": true,
58
+ "total_exposure_words": 30100000,
59
+ "track": "workshop_only",
60
+ "unique_corpus_words": 10000000,
61
+ "within_100m_conservative_budget": true,
62
+ "within_track_corpus_budget": false,
63
+ "within_track_exposure_budget": false
64
+ },
65
+ "dial": {
66
+ "condition": "shuffled",
67
+ "enabled": true,
68
+ "initial_aux_gradient_ratio": 0.25,
69
+ "pairs_per_microbatch": 128,
70
+ "projector_seed": 20260719,
71
+ "rank": 128,
72
+ "structured_init_fraction": 0.75,
73
+ "temperature": 0.07,
74
+ "withdrawal_words": 10000000
75
+ },
76
+ "factor_priors": {
77
+ "angular_margin": 0.1,
78
+ "dual": {
79
+ "enabled": false,
80
+ "learning_rate": 0.001,
81
+ "max_weight": 1.0,
82
+ "targets": [
83
+ 1.0,
84
+ 4.0,
85
+ 0.1
86
+ ]
87
+ },
88
+ "enabled": false,
89
+ "lambda_conceptual": 0.0,
90
+ "lambda_lexical": 0.0,
91
+ "lambda_syntax": 0.0,
92
+ "max_positions": 256,
93
+ "radial_margin": 0.02,
94
+ "radial_weight": 1.0,
95
+ "randomize": false,
96
+ "randomize_seed": 42,
97
+ "syntax_depth_weight": 0.1,
98
+ "syntax_distance_weight": 0.1,
99
+ "syntax_window": 32,
100
+ "temperature": 0.07,
101
+ "warmup_words": 1000000
102
+ },
103
+ "factorization": {
104
+ "context_window": 4,
105
+ "covariance_weight": 0.005,
106
+ "cross_covariance_weight": 0.002,
107
+ "enabled": false,
108
+ "gradient_budget": 0.1,
109
+ "gradient_ema_momentum": 0.95,
110
+ "graph_mode": "real",
111
+ "loss_enabled": false,
112
+ "max_gradient_scale": 1000.0,
113
+ "null_seed": 271828,
114
+ "pairs_per_microbatch": 128,
115
+ "ramp_words": 5000000,
116
+ "regularizer_words": 512,
117
+ "signal_mode": "online",
118
+ "temperature": 0.07,
119
+ "variance_weight": 0.02
120
+ },
121
+ "geometry": {
122
+ "curvature": 1.0,
123
+ "distance_margin": 0.1,
124
+ "enabled": false,
125
+ "lambda_radial": 0.02,
126
+ "lambda_related": 0.02,
127
+ "max_tokens": 256,
128
+ "radial_margin": 0.05,
129
+ "warmup_words": 5000000
130
+ },
131
+ "model": {
132
+ "absolute_positions": false,
133
+ "attention_dropout": 0.1,
134
+ "bos_token_id": 1,
135
+ "cognitive_readout_layer": 0,
136
+ "cognitive_readout_weight": 0.0,
137
+ "dial_mode": "shuffled",
138
+ "dial_projector_seed": 20260719,
139
+ "dial_rank": 128,
140
+ "direct_sum_dims": [],
141
+ "direct_sum_heads": [],
142
+ "direct_sum_intermediate_sizes": [],
143
+ "dropout": 0.1,
144
+ "eos_token_id": 2,
145
+ "expert_intermediate_size": null,
146
+ "experts_per_token": 1,
147
+ "factor_readout_dims": [
148
+ 128,
149
+ 128,
150
+ 128
151
+ ],
152
+ "factor_readout_mode": "",
153
+ "factor_readout_reflectors": 384,
154
+ "factor_readout_seed": 314159,
155
+ "future_offsets": [],
156
+ "geometry_curvature": 1.0,
157
+ "geometry_lexical_dim": 0,
158
+ "hidden_size": 384,
159
+ "initializer_range": 0.03227486121839514,
160
+ "intermediate_size": 1280,
161
+ "lexical_residual_buckets": 0,
162
+ "lexical_residual_dim": 0,
163
+ "lexical_residual_scale": 1.0,
164
+ "mask_token_id": 4,
165
+ "max_seq_len": 512,
166
+ "num_attention_heads": 6,
167
+ "num_experts": 1,
168
+ "num_hidden_layers": 12,
169
+ "pad_token_id": 3,
170
+ "position_buckets": 32,
171
+ "recurrent_steps": 1,
172
+ "residual_mixing": true,
173
+ "rtd_auxiliary": true,
174
+ "state_mixer_kernel": 0,
175
+ "structured_projection_dim": 0,
176
+ "tie_word_embeddings": false,
177
+ "use_alibi": false,
178
+ "use_rope": false,
179
+ "value_gating": true,
180
+ "vocab_size": 16384
181
+ },
182
+ "release": "TOLM",
183
+ "repository": {
184
+ "commit": "96cacfb55170ac664a5b124d6b6728c6fbc1c2e7",
185
+ "tracked_dirty": false
186
+ },
187
+ "structured_priors": {
188
+ "decay_end_words": 7000000,
189
+ "enabled": false,
190
+ "hold_until_words": 2000000,
191
+ "init_scale": 0.0,
192
+ "lambda_lexical": 0.0,
193
+ "lambda_orth": 0.0,
194
+ "lambda_syntax": 0.0,
195
+ "lexical_dim": 128,
196
+ "max_positions": 256,
197
+ "prior_mode": "contrastive",
198
+ "randomize": false,
199
+ "representation": "slices",
200
+ "require_aligned_artifacts": false,
201
+ "syntax_dim": 128,
202
+ "temperature": 0.07,
203
+ "warmup_words": 100000
204
+ },
205
+ "training": {
206
+ "adaptive_masking": {
207
+ "enabled": false,
208
+ "max_mask_prob": 4.0,
209
+ "min_mask_prob": 0.25,
210
+ "momentum": 0.99
211
+ },
212
+ "batch_size": 16,
213
+ "beta1": 0.9,
214
+ "beta2": 0.98,
215
+ "causal_fraction": 0.0,
216
+ "causal_noise_kind": "random",
217
+ "causal_noise_probability": 0.0,
218
+ "causal_unit": "segment",
219
+ "cooldown_fraction": 0.016,
220
+ "data2vec_layers": 4,
221
+ "data2vec_weight": 0.0,
222
+ "device": "cuda:0",
223
+ "document_curriculum": {
224
+ "enabled": false
225
+ },
226
+ "ema_decay": 0.9998,
227
+ "epsilon": 1e-08,
228
+ "exact_word_checkpoints": true,
229
+ "exposure_words": 30000000,
230
+ "final_lr_ratio": 0.1,
231
+ "frequency_aware_masking": {
232
+ "enabled": false,
233
+ "max_mask_prob": 4.0,
234
+ "min_mask_prob": 0.25,
235
+ "temperature": 1.0
236
+ },
237
+ "future_loss_weight": 0.0,
238
+ "gradient_clip": 2.0,
239
+ "label_smoothing": 0.0,
240
+ "learning_progress": {
241
+ "buckets_per_axis": 4,
242
+ "enabled": false,
243
+ "fast_momentum": 0.9,
244
+ "forgetting_weight": 1.0,
245
+ "slow_momentum": 0.99,
246
+ "uniform_floor": 0.3,
247
+ "window_documents": 4000
248
+ },
249
+ "learning_rate": 0.0035,
250
+ "learning_rate_schedule": "cosine",
251
+ "log_interval": 50,
252
+ "mask_probability_end": 0.5,
253
+ "mask_probability_schedule": "uniform",
254
+ "mask_probability_start": 0.15,
255
+ "mask_replace_probability": 0.8,
256
+ "mask_schedule": "complementary",
257
+ "masked_fraction": 1.0,
258
+ "masking_unit": "word",
259
+ "microbatch_tokens": 8192,
260
+ "mixed_precision": "bf16",
261
+ "num_workers": 0,
262
+ "objective_period": 16,
263
+ "objective_rng_reset_words": [],
264
+ "objective_sanity_window_steps": 512,
265
+ "objective_schedule": "coverage",
266
+ "optimizer": "lamb",
267
+ "packing_strategy": "dense",
268
+ "pin_memory": false,
269
+ "random_replace_probability": 0.1,
270
+ "recombine_probability": 0.0,
271
+ "recombine_strategy": "random",
272
+ "recovery_interval_words": 1000000,
273
+ "resource_memory": {
274
+ "enabled": false
275
+ },
276
+ "router_aux_weight": 0.0,
277
+ "save_steps_words": [
278
+ 1000000,
279
+ 2000000,
280
+ 3000000,
281
+ 4000000,
282
+ 5000000,
283
+ 6000000,
284
+ 7000000,
285
+ 8000000,
286
+ 9000000,
287
+ 10000000,
288
+ 20000000,
289
+ 30000000
290
+ ],
291
+ "schedule_total_words": 100000000,
292
+ "span_max_length": 3,
293
+ "stable_fraction": 0.9,
294
+ "telemetry": {
295
+ "enabled": true,
296
+ "gradient_checkpoints_words": [
297
+ 1000000,
298
+ 10000000,
299
+ 40000000
300
+ ],
301
+ "sampler_trace": true
302
+ },
303
+ "threads_per_process": 10,
304
+ "tokens_per_update": 16384,
305
+ "warmup_fraction": 0.016,
306
+ "weight_decay": 0.1,
307
+ "z_loss_weight": 0.0001
308
+ },
309
+ "variant": "dial_shuffled_s3",
310
+ "words_seen": 30000000
311
+ }