Ë
    g^(h<?  ã                   óŠ  — d dl Z d dlmZmZ d dlZd dlmZ d dlmZm	Z	 g d¢Z
 e j                  e«      Z	 dZ G d„ d«      Z G d	„ d
e«      Z e ej"                  d«      d„ «      Zd Z G d„ d«      Z G d„ d«      Zd„ Z	 	 ddeedf   deeeef      dedeeedf      deeeef      deee   ee   f   fd„Zdee   fd„Zy)é    N)ÚAnyÚOptional©Úmap_aggregate)Útree_flattenÚtree_unflatten)ÚTensorChunkSpecÚsplit_args_kwargs_into_chunksÚmerge_chunksFc                   ó   — e Zd ZdZd„ Zy)Ú_CustomReducera$  
    Custom reducer class that can be used to specify a custom operation that
    reduces losses of multiple microbatches into one value.

    Example:
    >>> # xdoctest: +SKIP
    >>> sum_reducer = _CustomReducer(
    >>>     torch.tensor(0.0),
    >>>     lambda a, b: a + b
    >>> )
    c                 ó    — || _         || _        y ©N)Ú
init_valueÚ	reduce_fn)Úselfr   r   s      úe/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributed/pipelining/microbatch.pyÚ__init__z_CustomReducer.__init__(   s   € Ø$ˆŒØ"ˆ�ó    N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   © r   r   r   r      s   „ ñ
ó#r   r   c                   ó   — e Zd Zy)Ú_LossReducerN©r   r   r   r   r   r   r   r   -   ó   „ Ør   r   g        c                 ó   — | |z   S r   r   )ÚaÚbs     r   ú<lambda>r"   1   s
   € ¸1¸q¹5€ r   c                   ón   — e Zd ZU dZd„ Zeed<   d„ Zd„ Ze	de
edf   fd„«       Ze	deeef   fd	„«       Zy
)r	   z2
    Class used to specify chunking of inputs
    c                 ó   — || _         y r   ©Ú	split_dim)r   r&   s     r   r   zTensorChunkSpec.__init__=   s	   € Ø"ˆ�r   r&   c                 ó|   — | j                   j                  › d| j                   j                  › d| j                  › d�S )Nú.ú(ú))Ú	__class__r   r   r&   ©r   s    r   Ú__repr__zTensorChunkSpec.__repr__B   s9   € à�~‰~×(Ñ(Ð)¨¨4¯>©>×+BÑ+BÐ*CÀ1ÀTÇ^Á^ÐDTÐTUÐVð	
