Reciprocal-Lattice Gauge-Covariant Bloch Network / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 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()