§
    ŠŠtj/4  ã                   ón  — d dl Z d dl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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! d dl"m#Z#m$Z$ d dl%m&Z& d dl'm(Z( d dl)m*Z*m+Z+m,Z, d dl-m.Z. d dl/m0Z0 d dl1m2Z2 d dl3m4Z4 e5e6e7ee8         dz  ee8         f         f         Z9dgZ:d+de8de6de6fd„Z;	 d,dej<        dz  defd„Z=dej>        de?fd„Z@	 d+ded ee8         de6dej>        fd!„ZAd"ede7e9ej<        dz  f         fd#„ZB G d$„ d%e¦  «        ZC	 d,d&ed'e6d(e(d)e!dz  def
d*„ZDdS )-é    N)ÚSequence)Úcast)Ú_get_device_module)ÚShardedTensor)ÚTensorProperties)ÚShard)ÚChunkShardingSpec)Úunflatten_state_dict)ÚDefaultLoadPlanner)ÚBytesStorageMetadataÚChunkStorageMetadataÚMetadataÚMetadataIndexÚSTATE_DICT_TYPEr   ÚTensorStorageMetadata)ÚLoadPlanÚLoadPlanner)Ú_create_read_itemsÚ create_read_items_for_chunk_list)Úload_state_dict)ÚStorageReader)Ú_element_wise_addÚ_element_wise_subÚ_normalize_device_info)Ú_get_default_group)Ú_create_chunk_sharded_tensor)Ú_remote_device)ÚDTensorÚ!load_sharded_optimizer_state_dictÚcudaÚglobal_rankÚdevice_typeÚreturnc                 ó¦   — |dk    rdS t          |¦  «        }|                     ¦   «         r%t          || |                     ¦   «         z  ¦  «        S dS )NÚcpu)r   Úis_availabler   Údevice_count)r!   r"   Údevice_modules      úd/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributed/checkpoint/optimizer.pyÚ_gen_rank_devicer*   8   sb   € Ø�eÒÐØˆuÝ& {Ñ3Ô3€MØ×!Ò!Ñ#Ô#ð 
Ý%Ø˜ }×'AÒ'AÑ'CÔ'CÑCñ
ô 
ð 	
ð ˆ5ó    Úpgc                 óv  ‡ ‡— t           j                             ‰ ¦  «        j        Š‰ €-ˆfd„t	          t          j        ¦   «         ¦  «        D ¦   «         }n.ˆ ˆfd„t	          ‰                      ¦   «         ¦  «        D ¦   «         }t          dt          t          t          t          z           |¦  «        ¬¦  «        S )Nc           	      ó<   •— g | ]}d |› dt          |‰¦  «        › �‘ŒS ©úrank:ú/)r*   )Ú.0ÚidxÚpg_device_types     €r)   ú
<listcomp>z(_create_colwise_spec.<locals>.<listcomp>H   sE   ø€ ð 
ð 
ð 
àð B�CÐAÐAÕ*¨3°Ñ?Ô?ÐAÐAð
ð 
ð 
r+   c                 ób   •— g | ]+}d |› dt          t          j        ‰|¦  «        ‰¦  «        › �‘Œ,S r/   )r*   ÚdistÚget_global_rank)r2   r3   r,   r4   s     €€r)   r5   z(_create_colwise_spec.<locals>.<listcomp>M   sR   ø€ ð 
ð 
ð 
àð \�CÐ[Ð[Õ*­4Ô+?ÀÀCÑ+HÔ+HÈ.ÑYÔYÐ[Ð[ð
ð 
ð 
r+   r   ©ÚdimÚ
placements)r7   Údistributed_c10dÚ_get_pg_default_deviceÚtypeÚrangeÚget_world_sizeÚsizer	   r   Úlistr   Ústr)r,   r;   r4   s   ` @r)   Ú_create_colwise_specrD   C   sÏ   øø€ õ Ô*×AÒAÀ"ÑEÔEÔJ€NØ	€zð
ð 
ð 
ð 
å�TÔ0Ñ2Ô2Ñ3Ô3ð
ñ 
ô 
ˆ
ˆ
ð

ð 
ð 
ð 
ð 
å˜RŸWšW™YœYÑ'Ô'ð
ñ 
ô 
ˆ
õ ØÝ��^­cÑ1Ô2°JÑ?Ô?ðñ ô ð r+   Úvalc                 ó&  — t          | ¦  «        t          u rŸt          |                      ¦   «         ¦  «        dk    rdS t          |                      ¦   «         d         j        ¦  «        t          u rdS t          |                      ¦   «         d         j        ¦  «        t
          u rt          d¦  «        ‚n[t          | ¦  «        t
          u rEt          | j        ¦  «        t
          u st          | j        ¦  «        t          u rt          d¦  «        ‚dS )Nr   FTz1Cannot handle DTensor nested inside ShardedTensorzCannot handle nested DTensor)r>   r   ÚlenÚlocal_shardsÚtensorr   Ú
ValueErrorÚ_local_tensor)rE   s    r)   Ú_is_nested_tensorrL   W   sî   € ÝˆC�y„y•MÐ!Ð!Ýˆs×ÒÑ!Ô!Ñ"Ô" aÒ'Ð'Ø�5Ý�× Ò Ñ"Ô" 1Ô%Ô,Ñ-Ô-µÐ>Ð>Ø�4Ý�× Ò Ñ"Ô" 1Ô%Ô,Ñ-Ô-µÐ8Ð8ÝÐPÑQÔQÐQð 9å	ˆc‰Œ•gÐ	Ð	ÝˆSÔÑÔ¥7Ð*Ð*­d°3Ô3DÑ.EÔ.EÍÐ.VÐ.VåÐ7Ñ8Ô8Ð8Øˆ5r+   ÚpropsrA   c                 óF  — |dk    r:t          t          j        t          |¦  «                             ¦   «         ¦  «        }n4t          j        |t          |¦  «                             ¦   «         ¦  «        }t          j        || j        | j        | j        | j	        |¬¦  «        S )Nr%   )rA   ÚdtypeÚlayoutÚrequires_gradÚ
pin_memoryÚdevice)
r   ÚtorchrS   r   Úcurrent_deviceÚemptyrO   rP   rQ   rR   )rM   rA   r"   rS   s       r)   Ú_alloc_tensorrW   f   s™   € ð �eÒÐÝ•e”lÕ$6°{Ñ$CÔ$C×$RÒ$RÑ$TÔ$TÑUÔUˆˆå”ØÕ+¨KÑ8Ô8×GÒGÑIÔIñ
ô 
ˆõ Œ;ØØŒkØŒ|ØÔ)ØÔ#Øðñ ô ð r+   Ú
state_dictc                 óÈ  — i }d}|                       ¦   «         D ]Æ\  }}d|                     ¦   «         f||<   t          |¦  «        r™t          |                     ¦   «         ¦  «        dk    st          d¦  «        ‚t          |t          ¦  «        st          d¦  «        ‚|                     ¦   «         d         }|j        j	        |j        j
        f||<   |j        j        }ŒÇ||fS )a+  
    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   )ÚitemsrA   rL   rG   rH   ÚAssertionErrorÚ
isinstancer   ÚmetadataÚshard_offsetsÚshard_sizesrI   Ú_process_group)rX   ÚspecsÚdp_pgÚkeyÚvalueÚshards         r)   Ú_get_state_dict_2d_layoutrg   z   só   € ð #%€EØ&*€EØ ×&Ò&Ñ(Ô(ð 0ð 0‰
ˆˆUØ˜EŸJšJ™LœLÐ)ˆˆc‰
Ý˜UÑ#Ô#ð 
	0Ý�u×)Ò)Ñ+Ô+Ñ,Ô,°Ò1Ð1Ý$Ð%LÑMÔMÐMÝ˜e¥]Ñ3Ô3ð MÝ$Ð%KÑLÔLÐLØ×&Ò&Ñ(Ô(¨Ô+ˆEà”Ô,Ø”Ô*ðˆE�#‰Jð ”LÔ/ˆEøð 	Øðð r+   c                   óž   ‡ — e Zd ZU eeef         ed<   eed<   eed<   deee	e
         f         ddfˆ fd„Zdefd„Zd	edej        fˆ fd
„Zˆ xZS )Ú_ReaderWithOffsetÚtranslationrX   r^   Úfqn_to_offsetr#   Nc                 óš   •— t          ¦   «                              ¦   «          || _        t          i ¦  «        | _        i | _        i | _        d S ©N)ÚsuperÚ__init__rk   r   r^   rX   rj   )Úselfrk   Ú	__class__s     €r)   ro   z_ReaderWithOffset.__init__£   sC   ø€ Ý‰Œ×ÒÑÔÐØ*ˆÔÝ  ™œˆŒØˆŒØˆÔÐÐr+   c           	      óÌ  — g }i | _         | j                             ¦   «         D �]²\  }}| j        j        |         }t          |t          ¦  «        s|t          |||¦  «        z  }ŒB|| j        vr|t          |||¦  «        z  }Œ`| j        |         }t          | 
                    ¦   «         ¦  «        dk    st          d¦  «        ‚| 
                    ¦   «         d         }t          t          j        t          |j        j        |¦  «        ¦  «        t          j        |j        j        ¦  «        ¬¦  «        g}t%          |t'          t(          |¦  «        |¦  «        }|D ]s}	|	j        j        €t          d¦  «        ‚t/          |	j        j        |¦  «        }
