File size: 2,819 Bytes
4bbdc93
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
"""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,
        )