import json, os, sys import numpy as np import torch import torch.nn as nn sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import train_model, evaluate, sweep_baseline, make_report from bench.protocol import DEFAULT_SEEDS from graph_track import get_dataset, META DEVICE = "cuda" if torch.cuda.is_available() else "cpu" class GeometricMessageNet(nn.Module): """Matched residual GNN; only the edge weighting changes.""" def __init__(self, weighted=True, alpha=.5, hidden=32, depth=4): super().__init__() self.weighted, self.alpha = weighted, alpha self.layers = nn.ModuleList([nn.Linear(5, hidden)] + [nn.Linear(hidden, hidden) for _ in range(depth - 1)]) self.head = nn.Sequential(nn.Linear(hidden, hidden), nn.ReLU(), nn.Linear(hidden, 1)) def propagation(self, x): xy = x[:, :, 2:4] d2 = ((xy[:, :, None] - xy[:, None, :]) ** 2).sum(-1) ids = d2.topk(6, dim=-1, largest=False).indices[:, :, 1:6] knn = torch.zeros_like(d2) knn.scatter_(2, ids, 1.0) knn = torch.maximum(knn, knn.transpose(1, 2)) w = knn * (d2.clamp_min(1e-6) if self.weighted else 1.0) return w / w.sum(-1, keepdim=True).clamp_min(1e-8) def forward(self, x): p = self.propagation(x) h = x for layer in self.layers: h = torch.relu(layer((1.0 - self.alpha) * h + self.alpha * torch.bmm(p, h))) root = x[:, :, 4:5] return self.head((h * root).sum(1) / root.sum(1).clamp_min(1.0)) def tensor_ds(seed): raw = get_dataset(seed, n_train=400, n_test=160) return {**raw, "xtr": torch.as_tensor(raw["xtr"], dtype=torch.float32), "ytr": torch.as_tensor(raw["ytr"], dtype=torch.float32).reshape(-1, 1), "xte": torch.as_tensor(raw["xte"], dtype=torch.float32), "yte": torch.as_tensor(raw["yte"], dtype=torch.float32).reshape(-1, 1)} def train_one(seed, weighted, lr, alpha=.5, epochs=18, return_model=False): np.random.seed(seed); torch.manual_seed(seed) ds = tensor_ds(seed) net = GeometricMessageNet(weighted=weighted, alpha=alpha) trained, metric, hist = train_model(net, ds, epochs=epochs, lr=lr, batch=128, weight_decay=0.0, log=lambda *_: None) if trained is None: return (float("nan"), None, ds) if return_model else float("nan") return (float(metric), trained, ds) if return_model else float(metric) def make_train_fn(cfg, weighted): return lambda seed: train_one(seed, weighted, cfg["lr"], cfg["alpha"]) def influence_signature(seed=0): result = {} for name, weighted in (("baseline", False), ("idea", True)): metric, net, ds = train_one(seed, weighted, .003, return_model=True) if net is None: result[name] = {"test_mse": float("nan"), "error": "training failed"} continue dev = next(net.parameters()).device x = ds["xte"][:1].to(dev).clone().requires_grad_(True) net.zero_grad(); net(x).sum().backward() observed = x.grad[0, :, 0].abs().detach().cpu().numpy(); observed /= observed.sum() + 1e-12 with torch.no_grad(): p = net.propagation(x.detach())[0].detach().cpu().numpy() root = int(np.argmax(ds["xte"][0, :, 4].numpy())) predicted = np.linalg.matrix_power(p, 4)[root] predicted /= predicted.sum() + 1e-12 result[name] = {"test_mse": metric, "predicted_l2": float(np.linalg.norm(predicted)), "observed_gradient_l2": float(np.linalg.norm(observed)), "predicted_observed_corr": float(np.corrcoef(observed, predicted)[0, 1])} b, i = result["baseline"], result["idea"] result["prediction"] = "distance weighting preserves more long-range influence" result["confirmed"] = bool(i["observed_gradient_l2"] > b["observed_gradient_l2"] + .01 and i["predicted_observed_corr"] > .5) return result def main(): grid = [{"lr": lr, "alpha": .5} for lr in (.001, .003, .01)] base = sweep_baseline(lambda cfg: make_train_fn(cfg, False), grid) trials = [{"cfg": cfg, "result": evaluate(make_train_fn(cfg, True), DEFAULT_SEEDS)} for cfg in grid] best = min(trials, key=lambda z: z["result"]["mean"]) rep = make_report("geometric_graph_diffusion", "GeometricMessageNet", base, best["result"], {"track_structure": "geometric graph node-field diffusion regression", "trained_model_measurements": influence_signature(), "idea_sweep": trials, "confirmed": False}) rep["custom_track"] = {"name": META["name"], "file": "graph_track.py", "domain": META["domain"]} rep["protocol_notes"] = {"paired_seeds": list(DEFAULT_SEEDS), "baseline_and_idea_same_architecture": True, "only_intervention": "edge weight d^2 versus unit edge weight", "device": DEVICE} os.makedirs("artifacts", exist_ok=True) with open("artifacts/bench_report.json", "w") as f: json.dump(rep, f, indent=2) print(json.dumps(rep, indent=2)) if __name__ == "__main__": main()