File size: 5,050 Bytes
4b5ca98
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
"""Check only staged public artifacts in an installed runtime, with no training."""

import argparse
import hashlib
import json
import os
from pathlib import Path
import subprocess


def sha(path):
    return hashlib.sha256(path.read_bytes()).hexdigest()


def verify(release, python, output, reference=None, candidate_name=None):
    release = release.resolve()
    output.mkdir(parents=True, exist_ok=False)
    receipt = json.loads((release / "release.json").read_text())
    env = dict(os.environ)
    for key in ("PYTHONPATH", "PYTHONHOME", "VIRTUAL_ENV"):
        env.pop(key, None)
    env.update(PYTHONDONTWRITEBYTECODE="1", OMP_NUM_THREADS="1", MKL_NUM_THREADS="1")
    results = []
    for arm, metadata in receipt["candidates"].items():
        if candidate_name is not None and arm != candidate_name:
            continue
        local = output / arm
        local.mkdir()
        candidate = release / "candidates" / arm
        for p in candidate.rglob("*"):
            if p.is_file():
                p.chmod(0o600)
        for p in (release / "examples").iterdir():
            p.chmod(0o600)
        assert sha(candidate / "consumer.json") == metadata["metadata_sha256"]
        base = [str(python), "-I", "-B", "-m", "laxmi_readout_runtime_v1",
                "--expected-package-sha256", receipt["package_sha256"],
                "--checkpoint", str(candidate / "inference"), "--metadata", str(candidate / "consumer.json"),
                "--metadata-sha256", metadata["metadata_sha256"]]

        def run(name, args):
            value = subprocess.run(base + args, cwd=output, env=env, capture_output=True, timeout=120)
            (local / f"{name}.stderr").write_bytes(value.stderr)
            log = local / f"{name}.json"
            log.write_bytes(value.stdout)
            log.chmod(0o600)
            if value.returncode:
                raise RuntimeError(f"{arm}/{name} failed: " + value.stderr.decode()[-1500:])
            return json.loads(value.stdout)

        state = local / "state.json"
        run("observe", ["observe", "--input", str(release / "examples/input.json"), "--out", str(state)])
        original = state.read_bytes()
        pred = run("predict", ["predict", "--state", str(state), "--input", str(release / "examples/input.json")])
        compared = run("compare", ["compare", "--state", str(state), "--input", str(release / "examples/compare.json")])
        assert pred["forecast"] == compared["forecast"][0]
        run("restore", ["restore-state", "--state", str(state)])
        run("save", ["save-state", "--state", str(state), "--out", str(local / "saved.json")])
        assert state.read_bytes() == original == (local / "saved.json").read_bytes()
        outcome = json.loads((release / "examples/outcome-template.json").read_text())
        outcome["prediction_file_sha256"] = sha(local / "predict.json")
        outcome["input_file_sha256"] = sha(release / "examples/input.json")
        target = local / "outcome.json"
        target.write_text(json.dumps(outcome))
        target.chmod(0o600)
        run("score", ["score-outcome", "--state", str(state), "--prediction", str(local / "predict.json"),
                      "--outcome", str(target), "--out", str(local / "comparison.json")])
        exact = None
        if reference:
            retained = json.loads((reference / f"{arm}-predict.log").read_text())
            assert pred["forecast"] == retained["forecast"]
            exact = True
        forecast_digest = hashlib.sha256(json.dumps(pred["forecast"], sort_keys=True).encode()).hexdigest()
        if "example_forecast_sha256" in metadata:
            assert forecast_digest == metadata["example_forecast_sha256"]
        results.append(dict(arm=arm, observe_predict_compare_restore_save_score=True,
                            exact_retained_forecast=exact,
                            forecast_sha256=forecast_digest))
        print(arm + ": all six operations passed", flush=True)
    if not results:
        raise ValueError("candidate not present in this release")
    summary = dict(status="passed", platform="macOS ARM64, Python 3.12; inspect interpreter for exact patch version",
                   scope="one retained synthetic depth-1 H1 example per candidate; installation/replay only",
                   records=results, model_acceptance=False, dependency_hash_lock=False)
    (output / "verification.json").write_text(json.dumps(summary, indent=2) + "\n")
    return summary


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--release", type=Path, required=True)
    parser.add_argument("--python", type=Path, required=True)
    parser.add_argument("--output", type=Path, required=True)
    parser.add_argument("--reference", type=Path)
    parser.add_argument("--candidate", choices=("17-original", "17-normalized", "43-original", "43-normalized"))
    args = parser.parse_args()
    verify(args.release, args.python.absolute(), args.output.resolve(), args.reference, args.candidate)