MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedASR / TransformerDecoder

Class TransformerDecoder

fireredasr/models/module/transformer_decoder.py:9–169  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class 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

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected