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mZmZmZ G dd deZe	 G dd deZe	 G d	d
 d
eZG dd deZdS )    N)ARCH_REGISTRY   )BaseNetwork
LIPEncoderSPADEResnetBlockget_nonspade_norm_layerc                       s<   e Zd ZdZd fd	d
	Zdd Zdd ZdddZ  ZS )SPADEGeneratorzGenerator with SPADEResBlock   @   F      spectralspadesyncbatch3x3Tc	           	         s  t    || _|| _|| _|| _d| _|d| j  | _| j| _|rft	
|d| j | j | j | _nt	j|d| j ddd| _td| j d| j || _td| j d| j || _td| j d| j || _t	td| j d| j |td| j d| j |td| j d| j |td| j d| j |g| _t	t	jd| j ddddt	jd| j ddddt	jd| j ddddt	jd| j ddddg| _t	jdd	| _d S )
N         r	   r   )padding      )Zscale_factor)super__init__nfinput_ncis_traintrain_phasescale_ratioswshnnZLinearfcConv2dr   head_0
g_middle_0
g_middle_1Z
ModuleListupsto_rgbsZUpsampleup	self	num_in_chnum_featZuse_vaeZz_dimZ	crop_sizeZnorm_gr   Zinit_train_phase	__class__ e/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/archs/hifacegan_arch.pyr      s6    	
"zSPADEGenerator.__init__c                 C   sN   |  dd \}}|d| j  |d| j   }}tj|||fd}| |S )z
        Encode input_tensor into feature maps, can be overridden in derived classes
        Default: nearest downsampling of 2**5 = 32 times
        Nr   )size)r/   r   FZinterpolater   )r'   input_tensorhwr   r   xr,   r,   r-   encode=   s    zSPADEGenerator.encodec                 C   s   |}|  |}| ||}| |}| ||}| ||}| jrN| jd }n
t| j}t	|D ]}| |}| j
| ||}q`| j|d  t|d}t|}|S )Nr   皙?)r5   r    r%   r!   r"   r   r   lenr$   ranger#   r0   
leaky_relutorchtanh)r'   r4   segphaseir,   r,   r-   forwardG   s    




zSPADEGenerator.forwardNr   progressivec           	      C   sv  |du r|  |S | jr$| jd }n
t| j}|dkrhtt|d| d}|g| |gd| |   }nl|dkrtt|d| d d}|gd|  }|||< n4|dkr|d| kr|  |S |gd|  }|||< | |d }| ||d }| 	|}| 
||d	 }| ||d }t|D ](}| 	|}| j| ||d|  }q$| j|d  t|d
}t|}|S )al  
        A helper class for subspace visualization. Input and seg are different images.
        For the first n levels (including encoder) we use input, for the rest we use seg.

        If mode = 'progressive', the output's like: AAABBB
        If mode = 'one_plug', the output's like:    AAABAA
        If mode = 'one_ablate', the output's like:  BBBABB
        Nr   r@   r   r   Zone_plugZ
one_ablater	   r   r6   )r?   r   r   r7   r$   maxminr5   r    r%   r!   r"   r8   r#   r0   r9   r:   r;   )	r'   Zinput_xr<   nmoder=   Z
guide_listr4   r>   r,   r,   r-   mixed_guidance_forward`   s8    







z%SPADEGenerator.mixed_guidance_forward)r	   r
   Fr   r   r   Tr	   )Nr   r@   )	__name__
__module____qualname____doc__r   r5   r?   rE   __classcell__r,   r,   r*   r-   r   
   s           0
