§
    ‚Štj7B  ã                  óô   — d dl mZ d dlZd dlmZ ddlmZ er
d dlZd dlm	Z	  e¦   «         r:d dlZd dlm	Z	 d dl
mZ d dlmZ  eed	¦  «        s eed
¦  «        rej        e_         G d„ d¦  «        Zdd„ZdS )é    )ÚannotationsN)ÚTYPE_CHECKINGé   )Úis_torch_available)ÚDTensor)Ú%compute_local_shape_and_global_offset)ÚShardÚlocal_shard_size_and_offsetÚ_local_shard_size_and_offsetc                  óN   — e Zd ZdZdd„Z	 d d!d„Zd"d„Zd#d„Zd$d„Zd%d„Z	d&d„Z
dS )'ÚDtensorShardOperationuú
  Shard-on-read: slice a full disk tensor down to this rank's local
    DTensor shard, for any combination of placements on a 1-D or n-D mesh.  It's on
    read because instructions are made so the cpu only fetches on disk the parts we want.

    Placements primer
    -----------------
    Each mesh dim carries one placement describing how it slices the tensor:

    | Placement                | Local data on each rank of the mesh dim         |
    |--------------------------|-------------------------------------------------|
    | Replicate                | full tensor (no slicing)                        |
    | Shard(d)                 | contiguous chunk of dim d (rows r*c .. (r+1)*c) |
    | _StridedShard(d, sf=N)   | one chunk from each of N groups along dim d,   |
    |                          | concatenated together (interleaved layout)      |

    Different scenarios of different placements
    ------------------------------------------
    Placement tuples are ordered outermost-first; for a 2-D (fsdp, tp) mesh
    the tuple is (fsdp_placement, tp_placement).

    | Scenario                                          | Placements                                  |
    |---------------------------------------------------|---------------------------------------------|
    | TP-only, non-fused (e.g. q_proj/k_proj/v_proj)    | [Shard(d)]                                   |
    | TP-only, fused gate/up                            | [_StridedShard(d, sf=2)]                     |
    | TP + FSDP, same tensor dim (contiguous TP case)   | [Shard(d), Shard(d)]                         |
    | TP + FSDP, same tensor dim (fused/interleaved TP) | [_StridedShard(d, sf=tp_size), Shard(d)]     |
    | TP + FSDP, different dims                         | [Shard(d1), Shard(d2)]                       |

    Loading (this class)
    --------------------
    During from_pretrained, each rank looks up every tensor key, but does not load the weight bytes yet (safetensors get_slice).
    When a weight is needed, shard_tensor indexes that slice to read only this rank's local shard through Dtensor logic which provides
    placement logics

    Depending on how the checkpoint was saved, the weight loader either gives us one big tensor or many small ones:

    1. One stacked tensor: the checkpoint has one key with all experts together, shaped [num_experts, in, out].
       The weight loader passes it straight through; we slice out this rank's piece.

    2. One tensor per expert:  the checkpoint has a separate key for each expert (expert 0, expert 1, â€¦), each shaped [in, out].
       The weight loader feeds them in one at a time. If this rank doesn't own a given expert,
       we skip it. Later, `MergeModulelist` will stack the owned expert we kept to create the rank's local shard
    Úparamr   c                óâ   — |j         | _         t          |j        ¦  «        | _        |j        | _        t          |j        | j         | j        ¦  «        \  }}|d         | _        |d         | _        d S ©Nr   )	Údevice_meshÚtupleÚ
