§
    ŠŠtj™  ã                  ó
  — U d dl mZ d dlZd dlZd dlmZmZ d dlmZ d dl	m
Z
mZmZmZmZmZ d dlZd dlmZ d dlmZ d dlmZ d d	lmZmZ erd d
lmZ d dlmZ  ej        e¦  «        Z G d„ de¦  «        Z e!e!e"df         edz  f         Z#de$d<    G d„ de%¦  «        Z& edd¬¦  «         G d„ d¦  «        ¦   «         Z' edd¬¦  «         G d„ de'¦  «        ¦   «         Z(e'e(z  Z)de$d<    ed¬¦  «         G d„ d¦  «        ¦   «         Z* edd¬¦  «         G d„ d¦  «        ¦   «         Z+ edd¬¦  «         G d„ d ¦  «        ¦   «         Z,djd&„Z-dkd*„Z. G d+„ d,¦  «        Z/ G d-„ d.e¦  «        Z0d/d0œdld3„Z1d4„ Z2	 dmdnd<„Z3	 dmdod>„Z4 ed¬¦  «         G d?„ d@¦  «        ¦   «         Z5dpdB„Z6eddCœdqdI„¦   «         Z7edrdM„¦   «         Z7d/dCœdsdO„Z7dtdudP„Z8e	 dvdwdS„¦   «         Z9e	 dvdxdU„¦   «         Z9	 dtdydW„Z9d/d/dXœdzd`„Z:dd/dXœd{db„Z;d|di„Z<dS )}é    )ÚannotationsN)Ú	dataclassÚfield)ÚEnum)ÚcastÚLiteralÚoverloadÚProtocolÚTYPE_CHECKINGÚ	TypeAlias)Úfx)Ú_MeshLayout)ÚDTensor©Útree_flattenÚtree_unflatten)Ú
DeviceMesh)Ú	Placementc                  ó   — e Zd ZdZd
d„Zd	S )ÚGetMeshCallbackzGCallback to create/retrieve a DeviceMesh from its cache key components.Úmesh_dim_namesútuple[str, ...]Úmesh_layoutú_MeshLayout | NoneÚreturnr   c                ó   — d S ©N© )Úselfr   r   s      úa/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributed/pipelining/_utils.pyÚ__call__zGetMeshCallback.__call__   s	   € ð �Só    N)r   r   r   r   r   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r!   r   r"   r    r   r      s.   € € € € € ØQÐQðð ð ð ð ð r"   r   .r   ÚMeshCacheKeyc                  ó   — e Zd ZdZdS )ÚPipeliningMetadataErrorz<Raised on metadata mismatches during pipeline communication.N)r#   r$   r%   r&   r   r"   r    r)   r)   ,   s   € € € € € ØFÐFÐFÐFr"   r)   T)ÚfrozenÚslotsc                  ód   — e Zd ZU dZded<   ded<   ded<   ded	<   edd„¦   «         Zdd„Zdd„ZdS )Ú_TensorMetazðTensor metadata for recv buffer allocation and validation.

    For plain tensors, these are the tensor's actual attributes.
    For DTensors, these are LOCAL shard attributes; global attributes
    are stored in :class:`_DTensorMeta`.
    ú
torch.SizeÚshapeútuple[int, ...]Ústrideztorch.dtypeÚdtypeÚboolÚrequires_gradÚtensorútorch.Tensorr   c                ó²   — t          | t          ¦  «        rt          d¦  «        ‚t          | j        |                      ¦   «         | j        | j        ¬¦  «        S )a  Create metadata from a plain tensor.

        Args:
            tensor: A plain ``torch.Tensor`` (not DTensor).

        Returns:
            Metadata capturing shape, stride, dtype, and requires_grad.

        Raises:
            TypeError: If ``tensor`` is a DTensor.
        zJExpected plain tensor, got DTensor. Use _DTensorMeta.from_dtensor instead.©r/   r1   r2   r4   )Ú
isinstancer   r)   r-   r/   r1   r2   r4   ©r5   s    r    Úfrom_tensorz_TensorMeta.from_tensor>   s_   € õ �f�gÑ&Ô&ð 	Ý)Ø\ñô ð õ Ø”,Ø—=’=‘?”?Ø”,Ø Ô.ð	
ñ 
ô 
ð 	
r"   Údeviceútorch.device | strc                ó‚   — t          | |¦  «        }|                     ¦   «         r|                     | j        ¦  «         |S )zÅReconstruct a tensor on ``device`` from this metadata.

        Args:
            device: Target device for the tensor.

        Returns:
            An empty strided tensor on ``device``.
        )Ú_make_tensor_from_metaÚis_floating_pointÚrequires_grad_r4   )r   r<   Úts      r    Ú	to_tensorz_TensorMeta.to_tensorV   sC   € õ # 4¨Ñ0Ô0ˆØ×ÒÑ Ô ð 	1Ø×Ò˜TÔ/Ñ0Ô0Ð0Øˆr"   Úotherú	list[str]c                óX  — | |k    rg S g }| j         |j         k    r%|                     d| j         › d|j         › �¦  «         | j        |j        k    r%|                     d| j        › d|j        › �¦  «         | j        |j        k    r%|                     d| j        › d|j        › �¦  «         |S )zÓReturn field-by-field differences with ``other``.

        Args:
            other: Metadata to compare against.

        Returns:
            List of human-readable difference strings (empty if equal).
        zshape mismatch: ú vs zstride mismatch: zdtype mismatch: )r/   Úappendr1   r2   ©r   rD   Údiffss      r    Úget_diffz_TensorMeta.get_diffd   sÀ   € ð �5Š=ˆ=ØˆIàˆØŒ:˜œÒ$Ð$Ø�LŠLÐI¨D¬JÐIÐI¸E¼KÐIÐIÑJÔJÐJØŒ;˜%œ,Ò&Ð&Ø�LŠLÐL¨T¬[ÐLÐL¸e¼lÐLÐLÑMÔMÐMØŒ:˜œÒ$Ð$Ø�LŠLÐI¨D¬JÐIÐI¸E¼KÐIÐIÑJÔJÐJð ˆr"   N)r5   r6   r   r-   )r<   r=   r   r6   ©rD   r-   r   rE   )	r#   r$   r%   r&   Ú__annotations__Ústaticmethodr;   rC   rK   r   r"   r    r-   r-   0   s˜   € € € € € € ðð ð ÐÐÑØÐÐÑØÐÐÑØÐÐÑàð
ð 
ð 
ñ „\ð
ð.ð ð ð ðð ð ð ð ð r"   r-   c                  ó   — e Zd ZU dZ ed„ ¬¦  «        Zded<    ed¬¦  «        Zded	<    ed¬¦  «        Zd
ed<    ed¬¦  «        Z	ded<    ed¬¦  «        Z
ded<   ed d„¦   «         Zed!d„¦   «         Zd"d„Zd#d„ZdS )$Ú_DTensorMetaaÔ  DTensor metadata extending :class:`_TensorMeta` with distribution info.

    Inherited fields (shape, stride, etc.) are LOCAL shard attributes.
    Additional fields capture global shape and placement information
    needed to reconstruct a :class:`DTensor` via ``DTensor.from_local()``.

    The :class:`DeviceMesh` is **not** stored (not serializable for P2P);
    it is looked up from :class:`_MeshCache` using
    ``(mesh_dim_names, mesh_layout)`` as the key.
    c                 ó*   — t          j        g ¦  «        S r   )ÚtorchÚSizer   r"   r    ú<lambda>z_DTensorMeta.<lambda>Š   s   € ½U¼ZÈ¹^¼^€ r"   )Údefault_factoryr.   Úglobal_shaper   )Údefaultr0   Úglobal_strideztuple[Placement, ...]Ú
