-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathinfer.py
More file actions
49 lines (43 loc) · 1.87 KB
/
Copy pathinfer.py
File metadata and controls
49 lines (43 loc) · 1.87 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
import os
import cv2
import torch
import numpy as np
from .model import UNet
@torch.no_grad()
def infer_image(weights_path: str, image_path: str, out_dir: str, thresh: float = 0.5):
os.makedirs(out_dir, exist_ok=True)
ckpt = torch.load(weights_path, map_location='cpu')
cfg = ckpt.get('cfg', {})
model = UNet(
in_channels=cfg.get('model',{}).get('in_channels',3),
out_channels=cfg.get('model',{}).get('out_channels',1),
features=tuple(cfg.get('model',{}).get('features',[64,128,256,512]))
)
model.load_state_dict(ckpt['model_state'])
model.eval()
img = cv2.imread(image_path, cv2.IMREAD_COLOR)
rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
h, w = rgb.shape[:2]
size = cfg.get('image_size', 512)
inp = cv2.resize(rgb, (size, size), interpolation=cv2.INTER_AREA).astype(np.float32)/255.0
inp = np.transpose(inp, (2,0,1))[None]
logits = model(torch.from_numpy(inp))
prob = torch.sigmoid(logits)[0,0].numpy()
mask = (prob > thresh).astype(np.uint8)*255
mask_up = cv2.resize(mask, (w, h), interpolation=cv2.INTER_NEAREST)
overlay = img.copy()
red = np.zeros_like(img)
red[:,:,2] = 255
overlay = cv2.addWeighted(overlay, 1.0, (red*(mask_up[...,None]>0)).astype(np.uint8), 0.3, 0)
base = os.path.splitext(os.path.basename(image_path))[0]
cv2.imwrite(os.path.join(out_dir, f"{base}_mask.png"), mask_up)
cv2.imwrite(os.path.join(out_dir, f"{base}_overlay.png"), overlay)
if __name__ == '__main__':
import argparse
parser = argparse.ArgumentParser()
parser.add_argument('--weights', type=str, required=True)
parser.add_argument('--image', type=str, required=True)
parser.add_argument('--out', type=str, default='outputs/preds')
parser.add_argument('--thresh', type=float, default=0.5)
args = parser.parse_args()
infer_image(args.weights, args.image, args.out, args.thresh)