a
    dv;                     @   sv   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m	Z	 d dl
mZ d dlmZ d dlmZ G dd	 d	ZdS )
    N)OrderedDict)deepcopy)DataParallelDistributedDataParallel)lr_scheduler)get_root_logger)master_onlyc                   @   s   e Zd ZdZdd Zdd Zdd Zdd	 Zd
d Zd:ddZ	dd Z
dd Zd;ddZdd Zdd Zdd Zdd Zdd Zed d! Zd"d# Zd$d% Zd<d'd(Zd)d* Zed=d,d-Zd>d/d0Zd?d1d2Zed3d4 Zd5d6 Zd7d8 Zd9S )@	BaseModelzBase model.c                 C   s<   || _ t|d dkrdnd| _|d | _g | _g | _d S )Nnum_gpur   cudacpuis_train)opttorchdevicer   
schedulers
optimizers)selfr    r   b/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/models/base_model.py__init__   s
    
zBaseModel.__init__c                 C   s   d S Nr   )r   datar   r   r   	feed_data   s    zBaseModel.feed_datac                 C   s   d S r   r   r   r   r   r   optimize_parameters   s    zBaseModel.optimize_parametersc                 C   s   d S r   r   r   r   r   r   get_current_visuals   s    zBaseModel.get_current_visualsc                 C   s   dS )z!Save networks and training state.Nr   )r   epochcurrent_iterr   r   r   save    s    zBaseModel.saveFc                 C   s0   | j d r| |||| n| |||| dS )a1  Validation function.

        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.
        distN)r   Zdist_validationZnondist_validation)r   Z
dataloaderr   Z	tb_loggerZsave_imgr   r   r   
validation$   s    	
zBaseModel.validationc                 C   s   t | dr|| jv rdS t | ds*t | _t }| jd d  D ]>\}}|dd}|dkrftdntd}t||d	d
||< qB|| j|< dS )zZInitialize the best metric results dict for recording the best metric value and iteration.best_metric_resultsNvalZmetricsbetterhigherz-infinf)r$   r#   iter)hasattrr"   dictr   itemsgetfloat)r   dataset_namerecordmetriccontentr$   Zinit_valr   r   r   _initialize_best_metric_results2   s    
z)BaseModel._initialize_best_metric_resultsc                 C   s   | j | | d dkrR|| j | | d kr|| j | | d< || j | | d< n:|| j | | d kr|| j | | d< || j | | d< d S )Nr$   r%   r#   r(   )r"   )r   r.   r0   r#   r   r   r   r   _update_best_metric_resultA   s    z$BaseModel._update_best_metric_result+?c                 C   s\   |  | j}t| }t| j }| D ](}|| j|j|| jd| d q.d S )N   )alpha)	get_bare_modelnet_gr*   Znamed_parametersZ	net_g_emakeysr   Zmul_Zadd_)r   Zdecayr8   Znet_g_paramsZnet_g_ema_paramskr   r   r   	model_emaK   s
    zBaseModel.model_emac                 C   s   | j S r   )log_dictr   r   r   r   get_current_logT   s    zBaseModel.get_current_logc                 C   sV   | | j}| jd r<| jdd}t|tj g|d}n| jd dkrRt|}|S )zModel to device. It also warps models with DistributedDataParallel
        or DataParallel.

        Args:
            net (nn.Module)
        r    find_unused_parametersF)Z
device_idsr>   r
   r5   )	tor   r   r,   r   r   r   Zcurrent_devicer   )r   netr>   r   r   r   model_to_deviceW   s    
zBaseModel.model_to_devicec                 K   s4   |dkr t jj||fi |}ntd| d|S )NAdamz
optimizer z is not supperted yet.)r   ZoptimrB   NotImplementedError)r   Z
optim_typeparamslrkwargs	optimizerr   r   r   get_optimizerg   s    zBaseModel.get_optimizerc                 C   s   | j d }|d d}|dv rL| jD ]"}| jtj|fi |d  q&nD|dkr| jD ]"}| jtj|fi |d  qZntd| ddS )	zSet up schedulers.Ztrain	schedulertype)ZMultiStepLRMultiStepRestartLRCosineAnnealingRestartLRz
Scheduler z is not implemented yet.N)	r   popr   r   appendr   rK   rL   rC   )r   Z	train_optZscheduler_typerG   r   r   r   setup_schedulersn   s    

