a
    d
<                     @   s  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 d dlmZ d dl	m
Z
 d dlmZmZmZ d dlZd dlmZmZmZ zd dlmZ W n ey   d dlmZ Y n0 zd dlZd	ZW n ey   d
ZY n0 ejdkrd dlmZ nd dlmZ d dlmZ d dlm Z  zBd dl!m"Z"m#Z#m$Z$m%Z%m&Z&m'Z' d dl(m)Z) ee$dedZ$d	Z*W n eyn   dZ$d
Z*Y n0 e+e,Z-g dZ.dZ/dZ0dZ1dZ2d@ddZ3dAddZ4dBddZ5dCdd Z6e7d!d"d#Z8ee7ej9f d$d%d&Z:e7e7d'd(d)Z;e7d*d+d,Z<e/fe7e7d'd-d.Z=dDe7ee> d/d0d1Z?dEe7ee> ee@ed2 f d3d4d5ZAdFe7e7ee7 ee7 e@e@ee> ee> ee@ed2 f d7	d8d9ZBe>e7d:d;d<ZCe7ee7 d=d>d?ZDdS )G    N)partial)Path)TemporaryDirectory)IterableOptionalUnion)
HASH_REGEXdownload_url_to_fileurlparse)get_dir)_get_torch_homeTF)      )Literal)__version__)filter_pretrained_cfg)create_repoget_hf_file_metadatahf_hub_download
hf_hub_urlrepo_type_and_id_from_hf_idupload_folder)EntryNotFoundErrortimm)Zlibrary_nameZlibrary_version)get_cache_dirdownload_cached_file
has_hf_hubhf_splitload_model_config_from_hfload_state_dict_from_hfsave_for_hfpush_to_hf_hubzpytorch_model.binzmodel.safetensorszopen_clip_pytorch_model.binzopen_clip_model.safetensors c                 C   sP   t drtd t }| s"dn| f} t jj|dg| R  }t j|dd |S )zf
    Returns the location of the directory where models are cached (and creates it if necessary).
    ZTORCH_MODEL_ZOOz@TORCH_MODEL_ZOO is deprecated, please use env TORCH_HOME instead ZcheckpointsT)exist_ok)osgetenv_loggerwarningr   pathjoinmakedirs)Z	child_dirZhub_dirZ	model_dirr#   r#   Y/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/timm/models/_hub.pyr   9   s    

r   c                 C   s   t | ttfr| \} }nt| }tj|j}tjt |}tj	|st
d| | d }|rt|}|r||dnd }t| |||d |S )NzDownloading: "{}" to {}
   )progress)
isinstancelisttupler
   r%   r)   basenamer*   r   existsr'   infoformatr   searchgroupr	   )url
check_hashr.   filenamepartscached_filehash_prefixrr#   r#   r,   r   H   s    

r   c           	      C   s   t | ttfr| \} }nt| }tj|j}tjt |}tj	|r|rt
|}|rf|dnd }|rt|dF}t|  }|d t| |krW d    dS W d    n1 s0    Y  dS dS )Nr-   rbFT)r/   r0   r1   r
   r%   r)   r2   r*   r   r3   r   r6   r7   openhashlibsha256read	hexdigestlen)	r8   r9   r:   r;   r<   r>   r=   fZhdr#   r#   r,   check_cached_fileY   s     

.rG   c                 C   s   t s| rtdt S )Nz^Hugging Face hub model specified but package not installed. Run `pip install huggingface_hub`.)_has_hf_hubRuntimeError)Z	necessaryr#   r#   r,   r   m   s
    r   )hf_idc                 C   sT   |  d}dt|  k r"dks,n J d|d }t|dkrH|d nd }||fS )N@r      zChf_hub id should only contain one @ character to identify revision.r-   )splitrE   )rJ   Z	rev_splithf_model_idhf_revisionr#   r#   r,   r   u   s
    
"r   )	json_filec                 C   s@   t | ddd}| }W d    n1 s,0    Y  t|S )Nr>   zutf-8)encoding)r@   rC   jsonloads)rQ   readertextr#   r#   r,   load_cfg_from_json~   s    &rW   )model_idr:   c                 C   s   t | \}}t|||dS )N)revision)r   r   )rX   r:   rO   rP   r#   r#   r,   download_from_hf   s    rZ   )rX   c                 C   s   t dsJ t| d}t|}d|vrj|}i }|d|d< |dd |d< d|v rb|d|d< ||d< |d }| |d< d	|d
< d|v r|d |d< d|v r|d|d< d|v r|d|d< |d }||fS )NTconfig.jsonpretrained_cfgarchitecturenum_featureslabelslabel_namesZ	hf_hub_idzhf-hubsourcenum_classeslabel_descriptions)r   rZ   rW   pop)rX   r<   	hf_configr\   
model_namer#   r#   r,   r      s,    
r   c              
   C   s   t dsJ t| \}}trt|D ]Z}zBt|||d}td|  d| d| d tjj	|ddW   S  t
