| 138 | |
| 139 | |
| 140 | class DefectModel(nn.Module): |
| 141 | def __init__(self, encoder, config, tokenizer, args): |
| 142 | super(DefectModel, self).__init__() |
| 143 | self.encoder = encoder |
| 144 | self.config = config |
| 145 | self.tokenizer = tokenizer |
| 146 | self.classifier = nn.Linear(config.hidden_size, 2) |
| 147 | self.args = args |
| 148 | |
| 149 | def get_t5_vec(self, source_ids): |
| 150 | attention_mask = source_ids.ne(self.tokenizer.pad_token_id) |
| 151 | outputs = self.encoder(input_ids=source_ids, attention_mask=attention_mask, |
| 152 | labels=source_ids, decoder_attention_mask=attention_mask, output_hidden_states=True) |
| 153 | hidden_states = outputs['decoder_hidden_states'][-1] |
| 154 | eos_mask = source_ids.eq(self.config.eos_token_id) |
| 155 | |
| 156 | if len(torch.unique(eos_mask.sum(1))) > 1: |
| 157 | print(eos_mask.sum(1)) |
| 158 | print(torch.unique(eos_mask.sum(1))) |
| 159 | raise ValueError("All examples must have the same number of <eos> tokens.") |
| 160 | vec = hidden_states[eos_mask, :].view(hidden_states.size(0), -1, |
| 161 | hidden_states.size(-1))[:, -1, :] |
| 162 | return vec |
| 163 | |
| 164 | def get_bart_vec(self, source_ids): |
| 165 | attention_mask = source_ids.ne(self.tokenizer.pad_token_id) |
| 166 | outputs = self.encoder(input_ids=source_ids, attention_mask=attention_mask, |
| 167 | labels=source_ids, decoder_attention_mask=attention_mask, output_hidden_states=True) |
| 168 | hidden_states = outputs['decoder_hidden_states'][-1] |
| 169 | eos_mask = source_ids.eq(self.config.eos_token_id) |
| 170 | |
| 171 | if len(torch.unique(eos_mask.sum(1))) > 1: |
| 172 | raise ValueError("All examples must have the same number of <eos> tokens.") |
| 173 | vec = hidden_states[eos_mask, :].view(hidden_states.size(0), -1, |
| 174 | hidden_states.size(-1))[:, -1, :] |
| 175 | return vec |
| 176 | |
| 177 | def get_roberta_vec(self, source_ids): |
| 178 | attention_mask = source_ids.ne(self.tokenizer.pad_token_id) |
| 179 | vec = self.encoder(input_ids=source_ids, attention_mask=attention_mask)[0][:, 0, :] |
| 180 | return vec |
| 181 | |
| 182 | def forward(self, source_ids=None, labels=None, weight=None): |
| 183 | # source_ids = source_ids.view(-1, self.args.max_source_length) |
| 184 | |
| 185 | if self.args.model_type == 'codet5': |
| 186 | vec = self.get_t5_vec(source_ids) |
| 187 | elif self.args.model_type == 'bart': |
| 188 | vec = self.get_bart_vec(source_ids) |
| 189 | elif self.args.model_type == 'roberta': |
| 190 | vec = self.get_roberta_vec(source_ids) |
| 191 | elif self.args.model_type == 't5': |
| 192 | vec = self.get_t5_vec(source_ids) |
| 193 | |
| 194 | logits = self.classifier(vec) |
| 195 | prob = nn.functional.softmax(logits) |
| 196 | |
| 197 | if labels is not None: |