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

Method __init__

src/dect_trainer.py:32–49  ·  view source on GitHub ↗
(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,
                 )

Source from the content-addressed store, hash-verified

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')

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected