§
    ‚ŠtjN ã                  ó&  — U d dl mZ d dlZd dlZd dlZd dlZd dlmZ ddlm	Z	 ddl
mZmZ ddlmZ ddlmZ  e¦   «         r)d dlZd dlmZ d d	lmZ ej                             ¦   «         Z ej        e¦  «        Zd
„ Z	 dndod„Zdpd„Zdqdrd„Z e¦   «         rEej        ej        ej         ej!        ej"        ej#        ej$        ej%        ej&        ej'        ej(        dœZ)dsd„Z*d „ Z+	 dtdud&„Z,dvdwd(„Z-d)„ Z. G d*„ d+ej/        j0        ¦  «        Z1 G d,„ d-ej/        j0        ¦  «        Z2 G d.„ d/ej/        j0        ¦  «        Z3 G d0„ d1ej/        j0        ¦  «        Z4 G d2„ d3ej/        j0        ¦  «        Z5d4„ Z6d5„ Z7d6„ Z8d7„ Z9d8„ Z:	 	 	 dndxd;„Z; G d<„ d=¦  «        Z< G d>„ d?e<¦  «        Z= G d@„ dAe<¦  «        Z> G dB„ dCe<¦  «        Z? G dD„ dEe<¦  «        Z@ G dF„ dGe<¦  «        ZA G dH„ dIe=¦  «        ZB G dJ„ dKeA¦  «        ZC G dL„ dMe<¦  «        ZD G dN„ dOe<¦  «        ZE G dP„ dQe<¦  «        ZF G dR„ dSe<¦  «        ZG G dT„ dUeG¦  «        ZH G dV„ dWe<¦  «        ZI G dX„ dYe<¦  «        ZJ G dZ„ d[e<¦  «        ZK G d\„ d]e¦  «        ZL eL¦   «         ZMd]eNd^<   dydc„ZOdzdf„ZPdg„ ZQdh„ ZRd{dl„ZSdm„ ZTdS )|é    )ÚannotationsN)Úreduceé   )ÚDistributedConfig)Úis_torch_greater_or_equalÚlogging)ÚGeneralInterface)Úis_torch_available)Únnc                óè   — t          t          j        d¦  «        rWt          t          j        j        d¦  «        r8t	          | t          j        j        j        ¦  «        r|                      ¦   «         S | S )a¢  Unwrap a `DTensor` to its local shard if needed; pass through otherwise.

    Custom kernels (CUTLASS, CuteDSL, Triton) take raw tensor pointers and don't
    understand `DTensor`, so weights wrapped by FSDP2 / EP need this unwrap before
    they can be fed to the kernel. ``to_local()`` is autograd-aware on the train
    path: backward rewraps the gradient as a DTensor matching each parameter's
    placements.
    ÚtensorÚDTensor)ÚhasattrÚtorchÚdistributedr   Ú
isinstancer   Úto_local)Úts    úg/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/integrations/tensor_parallel.pyr   r   (   s_   € õ �uÔ  (Ñ+Ô+ð  µ½Ô8IÔ8PÐR[Ñ0\Ô0\ð  Ý�a�Ô*Ô1Ô9Ñ:Ô:ð 	 Ø—:’:‘<”<ÐØ€Hó    Útp_planústr | dict[str, str] | NoneÚtp_sizeú
int | Nonec                óÐ  — |�| €t          d¦  «        ‚| �|�t          d¦  «        ‚|�€?t          d¦  «        st          d¦  «        ‚t          j                             ¦   «         j        }|dk    rt          d¦  «        ‚t          t          |¦  «        }t          j	         
                    ¦   «         sñ	 t          t          j        d         ¦  «        }t          t          j        d	         ¦  «        }t          t          j        d
         ¦  «        }dddddddœ}	|	                     |¦  «        }
t          j	                             |
||¬¦  «         t          t          |¦  «        }|dk    r|                     |¦  «         n"# t"          $ r}t          d¦  «        |‚d}~ww xY w|dk    r^|                     t          t          j        d	         ¦  «        ¦  «         |                     ¦   «         }t          j        ||¦  «        }|}nt          j        |¦  «        }|pi }|�|nt          j	                             ¦   «         }t          j	                             |j        |f¦  «        }nz|j        dk    r d|j        vrt          d¦  «        ‚|d         }|                     ¦   «         }t          j        |j        › dt          t          j        d	         ¦  «        › �¦  «        }|||fS )z¬
    Sets up the device mesh and initialized the backend for tensor parallelism.
    This function is called when the model is loaded and the TP plan is set to 'auto'.
    Nz-tp_plan has to be set when tp_size is passed.zY`tp_plan` and `device_map` are mutually exclusive. Choose either one for parallelization.z2.5z3Tensor parallel is only supported for `torch>=2.5`.Úmpsz3Tensor parallelism is not supported on MPS devices.ÚRANKÚ
LOCAL_RANKÚ
WORLD_SIZEÚncclÚglooÚxcclÚhcclÚneuronÚtpu_dist)ÚcudaÚcpuÚxpuÚhpur$   Útpu)ÚbackendÚrankÚ
world_sizer'   z†We tried to initialize torch.distributed for you, but it failed. Make sure you init torch distributed in your script to use `tp_plan`.é   ÚtpzsWhen using `tp_plan` and n-d `device_mesh`, it must contain a 'tp' dimension. Please provide a valid `device_mesh`.ú:)Ú
ValueErrorr   ÚOSErrorr   Ú_CÚ_get_acceleratorÚtypeÚRuntimeErrorÚgetattrr   Úis_initializedÚintÚosÚenvironÚgetÚinit_process_groupÚ
set_deviceÚ	ExceptionÚcurrent_deviceÚdeviceÚget_world_sizeÚinit_device_meshÚndimÚmesh_dim_namesÚsizeÚdevice_type)r   r   Údevice_meshÚ
device_maprG   r@   r,   Ú
local_rankr-   Úbackend_mapr+   ÚeÚindexÚ	tp_devices                 r   Úinitialize_tensor_parallelismrO   7   sì  € ð Ð˜w˜ÝÐHÑIÔIÐIØÐ˜zÐ5ÝÐtÑuÔuÐuØÑÝ(¨Ñ/Ô/ð 	QÝÐOÑPÔPÐPõ ”h×/Ò/Ñ1Ô1Ô6ˆØ˜%ÒÐÝÐTÑUÔUÐUÝ ¥¨Ñ4Ô4ˆÝÔ ×/Ò/Ñ1Ô1ð 	ðÝ�2œ: fÔ-Ñ.Ô.�Ý ¥¤¨LÔ!9Ñ:Ô:�
Ý ¥¤¨LÔ!9Ñ:Ô:�
ð #Ø!Ø!Ø!Ø&Ø%ðð �ð &Ÿ/š/¨+Ñ6Ô6�åÔ!×4Ò4¸WÈ4Ð\fÐ4ÑgÔgÐgÝ!(­°Ñ!<Ô!<�Ø %Ò'Ð'Ø"×-Ò-¨jÑ9Ô9Ð9øøåð ð ð ÝðWñô ð ðøøøøðøøøð ˜%ÒÐØ×%Ò%¥c­"¬*°\Ô*BÑ&CÔ&CÑDÔDÐDØ"×1Ò1Ñ3Ô3ˆEÝœ [°%Ñ8Ô8ˆIØ"ˆJˆJåœ [Ñ1Ô1ˆIØ$Ð*¨ˆJà$Ð0�'�'µeÔ6G×6VÒ6VÑ6XÔ6XˆÝÔ'×8Ò8¸¼È'ÈÑTÔTˆˆàÔ˜aÒÐØ˜;Ô5Ð5Ð5Ý ð<ñô ð ð & dÔ+ˆKØ×"Ò"Ñ$Ô$ˆÝ”\ [Ô%<Ð"^Ð"^½sÅ2Ä:ÈlÔC[Ñ?\Ô?\Ð"^Ð"^Ñ_Ô_ˆ
à�{ GÐ+Ð+s   Â4CF Æ
F!ÆFÆF!ÚnameÚstrÚreturnc                ó0   — t          j        dd„ | ¦  «        S )ag  
    Replace the numbers in the `name` by wildcards, only if they are in-between dots (`.`) or if they are between
    a dot (`.`) and the end of the string.
    This matches how modules are named/numbered when using a nn.ModuleList or nn.Sequential, but will NOT match
    numbers in a parameter name itself, e.g. if the param is named `"w1"` or `"w2"`.
    z\.\d+(\.|$)c                ó2   — d|                       d¦  «        z   S )Nz.*r.   ©Úgroup)Úms    r   ú<lambda>z2replace_layer_number_by_wildcard.<locals>.<lambda>†   s   € ¨D°1·7²7¸1±:´:Ñ,=€ r   )ÚreÚsub)rP   s    r   Ú replace_layer_number_by_wildcardr[      s   € õ Œ6�.Ð"=Ð"=¸tÑDÔDÐDr   TÚparameter_nameúdict[str, str]ú
str | Nonec                ó˜   — t          | ¦  «        }||v r||         S |r,d|v r(|                     dd¦  «        d         x}|v r||         S dS )a£  
    Get the TP style for a parameter from the TP plan.

    The TP plan is a dictionary that maps parameter names to TP styles.
    The parameter name can be a generic name with wildcards (e.g. "*.weight") or a specific name (e.g. "layer_1.weight").

    The `is_weight` is important because for weights, we want to support `.weights` and `.bias` cases seamlessly! but
    not parent classes for `post_init` calls
    ú.r.   r   N)r[   Úrsplit)r\   r   Ú	is_weightÚgeneric_param_nameÚmodule_names        r   Ú_get_parameter_tp_planre   ‰   su   € õ :¸.ÑIÔIÐØ˜WÐ$Ð$ØÐ)Ô*Ð*Ø	ð $�sÐ0Ð0Ð0ÐEW×E^ÒE^Ð_bÐdeÑEfÔEfÐghÔEiÐ6i°kÐnuÐ5uÐ5uØ�{Ô#Ð#Øˆ4r   )ÚBOOLÚU8ÚI8ÚI16ÚF16ÚBF16ÚI32ÚF32ÚF64ÚI64ÚF8_E4M3Ú
total_sizer9   Úblocksúint | list[int]ú	list[int]c                óæ   ‡— t          |t          ¦  «        r;t          |¦  «        }| |z  dk    sJ d| › d|› �¦   «         ‚| |z  Šˆfd„|D ¦   «         S | |z  dk    sJ d|› �¦   «         ‚| |z  }|g|z  S )aÃ  
    Convert block count or proportions to block sizes.

    This function accepts

    - The number of blocks (int), in which case the block size is
      total_size//blocks; or
    - A list of block sizes (list[int]).

    In the second case, if sum(blocks) < total_size, the ratios between
    the block sizes will be preserved. For instance, if blocks is
    [2, 1, 1] and total_size is 1024, the returned block sizes are
    [512, 256, 256].
    r   zCannot split z in proportional blocks: c                ó   •— g | ]}‰|z  ‘ŒS © rw   )Ú.0ÚblockÚ	part_sizes     €r   ú
<listcomp>z*_blocks_to_block_sizes.<locals>.<listcomp>Ã   s   ø€ Ð6Ð6Ð6 e�	˜EÑ!Ð6Ð6Ð6r   zPrepacked is not divisible by )r   ÚlistÚsum)rq   rr   Útotal_blocksÚsingle_sizerz   s       @r   Ú_blocks_to_block_sizesr€   °   s´   ø€ õ �&�$ÑÔð &Ý˜6‘{”{ˆØ˜LÑ(¨AÒ-Ð-Ð-Ð/l¸zÐ/lÐ/lÐdjÐ/lÐ/lÑ-Ô-Ð-Ø ,Ñ.ˆ	Ø6Ð6Ð6Ð6¨vÐ6Ñ6Ô6Ð6à˜FÑ" aÒ'Ð'Ð'Ð)RÈ&Ð)RÐ)RÑ'Ô'Ð'Ø  FÑ*ˆØˆ}˜vÑ%Ð%r   c                ó`  — | }|j         |         }|                     ¦   «         }t          |d¬¦  «        }g }	d}
|D ]2}||z  }||z  }|dz   |z  }|	t          |
|z   |
|z   ¦  «        z  }	|
|z  }
Œ3|                     ¦   «         }d}|dk    s|dk    r'|d                              t          j        ¦  «        }d	}|dk    r||	df         }nD|dk    s|d
k    r|dd…|	df         }n*|dk    s|dk    r|d|	f         }nt          d|› d�¦  «        ‚|r|S |                     t          |         ¦  «        S )uä  
    When weights are packed (gate_up_proj), we need to make sure each shard gets its correct share.
    So if you have: gate_proj       ( 16, 5120, 8190)
    and             up_proj         ( 16, 5120, 8190)
    packed as       gate_up_proj    ( 16, 5120, 2 * 8190)
    And you shard along the last dimension, you need to interleave the gate and up values:

    Now, if we shard along the last dimension across TP_size (Tensor Parallelism size), we must interleave the values from gate and up projections correctly.

    Let's take TP_size = 4 for an example:

    Packed tensor `gate_up_proj`
    ---------------------------------------------------------------
    [ G0  G1  G2  G3 | G4  G5  G6  G7 | ... | U0  U1  U2  U3 | U4  U5  U6  U7 | ... ]
     â†‘â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â†‘   â†‘â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â†‘        â†‘â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â†‘  â†‘â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â†‘
       Gate Slice 0      Gate Slice 1            Up Slice 0       Up Slice 1

    Explanation:
    - The first half of the tensor (left of the center) holds the gate_proj values.
    - The second half (right of the center) holds the up_proj values.
    - For TP=4, we divide each half into 4 slices. In this example, we show two slices for brevity.
    - Each shard receives one slice from the gate part and the corresponding slice from the up part.

    For instance:
    â€¢ Shard 0 gets: [ Gate Slice 0, Up Slice 0 ] = [ G0, G1, G2, G3, U0, U1, U2, U3 ]
    â€¢ Shard 1 gets: [ Gate Slice 1, Up Slice 1 ] = [ G4, G5, G6, G7, U4, U5, U6, U7 ]
    â€¢ â€¦ and so on.

    This ensures that each shard receives an equal portion of both gate and up projections, maintaining consistency across tensor parallelism.
    r   )rq   rr   r   r.   Frp   ÚF8_E5M2.TéþÿÿÿNéÿÿÿÿzUnsupported dim z", only dim 0, 1 or 2 are supported)
