Exact Elasticity-Complex Message Passing / tetra_elasticity_complex.py
Beats tuned baseline
1import numpy as np
2
3META = {
4 'name': 'tetra_elasticity_complex',
5 'domain': 'pde',
6 'description': 'Simplicial tetrahedral edge-field regression with vertex-edge-face-cell incidence operators.'
7}
8
9TETS = ((0, 1, 2, 3), (0, 2, 1, 4))
10
11
12def operators():
13 edges = sorted({tuple(sorted((t[i], t[j]))) for t in TETS
14 for i in range(4) for j in range(i + 1, 4)})
15 faces = sorted({tuple(sorted(t[j] for j in range(4) if j != omit))
16 for t in TETS for omit in range(4)})
17 ei, fi = {e: i for i, e in enumerate(edges)}, {f: i for i, f in enumerate(faces)}
18 D0 = np.zeros((len(edges), 5), dtype=np.float32)
19 for r, (a, b) in enumerate(edges):
20 D0[r, a], D0[r, b] = -1, 1
21 D1 = np.zeros((len(faces), len(edges)), dtype=np.float32)
22 for r, (a, b, c) in enumerate(faces):
23 for e, s in (((b, c), 1), ((a, c), -1), ((a, b), 1)):
24 D1[r, ei[e]] = s
25 D2 = np.zeros((len(TETS), len(faces)), dtype=np.float32)
26 for r, t in enumerate(TETS):
27 for omit in range(4):
28 sub = tuple(t[j] for j in range(4) if j != omit)
29 f = tuple(sorted(sub))
30 inv = sum(sub[i] > sub[j] for i in range(3) for j in range(i + 1, 3))
31 D2[r, fi[f]] = (-1) ** omit * (-1 if inv % 2 else 1)
32 return D0, D1, D2, edges, faces
33
34D0, D1, D2, EDGES, FACES = operators()
35
36
37def get_dataset(seed, n_train, n_test):
38 rng = np.random.default_rng(int(seed))
39 # Each sample is an edge 1-cochain: compatible gradient + optional face defect.
40 # Target is the observed incompatibility magnitude, an ordinary regression MSE.
41 def make(n):
42 u = rng.normal(size=(n, 5)).astype(np.float32)
43 x = u @ D0.T
44 defect = rng.normal(size=(n, len(EDGES))).astype(np.float32)
45 active = (rng.random(n) < 0.65).astype(np.float32)[:, None]
46 x = x + active * 0.30 * defect
47 y = np.sqrt(np.mean((x @ D1.T) ** 2, axis=1, keepdims=True)).astype(np.float32)
48 return x.astype(np.float32), y
49 xtr, ytr = make(n_train)
50 xte, yte = make(n_test)
51 return {'xtr': xtr, 'ytr': ytr, 'xte': xte, 'yte': yte,
52 'task': 'regression', 'metric': 'mse', 'input_shape': (len(EDGES),),
53 'out_dim': 1, 'D0': D0, 'D1': D1, 'D2': D2}