test-new-arch / configuration.py
Harley-ml's picture
Upload 5 files
d3eee5a verified
Raw
History Blame Contribute Delete
2.88 kB
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")