Tokle-SPAB-3M / configuration_tokle.py
anandselvadurai-techus's picture
Initial Commit - Tokle
4bbdc93 verified
Raw History Blame Contribute Delete
2.82 kB
"""Tokle config. Field names follow the training script (n_layer, d_model, ...);
attribute_map aliases them to the names Transformers and lm-eval look for.
"""
from transformers.configuration_utils import PretrainedConfig
class TokleConfig(PretrainedConfig):
"""Decoder-only LM with SPAB (Static Pairwise Attention Bias).
SPAB is a token-pair-indexed additive bias on the first layer's attention
logits. Values come from windowed PMI over the training corpus, hashed into
one flat table shared by every head and frozen before step 0 -- the only
thing trained here is spab scale, one number per head.
Two fields are easy to trip over. vocab_size (5056) is the padded embedding
matrix; real_vocab_size (5048) is what the tokenizer can actually emit, and
everything between them is untrained filler. use_cache is False because SPAB
hashes every id in the window, so there is nothing to cache incrementally.
"""
model_type = "tokle"
keys_to_ignore_at_inference = ["past_key_values"]
attribute_map = {
"num_hidden_layers": "n_layer",
"hidden_size": "d_model",
"num_attention_heads": "n_head",
"num_key_value_heads": "n_kv_head",
"intermediate_size": "ffn_hidden",
"max_position_embeddings": "max_seq_len",
"rms_norm_eps": "norm_eps",
}
def __init__(
self,
vocab_size=5056,
real_vocab_size=5048,
max_seq_len=512,
n_layer=9,
d_model=144,
n_head=3,
n_kv_head=1,
head_dim=48,
ffn_hidden=432,
rope_theta=10000.0,
norm_eps=1e-5,
spab_enabled=True,
spab_table_size=8388608,
spab_init_scale=0.1,
mask_padded_vocab_logits=True,
tie_word_embeddings=True,
bos_token_id=0,
eos_token_id=1,
pad_token_id=2,
use_cache=False,
**kwargs,
):
self.vocab_size = vocab_size
self.real_vocab_size = real_vocab_size
self.max_seq_len = max_seq_len
self.n_layer = n_layer
self.d_model = d_model
self.n_head = n_head
self.n_kv_head = n_kv_head
self.head_dim = head_dim
self.ffn_hidden = ffn_hidden
self.rope_theta = rope_theta
self.norm_eps = norm_eps
self.spab_enabled = spab_enabled
self.spab_table_size = spab_table_size
self.spab_init_scale = spab_init_scale
self.mask_padded_vocab_logits = mask_padded_vocab_logits
self.use_cache = use_cache # see class docstring: no KV cache with SPAB
super().__init__(
tie_word_embeddings=tie_word_embeddings,
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
pad_token_id=pad_token_id,
**kwargs,
)