from dataclasses import dataclass @dataclass class ViuAIConfig: vocab_size: int = 64003 d_model: int = 1280 n_layers: int = 24 n_heads: int = 20 n_kv_heads: int = 4 ffn_hidden: int = 3456 context_length: int = 2048 rope_theta: float = 10000.0 norm_eps: float = 1e-5 z_loss_weight: float = 0.0 # 0.0 default for SFT/Inference (use 1e-4 for pretraining) use_checkpoint: bool = True attn_dropout: float = 0.05 # 0.0 for pretraining, 0.05 for SFT resid_dropout: float = 0.05 # 0.0 for pretraining, 0.05 for SFT neftune_alpha: float = 5.0 # NEFTune noise scale for SFT quality boost @classmethod def pretrain(cls, **kwargs): """Standard configuration for base pretraining (no dropout, 1e-4 z-loss).""" defaults = dict(z_loss_weight=1e-4, attn_dropout=0.0, resid_dropout=0.0, neftune_alpha=0.0) defaults.update(kwargs) return cls(**defaults) @classmethod def sft(cls, **kwargs): """Optimized configuration for Supervised Fine-Tuning.""" defaults = dict(z_loss_weight=0.0, attn_dropout=0.05, resid_dropout=0.05, neftune_alpha=5.0) defaults.update(kwargs) return cls(**defaults) # Compatibility Alias ModelArgs = ViuAIConfig