§
    ‚ŠtjÇj  ã                  ó€  — d dl mZ d dlmZ d dlmZ ddlmZ ddlm	Z	 ddl
mZmZmZmZ dd	lmZ dd
lmZ  e¦   «         r8d dlZej                             e¦  «        Zej                             e¦  «        Z ej        e¦  «        Z	 	 d8d9d„Zd:d„Zd;d„Zd;d„Zd„ Zd „ Z e¦   «         rVej                              d!ed"d#¬$¦  «         ej         !                    d!e¦  «         ej         "                    d!ee¬%¦  «         d<d&„Z#d;d'„Z$	 	 d8d=d(„Z%d:d)„Z& G d*„ d+e	¦  «        Z' e'¦   «         Z(d>d-„Z)	 d?e(d.ddd.d/œd@d7„Z*dS )Aé    )Úannotations)ÚCallable)Úwrapsé   )Úlogging)ÚGeneralInterface)Úis_torch_availableÚis_torch_greater_or_equalÚis_torch_less_or_equalÚis_torchdynamo_compilingé   )Údeepgemm_bf16_experts_forward)Úsonicmoe_experts_forwardNFÚinputútorch.TensorÚweightÚbiasútorch.Tensor | NoneÚis_transposedÚboolÚreturnc                ó&  — |r<t          j        |                      d¦  «        |¦  «                             d¦  «        }n;t          j        ||                      d¦  «        ¦  «                             d¦  «        }|�|                     |¦  «         |S )a¶  Batched linear layer supporting optional bias and transposed weights.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (batch_size, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (batch_size, output_dim, input_dim) if transposed is `False`,
            else of shape (batch_size, input_dim, output_dim).
        bias (`torch.Tensor`, *optional*):
            Bias tensor of shape (batch_size, output_dim). Default is `None`.
        is_transposed (`bool`, *optional*, defaults to `False`):
            Whether the weight tensor is transposed.
    Returns:
        `torch.Tensor`: Output tensor of shape (batch_size, output_dim).
    r   éÿÿÿÿ)ÚtorchÚbmmÚ	unsqueezeÚsqueezeÚadd_)r   r   r   r   Úouts        ú[/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/integrations/moe.pyÚ_batched_linearr!   T   s�   € ð* ð AåŒi˜Ÿš¨Ñ*Ô*¨FÑ3Ô3×;Ò;¸AÑ>Ô>ˆˆõ Œi˜ §¢°Ñ 3Ô 3Ñ4Ô4×<Ò<¸RÑ@Ô@ˆàÐØ�Š�‰Œˆà€Jó    Úselfútorch.nn.ModuleÚhidden_statesÚtop_k_indexÚtop_k_weightsc                óž  — |                      d¦  «        }|                      d¦  «        }|                      d¦  «        }|                     |d¬¦  «        }|                     d¦  «        }|                     d¦  «        }	|	                     d| j        dz
  ¦  «        }	| j        r$| j        |	         }
| j        r| j        |	         nd }n#| j	        |	         }
| j        r| j
        |	         nd }t          ||
