d_k, d_v, T = 8, 8, 64
kk = jax.random.split(jax.random.key(2), 5)
q, k, v = (jax.random.normal(key, (T, d))
for key, d in zip(kk, (d_k, d_k, d_v)))
i_pre = jax.random.normal(kk[3], (T,)) # Gate pre-activations
f_pre = 2.0 + jax.random.normal(kk[4], (T,))
def mlstm_naive(q, k, v, i_pre, f_pre, dtype):
"""The unstabilized recurrence (NumPy), in a given precision."""
q, k, v = (np.asarray(x, dtype) for x in (q, k, v))
i_g, f_g = np.exp(np.asarray(i_pre, dtype)), np.exp(np.asarray(f_pre, dtype))
S, z = np.zeros((d_k, d_v), dtype), np.zeros(d_k, dtype)
outputs = []
for t in range(T):
S = f_g[t] * S + i_g[t] * np.outer(k[t], v[t])
z = f_g[t] * z + i_g[t] * k[t]
outputs.append(q[t] @ S / max(abs(q[t] @ z), 1.0))
return np.stack(outputs)
def mlstm_stabilized(q, k, v, i_pre, f_pre):
"""Carry m_t and rescale the past, as in online softmax."""
def step(carry, x):
S, z, m = carry
q_t, k_t, v_t, i_t, f_t = x
m_new = jnp.maximum(f_t + m, i_t)
f_g, i_g = jnp.exp(f_t + m - m_new), jnp.exp(i_t - m_new)
S = f_g * S + i_g * (k_t[:, None] * v_t[None, :])
z = f_g * z + i_g * k_t
o = q_t @ S / jnp.maximum(jnp.abs(q_t @ z), jnp.exp(-m_new))
return (S, z, m_new), o
init = (jnp.zeros((d_k, d_v)), jnp.zeros(d_k), -jnp.inf)
return jax.lax.scan(step, init, (q, k, v, i_pre, f_pre))[1]
exact = mlstm_naive(q, k, v, i_pre, f_pre, np.float64)
with np.errstate(over='ignore', invalid='ignore'): # The overflow is the point
naive = mlstm_naive(q, k, v, i_pre, f_pre, np.float32)
bad = ~np.isfinite(naive).all(-1) # argmax(bool) would report 0
first_bad = int(bad.argmax()) if bad.any() else 'none'
stab = mlstm_stabilized(q, k, v, i_pre, f_pre)
err = jnp.abs(stab - exact).max() / jnp.abs(exact).max()
print(f'float32 unstabilized: first non-finite output at step {first_bad}')
print(f'float32 stabilized vs float64: relative deviation {float(err):.2e}')
assert bool(jnp.isfinite(stab).all()) and err < 1e-3