placementsÚndimÚ
param_ndimr   ÚshapeÚ_axis0_offsetÚ_axis0_local_size)Úselfr   Úlocal_shapeÚoffsetss       úe/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/distributed/sharding_utils.pyÚ__init__zDtensorShardOperation.__init__R   sh   € Ø Ô,ˆÔÝ Ô 0Ñ1Ô1ˆŒØœ*ˆŒÝDÀUÄ[ÐRVÔRbÐdhÔdsÑtÔtÑˆ�Wð % QœZˆÔØ!,¨Q¤ˆÔÐÐó    NÚsourceútorch.TensorÚ
tensor_idxú
int | NoneÚreturnútorch.Tensor | Nonec                óZ  ‡ — t          |t          j        ¦  «        rt          |j        ¦  «        n|                     ¦   «         }d„ t          ‰ j        ¦  «        D ¦   «         }|�€Ï|s|d                              ||¬¦  «        S d„ |D ¦   «         }|D ]z\  }}	‰  	                    |¦  «        }
|
 
                    ¦   «         |
                     ¦   «         }}‰                      |	j        ¦  «        }||                              |	||f¦  «         Œ{d„ |D ¦   «         }t          |¦  «        D ]d\  }}||         }|D ]O\  }	}}|	                     ¦   «         r‰                      |||¦  «        }Œ2‰                      ||||	j        ¦  «        }ŒP|||<   Œet'          d„ |D ¦   «         ¦  «        }|r‰                      ||||¦  «        S g }|D ]E}t+          |¦  «        dk    r|d         nd	\  }}|                     t-          ||¦  «        ¦  «         ŒF|t/          |¦  «                                      ||¬¦  «        S ˆ fd
„|D ¦   «         }t'          d„ |D ¦   «         ¦  «        }‰ j        |cxk    o‰ j        ‰ j        z   k     nc }|r|sdS d„ |D ¦   «         }|D ]k\  }}}|dk    r_|dz
  }‰  	                    |¦  «        }
|
 
                    ¦   «         |
                     ¦   «         }}||                              ||f¦  «         Œld„ |D ¦   «         }t          |¦  «        D ]1\  }}||         }|D ]\  }}‰                      |||¦  «        }Œ|||<   Œ2g }|D ]4}|r|d         nd	\  }}|                     t-          ||¦  «        ¦  «         Œ5|t/          |¦  «                                      ||¬¦  «        S )aù  Return this rank's local shard of a checkpoint tensor.

        Two layouts (example param shape [N, in, out]):

        - tensor_idx is None: one stacked [N, in, out] tensor;
          slice every sharded dim (including axis 0).
        - tensor_idx given: one [in, out] tensor per expert;
          return None if this rank does not own that expert, else slice
          inner dims only. Surviving pieces are stacked by MergeModulelist
          into this rank's local [n_local, in, out] shard.
        c                ó<   — g | ]\  }}t          |d ¦  «        ¯||f‘ŒS ©Údim)Úhasattr)Ú.0Úmesh_dimÚ	placements      r   ú