"
"zBaseModel.setup_schedulersc                 C   s   t |ttfr|j}|S )zhGet bare model, especially under wrapping with
        DistributedDataParallel or DataParallel.
        )
isinstancer   r   module)r   r@   r   r   r   r7   {   s    zBaseModel.get_bare_modelc                 C   s   t |ttfr(|jj d|jjj }n
|jj }| |}t|}tt	dd |
 }t }|d| d|d || dS )zdPrint the str and parameter number of a network.

        Args:
            net (nn.Module)
        z - c                 S   s   |   S r   )Znumel)xr   r   r   <lambda>       z)BaseModel.print_network.<locals>.<lambda>z	Network: z, with parameters: z,dN)rP   r   r   	__class____name__rQ   r7   strsummap
parametersr   info)r   r@   Znet_cls_strZnet_strZ
net_paramsloggerr   r   r   print_network   s    

zBaseModel.print_networkc                 C   s8   t | j|D ]&\}}t |j|D ]\}}||d< q qdS )zSet learning rate for warm-up.

        Args:
            lr_groups_l (list): List for lr_groups, each for an optimizer.
        rE   N)zipr   param_groups)r   Zlr_groups_lrG   Z	lr_groupsparam_grouprE   r   r   r   _set_lr   s    zBaseModel._set_lrc                 C   s*   g }| j D ]}|dd |jD  q
|S )z;Get the initial lr, which is set by the scheduler.
        c                 S   s   g | ]}|d  qS )Z
initial_lrr   .0vr   r   r   
<listcomp>   rT   z*BaseModel._get_init_lr.<locals>.<listcomp>)r   rN   r_   )r   Zinit_lr_groups_lrG   r   r   r   _get_init_lr   s    
zBaseModel._get_init_lrr'   c                    sb    dkr| j D ]}|  q k r^|  }g }|D ]}| fdd|D  q4| | dS )u   Update learning rate.

        Args:
            current_iter (int): Current iteration.
            warmup_iter (int)： Warm-up iter numbers. -1 for no warm-up.
                Default： -1.
        r5   c                    s   g | ]}|   qS r   r   rb   r   warmup_iterr   r   re      rT   z2BaseModel.update_learning_rate.<locals>.<listcomp>N)r   steprf   rN   ra   )r   r   rh   rI   Zinit_lr_g_lZwarm_up_lr_lZ	init_lr_gr   rg   r   update_learning_rate   s    

