ViuAI-500M / code /config.py
ViuAI's picture
Update code/config.py with audited production fixes
5c5f44e verified
Raw
History Blame Contribute Delete
1.27 kB
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