import json import math import numpy as np SEED = 1198 RNG = np.random.default_rng(SEED) def delta_update(W, k, v, beta): pred = W @ k return W + beta * np.outer(v - pred, k), pred def stability_and_convergence(): # Choose k=sqrt(lambda)e_1, v=0, and W=e_1e_1^T. Then the exact # scalar error recurrence is e_(t+1)=(1-beta*lambda)e_t. products = np.linspace(0.05, 2.5, 491) # 100 iterations is enough to distinguish contraction from growth here. final_errors = np.abs(1.0 - products) ** 100 stable = final_errors < 1.0 observed_boundary = float(products[np.where(~stable)[0][0]]) rows = [] for beta, lam in [(0.25, 0.5), (0.5, 1.0), (0.8, 1.5), (1.0, 1.9), (1.0, 2.1), (1.5, 1.0)]: p = beta * lam k = np.array([math.sqrt(lam), 0.0]) W = np.eye(2) observed = [] for _ in range(8): W, _ = delta_update(W, k, np.zeros(2), beta) observed.append(abs(W[0, 0])) predicted = [abs(1-p)**t for t in range(1, 9)] rows.append({"beta": beta, "lambda": lam, "product": p, "predicted_final_error": predicted[-1], "observed_final_error": float(observed[-1]), "max_curve_error": float(max(abs(a-b) for a,b in zip(observed,predicted)))}) # For p in (0,1), t_epsilon=ceil(log(eps)/log(1-p)). step_rows = [] eps = 1e-4 for beta, lam in [(0.1, 1.0), (0.2, 1.0), (0.4, 1.0), (0.2, 2.0)]: p = beta * lam predicted_steps = math.ceil(math.log(eps) / math.log(abs(1-p))) W = np.eye(2); k = np.array([math.sqrt(lam), 0.0]); observed_steps = None for t in range(1, predicted_steps + 3): W, _ = delta_update(W, k, np.zeros(2), beta) if abs(W[0, 0]) <= eps: observed_steps = t; break step_rows.append({"beta": beta, "lambda": lam, "product": p, "predicted_steps_to_1e-4": predicted_steps, "observed_steps_to_1e-4": observed_steps}) return { "multiplier_max_abs_error": max(r["max_curve_error"] for r in rows), "predicted_absolute_stability_boundary_beta_lambda": 2.0, "observed_absolute_stability_boundary_beta_lambda": observed_boundary, "boundary_grid_resolution": float(products[1]-products[0]), "multiplier_sweep": rows, "convergence_step_sweep": step_rows, "prediction": "absolute convergence iff 0 < beta*lambda < 2; monotone convergence iff beta*lambda < 1" } def orthogonal_independence(): d = 8 W = np.zeros((d, d)); keys = np.eye(d) values = RNG.normal(size=(d, d)) before = W @ keys[1] W, _ = delta_update(W, keys[0], values[0], 0.73) after = W @ keys[1] return {"predicted_cross_key_change": 0.0, "observed_cross_key_change_norm": float(np.linalg.norm(after-before)), "predicted_written_key_error": 0.0, "observed_written_key_error": float(np.linalg.norm(W @ keys[0] - 0.73*values[0]))} def associative_recall(n_keys=32, dim=32, repeats=4, distractors=64): key_bank = RNG.normal(size=(n_keys, dim)); key_bank /= np.linalg.norm(key_bank, axis=1, keepdims=True) val_bank = RNG.normal(size=(n_keys, dim)) beta = 0.65; W = np.zeros((dim, dim)); write_keys=[]; write_vals=[] for _ in range(repeats): for i in RNG.permutation(n_keys): W, _ = delta_update(W, key_bank[i], val_bank[i], beta) write_keys.append(key_bank[i]); write_vals.append(val_bank[i]) for _ in range(distractors): k=RNG.normal(size=dim); k/=np.linalg.norm(k); v=RNG.normal(size=dim) W, _=delta_update(W,k,v,beta); write_keys.append(k); write_vals.append(v) K=np.asarray(write_keys); V=np.asarray(write_vals) delta_preds=np.asarray([W@k for k in key_bank]) logits=key_bank@K.T/0.08; logits-=logits.max(axis=1,keepdims=True) weights=np.exp(logits); weights/=weights.sum(axis=1,keepdims=True) attn_preds=weights@V def metrics(P): 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)))) dm,dc=metrics(delta_preds); am,ac=metrics(attn_preds) return {"delta_mse":dm,"attention_mse":am,"delta_cosine":dc,"attention_cosine":ac, "delta_state_floats":dim*dim,"attention_stored_write_floats":len(K)*dim*2, "settings":{"keys":n_keys,"dim":dim,"repeats":repeats,"distractors":distractors,"beta":beta}} def main(): result={"seed":SEED,"stability_and_convergence":stability_and_convergence(), "orthogonal_prediction":orthogonal_independence(),"associative_recall":associative_recall()} with open("results.json","w") as f: json.dump(result,f,indent=2) print(json.dumps(result,indent=2)) if __name__ == "__main__": main()