File size: 1,272 Bytes
cc131d8
 
 
 
da9bc4a
cc131d8
 
 
 
 
 
 
 
5c5f44e
cc131d8
6900e6a
 
7e252ea
663809b
5c5f44e
 
 
 
 
 
 
 
 
 
 
 
 
 
663809b
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
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