import json, time from pathlib import Path import numpy as np import torch from torch import nn SEED = 612 np.random.seed(SEED); torch.manual_seed(SEED) try: device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') except Exception: device = torch.device('cpu') class MaskedLowRank: def __init__(self, d, m, p, ridge=1e-3): self.d, self.m, self.p, self.ridge, self.U = d, m, p, ridge, None def fit(self, Y): n = Y.shape[0]; p = self.p S = (Y.T @ Y) / n diag = np.diag(S).copy() / p S = S / (p*p); np.fill_diagonal(S, diag); S = (S + S.T) / 2 _, V = np.linalg.eigh(S); self.U = V[:, -self.m:] def encode(self, X, M): # Batched version of (U_O' U_O + lambda I)^-1 U_O' x_O. A = self.U[None, :, :] * M[:, :, None] G = np.einsum('bdi,bdj->bij', A, A) rhs = np.einsum('bdi,bd->bi', A, X) lam = self.ridge * max(float(np.median(np.diagonal(G, axis1=1, axis2=2))), 1e-8) return np.linalg.solve(G + lam*np.eye(self.m), rhs).astype(np.float32) def orthobasis(d, r, rng): return np.linalg.qr(rng.normal(size=(d,r)))[0] def make_split(n, d, r, p, k, Q, task, rng): z = rng.normal(size=(n*k,r)).astype(np.float32) X = (z @ Q.T).astype(np.float32) C = rng.normal(size=(n*k,8)).astype(np.float32) M = rng.random(X.shape) < p utility = z @ task + .35*C[:,0]*z[:,0] + .12*rng.normal(size=n*k) return X, C, M, utility.reshape(n,k).argmax(1) class RankMLP(nn.Module): def __init__(self, f): super().__init__(); self.net=nn.Sequential(nn.Linear(f+8,32),nn.ReLU(),nn.Linear(32,1)) def forward(self,x,c): return self.net(torch.cat((x,c),1)).squeeze(1) def evaluate(model, feat, C, labels, n, k): with torch.no_grad(): logits=model(torch.as_tensor(feat,device=device),torch.as_tensor(C,device=device)).reshape(n,k) y=torch.as_tensor(labels,device=device) return float(nn.functional.cross_entropy(logits,y)), float((logits.argmax(1)==y).float().mean()) def run(p): 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 Q=orthobasis(d,r,rng); task=rng.normal(size=r).astype(np.float32) tr=make_split(ntr,d,r,p,k,Q,task,rng); va=make_split(nva,d,r,p,k,Q,task,rng) X,C,M,y=tr; XV,CV,MV,yv=va enc=MaskedLowRank(d,m,p); enc.fit(X*M) Z,ZV=enc.encode(X,M),enc.encode(XV,MV) controls={'zero_imputed':(d,X*M,XV*MV),'frozen_latent':(m,Z,ZV),'unfrozen_pca':(m,Z,ZV)} models={name:RankMLP(f).to(device) for name,(f,_,_) in controls.items()} opts={name:torch.optim.Adam(mod.parameters(),lr=.01) for name,mod in models.items()} hist={name:[] for name in models}; Ctr=np.asarray(C); Cval=np.asarray(CV) for ep in range(7): for name,mod in models.items(): if name=='unfrozen_pca': # Refitting at every epoch is the unfrozen control; frozen method fits once. e=MaskedLowRank(d,m,p); e.fit(X*M); F=e.encode(X,M); FV=e.encode(XV,MV) else: F,FV=controls[name][1:] mod.train(); opt=opts[name] logits=mod(torch.as_tensor(F,device=device),torch.as_tensor(Ctr,device=device)).reshape(ntr,k) loss=nn.functional.cross_entropy(logits,torch.as_tensor(y,device=device)) opt.zero_grad(); loss.backward(); opt.step(); mod.eval() hist[name].append(evaluate(mod,FV,Cval,yv,nva,k)) 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()} out['latent_ranker_flop_ratio']=(m+8)/(d+8) return out def math_check(): rng=np.random.default_rng(SEED); d,r,N=24,4,2500; p=.6 Q=orthobasis(d,r,rng); z=rng.normal(size=(N,r)).astype(np.float32); X=z@Q.T; M=rng.random(X.shape)