import numpy as np
from numpy_core import MultiHead

rng = np.random.default_rng(4)
query = rng.normal(size=(2, 4, 8))
source = rng.normal(size=(2, 6, 8))
layer = MultiHead(width=8, heads=2)
output, weights = layer(query, source)

assert output.shape == (2, 4, 8)
assert weights.shape == (2, 2, 4, 6)
np.testing.assert_allclose(weights.sum(-1), 1)
