a
    d5                     @   s  U 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	m
Z
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mZmZmZmZmZmZ d dlmZmZmZm Z m!Z!m"Z"m#Z# d d	l$m%Z%m&Z&m'Z' d d
l(m)Z)m*Z* d dl+m,Z, d dl-m.Z. d dl/m0Z0 d dl1m2Z2 e3e4Z5ej6e7d< g dZ8G dd deZ9G dd deZ:ee;ef e!edddZ<e	e e	e dddZ=ee;ef e>edddZ?e	e e
e	e e!f dddZ@e"e!d d!d"ZAeee>d#d$d%ZBejCee>d&d'd(ZDe	e e!e>d)d*d+ZEdS ),    N)ChainMap)reduce)ListTupleDictAnyUnioncast)narrow_tensor_by_index)ShardedTensor)SavePlannerLoadPlannerSavePlanLoadPlanReadItem	WriteItemWriteItemType)BytesStorageMetadataChunkStorageMetadataTensorStorageMetadataMetadataIndexMetadataSTATE_DICT_TYPESTORAGE_TYPES)_create_read_items_create_write_items"_create_default_metadata_only_plan)FLATTEN_MAPPINGflatten_state_dict)_flatten_sharded_tensors)dedup_tensors)find_state_dict_object)set_elementlogger)DefaultSavePlannerDefaultLoadPlannercreate_default_local_load_plancreate_default_global_load_plancreate_default_local_save_plancreate_default_global_save_planc                   @   s   e Zd ZU eed< deeeddddZeedddd	Ze	d
ddZ
ee	 eee	 ef dddZe	e	dddZeeejejf dddZeedddZeedddZdS )r$   mappingsTN)r   flatten_sharded_tensorsdedup_replicated_tensorsreturnc                 C   s   || _ || _|| _i | _d S N)r   r+   r,   r*   )selfr   r+   r,    r0   u/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/distributed/checkpoint/default_planner.py__init__G   s    zDefaultSavePlanner.__init__
state_dictis_coordinatorr-   c                 C   s2   | j rt |\}| _| jr"t|}|| _|| _d S r.   )r   r*   r+   r   r4   r5   )r/   r4   r5   r0   r0   r1   set_up_plannerR   s    z!DefaultSavePlanner.set_up_plannerr-   c                 C   s0   t | j| j}| jr$tj|| jd}|| _| jS )NZplanner_data)r(   r4   r5   r   dataclassesreplacer*   plan)r/   r;   r0   r0   r1   create_local_plan\   s    z$DefaultSavePlanner.create_local_plan	all_plansr-   c                 C   sr   | j rt|}t|\}}| jrHdd |D }tt| }tj||d}t||sZt	d|| _
|| _| j
| jfS )Nc                 S   s   g | ]
}|j qS r0   r8   ).0pr0   r0   r1   
<listcomp>s       z9DefaultSavePlanner.create_global_plan.<locals>.<listcomp>r8   zFailed to validate global plan)r,   r    r)   r   dictr   r9   r:   _validate_global_plan
ValueErrorglobal_planmetadata)r/   r>   rF   rG   Zplanner_data_dictZmerged_mappingsr0   r0   r1   create_global_planf   s    
z%DefaultSavePlanner.create_global_plannew_planr-   c                 C   s
   || _ |S r.   )r;   r/   rJ   r0   r0   r1   finish_plan   s    zDefaultSavePlanner.finish_plan)
write_itemr-   c                 C   s   |  |j}| ||S r.   )lookup_objectindextransform_object)r/   rM   objectr0   r0   r1   resolve_data   s    zDefaultSavePlanner.resolve_datarO   r-   c                 C   s   t | j|S zo
        This is an extension from the planner interface to make it easy to extend the default planner
        r!   r4   r/   rO   r0   r0   r1   rN      s    z DefaultSavePlanner.lookup_object)rM   rQ   c                 C   s(   |j tjkr$t }t|| |}|S rT   )typer   BYTE_IOioBytesIOtorchsave)r/   rM   rQ   bytesr0   r0   r1   rP      s
    z#DefaultSavePlanner.transform_object)TTT)__name__