ÚshaperF   r€   ÚrangeÚ	get_dtypeÚtor   Úfloat16r1   Ústr_to_dtype)ÚparamÚempty_paramrH   r,   ÚdimÚslice_rq   r-   Úblock_sizesÚtensors_slicesÚblock_offsetÚ
block_sizeÚshard_block_sizeÚstartÚstopÚslice_dtypeÚcastedr   s                     r   Úget_packed_weightsr˜   Ê   sŒ  € ð> €FØÔ" 3Ô'€JØ×!Ò!Ñ#Ô#€JÝ(°JÀqÐIÑIÔI€Kà€NØ€LØ!ð #ð #ˆ
Ø%¨Ñ3ÐØÐ'Ñ'ˆØ�q‘Ð,Ñ,ˆØ�% ¨uÑ 4°lÀTÑ6IÑJÔJÑJˆØ˜
Ñ"ˆˆà×"Ò"Ñ$Ô$€Kð €FØ�iÒÐ ;°)Ò#;Ð#;Ø˜”—’¥¤Ñ.Ô.ˆØˆà
ˆa‚x€xØ˜¨Ð+Ô,ˆˆØ	�Šˆ�S˜B’Y�YØ˜˜˜˜>¨3Ð.Ô/ˆˆØ	�Šˆ�S˜B’Y�YØ˜˜^Ð+Ô,ˆˆåÐS¨CÐSÐSÐSÑTÔTÐTàð 4Øˆà�yŠy� kÔ2Ñ3Ô3Ð3r   Úpacked_parameterútorch.TensorÚsharded_dimr-   Ú
num_blocksc                óÞ  — |dk    rt          d¦  «        ‚|dk    r|n	|| j        z   }| j        |         }||z  }||z  }| j        d|…         }| j        |dz   d…         }	 | j        g |¢|‘|‘|‘|	¢R Ž }
t	          |¦  «        }t	          |¦  «        dz   }t          t          |
j        ¦  «        ¦  «        }||         ||         c||<   ||<    |
j        |Ž }|                     | ¦  «        }|S )as  
    Reorders a tensor that was reconstructed from sharded packed weights into its canonical packed format.

    For example, if a weight was packed (e.g., gate_proj and up_proj) and then sharded,
    DTensor.full_tensor() might produce an interleaved layout like [G0, U0, G1, U1, ...]
    along the sharded dimension. This function reorders it to [G0, G1, ..., U0, U1, ...].
    This is an inverse operation to get_packed_weights.

    Args:
        reconstructed_tensor: The tensor reconstructed from DTensor (e.g., via .full_tensor().contiguous()).
        sharded_dim: The dimension index in the reconstructed_tensor that was originally sharded.
        world_size: The tensor parallel world size.
        num_packed_projs: The number of projections that were packed together (e.g., 2 for gate_up_proj).

    Returns:
        The reordered tensor in canonical packed format.
    r   z”Num blocks different from 2 is not supported yet. This is most likely a bug in your implementation as we only pack gate and up projections together.r   Nr.   )	r1   rD   r…   ÚviewÚlenr|   r†   ÚpermuteÚ
reshape_as)r™   r›   r-   rœ   Úactual_sharded_dimÚtotal_size_on_sharded_dimÚoriginal_block_size_on_dimÚshard_chunk_sizeÚprefix_shapeÚsuffix_shapeÚtensor_viewÚaxis_ws_absÚaxis_npp_absÚpermute_orderÚtensor_permutedÚfinal_ordered_tensors                   r   Úrepack_weightsr®     se  € ð0 �Q‚€Ýð cñ
ô 
ð 	
ð )4°qÒ(8Ð(8˜˜¸kÐL\ÔLaÑ>aÐØ 0Ô 6Ð7IÔ JÐØ!:¸jÑ!HÐØ1°ZÑ?Ðà#Ô)Ð*=Ð+=Ð*=Ô>€LØ#Ô)Ð*<¸qÑ*@Ð*BÐ*BÔC€Là'Ð"Ô'ð Ø	ðàðð 	ðð 	ð	ð
 
ðð ð €Kõ �lÑ#Ô#€KÝ�|Ñ$Ô$ qÑ(€Lå�˜{Ô/Ñ0Ô0Ñ1Ô1€MØ>KÈLÔ>YÐ[hÐitÔ[uÐ;€M�+Ñ ¨lÑ ;à)�kÔ)¨=Ð9€Oð +×5Ò5Ð6FÑGÔGÐàÐr   Ú
tensor_idxc                óÊ  — |j         }|j        }t          t          j        |¦  «        }t          | t          j        ¦  «        rt          | j        ¦  «        n|  	                    ¦   «         }	|dk     r||z   }| 
                    ¦   «         dk    r|dk    rt          |	¦  «        dk    rd}n3| 
                    ¦   «         dk    r|dk    rt          |	¦  «        dk    rd}t          j        |	|         |z  ¦  «        }
||
z  }t          ||
z   |	|         ¦  «        }||k    rt          d|› d|› �¦  «        ‚||k    rt          d|› d|› �¦  «        ‚|�l| 
                    ¦   «         dk    rT|dk    rNt          |	¦  «        dk    r;||cxk    r|k     rn n
| d	d	…         S t          j        g t          j        |¬
¦  «        S t%          d	¦  «        gt          |	¦  «        z  }||	|         k     rKt%          ||¦  «        ||<   | t'          |¦  «                 } t          | t          ¦  «        rd„ | D ¦   «         } | S d|	|<   t          j        t'          |	¦  «        t          j        ¬¦  «        S )a
  
    Generalized tensor sharding across a multi-dimensional device mesh.
    Extract only the fraction of the parameter owned by the given `rank` when the parameter would have gone sharding at provided `dim`.
    Extraction follows the pytorch `Shard` placement so that sharding and materializing back to full tensor follows `Shard` semantics.
    `Shard` follows torch.chunk style sharding of the tensor. We demonstrate some cases below on how sharding happens including some edge cases
    such as some ranks having an empty tensor as shard. Below implementation is robut to all these cases.

    Case (1)
    empty_param                 (16, 5120, 8190)
    dim                         0
    device_mesh.size()          4
    rank 0 gets					(4, 5120, 8190)			 (0 ... 4, 5120, 8190)
    rank 1 gets					(4, 5120, 8190)			 (4 ... 8, 5120, 8190)
    rank 2 gets					(4, 5120, 8190)			 (8 ... 12, 5120, 8190)
    rank 3 gets					(4, 5120, 8190)			 (12 ... 16, 5120, 8190)

    Case (2)
    empty_param                 (16, 5120, 8190)
    dim                         0
    device_mesh.size()          14
    rank 0 gets					(2, 5120, 8190)			 (0 ... 2, 5120, 8190)
    rank 1 gets					(2, 5120, 8190)			 (2 ... 4, 5120, 8190)
    rank 2 gets					(2, 5120, 8190)			 (4 ... 6, 5120, 8190)
    rank 3 gets					(2, 5120, 8190)			 (6 ... 8, 5120, 8190)
    rank 4 gets					(2, 5120, 8190)			 (8 ... 10, 5120, 8190)
    rank 5 gets					(2, 5120, 8190)			 (10 ... 12, 5120, 8190)
    rank 6 gets					(2, 5120, 8190)			 (12 ... 14, 5120, 8190)
    rank 7 gets					(2, 5120, 8190)			 (14 ... 16, 5120, 8190)
    rank 8 gets					(0, 5120, 8190)
    rank 9 gets					(0, 5120, 8190)
    rank 10 gets			    (0, 5120, 8190)
    rank 11 gets				(0, 5120, 8190)
    rank 12 gets				(0, 5120, 8190)
    rank 13 gets				(0, 5120, 8190)

    Case (3)
    empty_param                 (16, 5120, 8190)
    dim                         0
    device_mesh.size()          3
    rank 0 gets					(6, 5120, 8190)			 (0 ... 6, 5120, 8190)
    rank 1 gets					(6, 5120, 8190)			 (6 ... 12, 5120, 8190)
    rank 2 gets					(4, 5120, 8190)			 (12 ... 16, 5120, 8190)

    In case (2), empty shards are returned with appropriate dimension to allow for operations to work smoothly.
    Args:
        param (torch.Tensor): The tensor to shard.
        empty_param (torch.Tensor): A tensor used for shape reference.
        device_mesh (torch.Tensor): Shape [d_0, ..., d_n] representing the mesh.
        rank (int): Global rank of the current process/device.
        dim (int): Dimension along which to shard the tensor.
    r   é   r.   r   zdim z* is out of bounds for tensor of dimension zRank z  is out of bounds for mesh size N©ÚdtyperA   c                ó"   — g | ]}|d d …         ‘ŒS ©Nrw   )rx   Úps     r   r{   z$get_tensor_shard.<locals>.<listcomp>°  s    € Ð)Ð)Ð)˜a�Q�q�q�q”TÐ)Ð)Ð)r   )r³   )rD   r…   r   ÚoperatorÚmulr   r   ÚTensorr|   Ú	get_shaper�   rŸ   ÚmathÚceilÚminr1   ÚemptyÚint64ÚsliceÚtuple)r‹   rŒ   rH   r,   r�   r¯   Ú	param_dimÚ
mesh_shaper-   Úparam_shapeÚ
shard_sizer”   ÚendÚslice_indicess                 r   Úget_tensor_shardrÈ   O  sƒ  € ðh Ô €IØÔ"€JÝ�œ jÑ1Ô1€Jå'1°%½¼Ñ'FÔ'FÐ]•$�u”{Ñ#Ô#Ð#ÈEÏOÊOÑL]ÔL]€KØ
ˆQ‚w€wØ˜#‰oˆØ‡‚ÑÔ˜AÒÐ #¨¢( (­s°;Ñ/?Ô/?À1Ò/DÐ/DØˆˆØ	�ŠÑ	Ô	˜aÒ	Ð	 C¨1¢H Hµ°[Ñ1AÔ1AÀQÒ1FÐ1FØˆå”˜; sÔ+¨jÑ8Ñ9Ô9€JØ�:Ñ€EÝ
ˆe�jÑ  +¨cÔ"2Ñ
3Ô
3€Cà
ˆiÒÐÝÐZ ÐZÐZÈyÐZÐZÑ[Ô[Ð[àˆzÒÐÝÐS ÐSÐSÀzÐSÐSÑTÔTÐTð Ð +§/¢/Ñ"3Ô"3°qÒ"8Ð"8¸SÀAºX¸XÍ#ÈkÑJZÔJZÐ^_ÒJ_ÐJ_à�JÐ$Ð$Ò$Ð$ Ò$Ð$Ð$Ð$Ð$à˜˜˜”8ˆOå”;˜r­¬¸TÐBÑBÔBÐBå˜4‘[”[�M¥C¨Ñ$4Ô$4Ñ4€Màˆ{˜3ÔÒÐÝ" 5¨#Ñ.Ô.ˆ�cÑØ•e˜MÑ*Ô*Ô+ˆÝ�e�TÑ"Ô"ð 	*Ø)Ð) 5Ð)Ñ)Ô)ˆEØˆà€K�ÑÝŒ;•u˜[Ñ)Ô)µ´Ð=Ñ=Ô=Ð=r   c                ó0   — t          j        | |d¬¦  «        S )z9Split tensor along last dimension into world_size chunks.r„   ©r�   )r   Úchunk)Úxr-   s     r   Ú_split_along_last_dimrÍ   ·  s   € åŒ;�q˜*¨"Ð-Ñ-Ô-Ð-r   c                  ó>   — e Zd ZdZed„ ¦   «         Zed„ ¦   «         ZdS )Ú_AllReduceBackwardzRIdentity forward, all-reduce backward. Used before colwise layers (f in Megatron).c                ó   — || _         |S rµ   )rH   ©ÚctxrÌ   rH   s      r   Úforwardz_AllReduceBackward.forwardÔ  s   € à%ˆŒØˆr   c                óè   — | j         }|                     ¦   «         dk    r|d fS |                     ¦   «         }t          j        |t          j        j        |                     ¦   «         ¬¦  «         |d fS ©Nr.   ©ÚoprV   )rH   rF   Ú
contiguousÚdistÚ
all_reduceÚReduceOpÚSUMÚ	get_group)rÒ   Úgrad_outputrH   s      r   Úbackwardz_AllReduceBackward.backwardÙ  so   € à”oˆØ×ÒÑÔ Ò"Ð"Ø Ð$Ð$Ø!×,Ò,Ñ.Ô.ˆÝŒ˜­¬Ô(9À×AVÒAVÑAXÔAXÐYÑYÔYÐYØ˜DÐ Ð r   N©Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚstaticmethodrÓ   rß   rw   r   r   rÏ   rÏ   Ñ  sN   € € € € € Ø\Ð\àðð ñ „\ðð ð!ð !ñ „\ð!ð !ð !r   rÏ   c                  ó>   — e Zd ZdZed„ ¦   «         Zed„ ¦   «         ZdS )Ú_AllReduceForwardzQAll-reduce forward, identity backward. Used after rowwise layers (g in Megatron).c                óª   — |                      ¦   «         dk    r|S t          j        |t          j        j        |                     ¦   «         ¬¦  «         |S rÕ   )rF   rÙ   rÚ   rÛ   rÜ   rÝ   rÑ   s      r   rÓ   z_AllReduceForward.forwardæ  sK   € à×ÒÑÔ Ò"Ð"ØˆHÝŒ˜�dœmÔ/°{×7LÒ7LÑ7NÔ7NÐOÑOÔOÐOØˆr   c                ó
   — |d fS rµ   rw   )rÒ   rÞ   s     r   rß   z_AllReduceForward.backwardí  s   € à˜DÐ Ð r   Nrà   rw   r   r   rç   rç   ã  sN   € € € € € Ø[Ð[àðð ñ „\ðð ð!ð !ñ „\ð!ð !ð !r   rç   c                  ó>   — e Zd ZdZed„ ¦   «         Zed„ ¦   «         ZdS )Ú
_AllGatherz<All-gather forward, split backward. Gathers sharded outputs.c                ó®  ‡— || _         |                     ¦   «         }|dk    r‰S ‰                     ¦   «         dz
  }|                     ¦   «         }|                     ¦   «         }‰                     ¦   «         Šˆfd„t          |¦  «        D ¦   «         }‰||<   t          j        |‰|¬¦  «         t          j
        ||¬¦  «                             ¦   «         S )Nr.   c                ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS rw   ©r   Ú
