| |
| """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 |
|
|
| 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() |
|
|