diff --git a/docs/TRAINING.md b/docs/TRAINING.md index 3509d45..d5a270d 100644 --- a/docs/TRAINING.md +++ b/docs/TRAINING.md @@ -42,7 +42,7 @@ uv run --no-sync python -m micro_scout.train \ --output runs/minilm-v1 --resume runs/minilm-v1/last --device cuda ``` -Resume requires the same training configuration and dataset manifest. This command is intended for the project's own local optimizer states. Model weights use safetensors, and remote custom model code is disabled. +Resume requires the same training configuration, dataset manifest, and prepared-file hashes. A weights checksum prevents resuming mismatched weights and optimizer state after an incomplete checkpoint write. This command is intended for the project's own local optimizer states. Model weights use safetensors, and remote custom model code is disabled. ## Evaluation protocol diff --git a/src/micro_scout/train.py b/src/micro_scout/train.py index ab45072..3ab1c99 100644 --- a/src/micro_scout/train.py +++ b/src/micro_scout/train.py @@ -119,6 +119,12 @@ def train(data: Path, output: Path, config: dict, device: str, resume: Path | No optimizer_steps, skipped_optimizer_steps = 0, 0 if resume: state = torch.load(resume / "training_state.pt", map_location="cpu", weights_only=True) + weight_hashes = { + path.name: hashlib.sha256(path.read_bytes()).hexdigest() + for path in resume.glob("*.safetensors") + } + if weight_hashes != state["weights_sha256"]: + raise ValueError("Checkpoint weights and optimizer state do not match") if ( state["config"] != config or state["dataset_manifest_sha256"] != manifest_hash @@ -181,6 +187,10 @@ def train(data: Path, output: Path, config: dict, device: str, resume: Path | No "scaler": scaler.state_dict(), "torch_rng": torch.get_rng_state(), "cuda_rng": torch.cuda.get_rng_state_all() if use_cuda else [], + "weights_sha256": { + weight.name: hashlib.sha256(weight.read_bytes()).hexdigest() + for weight in path.glob("*.safetensors") + }, } torch.save(state, path / "training_state.pt.tmp") os.replace(path / "training_state.pt.tmp", path / "training_state.pt") diff --git a/tests/test_encoder.py b/tests/test_encoder.py index d1752ed..7b4c6b0 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -114,3 +114,8 @@ def test_training_updates_weights_and_can_resume(tiny_model, tmp_path): state = torch.load(output / "last/training_state.pt", weights_only=True) assert state["step"] == 2 assert json.loads((output / "result.json").read_text())["test_set_used_for_selection"] is False + with torch.no_grad(): + after.model.embeddings.word_embeddings.weight.add_(0.1) + after.save(output / "last") + with pytest.raises(ValueError, match="weights and optimizer state"): + train(data, output, config, "cpu", output / "last")