empty_like)rx   Ú_rÌ   s     €r   r{   z&_AllGather.forward.<locals>.<listcomp>  s$   ø€ ÐFÐFÐF¨q•uÔ'¨Ñ*Ô*ÐFÐFÐFr   rU   rÊ   ©rH   rF   r�   Úget_local_rankrÝ   rØ   r†   rÙ   Ú
all_gatherr   Úcat)rÒ   rÌ   rH   r-   Úlast_dimr,   rV   Útensor_lists    `      r   rÓ   z_AllGather.forwardõ  sÌ   ø€ à%ˆŒØ ×%Ò%Ñ'Ô'ˆ
à˜Š?ˆ?ØˆHà—5’5‘7”7˜Q‘;ˆØ×)Ò)Ñ+Ô+ˆØ×%Ò%Ñ'Ô'ˆà�LŠL‰NŒNˆØFÐFÐFÐFµE¸*Ñ4EÔ4EÐFÑFÔFˆØˆ�DÑÝŒ˜ Q¨eÐ4Ñ4Ô4Ð4ÝŒy˜¨(Ð3Ñ3Ô3×>Ò>Ñ@Ô@Ð@r   c                óÌ   — | j         }|                     ¦   «         }|dk    r|d fS |                     ¦   «         }t          ||¦  «        }||                              ¦   «         d fS ©Nr.   ©rH   rF   rò   rÍ   rØ   )rÒ   rÞ   rH   r-   r,   Úchunkss         r   rß   z_AllGather.backward  si   € à”oˆØ ×%Ò%Ñ'Ô'ˆ
à˜Š?ˆ?Ø Ð$Ð$à×)Ò)Ñ+Ô+ˆÝ& {°JÑ?Ô?ˆØ�dŒ|×&Ò&Ñ(Ô(¨$Ð.Ð.r   Nrà   rw   r   r   rë   rë   ò  sQ   € € € € € ØFÐFàðAð Añ „\ðAð" ð	/ð 	/ñ „\ð	/ð 	/ð 	/r   rë   c                  ó>   — e Zd ZdZed„ ¦   «         Zed„ ¦   «         ZdS )Ú_Splitz>Split forward, all-gather backward. Scatters replicated input.c                óÄ   — || _         |                     ¦   «         }|dk    r|S |                     ¦   «         }t          ||¦  «        }||                              ¦   «         S rø   rù   )rÒ   rÌ   rH   r-   r,   rú   s         r   rÓ   z_Split.forward  s^   € à%ˆŒØ ×%Ò%Ñ'Ô'ˆ
à˜Š?ˆ?ØˆHà×)Ò)Ñ+Ô+ˆÝ& q¨*Ñ5Ô5ˆØ�dŒ|×&Ò&Ñ(Ô(Ð(r   c                ó¶  ‡— | j         }|                     ¦   «         }|dk    r‰d fS ‰                     ¦   «         dz
  }|                     ¦   «         }|                     ¦   «         }‰                     ¦   «         Šˆfd„t          |¦  «        D ¦   «         }‰||<   t          j        |‰|¬¦  «         t          j
        ||¬¦  «                             ¦   «         d fS )Nr.   c                ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS rw   rî   ©rx   rð   rÞ   s     €r   r{   z#_Split.backward.<locals>.<listcomp>0  ó$   ø€ ÐPÐPÐP¸•uÔ'¨Ñ4Ô4ÐPÐPÐPr   rU   rÊ   rñ   ©rÒ   rÞ   rH   r-   rõ   r,   rV   rö   s    `      r   rß   z_Split.backward#  óÞ   ø€ à”oˆØ ×%Ò%Ñ'Ô'ˆ
à˜Š?ˆ?Ø Ð$Ð$à—?’?Ñ$Ô$ qÑ(ˆØ×)Ò)Ñ+Ô+ˆØ×%Ò%Ñ'Ô'ˆà!×,Ò,Ñ.Ô.ˆØPÐPÐPÐP½eÀJÑ>OÔ>OÐPÑPÔPˆØ'ˆ�DÑÝŒ˜ [¸Ð>Ñ>Ô>Ð>ÝŒy˜¨(Ð3Ñ3Ô3×>Ò>Ñ@Ô@À$ÐFÐFr   Nrà   rw   r   r   rü   rü     sS   € € € € € ØHÐHàð	)ð 	)ñ „\ð	)ð ðGð Gñ „\ðGð Gð Gr   rü   c                  ó>   — e Zd ZdZed„ ¦   «         Zed„ ¦   «         ZdS )Ú_ReduceScatterzCReduce-scatter forward, all-gather backward. For sequence parallel.c                óÂ  — || _         |                     ¦   «         }|dk    r|S |                     ¦   «         dz
  }|                     ¦   «         }t	          |                     ||¬¦  «        ¦  «        }t	          |j        ¦  «        }||xx         |z  cc<   t          j        ||j	        |j
        ¬¦  «        }t          j        ||t          j        j        |¬¦  «         |S )Nr.   rÊ   r²   rÖ   )rH   rF   r�   rÝ   r|   rË   r…   r   r¾   r³   rA   rÙ   Úreduce_scatterrÛ   rÜ   )	rÒ   rÌ   rH   r-   rõ   rV   Úinput_chunksÚoutput_shapeÚoutputs	            r   rÓ   z_ReduceScatter.forward9  sÎ   € à%ˆŒØ ×%Ò%Ñ'Ô'ˆ
à˜Š?ˆ?ØˆHà—5’5‘7”7˜Q‘;ˆØ×%Ò%Ñ'Ô'ˆå˜AŸGšG J°H˜GÑ=Ô=Ñ>Ô>ˆÝ˜AœG‘}”}ˆØ�XÐÐÔ :Ñ-ÐÐÑÝ”˜\°´ÀÄÐJÑJÔJˆåÔ˜F LµT´]Ô5FÈeÐTÑTÔTÐTØˆr   c                ó¶  ‡— | j         }|                     ¦   «         }|dk    r‰d fS ‰                     ¦   «         dz
  }|                     ¦   «         }|                     ¦   «         }‰                     ¦   «         Šˆfd„t          |¦  «        D ¦   «         }‰||<   t          j        |‰|¬¦  «         t          j
        ||¬¦  «                             ¦   «         d fS )Nr.   c                ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS rw   rî   r   s     €r   r{   z+_ReduceScatter.backward.<locals>.<listcomp>Y  r  r   rU   rÊ   rñ   r  s    `      r   rß   z_ReduceScatter.backwardL  r  r   Nrà   rw   r   r   r  r  6  sS   € € € € € ØMÐMàðð ñ „\ðð$ ðGð Gñ „\ðGð Gð Gr   r  c                ó8   — t                                | |¦  «        S )zAIdentity forward, all-reduce backward. Use before colwise layers.)rÏ   Úapply©rÌ   rH   s     r   Úall_reduce_backwardr  d  s   € å×#Ò# A {Ñ3Ô3Ð3r   c                ó8   — t                                | |¦  «        S )z@All-reduce forward, identity backward. Use after rowwise layers.)rç   r  r  s     r   Úall_reduce_forwardr  i  s   € å×"Ò" 1 kÑ2Ô2Ð2r   c                ó8   — t                                | |¦  «        S )z#All-gather forward, split backward.)rë   r  r  s     r   ró   ró   n  s   € å×Ò˜A˜{Ñ+Ô+Ð+r   c                ó8   — t                                | |¦  «        S )z#Split forward, all-gather backward.)rü   r  r  s     r   Úsplitr  s  s   € å�<Š<˜˜;Ñ'Ô'Ð'r   c                ó8   — t                                | |¦  «        S )z,Reduce-scatter forward, all-gather backward.)r  r  r  s     r   r  r  x  s   € å×Ò  ;Ñ/Ô/Ð/r   Úmoduleú	nn.Modulec                óx   ‡‡‡— ‰�|                       ˆˆfd„¦  «         ‰�|                      ˆˆfd„¦  «         | S )zž
    Copy pasted from torch's function but we remove the communications (partitioning)
    as well as buffer registering that is similarly not efficient.
    Nc                ó   •—  ‰| |‰¦  «        S rµ   rw   )ÚmodÚinputsrH   Úinput_fns     €€r   rX   z#distribute_module.<locals>.<lambda>ˆ  s   ø€ ¸X¸XÀcÈ6ÐS^Ñ=_Ô=_€ r   c                ó   •—  ‰| |‰¦  «        S rµ   rw   )r  r  ÚoutputsrH   Ú	output_fns      €€r   rX   z#distribute_module.<locals>.<lambda>Š  s   ø€ À)À)ÈCÐQXÐZeÑBfÔBf€ r   )Úregister_forward_pre_hookÚregister_forward_hook)r  rH   r  r   s    ```r   Údistribute_moduler#  }  sZ   øøø€ ð ÐØ×(Ò(Ð)_Ð)_Ð)_Ð)_Ð)_Ñ`Ô`Ð`ØÐØ×$Ò$Ð%fÐ%fÐ%fÐ%fÐ%fÑgÔgÐgØ€Mr   c                  ó`   — e Zd ZdZdZdZdZdd„Zd„ Zd„ Z		 ddd„Z
ddd„Zdd„Zdd„Zdd„ZdS )ÚTensorParallelLayerz.General tensor parallel layer for transformersNc                ó0   — || _         || _        || _        d S rµ   )r,   rH   rŒ   )ÚselfrH   r,   rŒ   s       r   Ú__init__zTensorParallelLayer.__init__•  s   € ØˆŒ	Ø&ˆÔØ&ˆÔÐÐr   c                ó   — t           ‚rµ   ©ÚNotImplementedError©r'  r  r  rH   s       r   Ú_prepare_input_fnz%TensorParallelLayer._prepare_input_fnš  ó   € Ý!Ð!r   c                ó   — t           ‚rµ   r*  ©r'  r  r  rH   s       r   Ú_prepare_output_fnz&TensorParallelLayer._prepare_output_fn�  r.  r   r‹   rš   r¯   r   rR   c                ó   — t           ‚rµ   r*  ©r'  r‹   r¯   rA   r³   s        r   Úshard_tensorz TensorParallelLayer.shard_tensor   s
   € õ "Ð!r   Ú r  r  Ú
layer_namerQ   c                ó   — dS )zHRaise if the module cannot be sharded with this style on the given mesh.Nrw   )r'  r  rH   r6  s       r   Úvalidate_modulez#TensorParallelLayer.validate_module¥  s   € àˆr   c                ó>   — t          ||| j        | j        ¦  «         d S rµ   )r#  r-  r1  ©r'  r  rH   Úkwargss       r   Úprepare_module_tpz%TensorParallelLayer.prepare_module_tp©  s0   € ÝØØØÔ"ØÔ#ñ		
ô 	
ð 	
ð 	
ð 	
r   Ú
full_shapeútuple[int, ...] | torch.Sizeútuple[int, ...]c                ó    — t          |¦  «        S )zé
        Compute the expected shape after TP sharding for a given full shape.

        Args:
            full_shape: The full (unsharded) parameter shape

        Returns:
            The expected sharded shape for this rank
        )rÁ   )r'  r=  s     r   Úget_expected_sharded_shapez.TensorParallelLayer.get_expected_sharded_shape±  s   € õ �ZÑ Ô Ð r   c                ó   — dS )zá
        Update module attributes (e.g. in_features, out_features) to reflect sharded dimensions.

        Args:
            module: The module to update

        Returns:
            None, update the module in-place
        Nrw   ©r'  r  s     r   Úupdate_module_attributesz,TensorParallelLayer.update_module_attributes¾  s	   € ð 	ˆr   ©NNN©r‹   rš   r¯   r   rR   rš   ©r5  ©r  r  r6  rQ   ©r  r  rR   r  ©r=  r>  rR   r?  ©r  r  )rá   râ   rã   rä   rH   r,   rŒ   r(  r-  r1  r4  r8  r<  rA  rD  rw   r   r   r%  r%  Ž  sÐ   € € € € € Ø8Ð8à€KØ€DØ€Kð'ð 'ð 'ð 'ð
