summaryrefslogtreecommitdiff
path: root/experiments/validate_crossover_artifact.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/validate_crossover_artifact.py')
-rw-r--r--experiments/validate_crossover_artifact.py73
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()