Download inference.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 5.03 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/1cca2ddeb5f7c0e6167c452f52541c7684182248/inference.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer@1cca2ddeb5f7c0e6167c452f52541c7684182248/inference.py
-
curl -L -o inference.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/1cca2ddeb5f7c0e6167c452f52541c7684182248/inference.py
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 | |
| 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() | |