"ð "ð "ð"ð "ð "ð VZð"ð "ð "ð "ð "ð
ð ð ð ð ð
ð 
ð 
ð 
ð!ð !ð !ð !ð
ð 
ð 
ð 
ð 
ð 
r   r%  c                  óX   ‡ — e Zd ZdZddˆ fd„Zddd„Zd„ Zd„ Z	 ddd„Zd d„Z	d!d„Z
ˆ xZS )"ÚColwiseParallelzÕ
    Column-wise parallel: weight is sharded on dim -2 (output features).
    Forward: input replicated -> output sharded on last dim.
    If gather_output=True, output is all-gathered to produce full tensor.
    FÚgather_outputÚboolc                óH   •—  t          ¦   «         j        di |¤Ž || _        d S ©Nrw   )Úsuperr(  rN  )r'  rN  r;  Ú	__class__s      €r   r(  zColwiseParallel.__init__Ò  s.   ø€ Ø�‰ŒÔÐ"Ð"˜6Ð"Ð"Ð"Ø*ˆÔÐÐr   r5  r  r  r6  rQ   c                óø   — t          |dd ¦  «        }| j        r]|�]||                     ¦   «         z  dk    rDt          d|› dt	          |¦  «        j        › d|› d|                     ¦   «         › d�	¦  «        ‚d S d S d S )NÚout_featuresr   ú`z` (z with out_features=zo) is sharded with 'colwise_gather_output', which requires out_features to be divisible by the number of ranks (z˜) to all-gather equal-size shards. Resize the weight (e.g. `model.resize_token_embeddings` for LM heads) or override this module's entry in the tp_plan.)r7   rN  rF   r1   r5   rá   )r'  r  rH   r6  rU  s        r   r8  zColwiseParallel.validate_moduleÖ  sÁ   € Ý˜v ~°tÑ<Ô<ˆØÔð 	 ,Ð":¸|Èk×N^ÒN^ÑN`ÔN`Ñ?`ÐdeÒ?eÐ?eÝðq�Jð qð q¥4¨¡<¤<Ô#8ð qð qÈ\ð qð qà×$Ò$Ñ&Ô&ðqð qð qñô ð ð	ð 	Ð":Ð":Ð?eÐ?er   c                ó:   — |r|d         n|}t          ||¦  «        S ©Nr   ©r  ©r'  r  r  rH   Úinput_tensors        r   r-  z!ColwiseParallel._prepare_input_fnà  s$   € Ø$*Ð6�v˜a”y�y°ˆÝ" <°Ñ=Ô=Ð=r   c                ó4   — | j         rt          ||¦  «        S |S rµ   )rN  ró   r0  s       r   r1  z"ColwiseParallel._prepare_output_fnä  s"   € ØÔð 	4Ý˜g {Ñ3Ô3Ð3Øˆr   Nr‹   rš   r¯   r   rR   c                ód  — t          |t          j        ¦  «        r|                     ¦   «         n t	          |                     ¦   «         ¦  «        }|dk    r#t          || j        | j        | j	        d¦  «        }n"t          || j        | j        | j	        d¦  «        }| 
                    ||¬¦  «        S ©Nr.   r„   rƒ   ©rA   r³   ©r   r   r¹   r�   rŸ   rº   rÈ   rŒ   rH   r,   rˆ   ©r'  r‹   r¯   rA   r³   r�   Ú	parameters          r   r4  zColwiseParallel.shard_tensoré  s™   € õ (¨­u¬|Ñ<Ô<ÐXˆe�iŠi‰kŒkˆkÅ#ÀeÇoÂoÑFWÔFWÑBXÔBXˆØ�!Š8ˆ8Ý(¨°Ô0@À$ÔBRÐTXÔT]Ð_aÑbÔbˆIˆIå(¨°Ô0@À$ÔBRÐTXÔT]Ð_aÑbÔbˆIØ�|Š| 6°ˆ|Ñ7Ô7Ð7r   r=  r>  r?  c                ób  — | j                              ¦   «         }t          |¦  «        }t          |¦  «        dk    rdnd}|dk     rt          |¦  «        |z   n|}t	          j        ||         |z  ¦  «        }| j        |z  }t          ||z   ||         ¦  «        }||z
  ||<   t          |¦  «        S )Nr.   r„   rƒ   r   )	rH   rF   r|   rŸ   r»   r¼   r,   r½   rÁ   ©r'  r=  r-   r…   r�   rÅ   r”   rÆ   s           r   rA  z*ColwiseParallel.get_expected_sharded_shapeô  s¬   € ØÔ%×*Ò*Ñ,Ô,ˆ
Ý�ZÑ Ô ˆå˜‘J”J !’O�Oˆbˆb¨ˆØ"%¨¢' '�c�%‰jŒj˜3ÑÐ¨sˆÝ”Y˜u Sœz¨JÑ6Ñ7Ô7ˆ
Ø”	˜JÑ&ˆÝ�%˜*Ñ$ e¨C¤jÑ1Ô1ˆØ˜5‘[ˆˆc‰
Ý�U‰|Œ|Ðr   c                óˆ   — | j         s8t          |d¦  «        r*|                      |j        f¦  «        d         |_        d S d S d S )NrU  r   )rN  r   rA  rU  rC  s     r   rD  z(ColwiseParallel.update_module_attributes   sb   € ð Ô!ð 	]¥g¨f°nÑ&EÔ&Eð 	]Ø"&×"AÒ"AÀ6ÔCVÐBXÑ"YÔ"YÐZ[Ô"\ˆFÔÐÐð	]ð 	]ð 	]ð 	]r   ©F)rN  rO  rG  rH  rE  rF  rJ  rK  )rá   râ   rã   rä   r(  r8  r-  r1  r4  rA  rD  Ú__classcell__©rS  s   @r   rM  rM  Ë  sÓ   ø€ € € € € ðð ð+ð +ð +ð +ð +ð +ð +ðð ð ð ð ð>ð >ð >ðð ð ð VZð	8ð 	8ð 	8ð 	8ð 	8ð
ð 
ð 
ð 
ð]ð ]ð ]ð ]ð ]ð ]ð ]ð ]r   rM  c                  ó,   — e Zd ZdZd„ Zd„ Zdd„Zd„ ZdS )ÚReplicatedWithGradAllReduceaM  
    Replicated parameter with gradient all-reduce.

    For parameters like q_norm/k_norm that sit between colwise and rowwise
    layers. The parameter is replicated (not sharded), but its gradient
    accumulates from local heads only in TP mode. This class registers a
    backward hook to all-reduce the parameter gradient.
    c                ó   — |S rµ   rw   r,  s       r   r-  z-ReplicatedWithGradAllReduce._prepare_input_fn  ó   € Øˆr   c                ó   — |S rµ   rw   r0  s       r   r1  z.ReplicatedWithGradAllReduce._prepare_output_fn  ó   € Øˆr   Nc                ó<   — |d                               ||¬¦  «        S ©N.r_  ©rˆ   r3  s        r   r4  z(ReplicatedWithGradAllReduce.shard_tensor  ó   € Ø�SŒz�}Š} F°%ˆ}Ñ8Ô8Ð8r   c                ó:   — |fd„}|                      |¦  «         d S )Nc                ól   — |                       ¦   «         D ]}|j        �t          |j        |¦  «         Œd S rµ   )Ú
parametersÚgradr  )r  Ú
grad_inputrÞ   Úmeshr‹   s        r   Ú_backward_hookzEReplicatedWithGradAllReduce.prepare_module_tp.<locals>._backward_hook  s@   € ØŸšÑ)Ô)ð 9ð 9�Ø”:Ð)Ý& u¤z°4Ñ8Ô8Ð8øð9ð 9r   )Úregister_full_backward_hook)r'  r  rH   r;  ry  s        r   r<  z-ReplicatedWithGradAllReduce.prepare_module_tp  s8   € ð ?Jð 	9ð 	9ð 	9ð 	9ð
 	×*Ò*¨>Ñ:Ô:Ð:Ð:Ð:r   rE  ©rá   râ   rã   rä   r-  r1  r4  r<  rw   r   r   rj  rj    s_   € € € € € ðð ðð ð ðð ð ð9ð 9ð 9ð 9ð;ð ;ð ;ð ;ð ;r   rj  c                  ó,   — e Zd ZdZd„ Zd„ Zdd„Zd„ ZdS )ÚAllReduceParallelzêAll-reduce a module's forward output across the TP mesh. Use as a declarative
    sync point at the boundary of a multi-arg module whose compute ends in a partial
    sum (e.g. the lightning indexer's score sum before its top-k).
    c                ó   — |S rµ   rw   r,  s       r   r-  z#AllReduceParallel._prepare_input_fn+  rl  r   c                ó"   — t          ||¦  «        S rµ   ©r  r0  s       r   r1  z$AllReduceParallel._prepare_output_fn.  s   € Ý! '¨;Ñ7Ô7Ð7r   Nc                ó<   — |d                               ||¬¦  «        S rp  rq  r3  s        r   r4  zAllReduceParallel.shard_tensor1  rr  r   c                ó4   — t          ||| j        ¬¦  «         d S ©N)r   )r#  r1  r:  s       r   r<  z#AllReduceParallel.prepare_module_tp4  s    € Ý˜& +¸Ô9PÐQÑQÔQÐQÐQÐQr   rE  r{  rw   r   r   r}  r}  %  sd   € € € € € ðð ð
ð ð ð8ð 8ð 8ð9ð 9ð 9ð 9ðRð Rð Rð Rð Rr   r}  c                  ó(   — e Zd ZdZd„ Zdd„Zdd„ZdS )ÚMlaKvAProjParallela#  
    For MLA attention used in DeepSeek-V2 style models (deepseek_v2, longcat_flash, glm_moe_dsa, glm4_moe_lite):
    kv_a_proj_with_mqa output is [kv_lora_rank + qk_rope_head_dim] (can have different naming but important thing
    to understand is that it is split)
    Example below (from modeling_longcat_flash.py):

    kv_a_proj_with_mqa
            |
            split
            /            k_pass    k_rot  <-- "bypasses kv_b_proj"
        |          |        (goes straight to attention,
    kv_a_layernorm |         never touches kv_b_proj)
        |          |
    kv_b_proj      |
    (colwise)      |
        |          |
        k_pass     k_rot
            \      /
               cat
                |
            key_states

    k_pass is passed to kv_b_proj (colwise) which has built-in all_reduce_backward so we don't have a partial gradient for it.
    However, k_rot goes straight to attention, never touches kv_b_proj. So we need to average gradient across all ranks otherwise we only get gradient for one rank (partial gradient).
    c                ó2  — t          |j        d¦  «        s%t          dt          |¦  «        j        › d�¦  «        ‚|j        j        }|                     |j        d         |z
  |gd¬¦  «        \  }}t          ||¦  «        }t          j
        ||gd¬¦  «        S )NÚqk_rope_head_dimzConfig for z· does not have `qk_rope_head_dim`. MlaKvAProjParallel requires `qk_rope_head_dim` to be defined in the model config. Please add it to the model's config or update the TP plan mapping.r„   rÊ   )r   ÚconfigÚAttributeErrorr5   rá   r‡  r  r…   r  r   rô   )r'  r  r
  rH   Úrope_dimÚpass_outputÚrope_outputs          r   r1  z%MlaKvAProjParallel._prepare_output_fnT  s­   € Ý�s”zÐ#5Ñ6Ô6ð 	Ý ðU�d 3™iœiÔ0ð Uð Uð Uñô ð ð
 ”:Ô.ˆØ#)§<¢<°´¸bÔ1AÀHÑ1LÈhÐ0WÐ]_ <Ñ#`Ô#`Ñ ˆ�[Ý)¨+°{ÑCÔCˆÝŒy˜+ {Ð3¸Ð<Ñ<Ô<Ð<r   Nc                ó<   — |d                               ||¬¦  «        S rp  rq  r3  s        r   r4  zMlaKvAProjParallel.shard_tensor`  rr  r   c                óB   — ||_         t          ||| j        ¬¦  «         d S rƒ  )rˆ  r#  r1  )r'  r  rH   rˆ  r;  s        r   r<  z$MlaKvAProjParallel.prepare_module_tpc  s'   € ØˆŒÝ˜& +¸Ô9PÐQÑQÔQÐQÐQÐQr   rE  rµ   )rá   râ   rã   rä   r1  r4  r<  rw   r   r   r…  r…  8  s[   € € € € € ðð ð6
=ð 
=ð 
=ð9ð 9ð 9ð 9ðRð Rð Rð Rð Rð Rr   r…  c                  óN   ‡ — e Zd ZdZddˆ fd„Zd„ Zd„ Z	 ddd„Zdd„Zdd„Z	ˆ xZ
S )ÚRowwiseParallelaž  
    Row-wise parallel: weight is sharded on dim -1 (input features).
    Forward: input (optionally split) -> output partial -> all-reduce to replicate.

    Args:
        split_input: If True, splits replicated input before matmul. Use when input
                     comes from a non-parallelizable operation (chunk/slice).
                     Default False (expects pre-sharded input from colwise layer).
    FÚsplit_inputrO  c                óH   •—  t          ¦   «         j        di |¤Ž || _        d S rQ  )rR  r(  r‘  )r'  r‘  r;  rS  s      €r   r(  zRowwiseParallel.__init__s  s.   ø€ Ø�‰ŒÔÐ"Ð"˜6Ð"Ð"Ð"Ø&ˆÔÐÐr   c                ó    — t          |d¦  «        r|j        �|j        |_        d |_        |r|d         n|}| j        rt	          ||¦  «        S |S )NÚbiasr   )r   r”  Ú_biasr‘  r  rZ  s        r   r-  z!RowwiseParallel._prepare_input_fnw  s^   € Ý�3˜ÑÔð 	 C¤HÐ$8ØœˆCŒIØˆCŒHà$*Ð6�v˜a”y�y°ˆàÔð 	4å˜ {Ñ3Ô3Ð3ØÐr   c                óh   — t          ||¦  «        }t          |d¦  «        r|j        �
