a
    d?                     @   s  d dl mZmZ d dlZd dlZd dlZd dlmZ d dlZd dlZd dl	Z	d dl
Z
d dlmZmZmZm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 dd
lmZmZmZ ddlmZmZm Z m!Z!m"Z"m#Z#m$Z$m%Z% d dl&m'Z' g dZ(eG dd dZ)eG dd dZ*dZ+ejejdddZ,e$edddZ-G dd deZ.G dd de.Z/G dd de.Z0e$e1ddd Z2ee$ eee$  d!d"d#Z3d$d% Z4ej5ej5e"e1e6d&d'd(Z7G d)d* d*eZ8G d+d, d,e	j9Z:G d-d. d.eZ;dS )/    )ABCabstractmethodN)	dataclass)ListUnionDictcast)Tensor)Future)Path   )MetadataMetadataIndex)StorageReaderStorageWriterWriteResult)LoadItemTypeLoadPlannerLoadPlanSavePlanSavePlannerReadItem	WriteItemWriteItemType)narrow_tensor_by_index)FileSystemWriterSlicedBufferedReaderFileSystemReaderc                   @   s*   e Zd ZU dZeed< eed< eed< dS )_StorageInfoz,
    This is the per entry storage info
    relative_pathoffsetlengthN)__name__
__module____qualname____doc__str__annotations__int r)   r)   p/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/distributed/checkpoint/filesystem.pyr   0   s   
r   c                   @   s   e Zd ZU eed< dS )_StoragePrefixprefixN)r"   r#   r$   r&   r'   r)   r)   r)   r*   r+   ;   s   
r+   z.distcp)tensorreturnc                 C   s,   |    } |   |  kr(|  } | S N)detachcpuZ_typed_storage_sizenumelclone)r-   r)   r)   r*   _trimC   s    r5   )itemr.   c                 C   s   t | j||dS )N)indexsize_in_bytesstorage_data)r   r7   )r6   r8   r9   r)   r)   r*   _result_from_write_itemJ   s    r:   c                   @   s,   e Zd Zedd Zdd Zedd ZdS )_TensorLoaderc                 C   s   d S r/   r)   selfsizeobjr)   r)   r*   addS   s    z_TensorLoader.addc                 C   s   d S r/   r)   r=   r)   r)   r*   start_loadingW   s    z_TensorLoader.start_loadingc                 C   s   d S r/   r)   rA   r)   r)   r*   valuesZ   s    z_TensorLoader.valuesN)r"   r#   r$   r   r@   rB   rC   r)   r)   r)   r*   r;   R   s
   
r;   c                   @   s,   e Zd Zdd Zdd Zdd Zdd Zd	S )
_SerialCpuLoaderc                 C   s   || _ g | _d S r/   )resolve_funitems)r=   rE   r)   r)   r*   __init__`   s    z_SerialCpuLoader.__init__c                 C   s   | j ||f d S r/   )rF   appendr<   r)   r)   r*   r@   d   s    z_SerialCpuLoader.addc                 C   s   d S r/   r)   rA   r)   r)   r*   rB   g   s    z_SerialCpuLoader.start_loadingc                 c   sP   | j D ]D\}}| | }| }|  | kr@| }||fV  qd S r/   )rF   rE   r0   r1   storager>   r3   r4   r=   _r?   r-   r)   r)   r*   rC   j   s    z_SerialCpuLoader.valuesN)r"   r#   r$   rG   r@   rB   rC   r)   r)   r)   r*   rD   _   s   rD   c                   @   sR   e Zd ZdddZedd Zdd Zd	d
 Zdd Zdd Z	dd Z
dd ZdS )_OverlappingCpuLoaderN@B c                 C   sd   || _ g | _|| _d| _t | _d| _d| _|p:t	j
 | _| jt	j
 kr`| jt	j
  d S )Nr   F)rE   rF   inflight_threshholdin_flight_datacollectionsdequecurrent_itemsidxstartedtorchcudaZcurrent_streamstreamZwait_stream)r=   rE   rW   rN   r)   r)   r*   rG   w   s    
z_OverlappingCpuLoader.__init__c                 C   s   | j t| jkS r/   )rS   lenrF   rA   r)   r)   r*   _done   s    z_OverlappingCpuLoader._donec                 C   sb   g }| j | jkr| j  | j | jkr^| j }|  j |d  |d   8  _ || q|S Nr   )	rO   rN   rW   synchronizerR   popleftr3   element_sizerH   )r=   drainedvalr)   r)   r*   _drain   s    

