Conditional spacetime-cluster sampler for rare neural trajectories / conditional_cluster_sampler.py

Failed on benchmark

Raw ⬇ ZIP
 1"""Exact conditional spacetime-cluster sampler for a binary Markov trajectory."""
 2import numpy as np
 3
 4
 5def transition(a, d):
 6    return np.array([[1-a, a], [d, 1-d]], dtype=float)
 7
 8
 9def exact_terminal_probability(a, d, T, x0=0, terminal=1):
10    P = transition(a, d); v = np.zeros(2); v[x0] = 1.
11    for _ in range(T): v = v @ P
12    return float(v[terminal])
13
14
15def sample_forward(rng, a, d, T, x0=0):
16    P = transition(a, d); x = np.empty(T+1, dtype=np.int8); x[0] = x0
17    for t in range(1, T+1): x[t] = rng.choice(2, p=P[x[t-1]])
18    return x
19
20
21def conditional_cluster(rng, a, d, T, terminal=1, x0=0):
22    """Heat-bath update of all interior spacetime sites, endpoints fixed."""
23    P = transition(a, d)
24    beta = np.zeros((T+1, 2)); beta[T, terminal] = 1.
25    for t in range(T-1, -1, -1): beta[t] = P @ beta[t+1]
26    x = np.empty(T+1, dtype=np.int8); x[0] = x0; x[T] = terminal
27    for t in range(1, T):
28        w = P[x[t-1]] * beta[t]
29        x[t] = rng.choice(2, p=w/w.sum())
30    return x
31
32
33def local_gibbs(rng, a, d, T, steps, terminal=1, x0=0, init=None):
34    P = transition(a, d)
35    x = conditional_cluster(rng, a, d, T, terminal, x0) if init is None else init.copy()
36    vals = np.empty(steps)
37    for k in range(steps):
38        t = rng.integers(1, T)
39        w = P[x[t-1]] * P[:, x[t+1]]
40        x[t] = rng.choice(2, p=w/w.sum())
41        vals[k] = x[1:T].sum()
42    return vals, x
43
44
45def enumerate_conditioned(a, d, T, terminal=1, x0=0):
46    P=transition(a,d); paths=[]; weights=[]
47    for mask in range(1 << (T-1)):
48        x=np.zeros(T+1,dtype=np.int8); x[0]=x0; x[T]=terminal
49        for t in range(1,T): x[t]=(mask>>(t-1))&1
50        w=np.prod([P[x[t-1],x[t]] for t in range(1,T+1)])
51        paths.append(x); weights.append(w)
52    weights=np.asarray(weights); return paths, weights/weights.sum()
53
54
55def ess(x):
56    x=np.asarray(x,float); x-=x.mean(); var=np.dot(x,x)/len(x)
57    if len(x)<3 or var < 1e-15: return float(len(x))
58    tau=1.
59    for lag in range(1,min(len(x)-1,2000)):
60        ac=np.dot(x[:-lag],x[lag:])/(len(x)-lag)/var
61        if ac <= 0: break
62        tau += 2*ac
63    return len(x)/tau
64
65
66def run_case(a, d=.5, T=10, n=200, proposals=20000, seed=123):
67    rng=np.random.default_rng(seed); Z=exact_terminal_probability(a,d,T)
68    valid=0
69    for _ in range(proposals): valid += int(sample_forward(rng,a,d,T)[-1] == 1)
70    cl=np.array([conditional_cluster(rng,a,d,T)[1:-1].sum() for _ in range(n)])
71    loc,_=local_gibbs(rng,a,d,T,n,init=conditional_cluster(rng,a,d,T))
72    paths,pw=enumerate_conditioned(a,d,T); counts=np.zeros(len(paths))
73    index={tuple(x):i for i,x in enumerate(paths)}
74    for _ in range(2000): counts[index[tuple(conditional_cluster(rng,a,d,T))]] += 1
75    tv=.5*np.abs(counts/2000-pw).sum()
76    exact_mean=sum(w*x[1:-1].sum() for x,w in zip(paths,pw))
77    return dict(a=a,Z=Z,proposal_acceptance=valid/proposals,
78      acceptance_pred=Z,rejection_cost_pred=1/Z,cluster_valid=1.,
79      cluster_mean=float(cl.mean()),exact_mean=float(exact_mean),tv=float(tv),
80      cluster_ess=ess(cl),local_ess=ess(loc))
81
82if __name__ == '__main__':
83    for r in [run_case(a) for a in (.05,.005,.0005)]: print(r)