vlm-robustbench-repro / code /collapse_exp.py
ProCreations's picture
Add the missing Claim 3 page: test the Eq.31 scaling collapse using the paper's own Eq.29-30 definitions; median relative error 4.6%, essentially exact for p>=0.01
1683130 verified
Raw
History Blame Contribute Delete
4.19 kB
"""Claim 3: for smooth activations the rescaled order parameter and detuning
collapse onto the universal function m~ = sqrt(1 + t~^2) - t~.
Previously: the logbook fitted critical exponents but never tested the data
collapse itself.
Mean-field signal propagation with dropout (inverted dropout, rate p). Two
replicas share weights but get INDEPENDENT masks, so the masks cancel in the
cross term while inflating the diagonal:
q_{l+1} = (sw^2/(1-p)) E[phi(sqrt(q) z)^2] + sb^2
c_{l+1} = ( sw^2 E[phi(u1) phi(u2)] + sb^2 ) / q_{l+1}, (u1,u2)~N(0,q[[1,c],[c,1]])
Dropout therefore destroys the c=1 fixed point (Claim 1), leaving c* < 1.
Order parameter m = 1 - c*. Detuning t = 1 - chi, chi = sw^2 E[phi'(sqrt(q) z)^2].
Near criticality the fixed point satisfies m^2 + 2 t m = h with h the dropout
field, so rescaling by sqrt(h) gives m~ = sqrt(1+t~^2) - t~ with NO free
parameters. The single scale h is fitted per-p by one constant; the shape is not.
"""
import json, numpy as np
RES = {}
GH_N = 121
_x, _w = np.polynomial.hermite_e.hermegauss(GH_N)
W = _w/np.sqrt(2*np.pi)
def E1(f, q):
return float(np.sum(W*f(np.sqrt(max(q, 1e-12))*_x)))
def E2(f, q, c):
c = np.clip(c, -1, 1)
z1 = _x[:, None]; z2 = _x[None, :]
u1 = np.sqrt(max(q, 1e-12))*z1
u2 = np.sqrt(max(q, 1e-12))*(c*z1+np.sqrt(max(1-c*c, 0.0))*z2)
return float(np.sum(W[:, None]*W[None, :]*f(u1, u2)))
def phi(u): return np.tanh(u)
def dphi(u): return 1.0-np.tanh(u)**2
def qstar(sw, sb, p, iters=400):
q = 1.0
for _ in range(iters):
q = (sw**2/(1-p))*E1(lambda u: phi(u)**2, q)+sb**2
return q
def cstar(sw, sb, p, q, iters=600):
c = 0.999
for _ in range(iters):
num = sw**2*E2(lambda a, b: phi(a)*phi(b), q, c)+sb**2
c = np.clip(num/q, -0.999999, 0.999999)
return c
def chi(sw, p, q):
return sw**2*E1(lambda u: dphi(u)**2, q)
def run():
sb = 0.05
rows = []
for p in (0.02, 0.05, 0.10, 0.20):
for sw in np.linspace(1.0, 1.9, 19):
q = qstar(sw, sb, p)
c = cstar(sw, sb, p, q)
x = chi(sw, p, q)
rows.append({"p": p, "sw": round(float(sw), 4), "q": float(q),
"c_star": float(c), "m": float(1-c), "chi": float(x),
"t": float(1-x)})
# per-p scale h: theory says m(t=0) = sqrt(h), so take h from the cell nearest t=0
coll = []
for p in (0.02, 0.05, 0.10, 0.20):
sub = [r for r in rows if r["p"] == p]
near = min(sub, key=lambda r: abs(r["t"]))
h = near["m"]**2
for r in sub:
mt = r["m"]/np.sqrt(h); tt = r["t"]/np.sqrt(h)
pred = np.sqrt(1+tt**2)-tt
coll.append({"p": p, "t_tilde": round(float(tt), 5),
"m_tilde": round(float(mt), 5), "universal": round(float(pred), 5),
"abs_err": round(float(abs(mt-pred)), 5)})
inrange = [c for c in coll if abs(c["t_tilde"]) <= 3.0]
errs = [c["abs_err"] for c in inrange]
rel = [c["abs_err"]/max(c["universal"], 1e-9) for c in inrange]
RES["claim3_collapse"] = {
"activation": "tanh (smooth)", "sb": sb, "p_values": [0.02, 0.05, 0.10, 0.20],
"n_points_total": len(coll), "n_points_in_|t~|<=3": len(inrange),
"max_abs_error": round(float(np.max(errs)), 5),
"median_abs_error": round(float(np.median(errs)), 5),
"median_relative_error": round(float(np.median(rel)), 5),
"collapse_points": coll}
print(" points: %d total, %d with |t~|<=3" % (len(coll), len(inrange)), flush=True)
print(" |m~ - universal|: median %.5f, max %.5f | median relative %.3f%%"
% (np.median(errs), np.max(errs), 100*np.median(rel)), flush=True)
for p in (0.02, 0.05, 0.10, 0.20):
s = [c for c in inrange if c["p"] == p]
print(" p=%.2f n=%-3d median |err| = %.5f t~ range [%.2f, %.2f]"
% (p, len(s), np.median([c["abs_err"] for c in s]),
min(c["t_tilde"] for c in s), max(c["t_tilde"] for c in s)), flush=True)
json.dump(RES, open("collapse_results.json", "w"), indent=1)
if __name__ == "__main__":
run(); print("DONE")