(query, key, value, attn_bias=None)
| 3 | |
| 4 | |
| 5 | def low_version_attention(query, key, value, attn_bias=None): |
| 6 | scale = 1 / query.shape[-1] ** 0.5 |
| 7 | query = query * scale |
| 8 | attn = torch.matmul(query, key.transpose(-2, -1)) |
| 9 | if attn_bias is not None: |
| 10 | attn = attn + attn_bias |
| 11 | attn = attn.softmax(-1) |
| 12 | return attn @ value |
| 13 | |
| 14 | |
| 15 | class Attention(torch.nn.Module): |