diff --git a/AGENTS.md b/AGENTS.md index 5a67604..4990c8f 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -90,10 +90,10 @@ If a relevant check cannot be run, state why and what remains unverified. ## Current Technical Status -Milestone 021 is the latest modeling evidence. Boundary-aware UTF-8 ByteBPE320 and ByteBPE512 reach -best full-validation BPC `2.0286` and `2.0083`, beating the corrected character (`2.0760`) and -BPE128 (`2.0976`) controls. ByteBPE512 overfits after step 1,750 and ends at BPC `2.2450`; future -work should test early stopping or regularization before increasing model scale. +Milestone 022 is the latest modeling evidence. Patience-3 validation early stopping reproduces +ByteBPE512's best full-validation BPC `2.0083`, stops at step 2,500, and roughly halves runtime. +Weight decay `0.01` reaches `2.0080`, too small a difference to claim improvement. Future work should +test the tokenizer result across seeds or corpus splits rather than add another narrow hyperparameter. ## Safety And Security diff --git a/README.md b/README.md index 63cc6e3..d6359a2 100644 --- a/README.md +++ b/README.md @@ -46,6 +46,7 @@ chronological 90/10 split unless noted otherwise. | [019](experiments/019-professionalization-and-corrected-evaluation.md) | Corrected character vs BPE128 | Full-validation best bits/character | `2.0760` char vs `2.0976` BPE128 | BPE128 shortened sequences but remained narrowly worse after removing leakage and coverage bias. | | [020](experiments/020-bpe-context-and-learning-rate.md) | BPE context/LR controls | Best BPE128 bits/character | `2.0976` remains best | Matched character context and `5e-4` LR changed timing/diversity but did not improve held-out BPC. | | [021](experiments/021-boundary-aware-byte-bpe.md) | Boundary-aware ByteBPE320/512 | Full-validation best bits/character | `2.0286` / `2.0083` | Both beat the corrected character and BPE128 controls; ByteBPE512 overfit after step 1750. | +| [022](experiments/022-early-stopping-and-regularization.md) | ByteBPE512 early stopping / weight decay | Actual steps and best BPC | `2500`, `2.0083` / `2.0080` | Early stopping halves runtime and limits final degradation; weight decay `0.01` is effectively neutral. | Token-level loss and perplexity are not directly comparable between character and BPE tokenizers because they predict different units. The tokenizer @@ -81,7 +82,7 @@ shapes link directly to the implementation and its tests. | Tokenization | Character, educational character-BPE, and lossless boundary-aware UTF-8 byte-BPE tokenizers. | | Model | Decoder-only GPT-style Transformer with causal self-attention. | | Evaluation | Uniform, unigram, and add-one smoothed bigram baselines. | -| Training | Config-driven training with validation loss, progress logging, checkpoints, metrics, summaries, and samples. | +| Training | Config-driven training with validation loss, optional early stopping, progress logging, checkpoints, metrics, summaries, and samples. | | Run records | Preserved run directories with copied dataset manifests and selected provenance fields in `summary.json`. | | Generation | `max_new_tokens`, `temperature`, `top_k`, `seed`, and greedy decoding. | | Tests | Focused tests for data, baselines, model shape, training artifacts, run utilities, and generation behavior. | diff --git a/configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop.yaml b/configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop.yaml new file mode 100644 index 0000000..508a7f0 --- /dev/null +++ b/configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop.yaml @@ -0,0 +1,38 @@ +data: + input_path: data/raw/input.txt + prepared_path: data/processed/corpus.txt + manifest_path: data/processed/corpus_manifest.json + tokenizer_path: data/processed/tokenizer_bytebpe512.json + tokenizer_type: byte_bpe + bpe_vocab_size: 512 + bpe_min_frequency: 2 + block_size: 37 + train_split: 0.9 + +model: + vocab_size: 512 + block_size: 37 + n_layer: 4 + n_head: 4 + n_embd: 128 + dropout: 0.1 + +train: + run_name: gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop + runs_dir: runs + batch_size: 27 + max_steps: 5000 + learning_rate: 0.001 + weight_decay: 0.0 + log_interval: 100 + eval_interval: 250 + eval_batches: null + early_stopping_patience: 3 + early_stopping_min_delta: 0.0 + sample_prompt: Once + sample_max_new_tokens: 100 + sample_temperature: 1.0 + sample_top_k: null + sample_seed: 1337 + sample_greedy: false + seed: 1337 diff --git a/configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop.yaml b/configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop.yaml new file mode 100644 index 0000000..a870a2e --- /dev/null +++ b/configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop.yaml @@ -0,0 +1,38 @@ +data: + input_path: data/raw/input.txt + prepared_path: data/processed/corpus.txt + manifest_path: data/processed/corpus_manifest.json + tokenizer_path: data/processed/tokenizer_bytebpe512.json + tokenizer_type: byte_bpe + bpe_vocab_size: 512 + bpe_min_frequency: 2 + block_size: 37 + train_split: 0.9 + +model: + vocab_size: 512 + block_size: 37 + n_layer: 4 + n_head: 4 + n_embd: 128 + dropout: 0.1 + +train: + run_name: gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop + runs_dir: runs + batch_size: 27 + max_steps: 5000 + learning_rate: 0.001 + weight_decay: 0.01 + log_interval: 100 + eval_interval: 250 + eval_batches: null + early_stopping_patience: 3 + early_stopping_min_delta: 0.0 + sample_prompt: Once + sample_max_new_tokens: 100 + sample_temperature: 1.0 + sample_top_k: null + sample_seed: 1337 + sample_greedy: false + seed: 1337 diff --git a/docs/codex/build-and-test.md b/docs/codex/build-and-test.md index 25a95f2..3e6d5ea 100644 --- a/docs/codex/build-and-test.md +++ b/docs/codex/build-and-test.md @@ -50,7 +50,7 @@ The individual commands remain canonical and are listed below. python -m pytest ``` -Current expected result after milestone 019+: 90 tests passing with at least 90% coverage. +Current expected result after milestone 022+: at least 121 tests passing with at least 90% coverage. ### Compile Check diff --git a/docs/codex/experiments.md b/docs/codex/experiments.md index b0a53dc..fcf908d 100644 --- a/docs/codex/experiments.md +++ b/docs/codex/experiments.md @@ -153,3 +153,7 @@ tokenizer design rather than continue fine tuning around BPE128. Milestone 021 altered tokenizer design with boundary-aware ByteBPE320/512. Best full-validation BPC improved to `2.0286` and `2.0083`; the 512-token model overfit sharply after step 1,750, so the next controlled question is early stopping or modest regularization rather than more scale. + +Milestone 022 adds patience-3 validation early stopping. It reproduces ByteBPE512's step-1,750 best +checkpoint and stops at step 2,500, roughly halving runtime. Weight decay `0.01` is effectively +neutral at best BPC `2.0080`; future work should test robustness across seeds or corpus splits. diff --git a/docs/experiments.md b/docs/experiments.md index f3570db..37a8c7e 100644 --- a/docs/experiments.md +++ b/docs/experiments.md @@ -28,6 +28,7 @@ not a replacement for the original reports. | [019 Professionalization and Corrected Evaluation](../experiments/019-professionalization-and-corrected-evaluation.md) | Scientific hardening | Corrects tokenizer leakage and validation coverage, hardens artifacts and quality gates, adds the theory handbook, and reruns the controlled comparison. | | [020 BPE Context and Learning Rate](../experiments/020-bpe-context-and-learning-rate.md) | Tokenizer diagnostics | Matching character context and lowering BPE learning rate did not beat the corrected BPE128 control or character model. | | [021 Boundary-Aware Byte BPE](../experiments/021-boundary-aware-byte-bpe.md) | Tokenizer design | Lossless boundary-aware ByteBPE320/512 beat both corrected controls on best BPC; ByteBPE512 reached `2.0083` but overfit early. | +| [022 Early Stopping and Regularization](../experiments/022-early-stopping-and-regularization.md) | Training control | Patience-3 stopping reproduced the step-1750 optimum and halved runtime; weight decay `0.01` was effectively neutral. | ## Topic Shortcuts @@ -74,3 +75,7 @@ reported, but it did not beat the character control. Milestone 021 adds a lossless UTF-8 byte fallback and whitespace-boundary-aware merges. ByteBPE320 reached best BPC `2.0286`; ByteBPE512 reached `2.0083`, the strongest corrected result so far. The 512-token run's final BPC rose to `2.2450`, making early best-checkpoint selection essential. + +Milestone 022 operationalizes that result. Patience-3 early stopping terminates at step 2,500 and +retains best BPC `2.0083`, while stopped-final BPC improves from `2.2450` to `2.0554`. Weight decay +`0.01` reaches best BPC `2.0080`; the `0.00025` difference is too small to interpret as a real gain. diff --git a/docs/training.md b/docs/training.md index a8494af..ac8a606 100644 --- a/docs/training.md +++ b/docs/training.md @@ -111,6 +111,12 @@ python scripts/show_run.py --run latest --run-name gptiny `show_run.py` prints run paths, dataset summary fields, the latest metric, and the saved sample. +`summary.json` records the configured `max_steps` ceiling and `actual_steps`. Optional +`early_stopping_patience` counts consecutive validation events without an improvement larger than +`early_stopping_min_delta`; `stopped_early`, `stop_reason`, and the terminal state make the decision +inspectable. The final checkpoint is the model at the stop step, while `best_checkpoint.pt` remains +the lowest observed validation-loss model. + ## Generation Controls ```bash @@ -158,6 +164,9 @@ Experiment 020 found that matching BPE character context and halving its learnin the experiment-019 BPE control. Experiment 021 changed the tokenizer itself: boundary-aware ByteBPE320 and ByteBPE512 reached best BPC `2.0286` and `2.0083`, both beating the character control, though ByteBPE512 overfit sharply after step 1,750. +Experiment 022 adds patience-3 early stopping: it ends at step 2,500, halves runtime, and avoids most +final-checkpoint degradation. Weight decay `0.01` changes best BPC by only `0.00025`, which is not +meaningful evidence of improvement. ## Artifact Policy diff --git a/experiments/022-early-stopping-and-regularization.md b/experiments/022-early-stopping-and-regularization.md new file mode 100644 index 0000000..f8832ba --- /dev/null +++ b/experiments/022-early-stopping-and-regularization.md @@ -0,0 +1,81 @@ +# 022 — Early Stopping and Regularization + +## Goal + +Turn milestone 021's early ByteBPE512 optimum into explicit, inspectable training behavior, then test +whether modest AdamW weight decay delays overfitting. The intervention stays narrow: the same seed, +tokenizer, model, data order, validation contract, and 5,000-step ceiling. + +## Implementation + +`TrainConfig` now accepts `early_stopping_patience` and `early_stopping_min_delta`. Patience counts +full validation events without a meaningful improvement. The lowest numerical validation checkpoint +is still saved independently. Run summaries add `actual_steps`, `stopped_early`, `stop_reason`, and +the terminal early-stopping state; the final checkpoint represents the actual stop step. + +Both experiment configs use patience 3, minimum delta 0, full validation every 250 steps, and the +milestone-021 ByteBPE512 setup. One keeps weight decay 0; the other changes only weight decay to +`0.01`. + +## Results + +| setting | actual / ceiling steps | best step | best loss | best BPC | stopped-final BPC | duration | +| --- | ---: | ---: | ---: | ---: | ---: | ---: | +| 021 control, no stopping | 5,000 / 5,000 | 1,750 | 2.375155 | 2.008269 | 2.244955 | 422.7s | +| early stopping | 2,500 / 5,000 | 1,750 | 2.375155 | 2.008269 | 2.055447 | 209.3s | +| early stopping + WD 0.01 | 2,500 / 5,000 | 1,750 | 2.374858 | 2.008018 | 2.053365 | 204.6s | + +The unregularized run reproduces every observed milestone-021 loss through step 2,500, showing that +the stopping feature does not perturb optimization. It cuts runtime by 50.5% and avoids most of the +late final-checkpoint degradation. This is an orchestration win, not a new best model. + +Weight decay improves best BPC by only `0.000251` (0.0125%) and stopped-final BPC by `0.002082`. +With one deterministic seed, those differences are practically neutral and do not justify changing +the default research conclusion. Both runs stop after the three non-improving evaluations at steps +2,000, 2,250, and 2,500. + +## Controlled Generation + +Prompt `Once`, 100 new tokens. Seeded decoding uses temperature 0.8, top-k 10, seed 1337. + +| setting/checkpoint | greedy distinct-2 | seeded distinct-2 | +| --- | ---: | ---: | +| early-stop best | 0.5217 | 0.5417 | +| early-stop final | 0.5723 | 0.5980 | +| WD 0.01 best | 0.4326 | 0.5439 | +| WD 0.01 final | 0.4253 | 0.6380 | + +Best-checkpoint likelihood and distinct-2 again disagree. Weight decay does not visibly solve +coherence or repetition; the samples remain locally plausible but semantically unstable. + +## Exact Commands + +```bash +uv run --frozen --extra dev python scripts/prepare_data.py --config configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop.yaml +uv run --frozen --extra dev python scripts/prepare_data.py --config configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop.yaml +uv run --frozen --extra dev python scripts/train.py --config configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop.yaml +uv run --frozen --extra dev python scripts/train.py --config configs/gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop.yaml +uv run --frozen --extra dev python scripts/show_run.py --run latest --run-name gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop +uv run --frozen --extra dev python scripts/show_run.py --run latest --run-name gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop +uv run --frozen --extra dev make check +``` + +Generation used `scripts/generate.py` for both runs and `best`/`final` checkpoint kinds, once with +`--greedy --diagnostics` and once with +`--temperature 0.8 --top-k 10 --seed 1337 --diagnostics`. + +Run paths: + +- `runs/gptiny_bytebpe512_5k_lr1e-3_ctx37_earlystop/2026-07-12_22-12-08` +- `runs/gptiny_bytebpe512_5k_lr1e-3_ctx37_wd0.01_earlystop/2026-07-12_22-16-28` + +## Limitations And Next Step + +- Patience is evaluated only every 250 steps, so stopping reacts with bounded delay. +- One seed and one chronological split cannot establish whether the tiny WD difference is stable. +- Early stopping observes the validation set repeatedly and is part of model selection. +- Generation diagnostics are surface statistics, not human evaluation. + +The next useful milestone is robustness rather than another single-run hyperparameter: repeat the +ByteBPE512 early-stopping control across multiple seeds and report mean, spread, stop-step variation, +and generation diagnostics without selecting the best seed. diff --git a/notes/04-training.md b/notes/04-training.md index e4e3b96..6770837 100644 --- a/notes/04-training.md +++ b/notes/04-training.md @@ -46,6 +46,30 @@ Using sampled validation adds estimator variance. Official configs therefore use deterministic validation; `eval_batches` remains available for quick experiments and records its coverage explicitly. +### Validation-based early stopping + +Let validation be observed at evaluation index (j), with loss (L_j). Given minimum meaningful +improvement (delta\geq0), maintain a reference (R) and stale count (q): + +\[ +(R,q)\leftarrow +\begin{cases} +(L_j,0), & L_j < R-\delta,\\ +(R,q+1), & \text{otherwise}. +\end{cases} +\] + +Training stops when (q\geq P), where (P) is patience measured in validation events—not gradient +steps or epochs. smaLLM still stores the numerically lowest validation checkpoint independently of +the early-stopping reference. This matters when improvements smaller than (delta) are real enough +to preserve but intentionally too small to reset patience. + +Early stopping is a sequential model-selection rule, not regularization: it limits exposure to +overfitting but does not change the objective or gradients before the stop. Its latency is bounded by +(P\times\texttt{eval_interval}), and sampled validation can make the stopping time noisy. Official +modeling configs therefore use full deterministic validation. Summaries record the step ceiling, +actual steps, stop reason, patience, minimum delta, and terminal stale count. + ## Determinism limits Setting Python and PyTorch seeds controls many random choices, but bitwise identity can still fail diff --git a/src/smallm/config.py b/src/smallm/config.py index c846588..a2ecb9d 100644 --- a/src/smallm/config.py +++ b/src/smallm/config.py @@ -1,6 +1,7 @@ from __future__ import annotations from dataclasses import dataclass +from math import isfinite from pathlib import Path from typing import Any @@ -10,6 +11,15 @@ yaml = None +def _is_bounded_finite_number(value: object) -> bool: + return ( + isinstance(value, int | float) + and not isinstance(value, bool) + and (not isinstance(value, float) or isfinite(value)) + and abs(value) <= 1_000_000_000 + ) + + @dataclass(frozen=True) class DataConfig: input_path: str = "data/raw/input.txt" @@ -77,6 +87,8 @@ class TrainConfig: log_interval: int = 10 eval_interval: int = 100 eval_batches: int | None = 5 + early_stopping_patience: int | None = None + early_stopping_min_delta: float = 0.0 sample_prompt: str = "Once" sample_max_new_tokens: int = 100 sample_temperature: float = 1.0 @@ -93,16 +105,35 @@ def __post_init__(self) -> None: raise ValueError(f"train.{name} must be positive") if self.eval_batches is not None and self.eval_batches <= 0: raise ValueError("train.eval_batches must be positive or null") - if self.learning_rate <= 0: - raise ValueError("train.learning_rate must be positive") - if self.weight_decay < 0: - raise ValueError("train.weight_decay must be non-negative") + if self.early_stopping_patience is not None and ( + not isinstance(self.early_stopping_patience, int) + or isinstance(self.early_stopping_patience, bool) + or self.early_stopping_patience <= 0 + ): + raise ValueError("train.early_stopping_patience must be a positive integer or null") + if ( + not isinstance(self.early_stopping_min_delta, int | float) + or isinstance(self.early_stopping_min_delta, bool) + or ( + isinstance(self.early_stopping_min_delta, float) + and not isfinite(self.early_stopping_min_delta) + ) + or self.early_stopping_min_delta < 0 + or self.early_stopping_min_delta > 1_000_000_000 + ): + raise ValueError( + "train.early_stopping_min_delta must be finite and between 0 and 1000000000" + ) + if not _is_bounded_finite_number(self.learning_rate) or self.learning_rate <= 0: + raise ValueError("train.learning_rate must be finite and positive") + if not _is_bounded_finite_number(self.weight_decay) or self.weight_decay < 0: + raise ValueError("train.weight_decay must be finite and non-negative") if not self.sample_prompt: raise ValueError("train.sample_prompt must not be empty") if self.sample_max_new_tokens < 0: raise ValueError("train.sample_max_new_tokens must be non-negative") - if self.sample_temperature <= 0: - raise ValueError("train.sample_temperature must be positive") + if not _is_bounded_finite_number(self.sample_temperature) or self.sample_temperature <= 0: + raise ValueError("train.sample_temperature must be finite and positive") if self.sample_top_k is not None and self.sample_top_k <= 0: raise ValueError("train.sample_top_k must be positive or null") diff --git a/src/smallm/training/artifacts.py b/src/smallm/training/artifacts.py index d12c0ee..596f056 100644 --- a/src/smallm/training/artifacts.py +++ b/src/smallm/training/artifacts.py @@ -48,7 +48,7 @@ def write_config_snapshot(path: str | Path, config: ExperimentConfig) -> None: for section, values in data.items(): lines.append(f"{section}:") for key, value in values.items(): - lines.append(f" {key}: {json.dumps(value, ensure_ascii=False)}") + lines.append(f" {key}: {json.dumps(value, ensure_ascii=False, allow_nan=False)}") lines.append("") atomic_write_text(output, "\n".join(lines).rstrip() + "\n") @@ -60,7 +60,7 @@ def __init__(self, path: str | Path) -> None: self._handle = self.path.open("w", encoding="utf-8") def write(self, record: dict[str, Any]) -> None: - self._handle.write(json.dumps(record, sort_keys=True) + "\n") + self._handle.write(json.dumps(record, sort_keys=True, allow_nan=False) + "\n") self._handle.flush() os.fsync(self._handle.fileno()) @@ -82,7 +82,7 @@ def __exit__( def write_json(path: str | Path, payload: dict[str, Any]) -> None: output = Path(path) output.parent.mkdir(parents=True, exist_ok=True) - atomic_write_text(output, json.dumps(payload, indent=2, sort_keys=True) + "\n") + atomic_write_text(output, json.dumps(payload, indent=2, sort_keys=True, allow_nan=False) + "\n") def load_dataset_manifest(path: str | Path) -> dict[str, Any]: diff --git a/src/smallm/training/trainer.py b/src/smallm/training/trainer.py index f4eb3dd..9dabd5f 100644 --- a/src/smallm/training/trainer.py +++ b/src/smallm/training/trainer.py @@ -3,7 +3,7 @@ import platform import sys from dataclasses import asdict, dataclass -from math import log +from math import isfinite, log from pathlib import Path from time import perf_counter from typing import Protocol, cast @@ -38,6 +38,12 @@ def source_character_count(self, token_ids: list[int]) -> int: ... def encode_with_character_counts(self, text: str) -> tuple[list[int], list[int]]: ... +def _require_finite_loss(value: float, *, label: str) -> float: + if not isfinite(value): + raise RuntimeError(f"{label} loss became non-finite; aborting training") + return value + + @torch.no_grad() def estimate_loss( model: GPT, @@ -55,7 +61,8 @@ def estimate_loss( y = y.to(device) _, loss = model(x, y) if loss is not None: - losses.extend([float(loss.item())] * y.numel()) + value = _require_finite_loss(float(loss.item()), label="estimated") + losses.extend([value] * y.numel()) if was_training: model.train() if not losses: @@ -81,6 +88,31 @@ def bits_per_character(self) -> float: return self.total_nll / (self.target_characters * log(2)) +@dataclass(frozen=True) +class EarlyStoppingState: + reference_loss: float | None = None + evaluations_without_improvement: int = 0 + + +def _update_early_stopping( + state: EarlyStoppingState, + *, + val_loss: float | None, + patience: int | None, + min_delta: float, +) -> tuple[EarlyStoppingState, bool]: + if val_loss is None or patience is None: + return state, False + _require_finite_loss(val_loss, label="validation") + if state.reference_loss is None or val_loss < state.reference_loss - min_delta: + return EarlyStoppingState(reference_loss=val_loss), False + updated = EarlyStoppingState( + reference_loss=state.reference_loss, + evaluations_without_improvement=state.evaluations_without_improvement + 1, + ) + return updated, updated.evaluations_without_improvement >= patience + + def _validation_starts(token_count: int, block_size: int, max_batches: int | None) -> list[int]: starts = list(range(0, token_count - 1, block_size)) if max_batches is None or max_batches >= len(starts): @@ -118,7 +150,8 @@ def evaluate_tokens( _, loss = model(x, y) assert loss is not None count = y.numel() - total_nll += float(loss.item()) * count + value = _require_finite_loss(float(loss.item()), label="validation") + total_nll += value * count target_tokens += count if character_counts is None: target_characters += tokenizer.source_character_count(y[0].tolist()) @@ -162,6 +195,7 @@ def _update_best_validation( ) -> tuple[float | None, int | None]: if val_loss is None: return best_loss, best_step + _require_finite_loss(val_loss, label="validation") if best_loss is None or val_loss < best_loss: return val_loss, step return best_loss, best_step @@ -276,20 +310,23 @@ def checkpoint_payload(checkpoint_step: int) -> dict[str, object]: start_time = perf_counter() tokens_seen = 0 final_evaluation: EvaluationResult | None = None + early_stopping_state = EarlyStoppingState() + stopped_early = False with MetricsWriter(run_dir / "metrics.jsonl") as metrics: - while step < config.train.max_steps: + while step < config.train.max_steps and not stopped_early: for x, y in train_loader: model.train() x = x.to(device) y = y.to(device) _, loss = model(x, y) assert loss is not None + loss_value = _require_finite_loss(float(loss.item()), label="training") optimizer.zero_grad(set_to_none=True) loss.backward() optimizer.step() step += 1 tokens_seen += x.numel() - final_loss = float(loss.item()) + final_loss = loss_value should_log = step % config.train.log_interval == 0 or step == config.train.max_steps should_eval = val_tokens.numel() > 1 and step % config.train.eval_interval == 0 if should_log or should_eval: @@ -306,6 +343,8 @@ def checkpoint_payload(checkpoint_step: int) -> dict[str, object]: character_counts=val_character_counts_tensor, ) val_loss = evaluation.loss if evaluation else None + if val_loss is not None: + _require_finite_loss(val_loss, label="validation") final_evaluation = evaluation final_val_loss = val_loss best_val_loss, best_val_step = _update_best_validation( @@ -317,6 +356,12 @@ def checkpoint_payload(checkpoint_step: int) -> dict[str, object]: if best_val_step == step: best_evaluation = evaluation save_checkpoint(best_checkpoint_path, checkpoint_payload(step)) + early_stopping_state, stopped_early = _update_early_stopping( + early_stopping_state, + val_loss=val_loss, + patience=config.train.early_stopping_patience, + min_delta=config.train.early_stopping_min_delta, + ) record = _metrics_record( step=step, config=config, @@ -337,7 +382,7 @@ def checkpoint_payload(checkpoint_step: int) -> dict[str, object]: tokens_per_second=cast(float, record["tokens_per_second"]), ) metrics.write(record) - if step >= config.train.max_steps: + if step >= config.train.max_steps or stopped_early: break needs_final_eval = ( val_tokens.numel() > 1 and step > 0 and step % config.train.eval_interval != 0 @@ -445,6 +490,16 @@ def checkpoint_payload(checkpoint_step: int) -> dict[str, object]: else max(0, val_tokens.numel() - 1), "validation_coverage": final_evaluation.coverage if final_evaluation else 0.0, "max_steps": config.train.max_steps, + "actual_steps": step, + "stopped_early": stopped_early, + "stop_reason": "early_stopping" if stopped_early else "max_steps", + "early_stopping": { + "patience": config.train.early_stopping_patience, + "min_delta": config.train.early_stopping_min_delta, + "evaluations_without_improvement": ( + early_stopping_state.evaluations_without_improvement + ), + }, "device": str(device), "environment": { "python": platform.python_version(), diff --git a/tests/test_config.py b/tests/test_config.py index 2d96bb9..efb38ab 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -99,11 +99,20 @@ def test_data_config_rejects_invalid_bpe_vocab_size(): (lambda: TrainConfig(log_interval=0), "log_interval"), (lambda: TrainConfig(eval_interval=0), "eval_interval"), (lambda: TrainConfig(eval_batches=0), "eval_batches"), + (lambda: TrainConfig(early_stopping_patience=0), "early_stopping_patience"), + (lambda: TrainConfig(early_stopping_patience=True), "early_stopping_patience"), + (lambda: TrainConfig(early_stopping_min_delta=-0.1), "early_stopping_min_delta"), + (lambda: TrainConfig(early_stopping_min_delta=float("nan")), "early_stopping_min_delta"), + (lambda: TrainConfig(early_stopping_min_delta=10**1000), "early_stopping_min_delta"), (lambda: TrainConfig(learning_rate=0), "learning_rate"), + (lambda: TrainConfig(learning_rate=float("nan")), "learning_rate"), + (lambda: TrainConfig(learning_rate=float("inf")), "learning_rate"), (lambda: TrainConfig(weight_decay=-1), "weight_decay"), + (lambda: TrainConfig(weight_decay=float("nan")), "weight_decay"), (lambda: TrainConfig(sample_prompt=""), "sample_prompt"), (lambda: TrainConfig(sample_max_new_tokens=-1), "sample_max_new_tokens"), (lambda: TrainConfig(sample_temperature=0), "sample_temperature"), + (lambda: TrainConfig(sample_temperature=float("inf")), "sample_temperature"), (lambda: TrainConfig(sample_top_k=0), "sample_top_k"), ( lambda: ExperimentConfig( diff --git a/tests/test_training_artifacts.py b/tests/test_training_artifacts.py index 0e02d59..78bf1b6 100644 --- a/tests/test_training_artifacts.py +++ b/tests/test_training_artifacts.py @@ -37,6 +37,18 @@ def test_metrics_writer_writes_jsonl(tmp_path): assert json.loads(path.read_text(encoding="utf-8")) == {"step": 1, "train_loss": 3.0} +def test_artifact_json_rejects_non_finite_numbers(tmp_path): + metrics_path = tmp_path / "metrics.jsonl" + with MetricsWriter(metrics_path) as metrics, pytest.raises(ValueError): + metrics.write({"train_loss": float("nan")}) + assert metrics_path.read_text(encoding="utf-8") == "" + + summary_path = tmp_path / "summary.json" + with pytest.raises(ValueError): + write_json(summary_path, {"loss": float("inf")}) + assert not summary_path.exists() + + def test_write_config_snapshot_and_summary_json(tmp_path): config_path = tmp_path / "config.yaml" summary_path = tmp_path / "summary.json" diff --git a/tests/test_training_eval.py b/tests/test_training_eval.py index c92588d..2db958f 100644 --- a/tests/test_training_eval.py +++ b/tests/test_training_eval.py @@ -1,5 +1,6 @@ import json +import pytest import torch from torch.utils.data import DataLoader @@ -9,8 +10,11 @@ from smallm.model import GPT, GPTConfig from smallm.training import load_checkpoint from smallm.training.trainer import ( + EarlyStoppingState, + EvaluationResult, _build_optimizer, _update_best_validation, + _update_early_stopping, _validation_starts, estimate_loss, evaluate_tokens, @@ -18,6 +22,38 @@ ) +def test_early_stopping_resets_on_meaningful_improvement(): + state = EarlyStoppingState() + state, stopped = _update_early_stopping(state, val_loss=2.0, patience=2, min_delta=0.01) + assert not stopped + state, stopped = _update_early_stopping(state, val_loss=1.995, patience=2, min_delta=0.01) + assert not stopped + assert state.evaluations_without_improvement == 1 + + state, stopped = _update_early_stopping(state, val_loss=1.98, patience=2, min_delta=0.01) + assert not stopped + assert state == EarlyStoppingState(reference_loss=1.98) + + +def test_early_stopping_triggers_at_patience_and_can_be_disabled(): + state = EarlyStoppingState(reference_loss=1.0) + state, stopped = _update_early_stopping(state, val_loss=1.1, patience=2, min_delta=0.0) + assert not stopped + state, stopped = _update_early_stopping(state, val_loss=1.2, patience=2, min_delta=0.0) + assert stopped + + unchanged, stopped = _update_early_stopping(state, val_loss=1.3, patience=None, min_delta=0.0) + assert unchanged == state + assert not stopped + + +def test_early_stopping_rejects_non_finite_validation_loss(): + with pytest.raises(RuntimeError, match="non-finite"): + _update_early_stopping( + EarlyStoppingState(), val_loss=float("nan"), patience=1, min_delta=0.0 + ) + + def test_estimate_loss_returns_loss_and_restores_train_mode(): model = GPT(GPTConfig(vocab_size=8, block_size=4, n_layer=1, n_head=1, n_embd=8)) model.train() @@ -180,6 +216,130 @@ def test_train_records_final_validation_when_max_steps_misses_eval_interval(tmp_ assert load_checkpoint(run_dir / "best_checkpoint.pt")["step"] == summary["best_val_step"] assert metrics[-1]["step"] == 3 assert metrics[-1]["val_loss"] == summary["final_val_loss"] + assert summary["actual_steps"] == 3 + assert summary["stopped_early"] is False + assert summary["stop_reason"] == "max_steps" + + +def test_train_stops_after_configured_non_improving_evaluations(tmp_path, monkeypatch): + prepared_path = tmp_path / "corpus.txt" + manifest_path = tmp_path / "corpus_manifest.json" + text = "Once upon a time\n" * 8 + prepared_path.write_text(text, encoding="utf-8") + manifest_path.write_text( + json.dumps( + { + "source_name": "test corpus", + "prepared_sha256": file_sha256(prepared_path), + "prepared_characters": len(text), + "train_split": 0.8, + } + ), + encoding="utf-8", + ) + + def constant_evaluation(*args, **kwargs): + return EvaluationResult( + loss=1.0, + total_nll=4.0, + target_tokens=4, + total_target_tokens=4, + target_characters=4, + mode="sampled", + ) + + monkeypatch.setattr("smallm.training.trainer.evaluate_tokens", constant_evaluation) + config = ExperimentConfig( + data=DataConfig( + prepared_path=str(prepared_path), + manifest_path=str(manifest_path), + tokenizer_path=str(tmp_path / "tokenizer.json"), + block_size=4, + train_split=0.8, + ), + model=ModelConfig(vocab_size=256, block_size=4, n_layer=1, n_head=1, n_embd=8), + train=TrainConfig( + run_name="early_stop", + runs_dir=str(tmp_path / "runs"), + batch_size=2, + max_steps=5, + log_interval=1, + eval_interval=1, + eval_batches=1, + early_stopping_patience=1, + sample_max_new_tokens=1, + ), + ) + + checkpoint_path = train(config) + summary = json.loads((checkpoint_path.parent / "summary.json").read_text(encoding="utf-8")) + + assert summary["actual_steps"] == 2 + assert summary["stopped_early"] is True + assert summary["stop_reason"] == "early_stopping" + assert summary["best_val_step"] == 1 + assert load_checkpoint(checkpoint_path)["step"] == 2 + + def nan_evaluation(*args, **kwargs): + return EvaluationResult( + loss=float("nan"), + total_nll=float("nan"), + target_tokens=4, + total_target_tokens=4, + target_characters=4, + mode="sampled", + ) + + monkeypatch.setattr("smallm.training.trainer.evaluate_tokens", nan_evaluation) + nan_config = ExperimentConfig( + data=config.data, + model=config.model, + train=TrainConfig( + run_name="nan_loss", + runs_dir=str(tmp_path / "runs"), + batch_size=2, + max_steps=2, + log_interval=1, + eval_interval=1, + eval_batches=1, + early_stopping_patience=1, + sample_max_new_tokens=1, + ), + ) + + with pytest.raises(RuntimeError, match="non-finite"): + train(nan_config) + nan_run_dirs = list((tmp_path / "runs" / "nan_loss").iterdir()) + assert len(nan_run_dirs) == 1 + assert not (nan_run_dirs[0] / "summary.json").exists() + + optimizer_steps = 0 + + def nan_forward(self, idx, targets=None): + return torch.empty(0), torch.tensor(float("nan"), requires_grad=True) + + def count_optimizer_step(self, closure=None): + nonlocal optimizer_steps + optimizer_steps += 1 + + monkeypatch.setattr(GPT, "forward", nan_forward) + monkeypatch.setattr(torch.optim.AdamW, "step", count_optimizer_step) + training_nan_config = ExperimentConfig( + data=config.data, + model=config.model, + train=TrainConfig( + run_name="nan_training_loss", + runs_dir=str(tmp_path / "runs"), + batch_size=2, + max_steps=1, + log_interval=1, + eval_interval=1, + sample_max_new_tokens=1, + ), + ) + with pytest.raises(RuntimeError, match="training loss became non-finite"): + train(training_nan_config) + assert optimizer_steps == 0 def test_train_can_use_bpe_tokenizer(tmp_path):