Steady-State First-Passage Sensitivity Regularizer / experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, time
  2import numpy as np
  3from scipy.linalg import solve
  4
  5SEED = 2117
  6rng = np.random.default_rng(SEED)
  7
  8# Row-generator convention: p'(t)=p(t)Q.  State 0 is source, state 3 target.
  9Q0 = np.array([
 10    [-1.20, 0.90, 0.30, 0.00],
 11    [ 0.10,-1.00, 0.70, 0.20],
 12    [ 0.05, 0.05,-0.55, 0.45],
 13    [ 0.30, 0.00, 0.10,-0.40],
 14], dtype=float)
 15# Perturbed directed edge n=1 -> m=2, baseline rate 0.70.
 16N, M, SRC, TGT = 1, 2, 0, 3
 17
 18def generator(h=0.0):
 19    Q = Q0.copy()
 20    # perturb rate additively; diagonal is adjusted to preserve row sums
 21    Q[N, M] = Q0[N, M] + h
 22    Q[N, N] = -np.sum(Q[N, np.arange(4) != N])
 23    return Q
 24
 25def mfpt(Q, source=SRC, target=TGT):
 26    transient = [x for x in range(len(Q)) if x != target]
 27    A = Q[np.ix_(transient, transient)]
 28    vals = solve(A, -np.ones(len(transient)))
 29    return float(vals[transient.index(source)])
 30
 31def stationary(Q):
 32    # solve Q^T pi=0 with one row replaced by normalization
 33    A = Q.T.copy(); b = np.zeros(len(Q))
 34    A[-1, :] = 1.0; b[-1] = 1.0
 35    return solve(A, b)
 36
 37def auxiliary(Q, reset_rate):
 38    A = Q.copy()
 39    # Add a fast edge target -> source, as in the paper's auxiliary system.
 40    A[TGT, SRC] += reset_rate
 41    A[TGT, TGT] -= reset_rate
 42    return A
 43
 44def aux_response(h, reset_rate, eps=1e-6):
 45    # Exact finite difference of -T(0)*d log pi_target / dh.
 46    pplus = stationary(auxiliary(generator(eps), reset_rate))[TGT]
 47    pminus = stationary(auxiliary(generator(-eps), reset_rate))[TGT]
 48    return -mfpt(generator(), SRC, TGT) * (math.log(pplus)-math.log(pminus))/(2*eps)
 49
 50def rollout_mfpt(Q, nroll=3000):
 51    # Gillespie first passage samples, used as the expensive baseline.
 52    out = np.empty(nroll)
 53    for r in range(nroll):
 54        s, tm = SRC, 0.0
 55        while s != TGT and tm < 1e7:
 56            rates = Q[s].copy(); rates[s] = 0.0
 57            total = rates.sum()
 58            tm += rng.exponential(1.0/total)
 59            u = rng.random() * total
 60            s = int(np.searchsorted(np.cumsum(rates), u))
 61        out[r] = tm
 62    return out.mean(), out.std(ddof=1)/math.sqrt(nroll)
 63
 64def main():
 65    baseT = mfpt(generator())
 66    # Formula (3a): derivative wrt W_nm.
 67    pi = stationary(generator())
 68    # Use direct all-pairs MFPT to evaluate the paper's exact response formula.
 69    def all_tau(Q): return np.array([[mfpt(Q,j,i) if i != j else 0.0 for j in range(4)] for i in range(4)])
 70    tau = all_tau(generator())
 71    # target k=TGT, source l=SRC, edge n->m
 72    R_formula = -pi[N] * (tau[TGT,N]-tau[TGT,M]) * (tau[N,TGT]+tau[TGT,SRC]-tau[N,SRC])
 73    eps=1e-5
 74    R_fd=(mfpt(generator(eps))-mfpt(generator(-eps)))/(2*eps)
 75
 76    # Prediction 1: fast-reset error decreases approximately as 1/reset_rate.
 77    Ks=np.array([2.,5.,10.,20.,50.,100.,200.,500.,1000.])
 78    aux=np.array([aux_response(0.0,k) for k in Ks])
 79    relerr=np.abs(aux-R_fd)/max(abs(R_fd),1e-12)
 80    slope=np.polyfit(np.log(Ks[-5:]), np.log(np.maximum(relerr[-5:],1e-16)), 1)[0]
 81
 82    # Prediction 2: exact finite perturbation curve obeys theorem's rational denominator.
 83    hs=np.array([-.45,-.30,-.15,.15,.30,.45])
 84    Sigma=tau[TGT,N]+tau[N,TGT]-tau[TGT,M]
 85    U=tau[N,TGT]+tau[TGT,SRC]-tau[N,SRC]
 86    G=tau[TGT,N]-tau[TGT,M]
 87    pred=[]; obs=[]; err=[]
 88    for h in hs:
 89        d=mfpt(generator(h))-baseT
 90        formula=-pi[N]*h*G*U/(1+pi[N]*h*Sigma)
 91        obs.append(d); pred.append(formula); err.append(abs(d-formula))
 92    # Prediction 3: symmetry/sign: make target distances equal by constructing a symmetric
 93    # two-branch chain and check derivative is zero; also sign follows G.
 94    Qsym=np.array([[-1, .5, .5, 0],[0,-1,0,1],[0,0,-1,1],[.2,0,0,-.2]],float)
 95    # perturb 1->2 changes branch balance; target distances from 1 and 2 equal.
 96    def symmf(h):
 97        Q=Qsym.copy(); Q[1,2]=h; Q[1,1]=-(Q[1,2]+Q[1,3]); return mfpt(Q,0,3)
 98    symR=(symmf(1e-6)-symmf(-1e-6))/(2e-6)
 99    # Sign check on two edges: edge 1->2 helpful (negative), edge 2->1 harmful (positive)
