Differentiable Physics-Equilibrium Projection / proper_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1"""
  2Проверка идеи "differentiable equilibrium projection" в ПРАВИЛЬНОЙ реализации.
  3
  4Задача (синтетическая, но нетривиальная):
  5  Истинное состояние подчиняется F(z; u) = z^3 + a*z - u = 0, где:
  6    - u дано как наблюдаемый вход (шумный),
  7    - a (истинная "жёсткость" / параметр системы) НЕ дан напрямую сети,
  8      а должен быть предсказан по косвенным признакам x (нелинейная скрытая
  9      зависимость + шум).
 10  Baseline: MLP предсказывает z напрямую по (x, u) - обычная регрессия.
 11  Idea:     MLP предсказывает a_hat по x, затем z получается через
 12            дифференцируемый Ньютон-солвер (ваш ImplicitCubic) для
 13            F(z; u) = z^3 + a_hat*z - u = 0.
 14            Здесь a_hat - это выход сети, и он ДЕЙСТВИТЕЛЬНО входит в уравнение
 15            (не сокращается алгебраически, как в багнутой версии), так что
 16            градиент течёт: dz/da_hat = -z/J по формуле implicit function theorem.
 17
 18Это честный тест: сеть предсказывает физически осмысленный параметр,
 19равновесие зависит от него нелинейно, и обучение идёт через implicit backward.
 20"""
 21import torch
 22import torch.nn as nn
 23import numpy as np
 24import copy
 25
 26torch.set_num_threads(4)
 27
 28# ---------- ваш код без изменений ----------
 29def newton_project(u, a, z0, steps=12, damping=1.0):
 30    z = z0
 31    for _ in range(steps):
 32        f = z**3 + a*z - u
 33        j = 3*z**2 + a
 34        z = z - damping*f/j
 35    return z
 36
 37class ImplicitCubic(torch.autograd.Function):
 38    @staticmethod
 39    def forward(ctx, u, a, z0, steps=20):
 40        with torch.no_grad():
 41            z = newton_project(u, a, z0, steps=steps)
 42        ctx.save_for_backward(z, a)
 43        return z
 44    @staticmethod
 45    def backward(ctx, grad_out):
 46        z, a = ctx.saved_tensors
 47        j = 3*z**2 + a
 48        return grad_out/j, (-grad_out*z/j), None, None
 49
 50def implicit_project(u, a, z0=None, steps=20):
 51    if z0 is None:
 52        z0 = torch.zeros_like(u)
 53    return ImplicitCubic.apply(u, a, z0, steps)
 54# --------------------------------------------
 55
 56def make_dataset(seed, n_train=400, n_test=160):
 57    rng = np.random.default_rng(seed)
 58    n = n_train + n_test
 59    # скрытый истинный параметр a в (0.2, 2.0), зависит от x нелинейно
 60    x_raw = rng.uniform(-2, 2, size=(n, 4)).astype(np.float32)
 61    true_a = (0.2 + 1.8 * (1 / (1 + np.exp(-(x_raw[:, 0]*0.8 + x_raw[:, 1]**2*0.3 - x_raw[:,2]*0.4)))))
 62    true_a = true_a.astype(np.float32)
 63    u = rng.uniform(-1.5, 1.5, size=n).astype(np.float32)
 64
 65    def solve_newton_np(u, a, z0=0.0, maxit=60, tol=1e-12):
 66        z = z0
 67        for _ in range(maxit):
 68            f = z**3 + a*z - u
 69            j = 3*z*z + a
 70            if abs(f) < tol:
 71                break
 72            z -= f / j
 73        return z
 74
 75    z_true = np.array([solve_newton_np(float(uu), float(aa)) for uu, aa in zip(u, true_a)], dtype=np.float32)
 76
 77    # добавим шум наблюдения к x (не к u), чтобы задача предсказания a была нетривиальной
 78    x_noisy = x_raw + rng.normal(0, 0.15, size=x_raw.shape).astype(np.float32)
 79
 80    X = torch.tensor(np.concatenate([x_noisy, u[:, None]], axis=1))
 81    U = torch.tensor(u[:, None])
 82    Z = torch.tensor(z_true[:, None])
 83    A_true = torch.tensor(true_a[:, None])
 84
 85    return {
 86        'Xtr': X[:n_train], 'Utr': U[:n_train], 'Ztr': Z[:n_train], 'Atr': A_true[:n_train],
 87        'Xte': X[n_train:], 'Ute': U[n_train:], 'Zte': Z[n_train:], 'Ate': A_true[n_train:],
 88    }
 89
 90class DirectMLP(nn.Module):
 91    """Baseline: предсказывает z напрямую по (x, u)."""
 92    def __init__(self):
 93        super().__init__()
 94        self.net = nn.Sequential(nn.Linear(5, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1))
 95    def forward(self, x, u):
 96        return self.net(torch.cat([x, u], dim=1))
 97
 98class AParamNet(nn.Module):
 99    """Idea: предсказывает a_hat по x (без u!), проекция решает физику."""
