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