torch.manual_seed(0)
d_k, d_v, T = 8, 8, 64
q, k, v = (torch.randn(T, d) for d in (d_k, d_k, d_v))
i_pre, f_pre = torch.randn(T), 2.0 + torch.randn(T) # Gate pre-activations
def mlstm_naive(q, k, v, i_pre, f_pre, dtype):
"""The unstabilized recurrence, in a given precision."""
q, k, v = q.to(dtype), k.to(dtype), v.to(dtype)
i_g, f_g = i_pre.to(dtype).exp(), f_pre.to(dtype).exp()
S, z = torch.zeros(d_k, d_v, dtype=dtype), torch.zeros(d_k, dtype=dtype)
outputs = []
for t in range(T):
S = f_g[t] * S + i_g[t] * (k[t][:, None] * v[t][None, :])
z = f_g[t] * z + i_g[t] * k[t]
outputs.append(q[t] @ S / (q[t] @ z).abs().clamp(min=1))
return torch.stack(outputs)
def mlstm_stabilized(q, k, v, i_pre, f_pre):
"""Carry m_t and rescale the past, as in online softmax."""
S, z = torch.zeros(d_k, d_v), torch.zeros(d_k)
m = torch.tensor(-torch.inf)
outputs = []
for t in range(T):
m_new = torch.maximum(f_pre[t] + m, i_pre[t])
f_g, i_g = torch.exp(f_pre[t] + m - m_new), torch.exp(i_pre[t] - m_new)
S = f_g * S + i_g * (k[t][:, None] * v[t][None, :])
z = f_g * z + i_g * k[t]
m = m_new
outputs.append(q[t] @ S
/ torch.maximum((q[t] @ z).abs(), torch.exp(-m)))
return torch.stack(outputs)
exact = mlstm_naive(q, k, v, i_pre, f_pre, torch.float64)
naive = mlstm_naive(q, k, v, i_pre, f_pre, torch.float32)
bad = ~torch.isfinite(naive).all(-1) # argmax(bool) would report 0
first_bad = int(bad.float().argmax()) if bool(bad.any()) else 'none'
stab = mlstm_stabilized(q, k, v, i_pre, f_pre)
err = (stab - exact.float()).abs().max() / exact.abs().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 torch.isfinite(stab).all() and err < 1e-3