"""ShadeNet-2 generator (standalone, no Lightning dependency). ParallelUNet: dual parallel encoders (vanilla UNet + frozen MobileNetV2), fused decoder -> 8ch intrinsic maps in [-1, 1]: [0:3] albedo | [3:4] relative depth (0=near) | [4:7] normal | [7:8] shading """ import torch import torch.nn as nn import torchvision from torchvision.models.feature_extraction import create_feature_extractor OUT_CH = 8 WIDTH_MULT = 1.45 def _num_groups(ch: int) -> int: g = min(32, ch) while ch % g != 0: g -= 1 return g 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): 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), bias=False), nn.GroupNorm(g, out_ch), nn.ELU(inplace=True), nn.Conv2d(out_ch, out_ch, (3, 1), padding=(1, 0), 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): super().__init__() self.pool = nn.AvgPool2d(2) self.conv = DoubleConv(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): super().__init__() skip_ch = skip_ch or in_ch self.up = nn.ConvTranspose2d(in_ch, out_ch, 2, stride=2) self.conv = DoubleConv(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 ParallelUNet(nn.Module): def __init__(self, in_ch=3, out_ch=OUT_CH, dropout=0.0, width_mult=WIDTH_MULT): super().__init__() w = float(width_mult) def C(n): return max(8, int(round(n * w / 8.0) * 8)) self.u_inc = DoubleConv(in_ch, C(64)) self.u_down1 = Down(C(64), C(128)) self.u_down2 = Down(C(128), C(256)) self.u_down3 = Down(C(256), C(512)) self.u_down4 = Down(C(512), C(256)) self.u_down5 = Down(C(256), C(512)) 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("m_shift", 1.0 - 2.0 * mean) self.register_buffer("m_scale", 1.0 / (2.0 * std)) mbnet = torchvision.models.mobilenet_v2(weights="IMAGENET1K_V1") self.mobile = 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.mobile.parameters(): p.requires_grad = False self.bottleneck_fusion = DoubleConv(C(512) + 320, C(512)) self.dropout = nn.Dropout2d(p=dropout) if dropout > 0 else nn.Identity() self.mem_u1 = ChannelLinear(C(64)) self.mem_u2 = ChannelLinear(C(128)) self.mem_u3 = ChannelLinear(C(256)) self.mem_u4 = ChannelLinear(C(512)) self.mem_u5 = ChannelLinear(C(256)) self.mem_u6 = ChannelLinear(C(512)) self.mem_b = ChannelLinear(C(512)) self.mem_d0 = ChannelLinear(C(256)) self.mem_d1 = ChannelLinear(C(256)) self.mem_d2 = ChannelLinear(C(256)) self.mem_d3 = ChannelLinear(C(128)) self.mem_d4 = ChannelLinear(C(64)) 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.head = DoubleConv(C(64) + 3, C(32)) self.head_out = nn.Conv2d(C(32), out_ch, 3, padding=1) def forward(self, x): u1 = self.mem_u1(self.u_inc(x)) u2 = self.mem_u2(self.u_down1(u1)) u3 = self.mem_u3(self.u_down2(u2)) u4 = self.mem_u4(self.u_down3(u3)) u5 = self.mem_u5(self.u_down4(u4)) u6 = self.mem_u6(self.u_down5(u5)) x_imagenet = (x + self.m_shift) * self.m_scale mf = self.mobile(x_imagenet) m0, m1, m2, m3, m4 = mf["m0"], mf["m1"], mf["m2"], mf["m3"], mf["m4"] b = self.mem_b(self.dropout( self.bottleneck_fusion(torch.cat([u6, m4], dim=1)))) d0 = self.mem_d0(self.dropout(self.up0(b, torch.cat([u5, m3], dim=1)))) d1 = self.mem_d1(self.dropout(self.up1(d0, torch.cat([u4, m2], dim=1)))) d2 = self.mem_d2(self.dropout(self.up2(d1, torch.cat([u3, m1], dim=1)))) d3 = self.mem_d3(self.dropout(self.up3(d2, torch.cat([u2, m0], dim=1)))) d4 = self.mem_d4(self.dropout(self.up4(d3, u1))) h = self.head(torch.cat([d4, x], dim=1)) return torch.tanh(self.head_out(h)) def load_shadenet2(checkpoint_path, device="cpu", use_ema=True, width_mult=WIDTH_MULT) -> ParallelUNet: """Build the generator and load shadenet2.ckpt (Lightning or raw format). Strips the `generator.` prefix from Lightning checkpoints and applies the EMA shadow (validated best) 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