"""Exact conditional spacetime-cluster sampler for a binary Markov trajectory.""" import numpy as np def transition(a, d): return np.array([[1-a, a], [d, 1-d]], dtype=float) def exact_terminal_probability(a, d, T, x0=0, terminal=1): P = transition(a, d); v = np.zeros(2); v[x0] = 1. for _ in range(T): v = v @ P return float(v[terminal]) def sample_forward(rng, a, d, T, x0=0): P = transition(a, d); x = np.empty(T+1, dtype=np.int8); x[0] = x0 for t in range(1, T+1): x[t] = rng.choice(2, p=P[x[t-1]]) return x def conditional_cluster(rng, a, d, T, terminal=1, x0=0): """Heat-bath update of all interior spacetime sites, endpoints fixed.""" P = transition(a, d) beta = np.zeros((T+1, 2)); beta[T, terminal] = 1. for t in range(T-1, -1, -1): beta[t] = P @ beta[t+1] x = np.empty(T+1, dtype=np.int8); x[0] = x0; x[T] = terminal for t in range(1, T): w = P[x[t-1]] * beta[t] x[t] = rng.choice(2, p=w/w.sum()) return x def local_gibbs(rng, a, d, T, steps, terminal=1, x0=0, init=None): P = transition(a, d) x = conditional_cluster(rng, a, d, T, terminal, x0) if init is None else init.copy() vals = np.empty(steps) for k in range(steps): t = rng.integers(1, T) w = P[x[t-1]] * P[:, x[t+1]] x[t] = rng.choice(2, p=w/w.sum()) vals[k] = x[1:T].sum() return vals, x def enumerate_conditioned(a, d, T, terminal=1, x0=0): P=transition(a,d); paths=[]; weights=[] for mask in range(1 << (T-1)): x=np.zeros(T+1,dtype=np.int8); x[0]=x0; x[T]=terminal for t in range(1,T): x[t]=(mask>>(t-1))&1 w=np.prod([P[x[t-1],x[t]] for t in range(1,T+1)]) paths.append(x); weights.append(w) weights=np.asarray(weights); return paths, weights/weights.sum() def ess(x): x=np.asarray(x,float); x-=x.mean(); var=np.dot(x,x)/len(x) if len(x)<3 or var < 1e-15: return float(len(x)) tau=1. for lag in range(1,min(len(x)-1,2000)): ac=np.dot(x[:-lag],x[lag:])/(len(x)-lag)/var if ac <= 0: break tau += 2*ac return len(x)/tau def run_case(a, d=.5, T=10, n=200, proposals=20000, seed=123): rng=np.random.default_rng(seed); Z=exact_terminal_probability(a,d,T) valid=0 for _ in range(proposals): valid += int(sample_forward(rng,a,d,T)[-1] == 1) cl=np.array([conditional_cluster(rng,a,d,T)[1:-1].sum() for _ in range(n)]) loc,_=local_gibbs(rng,a,d,T,n,init=conditional_cluster(rng,a,d,T)) paths,pw=enumerate_conditioned(a,d,T); counts=np.zeros(len(paths)) index={tuple(x):i for i,x in enumerate(paths)} for _ in range(2000): counts[index[tuple(conditional_cluster(rng,a,d,T))]] += 1 tv=.5*np.abs(counts/2000-pw).sum() exact_mean=sum(w*x[1:-1].sum() for x,w in zip(paths,pw)) return dict(a=a,Z=Z,proposal_acceptance=valid/proposals, acceptance_pred=Z,rejection_cost_pred=1/Z,cluster_valid=1., cluster_mean=float(cl.mean()),exact_mean=float(exact_mean),tv=float(tv), cluster_ess=ess(cl),local_ess=ess(loc)) if __name__ == '__main__': for r in [run_case(a) for a in (.05,.005,.0005)]: print(r)