a
    þdõ  ã                   @   sX   d Z ddlZddlZddlZddlZddlZddlmZmZ e 	e
¡ZG dd„ dƒZdS )z¡ Checkpoint Saver

Track top-n training checkpoints and maintain recovery checkpoints on specified intervals.

Hacked together by / Copyright 2020 Ross Wightman
é    Né   )Úunwrap_modelÚget_state_dictc                
   @   sZ   e Zd Zdddddddddef
dd„Zdd	d
„Zddd„Zddd„Zddd„Zdd„ Z	dS )ÚCheckpointSaverNÚ
checkpointZrecoveryÚ Fé
   c                 C   s   || _ || _|| _|| _|| _g | _d | _d | _d| _d| _	|| _
|	| _|| _|| _d| _|
| _|
rjtjntj| _|| _|| _| jdksŒJ ‚d S )Nr   z.pth.tarr   )ÚmodelÚ	optimizerÚargsÚ	model_emaÚ
amp_scalerÚcheckpoint_filesÚ
best_epochÚbest_metricÚcurr_recovery_fileÚlast_recovery_fileÚcheckpoint_dirÚrecovery_dirÚsave_prefixÚrecovery_prefixÚ	extensionÚ
decreasingÚoperatorÚltÚgtÚcmpÚmax_historyÚ	unwrap_fn)Úselfr	   r
   r   r   r   Zcheckpoint_prefixr   r   r   r   r   r   © r    úd/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/utils/checkpoint_saver.pyÚ__init__   s(    zCheckpointSaver.__init__c                 C   sÒ  |dksJ ‚t j | jd| j ¡}t j | jd| j ¡}|  |||¡ t j |¡r\t  |¡ t  ||¡ | j	rx| j	d nd }t
| j	ƒ| jk s¦|d u s¦|  ||d ¡r¶t
| j	ƒ| jkrÀ|  d¡ d | jt|ƒg¡| j }t j | j|¡}t  ||¡ | j	 ||f¡ t| j	dd„ | j d	| _	d
}| j	D ]}	|d |	¡7 }q*t |¡ |d ur¶| jd u sp|  || j¡r¶|| _|| _t j | jd| j ¡}
t j |
¡rªt  |
¡ t  ||
¡ | jd u rÆdS | j| jfS )Nr   ÚtmpÚlastéÿÿÿÿr   ú-c                 S   s   | d S )Nr   r    )Úxr    r    r!   Ú<lambda>Q   ó    z1CheckpointSaver.save_checkpoint.<locals>.<lambda>)ÚkeyÚreversezCurrent checkpoints:
z {}
Z
model_best)NN)ÚosÚpathÚjoinr   r   Ú_saveÚexistsÚunlinkÚrenamer   Úlenr   r   Ú_cleanup_checkpointsr   ÚstrÚlinkÚappendÚsortedr   ÚformatÚ_loggerÚinfor   r   )r   ÚepochÚmetricZtmp_save_pathZlast_save_pathZ
worst_fileÚfilenameÚ	save_pathZcheckpoints_strÚcZbest_save_pathr    r    r!   Úsave_checkpoint?   sF    
ÿÿ

þ

&
zCheckpointSaver.save_checkpointc                 C   s¤   |t | jƒj ¡ t| j| jƒ| j ¡ ddœ}| jd urL| jj|d< | j|d< | j	d urh| j	 ¡ || j	j
< | jd ur„t| j| jƒ|d< |d ur”||d< t ||¡ d S )Né   )r<   ÚarchÚ
state_dictr
   ÚversionrC   r   Zstate_dict_emar=   )Útyper	   Ú__name__Úlowerr   r   r
   rD   r   r   Zstate_dict_keyr   ÚtorchÚsave)r   r?   r<   r=   Z
save_stater    r    r!   r/   c   s     û



zCheckpointSaver._saver   c                 C   s¸   t t| jƒ|ƒ}| j| }|dk s0t| jƒ|kr4d S | j|d … }|D ]\}z"t d |¡¡ t |d ¡ W qF t	y  } zt 
d |¡¡ W Y d }~qFd }~0 0 qF| jd |… | _d S )Nr   zCleaning checkpoint: {}z(Exception '{}' while deleting checkpoint)Úminr3   r   r   r:   Údebugr9   r,   ÚremoveÚ	ExceptionÚerror)r   ZtrimZdelete_indexZ	to_deleteÚdÚer    r    r!   r4   v   s    
(z$CheckpointSaver._cleanup_checkpointsc              
   C   sÄ   |dksJ ‚d  | jt|ƒt|ƒg¡| j }tj  | j|¡}|  ||¡ tj | j	¡r²z"t
 d | j	¡¡ t | j	¡ W n8 ty° } z t
 d || j	¡¡ W Y d }~n
d }~0 0 | j| _	|| _d S )Nr   r&   zCleaning recovery: {}z Exception '{}' while removing {})r.   r   r5   r   r,   r-   r   r/   r0   r   r:   rL   r9   rM   rN   rO   r   )r   r<   Z	batch_idxr>   r?   rQ   r    r    r!   Úsave_recovery„   s     *zCheckpointSaver.save_recoveryc                 C   sB   t j | j| j¡}t |d | j ¡}t|ƒ}t|ƒr>|d S dS )NÚ*r   r   )	r,   r-   r.   r   r   Úglobr   r8   r3   )r   Zrecovery_pathÚfilesr    r    r!   Úfind_recovery’   s    zCheckpointSaver.find_recovery)N)N)r   )r   )
rG   Ú
__module__Ú__qualname__r   r"   rA   r/   r4   rR   rV   r    r    r    r!   r      s    ó
)
$


r   )Ú__doc__rT   r   r,   ÚloggingrI   r	   r   r   Ú	getLoggerrG   r:   r   r    r    r    r!   Ú<module>   s   
