Implicit Higher-Order TPR Memory / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, time
2from pathlib import Path
3import numpy as np
4
5SEED = 2952
6rng = np.random.default_rng(SEED)
7
8
9def softmax(x):
10 z = x - np.max(x, axis=-1, keepdims=True)
11 e = np.exp(z)
12 return e / e.sum(axis=-1, keepdims=True)
13
14
15def factorized_scores(objects, role_queries, filler_queries):
16 # objects: [N,R,D], role_queries: [m], filler_queries: [m,D]
17 sims = np.stack([objects[:, r, :] @ q for r, q in zip(role_queries, filler_queries)], axis=1)
18 return np.prod(sims, axis=1), sims
19
20
21def explicit_scores(objects, role_queries, filler_queries):
22 # Build the order-m query and O_t tensor product, then contract exactly.
23 # This is intentionally only used for the small algebraic sanity check.
24 m = len(role_queries)
25 vals = []
26 for obj in objects:
27 flat = obj.reshape(-1)
28 qs = []
29 for r, q in zip(role_queries, filler_queries):
30 e = np.zeros_like(obj)
31 e[r] = q
32 qs.append(e.reshape(-1))
33 tensor = flat
34 query = qs[0]
35 for k in range(1, m):
36 tensor = np.outer(tensor, flat).reshape(-1)
37 query = np.outer(query, qs[k]).reshape(-1)
38 vals.append(float(tensor @ query))
39 return np.asarray(vals)
40
41
42def make_problem(n, roles=4, dim=16, m=3, noise=0.08):
43 # Each object contains one normalized filler per role. One target object
44 # matches all query bindings; distractors match individual bindings only.
45 fillers = rng.normal(size=(roles, dim))
46 fillers /= np.linalg.norm(fillers, axis=1, keepdims=True)
47 target = rng.normal(size=(roles, dim))
48 target /= np.linalg.norm(target, axis=1, keepdims=True)
49 role_queries = np.array([0, 1, 2][:m])
50 filler_queries = target[role_queries]
51 objects = rng.normal(size=(n, roles, dim))
52 objects /= np.linalg.norm(objects, axis=2, keepdims=True)
53 objects[0, role_queries] = filler_queries + noise * rng.normal(size=(m, dim))
54 objects[0, role_queries] /= np.linalg.norm(objects[0, role_queries], axis=1, keepdims=True)
55 # Half of distractors contain one matching factor, making conjunction useful.
56 for i in range(1, n):
57 if i % 2 == 0:
58 k = i % m
59 objects[i, role_queries[k]] = filler_queries[k] + noise * rng.normal(size=dim)
60 objects[i, role_queries[k]] /= np.linalg.norm(objects[i, role_queries[k]])
61 return objects, role_queries, filler_queries
62
63
64def benchmark(n, trials=300, m=3):
65 hits_factor, hits_single, margins_factor, margins_single = [], [], [], []
66 t0 = time.perf_counter()
67 for _ in range(trials):
68 x, roles, qs = make_problem(n=n, m=m)
69 sf, sims = factorized_scores(x, roles, qs)
70 # Standard single-factor control: best one binding, equivalent to
71 # ordinary attention that cannot enforce the conjunction.
72 ss = sims[:, 0]
73 order = np.argsort(sf)[::-1]
74 order_s = np.argsort(ss)[::-1]
75 hits_factor.append(int(order[0] == 0))
76 hits_single.append(int(order_s[0] == 0))
77 margins_factor.append(sf[order[0]] - sf[order[1]])
78 margins_single.append(ss[order_s[0]] - ss[order_s[1]])
79 elapsed = time.perf_counter() - t0
80 return {
81 "n": n, "m": m, "trials": trials,
82 "factor_accuracy": float(np.mean(hits_factor)),
83 "single_accuracy": float(np.mean(hits_single)),
84 "factor_margin": float(np.mean(margins_factor)),
85 "single_margin": float(np.mean(margins_single)),
86 "seconds": elapsed,
87 # Explicit order-m tensor has N*D^(2m) scalar scale for order-2 objects;
88 # factorized retrieval stores objects and m query vectors: linear in N.
89 "factor_memory_units": int(n * 4 * 16),
90 "explicit_memory_units_relative": int(n * (4 * 16) ** m),
91 }
92
93
94def main():
95 # Algebraic verification with signed values and multiple orders.
96 max_err = 0.0
97 for m in (2, 3):
98 x = rng.normal(size=(4, 3, 3))
99 roles = np.arange(m)
100 qs = rng.normal(size=(m, 3))
101 a = factorized_scores(x, roles, qs)[0]
102 b = explicit_scores(x, roles, qs)
103 max_err = max(max_err, float(np.max(np.abs(a - b))))
104
105 results = {"identity_max_abs_error": max_err}
106 for n in (32, 128, 256):
107 results[f"n{n}_m2"] = benchmark(n, m=2)
108 results[f"n{n}_m3"] = benchmark(n, m=3)
109 # Verify output extraction is linear and soft retrieval is normalized.
110 x, roles, qs = make_problem(64, m=3)
111 scores, _ = factorized_scores(x, roles, qs)
112 attn = softmax(scores / 0.1)
113 out = np.einsum("n,nrd->rd", attn, x)
114 results["attention_sum"] = float(attn.sum())
115 results["output_shape"] = list(out.shape)
116 Path("results.json").write_text(json.dumps(results, indent=2))
117 print(json.dumps(results, indent=2))
118
119if __name__ == "__main__":
120 main()