ab_ext_binary16 / configuration_attn_ext.py
Bochkov's picture
Upload model files
7daf8d2 verified
Raw History Blame Contribute Delete
3.51 kB
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