"""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()