"z_OverlappingCpuLoader._drainc                 C   s   t j| j | js| j| jk r| j| j \}}|  jd7  _| |	 }|j
rd|jddd}n,|jt dkr|  | kr| }| j||f |  j| |  7  _qW d    n1 s0    Y  d S )Nr   r1   T)deviceZnon_blocking)rU   rV   rW   rY   rO   rN   rF   rS   rE   r0   is_cudatora   rI   r>   r3   r4   rR   rH   r]   rJ   r)   r)   r*   _refill   s&    
z_OverlappingCpuLoader._refillc                 C   s(   | j s
J t| jdkr"| j  | jS rZ   )rY   rX   rR   rW   r[   rA   r)   r)   r*   _finish   s    

z_OverlappingCpuLoader._finishc                 C   s"   | j rtd| j||f d S )Nz&cannot add items after loading started)rT   RuntimeErrorrF   rH   r<   r)   r)   r*   r@      s    z_OverlappingCpuLoader.addc                 C   s.   | j r
d S d| _ | jjdd d |   d S )NTc                 S   s   | d S rZ   r)   xr)   r)   r*   <lambda>       z5_OverlappingCpuLoader.start_loading.<locals>.<lambda>key)rT   rF   sortrd   rA   r)   r)   r*   rB      s
    z#_OverlappingCpuLoader.start_loadingc                 c   s<   |    | js*|  }|   |E d H  q|  E d H  d S r/   )rB   rY   r`   rd   re   )r=   r^   r)   r)   r*   rC      s    z_OverlappingCpuLoader.values)NrM   )r"   r#   r$   rG   propertyrY   r`   rd   re   r@   rB   rC   r)   r)   r)   r*   rL   v   s   


rL   c                 C   sB   d}| j d usJ | j jD ]}||9 }q| j jj}|tj| S Nr   )Ztensor_datar>   Z
propertiesdtyperU   _utilsZ_element_size)r6   r>   srp   r)   r)   r*   
_item_size   s    

rs   )rF   r.   c           	      C   s   | dkr|gS dd |D }dd |D }dd t | D }dd t | D }|jtdd t|D ]\}}|||   | qd|D ]>}tt|d	d
 dd }|| | ||  t|7  < q|S )Nr   c                 S   s   g | ]}|j tjkr|qS r)   typer   BYTE_IO.0wir)   r)   r*   
<listcomp>   rj   z+_split_by_size_and_type.<locals>.<listcomp>c                 S   s   g | ]}|j tjkr|qS r)   rt   rw   r)   r)   r*   rz      rj   c                 S   s   g | ]}g qS r)   r)   rx   rK   r)   r)   r*   rz      rj   c                 S   s   g | ]}d qS )r   r)   r{   r)   r)   r*   rz      rj   T)rl   reversec                 S   s   | d S ro   r)   rg   r)   r)   r*   ri      rj   z)_split_by_size_and_type.<locals>.<lambda>rk   r   )rangerm   rs   	enumeraterH   min)	ZbinsrF   bytes_wtensor_wZbucketsZbucket_sizesiry   rS   r)   r)   r*   _split_by_size_and_type   s    r   c                 C   s   |   }|jtjkr4t|tjs$J | |  n0t|t	j
sDJ |jt	dksXJ t	||  |   | }t||t|||S )Nr1   )tellru   r   rv   
isinstanceioBytesIOwrite	getbufferrU   r	   ra   saver:   r   )rW   data
write_itemstorage_keyr    r!   r)   r)   r*   _write_item   s    r   
file_queueresult_queueplannerrN   	use_fsyncc              	      sN  z0|   \}}}tj r:|dkr:t fdd|d}nt fdd}dd |D }	|	D ]}
|t|
|
 q\|  dd |D }g }t	|d	|}|D ]"}
 
