Ë
    *êñi±I  ã                   ó  — d dl Z d dlZd dlmZ d dlmZ d dlZd dlmZ d dl	m
Z
 d dlmZ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«      ej.                  «      Zd Z G d„ d«      Z G d„ d«      Zde
dedee
   fd„Zdej>                  dededeej>                     fd„Z d„ Z!	 	 d"de"edf   de#e$ef   dz  dede"edf   dz  de#e$ef   dz  de"ee"   ee#   f   fd „Z%dee   fd!„Z&y)#é    N)ÚSequence)ÚAny©Úmap_aggregate)Ú	BlockMask)Útree_flattenÚtree_mapÚ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      úi/var/www/pod-logistic/pod-ai/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   0   ó   „ Ør   r   g        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__E   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__J   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   ú<lambda>z,TensorChunkSpec.from_tuple.<locals>.<lambda>\   ó   € œ¨Ó,€ r   r   )r.   Úargs_chunk_specs     r   Ú
from_tuplezTensorChunkSpec.from_tupleM   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   r1   r2   s    r   r4   z+TensorChunkSpec.from_dict.<locals>.<lambda>n   r5   r   r   )r.   Úkwargs_chunk_specs     r   Ú	from_dictzTensorChunkSpec.from_dict`   s   € ô *ØÙ,ó
Ðð !Ð r   N)r   r   r   r   r   ÚintÚ__annotations__r+   r-   ÚstaticmethodÚtupler7   ÚdictÚstrr;   r   r   r   r   r   ;   se   … ñò#ð ƒNò
ò
4ð ðØ˜#˜s˜(‘Oòó ðð$ ð!Ø˜˜c˜‘Nò!ó ñ!r   r   c                   ó   — e Zd Zy)Ú
_ReplicateNr   r   r   r   rC   rC   t   r    r   rC   Ú
block_maskÚ
num_chunksÚreturnc                 óø  ‡ — ‰ j                   j                  d«      dk(  r‰ g|z  S ‰ j                   j                  d«      |k\  st        d«      ‚d}t        j                  ‰ j                   ||«      }t        j                  ‰ j
                  ||«      }‰ j                  �!t        j                  ‰ j                  ||«      ndg|z  }‰ j                  �!t        j                  ‰ j                  ||«      ndg|z  }g }d}t        |«      D ]o  }	ˆ fd„}
|j                  t        j                  ||	   ||	   ||	   ||	   ‰ j                   |
|«      ‰ j                  ¬«      «       |||	   j                  d«      z  }Œq |S )a	  Given a block mask, split the block mask along the batch dimension (dim0).

    Args:
        block_mask: Block mask to split
        num_chunks: Number of chunks to split the block mask into

    Returns:
        chunk_block_masks: List of chunked block masks
    r   é   z;Block mask has fewer batch size than the number of chunks. Nc                 ó   •‡ — ˆˆ fd„}|S )Nc                 ó^   •— t        j                  | ‰«      }‰j                  | |z   |||«      S r   )ÚtorchÚ	full_likeÚmask_mod)ÚbÚhÚq_idxÚkv_idxÚb_offsetrD   Úidxs        €€r   Úbatch_offset_mask_modzI_split_block_mask.<locals>.create_mask_mod.<locals>.batch_offset_mask_mod¤   s.   ø€ Ü Ÿ?™?¨1¨cÓ2�Ø!×*Ñ*¨1¨x©<¸¸EÀ6ÓJÐJr   r   )rS   rT   rD   s   ` €r   Úcreate_mask_modz*_split_block_mask.<locals>.create_mask_mod£   s   ù€ õKð )Ð(r   )Úkv_num_blocksÚ
