a
    d7                     @   s   d dl Z d dlZd dlZd dlmZ d dl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G dd dejZG dd dejZe G dd dejZdS )    N)default_init_weights)ARCH_REGISTRY)nn)
functionalc                   @   s   e Zd Zdd ZdS )NormStyleCodec                 C   s"   |t t j|d dddd  S )zNormalize the style codes.

        Args:
            x (Tensor): Style codes with shape (b, c).

        Returns:
            Tensor: Normalized tensor.
              T)Zdimkeepdim:0yE>)torchrsqrtmeanselfx r   j/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/gfpgan/archs/stylegan2_clean_arch.pyforward   s    	zNormStyleCode.forwardN)__name__
__module____qualname__r   r   r   r   r   r   
   s   r   c                       s2   e Zd ZdZd fdd	Zdd Zd	d
 Z  ZS )ModulatedConv2daX  Modulated Conv2d used in StyleGAN2.

    There is no bias in ModulatedConv2d.

    Args:
        in_channels (int): Channel number of the input.
        out_channels (int): Channel number of the output.
        kernel_size (int): Size of the convolving kernel.
        num_style_feat (int): Channel number of style features.
        demodulate (bool): Whether to demodulate in the conv layer. Default: True.
        sample_mode (str | None): Indicating 'upsample', 'downsample' or None. Default: None.
        eps (float): A value added to the denominator for numerical stability. Default: 1e-8.
    TNr
   c              	      s   t t|   || _|| _|| _|| _|| _|| _t	j
||dd| _t| jdddddd t	td||||t||d   | _|d | _d S )	NTbiasr   r   fan_inZlinearZscaleZ	bias_fillamodeZnonlinearityr   )superr   __init__in_channelsout_channelskernel_size
demodulatesample_modeepsr   Linear
modulationr   	Parameterr   randnmathsqrtweightpadding)r   r    r!   r"   num_style_featr#   r$   r%   	__class__r   r   r   '   s    zModulatedConv2d.__init__c           
      C   s  |j \}}}}| ||d|dd}| j| }| jrnt|dg d| j	 }|||| j
ddd }||| j
 || j| j}| jdkrtj|dddd}n| jdkrtj|d	ddd}|j \}}}}|d|| ||}tj||| j|d
}	|	j|| j
g|	j dd R  }	|	S )zForward function.

        Args:
            x (Tensor): Tensor with shape (b, c, h, w).
            style (Tensor): Tensor with shape (b, num_style_feat).

        Returns:
            Tensor: Modulated tensor after convolution.
        r   r   )r         upsamplebilinearFZscale_factorr   Zalign_cornersZ
downsampleg      ?)r-   groupsr2   )shaper'   viewr,   r#   r   r   powsumr%   r!   r"   r$   FinterpolateZconv2dr-   )
r   r   stylebchwr,   Zdemodoutr   r   r   r   A   s     

 

 zModulatedConv2d.forwardc                 C   s6   | j j d| j d| j d| j d| j d| j dS )Nz(in_channels=z, out_channels=z, kernel_size=z, demodulate=z, sample_mode=))r0   r   r    r!   r"   r#   r$   r   r   r   r   __repr__e   s    zModulatedConv2d.__repr__)TNr
   )r   r   r   __doc__r   r   rE   __classcell__r   r   r/   r   r      s      $r   c                       s,   e Zd ZdZd fdd	Zd	ddZ  ZS )
	StyleConva  Style conv used in StyleGAN2.

    Args:
        in_channels (int): Channel number of the input.
        out_channels (int): Channel number of the output.
        kernel_size (int): Size of the convolving kernel.
        num_style_feat (int): Channel number of style features.
        demodulate (bool): Whether demodulate in the conv layer. Default: True.
        sample_mode (str | None): Indicating 'upsample', 'downsample' or None. Default: None.
    TNc                    sb   t t|   t||||||d| _ttd| _	ttd|dd| _
tjddd| _d S )N)r#   r$   r   皙?TZnegative_slopeZinplace)r   rH   r   r   modulated_convr   r(   r   zerosr,   r   	LeakyReLUactivate)r   r    r!   r"   r.   r#   r$   r/   r   r   r   v   s    zStyleConv.__init__c           	      C   s`   |  ||d }|d u r:|j\}}}}||d|| }|| j|  }|| j }| |}|S )Ng;f?r   )rK   r7   Z	new_emptyZnormal_r,   r   rN   )	r   r   r=   noiserB   r>   _r@   rA   r   r   r   r   ~   s    

zStyleConv.forward)TN)Nr   r   r   rF   r   r   rG   r   r   r/   r   rH   j   s   rH   c                       s,   e Zd ZdZd fdd	Zd	ddZ  ZS )
ToRGBzTo RGB (image space) from features.

    Args:
        in_channels (int): Channel number of input.
        num_style_feat (int): Channel number of style features.
        upsample (bool): Whether to upsample. Default: True.
    Tc                    sF   t t|   || _t|dd|dd d| _tt	dddd| _
d S )Nr1   r   Fr"   r.   r#   r$   )r   rR   r   r3   r   rK   r   r(   r   rL   r   )r   r    r.   r3   r/   r   r   r      s    zToRGB.__init__Nc                 C   sB   |  ||}|| j }|dur>| jr6tj|dddd}|| }|S )a  Forward function.

        Args:
            x (Tensor): Feature tensor with shape (b, c, h, w).
            style (Tensor): Tensor with shape (b, num_style_feat).
            skip (Tensor): Base/skip tensor. Default: None.

        Returns:
            Tensor: RGB images.
        Nr   r4   Fr5   )rK   r   r3   r;   r<   )r   r   r=   skiprB   r   r   r   r      s    
zToRGB.forward)T)NrQ   r   r   r/   r   rR      s   rR   c                       s(   e Zd ZdZ fddZdd Z  ZS )ConstantInputzConstant input.

    Args:
        num_channel (int): Channel number of constant input.
        size (int): Spatial size of constant input.
    c                    s*   t t|   ttd|||| _d S Nr   )r   rU   r   r   r(   r   r)   r,   )r   Znum_channelsizer/   r   r   r      s    zConstantInput.__init__c                 C   s   | j |ddd}|S rV   )r,   repeat)r   batchrB   r   r   r   r      s    zConstantInput.forwardrQ   r   r   r/   r   rU      s   rU   c                       sD   e Zd ZdZd fdd	Zdd	 Zd
d Zdd ZdddZ  Z	S )StyleGAN2GeneratorCleana  Clean version of StyleGAN2 Generator.

    Args:
        out_size (int): The spatial size of outputs.
        num_style_feat (int): Channel number of style features. Default: 512.
        num_mlp (int): Layer number of MLP style layers. Default: 8.
        channel_multiplier (int): Channel multiplier for large networks of StyleGAN2. Default: 2.
        narrow (float): Narrow ratio for channels. Default: 1.0.
          r   r   c                    sN  t t|   || _t g}t|D ](}|tj||ddtj	dddg q$tj
| | _t| jdddddd	 td
| td
| td
| td
| td| | td| | td| | td| | td| | d	}|| _t|d dd| _t|d |d d|dd d| _t|d |dd| _tt|d| _| jd d d | _| jd d | _t | _t | _t | _|d }	t| jD ]<}
d|
d d  }dd||g}| jd|
 t j!|  qtd| jd D ]f}|d|   }| j"t|	|d|ddd | j"t||d|dd d | j"t||dd |}	qd S )NTr   rI   rJ   r   r   r   Z
leaky_relur   r[         @          )	48Z16Z32Z64Z128Z256Z512Z1024rb   r2   )rW   r1   rS   F)r3   r      rO   r3   )#r   rZ   r   r.   r   rangeextendr   r&   rM   Z
Sequential	style_mlpr   intchannelsrU   constant_inputrH   style_conv1rR   to_rgb1r*   loglog_size
num_layers
num_latentZ
ModuleListstyle_convsto_rgbsModulenoisesZregister_bufferr   r)   append)r   Zout_sizer.   Znum_mlpZchannel_multiplierZnarrowZstyle_mlp_layersiri   r    Z	layer_idx
resolutionr7   r!   r/   r   r   r      s    






z StyleGAN2GeneratorClean.__init__c                 C   sj   | j jj}tjdddd|dg}td| jd D ]4}tdD ]&}|tjddd| d| |d q<q0|S )zMake noise for noise injection.r   r2   devicer1   r   )rj   r,   ry   r   r)   re   rn   ru   )r   ry   rt   rv   rP   r   r   r   
make_noise  s    
&z"StyleGAN2GeneratorClean.make_noisec                 C   s
   |  |S )Nrg   r   r   r   r   
get_latent"  s    z"StyleGAN2GeneratorClean.get_latentc                 C   s0   t j|| j| jjjd}| |jddd}|S )Nrx   r   T)r	   )r   r)   r.   rj   r,   ry   rg   r   )r   rp   Z	latent_inlatentr   r   r   mean_latent%  s    z#StyleGAN2GeneratorClean.mean_latentFNTc	                    s@  |s fdd|D }|du rH|r0dg j  }n fddt j D }|dk rxg }	|D ]}
|	|||
|    qX|	}t|dkr j}|d jdk r|d dd|d}n|d }nnt|dkr*|du rt	d jd }|d dd|d}|d dd j| d}t
||gd} |jd } j||dddf |d d	} ||dddf }d}t jddd  jddd |ddd |ddd  jD ]l\}}}}}|||dd|f |d	}|||dd|d f |d	}|||dd|d f |}|d7 }q|}|r4||fS |dfS dS )
a  Forward function for StyleGAN2GeneratorClean.

        Args:
            styles (list[Tensor]): Sample codes of styles.
            input_is_latent (bool): Whether input is latent style. Default: False.
            noise (Tensor | None): Input noise or None. Default: None.
            randomize_noise (bool): Randomize noise, used when 'noise' is False. Default: True.
            truncation (float): The truncation ratio. Default: 1.
            truncation_latent (Tensor | None): The truncation latent tensor. Default: None.
            inject_index (int | None): The injection index for mixing noise. Default: None.
            return_latents (bool): Whether to return style latents. Default: False.
        c                    s   g | ]}  |qS r   r{   ).0srD   r   r   
<listcomp>A      z3StyleGAN2GeneratorClean.forward.<locals>.<listcomp>Nc                    s   g | ]}t  jd | qS rO   )getattrrt   )r   rv   rD   r   r   r   G  r   r   r   r1   r   r   )ro   re   ru   lenrp   ndimZ	unsqueezerX   randomrandintr   catrj   r7   rk   rl   ziprq   rr   )r   ZstylesZinput_is_latentrO   Zrandomize_noiseZ
truncationZtruncation_latentZinject_indexZreturn_latentsZstyle_truncationr=   r}   Zlatent1Zlatent2rB   rT   rv   Zconv1Zconv2Znoise1Znoise2Zto_rgbimager   rD   r   r   *  sJ    
 *zStyleGAN2GeneratorClean.forward)r[   r\   r   r   )FNTr   NNF)
r   r   r   rF   r   rz   r|   r~   r   rG   r   r   r/   r   rZ      s   
I       rZ   )r*   r   r   Zbasicsr.archs.arch_utilr   Zbasicsr.utils.registryr   r   Ztorch.nnr   r;   rs   r   r   rH   rR   rU   registerrZ   r   r   r   r   <module>   s   R#$