ó
    Eñi<t  ã                   óÌ  • S SK r S SKrS SKrS SKrS SKJrJrJr  S SKJ	r	J
r
JrJrJr  S SKrS SKJs  Jr  S SKJr  S SKJs  Jr  S SKJr  \R8                  " 5       (       d  \(       a  S SKJr  S SKJr  S SK J!r!J"r"J#r#  S SK$J%r%  S	\RL                  S
\RN                  S-  S\RP                  S-  S\	S\RL                  4
S jr)  SASSS
\RN                  S-  S\RP                  S-  S\RL                  4S jjr* " S S\+5      r,SSSSSSSS.S\	S\S\S\S
\RN                  S-  S\RP                  S-  S\-S\	S\.\/S4   S\-S \-S\0\1\	4   4S! jjr2SSSSSS".S#\0\1\	4   S
\RN                  S-  S\RP                  S-  S\-S\.\/S4   S\-S\0\1\	4   4S$ jjr3SSS%.S#\0\1\	4   S\.\/S4   S\-S\0\1\	4   4S& jjr4\Rj                  " 5         SBS#\0\1\	4   S'\0\1\	4   S \-S\-S\0\1\	4   4
S( jj5       r6\Rj                  " 5        SCS#\0\1\	4   S)\-S*\-S\0\1\	4   4S+ jj5       r7S#\0\1\	4   S,\0\1\	4   S\-4S- jr8 " S. S/\5      r9 SDS0\0\1\	4   S1\0\1\	4   S2\:\1   S\RP                  S
\RN                  S-  SS4S3 jjr; SDS1\0\1\	4   S2\:\1   S\RP                  S
\RN                  S-  SS4
S4 jjr<   SES0\0\1\	4   S1\0\1\	4   S\RP                  S
\RN                  S-  S5\-S\-SS4S6 jjr= SDS0\0\1\	4   S1\0\1\	4   S\RP                  S
\RN                  S-  SS4
S7 jjr>\\1\/4   r?\.\?S4   r@\0\1\@4   rA\0\1\	4   rB\\?\	4   rCS#\BS8\\@\	/S4   SS4S9 jrDS#\BS\.\B\A4   4S: jrES;\BS<\@S=\	SS4S> jrFS#\BS?\AS\B4S@ jrGg)Fé    N)ÚCallableÚMappingÚMutableMapping)ÚAnyÚcastÚ
NamedTupleÚTYPE_CHECKINGÚUnion)ÚAsyncCollectiveTensor)Údistributed_c10d)ÚShardedTensor)Údistribute_tensorÚDTensorÚ	Replicate)Ú%compute_local_shape_and_global_offsetÚobjÚpgÚdeviceÚcompanion_objÚreturnc                 ó   • U $ ©N© ©r   r   r   r   s       Ú`/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributed/_state_dict_utils.pyÚ_identity_funcr      s	   € ð €Jó    Úsharded_tensorr   c                 óx  • Uc  [         R                  " 5       n[        R                  " U5      nU R	                  5       nU R                  5       S   nU R                  5       R                  5       n[        R                  " XS-  5      U-  U-  nUc  [         R                  " U5      OUnU(       a„  US   R                  R                  5       n	U	R                  R                  UR                  :w  a  U	R                  U5      n	XyR                  5       -
  n
U
S:”  a  [        R                   " U	SU
/5      n	O["        R$                  " XpR&                  US9n	["        R(                  " Xs-  U	R&                  US9n[        R*                  " X¹US9  UR-                  SSU5      R/                  U R                  5       5      nU$ )Nr   )Údtyper   )Úgroup)r   Ú_get_default_groupÚdistÚget_world_sizeÚlocal_shardsÚsizeÚnumelÚmathÚceilÚ_get_pg_default_deviceÚtensorÚflattenr   ÚtypeÚtoÚFÚpadÚtorchÚzerosr    ÚemptyÚall_gather_into_tensorÚnarrowÚreshape)r   r   r   Ú
world_sizeÚshardsÚ
dim_0_sizeÚtensor_numelÚ
chunk_sizeÚ	pg_deviceÚlocal_tensorÚnum_paddingr+   s               r   Ú_all_gather_sharded_tensorr?       s}  € ð
 
�zÜ×0Ò0Ó2ˆÜ×$Ò$ RÓ(€JØ×(Ñ(Ó*€FØ×$Ñ$Ó& qÑ)€JØ!×&Ñ&Ó(×.Ñ.Ó0€LÜ—’˜:Ñ2Ó3°lÑBÀjÑP€Jà7=±~Ô×/Ò/°Ô3È6ð ö Ø˜a‘y×'Ñ'×/Ñ/Ó1ˆØ×Ñ×#Ñ# y§~¡~Ó5Ø'Ÿ?™?¨9Ó5ˆLØ ×#5Ñ#5Ó#7Ñ7ˆØ˜‹?ÜŸ5š5 °°;Ð/?Ó@ˆLøä—{’{Ø×2Ñ2¸9ñ
ˆô �[Š[ØÑØ× Ñ Øñ€Fô
 	×Ò ¸BÒ?à�]‰]˜1˜a Ó.×6Ñ6°~×7JÑ7JÓ7LÓM€FØ€Mr   c                   ó   • \ rS rSrSrg)ÚCompanionMismatchéF   r   N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__static_attributes__r   r   r   rA   rA   F   s   † Úr   rA   Fr   T©r   r   Úcpu_offloadr   Ú
ranks_onlyÚ
type_checkÚnon_blockingÚiter_objectÚsharded_tensor_funcÚdtensor_funcÚtensor_funcrI   rJ   .rK   rL   c                ó:  • [         R                  " S5      n[        U [        5      (       a  U" XXW5      nGOd[        U [        5      (       a  U" XXW5      nGOD[        U [         R
                  5      (       a  U" XXW5      nGO[        U [        [        [        [        [        R                  45      (       d  U c  U nGOß[        U [        5      (       aä  Ub£  [        U[        5      (       a4  [        UR                  5       5      [        U R                  5       5      :w  aZ  [        U[        5      (       a  SO7S[        UR                  5       5      < S[        U R                  5       5      < 3n[        U5      eU R!                  5        VVs0 s H   u  pïU[#        UUUUUUUUb  X~   OSUU	U
S9_M"     nnnOæ[        U [$        [&        45      (       a–  Ub9  [        U[$        [&        45      (       a  [)        U5      [)        U 5      :w  a  [        e[+        U 5       VVs/ s H!  u  nn[#        UUUUUUUUb  UU   OSUU	U
S9PM#     nnn[        U [&        5      (       a  ['        U5      nO5U	(       d  [,        R.                  " U 5      nO[1        S[3        U 5       35      eU(       a  [4        R6                  " U5      U;   Ga9  [        U[         R
                  5      (       Ga  U(       a  Uc  UR9                  U5      nUbù  [        U[        5      (       aE  [        U[        5      (       d  [;        S5      eUR<                  R?                  UR<                  U
S	9  O�[        U[        5      (       ay  [        U[        5      (       d  [;        S
5      e[+        URA                  5       5       H;  u  nnURB                  R?                  URA                  5       U   RB                  U
S	9  M=     OUR?                  XÊS	9  UnU$ [        U[        5      (       a  0 OSnU$ s  snnf s  snnf )a  Iterate through the state dict, applying the given functions to each tensor type.

Args:
    iter_object (Any): the target state_dict.
    sharded_tensor_func (Callable): the function to apply to ShardedTensor
    dtensor_func (Callable): the function to apply to DTensor
    tensor_func (Callable): the function to apply to Tensor
    pg (Optional[dist.ProcessGroup]): process group passed to tensor functions
    device (Optional[torch.device]): device passed to tensor functions
    cpu_offload (bool): whether to offload the tensors to CPU memory. This option is ignored
        if a companion_obj is supplied.
    companion_obj (Any): A companion object to the state dict. If this object
        is supplied, we attempt to copy the tensor to the companion object.
    ranks_only (Tuple[int, ...]): if this tuple is empty, all ranks will
        have the same state_dicts. Otherwise only ranks that in ``ranks_only``
        have the same state_dicts. Other ranks will get empty state_dicts.
    type_check (bool): check if the instance data type is a supported type
        that can be saved by DCP.  The current supported data types are
        torch.Tensor, DTensor, int, float, str, list, dict, None.
    non_blocking (bool): whether to use non-blocking copy when copying to the companion object.
ÚcpuNÚ zset(companion_obj.keys())=z set(iter_object.keys())=rH   zUnexpected value type z5ret must be a DTensor when companion_obj is a DTensor)rL   zAret must be a ShardedTensor when companion_obj is a ShardedTensor)"r1   r   Ú
isinstancer   r   ÚTensorÚintÚfloatÚstrÚbytesÚioÚBytesIOÚdictÚsetÚkeysrA   ÚitemsÚ_iterate_state_dictÚlistÚtupleÚlenÚ	enumerateÚcopyÚdeepcopyÚ
ValueErrorr-   r#   Úget_rankr.   ÚAssertionErrorÚ_local_tensorÚcopy_r%   r+   )rM   rN   rO   rP   r   r   rI   r   rJ   rK   rL   Ú
cpu_deviceÚretÚmsgÚkeyÚvalueÚidxÚvÚshards                      r   r`   r`   J   sÆ  € ôH —’˜eÓ$€JÜ�+œ}×-Ñ-Ù! +°6ÓIŠÜ	�K¤×	)Ñ	)Ù˜;¨FÓBŠÜ	�K¤§¡×	.Ñ	.Ù˜+¨6ÓAŠä�;¤¤e¬S´%¼¿¹Ð D×EÑEØÑàŠÜ	�K¤×	&Ñ	&ØÑ$Ü˜=¬$×/Ñ/Ü�=×%Ñ%Ó'Ó(¬C°×0@Ñ0@Ó0BÓ,CÓCô ˜m¬T×2Ñ2ñ à2œ˜M×.Ñ.Ó0Ó1Ñ3Ð3M´S¸×9IÑ9IÓ9KÓ5LÑ4NÐOð ô
 $ CÓ(Ð(ð  *×/Ñ/Ô1ô
ò 2‘
�ð Ô$ØØ#ØØØØØ'Ø4AÑ4M˜mÒ0ÐSWØ%Ø%Ø)ñò ñ 2ð 	ñ 
ˆô  
�K¤$¬ ×	/Ñ	/ØÑ$Ü˜=¬4´¨-×8Ñ8Ü�=Ó!¤S¨Ó%5Ó5ä#Ð#ô  $ KÔ0ô
ò 1‘��Qô  ØØ#ØØØØØ'Ø4AÑ4M˜m¨CÒ0ÐSWØ%Ø%Ø)ôñ 1ð 	ñ 
ô  �k¤5×)Ñ)Ü˜“*ˆCøÞÜ�mŠm˜KÓ(‰äÐ1´$°{Ó2CÐ1DÐEÓFÐFæœŸš rÓ*¨jÔ8Ü�cœ5Ÿ<™<×(Ò(Þ˜}Ñ4Ø—f‘f˜ZÓ(�àÑ(Ü˜m¬W×5Ñ5Ü% c¬7×3Ñ3Ü,ØSóð ð "×/Ñ/×5Ñ5Ø×)Ñ)¸ð 6ò ô   ¬}×=Ñ=Ü% c¬=×9Ñ9Ü,Ø_óð ô '0°×0JÑ0JÓ0LÖ&M™
˜˜UØŸ™×*Ñ*Ø×,Ñ,Ó.¨sÑ3×:Ñ:Èð +ó ò 'Nð
 "×'Ñ'¨Ð'ÑGØ#�ð
 €Jô ˜s¤D×)Ñ)‰b¨tˆð €JùóY
ùó.
s   Æ
'PÈ(P©r   r   rI   rJ   rK   Ú
state_dictc                ó8   • S nS n[        U UU[        UUUUUS9	$ )aÍ  
Given a state_dict, this API gathers all the ShardedTensors or DTensors in
the state_dict.


Args:
    state_dict (Dict[str, Any]): the target sharded state_dict.
    pg (Optional[dist.ProcessGroup]): the process group that is used to
        gather ShardedTensor. Note that gathering a DTensor will use
        the DeviceMesh. So this argument will be ignored when gathering a
        DTensor.
    device: (Optional[torch.device]): the device that is used to
        perform allgather for ShardedTensor. Note that gathering a DTensor
        will use the DeviceMesh. So this argument will be ignored when
        gathering a DTensor.
    cpu_offload (bool): whether to offload the tensors to CPU memory. The
        default value is False.
    ranks_only: (Tuple[int, ...]): if this tuple is empty, all ranks will
        have the same state_dicts. Otherwise only ranks that in ``ranks_only``
        have the same state_dicts. Other ranks will get empty state_dicts.
    type_check: (bool): check if the instance data type is a supported type
        that can be saved by DCP.  The current supported data types are
        torch.Tensor, DTensor, int, float, str, list, dict, None.

Returns:
    The gathered state dictionary.
c                 ó  • [         R                  " S5      n[        XU5      nU R                  5       (       a'  U R                  5       S   R                  R                  OUnUR                  U:w  a  UR                  U5      n U $ Un U $ )NrR   r   )r1   r   r?   r%   r+   r.   )rp   r   r   r   rl   Úoutput_tensorÚlocal_shard_devices          r   rN   Ú/_gather_state_dict.<locals>.sharded_tensor_funcú   s‹   € ô —\’\ %Ó(ˆ
Ü2°5¸fÓEˆð ×!Ñ!×#Ñ#ð ×ÑÓ  Ñ#×*Ñ*×1Ò1àð 	ð
 ×ÑÐ#5Ó5Ø!×$Ñ$Ð%7Ó8ˆEð ˆð "ˆEØˆr   c                 óˆ  • U R                   U R                  R                  :w  a%  U R                  U R                  R                  5      n U R                   Vs/ s H  n[        5       PM     nnU R                  U R                  US9n U R                  5       n [        U [        5      (       a  U R                  5       n U $ s  snf )N)Údevice_meshÚ
placements)r   r|   Údevice_typer.   r}   r   ÚredistributeÚto_localrT   r   Úwait)rp   r   r   r   Ú_r}   s         r   rO   Ú(_gather_state_dict.<locals>.dtensor_func  s¦   € Ø�<‰<˜5×,Ñ,×8Ñ8Ó8Ø—H‘H˜U×.Ñ.×:Ñ:Ó;ˆEð ,1×+;Ò+;Ó<Ò+; a”i–kÑ+;ˆ
Ð<Ø×"Ñ"Ø×)Ñ)Ø!ð #ð 
ˆð —‘Ó ˆÜ�eÔ2×3Ñ3Ø—J‘J“LˆEØˆùò =s   ÁB?rt   ©r`   r   )ru   r   r   rI   rJ   rK   rN   rO   s           r   Ú_gather_state_dictr…   Õ   s7   € òJò"ô* ØØØÜØØØØØñ
ð 
r   )rJ   rK   c                ó@   • [        U [        [        [        SSSUUS9	nU$ )aS  
Given a state_dict, this API offload all the tensors to CPU memory.

Args:
    state_dict (Dict[str, Any]): the target state_dict.
    pg (Optional[dist.ProcessGroup]): the process group that is used to
        gather ShardedTensor. Note that gathering a DTensor will use
        the DeviceMesh. So this argument will be ignored when gathering a
        DTensor.
    ranks_only: (Tuple[int, ...]): if this tuple is empty, all ranks will
        have the same state_dicts. Otherwise only ranks that in ``ranks_only``
        have the same state_dicts. Other ranks will get empty state_dicts.
    type_check: (bool): check if the instance data type is a supported type
        that can be saved by DCP.  The current supported data types are
        torch.Tensor, DTensor, int, float, str, list, dict, None.

Returns:
    The gathered state dictionary.
NTrt   r„   )ru   rJ   rK   rm   s       r   Ú_offload_state_dict_to_cpur‡   -  s0   € ô4 ØÜÜÜØØØØØñ
€Cð €Jr   Úcopy_state_dictc                 ó@   • [        U [        [        [        SSSSUUUS9$ )aO  
Copies all tensors in a given state dict into a different state_dict with the
same structure. Additionally, a copied state dict with the same value references
is returned. Editing the keys on this state dict will not affect the
passed in copy_state_dict (but the value references are the same).

.. warning::
    It is expected by this function that state_dict and copy_state_dict share
    the same structure and data types.

.. warning::
    The current supported data types are
        torch.Tensor, DTensor, int, float, str, list, dict, None.

Args:
    state_dict (Dict[str, Any]): the target state_dict.
    copy_state_dict (Dict[str, Any]):
        The state dict we are copying into. This state_dict must have exactly
         the same structure as the source `state_dict`.
    non_blocking: (bool): Whether copy ops should be performed asynchronously
    type_check (bool): check if the instance data type is a supported type
        that can be saved by DCP. The current supported data types are
        torch.Tensor, DTensor, int, float, str, list, dict, None.

Returns:
    State Dict copy
NFr   )r   r   rI   rJ   r   rK   rL   r„   )ru   rˆ   rL   rK   s       r   Ú_copy_state_dictrŠ   U  s3   € ôF ØÜÜÜØØØØØ%ØØ!ñð r   Ú
pin_memoryÚshare_memoryc                 óØ  ^^^• S[         R                  S[        R                  S-  S[         R                  S-  S[
        S[         R                  4
UU4S jjmS[        S[        R                  S-  S[         R                  S-  S[
        S[        4
U4S jjnS[        S[        R                  S-  S[         R                  S-  S[
        S[        4
U4S	 jjn[        U UUTSSS
SS
S9	nU$ )aÇ  
Given a state_dict, create another state_dict with the same structure and elements.
However, all tensors in the returned state_dict are new tensors on CPU. These
tensors can be placed on pin_memory or share_memory based on the provided arguments.

.. warning::
    Setting both `pin_memory` and `share_memory` to True significantly increases the
    latency of this method because of the nuances which require us to register memory
    as pinned directly as opposed to relying on the pin_memory cache allocator. This
    option should only be used for long lived tensors which are required to be shared.
    This is not the case as long as at least one of `pin_memory` or `share_memory` is
     set to False.

r   r   Nr   r‚   r   c                 óÌ  >• [        U R                  5       5      S:X  a  [        R                  " SU R                  S9$ U R                  5       S:X  d  U R                  5       S:X  a.  [        R                  " U SS9nT(       a  UR                  5       nU$ T(       aÈ  [        R                  " [        U R                  5       5      SU R                  06nUR                  5       nT(       ax  [        R                  " UR                  5       UR                  5       UR                  5       -  5        [        R                  " U[        R                   UR                  5       5        U$ T(       aE  [        R                  " [        U R                  5       5      SU R                  06R                  5       $ [        R                  " [        U R                  5       5      SU R                  06$ )Nr   )r    rR   ©r   r    )rc   r&   r1   r+   r    r'   Údata_ptrÚ
zeros_likeÚshare_memory_r3   rb   Úpin_memory_utilsr‹   Úelement_sizeÚweakrefÚfinalizeÚunpin_memory)r   r   r   r‚   Útr‹   rŒ   s        €€r   rP   Ú+_create_cpu_state_dict.<locals>.tensor_funcš  sA  ø€ ô ˆs�x‰x‹z‹?˜aÓÜ—<’< ¨¯©Ñ3Ð3ð
 �9‰9‹;˜!Ó˜sŸ|™|›~°Ó2Ü× Ò  ¨UÑ3ˆAÞØ—O‘OÓ%�ØˆHæÜ—’œU 3§8¡8£:Ó.Ð@°c·i±iÑ@ˆAØ—‘Ó!ˆAÞÜ ×+Ò+¨A¯J©J«L¸!¿'¹'»)ÀaÇnÁnÓFVÑ:VÔWÜ× Ò  Ô$4×$AÑ$AÀ1Ç:Á:Ã<ÔPàˆHÞÜ—;’;¤ c§h¡h£jÓ 1ÐC¸¿¹ÑC×NÑNÓPÐPä—;’;¤ c§h¡h£jÓ 1ÐC¸¿¹ÑCÐCr   c                 ó(  >• [        U R                  5       5      S:X  a  U $ U R                  [        R                  " S5      :w  a  [	        [
        U R                  SS95      nO[        R                  " U 5      nT" UR                  XS 5      Ul	        U$ )Nr   rR   r�   )
rc   r&   r   r1   r   r   r.   re   rf   rj   )r   r   r   r‚   rm   rP   s        €r   rO   Ú,_create_cpu_state_dict.<locals>.dtensor_func¹  sr   ø€ ô ˆs�x‰x‹z‹?˜aÓØˆJà�:‰:œŸš eÓ,Ó,Ü”w §¡¨e Ð 4Ó5‰Cä—-’- Ó$ˆCÙ'¨×(9Ñ(9¸2ÀtÓLˆÔØˆ
r   c                 ó*  >• U R                  5       (       d  U $ U R                  [        R                  " S5      :w  a  U R                  SS9nO[        R
                  " U 5      nUR                  5        H  nT" UR                  XS 5      Ul        M     U$ )NrR   r�   )r%   r   r1   r.   re   rf   r+   )r   r   r   r‚   rm   r8   rP   s         €r   rN   Ú3_create_cpu_state_dict.<locals>.sharded_tensor_funcÉ  sz   ø€ ð ×Ñ×!Ñ!ØˆJà�:‰:œŸš eÓ,Ó,Ø—&‘& �&Ð&‰Cä—-’- Ó$ˆCà×&Ñ&Ö(ˆFÙ'¨¯©°rÀ4ÓHˆFŽMñ )ð ˆ
r   Fr   rt   )	r1   rU   r#   ÚProcessGroupr   r   r   r   r`   )ru   r‹   rŒ   rO   rN   rm   rP   s    ``   @r   Ú_create_cpu_state_dictrŸ   ‡  s  ú€ ð&DÜ�\‰\ðDä×Ñ Ñ$ðDô —‘˜tÑ#ðDô ð	Dô
 
�‰÷Dð Dð>Üðä×Ñ Ñ$ðô —‘˜tÑ#ðô ð	ô
 
÷ð Üðä×Ñ Ñ$ðô —‘˜tÑ#ðô ð	ô
 
÷ô& ØØØØØØØØØñ
€Cð €Jr   Úcompared_state_dictc                 óü   • S[         R                  S[        R                  S-  S[         R                  S-  S[
        S[         R                  4
S jn [        U [        [        USSSS	USS
9
  g! [         a     gf = f)a  
Given two state_dicts, check if the structures are the same. And
if a [key, tensor] pair exist in one state_dict there must be
the a corresponding pait, [key, other_tensor], in the other state_dict,
where tensor and other_tensor have the same size and dtype.

Return the check result.
r   r   Nr   r   r   c                 óŠ   • UR                   U R                   :w  d"  UR                  5       U R                  5       :w  a  [        eU $ r   )r    r&   rA   r   s       r   rP   Ú1_check_state_dict_similarity.<locals>.tensor_func÷  s7   € ð ×Ñ #§)¡)Ó+¨}×/AÑ/AÓ/CÀsÇxÁxÃzÓ/QÜ#Ð#Øˆ
r   Fr   )r   r   rI   rJ   r   rK   T)	r1   rU   r#   rž   r   r   r`   r   rA   )ru   r    rP   s      r   Ú_check_state_dict_similarityr¤   ê  s•   € ðÜ�\‰\ðä×Ñ Ñ$ðô —‘˜tÑ#ðô ð	ô
 
�‰ôðÜØÜÜØØØØØØ-Øò	
ð øô ó Ùðús   ÁA. Á.
A;Á:A;c                   óR   • \ rS rSr% \R
                  \S'   \R                  \S'   Srg)Ú_TensorInfoi  r&   r    r   N)	rC   rD   rE   rF   r1   ÚSizeÚ__annotations__r    rG   r   r   r   r¦   r¦     s   ‡ Ø
�*‰*ÓØ�;‰;Ör   r¦   Úfull_state_dictÚlocal_state_dictr^   c                 óD  • Uc  [         R                  R                  5       nUR                  UR                   Vs1 s H  oUR                  iM     sn;   a  UOUR                  S   n/ nU HÛ  n[         R
                  " 5       S:X  aN  X   n[        U[        R                  5      (       d  [        S5      eUR                  5       R                  U5      n	O.X   n
[        R                  " U
R                  UU
R                  S9n	UR                  U	5        UR!                  U5      =nc  M¿  [        U["        5      (       a  X¹4OU	X'   MÝ     [%        U5      S:”  a  [         R&                  " XFSS5        O[         R(                  " US   SUS9  XS:w  a€  [+        X&5       Hq  u  pyUR!                  U5      =nc  M  [        U[,        5      (       a.  [        US   ["        5      (       a  US   U	R                  U5      4OU	R                  U5      X'   Ms     [/        XX45        g s  snf )Nr   z!full_state must be a torch.Tensor)r&   r   r    é   iô  ©Úsrcr!   )r#   r   r"   r-   Ú_device_typesrh   rT   r1   rU   ri   Údetachr.   r3   r&   r    ÚappendÚgetr   rc   Ú_broadcast_coalescedÚ	broadcastÚziprb   Ú_distribute_tensors)r©   rª   r^   r   r   r<   Útensorsro   Ú
full_stateÚfull_tensorÚtensor_infoÚlocal_states               r   Ú_broadcast_tensorsr¼     sð  € ð 
�zÜ×"Ñ"×5Ñ5Ó7ˆð �;‰;¸2×;KÒ;KÓLÒ;K¨iŸ>œ>Ñ;KÑLÓLñ 	à×Ñ˜aÑ ð ð #%€GÛˆÜ�=Š=‹?˜aÓØ(Ñ-ˆJÜ˜j¬%¯,©,×7Ñ7Ü$Ð%HÓIÐIØ$×+Ñ+Ó-×0Ñ0°Ó;‰Kà)Ñ.ˆKÜŸ+š+Ø ×%Ñ%Ø Ø!×'Ñ'ñˆKð 	�‰�{Ô#à+×/Ñ/°Ó4Ð4ˆKÑ=Ùô ˜+¤w×/Ñ/ð Ñ&àð 	Óñ' ô2 ˆ7ƒ|�aÓÜ×!Ò! "¨s°AÕ6ä�Š�w˜q‘z q°Ò3àÓÜ # DÖ 2ÑˆCØ/×3Ñ3°CÓ8Ð8�ÓEô # ;´×6Ñ6Ü& {°1¡~´w×?Ñ?ð ! ‘^ [§^¡^°FÓ%;Ñ<ð
 %Ÿ™¨Ó/ð !Ó%ñ !3ô Ð(°Õ;ùò_ Ms   »Hc           
      óð  • Uc  [         R                  R                  5       nU GHH  nU R                  U5      nUb  [        R
                  " U5      (       a  M5  US   nUS   n[        UR                  UR                  UR                  5      u  p‰[        X‰5       V
Vs/ s H  u  p«[        X»U
-   5      PM     nn
nUR                  (       ao  U[        U5         R                  5       R                  5       n[         R"                  " UUR                  UR                  UR                  UR%                  5       S9nO-UnUR'                  5       R)                  U[        U5         5        XàU'   GMK     g s  snn
f )Nr   r¬   )ÚshapeÚstride)r#   r   r"   r²   r1   Ú	is_tensorr   r¾   r|   r}   rµ   ÚsliceÚis_metarb   r°   Úcloner   Ú
from_localr¿   r€   rk   )rª   r^   r   r   ro   Ú_local_stater»   r¹   r¾   ÚoffsetÚ	cur_shapeÚ
cur_offsetÚslicesr=   rm   s                  r   r¶   r¶   V  sQ  € ð 
�zÜ×"Ñ"×5Ñ5Ó7ˆÜˆØ'×+Ñ+¨CÓ0ˆØÑ¤5§?¢?°<×#@Ñ#@Ùà" 1‘oˆØ" 1‘oˆä=Ø×Ñ˜{×6Ñ6¸×8NÑ8Nó
‰ˆô
 *-¨UÔ);ô
â);Ñ%�	ô �*¨9Ñ4Ö5Ù);ð 	ñ 
ð ××à&¤u¨V£}Ñ5×<Ñ<Ó>×DÑDÓFˆLô ×$Ò$ØØ×'Ñ'Ø×&Ñ&Ø!×'Ñ'Ø"×)Ñ)Ó+ñ‰Cð ˆCà�L‰L‹N× Ñ  ¬U°6«]Ñ!;Ô<Ø #˜Ôò? ùó
s   ÂE2Ústrictc                 ó4  • 0 n[         R                  " 5       S:X  aˆ  U R                  5        Ht  u  px[        R                  " U5      (       d  X†U'   M&  UR                  5       S:X  a  UR                  5       Xg'   MN  [        UR                  5       UR                  5      Xg'   Mv     U/n	[         R                  " U	SUS9  U	S   n/ n
[        UR                  5       5      n[        5       nUR                  5        H¸  u  pxUR                  U5        [        U[        5      (       d  Xq;   a  X�U'   M6  [         R                  " 5       S:X  a  X   Xg'   U
R                  U5        [!        U
5      S:¼  d  Mw  [#        XaX¢U5        U(       a  U
 H  nX   R                  5       X'   M     U
R%                  5         Mº     U(       a%  X¼-
  =n(       a  U H  nUR'                  U5        M     U
(       a3  [#        XaX¢U5        U(       a  U
 H  nX   R                  5       X'   M     g g g )Nr   r­   r¬   )r#   rh   r_   r1   rÀ   ÚdimrR   r¦   r&   r    Úbroadcast_object_listr]   r^   ÚaddrT   r±   rc   r¼   ÚclearÚpop)r©   rª   r   r   rÊ   rI   rm   ro   rp   Úbroadcast_listr^   Úlocal_state_dict_keysÚglobal_keysÚmissing_keyss                 r   Ú_broadcast_state_dictrÕ   €  sÅ  € ð €CÜ‡}‚}ƒ˜!ÓØ)×/Ñ/Ö1‰JˆCÜ—?’? 5×)Ñ)Ø �C“Ø—‘“ Ó!Ø Ÿ9™9›;�“ä& u§z¡z£|°U·[±[ÓA�“ñ 2ð �U€NÜ×Ò˜~°1¸BÒ?Ø
˜Ñ
€Cà€DÜÐ 0× 5Ñ 5Ó 7Ó8ÐÜ“%€KØ—i‘i–k‰
ˆØ�‰˜ÔÜ˜%¤×-Ñ-ØÓ&Ø(- Ñ%Ùä�=Š=‹?˜aÓØ&Ñ+ˆC‰Hà�‰�CÔäˆt‹9˜�>Ü˜s°dÀBÔGÞÛ�CØ,<Ñ,A×,EÑ,EÓ,GÐ$Ó)ñ  à�J‰JŽLñ# "ö& Ø1Ñ?Ð@ˆ<Õ@Û#�Ø ×$Ñ$ SÖ)ñ $ö Ü˜3°$ÀÔCÞÛ�Ø(8Ñ(=×(AÑ(AÓ(CÐ Ó%ò ð ð r   c                 óJ  • U R                  5        GH  u  pEX@;  a  M  [        R                  " U5      (       d  XQU'   M.  UR                  5       S:X  a  UR	                  5       X'   MV  [        U[        R                  5      (       d  [        S5      eUR                  U5      nUc  M–  [        U[        5      (       aB  [        UR                  5       R                  U5      UR                  UR                  5      X'   Mí  UR                  5       R                  U5      X'   GM     g )Nr   zvalue must be a torch.Tensor)r_   r1   rÀ   rÌ   rR   rT   rU   ri   r²   r   r   r°   r.   r|   r}   )r©   rª   r   r   ro   rp   r»   s          r   Ú_distribute_state_dictr×   »  sé   € ð &×+Ñ+×-‰
ˆØÓ%ÙÜ�Š˜u×%Ñ%Ø$)˜SÓ!Ø�Y‰Y‹[˜AÓØ$)§I¡I£KÐÓ!ä˜e¤U§\¡\×2Ñ2Ü$Ð%CÓDÐDØ*×.Ñ.¨sÓ3ˆKØÑ"ÙÜ˜K¬×1Ñ1Ü(9Ø—L‘L“N×%Ñ% fÓ-Ø×+Ñ+Ø×*Ñ*ó)Ð Ó%ð ).¯©«×(9Ñ(9¸&Ó(AÐ Ô%ò) .r   Úvisitorc                 óŽ   ^^• S[         S[        SS4UU4S jjmU R                  5        H  u  p#T" [        U5      4U5        M     g)zÃ
Invoke ``visitor`` for each value recursively in ``state_dict``.
Mapping, list, and tuple will be flattened and other value types are treated
as the terminal values and will invoke ``visitor``.
Úpathrp   r   Nc                 ó  >• [        U[        5      (       a0  UR                  5        H  u  p#T" U [        U5      4-   U5        M     g [        U[        [
        45      (       a!  [        U5       H  u  pCT" X4-   U5        M     g T" X5        g r   )rT   r   r_   rX   ra   rb   rd   )rÚ   rp   Úkrr   ÚiÚ_traverse_objrØ   s        €€r   rÞ   Ú+_traverse_state_dict.<locals>._traverse_objï  sq   ø€ Ü�eœW×%Ñ%ØŸ™ž‘�Ù˜d¤c¨!£f YÑ.°Ö2ò &ä˜¤¤e˜}×-Ñ-Ü! %Ö(‘�Ù˜d T™k¨1Ö-ò )ñ �DÕ r   )ÚOBJ_PATHr   r_   rX   )ru   rØ   ro   rp   rÞ   s    `  @r   Ú_traverse_state_dictrá   å  sI   ù€ ð!œHð !¬Sð !°T÷ !ð !ð !×&Ñ&Ö(‰
ˆÙ”s˜3“x�k 5Ö)ò )r   c                 óZ   ^^• 0 m0 mS[         S[        SS4UU4S jjn[        X5        TT4$ )ao  
Flatten ``state_dict`` made of nested dicts and lists into a top level dictionary.

Use ``unflatten_state_dict`` to revert this process.
Returns:
    A tuple with the flatten state_dict and a mapping from original to new state_dict.
N.B. The new keys are derived from the object paths, joined by dot.
    For example: ``{ 'a': {'b':...}}`` results in the key `a.b`.
rÚ   rp   r   Nc                 ó€   >• SR                  [        [        U 5      5      nUT;   a  [        SU 35      eUTU'   U TU'   g )NÚ.zduplicated flatten key )ÚjoinÚmaprX   rg   )rÚ   rp   Únew_fqnÚ	flattenedÚmappingss      €€r   Ú	flat_copyÚ&_flatten_state_dict.<locals>.flat_copy  sF   ø€ Ø—(‘(œ3œs D›>Ó*ˆØ�iÓÜÐ6°w°iÐ@ÓAÐAØ"ˆ	�'ÑØ ˆ�Òr   )rà   r   rá   )ru   rê   rè   ré   s     @@r   Ú_flatten_state_dictrì   ý  sC   ù€ ð "$€IØ "€Hð!œð !¬ð !°÷ !ð !ô ˜Ô/Ø�hÐÐr   Ú	root_dictrÚ   rp   c                 óÚ  • [        [        U 5      nS[        [           S[        SS4S jn[        S[        U5      5       Ho  nXS-
     nX   n[        U5      [        L a  0 O/ n[        U[        5      (       a!  [        [        UR                  Xh5      5      nMZ  U" X65        X6   c  XƒU'   X6   nMq     US   n[        U5      [        L a  U" [        [        [           U5      U5        X#U'   g)z>Set ``value`` in ``root_dict`` along the ``path`` object path.Úlstrq   r   Nc                 óh   • [        U 5      U::  a#  U R                  S 5        [        U 5      U::  a  M"  g g r   )rc   r±   )rï   rq   s     r   Úextend_listÚ!_set_element.<locals>.extend_list  s&   € Ü�#‹h˜#‹oØ�J‰J�tÔô �#‹h˜#�or   r¬   éÿÿÿÿ)r   ÚCONTAINER_TYPEra   r   rV   Úrangerc   r-   rX   rT   r   Ú
setdefault)	rí   rÚ   rp   Úcur_containerrñ   rÝ   Úprev_keyro   Údef_vals	            r   Ú_set_elementrú     så   € äœ¨Ó3€Mðœœc™ð ¬ð °ô ô �1”c˜$“iÖ ˆØ˜A™‘;ˆØ‰gˆÜ48¸³IÄÒ4D©bÈ"ˆä�m¤W×-Ñ-Ü Ü × 8Ñ 8¸Ó KóŠMñ
 ˜Ô0ØÑ&Ñ.Ø*1˜hÑ'Ø)Ñ3ŠMñ !ð  ˆr‰(€CÜˆCƒy”CÒÙ”Dœœc™ MÓ2°CÔ8à�#Òr   Úmappingc                 óZ   • 0 nU R                  5        H  u  p4[        X!U   U5        M     U$ )zaRestore the original nested state_dict according to ``mapping`` and the flattened ``state_dict``.)r_   rú   )ru   rû   Únestedro   rp   s        r   Ú_unflatten_state_dictrþ   6  s1   € ð !€FØ ×&Ñ&Ö(‰
ˆÜ�V S™\¨5Ö1ñ )à€Mr   )NN)FT)FFr   )NFF)Hre   rZ   r(   r•   Úcollections.abcr   r   r   Útypingr   r   r   r	   r
   r1   Útorch.cuda._pin_memory_utilsÚcudaÚ_pin_memory_utilsr“   Útorch.distributedÚdistributedr#   Útorch.nn.functionalÚnnÚ
functionalr/   Ú)torch.distributed._functional_collectivesr   Úis_availabler   Ú'torch.distributed._shard.sharded_tensorr   Útorch.distributed.tensorr   r   r   Útorch.distributed.tensor._utilsr   rU   rž   r   r   r?   Ú	ExceptionrA   Úboolrb   rV   r\   rX   r`   r…   r‡   Úno_gradrŠ   rŸ   r¤   r¦   ra   r¼   r¶   rÕ   r×   Ú	PATH_ITEMrà   ÚFLATTEN_MAPPINGÚSTATE_DICT_TYPErô   rá   rì   rú   rþ   r   r   r   Ú<module>r     sñ  ðã Û 	Û Û ß =Ñ =ß >Õ >ã ß 7Ð 7Ý  ß Ð Ý Kð ×Ò×Ñž-Ý2ÝEßNÑNÝUðØ	�‰ðà×Ñ˜DÑ ðð �L‰L˜4Ñðð ð	ð
 ‡\�\ôð $(Ø"&ñ#Ø#ð#à×Ñ˜DÑ ð#ð �L‰L˜4Ñð#ð ‡\�\õ	#ôL	˜	ô 	ð $(Ø"&ØØØ"$ØØòHØðHà!ðHð ðHð ð	Hð 	×Ñ˜DÑ ðHð �L‰L˜4ÑðHð ðHð ðHð �c˜3�h‘ðHð ðHð ðHð 
ˆ#ˆsˆ(�^õHð\ $(Ø"&ØØ"$ØòUØ�S˜#�X‘ðUð 	×Ñ˜DÑ ðUð �L‰L˜4Ñð	Uð
 ðUð �c˜3�h‘ðUð ðUð 
ˆ#ˆsˆ(�^õUðv #%Øò	%Ø�S˜#�X‘ð%ð �c˜3�h‘ð%ð ð	%ð
 
ˆ#ˆsˆ(�^õ%ðP ‡‚ƒð Øñ	.Ø�S˜#�X‘ð.à˜#˜s˜(‘^ð.ð ð.ð ð	.ð
 
ˆ#ˆsˆ(�^ô.ó ð.ðb ‡‚ƒàOTñ_Ø�S˜#�X‘ð_Ø,0ð_ØHLð_à	ˆ#ˆsˆ(�^ô_ó ð_ðD'Ø�S˜#�X‘ð'à˜c 3˜h™ð'ð 
ô'ôT�*ô ð $(ñ:<Ø˜#˜s˜(‘^ð:<à˜3 ˜8‘nð:<ð ˆs‰)ð:<ð �L‰Lð	:<ð
 	×Ñ˜DÑ ð:<ð 
õ:<ðB $(ñ	'$Ø˜3 ˜8‘nð'$à
ˆs‰)ð'$ð �L‰Lð'$ð 	×Ñ˜DÑ ð	'$ð
 
õ'$ð\ $(ØØñ8DØ˜#˜s˜(‘^ð8Dà˜3 ˜8‘nð8Dð �L‰Lð8Dð 	×Ñ˜DÑ ð	8Dð
 ð8Dð ð8Dð 
õ8Dð~ $(ñ	BØ˜#˜s˜(‘^ðBà˜3 ˜8‘nðBð �L‰LðBð 	×Ñ˜DÑ ð	Bð
 
õBðF �#�s�(‰O€	Ø�˜C�Ñ €Ø�s˜H�}Ñ%€Ø�s˜C�x‘.€Ø 	¨3 Ñ/€ð*Øð*à�x �o tÐ+Ñ,ð*ð 
ô*ð0Øðà
ˆ?˜OÐ+Ñ,ôð4˜Oð °8ð ÀCð ÈDô ð>ØðØ*9ðàõr   