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