kv_indicesÚfull_kv_num_blocksÚfull_kv_indicesÚ
BLOCK_SIZErM   Úseq_lengths)rV   ÚsizeÚAssertionErrorrK   Útensor_splitrW   rX   rY   ÚrangeÚappendr   Úfrom_kv_blocksrZ   r[   )rD   rE   Ú	batch_dimÚkv_num_blocks_chunksÚkv_indices_chunksÚfull_kv_num_blocks_chunksÚfull_kv_indices_chunksÚchunk_block_masksÚbatch_offsetÚ	chunk_idxrU   s   `          r   Ú_split_block_maskrj   x   s©  ø€ ð ×Ñ×$Ñ$ QÓ'¨1Ò,Øˆ|˜jÑ(Ð(à×#Ñ#×(Ñ(¨Ó+¨zÒ9ÜØIó
ð 	
ð €IÜ ×-Ñ-Ø× Ñ  *¨ióÐô ×*Ñ*¨:×+@Ñ+@À*ÈiÓXÐð ×(Ñ(Ð4ô 	×Ñ˜:×8Ñ8¸*ÀiÔPàˆV�jÑ ð ð ×%Ñ%Ð1ô 	×Ñ˜:×5Ñ5°zÀ9ÔMàˆV�jÑ ð ð ÐØ€LÜ˜:Ó&ò @ˆ	ô	)ð 	× Ñ Ü×$Ñ$Ø2°9Ñ=Ø,¨YÑ7Ø#<¸YÑ#GØ 6°yÑ AØ%×0Ñ0Ù(¨Ó6Ø&×2Ñ2ôô
	
ð 	Ð,¨YÑ7×<Ñ<¸QÓ?Ñ?‰ð)@ð* Ðr   ÚtensorÚspecc                 ó0  — | j                  |j                  «      |k\  s(t        d| j                  |j                  «      › d�«      ‚t        j                  | ||j                  «      }t
        s|S g }d}|D ]�  }t        j                  | «      }||j                  |j                  «      z   }t        ddd«      g|j                  z  }	t        ||«      |	|j                  <   |||	<   |j                  |«       ||j                  |j                  «      z  }ŒŸ |S )zýGiven a tensor, and a chunking spec, split the tensor.
    Args:

        tensor: Tensor to split
        spec: Chunking spec
        num_chunks: Number of chunks to split the tensor into

    Returns:
        chunk_tensors: List of chunked tensors
    zTensor size z is smaller than num_chunksr   N)
r\   r$   r]   rK   r^   Ú_debug_mask_minibatchesÚ
zeros_likeÚsliceÚndimr`   )
rk   rl   rE   Úchunk_tensorsÚexpanded_chunksÚsplit_dim_idxÚchunk_tensorÚnew_valÚ	upper_idxÚslice_indicess
             r   Ú_split_tensorry   ¹   s  € ð  �;‰;�t—~‘~Ó&¨*Ò4ÜØ˜6Ÿ;™; t§~¡~Ó6Ð7Ð7RÐSó
ð 	
ô ×&Ñ& v¨z¸4¿>¹>ÓJ€Må"ØÐà€OØ€MØ%ò 
;ˆÜ×"Ñ" 6Ó*ˆØ! L×$5Ñ$5°d·n±nÓ$EÑEˆ	ä˜t T¨4Ó0Ð1°G·L±LÑ@ˆÜ(-¨m¸YÓ(Gˆ�d—n‘nÑ%Ø!-ˆ�Ñà×Ñ˜wÔ'à˜×*Ñ*¨4¯>©>Ó:Ñ:‰ð
;ð Ðr   c           	      ó$  — | st        |«      D �cg c]  }i ‘Œ c}S t        | «      t        |«      k(  s?t        dt        | j	                  «       «      › dt        |j	                  «       «      › �«      ‚|€t        d«      ‚t        | d„ ¬«      \  }}t        |d„ ¬«      \  }}g }t        ||d¬«      D �]Z  \  }}	|	t        u st        |	t        «      r|j                  |«       Œ1t        |t        j                  «      rRt        |	t        «      st        d	t        |	«      › �«      ‚|j                  |j                  |	j                  «      «       Œ�t        |t         «      ržt        |	t        «      st        d	t        |	«      › �«      ‚|	j                  d
k(  st        d«      ‚|j"                  j                  d
«      dk(  r|j                  |«       �Œ|j                  |j"                  j                  d
«      «       �ŒKt%        d|	› d|› d�«      ‚ t'        g |¢|‘­Ž }
t        |
«      D �cg c]  }g ‘Œ }}t        ||d¬«      D ]¤  \  }}	g }|	t        u st        |	t        «      r|g|
z  }nWt        |t        j                  «      rt)        ||	|
«      }n/t        |t         «      rt+        ||
«      }nt%        d|	› d|› d�«      ‚t        ||d¬«      D ]  \  }}|j                  |«       Œ Œ¦ |D �cg c]  }t-        ||«      ‘Œ c}S c c}w c c}w 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
    zargs_dict.keys() = z args_chunk_spec.keys() = z.args_chunk_spec should have been set by callerc                 ó"   — t        | t        «      S r   ©Ú
