Harsha909/video-pose-normalization
0
1# Copyright (c) OpenMMLab. All rights reserved.
2from unittest import TestCase
3
4import numpy as np
5
6from mmpose.codecs import IntegralRegressionLabel # noqa: F401
7from mmpose.registry import KEYPOINT_CODECS
8
9
10class TestRegressionLabel(TestCase):
11
12 # name and configs of all test cases
13 def setUp(self) -> None:
14 self.configs = [
15 (
16 'ipr',
17 dict(
18 type='IntegralRegressionLabel',
19 input_size=(192, 256),
20 heatmap_size=(48, 64),
21 sigma=2),
22 ),
23 ]
24
25 # The bbox is usually padded so the keypoint will not be near the
26 # boundary
27 keypoints = (0.1 + 0.8 * np.random.rand(1, 17, 2)) * [192, 256]
28 keypoints = np.round(keypoints).astype(np.float32)
29 heatmaps = np.random.rand(17, 64, 48).astype(np.float32)
30 encoded_wo_sigma = np.random.rand(1, 17, 2)
31 keypoints_visible = np.ones((1, 17), dtype=np.float32)
32 self.data = dict(
33 keypoints=keypoints,
34 keypoints_visible=keypoints_visible,
35 heatmaps=heatmaps,
36 encoded_wo_sigma=encoded_wo_sigma)
37
38 def test_encode(self):
39 keypoints = self.data['keypoints']
40 keypoints_visible = self.data['keypoints_visible']
41
42 for name, cfg in self.configs:
43 codec = KEYPOINT_CODECS.build(cfg)
44
45 encoded = codec.encode(keypoints, keypoints_visible)
46 heatmaps = encoded['heatmaps']
47 keypoint_labels = encoded['keypoint_labels']
48 keypoint_weights = encoded['keypoint_weights']
49
50 self.assertEqual(heatmaps.shape, (17, 64, 48),
51 f'Failed case: "{name}"')
52 self.assertEqual(keypoint_labels.shape, (1, 17, 2),
53 f'Failed case: "{name}"')
54 self.assertEqual(keypoint_weights.shape, (1, 17),
55 f'Failed case: "{name}"')
56
57 def test_decode(self):
58 encoded_wo_sigma = self.data['encoded_wo_sigma']
59
60 for name, cfg in self.configs:
61 codec = KEYPOINT_CODECS.build(cfg)
62
63 keypoints, scores = codec.decode(encoded_wo_sigma)
64
65 self.assertEqual(keypoints.shape, (1, 17, 2),
66 f'Failed case: "{name}"')
67 self.assertEqual(scores.shape, (1, 17), f'Failed case: "{name}"')
68
69 def test_cicular_verification(self):
70 keypoints = self.data['keypoints']
71 keypoints_visible = self.data['keypoints_visible']
72
73 for name, cfg in self.configs:
74 codec = KEYPOINT_CODECS.build(cfg)
75
76 encoded = codec.encode(keypoints, keypoints_visible)
77 keypoint_labels = encoded['keypoint_labels']
78
79 _keypoints, _ = codec.decode(keypoint_labels)
80
81 self.assertTrue(
82 np.allclose(keypoints, _keypoints, atol=5.),
83 f'Failed case: "{name}"')
84 