File size: 2,871 Bytes
4dca198
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Run GridDehazeNet ONNX inference on a single image.
Output: side-by-side image (Hazy | Dehazed), result resized to original size.
"""
import argparse
import os

import numpy as np
import onnxruntime as ort
from PIL import Image, ImageDraw, ImageFont


def _get_font(size):
    for path in (
        "/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf",
        "/usr/share/fonts/truetype/liberation/LiberationSans-Bold.ttf",
        "/usr/share/fonts/truetype/freefont/FreeSansBold.ttf",
    ):
        if os.path.exists(path):
            return ImageFont.truetype(path, size)
    return ImageFont.load_default()


def _label_image(image, text, font):
    draw = ImageDraw.Draw(image)
    bbox = draw.textbbox((0, 0), text, font=font)
    w, h = bbox[2] - bbox[0], bbox[3] - bbox[1]
    pad = 6
    draw.rectangle([(0, 0), (w + pad * 2, h + pad * 2)], fill=(0, 0, 0, 180))
    draw.text((pad, pad), text, fill=(255, 255, 255), font=font)


def _tensor_to_image(tensor):
    arr = np.squeeze(tensor, axis=0).transpose(1, 2, 0)
    arr = np.clip(arr, 0.0, 1.0)
    return Image.fromarray((arr * 255.0 + 0.5).astype(np.uint8))


def main():
    parser = argparse.ArgumentParser(description="GridDehazeNet ONNX single-image inference")
    parser.add_argument("-onnx", default="griddehazenet_outdoor_1x3x480x640_npu_safe.onnx")
    parser.add_argument("-input", default="data/test/SOTS/outdoor/hazy/0001_0.8_0.2.jpg", help="image file")
    parser.add_argument("-output", default="onnx_result_concat.png")
    parser.add_argument("-height", type=int, default=480)
    parser.add_argument("-width", type=int, default=640)
    args = parser.parse_args()

    font = _get_font(24)
    session = ort.InferenceSession(args.onnx, providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
    input_name = session.get_inputs()[0].name
    output_name = session.get_outputs()[0].name

    original = Image.open(args.input).convert("RGB")
    orig_w, orig_h = original.size

    image = Image.open(args.input).convert("RGB")
    if image.size != (args.width, args.height):
        image = image.resize((args.width, args.height), Image.BICUBIC)
    arr = np.asarray(image).astype(np.float32) / 255.0
    arr = (arr - 0.5) / 0.5
    arr = arr.transpose(2, 0, 1)[None, ...]

    out = session.run([output_name], {input_name: arr})[0]
    dehazed = _tensor_to_image(out)
    if dehazed.size != (orig_w, orig_h):
        dehazed = dehazed.resize((orig_w, orig_h), Image.BICUBIC)

    hazy_labeled = original.copy()
    dehazed_labeled = dehazed.copy()
    _label_image(hazy_labeled, "Hazy", font)
    _label_image(dehazed_labeled, "Dehazed", font)

    concat = Image.new("RGB", (orig_w * 2, orig_h))
    concat.paste(hazy_labeled, (0, 0))
    concat.paste(dehazed_labeled, (orig_w, 0))
    concat.save(args.output)
    print(args.output)


if __name__ == "__main__":
    main()