<listcomp>z6DtensorShardOperation.shard_tensor.<locals>.<listcomp>k   sC   € ð 
ð 
ð 
Ù&9 h°	Õ[bÐclÐnsÑ[tÔ[tð
Ø�yÐ!ð
ð 
ð 
r   N.©ÚdeviceÚdtypec                ó   — g | ]}g ‘ŒS © r2   ©r*   Ú_s     r   r-   z6DtensorShardOperation.shard_tensor.<locals>.<listcomp>w   s   € Ð!;Ð!;Ð!;¨ "Ð!;Ð!;Ð!;r   c                ó   — g | ]}d |fg‘ŒS ©r   r2   ©r*   Úsizes     r   r-   z6DtensorShardOperation.shard_tensor.<locals>.<listcomp>   s   € ÐEÐEÐE° ! T  ÐEÐEÐEr   c              3  óF   K  — | ]\  }}|                      ¦   «          V — Œd S ©N)Úis_shard)r*   r4   r,   s      r   ú	<genexpr>z5DtensorShardOperation.shard_tensor.<locals>.<genexpr>‰   s5   è è € Ð#`Ð#`ÁÀÀI¨	×(:Ò(:Ñ(<Ô(<Ð$<Ð#`Ð#`Ð#`Ð#`Ð#`Ð#`r   r   )r   r   c                óP   •— g | ]"\  }}||‰                      |j        ¦  «        f‘Œ#S r2   )Ú_normalize_param_dimr(   )r*   r+   r,   r   s      €r   r-   z6DtensorShardOperation.shard_tensor.<locals>.<listcomp>™   sC   ø€ ð %
ð %
ð %
ÙPcÐPXÐZcˆX�y $×";Ò";¸I¼MÑ"JÔ"JÐKð%
ð %
ð %
r   c              3  ó*   K  — | ]\  }}}|d k    V — ŒdS )r   Nr2   )r*   r4   Ú	param_dims      r   r<   z5DtensorShardOperation.shard_tensor.<locals>.<genexpr>ž   s,   è è € Ð^Ð^±°°A°y˜i¨1šnÐ^Ð^Ð^Ð^Ð^Ð^r   c                ó   — g | ]}g ‘ŒS r2   r2   r3   s     r   r-   z6DtensorShardOperation.shard_tensor.<locals>.<listcomp>©   s   € Ð$>Ð$>Ð$>¨A RÐ$>Ð$>Ð$>r   é   c                ó   — g | ]}d |fg‘ŒS r6   r2   r7   s     r   r-   z6DtensorShardOperation.shard_tensor.<locals>.<listcomp>±   s   € Ð"HÐ"HÐ"H°4 Q¨ I ;Ð"HÐ"HÐ"Hr   )Ú
isinstanceÚtorchÚTensorÚlistr   Ú	get_shapeÚ	enumerater   ÚtoÚ_get_sub_meshÚget_local_rankr8   r>   r(   Úappendr;   Ú_compute_contiguous_sliceÚ_compute_strided_sliceÚsplit_factorÚanyÚ_slice_and_catÚlenÚslicer   r   r   )r   r   r!   r/   r0   Úsource_shapeÚdim_placementsÚplanned_ops_by_dimr+   r,   Úsub_meshÚrankÚ
world_sizeÚdim_idxÚintervals_by_dimÚplanned_opsÚ	intervalsÚhas_strided_shardÚslice_partsÚstartÚendÚnormalized_dim_placementsÚhas_axis0_shardÚowns_tensor_idxÚplanned_ops_by_source_dimr4   r@   Ú
source_dimÚintervals_by_source_dims   `                            r   Úshard_tensorz"DtensorShardOperation.shard_tensor\   s•  ø€ õ .8¸ÅÄÑ-MÔ-MÐe•t˜FœLÑ)Ô)Ð)ÐSY×ScÒScÑSeÔSeˆð
ð 
Ý=FÀtÄÑ=WÔ=Wð
ñ 
ô 
ˆð
 ÑØ!ð BØ˜c”{—~’~¨V¸5�~ÑAÔAÐAð
 "<Ð!;¨lÐ!;Ñ!;Ô!;ÐØ'5ð Rð RÑ#�˜)Ø×-Ò-¨hÑ7Ô7�Ø#+×#:Ò#:Ñ#<Ô#<¸h¿mºm¹o¼o�j�Ø×3Ò3°I´MÑBÔB�Ø" 7Ô+×2Ò2°I¸tÀZÐ3PÑQÔQÐQÐQð  FÐE¸ÐEÑEÔEÐÝ(1Ð2DÑ(EÔ(Eð 6ð 6Ñ$�˜Ø,¨WÔ5�	Ø3>ð uð uÑ/�I˜t ZØ ×)Ò)Ñ+Ô+ð uØ$(×$BÒ$BÀ9ÈdÐT^Ñ$_Ô$_˜	˜	à$(×$?Ò$?À	È4ÐQ[Ð]fÔ]sÑ$tÔ$t˜	˜	Ø,5Ð  Ñ)Ð)å #Ð#`Ð#`ÐQ_Ð#`Ñ#`Ô#`Ñ `Ô `Ðð !ð 	Qà×*Ò*¨6Ð3CÀVÈUÑSÔSÐSà �Ø!1ð :ð :�IÝ14°Y±´À!Ò1CÐ1C ¨1¤ È‘J�E˜3Ø×&Ò&¥u¨U°CÑ'8Ô'8Ñ9Ô9Ð9Ð9à�e KÑ0Ô0Ô1×4Ò4¸FÈ%Ð4ÑPÔPÐPð%