__module____qualname__r   __annotations__boolr2   r   r6   r   r<   r   r   r   rH   rL   r   r   r[   TensorrY   rZ   rR   r   r   rN   rP   r0   r0   r0   r1   r$   D   s.   
   
r$   c                   @   s   e Zd ZU dZeed< eed< d$eeddddZee	edd	d
dZ
edddZee ee dddZeedddZeejddddZedddZeejddddZeejddd Zeejd!d"d#ZdS )%r%   z
    DefaultLoadPlanner that adds multiple features on top of LoadPlanner.

    In particular it adds the following:

    flatten_state_dict: Handle state_dict with nested dicts
    flatten_sharded_tensors: For FSDP in 2D parallel mode
    original_state_dictr*   TN)r   r+   r-   c                 C   s   || _ || _i | _i | _d S r.   )r   r+   rd   r*   )r/   r   r+   r0   r0   r1   r2      s    zDefaultLoadPlanner.__init__)r4   rG   r5   r-   c                 C   s>   || _ | jrt|}| jr(t|\}| _|| _|| _|| _d S r.   )rd   r+   r   r   r*   r4   rG   r5   )r/   r4   rG   r5   r0   r0   r1   r6      s    z!DefaultLoadPlanner.set_up_plannerr7   c                 C   s   t | j| jS r.   )r&   r4   rG   )r/   r0   r0   r1   r<      s    z$DefaultLoadPlanner.create_local_plan)rF   r-   c                 C   s   t |S r.   )r'   )r/   rF   r0   r0   r1   rH      s    z%DefaultLoadPlanner.create_global_planrI   c                 C   s   |S r.   r0   rK   r0   r0   r1   rL      s    zDefaultLoadPlanner.finish_plan)	read_itemvaluer-   c                 C   s>   | j r&t| j| j|jj t| nt|| j|jj< d S r.   )	r   r"   rd   r*   
dest_indexfqnr[   loadr4   )r/   re   rf   r0   r0   r1   
load_bytes   s    zDefaultLoadPlanner.load_bytes)re   c                 C   s   |  |j}| ||S r.   )lookup_tensorrg   transform_tensorr/   re   tensorr0   r0   r1   resolve_tensor   s    z!DefaultLoadPlanner.resolve_tensor)re   rn   r-   c                 C   s   d S r.   r0   rm   r0   r0   r1   commit_tensor   s    z DefaultLoadPlanner.commit_tensorrS   c                 C   s   t | j|S rT   rU   rV   r0   r0   r1   rk      s    z DefaultLoadPlanner.lookup_tensor)re   rn   c                 C   s   t ||j|jS rT   )r
   Zdest_offsetslengthsrm   r0   r0   r1   rl      s    
z#DefaultLoadPlanner.transform_tensor)TT)r^   r_   r`   __doc__r   ra   r   rb   r2   r   r6   r   r<   r   rH   rL   r   rY   rZ   rj   ro   r[   rc   rp   r   rk   rl   r0   r0   r0   r1   r%      s.   
	  
r%   )r4   rG   r-   c                 C   s8   g }|   D ]"\}}|j| }|t|||7 }qt|S r.   )itemsstate_dict_metadatar   r   )r4   rG   requestsrh   objmdr0   r0   r1   r&      s
    

r&   r=   c                 C   s   | S )z
    Create global load plan used by DefaultLoadPlanner.

    The default load behavior involved no global coordination and this function
    currently doesn't change the local plans.
    r0   )r>   r0   r0   r1   r'     s    	r'   r3   c                 C   s:   g }|   D ]$\}}t|ts"|r|t||7 }qt|S )a  
    Create the ``SavePlan`` used by DefaultSavePlanner.

    On non-coordinator ranks, this function ignores tensors and non-tensor objects,
    only producing writes for ShardedTensor objects.

    On the coordinator rank, produce writes for all values.
    )rs   