r   c                       s*   e Zd ZdZd fd	d
	Zdd Z  ZS )	HiFaceGANzk
    HiFaceGAN: SPADEGenerator with a learnable feature encoder
    Current encoder design: LIPEncoder
    r	   r
   Fr   r   r   Tc	           	   
      s6   t  |||||||| t||| j| j| j| _d S N)r   r   r   r   r   r   lip_encoderr&   r*   r,   r-   r      s    	zHiFaceGAN.__init__c                 C   s
   |  |S rL   )rM   )r'   r1   r,   r,   r-   r5      s    zHiFaceGAN.encode)r	   r
   Fr   r   r   Tr	   )rF   rG   rH   rI   r   r5   rJ   r,   r,   r*   r-   rK      s           rK   c                       s2   e Zd ZdZd fdd		Zd
d Zdd Z  ZS )HiFaceGANDiscriminatora)  
    Inspired by pix2pixHD multiscale discriminator.
    Args:
        num_in_ch (int): Channel number of inputs. Default: 3.
        num_out_ch (int): Channel number of outputs. Default: 3.
        conditional_d (bool): Whether use conditional discriminator.
            Default: True.
        num_d (int): Number of Multiscale discriminators. Default: 3.
        n_layers_d (int): Number of downsample layers in each D. Default: 4.
        num_feat (int): Channel number of base intermediate features.
            Default: 64.
        norm_d (str): String to determine normalization layers in D.
            Choices: [spectral][instance/batch/syncbatch]
            Default: 'spectralinstance'.
        keep_features (bool): Keep intermediate features for matching loss, etc.
            Default: True.
    r	   Tr   r   r
   spectralinstancec	                    sT   t    || _|}	|r |	|7 }	t|D ]&}
t|	||||}| d|
 | q(d S )NZdiscriminator_)r   r   num_dr8   NLayerDiscriminator
add_module)r'   r(   Z
num_out_chZconditional_drP   
n_layers_dr)   norm_dkeep_featuresr   r>   Zsubnet_dr*   r,   r-   r      s    	
zHiFaceGANDiscriminator.__init__c                 C   s   t j|ddddgddS )Nr	   r   r   F)kernel_sizestrider   Zcount_include_pad)r0   Z
avg_pool2d)r'   r4   r,   r,   r-   
downsample   s    z!HiFaceGANDiscriminator.downsamplec                 C   s6   g }|   D ]$\}}||}|| | |}q|S rL   )Znamed_childrenappendrX   )r'   r4   result_Z_net_doutr,   r,   r-   r?      s    
zHiFaceGANDiscriminator.forward)r	   r	   Tr   r   r
   rO   T)rF   rG   rH   rI   r   rX   r?   rJ   r,   r,   r*   r-   rN      s           rN   c                       s(   e Zd ZdZ fddZdd Z  ZS )rQ   z@Defines the PatchGAN discriminator with the specified arguments.c              
      s  t    d}tt|d d }|}|| _t|}	tj|||d|dt	ddgg}
t
d|D ]T}|}t|d d}||d krdnd}|
|	tj|||||dt	ddgg7 }
qb|
tj|d|d|dgg7 }
t
t|
D ]"}| d	t| tj|
|   qd S )
Nr   g      ?r   )rV   rW   r   r6   Fr   r   model)r   r   intnpceilrU   r   r   r   Z	LeakyReLUr8   rB   r7   rR   strZ
Sequential)r'   r   rS   r)   rT   rU   kwZpadwr   Z
norm_layersequencerC   Znf_prevrW   r*   r,   r-   r      s$    
"

zNLayerDiscriminator.__init__c                 C   sH   |g}|   D ]}||d }|| q| jr<|dd  S |d S d S )Nr   )childrenrY   rU   )r'   r4   resultsZsubmodelZintermediate_outputr,   r,   r-   r?      s    zNLayerDiscriminator.forward)rF   rG   rH   rI   r   r?   rJ   r,   r,   r*   r-   rQ      s   rQ   )numpyr_   r:   Ztorch.nnr   Ztorch.nn.functionalZ
functionalr0   Zbasicsr.utils.registryr   Zhifacegan_utilr   r   r   r   r   registerrK   rN   rQ   r,   r,   r,   r-   <module>   s    6