placementsr   r   Nr   r   Údtensorr   r   c                ó  — | j         }t          | j        j        | j                             ¦   «         | j        | j        | j        |                      ¦   «         | j        j        |j	        rt          |j	        ¦  «        nd|j        ¬¦	  «	        S )zÅCreate metadata from a DTensor.

        Args:
            dtensor: The DTensor to extract metadata from.

        Returns:
            Metadata capturing both local and global attributes.
        r   )	r/   r1   r2   r4   rV   rX   rY   r   r   )Údevice_meshrP   Ú_local_tensorr/   r1   r2   r4   Ú_specrY   r   ÚtupleÚ_layout)rZ   r\   s     r    Úfrom_dtensorz_DTensorMeta.from_dtensor˜   s…   € ð Ô)ˆåàÔ'Ô-ØÔ(×/Ò/Ñ1Ô1Ø”-Ø!Ô/à œØ!Ÿ.š.Ñ*Ô*à”}Ô/à5@Ô5OÐW•�kÔ0Ñ1Ô1Ð1ÐUWà#Ô+ð
ñ 
ô 
ð 	
r"   r'   c                ó   — | j         | j        fS )z<Cache key ``(mesh_dim_names, mesh_layout)`` for mesh lookup.)r   r   ©r   s    r    Úmesh_cache_keyz_DTensorMeta.mesh_cache_keyµ   s   € ð Ô# TÔ%5Ð6Ð6r"   r<   r=   Úmeshr   c                óþ   — t          | |¦  «        }t          j        ||| j        | j        | j        d¬¦  «        }| j        r)|                     ¦   «         r|                     d¦  «        }t          t          |¦  «        S )zëReconstruct a DTensor on ``device`` with placements.

        Args:
            device: Target device for the local tensor.
            mesh: The ``DeviceMesh`` to attach.

        Returns:
            A DTensor on ``device``.
        F)r\   rY   r/   r1   Ú	run_checkT)
r?   r   Ú
from_localrY   rV   rX   r4   r@   rA   r   )r   r<   re   Úlocal_tensorÚdts        r    Ú
to_dtensorz_DTensorMeta.to_dtensorº   s‰   € õ .¨d°FÑ;Ô;ˆõ ÔØØØ”ØÔ#ØÔ%Øð
ñ 
ô 
ˆð Ôð 	) "×"6Ò"6Ñ"8Ô"8ð 	)Ø×"Ò" 4Ñ(Ô(ˆBÝ•G˜RÑ Ô Ð r"   rD   r-   rE   c                ó¶  — | |k    rg S t                                | |¦  «        }t          |t          ¦  «        �r
| j        |j        k    r%|                     d| j        › d|j        › �¦  «         | j        |j        k    r%|                     d| j        › d|j        › �¦  «         | j        |j        k    r%|                     d| j        › d|j        › �¦  «         | j        |j        k    r%|                     d| j        › d|j        › �¦  «         | j	        |j	        k    r%|                     d| j	        › d|j	        › �¦  «         n|                     d¦  «         |S )zçReturn field-by-field differences, including DTensor-specific fields.

        Args:
            other: Metadata to compare against.

        Returns:
            List of human-readable difference strings (empty if equal).
        zglobal_shape mismatch: rG   zglobal_stride mismatch: zplacements mismatch: zmesh_dim_names mismatch: zmesh_layout mismatch: z!type: _DTensorMeta vs _TensorMeta)
r-   rK   r9   rP   rV   rH   rX   rY   r   r   rI   s      r    rK   z_DTensorMeta.get_diffÓ   sž  € ð �5Š=ˆ=ØˆIõ
 ×$Ò$ T¨5Ñ1Ô1ˆõ �e�\Ñ*Ô*ñ 	>ØÔ  EÔ$6Ò6Ð6Ø—’ØY¨dÔ.?ÐYÐYÀUÔEWÐYÐYñô ð ð Ô! UÔ%8Ò8Ð8Ø—’Ø\¨tÔ/AÐ\Ð\ÀuÔGZÐ\Ð\ñô ð ð Œ %Ô"2Ò2Ð2Ø—’ØS¨D¬OÐSÐSÀÔAQÐSÐSñô ð ð Ô" eÔ&:Ò:Ð:Ø—’Ø_°Ô0CÐ_Ð_ÈÔI]Ð_Ð_ñô ð ð Ô 5Ô#4Ò4Ð4Ø—’ØV¨TÔ-=ÐVÐVÀ5ÔCTÐVÐVñô ð øð �LŠLÐ<Ñ=Ô=Ð=àˆr"   )rZ   r   r   rP   )r   r'   )r<   r=   re   r   r   r   rL   )r#   r$   r%   r&   r   rV   rM   rX   rY   r   r   rN   ra   Úpropertyrd   rk   rK   r   r"   r    rP   rP   |   sF  € € € € € € ð	ð 	ð  %˜uÐ5KÐ5KÐLÑLÔL€LÐLÐLÐLÑLØ%* U°2Ð%6Ñ%6Ô%6€MÐ6Ð6Ð6Ñ6ð ).¨Øð)ñ )ô )€Jð ð ð ñ ð
 ', e°BÐ&7Ñ&7Ô&7€NÐ7Ð7Ð7Ñ7Ø&+ eØð'ñ 'ô '€Kð ð ð ñ ð ð