||j        z   }|S )Nr•  )r  r   r•  r0  s       r   r1  z"RowwiseParallel._prepare_output_fnƒ  s;   € Ý$ W¨kÑ:Ô:ˆÝ�3˜Ñ Ô ð 	* S¤YÐ%:Ø ¤	Ñ)ˆGØˆr   Nr‹   rš   r¯   r   rR   c                ó0  — t          |t          j        ¦  «        r|                     ¦   «         n t	          |                     ¦   «         ¦  «        }|dk    r	|d         }n"t          || j        | j        | j	        d¦  «        }| 
                    ||¬¦  «        S )Nr.   .r„   r_  r`  ra  s          r   r4  zRowwiseParallel.shard_tensor‰  s‚   € õ (¨­u¬|Ñ<Ô<ÐXˆe�iŠi‰kŒkˆkÅ#ÀeÇoÂoÑFWÔFWÑBXÔBXˆØ�!Š8ˆ8Ø˜cœ
ˆIˆIå(¨°Ô0@À$ÔBRÐTXÔT]Ð_aÑbÔbˆIØ�|Š| 6°ˆ|Ñ7Ô7Ð7r   r=  r>  r?  c                ó|  — t          |¦  «        dk    rt          |¦  «        S | j                             ¦   «         }t	          |¦  «        }d}|dk     rt          |¦  «        |z   n|}t          j        ||         |z  ¦  «        }| j        |z  }t          ||z   ||         ¦  «        }||z
  ||<   t          |¦  «        S ©Nr.   r„   r   )	rŸ   rÁ   rH   rF   r|   r»   r¼   r,   r½   rd  s           r   rA  z*RowwiseParallel.get_expected_sharded_shape”  s¹   € åˆz‰?Œ?˜aÒÐÝ˜Ñ$Ô$Ð$ØÔ%×*Ò*Ñ,Ô,ˆ
Ý�ZÑ Ô ˆØˆØ"%¨¢' '�c�%‰jŒj˜3ÑÐ¨sˆÝ”Y˜u Sœz¨JÑ6Ñ7Ô7ˆ
Ø”	˜JÑ&ˆÝ�%˜*Ñ$ e¨C¤jÑ1Ô1ˆØ˜5‘[ˆˆc‰
Ý�U‰|Œ|Ðr   r  r  c                ó|   — t          |d¦  «        r+d|j        f}|                      |¦  «        d         |_        d S d S )NÚin_featuresr.   )r   r›  rA  )r'  r  r…   s      r   rD  z(RowwiseParallel.update_module_attributes¢  sQ   € Ý�6˜=Ñ)Ô)ð 	Kð ˜Ô*Ð+ˆEØ!%×!@Ò!@ÀÑ!GÔ!GÈÔ!JˆFÔÐÐð		Kð 	Kr   rf  )r‘  rO  rE  rF  rJ  rK  ©rá   râ   rã   rä   r(  r-  r1  r4  rA  rD  rg  rh  s   @r   r�  r�  h  sº   ø€ € € € € ðð ð'ð 'ð 'ð 'ð 'ð 'ð 'ð
ð 
ð 
ðð ð ð VZð	8ð 	8ð 	8ð 	8ð 	8ðð ð ð ðKð Kð Kð Kð Kð Kð Kð Kr   r�  c                  ó   — e Zd ZdZ	 d	d
d„ZdS )ÚPackedColwiseParallelz@Packed column-wise parallel for fused weights like gate_up_proj.Nr‹   rš   r¯   r   rR   c                ó  — t          |t          j        ¦  «        r|                     ¦   «         n t	          |                     ¦   «         ¦  «        }|dk    r#t          || j        | j        | j	        d¦  «        }nw|  
                    | j        j        ¦  «        }|t	          |¦  «        k     r#t          || j        | j        | j	        d¦  «        }n"t          || j        | j        | j	        d¦  «        }|                     ||¬¦  «        S r^  )r   r   r¹   r�   rŸ   rº   rÈ   rŒ   rH   r,   rA  r…   r˜   rˆ   )r'  r‹   r¯   rA   r³   r�   rb  Úexpected_shapes           r   r4  z"PackedColwiseParallel.shard_tensor­  së   € õ (¨­u¬|Ñ<Ô<ÐXˆe�iŠi‰kŒkˆkÅ#ÀeÇoÂoÑFWÔFWÑBXÔBXˆØ�!Š8ˆ8Ý(¨°Ô0@À$ÔBRÐTXÔT]Ð_aÑbÔbˆIˆIà!×<Ò<¸TÔ=MÔ=SÑTÔTˆNØ•S˜Ñ(Ô(Ò(Ð(õ -¨U°DÔ4DÀdÔFVÐX\ÔXaÐceÑfÔf�	�	õ /¨u°dÔ6FÈÔHXÐZ^ÔZcÐegÑhÔh�	Ø�|Š| 6°ˆ|Ñ7Ô7Ð7r   rE  rF  ©rá   râ   rã   rä   r4  rw   r   r   rž  rž  ª  s:   € € € € € ØJÐJð VZð8ð 8ð 8ð 8ð 8ð 8ð 8r   rž  c                  ó   — e Zd ZdZ	 d	d
d„ZdS )ÚPackedRowwiseParallelz=Packed row-wise parallel for fused weights like gate_up_proj.Nr‹   rš   r¯   r   rR   c                óˆ  — t          |t          j        ¦  «        r|                     ¦   «         n t	          |                     ¦   «         ¦  «        }|dk    r	|d         }nÎt          |t          j        ¦  «        r|j        n|                     ¦   «         }| j                             ¦   «         dk    r| j        j        d         nd}t	          |¦  «        dk    r|d         nd}	|	|k     r#t          || j        | j	        | j
        d¦  «        }n"t          || j        | j	        | j
        d¦  «        }|                     ||¬¦  «        S )Nr.   .r„   r   r_  )r   r   r¹   r�   rŸ   rº   r…   rŒ   rÈ   rH   r,   r˜   rˆ   )
r'  r‹   r¯   rA   r³   r�   rb  rÄ   Úexpected_packed_dimÚ
actual_dims
             r   r4  z"PackedRowwiseParallel.shard_tensorÃ  s*  € õ (¨­u¬|Ñ<Ô<ÐXˆe�iŠi‰kŒkˆkÅ#ÀeÇoÂoÑFWÔFWÑBXÔBXˆØ�!Š8ˆ8Ø˜cœ
ˆIˆIõ *4°E½5¼<Ñ)HÔ)HÐ_˜%œ+˜+ÈeÏoÊoÑN_ÔN_ˆKØ@DÔ@P×@TÒ@TÑ@VÔ@VÐZ[Ò@[Ð@[ $Ô"2Ô"8¸Ô"<Ð"<ÐabÐÝ,/°Ñ,<Ô,<ÀÒ,AÐ,A˜ Rœ˜ÀqˆJàÐ/Ò/Ð/å,¨U°DÔ4DÀdÔFVÐX\ÔXaÐceÑfÔf�	�	õ /¨u°dÔ6FÈÔHXÐZ^ÔZcÐegÑhÔh�	Ø�|Š| 6°ˆ|Ñ7Ô7Ð7r   rE  rF  r¡  rw   r   r   r£  r£  À  s:   € € € € € ØGÐGð VZð8ð 8ð 8ð 8ð 8ð 8ð 8r   r£  c                  óR   ‡ — e Zd ZdZddœdˆ fd„Zd„ Zd„ Z	 ddd„Zdd„Zdd„Z	ˆ xZ
S )ÚEmbeddingParallelzXEmbeddingParallel: shards embedding table, handles masked lookups for vocab parallelism.r   ©Úembedding_dim_shardingrª  r9   c               óH   •—  t          ¦   «         j        di |¤Ž || _        d S rQ  )rR  r(  rª  )r'  rª  r;  rS  s      €r   r(  zEmbeddingParallel.__init__Ý  s.   ø€ Ø�‰ŒÔÐ"Ð"˜6Ð"Ð"Ð"Ø&<ˆÔ#Ð#Ð#r   c                óø   — |r|d         n|}| j         dk    rb|                     ¦   «         }|j        j        d         }||z  }||z   }||k     ||k    z  }	|	|_        |                     ¦   «         |z
  }
d|
|	<   |
S |S rX  )rª  rò   Úweightr…   Ú_input_maskÚclone)r'  r  r  rH   r[  r,   Úper_partition_sizeÚvocab_start_indexÚvocab_end_indexÚ
input_maskÚmasked_inputs              r   r-  z#EmbeddingParallel._prepare_input_fná  s«   € Ø$*Ð6�v˜a”y�y°ˆð Ô&¨!Ò+Ð+Ø×-Ò-Ñ/Ô/ˆDð
 "%¤Ô!1°!Ô!4ÐØ $Ð'9Ñ 9ÐØ/Ð2DÑDˆOð 'Ð):Ò:¸|ÈÒ?^Ñ_ˆJØ(ˆCŒOð (×-Ò-Ñ/Ô/Ð2CÑCˆLØ'(ˆL˜Ñ$àÐàÐr   c                óÐ   — | j         dk    rLt          |d¦  «        r<|j        }|                     d¦  «        }||                      |j        ¦  «        z  }|`t          ||¦  «        S )Nr   r®  r„   )rª  r   r®  Ú	unsqueezerˆ   r³   r  )r'  r  r  rH   r³  Úmasks         r   r1  z$EmbeddingParallel._prepare_output_fnû  sh   € àÔ&¨!Ò+Ð+µ¸¸]Ñ0KÔ0KÐ+ØœˆJà×'Ò'¨Ñ+Ô+ˆDØ $ §
¢
¨7¬=Ñ 9Ô 9Ñ9ˆGØ�å! '¨;Ñ7Ô7Ð7r   Nr‹   rš   r¯   r   rR   c                ón  — t          |t          j        ¦  «        r|                     ¦   «         n t	          |                     ¦   «         ¦  «        }|dk    r#t          || j        | j        | j	        d¦  «        }n't          || j        | j        | j	        | j
        ¦  «        }|                     ||¬¦  «        S )Nr.   r„   r_  )r   r   r¹   r�   rŸ   rº   rÈ   rŒ   rH   r,   rª  rˆ   ra  s          r   r4  zEmbeddingParallel.shard_tensor  s¤   € õ (¨­u¬|Ñ<Ô<ÐXˆe�iŠi‰kŒkˆkÅ#ÀeÇoÂoÑFWÔFWÑBXÔBXˆØ�!Š8ˆ8Ý(¨°Ô0@À$ÔBRÐTXÔT]Ð_aÑbÔbˆIˆIå(ØØÔ ØÔ Ø”	ØÔ+ñô ˆIð �|Š| 6°ˆ|Ñ7Ô7Ð7r   r=  r>  r?  c                ól  — | j                              ¦   «         }t          |¦  «        }t          |¦  «        dk    rdn| j        }|dk     rt          |¦  «        |z   n|}t          j        ||         |z  ¦  «        }| j        |z  }t          ||z   ||         ¦  «        }||z
  ||<   t          |¦  «        S r™  )
