mini-unet-colorizer / colorize_onnx.py
User-2468's picture
Release v3.0: selected sub-4M semantic colorizer, inference and verified ONNX
1a9eb8a verified
Raw History Blame
1.28 kB
"""Minimal ONNX-only image example; fixed square network input.
For aspect-preserving guided upsampling and gamut compression use inference.py.
"""
import argparse
import numpy as np
import onnxruntime as ort
from PIL import Image,ImageOps
from skimage.color import rgb2lab,lab2rgb
def colorize(image,model='colorizer.onnx'):
image=ImageOps.exif_transpose(image).convert('RGB')
if image.width*image.height>12_000_000:raise ValueError('Maximum 12 megapixels')
light=rgb2lab(np.asarray(image,np.float32)/255)[...,0].astype(np.float32)
resized=np.asarray(Image.fromarray(light).resize((256,256),Image.Resampling.BILINEAR))
session=ort.InferenceSession(model,providers=['CPUExecutionProvider'])
ab=session.run(['ab'],{'L':(resized[None,None]/50-1).astype(np.float32)})[0][0]
ab=np.stack([np.asarray(Image.fromarray(channel).resize(image.size,Image.Resampling.BILINEAR)) for channel in ab],-1)
rgb=np.clip(lab2rgb(np.concatenate([light[...,None],ab],-1)),0,1)
return Image.fromarray(np.uint8(np.rint(rgb*255)))
if __name__=='__main__':
p=argparse.ArgumentParser();p.add_argument('input');p.add_argument('output');p.add_argument('--model',default='colorizer.onnx');a=p.parse_args()
colorize(Image.open(a.input),a.model).save(a.output)