1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
|
#!/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")
transformer_launch = copy.deepcopy(launch)
transformer_launch["num_jobs"] = 27
validate(transformer_launch, "transformer_p1_launch", "abc")
r2_launch = copy.deepcopy(transformer_launch)
r2_launch["stage"] = "r2"
validate(r2_launch, "resnet_r2_launch", "abc")
t2_launch = copy.deepcopy(transformer_launch)
t2_launch["stage"] = "t2"
validate(t2_launch, "transformer_t2_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()
|