a
    d?                  	   @   s  d dl Z d dlZd dlmZ d dlmZmZmZmZm	Z	 d dl
Z
d dl
mZmZ d dlmZ d dlmZmZ ddd	d
ZesdgZnd dlmZ eeejjjdddZG dd de
jjZG dd de
jjZG dd de
jjZ d?eeedddZ!d@eee"df edddZ#dAee$ed d!d"Z%dBee"ed$d%d&Z&G d'd( d(ejZ'G d)d* d*ejZ(G d+d, d,ejZ)G d-d. d.e)Z*eeed/d0d1Z+eeejeeee	e"ef f d2d3d4Z,dCee	ee"f ed6 ed7d8d9Z-dDeeed< ed6 eed=d>dZ.dS )E    N)
namedtuple)List
NamedTupleOptionalTupleUnion)Tensornn)Literal)_TORCHVISION_AVAILABLE_TORCHVISION_GREATER_EQUAL_0_13ZSqueezeNet1_1_WeightsZAlexNet_WeightsZVGG16_Weights)squeezenet1_1alexnetvgg16)learned_perceptual_image_patch_similarity)models)net
pretrainedreturnc                 C   sR   t r<|r(tt| ttt|  jdj}qNtt| ddj}ntt| |dj}|S )zGet torchvision network.

    Args:
        net: Name of network
        pretrained: If pretrained weights should be used

    )weightsN)r   )r   getattrtv_weight_mapZIMAGENET1K_V1features)r   r   pretrained_features r   l/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torchmetrics/functional/image/lpips.py_get_net0   s     r   c                       s<   e Zd ZdZdeedd fddZeedd	d
Z  Z	S )
SqueezeNetzSqueezeNet implementation.FTNrequires_gradr   r   c           
   	      s   t    td|}d| _g }tdtddtddtddtddtddtdd	g}|D ]6}tj }|D ]}|t	|||  qn|
| q\t|| _|s|  D ]
}	d
|	_qd S )Nr               
            F)super__init__r   N_slicesrangetorchr	   
Sequential
add_modulestrappend
ModuleListslices
parametersr    )
selfr    r   r   r3   Zfeature_rangesZfeature_rangeseqiparam	__class__r   r   r*   E   s    

:
zSqueezeNet.__init__xr   c                 C   s8   t dg d}g }| jD ]}||}|| q|| S )Process input.squeeze_output)relu1relu2relu3relu4relu5Zrelu6Zrelu7)r   r3   r1   )r5   r<   r>   ZrelusZslice_r   r   r   forwardW   s    
zSqueezeNet.forward)FT
__name__
__module____qualname____doc__boolr*   r   r   rD   __classcell__r   r   r9   r   r   B   s   r   c                       s<   e Zd ZdZdeedd fddZeedd	d
Z  Z	S )AlexnetzAlexnet implementation.FTNr   c                    s2  t    td|}tj | _tj | _tj | _tj | _	tj | _
d| _tdD ]}| jt|||  q^tddD ]}| jt|||  qtddD ]}| jt|||  qtddD ]}| j	t|||  qtddD ]}| j
t|||  q|s.|  D ]}d|_q d S )Nr   r#   r"   r$   r%   r'   Fr)   r*   r   r-   r	   r.   slice1slice2slice3slice4slice5r+   r,   r/   r0   r4   r    )r5   r    r   Zalexnet_pretrained_featuresr<   r8   r9   r   r   r*   e   s*    

zAlexnet.__init__r;   c           	      C   sd   |  |}|}| |}|}| |}|}| |}|}| |}|}tdg d}||||||S )r=   alexnet_outputs)r?   r@   rA   rB   rC   rN   rO   rP   rQ   rR   r   )	r5   r<   hZh_relu1Zh_relu2Zh_relu3Zh_relu4Zh_relu5rS   r   r   r   rD   }   s    




zAlexnet.forward)FTrE   r   r   r9   r   rL   b   s   rL   c                       s<   e Zd ZdZdeedd fddZeedd	d
Z  Z	S )Vgg16zVgg16 implementation.FTNr   c                    s2  t    td|}tj | _tj | _tj | _tj | _	tj | _
d| _tdD ]}| jt|||  q^tddD ]}| jt|||  qtddD ]}| jt|||  qtddD ]}| j	t|||  qtddD ]}| j
t|||  q|s.|  D ]}d|_q d S )	Nr   r#      	            FrM   )r5   r    r   Zvgg_pretrained_featuresr<   r8   r9   r   r   r*      s*    

