Conditional spacetime-cluster sampler for rare neural trajectories / conditional_cluster_sampler.py
Failed on benchmark
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)