| 33 | |
| 34 | |
| 35 | class TransformerDecoderLayerOptimal(nn.Module): |
| 36 | def __init__(self, d_model, nhead=8, dim_feedforward=2048, dropout=0.1, activation="relu", |
| 37 | layer_norm_eps=1e-5) -> None: |
| 38 | super(TransformerDecoderLayerOptimal, self).__init__() |
| 39 | self.norm1 = nn.LayerNorm(d_model, eps=layer_norm_eps) |
| 40 | self.dropout = nn.Dropout(dropout) |
| 41 | self.dropout1 = nn.Dropout(dropout) |
| 42 | self.dropout2 = nn.Dropout(dropout) |
| 43 | self.dropout3 = nn.Dropout(dropout) |
| 44 | |
| 45 | self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout) |
| 46 | |
| 47 | # Implementation of Feedforward model |
| 48 | self.linear1 = nn.Linear(d_model, dim_feedforward) |
| 49 | self.linear2 = nn.Linear(dim_feedforward, d_model) |
| 50 | |
| 51 | self.norm2 = nn.LayerNorm(d_model, eps=layer_norm_eps) |
| 52 | self.norm3 = nn.LayerNorm(d_model, eps=layer_norm_eps) |
| 53 | |
| 54 | self.activation = _get_activation_fn(activation) |
| 55 | |
| 56 | def __setstate__(self, state): |
| 57 | if 'activation' not in state: |
| 58 | state['activation'] = torch.nn.functional.relu |
| 59 | super(TransformerDecoderLayerOptimal, self).__setstate__(state) |
| 60 | |
| 61 | def forward(self, tgt: Tensor, memory: Tensor, tgt_mask: Optional[Tensor] = None, |
| 62 | memory_mask: Optional[Tensor] = None, |
| 63 | tgt_key_padding_mask: Optional[Tensor] = None, |
| 64 | memory_key_padding_mask: Optional[Tensor] = None) -> Tensor: |
| 65 | tgt = tgt + self.dropout1(tgt) |
| 66 | tgt = self.norm1(tgt) |
| 67 | tgt2 = self.multihead_attn(tgt, memory, memory)[0] |
| 68 | tgt = tgt + self.dropout2(tgt2) |
| 69 | tgt = self.norm2(tgt) |
| 70 | tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt)))) |
| 71 | tgt = tgt + self.dropout3(tgt2) |
| 72 | tgt = self.norm3(tgt) |
| 73 | return tgt |
| 74 | |
| 75 | |
| 76 | # @torch.jit.script |