Ë
    g^(h�9  ã                   ó”  — d dl Z d dlmZ d dlmZmZ d dlZd dlmc m	Z
 d dlmc mc mZ d dlmZmZ d dlmZ d dlmZmZmZmZ  e j2                  e«      Z G d„ de«      Zd	ed
edee   fd„Zed	ed
edee   fd„«       Zdddœdej@                  dedede!de!dej@                  fd„Z" G d„ dejF                  jH                  «      Z%y)é    N)Úcache)ÚcastÚ
NamedTuple)ÚDTensorSpecÚ
TensorMeta)Ú
DeviceMesh)ÚPartialÚ	PlacementÚ	ReplicateÚShardc                   ó<   — e Zd ZU eed<   eeef   ed<   ee   ed<   y)Ú_TransformInfoÚmesh_dimÚsrc_dst_placementsÚlogical_shapeN)Ú__name__Ú
__module__Ú__qualname__ÚintÚ__annotations__Útupler
   Úlist© ó    úd/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributed/tensor/_redistribute.pyr   r      s!   … ØƒMØ˜i¨Ð2Ñ3Ó3à˜‘9Ôr   r   Úsrc_specÚdst_specÚreturnc           	      óv  — g }| j                   }|j                  «       }|€J ‚t        | j                  «      }|g}|j                  dk(  r;|j                  t        d| j                  d   |j                  d   f|¬«      «       |S t        | j                  «      D ]¢  \  }}||   }	t        |t        «      rw||j                  dz
  k  sŒ.|j                  |¬«      }
|j                  |	|j                     |
||   «      \  }}t        |	«      }|||j                  <   |j                  |«       Œ’|j                  |	«       Œ¤ t        | j                  «      }t        |j                  «      }| j                  dkD  rôt        t!        t#        |«      «      «      D ]Ô  }||   }||   }t        |t        «      r‹|j                  }g g }}t        t%        ||«      «      D ]T  \  }\  }}||k\  r nG|j'                  |«      r|j                  |«       |j'                  |«      sŒD|j                  |«       ŒV ||k7  r
t)        «       }||k7  sŒ®|j                  t        |||f||   ¬«      «       |||<   ŒÖ t        t%        ||«      «      D ]5  \  }\  }}||k7  sŒ|j                  t        |||f||   ¬«      «       |||<   Œ7 |S )a�  
    Generate the transform infos from the source placements to the target placements.

    To transform from source to target placement it might have multiple steps, i.e. it
    might decompose Si -> Sj into Si -> R -> Sj.
    This would detect if there're mis-aligned/nested shardings between src/dst placements.
    E.g. Suppose the redistribution to perform is (Shard(0), Shard(0)) -> (Replicate(), Shard(0)),
    in this case Shard(0) -> Shard(0) for mesh dimension 1 actually needs resharding, because in
    the former is a nested-sharding of a tensor already already sharded dimension 0, whereras
    the latter is the first sharding on tensor dimension 0.
    é   r   )r   r   r   ©r   )Údevice_meshÚget_coordinater   ÚshapeÚndimÚappendr   Ú
