Download decision_model.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 5.7 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/1cca2ddeb5f7c0e6167c452f52541c7684182248/decision_model.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer@1cca2ddeb5f7c0e6167c452f52541c7684182248/decision_model.py
-
curl -L -o decision_model.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/1cca2ddeb5f7c0e6167c452f52541c7684182248/decision_model.py
5.7 kB
| """Single / multiple coherent palettes, optional compact global-attention encoder. | |
| Original experiment inspired by diverse colorization and multiple-choice learning; | |
| not a reproduction. All deployed parameters, including the selector, are counted. | |
| """ | |
| import json | |
| from pathlib import Path | |
| import torch | |
| from torch import nn | |
| import torch.nn.functional as F | |
| from safetensors.torch import save_file, load_file | |
| from semantic_model import SemanticColorizer | |
| class DecisionColorizer(SemanticColorizer): | |
| def __init__(self, modes=1, backbone='mobilenet', pretrained=False, | |
| encoder_revision=None, **kwargs): | |
| super().__init__(pretrained=False, **kwargs) | |
| self.modes = modes | |
| self.backbone = backbone | |
| d = self.config['width']; q = self.config['queries'] | |
| if backbone == 'mobilevit': | |
| import timm | |
| from huggingface_hub import hf_hub_download | |
| name = 'mobilevitv2_075.cvnets_in1k' | |
| overlay = None | |
| if pretrained: | |
| assert encoder_revision, 'Pin pretrained encoder revision' | |
| filename = hf_hub_download('timm/'+name, 'model.safetensors', revision=encoder_revision) | |
| overlay = {'file': filename, 'hf_hub_id': None} | |
| self.encoder = timm.create_model(name, pretrained=pretrained, | |
| pretrained_cfg_overlay=overlay, features_only=True, out_indices=(1,2,3,4)) | |
| self.lateral = nn.ModuleList([nn.Conv2d(c,d,1) for c in self.encoder.feature_info.channels()]) | |
| cfg = self.encoder.pretrained_cfg | |
| self.rgb_mean.copy_(torch.tensor(cfg['mean']).view(1,3,1,1)) | |
| self.rgb_std.copy_(torch.tensor(cfg['std']).view(1,3,1,1)) | |
| if modes > 1: | |
| self.mode_embeddings = nn.Parameter(torch.randn(modes,d)*.03) | |
| self.mode_offsets = nn.Parameter(torch.zeros(modes,q,2)) | |
| # Whole-image palette alternatives; no independently sampled pixels. | |
| with torch.no_grad(): | |
| a=torch.arange(modes)*2*torch.pi/modes | |
| self.mode_offsets[:,:,0]=.025*a.cos()[:,None] | |
| self.mode_offsets[:,:,1]=.025*a.sin()[:,None] | |
| self.mode_score = nn.Linear(d,modes) | |
| nn.init.zeros_(self.mode_score.weight);nn.init.zeros_(self.mode_score.bias) | |
| self.config.update(architecture='DecisionColorizer', format_version=3, | |
| modes=modes, backbone=backbone, encoder_revision=encoder_revision) | |
| count = sum(p.numel() for p in self.parameters()) | |
| assert count < 4_000_000, count | |
| def features(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 | |
| if self.backbone=='mobilevit': | |
| fs=self.encoder(x) | |
| else: | |
| fs=[] | |
| for i,layer in enumerate(self.encoder): | |
| x=layer(x) | |
| if i in [3,6,12,16]:fs.append(x) | |
| projected=[layer(f) for layer,f in zip(self.lateral,fs)] | |
| 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]) | |
| memory=self.memory_norm(torch.cat([F.adaptive_avg_pool2d(f,(8,8)).flatten(2).transpose(1,2) for f in projected[1:]],1)) | |
| 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) | |
| masks=torch.einsum('bqd,bdhw->bqhw',q,self.pixel(x))/(self.config['width']**.5) | |
| weights=F.softmax(masks.float(),dim=1) | |
| residual=2*torch.tanh(self.residual(x).float()) | |
| return q,weights,residual,memory.mean(1) | |
| def forward(self,L,return_details=False): | |
| q,weights,residual,context=self.features(L) | |
| if self.modes==1: | |
| palette=80*torch.tanh(self.palette(q))[:,None] | |
| scores=q.new_zeros((len(L),1)) | |
| else: | |
| mq=q[:,None]+self.mode_embeddings[None,:,None] | |
| palette=80*torch.tanh(self.palette(mq)+self.mode_offsets[None]) | |
| scores=self.mode_score(context) | |
| outputs=torch.einsum('bqhw,bkqc->bkchw',weights,palette.float())+residual[:,None] | |
| b,k,c,h,w=outputs.shape | |
| full=F.interpolate(outputs.reshape(b*k,c,h,w),size=(h*4,w*4),mode='bilinear',align_corners=False) | |
| full=full[...,:L.shape[-2],:L.shape[-1]].reshape(b,k,c,*L.shape[-2:]) | |
| default=full[torch.arange(b,device=L.device),scores.argmax(1)] | |
| if return_details: | |
| return {'output':default,'all':full,'scores':scores.float(), | |
| 'weights':weights,'palette':palette.float()} | |
| return default | |
| def initialize_from_v3(model,baseline): | |
| old=baseline.state_dict();state=model.state_dict() | |
| copied=[] | |
| for k,v in old.items(): | |
| if model.backbone=='mobilevit' and (k.startswith(('encoder.','lateral.')) or k in ['rgb_mean','rgb_std']):continue | |
| if k in state and v.shape==state[k].shape: | |
| state[k]=v;copied.append(k) | |
| model.load_state_dict(state,strict=True) | |
| return copied | |
| def save_decision(m,path): | |
| p=Path(path);p.mkdir(parents=True,exist_ok=True) | |
| save_file({k:v.detach().cpu().contiguous() for k,v in m.state_dict().items()},str(p/'model.safetensors')) | |
| (p/'config.json').write_text(json.dumps(m.config,indent=2)) | |
| def load_decision(path,device='cpu'): | |
| p=Path(path);cfg=json.loads((p/'config.json').read_text()) | |
| m=DecisionColorizer(**{k:cfg[k] for k in ['modes','backbone','width','queries','head','encoder_revision']}) | |
| m.load_state_dict(load_file(str(p/'model.safetensors')),strict=True) | |
| return m.to(device).eval() | |