File size: 3,699 Bytes
48094aa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
107
108
109
110
111
112
113
114
#!/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}")