Text Generation
Transformers
Safetensors
English
tokle
causal-lm
custom-architecture
slm
small-language-model
spab
custom_code
Instructions to use techdotus/Tokle-SPAB-3M with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use techdotus/Tokle-SPAB-3M with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="techdotus/Tokle-SPAB-3M", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("techdotus/Tokle-SPAB-3M", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use techdotus/Tokle-SPAB-3M with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "techdotus/Tokle-SPAB-3M" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "techdotus/Tokle-SPAB-3M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/techdotus/Tokle-SPAB-3M
- SGLang
How to use techdotus/Tokle-SPAB-3M with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "techdotus/Tokle-SPAB-3M" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "techdotus/Tokle-SPAB-3M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "techdotus/Tokle-SPAB-3M" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "techdotus/Tokle-SPAB-3M", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use techdotus/Tokle-SPAB-3M with Docker Model Runner:
docker model run hf.co/techdotus/Tokle-SPAB-3M
Initial Commit - Tokle
Browse files- LICENSE +21 -0
- README.md +131 -0
- config.json +32 -0
- configuration_tokle.py +80 -0
- generation_config.json +11 -0
- model.safetensors +3 -0
- modeling_tokle.py +379 -0
- special_tokens_map.json +23 -0
- tokenizer.json +0 -0
- tokenizer_config.json +0 -0
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
|
|
|