placementsÚ	enumerateÚ
isinstancer   ÚsizeÚ_local_shard_size_on_dimÚdimÚ
num_shardsÚreversedÚrangeÚlenÚzipÚis_shardr   )r   r   Útransform_infosr"   Úmy_coordinateÚinitial_logical_shapeÚmesh_dims_to_logical_shapeÚiÚsrcÚcurrent_logical_shapeÚmesh_dim_sizeÚlocal_shard_sizeÚ_Únew_logical_shapeÚcurrent_placementsÚtarget_placementsr   ÚcurrentÚtargetÚ	shard_dimÚcurrent_mesh_shardingÚtarget_mesh_shardingÚsÚps                           r   Ú_gen_transform_infos_non_cachedrG      s  € ð -/€Oà×&Ñ&€KØ×.Ñ.Ó0€MØÐ$Ð$Ð$ô ! §¡Ó0ÐØ"7Ð!8Ðà×Ñ˜1Òà×ÑÜØØ$,×$7Ñ$7¸Ñ$:¸H×<OÑ<OÐPQÑ<RÐ#SØ3ôô	
ð Ðô
 ˜H×/Ñ/Ó0ò E‰ˆˆ3Ø :¸1Ñ =ÐÜ�cœ5Ô!Ø�;×#Ñ# aÑ'Ó'à +× 0Ñ 0¸!Ð 0Ó <�Ø&)×&BÑ&BØ)¨#¯'©'Ñ2Ø!Ø! !Ñ$ó'Ñ#Ð  !ô
 %)Ð)>Ó$?Ð!Ø-=Ð! #§'¡'Ñ*Ø*×1Ñ1Ð2CÕDà&×-Ñ-Ð.CÕDðEô& ˜h×1Ñ1Ó2ÐÜ˜X×0Ñ0Ó1Ðà×Ñ˜QÒô
 !¤¤sÐ+=Ó'>Ó!?Ó@ò 	6ˆHØ(¨Ñ2ˆGØ& xÑ0ˆFô ˜&¤%Ô(à"ŸJ™J�	Ø>@À"Ð';Ð%Ü!*¬3Ð/AÐCTÓ+UÓ!Vò 7‘I�A‘v˜˜1Ø˜H’}ÙØ—z‘z )Ô,Ø-×4Ñ4°QÔ7Ø—z‘z )Õ,Ø,×3Ñ3°AÕ6ð7ð )Ð,@Ò@ô '›[�Fà˜&Ó Ø×&Ñ&Ü"Ø!)Ø,3°VÐ+<Ø&@ÀÑ&Jôôð 06Ð" 8Ò,ð=	6ôF (1ÜÐÐ 1Ó2ó(ò 2Ñ#ˆÑ#�7˜Fð �fÓØ×"Ñ"ÜØ%Ø(/°Ð'8Ø"<¸XÑ"Fôôð ,2Ð˜xÒ(ð2ð Ðr   c                 ó   — t        | |«      S ©N)rG   )r   r   s     r   Ú_gen_transform_infosrJ   ”   s   € ô
 +¨8°XÓ>Ð>r   F©Úasync_opÚis_backwardÚlocal_tensorÚcurrent_specÚtarget_specrL   rM   c                óž  — |j                   |j                   k7  rt        d«      ‚d}|j                   }|j                  «       }|€| S t        d„ |j                  D «       «      xs t        d„ |j                  D «       «      }|rt        ||«      }	nt        ||«      }	|	D �]v  }
|
j                  }|
j                  \  }}|j                  |¬«       ||k(  r| }Œ9t        j                  d|||«       |j                  «       r‡|j                  «       r%t        t        |«      }|j!                  | ||«      }�nÛ|j#                  «       r0t        t$        |«      }|j'                  | |||
j(                  «      }�n›t+        d|› d|› d	�«      ‚|j#                  «       rÜt        t$        |«      }|j                  «       r&t        t        |«      }|j-                  | |||«      }�n3|j                  «       r|j/                  | ||||   «      }�n
|j#                  «       s
J d
|› �«       ‚t        t$        |«      }|j0                  |j0                  k7  rÇ|j3                  | |||
j(                  |j0                  «      }n�|j                  «       r�|j                  «       r(t        t        |«      }|s|j5                  | ||«      n| }nU|j#                  «       rC|st+        d|› d|› d	�«      ‚t        t$        |«      }|j'                  | |||
j(                  «      }n| }|€J ‚|} �Œy |€J d«       ‚|s*t7        |t8        j:                  «      r|j=                  «       }|S )zÿ
    This redistribute the local tensor (torch.Tensor) from the current DTensorSpec to
    the target DTensorSpec, which involves the necessary collective calls to transform
    the local shard of the DTensor from its current spec to the target spec.
    z)Cross device mesh comm not supported yet!Nc              3   óP   K  — | ]  }t        |t        j                  «      –— Œ  y ­wrI   ©r)   ÚtorchÚSymInt©Ú.0rE   s     r   ú	<genexpr>z,redistribute_local_tensor.<locals>.<genexpr>¸   s   è ø€ ÒN°a”j ¤E§L¡L×1ÑNùó   ‚$&c              3   óP   K  — | ]  }t        |t        j                  «      –— Œ  y ­wrI   rS   rV   s     r   rX   z,redistribute_local_tensor.<locals>.<genexpr>¸   s"   è ø€ ò VØ()Œ