r   c                 ó"   — d| j                   › d�S )NzTensorChunkSpec(r*   r%   r,   s    r   Ú__str__zTensorChunkSpec.__str__G   s   € Ø! $§.¡.Ð!1°Ð3Ð3r   Ú
chunk_dims.c                 ó    — t        | d„ «      }|S )aŠ  
        A helper for creating a tuple of `TensorChunkSpec` from a tuple of chunk
        dimensions (int's).
        Example:
            >>> # xdoctest: +SKIP
            >>> # There are three positional arguments to the model, and
            >>> # we are chunking them along dimension 0, 0 and 1, respectively
            >>> args_chunk_spec = TensorChunkSpec.from_tuple((0, 0, 1))
        c                 ó   — t        | «      S r   ©r	   ©Údims    r   r"   z,TensorChunkSpec.from_tuple.<locals>.<lambda>Y   ó   € œ¨Ó,€ r   r   )r0   Úargs_chunk_specs     r   Ú
from_tuplezTensorChunkSpec.from_tupleJ   s   € ô (ØÙ,ó
ˆð Ðr   c                 ó    — t        | d„ «      }|S )a\  
        A helper for creating a dictionary of `TensorChunkSpec` from a
        dictionary of chunk dimensions (int's).
        Example:
            >>> # xdoctest: +SKIP
            >>> # Chunk dimension 0 for the "id" argument, 1 for the "mask" argument
            >>> kwargs_chunk_spec = TensorChunkSpec.from_dict({"id": 0, "mask": 1})
        c                 ó   — t        | «      S r   r3   r4   s    r   r"   z+TensorChunkSpec.from_dict.<locals>.<lambda>k   r6   r   r   )r0   Úkwargs_chunk_specs     r   Ú	from_dictzTensorChunkSpec.from_dict]   s   € ô *ØÙ,ó
Ðð !Ð r   N)r   r   r   r   r   ÚintÚ__annotations__r-   r/   ÚstaticmethodÚtupler8   ÚdictÚstrr<   r   r   r   r	   r	   8   se   … ñò#ð ƒNò
ò
4ð ðØ˜#˜s˜(‘Oòó ðð$ ð!Ø˜˜c˜‘Nò!ó ñ!r   r	   c                   ó   — e Zd Zy)Ú
_ReplicateNr   r   r   r   rD   rD   q   r   r   rD   c                 óô  — i }g }|}d}t        | «      t        |«      k(  s;J dt        | j                  «       «      › dt        |j                  «       «      › �«       ‚| j                  «       D �]G  \  }}t	        |«      \  }	}
|j                  |
«       ||   }|€J ‚t	        |«      \  }}t        |	«      t        |«      k7  rt        d|› d|› �«      ‚g }t        |	|«      D �]Ì  \  }}|t        u st        |t        j                  «      s|j                  |g|z  «       Œ?t        |t        «      �rqt        |t        j                  «      s
J |› d�«       ‚|j                  |j                  «      }||k  r9|r"t        j!                  d|› d	|› d
|› d�«       |}nt#        d|› d|› d|› d�«      ‚t        j$                  |||j                  «      }t&        r¸g }d}|D ]�  }t        j(                  |«      }||j                  |j                  «      z   }t+        ddd«      g|j,                  z  }t+        ||«      ||j                  <   |||<   |j                  |«       ||j                  |j                  «      z  }ŒŸ |j                  |«       n|j                  |«       d}�ŒÁt/        d|› �«      ‚ |||<   �ŒJ g }t1        |«      D ]D  }i }|j                  «       D ]  \  }}|D �cg c]  }||   ‘Œ	 }}|||<   Œ |j                  |«       ŒF g }|D ]b  } i }!t        |«      t        | «      k(  sJ ‚t        | j                  «       |«      D ]  \  \  }}}"t3        ||"«      |!|<   Œ |j                  |!«       Œd |S c c}w )aW  
    Given a dictionary of args, and a dictionary of chunking specs, shard the
    args according to the chunking specs.

    Args:
        args_dict: Dictionary of args
        args_chunk_spec: Dictionary of chunking specs
        num_chunks: Number of chunks to shard the args into

    Returns:
        args_split: List of sharded args
    Tzargs_dict.keys() = z args_chunk_spec.keys() = NzArgument value z9 did not have the same number of values as as chunk spec z is not a tensorz%Tensor size on chunking dimension is z', downsizing the number of chunks from z to r(   zArg z% on chunking dimension has a size of z$, smaller than the number of chunks zŒ. PiPPy cannot reduce the number of chunks because other arguments have bigger chunk-dimension sizes. Please adjust your num_chunks setting.r   FzUnrecognized chunk spec: )ÚlenÚlistÚkeysÚitemsr   ÚappendÚ