t1          j        |	j        t          j        |
¦  «        ¬¦  «        }|| j         |	j        <   Œt||z  }�Œ´t5          |¦  «        S )NrZ   z Expected exactly one local shardr   )ÚoffsetsÚsizesz"dest_index.offset must not be None)Úoffset)rj   rX   r[   r^   Ústate_dict_metadatar]   r   r   rk   rG   rH   r\   r   rT   ÚSizer   r_   r`   r   r   r   Ú
dest_indexru   r   ÚdataclassesÚreplacer   )rp   ÚrequestsÚfqnÚobjÚmdru   Úoriginal_shardÚlocal_chunksÚreqsÚriÚoriginal_offsetÚoriginal_indexs               r)   Úcreate_local_planz#_ReaderWithOffset.create_local_planª   sð  € ØˆØˆÔØœ×-Ò-Ñ/Ô/ð &	ñ &	‰HˆC�Ø”Ô2°3Ô7ˆBÝ˜c¥=Ñ1Ô1ð ØÕ.¨s°B¸Ñ<Ô<Ñ<�Øà˜$Ô,Ð,Ð,ØÕ.¨s°B¸Ñ<Ô<Ñ<�ØàÔ'¨Ô,ˆFå�s×'Ò'Ñ)Ô)Ñ*Ô*¨aÒ/Ð/Ý$Ð%GÑHÔHÐHØ ×-Ò-Ñ/Ô/°Ô2ˆNå$Ý!œJÝ)¨.Ô*AÔ*OÐQWÑXÔXñô õ  œ* ^Ô%<Ô%HÑIÔIð	ñ ô ðˆLõ 4Ø•TÕ/°Ñ4Ô4°lñô ˆDð
 ð Að A�Ø”=Ô'Ð/Ý(Ð)MÑNÔNÐNÝ"3°B´MÔ4HÈ&Ñ"QÔ"Q�Ý!,Ô!4Ø”M­%¬*°_Ñ*EÔ*Eð"ñ "ô "�ð 3A�Ô  ¤Ñ/Ð/à˜ÑˆH‰HÝ˜Ñ!Ô!Ð!r+   Úindexc                 óx   •— t          ¦   «                              | j                             ||¦  «        ¦  «        S rm   )rn   Úlookup_tensorrj   Úget)rp   r†   rq   s     €r)   rˆ   z_ReaderWithOffset.lookup_tensorÖ   s.   ø€ Ý‰wŒw×$Ò$ TÔ%5×%9Ò%9¸%ÀÑ%GÔ%GÑHÔHÐHr+   )Ú__name__Ú