�1”e—l‘l×#ñVùrY   r!   z)redistribute from %s to %s on mesh dim %szredistribute from z to z not supported yetz,Current placement should be shard but found zredistribute failed!)ÚmeshÚNotImplementedErrorr#   Úanyr$   rG   rJ   r   r   r*   ÚloggerÚdebugÚis_replicateÚ
is_partialr   r	   Ú_reduce_valuer2   r   Ú_to_replicate_tensorr   ÚRuntimeErrorÚ_reduce_shard_valueÚ_replicate_to_shardr,   Ú_to_new_shard_dimÚ_partition_valuer)   ÚfuncolÚAsyncCollectiveTensorÚwait)rN   rO   rP   rL   rM   Únew_local_tensorr"   r4   Úhas_symintsr3   Útransform_infor7   r@   rA   Úpartial_specÚcurrent_placementÚtarget_placementÚ
shard_specs                     r   Úredistribute_local_tensorrs   œ   sŠ  € ð ×Ñ˜K×,Ñ,Ò,ä!Ð"MÓNÐNàÐØ×#Ñ#€Kà×.Ñ.Ó0€MàÐð ÐäÑN¸<×;MÑ;MÔNÓNò ÔRUñ VØ-8×->Ñ->ôVó S€Kñ Ü9¸,ÈÓT‰ä.¨|¸[ÓIˆà)ó T(ˆØ×#Ñ#ˆØ(×;Ñ;‰ˆ�Ø×Ñ !ÐÔ$à�fÒà+ÐØä�‰Ð@À'È6ÐSTÔUà×ÑÔ à×!Ñ!Ô#Ü#¤G¨WÓ5�Ø#/×#=Ñ#=Ø  +¨qó$Ò ð ×!Ñ!Ô#Ü$(¬°Ó$8Ð!Ø#4×#IÑ#IØ  +¨q°.×2NÑ2Nó$Ò ô #Ø(¨¨	°°f°XÐ=OÐPóð ð �_‰_Ôä#¤E¨6Ó2ÐØ×!Ñ!Ô#Ü#¤G¨WÓ5�Ø#/×#CÑ#CØ  +¨qÐ2Bó$Ò ð ×%Ñ%Ô'à#3×#GÑ#GØ  +¨q°-ÀÑ2Bó$Ò ð ×'Ñ'Ô)ð ØBÀ7À)ÐLóÐ)ô "¤%¨Ó1�
Ø—>‘>Ð%5×%9Ñ%9Ò9Ø'1×'CÑ'CØ$Ø#ØØ&×4Ñ4Ø(×,Ñ,ó(Ñ$ð ×ÑÔ Ø×#Ñ#Ô%Ü#¤G¨VÓ4�ñ 'ð !×1Ñ1°,ÀÈQÔOà%ñ !ð
 ×!Ñ!Ô#Ù"Ü&Ø,¨W¨I°T¸&¸ÐASÐTóð ô %)¬°Ó$8Ð!Ø#4×#IÑ#IØ  +¨q°.×2NÑ2Nó$Ñ ð
 $0Ð àÐ+Ð+Ð+Ø'ŠðiT(ðl Ð'Ð?Ð)?Ó?Ð'áœ
Ð#3´V×5QÑ5QÔRØ+×0Ñ0Ó2ÐàÐr   c            
       óN   — e Zd Ze	 d
