wzf19947's picture
first commit
525e655
Raw
History Blame Contribute Delete
3.5 kB
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)