def masked_softmax(logits, key_mask):
"""Softmax attention logits over keys.
logits: [B,H,Q,K]; key_mask: [B,K].
Every row must contain at least one valid key.
"""
expanded_mask = key_mask[:, None, None, :]
masked_logits = jnp.where(
expanded_mask, logits, -jnp.inf
)
return stable_softmax(masked_logits, axis=-1)