isinstancer   ©Úxs    r   r4   z%_shard_dict_of_args.<locals>.<lambda>  s   € ¤Z°´9Ó%=€ r   ©Úis_leafc                 ó"   — t        | t        «      S r   r|   r~   s    r   r4   z%_shard_dict_of_args.<locals>.<lambda>  s   € ¬:°a¼Ó+C€ r   T©ÚstrictzExpected TensorChunkSpec, got r   z#BlockMask only supports split_dim=0rH   zUnsupported chunk spec: z and value: z combination.)r_   Úlenr]   ÚlistÚkeysr   ÚziprC   r}   r`   rK   ÚTensorr   Útyper\   r$   r   rV   Ú
ValueErrorÚminry   rj   r
   )Ú	args_dictr6   rE   Ú_ÚvaluesÚ	tree_specÚchunk_specsÚsplit_sizesÚvrl   Úresult_num_chunksÚflat_split_resultsÚv_splitsÚ_flat_split_resultÚ_v_splits                  r   Ú_shard_dict_of_argsr™   ã   s
  € ñ$ Ü! *Ó-Ö.�q’Ò.Ð.äˆy‹>œS Ó1Ò1ÜØ!¤$ y§~¡~Ó'7Ó"8Ð!9ð :(Ü(,¨_×-AÑ-AÓ-CÓ(DÐ'EðGó
ð 	
ð ÐÜÐMÓNÐNä$ØÑ=ôÑ€FˆIô "ØÑ!Cô�N€K�ð
 €KÜ�v˜{°4Ô8ó ‰ˆˆ4ð ”:Ñ¤¨D´*Ô!=Ø×Ñ˜zÕ*Ü˜œ5Ÿ<™<Ô(Ü˜d¤OÔ4Ü$Ð'EÄdÈ4ÃjÀ\Ð%RÓSÐSØ×Ñ˜qŸv™v d§n¡nÓ5Õ6Ü˜œ9Ô%Ü˜d¤OÔ4Ü$Ð'EÄdÈ4ÃjÀ\Ð%RÓSÐSØ—>‘> QÒ&Ü$Ð%JÓKÐKà�‰×#Ñ# AÓ&¨!Ò+Ø×"Ñ" :Ö.à×"Ñ" 1§?¡?×#7Ñ#7¸Ó#:Ö;äØ*¨4¨&°¸Q¸C¸}ÐMóð ð)ô. Ð5˜[Ð5¨*Ò5Ðä16Ð7HÓ1IÖ$J¨A¢RÐ$JÐÐ$JÜ�v˜{°4Ô8ò 0‰ˆˆ4Ø"$ˆØ”:Ñ¤¨D´*Ô!=Ø�sÐ.Ñ.‰HÜ˜œ5Ÿ<™<Ô(Ü$ Q¨Ð.?Ó@‰HÜ˜œ9Ô%Ü(¨Ð,=Ó>‰HäØ*¨4¨&°¸Q¸C¸}ÐMóð ô -0Ø °ô-
ò 	0Ñ(Ð ð ×%Ñ% hÕ/ñ	0ð0ð( #5öàô 	Ð)¨9Õ5òð ùò /ùòX %Kùò&s   �	LÈ)	LË-LÚargs.ÚkwargsÚchunksr6   r:   c                 ój  ‡	— |€i }d„ }|€t        || d„ ¬«      }|€t        ||d„ ¬«      }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
    c                 óv   — t        | t        j                  t        z  «      rt	        t
        «      S t        «       S r   )r}   rK   r‰   r   r   ÚDEFAULT_CHUNK_DIMrC   ©r“   s    r   Údefault_specz3split_args_kwargs_into_chunks.<locals>.default_specx  s)   € Ü�aœŸ™¬	Ñ1Ô2Ü"Ô#4Ó5Ð5ä“<Ðr   c                 ó"   — t        | t        «      S r   r|   r    s    r   r4   z/split_args_kwargs_into_chunks.<locals>.<lambda>€  s   € ´*¸QÄ	Ó2J€ r   r€   c                 ó"   — t        | t        «      S r   r|   r    s    r   r4   z/split_args_kwargs_into_chunks.<locals>.<lambda>…  s   € ´J¸qÄ)Ó4L€ r   z;args and kwargs are split into different number of chunks: z, c              3   ó(   •K  — | ]	  }‰|   –— Œ y ­wr   r   )Ú.0ÚiÚ
chunk_argss     €r   ú	<genexpr>z0split_args_kwargs_into_chunks.<locals>.<genexpr>§  s   øè ø€ Ò< ˆj˜�mÑ<ùs   ƒ)r	   r™   r@   Ú	enumerater…   ÚRuntimeErrorr?   r_   )rš   r›   rœ   r6   r:   r¡   Úargs_split_dictÚreal_num_chunksÚkwargs_splitr§   Ú
args_splits            ` r   r   r   ;  sW  ø€ ðp €~Øˆò ð ÐÜ"Ø˜$Ñ(Jô
ˆð Ð Ü$Ø˜&Ñ*Lô
Ðô *ÜŒY�t‹_ÓÜŒY�Ó'Ó(Øó€Oô
 ˜/Ó*€Oä&ØØØó€Lô ˆ<Ó˜?Ò*ô ˜lÓ+ˆä-Ü”˜4“Ó!Ü”˜?Ó+Ó,Øó
ˆô ˆ?Óœs <Ó0Ò0ÜØIÜ�?Ó#Ð$ B¤s¨<Ó'8Ð&9ð;ó
ð 	
ð *÷àô 	Ó<¤U¬3¨z«?Ó%;Ô<Õ<ð€Jð ð
 �|Ð#Ð#ùòs   Ã=-D0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 �]n  \  }
}t        |t        «      �r¢t        t        |«      «      D �cg c]
  }||   |
   ‘Œ }}t        �r@|d   j                  }|dd D ],  }|j                  |k(  rŒt        d|› d|j                  › �«      ‚ t        j                  t        j                  |dd	iŽt        |«      |j                  ¬
«      }g }d}t        |«      t        |«      k(  s#t        dt        |«      › dt        |«      › �«      ‚t!        ||d¬«      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Œt        d|› d||   |
   › �«      ‚ |	j                  |«       �Œq t1        |	|«      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 rH   zExpected shape z, got ÚdeviceÚmeta)Úsectionsr3   z6Expected len(partial_values) == len(meta_chunks), got z != Trƒ   r2   z	Expected )r   r   rŸ   r…   r‹   r`   r©   r}   r_   rn   Úshaper]   rK   r^   Úemptyr$   rˆ   r\   rp   rq   Úcatr   r   r   r
   )rœ   Ú
