pytorch torchvision 实现的keypoint rcnn使用python进行预测并显示预测结果
·
深度学习中,完成模型训练后,需要编写测试文件,读取输入数据和模型进行推理,对推理结果进行输出和显示。
这里参考如下教程进行模型的训练:
How to Train a Custom Keypoint Detection Model with PyTorch
代码是集合训练测试代码的jupyter notebook文件,这里自己整理出python的测试代码如下:
import torch
import torchvision
from torchvision.models.detection.rpn import AnchorGenerator
from torchvision.transforms import functional as F
import cv2
PATH = r'E:\git_rep\keypoint_rcnn_training_pytorch\keypointsrcnn_weights.pth'
def get_model(num_keypoints, weights_path=None):
anchor_generator = AnchorGenerator(sizes=(32, 64, 128, 256, 512),
aspect_ratios=(0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 4.0))
model = torchvision.models.detection.keypointrcnn_resnet50_fpn(pretrained=False,
pretrained_backbone=True,
num_keypoints=num_keypoints,
num_classes=2,
# Background is the first class, object is the second class
rpn_anchor_generator=anchor_generator)
if weights_path:
state_dict = torch.load(weights_path)
model.load_state_dict(state_dict)
return model
device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')
model = get_model(num_keypoints = 2)
model.to(device)
model.load_state_dict(torch.load(PATH), strict=True)
img_path = r'E:\datasets\24\2022-04-20\101_color_2.png'
# img_path = r'E:\datasets\24\2021-09-07\1_color_2.png'
img = cv2.imread(img_path)
# cv2.imshow('img', img)
# cv2.waitKey(0)
rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
# cv2.imshow('rgb', rgb)
# cv2.waitKey(0)
model.eval()
img = F.to_tensor(rgb)
# rgb = torch.Tensor(rgb).to(device)
# rgb = rgb.permute(2,0,1)
images = [img.to(device)]
# rgb = torchvision.transforms.ToPILImage(rgb)
with torch.no_grad():
output = model(images)
print("Predictions: \n", output)
import numpy as np
image = (images[0].permute(1, 2, 0).detach().cpu().numpy() * 255).astype(np.uint8)
scores = output[0]['scores'].detach().cpu().numpy()
high_scores_idxs = np.where(scores > 0.7)[0].tolist() # Indexes of boxes with scores > 0.7
post_nms_idxs = torchvision.ops.nms(output[0]['boxes'][high_scores_idxs], output[0]['scores'][high_scores_idxs],
0.3).cpu().numpy() # Indexes of boxes left after applying NMS (iou_threshold=0.3)
# Below, in output[0]['keypoints'][high_scores_idxs][post_nms_idxs] and output[0]['boxes'][high_scores_idxs][post_nms_idxs]
# Firstly, we choose only those objects, which have score above predefined threshold. This is done with choosing elements with [high_scores_idxs] indexes
# Secondly, we choose only those objects, which are left after NMS is applied. This is done with choosing elements with [post_nms_idxs] indexes
keypoints = []
for kps in output[0]['keypoints'][high_scores_idxs][post_nms_idxs].detach().cpu().numpy():
keypoints.append([list(map(int, kp[:2])) for kp in kps])
bboxes = []
for bbox in output[0]['boxes'][high_scores_idxs][post_nms_idxs].detach().cpu().numpy():
bboxes.append(list(map(int, bbox.tolist())))
keypoints_classes_ids2names = {0: 'Tragus', 1: 'Head'}
import matplotlib.pyplot as plt
def visualize(image, bboxes, keypoints, image_original=None, bboxes_original=None, keypoints_original=None):
fontsize = 18
for bbox in bboxes:
start_point = (bbox[0], bbox[1])
end_point = (bbox[2], bbox[3])
image = cv2.rectangle(image.copy(), start_point, end_point, (0, 255, 0), 2)
for kps in keypoints:
for idx, kp in enumerate(kps):
image = cv2.circle(image.copy(), tuple(kp), 5, (255, 0, 0), 10)
image = cv2.putText(image.copy(), " " + keypoints_classes_ids2names[idx], tuple(kp),
cv2.FONT_HERSHEY_SIMPLEX, 2, (255, 0, 0), 3, cv2.LINE_AA)
if image_original is None and keypoints_original is None:
plt.figure(figsize=(40, 40))
plt.imshow(image)
else:
for bbox in bboxes_original:
start_point = (bbox[0], bbox[1])
end_point = (bbox[2], bbox[3])
image_original = cv2.rectangle(image_original.copy(), start_point, end_point, (0, 255, 0), 2)
for kps in keypoints_original:
for idx, kp in enumerate(kps):
image_original = cv2.circle(image_original, tuple(kp), 5, (255, 0, 0), 10)
image_original = cv2.putText(image_original, " " + keypoints_classes_ids2names[idx], tuple(kp),
cv2.FONT_HERSHEY_SIMPLEX, 2, (255, 0, 0), 3, cv2.LINE_AA)
f, ax = plt.subplots(1, 2, figsize=(40, 20))
ax[0].imshow(image_original)
ax[0].set_title('Original image', fontsize=fontsize)
ax[1].imshow(image)
ax[1].set_title('Transformed image', fontsize=fontsize)
visualize(image, bboxes, keypoints)
更多推荐



所有评论(0)