r""" Args: query, key, value: map a query and a set of key-value pairs to an output. See "Attention Is All You Need" for more details. key_padding_mask: if provided, specified padding elements in the key will be ignored by the attention. When given a binar
(self, query: Tensor, key: Tensor, value: Tensor, key_padding_mask: Optional[Tensor] = None,
need_weights: bool = True, attn_mask: Optional[Tensor] = None)
| 424 | super(MultiheadAttention, self).__setstate__(state) |
| 425 | |
| 426 | def forward(self, query: Tensor, key: Tensor, value: Tensor, key_padding_mask: Optional[Tensor] = None, |
| 427 | need_weights: bool = True, attn_mask: Optional[Tensor] = None) -> Tuple[Tensor, Optional[Tensor]]: |
| 428 | r""" |
| 429 | Args: |
| 430 | query, key, value: map a query and a set of key-value pairs to an output. |
| 431 | See "Attention Is All You Need" for more details. |
| 432 | key_padding_mask: if provided, specified padding elements in the key will |
| 433 | be ignored by the attention. When given a binary mask and a value is True, |
| 434 | the corresponding value on the attention layer will be ignored. When given |
| 435 | a byte mask and a value is non-zero, the corresponding value on the attention |
| 436 | layer will be ignored |
| 437 | need_weights: output attn_output_weights. |
| 438 | attn_mask: 2D or 3D mask that prevents attention to certain positions. A 2D mask will be broadcasted for all |
| 439 | the batches while a 3D mask allows to specify a different mask for the entries of each batch. |
| 440 | |
| 441 | Shapes for inputs: |
| 442 | - query: :math:`(L, N, E)` where L is the target sequence length, N is the batch size, E is |
| 443 | the embedding dimension. |
| 444 | - key: :math:`(S, N, E)`, where S is the source sequence length, N is the batch size, E is |
| 445 | the embedding dimension. |
| 446 | - value: :math:`(S, N, E)` where S is the source sequence length, N is the batch size, E is |
| 447 | the embedding dimension. |
| 448 | - key_padding_mask: :math:`(N, S)` where N is the batch size, S is the source sequence length. |
| 449 | If a ByteTensor is provided, the non-zero positions will be ignored while the position |
| 450 | with the zero positions will be unchanged. If a BoolTensor is provided, the positions with the |
| 451 | value of ``True`` will be ignored while the position with the value of ``False`` will be unchanged. |
| 452 | - attn_mask: if a 2D mask: :math:`(L, S)` where L is the target sequence length, S is the |
| 453 | source sequence length. |
| 454 | |
| 455 | If a 3D mask: :math:`(N\cdot\text{num\_heads}, L, S)` where N is the batch size, L is the target sequence |
| 456 | length, S is the source sequence length. ``attn_mask`` ensure that position i is allowed to attend |
| 457 | the unmasked positions. If a ByteTensor is provided, the non-zero positions are not allowed to attend |
| 458 | while the zero positions will be unchanged. If a BoolTensor is provided, positions with ``True`` |
| 459 | is not allowed to attend while ``False`` values will be unchanged. If a FloatTensor |
| 460 | is provided, it will be added to the attention weight. |
| 461 | |
| 462 | Shapes for outputs: |
| 463 | - attn_output: :math:`(L, N, E)` where L is the target sequence length, N is the batch size, |
| 464 | E is the embedding dimension. |
| 465 | - attn_output_weights: :math:`(N, L, S)` where N is the batch size, |
| 466 | L is the target sequence length, S is the source sequence length. |
| 467 | """ |
| 468 | if not self._qkv_same_embed_dim: |
| 469 | return multi_head_attention_forward( |
| 470 | query, key, value, self.embed_dim, self.num_heads, |
| 471 | self.in_proj_weight, self.in_proj_bias, |
| 472 | self.bias_k, self.bias_v, self.add_zero_attn, |
| 473 | self.dropout, self.out_proj.weight, self.out_proj.bias, |
| 474 | training=self.training, |
| 475 | key_padding_mask=key_padding_mask, need_weights=need_weights, |
| 476 | attn_mask=attn_mask, use_separate_proj_weight=True, |
| 477 | q_proj_weight=self.q_proj_weight, k_proj_weight=self.k_proj_weight, |
| 478 | v_proj_weight=self.v_proj_weight) |
| 479 | else: |
| 480 | return multi_head_attention_forward( |
| 481 | query, key, value, self.embed_dim, self.num_heads, |
| 482 | self.in_proj_weight, self.in_proj_bias, |
| 483 | self.bias_k, self.bias_v, self.add_zero_attn, |
nothing calls this directly
no test coverage detected