ShadeNet-3.2-5M / model.py
singam96's picture
ShadeNet-3.2-5M (model): weights, ONNX, app + card
2d1eb0d verified
Raw History Blame Contribute Delete
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