§
    ŠŠtj[  ã                   ób  — d dl Z d dlZd dlmZ d dl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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)dS )$é    N)ÚSequence)ÚAnyÚcast)ÚDTensor©Ú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dS )Ú_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                 ó"   — || _         || _        d S ©N)Ú
init_valueÚ	reduce_fn)Úselfr   r   s      úe/var/www/html/CA-Chatbot/venv/lib/python3.11/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dS )Ú_LossReducerN©r   r   r   r   r   r   r    r    1   ó   € € € € € Ø€Dr   r    g        c                   óˆ   — 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d
S )r   z2
    Class used to specify chunking of inputs
    c                 ó   — || _         d S r   ©Ú	split_dim)r   r&   s     r   r   zTensorChunkSpec.__init__A   s   € Ø"ˆŒˆˆr   r&   c                 óJ   — | j         j        › d| j         j        › d| j        › d�S )Nú.ú(ú))Ú	__class__r   r   r&   ©r   s    r   Ú__repr__zTensorChunkSpec.__repr__F   s/   € àŒ~Ô(ÐVÐV¨4¬>Ô+BÐVÐVÀTÄ^ÐVÐVÐVð	
r   c                 ó   — d| j         › d�S )NzTensorChunkSpec(r*   r%   r,   s    r   Ú__str__zTensorChunkSpec.__str__K   s   € Ø3 $¤.Ð3Ð3Ð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   )r0   Úargs_chunk_specs     r   Ú
from_tuplezTensorChunkSpec.from_tupleN   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   r6   z+TensorChunkSpec.from_dict.<locals>.<lambda>o   r7   r   r   )r0   Úkwargs_chunk_specs     r   Ú	from_dictzTensorChunkSpec.from_dicta   s%   € õ *ØØ,Ð,ñ
ô 
Ðð !Ð r   N)r   r   r   r   r   ÚintÚ__annotations__r-   r/   ÚstaticmethodÚtupler9   ÚdictÚstrr=   r   r   r   r   r   <   s¸   € € € € € € ðð ð#ð #ð #ð €N€N�Nð
ð 
ð 
ð
4ð 4ð 4ð ðØ˜#˜s˜(”Oðð ð ñ „\ðð$ ð!Ø˜˜c˜”Nð!ð !ð !ñ „\ð!ð !ð !r   r   c                   ó   — e Zd ZdS )Ú
_ReplicateNr!   r   r   r   rE   rE   u   r"   r   rE   Ú
block_maskÚ
num_chunksÚreturnc                 óÞ  ‡ — ‰ j                              d¦  «        dk    r‰ g|z  S ‰ 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 ]~}	ˆ fd„}
| 	                    t          j        ||	         ||	         ||	         ||	         ‰ j         |
|¦  «        ‰ j        ¬¦  «        ¦  «         |||	                              d¦  «        z  }Œ|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                 ód   •— t          j        | ‰¦  «        }‰                     | |z   |||¦  «        S r   )ÚtorchÚ	full_likeÚmask_mod)ÚbÚhÚq_idxÚkv_idxÚb_offsetrF   Úidxs        €€r   Úbatch_offset_mask_modzI_split_block_mask.<locals>.create_mask_mod.<locals>.batch_offset_mask_mod¥   s2   ø€ Ý œ?¨1¨cÑ2Ô2�Ø!×*Ò*¨1¨x©<¸¸EÀ6ÑJÔJÐJr   r   )rU   rV   rF   s   ` €r   Úcreate_mask_modz*_split_block_mask.<locals>.create_mask_mod¤   s0   øø€ ðKð Kð Kð Kð Kð Kð )Ð(r   )Úkv_num_blocksÚ
kv_indicesÚfull_kv_num_blocksÚfull_kv_indicesÚ
BLOCK_SIZErO   Úseq_lengths)rX   ÚsizeÚAssertionErrorrM   Útensor_splitrY   rZ   r[   ÚrangeÚappendr	   Úfrom_kv_blocksr\   r]   )rF   rG   Ú	batch_dimÚkv_num_blocks_chunksÚkv_indices_chunksÚfull_kv_num_blocks_chunksÚfull_kv_indices_chunksÚchunk_block_masksÚbatch_offsetÚ	chunk_idxrW   s   `          r   Ú_split_block_maskrl   y   sÔ  ø€ ð Ô×$Ò$ QÑ'Ô'¨1Ò,Ð,Øˆ|˜jÑ(Ð(àÔ#×(Ò(¨Ñ+Ô+¨zÒ9Ð9ÝØIñ
ô 
ð 	
ð €IÝ Ô-ØÔ  *¨iñô Ðõ Ô*¨:Ô+@À*ÈiÑXÔXÐð Ô(Ð4õ 	Ô˜:Ô8¸*ÀiÑPÔPÐPàˆV�jÑ ð ð Ô%Ð1õ 	Ô˜:Ô5°zÀ9ÑMÔMÐMàˆV�jÑ ð ð ÐØ€LÝ˜:Ñ&Ô&ð @ð @ˆ	ð	)ð 	)ð 	)ð 	)ð 	)ð 	× Ò ÝÔ$Ø2°9Ô=Ø,¨YÔ7Ø#<¸YÔ#GØ 6°yÔ AØ%Ô0Ø(˜¨Ñ6Ô6Ø&Ô2ðñ ô ñ
	
