|
10 | 10 | from interface.core.project_state import ProjectStateMixin |
11 | 11 | from interface.core.project_state_apply import ProjectStateApplyMixin |
12 | 12 |
|
13 | | -from engine.config import DatasetConfig |
| 13 | +from engine.config import DatasetConfig, ModelConfig, TrainingConfig |
14 | 14 | from engine.data import Document |
15 | 15 | from engine.dataset_corpus import _StreamingCorpusBuilder |
16 | 16 | from engine.dataset_mixture import ( |
|
25 | 25 | token_dtype_for_vocab, |
26 | 26 | train_tokenizer, |
27 | 27 | ) |
28 | | -from engine.training import TokenDataset |
| 28 | +from engine.training import TokenDataset, train_model |
29 | 29 |
|
30 | 30 |
|
31 | 31 | class DiversityFilterTests(unittest.TestCase): |
@@ -181,5 +181,60 @@ def test_validation_stride_equals_context_length(self) -> None: |
181 | 181 | self.assertEqual(y1.tolist(), list(range(65, 129))) |
182 | 182 |
|
183 | 183 |
|
| 184 | +class TrainingDiagnosticsAndTelemetryTests(unittest.TestCase): |
| 185 | + """Tests for preflight warnings and telemetry emission enhancements.""" |
| 186 | + |
| 187 | + def test_head_dim_not_divisible_by_eight_detected(self) -> None: |
| 188 | + model_config = ModelConfig( |
| 189 | + vocab_size=32, |
| 190 | + context_length=16, |
| 191 | + embedding_size=560, |
| 192 | + head_count=8, |
| 193 | + layer_count=2, |
| 194 | + ) |
| 195 | + head_dim = model_config.embedding_size // model_config.head_count |
| 196 | + self.assertEqual(head_dim, 70) |
| 197 | + self.assertNotEqual(head_dim % 8, 0) |
| 198 | + |
| 199 | + def test_milestone_step_emits_event(self) -> None: |
| 200 | + with tempfile.TemporaryDirectory() as tmp_dir: |
| 201 | + model_config = ModelConfig( |
| 202 | + vocab_size=16, |
| 203 | + context_length=8, |
| 204 | + embedding_size=16, |
| 205 | + head_count=2, |
| 206 | + layer_count=1, |
| 207 | + dropout=0.0, |
| 208 | + ) |
| 209 | + training_config = TrainingConfig( |
| 210 | + output_dir=Path(tmp_dir), |
| 211 | + epochs=1, |
| 212 | + batch_size=1, |
| 213 | + learning_rate=1e-3, |
| 214 | + sample_stride=8, |
| 215 | + warmup_steps=0, |
| 216 | + eval_interval=0, |
| 217 | + save_interval=0, |
| 218 | + use_amp=False, |
| 219 | + precision="fp32", |
| 220 | + device="cpu", |
| 221 | + resume=False, |
| 222 | + early_stopping=False, |
| 223 | + ) |
| 224 | + events: list[dict] = [] |
| 225 | + train_model( |
| 226 | + model_config, |
| 227 | + training_config, |
| 228 | + [index % 16 for index in range(32)], |
| 229 | + [], |
| 230 | + pad_token_id=-1, |
| 231 | + progress=events.append, |
| 232 | + ) |
| 233 | + step_events = [e for e in events if e.get("event_type") == "step"] |
| 234 | + self.assertGreaterEqual(len(step_events), 1) |
| 235 | + self.assertIn("Step 1", step_events[0]["message"]) |
| 236 | + |
| 237 | + |
184 | 238 | if __name__ == "__main__": |
185 | 239 | unittest.main() |
| 240 | + |
0 commit comments