Role-Filler Attention / role_filler_experiment.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 2951
  6rng = np.random.default_rng(SEED)
  7
  8
  9def softmax(x, axis=-1):
 10    z = x - np.max(x, axis=axis, keepdims=True)
 11    e = np.exp(z)
 12    return e / e.sum(axis=axis, keepdims=True)
 13
 14
 15def entropy(p):
 16    return float(-(p * np.log(np.maximum(p, 1e-12))).sum(axis=-1).mean())
 17
 18
 19def math_check(R=5, D=7, N=9):
 20    # QR gives exactly orthonormal role vectors (up to numerical precision).
 21    Q, _ = np.linalg.qr(rng.normal(size=(D, R)))
 22    roles = Q.T                         # [R,D], rows orthonormal
 23    F = rng.normal(size=(N, R, D))
 24    O = np.einsum('rd,nre->nde', roles, F)  # [N,D,D]
 25    recovered = np.einsum('rd,nde->nre', roles, O)
 26    contraction_error = float(np.max(np.abs(recovered - F)))
 27
 28    # Query role m and filler from object 3 should retrieve object 3.
 29    m, target, source = 1, 3, 3
 30    qf = F[source, m]
 31    scores = F[:, m] @ qf
 32    probs = softmax(scores / 0.15)
 33    top = int(np.argmax(probs))
 34    out = probs @ F[:, target]
 35    retrieval_error = float(np.linalg.norm(out - F[source, target]))
 36
 37    # Wrong-role query is a control: it should not be systematically selective.
 38    wrong = (m + 1) % R
 39    wrong_probs = softmax((F[:, wrong] @ qf) / 0.15)
 40    return {
 41        'max_role_contraction_abs_error': contraction_error,
 42        'exact_query_top1': top == source,
 43        'exact_query_probability': float(probs[source]),
 44        'exact_rebinding_l2_error': retrieval_error,
 45        'exact_entropy': entropy(probs[None, :]),
 46        'mismatched_role_entropy': entropy(wrong_probs[None, :]),
 47    }
 48
 49
 50def make_episode(N=16, R=4, D=16, noise=0.08):
 51    # Each object has independent role fillers. The query is one role filler;
 52    # the desired answer is another role filler from that same object.
 53    F = rng.normal(size=(N, R, D)).astype(np.float32)
 54    src = int(rng.integers(N))
 55    m, target = 0, 1
 56    q = F[src, m] + noise * rng.normal(size=D).astype(np.float32)
 57    return F, q, src, m, target
 58
 59
 60def structured_trial(F, q, m, target, tau=0.25):
 61    s = F[:, m] @ q
 62    a = softmax(s / tau)
 63    y = a @ F[:, target]
 64    return int(np.argmax(s)), y, a
 65
 66
 67def dense_trial(F, q, target, tau=0.25):
 68    # Control: a standard flattened-object dot product, with the query copied
 69    # into every role block so it has the same input dimension as an object.
 70    flat = F.reshape(F.shape[0], -1)
 71    qflat = np.tile(q, F.shape[1])
 72    s = flat @ qflat
 73    a = softmax(s / tau)
 74    return int(np.argmax(s)), a @ F[:, target], a
 75
 76
 77def benchmark(episodes=2000):
 78    stats = {k: [] for k in ['structured_acc','dense_acc','structured_err','dense_err',
 79                             'structured_entropy','dense_entropy','wrong_entropy']}
 80    for _ in range(episodes):
 81        F, q, src, m, target = make_episode()
 82        si, sy, sa = structured_trial(F, q, m, target)
 83        di, dy, da = dense_trial(F, q, target)
 84        wrong = (m + 1) % F.shape[1]
 85        _, _, wa = structured_trial(F, q, wrong, target)
 86        stats['structured_acc'].append(si == src)
 87        stats['dense_acc'].append(di == src)
 88        stats['structured_err'].append(np.linalg.norm(sy - F[src,target]))
 89        stats['dense_err'].append(np.linalg.norm(dy - F[src,target]))
 90        stats['structured_entropy'].append(entropy(sa[None,:]))
 91        stats['dense_entropy'].append(entropy(da[None,:]))
 92        stats['wrong_entropy'].append(entropy(wa[None,:]))
 93    return {k: float(np.mean(v)) for k,v in stats.items()}
 94
 95
 96def main():
 97    np.set_printoptions(precision=6, suppress=True)
 98    result = {'seed': SEED, 'math_check': math_check(), 'benchmark': benchmark()}
 99    # Repeat the cheap benchmark with a different seed to check direction stability.
100    global rng
101    rng = np.random.default_rng(SEED + 1)
102    result['benchmark_repeat'] = benchmark(episodes=1000)
103    out = Path('results.json')
104    out.write_text(json.dumps(result, indent=2))
105    print(json.dumps(result, indent=2))
106
107if __name__ == '__main__':
108    main()