ô 
	
ð 
	
ð 	Ð,¨YÔ7×<Ò<¸QÑ?Ô?Ñ?ˆˆØÐr   ÚtensorÚspecc                 ó®  ‡‡‡‡‡— |                       ‰j        ¦  «        |k    s+t          d|                       ‰j        ¦  «        › d�¦  «        ‚t          | t          ¦  «        }|rø| j        Š| j        Š|                      ¦   «         }t          j	        ||‰j        ¦  «        }| j
        ‰j                 }t          ||¦  «        \  }}|                      ¦   «         Šg }	t          |¦  «        D ]m\  }
}t          | j
        ¦  «        }||
|k     rdndz   |‰j        <   |	                     t	          j        |‰‰t          j        |¦  «        ‰d¬¦  «        ¦  «         Œn|	}nt          j	        | |‰j        ¦  «        }| j        r | j        r|D ]}|                     ¦   «          Œt*          s|S dt          j        dt          j        d	t.          t          j        d
f         fˆfd„}|r_| j        Š| j        Š| j
        Š|                      ¦   «         Š ||                      ¦   «         gd„ |D ¦   «         ¢R Ž }ˆˆˆˆfd„|D ¦   «         S t           || g|¢R Ž ¦  «        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_chunksrJ   r   F©ÚshapeÚstrideÚ	run_checkÚorigÚchunksrH   .c                 ód  •— g }d}|D ]š}t          j        | ¦  «        }||                     ‰j        ¦  «        z   }t	          d ¦  «        g|j        z  }t	          ||¦  «        |‰j        <   |||<   |                     |¦  «         ||                     ‰j        ¦  «        z  }Œ›t          |¦  «        S )Nr   )rM   Ú
zeros_liker^   r&   ÚsliceÚndimrb   rA   )	rt   ru   ÚexpandedrU   ÚchunkÚnew_valÚupperÚslicesrn   s	           €r   Ú_expand_chunksz%_split_tensor.<locals>._expand_chunksû   s±   ø€ ð ˆØˆØð 	.ð 	.ˆEÝÔ& tÑ,Ô,ˆGØ˜%Ÿ*š* T¤^Ñ4Ô4Ñ4ˆEÝ#(¨¡;¤; -°'´,Ñ">ˆFÝ%*¨3°Ñ%6Ô%6ˆF�4”>Ñ"Ø#ˆG�F‰OØ�OŠO˜GÑ$Ô$Ð$Ø�5—:’:˜dœnÑ-Ô-Ñ-ˆCˆCÝ�X‰ŒÐr   c              3   ód   K  — | ]+}t          t          |¦  «                             ¦   «         V — Œ,d S r   )r   r   Úto_local)Ú.0Úcs     r   ú	<genexpr>z _split_tensor.<locals>.<genexpr>  s8   è è € ÐAÐA¨a�d•7˜AÑÔ×'Ò'Ñ)Ô)ÐAÐAÐAÐAÐAÐAr   c           
      óD   •— g | ]}t          j        |‰‰‰‰d ¬¦  «        ‘ŒS )Frp   )r   Ú
