#!/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}")