ð %
ð %
ð %
Øguð%
ñ %
ô %
Ð!õ
 Ð^Ð^ÐD]Ð^Ñ^Ô^Ñ^Ô^ˆØÔ,°
ÐhÐhÒhÐh¸TÔ=OÐRVÔRhÑ=hÒhÐhÐhÐhˆØð 	 ?ð 	Ø�4ð %?Ð$>°Ð$>Ñ$>Ô$>Ð!Ø&?ð 	Qð 	QÑ"ˆH�a˜Ø˜1Š}ˆ}Ø&¨™]�
Ø×-Ò-¨hÑ7Ô7�Ø#+×#:Ò#:Ñ#<Ô#<¸h¿mºm¹o¼o�j�Ø)¨*Ô5×<Ò<¸dÀJÐ=OÑPÔPÐPøà"HÐ"H¸<Ð"HÑ"HÔ"HÐÝ'0Ð1JÑ'KÔ'Kð 	<ð 	<Ñ#ˆJ˜Ø/°
Ô;ˆIØ$/ð Xð XÑ ��jØ ×:Ò:¸9ÀdÈJÑWÔW�	�	Ø2;Ð# JÑ/Ð/àˆØ0ð 	2ð 	2ˆIØ)2Ð>˜ 1œ˜¸‰JˆE�3Ø×Ò�u U¨CÑ0Ô0Ñ1Ô1Ð1Ð1à•e˜KÑ(Ô(Ô)×,Ò,°FÀ%Ð,ÑHÔHÐHr   r^   úlist[tuple[int, int]]rY   ÚintrZ   rP   c                ó:  — g }|D ]•\  }}t          j        ||z
  |z  ¦  «        }t          |¦  «        D ]f}	||	|z  z   }
t          |
|z   |¦  «        }||
z
  }|dk    r>t	          j        |||¦  «        \  }}|dk    r|
|z   }|                     |||z   f¦  «         ŒgŒ–|S r   )ÚmathÚceilÚrangeÚminr	   r
   rM   )r   r^   rY   rZ   rP   Úlocal_intervalsÚinterval_startÚinterval_endÚgroup_widthÚ	group_idxÚgroup_startÚ	group_endÚ	group_lenÚlocal_shard_sizeÚlocal_shard_offsetÚshard_starts                   r   rO   z,DtensorShardOperation._compute_strided_slice¿   sí   € ð ˆà,5ð 	^ð 	^Ñ(ˆN˜Låœ) \°NÑ%BÀlÑ$RÑSÔSˆKå" <Ñ0Ô0ð ^ð ^�	à,¨y¸;Ñ/FÑF�Ý ¨kÑ 9¸<ÑHÔH�	Ø%¨Ñ3�	à˜q’=�=å;@Ô;\Ø! :¨tñ<ô <Ñ8Ð$Ð&8ð (¨!Ò+Ð+à&1Ð4FÑ&F˜Ø'×.Ò.°¸[ÐK[Ñ=[Ð/\Ñ]Ô]Ð]øð^ð  Ðr   c                óô  — t          d„ |D ¦   «         ¦  «        }t          j        |||¦  «        \  }}||z   }|dk    rg S t          |¦  «        dk    r|d         \  }}	||z   ||z   fgS g }
d}|D ]0\  }}||z
  }|dk    r |
                     |||z   |f¦  «         ||z  }Œ1g }|
D ]S\  }}}t          ||¦  «        }t          ||¦  «        }||k     r'|||z
  z   }|||z
  z   }|                     ||f¦  «         ŒT|S )Nc              3  ó&   K  — | ]\  }}||z
  V — Œd S r:   r2   )r*   ra   rb   s      r   r<   zBDtensorShardOperation._compute_contiguous_slice.<locals>.<genexpr>ã   s*   è è € ÐEÐE©Z¨U°C˜S 5™[ÐEÐEÐEÐEÐEÐEr   r   rB   )Úsumr	   r
   rS   rM   Úmaxrp   )r   r^   rY   rZ   Úflat_total_lenÚlocal_flat_lenÚlocal_flat_startÚlocal_flat_endÚsource_startr4   Úflat_segmentsÚidxÚ