from_local)r‚   ÚtÚglobal_shapeÚglobal_strideÚmeshÚ
placementss     €€€€r   ú
<listcomp>z!_split_tensor.<locals>.<listcomp>  sO   ø€ ð 

ð 

ð 

ð õ ÔØØØØ"Ø$Øðñ ô ð

ð 

ð 

r   )r^   r&   r_   Ú
isinstancer   r‹   Údevice_meshr�   rM   r`   rq   Údivmodrr   Ú	enumerateÚlistrb   r†   ÚSizeÚrequires_gradÚis_leafÚretain_gradÚ_debug_mask_minibatchesÚTensorrA   )rm   rn   rG   Ú_is_dtensorÚlocal_tensorÚlocal_chunksÚglobal_split_sizeÚquotientÚ	remainderÚchunk_tensors_listÚiÚlocal_chunkÚchunk_shapeÚchunk_tensorsr{   r   Úlocal_expandedrˆ   r‰   rŠ   r‹   s    `               @@@@r   Ú_split_tensorr¤   º   sà  øøøøø€ ð  �;Š;�t”~Ñ&Ô&¨*Ò4Ð4ÝØS˜6Ÿ;š; t¤~Ñ6Ô6ÐSÐSÐSñ
ô 
ð 	
õ ˜V¥WÑ-Ô-€Kàð Oð Ô&ˆ
ØÔ!ˆØ—’Ñ(Ô(ˆÝÔ)¨,¸
ÀDÄNÑSÔSˆà"œL¨¬Ô8ÐÝ$Ð%6¸
ÑCÔCÑˆ�)ØŸš™œˆØ13ÐÝ'¨Ñ5Ô5ð 	ð 	‰NˆAˆ{Ý˜vœ|Ñ,Ô,ˆKØ*2¸1¸yº=¸=°a°aÈaÑ*PˆK˜œÑ'Ø×%Ò%ÝÔ"ØØØÝœ* [Ñ1Ô1Ø(Ø#ðñ ô ñ	ô 	ð 	ð 	ð 1CˆˆåÔ*¨6°:¸t¼~ÑNÔNˆð
 Ôð   ¤ð  Ø"ð 	 ð 	 ˆEØ×ÒÑÔÐÐå"ð ØÐðÝŒlðÝ%*¤\ðå	�uŒ|˜SÐ Ô	!ðð ð ð ð ð ð ð <ØÔ&ˆ
ØÔ!ˆØ”|ˆØŸš™œˆØ'˜Ø�OŠOÑÔð
àAÐA°=ÐAÑAÔAð
ð 
ð 
ˆð

ð 

ð 

ð 

ð 

ð 

ð 

ð $ð

ñ 

ô 

ð 
	
