"""Shared, checkpoint-compatible Mini U-Net. Parameters remain under 4M.""" import json import math from pathlib import Path import numpy as np import torch from torch import nn import torch.nn.functional as F from huggingface_hub import PyTorchModelHubMixin, snapshot_download from safetensors.torch import load_file, save_file 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), ) nn.init.zeros_(self.proj_out[-1].weight) 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): h, w = x.shape[-2:] x = F.pad(x, (0, (-w) % 8, 0, (-h) % 8), mode="replicate") 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)[..., :h, :w] def decode(self, logits, temperature: float = 0.38): if not math.isfinite(temperature) or temperature <= 0: raise ValueError("temperature must be finite and positive") probs_t = F.softmax(logits.float() / temperature, dim=1) return torch.einsum("bqhw,qc->bchw", probs_t, self.bin_centers) def load_model(source, revision=None, device="cpu"): path = Path(source) if not path.is_dir(): path = Path(snapshot_download(source, revision=revision, allow_patterns=["config.json", "model.safetensors"])) config = json.loads((path / "config.json").read_text()) state = load_file(str(path / "model.safetensors")) centers = torch.tensor(config["bin_centers"], dtype=torch.float32) if not torch.equal(centers, state["bin_centers"]): raise ValueError("Checkpoint config and state color bins differ; refusing ambiguous decode") model = SmallUNetColorizer(**config) model.load_state_dict(state, strict=True) model.to(device).eval() return model def save_model(model, path): path = Path(path); path.mkdir(parents=True, exist_ok=True) model.save_pretrained(path) # Mixin config can retain constructor bins; use the actual authoritative buffer. cfg = json.loads((path / "config.json").read_text()) cfg["bin_centers"] = model.bin_centers.detach().cpu().tolist() (path / "config.json").write_text(json.dumps(cfg, indent=2))