deeplabv3-imx / example.py
NXP
Release 2.0
7346868
Raw History Blame Contribute Delete
3.28 kB
#!/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()