CoolFace
Apppublic

Harsha909/video-pose-normalization

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
test_integral_regression_label.py84 linesDownload Raw Back to test_codecs
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