Compressed threshold-overlap Gram layer / experiment.py
Mechanism failed
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()