Influence-Adaptive Strategic Quantization / graph_strategic_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, random, sys
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import train_model, evaluate, sweep_baseline, make_report
  9
 10META = {'name': 'influence_graph_regression', 'domain': 'graph-nn',
 11        'description': 'Fixed-topology synthetic graph regression where the target depends on one-hop nonlinear neighbor aggregation.'}
 12
 13
 14def get_dataset(seed, n_train=400, n_test=120):
 15    rng = np.random.RandomState(seed)
 16    n, f = 10, 3
 17    A = rng.uniform(0.05, 1.0, size=(n, n)).astype('float32')
 18    A[rng.rand(n, n) < 0.52] = 0.0
 19    np.fill_diagonal(A, 0.0)
 20    A += np.eye(n, dtype='float32') * 0.15
 21    A /= A.sum(1, keepdims=True)
 22    def make(m, rs):
 23        z = rs.normal(size=(m, n, f)).astype('float32')
 24        scalar = np.tanh(z[:, :, 0] * 0.9 + z[:, :, 1] * 0.35)
 25        neigh = np.einsum('ij,bjf->bif', A, z)
 26        y = (0.65 * neigh[:, :, 0] + 0.30 * np.tanh(neigh[:, :, 1]) + 0.20 * scalar).mean(1)
 27        y += 0.04 * rs.normal(size=m)
 28        topo = np.broadcast_to(A[None, :, :], (m, n, n)).copy()
 29        return np.concatenate([z, topo], axis=2).astype('float32'), y.astype('float32')[:, None]
 30    xtr, ytr = make(n_train, rng)
 31    xte, yte = make(n_test, np.random.RandomState(seed + 5000))
 32    return {'xtr': xtr, 'ytr': ytr, 'xte': xte, 'yte': yte, 'task': 'regression',
 33            'metric': 'mse', 'out_dim': 1, 'input_shape': (n, n + f)}
 34
 35
 36def torch_dataset(seed, n_train=400, n_test=120):
 37    d = get_dataset(seed, n_train, n_test)
 38    out = dict(d)
 39    for k in ('xtr', 'ytr', 'xte', 'yte'):
 40        out[k] = torch.as_tensor(d[k], dtype=torch.float32)
 41    return out
 42
 43
 44class GraphMessageNet(nn.Module):
 45    def __init__(self, idea=False, kappa=2.0, eps=0.05, xmin=1., xmax=8.):
 46        super().__init__()
 47        self.idea, self.kappa, self.eps = idea, kappa, eps
 48        self.xmin, self.xmax = xmin, xmax
 49        self.node = nn.Linear(3, 24)
 50        self.msg_proj = nn.Linear(1, 24, bias=False)
 51        self.out = nn.Sequential(nn.Linear(24, 24), nn.ReLU(), nn.Linear(24, 1))
 52        self.last = {}
 53    def forward(self, x):
 54        n = x.shape[1]
 55        feat, A = x[:, :, :3], x[:, :, 3:3+n]
 56        h = torch.relu(self.node(feat))
 57        raw = torch.tanh(feat[:, :, 0] * 0.7 + feat[:, :, 1] * 0.2)
 58        influence = A.mean(1)[0]
 59        if self.idea:
 60            xi = torch.clamp(self.kappa / (self.eps + influence), self.xmin, self.xmax)
 61            msg = torch.clamp(raw * xi[None, :], -1., 1.)
 62        else:
 63            xi = torch.ones_like(influence)
 64            msg = raw
 65        agg = torch.einsum('bij,bj->bi', A, msg)
 66        pooled = h.mean(1) + self.msg_proj(agg.unsqueeze(-1)).mean(1)
 67        self.last = {'raw': raw.detach(), 'msg': msg.detach(), 'xi': xi.detach(), 'influence': influence.detach()}
 68        return self.out(pooled)
 69
 70
 71def run_one(seed, idea, lr, kappa=2.0, epochs=18):
 72    torch.manual_seed(10000 + seed); np.random.seed(10000 + seed); random.seed(10000 + seed)
 73    ds = torch_dataset(seed)
 74    model = GraphMessageNet(idea=idea, kappa=kappa)
 75    model, metric, _ = train_model(model, ds, epochs=epochs, lr=lr, batch=64,
 76                                   weight_decay=1e-4, log=lambda *_: None)
 77    if model is None or metric is None:
 78        return float('nan'), None
 79    dev = next(model.parameters()).device
 80    with torch.no_grad(): _ = model(ds['xte'].to(dev))
 81    return float(metric), model
 82
 83
 84def main():
 85    grid = [{'lr': 1e-3, 'kappa': 2.0}, {'lr': 3e-3, 'kappa': 2.0}, {'lr': 1e-2, 'kappa': 2.0}]
 86    def base_factory(cfg):
 87        return lambda seed: run_one(seed, False, cfg['lr'])[0]
 88    baseline = sweep_baseline(base_factory, grid)
 89    best_lr = baseline['best_cfg']['lr']
 90    idea_cfgs = [{'lr': best_lr, 'kappa': 1.0}, {'lr': best_lr, 'kappa': 2.0}, {'lr': best_lr, 'kappa': 4.0}]
 91    idea_trials = []
 92    for cfg in idea_cfgs:
 93        r = evaluate(lambda seed, c=cfg: run_one(seed, True, c['lr'], c['kappa'])[0])
 94        idea_trials.append({'cfg': cfg, **r})
 95    idea = min(idea_trials, key=lambda r: r['mean'])
 96    sig_rows = []
 97    for seed in range(8):
 98        _, m = run_one(seed, True, idea['cfg']['lr'], idea['cfg']['kappa'])
 99        z = m.last