zVgg16.__init__r;   c           	      C   sd   |  |}|}| |}|}| |}|}| |}|}| |}|}tdg d}||||||S )r=   vgg_outputs)Zrelu1_2Zrelu2_2Zrelu3_3Zrelu4_3Zrelu5_3rT   )	r5   r<   rU   Z	h_relu1_2Z	h_relu2_2Z	h_relu3_3Z	h_relu4_3Z	h_relu5_3r\   r   r   r   rD      s    




zVgg16.forward)FTrE   r   r   r9   r   rV      s   rV   T)in_tenskeepdimr   c                 C   s   | j ddg|dS )z1Spatial averaging over heigh and width of images.r"      r^   )mean)r]   r^   r   r   r   spatial_average   s    rb   @   rd   .)r]   out_hwr   c                 C   s   t j|ddd| S )z+Upsample input with bilinear interpolation.bilinearF)sizemodealign_corners)r	   ZUpsample)r]   re   r   r   r   upsam   s    rj   绽|=)in_featepsr   c                 C   s&   t t j| d ddd}| ||  S )zNormalize tensors.r"      T)Zdimr^   )r-   sqrtsum)rl   rm   Znorm_factorr   r   r   normalize_tensor   s    rq   rd   )r<   rg   r   c                 C   sN   | j d |kr4| j d |kr4tjjj| ||fddS tjjj| ||fdddS )zlhttps://github.com/toshas/torch-fidelity/blob/master/torch_fidelity/sample_similarity_lpips.py#L127C22-L132.area)rh   rf   F)rh   ri   )shaper-   r	   Z
functionalZinterpolate)r<   rg   r   r   r   resize_tensor   s    rv   c                       s6   e Zd ZdZdd fddZeedddZ  ZS )	ScalingLayerzScaling layer.N)r   c                    sb   t    | jdtg dd d d d d f dd | jdtg dd d d d d f dd d S )Nshift)gQgI+gMbȿF)
persistentscale)gZd;O?gy&1?g?)r)   r*   Zregister_bufferr-   r   )r5   r9   r   r   r*      s    
*zScalingLayer.__init__)inpr   c                 C   s   || j  | j S r=   )rx   rz   )r5   r{   r   r   r   rD      s    zScalingLayer.forward)rF   rG   rH   rI   r*   r   rD   rK   r   r   r9   r   rw      s   rw   c                       s>   e Zd ZdZdeeedd fddZeedd	d
Z  Z	S )NetLinLayerz,A single linear layer which does a 1x1 conv.rn   FN)chn_inchn_outuse_dropoutr   c              	      sH   t    |rt gng }|tj||dddddg7 }tj| | _d S )Nrn   r   F)ZstridepaddingZbias)r)   r*   r	   ZDropoutZConv2dr.   model)r5   r~   r   r   Zlayersr9   r   r   r*      s    
zNetLinLayer.__init__r;   c                 C   s
   |  |S r|   )r   )r5   r<   r   r   r   rD      s    zNetLinLayer.forward)rn   F)
rF   rG   rH   rI   intrJ   r*   r   rD   rK   r   r   r9   r   r}      s   	r}   c                       sn   e Zd Zdeed eeeeee eee dd
 fddZde	e	eee
e	ee	ee	 f f d	d
dZ  ZS )_LPIPSTalexFNr   vggsqueeze)
r   r   spatial	pnet_rand	pnet_tuner   
model_path	eval_moderesizer   c
              	      s  t    || _|| _|| _|| _|	| _t | _| jdv rJt	}
g d| _
n2| jdkrdt}
g d| _
n| jdkr|t}
g d| _
t| j
| _|
| j | jd| _t| j
d |d	| _t| j
d
 |d	| _t| j
d |d	| _t| j
d |d	| _t| j
d |d	| _| j| j| j| j| jg| _| jdkrdt| j
d |d	| _t| j
d |d	| _|  j| j| jg7  _t| j| _|r|du rtjtjt | jdd| d}| j!t"j#|dddd |r| $  dS )a  Initializes a perceptual loss torch.nn.Module.

        Args:
            pretrained: This flag controls the linear layers should be pretrained version or random
            net: Indicate backbone to use, choose between ['alex','vgg','squeeze']
            spatial: If input should be spatial averaged
            pnet_rand: If backbone should be random or use imagenet pre-trained weights
            pnet_tune: If backprop should be enabled
            use_dropout: If dropout layers should be added
            model_path: Model path to load pretained models from
            eval_mode: If network should be in evaluation mode
            resize: If input should be resized to this size

        )r   r   )rd            r   r   )rd        r   r   r   )rd   r   r   r   r   r   r   )r   r    r   )r   rn   r"   r_   rW   r#      Nz..zlpips_models/z.pthcpu)Zmap_locationF)strict)%r)   r*   Z	pnet_typer   r   r   r   rw   scaling_layerrV   ZchnsrL   r   lenLr   r}   Zlin0Zlin1Zlin2Zlin3Zlin4linsZlin5Zlin6r	   r2   ospathabspathjoininspectgetfileZload_state_dictr-   loadeval)r5   r   r   r   r   r   r   r   r   r   net_typer9   r   r   r*      sJ    





