q-Ary Influence Overlap Regularizer / qary_overlap_experiment.py

Mechanism failed

Raw ⬇ ZIP
  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()