mini-unet-colorizer / inference.py
User-2468's picture
Update validated colorizer weights, spatial decoder and deployment artifacts
704fa80 verified
Raw History Blame
2.98 kB
"""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()