""" Проверка идеи "differentiable equilibrium projection" в ПРАВИЛЬНОЙ реализации. Задача (синтетическая, но нетривиальная): Истинное состояние подчиняется F(z; u) = z^3 + a*z - u = 0, где: - u дано как наблюдаемый вход (шумный), - a (истинная "жёсткость" / параметр системы) НЕ дан напрямую сети, а должен быть предсказан по косвенным признакам x (нелинейная скрытая зависимость + шум). Baseline: MLP предсказывает z напрямую по (x, u) - обычная регрессия. Idea: MLP предсказывает a_hat по x, затем z получается через дифференцируемый Ньютон-солвер (ваш ImplicitCubic) для F(z; u) = z^3 + a_hat*z - u = 0. Здесь a_hat - это выход сети, и он ДЕЙСТВИТЕЛЬНО входит в уравнение (не сокращается алгебраически, как в багнутой версии), так что градиент течёт: dz/da_hat = -z/J по формуле implicit function theorem. Это честный тест: сеть предсказывает физически осмысленный параметр, равновесие зависит от него нелинейно, и обучение идёт через implicit backward. """ import torch import torch.nn as nn import numpy as np import copy torch.set_num_threads(4) # ---------- ваш код без изменений ---------- def newton_project(u, a, z0, steps=12, damping=1.0): z = z0 for _ in range(steps): f = z**3 + a*z - u j = 3*z**2 + a z = z - damping*f/j return z class ImplicitCubic(torch.autograd.Function): @staticmethod def forward(ctx, u, a, z0, steps=20): with torch.no_grad(): z = newton_project(u, a, z0, steps=steps) ctx.save_for_backward(z, a) return z @staticmethod def backward(ctx, grad_out): z, a = ctx.saved_tensors j = 3*z**2 + a return grad_out/j, (-grad_out*z/j), None, None def implicit_project(u, a, z0=None, steps=20): if z0 is None: z0 = torch.zeros_like(u) return ImplicitCubic.apply(u, a, z0, steps) # -------------------------------------------- def make_dataset(seed, n_train=400, n_test=160): rng = np.random.default_rng(seed) n = n_train + n_test # скрытый истинный параметр a в (0.2, 2.0), зависит от x нелинейно x_raw = rng.uniform(-2, 2, size=(n, 4)).astype(np.float32) 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))))) true_a = true_a.astype(np.float32) u = rng.uniform(-1.5, 1.5, size=n).astype(np.float32) def solve_newton_np(u, a, z0=0.0, maxit=60, tol=1e-12): z = z0 for _ in range(maxit): f = z**3 + a*z - u j = 3*z*z + a if abs(f) < tol: break z -= f / j return z z_true = np.array([solve_newton_np(float(uu), float(aa)) for uu, aa in zip(u, true_a)], dtype=np.float32) # добавим шум наблюдения к x (не к u), чтобы задача предсказания a была нетривиальной x_noisy = x_raw + rng.normal(0, 0.15, size=x_raw.shape).astype(np.float32) X = torch.tensor(np.concatenate([x_noisy, u[:, None]], axis=1)) U = torch.tensor(u[:, None]) Z = torch.tensor(z_true[:, None]) A_true = torch.tensor(true_a[:, None]) return { 'Xtr': X[:n_train], 'Utr': U[:n_train], 'Ztr': Z[:n_train], 'Atr': A_true[:n_train], 'Xte': X[n_train:], 'Ute': U[n_train:], 'Zte': Z[n_train:], 'Ate': A_true[n_train:], } class DirectMLP(nn.Module): """Baseline: предсказывает z напрямую по (x, u).""" def __init__(self): super().__init__() self.net = nn.Sequential(nn.Linear(5, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1)) def forward(self, x, u): return self.net(torch.cat([x, u], dim=1)) class AParamNet(nn.Module): """Idea: предсказывает a_hat по x (без u!), проекция решает физику.""" def __init__(self): super().__init__() self.net = nn.Sequential(nn.Linear(4, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1)) self.softplus = nn.Softplus() def forward(self, x, u): a_hat = self.softplus(self.net(x)) + 0.05 # a > 0 для устойчивости Ньютона z0 = torch.zeros_like(u) z = implicit_project(u.squeeze(1), a_hat.squeeze(1), z0.squeeze(1), steps=20).unsqueeze(1) return z, a_hat def train_direct(d, epochs=300, lr=0.01, seed=0): torch.manual_seed(seed) m = DirectMLP() opt = torch.optim.Adam(m.parameters(), lr=lr) x = d['Xtr'][:, :4]; u = d['Utr']; z = d['Ztr'] for _ in range(epochs): opt.zero_grad() pred = m(x, u) loss = ((pred - z)**2).mean() loss.backward(); opt.step() with torch.no_grad(): pred_te = m(d['Xte'][:, :4], d['Ute']) mse = ((pred_te - d['Zte'])**2).mean().item() return m, mse def train_idea(d, epochs=300, lr=0.01, seed=0): torch.manual_seed(seed) m = AParamNet() opt = torch.optim.Adam(m.parameters(), lr=lr) x = d['Xtr'][:, :4]; u = d['Utr']; z = d['Ztr'] grad_norms = [] for ep in range(epochs): opt.zero_grad() pred, a_hat = m(x, u) loss = ((pred - z)**2).mean() loss.backward() gnorm = sum(p.grad.norm().item()**2 for p in m.net.parameters() if p.grad is not None)**0.5 grad_norms.append(gnorm) opt.step() with torch.no_grad(): pred_te, a_hat_te = m(d['Xte'][:, :4], d['Ute']) mse = ((pred_te - d['Zte'])**2).mean().item() a_mse = ((a_hat_te - d['Ate'])**2).mean().item() return m, mse, grad_norms, a_mse def gradient_flow_and_ablation_check(m, d): """Механическая проверка: течёт ли градиент через сеть, и меняется ли выход при порче весов сети (та самая проверка из нашего обсуждения).""" x = d['Xte'][:, :4]; u = d['Ute'] with torch.no_grad(): out_A, a_A = m(x, u) m_broken = copy.deepcopy(m) for p in m_broken.net.parameters(): scale = p.data.std(unbiased=False).item() if p.data.numel() > 1 else p.data.abs().item() scale = scale if scale > 1e-8 else 1.0 p.data = torch.randn_like(p.data) * scale out_B, a_B = m_broken(x, u) delta = (out_A - out_B).abs().mean().item() out_std = out_A.std().item() return {'output_delta_after_breaking_network': delta, 'output_std': out_std, 'relative_delta': delta / (out_std + 1e-12)} def main(): seeds = range(8) direct_mses, idea_mses, a_mses = [], [], [] grad_flow_report = None for s in seeds: d = make_dataset(seed=s) _, mse_d = train_direct(d, seed=s) m_idea, mse_i, gnorms, a_mse = train_idea(d, seed=s) direct_mses.append(mse_d) idea_mses.append(mse_i) a_mses.append(a_mse) if s == 0: grad_flow_report = gradient_flow_and_ablation_check(m_idea, d) first_grad_norms = gnorms direct_mses = np.array(direct_mses) idea_mses = np.array(idea_mses) print("=== RESULTS (8 seeds) ===") print(f"Baseline (direct MLP, sees u+x) mean test MSE: {direct_mses.mean():.6e} std: {direct_mses.std():.6e}") print(f"Idea (predict a, Newton-project) mean test MSE: {idea_mses.mean():.6e} std: {idea_mses.std():.6e}") wins = int((idea_mses < direct_mses).sum()) print(f"Idea wins on {wins}/8 seeds") improvement = (direct_mses.mean() - idea_mses.mean()) / direct_mses.mean() * 100 print(f"Improvement: {improvement:.1f}%") print(f"Mean a_hat MSE vs true hidden parameter a: {np.mean(a_mses):.6e} (network is learning something real)") print() print("=== GRADIENT-FLOW / ABLATION CHECK (seed 0) ===") print(f"First 5 grad norms during training: {[round(g,5) for g in first_grad_norms[:5]]}") print(f"Last 5 grad norms during training: {[round(g,5) for g in first_grad_norms[-5:]]}") print(f"Ablation (network weights randomized) -> output delta: {grad_flow_report}") print() if grad_flow_report['relative_delta'] > 0.05 and first_grad_norms[0] > 1e-8: print("VERDICT: gradient DOES flow through the network, network output DOES affect final result.") print("This is a legitimate (non-degenerate) implementation of the idea.") else: print("VERDICT: still degenerate (should not happen with this code).") if __name__ == '__main__': main()