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