|
}|t|||
| q| D ]&\}}
|jrJ |t|||
| q|rt|  W d    n1 s0    Y  || qW n tjyH   Y n0 d S )
Nr   c                    s
     | S r/   resolve_datarg   r   r)   r*   ri   	  rj   z)_write_files_from_queue.<locals>.<lambda>)rN   c                    s
     | S r/   r   rg   r   r)   r*   ri     rj   c                 S   s   g | ]}|j tjkr|qS r)   rt   rw   r)   r)   r*   rz     s   z+_write_files_from_queue.<locals>.<listcomp>c                 S   s   g | ]}|j tjkr|qS r)   rt   rw   r)   r)   r*   rz     s   wb)
get_nowaitrU   rV   Zis_availablerL   rD   r@   rs   rB   openr   rH   r   rC   rb   osfsyncfilenoputqueueEmpty)r   r   r   rN   r   	file_namer   Zwrite_itemsloaderr   r   r   Zwrite_resultsrW   r   r-   r)   r   r*   _write_files_from_queue   sH    



.r   c                       s   e Zd ZdZdeeejf eee	e	dd fddZ
edd	d
dZeedddZee ee dddZeeeee  dddZeeee  ddddZ  ZS )r   aa  
    Basic implementation of StorageWriter using file IO.

    This implementation makes the following assumptions and simplifications:

    * The checkpoint path is an empty or non-existing directory.
    * File creation is atomic

    The checkpoint consist of one file per write request plus
    a `.metadata` file with the serialized metadata.

    Tr   逖 N)pathsingle_file_per_rank
sync_filesthread_countper_thread_copy_aheadr.   c                    s0   t    t|| _|| _|| _|| _|| _dS )a  
        Initialize the writer pointing to `path`

        Args:
            path: diretory where the checkpoint will be writen to.
            single_file_per_rank: Produce one file per rank instead of one file per tensor/blob. Default to True.
            sync_files : force files to be synced to permanent storage. Default to True.
            thread_count: Number of IO threads to use to write. Default to 1.
            per_thread_copy_ahead: How many bytes to copy from the GPU ahead of saving then. Default 10Mb.

        N. B. If sync_files is disabled, there's no guarantee that the checkpoint will be consistent in the case of a failure.
        N)superrG   r   r   r   r   r   r   )r=   r   r   r   r   r   	__class__r)   r*   rG   ?  s    

zFileSystemWriter.__init__)is_coordinatorr.   c                 C   s   d S r/   r)   )r=   r   r)   r)   r*   set_up_storage_writerZ  s    z&FileSystemWriter.set_up_storage_writerplanr.   c                 C   s   | j jddd |S )NT)parentsexist_ok)r   mkdirr=   r   r)   r)   r*   prepare_local_plan]  s    z#FileSystemWriter.prepare_local_planglobal_planr.   c                 C   s   dd t |D }|S )Nc                 S   s*   g | ]"\}}t j|td | ddqS )__rK   r9   )dataclassesreplacer+   )rx   r   r   r)   r)   r*   rz   d  s   z8FileSystemWriter.prepare_global_plan.<locals>.<listcomp>)r~   )r=   r   Z	new_plansr)   r)   r*   prepare_global_plana  s    z$FileSystemWriter.prepare_global_planr   r   r.   c                    s<  |j d  fdd}t }| jrXt| j|jD ] }| }|| j| ||f q4n*|jD ]"}| }|| j| ||gf q^t }g }	t	d| jD ]2}
t
jt|||| j| jfd}|  |	| qt|||| j| jd |	D ]}|  qg }z|| 7 }qW n* tjy6   t }|| | Y S 0 d S )Nr   c                     s   j    t }  d7  | S ro   )r,   DEFAULT_SUFFIX)r   Z
file_countZstorage_planr)   r*   gen_filer  s    z-FileSystemWriter.write_data.<locals>.gen_filer   )targetargsr   )r9   r   Queuer   r   r   rF   r   r   r}   	threadingThreadr   r   r   startrH   joinr   r   r
   
set_result)r=   r   r   r   r   Zbucketr   r6   r   threadsrK   tresfutr)   r   r*   
write_dataj  sV    



