a
    þdö-  ã                   @   s”   d dl Z d dlmZ d dlmZ d dlmZ d dlmZ d dl	m
Z
 d dlmZ d dlmZmZ d d	lmZ d
dlmZ e ¡ G dd„ deƒƒZdS )é    N)ÚOrderedDict)Úpath)Útqdm)Úbuild_network)Ú
build_loss)Úcalculate_metric)ÚimwriteÚ
tensor2img)ÚMODEL_REGISTRYé   )ÚSRModelc                       s^   e Zd ZdZdd„ Zdd„ Zdd„ Zedd	„ ƒZd
d„ Z	d‡ fdd„	Z
dd„ Zdd„ Z‡  ZS )ÚHiFaceGANModelz§HiFaceGAN model for generic-purpose face restoration.
    No prior modeling required, works for any degradations.
    Currently doesn't support EMA for inference.
    c                 C   sD  | j d }| dd¡| _| jdkr*tdƒ‚| j ¡  t| j d ƒ| _|  | j¡| _|  	| j¡ | d¡r€t
|d ƒ | j¡| _nd | _| d¡r¨t
|d ƒ | j¡| _nd | _| d¡rÐt
|d ƒ | j¡| _nd | _| jd u rò| jd u ròtd	ƒ‚| d
¡rt
|d
 ƒ | j¡| _| dd¡| _| dd¡| _|  ¡  |  ¡  d S )NÚtrainÚ	ema_decayr   z(HiFaceGAN does not support EMA now. PassZ	network_dZ	pixel_optZperceptual_optZfeature_matching_optz*Both pixel and perceptual losses are None.Zgan_optÚnet_d_itersr   Únet_d_init_iters)ÚoptÚgetr   ÚNotImplementedErrorÚnet_gr   r   Únet_dZmodel_to_deviceZprint_networkr   ÚtoZdeviceÚcri_pixÚcri_perceptualÚcri_featÚ
ValueErrorÚcri_ganr   r   Úsetup_optimizersZsetup_schedulers)ÚselfÚ	train_opt© r    úg/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/models/hifacegan_model.pyÚinit_training_settings   s2    





z%HiFaceGANModel.init_training_settingsc                 C   s†   | j d }|d  d¡}| j|| j ¡ fi |d ¤Ž| _| j | j¡ |d  d¡}| j|| j ¡ fi |d ¤Ž| _	| j | j	¡ d S )Nr   Zoptim_gÚtypeZoptim_d)
r   ÚpopZget_optimizerr   Ú
parametersÚoptimizer_gZ
optimizersÚappendr   Úoptimizer_d)r   r   Z
optim_typer    r    r!   r   ?   s    
  zHiFaceGANModel.setup_optimizersc                 C   sÒ   |j dd… \}}|j dd… |j dd… kr|tjj |||f¡}tjj |||f¡}tj||gdd}tj||gdd}	n$tj||gdd}tj||gdd}	tj||	gdd}
|  |
¡}|  |¡\}}||fS )a+  
        This is a conditional (on the input) discriminator
        In Batch Normalization, the fake and real images are
        recommended to be in the same batch to avoid disparate
        statistics in fake and real images.
        So both fake and real images are fed to D all at once.
        éþÿÿÿNr   ©Zdimr   )ÚshapeÚtorchÚnnZ
functionalZinterpolateÚcatr   Ú_divide_pred)r   Zinput_lqÚoutputZground_truthÚhÚwÚlqÚrealZfake_concatZreal_concatZfake_and_realZdiscriminator_outÚ	pred_fakeÚ	pred_realr    r    r!   ÚdiscriminateJ   s    
zHiFaceGANModel.discriminatec                 C   s|   t | ƒtkrHg }g }| D ],}| dd„ |D ƒ¡ | dd„ |D ƒ¡ qn,| d|  d¡d … }| |  d¡d d… }||fS )zÁ
        Take the prediction of fake and real images from the combined batch.
        The prediction contains the intermediate outputs of multiscale GAN,
        so it's usually a list
        c                 S   s"   g | ]}|d |  d¡d … ‘qS )Nr   é   ©Úsize©Ú.0Ztensorr    r    r!   Ú