source_endÚinterval_lenrq   Úinterval_flat_startÚinterval_flat_endÚoverlap_flat_startÚoverlap_flat_endÚsource_overlap_startÚsource_overlap_ends                        r   rN   z/DtensorShardOperation._compute_contiguous_sliceÚ   s…  € õ ÐEÐE¸9ÐEÑEÔEÑEÔEˆÝ+0Ô+LÈ^Ð]gÐimÑ+nÔ+nÑ(ˆÐ(Ø)¨NÑ:ˆà˜QÒÐØˆIõ ˆy‰>Œ>˜QÒÐØ'¨œl‰OˆL˜!Ø!Ð$4Ñ4°lÀ^Ñ6SÐTÐUÐUð ˆØˆØ(1ð 	$ð 	$Ñ$ˆL˜*Ø%¨Ñ4ˆLØ˜aÒÐØ×$Ò$ c¨3°Ñ+=¸|Ð%LÑMÔMÐMØ�|Ñ#�øð ˆØDQð 	Sð 	SÑ@ÐÐ!2°LÝ!$Ð%8Ð:JÑ!KÔ!KÐÝ"Ð#4°nÑEÔEÐØ!Ð$4Ò4Ð4Ø'3Ð7IÐL_Ñ7_Ñ'`Ð$Ø%1Ð5EÐH[Ñ5[Ñ%\Ð"Ø×&Ò&Ð(<Ð>PÐ'QÑRÔRÐRøàÐr   úlist[list[tuple[int, int]]]r/   útorch.device | str | int | Noner0   útorch.dtype | Nonec                óâ  — d„ t          |¦  «        D ¦   «         }t          |¦  «        dk    rt          d¦  «        ‚|r|d         nd }g }t          |¦  «        D ]\\  }}	||k    r#|                     t	          d ¦  «        ¦  «         Œ.|	d         \  }
}|                     t	          |
|¦  «        ¦  «         Œ]|€*|t          |¦  «                                      ||¬¦  «        S t          |¦  «        }g }||         D ]J\  }}g |d |…         ¢t	          ||¦  «        ‘||dz   d …         ¢R }|                     ||         ¦  «         ŒKt          j        ||¬¦  «                             ||¬¦  «        S )Nc                ó>   — g | ]\  }}t          |¦  «        d k    ¯|‘ŒS )rB   )rS   )r*   r[   Údim_intervalss      r   r-   z8DtensorShardOperation._slice_and_cat.<locals>.<listcomp>  s2   € ÐtÐtÐtÑ+A¨7°MÕ]`ÐanÑ]oÔ]oÐrsÒ]sÐ]s˜wÐ]sÐ]sÐ]sr   rB   zUCurrent shard-on-read only supports disjoint ranges on a single checkpoint dimension.r   r.   r'   )	rI   rS   Ú
ValueErrorrM   rT   r   rJ   rE   Úcat)r   r   r^   r/   r0   Úmulti_interval_dimsÚ
concat_dimÚbase_slicesr[   r”   ra   rb   Úbase_slices_tupleÚinterval_tensorsrr   rs   Úinterval_slicess                    r   rR   z$DtensorShardOperation._slice_and_cat  sÂ  € ð uÐtÅYÈyÑEYÔEYÐtÑtÔtÐÝÐ"Ñ#Ô# aÒ'Ð'õ ÐtÑuÔuÐuØ/BÐLÐ(¨Ô+Ð+Èˆ
àˆÝ&/°	Ñ&:Ô&:ð 	6ð 	6Ñ"ˆG�]Ø˜*Ò$Ð$à×"Ò"¥5¨¡;¤;Ñ/Ô/Ð/Ð/ð +¨1Ô-‘
��sØ×"Ò"¥5¨°Ñ#4Ô#4Ñ5Ô5Ð5Ð5ð ÐØ�% Ñ,Ô,Ô-×0Ò0¸ÀeÐ0ÑLÔLÐLõ " +Ñ.Ô.ÐØÐØ,5°jÔ,Að 	=ð 	=Ñ(ˆN˜LðØ" ; J ;Ô/ðå�n lÑ3Ô3ðð # :°¡>Ð#3Ð#3Ô4ðð ˆOð
 ×#Ò# F¨?Ô$;Ñ<Ô<Ð<Ð<åŒyÐ)¨zÐ:Ñ:Ô:×=Ò=ÀVÐSXÐ=ÑYÔYÐYr   r+   c                ój   — | j         j        dk    r| j         S | j         | j         j        |                  S )NrB   )r   r   Úmesh_dim_names)r   r+   s     r   rK   z#DtensorShardOperation._get_sub_mesh5  s5   € ØÔÔ  AÒ%Ð%ØÔ#Ð#ØÔ Ô 0Ô ?ÀÔ IÔJÐJr   r(   c                ó&   — |dk    r|n	| j         |z   S r   )r   )r   r(   s     r   r>   z*DtensorShardOperation._normalize_param_dim:  s   € à˜Q’h�hˆsˆs D¤O°cÑ$9Ð9r   )r   r   )NNN)r   r    r!   r"   r#   r$   )
r^   rj   rY   rk   rZ   rk   rP   rk   r#   rj   )r^   rj   rY   rk   rZ   rk   r#   rj   )
r   r    r^   r�   r/   r�   r0   r‘   r#   r    )r+   rk   )r(   rk   r#   rk   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ri   rO   rN   rR   rK   r>   r2   r   r   r   r   %   sÌ   € € € € € ð*ð *ðX0ð 0ð 0ð 0ð W[ðaIð aIð aIð aIð aIðFð ð ð ð60ð 0ð 0ð 0ðd'Zð 'Zð 'Zð 'ZðRKð Kð Kð Kð
:ð :ð :ð :ð :ð :r   r   Úlocal_tensorr    Úrefr   r#   c                ó¶   — t          j        |                      ¦   «         |j        |j        d|j        t          |                     ¦   «         ¦  «        ¬¦  «        S )zeWrap `local_tensor` as a DTensor that mirrors `ref`'s mesh, placements,
    global shape, and stride.F)Ú	run_checkr   Ústride)r   Ú
from_localÚ
contiguousr   r   r   r   r¨   )r¤   r¥   s     r   Ú_dtensor_from_local_liker«   ?  sR   € õ ÔØ×ÒÑ!Ô!ØŒØŒØØŒiÝ�S—Z’Z‘\”\Ñ"Ô"ðñ ô ð r   )r¤   r    r¥   r   r#   r   )Ú
__future__r   rm   Útypingr   Úutilsr   rE   Útorch.distributed.tensorr   Útorch.distributed.tensor._utilsr   Ú(torch.distributed.tensor.placement_typesr	   r)   r   r
   r   r«   r2   r   r   ú<module>r²      sC  ðð #Ð "Ð "Ð "Ð "Ð "à €€€Ø  Ð  Ð  Ð  Ð  Ð  à &Ð &Ð &Ð &Ð &Ð &ð ð 1Ø€L€L€LØ0Ð0Ð0Ð0Ð0Ð0àÐÑÔð OØ€L€L€LØ0Ð0Ð0Ð0Ð0Ð0ØUÐUÐUÐUÐUÐUØ>Ð>Ð>Ð>Ð>Ð>ð ˆ7�5Ð7Ñ8Ô8ð O¸W¸WÀUÐLjÑ=kÔ=kð OØ,1Ô,NˆÔ)ðW:ð W:ð W:ð W:ð W:ñ W:ô W:ð W:ðt
ð 
ð 
ð 
ð 
ð 
r   