"""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, )