CoolFace
Modelpublic

Lazyhope/python-clone-detection

sourceHugging Facemitupdated 4y agoView on Hugging Face
3likes74downloads
CloneDetectionModel.py97 linesDownload Raw Back to root
1"""2Original work:3https://github.com/sangHa0411/CloneDetection/blob/main/models/codebert.py#L1694 5Copyright (c) 2022 Sangha Park(sangha110495), Young Jin Ahn(snoop2head)6 7All credits to the original authors.8"""9import torch.nn as nn10from transformers import (11    RobertaPreTrainedModel,12    RobertaModel,13)14from transformers.modeling_outputs import SequenceClassifierOutput15 16 17class CloneDetectionModel(RobertaPreTrainedModel):18    _keys_to_ignore_on_load_missing = [r"position_ids"]19 20    def __init__(self, config):21        super().__init__(config)22        self.num_labels = config.num_labels23        self.config = config24 25        self.roberta = RobertaModel(config, add_pooling_layer=False)26        self.net = nn.Sequential(27            nn.Dropout(config.hidden_dropout_prob),28            nn.Linear(config.hidden_size, config.hidden_size),29            nn.ReLU(),30        )31        self.classifier = nn.Linear(config.hidden_size * 4, config.num_labels)32 33    def forward(34        self,35        input_ids=None,36        attention_mask=None,37        token_type_ids=None,38        position_ids=None,39        head_mask=None,40        inputs_embeds=None,41        labels=None,42        output_attentions=None,43        output_hidden_states=None,44        return_dict=None,45    ):46 47        return_dict = (48            return_dict if return_dict is not None else self.config.use_return_dict49        )50 51        outputs = self.roberta(52            input_ids,53            attention_mask=attention_mask,54            token_type_ids=token_type_ids,55            position_ids=position_ids,56            head_mask=head_mask,57            inputs_embeds=inputs_embeds,58            output_attentions=output_attentions,59            output_hidden_states=output_hidden_states,60            return_dict=return_dict,61        )62 63        hidden_states = outputs[0]64        batch_size, _, hidden_size = hidden_states.shape65 66        # CLS code1 SEP SEP code2 SEP67        cls_flag = input_ids == self.config.tokenizer_cls_token_id  # cls token68        sep_flag = input_ids == self.config.tokenizer_sep_token_id  # sep token69 70        special_token_states = hidden_states[cls_flag + sep_flag].view(71            batch_size, -1, hidden_size72        )  # (batch_size, 4, hidden_size)73        special_hidden_states = self.net(74            special_token_states75        )  # (batch_size, 4, hidden_size)76 77        pooled_output = special_hidden_states.view(78            batch_size, -179        )  # (batch_size, hidden_size * 4)80        logits = self.classifier(pooled_output)  # (batch_size, num_labels)81 82        loss = None83        if labels is not None:84            loss_fct = nn.CrossEntropyLoss()85            loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1))86 87        if not return_dict:88            output = (logits,) + outputs[2:]89            return ((loss,) + output) if loss is not None else output90 91        return SequenceClassifierOutput(92            loss=loss,93            logits=logits,94            hidden_states=outputs.hidden_states,95            attentions=outputs.attentions,96        )97