MCPcopy Create free account
hub / github.com/THUDM/P-tuning / evaluate

Method evaluate

LAMA/cli.py:117–140  ·  view source on GitHub ↗
(self, epoch_idx, evaluate_type)

Source from the content-addressed store, hash-verified

115 return relation, data_path_pre, data_path_post
116
117 def evaluate(self, epoch_idx, evaluate_type):
118 self.model.eval()
119 if evaluate_type == 'Test':
120 loader = self.test_loader
121 dataset = self.test_set
122 else:
123 loader = self.dev_loader
124 dataset = self.dev_set
125 with torch.no_grad():
126 self.model.eval()
127 hit1, loss = 0, 0
128 for x_hs, x_ts in loader:
129 if False and self.args.extend_data:
130 _loss, _hit1 = self.model.test_extend_data(x_hs, x_ts)
131 elif evaluate_type == 'Test':
132 _loss, _hit1, top10 = self.model(x_hs, x_ts, return_candidates=True)
133 else:
134 _loss, _hit1 = self.model(x_hs, x_ts)
135 hit1 += _hit1
136 loss += _loss.item()
137 hit1 /= len(dataset)
138 print("{} {} Epoch {} Loss: {} Hit@1:".format(self.args.relation_id, evaluate_type, epoch_idx,
139 loss / len(dataset)), hit1)
140 return loss, hit1
141
142 def get_task_name(self):
143 if self.args.only_evaluate:

Callers 1

trainMethod · 0.95

Calls 1

evalMethod · 0.80

Tested by

no test coverage detected