anandselvadurai-techus commited on
Commit
4bbdc93
·
verified ·
1 Parent(s): f02380f

Initial Commit - Tokle

Browse files
LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2026 tech.us
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
README.md CHANGED
@@ -1,3 +1,134 @@
1
  ---
 
 
2
  license: mit
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ language:
3
+ - en
4
  license: mit
5
+ library_name: transformers
6
+ pipeline_tag: text-generation
7
+ tags:
8
+ - text-generation
9
+ - causal-lm
10
+ - custom-architecture
11
+ - custom_code
12
+ - slm
13
+ - small-language-model
14
+ - spab
15
+ datasets:
16
+ - HuggingFaceFW/fineweb-edu
17
+ - HuggingFaceTB/cosmopedia
18
+ - agentlans/high-quality-english-sentences
19
+ - nampdn-ai/tiny-strange-textbooks
20
+ - armanc/ScienceQA
21
+ - nvidia/OpenMathInstruct-2
22
+ - microsoft/orca-math-word-problems-200k
23
  ---
24
+
25
+ # Tokle-3M
26
+
27
+ ## Model Summary
28
+
29
+ Tokle-3M is a decoder-only language model with 2.91M trainable parameters,
30
+ trained on 12B tokens. Its main architectural addition is SPAB (Static Pairwise
31
+ Attention Bias), a frozen table of token-pair association scores built from
32
+ Pointwise Mutual Information (PMI) over the training corpus and added to the
33
+ attention logits of the first layer.
34
+
35
+ For every query-key pair, SPAB hashes the two token IDs into the table, pulls
36
+ out their PMI value, multiplies it by a learned per-head scale, and adds it to
37
+ the attention logits before softmax. The bias ignores position and depends only
38
+ on which tokens are involved, so the model starts training already knowing
39
+ which tokens tend to co-occur. It only has to learn how much to trust that
40
+ prior.
41
+
42
+ ## Model Architecture
43
+
44
+ | Parameter | Value |
45
+ | --- | --- |
46
+ | Architecture | Custom decoder-only transformer + SPAB (`TokleForCausalLM`) |
47
+ | Layers | 9 |
48
+ | Hidden size (d_model) | 144 |
49
+ | Attention heads | 3 |
50
+ | KV heads (GQA) | 1 (multi-query attention) |
51
+ | Head dim | 48 |
52
+ | FFN intermediate size | 432 |
53
+ | Max sequence length | 512 |
54
+ | Trainable parameters | 2,908,947 |
55
+ | Frozen SPAB table | 8,388,608 (float32 buffer) |
56
+
57
+ ## How to use
58
+
59
+ This model uses a custom architecture, so it needs `trust_remote_code=True`.
60
+
61
+ ```python
62
+ import torch
63
+ from transformers import AutoModelForCausalLM, AutoTokenizer
64
+
65
+ model_id = "techdotus/Tokle-3M"
66
+ tok = AutoTokenizer.from_pretrained(model_id)
67
+ model = AutoModelForCausalLM.from_pretrained(model_id, trust_remote_code=True).eval()
68
+
69
+ ids = tok("The capital of France is", return_tensors="pt")
70
+ with torch.no_grad():
71
+ out = model.generate(**ids, max_new_tokens=32, do_sample=False) # greedy
72
+ print(tok.decode(out[0], skip_special_tokens=True))
73
+ ```
74
+
75
+ ## Benchmark Results
76
+
77
+ All scores are 0-shot acc_norm, using the Open SLM Leaderboard methodology.
78
+ Scores for the other models are from the Open SLM Leaderboard. **Bold** marks
79
+ the best result in each column.
80
+
81
+ | Model | Org | Params | Int Index | HellaSwag | ARC-Easy | ARC-Chal | PIQA | ArithMark-3 |
82
+ | --- | --- | --- | --- | --- | --- | --- | --- | --- |
83
+ | **Tokle-3M** | tech.us | 2.9M | **9.16** | 27.22% | **34.68%** | **24.49%** | 54.95% | **41.70%** |
84
+ | Ember-2 | SurjoLabs | 2.96Mx2 | 7.21 | 27.28% | 33.42% | 22.01% | 55.11% | 35.90% |
85
+ | BananaMind-2-Micro | BananaMind | 2.9M | 6.01 | **28.27%** | 33.12% | 21.93% | 53.21% | 34.00% |
86
+ | GPT-S-1.4M | Axiomic Labs | 1.4M | 5.40 | 26.89% | 31.57% | 21.93% | **55.17%** | 30.20% |
87
+
88
+ ## Training Data Details
89
+
90
+ We trained on a curated mixture with a strict cleaning pipeline that also
91
+ removed topics not useful for a model of this size.
92
+
93
+ | Source | Percentage |
94
+ | --- | --- |
95
+ | FineWeb-Edu | 43.1% |
96
+ | Cosmopedia | 24.3% |
97
+ | OpenMathInstruct-2 | 13.5% |
98
+ | Tiny Strange Textbooks | 9.0% |
99
+ | MegaScience (medicine & biology, custom curated) | 5.0% |
100
+ | High-Quality English Sentences | 3.0% |
101
+ | ScienceQA | 1.2% |
102
+ | Orca-Math Word Problems 200k | 0.9% |
103
+ | Total | 100% |
104
+
105
+ - **Tokenizer:** all data was tokenized with the model's 5,048-token BPE
106
+ tokenizer, and 1% was held out for validation.
107
+ - **Blending:** sources were blended per dataset using the weights above.
108
+
109
+ ## Limitations
110
+
111
+ - **Tiny model:** with ~2.9M trainable parameters and 144-dim hidden states,
112
+ generations are often repetitive, incoherent or factually wrong. The model is
113
+ a research artifact for studying small-scale LMs, not an assistant.
114
+ - **Short context:** 512 tokens maximum. RoPE tables are not built beyond that
115
+ length.
116
+ - **English only:** trained on English web, educational, synthetic and math text.
117
+ - **Not instruction-tuned or safety-aligned:** it may reproduce biases present
118
+ in web data.
119
+
120
+ ## Licenses
121
+
122
+ Model weights and code: MIT.
123
+
124
+ ## Citation
125
+
126
+ ```bibtex
127
+ @misc{tokle2026,
128
+ title = {{Tokle-3M}: Pointwise Mutual Information as an Inductive Bias for Self-Attention},
129
+ author = {{Tech.us Team}},
130
+ year = {2026},
131
+ publisher = {Hugging Face},
132
+ howpublished = {\url{https://huggingface.co/techdotus/Tokle-3M}}
133
+ }
134
+ ```
config.json ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "TokleForCausalLM"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_tokle.TokleConfig",
7
+ "AutoModel": "modeling_tokle.TokleModel",
8
+ "AutoModelForCausalLM": "modeling_tokle.TokleForCausalLM"
9
+ },
10
+ "bos_token_id": 0,
11
+ "d_model": 144,
12
+ "dtype": "float32",
13
+ "eos_token_id": 1,
14
+ "ffn_hidden": 432,
15
+ "head_dim": 48,
16
+ "mask_padded_vocab_logits": true,
17
+ "max_seq_len": 512,
18
+ "model_type": "tokle",
19
+ "n_head": 3,
20
+ "n_kv_head": 1,
21
+ "n_layer": 9,
22
+ "norm_eps": 1e-05,
23
+ "pad_token_id": 2,
24
+ "real_vocab_size": 5048,
25
+ "rope_theta": 10000.0,
26
+ "spab_enabled": true,
27
+ "spab_init_scale": 0.1,
28
+ "spab_table_size": 8388608,
29
+ "transformers_version": "4.57.6",
30
+ "use_cache": false,
31
+ "vocab_size": 5056
32
+ }
configuration_tokle.py ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tokle config. Field names follow the training script (n_layer, d_model, ...);
2
+ attribute_map aliases them to the names Transformers and lm-eval look for.
3
+ """
4
+ from transformers.configuration_utils import PretrainedConfig
5
+
6
+
7
+ class TokleConfig(PretrainedConfig):
8
+ """Decoder-only LM with SPAB (Static Pairwise Attention Bias).
9
+
10
+ SPAB is a token-pair-indexed additive bias on the first layer's attention
11
+ logits. Values come from windowed PMI over the training corpus, hashed into
12
+ one flat table shared by every head and frozen before step 0 -- the only
13
+ thing trained here is spab scale, one number per head.
14
+
15
+ Two fields are easy to trip over. vocab_size (5056) is the padded embedding
16
+ matrix; real_vocab_size (5048) is what the tokenizer can actually emit, and
17
+ everything between them is untrained filler. use_cache is False because SPAB
18
+ hashes every id in the window, so there is nothing to cache incrementally.
19
+ """
20
+
21
+ model_type = "tokle"
22
+ keys_to_ignore_at_inference = ["past_key_values"]
23
+
24
+ attribute_map = {
25
+ "num_hidden_layers": "n_layer",
26
+ "hidden_size": "d_model",
27
+ "num_attention_heads": "n_head",
28
+ "num_key_value_heads": "n_kv_head",
29
+ "intermediate_size": "ffn_hidden",
30
+ "max_position_embeddings": "max_seq_len",
31
+ "rms_norm_eps": "norm_eps",
32
+ }
33
+
34
+ def __init__(
35
+ self,
36
+ vocab_size=5056,
37
+ real_vocab_size=5048,
38
+ max_seq_len=512,
39
+ n_layer=9,
40
+ d_model=144,
41
+ n_head=3,
42
+ n_kv_head=1,
43
+ head_dim=48,
44
+ ffn_hidden=432,
45
+ rope_theta=10000.0,
46
+ norm_eps=1e-5,
47
+ spab_enabled=True,
48
+ spab_table_size=8388608,
49
+ spab_init_scale=0.1,
50
+ mask_padded_vocab_logits=True,
51
+ tie_word_embeddings=True,
52
+ bos_token_id=0,
53
+ eos_token_id=1,
54
+ pad_token_id=2,
55
+ use_cache=False,
56
+ **kwargs,
57
+ ):
58
+ self.vocab_size = vocab_size
59
+ self.real_vocab_size = real_vocab_size
60
+ self.max_seq_len = max_seq_len
61
+ self.n_layer = n_layer
62
+ self.d_model = d_model
63
+ self.n_head = n_head
64
+ self.n_kv_head = n_kv_head
65
+ self.head_dim = head_dim
66
+ self.ffn_hidden = ffn_hidden
67
+ self.rope_theta = rope_theta
68
+ self.norm_eps = norm_eps
69
+ self.spab_enabled = spab_enabled
70
+ self.spab_table_size = spab_table_size
71
+ self.spab_init_scale = spab_init_scale
72
+ self.mask_padded_vocab_logits = mask_padded_vocab_logits
73
+ self.use_cache = use_cache # see class docstring: no KV cache with SPAB
74
+ super().__init__(
75
+ tie_word_embeddings=tie_word_embeddings,
76
+ bos_token_id=bos_token_id,
77
+ eos_token_id=eos_token_id,
78
+ pad_token_id=pad_token_id,
79
+ **kwargs,
80
+ )
generation_config.json ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token_id": 0,
3
+ "do_sample": true,
4
+ "eos_token_id": 1,
5
+ "max_new_tokens": 64,
6
+ "pad_token_id": 2,
7
+ "temperature": 0.8,
8
+ "top_p": 0.95,
9
+ "transformers_version": "4.57.6",
10
+ "use_cache": false
11
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:082e3ae8de93edcbb4790e4917d089354ac79adc4acb937d41e3277ef8f59ec5
3
+ size 45201124
modeling_tokle.py ADDED
@@ -0,0 +1,379 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tokle: decoder-only transformer (RMSNorm, RoPE, GQA, SwiGLU) plus SPAB.
2
+
3
+ Math is unchanged from the training-time model.py -- same norms, same half-split
4
+ RoPE, same hash -- so an exported checkpoint scores identically. Only the module
5
+ names were changed to the usual Transformers layout.
6
+
7
+ Works on Transformers 4.x and 5.x. The v5 differences are marked inline; they
8
+ are not cosmetic, see the note above _M1.
9
+ """
10
+ import math
11
+ from typing import Optional, Tuple, Union
12
+
13
+ import torch
14
+ import torch.nn as nn
15
+ import torch.nn.functional as F
16
+ from transformers import __version__ as _transformers_version
17
+ from transformers.generation import GenerationMixin
18
+ from transformers.modeling_outputs import BaseModelOutput, CausalLMOutput
19
+ from transformers.modeling_utils import PreTrainedModel
20
+
21
+ try: # as a Transformers dynamic module (trust_remote_code)
22
+ from .configuration_tokle import TokleConfig
23
+ except ImportError: # running the files straight out of a checkout
24
+ from configuration_tokle import TokleConfig
25
+
26
+
27
+ # ----------------------------------------------------------------------------- norms and rope
28
+ class TokleRMSNorm(nn.Module):
29
+ def __init__(self, d, eps=1e-5):
30
+ super().__init__()
31
+ self.eps = eps
32
+ self.weight = nn.Parameter(torch.ones(d))
33
+
34
+ def forward(self, x):
35
+ # Always reduce in fp32: in fp16 the mean of squares underflows on the
36
+ # long tail and the norm comes back slightly wrong.
37
+ dt = x.dtype
38
+ x = x.float()
39
+ x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
40
+ return (x * self.weight.float()).to(dt)
41
+
42
+ def extra_repr(self):
43
+ return f"{tuple(self.weight.shape)}, eps={self.eps}"
44
+
45
+
46
+ def build_rope(head_dim, max_seq_len, theta, device=None):
47
+ inv = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32,
48
+ device=device) / head_dim))
49
+ t = torch.arange(max_seq_len, dtype=torch.float32, device=device)
50
+ f = torch.outer(t, inv)
51
+ return torch.cos(f), torch.sin(f)
52
+
53
+
54
+ def apply_rope(x, cos, sin):
55
+ """Half-split rotation (first half against second), not the interleaved
56
+ variant. Training used this one; swapping it silently degrades the model."""
57
+ b, h, t, d = x.shape
58
+ x1, x2 = x.float().chunk(2, dim=-1)
59
+ c = cos[:t].view(1, 1, t, d // 2).to(x1.device)
60
+ s = sin[:t].view(1, 1, t, d // 2).to(x1.device)
61
+ return torch.cat([x1 * c - x2 * s, x1 * s + x2 * c], dim=-1).to(x.dtype)
62
+
63
+
64
+ # ----------------------------------------------------------------------------- spab
65
+ # Mixing constants for the pair hash. Plain ints, not buffers, and that matters:
66
+ # buffers here would be non-persistent, and Transformers 5 builds the model on
67
+ # the meta device and fills only what the checkpoint contains. Non-persistent
68
+ # buffers come back as uninitialised memory -- in practice zeros, which sends
69
+ # every token pair to slot 0 and quietly flattens the whole bias. Loads fine,
70
+ # generates garbage. Same reason the rope tables below are not buffers either.
71
+ _M1 = -7046029254386353131
72
+ _M2 = -4417276706812531889
73
+ _M3 = -4658895280553007687
74
+
75
+
76
+ class SPABBias(nn.Module):
77
+ """Static Pairwise Attention Bias, keyed by (source id, target id).
78
+
79
+ Hashing the pair into one flat table costs a single lookup and keeps storage
80
+ fixed -- a real vocab^2 bias matrix would be 25M entries here and would grow
81
+ with the vocabulary. The table holds windowed PMI from the training corpus
82
+ and never moves; `scale` (one per head) is the only trained tensor.
83
+
84
+ Position plays no part in the lookup, which is what sets SPAB apart from
85
+ ALiBi or T5-style biases.
86
+ """
87
+
88
+ is_spab = True
89
+
90
+ def __init__(self, config: TokleConfig):
91
+ super().__init__()
92
+ self.enabled = config.spab_enabled
93
+ self.table_size = config.spab_table_size
94
+ self.n_head = config.n_head
95
+ init = config.spab_init_scale if config.spab_enabled else 0.0
96
+ self.scale = nn.Parameter(torch.full((config.n_head,), init))
97
+ if not config.spab_enabled:
98
+ self.scale.requires_grad_(False)
99
+ self.register_buffer("table",
100
+ torch.zeros(config.spab_table_size, dtype=torch.float32),
101
+ persistent=True)
102
+
103
+ def hash(self, i, j):
104
+ # Murmur-style mix. Relies on int64 overflow wrapping, so keep it in
105
+ # int64 throughout -- the masks after each shift are what make the
106
+ # result reproducible across devices.
107
+ h = i.to(torch.int64) * _M1 + j.to(torch.int64) * _M2
108
+ h = h ^ ((h >> 29) & 0x7FFFFFFFF)
109
+ h = h * _M3
110
+ h = h ^ ((h >> 32) & 0xFFFFFFFF)
111
+ return (h & 0x7FFFFFFFFFFFFFFF) % self.table_size
112
+
113
+ def forward(self, idx):
114
+ if not self.enabled:
115
+ return None
116
+ # Materialises b*t*t int64 pairs, so this is the memory high-water mark
117
+ # of a forward pass. Fine at t=512; watch it if max_seq_len ever grows.
118
+ b, t = idx.shape
119
+ src = idx.unsqueeze(1).expand(b, t, t)
120
+ dst = idx.unsqueeze(2).expand(b, t, t)
121
+ vals = self.table[self.hash(src, dst)]
122
+ return self.scale.view(1, -1, 1, 1) * vals.unsqueeze(1)
123
+
124
+
125
+ # ----------------------------------------------------------------------------- blocks
126
+ class TokleAttention(nn.Module):
127
+ def __init__(self, config: TokleConfig, layer_idx: int):
128
+ super().__init__()
129
+ self.layer_idx = layer_idx
130
+ self.n_head = config.n_head
131
+ self.n_kv_head = config.n_kv_head
132
+ self.head_dim = config.head_dim
133
+ self.rep = config.n_head // config.n_kv_head
134
+ self.q_proj = nn.Linear(config.d_model, config.n_head * config.head_dim, bias=False)
135
+ self.k_proj = nn.Linear(config.d_model, config.n_kv_head * config.head_dim, bias=False)
136
+ self.v_proj = nn.Linear(config.d_model, config.n_kv_head * config.head_dim, bias=False)
137
+ self.o_proj = nn.Linear(config.n_head * config.head_dim, config.d_model, bias=False)
138
+ self.q_norm = TokleRMSNorm(config.head_dim, config.norm_eps)
139
+ self.k_norm = TokleRMSNorm(config.head_dim, config.norm_eps)
140
+
141
+ def forward(self, x, cos, sin, attn_mask, spab_bias):
142
+ b, t, _ = x.shape
143
+ q = self.q_proj(x).view(b, t, self.n_head, self.head_dim).transpose(1, 2)
144
+ k = self.k_proj(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
145
+ v = self.v_proj(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
146
+ # QK-norm before rope, as trained.
147
+ q = apply_rope(self.q_norm(q), cos, sin)
148
+ k = apply_rope(self.k_norm(k), cos, sin)
149
+ k = k.repeat_interleave(self.rep, dim=1)
150
+ v = v.repeat_interleave(self.rep, dim=1)
151
+ if spab_bias is None and attn_mask is None:
152
+ # Nothing to add: let SDPA build the causal mask itself, which is
153
+ # the fast kernel path. This is what layers 1..8 normally hit.
154
+ o = F.scaled_dot_product_attention(q, k, v, is_causal=True)
155
+ else:
156
+ mask = attn_mask
157
+ if spab_bias is not None:
158
+ causal = torch.ones(t, t, dtype=torch.bool, device=x.device).triu(1)
159
+ bias = spab_bias.masked_fill(causal, float("-inf"))
160
+ mask = bias if mask is None else bias + mask
161
+ o = F.scaled_dot_product_attention(q, k, v, attn_mask=mask.to(q.dtype))
162
+ return self.o_proj(o.transpose(1, 2).reshape(b, t, -1))
163
+
164
+
165
+ class TokleMLP(nn.Module):
166
+ def __init__(self, config: TokleConfig):
167
+ super().__init__()
168
+ self.gate_proj = nn.Linear(config.d_model, config.ffn_hidden, bias=False)
169
+ self.up_proj = nn.Linear(config.d_model, config.ffn_hidden, bias=False)
170
+ self.down_proj = nn.Linear(config.ffn_hidden, config.d_model, bias=False)
171
+
172
+ def forward(self, x):
173
+ return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
174
+
175
+
176
+ class TokleDecoderLayer(nn.Module):
177
+ def __init__(self, config: TokleConfig, layer_idx: int):
178
+ super().__init__()
179
+ self.input_layernorm = TokleRMSNorm(config.d_model, config.norm_eps)
180
+ self.self_attn = TokleAttention(config, layer_idx)
181
+ self.post_attention_layernorm = TokleRMSNorm(config.d_model, config.norm_eps)
182
+ self.mlp = TokleMLP(config)
183
+
184
+ def forward(self, x, cos, sin, attn_mask, spab_bias):
185
+ x = x + self.self_attn(self.input_layernorm(x), cos, sin, attn_mask, spab_bias)
186
+ return x + self.mlp(self.post_attention_layernorm(x))
187
+
188
+
189
+ # ----------------------------------------------------------------------------- model
190
+ class ToklePreTrainedModel(PreTrainedModel):
191
+ config_class = TokleConfig
192
+ base_model_prefix = "model"
193
+ supports_gradient_checkpointing = True
194
+ _no_split_modules = ["TokleDecoderLayer"]
195
+ _supports_sdpa = True
196
+
197
+ def _init_weights(self, module):
198
+ if isinstance(module, nn.Linear):
199
+ nn.init.normal_(module.weight, 0.0, 0.02)
200
+ if module.bias is not None:
201
+ nn.init.zeros_(module.bias)
202
+ elif isinstance(module, nn.Embedding):
203
+ nn.init.normal_(module.weight, 0.0, 0.02)
204
+ elif isinstance(module, TokleRMSNorm):
205
+ nn.init.ones_(module.weight)
206
+ elif isinstance(module, SPABBias):
207
+ nn.init.constant_(
208
+ module.scale, self.config.spab_init_scale if self.config.spab_enabled else 0.0)
209
+ # Residual-output projections get scaled down by depth so the residual
210
+ # stream does not blow up at init (GPT-2 trick).
211
+ if isinstance(module, (TokleAttention, TokleMLP)):
212
+ out = module.o_proj if isinstance(module, TokleAttention) else module.down_proj
213
+ nn.init.normal_(out.weight, 0.0, 0.02 / math.sqrt(2 * self.config.n_layer))
214
+
215
+
216
+ class TokleModel(ToklePreTrainedModel):
217
+ def __init__(self, config: TokleConfig):
218
+ super().__init__(config)
219
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.d_model)
220
+ self.spab = SPABBias(config)
221
+ self.layers = nn.ModuleList(
222
+ [TokleDecoderLayer(config, i) for i in range(config.n_layer)])
223
+ self.norm = TokleRMSNorm(config.d_model, config.norm_eps)
224
+ # Plain dict, not a buffer -- see the _M1 note. Keyed by device so
225
+ # .to()/.cuda() need not move anything.
226
+ self._rope_cache = {}
227
+ self.gradient_checkpointing = False
228
+ self.post_init()
229
+
230
+ def get_input_embeddings(self):
231
+ return self.embed_tokens
232
+
233
+ def set_input_embeddings(self, value):
234
+ self.embed_tokens = value
235
+
236
+ def rope(self, device):
237
+ key = str(device)
238
+ if key not in self._rope_cache:
239
+ self._rope_cache[key] = build_rope(
240
+ self.config.head_dim, self.config.max_seq_len,
241
+ self.config.rope_theta, device=device)
242
+ return self._rope_cache[key]
243
+
244
+ def _padding_mask(self, attention_mask, dtype):
245
+ """2D mask -> additive (b, 1, 1, t), or None when nothing is padded so
246
+ the caller can take the is_causal fast path."""
247
+ if attention_mask is None:
248
+ return None
249
+ attention_mask = attention_mask.to(torch.bool)
250
+ if bool(attention_mask.all()):
251
+ return None
252
+ m = torch.zeros(attention_mask.shape, dtype=dtype, device=attention_mask.device)
253
+ m = m.masked_fill(~attention_mask, torch.finfo(dtype).min)
254
+ return m[:, None, None, :]
255
+
256
+ def forward(
257
+ self,
258
+ input_ids: Optional[torch.LongTensor] = None,
259
+ attention_mask: Optional[torch.Tensor] = None,
260
+ inputs_embeds: Optional[torch.FloatTensor] = None,
261
+ output_hidden_states: Optional[bool] = None,
262
+ return_dict: Optional[bool] = None,
263
+ **kwargs,
264
+ ) -> Union[Tuple, BaseModelOutput]:
265
+ # getattr, not config.use_return_dict: that property is deprecated in
266
+ # Transformers 5 and warns on every call.
267
+ return_dict = (return_dict if return_dict is not None
268
+ else getattr(self.config, "return_dict", True))
269
+ output_hidden_states = (output_hidden_states if output_hidden_states is not None
270
+ else self.config.output_hidden_states)
271
+ if (input_ids is None) == (inputs_embeds is None):
272
+ raise ValueError("Pass exactly one of input_ids or inputs_embeds.")
273
+ if inputs_embeds is not None and self.config.spab_enabled:
274
+ raise ValueError(
275
+ "inputs_embeds is not supported while SPAB is enabled: the bias is "
276
+ "indexed by token id, so input_ids are required.")
277
+
278
+ x = self.embed_tokens(input_ids) if inputs_embeds is None else inputs_embeds
279
+ t = x.shape[1]
280
+ if t > self.config.max_seq_len:
281
+ raise ValueError(
282
+ f"Sequence length {t} exceeds max_seq_len {self.config.max_seq_len}.")
283
+
284
+ cos, sin = self.rope(x.device)
285
+ bias = self.spab(input_ids) if input_ids is not None else None
286
+
287
+ # Fold the causal mask in here once rather than per layer. Under causal
288
+ # attention with right padding, position 0 always sees itself, so no row
289
+ # ends up fully masked (which would give NaNs out of softmax).
290
+ pad = self._padding_mask(attention_mask, x.dtype)
291
+ if pad is not None:
292
+ causal = torch.ones(t, t, dtype=torch.bool, device=x.device).triu(1)
293
+ pad = pad.masked_fill(causal, torch.finfo(x.dtype).min)
294
+
295
+ hidden_states = () if output_hidden_states else None
296
+ for i, layer in enumerate(self.layers):
297
+ if output_hidden_states:
298
+ hidden_states += (x,)
299
+ layer_bias = bias if i == 0 else None # SPAB is layer 0 only
300
+ if self.gradient_checkpointing and self.training:
301
+ x = self._gradient_checkpointing_func(
302
+ layer.__call__, x, cos, sin, pad, layer_bias)
303
+ else:
304
+ x = layer(x, cos, sin, pad, layer_bias)
305
+ x = self.norm(x)
306
+ if output_hidden_states:
307
+ hidden_states += (x,)
308
+
309
+ if not return_dict:
310
+ return tuple(v for v in (x, hidden_states) if v is not None)
311
+ return BaseModelOutput(last_hidden_state=x, hidden_states=hidden_states)
312
+
313
+
314
+ # Transformers 5 turned _tied_weights_keys from a list of names into a
315
+ # {tied: source} mapping and calls .keys() on it, so the old form raises.
316
+ _TRANSFORMERS_V5 = int(_transformers_version.split(".")[0]) >= 5
317
+
318
+
319
+ class TokleForCausalLM(ToklePreTrainedModel, GenerationMixin):
320
+ _tied_weights_keys = ({"lm_head.weight": "model.embed_tokens.weight"}
321
+ if _TRANSFORMERS_V5 else ["lm_head.weight"])
322
+
323
+ def __init__(self, config: TokleConfig):
324
+ super().__init__(config)
325
+ self.model = TokleModel(config)
326
+ self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
327
+ self.post_init()
328
+
329
+ def get_input_embeddings(self):
330
+ return self.model.embed_tokens
331
+
332
+ def set_input_embeddings(self, value):
333
+ self.model.embed_tokens = value
334
+
335
+ def get_output_embeddings(self):
336
+ return self.lm_head
337
+
338
+ def set_output_embeddings(self, new_embeddings):
339
+ self.lm_head = new_embeddings
340
+
341
+ def get_decoder(self):
342
+ return self.model
343
+
344
+ def forward(
345
+ self,
346
+ input_ids: Optional[torch.LongTensor] = None,
347
+ attention_mask: Optional[torch.Tensor] = None,
348
+ inputs_embeds: Optional[torch.FloatTensor] = None,
349
+ labels: Optional[torch.LongTensor] = None,
350
+ output_hidden_states: Optional[bool] = None,
351
+ return_dict: Optional[bool] = None,
352
+ **kwargs,
353
+ ) -> Union[Tuple, CausalLMOutput]:
354
+ return_dict = (return_dict if return_dict is not None
355
+ else getattr(self.config, "return_dict", True))
356
+ out = self.model(
357
+ input_ids=input_ids,
358
+ attention_mask=attention_mask,
359
+ inputs_embeds=inputs_embeds,
360
+ output_hidden_states=output_hidden_states,
361
+ return_dict=True,
362
+ )
363
+ logits = self.lm_head(out.last_hidden_state)
364
+ # Rows past real_vocab_size only pad the matrix to a multiple of 64.
365
+ # They never saw a gradient and no target ever points at them, so keep
366
+ # them out of sampling and out of any log-likelihood the harness sums.
367
+ rv = self.config.real_vocab_size
368
+ if self.config.mask_padded_vocab_logits and rv is not None and rv < self.config.vocab_size:
369
+ logits = logits.clone()
370
+ logits[..., rv:] = float("-inf")
371
+
372
+ loss = None
373
+ if labels is not None:
374
+ loss = self.loss_function(logits=logits, labels=labels,
375
+ vocab_size=self.config.vocab_size, **kwargs)
376
+
377
+ if not return_dict:
378
+ return (loss, logits) if loss is not None else (logits,)
379
+ return CausalLMOutput(loss=loss, logits=logits, hidden_states=out.hidden_states)
special_tokens_map.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token": {
3
+ "content": "<s>",
4
+ "lstrip": false,
5
+ "normalized": false,
6
+ "rstrip": false,
7
+ "single_word": false
8
+ },
9
+ "eos_token": {
10
+ "content": "</s>",
11
+ "lstrip": false,
12
+ "normalized": false,
13
+ "rstrip": false,
14
+ "single_word": false
15
+ },
16
+ "pad_token": {
17
+ "content": "<pad>",
18
+ "lstrip": false,
19
+ "normalized": false,
20
+ "rstrip": false,
21
+ "single_word": false
22
+ }
23
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
The diff for this file is too large to render. See raw diff