-
Notifications
You must be signed in to change notification settings - Fork 11
Expand file tree
/
Copy pathpredict_visualize.py
More file actions
105 lines (79 loc) · 3.28 KB
/
Copy pathpredict_visualize.py
File metadata and controls
105 lines (79 loc) · 3.28 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
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
import datetime
import os
import time
import torch
import torch.utils.data
from torch import nn
import torchvision
import torchvision.models.detection
import torchvision.models.detection.mask_rcnn
from torchvision import transforms
from coco_utils import get_coco, get_coco_kp
from group_by_aspect_ratio import GroupedBatchSampler, create_aspect_ratio_groups
from engine import train_one_epoch, evaluate
import utils
import transforms as T
from PIL import Image
from plot import plot_poses
import numpy as np
def get_dataset(name, image_set, transform):
paths = {
"coco": ('/home/hzj/data/COCO2017/', get_coco, 91),
"coco_kp": ('/home/hzj/data/COCO2017/', get_coco_kp, 2)
}
p, ds_fn, num_classes = paths[name]
ds = ds_fn(p, image_set=image_set, transforms=transform)
return ds, num_classes
def get_transform(train):
transforms = []
transforms.append(T.ToTensor())
if train:
transforms.append(T.RandomHorizontalFlip(0.5))
return T.Compose(transforms)
def main():
device = torch.device("cuda:0")
# Data loading code
print("Loading data")
#dataset, num_classes = get_dataset(args.dataset, "train", get_transform(train=True))
dataset_test, num_classes = get_dataset("coco_kp", "val", get_transform(train=False))
print("Creating data loaders")
#train_sampler = torch.utils.data.RandomSampler(dataset)
test_sampler = torch.utils.data.SequentialSampler(dataset_test)
#train_batch_sampler = torch.utils.data.BatchSampler(
# train_sampler, args.batch_size, drop_last=True)
#data_loader = torch.utils.data.DataLoader(
# dataset, batch_sampler=train_batch_sampler, num_workers=args.workers,
# collate_fn=utils.collate_fn)
data_loader_test = torch.utils.data.DataLoader(
dataset_test, batch_size=1,
sampler=test_sampler, num_workers=4,
collate_fn=utils.collate_fn)
print("Creating model")
model = torchvision.models.detection.__dict__['keypointrcnn_resnet50_fpn'](num_classes=num_classes,
pretrained=True)
model.to(device)
#checkpoint = torch.load(args.resume, map_location='cpu')
#model_without_ddp.load_state_dict(checkpoint['model'])
#optimizer.load_state_dict(checkpoint['optimizer'])
#lr_scheduler.load_state_dict(checkpoint['lr_scheduler'])
model.eval()
detect_threshold = 0.7
keypoint_score_threshold = 2
with torch.no_grad():
for i in range(20):
img,_ = dataset_test[i]
prediction = model([img.to(device)])
keypoints = prediction[0]['keypoints'].cpu().numpy()
scores = prediction[0]['scores'].cpu().numpy()
keypoints_scores = prediction[0]['keypoints_scores'].cpu().numpy()
idx = np.where(scores>detect_threshold)
keypoints = keypoints[idx]
keypoints_scores = keypoints_scores[idx]
for j in range(keypoints.shape[0]):
for num in range(17):
if keypoints_scores[j][num]<keypoint_score_threshold:
keypoints[j][num]=[0,0,0]
img = img.mul(255).permute(1, 2, 0).byte().numpy()
plot_poses(img,keypoints,save_name='./result/'+str(i)+'.jpg')
if __name__ == "__main__":
main()