Kac-Ward Exact Teacher for Autoregressive Samplers / ising_teacher_track.py
Mechanism confirmed, baseline not beaten
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}