õ �N�N 6Ð:¨MÐ:Ð:Ð:Ñ;Ô;Ð;r   c           	      óö  ‡— | sd„ t          |¦  «        D ¦   «         S t          | ¦  «        t          |¦  «        k    sSt          dt          |                      ¦   «         ¦  «        › dt          |                     ¦   «         ¦  «        › �¦  «        ‚|€t          d¦  «        ‚t          | d„ ¬¦  «        \  }Št          |d„ ¬¦  «        \  }}g }t          ||d	¬
¦  «        D �]‘\  }}|t          u st          |t          ¦  «        r| 	                    |¦  «         Œ:t          |t          j        ¦  «        rbt          |t          ¦  «        st          dt          |¦  «        › �¦  «        ‚| 	                    |                     |j        ¦  «        ¦  «         Œ¶t          |t           ¦  «        r²t          |t          ¦  «        st          dt          |¦  «        › �¦  «        ‚|j        dk    st          d¦  «        ‚|j                             d¦  «        dk    r| 	                    |¦  «         �ŒN| 	                    |j                             d¦  «        ¦  «         �Œ}t%          d|› d|› d�¦  «        ‚t'          g |¢|‘R Ž }	d„ t          |	¦  «        D ¦   «         }
t          ||d	¬
¦  «        D ]Á\  }}g }|t          u st          |t          ¦  «        r|g|	z  }nht          |t          j        ¦  «        rt)          |||	¦  «        }n<t          |t           ¦  «        rt+          ||	¦  «        }nt%          d|› d|› d�¦  «        ‚t          |
|d	¬
¦  «        D ]\  }}| 	                    |¦  «         ŒŒÂˆfd„|
D ¦   «         S )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
    c                 ó   — g | ]}i ‘ŒS r   r   ©r‚   Ú_s     r   rŒ   z'_shard_dict_of_args.<locals>.<listcomp>5  s   € Ð.Ð.Ð.�q�Ð.Ð.Ð.r   zargs_dict.keys() = z args_chunk_spec.keys() = Nz.args_chunk_spec should have been set by callerc                 ó,   — t          | t          ¦  «        S r   ©r�   r	   ©Úxs    r   r6   z%_shard_dict_of_args.<locals>.<lambda>@  s   € ¥Z°µ9Ñ%=Ô%=€ r   ©r”   c                 ó,   — t          | t          ¦  «        S r   rª   r«   s    r   r6   z%_shard_dict_of_args.<locals>.<lambda>C  s   € ­:°a½Ñ+CÔ+C€ r   T©ÚstrictzExpected TensorChunkSpec, got r   z#BlockMask only supports split_dim=0rJ   zUnsupported chunk spec: z and value: z combination.c                 ó   — g | ]}g ‘ŒS r   r   r§   s     r   rŒ   z'_shard_dict_of_args.<locals>.<listcomp>a  s   € Ð$JÐ$JÐ$J¨A RÐ$JÐ$JÐ$Jr   c                 ó0   •— g | ]}t          |‰¦  «        ‘ŒS r   )r   )r‚   Ú_flat_split_resultÚ	tree_specs     €r   rŒ   z'_shard_dict_of_args.<locals>.<listcomp>t  s4   ø€ ð ð ð àõ 	Ð)¨9Ñ5Ô5ðð ð r   )ra   Úlenr_   r‘   Úkeysr
   ÚziprE   r�   rb   rM   r—   r   Útyper^   r&   r	   rX   Ú
