a
    dQ-                     @   s  d dl Z d dlZd dlmZmZmZmZmZmZm	Z	 d dl
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 d dlmZ d dlm  mZ d dlmZmZ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(m)Z) d dl*m+Z+ d dl,m-Z- d dl.m/Z/m0Z0 ee1eeee2  ee2 f f Z3dgZ4e2e1dddZ5d$eej6 edddZ7ej8e9dddZ:eee2 ej8dddZ;eee3eej6 f dddZ<G d d! d!e)Z=ee1ej>ed"d#dZ?dS )%    N)DictListOptionalSequenceTupleUnioncast)LoadPlan)ShardedTensor)TensorProperties)Shard)ChunkShardingSpec)BytesStorageMetadataMetadataMetadataIndexSTATE_DICT_TYPETensorStorageMetadata)_create_sharded_read_items_create_read_items)_remote_device)DTensor)DefaultLoadPlanner)_shard_tensor)unflatten_state_dict)_element_wise_add_element_wise_sub!load_sharded_optimizer_state_dict)global_rankreturnc                 C   s"   t j rd| t j   S dS )Nzcuda:cpu)torchcudaZis_availabledevice_count)r    r#   o/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/distributed/checkpoint/optimizer.py_gen_rank_device4   s    
r%   )pgr   c                    sX    d u r dd t t D }n fddt   D }tdttttt	f  |dS )Nc                 S   s    g | ]}d | dt | qS rank:/)r%   .0idxr#   r#   r$   
<listcomp>>   s   z(_create_colwise_spec.<locals>.<listcomp>c              
      s(   g | ] }d | dt t | qS r'   )r%   distZget_global_rankr*   r&   r#   r$   r-   C   s   r   Zdim
placements)
ranger.   get_world_sizesizer   r   r   r   r   str)r&   r1   r#   r/   r$   _create_colwise_spec:   s    


r6   )valr   c                 C   s   t | tu rZt|  dkr dS t |  d jtu r:dS t |  d jtu rtdn0t | tu rt | jtu st | jtu rtddS )Nr   FTz2Cannot handle DTensor nested insided ShardedTensorzCannot handle nested DTensor)typer
   lenlocal_shardstensorr   
ValueErrorZ_local_tensor)r7   r#   r#   r$   _is_nested_tensorM   s     r=   )propsr4   r   c              
   C   s,   t j|| j| j| j| jtt jt j	 dS )N)r4   dtypelayoutrequires_grad
pin_memorydevice)
r    emptyr?   r@   rA   rB   r   rC   r!   Zcurrent_device)r>   r4   r#   r#   r$   _alloc_tensor_   s    rE   )
state_dictr   c                 C   s   i }d}|   D ]r\}}d| f||< t|rt| dksHJ dt|tsZJ d| d }|jj|jj	f||< |j
j}q||fS )a5  
    We have to load the right TP slice of the optimizer state.
    This is not easy since the per-tensor slicing can't be inferred from checkpoint metadata.
    We take advantage of the model state_dict producing a sliced ST to figure out what we need to load.
    This is pretty fragile and it might be easier for FSDP to compute this info for us.
    Returns a dictionary where keys are the same of the state_dict and the value is a tuple of
    (offset, size) for the current rank TP slice.
    N.B. The state_dict *MUST* come from FSDP.sharded_state_dict.
    N   z%Cannot handle ST with multiple shardsz$Can only handle nested ShardedTensorr   )itemsr4   r=   r9   r:   
isinstancer
   metadatashard_offsetsshard_sizesr;   Z_process_group)rF   specsdp_pgkeyvalueZshardr#   r#   r$   _get_state_dict_2d_layoutj   s,    
rQ   c                       sv   e Zd ZU eeef ed< eed< eed< eee	e
 f dd fddZedd	d
Zeejd fddZ  ZS )_ReaderWithOffsettranslationrF   rJ   N)fqn_to_offsetr   c                    s*   t    || _ti | _i | _i | _d S N)super__init__rT   r   rJ   rF   rS   )selfrT   	__class__r#   r$   rW      s
    

z_ReaderWithOffset.__init__)r   c                 C   s"  g }i | _ | j D ]\}}| jj| }t|tsF|t|||7 }q|| jvrb|t|||7 }q| j| }t	|
 dksJ |
 d }t|j}t|j||_t|j|g}t|tt||}	|	D ]D}
|
jjd usJ t|
jj|}tj|
jt|d}|| j |
j< q||	7 }qt|S )NrG   r   )offset)rS   rF   rH   rJ   state_dict_metadatarI   r
   r   rT   r9   r:   copydeepcopyr   rK   r   r;   r   r   r   Z
