File size: 3,233 Bytes
1ad01ce
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Prove the patched graph is numerically identical to the original.

This is the ONLY thing standing between a subtly wrong graph and production: the
patched ONNX is our own artifact, no upstream validates it, and a Split tree that
reorders or mis-sizes a group would still produce a plausible-looking matte.
A quality metric cannot catch that. Bit-level agreement can.

`patch_split.py` and `patch_deform.py` perform structural rewrites with no numerical content, so the
correct expectation is EXACT equality, not "close enough". The 0.9999 gate exists
only to absorb non-determinism in CPU kernel scheduling.

    python verify_patch.py orig.onnx patched.onnx img1.jpg img2.jpg ...
"""

import argparse
import sys

import numpy as np
import onnxruntime as ort
from PIL import Image

SIZE = 1024
MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)


def preprocess(path):
    # fit='fill' - the aspect ratio is squashed, matching what the browser does
    # by drawing into a square canvas. A different framing is a different input.
    img = Image.open(path).convert("RGB").resize((SIZE, SIZE), Image.BILINEAR)
    x = np.asarray(img, dtype=np.float32) / 255.0
    x = (x - MEAN) / STD
    return x.transpose(2, 0, 1)[None].astype(np.float32)


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("orig")
    ap.add_argument("patched")
    ap.add_argument("images", nargs="+")
    ap.add_argument("--min-correlation", type=float, default=0.9999)
    args = ap.parse_args()

    opts = ort.SessionOptions()
    opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_DISABLE_ALL
    a = ort.InferenceSession(args.orig, opts, providers=["CPUExecutionProvider"])
    b = ort.InferenceSession(args.patched, opts, providers=["CPUExecutionProvider"])

    name_a = a.get_inputs()[0].name
    name_b = b.get_inputs()[0].name
    print(f"inputs: {name_a} / {name_b}")

    worst_corr, worst_absdiff, failed = 1.0, 0.0, False
    for path in args.images:
        x = preprocess(path)
        oa = np.asarray(a.run(None, {name_a: x})[0], dtype=np.float64).ravel()
        ob = np.asarray(b.run(None, {name_b: x})[0], dtype=np.float64).ravel()

        if oa.shape != ob.shape:
            print(f"FAIL {path}: shape {oa.shape} vs {ob.shape}")
            failed = True
            continue

        absdiff = float(np.max(np.abs(oa - ob)))
        corr = 1.0 if absdiff == 0.0 else float(np.corrcoef(oa, ob)[0, 1])
        worst_corr = min(worst_corr, corr)
        worst_absdiff = max(worst_absdiff, absdiff)
        exact = "EXACT" if absdiff == 0.0 else f"max|diff| {absdiff:.3e}"
        print(f"  {path.split('/')[-1]:50s} corr {corr:.8f}  {exact}")

    print(f"\nworst correlation {worst_corr:.8f}, worst max|diff| {worst_absdiff:.3e}")
    if failed or worst_corr < args.min_correlation:
        print("REJECTED — do not ship this graph.")
        return 1
    if worst_absdiff == 0.0:
        print("ACCEPTED — bit-identical, as a structural rewrite should be.")
    else:
        print("ACCEPTED — within tolerance, but NOT bit-identical; investigate before shipping.")
    return 0


if __name__ == "__main__":
    sys.exit(main())