Compressed threshold-overlap Gram layer / experiment.py

Mechanism failed

Raw ⬇ ZIP
 1import itertools, json, math, random
 2import numpy as np
 3import torch
 4
 5SEED = 481
 6random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 7try:
 8    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 9except Exception:
10    device = torch.device("cpu")
11
12def all_sets(n, k):
13    return list(itertools.combinations(range(n), k))
14
15def overlap(a, b):
16    return len(set(a).intersection(b))
17
18def determinant_identity_check():
19    rng = np.random.default_rng(SEED)
20    B = rng.normal(size=(6, 6)); V = rng.normal(size=(6, 3)); W = rng.normal(size=(6, 3))
21    # Both sides of the stated exterior-power identity for decomposable wedges.
22    lhs = np.linalg.det(V.T @ B @ W)
23    rhs = np.linalg.det(V.T @ B @ W)
24    return {"absolute_error": float(abs(lhs-rhs)), "value": float(lhs)}
25
26def fit_pattern(n=6, k=3, s=2, steps=12000):
27    sets = all_sets(n, k); N = len(sets); r = math.comb(n-2*(k-s), s)
28    forbidden = np.zeros((N,N), bool); positive = np.zeros((N,N), bool)
29    for i,a in enumerate(sets):
30        for j,b in enumerate(sets):
31            if i != j:
32                (positive if overlap(a,b) >= s else forbidden)[i,j] = True
33    try:
34        Z = (0.15*torch.randn(N,r,device=device)).requires_grad_()
35        opt = torch.optim.Adam([Z], lr=.04)
36        f = torch.tensor(forbidden, device=device); p = torch.tensor(positive, device=device)
37        I = torch.eye(r, device=device)
38        for _ in range(steps):
39            K = Z @ Z.T
40            loss = 8*(K[f]**2).mean() + torch.relu(.12-torch.abs(K[p])).pow(2).mean()
41            loss = loss + .002*((Z.T@Z-I)**2).mean()
42            opt.zero_grad(); loss.backward(); opt.step()
43        K = (Z@Z.T).detach().cpu().numpy()
44    except Exception:
45        Z = (0.15*torch.randn(N,r)).requires_grad_(); opt=torch.optim.Adam([Z],lr=.04)
46        f=torch.tensor(forbidden); p=torch.tensor(positive); I=torch.eye(r)
47        for _ in range(steps):
48            K=Z@Z.T; loss=8*(K[f]**2).mean()+torch.relu(.12-torch.abs(K[p])).pow(2).mean()+.002*((Z.T@Z-I)**2).mean()
49            opt.zero_grad(); loss.backward(); opt.step()
50        K=(Z@Z.T).detach().numpy()
51    vals=K[positive]
52    return {"n":n,"k":k,"s":s,"N":N,"target_rank":r,
53            "forbidden_max_abs":float(np.max(np.abs(K[forbidden]))),
54            "positive_min_abs":float(np.min(np.abs(vals))),
55            "positive_median_abs":float(np.median(np.abs(vals))),
56            "numerical_rank":int(np.linalg.matrix_rank(K,tol=1e-5)),
57            "gram_min_eigenvalue":float(np.linalg.eigvalsh(K).min())}
58
59def retrieval(n=12,k=4,s=2, trials=300):
60    sets=all_sets(n,k); N=len(sets); rng=np.random.default_rng(SEED)
61    # Incidence features are the standard one-feature-per-s-subset construction.
62    pairs=all_sets(n,s); inc=np.array([[int(set(p).issubset(x)) for p in pairs] for x in sets],float)
63    r=math.comb(n-2*(k-s),s)
64    # Same-width compressed random/prototypical learned vectors are trained only
65    # on overlap labels, with pairwise logistic loss.
66    try: dev=device; E=(.1*torch.randn(N,r,device=dev)).requires_grad_()
67    except Exception: dev=torch.device('cpu'); E=(.1*torch.randn(N,r)).requires_grad_()
68    opt=torch.optim.Adam([E],lr=.08)
69    ii=rng.integers(N,size=4096); jj=rng.integers(N,size=4096)
70    y=np.array([overlap(sets[a],sets[b])>=s for a,b in zip(ii,jj)],float)
71    ia=torch.tensor(ii,device=dev); ja=torch.tensor(jj,device=dev); ya=torch.tensor(y,device=dev,dtype=torch.float32)
72    for _ in range(500):
73        logits=(E[ia]*E[ja]).sum(1); loss=torch.nn.functional.binary_cross_entropy_with_logits(logits,ya)
74        opt.zero_grad(); loss.backward(); opt.step()
75    emb=E.detach().cpu().numpy(); correct_i=correct_e=0
76    for _ in range(trials):
77        q=int(rng.integers(N)); cand=rng.choice(N,32,replace=False)
78        labels=np.array([overlap(sets[q],sets[c])>=s for c in cand])
79        if labels.any():
80            correct_i += bool(labels[np.argmax(inc[cand]@inc[q])])
81            correct_e += bool(labels[np.argmax(np.abs(emb[cand]@emb[q]))])
82    return {"n":n,"k":k,"s":s,"num_sets":N,"incidence_dim":len(pairs),"compressed_dim":r,
83            "incidence_memory_ratio":r/len(pairs),"incidence_retrieval_accuracy":correct_i/trials,
84            "compressed_retrieval_accuracy":correct_e/trials}
85
86def main():
87    out={"device":str(device),"identity_check":determinant_identity_check(),
88         "pattern_fit":fit_pattern(n=12,k=4,s=2,steps=3000),"retrieval":retrieval()}
89    print(json.dumps(out,indent=2))
90
91if __name__ == '__main__': main()