ValueErrorÚminr¤   rl   )Ú	args_dictr8   rG   ÚvaluesÚchunk_specsr¨   Úsplit_sizesÚvrn   Úresult_num_chunksÚflat_split_resultsÚv_splitsr³   Ú_v_splitr´   s                 @r   Ú_shard_dict_of_argsrÄ   "  s  ø€ ð$ ð /Ø.Ð.�E *Ñ-Ô-Ð.Ñ.Ô.Ð.åˆy‰>Œ>�S Ñ1Ô1Ò1Ð1ÝðG¥$ y§~¢~Ñ'7Ô'7Ñ"8Ô"8ð Gð GÝ(,¨_×-AÒ-AÑ-CÔ-CÑ(DÔ(DðGð Gñ
ô 
ð 	
ð ÐÝÐMÑNÔNÐNå$ØÐ=Ð=ðñ ô Ñ€FˆIõ "ØÐ!CÐ!Cðñ ô �N€K�ð
 €KÝ�v˜{°4Ð8Ñ8Ô8ð ñ ‰ˆˆ4ð •:ÐÐ¥¨Dµ*Ñ!=Ô!=ÐØ×Ò˜zÑ*Ô*Ð*Ð*Ý˜�5œ<Ñ(Ô(ð 	Ý˜d¥OÑ4Ô4ð TÝ$Ð%RÅdÈ4ÁjÄjÐ%RÐ%RÑSÔSÐSØ×Ò˜qŸvšv d¤nÑ5Ô5Ñ6Ô6Ð6Ð6Ý˜�9Ñ%Ô%ð 	Ý˜d¥OÑ4Ô4ð TÝ$Ð%RÅdÈ4ÁjÄjÐ%RÐ%RÑSÔSÐSØ”> QÒ&Ð&Ý$Ð%JÑKÔKÐKàŒ×#Ò# AÑ&Ô&¨!Ò+Ð+Ø×"Ò" :Ñ.Ô.Ð.Ñ.à×"Ò" 1¤?×#7Ò#7¸Ñ#:Ô#:Ñ;Ô;Ð;Ñ;åØM¨4ÐMÐM¸QÐMÐMÐMñô ð õ Ð5˜[Ð5¨*Ð5Ð5Ð5Ðà$JÐ$JµÐ7HÑ1IÔ1IÐ$JÑ$JÔ$JÐÝ�v˜{°4Ð8Ñ8Ô8ð 0ð 0‰ˆˆ4Ø"$ˆØ•:ÐÐ¥¨Dµ*Ñ!=Ô!=ÐØ�sÐ.Ñ.ˆHˆHÝ˜�5œ<Ñ(Ô(ð 	Ý$ Q¨Ð.?Ñ@Ô@ˆHˆHÝ˜�9Ñ%Ô%ð 	Ý(¨Ð,=Ñ>Ô>ˆHˆHåØM¨4ÐMÐM¸QÐMÐMÐMñô ð õ -0Ø °ð-
ñ -
ô -
ð 	0ð 	0Ñ(Ð ð ×%Ò% hÑ/Ô/Ð/Ð/ð	0ð
ð ð ð à"4ðñ ô ð r   Úargs.Úkwargsru   r8   r<   c                 óº  — |€i }d„ }|€t          || d„ ¬¦  «        }|€t          ||d„ ¬¦  «        }t          t          t          | ¦  «        ¦  «        t          t          |¦  «        ¦  «        |¦  «        }t	          |¦  «        }t          |||¦  «        }t	          |¦  «        |k     rTt	          |¦  «        }t          t          t          | ¦  «        ¦  «        t          t          |¦  «        ¦  «        |¦  «        }t	          |¦  «        t	          |¦  «        k    r/t          dt	          |¦  «        › dt	          |¦  «        › �¦  «        ‚d„ |D ¦   «         }	|	|fS )	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
    Nc                 óŠ   — t          | t          j        t          z  ¦  «        rt	          t
          ¦  «        S t          ¦   «         S r   )r�   rM   r—   r	   r   ÚDEFAULT_CHUNK_DIMrE   ©r¿   s    r   Údefault_specz3split_args_kwargs_into_chunks.<locals>.default_spec·  s4   € Ý�a�œ­	Ñ1Ñ2Ô2ð 	 Ý"Õ#4Ñ5Ô5Ð5å‘<”<Ðr   c                 ó,   — t          | t          ¦  «        S r   rª   rÊ   s    r   r6   z/split_args_kwargs_into_chunks.<locals>.<lambda>¿  s   € µ*¸QÅ	Ñ2JÔ2J€ r   r­   c                 ó,   — t          | t          ¦  «        S r   rª   rÊ   s    r   r6   z/split_args_kwargs_into_chunks.<locals>.<lambda>Ä  s   € µJ¸qÅ)Ñ4LÔ4L€ r   z;args and kwargs are split into different number of chunks: z, c           
      óz   ‡— g | ]7Št          ˆfd „t          t          ‰¦  «        ¦  «        D ¦   «         ¦  «        ‘Œ8S )c              3   ó(   •K  — | ]}‰|         V — Œd S r   r   )r‚   rŸ   Ú
