Progressive rollout-consistency training / experiment.py
Mechanism failed
1import json
2import math
3import random
4import numpy as np
5import torch
6from torch import nn
7
8SEED = 7
9random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
10torch.set_num_threads(4)
11device = torch.device('cpu')
12DT = 0.08
13DIM = 2
14
15
16def rk4(fun, x, dt=DT):
17 k1 = dt * fun(x)
18 k2 = dt * fun(x + 0.5*k1)
19 k3 = dt * fun(x + 0.5*k2)
20 k4 = dt * fun(x + k3)
21 return x + (k1 + 2*k2 + 2*k3 + k4) / 6.0
22
23
24def true_field(x):
25 # Stable nonlinear oscillator with state-dependent frequency and damping.
26 q, p = x[..., 0], x[..., 1]
27 return torch.stack((p, -0.8*q - 0.15*p - 0.18*q**3), dim=-1)
28
29
30def make_data(ntraj=96, length=72, noise=0.002):
31 g = torch.Generator().manual_seed(SEED)
32 x = torch.zeros(ntraj, length+1, DIM)
33 x[:, 0] = torch.empty(ntraj, DIM).uniform_(-1.3, 1.3, generator=g)
34 with torch.no_grad():
35 for t in range(length):
36 x[:, t+1] = rk4(true_field, x[:, t])
37 # Observational noise is applied to training targets, while evaluation uses clean trajectories.
38 noisy = x + noise * torch.randn(x.shape, generator=g)
39 return x, noisy
40
41
42class Field(nn.Module):
43 def __init__(self):
44 super().__init__()
45 self.net = nn.Sequential(nn.Linear(2, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 2))
46 def forward(self, x):
47 return self.net(x)
48
49
50def rollout(model, x0, horizon):
51 xs = [x0]
52 x = x0
53 for _ in range(horizon):
54 x = rk4(model, x)
55 xs.append(x)
56 return torch.stack(xs, dim=1)
57
58
59def l1_penalty(model):
60 return sum(p.abs().sum() for p in model.parameters())
61
62
63def train_one_step(train_noisy, epochs=105, batch=24):
64 model = Field().to(device)
65 opt = torch.optim.Adam(model.parameters(), lr=3e-3)
66 n = train_noisy.shape[0]
67 model.train()
68 for ep in range(epochs):
69 order = torch.randperm(n)
70 for start in range(0, n, batch):
71 ids = order[start:start+batch]
72 # Random teacher-forced state and its next noisy observation.
73 t = torch.randint(0, train_noisy.shape[1]-1, (len(ids),))
74 x0 = train_noisy[ids, t]
75 target = train_noisy[ids, t+1]
76 pred = rk4(model, x0)
77 loss = (pred-target).abs().mean() + 1e-6*l1_penalty(model)
78 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); opt.step()
79 return model
80
81
82def train_progressive(train_noisy, schedule=(1,2,4,8), epochs_per_phase=27, batch=24):
83 model = Field().to(device)
84 opt = torch.optim.Adam(model.parameters(), lr=3e-3)
85 n = train_noisy.shape[0]
86 model.train()
87 for H in schedule:
88 for ep in range(epochs_per_phase):
89 order = torch.randperm(n)
90 for start in range(0, n, batch):
91 ids = order[start:start+batch]
92 max_t = train_noisy.shape[1] - 1 - H
93 t = torch.randint(0, max_t+1, (len(ids),))
94 target = train_noisy[ids[:, None], t[:, None] + torch.arange(1,H+1)[None, :]]
95 x = train_noisy[ids, t]
96 pred = rollout(model, x, H)[:, 1:]
97 loss = (pred-target).abs().mean() + 1e-6*l1_penalty(model)
98 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); opt.step()
99 return model
100
101
102def evaluate(model, clean, horizons=(8,16,32,64)):
103 model.eval(); out={}
104 with torch.no_grad():
105 for H in horizons:
106 pred = rollout(model, clean[:,0], H)[:,1:]
107 target = clean[:,1:H+1]
108 err = (pred-target).abs().mean().item()
109 maxnorm = pred.norm(dim=-1).max().item()
110 divergent = (pred.norm(dim=-1).max(dim=1).values > 4.0).float().mean().item()
111 out[str(H)] = {'mae':err, 'max_norm':maxnorm, 'divergence_rate':divergent}
112 return out
113
114
115def rk4_sanity():
116 # For x'=A x, RK4 error should decrease by ~16 when dt is halved.
117 A = torch.tensor([[-0.2, 1.0],[-1.4,-0.3]], dtype=torch.float64)
118 x0 = torch.tensor([[1.1,-0.4]], dtype=torch.float64)
119 def f(x): return x @ A.T
120 # Same physical time, reference from a very fine RK4 integration.
121 with torch.no_grad():
122 ref=x0.clone()
123 for _ in range(40000): ref=rk4(f, ref, 0.00002)
124 errs=[]
125 for dt, steps in [(0.16,5),(0.08,10),(0.04,20)]:
126 y=x0.clone()
127 for _ in range(steps): y=rk4(f,y,dt)
128 errs.append(float((y-ref).abs().max()))
129 ratios=[errs[i]/errs[i+1] for i in range(2)]
130 return {'errors_dt_0.16_0.08_0.04':errs, 'halving_error_ratios':ratios, 'passes_fourth_order_signal': all(r>8 for r in ratios)}
131
132
133def main():
134 sanity = rk4_sanity()
135 clean, noisy = make_data()
136 # Equal optimizer epochs and same architecture; progressive phases expose the model to all horizons.
137 baseline = train_one_step(noisy, epochs=225)
138 progressive = train_progressive(noisy, schedule=(1,2,4,8), epochs_per_phase=15)
139 result = {
140 'seed': SEED, 'dt': DT, 'train_trajectories': len(clean),
141 'rk4_sanity': sanity,
142 'baseline': evaluate(baseline, clean),
143 'progressive': evaluate(progressive, clean),
144 }
145 print(json.dumps(result, indent=2))
146
147if __name__ == '__main__': main()