MCPcopy Create free account
hub / github.com/ZhengPeng7/BiRefNet / inference

Function inference

inference.py:18–41  ·  view source on GitHub ↗
(model, data_loader_test, pred_root, method, testset, device=0)

Source from the content-addressed store, hash-verified

16
17
18def 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
44def main(args):

Callers 2

validFunction · 0.90
mainFunction · 0.85

Calls 2

save_tensor_imgFunction · 0.90
trainMethod · 0.80

Tested by

no test coverage detected