Kac-Ward Exact Teacher for Autoregressive Samplers / ising_teacher_track.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import numpy as np
 2
 3META = {
 4    'name': 'planar_ising_autoregressive',
 5    'domain': 'sequence-level sampling',
 6    'description': 'Planar 4x4 Ising raster-prefix conditional prediction with exact enumerated conditional probabilities.'
 7}
 8L, N, BETA = 4, 16, 0.65
 9EDGES = ([(r*L+c, r*L+c+1) for r in range(L) for c in range(L-1)] +
10         [(r*L+c, (r+1)*L+c) for r in range(L-1) for c in range(L)])
11_CACHE = None
12
13def _tables():
14    global _CACHE
15    if _CACHE is not None:
16        return _CACHE
17    states = np.array(np.meshgrid(*([[-1, 1]] * N))).T.reshape(-1, N).astype(np.int8)
18    logw = np.zeros(len(states), dtype=np.float64)
19    for u, v in EDGES:
20        logw += BETA * states[:, u] * states[:, v]
21    logw -= logw.max()
22    probs = np.exp(logw); probs /= probs.sum()
23    table = {}
24    for i in range(N):
25        groups = {}
26        for j, s in enumerate(states):
27            key = tuple(int(v) for v in s[:i])
28            z = groups.setdefault(key, [0.0, 0.0])
29            z[0 if s[i] < 0 else 1] += probs[j]
30        for key, z in groups.items():
31            table[(i, key)] = z[1] / (z[0] + z[1])
32    _CACHE = (states, probs, table)
33    return _CACHE
34
35def get_dataset(seed, n_train, n_test):
36    states, probs, table = _tables()
37    rng = np.random.RandomState(seed)
38    def make(n):
39        ix = rng.choice(len(states), size=n, p=probs)
40        pos = rng.randint(0, N, size=n)
41        x = np.zeros((n, 2*N), dtype=np.float32)
42        y = np.zeros((n, 1), dtype=np.float32)
43        q = np.zeros((n, 1), dtype=np.float32)
44        for k, j in enumerate(ix):
45            s = states[j]; i = int(pos[k])
46            x[k, :i] = s[:i]; x[k, N+i] = 1.0
47            q[k, 0] = table[(i, tuple(int(v) for v in s[:i]))]
48            y[k, 0] = float(rng.rand() < q[k, 0])
49        return x, y, q
50    xtr, ytr, qtr = make(n_train); xte, yte, qte = make(n_test)
51    return {'xtr':xtr, 'ytr':ytr, 'xte':xte, 'yte':yte,
52            'qtr':qtr, 'qte':qte, 'task':'regression', 'metric':'mse', 'out_dim':1}