| from __future__ import annotations | |
| from transformers import PretrainedConfig | |
| class CustomTransformerConfig(PretrainedConfig): | |
| model_type = "custom_transformer" | |
| def __init__( | |
| self, | |
| vocab_size: int = 32000, | |
| hidden_size: int = 512, | |
| num_hidden_layers: int = 8, | |
| num_attention_heads: int = 8, | |
| num_key_value_heads: int = 2, | |
| ffn_hidden_size: int = 1365, | |
| max_position_embeddings: int = 2048, | |
| rope_theta: float = 10000.0, | |
| norm_eps: float = 1e-5, | |
| dropout: float = 0.0, | |
| qk_bias: bool = True, | |
| use_head_gating: bool = True, | |
| attn_res_mode: str = "full", | |
| attn_res_block_size: int = 4, | |
| tie_word_embeddings: bool = True, | |
| pad_token_id: int | None = None, | |
| bos_token_id: int | None = None, | |
| eos_token_id: int | None = None, | |
| **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, | |
| ) | |
| self.vocab_size = int(vocab_size) | |
| self.hidden_size = int(hidden_size) | |
| self.num_hidden_layers = int(num_hidden_layers) | |
| self.num_attention_heads = int(num_attention_heads) | |
| self.num_key_value_heads = int(num_key_value_heads) | |
| self.ffn_hidden_size = int(ffn_hidden_size) | |
| self.max_position_embeddings = int(max_position_embeddings) | |
| self.rope_theta = float(rope_theta) | |
| self.norm_eps = float(norm_eps) | |
| self.dropout = float(dropout) | |
| self.qk_bias = bool(qk_bias) | |
| self.use_head_gating = bool(use_head_gating) | |
| self.attn_res_mode = str(attn_res_mode) | |
| self.attn_res_block_size = int(attn_res_block_size) | |
| self.tie_word_embeddings = bool(tie_word_embeddings) | |
| self.use_cache = False | |
| if self.attn_res_mode not in ("full", "block", "none"): | |
| raise ValueError("attn_res_mode must be one of: 'full', 'block', 'none'") | |
| if self.num_attention_heads <= 0: | |
| raise ValueError("num_attention_heads must be positive") | |
| if self.num_key_value_heads <= 0: | |
| raise ValueError("num_key_value_heads must be positive") | |
| if self.hidden_size % self.num_attention_heads != 0: | |
| raise ValueError("hidden_size must be divisible by num_attention_heads") | |
| if self.num_attention_heads % self.num_key_value_heads != 0: | |
| raise ValueError("num_attention_heads must be divisible by num_key_value_heads") | |
| if self.ffn_hidden_size < 1: | |
| raise ValueError("ffn_hidden_size must be positive") | |
| if self.max_position_embeddings <= 0: | |
| raise ValueError("max_position_embeddings must be positive") | |