def delta_recurrent(Q, K, V, beta):
"""The delta rule in the error form: read, subtract, write, read out."""
def step(S, qkvb):
q, k, v, b = qkvb
error = v - S.T @ k # What the memory got wrong
S = S + b * k[:, None] * error[None, :]
return S, S.T @ q
S0 = jnp.zeros((Q.shape[-1], V.shape[-1]))
return jax.lax.scan(step, S0, (Q, K, V, beta))[1]
def delta_recurrent_matrix(Q, K, V, beta):
"""The same update as transition-then-write, for the family template."""
def step(S, qkvb):
q, k, v, b = qkvb
S = (jnp.eye(len(k)) - b * jnp.outer(k, k)) @ S + b * jnp.outer(k, v)
return S, S.T @ q
S0 = jnp.zeros((Q.shape[-1], V.shape[-1]))
return jax.lax.scan(step, S0, (Q, K, V, beta))[1]
d_k, d_v, T = 64, 64, 512
rng_keys = jax.random.split(jax.random.key(0), 4)
Q, K, V = (jax.random.normal(k, (T, d))
for k, d in zip(rng_keys, (d_k, d_k, d_v)))
K = K / jnp.maximum(jnp.linalg.norm(K, axis=-1, keepdims=True),
1e-12) # Unit keys; F.normalize contract
beta = jax.nn.sigmoid(jax.random.normal(rng_keys[3], (T,)))
with jax.default_matmul_precision('highest'):
y_delta = delta_recurrent(Q, K, V, beta)
err = jnp.abs(y_delta - delta_recurrent_matrix(Q, K, V, beta)).max() \
/ jnp.abs(y_delta).max()
print(f'error form vs matrix form: relative deviation {float(err):.2e}')
assert err < 1e-5 # One rule, two readings