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)