|| j        ¬¦  «        }| j        r|                      |¦  «        }n|                      |¦  «        }| j        |	         }
| j        r| j        |	         nd }t          ||
|| j        ¬¦  «        }||                     d¦  «        z  }|                     |||¦  «                             d¬¦  «        }|                     |j        ¦  «        S )Nr   r   ©Údimr   ©r   r   )ÚsizeÚrepeat_interleaveÚreshapeÚclampÚnum_expertsÚhas_gateÚgate_up_projÚhas_biasÚgate_up_proj_biasÚup_projÚup_proj_biasr!   r   Ú_apply_gateÚact_fnÚ	down_projÚdown_proj_biasr   ÚviewÚsumÚtoÚdtype)r#   r%   r&   r'   Ú	num_top_kÚ
num_tokensÚ
hidden_dimÚselected_hidden_statesÚsample_weightsÚ
expert_idsÚselected_weightsÚselected_biasesÚproj_outÚweighted_outÚfinal_hidden_statess                  r    Úbatched_mm_experts_forwardrJ   v   sú  € ð × Ò  Ñ$Ô$€IØ×#Ò# AÑ&Ô&€JØ×#Ò# BÑ'Ô'€Jð +×<Ò<¸YÈAÐ<ÑNÔNÐØ"×*Ò*¨2Ñ.Ô.€NØ×$Ò$ RÑ(Ô(€Jð ×!Ò! ! TÔ%5¸Ñ%9Ñ:Ô:€Jð „}ð SØÔ,¨ZÔ8ÐØ@DÄÐW˜$Ô0°Ô<Ð<ÐSWˆˆàœ<¨
Ô3ÐØ;?¼=ÐR˜$Ô+¨JÔ7Ð7Èdˆõ ØÐ 0°ÐVZÔVhðñ ô €Hð
 „}ð )à×#Ò# HÑ-Ô-ˆˆð —;’;˜xÑ(Ô(ˆð ”~ jÔ1ÐØ9=¼ÐP�dÔ)¨*Ô5Ð5ÈD€Oõ ØÐ"¨ÈÔHZðñ ô €Hð
 ˜n×6Ò6°rÑ:Ô:Ñ:€Lð '×+Ò+¨J¸	À:ÑNÔN×RÒRÐWXÐRÑYÔYÐà×!Ò! -Ô"5Ñ6Ô6Ð6r"   Úoffsc                óT  — t          j        |                      d¦  «        |                     d¦  «        | j        | j        ¬¦  «        }d}t          |                     ¦   «         ¦  «        D ];\  }}||k    rŒt          j        | ||…         ||         |||…         ¬¦  «         |}Œ<|S )a(  
    Fallback grouped matrix multiplication used when `torch.nn.functional.grouped_mm` and `torch._grouped_mm`
    are unavailable or incompatible with `torch.compile` (e.g. non-bfloat16 weights).

    Args:
        input (`torch.Tensor`): Input of shape (S, input_dim), sorted by expert id.
        weight (`torch.Tensor`): Expert weights of shape (num_experts, input_dim, output_dim).
        offs (`torch.Tensor`): Cumulative token counts per expert of shape (num_experts,).
    Returns:
        `torch.Tensor`: Output of shape (S, output_dim).
    r   r   ©Údevicer>   ©r   )r   Úzerosr,   rN   r>   Ú	enumerateÚtolistÚmm)r   r   rK   ÚoutputÚstartÚiÚends          r    Ú_grouped_mm_fallbackrX   ¹   s¥   € õ Œ[˜Ÿš A™œ¨¯ª°A©¬¸u¼|ÐSXÔS^Ð_Ñ_Ô_€Fà€Eõ ˜DŸKšK™MœMÑ*Ô*ð ð ‰ˆˆ3Ø�CŠ<ˆ<ØÝŒ��u˜S�yÔ! 6¨!¤9°&¸¸s¸Ô2CÐDÑDÔDÐDØˆˆà€Mr"   c                óÆ  — |                       ¦   «         dk    sJ dt          | j        ¦  «        › �¦   «         ‚|                      ¦   «         dk    sJ dt          |j        ¦  «        › �¦   «         ‚|                      ¦   «         dk    sJ dt          |j        ¦  «        › �¦   «         ‚|                     d¦  «        |                     d¦  «        k    s6J d|                     d¦  «        › d	|                     d¦  «        › �¦   «         ‚|                      d¦  «        |                     d¦  «        k    s6J d
|                      d¦  «        › d|                     d¦  «        › �¦   «         ‚|j        t
          j        t
          j        fv sJ d|j        › �¦   «         ‚t          j        |                      d¦  «        |                     d¦  «        | j	        | j        ¬¦  «        S )zRShape/dtype inference stub for `_grouped_mm_fallback` required by `torch.compile`.r   z+input must be 2D (S, input_dim), got shape é   zBweight must be 3D (num_experts, input_dim, output_dim), got shape r   z*offs must be 1D (num_experts,), got shape r   zoffs length z must match number of experts zinput_dim mismatch: input has z, weight has z$offs must be an integer tensor, got rM   )
