from dataclasses import dataclass, field @dataclass class ModelConfig: # Architecture vocab_size: int = 32000 hidden_size: int = 512 num_layers: int = 6 num_heads: int = 8 ffn_intermediate_size: int = 1380 # SwiGLU max_seq_len: int = 512 dropout: float = 0.1 # IPA ipa_dim: int = 24 ipa_lambda: float = 1.0 ipa_mask_ratio: float = 0.15 # Ablation variant # "baseline" → no IPA # "ipa_add" → addition, no gate, no auxiliary loss # "ipa_gate" → gated fusion, no auxiliary loss # "ipa_full" → gated fusion + auxiliary loss variant: str = "ipa_full" # Training learning_rate: float = 3e-4 warmup_steps: int = 1000 weight_decay: float = 0.1 batch_size: int = 32 grad_clip: float = 1.0 # Data languages: list = field(default_factory=lambda: ["en", "nl", "mandarin"]) byte_premiums: dict = field(default_factory=lambda: { "en": 1.0, "nl": 1.0516, "mandarin": 0.9894 }) # BabyLM budget in byte-premium-adjusted whitespace-separated words. # STRICT-SMALL track: 100M words. MULTILINGUAL track: 1B words. word_budget: int = 100_000_000 # Paths data_dir: str = "/N/slate/partkaew/BigRed200/babylm2026/data" output_dir: str = "checkpoints" # Checkpointing milestones (Byte-Premium-adjusted word counts) checkpoint_milestones: list = field(default_factory=lambda: [ *range(1_000_000, 10_000_000, 1_000_000), *range(10_000_000, 100_000_000, 10_000_000), *range(100_000_000, 1_000_000_000, 100_000_000), ]) def validate(self): assert self.variant in {"baseline", "ipa_add", "ipa_gate", "ipa_full"}, \ f"Unknown variant: {self.variant}" assert self.num_heads % 1 == 0 and self.hidden_size % self.num_heads == 0, \ "hidden_size must be divisible by num_heads" assert all(lang in self.byte_premiums for lang in self.languages), \ "All languages must have a byte premium defined" return self if __name__ == "__main__": cfg = ModelConfig().validate() print(cfg) print(f"Checkpoint milestones: {len(cfg.checkpoint_milestones)} total")