mini-unet-colorizer / quantized /base93-v3 /semantic_model.py
User-2468's picture
Add calibrated Base93 production-v3 export under 5 MB with decoders and visual validation
915de37 verified
Raw History Blame Contribute Delete
5.57 kB
"""Compact pretrained semantic colorizer; dense and shared-palette variants."""
import json
from pathlib import Path
import torch
from torch import nn
import torch.nn.functional as F
from torchvision.models import mobilenet_v3_large, MobileNet_V3_Large_Weights
from safetensors.torch import save_file, load_file
class QueryBlock(nn.Module):
def __init__(self,d=96):
super().__init__()
self.self_attn=nn.MultiheadAttention(d,4,batch_first=True,dropout=0)
self.cross_attn=nn.MultiheadAttention(d,4,batch_first=True,dropout=0)
self.norms=nn.ModuleList([nn.LayerNorm(d) for _ in range(3)])
self.ff=nn.Sequential(nn.Linear(d,2*d),nn.GELU(),nn.Linear(2*d,d))
def forward(self,q,memory):
x=self.norms[0](q);q=q+self.self_attn(x,x,x,need_weights=False)[0]
x=self.norms[1](q);q=q+self.cross_attn(x,memory,memory,need_weights=False)[0]
return q+self.ff(self.norms[2](q))
def refine(d):
return nn.Sequential(nn.Conv2d(d,d,3,padding=1,bias=False),nn.GroupNorm(8,d),nn.SiLU())
class SemanticColorizer(nn.Module):
def __init__(self,head='palette',pretrained=False,width=128,queries=16):
super().__init__()
if head not in ['palette','dense']:raise ValueError(head)
self.config={'architecture':'SemanticColorizer','head':head,'width':width,'queries':queries,'format_version':1}
self.encoder=mobilenet_v3_large(weights=MobileNet_V3_Large_Weights.IMAGENET1K_V2 if pretrained else None,progress=False).features
self.lateral=nn.ModuleList([nn.Conv2d(c,width,1) for c in [24,40,112,960]])
self.refine=nn.ModuleList([refine(width) for _ in range(3)])
self.register_buffer('rgb_mean',torch.tensor([.485,.456,.406]).view(1,3,1,1))
self.register_buffer('rgb_std',torch.tensor([.229,.224,.225]).view(1,3,1,1))
if head=='palette':
self.queries=nn.Parameter(torch.randn(queries,width)*.2)
self.query_blocks=nn.ModuleList([QueryBlock(width) for _ in range(2)])
self.memory_norm=nn.LayerNorm(width)
self.query_norm=nn.LayerNorm(width)
self.pixel=nn.Conv2d(width,width,1)
self.palette=nn.Sequential(nn.Linear(width,width),nn.GELU(),nn.Linear(width,2))
self.residual=nn.Conv2d(width,2,1)
nn.init.normal_(self.palette[-1].weight,std=.01);nn.init.zeros_(self.palette[-1].bias)
nn.init.zeros_(self.residual.weight);nn.init.zeros_(self.residual.bias)
else:
self.dense=nn.Sequential(refine(width),nn.Conv2d(width,2,1))
nn.init.normal_(self.dense[-1].weight,std=.01);nn.init.zeros_(self.dense[-1].bias)
count=sum(p.numel() for p in self.parameters())
if count>=4_000_000:raise ValueError(f'Parameter budget exceeded: {count}')
@staticmethod
def neutral_rgb(L):
light=(L.float()*50+50).clamp(0,100)
y=torch.where(light>8,((light+16)/116)**3,light/903.296296)
g=torch.where(y<=.0031308,12.92*y,1.055*y.clamp_min(1e-8).pow(1/2.4)-.055)
return g.expand(-1,3,-1,-1)
def forward(self,L):
h,w=L.shape[-2:]
x=F.pad(self.neutral_rgb(L),(0,(-w)%32,0,(-h)%32),mode='replicate')
x=(x-self.rgb_mean)/self.rgb_std
features=[]
for i,layer in enumerate(self.encoder):
x=layer(x)
if i in [3,6,12,16]:features.append(x)
projected=[layer(f) for layer,f in zip(self.lateral,features)]
x=projected[-1]
for i in range(2,-1,-1):
x=self.refine[2-i](F.interpolate(x,size=projected[i].shape[-2:],mode='bilinear',align_corners=False)+projected[i])
if self.config['head']=='palette':
memory=torch.cat([F.adaptive_avg_pool2d(f,(8,8)).flatten(2).transpose(1,2) for f in projected[1:]],1)
memory=self.memory_norm(memory)
q=self.queries[None].expand(L.shape[0],-1,-1)
for block in self.query_blocks:q=block(q,memory)
q=self.query_norm(q)
palette=80*torch.tanh(self.palette(q))
masks=torch.einsum('bqd,bdhw->bqhw',q,self.pixel(x))/(self.config['width']**.5)
weights=F.softmax(masks.float(),dim=1)
ab=torch.einsum('bqhw,bqc->bchw',weights,palette.float())+2*torch.tanh(self.residual(x).float())
else:ab=80*torch.tanh(self.dense(x).float())
return F.interpolate(ab,size=(x.shape[-2]*4,x.shape[-1]*4),mode='bilinear',align_corners=False)[...,:h,:w]
def decode(self,z,temperature=.38):return z.float()
def save_semantic(model,path):
path=Path(path);path.mkdir(parents=True,exist_ok=True)
save_file({k:v.detach().cpu().contiguous() for k,v in model.state_dict().items()},str(path/'model.safetensors'))
(path/'config.json').write_text(json.dumps(model.config,indent=2))
(path/'README.md').write_text('# Experimental semantic colorizer\n\nHead: '+model.config['head']+'. Parameters: '+str(sum(p.numel() for p in model.parameters()))+'.\n\nNot approved for production. Requires semantic_model.py; incompatible with the old U-Net loader. Input is Lab lightness normalized to [-1,1]; output is Lab ab. See the run protocol, provenance, selection and visual comparisons. Predictions are plausible colors, not recovered historical truth.\n')
def load_semantic(path,device='cpu'):
path=Path(path);cfg=json.loads((path/'config.json').read_text())
model=SemanticColorizer(**{k:cfg[k] for k in ['head','width','queries']})
model.load_state_dict(load_file(str(path/'model.safetensors')),strict=True)
return model.to(device).eval()