ð 
ð 
ñ „\ð
ð8 ð7ð 7ð 7ñ „Xð7ð!ð !ð !ð !ð2*ð *ð *ð *ð *ð *r"   rP   Ú
TensorMeta)r+   c                  ód   — e Zd ZU dZdZded<   dZded<   dZded<   dZded<   dd„Z	dd„Z
dd„ZdS )Ú
_StageMetazPConsolidated tensor metadata for a pipeline stage's forward and backward passes.Nútuple[TensorMeta, ...] | NoneÚinputsÚoutputsú$tuple[TensorMeta | None, ...] | NoneÚinput_gradsÚoutput_gradsr   r3   c                ód   — t          d„ | j        | j        | j        | j        fD ¦   «         ¦  «        S )z)Check if any metadata field is populated.c              3  ó   K  — | ]}|d uV — Œ	d S r   r   )Ú.0Úvs     r    ú	<genexpr>z%_StageMeta.has_any.<locals>.<genexpr>  s:   è è € ð 
ð 
àð �TˆMð
ð 
ð 
ð 
ð 
ð 
r"   )Úanyrr   rs   ru   rv   rc   s    r    Úhas_anyz_StageMeta.has_any  sC   € åð 
ð 
à”k 4¤<°Ô1AÀ4ÔCTÐUð
ñ 
ô 
ñ 
ô 
ð 	
r"   c                ód   — | j         | j        fD ] }|rt          d„ |D ¦   «         ¦  «        r dS Œ!dS )z3Check if any input/output metadata is DTensor type.c              3  óD   K  — | ]}|¯t          |t          ¦  «        V — Œd S r   )r9   rP   )ry   Úms     r    r{   z*_StageMeta.has_dtensors.<locals>.<genexpr>  s1   è è € ÐMÐM¸QÈ1ÐM�Z¨­<Ñ8Ô8ÐMÐMÐMÐMÐMÐMr"   TF)rr   rs   r|   )r   Úmetass     r    Úhas_dtensorsz_StageMeta.has_dtensors  sM   € à”k 4¤<Ð0ð 	ð 	ˆEØð �ÐMÐMÀ%ÐMÑMÔMÑMÔMð Ø�t�tøØˆur"   c                ó&   — | j         duo| j        duS )z-Check if forward metadata is fully populated.N)rr   rs   rc   s    r    Úis_complete_for_forwardz"_StageMeta.is_complete_for_forward  s   € àŒ{ $Ð&ÐC¨4¬<¸tÐ+CÐCr"   )r   r3   )r#   r$   r%   r&   rr   rM   rs   ru   rv   r}   r‚   r„   r   r"   r    rp   rp     s¢   € € € € € € àZÐZà,0€FÐ0Ð0Ð0Ñ0Ø-1€GÐ1Ð1Ð1Ñ1Ø8<€KÐ<Ð<Ð<Ñ<Ø9=€LÐ=Ð=Ð=Ñ=ð
ð 
ð 
ð 
ðð ð ð ðDð Dð Dð Dð Dð Dr"   rp   c                  ó   — e Zd ZU dZded<   dS )Ú_StageForwardMetazLForward metadata transmitted from stage *i* to stage *i+1* during inference.útuple[TensorMeta, ...]Úforward_metasN©r#   r$   r%   r&   rM   r   r"   r    r†   r†   "  s$   € € € € € € àVÐVà)Ð)Ð)Ñ)Ð)Ð)r"   r†   c                  ó   — e Zd ZU dZded<   dS )Ú_StageBackwardMetauº   Backward metadata transmitted from stage *i* to stage *i-1* during inference.

    Gradient placements may differ from forward activations
    (e.g., ``Replicate`` â†’ ``Partial``).
    útuple[TensorMeta | None, ...]Úbackward_metasNr‰   r   r"   r    r‹   r‹   )  s4   € € € € € € ðð ðð ð ñ ð ð r"   r‹   Úmetar<   r=   r   r6   c                óP   — t          j        | j        | j        | j        |¬¦  «        S )zÙCreate a tensor from metadata.

    Args:
        meta: Metadata with shape, stride, and dtype.
        device: Target device for the tensor.

    Returns:
        Empty tensor preserving the exact memory layout.
    )Úsizer1   r2   r<   )rR   Úempty_stridedr/   r1   r2   )rŽ   r<   s     r    r?   r?   6  s0   € õ ÔØŒZØŒ{ØŒjØð	ñ ô ð r"   Útensor_metasr‡   rŒ   c                óB   ‡— dd„Št          ˆfd„| D ¦   «         ¦  «        S )zÑDerive gradient metadata from tensor metadata.

    Returns metadata with the same shape/stride/dtype but ``requires_grad=False``.
    Entries where the source has ``requires_grad=False`` become ``None``.
    r€   rn   r   úTensorMeta | Nonec                ó†   — | j         sd S t          | t          ¦  «        rd S t          | j        | j        | j        d¬¦  «        S )NFr8   )r4   r9   rP   r-   r/   r1   r2   )r€   s    r    Ú
derive_onez&_derive_grad_metas.<locals>.derive_oneT  sQ   € ØŒð 	Ø�4Ý�a�Ñ&Ô&ð 	Ø�4ÝØ”'Ø”8Ø”'Øð	
ñ 
ô 
ð 	
r"   c              3  ó.   •K  — | ]} ‰|¦  «        V — Œd S r   r   )ry   r€   r–   s     €r    r{   z%_derive_grad_metas.<locals>.<genexpr>`  s+   øè è € Ð5Ð5 1��˜A‘”Ð5Ð5Ð5Ð5Ð5Ð5r"   )r€   rn   r   r”   )r_   )r’   r–   s    @r    Ú_derive_grad_metasr˜   K  s<   ø€ ð

ð 

ð 

ð 

õ Ð5Ð5Ð5Ð5¨Ð5Ñ5Ô5Ñ5Ô5Ð5r"   c                  óD   — e Zd ZdZddd„Zdd„Zdd„Zdd„Zdd„Zdd„Z	dS )Ú
_MeshCachezæCache for :class:`DeviceMesh` objects keyed by ``(mesh_dim_names, mesh_layout)``.

    Assumes all pipeline stages share the same rank tensor (true for
    TorchTitan-style frameworks where meshes derive from a common world).
    NÚget_mesh_cbúGetMeshCallback | Noner   ÚNonec                ó"   — i | _         || _        d S r   )Ú_cacheÚ_get_mesh_cb)r   r›   s     r    Ú__init__z_MeshCache.__init__j  s   € Ø68ˆŒØ'ˆÔÐÐr"   Úkeyr'   r   c                óæ   — || j         v r| j         |         S |\  }}| j        €t          d|› d|› d�¦  «        ‚|                      ||¦  «        }|€t          d|› d|› d�¦  «        ‚|| j         |<   |S )a  Return a cached mesh, or create one via the callback.

        Args:
            key: Cache key ``(mesh_dim_names, mesh_layout)``.

        Returns:
            The ``DeviceMesh``.

        Raises:
            PipeliningMetadataError: If not cached and no callback provided.
        Nz+Mesh not found in cache for mesh_dim_names=z, mesh_layout=z`, and no get_mesh callback provided. Provide a get_mesh callback or use DTensors in static mode.z>Mesh lookup failed: callback returned None for mesh_dim_names=z6. Ensure all stages use meshes from the same universe.)rŸ   r    r)   )r   r¢   r   r   re   s        r    Úget_meshz_MeshCache.get_meshn  s×   € ð �$”+ÐÐØ”;˜sÔ#Ð#à&)Ñ#ˆ˜àÔÐ$Ý)ðO¸nð Oð OØ*ðOð Oð Oñô ð ð × Ò  °Ñ=Ô=ˆØˆ<Ý)ðHØ"0ðHð HØ@KðHð Hð Hñô ð ð
  ˆŒ�CÑØˆr"   re   c                ó   — || j         |<   dS )zAdd a mesh to the cache.N©rŸ   )r   r¢   re   s      r    Úputz_MeshCache.put�  s   € àˆŒ�CÑÐÐr"   Útensorsútuple[torch.Tensor | None, ...]c                ó¾   — |D ]Y}t          |t          ¦  «        rB|j        }|j        rt	          |j        ¦  «        nd}|j        }||f}|| j        vr
|| j        |<   ŒZdS )zJExtract and cache meshes from any :class:`DTensor` instances in *tensors*.r   N)r9   r   r\   r   r_   r`   rŸ   )r   r¨   r5   re   Ú	dim_namesr   r¢   s          r    Úupdate_from_tensorsz_MeshCache.update_from_tensors”  s~   € àð 	,ð 	,ˆFÝ˜&¥'Ñ*Ô*ð ,ØÔ)�Ø:>Ô:MÐU�E $Ô"5Ñ6Ô6Ð6ÐSU�	Ø"œl�Ø  +Ð.�Ø˜dœkÐ)Ð)Ø'+�D”K Ñ$øð	,ð 	,r"   r3   c                ó   — || j         v S r   r¦   )r   r¢   s     r    Ú__contains__z_MeshCache.__contains__Ÿ  s   € Ø�d”kÐ!Ð!r"   Úintc                ó*   — t          | j        ¦  «        S r   )ÚlenrŸ   rc   s    r    Ú__len__z_MeshCache.__len__¢  s   € Ý�4”;ÑÔÐr"   r   )r›   rœ   r   r�   )r¢   r'   r   r   )r¢   r'   re   r   r   r�   )r¨   r©   r   r�   )r¢   r'   r   r3   )r   r¯   )
r#   r$   r%   r&   r¡   r¤   r§   r¬   r®   r²   r   r"   r    rš   rš   c  sœ   € € € € € ðð ð(ð (ð (ð (ð (ð ð  ð  ð  ðD ð  ð  ð  ð	,ð 	,ð 	,ð 	,ð"ð "ð "ð "ð ð  ð  ð  ð  ð  r"   rš   c                  ó2   — e Zd ZdZdZdZedd	„¦   «         Zd
S )ÚInferenceModeaÑ  Pipeline-level metadata inference mode, determined collectively across all PP ranks.

    The mode is set by the schedule (not individual stages) because
    ``has_backward`` is only known at schedule creation time and all
    stages must agree to avoid P2P hangs.

    .. attribute:: STATIC

        All stages have sufficient metadata; runtime inference is skipped.

    .. attribute:: DYNAMIC

        At least one stage requires runtime metadata inference.
    ÚstaticÚdynamicrŽ   rp   Ústage_has_backwardr3   r   c                ó†   — |                      ¦   «         sdS |                     ¦   «         sdS |sdS |j        �|j        €dS dS )a'  Determine whether dynamic metadata inference is needed for a stage.

        Args:
            meta: Stage metadata from user-provided args.
            stage_has_backward: Whether a backward pass will be performed.

        Returns:
            ``True`` if dynamic inference is needed.
        TF)r„   r‚   ru   rv   )ÚclsrŽ   r·   s      r    Úneeds_dynamiczInferenceMode.needs_dynamic¾  sf   € ð ×+Ò+Ñ-Ô-ð 	Ø�4ð × Ò Ñ"Ô"ð 	Ø�5ð "ð 	Ø�5ð ÔÐ# tÔ'8Ð'@Ø�4ð ˆur"   N)rŽ   rp   r·   r3   r   r3   )r#   r$   r%   r&   ÚSTATICÚDYNAMICÚclassmethodrº   r   r"   r    r´   r´   «  sH   € € € € € ðð ð €FØ€Gàðð ð ñ „[ðð ð r"   r´   F©Údetachr¿   r3   c               ón   — t          | ¦  «        \  }}|r d„ |D ¦   «         }t          ||¦  «        }||fS |S )a;  Flatten ``args`` into a list, optionally detaching tensors.

    Args:
        args: Nested arguments to flatten.
        detach: If ``True``, detach tensors while preserving ``requires_grad``.

    Returns:
        ``(new_args, flat_detached_args)`` when ``detach=True``;
        ``flat_args`` list otherwise.
    c                óž   — g | ]J}t          |t          j        ¦  «        r,|                     ¦   «                              |j        ¦  «        n|‘ŒKS r   )r9   rR   ÚTensorr¿   rA   r4   )ry   Úas     r    ú
<listcomp>z flatten_args.<locals>.<listcomp>ñ  s[   € ð 
ð 
ð 
ð õ ˜!�Uœ\Ñ*Ô*ðˆA�HŠH‰JŒJ×%Ò% a¤oÑ6Ô6Ð6àð
ð 
ð 
r"   r   )Úargsr¿   Ú	flat_argsÚtreespecÚflat_detachedÚnew_argss         r    Úflatten_argsrÊ   ã  s`   € õ ' tÑ,Ô,Ñ€Iˆxàð 'ð
ð 
ð ð	
ñ 
ô 
ˆõ " -°Ñ:Ô:ˆØ˜Ð&Ð&àÐr"   c                ó$   — t          | d¬¦  «        S )zHFlatten and detach. Deprecated: use ``flatten_args(args, detach=True)``.Tr¾   )rÊ   )rÅ   s    r    Úflatten_args_detachrÌ   þ  s   € å˜ TÐ*Ñ*Ô*Ð*r"   ÚloopÚpp_sizer¯   Ú
num_stagesÚstyleÚstrúdict[int, int]c                ó8  — i }|dk    rt          |¦  «        D ]
}|| z  ||<   Œnv|dk    r]|| z  dk    rt          d|› d| › d�¦  «        ‚d}t          |¦  «        D ]+}|||<   |dz   | z  dk    rŒ|| z  dz  dk    r|dz  }Œ&|dz  }Œ,nt          d	|› d
�¦  «        ‚|S )zá
    Compute the stage id to rank mapping for either a looped or V-style schedule.

    Most commonly num_stages == pp_size * 2, but this function can be used to
    compute the mapping for any number of stages per rank.
    rÍ   rz   r   znum_stages z% must be evenly divisible by pp_size z for V schedulesé   é   zStyle z is not supported.)ÚrangeÚ
ValueError)rÎ   rÏ   rÐ   ÚmappingÚstage_indexÚ
rank_indexs         r    Úgenerate_stage_to_rank_mappingrÛ     s  € ð €GØ�‚€Ý  Ñ,Ô,ð 	9ð 	9ˆKØ#.°Ñ#8ˆG�KÑ Ð ð	9à	�#ŠˆØ˜Ñ 1Ò$Ð$ÝØh˜jÐhÐhÈwÐhÐhÐhñô ð ð ˆ
Ý  Ñ,Ô,ð 	 ð 	 ˆKØ#-ˆG�KÑ à˜a‘ 7Ñ*¨aÒ/Ð/ØØ˜wÑ&¨!Ñ+¨qÒ0Ð0Ø˜a‘�
�
à˜a‘�
�
ð	 õ Ð; %Ð;Ð;Ð;Ñ<Ô<Ð<Ø€Nr"   údict[int, list[int]]c                óþ   — t          | ||¦  «        }i }|                     ¦   «         D ])\  }}||vrg ||<   ||                              |¦  «         Œ*|                     ¦   «         D ]}|                     ¦   «          Œ|S )a  
    Compute the rank to stage id mapping for either a looped or V-style schedule.

    This function inverts the stage_to_rank_mapping to get which stages are assigned to each rank.

    Returns a dictionary mapping rank -> list of stage indices assigned to that rank.
    )rÛ   ÚitemsrH   ÚvaluesÚsort)rÎ   rÏ   rÐ   Ústage_to_rankÚrank_to_stagesÚstage_idÚrankÚstagess           r    Úgenerate_rank_to_stage_mappingræ   %  sž   € õ 3°7¸JÈÑNÔN€Mð ,.€NØ'×-Ò-Ñ/Ô/ð .ð .‰ˆ�$Ø�~Ð%Ð%Ø#%ˆN˜4Ñ Ø�tÔ×#Ò# HÑ-Ô-Ð-Ð-ð !×'Ò'Ñ)Ô)ð ð ˆØ�Š‰ŒˆˆàÐr"   c                  ó2   — e Zd ZU dZded<   ded<   ded<   dS )	ÚPipeInfoz>
    Captures information for a pipeline (`Pipe` object).
    zfx.GraphÚgraphr¯   rÏ   r3   Úhas_loss_and_backwardNr‰   r   r"   r    rè   rè   ?  s<   € € € € € € ðð ð €O€O�OØ€O€O�OØÐÐÑÐÐr"   rè   r5   c                ó”   — t          | t          ¦  «        rt                               | ¦  «        S t                               | ¦  «        S )aª  Extract metadata from a tensor.

    Handles both plain Tensor and DTensor correctly: DTensors are
    dispatched to ``_DTensorMeta.from_dtensor`` which captures local
    shard attributes plus global shape/placement info, while plain
    tensors use ``_TensorMeta.from_tensor``.

    Args:
        tensor: A plain tensor or DTensor.

    Returns:
        ``_TensorMeta`` for plain tensors, ``_DTensorMeta`` for DTensors.
    )r9   r   rP   ra   r-   r;   r:   s    r    Úextract_tensor_metarì   O  s>   € õ �&�'Ñ"Ô"ð /Ý×(Ò(¨Ñ0Ô0Ð0å×&Ò& vÑ.Ô.Ð.r"   )Ú
