| 29 | |
| 30 | |
| 31 | class ARDecoder(nn.Module): |
| 32 | |
| 33 | def __init__( |
| 34 | self, |
| 35 | in_channels, |
| 36 | out_channels, |
| 37 | nhead=None, |
| 38 | num_decoder_layers=6, |
| 39 | max_len=25, |
| 40 | attention_dropout_rate=0.0, |
| 41 | residual_dropout_rate=0.1, |
| 42 | scale_embedding=True, |
| 43 | ): |
| 44 | super(ARDecoder, self).__init__() |
| 45 | self.out_channels = out_channels |
| 46 | self.ignore_index = out_channels - 1 |
| 47 | self.bos = out_channels - 2 |
| 48 | self.eos = 0 |
| 49 | self.max_len = max_len |
| 50 | d_model = in_channels |
| 51 | dim_feedforward = d_model * 4 |
| 52 | nhead = nhead if nhead is not None else d_model // 32 |
| 53 | self.embedding = Embeddings( |
| 54 | d_model=d_model, |
| 55 | vocab=self.out_channels, |
| 56 | padding_idx=0, |
| 57 | scale_embedding=scale_embedding, |
| 58 | ) |
| 59 | self.pos_embed = nn.Parameter(torch.zeros([1, max_len + 1, d_model], |
| 60 | dtype=torch.float32), |
| 61 | requires_grad=True) |
| 62 | trunc_normal_(self.pos_embed, std=0.02) |
| 63 | self.decoder = nn.ModuleList([ |
| 64 | TransformerBlock( |
| 65 | d_model, |
| 66 | nhead, |
| 67 | dim_feedforward, |
| 68 | attention_dropout_rate, |
| 69 | residual_dropout_rate, |
| 70 | with_self_attn=True, |
| 71 | with_cross_attn=False, |
| 72 | ) for i in range(num_decoder_layers) |
| 73 | ]) |
| 74 | |
| 75 | self.tgt_word_prj = nn.Linear(d_model, |
| 76 | self.out_channels - 2, |
| 77 | bias=False) |
| 78 | self.apply(self._init_weights) |
| 79 | |
| 80 | def _init_weights(self, m): |
| 81 | if isinstance(m, nn.Linear): |
| 82 | nn.init.xavier_normal_(m.weight) |
| 83 | if m.bias is not None: |
| 84 | nn.init.zeros_(m.bias) |
| 85 | |
| 86 | def forward_train(self, src, tgt): |
| 87 | tgt = tgt[:, :-1] |
| 88 | |