a
    
du'                     @   s  d dl mZ d dlZd dlmZ d dlm  mZ d dlmZ d dl	Z
ddlmZ d dlZd dlZdddZ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jZG dd dejZG dd dejZG dd deZG dd deZdd ZdS )     )absolute_importN)Variable   )pretrained_networksTc                 C   s   | j ddg|dS )N      keepdim)mean)in_tensr	    r   T/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/lpips/lpips.pyspatial_average   s    r   @   r   c                 C   s*   | j d | j d  }}tj|ddd| S )Nr   r   ZbilinearF)sizemodeZalign_corners)shapennZUpsample)r   out_HWZin_HZin_Wr   r   r   upsample   s    r   c                       s(   e Zd Zd
 fdd	Zddd	Z  ZS )LPIPSTalex0.1FNc              	      s4  t t|   |r6td|rdnd|||r,dndf  || _|| _|| _|| _|| _|| _	t
 | _| jdv r~tj}g d| _n6| jdkrtj}g d	| _n| jd
krtj}g d| _t| j| _|| j | jd| _|r"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rt| jd |d| _t| jd |d| _|  j| j| jg7  _t| j| _|r"|	du rddl}ddl }|j!"|j!#|$| jdd||f }	|r
td|	  | j%t&j'|	dddd |
r0| (  dS )a?   Initializes a perceptual loss torch.nn.Module

        Parameters (default listed first)
        ---------------------------------
        lpips : bool
            [True] use linear layers on top of base/trunk network
            [False] means no linear layers; each layer is averaged together
        pretrained : bool
            This flag controls the linear layers, which are only in effect when lpips=True above
            [True] means linear layers are calibrated with human perceptual judgments
            [False] means linear layers are randomly initialized
        pnet_rand : bool
            [False] means trunk loaded with ImageNet classification weights
            [True] means randomly initialized trunk
        net : str
            ['alex','vgg','squeeze'] are the base/trunk networks available
        version : str
            ['v0.1'] is the default and latest
            ['v0.0'] contained a normalization bug; corresponds to old arxiv v1 (https://arxiv.org/abs/1801.03924v1)
        model_path : 'str'
            [None] is default and loads the pretrained weights from paper https://arxiv.org/abs/1801.03924v1

        The following parameters should only be changed if training the network

        eval_mode : bool
            [True] is for test mode (default)
            [False] is for training mode
        pnet_tune
            [False] tune the base/trunk network
            [True] keep base/trunk frozen
        use_dropout : bool
            [True] to use dropout when training linear layers
            [False] for no dropout when training linear layers
        z@Setting up [%s] perceptual loss: trunk [%s], v[%s], spatial [%s]r   Zbaselineonoff)Zvggvgg16)r            r   r   )r        r   r   Zsqueeze)r   r   r   r!   r!   r   r   )
pretrainedZrequires_gradr   )use_dropoutr   r   r            Nz..zweights/v%s/%s.pthzLoading model from: %scpu)Zmap_locationF)strict))superr   __init__printZ	pnet_type	pnet_tune	pnet_randspatiallpipsversionScalingLayerscaling_layerpnr   ZchnsZalexnetZ
squeezenetlenLnetNetLinLayerZlin0Zlin1Zlin2Zlin3Zlin4linsZlin5Zlin6r   Z
ModuleListinspectospathabspathjoingetfileZload_state_dicttorchloadeval)selfr"   r6   r0   r/   r.   r-   r,   r#   Z
model_pathZ	eval_modeverboseZnet_typer9   r:   	__class__r   r   r*      sZ    %




(zLPIPS.__init__c                    sz  |rd d d| d }j dkr:|fn|f\}}j|j| }}i i i   }	}
 tjD ]B}t|| t||  |	|< |
