"""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()