__module__Ú__qualname__Údictr   Ú__annotations__r   r   rC   r   Úintro   r   r…   rT   ÚTensorrˆ   Ú__classcell__)rq   s   @r)   ri   ri   �   sß   ø€ € € € € € Ø�m ]Ð2Ô3Ð3Ð3Ñ3ØÐÐÑàÐÐÑð d¨3°¸´Ð+=Ô&>ð À4ð ð ð ð ð ð ð*" 8ð *"ð *"ð *"ð *"ðXI =ð I°U´\ð Ið Ið Ið Ið Ið Ið Ið Ið Ið Ir+   ri   Úmodel_state_dictÚoptimizer_keyÚstorage_readerÚplannerc                 óR  — |                      ¦   «         }t          | ¦  «        \  }}t          j                             |¦  «        j        }t          |¦  «        }|€wg }	t          t          j        ¦   «         ¦  «        D ]B}
t          ||
| 
                    ¦   «         z  ¦  «        }|	                     d|
› d|› �¦  «         ŒCt          d|	¬¦  «        }nt          |¦  «        }i }i }|j                             ¦   «         D �]n\  }}|j        |         }|d         |k    rŒ t#          |t$          ¦  «        rd||<   Œ;|j                             ¦   «         dk    rt+          |j        |j        |¦  «        ||<   Œw|€qt/          t+          |j        |j        |¦  «        t          j        ¦   «         t          j        ¦   «         | 
                    ¦   «         t3          ¦   «         ¬¦  «        ||<   Œê|d	         }|                     |d|j        f¦  «        d         }t7          |j        j        |j        j        |j        j        |j        j        |j        j         ¬
