File size: 3,495 Bytes
525e655
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import argparse
import os
import time

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


def preprocess(image_path, height, width, resize):
	original_image = Image.open(image_path).convert('RGB')
	original_size = original_image.size
	if resize:
		model_image = original_image.resize((width, height), Image.BILINEAR)
	else:
		if original_image.size[0] < width or original_image.size[1] < height:
			raise ValueError('Image is smaller than ONNX input size: {}'.format(image_path))
		model_image = original_image.crop((0, 0, width, height))

	image_np = np.asarray(model_image).astype(np.float32) / 255.0
	image_np = image_np.transpose(2, 0, 1)[None, :, :, :]
	return original_image, original_size, image_np.astype(np.float32)


def postprocess(enhanced_np, original_size):
	enhanced_np = np.clip(enhanced_np[0].transpose(1, 2, 0), 0.0, 1.0)
	enhanced_image = Image.fromarray((enhanced_np * 255.0).astype(np.uint8))
	return enhanced_image.resize(original_size, Image.BILINEAR)


def add_label(image, text):
	label_height = 32
	canvas = Image.new('RGB', (image.width, image.height + label_height), color=(0, 0, 0))
	canvas.paste(image, (0, label_height))
	draw = ImageDraw.Draw(canvas)
	draw.text((10, 8), text, fill=(255, 255, 255))
	return canvas


def save_compare(original_image, enhanced_image, result_path):
	original_labeled = add_label(original_image, 'Original')
	enhanced_labeled = add_label(enhanced_image, 'Enhanced')
	compare_image = Image.new('RGB', (original_labeled.width + enhanced_labeled.width, original_labeled.height))
	compare_image.paste(original_labeled, (0, 0))
	compare_image.paste(enhanced_labeled, (original_labeled.width, 0))
	result_dir = os.path.dirname(result_path)
	if result_dir and not os.path.exists(result_dir):
		os.makedirs(result_dir)
	compare_image.save(result_path)


def build_onnx_session(onnx_path):
	providers = ['CUDAExecutionProvider', 'CPUExecutionProvider']
	available = ort.get_available_providers()
	providers = [provider for provider in providers if provider in available]
	return ort.InferenceSession(onnx_path, providers=providers)


def infer(config):
	if not os.path.isfile(config.input):
		raise ValueError('Input image does not exist: {}'.format(config.input))

	session = build_onnx_session(config.onnx)
	input_name = session.get_inputs()[0].name
	input_shape = session.get_inputs()[0].shape
	height = int(input_shape[2]) if config.height <= 0 else config.height
	width = int(input_shape[3]) if config.width <= 0 else config.width

	original_image, original_size, input_np = preprocess(config.input, height, width, bool(config.resize))
	start = time.time()
	onnx_outputs = session.run(None, {input_name: input_np})
	onnx_enhanced = onnx_outputs[0]
	elapsed = time.time() - start

	enhanced_image = postprocess(onnx_enhanced, original_size)
	save_compare(original_image, enhanced_image, config.output)

	print('Input image:', config.input)
	print('Output image:', config.output)
	print('ONNX time:', elapsed)


if __name__ == '__main__':
	parser = argparse.ArgumentParser()
	parser.add_argument('--onnx', type=str, default='zerodcepp_512_sf8.onnx')
	parser.add_argument('--input', type=str, default='data/test_data/real/11_0_.png')
	parser.add_argument('--output', type=str, default='onnx_res.jpg')
	parser.add_argument('--height', type=int, default=512)
	parser.add_argument('--width', type=int, default=512)
	parser.add_argument('--resize', type=int, default=1)
	config = parser.parse_args()

	infer(config)