Role-Filler Attention / role_filler_experiment.py
Failed on benchmark
1import json, math, random
2from pathlib import Path
3import numpy as np
4
5SEED = 2951
6rng = np.random.default_rng(SEED)
7
8
9def softmax(x, axis=-1):
10 z = x - np.max(x, axis=axis, keepdims=True)
11 e = np.exp(z)
12 return e / e.sum(axis=axis, keepdims=True)
13
14
15def entropy(p):
16 return float(-(p * np.log(np.maximum(p, 1e-12))).sum(axis=-1).mean())
17
18
19def math_check(R=5, D=7, N=9):
20 # QR gives exactly orthonormal role vectors (up to numerical precision).
21 Q, _ = np.linalg.qr(rng.normal(size=(D, R)))
22 roles = Q.T # [R,D], rows orthonormal
23 F = rng.normal(size=(N, R, D))
24 O = np.einsum('rd,nre->nde', roles, F) # [N,D,D]
25 recovered = np.einsum('rd,nde->nre', roles, O)
26 contraction_error = float(np.max(np.abs(recovered - F)))
27
28 # Query role m and filler from object 3 should retrieve object 3.
29 m, target, source = 1, 3, 3
30 qf = F[source, m]
31 scores = F[:, m] @ qf
32 probs = softmax(scores / 0.15)
33 top = int(np.argmax(probs))
34 out = probs @ F[:, target]
35 retrieval_error = float(np.linalg.norm(out - F[source, target]))
36
37 # Wrong-role query is a control: it should not be systematically selective.
38 wrong = (m + 1) % R
39 wrong_probs = softmax((F[:, wrong] @ qf) / 0.15)
40 return {
41 'max_role_contraction_abs_error': contraction_error,
42 'exact_query_top1': top == source,
43 'exact_query_probability': float(probs[source]),
44 'exact_rebinding_l2_error': retrieval_error,
45 'exact_entropy': entropy(probs[None, :]),
46 'mismatched_role_entropy': entropy(wrong_probs[None, :]),
47 }
48
49
50def make_episode(N=16, R=4, D=16, noise=0.08):
51 # Each object has independent role fillers. The query is one role filler;
52 # the desired answer is another role filler from that same object.
53 F = rng.normal(size=(N, R, D)).astype(np.float32)
54 src = int(rng.integers(N))
55 m, target = 0, 1
56 q = F[src, m] + noise * rng.normal(size=D).astype(np.float32)
57 return F, q, src, m, target
58
59
60def structured_trial(F, q, m, target, tau=0.25):
61 s = F[:, m] @ q
62 a = softmax(s / tau)
63 y = a @ F[:, target]
64 return int(np.argmax(s)), y, a
65
66
67def dense_trial(F, q, target, tau=0.25):
68 # Control: a standard flattened-object dot product, with the query copied
69 # into every role block so it has the same input dimension as an object.
70 flat = F.reshape(F.shape[0], -1)
71 qflat = np.tile(q, F.shape[1])
72 s = flat @ qflat
73 a = softmax(s / tau)
74 return int(np.argmax(s)), a @ F[:, target], a
75
76
77def benchmark(episodes=2000):
78 stats = {k: [] for k in ['structured_acc','dense_acc','structured_err','dense_err',
79 'structured_entropy','dense_entropy','wrong_entropy']}
80 for _ in range(episodes):
81 F, q, src, m, target = make_episode()
82 si, sy, sa = structured_trial(F, q, m, target)
83 di, dy, da = dense_trial(F, q, target)
84 wrong = (m + 1) % F.shape[1]
85 _, _, wa = structured_trial(F, q, wrong, target)
86 stats['structured_acc'].append(si == src)
87 stats['dense_acc'].append(di == src)
88 stats['structured_err'].append(np.linalg.norm(sy - F[src,target]))
89 stats['dense_err'].append(np.linalg.norm(dy - F[src,target]))
90 stats['structured_entropy'].append(entropy(sa[None,:]))
91 stats['dense_entropy'].append(entropy(da[None,:]))
92 stats['wrong_entropy'].append(entropy(wa[None,:]))
93 return {k: float(np.mean(v)) for k,v in stats.items()}
94
95
96def main():
97 np.set_printoptions(precision=6, suppress=True)
98 result = {'seed': SEED, 'math_check': math_check(), 'benchmark': benchmark()}
99 # Repeat the cheap benchmark with a different seed to check direction stability.
100 global rng
101 rng = np.random.default_rng(SEED + 1)
102 result['benchmark_repeat'] = benchmark(episodes=1000)
103 out = Path('results.json')
104 out.write_text(json.dumps(result, indent=2))
105 print(json.dumps(result, indent=2))
106
107if __name__ == '__main__':
108 main()