allow_noner¨   útuple[torch.Tensor, ...] | Nonerí   úLiteral[False]rq   c               ó   — d S r   r   ©r¨   rí   s     r    Úextract_tensor_metasrò   c  s	   € ð
 %( Cr"   ú&tuple[torch.Tensor | None, ...] | NoneúLiteral[True]rt   c               ó   — d S r   r   rñ   s     r    rò   rò   k  s	   € ð
 ,/¨3r"   úAtuple[torch.Tensor | None, ...] | tuple[torch.Tensor, ...] | Nonec               ó  — | €dS g }d}| D ]V}t          |t          j        ¦  «        r#|                     t	          |¦  «        ¦  «         Œ?d}|                     d¦  «         ŒW|s|rt          d¦  «        ‚t          |¦  «        S )aŠ  Extract metadata from a tuple of tensors.

    Args:
        tensors: Tuple of tensors (may include ``None`` when ``allow_none=True``).
        allow_none: If ``True``, preserve ``None`` elements (for gradients).

    Returns:
        Tuple of ``TensorMeta``, or ``None`` if ``tensors`` is ``None``.

    Raises:
        PipeliningMetadataError: If ``None`` found and ``allow_none=False``.
    NFTz_None values are not allowed in tensor metadata tuples. Use allow_none=True for optional values.)r9   rR   rÂ   rH   rì   r)   r_   )r¨   rí   Úmetas_with_noneÚhas_nonerB   s        r    rò   rò   s  s­   € ð" €Øˆtà/1€OØ€HØð )ð )ˆÝ�a�œÑ&Ô&ð 	)Ø×"Ò"Õ#6°qÑ#9Ô#9Ñ:Ô:Ð:Ð:àˆHØ×"Ò" 4Ñ(Ô(Ð(Ð(Øð 
˜(ð 
Ý%ð7ñ
ô 
ð 	
õ �Ñ!Ô!Ð!r"   c                óˆ   — |r|                       ¦   «         n| }t          |t          ¦  «        r|                     ¦   «         S |S )u£  Convert a DTensor to its local shard, or return a plain tensor as-is.

    When ``detach=True``, the tensor is detached before conversion â€”
    this applies to both DTensors and plain tensors.

    Args:
        tensor: A tensor that may be a DTensor.
        detach: If ``True``, detach before ``to_local()`` to avoid
            redistribution during backward.

    Returns:
        The local tensor component.
    )r¿   r9   r   Úto_local)r5   r¿   Úmaybe_detached_tensors      r    Úto_local_if_dtensorrý   —  sF   € ð 06ÐA˜FŸMšM™OœO˜O¸6ÐÝÐ'­Ñ1Ô1ð 0Ø$×-Ò-Ñ/Ô/Ð/Ø Ð r"   rÅ   úCtorch.Tensor | tuple[torch.Tensor, ...] | list[torch.Tensor] | Nonec                ó   — d S r   r   ©rÅ   rí   s     r    Úvalidate_and_normalize_to_tupler  «  s	   € ð '* cr"   úQtorch.Tensor | tuple[torch.Tensor | None, ...] | list[torch.Tensor | None] | Nonec                ó   — d S r   r   r   s     r    r  r  ²  s	   € ð .1¨Sr"   ú�torch.Tensor | tuple[torch.Tensor, ...] | tuple[torch.Tensor | None, ...] | list[torch.Tensor] | list[torch.Tensor | None] | Nonec           	     óð  — | €dS t          | t          j        ¦  «        r| fS t          | t          t          f¦  «        r•t          | ¦  «        D ]_\  }}|€|st          d|› d�¦  «        ‚Œt          |t          j        ¦  «        s(t          d|› dt          |¦  «        j        › d�¦  «        ‚Œ`t          | t          ¦  «        rt          | ¦  «        n| S t          dt          | ¦  «        j        › d�¦  «        ‚)a¬  Normalize ``args`` to a tuple and validate that all elements are tensors.

    Args:
        args: A single tensor, tuple/list of tensors, or ``None``.
        allow_none: If ``True``, permit ``None`` elements (for gradients).

    Returns:
        Tuple of tensors, or ``None`` if ``args`` is ``None``.

    Raises:
        PipeliningMetadataError: On non-tensor values
            (or ``None`` when ``allow_none=False``).
    Nz
Stage arg[zF] is None. Stage args must be tensors. Use kwargs for optional values.z] has type zC. All stage args must be tensors. Use kwargs for non-tensor inputs.z<Stage args must be a tensor, tuple, or list of tensors, got ú.)	r9   rR   rÂ   r_   ÚlistÚ	enumerater)   Útyper#   )rÅ   rí   ÚiÚargs       r    r  r  ¼  s=  € ð, €|ØˆtÝ	�D�%œ,Ñ	'Ô	'ð 
ØˆwˆÝ	�D�5¥$˜-Ñ	(Ô	(ð 
Ý ‘o”oð 	ð 	‰FˆAˆsØˆ{Ø!ð Ý1ðW Qð Wð Wð Wñô ð ð Ý˜c¥5¤<Ñ0Ô0ð Ý-ðY ð Yð Y­t°C©y¬yÔ/Að Yð Yð Yñô ð ðõ )¨­tÑ4Ô4Ð>�u�T‰{Œ{ˆ{¸$Ð>å%ØaÍ4ÐPTÉ:Ì:ÔK^ÐaÐaÐañ
ô 
ð 	
r"   ©Úraise_on_mismatchÚwarn_on_mismatchÚdescÚexpectedÚactualútorch.Tensor | TensorMetar  r  rE   c               ór  — t          |t          j        ¦  «        rt          |¦  «        }n|}t	          |¦  «        t	          |¦  «        urudt	          |¦  «        j        › dt	          |¦  «        j        › �g}|rt          | › d|d         › �¦  «        ‚|r(t          j        | › d|d         › d�t          d¬¦  «         |S | 
                    |¦  «        }|r`|r't          | › dd	                     |¦  «        › �¦  «        ‚|r5t          j        | › d
d	                     |¦  «        › d�t          d¬¦  «         |S )al  
    Compare expected metadata against actual tensor or metadata.

    This is the unified validation/comparison function that uses get_diff() from
    metadata classes. Works with both plain tensors and DTensors.

    For plain tensors: compares shape/stride/dtype/requires_grad.
    For DTensors: compares all properties including global shape and placements.

    Args:
        desc: Description for error/warning messages.
        expected: Expected tensor metadata (_TensorMeta or _DTensorMeta).
        actual: Actual tensor or metadata to compare against.
        raise_on_mismatch: If True, raise PipeliningMetadataError on mismatch.
        warn_on_mismatch: If True, issue a warning on mismatch.

    Returns:
        List of differences (empty if metadata matches).

    Raises:
        PipeliningMetadataError: If raise_on_mismatch=True and differences exist.
    ztype: expected ú, got ú: r   z: Metadata type mismatch. z.. Using dynamically inferred metadata instead.rÕ   ©Ú
