stackcraft-clef-flash-lora / code /scripts /select_checkpoint.py
nima1's picture
Publish verified checkpoint and losslessly compressed study evidence
4be6a52 verified
Raw History Blame Contribute Delete
11.1 kB
"""Select only the two preregistered epochs using complete validation evidence."""
from __future__ import annotations
import argparse
import hashlib
import json
import shutil
from pathlib import Path
from typing import Any
from evaluate_clef import (
ENCODING_VERSION,
FINAL_TEST_SEEDS,
checkpoint_hashes,
file_hash,
validate_selection,
validation_ineligibility,
)
from stackcraft.data import audit_dataset, canonical_json
from stackcraft.evaluation import summarize_positions
from stackcraft.players import Decision
RULE = (
"lowest finite mean target NLL over all 215 validation positions; "
"exact tie chooses earlier epoch"
)
PLAYERS = ["base", "base-fp32", "trained", "random", "heuristic"]
TRAINING_CONFIG = {
"mode": "lora",
"epochs": 2,
"seed": 42,
"rank": 4,
"learning_rate": 1e-5,
"batch_size": 1,
"gradient_accumulation": 8,
"max_length": 4096,
"optimizer": "AdamW",
"weight_decay": 0.01,
"clip_gradient_norm": 1.0,
"label_smoothing": 0.05,
"brier_weight": 0.1,
}
def load_validation(dataset: Path) -> tuple[list[dict[str, Any]], str]:
manifest = json.loads((dataset / "manifest.json").read_text())
rows = {
split: [json.loads(line) for line in (dataset / f"{split}.jsonl").read_text().splitlines()]
for split in ("train", "validation")
}
audit_dataset(rows, manifest)
if len(rows["validation"]) != 215:
raise ValueError("selection requires all 215 frozen validation positions")
return rows["validation"], file_hash(dataset / "manifest.json")
def inspect_candidate(
checkpoint: Path,
report_path: Path,
epoch: int,
rows: list[dict[str, Any]],
manifest_sha256: str,
) -> dict[str, Any]:
hashes = checkpoint_hashes(checkpoint)
metadata = json.loads((checkpoint / "training_config.json").read_text())
report = json.loads(report_path.read_text())
extra = metadata.get("extra", {})
if extra.get("epoch") != epoch or extra.get("dataset_manifest_sha256") != manifest_sha256:
raise ValueError("candidate epoch/dataset differs from preregistered study")
if any(extra.get("config", {}).get(key) != value for key, value in TRAINING_CONFIG.items()):
raise ValueError("candidate training configuration differs from preregistration")
if (
extra.get("test_trajectories_used") is not False
or extra.get("validation_used_for_training") is not False
):
raise ValueError("candidate training metadata must exclude test and validation use")
if (
report.get("checkpoint_sha256") != hashes
or report.get("dataset_manifest_sha256") != manifest_sha256
):
raise ValueError("validation report is not bound to this checkpoint and dataset")
if (
report.get("positions") != 215
or report.get("mode") != "positions"
or report.get("split") != "validation"
):
raise ValueError("candidate report must cover the complete validation split")
trained = report.get("players", {}).get("trained")
if not isinstance(trained, dict):
raise ValueError("candidate report lacks trained predictions")
predictions_path = report_path.parent / trained["predictions_file"]
if file_hash(predictions_path) != trained["predictions_sha256"]:
raise ValueError("raw prediction hash differs from validation report")
events = [json.loads(line) for line in predictions_path.read_text().splitlines()]
if [event.get("id") for event in events] != [row["id"] for row in rows]:
raise ValueError("raw predictions omit, duplicate, reorder or add validation positions")
predictions = {}
for row, event in zip(rows, events, strict=True):
if event.get("target_action_id") != row["action_id"]:
raise ValueError("prediction target differs from frozen teacher label")
predictions[row["id"]] = None if "error" in event else Decision(**event["decision"])
recomputed = summarize_positions(rows, predictions)
if trained.get("metrics") != recomputed:
raise ValueError("reported validation metrics differ from raw predictions")
reasons = validation_ineligibility(report, metadata)
return {
"key": f"epoch-{epoch:02d}",
"epoch": epoch,
"checkpoint_sha256": hashes,
"metrics": recomputed,
"ineligibility_reasons": reasons,
"checkpoint": checkpoint,
"report_path": report_path,
"report": report,
"metadata": metadata,
}
def select_checkpoint(
candidates: list[tuple[Path, Path]], dataset: Path, output: Path
) -> dict[str, Any]:
if len(candidates) != 2:
raise ValueError("exactly epoch 01 and epoch 02 candidates are required")
if output.exists():
raise FileExistsError("selection output already exists")
rows, manifest_hash = load_validation(dataset)
inspected = [
inspect_candidate(checkpoint, report, index + 1, rows, manifest_hash)
for index, (checkpoint, report) in enumerate(candidates)
]
first, second = inspected
for key in ("source_hashes", "config", "dataset_split_sha256", "dataset_counts"):
if first["metadata"]["extra"].get(key) != second["metadata"]["extra"].get(key):
raise ValueError(f"candidate training provenance differs: {key}")
for key in ("source_hashes", "installed_versions", "encoding_version", "base_revision"):
if first["report"].get("provenance", {}).get(key) != second["report"].get(
"provenance", {}
).get(key):
raise ValueError(f"candidate evaluation provenance differs: {key}")
runtime = [
candidate["report"]["players"]["trained"]["runtime_config"] for candidate in inspected
]
for key in (
"model_id",
"base_revision",
"encoding_version",
"source_sha256",
"max_length",
"dtype",
"device",
"head_dtype",
"versions",
):
if runtime[0].get(key) != runtime[1].get(key):
raise ValueError(f"candidate evaluation runtime differs: {key}")
output.mkdir(parents=True)
preserved: list[dict[str, Any]] = []
for candidate in inspected:
folder = output / candidate["key"]
folder.mkdir()
shutil.copyfile(candidate["report_path"], folder / "report.json")
shutil.copyfile(
candidate["checkpoint"] / "training_config.json", folder / "training_config.json"
)
for player in candidate["report"]["players"].values():
relative = Path(player["predictions_file"])
if relative.is_absolute() or ".." in relative.parts:
raise ValueError("prediction evidence must use a contained relative path")
source = candidate["report_path"].parent / relative
if file_hash(source) != player["predictions_sha256"]:
raise ValueError("diagnostic prediction file hash mismatch")
target = folder / relative
target.parent.mkdir(parents=True, exist_ok=True)
shutil.copyfile(source, target)
preserved.append(
{
"key": candidate["key"],
"epoch": candidate["epoch"],
"checkpoint_sha256": candidate["checkpoint_sha256"],
"checkpoint_origin": str(candidate["checkpoint"].resolve()),
"checkpoint_metadata": {
"path": f"{candidate['key']}/training_config.json",
"sha256": file_hash(folder / "training_config.json"),
},
"validation_evidence": {
"path": f"{candidate['key']}/report.json",
"sha256": file_hash(folder / "report.json"),
},
"metrics": candidate["metrics"],
"ineligibility_reasons": candidate["ineligibility_reasons"],
}
)
eligible = [candidate for candidate in preserved if not candidate["ineligibility_reasons"]]
audit = {
"decision_rule": RULE,
"dataset_manifest_sha256": manifest_hash,
"candidates": preserved,
"status": "selected" if eligible else "no-eligible-candidate",
}
(output / "selection-audit.json").write_text(
json.dumps(audit, indent=2, allow_nan=False) + "\n"
)
if not eligible:
raise ValueError(
"both epoch candidates are ineligible; evidence preserved, no test selection"
)
selected = min(
eligible, key=lambda candidate: (candidate["metrics"]["mean_nll"], candidate["epoch"])
)
selection = {
"schema_version": 1,
"selection_split": "validation",
"selected_key": selected["key"],
"decision_rule": RULE,
"checkpoint_sha256": selected["checkpoint_sha256"],
"validation_metrics": selected["metrics"],
"validation_evidence": selected["validation_evidence"],
"candidates": preserved,
"dataset_manifest_sha256": manifest_hash,
"test_seeds": list(FINAL_TEST_SEEDS),
"test_seeds_sha256": hashlib.sha256(
canonical_json(list(FINAL_TEST_SEEDS)).encode()
).hexdigest(),
"max_pieces": 200,
"encoding_version": ENCODING_VERSION,
"max_length": 4096,
"players": PLAYERS,
"head_dtypes": {"base": "bfloat16", "base-fp32": "float32", "trained": "float32"},
"bootstrap_samples": 10000,
"bootstrap_seed": 2026,
"max_error_rate": 0.0,
}
selection_path = output / "selection.json"
selection_path.write_text(json.dumps(selection, indent=2, allow_nan=False) + "\n")
args = argparse.Namespace(
final_test=True,
seeds=FINAL_TEST_SEEDS,
max_pieces=200,
players=PLAYERS,
max_length=4096,
selection_file=selection_path,
checkpoint=Path(selected["checkpoint_origin"]),
)
validate_selection(args, selected["checkpoint_sha256"])
return selection
def main(argv: list[str] | None = None) -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--candidate01", nargs=2, type=Path, metavar=("CHECKPOINT", "REPORT"), required=True
)
parser.add_argument(
"--candidate02", nargs=2, type=Path, metavar=("CHECKPOINT", "REPORT"), required=True
)
parser.add_argument("--dataset", type=Path, default=Path("data/study-v1"))
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args(argv)
try:
selection = select_checkpoint(
[tuple(args.candidate01), tuple(args.candidate02)], args.dataset, args.output
)
except (ValueError, OSError, KeyError) as error:
parser.error(str(error))
print(
json.dumps(
{
"selected_key": selection["selected_key"],
"mean_nll": selection["validation_metrics"]["mean_nll"],
"output": str(args.output),
}
)
)
return 0
if __name__ == "__main__":
raise SystemExit(main())