import json, math, random import numpy as np import torch from torch import nn SEED=2287 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) DT=np.complex128 def math_checks(): a=1.0; K=2*np.pi/a x=np.linspace(0,a,4097,endpoint=False) # A smooth periodic factor with known spectral derivatives. f=np.exp(0.35j*np.cos(K*x)+0.12j*np.sin(2*K*x)) fp=f*(-0.35j*K*np.sin(K*x)+0.24j*K*np.cos(2*K*x)) fpp=f*((-0.35j*K*np.sin(K*x)+0.24j*K*np.cos(2*K*x))**2-0.35j*K*K*np.cos(K*x)-0.48j*K*K*np.sin(2*K*x)) q=0.73; k=2.4 L=fpp+2j*q*fp+(k*k-q*q)*f rows=[] for m in [-3,-2,-1,0,1,2,3]: fs=f*np.exp(-1j*m*K*x) # exact transformed derivatives (avoids finite-difference artifacts) fps=(fp-1j*m*K*f)*np.exp(-1j*m*K*x) fpps=(fpp-2j*m*K*fp-(m*K)**2*f)*np.exp(-1j*m*K*x) Ls=fpps+2j*(q+m*K)*fps+(k*k-(q+m*K)**2)*fs phase_err=float(np.max(np.abs(fs*np.exp(1j*m*K*x)-f))) ratio=float(np.linalg.norm(Ls-np.exp(-1j*m*K*x)*L)/np.linalg.norm(L)) # If one incorrectly keeps the periodic factor unchanged, this is the alias error. naive=float(np.sqrt(np.mean(np.abs(f-np.exp(-1j*m*K*x)*f)**2))) if m else 0.0 rows.append({'m':m,'phase_identity_max':phase_err,'operator_relative_error':ratio,'naive_factor_rms':naive}) nonzero=[r for r in rows if r['m']] return {'predictions':{ 'phase identity': 'max error should be machine precision for every integer m', 'operator covariance': 'relative residual should be machine precision for every integer m', 'naive alias mismatch': 'RMS should be sqrt(2)=%.8f for nonzero integer m (unit-modulus f)'%math.sqrt(2)}, 'rows':rows, 'max_phase_error':max(r['phase_identity_max'] for r in rows), 'max_operator_error':max(r['operator_relative_error'] for r in rows), 'naive_nonzero_mean':float(np.mean([r['naive_factor_rms'] for r in nonzero]))} class MLP(nn.Module): def __init__(self): super().__init__(); self.net=nn.Sequential(nn.Linear(3,64),nn.Tanh(),nn.Linear(64,64),nn.Tanh(),nn.Linear(64,2)) def forward(self,z): return self.net(z) def target(c,q,x): # periodic factor in the first Brillouin zone; shifted factors obey exact gauge. amp=np.exp(0.30j*c*np.cos(2*np.pi*x)+0.10j*c*np.sin(4*np.pi*x)) return amp*(1+0.12*q*np.cos(2*np.pi*x)+0.05j*q*np.sin(2*np.pi*x)) def train_and_test(): torch.manual_seed(SEED) n=96; nx=32 c=torch.linspace(-1,1,n).repeat_interleave(nx) x=torch.linspace(0,1,nx+1)[:-1].repeat(n) q0=(torch.rand(n*nx)*2-1)*math.pi y=target(c.numpy(),q0.numpy(),x.numpy()) Y=torch.tensor(np.stack([y.real,y.imag],1),dtype=torch.float32) # Same training data/model size: baseline sees raw q, proposed model sees canonical q. raw=MLP(); canon=MLP(); opt1=torch.optim.Adam(raw.parameters(),lr=3e-3); opt2=torch.optim.Adam(canon.parameters(),lr=3e-3) def feat(c_,x_,q_): return torch.stack([c_,x_,q_],1) for step in range(900): p1=raw(feat(c,x,q0)); p2=canon(feat(c,x,q0)) l1=((p1-Y)**2).mean(); l2=((p2-Y)**2).mean() opt1.zero_grad(); l1.backward(); opt1.step(); opt2.zero_grad(); l2.backward(); opt2.step() # Held-out material/q points and reciprocal aliases. ct=torch.tensor(np.linspace(-.93,.93,17),dtype=torch.float32).repeat_interleave(nx) xt=torch.linspace(0,1,nx+1)[:-1].repeat(17) qt=torch.tensor(np.linspace(-.85*np.pi,.85*np.pi,17),dtype=torch.float32).repeat_interleave(nx) errs=[] with torch.no_grad(): for m in [-3,-2,-1,0,1,2,3]: qs=qt+m*2*math.pi yy=target(ct.numpy(),qt.numpy(),xt.numpy())*np.exp(-1j*m*2*np.pi*xt.numpy()) truth=torch.tensor(np.stack([yy.real,yy.imag],1),dtype=torch.float32) pr=raw(feat(ct,xt,qs)); pc=canon(feat(ct,xt,qt)) # canonical wrapper applies the gauge phase to the periodic-factor output. ph=torch.tensor(np.exp(-1j*m*2*np.pi*xt.numpy())) z=pc[:,0].numpy()+1j*pc[:,1].numpy(); z=z*ph.numpy() wrapped=torch.tensor(np.stack([z.real,z.imag],1),dtype=torch.float32) er=float(torch.sqrt(((pr-truth)**2).mean()).item()) ei=float(torch.sqrt(((wrapped-truth)**2).mean()).item()) errs.append({'m':m,'raw_rmse':er,'canonical_gauge_rmse':ei}) return {'training_steps':900,'test_rows':errs, 'raw_nonzero_mean':float(np.mean([r['raw_rmse'] for r in errs if r['m']])), 'canonical_nonzero_mean':float(np.mean([r['canonical_gauge_rmse'] for r in errs if r['m']]))} def main(): out={'math':math_checks(),'mlp':train_and_test()} print(json.dumps(out,indent=2)) if __name__=='__main__': main()