| 7 | |
| 8 | |
| 9 | class TransformerDecoder(nn.Module): |
| 10 | def __init__( |
| 11 | self, sos_id, eos_id, pad_id, odim, |
| 12 | n_layers, n_head, d_model, |
| 13 | residual_dropout=0.1, pe_maxlen=5000): |
| 14 | super().__init__() |
| 15 | self.INF = 1e10 |
| 16 | # parameters |
| 17 | self.pad_id = pad_id |
| 18 | self.sos_id = sos_id |
| 19 | self.eos_id = eos_id |
| 20 | self.n_layers = n_layers |
| 21 | |
| 22 | # Components |
| 23 | self.tgt_word_emb = nn.Embedding(odim, d_model, padding_idx=self.pad_id) |
| 24 | self.positional_encoding = PositionalEncoding(d_model, max_len=pe_maxlen) |
| 25 | self.dropout = nn.Dropout(residual_dropout) |
| 26 | |
| 27 | self.layer_stack = nn.ModuleList() |
| 28 | for l in range(n_layers): |
| 29 | block = DecoderLayer(d_model, n_head, residual_dropout) |
| 30 | self.layer_stack.append(block) |
| 31 | |
| 32 | self.tgt_word_prj = nn.Linear(d_model, odim, bias=False) |
| 33 | self.layer_norm_out = nn.LayerNorm(d_model) |
| 34 | |
| 35 | self.tgt_word_prj.weight = self.tgt_word_emb.weight |
| 36 | self.scale = (d_model ** 0.5) |
| 37 | |
| 38 | def batch_beam_search(self, encoder_outputs, src_masks, |
| 39 | beam_size=1, nbest=1, decode_max_len=0, |
| 40 | softmax_smoothing=1.0, length_penalty=0.0, eos_penalty=1.0): |
| 41 | B = beam_size |
| 42 | N, Ti, H = encoder_outputs.size() |
| 43 | device = encoder_outputs.device |
| 44 | maxlen = decode_max_len if decode_max_len > 0 else Ti |
| 45 | assert eos_penalty > 0.0 and eos_penalty <= 1.0 |
| 46 | |
| 47 | # Init |
| 48 | encoder_outputs = encoder_outputs.unsqueeze(1).repeat(1, B, 1, 1).view(N*B, Ti, H) |
| 49 | src_mask = src_masks.unsqueeze(1).repeat(1, B, 1, 1).view(N*B, -1, Ti) |
| 50 | ys = torch.ones(N*B, 1).fill_(self.sos_id).long().to(device) |
| 51 | caches: List[Optional[Tensor]] = [] |
| 52 | for _ in range(self.n_layers): |
| 53 | caches.append(None) |
| 54 | scores = torch.tensor([0.0] + [-self.INF]*(B-1)).float().to(device) |
| 55 | scores = scores.repeat(N).view(N*B, 1) |
| 56 | is_finished = torch.zeros_like(scores) |
| 57 | |
| 58 | # Autoregressive Prediction |
| 59 | for t in range(maxlen): |
| 60 | tgt_mask = self.ignored_target_position_is_0(ys, self.pad_id) |
| 61 | |
| 62 | dec_output = self.dropout( |
| 63 | self.tgt_word_emb(ys) * self.scale + |
| 64 | self.positional_encoding(ys)) |
| 65 | |
| 66 | i = 0 |