Implicit Higher-Order TPR Memory / experiment.py

Mechanism confirmed, baseline not beaten

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