| 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: |