tmnam20/code-summarization
1
1# Copyright (c) Microsoft Corporation.2# Licensed under the MIT license.3 4import torch5import torch.nn as nn6import torch7from torch.autograd import Variable8import copy9 10 11class Seq2Seq(nn.Module):12 """13 Build Seqence-to-Sequence.14 15 Parameters:16 17 * `encoder`- encoder of seq2seq model. e.g. roberta18 * `decoder`- decoder of seq2seq model. e.g. transformer19 * `config`- configuration of encoder model.20 * `beam_size`- beam size for beam search.21 * `max_length`- max length of target for beam search.22 * `sos_id`- start of symbol ids in target for beam search.23 * `eos_id`- end of symbol ids in target for beam search.24 """25 26 def __init__(27 self,28 encoder,29 decoder,30 config,31 beam_size=None,32 max_length=None,33 sos_id=None,34 eos_id=None,35 ):36 super(Seq2Seq, self).__init__()37 self.encoder = encoder38 self.decoder = decoder39 self.config = config40 self.register_buffer("bias", torch.tril(torch.ones(2048, 2048)))41 self.dense = nn.Linear(config.hidden_size, config.hidden_size)42 self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)43 self.lsm = nn.LogSoftmax(dim=-1)44 self.tie_weights()45 46 self.beam_size = beam_size47 self.max_length = max_length48 self.sos_id = sos_id49 self.eos_id = eos_id50 51 def _tie_or_clone_weights(self, first_module, second_module):52 """Tie or clone module weights depending of weither we are using TorchScript or not"""53 if self.config.torchscript:54 first_module.weight = nn.Parameter(second_module.weight.clone())55 else:56 first_module.weight = second_module.weight57 58 def tie_weights(self):59 """Make sure we are sharing the input and output embeddings.60 Export to TorchScript can't handle parameter sharing so we are cloning them instead.61 """62 self._tie_or_clone_weights(63 self.lm_head, self.encoder.embeddings.word_embeddings64 )65 66 def forward(67 self,68 source_ids=None,69 source_mask=None,70 target_ids=None,71 target_mask=None,72 args=None,73 ):74 outputs = self.encoder(source_ids, attention_mask=source_mask)75 encoder_output = outputs[0].permute([1, 0, 2]).contiguous()76 if target_ids is not None:77 attn_mask = -1e4 * (78 1 - self.bias[: target_ids.shape[1], : target_ids.shape[1]]79 )80 tgt_embeddings = (81 self.encoder.embeddings(target_ids).permute([1, 0, 2]).contiguous()82 )83 out = self.decoder(84 tgt_embeddings,85 encoder_output,86 tgt_mask=attn_mask,87 memory_key_padding_mask=(1 - source_mask).bool(),88 )89 hidden_states = torch.tanh(self.dense(out)).permute([1, 0, 2]).contiguous()90 lm_logits = self.lm_head(hidden_states)91 # Shift so that tokens < n predict n92 active_loss = target_mask[..., 1:].ne(0).view(-1) == 193 shift_logits = lm_logits[..., :-1, :].contiguous()94 shift_labels = target_ids[..., 1:].contiguous()95 # Flatten the tokens96 loss_fct = nn.CrossEntropyLoss(ignore_index=-1)97 loss = loss_fct(98 shift_logits.view(-1, shift_logits.size(-1))[active_loss],99 shift_labels.view(-1)[active_loss],100 )101 102 outputs = loss, loss * active_loss.sum(), active_loss.sum()103 return outputs104 else:105 # Predict106 preds = []107 try:108 zero = torch.cuda.LongTensor(1).fill_(0)109 except Exception as e:110 zero = torch.LongTensor(1).fill_(0)111 for i in range(source_ids.shape[0]):112 context = encoder_output[:, i : i + 1]113 context_mask = source_mask[i : i + 1, :]114 beam = Beam(self.beam_size, self.sos_id, self.eos_id)115 input_ids = beam.getCurrentState()116 context = context.repeat(1, self.beam_size, 1)117 context_mask = context_mask.repeat(self.beam_size, 1)118 for _ in range(self.max_length):119 if beam.done():120 break121 attn_mask = -1e4 * (122 1 - self.bias[: input_ids.shape[1], : input_ids.shape[1]]123 )124 tgt_embeddings = (125 self.encoder.embeddings(input_ids)126 .permute([1, 0, 2])127 .contiguous()128 )129 out = self.decoder(130 tgt_embeddings,131 context,132 tgt_mask=attn_mask,133 memory_key_padding_mask=(1 - context_mask).bool(),134 )135 out = torch.tanh(self.dense(out))136 hidden_states = out.permute([1, 0, 2]).contiguous()[:, -1, :]137 out = self.lsm(self.lm_head(hidden_states)).data138 beam.advance(out)139 input_ids.data.copy_(140 input_ids.data.index_select(0, beam.getCurrentOrigin())141 )142 input_ids = torch.cat((input_ids, beam.getCurrentState()), -1)143 hyp = beam.getHyp(beam.getFinal())144 pred = beam.buildTargetTokens(hyp)[: self.beam_size]145 pred = [146 torch.cat(147 [x.view(-1) for x in p] + [zero] * (self.max_length - len(p))148 ).view(1, -1)149 for p in pred150 ]151 preds.append(torch.cat(pred, 0).unsqueeze(0))152 153 preds = torch.cat(preds, 0)154 return preds155 156 157class Beam(object):158 def __init__(self, size, sos, eos):159 self.size = size160 if torch.cuda.is_available():161 self.tt = torch.cuda162 else:163 self.tt = torch164 # The score for each translation on the beam.165 self.scores = self.tt.FloatTensor(size).zero_()166 # The backpointers at each time-step.167 self.prevKs = []168 # The outputs at each time-step.169 self.nextYs = [self.tt.LongTensor(size).fill_(0)]170 self.nextYs[0][0] = sos171 # Has EOS topped the beam yet.172 self._eos = eos173 self.eosTop = False174 # Time and k pair for finished.175 self.finished = []176 177 def getCurrentState(self):178 "Get the outputs for the current timestep."179 batch = self.tt.LongTensor(self.nextYs[-1]).view(-1, 1)180 return batch181 182 def getCurrentOrigin(self):183 "Get the backpointers for the current timestep."184 return self.prevKs[-1]185 186 def advance(self, wordLk):187 """188 Given prob over words for every last beam `wordLk` and attention189 `attnOut`: Compute and update the beam search.190 191 Parameters:192 193 * `wordLk`- probs of advancing from the last step (K x words)194 * `attnOut`- attention at the last step195 196 Returns: True if beam search is complete.197 """198 numWords = wordLk.size(1)199 200 # Sum the previous scores.201 if len(self.prevKs) > 0:202 beamLk = wordLk + self.scores.unsqueeze(1).expand_as(wordLk)203 204 # Don't let EOS have children.205 for i in range(self.nextYs[-1].size(0)):206 if self.nextYs[-1][i] == self._eos:207 beamLk[i] = -1e20208 else:209 beamLk = wordLk[0]210 flatBeamLk = beamLk.view(-1)211 bestScores, bestScoresId = flatBeamLk.topk(self.size, 0, True, True)212 213 self.scores = bestScores214 215 # bestScoresId is flattened beam x word array, so calculate which216 # word and beam each score came from217 prevK = bestScoresId // numWords218 self.prevKs.append(prevK)219 self.nextYs.append((bestScoresId - prevK * numWords))220 221 for i in range(self.nextYs[-1].size(0)):222 if self.nextYs[-1][i] == self._eos:223 s = self.scores[i]224 self.finished.append((s, len(self.nextYs) - 1, i))225 226 # End condition is when top-of-beam is EOS and no global score.227 if self.nextYs[-1][0] == self._eos:228 self.eosTop = True229 230 def done(self):231 return self.eosTop and len(self.finished) >= self.size232 233 def getFinal(self):234 if len(self.finished) == 0:235 self.finished.append((self.scores[0], len(self.nextYs) - 1, 0))236 self.finished.sort(key=lambda a: -a[0])237 if len(self.finished) != self.size:238 unfinished = []239 for i in range(self.nextYs[-1].size(0)):240 if self.nextYs[-1][i] != self._eos:241 s = self.scores[i]242 unfinished.append((s, len(self.nextYs) - 1, i))243 unfinished.sort(key=lambda a: -a[0])244 self.finished += unfinished[: self.size - len(self.finished)]245 return self.finished[: self.size]246 247 def getHyp(self, beam_res):248 """249 Walk back to construct the full hypothesis.250 """251 hyps = []252 for _, timestep, k in beam_res:253 hyp = []254 for j in range(len(self.prevKs[:timestep]) - 1, -1, -1):255 hyp.append(self.nextYs[j + 1][k])256 k = self.prevKs[j][k]257 hyps.append(hyp[::-1])258 return hyps259 260 def buildTargetTokens(self, preds):261 sentence = []262 for pred in preds:263 tokens = []264 for tok in pred:265 if tok == self._eos:266 break267 tokens.append(tok)268 sentence.append(tokens)269 return sentence270 