stacklevelz; z: Metadata mismatch. )r9   rR   rÂ   rì   r	  r#   r)   ÚwarningsÚwarnÚUserWarningrK   Újoin)r  r  r  r  r  Úactual_metaÚ	type_diffrJ   s           r    Úvalidate_metadatar  ñ  s¡  € õ> �&�%œ,Ñ'Ô'ð Ý)¨&Ñ1Ô1ˆˆàˆõ ˆH�~„~�T +Ñ.Ô.Ð.Ð.àY�d 8™nœnÔ5ÐYÐY½TÀ+Ñ=NÔ=NÔ=WÐYÐYð
ˆ	ð ð 	EÝ)¨TÐ*CÐ*C°Y¸q´\Ð*CÐ*CÑDÔDÐDØð 	ÝŒMØð @ð @°9¸Q´<ð @ð @ð @åØð	ñ ô ð ð Ðð ×Ò˜kÑ*Ô*€Eàð 	Øð 	IÝ)¨TÐ*GÐ*G°T·Y²Y¸uÑ5EÔ5EÐ*GÐ*GÑHÔHÐHØð 	ÝŒMØð @ð @¨d¯iªi¸Ñ.>Ô.>ð @ð @ð @åØð	ñ ô ð ð €Lr"   ú,tuple[torch.Tensor | TensorMeta | None, ...]c               ój  — t          |¦  «        t          |¦  «        k    rV| › dt          |¦  «        › dt          |¦  «        › �}|rt          |¦  «        ‚|rt          j        |t          d¬¦  «         |gS g }t          t          ||d¬¦  «        ¦  «        D ]š\  }\  }}	|€|	€Œ|�|	€Z| › d|› d	|€d
nd› d|	€d
nd› �}|rt          |¦  «        ‚|rt          j        |t          d¬¦  «         |                     |¦  «         Œkt          | › d|› d�||	||¬¦  «        }
| 	                    |
¦  «         Œ›|S )a2  Validate metadata for a tuple of tensors element-wise.

    Args:
        desc: Description prefix for error/warning messages.
        expected: Tuple of expected metadata (may include ``None`` for grads).
        actual: Tuple of actual tensors or metadata to compare against.
        raise_on_mismatch: If ``True``, raise on the first mismatch.
        warn_on_mismatch: If ``True``, issue warnings for mismatches.

    Returns:
        Aggregated list of difference strings.

    Raises:
        PipeliningMetadataError: If lengths differ or on mismatch.
    z: expected z tensors, got rÕ   r  T©ÚstrictNú[z]: expected r�   Úmetadatar  ú]r  )
r±   r)   r  r  r  r  ÚziprH   r  Úextend)r  r  r  r  r  ÚmsgÚ	all_diffsr
  ÚexpÚactrJ   s              r    Úvalidate_tensors_metadatar,  6  s©  € õ. ˆ8�}„}�˜F™œÒ#Ð#ØÐLÐL¥# h¡-¤-ÐLÐL½sÀ6¹{¼{ÐLÐLˆØð 	/Ý)¨#Ñ.Ô.Ð.Øð 	:ÝŒM˜#�{°qÐ9Ñ9Ô9Ð9Øˆuˆà€IÝ"¥3 x°ÀÐ#EÑ#EÔ#EÑFÔFð  ð  ‰ˆ‰:ˆC�Øˆ;˜3˜;ØØˆ;˜#˜+àð ?ð ?˜!ð ?ð ?°3°;¨¨ÀJð ?ð ?Ø!$ �v�v°*ð?ð ?ð ð !ð 3Ý-¨cÑ2Ô2Ð2Øð >Ý”˜c¥;¸1Ð=Ñ=Ô=Ð=Ø×Ò˜SÑ!Ô!Ð!ØÝ!ØˆNˆN�aˆNˆNˆNØØØ/Ø-ð
