User-2468's picture
Update validated colorizer weights, spatial decoder and deployment artifacts
704fa80 verified
Raw History Blame
4.94 kB
"""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))