"""Aspect-preserving inference; infer chroma globally and retain original luminance.""" import argparse import math from pathlib import Path import numpy as np from PIL import Image, ImageOps import torch import torch.nn.functional as F from skimage.color import rgb2lab, lab2rgb from model import load_model from spatial import guided_chroma @torch.inference_mode() def colorize(model, image, size=256, temperature=0.38, saturation=1.0, flip_tta=False, guided_radius=8, guided_epsilon=.001): if size < 8 or not math.isfinite(saturation) or saturation < 0: raise ValueError("size must be >=8 and saturation finite and nonnegative") image = ImageOps.exif_transpose(image).convert("RGB") rgb = np.asarray(image, dtype=np.float32) / 255.0 luminance = rgb2lab(rgb)[..., 0].astype(np.float32) h, w = luminance.shape scale = min(size / max(h, w), 1.0) target = (max(8, round(h*scale)), max(8, round(w*scale))) device = next(model.parameters()).device L = torch.from_numpy(luminance)[None, None].to(device) / 50 - 1 small = F.interpolate(L, size=target, mode="bilinear", align_corners=False, antialias=True) logits = model(small) if flip_tta: logits = (logits + model(small.flip(-1)).flip(-1)) * 0.5 ab = model.decode(logits, temperature) ab = guided_chroma(small, ab, guided_radius, guided_epsilon) ab = F.interpolate(ab, size=(h,w), mode="bilinear", align_corners=False) ab = ab[0].permute(1,2,0).cpu().numpy() * saturation lab = np.concatenate([luminance[...,None], ab], axis=-1) result = np.clip(lab2rgb(lab), 0, 1) return Image.fromarray(np.rint(result * 255).astype(np.uint8)) def main(): p = argparse.ArgumentParser(__doc__) p.add_argument("images", nargs="+") p.add_argument("--model", required=True) p.add_argument("--revision", default=None) p.add_argument("--size", type=int, default=256) p.add_argument("--temperature", type=float, default=.38) p.add_argument("--saturation", type=float, default=1) p.add_argument("--flip-tta", action="store_true") p.add_argument("--guided-radius",type=int,default=8,help="Chroma smoothing radius at model resolution;0 disables") p.add_argument("--guided-epsilon",type=float,default=.001) p.add_argument("--output-dir", default="colorized") p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") a = p.parse_args() torch.set_num_threads(min(torch.get_num_threads(),4)) model = load_model(a.model, a.revision, a.device) out = Path(a.output_dir); out.mkdir(parents=True, exist_ok=True) for file in a.images: with Image.open(file) as im: result = colorize(model, im, a.size, a.temperature, a.saturation, a.flip_tta, a.guided_radius,a.guided_epsilon) dest = out / (Path(file).stem + "_colorized.png") result.save(dest); print(dest) if __name__ == "__main__": main()