import json import math import random from itertools import product from pathlib import Path import numpy as np import torch def set_seed(seed=7): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) def exact_lemma_check(q=3, n=4, seed=7): """Exhaustively verify Lemma 3.1 for an arbitrary S-dependent law.""" rng = np.random.default_rng(seed) subsets = [np.array([(mask >> i) & 1 for i in range(n)], dtype=bool) for mask in range(1 << n)] pS = rng.dirichlet(np.ones(len(subsets))) cond = [] for s in subsets: active = int(s.sum()) cond.append(rng.dirichlet(np.ones(q ** active))) states = list(product(range(q), repeat=n)) mu = np.zeros(len(states)) nu = np.full(len(states), q ** (-n)) for si, s in enumerate(subsets): active_ix = np.flatnonzero(s) for xi, x in enumerate(states): code = 0 for pos in active_ix: code = code * q + x[pos] mu[xi] += pS[si] * cond[si][code] * q ** (-(n - len(active_ix))) chi2 = float(np.sum((mu - nu) ** 2 / nu)) moment = sum(pS[i] * pS[j] * q ** int(np.logical_and(s, t).sum()) for i, s in enumerate(subsets) for j, t in enumerate(subsets)) return {"chi2": chi2, "rhs": moment - 1.0, "slack": moment - 1.0 - chi2, "holds": bool(chi2 <= moment - 1.0 + 1e-10)} def bernoulli_moment(p, q): """For independent Bernoulli supports, E[q^|S cap S'|].""" return float(np.prod(1.0 + p * p * (q - 1.0))) def moment_from_masks(a, b, q): overlap = (a * b).sum(dim=-1) return torch.exp(math.log(q) * overlap).mean() - 1.0 def optimize(seed=7, q=4, mode="qary", steps=500, n=16, heads=4, k=4): """Toy routing optimization with hard forward supports and soft backward.""" set_seed(seed) device = "cuda" if torch.cuda.is_available() else "cpu" try: logits = torch.randn(heads, n, device=device, requires_grad=True) opt = torch.optim.Adam([logits], lr=0.06) target = torch.zeros(heads, n, device=device) target[:, :k] = 1.0 def sample(): u = torch.rand_like(logits).clamp(1e-5, 1-1e-5) g = -torch.log(-torch.log(u)) z = (logits + g) / 0.7 idx = z.topk(k, dim=-1).indices hard = torch.zeros_like(logits).scatter(1, idx, 1.0) soft = k * torch.softmax(z, dim=-1) return hard + soft - soft.detach(), hard for _ in range(steps): a, ah = sample(); b, bh = sample() task = ((torch.sigmoid(logits) - target) ** 2).mean() overlap = (a * b).sum(-1) if mode == "qary": reg = torch.exp(math.log(q) * overlap).mean() - 1.0 elif mode == "linear": reg = overlap.mean() else: reg = torch.zeros((), device=device) loss = task + (0.20 if mode != "none" else 0.0) * reg opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): ovs=[] for _ in range(200): _, aa=sample(); _, bb=sample() ovs.extend(((aa*bb).sum(-1)).cpu().numpy().tolist()) vals=np.asarray(ovs) return {"mode":mode,"q":q,"device":device,"final_task":float(task), "mean_overlap":float(vals.mean()), "severe_overlap_ge3":float(np.mean(vals>=3)), "moment_minus_1":float(np.mean(q**vals)-1.0)} except Exception as exc: if device == "cuda": torch.cuda.empty_cache() # Explicit CPU retry satisfies the shared-GPU fallback requirement. old = torch.cuda.is_available torch.cuda.is_available = lambda: False try: return optimize(seed, q, mode, steps, n, heads, k) finally: torch.cuda.is_available = old raise def main(): set_seed(7) checks = [exact_lemma_check(q=q, n=4, seed=7 + q) for q in (2, 4, 8)] analytic = {str(q): {"exact": float(bernoulli_moment(np.full(12, .5), q)), "low_collision_p": float(bernoulli_moment(np.full(12, .2), q))} for q in (2, 4, 8)} results = [] for q in (2, 4, 8): for mode in ("none", "linear", "qary"): results.append(optimize(seed=19, q=q, mode=mode)) out = {"lemma_checks": checks, "analytic_moments": analytic, "toy_results": results} Path("results.json").write_text(json.dumps(out, indent=2)) print(json.dumps(out, indent=2)) if __name__ == "__main__": main()