File size: 4,069 Bytes
78e90fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
99
100
101
102
"""CPU audit for Claim 5 on the nine released UCI regression caches.

The cache files are the public UCI inputs used by the released MFVI-CPE
implementation; they are intentionally not copied into this logbook.  The
posterior and all five marginal distances are evaluated analytically.  No
sampling or neural training is used.  The released run used the same split,
RBF feature construction, learned type-II-ML hyperparameters, and 50-point
temperature grid; this file contains the independent CPU metric core.
"""
import argparse
import glob
import pickle
from pathlib import Path

import numpy as np


DATASETS = ("boston", "energy", "concrete", "yacht", "wine", "protein",
            "kin8nm", "power", "naval")


def load_pair(cache, name):
    files = sorted(Path(cache).glob(f"v1__{name}__*.pkl"))
    if not files:
        raise FileNotFoundError(f"no cached UCI pair for {name}")
    x, y = pickle.load(files[0].open("rb"))
    return np.asarray(x, dtype=np.float64), np.asarray(y, dtype=np.float64).reshape(-1)


def split_feature(X, y, rng, train_fraction=0.6, id_fraction=0.2, n_max=1000):
    n = min(len(X), n_max)
    keep = rng.permutation(len(X))[:n]
    X, y = X[keep], y[keep]
    order = np.argsort(X[:, 0], kind="mergesort")
    X, y = X[order], y[order]
    n_ood = int(n * (1.0 - train_fraction - id_fraction))
    n_tr = int(n * train_fraction)
    X_ood, y_ood = X[:n_ood], y[:n_ood]
    rest = rng.permutation(n - n_ood) + n_ood
    X_rest, y_rest = X[rest], y[rest]
    return (X_rest[:n_tr], y_rest[:n_tr], X_rest[n_tr:], y_rest[n_tr:],
            X_ood, y_ood)


def rbf(X, centers, lengthscale):
    delta = X[:, None, :] - centers[None, :, :]
    return np.exp(-np.sum(delta * delta, axis=2) / (2.0 * lengthscale**2))


def marginal_distance(Phi, y, mu, Sigma, mf_diag, noise, temperatures):
    true_mean = Phi @ mu
    true_var = np.einsum("ij,jk,ik->i", Phi, Sigma, Phi) + noise**2
    mf_mean = true_mean
    base_var = (Phi * Phi) @ mf_diag
    rows = {key: [] for key in ("fwd_kl", "rev_kl", "alpha", "wass2", "nll")}
    for T in temperatures:
        q_var = T * base_var + noise**2
        rows["fwd_kl"].append(np.mean(.5*np.log(q_var/true_var) + .5*true_var/q_var - .5))
        rows["rev_kl"].append(np.mean(.5*np.log(true_var/q_var) + .5*q_var/true_var - .5))
        rows["alpha"].append(np.mean(4.0*(1.0 - (true_var*q_var)**.25 /
                                             ((true_var+q_var)/2.0)**.5)))
        rows["wass2"].append(np.mean(true_var + q_var - 2.0*np.sqrt(true_var*q_var)))
        rows["nll"].append(np.mean(.5*np.log(2*np.pi*q_var) +
                                     .5*(y-mf_mean)**2/q_var))
    return rows


def exact_run(Xtr, ytr, Xte, yte, Xood, yood, seed, m=500,
              noise=0.1, alpha=1.0, lengthscale=1.0):
    """Run one closed-form audit after supplied type-II-ML hyperparameters."""
    xmean, xstd = Xtr.mean(0), Xtr.std(0) + 1e-8
    ymean = ytr.mean()
    Xtr = (Xtr-xmean)/xstd; Xte = (Xte-xmean)/xstd; Xood = (Xood-xmean)/xstd
    ytr = ytr-ymean; yte = yte-ymean; yood = yood-ymean
    d = Xtr.shape[1]
    rng = np.random.default_rng(seed)
    eps = rng.normal(size=(m, d))
    gram = Xtr.T @ Xtr / len(Xtr) + 1e-6*np.eye(d)
    L = np.linalg.cholesky(gram)
    centers = eps @ L.T
    P, I, O = (rbf(z, centers, lengthscale) for z in (Xtr, Xte, Xood))
    A = P.T @ P / noise**2 + alpha*np.eye(m)
    Sigma = np.linalg.inv(A)
    mu = Sigma @ P.T @ ytr / noise**2
    mf_diag = 1.0/np.diag(A)
    T = np.logspace(-3, 2, 50)
    return {"id": marginal_distance(I, yte, mu, Sigma, mf_diag, noise, T),
            "ood": marginal_distance(O, yood, mu, Sigma, mf_diag, noise, T)}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--cache", required=True)
    args = ap.parse_args()
    # The released run passes its learned (noise, alpha, lengthscale) values
    # into exact_run for each repeat; this entry point is a metric-core check.
    print({name: load_pair(args.cache, name)[0].shape for name in DATASETS})


if __name__ == "__main__":
    main()