100        raw = z['raw'].cpu().numpy().ravel(); msg = z['msg'].cpu().numpy().ravel()
101        xi0 = z['xi'].cpu().numpy().ravel(); a0 = z['influence'].cpu().numpy().ravel()
102        reps = raw.size // xi0.size
103        xi = np.tile(xi0, reps); a = np.tile(a0, reps)
104        unsat = np.abs(raw * xi) < .98
105        slope = float(np.mean(np.abs(msg[unsat]) / (np.abs(raw[unsat]) + 1e-8))) if unsat.any() else 0.
106        sig_rows.append({'slope_observed': slope, 'xi_mean': float(xi.mean()),
107                         'sat_fraction': float((np.abs(raw * xi) >= .98).mean()),
108                         'bounded_max': float(np.abs(msg).max()),
109                         'corr_influence_xi': float(np.corrcoef(a, xi)[0, 1])})
110    predicted = float(np.mean([r['xi_mean'] for r in sig_rows])); observed = float(np.mean([r['slope_observed'] for r in sig_rows]))
111    relerr = abs(observed-predicted)/(abs(predicted)+1e-8)
112    signature = {'prediction': 'trained unsaturated message gain tracks xi=kappa/(epsilon+a), and messages stay in [-1,1]',
113                 'predicted_mean_xi': predicted, 'observed_mean_gain': observed,
114                 'relative_gain_error': relerr, 'mean_saturation_fraction': float(np.mean([r['sat_fraction'] for r in sig_rows])),
115                 'max_bounded_message': float(max(r['bounded_max'] for r in sig_rows)), 'per_seed': sig_rows,
116                 'confirmed': bool(relerr < .15 and max(r['bounded_max'] for r in sig_rows) <= 1.00001)}
117    rep = make_report('influence_graph_regression', 'custom_graph_message_net', baseline,
118                      {'cfg': idea['cfg'], 'mean': idea['mean'], 'std': idea['std'], 'per_seed': idea['per_seed'], 'n': idea['n']}, signature)
119    rep['custom_track'] = {'name': META['name'], 'file': 'graph_strategic_bench.py', 'domain': META['domain']}
120    Path('bench_report.json').write_text(json.dumps(rep, indent=2)); print(json.dumps(rep, indent=2))
121
122if __name__ == '__main__': main()