zBaseModel.update_learning_ratec                 C   s   dd | j d jD S )Nc                 S   s   g | ]}|d  qS )rE   r   )rc   r`   r   r   r   re      rT   z7BaseModel.get_current_learning_rate.<locals>.<listcomp>r   )r   r_   r   r   r   r   get_current_learning_rate   s    z#BaseModel.get_current_learning_raterD   c              
   C   s  |dkrd}| d| d}t j| jd d |}t|trB|n|g}t|trV|n|g}t|t|kstJ di }t||D ]V\}}	| |}|	 }
|

 D ]*\}}|dr|d	d
 }| |
|< q|
||	< qd}|dkrrz|zt|| W nP tyJ } z6t }|d| d|d   td W Y d
}~nd
}~0 0 W |d8 }qrW |d8 }q|d8 }0 q|dkr|d| d d
S )a@  Save networks.

        Args:
            net (nn.Module | list[nn.Module]): Network(s) to be saved.
            net_label (str): Network label.
            current_iter (int): Current iter number.
            param_key (str | list[str]): The parameter key(s) to save network.
                Default: 'params'.
        r'   Zlatest_z.pthpathmodelsz4The lengths of net and param_key should be the same.module.   N   r   zSave model error: , remaining retry times: r5   Still cannot save . Just ignore it.)osrm   joinr   rP   listlenr^   r7   
state_dictr+   
startswithr   r   r   	Exceptionr   warningtimesleep)r   r@   Z	net_labelr   	param_keysave_filename	save_path	save_dictZnet_Z
param_key_ry   keyparamretryer\   r   r   r   save_network   s<    



 
zBaseModel.save_networkTc           
   
   C   s  |  |}| }t| }t| }t }||kr|d tt|| D ]}|d|  qR|d tt|| D ]}|d|  q|s||@ }|D ]V}	||	  ||	  kr|d|	 d||	 j	 d||	 j	  |
|	||	d < qdS )	a  Print keys with different name or different size when loading models.

        1. Print keys with different names.
        2. If strict=False, print the same key but with different tensor size.
            It also ignore these keys with different sizes (not load).

        Args:
            crt_net (torch model): Current network.
            load_net (dict): Loaded network.
            strict (bool): Whether strictly loaded. Default: True.
        zCurrent net - loaded net:z  zLoaded net - current net:zSize different, ignore [z]: crt_net: z; load_net: z.ignoreN)r7   ry   setr9   r   r|   sortedrw   sizeshaperM   )
r   Zcrt_netload_netstrictZcrt_net_keysZload_net_keysr\   rd   Zcommon_keysr:   r   r   r   _print_different_keys_loading   s,    


z'BaseModel._print_different_keys_loadingc           	   	   C   s   t  }| |}tj|dd d}|durP||vrHd|v rHd}|d || }|d|jj d| d	| d
 t| D ],\}}|	dr~|||dd < |
| q~| ||| |j||d dS )aY  Load network.

        Args:
            load_path (str): The path of networks to be loaded.
            net (nn.Module): Network.
            strict (bool): Whether strictly loaded.
            param_key (str): The parameter key of loaded network. If set to
                None, use the root 'path'.
                Default: 'params'.
        c                 S   s   | S r   r   )Zstoragelocr   r   r   rS   "  rT   z(BaseModel.load_network.<locals>.<lambda>)Zmap_locationNrD   z/Loading: params_ema does not exist, use params.zLoading z model from z, with param key: [z].ro   rp   )r   )r   r7   r   loadr[   rU   rV   r   r+   rz   rM   r   load_state_dict)	r   r@   Z	load_pathr   r   r\   r   r:   rd   r   r   r   load_network  s    

"
zBaseModel.load_networkc              
   C   s,  |dkr(||g g d}| j D ]}|d |  q| jD ]}|d |  q<| d}tj| jd d |}d}|d	krzzzt	|| W nN t
y }	 z6t }
|
d
|	 d|d   td W Y d}	~	nd}	~	0 0 W |d8 }qW |d8 }qz|d8 }0 qz|d	kr(|
d| d dS )zSave training states during training, which will be used for
        resuming.

        Args:
            epoch (int): Current epoch.
            current_iter (int): Current iteration.
        r'   )r   r(   r   r   r   r   z.staterm   Ztraining_statesrq   r   zSave training state error: rr   r5   Nrs   rt   )r   rN   ry   r   ru   rm   rv   r   r   r   r{   r   r|   r}   r~   )r   r   r   stateosr   r   r   r   r\   r   r   r   save_training_state1  s,    	




 
zBaseModel.save_training_statec                 C   s   |d }|d }t |t | jks*J dt |t | jksDJ dt|D ]\}}| j| | qLt|D ]\}}| j| | qndS )zReload the optimizers and schedulers for resumed training.

        Args:
            resume_state (dict): Resume state.
        r   r   zWrong lengths of optimizerszWrong lengths of schedulersN)rx   r   r   	enumerater   )r   Zresume_stateZresume_optimizersZresume_schedulersir   r   r   r   r   resume_trainingT  s    zBaseModel.resume_trainingc                 C   s   t   | jd rg }g }| D ]\}}|| || q$t |d}t jj|dd | jd dkrz|| jd  }dd t||D }t	 }| D ]\}}|
  ||< q|W  d   S 1 s0    Y  dS )	zreduce loss dict.

        In distributed training, it averages the losses among different GPUs .

        Args:
            loss_dict (OrderedDict): Loss dict.
        r    r   )dstZrankZ
world_sizec                 S   s   i | ]\}}||qS r   r   )rc   r   Zlossr   r   r   
<dictcomp>v  rT   z.BaseModel.reduce_loss_dict.<locals>.<dictcomp>N)r   Zno_gradr   r+   rN   stackZdistributedreducer^   r   meanitem)r   Z	loss_dictr9   Zlossesnamevaluer<   r   r   r   reduce_loss_dictc  s     


zBaseModel.reduce_loss_dictN)F)r4   )r'   )rD   )T)TrD   )rV   
__module____qualname____doc__r   r   r   r   r   r!   r2   r3   r;   r=   rA   rH   rO   r7   r   r]   ra   rf   rj   rk   r   r   r   r   r   r   r   r   r   r   r	      s:   


	


.
#

"r	   )ru   r}   r   collectionsr   copyr   Ztorch.nn.parallelr   r   Zbasicsr.modelsr   Zbasicsr.utilsr   Zbasicsr.utils.dist_utilr   r	   r   r   r   r   <module>   s   