Implements Figure 2
(self, query, key, value, mask=None)
| 36 | self.dropout = torch.nn.Dropout(p=dropout) |
| 37 | |
| 38 | def forward(self, query, key, value, mask=None): |
| 39 | "Implements Figure 2" |
| 40 | if mask is not None: |
| 41 | # Same mask applied to all h heads. |
| 42 | mask = mask.unsqueeze(1) |
| 43 | nbatches = query.size(0) |
| 44 | |
| 45 | # 1) Do all the linear projections in batch from d_model => h x d_k |
| 46 | query, key, value = \ |
| 47 | [l(x).view(nbatches, -1, self.h, self.d_k).transpose(1, 2) |
| 48 | for l, x in zip(self.linears, (query, key, value))] |
| 49 | |
| 50 | # 2) Apply attention on all the projected vectors in batch. |
| 51 | x, self.attn = attention(query, key, value, mask=mask, |
| 52 | dropout=self.dropout) |
| 53 | |
| 54 | # 3) "Concat" using a view and apply a final linear. |
| 55 | x = x.transpose(1, 2).contiguous() \ |
| 56 | .view(nbatches, -1, self.h * self.d_k) |
| 57 | return self.linears[-1](x) |
| 58 | |
| 59 | |
| 60 | class SelfAttention(torch.nn.Module): |
nothing calls this directly
no test coverage detected