<listcomp>l   ó    z/HiFaceGANModel._divide_pred.<locals>.<listcomp>c                 S   s"   g | ]}||  d ¡d d… ‘qS )r   r8   Nr9   r;   r    r    r!   r=   m   r>   Nr   r8   )r#   Úlistr'   r:   )ÚpredZfaker4   Úpr    r    r!   r/   a   s    zHiFaceGANModel._divide_predc                 C   sâ  | j  ¡ D ]
}d|_q
| j ¡  |  | j¡| _d}tƒ }|| j	 dkr2|| j
kr2| jrz|  | j| j¡}||7 }||d< | jrÄ|  | j| j¡\}}|d ur¬||7 }||d< |d urÄ||7 }||d< |  | j| j| j¡\}}	| j|ddd}
||
7 }|
|d< | jr |  ||	¡}||7 }||d	< | ¡  | j ¡  | j  ¡ D ]}d|_q<| j ¡  |  | j| j ¡ | j¡\}}	| j|	ddd}||d
< | j|ddd}||d< || d }| ¡  | j ¡  |  |¡| _| jdkrÞtdƒ d S )NFr   Úl_g_pixÚ
l_g_percepÚ	l_g_styleT)Zis_discÚl_g_ganÚl_g_featÚl_d_realÚl_d_faker8   z(HiFaceGAN does not support EMA now. pass)r   r%   Zrequires_gradr&   Z	zero_gradr   r3   r0   r   r   r   r   Úgtr   r7   r   r   ZbackwardÚstepr(   ÚdetachZreduce_loss_dictZlog_dictr   Úprint)r   Úcurrent_iterrA   Z	l_g_totalZ	loss_dictrB   rC   rD   r5   r6   rE   rF   rG   rH   Z	l_d_totalr    r    r!   Úoptimize_parameterst   sV    




z"HiFaceGANModel.optimize_parametersFc                    sV   | j d d dv r| j ¡  | j d r8|  ||||¡ ntdƒ tƒ  ||||¡ dS )a²  
        Warning: HiFaceGAN requires train() mode even for validation
        For more info, see https://github.com/Lotayou/Face-Renovation/issues/31

        Args:
            dataloader (torch.utils.data.DataLoader): Validation dataloader.
            current_iter (int): Current iteration.
            tb_logger (tensorboard logger): Tensorboard logger.
            save_img (bool): Whether to save images. Default: False.
        Z	network_gr#   )Z	HiFaceGANZSPADEGeneratorÚdistzwIn HiFaceGANModel: The new metrics package is under development.Using super method now (Only PSNR & SSIM are supported)N)r   r   r   Zdist_validationrL   ÚsuperÚnondist_validation)r   Ú
dataloaderrM   Ú	tb_loggerÚsave_img©Ú	__class__r    r!   Ú
validationÂ   s    

zHiFaceGANModel.validationc              	   C   sþ  |j jd }| jd  d¡du}|r4tƒ | _g }g }tt|ƒdd}	|D ]@}
t t 	|
d d ¡¡d }|  
|
¡ |  ¡  |  ¡ }| |d	 ¡ d
|v r¬| |d
 ¡ | `| `| `tj ¡  |rp| jd rôt | jd d ||› d|› d¡}nj| jd d r4t | jd d ||› d| jd d › d¡}n*t | jd d ||› d| jd › d¡}tt|d	 ƒ|ƒ |	 d¡ |	 d|› ¡ qH|	 ¡  |rútj|dd}tj|dd}| jd d  ¡ D ]"\}}tt||d|ƒ| j|< qÈ|  |||¡ dS )a¥  
        TODO: Validation using updated metric system
        The metrics are now evaluated after all images have been tested
        This allows batch processing, and also allows evaluation of
        distributional metrics, such as:

        @ Frechet Inception Distance: FID
        @ Maximum Mean Discrepancy: MMD

        Warning:
            Need careful batch management for different inference settings.

        ÚnameÚvalZmetricsNÚimage)ÚtotalÚunitZlq_pathr   ÚresultrI   Zis_trainr   ZvisualizationÚ_z.pngÚsuffixr   zTest r*   )Úsr_packÚgt_pack)Zdatasetr   r   ÚdictZmetric_resultsr   ÚlenÚospÚsplitextÚbasenameZ	feed_dataÚtestZget_current_visualsr'   rI   r3   r0   r,   ZcudaZempty_cacheÚjoinr   r	   ÚupdateÚset_descriptionÚcloser.   Úitemsr   Z_log_validation_metric_values)r   rR   rM   rS   rT   Zdataset_nameZwith_metricsZ
sr_tensorsZ
gt_tensorsZpbarZval_dataZimg_nameZvisualsZsave_img_pathr`   ra   rX   Zopt_r    r    r!   rQ   Ø   sR    



ÿÿÿ
z!HiFaceGANModel.nondist_validationc                 C   sB   t | dƒrtdƒ |  | jd|¡ |  | jd|¡ |  ||¡ d S )NZ	net_g_emaz<HiFaceGAN does not support EMA now. Fallback to normal mode.r   r   )ÚhasattrrL   Zsave_networkr   r   Zsave_training_state)r   ÚepochrM   r    r    r!   Úsave  s
    
zHiFaceGANModel.save)F)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r"   r   r7   Ústaticmethodr/   rN   rW   rQ   ro   Ú__classcell__r    r    rU   r!   r      s   *
NBr   )r,   Úcollectionsr   Úosr   rd   r   Zbasicsr.archsr   Zbasicsr.lossesr   Zbasicsr.metricsr   Zbasicsr.utilsr   r	   Zbasicsr.utils.registryr
   Zsr_modelr   Úregisterr   r    r    r    r!   Ú<module>   s   