q-Ary Influence Overlap Regularizer / qary_overlap_experiment.py
Mechanism failed
1import json
2import math
3import random
4from itertools import product
5from pathlib import Path
6
7import numpy as np
8import torch
9
10
11def set_seed(seed=7):
12 random.seed(seed)
13 np.random.seed(seed)
14 torch.manual_seed(seed)
15
16
17def exact_lemma_check(q=3, n=4, seed=7):
18 """Exhaustively verify Lemma 3.1 for an arbitrary S-dependent law."""
19 rng = np.random.default_rng(seed)
20 subsets = [np.array([(mask >> i) & 1 for i in range(n)], dtype=bool)
21 for mask in range(1 << n)]
22 pS = rng.dirichlet(np.ones(len(subsets)))
23 cond = []
24 for s in subsets:
25 active = int(s.sum())
26 cond.append(rng.dirichlet(np.ones(q ** active)))
27 states = list(product(range(q), repeat=n))
28 mu = np.zeros(len(states))
29 nu = np.full(len(states), q ** (-n))
30 for si, s in enumerate(subsets):
31 active_ix = np.flatnonzero(s)
32 for xi, x in enumerate(states):
33 code = 0
34 for pos in active_ix:
35 code = code * q + x[pos]
36 mu[xi] += pS[si] * cond[si][code] * q ** (-(n - len(active_ix)))
37 chi2 = float(np.sum((mu - nu) ** 2 / nu))
38 moment = sum(pS[i] * pS[j] * q ** int(np.logical_and(s, t).sum())
39 for i, s in enumerate(subsets) for j, t in enumerate(subsets))
40 return {"chi2": chi2, "rhs": moment - 1.0,
41 "slack": moment - 1.0 - chi2,
42 "holds": bool(chi2 <= moment - 1.0 + 1e-10)}
43
44
45def bernoulli_moment(p, q):
46 """For independent Bernoulli supports, E[q^|S cap S'|]."""
47 return float(np.prod(1.0 + p * p * (q - 1.0)))
48
49
50def moment_from_masks(a, b, q):
51 overlap = (a * b).sum(dim=-1)
52 return torch.exp(math.log(q) * overlap).mean() - 1.0
53
54
55def optimize(seed=7, q=4, mode="qary", steps=500, n=16, heads=4, k=4):
56 """Toy routing optimization with hard forward supports and soft backward."""
57 set_seed(seed)
58 device = "cuda" if torch.cuda.is_available() else "cpu"
59 try:
60 logits = torch.randn(heads, n, device=device, requires_grad=True)
61 opt = torch.optim.Adam([logits], lr=0.06)
62 target = torch.zeros(heads, n, device=device)
63 target[:, :k] = 1.0
64 def sample():
65 u = torch.rand_like(logits).clamp(1e-5, 1-1e-5)
66 g = -torch.log(-torch.log(u))
67 z = (logits + g) / 0.7
68 idx = z.topk(k, dim=-1).indices
69 hard = torch.zeros_like(logits).scatter(1, idx, 1.0)
70 soft = k * torch.softmax(z, dim=-1)
71 return hard + soft - soft.detach(), hard
72 for _ in range(steps):
73 a, ah = sample(); b, bh = sample()
74 task = ((torch.sigmoid(logits) - target) ** 2).mean()
75 overlap = (a * b).sum(-1)
76 if mode == "qary":
77 reg = torch.exp(math.log(q) * overlap).mean() - 1.0
78 elif mode == "linear":
79 reg = overlap.mean()
80 else:
81 reg = torch.zeros((), device=device)
82 loss = task + (0.20 if mode != "none" else 0.0) * reg
83 opt.zero_grad(); loss.backward(); opt.step()
84 with torch.no_grad():
85 ovs=[]
86 for _ in range(200):
87 _, aa=sample(); _, bb=sample()
88 ovs.extend(((aa*bb).sum(-1)).cpu().numpy().tolist())
89 vals=np.asarray(ovs)
90 return {"mode":mode,"q":q,"device":device,"final_task":float(task),
91 "mean_overlap":float(vals.mean()),
92 "severe_overlap_ge3":float(np.mean(vals>=3)),
93 "moment_minus_1":float(np.mean(q**vals)-1.0)}
94 except Exception as exc:
95 if device == "cuda":
96 torch.cuda.empty_cache()
97 # Explicit CPU retry satisfies the shared-GPU fallback requirement.
98 old = torch.cuda.is_available
99 torch.cuda.is_available = lambda: False
100 try: return optimize(seed, q, mode, steps, n, heads, k)
101 finally: torch.cuda.is_available = old
102 raise
103
104
105def main():
106 set_seed(7)
107 checks = [exact_lemma_check(q=q, n=4, seed=7 + q) for q in (2, 4, 8)]
108 analytic = {str(q): {"exact": float(bernoulli_moment(np.full(12, .5), q)),
109 "low_collision_p": float(bernoulli_moment(np.full(12, .2), q))}
110 for q in (2, 4, 8)}
111 results = []
112 for q in (2, 4, 8):
113 for mode in ("none", "linear", "qary"):
114 results.append(optimize(seed=19, q=q, mode=mode))
115 out = {"lemma_checks": checks, "analytic_moments": analytic, "toy_results": results}
116 Path("results.json").write_text(json.dumps(out, indent=2))
117 print(json.dumps(out, indent=2))
118
119
120if __name__ == "__main__":
121 main()