MCPcopy Create free account
hub / github.com/DLVulDet/PrimeVul / DecoderClassifier

Class DecoderClassifier

os_expr/model.py:205–245  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

203
204
205class DecoderClassifier(nn.Module):
206 def __init__(self, encoder, config, tokenizer, args):
207 super(DecoderClassifier, self).__init__()
208 self.encoder = encoder
209 self.config=config
210 self.tokenizer=tokenizer
211 self.args=args
212 self.classifier = nn.Linear(config.hidden_size, 2)
213
214 def forward(self, input_ids=None, labels=None, weight=None):
215 attention_mask = input_ids.ne(self.tokenizer.pad_token_id)
216 outputs = self.encoder(input_ids, attention_mask=attention_mask)
217 hidden_states = outputs[0]
218 logits = self.classifier(hidden_states)
219
220 batch_size = input_ids.size(0)
221 if self.config.pad_token_id is None and batch_size != 1:
222 raise ValueError("Cannot handle batch sizes > 1 if no padding token is defined.")
223 if self.config.pad_token_id is None:
224 sequence_lengths = -1
225 else:
226 if input_ids is not None:
227 # if no pad token found, use modulo instead of reverse indexing for ONNX compatibility
228 sequence_lengths = torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1
229 sequence_lengths = sequence_lengths % input_ids.shape[-1]
230 sequence_lengths = sequence_lengths.to(logits.device)
231 else:
232 sequence_lengths = -1
233
234 pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths]
235 prob = nn.functional.softmax(pooled_logits, dim=-1)
236
237
238 if labels is not None:
239 labels = labels.to(logits.device)
240 loss_fct = nn.CrossEntropyLoss(weight=weight)
241
242 loss = loss_fct(pooled_logits.view(-1, 2), labels.view(-1))
243 return loss, prob
244 else:
245 return prob

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected