(model, data_loader_test, pred_root, method, testset, device=0)
| 16 | |
| 17 | |
| 18 | def inference(model, data_loader_test, pred_root, method, testset, device=0): |
| 19 | model_training = model.training |
| 20 | if model_training: |
| 21 | model.eval() |
| 22 | for batch in tqdm(data_loader_test, total=len(data_loader_test)) if 1 or config.verbose_eval else data_loader_test: |
| 23 | inputs = batch[0].to(device) |
| 24 | # gts = batch[1].to(device) |
| 25 | label_paths = batch[-1] |
| 26 | with torch.no_grad(): |
| 27 | scaled_preds = model(inputs)[-1].sigmoid() |
| 28 | |
| 29 | os.makedirs(os.path.join(pred_root, method, testset), exist_ok=True) |
| 30 | |
| 31 | for idx_sample in range(scaled_preds.shape[0]): |
| 32 | res = nn.functional.interpolate( |
| 33 | scaled_preds[idx_sample].unsqueeze(0), |
| 34 | size=cv2.imread(label_paths[idx_sample], cv2.IMREAD_GRAYSCALE).shape[:2], |
| 35 | mode='bilinear', |
| 36 | align_corners=True |
| 37 | ) |
| 38 | save_tensor_img(res, os.path.join(os.path.join(pred_root, method, testset), label_paths[idx_sample].replace('\\', '/').split('/')[-1])) # test set dir + file name |
| 39 | if model_training: |
| 40 | model.train() |
| 41 | return None |
| 42 | |
| 43 | |
| 44 | def main(args): |
no test coverage detected