a
    d|4                     @   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	m
Z
mZmZmZmZ G dd deZe G d	d
 d
ejZdS )    N)ARCH_REGISTRY)nn   )
ResUpBlock)	ConvLayerEqualConv2dEqualLinearResBlockScaledLeakyReLUStyleGAN2GeneratorBilinearc                       s,   e Zd ZdZd fdd		ZdddZ  ZS )StyleGAN2GeneratorBilinearSFTa8  StyleGAN2 Generator with SFT modulation (Spatial Feature Transform).

    It is the bilinear version. It does not use the complicated UpFirDnSmooth function that is not friendly for
    deployment. It can be easily converted to the clean version: StyleGAN2GeneratorCSFT.

    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.
        lr_mlp (float): Learning rate multiplier for mlp layers. Default: 0.01.
        narrow (float): The narrow ratio for channels. Default: 1.
        sft_half (bool): Whether to apply SFT on half of the input channels. Default: False.
             {Gz?r   Fc                    s&   t t| j||||||d || _d S )N)num_style_featnum_mlpchannel_multiplierlr_mlpnarrow)superr   __init__sft_half)selfout_sizer   r   r   r   r   r   	__class__ j/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/gfpgan/archs/gfpgan_bilinear_arch.pyr      s    
z&StyleGAN2GeneratorBilinearSFT.__init__NTc
                    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 ]\}}}}}|||dd|f |d	}|t|k rX jr@t
j|t|dd dd
\}}|||d   ||  }t
j||gdd
}n|||d   ||  }|||dd|d f |d	}|||dd|d f |}|d7 }q|}|	r||fS |dfS dS )a  Forward function for StyleGAN2GeneratorBilinearSFT.

        Args:
            styles (list[Tensor]): Sample codes of styles.
            conditions (list[Tensor]): SFT conditions to generators.
            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   )Z	style_mlp).0sr   r   r   
<listcomp>F       z9StyleGAN2GeneratorBilinearSFT.forward.<locals>.<listcomp>Nc                    s   g | ]}t  jd | qS noise)getattrZnoises)r   ir!   r   r   r"   L   r#   r   r      r   r$   )Zdim)Z
num_layersrangeappendlenZ
num_latentndimZ	unsqueezerepeatrandomrandinttorchcatZconstant_inputshapeZstyle_conv1Zto_rgb1zipZstyle_convsZto_rgbsr   splitintsize)r   Zstyles
conditionsinput_is_latentr%   randomize_noiseZ
truncationZtruncation_latentZinject_indexreturn_latentsZstyle_truncationstyleZlatentZlatent1Zlatent2outskipr'   Zconv1Zconv2Znoise1Znoise2Zto_rgbZout_sameZout_sftimager   r!   r   forward-   sV    
 *"z%StyleGAN2GeneratorBilinearSFT.forward)r   r   r   r   r   F)FNTr   NNF__name__
__module____qualname____doc__r   r?   __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 )GFPGANBilineara  The GFPGAN architecture: Unet + StyleGAN2 decoder with SFT.

    It is the bilinear version and it does not use the complicated UpFirDnSmooth function that is not friendly for
    deployment. It can be easily converted to the clean version: GFPGANv1Clean.


    Ref: GFP-GAN: Towards Real-World Blind Face Restoration with Generative Facial Prior.

    Args:
        out_size (int): The spatial size of outputs.
        num_style_feat (int): Channel number of style features. Default: 512.
        channel_multiplier (int): Channel multiplier for large networks of StyleGAN2. Default: 2.
        decoder_load_path (str): The path to the pre-trained decoder model (usually, the StyleGAN2). Default: None.
        fix_decoder (bool): Whether to fix the decoder. Default: True.

        num_mlp (int): Layer number of MLP style layers. Default: 8.
        lr_mlp (float): Learning rate multiplier for mlp layers. Default: 0.01.
        input_is_latent (bool): Whether input is latent style. Default: False.
        different_w (bool): Whether to use different latent w for different layers. Default: False.
        narrow (float): The narrow ratio for channels. Default: 1.
        sft_half (bool): Whether to apply SFT on half of the input channels. Default: False.
    r   r   NTr   r   Fc                    sR  t t|   || _|	| _|| _|
d }td| td| td| td| td| | td| | td| | td| | td| | d	}tt|d	| _	d	tt|d	 }t
d
||  dddd| _||  }t | _t| j	d	dD ],}|d	|d    }| jt|| |}qt
||d d
ddd| _|d }t | _td
| j	d D ]*}|d	|   }| jt|| |}qft | _td
| j	d D ].}| jt|d	|   d
dddddd q|	r tt|d	d	 d	 | }n|}t|d d d |dddd d| _t||||||
|d| _|rb| jtj|dd dd  |r| j D ]\}}d|_qrt | _ t | _!td
| j	d D ]}|d	|   }|r|}n|d	 }| j t"t||d
dddddt#dt||d
ddddd | j!t"t||d
dddddt#dt||d
ddddd qd S )Ng      ?r         @          )	48Z16Z32Z64Z128Z256Z512Z1024r   r(   r   T)biasactivaterL   r   )ZstridepaddingrN   bias_init_val   )rN   rR   Zlr_mulZ
activation)r   r   r   r   r   r   r   c                 S   s   | S )Nr   )Zstoragelocr   r   r   <lambda>   r#   z)GFPGANBilinear.__init__.<locals>.<lambda>)Zmap_locationZ
params_emaFg?)$r   rF   r   r8   different_wr   r5   mathloglog_sizer   conv_body_firstr   Z
ModuleListconv_body_downr)   r*   r	   
final_convconv_body_upr   toRGBr   r   final_linearr   stylegan_decoderZload_state_dictr0   loadZnamed_parametersZrequires_gradcondition_scalecondition_shiftZ
Sequentialr
   )r   r   r   r   Zdecoder_load_pathZfix_decoderr   r   r8   rV   r   r   Zunet_narrowZchannelsZfirst_out_sizeZin_channelsr'   Zout_channelsZlinear_out_channel_paramZsft_out_channelsr   r   r   r      s    







,



zGFPGANBilinear.__init__c                 C   s$  g }g }g }|  |}t| jd D ]}	| j|	 |}|d| q$| |}| ||dd}
| j	r|
|
dd| j
}
t| jd D ]n}	|||	  }| j|	 |}| j|	 |}||  | j|	 |}||  |r|| j|	 | q| j|
g||| j|d\}}||fS )al  Forward function for GFPGANBilinear.

        Args:
            x (Tensor): Input images.
            return_latents (bool): Whether to return style latents. Default: False.
            return_rgb (bool): Whether return intermediate rgb images. Default: True.
            randomize_noise (bool): Randomize noise, used when 'noise' is False. Default: True.
        r   r   rP   )r:   r8   r9   )rZ   r)   rY   r[   insertr\   r_   viewr6   rV   r   r]   rb   r*   clonerc   r^   r`   r8   )r   xr:   Z
return_rgbr9   r7   Z
unet_skipsZout_rgbsZfeatr'   Z
style_codeZscaleshiftr>   rd   r   r   r   r?     s6    	


zGFPGANBilinear.forward)
r   r   NTr   r   FFr   F)FTTr@   r   r   r   r   rF      s             lrF   )rW   r.   r0   Zbasicsr.utils.registryr   r   Zgfpganv1_archr   Zstylegan2_bilinear_archr   r   r   r	   r
   r   r   registerModulerF   r   r   r   r   <module>   s    w