File size: 7,497 Bytes
685e018
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
"""Typed experiment configuration loaded from YAML."""

from __future__ import annotations

from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any

import yaml


MODEL_BACKBONES = ("project", "hf-qwen3")


@dataclass(frozen=True)
class ModelConfig:
    """Architecture for the bidirectional token denoiser."""

    vocab_size: int
    mask_token_id: int
    max_seq_len: int
    d_model: int
    n_layers: int
    n_heads: int
    d_ff: int
    dropout: float = 0.1
    tie_embeddings: bool = True
    activation_checkpointing: bool = False
    use_flex_attention: bool = False
    compile_backbone: bool = False
    forbidden_output_token_ids: tuple[int, ...] = (0, 1, 2, 4)
    backbone: str = "project"
    pretrained_path: str | None = None
    # First id of the untrained tail of a pretrained embedding matrix; logits from that id
    # upward are masked without enumerating hundreds of forbidden ids.
    forbidden_output_from: int | None = None

    def __post_init__(self) -> None:
        if self.vocab_size <= 1:
            raise ValueError("vocab_size must be greater than one")
        if self.backbone not in MODEL_BACKBONES:
            raise ValueError(f"backbone must be one of {MODEL_BACKBONES}")
        if self.backbone != "project" and self.pretrained_path is None:
            raise ValueError("pretrained backbones require pretrained_path")
        if self.forbidden_output_from is not None and not (
            0 < self.forbidden_output_from <= self.vocab_size
        ):
            raise ValueError("forbidden_output_from must be inside the vocabulary")
        if not 0 <= self.mask_token_id < self.vocab_size:
            raise ValueError("mask_token_id must be inside the vocabulary")
        if self.max_seq_len <= 0:
            raise ValueError("max_seq_len must be positive")
        if self.d_model <= 0 or self.n_layers <= 0 or self.n_heads <= 0 or self.d_ff <= 0:
            raise ValueError("all model dimensions must be positive")
        if self.d_model % self.n_heads != 0:
            raise ValueError("d_model must be divisible by n_heads")
        if not 0.0 <= self.dropout < 1.0:
            raise ValueError("dropout must be in [0, 1)")
        if not isinstance(self.activation_checkpointing, bool):
            raise ValueError("activation_checkpointing must be a boolean")
        if not isinstance(self.use_flex_attention, bool):
            raise ValueError("use_flex_attention must be a boolean")
        if not isinstance(self.compile_backbone, bool):
            raise ValueError('compile_backbone must be a boolean')

        forbidden_ids = tuple(self.forbidden_output_token_ids)
        if any(
            isinstance(token_id, bool) or not isinstance(token_id, int)
            for token_id in forbidden_ids
        ):
            raise ValueError("forbidden_output_token_ids must contain integers")
        if len(set(forbidden_ids)) != len(forbidden_ids):
            raise ValueError("forbidden_output_token_ids must not contain duplicates")
        if any(not 0 <= token_id < self.vocab_size for token_id in forbidden_ids):
            raise ValueError("forbidden output token ids must be inside the vocabulary")
        if self.mask_token_id not in forbidden_ids:
            raise ValueError("mask_token_id must be a forbidden output token")
        if self.backbone == "project" and 3 in forbidden_ids:
            raise ValueError("EOS token id 3 must remain an allowed output token")
        object.__setattr__(self, "forbidden_output_token_ids", forbidden_ids)


@dataclass(frozen=True)
class TrainingConfig:
    """Single-device training settings for the first research iteration."""

    train_data: str
    tokenizer: str
    output_dir: str = "outputs/run"
    val_data: str | None = None
    batch_size: int = 32
    gradient_accumulation_steps: int = 1
    max_steps: int = 10_000
    learning_rate: float = 3e-4
    min_learning_rate: float = 3e-5
    warmup_steps: int = 500
    weight_decay: float = 0.1
    grad_clip: float = 1.0
    mask_eps: float = 1e-3
    seed: int = 1337
    device: str = "auto"
    precision: str = "auto"
    optimizer: str = "adamw"
    optimizer_min_8bit_size: int = 4096
    optimizer_embedding_32bit: bool = True
    require_fused_attention: bool = False
    save_inference_checkpoint: bool = False
    num_workers: int = 0
    log_interval: int = 10
    eval_interval: int = 500
    eval_batches: int = 20
    save_interval: int = 500
    keep_last_checkpoints: int = 3

    def __post_init__(self) -> None:
        if self.batch_size <= 0 or self.gradient_accumulation_steps <= 0:
            raise ValueError("batch sizes must be positive")
        if self.max_steps <= 0:
            raise ValueError("max_steps must be positive")
        if not 0.0 < self.learning_rate:
            raise ValueError("learning_rate must be positive")
        if not 0.0 <= self.min_learning_rate <= self.learning_rate:
            raise ValueError("min_learning_rate must be between zero and learning_rate")
        if not 0 <= self.warmup_steps < self.max_steps:
            raise ValueError("warmup_steps must be non-negative and less than max_steps")
        if self.weight_decay < 0.0:
            raise ValueError("weight_decay must be non-negative")
        if self.grad_clip <= 0.0:
            raise ValueError("grad_clip must be positive")
        if not 0.0 < self.mask_eps < 1.0:
            raise ValueError("mask_eps must be in (0, 1)")
        if self.num_workers < 0:
            raise ValueError("num_workers must be non-negative")
        if self.log_interval <= 0 or self.eval_interval <= 0 or self.save_interval <= 0:
            raise ValueError("log, eval, and save intervals must be positive")
        if self.eval_batches <= 0:
            raise ValueError("eval_batches must be positive")
        if self.keep_last_checkpoints < 0:
            raise ValueError("keep_last_checkpoints must be non-negative")
        if self.precision not in {"auto", "float32", "bfloat16", "float16"}:
            raise ValueError("precision must be auto, float32, bfloat16, or float16")
        if self.optimizer not in {"adamw", "adamw8bit"}:
            raise ValueError("optimizer must be adamw or adamw8bit")
        if self.optimizer_min_8bit_size <= 0:
            raise ValueError("optimizer_min_8bit_size must be positive")
        if not isinstance(self.optimizer_embedding_32bit, bool):
            raise ValueError("optimizer_embedding_32bit must be a boolean")
        if not isinstance(self.require_fused_attention, bool):
            raise ValueError("require_fused_attention must be a boolean")
        if not isinstance(self.save_inference_checkpoint, bool):
            raise ValueError("save_inference_checkpoint must be a boolean")


@dataclass(frozen=True)
class ExperimentConfig:
    model: ModelConfig
    training: TrainingConfig

    def to_dict(self) -> dict[str, Any]:
        return asdict(self)


def load_config(path: str | Path) -> ExperimentConfig:
    """Load and validate an experiment YAML file."""

    config_path = Path(path)
    with config_path.open("r", encoding="utf-8") as handle:
        raw = yaml.safe_load(handle)
    if not isinstance(raw, dict) or "model" not in raw or "training" not in raw:
        raise ValueError("config must contain top-level model and training mappings")
    return ExperimentConfig(
        model=ModelConfig(**raw["model"]),
        training=TrainingConfig(**raw["training"]),
    )