Download model.py from singam96/ShadeNet-2-20M: direct link, hf CLI and curl.
- Browser
- Download file 6.56 kB
-
https://huggingface.co/singam96/ShadeNet-2-20M/resolve/main/model.py
- Command line
-
hf download hf://singam96/ShadeNet-2-20M/model.py
-
curl -L -o model.py https://huggingface.co/singam96/ShadeNet-2-20M/resolve/main/model.py
6.56 kB
| """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 | |