Bochkov commited on
Commit
7daf8d2
·
verified ·
1 Parent(s): 8bb0c8b

Upload model files

Browse files
README.md ADDED
@@ -0,0 +1,320 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ library_name: transformers
3
+ pipeline_tag: text-generation
4
+ language:
5
+ - en
6
+ tags:
7
+ - causal-lm
8
+ - base-model
9
+ - custom-code
10
+ - safetensors
11
+ - research
12
+ - fixed-token-codes
13
+ - frozen-input-representations
14
+ ---
15
+
16
+ # AB-EXT Binary16 — 1.711B parameters, 100B-token target
17
+
18
+ A **base pretrained decoder-only causal language model** released for
19
+ research on trainable input embedding tables and fixed token identities.
20
+
21
+ It is not an instruction-tuned or preference-optimized assistant.
22
+
23
+ ## Research question
24
+
25
+ Can a shared contextual network learn useful language-modeling behavior
26
+ without independently trainable token-specific input vectors?
27
+
28
+ The controlled family contains a learned-input model, a canonical
29
+ 16-bit-code model, and an invertibly recoded GF2 model. They share the
30
+ contextual backbone and output-head architecture, but not the same total
31
+ trainable parameter count.
32
+
33
+ **The results support viability, not performance equivalence.**
34
+ The fixed-code models retain substantial capability while the learned
35
+ model performs better on several informative evaluations.
36
+
37
+ ## Model specification
38
+
39
+ | Property | Value |
40
+ |---|---|
41
+ | Input mode | `binary16` |
42
+ | Trainable parameters | 1,711,376,384 |
43
+ | Trainable input parameters | 0 |
44
+ | Trainable body parameters, excluding input/output | 1,610,713,088 |
45
+ | Trainable untied output-head parameters | 100,663,296 |
46
+ | Persistent input-buffer values | 786,432 |
47
+ | Hidden width | 2048 |
48
+ | Decoder blocks | 24 |
49
+ | Attention heads | 32 |
50
+ | FFN intermediate width | 8192 |
51
+ | Training context | 2048 tokens |
52
+ | Position encoding | RoPE |
53
+ | Normalization / activation | RMSNorm / SwiGLU |
54
+ | Tokenizer source | `HuggingFaceTB/SmolLM2-1.7B` |
55
+ | Exported tokenizer revision | effd688a12921b4cc83e3312b6feb579f70f9c71 |
56
+ | Evaluated training runs for this interface | One |
57
+
58
+ The stored tensor-value count includes buffers and must not be reported
59
+ as the trainable parameter count.
60
+
61
+ ## Input representation
62
+
63
+
64
+ Each token ID is represented by its canonical little-endian binary code:
65
+
66
+ $$
67
+ c(t)_j =
68
+ \left\lfloor \frac{t}{2^j} \right\rfloor \bmod 2,
69
+ \qquad j=0,\ldots,15.
70
+ $$
71
+
72
+ Since the vocabulary contains 49,152 entries, an injective fixed-length
73
+ binary code requires 16 bits.
74
+
75
+ The code is repeated 128 times to width 2048:
76
+
77
+ $$
78
+ x(t)=
79
+ \underbrace{c(t)\Vert\cdots\Vert c(t)}_{128\text{ copies}}.
80
+ $$
81
+
82
+ There are **zero trainable input-interface parameters** and no additional
83
+ trainable input projection before the standard backbone.
84
+ The backbone and the untied output vocabulary projection remain trainable.
85
+
86
+ The evaluated implementation stores the 49,152-by-16 codebook as a
87
+ persistent non-trainable buffer. It is not literally lookup-free.
88
+ “Minimal” describes fixed-length binary identity width, not the storage
89
+ of the complete model or an entropy-optimal token code.
90
+
91
+
92
+
93
+ ## Training
94
+
95
+ - **Training tokens:** Approximately 100B according to the run report; target budget 100,000,000,000 prediction targets.
96
+ - **Training precision:** FP32 parameters with BF16 autocast in the supplied trainer.
97
+ - **Reported recipe:** AdamW; peak learning rate 0.00015; minimum
98
+ scheduled learning rate 0.00001; 2000 warmup steps; cosine decay;
99
+ weight decay 0.01; betas 0.9 and 0.95; gradient clipping 1.0.
100
+ - **Reported launch geometry:** two GPUs per run, microbatch eight per
101
+ GPU, eight accumulation steps, sequence length 2048.
102
+
103
+ The launch geometry corresponds to 262,144 prediction targets per
104
+ optimizer step. Exact final counts must come from the checkpoint,
105
+ not from the requested budget.
106
+
107
+ The supplied sampler selects within-document windows from eligible
108
+ documents of at least 2049 tokens. Sampling can repeat or overlap
109
+ windows; 100B processed targets does not imply 100B unique corpus tokens.
110
+ The original trainer does not fully restore per-rank sampling state
111
+ on resume. A shared recipe alone does not establish identical realized
112
+ sample order across interrupted runs.
113
+
114
+ The model weights were NOT initialized from SmolLM2.
115
+ SmolLM2 supplies tokenizer artifacts, not pretrained model weights.
116
+
117
+ ## Evaluation results
118
+
119
+ These scores are transcribed from the supplied completed evaluation
120
+ summary; the generator does not rerun benchmarks. Raw, unrounded
121
+ harness outputs remain authoritative.
122
+
123
+ Accuracy entries are percentages. Their reported `±` values are
124
+ evaluation standard errors, **not variation across training seeds**.
125
+ Perplexities and bits per byte are not percentages.
126
+
127
+ | Metric | Shots | Result |
128
+ |---|---:|---:|
129
+ | HellaSwag acc (%) | 0 | 40.70 ± 0.49 |
130
+ | HellaSwag acc_norm (%) | 0 | 52.40 ± 0.50 |
131
+ | ARC-Easy acc (%) | 0 | 67.59 ± 0.96 |
132
+ | ARC-Easy acc_norm (%) | 0 | 61.53 ± 1.00 |
133
+ | ARC-Challenge acc (%) | 0 | 32.34 ± 1.37 |
134
+ | ARC-Challenge acc_norm (%) | 0 | 34.04 ± 1.38 |
135
+ | PIQA acc (%) | 0 | 70.51 ± 1.06 |
136
+ | PIQA acc_norm (%) | 0 | 71.11 ± 1.06 |
137
+ | WinoGrande acc (%) | 0 | 55.33 ± 1.40 |
138
+ | OpenBookQA acc (%) | 0 | 28.20 ± 2.01 |
139
+ | OpenBookQA acc_norm (%) | 0 | 38.00 ± 2.17 |
140
+ | CommonsenseQA acc (%) | 0 | 20.56 ± 1.16 |
141
+ | MMLU acc (%) | 0 | 25.88 ± 0.37 |
142
+ | MMLU acc (%; some prompts truncated) | 5 | 25.57 ± 0.37 |
143
+ | LAMBADA accuracy (%) | 0 | 42.75 ± 0.69 |
144
+ | LAMBADA perplexity ↓ | 0 | 17.91 ± 0.62 |
145
+ | WikiText word perplexity ↓ | — | 18.58 |
146
+ | WikiText byte perplexity ↓ | — | 1.73 |
147
+ | WikiText bits/byte ↓ | — | 0.79 |
148
+
149
+ ### Audit and coverage limitations
150
+
151
+ The supplied audit reports identical sample/prompt multisets across
152
+ all six models in each completed task group.
153
+
154
+ For MMLU 5-shot, **1,508 / 56,168 candidate log-likelihood requests**
155
+ were marked as truncated for each model, approximately 2.68%.
156
+ These are candidate requests, not necessarily distinct questions.
157
+ The displayed MMLU 5-shot score therefore includes truncated prompts.
158
+
159
+ No truncations were reported for the other groups by that audit.
160
+ For WikiText rolling likelihood, this does not mean that whole documents
161
+ fit into one model context: rolling windowing is part of scoring.
162
+
163
+ <details>
164
+ <summary>Full six-model comparison</summary>
165
+
166
+ | Metric | AB-EXT Learned | AB-EXT Binary16 | AB-EXT GF2 | SmolLM2-135M | SmolLM2-360M | SmolLM2-1.7B |
167
+ |---|---:|---:|---:|---:|---:|---:|
168
+ | HellaSwag acc (%); shots=0 | 44.21 ± 0.50 | 40.70 ± 0.49 | 40.24 ± 0.49 | 35.36 ± 0.48 | 43.05 ± 0.49 | 53.38 ± 0.50 |
169
+ | HellaSwag acc_norm (%); shots=0 | 57.79 ± 0.49 | 52.40 ± 0.50 | 51.44 ± 0.50 | 43.02 ± 0.49 | 56.28 ± 0.50 | 71.43 ± 0.45 |
170
+ | ARC-Easy acc (%); shots=0 | 71.63 ± 0.92 | 67.59 ± 0.96 | 66.84 ± 0.97 | 64.44 ± 0.98 | 70.24 ± 0.94 | 77.86 ± 0.85 |
171
+ | ARC-Easy acc_norm (%); shots=0 | 66.04 ± 0.97 | 61.53 ± 1.00 | 60.73 ± 1.00 | 58.75 ± 1.01 | 68.18 ± 0.96 | 73.36 ± 0.91 |
172
+ | ARC-Challenge acc (%); shots=0 | 35.92 ± 1.40 | 32.34 ± 1.37 | 30.55 ± 1.35 | 28.07 ± 1.31 | 36.26 ± 1.40 | 44.37 ± 1.45 |
173
+ | ARC-Challenge acc_norm (%); shots=0 | 37.63 ± 1.42 | 34.04 ± 1.38 | 34.22 ± 1.39 | 29.61 ± 1.33 | 38.05 ± 1.42 | 47.27 ± 1.46 |
174
+ | PIQA acc (%); shots=0 | 72.69 ± 1.04 | 70.51 ± 1.06 | 71.16 ± 1.06 | 68.44 ± 1.08 | 71.38 ± 1.05 | 76.99 ± 0.98 |
175
+ | PIQA acc_norm (%); shots=0 | 72.14 ± 1.05 | 71.11 ± 1.06 | 72.14 ± 1.05 | 68.39 ± 1.08 | 71.82 ± 1.05 | 77.20 ± 0.98 |
176
+ | WinoGrande acc (%); shots=0 | 58.56 ± 1.38 | 55.33 ± 1.40 | 55.01 ± 1.40 | 52.57 ± 1.40 | 59.35 ± 1.38 | 65.98 ± 1.33 |
177
+ | OpenBookQA acc (%); shots=0 | 27.60 ± 2.00 | 28.20 ± 2.01 | 25.20 ± 1.94 | 22.00 ± 1.85 | 24.80 ± 1.93 | 32.20 ± 2.09 |
178
+ | OpenBookQA acc_norm (%); shots=0 | 37.80 ± 2.17 | 38.00 ± 2.17 | 36.80 ± 2.16 | 32.60 ± 2.10 | 37.80 ± 2.17 | 44.40 ± 2.22 |
179
+ | CommonsenseQA acc (%); shots=0 | 19.82 ± 1.14 | 20.56 ± 1.16 | 19.74 ± 1.14 | 19.90 ± 1.14 | 21.05 ± 1.17 | 41.69 ± 1.41 |
180
+ | MMLU acc (%); shots=0 | 25.32 ± 0.37 | 25.88 ± 0.37 | 26.11 ± 0.37 | 24.25 ± 0.36 | 25.47 ± 0.37 | 48.40 ± 0.41 |
181
+ | MMLU acc (%; some prompts truncated); shots=5 | 25.48 ± 0.37 | 25.57 ± 0.37 | 24.66 ± 0.36 | 25.15 ± 0.36 | 25.03 ± 0.37 | 50.06 ± 0.41 |
182
+ | LAMBADA accuracy (%); shots=0 | 47.72 ± 0.70 | 42.75 ± 0.69 | 42.29 ± 0.69 | 42.97 ± 0.69 | 53.31 ± 0.70 | 67.51 ± 0.65 |
183
+ | LAMBADA perplexity ↓; shots=0 | 12.88 ± 0.42 | 17.91 ± 0.62 | 18.47 ± 0.63 | 19.06 ± 0.63 | 9.38 ± 0.27 | 4.44 ± 0.10 |
184
+ | WikiText word perplexity ↓; shots=— | 16.50 | 18.58 | 19.03 | 23.14 | 17.12 | 11.62 |
185
+ | WikiText byte perplexity ↓; shots=— | 1.69 | 1.73 | 1.73 | 1.80 | 1.70 | 1.58 |
186
+ | WikiText bits/byte ↓; shots=— | 0.76 | 0.79 | 0.79 | 0.85 | 0.77 | 0.66 |
187
+
188
+ </details>
189
+
190
+ ### How to interpret SmolLM2 comparisons
191
+
192
+ All scores above are from the supplied local evaluation summary, not
193
+ copied leaderboard scores.
194
+
195
+ The SmolLM2 technical report gives approximate training budgets of:
196
+
197
+ | External reference | Published budget | Relative to 100B |
198
+ |---|---:|---:|
199
+ | SmolLM2-135M | 2T tokens | 20× |
200
+ | SmolLM2-360M | 4T tokens | 40× |
201
+ | SmolLM2-1.7B | 11T tokens | 110× |
202
+
203
+ Source: https://arxiv.org/abs/2502.02737
204
+
205
+ These models differ in architecture, size, data, training schedule,
206
+ and compute. They are quality references, **not matched controls** and
207
+ not proof of a sample-efficiency advantage.
208
+
209
+ ## Usage
210
+
211
+ Review the custom Python files before enabling `trust_remote_code=True`.
212
+ Use a tested Transformers version and pin the Hub revision for
213
+ reproducible deployment.
214
+
215
+ ```python
216
+ import torch
217
+ from transformers import AutoTokenizer, AutoModelForCausalLM
218
+
219
+ model_id = 'E6E831728/ab_ext_binary16'
220
+ # For published Hub use, pin revision to a reviewed commit.
221
+ revision = None
222
+
223
+ tokenizer = AutoTokenizer.from_pretrained(
224
+ model_id,
225
+ revision=revision,
226
+ )
227
+
228
+ model = AutoModelForCausalLM.from_pretrained(
229
+ model_id,
230
+ revision=revision,
231
+ trust_remote_code=True,
232
+ dtype=torch.bfloat16,
233
+ ).to("cuda").eval()
234
+
235
+ inputs = tokenizer(
236
+ "Gravity is",
237
+ return_tensors="pt",
238
+ add_special_tokens=False,
239
+ return_attention_mask=True,
240
+ ).to("cuda")
241
+
242
+ pad_id = tokenizer.pad_token_id
243
+ if pad_id is None:
244
+ pad_id = tokenizer.eos_token_id
245
+
246
+ with torch.inference_mode():
247
+ output = model.generate(
248
+ input_ids=inputs["input_ids"],
249
+ attention_mask=inputs["attention_mask"],
250
+ max_new_tokens=32,
251
+ do_sample=False,
252
+ use_cache=False,
253
+ eos_token_id=tokenizer.eos_token_id,
254
+ pad_token_id=pad_id,
255
+ )
256
+
257
+ print(tokenizer.decode(output[0], skip_special_tokens=True))
258
+ ```
259
+
260
+ The implementation does not provide a KV cache.
261
+ The trained context is 2048 tokens; the supplied generation adapter
262
+ uses a sliding window when the context grows beyond its limit.
263
+ This is not evidence of trained long-context capability.
264
+
265
+ ### Loss API
266
+
267
+ The original training model consumes already-shifted targets.
268
+ The HF runtime is intended to expose the usual causal-LM convention
269
+ with an internal label shift. Do not pass already-shifted labels to
270
+ such a runtime.
271
+
272
+ Before fine-tuning, verify the actual runtime's loss implementation.
273
+ Forward-logit equivalence does not by itself test label conventions.
274
+
275
+ ## Verification and integrity
276
+
277
+ The supplied verification logs report:
278
+
279
+ - successful BF16 loading and generation for all three releases;
280
+ - exactly matching original/exported forward logits on four short
281
+ prompts for each model, in the tested verification configuration.
282
+
283
+ These are smoke and implementation-parity checks, not an exhaustive
284
+ test across padding, context lengths, dtypes, or generation modes.
285
+
286
+ - Weight file: `model.safetensors`
287
+ - Weight SHA-256: `c3e78f24f03dff36b6174985dc5189e8ce37a52685d022ade36fa68c3239801b`
288
+ - Stored tensor values: 1,712,162,816
289
+ - Stored values by dtype: `{"F32": 1712162816}`
290
+ - Trainable parameter count: 1,711,376,384
291
+ - Persistent input-buffer values: 786,432
292
+
293
+ This card update does not modify the weights, tokenizer, model code,
294
+ or configuration.
295
+
296
+ ## Limitations and intended use
297
+
298
+ - Research use and text completion; not a validated high-stakes assistant.
299
+ - One evaluated training run per input interface at this scale.
300
+ - Fixed-code and learned-input models are backbone-matched, not
301
+ total-parameter-matched.
302
+ - No measured runtime or energy advantage is established by parameter
303
+ counts alone.
304
+ - One GF2 recoding does not establish invariance to arbitrary codes.
305
+ - The output vocabulary matrix remains trainable and token-specific.
306
+ - Input-code structure is not fitted to the pretraining objective, but
307
+ the tokenizer and its ID assignment can contain corpus-derived structure.
308
+ - Benchmark contamination has not been independently certified absent.
309
+ - Generated text can be false, biased, or harmful.
310
+
311
+ - Zero/one coding maps token ID zero to a zero input vector; the
312
+ zero-offset GF2 transform preserves it. In the supplied bias-free
313
+ architecture, a context made entirely of zero-code tokens gives
314
+ uniform logits. This does not apply to arbitrary contexts ending
315
+ in that token.
316
+
317
+
318
+ ## Attribution and licensing
319
+
320
+ Tokenizer artifacts are sourced from `HuggingFaceTB/SmolLM2-1.7B`.
config.json ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "AttnExtForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "auto_map": {
7
+ "AutoConfig": "configuration_attn_ext.AttnExtConfig",
8
+ "AutoModelForCausalLM": "modeling_attn_ext.AttnExtForCausalLM"
9
+ },
10
+ "binary_dim": 16,
11
+ "binary_encoding": "zero_one",
12
+ "binary_repeat": 128,
13
+ "binary_scale": 1.0,
14
+ "block_size": 2048,
15
+ "bos_token_id": 0,
16
+ "code_seed": 12345,
17
+ "d_model": 2048,
18
+ "dropout": 0.0,
19
+ "eos_token_id": 0,
20
+ "ffn_multiplier": 4.0,
21
+ "head_dim": 64,
22
+ "hidden_size": 2048,
23
+ "initializer_range": 0.02,
24
+ "input_mode": "binary16",
25
+ "is_decoder": true,
26
+ "is_encoder_decoder": false,
27
+ "max_position_embeddings": 2048,
28
+ "min_col_weight": 4,
29
+ "min_row_weight": 4,
30
+ "mlp_bias": false,
31
+ "model_type": "attn_ext",
32
+ "multiple_of": 256,
33
+ "n_head": 32,
34
+ "n_layer": 24,
35
+ "num_attention_heads": 32,
36
+ "num_hidden_layers": 24,
37
+ "pad_token_id": null,
38
+ "rms_norm_eps": 1e-05,
39
+ "rope_theta": 10000.0,
40
+ "tie_word_embeddings": false,
41
+ "transformers_version": null,
42
+ "use_cache": false,
43
+ "vocab_size": 49152
44
+ }
configuration_attn_ext.py ADDED
@@ -0,0 +1,113 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import PretrainedConfig
2
+
3
+
4
+ class AttnExtConfig(PretrainedConfig):
5
+ model_type = "attn_ext"
6
+ keys_to_ignore_at_inference = ["past_key_values"]
7
+
8
+ def __init__(
9
+ self,
10
+ vocab_size=49152,
11
+ d_model=2048,
12
+ n_layer=24,
13
+ n_head=32,
14
+ ffn_multiplier=4.0,
15
+ multiple_of=256,
16
+ block_size=2048,
17
+ rope_theta=10000.0,
18
+ dropout=0.0,
19
+ rms_norm_eps=1e-5,
20
+ initializer_range=0.02,
21
+ attention_bias=False,
22
+ mlp_bias=False,
23
+ input_mode="learned",
24
+ binary_dim=16,
25
+ binary_encoding="zero_one",
26
+ binary_scale=1.0,
27
+ code_seed=12345,
28
+ min_row_weight=4,
29
+ min_col_weight=4,
30
+ pad_token_id=None,
31
+ bos_token_id=None,
32
+ eos_token_id=None,
33
+ tie_word_embeddings=False,
34
+ use_cache=False,
35
+ **kwargs,
36
+ ):
37
+ super().__init__(
38
+ pad_token_id=pad_token_id,
39
+ bos_token_id=bos_token_id,
40
+ eos_token_id=eos_token_id,
41
+ tie_word_embeddings=tie_word_embeddings,
42
+ **kwargs,
43
+ )
44
+
45
+ if d_model % n_head != 0:
46
+ raise ValueError("d_model must be divisible by n_head")
47
+
48
+ head_dim = d_model // n_head
49
+ if head_dim % 2 != 0:
50
+ raise ValueError("RoPE requires an even head dimension")
51
+
52
+ if input_mode not in {"learned", "binary16", "gf2"}:
53
+ raise ValueError(
54
+ "input_mode must be learned, binary16, or gf2"
55
+ )
56
+
57
+ if input_mode != "learned":
58
+ if binary_dim != 16:
59
+ raise ValueError("Frozen-code models require binary_dim=16")
60
+ if vocab_size > 2**binary_dim:
61
+ raise ValueError("Vocabulary does not fit in 16 bits")
62
+ if d_model % binary_dim != 0:
63
+ raise ValueError(
64
+ "d_model must be divisible by binary_dim"
65
+ )
66
+ if tie_word_embeddings:
67
+ raise ValueError(
68
+ "Frozen input codes cannot be tied to lm_head"
69
+ )
70
+
71
+ if binary_encoding not in {"zero_one", "bipolar"}:
72
+ raise ValueError(
73
+ "binary_encoding must be zero_one or bipolar"
74
+ )
75
+
76
+ self.vocab_size = vocab_size
77
+
78
+ self.d_model = d_model
79
+ self.hidden_size = d_model
80
+
81
+ self.n_layer = n_layer
82
+ self.num_hidden_layers = n_layer
83
+
84
+ self.n_head = n_head
85
+ self.num_attention_heads = n_head
86
+ self.head_dim = head_dim
87
+
88
+ self.ffn_multiplier = ffn_multiplier
89
+ self.multiple_of = multiple_of
90
+
91
+ self.block_size = block_size
92
+ self.max_position_embeddings = block_size
93
+ self.rope_theta = rope_theta
94
+
95
+ self.dropout = dropout
96
+ self.rms_norm_eps = rms_norm_eps
97
+ self.initializer_range = initializer_range
98
+ self.attention_bias = attention_bias
99
+ self.mlp_bias = mlp_bias
100
+
101
+ self.input_mode = input_mode
102
+ self.binary_dim = binary_dim
103
+ self.binary_encoding = binary_encoding
104
+ self.binary_scale = binary_scale
105
+ self.binary_repeat = d_model // binary_dim
106
+
107
+ self.code_seed = code_seed
108
+ self.min_row_weight = min_row_weight
109
+ self.min_col_weight = min_col_weight
110
+
111
+ self.use_cache = use_cache
112
+ self.is_decoder = True
113
+ self.is_encoder_decoder = False
generation_config.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token_id": 0,
3
+ "do_sample": false,
4
+ "eos_token_id": 0,
5
+ "pad_token_id": null,
6
+ "transformers_version": null,
7
+ "use_cache": false
8
+ }
merges.txt ADDED
The diff for this file is too large to render. See raw diff
 
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c3e78f24f03dff36b6174985dc5189e8ce37a52685d022ade36fa68c3239801b
3
+ size 6848674776
modeling_attn_ext.py ADDED
@@ -0,0 +1,733 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from typing import Optional
3
+
4
+ import torch
5
+ import torch.utils.checkpoint
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+
9
+ from transformers import PreTrainedModel
10
+ from transformers.generation import GenerationMixin
11
+ from transformers.modeling_outputs import CausalLMOutputWithPast
12
+
13
+ from .configuration_attn_ext import AttnExtConfig
14
+
15
+
16
+ def round_up(value: int, multiple: int) -> int:
17
+ return multiple * math.ceil(value / multiple)
18
+
19
+
20
+ class RMSNorm(nn.Module):
21
+ def __init__(self, dim: int, eps: float):
22
+ super().__init__()
23
+ self.weight = nn.Parameter(torch.ones(dim))
24
+ self.eps = eps
25
+
26
+ def forward(self, x):
27
+ dtype = x.dtype
28
+ xf = x.float()
29
+ xf = xf * torch.rsqrt(
30
+ xf.pow(2).mean(dim=-1, keepdim=True) + self.eps
31
+ )
32
+ return (xf * self.weight.float()).to(dtype)
33
+
34
+
35
+ def rotate_half(x):
36
+ x1 = x[..., ::2]
37
+ x2 = x[..., 1::2]
38
+ return torch.stack((-x2, x1), dim=-1).flatten(-2)
39
+
40
+
41
+ class RotaryEmbedding(nn.Module):
42
+ def __init__(self, dim, max_position, theta):
43
+ super().__init__()
44
+
45
+ inv_freq = 1.0 / (
46
+ theta
47
+ ** (
48
+ torch.arange(0, dim, 2, dtype=torch.float32)
49
+ / dim
50
+ )
51
+ )
52
+
53
+ positions = torch.arange(
54
+ max_position,
55
+ dtype=torch.float32,
56
+ )
57
+
58
+ frequencies = torch.outer(positions, inv_freq)
59
+ embedding = torch.repeat_interleave(
60
+ frequencies,
61
+ repeats=2,
62
+ dim=-1,
63
+ )
64
+
65
+ self.register_buffer(
66
+ "cos_cached",
67
+ embedding.cos(),
68
+ persistent=False,
69
+ )
70
+ self.register_buffer(
71
+ "sin_cached",
72
+ embedding.sin(),
73
+ persistent=False,
74
+ )
75
+
76
+ def forward(self, q, k, position_ids=None):
77
+ sequence_length = q.shape[-2]
78
+
79
+ if position_ids is None:
80
+ cos = self.cos_cached[:sequence_length][
81
+ None, None, :, :
82
+ ]
83
+ sin = self.sin_cached[:sequence_length][
84
+ None, None, :, :
85
+ ]
86
+ else:
87
+ cos = self.cos_cached[position_ids][:, None, :, :]
88
+ sin = self.sin_cached[position_ids][:, None, :, :]
89
+
90
+ cos = cos.to(device=q.device, dtype=q.dtype)
91
+ sin = sin.to(device=q.device, dtype=q.dtype)
92
+
93
+ q = q * cos + rotate_half(q) * sin
94
+ k = k * cos + rotate_half(k) * sin
95
+ return q, k
96
+
97
+
98
+ class CausalSelfAttention(nn.Module):
99
+ def __init__(self, config):
100
+ super().__init__()
101
+
102
+ self.d_model = config.d_model
103
+ self.n_head = config.n_head
104
+ self.head_dim = config.head_dim
105
+ self.dropout_p = config.dropout
106
+
107
+ self.q_proj = nn.Linear(
108
+ config.d_model,
109
+ config.d_model,
110
+ bias=config.attention_bias,
111
+ )
112
+ self.k_proj = nn.Linear(
113
+ config.d_model,
114
+ config.d_model,
115
+ bias=config.attention_bias,
116
+ )
117
+ self.v_proj = nn.Linear(
118
+ config.d_model,
119
+ config.d_model,
120
+ bias=config.attention_bias,
121
+ )
122
+ self.o_proj = nn.Linear(
123
+ config.d_model,
124
+ config.d_model,
125
+ bias=config.attention_bias,
126
+ )
127
+
128
+ self.rope = RotaryEmbedding(
129
+ config.head_dim,
130
+ config.block_size,
131
+ config.rope_theta,
132
+ )
133
+
134
+ def forward(
135
+ self,
136
+ x,
137
+ attention_mask=None,
138
+ position_ids=None,
139
+ ):
140
+ batch_size, sequence_length, channels = x.shape
141
+
142
+ q = self.q_proj(x).view(
143
+ batch_size,
144
+ sequence_length,
145
+ self.n_head,
146
+ self.head_dim,
147
+ ).transpose(1, 2)
148
+
149
+ k = self.k_proj(x).view(
150
+ batch_size,
151
+ sequence_length,
152
+ self.n_head,
153
+ self.head_dim,
154
+ ).transpose(1, 2)
155
+
156
+ v = self.v_proj(x).view(
157
+ batch_size,
158
+ sequence_length,
159
+ self.n_head,
160
+ self.head_dim,
161
+ ).transpose(1, 2)
162
+
163
+ q, k = self.rope(
164
+ q,
165
+ k,
166
+ position_ids=position_ids,
167
+ )
168
+
169
+ dropout_p = self.dropout_p if self.training else 0.0
170
+
171
+ if attention_mask is None or bool(attention_mask.all()):
172
+ output = F.scaled_dot_product_attention(
173
+ q,
174
+ k,
175
+ v,
176
+ attn_mask=None,
177
+ dropout_p=dropout_p,
178
+ is_causal=True,
179
+ )
180
+ else:
181
+ if attention_mask.shape != (
182
+ batch_size,
183
+ sequence_length,
184
+ ):
185
+ raise ValueError(
186
+ "attention_mask must have shape "
187
+ f"{(batch_size, sequence_length)}"
188
+ )
189
+
190
+ causal = torch.ones(
191
+ sequence_length,
192
+ sequence_length,
193
+ device=x.device,
194
+ dtype=torch.bool,
195
+ ).tril()
196
+
197
+ allowed = (
198
+ causal[None, None, :, :]
199
+ & attention_mask[:, None, None, :].bool()
200
+ )
201
+
202
+ output = F.scaled_dot_product_attention(
203
+ q,
204
+ k,
205
+ v,
206
+ attn_mask=allowed,
207
+ dropout_p=dropout_p,
208
+ is_causal=False,
209
+ )
210
+
211
+ output = output.transpose(1, 2).contiguous().view(
212
+ batch_size,
213
+ sequence_length,
214
+ channels,
215
+ )
216
+
217
+ return self.o_proj(output)
218
+
219
+
220
+ class SwiGLU(nn.Module):
221
+ def __init__(self, config):
222
+ super().__init__()
223
+
224
+ hidden_dim = round_up(
225
+ int(config.ffn_multiplier * config.d_model),
226
+ config.multiple_of,
227
+ )
228
+
229
+ self.gate_proj = nn.Linear(
230
+ config.d_model,
231
+ hidden_dim,
232
+ bias=config.mlp_bias,
233
+ )
234
+ self.up_proj = nn.Linear(
235
+ config.d_model,
236
+ hidden_dim,
237
+ bias=config.mlp_bias,
238
+ )
239
+ self.down_proj = nn.Linear(
240
+ hidden_dim,
241
+ config.d_model,
242
+ bias=config.mlp_bias,
243
+ )
244
+ self.dropout = nn.Dropout(config.dropout)
245
+
246
+ def forward(self, x):
247
+ x = F.silu(self.gate_proj(x)) * self.up_proj(x)
248
+ return self.dropout(self.down_proj(x))
249
+
250
+
251
+ class TransformerBlock(nn.Module):
252
+ def __init__(self, config):
253
+ super().__init__()
254
+
255
+ self.input_norm = RMSNorm(
256
+ config.d_model,
257
+ config.rms_norm_eps,
258
+ )
259
+ self.post_attention_norm = RMSNorm(
260
+ config.d_model,
261
+ config.rms_norm_eps,
262
+ )
263
+
264
+ self.attention = CausalSelfAttention(config)
265
+ self.mlp = SwiGLU(config)
266
+
267
+ def forward(
268
+ self,
269
+ x,
270
+ attention_mask=None,
271
+ position_ids=None,
272
+ ):
273
+ x = x + self.attention(
274
+ self.input_norm(x),
275
+ attention_mask=attention_mask,
276
+ position_ids=position_ids,
277
+ )
278
+
279
+ x = x + self.mlp(
280
+ self.post_attention_norm(x)
281
+ )
282
+
283
+ return x
284
+
285
+
286
+ def canonical_binary_codebook(
287
+ vocab_size,
288
+ bits,
289
+ encoding,
290
+ ):
291
+ token_ids = torch.arange(
292
+ vocab_size,
293
+ dtype=torch.int64,
294
+ )
295
+ shifts = torch.arange(
296
+ bits,
297
+ dtype=torch.int64,
298
+ )
299
+
300
+ codebook = (
301
+ (token_ids[:, None] >> shifts[None, :]) & 1
302
+ ).to(torch.float32)
303
+
304
+ if encoding == "bipolar":
305
+ codebook = codebook.mul(2.0).sub(1.0)
306
+
307
+ return codebook.contiguous()
308
+
309
+
310
+ def gf2_rank(matrix):
311
+ matrix = matrix.detach().cpu().to(
312
+ torch.uint8
313
+ ).clone()
314
+ matrix &= 1
315
+
316
+ rows, columns = matrix.shape
317
+ rank = 0
318
+
319
+ for column in range(columns):
320
+ pivot = None
321
+
322
+ for row in range(rank, rows):
323
+ if int(matrix[row, column]) == 1:
324
+ pivot = row
325
+ break
326
+
327
+ if pivot is None:
328
+ continue
329
+
330
+ if pivot != rank:
331
+ temporary = matrix[rank].clone()
332
+ matrix[rank] = matrix[pivot]
333
+ matrix[pivot] = temporary
334
+
335
+ for row in range(rows):
336
+ if row != rank and int(
337
+ matrix[row, column]
338
+ ) == 1:
339
+ matrix[row] ^= matrix[rank]
340
+
341
+ rank += 1
342
+
343
+ if rank == rows:
344
+ break
345
+
346
+ return rank
347
+
348
+
349
+ def make_invertible_gf2_matrix(
350
+ bits,
351
+ seed,
352
+ min_row_weight,
353
+ min_col_weight,
354
+ ):
355
+ generator = torch.Generator(device="cpu")
356
+ generator.manual_seed(seed)
357
+
358
+ for _ in range(1_000_000):
359
+ matrix = torch.randint(
360
+ 0,
361
+ 2,
362
+ (bits, bits),
363
+ generator=generator,
364
+ dtype=torch.uint8,
365
+ )
366
+
367
+ if bool(
368
+ torch.any(
369
+ matrix.sum(dim=1) < min_row_weight
370
+ )
371
+ ):
372
+ continue
373
+
374
+ if bool(
375
+ torch.any(
376
+ matrix.sum(dim=0) < min_col_weight
377
+ )
378
+ ):
379
+ continue
380
+
381
+ if gf2_rank(matrix) == bits:
382
+ return matrix.contiguous()
383
+
384
+ raise RuntimeError(
385
+ "Could not construct an invertible GF(2) matrix"
386
+ )
387
+
388
+
389
+ def gf2_binary_codebook(config):
390
+ source = canonical_binary_codebook(
391
+ config.vocab_size,
392
+ config.binary_dim,
393
+ "zero_one",
394
+ ).to(torch.uint8)
395
+
396
+ matrix = make_invertible_gf2_matrix(
397
+ bits=config.binary_dim,
398
+ seed=config.code_seed,
399
+ min_row_weight=config.min_row_weight,
400
+ min_col_weight=config.min_col_weight,
401
+ )
402
+
403
+ shift = torch.zeros(
404
+ config.binary_dim,
405
+ dtype=torch.uint8,
406
+ )
407
+
408
+ codebook = (
409
+ source.to(torch.int16)
410
+ @ matrix.to(torch.int16).T
411
+ ).remainder(2).to(torch.uint8)
412
+
413
+ codebook = codebook ^ shift
414
+
415
+ if config.binary_encoding == "bipolar":
416
+ codebook = (
417
+ codebook.float().mul(2.0).sub(1.0)
418
+ )
419
+ else:
420
+ codebook = codebook.float()
421
+
422
+ return (
423
+ codebook.contiguous(),
424
+ matrix.contiguous(),
425
+ shift.contiguous(),
426
+ )
427
+
428
+
429
+ class FixedBinaryEmbedding(nn.Module):
430
+ def __init__(self, config):
431
+ super().__init__()
432
+
433
+ if config.input_mode == "binary16":
434
+ codebook = canonical_binary_codebook(
435
+ config.vocab_size,
436
+ config.binary_dim,
437
+ config.binary_encoding,
438
+ )
439
+ matrix = None
440
+ shift = None
441
+
442
+ elif config.input_mode == "gf2":
443
+ codebook, matrix, shift = (
444
+ gf2_binary_codebook(config)
445
+ )
446
+
447
+ else:
448
+ raise ValueError(
449
+ "FixedBinaryEmbedding requires a "
450
+ "frozen-code input mode"
451
+ )
452
+
453
+ self.register_buffer(
454
+ "codebook",
455
+ codebook,
456
+ persistent=True,
457
+ )
458
+
459
+ if matrix is not None:
460
+ self.register_buffer(
461
+ "A_gf2",
462
+ matrix,
463
+ persistent=True,
464
+ )
465
+ self.register_buffer(
466
+ "b_gf2",
467
+ shift,
468
+ persistent=True,
469
+ )
470
+
471
+ self.repeat = config.binary_repeat
472
+ self.binary_scale = config.binary_scale
473
+
474
+ @property
475
+ def weight(self):
476
+ return self.codebook
477
+
478
+ def forward(self, input_ids):
479
+ code = self.codebook[input_ids.long()]
480
+
481
+ output = code.repeat(
482
+ *([1] * (code.ndim - 1)),
483
+ self.repeat,
484
+ )
485
+
486
+ if self.binary_scale != 1.0:
487
+ output = output * self.binary_scale
488
+
489
+ return output
490
+
491
+
492
+ class AttnExtPreTrainedModel(PreTrainedModel):
493
+ config_class = AttnExtConfig
494
+ base_model_prefix = "attn_ext"
495
+ supports_gradient_checkpointing = True
496
+ _supports_sdpa = True
497
+ _no_split_modules = ["TransformerBlock"]
498
+
499
+ def _init_weights(self, module):
500
+ if isinstance(module, nn.Linear):
501
+ nn.init.normal_(
502
+ module.weight,
503
+ mean=0.0,
504
+ std=self.config.initializer_range,
505
+ )
506
+ if module.bias is not None:
507
+ nn.init.zeros_(module.bias)
508
+
509
+ elif isinstance(module, nn.Embedding):
510
+ nn.init.normal_(
511
+ module.weight,
512
+ mean=0.0,
513
+ std=self.config.initializer_range,
514
+ )
515
+
516
+
517
+ class AttnExtForCausalLM(
518
+ AttnExtPreTrainedModel,
519
+ GenerationMixin,
520
+ ):
521
+ main_input_name = "input_ids"
522
+
523
+ def __init__(self, config):
524
+ super().__init__(config)
525
+
526
+ if config.input_mode == "learned":
527
+ self.token_embeddings = nn.Embedding(
528
+ config.vocab_size,
529
+ config.d_model,
530
+ )
531
+ else:
532
+ self.token_embeddings = FixedBinaryEmbedding(
533
+ config
534
+ )
535
+
536
+ self.layers = nn.ModuleList(
537
+ [
538
+ TransformerBlock(config)
539
+ for _ in range(config.n_layer)
540
+ ]
541
+ )
542
+
543
+ self.final_norm = RMSNorm(
544
+ config.d_model,
545
+ config.rms_norm_eps,
546
+ )
547
+
548
+ self.lm_head = nn.Linear(
549
+ config.d_model,
550
+ config.vocab_size,
551
+ bias=False,
552
+ )
553
+
554
+ self.gradient_checkpointing = False
555
+ self.post_init()
556
+
557
+ residual_std = (
558
+ config.initializer_range
559
+ / math.sqrt(2 * config.n_layer)
560
+ )
561
+
562
+ for layer in self.layers:
563
+ nn.init.normal_(
564
+ layer.attention.o_proj.weight,
565
+ mean=0.0,
566
+ std=residual_std,
567
+ )
568
+ nn.init.normal_(
569
+ layer.mlp.down_proj.weight,
570
+ mean=0.0,
571
+ std=residual_std,
572
+ )
573
+
574
+ def get_input_embeddings(self):
575
+ return self.token_embeddings
576
+
577
+ def set_input_embeddings(self, value):
578
+ if self.config.input_mode != "learned":
579
+ raise RuntimeError(
580
+ "Frozen input codes cannot be replaced "
581
+ "through set_input_embeddings"
582
+ )
583
+ self.token_embeddings = value
584
+
585
+ def get_output_embeddings(self):
586
+ return self.lm_head
587
+
588
+ def set_output_embeddings(self, value):
589
+ self.lm_head = value
590
+
591
+ def prepare_inputs_for_generation(
592
+ self,
593
+ input_ids,
594
+ attention_mask=None,
595
+ **kwargs,
596
+ ):
597
+ if input_ids.shape[1] > self.config.block_size:
598
+ input_ids = input_ids[
599
+ :, -self.config.block_size:
600
+ ]
601
+
602
+ if attention_mask is not None:
603
+ attention_mask = attention_mask[
604
+ :, -self.config.block_size:
605
+ ]
606
+
607
+ position_ids = None
608
+
609
+ if attention_mask is not None:
610
+ position_ids = (
611
+ attention_mask.long().cumsum(-1) - 1
612
+ )
613
+ position_ids.masked_fill_(
614
+ attention_mask == 0,
615
+ 0,
616
+ )
617
+
618
+ return {
619
+ "input_ids": input_ids,
620
+ "attention_mask": attention_mask,
621
+ "position_ids": position_ids,
622
+ "use_cache": False,
623
+ }
624
+
625
+ def forward(
626
+ self,
627
+ input_ids=None,
628
+ attention_mask=None,
629
+ labels=None,
630
+ position_ids=None,
631
+ inputs_embeds=None,
632
+ use_cache=None,
633
+ return_dict=None,
634
+ **kwargs,
635
+ ):
636
+ if input_ids is None and inputs_embeds is None:
637
+ raise ValueError(
638
+ "input_ids or inputs_embeds is required"
639
+ )
640
+
641
+ if inputs_embeds is not None:
642
+ x = inputs_embeds
643
+ batch_size, sequence_length, _ = x.shape
644
+ else:
645
+ batch_size, sequence_length = input_ids.shape
646
+ x = self.token_embeddings(input_ids)
647
+
648
+ if sequence_length > self.config.block_size:
649
+ raise ValueError(
650
+ f"Sequence length {sequence_length} exceeds "
651
+ f"block_size={self.config.block_size}"
652
+ )
653
+
654
+ if attention_mask is not None:
655
+ expected = (batch_size, sequence_length)
656
+ if attention_mask.shape != expected:
657
+ raise ValueError(
658
+ f"attention_mask must have shape {expected}"
659
+ )
660
+
661
+ # HF_EXPORT_INPUT_DTYPE_FIX
662
+ # Frozen floating-point buffers may remain FP32 after loading.
663
+ # Match the residual stream to the backbone parameter dtype.
664
+ x = x.to(dtype=self.layers[0].attention.q_proj.weight.dtype)
665
+
666
+ for layer in self.layers:
667
+ if self.gradient_checkpointing and self.training:
668
+ def custom_forward(hidden_states, current_layer=layer):
669
+ return current_layer(
670
+ hidden_states,
671
+ attention_mask=attention_mask,
672
+ position_ids=position_ids,
673
+ )
674
+
675
+ x = torch.utils.checkpoint.checkpoint(
676
+ custom_forward,
677
+ x,
678
+ use_reentrant=False,
679
+ )
680
+ else:
681
+ x = layer(
682
+ x,
683
+ attention_mask=attention_mask,
684
+ position_ids=position_ids,
685
+ )
686
+
687
+ x = self.final_norm(x)
688
+ logits = self.lm_head(x)
689
+
690
+ loss = None
691
+
692
+ if labels is not None:
693
+ if labels.shape != (
694
+ batch_size,
695
+ sequence_length,
696
+ ):
697
+ raise ValueError(
698
+ "labels must have the same shape as input_ids"
699
+ )
700
+
701
+ shift_logits = logits[:, :-1, :].contiguous()
702
+ shift_labels = labels[:, 1:].contiguous().clone()
703
+
704
+ if attention_mask is not None:
705
+ shift_labels.masked_fill_(
706
+ attention_mask[:, 1:].eq(0),
707
+ -100,
708
+ )
709
+
710
+ loss = F.cross_entropy(
711
+ shift_logits.float().view(
712
+ -1,
713
+ self.config.vocab_size,
714
+ ),
715
+ shift_labels.view(-1),
716
+ ignore_index=-100,
717
+ )
718
+
719
+ return_dict = (
720
+ self.config.use_return_dict
721
+ if return_dict is None
722
+ else return_dict
723
+ )
724
+
725
+ if not return_dict:
726
+ output = (logits,)
727
+ return ((loss,) + output) if loss is not None else output
728
+
729
+ return CausalLMOutputWithPast(
730
+ loss=loss,
731
+ logits=logits,
732
+ past_key_values=None,
733
+ )
special_tokens_map.json ADDED
@@ -0,0 +1,42 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ "<|endoftext|>",
4
+ "<|im_start|>",
5
+ "<|im_end|>",
6
+ "<repo_name>",
7
+ "<reponame>",
8
+ "<file_sep>",
9
+ "<filename>",
10
+ "<gh_stars>",
11
+ "<issue_start>",
12
+ "<issue_comment>",
13
+ "<issue_closed>",
14
+ "<jupyter_start>",
15
+ "<jupyter_text>",
16
+ "<jupyter_code>",
17
+ "<jupyter_output>",
18
+ "<jupyter_script>",
19
+ "<empty_output>"
20
+ ],
21
+ "bos_token": {
22
+ "content": "<|endoftext|>",
23
+ "lstrip": false,
24
+ "normalized": false,
25
+ "rstrip": false,
26
+ "single_word": false
27
+ },
28
+ "eos_token": {
29
+ "content": "<|endoftext|>",
30
+ "lstrip": false,
31
+ "normalized": false,
32
+ "rstrip": false,
33
+ "single_word": false
34
+ },
35
+ "unk_token": {
36
+ "content": "<|endoftext|>",
37
+ "lstrip": false,
38
+ "normalized": false,
39
+ "rstrip": false,
40
+ "single_word": false
41
+ }
42
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,168 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "added_tokens_decoder": {
4
+ "0": {
5
+ "content": "<|endoftext|>",
6
+ "lstrip": false,
7
+ "normalized": false,
8
+ "rstrip": false,
9
+ "single_word": false,
10
+ "special": true
11
+ },
12
+ "1": {
13
+ "content": "<|im_start|>",
14
+ "lstrip": false,
15
+ "normalized": false,
16
+ "rstrip": false,
17
+ "single_word": false,
18
+ "special": true
19
+ },
20
+ "2": {
21
+ "content": "<|im_end|>",
22
+ "lstrip": false,
23
+ "normalized": false,
24
+ "rstrip": false,
25
+ "single_word": false,
26
+ "special": true
27
+ },
28
+ "3": {
29
+ "content": "<repo_name>",
30
+ "lstrip": false,
31
+ "normalized": false,
32
+ "rstrip": false,
33
+ "single_word": false,
34
+ "special": true
35
+ },
36
+ "4": {
37
+ "content": "<reponame>",
38
+ "lstrip": false,
39
+ "normalized": false,
40
+ "rstrip": false,
41
+ "single_word": false,
42
+ "special": true
43
+ },
44
+ "5": {
45
+ "content": "<file_sep>",
46
+ "lstrip": false,
47
+ "normalized": false,
48
+ "rstrip": false,
49
+ "single_word": false,
50
+ "special": true
51
+ },
52
+ "6": {
53
+ "content": "<filename>",
54
+ "lstrip": false,
55
+ "normalized": false,
56
+ "rstrip": false,
57
+ "single_word": false,
58
+ "special": true
59
+ },
60
+ "7": {
61
+ "content": "<gh_stars>",
62
+ "lstrip": false,
63
+ "normalized": false,
64
+ "rstrip": false,
65
+ "single_word": false,
66
+ "special": true
67
+ },
68
+ "8": {
69
+ "content": "<issue_start>",
70
+ "lstrip": false,
71
+ "normalized": false,
72
+ "rstrip": false,
73
+ "single_word": false,
74
+ "special": true
75
+ },
76
+ "9": {
77
+ "content": "<issue_comment>",
78
+ "lstrip": false,
79
+ "normalized": false,
80
+ "rstrip": false,
81
+ "single_word": false,
82
+ "special": true
83
+ },
84
+ "10": {
85
+ "content": "<issue_closed>",
86
+ "lstrip": false,
87
+ "normalized": false,
88
+ "rstrip": false,
89
+ "single_word": false,
90
+ "special": true
91
+ },
92
+ "11": {
93
+ "content": "<jupyter_start>",
94
+ "lstrip": false,
95
+ "normalized": false,
96
+ "rstrip": false,
97
+ "single_word": false,
98
+ "special": true
99
+ },
100
+ "12": {
101
+ "content": "<jupyter_text>",
102
+ "lstrip": false,
103
+ "normalized": false,
104
+ "rstrip": false,
105
+ "single_word": false,
106
+ "special": true
107
+ },
108
+ "13": {
109
+ "content": "<jupyter_code>",
110
+ "lstrip": false,
111
+ "normalized": false,
112
+ "rstrip": false,
113
+ "single_word": false,
114
+ "special": true
115
+ },
116
+ "14": {
117
+ "content": "<jupyter_output>",
118
+ "lstrip": false,
119
+ "normalized": false,
120
+ "rstrip": false,
121
+ "single_word": false,
122
+ "special": true
123
+ },
124
+ "15": {
125
+ "content": "<jupyter_script>",
126
+ "lstrip": false,
127
+ "normalized": false,
128
+ "rstrip": false,
129
+ "single_word": false,
130
+ "special": true
131
+ },
132
+ "16": {
133
+ "content": "<empty_output>",
134
+ "lstrip": false,
135
+ "normalized": false,
136
+ "rstrip": false,
137
+ "single_word": false,
138
+ "special": true
139
+ }
140
+ },
141
+ "additional_special_tokens": [
142
+ "<|endoftext|>",
143
+ "<|im_start|>",
144
+ "<|im_end|>",
145
+ "<repo_name>",
146
+ "<reponame>",
147
+ "<file_sep>",
148
+ "<filename>",
149
+ "<gh_stars>",
150
+ "<issue_start>",
151
+ "<issue_comment>",
152
+ "<issue_closed>",
153
+ "<jupyter_start>",
154
+ "<jupyter_text>",
155
+ "<jupyter_code>",
156
+ "<jupyter_output>",
157
+ "<jupyter_script>",
158
+ "<empty_output>"
159
+ ],
160
+ "bos_token": "<|endoftext|>",
161
+ "clean_up_tokenization_spaces": false,
162
+ "eos_token": "<|endoftext|>",
163
+ "extra_special_tokens": {},
164
+ "model_max_length": 8192,
165
+ "tokenizer_class": "GPT2Tokenizer",
166
+ "unk_token": "<|endoftext|>",
167
+ "vocab_size": 49152
168
+ }
vocab.json ADDED
The diff for this file is too large to render. See raw diff