Team Ai
Apppublic

p4vv37/CodeBERT_CodeReviewer

sourceHugging Facemitupdated 2y agoView on Hugging Face
13likes
app.py289 linesDownload Raw Back to root
1import gradio as gr2import requests3from torch import nn4from torch.nn import CrossEntropyLoss5from transformers import AutoTokenizer, T5ForConditionalGeneration, AutoModelForSeq2SeqLM, T5Config6import torch7 8MAX_SOURCE_LENGTH = 5129 10 11class ReviewerModel(T5ForConditionalGeneration):12 13    def __init__(self, config):14        super().__init__(config)15        self.cls_head = nn.Linear(self.config.d_model, 2, bias=True)16        self.init()17 18    def init(self):19        nn.init.xavier_uniform_(self.lm_head.weight)20        factor = self.config.initializer_factor21        self.cls_head.weight.data.normal_(mean=0.0, \22                                          std=factor * ((self.config.d_model) ** -0.5))23        self.cls_head.bias.data.zero_()24 25    def forward(26            self, *argv, **kwargs27    ):28        r"""29        Doc from Huggingface transformers:30        labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size,)`, `optional`):31            Labels for computing the sequence classification/regression loss. Indices should be in :obj:`[-100, 0, ...,32            config.vocab_size - 1]`. All labels set to ``-100`` are ignored (masked), the loss is only computed for33            labels in ``[0, ..., config.vocab_size]``34        Returns:35        Examples::36            >>> from transformers import T5Tokenizer, T5ForConditionalGeneration37            >>> tokenizer = T5Tokenizer.from_pretrained('t5-small')38            >>> model = T5ForConditionalGeneration.from_pretrained('t5-small')39            >>> # training40            >>> input_ids = tokenizer('The <extra_id_0> walks in <extra_id_1> park', return_tensors='pt').input_ids41            >>> labels = tokenizer('<extra_id_0> cute dog <extra_id_1> the <extra_id_2>', return_tensors='pt').input_ids42            >>> outputs = model(input_ids=input_ids, labels=labels)43            >>> loss = outputs.loss44            >>> logits = outputs.logits45            >>> # inference46            >>> input_ids = tokenizer("summarize: studies have shown that owning a dog is good for you", return_tensors="pt").input_ids  # Batch size 147            >>> outputs = model.generate(input_ids)48            >>> print(tokenizer.decode(outputs[0], skip_special_tokens=True))49            >>> # studies have shown that owning a dog is good for you.50        """51        if "cls" in kwargs:52            assert (53                    "input_ids" in kwargs and \54                    "labels" in kwargs and \55                    "attention_mask" in kwargs56            )57            return self.cls(58                input_ids=kwargs["input_ids"],59                labels=kwargs["labels"],60                attention_mask=kwargs["attention_mask"],61            )62        if "input_labels" in kwargs:63            assert (64                    "input_ids" in kwargs and \65                    "input_labels" in kwargs and \66                    "decoder_input_ids" in kwargs and \67                    "attention_mask" in kwargs and \68                    "decoder_attention_mask" in kwargs69            ), "Please give these arg keys."70            input_ids = kwargs["input_ids"]71            input_labels = kwargs["input_labels"]72            decoder_input_ids = kwargs["decoder_input_ids"]73            attention_mask = kwargs["attention_mask"]74            decoder_attention_mask = kwargs["decoder_attention_mask"]75            if "encoder_loss" not in kwargs:76                encoder_loss = True77            else:78                encoder_loss = kwargs["encoder_loss"]79            return self.review_forward(input_ids, input_labels, decoder_input_ids, attention_mask,80                                       decoder_attention_mask, encoder_loss)81        return super().forward(*argv, **kwargs)82 83    def cls(84            self,85            input_ids,86            labels,87            attention_mask,88    ):89        encoder_outputs = self.encoder( \90            input_ids=input_ids,91            attention_mask=attention_mask,92            output_attentions=False,93            return_dict=False94        )95        hidden_states = encoder_outputs[0]96        first_hidden = hidden_states[:, 0, :]97        first_hidden = nn.Dropout(0.3)(first_hidden)98        logits = self.cls_head(first_hidden)99        loss_fct = CrossEntropyLoss()100        if labels != None:101            loss = loss_fct(logits, labels)102            return loss103        return logits104 105    def review_forward(106            self,107            input_ids,108            input_labels,109            decoder_input_ids,110            attention_mask,111            decoder_attention_mask,112            encoder_loss=True113    ):114        encoder_outputs = self.encoder( \115            input_ids=input_ids,116            attention_mask=attention_mask,117            output_attentions=False,118            return_dict=False119        )120        hidden_states = encoder_outputs[0]121        decoder_inputs = self._shift_right(decoder_input_ids)122        # Decode123        decoder_outputs = self.decoder(124            input_ids=decoder_inputs,125            attention_mask=decoder_attention_mask,126            encoder_hidden_states=hidden_states,127            encoder_attention_mask=attention_mask,128            output_attentions=False,129            return_dict=False130        )131        sequence_output = decoder_outputs[0]132        if self.config.tie_word_embeddings:  # this is True default133            sequence_output = sequence_output * (self.model_dim ** -0.5)134        if encoder_loss:135            # print(self.encoder.get_input_embeddings().weight.shape)136            cls_logits = nn.functional.linear(hidden_states, self.encoder.get_input_embeddings().weight)137            # cls_logits = self.cls_head(hidden_states)138        lm_logits = self.lm_head(sequence_output)139        if decoder_input_ids is not None:140            lm_loss_fct = CrossEntropyLoss(ignore_index=0)  # Warning: PAD_ID should be 0141            loss = lm_loss_fct(lm_logits.view(-1, lm_logits.size(-1)), decoder_input_ids.view(-1))142            if encoder_loss and input_labels is not None:143                cls_loss_fct = CrossEntropyLoss(ignore_index=-100)144                loss += cls_loss_fct(cls_logits.view(-1, cls_logits.size(-1)), input_labels.view(-1))145            return loss146        return cls_logits, lm_logits147 148 149def prepare_models():150    tokenizer = AutoTokenizer.from_pretrained("microsoft/codereviewer")151 152    tokenizer.special_dict = {153        f"<e{i}>": tokenizer.get_vocab()[f"<e{i}>"] for i in range(99, -1, -1)154    }155    tokenizer.mask_id = tokenizer.get_vocab()["<mask>"]156    tokenizer.bos_id = tokenizer.get_vocab()["<s>"]157    tokenizer.pad_id = tokenizer.get_vocab()["<pad>"]158    tokenizer.eos_id = tokenizer.get_vocab()["</s>"]159    tokenizer.msg_id = tokenizer.get_vocab()["<msg>"]160    tokenizer.keep_id = tokenizer.get_vocab()["<keep>"]161    tokenizer.add_id = tokenizer.get_vocab()["<add>"]162    tokenizer.del_id = tokenizer.get_vocab()["<del>"]163    tokenizer.start_id = tokenizer.get_vocab()["<start>"]164    tokenizer.end_id = tokenizer.get_vocab()["<end>"]165 166    config = T5Config.from_pretrained("microsoft/codereviewer")167    model = ReviewerModel.from_pretrained("microsoft/codereviewer", config=config)168 169    model.eval()170    return tokenizer, model171 172 173def pad_assert(tokenizer, source_ids):174    source_ids = source_ids[:MAX_SOURCE_LENGTH - 2]175    source_ids = [tokenizer.bos_id] + source_ids + [tokenizer.eos_id]176    pad_len = MAX_SOURCE_LENGTH - len(source_ids)177    source_ids += [tokenizer.pad_id] * pad_len178    assert len(source_ids) == MAX_SOURCE_LENGTH, "Not equal length."179    return source_ids180 181 182def encode_diff(tokenizer, diff, msg, source):183    difflines = diff.split("\n")[1:]  # remove start @@184    difflines = [line for line in difflines if len(line.strip()) > 0]185    map_dic = {"-": 0, "+": 1, " ": 2}186 187    def f(s):188        if s in map_dic:189            return map_dic[s]190        else:191            return 2192 193    labels = [f(line[0]) for line in difflines]194    difflines = [line[1:].strip() for line in difflines]195    inputstr = "<s>" + source + "</s>"196    inputstr += "<msg>" + msg197    for label, line in zip(labels, difflines):198        if label == 1:199            inputstr += "<add>" + line200        elif label == 0:201            inputstr += "<del>" + line202        else:203            inputstr += "<keep>" + line204    source_ids = tokenizer.encode(inputstr, max_length=MAX_SOURCE_LENGTH, truncation=True)[1:-1]205    source_ids = pad_assert(tokenizer, source_ids)206    return source_ids207 208 209class FileDiffs(object):210    def __init__(self, diff_string):211        diff_array = diff_string.split("\n")212        self.file_name = diff_array[0]213        self.file_path = self.file_name.split("a/", 1)[1].rsplit("b/", 1)[0]214        self.diffs = list()215        for line in diff_array[4:]:216            if line.startswith("@@"):217                self.diffs.append(str())218            self.diffs[-1] += "\n" + line219 220 221def review_commit(user="p4vv37", repository="ueflow", commit="610a8c7b02b946bc9e5e26e6dacbba0e2abba259"):222    tokenizer, model = prepare_models()223 224    # Get diff and commit metadata from GitHub API225    commit_metadata = requests.get(F"https://api.github.com/repos/{user}/{repository}/commits/{commit}").json()226    msg = commit_metadata["commit"]["message"]227    diff_data = requests.get(F"https://api.github.com/repos/{user}/{repository}/commits/{commit}",228                             headers={"Accept": "application/vnd.github.diff"})229    code_diff = diff_data.text230 231    # Parse diff into FileDiffs objects232    files_diffs = list()233    for file in code_diff.split("diff --git"):234        if len(file) > 0:235            fd = FileDiffs(file)236            files_diffs.append(fd)237 238    # Generate comments for each diff239    output = ""240    for fd in files_diffs:241        output += F"File:{fd.file_path}\n"242        source = requests.get(F"https://raw.githubusercontent.com/{user}/{repository}/^{commit}/{fd.file_path}").text243 244        for diff in fd.diffs:245            inputs = torch.tensor([encode_diff(tokenizer, diff, msg, source)], dtype=torch.long).to("cpu")246            inputs_mask = inputs.ne(tokenizer.pad_id)247            logits = model(248                input_ids=inputs,249                cls=True,250                attention_mask=inputs_mask,251                labels=None,252                use_cache=True,253                num_beams=5,254                early_stopping=True,255                max_length=100256            )257            needs_review = torch.argmax(logits, dim=-1).cpu().numpy()[0]258            if not needs_review:259                continue260            preds = model.generate(inputs,261                                   attention_mask=inputs_mask,262                                   use_cache=True,263                                   num_beams=5,264                                   early_stopping=True,265                                   max_length=100,266                                   num_return_sequences=2267                                   )268            preds = list(preds.cpu().numpy())269            pred_nls = [tokenizer.decode(_id[2:], skip_special_tokens=True, clean_up_tokenization_spaces=False)270                        for _id in preds]271            output += diff + "\n#######\nComment:\n#######\n" + pred_nls[0] + "\n#######\n"272    return output273 274 275description = "An interface for running " \276              "\"Microsoft CodeBERT CodeReviewer: Pre-Training for Automating Code Review Activities.\" " \277              "(microsoft/codereviewer) on GitHub commits."278examples = [279    ["p4vv37", "ueflow", "610a8c7b02b946bc9e5e26e6dacbba0e2abba259"],280    ["microsoft", "vscode", "378b0d711f6b82ac59b47fb246906043a6fb995a"],281]282iface = gr.Interface(fn=review_commit,283                     description=description,284                     inputs=["text", "text", "text"],285                     outputs="text",286                     examples=examples,287                     cache_examples=False)288iface.launch()289