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