File size: 3,513 Bytes
7daf8d2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
from transformers import PretrainedConfig


class AttnExtConfig(PretrainedConfig):
    model_type = "attn_ext"
    keys_to_ignore_at_inference = ["past_key_values"]

    def __init__(

        self,

        vocab_size=49152,

        d_model=2048,

        n_layer=24,

        n_head=32,

        ffn_multiplier=4.0,

        multiple_of=256,

        block_size=2048,

        rope_theta=10000.0,

        dropout=0.0,

        rms_norm_eps=1e-5,

        initializer_range=0.02,

        attention_bias=False,

        mlp_bias=False,

        input_mode="learned",

        binary_dim=16,

        binary_encoding="zero_one",

        binary_scale=1.0,

        code_seed=12345,

        min_row_weight=4,

        min_col_weight=4,

        pad_token_id=None,

        bos_token_id=None,

        eos_token_id=None,

        tie_word_embeddings=False,

        use_cache=False,

        **kwargs,

    ):
        super().__init__(
            pad_token_id=pad_token_id,
            bos_token_id=bos_token_id,
            eos_token_id=eos_token_id,
            tie_word_embeddings=tie_word_embeddings,
            **kwargs,
        )

        if d_model % n_head != 0:
            raise ValueError("d_model must be divisible by n_head")

        head_dim = d_model // n_head
        if head_dim % 2 != 0:
            raise ValueError("RoPE requires an even head dimension")

        if input_mode not in {"learned", "binary16", "gf2"}:
            raise ValueError(
                "input_mode must be learned, binary16, or gf2"
            )

        if input_mode != "learned":
            if binary_dim != 16:
                raise ValueError("Frozen-code models require binary_dim=16")
            if vocab_size > 2**binary_dim:
                raise ValueError("Vocabulary does not fit in 16 bits")
            if d_model % binary_dim != 0:
                raise ValueError(
                    "d_model must be divisible by binary_dim"
                )
            if tie_word_embeddings:
                raise ValueError(
                    "Frozen input codes cannot be tied to lm_head"
                )

        if binary_encoding not in {"zero_one", "bipolar"}:
            raise ValueError(
                "binary_encoding must be zero_one or bipolar"
            )

        self.vocab_size = vocab_size

        self.d_model = d_model
        self.hidden_size = d_model

        self.n_layer = n_layer
        self.num_hidden_layers = n_layer

        self.n_head = n_head
        self.num_attention_heads = n_head
        self.head_dim = head_dim

        self.ffn_multiplier = ffn_multiplier
        self.multiple_of = multiple_of

        self.block_size = block_size
        self.max_position_embeddings = block_size
        self.rope_theta = rope_theta

        self.dropout = dropout
        self.rms_norm_eps = rms_norm_eps
        self.initializer_range = initializer_range
        self.attention_bias = attention_bias
        self.mlp_bias = mlp_bias

        self.input_mode = input_mode
        self.binary_dim = binary_dim
        self.binary_encoding = binary_encoding
        self.binary_scale = binary_scale
        self.binary_repeat = d_model // binary_dim

        self.code_seed = code_seed
        self.min_row_weight = min_row_weight
        self.min_col_weight = min_col_weight

        self.use_cache = use_cache
        self.is_decoder = True
        self.is_encoder_decoder = False