import itertools, math, json, random import numpy as np import torch from torch import nn SEED=180 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) def partitions(n, mx=None): if n==0: return [()] if mx is None or mx>n: mx=n out=[] for k in range(mx,0,-1): for rest in partitions(n-k,k): out.append((k,)+rest) return out def matchings(items): items=list(items) if not items: yield (); return a=items[0] for j in range(1,len(items)): b=items[j] for rest in matchings(items[1:j]+items[j+1:]): yield ((a,b),)+rest def delta_for_partition(lam,n): pairs=[]; pos=0 for k in lam: cyc=list(range(pos,pos+k)); pos+=k for j,i in enumerate(cyc): pairs.append((i,n+cyc[(j+1)%k])) return tuple(pairs) def epsilon(n): return tuple((i,n+i) for i in range(n)) def component_type(m1,m2,n): adj=[[] for _ in range(2*n)] for a,b in list(m1)+list(m2): adj[a].append(b); adj[b].append(a) seen=set(); sizes=[] for s in range(2*n): if s not in seen: stack=[s]; seen.add(s); c=0 while stack: u=stack.pop(); c+=1 for v in adj[u]: if v not in seen: seen.add(v); stack.append(v) sizes.append(c//2) return tuple(sorted(sizes,reverse=True)) def coefficient_table(n): ps=partitions(n); idx={p:i for i,p in enumerate(ps)} tab=np.zeros((len(ps),len(ps),len(ps)),dtype=np.int64) eps=epsilon(n); ms=list(matchings(range(2*n))) for li,lam in enumerate(ps): dl=delta_for_partition(lam,n) for d in ms: tab[li,idx[component_type(d,dl,n)],idx[component_type(d,eps,n)]]+=1 return ps,tab,len(ms) def traces_products(M, ps): vals={}; P=torch.eye(M.shape[-1],device=M.device,dtype=M.dtype) for k in range(1,max(max(p) for p in ps)+1): P=P@M; vals[k]=torch.diagonal(P,dim1=-2,dim2=-1).sum(-1) return torch.stack([torch.prod(torch.stack([vals[k] for k in p],-1),-1) for p in ps],-1) class WickLayer(nn.Module): def __init__(self,d,q,n,hidden=24): super().__init__(); self.q=q self.a=nn.Sequential(nn.Linear(d,hidden),nn.Tanh(),nn.Linear(hidden,q*q)) self.b=nn.Sequential(nn.Linear(d,hidden),nn.Tanh(),nn.Linear(hidden,q*q)) ps,tab,_=coefficient_table(n); self.ps=ps self.register_buffer('coef',torch.tensor(tab,dtype=torch.float32)) self.head=nn.Sequential(nn.LayerNorm(len(ps)),nn.Linear(len(ps),16),nn.Tanh(),nn.Linear(16,1)) def forward(self,h): B=h.shape[0]; q=self.q U=self.a(h).reshape(B,q,q); V=self.b(h).reshape(B,q,q) eye=torch.eye(q,device=h.device,dtype=h.dtype) A=U@U.transpose(-1,-2)+.1*eye; C=V@V.transpose(-1,-2)+.1*eye pa=traces_products(A,self.ps); pb=traces_products(C,self.ps) F=torch.einsum('lmn,bm,bn->bl',self.coef,pa,pb) return self.head(F).squeeze(-1) class DeepSets(nn.Module): def __init__(self,d,width=32): super().__init__() self.phi=nn.Sequential(nn.Linear(d,width),nn.Tanh(),nn.Linear(width,width),nn.Tanh()) self.rho=nn.Sequential(nn.Linear(width,16),nn.Tanh(),nn.Linear(16,1)) def forward(self,x): return self.rho(self.phi(x).mean(1)).squeeze(-1) class WickModel(nn.Module): def __init__(self,d): super().__init__(); self.enc=nn.Sequential(nn.Linear(d,24),nn.Tanh(),nn.Linear(24,16),nn.Tanh()); self.w=WickLayer(16,8,3) def forward(self,x): return self.w(self.enc(x).mean(1)) def pmat(M,p): return np.prod([np.trace(np.linalg.matrix_power(M,k)) for k in p]) def math_check(): n=3; ps,tab,nm=coefficient_table(n) A=np.array([[1.2,.2,-.1],[.2,.8,.15],[-.1,.15,1.1]]) B=np.array([[.9,-.1,.2],[-.1,1.3,.05],[.2,.05,.7]]) exact=np.array([sum(tab[li,mi,ni]*pmat(A,mu)*pmat(B,nu) for mi,mu in enumerate(ps) for ni,nu in enumerate(ps)) for li in range(len(ps))]) rng=np.random.default_rng(SEED); vals=[] for _ in range(30000): Z=rng.standard_normal((3,3)); M=A@Z@B@Z.T vals.append([pmat(M,p) for p in ps]) mc=np.mean(vals,0); rel=np.max(np.abs(mc-exact)/(1+np.abs(exact))) return {'degree':n,'partitions':ps,'matchings':nm,'exact':exact.tolist(),'mc':mc.tolist(),'max_relative_error':float(rel)} def train(model,x,y,steps,device): model.to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3); lossfn=nn.MSELoss(); model.train() for _ in range(steps): opt.zero_grad(); loss=lossfn(model(x),y); loss.backward(); opt.step() return model def benchmark(): rng=np.random.default_rng(SEED); N=600; sets=5; d=4 x=rng.normal(size=(N,sets,d)).astype('float32'); s=x.sum(1) target=(.7*s[:,0]**2+.35*s[:,1]*s[:,2]+.25*np.sum(x[:,:,3]**2,1)+.15*np.prod(s[:,:3],1)).astype('float32') target += .03*rng.normal(size=N).astype('float32') perm=rng.permutation(N); tr=perm[:60]; te=perm[60:] device='cuda' if torch.cuda.is_available() else 'cpu' try: xt=torch.tensor(x); yt=torch.tensor(target); trainx=xt[tr].to(device); trainy=yt[tr].to(device); testx=xt[te].to(device); testy=yt[te].to(device) torch.manual_seed(SEED); base=train(DeepSets(d,32),trainx,trainy,700,device) torch.manual_seed(SEED); idea=train(WickModel(d),trainx,trainy,700,device) except Exception as e: if device!='cuda': raise device='cpu'; xt=torch.tensor(x); yt=torch.tensor(target); trainx=xt[tr]; trainy=yt[tr]; testx=xt[te]; testy=yt[te] torch.manual_seed(SEED); base=train(DeepSets(d,32),trainx,trainy,700,device) torch.manual_seed(SEED); idea=train(WickModel(d),trainx,trainy,700,device) base.eval(); idea.eval() with torch.no_grad(): bm=float(torch.mean((base(testx)-testy)**2).sqrt()); im=float(torch.mean((idea(testx)-testy)**2).sqrt()) xx=testx[:8]; p=torch.randperm(sets,device=testx.device) bi=float(torch.max(torch.abs(base(xx)-base(xx[:,p])))); ii=float(torch.max(torch.abs(idea(xx)-idea(xx[:,p])))) return {'device':device,'train_examples':len(tr),'test_rmse_deepsets':bm,'test_rmse_wick':im,'perm_diff_deepsets':bi,'perm_diff_wick':ii,'params_deepsets':sum(p.numel() for p in base.parameters()),'params_wick':sum(p.numel() for p in idea.parameters())} if __name__=='__main__': print(json.dumps({'math':math_check(),'benchmark':benchmark()},indent=2))