Download spatial.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 1.6 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/915de375b4f97d217f62fedcf467e8a4e52b3094/spatial.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer@915de375b4f97d217f62fedcf467e8a4e52b3094/spatial.py
-
curl -L -o spatial.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/915de375b4f97d217f62fedcf467e8a4e52b3094/spatial.py
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) | |