rH   rF   r|   rŸ   rª  r»   r¼   r,   r½   rÁ   rd  s           r   rA  z,EmbeddingParallel.get_expected_sharded_shape  s±   € ØÔ%×*Ò*Ñ,Ô,ˆ
Ý�ZÑ Ô ˆõ ˜‘J”J !’O�Oˆbˆb¨Ô)DˆØ"%¨¢' '�c�%‰jŒj˜3ÑÐ¨sˆÝ”Y˜u Sœz¨JÑ6Ñ7Ô7ˆ
Ø”	˜JÑ&ˆÝ�%˜*Ñ$ e¨C¤jÑ1Ô1ˆØ˜5‘[ˆˆc‰
Ý�U‰|Œ|Ðr   r  r  c                ó  — t          |d¦  «        r1| j        dk    r&|                      |j        f¦  «        d         |_        t          |d¦  «        r3| j        dk    r*|                      |j        f¦  «        d         |_        d S d S d S )NÚnum_embeddingsr   Úembedding_dimr.   )r   rª  rA  r»  r¼  rC  s     r   rD  z*EmbeddingParallel.update_module_attributes$  s    € Ý�6Ð+Ñ,Ô,ð 	a°Ô1LÐPQÒ1QÐ1QØ$(×$CÒ$CÀVÔEZÐD\Ñ$]Ô$]Ð^_Ô$`ˆFÔ!Ý�6˜?Ñ+Ô+ð 	_°Ô0KÈqÒ0PÐ0PØ#'×#BÒ#BÀFÔDXÐCZÑ#[Ô#[Ð\]Ô#^ˆFÔ Ð Ð ð	_ð 	_Ð0PÐ0Pr   )rª  r9   rE  rF  rJ  rK  rœ  rh  s   @r   r¨  r¨  Ú  s¾   ø€ € € € € ØbÐbà89ð =ð =ð =ð =ð =ð =ð =ð =ðð ð ð4	8ð 	8ð 	8ð VZð8ð 8ð 8ð 8ð 8ð"ð ð ð ð_ð _ð _ð _ð _ð _ð _ð _r   r¨  c                  ó>   ‡ — e Zd ZdZddˆ fd„Zd	„ Zd
„ Z	 ddd„Zˆ xZS )ÚSequenceParallelzd
    Sequence Parallel: input/output sharded on sequence dimension.
    Weights are replicated.
    r.   FÚsequence_dimr9   Úuse_local_outputrO  c                óH   •—  t          ¦   «         j        di |¤Ž || _        d S rQ  )rR  r(  r¿  )r'  r¿  rÀ  Úuse_dtensorr;  rS  s        €r   r(  zSequenceParallel.__init__1  s.   ø€ Ø�‰ŒÔÐ"Ð"˜6Ð"Ð"Ð"Ø(ˆÔÐÐr   c                ó:   — |r|d         n|}t          ||¦  «        S rX  )ró   rZ  s        r   r-  z"SequenceParallel._prepare_input_fn5  s&   € Ø$*Ð6�v˜a”y�y°ˆõ ˜,¨Ñ4Ô4Ð4r   c                ó"   — t          ||¦  «        S rµ   )r  r0  s       r   r1  z#SequenceParallel._prepare_output_fn;  s   € Ý˜g {Ñ3Ô3Ð3r   Nr‹   rš   r¯   r   rR   c                ó<   — |d                               ||¬¦  «        S rp  rq  r3  s        r   r4  zSequenceParallel.shard_tensor>  ó   € ð �SŒz�}Š} F°%ˆ}Ñ8Ô8Ð8r   )r.   FF)r¿  r9   rÀ  rO  rE  rF  ©	rá   râ   rã   rä   r(  r-  r1  r4  rg  rh  s   @r   r¾  r¾  +  sŠ   ø€ € € € € ðð ð
)ð )ð )ð )ð )ð )ð )ð5ð 5ð 5ð4ð 4ð 4ð VZð9ð 9ð 9ð 9ð 9ð 9ð 9ð 9ð 9r   r¾  c                  ó>   ‡ — e Zd ZdZˆ fd„Z	 ddd	„Zdd„Zdd„Zˆ xZS )ÚGroupedGemmParallelzb
    Applies Expert Parallelism to MoE experts by loading the correct experts on each device.
    c                ó:   •—  t          ¦   «         j        di |¤Ž d S rQ  ©rR  r(  ©r'  r;  rS  s     €r   r(  zGroupedGemmParallel.__init__I  ó&   ø€ Ø�‰ŒÔÐ"Ð"˜6Ð"Ð"Ð"Ð"Ð"r   Nr‹   rš   r¯   r   rR   c                ó¢  — | j         j        d         }|| j                             ¦   «         z  dk    r-t	          d|› d| j                             ¦   «         › d�¦  «        ‚|| j                             ¦   «         z  }|}| j        |z  }| j        dz   |z  }	t          |t          j        ¦  «        s| 	                    ¦   «         n|j        }
|�.||cxk    r|	k     r!n n|d d …          
                    |¬¦  «        S |€|||	…          
                    ||¬¦  «        S t          |
¦  «        dk    r|�d S |d d …          
                    ||¬¦  «        S )Nr   zAGlobal number of experts must be divisible by number of devices: ú % ú != 0r.   )rA   r_  )rŒ   r…   rH   rF   r1   r,   r   r   r¹   rº   rˆ   rŸ   )r'  r‹   r¯   rA   r³   Úglobal_num_expertsÚlocal_num_expertsrÅ   r”   rÆ   r…   s              r   r4  z GroupedGemmParallel.shard_tensorL  s•  € ð "Ô-Ô3°AÔ6ÐØ Ô 0× 5Ò 5Ñ 7Ô 7Ñ7¸1Ò<Ð<Ýð JÐTfð  Jð  JÐkoÔk{÷  lAò  lAñ  lCô  lCð  Jð  Jð  Jñô ð ð /°$Ô2B×2GÒ2GÑ2IÔ2IÑIÐØ&ˆ
Ø”	˜JÑ&ˆØŒy˜1‰} 
Ñ*ˆå)3°E½5¼<Ñ)HÔ)HÐY�—’Ñ!Ô!Ð!ÈeÌkˆØÐ! e¨zÐ&?Ð&?Ò&?Ð&?¸CÒ&?Ð&?Ð&?Ð&?Ð&?à˜˜˜”8—;’; f�;Ñ-Ô-Ð-ØÐØ˜˜s˜Ô#×&Ò&¨f¸EÐ&ÑBÔBÐBÝ�‰ZŒZ˜1Š_ˆ_ Ð!7Ø�4à˜˜˜”8—;’; f°E�;Ñ:Ô:Ð:r   r=  r>  r?  c                ó�   — | j                              ¦   «         }t          |¦  «        }|d         |z  }||d<   t          |¦  «        S rX  )rH   rF   r|   rÁ   )r'  r=  r-   r…   rÒ  s        r   rA  z.GroupedGemmParallel.get_expected_sharded_shaped  sG   € àÔ%×*Ò*Ñ,Ô,ˆ
Ý�ZÑ Ô ˆØ! !œH¨
Ñ2ÐØ$ˆˆa‰Ý�U‰|Œ|Ðr   r  r  c                óŒ   — t          |d¦  «        r3|                      | j        j        d         f¦  «        d         |_        d S d S )NÚnum_expertsr   )r   rA  rŒ   r…   rÕ  rC  s     r   rD  z,GroupedGemmParallel.update_module_attributesl  sR   € Ý�6˜=Ñ)Ô)ð 	bØ!%×!@Ò!@À$ÔBRÔBXÐYZÔB[ÐA]Ñ!^Ô!^Ð_`Ô!aˆFÔÐÐð	bð 	br   rE  rF  rJ  rK  )	rá   râ   rã   rä   r(  r4  rA  rD  rg  rh  s   @r   rÉ  rÉ  D  s’   ø€ € € € € ðð ð#ð #ð #ð #ð #ð VZð;ð ;ð ;ð ;ð ;ð0ð ð ð ðbð bð bð bð bð bð bð br   rÉ  c                  ó:   ‡ — e Zd ZdZˆ fd„Zd„ Zd„ Z	 ddd„Zˆ xZS )ÚRouterParallelzQ
    Allows to reshape the router scores to support running expert parallel.
    c                ó:   •—  t          ¦   «         j        di |¤Ž d S rQ  rË  rÌ  s     €r   r(  zRouterParallel.__init__v  rÍ  r   c                ó   — |S rµ   rw   r,  s       r   r-  z RouterParallel._prepare_input_fny  rl  r   c                ó¤  — |                      ¦   «         |                     ¦   «         }}t          |dd¦  «        }|€ t          t          |dd¦  «        dd¦  «        }|€%t          dt	          |¦  «        j        › d�¦  «        ‚||z  dk    rt          d|› d|› d	�¦  «        ‚||z  }|^}}	}
}|
|z  |k    }|	                     |d
¦  «        }	|
                     |d¦  «        }
|dk    rt          j	        |
|¦  «        }
n2|
                     |
dk    d¦  «                             |
dk     d¦  «        }
|
                     |
dk    |¦  «        }
||	|
g|¢R S )u	  
        Remap global expert indices to local and zero out non-local scores.

        Example: 4 tokens, top_k=4, 128 experts, EP=8. num_local_experts = 128/8 = 16.

        Router produces (all ranks see the same values):
            router_scores:  (4, 4)  â€” top-k routing weights
            router_indices: (4, 4)  â€” global expert IDs
                [ 52,  42, 119,  67],
                [102,  89,  61,  40],
                [ 82, 103,   4,  34],
                [ 93,  23, 109,  11],

        Each index maps to a rank: index // 16 gives the owning rank.
                [3, 2, 7, 4],
                [6, 5, 3, 2],
                [5, 6, 0, 2],
                [5, 1, 6, 0],

        For rank 0 (owns experts 0-15), we remap local indices with fmod and
        fill non-local with sentinel=16 (used for one_hot masking):
            router_indices (rank 0):
                [ 16, 16, 16, 16],
                [ 16, 16, 16, 16],
                [ 16, 16,  4, 16],
                [ 16, 16, 16, 11],

        Scores for non-local experts are zeroed out via masked_fill:
            router_scores (rank 0):
                [0.0, 0.0, 0.0, 0.0],
                [0.0, 0.0, 0.0, 0.0],
                [0.0, 0.0, 0.3, 0.0],    â†� only expert 4 (local) keeps its score
                [0.0, 0.0, 0.0, 0.1],    â†� only expert 11 (local) keeps its score

        both router_scores and router_indices stay (seq, top_k) shape.
        They are paired element-wise: scores[i] is the weight for indices[i].
        All expert forward implementations (grouped_mm, batched_mm, eager) flatten
        both with reshape(-1) and rely on this pairing. Changing the shape of one
        without the other breaks routing!

        Each rank believes it is alone and computes only its part of the hidden states.
        The sentinel index (num_local_experts) is skipped by one_hot encoding or clamped
        + masked in grouped_mm/batched_mm. After the expert forward, an all_reduce sums
        partial outputs across EP ranks to produce the full result.
        rÕ  Nrˆ  zRouter module z. is missing num_experts and config.num_expertsr   z>The number of experts must be divisible by number of ep_size: rÏ  rÐ  g        r„   r.   )
rò   rF   r7   r‰  r5   rá   r1   Úmasked_fillr   Úfmod)r'  r  r  rH   Úep_rankÚep_sizerÕ  Únum_local_expertsÚrouter_logitsÚrouter_scoresÚrouter_indicesÚextra_outputsÚnon_local_masks                r   r1  z!RouterParallel._prepare_output_fn|  s“  € ð\ '×5Ò5Ñ7Ô7¸×9IÒ9IÑ9KÔ9K�ˆÝ˜c =°$Ñ7Ô7ˆØÐÝ!¥'¨#¨x¸Ñ">Ô">ÀÈtÑTÔTˆKØÐÝ Ð!tµ$°s±)´)Ô2DÐ!tÐ!tÐ!tÑuÔuÐuà˜Ñ  AÒ%Ð%ÝØoÐQ\ÐoÐoÐahÐoÐoÐoñô ð ð (¨7Ñ2ÐàGNÐDˆ�} n°}Ø(Ð,=Ñ=À'ÒIˆØ%×1Ò1°.À#ÑFÔFˆØ'×3Ò3°NÀBÑGÔGˆà˜qÒ Ð Ý"œZ¨Ð8IÑJÔJˆNˆNà+×7Ò7¸ÈÒ8JÈAÑNÔN×ZÒZÐ[iÐlmÒ[mÐoqÑrÔrˆNØ'×3Ò3°NÀbÒ4HÐJ[Ñ\Ô\ˆØ˜m¨^ÐK¸mÐKÐKÐKr   Nr‹   rš   r¯   r   rR   c                ó<   — |d                               ||¬¦  «        S rp  rq  r3  s        r   r4  zRouterParallel.shard_tensorÃ  rÆ  r   rE  rF  rÇ  rh  s   @r   r×  r×  q  s‡   ø€ € € € € ðð ð#ð #ð #ð #ð #ðð ð ðELð ELð ELðP VZð9ð 9ð 9ð 9ð 9ð 9ð 9ð 9ð 9r   r×  c                  ó   — e Zd ZdZd„ ZdS )ÚRouterParallelMegaMoeu‹  Router TP plan used with DeepGEMM Mega MoE.

    Mega MoE handles EP dispatch inside the kernel and wants raw global expert ids
    with unmasked routing weights, so the router doesn't pre-shard per EP rank like
    `RouterParallel._prepare_output_fn` does. The quantizer's `update_tp_plan` swaps
    `"ep_router"` â†’ `"megamoe_router"` when `experts_implementation == "deepgemm_megamoe"`.
    c                ó   — |S rµ   rw   r0  s       r   r1  z(RouterParallelMegaMoe._prepare_output_fnÒ  rn  r   N)rá   râ   rã   rä   r1  rw   r   r   rç  rç  É  s-   € € € € € ðð ðð ð ð ð r   rç  c                  ó:   ‡ — e Zd ZdZˆ fd„Zd„ Zd„ Z	 ddd„Zˆ xZS )ÚMoeTensorParalellExpertsa7  
    Note: For tensor parallel, the MoEExpertsParallel TP layer handles gradient sync:
        - all_reduce_backward on hidden_states (for colwise gate_up_proj gradient)
        - all_reduce_backward on top_k_weights (for router gradient)
        - all_reduce_forward on output (for partial expert outputs)
    c                ó:   •—  t          ¦   «         j        di |¤Ž d S rQ  rË  rÌ  s     €r   r(  z!MoeTensorParalellExperts.__init__Þ  rÍ  r   c                ó|   — |d         }|d         }|d         }t          ||¦  «        }t          ||¦  «        }|||fS ©Nr   r.   r   rY  ©r'  r  r  rH   Úhidden_statesÚtop_k_indexÚtop_k_weightss          r   r-  z*MoeTensorParalellExperts._prepare_input_fná  sL   € à˜qœ	ˆØ˜Q”iˆØ˜qœ	ˆõ ,¨M¸;ÑGÔGˆõ
 ,¨M¸;ÑGÔGˆà˜k¨=Ð8Ð8r   c                ó"   — t          ||¦  «        S rµ   r€  r0  s       r   r1  z+MoeTensorParalellExperts._prepare_output_fnñ  s   € å! '¨;Ñ7Ô7Ð7r   Nr‹   rš   r¯   r   rR   c                ó<   — |d                               ||¬¦  «        S rp  rq  r3  s        r   r4  z%MoeTensorParalellExperts.shard_tensorõ  s   € ð
 �SŒz�}Š} F°%ˆ}Ñ8Ô8Ð8r   rE  rF  rÇ  rh  s   @r   rê  rê  Ö  s€   ø€ € € € € ðð ð#ð #ð #ð #ð #ð9ð 9ð 9ð 8ð 8ð 8ð
 VZð9ð 9ð 9ð 9ð 9ð 9ð 9ð 9ð 9r   rê  c                  ó*   — e Zd ZdZd„ Zd„ Z	 ddd
