File size: 5,230 Bytes
4be6a52 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 | """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
@pytest.fixture
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)
@pytest.mark.parametrize(
"arguments", [["--epochs", "3"], ["--learning-rate", "nan"], ["--accumulation", "0"]]
)
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)
|