Compute 'Scaled Dot Product Attention
(query, key, value, mask=None, dropout=None)
| 6 | |
| 7 | |
| 8 | def attention(query, key, value, mask=None, dropout=None): |
| 9 | "Compute 'Scaled Dot Product Attention'" |
| 10 | d_k = query.size(-1) |
| 11 | scores = torch.matmul(query, key.transpose(-2, -1)) \ |
| 12 | / math.sqrt(d_k) |
| 13 | if mask is not None: |
| 14 | scores = scores.masked_fill(mask == 0, -1e9) |
| 15 | p_attn = F.softmax(scores, dim = -1) |
| 16 | if dropout is not None: |
| 17 | p_attn = dropout(p_attn) |
| 18 | return torch.matmul(p_attn, value), p_attn |
| 19 | |
| 20 | |
| 21 | def clones(module, N): |