Mindcraft's picture
Publish frozen reproduction executor image source
17d5066 verified
Raw
History Blame Contribute Delete
7.77 kB
#!/usr/bin/env python3
"""Scientific training contract gate for one registered isolated seed."""
from __future__ import annotations
import argparse
import json
import os
import pickle
import subprocess
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from repro_control.archives import safe_extract_tar_gz
from repro_control.artifacts import finalize_attempt, verify_read_back
from repro_control.configs import (
config_sha256,
load_json_config,
registered_config_for,
verify_before_optimizer,
)
from repro_control.hashing import atomic_write_json, sha256_file
from repro_control.heartbeat import heartbeat, marker
from repro_control.interventions import contains_restriction_map
from repro_control.runtime import (
configure_scientific_runtime,
load_job_manifest,
require_control_freeze,
require_jax_platform,
verify_science_spec,
)
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--science-spec", type=Path, required=True)
parser.add_argument("--job-manifest", type=Path, required=True)
parser.add_argument("--control-dir", type=Path, required=True)
parser.add_argument("--freeze-sha256", required=True)
parser.add_argument("--data-dir", type=Path)
parser.add_argument("--train-archive", type=Path)
parser.add_argument("--output-dir", type=Path)
parser.add_argument("--dry-run", action="store_true")
args = parser.parse_args()
spec = verify_science_spec(args.science_spec)
manifest = load_job_manifest(args.job_manifest, freeze_sha256=args.freeze_sha256)
if manifest["job_class"] != "SCIENTIFIC_TRAIN":
raise SystemExit("training entrypoint requires SCIENTIFIC_TRAIN")
if manifest["logical_id"] not in spec["run_identities"]:
raise SystemExit("unregistered training logical identity")
require_control_freeze(args.control_dir, args.freeze_sha256, manifest["science_spec_sha256"])
configure_scientific_runtime()
if args.dry_run:
print(
json.dumps(
{
"contract_verified": True,
"logical_id": manifest["logical_id"],
"outcomes": {},
}
)
)
return 0
if args.data_dir is None or args.output_dir is None:
raise SystemExit("--data-dir and --output-dir are required for execution")
if args.train_archive:
safe_extract_tar_gz(args.train_archive, args.data_dir)
require_jax_platform("gpu")
marker("GPU_READY", manifest["logical_id"])
logical_id = manifest["logical_id"]
seed = manifest["seed"]
expected_path = args.control_dir / "registered-configs" / f"{logical_id}.json"
expected = load_json_config(expected_path)
expected_sha256 = config_sha256(expected)
if expected_sha256 != manifest["hashes"]["config"]:
raise SystemExit("registered config hash differs from the Job manifest")
independently_composed = registered_config_for(
Path(__file__).resolve().parents[1],
logical_id,
data_root=str(args.data_dir),
)
verify_before_optimizer(
independently_composed,
expected,
expected_sha256=expected_sha256,
)
args.output_dir.mkdir(parents=True, exist_ok=False)
with heartbeat(f"SCIENTIFIC_TRAIN:{logical_id}"):
if logical_id.startswith("C3-MNIST-CNN-"):
from repro_control.cnn_training import train_fixed_cnn
train_fixed_cnn(
Path(expected["data"]["dir"]),
args.output_dir,
seed=seed,
registered_config=expected,
registered_config_sha256=expected_sha256,
)
else:
root = Path(__file__).resolve().parents[1]
if logical_id.startswith("C1-SUD-MPNN225-"):
experiment = "sudoku_mpnn"
overrides = [
"model.d_v=225",
"training.lr=1.7e-3",
"training.epochs=10",
]
elif logical_id.startswith("C2-MAZE-MPNN84-"):
experiment = "maze_mpnn"
overrides = []
elif logical_id.startswith("C4-SUD-IDENTITY-"):
experiment = "sudoku_sheaf"
overrides = ["model.rm_init=identity", "+model.rm_constant=true"]
elif logical_id.startswith("C4-MAZE-QUADRATIC-"):
experiment = "maze_sheaf"
overrides = ["model.objective_mode=quadratic"]
else:
raise SystemExit("no registered training backend")
command = [
sys.executable,
str(root / "scripts" / "train.py"),
f"+experiment={experiment}",
f"training.seed={seed}",
f"data.dir={expected['data']['dir']}",
"data.val_splits=[]",
"wandb.mode=disabled",
f"hydra.run.dir={args.output_dir}",
*overrides,
]
env = os.environ.copy()
env.update(
REPRO_EXPECTED_CONFIG_PATH=str(expected_path),
REPRO_EXPECTED_CONFIG_SHA256=expected_sha256,
)
subprocess.run(command, check=True, env=env)
checkpoint_path = args.output_dir / "checkpoint.pkl"
with checkpoint_path.open("rb") as handle:
checkpoint = pickle.load(handle)
if handle.read(1):
raise SystemExit("checkpoint contains trailing bytes")
if checkpoint.get("ema_params") is None:
raise SystemExit("training checkpoint lacks EMA parameters")
import jax
parameter_count = sum(
int(leaf.size) for leaf in jax.tree_util.tree_leaves(checkpoint["params"])
)
identity_parity = None
if logical_id.startswith("C4-SUD-IDENTITY-"):
parity_path = args.output_dir / "identity-parity.json"
if not parity_path.is_file():
raise SystemExit("identity training lacks the pre-step parity receipt")
identity_parity = json.loads(parity_path.read_text())
if not all(
identity_parity.get(name) is True
for name in (
"common_parameter_leaves_bit_identical",
"common_optimizer_leaves_bit_identical",
"counters_bit_identical",
"checked_before_first_optimizer_step",
)
):
raise SystemExit("identity pre-step parity receipt is not affirmative")
if contains_restriction_map(checkpoint["params"]):
raise SystemExit("identity checkpoint unexpectedly contains restriction-map parameters")
receipt = {
"format": 1,
"logical_id": logical_id,
"attempt_id": manifest["attempt_id"],
"seed": seed,
"registered_config_sha256": expected_sha256,
"checkpoint_sha256": sha256_file(checkpoint_path),
"history_sha256": sha256_file(args.output_dir / "history.json"),
"parameter_count": parameter_count,
"ema": True,
"fixed_final_epoch": int(expected["training"]["epochs"]) - 1,
"identity_common_parity": identity_parity,
"outcomes": {},
}
atomic_write_json(args.output_dir / "training-receipt.json", receipt)
finalize_attempt(
args.output_dir,
logical_id=logical_id,
attempt_id=manifest["attempt_id"],
expected_outputs=manifest["expected_outputs"],
)
verify_read_back(
args.output_dir,
logical_id=logical_id,
attempt_id=manifest["attempt_id"],
)
marker("DONE", f"{logical_id} {manifest['attempt_id']}")
return 0
if __name__ == "__main__":
raise SystemExit(main())