File size: 3,283 Bytes
7346868
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
# Copyright 2023-2024,2026 NXP
# SPDX-License-Identifier: MIT

import argparse
import time
from random import seed, randint

import cv2
import numpy as np
import matplotlib.pyplot as plt

seed(1337)

N_CLASSES = 21
COLORS = [(0, 0, 0)]
COLORS += [(randint(30, 254), randint(30, 254), randint(30, 254))
           for _ in range(N_CLASSES - 1)]


def load_image(filename):
    orig_image = cv2.imread(filename, 1)
    image = cv2.cvtColor(orig_image, cv2.COLOR_BGR2RGB)
    image = cv2.resize(image, (513, 513))
    image = image[..., ::-1]
    image = np.expand_dims(image, axis=0)
    image = (image - 127.5) / 127.5
    return orig_image, image


def run_inference(interpreter, image):
    import tflite_runtime.interpreter as tflite  # noqa: F401 (imported for side-effects check)
    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 / scale + zero_point).astype(np.int8)
    else:
        image = image.astype(np.float32)

    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='DeepLabV3 semantic segmentation demo')
    parser.add_argument('-m', '--model', default='original_model/deeplabv3_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, ...]
    end = time.time()
    print("Inference time: {:.1f} ms".format((end - start) * 1000))

    out = np.argmax(out, axis=-1)

    display_image = np.zeros((out.shape[0], out.shape[1], 3))
    for i in range(N_CLASSES):
        display_image[out == i] = COLORS[i]

    orig_size = orig_image.shape[0:2]
    mask_resized = cv2.resize(display_image, (orig_size[1], orig_size[0]))

    fig, ax = plt.subplots()
    ax.imshow(np.flip(orig_image, axis=-1))
    ax.imshow(mask_resized.astype(np.int8), alpha=0.7)
    ax.axis('off')

    if args.output:
        plt.savefig(args.output, bbox_inches='tight', pad_inches=0)
        print("Output saved to", args.output)
    else:
        plt.show()


if __name__ == '__main__':
    main()