from __future__ import annotations import math from copy import deepcopy from pathlib import Path from typing import Any, ClassVar import yaml from transformers import PreTrainedConfig # Rotary base frequency. 500000 follows Llama-3 rather than Llama-2's 10000: # it trades a little short-range frequency resolution for enough headroom to # extend context past the frozen 4K presets later. RoPE scaling is not # implemented, so this value is fixed for the lifetime of a pretrained model. DEFAULT_ROPE_THETA = 500_000.0 class NeuronLMConfig(PreTrainedConfig): """Configuration for the NeuronLM decoder-only language model. ``rope_theta`` remains accepted as a compatibility alias, but new configurations serialize rotary settings through Transformers v5's ``rope_parameters`` field. """ model_type = "neuron_lm" keys_to_ignore_at_inference = ["past_key_values"] # Frozen 4K-context training presets live as YAML files under this # directory (one `.yaml` file per preset), not in this module. # Add or change a preset by editing/adding a YAML file, not this class. PRESETS_DIR: ClassVar[Path] = Path("configs/model") def __init__( self, vocab_size: int = 32_000, hidden_size: int = 768, intermediate_size: int = 2_048, num_hidden_layers: int = 12, num_attention_heads: int = 12, num_key_value_heads: int | None = None, max_position_embeddings: int = 2_048, rope_parameters: dict[str, Any] | None = None, rope_theta: float | None = None, rms_norm_eps: float = 1e-5, use_qk_norm: bool = True, attention_dropout: float = 0.0, residual_dropout: float = 0.0, initializer_range: float = 0.02, tie_word_embeddings: bool = True, use_cache: bool = True, # =0, =1, =2, =3 is the fixed special-token order # every NeuronLM tokenizer trains with (pretokenization.py's # DEFAULT_SPECIAL_TOKENS). Defaulting here means .generate() stops at # EOS without the caller having to pass it explicitly, and it can # still be overridden for a tokenizer with a different layout. pad_token_id: int | None = 3, bos_token_id: int | None = 0, eos_token_id: int | list[int] | None = 1, **kwargs: Any, ) -> None: if rope_parameters is not None and not isinstance( rope_parameters, dict, ): raise TypeError( "rope_parameters must be a dictionary or None, " f"got {type(rope_parameters).__name__}" ) resolved_rope_parameters = deepcopy(rope_parameters) or {} if ( rope_theta is not None and "rope_theta" in resolved_rope_parameters and resolved_rope_parameters["rope_theta"] != rope_theta ): raise ValueError( "rope_theta and rope_parameters['rope_theta'] disagree: " f"{rope_theta!r} != " f"{resolved_rope_parameters['rope_theta']!r}" ) resolved_rope_parameters.setdefault( "rope_theta", DEFAULT_ROPE_THETA if rope_theta is None else rope_theta, ) resolved_rope_parameters.setdefault("rope_type", "default") self.vocab_size = vocab_size self.hidden_size = hidden_size self.intermediate_size = intermediate_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.num_key_value_heads = ( num_attention_heads if num_key_value_heads is None else num_key_value_heads ) self.max_position_embeddings = max_position_embeddings self.rope_parameters = resolved_rope_parameters self.rms_norm_eps = rms_norm_eps self.use_qk_norm = use_qk_norm self.attention_dropout = attention_dropout self.residual_dropout = residual_dropout self.initializer_range = initializer_range self.is_decoder = True self.is_encoder_decoder = False self.is_causal = True self.use_cache = use_cache self._validate_dimensions() kwargs.update( { "pad_token_id": pad_token_id, "bos_token_id": bos_token_id, "eos_token_id": eos_token_id, "tie_word_embeddings": tie_word_embeddings, } ) super().__init__(**kwargs) self._validate() self.validate_rope() @classmethod def available_presets(cls, presets_dir: str | Path | None = None) -> list[str]: """List preset names discoverable as YAML files in ``presets_dir``.""" directory = Path(presets_dir) if presets_dir is not None else cls.PRESETS_DIR if not directory.is_dir(): return [] return sorted(path.stem for path in directory.glob("*.yaml")) @classmethod def _load_preset(cls, name: str, presets_dir: str | Path | None) -> dict[str, Any]: directory = Path(presets_dir) if presets_dir is not None else cls.PRESETS_DIR preset_path = directory / f"{name}.yaml" try: with preset_path.open("r", encoding="utf-8") as handle: preset = yaml.safe_load(handle) except OSError as error: available = ", ".join(cls.available_presets(directory)) raise ValueError( f"Unknown NeuronLM preset {name!r}; available presets: {available}" ) from error if not isinstance(preset, dict): raise ValueError( f"preset file {preset_path} must contain a YAML mapping of " "structural fields" ) return preset @classmethod def from_preset( cls, name: str, *, vocab_size: int = 32_000, presets_dir: str | Path | None = None, **overrides: Any, ) -> NeuronLMConfig: """Construct one of the frozen 4K-context training presets. Presets are loaded from YAML files under ``presets_dir`` (defaults to ``cls.PRESETS_DIR``), one file per preset named ``.yaml``. """ preset = deepcopy(cls._load_preset(name, presets_dir)) structural_fields = set(preset) conflicting = structural_fields.intersection(overrides) if conflicting: names = ", ".join(sorted(conflicting)) raise ValueError( f"Preset {name!r} has frozen structural fields and cannot " f"override: {names}" ) return cls(vocab_size=vocab_size, **preset, **overrides) def _validate(self) -> None: self._validate_dimensions() if self.head_dim % 2 != 0: raise ValueError( f"RoPE requires an even head dimension, got head_dim={self.head_dim}" ) if self.rope_parameters.get("rope_type") != "default": raise ValueError( "NeuronLM currently supports only default RoPE; context " "extrapolation methods are intentionally deferred" ) if not _is_positive_finite_number(self.rope_theta): raise ValueError(f"rope_theta must be positive, got {self.rope_theta}") if not _is_positive_finite_number(self.rms_norm_eps): raise ValueError(f"rms_norm_eps must be positive, got {self.rms_norm_eps}") if type(self.use_qk_norm) is not bool: raise ValueError(f"use_qk_norm must be a boolean, got {self.use_qk_norm!r}") for name, value in { "attention_dropout": self.attention_dropout, "residual_dropout": self.residual_dropout, }.items(): if ( isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(float(value)) or not 0.0 <= value < 1.0 ): raise ValueError(f"{name} must be in [0, 1), got {value}") if not _is_positive_finite_number(self.initializer_range): raise ValueError( f"initializer_range must be positive, got {self.initializer_range}" ) def _validate_dimensions(self) -> None: positive_int_fields = { "vocab_size": self.vocab_size, "hidden_size": self.hidden_size, "intermediate_size": self.intermediate_size, "num_hidden_layers": self.num_hidden_layers, "num_attention_heads": self.num_attention_heads, "num_key_value_heads": self.num_key_value_heads, "max_position_embeddings": self.max_position_embeddings, } for name, value in positive_int_fields.items(): if type(value) is not int or value <= 0: raise ValueError(f"{name} must be a positive integer, got {value!r}") if self.hidden_size % self.num_attention_heads != 0: raise ValueError( f"hidden_size={self.hidden_size} must be divisible by " f"num_attention_heads={self.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, got " f"{self.num_attention_heads} and " f"{self.num_key_value_heads}" ) @property def head_dim(self) -> int: return self.hidden_size // self.num_attention_heads @property def qkv_projection_size(self) -> int: return (self.num_attention_heads + 2 * self.num_key_value_heads) * self.head_dim @property def rope_theta(self) -> float: return float(self.rope_parameters["rope_theta"]) @property def d_model(self) -> int: return self.hidden_size @property def d_ff(self) -> int: return self.intermediate_size @property def num_layers(self) -> int: return self.num_hidden_layers @property def num_heads(self) -> int: return self.num_attention_heads @property def context_length(self) -> int: return self.max_position_embeddings def _is_positive_finite_number(value: Any) -> bool: return ( not isinstance(value, bool) and isinstance(value, (int, float)) and math.isfinite(float(value)) and value > 0 )