Rank-One Delta Associative Memory / experiment.py
Mechanism confirmed, baseline not beaten
1import json
2import math
3import numpy as np
4
5SEED = 1198
6RNG = np.random.default_rng(SEED)
7
8
9def delta_update(W, k, v, beta):
10 pred = W @ k
11 return W + beta * np.outer(v - pred, k), pred
12
13
14def stability_and_convergence():
15 # Choose k=sqrt(lambda)e_1, v=0, and W=e_1e_1^T. Then the exact
16 # scalar error recurrence is e_(t+1)=(1-beta*lambda)e_t.
17 products = np.linspace(0.05, 2.5, 491)
18 # 100 iterations is enough to distinguish contraction from growth here.
19 final_errors = np.abs(1.0 - products) ** 100
20 stable = final_errors < 1.0
21 observed_boundary = float(products[np.where(~stable)[0][0]])
22 rows = []
23 for beta, lam in [(0.25, 0.5), (0.5, 1.0), (0.8, 1.5),
24 (1.0, 1.9), (1.0, 2.1), (1.5, 1.0)]:
25 p = beta * lam
26 k = np.array([math.sqrt(lam), 0.0])
27 W = np.eye(2)
28 observed = []
29 for _ in range(8):
30 W, _ = delta_update(W, k, np.zeros(2), beta)
31 observed.append(abs(W[0, 0]))
32 predicted = [abs(1-p)**t for t in range(1, 9)]
33 rows.append({"beta": beta, "lambda": lam, "product": p,
34 "predicted_final_error": predicted[-1],
35 "observed_final_error": float(observed[-1]),
36 "max_curve_error": float(max(abs(a-b) for a,b in zip(observed,predicted)))})
37 # For p in (0,1), t_epsilon=ceil(log(eps)/log(1-p)).
38 step_rows = []
39 eps = 1e-4
40 for beta, lam in [(0.1, 1.0), (0.2, 1.0), (0.4, 1.0), (0.2, 2.0)]:
41 p = beta * lam
42 predicted_steps = math.ceil(math.log(eps) / math.log(abs(1-p)))
43 W = np.eye(2); k = np.array([math.sqrt(lam), 0.0]); observed_steps = None
44 for t in range(1, predicted_steps + 3):
45 W, _ = delta_update(W, k, np.zeros(2), beta)
46 if abs(W[0, 0]) <= eps:
47 observed_steps = t; break
48 step_rows.append({"beta": beta, "lambda": lam, "product": p,
49 "predicted_steps_to_1e-4": predicted_steps,
50 "observed_steps_to_1e-4": observed_steps})
51 return {
52 "multiplier_max_abs_error": max(r["max_curve_error"] for r in rows),
53 "predicted_absolute_stability_boundary_beta_lambda": 2.0,
54 "observed_absolute_stability_boundary_beta_lambda": observed_boundary,
55 "boundary_grid_resolution": float(products[1]-products[0]),
56 "multiplier_sweep": rows,
57 "convergence_step_sweep": step_rows,
58 "prediction": "absolute convergence iff 0 < beta*lambda < 2; monotone convergence iff beta*lambda < 1"
59 }
60
61
62def orthogonal_independence():
63 d = 8
64 W = np.zeros((d, d)); keys = np.eye(d)
65 values = RNG.normal(size=(d, d))
66 before = W @ keys[1]
67 W, _ = delta_update(W, keys[0], values[0], 0.73)
68 after = W @ keys[1]
69 return {"predicted_cross_key_change": 0.0,
70 "observed_cross_key_change_norm": float(np.linalg.norm(after-before)),
71 "predicted_written_key_error": 0.0,
72 "observed_written_key_error": float(np.linalg.norm(W @ keys[0] - 0.73*values[0]))}
73
74
75def associative_recall(n_keys=32, dim=32, repeats=4, distractors=64):
76 key_bank = RNG.normal(size=(n_keys, dim)); key_bank /= np.linalg.norm(key_bank, axis=1, keepdims=True)
77 val_bank = RNG.normal(size=(n_keys, dim))
78 beta = 0.65; W = np.zeros((dim, dim)); write_keys=[]; write_vals=[]
79 for _ in range(repeats):
80 for i in RNG.permutation(n_keys):
81 W, _ = delta_update(W, key_bank[i], val_bank[i], beta)
82 write_keys.append(key_bank[i]); write_vals.append(val_bank[i])
83 for _ in range(distractors):
84 k=RNG.normal(size=dim); k/=np.linalg.norm(k); v=RNG.normal(size=dim)
85 W, _=delta_update(W,k,v,beta); write_keys.append(k); write_vals.append(v)
86 K=np.asarray(write_keys); V=np.asarray(write_vals)
87 delta_preds=np.asarray([W@k for k in key_bank])
88 logits=key_bank@K.T/0.08; logits-=logits.max(axis=1,keepdims=True)
89 weights=np.exp(logits); weights/=weights.sum(axis=1,keepdims=True)
90 attn_preds=weights@V
91 def metrics(P):
92 return (float(np.mean((P-val_bank)**2)), float(np.mean(np.sum(P*val_bank,axis=1)/(np.linalg.norm(P,axis=1)*np.linalg.norm(val_bank,axis=1)+1e-12))))
93 dm,dc=metrics(delta_preds); am,ac=metrics(attn_preds)
94 return {"delta_mse":dm,"attention_mse":am,"delta_cosine":dc,"attention_cosine":ac,
95 "delta_state_floats":dim*dim,"attention_stored_write_floats":len(K)*dim*2,
96 "settings":{"keys":n_keys,"dim":dim,"repeats":repeats,"distractors":distractors,"beta":beta}}
97
98
99def main():
100 result={"seed":SEED,"stability_and_convergence":stability_and_convergence(),
101 "orthogonal_prediction":orthogonal_independence(),"associative_recall":associative_recall()}
102 with open("results.json","w") as f: json.dump(result,f,indent=2)
103 print(json.dumps(result,indent=2))
104
105if __name__ == "__main__": main()