r*   ÚtupleÚshaper,   r>   r   Úint32Úint64ÚemptyrN   ©r   r   rK   s      r    Ú_grouped_mm_fallback_fakera   Ó   s³  € à�9Š9‰;Œ;˜!ÒÐÐÐ_Í5ÐQVÔQ\ÑK]ÔK]Ð_Ð_ÑÔÐØ�:Š:‰<Œ<˜1ÒÐÐØbÍUÐSYÔS_ÑM`ÔM`ÐbÐbñ ÔÐð �8Š8‰:Œ:˜Š?ˆ?ˆ?Ð\ÍÈtÌzÑIZÔIZÐ\Ð\‰?Œ?ˆ?Ø�9Š9�Q‰<Œ<˜6Ÿ;š; q™>œ>Ò)Ð)Ð)Ð+v¸$¿)º)ÀA¹,¼,Ð+vÐ+vÐfl×fqÒfqÐrsÑftÔftÐ+vÐ+vÑ)Ô)Ð)Ø�:Š:�a‰=Œ=˜FŸKšK¨™NœNÒ*Ð*Ð*ØU¨¯ª°A©¬ÐUÐUÀVÇ[Â[ÐQRÁ^Ä^ÐUÐUñ +Ô*Ð*ð Œ:�%œ+¥u¤{Ð3Ð3Ð3Ð3Ð5hÐ\`Ô\fÐ5hÐ5hÑ3Ô3Ð3ÝŒ;�u—z’z !‘}”} f§k¢k°!¡n¤n¸U¼\ÐQVÔQ\Ð]Ñ]Ô]Ð]r"   c                ód   — |                       |d         |d         ¦  «         |d         | _        dS )zjSaves input and weight for backward; offs is stored directly as it is a non-differentiable integer tensor.r   r   r   N)Úsave_for_backwardrK   )ÚctxÚinputsrT   s      r    Ú"_grouped_mm_fallback_setup_contextrf   â   s/   € à×Ò˜& œ) V¨A¤YÑ/Ô/Ð/Ø�aŒy€C„H€H€Hr"   c                ó¦  — | j         \  }}t          j        |¦  «        }t          j        |¦  «        }d}t          | j                             ¦   «         ¦  «        D ]r\  }}||k    rŒt          j        |||…         ||         j        |||…         ¬¦  «         t          j        |||…         j        |||…         ||         ¬¦  «         |}Œs||dfS )zuBackward pass for `_grouped_mm_fallback`. Computes grad_input and grad_weight per expert group; offs has no gradient.r   rO   N)Úsaved_tensorsr   Ú
zeros_likerQ   rK   rR   rS   ÚT)	rd   Úgrad_outputr   r   Ú
grad_inputÚgrad_weightrU   rV   rW   s	            r    Ú_grouped_mm_fallback_backwardrn   è   sÛ   € àÔ%�M€Eˆ6ÝÔ! %Ñ(Ô(€JÝÔ" 6Ñ*Ô*€Kà€Eõ ˜CœHŸOšOÑ-Ô-Ñ.Ô.ð ð ‰ˆˆ3Ø�CŠ<ˆ<ØÝŒ�˜U 3˜YÔ'¨°¬¬¸*ÀUÈ3ÀYÔ:OÐPÑPÔPÐPÝŒ��u˜S�yÔ!Ô# [°°s°Ô%;ÀÈQÄÐPÑPÔPÐPØˆˆà�{ DÐ(Ð(r"   z!transformers::grouped_mm_fallback© z4(Tensor input, Tensor weight, Tensor offs) -> Tensor)Úmutates_argsÚschema)Úsetup_contextc                óB  — t          ¦   «         r|j        t          j        k    sx|j        j        dk    rGt          dd¬¦  «        r6|                     ¦   «         dz  dk    s<|                      ¦   «         dz  dk    s!|j        j        dk    rt          dd¬¦  «        rdS |j        j        d	k    r¿t          t          j	        j
        d
¦  «        r(t          j                             |j        ¦  «        dk    S t          t          d¦  «        rat          dd¬¦  «        r(t          j                             |j        ¦  «        dk    S t          j                             |j        ¦  «        dk    S dS t          t          j	        j
        d
¦  «        pt          t          d¦  «        S )a  
    Check if torch.nn.functional.grouped_mm or torch._grouped_mm can be used based on availability and compatibility with torch.compile.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (S, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (num_experts, input_dim, output_dim).
        offs (`torch.Tensor`):
            Offsets tensor indicating the boundaries of each group in the input tensor.
    Returns:
        `bool`: True if grouped_mm can be used, False otherwise.
    Úcpuz2.10.0T)Ú
accept_devé   r   z2.8.0FÚcudaÚ
grouped_mm)é   r   Ú_grouped_mmz2.9)é	   r   )r   r>   r   Úbfloat16rN   Útyper   Údata_ptrÚhasattrÚnnÚ
functionalrw   Úget_device_capabilityr
   r`   s      r    Ú_can_use_grouped_mmrƒ   
  sx  € õ  
"Ñ	#Ô	#ðØ(.¬½¼Ò(FÐ(FØŒ=Ô Ò&Ð&Ý" 8¸Ð=Ñ=Ô=ð 'à�_Š_ÑÔ Ñ# qÒ(Ð(¨E¯NªNÑ,<Ô,<¸rÑ,AÀQÒ,FÐ,FØŒ=Ô Ò&Ð&Ý" 7°tÐ<Ñ<Ô<ð 'ð ˆuð
 „}Ô˜VÒ#Ð#Ý•5”8Ô&¨Ñ5Ô5ð 	MÝ”:×3Ò3°F´MÑBÔBÀfÒLÐLÝ•5˜-Ñ(Ô(ð 	QÝ(¨¸4Ð@Ñ@Ô@ð QÝ”z×7Ò7¸¼ÑFÔFÈ&ÒPÐPå”z×7Ò7¸¼ÑFÔFÈ&ÒPÐPàˆuå•5”8Ô&¨Ñ5Ô5ÐV½ÅÈÑ9VÔ9VÐVr"   c                ó¶  — t          | ||¦  «        r¢t          t          j        j        d¦  «        r?t          j        j                             |                      |j        ¦  «        ||¬¦  «        S t          t          d¦  «        r/t          j        |                      |j        ¦  «        ||¬¦  «        S t          j	        j
                             | ||¬¦  «        S )a  Grouped matrix multiplication dispatcher that uses torch.nn.functional.grouped_mm if available, else falls back to torch._grouped_mm.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (S, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (num_experts, input_dim, output_dim).
        offs (`torch.Tensor`):
            Offsets tensor indicating the boundaries of each group in the input tensor.
    Returns:
        `torch.Tensor`: Output tensor of shape (S, output_dim).
    rx   ©rK   rz   )rƒ   r   r   r€   r�   rx   r=   r>   rz   ÚopsÚtransformersÚgrouped_mm_fallbackr`   s      r    rz   rz   :  s»   € õ$ ˜5 &¨$Ñ/Ô/ð Põ
 •5”8Ô&¨Ñ5Ô5ð 	PÝ”8Ô&×1Ò1°%·(²(¸6¼<Ñ2HÔ2HÈ&ÐW[Ð1Ñ\Ô\Ð\Ý•U˜MÑ*Ô*ð 	PÝÔ$ U§X¢X¨f¬lÑ%;Ô%;¸VÈ$ÐOÑOÔOÐOåŒ9Ô!×5Ò5°e¸VÈ$Ð5ÑOÔOÐOr"   c                óª   — |rt          | ||¬¦  «        }n&t          | |                     dd¦  «        |¬¦  «        }|�|                     |¦  «         |S )a  Grouped linear layer supporting optional bias and transposed weights.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (S, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (num_experts, input_dim, output_dim) if `is_transposed`,
            else of shape (num_experts, output_dim, input_dim).
        offs (`torch.Tensor`):
            Offsets tensor indicating the boundaries of each group in the input tensor.
        bias (`torch.Tensor`, *optional*):
            Bias tensor of shape (num_experts, output_dim). Default is `None`.
        is_transposed (`bool`, *optional*, defaults to `False`):
            Whether the weight tensor is transposed.
    Returns:
        `torch.Tensor`: Output tensor of shape (S, output_dim).
    r…   éþÿÿÿr   )rz   Ú	transposer   )r   r   rK   r   r   r   s         r    Ú_grouped_linearrŒ   Y  sc   € ð0 ð Få˜% ¨dÐ3Ñ3Ô3ˆˆõ ˜% ×!1Ò!1°"°bÑ!9Ô!9ÀÐEÑEÔEˆàÐà�Š�‰Œˆà€Jr"   c                óÄ  — |j         }|                     d¦  «        }|                     d¦  «        }|                     d¦  «        }|                     d¦  «        }|                     d¦  «        }	t          j        |	¦  «        \  }
}|||z           }||         }|j        dv r|
                     ¦   «         n|
                     ¦   «         }t          j        || j	        d| j	        dz
  ¬¦  «        }t          j
        |dt          j        ¬¦  «        }|
| j	        k                         d¦  «        }|
                     | j	        dz
  ¬¦  «         | j        r| j        }| j        r| j        |
         nd }n| j        }| j        r| j        |
         nd }|                     |d¦  «         t+          ||||| j        ¬	¦  «        }| j        r|                      |¦  «        }n|                      |¦  «        }| j        }| j        r| j        |
         nd }t+          ||||| j        ¬	¦  «        }||                     d¦  «        z  }|                     |d¦  «         t          j        |¦  «        }t          j        |                     d¦  «        |¬
¦  «        ||<   ||         }|                     |||¦  «                             d¬¦  «        }|                     |j         ¦  «        S )Nr   r   )rt   Úmpsr   )ÚbinsÚminÚmax)r*   r>   )r‘   g        r+   )rN   r)   )!rN   r,   r.   r   Úsortr}   ÚfloatÚintÚhistcr0   Úcumsumr]   r   Úclamp_r1   r2   r3   r4   r5   r6   Úmasked_fill_rŒ   r   r7   r8   r9   r:   Ú
empty_likeÚaranger;   r<   r=   r>   )r#   r%   r&   r'   rN   r?   r@   rA   rC   rD   Úexpert_ids_gÚpermÚselected_hidden_states_gÚsample_weights_gÚhistc_inputÚtokens_per_expertÚoffsetsÚsentinel_maskrE   rF   rG   rH   Úinv_permrI   s                           r    Úgrouped_mm_experts_forwardr¤     s  € ð Ô!€FØ× Ò  Ñ$Ô$€IØ×#Ò# AÑ&Ô&€JØ×#Ò# BÑ'Ô'€Jð #×*Ò*¨2Ñ.Ô.€NØ×$Ò$ RÑ(Ô(€Jõ œ JÑ/Ô/Ñ€L�$Ø,¨T°YÑ->Ô?ÐØ% dÔ+Ðð +1¬+¸Ð*GÐ*G�,×$Ò$Ñ&Ô&Ð&È\×M]ÒM]ÑM_ÔM_€KÝœ K°dÔ6FÈAÐSWÔScÐfgÑSgÐhÑhÔhÐÝŒlÐ,°!½5¼;ÐGÑGÔG€Gð" " TÔ%5Ò5×@Ò@ÀÑDÔD€MØ×Ò˜DÔ,¨qÑ0ÐÑ1Ô1Ð1ð „}ð UØÔ,ÐØBFÄ-ÐY˜$Ô0°Ô>Ð>ÐUYˆˆàœ<ÐØ=A¼]ÐT˜$Ô+¨LÔ9Ð9ÐPTˆð ×)Ò)¨-¸Ñ=Ô=Ð=õ Ø Ð"2°GÀ/ÐaeÔasðñ ô €Hð
 „}ð )à×#Ò# HÑ-Ô-ˆˆð —;’;˜xÑ(Ô(ˆð ”~ÐØ;?¼=ÐR�dÔ)¨,Ô7Ð7Èd€Oõ ØÐ" G°/ÐQUÔQcðñ ô €Hð
 Ð.×8Ò8¸Ñ<Ô<Ñ<€Lð ×Ò˜m¨SÑ1Ô1Ð1õ Ô Ñ%Ô%€HÝ”\ $§)¢)¨A¡,¤,°vÐ>Ñ>Ô>€HˆT�NØ Ô)€Lð '×+Ò+¨J¸	À:ÑNÔN×RÒRÐWXÐRÑYÔYÐà×!Ò! -Ô"5Ñ6Ô6Ð6r"   c                  ó2   ‡ — e Zd ZdZeeeedœZd	ˆ fd„Z	ˆ xZ
S )
ÚExpertsInterfacez;Interface for registering custom experts forward functions.)ÚdeepgemmÚ
batched_mmrx   ÚsonicmoeÚexperts_implementationÚstrÚdefaultr   r   c                ó¼   •— |€t                                d¦  «         n|dk    r|| vrt          d|› d�¦  «        ‚t          ¦   «                              ||¦  «        S )zfReturn the requested `experts_implementation`. Also strictly check its validity, and raise if invalid.Na
  You tried to access the `ExpertsInterface` with a `config._experts_implementation` set to `None`. This is expected if you use an Expert Module as a standalone Module. If this is not the case, something went wrong with the dispatch of `config._experts_implementation`Úeagerú`zL` is not a valid experts implementation registered in the `ExpertsInterface`)ÚloggerÚwarning_onceÚKeyErrorÚsuperÚget)r#   rª   r¬   Ú	__class__s      €r    Úget_interfacezExpertsInterface.get_interfaceñ  s�   ø€ à!Ð)Ý×ÒðNñô ð ð ð
 $ wÒ.Ð.Ð3IÐQUÐ3UÐ3UÝØxÐ*ÐxÐxÐxñô ð õ ‰wŒw�{Š{Ð1°7Ñ;Ô;Ð;r"   )rª   r«   r¬   r   r   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   rJ   r¤   r   Ú_global_mappingr¶   Ú__classcell__)rµ   s   @r    r¦   r¦   ç  s]   ø€ € € € € ØEÐEð 2Ø0Ø0Ø,ð	ð €Oð<ð <ð <ð <ð <ð <ð <ð <ð <ð <r"   r¦   Úgate_up_outc                óf   — |                      dd¬¦  «        \  }}|                      |¦  «        |z  S )a›  
    Default gating mechanism: splits the gate_up_out into gate and up parts,
    applies the activation function to the gate part, and multiplies it with the up part.
    Args:
        gate_up_out (`torch.Tensor`):
            The output tensor from the gate and up projection of shape (S, 2 * intermediate_dim).
    Returns:
        `torch.Tensor`: The gated output tensor of shape (S, intermediate_dim).
    r   r   r)   )Úchunkr8   )r#   r½   ÚgateÚups       r    Ú_default_apply_gaterÂ     s7   € ð × Ò  ¨Ð Ñ+Ô+�H€Dˆ"Ø�;Š;�tÑÔ˜rÑ!Ð!r"   T)Úexperts_interfaceÚis_concatenatedr   r3   r1   Úexperts_classútype[torch.nn.Module] | NonerÃ   rÄ   r3   r1   útype[torch.nn.Module]c               ó>   ‡‡‡‡‡— dˆˆˆˆˆfd„}| � || ¦  «        S |S )a¤  Decorator to modify experts class to support different experts implementations.

    Args:
        experts_class (`type[torch.nn.Module]`, *optional*):
            The experts class to modify. If not provided, returns a decorator that can be applied to the class.
        experts_interface (`ExpertsInterface`, *optional*, defaults to `ALL_EXPERTS_FUNCTIONS`):
            The experts interface to use for dispatching the forward method.
        is_concatenated (`bool`, *optional*, defaults to `True`):
            Whether the expert weights are stored in concatenated layout [gate;up]
            or interleaved layout [gate0, up0, gate1, up1, ...].
        is_transposed (`bool`, *optional*, defaults to `False`):
            Whether the expert weights are stored in transposed format.
        has_bias (`bool`, *optional*, defaults to `False`):
            Whether the expert layers include bias terms or not.
        has_gate (`bool`, *optional*, defaults to `True`):
            Whether the experts use a gating mechanism or not.
            Whether it has gate_up_proj weights or just up_proj weights.

    Returns:
        `type[torch.nn.Module]`: The modified experts class.
    rÅ   rÇ   r   c                óî   •‡‡— | j         Š| j        Št          ‰¦  «        ˆˆˆˆ	ˆfd„¦   «         }t          ‰¦  «        ˆˆfd„¦   «         }t          | d¦  «        st          | _        || _         || _        | S )Nc                óh   •—  ‰| |g|¢R i |¤Ž || _         ‰| _        ‰| _        ‰| _        ‰| _        d S ©N)Úconfigr1   r3   r   rÄ   )	r#   rÌ   ÚargsÚkwargsr3   r1   rÄ   r   Úoriginal_inits	       €€€€€r    Ú__init__z=use_experts_implementation.<locals>.wrapper.<locals>.__init__4  sP   ø€ àˆM˜$ Ð8¨Ð8Ð8Ð8°Ð8Ð8Ð8Ø ˆDŒKØ$ˆDŒMØ$ˆDŒMØ!.ˆDÔØ#2ˆDÔ Ð Ð r"   c                ó\   •— ‰                      | j        j        ‰¦  «        } || g|¢R i |¤ŽS rË   )r¶   rÌ   Ú_experts_implementation)r#   rÍ   rÎ   Úexperts_forwardrÃ   Úoriginal_forwards       €€r    Úforwardz<use_experts_implementation.<locals>.wrapper.<locals>.forward=  s>   ø€ à/×=Ò=¸d¼kÔ>aÐcsÑtÔtˆOØ"�? 4Ð9¨$Ð9Ð9Ð9°&Ð9Ð9Ð9r"   r7   )rÐ   rÕ   r   r   rÂ   r7   )
rÅ   rÐ   rÕ   rÔ   rÏ   rÃ   r3   r1   rÄ   r   s
      @@€€€€€r    Úwrapperz+use_experts_implementation.<locals>.wrapper0  s¼   øøø€ Ø%Ô.ˆØ(Ô0Ðå	ˆ}Ñ	Ô	ð	3ð 	3ð 	3ð 	3ð 	3ð 	3ð 	3ð 	3ñ 
Ô	ð	3õ 
ÐÑ	 Ô	 ð	:ð 	:ð 	:ð 	:ð 	:ñ 
!Ô	 ð	:õ �} mÑ4Ô4ð 	<Ý(;ˆMÔ%à!)ˆÔØ 'ˆÔØÐr"   N)rÅ   rÇ   r   rÇ   ro   )rÅ   rÃ   rÄ   r   r3   r1   rÖ   s    ````` r    Úuse_experts_implementationr×     sV   øøøøø€ ð>ð ð ð ð ð ð ð ð ð ð2 Ð Øˆw�}Ñ%Ô%Ð%à€Nr"   )NF)
r   r   r   r   r   r   r   r   r   r   )
r#   r$   r%   r   r&   r   r'   r   r   r   )r   r   r   r   rK   r   r   r   )r   r   r   r   rK   r   r   r   )r   r   r   r   rK   r   r   r   r   r   r   r   )r½   r   r   r   rË   )rÅ   rÆ   rÃ   r¦   rÄ   r   r   r   r3   r   r1   r   r   rÇ   )+Ú
__future__r   Úcollections.abcr   Ú	functoolsr   Úutilsr   Úutils.genericr   Úutils.import_utilsr	   r
   r   r   r§   r   r©   r   r   Ú_dynamoÚassume_constant_resultÚ
