Phase-Polytope Robust Neural Dynamics / phase_polytope_experiment.py
Failed on benchmark
1import json
2import random
3from pathlib import Path
4import numpy as np
5import torch
6from torch import nn
7
8SEED = 624
9M = 4
10DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
11
12
13def seed_all(seed=SEED):
14 random.seed(seed)
15 np.random.seed(seed)
16 torch.manual_seed(seed)
17 if torch.cuda.is_available():
18 torch.cuda.manual_seed_all(seed)
19
20
21def variance_check():
22 rng = np.random.default_rng(SEED)
23 sigma = 1.7
24 n = 300000
25 rows = []
26 for m in [1, 2, 4, 8, 16]:
27 errors = rng.normal(0.0, sigma, size=(n, m))
28 observed = float(np.var(errors.mean(axis=1)))
29 predicted = sigma * sigma / m
30 rows.append({"M": m, "observed": observed, "predicted": predicted,
31 "ratio": observed / predicted})
32 return rows
33
34
35def oscillator_step(x, dt=0.08):
36 q, v = x[:, 0], x[:, 1]
37 q_next = q + dt * v
38 v_next = v + dt * (-0.8 * q - 0.15 * v + 0.10 * q ** 3)
39 return torch.stack([q_next, v_next], dim=1)
40
41
42def make_data(n, seed):
43 generator = torch.Generator().manual_seed(seed)
44 x = torch.empty(n, 2).uniform_(-1.2, 1.2, generator=generator)
45 y = oscillator_step(x)
46 return x, y
47
48
49class PhaseNet(nn.Module):
50 def __init__(self, m=M):
51 super().__init__()
52 self.feature = nn.Sequential(
53 nn.Linear(2, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh()
54 )
55 self.heads = nn.ModuleList([nn.Linear(32, 2) for _ in range(m)])
56
57 def forward(self, x):
58 h = self.feature(x)
59 return torch.stack([head(h) for head in self.heads], dim=1)
60
61
62def train(mode, x_train, y_train, phase_offsets, steps=500):
63 model = PhaseNet(M).to(DEVICE)
64 optimizer = torch.optim.Adam(model.parameters(), lr=2e-3)
65 mse = nn.MSELoss()
66 x_train, y_train = x_train.to(DEVICE), y_train.to(DEVICE)
67 batch_size = 128
68 for step in range(steps):
69 idx = torch.randint(0, len(x_train), (batch_size,), device=DEVICE)
70 pred = model(x_train[idx])
71 target = y_train[idx]
72 # Simulate cyclic phase alignment by assigning each head a small,
73 # fixed offset in its phase-specific target.
74 targets = target[:, None, :].expand(-1, M, -1) + phase_offsets[None, :, :]
75 losses = ((pred - targets) ** 2).mean(dim=(0, 2))
76 centroid_loss = ((pred.mean(dim=1) - target) ** 2).mean()
77 if mode == "single":
78 loss = losses[0]
79 elif mode == "centroid":
80 loss = centroid_loss
81 elif mode == "robust":
82 worst = losses.max()
83 # Output-space convex mixture, choosing the worst vertex exactly.
84 loss = centroid_loss + 0.5 * worst
85 else:
86 raise ValueError(mode)
87 optimizer.zero_grad(set_to_none=True)
88 loss.backward()
89 optimizer.step()
90 return model
91
92
93def evaluate(model, x, y, phase_offsets):
94 model.eval()
95 with torch.no_grad():
96 pred = model(x.to(DEVICE))
97 target = y.to(DEVICE)
98 targets = target[:, None, :].expand(-1, M, -1) + phase_offsets[None, :, :]
99 vertex_mse = ((pred - targets) ** 2).mean(dim=(0, 2))
100 centroid_mse = ((pred.mean(dim=1) - target) ** 2).mean()
101 spread = pred.var(dim=1, unbiased=False).mean()
102 return {
103 "centroid_rmse": float(torch.sqrt(centroid_mse).cpu()),
104 "max_vertex_rmse": float(torch.sqrt(vertex_mse.max()).cpu()),
105 "mean_vertex_rmse": float(torch.sqrt(vertex_mse.mean()).cpu()),
106 "vertex_loss_spread": float(vertex_mse.max().cpu() - vertex_mse.min().cpu()),
107 "prediction_spread": float(spread.cpu()),
108 }
109
110
111def main():
112 seed_all()
113 x_train, y_train = make_data(4096, SEED + 1)
114 x_test, y_test = make_data(2048, SEED + 2)
115 # Small phase-dependent observation shifts make the heads genuinely distinct.
116 offsets = torch.tensor([
117 [-0.030, 0.020], [0.020, -0.025], [0.015, 0.030], [-0.025, -0.020]
118 ], dtype=torch.float32, device=DEVICE)
119 results = {"device": DEVICE, "variance_check": variance_check(), "models": {}}
120 for mode in ["single", "centroid", "robust"]:
121 seed_all(SEED + {"single": 1, "centroid": 2, "robust": 3}[mode])
122 model = train(mode, x_train, y_train, offsets)
123 results["models"][mode] = evaluate(model, x_test, y_test, offsets)
124 print(json.dumps(results, indent=2, sort_keys=True))
125 Path("phase_polytope_results.json").write_text(json.dumps(results, indent=2, sort_keys=True))
126
127
128if __name__ == "__main__":
129 try:
130 main()
131 except (RuntimeError, torch.cuda.CudaError) as exc:
132 if DEVICE == "cuda":
133 print("CUDA failed; rerun on CPU:", repr(exc))
134 torch.cuda.empty_cache()
135 torch.cuda.is_available = lambda: False
136 main()
137 else:
138 raise