remg1997's picture
Upload folder using huggingface_hub
4784dd3 verified
Raw
History Blame Contribute Delete
4.77 kB
from __future__ import annotations
from typing import Optional
from transformers import PretrainedConfig
from babylm.config.schema import ModelConfig
class ModernBertSmallConfig(PretrainedConfig):
"""HF-compatible config for the from-scratch ModernBERT-Small architecture.
Bridges to the pydantic `ModelConfig` (used for JSON validation/CLI) via `to_hf_config`,
while giving `save_pretrained`/`from_pretrained`/`AutoConfig` compatibility.
"""
model_type = "modernbert_small"
def __init__(
self,
vocab_size: int = 30522,
hidden_size: int = 384,
num_hidden_layers: int = 16,
num_attention_heads: int = 6,
intermediate_size: int = 576,
hidden_activation: str = "gelu",
attention_pattern: str = "alternating",
global_attn_every_n_layers: int = 3,
global_rope_theta: float = 160000.0,
local_rope_theta: float = 10000.0,
local_attention_window: int = 128,
norm_eps: float = 1e-5,
attention_dropout: float = 0.0,
mlp_dropout: float = 0.0,
embedding_dropout: float = 0.0,
embedding_type: str = "standard",
factorized_linear_bottleneck_dim: int = 128,
factorized_mlp_bottleneck_dim: int = 128,
sttp_rank: int = 8,
sttp_num_factors: int = 3,
pete_basis_dim: int = 384,
compositional_bottleneck_dim: int = 128,
compositional_num_buckets: int = 8192,
compositional_min_ngram: int = 1,
compositional_max_ngram: int = 4,
compositional_residual_mode: str = "none",
compositional_frequency_tau: float = 100.0,
compositional_frequency_shuffle_seed: int = 42,
geometry_penalty_mode: str = "none",
geometry_penalty_weight: float = 0.01,
geometry_sample_size: int = 2048,
geometry_sampling_mode: str = "zipfian",
compositional_token_strings: Optional[list[str]] = None,
token_frequencies: Optional[list[int]] = None,
special_token_ids: Optional[list[int]] = None,
initializer_range: float = 0.02,
pad_token_id: int = 0,
tie_word_embeddings: bool = True,
**kwargs,
) -> None:
super().__init__(pad_token_id=pad_token_id, tie_word_embeddings=tie_word_embeddings, **kwargs)
self.vocab_size = vocab_size
self.hidden_size = hidden_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.intermediate_size = intermediate_size
self.hidden_activation = hidden_activation
self.attention_pattern = attention_pattern
self.global_attn_every_n_layers = global_attn_every_n_layers
self.global_rope_theta = global_rope_theta
self.local_rope_theta = local_rope_theta
self.local_attention_window = local_attention_window
self.norm_eps = norm_eps
self.attention_dropout = attention_dropout
self.mlp_dropout = mlp_dropout
self.embedding_dropout = embedding_dropout
self.embedding_type = embedding_type
self.factorized_linear_bottleneck_dim = factorized_linear_bottleneck_dim
self.factorized_mlp_bottleneck_dim = factorized_mlp_bottleneck_dim
self.sttp_rank = sttp_rank
self.sttp_num_factors = sttp_num_factors
self.pete_basis_dim = pete_basis_dim
self.compositional_bottleneck_dim = compositional_bottleneck_dim
self.compositional_num_buckets = compositional_num_buckets
self.compositional_min_ngram = compositional_min_ngram
self.compositional_max_ngram = compositional_max_ngram
self.compositional_residual_mode = compositional_residual_mode
self.compositional_frequency_tau = compositional_frequency_tau
self.compositional_frequency_shuffle_seed = compositional_frequency_shuffle_seed
self.geometry_penalty_mode = geometry_penalty_mode
self.geometry_penalty_weight = geometry_penalty_weight
self.geometry_sample_size = geometry_sample_size
self.geometry_sampling_mode = geometry_sampling_mode
self.compositional_token_strings = compositional_token_strings
self.token_frequencies = token_frequencies
self.special_token_ids = special_token_ids or []
self.initializer_range = initializer_range
def to_hf_config(model_cfg: ModelConfig, pad_token_id: Optional[int] = None) -> ModernBertSmallConfig:
"""Builds a ModernBertSmallConfig from the pydantic ModelConfig, optionally overriding
pad_token_id with the value discovered from a trained tokenizer."""
data = model_cfg.model_dump()
if pad_token_id is not None:
data["pad_token_id"] = pad_token_id
return ModernBertSmallConfig(**data)