a
    d"                     @   s   d dl Z d dlZddlmZ ddlmZmZmZ ddlm	Z	 e	derd dl
m  mZ d dlmZmZ d d	lmZ d d
lmZ d dlmZ eeZdddZdddZdddZdddZdS )    N   )
get_logger   )FSDP_PYTORCH_VERSION
MODEL_NAMEOPTIMIZER_NAME)is_torch_versionz>=)DefaultLoadPlannerDefaultSavePlanner)!load_sharded_optimizer_state_dict)FullyShardedDataParallel)StateDictTypec           	      C   s  t j|dd t|| j| j| j| | }| jtjkr|dkrNt	 dnt	 d| d}t j
||}|jdkrtd|  t|| td|  n| jtjkr |dkrt	 d|j dnt	 d| d|j d}t j
||}td|  t|| td|  nt| jtjkrt j
|t	 d| }t j|dd td|  d	|i}tj|t|t d
 td|  W d    n1 s0    Y  d S )NTexist_okr   .bin_zSaving model to zModel saved to _rankmodel
state_dictZstorage_writerplanner)osmakedirsFSDPstate_dict_typestate_dict_configoptim_state_dict_configr   r   FULL_STATE_DICTr   pathjoinprocess_indexloggerinfotorchsaveLOCAL_STATE_DICTSHARDED_STATE_DICTdist_cpsave_state_dictFileSystemWriterr
   )	fsdp_pluginacceleratorr   
output_dirmodel_indexr   weights_nameZoutput_model_fileckpt_dir r0   d/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/accelerate/utils/fsdp_utils.pysave_fsdp_model"   s@    "
r2   c           	      C   s  |   t|| j| j| j | jtjkrt|tkrb|jdkrb| j	sRt
dW d    d S |dkrtt dnt d| d}tj||}td|  t|}td|  n| jtjkr8|dkrt d|j dnt d| d|j d}tj||}td|  t|}td|  n| jtjkrt |vrjtj|t d| n|}td|  d| i}tj|t|t d	 |d }td|  || W d    n1 s0    Y  d S )
Nr   zzSet the `sync_module_states` flag to `True` so that model states are synced across processes when initializing FSDP objectr   r   zLoading model from zModel loaded from r   r   )r   storage_readerr   )wait_for_everyoner   r   r   r   r   r   typer    Zsync_module_states
ValueErrorr   r   r   r   r!   r"   r#   loadr%   r&   r   r'   load_state_dictFileSystemReaderr	   )	r*   r+   r   	input_dirr-   r.   Zinput_model_filer   r/   r0   r0   r1   load_fsdp_modelG   sP    "

r;   c           
      C   s,  t j|dd t|| j| j| j t||}| jtjkr|j	dkr|dkrZt
 dnt
 d| d}t j||}td|  t|| td|  nbt j|t
 d| }	t j|	dd td|	  tjd|it|	t d	 td|	  W d    n1 s0    Y  d S )
NTr   r   r   r   zSaving Optimizer state to zOptimizer state saved in 	optimizerr   )r   r   r   r   r   r   Zoptim_state_dictr   r   r    r   r   r   r!   r"   r#   r$   r'   r(   r)   r
   )
r*   r+   r<   r   r,   optimizer_indexoptim_stateZoptim_state_nameZoutput_optimizer_filer/   r0   r0   r1   save_fsdp_optimizeru   s,    
 r?   c                 C   s.  |   t|| j| j| j | jtjkrd }|dkrBt dnt d| d}tj	
||}td|  t|}td|  nft |vrtj	
|t d| n|}	td|	  t| dt|	d}|d }td	|	  t|||}
||
 W d    n1 s 0    Y  d S )
Nr   r   r   zLoading Optimizer state from zOptimizer state loaded from zLoading Optimizer from r<   )Zmodel_state_dictZoptimizer_keyr3   zOptimizer loaded from )r4   r   r   r   r   r   r   r   r   r   r   r!   r"   r#   r7   r   r   r'   r9   Zoptim_state_dict_to_loadr8   )r*   r+   r<   r   r:   r=   r>   Zoptimizer_nameZinput_optimizer_filer/   Zflattened_osdr0   r0   r1   load_fsdp_optimizer   s4     
r@   )r   )r   )r   )r   )r   r#   loggingr   	constantsr   r   r   versionsr   Ztorch.distributed.checkpointZdistributed
checkpointr'   Z,torch.distributed.checkpoint.default_plannerr	   r
   Z&torch.distributed.checkpoint.optimizerr   Z2torch.distributed.fsdp.fully_sharded_data_parallelr   r   r   __name__r!   r2   r;   r?   r@   r0   r0   r0   r1   <module>   s   

%
.
