| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
| from pathlib import Path |
| import json |
|
|
|
|
| @dataclass |
| class ModelConfig: |
| vocab_size: int = 48000 |
| max_seq_len: int = 2048 |
|
|
| hidden_size: int = 768 |
| intermediate_size: int = 2304 |
| num_layers: int = 24 |
| num_heads: int = 12 |
| num_kv_heads: int = 4 |
|
|
| rope_theta: float = 10_000.0 |
|
|
| rms_norm_eps: float = 1.0e-5 |
| qk_norm: bool = True |
|
|
| initializer_range: float = 0.02 |
|
|
| eos_token_id: int = 0 |
| pad_token_id: int = 1 |
| bos_token_id: int = 3 |
|
|
| tie_word_embeddings: bool = True |
|
|
| z_loss_coef: float = 1.0e-4 |
| attn_dropout: float = 0.0 |
| resid_dropout: float = 0.0 |
|
|
| @property |
| def head_dim(self) -> int: |
| assert self.hidden_size % self.num_heads == 0 |
| return self.hidden_size // self.num_heads |
|
|
| @property |
| def kv_groups(self) -> int: |
| assert self.num_heads % self.num_kv_heads == 0 |
| return self.num_heads // self.num_kv_heads |
|
|
| @classmethod |
| def load(cls, path: str | Path) -> "ModelConfig": |
| d = json.loads(Path(path).read_text(encoding="utf-8")) |
| valid = {f for f in cls.__dataclass_fields__} |
| return cls(**{k: v for k, v in d.items() if k in valid}) |
|
|