d_mem, h_mem, num_pairs = 32, 128, 16
ks = jax.random.split(jax.random.key(3), 4)
k_pairs = jax.random.normal(ks[0], (num_pairs, d_mem))
k_pairs /= jnp.linalg.norm(k_pairs, axis=1, keepdims=True)
v_pairs = jax.random.normal(ks[1], (num_pairs, d_mem))
v_pairs /= jnp.linalg.norm(v_pairs, axis=1, keepdims=True)
order = jax.random.permutation(ks[2], jnp.tile(jnp.arange(num_pairs), 4))
params = {'W1': jax.random.normal(ks[3], (d_mem, h_mem)), # Unit-var hidden
'b1': jnp.zeros(h_mem),
'W2': jnp.zeros((h_mem, d_mem)), # Start empty: retrieve 0
'b2': jnp.zeros(d_mem)}
def memory(params, k):
hidden = jax.nn.gelu(k @ params['W1'] + params['b1'])
return hidden @ params['W2'] + params['b2']
def recall_loss(params, k, v):
return ((memory(params, k) - v)**2).sum()
grad_fn = jax.grad(recall_loss)
def retrieval_mse(params):
return float(((memory(params, k_pairs) - v_pairs)**2).mean())
before = retrieval_mse(params)
theta, eta = 0.005, 0.5
def write(carry, t):
params, velocity = carry
g = grad_fn(params, k_pairs[t], v_pairs[t]) # Surprise
velocity = jax.tree.map(lambda u, gg: eta * u - theta * gg, velocity, g)
params = jax.tree.map(lambda p, u: p + u, params, velocity)
return (params, velocity), None
velocity = jax.tree.map(jnp.zeros_like, params)
(params, _), _ = jax.lax.scan(write, (params, velocity), order)
after = retrieval_mse(params)
print(f'retrieval MSE: empty memory {before:.4f} -> after the stream '
f'{after:.4f}')
assert after < before / 5