chunk_specÚspec_flattenedÚflatten_specÚchunk0_flatÚchunks_flattenedÚchunkÚchunk_flattenedrŽ   Úargs_flattenedÚarg_idxÚargri   Úpartial_valuesÚoverall_shapeÚvalÚmeta_chunksÚvalues_to_catÚchunk_start_idxÚpartial_valueÚ
meta_chunkÚchunk_end_idxrx   ÚslicedÚreduced_valÚvalues                             r   r   r   ®  sn  € ð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ó <)‰ˆ�Ü�cœ?Õ+ô "'¤sÐ+;Ó'<Ó!=öàð ! Ñ+¨GÓ4ðˆNð ö
 'à .¨qÑ 1× 7Ñ 7�Ø)¨!¨"Ð-ò �CØŸ9™9¨Ó5Ü,Ø-¨m¨_¸FÀ3Ç9Á9À+ÐNóð ðô
 $×0Ñ0Ü—K‘K Ð>°vÑ>Ü  Ó0ØŸ™ô�ð !#�Ø"#�Ü˜>Ó*¬c°+Ó.>Ò>Ü(ØPÔQTÐUcÓQdÐPeÐeiÔjmÐnyÓjzÐi{Ð|óð ô 25Ø" K¸ô2ò 
4Ñ-�M :ð %4°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Ð*:Ó&;Ó<ò �	Ø'¨	Ñ2°7Ñ;¸uÓDÜ(Ø# E 7¨&Ð1AÀ)Ñ1LÈWÑ1UÐ0VÐWóð ðð
 ×!Ñ! %Ö(ðy<)ô~ ˜.¨,Ó7Ð7ùò{s   Ã
L)NN)'ÚloggingÚoperatorÚcollections.abcr   Útypingr   rK   Útorch.fx.noder   Ú!torch.nn.attention.flex_attentionr   Útorch.utils._pytreer   r	   r
   Ú__all__Ú	getLoggerr   Úloggerrn   r   r   rk   ÚaddÚsum_reducerrŸ   r   rC   r<   r†   rj   r‰   ry   r™   r?   r@   rA   r   r   r   r   r   ú<module>rØ      s¡  ðó Û Ý $Ý ã Ý 'Ý 7ß FÑ Fò€ð 
ˆ×	Ñ	˜8Ó	$€ðð
  Ð ÷#ñ #ô$	�>ô 	ñ ˜<˜5Ÿ<™<¨Ó,¨h¯l©lÓ;€ð Ð ÷5!ñ 5!÷r	ñ 	ð>Øð>àð>ð 
ˆ)�_ó>ðB'Ø�L‰Lð'à
ð'ð ð'ð ˆe�l‰lÑó	'òTUðx ;?Ø;?ñp$Ø
��S�‰/ðp$à��c�‰N˜TÑ!ðp$ð ðp$ð ˜?¨CÐ/Ñ0°4Ñ7ð	p$ð
 ˜C Ð0Ñ1°DÑ8ðp$ð ˆ4�‰;˜˜T™
Ð"Ñ#óp$ðfC8Ø�‰IôC8r   