miguelcsx commited on
Commit
16a9a37
·
verified ·
1 Parent(s): 6ff963e

main chck_100M

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