(self, setting, test=0)
| 157 | return self.model |
| 158 | |
| 159 | def test(self, setting, test=0): |
| 160 | test_data, test_loader = self._get_data(flag='TEST') |
| 161 | if test: |
| 162 | print('loading model') |
| 163 | self.model.load_state_dict(torch.load(os.path.join('./checkpoints/' + setting, 'checkpoint.pth'))) |
| 164 | |
| 165 | preds = [] |
| 166 | trues = [] |
| 167 | folder_path = './test_results/' + setting + '/' |
| 168 | if not os.path.exists(folder_path): |
| 169 | os.makedirs(folder_path) |
| 170 | |
| 171 | self.model.eval() |
| 172 | with torch.no_grad(): |
| 173 | for i, (batch_x, label, padding_mask) in enumerate(test_loader): |
| 174 | batch_x = batch_x.float().to(self.device) |
| 175 | padding_mask = padding_mask.float().to(self.device) |
| 176 | label = label.to(self.device) |
| 177 | |
| 178 | outputs = self.model(batch_x, padding_mask, None, None) |
| 179 | |
| 180 | preds.append(outputs.detach()) |
| 181 | trues.append(label) |
| 182 | |
| 183 | preds = torch.cat(preds, 0) |
| 184 | trues = torch.cat(trues, 0) |
| 185 | print('test shape:', preds.shape, trues.shape) |
| 186 | |
| 187 | probs = torch.nn.functional.softmax(preds) # (total_samples, num_classes) est. prob. for each class and sample |
| 188 | predictions = torch.argmax(probs, dim=1).cpu().numpy() # (total_samples,) int class index for each sample |
| 189 | trues = trues.flatten().cpu().numpy() |
| 190 | accuracy = cal_accuracy(predictions, trues) |
| 191 | |
| 192 | # result save |
| 193 | folder_path = './results/' + setting + '/' |
| 194 | if not os.path.exists(folder_path): |
| 195 | os.makedirs(folder_path) |
| 196 | |
| 197 | print('accuracy:{}'.format(accuracy)) |
| 198 | file_name='result_classification.txt' |
| 199 | f = open(os.path.join(folder_path,file_name), 'a') |
| 200 | f.write(setting + " \n") |
| 201 | f.write('accuracy:{}'.format(accuracy)) |
| 202 | f.write('\n') |
| 203 | f.write('\n') |
| 204 | f.close() |
| 205 | return |
nothing calls this directly
no test coverage detected