r"""A runner for DecT This class is specially implemented for classification. Args: model (:obj:`PromptForClassification`): One ``PromptForClassification`` object. train_dataloader (:obj:`PromptDataloader`, optional): The dataloader to bachify and process the training data.
| 20 | from openprompt.utils.logging import logger |
| 21 | |
| 22 | class DecTRunner(object): |
| 23 | r"""A runner for DecT |
| 24 | This class is specially implemented for classification. |
| 25 | |
| 26 | Args: |
| 27 | model (:obj:`PromptForClassification`): One ``PromptForClassification`` object. |
| 28 | train_dataloader (:obj:`PromptDataloader`, optional): The dataloader to bachify and process the training data. |
| 29 | valid_dataloader (:obj:`PromptDataloader`, optionla): The dataloader to bachify and process the val data. |
| 30 | test_dataloader (:obj:`PromptDataloader`, optional): The dataloader to bachify and process the test data. |
| 31 | """ |
| 32 | def __init__(self, |
| 33 | model: PromptForClassification, |
| 34 | train_dataloader: Optional[PromptDataLoader] = None, |
| 35 | valid_dataloader: Optional[PromptDataLoader] = None, |
| 36 | test_dataloader: Optional[PromptDataLoader] = None, |
| 37 | calibrate_dataloader: Optional[PromptDataLoader] = None, |
| 38 | id2label: Optional[Dict] = None, |
| 39 | verbalizer = None, |
| 40 | ): |
| 41 | self.model = model.cuda() |
| 42 | self.train_dataloader = train_dataloader |
| 43 | self.valid_dataloader = valid_dataloader |
| 44 | self.test_dataloader = test_dataloader |
| 45 | self.calibrate_dataloader = calibrate_dataloader |
| 46 | self.loss_function = torch.nn.CrossEntropyLoss() |
| 47 | self.id2label = id2label |
| 48 | self.verbalizer = verbalizer |
| 49 | self.clean = True |
| 50 | |
| 51 | def inference_step(self, batch, batch_idx): |
| 52 | label = batch.pop('label') |
| 53 | logits = self.model(batch) |
| 54 | pred = torch.argmax(logits, dim=-1) |
| 55 | return pred.cpu().tolist(), label.cpu().tolist() |
| 56 | |
| 57 | def inference_epoch(self, split: str): |
| 58 | outputs = [] |
| 59 | scores = {} |
| 60 | self.model.eval() |
| 61 | with torch.no_grad(): |
| 62 | data_loader = self.valid_dataloader if split=='validation' else self.test_dataloader |
| 63 | model_preds, preds, labels = self.verbalizer.test(self.model, data_loader) |
| 64 | # zs_score = accuracy_score(labels, model_preds) |
| 65 | score = accuracy_score(labels, preds) |
| 66 | scores = {"dect acc": score} |
| 67 | return scores |
| 68 | |
| 69 | def inference_epoch_end(self, outputs): |
| 70 | preds = [] |
| 71 | labels = [] |
| 72 | for pred, label in outputs: |
| 73 | preds.extend(pred) |
| 74 | labels.extend(label) |
| 75 | |
| 76 | score = accuracy_score(labels, preds) |
| 77 | return score |
| 78 | |
| 79 | def training_step(self, batch, batch_idx): |