def scaled_dot_product_attention(q, k, v, key_mask):
"""q: [B,H,T,Dk]; k: [B,H,S,Dk]; v: [B,H,S,Dv]."""
logits = jnp.einsum("bhtd,bhsd->bhts", q, k)
logits = logits / jnp.sqrt(q.shape[-1])
weights = masked_softmax(logits, key_mask)
output = jnp.einsum("bhts,bhse->bhte", weights, v)
return output, weights