diff options
Diffstat (limited to 'experiments/oral_a_dynamic_scaling_v2_smoke.py')
| -rw-r--r-- | experiments/oral_a_dynamic_scaling_v2_smoke.py | 73 |
1 files changed, 73 insertions, 0 deletions
diff --git a/experiments/oral_a_dynamic_scaling_v2_smoke.py b/experiments/oral_a_dynamic_scaling_v2_smoke.py new file mode 100644 index 0000000..2fa72dd --- /dev/null +++ b/experiments/oral_a_dynamic_scaling_v2_smoke.py @@ -0,0 +1,73 @@ +#!/usr/bin/env python3 +"""Endpoint-free audit of the oral-A recovery-v2 frozen funnel.""" +import copy +import json +import os +import tempfile + +from analyze_oral_a_dynamic_scaling import oral_a_checks +from oral_a_dynamic_scaling import DEPTHS, METHODS, SEEDS, jobs +from oral_a_dynamic_scaling_v2 import require_prerequisites + + +ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + + +def main(): + generated = jobs("cuda", "results/oral_a_dynamic_scaling_v2") + assert len(generated) == 50 + keys = [(method, depth, seed) for method, depth, seed, _, _ in generated] + assert len(set(keys)) == 50 + expected = { + (method, depth, seed) + for depth in DEPTHS + for seed in SEEDS + for method in METHODS + if not (depth == 20 and method in ("clean_kp", "dynamic")) + } + assert set(keys) == expected + for method, depth, seed, path, command in generated: + assert path.endswith(f"{method}_d{depth}_s{seed}.json") + assert "--epochs" in command and command[command.index("--epochs")+1] == "200" + assert "--eval_every" in command + assert command[command.index("--eval_every")+1] == "0" + assert "--eval_split" in command + assert command[command.index("--eval_split")+1] == "test" + + passing = { + "bp": {20: [0.91] * 5, 32: [0.92] * 5, 56: [0.93] * 5}, + "dfa": {20: [0.80] * 5, 32: [0.81] * 5, 56: [0.84] * 5}, + "clean_kp": {20: [0.91] * 5, 32: [0.92] * 5, 56: [0.93] * 5}, + "dynamic": {20: [0.91] * 5, 32: [0.92] * 5, 56: [0.93] * 5}, + } + alignments = {20: [0.95] * 5, 32: [0.93] * 5, 56: [0.90] * 5} + checks, _ = oral_a_checks(passing, alignments, []) + assert all(checks.values()) + failing = copy.deepcopy(passing) + failing["dynamic"][56] = [0.86] * 5 + checks, _ = oral_a_checks(failing, alignments, []) + assert not all(checks.values()) + + d4_path = os.path.join( + ROOT, "results", "kp_dynamic_projection_confirmation_gate.json") + bci_path = os.path.join( + ROOT, "results", "bci_v2_calibrated_confirmation_gate.json") + with open(bci_path) as handle: + old_gate = json.load(handle) + old_gate["protocol"] = "oral_b_td_confirmation_v1" + with tempfile.TemporaryDirectory() as tmp: + fake = os.path.join(tmp, "old_r2.json") + with open(fake, "w") as handle: + json.dump(old_gate, handle) + try: + require_prerequisites(d4_path, fake) + except ValueError: + pass + else: + raise AssertionError("legacy oral-B gate must not open recovery v2") + + print("oral-A recovery-v2 endpoint-free smoke passed") + + +if __name__ == "__main__": + main() |
