SabaPivot's picture
download
raw
6.35 kB
#!/usr/bin/env python3
"""Independent executable audit of TFMixer equations 1--14.
This is not the unreleased benchmark implementation. It implements the two
novel modules exactly at the equation level and checks shapes, masking,
gradient flow, reconstruction, fixed-capacity query aggregation, and dual
mixing on irregular multivariate observations.
"""
from pathlib import Path
import json
import math
import torch
from torch import nn
torch.manual_seed(260200582)
torch.set_num_threads(2)
class LearnableNUDFT(nn.Module):
def __init__(self, n_variables=3, k=6, hidden=32):
super().__init__()
# Positive, learnable frequencies. Softplus avoids invalid negatives.
init = torch.linspace(0.15, 1.4, k)
self.raw_omega = nn.Parameter(torch.log(torch.expm1(init)))
self.refine = nn.Sequential(nn.Linear(2*k, hidden), nn.GELU(),
nn.Linear(hidden, 2*k))
self.k = k
@property
def omega(self):
return torch.nn.functional.softplus(self.raw_omega)
def forward(self, t, v, mask):
# t:[B,L], v/mask:[B,N,L]; equations 1--4.
phase = 2 * math.pi * t[:, None, :, None] * self.omega[None,None,None,:]
z = mask.sum(-1, keepdim=True).clamp_min(1.0)
weighted = (mask*v)[..., None]
real = (weighted * phase.cos()).sum(2) / z
imag = -(weighted * phase.sin()).sum(2) / z
raw = torch.cat([real, imag], -1)
refined = self.refine(raw)
return raw, refined
def inverse(self, coeff, q):
# q:[B,H], coeff:[B,N,2K]; equation 14.
real, imag = coeff[...,:self.k], coeff[...,self.k:]
phase = 2 * math.pi * q[:,None,:,None] * self.omega[None,None,None,:]
return (real[:,:,None,:]*phase.cos() - imag[:,:,None,:]*phase.sin()).sum(-1)
class ContinuousTimeEmbedding(nn.Module):
def __init__(self, d=8):
super().__init__(); self.w=nn.Parameter(torch.randn(d)); self.a=nn.Parameter(torch.zeros(d))
def forward(self,t):
y=torch.sin(t[...,None]*self.w+self.a)
y[...,0]=t*self.w[0]+self.a[0]
return y
class QueryPatchMixer(nn.Module):
def __init__(self,n_variables=3,p=12,w=4,d=16):
super().__init__(); self.p=p; self.w=w; self.d=d; self.n=n_variables
self.time=ContinuousTimeEmbedding(8)
self.patch_proj=nn.Linear(10,d) # mean time embedding + mean value + mask
self.query=nn.Parameter(torch.randn(w,d)/math.sqrt(d))
self.pos=nn.Parameter(torch.randn(p,d)/math.sqrt(d))
self.patch_mlp=nn.Sequential(nn.Linear(w,2*w),nn.GELU(),nn.Linear(2*w,w))
self.var_mlp=nn.Sequential(nn.Linear(n_variables,2*n_variables),nn.GELU(),nn.Linear(2*n_variables,n_variables))
self.norm1=nn.LayerNorm(d); self.norm2=nn.LayerNorm(d)
def forward(self,t,v,mask):
# Equal time windows, variable observation counts: transformable patches.
b,n,l=v.shape; edges=torch.linspace(0,1,self.p+1,device=t.device)
patches=[]; counts=[]
te=self.time(t)
for j in range(self.p):
inwin=((t>=edges[j]) & (t<(edges[j+1] if j+1<self.p else 1.00001))).float()
m=mask*inwin[:,None,:]; z=m.sum(-1,keepdim=True).clamp_min(1)
meanv=(m*v).sum(-1,keepdim=True)/z
meant=(m[...,None]*te[:,None,:,:]).sum(-2)/z
present=(m.sum(-1,keepdim=True)>0).float()
patches.append(self.patch_proj(torch.cat([meant,meanv,present],-1)))
counts.append(m.sum(-1))
h=torch.stack(patches,2) # [B,N,P,D]
logits=torch.einsum('wd,bnpd->bnwp',self.query,h+self.pos[None,None])/math.sqrt(self.d)
attn=logits.softmax(-1)
hw=torch.einsum('bnwp,bnpd->bnwd',attn,h) # fixed W tokens
# Equation 12: mix W, then equation 13: mix N.
x=hw.permute(0,1,3,2); x=self.norm1((x+self.patch_mlp(x)).permute(0,1,3,2))
y=x.permute(0,2,3,1); y=self.norm2((y+self.var_mlp(y)).permute(0,3,1,2))
return y,attn,torch.stack(counts,-1)
def make_batch(b=24,n=3,l=120):
t=torch.sort(torch.rand(b,l),-1).values
f=torch.tensor([0.7,1.1,1.35])
phase=torch.tensor([0.1,0.7,1.2])
v=torch.sin(2*math.pi*f[None,:,None]*t[:,None,:]+phase[None,:,None])
v=v+0.25*torch.sin(2*math.pi*0.2*t[:,None,:])+0.03*torch.randn_like(v)
# Different sampling densities per variable and a fully empty patch edge case.
probs=torch.tensor([0.85,0.55,0.30])[None,:,None]
mask=(torch.rand_like(v)<probs).float(); mask[:,:,48:58]=0
return t,v,mask
t,v,mask=make_batch(); model=LearnableNUDFT(); opt=torch.optim.Adam(model.parameters(),lr=0.02)
losses=[]
for step in range(301):
_,coeff=model(t,v,mask); recon=model.inverse(coeff,t)
loss=(((recon-v)*mask)**2).sum()/mask.sum()
opt.zero_grad(); loss.backward(); opt.step(); losses.append(float(loss.detach()))
local=QueryPatchMixer(); hw,attn,counts=local(t,v,mask)
probe=(hw.square().mean()+attn.square().mean()); probe.backward()
grad_norm=float(local.query.grad.norm())
result={
'scope':'Equation-level independent implementation; not a real-benchmark rerun.',
'input_shape':list(v.shape), 'valid_observations':int(mask.sum()),
'nudft_raw_shape':[24,3,12], 'nudft_refined_shape':list(coeff.shape),
'inverse_shape':list(recon.shape), 'loss_initial':losses[0],
'loss_final':losses[-1], 'loss_reduction_percent':100*(1-losses[-1]/losses[0]),
'learned_frequencies':[float(x) for x in model.omega.detach()],
'query_tokens_shape':list(hw.shape), 'attention_shape':list(attn.shape),
'attention_rows_sum_max_error':float((attn.sum(-1)-1).abs().max()),
'min_patch_observations':float(counts.min()), 'max_patch_observations':float(counts.max()),
'query_gradient_norm':grad_norm,
'checks':{
'masked_nudft_and_inverse_train':losses[-1] < losses[0]*0.4,
'fixed_capacity_W_tokens':list(hw.shape)==[24,3,4,16],
'attention_normalized':float((attn.sum(-1)-1).abs().max())<1e-6,
'empty_patch_handled':float(counts.min())==0.0,
'query_is_learnable':grad_norm>0,
}
}
assert all(result['checks'].values()),result
out=Path(__file__).resolve().parent/'outputs'; out.mkdir(exist_ok=True)
(out/'module_verification.json').write_text(json.dumps(result,indent=2)+'\n')
print(json.dumps(result,indent=2))

Xet Storage Details

Size:
6.35 kB
·
Xet hash:
a84e89d224fa6fa82bb83f27bb4ee2d664f00f261e3d43e678ebe5fcf81f7bca

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.