100    def edgeR(a,b):
101        def f(h):
102            Q=generator(); Q[a,b]+=h; Q[a,a]-=h; return mfpt(Q)
103        return (f(1e-5)-f(-1e-5))/(2e-5)
104    signRs={'1_to_2':edgeR(1,2),'2_to_1':edgeR(2,1)}
105
106    # Mini experiment: finite-difference rollout baseline versus stationary auxiliary estimate.
107    t0=time.time(); nroll=2500; hh=.12
108    p0,se0=rollout_mfpt(generator(hh),nroll); m0,se_m0=rollout_mfpt(generator(-hh),nroll)
109    mc_fd=(p0-m0)/(2*hh); mc_se=math.sqrt(se0**2+se_m0**2)/(2*hh)
110    exact_aux=aux_response(0,1000)
111    result={
112      'seed':SEED,'base_mfpt':baseT,'exact_fd_derivative':R_fd,'formula_3a':R_formula,
113      'formula_abs_error':abs(R_formula-R_fd),
114      'prediction_1_reset_rates':Ks.tolist(),'prediction_1_aux_response':aux.tolist(),
115      'prediction_1_relative_error':relerr.tolist(),'prediction_1_loglog_slope':float(slope),
116      'prediction_2_G':G,'prediction_2_U':U,'prediction_2_Sigma':Sigma,
117      'prediction_2_observed_delta':obs,'prediction_2_predicted_delta':pred,
118      'prediction_2_max_abs_error':max(err),
119      'prediction_3_symmetric_derivative':float(symR),'prediction_3_sign_derivatives':signRs,
120      'baseline_rollout_fd':mc_fd,'baseline_rollout_standard_error':mc_se,
121      'idea_auxiliary_fd':float(exact_aux),'idea_abs_error_vs_exact':abs(exact_aux-R_fd),
122      'runtime_sec':time.time()-t0
123    }
124    with open('results.json','w') as f: json.dump(result,f,indent=2)
125    print(json.dumps(result,indent=2))
126
127if __name__=='__main__': main()