isinstancer   r   r   )r4   r5   ru   rh   rv   r0   r0   r1   r(     s
    r(   c           	      C   s  i }g }| D ]}g }|j D ]}|jtjks:|jj|vs:J |jtjkr`t ||jj< || q|j	dusnJ t
t||jjt|j	j|j	jg d}tj|jt|jd}tj||d}|| |j	jdusJ d|jj d|j|j	j q|tj||d q|t|fS )a  
    Create the global plan and metadata used by DefaultSavePlanner.

    Metadata is produced by concatenating the metadata of all ``WriteItem`` from the supplied plans.

    The only global planning change is to update index hints in all ``MetadataIndex`` objects.
    N)
propertiessizechunks)rO   zZ
                    Cannot create MD for tensor without bounds.
                    FQN: 
                )rs   )rs   rW   r   ZSHARDrO   rh   rX   r   appendZtensor_datar	   r   
setdefaultry   rz   r9   r:   lenr{   chunkr   )	r>   rw   Z	new_plansr;   Z	new_itemsitemZ	tensor_mdZ	new_indexZnew_itemr0   r0   r1   r)   !  sH    



r)   )r4   r-   c                 C   s   t | }t|g\}}|S )z^
    Return the ``Metadata`` if DefaultSavePlanner was used to checkpoint ``state_dict``.
    )r   r)   )r4   r;   _rw   r0   r0   r1   _create_default_local_metadataT  s    r   )box0box1r-   c                 C   sd   t | j}t|D ]L}| j| |j| |j|  kr: dS |j| | j| | j|  kr dS qdS )zC
    Checks if two boxes overlap. Tuples are (offset, lengths)
    FT)r   offsetsrangesizes)r   r   Zndimsir0   r0   r1   _check_box_overlap]  s    
r   )outer_box_size	inner_boxr-   c                 C   s`   t t| D ]N}|j| dk r$ dS |j| dk r8 dS |j| |j|  | | kr dS qdS )Nr   FT)r   r   r   r   )r   r   r   r0   r0   r1   _check_box_boundsr  s    r   )rF   rG   r-   c           
      C   s  d}|j  D ]\}}t|tr"qt|jdkr2qd}t|jD ]\}}t|j|sxt	
d| d|j d| d d}|ttj|jd7 }|j|d d  D ].}t||rt	
d	| d
| d|  d}qq@ttj|jd}	||	krt	
d| d|	 d| d d}q|S )NTr   z
                        key:z> has out of bounds chunk:
                        tensor-size:z chunk: z
                    F   zkey:z has overlapping chunks:  z
                    key:z1 invalid fill tensor-volume:
                    z chunks-volume: r|   )rt   rs   rx   r   r   rz   	enumerater{   r   r#   warningr   operatormulr   r   )
rF   rG   all_goodkeyrf   Zchunks_volumeZ	chunk_idxZchunk0Zchunk1Ztensor_volumer0   r0   r1   rD     sN    

rD   )Fr9   rY   loggingr   collectionsr   	functoolsr   typingr   r   r   r   r   r	   r[   Ztorch.distributed._shard._utilsr
   Z'torch.distributed._shard.sharded_tensorr   Z$torch.distributed.checkpoint.plannerr   r   r   r   r   r   r   Z%torch.distributed.checkpoint.metadatar   r   r   r   r   r   r   Z,torch.distributed.checkpoint.planner_helpersr   r   r   Z)torch.distributed.checkpoint._nested_dictr   r   Z2torch.distributed.checkpoint._sharded_tensor_utilsr   Z+torch.distributed.checkpoint._dedup_tensorsr    Z"torch.distributed.checkpoint.utilsr!   Z&torch.distributed.checkpoint._traverser"   	getLogger__file__r#   Loggerra   __all__r$   r%   strr&   r'   rb   r(   r)   r   r   Sizer   rD   r0   r0   r0   r1   <module>   sX   
 $
$
XS
3
