a
    d4                  	   @   s8  d Z ddlZddlZddlmZ ddlmZmZmZm	Z	m
Z
 ddlZzddlZdZW n eyj   dZY n0 eeZg dZeeef eeef dd	d
Zdeee
eejf eeef dddZdejjeee
eejf eee	e dddZdeeef ejjedddZdejjeejjeedddZdS )zi Model creation / weight loading / state_dict helpers

Hacked together by / Copyright 2020 Ross Wightman
    N)OrderedDict)AnyCallableDictOptionalUnionTF)clean_state_dictload_state_dictload_checkpointremap_state_dictresume_checkpoint)
state_dictreturnc                 C   s<   i }|   D ]*\}}|dr*|dd  n|}|||< q|S )Nzmodule.   )items
startswith)r   Zcleaned_state_dictkvname r   ]/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/_helpers.pyr      s
    
r   cpu)checkpoint_pathuse_emadevicer   c                 C   s   | rt j| rt| dr<ts*J dtjj| |d}ntj	| |d}d}t
|tr|rr|dd d urrd}n4|r|dd d urd}nd|v rd}nd	|v rd	}t|r|| n|}td
||  |S td|  t d S )Nz.safetensorsz-`pip install safetensors` to use .safetensorsr   Zmap_location Zstate_dict_emaZ	model_emar   modelzLoaded {} from checkpoint '{}'No checkpoint found at '{}')ospathisfilestrendswith_has_safetensorssafetensorstorchZ	load_fileload
isinstancedictgetr   _loggerinfoformaterrorFileNotFoundError)r   r   r   
checkpointstate_dict_keyr   r   r   r   r	      s(    
r	   )r   r   r   r   strictremap	filter_fnc           	      C   sx   t j|d  dv r:t| dr.| | ntdd S t|||d}|rXt|| }n|rf||| }| j||d}|S )N)z.npzz.npyload_pretrainedz"Model cannot load numpy checkpointr   )r3   )	r    r!   splitextlowerhasattrr7   NotImplementedErrorr	   r   )	r   r   r   r   r3   r4   r5   r   Zincompatible_keysr   r   r   r
   >   s    	

r
   )r   r   allow_reshapec                 C   s   i }t |  |  D ]\\}}\}}| | ks`J d| d|j d| d|j d	|j|jkr|r~||j}n*dsJ d| d|j d| d|j d	|||< q|S )z remap checkpoint by iterating over state dicts in order (ignoring original keys).
    This assumes models (and originating state dict) were created with params registered in same order.
    zTensor size mismatch z: z vs z. Remap failed.FzTensor shape mismatch )zipr   r   ZnumelshapeZreshape)r   r   r<   Zout_dictkavakbZvbr   r   r   r   X   s    &6*
r   )r   r   	optimizerloss_scalerlog_infoc                 C   s8  d }t j|rtj|dd}t|trd|v r|r@td t	|d }| 
| |d urd|v r|rttd |
|d  |d ur|j|v r|rtd |
||j  d|v r|d }d	|v r|d	 d
kr|d
7 }|rtd||d  n | 
| |rtd| |S td| t d S )Nr   r   r   z(Restoring model state from checkpoint...rB   z,Restoring optimizer state from checkpoint...z2Restoring AMP loss scaler state from checkpoint...epochversion   z!Loaded checkpoint '{}' (epoch {})zLoaded checkpoint '{}'r   )r    r!   r"   r'   r(   r)   r*   r,   r-   r   r	   r2   r.   r/   r0   )r   r   rB   rC   rD   Zresume_epochr1   r   r   r   r   r   l   s8    




r   )Tr   )Tr   TFN)T)NNT) __doc__loggingr    collectionsr   typingr   r   r   r   r   r'   Zsafetensors.torchr&   r%   ImportError	getLogger__name__r,   __all__r#   r   boolr   r	   nnModuler
   r   ZoptimZ	Optimizerr   r   r   r   r   <module>   sf   

   
"      
   