File size: 5,570 Bytes
915de37
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
"""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()