ñ 
ô 
ˆð 	×Ò˜ÑÔÐÐØÐr"   rÙ   útuple[torch.Tensor, ...]Úgradsr©   Úis_inputr�   c                óÚ  — |rdnd}|› d�}|› d�}t          |¦  «        t          |¦  «        k    r9t          d| › d|› dt          |¦  «        › d|› dt          |¦  «        › d	�¦  «        ‚t          t          ||d
¬¦  «        ¦  «        D ]á\  }\  }}	|j        s6|	�4t          d| › d|› d|› d|› d|› dt          |	¦  «        j        › d�¦  «        ‚|j        r.|	€,t          j        d| › d|› d|› d|› d|› d�t          d¬¦  «         t          |t          ¦  «        rR|j        rK|	�It          |	t          ¦  «        s4t          d| › d|› d|› d|› d|› dt          |	¦  «        j        › d�¦  «        ‚ŒâdS )u/  
    Validate the argsâ†”grads contract for static mode.

    Enforces four rules for each (arg, grad) pair:
      1. len(args) must equal len(grads).
      2. If arg.requires_grad is False, grad must be None.
      3. If arg.requires_grad is True and grad is None, emit a warning
         (this is legal at pipeline boundaries but may indicate a bug).
      4. If arg is a DTensor with requires_grad=True and grad is not None,
         grad must also be a DTensor.

    Args:
        stage_index: The stage index for error messages.
        args: Tuple of forward tensors.
        grads: Tuple of gradient tensors (can include None).
        is_input: True for input_args/input_grads, False for output_args/output_grads.

    Raises:
        PipeliningMetadataError: If any hard rule (1, 2, or 4) is violated.
    ÚinputÚoutputÚ_argsÚ_gradszStage r  z	 length (z) does not match zo). Each forward tensor must have a corresponding gradient entry (use None for tensors that don't require grad).Tr!  Nr#  z] has requires_grad=False, but z] is not None (zE). Non-differentiable tensors must have None as their gradient entry.z] has requires_grad=True, but zT] is None. This is legal at pipeline boundaries but may indicate a missing gradient.rÕ   r  z,] is a DTensor with requires_grad=True, but z] is za, expected DTensor or None. DTensor gradients may have different placements than forward tensors.)r±   r)   r  r&  r4   r	  r#   r  r  r  r9   r   )
rÙ   rÅ   r.  r/  ÚkindÚ	args_nameÚ
grads_namer
  r  Úgrads
             r    Ú'validate_static_arg_grad_correspondencer9  o  s¯  € ð4 Ð,ˆ7ˆ7 H€DØ���€IØ���€Jõ ˆ4�y„y•C˜‘J”JÒÐÝ%ð\�[ð \ð \ Jð \ð \½¸U¹¼ð \ð \Øð\ð \Ý#& t¡9¤9ð\ð \ð \ñ
ô 
ð 	
õ $¥C¨¨e¸DÐ$AÑ$AÔ$AÑBÔBð ð ‰ˆ‰;ˆC�àÔ ð 	 TÐ%5Ý)ðV˜ð Vð V¨	ð Vð V°Að Vð VØ!ðVð VØ$%ðVð VÝ6:¸4±j´jÔ6IðVð Vð Vñô ð ð Ôð 	  ÝŒMð8˜ð 8ð 8¨	ð 8ð 8°Að 8ð 8Ø!ð8ð 8Ø$%ð8ð 8ð 8õ Øðñ ô ð õ �s�GÑ$Ô$ð
	àÔ!ð
	ð Ð Ý˜t¥WÑ-Ô-ð !õ *ðY˜ð Yð Y¨	ð Yð Y°Að Yð YØ!ðYð YØ$%ðYð YÝ,0°©J¬JÔ,?ðYð Yð Yñô ð øð5ð r"   )rŽ   r-   r<   r=   r   r6   )r’   r‡   r   rŒ   )r¿   r3   )rÍ   )rÎ   r¯   rÏ   r¯   rÐ   rÑ   r   rÒ   )rÎ   r¯   rÏ   r¯   rÐ   rÑ   r   rÜ   )r5   r6   r   rn   )r¨   rî   rí   rï   r   rq   )r¨   ró   rí   rô   r   rt   )r¨   rö   rí   r3   r   rt   )F)r5   r6   r¿   r3   r   r6   ).)rÅ   rþ   rí   rï   r   rî   )rÅ   r  rí   rô   r   ró   )rÅ   r  rí   r3   r   rö   )r  rÑ   r  rn   r  r  r  r3   r  r3   r   rE   )r  rÑ   r  rŒ   r  r  r  r3   r  r3   r   rE   )
rÙ   r¯   rÅ   r-  r.  r©   r/  r3   r   r�   )=Ú
__future__r   Úloggingr  Údataclassesr   r   Úenumr   Útypingr   r   r	   r
   r   r   rR   r   Útorch.distributed._mesh_layoutr   Útorch.distributed.tensorr   Útorch.utils._pytreer   r   Útorch.distributed.device_meshr   Ú(torch.distributed.tensor.placement_typesr   Ú	getLoggerr#   Úloggerr   r_   rÑ   r'   rM   ÚRuntimeErrorr)   r-   rP   rn   rp   r†   r‹   r?   r˜   rš   r´   rÊ   rÌ   rÛ   ræ   rè   rì   rò   rý   r  r  r,  r9  r   r"   r    ú<module>rG     sî  ðð #Ð "Ð "Ð "Ð "Ð "Ð "à €€€Ø €€€Ø (Ð (Ð (Ð (Ð (Ð (Ð (Ð (Ø Ð Ð Ð Ð Ð Ø NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ NÐ Nà €€€Ø Ð Ð Ð Ð Ð Ø 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø ,Ð ,Ð ,Ð ,Ð ,Ð ,Ø <Ð <Ð <Ð <Ð <Ð <Ð <Ð <ð ð CØ8Ð8Ð8Ð8Ð8Ð8ØBÐBÐBÐBÐBÐBð 
ˆÔ	˜8Ñ	$Ô	$€ðð ð ð ð �hñ ô ð ð    c¨3 h¤°¸tÑ1CÐ CÔD€Ð DÐ DÐ DÑ DðGð Gð Gð Gð G˜lñ Gô Gð Gð €�$˜dÐ#Ñ#Ô#ðHð Hð Hð Hð Hñ Hô Hñ $Ô#ðHðV €�$˜dÐ#Ñ#Ô#ð@ð @ð @ð @ð @�;ñ @ô @ñ $Ô#ð@ðH $ lÑ2€
Ð 2Ð 2Ð 2Ñ 2ð
 €�ÐÑÔðDð Dð Dð Dð Dñ Dô Dñ ÔðDð6 €�$˜dÐ#Ñ#Ô#ð*ð *ð *ð *ð *ñ *ô *ñ $Ô#ð*ð €�$˜dÐ#Ñ#Ô#ð	ð 	ð 	ð 	ð 	ñ 	ô 	ñ $Ô#ð	ðð ð ð ð*6ð 6ð 6ð 6ð0@ ð @ ð @ ð @ ð @ ñ @ ô @ ð @ ðP0ð 0ð 0ð 0ð 0�Dñ 0ô 0ð 0ðp */ð ð ð ð ð ð ð6+ð +ð +ð 17ðð ð ð ð ðF 17ðð ð ð ð ð4 €�ÐÑÔð ð  ð  ð  ð  ñ  ô  ñ Ôð ð/ð /ð /ð /ð( 
ð "%ð(ð (ð (ð (ð (ñ 
„ð(ð 
ð/ð /ð /ñ 
„ð/ð ð!"ð !"ð !"ð !"ð !"ð !"ðH!ð !ð !ð !ð !ð( 
ð "%ð*ð *ð *ð *ñ 
„ð*ð 
ð !$ð1ð 1ð 1ð 1ñ 
„ð1ð  ð-
ð -
ð -
ð -
ð -
ðt $Ø"ðBð Bð Bð Bð Bð BðT #Ø"ð6ð 6ð 6ð 6ð 6ð 6ðrDð Dð Dð Dð Dð Dr"   