chunk_argss     €r   r„   z;split_args_kwargs_into_chunks.<locals>.<listcomp>.<genexpr>æ  s'   øè è € Ð<Ð< ˆj˜ŒmÐ<Ð<Ð<Ð<Ð<Ð<r   )rA   ra   rµ   )r‚   rÐ   s    @r   rŒ   z1split_args_kwargs_into_chunks.<locals>.<listcomp>å  sT   ø€ ð ð ð àõ 	Ð<Ð<Ð<Ð<¥U­3¨z©?¬?Ñ%;Ô%;Ð<Ñ<Ô<Ñ<Ô<ðð ð r   )r   rÄ   rB   r�   rµ   ÚRuntimeError)
rÅ   rÆ   ru   r8   r<   rË   Úargs_split_dictÚreal_num_chunksÚkwargs_splitÚ
args_splits
             r   r   r   z  s¯  € ðp €~Øˆð ð  ð  ð ÐÝ"Ø˜$Ð(JÐ(Jð
ñ 
ô 
ˆð Ð Ý$Ø˜&Ð*LÐ*Lð
ñ 
ô 
Ðõ *Ý�Y�t‰_Œ_ÑÔÝ�Y�Ñ'Ô'Ñ(Ô(Øñô €Oõ
 ˜/Ñ*Ô*€Oå&ØØØñô €Lõ ˆ<ÑÔ˜?Ò*Ð*õ ˜lÑ+Ô+ˆå-Ý•˜4‘”Ñ!Ô!Ý•˜?Ñ+Ô+Ñ,Ô,Øñ
ô 
ˆõ ˆ?ÑÔ�s <Ñ0Ô0Ò0Ð0Ýð;Ý�?Ñ#Ô#ð;ð ;Ý'*¨<Ñ'8Ô'8ð;ð ;ñ
ô 
ð 	
ð
ð à)ðñ ô €Jð
 �|Ð#Ð#r   c                 óÞ	  ‡‡ ‡!— |�t          |¦  «        \  }}n=t          | d         ¦  «        \  }}t          t          ¦  «        gt          |¦  «        z  }g Š!| D ]^}t          |¦  «        \  }}t          |¦  «        t          |¦  «        k    rt	          d|› d|› �¦  «        ‚‰!                     |¦  «         Œ_g }t          |¦  «        D �]\  Š Št          ‰t          ¦  «        �rˆ ˆ!fd„t          t          ‰!¦  «        ¦  «        D ¦   «         }	t          �rQ|	d         j
        }
|	dd…         D ]'}|j
        |