get_loggerr·   r°   r!   rJ   rX   ra   rf   rn   ÚlibraryÚ	custom_opÚregister_fakeÚregister_autogradrƒ   rz   rŒ   r¤   r¦   ÚALL_EXPERTS_FUNCTIONSrÂ   r×   ro   r"   r    ú<module>ræ      s)  ðð #Ð "Ð "Ð "Ð "Ð "à $Ð $Ð $Ð $Ð $Ð $Ø Ð Ð Ð Ð Ð à Ð Ð Ð Ð Ð Ø ,Ð ,Ð ,Ð ,Ð ,Ð ,ðð ð ð ð ð ð ð ð ð ð ð ð 4Ð 3Ð 3Ð 3Ð 3Ð 3Ø .Ð .Ð .Ð .Ð .Ð .ð ÐÑÔð ZØ€L€L€Lð
 !&¤× DÒ DÐE^Ñ _Ô _ÐØ"œ]×AÒAÐBXÑYÔYÐð 
ˆÔ	˜HÑ	%Ô	%€ð\ !%Øð	ð ð ð ð ðD=7ð =7ð =7ð =7ðFð ð ð ð4^ð ^ð ^ð ^ðð ð ð)ð )ð )ð& ÐÑÔð Ø	„M×ÒØ+ØØØEð	 ñ ô ð ð 
„M×ÒÐ CÐE^Ñ_Ô_Ð_Ø	„M×#Ò#Ø+Ø%Ø8ð $ñ ô ð ð-Wð -Wð -Wð -Wð`Pð Pð Pð PðF !%Øð#ð #ð #ð #ð #ðLe7ð e7ð e7ð e7ðP<ð <ð <ð <ð <Ð'ñ <ô <ð <ð2 )Ð(Ñ*Ô*Ð ð"ð "ð "ð "ð 37ð;ð +@Ø ØØØð;ð ;ð ;ð ;ð ;ð ;ð ;ð ;r"   