a
    d%                     @   s   d dl Z d dlZd dlmZ d dlm  mZ d dlmZ d dlm	Z	 G dd dej
ZG dd dej
ZG dd	 d	ej
ZdddZG dd dej
ZG dd dej
ZG dd deZdddZdS )    N)init)spectral_normc                       s$   e Zd Z fddZdd Z  ZS )SPADEc           	         s  t    |dsJ td|}t|d}t|d}|dkrVt	|| _
nF|dkrttd t	|| _
n(|dkrtj|d	d
| _
nt| d|dkrdn|}|d }ttj||||dt | _tj||||d	d| _tj||||d	d| _d S )NZspadezspade(\D+)(\d)x\d      instanceZ	syncbatch\SyncBatchNorm is currently not supported under single-GPU mode, switch to "instance" insteadbatchFZaffinez2 is not a recognized param-free norm type in SPADE   kernel_sizepadding)r   r   bias)super__init__
startswithresearchstrgroupintnnInstanceNorm2dparam_free_normprintBatchNorm2d
ValueError
SequentialConv2dReLU
mlp_shared	mlp_gammamlp_beta)	selfZconfig_textZnorm_ncZlabel_ncparsedZparam_free_norm_typeksZnhiddenpw	__class__ e/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/archs/hifacegan_util.pyr      s$    
 zSPADE.__init__c                 C   sT   |  |}tj|| dd  dd}| |}| |}| |}|| | }|S )Nr   Znearest)sizemode)r   FZinterpolater,   r!   r"   r#   )r$   xZsegmap
normalizedZactvgammabetaoutr*   r*   r+   forward(   s    



zSPADE.forward)__name__
__module____qualname__r   r4   __classcell__r*   r*   r(   r+   r      s   r   c                       s:   e Zd ZdZd fdd	Zdd Zdd	 Zd
d Z  ZS )SPADEResnetBlocka  
    ResNet block that uses SPADE. It differs from the ResNet block of pix2pixHD in that
    it takes in the segmentation map as input, learns the skip connection if necessary,
    and applies normalization first and then convolution.
    This architecture seemed like a standard architecture for unconditional or
    class-conditional GAN architecture using residual block.
    The code was inspired from https://github.com/LMescheder/GAN_stability.
    spectralspadesyncbatch3x3   c                    s   t    ||k| _t||}tj||ddd| _tj||ddd| _| jr`tj||ddd| _d|v rt	| j| _t	| j| _| jrt	| j| _|
dd}t|||| _t|||| _| jrt|||| _d S )Nr;   r   r   F)r   r   spectral )r   r   learned_shortcutminr   r   conv_0conv_1conv_sr   replacer   norm_0norm_1norm_s)r$   ZfinZfoutZnorm_gZsemantic_ncZfmiddleZspade_config_strr(   r*   r+   r   C   s"    


zSPADEResnetBlock.__init__c                 C   sH   |  ||}| | | ||}| | | ||}|| }|S N)shortcutr@   actrD   rA   rE   )r$   r/   segx_sZdxr3   r*   r*   r+   r4   _   s
    zSPADEResnetBlock.forwardc                 C   s"   | j r| | ||}n|}|S rG   )r>   rB   rF   )r$   r/   rJ   rK   r*   r*   r+   rH   f   s    zSPADEResnetBlock.shortcutc                 C   s   t |dS )Ng?)r.   Z
leaky_relur$   r/   r*   r*   r+   rI   m   s    zSPADEResnetBlock.act)r:   r;   )	r5   r6   r7   __doc__r   r4   rH   rI   r8   r*   r*   r(   r+   r9   9   s
   	r9   c                   @   s"   e Zd ZdZd	ddZdd ZdS )
BaseNetworkz8 A basis for hifacegan archs with custom initialization normal{Gz?c                    s@    fdd}|  | |  D ]}t|dr |  q d S )Nc                    sp  | j j}|ddkrdt| dr<| jd ur<t| jjd  t| dr`| jd ur`t	| jjd nt| drl|ddks|ddkrld	krt| jjd  nd
krtj
| jj d n~dkrtj| jjdd nbdkrtj| jjddd nBdkr tj| jj d n$dkr4|   ntd dt| drl| jd urlt	| jjd d S )Nr   weightg      ?r           ZConvZLinearrO   Zxavier)gainZxavier_uniformZkaimingr   Zfan_in)ar-   Z
orthogonalnonezinitialization method [z] is not implemented)r)   r5   findhasattrrR   r   Znormal_datar   Z	constant_Zxavier_normal_Zxavier_uniform_Zkaiming_normal_Zorthogonal_Zreset_parametersNotImplementedError)m	classnamerT   	init_typer*   r+   	init_funcv   s,    *



z+BaseNetwork.init_weights.<locals>.init_funcinit_weights)applychildrenrX   r`   )r$   r^   rT   r_   r[   r*   r]   r+   r`   t   s
    

zBaseNetwork.init_weightsc                 C   s   d S rG   r*   rL   r*   r*   r+   r4      s    zBaseNetwork.forwardN)rO   rP   )r5   r6   r7   rM   r`   r4   r*   r*   r*   r+   rN   q   s   
"rN   r;   r   r   c                 C   s,   |  }t| | |||t|||| S rG   )expr.   Z
avg_pool2d)r/   logitkernelstrider   rR   r*   r*   r+   lip2d   s    rg   c                   @   s   e Zd ZdZdd ZdS )SoftGateg      (@c                 C   s   t || jS rG   )torchZsigmoidmulCOEFFrL   r*   r*   r+   r4      s    zSoftGate.forwardN)r5   r6   r7   rk   r4   r*   r*   r*   r+   rh      s   rh   c                       s,   e Zd Z fddZdd Zdd Z  ZS )SimplifiedLIPc              	      s>   t t|   ttj||ddddtj|ddt | _d S )Nr;   r   F)r   r   Tr
   )	r   rl   r   r   r   r   r   rh   rd   )r$   Zchannelsr(   r*   r+   r      s
    zSimplifiedLIP.__init__c                 C   s   | j d jjd d S )Nr   rS   )rd   rR   rY   Zfill_)r$   r*   r*   r+   
