Exact Elasticity-Complex Message Passing / complex_mvp.py
Beats tuned baseline
1import random
2import numpy as np
3import torch
4from torch import nn
5from scipy.spatial import Delaunay
6
7SEED = 113
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9
10def make_mesh(nx=2, ny=2, nz=1):
11 pts = np.array([(i/nx, j/ny, k/max(1,nz))
12 for k in range(nz+1) for j in range(ny+1) for i in range(nx+1)], float)
13 tet = Delaunay(pts).simplices.copy()
14 for q, t in enumerate(tet):
15 a,b,c,d = pts[t]
16 if np.linalg.det(np.stack([b-a,c-a,d-a])) < 0:
17 tet[q,[0,1]] = tet[q,[1,0]]
18 return pts, tet
19
20def build_incidence(tet, nvert):
21 edges = sorted({tuple(sorted((int(t[i]),int(t[j]))))
22 for t in tet for i in range(4) for j in range(i+1,4)})
23 faces = sorted({tuple(sorted((int(t[i]),int(t[j]),int(t[k]))))
24 for t in tet for i in range(4) for j in range(i+1,4)
25 for k in range(j+1,4)})
26 ei = {e:i for i,e in enumerate(edges)}
27 fi = {f:i for i,f in enumerate(faces)}
28 D0 = np.zeros((len(edges), nvert))
29 for r,(a,b) in enumerate(edges): D0[r,a],D0[r,b] = -1,1
30 D1 = np.zeros((len(faces),len(edges)))
31 for r,f in enumerate(faces):
32 for omit in range(3):
33 sub = tuple(f[j] for j in range(3) if j != omit)
34 e = tuple(sorted(sub)); sign = (-1)**omit
35 if sub != e: sign *= -1
36 D1[r,ei[e]] += sign
37 D2 = np.zeros((len(tet),len(faces)))
38 for r,t in enumerate(tet):
39 # positively oriented tetrahedron boundary: omit vertex with alternating sign.
40 for omit in range(4):
41 sub = tuple(int(t[j]) for j in range(4) if j != omit)
42 f = tuple(sorted(sub)); sign = (-1)**omit
43 inv = sum(sub[i] > sub[j] for i in range(3) for j in range(i+1,3))
44 if inv % 2: sign *= -1
45 D2[r,fi[f]] += sign
46 return D0,D1,D2,edges,faces
47
48class ExactComplex(nn.Module):
49 def __init__(self, D0, D1, D2):
50 super().__init__()
51 self.register_buffer('D0', torch.tensor(D0,dtype=torch.float32))
52 self.register_buffer('D1', torch.tensor(D1,dtype=torch.float32))
53 self.register_buffer('D2', torch.tensor(D2,dtype=torch.float32))
54 self.net = nn.Sequential(nn.Linear(1,16),nn.Tanh(),nn.Linear(16,1))
55 def forward(self, x):
56 # Exact compatible signal is annihilated by D1 after D0; predict a scalar
57 # response from a defect magnitude, with no learned cross-order operator.
58 defect = self.D1 @ x
59 return self.net((defect.pow(2).mean().sqrt()).reshape(1,1)).reshape(1)
60
61class Unconstrained(nn.Module):
62 def __init__(self, ne, nf):
63 super().__init__()
64 self.map = nn.Parameter(torch.randn(nf,ne)*0.15)
65 self.net = nn.Sequential(nn.Linear(1,16),nn.Tanh(),nn.Linear(16,1))
66 def forward(self,x):
67 g = self.map @ x
68 return self.net((g.pow(2).mean().sqrt()).reshape(1,1)).reshape(1)
69
70def run():
71 pts,tet = make_mesh(); D0,D1,D2,edges,faces = build_incidence(tet,len(pts))
72 r10=np.linalg.norm(D1@D0)/(np.linalg.norm(D0)+1e-12)
73 r21=np.linalg.norm(D2@D1)/(np.linalg.norm(D1)+1e-12)
74 rng=np.random.default_rng(SEED)
75 # Direct compatible-field test: D0 u must be killed by D1. A same-size
76 # unconstrained cross-order map has no such algebraic guarantee.
77 compatible_leaks=[]; unconstrained_leaks=[]
78 A=rng.normal(size=(D1.shape[0],D0.shape[0]))
79 for _ in range(100):
80 u=rng.normal(size=D0.shape[1]); e=D0@u
81 compatible_leaks.append(np.linalg.norm(D1@e)/(np.linalg.norm(e)+1e-12))
82 unconstrained_leaks.append(np.linalg.norm(A@e)/(np.linalg.norm(e)+1e-12))
83 exact_leak=float(np.mean(compatible_leaks)); uncon_leak=float(np.mean(unconstrained_leaks))
84 # Compatible vertex displacement fields versus random edge perturbations.
85 xs=[]; ys=[]
86 for _ in range(160):
87 u=rng.normal(size=len(pts)); compatible=D0@u
88 incompatible=compatible + (0.35 if _%2 else 0.0)*rng.normal(size=len(edges))
89 x=torch.tensor(incompatible,dtype=torch.float32)
90 xs.append(x); ys.append(float(_%2))
91 def train(model):
92 opt=torch.optim.Adam(model.parameters(),lr=0.02)
93 for _ in range(250):
94 loss=0.
95 for x,y in zip(xs,ys):
96 pred=model(x); loss=loss+(pred-y)**2
97 opt.zero_grad(); loss.backward(); opt.step()
98 with torch.no_grad():
99 pred=np.array([float(model(x)) for x in xs])
100 return float(np.mean((pred-np.array(ys))**2)), float(np.mean((pred[:80]<.5)==(np.array(ys[:80])<.5)))
101 exact=train(ExactComplex(D0,D1,D2)); uncon=train(Unconstrained(len(edges),len(faces)))
102 print({'vertices':len(pts),'edges':len(edges),'faces':len(faces),'cells':len(tet),
103 'relative_D1D0':r10,'relative_D2D1':r21,
104 'compatible_D1D0_leak':exact_leak,
105 'unconstrained_leak':uncon_leak,
106 'exact_mse_accuracy':exact,'unconstrained_mse_accuracy':uncon})
107
108if __name__ == '__main__': run()