y|   Y q$0 q$t|||d	}td|  d
| d tj|ddS )NTrepo_idr:   rY   [z"] Safe alternative available for 'z' (as 'z&'). Loading weights using safetensors.cpu)Zdevice)r:   rY   z"] Safe alternative not found for 'z)'. Loading weights using default pytorch.)Zmap_location)r   r   _has_safetensors_get_safe_alternativesr   r'   r4   safetensorstorchZ	load_filer   debugload)rX   r:   rO   rP   Zsafe_filenameZcached_safe_filer<   r#   r#   r,   r      s"    r   )config_pathmodel_configc           	      C   s@  |pi }i }t | jddd}|d|d< |d| j|d< |d| j|d< |dt| dd }t|trx|rx||d< d|v rt	
d |d	|d |d	d }|rt|tttfsJ ||d	< |d
d }|rt|tsJ ||d
< ||d< || |d }tj||dd W d    n1 s20    Y  d S )NT)Zremove_sourceZremove_nullr]   rb   r^   Zglobal_poolr_   z'labels' as a config field for is deprecated. Please use 'label_names' and 'label_descriptions'. Renaming provided 'labels' field to 'label_names'.r`   rc   r\   wrL   )indent)r   r\   rd   getrb   r^   getattrr/   strr'   r(   
setdefaultdictr0   r1   updater@   rS   dump)	modelrq   rr   re   r\   Zglobal_pool_typer`   rc   rF   r#   r#   r,   save_config_for_hf   s4    
r}   both)save_directoryrr   safe_serializationc                 C   s   t dsJ t|}|jddd |  }|du s:|dkrXtsFJ dtj||t  |du sh|dkrxt	||t
  |d }t| ||d d S )NT)r$   parentsr~   z-`pip install safetensors` to use .safetensorsFr[   )rr   )r   r   mkdirZ
state_dictrk   rm   rn   Z	save_fileHF_SAFE_WEIGHTS_NAMEsaveHF_WEIGHTS_NAMEr}   )r|   r   rr   r   Ztensorsrq   r#   r#   r,   r       s    r    	Add model)	rh   commit_messagetokenrY   private	create_prrr   
model_cardr   c
                 C   s   t |||dd}
t|
\}}}| d| }ztt|d|d d}W n ty\   d}Y n0 t l}t| |||	d |s|pi }|dd }t|d }t	||}|
| t|||||d	W  d
   S 1 s0    Y  d
S )a5  
    Arguments:
        (...)
        safe_serialization (`bool` or `"both"`, *optional*, defaults to `False`):
            Whether to save the model using `safetensors` or the traditional PyTorch way (that uses `pickle`).
            Can be set to `"both"` in order to push both safe and unsafe weights.
    T)r   r   r$   /z	README.mdrg   F)rr   r   rM   )rh   Zfolder_pathrY   r   r   N)r   r   r   r   r   r   r    rN   r   generate_readme
write_textr   )r|   rh   r   r   rY   r   r   rr   r   r   repo_url_Z
repo_owner	repo_nameZ
has_readmetmpdirrf   Zreadme_pathreadme_textr#   r#   r,   r!     s.    


r!   )r   rf   c           
      C   s  d}|d7 }|d7 }|d|  dd d7 }d| v rd	| d v r|d
7 }t| d d	 ttfr| d d	 D ]}|d|  d7 }qnn|d| d d	   d7 }d| d v rt| d d ttfr| d d D ]}|d|  d7 }qn|d| d d   d7 }|d7 }|d| d7 }d| v rD|d| d  d7 }d| v r
|d7 }| d  D ]\}}t|ttfr|d| d7 }|D ]}|d| d7 }qn^t|tr|d| d7 }| D ] \}}|d| d| d7 }qn|d| d| d7 }qbd| v r0|d7 }|| d 7 }|d7 }d| v rV|d7 }|| d 7 }|d7 }d| v r|d7 }t| d ttfs| d g}n| d }|D ]}	|d|	 d7 }q|S )Nz---
z$tags:
- image-classification
- timm
zlibrary_name: timm
z	license: licensez
apache-2.0
detailsZDatasetz
datasets:
z- zPretrain Datasetz# Model card for descriptionz
## Model Details
z- **z:**
z  - z: z:** usagez
## Model Usage
Z
comparisonz
## Model Comparison
Zcitationz
## Citation
z
```bibtex
z
```
)ru   r/   r1   r0   loweritemsry   )
r   rf   r   dkvviZkiZ	citationscr#   r#   r,   r   :  s^    




r   )r:   returnc                 c   s:   | t krtV  | t tfvr6| dr6| dd d V  dS )aI  Returns potential safetensors alternatives for a given filename.

    Use case:
        When downloading a model from the Huggingface Hub, we first look if a .safetensors file exists and if yes, we use it.
        Main use case is filename "pytorch_model.bin" => check for "model.safetensors" or "pytorch_model.safetensors".
    z.binNz.safetensors)r   r   HF_OPEN_CLIP_WEIGHTS_NAMEendswith)r:   r#   r#   r,   rl   r  s    rl   )r"   )TF)T)F)N)NF)r   NNFFNNF)ErA   rS   loggingr%   sys	functoolsr   pathlibr   tempfiler   typingr   r   r   rn   Z	torch.hubr   r	   r
   r   ImportErrorr   Zsafetensors.torchrm   rk   version_infor   Ztyping_extensionsr   r   Ztimm.models._pretrainedr   Zhuggingface_hubr   r   r   r   r   r   Zhuggingface_hub.utilsr   rH   	getLogger__name__r'   __all__r   r   r   ZHF_OPEN_CLIP_SAFE_WEIGHTS_NAMEr   r   rG   r   rw   r   PathLikerW   rZ   r   r   ry   r}   boolr    r!   r   rl   r#   r#   r#   r,   <module>   s   

 





	" -          98