diff options
Diffstat (limited to 'experiments/crossover_artifact_validation_smoke.py')
| -rw-r--r-- | experiments/crossover_artifact_validation_smoke.py | 72 |
1 files changed, 72 insertions, 0 deletions
diff --git a/experiments/crossover_artifact_validation_smoke.py b/experiments/crossover_artifact_validation_smoke.py new file mode 100644 index 0000000..e267832 --- /dev/null +++ b/experiments/crossover_artifact_validation_smoke.py @@ -0,0 +1,72 @@ +#!/usr/bin/env python3 +"""Deterministic checks for dependency-free handoff validation.""" +import copy + +from validate_crossover_artifact import validate + + +def must_fail(report, kind, expected_commit=None): + try: + validate(report, kind, expected_commit) + except (AssertionError, KeyError, ValueError): + return + raise AssertionError(f"invalid {kind} artifact was accepted") + + +def selector(stage, count): + return { + "gate": "pass", + "stage": stage, + "num_expected_records": count, + "num_audited_records": count, + "missing_experiments": [], + "selected": {str(index): {} for index in range(9)}, + "test_policy": "none", + } + + +def audit(stage, count): + return { + "audit_status": "passed", + "stage": stage, + "complete_grid": True, + "failure_retaining": True, + "num_audited_cells": count, + "test_policy": "none", + } + + +def main(): + launch = { + "stage": "p1", + "source": {"git_commit": "abc"}, + "num_jobs": 19, + "allowed_physical_gpus": [5, 7], + } + validate(launch, "resnet_p1_launch", "abc") + bad_launch = copy.deepcopy(launch) + bad_launch["allowed_physical_gpus"] = [0, 1] + must_fail(bad_launch, "resnet_p1_launch", "abc") + + resnet_selector = selector("resnet_crossover_p1", 19) + transformer_selector = selector("transformer_crossover_p1", 27) + validate(resnet_selector, "resnet_p1_selector") + validate(transformer_selector, "transformer_p1_selector") + bad_selector = copy.deepcopy(transformer_selector) + bad_selector["selected"].pop("0") + must_fail(bad_selector, "transformer_p1_selector") + + resnet = audit("resnet_crossover_r2", 27) + transformer = audit("transformer_crossover_t2", 27) + crossover = audit("crossover_x1_81", 81) + validate(resnet, "resnet_r2_audit") + validate(transformer, "transformer_t2_audit") + validate(crossover, "crossover_81_audit") + bad_audit = copy.deepcopy(crossover) + bad_audit["num_audited_cells"] = 80 + must_fail(bad_audit, "crossover_81_audit") + print("crossover artifact validation smoke: all checks passed") + + +if __name__ == "__main__": + main() |
