repro-formal-problem-solving / code /claim5_native_benchmark.py
ProCreations's picture
Repair claim 5 with native benchmark execution
35042ba
Raw
History Blame Contribute Delete
10.6 kB
#!/usr/bin/env python3
"""Bounded native rerun of the paper's two synthetic WIRE experiments.
The graph attention module is imported from the pinned Graph-RoPE commit
4ac067eb38272543b0cdd7591d630399ff37bce4. The two data generators and the
small training harness follow the paper's Section 4.1 / Appendix A.1
protocols: 5x5 edge-deleted grids for the monochromatic-subgraph task and
10-node Watts--Strogatz graphs for shortest-path prediction. The default
route keeps the paper's graph sizes, four-layer width-32 single-head
transformer, dropout, AdamW hyperparameters, and spectral-coordinate
conditions, while bounding examples, epochs, and seeds for CPU execution.
"""
from __future__ import annotations
import argparse
import json
import math
import os
import random
import sys
import time
import types
from dataclasses import dataclass
from pathlib import Path
import networkx as nx
import numpy as np
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
def import_pinned_graphrope(repo: Path):
"""Load the official module without importing GraphGPS optional plugins."""
root = str(repo)
graphgps = types.ModuleType("graphgps")
graphgps.__path__ = [str(repo / "graphgps")]
sys.modules["graphgps"] = graphgps
layer = types.ModuleType("graphgps.layer")
layer.__path__ = [str(repo / "graphgps" / "layer")]
sys.modules["graphgps.layer"] = layer
from graphgps.layer.graphrope import GraphRoPE # noqa: WPS433
return GraphRoPE
def seed_all(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def laplace_features(g: nx.Graph, max_freqs: int = 10) -> np.ndarray:
"""Return the first nonconstant unnormalised-Laplacian eigenvectors."""
a = nx.to_numpy_array(g, nodelist=range(g.number_of_nodes()), dtype=float)
lap = np.diag(a.sum(axis=1)) - a
_, vecs = np.linalg.eigh(lap)
pe = vecs[:, 1 : 1 + max_freqs]
if pe.shape[1] < max_freqs:
pe = np.pad(pe, ((0, 0), (0, max_freqs - pe.shape[1])))
return pe.astype(np.float32)
def largest_monochromatic_component(g: nx.Graph, colors: np.ndarray) -> int:
best = 0
for color in np.unique(colors):
nodes = np.flatnonzero(colors == color).tolist()
if nodes:
best = max(best, max((len(c) for c in nx.connected_components(g.subgraph(nodes))), default=0))
return best
def make_monochromatic(n: int, deleted_edges: int, seed: int, max_freqs: int = 10):
rng = np.random.default_rng(seed)
base = nx.grid_2d_graph(5, 5)
base = nx.convert_node_labels_to_integers(base, ordering="sorted")
edges = list(base.edges())
removed = rng.choice(len(edges), size=deleted_edges, replace=False)
g = base.copy()
g.remove_edges_from([edges[int(i)] for i in removed])
colors = rng.integers(0, 4, size=n, dtype=np.int64)
pe = laplace_features(g, max_freqs)
x = np.concatenate([pe, np.eye(4, dtype=np.float32)[colors]], axis=1)
y = np.float32(largest_monochromatic_component(g, colors) / n)
return x, pe, y
def make_shortest(seed: int, max_freqs: int = 10):
rng = np.random.default_rng(seed)
graph_seed = int(rng.integers(0, 2**31 - 1))
while True:
g = nx.watts_strogatz_graph(10, 2, 0.6, seed=graph_seed)
if nx.is_connected(g):
break
graph_seed += 1
source, target = rng.choice(10, size=2, replace=False).tolist()
pe = laplace_features(g, max_freqs)
marks = np.zeros((10, 2), dtype=np.float32)
marks[source, 0] = 1.0
marks[target, 1] = 1.0
x = np.concatenate([pe, marks], axis=1)
y = np.float32(nx.shortest_path_length(g, source, target) / 10.0)
return x, pe, y
@dataclass
class TaskData:
x: torch.Tensor
pe: torch.Tensor
y: torch.Tensor
def build_data(task: str, deleted_edges: int | None, count: int, seed: int) -> TaskData:
rows, pes, ys = [], [], []
for i in range(count):
if task == "monochromatic":
x, pe, y = make_monochromatic(25, int(deleted_edges), seed + i)
else:
x, pe, y = make_shortest(seed + i)
rows.append(x)
pes.append(pe)
ys.append(y)
return TaskData(
torch.from_numpy(np.stack(rows)),
torch.from_numpy(np.stack(pes)),
torch.tensor(ys, dtype=torch.float32),
)
class WIREBlock(nn.Module):
def __init__(self, GraphRoPE, hidden: int, m: int, dropout: float, omega_init: str):
super().__init__()
self.m = m
self.attn = GraphRoPE(
k=max(1, m),
d=hidden,
num_heads=1,
dropout=dropout,
enable=m > 0,
init_omega=omega_init,
attn_type="Full",
)
self.norm1 = nn.LayerNorm(hidden)
self.norm2 = nn.LayerNorm(hidden)
self.ff = nn.Sequential(
nn.Linear(hidden, hidden),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(hidden, hidden),
nn.Dropout(dropout),
)
def forward(self, x: torch.Tensor, pe: torch.Tensor) -> torch.Tensor:
batch_size, nodes, _ = x.shape
flat_batch = torch.arange(batch_size, device=x.device).repeat_interleave(nodes)
batch = types.SimpleNamespace(x=x.reshape(batch_size * nodes, -1), batch=flat_batch)
if self.m > 0:
batch.t = pe[:, :, : self.m].reshape(batch_size * nodes, self.m)
attn = self.attn(batch).reshape(batch_size, nodes, -1)
x = self.norm1(x + attn)
return self.norm2(x + self.ff(x))
class NativeWIRERegressor(nn.Module):
def __init__(self, GraphRoPE, input_dim: int, m: int, hidden: int = 32, layers: int = 4, omega_init: str = "zero"):
super().__init__()
self.input = nn.Linear(input_dim, hidden)
self.blocks = nn.ModuleList([WIREBlock(GraphRoPE, hidden, m, 0.2, omega_init) for _ in range(layers)])
self.head = nn.Linear(hidden, 1)
def forward(self, x: torch.Tensor, pe: torch.Tensor) -> torch.Tensor:
h = self.input(x)
for block in self.blocks:
h = block(h, pe)
return self.head(h.mean(dim=1)).squeeze(-1)
def train_one(GraphRoPE, task: str, setting: str, m: int, seed: int, args) -> dict:
seed_all(seed)
deleted = int(setting) if task == "monochromatic" else None
train = build_data(task, deleted, args.train_examples, 100000 + seed * 10000 + (deleted or 0) * 100)
test = build_data(task, deleted, args.test_examples, 200000 + seed * 10000 + (deleted or 0) * 100)
loader = DataLoader(TensorDataset(train.x, train.pe, train.y), batch_size=16, shuffle=True)
input_dim = train.x.shape[-1]
model = NativeWIRERegressor(GraphRoPE, input_dim=input_dim, m=m, omega_init=args.omega_init)
opt = torch.optim.AdamW(model.parameters(), lr=2e-4, weight_decay=1e-4)
sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=args.epochs, eta_min=2e-6)
best = float("inf")
for epoch in range(args.epochs):
model.train()
for xb, pb, yb in loader:
opt.zero_grad(set_to_none=True)
loss = nn.functional.mse_loss(model(xb, pb), yb)
loss.backward()
opt.step()
sched.step()
model.eval()
with torch.no_grad():
pred = model(test.x, test.pe)
rmse = torch.sqrt(torch.mean((pred - test.y) ** 2)).item()
best = min(best, rmse)
return {
"task": task,
"setting": setting,
"m": m,
"seed": seed,
"train_examples": args.train_examples,
"test_examples": args.test_examples,
"epochs": args.epochs,
"best_normalized_test_rmse": best,
"parameters": sum(p.numel() for p in model.parameters()),
}
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--repo", type=Path, required=True)
ap.add_argument("--out", type=Path, required=True)
ap.add_argument("--train-examples", type=int, default=512)
ap.add_argument("--test-examples", type=int, default=256)
ap.add_argument("--epochs", type=int, default=30)
ap.add_argument("--seeds", type=int, nargs="+", default=[0, 1])
ap.add_argument("--omega-init", choices=["zero", "uniform", "orthogonal", "none"], default="zero")
ap.add_argument("--mono-settings", nargs="+", default=["0", "5", "10", "15"])
ap.add_argument("--ms", nargs="+", type=int, default=[0, 3, 5, 10])
ap.add_argument("--skip-shortest", action="store_true")
args = ap.parse_args()
torch.set_num_threads(min(4, os.cpu_count() or 1))
GraphRoPE = import_pinned_graphrope(args.repo)
started = time.time()
rows = []
tasks = [("monochromatic", args.mono_settings)]
if not args.skip_shortest:
tasks.append(("shortest", ["watts_strogatz_p0.6"]))
for task, settings in tasks:
for setting in settings:
for m in args.ms:
for seed in args.seeds:
print(f"running task={task} setting={setting} m={m} seed={seed}", flush=True)
row = train_one(GraphRoPE, task, setting, m, seed, args)
rows.append(row)
print(f" best normalized test RMSE={row['best_normalized_test_rmse']:.6f}", flush=True)
summary = {
"protocol": {
"paper": "arXiv:2509.22259v1, Section 4.1 and Appendix A.1",
"official_repository": "https://github.com/cederikhoefs/Graph-RoPE",
"official_commit": "4ac067eb38272543b0cdd7591d630399ff37bce4",
"graph_sizes": {"monochromatic": "5x5 grid", "shortest": "10-node Watts-Strogatz, k=2, p=0.6"},
"model": "4-layer single-head transformer, hidden=32, MLP=32, dropout=0.2",
"optimizer": "AdamW(lr=2e-4, weight_decay=1e-4), cosine eta_min=2e-6",
"omega_initialization": args.omega_init,
"budget": {"train_examples": args.train_examples, "test_examples": args.test_examples, "epochs": args.epochs, "seeds": args.seeds},
"normalization": "RMSE divided by graph size (25 for monochromatic, 10 for shortest path)",
"baseline": "m=0, same model and data with WIRE disabled",
},
"rows": rows,
"runtime_seconds": time.time() - started,
}
args.out.parent.mkdir(parents=True, exist_ok=True)
args.out.write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8")
print(f"wrote {args.out} ({len(rows)} runs, {summary['runtime_seconds']:.1f}s)")
if __name__ == "__main__":
main()