k    s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 ]s\  }}||                     ‰j        ¦  «        z   }t%          ddd¦  «        g|j        z  }t%          ||¦  «        |‰j        <   ||         }|                     |¦  «         |}Œtn|	}d„ |D ¦   «         }t)          |¦  «        �r=t+          |¦  «        st          d¦  «        ‚|d         j        }t          |dd…         d¦  «        D ]-\  }}|j        |k    rt          d|› d|› d|j        › �¦  «        ‚Œ.|d         j        }t          j        d„ |D ¦   «         ‰j        ¬¦  «        }t3          |d         j
        ¦  «        }t5          ˆfd„|D ¦   «         ¦  «        |‰j        <   t          j        |¦  «        }|d                              ¦   «         }|                     t;          j        |||||d¬¦  «        ¦  «         �Œù|                     t          j        |‰j        ¬¦  «        ¦  «         �Œ)t          ‰t>          ¦  «        r_‰j         }t          t          ‰!¦  «        ¦  «        D ]$}‰ !                    |‰!|         ‰          ¦  «        }Œ%|                     |¦  «         �Œ�‰!d         ‰          }t          dt          ‰!¦  «        ¦  «        D ]5}‰!|         ‰          |k    s!t          d|› d‰!|         ‰          › �¦  «        ‚Œ6|                     |¦  «         �ŒtE          ||¦  «        S )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 c                 ó,   •— g | ]}‰|         ‰         ‘ŒS r   r   )r‚   rk   Úarg_idxÚchunks_flatteneds     €€r   rŒ   z merge_chunks.<locals>.<listcomp>3  s3   ø€ ð ð ð àð ! Ô+¨GÔ4ðð ð r   rJ   zExpected shape z, got ÚdeviceÚmeta)Úsectionsr5   z6Expected len(partial_values) == len(meta_chunks), got z != Tr¯   c                 ó8   — g | ]}t          |t          ¦  «        ‘ŒS r   )r�   r   ©r‚   r¿   s     r   rŒ   z merge_chunks.<locals>.<listcomp>^  s"   € ÐKÐKÐK¸�Z¨­7Ñ3Ô3ÐKÐKÐKr   zRmerge_chunks: expected all values to be DTensors or none to be DTensors, got a mixz*merge_chunks: placement mismatch at chunk z: expected c                 ó6   — g | ]}|                      ¦   «         ‘ŒS r   )r�   rÞ   s     r   rŒ   z merge_chunks.<locals>.<listcomp>o  s    € Ð9Ð9Ð9 a�Q—Z’Z‘\”\Ð9Ð9Ð9r   r4   c              3   ó<   •K  — | ]}|j         ‰j                 V — Œd S r   )rq   r&   )r‚   r¿   Úargs     €r   r„   zmerge_chunks.<locals>.<genexpr>s  s=   øè è € ð /ð /Ø/0�A”G˜CœMÔ*ð/ð /ð /ð /ð /ð /r   Frp   z	Expected )#r
   r   rÉ   rµ   r¹   rb   r�   r�   ra   r–   rq   r_   rM   r`   Úemptyr&   r·   r^   rx   ry   ÚanyÚallr‹   rŽ   Úcatr‘   Úsumr’   rr   r   r†   r   r   r   r   )"ru   Ú
chunk_specÚspec_flattenedÚflatten_specÚchunk0_flatr{   Úchunk_flattenedr¨   Úargs_flattenedÚpartial_valuesÚoverall_shapeÚvalÚmeta_chunksÚvalues_to_catÚchunk_start_idxÚpartial_valueÚ
meta_chunkÚchunk_end_idxÚslice_indicesÚslicedÚdtensor_flagsr‹   rŸ   r¿   rŠ   Ú	local_catÚ	cat_shapeÚ
cat_strideÚreduced_valrk   Úvaluerá   rØ   rÙ   s"                                  @@@r   r   r   í  s  øøø€ ðZ ÐÝ'3°JÑ'?Ô'?Ñ$ˆ˜˜õ %1°¸´Ñ$;Ô$;Ñ!ˆ�\Ý)Õ*;Ñ<Ô<Ð=ÅÀKÑ@PÔ@PÑPˆð Ðàð 1ð 1ˆÝ)¨%Ñ0Ô0Ñˆ˜ÝˆÑÔ¥3 ~Ñ#6Ô#6Ò6Ð6ÝÐS eÐSÐSÀzÐSÐSÑTÔTÐTà×Ò Ñ0Ô0Ð0Ð0ð
 !#€NÝ! .Ñ1Ô1ð c)ñ c)‰ˆ�Ý�c�?Ñ+Ô+ñ b	)ðð ð ð ð å!&¥sÐ+;Ñ'<Ô'<Ñ!=Ô!=ðñ ô ˆNõ
 'ñ "/à .¨qÔ 1Ô 7�Ø)¨!¨"¨"Ô-ð ð �CØœ9¨Ò5Ð5Ý,ØN¨mÐNÐNÀ3Ä9ÐNÐNñô ð ð 6õ $Ô0Ý”K Ð>°vÐ>Ð>Ý  Ñ0Ô0Øœðñ ô �ð !#�Ø"#�Ý˜>Ñ*Ô*­c°+Ñ.>Ô.>Ò>Ð>Ý(Ø|ÕQTÐUcÑQdÔQdÐ|Ð|ÕjmÐnyÑjzÔjzÐ|Ð|ñô ð õ 25Ø" K¸ð2ñ 2ô 2ð 
4ð 
4Ñ-�M :ð %4°j·o²oÀcÄmÑ6TÔ6TÑ$T�Må%*¨4°°tÑ%<Ô%<Ð$=ÀÔ@RÑ$R�MÝ38¸È-Ñ3XÔ3X�M #¤-Ñ0Ø*¨=Ô9�FØ!×(Ò(¨Ñ0Ô0Ð0à&3�O�Oð
4ð !/�ð LÐK¸]ÐKÑKÔKˆMÝ�=Ñ!Ô!ñ $SÝ˜=Ñ)Ô)ð Ý(ð9ñô ð ð
 +¨1Ô-Ô8�
Ý% m°A°B°BÔ&7¸Ñ;Ô;ð ð ‘D�A�qØ”| zÒ1Ð1Ý,ðIÈð Ið IØ(2ðIð IØ:;¼,ðIð Iñô ð ð 2ð
 % QÔ'Ô3�Ý!œIØ9Ð9¨=Ð9Ñ9Ô9Øœðñ ô �	õ ! ¨qÔ!1Ô!7Ñ8Ô8�	Ý+.ð /ð /ð /ð /Ø4Að/ñ /ô /ñ ,ô ,�	˜#œ-Ñ(õ "œJ yÑ1Ô1�	Ø*¨1Ô-×4Ò4Ñ6Ô6�
Ø×%Ò%ÝÔ&Ø!ØØ"Ø'Ø)Ø"'ðñ ô ñ	ô 	ð 	ñ 	ð ×%Ò%¥e¤i°À3Ä=Ð&QÑ&QÔ&QÑRÔRÐRÑRÝ˜�^Ñ,Ô,ð 	)Øœ.ˆKå"¥3Ð'7Ñ#8Ô#8Ñ9Ô9ð ð �	Ø!ŸmšmØÐ!1°)Ô!<¸WÔ!Eñô ��ð ×!Ò! +Ñ.Ô.Ð.Ñ.à$ QÔ'¨Ô0ˆEÝ" 1¥cÐ*:Ñ&;Ô&;Ñ<Ô<ð ð �	Ø'¨	Ô2°7Ô;¸uÒDÐDÝ(ØW EÐWÐWÐ1AÀ)Ô1LÈWÔ1UÐWÐWñô ð ð Eð ×!Ò! %Ñ(Ô(Ð(Ñ(õ ˜.¨,Ñ7Ô7Ð7r   )NN)*ÚloggingÚoperatorÚcollections.abcr   Útypingr   r   rM   Útorch.distributed.tensorr   Útorch.fx.noder   Ú!torch.nn.attention.flex_attentionr	   Útorch.utils._pytreer
   r   r   Ú__all__Ú	getLoggerr   Úloggerr–   r   r    rm   ÚaddÚsum_reducerrÉ   r   rE   r>   r‘   rl   r—   r¤   rÄ   rA   rB   rC   r   r   r   r   r   ú<module>r     s	  ðð €€€Ø €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø Ð Ð Ð Ð Ð Ð Ð à €€€Ø ,Ð ,Ð ,Ð ,Ð ,Ð ,Ø 'Ð 'Ð 'Ð 'Ð 'Ð 'Ø 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø FÐ FÐ FÐ FÐ FÐ FÐ FÐ FÐ FÐ Fðð ð €ð 
ˆÔ	˜8Ñ	$Ô	$€ðð
  Ð ð#ð #ð #ð #ð #ñ #ô #ð #ð$	ð 	ð 	ð 	ð 	�>ñ 	ô 	ð 	ð ˆl˜<˜5œ<¨Ñ,Ô,¨h¬lÑ;Ô;€ð Ð ð5!ð 5!ð 5!ð 5!ð 5!ñ 5!ô 5!ð 5!ðr	ð 	ð 	ð 	ð 	ñ 	ô 	ð 	ð>Øð>àð>ð 
ˆ)„_ð>ð >ð >ð >ðBe<ØŒLðe<à
ðe<ð ðe<ð ˆeŒlÔð	e<ð e<ð e<ð e<ðPUð Uð Uðx ;?Ø;?ðp$ð p$Ø
��S�Œ/ðp$à��c�ŒN˜TÑ!ðp$ð ðp$ð ˜?¨CÐ/Ô0°4Ñ7ð	p$ð
 ˜C Ð0Ô1°DÑ8ðp$ð ˆ4�Œ;˜˜Tœ
Ð"Ô#ðp$ð p$ð p$ð p$ðfj8Ø�ŒIðj8ð j8ð j8ð j8ð j8ð j8r   