Download example.py from nxp/midas-v2-imx: direct link, hf CLI and curl.
- Browser
- Download file 3.7 kB
-
https://huggingface.co/nxp/midas-v2-imx/resolve/main/example.py
- Command line
-
hf download hf://nxp/midas-v2-imx/example.py
-
curl -L -o example.py https://huggingface.co/nxp/midas-v2-imx/resolve/main/example.py
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}") | |