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