"""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