221 lines
6.6 KiB
Python
221 lines
6.6 KiB
Python
#!/usr/bin/env python3
|
|
"""Run the frozen Round 08 matrix with at most two isolated processes."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import os
|
|
import subprocess
|
|
import time
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
|
|
VARIANTS = (
|
|
"uniform_group_6_forward",
|
|
"uniform_group_7_forward",
|
|
"uniform_groups_6_7_forward",
|
|
"uniform_group_7_mlp_forward",
|
|
)
|
|
SEEDS = (2026073001, 2026073002, 2026073003)
|
|
|
|
|
|
def parse_args() -> argparse.Namespace:
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--python", type=Path, required=True)
|
|
parser.add_argument("--cache-dir", type=Path, required=True)
|
|
parser.add_argument("--parent-manifest", type=Path, required=True)
|
|
parser.add_argument("--study-manifest", type=Path, required=True)
|
|
parser.add_argument("--output-dir", type=Path, required=True)
|
|
parser.add_argument(
|
|
"--phase", choices=("formal", "replay", "all"), default="all"
|
|
)
|
|
parser.add_argument("--concurrency", type=int, default=2)
|
|
return parser.parse_args()
|
|
|
|
|
|
def cell_output(
|
|
output_dir: Path, variant: str, seed: int, run_kind: str
|
|
) -> Path:
|
|
return output_dir / (
|
|
f"{run_kind}-{variant}-seed-{seed}.json"
|
|
)
|
|
|
|
|
|
def command_for(
|
|
args: argparse.Namespace, variant: str, seed: int, run_kind: str
|
|
) -> list[str]:
|
|
runner = Path(__file__).resolve().parent / "train.py"
|
|
return [
|
|
str(args.python),
|
|
str(runner),
|
|
"--variant",
|
|
variant,
|
|
"--study-manifest",
|
|
str(args.study_manifest),
|
|
"--run-kind",
|
|
run_kind,
|
|
"--architecture",
|
|
"block",
|
|
"--depth",
|
|
"32",
|
|
"--seed",
|
|
str(seed),
|
|
"--cache-dir",
|
|
str(args.cache_dir),
|
|
"--manifest",
|
|
str(args.parent_manifest),
|
|
"--output",
|
|
str(cell_output(args.output_dir, variant, seed, run_kind)),
|
|
]
|
|
|
|
|
|
def validate_manifest(args: argparse.Namespace) -> None:
|
|
manifest = json.loads(args.study_manifest.read_text())
|
|
if (
|
|
manifest["status"] != "frozen-before-model-output"
|
|
or tuple(manifest["variants"]) != VARIANTS
|
|
or tuple(manifest["formal_seeds"]) != SEEDS
|
|
or manifest["concurrency_maximum"] != 2
|
|
):
|
|
raise RuntimeError("study manifest matrix/concurrency drift")
|
|
if args.concurrency < 1 or args.concurrency > 2:
|
|
raise ValueError("the frozen protocol permits one or two processes")
|
|
|
|
|
|
def stop_processes(items: list[dict[str, Any]]) -> None:
|
|
for item in items:
|
|
if item["process"].poll() is None:
|
|
item["process"].terminate()
|
|
for item in items:
|
|
process = item["process"]
|
|
if process.poll() is not None:
|
|
continue
|
|
try:
|
|
process.wait(timeout=10)
|
|
except subprocess.TimeoutExpired:
|
|
process.kill()
|
|
process.wait()
|
|
|
|
|
|
def quarantine_failed_output(
|
|
output_dir: Path, variant: str, seed: int, run_kind: str
|
|
) -> str | None:
|
|
output = cell_output(output_dir, variant, seed, run_kind)
|
|
if not output.exists():
|
|
return None
|
|
failed = output.with_suffix(".failed.json")
|
|
if failed.exists():
|
|
failed = output.with_suffix(f".failed-{time.time_ns()}.json")
|
|
output.replace(failed)
|
|
return str(failed)
|
|
|
|
|
|
def run_cells(
|
|
args: argparse.Namespace,
|
|
cells: list[tuple[str, int, str]],
|
|
) -> None:
|
|
args.output_dir.mkdir(parents=True, exist_ok=True)
|
|
for variant, seed, run_kind in cells:
|
|
output = cell_output(args.output_dir, variant, seed, run_kind)
|
|
if output.exists():
|
|
raise FileExistsError(
|
|
f"refusing to overwrite existing result: {output}"
|
|
)
|
|
|
|
environment = dict(os.environ)
|
|
environment["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
|
|
pending = list(cells)
|
|
running: list[dict[str, Any]] = []
|
|
completed = 0
|
|
while pending or running:
|
|
while pending and len(running) < args.concurrency:
|
|
variant, seed, run_kind = pending.pop(0)
|
|
command = command_for(args, variant, seed, run_kind)
|
|
process = subprocess.Popen(command, env=environment)
|
|
running.append(
|
|
{
|
|
"identity": (variant, seed, run_kind),
|
|
"process": process,
|
|
"started": time.monotonic(),
|
|
}
|
|
)
|
|
print(
|
|
json.dumps(
|
|
{
|
|
"event": "cell_started",
|
|
"variant": variant,
|
|
"seed": seed,
|
|
"run_kind": run_kind,
|
|
"pid": process.pid,
|
|
"active": len(running),
|
|
"remaining": len(pending),
|
|
},
|
|
sort_keys=True,
|
|
),
|
|
flush=True,
|
|
)
|
|
time.sleep(1)
|
|
survivors = []
|
|
for item in running:
|
|
return_code = item["process"].poll()
|
|
if return_code is None:
|
|
survivors.append(item)
|
|
continue
|
|
variant, seed, run_kind = item["identity"]
|
|
elapsed = time.monotonic() - item["started"]
|
|
if return_code != 0:
|
|
stop_processes(
|
|
[candidate for candidate in running if candidate is not item]
|
|
)
|
|
quarantined = quarantine_failed_output(
|
|
args.output_dir, variant, seed, run_kind
|
|
)
|
|
raise RuntimeError(
|
|
f"cell failed: {variant}/{seed}/{run_kind}: {return_code}; "
|
|
f"quarantined_output={quarantined}"
|
|
)
|
|
completed += 1
|
|
print(
|
|
json.dumps(
|
|
{
|
|
"event": "cell_completed",
|
|
"variant": variant,
|
|
"seed": seed,
|
|
"run_kind": run_kind,
|
|
"elapsed_seconds": elapsed,
|
|
"completed": completed,
|
|
"total": len(cells),
|
|
},
|
|
sort_keys=True,
|
|
),
|
|
flush=True,
|
|
)
|
|
running = survivors
|
|
|
|
|
|
def main() -> None:
|
|
args = parse_args()
|
|
validate_manifest(args)
|
|
formal = [
|
|
(variant, seed, "formal")
|
|
for variant in VARIANTS
|
|
for seed in SEEDS
|
|
]
|
|
replay = [
|
|
("uniform_groups_6_7_forward", 2026073001, "replay")
|
|
]
|
|
cells = (
|
|
formal
|
|
if args.phase == "formal"
|
|
else replay
|
|
if args.phase == "replay"
|
|
else formal + replay
|
|
)
|
|
run_cells(args, cells)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|