Phase-Polytope Robust Neural Dynamics / phase_polytope_experiment.py

Failed on benchmark

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