|< |	| |
|  d  |< qzjrjr fddtjD }n fddtjD }n<jr* fddtjD }n fddtjD }d	}tjD ]}||| 7 }qP|rr||fS |S d S )
Nr   r   r   c                    s0   g | ](}t j|  | jd d dqS )r   Nr   )r   r8   r   .0kkdiffsin0rB   r   r   
<listcomp>       z!LPIPS.forward.<locals>.<listcomp>c                    s&   g | ]}t j|  | d dqS )Tr   )r   r8   rG   )rK   rB   r   r   rM      rN   c                    s0   g | ](}t  | jd ddjdd dqS )r   Tdimr	   r   NrF   )r   sumr   rG   )rK   rL   r   r   rM      rN   c                    s&   g | ]}t  | jd ddddqS )r   TrO   r   )r   rQ   rG   )rK   r   r   rM      rN   r   )	r0   r2   r6   forwardranger5   r/   Znormalize_tensorr.   )rB   rL   in1retPerLayer	normalizeZ	in0_inputZ	in1_inputZouts0Zouts1Zfeats0Zfeats1rI   resvallr   rJ   r   rR   p   s,    *&zLPIPS.forward)Tr   r   TFFFTNTT)FF__name__
__module____qualname__r*   rR   __classcell__r   r   rD   r   r      s     Yr   c                       s$   e Zd Z fddZdd Z  ZS )r1   c                    s^   t t|   | dtg dd d d d d f  | dtg dd d d d d f  d S )Nshift)gQgI+gMbȿscale)gZd;O?gy&1?g?)r)   r1   r*   Zregister_bufferr?   Tensor)rB   rD   r   r   r*      s    &zScalingLayer.__init__c                 C   s   || j  | j S N)r_   r`   )rB   inpr   r   r   rR      s    zScalingLayer.forwardrZ   r   r   rD   r   r1      s   r1   c                       s*   e Zd ZdZd fdd	Zdd Z  ZS )	r7   z- A single linear layer which does a 1x1 conv r   Fc              	      sL   t t|   |rt gng }|tj||dddddg7 }tj| | _d S )Nr   r   FZstridepaddingZbias)r)   r7   r*   r   ZDropoutConv2d
Sequentialmodel)rB   Zchn_inZchn_outr#   layersrD   r   r   r*      s    zNetLinLayer.__init__c                 C   s
   |  |S rb   )rh   )rB   xr   r   r   rR      s    zNetLinLayer.forward)r   Fr[   r\   r]   __doc__r*   rR   r^   r   r   rD   r   r7      s   r7   c                       s,   e Zd ZdZd	 fdd	Zd
ddZ  ZS )Dist2LogitLayerzc takes 2 distances, puts through fc layers, spits out value between [0,1] (if use_sigmoid is True)     Tc              	      s   t t|   tjd|dddddg}|tddg7 }|tj||dddddg7 }|tddg7 }|tj|ddddddg7 }|r|t g7 }tj| | _d S )Nr%   r   r   Trd   g?)	r)   rm   r*   r   rf   Z	LeakyReLUZSigmoidrg   rh   )rB   chn_midZuse_sigmoidri   rD   r   r   r*      s    zDist2LogitLayer.__init__皙?c              
   C   s4   | j tj|||| |||  |||  fddS )Nr   rP   )rh   rR   r?   cat)rB   d0d1Zepsr   r   r   rR      s    zDist2LogitLayer.forward)rn   T)rp   rk   r   r   rD   r   rm      s   rm   c                       s&   e Zd Zd fdd	Zdd Z  ZS )BCERankingLossrn   c                    s*   t t|   t|d| _tj | _d S )N)ro   )	r)   ru   r*   rm   r6   r?   r   ZBCELossloss)rB   ro   rD   r   r   r*      s    zBCERankingLoss.__init__c                 C   s*   |d d }| j ||| _| | j|S )N      ?g       @)r6   rR   Zlogitrv   )rB   rs   rt   ZjudgeZperr   r   r   rR      s    zBCERankingLoss.forward)rn   rZ   r   r   rD   r   ru      s   ru   c                       s   e Zd Zd fdd	Z  ZS )FakeNetTLabc                    s   t t|   || _|| _d S rb   )r)   rx   r*   use_gpu
colorspace)rB   rz   r{   rD   r   r   r*      s    zFakeNet.__init__)Try   )r[   r\   r]   r*   r^   r   r   rD   r   rx      s   rx   c                   @   s   e Zd ZdddZdS )L2Nc           
      C   s   |  d dksJ | jdkrz|  \}}}}tjtjtj|| d dd|d||dd|dd|dd|}|S | jdkrtjttj|j	dd	ttj|j	dd	d
d
d}tt|f}	| jr|	 }	|	S d S )Nr   r   RGBr   rq   r   ry   FZto_norm      Y@rS   float)r   r{   r?   r
   viewr/   l2	tensor2nptensor2tensorlabdataastyper   ra   rz   cuda)
rB   rL   rT   rU   NCXYvalueret_varr   r   r   rR      s    
H
z
L2.forward)Nr[   r\   r]   rR   r   r   r   r   r|      s   r|   c                   @   s   e Zd ZdddZdS )DSSIMNc                 C   s   |  d dksJ | jdkrNtjdt|j dt|j ddd}nD| jdkrtjttj|jd	d
ttj|jd	d
ddd}t	t
|f}| jr| }|S )Nr   r   r}   rw   g     o@r   r   ry   Fr~   r   )r   r{   r/   ZdssimZ	tensor2imr   r   r   r   r   r?   ra   rz   r   )rB   rL   rT   rU   r   r   r   r   r   rR      s    
0
zDSSIM.forward)Nr   r   r   r   r   r      s   r   c                 C   s8   d}|   D ]}|| 7 }qtd|  td|  d S )Nr   ZNetworkzTotal number of parameters: %d)
parametersZnumelr+   )r6   Z
num_paramsparamr   r   r   print_network   s
    
r   )T)r   )
__future__r   r?   Ztorch.nnr   Ztorch.nn.initinitZtorch.autogradr   numpynp r   r3   r/   r   r   Moduler   r1   r7   rm   ru   rx   r|   r   r   r   r   r   r   <module>   s&   

}
