| 76 | |
| 77 | |
| 78 | class TransformerEncoder(nn.Module): |
| 79 | def __init__(self, encoder_layer, num_layers, norm=None): |
| 80 | super().__init__() |
| 81 | self.layers = _get_clones(encoder_layer, num_layers) |
| 82 | self.num_layers = num_layers |
| 83 | self.norm = norm |
| 84 | |
| 85 | def forward( |
| 86 | self, |
| 87 | src, |
| 88 | mask: Optional[Tensor] = None, |
| 89 | src_key_padding_mask: Optional[Tensor] = None, |
| 90 | pos: Optional[Tensor] = None, |
| 91 | ): |
| 92 | output = src |
| 93 | |
| 94 | for layer in self.layers: |
| 95 | output = layer( |
| 96 | output, src_mask=mask, src_key_padding_mask=src_key_padding_mask, pos=pos |
| 97 | ) |
| 98 | |
| 99 | if self.norm is not None: |
| 100 | output = self.norm(output) |
| 101 | |
| 102 | return output |
| 103 | |
| 104 | |
| 105 | class TransformerDecoder(nn.Module): |