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