File size: 1,276 Bytes
1a9eb8a
 
 
 
704fa80
 
 
 
 
 
1a9eb8a
704fa80
1a9eb8a
 
 
 
 
 
 
 
704fa80
 
1a9eb8a
 
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
"""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)