(self,
model: PromptForClassification,
train_dataloader: Optional[PromptDataLoader] = None,
valid_dataloader: Optional[PromptDataLoader] = None,
test_dataloader: Optional[PromptDataLoader] = None,
calibrate_dataloader: Optional[PromptDataLoader] = None,
id2label: Optional[Dict] = None,
verbalizer = None,
)
| 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') |
nothing calls this directly
no outgoing calls
no test coverage detected