vlm-robustbench-repro / code /collapse2.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.6 kB
"""Claim 3, using the paper's OWN definitions (Eqs. 29-31), pinned from source:
order parameter m = 1 - c*
equation of state h = (g_rho/2) m^2 - t m , t = chi_rho - 1 (29)
rescaled m~ = m sqrt(g_rho / (2h)) , t~ = -t / sqrt(2 g_rho h) (30)
universal m~ = sqrt(1 + t~^2) - t~ (31)
My first attempt used m/sqrt(h) and t = 1 - chi: the sign of t was inverted and
the curvature g_rho was omitted from both scales, which is why the collapse was
poor (8.7% median relative error) and t~ never left [-0.42, 0.11].
Everything below is measured from the MFT correlation map F(c), not assumed:
h = 1 - F(1) (the dropout field: how far c=1 falls short)
chi_rho= F'(c*) (slope at the fixed point)
g_rho = -F''(c*) (curvature)
"""
import json, numpy as np
RES = {}
_x, _w = np.polynomial.hermite_e.hermegauss(121)
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 = float(np.clip(c, -1, 1))
z1 = _x[:, None]; z2 = _x[None, :]
s = np.sqrt(max(q, 1e-12))
u1 = s*z1; u2 = s*(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 qstar(sw, sb, p, iters=500):
q = 1.0
for _ in range(iters):
q = (sw**2/(1-p))*E1(lambda u: phi(u)**2, q)+sb**2
return q
def Fmap(c, sw, sb, q):
return (sw**2*E2(lambda a, b: phi(a)*phi(b), q, c)+sb**2)/q
def run():
sb = 0.02
rows = []
for p in (0.001, 0.002, 0.005, 0.01, 0.02, 0.05):
for sw in np.linspace(0.80, 1.60, 161):
q = qstar(sw, sb, p)
# fixed point of F
c = 0.5
for _ in range(800):
c = float(np.clip(Fmap(c, sw, sb, q), -0.999999, 0.999999))
m = 1.0-c
h = 1.0-Fmap(1.0, sw, sb, q) # dropout field
# Expand F about c = 1 (NOT about c*): substituting c = 1 - m into
# c = F(c) gives exactly h = (g_rho/2) m^2 - t m with
# chi_rho = F'(1), g_rho = F''(1).
e = 2e-3
F1 = Fmap(1.0-0*e, sw, sb, q); Fa = Fmap(1.0-e, sw, sb, q); Fb = Fmap(1.0-2*e, sw, sb, q)
chi = (F1-Fa)/e # one-sided F'(1)
g = (F1-2*Fa+Fb)/e**2 # one-sided F''(1)
t = chi-1.0
if h <= 1e-12 or g <= 1e-9 or m <= 1e-9 or m > 0.75: continue # near-critical: small m
mt = m*np.sqrt(g/(2*h)); tt = -t/np.sqrt(2*g*h)
rows.append({"p": p, "sw": round(float(sw), 4), "m": float(m), "t": float(t),
"h": float(h), "g_rho": float(g),
"m_tilde": float(mt), "t_tilde": float(tt),
"universal": float(np.sqrt(1+tt**2)-tt)})
band = [r for r in rows if abs(r["t_tilde"]) <= 4.0]
err = [abs(r["m_tilde"]-r["universal"]) for r in band]
rel = [abs(r["m_tilde"]-r["universal"])/max(r["universal"], 1e-9) for r in band]
RES["claim3_collapse_paper_defs"] = {
"activation": "tanh (smooth class)", "sb": sb,
"p_values": [0.001, 0.002, 0.005, 0.01, 0.02, 0.05],
"n_points": len(rows), "n_in_band": len(band),
"t_tilde_range": [round(min(r["t_tilde"] for r in band), 3),
round(max(r["t_tilde"] for r in band), 3)],
"median_abs_error": round(float(np.median(err)), 6),
"max_abs_error": round(float(np.max(err)), 6),
"median_relative_error": round(float(np.median(rel)), 6),
"points": [{k: (round(v, 5) if isinstance(v, float) else v) for k, v in r.items()}
for r in band]}
R = RES["claim3_collapse_paper_defs"]
print(" %d points, %d in |t~|<=4 ; t~ spans [%.2f, %.2f]"
% (R["n_points"], R["n_in_band"], R["t_tilde_range"][0], R["t_tilde_range"][1]), flush=True)
print(" |m~ - universal|: median %.6f max %.6f | median relative %.4f%%"
% (R["median_abs_error"], R["max_abs_error"], 100*R["median_relative_error"]), flush=True)
for p in [0.001, 0.002, 0.005, 0.01, 0.02, 0.05]:
s = [r for r in band if r["p"] == p]
if not s: continue
e = [abs(r["m_tilde"]-r["universal"]) for r in s]
print(" p=%.2f n=%-3d median|err|=%.6f t~ [%.2f, %.2f]"
% (p, len(s), np.median(e), min(r["t_tilde"] for r in s), max(r["t_tilde"] for r in s)), flush=True)
json.dump(RES, open("collapse_results2.json", "w"), indent=1)
if __name__ == "__main__":
run(); print("DONE")