class MultiHeadAttention(nnx.Module):
"""Multi-head attention."""
def __init__(self, num_hiddens, num_heads, dropout, bias=False, rngs=None):
assert num_hiddens % num_heads == 0, 'heads must divide num_hiddens'
rngs = nnx.Rngs(params=0, dropout=1) if rngs is None else rngs
self.num_hiddens, self.num_heads = num_hiddens, num_heads
self.attention = d2l.DotProductAttention(dropout, rngs=rngs)
self.W_q = nnx.Linear(num_hiddens, num_hiddens, use_bias=bias,
rngs=rngs)
self.W_k = nnx.Linear(num_hiddens, num_hiddens, use_bias=bias,
rngs=rngs)
self.W_v = nnx.Linear(num_hiddens, num_hiddens, use_bias=bias,
rngs=rngs)
self.W_o = nnx.Linear(num_hiddens, num_hiddens, use_bias=bias,
rngs=rngs)
def __call__(self, queries, keys, values, valid_lens):
# Shape of queries, keys, or values:
# (batch_size, no. of queries or key-value pairs, num_hiddens)
# Shape of valid_lens: (batch_size,) or (batch_size, no. of queries)
queries = self.transpose_qkv(self.W_q(queries))
keys = self.transpose_qkv(self.W_k(keys))
values = self.transpose_qkv(self.W_v(values))
if valid_lens is not None:
# On axis 0, copy the first item (scalar or vector) for num_heads
# times, then copy the next item, and so on
valid_lens = jnp.repeat(valid_lens, self.num_heads, axis=0)
# Shape of output: (batch_size * num_heads, no. of queries,
# num_hiddens / num_heads)
output, attention_weights = self.attention(
queries, keys, values, valid_lens)
# Shape of output_concat: (batch_size, no. of queries, num_hiddens)
output_concat = self.transpose_output(output)
# NNX idiom: return (output, weights); PyTorch returns output only
return self.W_o(output_concat), attention_weights