Download model.py from singam96/ShadeNet-3.2-5M: direct link, hf CLI and curl.
- Browser
- Download file 12.3 kB
-
https://huggingface.co/singam96/ShadeNet-3.2-5M/resolve/main/model.py
- Command line
-
hf download hf://singam96/ShadeNet-3.2-5M/model.py
-
curl -L -o model.py https://huggingface.co/singam96/ShadeNet-3.2-5M/resolve/main/model.py
12.3 kB
| """ShadeNet-3.2 generator (standalone, no Lightning dependency). | |
| ParallelUNet v3: dual parallel encoders (vanilla UNet + frozen MobileNetV2), | |
| depthwise-separable factorized convs, full H/32 bottleneck, and a patch | |
| dictionary output tail -> 8ch intrinsic maps in [-1, 1]: | |
| [0:3] albedo | [3:4] relative depth (0=near) | [4:7] normal | [7:8] shading | |
| """ | |
| import math | |
| import torch | |
| import torch.nn as nn | |
| import torchvision | |
| from torchvision.models.feature_extraction import create_feature_extractor | |
| OUT_CH = 8 | |
| WIDTH_MULT = 0.9 | |
| def _num_groups(ch: int) -> int: | |
| g = min(32, ch) | |
| while ch % g != 0: | |
| g -= 1 | |
| return g | |
| def _scale(width_mult): | |
| """Channel scaler: width_mult rounded to multiples of 8 (GroupNorm-safe).""" | |
| w = float(width_mult) | |
| def C(n): | |
| return max(8, int(round(n * w / 8.0) * 8)) | |
| return C | |
| class ChannelLinear(nn.Module): | |
| """Per-channel learnable affine: y = x * weight + bias.""" | |
| def __init__(self, channels: int, init_scale: float = 0.01): | |
| super().__init__() | |
| self.weight = nn.Parameter( | |
| torch.empty(1, channels, 1, 1).uniform_(-init_scale, init_scale) + 1.0 | |
| ) | |
| self.bias = nn.Parameter( | |
| torch.empty(1, channels, 1, 1).uniform_(-init_scale, init_scale) | |
| ) | |
| def forward(self, x): | |
| return x * self.weight + self.bias | |
| class DoubleConv(nn.Module): | |
| """Factorized (1,3)+(3,1) conv pair with GN+ELU, reflect padding.""" | |
| def __init__(self, in_ch, out_ch): | |
| super().__init__() | |
| g = _num_groups(out_ch) | |
| self.conv = nn.Sequential( | |
| nn.Conv2d(in_ch, out_ch, (1, 3), padding=(0, 1), | |
| padding_mode="reflect", bias=False), | |
| nn.GroupNorm(g, out_ch), | |
| nn.ELU(inplace=True), | |
| nn.Conv2d(out_ch, out_ch, (3, 1), padding=(1, 0), | |
| padding_mode="reflect", bias=False), | |
| nn.GroupNorm(g, out_ch), | |
| nn.ELU(inplace=True), | |
| ) | |
| def forward(self, x): | |
| return self.conv(x) | |
| class DSDoubleConv(nn.Module): | |
| """Depthwise-separable DoubleConv: same geometry, ~3x fewer params.""" | |
| def __init__(self, in_ch, out_ch): | |
| super().__init__() | |
| g = _num_groups(out_ch) | |
| self.conv = nn.Sequential( | |
| nn.Conv2d(in_ch, in_ch, (1, 3), padding=(0, 1), | |
| padding_mode="reflect", groups=in_ch, bias=False), | |
| nn.Conv2d(in_ch, out_ch, 1, bias=False), | |
| nn.GroupNorm(g, out_ch), | |
| nn.ELU(inplace=True), | |
| nn.Conv2d(out_ch, out_ch, (3, 1), padding=(1, 0), | |
| padding_mode="reflect", groups=out_ch, bias=False), | |
| nn.Conv2d(out_ch, out_ch, 1, bias=False), | |
| nn.GroupNorm(g, out_ch), | |
| nn.ELU(inplace=True), | |
| ) | |
| def forward(self, x): | |
| return self.conv(x) | |
| class Down(nn.Module): | |
| def __init__(self, in_ch, out_ch, block=DSDoubleConv): | |
| super().__init__() | |
| self.pool = nn.AvgPool2d(2) | |
| self.conv = block(in_ch, out_ch) | |
| def forward(self, x): | |
| return self.conv(self.pool(x)) | |
| class Up(nn.Module): | |
| def __init__(self, in_ch, out_ch, skip_ch=None, block=DSDoubleConv): | |
| super().__init__() | |
| skip_ch = skip_ch or in_ch | |
| self.up = nn.ConvTranspose2d(in_ch, out_ch, 2, stride=2) | |
| self.conv = block(out_ch + skip_ch, out_ch) | |
| def forward(self, x, skip): | |
| x = self.up(x) | |
| return self.conv(torch.cat([skip, x], dim=1)) | |
| class PatchDictionaryBias(nn.Module): | |
| """Soft patch-dictionary prior: tile -> softmax over atoms -> blend.""" | |
| def __init__(self, channels=OUT_CH, patch=16, n_atoms=1024, hidden=128, | |
| deblock=True, blur_radius=2): | |
| super().__init__() | |
| self.patch = int(patch) | |
| self.n_atoms = int(n_atoms) | |
| self.atoms = nn.Parameter( | |
| torch.zeros(n_atoms, channels, patch, patch)) | |
| self.desc = nn.Conv2d(channels, channels, 4, stride=4, | |
| groups=channels, bias=False) | |
| d = channels * (patch // 4) * (patch // 4) | |
| self.addr = nn.Sequential( | |
| nn.Linear(d, hidden), | |
| nn.ELU(inplace=True), | |
| nn.Linear(hidden, n_atoms), | |
| ) | |
| self.blend = nn.Conv2d(2 * channels, channels, 1, bias=True) | |
| self.blur_recon = bool(deblock) | |
| self.blur_radius = int(blur_radius) | |
| r = self.blur_radius | |
| row = torch.tensor([math.comb(2 * r, i) for i in range(2 * r + 1)], | |
| dtype=torch.float32) | |
| row = row / row.sum() | |
| k2 = (row[:, None] * row[None, :])[None, None].expand( | |
| channels, 1, -1, -1).contiguous() | |
| self.register_buffer("blur_k", k2) | |
| def _gauss_blur(self, t): | |
| r = self.blur_radius | |
| t = torch.nn.functional.pad(t, (r, r, r, r), mode="reflect") | |
| return torch.nn.functional.conv2d(t, self.blur_k, groups=t.shape[1]) | |
| def _deblock(self, recon, b, c, nh, nw, p): | |
| img = (recon.reshape(b, nh, nw, c, p, p) | |
| .permute(0, 3, 1, 4, 2, 5) | |
| .reshape(b, c, nh * p, nw * p)) | |
| r = self.blur_radius | |
| blurred = self._gauss_blur(img) | |
| H, W = img.shape[-2:] | |
| ih = torch.arange(H, device=img.device) | |
| iw = torch.arange(W, device=img.device) | |
| mh = ((ih % p) < r) | ((ih % p) >= p - r) | |
| mw = ((iw % p) < r) | ((iw % p) >= p - r) | |
| mask = (mh[:, None] | mw[None, :]).to(img.dtype)[None, None] | |
| img = img + mask * (blurred - img) | |
| return (img.reshape(b, c, nh, p, nw, p) | |
| .permute(0, 2, 4, 1, 3, 5) | |
| .reshape(b * nh * nw, c, p, p)) | |
| def forward(self, x): | |
| b, c, h, w = x.shape | |
| p = self.patch | |
| nh, nw = h // p, w // p | |
| tiles = (x.reshape(b, c, nh, p, nw, p) | |
| .permute(0, 2, 4, 1, 3, 5) | |
| .reshape(b * nh * nw, c, p, p)) | |
| feat = self.desc(tiles).flatten(1) | |
| wts = torch.softmax(self.addr(feat), dim=1) | |
| recon = torch.einsum("na,achw->nchw", wts, self.atoms) | |
| if self.blur_recon: | |
| recon = self._deblock(recon, b, c, nh, nw, p) | |
| both = torch.cat([tiles, recon], dim=1) | |
| out = self.blend(both).reshape(b, nh, nw, c, p, p) | |
| return out.permute(0, 3, 1, 4, 2, 5).reshape(b, c, nh * p, nw * p) | |
| class VanillaEncoder(nn.Module): | |
| def __init__(self, in_ch=3, width_mult=1.0): | |
| super().__init__() | |
| C = _scale(width_mult) | |
| self.inc = DoubleConv(in_ch, C(64)) | |
| self.down1 = Down(C(64), C(128)) | |
| self.down2 = Down(C(128), C(256)) | |
| self.down3 = Down(C(256), C(512)) | |
| self.down4 = Down(C(512), C(256)) | |
| self.down5 = Down(C(256), C(512)) | |
| self.mem1 = ChannelLinear(C(64)) | |
| self.mem2 = ChannelLinear(C(128)) | |
| self.mem3 = ChannelLinear(C(256)) | |
| self.mem4 = ChannelLinear(C(512)) | |
| self.mem5 = ChannelLinear(C(256)) | |
| self.mem6 = ChannelLinear(C(512)) | |
| def forward(self, x): | |
| u1 = self.mem1(self.inc(x)) | |
| u2 = self.mem2(self.down1(u1)) | |
| u3 = self.mem3(self.down2(u2)) | |
| u4 = self.mem4(self.down3(u3)) | |
| u5 = self.mem5(self.down4(u4)) | |
| u6 = self.mem6(self.down5(u5)) | |
| return u1, u2, u3, u4, u5, u6 | |
| class MobileEncoder(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1) | |
| std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1) | |
| self.register_buffer("shift", 1.0 - 2.0 * mean) | |
| self.register_buffer("scale", 1.0 / (2.0 * std)) | |
| mbnet = torchvision.models.mobilenet_v2(weights="IMAGENET1K_V1") | |
| self.trunk = create_feature_extractor( | |
| mbnet, | |
| return_nodes={ | |
| "features.0": "m0", "features.2": "m1", "features.6": "m2", | |
| "features.13": "m3", "features.17": "m4", | |
| }, | |
| ) | |
| for p in self.trunk.parameters(): | |
| p.requires_grad = False | |
| def forward(self, x): | |
| feats = self.trunk((x + self.shift) * self.scale) | |
| return feats["m0"], feats["m1"], feats["m2"], feats["m3"], feats["m4"] | |
| class FusionBottleneck(nn.Module): | |
| def __init__(self, width_mult=1.0): | |
| super().__init__() | |
| C = _scale(width_mult) | |
| self.m_proj = nn.Conv2d(320, C(96), 1, bias=False) | |
| self.fusion = DSDoubleConv(C(512) + C(96), C(512)) | |
| self.mem = ChannelLinear(C(512)) | |
| def forward(self, u6, m4): | |
| return self.mem(self.fusion(torch.cat([u6, self.m_proj(m4)], dim=1))) | |
| class FusedDecoder(nn.Module): | |
| def __init__(self, width_mult=1.0): | |
| super().__init__() | |
| C = _scale(width_mult) | |
| self.up0 = Up(C(512), C(256), skip_ch=C(256) + 96) | |
| self.up1 = Up(C(256), C(256), skip_ch=C(512) + 32) | |
| self.up2 = Up(C(256), C(256), skip_ch=C(256) + 24) | |
| self.up3 = Up(C(256), C(128), skip_ch=C(128) + 32) | |
| self.up4 = Up(C(128), C(64), skip_ch=C(64)) | |
| self.mem0 = ChannelLinear(C(256)) | |
| self.mem1 = ChannelLinear(C(256)) | |
| self.mem2 = ChannelLinear(C(256)) | |
| self.mem3 = ChannelLinear(C(128)) | |
| self.mem4 = ChannelLinear(C(64)) | |
| def forward(self, b, uni, mob): | |
| u1, u2, u3, u4, u5 = uni | |
| m0, m1, m2, m3 = mob | |
| d0 = self.mem0(self.up0(b, torch.cat([u5, m3], dim=1))) | |
| d1 = self.mem1(self.up1(d0, torch.cat([u4, m2], dim=1))) | |
| d2 = self.mem2(self.up2(d1, torch.cat([u3, m1], dim=1))) | |
| d3 = self.mem3(self.up3(d2, torch.cat([u2, m0], dim=1))) | |
| return self.mem4(self.up4(d3, u1)) | |
| class IntrinsicHead(nn.Module): | |
| def __init__(self, out_ch=OUT_CH, width_mult=1.0): | |
| super().__init__() | |
| C = _scale(width_mult) | |
| self.head = DoubleConv(C(64) + 3, C(32)) | |
| self.head_out = nn.Conv2d(C(32), out_ch, 3, padding=1, | |
| padding_mode="reflect") | |
| def forward(self, d4, x): | |
| return self.head_out(self.head(torch.cat([d4, x], dim=1))) | |
| class DictTail(nn.Module): | |
| """Patch-dictionary prior (blends internally, deblocked) -> tanh.""" | |
| def __init__(self, out_ch=OUT_CH, deblock=True, dict_atoms=1024): | |
| super().__init__() | |
| self.dict = PatchDictionaryBias(channels=out_ch, patch=16, | |
| n_atoms=dict_atoms, deblock=deblock) | |
| def forward(self, logits): | |
| return torch.tanh(self.dict(logits)) | |
| class ParallelUNet(nn.Module): | |
| def __init__(self, in_ch=3, out_ch=OUT_CH, width_mult=WIDTH_MULT, | |
| bias_res=384, deblock=True, dict_atoms=1024): | |
| super().__init__() | |
| self.encoder = VanillaEncoder(in_ch, width_mult) | |
| self.mobile = MobileEncoder() | |
| self.bottleneck = FusionBottleneck(width_mult) | |
| self.decoder = FusedDecoder(width_mult) | |
| self.head = IntrinsicHead(out_ch, width_mult) | |
| self.tail = DictTail(out_ch, deblock=deblock, dict_atoms=dict_atoms) | |
| def forward(self, x): | |
| u1, u2, u3, u4, u5, u6 = self.encoder(x) | |
| m0, m1, m2, m3, m4 = self.mobile(x) | |
| b = self.bottleneck(u6, m4) | |
| d4 = self.decoder(b, (u1, u2, u3, u4, u5), (m0, m1, m2, m3)) | |
| return self.tail(self.head(d4, x)) | |
| def load_shadenet32(checkpoint_path, device="cpu", use_ema=True, | |
| width_mult=WIDTH_MULT) -> ParallelUNet: | |
| """Build the generator and load shadenet32.ckpt (Lightning or raw format). | |
| Strips the `generator.` prefix from Lightning checkpoints and applies the | |
| EMA shadow unless use_ema=False. | |
| """ | |
| model = ParallelUNet(out_ch=OUT_CH, width_mult=width_mult) | |
| ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False) | |
| sd = ckpt.get("state_dict", ckpt) | |
| if any(k.startswith("generator.") for k in sd): | |
| sd = {k[len("generator."):]: v for k, v in sd.items() | |
| if k.startswith("generator.")} | |
| model.load_state_dict(sd, strict=False) | |
| if use_ema: | |
| ema = ckpt.get("ema_generator") or {} | |
| if ema: | |
| model.load_state_dict(ema, strict=False) | |
| print(f"Using EMA weights ({len(ema)} tensors).") | |
| else: | |
| print("No EMA shadow in checkpoint, using raw weights.") | |
| model.eval().to(device) | |
| return model | |