init_layer   s    zSimplifiedLIP.init_layerc                 C   s   t || |}|S rG   )rg   rd   )r$   r/   fracr*   r*   r+   r4      s    zSimplifiedLIP.forward)r5   r6   r7   r   rm   r4   r8   r*   r*   r(   r+   rl      s   rl   c                       s.   e Zd ZdZejf fdd	Zdd Z  ZS )
LIPEncoderz<Local Importance-based Pooling (Ziteng Gao et.al.,ICCV 2019)c              
      s   t    || _|| _d| _d}|d d }tj|||d|dd||t g}	d}
t|D ]l}t	|
d | j}|	t
||
 tj||
 || |d|d||| g7 }	|}
||d k r\|	tjdd	g7 }	q\tj|	 | _d S )
N   r;   r   r   F)rf   r   r   )rf   r   T)Zinplace)r   r   swshZ	max_ratior   r   r    ranger?   rl   r   model)r$   Zinput_ncZngfrq   rr   Zn_2xdown
norm_layerkwr'   rt   Z	cur_ratioiZ
next_ratior(   r*   r+   r      s,    


zLIPEncoder.__init__c                 C   s
   |  |S rG   )rt   rL   r*   r*   r+   r4      s    zLIPEncoder.forward)	r5   r6   r7   rM   r   r   r   r4   r8   r*   r*   r(   r+   ro      s   ro   r   c                    s"   dd   fdd}t d |S )Nc                 S   s    t | drt| dS | jdS )NZout_channelsr   )rX   getattrrR   r,   )layerr*   r*   r+   get_out_channel   s    

z0get_nonspade_norm_layer.<locals>.get_out_channelc                    s    dr"t| } tdd  }|dks6t|dkr:| S t| dd d ur`t| d | dd  |dkr|tj | dd}nP|dkrtd	 tj	 | d
d}n,|dkrtj	 | d
d}nt
d| dt| |S )Nr<   rV   r   r   r	   Tr
   Z
sync_batchr   Fr   znormalization layer z is not recognized)r   r   lenrx   delattrZregister_parameterr   r   r   r   r   r   )ry   Zsubnorm_typeru   rz   	norm_typer*   r+   add_norm_layer   s"    

z/get_nonspade_norm_layer.<locals>.add_norm_layerzKThis is a legacy from nvlabs/SPADE, and will be removed in future versions.)r   )r~   r   r*   r}   r+   get_nonspade_norm_layer   s    r   )r;   r   r   )r   )r   ri   Ztorch.nnr   Ztorch.nn.functionalZ
functionalr.   r   Ztorch.nn.utilsr   Moduler   r9   rN   rg   rh   rl   ro   r   r*   r*   r*   r+   <module>   s   -8)
#