spec-b300 / source /tests /unit /train /test_save_train_command.py
khazic's picture
Archive three-epoch run: logs and provenance part 5
2dd5f57 verified
Raw History Blame Contribute Delete
6.79 kB
import os
import sys
from pathlib import Path
from unittest.mock import patch
from speculators.train.checkpointer import SingleGPUCheckpointer
from speculators.train.utils import save_train_command
# ---------------------------------------------------------------------------
# save_train_command tests
# ---------------------------------------------------------------------------
class TestSaveTrainCommand:
def test_creates_file(self, tmp_path: Path):
save_train_command(str(tmp_path))
assert (tmp_path / "train_command.txt").exists()
def test_creates_directory_if_missing(self, tmp_path: Path):
save_path = tmp_path / "nested" / "dir"
save_train_command(str(save_path))
assert (save_path / "train_command.txt").exists()
def test_contains_sys_argv(self, tmp_path: Path):
with patch.object(sys, "argv", ["scripts/train.py", "--lr", "1e-4"]):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "scripts/train.py --lr 1e-4" in content
def test_header_has_timestamp(self, tmp_path: Path):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "# Timestamp:" in content
def test_header_has_git_sha(self, tmp_path: Path):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "# Git SHA:" in content
def test_header_has_world_size(self, tmp_path: Path):
with patch.dict(os.environ, {"WORLD_SIZE": "8"}):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "# World size: 8" in content
def test_world_size_defaults_to_1(self, tmp_path: Path):
env = os.environ.copy()
env.pop("WORLD_SIZE", None)
with patch.dict(os.environ, env, clear=True):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "# World size: 1" in content
def test_header_has_package_versions(self, tmp_path: Path):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
for pkg in ("speculators", "transformers", "torch"):
assert f"# {pkg}:" in content
def test_git_sha_fallback_on_error(self, tmp_path: Path):
with patch(
"speculators.train.utils.git_sha",
return_value="unknown",
):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "# Git SHA: unknown" in content
def test_quotes_args_with_spaces(self, tmp_path: Path):
with patch.object(sys, "argv", ["train.py", "--path", "/has spaces/dir"]):
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "'/has spaces/dir'" in content
def test_no_leftover_tmp_files(self, tmp_path: Path):
save_train_command(str(tmp_path))
tmp_files = [
f for f in tmp_path.iterdir() if f.name.startswith(".train_command_")
]
assert tmp_files == []
def test_overwrites_existing(self, tmp_path: Path):
(tmp_path / "train_command.txt").write_text("old content")
save_train_command(str(tmp_path))
content = (tmp_path / "train_command.txt").read_text()
assert "old content" not in content
assert "# Timestamp:" in content
# ---------------------------------------------------------------------------
# speculators.patch tests
# ---------------------------------------------------------------------------
class TestSpeculatorsPatch:
def test_creates_patch_file(self, tmp_path: Path):
save_train_command(str(tmp_path))
assert (tmp_path / "speculators.patch").exists()
def test_patch_contains_repo_header(self, tmp_path: Path):
save_train_command(str(tmp_path))
content = (tmp_path / "speculators.patch").read_text()
assert content.startswith("# repo: ")
def test_patch_contains_sha(self, tmp_path: Path):
save_train_command(str(tmp_path))
content = (tmp_path / "speculators.patch").read_text()
first_line = content.split("\n")[0]
assert "(" in first_line
assert ")" in first_line
def test_no_patch_when_no_repo(self, tmp_path: Path):
with patch(
"speculators.train.utils.find_repo_root",
return_value=None,
):
save_train_command(str(tmp_path))
assert not (tmp_path / "speculators.patch").exists()
def test_patch_failure_does_not_block(self, tmp_path: Path):
with patch(
"speculators.train.utils.git_diff",
side_effect=OSError("git broke"),
):
save_train_command(str(tmp_path))
assert (tmp_path / "train_command.txt").exists()
assert not (tmp_path / "speculators.patch").exists()
# ---------------------------------------------------------------------------
# _copy_train_command tests (checkpointer)
# ---------------------------------------------------------------------------
class TestCopyTrainCommand:
def test_copies_into_epoch_dir(self, tmp_path: Path):
src_content = "# test content\ntrain.py --lr 1e-4\n"
(tmp_path / "train_command.txt").write_text(src_content)
(tmp_path / "0").mkdir()
cp = SingleGPUCheckpointer(str(tmp_path))
cp._copy_train_command(0)
copied = tmp_path / "0" / "train_command.txt"
assert copied.exists()
assert copied.read_text() == src_content
def test_noop_when_source_missing(self, tmp_path: Path):
(tmp_path / "0").mkdir()
cp = SingleGPUCheckpointer(str(tmp_path))
cp._copy_train_command(0)
assert not (tmp_path / "0" / "train_command.txt").exists()
def test_copies_into_string_epoch(self, tmp_path: Path):
(tmp_path / "train_command.txt").write_text("content")
(tmp_path / "interrupted").mkdir()
cp = SingleGPUCheckpointer(str(tmp_path))
cp._copy_train_command("interrupted")
assert (tmp_path / "interrupted" / "train_command.txt").exists()
def test_cleanup_keep_only_best_preserves_train_command(self, tmp_path: Path):
(tmp_path / "train_command.txt").write_text("content")
(tmp_path / "0").mkdir()
(tmp_path / "1").mkdir()
(tmp_path / "1" / "model.safetensors").touch()
cp = SingleGPUCheckpointer(str(tmp_path))
cp.update_best_symlink(1)
cp.cleanup_keep_only_best(1)
assert (tmp_path / "train_command.txt").exists()
assert not (tmp_path / "0").exists()