File size: 4,182 Bytes
a948f72
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
#!/usr/bin/env python3
# Copyright 2022-2024,2026 NXP
# SPDX-License-Identifier: MIT

import argparse
import time

import cv2
import numpy as np

LINE_COLOR = (255, 128, 0)
POINT_COLOR = (0, 0, 255)

# Keypoint definitions:
# https://github.com/tensorflow/tfjs-models/tree/master/pose-detection#keypoint-diagram
keypoints_def = [
    {'label': 'nose',           'connections': [1, 2]},
    {'label': 'left_eye',       'connections': [0, 3]},
    {'label': 'right_eye',      'connections': [0, 4]},
    {'label': 'left_ear',       'connections': [1]},
    {'label': 'right_ear',      'connections': [2]},
    {'label': 'left_shoulder',  'connections': [6, 7, 11]},
    {'label': 'right_shoulder', 'connections': [5, 8, 12]},
    {'label': 'left_elbow',     'connections': [5, 9]},
    {'label': 'right_elbow',    'connections': [6, 10]},
    {'label': 'left_wrist',     'connections': [7]},
    {'label': 'right_wrist',    'connections': [8]},
    {'label': 'left_hip',       'connections': [5, 12, 13]},
    {'label': 'right_hip',      'connections': [6, 11, 14]},
    {'label': 'left_knee',      'connections': [11, 15]},
    {'label': 'right_knee',     'connections': [12, 16]},
    {'label': 'left_ankle',     'connections': [13]},
    {'label': 'right_ankle',    'connections': [14]},
]

connections = [(i, j) for i in range(len(keypoints_def))
               for j in keypoints_def[i]['connections']]


def load_image(filename):
    orig_image = cv2.imread(filename, 1)
    image = cv2.cvtColor(orig_image, cv2.COLOR_BGR2RGB)
    image = cv2.resize(image, (192, 192))
    image = np.expand_dims(image, axis=0)
    return orig_image, image


def run_inference(interpreter, image):
    input_details = interpreter.get_input_details()
    output_details = interpreter.get_output_details()

    # Handle int8 quantized input
    input_dtype = input_details[0]['dtype']
    if input_dtype == np.int8:
        scale, zero_point = input_details[0]['quantization']
        image = (image.astype(np.float32) / scale + zero_point).astype(np.int8)
    else:
        image = image.astype(np.uint8)

    interpreter.set_tensor(input_details[0]['index'], image)
    interpreter.invoke()
    out = interpreter.get_tensor(output_details[0]['index'])

    # Dequantize int8 output if needed
    output_dtype = output_details[0]['dtype']
    if output_dtype == np.int8:
        scale, zero_point = output_details[0]['quantization']
        out = (out.astype(np.float32) - zero_point) * scale

    return out.astype(np.float32)


def main():
    parser = argparse.ArgumentParser(description='MoveNet single-pose Lightning demo')
    parser.add_argument('-m', '--model', default='original_model/movenet_quant.tflite',
                        help='Path to TFLite model file')
    parser.add_argument('-i', '--input', default='example_input.jpg',
                        help='Path to input image')
    parser.add_argument('-o', '--output', default=None,
                        help='Path to save output image (optional)')
    args = parser.parse_args()

    try:
        import tflite_runtime.interpreter as tflite
        interpreter = tflite.Interpreter(args.model)
    except ImportError:
        import tensorflow as tf
        interpreter = tf.lite.Interpreter(args.model)

    interpreter.allocate_tensors()

    orig_image, processed_image = load_image(args.input)

    start = time.time()
    out = run_inference(interpreter, processed_image)[0, 0, ...]
    end = time.time()
    print("Inference time: {:.1f} ms".format((end - start) * 1000))

    h, w = orig_image.shape[0:2]
    out[:, 0] *= h
    out[:, 1] *= w

    for c in connections:
        i, j = c
        cv2.line(orig_image,
                 (int(out[i, 1]), int(out[i, 0])),
                 (int(out[j, 1]), int(out[j, 0])),
                 LINE_COLOR, 5)

    for i in range(out.shape[0]):
        cv2.circle(orig_image,
                   (int(out[i, 1]), int(out[i, 0])), 5, POINT_COLOR, 10)

    if args.output:
        cv2.imwrite(args.output, orig_image)
        print("Output saved to", args.output)
    else:
        cv2.imshow('MoveNet', orig_image)
        cv2.waitKey()


if __name__ == '__main__':
    main()