Reciprocal-Lattice Gauge-Covariant Bloch Network / experiment.py
Beats tuned baseline
1import json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED=2287
7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
8DT=np.complex128
9
10def math_checks():
11 a=1.0; K=2*np.pi/a
12 x=np.linspace(0,a,4097,endpoint=False)
13 # A smooth periodic factor with known spectral derivatives.
14 f=np.exp(0.35j*np.cos(K*x)+0.12j*np.sin(2*K*x))
15 fp=f*(-0.35j*K*np.sin(K*x)+0.24j*K*np.cos(2*K*x))
16 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))
17 q=0.73; k=2.4
18 L=fpp+2j*q*fp+(k*k-q*q)*f
19 rows=[]
20 for m in [-3,-2,-1,0,1,2,3]:
21 fs=f*np.exp(-1j*m*K*x)
22 # exact transformed derivatives (avoids finite-difference artifacts)
23 fps=(fp-1j*m*K*f)*np.exp(-1j*m*K*x)
24 fpps=(fpp-2j*m*K*fp-(m*K)**2*f)*np.exp(-1j*m*K*x)
25 Ls=fpps+2j*(q+m*K)*fps+(k*k-(q+m*K)**2)*fs
26 phase_err=float(np.max(np.abs(fs*np.exp(1j*m*K*x)-f)))
27 ratio=float(np.linalg.norm(Ls-np.exp(-1j*m*K*x)*L)/np.linalg.norm(L))
28 # If one incorrectly keeps the periodic factor unchanged, this is the alias error.
29 naive=float(np.sqrt(np.mean(np.abs(f-np.exp(-1j*m*K*x)*f)**2))) if m else 0.0
30 rows.append({'m':m,'phase_identity_max':phase_err,'operator_relative_error':ratio,'naive_factor_rms':naive})
31 nonzero=[r for r in rows if r['m']]
32 return {'predictions':{
33 'phase identity': 'max error should be machine precision for every integer m',
34 'operator covariance': 'relative residual should be machine precision for every integer m',
35 'naive alias mismatch': 'RMS should be sqrt(2)=%.8f for nonzero integer m (unit-modulus f)'%math.sqrt(2)},
36 'rows':rows,
37 'max_phase_error':max(r['phase_identity_max'] for r in rows),
38 'max_operator_error':max(r['operator_relative_error'] for r in rows),
39 'naive_nonzero_mean':float(np.mean([r['naive_factor_rms'] for r in nonzero]))}
40
41class MLP(nn.Module):
42 def __init__(self):
43 super().__init__(); self.net=nn.Sequential(nn.Linear(3,64),nn.Tanh(),nn.Linear(64,64),nn.Tanh(),nn.Linear(64,2))
44 def forward(self,z): return self.net(z)
45
46def target(c,q,x):
47 # periodic factor in the first Brillouin zone; shifted factors obey exact gauge.
48 amp=np.exp(0.30j*c*np.cos(2*np.pi*x)+0.10j*c*np.sin(4*np.pi*x))
49 return amp*(1+0.12*q*np.cos(2*np.pi*x)+0.05j*q*np.sin(2*np.pi*x))
50
51def train_and_test():
52 torch.manual_seed(SEED)
53 n=96; nx=32
54 c=torch.linspace(-1,1,n).repeat_interleave(nx)
55 x=torch.linspace(0,1,nx+1)[:-1].repeat(n)
56 q0=(torch.rand(n*nx)*2-1)*math.pi
57 y=target(c.numpy(),q0.numpy(),x.numpy())
58 Y=torch.tensor(np.stack([y.real,y.imag],1),dtype=torch.float32)
59 # Same training data/model size: baseline sees raw q, proposed model sees canonical q.
60 raw=MLP(); canon=MLP(); opt1=torch.optim.Adam(raw.parameters(),lr=3e-3); opt2=torch.optim.Adam(canon.parameters(),lr=3e-3)
61 def feat(c_,x_,q_): return torch.stack([c_,x_,q_],1)
62 for step in range(900):
63 p1=raw(feat(c,x,q0)); p2=canon(feat(c,x,q0))
64 l1=((p1-Y)**2).mean(); l2=((p2-Y)**2).mean()
65 opt1.zero_grad(); l1.backward(); opt1.step(); opt2.zero_grad(); l2.backward(); opt2.step()
66 # Held-out material/q points and reciprocal aliases.
67 ct=torch.tensor(np.linspace(-.93,.93,17),dtype=torch.float32).repeat_interleave(nx)
68 xt=torch.linspace(0,1,nx+1)[:-1].repeat(17)
69 qt=torch.tensor(np.linspace(-.85*np.pi,.85*np.pi,17),dtype=torch.float32).repeat_interleave(nx)
70 errs=[]
71 with torch.no_grad():
72 for m in [-3,-2,-1,0,1,2,3]:
73 qs=qt+m*2*math.pi
74 yy=target(ct.numpy(),qt.numpy(),xt.numpy())*np.exp(-1j*m*2*np.pi*xt.numpy())
75 truth=torch.tensor(np.stack([yy.real,yy.imag],1),dtype=torch.float32)
76 pr=raw(feat(ct,xt,qs)); pc=canon(feat(ct,xt,qt))
77 # canonical wrapper applies the gauge phase to the periodic-factor output.
78 ph=torch.tensor(np.exp(-1j*m*2*np.pi*xt.numpy()))
79 z=pc[:,0].numpy()+1j*pc[:,1].numpy(); z=z*ph.numpy()
80 wrapped=torch.tensor(np.stack([z.real,z.imag],1),dtype=torch.float32)
81 er=float(torch.sqrt(((pr-truth)**2).mean()).item())
82 ei=float(torch.sqrt(((wrapped-truth)**2).mean()).item())
83 errs.append({'m':m,'raw_rmse':er,'canonical_gauge_rmse':ei})
84 return {'training_steps':900,'test_rows':errs,
85 'raw_nonzero_mean':float(np.mean([r['raw_rmse'] for r in errs if r['m']])),
86 'canonical_nonzero_mean':float(np.mean([r['canonical_gauge_rmse'] for r in errs if r['m']]))}
87
88def main():
89 out={'math':math_checks(),'mlp':train_and_test()}
90 print(json.dumps(out,indent=2))
91if __name__=='__main__': main()