„ZdS )ÚMoeTensorParalellMegaMoeExpertsuI  TP layer for DeepGEMM Mega MoE experts.

    Mega MoE is inference-only (the kernel has no backward) and handles EP dispatch +
    combine + per-rank token sharding internally â€” so we skip the gradient-sync hooks
    that the regular `MoeTensorParalellExperts` would apply, and we forward the EP
    `process_group` into the module so the symm-buffer rendezvous can run on first
    forward. The quantizer's `update_tp_plan` swaps the experts plan key from
    `"moe_tp_experts"` to `"megamoe_experts"` when
    `from_pretrained(..., experts_implementation="deepgemm_megamoe")`.
    c                ób   — |d         |d         |d         }}}||||                      ¦   «         fS rí  )rÝ   rî  s          r   r-  z1MoeTensorParalellMegaMoeExperts._prepare_input_fn	  s7   € Ø4:¸1´I¸vÀa¼yÈ&ÐQRÌ) M�{ˆØ˜k¨=¸+×:OÒ:OÑ:QÔ:QÐQÐQr   c                ó   — |S rµ   rw   r0  s       r   r1  z2MoeTensorParalellMegaMoeExperts._prepare_output_fn  s   € àˆr   Nr‹   rš   r¯   r   rR   c                ó<   — |d                               ||¬¦  «        S rp  rq  r3  s        r   r4  z,MoeTensorParalellMegaMoeExperts.shard_tensor  rÆ  r   rE  rF  )rá   râ   rã   rä   r-  r1  r4  rw   r   r   rõ  rõ  ý  s_   € € € € € ð	ð 	ðRð Rð Rðð ð ð
 VZð9ð 9ð 9ð 9ð 9ð 9ð 9r   rõ  c                  ó&   — e Zd ZdZd„ Zdd„Zd„ ZdS )ÚMoeIdentityExpertParallelaM  
    TP class for zero/identity experts in MoE layers.

    Under TP, the parent MoeTensorParalellExperts does all_reduce_forward (sum)
    on the expert module output. Identity experts produce the same output on
    every rank, so the sum gives world_size * output. This class divides the
    input by world_size to compensate.
    c                óH   — |r|d         n|}||                      ¦   «         z  S rX  )rF   rZ  s        r   r-  z+MoeIdentityExpertParallel._prepare_input_fn!  s+   € Ø$*Ð6�v˜a”y�y°ˆà˜k×.Ò.Ñ0Ô0Ñ0Ð0r   Nc                ó<   — |d                               ||¬¦  «        S rp  rq  r3  s        r   r4  z&MoeIdentityExpertParallel.shard_tensor&  rr  r   c                ó4   — t          ||| j        ¬¦  «         d S )N)r  )r#  r-  r:  s       r   r<  z+MoeIdentityExpertParallel.prepare_module_tp)  s    € Ý˜& +¸Ô8NÐOÑOÔOÐOÐOÐOr   rE  )rá   râ   rã   rä   r-  r4  r<  rw   r   r   rú  rú    sU   € € € € € ðð ð1ð 1ð 1ð
9ð 9ð 9ð 9ðPð Pð Pð Pð Pr   rú  c                  óN  — e Zd ZU  e¦   «         rÓerÑi d ed¬¦  «        “d ed¬¦  «        “d ed¬¦  «        “d	 e¦   «         “d
 e¦   «         “d ed¬¦  «        “d e¦   «         “d e	¦   «         “d e
¦   «         “d e¦   «         “d e¦   «         “d e¦   «         “d e¦   «         “d e¦   «         “d e¦   «         “d e¦   «         “d e¦   «         “d e¦   «         i¥ni ZdddddddddddddœZded<   dddddddddddddœZded<   ed&d$„¦   «         Zed&d%„¦   «         ZdS )'ÚParallelInterfaceÚembedding_rowwiser   r©  Úembedding_colwiser.   Úcolwise_gather_outputT)rN  ÚcolwiseÚrowwiseÚrowwise_split_input)r‘  Úpacked_colwiseÚpacked_rowwiseÚsequence_parallelÚgrouped_gemmÚ	ep_routerÚmegamoe_routerÚmoe_tp_expertsÚmegamoe_expertsÚmoe_identity_expertÚreplicated_with_grad_allreduceÚmla_kv_a_projrÚ   rƒ   r„   N)r  r  r  r  r  r  r   r  r  r  r  rÚ   zdict[str, int | None]Úplan_to_weight_dimÚplan_to_bias_dimÚkeyrQ   Úvaluer   c                ó   — || j         |<   d S rµ   )r  ©Úclsr  r  s      r   Úregister_plan_to_weight_dimz-ParallelInterface.register_plan_to_weight_dimk  s   € à&+ˆÔ˜sÑ#Ð#Ð#r   c                ó   — || j         |<   d S rµ   )r  r  s      r   Úregister_plan_to_bias_dimz+ParallelInterface.register_plan_to_bias_dimo  s   € à$)ˆÔ˜SÑ!Ð!Ð!r   )r  rQ   r  r   )rá   râ   rã   r
   Ú_torch_distributed_availabler¨  rM  r�  rž  r£  r¾  rÉ  r×  rç  rê  rõ  rú  rj  r…  r}  Ú_global_mappingr  Ú__annotations__r  Úclassmethodr  r  rw   r   r   rÿ  rÿ  -  s­  € € € € € € ð0 ÐÑÔð)	ð( %Að)	ð 	
ØÐ!2Ð!2È!Ð!LÑ!LÔ!Lð	
àÐ!2Ð!2È!Ð!LÑ!LÔ!Lð	
ð $ _ _À4Ð%HÑ%HÔ%Hð	
ð ��Ñ(Ô(ð		
ð
 ��Ñ(Ô(ð	
ð " ? ?¸tÐ#DÑ#DÔ#Dð	
ð Ð3Ð3Ñ5Ô5ð	
ð Ð3Ð3Ñ5Ô5ð	
ð  Ð!1Ð!1Ñ!3Ô!3ð	
ð Ð/Ð/Ñ1Ô1ð	
ð ˜˜Ñ)Ô)ð	
ð Ð3Ð3Ñ5Ô5ð	
ð Ð6Ð6Ñ8Ô8ð	
ð Ð>Ð>Ñ@Ô@ð	
ð "Ð#<Ð#<Ñ#>Ô#>ð	
ð  -Ð.IÐ.IÑ.KÔ.Kð!	
ð" Ð/Ð/Ñ1Ô1ð#	
ð$ Ð+Ð+Ñ-Ô-ð%	
ð 	
ð 	
ð* ð- ð: Ø!#ØØØ!ØØØØ!Ø*.ØØð1ð 1Ðð ð ð ñ ð" Ø!#ØØØ#ØØ!Ø!Ø!Ø*.ØØð/ð /Ðð ð ð ñ ð ð,ð ,ð ,ñ „[ð,ð ð*ð *ð *ñ „[ð*ð *ð *r   rÿ  ÚALL_PARALLEL_STYLESÚlocal_tensorÚ	shard_dimrH   údist.device_mesh.DeviceMeshc                óD  ‡ — |                      ¦   «         }d|j        pi v r|                     d¦  «        nd}|dk     r
‰ j        |z   }ˆ fd„t	          |¦  «        D ¦   «         }t          j        |‰                      ¦   «         |¬¦  «         t          j	        ||¬¦  «        S )a~  
    All-gather a sharded tensor along the specified dimension to reconstruct the full tensor.

    Args:
        local_tensor: The local shard of the tensor on this rank
        shard_dim: The dimension along which the tensor was sharded
        device_mesh: The device mesh for distributed communication

    Returns:
        The full reconstructed tensor (same on all ranks)
    r/   Nr   c                ó8   •— g | ]}t          j        ‰¦  «        ‘ŒS rw   rî   )rx   rð   r   s     €r   r{   z&gather_full_tensor.<locals>.<listcomp>“  s$   ø€ ÐRÐRÐR¸1�Ô(¨Ñ6Ô6ÐRÐRÐRr   rU   rÊ   )
rF   rE   rÝ   rD   r†   rÙ   ró   rØ   r   rô   )r   r!  rH   r-   Úprocess_groupÚgathered_tensorss   `     r   Úgather_full_tensorr'  |  s»   ø€ ð ×!Ò!Ñ#Ô#€Jà37¸KÔ<VÐ<\ÐZ\Ð3]Ð3]�K×)Ò)¨$Ñ/Ô/Ð/Ðcg€Mð �1‚}€}Ø Ô%¨	Ñ1ˆ	ð SÐRÐRÐRÅÀjÑ@QÔ@QÐRÑRÔRÐÝ„OÐ$ l×&=Ò&=Ñ&?Ô&?À}ÐUÑUÔUÐUõ Œ9Ð%¨9Ð5Ñ5Ô5Ð5r   Ú
state_dictúdict[str, torch.Tensor]c                óê  — t           j        }t           j        }i }|                      ¦   «         D �]B\  }}d|v r|                     dd¦  «        d         n|}	d|v r|                     dd¦  «        d         nd}
t          j        dd|	¦  «        }t          j        dd|¦  «        }d}||v r	||         }n9||v r	||         }n,d|v r(|                     dd¦  «        d         }||v r||         }|�||vr|||<   ŒÊ|
dk    r|                     |¦  «        }n|                     |¦  «        }|€|||<   �Œt          |||¦  «        }|dv rt          |||d	¦  «        }| 
                    ¦   «         ||<   �ŒD|S )
a*  
    Gather sharded tensors to reconstruct full tensors for saving.

    This function all-gathers each sharded tensor along its shard dimension
    to reconstruct the full unsharded tensor for checkpoint saving.

    Args:
        state_dict: The model state dict with local sharded tensors
        tp_plan: The tensor parallel plan mapping layer patterns to shard styles
        device_mesh: The device mesh for distributed communication
        tp_size: The tensor parallel world size

    Returns:
        State dict with full (gathered) tensors
    r`   r.   r   Nú\d+Ú*r”  )r  r  r   )r  r  r  Úitemsra   rY   rZ   r<   r'  r®   rØ   )r(  r   rH   r   r  r  Úresultr  r   Ú
param_nameÚ
param_typerc   Úgeneric_full_keyÚcurrent_planÚparent_param_namer!  Úfull_tensors                    r   Úgather_state_dict_for_saver5  š  sÒ  € õ, -Ô?ÐÝ*Ô;Ðà€FØ!×'Ò'Ñ)Ô)ð (/ñ (/‰ˆˆVà.1°S¨j¨j�S—Z’Z  QÑ'Ô'¨Ô*Ð*¸cˆ
Ø.1°S¨j¨j�S—Z’Z  QÑ'Ô'¨Ô*Ð*¸dˆ
ÝœV F¨C°Ñ<Ô<Ðåœ6 &¨#¨sÑ3Ô3Ðð ˆØ˜wÐ&Ð&à"Ð#3Ô4ˆLˆLØ 7Ð*Ð*Ø"Ð#5Ô6ˆLˆLØÐ&Ð&Ð&Ø 2× 9Ò 9¸#¸qÑ AÔ AÀ!Ô DÐØ  GÐ+Ð+Ø&Ð'8Ô9�àÐ <Ð7IÐ#IÐ#Ià ˆF�3‰KØð ˜ÒÐØ(×,Ò,¨\Ñ:Ô:ˆIˆIà*×.Ò.¨|Ñ<Ô<ˆIàÐà ˆF�3‰KÙõ )¨°¸KÑHÔHˆØÐ?Ð?Ð?Ý(¨°iÀÈ!ÑLÔLˆKØ!×,Ò,Ñ.Ô.ˆˆs‰‰à€Mr   c           	     ó>  ‡‡— ‰�˜t           ‰         }|                     ‰||¦  «         	 |                     ‰|| j        ¬¦  «         n:# t          $ r-}t
                               d|› d‰› d|› �¦  «         Y d}~nd}~ww xY w‰‰_        |‰_        ˆˆfd„‰_	        dS dS )a  
    This function is called in `PretrainedModel.post_init()`. It is responsible of adding hooks
    to the modules of the `model`, based on the `PretrainedModel._tp_plan`.

    This is the place where we add the `pre_forward` and `post_forwards` hooks. These are defined
    for each `TensorParallelLayer` as `_prepare_input_fn` and `_prepare_output_fn`.

    Args:
        model (`PretrainedModel`): The model containing the modules.
        module (`nn.Module`): The current module to which we want to add the hooks.
        current_module_plan (`str` or `None`): The tensor parallel plan for the current module, if any.
        layer_name (`str`): The qualified name of the current module.
        device_mesh (`dist.device_mesh.DeviceMesh`): The device mesh for distributed communication.

    N)rˆ  úTrying to prepare ú0, but it's not supported. Corresponding module: z Fix it's TP plan: c                 ó6   •— ‰                      ¦   «         › d‰ › �S )Nz

