mini-unet-colorizer / spatial.py
User-2468's picture
Update validated colorizer weights, spatial decoder and deployment artifacts
704fa80 verified
Raw History Blame
1.6 kB
"""Controlled spatial decoding alternatives; no learned parameters."""
import math
import torch
import torch.nn.functional as F
def box_mean(x,radius):
# Border-normalized local windows, also work on images smaller than kernel.
return F.avg_pool2d(x,2*radius+1,stride=1,padding=radius,count_include_pad=False)
def guided_chroma(L,ab,radius=4,epsilon=.001):
"""Scalar luminance-guided local linear filter (He et al., ECCV 2010).
L is the network's [-1,1] luminance. Epsilon is in [0,1] luminance units.
Uses luminance only; reference colors never enter inference.
"""
if not isinstance(radius,int) or radius<0 or not math.isfinite(epsilon) or epsilon<=0:
raise ValueError('Invalid guided-filter radius or epsilon')
if radius==0:return ab
I=(L.float()+1)/2;p=ab.float()
mi=box_mean(I,radius);mp=box_mean(p,radius)
var=(box_mean(I*I,radius)-mi*mi).clamp_min(0)
cov=box_mean(I*p,radius)-mi*mp
a=cov/(var+epsilon);b=mp-a*mi
return box_mean(a,radius)*I+box_mean(b,radius)
def spatial_decode(model,logits,L,temperature=.38,pool=1,radius=0,epsilon=.001):
if pool<1 or not isinstance(pool,int):raise ValueError('pool must be a positive integer')
if pool>1:
# Average evidence before annealing, then upsample chroma, not RGB.
small=F.avg_pool2d(logits,pool,ceil_mode=True,count_include_pad=False)
ab=model.decode(small,temperature)
ab=F.interpolate(ab,size=L.shape[-2:],mode='bilinear',align_corners=False)
else:ab=model.decode(logits,temperature)
return guided_chroma(L,ab,radius,epsilon)