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

Method forward

os_expr/model.py:182–202  ·  view source on GitHub ↗
(self, source_ids=None, labels=None, weight=None)

Source from the content-addressed store, hash-verified

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:
198 loss_fct = nn.CrossEntropyLoss(weight=weight)
199 loss = loss_fct(logits, labels)
200 return loss, prob
201 else:
202 return prob
203
204
205class DecoderClassifier(nn.Module):

Callers

nothing calls this directly

Calls 3

get_t5_vecMethod · 0.95
get_bart_vecMethod · 0.95
get_roberta_vecMethod · 0.95

Tested by

no test coverage detected