ValueErrorÚziprD   Ú
isinstanceÚtorchÚTensorr	   Úsizer&   ÚloggerÚwarningÚRuntimeErrorÚtensor_splitÚ_debug_mask_minibatchesÚ
zeros_likeÚsliceÚndimÚ	TypeErrorÚranger   )#Ú	args_dictr7   Ú
num_chunksÚargs_sharded_replicatedÚ	arg_specsÚreal_num_chunksÚfirst_tensorÚarg_keyÚargÚflatÚspecÚ
chunk_specÚchunk_spec_flatÚ_Úsharded_arg_flatÚvÚchunk_vÚv_split_dim_sizeÚchunk_tensorsÚexpanded_chunksÚsplit_dim_idxÚchunk_tensorÚnew_valÚ	upper_idxÚslice_indicesÚchunks_flatÚ	chunk_idxÚ
chunk_argsÚkeyÚv_flatÚarg_single_chunkÚ
args_splitÚchunkÚper_chunk_argsÚarg_specs#                                      r   Ú_shard_dict_of_argsr}   u   s  € ð( !ÐØ€Ià €OØ€Läˆy‹>œS Ó1Ò1ð Ø
œd 9§>¡>Ó#3Ó4Ð5Ð5OÔPTÐUd×UiÑUiÓUkÓPlÐOmÐnóÐ1ð "Ÿ™Ó)ó I<‰ˆ�Ü! #Ó&‰
ˆˆdØ×Ñ˜Ôà$ WÑ-ˆ
ØÐ%Ð%Ð%Ü)¨*Ó5Ñˆ˜Üˆt‹9œ˜OÓ,Ò,ÜØ! # ð '+Ø+5¨,ð8óð ð
 Ðä˜d OÓ4ó 8	G‰JˆAˆwØœ*Ñ$¬J°q¼%¿,¹,Ô,GØ ×'Ñ'¨¨¨oÑ(=Õ>Ü˜G¤_Õ5ô " !¤U§\¡\Ô2ÐJ°q°cÐ9IÐ4JÓJÐ2à#$§6¡6¨'×*;Ñ*;Ó#<Ð Ø# oÒ5Ù#ô Ÿ™ØCÐDTÐCUð VDØDNÀ<ÈtÐTdÐSeÐefðhôð +;™ä*Ø" 7 )Ð+PÐQaÐPbð cAØAKÀð MEðEóð ô !&× 2Ñ 2Ø�¨×(9Ñ(9ó!�õ +Ø&(�Oà$%�MØ(5ò N˜Ü"'×"2Ñ"2°1Ó"5˜Ø$1°L×4EÑ4EÀg×FWÑFWÓ4XÑ$X˜	ä).¨t°T¸4Ó)@Ð(AÀGÇLÁLÑ(P˜Ü;@Ø)¨9ó<˜ g×&7Ñ&7Ñ8ð 2>˜ Ñ.à'×.Ñ.¨wÔ7à%¨×):Ñ):¸7×;LÑ;LÓ)MÑM™ðNð %×+Ñ+¨OÕ<à$×+Ñ+¨MÔ:à$’äÐ";¸G¸9Ð EÓFÐFðq8	Gðt ,<Ð Ó(ðSI<ðX €KÜ˜?Ó+ò 'ˆ	Øˆ
Ø/×5Ñ5Ó7ò 	/‰HˆC�Ø@CÖD°f  yÓ 1ÐDÐÐDØ.ˆJ�sŠOð	/ð 	×Ñ˜:Õ&ð'ð €Jàò *ˆØˆÜ�9‹~¤ U£Ò+Ð+Ð+Ü$'¨¯©«°yÓ$Aò 	@Ñ ‰JˆS�#˜Ü"0°°hÓ"?ˆN˜3Òð	@à×Ñ˜.Õ)ð*ð Ðùò  Es   Ë"M5Úargs.ÚkwargsÚchunksr7   r;   Úreturnc                 ó¦  ‡— |€i }|€t        t        «      ft        | «      z  }|€#t        j	                  |t        t        «      «      }t        t        t        | «      «      t        t        |«      «      |«      }t        |«      }t        |||«      }t        |«      |k  r<t        |«      }t        t        t        | «      «      t        t        |«      «      |«      }t        |«      t        |«      k7  r#t        dt        |«      › dt        |«      › �«      ‚|D �‡cg c](  Št        ˆfd„t        t        ‰«      «      D «       «      ‘Œ* }	}|	|fS c c}w )a  
    Given a sequence of args and kwargs, split them into a number of chunks
    according to  their respective chunking specs.

    Args:
        args: Tuple of args
        kwargs: Dict of kwargs
        chunks: Number of chunks to split the args and kwargs into
        args_chunk_spec: chunking specs for args, in same shape as args
        kwargs_chunk_spec: chunking specs for kwargs, in same shape as kwargs

    Returns:
        args_split: List of sharded args
        kwargs_split: List of sharded kwargs
    z;args and kwargs are split into different number of chunks: z, c              3   ó(   •K  — | ]	  }‰|   –— Œ y ­wr   r   )Ú.0Úiru   s     €r   ú	<genexpr>z0split_args_kwargs_into_chunks.<locals>.<genexpr>V  s   øè ø€ Ò< ˆj˜�mÑ<ùs   ƒ)
