def multi_head_attention(
    x, key_mask, w_q, w_k, w_v, w_o, n_heads
):
    """x: [B,T,Dmodel]."""
    b, t, _ = x.shape
    d_k = w_q.shape[-1] // n_heads
    d_v = w_v.shape[-1] // n_heads

    q = x @ w_q
    k = x @ w_k
    v = x @ w_v

    q = q.reshape(b, t, n_heads, d_k).transpose(0, 2, 1, 3)
    k = k.reshape(b, t, n_heads, d_k).transpose(0, 2, 1, 3)
    v = v.reshape(b, t, n_heads, d_v).transpose(0, 2, 1, 3)

    head_output, weights = scaled_dot_product_attention(
        q, k, v, key_mask
    )

    merged = head_output.transpose(0, 2, 1, 3)
    merged = merged.reshape(b, t, n_heads * d_v)
    output = merged @ w_o
    return output, weights