Lazyhope/python-clone-detection
374
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 