"""HF config for MetaLLM / Vishvakarma (custom arch: NoPE-every-N, QK-norm, GQA, SwiGLU). Field names intentionally mirror configs/model.py:ModelConfig so this object can be passed directly to metallm_core.MetaLLMv2 as its `cfg` (duck-typed). """ from transformers import PretrainedConfig class MetaLLMConfig(PretrainedConfig): model_type = "metallm" def __init__( self, vocab_size: int = 40000, n_layers: int = 26, d_model: int = 1792, n_heads: int = 28, n_kv_heads: int = 14, d_ff: int = 4864, max_seq_len: int = 2048, rope_theta: float = 500_000.0, norm_eps: float = 1e-5, tie_embeddings: bool = True, qk_norm: bool = True, z_loss_weight: float = 0.0, # inference shim: aux loss unused, kept for fidelity rope_fp32: bool = True, doc_mask: bool = True, # equals plain causal for single-document prompts attn_impl: str = "sdpa", # "sdpa" is the portable inference default nope_every: int = 4, bos_id: int = 1, **kwargs, ): self.vocab_size = vocab_size self.n_layers = n_layers self.d_model = d_model self.n_heads = n_heads self.n_kv_heads = n_kv_heads self.d_ff = d_ff self.max_seq_len = max_seq_len self.rope_theta = rope_theta self.norm_eps = norm_eps self.tie_embeddings = tie_embeddings self.qk_norm = qk_norm self.z_loss_weight = z_loss_weight self.rope_fp32 = rope_fp32 self.doc_mask = doc_mask self.attn_impl = attn_impl self.nope_every = nope_every self.bos_id = bos_id # standard-name aliases: transformers>=5.13 core reads these directly self.num_hidden_layers = n_layers self.hidden_size = d_model self.num_attention_heads = n_heads self.num_key_value_heads = n_kv_heads self.max_position_embeddings = max_seq_len kwargs.setdefault("tie_word_embeddings", tie_embeddings) kwargs.setdefault("bos_token_id", bos_id) super().__init__(**kwargs) @property def head_dim(self) -> int: return self.d_model // self.n_heads