ó
    EñiŒ5  ã                   ó�  • S SK r S SKJrJr  S SKrS SKJr  S SKJs  J	s  J
r  S SKJs  Jr  S SKJrJrJrJr  S SKJr  S SKJr  S SKJr  S SKJr  S SKJr  S S	KJr  S S
KJ r J!r!J"r"Jr#  S SK$J%r%J&r&  S/r'S\!S\(\RR                  \RR                  4   4S jr*S\!S\+S\(\RR                  \RR                  4   4S jr,S\!S\(\RR                  \RR                  4   4S jr-S\!S\+S\4S jr.S\!S\R^                  S\4S jr0S\!S\R^                  4S jr1S\Rd                  S\Rf                  S\+S\Rd                  4S jr4S\Rf                  S\+S\+S\+S\R^                  S\Rf                  4S jr5S\Rf                  S\+S \ S\!4S! jr6S\Rf                  S\(\Rf                  \7\   4   4S" jr8S\!S#\ S-  S\Rf                  4S$ jr9 " S% S\5      r:g)&é    N)ÚAnyÚcast)ÚShardÚShardedTensorÚShardedTensorMetadataÚTensorProperties)ÚShardMetadata)ÚChunkShardingSpec)Ú_set_fsdp_flattened)ÚFSDPExtensions)Ú_create_chunk_sharded_tensor)Ú_remote_device)Ú
DeviceMeshÚDTensorÚ	Replicater   )Ú_flatten_tensorÚ_unflatten_tensorÚDTensorExtensionsÚtensorÚreturnc                 óÆ  • U R                   nUR                  S:X  d   S5       eU R                  S   nS/[        U R	                  5       5      -  nUR	                  SS9nU R                  S   R                  5       (       a2  [        [        U5      R                  nU R	                  U5      U-  nXcU'   [        R                  " U5      U R                  R	                  5       4$ )Né   ú&Only 1D DeviceMeshes currently handledr   )Úmesh_dim)Údevice_meshÚndimÚ
placementsÚlenÚsizeÚis_shardr   ÚDShardÚdimÚtorchÚSizeÚ_local_tensor)r   r   Ú	placementÚoffsetsÚ
num_chunksÚ	shard_dimÚ
chunk_sizes          Úc/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributed/tensor/parallel/fsdp.pyÚ_get_boxr,      sË   € Ø×$Ñ$€KØ×Ñ˜qÓ ÐJÐ"JÓJÐ à×!Ñ! !Ñ$€IØˆc”C˜Ÿ™›Ó&Ñ&€GØ×!Ñ!¨1Ð!Ð-€Jà×Ñ˜Ñ×$Ñ$×&Ñ&Üœ Ó+×/Ñ/ˆ	Ø—[‘[ Ó+¨zÑ9ˆ
Ø'�	Ñä�JŠJ�wÓ ×!5Ñ!5×!:Ñ!:Ó!<Ð=Ð=ó    Úidxc                 ó|   • [        U 5      u  p#[        R                  " U Vs/ s H  oDU-  PM	     sn5      U4$ s  snf ©N)r,   r#   r$   )r   r.   r'   r   Úvals        r+   Ú_get_box_forr2   /   s6   € Ü˜VÓ$�M€GÜ�JŠJ©WÓ5ªW c˜cœ	©WÑ5Ó6¸Ð=Ð=ùÒ5s   ¢9c                 ó`   • U R                   nUR                  5       nUc   e[        XS   5      $ )Nr   )r   Úget_coordinater2   )r   r   Úcoords      r+   Ú_get_local_boxr6   4   s6   € Ø×$Ñ$€KØ×&Ñ&Ó(€EØÑÐÐÜ˜ a¡Ó)Ð)r-   ÚdtÚcurrent_rankc                 óÐ   • U R                   nUR                  S:X  d   S5       e[        U 5      u  p4[        [	        U5      [	        U5      SU SU R
                  R                   3S9$ )Nr   r   úrank:Ú/©Úshard_offsetsÚshard_sizesr&   )r   r   r6   r	   Úlistr%   Údevice)r7   r8   Úmeshr'   Úsizess        r+   Ú_create_shard_md_from_dtrC   ;   se   € Ø�>‰>€DØ�9‰9˜‹>ÐCÐCÓCˆ>ä# BÓ'�N€GÜÜ˜7“mÜ˜“KØ˜,˜ q¨×)9Ñ)9×)@Ñ)@Ð(AÐBñð r-   Údt_pgc                 ó
  • / n[         R                  " U5      nUS:”  a  SOSnU R                  S   R                  5       (       a  UR	                  5       nOSn[        U5       H^  n[        X5      u  pxUR                  [        [        U5      [        U5      SUS:”  a  UOU SU R                  R                   3S95        M`     [        UU R	                  5       [        U R                  U R                  U R                   S9S9$ )Nr   r   r:   r;   r<   )ÚdtypeÚlayoutÚrequires_grad)Úshards_metadatar   Útensor_properties)ÚdistÚget_rankr   r    r   Úranger2   Úappendr	   r?   r%   r@   r   r   rF   rG   rH   )	r7   rD   Ú	shards_mdÚmy_rankÚscapegoat_rankÚshard_countÚir'   rB   s	            r+   Ú!_create_sharded_tensor_md_from_dtrT   G   së   € ð €IÜ�mŠm˜EÓ"€GØ! A›+‘Q¨1€Nà	‡}�}�QÑ× Ñ ×"Ñ"Ø—j‘j“l‰àˆä�;ÖˆÜ% bÓ,‰ˆØ×ÑÜÜ" 7›mÜ  ›Kà¨a°!«e™N¸ÐAÀÀ2×CSÑCS×CZÑCZÐB[Ð\ñ	ö	