r	   ÚDEFAULT_CHUNK_DIMrF   rA   Úfromkeysr}   Ú	enumeraterS   r@   rZ   )
r~   r   r€   r7   r;   Úargs_split_dictr_   Úkwargs_splitru   ry   s
           ` r   r
   r
   ô   sW  ø€ ðp €~Øˆð ÐÜ*Ô+<Ó=Ð?Ä#ÀdÃ)ÑKˆàÐ Ü ŸM™M¨&´/ÔBSÓ2TÓUÐä)ÜŒY�t‹_ÓÜŒY�Ó'Ó(Øó€Oô
 ˜/Ó*€Oä&ØØØó€Lô ˆ<Ó˜?Ò*ô ˜lÓ+ˆä-Ü”˜4“Ó!Ü”˜?Ó+Ó,Øó
ˆô ˆ?Óœs <Ó0Ò0ÜØIÜ�?Ó#Ð$ B¤s¨<Ó'8Ð&9ð;ó
ð 	
ð *÷àô 	Ó<¤U¬3¨z«?Ó%;Ô<Õ<ð€Jð ð
 �|Ð#Ð#ùòs   Ä-Ec                 óœ  — |�t        |«      \  }}n-t        | d   «      \  }}t        t        «      gt        |«      z  }g }| D ]I  }t        |«      \  }}t        |«      t        |«      k7  rt	        d|› d|› �«      ‚|j                  |«       ŒK g }	t        |«      D �]  \  }
}t        |t        «      �rft        t        |«      «      D �cg c]
  }||   |
   ‘Œ }}t        �r|d   j                  }|dd D ]  }|j                  |k(  rŒJ ‚ t        j                  t        j                  |ddiŽt        |«      |j                  ¬«      }g }d}t        |«      t        |«      k(  sJ ‚t        ||«      D ]o  \  }}||j!                  |j                  «      z   }t#        ddd«      g|j$                  z  }t#        ||«      ||j                  <   ||   }|j                  |«       |}Œq n|}|	j                  t        j&                  ||j                  ¬	«      «       �Œ~t        |t(        «      rP|j*                  }t        t        |«      «      D ]  }|j-                  |||   |
   «      }Œ |	j                  |«       �ŒÞ|d   |
   }t        dt        |«      «      D ]  }||   |
   |k(  rŒJ ‚ |	j                  |«       �Œ  t/        |	|«      S c c}w )
zæ
    Given a list of chunks, merge them into a single value according to
    the chunk spec.

    Args:
        chunks: list of chunks
        chunk_spec: Chunking spec for the chunks

    Returns:
        value: Merged value
    Nr   zChunk z did not match chunk spec é   ÚdeviceÚmeta)Úsectionsr5   r4   )r   r	   r‡   rF   rK   rJ   r‰   rM   rZ   rU   ÚshaperN   rT   Úemptyr&   rL   rP   rW   rX   Úcatr   r   r   r   )r€   re   Úspec_flattenedÚflatten_specÚchunk0_flatÚchunks_flattenedrz   Úchunk_flattenedrg   Úargs_flattenedÚarg_idxrb   rt   Úpartial_valuesÚoverall_shapeÚvalÚmeta_chunksÚvalues_to_catÚchunk_start_idxÚpartial_valueÚ
meta_chunkÚchunk_end_idxrr   ÚslicedÚreduced_valÚvalues                             r   r   r   ]  sý  € ðZ ÐÜ'3°JÓ'?Ñ$ˆ™ô %1°¸±Ó$;Ñ!ˆ�\Ü)Ô*;Ó<Ð=ÄÀKÓ@PÑPˆð Ðàò 1ˆÜ)¨%Ó0Ñˆ˜ÜˆÓ¤3 ~Ó#6Ò6Ü˜v e WÐ,FÀzÀlÐSÓTÐTà×Ñ Õ0ð1ð €NÜ! .Ó1ó 0)‰ˆ�Ü�cœ?Õ+ô "'¤sÐ+;Ó'<Ó!=öàð ! Ñ+¨GÓ4ðˆNð ö
 'à .¨qÑ 1× 7Ñ 7�Ø)¨!¨"Ð-ò 6�CØŸ9™9¨Ó5Ð5Ð5ð6ä#×0Ñ0Ü—K‘K Ð>°vÑ>Ü  Ó0ØŸ™ô�ð !#�Ø"#�Ü˜>Ó*¬c°+Ó.>Ò>Ð>Ð>Ü14°^À[Ó1Qò 4Ñ-�M :Ø$3°j·o±oÀcÇmÁmÓ6TÑ$T�Mä%*¨4°°tÓ%<Ð$=À×@RÑ@RÑ$R�MÜ38¸È-Ó3X�M #§-¡-Ñ0Ø*¨=Ñ9�FØ!×(Ñ(¨Ô0à&3‘Oñ4ð !/�à×!Ñ!¤%§)¡)¨M¸s¿}¹}Ô"MÖNÜ˜œ^Ô,ØŸ.™.ˆKä"¤3Ð'7Ó#8Ó9ò �	Ø!Ÿm™mØÐ!1°)Ñ!<¸WÑ!Eó‘ðð
 ×!Ñ! +Ö.à$ QÑ'¨Ñ0ˆEÜ" 1¤cÐ*:Ó&;Ó<ò E�	Ø'¨	Ñ2°7Ñ;¸uÓDÐDÐDðEà×!Ñ! %Ö(ða0)ôf ˜.¨,Ó7Ð7ùòcs   Ã
K	)NN)ÚloggingÚtypingr   r   rN   Útorch.fx.noder   Útorch.utils._pytreer   r   Ú__all__Ú	getLoggerr   rQ   rU   r   r   ÚtensorÚsum_reducerr‡   r	   rD   r}   r@   rA   rB   r=   rG   r
   r   r   r   r   ú<module>r¯      s8  ðó ß  ã Ý 'ß <ò€ð 
ˆ×	Ñ	˜8Ó	$€ðð
  Ð ÷#ñ #ô$	�>ô 	ñ ˜<˜5Ÿ<™<¨Ó,Ñ.@ÓA€ð Ð ÷5!ñ 5!÷r	ñ 	ò|ðF >BØ>Bñf$Ø
��S�‰/ðf$à�T˜#˜s˜(‘^Ñ$ðf$ð ðf$ð ˜e O°SÐ$8Ñ9Ñ:ð	f$ð
    S¨/Ð%9Ñ :Ñ;ðf$ð ˆ4�‰;˜˜T™
Ð"Ñ#óf$ðRw8Ø�‰Iôw8r   