¦  «        }| !                    tE          j#        |¦  «        |¦  «        }g }t          j        |¦  «        }|j$        D ]p}tK          tL          |j'        ¦  «         (                    ¦   «         |k    rŒ3|                     tS          t+          |j        |j*        |¦  «        |¬¦  «        ¦  «         ŒqtW          j,        |||¬¦  «        }||v r=||         d         �/tK          tZ          t\                   ||         d         ¦  «        ||<   |||<   �Œpt_          |||�ta          |¦  «        n|¬¦  «         tc          ||j        ¦  «        }|S )aø  
    Load a state_dict in conjunction with FSDP sharded optimizer state.

    This is the current recommended way to checkpoint FSDP.

    Examples::

    >>> # 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)
    >>> # Save
    >>> with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT):
    >>>     state_dict = {
    >>>         "optimizer": FSDP.optim_state_dict(model, optim),
    >>>         "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 = dist_cp.load_sharded_optimizer_state_dict(
    >>>         model_state_dict,
    >>>         optimizer_key="optimizer",
    >>>         storage_reader=dist_cp.FileSystemReader("checkpoint"),
    >>>     )
    >>>
    >>>     flattened_osd = FSDP.optim_state_dict_to_load(
    >>>        model, optim, optim_state["optimizer"]
    >>>     )
    >>>
    >>>     optim.load_state_dict(flattened_osd)
    Nr0   r1   r   r9   z
