Bounded Commuting Cochain Layer / cochain_mvp.py
Mechanism confirmed, baseline not beaten
1import json, random
2from pathlib import Path
3import numpy as np
4import torch
5
6SEED = 17
7np.random.seed(SEED)
8random.seed(SEED)
9torch.manual_seed(SEED)
10
11
12def weighted_qr(A, masses):
13 """Return Q with Q.T M Q=I, preserving the span of A."""
14 sqrt_m = np.sqrt(masses)
15 q, _ = np.linalg.qr(sqrt_m[:, None] * A)
16 return q / sqrt_m[:, None]
17
18
19def projector(Q, masses):
20 return Q @ Q.T @ np.diag(masses)
21
22
23def mass_opnorm(A, masses_in, masses_out):
24 """Induced M_in -> M_out operator norm."""
25 scaled = np.diag(np.sqrt(masses_out)) @ A @ np.diag(1.0 / np.sqrt(masses_in))
26 return float(np.linalg.svd(scaled, compute_uv=False)[0])
27
28
29def build_complex():
30 # Oriented filled triangle. d0 is edge-minus-vertex incidence and d1 is
31 # the oriented face boundary, so d1 d0=0.
32 d0 = np.array([[-1, 1, 0], [-1, 0, 1], [0, -1, 1]], dtype=float)
33 d1 = np.array([[1, -1, 1]], dtype=float)
34 return d0, d1
35
36
37def math_check():
38 d0, d1 = build_complex()
39 m0 = np.array([1.0, 2.0, 1.0])
40 m1 = np.array([1.0, 1.5, 2.0])
41 # P0=I and P1 projects onto exact 1-cochains im(d0). This removes the
42 # one-dimensional cycle component while preserving d(P0 z)=P1(dz).
43 Q1 = weighted_qr(d0[:, :2], m1)
44 P0 = np.eye(3)
45 P1 = projector(Q1, m1)
46 # Same-rank independent control (P0=I, random-looking edge subspace).
47 R1 = weighted_qr(np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]), m1)
48 C1 = projector(R1, m1)
49
50 rows = []
51 for name, p1 in [("compatible", P1), ("independent", C1)]:
52 idem0 = mass_opnorm(P0 @ P0 - P0, m0, m0)
53 idem1 = mass_opnorm(p1 @ p1 - p1, m1, m1)
54 defect01 = mass_opnorm(d0 @ P0 - p1 @ d0, m0, m1)
55 defect12 = mass_opnorm(d1 @ p1, m1, np.ones(1))
56 bound0 = mass_opnorm(P0, m0, m0)
57 bound1 = mass_opnorm(p1, m1, m1)
58 rows.append(dict(name=name, idempotence_k0=idem0, idempotence_k1=idem1,
59 commutation_defect_d0=defect01,
60 commutation_defect_d1=defect12,
61 bound_k0=bound0, bound_k1=bound1))
62 assert np.max(np.abs(d1 @ d0)) < 1e-12
63 return d0, m0, m1, P0, P1, rows
64
65
66def train_experiment(P0, P1, d0, steps=250, seed=SEED):
67 """Tiny fixed-seed regression: noisy node/edge features -> exact cochains."""
68 torch.manual_seed(seed)
69 n, channels = 96, 3
70 clean0 = torch.randn(n, 3, channels)
71 clean1 = torch.einsum("ei,nic->nec", torch.tensor(d0, dtype=torch.float32), clean0)
72 noisy0 = clean0 + 0.55 * torch.randn_like(clean0)
73 noisy1 = clean1 + 0.55 * torch.randn_like(clean1)
74 w0 = torch.nn.Parameter(torch.randn(channels, channels) * 0.2)
75 w1 = torch.nn.Parameter(torch.randn(channels, channels) * 0.2)
76 opt = torch.optim.Adam([w0, w1], lr=0.04)
77 t0 = torch.tensor(P0, dtype=torch.float32)
78 t1 = torch.tensor(P1, dtype=torch.float32)
79 for _ in range(steps):
80 y0, y1 = noisy0 @ w0, noisy1 @ w1
81 z0 = torch.einsum("ij,njc->nic", t0, y0)
82 z1 = torch.einsum("ij,njc->nic", t1, y1)
83 loss = ((z0 - clean0) ** 2).mean() + ((z1 - clean1) ** 2).mean()
84 opt.zero_grad(); loss.backward(); opt.step()
85 return float(((z0 - clean0) ** 2).mean() + ((z1 - clean1) ** 2).mean())
86
87
88def main():
89 d0, m0, m1, P0, P1, checks = math_check()
90 seeds = [17, 23, 41, 59, 71]
91 baseline_vals = [train_experiment(np.eye(3), np.eye(3), d0, seed=s) for s in seeds]
92 idea_vals = [train_experiment(P0, P1, d0, seed=s) for s in seeds]
93 baseline, idea = baseline_vals[0], idea_vals[0]
94 result = {
95 "seed": SEED,
96 "repeat_seeds": seeds,
97 "incidence_d1_d0_max": 0.0,
98 "checks": checks,
99 "denoising_mse": {"baseline_no_projection": baseline,
100 "bounded_commuting_layer": idea},
101 "relative_mse_reduction": (baseline - idea) / baseline,
102 "repeat_mse_mean_std": {
103 "baseline_no_projection": [float(np.mean(baseline_vals)), float(np.std(baseline_vals))],
104 "bounded_commuting_layer": [float(np.mean(idea_vals)), float(np.std(idea_vals))],
105 "relative_reduction_mean": float(np.mean((np.array(baseline_vals)-np.array(idea_vals))/np.array(baseline_vals)))}
106 }
107 Path("results.json").write_text(json.dumps(result, indent=2))
108 print(json.dumps(result, indent=2))
109
110
111if __name__ == "__main__":
112 main()