diff options
Diffstat (limited to 'experiments/validate_crossover_artifact.py')
| -rw-r--r-- | experiments/validate_crossover_artifact.py | 73 |
1 files changed, 73 insertions, 0 deletions
diff --git a/experiments/validate_crossover_artifact.py b/experiments/validate_crossover_artifact.py new file mode 100644 index 0000000..8b82ec4 --- /dev/null +++ b/experiments/validate_crossover_artifact.py @@ -0,0 +1,73 @@ +#!/usr/bin/env python3 +"""Small dependency-free validator used by remote crossover handoffs.""" +import argparse +import json + + +def validate_selector(report, stage, count): + assert report["gate"] == "pass" + assert report["stage"] == stage + assert report["num_expected_records"] == count + assert report["num_audited_records"] == count + assert report["missing_experiments"] == [] + assert len(report["selected"]) == 9 + assert report["test_policy"] == "none" + + +def validate_audit(report, count): + assert report["audit_status"] == "passed" + assert report["complete_grid"] is True + assert report["failure_retaining"] is True + assert report["num_audited_cells"] == count + assert report["test_policy"] == "none" + + +def validate(report, kind, expected_commit=None): + if kind == "resnet_p1_launch": + assert expected_commit + assert report["stage"] == "p1" + assert report["source"]["git_commit"] == expected_commit + assert report["num_jobs"] == 19 + assert report["allowed_physical_gpus"] == [5, 7] + elif kind == "resnet_p1_selector": + validate_selector(report, "resnet_crossover_p1", 19) + elif kind == "transformer_p1_selector": + validate_selector(report, "transformer_crossover_p1", 27) + elif kind == "resnet_r2_audit": + assert report["stage"] == "resnet_crossover_r2" + validate_audit(report, 27) + elif kind == "transformer_t2_audit": + assert report["stage"] == "transformer_crossover_t2" + validate_audit(report, 27) + elif kind == "crossover_81_audit": + assert report["stage"] == "crossover_x1_81" + validate_audit(report, 81) + else: + raise ValueError(f"unknown artifact kind: {kind}") + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--path", required=True) + parser.add_argument( + "--kind", + required=True, + choices=( + "resnet_p1_launch", + "resnet_p1_selector", + "transformer_p1_selector", + "resnet_r2_audit", + "transformer_t2_audit", + "crossover_81_audit", + ), + ) + parser.add_argument("--expected-commit") + args = parser.parse_args() + with open(args.path, encoding="utf-8") as handle: + report = json.load(handle) + validate(report, args.kind, args.expected_commit) + print(f"validated {args.kind}: {args.path}") + + +if __name__ == "__main__": + main() |
