import json, math, random import numpy as np SEED = 996 np.random.seed(SEED); random.seed(SEED) def sat(u): return np.clip(u, -1.0, 1.0) def xi_law(a, kappa=2.0, eps=0.05, xmin=1.0, xmax=8.0): return np.clip(kappa/(eps + np.asarray(a)), xmin, xmax) def midpoint_quantize(x, L): # Uniform interval quantizer on [-1,1], midpoint reconstruction. edges = np.linspace(-1, 1, int(L)+1) idx = np.clip(np.searchsorted(edges, x, side='right')-1, 0, L-1) return (edges[idx] + edges[idx+1])/2 def math_checks(): # Prediction 1: saturation begins exactly at |x|=1/xi. xis = [1.25, 2., 4., 7.] threshold_err = [] for z in xis: x = np.linspace(0, 1, 200001) y = sat(z*x) first = x[np.argmax(y >= 1-1e-10)] threshold_err.append(abs(first - 1/z)) # Prediction 2: unsaturated local slope is xi and decreases with influence. aa = np.linspace(0.05, 2.0, 20) zz = xi_law(aa) slopes = [] for z in zz: x = np.linspace(-0.05, 0.05, 101) slopes.append(np.polyfit(x, sat(z*x), 1)[0]) slope_err = float(np.max(np.abs(np.asarray(slopes)-zz))) monotone = bool(np.all(np.diff(zz) <= 1e-12)) # Prediction 3: increasing strategic pressure reduces adaptive credible bins. # The proposed design maps xi to a variable number of bins, capped at binary. kappas = np.linspace(.5, 8., 16) a = .30 bins = [max(2, int(np.rint(8/xi_law(a, kappa=k)))) for k in kappas] occupied = [] x = np.linspace(-1, 1, 20001) for k, L in zip(kappas, bins): z = xi_law(a, kappa=k) occupied.append(len(np.unique(midpoint_quantize(sat(z*x), L)))) return { 'threshold_max_abs_error': float(max(threshold_err)), 'threshold_observed': [float(1/z) for z in xis], 'threshold_errors': [float(v) for v in threshold_err], 'influence_xi_at_low_high_a': [float(zz[0]), float(zz[-1])], 'influence_xi_monotone_decreasing': monotone, 'slope_max_abs_error_vs_xi': slope_err, 'kappas': [float(v) for v in kappas], 'adaptive_bins': bins, 'occupied_bins': occupied, 'binary_reached': bool(min(bins)==2 and occupied[-1] <= 2), 'formula_bounds_ok': bool(np.all(np.abs(sat(zz*x[:len(zz)])) <= 1+1e-12)) } def gnn_experiment(): try: import torch torch.manual_seed(SEED) device = 'cuda' if torch.cuda.is_available() else 'cpu' try: if device == 'cuda': torch.cuda.empty_cache() except Exception: device='cpu' n=96; d=12; c=2 # Homophilic SBM, undirected with self loops, row normalized. y=torch.arange(n, device=device)%2 rng=np.random.default_rng(SEED) A=np.zeros((n,n),dtype=np.float32) for i in range(n): for j in range(i+1,n): p=.24 if int(y[i])==int(y[j]) else .035 if rng.random() larger xi -> fewer bins L=torch.clamp(torch.round(8./xi),min=2,max=8).long() vals=[] for i in range(n): vals.append(midpoint_quantize(msg[i].detach().cpu().numpy(),int(L[i]))) msg=torch.tensor(np.asarray(vals),device=device,dtype=X.dtype) h=torch.relu(h + At @ self.wm(msg[:,None])) return self.out(h), raw, msg, xi def train_model(model, strategic=False, quant=False): opt=torch.optim.Adam(model.parameters(),lr=.025,weight_decay=1e-3) for step in range(180): opt.zero_grad() out,*_=model() if strategic else (model(),) loss=torch.nn.functional.cross_entropy(out[tr],y[tr]); loss.backward(); opt.step() with torch.no_grad(): if strategic: out,raw,msg,xi=model(eval_quant=quant) else: out=model(); raw=msg=xi=None acc=(out[va].argmax(1)==y[va]).float().mean().item() return acc, raw, msg, xi b=Base().to(device); bacc,*_=train_model(b) s=Strategic(False).to(device); sacc,sraw,smsg,sxi=train_model(s,True,False) q=Strategic(True).to(device); qacc,qraw,qmsg,qxi=train_model(q,True,True) with torch.no_grad(): satur=np.abs((sxi*sraw).cpu().numpy())>=1-1e-5 bits=np.log2(np.maximum(2,np.rint(8/sxi.cpu().numpy()))) ent=None # communication summary uses adaptive assigned bits, not tensor storage. return {'device':device,'nodes':n,'train_steps':180, 'validation_accuracy':{'baseline_gcn':bacc,'strategic_clipped':sacc,'strategic_adaptive_quantized':qacc}, 'strategic_saturation_fraction':float(satur.mean()), 'mean_adaptive_bits':float(bits.mean()),'min_max_xi':[float(sxi.min()),float(sxi.max())], 'bounded_message_max_abs':float(torch.max(torch.abs(smsg)).item()), 'mean_influence':float(sxi.numel() and (At.sum(0).mean()).item())} except Exception as e: return {'error':repr(e),'fallback':'math checks still valid'} def main(): result={'seed':SEED,'math_checks':math_checks(),'gnn_experiment':gnn_experiment()} print(json.dumps(result,indent=2,sort_keys=True)) if __name__=='__main__': main()