MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / test

Method test

exp/exp_classification.py:159–205  ·  view source on GitHub ↗
(self, setting, test=0)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 3

_get_dataMethod · 0.95
cal_accuracyFunction · 0.90
loadMethod · 0.80

Tested by

no test coverage detected