# /// script # requires-python = ">=3.10" # dependencies = [ # "torch>=2.3", # "torchvision>=0.18", # "huggingface_hub>=0.24", # "safetensors>=0.4", # "scikit-image>=0.22", # "pillow>=10.0", # "numpy", # ] # /// """ Colorize photos with a SmallUNetColorizer checkpoint trained by colorize_train.py (classification-head version: predicts a distribution over quantized Lab ab bins per pixel, decoded with an annealed mean). Usage: uv run colorize_image.py --model User-2468/mini-unet-colorizer photo.jpg uv run colorize_image.py --model ./local_checkpoint --temperature 0.2 a.jpg b.jpg --temperature is the main "vividness" knob: lower values weight the decode toward the most likely color bin (more saturated, can be a bit blotchy); higher values move toward the full expectation over the distribution (smoother, but can drift back toward desaturated -- the same hedging effect plain regression had). 0.38 (the default) is a reasonable middle ground. --saturation-boost is an optional *additional* post-decode multiplier on top of that, for further hand-tuning after picking a temperature. """ import argparse from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from huggingface_hub import PyTorchModelHubMixin from PIL import Image from skimage.color import lab2rgb, rgb2lab def double_conv(in_ch, out_ch): return nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) class DilatedContextBlock(nn.Module): def __init__(self, channels, mid_ch=96, dilations=(2, 4, 8)): super().__init__() self.proj_in = nn.Sequential( nn.Conv2d(channels, mid_ch, 1, bias=False), nn.BatchNorm2d(mid_ch), nn.ReLU(inplace=True), ) layers = [] for d in dilations: layers += [ nn.Conv2d(mid_ch, mid_ch, 3, padding=d, dilation=d, bias=False), nn.BatchNorm2d(mid_ch), nn.ReLU(inplace=True), ] self.dilated = nn.Sequential(*layers) self.proj_out = nn.Sequential( nn.Conv2d(mid_ch, channels, 1, bias=False), nn.BatchNorm2d(channels), ) self.relu = nn.ReLU(inplace=True) def forward(self, x): y = self.proj_in(x) y = self.dilated(y) y = self.proj_out(y) return self.relu(x + y) class SmallUNetColorizer( nn.Module, PyTorchModelHubMixin, pipeline_tag="image-to-image", license="apache-2.0", tags=["colorization", "unet", "image-to-image", "classification"], ): def __init__(self, bin_centers, in_ch: int = 1, base: int = 44, context_mid_ch: int = 96, context_dilations=(2, 4, 8)): super().__init__() self.in_ch, self.base = in_ch, base num_bins = len(bin_centers) self.num_bins = num_bins self.register_buffer("bin_centers", torch.tensor(bin_centers, dtype=torch.float32)) self.enc1 = double_conv(in_ch, base) self.enc2 = double_conv(base, base * 2) self.enc3 = double_conv(base * 2, base * 4) self.enc4 = double_conv(base * 4, base * 8) self.pool = nn.MaxPool2d(2) self.context = DilatedContextBlock(base * 8, mid_ch=context_mid_ch, dilations=tuple(context_dilations)) self.up3 = nn.ConvTranspose2d(base * 8, base * 4, 2, stride=2) self.dec3 = double_conv(base * 8, base * 4) self.up2 = nn.ConvTranspose2d(base * 4, base * 2, 2, stride=2) self.dec2 = double_conv(base * 4, base * 2) self.up1 = nn.ConvTranspose2d(base * 2, base, 2, stride=2) self.dec1 = double_conv(base * 2, base) self.out_conv = nn.Conv2d(base, num_bins, 1) def forward(self, x): e1 = self.enc1(x) e2 = self.enc2(self.pool(e1)) e3 = self.enc3(self.pool(e2)) e4 = self.context(self.enc4(self.pool(e3))) d3 = self.dec3(torch.cat([self.up3(e4), e3], dim=1)) d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1)) d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1)) return self.out_conv(d1) def decode(self, logits, temperature: float = 0.38): logp = F.log_softmax(logits, dim=1) probs_t = F.softmax(logp / temperature, dim=1) return torch.einsum("bqhw,qc->bchw", probs_t, self.bin_centers) def colorize(model, img, size, temperature, saturation_boost, device): img = img.convert("RGB").resize((size, size)) arr = np.asarray(img).astype(np.float32) / 255.0 lab = rgb2lab(arr).astype(np.float32) L = torch.from_numpy(lab[:, :, 0:1] / 50.0 - 1.0).permute(2, 0, 1)[None].to(device) with torch.no_grad(): logits = model(L) ab = model.decode(logits, temperature=temperature)[0].permute(1, 2, 0).cpu().numpy() ab = np.clip(ab * saturation_boost, -128, 127) L_out = (L[0, 0].cpu().numpy() + 1.0) * 50.0 lab_out = np.concatenate([L_out[:, :, None], ab], axis=-1) rgb_out = np.clip(lab2rgb(lab_out), 0, 1) return Image.fromarray((rgb_out * 255).astype(np.uint8)) def main(): p = argparse.ArgumentParser(description=__doc__) p.add_argument("images", nargs="+", help="Path(s) to input photo(s)") p.add_argument("--model", required=True, help="Hub model id or local checkpoint path") p.add_argument("--size", type=int, default=256, help="Resize input to this square size") p.add_argument("--temperature", type=float, default=0.38, help="Annealed-mean decode temperature. Lower = more vivid/mode-like, " "higher = smoother but can desaturate again. Try 0.15-0.6.") p.add_argument("--saturation-boost", type=float, default=1.0, help="Extra multiplier on the decoded ab, applied after temperature. " "1.0 = no extra boost.") p.add_argument("--output-dir", default="./colorized") args = p.parse_args() device = "cuda" if torch.cuda.is_available() else "cpu" print(f"Loading {args.model} on {device} ...") model = SmallUNetColorizer.from_pretrained(args.model).to(device).eval() print(f"({model.num_bins} color bins)") out_dir = Path(args.output_dir) out_dir.mkdir(parents=True, exist_ok=True) for path in args.images: img = Image.open(path) result = colorize(model, img, args.size, args.temperature, args.saturation_boost, device) out_path = out_dir / f"{Path(path).stem}_colorized.png" result.save(out_path) print(f"{path} -> {out_path}") if __name__ == "__main__": main()