TP Plan: )Ú__repr__)Úcurrent_module_planr  s   €€r   rX   z5add_tensor_parallel_hooks_to_module.<locals>.<lambda>  s    ø€  V§_¢_Ñ%6Ô%6Ð"XÐ"XÐCVÐ"XÐ"X€ r   )
r  r8  r<  rˆ  r+  ÚloggerÚwarningÚ_hf_tp_planÚ_hf_device_meshr:  )Úmodelr  r;  r6  rH   Útp_layerrL   s    ``    r   Ú#add_tensor_parallel_hooks_to_modulerB  á  sþ   øø€ ð, Ð&Ý&Ð':Ô;ˆØ× Ò  ¨°jÑAÔAÐAð	Ø×&Ò& v¨{À5Ä<Ð&ÑPÔPÐPÐPøÝ"ð 	ð 	ð 	Ý�NŠNð Zð ð Ðagð ð Øðð ñô ð ð ð ð ð ð øøøøð	øøøð 1ˆÔØ!,ˆÔØXÐXÐXÐXÐXˆŒˆˆð 'Ð&s   ªA Á
A?Á#A:Á:A?c                ó®  — d|v r|                      dd¦  «        n|\  }}	| j        pi }
|                      |¦  «        }t          |¦  «        }t	          ||
¦  «        }t          j        ¦   «         dk    rA|€t                               d|› d�¦  «         n t                               d|› d|› �¦  «         d}|�…	 t          |         }||_
        ||_        ||_        |                     |d||¬¦  «        }|r|                     ¦   «         }nO# t          $ r%}t!          d	|› d
|› d|› d|› �¦  «         Y d}~n%d}~ww xY w|dd…                              |¦  «        }t%          |t&          j        j        ¦  «        s3t&          j                             ||                     ¦   «         ¬¦  «        }t/          ||	|¦  «         |�|                     |¦  «         |S )aÎ  
    This function is called in `from_pretrained` when loading a model's checkpoints.
    It receives the pointer to the parameter (or the parameter itself) and takes care of "sharding".
    All process run this function, so they just load the partition of the tensor that they require.

    Main uses cases:
    - column / rowise parallelism, you just shard all the weights of the layer (weight and bias)
    - packed layers: you slice the weights, then shard like above
    - custom operation:
        - you want to add an all-gather at the end of a local layer.
        - you want to have a layer that is isolated from the rest of the world (because torch.DTensor does not work well with `.view` for instance)

    r`   r.   r   NzTensor sharding plan for z+ not found, using default 'replicate' plan.z: )r¯   r³   rA   r7  r8  z" Fix it's TP plan, current layer: z : )Úrequires_grad)ra   r   Úget_submoduler9   re   rÙ   Úget_rankr<  Úinfor  rŒ   rH   r,   r4  rØ   r+  Úprintrˆ   r   r   r   Ú	ParameterÚis_floating_pointÚsetattrrD  )r@  r‹   rŒ   r\   Úparam_casting_dtypeÚis_contiguousr,   rH   r/  r0  r   Úmodule_to_tpÚcurrent_shard_planrA  rL   s                  r   Úshard_and_distribute_modulerP    sO  € ð  ?BÀ^Ð>SÐ>S˜^×2Ò2°3¸Ñ:Ô:Ð:ÐYgÑ€J�
ØŒmÐ!˜r€GØ×&Ò& zÑ2Ô2€LÝˆt‰9Œ9€DÝ/°ÀÑHÔHÐå„}�„˜!ÒÐØÐ%Ý�KŠKÐk°JÐkÐkÐkÑlÔlÐlÐlå�KŠKÐV°JÐVÐVÐBTÐVÐVÑWÔWÐWà€HØÐ%ð	Ý*Ð+=Ô>ˆHØ#.ˆHÔ Ø#.ˆHÔ Ø ˆHŒMØ×)Ò)¨%¸DÐH[ÐdhÐ)ÑiÔiˆEØð +Ø×(Ò(Ñ*Ô*�øøÝ"ð 	ð 	ð 	Ýð f ^ð  fð  fÐeqð  fð  fð  V^ð  fð  fð  cdð  fð  fñô ð ð ð ð ð ð øøøøð	øøøð
 �a�a�a”—’Ð/Ñ0Ô0ˆõ �e�UœXÔ/Ñ0Ô0ð YÝ”×"Ò" 5¸×8UÒ8UÑ8WÔ8WÐ"ÑXÔXˆÝˆL˜* eÑ,Ô,Ð,ØÐØ×)Ò)¨,Ñ7Ô7Ð7Ø€Ls   Â:AD Ä
D;ÄD6Ä6D;Úexpected_keysú	list[str]údict[str, str] | Nonec                óÎ  — |€dS d„ | D ¦   «         }t          |¦  «        }|                     ¦   «         }|D ]¹}d|v r|                     dd¦  «        d         n|}t          j        dd|¦  «        }||v r,|                     |d¦  «         |                     |¦  «         Œjd|v rK|                     dd¦  «        d         x}|v r+|                     |d¦  «         |                     |¦  «         Œºt          |¦  «        dk    rt           	                    d|› �¦  «         t          |¦  «        dk    r2t           	                    d	d
 
                    |¦  «        › �¦  «         dS dS )z�
    Verify the TP plan of the model, log a warning if the layers that were not sharded and the rules that were not applied.
    Nc                ó,   — h | ]}t          |¦  «        ’ŒS rw   )r[   )rx   r  s     r   ú	<setcomp>z!verify_tp_plan.<locals>.<setcomp>F  s!   € ÐSÐSÐS¸cÕ4°SÑ9Ô9ÐSÐSÐSr   r`   r.   r   r+  r,  z>The following TP rules were not applied on any of the layers: z'The following layers were not sharded: z, )ÚsetÚcopyra   rY   rZ   ÚpopÚdiscardrŸ   r<  r=  Újoin)	rQ  r   Úgeneric_keysÚunsharded_layersÚunused_rulesr  r/  rc   r3  s	            r   Úverify_tp_planr_  >  s˜  € ð
 €ØˆàSÐSÀ]ÐSÑSÔS€LÝ˜<Ñ(Ô(ÐØ—<’<‘>”>€Làð 	*ð 	*ˆØ.1°S¨j¨j�S—Z’Z  QÑ'Ô'¨Ô*Ð*¸cˆ
ÝœV F¨C°Ñ<Ô<Ðà Ð(Ð(Ø×ÒÐ/°Ñ6Ô6Ð6Ø×$Ò$ SÑ)Ô)Ð)Ð)ØÐ&Ð&Ð&ÐAS×AZÒAZÐ[^Ð`aÑAbÔAbÐcdÔAeÐ,eÐ,=ÐjqÐ+qÐ+qØ×ÒÐ.°Ñ5Ô5Ð5Ø×$Ò$ SÑ)Ô)Ð)øå
ˆ<ÑÔ˜1ÒÐÝ�ŠÐfÐXdÐfÐfÑgÔgÐgÝ
ÐÑÔ˜qÒ Ð Ý�ŠÐ^ÀÇÂÐK[ÑA\ÔA\Ð^Ð^Ñ_Ô_Ð_Ð_Ð_ð !Ð r   c                ó
  — || _         || _        |�5t          |t          ¦  «        rt	          j        |¦  «        }|| j        _        t          |t          ¦  «        r|| _        | j        }|�˜t          r‘| 
                    ¦   «         D ]%}|t          vrt          d|› dt          › �¦  «        ‚Œ&|                      ¦   «         D ]B\  }}t          |dd¦  «        s%t          ||d¬¦  «        }	t!          | ||	||¦  «         d|_        ŒC| S )z,Distribute a model according to the TP plan.Nz"Unsupported tensor parallel style z. Supported styles are Ú
_is_hookedF)r\   r   rb   T)Ú_tp_sizeÚ_device_meshr   Údictr   Ú	from_dictrˆ  Údistributed_configr   r  Úvaluesr  r1   Únamed_modulesr7   re   rB  ra  )
r@  r   rf  rH   r   Ú
model_planÚvrP   r  Úplans
             r   Údistribute_modelrl  [  s<  € à€E„NØ$€EÔØÐ%ÝÐ(­$Ñ/Ô/ð 	QÝ!2Ô!<Ð=OÑ!PÔ!PÐØ*<ˆŒÔ'å�'�4Ñ Ô ð  ØˆŒØ”€JØÐÕ">ÐØ×"Ò"Ñ$Ô$ð 	wð 	wˆAØÕ+Ð+Ð+Ý Ð!uÀaÐ!uÐ!uÕ`sÐ!uÐ!uÑvÔvÐvð ,à!×/Ò/Ñ1Ô1ð 
	%ð 
	%‰LˆD�&Ý˜6 <°Ñ7Ô7ð Ý-¸TÈ:ÐafÐgÑgÔg�Ý3ØØØØØñô ð ð !%ˆFÔÐØ€Lr   rE  )r   r   r   r   )rP   rQ   rR   rQ   )T)r\   rQ   r   r]   rR   r^   )rq   r9   rr   rs   rR   rt   )r   )
r™   rš   r›   r9   r-   r9   rœ   r9   rR   rš   rµ   )r¯   r   rI  )r   rš   r!  r9   rH   r"  rR   rš   )r(  r)  r   r]   r   r9   rR   r)  )rQ  rR  r   rS  )UÚ
__future__r   r»   r·   r:   rY   Ú	functoolsr   r   r   Úutilsr   r   Úutils.genericr	   Úutils.import_utilsr
   r   Útorch.distributedrÙ   r   Úis_availabler  Ú
get_loggerrá   r<  r   rO   r[   re   rO  Úuint8Úint8Úint16r‰   Úbfloat16Úint32Úfloat32Úfloat64r¿   Úfloat8_e4m3fnrŠ   r€   r˜   r®   rÈ   rÍ   ÚautogradÚFunctionrÏ   rç   rë   rü   r  r  r  ró   r  r  r#  r%  rM  rj  r}  r…  r�  rž  r£  r¨  r¾  rÉ  r×  rç  rê  rõ  rú  rÿ  r  r  r'  r5  rB  rP  r_  rl  rw   r   r   ú<module>r     sv  ðð #Ð "Ð "Ð "Ð "Ð "Ð "à €€€Ø €€€Ø 	€	€	€	Ø 	€	€	€	Ø Ð Ð Ð Ð Ð à +Ð +Ð +Ð +Ð +Ð +Ø 6Ð 6Ð 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø ,Ð ,Ð ,Ð ,Ð ,Ð ,Ø 3Ð 3Ð 3Ð 3Ð 3Ð 3ð ÐÑÔð DØ€L€L€LØ$Ð$Ð$Ð$Ð$Ð$ØÐÐÐÐÐð $)Ô#4×#AÒ#AÑ#CÔ#CÐ ð 
ˆÔ	˜HÑ	%Ô	%€ðð ð ð  dhðE,ð E,ð E,ð E,ð E,ðPEð Eð Eð Eðð ð ð ð ð. ÐÑÔð à”
ØŒkØŒjØŒ{ØŒ}Ø”ØŒ{ØŒ}ØŒ}ØŒ{ØÔ&ðð €Lð&ð &ð &ð &ð4A4ð A4ð A4ðP ð	> ð > ð > ð > ð > ðBe>ð e>ð e>ð e>ð e>ðP.ð .ð .ð4!ð !ð !ð !ð !˜œÔ0ñ !ô !ð !ð$!ð !ð !ð !ð !˜œÔ/ñ !ô !ð !ð/ð /ð /ð /ð /�”Ô(ñ /ô /ð /ðDGð Gð Gð Gð GˆUŒ^Ô$ñ Gô Gð GðD&Gð &Gð &Gð &Gð &G�U”^Ô,ñ &Gô &Gð &Gð\4ð 4ð 4ð
3ð 3ð 3ð
,ð ,ð ,ð
(ð (ð (ð
0ð 0ð 0ð ØØð	ð ð ð ð ð":ð :ð :ð :ð :ñ :ô :ð :ðz9]ð 9]ð 9]ð 9]ð 9]Ð)ñ 9]ô 9]ð 9]ðx;ð ;ð ;ð ;ð ;Ð"5ñ ;ô ;ð ;ð<Rð Rð Rð Rð RÐ+ñ Rô Rð Rð&-Rð -Rð -Rð -Rð -RÐ,ñ -Rô -Rð -Rð`?Kð ?Kð ?Kð ?Kð ?KÐ)ñ ?Kô ?Kð ?KðD8ð 8ð 8ð 8ð 8˜Oñ 8ô 8ð 8ð,8ð 8ð 8ð 8ð 8˜Oñ 8ô 8ð 8ð4N_ð N_ð N_ð N_ð N_Ð+ñ N_ô N_ð N_ðb9ð 9ð 9ð 9ð 9Ð*ñ 9ô 9ð 9ð2*bð *bð *bð *bð *bÐ-ñ *bô *bð *bðZU9ð U9ð U9ð U9ð U9Ð(ñ U9ô U9ð U9ðp
ð 
ð 
ð 
ð 
˜Nñ 
ô 
ð 
ð$9ð $9ð $9ð $9ð $9Ð2ñ $9ô $9ð $9ðN9ð 9ð 9ð 9ð 9Ð&9ñ 9ô 9ð 9ð4Pð Pð Pð Pð PÐ 3ñ Pô Pð Pð,D*ð D*ð D*ð D*ð D*Ð(ñ D*ô D*ð D*ðN *;Ð):Ñ)<Ô)<Ð Ð <Ð <Ð <Ñ <ð6ð 6ð 6ð 6ð<Dð Dð Dð DðN#Yð #Yð #YðL4ð 4ð 4ðn`ð `ð `ð `ð:ð ð ð ð r   