Progressive rollout-consistency training / experiment.py

Mechanism failed

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