(self, source_ids=None, labels=None, weight=None)
| 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 | |
| 205 | class DecoderClassifier(nn.Module): |
nothing calls this directly
no test coverage detected