mini-unet-colorizer / decision_model.py
User-2468's picture
Document final training decision and organize research history
1cca2dd verified
Raw History Blame
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()