dest_indexr[   r   dataclassesreplacer    Sizer	   )rX   requestsZfqnobjZmdr[   Zoriginal_shardshard_mdr:   reqsZwiZoriginal_offsetZoriginal_indexr#   r#   r$   create_local_plan   s@    



z#_ReaderWithOffset.create_local_plan)indexr   c                    s   t  | j||S rU   )rV   lookup_tensorrS   get)rX   rg   rY   r#   r$   rh      s    z_ReaderWithOffset.lookup_tensor)__name__
__module____qualname__r   r   __annotations__r   r   r5   r   intrW   r	   rf   r    Tensorrh   __classcell__r#   r#   rY   r$   rR      s   
 )rR   )model_state_dictoptimizer_keystorage_readerr   c                 C   s  |  }t| \}}|du r<tddd tt D d}nt|}i }i }|j D ]J\}	}
|j	|	 }|d |krxqVt
|
trd||	< qV|
j dkrt|
j|
j||	< qV|du rtt|
j|
j|||	< qV|d }||d|
jfd }|t||
j}g }t|}|jD ]<}tt|j |kr4q|tt|
j|j|d	 qtj|||d
}||v r|| d durtt t! || d ||	< |||	< qVt"j#|||durt$|ndd t%||j	}|S )a  
    Loads a state_dict to be used in conjuntion with FSDP sharded optimizer state.
    This is the current recommended way to checkpoint is FSDP
    >>> # xdoctest: +SKIP
    >>> import torch.distributed.checkpoint as dist_cp
    >>> # Save
    >>> model: torch.nn.Model
    >>> optim_params = model.parameters()
    >>> optim = torch.optim.SGD(optim_params, lr=0.01)
    >>>
    >>> with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
    >>>     state_dict = {
    >>>         "optimizer": FSDP.sharded_optim_state_dict(model, optim, optim_params),
    >>>         "model": model.state_dict()
    >>>     }
    >>>     dist_cp.save_state_dict(
    >>>         state_dict=optim_state,
    >>>         storage_writer=dist_cp.FileSystemWriter("checkpoint"),
    >>>         planner=dist_cp.DefaultSavePlanner(),
    >>>     )
    >>>
    >>> # Load
    >>> with FSDP.state_dict_type(model_tp, StateDictType.SHARDED_STATE_DICT):
    >>>     model_state_dict = model_tp.state_dict()
    >>>     checkpoint = {
    >>>         "model": model_state_dict
    >>>     }
    >>>     dist_cp.load_state_dict(
    >>>         state_dict=checkpoint,
    >>>         storage_reader=dist_cp.FileSystemReader(checkpoint_file),
    >>>         planner=dist_cp.DefaultLoadPlanner(),
    >>>     )
    >>>     model.load_state_dict(checkpoint["model_state"])
    >>>
    >>>     optim_state = sp_cp.load_sharded_optimizer_state_dict(
    >>>         model_state_dict,
    >>>         optimizer_key="optimizer",
    >>>         storage_reader=dist_cp.FileSystemReader("checkpoint"),
    >>>     )
    >>>
    >>>     flattened_osd = FSDP.flatten_sharded_optim_state_dict(
    >>>        optim_state["optimizer"], model, optim
    >>>     )
    >>>
    >>>     optim.load_state_dict(flattened_osd)
    Nr   c                 S   s&   g | ]}d | d|t j   qS )r(   z/cuda:)r    r!   r"   )r+   ir#   r#   r$   r-     s   z5load_sharded_optimizer_state_dict.<locals>.<listcomp>r0   z
<bytes_io>rG      )r;   rJ   )Zprocess_group)rF   rs   Zplanner)&Zread_metadatarQ   r   r2   r.   r3   r6   r\   rH   Zplanner_datarI   r   r4   ZnumelrE   Z
propertiesr   ri   Zbuild_metadatar    ra   Zget_rankZshards_metadatar   r   Z	placementZrankappendr   rL   r
   Z+_init_from_local_shards_and_global_metadatar   rn   dist_cpZload_state_dictrR   r   )rq   rr   rs   rJ   Zlayout_specsrN   Zsharding_specrF   rT   rO   rP   Zkey_pathZspec_keyZ
alloc_sizeZst_mdr:   Zcurrent_rankrd   str#   r#   r$   r      s    3





	
)N)@r]   r_   typingr   r   r   r   r   r   r   Z$torch.distributed.checkpoint.plannerr	   r    Ztorch.distributedZdistributedr.   Z+torch.distributed._shard.sharded_tensor.apir
   Z0torch.distributed._shard.sharded_tensor.metadatar   Z-torch.distributed._shard.sharded_tensor.shardr   Z:torch.distributed._shard.sharding_spec.chunk_sharding_specr   Ztorch.distributed.checkpoint
checkpointrw   Z%torch.distributed.checkpoint.metadatar   r   r   r   r   Z,torch.distributed.checkpoint.planner_helpersr   r   Ztorch.distributed.remote_devicer   Ztorch.distributed._tensorr   Z,torch.distributed.checkpoint.default_plannerr   Ztorch.distributed._shard.apir   Z)torch.distributed.checkpoint._nested_dictr   Z"torch.distributed.checkpoint.utilsr   r   r5   rn   ZSTATE_DICT_2D_LAYOUT__all__r%   ZProcessGroupr6   ro   boolr=   rE   rQ   rR   ZStorageReaderr   r#   r#   r#   r$   <module>   sL   $  $: