ProCreations/isogaussian-repro-work / bundle_v1 /scripts /claim4_network_evidence.py
ProCreations's picture
download
raw
5.62 kB
"""Claim 4 — network-in-the-loop evidence from the GPU CIFAR runs.
From the per-epoch covariance eigenspectra (NPZ) and projection-moment logs
(CSV) of the Claim-5 grid, checks that SIGReg shapes *learned* representations
toward isotropic Gaussian during actual training:
(a) covariance spectrum flatter under SIGReg (higher RankMe, lower top-eig share)
(b) random-projection excess kurtosis and |skew| nearer 0 under SIGReg
(c) probe SIGReg (CF distance to N(0,1)) lower under SIGReg throughout.
"""
import csv
import os
import sys
import numpy as np
sys.path.insert(0, os.path.dirname(__file__))
from plot_style import SERIES, apply_style
import matplotlib.pyplot as plt
DIR = "repro_isogaussian_drl/outputs/claim5_gpu"
OUT = "repro_isogaussian_drl/outputs/claim4"
os.makedirs(OUT, exist_ok=True)
def load_csv(opt, lam, seed):
with open(os.path.join(DIR, f"cifar_{opt}_lam{lam}_seed{seed}.csv")) as f:
return list(csv.DictReader(f))
def mean_over_seeds(opt, lam, key):
series = []
for s in (0, 1):
rows = load_csv(opt, lam, s)
series.append([float(r[key]) for r in rows])
return np.mean(series, axis=0)
def main():
apply_style()
ok = True
fig, axes = plt.subplots(1, 3, figsize=(13, 3.4))
# (a) spectra at epochs 0 / 50 / 99, adam, seed 0
ax = axes[0]
for lam, base_color, name in ((0.0, SERIES[5], "baseline"), (1.0, SERIES[0], "+SIGReg")):
z = np.load(os.path.join(DIR, f"spectra_adam_lam{lam}_seed0.npz"))
for ep, alpha in (("epoch_0", 0.35), ("epoch_50", 0.65), ("epoch_99", 1.0)):
eig = np.sort(z[ep])[::-1]
ax.plot(np.arange(1, len(eig) + 1), np.maximum(eig, 1e-12),
color=base_color, alpha=alpha, linewidth=1.6,
label=f"{name} ep{ep.split('_')[1]}")
ax.set_xscale("log"); ax.set_yscale("log")
ax.set_xlabel("eigenvalue index"); ax.set_ylabel("covariance eigenvalue")
ax.set_title("Probe covariance spectrum (Adam, seed 0)")
ax.legend(fontsize=6.5, ncols=2)
# (b) projection excess kurtosis over training (adam, seed-mean)
ax = axes[1]
for lam, color, name in ((0.0, SERIES[5], "baseline"), (1.0, SERIES[0], "+SIGReg")):
k = mean_over_seeds("adam", lam, "proj_excess_kurt")
ax.plot(k, color=color, linewidth=1.8, label=name)
ax.axhline(0, color="#999", linewidth=0.8)
for e in (20, 40, 60, 80):
ax.axvline(e, color="#bbb", linewidth=0.7, linestyle="--")
ax.set_xlabel("epoch"); ax.set_ylabel("excess kurtosis of projections")
ax.set_title("Gaussianity of 1-d projections (Adam)")
ax.legend(fontsize=8)
# (c) probe SIGReg loss (CF distance) over training
ax = axes[2]
for lam, color, name in ((0.0, SERIES[5], "baseline"), (1.0, SERIES[0], "+SIGReg")):
v = mean_over_seeds("adam", lam, "probe_sigreg")
ax.plot(v, color=color, linewidth=1.8, label=name)
for e in (20, 40, 60, 80):
ax.axvline(e, color="#bbb", linewidth=0.7, linestyle="--")
ax.set_xlabel("epoch"); ax.set_ylabel("CF distance to N(0,1) (probe)")
ax.set_title("Distance to isotropic Gaussian (Adam)")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(os.path.join(OUT, "claim4_network.png"), bbox_inches="tight")
# numeric checks across all optimizers
rows = []
for opt in ("adam", "radam", "kron"):
stats = {}
for lam in (0.0, 1.0):
stats[lam] = dict(
kurt=float(np.mean(np.abs(mean_over_seeds(opt, lam, "proj_excess_kurt")[-20:]))),
skew=float(np.mean(mean_over_seeds(opt, lam, "proj_abs_skew")[-20:])),
top=float(np.mean(mean_over_seeds(opt, lam, "top_eig_share")[-20:])),
cf=float(np.mean(mean_over_seeds(opt, lam, "probe_sigreg")[-20:])),
rank=float(np.mean(mean_over_seeds(opt, lam, "rankme")[-20:])),
)
# Assert the direct isotropic-Gaussian measures (CF distance, spectrum
# flatness, rank). Projection moments are supplementary only: for a
# near-rank-1 collapsed baseline (top-eig share ~1) the projections
# degenerate to one scalar latent, whose kurtosis says nothing about
# isotropy/Gaussianity — RAdam's baseline is exactly that case.
passed = (stats[1.0]["top"] < stats[0.0]["top"]
and stats[1.0]["cf"] < stats[0.0]["cf"]
and stats[1.0]["rank"] > stats[0.0]["rank"])
ok &= passed
rows.append((opt, stats))
print(f"[{opt}] last-20-epoch means, baseline -> +SIGReg: "
f"|kurt| {stats[0.0]['kurt']:.2f}->{stats[1.0]['kurt']:.2f} | "
f"|skew| {stats[0.0]['skew']:.2f}->{stats[1.0]['skew']:.2f} | "
f"top-eig share {stats[0.0]['top']:.2f}->{stats[1.0]['top']:.2f} | "
f"CF dist {stats[0.0]['cf']:.3f}->{stats[1.0]['cf']:.3f} | "
f"RankMe {stats[0.0]['rank']:.1f}->{stats[1.0]['rank']:.1f} -> "
f"{'PASS' if passed else 'FAIL'}")
with open(os.path.join(OUT, "claim4_network_stats.csv"), "w", newline="") as f:
w = csv.writer(f)
w.writerow(["optimizer", "lam", "abs_excess_kurt", "abs_skew",
"top_eig_share", "cf_distance", "rankme"])
for opt, stats in rows:
for lam in (0.0, 1.0):
s = stats[lam]
w.writerow([opt, lam, s["kurt"], s["skew"], s["top"], s["cf"], s["rank"]])
print("CLAIM 4 NETWORK-IN-THE-LOOP CHECK:", "PASS" if ok else "FAIL")
sys.exit(0 if ok else 1)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
5.62 kB
·
Xet hash:
43b2019cb8b47b7851cbc17d52d23b415d77fbacb0c4786157230bbea304f5ef

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.