zFileSystemWriter.write_data)metadataresultsr.   c                 C   s   t  }|D ]}|dd |D  q
||_| jd d*}t|| t|	  W d    n1 sj0    Y  | jd 
| jd  d S )Nc                 S   s   i | ]}|j |jqS r)   )r7   r9   )rx   wrr)   r)   r*   
<dictcomp>  rj   z+FileSystemWriter.finish.<locals>.<dictcomp>z.metadata.tmpr   	.metadata)dictupdater9   r   r   pickledumpr   r   r   rename)r=   r   r   Z
storage_mdZwr_listmetadata_filer)   r)   r*   finish  s    ,zFileSystemWriter.finish)TTr   r   )r"   r#   r$   r%   r   r&   r   PathLikeboolr(   rG   r   r   r   r   r   r   r
   r   r   r   r   __classcell__r)   r)   r   r*   r   1  s2       
Br   c                       sV   e Zd Zejeed fddZejfeeed fddZ	ed fdd	Z
  ZS )
r   )base_streamr    rX   c                    s&   t  | || _|| _| d d S rZ   )r   rG   r    rX   seek)r=   r   r    rX   r   r)   r*   rG     s    zSlicedBufferedReader.__init__)_SlicedBufferedReader__offset_SlicedBufferedReader__whencer.   c                    sD   |t jkr| j| }n |t jkr6t j}| j| j | }t ||S r/   )r   SEEK_SETr    SEEK_ENDrX   r   r   )r=   r   r   r   r)   r*   r     s    

zSlicedBufferedReader.seekr.   c                    s   t   | j S r/   )r   r   r    rA   r   r)   r*   r     s    zSlicedBufferedReader.tell)r"   r#   r$   r   	RawIOBaser(   rG   r   r   r   r   r   r)   r)   r   r*   r     s   r   c                       s   e Zd Zeeejf dd fddZedddZ	e
eed dd	d
ZedddZeeddddZe
e
dddZee
 ee
 dddZ  ZS )r   N)r   r.   c                    s    t    t|| _t | _d S r/   )r   rG   r   r   r   r9   )r=   r   r   r)   r*   rG     s    

zFileSystemReader.__init__)sinfoc                 C   s   t tj| dd|j|jS )NF)closefd)r   r   FileIOr   r    r!   )r=   filer   r)   r)   r*   _slice_file  s    zFileSystemReader._slice_filer   c                 C   sf  t  }|jD ](}| j|j }|j}||g | q| D ]\}}| j| d}	|D ]}
| j|
j }| 	|	|}|
j
tjkrt||j}|d ||
| q^tttj|dd}t||
j|
j}||
 }| | ksJ d|
j d|  d|  || ||
| q^W d    q>1 sF0    Y  q>t }| d  |S )Nrbr   r1   )Zmap_locationzreq z mismatch sizes z vs )!r   rF   r9   Zstorage_indexr   
setdefaultrH   r   r   r   ru   r   rv   r   r   readr!   r   Z
load_bytesr   r	   rU   loadr   Zstorage_offsetslengthsZresolve_tensorr0   r>   Zcopy_Zcommit_tensorr
   r   )r=   r   r   Zper_fileZ	read_itemZitem_mdr   r   reqsr   reqZ
file_slicebytesr-   Ztarget_tensorr   r)   r)   r*   	read_data  s<    



0
zFileSystemReader.read_datar   c                 C   s>   | j d d}t|W  d    S 1 s00    Y  d S )Nr   r   )r   r   r   r   )r=   r   r)   r)   r*   read_metadata  s    zFileSystemReader.read_metadata)r   r   r.   c                 C   s   |j | _ | j d usJ d S r/   r   )r=   r   r   r)   r)   r*   set_up_storage_reader  s    z&FileSystemReader.set_up_storage_readerr   c                 C   s   |S r/   r)   r   r)   r)   r*   r     s    z#FileSystemReader.prepare_local_planr   c                 C   s   |S r/   r)   )r=   r   r)   r)   r*   r   	  s    z$FileSystemReader.prepare_global_plan)r"   r#   r$   r   r&   r   r   rG   r   r   r   r   r
   r   r   r   r   r   r   r   r   r   r)   r)   r   r*   r     s   &r   )<abcr   r   r   r   rP   r   r   r   r   r   typingr   r   r   r   rU   r	   Ztorch.futuresr
   pathlibr   r   r   r   rI   r   r   r   r   r   r   r   r   r   r   r   r   Ztorch.distributed._shard._utilsr   __all__r   r+   r   r5   r:   r;   rD   rL   r(   rs   r   r   r   r   r   r   BufferedReaderr   r   r)   r)   r)   r*   <module>   sZ   (
O
6 	