Team Ai
Apppublic

tmnam20/code-summarization

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
model.py270 linesDownload Raw Back to root
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