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)