dddedeedf   defd„«       Zedd„«       Z	y	)ÚRedistributeÚinputúdtensor.DTensorr"   r'   .rL   c                 ó0  — |j                   }|| _        || _        |j                  |k7  r>t	        |||j                   j
                  ¬«      }|j                  }t        ||||¬«      }n|j                  }|}t        j                  |||j                  ¬«      S )N©Útensor_meta)rL   ©Úrequires_grad)Ú_specrO   rL   r'   r   rz   Ú_local_tensorrs   ÚdtensorÚDTensorr|   )	Úctxrv   r"   r'   rL   rO   rP   rN   Úoutputs	            r   ÚforwardzRedistribute.forward  s—   € ð —{‘{ˆØ'ˆÔØˆŒà×"Ñ" jÒ0Ü%Ø˜Z°U·[±[×5LÑ5LôˆKð !×.Ñ.ˆLÜ.Ø˜l¨KÀ(ô‰Fð
 ×(Ñ(ˆFØ&ˆKä�‰ØØØ×-Ñ-ô
ð 	
r   c           	      ó  — | j                   }|j                  }| j                  }|j                  }t	        ||||d¬«      }g }|j
                  D ]=  }|j                  «       r|j                  t        «       «       Œ-|j                  |«       Œ? t        |j                  t        |«      t        |j                  |j                  «       |j                  ¬«      ¬«      }	t!        j"                  ||	|j$                  ¬«      }
|
d d d fS )NTrK   )r$   ÚstrideÚdtypery   r{   )rO   r}   rL   r~   rs   r'   ra   r&   r   r   r"   r   r   r$   r…   r†   r   r€   r|   )r�   Úgrad_outputÚprevious_specrO   rL   rN   r‚   Únormalized_placementsÚprevious_placementÚspecÚoutput_dtensors              r   ÚbackwardzRedistribute.backward@  s  € à×(Ñ(ˆØ"×(Ñ(ˆØ—<‘<ˆà"×0Ñ0ˆÜ*ØØØØØô
ˆð 24ÐØ"/×":Ñ":ò 	AÐØ!×,Ñ,Ô.à%×,Ñ,¬Y«[Õ9à%×,Ñ,Ð-?Õ@ð	Aô Ø×%Ñ%ÜÐ'Ó(Ü"Ø!×'Ñ'Ø"×)Ñ)Ó+Ø!×'Ñ'ôô
ˆô !Ÿ™ØØØ%×3Ñ3ô
ˆð ØØØð	
ð 	
r   N)F)r‡   rw   )
r   r   r   Ústaticmethodr   r   r
   Úboolrƒ   r�   r   r   r   ru   ru     s_   „ Øð ñ
ð !ð
ð  ð	
ð
 ˜) S˜.Ñ)ð
ð ò
ó ð
ð@ ò*
ó ñ*
r   ru   )&ÚloggingÚ	functoolsr   Útypingr   r   rT   Ú)torch.distributed._functional_collectivesÚdistributedÚ_functional_collectivesri   Útorch.distributed.tensor._apiÚtensorÚ_apir   Ú&torch.distributed.tensor._dtensor_specr   r   Ú$torch.distributed.tensor.device_meshr   Ú(torch.distributed.tensor.placement_typesr	   r
   r   r   Ú	getLoggerr   r^   r   r   rG   rJ   ÚTensorr�   rs   ÚautogradÚFunctionru   r   r   r   ú<module>r       s  ðó Ý ß #ã ß :Ð :ß /Ó /ß JÝ ;÷ó ð 
ˆ×	Ñ	˜8Ó	$€ô�Zô ðsØðsàðsð 
ˆ.Ñósðl ð?Øð?àð?ð 
ˆ.Ñò?ó ð?ð ØòØ—,‘,ðàðð ðð
 ðð ðð ‡\�\óôDM
�5—>‘>×*Ñ*õ M
r   