ñ  ô !Ø!Ø�W‰W‹YÜ*Ø—(‘(Ø—9‘9Ø×*Ñ*ñ
ñ	ð 	r-   c                 óh   • U R                   nUR                  S:X  d   S5       eUR                  5       $ )Nr   r   )r   r   Ú	get_group)r7   rA   s     r+   Ú
_get_dt_pgrW   n   s.   € Ø�>‰>€DØ�9‰9˜‹>ÐCÐCÓCˆ>Ø�>‰>ÓÐr-   ÚspecÚrankc                 ó@  • [        U [        5      (       d  U $ SnU R                   HK  n[        [        U5      nUR                  5       U:X  d  M)  UR                  5       UR                  :w  d  MI  Sn  O   U(       a¢  [        R                  " U 5      n [        U R                  5       Hs  u  pV[        [        U5      nUR                  5       U:X  d  M+  UR                  5       UR                  :w  d  MK  [	        SU SUR                   35      U R                  U'   Mu     U $ )zÈ
Rewrite ``spec`` to match the device of ``tensor``.

FSDP.sharded_optim_state_dict sneakly ships optimizer state to CPU so if the original ShardingSpec
produces CUDA metadata, ST construction bombs.
FTr:   r;   )
Ú
isinstancer
   r   r   r   rY   r@   ÚcopyÚdeepcopyÚ	enumerate)rX   r   rY   ÚrewriteÚprS   r&   s          r+   Ú_rewrite_spec_if_neededra   t   sá   € ô �dÔ-×.Ñ.Øˆð €GØ�_Œ_ˆÜ” Ó#ˆØ�6‰6‹8�tÕ §¡£
¨f¯m©mÕ ;ØˆGÙñ	 ö
 Ü�}Š}˜TÓ"ˆä% d§o¡oÖ6‰LˆAÜœ^¨YÓ7ˆIØ�~‰~Ó 4Õ'¨I×,<Ñ,<Ó,>À&Ç-Á-Õ,Oä%3°e¸D¸6ÀÀ6Ç=Á=À/Ð4RÓ%S�—‘ Ó"ñ	 7ð €Kr-   Ú
world_sizeÚnum_devices_per_nodeÚpgc           	      óš  • [        U 5      [        L aÔ  [        U R                  5       5      S:X  d   eU R	                  5       n[        UUUUU5      nU R                  5       S   n[        U[        R                  " UR                  5      5      /n[        R                  " U R                  5       5      n	SU	R                  l        [        R                  " UU	U R                  SS9n
U
$ [        U 5      [        L aÅ  U R                  nUR                   S:X  d   S5       eU R"                  n[        UUU[$        R&                  R)                  5       U5      n[+        U 5      n[        U[-        U [.        R0                  " U5      5      5      /n[3        X5      n	SU	R                  l        [        R                  " UU	USS9n
U
$ [        U UUUU5      $ )Nr   r   F)Úsharded_tensor_metadataÚprocess_groupÚ
init_rrefsr   )Útyper   r   Úlocal_shardsÚlocal_tensorr   r   r\   r]   ÚmetadatarJ   rH   Ú+_init_from_local_shards_and_global_metadataÚ_process_groupr   r   r   r%   r#   ÚacceleratorÚdevice_countrW   rC   rK   rL   rT   )r   rY   rb   rc   rd   Úinner_paramÚinner_stÚouter_local_shardÚshardsÚst_metaÚst_outerr   rD   s                r+   Ú_chunk_tensorrw   “   sÅ  € ô ˆFƒ|”}Ò$Ü�6×&Ñ&Ó(Ó)¨QÓ.Ð.Ð.à×)Ñ)Ó+ˆÜ/ØØØØ Øó
ˆð #×/Ñ/Ó1°!Ñ4Ðä�(œDŸMšMÐ*;×*DÑ*DÓEÓFð
ˆô —-’- §¡Ó 1Ó2ˆØ27ˆ×!Ñ!Ô/ä ×LÒLØØ$+Ø ×/Ñ/Øñ	
ˆð ˆÜ	ˆf‹œÒ	 Ø×(Ñ(ˆØ×Ñ 1Ó$ÐNÐ&NÓNÐ$à×*Ñ*ˆä/ØØØÜ×Ñ×*Ñ*Ó,Øó
ˆô ˜6Ó"ˆô �(Ô4°V¼T¿]º]È5Ó=QÓRÓSð
ˆô 4°FÓBˆØ27ˆ×!Ñ!Ô/ä ×LÒLØØ$+ØØñ	
ˆð ˆä+ØØØØ Øó
ð 	
r-   r   c                 óê  • Ub  UR                  5       OSnUc  [        S5      eUR                  S:  a  [        SUR                   S3S5      eU R                  5       R	                  5       n [        U [        R                  5      (       a¡  [        U [        5      (       dŒ  [        UR                  5       Vs/ s H  n[        5       PM     nn[        UR                  5       Vs/ s H  n[        5       PM     nn[        S5      US'   [        R                  " XUSS	9R                  UUS
9$ U R                  nUS   nU R                  5       n [        UR                  5       Vs/ s H  n[        5       PM     nnX…S'   [        UR                  5       V	s/ s H  n	[        5       PM     nn	[        S5      US'   X†S'   [        R                  " XUSS	9R                  UUS
9$ s  snf s  snf s  snf s  sn	f )z�
Shard a tensor to chunks along the first dimension.

The local rank will gets its corresponding chunk as the local tensor to create a DTensor.
Nz4No parent device_mesh is found for FSDP device_mesh.é   z!Found parent device_mesh of ndim=Ú,zbut meshes must be at least 2D.r   F)Ú	run_check©r   r   éÿÿÿÿéþÿÿÿ)Ú_get_root_meshÚRuntimeErrorr   ÚdetachÚcloner[   r#   ÚTensorr   rM   r   r!   Ú
from_localÚredistributer   Úto_local)
r   rY   r   Ú	root_meshÚ_Úreplicate_placementsÚshard_placementsÚtp_placementsÚtp_placementrS   s
             r+   Ú_chunk_dtensorr�   Ý   sÜ  € ð 1<Ñ0G�×*Ñ*Ô,ÈT€IØÑÜÐQÓRÐRØ‡~�~˜ÓÜØ/°	·±Ð/?¸qÐAØ-ó
ð 	
ð �]‰]‹_×"Ñ"Ó$€Fô
 �&œ%Ÿ,™,×'Ñ'´
¸6Ä7×0KÑ0Kô 6;¸9¿>¹>Ô5JÓKÒ5J°¤	¦Ñ5JÐÐKÜ16°y·~±~Ô1FÓGÒ1F¨AœIžKÑ1FÐÐGÜ$ Q›iÐ˜Ñä×!Ò!ØÐ3¸uñ
ç
‰,Ø!Ø'ð ð 
ð	
ð ×)Ñ)ˆØ$ QÑ'ˆà—‘Ó"ˆô 6;¸9¿>¹>Ô5JÓKÒ5J°¤	¦Ñ5JÐÐKØ#/˜RÑ Ü16°y·~±~Ô1FÓGÒ1F¨AœIžKÑ1FÐÐGÜ% a›yÐ˜ÑØ+˜Ñä×!Ò!ØÐ3¸uñ
ç
‰,Ø!Ø'ð ð 
ð	
ùò9  LùÚGùò*  LùâGs   Â7G!Ã$G&Å$G+ÆG0c                 ó  • [        [        U 5      R                  5       n[        U5      S:X  a@  [	        US   R
                  5      [        L a!  US   R
                  nUR                  5       nUn U [        U5      S:”  a  U4$ / 4$ )Nr   r   )r   r   rj   r   ri   r   )r   rt   Úinner_tensors      r+   Ú_pre_load_state_dictr�     sz   € ô ”- Ó(×5Ñ5Ó7€FÜ
ˆ6ƒ{�aÓœD ¨¡×!1Ñ!1Ó2´mÒCØ˜a‘y×'Ñ'ˆØ×*Ñ*Ó,ˆØˆàœc &›k¨A›o�FÐ6Ð6°2Ð6Ð6r-   Úparent_meshc                 ó  • XR                   :X  d   e[        [        R                  " U R                  5      5      n[        [        U5      S-
  5       H  n[        5       X#'   M     U R                  U R                   US9n U R                  5       $ )zGAll gather a DTensor in its FSDP dimension and return the local tensor.r   r|   )
r   r?   r\   r]   r   rM   r   r   r…   r†   )r   r‘   r   rS   s       r+   Ú_all_gather_dtensorr“   *  s‚   € ð
 ×,Ñ,Ó,Ð,Ð,ä”d—m’m F×$5Ñ$5Ó6Ó7€Jô ”3�z“? QÑ&Ö'ˆÜ!›ˆ
‹ñ (à× Ñ Ø×&Ñ&Øð !ð €Fð
 �?‰?ÓÐr-   c                   óö  ^ • \ rS rSrSrSU 4S jjrS\R                  S\\R                  \	S-  4   4S jr
S\R                  S\	S\R                  4S	 jr SS\R                  S
\S\S\S\R                  S\R                  S-  S\R                  4S jjrS\R                  S
\S\S\R                  4S jrS\R                  S\\R                  \\   4   4S jrS\S\S-  S\R                  4S jrSrU =r$ )r   i>  zÞ
DTensorExtension is the TensorFlattener extension needed for 2D FSDP + TP.

This is the implementation for FSDPExtensions defined in
https://github.com/pytorch/pytorch/blob/main/torch/distributed/fsdp/_fsdp_extensions.py
r   Nc                 ó˜   >• [         TU ]  5         S U l        Xl        [        R
                  R                  U R                  5      U l        g r0   )ÚsuperÚ__init__Úcompute_streamÚdevice_handler#   Ú_dynamoÚdisableÚpost_unflatten_transform)Úselfr™   Ú	__class__s     €r+   r—   ÚDTensorExtensions.__init__F  s>   ø€ Ü‰ÑÔØ"ˆÔØ*Ôô ).¯©×(=Ñ(=Ø×)Ñ)ó)
ˆÕ%r-   r   c                 ó   • [        U5      $ r0   )r   ©r�   r   s     r+   Úpre_flatten_transformÚ'DTensorExtensions.pre_flatten_transformP  s   € ô ˜vÓ&Ð&r-   Úparam_extensionc                 ó"  • U R                   =(       d    U R                  R                  5       nU R                  R                  U5         [	        UUU R                  U R                   S9n[        U5        UsS S S 5        $ ! , (       d  f       g = f)N)r™   r˜   )r˜   r™   Úcurrent_streamÚstreamr   r   )r�   r   r¤   r§   Úresults        r+   rœ   Ú*DTensorExtensions.post_unflatten_transformV  st   € ð ×$Ñ$×K¨×(:Ñ(:×(IÑ(IÓ(KˆØ×Ñ×&Ñ& vÕ.ô 'ØØØ"×0Ñ0Ø#×2Ñ2ñ	ˆFô   Ô'Ø÷ /×.×.ús   Á	-B Â 
BrY   rb   rc   rd   r@   c                 ó   • [        XX4U5      $ r0   )rw   )r�   r   rY   rb   rc   rd   r@   s          r+   Úchunk_tensorÚDTensorExtensions.chunk_tensori  s   € ô ˜V¨:ÈRÓPÐPr-   r   c                 ó   • [        XU5      $ r0   )r�   )r�   r   rY   r   s       r+   Úchunk_dtensorÚDTensorExtensions.chunk_dtensort  s   € ô ˜f¨KÓ8Ð8r-   c                 ó   • [        U5      $ r0   )r�   r¡   s     r+   Úpre_load_state_dict_transformÚ/DTensorExtensions.pre_load_state_dict_transform|  s   € ô $ FÓ+Ð+r-   r‘   c                 ó   • [        X5      $ r0   )r“   )r�   r   r‘   s      r+   Úall_gather_dtensorÚ$DTensorExtensions.all_gather_dtensor‚  s   € ô
 # 6Ó7Ð7r-   )r˜   r™   rœ   )r   Nr0   )Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r—   r#   rƒ   Útupler   r¢   rœ   ÚintrK   ÚProcessGroupr@   r«   r   r®   r?   r   r±   r   r´   Ú__static_attributes__Ú__classcell__)rž   s   @r+   r   r   >  s[  ø† ñ÷
ð'à—‘ð'ð 
ˆu�|‰|˜S 4™ZÐ'Ñ	(ô'ðØ—l‘lðØ58ðà	�‰ôð4 '+ñ	Qà—‘ð	Qð ð	Qð ð		Qð
 "ð	Qð ×Ñð	Qð —‘˜tÑ#ð	Qð 
�‰õ	Qð9à—‘ð9ð ð9ð  ð	9ð
 
�‰ô9ð,à—‘ð,ð 
ˆu�|‰|˜T %™[Ð(Ñ	)ô,ð8àð8ð   $Ñ&ð8ð 
�‰÷	8ò 8r-   );r\   Útypingr   r   r#   Útorch.distributedÚdistributedrK   Ú&torch.distributed._shard.sharding_specÚ_shardÚsharding_specÚ
shard_specÚ"torch.distributed.distributed_c10dÚdistributed_c10dÚc10dÚ'torch.distributed._shard.sharded_tensorr   r   r   r   r	   Ú:torch.distributed._shard.sharding_spec.chunk_sharding_specr
   Ú$torch.distributed.fsdp._common_utilsr   Ú'torch.distributed.fsdp._fsdp_extensionsr   Ú#torch.distributed.fsdp._shard_utilsr   Útorch.distributed.remote_devicer   Útorch.distributed.tensorr   r   r   r!   Ú6torch.distributed.tensor.parallel._data_parallel_utilsr   r   Ú__all__r»   r$   r,   r¼   r2   r6   rC   r½   rT   rW   ÚShardingSpecrƒ   ra   rw   r�   r?   r�   r“   r   © r-   r+   Ú<module>rÕ      s5  ðã ß ã Ý  ß ;Ó ;ß 1Ð 1÷ó õ AÝ XÝ DÝ BÝ LÝ :ß TÓ T÷ð Ð
€ð>�Wð >  u§z¡z°5·:±:Ð'=Ñ!>ô >ð >˜ð > sð >¨u°U·Z±ZÀÇÁÐ5KÑ/Lô >ð
*˜7ð * u¨U¯Z©Z¸¿¹Ð-CÑ'Dô *ð	 ð 	¸ð 	Àô 	ð$Øð$Ø×)Ñ)ð$àô$ðN�7ð ˜t×0Ñ0ô ðØ
×
!Ñ
!ðØ+0¯<©<ðØ?Bðà×Ñôð>G
Ø�L‰LðG
à
ðG
ð ðG
ð ð	G
ð
 	×ÑðG
ð ‡\�\ôG
ðT>
Ø�L‰Lð>
à
ð>
ð ð>
ð ô	>
ðB	7Ø�L‰Lð	7à
ˆ5�<‰<˜˜e™Ð$Ñ%ô	7ðØðà˜dÑ"ðð ‡\�\ôô(I8˜õ I8r-   