ó
    EñiÜ  ã                   óV  • S SK r S SKrS SKrS SK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  S SKJrJrJrJr  S r SS\R.                  S	\S
\S\S\R2                  S\R4                  S-  S\4S jjrS\R.                  S	\S\S\4S jrS\S\S-  S\R.                  4S jrg)é    N)Ú_get_device_module)Údistributed_c10d)ÚShardÚShardedTensorÚShardedTensorMetadataÚTensorProperties)ÚShardMetadata)Ú
DeviceMeshÚDTensorÚ	Replicater   c                 óÀ   • UR                  5       S:X  a  SU  SU 3$ UR                  5       S:X  a"  SU  SU S[        U5      R                  5        3$ SU  SU SX-   3$ )NÚcpuzrank:Ú/ÚhpuÚ:)Úlowerr   Úcurrent_device)ÚrankÚdevice_typeÚnum_devices_per_nodes      Ú`/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributed/fsdp/_shard_utils.pyÚ_get_remote_device_strr      s}   € Ø×ÑÓ˜eÓ#Ø�t�f˜A˜k˜]Ð+Ð+Ø	×	Ñ	Ó	 Ó	%Ø�t�f˜A˜k˜]¨!Ô,>¸{Ó,K×,ZÑ,ZÓ,\Ð+]Ð^Ð^à�t�f˜A˜k˜]¨!¨DÑ,GÐ+HÐIÐIó    Útensorr   Ú
world_sizer   ÚpgÚdeviceÚreturnc                 ój  • U R                  USS9n[        U5      U:”  a{  Xa   R                  5       nU R                  5        Vs/ s H  nSPM     n	n[        R
                  " U R                  5       S   U-  5      U-  U	S'   [        R                  " XyU5      /n
O/ n
U Vs/ s H  n[        UR                  5       5      PM     nnS/[        [        R                  " U Vs/ s H  oÝS   PM	     sn5      5      SS -   nS/[        US   5      S-
  -  n	U Vs/ s H  oÿ/U	-   PM
     nnUc   [        R                  " U5      R                  OUR                  n[        [        U5      5       Vs/ s H%  n[        [         R"                  " UU5      UU5      PM'     nn[        U5      [        U5      :w  d  [        U5      [        U5      :w  a/  [%        S[        U5       S[        U5       S[        U5       35      e['        UUU5       VVVs/ s H  u  nnn[)        UUU5      PM     nnnn[+        UU R                  5       [-        U R.                  U R0                  S[2        R4                  U R7                  5       S	9S
9n[8        R:                  " U
UUS9$ s  snf s  snf s  snf s  snf s  snf s  snnnf )z”
Shard a tensor to chunks along the first dimension. The local rank will gets its
corresponding chunk as the local shard to create a ShardedTensor.
r   )ÚdimNéÿÿÿÿé   zQExpected chunk_sizes, chunk_offsets, and placements to have the same length, got z, F)ÚdtypeÚlayoutÚrequires_gradÚmemory_formatÚ
pin_memory)Úshards_metadataÚsizeÚtensor_properties)Úsharded_tensor_metadataÚprocess_group)ÚchunkÚlenÚcloner)   ÚmathÚceilr   Úfrom_tensor_and_offsetsÚlistÚ	itertoolsÚ
accumulater   Ú_get_pg_default_deviceÚtypeÚranger   ÚdistÚget_global_rankÚAssertionErrorÚzipr	   r   r   r#   r$   ÚtorchÚcontiguous_formatÚ	is_pinnedr   Ú+_init_from_local_shards_and_global_metadata)r   r   r   r   r   r   ÚchunksÚlocal_shardÚ_ÚoffsetsÚlocal_shardsr-   Úchunk_sizesÚ
chunk_sizeÚdim0_offsetsÚd0Úchunk_offsetsr   ÚrÚ
placementsÚoffsetr)   Ú	placementÚshard_metadatar+   s                            r   Ú_create_chunk_sharded_tensorrP      s­  € ð �\‰\˜*¨!ˆ\Ð,€FÜ
ˆ6ƒ{�TÓØ‘l×(Ñ(Ó*ˆØ$Ÿk™kœmÓ,šm˜“1™mˆÐ,Ü—Y’Y˜vŸ{™{›}¨QÑ/°*Ñ<Ó=ÀÑDˆ�‰
Ü×5Ò5°kÈDÓQÐR‰àˆñ 4:Ó:²6¨%”4˜Ÿ
™
›Ö%±6€KÐ:Ø�3œÜ×Ò¹kÓJºk°
¨œm¹kÑJÓKóà	€rðñ €Lð ˆc”S˜ Q™Ó(¨1Ñ,Ñ-€GÙ.:Ó;ªl¨�T˜G”^©l€MÐ;ð ‰>ô 	×/Ò/°Ó3×8Ò8à�[‰[ð ô ”s˜;Ó'Ô(óò )ˆAô 	Ü× Ò   QÓ'ØØ ö	
ñ
 )ð ð ô ˆ;Óœ3˜}Ó-Ó-´°[Ó1AÄSÈÃ_Ó1TÜðÜ�{Ó#Ð$ B¤s¨=Ó'9Ð&:¸"¼SÀ»_Ð<MðOó
ð 	
ô (+¨=¸+ÀzÔ'Rõâ'RÑ#ˆF�D˜)ô 	�f˜d IÖ.Ù'Rð ò ô 4Ø&Ø�[‰[‹]Ü*Ø—,‘,Ø—=‘=ØÜ×1Ñ1Ø×'Ñ'Ó)ñ
ñ
Ðô ×DÒDØÐ.EÐUWñð ùò] -ùò ;ùâJùò <ùòùôs$   ÁJÂ!#JÃ!JÄJ$Å.,J)ÈJ.Údevice_meshc                 óh  • U R                  5       R                  5       n [        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S9$ 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.
r   r!   F)Ú	run_check)rL   )	Údetachr/   r8   Úndimr   ÚDShardr   Ú
from_localÚredistribute)r   r   rQ   rC   Úreplicate_placementsÚshard_placementss         r   Ú_create_chunk_dtensorr[   _   s§   € ð �]‰]‹_×"Ñ"Ó$€Fô 27°{×7GÑ7GÔ1HÓIÒ1H¨AœIžKÑ1HÐÐIÜ-2°;×3CÑ3CÔ-DÓEÒ-D¨œ	žÑ-DÐÐEÜ! !›9Ð�RÑä×ÒØÐ1¸Uñç�lØ#ð ð ðùò	 JùÚEs   ¶B*Á#B/Ú	root_meshc                 óö   • XR                   :w  a  [        S5      e[        [        R                  " U R
                  5      5      n[        5       US'   U R                  U R                   US9n U R                  5       $ )zL
All gather a DTensor in its sharded dimension and return the local tensor.
z2The device mesh of a tensor should be a root mesh.r!   )rQ   rL   )	rQ   r;   r3   ÚcopyÚdeepcopyrL   r   rX   Úto_local)r   r\   rL   s      r   Ú_all_gather_dtensorra   x   sr   € ð ×&Ñ&Ó&ÜÐQÓRÐRä”d—m’m F×$5Ñ$5Ó6Ó7€Jô “[€Jˆr�NØ× Ñ Ø×&Ñ&Øð !ð €Fð
 �?‰?ÓÐr   )N)r^   r4   r0   r=   Útorch.distributedÚdistributedr9   Útorch._utilsr   r   Ú'torch.distributed._shard.sharded_tensorr   r   r   r   Ú&torch.distributed._shard.sharding_specr	   Útorch.distributed.tensorr
   r   r   rV   r   ÚTensorÚintÚProcessGroupr   rP   r[   ra   © r   r   Ú<module>rl      së   ðã Û Û ã Ý  Ý +Ý .÷ó õ Aß TÓ TòJð #'ñ?Ø�L‰Lð?à
ð?ð ð?ð ð	?ð
 	×Ñð?ð �L‰L˜4Ñð?ð õ?ðDØ�L‰Lðà
ðð ðð ô	ð2Øðà˜DÑ ðð ‡\�\õr   