THE CODE READER / PYTHON
Download raw .py ↓example-8d592b02a427.py
Your snippet, in context. Explore the file or visualize its recorded example.
1import numpy as np
2from numpy.typing import NDArray
3
4
5def phi(x: NDArray[np.float64]) -> NDArray[np.float64]:
6 # exp(min(x, 0)) avoids overflowing the unused branch.
7 return np.where(x >= 0, x + 1, np.exp(np.minimum(x, 0)))
8
9
10def linear_attention(q: NDArray[np.float64], k: NDArray[np.float64], v: NDArray[np.float64]) -> NDArray[np.float64]:
11 qf, kf = phi(q), phi(k)
12 state = kf.T @ v # (features, value_dim)
13 normalizer = kf.sum(axis=0) # (features,)
14 denominator = qf @ normalizer
15 return (qf @ state) / denominator[:, None].clip(1e-8)
16
17
18rng = np.random.default_rng(7)
19q, k = rng.normal(size=(2, 12, 8))
20v = rng.normal(size=(12, 6))
21
22pairwise = phi(q) @ phi(k).T
23reference = (pairwise @ v) / pairwise.sum(-1, keepdims=True)
24output = linear_attention(q, k, v)
25
26np.testing.assert_allclose(output, reference, atol=1e-10)
27print(output.shape) # (12, 6)