Download quantized/base93-v3/semantic_model.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 5.57 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/main/quantized/base93-v3/semantic_model.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer/quantized/base93-v3/semantic_model.py
-
curl -L -o semantic_model.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/main/quantized/base93-v3/semantic_model.py
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}') | |
| 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() | |