Download code/tests/test_train_cli.py from nima1/stackcraft-clef-flash-lora: direct link, hf CLI and curl.
- Browser
- Download file 5.23 kB
-
https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/tests/test_train_cli.py
- Command line
-
hf download hf://nima1/stackcraft-clef-flash-lora/code/tests/test_train_cli.py
-
curl -L -o test_train_cli.py https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/tests/test_train_cli.py
5.23 kB
| """Training orchestration checks; no GPU or pretrained weights are loaded.""" | |
| import hashlib | |
| import importlib.util | |
| import json | |
| import sys | |
| from pathlib import Path | |
| from types import SimpleNamespace | |
| import pytest | |
| def script(): | |
| path = Path(__file__).resolve().parents[1] / "scripts/train_clef.py" | |
| spec = importlib.util.spec_from_file_location("stackcraft_test_train_cli", path) | |
| assert spec is not None and spec.loader is not None | |
| module = importlib.util.module_from_spec(spec) | |
| sys.modules[spec.name] = module | |
| spec.loader.exec_module(module) | |
| return module | |
| def test_epoch_covers_every_training_row_and_correct_partial_group(script): | |
| first = script.accumulation_groups(827, 8, epoch=1) | |
| assert len(first) == 104 | |
| assert len(first[-1]) == 3 | |
| assert sorted(index for group in first for index in group) == list(range(827)) | |
| assert first == script.accumulation_groups(827, 8, epoch=1) | |
| assert first != script.accumulation_groups(827, 8, epoch=2) | |
| def test_invalid_configuration_fails_before_any_model_load(script, tmp_path, arguments): | |
| with pytest.raises(SystemExit) as error: | |
| script.main(["--output", str(tmp_path / "new"), *arguments]) | |
| assert error.value.code == 2 | |
| assert not (tmp_path / "new").exists() | |
| def test_existing_output_is_never_overwritten(script, tmp_path): | |
| with pytest.raises(SystemExit) as error: | |
| script.main(["--output", str(tmp_path)]) | |
| assert error.value.code == 2 | |
| def test_only_fixed_train_and_validation_are_read_and_audited(script, tmp_path, monkeypatch): | |
| train = [{"id": "train-a"}, {"id": "train-b"}] | |
| validation = [{"id": "validation-a"}] | |
| hashes = {} | |
| for split, rows in (("train", train), ("validation", validation)): | |
| content = "".join(json.dumps(row) + "\n" for row in rows) | |
| (tmp_path / f"{split}.jsonl").write_text(content) | |
| hashes[split] = hashlib.sha256(content.encode()).hexdigest() | |
| (tmp_path / "test.jsonl").write_text("this file must never be opened as JSON") | |
| manifest = {"source_commit": "test-source", "config_sha256": "test-config"} | |
| (tmp_path / "manifest.json").write_text(json.dumps(manifest)) | |
| monkeypatch.setattr(script, "STUDY_HASHES", hashes) | |
| monkeypatch.setattr(script, "STUDY_COUNTS", {"train": 2, "validation": 1}) | |
| calls = [] | |
| monkeypatch.setattr(script, "audit_dataset", lambda records, manifest: calls.append(records)) | |
| loaded, metadata = script.load_study(tmp_path) | |
| assert loaded == train | |
| assert calls == [{"train": train, "validation": validation}] | |
| assert metadata["test_trajectories_used"] is False | |
| assert metadata["validation_used_for_training"] is False | |
| (tmp_path / "train.jsonl").write_text('{"id":"tampered"}\n') | |
| with pytest.raises(ValueError, match="frozen study-v1 SHA256"): | |
| script.load_study(tmp_path) | |
| def test_native_loss_accumulation_matches_actual_group_mean(script, tmp_path, monkeypatch): | |
| torch = pytest.importorskip("torch") | |
| from stackcraft.training import decision_loss | |
| class TinyModel(torch.nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.logits = torch.nn.Parameter(torch.zeros(2)) | |
| self.language_model = torch.nn.Identity() | |
| def forward(self, batch): | |
| return [[self.logits]] | |
| encoded = SimpleNamespace( | |
| input_ids=(1, 2), | |
| questions=(SimpleNamespace(question_type=1, option_ids=("r0x0", "r0x1")),), | |
| ) | |
| monkeypatch.setattr(script, "row_observation", lambda row: row) | |
| monkeypatch.setattr(script, "encode_observation", lambda *args: encoded) | |
| model = TinyModel() | |
| reference = TinyModel() | |
| rows = [{"id": f"row-{index}", "action_id": "r0x1"} for index in range(5)] | |
| player = SimpleNamespace( | |
| model=model, | |
| processor=SimpleNamespace(tokenizer=SimpleNamespace(pad_token_id=0)), | |
| native=SimpleNamespace(collate_records=lambda *args: {}), | |
| max_length=4096, | |
| ) | |
| optimizer = torch.optim.SGD(model.parameters(), lr=0.1) | |
| expected_optimizer = torch.optim.SGD(reference.parameters(), lr=0.1) | |
| for _ in range(2): | |
| expected_optimizer.zero_grad(set_to_none=True) | |
| decision_loss(reference.logits, encoded, "r0x1").backward() | |
| torch.nn.utils.clip_grad_norm_(reference.parameters(), 1.0, error_if_nonfinite=True) | |
| expected_optimizer.step() | |
| outcome = script.train_epoch( | |
| player, rows, optimizer, epoch=1, accumulation=3, mode="head", output=tmp_path | |
| ) | |
| torch.testing.assert_close(model.logits, reference.logits, rtol=0, atol=1e-8) | |
| assert outcome["examples"] == 5 | |
| assert outcome["optimizer_steps"] == 2 | |
| events = [ | |
| json.loads(line) for line in (tmp_path / "epoch-01-events.jsonl").read_text().splitlines() | |
| ] | |
| microsteps = [event for event in events if event["event"] == "microstep"] | |
| assert {event["row_id"] for event in microsteps} == {row["id"] for row in rows} | |
| assert [event["accumulation_group_size"] for event in microsteps] == [3, 3, 3, 2, 2] | |
| assert all(event["peak_allocated_bytes"] == 0 for event in microsteps) | |