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