import json, math, time from pathlib import Path import numpy as np SEED = 2952 rng = np.random.default_rng(SEED) def softmax(x): z = x - np.max(x, axis=-1, keepdims=True) e = np.exp(z) return e / e.sum(axis=-1, keepdims=True) def factorized_scores(objects, role_queries, filler_queries): # objects: [N,R,D], role_queries: [m], filler_queries: [m,D] sims = np.stack([objects[:, r, :] @ q for r, q in zip(role_queries, filler_queries)], axis=1) return np.prod(sims, axis=1), sims def explicit_scores(objects, role_queries, filler_queries): # Build the order-m query and O_t tensor product, then contract exactly. # This is intentionally only used for the small algebraic sanity check. m = len(role_queries) vals = [] for obj in objects: flat = obj.reshape(-1) qs = [] for r, q in zip(role_queries, filler_queries): e = np.zeros_like(obj) e[r] = q qs.append(e.reshape(-1)) tensor = flat query = qs[0] for k in range(1, m): tensor = np.outer(tensor, flat).reshape(-1) query = np.outer(query, qs[k]).reshape(-1) vals.append(float(tensor @ query)) return np.asarray(vals) def make_problem(n, roles=4, dim=16, m=3, noise=0.08): # Each object contains one normalized filler per role. One target object # matches all query bindings; distractors match individual bindings only. fillers = rng.normal(size=(roles, dim)) fillers /= np.linalg.norm(fillers, axis=1, keepdims=True) target = rng.normal(size=(roles, dim)) target /= np.linalg.norm(target, axis=1, keepdims=True) role_queries = np.array([0, 1, 2][:m]) filler_queries = target[role_queries] objects = rng.normal(size=(n, roles, dim)) objects /= np.linalg.norm(objects, axis=2, keepdims=True) objects[0, role_queries] = filler_queries + noise * rng.normal(size=(m, dim)) objects[0, role_queries] /= np.linalg.norm(objects[0, role_queries], axis=1, keepdims=True) # Half of distractors contain one matching factor, making conjunction useful. for i in range(1, n): if i % 2 == 0: k = i % m objects[i, role_queries[k]] = filler_queries[k] + noise * rng.normal(size=dim) objects[i, role_queries[k]] /= np.linalg.norm(objects[i, role_queries[k]]) return objects, role_queries, filler_queries def benchmark(n, trials=300, m=3): hits_factor, hits_single, margins_factor, margins_single = [], [], [], [] t0 = time.perf_counter() for _ in range(trials): x, roles, qs = make_problem(n=n, m=m) sf, sims = factorized_scores(x, roles, qs) # Standard single-factor control: best one binding, equivalent to # ordinary attention that cannot enforce the conjunction. ss = sims[:, 0] order = np.argsort(sf)[::-1] order_s = np.argsort(ss)[::-1] hits_factor.append(int(order[0] == 0)) hits_single.append(int(order_s[0] == 0)) margins_factor.append(sf[order[0]] - sf[order[1]]) margins_single.append(ss[order_s[0]] - ss[order_s[1]]) elapsed = time.perf_counter() - t0 return { "n": n, "m": m, "trials": trials, "factor_accuracy": float(np.mean(hits_factor)), "single_accuracy": float(np.mean(hits_single)), "factor_margin": float(np.mean(margins_factor)), "single_margin": float(np.mean(margins_single)), "seconds": elapsed, # Explicit order-m tensor has N*D^(2m) scalar scale for order-2 objects; # factorized retrieval stores objects and m query vectors: linear in N. "factor_memory_units": int(n * 4 * 16), "explicit_memory_units_relative": int(n * (4 * 16) ** m), } def main(): # Algebraic verification with signed values and multiple orders. max_err = 0.0 for m in (2, 3): x = rng.normal(size=(4, 3, 3)) roles = np.arange(m) qs = rng.normal(size=(m, 3)) a = factorized_scores(x, roles, qs)[0] b = explicit_scores(x, roles, qs) max_err = max(max_err, float(np.max(np.abs(a - b)))) results = {"identity_max_abs_error": max_err} for n in (32, 128, 256): results[f"n{n}_m2"] = benchmark(n, m=2) results[f"n{n}_m3"] = benchmark(n, m=3) # Verify output extraction is linear and soft retrieval is normalized. x, roles, qs = make_problem(64, m=3) scores, _ = factorized_scores(x, roles, qs) attn = softmax(scores / 0.1) out = np.einsum("n,nrd->rd", attn, x) results["attention_sum"] = float(attn.sum()) results["output_shape"] = list(out.shape) Path("results.json").write_text(json.dumps(results, indent=2)) print(json.dumps(results, indent=2)) if __name__ == "__main__": main()