THE CODE READER / PYTHON

example-8d592b02a427.py

Your snippet, in context. Explore the file or visualize its recorded example.

Download raw .py ↓
Complete file
PYTHON / LINE NUMBERS
 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)