import numpy as np
from numpy.typing import NDArray


def phi(x: NDArray[np.float64]) -> NDArray[np.float64]:
    # exp(min(x, 0)) avoids overflowing the unused branch.
    return np.where(x >= 0, x + 1, np.exp(np.minimum(x, 0)))


def linear_attention(q: NDArray[np.float64], k: NDArray[np.float64], v: NDArray[np.float64]) -> NDArray[np.float64]:
    qf, kf = phi(q), phi(k)
    state = kf.T @ v               # (features, value_dim)
    normalizer = kf.sum(axis=0)    # (features,)
    denominator = qf @ normalizer
    return (qf @ state) / denominator[:, None].clip(1e-8)


rng = np.random.default_rng(7)
q, k = rng.normal(size=(2, 12, 8))
v = rng.normal(size=(12, 6))

pairwise = phi(q) @ phi(k).T
reference = (pairwise @ v) / pairwise.sum(-1, keepdims=True)
output = linear_attention(q, k, v)

np.testing.assert_allclose(output, reference, atol=1e-10)
print(output.shape)  # (12, 6)
