mini-unet-colorizer / inference.py
User-2468's picture
Release v3.0: selected sub-4M semantic colorizer, inference and verified ONNX
1a9eb8a verified
Raw History Blame
5.03 kB
"""Aspect-preserving colourisation with a compact semantic model.
Input and output are PIL images. No teacher, critic or second learned model is
loaded. Original lightness and alpha are retained; out-of-gamut chroma is reduced.
"""
import json
import math
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image, ImageOps
from skimage.color import rgb2lab
from semantic_model import load_semantic
MAX_PIXELS = 12_000_000
def load_colorizer(path, device='cpu'):
model = load_semantic(path, device)
count = sum(p.numel() for p in model.parameters())
if count >= 4_000_000:
raise ValueError('Model exceeds the four-million-parameter limit')
return model
def box_mean(x, r):
return F.avg_pool2d(x, 2*r+1, 1, r, count_include_pad=False)
def prepare(image, size=256):
if not isinstance(image, Image.Image):
raise TypeError('Expected a PIL image')
if not 128 <= int(size) <= 512:
raise ValueError('Input size must be between 128 and 512')
image = ImageOps.exif_transpose(image)
if image.width * image.height > MAX_PIXELS:
raise ValueError('Please resize the image to at most 12 megapixels')
alpha = np.asarray(image.getchannel('A')).copy() if 'A' in image.getbands() else None
rgb = np.asarray(image.convert('RGB'), dtype=np.float32) / 255
light = rgb2lab(rgb)[..., 0].astype(np.float32)
scale = min(int(size)/max(image.size), 1.)
shape = (max(8, round(image.height*scale)), max(8, round(image.width*scale)))
x = torch.from_numpy(light)[None,None]/50-1
small = F.interpolate(x, size=shape, mode='bilinear', align_corners=False, antialias=True)
return light, alpha, small
@torch.inference_mode()
def chroma_coefficients(model, small, radius=8):
if radius not in (0,4,8,12,16):
raise ValueError('Unsupported smoothing radius')
device = next(model.parameters()).device
L = small.to(device)
ab = model(L).float()
if not torch.isfinite(ab).all():
raise RuntimeError('Model returned non-finite colours')
if radius == 0:
return torch.zeros_like(ab).cpu(), ab.cpu()
guide = (L.float()+1)/2
mi, mp = box_mean(guide,radius), box_mean(ab,radius)
var = (box_mean(guide*guide,radius)-mi*mi).clamp_min(0)
cov = box_mean(guide*ab,radius)-mi*mp
a = cov/(var+.001)
b = mp-a*mi
# Coefficients are upsampled, then evaluated against original-resolution L.
return box_mean(a,radius).cpu(), box_mean(b,radius).cpu()
def _linear_rgb(light, ab):
fy = (light+16)/116
f = np.stack([fy+ab[...,0]/500,fy,fy-ab[...,1]/200],axis=-1)
xyz = np.where(f>6/29,f**3,(f-4/29)*(3*(6/29)**2))
xyz *= np.array([.95047,1.,1.08883],np.float32)
matrix = np.array([[3.24048134,-1.53715152,-.49853633],[-.96925495,1.87599,.04155593],[.05564664,-.20404134,1.05731107]],np.float32)
return xyz @ matrix.T
def render(light, alpha, coefficients, saturation=1.):
saturation = float(saturation)
if not math.isfinite(saturation) or not 0 <= saturation <= 1.5:
raise ValueError('Colour strength must be between 0 and 1.5')
a,b = [F.interpolate(v.float(),size=light.shape,mode='bilinear',align_corners=False)[0].permute(1,2,0).numpy() for v in coefficients]
ab = (a*(light[...,None]/100)+b)*saturation
# Binary-search chroma compression retains Lab hue and lightness.
linear = _linear_rgb(light,ab)
invalid = ((linear < -1e-5)|(linear > 1+1e-5)).any(-1)
if invalid.any():
L = light[invalid]; colors=ab[invalid];lo=np.zeros(len(L),np.float32);hi=np.ones(len(L),np.float32)
for _ in range(9):
mid=(lo+hi)/2; candidate=_linear_rgb(L,colors*mid[:,None])
valid=((candidate>=-1e-5)&(candidate<=1+1e-5)).all(-1)
lo=np.where(valid,mid,lo);hi=np.where(valid,hi,mid)
ab[invalid]=colors*lo[:,None]
linear[invalid]=_linear_rgb(L,ab[invalid])
linear=np.clip(linear,0,1)
rgb=np.where(linear<=.0031308,12.92*linear,1.055*np.power(linear,1/2.4)-.055)
pixels=np.uint8(np.clip(np.rint(rgb*255),0,255))
if alpha is not None: pixels=np.concatenate([pixels,alpha[...,None]],axis=-1)
return Image.fromarray(pixels)
def colorize(model, image, size=256, radius=8, saturation=1.):
light,alpha,small=prepare(image,size)
return render(light,alpha,chroma_coefficients(model,small,radius),saturation)
def main():
import argparse
parser=argparse.ArgumentParser(description='Compact photo colouriser')
parser.add_argument('input');parser.add_argument('output')
parser.add_argument('--model',default='.');parser.add_argument('--device',default='cpu')
parser.add_argument('--size',type=int,default=256);parser.add_argument('--saturation',type=float,default=1.)
args=parser.parse_args()
model=load_colorizer(args.model,args.device)
with Image.open(args.input) as image:
colorize(model,image,args.size,saturation=args.saturation).save(args.output)
if __name__=='__main__':main()