<bytes_io>rZ   )ÚrankÚ
world_sizeÚnum_devices_per_noder,   é   )rO   rP   rQ   Úmemory_formatrR   )rI   r^   )Úprocess_group)rX   r”   r•   )2Úread_metadatarg   r7   r<   r=   r>   r   r?   r@   r   r'   Úappendr	   rD   rv   r[   Úplanner_datar]   r   rA   ÚnumelrW   Ú
propertiesr   Úget_rankr   r‰   ÚShardTensorPropertiesrO   rP   rQ   r›   rR   Úbuild_metadatarT   rw   Úshards_metadatar   r   Ú	placementr—   r   r`   r   Ú+_init_from_local_shards_and_global_metadatar   r�   r   ri   r
   )r’   r“   r”   r•   r^   Úlayout_specsrc   Údp_pg_device_typer(   r;   ÚiÚdevice_infoÚsharding_specrX   rk   rd   re   Úkey_pathÚspec_keyÚ
alloc_sizer¡   Úst_mdrH   Úcurrent_rankÚshard_mdÚsts                             r)   r   r   Ú   sÚ  € ðp ×+Ò+Ñ-Ô-€Hå3Ð4DÑEÔEÑ€L�%ÝÔ-×DÒDÀUÑKÔKÔPÐÝ&Ð'8Ñ9Ô9€Mà€}Øˆ
Ý•tÔ*Ñ,Ô,Ñ-Ô-ð 	9ð 	9ˆAÝ0Ø! 1 }×'AÒ'AÑ'CÔ'CÑ#Cñô ˆKð ×ÒÐ7 aÐ7Ð7¨+Ð7Ð7Ñ8Ô8Ð8Ð8Ý)¨a¸JÐGÑGÔGˆˆå,¨UÑ3Ô3ˆð #%€Jà.0€MØÔ2×8Ò8Ñ:Ô:ð 8!ñ 8!‰
ˆˆUØÔ(¨Ô-ˆØ�AŒ;˜-Ò'Ð'Øå�eÕ1Ñ2Ô2ð 	Ø*ˆJ�s‰OØð Œ:×ÒÑÔ Ò"Ð"Ý+ØÔ  %¤*Ð.?ñô ˆJ�s‰OˆOð ˆ]Ý:Ý˜eÔ.°´
Ð<MÑNÔNÝ”]‘_”_ÝÔ.Ñ0Ô0Ø%2×%?Ò%?Ñ%AÔ%AÝ%Ñ'Ô'ðñ ô ˆJ�s‰OˆOð   ”{ˆHØ%×)Ò)¨(°T¸5¼:Ð4FÑGÔGÈÔJˆJå.ØÔ&Ô,ØÔ'Ô.Ø#Ô.Ô<Ø#Ô.Ô<Ø Ô+Ô6ðñ ô ˆJð "×0Ò0µ´¸JÑ1GÔ1GÈÑTÔTˆEØˆLÝœ=¨Ñ/Ô/ˆLØ!Ô1ð 
ð 
�Ý�¨Ô(:Ñ;Ô;×@Ò@ÑBÔBÀlÒRÐRØØ×#Ò#ÝÝ,Ø!Ô,¨hÔ.BÐDUñ ô  ð "*ð	ñ ô ñô ð ð õ ÔJØ˜e°5ðñ ô ˆBð ˜<Ð'Ð'¨L¸Ô,BÀ1Ô,EÐ,QÝ%)­(µ3¬-¸ÀhÔ9OÐPQÔ9RÑ%SÔ%S�˜cÑ"à ˆJ�s‰O‰Oõ ØØ%à49Ð4EÕ! -Ñ0Ô0Ð0È7ð	ñ ô ð õ & j°(Ô2GÑHÔH€JàÐr+   )r    rm   )Ery   Úcollections.abcr   Útypingr   rT   Útorch.distributedÚdistributedr7   Útorch._utilsr   Ú+torch.distributed._shard.sharded_tensor.apir   Ú0torch.distributed._shard.sharded_tensor.metadatar   r£   Ú-torch.distributed._shard.sharded_tensor.shardr   Ú:torch.distributed._shard.sharding_spec.chunk_sharding_specr	   Ú)torch.distributed.checkpoint._nested_dictr
   Ú,torch.distributed.checkpoint.default_plannerr   Ú%torch.distributed.checkpoint.metadatar   r   r   r   r   r   Ú$torch.distributed.checkpoint.plannerr   r   Ú,torch.distributed.checkpoint.planner_helpersr   r   Ú.torch.distributed.checkpoint.state_dict_loaderr   Ú$torch.distributed.checkpoint.storager   Ú"torch.distributed.checkpoint.utilsr   r   r   Ú"torch.distributed.distributed_c10dr   Ú#torch.distributed.fsdp._shard_utilsr   Útorch.distributed.remote_devicer   Útorch.distributed.tensorr   r�   rC   Útupler�   ÚSTATE_DICT_2D_LAYOUTÚ__all__r*   ÚProcessGrouprD   r�   ÚboolrL   rW   rg   ri   r   © r+   r)   ú<module>rÏ      sö  ðð Ð Ð Ð Ø $Ð $Ð $Ð $Ð $Ð $Ø Ð Ð Ð Ð Ð à €€€Ø  Ð  Ð  Ð  Ð  Ð  Ø +Ð +Ð +Ð +Ð +Ð +Ø EÐ EÐ EÐ EÐ EÐ Eðð ð ð ð ð ð @Ð ?Ð ?Ð ?Ð ?Ð ?Ø XÐ XÐ XÐ XÐ XÐ XØ JÐ JÐ JÐ JÐ JÐ JØ KÐ KÐ KÐ KÐ KÐ Kðð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð GÐ FÐ FÐ FÐ FÐ FÐ FÐ Fðð ð ð ð ð ð ð ð KÐ JÐ JÐ JÐ JÐ JØ >Ð >Ð >Ð >Ð >Ð >ðð ð ð ð ð ð ð ð ð ð
 BÐ AÐ AÐ AÐ AÐ AØ LÐ LÐ LÐ LÐ LÐ LØ :Ð :Ð :Ð :Ð :Ð :Ø ,Ð ,Ð ,Ð ,Ð ,Ð ,ð ˜C  x°¤}°tÑ';¸XÀc¼]Ð'JÔ!KÐKÔLÐ ð
 (ð€ð
ð  #ð °Cð ÀSð ð ð ð ð $(ðð ØÔ˜DÑ ðàðð ð ð ð(˜5œ<ð ¨Dð ð ð ð ð  FLðð ØðØ#+¨C¤=ðØ?Bðà
„\ðð ð ð ð( Øð à
Ð Ô!2°TÑ!9Ð9Ô:ð ð  ð  ð  ðF:Ið :Ið :Ið :Ið :IÐ*ñ :Iô :Ið :IðB #'ð	Qð QØ%ðQàðQð "ðQð ˜4Ñð	Qð
 ðQð Qð Qð Qð Qð Qr+   