Rank-One Delta Associative Memory / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()