midas-v2-imx / example.py
NXP
Release 2.0
48094aa
Raw History Blame Contribute Delete
3.7 kB
#!/usr/bin/env python3
# Copyright 2023-2024,2026 NXP
# SPDX-License-Identifier: MIT
from argparse import ArgumentParser
import cv2
import numpy as np
import matplotlib.pyplot as plt
import tensorflow as tf
import time
WIDTH = 256
HEIGHT = 256
parser = ArgumentParser(description="MiDaS v2.1 Small depth estimation inference example")
parser.add_argument(
"-m", "--model",
help="Path to a .tflite model file.",
type=str,
default="original_model/midas_v2_1_small_quant.tflite")
parser.add_argument(
"-i", "--input",
help="Path to an input image file.",
type=str,
default="example_input.jpg")
parser.add_argument(
"-o", "--output",
help="Path to save the output depth map image.",
type=str,
default="example_output.jpg")
args = parser.parse_args()
def load_image(filename):
"""Load and preprocess an image for MiDaS inference."""
image = cv2.imread(filename, cv2.IMREAD_COLOR)
if image is None:
raise FileNotFoundError(f"Could not read image: {filename}")
orig_height, orig_width = image.shape[:2]
image = cv2.resize(image, (WIDTH, HEIGHT), interpolation=cv2.INTER_CUBIC)
image = image.astype(np.float32) / 255.0
image_input = np.expand_dims(image, axis=0)
return image_input, orig_height, orig_width
def run_inference(interpreter, image):
"""Run TFLite inference and return the raw output tensor."""
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()
# Handle quantized (int8) models: scale input if needed
input_scale = input_details[0].get('quantization_parameters', {}).get('scales', [1.0])
input_zero_point = input_details[0].get('quantization_parameters', {}).get('zero_points', [0])
if input_details[0]['dtype'] == np.int8:
scale = input_scale[0] if len(input_scale) > 0 else 1.0
zero_point = input_zero_point[0] if len(input_zero_point) > 0 else 0
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 output if int8
if output_details[0]['dtype'] == np.int8:
out_scale = output_details[0].get('quantization_parameters', {}).get('scales', [1.0])
out_zp = output_details[0].get('quantization_parameters', {}).get('zero_points', [0])
scale = out_scale[0] if len(out_scale) > 0 else 1.0
zero_point = out_zp[0] if len(out_zp) > 0 else 0
out = (out.astype(np.float32) - zero_point) * scale
return out.astype(np.float32)
def post_process(output, orig_height, orig_width):
"""Resize depth map to original image size and normalize to [0, 1]."""
disp = output[0, ..., 0] if output.ndim == 4 else output[0]
disp = cv2.resize(disp, (orig_width, orig_height), interpolation=cv2.INTER_CUBIC)
disp_min = disp.min()
disp_max = disp.max()
if disp_max - disp_min > 1e-6:
disp = (disp - disp_min) / (disp_max - disp_min)
else:
disp.fill(0.5)
return disp
# Load model
interpreter = tf.lite.Interpreter(args.model)
interpreter.allocate_tensors()
# Load and preprocess input image
image_input, orig_height, orig_width = load_image(args.input)
# Run inference
start = time.time()
out = run_inference(interpreter, image_input)
elapsed = time.time() - start
print(f"Inference time: {elapsed * 1000:.1f} ms")
# Post-process and save result
disp = post_process(out, orig_height, orig_width)
plt.imsave(args.output, disp, vmin=0, vmax=1, cmap='inferno')
print(f"Depth map saved to: {args.output}")