Limen0.2B / config.py
ucr-max's picture
Release Limen0.2B
dc64c03 verified
Raw
History Blame Contribute Delete
2.7 kB
"""Standalone Transformers configuration for Limen0.2B."""
from __future__ import annotations
from transformers import PretrainedConfig
DEFAULT_VOCAB_SIZE = 16_384
DEFAULT_HIDDEN_SIZE = 768
DEFAULT_NUM_HIDDEN_LAYERS = 35
DEFAULT_NUM_ATTENTION_HEADS = 6
DEFAULT_NUM_KEY_VALUE_HEADS = 2
DEFAULT_HEAD_DIM = DEFAULT_HIDDEN_SIZE // DEFAULT_NUM_ATTENTION_HEADS
DEFAULT_INTERMEDIATE_SIZE = DEFAULT_HIDDEN_SIZE * 5 // 2
DEFAULT_BLOCK_SIZE = 1024
DEFAULT_ROPE_THETA = 100_000.0
class GPTConfig(PretrainedConfig):
"""Configuration for the Limen0.2B decoder-only language model."""
model_type = "gpt"
def __init__(
self,
vocab_size: int = DEFAULT_VOCAB_SIZE,
hidden_size: int = DEFAULT_HIDDEN_SIZE,
num_hidden_layers: int = DEFAULT_NUM_HIDDEN_LAYERS,
num_attention_heads: int = DEFAULT_NUM_ATTENTION_HEADS,
num_key_value_heads: int | None = DEFAULT_NUM_KEY_VALUE_HEADS,
intermediate_size: int | None = DEFAULT_INTERMEDIATE_SIZE,
head_dim: int | None = None,
block_size: int = DEFAULT_BLOCK_SIZE,
rope_theta: float = DEFAULT_ROPE_THETA,
rms_norm_eps: float = 1e-6,
xsa_projection: bool = True,
tie_word_embeddings: bool = True,
labels_are_shifted: bool = False,
**kwargs,
):
if num_key_value_heads is None:
num_key_value_heads = num_attention_heads
if head_dim is None:
if hidden_size % num_attention_heads != 0:
raise ValueError("hidden_size must be divisible by num_attention_heads")
head_dim = hidden_size // num_attention_heads
if intermediate_size is None:
intermediate_size = hidden_size * 4
if num_attention_heads % num_key_value_heads != 0:
raise ValueError("num_attention_heads must be divisible by num_key_value_heads")
if head_dim % 2 != 0:
raise ValueError("head_dim must be even for RoPE")
super().__init__(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.intermediate_size = int(intermediate_size)
self.head_dim = int(head_dim)
self.block_size = int(block_size)
self.max_position_embeddings = int(block_size)
self.rope_theta = float(rope_theta)
self.rms_norm_eps = float(rms_norm_eps)
self.xsa_projection = bool(xsa_projection)
self.labels_are_shifted = bool(labels_are_shifted)