Ë
    f^(hd5  ã                   óŠ  — d dl Z d dlmZmZmZ d dlZd dlmZ d dl	mc m
c mZ d dlmc mZ d dlmZmZ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 d d
l m!Z! d dl"m#Z#m$Z$m%Z%mZ& d dl'm(Z(m)Z) dgZ*de$de+ejX                  ejX                  f   fd„Z-de$de.de+ejX                  ejX                  f   fd„Z/de$de+ejX                  ejX                  f   fd„Z0de$de.defd„Z1de$dejd                  defd„Z3de$dejd                  fd„Z4dejj                  dejl                  de.dejj                  fd„Z7dejl                  de.de.de.dejd                  dejl                  fd „Z8dejl                  de.d!e#de$fd"„Z9dejl                  de+ejl                  e:e   f   fd#„Z;de$d$ee#   dejl                  fd%„Z< G d&„ de«      Z=y)'é    N)ÚAnyÚcastÚOptional)ÚShardÚShardedTensorÚShardedTensorMetadataÚTensorProperties)ÚShardMetadata)ÚChunkShardingSpec)Ú_mesh_resources)Ú_set_fsdp_flattened)ÚFSDPExtensions)Ú_create_chunk_sharded_tensor)Ú_remote_device)Ú
DeviceMeshÚDTensorÚ	Replicater   )Ú_flatten_tensorÚ_unflatten_tensorÚDTensorExtensionsÚtensorÚreturnc                 óÀ  — | j                   }|j                  dk(  sJ d«       ‚| j                  d   }dgt        | j	                  «       «      z  }|j	                  d¬«      }| j                  d   j                  «       r3t        t        |«      j                  }| j	                  |«      |z  }|||<   t        j                  |«      | j                  j	                  «       fS )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          úd/var/www/skyplay_api_hub/venv/lib/python3.12/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                 óx   — t        | «      \  }}t        j                  |D �cg c]  }||z  ‘Œ	 c}«      |fS c c}w ©N)r.   r%   r&   )r   r0   r)   r!   Úvals        r-   Ú_get_box_forr4   0   s6   € Ü˜VÓ$�M€GˆTÜ�J‰J¨WÖ5 c˜˜c›	Ò5Ó6¸Ð=Ð=ùÒ5s   ¢7c                 ó`   — | j                   }|j                  «       }|€J ‚t        | |d   «      S )Nr   )r   Úget_coordinater4   )r   r   Úcoords      r-   Ú_get_local_boxr8   5   s8   € Ø×$Ñ$€KØ×&Ñ&Ó(€EØÐÐÐÜ˜  a¡Ó)Ð)r/   ÚdtÚcurrent_rankc                 óÖ   — | j                   }|j                  dk(  sJ d«       ‚t        | «      \  }}t        t	        |«      t	        |«      d|› d| j
                  j                  › �¬«      S )Nr   r   úrank:ú/©Úshard_offsetsÚshard_sizesr(   )r   r   r8   r
   Úlistr'   Údevice)r9   r:   Úmeshr)   Úsizess        r-   Ú_create_shard_md_from_dtrE   <   sg   € Ø�>‰>€DØ�9‰9˜Š>ÐCÐCÓCˆ>ä# BÓ'�N€GˆUÜÜ˜7“mÜ˜“KØ˜,˜ q¨×)9Ñ)9×)@Ñ)@Ð(AÐBôð r/   Údt_pgc                 ó  — g }t        j                  |«      }|dkD  rdnd}| j                  d   j                  «       r|j	                  «       }nd}t        |«      D ]a  }t        | |«      \  }}|j                  t        t        |«      t        |«      d|dkD  r|n|› d| j                  j                  › �¬«      «       Œc t        || j	                  «       t        | j                  | j                  | j                   ¬«      ¬«      S )Nr   r   r<   r=   r>   )ÚdtypeÚlayoutÚrequires_grad)Úshards_metadatar!   Útensor_properties)ÚdistÚget_rankr   r"   r!   Úranger4   Úappendr
   rA   r'   rB   r   r	   rH   rI   rJ   )	r9   rF   Ú	shards_mdÚmy_rankÚscapegoat_rankÚshard_countÚir)   rD   s	            r-   Ú!_create_sharded_tensor_md_from_dtrV   H   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                 óf   — | j                   }|j                  dk(  sJ d«       ‚|j                  «       S )Nr   r   )r   r   Ú	get_group)r9   rC   s     r-   Ú
_get_dt_pgrY   o   s.   € Ø�>‰>€DØ�9‰9˜Š>ÐCÐCÓCˆ>Ø�>‰>ÓÐr/   ÚspecÚrankc                 ó  — t        | t        «      s| S d}| j                  D ]G  }t        t        |«      }|j                  «       |k(  sŒ'|j                  «       |j                  k7  sŒEd} n |rœt        j                  | «      } t        | j                  «      D ]o  \  }}t        t        |«      }|j                  «       |k(  sŒ*|j                  «       |j                  k7  sŒHt	        d|› d|j                  › �«      | j                  |<   Œq | S )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   r[   rB   ÚcopyÚdeepcopyÚ	enumerate)rZ   r   r[   ÚrewriteÚprU   r(   s          r-   Ú_rewrite_spec_if_neededrc   u   sì   € ô �dÔ-Ô.Øˆð €GØ�_‰_ò ˆÜ” Ó#ˆØ�6‰6‹8�tÓ §¡£
¨f¯m©mÓ ;ØˆGÙð	ñ
 Ü�}‰}˜TÓ"ˆÜ% d§o¡oÓ6ò 	T‰LˆAˆyÜœ^¨YÓ7ˆIØ�~‰~Ó 4Ó'¨I×,<Ñ,<Ó,>À&Ç-Á-Ó,OÜ%3°e¸D¸6ÀÀ6Ç=Á=À/Ð4RÓ%S�—‘ Ò"ð	Tð
 €Kr/   Ú
world_sizeÚnum_devices_per_nodeÚpgc           	      ó–  — t        | «      t        u rÓt        | j                  «       «      dk(  sJ ‚| j	                  «       }t        |||||«      }| j                  «       d   }t        |t        j                  |j                  «      «      g}t        j                  | j                  «       «      }	d|	j                  _        t        j                  ||	| j                  d¬«      }
|
S t        | «      t        u rÆ| j                  }|j                   dk(  sJ d«       ‚| j"                  }t        |||t$        j&                  j)                  «       |«      }t+        | «      }t        |t-        | t/        j0                  |«      «      «      g}t3        | |«      }	d|	j                  _        t        j                  ||	|d¬«      }
|
S t        | ||||«      S )Nr   r   F)Úsharded_tensor_metadataÚprocess_groupÚ
init_rrefsr   )Útyper   r    Úlocal_shardsÚlocal_tensorr   r   r^   r_   ÚmetadatarL   rJ   Ú+_init_from_local_shards_and_global_metadataÚ_process_groupr   r   r   r'   r%   ÚacceleratorÚdevice_countrY   rE   rM   rN   rV   )r   r[   rd   re   rf   Úinner_paramÚinner_stÚouter_local_shardÚshardsÚst_metaÚst_outerr   rF   s                r-   Ú_chunk_tensorry   ’   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¸EÓBˆØ27ˆ×!Ñ!Ô/ä ×LÑLØØ$+ØØô	
ˆð ˆä+ØØØØ Øó
ð 	
r/   r   c                 óÖ  — t        j                  |«      }|€t        d«      ‚|j                  dk  rt        d|j                  › d�d«      ‚| j	                  «       j                  «       } t        | t        j                  «      rœt        | t        «      sŒt        |j                  «      D �cg c]  }t        «       ‘Œ }}t        |j                  «      D �cg c]  }t        «       ‘Œ }}t        d«      |d<   t        j                  | ||d¬«      j                  ||¬	«      S | j                  }|d   }| j!                  «       } t        |j                  «      D �cg c]  }t        «       ‘Œ }}||d
<   t        |j                  «      D �	cg c]  }	t        «       ‘Œ }}	t        d«      |d<   ||d
<   t        j                  | ||d¬«      j                  ||¬	«      S c c}w c c}w c c}w c c}	w )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.
    z4No 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   éÿÿÿÿéþÿÿÿ)r   Úget_root_meshÚRuntimeErrorr   ÚdetachÚcloner]   r%   ÚTensorr   rO   r   r#   Ú
from_localÚredistributer   Úto_local)
r   r[   r   Ú	root_meshÚ_Úreplicate_placementsÚshard_placementsÚtp_placementsÚtp_placementrU   s
             r-   Ú_chunk_dtensorr�   Ü   sÏ  € ô  ×-Ñ-¨kÓ:€IØÐÜÐQÓRÐRØ‡~�~˜ÒÜØ/°	·±Ð/?¸qÐAØ-ó
ð 	
ð �]‰]‹_×"Ñ"Ó$€Fô
 �&œ%Ÿ,™,Ô'´
¸6Ä7Ô0Kô 6;¸9¿>¹>Ó5JÖK°¤	¥ÐKÐÐKÜ16°y·~±~Ó1FÖG¨AœI�KÐGÐÐGÜ$ Q›iÐ˜Ñä×!Ñ!Ø�IÐ3¸uô
ç
‰,Ø!Ø'ð ó 
ð	
ð ×)Ñ)ˆØ$ QÑ'ˆà—‘Ó"ˆô 6;¸9¿>¹>Ó5JÖK°¤	¥ÐKÐÐKØ#/Ð˜RÑ Ü16°y·~±~Ó1FÖG¨AœI�KÐGÐÐGÜ% a›yÐ˜ÑØ+Ð˜Ñä×!Ñ!Ø�IÐ3¸uô
ç
‰,Ø!Ø'ð ó 
ð	
ùò9  LùÚGùò*  LùâGs   Â+GÃGÅG!ÆG&c                 ó  — t        t        | «      j                  «       }t        |«      dk(  r?t	        |d   j
                  «      t        u r!|d   j
                  }|j                  «       }|} | t        |«      dkD  r|fS g fS )Nr   r   )r   r   rl   r    rk   r   )r   rv   Ú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                 ó"  — || j                   k(  sJ ‚t        t        j                  | j                  «      «      }t        dt        |«      dz
  «      D ]  }t        «       ||<   Œ | j                  | j                   |¬«      } | j                  «       S )zGAll gather a DTensor in its FSDP dimension and return the local tensor.r   r   r~   )
r   rA   r^   r_   r   rO   r    r   r‡   rˆ   )r   r“   r   rU   s       r-   Ú_all_gather_dtensorr•   )  s�   € ð
 ˜&×,Ñ,Ò,Ð,Ð,ä”d—m‘m F×$5Ñ$5Ó6Ó7€Jô �1”c˜*“o¨Ñ)Ó*ò $ˆÜ!›ˆ
�1Šð$à× Ñ Ø×&Ñ&Øð !ó €Fð
 �?‰?ÓÐr/   c                   óÜ  ‡ — e Zd ZdZdˆ fd„Zdej                  deej                  ee	   f   fd„Z
dej                  de	dej                  fd„Z	 ddej                  ded	ed
edej                  deej                     dej                  fd„Zdej                  dededej                  fd„Zdej                  deej                  ee   f   fd„Zdedee   dej                  fd„Zˆ xZS )r   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   c                 óš   •— t         ‰| �  «        d | _        || _        t        j
                  j                  | j                  «      | _        y r2   )ÚsuperÚ__init__Úcompute_streamÚdevice_handler%   Ú_dynamoÚdisableÚpost_unflatten_transform)Úselfr›   Ú	__class__s     €r-   r™   zDTensorExtensions.__init__E  s@   ø€ Ü‰ÑÔØ"ˆÔØ*ˆÔô ).¯©×(=Ñ(=Ø×)Ñ)ó)
ˆÕ%r/   r   c                 ó   — t        |«      S r2   )r   ©rŸ   r   s     r-   Úpre_flatten_transformz'DTensorExtensions.pre_flatten_transformO  s   € ô ˜vÓ&Ð&r/   Úparam_extensionc                 ó  — | j                   xs | j                  j                  «       }| j                  j                  |«      5  t	        ||| j                  | j                   ¬«      }t        |«       |cd d d «       S # 1 sw Y   y xY w)N)r›   rš   )rš   r›   Úcurrent_streamÚstreamr   r   )rŸ   r   r¤   r§   Úresults        r-   rž   z*DTensorExtensions.post_unflatten_transformU  s}   € ð ×$Ñ$ÒK¨×(:Ñ(:×(IÑ(IÓ(KˆØ×Ñ×&Ñ& vÓ.ñ 	ô 'ØØØ"×0Ñ0Ø#×2Ñ2ô	ˆFô   Ô'Ø÷	÷ 	ò 	ús   Á0A>Á>Br[   rd   re   rf   rB   c                 ó    — t        |||||«      S r2   )ry   )rŸ   r   r[   rd   re   rf   rB   s          r-   Úchunk_tensorzDTensorExtensions.chunk_tensorh  s   € ô ˜V T¨:Ð7KÈRÓPÐPr/   r   c                 ó   — t        |||«      S r2   )r�   )rŸ   r   r[   r   s       r-   Úchunk_dtensorzDTensorExtensions.chunk_dtensors  s   € ô ˜f d¨KÓ8Ð8r/   c                 ó   — t        |«      S r2   )r’   r¢   s     r-   Úpre_load_state_dict_transformz/DTensorExtensions.pre_load_state_dict_transform{  s   € ô $ FÓ+Ð+r/   r“   c                 ó   — t        ||«      S r2   )r•   )rŸ   r   r“   s      r-   Úall_gather_dtensorz$DTensorExtensions.all_gather_dtensor�  s   € ô
 # 6¨;Ó7Ð7r/   )r   Nr2   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r™   r%   r…   Útupler   r   r£   rž   ÚintrM   ÚProcessGrouprB   rª   r   r¬   rA   r   r®   r   r°   Ú__classcell__)r    s   @r-   r   r   =  sV  ø„ ñõ
ð'à—‘ð'ð 
ˆu�|‰|˜X c™]Ð*Ñ	+ó'ðØ—l‘lðØ58ðà	�‰óð4 *.ñ	Qà—‘ð	Qð ð	Qð ð		Qð
 "ð	Qð ×Ñð	Qð ˜Ÿ™Ñ&ð	Qð 
�‰ó	Qð9à—‘ð9ð ð9ð  ð	9ð
 
�‰ó9ð,à—‘ð,ð 
ˆu�|‰|˜T %™[Ð(Ñ	)ó,ð8àð8ð ˜jÑ)ð8ð 
�‰÷	8r/   )>r^   Útypingr   r   r   r%   Útorch.distributedÚdistributedrM   Ú&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.device_meshr   Ú$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¶   r4   r8   rE   r·   rV   rY   ÚShardingSpecr…   rc   ry   r�   rA   r’   r•   r   © r/   r-   ú<module>rÏ      s;  ðã ß &Ñ &ã Ý  ß ;Ó ;ß 1Ð 1÷ó õ AÝ XÝ 9Ý 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ðØðà˜*Ñ%ðð ‡\�\óô(I8˜õ I8r/   