z_LPIPS.__init__)in0in1retperlayer	normalizer   c              	   C   sR  |rd| d }d| d }|  ||  | }}| jd urXt|| jd}t|| jd}| j|| j| }}i i i   }	}
}t| jD ]>}t|| t||  |	|< |
|< |	| |
|  d ||< qg }t| jD ]\}| jr|	t
| j| || t|jdd  d q|	t| j| || dd qt|}|rN||fS |S )Nr"   rn   )rg   )re   Tr`   )r   r   rv   r   rD   r,   r   rq   r   r1   rj   r   tupleru   rb   rp   )r5   r   r   r   r   Z	in0_inputZ	in1_inputZouts0Zouts1Zfeats0Zfeats1Zdiffskkresvalr   r   r   rD   5  s*    
"0"z_LPIPS.forward)	Tr   FFFTNTN)FF)rF   rG   rH   rJ   r
   r   r0   r   r*   r   r   r   r   rD   rK   r   r   r9   r   r      s4            H r   c                       s(   e Zd ZdZed d fddZ  ZS )_NoTrainLpipsz8Wrapper to make sure LPIPS never leaves evaluation mode.)rh   r   c                    s   t  dS )z.Force network to always be in evaluation mode.F)r)   train)r5   rh   r9   r   r   r   [  s    z_NoTrainLpips.train)rF   rG   rH   rI   rJ   r   rK   r   r   r9   r   r   X  s   r   )imgr   r   c                 C   sD   |r|   dko&|  dkn
|  dk}| jdkoB| jd dkoB|S )z1Check that input is a valid image to the network.g      ?g        rr   rW   rn   r_   )maxminndimru   )r   r   Zvalue_checkr   r   r   
_valid_img`  s    (r   )img1img2r   r   r   c                 C   s   t | |rt ||shtd| j d|j d|  |  g d| | g d|rXddgnddg d|| ||d	 }|| jd fS )
NzeExpected both input arguments to be normalized tensors with shape [N, 3, H, W]. Got input with shape z and z and values in range z+ when all values are expected to be in the r   rn   rr   z range.)r   )r   
ValueErrorru   r   r   r   )r   r   r   r   lossr   r   r   _lpips_updatef  s     r   ra   )rp   ra   )
sum_scorestotal	reductionr   c                 C   s   |dkr| | S | S )Nra   r   )r   r   r   r   r   r   _lpips_computer  s    r   r   Fr   )r   r   r   r   r   r   c                 C   s,   t |d}t| |||\}}t| ||S )a  The Learned Perceptual Image Patch Similarity (`LPIPS_`) calculates the perceptual similarity between two images.

    LPIPS essentially computes the similarity between the activations of two image patches for some pre-defined network.
    This measure has been shown to match human perception well. A low LPIPS score means that image patches are
    perceptual similar.

    Both input image patches are expected to have shape ``(N, 3, H, W)``. The minimum size of `H, W` depends on the
    chosen backbone (see `net_type` arg).

    Args:
        img1: first set of images
        img2: second set of images
        net_type: str indicating backbone network type to use. Choose between `'alex'`, `'vgg'` or `'squeeze'`
        reduction: str indicating how to reduce over the batch dimension. Choose between `'sum'` or `'mean'`.
        normalize: by default this is ``False`` meaning that the input is expected to be in the [-1,1] range. If set
            to ``True`` will instead expect input to be in the ``[0,1]`` range.

    Example:
        >>> import torch
        >>> _ = torch.manual_seed(123)
        >>> from torchmetrics.functional.image.lpips import learned_perceptual_image_patch_similarity
        >>> img1 = (torch.rand(10, 3, 100, 100) * 2) - 1
        >>> img2 = (torch.rand(10, 3, 100, 100) * 2) - 1
        >>> learned_perceptual_image_patch_similarity(img1, img2, net_type='squeeze')
        tensor(0.1008, grad_fn=<DivBackward0>)

    )r   )r   r   r   rp   )r   r   r   r   r   r   r   r   r   r   r   r   v  s    "
)T)rc   )rk   )rd   )ra   )r   ra   F)/r   r   collectionsr   typingr   r   r   r   r   r-   r   r	   Ztyping_extensionsr
   Ztorchmetrics.utilities.importsr   r   r   Z__doctest_skip__Ztorchvisionr   r   r0   rJ   modules	containerr.   r   Moduler   rL   rV   rb   r   rj   floatrq   rv   rw   r}   r   r   r   r   r   r   r   r   r   r   <module>   sP    ++k("   