Epoch-Frozen Masked Low-Rank Candidate Encoder / experiment.py
Beats tuned baseline
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))