Buckets:
| #!/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 | |
| 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.