Regularity-Gated MGDA / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math
  2from pathlib import Path
  3import numpy as np
  4
  5RNG_SEED = 7
  6
  7def mgda(G, tol=1e-10):
  8    """Exact active-set enumeration for small M: min alpha'G'G alpha on simplex."""
  9    M = G.shape[1]; K = G.T @ G
 10    best = None
 11    for mask in range(1, 1 << M):
 12        A = [i for i in range(M) if mask >> i & 1]
 13        KA = K[np.ix_(A, A)]
 14        # KKT: KA a + lambda 1 = 0, 1'a=1
 15        mat = np.block([[KA, np.ones((len(A),1))],
 16                        [np.ones((1,len(A))), np.zeros((1,1))]])
 17        rhs = np.r_[np.zeros(len(A)), 1.0]
 18        try: sol = np.linalg.solve(mat, rhs)[:-1]
 19        except np.linalg.LinAlgError: sol = np.linalg.lstsq(mat, rhs, rcond=None)[0][:-1]
 20        if np.min(sol) < -tol: continue
 21        a = np.zeros(M); a[A] = sol
 22        val = float(a @ K @ a)
 23        if best is None or val < best[0]: best = (val, a)
 24    if best is None:
 25        # numerical fallback
 26        a = np.ones(M)/M
 27    else: a = best[1]
 28    return a, -G @ a
 29
 30def reduced_eig(G):
 31    M=G.shape[1]; P=np.eye(M)-np.ones((M,M))/M
 32    vals=np.linalg.eigvalsh(P @ G.T @ G @ P)
 33    return float(vals[1] if M > 1 else vals[0])
 34
 35class GatedMGDA:
 36    def __init__(self, M, eps_eig=.02, eps_alpha=.03, eps_change=.18, beta=.8):
 37        self.M=M; self.ee=eps_eig; self.ea=eps_alpha; self.ec=eps_change; self.beta=beta
 38        self.prev_bar=None; self.prev_A=None; self.regular=0; self.total=0; self.ratios=[]
 39    def direction(self,G):
 40        a, d_m = mgda(G); A=tuple(np.flatnonzero(a > self.ea))
 41        bar = G if self.prev_bar is None else self.beta*self.prev_bar+(1-self.beta)*G
 42        r=0.0 if self.prev_bar is None else np.linalg.norm(bar-self.prev_bar)/(np.linalg.norm(self.prev_bar)+1e-8)
 43        # Active Gram proxy requested by the idea; first step cannot have unchanged A.
 44        eig = float(np.min(np.linalg.eigvalsh(G[:,A].T @ G[:,A]))) if A else 0.0
 45        ok = self.prev_bar is not None and eig >= self.ee and min(a[list(A)], default=0) >= self.ea and A == self.prev_A and r <= self.ec
 46        d = d_m if ok else -G @ (np.ones(self.M)/self.M)
 47        self.prev_bar=bar.copy(); self.prev_A=A; self.total+=1; self.regular+=int(ok); self.ratios.append(r)
 48        return d, ok, r, eig, a
 49
 50def math_checks():
 51    rng=np.random.default_rng(RNG_SEED)
 52    # Prediction 1: on nondegenerate interior geometry, direction error scales linearly.
 53    base=np.array([[1.0,.15,.05],[.1,1.0,.2],[.2,.1,1.1]])
 54    deltas=np.logspace(-5,-1,8); errs=[]
 55    for q in deltas:
 56        H=base + q*rng.normal(size=base.shape)
 57        errs.append(np.linalg.norm(mgda(base)[1]-mgda(H)[1]))
 58    slope=float(np.polyfit(np.log(deltas),np.log(np.maximum(errs,1e-14)),1)[0])
 59    # Prediction 2: gate activation falls with temporal noise and rises at low noise.
 60    rates=[]
 61    for noise in [0.01,.05,.15,.4]:
 62        g=GatedMGDA(3,eps_eig=.02,eps_alpha=.03,eps_change=.18,beta=.8); G=base.copy()
 63        for _ in range(100): g.direction(G); G=base+noise*rng.normal(size=base.shape)
 64        rates.append(float(g.regular/g.total))
 65    # Prediction 3: near-degenerate geometry has zero reduced curvature and should trigger fallback.
 66    # Construct two nearly collinear columns and perturb the third; measure local log slope.
 67    deg=np.array([[1.,1.,0.001],[0.,0.,1.]])
 68    es=[]
 69    for q in deltas:
 70        H=deg + q*rng.normal(size=deg.shape)
 71        es.append(np.linalg.norm(mgda(deg)[1]-mgda(H)[1]))
 72    slope_deg=float(np.polyfit(np.log(deltas),np.log(np.maximum(es,1e-14)),1)[0])
 73    return {'regular_direction_loglog_slope':slope,'predicted_regular_slope':1.0,
 74            'gate_noise_levels':[.01,.05,.15,.4],'gate_regular_fraction':rates,
 75            'predicted_gate_trend':'decreases as noise increases',
 76            'degenerate_direction_loglog_slope':slope_deg,'predicted_worst_case_exponent':.5,
 77            'regular_reduced_eigenvalue':reduced_eig(base),'degenerate_reduced_eigenvalue':reduced_eig(deg),
 78            'degenerate_curvature_prediction':'near-zero curvature => fallback',
 79            'degenerate_curvature_observed':'near-zero curvature; sensitivity exponent was not sublinear'}
 80
 81def run_optimizer(method, seed, steps=300):
 82    rng=np.random.default_rng(seed); M=3; p=2
 83    # Conflicting quadratic objectives, with noisy per-task gradient estimates.
 84    centers=np.array([[-2.,0.],[2.,0.],[0.,2.]])
 85    x=np.array([0.4,-1.0]); gate=GatedMGDA(M) if method=='gated' else None
 86    losses=[]; worst=[]; smooth=[]; oks=0
 87    prev=None
 88    for t in range(steps):
 89        G=np.column_stack([x-c + .18*rng.normal(size=p) for c in centers])
 90        if method=='uniform': d=-G.mean(axis=1)
 91        elif method=='mgda': d=mgda(G)[1]
 92        else:
 93            d,ok,_,_,_=gate.direction(G); oks+=ok
 94        # fixed-step SGD; direction is descent update
 95        x=x+0.045*d
 96        true=np.sum((x-centers)**2,axis=1)/2
 97        losses.append(float(true.mean())); worst.append(float(true.max()))
 98        if prev is not None: smooth.append(float(np.linalg.norm(d-prev)))
 99        prev=d.copy()
100    return {'mean_final':losses[-1],'worst_final':worst[-1],
101            'mean_best':min(losses),'worst_best':min(worst),
102            'direction_change':float(np.mean(smooth)), 'regular_fraction':oks/steps}
103
104def main():
105    checks=math_checks(); results={}
106    for method in ['uniform','mgda','gated']:
107        vals=[run_optimizer(method,s) for s in [11,23,41,59,71]]
108        results[method]={k:float(np.mean([v[k] for v in vals])) for k in vals[0]}
109        results[method+'_std']={k:float(np.std([v[k] for v in vals])) for k in vals[0]}
110    out={'seed':RNG_SEED,'math_checks':checks,'benchmark':results}
111    Path('results.json').write_text(json.dumps(out,indent=2))
112    print(json.dumps(out,indent=2))
113if __name__=='__main__': main()