Epoch-Frozen Masked Low-Rank Candidate Encoder / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 1import json, time
 2from pathlib import Path
 3import numpy as np
 4import torch
 5from torch import nn
 6
 7SEED = 612
 8np.random.seed(SEED); torch.manual_seed(SEED)
 9try:
10    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
11except Exception:
12    device = torch.device('cpu')
13
14class MaskedLowRank:
15    def __init__(self, d, m, p, ridge=1e-3):
16        self.d, self.m, self.p, self.ridge, self.U = d, m, p, ridge, None
17    def fit(self, Y):
18        n = Y.shape[0]; p = self.p
19        S = (Y.T @ Y) / n
20        diag = np.diag(S).copy() / p
21        S = S / (p*p); np.fill_diagonal(S, diag); S = (S + S.T) / 2
22        _, V = np.linalg.eigh(S); self.U = V[:, -self.m:]
23    def encode(self, X, M):
24        # Batched version of (U_O' U_O + lambda I)^-1 U_O' x_O.
25        A = self.U[None, :, :] * M[:, :, None]
26        G = np.einsum('bdi,bdj->bij', A, A)
27        rhs = np.einsum('bdi,bd->bi', A, X)
28        lam = self.ridge * max(float(np.median(np.diagonal(G, axis1=1, axis2=2))), 1e-8)
29        return np.linalg.solve(G + lam*np.eye(self.m), rhs).astype(np.float32)
30
31def orthobasis(d, r, rng):
32    return np.linalg.qr(rng.normal(size=(d,r)))[0]
33
34def make_split(n, d, r, p, k, Q, task, rng):
35    z = rng.normal(size=(n*k,r)).astype(np.float32)
36    X = (z @ Q.T).astype(np.float32)
37    C = rng.normal(size=(n*k,8)).astype(np.float32)
38    M = rng.random(X.shape) < p
39    utility = z @ task + .35*C[:,0]*z[:,0] + .12*rng.normal(size=n*k)
40    return X, C, M, utility.reshape(n,k).argmax(1)
41
42class RankMLP(nn.Module):
43    def __init__(self, f):
44        super().__init__(); self.net=nn.Sequential(nn.Linear(f+8,32),nn.ReLU(),nn.Linear(32,1))
45    def forward(self,x,c): return self.net(torch.cat((x,c),1)).squeeze(1)
46
47def evaluate(model, feat, C, labels, n, k):
48    with torch.no_grad():
49        logits=model(torch.as_tensor(feat,device=device),torch.as_tensor(C,device=device)).reshape(n,k)
50        y=torch.as_tensor(labels,device=device)
51        return float(nn.functional.cross_entropy(logits,y)), float((logits.argmax(1)==y).float().mean())
52
53def run(p):
54    rng=np.random.default_rng(SEED+int(p*100)); d=r=m=32 if False else (48); r=m=6; k=8; ntr,nva=320,120
55    Q=orthobasis(d,r,rng); task=rng.normal(size=r).astype(np.float32)
56    tr=make_split(ntr,d,r,p,k,Q,task,rng); va=make_split(nva,d,r,p,k,Q,task,rng)
57    X,C,M,y=tr; XV,CV,MV,yv=va
58    enc=MaskedLowRank(d,m,p); enc.fit(X*M)
59    Z,ZV=enc.encode(X,M),enc.encode(XV,MV)
60    controls={'zero_imputed':(d,X*M,XV*MV),'frozen_latent':(m,Z,ZV),'unfrozen_pca':(m,Z,ZV)}
61    models={name:RankMLP(f).to(device) for name,(f,_,_) in controls.items()}
62    opts={name:torch.optim.Adam(mod.parameters(),lr=.01) for name,mod in models.items()}
63    hist={name:[] for name in models}; Ctr=np.asarray(C); Cval=np.asarray(CV)
64    for ep in range(7):
65        for name,mod in models.items():
66            if name=='unfrozen_pca':
67                # Refitting at every epoch is the unfrozen control; frozen method fits once.
68                e=MaskedLowRank(d,m,p); e.fit(X*M); F=e.encode(X,M); FV=e.encode(XV,MV)
69            else: F,FV=controls[name][1:]
70            mod.train(); opt=opts[name]
71            logits=mod(torch.as_tensor(F,device=device),torch.as_tensor(Ctr,device=device)).reshape(ntr,k)
72            loss=nn.functional.cross_entropy(logits,torch.as_tensor(y,device=device))
73            opt.zero_grad(); loss.backward(); opt.step(); mod.eval()
74            hist[name].append(evaluate(mod,FV,Cval,yv,nva,k))
75    out={name:{'final_loss':h[-1][0],'final_accuracy':h[-1][1],'best_accuracy':max(x[1] for x in h),'val_loss_std_last4':float(np.std([x[0] for x in h[-4:]]))} for name,h in hist.items()}
76    out['latent_ranker_flop_ratio']=(m+8)/(d+8)
77    return out
78
79def math_check():
80    rng=np.random.default_rng(SEED); d,r,N=24,4,2500; p=.6
81    Q=orthobasis(d,r,rng); z=rng.normal(size=(N,r)).astype(np.float32); X=z@Q.T; M=rng.random(X.shape)<p
82    e=MaskedLowRank(d,r,p); e.fit(X*M); cos=np.linalg.svd(Q.T@e.U,compute_uv=False)
83    Z=e.encode(X[:500],M[:500]); Xhat=Z@e.U.T
84    rel=np.mean(np.linalg.norm(X[:500]-Xhat,axis=1)/(np.linalg.norm(X[:500],axis=1)+1e-8))
85    # Empirical covariance correction check: corrected expectation should match XX'/N.
86    S=(X*M).T@(X*M)/N; corr=S/(p*p); np.fill_diagonal(corr,np.diag(S)/p)
87    coverr=np.linalg.norm(corr-X.T@X/N)/np.linalg.norm(X.T@X/N)
88    return {'principal_cosines':cos.tolist(),'relative_reconstruction_error':float(rel),'relative_covariance_error':float(coverr)}
89
90if __name__=='__main__':
91    t=time.time(); result={'device':str(device),'math_check':math_check(),'experiments':{str(p):run(p) for p in (.25,.5,.75)}}
92    Path('results.json').write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2)); print('elapsed_sec',round(time.time()-t,2))