100    def __init__(self):
101        super().__init__()
102        self.net = nn.Sequential(nn.Linear(4, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1))
103        self.softplus = nn.Softplus()
104    def forward(self, x, u):
105        a_hat = self.softplus(self.net(x)) + 0.05  # a > 0 для устойчивости Ньютона
106        z0 = torch.zeros_like(u)
107        z = implicit_project(u.squeeze(1), a_hat.squeeze(1), z0.squeeze(1), steps=20).unsqueeze(1)
108        return z, a_hat
109
110def train_direct(d, epochs=300, lr=0.01, seed=0):
111    torch.manual_seed(seed)
112    m = DirectMLP()
113    opt = torch.optim.Adam(m.parameters(), lr=lr)
114    x = d['Xtr'][:, :4]; u = d['Utr']; z = d['Ztr']
115    for _ in range(epochs):
116        opt.zero_grad()
117        pred = m(x, u)
118        loss = ((pred - z)**2).mean()
119        loss.backward(); opt.step()
120    with torch.no_grad():
121        pred_te = m(d['Xte'][:, :4], d['Ute'])
122        mse = ((pred_te - d['Zte'])**2).mean().item()
123    return m, mse
124
125def train_idea(d, epochs=300, lr=0.01, seed=0):
126    torch.manual_seed(seed)
127    m = AParamNet()
128    opt = torch.optim.Adam(m.parameters(), lr=lr)
129    x = d['Xtr'][:, :4]; u = d['Utr']; z = d['Ztr']
130    grad_norms = []
131    for ep in range(epochs):
132        opt.zero_grad()
133        pred, a_hat = m(x, u)
134        loss = ((pred - z)**2).mean()
135        loss.backward()
136        gnorm = sum(p.grad.norm().item()**2 for p in m.net.parameters() if p.grad is not None)**0.5
137        grad_norms.append(gnorm)
138        opt.step()
139    with torch.no_grad():
140        pred_te, a_hat_te = m(d['Xte'][:, :4], d['Ute'])
141        mse = ((pred_te - d['Zte'])**2).mean().item()
142        a_mse = ((a_hat_te - d['Ate'])**2).mean().item()
143    return m, mse, grad_norms, a_mse
144
145def gradient_flow_and_ablation_check(m, d):
146    """Механическая проверка: течёт ли градиент через сеть, и меняется ли выход
147    при порче весов сети (та самая проверка из нашего обсуждения)."""
148    x = d['Xte'][:, :4]; u = d['Ute']
149    with torch.no_grad():
150        out_A, a_A = m(x, u)
151        m_broken = copy.deepcopy(m)
152        for p in m_broken.net.parameters():
153            scale = p.data.std(unbiased=False).item() if p.data.numel() > 1 else p.data.abs().item()
154            scale = scale if scale > 1e-8 else 1.0
155            p.data = torch.randn_like(p.data) * scale
156        out_B, a_B = m_broken(x, u)
157    delta = (out_A - out_B).abs().mean().item()
158    out_std = out_A.std().item()
159    return {'output_delta_after_breaking_network': delta,
160            'output_std': out_std,
161            'relative_delta': delta / (out_std + 1e-12)}
162
163def main():
164    seeds = range(8)
165    direct_mses, idea_mses, a_mses = [], [], []
166    grad_flow_report = None
167    for s in seeds:
168        d = make_dataset(seed=s)
169        _, mse_d = train_direct(d, seed=s)
170        m_idea, mse_i, gnorms, a_mse = train_idea(d, seed=s)
171        direct_mses.append(mse_d)
172        idea_mses.append(mse_i)
173        a_mses.append(a_mse)
174        if s == 0:
175            grad_flow_report = gradient_flow_and_ablation_check(m_idea, d)
176            first_grad_norms = gnorms
177
178    direct_mses = np.array(direct_mses)
179    idea_mses = np.array(idea_mses)
180
181    print("=== RESULTS (8 seeds) ===")
182    print(f"Baseline (direct MLP, sees u+x)   mean test MSE: {direct_mses.mean():.6e}  std: {direct_mses.std():.6e}")
183    print(f"Idea (predict a, Newton-project)  mean test MSE: {idea_mses.mean():.6e}  std: {idea_mses.std():.6e}")
184    wins = int((idea_mses < direct_mses).sum())
185    print(f"Idea wins on {wins}/8 seeds")
186    improvement = (direct_mses.mean() - idea_mses.mean()) / direct_mses.mean() * 100
187    print(f"Improvement: {improvement:.1f}%")
188    print(f"Mean a_hat MSE vs true hidden parameter a: {np.mean(a_mses):.6e}  (network is learning something real)")
189    print()
190    print("=== GRADIENT-FLOW / ABLATION CHECK (seed 0) ===")
191    print(f"First 5 grad norms during training: {[round(g,5) for g in first_grad_norms[:5]]}")
192    print(f"Last 5 grad norms during training:  {[round(g,5) for g in first_grad_norms[-5:]]}")
193    print(f"Ablation (network weights randomized) -> output delta: {grad_flow_report}")
194    print()
195    if grad_flow_report['relative_delta'] > 0.05 and first_grad_norms[0] > 1e-8:
196        print("VERDICT: gradient DOES flow through the network, network output DOES affect final result.")
197        print("This is a legitimate (non-degenerate) implementation of the idea.")
198    else:
199        print("VERDICT: still degenerate (should not happen with this code).")
200
201if __name__ == '__main__':
202    main()