MCPcopy Create free account
hub / github.com/OpenBMB/DecT / DecTRunner

Class DecTRunner

src/dect_trainer.py:22–102  ·  view source on GitHub ↗

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.

Source from the content-addressed store, hash-verified

20from openprompt.utils.logging import logger
21
22class 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):

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected