File size: 5,028 Bytes
1a9eb8a
 
 
 
 
 
704fa80
 
 
 
 
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
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
"""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()