ShadeNet-2-20M / model.py
singam96's picture
ShadeNet-2 20M: weights (EMA), ONNX fp32+int8, app + card
21e2ec3 verified
Raw History Blame Contribute Delete
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