Download code/scripts/select_checkpoint.py from nima1/stackcraft-clef-flash-lora: direct link, hf CLI and curl.
- Browser
- Download file 11.1 kB
-
https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/scripts/select_checkpoint.py
- Command line
-
hf download hf://nima1/stackcraft-clef-flash-lora/code/scripts/select_checkpoint.py
-
curl -L -o select_checkpoint.py https://huggingface.co/nima1/stackcraft-clef-flash-lora/resolve/main/code/scripts/select_checkpoint.py
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()) | |