a
    þdË  ã                   @   sl   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 e
 ¡ G d	d
„ d
eƒƒZdS )é    N)ÚOrderedDict)Úbuild_network)Ú
build_loss)Úget_root_logger)ÚMODEL_REGISTRYé   )ÚSRModelc                   @   s0   e Zd ZdZdd„ Zdd„ Zdd„ Zdd	„ Zd
S )Ú
SRGANModelz.SRGAN model for single image super-resolution.c                 C   sþ  | j d }| dd¡| _| jdkr¢tƒ }| d| j› ¡ t| j d ƒ | j¡| _| j d  dd ¡}|d urŽ|  	| j|| j d  dd	¡d
¡ n
|  
d¡ | j ¡  t| j d ƒ| _|  | j¡| _|  | j¡ | j d  dd ¡}|d ur| j d  dd¡}|  	| j|| j d  dd	¡|¡ | j ¡  | j ¡  | d¡rRt|d ƒ | j¡| _nd | _| d¡r|t|d ƒ | j¡| _nd | _| d¡r¦t|d ƒ | j¡| _nd | _| d¡rÎt|d ƒ | j¡| _| dd¡| _| dd¡| _|  ¡  |  ¡  d S )NÚtrainÚ	ema_decayr   z+Use Exponential Moving Average with decay: Z	network_gÚpathZpretrain_network_gZstrict_load_gTÚ
params_emaZ	network_dZpretrain_network_dZparam_key_dÚparamsZstrict_load_dZ	pixel_optZldl_optZperceptual_optZgan_optÚnet_d_itersr   Únet_d_init_iters)ÚoptÚgetr   r   Úinfor   ÚtoZdeviceÚ	net_g_emaZload_networkÚ	model_emaÚevalÚnet_dZmodel_to_deviceZprint_networkÚnet_gr
   r   Úcri_pixZcri_ldlÚcri_perceptualÚcri_ganr   r   Úsetup_optimizersZsetup_schedulers)ÚselfÚ	train_optÚloggerZ	load_pathÚ	param_key© r"   úc/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/models/srgan_model.pyÚinit_training_settings   sF    

"


 

z!SRGANModel.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   J   s    
  zSRGANModel.setup_optimizersc                 C   sÔ  | j  ¡ D ]
}d|_q
| j ¡  |  | j¡| _d}tƒ }|| j	 dkrþ|| j
krþ| jrv|  | j| j¡}||7 }||d< | jrÀ|  | j| j¡\}}|d ur¨||7 }||d< |d urÀ||7 }||d< |   | j¡}| j|ddd}	||	7 }|	|d< | ¡  | j ¡  | j  ¡ D ]}d|_q| j ¡  |   | j¡}
| j|
ddd}||d	< t |
 ¡ ¡|d
< | ¡  |   | j ¡ ¡}| j|ddd}||d< t | ¡ ¡|d< | ¡  | j ¡  |  |¡| _| jdkrÐ| j| jd d S )NFr   Úl_g_pixÚ
l_g_percepÚ	l_g_styleT)Zis_discÚl_g_ganÚl_d_realZ
out_d_realÚl_d_fakeZ
out_d_fake)Zdecay)r   r'   Zrequires_gradr(   Z	zero_gradr   ZlqÚoutputr   r   r   r   Úgtr   r   ZbackwardÚstepr*   ÚtorchÚmeanÚdetachZreduce_loss_dictZlog_dictr   r   )r   Úcurrent_iterÚpZ	l_g_totalZ	loss_dictr+   r,   r-   Zfake_g_predr.   Zreal_d_predr/   Zfake_d_predr0   r"   r"   r#   Úoptimize_parametersU   sT    




zSRGANModel.optimize_parametersc                 C   sZ   t | dƒr*| j| j| jgd|ddgd n|  | jd|¡ |  | jd|¡ |  ||¡ d S )Nr   r   r   r   )r!   r   )ÚhasattrZsave_networkr   r   r   Zsave_training_state)r   Úepochr7   r"   r"   r#   Úsave   s
    
 zSRGANModel.saveN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r$   r   r9   r<   r"   r"   r"   r#   r	      s
   ;:r	   )r4   Úcollectionsr   Zbasicsr.archsr   Zbasicsr.lossesr   Zbasicsr.utilsr   Zbasicsr.utils.registryr   Zsr_modelr   Úregisterr	   r"   r"   r"   r#   Ú<module>   s   